diff --git a/.dockerignore b/.dockerignore index 93a95eeb5..28cd2806a 100644 --- a/.dockerignore +++ b/.dockerignore @@ -9,6 +9,10 @@ target/ frontend/node_modules/ frontend/dist/ frontend/.vite/ +aether-vscodex/web/node_modules/ +aether-vscodex/web/dist/ +aether-vscodex/vscode-extension/node_modules/ +aether-vscodex/vscode-extension/dist/ # Development .git/ diff --git a/.env.example b/.env.example index a529dd673..c50116a15 100644 --- a/.env.example +++ b/.env.example @@ -15,6 +15,11 @@ APP_PORT=8084 # APP_IMAGE=ghcr.io/fawney19/aether:beta # APP_IMAGE=ghcr.io/fawney19/aether:0.7.0-rc.1 +# Compose 应用容器的非 root 数字身份。Single Node 的 ./data 必须归该身份所有; +# install.sh 会自动写入安装用户的 UID/GID 并迁移旧数据目录。 +AETHER_CONTAINER_UID=65532 +AETHER_CONTAINER_GID=65532 + # API Key 前缀(默认 sk) API_KEY_PREFIX=sk @@ -35,12 +40,16 @@ DB_HOST=localhost DB_PORT=5432 DB_USER=postgres DB_NAME=aether -DB_PASSWORD=aether +DB_PASSWORD= # Redis 配置 REDIS_HOST=localhost REDIS_PORT=6379 -REDIS_PASSWORD=aether +REDIS_PASSWORD= + +# 可选 MySQL profile 的应用用户与 root 密码 +MYSQL_PASSWORD= +MYSQL_ROOT_PASSWORD= # JWT密钥(使用 ./generate_keys.sh 生成) # 用于用户登录 token 签名,更换后所有用户需重新登录 @@ -50,8 +59,12 @@ JWT_SECRET_KEY=change-this-to-a-secure-random-string # 注意:更换此密钥后需要在管理面板重新配置所有 Provider API Key ENCRYPTION_KEY=change-this-to-another-secure-random-string +# S3 备份的独立加密密钥(推荐)。未配置时为兼容旧部署,会回退到 ENCRYPTION_KEY。 +# 密钥轮换前必须保留旧值,离线恢复工具需要它解密历史备份。 +# AETHER_BACKUP_ENCRYPTION_KEY=change-this-to-a-dedicated-secure-random-string + # 启动自举管理员(仅在当前库里还没有活动管理员时生效) -# 手动部署时取消注释并设置;install.sh 首次生成配置时会提示输入。 +# 首次启动前必须设置 ADMIN_PASSWORD;install.sh 首次生成配置时会提示输入。 ADMIN_EMAIL=admin@example.com ADMIN_USERNAME=admin123456 # ADMIN_PASSWORD= @@ -63,21 +76,51 @@ ADMIN_USERNAME=admin123456 # Docker/Nginx 位于独立容器时,请按实际容器网络设置,例如:172.16.0.0/12。 # AETHER_TRUSTED_PROXY_CIDRS=127.0.0.0/8,::1/128,172.16.0.0/12 -# docker compose 下 app 启动前自动执行 pending migration/backfill(默认 true) -# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true +# VS Code Codex 云端协同(仅在叠加 aether-vscodex/docker-compose.aether.yml 时需要) +# 内部 token 至少 24 字节,建议使用:openssl rand -base64 32 +# AETHER_VSCODEX_INTERNAL_TOKEN=replace-with-a-long-random-secret +# AETHER_VSCODEX_PUBLIC_WS_URL=wss://aether.example.com/api/vscodex/ws +# AETHER_VSCODEX_ALLOWED_ORIGINS=https://aether.example.com + +# 启动时的数据库准备策略:auto(默认)或 verify-only +# AETHER_GATEWAY_DATABASE_MODE=auto # PostgreSQL 连接池配置(默认每核 4 条、总池至少 32 条且最多 100 条;多实例部署应显式分配每实例预算) # AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12 # AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80 # AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048 # AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB=256 -# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=120000 -# 可选的 Payload 上限(MiB);默认及 0 均表示不限制。 -# AETHER_MAX_REQUEST_BODY_MB=0 +# 请求体完整读取总超时默认关闭;确需限制时配置 1000-600000 毫秒的非零值。 +# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=0 +# 单请求解压后 Payload 上限(MiB),默认 256;显式设为 0 才表示不限制。 +# AETHER_MAX_REQUEST_BODY_MB=256 # AETHER_GATEWAY_SECURITY_CACHE_TTL_MS=1000 -# AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB=0 -# AETHER_MAX_INTERNAL_BUFFERED_BODY_MB=0 +# AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB=64 +# AETHER_MAX_INTERNAL_BUFFERED_BODY_MB=64 # AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY=1024 +# Tunnel relay 使用的独立 HMAC 密钥。启用 HTTP tunnel relay 或多网关 owner 转发时必须配置, +# 所有网关实例必须使用同一个至少 32 字节的随机值;不要复用 JWT 或数据加密密钥。 +# AETHER_TUNNEL_RELAY_AUTH_SECRET= +# 旧版 /api/internal/gateway/* 控制面默认关闭。确需独立服务调用时,配置至少 32 字节的 +# 独立 HMAC 密钥;不要复用 JWT、数据加密或 tunnel relay 密钥。多节点必须使用相同值和共享 Redis。 +# AETHER_INTERNAL_GATEWAY_AUTH_SECRET= +# 远程 relay 地址必须使用 HTTPS;HTTP 仅允许 localhost 或回环 IP。 +# AETHER_TUNNEL_RELAY_BASE_URL=https://gateway-a.example.com +# 跨网关 relay 解析到受控私有地址时才显式开启;默认关闭以防止被篡改的 attachment +# 记录诱导网关向内网转发 relay 凭据。该开关不放宽普通 provider 的目标地址策略。 +# AETHER_TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS=false +# 更推荐按 relay 主机名精确放行私网部署(逗号分隔,大小写不敏感);不支持通配符/后缀。 +# AETHER_TUNNEL_RELAY_PRIVATE_HOST_ALLOWLIST=gateway-a.internal,gateway-b.internal +# Bark 自建服务默认仅允许公网 HTTPS。确需明文 HTTP 或内网目标时分别显式开启: +# AETHER_BARK_ALLOW_HTTP=false +# AETHER_BARK_ALLOW_PRIVATE_TARGETS=false + +# 可选 Provider OAuth 客户端。使用 Gemini CLI / Antigravity 浏览器授权时必须配置 +# 对应的 client secret;client ID 未配置时使用内置的公开 native-app client ID。 +# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID= +# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET= +# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID= +# AETHER_ANTIGRAVITY_OAUTH_CLIENT_SECRET= # PostgreSQL 容器调优:docker-compose.yml 已内置通用默认值,通常不用配置。 # 只有在 Postgres 独占大内存、或压测显示 DB 缓存/排序/维护任务成为瓶颈时再覆盖。 diff --git a/.github/workflows/build-tunnel.yml b/.github/workflows/build-tunnel.yml index 62b5770dd..78b4fa514 100644 --- a/.github/workflows/build-tunnel.yml +++ b/.github/workflows/build-tunnel.yml @@ -6,7 +6,8 @@ on: workflow_dispatch: permissions: - contents: write + actions: read + contents: read concurrency: group: build-tunnel-${{ github.ref }} @@ -17,7 +18,7 @@ jobs: runs-on: ubuntu-latest if: startsWith(github.ref, 'refs/tags/') steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Ensure tunnel tag matches Cargo version shell: bash @@ -78,10 +79,10 @@ jobs: use_cross: false steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: targets: ${{ matrix.target }} @@ -89,14 +90,14 @@ jobs: run: rustup target add ${{ matrix.target }} - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: workspaces: apps/aether-tunnel -> target key: ${{ matrix.target }} - name: Install cross if: matrix.use_cross - uses: taiki-e/install-action@cross + uses: taiki-e/install-action@1ae7257be536a92d9218a6b343dc6e6ba650f7e1 # cross - name: Build working-directory: apps/aether-tunnel @@ -122,9 +123,10 @@ jobs: run: | cd target/${{ matrix.target }}/release 7z a ../../../aether-tunnel-${{ matrix.name }}.zip aether-tunnel.exe + tar czf ../../../aether-tunnel-${{ matrix.name }}.tar.gz aether-tunnel.exe - name: Upload artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: aether-tunnel-${{ matrix.name }} path: | @@ -137,9 +139,14 @@ jobs: needs: build runs-on: ubuntu-latest if: startsWith(github.ref, 'refs/tags/') + permissions: + actions: read + attestations: write + contents: write + id-token: write steps: - name: Download all artifacts - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: merge-multiple: true path: artifacts @@ -148,6 +155,20 @@ jobs: working-directory: artifacts run: sha256sum aether-tunnel-* > SHA256SUMS.txt + - name: Attest tunnel release provenance + id: attest-release + uses: actions/attest@1e69f48acb82d1966a394da916b4c1698aa569d6 # v4.2.2 + with: + subject-path: | + artifacts/aether-tunnel-*.tar.gz + artifacts/aether-tunnel-*.zip + artifacts/SHA256SUMS.txt + + - name: Bundle tunnel release provenance + env: + ATTESTATION_BUNDLE: ${{ steps.attest-release.outputs.bundle-path }} + run: install -m 0644 "${ATTESTATION_BUNDLE}" artifacts/AETHER_TUNNEL_RELEASE_PROVENANCE.sigstore.json + - name: Delete stale draft releases for tag env: GH_TOKEN: ${{ github.token }} @@ -170,12 +191,13 @@ jobs: done <<< "${draft_ids}" - name: Create GitHub Release - uses: softprops/action-gh-release@v2 + uses: softprops/action-gh-release@3bb12739c298aeb8a4eeaf626c5b8d85266b0e65 # v2 with: name: "${{ github.ref_name }}" generate_release_notes: true files: | artifacts/aether-tunnel-* + artifacts/AETHER_TUNNEL_RELEASE_PROVENANCE.sigstore.json artifacts/SHA256SUMS.txt fail_on_unmatched_files: true @@ -183,8 +205,10 @@ jobs: needs: release runs-on: ubuntu-latest if: startsWith(github.ref, 'refs/tags/') + permissions: + contents: write steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: main diff --git a/.github/workflows/deploy-pages.yml b/.github/workflows/deploy-pages.yml index b57dfeb1b..d579c2db4 100644 --- a/.github/workflows/deploy-pages.yml +++ b/.github/workflows/deploy-pages.yml @@ -7,8 +7,6 @@ on: permissions: contents: read - pages: write - id-token: write concurrency: group: pages @@ -46,14 +44,22 @@ jobs: if: needs.preflight.outputs.deploy_pages == 'true' runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Setup Node.js - uses: actions/setup-node@v5 + uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5 with: node-version: '22' cache: 'npm' - cache-dependency-path: frontend/package-lock.json + cache-dependency-path: | + frontend/package-lock.json + aether-vscodex/web/package-lock.json + + - name: Build aether-vscodex web + working-directory: aether-vscodex/web + run: | + npm ci + npm run build - name: Install dependencies working-directory: frontend @@ -69,10 +75,10 @@ jobs: run: cp frontend/dist/index.html frontend/dist/404.html - name: Setup Pages - uses: actions/configure-pages@v5 + uses: actions/configure-pages@983d7736d9b0ae728b81ab479565c72886d7745b # v5 - name: Upload artifact - uses: actions/upload-pages-artifact@v3 + uses: actions/upload-pages-artifact@56afc609e74202658d3ffba0e8f6dda462b719fa # v3 with: path: frontend/dist @@ -82,7 +88,10 @@ jobs: url: ${{ steps.deployment.outputs.page_url }} runs-on: ubuntu-latest needs: build + permissions: + id-token: write + pages: write steps: - name: Deploy to GitHub Pages id: deployment - uses: actions/deploy-pages@v4 + uses: actions/deploy-pages@d6db90164ac5ed86f2b6aed7e0febac5b3c0c03e # v4 diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index 9cedd955e..c216dbfd6 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -25,7 +25,6 @@ env: CARGO_PROFILE_TEST_DEBUG: '0' CARGO_TERM_COLOR: always RUST_BACKTRACE: '1' - GHCR_IMAGE: ghcr.io/fawney19/aether jobs: source: @@ -36,6 +35,7 @@ jobs: sha: ${{ steps.snapshot.outputs.sha }} short_sha: ${{ steps.snapshot.outputs.short_sha }} date: ${{ steps.snapshot.outputs.date }} + ghcr_image: ${{ steps.snapshot.outputs.ghcr_image }} steps: - name: Require main branch id: snapshot @@ -49,9 +49,13 @@ jobs: fi sha="${GITHUB_SHA}" + # Docker 镜像仓库名必须全小写;GitHub owner 可能保留大写,先统一规范化。 + repository_owner="${GITHUB_REPOSITORY%%/*}" + repository_owner="${repository_owner,,}" echo "sha=${sha}" >> "${GITHUB_OUTPUT}" echo "short_sha=${sha:0:7}" >> "${GITHUB_OUTPUT}" echo "date=$(date -u +'%Y-%m-%d')" >> "${GITHUB_OUTPUT}" + echo "ghcr_image=ghcr.io/${repository_owner}/aether" >> "${GITHUB_OUTPUT}" echo "Building main at ${sha}." # Keep the scheduled backend coverage in one place so it cannot drift from PR CI. @@ -66,12 +70,12 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 90 steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Install pinned Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 @@ -79,13 +83,13 @@ jobs: run: rustc -Vv - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: nightly-rust-1.95-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Check all workspace targets env: @@ -112,16 +116,26 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 30 steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Setup Node.js - uses: actions/setup-node@v5 + uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5 with: node-version: '22' cache: npm - cache-dependency-path: frontend/package-lock.json + cache-dependency-path: | + frontend/package-lock.json + aether-vscodex/web/package-lock.json + + # The frontend prebuild synchronizes the embedded VSCodex UI by running + # its build from a separate package. Install that package explicitly so + # vue-tsc can resolve vite/client, vitest/globals, and node types in a + # clean runner. + - name: Install VSCodex web dependencies + working-directory: aether-vscodex/web + run: npm ci - name: Install dependencies working-directory: frontend @@ -147,7 +161,7 @@ jobs: run: npm run build - name: Upload frontend artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: nightly-frontend-dist path: frontend/dist/ @@ -161,12 +175,12 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 10 steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Setup Node.js - uses: actions/setup-node@v5 + uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5 with: node-version: '22' @@ -250,25 +264,25 @@ jobs: os: macos-15 use_cross: false steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Install pinned Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 targets: ${{ matrix.target }} - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: nightly-release-${{ matrix.target }} workspaces: . -> target - name: Install cross if: matrix.use_cross - uses: taiki-e/install-action@cross + uses: taiki-e/install-action@1ae7257be536a92d9218a6b343dc6e6ba650f7e1 # cross - name: Build release binary env: @@ -285,7 +299,7 @@ jobs: fi - name: Upload binary artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: nightly-gateway-${{ matrix.platform }}-${{ matrix.arch }} path: target/${{ matrix.target }}/release/aether-gateway @@ -298,17 +312,19 @@ jobs: needs: [source, checks, build] if: ${{ needs.checks.result == 'success' && needs.build.result == 'success' }} runs-on: ubuntu-latest + env: + GHCR_IMAGE: ${{ needs.source.outputs.ghcr_image }} permissions: actions: read contents: read packages: write steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Download Linux binaries and frontend - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: pattern: nightly-* path: artifacts @@ -325,20 +341,20 @@ jobs: cp -R artifacts/nightly-frontend-dist/. dist/frontend/ - name: Set up QEMU - uses: docker/setup-qemu-action@v3 + uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3 - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 + uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6 # v3 - name: Log in to GHCR - uses: docker/login-action@v3 + uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3 with: registry: ghcr.io username: ${{ github.actor }} password: ${{ secrets.GITHUB_TOKEN }} - name: Build and push nightly image - uses: docker/build-push-action@v6 + uses: docker/build-push-action@10e90e3645eae34f1e60eeb005ba3a3d33f178e8 # v6 with: context: . file: ./Dockerfile.app @@ -362,12 +378,12 @@ jobs: actions: read contents: read steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 with: ref: ${{ needs.source.outputs.sha }} - name: Download nightly artifacts - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: pattern: nightly-* path: artifacts @@ -424,7 +440,7 @@ jobs: done - name: Upload nightly package artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: nightly-release-assets path: release-assets/* @@ -442,7 +458,7 @@ jobs: contents: write steps: - name: Download nightly package artifact - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: name: nightly-release-assets path: release-assets @@ -456,6 +472,7 @@ jobs: SOURCE_SHA: ${{ needs.source.outputs.sha }} SOURCE_SHORT_SHA: ${{ needs.source.outputs.short_sha }} RELEASE_DATE: ${{ needs.source.outputs.date }} + GHCR_IMAGE: ${{ needs.source.outputs.ghcr_image }} run: | set -euo pipefail diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 33256e04d..9f453bc86 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -6,8 +6,8 @@ on: workflow_dispatch: permissions: - contents: write - packages: write + actions: read + contents: read concurrency: group: release-aether-${{ github.ref }} @@ -70,14 +70,22 @@ jobs: needs: preflight runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Setup Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 with: node-version: 22 cache: npm - cache-dependency-path: frontend/package-lock.json + cache-dependency-path: | + frontend/package-lock.json + aether-vscodex/web/package-lock.json + + - name: Build aether-vscodex web + working-directory: aether-vscodex/web + run: | + npm ci + npm run build - name: Install & build working-directory: frontend @@ -86,13 +94,78 @@ jobs: npm run build - name: Upload frontend artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: frontend-dist path: frontend/dist/ if-no-files-found: error retention-days: 1 + vscodex: + name: Build VS Code Codex extension + needs: preflight + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 + + - name: Setup Node.js + uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 + with: + node-version: 22 + cache: npm + cache-dependency-path: | + aether-vscodex/package-lock.json + aether-vscodex/web/package-lock.json + aether-vscodex/vscode-extension/package-lock.json + + - name: Install module test dependencies + working-directory: aether-vscodex + run: npm ci + + - name: Build the embedded Web UI + working-directory: aether-vscodex/web + run: | + npm ci + npm run build + + - name: Install extension dependencies + working-directory: aether-vscodex/vscode-extension + run: npm ci + + - name: Check and compile the extension + working-directory: aether-vscodex/vscode-extension + run: | + npm run check + npm run build + + - name: Run module tests + working-directory: aether-vscodex + run: npm test + + - name: Run Web UI tests + working-directory: aether-vscodex/web + run: npm test + + - name: Package VSIX + working-directory: aether-vscodex/vscode-extension + shell: bash + run: | + set -euo pipefail + version="$(node -p "require('./package.json').version")" + npx --yes @vscode/vsce package --no-update-package-json --allow-missing-repository + source_vsix="codex-remote-collab-${version}.vsix" + test -f "${source_vsix}" + mv "${source_vsix}" "aether-vscodex-${version}.vsix" + unzip -l "aether-vscodex-${version}.vsix" | grep 'extension/node_modules/ws/index.js' >/dev/null + + - name: Upload VSIX artifact + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 + with: + name: aether-vscodex-vsix + path: aether-vscodex/vscode-extension/aether-vscodex-*.vsix + if-no-files-found: error + retention-days: 7 + build: name: Build ${{ matrix.name }} needs: preflight @@ -126,22 +199,22 @@ jobs: os: macos-15 use_cross: false steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: targets: ${{ matrix.target }} - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: release-${{ matrix.target }} workspaces: . -> target - name: Install cross if: matrix.use_cross - uses: taiki-e/install-action@cross + uses: taiki-e/install-action@1ae7257be536a92d9218a6b343dc6e6ba650f7e1 # cross - name: Build env: @@ -157,7 +230,7 @@ jobs: fi - name: Upload binary artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: aether-gateway-${{ matrix.platform }}-${{ matrix.arch }} path: target/${{ matrix.target }}/release/aether-gateway @@ -169,11 +242,17 @@ jobs: needs: [preflight, frontend, build] if: needs.preflight.outputs.publish == 'true' runs-on: ubuntu-latest + permissions: + actions: read + attestations: write + contents: read + id-token: write + packages: write steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Download all artifacts - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: path: artifacts @@ -186,27 +265,27 @@ jobs: cp -r artifacts/frontend-dist dist/frontend - name: Set up QEMU - uses: docker/setup-qemu-action@v3 + uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3 - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 + uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3 - name: Log in to GHCR - uses: docker/login-action@v3 + uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3 with: registry: ${{ env.REGISTRY }} username: ${{ github.actor }} password: ${{ secrets.GITHUB_TOKEN }} - name: Log in to Docker Hub - uses: docker/login-action@v3 + uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3 with: username: ${{ secrets.DOCKERHUB_USERNAME }} password: ${{ secrets.DOCKERHUB_TOKEN }} - name: Extract metadata id: meta - uses: docker/metadata-action@v5 + uses: docker/metadata-action@c299e40c65443455700f0fdfc63efafe5b349051 # v5 with: images: | ${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }} @@ -222,7 +301,8 @@ jobs: latest=false - name: Build and push - uses: docker/build-push-action@v6 + id: push + uses: docker/build-push-action@10e90e3645eae34f1e60eeb005ba3a3d33f178e8 # v6 with: context: . file: ./Dockerfile.app @@ -231,15 +311,36 @@ jobs: labels: ${{ steps.meta.outputs.labels }} platforms: linux/amd64,linux/arm64 + - name: Attest GHCR image provenance + uses: actions/attest@1e69f48acb82d1966a394da916b4c1698aa569d6 # v4.2.2 + with: + subject-name: ${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }} + subject-digest: ${{ steps.push.outputs.digest }} + push-to-registry: true + create-storage-record: false + + - name: Attest Docker Hub image provenance + uses: actions/attest@1e69f48acb82d1966a394da916b4c1698aa569d6 # v4.2.2 + with: + subject-name: docker.io/${{ env.DOCKERHUB_IMAGE }} + subject-digest: ${{ steps.push.outputs.digest }} + push-to-registry: true + create-storage-record: false + package: name: Release tarballs needs: [preflight, frontend, build] runs-on: ubuntu-latest + permissions: + actions: read + attestations: write + contents: read + id-token: write steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Download all artifacts - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: path: artifacts @@ -289,8 +390,24 @@ jobs: chmod +x release-assets/install.sh (cd release-assets && sha256sum *.tar.gz > SHA256SUMS) + - name: Attest release package provenance + id: attest-release + if: needs.preflight.outputs.publish == 'true' + uses: actions/attest@1e69f48acb82d1966a394da916b4c1698aa569d6 # v4.2.2 + with: + subject-path: | + release-assets/*.tar.gz + release-assets/install.sh + release-assets/SHA256SUMS + + - name: Bundle release package provenance + if: needs.preflight.outputs.publish == 'true' + env: + ATTESTATION_BUNDLE: ${{ steps.attest-release.outputs.bundle-path }} + run: install -m 0644 "${ATTESTATION_BUNDLE}" release-assets/AETHER_RELEASE_PROVENANCE.sigstore.json + - name: Upload release package artifact - uses: actions/upload-artifact@v5 + uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5 with: name: release-assets path: release-assets/* @@ -299,16 +416,25 @@ jobs: github-release: name: GitHub Release assets - needs: [preflight, docker, package] + needs: [preflight, docker, package, vscodex] if: needs.preflight.outputs.publish == 'true' runs-on: ubuntu-latest + permissions: + actions: read + contents: write steps: - name: Download release package artifact - uses: actions/download-artifact@v5 + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 with: name: release-assets path: release-assets + - name: Download VSIX artifact + uses: actions/download-artifact@634f93cb2916e3fdff6788551b99b062d0335ce0 # v5 + with: + name: aether-vscodex-vsix + path: release-assets + - name: Delete stale draft releases for tag env: GH_TOKEN: ${{ github.token }} @@ -331,12 +457,14 @@ jobs: done <<< "${draft_ids}" - name: Publish GitHub Release assets - uses: softprops/action-gh-release@v2 + uses: softprops/action-gh-release@3bb12739c298aeb8a4eeaf626c5b8d85266b0e65 # v2 with: generate_release_notes: true prerelease: ${{ needs.preflight.outputs.prerelease }} make_latest: ${{ needs.preflight.outputs.make_latest }} files: | release-assets/*.tar.gz + release-assets/AETHER_RELEASE_PROVENANCE.sigstore.json release-assets/SHA256SUMS release-assets/install.sh + release-assets/*.vsix diff --git a/.github/workflows/rust-ci.yml b/.github/workflows/rust-ci.yml index 6f115c467..d5563c019 100644 --- a/.github/workflows/rust-ci.yml +++ b/.github/workflows/rust-ci.yml @@ -11,6 +11,23 @@ on: - "Cargo.lock" - "crates/**" - "apps/**" + - "install.sh" + - "deploy.sh" + - "update.sh" + - "generate_keys.sh" + - ".env.example" + - "README.md" + - "Dockerfile.app" + - "docker-compose.yml" + - "docker-compose.single-node.yml" + - "tests/install_*_test.sh" + - "tests/deploy_*_test.sh" + - "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" pull_request: @@ -19,6 +36,23 @@ on: - "Cargo.lock" - "crates/**" - "apps/**" + - "install.sh" + - "deploy.sh" + - "update.sh" + - "generate_keys.sh" + - ".env.example" + - "README.md" + - "Dockerfile.app" + - "docker-compose.yml" + - "docker-compose.single-node.yml" + - "tests/install_*_test.sh" + - "tests/deploy_*_test.sh" + - "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" @@ -36,14 +70,34 @@ env: CARGO_TERM_COLOR: always jobs: + shell_security: + name: Shell security fixtures + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 + + - name: Run installer and supply-chain fixtures + shell: bash + run: | + bash tests/deploy_state_safety_test.sh + bash tests/install_archive_safety_test.sh + bash tests/install_container_runtime_security_test.sh + bash tests/install_current_release_link_test.sh + bash tests/install_local_bundle_safety_test.sh + bash tests/install_privileged_write_safety_test.sh + bash tests/install_source_trust_test.sh + bash tests/release_supply_chain_test.sh + bash tests/update_compose_safety_test.sh + bash tests/tunnel_installer_config_security_test.sh + fmt: name: Format runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 components: rustfmt @@ -55,22 +109,22 @@ jobs: name: Clippy (Gateway) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 components: clippy - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Clippy env: @@ -89,22 +143,22 @@ jobs: name: Clippy (Data) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 components: clippy - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Clippy env: @@ -123,22 +177,22 @@ jobs: name: Clippy (Workspace Rest) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: 1.95.0 components: clippy - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Clippy env: @@ -175,28 +229,28 @@ jobs: name: Test (Gateway) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Setup mold - uses: rui314/setup-mold@v1 + uses: rui314/setup-mold@7e4f20ad28a2e8ca6fd0892ccf72e2abb706b9c3 # v1 - name: Install nextest - uses: taiki-e/install-action@nextest + uses: taiki-e/install-action@d5f9268ff7620505a81ada10ddf18cdd72240185 # nextest - name: Test lib env: @@ -225,25 +279,25 @@ jobs: name: Test (Data) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Install nextest - uses: taiki-e/install-action@nextest + uses: taiki-e/install-action@d5f9268ff7620505a81ada10ddf18cdd72240185 # nextest - name: Test env: @@ -270,19 +324,19 @@ jobs: - sqlite - all-drivers steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Check selected data driver env: @@ -301,25 +355,25 @@ jobs: name: Test (Workspace Rest) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Install nextest - uses: taiki-e/install-action@nextest + uses: taiki-e/install-action@d5f9268ff7620505a81ada10ddf18cdd72240185 # nextest - name: Test env: @@ -345,22 +399,22 @@ jobs: - aether-data-mysql - aether-data-sqlite steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Install nextest - uses: taiki-e/install-action@nextest + uses: taiki-e/install-action@d5f9268ff7620505a81ada10ddf18cdd72240185 # nextest - name: Test adapter env: @@ -379,19 +433,19 @@ jobs: name: Test (Integration Scenarios) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Test scenario binaries and end-to-end suites env: @@ -434,22 +488,22 @@ jobs: name: Data DB Smoke (SQLite) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Run SQLite data smoke tests env: @@ -482,22 +536,22 @@ jobs: --health-timeout=5s --health-retries=20 steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Add PostgreSQL server binaries to PATH run: echo "$(pg_config --bindir)" >> "$GITHUB_PATH" @@ -568,22 +622,22 @@ jobs: --health-timeout=5s --health-retries=20 steps: - - uses: actions/checkout@v5 + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Show Rust toolchain run: rustup show active-toolchain - name: Rust cache - uses: Swatinem/rust-cache@v2 + uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: shared-key: rust-ci-${{ runner.os }} workspaces: . -> target - name: Setup sccache - uses: mozilla-actions/sccache-action@v0.0.9 + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad # v0.0.9 - name: Run MySQL migration smoke test env: @@ -674,6 +728,7 @@ jobs: - clippy - test - data_db_smoke + - shell_security if: ${{ always() }} steps: - name: Verify required jobs @@ -681,7 +736,8 @@ jobs: if [ "${{ needs.fmt.result }}" != "success" ] || \ [ "${{ needs.clippy.result }}" != "success" ] || \ [ "${{ needs.test.result }}" != "success" ] || \ - [ "${{ needs.data_db_smoke.result }}" != "success" ]; then + [ "${{ needs.data_db_smoke.result }}" != "success" ] || \ + [ "${{ needs.shell_security.result }}" != "success" ]; then echo "Rust CI failed" exit 1 fi diff --git a/.gitignore b/.gitignore index fab43389a..35ed543b2 100644 --- a/.gitignore +++ b/.gitignore @@ -248,3 +248,5 @@ src/_version.py analysis/ new-api/ apps/aether-tunnel/aether-tunnel.toml +# Generated by frontend/scripts/sync-vscodex.mjs. +frontend/public/aether-vscodex/ diff --git a/Cargo.lock b/Cargo.lock index 32ff749e6..95e1f8529 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -14,7 +14,7 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "generic-array", ] @@ -26,7 +26,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" dependencies = [ "cfg-if", "cipher", - "cpufeatures", + "cpufeatures 0.2.17", ] [[package]] @@ -55,11 +55,11 @@ dependencies = [ "aether-provider-pool", "aether-provider-transport", "axum", - "base64 0.22.1", + "base64", "chrono", "http", "regex", - "reqwest", + "reqwest 0.12.28", "semver", "serde", "serde_json", @@ -82,7 +82,7 @@ name = "aether-ai-formats" version = "0.1.0" dependencies = [ "aether-contracts", - "base64 0.22.1", + "base64", "http", "regex", "serde", @@ -103,7 +103,7 @@ dependencies = [ "aether-pool-core", "aether-scheduler-core", "async-trait", - "base64 0.22.1", + "base64", "http", "serde", "serde_json", @@ -133,7 +133,7 @@ name = "aether-contracts" version = "0.1.0" dependencies = [ "aes-gcm", - "base64 0.22.1", + "base64", "bytes", "flate2", "hmac", @@ -141,6 +141,7 @@ dependencies = [ "serde_json", "sha2", "thiserror 2.0.18", + "url", ] [[package]] @@ -148,7 +149,8 @@ name = "aether-crypto" version = "0.1.0" dependencies = [ "aes", - "base64 0.22.1", + "aws-lc-rs", + "base64", "cbc", "hmac", "pbkdf2", @@ -189,14 +191,20 @@ name = "aether-data-contracts" version = "0.1.0" dependencies = [ "aether-ai-formats", + "aether-contracts", "aether-routing-core", "async-trait", + "base64", + "bcrypt", "chrono", + "chrono-tz", "serde", "serde_json", "sha2", "thiserror 2.0.18", "tokio", + "url", + "uuid", ] [[package]] @@ -322,8 +330,9 @@ dependencies = [ "aether-wallet", "async-stream", "async-trait", + "aws-lc-rs", "axum", - "base64 0.22.1", + "base64", "bcrypt", "brotli", "bytes", @@ -340,13 +349,13 @@ dependencies = [ "hyper-util", "ldap3", "libc", - "md-5", + "md-5 0.10.6", "object_store", "parking_lot", + "percent-encoding", "regex", - "reqwest", - "rsa", - "rustls 0.23.37", + "reqwest 0.12.28", + "rustls", "semver", "serde", "serde_json", @@ -360,6 +369,7 @@ dependencies = [ "tikv-jemalloc-sys", "tikv-jemallocator", "tokio", + "tokio-tungstenite 0.28.0", "tokio-util", "tower", "tower-http", @@ -412,7 +422,7 @@ version = "0.1.0" dependencies = [ "aether-admission-core", "aether-contracts", - "base64 0.22.1", + "base64", "bytes", "http", "serde", @@ -434,8 +444,10 @@ dependencies = [ name = "aether-http" version = "0.1.0" dependencies = [ - "reqwest", + "reqwest 0.12.28", "serde", + "tokio", + "url", ] [[package]] @@ -453,7 +465,7 @@ dependencies = [ "axum", "futures-util", "http", - "reqwest", + "reqwest 0.12.28", "serde", "serde_json", "sha2", @@ -475,7 +487,7 @@ dependencies = [ "futures-util", "http", "libc", - "reqwest", + "reqwest 0.12.28", "serde", "serde_json", "sysinfo", @@ -488,16 +500,17 @@ version = "0.1.0" dependencies = [ "aether-ai-formats", "aether-contracts", + "aether-crypto", "aether-data-contracts", "aether-provider-transport", "aether-scheduler-core", "async-trait", - "base64 0.22.1", + "aws-lc-rs", + "base64", "regex", - "rsa", "serde_json", - "sha2", "tokio", + "url", "uuid", ] @@ -507,9 +520,9 @@ version = "0.1.0" dependencies = [ "aether-contracts", "async-trait", - "base64 0.22.1", + "base64", "http", - "reqwest", + "reqwest 0.12.28", "serde", "serde_json", "sha2", @@ -539,6 +552,7 @@ dependencies = [ name = "aether-provider-pool" version = "0.1.0" dependencies = [ + "aether-contracts", "aether-data-contracts", "aether-pool-core", "aether-provider-transport", @@ -555,19 +569,20 @@ dependencies = [ "aether-contracts", "aether-crypto", "aether-data-contracts", + "aether-http", "aether-oauth", "aether-runtime-state", "aether-video-tasks-core", "async-trait", + "aws-lc-rs", "axum", - "base64 0.22.1", + "base64", "chrono", "crypto_box", "ed25519-dalek", "http", "regex", - "reqwest", - "rsa", + "reqwest 0.12.28", "serde", "serde_json", "sha2", @@ -596,6 +611,7 @@ dependencies = [ "axum", "chrono", "futures-util", + "libc", "serde_json", "sha2", "thiserror 2.0.18", @@ -670,12 +686,14 @@ dependencies = [ name = "aether-testkit" version = "0.1.0" dependencies = [ + "aether-contracts", "aether-data", "aether-gateway", "aether-loadtools", "aether-runtime", "aether-runtime-state", "axum", + "http", "sqlx", "tokio", ] @@ -693,7 +711,7 @@ dependencies = [ "anyhow", "arc-swap", "axum", - "base64 0.22.1", + "base64", "bytes", "clap", "crossterm 0.28.1", @@ -705,8 +723,9 @@ dependencies = [ "hyper-util", "libc", "ratatui", - "reqwest", - "rustls 0.23.37", + "reqwest 0.12.28", + "rustls", + "semver", "serde", "serde_json", "sha2", @@ -715,7 +734,7 @@ dependencies = [ "tar", "thiserror 2.0.18", "tokio", - "tokio-rustls 0.26.4", + "tokio-rustls", "tokio-tungstenite 0.24.0", "toml", "tower-service", @@ -743,7 +762,7 @@ dependencies = [ "aether-data-contracts", "aether-runtime-state", "async-trait", - "base64 0.22.1", + "base64", "futures-util", "serde", "serde_json", @@ -756,6 +775,7 @@ name = "aether-video-tasks-core" version = "0.1.0" dependencies = [ "aether-contracts", + "aether-crypto", "aether-data-contracts", "async-trait", "serde", @@ -858,7 +878,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -869,14 +889,23 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] name = "anyhow" -version = "1.0.102" +version = "1.0.103" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" + +[[package]] +name = "approx" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cab112f0a86d568ea0e627cc1d6be74a1e9cd55214684db5561995f6dad897c6" +dependencies = [ + "num-traits", +] [[package]] name = "arc-swap" @@ -889,9 +918,9 @@ dependencies = [ [[package]] name = "asn1-rs" -version = "0.5.2" +version = "0.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f6fd5ddaf0351dff5b8da21b2fb4ff8e08ddd02857f0bf69c47639106c0fff0" +checksum = "b7f43a50ac4fdca5df8e885c21b835997f0a1cdee65494a6847694a98652d9d8" dependencies = [ "asn1-rs-derive", "asn1-rs-impl", @@ -899,31 +928,31 @@ dependencies = [ "nom", "num-traits", "rusticata-macros", - "thiserror 1.0.69", + "thiserror 2.0.18", "time", ] [[package]] name = "asn1-rs-derive" -version = "0.4.0" +version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "726535892e8eae7e70657b4c8ea93d26b8553afb1ce617caee529ef96d7dee6c" +checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c" dependencies = [ "proc-macro2", "quote", - "syn 1.0.109", - "synstructure 0.12.6", + "syn 2.0.117", + "synstructure", ] [[package]] name = "asn1-rs-impl" -version = "0.1.0" +version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2777730b2039ac0f95f093556e61b6d26cebed5393ca6f152717777cec3a42ed" +checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7" dependencies = [ "proc-macro2", "quote", - "syn 1.0.109", + "syn 2.0.117", ] [[package]] @@ -1030,7 +1059,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b52af3cb4058c895d37317bb27508dccc8e5f2d39454016b297bf4a400597b8" dependencies = [ "axum-core", - "base64 0.22.1", + "base64", "bytes", "form_urlencoded", "futures-util", @@ -1087,12 +1116,6 @@ dependencies = [ "fastrand", ] -[[package]] -name = "base64" -version = "0.21.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9d297deb1925b89f2ccc13d7635fa0714f12c87adce1c75356b39ca9b7178567" - [[package]] name = "base64" version = "0.22.1" @@ -1111,7 +1134,7 @@ version = "0.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2b1866ecef4f2d06a0bb77880015fdf2b89e25a1c2e5addacb87e459c86dc67e" dependencies = [ - "base64 0.22.1", + "base64", "blowfish", "getrandom 0.2.17", "subtle", @@ -1137,7 +1160,7 @@ version = "0.72.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "993776b509cfb49c750f11b8f07a46fa23e0a1386ffc01fb1e7d343efc387895" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "cexpr", "clang-sys", "itertools 0.13.0", @@ -1172,9 +1195,9 @@ checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" [[package]] name = "bitflags" -version = "2.11.0" +version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" dependencies = [ "serde_core", ] @@ -1185,7 +1208,7 @@ version = "0.10.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe" dependencies = [ - "digest", + "digest 0.10.7", ] [[package]] @@ -1197,6 +1220,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" +dependencies = [ + "hybrid-array", +] + [[package]] name = "block-padding" version = "0.3.3" @@ -1234,7 +1266,7 @@ version = "5.0.0-alpha.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "183ccc3854411c035410dcdbffafca62084f3a6c33f013c77e83c025d2a08a28" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "boring-sys2", "foreign-types", "libc", @@ -1268,6 +1300,12 @@ version = "3.20.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" +[[package]] +name = "by_address" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "64fa3c856b712db6612c019f14756e64e4bcea13337a6b33b696333a9eaa2d06" + [[package]] name = "bytemuck" version = "1.25.0" @@ -1337,6 +1375,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.1", + "rand_core 0.10.1", +] + [[package]] name = "chrono" version = "0.4.44" @@ -1367,7 +1416,7 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "inout", "zeroize", ] @@ -1483,15 +1532,6 @@ version = "0.4.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "75984efb6ed102a0d42db99afb6c1948f0380d1d91808d5529916e6c08b49d8d" -[[package]] -name = "concurrent-queue" -version = "2.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" -dependencies = [ - "crossbeam-utils", -] - [[package]] name = "const-oid" version = "0.9.6" @@ -1507,16 +1547,6 @@ dependencies = [ "unicode-segmentation", ] -[[package]] -name = "core-foundation" -version = "0.9.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" -dependencies = [ - "core-foundation-sys", - "libc", -] - [[package]] name = "core-foundation" version = "0.10.1" @@ -1542,6 +1572,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ca28b0ae3115b884660db4118d803791fd6756b6e88f39c0f3f7859060d7566" +dependencies = [ + "libc", +] + [[package]] name = "crc" version = "3.4.0" @@ -1557,6 +1596,16 @@ version = "2.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19d374276b40fb8bbdee95aef7c7fa6b5316ec764510eb64b8dd0e2ed0d7e7f5" +[[package]] +name = "crc-fast" +version = "1.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e75b2483e97a5a7da73ac68a05b629f9c53cff58d8ed1c77866079e18b00dba5" +dependencies = [ + "digest 0.10.7", + "spin 0.10.1", +] + [[package]] name = "crc32fast" version = "1.5.0" @@ -1566,6 +1615,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "critical-section" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" + [[package]] name = "crossbeam-deque" version = "0.8.6" @@ -1578,9 +1633,9 @@ dependencies = [ [[package]] name = "crossbeam-epoch" -version = "0.9.18" +version = "0.9.20" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f" dependencies = [ "crossbeam-utils", ] @@ -1606,7 +1661,7 @@ version = "0.28.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "829d955a0bb380ef178a640b91779e3987da38c9aea133b20614cfed8cdea9c6" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "crossterm_winapi", "mio", "parking_lot", @@ -1622,7 +1677,7 @@ version = "0.29.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d8b9f2e4c67f833b660cdb0a3523065869fb35570177239812ed4c905aeff87b" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "crossterm_winapi", "derive_more", "document-features", @@ -1654,6 +1709,15 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "hybrid-array", +] + [[package]] name = "crypto_box" version = "0.9.1" @@ -1710,9 +1774,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "curve25519-dalek-derive", - "digest", + "digest 0.10.7", "fiat-crypto", "rustc_version", "subtle", @@ -1803,9 +1867,9 @@ dependencies = [ [[package]] name = "der-parser" -version = "8.2.0" +version = "10.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dbd676fbbab537128ef0278adb5576cf363cff6aa22a7b24effe97347cfab61e" +checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6" dependencies = [ "asn1-rs", "displaydoc", @@ -1852,12 +1916,22 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", + "block-buffer 0.10.4", "const-oid", - "crypto-common", + "crypto-common 0.1.7", "subtle", ] +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.1", + "crypto-common 0.2.2", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -1936,7 +2010,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -1961,11 +2035,10 @@ dependencies = [ [[package]] name = "event-listener" -version = "5.4.1" +version = "5.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e13b66accf52311f30a0db42147dadea9850cb48cd070028831ae5f5d4b856ab" +checksum = "5a23add41df1562121a9393cb065eab5146a1242410f23a644851e90cfd669d2" dependencies = [ - "concurrent-queue", "parking", "pin-project-lite", ] @@ -2050,7 +2123,7 @@ checksum = "da0e4dd2a88388a1f4ccc7c9ce104604dab68d9f408dc34cd45823d5a9069095" dependencies = [ "futures-core", "futures-sink", - "spin 0.9.8", + "spin 0.9.9", ] [[package]] @@ -2269,6 +2342,7 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", + "rand_core 0.10.1", "wasip2", "wasip3", ] @@ -2291,9 +2365,9 @@ checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" [[package]] name = "h2" -version = "0.4.13" +version = "0.4.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2f44da3a8150a6703ed5d34e164b875fd14c2cdab9af1252a9a1020bde2bdc54" +checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27" dependencies = [ "atomic-waker", "bytes", @@ -2342,6 +2416,17 @@ dependencies = [ "foldhash 0.2.0", ] +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash 0.2.0", +] + [[package]] name = "hashlink" version = "0.10.0" @@ -2378,7 +2463,7 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" dependencies = [ - "digest", + "digest 0.10.7", ] [[package]] @@ -2467,6 +2552,15 @@ version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "135b12329e5e3ce057a9f972339ea52bc954fe1e9358ef27f95e89716fbc5424" +[[package]] +name = "hybrid-array" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "707114b52a152fa7bdb290cd7cd5912d9467273b6d74e21b8d81aca1f8533f6b" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.8.1" @@ -2499,11 +2593,10 @@ dependencies = [ "http", "hyper", "hyper-util", - "rustls 0.23.37", - "rustls-native-certs 0.8.3", + "rustls", "rustls-pki-types", "tokio", - "tokio-rustls 0.26.4", + "tokio-rustls", "tower-service", "webpki-roots 1.0.6", ] @@ -2514,7 +2607,7 @@ version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" dependencies = [ - "base64 0.22.1", + "base64", "bytes", "futures-channel", "futures-util", @@ -2525,7 +2618,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.3", + "socket2 0.5.10", "tokio", "tower-layer", "tower-service", @@ -2754,12 +2847,70 @@ dependencies = [ "either", ] +[[package]] +name = "itertools" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b4baf93f58d4425749ca49a51c50ebab072c5df6994d08fed93541c331481dc" +dependencies = [ + "either", +] + [[package]] name = "itoa" version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror 2.0.18", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn 2.0.117", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.117", +] + [[package]] name = "jobserver" version = "0.1.34" @@ -2803,14 +2954,14 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" dependencies = [ - "spin 0.9.8", + "spin 0.9.9", ] [[package]] name = "lber" -version = "0.4.2" +version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2df7f9fd9f64cf8f59e1a4a0753fe7d575a5b38d3d7ac5758dcee9357d83ef0a" +checksum = "cbcf559624bfd9fe8d488329a8959766335a43a9b8b2cdd6a2c379fca02909a5" dependencies = [ "bytes", "nom", @@ -2818,25 +2969,23 @@ dependencies = [ [[package]] name = "ldap3" -version = "0.11.5" +version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "166199a8207874a275144c8a94ff6eed5fcbf5c52303e4d9b4d53a0c7ac76554" +checksum = "01fe89f5e7cfb7e4701e3a38ff9f00358e026a9aee940355d88ee9d81e5c7503" dependencies = [ "async-trait", "bytes", "futures", "futures-util", - "lazy_static", "lber", "log", "nom", "percent-encoding", - "ring 0.16.20", - "rustls 0.21.12", - "rustls-native-certs 0.6.3", - "thiserror 1.0.69", + "rustls", + "rustls-native-certs", + "thiserror 2.0.18", "tokio", - "tokio-rustls 0.24.1", + "tokio-rustls", "tokio-stream", "tokio-util", "url", @@ -2877,7 +3026,7 @@ version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "libc", "plain", "redox_syscall 0.7.3", @@ -2900,7 +3049,7 @@ version = "0.3.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5f4de44e98ddbf09375cbf4d17714d18f39195f4f4894e8524501726fd9a8a4a" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", ] [[package]] @@ -2944,11 +3093,11 @@ checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" [[package]] name = "lru" -version = "0.16.3" +version = "0.18.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" +checksum = "0d317b4b9eb398e6acce275758ec6125535505e7a146fb1a9b8bda2451b0ff4c" dependencies = [ - "hashbrown 0.16.1", + "hashbrown 0.17.1", ] [[package]] @@ -2989,7 +3138,17 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" dependencies = [ "cfg-if", - "digest", + "digest 0.10.7", +] + +[[package]] +name = "md-5" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69b6441f590336821bb897fb28fc622898ccceb1d6cea3fde5ea86b090c4de98" +dependencies = [ + "cfg-if", + "digest 0.11.3", ] [[package]] @@ -3063,7 +3222,7 @@ version = "0.29.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "71e2746dc3a24dd78b3cfcb7be93368c6de9963d30f43a6a73998a9cf4b17b46" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "cfg-if", "cfg_aliases", "libc", @@ -3095,7 +3254,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -3119,7 +3278,7 @@ dependencies = [ "num-integer", "num-iter", "num-traits", - "rand 0.8.5", + "rand 0.8.6", "smallvec", "zeroize", ] @@ -3182,28 +3341,32 @@ dependencies = [ [[package]] name = "object_store" -version = "0.12.5" +version = "0.14.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fbfbfff40aeccab00ec8a910b57ca8ecf4319b335c542f2edcd19dd25a1e2a00" +checksum = "d354792e39fa5f0009e47623cf8b15b099bf9a652fa55c6f817fe28ac84fea50" dependencies = [ "async-trait", - "base64 0.22.1", + "aws-lc-rs", + "base64", "bytes", "chrono", + "crc-fast", "form_urlencoded", - "futures", + "futures-channel", + "futures-core", + "futures-util", "http", "http-body-util", "humantime", "hyper", - "itertools 0.14.0", - "md-5", + "itertools 0.15.0", + "md-5 0.11.0", "parking_lot", "percent-encoding", "quick-xml", - "rand 0.9.2", - "reqwest", - "ring 0.17.14", + "rand 0.10.2", + "reqwest 0.13.4", + "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", @@ -3217,9 +3380,9 @@ dependencies = [ [[package]] name = "oid-registry" -version = "0.6.1" +version = "0.8.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9bedf36ffb6ba96c2eb7144ef6270557b52e54b20c0a8e1eb2ff99a6c6959bff" +checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7" dependencies = [ "asn1-rs", ] @@ -3253,12 +3416,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "openssl-probe" -version = "0.1.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e" - [[package]] name = "openssl-probe" version = "0.2.1" @@ -3274,6 +3431,39 @@ dependencies = [ "num-traits", ] +[[package]] +name = "palette" +version = "0.7.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddeed8580d347d2abf3dcf06a5f0b3dc020258338526b277847cd4248a70fc64" +dependencies = [ + "approx", + "libm", + "palette_derive", + "palette_math", +] + +[[package]] +name = "palette_derive" +version = "0.7.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88537020289b719d81be994ccf1bbf4990f477e2f69ee52fe3e45f43a02e56be" +dependencies = [ + "by_address", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "palette_math" +version = "0.7.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e6eb142958d64335fb0e345c5b9ead2ecd6fc438c307e9d7d3c4fd428dbaf12" +dependencies = [ + "libm", +] + [[package]] name = "parking" version = "2.2.1" @@ -3309,7 +3499,7 @@ version = "0.12.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" dependencies = [ - "digest", + "digest 0.10.7", "hmac", ] @@ -3407,7 +3597,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d" dependencies = [ "phf_shared 0.11.3", - "rand 0.8.5", + "rand 0.8.6", ] [[package]] @@ -3492,7 +3682,7 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" dependencies = [ - "cpufeatures", + "cpufeatures 0.2.17", "opaque-debug", "universal-hash", ] @@ -3504,7 +3694,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "opaque-debug", "universal-hash", ] @@ -3560,9 +3750,9 @@ dependencies = [ [[package]] name = "quick-xml" -version = "0.38.4" +version = "0.41.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b66c2058c55a409d601666cffe35f04333cf1013010882cec174a7467cd4e21c" +checksum = "e660451e55124f798a69a5af3f49ccfbefbd41910eefd25caf2393e1f3473ec1" dependencies = [ "memchr", "serde", @@ -3580,8 +3770,8 @@ dependencies = [ "quinn-proto", "quinn-udp", "rustc-hash", - "rustls 0.23.37", - "socket2 0.6.3", + "rustls", + "socket2 0.5.10", "thiserror 2.0.18", "tokio", "tracing", @@ -3590,17 +3780,18 @@ dependencies = [ [[package]] name = "quinn-proto" -version = "0.11.14" +version = "0.11.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "434b42fec591c96ef50e21e886936e66d3cc3f737104fdb9b737c40ffb94c098" +checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e" dependencies = [ + "aws-lc-rs", "bytes", "getrandom 0.3.4", "lru-slab", - "rand 0.9.2", - "ring 0.17.14", + "rand 0.9.3", + "ring", "rustc-hash", - "rustls 0.23.37", + "rustls", "rustls-pki-types", "slab", "thiserror 2.0.18", @@ -3618,7 +3809,7 @@ dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2 0.6.3", + "socket2 0.5.10", "tracing", "windows-sys 0.60.2", ] @@ -3646,9 +3837,9 @@ checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" [[package]] name = "rand" -version = "0.8.5" +version = "0.8.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404" +checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" dependencies = [ "libc", "rand_chacha 0.3.1", @@ -3657,14 +3848,25 @@ dependencies = [ [[package]] name = "rand" -version = "0.9.2" +version = "0.9.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6db2770f06117d490610c7488547d543617b21bfa07796d7a12f6f1bd53850d1" +checksum = "7ec095654a25171c2124e9e3393a930bddbffdc939556c914957a4c3e0a87166" dependencies = [ "rand_chacha 0.9.0", "rand_core 0.9.5", ] +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.2", + "rand_core 0.10.1", +] + [[package]] name = "rand_chacha" version = "0.3.1" @@ -3704,32 +3906,42 @@ dependencies = [ ] [[package]] -name = "ratatui" -version = "0.30.0" +name = "rand_core" +version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d1ce67fb8ba4446454d1c8dbaeda0557ff5e94d39d5e5ed7f10a65eb4c8266bc" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + +[[package]] +name = "ratatui" +version = "0.30.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3274ba0a2c5e1bcad2a2005d20f4dc59dad26b2eb0940fb094500dba4099d57d" dependencies = [ "instability", "ratatui-core", "ratatui-crossterm", "ratatui-macros", + "ratatui-termina", "ratatui-termwiz", "ratatui-widgets", + "serde", ] [[package]] name = "ratatui-core" -version = "0.1.0" +version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5ef8dea09a92caaf73bff7adb70b76162e5937524058a7e5bff37869cbbec293" +checksum = "cbb175c433c8e28a809d1f5773a2ae96e68c0ce40db865cbab1020bf33ae479c" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "compact_str", - "hashbrown 0.16.1", - "indoc", + "critical-section", + "hashbrown 0.17.1", "itertools 0.14.0", "kasuari", "lru", + "palette", + "serde", "strum", "thiserror 2.0.18", "unicode-segmentation", @@ -3739,9 +3951,9 @@ dependencies = [ [[package]] name = "ratatui-crossterm" -version = "0.1.0" +version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "577c9b9f652b4c121fb25c6a391dd06406d3b092ba68827e6d2f09550edc54b3" +checksum = "567584a3b0e6a8203c23de40b4861497266725eb5363dbfd18a1edd603cca9f0" dependencies = [ "cfg-if", "crossterm 0.29.0", @@ -3751,19 +3963,30 @@ dependencies = [ [[package]] name = "ratatui-macros" -version = "0.7.0" +version = "0.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a7f1342a13e83e4bb9d0b793d0ea762be633f9582048c892ae9041ef39c936f4" +checksum = "ed7dc68daa7498a43e4d68e0eb078427e10c38fbcfbb1e42d955f1fa2140d814" dependencies = [ "ratatui-core", "ratatui-widgets", ] [[package]] -name = "ratatui-termwiz" +name = "ratatui-termina" version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0f76fe0bd0ed4295f0321b1676732e2454024c15a35d01904ddb315afd3d545c" +checksum = "c0bf912d9e66f057a759d92e386a280ea886b352ab757d6ac4d653c7ed2c43c2" +dependencies = [ + "instability", + "ratatui-core", + "termina", +] + +[[package]] +name = "ratatui-termwiz" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "faf03e0380b7744054d6cb74224fe3adf062a029754933f575ca1e3b4c2ce977" dependencies = [ "ratatui-core", "termwiz", @@ -3771,17 +3994,18 @@ dependencies = [ [[package]] name = "ratatui-widgets" -version = "0.3.0" +version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d7dbfa023cd4e604c2553483820c5fe8aa9d71a42eea5aa77c6e7f35756612db" +checksum = "66e3d19bcc9130ca376277d93b60767ff121ace3be06f5f95f81dd68956407d1" dependencies = [ - "bitflags 2.11.0", - "hashbrown 0.16.1", + "bitflags 2.13.1", + "hashbrown 0.17.1", "indoc", "instability", "itertools 0.14.0", "line-clipping", "ratatui-core", + "serde", "strum", "time", "unicode-segmentation", @@ -3837,7 +4061,7 @@ version = "0.5.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", ] [[package]] @@ -3846,7 +4070,7 @@ version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce70a74e890531977d37e532c34d45e9055d2409ed08ddba14529471ed0be16" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", ] [[package]] @@ -3884,7 +4108,7 @@ version = "0.12.28" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ - "base64 0.22.1", + "base64", "bytes", "futures-core", "futures-util", @@ -3900,15 +4124,14 @@ dependencies = [ "percent-encoding", "pin-project-lite", "quinn", - "rustls 0.23.37", - "rustls-native-certs 0.8.3", + "rustls", "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", "sync_wrapper", "tokio", - "tokio-rustls 0.26.4", + "tokio-rustls", "tokio-util", "tower", "tower-http", @@ -3916,24 +4139,48 @@ dependencies = [ "url", "wasm-bindgen", "wasm-bindgen-futures", - "wasm-streams", + "wasm-streams 0.4.2", "web-sys", "webpki-roots 1.0.6", ] [[package]] -name = "ring" -version = "0.16.20" +name = "reqwest" +version = "0.13.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3053cf52e236a3ed746dfc745aa9cacf1b791d846bdaf412f60a8d7d6e17c8fc" +checksum = "219c5811de6525e5416c7d5d53bb656d3afdbc6c5af816e0802bcfa42dbdc1c3" dependencies = [ - "cc", - "libc", - "once_cell", - "spin 0.5.2", - "untrusted 0.7.1", + "base64", + "bytes", + "futures-core", + "futures-util", + "h2", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "rustls-platform-verifier", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tokio-util", + "tower", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams 0.5.0", "web-sys", - "winapi", ] [[package]] @@ -3946,7 +4193,7 @@ dependencies = [ "cfg-if", "getrandom 0.2.17", "libc", - "untrusted 0.9.0", + "untrusted", "windows-sys 0.52.0", ] @@ -3957,7 +4204,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b8573f03f5883dcaebdfcf4725caa1ecb9c15b2ef50c43a07b816e06799bb12d" dependencies = [ "const-oid", - "digest", + "digest 0.10.7", "num-bigint-dig", "num-integer", "num-traits", @@ -4000,7 +4247,7 @@ version = "0.38.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "errno", "libc", "linux-raw-sys 0.4.15", @@ -4013,23 +4260,11 @@ version = "1.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.61.2", -] - -[[package]] -name = "rustls" -version = "0.21.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3f56a14d1f48b391359b22f731fd4bd7e43c97f3c50eee276f3aa09c94784d3e" -dependencies = [ - "log", - "ring 0.17.14", - "rustls-webpki 0.101.7", - "sct", + "windows-sys 0.59.0", ] [[package]] @@ -4041,44 +4276,23 @@ dependencies = [ "aws-lc-rs", "log", "once_cell", - "ring 0.17.14", + "ring", "rustls-pki-types", - "rustls-webpki 0.103.9", + "rustls-webpki", "subtle", "zeroize", ] -[[package]] -name = "rustls-native-certs" -version = "0.6.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00" -dependencies = [ - "openssl-probe 0.1.6", - "rustls-pemfile", - "schannel", - "security-framework 2.11.1", -] - [[package]] name = "rustls-native-certs" version = "0.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "612460d5f7bea540c490b2b6395d8e34a953e52b491accd6c86c8164c5932a63" dependencies = [ - "openssl-probe 0.2.1", + "openssl-probe", "rustls-pki-types", "schannel", - "security-framework 3.7.0", -] - -[[package]] -name = "rustls-pemfile" -version = "1.0.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1c74cae0a4cf6ccbbf5f359f08efdf8ee7e1dc532573bf0db71968cb56b1448c" -dependencies = [ - "base64 0.21.7", + "security-framework", ] [[package]] @@ -4092,25 +4306,42 @@ dependencies = [ ] [[package]] -name = "rustls-webpki" -version = "0.101.7" +name = "rustls-platform-verifier" +version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b6275d1ee7a1cd780b64aca7726599a1dbc893b1e64144529e55c3c2f745765" +checksum = "26d1e2536ce4f35f4846aa13bff16bd0ff40157cdb14cc056c7b14ba41233ba0" dependencies = [ - "ring 0.17.14", - "untrusted 0.9.0", + "core-foundation", + "core-foundation-sys", + "jni", + "log", + "once_cell", + "rustls", + "rustls-native-certs", + "rustls-platform-verifier-android", + "rustls-webpki", + "security-framework", + "security-framework-sys", + "webpki-root-certs", + "windows-sys 0.59.0", ] [[package]] -name = "rustls-webpki" -version = "0.103.9" +name = "rustls-platform-verifier-android" +version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53" +checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" dependencies = [ "aws-lc-rs", - "ring 0.17.14", + "ring", "rustls-pki-types", - "untrusted 0.9.0", + "untrusted", ] [[package]] @@ -4134,6 +4365,15 @@ dependencies = [ "cipher", ] +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + [[package]] name = "schannel" version = "0.1.29" @@ -4160,37 +4400,14 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" -[[package]] -name = "sct" -version = "0.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "da046153aa2352493d6cb7da4b6e5c0c057d8a1d0a9aa8560baffdd945acd414" -dependencies = [ - "ring 0.17.14", - "untrusted 0.9.0", -] - -[[package]] -name = "security-framework" -version = "2.11.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "897b2245f0b511c87893af39b033e5ca9cce68824c4d7e7630b5a1d339658d02" -dependencies = [ - "bitflags 2.11.0", - "core-foundation 0.9.4", - "core-foundation-sys", - "libc", - "security-framework-sys", -] - [[package]] name = "security-framework" version = "3.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" dependencies = [ - "bitflags 2.11.0", - "core-foundation 0.10.1", + "bitflags 2.13.1", + "core-foundation", "core-foundation-sys", "libc", "security-framework-sys", @@ -4295,8 +4512,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", ] [[package]] @@ -4312,8 +4529,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", ] [[package]] @@ -4368,7 +4585,7 @@ version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" dependencies = [ - "digest", + "digest 0.10.7", "rand_core 0.6.4", ] @@ -4378,6 +4595,22 @@ version = "0.3.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e320a6c5ad31d271ad523dcf3ad13e2767ad8b1cb8f047f75a8aeaf8da139da2" +[[package]] +name = "simd_cesu8" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "siphasher" version = "1.0.2" @@ -4416,24 +4649,24 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] name = "spin" -version = "0.5.2" +version = "0.9.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e63cff320ae2c57904679ba7cb63280a3dc4613885beafb148ee7bf9aa9042d" - -[[package]] -name = "spin" -version = "0.9.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6980e8d7511241f8acf4aebddbb1ff938df5eebe98691418c4468d0b72a96a67" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" dependencies = [ "lock_api", ] +[[package]] +name = "spin" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "023a211cb3138dbc438680b32560ad89f699977624c9f8dbb95a47d5b4c07dd3" + [[package]] name = "spki" version = "0.7.3" @@ -4463,7 +4696,7 @@ version = "0.8.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ee6798b1838b6a0f69c007c133b8df5866302197e404e8b6ee8ed3e3a5e68dc6" dependencies = [ - "base64 0.22.1", + "base64", "bigdecimal", "bytes", "chrono", @@ -4482,7 +4715,7 @@ dependencies = [ "memchr", "once_cell", "percent-encoding", - "rustls 0.23.37", + "rustls", "serde", "serde_json", "sha2", @@ -4540,14 +4773,14 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "aa003f0038df784eb8fecbbac13affe3da23b45194bd57dba231c8f48199c526" dependencies = [ "atoi", - "base64 0.22.1", + "base64", "bigdecimal", - "bitflags 2.11.0", + "bitflags 2.13.1", "byteorder", "bytes", "chrono", "crc", - "digest", + "digest 0.10.7", "dotenvy", "either", "futures-channel", @@ -4560,11 +4793,11 @@ dependencies = [ "hmac", "itoa", "log", - "md-5", + "md-5 0.10.6", "memchr", "once_cell", "percent-encoding", - "rand 0.8.5", + "rand 0.8.6", "rsa", "serde", "sha1", @@ -4584,9 +4817,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "db58fcd5a53cf07c184b154801ff91347e4c30d17a3562a635ff028ad5deda46" dependencies = [ "atoi", - "base64 0.22.1", + "base64", "bigdecimal", - "bitflags 2.11.0", + "bitflags 2.13.1", "byteorder", "chrono", "crc", @@ -4601,11 +4834,11 @@ dependencies = [ "home", "itoa", "log", - "md-5", + "md-5 0.10.6", "memchr", "num-bigint", "once_cell", - "rand 0.8.5", + "rand 0.8.6", "serde", "serde_json", "sha2", @@ -4673,18 +4906,18 @@ checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" [[package]] name = "strum" -version = "0.27.2" +version = "0.28.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "af23d6f6c1a224baef9d3f61e287d2761385a5b88fdab4eb4c6f11aeb54c4bcf" +checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" dependencies = [ "strum_macros", ] [[package]] name = "strum_macros" -version = "0.27.2" +version = "0.28.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7695ce3845ea4b33927c055a39dc438a45b059f7c1b3d91d38d10355fb8cbca7" +checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" dependencies = [ "heck", "proc-macro2", @@ -4729,18 +4962,6 @@ dependencies = [ "futures-core", ] -[[package]] -name = "synstructure" -version = "0.12.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f36bdaa60a83aca3921b5259d5400cbf5e90fc51931376a9bd4a0eb79aa7210f" -dependencies = [ - "proc-macro2", - "quote", - "syn 1.0.109", - "unicode-xid", -] - [[package]] name = "synstructure" version = "0.13.2" @@ -4777,6 +4998,19 @@ dependencies = [ "xattr", ] +[[package]] +name = "termina" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9048a889effe34a5cddee0af7f53285198b16dca3be510858d38dfdb3e62a04e" +dependencies = [ + "bitflags 2.13.1", + "parking_lot", + "rustix 1.1.4", + "signal-hook", + "windows-sys 0.60.2", +] + [[package]] name = "terminfo" version = "0.9.0" @@ -4805,8 +5039,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4676b37242ccbd1aabf56edb093a4827dc49086c0ffd764a5705899e0f35f8f7" dependencies = [ "anyhow", - "base64 0.22.1", - "bitflags 2.11.0", + "base64", + "bitflags 2.13.1", "fancy-regex", "filedescriptor", "finl_unicode", @@ -5005,23 +5239,13 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "tokio-rustls" -version = "0.24.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c28327cf380ac148141087fbfb9de9d7bd4e84ab5d2c28fbc911d753de8a7081" -dependencies = [ - "rustls 0.21.12", - "tokio", -] - [[package]] name = "tokio-rustls" version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ - "rustls 0.23.37", + "rustls", "tokio", ] @@ -5056,10 +5280,10 @@ checksum = "edc5f74e248dc973e0dbb7b74c7e0d6fcc301c694ff50049504004ef4d0cdcd9" dependencies = [ "futures-util", "log", - "rustls 0.23.37", + "rustls", "rustls-pki-types", "tokio", - "tokio-rustls 0.26.4", + "tokio-rustls", "tungstenite 0.24.0", "webpki-roots 0.26.11", ] @@ -5072,10 +5296,10 @@ checksum = "d25a406cddcc431a75d3d9afc6a7c0f7428d4891dd973e4d54c56b46127bf857" dependencies = [ "futures-util", "log", - "rustls 0.23.37", + "rustls", "rustls-pki-types", "tokio", - "tokio-rustls 0.26.4", + "tokio-rustls", "tungstenite 0.28.0", "webpki-roots 0.26.11", ] @@ -5157,7 +5381,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8" dependencies = [ "async-compression", - "bitflags 2.11.0", + "bitflags 2.13.1", "bytes", "futures-core", "futures-util", @@ -5284,8 +5508,8 @@ dependencies = [ "http", "httparse", "log", - "rand 0.8.5", - "rustls 0.23.37", + "rand 0.8.6", + "rustls", "rustls-pki-types", "sha1", "thiserror 1.0.69", @@ -5303,8 +5527,8 @@ dependencies = [ "http", "httparse", "log", - "rand 0.9.2", - "rustls 0.23.37", + "rand 0.9.3", + "rustls", "rustls-pki-types", "sha1", "thiserror 2.0.18", @@ -5333,9 +5557,9 @@ dependencies = [ [[package]] name = "typenum" -version = "1.19.0" +version = "1.20.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" [[package]] name = "ucd-trie" @@ -5411,16 +5635,10 @@ version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "subtle", ] -[[package]] -name = "untrusted" -version = "0.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a156c684c91ea7d62626509bce3cb4e1d9ed5c4d978f7b4352658f96a4c26b4a" - [[package]] name = "untrusted" version = "0.9.0" @@ -5498,6 +5716,16 @@ dependencies = [ "utf8parse", ] +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + [[package]] name = "want" version = "0.3.1" @@ -5631,13 +5859,26 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasm-streams" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + [[package]] name = "wasmparser" version = "0.244.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" dependencies = [ - "bitflags 2.11.0", + "bitflags 2.13.1", "hashbrown 0.15.5", "indexmap", "semver", @@ -5788,6 +6029,15 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.48.0", +] + [[package]] name = "winapi-x86_64-pc-windows-gnu" version = "0.4.0" @@ -6151,7 +6401,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" dependencies = [ "anyhow", - "bitflags 2.11.0", + "bitflags 2.13.1", "indexmap", "log", "serde", @@ -6239,9 +6489,9 @@ checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" [[package]] name = "x509-parser" -version = "0.15.1" +version = "0.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7069fba5b66b9193bd2c5d3d4ff12b839118f6bcbef5328efafafb5395cf63da" +checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202" dependencies = [ "asn1-rs", "data-encoding", @@ -6250,7 +6500,7 @@ dependencies = [ "nom", "oid-registry", "rusticata-macros", - "thiserror 1.0.69", + "thiserror 2.0.18", "time", ] @@ -6284,7 +6534,7 @@ dependencies = [ "proc-macro2", "quote", "syn 2.0.117", - "synstructure 0.13.2", + "synstructure", ] [[package]] @@ -6325,7 +6575,7 @@ dependencies = [ "proc-macro2", "quote", "syn 2.0.117", - "synstructure 0.13.2", + "synstructure", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index ae184e1cb..7763dffe8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -101,6 +101,7 @@ aether-runtime = { path = "crates/aether-runtime/base" } aether-testkit = { path = "crates/aether-testing/testkit" } aes = "0.8" aes-gcm = "0.10" +aws-lc-rs = { version = "1.16.2", default-features = false, features = ["alloc", "aws-lc-sys"] } async-stream = "0.3" async-trait = "0.1" axum = "0.8" @@ -117,8 +118,9 @@ flate2 = "1" futures-util = "0.3" hmac = "0.12" http = "1" -object_store = { version = "0.12", default-features = false, features = ["aws"] } +object_store = { version = "0.14.1", default-features = false, features = ["aws"] } pbkdf2 = { version = "0.12", default-features = false, features = ["hmac"] } +percent-encoding = "2" reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "http2", "socks"] } redis = { version = "0.28", default-features = false, features = ["tokio-comp", "script", "streams", "connection-manager"] } regex = "1" diff --git a/Dockerfile.app b/Dockerfile.app index 901af65ed..f17913c5b 100644 --- a/Dockerfile.app +++ b/Dockerfile.app @@ -10,20 +10,24 @@ # --- layout stage: create /opt/aether directory structure with symlink --- # distroless has no shell, so we use busybox to set up the symlink. -FROM busybox:1.37-musl AS layout +FROM busybox:1.37.0-musl@sha256:fc6dddc4c44b1bfe37f41cae8e67d1693828e8f42a91862816d7953e2c9d3f23 AS layout ARG TARGETARCH RUN mkdir -p /opt/aether/releases/image/bin /opt/aether/releases/image/frontend /opt/aether/logs COPY dist/aether-gateway-${TARGETARCH} /opt/aether/releases/image/bin/aether-gateway -RUN chmod 0755 /opt/aether/releases/image/bin/aether-gateway COPY dist/frontend/ /opt/aether/releases/image/frontend/ +# Keep the immutable release root-owned while guaranteeing that the runtime +# identity can traverse and read every packaged asset. +RUN chmod -R u=rwX,go=rX /opt/aether/releases/image \ + && chmod 0755 /opt/aether/releases/image/bin/aether-gateway + RUN ln -s /opt/aether/releases/image /opt/aether/current # --- final stage: distroless runtime --- -FROM gcr.io/distroless/static-debian12 +FROM gcr.io/distroless/static-debian12@sha256:6447365a6337c3732f412d1b74357b30a633831955b2bc45552b0086be907687 COPY --from=layout /opt/aether /opt/aether @@ -31,6 +35,7 @@ WORKDIR /opt/aether ENV RUST_LOG=aether_gateway=info \ APP_PORT=8084 \ + HOME=/tmp/aether-home \ AETHER_UPDATE_STRATEGY=docker \ AETHER_GATEWAY_STATIC_DIR=/opt/aether/current/frontend @@ -39,5 +44,5 @@ EXPOSE 8084 HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \ CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"] -USER root +USER 65532:65532 ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"] diff --git a/Dockerfile.app.local b/Dockerfile.app.local index 6c0971e6d..3067d37e4 100644 --- a/Dockerfile.app.local +++ b/Dockerfile.app.local @@ -11,6 +11,15 @@ FROM ${NODE_BASE_IMAGE} AS frontend-builder ARG AETHER_BUILD_VERSION ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \ AETHER_VERSION=${AETHER_BUILD_VERSION} +WORKDIR /app/aether-vscodex/web +COPY aether-vscodex/web/package*.json ./ +RUN --mount=type=cache,id=aether-vscodex-npm-cache,target=/root/.npm,sharing=locked \ + npm config set registry https://registry.npmmirror.com && \ + npm ci --no-audit --no-fund +COPY aether-vscodex/public /app/aether-vscodex/public +COPY aether-vscodex/web/ ./ +RUN npm run build + WORKDIR /app/frontend COPY frontend/package*.json ./ RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \ diff --git a/Dockerfile.app.release-local b/Dockerfile.app.release-local index 8eddbd57a..169d79184 100644 --- a/Dockerfile.app.release-local +++ b/Dockerfile.app.release-local @@ -11,6 +11,15 @@ FROM ${NODE_BASE_IMAGE} AS frontend-builder ARG AETHER_BUILD_VERSION ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \ AETHER_VERSION=${AETHER_BUILD_VERSION} +WORKDIR /app/aether-vscodex/web +COPY aether-vscodex/web/package*.json ./ +RUN --mount=type=cache,id=aether-vscodex-npm-cache,target=/root/.npm,sharing=locked \ + npm config set registry https://registry.npmmirror.com && \ + npm ci --no-audit --no-fund +COPY aether-vscodex/public /app/aether-vscodex/public +COPY aether-vscodex/web/ ./ +RUN npm run build + WORKDIR /app/frontend COPY frontend/package*.json ./ RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \ diff --git a/Makefile b/Makefile index 5d5978c37..3c3de9d1b 100644 --- a/Makefile +++ b/Makefile @@ -6,7 +6,7 @@ DEV_RUST_LOG := $(RUST_LOG) endif export DEV_RUST_LOG -.PHONY: dev dev-backend dev-frontend migration backfill +.PHONY: dev dev-backend dev-frontend db-status db-prepare migration backfill define DEV_BACKEND_SCRIPT set -euo pipefail @@ -20,6 +20,13 @@ set -a source .env set +a +if [[ -n "$${ADMIN_EMAIL:-}" || -n "$${ADMIN_USERNAME:-}" || -n "$${ADMIN_PASSWORD:-}" ]]; then + if [[ -z "$${ADMIN_USERNAME:-}" || -z "$${ADMIN_PASSWORD:-}" ]]; then + echo "=> 管理员自举配置不完整,请在 .env 中设置 ADMIN_USERNAME 和 ADMIN_PASSWORD" + exit 1 + fi +fi + dotenv_has_key() { local key="$$1" grep -Eq "^[[:space:]]*$${key}=" .env @@ -201,12 +208,17 @@ print_startup_failure_hint() { if [ -n "$${log_file}" ] && [ -f "$${log_file}" ]; then if grep -Eq "database schema is behind" "$${log_file}"; then - echo "=> 检测到数据库 schema 落后,请执行: make migration" + echo "=> 检测到数据库尚未准备完成,请执行: make db-prepare" return fi if grep -Eq "database backfills are behind" "$${log_file}"; then - echo "=> 检测到待执行 backfills,请执行: make backfill" + echo "=> 检测到数据库尚未准备完成,请执行: make db-prepare" + return + fi + + if grep -Eq "bootstrap admin env is partially configured.*ADMIN_PASSWORD" "$${log_file}"; then + echo "=> 首次启动需要管理员密码,请在 .env 中设置 ADMIN_PASSWORD" return fi fi @@ -344,6 +356,9 @@ if ! ensure_dev_infra; then exit 1 fi +echo "=> 编译 aether-gateway..." +cargo build -p aether-gateway --bin aether-gateway + GATEWAY_PID="" GATEWAY_LOG_DIR="" GATEWAY_LOG_FILE="" @@ -352,8 +367,8 @@ create_gateway_log_file echo "=> 启动 aether-gateway (Rust frontdoor: 0.0.0.0:$${APP_PORT})..." echo "=> 日志过滤: $${RUST_LOG}" -echo "=> 执行命令: cargo run -p aether-gateway -- --app-port $${APP_PORT}" -cargo run -p aether-gateway -- --app-port "$${APP_PORT}" > >( +echo "=> 执行命令: target/debug/aether-gateway --app-port $${APP_PORT}" +target/debug/aether-gateway --app-port "$${APP_PORT}" > >( tee -a "$${GATEWAY_LOG_FILE}" ) 2>&1 & GATEWAY_PID=$$! @@ -444,7 +459,7 @@ if [ -f .env ]; then fi export APP_PORT="$${APP_PORT:-8084}" -echo "=> 启动后端: RUST_LOG=$${DEV_RUST_LOG} cargo run -p aether-gateway -- --app-port $${APP_PORT:-8084}" +echo "=> 启动后端: 先编译 aether-gateway,再运行 target/debug/aether-gateway --app-port $${APP_PORT:-8084}" /bin/bash -euo pipefail -c "$$DEV_BACKEND_SCRIPT" & backend_pid=$$! @@ -494,8 +509,14 @@ export DEV_SCRIPT define DB_TASK_SCRIPT set -euo pipefail -if [ -z "$${DB_TASK_FLAG:-}" ] || [ -z "$${DB_TASK_LABEL:-}" ]; then - echo "=> 内部错误: DB_TASK_FLAG / DB_TASK_LABEL 未设置" +if [ -z "$${DB_TASK_COMMAND:-}" ] || [ -z "$${DB_TASK_LABEL:-}" ]; then + echo "=> 内部错误: DB_TASK_COMMAND / DB_TASK_LABEL 未设置" + exit 1 +fi + +read -r -a db_task_args <<< "$${DB_TASK_COMMAND}" +if [ "$${#db_task_args[@]}" -eq 0 ]; then + echo "=> 内部错误: DB_TASK_COMMAND 为空" exit 1 fi @@ -546,8 +567,8 @@ if ! command -v cargo >/dev/null 2>&1; then exit 1 fi -echo "=> 执行 $${DB_TASK_LABEL}: cargo run -p aether-gateway -- $${DB_TASK_FLAG}" -exec cargo run -p aether-gateway -- "$${DB_TASK_FLAG}" +echo "=> 执行 $${DB_TASK_LABEL}: cargo run -p aether-gateway --bin aether-gateway -- $${db_task_args[*]}" +exec cargo run -p aether-gateway --bin aether-gateway -- "$${db_task_args[@]}" endef export DB_TASK_SCRIPT @@ -560,8 +581,14 @@ dev-backend: dev-frontend: @cd frontend && npm run dev +db-status: + @DB_TASK_COMMAND="db status" DB_TASK_LABEL="数据库状态检查" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" + +db-prepare: + @DB_TASK_COMMAND="db prepare" DB_TASK_LABEL="数据库准备" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" + migration: - @DB_TASK_FLAG=--migrate DB_TASK_LABEL="数据库迁移" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" + @DB_TASK_COMMAND="--migrate" DB_TASK_LABEL="数据库迁移" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" backfill: - @DB_TASK_FLAG=--apply-backfills DB_TASK_LABEL="数据库 backfill" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" + @DB_TASK_COMMAND="--apply-backfills" DB_TASK_LABEL="数据库 backfill" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT" diff --git a/README.md b/README.md index 73e9100ab..425cd6a30 100644 --- a/README.md +++ b/README.md @@ -44,17 +44,30 @@ cd Aether # 2. 配置环境变量 cp .env.example .env -# 生成 JWT_SECRET_KEY / ENCRYPTION_KEY, 并填入 .env +# .env 包含数据库、JWT 和数据加密密钥,先限制为仅当前用户可读写 +chmod 600 .env +# 生成 JWT / 加密 / Postgres / Redis / MySQL 独立随机密钥,并填入 .env ./generate_keys.sh # 编辑 .env 设置 ADMIN_PASSWORD # 3. 首次部署 / 更新 (从以下部署形态任选其一) -# Postgres + Redis (适用于企业或多人使用) +# Postgres + Redis (推荐) docker compose pull && docker compose up -d -# Single Node (适用于个人用户或朋友分享) +# Single Node:默认容器身份为 65532:65532,先停止旧容器并检查/迁移 SQLite bind 目录 +docker compose -f docker-compose.single-node.yml stop app +mkdir -p data +test -z "$(find data ! -type d ! -type f -print -quit)" || { echo "data 中存在 symlink/FIFO/socket/device,拒绝迁移" >&2; exit 1; } +test -z "$(find data -type f -links +1 -print -quit)" || { echo "data 中存在硬链接,拒绝迁移" >&2; exit 1; } +sudo chown -R -P 65532:65532 ./data +sudo find data -type d -exec chmod 0700 {} + +sudo find data -type f -exec chmod 0600 {} + docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d ``` +应用镜像默认以固定非 root 身份 `65532:65532` 运行;Compose 同时移除全部 Linux capabilities、禁止提权、启用只读根文件系统,并仅提供带 `nosuid,nodev,noexec` 的 `/tmp` 临时文件系统。若宿主机不适合使用固定 UID/GID,可在 `.env` 中把 `AETHER_CONTAINER_UID` / `AETHER_CONTAINER_GID` 改成其他非零数字身份,并让 Single Node 的 `./data` 归该身份所有。`install.sh --mode compose-single-node` 会按安装用户自动生成这两个值;使用 `sudo` 运行时会采用原调用用户身份,并安全迁移已有 SQLite 数据。 + +从旧版 root 容器升级 Single Node 时,必须先停止旧 `app` 容器,再在第一次启动新版 Compose 前完成一次数据目录迁移;安装器检测到容器仍在运行会拒绝迁移,避免并发改写造成检查竞态。停止容器后用 `sudo` 重新执行一键安装器会自动处理;非 root 安装器发现旧数据所有权不匹配时会拒绝启动并提示迁移,不会放宽目录权限。手工部署且仍使用默认身份时执行上面的检查、`chown` 和 `find ... chmod` 命令即可。迁移只改变 `./data` 的所有权和权限,不会删除数据库、WAL 或备份文件。 + ### 一键更新 Docker Compose 部署后,可在部署目录直接执行: @@ -69,10 +82,24 @@ Docker Compose 部署后,可在部署目录直接执行: ./update.sh --mode single-node ``` -仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免正式发布镜像切换到非 root 用户后再被宿主机挂载日志目录的权限问题拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,并额外挂载一个容器用户可写的目录到 `/opt/aether/logs`。 +仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免非 root 用户被宿主机日志目录权限拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,额外挂载目录到 `/opt/aether/logs`,并让它归 `.env` 中配置的容器 UID/GID 所有;只读根文件系统不会阻止显式可写挂载。 管理后台右上角“版本信息”会检测新版本。Docker Compose 部署只提示版本,实际更新继续执行 `./update.sh`;systemd / launchd / 二进制部署才使用后台自更新,流程是下载对应平台的 GitHub Release 包、强制校验 `SHA256SUMS`、解压到 `/opt/aether/releases/`,再切换 `/opt/aether/current` 并退出进程,交给 systemd / launchd 拉起新版本。 +正式 Release 还会发布由 GitHub Actions OIDC / Sigstore 签发的 SLSA build provenance。需要验证发布者身份时,下载目标 tarball 和 `AETHER_RELEASE_PROVENANCE.sigstore.json`,并把 `TAG` 设置为对应 Release tag: + +```bash +gh attestation verify "aether-${TAG}-linux-amd64.tar.gz" \ + --repo fawney19/Aether \ + --signer-workflow fawney19/Aether/.github/workflows/release.yml \ + --source-ref "refs/tags/${TAG}" \ + --bundle AETHER_RELEASE_PROVENANCE.sigstore.json +``` + +`docker-compose.yml` 中的官方 PostgreSQL、Redis 和 MySQL 镜像均固定到多架构 OCI index digest。升级这些依赖时应在发布变更中显式更新 digest,避免同名 tag 在无人审查的情况下改变部署内容。 + +正式发布到 GHCR 和 Docker Hub 的多架构 Aether 镜像也带有同一 GitHub Actions OIDC / Sigstore provenance;生产 `Dockerfile.app` 的 BusyBox 与 Distroless 基础镜像同样固定到多架构 OCI index digest。 + 源码或本地构建版本不会启用后台在线更新,请继续使用源码更新流程。Docker Compose 用户如果希望“容器重建后也保持镜像层面的新版本”,仍建议定期运行 `./update.sh` 拉取并重建 app 镜像。服务器访问 GitHub 需要代理时,可设置 `AETHER_UPDATE_PROXY_URL`,也兼容 `UPDATE_PROXY_URL`、`HTTPS_PROXY`、`ALL_PROXY`、`HTTP_PROXY` 以及 `NO_PROXY`。共享出口触发 GitHub API 限流时,可设置只读 `AETHER_UPDATE_GITHUB_TOKEN`,也兼容 `GITHUB_TOKEN` / `GH_TOKEN`。下载总超时默认 600 秒,连续无响应/无数据默认 30 秒,可通过 `AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS` 和 `AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS` 调整。 标准 Docker Compose 使用 Docker named volumes 存放 Postgres/Redis/MySQL 数据;Single Node 使用部署目录下的 `./data` 存放 SQLite 数据。 @@ -128,6 +155,7 @@ Docker Compose 用户可在部署目录的 `.env` 中设置 `APP_IMAGE=ghcr.io/f ## 本地开发 依赖 Docker、Rust toolchain、Node.js 和 make。 +首次启动前需要在 `.env` 中设置 `ADMIN_PASSWORD`,用于创建本地管理员。 ```bash make dev @@ -135,6 +163,18 @@ make dev `make dev` 会同时启动后端 `aether-gateway` 和前端 `frontend` 的 Vite dev server。需要单独启动时可使用 `make dev-backend` 或 `make dev-frontend`。 Postgres / Redis 本地依赖未就绪时,`make dev` 会自动执行 `docker compose up -d postgres redis`。 +`make dev` 会先完成后端编译,再开始计算服务健康检查超时。数据库 schema 和必要的派生数据准备也会在启动时自动完成;通常不需要手动区分 migration 与 backfill。升级不会主动重写或清除已有业务历史记录,新写入会直接遵循当前的数据持久化策略。排查或部署前预执行时可使用: + +```bash +make db-status +make db-prepare +``` + +## Codex 远程协同 + +`aether-vscodex/` 是独立的 VS Code Codex 协同模块:同步模式跟随 VS Code 官方 Codex 面板当前会话且不另起进程;异步模式使用独立 app-server,让浏览器自行列出、恢复、新建和切换会话。两种模式都能从本机 URL 或 Aether 云端查看输出、发送消息和处理授权,模块内的 Vue 前端提供中英文界面。 + +安装、云端配对和安全边界请参阅 [`aether-vscodex/README.md`](aether-vscodex/README.md)。 ## Aether Tunnel (可选) @@ -159,21 +199,42 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙 - `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时 SQLite 固定 `1/1`,Postgres/MySQL 按每核 `4` 条自动推导,总池范围为 `32-100`。该预算按进程计算,多实例部署应按数据库连接上限显式分配 - `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调 - `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB` -- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms` -- `AETHER_MAX_REQUEST_BODY_MB`:可选的单请求解压后请求体上限;未配置或设为 `0` 时不限制 -- `AETHER_MAX_INTERNAL_BUFFERED_BODY_MB`:可选的 heartbeat、管理探测等内部整包响应体上限;未配置或设为 `0` 时不限制 +- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:可选的请求体完整读取超时;默认或显式设为 `0` 时关闭,非零值限制在 `1000-600000ms` +- `AETHER_MAX_REQUEST_BODY_MB`:单请求解压后请求体上限,默认 `256MB`;显式设为 `0` 表示不再收紧默认值,但仍受 `256MB` 安全硬上限约束 +- `AETHER_MAX_INTERNAL_BUFFERED_BODY_MB`:heartbeat、管理探测等内部整包响应体上限,默认 `64MB`;显式设为 `0` 表示不再收紧默认值,但仍受 `256MB` 安全硬上限约束 - `AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY`:隧道节点状态上报队列容量,默认 `1024`;满载时拒绝新事件,避免控制面故障导致无界内存增长 +- `AETHER_TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS`:跨网关 owner relay 解析到私有/保留地址时的显式运维开关,默认关闭;仅当多网关 relay URL 是受控的内网 HTTPS 地址时设置为 `true`。它不改变普通 provider 请求的 DNS/代理策略,也不允许明文 HTTP 非 loopback relay +- `AETHER_TUNNEL_RELAY_PRIVATE_HOST_ALLOWLIST`:更窄的 owner relay 私网例外,填写逗号分隔的精确主机名(例如 `gateway-a.internal,gateway-b.internal`,忽略大小写和末尾点);仅这些主机解析出的私有地址会被允许,并且请求仍使用解析后地址 pin。不要填写通配符或 `.internal` 这类后缀 +- `AETHER_INTERNAL_GATEWAY_AUTH_SECRET`:旧版 `/api/internal/gateway/*` 高权限控制面的独立 HMAC 密钥,至少 `32` 字节;未配置时该控制面返回 `404`。不要复用 JWT、数据加密或 tunnel relay 密钥,多节点必须使用同一值及共享 Redis 防重放 - `AETHER_GATEWAY_SECURITY_CACHE_TTL_MS`:IP 黑白名单本地缓存时间,默认 `1000ms`,写操作会主动失效相关缓存 -- `AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB`:可选的 PII 恢复同步响应缓冲上限;未配置或设为 `0` 时不限制 +- `AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB`:PII 恢复同步响应缓冲上限,默认 `64MB`;显式设为 `0` 表示不再收紧默认值,但仍受 `256MB` 安全硬上限约束 - `REDIS_URL`:Redis 连接串;仅 Postgres + Redis 的 Docker Compose 部署需要配置 - `AETHER_RUNTIME_BACKEND=memory|redis`:运行时缓存/协调后端。SQLite 默认用 `memory`,不会连接 Redis;多节点部署和需要跨 gateway 重启恢复 OpenAI Responses continuation history 的部署必须使用共享 Redis -- `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:常规启动前自动执行挂起的 schema migration 和 backfill;仓库自带的 `docker-compose.yml` 默认开启 +- `AETHER_GATEWAY_DATABASE_MODE=auto|verify-only`:数据库启动策略,默认 `auto`,自动完成挂起的 schema migration 和 backfill;`verify-only` 仅检查并在数据库落后时拒绝启动 +- `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:旧版兼容开关;新配置请使用 `AETHER_GATEWAY_DATABASE_MODE` - `JWT_SECRET_KEY` / `ENCRYPTION_KEY`:认证和敏感数据加密所需密钥 +- `AETHER_BACKUP_ENCRYPTION_KEY`:推荐的 S3 备份独立加密密钥;缺省回退到 `ENCRYPTION_KEY`。新备份使用带 key ID 的 AES-256-GCM v2 envelope,轮换前必须保留旧密钥 - `API_KEY_PREFIX`:用户和管理员新建 API Key 时使用的前缀,默认 `sk` - `ADMIN_USERNAME` / `ADMIN_PASSWORD` / `ADMIN_EMAIL`:首次启动时自举首个本地管理员;`install.sh` 会提示输入管理员密码 - `CORS_ORIGINS` / `CORS_ALLOW_CREDENTIALS`:前端跨域来源控制;如果要跨域带登录 Cookie,`CORS_ORIGINS` 不能写 `*` - `RUST_LOG`:Rust 日志过滤,例如 `aether_gateway=info`、`aether_gateway=debug,sqlx=warn` -- Docker Compose 的 `DB_PASSWORD` / `REDIS_PASSWORD` 默认使用 `aether` +- `DB_PASSWORD` / `REDIS_PASSWORD` / `MYSQL_PASSWORD` / `MYSQL_ROOT_PASSWORD`:Docker Compose 后端密码,首次安装时分别随机生成;手工部署必须替换示例占位值,不要互相复用 + +### S3 备份离线恢复 + +先从 S3 下载完整的 `.json.zst.aes256gcm` 对象,再使用原始的完整 S3 object key 做认证解密。恢复工具只验证并输出本地 JSON,不会直接写数据库;数据库导入仍应在维护窗口通过管理端完成。 + +```bash +AETHER_BACKUP_ENCRYPTION_KEY='原备份密钥' \ + cargo run -p aether-gateway --bin aether-backup-restore -- \ + --input ./backup.json.zst.aes256gcm \ + --object-key 'aether/backups/aether-data-backup-20260822-010000.json.zst.aes256gcm' \ + --output ./restored-backup.json +``` + +工具默认拒绝覆盖,输出采用原子写并在 Unix 上设置为 `0600`;Unix 可用 `--overwrite` 原子替换,Windows 为避免非原子删除窗口会要求选择新输出路径。密钥不能作为命令行参数。可使用 `AETHER_BACKUP_ENCRYPTION_KEY`、兼容用 `AETHER_GATEWAY_DATA_ENCRYPTION_KEY` / `ENCRYPTION_KEY`、受保护的 `--key-file`,或 `AETHER_BACKUP_KEYRING_FILE`。Keyring JSON 格式为 `{"version":1,"keys":["当前或历史 v2 secret"],"legacy_v1":["旧 v1 secret"]}`;条目也可写成 `{"secret":"..."}`(兼容字段名 `key`)。也可由 `AETHER_BACKUP_HISTORICAL_KEYS_JSON` 提供同一结构。密钥文件必须是非符号链接的普通文件,Unix 下权限需为 `0600` 或更严格。 + +默认限制密文为 `512MiB`、解压后 JSON 为 `1GiB`,可通过受限的 `--max-encrypted-mib` / `--max-json-mib` 调整。网关最多扫描同一备份前缀下 10,000 个对象,并且不会自动删除 S3 对象:`backup_s3_retention_count` 只用于报告超出保留数量的清理候选。旧明文备份在创建并验证加密副本后仍会保留,必须通过 bucket lifecycle 或支持版本条件的外部清理工具移除;启用 Versioning 时还需清理 noncurrent versions,Object Lock/retention 可能阻止物理删除。 --- diff --git a/aether-vscodex/.dockerignore b/aether-vscodex/.dockerignore new file mode 100644 index 000000000..9739dbad4 --- /dev/null +++ b/aether-vscodex/.dockerignore @@ -0,0 +1,10 @@ +.git +.github +node_modules +test +fixtures +vscode-extension +*.vsix +coverage +data +.DS_Store diff --git a/aether-vscodex/.gitignore b/aether-vscodex/.gitignore new file mode 100644 index 000000000..890f9f2b7 --- /dev/null +++ b/aether-vscodex/.gitignore @@ -0,0 +1,8 @@ +node_modules/ +vscode-extension/node_modules/ +vscode-extension/dist/ +data/ +coverage/ +*.vsix +.DS_Store +*.log diff --git a/aether-vscodex/Dockerfile b/aether-vscodex/Dockerfile new file mode 100644 index 000000000..5e8b74fc5 --- /dev/null +++ b/aether-vscodex/Dockerfile @@ -0,0 +1,26 @@ +FROM node:22-alpine + +ENV NODE_ENV=production \ + HOST=0.0.0.0 \ + PORT=8788 \ + AETHER_VSCODEX_DATA_DIR=/var/lib/aether-vscodex + +WORKDIR /app + +COPY package.json package-lock.json ./ +RUN npm ci --omit=dev && npm cache clean --force + +COPY cloud ./cloud +COPY relay ./relay +COPY public ./public + +RUN mkdir -p /var/lib/aether-vscodex && chown -R node:node /var/lib/aether-vscodex /app + +USER node + +EXPOSE 8788 + +HEALTHCHECK --interval=30s --timeout=5s --start-period=5s --retries=3 \ + CMD node -e "fetch('http://127.0.0.1:8788/healthz').then(r=>{if(!r.ok)process.exit(1)}).catch(()=>process.exit(1))" + +CMD ["node", "cloud/server.js"] diff --git a/aether-vscodex/README.md b/aether-vscodex/README.md new file mode 100644 index 000000000..035e4309a --- /dev/null +++ b/aether-vscodex/README.md @@ -0,0 +1,283 @@ +# aether-vscodex + +这个项目让浏览器从本机 URL 或 Aether 云端查看、输入并处理 Codex 会话,提供两种 +可随时切换的控制模式。默认的**同步模式**通过官方扩展使用的本机 IPC socket,严格 +跟随 VS Code Codex 面板当前会话,不启动另一个 `codex` 进程;**异步模式**由伴随扩展 +启动独立 app-server,网页可以自行列出、恢复、新建和切换会话。 + +同一个伴随扩展可同时连接两个互不替代的通道:本机 loopback 控制台和部署在 +Aether 中的云端控制台。本机通道默认免密码且只能从本机访问;云端通道使用 +Aether 登录鉴权、一次性浏览器票据和独立设备凭据;父页面不会通过协议把 Aether JWT +传给 iframe 或 Node sidecar。iframe 是随 Aether 一起发布的同源受信代码,不应被视为 +隔离不受信内容的安全边界。 + +`vscode-extension/codex-remote-collab-0.4.0.vsix` 安装到 VS Code 后,会作为官方 +`openai.chatgpt` Codex 扩展的伴随扩展,并自动托管只监听本机的 relay。开发时仍可 +单独运行 `relay/server.js`。不要卸载或替换官方 Codex 扩展。 + +## 工作方式 + +```text +官方 VS Code Codex 会话 + │ 本机私有 IPC(只在 VS Code 所在机器上) + ▼ + ┌── 本机 relay ── http://127.0.0.1:8787 +VS Code aether-vscodex 扩展 ────┤ + └── Aether gateway ── 用户/设备隔离的云端 relay +``` + +控制模式与传输通道是两个独立维度:切换同步/异步不会重连本地或云端 relay。本机和 +Aether 控制页连接到同一台 VS Code 主机时,会看到同一个当前模式。 + +| 控制模式 | 会话所有者 | 网页会话导航 | +| --- | --- | --- | +| 同步 | 官方 VS Code Codex 面板 | 禁止网页自行切换;自动跟随 VS Code | +| 异步 | 扩展启动的独立 app-server | 可列出、恢复、新建和切换会话 | + +浏览器的 `operator` 可以发送任务、继续/中断当前 turn,并处理 Codex 的审批、 +用户输入和 MCP elicitation;`viewer` 只能查看事件和输出。远程浏览器不接触 +VS Code 的 SecretStorage,也不直接连接 IPC socket。 + +## 前端结构 + +Aether 页面使用仓库既有的 Vue 3、TypeScript、Vite 和 i18n。独立控制台也提供 +Vue/Vite 源码入口,但当前高保真的会话渲染与协议状态机作为兼容运行时保留,构建到 +`public/` 后同时供本机 URL 和 Aether 同源 iframe 使用。这样不需要一次性重写并丢失 +命令展开、滚动锚点、思考状态、Markdown、子代理、模型和权限菜单等已有行为。 + +界面支持 `zh-CN` 与 `en-US`。Aether 的语言和深浅色主题会通过经过来源校验的 +`postMessage` 同步给 iframe;VS Code 命令与设置说明使用 `package.nls` 本地化。 + +## Aether 云端部署 + +云端模式由 Aether gateway 和独立 Node sidecar 组成。sidecar 只在 Compose 内网暴露 +8788,公网的 HTTP、配对交换和 WebSocket 都经 Aether gateway: + +```text +GET /api/users/me/vscodex/devices +POST /api/users/me/vscodex/pairings +DELETE /api/users/me/vscodex/devices/:device_id +POST /api/users/me/vscodex/ws-tickets +POST /api/vscodex/pair +WS /api/vscodex/ws +``` + +生成至少 32 字节的内部令牌,并按 Aether 的公开 HTTPS 地址设置变量: + +```sh +export AETHER_VSCODEX_INTERNAL_TOKEN="$(openssl rand -base64 32)" +export AETHER_VSCODEX_PUBLIC_WS_URL="wss://aether.example.com/api/vscodex/ws" +export AETHER_VSCODEX_ALLOWED_ORIGINS="https://aether.example.com" + +docker compose \ + -f docker-compose.yml \ + -f docker-compose.local.yml \ + -f aether-vscodex/docker-compose.aether.yml \ + up -d --build +``` + +源码部署必须包含 `docker-compose.local.yml`,以保证 gateway、前端和 sidecar 来自同一份 +checkout。使用发布镜像时可以去掉该文件,但 `APP_IMAGE` 必须固定为包含相同 +`aether-vscodex` 协议版本的 Aether 镜像,不能把当前 sidecar 与旧的 `latest` gateway 混用。 + +首次使用源码 Compose 前先构建控制台;正式 Aether 发布流程与 Dockerfile 已自动执行 +同一步骤: + +```sh +npm --prefix aether-vscodex/web ci +npm --prefix aether-vscodex/web run build +``` + +第一阶段 sidecar 是有状态单副本:设备凭据的 scrypt 哈希保存在 +`vscodex_data`,短期配对码、60 秒一次性浏览器票据和在线房间保存在内存。不要在未引入 +共享连接目录前横向扩容 sidecar。 + +登录 Aether 后打开“Codex 远程控制”,生成一次性配对码。然后在 VS Code 命令面板执行 +**Codex Remote: Pair with Aether**,填写 Aether 地址和配对码。插件会把设备凭据写入 +VS Code SecretStorage,并同时保持本机控制台连接。 + +## 快速开始 + +前提:Node.js 20+;官方 `openai.chatgpt` VS Code 扩展已安装并登录;目标会话 +已经在 VS Code 的 Codex 面板中打开。VS Code 和 relay 必须以同一个操作系统用户 +运行,因为 IPC socket 是本机文件。 + +1. 安装依赖并构建伴随扩展: + + ```sh + npm --prefix vscode-extension install + npm --prefix vscode-extension run build + ``` + + 本机 `ws://` 地址会由扩展自动启动 relay;loopback 模式默认不需要 token,且 + `host` 模式不会启动 `codex app-server`。 + +2. 安装 `vscode-extension/codex-remote-collab-0.4.0.vsix`(或在扩展目录先 + `npm run build` 再用 `npx --yes @vscode/vsce package` 打包),然后在 VS Code + 执行 **Developer: Reload Window**。 + +3. 在 VS Code 设置中填写: + + ```json + { + "codexRemoteCollab.localRelayUrl": "ws://127.0.0.1:8787/v1/connect", + "codexRemoteCollab.controlMode": "sync", + "codexRemoteCollab.autoDiscoverThread": true, + "codexRemoteCollab.autoStart": true + } + ``` + +4. 执行一次 **Developer: Reload Window** 后,扩展会自动找到最近的、仍由官方 + VS Code Codex owner 持有的会话,并把已有输出同步到 relay;如果没有自动启动, + 无需手动启动或断开。右下角状态项只用于显示状态并打开 Web。需要精确指定会话时,执行 + **Codex Remote: Set Existing Thread ID**;留空则恢复自动发现。 + 官方 Codex 面板切换会话时,Web 默认会在新会话快照就绪后自动跟随;正在执行或等待 + 授权的旧会话会先保持附着,结束后再安全切换。 + +5. 浏览器打开 `http://127.0.0.1:8787`,页面会自动以本机 operator 身份连接, + 不需要输入密码。 + +如果页面显示“等待 VS Code 主机连接”,先确认 relay 地址与扩展设置的端口完全一致, +然后在 VS Code 执行一次 **Developer: Reload Window**。同步模式必须在官方 Codex +面板已经打开至少一个会话后才能发现 owner;通常不需要手工填写 +`codexRemoteCollab.threadId`,留空会自动选择最近的可用会话。若之前填写过已经关闭的 +thread ID,清空该设置后再重载窗口。 + +### 发布与下载插件 + +正式发布时不需要用户在本地编译。仓库的 `.github/workflows/release.yml` 在推送 +`vX.Y.Z`、`vX.Y.Z-beta.N` 或 `vX.Y.Z-rc.N` 标签时,会在 GitHub Actions 中完成 Web +前端构建、扩展编译和 VSIX 打包,并把 +`aether-vscodex-.vsix` 附加到对应的 GitHub Release。用户从 Release +页面下载该 VSIX,在 VS Code 的扩展视图中选择“从 VSIX 安装...”即可;安装后执行一次 +**Developer: Reload Window**。 + +手动运行该 workflow 时,VSIX 会作为 `aether-vscodex-vsix` Actions artifact 提供下载, +但不会创建 GitHub Release。源码目录中的 VSIX 只用于本地开发验证,不是用户发布渠道。 + +如果命令面板提示 `command 'codexRemoteCollab.start' not found`,通常是旧版 +VSIX 激活失败(旧包可能没有包含 `ws` 运行依赖)。请安装当前的 +`codex-remote-collab-0.4.0.vsix` 并使用 `--force` 覆盖旧版本,然后执行一次 +**Developer: Reload Window**: + +```sh +code --install-extension vscode-extension/codex-remote-collab-0.4.0.vsix --force +``` + +也可以在 **Output → Codex Remote Collaboration** 中确认没有 +`Cannot find module 'ws'`;出现该错误时,说明扩展尚未成功激活。 + +网页现在按官方 Codex Webview 的会话模型展示:历史和实时输出在中间消息流,用户、 +助手、reasoning、命令输出分别投影为对应的消息项;助手内容支持安全的 Markdown、 +代码块和复制操作,reasoning/命令活动可折叠。底部 composer 使用可编辑富文本区域, +回车发送、Shift+Enter 换行;审批和用户输入会以内嵌 card 出现在会话流中,支持风险 +标记、输入控件、授权范围和明确的允许/拒绝动作。附着适配器会额外发送可选的 +`messages` 角色投影,旧版 host 没有该字段时网页仍回退到纯文本快照。 + +页面打开后自动连接并在断线后重连,不再需要手动点击“连接”或“断开”。同步模式下 +会话列表、返回历史和新建入口会被禁用,所有输入都发送到 VS Code 当前会话。这里复刻的是从本机已安装 +官方 bundle 审计出的布局、状态和交互;官方 bundle 依赖 VS Code 私有 Webview API, +不能安全地直接作为 iframe 嵌入浏览器。 + +底部的“同步 / 异步”分段控件发送 `control/mode/set`。当前 turn 正在执行或存在待处理 +授权、用户输入时,主机拒绝切换;候选适配器启动失败时保留原模式和原会话。切入异步 +模式后,页面顶部会恢复会话历史、新建和选择入口;`session/list` 映射到 +`thread/list`,选择会话使用 `thread/resume` 并水合完整历史,新建会话使用 +`thread/start`。切回同步模式会关闭独立 app-server,并重新以 VS Code 面板为唯一 +会话导航来源。 + +### 认证(可选) + +如果以后需要保护 relay,可显式开启认证;本机流程默认不需要这些变量: + +```sh +CODEX_REMOTE_AUTH=required \ +CODEX_REMOTE_HOST_TOKEN='host-only-secret' \ +CODEX_REMOTE_TOKEN='browser-operator-secret' \ +CODEX_REMOTE_VIEW_TOKEN='browser-viewer-secret' \ +CODEX_REMOTE_MODE=host npm start +``` + +认证开启后,Host token 填在 VS Code 扩展中,Operator/Viewer token 填在浏览器中。 + +## `spawn codex ENOENT` 是什么 + +这个错误只表示某处正在尝试启动**独立**的 `codex app-server`,但 VS Code 图形 +进程的 `PATH` 找不到可执行文件。对于本项目默认的同步模式,不会调用 +`spawn codex`,因此不需要通过设置 `codexCommand` 来修复它。 + +只有切换到异步模式(或仍使用旧版兼容设置)才需要独立可执行文件: + +```json +"codexRemoteCollab.controlMode": "async" +``` + +扩展会优先解析 `codexRemoteCollab.codexCommand`,并可回退到官方 Codex 扩展内置的 +可执行文件;`codexRemoteCollab.codexArgs` 默认是 `["app-server", "--stdio"]`。 +旧 `mode=attach/spawn` 会分别迁移为 `sync/async`。 + +## Relay 模式 + +### `host`(推荐) + +relay 只负责认证、事件缓存和转发;VS Code 扩展通过私有 IPC 附着官方 Codex +会话。必须先打开目标会话;本机 loopback 默认不需要 host token,只有显式开启认证时 +才把 host token 提供给扩展。 + +### `embedded`(旧的独立进程模式) + +只有显式设置 `CODEX_REMOTE_MODE=embedded` 时,relay 才会启动自己的 +`codex app-server --stdio`,适合测试页面和公开 app-server 协议;它与 VS Code +当前会话无关: + +```sh +CODEX_REMOTE_MODE=embedded CODEX_CWD="$PWD" npm start +``` + +`CODEX_BIN` 可指定独立进程的可执行文件;`CODEX_ARGS_JSON` 可覆盖其参数。不要 +把这些设置误认为 attach 模式的必要配置。 + +## HTTP API + +认证开启时,除 `/api/health` 外的 `/api/*` 都需要 +`Authorization: Bearer ` 或 `X-Codex-Token`;本机免认证 +模式下 loopback 请求直接作为 operator 处理。 + +```text +GET /api/health +GET /api/state +GET /api/events?fromSeq=0 +POST /api/command {"commandId":"...","method":"turn/start","params":{...}} +POST /api/respond {"requestId":"...","result":{...}} +``` + +host 模式下,同步控制会拒绝 `thread/start` 和网页会话导航;异步控制会把它们转给 +独立 app-server。浏览器使用 `threadId` 发送 `turn/start`、`turn/steer` 或 +`turn/interrupt`。认证开启时写操作和 +响应请求必须使用 operator token;本机免认证模式下 loopback operator 可直接操作。 + +## 私有协议和限制 + +- IPC follower 协议是官方 VS Code 扩展的私有、带版本号实现,不是公开 API;官方 + 扩展升级后可能需要同步适配。启用 `codexRemoteCollab.ipcStrictVersions` + 时,未知 stream 版本会让连接报错而不是猜测执行。 +- 自动发现只把本地 rollout 元数据当作候选,最终仍通过 IPC owner discovery + 验证;生产或多会话场景建议设置明确的 `threadId`。 +- relay 默认只监听 loopback,且 loopback 默认免认证;这意味着同一台机器上能访问 + loopback 的本地进程都可能控制会话,不要把它反向代理或暴露到外部。如果开启 token + 认证,token 是 bearer secret。高风险授权默认被 host policy 拒绝,只有显式设置 + `codexRemoteCollab.allowHighRiskApprovals=true` 才允许。 +- 输出会做常见 token/密码脱敏,但不能识别所有秘密;不要把凭据发送给 Codex。 +- 当前 UI 控制一个 host 会话,不提供多人同时编辑或文件同步。 + +## 测试 + +根目录测试使用假的 stdio app-server,不会向真实 Codex 发送任务: + +```sh +npm test +cd vscode-extension && npm run check && npm run build +``` + +要验证真实附着,只读地打开官方 VS Code 会话后启动 bridge;不要在验证脚本中 +调用 `turn/start`,除非你确实要向该会话发送任务。 diff --git a/aether-vscodex/cloud/server.js b/aether-vscodex/cloud/server.js new file mode 100644 index 000000000..08724945e --- /dev/null +++ b/aether-vscodex/cloud/server.js @@ -0,0 +1,727 @@ +"use strict"; + +const crypto = require("node:crypto"); +const fs = require("node:fs"); +const http = require("node:http"); +const net = require("node:net"); +const path = require("node:path"); +const { URL } = require("node:url"); +const { WebSocket, WebSocketServer } = require("ws"); + +const { CodexRelay } = require("../relay/server.js"); + +const MAX_JSON_BYTES = 64 * 1024; +const MAX_WS_BYTES = 16 * 1024 * 1024; +const DEFAULT_PAIRING_TTL_MS = 10 * 60 * 1000; +const DEFAULT_TICKET_TTL_MS = 60 * 1000; +const DEFAULT_ROOM_IDLE_MS = 30 * 60 * 1000; + +class DeviceStore { + constructor(filePath) { + this.filePath = filePath; + this.data = { version: 1, devices: [] }; + this.load(); + } + + load() { + try { + const parsed = JSON.parse(fs.readFileSync(this.filePath, "utf8")); + if (parsed?.version !== 1 || !Array.isArray(parsed.devices)) throw new Error("unsupported device store format"); + this.data = parsed; + } catch (error) { + if (error?.code !== "ENOENT") throw error; + fs.mkdirSync(path.dirname(this.filePath), { recursive: true, mode: 0o700 }); + this.persist(); + } + } + + list(userId, connectedDeviceIds = new Set()) { + return this.data.devices + .filter((device) => device.user_id === userId && !device.revoked_at) + .map((device) => publicDevice(device, connectedDeviceIds.has(device.id))); + } + + create(userId, name) { + const id = crypto.randomUUID(); + const secret = crypto.randomBytes(32).toString("base64url"); + const salt = crypto.randomBytes(16).toString("base64url"); + const now = new Date().toISOString(); + const device = { + id, + user_id: userId, + name: normalizeName(name), + secret_salt: salt, + secret_hash: deriveSecret(secret, salt), + created_at: now, + last_seen_at: null, + revoked_at: null, + }; + this.data.devices.push(device); + this.persist(); + return { device: publicDevice(device, false), token: `avx1.${id}.${secret}` }; + } + + authenticate(token) { + const parsed = parseDeviceToken(token); + if (!parsed) return null; + const device = this.data.devices.find((candidate) => candidate.id === parsed.id && !candidate.revoked_at); + if (!device) return null; + const actual = Buffer.from(deriveSecret(parsed.secret, device.secret_salt), "base64url"); + const expected = Buffer.from(device.secret_hash, "base64url"); + if (actual.length !== expected.length || !crypto.timingSafeEqual(actual, expected)) return null; + return device; + } + + get(userId, deviceId) { + return this.data.devices.find((device) => device.user_id === userId && device.id === deviceId && !device.revoked_at) || null; + } + + touch(deviceId) { + const device = this.data.devices.find((candidate) => candidate.id === deviceId && !candidate.revoked_at); + if (!device) return; + device.last_seen_at = new Date().toISOString(); + this.persist(); + } + + revoke(userId, deviceId) { + const device = this.get(userId, deviceId); + if (!device) return false; + device.revoked_at = new Date().toISOString(); + this.persist(); + return true; + } + + persist() { + fs.mkdirSync(path.dirname(this.filePath), { recursive: true, mode: 0o700 }); + const temporary = `${this.filePath}.${process.pid}.${crypto.randomBytes(4).toString("hex")}.tmp`; + fs.writeFileSync(temporary, `${JSON.stringify(this.data, null, 2)}\n`, { mode: 0o600 }); + fs.renameSync(temporary, this.filePath); + } +} + +class EphemeralCredentials { + constructor(options = {}) { + this.pairingTtlMs = options.pairingTtlMs || DEFAULT_PAIRING_TTL_MS; + this.ticketTtlMs = options.ticketTtlMs || DEFAULT_TICKET_TTL_MS; + this.pairings = new Map(); + this.tickets = new Map(); + } + + createPairing(userId, requestedName) { + const code = pairingCode(); + const record = { + id: crypto.randomUUID(), + code, + user_id: userId, + requested_name: normalizeName(requestedName), + expires_at_ms: Date.now() + this.pairingTtlMs, + }; + this.pairings.set(normalizePairingCode(code), record); + return record; + } + + consumePairing(code) { + const key = normalizePairingCode(code); + const record = this.pairings.get(key); + this.pairings.delete(key); + if (!record || record.expires_at_ms <= Date.now()) return null; + return record; + } + + createTicket(userId, deviceId) { + const ticket = `avt1.${crypto.randomBytes(32).toString("base64url")}`; + this.tickets.set(ticket, { + user_id: userId, + device_id: deviceId, + expires_at_ms: Date.now() + this.ticketTtlMs, + }); + return ticket; + } + + consumeTicket(ticket) { + const record = this.tickets.get(ticket); + this.tickets.delete(ticket); + if (!record || record.expires_at_ms <= Date.now()) return null; + return record; + } + + cleanup() { + const now = Date.now(); + for (const [key, record] of this.pairings) if (record.expires_at_ms <= now) this.pairings.delete(key); + for (const [key, record] of this.tickets) if (record.expires_at_ms <= now) this.tickets.delete(key); + } +} + +class RoomManager { + constructor(options = {}) { + this.rooms = new Map(); + this.pendingRooms = new Map(); + this.revokedRoomKeys = new Set(); + this.idleMs = options.idleMs || DEFAULT_ROOM_IDLE_MS; + } + + key(userId, deviceId) { + return `${encodeURIComponent(userId)}:${deviceId}`; + } + + async get(userId, deviceId) { + const key = this.key(userId, deviceId); + if (this.revokedRoomKeys.has(key)) throw httpError(401, "device revoked"); + let room = this.rooms.get(key); + if (!room && this.pendingRooms.has(key)) room = await this.pendingRooms.get(key); + if (!room) { + const creating = this.createRoom(key, userId, deviceId); + this.pendingRooms.set(key, creating); + try { + room = await creating; + } finally { + this.pendingRooms.delete(key); + } + } + if (this.revokedRoomKeys.has(key)) throw httpError(401, "device revoked"); + room.lastActiveMs = Date.now(); + return room; + } + + async createRoom(key, userId, deviceId) { + const hostToken = randomToken(); + const operatorToken = randomToken(); + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + authRequired: true, + hostToken, + operatorToken, + viewerToken: randomToken(), + }); + await relay.start(); + if (this.revokedRoomKeys.has(key)) { + await relay.stop().catch(() => undefined); + throw httpError(401, "device revoked"); + } + const address = relay.address(); + const room = { + key, + userId, + deviceId, + relay, + hostToken, + operatorToken, + baseUrl: `ws://127.0.0.1:${address.port}`, + connections: 0, + lastActiveMs: Date.now(), + }; + this.rooms.set(key, room); + return room; + } + + connectedDeviceIds(userId) { + return new Set([...this.rooms.values()] + .filter((room) => room.userId === userId && room.relay.state.hostConnected) + .map((room) => room.deviceId)); + } + + retain(room) { + room.connections += 1; + room.lastActiveMs = Date.now(); + } + + release(room) { + room.connections = Math.max(0, room.connections - 1); + room.lastActiveMs = Date.now(); + } + + async cleanup() { + const now = Date.now(); + for (const [key, room] of this.rooms) { + if (room.connections > 0 || now - room.lastActiveMs < this.idleMs) continue; + this.rooms.delete(key); + await room.relay.stop(); + } + } + + async revoke(userId, deviceId) { + const key = this.key(userId, deviceId); + this.revokedRoomKeys.add(key); + const pending = this.pendingRooms.get(key); + if (pending) await pending.catch(() => undefined); + const room = this.rooms.get(key); + if (!room) return; + this.rooms.delete(key); + await room.relay.stop(); + } + + async stop() { + await Promise.allSettled([...this.pendingRooms.values()]); + this.pendingRooms.clear(); + const rooms = [...this.rooms.values()]; + this.rooms.clear(); + this.revokedRoomKeys.clear(); + await Promise.allSettled(rooms.map((room) => room.relay.stop())); + } +} + +class AetherVscodexCloudServer { + constructor(options = {}) { + this.host = options.host || process.env.HOST || "127.0.0.1"; + this.port = parsePort(options.port ?? process.env.PORT, 8788); + this.internalToken = options.internalToken || process.env.AETHER_VSCODEX_INTERNAL_TOKEN || ""; + this.publicWsUrl = options.publicWsUrl || process.env.AETHER_VSCODEX_PUBLIC_WS_URL || ""; + this.allowedOrigins = normalizeOrigins(options.allowedOrigins ?? process.env.AETHER_VSCODEX_ALLOWED_ORIGINS); + const dataDir = options.dataDir || process.env.AETHER_VSCODEX_DATA_DIR || path.join(process.cwd(), "data"); + this.store = options.store || new DeviceStore(path.join(dataDir, "devices.json")); + this.credentials = options.credentials || new EphemeralCredentials(options); + this.rooms = options.rooms || new RoomManager(options); + this.exchangeAttempts = new Map(); + this.httpServer = null; + this.wsServer = null; + this.cleanupTimer = null; + } + + async start() { + if (!this.internalToken) throw new Error("AETHER_VSCODEX_INTERNAL_TOKEN is required"); + if (Buffer.byteLength(this.internalToken, "utf8") < 24) throw new Error("AETHER_VSCODEX_INTERNAL_TOKEN must contain at least 24 bytes"); + if (!this.publicWsUrl) throw new Error("AETHER_VSCODEX_PUBLIC_WS_URL is required"); + validatePublicWsUrl(this.publicWsUrl); + if (!isLoopbackHost(this.host) && this.allowedOrigins.size === 0) { + throw new Error("AETHER_VSCODEX_ALLOWED_ORIGINS is required when binding outside loopback"); + } + this.httpServer = http.createServer((request, response) => { + void this.handleHttp(request, response).catch((error) => { + jsonResponse(response, error.statusCode || 500, { error: error.expose ? error.message : "internal server error" }); + }); + }); + this.wsServer = new WebSocketServer({ noServer: true, maxPayload: MAX_WS_BYTES }); + this.httpServer.on("upgrade", (request, socket, head) => this.handleUpgrade(request, socket, head)); + this.cleanupTimer = setInterval(() => { + this.credentials.cleanup(); + this.cleanupExchangeAttempts(); + void this.rooms.cleanup(); + }, 30_000); + this.cleanupTimer.unref(); + await new Promise((resolve, reject) => { + const onError = (error) => reject(error); + this.httpServer.once("error", onError); + this.httpServer.listen(this.port, this.host, () => { + this.httpServer.off("error", onError); + resolve(); + }); + }); + return this.address(); + } + + address() { + const address = this.httpServer.address(); + if (!address || typeof address === "string") return { host: this.host, port: this.port }; + return { host: address.address, port: address.port }; + } + + async stop() { + if (this.cleanupTimer) clearInterval(this.cleanupTimer); + this.cleanupTimer = null; + if (this.wsServer) { + for (const client of this.wsServer.clients) client.close(1001, "server shutting down"); + await new Promise((resolve) => this.wsServer.close(() => resolve())); + } + if (this.httpServer) await new Promise((resolve) => this.httpServer.close(() => resolve())); + this.wsServer = null; + this.httpServer = null; + await this.rooms.stop(); + } + + async handleHttp(request, response) { + const requestUrl = new URL(request.url || "/", "http://sidecar.local"); + if (request.method === "GET" && requestUrl.pathname === "/healthz") { + jsonResponse(response, 200, { ok: true, service: "aether-vscodex", mode: "single-replica" }); + return; + } + if (request.method === "POST" && requestUrl.pathname === "/v1/pairings/exchange") { + this.enforceExchangeRate(request); + const body = await readJson(request); + const pairing = this.credentials.consumePairing(body.code); + if (!pairing) throw httpError(400, "invalid or expired pairing code"); + const created = this.store.create(pairing.user_id, body.name || pairing.requested_name); + jsonResponse(response, 201, { + device_id: created.device.id, + device_name: created.device.name, + device_token: created.token, + ws_url: this.publicWsUrl, + }); + return; + } + + const match = requestUrl.pathname.match(/^\/internal\/v1\/users\/([^/]+)\/(devices|pairings|ws-tickets)(?:\/([^/]+))?$/); + if (!match) { + jsonResponse(response, 404, { error: "not found" }); + return; + } + this.requireInternalAuth(request); + const userId = decodeURIComponent(match[1]); + const resource = match[2]; + const resourceId = match[3] ? decodeURIComponent(match[3]) : null; + if (!userId || userId.length > 256) throw httpError(400, "invalid user id"); + + if (request.method === "GET" && resource === "devices" && !resourceId) { + jsonResponse(response, 200, { devices: this.store.list(userId, this.rooms.connectedDeviceIds(userId)) }); + return; + } + if (request.method === "POST" && resource === "pairings" && !resourceId) { + const body = await readJson(request); + const pairing = this.credentials.createPairing(userId, body.name); + jsonResponse(response, 201, { + pairing_id: pairing.id, + code: pairing.code, + expires_at: new Date(pairing.expires_at_ms).toISOString(), + }); + return; + } + if (request.method === "DELETE" && resource === "devices" && resourceId) { + if (!this.store.revoke(userId, resourceId)) throw httpError(404, "device not found"); + await this.rooms.revoke(userId, resourceId); + response.writeHead(204, { "Cache-Control": "no-store" }); + response.end(); + return; + } + if (request.method === "POST" && resource === "ws-tickets" && !resourceId) { + const body = await readJson(request); + const deviceId = typeof body.device_id === "string" ? body.device_id : ""; + if (!deviceId || !this.store.get(userId, deviceId)) throw httpError(404, "device not found"); + jsonResponse(response, 201, { + ticket: this.credentials.createTicket(userId, deviceId), + ws_url: "/api/vscodex/ws", + expires_in: Math.floor(this.credentials.ticketTtlMs / 1000), + }); + return; + } + jsonResponse(response, 405, { error: "method not allowed" }, { Allow: allowedMethod(resource, resourceId) }); + } + + requireInternalAuth(request) { + if (!this.hasInternalAuth(request)) throw httpError(401, "unauthorized"); + } + + hasInternalAuth(request) { + const authorization = String(request.headers.authorization || ""); + const token = authorization.startsWith("Bearer ") ? authorization.slice(7) : ""; + return secureEqual(token, this.internalToken); + } + + enforceExchangeRate(request) { + const address = this.exchangeRateAddress(request); + const now = Date.now(); + const attempts = (this.exchangeAttempts.get(address) || []).filter((time) => now - time < 60_000); + if (attempts.length >= 10) throw httpError(429, "too many pairing attempts"); + attempts.push(now); + this.exchangeAttempts.set(address, attempts); + } + + exchangeRateAddress(request) { + if (this.hasInternalAuth(request)) { + const forwardedAddress = singleHeaderValue(request, "x-aether-client-ip")?.trim(); + if (forwardedAddress && net.isIP(forwardedAddress)) return forwardedAddress; + } + return request.socket.remoteAddress || "unknown"; + } + + cleanupExchangeAttempts() { + const now = Date.now(); + for (const [address, attempts] of this.exchangeAttempts) { + const active = attempts.filter((time) => now - time < 60_000); + if (active.length) this.exchangeAttempts.set(address, active); + else this.exchangeAttempts.delete(address); + } + } + + handleUpgrade(request, socket, head) { + const requestUrl = new URL(request.url || "/", "http://sidecar.local"); + if (requestUrl.pathname !== "/api/vscodex/ws" && requestUrl.pathname !== "/v1/connect") { + rejectUpgrade(socket, 404, "Not Found"); + return; + } + const origin = request.headers.origin; + if (origin && this.allowedOrigins.size > 0 && !this.allowedOrigins.has(normalizeOrigin(origin))) { + rejectUpgrade(socket, 403, "Forbidden"); + return; + } + this.wsServer.handleUpgrade(request, socket, head, (webSocket) => { + this.wsServer.emit("connection", webSocket, request); + this.handleWebSocket(webSocket, request); + }); + } + + handleWebSocket(socket, request) { + let hello = null; + let token = ""; + let upstream = null; + let authenticating = false; + let room = null; + const queued = []; + const authTimer = setTimeout(() => socket.close(1008, "authentication required"), 10_000); + authTimer.unref(); + + const connectUpstream = async () => { + if (authenticating || upstream || !hello || !token) return; + authenticating = true; + let identity; + let upstreamToken; + if (hello.clientType === "host") { + const device = this.store.authenticate(token); + if (!device) throw httpError(401, "invalid device credential"); + identity = { userId: device.user_id, deviceId: device.id }; + room = await this.rooms.get(identity.userId, identity.deviceId); + upstreamToken = room.hostToken; + this.store.touch(device.id); + } else { + const ticket = this.credentials.consumeTicket(token); + if (!ticket || !this.store.get(ticket.user_id, ticket.device_id)) throw httpError(401, "invalid or expired browser ticket"); + identity = { userId: ticket.user_id, deviceId: ticket.device_id }; + room = await this.rooms.get(identity.userId, identity.deviceId); + upstreamToken = room.operatorToken; + } + this.rooms.retain(room); + upstream = new WebSocket(`${room.baseUrl}${hello.clientType === "host" ? "/v1/connect" : "/ws"}`, { + maxPayload: MAX_WS_BYTES, + }); + upstream.once("open", () => { + if (socket.readyState !== WebSocket.OPEN) { + upstream.close(); + return; + } + upstream.send(JSON.stringify(hello)); + upstream.send(JSON.stringify(hello.clientType === "host" + ? { v: 1, kind: "auth", accessToken: upstreamToken } + : { type: "auth", token: upstreamToken })); + for (const frame of queued.splice(0)) upstream.send(frame); + }); + upstream.on("message", (data, isBinary) => { + if (socket.readyState === WebSocket.OPEN) socket.send(data, { binary: isBinary }); + }); + upstream.on("close", (code, reason) => { + if (socket.readyState === WebSocket.OPEN) socket.close(validCloseCode(code) ? code : 1011, reason.toString().slice(0, 120) || "relay closed"); + }); + upstream.on("error", () => { + if (socket.readyState === WebSocket.OPEN) socket.close(1011, "relay unavailable"); + }); + clearTimeout(authTimer); + }; + + socket.on("message", (data, isBinary) => { + if (isBinary) { + socket.close(1003, "JSON text frames only"); + return; + } + if (upstream) { + const text = data.toString("utf8"); + if (upstream.readyState === WebSocket.OPEN) upstream.send(text); + else queued.push(text); + return; + } + let message; + try { + message = JSON.parse(data.toString("utf8")); + } catch { + socket.close(1007, "invalid JSON"); + return; + } + if (message?.kind === "hello") { + if (Number(message.protocol || 1) !== 1) { + socket.close(1002, "unsupported protocol"); + return; + } + hello = { + v: 1, + kind: "hello", + clientType: message.clientType === "host" ? "host" : "web", + protocol: 1, + ...(typeof message.sessionId === "string" ? { sessionId: message.sessionId } : {}), + ...(Number.isFinite(Number(message.lastSeq)) ? { lastSeq: Number(message.lastSeq) } : {}), + }; + } else if (message?.kind === "auth" || message?.type === "auth") { + token = typeof message.accessToken === "string" ? message.accessToken : typeof message.token === "string" ? message.token : ""; + } else { + socket.close(1002, "hello and auth required"); + return; + } + void connectUpstream().catch(() => socket.close(1008, "authentication failed")); + }); + socket.on("close", () => { + clearTimeout(authTimer); + if (upstream && upstream.readyState < WebSocket.CLOSING) upstream.close(); + if (room) this.rooms.release(room); + }); + socket.on("error", () => {}); + } +} + +function parseDeviceToken(token) { + const match = /^avx1\.([0-9a-f-]{36})\.([A-Za-z0-9_-]{32,})$/.exec(String(token || "")); + return match ? { id: match[1], secret: match[2] } : null; +} + +function deriveSecret(secret, salt) { + return crypto.scryptSync(secret, Buffer.from(salt, "base64url"), 32).toString("base64url"); +} + +function publicDevice(device, connected) { + return { + id: device.id, + name: device.name, + connected, + created_at: device.created_at, + last_seen_at: device.last_seen_at, + }; +} + +function normalizeName(value) { + const name = typeof value === "string" ? value.trim().replace(/\s+/g, " ").slice(0, 80) : ""; + return name || "VS Code"; +} + +function pairingCode() { + const alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"; + const bytes = crypto.randomBytes(8); + let result = ""; + for (let index = 0; index < 8; index += 1) result += alphabet[bytes[index] % alphabet.length]; + return `${result.slice(0, 4)}-${result.slice(4)}`; +} + +function normalizePairingCode(value) { + return String(value || "").toUpperCase().replace(/[^A-Z2-9]/g, ""); +} + +function randomToken() { + return crypto.randomBytes(32).toString("base64url"); +} + +function secureEqual(left, right) { + const a = Buffer.from(String(left || "")); + const b = Buffer.from(String(right || "")); + return a.length === b.length && crypto.timingSafeEqual(a, b); +} + +function singleHeaderValue(request, name) { + const distinctValues = request.headersDistinct?.[name]; + if (Array.isArray(distinctValues)) return distinctValues.length === 1 ? distinctValues[0] : null; + const value = request.headers[name]; + return typeof value === "string" ? value : null; +} + +function parsePort(value, fallback) { + const parsed = Number(value ?? fallback); + if (!Number.isInteger(parsed) || parsed < 0 || parsed > 65535) throw new Error("invalid port"); + return parsed; +} + +function normalizeOrigins(value) { + const values = Array.isArray(value) ? value : String(value || "").split(","); + return new Set(values.map(normalizeOrigin).filter(Boolean)); +} + +function normalizeOrigin(value) { + try { + return new URL(String(value).trim()).origin.toLowerCase(); + } catch { + return ""; + } +} + +function validatePublicWsUrl(value) { + let url; + try { + url = new URL(value); + } catch { + throw new Error("AETHER_VSCODEX_PUBLIC_WS_URL must be an absolute WebSocket URL"); + } + if (url.protocol !== "wss:" && !(url.protocol === "ws:" && isLoopbackHost(url.hostname))) { + throw new Error("AETHER_VSCODEX_PUBLIC_WS_URL must use wss:// outside loopback"); + } +} + +function isLoopbackHost(value) { + const host = String(value || "").replace(/^\[|\]$/g, "").toLowerCase(); + return host === "127.0.0.1" || host === "localhost" || host === "::1"; +} + +function validCloseCode(code) { + return code === 1000 || (code >= 1001 && code <= 1014 && ![1004, 1005, 1006].includes(code)) || (code >= 3000 && code <= 4999); +} + +function readJson(request) { + return new Promise((resolve, reject) => { + let size = 0; + const chunks = []; + request.on("data", (chunk) => { + size += chunk.length; + if (size > MAX_JSON_BYTES) { + reject(httpError(413, "request body too large")); + request.destroy(); + return; + } + chunks.push(chunk); + }); + request.on("end", () => { + try { + const value = JSON.parse(Buffer.concat(chunks).toString("utf8") || "{}"); + if (!value || typeof value !== "object" || Array.isArray(value)) throw new Error(); + resolve(value); + } catch { + reject(httpError(400, "invalid JSON body")); + } + }); + request.on("error", reject); + }); +} + +function jsonResponse(response, statusCode, body, extraHeaders = {}) { + if (response.headersSent) return; + const payload = Buffer.from(JSON.stringify(body)); + response.writeHead(statusCode, { + "Content-Type": "application/json; charset=utf-8", + "Content-Length": payload.length, + "Cache-Control": "no-store", + ...extraHeaders, + }); + response.end(payload); +} + +function rejectUpgrade(socket, status, reason) { + socket.write(`HTTP/1.1 ${status} ${reason}\r\nConnection: close\r\n\r\n`); + socket.destroy(); +} + +function httpError(statusCode, message) { + return Object.assign(new Error(message), { statusCode, expose: statusCode < 500 }); +} + +function allowedMethod(resource, resourceId) { + if (resource === "devices" && resourceId) return "DELETE"; + if (resource === "devices") return "GET"; + return "POST"; +} + +async function main() { + const server = new AetherVscodexCloudServer(); + const address = await server.start(); + process.stdout.write(`Aether VS Codex sidecar listening on ${address.host}:${address.port}\n`); + const shutdown = async () => { + await server.stop(); + process.exit(0); + }; + process.once("SIGINT", shutdown); + process.once("SIGTERM", shutdown); +} + +if (require.main === module) { + main().catch((error) => { + process.stderr.write(`${error.stack || error}\n`); + process.exitCode = 1; + }); +} + +module.exports = { + AetherVscodexCloudServer, + DeviceStore, + EphemeralCredentials, + RoomManager, +}; diff --git a/aether-vscodex/docker-compose.aether.yml b/aether-vscodex/docker-compose.aether.yml new file mode 100644 index 000000000..2ebffc44b --- /dev/null +++ b/aether-vscodex/docker-compose.aether.yml @@ -0,0 +1,37 @@ +services: + app: + environment: + AETHER_VSCODEX_ENABLED: "true" + AETHER_VSCODEX_INTERNAL_URL: http://vscodex:8788 + AETHER_VSCODEX_INTERNAL_TOKEN: ${AETHER_VSCODEX_INTERNAL_TOKEN:?set AETHER_VSCODEX_INTERNAL_TOKEN} + AETHER_VSCODEX_PUBLIC_WS_URL: ${AETHER_VSCODEX_PUBLIC_WS_URL:?set AETHER_VSCODEX_PUBLIC_WS_URL} + depends_on: + vscodex: + condition: service_healthy + volumes: + - ./aether-vscodex/web/dist:/opt/aether/releases/image/frontend/aether-vscodex:ro + + vscodex: + build: + context: ./aether-vscodex + image: ${AETHER_VSCODEX_IMAGE:-aether-vscodex:local} + environment: + HOST: 0.0.0.0 + PORT: 8788 + AETHER_VSCODEX_INTERNAL_TOKEN: ${AETHER_VSCODEX_INTERNAL_TOKEN:?set AETHER_VSCODEX_INTERNAL_TOKEN} + AETHER_VSCODEX_PUBLIC_WS_URL: ${AETHER_VSCODEX_PUBLIC_WS_URL:?set AETHER_VSCODEX_PUBLIC_WS_URL} + AETHER_VSCODEX_ALLOWED_ORIGINS: ${AETHER_VSCODEX_ALLOWED_ORIGINS:?set AETHER_VSCODEX_ALLOWED_ORIGINS} + AETHER_VSCODEX_DATA_DIR: /var/lib/aether-vscodex + expose: + - "8788" + volumes: + - vscodex_data:/var/lib/aether-vscodex + logging: + driver: local + options: + max-size: "50m" + max-file: "3" + restart: unless-stopped + +volumes: + vscodex_data: diff --git a/aether-vscodex/docs/cloud-security.md b/aether-vscodex/docs/cloud-security.md new file mode 100644 index 000000000..3608afe89 --- /dev/null +++ b/aether-vscodex/docs/cloud-security.md @@ -0,0 +1,20 @@ +# Cloud security model + +## Trust boundaries + +- Aether authenticates browser HTTP requests and resolves the user ID. The client never supplies a trusted user ID. +- The Node sidecar never receives an Aether access token or JWT signing key. +- A VS Code installation receives one revocable device credential. Only its scrypt hash is persisted. +- An iframe receives a random, one-time WebSocket ticket with a 60-second lifetime. Tickets are sent in an auth frame, never in a URL. +- The embedded UI is trusted, same-origin Aether code. `allow-same-origin` is required by the current integration, so the iframe is not a sandbox boundary for untrusted content even though the parent does not post its JWT into the frame. +- Relay state is isolated by `(user_id, device_id)`. A browser ticket and host credential must resolve to the same room. + +## Network boundary + +Run the sidecar on the private Compose network. Do not publish port 8788. Aether gateway is the only public HTTP and WebSocket entry point and authenticates internal API calls with `AETHER_VSCODEX_INTERNAL_TOKEN`. + +`AETHER_VSCODEX_ALLOWED_ORIGINS` must contain the exact public Aether origin when the sidecar binds outside loopback. Public deployments must use HTTPS/WSS. + +## Current scaling limit + +The first release intentionally runs one sidecar replica. Pairing codes, browser tickets, and the live connection directory are process-local. Before adding replicas, move those records to a shared atomic store and add sticky or distributed WebSocket room routing. diff --git a/aether-vscodex/fixtures/fake-app-server.cjs b/aether-vscodex/fixtures/fake-app-server.cjs new file mode 100644 index 000000000..2001db255 --- /dev/null +++ b/aether-vscodex/fixtures/fake-app-server.cjs @@ -0,0 +1,59 @@ +"use strict"; + +const readline = require("node:readline"); + +let threadNumber = 0; +let turnNumber = 0; +let activeThread = null; +let activeTurn = null; + +function send(message) { + process.stdout.write(`${JSON.stringify(message)}\n`); +} + +const input = readline.createInterface({ input: process.stdin }); +input.on("line", (line) => { + let request; + try { request = JSON.parse(line); } catch { return; } + if (request.method === "initialize") { + send({ id: request.id, result: { userAgent: "fake", codexHome: "/tmp/codex" } }); + send({ method: "remoteControl/status/changed", params: { status: "disabled" } }); + return; + } + if (request.method === "thread/start") { + activeThread = `thread-${++threadNumber}`; + send({ id: request.id, result: { thread: { id: activeThread }, cwd: request.params?.cwd || "/tmp" } }); + send({ method: "thread/started", params: { thread: { id: activeThread } } }); + return; + } + if (request.method === "turn/start") { + activeTurn = `turn-${++turnNumber}`; + send({ id: request.id, result: { turn: { id: activeTurn } } }); + send({ method: "turn/started", params: { threadId: request.params.threadId, turn: { id: activeTurn } } }); + const text = request.params.input?.[0]?.text || ""; + send({ method: "item/agentMessage/delta", params: { threadId: request.params.threadId, turnId: activeTurn, itemId: "item-1", delta: `echo: ${text}` } }); + if (text.includes("approve")) { + send({ id: 9001, method: "item/commandExecution/requestApproval", params: { threadId: request.params.threadId, turnId: activeTurn, itemId: "item-2", command: "echo approval" } }); + } else { + send({ method: "turn/completed", params: { threadId: request.params.threadId, turn: { id: activeTurn } } }); + activeTurn = null; + } + return; + } + if (request.method === "turn/steer") { + send({ id: request.id, result: { turn: { id: activeTurn } } }); + send({ method: "item/agentMessage/delta", params: { delta: `steered: ${request.params.input?.[0]?.text || ""}` } }); + return; + } + if (request.method === "turn/interrupt") { + send({ id: request.id, result: {} }); + send({ method: "turn/completed", params: { threadId: request.params.threadId, turn: { id: request.params.turnId } } }); + activeTurn = null; + return; + } + if (request.id === 9001 && (request.result || request.error)) { + send({ method: "item/agentMessage/delta", params: { delta: `approval response: ${JSON.stringify(request.result || request.error)}` } }); + send({ method: "turn/completed", params: { threadId: activeThread, turn: { id: activeTurn } } }); + activeTurn = null; + } +}); diff --git a/aether-vscodex/package-lock.json b/aether-vscodex/package-lock.json new file mode 100644 index 000000000..97286de97 --- /dev/null +++ b/aether-vscodex/package-lock.json @@ -0,0 +1,39 @@ +{ + "name": "aether-vscodex", + "version": "0.4.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "aether-vscodex", + "version": "0.4.0", + "dependencies": { + "ws": "^8.18.3" + }, + "engines": { + "node": ">=20" + } + }, + "node_modules/ws": { + "version": "8.21.3", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.3.tgz", + "integrity": "sha512-201TZ/kPWxoPr/OKWjquZR1SWKXcvxdH+e1xrx89b3YbmzLMFCLfnaG1HFIgWzJOEWZ7MvpK++odZufgYR50Rw==", + "license": "MIT", + "engines": { + "node": ">=10.0.0" + }, + "peerDependencies": { + "bufferutil": "^4.0.1", + "utf-8-validate": ">=5.0.2" + }, + "peerDependenciesMeta": { + "bufferutil": { + "optional": true + }, + "utf-8-validate": { + "optional": true + } + } + } + } +} diff --git a/aether-vscodex/package.json b/aether-vscodex/package.json new file mode 100644 index 000000000..831fe6aa0 --- /dev/null +++ b/aether-vscodex/package.json @@ -0,0 +1,23 @@ +{ + "name": "aether-vscodex", + "version": "0.4.0", + "private": true, + "description": "Synchronous VS Code Codex mirroring and asynchronous Codex control for local Web and Aether", + "type": "commonjs", + "main": "relay/server.js", + "scripts": { + "start": "node relay/server.js", + "start:cloud": "node cloud/server.js", + "build:web": "npm --prefix web run build", + "build:extension": "npm --prefix vscode-extension run build", + "build": "npm run build:web && npm run build:extension", + "test": "node --test test/*.test.js", + "test:web": "npm --prefix web test" + }, + "engines": { + "node": ">=20" + }, + "dependencies": { + "ws": "^8.18.3" + } +} diff --git a/aether-vscodex/public/app.js b/aether-vscodex/public/app.js new file mode 100644 index 000000000..e6947d26b --- /dev/null +++ b/aether-vscodex/public/app.js @@ -0,0 +1,8199 @@ +(function () { + "use strict"; + + const $ = (id) => document.getElementById(id); + const embedBridge = window.AetherVscodexEmbed; + const embeddedInAether = Boolean(embedBridge?.active); + const i18n = window.VscodexI18n; + const t = (value) => i18n?.t ? i18n.t(value) : String(value ?? ""); + const uiLocale = () => i18n?.locale?.() || "zh-CN"; + // Translate renderer-owned labels while leaving values supplied by the host + // (conversation titles, file paths, commands, and message text) untouched. + const uiText = (zh, en) => uiLocale() === "en-US" ? en : zh; + const uiWithRaw = (zhPrefix, enPrefix, raw, zhSuffix = "", enSuffix = "") => + `${uiText(zhPrefix, enPrefix)}${String(raw ?? "")}${uiText(zhSuffix, enSuffix)}`; + const appStatusLabel = (value) => { + const normalized = String(value || "").trim().toLowerCase(); + const labels = { + ready: "已连接", + online: "已连接", + offline: "VS Code 主机未连接", + waiting_for_host: "等待 VS Code 主机连接", + app_not_ready: "等待 VS Code 主机连接", + starting: "正在连接", + connecting: "正在连接", + stopped: "已停止", + }; + return t(labels[normalized] || value || "未知"); + }; + const state = { + ws: null, + token: "", + role: null, + appReady: false, + lastSeq: 0, + // A control snapshot is authoritative for every event up to this + // sequence. Replayed notifications from the subscribe handshake must not + // resurrect an already-finished turn or duplicate its transcript. + lastSnapshotSeq: 0, + awaitingSnapshot: false, + threadId: "", + turnId: "", + attachMode: false, + authRequired: null, + outputSynced: false, + structuredMessages: [], + requests: new Map(), + responding: new Set(), + commandResults: new Set(), + reconnectTimer: null, + embedTicket: "", + embedWsUrl: "", + embedDeviceId: "", + embedTicketRequested: false, + embedStopped: false, + activeAssistantBody: null, + activeAssistantStream: null, + activeAssistantText: "", + pendingUserText: "", + retiredTurnIds: new Set(), + syncedThreadId: null, + snapshotNoticeShown: false, + // Live work items are rendered as updateable transcript entries. The + // adapter may emit item lifecycle notifications or only output chunks; + // keeping a small client-side index lets both forms converge on one row. + activities: new Map(), + commandDisclosure: new Map(), + activitySequence: 0, + activityTimer: null, + activeAssistantActivityKey: null, + turnStartedAt: null, + // The official worked-for row measures from the first work item until + // the final assistant response starts. Keep this separate from the + // overall turn clock because the latter also includes queue/approval time. + turnWorkStartedAt: null, + finalAssistantStartedAt: null, + turnStatus: "idle", + lastTurnDurationMs: null, + lastWorkedDurationMs: null, + workedDurationMs: null, + currentActivity: "idle", + currentActivityStartedAt: null, + currentActivityDurationMs: null, + currentActivityTurnId: "", + currentModel: "", + // Empty means the host has not reported an effort yet; null is an + // authoritative "use the model default" value and must not be serialized + // as medium on the next turn. + currentEffort: "", + sandboxPolicy: "workspace-write", + approvalPolicy: "on-request", + tokenUsage: null, + availableModels: [], + subagents: [], + controlMode: "sync", + modeEpoch: -1, + capabilities: { + followsVscodeRoute: true, + sessionList: false, + sessionSelect: false, + sessionCreate: false, + threadSettings: false, + }, + modeSnapshotReady: false, + modeSwitching: false, + modeCommandId: "", + requestedControlMode: "", + modeRequestEpoch: -1, + sessions: [], + sessionPickerOpen: false, + sessionSearch: "", + sessionFocusedId: "", + sessionListLoading: false, + sessionListError: "", + sessionListCommandId: "", + sessionSelectCommandId: "", + newSessionCommandId: "", + sessionSelectedThreadId: "", + sessionSwitching: false, + // Keep the previous view/title visible until the host confirms the target + // with an authoritative session snapshot. A target can be listed as + // attachable and still time out during owner hand-off; dropping the old + // DOM at `session.switching` would leave the browser blank in that case. + sessionSwitchContext: null, + modelUpdatePending: false, + modelAdvancedOpen: false, + // The official composer keeps the background-agent disclosure closed on + // first render; the @ hint appears only after the reader expands it. + subagentsCollapsed: true, + subagentsExpanded: { active: false, done: false }, + lastRenderedDateKey: "", + lastRenderedTimestamp: null, + lastRenderedRole: "", + hasRenderedUser: false, + lastDateSeparatorTimestamp: null, + turnDividers: new Map(), + // Preserve an explicit worked-for toggle across authoritative snapshot + // rebuilds. Unset entries follow the official default: the latest turn is + // open while older completed turns remain compact. + turnExpansion: new Map(), + // Legacy attach snapshots may omit turnId on individual projected items. + // Keep the derived association by object identity for the duration of a + // snapshot so all grouping/timing paths use the same anonymous turn key. + structuredTurnKeys: new WeakMap(), + liveActivityKey: null, + pendingUserArticle: null, + outputDistanceFromBottom: 0, + timelineAnchorLockUntil: 0, + timelineAnchorCancel: null, + timelineRevealCancel: null, + }; + const RESPONDABLE_METHODS = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", + "item/permissions/requestApproval", + "item/tool/requestUserInput", + "mcpServer/elicitation/request", + "applyPatchApproval", + "execCommandApproval", + ]); + + const requestKey = (requestId) => `${typeof requestId}:${String(requestId)}`; + + function normalizeControlMode(value) { + const normalized = String(value || "").trim().toLowerCase(); + return normalized === "sync" || normalized === "async" ? normalized : ""; + } + + function normalizedControlCapabilities(value) { + const source = isRecord(value) ? value : {}; + return { + followsVscodeRoute: source.followsVscodeRoute === true, + sessionList: source.sessionList === true, + sessionSelect: source.sessionSelect === true, + sessionCreate: source.sessionCreate === true, + threadSettings: source.threadSettings === true, + }; + } + + function sessionControlAllowed(capability) { + return state.modeSnapshotReady + && state.controlMode === "async" + && state.capabilities[capability] === true; + } + + function threadSettingsAllowed() { + return state.modeSnapshotReady && state.capabilities.threadSettings === true; + } + + function controlModeChangeBlocked() { + return !state.modeSnapshotReady + || state.modeSwitching + || state.sessionSwitching + || !state.appReady + || !state.ws + || state.ws.readyState !== WebSocket.OPEN + || !["operator", "owner", "host"].includes(String(state.role || "")) + || Boolean(state.turnId) + || state.turnStartedAt !== null + || state.turnStatus === "active" + || state.turnStatus === "waiting" + || Boolean(state.pendingUserText) + || state.requests.size > 0 + || state.responding.size > 0 + || Boolean(state.newSessionCommandId) + || Boolean(state.sessionListCommandId) + || Boolean(state.sessionSelectCommandId) + || state.modelUpdatePending; + } + + function clearControlModeRequest() { + state.modeSwitching = false; + state.modeCommandId = ""; + state.requestedControlMode = ""; + state.modeRequestEpoch = -1; + } + + function renderControlMode() { + const control = $("controlModeSwitch"); + if (!control) return; + const listAllowed = sessionControlAllowed("sessionList"); + const createAllowed = sessionControlAllowed("sessionCreate"); + const settingsAllowed = threadSettingsAllowed(); + const blocked = controlModeChangeBlocked(); + control.dataset.mode = state.controlMode; + control.dataset.epoch = String(state.modeEpoch); + control.dataset.switching = String(state.modeSwitching); + control.setAttribute("aria-label", t("控制模式")); + control.setAttribute("aria-busy", String(state.modeSwitching)); + for (const button of control.querySelectorAll("[data-control-mode]")) { + const mode = normalizeControlMode(button.dataset.controlMode); + const current = mode === state.controlMode; + const pending = state.modeSwitching && mode === state.requestedControlMode; + button.textContent = t(mode === "async" ? "异步" : "同步"); + button.title = t(mode === "async" ? "异步模式可独立管理会话" : "同步模式跟随 VS Code 当前会话"); + button.setAttribute("aria-pressed", String(current)); + button.dataset.pending = String(pending); + button.disabled = blocked || current; + } + + const sessionPickerButton = $("sessionPickerButton"); + if (sessionPickerButton) { + sessionPickerButton.disabled = !listAllowed || state.modeSwitching || state.sessionSwitching; + sessionPickerButton.setAttribute("aria-disabled", String(sessionPickerButton.disabled)); + } + for (const id of ["backButton", "historyButton"]) { + const button = $(id); + if (button) button.hidden = !listAllowed; + } + const newSessionButton = $("newSessionButton"); + if (newSessionButton) newSessionButton.hidden = !createAllowed; + const sessionsMenuItem = document.querySelector('[data-menu-action="sessions"]'); + if (sessionsMenuItem) sessionsMenuItem.hidden = !listAllowed; + const refresh = $("sessionPickerRefresh"); + if (refresh) refresh.disabled = !listAllowed || state.modeSwitching; + if (!listAllowed && state.sessionPickerOpen) setSessionPicker(false); + + for (const action of document.querySelectorAll('[data-settings-action="model"], [data-settings-action="permission"]')) { + action.disabled = !settingsAllowed || state.modeSwitching || state.sessionSwitching; + } + for (const id of ["modelPickerButton", "permissionChip"]) { + const button = $(id); + if (button) button.disabled = !settingsAllowed || state.modeSwitching || state.sessionSwitching; + } + if (!settingsAllowed) { + setModelMenu(false); + setPermissionMenu(false); + } + } + + function applyControlModeSnapshot(metadata) { + if (!isRecord(metadata)) return false; + const mode = normalizeControlMode(metadata.controlMode); + const epoch = finiteNumber(metadata.modeEpoch); + if (!mode || epoch === null || (state.modeSnapshotReady && epoch < state.modeEpoch)) return false; + const wasSwitching = state.modeSwitching; + const requestedMode = state.requestedControlMode; + const requestEpoch = state.modeRequestEpoch; + state.controlMode = mode; + state.modeEpoch = epoch; + state.capabilities = normalizedControlCapabilities(metadata.capabilities); + state.modeSnapshotReady = true; + const requestResolved = wasSwitching && (epoch > requestEpoch || mode === requestedMode); + if (requestResolved) { + clearControlModeRequest(); + setConversationStatus(mode === requestedMode ? "控制模式已切换" : "控制模式切换失败", mode === requestedMode ? "ready" : "warning"); + } + updateIds(); + return true; + } + + function resolveRequestKey(requestId) { + const exact = requestKey(requestId); + if (state.requests.has(exact)) return exact; + const text = String(requestId); + const candidates = [...state.requests] + .filter(([, request]) => String(request.requestId) === text) + .map(([key]) => key); + return candidates.length === 1 ? candidates[0] : exact; + } + + function shouldFollowOutput(output) { + return output.scrollHeight - output.scrollTop - output.clientHeight <= 24; + } + + const disclosureAnimations = new WeakMap(); + const activityAnimations = new WeakMap(); + const commandAnimations = new WeakMap(); + + function motionDuration(value = 220) { + try { + if (window.matchMedia?.("(prefers-reduced-motion: reduce)")?.matches) return 0; + } catch { + // Older embedded webviews may not expose matchMedia. + } + return Math.max(0, Number(value) || 0); + } + + function motionClock() { + return typeof performance === "object" && typeof performance.now === "function" + ? performance.now() + : Date.now(); + } + + function scheduleFrame(callback) { + if (typeof requestAnimationFrame === "function") return requestAnimationFrame(callback); + return setTimeout(callback, 0); + } + + function cancelFrame(frame) { + if (frame === null || frame === undefined) return; + if (typeof cancelAnimationFrame === "function") cancelAnimationFrame(frame); + clearTimeout(frame); + } + + function runMeasuredCssTransition(element, from, to, duration, onFinish) { + if (!element) return null; + const priorTransition = element.style.transition; + let frame = null; + let timer = null; + let cancelled = false; + const setStyles = (values) => { + for (const [property, value] of Object.entries(values)) element.style[property] = String(value); + }; + const finish = () => { + if (cancelled) return; + cancelled = true; + cancelFrame(frame); + frame = null; + if (timer !== null) clearTimeout(timer); + timer = null; + element.style.transition = priorTransition; + onFinish?.(); + }; + const controller = { + cancel() { + if (cancelled) return; + cancelled = true; + cancelFrame(frame); + frame = null; + if (timer !== null) clearTimeout(timer); + timer = null; + element.style.transition = priorTransition; + }, + }; + const properties = Object.keys(to).map((property) => property.replace(/[A-Z]/g, (letter) => `-${letter.toLowerCase()}`)); + element.style.transition = properties + .map((property) => `${property} ${duration}ms cubic-bezier(.33,1,.68,1)`) + .join(", "); + setStyles(from); + // Force a layout boundary before moving to the measured target. This is + // the fallback for embedded Chromium builds without Element.animate(). + void element.offsetHeight; + frame = scheduleFrame(() => { + frame = null; + if (cancelled) return; + setStyles(to); + }); + timer = setTimeout(finish, duration + 55); + return controller; + } + + // Keep the element nearest the reader at the same viewport coordinate while + // a disclosure or streamed row changes height. This mirrors the official + // preserve-timeline-anchor-position helper and avoids bottom-relative jumps + // when the reader is inspecting an earlier turn. + function preserveTimelineAnchor(anchor, duration = 250) { + const target = anchor?.nodeType === 1 ? anchor : null; + const output = target?.closest?.(".chat-scroll") || $("output"); + if (!target || !output || !target.isConnected || !output.isConnected) return; + if (typeof state.timelineAnchorCancel === "function") state.timelineAnchorCancel(); + const initialTop = target.getBoundingClientRect().top; + if (!Number.isFinite(initialTop)) return; + const deadline = motionClock() + duration + 120; + state.timelineAnchorLockUntil = Math.max(state.timelineAnchorLockUntil || 0, deadline); + let frame = null; + let timer = null; + let disposed = false; + const adjust = () => { + if (disposed || !target.isConnected || !output.isConnected) return; + const delta = target.getBoundingClientRect().top - initialTop; + if (Number.isFinite(delta) && Math.abs(delta) >= 0.1) { + const maxScroll = Math.max(0, output.scrollHeight - output.clientHeight); + output.scrollTop = Math.max(0, Math.min(maxScroll, output.scrollTop + delta)); + updateScrollToBottom(output); + } + // A disclosure can change height without producing a ResizeObserver + // callback on older Chromium builds. Keep one frame queued for the full + // measured transition so the scrollbar and reader anchor move together. + if (motionClock() < deadline) schedule(); + }; + const schedule = () => { + if (disposed || frame !== null) return; + frame = scheduleFrame(() => { + frame = null; + adjust(); + }); + }; + const immediate = () => { + if (frame !== null) { + cancelFrame(frame); + frame = null; + } + adjust(); + schedule(); + }; + const observed = target.closest("[data-turn-key]") || target.closest(".message") || target; + let observer = null; + if (typeof ResizeObserver === "function") { + observer = new ResizeObserver(immediate); + observer.observe(observed); + if (observed !== target) observer.observe(target); + } + schedule(); + const finish = () => { + if (disposed) return; + disposed = true; + cancelFrame(frame); + frame = null; + if (timer !== null) clearTimeout(timer); + timer = null; + observer?.disconnect(); + if (state.timelineAnchorCancel === finish) state.timelineAnchorCancel = null; + }; + state.timelineAnchorCancel = finish; + timer = setTimeout(finish, Math.max(250, duration + 150)); + } + + function setDisclosureBodyState(details, expanded) { + const body = details?.querySelector?.(":scope > .details-body"); + if (!body) return; + body.style.display = "block"; + body.style.height = expanded ? "auto" : "0px"; + body.style.opacity = expanded ? "1" : "0"; + body.style.overflow = expanded ? "" : "hidden"; + body.style.pointerEvents = expanded ? "auto" : "none"; + body.dataset.disclosureState = expanded ? "expanded" : "collapsed"; + details.dataset.expanded = String(Boolean(expanded)); + details.querySelector("summary")?.setAttribute("aria-expanded", String(Boolean(expanded))); + } + + function animateDisclosureBody(details, expanded, options = {}) { + const body = details?.querySelector?.(":scope > .details-body"); + if (!body) return; + const previous = disclosureAnimations.get(body); + if (previous) { + try { previous.commitStyles?.(); } catch { /* animation may already be finished */ } + previous.cancel(); + } + disclosureAnimations.delete(body); + const next = Boolean(expanded); + details.dataset.expanded = String(next); + details.dataset.animating = "true"; + body.style.display = "block"; + const duration = motionDuration(options.duration ?? 220); + const computed = getComputedStyle(body); + const currentHeight = Math.max(0, body.getBoundingClientRect().height || 0); + const currentOpacity = Number.parseFloat(computed.opacity); + const fromOpacity = Number.isFinite(currentOpacity) ? currentOpacity : (next ? 0 : 1); + if (options.immediate || duration === 0) { + setDisclosureBodyState(details, next); + details.dataset.animating = "false"; + return; + } + let targetHeight = 0; + if (next) { + body.style.height = "auto"; + body.style.opacity = "1"; + body.style.overflow = "hidden"; + targetHeight = Math.max(0, body.getBoundingClientRect().height || body.scrollHeight || 0); + body.style.height = `${currentHeight}px`; + } else { + body.style.height = `${currentHeight}px`; + body.style.opacity = String(fromOpacity); + body.style.overflow = "hidden"; + } + if (Math.abs(targetHeight - currentHeight) < 0.5 && (next ? fromOpacity >= 0.99 : fromOpacity <= 0.01)) { + setDisclosureBodyState(details, next); + details.dataset.animating = "false"; + return; + } + if (typeof body.animate !== "function") { + let transition; + transition = runMeasuredCssTransition( + body, + { height: `${currentHeight}px`, opacity: fromOpacity }, + { height: `${targetHeight}px`, opacity: next ? 1 : 0 }, + duration, + () => { + if (disclosureAnimations.get(body) !== transition) return; + disclosureAnimations.delete(body); + setDisclosureBodyState(details, next); + details.dataset.animating = "false"; + }, + ); + if (transition) disclosureAnimations.set(body, transition); + return; + } + const animation = body.animate([ + { height: `${currentHeight}px`, opacity: fromOpacity }, + { height: `${targetHeight}px`, opacity: next ? 1 : 0 }, + ], { duration, easing: "cubic-bezier(.33,1,.68,1)", fill: "forwards" }); + disclosureAnimations.set(body, animation); + animation.addEventListener("finish", () => { + if (disclosureAnimations.get(body) !== animation) return; + disclosureAnimations.delete(body); + setDisclosureBodyState(details, next); + details.dataset.animating = "false"; + // `fill: forwards` keeps the animation in the cascade after finish. + // Release it only after the stable inline state is written, otherwise a + // later open can still measure the previous collapsed height (zero). + animation.cancel(); + }, { once: true }); + animation.addEventListener("cancel", () => { + if (disclosureAnimations.get(body) === animation) disclosureAnimations.delete(body); + }, { once: true }); + } + + function installDisclosure(details, initiallyOpen) { + if (!details || details.dataset.disclosureInstalled === "true") return; + const summary = details.querySelector(":scope > summary"); + const body = details.querySelector(":scope > .details-body"); + if (!summary || !body) return; + details.dataset.disclosureInstalled = "true"; + // Keep the native details element mounted. The summary still supplies the + // familiar keyboard/focus semantics, while the body itself is animated + // with measured height so a close does not jump the transcript. + details.open = true; + setDisclosureBodyState(details, Boolean(initiallyOpen)); + summary.addEventListener("click", (event) => { + event.preventDefault(); + setDetailsExpanded(details, !isDetailsExpanded(details)); + }, true); + details.addEventListener("toggle", () => { + // A browser/plugin script may still assign `.open = false`; restore the + // mounted shell and leave the visual state under our measured body. + if (!details.open && !details.__codexRestoring) { + details.__codexRestoring = true; + details.open = true; + details.__codexRestoring = false; + } + }); + } + + function isDetailsExpanded(details) { + if (!details) return false; + if (details.dataset.expanded !== undefined) return details.dataset.expanded === "true"; + return details.open === true; + } + + function setDetailsExpanded(details, expanded, options = {}) { + if (!details) return; + const next = Boolean(expanded); + const previous = isDetailsExpanded(details); + const installed = details.dataset.disclosureInstalled === "true"; + if (!installed) { + details.open = next; + return; + } + if (previous === next && !options.force) { + if (options.immediate) animateDisclosureBody(details, next, { immediate: true }); + return; + } + details.dataset.expanded = String(next); + if (options.preserve !== false && !options.immediate) preserveTimelineAnchor(details, options.duration ?? 230); + animateDisclosureBody(details, next, { immediate: Boolean(options.immediate), duration: options.duration }); + if (next && !options.immediate && options.reveal !== false) { + scheduleTimelineReveal(details.querySelector(":scope > .details-body") || details, options.duration ?? 220); + } + } + + function animateActivityArticle(article, expanded, options = {}) { + if (!article) return; + const previous = activityAnimations.get(article); + if (previous) { + try { previous.commitStyles?.(); } catch { /* animation may already be finished */ } + previous.cancel(); + } + activityAnimations.delete(article); + const next = Boolean(expanded); + const duration = motionDuration(options.duration ?? 230); + const wasCollapsed = article.classList.contains("turn-collapsed"); + article.dataset.turnExpanded = String(next); + + // `worked-for` is an outer disclosure. The official renderer removes the + // whole activity group from layout while it is collapsed, rather than + // leaving one summary row per command. Animate the article shell itself; + // nested command/read disclosures keep their own expanded state. + if (next) article.classList.remove("turn-collapsed"); + article.style.display = ""; + article.style.height = ""; + article.style.opacity = ""; + article.style.visibility = ""; + article.style.pointerEvents = ""; + article.style.overflow = ""; + const currentHeight = Math.max(0, article.getBoundingClientRect().height || 0); + const naturalHeight = Math.max(0, article.scrollHeight || currentHeight); + const fromHeight = next && wasCollapsed ? 0 : currentHeight; + const targetHeight = next ? naturalHeight : 0; + const fromOpacity = next ? (wasCollapsed ? 0 : 1) : 1; + const targetOpacity = next ? 1 : 0; + const finish = () => { + if (next) { + article.classList.remove("turn-collapsed"); + article.style.height = ""; + article.style.opacity = ""; + article.style.visibility = ""; + article.style.pointerEvents = ""; + article.style.overflow = ""; + } else { + article.classList.add("turn-collapsed"); + article.style.height = "0px"; + article.style.opacity = "0"; + article.style.visibility = "hidden"; + article.style.pointerEvents = "none"; + article.style.overflow = "hidden"; + } + }; + if (options.immediate || duration === 0 || Math.abs(targetHeight - fromHeight) < 0.5) { + finish(); + return; + } + article.style.height = `${fromHeight}px`; + article.style.opacity = String(fromOpacity); + article.style.overflow = "hidden"; + if (typeof article.animate !== "function") { + let transition; + transition = runMeasuredCssTransition( + article, + { height: `${fromHeight}px`, opacity: fromOpacity }, + { height: `${targetHeight}px`, opacity: targetOpacity }, + duration, + () => { + if (activityAnimations.get(article) !== transition) return; + activityAnimations.delete(article); + finish(); + }, + ); + if (transition) activityAnimations.set(article, transition); + return; + } + const animation = article.animate([ + { height: `${fromHeight}px`, opacity: fromOpacity }, + { height: `${targetHeight}px`, opacity: targetOpacity }, + ], { duration, easing: "cubic-bezier(.33,1,.68,1)", fill: "forwards" }); + activityAnimations.set(article, animation); + animation.addEventListener("finish", () => { + if (activityAnimations.get(article) !== animation) return; + activityAnimations.delete(article); + finish(); + animation.cancel(); + }, { once: true }); + animation.addEventListener("cancel", () => { + if (activityAnimations.get(article) === animation) activityAnimations.delete(article); + }, { once: true }); + } + + function animateCommandRow(commandRow, expanded, options = {}) { + if (!commandRow) return; + const previous = commandAnimations.get(commandRow); + if (previous) { + try { previous.commitStyles?.(); } catch { /* animation may already be finished */ } + previous.cancel(); + } + commandAnimations.delete(commandRow); + const next = Boolean(expanded); + const duration = motionDuration(options.duration ?? 190); + const fromHeight = Math.max(0, commandRow.getBoundingClientRect().height || 0); + commandRow.dataset.expanded = String(next); + commandRow.setAttribute("aria-expanded", String(next)); + commandRow.style.overflow = "hidden"; + commandRow.style.height = "auto"; + const targetHeight = Math.max(0, commandRow.getBoundingClientRect().height || 0); + if (options.immediate || duration === 0 || Math.abs(targetHeight - fromHeight) < 0.5) { + commandRow.style.height = ""; + commandRow.style.overflow = ""; + return; + } + if (typeof commandRow.animate !== "function") { + let transition; + transition = runMeasuredCssTransition( + commandRow, + { height: `${fromHeight}px` }, + { height: `${targetHeight}px` }, + duration, + () => { + if (commandAnimations.get(commandRow) !== transition) return; + commandAnimations.delete(commandRow); + commandRow.style.height = ""; + commandRow.style.overflow = ""; + }, + ); + if (transition) commandAnimations.set(commandRow, transition); + return; + } + commandRow.style.height = `${fromHeight}px`; + const animation = commandRow.animate([ + { height: `${fromHeight}px` }, + { height: `${targetHeight}px` }, + ], { duration, easing: "cubic-bezier(.33,1,.68,1)", fill: "forwards" }); + commandAnimations.set(commandRow, animation); + animation.addEventListener("finish", () => { + if (commandAnimations.get(commandRow) !== animation) return; + commandAnimations.delete(commandRow); + commandRow.style.height = ""; + commandRow.style.overflow = ""; + animation.cancel(); + }, { once: true }); + animation.addEventListener("cancel", () => { + if (commandAnimations.get(commandRow) === animation) commandAnimations.delete(commandRow); + }, { once: true }); + } + + function updateScrollToBottom(output = $("output")) { + const button = $("scrollToBottom"); + if (!button || !output) return; + const distance = Math.max(0, output.scrollHeight - output.scrollTop - output.clientHeight); + state.outputDistanceFromBottom = distance; + const visible = distance > 24; + const working = state.turnStartedAt !== null || state.turnStatus === "active" || state.turnStatus === "waiting"; + button.dataset.visible = String(visible); + button.dataset.working = String(working); + button.setAttribute("aria-label", t(working ? "正在工作,回到最新消息" : "回到最新消息")); + button.setAttribute("aria-hidden", String(!visible)); + button.tabIndex = visible ? 0 : -1; + } + + function scrollOutput(output, force = false) { + if (!output) return; + if (force || shouldFollowOutput(output)) { + if (typeof output.scrollTo === "function") output.scrollTo({ top: output.scrollHeight, behavior: "auto" }); + else output.scrollTop = output.scrollHeight; + } + updateScrollToBottom(output); + } + + function updateScrollPadding() { + const output = $("output"); + const panel = document.querySelector(".chat-panel"); + const composer = $("messageForm"); + if (!output || !panel || !composer) return; + const outputRect = output.getBoundingClientRect(); + const composerRect = composer.getBoundingClientRect(); + // The composer is an overlay in the official panel. Reserve only the + // portion that actually covers the scroll viewport, plus a small gap. + const overlap = Math.max(0, Math.ceil(outputRect.bottom - composerRect.top)); + const reserve = Math.max(72, overlap + 16); + output.style.setProperty("--thread-scroll-padding-bottom", `${reserve}px`); + panel.style.setProperty("--thread-scroll-padding-bottom", `${reserve}px`); + updateScrollToBottom(output); + } + + function timelineVisibleBounds(output) { + if (!output) return null; + const outputRect = output.getBoundingClientRect(); + const composer = $("messageForm"); + const composerRect = composer?.getBoundingClientRect?.(); + const top = outputRect.top + 8; + const composerTop = composerRect && Number.isFinite(composerRect.top) ? composerRect.top - 10 : outputRect.bottom - 8; + const bottom = Math.min(outputRect.bottom - 8, composerTop); + return { top, bottom: Math.max(top, bottom) }; + } + + // Keep an expanding activity row inside the portion of the transcript that + // is actually readable above the composer. This runs across the measured + // height animation because one layout pass is not enough when streamed + // command output arrives at the same time. + function ensureTimelineVisible(target) { + const element = target?.nodeType === 1 ? target : null; + const output = element?.closest?.(".chat-scroll") || $("output"); + if (!element || !output || !element.isConnected || !output.isConnected) return; + const bounds = timelineVisibleBounds(output); + if (!bounds) return; + const rect = element.getBoundingClientRect(); + let delta = 0; + if (rect.top < bounds.top) delta = rect.top - bounds.top; + else if (rect.bottom > bounds.bottom) delta = rect.bottom - bounds.bottom; + if (!Number.isFinite(delta) || Math.abs(delta) < 0.25) return; + const maxScroll = Math.max(0, output.scrollHeight - output.clientHeight); + output.scrollTop = Math.max(0, Math.min(maxScroll, output.scrollTop + delta)); + updateScrollToBottom(output); + } + + function scheduleTimelineReveal(target, duration = 220) { + const element = target?.nodeType === 1 ? target : null; + if (!element) return; + if (typeof state.timelineRevealCancel === "function") state.timelineRevealCancel(); + const deadline = motionClock() + motionDuration(duration) + 90; + let frame = null; + let timer = null; + let disposed = false; + const tick = () => { + if (disposed) return; + frame = null; + ensureTimelineVisible(element); + if (motionClock() < deadline) frame = scheduleFrame(tick); + }; + const cancel = () => { + if (disposed) return; + disposed = true; + cancelFrame(frame); + frame = null; + if (timer !== null) clearTimeout(timer); + timer = null; + if (state.timelineRevealCancel === cancel) state.timelineRevealCancel = null; + }; + state.timelineRevealCancel = cancel; + tick(); + timer = setTimeout(cancel, Math.max(180, motionDuration(duration) + 130)); + } + + function comparableText(value) { + return String(value ?? "").replace(/\s+/g, " ").trim(); + } + + function hasRenderedMessage(text, turnId, role) { + const target = comparableText(text); + if (!target) return false; + const output = $("output"); + if (!output) return false; + return [...output.querySelectorAll(`.message.${role}`)].some((article) => { + if (turnId && article.dataset.turnId && article.dataset.turnId !== String(turnId)) return false; + return comparableText(article.dataset.rawText) === target; + }); + } + + function hasRenderedCompletedMessage(text, turnId, role) { + const target = comparableText(text); + if (!target) return false; + const output = $("output"); + if (!output) return false; + return [...output.querySelectorAll(`.message.${role}`)].some((article) => { + if (article.classList.contains("streaming")) return false; + if (turnId && article.dataset.turnId && article.dataset.turnId !== String(turnId)) return false; + const raw = comparableText(article.dataset.rawText); + return raw === target || raw.includes(target) || target.includes(raw); + }); + } + + function appendInlineMarkdown(parent, source) { + const pattern = /(\[[^\]]+\]\(https?:\/\/[^)\s]+\)|`[^`\n]+`|\*\*[^*\n]+\*\*|__[^_\n]+__|~~[^~\n]+~~|\*[^*\n]+\*|_[^_\n]+_)/g; + let cursor = 0; + const appendText = (value) => { + const parts = String(value).split("\n"); + parts.forEach((part, index) => { + if (part) parent.append(document.createTextNode(part)); + if (index < parts.length - 1) parent.append(document.createElement("br")); + }); + }; + for (const match of String(source).matchAll(pattern)) { + if (match.index > cursor) appendText(String(source).slice(cursor, match.index)); + const token = match[0]; + if (token.startsWith("[") && token.endsWith(")")) { + const split = token.match(/^\[([^\]]+)\]\((https?:\/\/[^)\s]+)\)$/); + if (split) { + const link = document.createElement("a"); + link.href = split[2]; + link.target = "_blank"; + link.rel = "noopener noreferrer"; + link.textContent = split[1]; + parent.append(link); + } else appendText(token); + } else if (token.startsWith("`") && token.endsWith("`")) { + const code = document.createElement("code"); + code.textContent = token.slice(1, -1); + parent.append(code); + } else if (token.startsWith("**") || token.startsWith("__")) { + const strong = document.createElement("strong"); + strong.textContent = token.slice(2, -2); + parent.append(strong); + } else if (token.startsWith("~~")) { + const deleted = document.createElement("del"); + deleted.textContent = token.slice(2, -2); + parent.append(deleted); + } else if (token.startsWith("*") || token.startsWith("_")) { + const emphasis = document.createElement("em"); + emphasis.textContent = token.slice(1, -1); + parent.append(emphasis); + } else appendText(token); + cursor = match.index + token.length; + } + if (cursor < String(source).length) appendText(String(source).slice(cursor)); + } + + function tableCells(line) { + let value = String(line || "").trim(); + if (value.startsWith("|")) value = value.slice(1); + if (value.endsWith("|")) value = value.slice(0, -1); + return value.split("|").map((cell) => cell.trim()); + } + + function isTableDivider(line) { + const cells = tableCells(line); + return cells.length > 0 && cells.every((cell) => /^:?-{3,}:?$/.test(cell)); + } + + function appendTable(container, headerLine, bodyLines) { + const table = document.createElement("table"); + const thead = document.createElement("thead"); + const header = document.createElement("tr"); + for (const cell of tableCells(headerLine)) { + const th = document.createElement("th"); + appendInlineMarkdown(th, cell); + header.append(th); + } + thead.append(header); + table.append(thead); + const tbody = document.createElement("tbody"); + for (const line of bodyLines) { + const row = document.createElement("tr"); + for (const cell of tableCells(line)) { + const td = document.createElement("td"); + appendInlineMarkdown(td, cell); + row.append(td); + } + tbody.append(row); + } + table.append(tbody); + container.append(table); + } + + /** + * Render the small, safe Markdown subset used by Codex messages. The + * official webview uses a full Markdown/ProseMirror pipeline; the relay + * intentionally keeps this browser-side renderer dependency-free and never + * assigns untrusted text to innerHTML. + */ + function renderMarkdown(container, source) { + container.replaceChildren(); + const lines = String(source ?? "").replace(/\r\n?/g, "\n").split("\n"); + let index = 0; + const addParagraph = (paragraph) => { + if (!paragraph.length) return; + const element = document.createElement("p"); + appendInlineMarkdown(element, paragraph.join("\n")); + container.append(element); + }; + while (index < lines.length) { + const line = lines[index]; + if (!line.trim()) { index += 1; continue; } + const fence = line.match(/^\s*```\s*([\w.+-]*)\s*$/); + if (fence) { + index += 1; + const codeLines = []; + while (index < lines.length && !/^\s*```\s*$/.test(lines[index])) codeLines.push(lines[index++]); + if (index < lines.length) index += 1; + const pre = document.createElement("pre"); + const code = document.createElement("code"); + if (fence[1]) code.dataset.language = fence[1]; + code.textContent = codeLines.join("\n"); + pre.append(code); + container.append(pre); + continue; + } + if (index + 1 < lines.length && line.includes("|") && isTableDivider(lines[index + 1])) { + const body = []; + index += 2; + while (index < lines.length && lines[index].trim() && lines[index].includes("|")) body.push(lines[index++]); + appendTable(container, line, body); + continue; + } + if (/^\s*(?:---+|___+|\*\*\*+)\s*$/.test(line)) { + container.append(document.createElement("hr")); + index += 1; + continue; + } + const heading = line.match(/^\s*(#{1,3})\s+(.+?)\s*#*$/); + if (heading) { + const element = document.createElement(`h${heading[1].length}`); + appendInlineMarkdown(element, heading[2]); + container.append(element); + index += 1; + continue; + } + if (/^\s*>\s?/.test(line)) { + const quote = document.createElement("blockquote"); + while (index < lines.length && /^\s*>\s?/.test(lines[index])) { + const paragraph = document.createElement("p"); + appendInlineMarkdown(paragraph, lines[index].replace(/^\s*>\s?/, "")); + quote.append(paragraph); + index += 1; + } + container.append(quote); + continue; + } + const list = line.match(/^\s*([-*+]|\d+[.)])\s+(.+)$/); + if (list) { + const ordered = /^\d/.test(list[1]); + const listElement = document.createElement(ordered ? "ol" : "ul"); + while (index < lines.length) { + const item = lines[index].match(/^\s*([-*+]|\d+[.)])\s+(.+)$/); + if (!item || /^\d/.test(item[1]) !== ordered) break; + const li = document.createElement("li"); + const task = item[2].match(/^\[([ xX])\]\s+(.+)$/); + if (task) { + li.className = "task-list-item"; + const checkbox = document.createElement("input"); + checkbox.type = "checkbox"; + checkbox.checked = task[1].toLowerCase() === "x"; + checkbox.disabled = true; + checkbox.setAttribute("aria-label", t(checkbox.checked ? "已完成" : "未完成")); + li.append(checkbox); + appendInlineMarkdown(li, task[2]); + } else appendInlineMarkdown(li, item[2]); + listElement.append(li); + index += 1; + } + container.append(listElement); + continue; + } + const paragraph = [line]; + index += 1; + while (index < lines.length + && lines[index].trim() + && !/^\s*```/.test(lines[index]) + && !/^\s*(#{1,3})\s+/.test(lines[index]) + && !/^\s*>\s?/.test(lines[index]) + && !/^\s*([-*+]|\d+[.)])\s+/.test(lines[index])) { + paragraph.push(lines[index++]); + } + addParagraph(paragraph); + } + } + + function renderMessageBody(body, text, role, tone, kind) { + const value = String(text ?? ""); + body.dataset.rawText = value; + // User-authored prompts use the same safe Markdown subset as assistant + // messages. Tool/status rows stay literal so command output cannot be + // mistaken for formatted content. + const markdownRole = role === "assistant" || role === "user" + || kind === "reasoning" || kind === "plan" || kind === "subagent" || kind === "commentary"; + body.classList.toggle("markdown-body", markdownRole && tone !== "meta" && tone !== "error" && kind !== "tool"); + body.classList.toggle("diff-body", kind === "edit"); + if (kind === "edit") { + body.replaceChildren(); + const output = document.createElement("pre"); + output.className = "diff-output"; + for (const line of value.replace(/\r\n?/g, "\n").split("\n")) { + const row = document.createElement("span"); + row.className = line.startsWith("+") && !line.startsWith("+++") + ? "diff-line added" + : line.startsWith("-") && !line.startsWith("---") + ? "diff-line removed" + : line.startsWith("@@") || line.startsWith("diff ") || line.startsWith("[") + ? "diff-line context" + : "diff-line"; + row.textContent = line || " "; + output.append(row); + } + body.append(output); + } else if (kind === "tool" || kind === "read") body.textContent = value; + else if (markdownRole && tone !== "meta" && tone !== "error") renderMarkdown(body, value); + else body.textContent = value; + } + + function createTerminalIcon() { + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.className.baseVal = "activity-summary-icon terminal-icon"; + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + const frame = document.createElementNS("http://www.w3.org/2000/svg", "rect"); + frame.setAttribute("x", "2.25"); + frame.setAttribute("y", "2.75"); + frame.setAttribute("width", "11.5"); + frame.setAttribute("height", "10.5"); + frame.setAttribute("rx", "1.4"); + const prompt = document.createElementNS("http://www.w3.org/2000/svg", "path"); + prompt.setAttribute("d", "m4.5 6 2 2-2 2"); + const cursor = document.createElementNS("http://www.w3.org/2000/svg", "path"); + cursor.setAttribute("d", "M8.5 10h2.5"); + svg.append(frame, prompt, cursor); + return svg; + } + + function createSubagentIcon() { + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.className.baseVal = "activity-summary-icon subagent-summary-icon"; + svg.setAttribute("viewBox", "0 0 24 24"); + svg.setAttribute("aria-hidden", "true"); + // The installed Codex webview uses the filled blossom mark for both + // sub-agent activity rows and the composer agent rail. Keep the path + // inline so the standalone relay page does not depend on VS Code's URI + // resolver or on an extension-owned asset path. + const mark = document.createElementNS("http://www.w3.org/2000/svg", "path"); + mark.setAttribute("fill", "currentColor"); + mark.setAttribute("d", "M13.795 23.856q-1.188 0-2.256-.448a6.1 6.1 0 0 1-1.9-1.247 5.8 5.8 0 0 1-1.875.306 5.8 5.8 0 0 1-2.944-.777 6.1 6.1 0 0 1-2.184-2.12q-.807-1.34-.808-2.99 0-.682.19-1.482a6.3 6.3 0 0 1-1.472-2.002 5.76 5.76 0 0 1 .024-4.85q.546-1.177 1.52-2.024a5.5 5.5 0 0 1 2.303-1.2A5.55 5.55 0 0 1 5.485 2.62 6.06 6.06 0 0 1 7.575.925 5.85 5.85 0 0 1 10.21.313q1.187 0 2.255.447a6.1 6.1 0 0 1 1.9 1.248 5.8 5.8 0 0 1 1.875-.306q1.59 0 2.944.776a5.9 5.9 0 0 1 2.16 2.12q.832 1.34.832 2.99 0 .682-.19 1.483a6.2 6.2 0 0 1 1.472 2.024q.522 1.13.522 2.378 0 1.272-.546 2.449a6.1 6.1 0 0 1-1.543 2.048 5.45 5.45 0 0 1-2.28 1.177 5.4 5.4 0 0 1-1.115 2.402 5.8 5.8 0 0 1-2.066 1.695 5.85 5.85 0 0 1-2.635.612M7.93 20.913q1.188 0 2.066-.495l4.463-2.542a.52.52 0 0 0 .238-.448v-2.024L8.95 18.676a.97.97 0 0 1-1.044 0L3.419 16.11a.7.7 0 0 1-.024.165v.282q0 1.201.57 2.213.594.99 1.639 1.554 1.044.59 2.326.589m.238-3.838q.143.07.26.07a.46.46 0 0 0 .238-.07l1.781-1.012-5.722-3.296q-.522-.306-.522-.918v-5.11a4.27 4.27 0 0 0-1.9 1.602 4.13 4.13 0 0 0-.712 2.354q0 1.155.594 2.213.593 1.06 1.543 1.601zm5.627 5.227q1.258 0 2.279-.565a4.25 4.25 0 0 0 1.614-1.554q.594-.99.594-2.213v-5.085q0-.283-.237-.424l-1.805-1.036v6.568q0 .613-.522.919l-4.487 2.566q1.163.825 2.564.824m.902-8.617v-3.202l-2.683-1.507-2.707 1.507v3.202l2.707 1.507zm-6.933-7.51q0-.612.522-.918l4.488-2.567a4.34 4.34 0 0 0-2.564-.824q-1.26 0-2.28.565a4.25 4.25 0 0 0-1.614 1.554q-.57.99-.57 2.213v5.062q0 .283.237.447l1.781 1.036zm12.061 11.253a4.13 4.13 0 0 0 1.876-1.6 4.2 4.2 0 0 0 .712-2.355q0-1.154-.593-2.213-.594-1.06-1.544-1.6l-4.44-2.543q-.142-.095-.26-.071a.46.46 0 0 0-.238.07l-1.78.99 5.745 3.319q.26.141.38.377a.9.9 0 0 1 .142.518zm-4.772-11.96q.522-.33 1.045 0l4.51 2.614v-.424q0-1.13-.57-2.142a4.1 4.1 0 0 0-1.59-1.648q-1.02-.613-2.374-.613-1.187 0-2.066.495L9.545 6.292a.52.52 0 0 0-.238.448v2.025z"); + svg.append(mark); + return svg; + } + + function createReadIcon() { + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.className.baseVal = "activity-summary-icon read-summary-icon"; + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", "M2.5 4.5h4l1.2 1.4h5.8v6.2c0 .8-.5 1.4-1.3 1.4H3.8c-.8 0-1.3-.6-1.3-1.4V4.5Z"); + const top = document.createElementNS("http://www.w3.org/2000/svg", "path"); + top.setAttribute("d", "M2.5 6h11"); + svg.append(path, top); + return svg; + } + + // The official sub-agent activity renderer keeps the agent name in a + // bordered chip and places the lifecycle text beside it. Keeping those two + // nodes separate also lets a running row update only its status without + // replacing the blossom icon or the accessible label. + function setSubagentSummary(summary, name, statusText = "") { + if (!summary) return; + let chip = summary.querySelector(".subagent-summary-chip"); + let label = chip?.querySelector(".subagent-summary-label"); + let status = summary.querySelector(".subagent-summary-status"); + if (!chip || !label || !status) { + chip = document.createElement("span"); + chip.className = "subagent-summary-chip"; + const icon = createSubagentIcon(); + label = document.createElement("span"); + label.className = "subagent-summary-label"; + chip.append(icon, label); + status = document.createElement("span"); + status.className = "subagent-summary-status"; + summary.replaceChildren(chip, status); + } + label.textContent = String(name || t("子代理")); + status.textContent = statusText ? t(statusText) : ""; + summary.setAttribute("aria-label", [label.textContent, status.textContent].filter(Boolean).join(" ")); + } + + function setActivitySummary(summary, text, kind = "") { + if (!summary) return; + if (kind === "subagent") { + const value = String(text || ""); + const match = value.match(/^(.*?)(?:\s+(已开始工作|已完成|失败|已中断|等待中|处理中|Started working|Completed|Failed|Interrupted|Waiting|Working))$/); + setSubagentSummary(summary, match ? match[1] : value || t("子代理"), match ? t(match[2]) : ""); + return; + } + if (kind !== "tool" && kind !== "read" && kind !== "subagent") { + summary.textContent = String(text || ""); + return; + } + let icon = summary.querySelector(".activity-summary-icon"); + let label = summary.querySelector(".activity-summary-label"); + if (!icon || !label) { + icon = kind === "subagent" ? createSubagentIcon() : kind === "read" ? createReadIcon() : createTerminalIcon(); + label = document.createElement("span"); + label.className = "activity-summary-label"; + summary.replaceChildren(icon, label); + } + label.textContent = String(text || ""); + } + + function stripShellQuotes(value) { + let text = String(value || "").trim(); + let changed = true; + while (changed) { + changed = false; + if (text.startsWith("$'") && text.endsWith("'")) { + text = text.slice(2, -1).replace(/\\'/g, "'"); + changed = true; + } else if ((text.startsWith("'") && text.endsWith("'")) + || (text.startsWith('"') && text.endsWith('"'))) { + const quoted = text.startsWith('"'); + text = text.slice(1, -1); + if (quoted) text = text.replace(/\\"/g, '"'); + changed = true; + } + } + return text.trim(); + } + + // The official command renderer hides the shell bootstrap used by the IPC + // runner (`/bin/zsh -lc '…'`) and shows the actual user command instead. + function terminalCommandText(value) { + const command = stripShellQuotes(value); + const match = command.match(/^(?:.*[/\\])?(?:bash|cmd(?:\.exe)?|fish|powershell(?:\.exe)?|pwsh(?:\.exe)?|sh|zsh)\s+-lc\s+([\s\S]+)$/i); + return match ? stripShellQuotes(match[1]) : command; + } + + function addMessageActions(article, enabled = true) { + // Tool/reasoning rows have their own disclosure controls and should not + // grow a second action rail. Editing an earlier turn is intentionally not + // exposed until the relay can perform the official branch/edit operation. + if (!enabled || (!article.classList.contains("user") && !article.classList.contains("assistant"))) return null; + if (article.classList.contains("assistant") && article.classList.contains("streaming")) return null; + const actions = document.createElement("div"); + actions.className = "message-actions"; + const copy = document.createElement("button"); + copy.type = "button"; + copy.className = "message-action"; + copy.title = t("复制消息"); + copy.setAttribute("aria-label", t("复制消息")); + const setCopyIcon = (copied = false) => { + copy.replaceChildren(); + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + if (copied) { + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", "m3.5 8.2 2.7 2.7 6.3-6.3"); + svg.append(path); + } else { + const back = document.createElementNS("http://www.w3.org/2000/svg", "path"); + back.setAttribute("d", "M5.5 5.5V4c0-.8.7-1.5 1.5-1.5h5c.8 0 1.5.7 1.5 1.5v5c0 .8-.7 1.5-1.5 1.5h-1.5"); + const front = document.createElementNS("http://www.w3.org/2000/svg", "rect"); + front.setAttribute("x", "2.5"); + front.setAttribute("y", "5.5"); + front.setAttribute("width", "7.5"); + front.setAttribute("height", "7.5"); + front.setAttribute("rx", "1.2"); + svg.append(back, front); + } + copy.append(svg); + }; + setCopyIcon(); + copy.addEventListener("click", async () => { + const value = article.dataset.rawText || ""; + try { + await navigator.clipboard?.writeText(value); + setCopyIcon(true); + setTimeout(() => { setCopyIcon(false); }, 1_200); + } catch { + copy.textContent = "!"; + } + }); + actions.append(copy); + article.append(actions); + return actions; + } + + function appendMessage(text, role = "assistant", tone = "text", meta = "", options = {}) { + if (text === undefined || text === null) return null; + const output = $("output"); + const follow = shouldFollowOutput(output); + const article = document.createElement("article"); + article.className = `message ${role}${tone && tone !== "text" ? ` ${tone}` : ""}`; + if (options.kind) article.dataset.kind = options.kind; + if (options.status) article.dataset.status = String(options.status); + if (options.turnId) article.dataset.turnId = String(options.turnId); + if (options.structuredKey) article.dataset.structuredKey = String(options.structuredKey); + if (options.activityKey) article.dataset.activityKey = String(options.activityKey); + if (options.itemId !== undefined && options.itemId !== null) article.dataset.itemId = String(options.itemId); + if (options.itemType) article.dataset.itemType = String(options.itemType); + if (options.agentThreadId) article.dataset.agentThreadId = String(options.agentThreadId); + if (options.timestamp) { + const parsedTimestamp = timestampMs(options.timestamp); + if (parsedTimestamp !== null) article.dataset.timestamp = String(parsedTimestamp); + } + article.dataset.rawText = String(text); + const content = document.createElement("div"); + content.className = "message-content"; + let body = document.createElement("div"); + body.className = "message-body"; + let details = null; + const initiallyOpen = options.open !== false; + if (options.collapsible) { + details = document.createElement("details"); + details.className = `message-details ${options.kind || ""}`; + const summary = document.createElement("summary"); + const summaryText = options.summary || options.label || t("详情"); + if (options.kind === "tool" || options.kind === "read" || options.kind === "subagent") { + setActivitySummary(summary, summaryText, options.kind); + if (options.command) summary.title = terminalCommandText(options.command); + } else { + summary.textContent = summaryText; + if (options.command && !options.summary) { + // Translate only the renderer-owned label. The command is host data + // and must remain byte-for-byte unchanged in every locale. + summary.textContent = `${t(options.label || "命令")} · ${String(options.command)}`; + summary.title = String(options.command); + } + } + body.classList.add("details-body"); + details.append(summary, body); + content.append(details); + } else content.append(body); + renderMessageBody(body, text, role, tone, options.kind); + article.append(content); + const actions = addMessageActions(article, options.showActions !== false); + const metaParts = []; + if (meta) metaParts.push(String(meta)); + if (options.showTimestamp === true && options.timestamp) { + const timestamp = formatMessageTime(options.timestamp); + if (timestamp) metaParts.push(timestamp); + } + if (options.showDuration === true && options.durationMs !== undefined && options.durationMs !== null) { + const duration = formatDuration(options.durationMs); + if (duration) metaParts.push(uiWithRaw("用时 ", "Worked for ", duration)); + } + if (metaParts.length) { + const stamp = document.createElement("div"); + stamp.className = "message-meta"; + stamp.textContent = metaParts.join(" · "); + if (actions) actions.prepend(stamp); + else content.append(stamp); + } + output.append(article); + if (details) installDisclosure(details, initiallyOpen); + if (follow) scrollOutput(output, true); + return { article, content: body, wrapper: content }; + } + + function appendDateSeparator(timestamp, options = {}) { + const parsedTimestamp = timestampMs(timestamp); + const role = String(options.role || "user"); + // The official timestamp projection treats an item without a usable time + // as an adjacency break. Do not carry the previous role/time across an + // untimestamped item, otherwise a later assistant message can inherit an + // unrelated 10-minute/1-hour gap. + if (parsedTimestamp === null) { + state.lastRenderedDateKey = ""; + state.lastRenderedTimestamp = null; + state.lastRenderedRole = ""; + if (role === "user") state.hasRenderedUser = true; + return; + } + const previousTimestamp = state.lastRenderedTimestamp; + const previousRole = state.lastRenderedRole; + const hour = 60 * 60 * 1000; + const tenMinutes = 10 * 60 * 1000; + const gap = previousTimestamp === null ? null : parsedTimestamp - previousTimestamp; + // This mirrors the official timestamps projection: a first/next user turn + // is separated only after a substantial pause, while consecutive assistant + // entries can be separated after a shorter gap. + const threshold = previousRole === "assistant" + ? role === "user" ? hour : tenMinutes + : Infinity; + const firstUserIsOld = !state.hasRenderedUser && role === "user" + && Date.now() - parsedTimestamp > hour; + // The official local composer always gives the first user message a + // centered date/time anchor, even when that message was sent today. It is + // also the visual boundary that separates a freshly attached history from + // the input composer, so do not hide it merely because the turn is recent. + const firstUserTurn = !state.hasRenderedUser && role === "user"; + const show = options.force === true || options.breaksPreviousAdjacency === true || firstUserIsOld + || firstUserTurn + || (gap !== null && gap > 0 && previousRole === "assistant" && gap > threshold); + if (show && state.lastDateSeparatorTimestamp !== parsedTimestamp) { + const label = formatMessageDate(parsedTimestamp); + if (label) { + const separator = document.createElement("div"); + separator.className = "date-separator"; + separator.setAttribute("role", "separator"); + separator.setAttribute("aria-label", label); + const time = document.createElement("time"); + time.dateTime = new Date(parsedTimestamp).toISOString(); + const splitAt = label.lastIndexOf(" "); + if (splitAt > 0 && splitAt < label.length - 1) { + const dateLabel = document.createElement("span"); + dateLabel.className = "date-label"; + dateLabel.textContent = label.slice(0, splitAt); + const timeLabel = document.createElement("span"); + timeLabel.className = "date-time"; + timeLabel.textContent = label.slice(splitAt + 1); + time.append(dateLabel, " ", timeLabel); + } else time.textContent = label; + separator.append(time); + $("output").append(separator); + state.lastDateSeparatorTimestamp = parsedTimestamp; + } + } + state.lastRenderedDateKey = messageDateKey(parsedTimestamp); + state.lastRenderedTimestamp = parsedTimestamp; + state.lastRenderedRole = role; + if (role === "user") state.hasRenderedUser = true; + } + + function turnDividerLabel(status, durationMs) { + const duration = elapsedDuration(durationMs); + const normalized = normalizeActivityStatus(status, "completed"); + if (normalized === "interrupted") return duration + ? uiWithRaw("你在 ", "You stopped after ", duration, " 后停止了", "") + : uiText("你停止了工作", "You stopped working"); + if (normalized === "failed") return duration + ? uiWithRaw("执行失败 · ", "Action failed · ", duration) + : uiText("执行失败", "Action failed"); + if (normalized === "inProgress") { + const visible = elapsedDuration(durationMs); + return visible ? uiWithRaw("用时 ", "Worked for ", visible) : uiText("正在处理", "Working"); + } + return duration ? uiWithRaw("用时 ", "Worked for ", duration) : uiText("已完成", "Completed"); + } + + // A worked-for disclosure only has meaning when the turn owns at least one + // concrete activity row. Keep this cleanup centralized so a transient + // status-only row cannot leave an empty divider behind after it is retired. + function removeEmptyTurnDivider(turnId) { + const key = String(turnId || ""); + if (!key) return; + const output = $("output"); + if (!output) return; + const hasActivity = [...output.querySelectorAll(".message.activity")] + .some((entry) => entry.dataset.turnId === key); + if (hasActivity) return; + const divider = state.turnDividers.get(key) + || [...output.querySelectorAll(".turn-divider")] + .find((entry) => entry.dataset.turnId === key) + || null; + if (divider?.parentNode) divider.remove(); + state.turnDividers.delete(key); + } + + function retireActivity(activity) { + if (!activity) return; + const key = String(activity.key || ""); + const turnId = String(activity.turnId || ""); + if (key && state.activities.get(key) === activity) state.activities.delete(key); + if (key) state.commandDisclosure.delete(key); + if (state.liveActivityKey === key) state.liveActivityKey = null; + if (state.activeAssistantActivityKey === key) { + state.activeAssistantBody = null; + state.activeAssistantStream = null; + state.activeAssistantText = ""; + state.activeAssistantActivityKey = null; + } + if (activity.article?.parentNode) activity.article.remove(); + removeEmptyTurnDivider(turnId); + stopActivityTimerIfIdle(); + updateScrollToBottom($("output")); + } + + function retireStatusOnlyActivities(turnId = "") { + const key = String(turnId || ""); + for (const activity of [...state.activities.values()]) { + if (activity.statusOnly !== true || activity.concrete === true) continue; + if (key && activity.turnId && activity.turnId !== key) continue; + retireActivity(activity); + } + } + + function appendTurnDivider(turnId, status = "completed", durationMs, beforeArticle = null, options = {}) { + const key = String(turnId || `anonymous-${state.activitySequence}`); + const output = $("output"); + const autoPosition = beforeArticle === null; + const turnEntries = [...output.querySelectorAll(".message")] + .filter((entry) => entry.dataset.turnId === key); + // The official worked-for disclosure owns the activity portion of a turn. + // A final assistant answer remains visible below the disclosure; commentary + // and concrete work rows are the entries that collapse underneath it. + const firstTurnActivity = turnEntries.find((entry) => entry.classList.contains("activity")) || null; + // Do not manufacture an empty worked-for row for a final-only turn. This + // can happen when completion metadata arrives before any item lifecycle + // event, or when a transient status row has just been retired. + if (!firstTurnActivity) { + removeEmptyTurnDivider(key); + return null; + } + const firstTurnContent = firstTurnActivity + || turnEntries.find((entry) => !entry.classList.contains("user")) + || null; + if (firstTurnContent) beforeArticle = firstTurnContent; + else if (autoPosition) beforeArticle = turnEntries.at(-1) || null; + let divider = state.turnDividers.get(key); + if (!divider) { + divider = [...output.querySelectorAll(".turn-divider")].find((entry) => entry.dataset.turnId === key) || null; + } + if (!divider) { + divider = document.createElement("div"); + divider.className = "turn-divider"; + divider.dataset.turnId = key; + const button = document.createElement("button"); + button.type = "button"; + button.className = "turn-divider-toggle"; + const initialStatus = normalizeActivityStatus(status, "completed"); + const rememberedExpansion = state.turnExpansion.get(key); + const initialExpanded = rememberedExpansion !== undefined + ? rememberedExpansion + : options.defaultExpanded === true || initialStatus === "inProgress"; + if (rememberedExpansion !== undefined || options.defaultExpanded === true) { + button.dataset.userToggled = rememberedExpansion !== undefined ? "true" : "false"; + } + button.setAttribute("aria-expanded", String(initialExpanded)); + const label = document.createElement("span"); + label.className = "turn-divider-label"; + const icon = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + icon.setAttribute("viewBox", "0 0 16 16"); + icon.setAttribute("aria-hidden", "true"); + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", "m6 3 5 5-5 5"); + icon.append(path); + button.append(label, icon); + const rule = document.createElement("div"); + rule.className = "turn-divider-rule"; + divider.append(button, rule); + button.addEventListener("click", () => { + const expanded = button.getAttribute("aria-expanded") === "true"; + const next = !expanded; + // Hydration may choose an expanded latest turn for parity with the + // official panel. Once the reader explicitly toggles it, preserve + // that choice across subsequent structured snapshots. + button.dataset.userToggled = "true"; + state.turnExpansion.set(key, next); + button.setAttribute("aria-expanded", String(next)); + setTurnActivityVisibility(key, next); + if (next) { + const firstVisibleEntry = [...output.querySelectorAll(".message.activity")] + .filter((entry) => entry.dataset.turnId === key) + .find((entry) => !entry.classList.contains("turn-collapsed")); + if (firstVisibleEntry) scheduleTimelineReveal(firstVisibleEntry, 280); + } + }); + if (beforeArticle && beforeArticle.parentNode === output) output.insertBefore(divider, beforeArticle); + else output.append(divider); + state.turnDividers.set(key, divider); + } else if (beforeArticle === false && divider.parentNode === output && output.lastElementChild !== divider) { + // An additional activity can arrive after an already-created live + // divider. Explicit `false` means place the divider at the current tail. + output.append(divider); + } else if (beforeArticle && beforeArticle.parentNode === output && divider !== beforeArticle) { + // Completion can arrive after the final assistant article. Move the + // existing disclosure instead of creating a second one at the tail. + output.insertBefore(divider, beforeArticle); + } + const label = divider.querySelector(".turn-divider-label"); + if (durationMs !== null && durationMs !== undefined && finiteNumber(durationMs) !== null) { + divider.dataset.durationMs = String(Math.max(0, Number(durationMs))); + } + const storedDuration = finiteNumber(durationMs, divider.dataset.durationMs); + const liveDuration = durationMs === null || durationMs === undefined + ? key === state.turnId && state.turnWorkStartedAt !== null + ? Math.max(0, Date.now() - state.turnWorkStartedAt) + : storedDuration + : storedDuration; + if (label) label.textContent = turnDividerLabel(status, liveDuration); + const normalizedStatus = normalizeActivityStatus(status, "completed"); + divider.dataset.status = normalizedStatus; + const toggle = divider.querySelector(".turn-divider-toggle"); + if (options.defaultExpanded === true && toggle?.dataset.userToggled !== "true") { + toggle.dataset.defaultExpanded = "true"; + } + if (normalizedStatus === "inProgress" && toggle?.dataset.userToggled !== "true") { + toggle.setAttribute("aria-expanded", "true"); + } + const rememberedExpansion = state.turnExpansion.get(key); + // The latest completed turn is expanded on first render in the official + // panel. A remembered click always wins, including an explicit collapse. + const expanded = rememberedExpansion !== undefined + ? rememberedExpansion + : toggle?.dataset.userToggled === "true" + ? toggle.getAttribute("aria-expanded") === "true" + : normalizedStatus === "inProgress" || toggle?.dataset.defaultExpanded === "true"; + setTurnActivityVisibility(key, expanded); + return divider; + } + + function appendCompletedTurnDivider(turnId, status = "completed", durationMs) { + const key = String(turnId || ""); + if (!key) return null; + const output = $("output"); + if (!output) return null; + const hasTurnMessage = [...output.querySelectorAll(".message")] + .some((article) => article.dataset.turnId === key); + if (!hasTurnMessage) return null; + const hasTurnActivity = [...output.querySelectorAll(".message.activity")] + .some((article) => article.dataset.turnId === key); + // A final assistant article by itself is not a worked-for group. The + // official UI shows its duration in the assistant block without an empty + // disclosure row above it. + if (!hasTurnActivity) return null; + const firstContent = [...output.querySelectorAll(".message")] + .find((entry) => entry.dataset.turnId === key && entry.classList.contains("activity")) + || [...output.querySelectorAll(".message")] + .find((entry) => entry.dataset.turnId === key && !entry.classList.contains("user")) + || null; + return appendTurnDivider(key, status, durationMs, firstContent); + } + + // Structured snapshots can append bookkeeping or collaboration items after + // the final assistant item. Reconcile the outer worked-for boundary after + // hydration so every activity row remains in the same disclosure group. + function reconcileTurnDividers() { + const output = $("output"); + if (!output) return; + const turnIds = new Set( + [...output.querySelectorAll(".message.activity")] + .map((article) => article.dataset.turnId) + .filter(Boolean), + ); + for (const key of turnIds) { + const activities = [...output.querySelectorAll(".message.activity")] + .filter((article) => article.dataset.turnId === key); + if (!activities.length) continue; + let divider = state.turnDividers.get(key) + || [...output.querySelectorAll(".turn-divider")].find((entry) => entry.dataset.turnId === key) + || null; + const finalAssistant = [...output.querySelectorAll(".message")] + .filter((article) => article.dataset.turnId === key + && !article.classList.contains("activity") + && article.dataset.kind === "assistant") + .at(-1) + || null; + const activeTurn = key === state.turnId + && state.turnStartedAt !== null + && ["active", "waiting"].includes(state.turnStatus); + if (!divider) { + divider = appendTurnDivider(key, activeTurn ? "inProgress" : "completed", state.workedDurationMs, activities[0]); + } else { + if (activities[0].parentNode === output && divider.parentNode === output && divider !== activities[0]) { + output.insertBefore(divider, activities[0]); + } + const toggle = divider.querySelector(".turn-divider-toggle"); + if (toggle && toggle.dataset.userToggled !== "true") { + const rememberedExpansion = state.turnExpansion.get(key); + toggle.setAttribute("aria-expanded", String(rememberedExpansion !== undefined + ? rememberedExpansion + : activeTurn)); + } + const status = activeTurn ? "inProgress" : normalizeActivityStatus(divider.dataset.status, "completed"); + divider.dataset.status = status; + setTurnActivityVisibility(key, toggle?.getAttribute("aria-expanded") === "true"); + } + // Keep a divider before the activity group even when a final answer was + // rendered first. The answer itself intentionally stays outside it. + if (divider && activities[0].parentNode === output && divider.nextSibling !== activities[0]) { + output.insertBefore(divider, activities[0]); + } + if (finalAssistant && divider) divider.dataset.hasFinalAssistant = "true"; + if (finalAssistant) { + // Keep trailing bookkeeping/collaboration rows inside the same outer + // group. Moving only rows that currently follow the final response + // preserves the chronological order of the already-correct prefix. + const trailing = activities.filter((activity) => ( + activity.compareDocumentPosition(finalAssistant) & Node.DOCUMENT_POSITION_PRECEDING + )); + for (const activity of trailing) output.insertBefore(activity, finalAssistant); + // The moves above can carry the first activity past the existing + // disclosure. Re-anchor it after every move so the worked-for button + // always remains the first node in the activity group. + if (activities[0].parentNode === output && divider.parentNode === output) { + output.insertBefore(divider, activities[0]); + } + } + } + } + + function expandLatestTurnActivity() { + const output = $("output"); + if (!output) return; + const dividers = [...output.querySelectorAll(".turn-divider")]; + const divider = dividers.at(-1); + // The latest worked-for group is expanded after a completed snapshot so + // the reader can see the activity stream immediately. Older groups stay + // compact, and an explicit user toggle always wins. + if (!divider) return; + const toggle = divider.querySelector(".turn-divider-toggle"); + if (!toggle || toggle.dataset.userToggled === "true") return; + const key = divider.dataset.turnId || ""; + if (!key || ![...output.querySelectorAll(".message.activity")].some((entry) => entry.dataset.turnId === key)) return; + const rememberedExpansion = state.turnExpansion.get(key); + const shouldExpand = rememberedExpansion !== undefined + ? rememberedExpansion + : divider.dataset.status === "inProgress" || divider === dividers.at(-1); + if (!shouldExpand) return; + if (toggle.getAttribute("aria-expanded") === "true") return; + toggle.setAttribute("aria-expanded", "true"); + setTurnActivityVisibility(key, true); + } + + function setTurnActivityVisibility(turnId, expanded) { + const key = String(turnId || ""); + if (!key) return; + const output = $("output"); + if (!output) return; + const divider = state.turnDividers.get(key) + || [...output.querySelectorAll(".turn-divider")].find((entry) => entry.dataset.turnId === key) + || null; + // Keep the final assistant response and its timestamp/actions mounted. + // Only rows explicitly marked as activity participate in the outer turn + // disclosure, matching the official conversation-blocks split. + const activities = [...output.querySelectorAll(".message.activity")] + .filter((entry) => entry.dataset.turnId === key); + if (!activities.length) return; + const knownState = activities.some((activity) => activity.dataset.turnExpanded !== undefined); + if (!knownState && !expanded) { + for (const activity of activities) { + activity.dataset.turnExpanded = "false"; + animateActivityArticle(activity, false, { immediate: true }); + } + return; + } + if (divider || activities[0]) preserveTimelineAnchor(divider || activities[0], 280); + for (const activity of activities) { + const previous = activity.dataset.turnExpanded === "true"; + activity.dataset.turnExpanded = String(Boolean(expanded)); + // The outer worked-for disclosure controls every non-user entry. Each + // nested command/reasoning disclosure keeps its own state, matching the + // official panel where opening a completed turn does not open every + // terminal body at once. + animateActivityArticle(activity, expanded, { + immediate: !knownState && !previous, + fromHeight: expanded && !previous ? 0 : undefined, + preserve: false, + }); + } + if (expanded) revealTerminalOutputs(key); + } + + function revealTerminalOutputs(turnId) { + const key = String(turnId || ""); + if (!key) return; + requestAnimationFrame(() => { + const output = $("output"); + if (!output) return; + for (const activity of output.querySelectorAll(".message.activity")) { + if (activity.dataset.turnId !== key) continue; + const terminal = activity.querySelector(".terminal-output"); + if (!terminal || terminal.scrollHeight <= terminal.clientHeight) continue; + // Hidden turn sections have no layout while hydrated. Once revealed, + // show the newest terminal lines just like the official reverse list. + terminal.scrollTop = terminal.scrollHeight; + updateTerminalOutputFade(terminal); + } + }); + } + + function ensureLiveTurnDivider(turnId) { + const key = String(turnId || ""); + if (!key) return null; + const output = $("output"); + const entries = [...output.querySelectorAll(".message")] + .filter((entry) => entry.dataset.turnId === key); + const firstContent = entries.find((entry) => entry.classList.contains("activity")) + || entries.find((entry) => !entry.classList.contains("user")) + || null; + if (firstContent) return appendTurnDivider(key, "inProgress", null, firstContent); + return appendTurnDivider(key, "inProgress", null, false); + } + + const isRecord = (value) => Boolean(value && typeof value === "object" && !Array.isArray(value)); + + function finiteNumber(...values) { + for (const value of values) { + if (value === null || value === undefined || value === "" || typeof value === "boolean") continue; + const number = Number(value); + if (Number.isFinite(number)) return number; + } + return null; + } + + function timestampMs(...values) { + let value = null; + for (const candidate of values) { + if (candidate === null || candidate === undefined || candidate === "" || typeof candidate === "boolean") continue; + if (typeof candidate === "string" && candidate.trim() && !/^\d+(?:\.\d+)?$/.test(candidate.trim())) { + const parsed = Date.parse(candidate); + if (Number.isFinite(parsed)) { value = parsed; break; } + } + const numeric = Number(candidate); + if (Number.isFinite(numeric)) { value = numeric; break; } + } + if (value === null || value <= 0) return null; + // Accept ISO epoch seconds from older followers as well as the current + // app-server's millisecond fields. + return value < 100_000_000_000 ? value * 1_000 : value; + } + + function workedDurationFor(value, fallback = null) { + const source = isRecord(value) ? value : {}; + const nestedTurn = isRecord(source.turn) ? source.turn : {}; + const nestedStatus = isRecord(source.status) ? source.status : {}; + const explicit = finiteNumber( + source.workedDurationMs, + source.workDurationMs, + source.workedForMs, + source.turnWorkedDurationMs, + source.worked_for_ms, + nestedTurn.workedDurationMs, + nestedTurn.workDurationMs, + nestedTurn.workedForMs, + nestedStatus.workedDurationMs, + nestedStatus.workDurationMs, + ); + if (explicit !== null) return Math.max(0, explicit); + const started = timestampMs( + source.firstTurnWorkItemStartedAtMs, + source.workStartedAtMs, + source.turnWorkStartedAtMs, + nestedTurn.firstTurnWorkItemStartedAtMs, + nestedTurn.workStartedAtMs, + nestedTurn.turnStartedAtMs, + ); + const completed = timestampMs( + source.finalAssistantStartedAtMs, + source.workCompletedAtMs, + nestedTurn.finalAssistantStartedAtMs, + nestedTurn.workCompletedAtMs, + source.completedAtMs, + ); + if (started !== null && completed !== null) return Math.max(0, completed - started); + return finiteNumber(fallback); + } + + function formatDuration(value) { + const duration = finiteNumber(value); + if (duration === null || duration < 0) return ""; + if (uiLocale() === "en-US") { + if (duration < 1_000) return `${Math.round(duration)}ms`; + const totalSeconds = Math.floor(duration / 1_000); + const minutes = Math.floor(totalSeconds / 60); + const seconds = totalSeconds % 60; + if (!minutes) return `${seconds}s`; + return seconds ? `${minutes}m ${seconds}s` : `${minutes}m`; + } + if (duration < 1_000) return `${Math.round(duration)}毫秒`; + const totalSeconds = Math.floor(duration / 1_000); + const minutes = Math.floor(totalSeconds / 60); + const seconds = totalSeconds % 60; + if (!minutes) return `${seconds}秒`; + // The Chinese locale in the official panel uses the long form for one + // minute and the compact form for longer durations. + const minuteLabel = minutes === 1 ? "1分钟" : `${minutes}分`; + return seconds ? `${minuteLabel}${seconds}秒` : minuteLabel; + } + + // Working indicators in the official panel stay textual until a full + // second has elapsed. This avoids the visually noisy "0ms"/"1ms" state on + // every newly-created activity row while preserving precise durations once + // an operation is complete. + function elapsedDuration(value) { + const duration = finiteNumber(value); + return duration !== null && duration >= 1_000 ? formatDuration(duration) : ""; + } + + function formatMessageTime(value) { + const timestamp = timestampMs(value); + if (timestamp === null) return ""; + try { + return new Intl.DateTimeFormat(uiLocale(), { hour: "2-digit", minute: "2-digit" }).format(new Date(timestamp)); + } catch { + return ""; + } + } + + function formatMessageDate(value, nowValue = Date.now()) { + const timestamp = timestampMs(value); + if (timestamp === null) return ""; + try { + const date = new Date(timestamp); + const now = new Date(nowValue); + // Match the official timestamp separator: compare calendar dates while + // avoiding DST changes in the local timezone. + const dateDay = Date.UTC(date.getFullYear(), date.getMonth(), date.getDate()); + const nowDay = Date.UTC(now.getFullYear(), now.getMonth(), now.getDate()); + const dayDifference = Math.max(0, Math.round((nowDay - dateDay) / 86_400_000)); + const locale = uiLocale(); + const time = new Intl.DateTimeFormat(locale, { hour: "numeric", minute: "2-digit" }).format(date); + if (dayDifference <= 1) { + try { + const relative = new Intl.RelativeTimeFormat(locale, { numeric: "auto" }).format(-Math.max(dayDifference, 0), "day"); + return `${relative} ${time}`; + } catch { + return `${t(dayDifference === 1 ? "昨天" : "今天")} ${time}`; + } + } + if (dayDifference <= 7 && dayDifference > 0) { + const weekday = new Intl.DateTimeFormat(locale, { weekday: "long" }).format(date); + return `${weekday} ${time}`; + } + const datePart = dayDifference <= 365 + ? new Intl.DateTimeFormat(locale, { month: "short", day: "numeric", weekday: "short" }).format(date) + : new Intl.DateTimeFormat(locale, { year: "numeric", month: "short", day: "numeric" }).format(date); + return `${datePart} ${time}`; + } catch { + return ""; + } + } + + function messageDateKey(value) { + const timestamp = timestampMs(value); + if (timestamp === null) return ""; + const date = new Date(timestamp); + return `${date.getFullYear()}-${date.getMonth()}-${date.getDate()}`; + } + + function normalizeActivityStatus(value, fallback = "inProgress") { + const normalized = String(value ?? fallback).replace(/[\s_-]+/g, "").toLowerCase(); + if (["inprogress", "running", "started", "active", "pending"].includes(normalized)) return "inProgress"; + if (["completed", "complete", "success", "succeeded", "done"].includes(normalized)) return "completed"; + if (["failed", "failure", "error"].includes(normalized)) return "failed"; + if (["declined", "denied", "rejected"].includes(normalized)) return "declined"; + if (["interrupted", "cancelled", "canceled", "aborted"].includes(normalized)) return "interrupted"; + return fallback; + } + + function activityStatusLabel(status) { + if (status === "inProgress") return "进行中"; + if (status === "completed") return "已完成"; + if (status === "failed") return "失败"; + if (status === "declined") return "已拒绝"; + if (status === "interrupted") return "已中断"; + return String(status || ""); + } + + const isRunningActivity = (activity) => activity && activity.status === "inProgress"; + + function eventParams(payload) { + if (!isRecord(payload)) return {}; + return isRecord(payload.params) ? payload.params : payload; + } + + function eventThreadId(payload) { + const params = eventParams(payload); + return payload?.threadId || params.threadId || params.thread?.id || params.turn?.threadId || ""; + } + + function eventTurnId(payload) { + const params = eventParams(payload); + return payload?.turnId || params.turnId || params.turn?.id || params.item?.turnId || ""; + } + + function itemFromPayload(payload) { + const params = eventParams(payload); + if (isRecord(params.item)) return params.item; + if (isRecord(payload?.item)) return payload.item; + return isRecord(params) && (params.type || params.kind) ? params : {}; + } + + function normalizedItemType(item) { + return String(item?.type || item?.kind || "").replace(/[\s/_.-]+/g, "").toLowerCase(); + } + + function activityKindForItem(item) { + const type = normalizedItemType(item); + if (isReadActivity(item)) return "read"; + if (!type) return ""; + if (type.includes("subagent") || type.includes("collabagent")) return "subagent"; + if (type.includes("filechange") || type.includes("patch") || type.includes("edit")) return "edit"; + if (type.includes("reasoning") || type.includes("contextcompaction")) return "reasoning"; + if (type.includes("plan") || type.includes("reviewmode")) return "plan"; + if (type.includes("command") || type.includes("exec") || type.includes("process")) return "tool"; + if (type.includes("tool") || type.includes("mcp") || type.includes("websearch") || type.includes("imageview")) return "tool"; + if (type.includes("usermessage")) return "user"; + if (type.includes("agentmessage") || type.includes("assistantmessage")) return "assistant"; + return ""; + } + + function activityLabelForItem(item, kind) { + const type = normalizedItemType(item); + if (kind === "subagent" || type.includes("subagent") || type.includes("collabagent")) { + return firstString(item.displayName, item.agentPath, item.action) || "子代理"; + } + if (type.includes("contextcompaction")) return "整理上下文"; + if (type.includes("websearch")) return "搜索"; + if (type.includes("imageview")) return "查看图像"; + if (type.includes("mcp")) return "MCP 工具"; + if (type.includes("dynamictool")) return "工具"; + if (kind === "edit") return "编辑文件"; + if (kind === "read") return "读取文件"; + if (kind === "plan") return "计划"; + if (kind === "reasoning") return "思考"; + if (kind === "commentary") return "工作说明"; + if (kind === "tool") return "运行命令"; + return "执行步骤"; + } + + function historyActivitySummary(item, kind, status, duration) { + const elapsed = duration ? ` · ${duration}` : ""; + if (kind === "subagent") { + const name = firstString(item.displayName, item.agentPath, item.agentThreadId ? `thread ${item.agentThreadId}` : "") || t("子代理"); + const uiStatus = firstString(item.displayStatus, item.activityKind); + if (status === "inProgress" || uiStatus === "active" || uiStatus === "updated") return `${name} ${t("已开始工作")}`; + if (status === "failed") return `${name} ${t("失败")}`; + if (status === "interrupted") return `${name} ${t("已中断")}`; + return `${name} ${t("已完成")}`; + } + if (kind === "tool" || kind === "read") { + const command = terminalCommandText(commandText(item)); + const actionSource = String(item.label || "工具"); + const action = t(actionSource); + const path = firstString(...readPathList(item)); + if (kind === "read" && status === "inProgress") return readSummaryLabel(item, status); + if (kind === "read" && !command) { + return readSummaryLabel(item, status); + } + if (kind === "read" && status !== "inProgress") { + // Parsed read/search actions are coalesced by the official renderer + // into one compact exploration row rather than a path plus a second + // timed command row. + if (command) return status === "failed" ? t("读取文件运行命令失败") : status === "interrupted" ? t("已停止读取文件运行命令") : t("已读取文件运行了命令"); + return readSummaryLabel(item, status); + } + if (status === "inProgress") return command + ? uiWithRaw("正在运行 ", "Running ", command) + : uiLocale() === "en-US" ? `Running ${action}` : `正在${action}`; + if (status === "failed") return command + ? `${uiWithRaw("命令运行失败 · ", "Command failed · ", command)}${elapsed}` + : `${action}${uiText("失败", " failed")}${elapsed}`; + if (status === "interrupted") return command + ? `${uiWithRaw("已停止 ", "Stopped ", command)}${elapsed}` + : `${uiText("已停止", "Stopped ")}${action}${elapsed}`; + if (command) return duration + ? `${uiWithRaw("已在 ", "Ran ", command, " 内运行 ", " in ")}${duration}` + : uiWithRaw("已运行 ", "Ran ", command); + return uiLocale() === "en-US" + ? `Ran ${action}${elapsed}` + : `已${action}${elapsed}`; + } + if (kind === "edit") return status === "inProgress" ? t("正在编辑文件") : `${t("编辑了文件")}${elapsed}`; + if (kind === "reasoning") return status === "inProgress" ? t("正在思考") : duration ? uiWithRaw("已思考 ", "Thought for ", duration) : t("已完成思考"); + if (kind === "plan") return status === "inProgress" ? t("正在制定计划") : `${t("已完成计划")}${elapsed}`; + if (kind === "commentary") { + if (status === "inProgress") return t("正在处理"); + return t("工作说明"); + } + return ""; + } + + function textFromValue(value, depth = 0) { + if (typeof value === "string") return value; + if (typeof value === "number" || typeof value === "boolean") return String(value); + if (Array.isArray(value)) { + const parts = value.map((entry) => textFromValue(entry, depth + 1)).filter(Boolean); + return parts.length ? parts.join("\n") : ""; + } + if (!isRecord(value) || depth > 3) return ""; + for (const key of ["text", "value", "output", "stdout", "stderr", "delta", "summary", "message"]) { + const text = textFromValue(value[key], depth + 1); + if (text) return text; + } + return ""; + } + + function displayValue(value) { + const text = textFromValue(value); + if (text) return text; + if (value === undefined || value === null) return ""; + try { return JSON.stringify(value, null, 2); } catch { return String(value); } + } + + function commandText(item) { + if (Array.isArray(item.command)) return item.command.map(String).join(" "); + if (typeof item.command === "string") return item.command; + if (typeof item.commandLine === "string") return item.commandLine; + if (Array.isArray(item.commandActions)) { + return item.commandActions + .map((action) => isRecord(action) ? action.command || action.description || "" : "") + .filter(Boolean) + .join("\n"); + } + return ""; + } + + function fileChangesText(changes) { + if (!Array.isArray(changes)) return displayValue(changes); + return changes.map((change) => { + if (!isRecord(change)) return displayValue(change); + const path = change.path || change.file || change.filePath || change.name || "文件"; + const kind = change.kind || change.type || change.status || ""; + const diff = displayValue(change.diff || change.patch || change.output || change.text); + const heading = `${kind ? `[${kind}] ` : ""}${path}`; + return diff ? `${heading}\n${diff}` : heading; + }).filter(Boolean).join("\n\n"); + } + + function planText(value) { + const plan = Array.isArray(value) ? value : value?.plan || value?.steps; + if (!Array.isArray(plan)) return displayValue(value?.text || value); + return plan.map((entry) => { + if (!isRecord(entry)) return `- ${displayValue(entry)}`; + const status = normalizeActivityStatus(entry.status, "pending"); + const marker = status === "completed" ? "[x]" : status === "inProgress" ? "[~]" : "[ ]"; + const step = entry.step || entry.text || entry.title || entry.description || "步骤"; + return `${marker} ${step}`; + }).join("\n"); + } + + function activityHeader(item, kind) { + // Terminal activities render command/cwd as separate fields below. Keep + // this helper for non-terminal activity kinds and future item types. + return ""; + } + + function activityOutput(item, kind) { + if (kind === "reasoning") return displayValue(item.summary || item.content || item.text); + if (kind === "plan") return planText(item.plan || item.steps || item.content || item.text); + if (kind === "edit") return fileChangesText(item.changes || item.files || item.diff || item.patch || item.output || item.text); + if (kind === "tool" || kind === "read") { + if (kind === "read") { + // Paths are rendered as their own compact exploration rows. Keep only + // the actual file content here; adapter summaries such as + // "已读取 ..." must not be repeated inside a Shell block. + let value = displayValue(item.aggregatedOutput ?? item.output ?? item.stdout ?? item.stderr ?? item.result); + const summary = firstString(item.text, item.summary); + if (summary && value === summary) return ""; + if (summary && value.startsWith(`${summary}\n`)) value = value.slice(summary.length + 1); + return value; + } + return displayValue(item.aggregatedOutput ?? item.output ?? item.stdout ?? item.stderr ?? item.result ?? item.error ?? item.text); + } + return displayValue(item.text || item.content || item.output); + } + + function activityKey(payload, item, kind, explicitKey) { + if (explicitKey) return explicitKey; + const params = eventParams(payload); + const id = item.id ?? params.itemId ?? payload?.itemId; + const threadId = eventThreadId(payload) || state.threadId || "thread"; + const turnId = eventTurnId(payload) || state.turnId || "turn"; + if (id !== undefined && id !== null) return `${threadId}:${turnId}:${kind}:${typeof id}:${String(id)}`; + state.activitySequence += 1; + return `${threadId}:${turnId}:${kind}:anonymous:${state.activitySequence}`; + } + + function ensureActivity(key, config = {}) { + const kind = config.kind || "reasoning"; + let existing = state.activities.get(key); + if (!existing) { + const requestedItemId = config.itemId === undefined || config.itemId === null ? "" : String(config.itemId); + const sameTurn = (entry) => entry.kind === kind + && (!config.turnId || !entry.turnId || entry.turnId === String(config.turnId)); + const entries = [...state.activities.entries()].reverse(); + const exact = requestedItemId + ? entries.find(([, entry]) => sameTurn(entry) && entry.itemId === requestedItemId) + : null; + // Status snapshots may arrive before the item lifecycle notification. + // When the concrete item appears, adopt the still-running anonymous row + // instead of appending a second "正在思考/读取/编辑" entry. + const anonymous = entries.find(([, entry]) => sameTurn(entry) + && entry.anonymous + && isRunningActivity(entry)); + const unkeyed = !requestedItemId + ? entries.find(([, entry]) => sameTurn(entry) && !entry.itemId && isRunningActivity(entry)) + : null; + const match = exact || anonymous || unkeyed; + if (match) { + const [oldKey, candidate] = match; + existing = candidate; + if (oldKey !== key) { + state.activities.delete(oldKey); + existing.key = key; + existing.article.dataset.activityKey = key; + if (state.liveActivityKey === oldKey) state.liveActivityKey = key; + if (state.activeAssistantActivityKey === oldKey) state.activeAssistantActivityKey = key; + } + if (requestedItemId && existing.anonymous) { + existing.itemId = requestedItemId; + existing.anonymous = false; + } + state.activities.set(key, existing); + } + } + if (existing) { + // A lifecycle event or real output upgrades a transient status row into + // a concrete activity. Once upgraded, it must survive turn completion; + // status-only rows are retired when the host reports a terminal state. + if (config.concrete === true) { + existing.concrete = true; + existing.statusOnly = false; + // Keep the anonymous marker for id-less lifecycle streams so the + // terminal status path can still close them when the host omits an + // explicit item/completed event. An identified item is authoritative. + if (config.itemId !== undefined && config.itemId !== null && String(config.itemId)) { + existing.anonymous = false; + } + } else if (config.statusOnly === true && existing.concrete !== true) { + existing.statusOnly = true; + } + if (config.label) existing.label = config.label; + if (config.itemId !== undefined) existing.itemId = String(config.itemId); + if (config.command) existing.command = kind === "tool" || kind === "read" ? terminalCommandText(config.command) : String(config.command); + if (config.filePath !== undefined) existing.filePath = String(config.filePath || ""); + if (config.cwd !== undefined) existing.cwd = String(config.cwd || ""); + if (config.shellName) existing.shellName = String(config.shellName); + if (config.agentThreadId !== undefined) existing.agentThreadId = String(config.agentThreadId || ""); + if (config.displayName !== undefined) existing.displayName = String(config.displayName || ""); + if (config.objective !== undefined) existing.objective = String(config.objective || ""); + if (config.activityKind !== undefined) existing.activityKind = String(config.activityKind || ""); + if (config.displayStatus !== undefined) existing.displayStatus = String(config.displayStatus || ""); + if (config.model !== undefined) existing.model = String(config.model || ""); + if (config.action !== undefined) existing.action = String(config.action || ""); + if (config.prompt !== undefined) existing.prompt = String(config.prompt || ""); + if (config.senderThreadId !== undefined) existing.senderThreadId = String(config.senderThreadId || ""); + if (config.receiverThreadIds !== undefined) existing.receiverThreadIds = Array.isArray(config.receiverThreadIds) ? config.receiverThreadIds.map(String) : []; + if (config.agentsStates !== undefined) existing.agentsStates = isRecord(config.agentsStates) ? config.agentsStates : {}; + if (config.canInteract !== undefined) existing.canInteract = config.canInteract !== false; + if (config.exitCode !== undefined) existing.exitCode = finiteNumber(config.exitCode); + if (config.turnId) existing.turnId = String(config.turnId); + if (config.startedAt) existing.startedAt = timestampMs(config.startedAt) || existing.startedAt; + if (existing.agentThreadId) existing.article.dataset.agentThreadId = existing.agentThreadId; + else delete existing.article.dataset.agentThreadId; + if (kind === "tool" || kind === "read" || kind === "subagent" || kind === "commentary") renderActivityText(existing); + refreshActivity(existing); + if (existing.turnId && state.turnStartedAt !== null && ["active", "waiting"].includes(state.turnStatus)) { + ensureLiveTurnDivider(existing.turnId); + } + return existing; + } + const role = kind === "commentary" + ? "assistant" + : kind === "tool" || kind === "read" || kind === "edit" || kind === "subagent" ? "tool" : "system"; + const messageKind = kind === "tool" || kind === "read" ? "tool" : kind === "reasoning" ? "reasoning" : kind === "plan" ? "plan" : kind; + const message = appendMessage("", role, "activity", "", { + kind: messageKind, + label: config.label || "执行步骤", + command: config.command ? (kind === "tool" || kind === "read" ? terminalCommandText(config.command) : String(config.command)) : "", + turnId: config.turnId || "", + agentThreadId: config.agentThreadId || "", + // Commentary is an assistant paragraph in the official transcript, not + // a nested disclosure. The outer worked-for group still owns its layout. + collapsible: kind !== "commentary", + showActions: false, + // Keep commentary visible while a turn is open; command/read/reasoning + // bodies remain independently collapsible. + open: config.open === true || kind === "commentary", + }); + if (!message) return null; + const activity = { + key, + kind, + role, + messageKind, + label: config.label || "执行步骤", + command: config.command ? (kind === "tool" || kind === "read" ? terminalCommandText(config.command) : String(config.command)) : "", + cwd: config.cwd ? String(config.cwd) : "", + shellName: config.shellName ? String(config.shellName) : "Shell", + agentThreadId: config.agentThreadId ? String(config.agentThreadId) : "", + displayName: config.displayName ? String(config.displayName) : "", + objective: config.objective ? String(config.objective) : "", + activityKind: config.activityKind ? String(config.activityKind) : "", + displayStatus: config.displayStatus ? String(config.displayStatus) : "", + model: config.model ? String(config.model) : "", + action: config.action ? String(config.action) : "", + filePath: config.filePath ? String(config.filePath) : "", + prompt: config.prompt ? String(config.prompt) : "", + senderThreadId: config.senderThreadId ? String(config.senderThreadId) : "", + receiverThreadIds: Array.isArray(config.receiverThreadIds) ? config.receiverThreadIds.map(String) : [], + agentsStates: isRecord(config.agentsStates) ? config.agentsStates : {}, + canInteract: config.canInteract !== false, + exitCode: config.exitCode === undefined ? null : finiteNumber(config.exitCode), + itemId: config.itemId === undefined ? "" : String(config.itemId), + threadId: config.threadId || state.threadId || "", + turnId: config.turnId || state.turnId || "", + startedAt: timestampMs(config.startedAt) || Date.now(), + finishedAt: null, + durationMs: null, + durationExplicit: false, + status: normalizeActivityStatus(config.status, "inProgress"), + headerText: "", + outputText: "", + anonymous: Boolean(config.anonymous), + concrete: config.concrete === true, + statusOnly: config.statusOnly === true, + article: message.article, + body: message.content, + wrapper: message.wrapper, + details: message.article.querySelector("details"), + summary: message.article.querySelector("summary"), + }; + activity.article.dataset.activityKey = key; + activity.article.dataset.activityKind = kind; + if (activity.agentThreadId) activity.article.dataset.agentThreadId = activity.agentThreadId; + state.activities.set(key, activity); + while (state.activities.size > 500) { + const removable = [...state.activities].find(([, entry]) => !isRunningActivity(entry)); + if (!removable) break; + state.activities.delete(removable[0]); + } + renderActivityText(activity); + refreshActivity(activity); + if (activity.turnId && state.turnStartedAt !== null && ["active", "waiting"].includes(state.turnStatus)) { + ensureLiveTurnDivider(activity.turnId); + } + ensureActivityTimer(); + return activity; + } + + function activityText(activity) { + if (activity.kind === "tool" || activity.kind === "read") { + const command = terminalCommandText(activity.command); + const commandLine = command ? `$ ${command}` : ""; + return [commandLine, activity.outputText].filter(Boolean).join("\n"); + } + if (activity.kind === "subagent") { + const name = activity.displayName || activity.label || t("子代理"); + const objective = activity.objective || activity.outputText; + return [name, objective].filter(Boolean).join("\n"); + } + if (activity.headerText && activity.outputText) return `${activity.headerText}\n\n${activity.outputText}`; + return activity.headerText || activity.outputText || ""; + } + + function normalizeSubagentActionStatus(value) { + const normalized = String(value || "").replace(/[\s_-]+/g, "").toLowerCase(); + if (["pendinginit", "pending", "waiting"].includes(normalized)) return "waiting"; + if (["running", "working", "active", "started", "interacted", "updated", "inprogress"].includes(normalized)) return "working"; + if (["completed", "complete", "done", "interrupted", "shutdown"].includes(normalized)) return "done"; + if (["errored", "error", "failed", "notfound"].includes(normalized)) return "failed"; + return "waiting"; + } + + function renderSubagentBody(activity) { + const body = activity?.body; + if (!body) return; + body.replaceChildren(); + body.classList.remove("terminal-body", "diff-body"); + body.classList.add("subagent-body"); + + const prompt = firstString(activity.prompt, activity.objective, activity.outputText); + if (prompt) { + const promptNode = document.createElement("div"); + promptNode.className = "subagent-prompt markdown-body"; + renderMarkdown(promptNode, prompt); + body.append(promptNode); + } + if (activity.model || activity.action) { + const meta = document.createElement("div"); + meta.className = "subagent-action-meta"; + if (activity.action) { + const action = document.createElement("span"); + action.textContent = activity.action === "spawnAgent" ? t("启动子代理") + : activity.action === "sendInput" ? t("发送输入") + : activity.action === "resumeAgent" ? t("恢复子代理") + : activity.action === "closeAgent" ? t("关闭子代理") : activity.action; + meta.append(action); + } + if (activity.model) { + const model = document.createElement("span"); + model.className = "subagent-model"; + model.textContent = activity.model; + meta.append(model); + } + body.append(meta); + } + + const states = isRecord(activity.agentsStates) ? activity.agentsStates : {}; + const receiverIds = Array.isArray(activity.receiverThreadIds) ? activity.receiverThreadIds : []; + const ids = [...new Set([...receiverIds, ...Object.keys(states)])].filter(Boolean); + if (!ids.length) return; + const rows = document.createElement("div"); + rows.className = "subagent-action-rows"; + for (const threadId of ids) { + const raw = isRecord(states[threadId]) ? states[threadId] : {}; + const status = normalizeSubagentActionStatus(raw.status); + const row = document.createElement("div"); + row.className = "subagent-action-row"; + row.dataset.status = status; + const icon = document.createElement("span"); + icon.className = "subagent-action-icon"; + icon.setAttribute("aria-hidden", "true"); + const label = document.createElement("span"); + label.className = "subagent-action-label"; + label.textContent = threadId === activity.agentThreadId + ? firstString(activity.displayName, threadId) + : `thread ${threadId}`; + const statusNode = document.createElement("span"); + statusNode.className = "subagent-action-status"; + statusNode.textContent = subagentStatusLabel(status); + row.append(icon, label, statusNode); + const message = firstString(raw.message, raw.statusMessage); + if (message) { + const note = document.createElement("div"); + note.className = "subagent-action-note"; + note.textContent = message; + row.append(note); + } + rows.append(row); + } + body.append(rows); + } + + function terminalOutputText(activity) { + const value = String(activity?.outputText || ""); + // A few older bridge payloads included an exit-code suffix in the output + // string. Strip only that exact synthetic line; real command output stays + // untouched. + return normalizeTerminalOutput(value.replace(/\n?exit code:\s*-?\d+\s*$/i, "")); + } + + function normalizeTerminalOutput(value) { + // Commands often use carriage returns/backspaces for progress updates and + // ANSI SGR sequences for color. Resolve the control characters before the + // lightweight renderer turns the result into safe DOM nodes. + const stripped = String(value || "") + .replace(/\x1b\][^\x07]*(?:\x07|\x1b\\)/g, "") + .replace(/\x1b\[(?![0-9;]*m)[0-?]*[ -/]*[@-~]/g, ""); + return stripped.replace(/\r\n/g, "\n").split("\n").map((line) => { + const cells = []; + let cursor = 0; + for (const character of line) { + if (character === "\r") { cursor = 0; continue; } + if (character === "\b") { cursor = Math.max(0, cursor - 1); continue; } + cells[cursor] = character; + cursor += 1; + } + return cells.join(""); + }).join("\n"); + } + + function renderTerminalOutput(parent, value) { + parent.replaceChildren(); + const text = String(value || ""); + const sgr = /\x1b\[([0-9;]*)m/g; + let cursor = 0; + const style = { fg: "", bg: "", bold: false, dim: false, italic: false, underline: false, strike: false }; + const appendSegment = (segment) => { + if (!segment) return; + const classes = []; + if (style.fg) classes.push(`ansi-${style.fg}-fg`); + if (style.bg) classes.push(`ansi-${style.bg}-bg`); + if (style.bold) classes.push("ansi-bold"); + if (style.dim) classes.push("ansi-dim"); + if (style.italic) classes.push("ansi-italic"); + if (style.underline) classes.push("ansi-underline"); + if (style.strike) classes.push("ansi-strikethrough"); + if (!classes.length) parent.append(document.createTextNode(segment)); + else { + const span = document.createElement("span"); + span.className = classes.join(" "); + span.textContent = segment; + parent.append(span); + } + }; + const applySgr = (codes) => { + const values = codes.length ? codes : [0]; + for (const code of values) { + if (code === 0) Object.assign(style, { fg: "", bg: "", bold: false, dim: false, italic: false, underline: false, strike: false }); + else if (code === 1) style.bold = true; + else if (code === 2) style.dim = true; + else if (code === 3) style.italic = true; + else if (code === 4) style.underline = true; + else if (code === 9) style.strike = true; + else if (code === 22) { style.bold = false; style.dim = false; } + else if (code === 23) style.italic = false; + else if (code === 24) style.underline = false; + else if (code === 29) style.strike = false; + else if (code === 39) style.fg = ""; + else if (code === 49) style.bg = ""; + else if (code >= 30 && code <= 37) style.fg = ["black", "red", "green", "yellow", "blue", "magenta", "cyan", "white"][code - 30]; + else if (code >= 90 && code <= 97) style.fg = ["bright-black", "bright-red", "bright-green", "bright-yellow", "bright-blue", "bright-magenta", "bright-cyan", "bright-white"][code - 90]; + else if (code >= 40 && code <= 47) style.bg = ["black", "red", "green", "yellow", "blue", "magenta", "cyan", "white"][code - 40]; + else if (code >= 100 && code <= 107) style.bg = ["bright-black", "bright-red", "bright-green", "bright-yellow", "bright-blue", "bright-magenta", "bright-cyan", "bright-white"][code - 100]; + } + }; + for (const match of text.matchAll(sgr)) { + appendSegment(text.slice(cursor, match.index)); + applySgr(match[1] ? match[1].split(";").map((part) => Number(part) || 0) : [0]); + cursor = match.index + match[0].length; + } + appendSegment(text.slice(cursor)); + } + + function updateTerminalOutputFade(output) { + if (!output) return; + const overflow = output.scrollHeight - output.clientHeight; + output.dataset.fadeTop = String(output.scrollTop > 1); + output.dataset.fadeBottom = String(overflow - output.scrollTop > 1); + } + + function terminalStatusText(activity) { + if (!activity) return ""; + // File-read rows do not have a process exit code. The official renderer + // closes them with the read summary, not a misleading "unknown" status. + if (activity.kind === "read" && (activity.exitCode === null || activity.exitCode === undefined)) return ""; + if (activity.status === "inProgress") return ""; + if (activity.status === "interrupted") return t("已停止"); + if (activity.status === "failed" || activity.status === "declined") { + return activity.exitCode === null || activity.exitCode === undefined + ? t("退出码 未知") + : t(`退出码 ${activity.exitCode}`); + } + if (activity.status === "completed") { + if (activity.exitCode === 0) return t("成功"); + if (activity.exitCode !== null && activity.exitCode !== undefined) return t(`退出码 ${activity.exitCode}`); + return t("退出码 未知"); + } + return ""; + } + + function appendTerminalCheckIcon(parent) { + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.className.baseVal = "terminal-status-icon"; + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", "m3.5 8.2 2.7 2.7 6.3-6.3"); + svg.append(path); + parent.append(svg); + } + + function setTerminalActionIcon(button, kind = "copy") { + if (!button) return; + button.replaceChildren(); + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + if (kind === "check") { + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", "m3.5 8.2 2.7 2.7 6.3-6.3"); + svg.append(path); + } else if (kind === "collapse" || kind === "expand") { + const path = document.createElementNS("http://www.w3.org/2000/svg", "path"); + path.setAttribute("d", kind === "collapse" ? "m5 6 3 3 3-3" : "m5 10 3-3 3 3"); + svg.append(path); + } else { + const back = document.createElementNS("http://www.w3.org/2000/svg", "path"); + back.setAttribute("d", "M5.5 5.5V4c0-.8.7-1.5 1.5-1.5h5c.8 0 1.5.7 1.5 1.5v5c0 .8-.7 1.5-1.5 1.5h-1.5"); + const front = document.createElementNS("http://www.w3.org/2000/svg", "rect"); + front.setAttribute("x", "2.5"); + front.setAttribute("y", "5.5"); + front.setAttribute("width", "7.5"); + front.setAttribute("height", "7.5"); + front.setAttribute("rx", "1.2"); + svg.append(back, front); + } + button.append(svg); + } + + function createTerminalAction(kind, title) { + const button = document.createElement("button"); + button.type = "button"; + button.className = "terminal-action"; + button.title = t(title); + button.setAttribute("aria-label", t(title)); + setTerminalActionIcon(button, kind); + return button; + } + + async function copyTerminalValue(value, button) { + try { + await navigator.clipboard?.writeText(String(value || "")); + setTerminalActionIcon(button, "check"); + button.dataset.copied = "true"; + setTimeout(() => { + if (!button.isConnected) return; + setTerminalActionIcon(button, "copy"); + button.dataset.copied = "false"; + }, 1_500); + } catch { + button.dataset.copied = "false"; + } + } + + function createReadPathIcon() { + const svg = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + svg.className.baseVal = "read-path-icon"; + svg.setAttribute("viewBox", "0 0 16 16"); + svg.setAttribute("aria-hidden", "true"); + const folder = document.createElementNS("http://www.w3.org/2000/svg", "path"); + folder.setAttribute("d", "M2.5 4.5h3l1.2 1.4h6.8v6.2a1 1 0 0 1-1 1h-9a1 1 0 0 1-1-1z"); + const top = document.createElementNS("http://www.w3.org/2000/svg", "path"); + top.setAttribute("d", "M2.5 4.5v-1h3l1.1 1h5.9"); + svg.append(folder, top); + return svg; + } + + function renderReadBody(activity) { + const body = activity?.body; + if (!body) return; + body.replaceChildren(); + body.classList.remove("terminal-body", "diff-body", "markdown-body"); + body.classList.add("read-body"); + const paths = String(activity.filePath || "").split("\n").map((value) => value.trim()).filter(Boolean); + if (paths.length) { + const list = document.createElement("div"); + list.className = "read-path-list"; + for (const path of paths) { + const row = document.createElement("div"); + row.className = "read-path-row"; + row.append(createReadPathIcon()); + const value = document.createElement("span"); + value.textContent = path; + value.title = path; + row.append(value); + list.append(row); + } + body.append(list); + } + const output = terminalOutputText(activity); + if (output) { + const outputBlock = document.createElement("pre"); + outputBlock.className = "read-output"; + outputBlock.textContent = output; + body.append(outputBlock); + } + if (!paths.length && !output) { + const empty = document.createElement("span"); + empty.className = "read-empty"; + empty.textContent = t(activity.status === "inProgress" ? "正在读取文件" : "读取完成"); + body.append(empty); + } + } + + function renderTerminalBody(activity) { + const body = activity?.body; + if (!body) return; + if (activity.kind === "read" && !terminalCommandText(activity.command)) { + renderReadBody(activity); + return; + } + const previousOutput = body.querySelector(".terminal-output"); + const previousScrollTop = previousOutput?.scrollTop || 0; + const previousScrollLeft = previousOutput?.scrollLeft || 0; + const previousAtBottom = previousOutput + ? previousOutput.scrollHeight - previousOutput.scrollTop - previousOutput.clientHeight <= 2 + : true; + body.replaceChildren(); + body.classList.add("terminal-body"); + body.classList.remove("markdown-body", "diff-body"); + + const shell = document.createElement("div"); + shell.className = "terminal-shell"; + const command = terminalCommandText(activity.command); + const output = terminalOutputText(activity); + + const shellHeader = document.createElement("div"); + shellHeader.className = "terminal-shell-header"; + const shellLabel = document.createElement("div"); + shellLabel.className = "terminal-shell-label"; + shellLabel.textContent = activity.shellName || "Shell"; + if (activity.cwd) shellLabel.title = `cwd\\n${activity.cwd}`; + // The official embedded shell uses a lightweight label row. The parent + // activity disclosure owns collapse; command and output each expose their + // own copy action on hover. + shellHeader.append(shellLabel); + shell.append(shellHeader); + + if (activity.kind === "read" && activity.filePath && !command) { + const pathRow = document.createElement("div"); + pathRow.className = "terminal-file-path"; + const paths = activity.filePath.split("\n").filter(Boolean); + pathRow.textContent = paths.length > 1 ? `${paths[0]} (+${paths.length - 1})` : paths[0]; + pathRow.title = activity.filePath; + shell.append(pathRow); + } + + if (command) { + const commandRow = document.createElement("div"); + commandRow.className = "terminal-command-line"; + const prompt = document.createElement("span"); + prompt.className = "terminal-prompt"; + prompt.textContent = "$"; + const commandCode = document.createElement("code"); + commandCode.textContent = command; + const commandChevron = document.createElementNS("http://www.w3.org/2000/svg", "svg"); + commandChevron.classList.add("terminal-command-chevron"); + commandChevron.setAttribute("viewBox", "0 0 16 16"); + commandChevron.setAttribute("aria-hidden", "true"); + const commandChevronPath = document.createElementNS("http://www.w3.org/2000/svg", "path"); + commandChevronPath.setAttribute("d", "m4.5 6 3.5 3.5L11.5 6"); + commandChevron.append(commandChevronPath); + const commandExpanded = activity.commandExpanded === true + || state.commandDisclosure.get(activity.key) === true; + commandRow.dataset.expanded = String(commandExpanded); + commandRow.setAttribute("role", "button"); + commandRow.setAttribute("tabindex", "0"); + commandRow.setAttribute("aria-expanded", String(commandExpanded)); + commandRow.setAttribute("aria-label", `$ ${command}`); + const toggleCommand = (event) => { + if (event.target.closest(".terminal-action")) return; + event.preventDefault(); + const expanded = commandRow.dataset.expanded === "true"; + const next = !expanded; + activity.commandExpanded = next; + if (activity.key) state.commandDisclosure.set(activity.key, next); + preserveTimelineAnchor(commandRow, 210); + animateCommandRow(commandRow, next); + if (next) scheduleTimelineReveal(commandRow, 190); + }; + commandRow.addEventListener("click", toggleCommand); + commandRow.addEventListener("keydown", (event) => { + if (event.key === "Enter" || event.key === " ") toggleCommand(event); + }); + const copyCommand = createTerminalAction("copy", "复制命令"); + copyCommand.classList.add("terminal-command-action"); + copyCommand.addEventListener("click", (event) => { + event.stopPropagation(); + copyTerminalValue(command, copyCommand); + }); + commandRow.append(prompt, commandCode, commandChevron, copyCommand); + shell.append(commandRow); + } + + const outputWrap = document.createElement("div"); + outputWrap.className = "terminal-output-wrap"; + if (output) { + const outputBlock = document.createElement("pre"); + outputBlock.className = "terminal-output"; + const outputContent = document.createElement("div"); + outputContent.className = "terminal-output-content"; + renderTerminalOutput(outputContent, output); + outputBlock.append(outputContent); + // Activity deltas rebuild the lightweight DOM. Preserve the reader's + // position, while still following the tail when they were already at + // the bottom of the terminal stream. + outputBlock.scrollLeft = previousScrollLeft; + requestAnimationFrame(() => { + if (previousAtBottom) outputBlock.scrollTop = outputBlock.scrollHeight; + else outputBlock.scrollTop = previousScrollTop; + outputBlock.scrollLeft = previousScrollLeft; + updateTerminalOutputFade(outputBlock); + }); + outputBlock.addEventListener("scroll", () => updateTerminalOutputFade(outputBlock), { passive: true }); + const copyOutput = createTerminalAction("copy", "复制输出"); + copyOutput.classList.add("terminal-output-action"); + copyOutput.addEventListener("click", (event) => { + event.stopPropagation(); + copyTerminalValue(output, copyOutput); + }); + outputWrap.append(outputBlock, copyOutput); + } else if (activity.status !== "inProgress") { + const outputBlock = document.createElement("pre"); + outputBlock.className = "terminal-output terminal-output-empty"; + const outputContent = document.createElement("div"); + outputContent.className = "terminal-output-content terminal-no-output"; + outputContent.textContent = t("无输出"); + outputBlock.append(outputContent); + outputWrap.append(outputBlock); + } + shell.append(outputWrap); + + const statusText = terminalStatusText(activity); + const footer = document.createElement("div"); + footer.className = "terminal-footer"; + footer.dataset.status = activity.status; + if (statusText) { + const status = document.createElement("span"); + status.className = "terminal-status"; + if (activity.status === "completed" && activity.exitCode === 0) appendTerminalCheckIcon(status); + status.append(document.createTextNode(statusText)); + footer.append(status); + } + shell.append(footer); + body.append(shell); + } + + function renderActivityText(activity) { + const text = activityText(activity); + const visibleText = text || t(isRunningActivity(activity) ? "等待输出…" : activityStatusLabel(activity.status)); + if (activity.kind === "tool" || activity.kind === "read") renderTerminalBody(activity); + else if (activity.kind === "subagent") renderSubagentBody(activity); + else renderMessageBody(activity.body, visibleText, activity.role, "activity", activity.messageKind); + activity.article.dataset.rawText = text; + } + + function setActivityContent(activity, header, output, append = false) { + if (!activity) return; + if (header !== undefined && header !== null && String(header)) activity.headerText = String(header); + if (output !== undefined && output !== null) { + activity.outputText = append ? `${activity.outputText}${String(output)}` : String(output); + } + renderActivityText(activity); + } + + function activityElapsed(activity) { + if (activity.durationMs !== null) return Math.max(0, activity.durationMs); + if (!activity.startedAt) return null; + // A completed legacy row without an end marker has an unknown duration; + // showing the current wall clock makes old work appear to keep running. + if (!isRunningActivity(activity) && !activity.finishedAt) return null; + const end = activity.finishedAt || Date.now(); + return Math.max(0, end - activity.startedAt); + } + + function refreshActivity(activity) { + if (!activity?.summary) return; + const elapsed = activityElapsed(activity); + // The official activity rows only show an item duration when the owner + // supplied one. A timestamp interval inferred while hydrating history is + // useful for the turn clock, but should not become a misleading per-row + // wall-clock label. + const duration = activity.durationExplicit === true ? elapsedDuration(elapsed) : ""; + const label = t(activity.label || "执行步骤"); + const setSummary = (value) => setActivitySummary(activity.summary, value, activity.kind); + // Recompute the summary in the active locale. Dynamic command/path/name + // values are appended as raw strings by historyActivitySummary(). + if (activity.kind === "subagent") { + const name = activity.displayName || activity.label || t("子代理"); + const statusText = activity.status === "inProgress" ? t("已开始工作") + : activity.status === "completed" ? t("已完成") + : activity.status === "failed" ? t("失败") + : activity.status === "interrupted" ? t("已中断") : ""; + setSubagentSummary(activity.summary, name, statusText); + } else { + const localizedSummary = historyActivitySummary(activity, activity.kind, activity.status, duration); + if (localizedSummary) setSummary(localizedSummary); + else if (activity.status === "inProgress") setSummary(label); + else if (activity.status === "failed") setSummary(`${label} · ${t("失败")}`); + else if (activity.status === "interrupted") setSummary(`${label} · ${t("已中断")}`); + else setSummary(label); + } + activity.article.dataset.status = activity.status; + } + + function finishActivity(activity, status = "completed", durationMs, finishedAt) { + if (!activity) return; + // `updateLiveActivity` may create a short-lived row before the host emits + // an item lifecycle event. It is useful while the turn is live, but the + // official transcript does not retain a blank "思考/编辑/读取" item once + // the turn ends. A concrete lifecycle row (or a row that received real + // output) has already cleared `statusOnly` and follows the normal path. + if (activity.statusOnly === true && activity.concrete !== true) { + retireActivity(activity); + return; + } + activity.status = normalizeActivityStatus(status, "completed"); + activity.finishedAt = timestampMs(finishedAt) || Date.now(); + const explicitDuration = finiteNumber(durationMs); + activity.durationExplicit = explicitDuration !== null; + const inferredDuration = activity.startedAt === null + ? null + : Math.max(0, activity.finishedAt - activity.startedAt); + activity.durationMs = explicitDuration === null + ? inferredDuration + : Math.max(0, explicitDuration); + activity.article.classList.remove("streaming"); + renderActivityText(activity); + refreshActivity(activity); + if (activity.details && (activity.kind === "reasoning" || activity.kind === "plan" || activity.kind === "subagent")) { + setDetailsExpanded(activity.details, false); + } + if (state.activeAssistantActivityKey === activity.key) { + state.activeAssistantBody = null; + state.activeAssistantStream = null; + state.activeAssistantText = ""; + state.activeAssistantActivityKey = null; + } + stopActivityTimerIfIdle(); + } + + function latestRunningActivity(kind, context = {}) { + const itemId = context.itemId === undefined || context.itemId === null ? "" : String(context.itemId); + const turnId = context.turnId || state.turnId || ""; + const activities = [...state.activities.values()].reverse(); + if (itemId) { + const exact = activities.find((activity) => activity.itemId === itemId && isRunningActivity(activity)); + if (exact) return exact; + } + const acceptedKinds = kind === "reasoning" ? new Set(["reasoning", "plan"]) : new Set([kind]); + return activities.find((activity) => isRunningActivity(activity) + && acceptedKinds.has(activity.kind) + && (!turnId || !activity.turnId || activity.turnId === turnId)) || null; + } + + function appendActivityChunk(activity, text) { + if (!activity || !text) return; + // A stream chunk is evidence of a real work item, even when the host + // omitted its item id. It should not be mistaken for a status-only row at + // turn completion. + activity.concrete = true; + activity.statusOnly = false; + activity.status = "inProgress"; + setActivityContent(activity, undefined, text, true); + activity.article.classList.add("streaming"); + if ((activity.kind === "reasoning" || activity.kind === "plan" || activity.kind === "commentary") && activity.details) { + // Reasoning is the one activity the official transcript expands while + // it is streaming; completed reasoning collapses again in finishActivity. + setDetailsExpanded(activity.details, true); + } + refreshActivity(activity); + ensureActivityTimer(); + } + + function handleItemLifecycle(phase, payload) { + const params = eventParams(payload); + const item = itemFromPayload(payload); + const kind = activityKindForItem(item); + if (!kind || kind === "user" || kind === "assistant") return false; + const itemId = item.id ?? params.itemId ?? payload?.itemId; + const duration = finiteNumber(item.durationMs, params.durationMs); + const finishedAt = timestampMs(item.completedAtMs, item.finishedAtMs, params.completedAtMs, params.emittedAtMs); + const startedAt = timestampMs(item.startedAtMs, item.startedAt, params.startedAtMs) + || (duration !== null && finishedAt ? finishedAt - duration : null); + const fallbackStatus = phase === "completed" ? "completed" : "inProgress"; + const activity = ensureActivity(activityKey(payload, item, kind), { + kind, + label: activityLabelForItem(item, kind), + command: kind === "tool" || kind === "read" ? terminalCommandText(commandText(item)) : "", + filePath: kind === "read" ? readPathList(item).join("\n") : "", + cwd: (kind === "tool" || kind === "read") && typeof item.cwd === "string" ? item.cwd : "", + shellName: (kind === "tool" || kind === "read") && typeof item.shellName === "string" ? item.shellName : "Shell", + agentThreadId: firstString(item.agentThreadId, item.childThreadId, item.threadId), + displayName: firstString(item.displayName, item.agentNickname, item.agentName, item.agentPath), + objective: firstString(item.objective, item.prompt, item.statusMessage, item.message), + activityKind: firstString(item.activityKind, item.kind), + displayStatus: firstString(item.displayStatus, item.status), + model: firstString(item.model, item.modelId), + action: firstString(item.action, item.tool), + prompt: item.prompt === null ? "" : firstString(item.prompt), + senderThreadId: firstString(item.senderThreadId), + receiverThreadIds: Array.isArray(item.receiverThreadIds) + ? item.receiverThreadIds.map(String) + : Array.isArray(item.receiverThreads) ? item.receiverThreads.map(String) : [], + agentsStates: isRecord(item.agentsStates) ? item.agentsStates : {}, + canInteract: item.canInteract !== false, + exitCode: kind === "tool" || kind === "read" ? finiteNumber(item.exitCode, item.exit_code) : undefined, + itemId, + threadId: eventThreadId(payload), + turnId: eventTurnId(payload), + startedAt, + status: normalizeActivityStatus(item.status, fallbackStatus), + concrete: true, + anonymous: itemId === undefined || itemId === null, + }); + if (!activity) return true; + if (kind === "tool" || kind === "read") { + const command = terminalCommandText(commandText(item)); + if (command) activity.command = command; + if (kind === "read") activity.filePath = readPathList(item).join("\n") || activity.filePath; + if (typeof item.cwd === "string") activity.cwd = item.cwd; + if (typeof item.shellName === "string" && item.shellName) activity.shellName = item.shellName; + const exitCode = finiteNumber(item.exitCode, item.exit_code); + if (exitCode !== null) activity.exitCode = exitCode; + } + if (kind === "subagent") { + activity.agentThreadId = firstString(item.agentThreadId, item.childThreadId, item.threadId, activity.agentThreadId); + activity.displayName = firstString(item.displayName, item.agentNickname, item.agentName, item.agentPath, activity.displayName); + activity.objective = firstString(item.objective, item.prompt, item.statusMessage, item.message, activity.objective); + activity.activityKind = firstString(item.activityKind, item.kind, activity.activityKind); + activity.displayStatus = firstString(item.displayStatus, item.status, activity.displayStatus); + activity.model = firstString(item.model, item.modelId, activity.model); + activity.action = firstString(item.action, item.tool, activity.action); + activity.prompt = item.prompt === null ? "" : firstString(item.prompt, activity.prompt); + activity.senderThreadId = firstString(item.senderThreadId, activity.senderThreadId); + if (Array.isArray(item.receiverThreadIds)) activity.receiverThreadIds = item.receiverThreadIds.map(String); + else if (Array.isArray(item.receiverThreads)) activity.receiverThreadIds = item.receiverThreads.map(String); + if (isRecord(item.agentsStates)) activity.agentsStates = item.agentsStates; + if (item.canInteract !== undefined) activity.canInteract = item.canInteract !== false; + if (activity.agentThreadId) activity.article.dataset.agentThreadId = activity.agentThreadId; + else delete activity.article.dataset.agentThreadId; + } + const header = activityHeader(item, kind); + const output = activityOutput(item, kind); + if (header || output) setActivityContent(activity, header, output); + if (phase === "completed") { + const status = normalizeActivityStatus(item.status, item.error ? "failed" : "completed"); + finishActivity(activity, status, duration, finishedAt); + } else { + activity.status = normalizeActivityStatus(item.status, "inProgress"); + refreshActivity(activity); + } + return true; + } + + function handlePlanUpdate(payload) { + const params = eventParams(payload); + const steps = Array.isArray(params.plan) ? params.plan : Array.isArray(params.steps) ? params.steps : null; + const text = planText(steps || params.text || params); + if (!text) return false; + const threadId = eventThreadId(payload) || state.threadId || "thread"; + const turnId = eventTurnId(payload) || state.turnId || "turn"; + const activity = ensureActivity(`plan:${threadId}:${turnId}`, { + kind: "plan", + label: "计划", + threadId, + turnId, + status: "inProgress", + concrete: true, + }); + setActivityContent(activity, "", text); + const statuses = (steps || []).map((step) => normalizeActivityStatus(step?.status, "pending")); + if (statuses.length && statuses.every((status) => status === "completed")) finishActivity(activity, "completed"); + else refreshActivity(activity); + return true; + } + + function handleDiffUpdate(payload) { + const params = eventParams(payload); + const diff = fileChangesText(params.changes || params.diff || params.patch || params.delta || params.text || params.output); + if (!diff) return false; + const threadId = eventThreadId(payload) || state.threadId || "thread"; + const turnId = eventTurnId(payload) || state.turnId || "turn"; + const activity = ensureActivity(`diff:${threadId}:${turnId}`, { + kind: "edit", + label: "文件变更", + threadId, + turnId, + status: "inProgress", + concrete: true, + }); + setActivityContent(activity, "", diff); + refreshActivity(activity); + return true; + } + + function finishActivitiesForTurn(turnId, status) { + for (const activity of state.activities.values()) { + if (!isRunningActivity(activity)) continue; + if (turnId && activity.turnId && activity.turnId !== turnId) continue; + finishActivity(activity, status); + } + } + + function turnStatusLabel(status) { + if (status === "waiting") return t("等待授权"); + if (status === "generating") return t("正在生成"); + if (status === "interrupted") return t("已中断"); + if (status === "failed") return t("失败"); + if (status === "completed") return t("已完成"); + return t("正在工作"); + } + + function statusActivityLabel(activity, flags = []) { + const normalized = String(activity || "").replace(/[\s-]+/g, "_").toLowerCase(); + if (normalized === "waiting_approval" || normalized === "waiting_for_approval" || flags.some((flag) => /approval|permission/.test(String(flag)))) return t("等待授权"); + if (normalized === "waiting_input" || normalized === "waiting_for_user_input" || flags.some((flag) => /user.?input/.test(String(flag)))) return t("正在等待你的回答"); + if (normalized === "thinking" || normalized === "reasoning") return t("正在思考"); + if (normalized === "editing" || normalized === "edit") return t("正在编辑文件"); + if (normalized === "reading" || normalized === "reading_file" || normalized === "file_read") return t("正在读取文件"); + if (normalized === "running" || normalized === "running_command" || normalized === "tool") return t("正在运行命令"); + if (normalized === "searching" || normalized === "searching_web") return t("正在搜索网页"); + if (normalized === "responding" || normalized === "generating") return t("正在生成"); + if (normalized === "failed") return t("执行失败"); + if (normalized === "interrupted") return t("已中断"); + if (normalized === "completed") return t("已完成"); + return normalized && normalized !== "idle" ? t("处理中") : ""; + } + + function updateLiveActivity(activity, startedAt, durationMs, flags = [], turnId = "") { + const element = $("liveActivity"); + if (!element) return; + const normalized = String(activity || "idle").replace(/[\s-]+/g, "_").toLowerCase(); + const terminal = ["idle", "ready", "completed", "failed", "interrupted", "cancelled", "canceled"].includes(normalized); + state.currentActivity = normalized; + state.currentActivityStartedAt = timestampMs(startedAt) || state.currentActivityStartedAt; + state.currentActivityDurationMs = finiteNumber(durationMs); + state.currentActivityTurnId = turnId || state.currentActivityTurnId || ""; + if (terminal) { + const live = state.liveActivityKey ? state.activities.get(state.liveActivityKey) : null; + // Concrete lifecycle rows are finalized by item/turn completion events. + // A status transition may only close an anonymous streaming fallback; + // closing a concrete row here can race the owner snapshot and erase its + // command/edit label before the real completion payload arrives. + if (live?.anonymous && isRunningActivity(live)) { + finishActivity(live, normalized === "failed" ? "failed" : normalized === "interrupted" ? "interrupted" : "completed", durationMs); + } + state.liveActivityKey = null; + element.hidden = true; + element.dataset.activity = normalized; + element.dataset.active = "false"; + return; + } + const label = statusActivityLabel(normalized, flags); + if (!label) { + element.hidden = true; + return; + } + // The live status node remains available to assistive technology, while + // the visible transcript is the source of truth for work-in-progress rows. + element.hidden = false; + element.dataset.activity = normalized; + element.dataset.active = "true"; + const labelElement = element.querySelector(".activity-label"); + const elapsedElement = element.querySelector(".activity-elapsed"); + const activeTranscript = latestRunningActivity( + normalized === "thinking" || normalized === "reasoning" ? "reasoning" + : normalized === "editing" || normalized === "edit" ? "edit" + : normalized === "reading" || normalized === "reading_file" || normalized === "file_read" ? "read" + : normalized === "searching" || normalized === "searching_web" ? "tool" + : normalized === "running" || normalized === "running_command" || normalized === "tool" ? "tool" : "", + { turnId: turnId || state.turnId || "" }, + ); + let visibleLabel = label; + if (activeTranscript) { + if (activeTranscript.kind === "tool" && activeTranscript.command) { + visibleLabel = uiWithRaw("正在运行 ", "Running ", terminalCommandText(activeTranscript.command)); + } else if (activeTranscript.kind === "edit") visibleLabel = t("正在编辑文件"); + else if (activeTranscript.kind === "read") visibleLabel = t("正在读取文件"); + else if (activeTranscript.kind === "reasoning") visibleLabel = t("正在思考"); + else if (activeTranscript.label) visibleLabel = activeTranscript.label; + } + if (labelElement) labelElement.textContent = visibleLabel; + const elapsed = state.currentActivityStartedAt === null + ? state.currentActivityDurationMs + : Math.max(0, Date.now() - state.currentActivityStartedAt); + if (elapsedElement) elapsedElement.textContent = elapsedDuration(elapsed); + + // Status snapshots are projections of the current turn, not work items. + // Bind to an existing concrete lifecycle row when possible. For the few + // builds that publish a status before its item, the fallback below creates + // a transient, explicitly `statusOnly` row that is removed at completion; + // waiting approvals and ordinary status ticks never become history items. + const effectiveTurnId = turnId || state.turnId || ""; + const transcriptKind = normalized === "thinking" || normalized === "reasoning" + ? "reasoning" + : normalized === "editing" || normalized === "edit" + ? "edit" + : normalized === "reading" || normalized === "reading_file" || normalized === "file_read" + ? "read" + : normalized === "running" || normalized === "running_command" || normalized === "tool" + ? "tool" + : normalized === "searching" || normalized === "searching_web" + ? "tool" + : ""; + let transcript = transcriptKind + ? latestRunningActivity(transcriptKind, { turnId: effectiveTurnId }) + : null; + if (!transcript && effectiveTurnId && (transcriptKind === "reasoning" || transcriptKind === "edit" || transcriptKind === "read")) { + // A few official builds publish the turn activity before the first + // reasoning/diff item. Materialize one stable anonymous row so the + // reader sees "正在思考"/"正在编辑文件" immediately; a later concrete + // lifecycle item is merged into the same visual stream by kind/turn. + const key = `status:${state.threadId || "thread"}:${effectiveTurnId}:${transcriptKind}`; + transcript = ensureActivity(key, { + kind: transcriptKind, + label: transcriptKind === "reasoning" ? "思考" : transcriptKind === "edit" ? "编辑文件" : "读取文件", + threadId: state.threadId, + turnId: effectiveTurnId, + startedAt: state.currentActivityStartedAt || state.turnStartedAt, + status: "inProgress", + anonymous: true, + statusOnly: true, + open: false, + }); + } + if (transcript) { + state.liveActivityKey = transcript.key; + // Keep the concrete item label (command, tool name, or edit summary) + // supplied by its lifecycle event. The global status must not overwrite + // it with a generic "正在运行命令"/"正在思考" label. + transcript.status = "inProgress"; + if (state.currentActivityStartedAt && !transcript.startedAt) transcript.startedAt = state.currentActivityStartedAt; + refreshActivity(transcript); + if (effectiveTurnId) ensureLiveTurnDivider(effectiveTurnId); + } else { + state.liveActivityKey = null; + } + } + + function refreshTurnClock() { + if (state.turnStartedAt === null) return; + const elapsed = Math.max(0, Date.now() - state.turnStartedAt); + const workElapsed = state.turnWorkStartedAt === null + ? elapsed + : Math.max(0, Date.now() - state.turnWorkStartedAt); + const activity = state.currentActivity && state.currentActivity !== "idle" + ? statusActivityLabel(state.currentActivity) + : turnStatusLabel(state.turnStatus); + const visibleElapsed = elapsedDuration(workElapsed); + setConversationStatus(visibleElapsed ? `${activity} · ${visibleElapsed}` : activity, state.turnStatus === "waiting" ? "warning" : "active"); + const divider = state.turnDividers.get(state.turnId) + || [...$("output").querySelectorAll(".turn-divider")] + .find((entry) => entry.dataset.turnId === state.turnId); + if (divider?.dataset.status === "inProgress") { + const label = divider.querySelector(".turn-divider-label"); + if (label) label.textContent = turnDividerLabel("inProgress", workElapsed); + } + updateLiveActivity(state.currentActivity || "running", state.currentActivityStartedAt || state.turnStartedAt, null, [], state.currentActivityTurnId || state.turnId); + } + + function startTurnClock(turnId, startedAt, elapsedMs) { + const nextTurnId = turnId || state.turnId || ""; + const explicitStartedAt = timestampMs(startedAt); + const explicitElapsed = finiteNumber(elapsedMs); + if (nextTurnId && state.turnId && nextTurnId !== state.turnId) { + state.turnStartedAt = null; + state.turnWorkStartedAt = null; + state.finalAssistantStartedAt = null; + state.workedDurationMs = null; + state.currentActivityStartedAt = null; + state.currentActivityDurationMs = null; + state.currentActivity = "running"; + } + if (nextTurnId) state.turnId = nextTurnId; + if (state.turnStartedAt === null) { + state.turnStartedAt = explicitStartedAt + || (explicitElapsed !== null ? Date.now() - Math.max(0, explicitElapsed) : Date.now()); + } + state.turnStatus = "active"; + state.lastTurnDurationMs = null; + state.lastWorkedDurationMs = null; + if (state.turnWorkStartedAt === null) state.turnWorkStartedAt = state.turnStartedAt; + const output = $("output"); + if (output && [...output.querySelectorAll(".message.activity")] + .some((article) => article.dataset.turnId === nextTurnId)) ensureLiveTurnDivider(nextTurnId); + refreshTurnClock(); + ensureActivityTimer(); + } + + function stopTurnClock(status = "completed", durationMs, finishedAt, workedDurationMs) { + // Some completion envelopes do not carry an item lifecycle event and do + // not leave `liveActivityKey` pointing at the transient status row. Sweep + // those rows here as a final safety net before clearing the turn clock. + retireStatusOnlyActivities(state.turnId); + const end = timestampMs(finishedAt) || Date.now(); + const explicitDuration = finiteNumber(durationMs); + const elapsed = explicitDuration !== null + ? Math.max(0, explicitDuration) + : state.turnStartedAt === null ? null : Math.max(0, end - state.turnStartedAt); + // A hydrated terminal snapshot can arrive immediately after metadata. In + // that case the metadata duration is authoritative even though the + // terminal status envelope does not repeat it. + const authoritativeWorkedDuration = finiteNumber( + workedDurationMs, + state.workedDurationMs, + state.lastWorkedDurationMs, + ); + const worked = workedDurationFor({ + workedDurationMs: authoritativeWorkedDuration, + firstTurnWorkItemStartedAtMs: state.turnWorkStartedAt, + finalAssistantStartedAtMs: state.finalAssistantStartedAt, + completedAtMs: end, + }, state.turnWorkStartedAt === null ? elapsed : Math.max(0, end - state.turnWorkStartedAt)); + state.turnStartedAt = null; + state.turnWorkStartedAt = null; + state.turnStatus = status; + state.lastTurnDurationMs = elapsed; + state.lastWorkedDurationMs = worked; + state.workedDurationMs = worked; + const duration = elapsedDuration(worked ?? elapsed); + setConversationStatus(`${turnStatusLabel(status)}${duration ? ` · ${duration}` : ""}`, status === "completed" ? "ready" : "warning"); + state.currentActivity = status; + state.currentActivityStartedAt = null; + state.currentActivityDurationMs = worked ?? elapsed; + updateLiveActivity(status, null, worked ?? elapsed, [], state.turnId); + stopActivityTimerIfIdle(); + } + + function turnStatusFromValue(value, fallback = "active") { + const normalized = normalizeActivityStatus(value, fallback); + if (normalized === "inProgress") return "active"; + if (normalized === "declined" || normalized === "interrupted") return "interrupted"; + return normalized; + } + + function applyStatusSnapshot(payload, options = {}) { + const status = isRecord(payload?.status) ? payload.status : {}; + const metadata = isRecord(payload?.metadata) ? payload.metadata : {}; + const rawTurnStatus = payload?.turnStatus ?? status.turnStatus ?? metadata.turnStatus ?? payload?.state; + const activity = String(payload?.activity ?? status.activity ?? metadata.activity ?? "").toLowerCase(); + const flags = [ + ...(Array.isArray(payload?.activeFlags) ? payload.activeFlags : []), + ...(Array.isArray(status.activeFlags) ? status.activeFlags : []), + ...(Array.isArray(metadata.activeFlags) ? metadata.activeFlags : []), + ].map((flag) => String(flag).toLowerCase()); + const turnId = payload?.turnId || state.turnId || ""; + const projectedWorkedDuration = workedDurationFor(payload, + workedDurationFor(status, workedDurationFor(metadata, null))); + const projectedWorkStart = timestampMs( + payload?.firstTurnWorkItemStartedAtMs, + payload?.workStartedAtMs, + status.firstTurnWorkItemStartedAtMs, + status.workStartedAtMs, + metadata.firstTurnWorkItemStartedAtMs, + metadata.workStartedAtMs, + ); + const projectedFinalAssistantStart = timestampMs( + payload?.finalAssistantStartedAtMs, + status.finalAssistantStartedAtMs, + metadata.finalAssistantStartedAtMs, + ); + if (projectedWorkStart !== null) state.turnWorkStartedAt = projectedWorkStart; + if (projectedFinalAssistantStart !== null) state.finalAssistantStartedAt = projectedFinalAssistantStart; + if (projectedWorkedDuration !== null) state.workedDurationMs = projectedWorkedDuration; + const rawNormalized = String(rawTurnStatus || "").replace(/[\s-]+/g, "_").toLowerCase(); + const normalized = turnStatusFromValue(rawTurnStatus || (turnId ? "active" : "idle"), turnId ? "active" : "idle"); + const terminal = ["completed", "complete", "done", "failed", "error", "interrupted", "cancelled", "canceled"].includes(rawNormalized) + || ["completed", "failed", "interrupted"].includes(normalized); + const effectiveActivity = activity || (flags.some((flag) => /approval|permission/.test(flag)) ? "waiting_approval" : turnId ? "running" : normalized); + const explicitlyActive = !terminal && (Boolean(turnId) + || ["active", "running", "working", "inprogress", "thinking", "reasoning", "editing", "edit", "reading", "readingfile", "fileread", "searching", "responding", "generating"].includes(activity.replace(/[\s_-]+/g, "")) + || normalized === "active"); + if (explicitlyActive) { + startTurnClock( + turnId, + payload?.startedAtMs ?? status.startedAtMs ?? metadata.startedAtMs, + payload?.elapsedMs ?? status.elapsedMs ?? metadata.elapsedMs, + ); + state.turnStatus = flags.some((flag) => /approval|permission|input|waiting/.test(flag)) ? "waiting" : "active"; + state.currentActivity = effectiveActivity; + updateLiveActivity(effectiveActivity, payload?.startedAtMs ?? status.startedAtMs ?? metadata.startedAtMs ?? state.turnStartedAt, payload?.durationMs ?? status.durationMs ?? metadata.durationMs, flags, turnId); + refreshTurnClock(); + return; + } + if (options.allowTerminal === false) return; + const duration = payload?.durationMs ?? status.durationMs ?? metadata.durationMs; + if (["completed", "interrupted", "failed"].includes(normalized) || terminal) { + stopTurnClock(normalized, duration, payload?.completedAtMs ?? status.completedAtMs ?? metadata.completedAtMs, projectedWorkedDuration); + } + else if (state.turnStartedAt === null && options.showIdle !== false) setConversationStatus("ready"); + } + + function snapshotStatusProjection(snapshot = {}, appState = {}) { + const snapshotRecord = isRecord(snapshot) ? snapshot : {}; + const stateRecord = isRecord(appState) ? appState : {}; + const statusRecord = [ + snapshotRecord.executionStatus, + snapshotRecord.status, + stateRecord.executionStatus, + stateRecord.status, + ].find(isRecord) || {}; + const rawTurnStatus = firstDefined( + snapshotRecord.turnStatus, + statusRecord.turnStatus, + statusRecord.status, + stateRecord.turnStatus, + typeof stateRecord.status === "string" ? stateRecord.status : undefined, + ); + const rawActivity = firstString( + snapshotRecord.activity, + statusRecord.activity, + stateRecord.activity, + ).toLowerCase(); + const rawNormalized = String(rawTurnStatus || "").replace(/[\s-]+/g, "_").toLowerCase(); + const fallbackStatus = rawActivity === "completed" || rawActivity === "complete" || rawActivity === "done" + ? "completed" + : rawActivity === "failed" || rawActivity === "error" + ? "failed" + : rawActivity === "interrupted" || rawActivity === "cancelled" || rawActivity === "canceled" + ? "interrupted" + : snapshotRecord.turnId || stateRecord.activeTurnId ? "active" : "idle"; + const normalized = turnStatusFromValue(rawTurnStatus ?? fallbackStatus, fallbackStatus); + const terminal = ["completed", "complete", "done", "failed", "error", "interrupted", "cancelled", "canceled"].includes(rawNormalized) + || ["completed", "failed", "interrupted"].includes(normalized) + || ["completed", "failed", "interrupted"].includes(rawActivity); + const activeActivity = ["active", "running", "working", "inprogress", "thinking", "reasoning", "editing", "edit", "reading", "reading_file", "file_read", "searching", "searching_web", "responding", "generating"].includes(rawActivity.replace(/[\s-]+/g, "_")); + const explicitTurnId = snapshotRecord.turnId !== undefined + ? snapshotRecord.turnId + : stateRecord.activeTurnId; + return { + normalized, + terminal, + hasActiveTurn: Boolean(explicitTurnId) || activeActivity || normalized === "active" || normalized === "waiting", + durationMs: finiteNumber( + snapshotRecord.durationMs, + snapshotRecord.elapsedMs, + statusRecord.durationMs, + statusRecord.elapsedMs, + stateRecord.durationMs, + stateRecord.elapsedMs, + ), + workedDurationMs: workedDurationFor(snapshotRecord, + workedDurationFor(statusRecord, workedDurationFor(stateRecord, null))), + firstTurnWorkItemStartedAtMs: timestampMs( + snapshotRecord.firstTurnWorkItemStartedAtMs, + snapshotRecord.workStartedAtMs, + statusRecord.firstTurnWorkItemStartedAtMs, + statusRecord.workStartedAtMs, + stateRecord.firstTurnWorkItemStartedAtMs, + stateRecord.workStartedAtMs, + ), + finalAssistantStartedAtMs: timestampMs( + snapshotRecord.finalAssistantStartedAtMs, + statusRecord.finalAssistantStartedAtMs, + stateRecord.finalAssistantStartedAtMs, + ), + completedAtMs: timestampMs(snapshotRecord.completedAtMs, statusRecord.completedAtMs, stateRecord.completedAtMs), + }; + } + + // A late authoritative snapshot can arrive after replayed lifecycle events. + // Once it says there is no active turn, clear only the transient projection; + // hydrated history remains intact and a locally queued user message is kept. + function reconcileSnapshotTerminalState(snapshot = {}, appState = {}) { + const projection = snapshotStatusProjection(snapshot, appState); + if (!projection.terminal || projection.hasActiveTurn || state.pendingUserText) return; + const hadTransientTurn = Boolean( + state.turnId + || state.turnStartedAt !== null + || state.currentActivity === "active" + || state.currentActivity === "running" + || state.currentActivity === "working" + || state.currentActivity === "thinking" + || state.currentActivity === "editing" + || state.currentActivity === "generating" + || [...state.activities.values()].some(isRunningActivity), + ); + if (!hadTransientTurn) return; + const finishedTurnId = state.turnId; + finishAssistantStream(); + finishActivitiesForTurn(finishedTurnId, projection.normalized); + stopTurnClock(projection.normalized, projection.durationMs, projection.completedAtMs, projection.workedDurationMs); + state.turnId = ""; + state.currentActivityTurnId = ""; + state.currentActivityStartedAt = null; + state.currentActivityDurationMs = projection.durationMs; + updateIds(); + } + + function appendFileChangeChunk(payload, text) { + if (!text) return false; + const params = eventParams(payload); + const threadId = eventThreadId(payload) || state.threadId || "thread"; + const turnId = eventTurnId(payload) || state.turnId || "turn"; + const itemId = params.itemId ?? payload?.itemId; + let activity = latestRunningActivity("edit", { itemId, turnId }); + if (!activity) { + const key = itemId === undefined + ? `diff:${threadId}:${turnId}` + : `${threadId}:${turnId}:edit:${typeof itemId}:${String(itemId)}`; + activity = ensureActivity(key, { + kind: "edit", + label: "编辑文件", + itemId, + threadId, + turnId, + status: "inProgress", + concrete: true, + }); + } + appendActivityChunk(activity, text); + return true; + } + + function refreshElapsedDisplays() { + refreshTurnClock(); + for (const activity of state.activities.values()) if (isRunningActivity(activity)) refreshActivity(activity); + refreshSubagentElapsed(); + if (state.currentActivity && state.currentActivity !== "idle") { + updateLiveActivity( + state.currentActivity, + state.currentActivityStartedAt || state.turnStartedAt, + state.currentActivityDurationMs, + [], + state.currentActivityTurnId || state.turnId, + ); + } + stopActivityTimerIfIdle(); + } + + function ensureActivityTimer() { + if (state.activityTimer !== null) return; + state.activityTimer = setInterval(refreshElapsedDisplays, 250); + } + + function stopActivityTimerIfIdle() { + if (state.turnStartedAt !== null || [...state.activities.values()].some(isRunningActivity)) return; + if (state.activityTimer !== null) clearInterval(state.activityTimer); + state.activityTimer = null; + } + + function renderEmptyOutput() { + const output = $("output"); + output.replaceChildren(); + output.dataset.outputTail = ""; + state.activeAssistantBody = null; + state.activeAssistantStream = null; + state.activeAssistantText = ""; + state.activeAssistantActivityKey = null; + state.activities.clear(); + state.turnDividers.clear(); + state.liveActivityKey = null; + state.pendingUserArticle = null; + state.lastRenderedDateKey = ""; + state.lastRenderedTimestamp = null; + state.lastRenderedRole = ""; + state.hasRenderedUser = false; + state.lastDateSeparatorTimestamp = null; + state.outputDistanceFromBottom = 0; + state.structuredMessages = []; + stopActivityTimerIfIdle(); + } + + function appendOutput(text, tone) { + if (!text) return; + const visibleText = tone === "error" || tone === "meta" ? t(text) : text; + if (tone === "meta") { + setConversationStatus(String(visibleText)); + return; + } + finishAssistantStream(); + const role = tone === "error" ? "error" : tone === "meta" ? "system" : "assistant"; + appendMessage(visibleText, role, tone || "text", tone === "meta" ? t("状态") : ""); + } + + function setConversationStatus(text, tone = "ready") { + const value = t(text || ""); + const status = $("appState"); + if (status) { + status.textContent = value; + status.dataset.tone = tone; + } + const hint = $("outputHint"); + if (hint && value) hint.textContent = value; + } + + function sessionCommandMethod(value) { + return String(value || "").trim().replace(/\./g, "/").toLowerCase(); + } + + function sessionErrorMessage(value, fallback = "会话操作失败") { + const source = isRecord(value) ? value : {}; + const error = isRecord(source.error) ? source.error : {}; + const code = firstString(source.code, error.code).toLowerCase(); + const message = value instanceof Error + ? value.message + : firstString(source.message, error.message, typeof value === "string" ? value : ""); + const normalized = `${code} ${message}`.toLowerCase(); + if (/app_not_ready|app-server is not ready|waiting_for_host/.test(normalized)) return "等待 VS Code 主机连接"; + if (/host_unavailable|host_disconnected|vscode host is disconnected/.test(normalized)) return "VS Code 主机未连接"; + if (/mode_switch_pending/.test(normalized)) return "正在切换控制模式"; + if (/mode_busy|cannot switch control mode/.test(normalized)) return "当前任务或请求完成后才能切换控制模式"; + if (/session_busy|turn_active|running turn|pending request/.test(normalized)) return "当前任务结束或请求处理后才能切换"; + if (/timed out waiting for a snapshot|snapshot from vscode|找不到会话.*owner|no live vscode owner/.test(normalized)) { + return "目标会话没有返回 VS Code 快照,请先在官方 Codex 面板打开它"; + } + if (/method_not_allowed/.test(normalized)) return "当前 relay 版本不支持此会话操作,请重启 relay"; + return message || fallback; + } + + function setSessionSwitchingVisual(switching) { + const active = Boolean(switching); + const panel = document.querySelector(".chat-panel"); + if (panel) panel.dataset.sessionSwitching = String(active); + const output = $("output"); + if (output) output.setAttribute("aria-busy", String(active)); + updateIds(); + renderRequests(); + } + + function sessionSwitchTargetTitle(threadId) { + const id = String(threadId || ""); + if (!id || !Array.isArray(state.sessions)) return ""; + const entry = state.sessions.find((candidate) => sessionEntryId(candidate) === id); + return entry ? sessionEntryTitle(entry) : ""; + } + + function beginSessionSwitchContext(threadId, title = "") { + const targetThreadId = String(threadId || ""); + if (state.sessionSwitchContext) { + if (!targetThreadId || state.sessionSwitchContext.targetThreadId === targetThreadId) { + return state.sessionSwitchContext; + } + // A newer VS Code navigation can supersede an in-flight target. Retain + // the original fallback, but reset both completion gates for the new + // target so an acknowledgement/snapshot from the older route cannot + // unlock the composer. + state.sessionSwitchContext.targetThreadId = targetThreadId; + state.sessionSwitchContext.targetTitle = String(title || sessionSwitchTargetTitle(targetThreadId) || ""); + state.sessionSwitchContext.targetSnapshotReady = false; + state.sessionSwitchContext.selectedAckReady = false; + return state.sessionSwitchContext; + } + const titleNode = $("threadTitle"); + const previousThreadId = String(state.threadId || state.syncedThreadId || ""); + state.sessionSwitchContext = { + previousThreadId, + previousTitle: titleNode?.textContent || "Codex", + targetThreadId, + targetTitle: String(title || sessionSwitchTargetTitle(targetThreadId) || ""), + targetSnapshotReady: false, + selectedAckReady: false, + }; + return state.sessionSwitchContext; + } + + function finishSessionSwitchContext() { + state.sessionSwitchContext = null; + setSessionSwitchingVisual(false); + } + + function restoreSessionSwitchContext() { + const context = state.sessionSwitchContext; + if (!context) { + state.sessionSwitching = false; + state.sessionSelectCommandId = ""; + setSessionSwitchingVisual(false); + return; + } + const titleNode = $("threadTitle"); + if (titleNode && context.previousTitle) titleNode.textContent = context.previousTitle; + if (context.previousThreadId) { + state.threadId = context.previousThreadId; + state.sessionSelectedThreadId = context.previousThreadId; + } else { + state.sessionSelectedThreadId = ""; + } + state.sessionSwitching = false; + state.sessionSelectCommandId = ""; + finishSessionSwitchContext(); + } + + function failSessionSwitch(value, fallback = "会话切换失败") { + // A target projection may arrive before the adapter's final owner check. + // A later failure must still restore the old routing context; treating the + // early snapshot as success strands the browser on an unconfirmed target. + restoreSessionSwitchContext(); + return sessionErrorMessage(value, fallback); + } + + function finishSessionSwitchIfReady() { + const context = state.sessionSwitchContext; + if (!context || !context.targetSnapshotReady || !context.selectedAckReady) return false; + const targetThreadId = String(context.targetThreadId || ""); + if (!targetThreadId || state.syncedThreadId !== targetThreadId) return false; + const titleNode = $("threadTitle"); + if (context.targetTitle && titleNode) titleNode.textContent = context.targetTitle; + state.sessionSelectedThreadId = ""; + state.sessionSwitching = false; + state.sessionSelectCommandId = ""; + finishSessionSwitchContext(); + syncSessionActive(targetThreadId); + return true; + } + + function sessionEntryId(entry) { + if (!isRecord(entry)) return ""; + const thread = isRecord(entry.thread) ? entry.thread : {}; + return firstString(entry.threadId, entry.conversationId, entry.conversation_id, entry.id, + thread.threadId, thread.conversationId, thread.conversation_id, thread.id); + } + + function sessionEntryTitle(entry) { + if (!isRecord(entry)) return ""; + const thread = isRecord(entry.thread) ? entry.thread : {}; + return firstString(entry.title, entry.name, entry.preview, entry.firstUserMessage, entry.first_user_message, + entry.threadTitle, entry.thread_name, thread.title, thread.name, thread.preview, thread.thread_name); + } + + function sessionEntryCwd(entry) { + if (!isRecord(entry)) return ""; + const thread = isRecord(entry.thread) ? entry.thread : {}; + return firstString(entry.cwd, entry.workspace, entry.workspacePath, entry.workspace_path, + thread.cwd, thread.workspace, thread.workspacePath, thread.workspace_path); + } + + function sessionEntryUpdatedAt(entry) { + if (!isRecord(entry)) return null; + const thread = isRecord(entry.thread) ? entry.thread : {}; + // App-server history is ordered by recency_at. Older relay versions only + // expose updatedAt, so keep those fields as a compatibility fallback. + return timestampMs( + entry.recencyAtMs, entry.recencyAt, entry.recency_at_ms, entry.recency_at, + thread.recencyAtMs, thread.recencyAt, thread.recency_at_ms, thread.recency_at, + entry.updatedAtMs, entry.updatedAt, entry.lastUpdatedAtMs, entry.lastUpdatedAt, + entry.updated_at_ms, entry.updated_at, entry.last_updated_at, + entry.mtime, entry.modifiedAt, thread.updatedAtMs, thread.updatedAt, + thread.updated_at_ms, thread.updated_at, + ); + } + + function sessionEntryStatus(entry) { + if (!isRecord(entry)) return { kind: "idle", label: "", active: false, attention: false, unread: false }; + const thread = isRecord(entry.thread) ? entry.thread : {}; + const nestedStatus = isRecord(entry.status) ? entry.status : {}; + const nestedExecutionStatus = isRecord(entry.executionStatus) ? entry.executionStatus : {}; + const nestedThreadStatus = isRecord(thread.status) ? thread.status : {}; + let rawStatus = firstString( + entry.activity, entry.activityStatus, entry.status, entry.executionStatus, entry.turnStatus, + entry.threadRuntimeStatus, entry.thread_runtime_status, entry.runtimeStatus, entry.lastTurnStatus, + entry.last_turn_status, entry.phase, entry.state, + nestedStatus.activity, nestedStatus.type, nestedStatus.status, nestedStatus.kind, + nestedExecutionStatus.activity, nestedExecutionStatus.type, nestedExecutionStatus.status, nestedExecutionStatus.kind, + nestedThreadStatus.activity, nestedThreadStatus.type, nestedThreadStatus.status, nestedThreadStatus.kind, + thread.threadRuntimeStatus, thread.thread_runtime_status, thread.turnStatus, thread.lastTurnStatus, thread.state, + ).replace(/([a-z])([A-Z])/g, "$1_$2").toLowerCase().replace(/[\s-]+/g, "_"); + if ((!rawStatus || rawStatus === "idle") && entry.active) { + rawStatus = firstString( + state.currentActivity !== "idle" ? state.currentActivity : "", + state.turnStatus, + state.turnId ? "working" : "", + ).replace(/([a-z])([A-Z])/g, "$1_$2").toLowerCase().replace(/[\s-]+/g, "_"); + } + const unread = [ + entry.hasUnreadTurn, entry.has_unread_turn, entry.unread, entry.isUnread, + entry.needsAttention, entry.needs_attention, thread.hasUnreadTurn, thread.has_unread_turn, + ].some((value) => value === true || value === 1 || ["true", "1", "yes"].includes(String(value || "").toLowerCase())); + let kind = "idle"; + if (entry.isApproval === true || /approval|permission|request_approval|awaiting_authorization|needs_authorization|requires_action/.test(rawStatus)) kind = "approval"; + else if (entry.isWaiting === true || /needs?_?input|waiting_for_input|pending_input|pending|queued/.test(rawStatus)) kind = "waiting"; + else if (entry.isEditing === true || /edit|apply_patch|file_change|writing/.test(rawStatus)) kind = "editing"; + else if (entry.isThinking === true || /think|reason/.test(rawStatus)) kind = "thinking"; + else if (entry.isWorking === true || entry.isRunning === true || /run|stream|working|in_progress|active|busy|generat/.test(rawStatus)) kind = "working"; + else if (/error|fail|cancel|interrupt/.test(rawStatus)) kind = "error"; + else if (unread) kind = "unread"; + const labels = { + approval: "等待授权", + waiting: "等待输入", + editing: "编辑中", + thinking: "思考中", + working: "进行中", + error: "异常", + unread: "未读", + idle: "", + }; + return { + kind, + label: labels[kind] || "", + active: ["approval", "waiting", "editing", "thinking", "working"].includes(kind), + attention: ["approval", "waiting", "error", "unread"].includes(kind) || unread, + unread, + }; + } + + function sessionEntrySearchText(entry) { + if (!isRecord(entry)) return ""; + const thread = isRecord(entry.thread) ? entry.thread : {}; + const status = sessionEntryStatus(entry); + return [ + sessionEntryTitle(entry), sessionEntryCwd(entry), sessionEntryId(entry), + status.label, entry.mode, entry.source, entry.threadSource, thread.mode, + ].filter((value) => typeof value === "string" && value.trim()).join(" ").toLowerCase(); + } + + function sessionOptionDomId(threadId) { + try { + return `session-option-${encodeURIComponent(String(threadId)).replace(/%/g, "_")}`; + } catch { + return `session-option-${String(threadId).replace(/[^a-z0-9_-]/gi, "_")}`; + } + } + + function sessionEntryIsActive(entry, activeId = "") { + if (!isRecord(entry)) return false; + const id = sessionEntryId(entry); + return id === activeId || entry.active === true || entry.active === 1 + || ["true", "1", "yes"].includes(String(entry.active || "").toLowerCase()); + } + + function sessionEntryIsAvailable(entry) { + return isRecord(entry) + && entry.available !== false + && entry.canAttach !== false + && entry.attachable !== false; + } + + function sessionIsSelectable(entry, activeId, canSwitch) { + if (!isRecord(entry)) return false; + const id = sessionEntryId(entry); + const available = sessionEntryIsAvailable(entry); + const current = sessionEntryIsActive(entry, activeId); + return Boolean(id && available && !current && canSwitch && !state.sessionSwitching); + } + + function filteredSessionEntries() { + const source = Array.isArray(state.sessions) + ? state.sessions + .filter(isRecord) + .filter((entry) => !state.attachMode || sessionEntryIsAvailable(entry)) + .map((entry, index) => ({ entry, index })) + : []; + source.sort((left, right) => { + const rightTime = sessionEntryUpdatedAt(right.entry) ?? 0; + const leftTime = sessionEntryUpdatedAt(left.entry) ?? 0; + return rightTime - leftTime || left.index - right.index; + }); + const ordered = source.map(({ entry }) => entry); + const query = String(state.sessionSearch || "").trim().toLowerCase(); + return query ? ordered.filter((entry) => sessionEntrySearchText(entry).includes(query)) : ordered; + } + + function sessionPathLabel(value) { + const text = String(value || "").trim(); + if (!text) return t("本地会话"); + const parts = text.split(/[\\/]+/).filter(Boolean); + return parts.length > 1 ? `${t("工作区")} · ${parts[parts.length - 1]}` : text; + } + + function sessionTimeLabel(value) { + const timestamp = timestampMs(value); + if (timestamp === null) return ""; + try { + const date = new Date(timestamp); + const now = new Date(); + const day = Date.UTC(date.getFullYear(), date.getMonth(), date.getDate()); + const today = Date.UTC(now.getFullYear(), now.getMonth(), now.getDate()); + const difference = Math.round((today - day) / 86_400_000); + const locale = uiLocale(); + const time = new Intl.DateTimeFormat(locale, { hour: "2-digit", minute: "2-digit" }).format(date); + if (difference === 0) return time; + if (difference === 1) return `${t("昨天")} ${time}`; + if (difference > 1 && difference < 7) return `${new Intl.DateTimeFormat(locale, { weekday: "short" }).format(date)} ${time}`; + return `${new Intl.DateTimeFormat(locale, { month: "numeric", day: "numeric" }).format(date)} ${time}`; + } catch { + return ""; + } + } + + function renderSessionPicker() { + const list = $("sessionList"); + const status = $("sessionPickerStatus"); + const picker = $("sessionPicker"); + if (!list || !status || !picker) return; + renderControlMode(); + const activeId = String(state.threadId || state.syncedThreadId || ""); + const canSwitch = sessionControlAllowed("sessionSelect") + && state.appReady + && state.ws?.readyState === WebSocket.OPEN + && !state.turnId + && state.requests.size === 0; + const allSessions = Array.isArray(state.sessions) + ? state.sessions.filter(isRecord).filter((entry) => !state.attachMode || sessionEntryIsAvailable(entry)) + : []; + const sessions = filteredSessionEntries(); + const query = String(state.sessionSearch || "").trim(); + const selectableIds = sessions + .filter((entry) => sessionIsSelectable(entry, activeId, canSwitch)) + .map((entry) => sessionEntryId(entry)); + if (!selectableIds.includes(state.sessionFocusedId)) { + state.sessionFocusedId = selectableIds[0] || ""; + } + const searchInput = $("sessionSearchInput"); + const searchClear = $("sessionSearchClear"); + if (searchInput) { + if (searchInput.value !== state.sessionSearch) searchInput.value = state.sessionSearch; + searchInput.setAttribute("aria-expanded", String(state.sessionPickerOpen)); + searchInput.setAttribute("aria-activedescendant", state.sessionFocusedId ? sessionOptionDomId(state.sessionFocusedId) : ""); + } + if (searchClear) searchClear.hidden = !query; + list.replaceChildren(); + picker.dataset.switching = String(state.sessionSwitching); + list.setAttribute("aria-activedescendant", state.sessionFocusedId ? sessionOptionDomId(state.sessionFocusedId) : ""); + status.dataset.tone = state.sessionListError ? "warning" : ""; + if (state.sessionListLoading && !sessions.length) status.textContent = t("正在读取会话…"); + else if (state.sessionSwitching) { + const targetTitle = state.sessionSwitchContext?.targetTitle; + status.textContent = targetTitle + ? uiLocale() === "en-US" ? `Switching to “${targetTitle}”...` : `正在切换到「${targetTitle}」…` + : t("正在切换会话…"); + } + else if (state.sessionListError) status.textContent = t(state.sessionListError); + else if (!canSwitch && sessions.length > 1) status.textContent = t("当前任务结束或请求处理后才能切换"); + else if (query) status.textContent = uiLocale() === "en-US" + ? `${sessions.length}/${allSessions.length} conversations` + : `${sessions.length}/${allSessions.length} 个会话`; + else status.textContent = sessions.length + ? (uiLocale() === "en-US" ? `${sessions.length} conversations` : `${sessions.length} 个会话`) + : ""; + + if (!sessions.length && !state.sessionListLoading) { + const empty = document.createElement("div"); + empty.className = "session-list-empty"; + const listError = String(state.sessionListError || ""); + empty.textContent = listError === "等待 VS Code 主机连接" + ? t("等待 VS Code 伴随扩展连接") + : listError === "VS Code 主机未连接" + ? t("VS Code 伴随扩展未连接") + : listError === "等待 relay 连接" + ? t("等待 relay 连接") + : listError + ? t("无法读取会话") + : query + ? t("没有匹配的会话") + : t(state.attachMode ? "没有可附加的会话" : "没有可控制的会话"); + list.append(empty); + return; + } + for (const raw of sessions) { + if (!isRecord(raw)) continue; + const id = sessionEntryId(raw); + if (!id) continue; + const title = sessionEntryTitle(raw) || `${t("会话")} ${id.slice(0, 8)}`; + const cwd = sessionEntryCwd(raw); + const updated = sessionEntryUpdatedAt(raw); + const current = sessionEntryIsActive(raw, activeId); + const available = sessionEntryIsAvailable(raw); + const statusInfo = sessionEntryStatus(raw); + const switchingTarget = state.sessionSwitching && id === state.sessionSelectedThreadId; + const selectable = sessionIsSelectable(raw, activeId, canSwitch); + const option = document.createElement("button"); + option.type = "button"; + option.className = "session-option"; + option.id = sessionOptionDomId(id); + option.dataset.available = String(available); + option.dataset.threadId = id; + option.dataset.status = statusInfo.kind; + option.dataset.unread = String(statusInfo.unread); + option.dataset.switching = String(switchingTarget); + option.dataset.focused = String(id === state.sessionFocusedId); + option.setAttribute("role", "option"); + option.setAttribute("aria-selected", String(current)); + option.setAttribute("aria-disabled", String(!selectable)); + option.disabled = !selectable; + const titleNode = document.createElement("span"); + titleNode.className = "session-option-title"; + titleNode.textContent = title; + titleNode.title = title; + const timeNode = document.createElement("span"); + timeNode.className = "session-option-time"; + timeNode.textContent = sessionTimeLabel(updated); + const metaNode = document.createElement("span"); + metaNode.className = "session-option-meta"; + metaNode.textContent = sessionPathLabel(cwd); + metaNode.title = cwd || title; + const stateNode = document.createElement("span"); + stateNode.className = "session-option-state"; + const stateLabels = []; + if (switchingTarget) stateLabels.push(t("正在切换")); + else if (!available) stateLabels.push(t("未打开")); + else if (current) stateLabels.push(t("当前")); + else if (!statusInfo.label) stateLabels.push(t("可切换")); + if (available && statusInfo.label && (!current || statusInfo.active || statusInfo.attention)) stateLabels.push(t(statusInfo.label)); + const stateDot = document.createElement("span"); + stateDot.className = "session-status-dot"; + stateDot.setAttribute("aria-hidden", "true"); + stateNode.append(stateDot, document.createTextNode(stateLabels.join(" · "))); + option.append(titleNode, timeNode, metaNode, stateNode); + option.setAttribute("aria-label", `${title}, ${stateLabels.join(", ") || t("会话")}`); + option.addEventListener("mouseenter", () => { + if (option.disabled) return; + state.sessionFocusedId = id; + for (const peer of list.querySelectorAll(".session-option")) peer.dataset.focused = String(peer.dataset.threadId === id); + list.setAttribute("aria-activedescendant", sessionOptionDomId(id)); + }); + if (!option.disabled) option.addEventListener("click", () => { + state.sessionFocusedId = id; + selectSession(id, title); + }); + list.append(option); + } + } + + function sessionPickerOptions() { + return [...document.querySelectorAll("#sessionList .session-option")] + .filter((option) => !option.disabled && option.dataset.threadId); + } + + function setSessionFocus(threadId, { scroll = true } = {}) { + const id = String(threadId || ""); + state.sessionFocusedId = id; + const list = $("sessionList"); + if (!list) return; + const options = list.querySelectorAll(".session-option"); + for (const option of options) option.dataset.focused = String(option.dataset.threadId === id); + list.setAttribute("aria-activedescendant", id ? sessionOptionDomId(id) : ""); + $("sessionSearchInput")?.setAttribute("aria-activedescendant", id ? sessionOptionDomId(id) : ""); + if (scroll) { + const target = [...list.querySelectorAll(".session-option")] + .find((option) => option.dataset.threadId === id); + target?.scrollIntoView?.({ block: "nearest" }); + } + } + + function moveSessionFocus(delta) { + const options = sessionPickerOptions(); + if (!options.length) return false; + let index = options.findIndex((option) => option.dataset.threadId === state.sessionFocusedId); + if (index < 0) index = delta >= 0 ? -1 : 0; + index = (index + delta + options.length) % options.length; + setSessionFocus(options[index].dataset.threadId); + return true; + } + + function activateFocusedSession() { + const id = state.sessionFocusedId; + if (!id) return false; + const entry = (Array.isArray(state.sessions) ? state.sessions : []) + .find((candidate) => sessionEntryId(candidate) === id); + if (!entry) return false; + const activeId = String(state.threadId || state.syncedThreadId || ""); + const canSwitch = !state.turnId && state.requests.size === 0; + if (!sessionIsSelectable(entry, activeId, canSwitch)) return false; + selectSession(id, sessionEntryTitle(entry)); + return true; + } + + function handleSessionPickerKeydown(event) { + if (!state.sessionPickerOpen) return; + if (event.key === "ArrowDown") { + if (moveSessionFocus(1)) event.preventDefault(); + return; + } + if (event.key === "ArrowUp") { + if (moveSessionFocus(-1)) event.preventDefault(); + return; + } + if (event.key === "Home") { + const options = sessionPickerOptions(); + if (options.length) { + setSessionFocus(options[0].dataset.threadId); + event.preventDefault(); + } + return; + } + if (event.key === "End") { + const options = sessionPickerOptions(); + if (options.length) { + setSessionFocus(options[options.length - 1].dataset.threadId); + event.preventDefault(); + } + return; + } + if (event.key === "Enter") { + if (activateFocusedSession()) event.preventDefault(); + return; + } + if (event.key === "Escape") { + event.preventDefault(); + setSessionPicker(false); + } + } + + function setSessionPicker(open) { + const picker = $("sessionPicker"); + const button = $("sessionPickerButton"); + if (!picker || !button) return; + const next = Boolean(open) && sessionControlAllowed("sessionList") && !state.modeSwitching; + picker.hidden = !next; + state.sessionPickerOpen = next; + button.setAttribute("aria-expanded", String(next)); + if (next) { + state.sessionSearch = ""; + state.sessionFocusedId = ""; + $("panelMenu").hidden = true; + $("detailsPopover").hidden = true; + renderSessionPicker(); + requestSessionList(); + scheduleFrame(() => { + if (state.sessionPickerOpen) $("sessionSearchInput")?.focus(); + }); + } else { + state.sessionSearch = ""; + state.sessionFocusedId = ""; + const input = $("sessionSearchInput"); + if (input) { + input.value = ""; + input.setAttribute("aria-expanded", "false"); + input.setAttribute("aria-activedescendant", ""); + } + const clear = $("sessionSearchClear"); + if (clear) clear.hidden = true; + if (picker.contains(document.activeElement)) button.focus(); + } + } + + function openSessionHistory() { + if (!sessionControlAllowed("sessionList")) return; + setSessionPicker(true); + } + + function requestControlMode(value) { + const mode = normalizeControlMode(value); + if (!mode || mode === state.controlMode || controlModeChangeBlocked()) return; + closePopovers(); + state.modeSwitching = true; + state.requestedControlMode = mode; + state.modeRequestEpoch = state.modeEpoch; + setConversationStatus("正在切换控制模式", "active"); + updateIds(); + try { + state.modeCommandId = command("control/mode/set", { mode }); + } catch (error) { + clearControlModeRequest(); + setConversationStatus(error?.message || "控制模式切换失败", "warning"); + updateIds(); + } + } + + function requestNewSession() { + closePopovers(); + if (!sessionControlAllowed("sessionCreate")) { + setConversationStatus("同步模式下会话管理由 VS Code 控制", "warning"); + return; + } + if (!state.ws || state.ws.readyState !== WebSocket.OPEN) { + setConversationStatus("等待 relay 连接", "warning"); + return; + } + if (!state.appReady) { + setConversationStatus("等待 VS Code 主机连接", "warning"); + return; + } + if (state.role !== "operator" && state.role !== "owner" && state.role !== "host") { + setConversationStatus("当前角色不能创建会话", "warning"); + return; + } + if (state.newSessionCommandId) return; + try { + state.newSessionCommandId = command("session/new", {}); + setConversationStatus("正在创建新会话", "active"); + updateIds(); + } catch (error) { + state.newSessionCommandId = ""; + setConversationStatus(sessionErrorMessage(error, "无法创建新会话"), "warning"); + updateIds(); + } + } + + function requestSessionList() { + if (!sessionControlAllowed("sessionList")) return; + if (state.sessionListLoading) { + renderSessionPicker(); + return; + } + if (!state.ws || state.ws.readyState !== WebSocket.OPEN) { + state.sessionListError = "等待 relay 连接"; + renderSessionPicker(); + return; + } + if (!state.appReady) { + state.sessionListError = "等待 VS Code 主机连接"; + renderSessionPicker(); + return; + } + state.sessionListLoading = true; + state.sessionListError = ""; + renderSessionPicker(); + try { + state.sessionListCommandId = command("session/list", {}); + } catch (error) { + state.sessionListLoading = false; + state.sessionListError = sessionErrorMessage(error, "无法读取会话"); + renderSessionPicker(); + } + } + + function selectSession(threadId, title = "") { + const id = String(threadId || "").trim(); + if (!sessionControlAllowed("sessionSelect")) return; + if (!id || state.sessionSwitching || id === state.threadId) return; + if (state.turnId || state.requests.size) { + state.sessionListError = "当前任务仍在运行或等待授权,暂不能切换"; + renderSessionPicker(); + return; + } + const context = beginSessionSwitchContext(id, title); + state.sessionSwitching = true; + state.sessionListError = ""; + setSessionSwitchingVisual(true); + renderSessionPicker(); + try { + state.sessionSelectCommandId = command("session/select", { threadId: id }); + setConversationStatus( + context.targetTitle + ? uiLocale() === "en-US" ? `Switching to “${context.targetTitle}”` : `正在切换到「${context.targetTitle}」` + : t("正在切换会话"), + "active", + ); + } catch (error) { + restoreSessionSwitchContext(); + state.sessionListError = sessionErrorMessage(error, "会话切换失败"); + renderSessionPicker(); + } + } + + function applySessionListResult(result) { + const body = isRecord(result) ? result : {}; + const source = Array.isArray(result) ? result : Array.isArray(body.sessions) ? body.sessions : Array.isArray(body.threads) ? body.threads : []; + const activeId = firstString(body.activeThreadId, body.threadId); + state.sessions = source + .filter((entry) => isRecord(entry)) + .filter((entry) => !state.attachMode || sessionEntryIsAvailable(entry)) + .map((entry) => ({ ...entry })); + if (activeId) state.sessions = state.sessions.map((entry) => ({ ...entry, active: sessionEntryId(entry) === activeId || entry.active === true })); + state.sessionListLoading = false; + state.sessionListCommandId = ""; + state.sessionListError = ""; + renderSessionPicker(); + } + + function applySessionSelectResult(result) { + const body = isRecord(result) ? result : {}; + const selected = firstString(body.threadId, body.activeThreadId); + state.sessionSelectedThreadId = selected; + if (selected && state.sessions.length) { + state.sessions = state.sessions.map((entry) => ({ ...entry, active: sessionEntryId(entry) === selected })); + } + // A command result confirms only that the command completed. The switch + // itself stays fenced until both the sequenced `session.selected` event + // and the target's authoritative projection have independently arrived. + // This also keeps automatic VS Code navigation and picker navigation on + // the same completion path. + if (state.sessionSwitchContext) { + state.sessionSwitching = true; + setSessionSwitchingVisual(true); + const context = state.sessionSwitchContext; + setConversationStatus( + context.targetSnapshotReady ? "正在确认会话" : "正在加载会话", + "active", + ); + } else { + state.sessionSwitching = false; + state.sessionSelectCommandId = ""; + finishSessionSwitchContext(); + } + renderSessionPicker(); + if (!state.sessionSwitching) { + setSessionPicker(false); + setConversationStatus("会话已切换", "ready"); + } + requestRefresh(); + } + + function syncSessionActive(threadId) { + const activeId = String(threadId || ""); + if (!activeId || !Array.isArray(state.sessions) || !state.sessions.length) return; + state.sessions = state.sessions.map((entry) => ({ + ...entry, + active: sessionEntryId(entry) === activeId, + })); + renderSessionPicker(); + } + + function messageText(item) { + if (!isRecord(item)) return ""; + const direct = [item.text, item.content, item.message, item.summary, item.output] + .find((value) => typeof value === "string" && value.length); + if (direct) return direct; + if (Array.isArray(item.content)) { + return item.content.map((part) => { + if (typeof part === "string") return part; + if (!isRecord(part)) return ""; + return part.text || part.content || part.value || ""; + }).filter(Boolean).join("\n"); + } + // Official collaboration records are often state-only projections: they + // carry an agent id/name and lifecycle kind, but no user-facing `text`. + // Keep those rows in the transcript so the sub-agent disclosure and panel + // can be hydrated from the same snapshot. + if (historyKind(item) === "subagent") { + const agent = isRecord(item.agent) ? item.agent : {}; + const name = firstString( + item.displayName, + item.agentNickname, + item.agentName, + item.agentPath, + agent.displayName, + agent.name, + item.agentThreadId ? `thread ${item.agentThreadId}` : "子代理", + ); + const objective = firstString( + item.objective, + item.prompt, + item.statusMessage, + item.description, + agent.objective, + agent.prompt, + ); + const action = firstString(item.action, item.tool, item.activityKind, agent.action); + const rawStatus = firstString(item.displayStatus, item.status, item.state, agent.status); + const normalizedStatus = normalizeSubagentStatus(rawStatus); + const statusText = subagentStatusLabel(normalizedStatus); + if (objective && name) return `${name}:${objective}`; + if (action && name) return `${name} · ${action}`; + if (name) return `${name} · ${statusText}`; + return `子代理 · ${statusText}`; + } + if (isReadActivity(item)) { + return firstString(...readPathList(item), commandText(item), "读取文件"); + } + if (/(?:command|exec|process|tool)/i.test(String(item.type || item.kind || ""))) { + return firstString(commandText(item), "工具输出"); + } + const fallback = [item.reasoning, item.plan, item.steps, item.diff, item.patch, item.description] + .map((value) => displayValue(value)) + .find((value) => value); + if (fallback) return fallback; + return ""; + } + + function historyKind(item) { + const sourceKind = String(item?.kind || item?.type || "assistant").toLowerCase(); + if (item?.uiType === "subagent-activity" || item?.uiType === "multi-agent-action" + || sourceKind.includes("subagent") || sourceKind.includes("collabagent")) return "subagent"; + if (item?.role === "user" || sourceKind.includes("user")) return "user"; + if (isReadActivity(item)) return "read"; + if (sourceKind.includes("edit") || sourceKind.includes("filechange") || sourceKind.includes("patch")) return "edit"; + if (sourceKind.includes("tool") || sourceKind.includes("command") || sourceKind.includes("exec") || sourceKind.includes("process") + || sourceKind.includes("websearch") || sourceKind.includes("mcp") || sourceKind.includes("imageview") + || sourceKind.includes("generatedimage") || sourceKind.includes("dynamictool") + || sourceKind.includes("permissionrequest") || sourceKind.includes("userinput")) return "tool"; + if (sourceKind.includes("reasoning") || sourceKind.includes("contextcompaction") || sourceKind.includes("approvalreview")) return "reasoning"; + if (sourceKind.includes("plan") || sourceKind.includes("todo")) return "plan"; + return "assistant"; + } + + function normalizedAssistantPhase(item) { + return String(item?.phase || item?.messagePhase || "") + .replace(/[\s-]+/g, "_") + .toLowerCase(); + } + + // Agent commentary is a work item in the official transcript. It belongs + // inside the worked-for disclosure, while the final answer stays outside it. + // Older snapshots omit `phase`, so use the last item in the turn as the + // fallback final answer boundary. + function isAssistantCommentary(item, index, turnKey, turnLastIndex) { + if (historyKind(item) !== "assistant") return false; + const phase = normalizedAssistantPhase(item); + if (phase === "final_answer" || phase === "finalanswer") return false; + if (phase === "commentary" || phase === "analysis" || phase === "reasoning" || phase === "thinking") return true; + return turnLastIndex instanceof Map && turnLastIndex.get(turnKey) !== index; + } + + function structuredTurnFinalAssistantIndexes(messages) { + const result = new Map(); + const explicit = new Set(); + if (!Array.isArray(messages)) return result; + messages.forEach((item, index) => { + if (!isRecord(item) || historyKind(item) !== "assistant") return; + const key = structuredMessageTurn(item, index); + const phase = normalizedAssistantPhase(item); + if (phase === "final_answer" || phase === "finalanswer") { + result.set(key, index); + explicit.add(key); + } else if (!explicit.has(key)) { + // A bookkeeping/sub-agent item can follow the answer in newer + // snapshots. Choose the last assistant message, not the last item. + result.set(key, index); + } + }); + return result; + } + + function structuredDisplayKind(item, index, turnKey, turnLastIndex) { + const kind = historyKind(item); + return kind === "assistant" && isAssistantCommentary(item, index, turnKey, turnLastIndex) + ? "commentary" + : kind; + } + + function isCollapsibleKind(kind) { + return kind === "tool" || kind === "read" || kind === "edit" || kind === "reasoning" + || kind === "plan" || kind === "subagent" || kind === "commentary"; + } + + function isReadActivity(item) { + if (!isRecord(item)) return false; + const type = normalizedItemType(item); + const semantic = String(item.activityKind || item.uiType || item.operation || item.commandType || "") + .replace(/[\s_./-]+/g, "") + .toLowerCase(); + const parsed = item.parsedCmd || item.parsedCommand || item.parsedCommandType; + const parsedType = typeof parsed === "string" + ? parsed + : isRecord(parsed) ? firstString(parsed.type, parsed.kind, parsed.operation) : ""; + const parsedNormalized = String(parsedType).replace(/[\s_./-]+/g, "").toLowerCase(); + return semantic.includes("fileread") || semantic.includes("readfile") || semantic === "read" + || type.includes("fileread") || type.includes("readfile") || type === "read" + || parsedNormalized === "read" || parsedNormalized === "fileread"; + } + + function readPathList(item) { + if (!isRecord(item)) return []; + const paths = []; + const visit = (value) => { + if (typeof value === "string" && value.trim()) paths.push(value.trim()); + else if (Array.isArray(value)) value.forEach(visit); + else if (isRecord(value)) visit(value.path ?? value.file ?? value.filePath ?? value.name); + }; + [item.path, item.file, item.filePath, item.filename, item.name, item.readPath, item.readPaths, item.files, item.paths].forEach(visit); + return [...new Set(paths)]; + } + + function readSummaryLabel(item, status, duration = "") { + const paths = [...new Set(readPathList(item).flatMap((value) => String(value).split(/\r?\n/).map((part) => part.trim()).filter(Boolean)))]; + const suffix = duration ? ` · ${duration}` : ""; + if (status === "inProgress") return paths.length === 1 + ? uiWithRaw("正在读取 ", "Reading ", paths[0]) + : t("正在读取文件"); + if (status === "failed") return paths.length === 1 + ? `${uiWithRaw("读取失败 · ", "Failed to read ", paths[0])}${suffix}` + : `${t("读取文件失败")}${suffix}`; + if (status === "interrupted") return paths.length === 1 + ? `${uiWithRaw("已停止读取 ", "Stopped reading ", paths[0])}${suffix}` + : `${t("已停止读取文件")}${suffix}`; + if (paths.length === 1) return `${uiWithRaw("已读取 ", "Read ", paths[0])}${suffix}`; + if (paths.length > 1) { + const count = String(paths.length); + return `${uiText("已读取这些内容 · ", "Read these items · ")}${count}${uiText(" 个文件", " files")}${suffix}`; + } + return `${t("已读取文件")}${suffix}`; + } + + function structuredMessageKey(item, index) { + const id = item?.id ?? item?.itemId; + if (id !== undefined && id !== null && String(id)) return `id:${String(id)}`; + const turn = item?.turnId ? String(item.turnId) : "item"; + return `turn:${turn}:${historyKind(item)}:${index}`; + } + + function structuredMessageRole(item, kind) { + return item.role === "user" || kind === "user" + ? "user" + : item.role === "tool" || kind === "tool" || kind === "read" || kind === "edit" || kind === "subagent" + ? "tool" + : item.role === "reasoning" || kind === "reasoning" || kind === "plan" + ? "system" + : item.role === "error" + ? "error" + : "assistant"; + } + + function structuredMessageStatus(item) { + if (item.status !== undefined && item.status !== null && item.status !== "") { + return normalizeActivityStatus(item.status, String(item.status)); + } + if (item.completed === true) return "completed"; + if (item.completed === false) return "inProgress"; + if (item.turnStatus !== undefined && item.turnStatus !== null && item.turnStatus !== "") { + return normalizeActivityStatus(item.turnStatus, String(item.turnStatus)); + } + return ""; + } + + function explicitStructuredTurnId(item) { + if (!isRecord(item)) return ""; + const nestedTurn = isRecord(item.turn) ? item.turn : {}; + const value = [ + item.turnId, + item.conversationTurnId, + item.conversation_turn_id, + item.turn_id, + nestedTurn.id, + nestedTurn.turnId, + ].find((candidate) => candidate !== undefined && candidate !== null && String(candidate).trim()); + return value === undefined ? "" : String(value); + } + + function deriveStructuredTurnKeys(messages) { + if (!Array.isArray(messages)) return new Map(); + const keys = new Map(); + let anonymousNumber = 0; + let currentKey = ""; + let currentHasUser = false; + let currentEnded = false; + const makeAnonymousKey = () => `anonymous:${anonymousNumber++}`; + messages.forEach((item, index) => { + if (!isRecord(item)) return; + const explicit = explicitStructuredTurnId(item); + const kind = historyKind(item); + if (explicit) { + currentKey = explicit; + currentHasUser = kind === "user"; + currentEnded = false; + keys.set(item, currentKey); + return; + } + if (!currentKey || kind === "user" && (currentHasUser || currentEnded)) { + currentKey = makeAnonymousKey(); + currentHasUser = false; + currentEnded = false; + } + if (kind === "user") currentHasUser = true; + keys.set(item, currentKey); + // A final-answer marker is the only reliable boundary in a legacy + // projection. Do not use turnStatus here: adapters attach the terminal + // turn status to every item in the turn. + const phase = normalizedAssistantPhase(item); + if (kind === "assistant" && (phase === "final_answer" || phase === "finalanswer")) currentEnded = true; + // Keep the index in the map as a debugging fallback for primitive array + // entries, while object identity remains the canonical lookup. + if (index < 0) keys.set(item, currentKey); + }); + state.structuredTurnKeys = new WeakMap(); + for (const [item, key] of keys) state.structuredTurnKeys.set(item, key); + return keys; + } + + function structuredMessageTurn(item, index) { + const explicit = explicitStructuredTurnId(item); + if (explicit) return explicit; + if (isRecord(item)) { + const derived = state.structuredTurnKeys?.get(item); + if (derived) return derived; + } + return `item:${index}`; + } + + function structuredActivityKey(item, kind, turnKey, index) { + const itemId = item?.itemId ?? item?.id; + const threadKey = item?.threadId || state.threadId || "thread"; + if (itemId !== undefined && itemId !== null && String(itemId)) { + return `${threadKey}:${turnKey || "turn"}:${kind}:${typeof itemId}:${String(itemId)}`; + } + return `structured:${structuredMessageKey(item, index)}`; + } + + function structuredTurnWorkedDurations(messages) { + const result = new Map(); + const explicitKeys = new Set(); + const starts = new Map(); + const ends = new Map(); + if (!Array.isArray(messages)) return result; + deriveStructuredTurnKeys(messages); + const finalAssistantIndexes = structuredTurnFinalAssistantIndexes(messages); + messages.forEach((item, index) => { + const key = structuredMessageTurn(item, index); + const explicit = finiteNumber( + item?.workedDurationMs, + item?.workDurationMs, + item?.workedForMs, + item?.turnWorkedDurationMs, + item?.workedFor?.durationMs, + ); + if (explicit !== null) { + result.set(key, Math.max(0, explicit)); + explicitKeys.add(key); + } + const workStart = timestampMs( + item?.firstTurnWorkItemStartedAtMs, + item?.workStartedAtMs, + item?.turnWorkStartedAtMs, + item?.turnStartedAtMs, + item?.workedFor?.startedAtMs, + ); + if (workStart !== null && !starts.has(key)) starts.set(key, workStart); + const kind = historyKind(item); + if (kind !== "user") { + const itemStart = timestampMs(item?.startedAtMs, item?.startedAt, item?.createdAtMs); + if (itemStart !== null && kind !== "assistant") { + // Activity timestamps are the closest equivalent to the official + // first-turn-work-item marker in older follower snapshots. + if (!starts.has(key)) starts.set(key, itemStart); + } + } + const explicitAssistantStart = timestampMs( + item?.finalAssistantStartedAtMs, + item?.workedFor?.completedAtMs, + ); + // The official worked-for clock ends when the final assistant response + // starts. Earlier assistant commentary is part of the work group and + // must not extend the duration; this was the source of the recurring + // three-to-five-second discrepancy in legacy snapshots. + const isFinalAssistant = kind === "assistant" && finalAssistantIndexes.get(key) === index; + const assistantStart = explicitAssistantStart !== null + ? explicitAssistantStart + : isFinalAssistant ? timestampMs(item?.startedAtMs, item?.startedAt, item?.createdAtMs) : null; + if (assistantStart !== null) ends.set(key, Math.max(ends.get(key) || 0, assistantStart)); + }); + for (const [key, start] of starts) { + if (result.has(key)) continue; + const end = ends.get(key); + if (end !== undefined && end >= start) result.set(key, end - start); + } + // Some older host snapshots keep the authoritative worked-for duration in + // session metadata rather than repeating it on each projected item. Apply + // it to the most recent turn only; explicit per-item values always win. + const metadataDuration = finiteNumber(state.workedDurationMs, state.lastWorkedDurationMs); + if (metadataDuration !== null && messages.length) { + const keysInOrder = messages + .filter(isRecord) + .map((item, index) => structuredMessageTurn(item, index)); + const target = state.turnId && keysInOrder.includes(state.turnId) + ? state.turnId + : keysInOrder.at(-1); + // Session metadata is authoritative over timestamp inference, but an + // explicit per-turn/item value from the host remains the strongest + // signal. + if (target && !explicitKeys.has(target)) result.set(target, Math.max(0, metadataDuration)); + } + // Older attach snapshots may expose the final-answer start separately + // from the duration. Use it only for the corresponding most recent turn. + const metadataFinal = timestampMs(state.finalAssistantStartedAt); + if (metadataFinal !== null && messages.length) { + const keyed = messages.filter(isRecord).map((item, index) => structuredMessageTurn(item, index)); + const target = state.turnId && keyed.includes(state.turnId) ? state.turnId : keyed.at(-1); + const start = target ? starts.get(target) : null; + if (target && start !== undefined && !explicitKeys.has(target) && !result.has(target) && metadataFinal >= start) { + result.set(target, metadataFinal - start); + } + } + return result; + } + + function structuredActivityTiming(messages, index, turnKey, status) { + const item = Array.isArray(messages) ? messages[index] : null; + if (!isRecord(item)) return { durationMs: null, finishedAt: null }; + const explicitDuration = finiteNumber(item.durationMs, item.elapsedMs); + const startedAt = timestampMs(item.startedAtMs, item.startedAt, item.createdAtMs); + const explicitFinishedAt = timestampMs(item.completedAtMs, item.completedAt, item.finishedAtMs, item.finishedAt); + if (explicitDuration !== null) { + return { + durationMs: explicitDuration, + finishedAt: explicitFinishedAt ?? (startedAt === null ? null : startedAt + explicitDuration), + }; + } + if (explicitFinishedAt !== null && startedAt !== null) { + return { durationMs: Math.max(0, explicitFinishedAt - startedAt), finishedAt: explicitFinishedAt }; + } + // Hydrated legacy rows sometimes expose only start times. The next item in + // the same turn is the closest end marker (and matches the official + // projection's item interval). Never infer an end for an explicitly live + // row, and never let a later turn's timestamp inflate this activity. + if (status !== "inProgress" && startedAt !== null && Array.isArray(messages)) { + for (let cursor = index + 1; cursor < messages.length; cursor += 1) { + const next = messages[cursor]; + if (!isRecord(next)) continue; + if (structuredMessageTurn(next, cursor) !== turnKey) break; + const nextStartedAt = timestampMs(next.startedAtMs, next.startedAt, next.createdAtMs); + if (nextStartedAt === null) continue; + if (nextStartedAt >= startedAt) return { + durationMs: nextStartedAt - startedAt, + finishedAt: nextStartedAt, + }; + break; + } + } + return { durationMs: null, finishedAt: explicitFinishedAt }; + } + + function indexStructuredActivity(article, item, kind, turnKey, index, messages = null) { + if (!article || !isCollapsibleKind(kind)) return null; + const key = article.dataset.activityKey || structuredActivityKey(item, kind, turnKey, index); + article.dataset.activityKey = key; + article.dataset.activityKind = kind; + const existing = state.activities.get(key); + const turnStatus = normalizeActivityStatus(item.turnStatus, ""); + const status = structuredMessageStatus(item) || (turnStatus === "inProgress" ? "inProgress" : "completed"); + const timing = structuredActivityTiming(messages, index, turnKey, status); + const durationMs = timing.durationMs; + const startedAt = timestampMs(item.startedAtMs, item.startedAt, item.createdAtMs); + const activity = existing || { + key, + kind, + role: kind === "commentary" + ? "assistant" + : kind === "tool" || kind === "read" || kind === "edit" || kind === "subagent" ? "tool" : "system", + messageKind: kind === "tool" || kind === "read" ? "tool" : kind, + label: item.label || activityLabelForItem(item, kind), + command: kind === "tool" || kind === "read" ? terminalCommandText(commandText(item)) : commandText(item), + filePath: kind === "read" ? readPathList(item).join("\n") : "", + cwd: (kind === "tool" || kind === "read") && typeof item.cwd === "string" ? item.cwd : "", + shellName: (kind === "tool" || kind === "read") && typeof item.shellName === "string" && item.shellName ? item.shellName : "Shell", + agentThreadId: kind === "subagent" ? firstString(item.agentThreadId, item.childThreadId, item.threadId) : "", + displayName: kind === "subagent" ? firstString(item.displayName, item.agentNickname, item.agentName, item.agentPath) : "", + objective: kind === "subagent" ? firstString(item.objective, item.prompt, item.statusMessage, item.message) : "", + activityKind: kind === "subagent" ? firstString(item.activityKind, item.kind) : "", + displayStatus: kind === "subagent" ? firstString(item.displayStatus, item.status) : "", + model: kind === "subagent" ? firstString(item.model, item.modelId) : "", + action: kind === "subagent" ? firstString(item.action, item.tool) : "", + prompt: kind === "subagent" && item.prompt !== null ? firstString(item.prompt) : "", + senderThreadId: kind === "subagent" ? firstString(item.senderThreadId) : "", + receiverThreadIds: kind === "subagent" + ? Array.isArray(item.receiverThreadIds) + ? item.receiverThreadIds.map(String) + : Array.isArray(item.receiverThreads) ? item.receiverThreads.map(String) : [] + : [], + agentsStates: kind === "subagent" && isRecord(item.agentsStates) ? item.agentsStates : {}, + canInteract: item.canInteract !== false, + exitCode: kind === "tool" || kind === "read" ? finiteNumber(item.exitCode, item.exit_code) : null, + itemId: item.itemId === undefined || item.itemId === null ? "" : String(item.itemId), + threadId: item.threadId || state.threadId || "", + turnId: turnKey || "", + startedAt: startedAt || null, + finishedAt: null, + durationMs: null, + durationExplicit: finiteNumber(item.durationMs, item.elapsedMs) !== null, + status, + headerText: "", + outputText: "", + anonymous: false, + concrete: true, + statusOnly: false, + article, + body: article.querySelector(".details-body") || article.querySelector(".message-body"), + wrapper: article.querySelector(".message-content"), + details: article.querySelector("details"), + summary: article.querySelector("summary"), + }; + activity.key = key; + activity.kind = kind; + activity.messageKind = kind === "tool" || kind === "read" ? "tool" : kind; + activity.itemId = item.itemId === undefined || item.itemId === null ? activity.itemId || "" : String(item.itemId); + // Structured history is an authoritative concrete projection. If a live + // status row occupied the same slot during reconnect, promote it here so + // the next terminal update does not discard the hydrated item. + activity.concrete = true; + activity.statusOnly = false; + activity.anonymous = false; + activity.turnId = turnKey || activity.turnId || ""; + activity.threadId = item.threadId || activity.threadId || state.threadId || ""; + activity.label = item.label || activity.label || activityLabelForItem(item, kind); + activity.command = kind === "tool" || kind === "read" + ? terminalCommandText(commandText(item)) || activity.command || "" + : commandText(item) || activity.command || ""; + if (kind === "tool" || kind === "read") { + if (kind === "read") activity.filePath = readPathList(item).join("\n") || activity.filePath; + activity.cwd = typeof item.cwd === "string" ? item.cwd : activity.cwd || ""; + activity.shellName = typeof item.shellName === "string" && item.shellName + ? item.shellName + : activity.shellName || "Shell"; + activity.exitCode = item.exitCode === undefined || item.exitCode === null + ? activity.exitCode ?? null + : finiteNumber(item.exitCode, item.exit_code); + } + if (kind === "subagent") { + activity.agentThreadId = firstString(item.agentThreadId, item.childThreadId, item.threadId, activity.agentThreadId); + activity.displayName = firstString(item.displayName, item.agentNickname, item.agentName, item.agentPath, activity.displayName); + activity.objective = firstString(item.objective, item.prompt, item.statusMessage, item.message, activity.objective); + activity.activityKind = firstString(item.activityKind, item.kind, activity.activityKind); + activity.displayStatus = firstString(item.displayStatus, item.status, activity.displayStatus); + activity.model = firstString(item.model, item.modelId, activity.model); + activity.action = firstString(item.action, item.tool, activity.action); + if (item.prompt !== null && item.prompt !== undefined) activity.prompt = firstString(item.prompt, activity.prompt); + activity.senderThreadId = firstString(item.senderThreadId, activity.senderThreadId); + if (Array.isArray(item.receiverThreadIds)) activity.receiverThreadIds = item.receiverThreadIds.map(String); + else if (Array.isArray(item.receiverThreads)) activity.receiverThreadIds = item.receiverThreads.map(String); + if (isRecord(item.agentsStates)) activity.agentsStates = item.agentsStates; + if (item.canInteract !== undefined) activity.canInteract = item.canInteract !== false; + if (activity.agentThreadId) article.dataset.agentThreadId = activity.agentThreadId; + else delete article.dataset.agentThreadId; + } + activity.startedAt = startedAt || activity.startedAt || null; + if (durationMs !== null || !existing) activity.durationMs = durationMs; + activity.durationExplicit = finiteNumber(item.durationMs, item.elapsedMs) !== null; + activity.status = status; + activity.finishedAt = timing.finishedAt || activity.finishedAt || null; + activity.article = article; + activity.body = article.querySelector(".details-body") || article.querySelector(".message-body"); + activity.wrapper = article.querySelector(".message-content"); + activity.details = article.querySelector("details"); + activity.summary = article.querySelector("summary"); + activity.outputText = kind === "tool" || kind === "read" ? activityOutput(item, kind) : messageText(item); + state.activities.set(key, activity); + // Hydrated rows may have been created before lifecycle metadata arrived. + // Recompute the disclosure label from the normalized status/timestamps so + // history and live updates cannot leave stale `0ms` or generic text behind. + renderActivityText(activity); + refreshActivity(activity); + return activity; + } + + function updateStructuredMessageArticle(article, item, index, turnLastIndex, turnHasActivity = new Set(), turnWorkedDurations = new Map()) { + const itemText = messageText(item); + if (!itemText) return false; + const sourceKind = historyKind(item); + const turnKey = structuredMessageTurn(item, index); + const kind = structuredDisplayKind(item, index, turnKey, turnLastIndex); + const role = structuredMessageRole(item, kind); + const status = structuredMessageStatus(item) || "completed"; + const durationMs = item.durationMs ?? item.elapsedMs; + const duration = elapsedDuration(durationMs); + const body = article.querySelector(".message-body"); + if (!body) return false; + const details = article.querySelector("details"); + const wasOpen = details ? isDetailsExpanded(details) : undefined; + article.dataset.rawText = String(itemText); + article.dataset.turnId = turnKey; + article.dataset.kind = kind; + if (item.itemId !== undefined && item.itemId !== null) article.dataset.itemId = String(item.itemId); + if (item.itemType) article.dataset.itemType = String(item.itemType); + const agentThreadId = kind === "subagent" ? firstString(item.agentThreadId, item.childThreadId, item.threadId) : ""; + if (agentThreadId) article.dataset.agentThreadId = agentThreadId; + else delete article.dataset.agentThreadId; + if (item.startedAtMs || item.completedAtMs) { + const parsed = timestampMs(item.startedAtMs ?? item.completedAtMs); + if (parsed !== null) article.dataset.timestamp = String(parsed); + } + if (status) article.dataset.status = status; + const isActivity = isCollapsibleKind(kind); + article.classList.toggle("streaming", status === "inProgress" && !isActivity); + renderMessageBody(body, itemText, role, isActivity ? "activity" : "history", kind); + if (details) { + const summary = details.querySelector("summary"); + const activitySummary = historyActivitySummary(item, kind, status, duration); + if (summary && activitySummary) setActivitySummary(summary, activitySummary, kind); + if (wasOpen !== undefined) setDetailsExpanded(details, wasOpen, { immediate: true, preserve: false }); + } + const finalAssistant = sourceKind === "assistant" && role === "assistant" && kind === "assistant" + && (item.phase === "final_answer" || item.phase === "final-answer" || !isAssistantCommentary(item, index, turnKey, turnLastIndex)); + if (finalAssistant && status !== "inProgress" && !article.classList.contains("streaming") && !article.querySelector(".message-actions")) { + addMessageActions(article, true); + } + const terminalTurn = ["completed", "complete", "done", "failed", "error", "interrupted", "cancelled", "canceled"] + .includes(String(item.turnStatus || "").replace(/[\s-]+/g, "_").toLowerCase()); + const workedDuration = turnWorkedDurations.get(turnKey) + ?? workedDurationFor(item, null); + if (finalAssistant && (terminalTurn || workedDuration !== null || (durationMs !== undefined && status !== "inProgress"))) { + appendTurnDivider(turnKey, item.turnStatus || status || "completed", workedDuration ?? durationMs, article); + } + return true; + } + + function reconcileStructuredOutput(text, structuredMessages) { + const output = $("output"); + if (!output || !Array.isArray(structuredMessages) || !structuredMessages.length) return false; + const articles = [...output.querySelectorAll(".message[data-structured-key]")]; + if (articles.length !== structuredMessages.length) return false; + const keys = structuredMessages.map((item, index) => structuredMessageKey(item, index)); + if (articles.some((article, index) => article.dataset.structuredKey !== keys[index])) return false; + const follow = shouldFollowOutput(output); + const turnLastIndex = structuredTurnFinalAssistantIndexes(structuredMessages); + const turnHasActivity = new Set(); + const turnWorkedDurations = structuredTurnWorkedDurations(structuredMessages); + structuredMessages.forEach((item, index) => { + const turnKey = structuredMessageTurn(item, index); + if (isCollapsibleKind(structuredDisplayKind(item, index, turnKey, turnLastIndex))) turnHasActivity.add(turnKey); + }); + for (const [index, item] of structuredMessages.entries()) { + if (!updateStructuredMessageArticle(articles[index], item, index, turnLastIndex, turnHasActivity, turnWorkedDurations)) return false; + const turnKey = structuredMessageTurn(item, index); + const kind = structuredDisplayKind(item, index, turnKey, turnLastIndex); + if (isCollapsibleKind(kind)) { + indexStructuredActivity(articles[index], item, kind, turnKey, index, structuredMessages); + } + } + reconcileTurnDividers(); + expandLatestTurnActivity(); + if (typeof text === "string") output.dataset.outputTail = text; + state.structuredMessages = structuredMessages.slice(); + state.outputSynced = true; + if (state.turnId && state.turnStartedAt !== null && ["active", "waiting"].includes(state.turnStatus) + && [...output.querySelectorAll(".message.activity")].some((article) => article.dataset.turnId === state.turnId)) { + ensureLiveTurnDivider(state.turnId); + } + if (follow) scrollOutput(output, true); + updateScrollToBottom(output); + return true; + } + + function replaceOutput(text, structuredMessages) { + const output = $("output"); + const preserveBottom = !state.outputSynced || !output || shouldFollowOutput(output); + const preservedDistanceFromBottom = output + ? Math.max(0, output.scrollHeight - output.scrollTop - output.clientHeight) + : 0; + const pendingUserText = state.pendingUserText; + const snapshotHasPendingUser = Boolean(pendingUserText && ( + text.includes(`> ${pendingUserText}`) + || (Array.isArray(structuredMessages) && structuredMessages.some((item) => { + const kind = historyKind(item); + return kind === "user" && comparableText(messageText(item)) === comparableText(pendingUserText); + })) + )); + renderEmptyOutput(); + if (typeof text === "string") output.dataset.outputTail = text; + state.structuredMessages = Array.isArray(structuredMessages) ? structuredMessages.slice() : []; + if (Array.isArray(structuredMessages) && structuredMessages.length) { + // Older follower snapshots may omit the top-level subagents projection. + // Derive the small panel model from the same collaboration items used by + // the transcript so those agents remain discoverable after reconnect. + const suppliedSubagents = state.subagents; + const derivedSubagents = deriveSubagentsFromMessages(structuredMessages); + if (derivedSubagents.length) { + state.subagents = mergeSubagentProjections(derivedSubagents, suppliedSubagents); + renderSubagents(); + } else if (state.subagents.length) { + // A complete snapshot with no collaboration items is authoritative for + // the transcript. Do not carry a finished agent from the previous + // thread into the new composer. + state.subagents = []; + renderSubagents(); + } + const turnLastIndex = structuredTurnFinalAssistantIndexes(structuredMessages); + const turnHasActivity = new Set(); + const turnWorkedDurations = structuredTurnWorkedDurations(structuredMessages); + let previousTurnKey = ""; + structuredMessages.forEach((item, index) => { + const key = structuredMessageTurn(item, index); + if (isCollapsibleKind(structuredDisplayKind(item, index, key, turnLastIndex))) turnHasActivity.add(key); + }); + structuredMessages.forEach((item, index) => { + const itemText = messageText(item); + if (!itemText) return; + const sourceKind = historyKind(item); + const turnKey = structuredMessageTurn(item, index); + const kind = structuredDisplayKind(item, index, turnKey, turnLastIndex); + const role = structuredMessageRole(item, kind); + const status = structuredMessageStatus(item) || "completed"; + const durationMs = item.durationMs ?? item.elapsedMs; + const duration = elapsedDuration(durationMs); + const timestamp = item.startedAtMs ?? item.completedAtMs; + if (role === "user" || role === "assistant") { + appendDateSeparator(timestamp, { + role, + turnId: turnKey, + turnStart: turnKey !== previousTurnKey, + breaksPreviousAdjacency: item.breaksPreviousAdjacency === true, + }); + } + const finalAssistant = sourceKind === "assistant" && role === "assistant" && kind === "assistant" + && (item.phase === "final_answer" || item.phase === "final-answer" || !isAssistantCommentary(item, index, turnKey, turnLastIndex)); + const terminalTurn = ["completed", "complete", "done", "failed", "error", "interrupted", "cancelled", "canceled"] + .includes(String(item.turnStatus || "").replace(/[\s-]+/g, "_").toLowerCase()); + const workedDuration = turnWorkedDurations.get(turnKey) ?? workedDurationFor(item, null); + const activitySummary = historyActivitySummary(item, kind, status, duration) || undefined; + const collapsible = isCollapsibleKind(kind); + const activityKey = collapsible ? structuredActivityKey(item, kind, turnKey, index) : ""; + const message = appendMessage(itemText, role, collapsible ? "activity" : "history", "", { + kind, + status, + label: item.label || (kind === "reasoning" ? "思考" : kind === "plan" ? "计划" : kind === "edit" ? "文件变更" : kind === "read" ? "读取文件" : kind === "tool" ? "工具输出" : kind === "subagent" ? "子代理" : kind === "commentary" ? "工作说明" : ""), + summary: activitySummary, + command: item.command || item.commandLine, + turnId: turnKey, + structuredKey: structuredMessageKey(item, index), + activityKey, + itemId: item.itemId, + itemType: item.itemType, + agentThreadId: kind === "subagent" ? firstString(item.agentThreadId, item.childThreadId, item.threadId) : "", + timestamp, + showTimestamp: role === "user" || role === "assistant", + showActions: role === "user" || (role === "assistant" && finalAssistant && kind !== "commentary"), + collapsible, + open: kind === "commentary" || (status === "inProgress" && (kind === "reasoning" || kind === "plan")), + }); + if (message && collapsible) indexStructuredActivity(message.article, item, kind, turnKey, index, structuredMessages); + if (finalAssistant && (terminalTurn || workedDuration !== null || (durationMs !== undefined && status !== "inProgress"))) { + appendTurnDivider(turnKey, item.turnStatus || status || "completed", workedDuration ?? durationMs, message?.article || null); + } + previousTurnKey = turnKey; + }); + reconcileTurnDividers(); + expandLatestTurnActivity(); + } else if (text) { + // The attach adapter prefixes user items with `> ` and separates items + // with blank lines. Use that stable marker to recreate the two sides of + // the conversation without interpreting arbitrary Markdown as HTML. + const chunks = text.split(/\n{2,}/).map((chunk) => chunk.trim()).filter(Boolean); + let expectUser = true; + for (const chunk of chunks) { + const markedUser = chunk.startsWith("> "); + if (markedUser && expectUser) { + appendMessage(chunk.slice(2), "user", "history", "", { kind: "user" }); + expectUser = false; + } else { + appendMessage(chunk, "assistant", "history", "", { kind: "assistant" }); + expectUser = true; + } + } + } + if (pendingUserText && !snapshotHasPendingUser) { + appendDateSeparator(Date.now(), { role: "user", force: true }); + appendMessage(pendingUserText, "user", "streaming", "", { kind: "user", timestamp: Date.now(), showTimestamp: true }); + } + state.pendingUserText = snapshotHasPendingUser ? "" : pendingUserText; + state.outputSynced = true; + if (state.turnId && state.turnStartedAt !== null && ["active", "waiting"].includes(state.turnStatus) + && [...output.querySelectorAll(".message.activity")].some((article) => article.dataset.turnId === state.turnId)) { + ensureLiveTurnDivider(state.turnId); + } + if (preserveBottom) scrollOutput(output, true); + else if (output) { + // Rebuilding a long history changes scrollHeight. Preserve the user's + // distance from the bottom just like the official thread scroll layout, + // so a background sync does not yank them away from the message they are + // reading. + output.scrollTop = Math.max(0, output.scrollHeight - output.clientHeight - preservedDistanceFromBottom); + updateScrollToBottom(output); + } + } + + function applyStructuredMessagesPatch(patch) { + if (!isRecord(patch) || !Array.isArray(state.structuredMessages) || !Array.isArray(patch.messages)) return null; + const start = Number(patch.start); + const deleteCount = Number(patch.deleteCount); + if (!Number.isInteger(start) || start < 0 || start > state.structuredMessages.length + || !Number.isInteger(deleteCount) || deleteCount < 0 + || start + deleteCount > state.structuredMessages.length) return null; + return [ + ...state.structuredMessages.slice(0, start), + ...patch.messages, + ...state.structuredMessages.slice(start + deleteCount), + ]; + } + + function appendOutputChunk(text, stream = "codex", context = {}) { + if (!text) return; + const output = $("output"); + const follow = shouldFollowOutput(output); + const normalizedStream = String(stream || "codex").toLowerCase(); + if (normalizedStream === "codex") { + const userItem = text.match(/^(?:\n{2,})?> ([^\n]+)\n?$/); + if (userItem) { + const userText = userItem[1].trim(); + const eventTurn = context.turnId || state.turnId || ""; + finishAssistantStream(); + appendDateSeparator(context.timestamp || Date.now(), { role: "user", turnId: eventTurn, turnStart: true }); + if (!hasRenderedMessage(userText, eventTurn, "user")) { + appendMessage(userText, "user", "history", "", { + kind: "user", + turnId: eventTurn, + timestamp: context.timestamp || Date.now(), + showTimestamp: true, + }); + } + state.pendingUserText = ""; + state.outputSynced = true; + return; + } + } + if (normalizedStream === "codex" && state.pendingUserText) { + const marker = `> ${state.pendingUserText}`; + if (text.includes(marker)) { + text = text.replace(marker, "").replace(/^\n{1,2}/, ""); + state.pendingUserText = ""; + if (!text) return; + } + } + + const eventTurnId = context.turnId || state.turnId || ""; + if (state.outputSynced && eventTurnId && normalizedStream === "codex" + && hasRenderedCompletedMessage(text, eventTurnId, "assistant")) { + // A host replay often sends the same output delta immediately after an + // authoritative snapshot. Do not append a second assistant bubble. + state.pendingUserText = ""; + return; + } + + const inferredActivity = normalizedStream === "reasoning" + ? "thinking" + : normalizedStream === "read" || normalizedStream === "reading" + ? "reading" + : normalizedStream === "stdout" || normalizedStream === "stderr" + ? "running" + : "generating"; + if (!state.currentActivity || ["idle", "completed", "failed", "interrupted"].includes(state.currentActivity)) { + state.currentActivity = inferredActivity; + } + if (state.turnStartedAt === null) startTurnClock(context.turnId || state.turnId, null, null); + updateLiveActivity( + state.currentActivity, + state.currentActivityStartedAt || state.turnStartedAt, + null, + [], + context.turnId || state.turnId, + ); + const activityText = statusActivityLabel(state.currentActivity) || t("正在生成"); + const outputElapsed = elapsedDuration(Math.max(0, Date.now() - (state.turnStartedAt || Date.now()))); + setConversationStatus(outputElapsed ? `${activityText} · ${outputElapsed}` : activityText, "active"); + + // Lifecycle notifications create the canonical row first. Output deltas + // from app-server versions that omit itemId are attached to the newest + // running row of the corresponding kind. + const streamKind = normalizedStream === "reasoning" + ? "reasoning" + : normalizedStream === "read" || normalizedStream === "reading" + ? "read" + : normalizedStream === "stdout" || normalizedStream === "stderr" + ? "tool" + : context.kind || ""; + const activity = streamKind + ? latestRunningActivity(streamKind, context) + : null; + if (activity) { + appendActivityChunk(activity, text); + state.activeAssistantBody = activity.body; + state.activeAssistantStream = normalizedStream; + state.activeAssistantText = activity.outputText; + state.activeAssistantActivityKey = activity.key; + if (follow) scrollOutput(output, true); + state.outputSynced = true; + return; + } + if (state.activeAssistantStream !== normalizedStream) finishAssistantStream(); + if (!state.activeAssistantBody && normalizedStream === "codex") text = text.replace(/^\n{2,}/, ""); + if (!state.activeAssistantBody || !output.contains(state.activeAssistantBody)) { + const role = normalizedStream === "stderr" + ? "error" + : normalizedStream === "stdout" || normalizedStream === "read" || normalizedStream === "reading" + ? "tool" + : normalizedStream === "reasoning" + ? "system" + : "assistant"; + const label = normalizedStream === "reasoning" ? "思考" + : normalizedStream === "stdout" ? "命令输出" + : normalizedStream === "read" || normalizedStream === "reading" ? "读取文件" : ""; + const kind = normalizedStream === "reasoning" ? "reasoning" + : normalizedStream === "stdout" ? "tool" + : normalizedStream === "read" || normalizedStream === "reading" || context.kind === "read" ? "read" : "assistant"; + let message; + if (kind === "reasoning" || kind === "tool" || kind === "read") { + state.activitySequence += 1; + const key = `stream:${state.threadId || "thread"}:${state.turnId || "turn"}:${normalizedStream}:${state.activitySequence}`; + const streamActivity = ensureActivity(key, { + kind, + label: label || (kind === "tool" ? "工具输出" : kind === "read" ? "读取文件" : "思考"), + threadId: state.threadId, + turnId: eventTurnId || state.turnId, + status: "inProgress", + anonymous: true, + concrete: true, + }); + message = streamActivity && { article: streamActivity.article, content: streamActivity.body }; + state.activeAssistantActivityKey = streamActivity?.key || null; + } else { + appendDateSeparator(context.timestamp || Date.now(), { role: "assistant", turnId: eventTurnId, turnStart: false }); + message = appendMessage("", role, "streaming", "", { + kind, + label, + turnId: eventTurnId, + timestamp: context.timestamp || Date.now(), + showTimestamp: true, + showActions: false, + collapsible: false, + }); + } + state.activeAssistantBody = message?.content || null; + state.activeAssistantStream = normalizedStream; + state.activeAssistantText = ""; + } + if (state.activeAssistantActivityKey) { + const streamActivity = state.activities.get(state.activeAssistantActivityKey); + if (streamActivity) { + appendActivityChunk(streamActivity, text); + state.activeAssistantText = streamActivity.outputText; + if (follow) scrollOutput(output, true); + state.outputSynced = true; + return; + } + } + if (state.activeAssistantBody) { + state.activeAssistantText += text; + const article = state.activeAssistantBody.closest(".message"); + const role = article?.classList.contains("system") ? "system" : article?.classList.contains("error") ? "error" : "assistant"; + const kind = article?.dataset.kind || "assistant"; + renderMessageBody(state.activeAssistantBody, state.activeAssistantText, role, "streaming", kind); + if (article) article.dataset.rawText = state.activeAssistantText; + article?.classList.add("streaming"); + } + if (follow) scrollOutput(output, true); + state.outputSynced = true; + } + + function finishAssistantStream() { + const article = state.activeAssistantBody?.closest(".message"); + article?.classList.remove("streaming"); + if (article?.classList.contains("assistant") && !article.querySelector(".message-actions")) addMessageActions(article, true); + if (article?.dataset.kind === "reasoning") { + const details = article.querySelector("details"); + if (details) setDetailsExpanded(details, false); + } + const activity = state.activeAssistantActivityKey + ? state.activities.get(state.activeAssistantActivityKey) + : null; + if (activity?.anonymous && isRunningActivity(activity)) finishActivity(activity, "completed"); + state.activeAssistantBody = null; + state.activeAssistantStream = null; + state.activeAssistantText = ""; + state.activeAssistantActivityKey = null; + } + + function eventBelongsToCurrentTurn(payload) { + if (!payload || typeof payload !== "object") return true; + const threadId = eventThreadId(payload); + const turnId = eventTurnId(payload); + if (threadId && state.threadId && threadId !== state.threadId) return false; + if (turnId && state.turnId && turnId !== state.turnId) return false; + if (turnId && state.retiredTurnIds.has(turnId)) return false; + return true; + } + + function eventMessage(payload) { + if (!payload || typeof payload !== "object") return "未知错误"; + const candidates = [ + payload.message, + payload.error?.message, + payload.error, + payload.params?.message, + payload.params?.error, + payload.text, + ]; + const value = candidates.find((entry) => typeof entry === "string" && entry.trim()); + if (value) return value; + try { return JSON.stringify(payload); } catch { return "未知错误"; } + } + + function setAttachMode(value) { + state.attachMode = Boolean(value); + document.body.classList.toggle("attach-mode", state.attachMode); + $("sessionMode").textContent = t(state.attachMode + ? "已附着 VS Code 当前 Codex 会话;输入、输出和授权都回到同一个会话。" + : "当前为独立 app-server 模式。"); + $("startThreadButton").textContent = t(state.attachMode ? "已附着现有会话" : "启动新 thread"); + const modeLabel = $("modeLabel"); + if (modeLabel) modeLabel.textContent = t(state.attachMode ? "本地模式" : "独立模式"); + const popoverMode = $("popoverMode"); + if (popoverMode) popoverMode.textContent = t(state.attachMode ? "本地模式" : "独立模式"); + } + + function firstString(...values) { + return values.find((value) => typeof value === "string" && value.trim())?.trim() || ""; + } + + // Unlike firstString, settings projections need to preserve an explicit + // null. Codex uses null for a model with no reasoning selector; treating it + // as "missing" leaves the browser showing the previous turn's effort. + function firstDefined(...values) { + return values.find((value) => value !== undefined); + } + + function modelDisplayName(model, effort) { + const value = String(model || "").trim(); + if (!value) return ""; + const words = value + .replace(/^gpt[-_]/i, "") + .replace(/[-_]+/g, " ") + .split(/\s+/) + .filter(Boolean) + .map((word) => /^\d+(?:\.\d+)*$/.test(word) ? word : `${word[0].toUpperCase()}${word.slice(1)}`); + const effortValue = String(effort || "").trim(); + if (effortValue && !words.some((word) => word.toLowerCase() === effortValue.toLowerCase())) { + words.push(`${effortValue[0].toUpperCase()}${effortValue.slice(1)}`); + } + return words.join(" "); + } + + // Keep this tiny fallback aligned with the model power choices shipped in + // the installed official webview. A host that exposes `availableModels` + // always wins; these entries only make the picker useful while an older + // private IPC build has no model/list projection. + const FALLBACK_MODELS = [ + { model: "gpt-5.6-sol", displayName: "5.6 Sol", description: "通用 Codex 模型", efforts: ["low", "medium", "high", "xhigh"] }, + { model: "gpt-5.6-terra", displayName: "5.6 Terra", description: "平衡速度与推理", efforts: ["low", "medium", "high", "xhigh"] }, + ]; + // Labels used by the shipped Work composer. The compact trigger shows the + // model plus its selected reasoning preset (for example `5.6 Sol 标准`). + const EFFORT_LABELS = { none: "默认", minimal: "极低", low: "轻度", medium: "标准", high: "深度", xhigh: "极高", max: "最大", ultra: "Ultra" }; + + // The compact Work picker is a power control. Keep the mapping deterministic + // so keyboard/range input still goes through the same effort protocol used by + // the advanced picker. + function modelPowerOptions(model) { + const efforts = Array.isArray(model?.efforts) ? model.efforts.filter(Boolean) : []; + return efforts.map((effort) => ({ effort, label: t(EFFORT_LABELS[effort] || effort) })); + } + + function normalizeModelOption(value) { + if (typeof value === "string") { + const fallback = FALLBACK_MODELS.find((entry) => entry.model === value); + return fallback ? { ...fallback } : { model: value, displayName: modelDisplayName(value), description: "可用模型", efforts: [] }; + } + if (!isRecord(value)) return null; + const model = firstString(value.model, value.id, value.slug, value.name); + if (!model) return null; + const hasSupported = Array.isArray(value.supportedReasoningEfforts) || Array.isArray(value.efforts); + const supported = Array.isArray(value.supportedReasoningEfforts) + ? value.supportedReasoningEfforts.map((entry) => typeof entry === "string" ? entry : firstString(entry?.reasoningEffort, entry?.effort)).filter(Boolean) + : Array.isArray(value.efforts) ? value.efforts.filter((entry) => typeof entry === "string") : []; + const fallback = FALLBACK_MODELS.find((entry) => entry.model === model); + return { + model, + displayName: firstString(value.displayName, value.label, fallback?.displayName, modelDisplayName(model)), + description: firstString(value.description, fallback?.description), + // Unknown catalog entries must not inherit a fabricated effort list. The + // official advanced picker disables the power controls until the host + // reports capabilities for that model. + efforts: hasSupported ? supported : fallback?.efforts || [], + defaultReasoningEffort: firstString(value.defaultReasoningEffort, value.defaultEffort, fallback?.efforts?.[1]), + hidden: value.hidden === true, + }; + } + + function normalizedModelOptions() { + let source; + if (state.availableModels.length) source = state.availableModels; + else if (state.currentModel && FALLBACK_MODELS.some((entry) => entry.model === state.currentModel)) source = FALLBACK_MODELS; + else if (state.currentModel) source = [state.currentModel]; + else return []; + const seen = new Set(); + const result = []; + for (const entry of source) { + const option = normalizeModelOption(entry); + if (!option || option.hidden || seen.has(option.model)) continue; + seen.add(option.model); + result.push(option); + } + if (state.currentModel && !seen.has(state.currentModel)) { + result.unshift(normalizeModelOption(state.currentModel)); + } + return result; + } + + function currentModelOption() { + return normalizedModelOptions().find((entry) => entry.model === state.currentModel) || normalizedModelOptions()[0]; + } + + function renderModelPicker() { + const button = $("modelPickerButton"); + const label = $("modelLabel"); + const effortLabelNode = $("modelEffortLabel"); + const menu = $("modelMenu"); + const modelOptions = $("modelOptions"); + const effortOptions = $("effortOptions"); + const powerView = $("modelPowerView"); + const advancedView = $("modelAdvancedView"); + const advancedToggle = $("modelAdvancedToggle"); + const powerSlider = $("modelPowerSlider"); + const powerValue = $("modelPowerValue"); + if (!button || !label || !effortLabelNode || !menu || !modelOptions || !effortOptions) return; + const current = currentModelOption(); + if (!current) { + button.hidden = true; + return; + } + button.hidden = false; + const selectedEffort = state.currentEffort === null + ? "" + : state.currentEffort || current.defaultReasoningEffort || current.efforts?.[1] || current.efforts?.[0] || ""; + const effortLabel = selectedEffort ? t(EFFORT_LABELS[selectedEffort] || selectedEffort) : ""; + label.textContent = current.displayName || modelDisplayName(current.model); + effortLabelNode.textContent = effortLabel; + effortLabelNode.hidden = !effortLabel; + const settingsModelValue = $("settingsModelValue"); + if (settingsModelValue) settingsModelValue.textContent = [label.textContent, effortLabel].filter(Boolean).join(" ") || t("默认"); + const triggerLabel = [label.textContent, effortLabel].filter(Boolean).join(" "); + button.setAttribute("aria-label", uiLocale() === "en-US" + ? `Current model: ${triggerLabel}; change model` + : `当前模型 ${triggerLabel},切换模型`); + button.setAttribute("title", uiLocale() === "en-US" + ? `Change model (current: ${triggerLabel})` + : `切换模型(当前 ${triggerLabel})`); + modelOptions.replaceChildren(); + for (const option of normalizedModelOptions()) { + const item = document.createElement("button"); + item.type = "button"; + item.className = "model-option"; + item.setAttribute("role", "option"); + item.setAttribute("aria-selected", String(option.model === state.currentModel)); + item.dataset.model = option.model; + const name = document.createElement("span"); + name.className = "model-option-name"; + name.textContent = option.displayName || option.model; + const check = document.createElement("span"); + check.className = "model-option-check"; + check.textContent = "✓"; + const description = document.createElement("span"); + description.className = "model-option-description"; + description.textContent = option.description ? t(option.description) : option.model; + item.append(name, check, description); + item.addEventListener("click", () => selectModel(option.model)); + modelOptions.append(item); + } + effortOptions.replaceChildren(); + const efforts = current.efforts?.length ? [...current.efforts] : []; + // Ultra is capability-gated in the official picker. Preserve it when the + // host explicitly reports the current setting, but do not advertise it to + // every fallback session. + if (state.currentEffort === "ultra" && current.model === "gpt-5.6-sol" && !efforts.includes("ultra")) efforts.push("ultra"); + const selectedMenuEffort = state.currentEffort === null + ? "" + : state.currentEffort || current.defaultReasoningEffort || efforts[1] || efforts[0] || ""; + if (!efforts.length) { + const empty = document.createElement("div"); + empty.className = "effort-empty"; + empty.textContent = t("此模型使用默认推理强度"); + effortOptions.append(empty); + } + for (const effort of efforts) { + const item = document.createElement("button"); + item.type = "button"; + item.className = "effort-option"; + item.setAttribute("role", "option"); + item.setAttribute("aria-selected", String(effort === selectedMenuEffort)); + item.dataset.effort = effort; + const name = document.createElement("span"); + name.textContent = t(EFFORT_LABELS[effort] || effort); + const check = document.createElement("span"); + check.className = "effort-option-check"; + check.textContent = "✓"; + item.append(name, check); + item.title = effort; + item.addEventListener("click", () => selectEffort(effort)); + effortOptions.append(item); + } + const powerOptions = modelPowerOptions(current); + if (powerView && advancedView) { + powerView.hidden = state.modelAdvancedOpen; + advancedView.hidden = !state.modelAdvancedOpen; + } + if (advancedToggle) { + advancedToggle.textContent = t(state.modelAdvancedOpen ? "简洁" : "高级"); + advancedToggle.setAttribute("aria-label", t(state.modelAdvancedOpen ? "返回简洁模型选择" : "显示高级模型选项")); + } + if (powerSlider) { + const hasPower = powerOptions.length > 0; + powerSlider.disabled = !hasPower; + powerSlider.min = "0"; + powerSlider.max = String(Math.max(0, powerOptions.length - 1)); + const selectedPower = powerOptions.findIndex((entry) => entry.effort === selectedMenuEffort); + powerSlider.value = String(selectedPower >= 0 ? selectedPower : Math.min(1, Math.max(0, powerOptions.length - 1))); + powerSlider.setAttribute("aria-valuetext", powerOptions[Number(powerSlider.value)]?.label || t("默认")); + powerSlider.style.setProperty("--power-position", powerOptions.length > 1 + ? `${Number(powerSlider.value) / (powerOptions.length - 1) * 100}%` + : "0%"); + if (powerValue) powerValue.textContent = powerOptions[Number(powerSlider.value)]?.label || t("默认"); + } else if (powerValue) powerValue.textContent = ""; + $("modelPicker")?.toggleAttribute("data-pending", state.modelUpdatePending); + renderControlMode(); + } + + function setModelMenu(open) { + const button = $("modelPickerButton"); + const menu = $("modelMenu"); + if (!button || !menu) return; + const next = Boolean(open) && threadSettingsAllowed() && !state.modeSwitching && !state.sessionSwitching; + menu.hidden = !next; + button.setAttribute("aria-expanded", String(next)); + if (next) { + state.modelAdvancedOpen = false; + renderModelPicker(); + } + } + + function setPermissionMenu(open) { + const button = $("permissionChip"); + const menu = $("permissionMenu"); + if (!button || !menu) return; + const next = Boolean(open) && threadSettingsAllowed() && !state.modeSwitching && !state.sessionSwitching; + menu.hidden = !next; + button.setAttribute("aria-expanded", String(next)); + if (next) renderPermissionMenu(); + } + + function permissionSandboxLabel(value) { + const label = { + "read-only": "只读", + "workspace-write": "工作区写入", + "danger-full-access": "完全访问", + }[normalizeSandboxName(value)] || String(value || "工作区写入"); + return t(label); + } + + function renderPermissionMenu() { + const menu = $("permissionMenu"); + const chip = $("permissionChip"); + const label = $("permissionLabel"); + if (!menu || !chip || !label) return; + const mode = state.sandboxPolicy === "danger-full-access" && state.approvalPolicy === "never" + ? "full" + : state.sandboxPolicy === "read-only" + ? "readonly" + : state.approvalPolicy === "untrusted" || state.approvalPolicy === "never" ? "auto" : "ask"; + label.textContent = t({ + ask: "需要时询问", + auto: "由 Codex 审批", + full: "完全访问", + readonly: "只读", + }[mode]); + const settingsPermissionValue = $("settingsPermissionValue"); + if (settingsPermissionValue) settingsPermissionValue.textContent = label.textContent; + chip.setAttribute("aria-label", uiLocale() === "en-US" + ? `Change permissions; current: ${label.textContent}` + : `修改权限,当前为${label.textContent}`); + chip.setAttribute("title", uiLocale() === "en-US" + ? `Change permissions (current: ${label.textContent})` + : `修改权限(当前:${label.textContent})`); + menu.querySelectorAll("[data-permission-mode]").forEach((item) => { + item.setAttribute("aria-checked", String(item.dataset.permissionMode === mode)); + }); + } + + function selectPermissionMode(mode) { + if (!threadSettingsAllowed()) { + setConversationStatus("当前模式不支持修改会话设置", "warning"); + return; + } + if (mode === "custom") { + setPermissionMenu(false); + setConversationStatus("自定义权限由 config.toml 管理", "ready"); + return; + } + if (mode === "full" && !(state.sandboxPolicy === "danger-full-access" && state.approvalPolicy === "never")) { + const confirm = $("permissionConfirm"); + if (confirm) { + confirm.hidden = false; + confirm.dataset.pendingMode = mode; + setPermissionMenu(false); + $("permissionConfirmAccept")?.focus(); + } + return; + } + applyPermissionMode(mode); + } + + function applyPermissionMode(mode) { + if (!threadSettingsAllowed()) return; + const presets = { + ask: { sandboxPolicy: "workspace-write", approvalPolicy: "on-request" }, + auto: { sandboxPolicy: "workspace-write", approvalPolicy: "untrusted" }, + full: { sandboxPolicy: "danger-full-access", approvalPolicy: "never" }, + readonly: { sandboxPolicy: "read-only", approvalPolicy: "on-request" }, + }; + const preset = presets[mode]; + if (!preset) return; + state.sandboxPolicy = preset.sandboxPolicy; + state.approvalPolicy = preset.approvalPolicy; + const sandboxInput = $("sandboxInput"); + const approvalInput = $("approvalInput"); + if (sandboxInput && [...sandboxInput.options].some((option) => option.value === state.sandboxPolicy)) sandboxInput.value = state.sandboxPolicy; + if (approvalInput && [...approvalInput.options].some((option) => option.value === state.approvalPolicy)) approvalInput.value = state.approvalPolicy; + renderPermissionMenu(); + updateThreadSettings(undefined, undefined, { + sandboxPolicy: state.sandboxPolicy, + approvalPolicy: state.approvalPolicy, + permissions: state.sandboxPolicy === "read-only" + ? ":read-only" + : state.sandboxPolicy === "danger-full-access" ? ":danger-full-access" : ":workspace", + approvalsReviewer: "user", + }); + setPermissionMenu(false); + } + + function selectPermissionSetting(kind, value) { + if (!threadSettingsAllowed()) return; + if (!value) return; + if (kind === "sandbox") { + state.sandboxPolicy = normalizeSandboxName(value); + const input = $("sandboxInput"); + if (input && [...input.options].some((option) => option.value === state.sandboxPolicy)) input.value = state.sandboxPolicy; + } else { + state.approvalPolicy = String(value); + const input = $("approvalInput"); + if (input && [...input.options].some((option) => option.value === state.approvalPolicy)) input.value = state.approvalPolicy; + } + renderPermissionMenu(); + updateThreadSettings(undefined, undefined, { + sandboxPolicy: state.sandboxPolicy, + approvalPolicy: state.approvalPolicy, + permissions: state.sandboxPolicy === "read-only" + ? ":read-only" + : state.sandboxPolicy === "danger-full-access" ? ":danger-full-access" : ":workspace", + approvalsReviewer: "user", + }); + setPermissionMenu(false); + } + + function normalizeUsage(value) { + if (!isRecord(value)) return null; + const source = isRecord(value.contextWindow) ? value.contextWindow : isRecord(value.context) ? value.context : value; + const total = isRecord(value.total) ? value.total : isRecord(source.total) ? source.total : {}; + const last = isRecord(value.last) ? value.last : isRecord(source.last) ? source.last : {}; + const breakdownTotal = (entry) => { + if (!isRecord(entry)) return null; + const explicit = finiteNumber(entry.totalTokens, entry.total_tokens, entry.tokens, entry.used, entry.inputTokens); + if (explicit !== null) return explicit; + const parts = [ + finiteNumber(entry.inputTokens, entry.input_tokens), + finiteNumber(entry.cachedInputTokens, entry.cached_input_tokens), + finiteNumber(entry.cacheWriteInputTokens, entry.cache_write_input_tokens), + finiteNumber(entry.outputTokens, entry.output_tokens), + finiteNumber(entry.reasoningOutputTokens, entry.reasoning_output_tokens), + ].filter((entry) => entry !== null); + return parts.length ? parts.reduce((sum, entry) => sum + entry, 0) : null; + }; + const used = finiteNumber( + source.used, + source.usedTokens, + source.inputTokens, + source.input_tokens, + source.tokensUsed, + source.totalTokens, + value.usedTokens, + value.inputTokens, + breakdownTotal(last), + breakdownTotal(total), + ); + const limit = finiteNumber( + value.modelContextWindow, + value.model_context_window, + source.limit, + source.max, + source.maxTokens, + typeof source.contextWindow === "number" ? source.contextWindow : null, + value.limit, + value.maxTokens, + ); + const remaining = finiteNumber(source.remaining, source.remainingTokens, value.remainingTokens); + const percent = finiteNumber(source.percent, source.percentage, value.percent, value.percentage); + if (used === null && limit === null && remaining === null && percent === null) return null; + const computedPercent = percent !== null + ? Math.max(0, Math.min(100, percent)) + : used !== null && limit !== null && limit > 0 ? Math.max(0, Math.min(100, used / limit * 100)) + : used !== null && remaining !== null && used + remaining > 0 ? used / (used + remaining) * 100 : 0; + const clampedUsed = used !== null && limit !== null && limit > 0 ? Math.min(used, limit) : used; + const computedRemaining = remaining !== null + ? remaining + : clampedUsed !== null && limit !== null ? Math.max(limit - clampedUsed, 0) : null; + return { + used: clampedUsed, + limit, + remaining: computedRemaining, + percent: computedPercent, + totalTokens: breakdownTotal(total), + lastTokens: breakdownTotal(last), + }; + } + + function renderUsage() { + const picker = $("usagePicker"); + const button = $("usageButton"); + const label = $("usageLabel"); + const ring = $("usageRing"); + const summary = $("usageSummary"); + const details = $("usageDetails"); + const bar = $("usageMeterBar"); + if (!picker || !button || !label || !ring || !summary || !details || !bar) return; + const usage = normalizeUsage(state.tokenUsage); + picker.hidden = !usage; + if (!usage) return; + const percent = Math.round(usage.percent); + label.textContent = `${percent}%`; + ring.style.setProperty("--usage-percent", `${percent}%`); + ring.dataset.level = percent >= 90 ? "critical" : percent >= 70 ? "warning" : "normal"; + const remainingPercent = Math.max(0, 100 - percent); + button.title = uiLocale() === "en-US" + ? `Context used ${percent}% (${remainingPercent}% remaining)` + : `上下文已使用 ${percent}%(剩余 ${remainingPercent}%)`; + button.setAttribute("aria-label", button.title); + summary.textContent = usage.limit !== null + ? `${usage.used ?? 0} / ${usage.limit} tokens (${percent}%)` + : uiLocale() === "en-US" ? `${percent}% used` : `${percent}% 已使用`; + bar.style.width = `${percent}%`; + details.textContent = [ + usage.remaining !== null ? (uiLocale() === "en-US" ? `${usage.remaining} tokens remaining` : `剩余 ${usage.remaining} tokens`) : "", + usage.used !== null ? (uiLocale() === "en-US" ? `Current context: ${usage.used} tokens` : `当前上下文 ${usage.used} tokens`) : "", + usage.lastTokens !== null && usage.lastTokens !== usage.used ? (uiLocale() === "en-US" ? `Latest request: ${usage.lastTokens} tokens` : `最近请求 ${usage.lastTokens} tokens`) : "", + usage.totalTokens !== null && usage.totalTokens !== usage.used ? (uiLocale() === "en-US" ? `Total: ${usage.totalTokens} tokens` : `累计 ${usage.totalTokens} tokens`) : "", + ].filter(Boolean).join("\n"); + } + + function setUsageMenu(open) { + const button = $("usageButton"); + const menu = $("usageMenu"); + if (!button || !menu) return; + const next = Boolean(open); + menu.hidden = !next; + button.setAttribute("aria-expanded", String(next)); + } + + function updateThreadSettings(model, effort, extra = {}) { + if (!threadSettingsAllowed() || !state.threadId || state.role !== "operator") return; + const threadSettings = {}; + if (model) threadSettings.model = model; + if (effort !== undefined) threadSettings.effort = effort; + if (isRecord(extra)) { + for (const [key, value] of Object.entries(extra)) { + if (value !== undefined) threadSettings[key] = value; + } + } + if (!Object.keys(threadSettings).length) return; + state.modelUpdatePending = true; + if (model) state.currentModel = model; + if (effort !== undefined) state.currentEffort = effort === null ? null : effort || ""; + renderModelPicker(); + try { + command("thread/settings/update", { threadId: state.threadId, threadSettings }); + } catch (error) { + state.modelUpdatePending = false; + appendOutput(error.message || "无法更新模型设置", "error"); + renderModelPicker(); + } + } + + function currentEffortParams() { + if (state.currentEffort === null) return { effort: null }; + return state.currentEffort ? { effort: state.currentEffort } : {}; + } + + function selectModel(model) { + const option = normalizedModelOptions().find((entry) => entry.model === model); + const effort = option?.efforts?.length + ? option.efforts.includes(state.currentEffort) + ? state.currentEffort + : option.defaultReasoningEffort || option.efforts[1] || option.efforts[0] + : null; + updateThreadSettings(model, effort); + setModelMenu(false); + } + + function selectEffort(effort) { + updateThreadSettings(undefined, effort); + setModelMenu(false); + } + + function selectPowerIndex(value) { + const current = currentModelOption(); + const options = modelPowerOptions(current); + if (!options.length) return; + const index = Math.max(0, Math.min(options.length - 1, Number(value) || 0)); + const option = options[index]; + if (!option) return; + updateThreadSettings(undefined, option.effort); + // The official power control remains open while keyboard arrows adjust it. + state.modelAdvancedOpen = false; + renderModelPicker(); + } + + function normalizeSubagentStatus(value) { + const raw = isRecord(value) ? firstString(value.status, value.type, value.state) : value; + const normalized = String(raw || "").replace(/[\s_-]+/g, "").toLowerCase(); + if (["running", "working", "interacted", "updated", "inprogress", "active", "started"].includes(normalized)) return "working"; + if (["pending", "pendinginit", "waiting", "queued", "waitingforinput", "awaitinginstruction", "waitingforinstruction", "needsinput"].includes(normalized)) return "waiting"; + if (["failed", "errored", "error", "notfound"].includes(normalized)) return "failed"; + if (["completed", "complete", "done", "interrupted", "shutdown", "cancelled", "canceled"].includes(normalized)) return "done"; + return normalized || "waiting"; + } + + function subagentStatusLabel(status) { + // These are the localized equivalents of the official background-agent + // rows ("is working", "is awaiting instruction", "is done"). + return t({ waiting: "正在等待指示", working: "正在工作", done: "已完成", failed: "失败" }[status] || status); + } + + function subagentThreadId(entry) { + return firstString(entry.threadId, entry.agentThreadId, entry.childThreadId); + } + + function subagentElapsed(entry) { + const startedAt = timestampMs(entry.startedAtMs, entry.startedAt, entry.createdAtMs, entry.createdAt); + if (startedAt === null) return ""; + const status = normalizeSubagentStatus(entry.status); + const completedAt = timestampMs(entry.completedAtMs, entry.completedAt, entry.finishedAtMs, entry.finishedAt, entry.lastAssistantMessageAtMs); + const end = status === "working" || status === "waiting" ? Date.now() : completedAt; + if (end === null || end === undefined) return ""; + return elapsedDuration(Math.max(0, end - startedAt)); + } + + function focusSubagentActivity(threadId) { + if (!threadId) return; + const output = $("output"); + const target = output + ? [...output.querySelectorAll(".message[data-agent-thread-id]")] + .reverse() + .find((article) => article.dataset.agentThreadId === threadId) + : null; + if (!target) return; + const turnId = target.dataset.turnId; + const turnActivity = turnId ? state.turnDividers.get(turnId) : null; + const turnToggle = turnActivity?.querySelector(".turn-divider-toggle"); + if (turnToggle?.getAttribute("aria-expanded") === "false") { + turnToggle.setAttribute("aria-expanded", "true"); + setTurnActivityVisibility(turnId, true); + } + target.scrollIntoView({ block: "center", behavior: "smooth" }); + target.classList.remove("subagent-focus"); + void target.offsetWidth; + target.classList.add("subagent-focus"); + setTimeout(() => target.classList.remove("subagent-focus"), 1_200); + } + + function createSubagentPanelRow(entry) { + const status = normalizeSubagentStatus(entry.status); + const threadId = subagentThreadId(entry); + const row = document.createElement(threadId ? "button" : "div"); + if (threadId) row.type = "button"; + row.className = "subagent-row"; + row.dataset.status = status; + if (threadId) row.dataset.agentThreadId = threadId; + const startedAt = timestampMs(entry.startedAtMs, entry.startedAt, entry.createdAtMs, entry.createdAt); + const completedAt = timestampMs( + entry.completedAtMs, + entry.completedAt, + entry.finishedAtMs, + entry.finishedAt, + entry.lastAssistantMessageAtMs, + entry.recencyAtMs, + entry.recencyAt, + ); + if (startedAt !== null) row.dataset.startedAt = String(startedAt); + if (completedAt !== null) row.dataset.completedAt = String(completedAt); + + const icon = document.createElement("span"); + icon.className = "subagent-icon"; + icon.setAttribute("aria-hidden", "true"); + icon.append(createSubagentIcon()); + const copy = document.createElement("span"); + copy.className = "subagent-copy"; + const name = document.createElement("span"); + name.className = "subagent-name"; + name.textContent = firstString(entry.displayName, entry.agentNickname, entry.name, entry.agentPath) || t("子代理"); + copy.append(name); + const stateLabel = document.createElement("span"); + stateLabel.className = "subagent-status"; + const statusText = document.createElement("span"); + statusText.className = "subagent-status-text"; + statusText.textContent = subagentStatusLabel(status); + stateLabel.append(statusText); + const diff = isRecord(entry.diffStats) ? entry.diffStats : isRecord(entry.diff_stats) ? entry.diff_stats : null; + const added = finiteNumber(diff?.linesAdded, diff?.added); + const removed = finiteNumber(diff?.linesRemoved, diff?.removed); + if (added !== null || removed !== null) { + const diffLabel = document.createElement("span"); + diffLabel.className = "subagent-diff-stats"; + diffLabel.textContent = `+${Math.max(0, added || 0)} -${Math.max(0, removed || 0)}`; + stateLabel.append(diffLabel); + } + const elapsed = document.createElement("span"); + elapsed.className = "subagent-elapsed"; + elapsed.textContent = subagentElapsed(entry); + // The composer keeps the row compact, but exposes timing in the tooltip + // and makes the elapsed node available for live updates. + if (elapsed.textContent) stateLabel.append(elapsed); + row.append(icon, copy, stateLabel); + const objective = firstString(entry.statusMessage, entry.objective, entry.prompt, entry.role, entry.agentRole); + const model = firstString(entry.spawnModel, entry.model, entry.modelId); + const tooltipParts = [objective, model ? (uiLocale() === "en-US" ? `Using ${model}` : `使用 ${model}`) : "", elapsed.textContent ? (uiLocale() === "en-US" ? `Elapsed: ${elapsed.textContent}` : `已用时 ${elapsed.textContent}`) : ""].filter(Boolean); + if (tooltipParts.length) row.title = `${name.textContent}: ${tooltipParts.join(" · ")}`; + if (threadId) { + row.title = tooltipParts.length + ? `${name.textContent}: ${tooltipParts.join(" · ")}` + : uiLocale() === "en-US" ? `View ${name.textContent}'s activity` : `查看 ${name.textContent} 的活动`; + row.setAttribute("aria-label", `${name.textContent}, ${subagentStatusLabel(status)}`); + row.addEventListener("click", () => focusSubagentActivity(threadId)); + } + return row; + } + + function deriveSubagentsFromMessages(messages) { + if (!Array.isArray(messages)) return []; + deriveStructuredTurnKeys(messages); + const entries = new Map(); + let anonymousIndex = 0; + const upsert = (item, threadId = "", itemIndex = -1) => { + const agent = isRecord(item.agent) ? item.agent : {}; + const displayName = firstString(item.displayName, item.agentNickname, item.agentName, item.agentPath, agent.displayName, agent.name); + const key = threadId || displayName || firstString(item.agentPath, item.action, item.activityKind) || `anonymous-${anonymousIndex++}`; + const existing = entries.get(key) || { threadId: threadId || "", displayName: null, prompt: null, objective: null, status: "working", statusMessage: null, canInteract: false }; + const activityKind = firstString(item.activityKind, item.kind, agent.activityKind); + const rawStatus = firstString(item.displayStatus, item.status, item.state, agent.status) + || (activityKind === "completed" ? "completed" : activityKind === "interrupted" ? "interrupted" : "working"); + const normalizedStatus = normalizeSubagentStatus(rawStatus); + existing.threadId = existing.threadId || threadId; + existing.displayName = displayName || existing.displayName; + existing.agentPath = firstString(item.agentPath, agent.agentPath, existing.agentPath); + existing.prompt = firstString(item.prompt, agent.prompt, existing.prompt) || null; + existing.objective = firstString(item.objective, item.statusMessage, item.prompt, agent.objective, existing.objective) || null; + existing.statusMessage = firstString(item.statusMessage, agent.statusMessage, existing.statusMessage) || null; + existing.status = normalizedStatus; + existing.canInteract = item.canInteract !== undefined ? item.canInteract !== false : existing.canInteract; + const startedAt = timestampMs(item.startedAtMs, item.startedAt, item.createdAtMs, agent.startedAtMs); + const completedAt = timestampMs( + item.completedAtMs, + item.completedAt, + item.finishedAtMs, + item.finishedAt, + item.lastAssistantMessageAtMs, + item.recencyAtMs, + item.recencyAt, + agent.completedAtMs, + agent.lastAssistantMessageAtMs, + agent.recencyAtMs, + ); + if (startedAt !== null) existing.startedAtMs = existing.startedAtMs ?? startedAt; + if (completedAt !== null) existing.completedAtMs = completedAt; + if (existing.completedAtMs === undefined && normalizedStatus === "done" && itemIndex >= 0) { + const turnKey = structuredMessageTurn(item, itemIndex); + for (let cursor = itemIndex + 1; cursor < messages.length; cursor += 1) { + const next = messages[cursor]; + if (!isRecord(next)) continue; + if (structuredMessageTurn(next, cursor) !== turnKey) break; + const nextStartedAt = timestampMs(next.startedAtMs, next.startedAt, next.createdAtMs); + if (nextStartedAt !== null) { + existing.completedAtMs = nextStartedAt; + break; + } + } + } + if (item.model !== undefined || agent.model !== undefined) existing.model = firstString(item.model, agent.model, existing.model) || null; + entries.set(key, existing); + }; + messages.forEach((item, index) => { + if (!isRecord(item) || historyKind(item) !== "subagent") return; + const receivers = Array.isArray(item.receiverThreadIds) + ? item.receiverThreadIds.map(String) + : Array.isArray(item.receiverThreads) ? item.receiverThreads.map(String) : []; + const states = isRecord(item.agentsStates) ? Object.keys(item.agentsStates) : []; + const ids = [...new Set([ + firstString(item.agentThreadId, item.childThreadId), + ...receivers, + ...states, + ].filter(Boolean))]; + if (!ids.length) upsert(item, "", index); + else ids.forEach((id) => upsert(item, id, index)); + }); + return [...entries.values()]; + } + + function subagentIdentity(entry) { + if (!isRecord(entry)) return ""; + return firstString( + entry.threadId, + entry.agentThreadId, + entry.childThreadId, + entry.displayName, + entry.agentNickname, + entry.name, + entry.agentPath, + ); + } + + // A snapshot can carry a rich top-level subagent projection while the + // message list only contains the compact activity item. Merge matching + // identities so timestamps, model names, and interaction flags survive the + // transcript re-render without carrying stale agents across threads. + function mergeSubagentProjections(derived, supplied) { + const rich = Array.isArray(supplied) ? supplied.filter(isRecord) : []; + if (!rich.length) return derived; + const byIdentity = new Map(rich + .map((entry) => [subagentIdentity(entry), entry]) + .filter(([identity]) => Boolean(identity))); + const matched = new Set(); + const merged = derived.map((entry) => { + const identity = subagentIdentity(entry); + const source = identity ? byIdentity.get(identity) : undefined; + if (!source) return entry; + matched.add(source); + return { ...entry, ...source }; + }); + // Keep a rich top-level entry when the transcript projection is absent + // (for example while a newly spawned agent has not emitted its first + // activity item yet), but discard unrelated entries from an older thread. + if (!merged.length) return rich; + for (const source of rich) { + if (matched.has(source)) continue; + const identity = subagentIdentity(source); + if (identity && !merged.some((entry) => subagentIdentity(entry) === identity) + && normalizeSubagentStatus(source.status) !== "done") merged.push(source); + } + return merged; + } + + function refreshSubagentElapsed() { + for (const row of document.querySelectorAll(".subagent-row[data-started-at]")) { + const startedAt = timestampMs(row.dataset.startedAt); + if (startedAt === null) continue; + const completedAt = timestampMs(row.dataset.completedAt); + const status = normalizeSubagentStatus(row.dataset.status); + const end = status === "working" || status === "waiting" ? Date.now() : completedAt; + const label = row.querySelector(".subagent-elapsed"); + if (label) label.textContent = end === null || end === undefined ? "" : elapsedDuration(Math.max(0, end - startedAt)); + } + } + + function renderSubagents() { + const panel = $("subagentsPanel"); + const list = $("subagentsList"); + const count = $("subagentsCount"); + if (!panel || !list || !count) return; + const entries = Array.isArray(state.subagents) ? state.subagents.filter(isRecord) : []; + const visibleEntries = entries.filter((entry) => firstString( + entry.displayName, + entry.agentNickname, + entry.name, + entry.agentPath, + entry.threadId, + entry.agentThreadId, + entry.childThreadId, + entry.objective, + )); + panel.hidden = visibleEntries.length === 0; + if (!visibleEntries.length) { + count.textContent = ""; + if (isRecord(state.subagentsExpanded)) { + state.subagentsExpanded.active = false; + state.subagentsExpanded.done = false; + } + list.replaceChildren(); + return; + } + const normalizedEntries = visibleEntries.map((entry) => ({ entry, status: normalizeSubagentStatus(entry.status) })); + // The official composer uses one compact disclosure row. It does not + // split the list into active/done sections or show per-agent wall-clock + // durations; status is rendered inline beside each display name. + const title = $("subagentsToggle")?.querySelector(".subagents-title"); + if (title) title.textContent = uiLocale() === "en-US" + ? `${normalizedEntries.length} background agents${state.subagentsCollapsed ? "" : " · @ to mention agents"}` + : `${normalizedEntries.length} 个后台代理${state.subagentsCollapsed ? "" : " · @ 可标记代理"}`; + count.textContent = ""; + panel.dataset.collapsed = String(state.subagentsCollapsed); + const toggle = $("subagentsToggle"); + if (toggle) toggle.setAttribute("aria-expanded", String(!state.subagentsCollapsed)); + list.setAttribute("aria-hidden", String(state.subagentsCollapsed)); + list.inert = state.subagentsCollapsed; + list.replaceChildren(); + const rows = document.createElement("div"); + rows.className = "subagent-section-rows"; + for (const { entry } of normalizedEntries) rows.append(createSubagentPanelRow(entry)); + list.append(rows); + } + + function normalizeSandboxName(value) { + const normalized = String(value || "") + .replace(/^:+/, "") + .replace(/([a-z0-9])([A-Z])/g, "$1-$2") + .replace(/[\s_]+/g, "-") + .toLowerCase(); + if (normalized === "dangerfullaccess" || normalized === "full-access" || normalized === "fullaccess") return "danger-full-access"; + if (normalized === "workspacewrite") return "workspace-write"; + if (normalized === "workspace" || normalized === "write") return "workspace-write"; + if (normalized === "readonly" || normalized === "read") return "read-only"; + return normalized; + } + + function sandboxFromPermissions(value) { + if (typeof value === "string") { + const normalized = normalizeSandboxName(value); + if (["read-only", "workspace-write", "danger-full-access"].includes(normalized)) return normalized; + return ""; + } + if (!isRecord(value)) return ""; + const explicit = firstString(value.sandboxPolicy, value.sandbox, value.mode, value.profile, value.type); + if (explicit) { + const normalized = normalizeSandboxName(explicit); + if (["read-only", "workspace-write", "danger-full-access"].includes(normalized)) return normalized; + } + if (value.dangerFullAccess === true || value.fullAccess === true || value.full_access === true) return "danger-full-access"; + const fileSystem = isRecord(value.fileSystem) ? value.fileSystem : isRecord(value.file_system) ? value.file_system : value; + if (fileSystem.write === true || fileSystem.workspaceWrite === true || fileSystem.workspace_write === true) return "workspace-write"; + if (fileSystem.readOnly === true || fileSystem.read_only === true || fileSystem.read === true) return "read-only"; + return ""; + } + + function normalizeApprovalPolicy(value) { + const candidate = isRecord(value) + ? firstString(value.policy, value.mode, value.type, value.approvalPolicy, value.approval_policy) + : firstString(value); + const normalized = String(candidate || "").replace(/[\s_]+/g, "-").toLowerCase(); + return ["on-request", "never", "untrusted"].includes(normalized) ? normalized : ""; + } + + // Model metadata belongs to the attached thread. Clear it before applying + // a snapshot for a different thread so an older catalog cannot leak into a + // newly selected conversation that does not expose model data. + function resetSessionModelMetadata() { + state.currentModel = ""; + state.currentEffort = ""; + state.availableModels = []; + state.modelUpdatePending = false; + state.tokenUsage = null; + state.workedDurationMs = null; + state.lastWorkedDurationMs = null; + state.turnWorkStartedAt = null; + state.finalAssistantStartedAt = null; + state.sandboxPolicy = "workspace-write"; + state.approvalPolicy = "on-request"; + const modelInput = $("modelInput"); + if (modelInput) modelInput.value = ""; + const label = $("modelLabel"); + if (label) { + label.textContent = ""; + label.hidden = true; + } + setModelMenu(false); + renderModelPicker(); + renderPermissionMenu(); + renderUsage(); + } + + function prepareForSessionSnapshot(threadId) { + const incoming = typeof threadId === "string" ? threadId : ""; + const previous = state.syncedThreadId !== null ? state.syncedThreadId : state.threadId || ""; + if (previous !== incoming && (previous || incoming)) { + resetSessionModelMetadata(); + state.turnExpansion.clear(); + } + if (incoming) syncSessionActive(incoming); + return previous !== incoming; + } + + function snapshotHistoryComplete(...sources) { + const seen = new Set(); + const visit = (value, depth = 0) => { + if (!isRecord(value) || seen.has(value) || depth > 3) return undefined; + seen.add(value); + if (typeof value.historyComplete === "boolean") return value.historyComplete; + for (const key of ["metadata", "sessionMetadata", "state", "snapshot"]) { + const nested = visit(value[key], depth + 1); + if (typeof nested === "boolean") return nested; + } + return undefined; + }; + for (const source of sources) { + const result = visit(source); + if (typeof result === "boolean") return result; + } + return undefined; + } + + function projectionTargetThreadId() { + const switchingTarget = firstString(state.sessionSwitchContext?.targetThreadId); + if (switchingTarget) return switchingTarget; + if (state.sessionSelectedThreadId && state.sessionSelectedThreadId !== state.syncedThreadId) { + return state.sessionSelectedThreadId; + } + return firstString(state.threadId, state.syncedThreadId); + } + + function outputProjectionAllowed(threadId) { + const incoming = String(threadId || ""); + const expected = projectionTargetThreadId(); + return !expected || !incoming || incoming === expected; + } + + function hasVisibleOutputProjection() { + const output = $("output"); + return Boolean( + state.structuredMessages.length + || output?.dataset.outputTail + || output?.querySelector(".message") + ); + } + + function finishSessionSnapshotCommit(threadId, authoritativeSnapshot = false) { + const incoming = String(threadId || ""); + if (!incoming) return; + const context = state.sessionSwitchContext; + if (authoritativeSnapshot && context?.targetThreadId === incoming) { + context.targetSnapshotReady = true; + finishSessionSwitchIfReady(); + } + syncSessionActive(incoming); + } + + // Transcript, active-thread identity, and switch completion are one commit. + // This prevents a metadata-only/placeholder reconnect snapshot from clearing + // a fully rendered thread or completing a switch before its history arrives. + function commitOutputProjection(threadId, text, structuredMessages, options = {}) { + const incoming = String(threadId || ""); + if (!outputProjectionAllowed(incoming)) return false; + const hasContent = Boolean( + (typeof text === "string" && text.length) + || (Array.isArray(structuredMessages) && structuredMessages.length) + ); + const historyComplete = options.historyComplete; + if (!hasContent && historyComplete !== true) { + // All relay control snapshots have projection-shaped placeholder fields. + // Until the host explicitly says history loading is complete, an empty + // projection is not authoritative—even when the current DOM is empty. + return false; + } + const changedThread = state.syncedThreadId !== null && state.syncedThreadId !== incoming; + prepareForSessionSnapshot(incoming); + if (changedThread) { + state.outputSynced = false; + state.snapshotNoticeShown = false; + } + replaceOutput(typeof text === "string" ? text : "", structuredMessages); + state.syncedThreadId = incoming; + if (incoming) state.threadId = incoming; + finishSessionSnapshotCommit(incoming, options.authoritativeSnapshot === true); + return true; + } + + function applySessionMetadata(metadata, snapshotState = {}) { + const meta = isRecord(metadata) ? metadata : {}; + const stateSnapshot = isRecord(snapshotState) ? snapshotState : {}; + const thread = isRecord(meta.thread) ? meta.thread : isRecord(stateSnapshot.thread) ? stateSnapshot.thread : {}; + const settings = isRecord(meta.threadSettings) + ? meta.threadSettings + : isRecord(meta.latestThreadSettings) + ? meta.latestThreadSettings + : isRecord(meta.settings) + ? meta.settings + : isRecord(stateSnapshot.threadSettings) ? stateSnapshot.threadSettings : {}; + const title = firstString(meta.title, meta.threadTitle, meta.name, thread.title, thread.name, thread.preview, stateSnapshot.title, stateSnapshot.threadTitle); + if (title) $("threadTitle").textContent = title; + + const cwd = firstString(meta.cwd, settings.cwd, stateSnapshot.cwd); + if (cwd) $("cwdInput").value = cwd; + const modelValue = meta.latestModel ?? meta.model ?? meta.modelName ?? meta.modelId + ?? settings.model ?? settings.modelName ?? stateSnapshot.latestModel ?? stateSnapshot.model ?? stateSnapshot.modelName; + const model = typeof modelValue === "object" && modelValue !== null + ? firstString(modelValue.name, modelValue.id, modelValue.slug) + : firstString(modelValue); + if (model) { + $("modelInput").value = model; + const modelLabel = $("modelLabel"); + const effortValue = firstDefined( + meta.latestReasoningEffort, + meta.effort, + settings.effort, + stateSnapshot.latestReasoningEffort, + stateSnapshot.effort, + ); + state.currentModel = model; + if (effortValue !== undefined) { + state.currentEffort = effortValue === null ? null : firstString(effortValue); + } + if (modelLabel) { modelLabel.textContent = modelDisplayName(model); modelLabel.hidden = false; } + } else if (modelValue !== undefined) { + const modelLabel = $("modelLabel"); + if (modelLabel) modelLabel.hidden = true; + } + const modelsValue = meta.availableModels ?? meta.models ?? stateSnapshot.availableModels ?? stateSnapshot.models; + if (Array.isArray(modelsValue)) state.availableModels = modelsValue.filter((entry) => typeof entry === "string" || isRecord(entry)); + const subagentsValue = meta.subagents ?? stateSnapshot.subagents; + if (Array.isArray(subagentsValue)) state.subagents = subagentsValue; + const usageValue = meta.tokenUsage ?? meta.latestTokenUsageInfo ?? meta.contextUsage ?? meta.usage + ?? stateSnapshot.tokenUsage ?? stateSnapshot.latestTokenUsageInfo ?? stateSnapshot.contextUsage ?? stateSnapshot.usage; + if (usageValue !== undefined) state.tokenUsage = usageValue; + const metadataWorkedDuration = finiteNumber( + meta.workedDurationMs, + meta.workDurationMs, + meta.workedForMs, + meta.workedFor?.durationMs, + settings.workedDurationMs, + stateSnapshot.workedDurationMs, + stateSnapshot.workDurationMs, + ); + if (metadataWorkedDuration !== null) { + state.workedDurationMs = Math.max(0, metadataWorkedDuration); + state.lastWorkedDurationMs = Math.max(0, metadataWorkedDuration); + } + const metadataWorkStart = timestampMs( + meta.firstTurnWorkItemStartedAtMs, + meta.firstWorkItemStartedAtMs, + meta.workStartedAtMs, + meta.workedFor?.startedAtMs, + settings.firstTurnWorkItemStartedAtMs, + stateSnapshot.firstTurnWorkItemStartedAtMs, + stateSnapshot.workStartedAtMs, + ); + if (metadataWorkStart !== null) state.turnWorkStartedAt = metadataWorkStart; + const metadataFinalStart = timestampMs( + meta.finalAssistantStartedAtMs, + meta.workedFor?.completedAtMs, + settings.finalAssistantStartedAtMs, + stateSnapshot.finalAssistantStartedAtMs, + ); + if (metadataFinalStart !== null) state.finalAssistantStartedAt = metadataFinalStart; + renderModelPicker(); + renderSubagents(); + renderUsage(); + + const permissionsValue = meta.permissions ?? meta.currentPermissions ?? settings.permissions + ?? stateSnapshot.permissions ?? stateSnapshot.currentPermissions; + const sandboxValue = meta.sandboxPolicy ?? meta.sandbox ?? settings.sandboxPolicy ?? settings.sandbox + ?? stateSnapshot.sandboxPolicy ?? stateSnapshot.sandbox ?? sandboxFromPermissions(permissionsValue); + const sandbox = typeof sandboxValue === "object" && sandboxValue !== null + ? firstString(sandboxValue.type, sandboxValue.mode, sandboxValue.policy) + : firstString(sandboxValue); + if (sandbox) { + const normalizedSandbox = normalizeSandboxName(sandbox); + state.sandboxPolicy = normalizedSandbox; + const sandboxInput = $("sandboxInput"); + if (sandboxInput && [...sandboxInput.options].some((option) => option.value === normalizedSandbox)) sandboxInput.value = normalizedSandbox; + const permission = $("permissionChip"); + const permissionLabel = $("permissionLabel"); + if (permission && permissionLabel) { + permissionLabel.textContent = permissionSandboxLabel(normalizedSandbox); + permission.hidden = false; + } + } else { + const permission = $("permissionChip"); + if (permission) permission.hidden = false; + } + const approval = normalizeApprovalPolicy( + meta.approvalPolicy !== undefined ? meta.approvalPolicy + : settings.approvalPolicy !== undefined ? settings.approvalPolicy + : stateSnapshot.approvalPolicy, + ); + if (approval) { + state.approvalPolicy = approval; + const approvalInput = $("approvalInput"); + if (approvalInput && [...approvalInput.options].some((option) => option.value === approval)) approvalInput.value = approval; + } + renderPermissionMenu(); + const mode = firstString(meta.mode, stateSnapshot.mode); + const modeLabel = $("modeLabel"); + if (modeLabel && mode) modeLabel.textContent = t(/cloud|remote/i.test(mode) ? "云端模式" : "本地模式"); + const popoverCwd = $("popoverCwd"); + const popoverMode = $("popoverMode"); + if (popoverCwd && cwd) popoverCwd.textContent = cwd; + if (popoverMode && mode) popoverMode.textContent = t(/cloud|remote/i.test(mode) ? "云端模式" : "本地模式"); + } + + const setConnection = (kind, text) => { + $("connectionDot").className = `dot ${kind}`; + $("connectionText").textContent = t(text); + if (embeddedInAether) embedBridge.reportState(kind, { message: String(text || "") }); + }; + + function setAuthRequired(value) { + if (typeof value !== "boolean") return; + state.authRequired = value; + document.body.classList.toggle("local-no-auth", !value); + const label = $("tokenLabel"); + const input = $("tokenInput"); + if (!label || !input) return; + label.textContent = t(value ? "访问 token(认证模式)" : "本机连接(无需 token)"); + input.placeholder = t(value ? "粘贴 relay 启动时打印的 token" : "本机模式无需填写;认证模式再填写"); + input.setAttribute("aria-label", t(value ? "relay access token" : "本地连接无需 token")); + } + + const authHeaders = () => ({ Authorization: `Bearer ${state.token}`, "Content-Type": "application/json" }); + + function attachPendingUserToTurn(turnId) { + const key = String(turnId || ""); + if (!key) return; + const article = state.pendingUserArticle; + if (article && article.isConnected) article.dataset.turnId = key; + state.pendingUserArticle = null; + } + + function sendFrame(frame) { + if (!state.ws || state.ws.readyState !== WebSocket.OPEN) throw new Error(t("WebSocket 未连接")); + state.ws.send(JSON.stringify(frame)); + } + + function command(method, params) { + const commandId = `web-${crypto.randomUUID()}`; + sendFrame({ type: "command", commandId, method, params }); + if (method === "turn/start" || method === "turn/steer") { + const text = (params?.input || []) + .map((item) => typeof item === "string" ? item : item?.text) + .filter(Boolean) + .join("\n"); + finishAssistantStream(); + state.pendingUserText = text; + appendDateSeparator(Date.now(), { role: "user", turnId: state.turnId || params?.expectedTurnId || "", turnStart: true }); + const userMessage = appendMessage(text || "(空消息)", "user", "text", "", { + kind: "user", + turnId: state.turnId || params?.expectedTurnId || "", + timestamp: Date.now(), + showTimestamp: true, + }); + if (method === "turn/start") state.pendingUserArticle = userMessage?.article || null; + } else if (method === "thread/settings/update") { + // Keep settings changes quiet in the transcript. The model picker has + // its own pending state and the authoritative snapshot will update the + // label once the official follower accepts the request. + setConversationStatus("正在更新模型设置", "active"); + } else if (["session/list", "session/select", "control/mode/set"].includes(sessionCommandMethod(method))) { + // Session navigation is shell UI state. It must not appear as a Codex + // message or alter the current turn's activity timeline. + } else { + appendOutput(`${method} (${commandId})`, "meta"); + } + return commandId; + } + + function composerText() { + const editor = $("messageInput"); + if (!editor) return ""; + return String(editor.innerText || editor.textContent || "") + .replace(/\u00a0/g, " ") + .replace(/\n{3,}/g, "\n\n") + .trim(); + } + + function clearComposer() { + const editor = $("messageInput"); + if (!editor) return; + editor.replaceChildren(); + editor.style.height = ""; + resizeComposer(); + updateIds(); + } + + function resizeComposer() { + const editor = $("messageInput"); + if (!editor) return; + editor.style.height = "auto"; + const maxHeight = Math.max(40, Math.round(window.innerHeight * 0.25)); + const nextHeight = Math.min(Math.max(editor.scrollHeight, 40), maxHeight); + editor.style.height = `${nextHeight}px`; + updateScrollPadding(); + } + + function defaultResponse(request) { + const method = request.method; + // Rendering a request must never imply consent. The operator still has + // to press an explicit action button. + if (method === "item/commandExecution/requestApproval") return { decision: "decline" }; + if (method === "item/fileChange/requestApproval") return { decision: "decline" }; + if (method === "item/permissions/requestApproval") { + return normalizePermissionResponse({}); + } + if (method === "applyPatchApproval" || method === "execCommandApproval") return { decision: { denied: { rejection: "默认拒绝,请明确允许" } } }; + if (method === "item/tool/requestUserInput") { + const answers = {}; + for (const question of request.params?.questions || []) answers[question.id] = { answers: [""] }; + return { answers }; + } + if (method === "mcpServer/elicitation/request") return { action: "decline", content: null, _meta: null }; + return {}; + } + + function normalizePermissionResponse(rawPermissions, scope, strictAutoReview) { + const permissions = {}; + if (rawPermissions && typeof rawPermissions === "object" && !Array.isArray(rawPermissions)) { + for (const [key, value] of Object.entries(rawPermissions)) { + if (value && typeof value === "object" && !Array.isArray(value)) permissions[key] = value; + } + } + const response = { permissions, scope: scope === "session" ? "session" : "turn" }; + if (typeof strictAutoReview === "boolean") response.strictAutoReview = strictAutoReview; + return response; + } + + function requestSummary(request) { + const params = request.params || {}; + if (Array.isArray(params.commandActions)) { + const commands = params.commandActions + .map((action) => action && typeof action === "object" ? action.command || action.description : "") + .filter(Boolean); + if (commands.length) return commands.join("\n"); + } + if (request.summary) return request.summary; + if (Array.isArray(params.command)) return params.command.join(" "); + if (params.command) return params.command; + if (params.reason) return params.reason; + if (Array.isArray(params.questions)) return params.questions.map((q) => q.question).join(" / "); + if (params.message) return params.message; + return t("需要远程确认或输入"); + } + + function requestTitle(request) { + const method = String(request.method || ""); + if (method === "item/commandExecution/requestApproval" || method === "execCommandApproval") return t("允许运行命令?"); + if (method === "item/fileChange/requestApproval" || method === "applyPatchApproval") return t("允许修改文件?"); + if (method === "item/permissions/requestApproval") return t("需要扩大权限"); + if (method === "item/tool/requestUserInput") return t("Codex 需要你的回答"); + if (method === "mcpServer/elicitation/request") return t("需要外部服务确认"); + return t("Codex 请求确认"); + } + + function requestRisk(request) { + const risk = String(request.risk || "medium").toLowerCase(); + return t(risk === "high" ? "高风险" : risk === "low" ? "低风险" : "需确认"); + } + + function requestCommand(request) { + const params = request.params || {}; + if (Array.isArray(params.commandActions)) { + return params.commandActions + .map((action) => action && typeof action === "object" ? action.command || action.description : "") + .filter(Boolean) + .join("\n"); + } + if (Array.isArray(params.command)) return params.command.join(" "); + return typeof params.command === "string" ? params.command : ""; + } + + function questionOptions(question) { + const options = question?.options || question?.choices || question?.enum; + if (!Array.isArray(options)) return []; + return options.map((option) => { + if (typeof option === "string") return { label: option, value: option }; + if (option && typeof option === "object") { + const value = option.value ?? option.id ?? option.label ?? option.name; + const label = option.label ?? option.name ?? value; + return { label: String(label ?? ""), value: String(value ?? "") }; + } + return null; + }).filter((option) => option && option.value); + } + + function renderQuestionFields(container, request) { + const questions = Array.isArray(request.params?.questions) ? request.params.questions : []; + container.replaceChildren(); + if (!questions.length) { + container.hidden = true; + return; + } + container.hidden = false; + questions.forEach((question, index) => { + if (!question || typeof question !== "object") return; + const field = document.createElement("label"); + field.className = "request-question"; + field.dataset.questionId = String(question.id ?? question.key ?? index); + const prompt = document.createElement("span"); + const questionPrompt = question.question ?? question.prompt ?? question.label; + prompt.textContent = questionPrompt === undefined || questionPrompt === null ? t("请输入") : String(questionPrompt); + field.append(prompt); + const options = questionOptions(question); + let control; + if (options.length) { + control = document.createElement("select"); + options.forEach((option) => { + const item = document.createElement("option"); + item.value = option.value; + item.textContent = option.label; + control.append(item); + }); + } else { + control = document.createElement("input"); + control.type = question.secret ? "password" : "text"; + control.placeholder = String(question.placeholder ?? ""); + } + control.className = "request-answer"; + control.dataset.answerId = field.dataset.questionId; + field.append(control); + container.append(field); + }); + } + + function inputResponseFromCard(article) { + const answers = {}; + for (const field of article.querySelectorAll(".request-question")) { + const id = field.dataset.questionId; + const value = field.querySelector(".request-answer")?.value ?? ""; + answers[id] = { answers: [value] }; + } + return { answers }; + } + + function requestResponseFromCard(request, article, responseBox, action) { + if (request.method === "item/tool/requestUserInput" && action === "allow") { + // Prefer explicit edits in the JSON editor; otherwise collect the + // first-class answer controls rendered for each question. + if (responseBox.value !== article.dataset.defaultResponse) { + try { + const parsed = JSON.parse(responseBox.value); + if (parsed && typeof parsed === "object") return parsed; + } catch { + return null; + } + } + return inputResponseFromCard(article); + } + if (action === "allow") { + const response = allowResponse(request, responseBox.value); + const scope = article.querySelector(".request-scope")?.value; + if (request.method === "item/permissions/requestApproval" && response && typeof response === "object" && scope) response.scope = scope; + return response; + } + if (action === "deny") return denyResponse(request); + try { + const parsed = JSON.parse(responseBox.value); + return parsed && typeof parsed === "object" ? parsed : {}; + } catch { + return null; + } + } + + function buildRequestCard(request) { + const fragment = $("requestTemplate").content.cloneNode(true); + const article = fragment.querySelector(".request"); + const method = String(request.method || "unknown"); + const risk = requestRisk(request); + article.dataset.requestKey = requestKey(request.requestId); + article.dataset.risk = String(request.risk || "medium").toLowerCase(); + fragment.querySelector(".request-method").textContent = requestTitle(request); + fragment.querySelector(".request-id").textContent = `#${request.requestId}`; + const riskNode = fragment.querySelector(".request-risk"); + riskNode.textContent = risk; + riskNode.dataset.risk = article.dataset.risk; + fragment.querySelector(".request-summary").textContent = requestSummary(request); + const commandNode = fragment.querySelector(".request-command"); + const commandText = requestCommand(request); + commandNode.textContent = commandText; + commandNode.hidden = !commandText; + fragment.querySelector(".request-json").textContent = JSON.stringify(request.params || {}, null, 2); + const responseBox = fragment.querySelector(".request-response"); + responseBox.value = JSON.stringify(defaultResponse(request), null, 2); + article.dataset.defaultResponse = responseBox.value; + const scopeWrap = fragment.querySelector(".request-scope-wrap"); + if (request.method === "item/permissions/requestApproval") scopeWrap.hidden = false; + renderQuestionFields(fragment.querySelector(".request-questions"), request); + const allow = fragment.querySelector(".request-allow"); + const deny = fragment.querySelector(".request-deny"); + const send = fragment.querySelector(".request-send"); + allow.textContent = t(method === "item/tool/requestUserInput" ? "提交回答" : method === "item/permissions/requestApproval" ? "允许" : "允许一次"); + send.textContent = t("发送自定义响应"); + allow.addEventListener("click", () => { + const result = requestResponseFromCard(request, article, responseBox, "allow"); + if (result) respond(request.requestId, JSON.stringify(result)); + else appendOutput("自定义响应不是有效 JSON", "error"); + }); + deny.addEventListener("click", () => respond(request.requestId, JSON.stringify(requestResponseFromCard(request, article, responseBox, "deny")))); + send.addEventListener("click", () => { + const result = requestResponseFromCard(request, article, responseBox, "custom"); + if (result) respond(request.requestId, JSON.stringify(result)); + else appendOutput("响应不是有效 JSON", "error"); + }); + if (state.sessionSwitching || state.modeSwitching || state.role !== "operator" || !RESPONDABLE_METHODS.has(method) || state.responding.has(requestKey(request.requestId))) { + for (const button of fragment.querySelectorAll("button")) button.disabled = true; + responseBox.disabled = true; + for (const control of fragment.querySelectorAll("input, select")) control.disabled = true; + } + return fragment; + } + + function renderRequests() { + const container = $("inlineRequests"); + const legacyContainer = $("requests"); + container.replaceChildren(); + if (legacyContainer) legacyContainer.replaceChildren(); + const requests = [...state.requests.values()]; + $("requestCount").textContent = String(requests.length); + $("factRequests").textContent = String(requests.length); + renderControlMode(); + const panel = $("requestsPanel"); + if (!requests.length) { + container.className = "inline-requests empty"; + if (legacyContainer) { + legacyContainer.className = "requests empty"; + legacyContainer.textContent = t("暂无待处理请求"); + } + if (panel) panel.open = false; + return; + } + container.className = "inline-requests"; + for (const request of requests) { + const fragment = buildRequestCard(request); + container.append(fragment); + } + if (panel) panel.open = false; + } + + function denyResponse(request) { + if (request.method === "item/commandExecution/requestApproval" || request.method === "item/fileChange/requestApproval") return { decision: "decline" }; + if (request.method === "item/permissions/requestApproval") return normalizePermissionResponse({}); + if (request.method === "applyPatchApproval" || request.method === "execCommandApproval") return { decision: { denied: { rejection: "远程参与者拒绝" } } }; + if (request.method === "mcpServer/elicitation/request") return { action: "decline", content: null, _meta: null }; + if (request.method === "item/tool/requestUserInput") return { answers: {} }; + return { decision: "decline" }; + } + + function allowResponse(request, raw) { + const method = request.method; + if (method === "item/commandExecution/requestApproval" || method === "item/fileChange/requestApproval") return { decision: "accept" }; + if (method === "item/permissions/requestApproval") { + if (raw === undefined) return normalizePermissionResponse(request.params?.permissions); + try { + const parsed = JSON.parse(raw); + return parsed && typeof parsed === "object" + ? normalizePermissionResponse(parsed.permissions, parsed.scope, parsed.strictAutoReview) + : normalizePermissionResponse({}); + } catch { + return normalizePermissionResponse(request.params?.permissions); + } + } + if (method === "applyPatchApproval" || method === "execCommandApproval") return { decision: "approved" }; + if (method === "mcpServer/elicitation/request") return { action: "accept", content: null, _meta: null }; + // User-input requests need the operator's edited answers. Keep the JSON + // editor as the source of truth and fail closed if it is malformed. + try { + const parsed = JSON.parse(raw); + return parsed && typeof parsed === "object" ? parsed : { answers: {} }; + } catch { + return { answers: {} }; + } + } + + function respond(requestId, raw) { + if (state.sessionSwitching) return; + const key = requestKey(requestId); + if (state.responding.has(key)) return; + let result; + try { result = JSON.parse(raw); } catch { appendOutput("响应不是有效 JSON", "error"); return; } + try { sendFrame({ type: "respond", requestId, result }); } catch (error) { appendOutput(error.message, "error"); return; } + state.responding.add(key); + renderRequests(); + } + + function updateIds() { + $("threadId").textContent = state.threadId || "-"; + $("turnId").textContent = state.turnId || "-"; + const popoverThread = $("popoverThread"); + if (popoverThread) popoverThread.textContent = state.threadId || "-"; + const hasThread = Boolean(state.threadId); + const hasTurn = Boolean(state.turnId); + const hasComposerText = Boolean(composerText()); + const switching = Boolean(state.sessionSwitching || state.modeSwitching); + document.body.classList.toggle("turn-active", hasTurn); + $("startThreadButton").disabled = state.attachMode || !state.appReady || state.role !== "operator"; + const newSessionButton = $("newSessionButton"); + if (newSessionButton) { + newSessionButton.disabled = !sessionControlAllowed("sessionCreate") || switching || !state.appReady + || Boolean(state.newSessionCommandId) + || !["operator", "owner", "host"].includes(String(state.role || "")); + newSessionButton.setAttribute("aria-busy", String(Boolean(state.newSessionCommandId))); + } + $("startTurnButton").disabled = switching || !state.appReady || !hasThread || hasTurn || !hasComposerText || state.role !== "operator"; + $("steerButton").disabled = switching || !state.appReady || !hasTurn || !hasComposerText || state.role !== "operator"; + $("interruptButton").disabled = switching || !state.appReady || !hasTurn || state.role !== "operator"; + const messageInput = $("messageInput"); + if (messageInput) { + messageInput.contentEditable = switching ? "false" : "true"; + messageInput.setAttribute("aria-disabled", String(switching)); + } + for (const id of ["modelPickerButton", "permissionChip", "composerPlusButton"]) { + const control = $(id); + if (control) control.disabled = switching; + } + const send = $("startTurnButton"); + if (send) { + send.title = t(hasTurn ? "发送 Steer" : "发送消息"); + send.setAttribute("aria-label", t(hasTurn ? "发送 Steer" : "发送消息")); + } + const settingsModelValue = $("settingsModelValue"); + if (settingsModelValue) settingsModelValue.textContent = state.currentModel ? modelDisplayName(state.currentModel, state.currentEffort) : t("默认"); + const settingsPermissionValue = $("settingsPermissionValue"); + if (settingsPermissionValue) settingsPermissionValue.textContent = permissionSandboxLabel(state.sandboxPolicy); + renderControlMode(); + } + + function protocolMethodForEvent(event, payload) { + if (typeof payload?.method === "string") return payload.method; + if (typeof payload?.params?.method === "string") return payload.params.method; + const type = String(event?.type || ""); + if (type === "item.started") return "item/started"; + if (type === "item.completed") return "item/completed"; + if (type === "thread.started") return "thread/started"; + if (type === "turn.started") return "turn/started"; + if (type === "turn.completed") return "turn/completed"; + if (type === "turn.plan.updated") return "turn/plan/updated"; + if (type === "turn.diff.updated") return "turn/diff/updated"; + return ""; + } + + function handleProtocolNotification(method, payload, eventType = "") { + const params = eventParams(payload); + if (method === "item/started" || eventType === "item.started") { + if (!eventBelongsToCurrentTurn(payload)) return true; + handleItemLifecycle("started", params); + return true; + } + if (method === "item/completed" || eventType === "item.completed") { + if (!eventBelongsToCurrentTurn(payload)) return true; + handleItemLifecycle("completed", params); + return true; + } + if (method === "turn/plan/updated" || method === "turn/plan/update" || eventType === "turn.plan.updated") { + if (!eventBelongsToCurrentTurn(payload)) return true; + handlePlanUpdate(params); + return true; + } + if (method === "turn/diff/updated" || method === "turn/diff/update" || eventType === "turn.diff.updated") { + if (!eventBelongsToCurrentTurn(payload)) return true; + handleDiffUpdate(params); + return true; + } + if (/^(?:item\/)?fileChange\/outputDelta$/.test(method) + || method === "item/fileChange/delta") { + if (!eventBelongsToCurrentTurn(payload)) return true; + appendFileChangeChunk(payload, textFromValue(params.delta ?? params.text ?? params.output ?? params.chunk)); + return true; + } + if (/^(?:item\/)?(?:fileRead|readFile|fileReadOutput)\/(?:outputDelta|delta|textDelta)$/.test(method) + || method === "item/fileRead/outputDelta" + || method === "item/fileRead/delta") { + if (!eventBelongsToCurrentTurn(payload)) return true; + appendOutputChunk(textFromValue(params.delta ?? params.text ?? params.output ?? params.chunk ?? params.content), "read", { + kind: "read", + itemId: params.itemId, + turnId: eventTurnId(payload), + }); + return true; + } + if (method === "item/commandExecution/outputDelta" + || method === "command/exec/outputDelta" + || method === "process/outputDelta") { + if (!eventBelongsToCurrentTurn(payload)) return true; + const stream = params.stream || params.channel || (params.stderr ? "stderr" : "stdout"); + appendOutputChunk(textFromValue(params.delta ?? params.text ?? params.output ?? params.chunk), stream, { + kind: "tool", + itemId: params.itemId, + turnId: eventTurnId(payload), + }); + return true; + } + if (method === "item/reasoning/summaryTextDelta" + || method === "item/reasoning/textDelta" + || method === "item/plan/delta") { + if (!eventBelongsToCurrentTurn(payload)) return true; + appendOutputChunk(textFromValue(params.delta ?? params.text), "reasoning", { + kind: method === "item/plan/delta" ? "plan" : "reasoning", + itemId: params.itemId, + turnId: eventTurnId(payload), + }); + return true; + } + if (method === "turn/started") { + const turn = isRecord(params.turn) ? params.turn : params; + const startedTurnId = turn.id || params.turnId; + const workStart = timestampMs( + turn.firstTurnWorkItemStartedAtMs, + turn.workStartedAtMs, + params.firstTurnWorkItemStartedAtMs, + params.workStartedAtMs, + ); + if (workStart !== null) state.turnWorkStartedAt = workStart; + startTurnClock(startedTurnId, turn.startedAtMs || params.startedAtMs, turn.elapsedMs || params.elapsedMs); + attachPendingUserToTurn(startedTurnId); + return true; + } + if (method === "turn/completed") { + const turn = isRecord(params.turn) ? params.turn : params; + const turnId = turn.id || params.turnId || ""; + if (!eventBelongsToCurrentTurn(payload)) return true; + const status = normalizeActivityStatus(turn.status || params.status, "completed"); + const duration = turn.durationMs || params.durationMs; + const worked = workedDurationFor(turn, workedDurationFor(params, null)); + const finalAssistantStart = timestampMs(turn.finalAssistantStartedAtMs, params.finalAssistantStartedAtMs); + if (finalAssistantStart !== null) state.finalAssistantStartedAt = finalAssistantStart; + finishActivitiesForTurn(turnId, status); + finishAssistantStream(); + const completedTurnId = turnId || state.turnId || ""; + stopTurnClock(status, duration, turn.completedAtMs || params.completedAtMs, worked); + appendCompletedTurnDivider(completedTurnId, status, worked ?? state.lastWorkedDurationMs ?? duration ?? state.lastTurnDurationMs); + reconcileTurnDividers(); + if (completedTurnId) { + state.retiredTurnIds.add(completedTurnId); + if (state.retiredTurnIds.size > 100) state.retiredTurnIds.delete(state.retiredTurnIds.values().next().value); + } + state.pendingUserText = ""; + state.turnId = ""; + updateIds(); + return true; + } + return false; + } + + function handleEvent(event) { + const eventSeq = finiteNumber(event.seq); + // The relay sends a buffered event stream followed by one control + // `session.snapshot`. Treat that control frame as the only baseline during + // the handshake; an older host `session.snapshot` event in the buffer is + // just historical data and may describe a turn that already ended. + if (state.awaitingSnapshot) return; + // `subscribe` may replay events that are already represented by the + // following authoritative snapshot. Ignore those frames entirely so a + // stale task.started/task.status cannot reopen the composer state. + if (eventSeq !== null && state.lastSnapshotSeq > 0 && eventSeq <= state.lastSnapshotSeq) return; + if (eventSeq !== null && eventSeq > state.lastSeq) state.lastSeq = eventSeq; + $("latestSeq").textContent = `seq ${state.lastSeq}`; + $("lastEvent").textContent = `${event.seq || "-"} / ${event.type || "event"}`; + const payload = isRecord(event.payload) ? { ...event.payload } : {}; + // Older bridge versions put identity/status fields on the event envelope + // instead of inside payload. Normalize both shapes before routing so a + // terminal event cannot be attributed to the wrong turn. + for (const key of [ + "threadId", "turnId", "requestId", "method", "params", "status", "executionStatus", "activity", "turnStatus", "activeFlags", + "controlMode", "mode", "targetMode", "modeEpoch", "capabilities", + "startedAtMs", "durationMs", "completedAtMs", "workedDurationMs", "workDurationMs", "workedForMs", + "firstTurnWorkItemStartedAtMs", "firstWorkItemStartedAtMs", "workStartedAtMs", "finalAssistantStartedAtMs", + ]) { + if (payload[key] === undefined && event[key] !== undefined) payload[key] = event[key]; + } + const usagePayload = payload.tokenUsage ?? payload.latestTokenUsageInfo ?? payload.contextUsage ?? payload.usage + ?? payload.params?.tokenUsage ?? payload.params?.latestTokenUsageInfo ?? payload.params?.contextUsage; + if (usagePayload !== undefined) { + state.tokenUsage = usagePayload; + renderUsage(); + } + const protocolMethod = protocolMethodForEvent(event, payload); + // Do not let a replayed terminal event from an older turn overwrite the + // timer/status of a newer turn. The protocol handler performs the same + // check, but status is normally applied before routing the event. + const staleTerminalEvent = (event.type === "task.finished" + || event.type === "task.cancelled" + || protocolMethod === "turn/completed") + && !eventBelongsToCurrentTurn(payload); + // Status is carried both as a typed relay field and inside the payload for + // older clients. Apply it before routing the event so a normal output + // delta cannot hide an active thinking/editing/approval state. + const hasExecutionStatus = isRecord(payload.executionStatus) + || isRecord(event.status) + || isRecord(payload.status) + || typeof payload.activity === "string" + || typeof payload.turnStatus === "string" + || Array.isArray(payload.activeFlags) + || payload.startedAtMs !== undefined + || payload.durationMs !== undefined; + if (hasExecutionStatus && !staleTerminalEvent) { + applyStatusSnapshot({ + ...payload, + status: payload.executionStatus || event.status || payload.status, + }, { allowTerminal: true, showIdle: false }); + } + if (protocolMethod && handleProtocolNotification(protocolMethod, payload, event.type)) { + if (event.type === "item.started" + || event.type === "item.completed" + || event.type === "task.status" + || protocolMethod === "turn/completed") updateIds(); + return; + } + if (event.type === "control.mode.switching" || event.type === "control.mode.changed") { + const requestedMode = normalizeControlMode(firstString(payload.controlMode, payload.mode, payload.targetMode)); + if (requestedMode && (requestedMode !== state.controlMode || state.modeSwitching)) { + state.modeSwitching = true; + state.requestedControlMode = requestedMode; + if (state.modeRequestEpoch < 0) state.modeRequestEpoch = state.modeEpoch; + setConversationStatus("正在切换控制模式", "active"); + updateIds(); + } + return; + } + if (event.type === "session.switching") { + state.sessionSwitching = true; + const targetThreadId = firstString(payload.targetThreadId, payload.threadId); + let switchContext = state.sessionSwitchContext; + if (targetThreadId) { + state.sessionSelectedThreadId = targetThreadId; + switchContext = beginSessionSwitchContext(targetThreadId); + // Route subsequent target events against the requested owner while + // the previous transcript remains mounted as the visual fallback. + state.threadId = targetThreadId; + if (switchContext?.targetTitle) { + setConversationStatus( + uiLocale() === "en-US" + ? `Switching to “${switchContext.targetTitle}”` + : `正在切换到「${switchContext.targetTitle}」`, + "active", + ); + } + } + finishAssistantStream(); + // Keep the previous transcript mounted until the target's authoritative + // snapshot arrives. The bridge may emit `session.switching` before + // owner discovery/follow completes; clearing here made a failed switch + // look like an empty conversation and left the user with no way to tell + // whether the target had actually loaded. + setSessionSwitchingVisual(true); + if (!switchContext?.targetTitle) setConversationStatus("正在切换会话", "active"); + renderSessionPicker(); + return; + } + if (event.type === "session.selected") { + const selectedThreadId = firstString(payload.threadId, payload.activeThreadId); + const switchContext = state.sessionSwitchContext; + if (payload.failed === true) { + // VS Code-driven attachment changes have no browser command result. + // The adapter therefore publishes an explicit failed selection for + // the previous thread. Roll routing/title back even if a target + // snapshot was already rendered while the adapter validated its owner. + const previousThreadId = firstString(switchContext?.previousThreadId, selectedThreadId); + restoreSessionSwitchContext(); + if (previousThreadId) { + state.threadId = previousThreadId; + state.sessionSelectedThreadId = previousThreadId; + syncSessionActive(previousThreadId); + } + state.sessionListError = sessionErrorMessage(payload, "会话切换失败,已恢复原会话"); + setConversationStatus(state.sessionListError, "warning"); + renderSessionPicker(); + return; + } + + const matchesTarget = Boolean(switchContext + && selectedThreadId + && switchContext.targetThreadId === selectedThreadId); + if (selectedThreadId && (!switchContext || matchesTarget)) { + state.sessionSelectedThreadId = selectedThreadId; + syncSessionActive(selectedThreadId); + } + // This is only one half of the switch commit. An acknowledgement for a + // superseded target is ignored, and a matching acknowledgement keeps all + // input disabled until the target's authoritative snapshot is committed. + if (matchesTarget) { + switchContext.selectedAckReady = true; + if (finishSessionSwitchIfReady()) { + setConversationStatus("会话已切换", "ready"); + } else { + state.sessionSwitching = true; + setSessionSwitchingVisual(true); + setConversationStatus("正在加载会话", "active"); + } + } else if (!switchContext) { + state.sessionSwitching = false; + finishSessionSwitchContext(); + } + renderSessionPicker(); + return; + } + if (/error|warning/i.test(String(event.type || "")) + || ["error", "warning"].includes(String(payload.method || "").toLowerCase())) { + finishAssistantStream(); + appendOutput(eventMessage(payload), /warning/i.test(String(event.type || "")) ? "meta" : "error"); + return; + } + if (event.type === "connection.opened" || event.type === "app.ready") { + state.appReady = true; + $("appState").textContent = appStatusLabel("ready"); + $("factApp").textContent = appStatusLabel("ready"); + } + if (event.type === "connection.closed" || event.type === "host.disconnected" || event.type === "app.exited") { + finishAssistantStream(); + const disconnectedTurnId = state.turnId; + if (disconnectedTurnId) finishActivitiesForTurn(disconnectedTurnId, "interrupted"); + state.pendingUserText = ""; + state.appReady = false; + state.snapshotNoticeShown = false; + state.turnId = ""; + state.turnStartedAt = null; + state.turnWorkStartedAt = null; + state.finalAssistantStartedAt = null; + state.workedDurationMs = null; + state.lastWorkedDurationMs = null; + state.currentActivity = "idle"; + state.currentActivityStartedAt = null; + state.currentActivityDurationMs = null; + state.currentActivityTurnId = ""; + state.subagents = []; + renderSubagents(); + updateLiveActivity("idle"); + state.sessionListLoading = false; + state.sessionListCommandId = ""; + state.newSessionCommandId = ""; + state.sessions = []; + state.sessionFocusedId = ""; + if (state.sessionSwitching) { + // An explicit host disconnect aborts an in-flight hand-off. Restore + // the previous thread identity while keeping its transcript mounted. + restoreSessionSwitchContext(); + } + state.sessionSelectedThreadId = ""; + state.sessionListError = "VS Code 主机未连接"; + renderSessionPicker(); + $("appState").textContent = appStatusLabel("offline"); + $("factApp").textContent = appStatusLabel("offline"); + updateIds(); + updateScrollToBottom($("output")); + } + if (event.type === "output.snapshot") { + const snapshotThreadId = firstString(payload.threadId, state.threadId, state.syncedThreadId); + if (!outputProjectionAllowed(snapshotThreadId)) return; + const committed = commitOutputProjection(snapshotThreadId, typeof payload.text === "string" ? payload.text : "", payload.messages, { + historyComplete: snapshotHistoryComplete(payload), + authoritativeSnapshot: true, + }); + if (!committed) return; + if (Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } + return; + } + if ((event.type === "output.delta" || event.type === "output.chunk") + && (payload.text || Array.isArray(payload.messages) || isRecord(payload.messagesPatch))) { + const projectionThreadId = firstString(payload.threadId, state.threadId, state.syncedThreadId); + if (!outputProjectionAllowed(projectionThreadId)) return; + if (!eventBelongsToCurrentTurn(payload)) return; + const projectionMatches = !projectionThreadId || projectionThreadId === state.syncedThreadId; + // Attach-mode adapters include the complete role-aware projection on a + // delta. Re-rendering that projection keeps reasoning, tools, edits and + // assistant text in their canonical item boundaries while retaining the + // legacy append-only text field for older bridges. + if (Array.isArray(payload.messages)) { + if (!projectionMatches) { + const committed = commitOutputProjection( + projectionThreadId, + typeof payload.outputTail === "string" ? payload.outputTail : typeof payload.text === "string" ? payload.text : "", + payload.messages, + { historyComplete: snapshotHistoryComplete(payload) }, + ); + if (committed && Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } + return; + } + if (Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } + if (payload.structureChanged === false && reconcileStructuredOutput( + typeof payload.outputTail === "string" ? payload.outputTail : payload.text, + payload.messages, + )) return; + replaceOutput( + typeof payload.outputTail === "string" ? payload.outputTail : typeof payload.text === "string" ? payload.text : "", + payload.messages, + ); + return; + } + if (isRecord(payload.messagesPatch)) { + // A suffix patch has meaning only relative to the same thread's + // authoritative baseline. Never apply it to the transcript retained + // while another session is still loading. + if (!projectionMatches) return; + if (Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } + const patchedMessages = applyStructuredMessagesPatch(payload.messagesPatch); + if (patchedMessages) { + const outputTail = typeof payload.outputTail === "string" + ? payload.outputTail + : typeof payload.text === "string" + ? `${$("output")?.dataset.outputTail || ""}${payload.text}`.slice(-32_000) + : $("output")?.dataset.outputTail || ""; + if (reconcileStructuredOutput(outputTail, patchedMessages)) return; + replaceOutput(outputTail, patchedMessages); + return; + } + } + if (!projectionMatches) return; + if (Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } + appendOutputChunk(payload.text, payload.stream, { + kind: payload.kind, + itemId: payload.itemId, + turnId: eventTurnId(payload), + timestamp: payload.timestamp || payload.startedAtMs, + }); + return; + } + if (event.type === "app.stderr") { + finishAssistantStream(); + appendOutput(payload.text, "meta"); + return; + } + if (event.type === "approval.requested" || event.type === "input.requested" || event.type === "server.requested" || event.type === "server.request") { + state.requests.set(requestKey(payload.requestId), { + requestId: payload.requestId, + method: payload.method, + params: payload.params || payload, + ...(payload.risk ? { risk: payload.risk } : {}), + ...(payload.summary ? { summary: payload.summary } : {}), + ...(payload.commandHash ? { commandHash: payload.commandHash } : {}), + ...(payload.createdAt ? { createdAt: payload.createdAt } : {}), + ...(payload.expiresAt ? { expiresAt: payload.expiresAt } : {}), + }); + renderRequests(); + const inlineRequests = $("inlineRequests"); + if (inlineRequests) inlineRequests.scrollTop = inlineRequests.scrollHeight; + finishAssistantStream(); + // The inline request card is the canonical representation. A separate + // plain-text transcript line would duplicate the approval/input prompt + // and could be mistaken for Codex output. + const waitingActivity = payload.method === "item/tool/requestUserInput" + || payload.method === "mcpServer/elicitation/request" + ? "waiting_input" + : "waiting_approval"; + state.currentActivity = waitingActivity; + setConversationStatus(statusActivityLabel(waitingActivity), "warning"); + return; + } + if (event.type === "server.responded" || event.type === "approval.resolved" || event.type === "input.resolved" || event.type === "approval.expired" || event.type === "input.expired") { + const key = resolveRequestKey(payload.requestId); + state.responding.delete(key); + state.requests.delete(key); + renderRequests(); + return; + } + if (event.type === "task.status") { + applyStatusSnapshot(payload, { allowTerminal: true, showIdle: false }); + updateIds(); + return; + } + if (event.type === "task.started") { + finishAssistantStream(); + if (payload.turnId) state.turnId = payload.turnId; + if (payload.threadId) state.threadId = payload.threadId; + attachPendingUserToTurn(payload.turnId || state.turnId); + applyStatusSnapshot(payload, { allowTerminal: false }); + if (!state.currentActivity || state.currentActivity === "idle") state.currentActivity = "running"; + updateLiveActivity(state.currentActivity, state.turnStartedAt, null, payload.activeFlags || [], payload.turnId || state.turnId); + updateIds(); + return; + } + if (event.type === "task.finished" || event.type === "task.cancelled") { + if (!eventBelongsToCurrentTurn(payload)) return; + finishAssistantStream(); + const reportedStatus = payload.turnStatus + ?? payload.status?.turnStatus + ?? payload.status?.status + ?? payload.status + ?? payload.executionStatus?.turnStatus + ?? payload.executionStatus?.status + ?? payload.executionStatus; + const finalStatus = event.type === "task.cancelled" + ? "interrupted" + : normalizeActivityStatus(reportedStatus, "completed"); + const finishedTurnId = eventTurnId(payload) || state.turnId || ""; + applyStatusSnapshot({ ...payload, turnStatus: finalStatus }, { allowTerminal: false }); + state.pendingUserText = ""; + const worked = workedDurationFor(payload, null); + stopTurnClock(finalStatus, payload.durationMs, payload.completedAtMs, worked); + appendCompletedTurnDivider(finishedTurnId, finalStatus, worked ?? state.lastWorkedDurationMs ?? payload.durationMs ?? state.lastTurnDurationMs); + reconcileTurnDividers(); + if (finishedTurnId) state.retiredTurnIds.add(finishedTurnId); + if (state.retiredTurnIds.size > 100) state.retiredTurnIds.delete(state.retiredTurnIds.values().next().value); + if (payload.threadId) state.threadId = payload.threadId; + state.pendingUserText = ""; + state.turnId = ""; + updateIds(); + return; + } + if (event.type === "session.created" && payload.thread?.id) { + state.threadId = payload.thread.id; + if (state.syncedThreadId === null || state.syncedThreadId === payload.thread.id) { + prepareForSessionSnapshot(payload.thread.id); + applySessionMetadata({ ...(isRecord(payload.metadata) ? payload.metadata : {}), ...payload }, payload); + } else { + setConversationStatus("正在加载会话", "active"); + } + updateIds(); + return; + } + if (event.type === "session.snapshot") { + state.awaitingSnapshot = false; + state.appReady = true; + const snapshotThreadId = payload.threadId || ""; + applyControlModeSnapshot(payload.metadata); + const waitingForSession = payload.state === "waiting_for_host" + || (isRecord(payload.metadata) && payload.metadata.waitingForSession === true); + const eventSnapshotSeq = finiteNumber(payload.latestSeq, event.seq); + if (eventSnapshotSeq !== null) state.lastSnapshotSeq = eventSnapshotSeq; + const adapterName = payload.metadata && payload.metadata.adapter; + if (adapterName === "codex-ipc-follower") setAttachMode(true); + else if (adapterName) setAttachMode(false); + const hasProjection = typeof payload.outputTail === "string" || Array.isArray(payload.messages); + const projectionCommitted = hasProjection && commitOutputProjection( + snapshotThreadId, + typeof payload.outputTail === "string" ? payload.outputTail : "", + payload.messages, + { + historyComplete: snapshotHistoryComplete(payload, payload.state), + authoritativeSnapshot: true, + }, + ); + const appliesToView = projectionCommitted || snapshotThreadId === state.syncedThreadId; + if (!appliesToView) { + // Keep the retained transcript and loading state until a matching, + // non-placeholder snapshot is available for the selected thread. + updateIds(); + return; + } + prepareForSessionSnapshot(snapshotThreadId); + applySessionMetadata({ ...(isRecord(payload.metadata) ? payload.metadata : {}), ...payload }, payload.state); + if (Array.isArray(payload.subagents)) { + state.subagents = payload.subagents; + renderSubagents(); + } else if (isRecord(payload.state) && Array.isArray(payload.state.subagents)) { + state.subagents = payload.state.subagents; + renderSubagents(); + } + if (payload.threadId !== undefined) state.threadId = payload.threadId || ""; + if (payload.turnId !== undefined) state.turnId = payload.turnId || ""; + applyStatusSnapshot({ + ...payload, + ...(payload.executionStatus ? { status: payload.executionStatus } : {}), + turnId: payload.turnId !== undefined ? payload.turnId : state.turnId, + }, { allowTerminal: true, showIdle: false }); + if (payload.metadata && typeof payload.metadata.cwd === "string") $("cwdInput").value = payload.metadata.cwd; + reconcileSnapshotTerminalState(payload, payload.state); + if (Array.isArray(payload.pendingRequests)) { + state.requests = new Map(payload.pendingRequests + .filter((request) => request && request.requestId !== undefined) + .map((request) => [requestKey(request.requestId), request])); + state.responding.clear(); + renderRequests(); + } + if (waitingForSession) { + setConversationStatus("等待在 VS Code 中打开 Codex 会话", "active"); + state.snapshotNoticeShown = false; + } else if (!state.snapshotNoticeShown && state.turnStartedAt === null && (!state.currentActivity || state.currentActivity === "idle")) { + setConversationStatus("ready"); + state.snapshotNoticeShown = true; + } + updateIds(); + return; + } + if (event.type === "session.closed") { + finishAssistantStream(); + state.pendingUserText = ""; + state.appReady = false; + state.sessions = []; + state.sessionFocusedId = ""; + state.sessionSelectedThreadId = ""; + state.threadId = ""; + state.turnId = ""; + const closedTitle = $("threadTitle"); + if (closedTitle) closedTitle.textContent = "Codex"; + state.turnStartedAt = null; + state.turnWorkStartedAt = null; + state.finalAssistantStartedAt = null; + state.workedDurationMs = null; + state.lastWorkedDurationMs = null; + state.currentActivity = "idle"; + state.currentActivityStartedAt = null; + state.subagents = []; + resetSessionModelMetadata(); + renderSubagents(); + updateLiveActivity("idle"); + state.sessionSwitching = false; + state.sessionSelectCommandId = ""; + finishSessionSwitchContext(); + setConversationStatus("会话已关闭", "warning"); + state.sessionListError = "VS Code 会话已关闭"; + renderSessionPicker(); + updateIds(); + return; + } + if (event.type === "app.notification") { + if (payload.method === "thread/started" && payload.params?.thread?.id) state.threadId = payload.params.thread.id; + if (payload.method === "turn/started" && payload.params?.turn?.id) { + state.turnId = payload.params.turn.id; + attachPendingUserToTurn(state.turnId); + } + if (payload.method === "turn/completed") { + state.turnId = ""; + finishAssistantStream(); + } + if (payload.text) { + finishAssistantStream(); + appendOutput(payload.text); + } + else if (["thread/started", "turn/started", "turn/completed", "thread/status/changed"].includes(payload.method)) appendOutput(`${payload.method}`, "meta"); + updateIds(); + return; + } + if (event.type === "thread.started" && payload.params?.thread?.id) { + state.threadId = payload.params.thread.id; + updateIds(); + return; + } + if (event.type === "turn.started" && payload.params?.turn?.id) { + state.turnId = payload.params.turn.id; + attachPendingUserToTurn(state.turnId); + updateIds(); + return; + } + if (event.type === "turn.completed") { + finishAssistantStream(); + const completedTurnId = eventTurnId(payload) || state.turnId || ""; + const status = normalizeActivityStatus(payload.status, "completed"); + const worked = workedDurationFor(payload, null); + stopTurnClock(status, payload.durationMs, payload.completedAtMs, worked); + appendCompletedTurnDivider(completedTurnId, status, worked ?? state.lastWorkedDurationMs ?? payload.durationMs ?? state.lastTurnDurationMs); + reconcileTurnDividers(); + state.turnId = ""; + updateIds(); + return; + } + if (event.type === "command.result") { + const commandId = payload.commandId; + // The relay sends a sequenced event to every subscriber and a direct + // acknowledgement to the originating browser. Render a result once. + if (commandId) { + const key = String(commandId); + if (state.commandResults.has(key)) return; + state.commandResults.add(key); + if (state.commandResults.size > 2_000) state.commandResults.delete(state.commandResults.values().next().value); + } + const commandMethodName = sessionCommandMethod(payload.method); + if (commandMethodName === "control/mode/set") { + state.modeCommandId = ""; + if (!payload.ok) { + clearControlModeRequest(); + setConversationStatus(sessionErrorMessage(payload, "控制模式切换失败"), "warning"); + } else if (state.modeSwitching) { + // The command result is only an acknowledgement. Keep the requested + // segment pending until a newer authoritative snapshot supplies the + // resulting modeEpoch and capabilities. + setConversationStatus("正在切换控制模式", "active"); + } + updateIds(); + return; + } + if (commandMethodName === "session/list" || commandMethodName === "thread/list") { + if (!payload.ok) { + state.sessionListLoading = false; + state.sessionListError = eventMessage(payload); + state.sessionListError = sessionErrorMessage(payload, state.sessionListError); + renderSessionPicker(); + } else { + applySessionListResult(payload.result ?? payload); + } + return; + } + if (commandMethodName === "session/select" || commandMethodName === "thread/select") { + if (!payload.ok) { + state.sessionListError = failSessionSwitch(payload, eventMessage(payload)); + renderSessionPicker(); + setConversationStatus(state.sessionListError, "warning"); + requestRefresh(); + } else { + applySessionSelectResult(payload.result ?? payload); + } + return; + } + if (commandMethodName === "session/new" || commandMethodName === "thread/new") { + state.newSessionCommandId = ""; + if (!payload.ok) { + const detail = eventMessage(payload); + setConversationStatus( + uiLocale() === "en-US" + ? `Unable to create a new conversation: ${detail}` + : `新会话创建失败:${detail}`, + "warning", + ); + } else { + setConversationStatus("新会话已在 VS Code 中打开", "ready"); + // The official command opens the new panel asynchronously. Give the + // host a short window to publish its rollout/owner, then refresh the + // same history menu so it can be selected without leaving the web UI. + openSessionHistory(); + let refreshAttempts = 0; + const refreshNewSession = () => { + if (!state.sessionPickerOpen || refreshAttempts >= 4) return; + refreshAttempts += 1; + if (!state.sessionListLoading) requestSessionList(); + setTimeout(refreshNewSession, 700); + }; + setTimeout(refreshNewSession, 500); + } + updateIds(); + return; + } + if (payload.method === "thread/settings/update") { + state.modelUpdatePending = false; + if (!payload.ok) { + const detail = eventMessage(payload); + appendOutput( + uiLocale() === "en-US" + ? `Unable to update model settings: ${detail}` + : `模型设置更新失败:${detail}`, + "error", + ); + } else setConversationStatus(t("模型设置已更新"), "ready"); + renderModelPicker(); + return; + } + if (!payload.ok) { + const methodLabel = payload.method || t("命令"); + const uncertainty = payload.uncertain + ? uiText("(执行状态未知,请等待主机恢复)", " (execution status unknown; wait for the host to recover)") + : ""; + appendOutput(`${methodLabel}: ${JSON.stringify(payload.error)}${uncertainty}`, "error"); + } + else { + const result = payload.result || {}; + if (payload.method === "thread/start" && result.thread?.id) state.threadId = result.thread.id; + if (payload.method === "turn/start" && result.turn?.id) state.turnId = result.turn.id; + if (payload.method === "turn/start" && result.turn?.id) attachPendingUserToTurn(result.turn.id); + if (payload.method === "thread/start") { + applySessionMetadata({ ...result, ...(isRecord(result.thread) ? result.thread : {}) }, result); + } + finishAssistantStream(); + appendOutput(`${payload.method || "命令"} 完成`, "meta"); + updateIds(); + } + } + } + + function handleMessage(message) { + if (message.type === "auth.ok") { + state.role = message.role; + setAuthRequired(message.authRequired); + $("roleBadge").textContent = message.role; + $("roleBadge").className = `badge ${message.role === "operator" ? "" : "warning"}`; + setConnection("pending", "同步中"); + state.awaitingSnapshot = true; + sendFrame({ type: "subscribe", fromSeq: state.lastSeq }); + return; + } + // A host event can legitimately have the same type as a relay control + // frame (notably `session.snapshot`). Route the envelope by `kind` first + // so its payload is not mistaken for the compact control shape below. + if (message.kind === "event") { + handleEvent(message); + return; + } + if (message.type === "session.snapshot") { + state.awaitingSnapshot = false; + const snapshot = message.snapshot || {}; + const appState = snapshot.state || {}; + const snapshotThreadId = appState.activeThreadId || ""; + const controlMetadata = { + ...(isRecord(snapshot.metadata) ? snapshot.metadata : {}), + ...(isRecord(appState.sessionMetadata) ? appState.sessionMetadata : {}), + }; + applyControlModeSnapshot(controlMetadata); + const waitingForSession = controlMetadata.waitingForSession === true + || (!snapshotThreadId && controlMetadata.attachReady === false); + $("appState").textContent = appStatusLabel(appState.app); + $("factApp").textContent = appState.app ? appStatusLabel(appState.app) : "-"; + $("factClients").textContent = String((snapshot.clients || []).length); + state.appReady = appState.app === "ready" || appState.initialized === true; + if (appState.mode === "host") setAttachMode(true); + if (appState.mode === "embedded") setAttachMode(false); + const snapshotSeq = finiteNumber(snapshot.latestSeq); + if (snapshotSeq !== null) { + state.lastSeq = snapshotSeq; + state.lastSnapshotSeq = snapshotSeq; + } + // The control snapshot is authoritative for routing, but its transcript + // fields can still be placeholders while VS Code is loading history. + // Route target events immediately without replacing the retained view. + if (snapshotThreadId || state.syncedThreadId === null) state.threadId = snapshotThreadId; + state.turnId = appState.activeTurnId || ""; + state.requests = new Map((snapshot.pendingRequests || []).map((request) => [requestKey(request.requestId), request])); + state.responding.clear(); + // `subscribe` replays buffered events before sending this control + // snapshot. The snapshot is authoritative, so reconcile once at the + // end of the replay instead of leaving transient duplicate bubbles. + const hasProjection = typeof snapshot.outputTail === "string" || Array.isArray(snapshot.messages); + const projectionCommitted = hasProjection && commitOutputProjection( + snapshotThreadId, + typeof snapshot.outputTail === "string" ? snapshot.outputTail : "", + snapshot.messages, + { + historyComplete: snapshotHistoryComplete(snapshot, snapshot.metadata, appState), + authoritativeSnapshot: true, + }, + ); + const appliesToView = projectionCommitted || snapshotThreadId === state.syncedThreadId; + if (appliesToView) { + prepareForSessionSnapshot(snapshotThreadId); + applySessionMetadata({ + ...controlMetadata, + ...appState, + }, appState); + if (Array.isArray(snapshot.subagents)) { + state.subagents = snapshot.subagents; + renderSubagents(); + } else if (Array.isArray(appState.subagents)) { + state.subagents = appState.subagents; + renderSubagents(); + } + applyStatusSnapshot({ + ...(snapshot.status ? { status: snapshot.status } : {}), + ...(snapshot.executionStatus ? { status: snapshot.executionStatus } : {}), + ...(snapshot.state && typeof snapshot.state === "object" ? snapshot.state : {}), + turnId: state.turnId, + }, { allowTerminal: true, showIdle: false }); + reconcileSnapshotTerminalState(snapshot, appState); + } else if (state.sessionSwitching || hasVisibleOutputProjection()) { + setConversationStatus("正在加载会话", "active"); + } + renderRequests(); + updateIds(); + setConnection("online", "已连接"); + if (waitingForSession) { + setConversationStatus("等待在 VS Code 中打开 Codex 会话", "active"); + state.snapshotNoticeShown = false; + } else { + $("outputHint").textContent = `${appStatusLabel(appState.app || "app-server")} / ${state.role}`; + } + if (state.sessionPickerOpen && state.appReady) requestSessionList(); + return; + } + if (message.type === "resync.required") { + appendOutput("事件窗口已过期,请以当前快照为准", "error"); + return; + } + if (message.type === "response.accepted") { + const key = resolveRequestKey(message.requestId); + state.responding.delete(key); + state.requests.delete(key); + renderRequests(); + appendOutput(`请求 #${message.requestId} 已提交`, "meta"); + return; + } + if (message.type === "response.pending") { + appendOutput(`请求 #${message.requestId} 已发送,等待 VS Code 主机确认`, "meta"); + return; + } + if (message.type === "command.accepted") return; + if (message.type === "command.rejected" || message.type === "response.rejected") { + if (message.requestId !== undefined) state.responding.delete(resolveRequestKey(message.requestId)); + const rejectedId = String(message.commandId || ""); + if (rejectedId && rejectedId === String(state.modeCommandId || "")) { + clearControlModeRequest(); + setConversationStatus(sessionErrorMessage(message, "控制模式切换失败"), "warning"); + updateIds(); + return; + } + if (rejectedId && rejectedId === String(state.newSessionCommandId || "")) { + state.newSessionCommandId = ""; + const detail = message.message || message.code || t("未知错误"); + setConversationStatus( + uiLocale() === "en-US" + ? `Unable to create a new conversation: ${detail}` + : `新会话创建失败:${detail}`, + "warning", + ); + updateIds(); + return; + } + if (rejectedId && rejectedId === String(state.sessionListCommandId || "")) { + state.sessionListLoading = false; + state.sessionListCommandId = ""; + state.sessionListError = sessionErrorMessage(message, "无法读取会话"); + renderSessionPicker(); + return; + } + if (rejectedId && rejectedId === String(state.sessionSelectCommandId || "")) { + state.sessionListError = failSessionSwitch(message, "会话切换失败"); + renderSessionPicker(); + setConversationStatus(state.sessionListError, "warning"); + requestRefresh(); + return; + } + appendOutput(`${message.code}: ${message.message}`, "error"); + renderRequests(); + return; + } + if (message.type === "command.result") { + handleEvent({ type: "command.result", seq: message.seq, payload: message }); + return; + } + if (message.type === "error") appendOutput(message.message || "relay error", "error"); + } + + function embeddedSocketUrl(value) { + if (!embeddedInAether) return ""; + const raw = String(value || "").trim(); + if (!raw) return ""; + try { + const url = new URL(raw, location.href); + const expectedProtocol = location.protocol === "https:" ? "wss:" : "ws:"; + if (url.origin !== `${location.protocol}//${location.host}` && url.origin !== `${expectedProtocol}//${location.host}`) { + throw new Error(t("云端连接地址必须与当前页面同源")); + } + if (url.protocol === "http:") url.protocol = "ws:"; + if (url.protocol === "https:") url.protocol = "wss:"; + if (url.protocol !== "ws:" && url.protocol !== "wss:") throw new Error(t("云端连接配置无效")); + return url.href; + } catch (error) { + appendOutput(error?.message || t("云端连接配置无效"), "error"); + embedBridge.reportState("error", { code: "invalid_ws_url", message: error?.message || "invalid ws url" }); + return ""; + } + } + + function requestEmbedTicket(reason = "missing") { + if (!embeddedInAether || state.embedStopped || state.embedTicketRequested) return; + state.embedTicketRequested = true; + setConversationStatus(t("正在获取新的连接凭证"), "active"); + embedBridge.requestTicket({ reason, deviceId: state.embedDeviceId || undefined }); + } + + function disconnectEmbedded(reason = "parent") { + if (!embeddedInAether) return; + state.embedStopped = true; + state.embedTicket = ""; + state.embedTicketRequested = false; + clearTimeout(state.reconnectTimer); + state.reconnectTimer = null; + const socket = state.ws; + state.ws = null; + if (socket && socket.readyState <= WebSocket.OPEN) socket.close(1000, reason); + setConnection("offline", "云端连接已断开"); + setConversationStatus(t("父页面已断开连接"), "warning"); + } + + function applyEmbedConnection(message) { + if (!embeddedInAether) return; + if (message.locale) i18n?.setLocale?.(message.locale, { persist: false }); + const ticket = typeof message.ticket === "string" ? message.ticket.trim() : ""; + const wsUrl = embeddedSocketUrl(message.wsUrl || "/api/vscodex/ws"); + if (!ticket || !wsUrl) { + appendOutput(t("云端连接配置无效"), "error"); + embedBridge.reportState("error", { code: "invalid_connection_config", message: "ticket and wsUrl are required" }); + requestEmbedTicket("invalid"); + return; + } + state.embedStopped = false; + state.embedTicketRequested = false; + state.embedTicket = ticket; + state.embedWsUrl = wsUrl; + state.embedDeviceId = typeof message.deviceId === "string" ? message.deviceId : state.embedDeviceId; + setAuthRequired(true); + document.body.classList.add("embed-aether"); + setConversationStatus(t("正在连接云端会话"), "active"); + if (state.ws && state.ws.readyState <= WebSocket.OPEN) { + const socket = state.ws; + state.ws = null; + socket.close(1000, "connection replaced"); + } + connect(); + } + + function connect() { + if (state.ws && state.ws.readyState <= WebSocket.OPEN) return; + if (embeddedInAether) { + if (state.embedStopped) return; + if (!state.embedTicket || !state.embedWsUrl) { requestEmbedTicket("missing"); return; } + state.token = state.embedTicket; + } else state.token = $("tokenInput").value.trim(); + if (!state.token && state.authRequired === true) { appendOutput("当前 relay 需要 token", "error"); return; } + setConnection("pending", "连接中"); + const protocol = location.protocol === "https:" ? "wss:" : "ws:"; + let socket; + try { + socket = new WebSocket(embeddedInAether ? state.embedWsUrl : `${protocol}//${location.host}/ws`); + } catch (error) { + if (embeddedInAether) { + state.embedTicket = ""; + requestEmbedTicket("socket-error"); + } + appendOutput(error?.message || "WebSocket 未连接", "error"); + return; + } + const connectionTicket = embeddedInAether ? state.embedTicket : ""; + if (embeddedInAether) { + state.embedTicket = ""; + state.token = ""; + } + state.ws = socket; + socket.addEventListener("open", () => { + if (state.ws !== socket) return; + // Send a hello even when local auth is disabled so the relay can assign + // the browser role without requiring a dummy password. + socket.send(JSON.stringify({ v: 1, kind: "hello", clientType: "web", protocol: 1 })); + // A loopback relay authenticates on hello. Do not send a stale token as + // a second frame after that handshake, because it is already complete. + const token = embeddedInAether ? connectionTicket : state.token; + if (token && state.authRequired !== false) socket.send(JSON.stringify({ type: "auth", token })); + }); + socket.addEventListener("message", (event) => { + if (state.ws !== socket) return; + try { handleMessage(JSON.parse(event.data)); } catch { appendOutput("收到无法解析的 relay 消息", "error"); } + }); + socket.addEventListener("close", (event) => { + if (state.ws !== socket) return; + setConnection("offline", event.code === 1008 ? "认证失败,准备重连" : "准备重连"); + state.appReady = false; + state.turnId = ""; + state.sessionListLoading = false; + state.sessionListCommandId = ""; + state.newSessionCommandId = ""; + clearControlModeRequest(); + state.sessions = []; + state.sessionFocusedId = ""; + // A transport reconnect may resume the same owner hand-off. Preserve + // its previous/target identities so the next control placeholder cannot + // be mistaken for an authoritative empty target transcript. + if (!state.sessionSwitching) state.sessionSelectedThreadId = ""; + state.sessionListError = "等待 relay 连接"; + updateIds(); + state.ws = null; + renderSessionPicker(); + clearTimeout(state.reconnectTimer); + state.reconnectTimer = null; + if (embeddedInAether) { + if (!state.embedStopped) requestEmbedTicket(event.code === 1008 ? "ticket-rejected" : "disconnected"); + } else state.reconnectTimer = setTimeout(connect, 3000); + }); + socket.addEventListener("error", () => { + // The close handler owns retrying. Keep transient socket errors in the + // connection indicator instead of adding noisy messages to the turn. + if (state.ws === socket) setConnection("offline", "重连中"); + }); + } + + function closePopovers() { + for (const id of ["panelMenu", "detailsPopover", "sessionPicker", "composerPlusMenu"]) { + const element = $(id); + if (element) element.hidden = true; + } + state.sessionPickerOpen = false; + state.sessionSearch = ""; + state.sessionFocusedId = ""; + $("sessionPickerButton")?.setAttribute("aria-expanded", "false"); + const sessionSearchInput = $("sessionSearchInput"); + if (sessionSearchInput) { + sessionSearchInput.value = ""; + sessionSearchInput.setAttribute("aria-expanded", "false"); + sessionSearchInput.setAttribute("aria-activedescendant", ""); + } + $("sessionSearchClear")?.setAttribute("hidden", ""); + setModelMenu(false); + setPermissionMenu(false); + setUsageMenu(false); + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + const confirm = $("permissionConfirm"); + if (confirm) { confirm.hidden = true; delete confirm.dataset.pendingMode; } + } + + function requestRefresh() { + if (state.ws && state.ws.readyState === WebSocket.OPEN) { + try { sendFrame({ type: "subscribe", fromSeq: state.lastSeq }); } catch { connect(); } + } else connect(); + } + + document.querySelectorAll("[data-panel-action]").forEach((button) => { + button.addEventListener("click", () => { + const action = button.dataset.panelAction; + if (action === "back" || action === "history") { + setSessionPicker(!state.sessionPickerOpen); + return; + } + if (action === "new-session") { + requestNewSession(); + return; + } + if (action === "expand") { + document.body.classList.toggle("panel-expanded"); + return; + } + if (action === "close") { + document.body.classList.add("panel-hidden"); + const restore = $("restorePanel"); + if (restore) restore.hidden = false; + closePopovers(); + return; + } + if (action === "refresh") { closePopovers(); requestRefresh(); return; } + if (action === "menu") { + const menu = $("panelMenu"); + const details = $("detailsPopover"); + if (details) details.hidden = true; + if (menu) menu.hidden = !menu.hidden; + return; + } + if (action === "settings") { + const details = $("detailsPopover"); + const menu = $("panelMenu"); + if (menu) menu.hidden = true; + if (details) details.hidden = !details.hidden; + updateIds(); + } + }); + }); + $("sessionPickerButton")?.addEventListener("click", () => { + setSessionPicker(!state.sessionPickerOpen); + }); + $("controlModeSwitch")?.querySelectorAll("[data-control-mode]").forEach((button) => { + button.addEventListener("click", () => requestControlMode(button.dataset.controlMode)); + }); + $("sessionPickerRefresh")?.addEventListener("click", (event) => { + event.stopPropagation(); + requestSessionList(); + }); + $("sessionSearchInput")?.addEventListener("input", (event) => { + const input = event.currentTarget; + if (!(input instanceof HTMLInputElement)) return; + state.sessionSearch = input.value; + state.sessionFocusedId = ""; + renderSessionPicker(); + }); + $("sessionSearchInput")?.addEventListener("keydown", handleSessionPickerKeydown); + $("sessionList")?.addEventListener("keydown", handleSessionPickerKeydown); + $("sessionSearchClear")?.addEventListener("click", (event) => { + event.preventDefault(); + event.stopPropagation(); + state.sessionSearch = ""; + state.sessionFocusedId = ""; + renderSessionPicker(); + $("sessionSearchInput")?.focus(); + }); + $("restorePanel")?.addEventListener("click", () => { + document.body.classList.remove("panel-hidden"); + $("restorePanel").hidden = true; + }); + $("panelMenu")?.querySelectorAll("[data-menu-action]").forEach((button) => { + button.addEventListener("click", (event) => { + if (button.dataset.menuAction === "sessions") { + // The document-level outside-click handler runs in the same bubble + // phase. Keep the picker open when it is launched from this menu. + event.stopPropagation(); + closePopovers(); + setSessionPicker(true); + return; + } else if (button.dataset.menuAction === "clear") { + renderEmptyOutput(); + state.outputSynced = false; + } else if (button.dataset.menuAction === "refresh") requestRefresh(); + else if (button.dataset.menuAction === "expand") document.body.classList.toggle("panel-expanded"); + else if (button.dataset.menuAction === "close") { + document.body.classList.add("panel-hidden"); + const restore = $("restorePanel"); + if (restore) restore.hidden = false; + } + closePopovers(); + }); + }); + $("detailsPopover")?.querySelectorAll("[data-settings-action]").forEach((button) => { + button.addEventListener("click", (event) => { + event.stopPropagation(); + const action = button.dataset.settingsAction; + closePopovers(); + if (action === "model") { + setModelMenu(true); + $("modelPickerButton")?.focus(); + } else if (action === "permission") { + setPermissionMenu(true); + $("permissionChip")?.focus(); + } + }); + }); + document.addEventListener("click", (event) => { + const target = event.target; + if (!(target instanceof Element)) return; + if (!target.closest(".model-picker")) setModelMenu(false); + if (!target.closest(".permission-menu, #permissionChip")) setPermissionMenu(false); + if (!target.closest(".usage-menu, #usageButton")) setUsageMenu(false); + if (!target.closest(".composer-plus-menu, #composerPlusButton")) { + const menu = $("composerPlusMenu"); + if (menu) menu.hidden = true; + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + } + if (!target.closest("[data-panel-action], .panel-popover, .composer-popover, .composer-icon-button, .permission-chip, .usage-button, #sessionPickerButton")) closePopovers(); + }); + document.addEventListener("keydown", (event) => { + if (event.key === "Escape" && state.sessionPickerOpen) { + event.preventDefault(); + closePopovers(); + $("sessionPickerButton")?.focus(); + } + }); + + $("tokenInput").addEventListener("keydown", (event) => { + if (event.key === "Enter") { + event.preventDefault(); + connect(); + } + }); + $("tokenInput").addEventListener("change", connect); + $("localeSelect")?.addEventListener("change", (event) => { + if (!embeddedInAether) i18n?.setLocale?.(event.target.value, { persist: true }); + }); + window.addEventListener("aether-vscodex:locale", () => { + for (const activity of state.activities.values()) { + renderActivityText(activity); + refreshActivity(activity); + } + for (const [turnId, divider] of state.turnDividers) { + const label = divider.querySelector(".turn-divider-label"); + if (!label) continue; + const duration = finiteNumber(divider.dataset.durationMs); + label.textContent = turnDividerLabel(divider.dataset.status || "completed", duration); + } + setAttachMode(state.attachMode); + setAuthRequired(state.authRequired); + renderSessionPicker(); + renderModelPicker(); + renderPermissionMenu(); + renderUsage(); + renderSubagents(); + updateIds(); + renderRequests(); + for (const checkbox of document.querySelectorAll(".task-list-item input[type=checkbox]")) { + checkbox.setAttribute("aria-label", t(checkbox.checked ? "已完成" : "未完成")); + } + }); + $("clearOutputButton").addEventListener("click", () => { + renderEmptyOutput(); + state.outputSynced = false; + }); + $("startThreadButton").addEventListener("click", () => { + if (state.attachMode) return; + const params = { cwd: $("cwdInput").value.trim() || undefined, sandbox: $("sandboxInput").value, approvalPolicy: $("approvalInput").value }; + if ($("modelInput").value.trim()) params.model = $("modelInput").value.trim(); + command("thread/start", params); + }); + $("startTurnButton").addEventListener("click", () => { + const text = composerText(); + if (!text || !state.threadId) return; + command("turn/start", { + threadId: state.threadId, + input: [{ type: "text", text, text_elements: [] }], + ...(state.currentModel ? { model: state.currentModel } : {}), + ...currentEffortParams(), + }); + clearComposer(); + }); + $("steerButton").addEventListener("click", () => { + const text = composerText(); + if (!text || !state.threadId || !state.turnId) return; + command("turn/steer", { + threadId: state.threadId, + expectedTurnId: state.turnId, + input: [{ type: "text", text, text_elements: [] }], + ...(state.currentModel ? { model: state.currentModel } : {}), + ...currentEffortParams(), + }); + clearComposer(); + }); + $("interruptButton").addEventListener("click", () => { + if (state.threadId && state.turnId) command("turn/interrupt", { threadId: state.threadId, turnId: state.turnId }); + }); + $("messageInput").addEventListener("keydown", (event) => { + if (event.key !== "Enter" || event.shiftKey || event.isComposing) return; + event.preventDefault(); + const button = state.turnId ? $("steerButton") : $("startTurnButton"); + if (button && !button.disabled) button.click(); + }); + $("messageInput").addEventListener("input", () => { + resizeComposer(); + updateIds(); + }); + $("messageInput").addEventListener("paste", (event) => { + event.preventDefault(); + const text = event.clipboardData?.getData("text/plain") || ""; + if (!text) return; + const editor = $("messageInput"); + const selection = window.getSelection(); + if (editor && selection && selection.rangeCount) { + const range = selection.getRangeAt(0); + if (editor.contains(range.commonAncestorContainer)) { + range.deleteContents(); + const node = document.createTextNode(text); + range.insertNode(node); + range.setStartAfter(node); + range.collapse(true); + selection.removeAllRanges(); + selection.addRange(range); + } else editor.append(document.createTextNode(text)); + } else editor?.append(document.createTextNode(text)); + resizeComposer(); + updateIds(); + }); + $("modelPickerButton")?.addEventListener("click", (event) => { + event.stopPropagation(); + const menu = $("modelMenu"); + setPermissionMenu(false); + setUsageMenu(false); + const plus = $("composerPlusMenu"); + if (plus) plus.hidden = true; + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + setModelMenu(Boolean(menu?.hidden)); + }); + $("modelAdvancedToggle")?.addEventListener("click", (event) => { + event.stopPropagation(); + state.modelAdvancedOpen = !state.modelAdvancedOpen; + renderModelPicker(); + if (state.modelAdvancedOpen) $("modelAdvancedBack")?.focus(); + else $("modelPowerSlider")?.focus(); + }); + $("modelAdvancedBack")?.addEventListener("click", (event) => { + event.stopPropagation(); + state.modelAdvancedOpen = false; + renderModelPicker(); + $("modelPowerSlider")?.focus(); + }); + $("modelPowerSlider")?.addEventListener("input", (event) => { + selectPowerIndex(event.currentTarget.value); + }); + $("composerPlusButton")?.addEventListener("click", (event) => { + event.stopPropagation(); + const menu = $("composerPlusMenu"); + if (!menu) return; + const next = menu.hidden; + menu.hidden = !next; + $("composerPlusButton").setAttribute("aria-expanded", String(next)); + if (next) { + setPermissionMenu(false); + setUsageMenu(false); + setModelMenu(false); + } + }); + $("permissionChip")?.addEventListener("click", (event) => { + event.stopPropagation(); + const menu = $("permissionMenu"); + setPermissionMenu(Boolean(menu?.hidden)); + if (!menu?.hidden) { + const plus = $("composerPlusMenu"); + if (plus) plus.hidden = true; + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + setUsageMenu(false); + setModelMenu(false); + } + }); + $("usageButton")?.addEventListener("click", (event) => { + event.stopPropagation(); + const menu = $("usageMenu"); + setPermissionMenu(false); + setModelMenu(false); + const plus = $("composerPlusMenu"); + if (plus) plus.hidden = true; + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + setUsageMenu(Boolean(menu?.hidden)); + if (!menu?.hidden) { + setPermissionMenu(false); + setModelMenu(false); + } + }); + const insertComposerText = (text) => { + const editor = $("messageInput"); + if (!editor || !text) return; + editor.focus(); + const selection = window.getSelection(); + if (selection && selection.rangeCount) { + const range = selection.getRangeAt(0); + if (editor.contains(range.commonAncestorContainer)) { + range.deleteContents(); + const node = document.createTextNode(text); + range.insertNode(node); + range.setStartAfter(node); + range.collapse(true); + selection.removeAllRanges(); + selection.addRange(range); + } else editor.append(document.createTextNode(text)); + } else editor.append(document.createTextNode(text)); + resizeComposer(); + updateIds(); + }; + $("composerPlusMenu")?.querySelectorAll("[data-composer-action]").forEach((button) => { + button.addEventListener("click", () => { + const action = button.dataset.composerAction; + if (action === "attach") { + const input = $("attachmentInput"); + if (input) { input.accept = ".txt,.md,.json,.js,.ts,.tsx,.jsx,.css,.html,.yml,.yaml,.xml,.py,.go,.rs,.java,.c,.cpp,.h"; input.click(); } + } else if (action === "photo") { + const input = $("attachmentInput"); + if (input) { input.accept = "image/*"; input.click(); } + } else if (action === "workspace") { + insertComposerText("\n\n@workspace "); + setConversationStatus("已添加工作区上下文", "ready"); + } else if (action === "web-search") { + insertComposerText("\n\n/web-search "); + setConversationStatus("已添加网页搜索", "ready"); + } + const menu = $("composerPlusMenu"); + if (menu) menu.hidden = true; + $("composerPlusButton")?.setAttribute("aria-expanded", "false"); + }); + }); + $("attachmentInput")?.addEventListener("change", async (event) => { + const files = [...(event.target?.files || [])]; + for (const file of files) { + try { + if (file.type.startsWith("image/")) { + insertComposerText(`\n\n[图片附件:${file.name}]\n`); + continue; + } + const text = await file.text(); + const clipped = text.length > 80_000 ? `${text.slice(0, 80_000)}\n${t("…(文件已截断)")}` : text; + const fence = "```"; + insertComposerText(`\n\n### ${file.name}\n\n${fence}\n${clipped}\n${fence}\n`); + } catch { + setConversationStatus(`无法读取 ${file.name}`, "warning"); + } + } + event.target.value = ""; + }); + $("permissionMenu")?.querySelectorAll("[data-permission-mode], [data-sandbox], [data-approval]").forEach((button) => { + button.addEventListener("click", () => { + if (button.dataset.permissionMode) selectPermissionMode(button.dataset.permissionMode); + else if (button.dataset.sandbox) selectPermissionSetting("sandbox", button.dataset.sandbox); + else if (button.dataset.approval) selectPermissionSetting("approval", button.dataset.approval); + }); + }); + $("permissionConfirmCancel")?.addEventListener("click", () => { + const confirm = $("permissionConfirm"); + if (confirm) { confirm.hidden = true; delete confirm.dataset.pendingMode; } + }); + $("permissionConfirmAccept")?.addEventListener("click", () => { + const confirm = $("permissionConfirm"); + const mode = confirm?.dataset.pendingMode || "full"; + if (confirm) { confirm.hidden = true; delete confirm.dataset.pendingMode; } + applyPermissionMode(mode); + }); + $("subagentsToggle")?.addEventListener("click", () => { + state.subagentsCollapsed = !state.subagentsCollapsed; + renderSubagents(); + }); + $("output").addEventListener("scroll", () => updateScrollToBottom($("output")), { passive: true }); + $("scrollToBottom")?.addEventListener("click", () => { + const output = $("output"); + if (output) output.scrollTo({ top: output.scrollHeight, behavior: "smooth" }); + }); + if (typeof ResizeObserver === "function") { + const output = $("output"); + const observedContent = new Set(); + const observeTranscriptContent = () => { + const children = new Set(output ? [...output.children] : []); + for (const child of observedContent) { + if (children.has(child)) continue; + layoutObserver.unobserve(child); + observedContent.delete(child); + } + for (const child of children) { + if (observedContent.has(child)) continue; + observedContent.add(child); + layoutObserver.observe(child); + } + }; + const layoutObserver = new ResizeObserver(() => { + const distance = state.outputDistanceFromBottom; + const following = distance <= 24; + const anchorLocked = motionClock() < (state.timelineAnchorLockUntil || 0); + updateScrollPadding(); + // Preserve the reader's distance from the bottom when a streamed item or + // the composer changes height. The official thread layout uses the same + // bottom-relative anchor instead of allowing content to jump. + if (output && !anchorLocked) { + if (following) scrollOutput(output, true); + else output.scrollTop = Math.max(0, output.scrollHeight - output.clientHeight - distance); + updateScrollToBottom(output); + } else if (output) updateScrollToBottom(output); + observeTranscriptContent(); + }); + layoutObserver.observe($("messageForm")); + layoutObserver.observe($("inlineRequests")); + if (output) layoutObserver.observe(output); + observeTranscriptContent(); + if (typeof MutationObserver === "function" && output) { + const childObserver = new MutationObserver(observeTranscriptContent); + childObserver.observe(output, { childList: true }); + } + } + window.addEventListener("resize", updateScrollPadding, { passive: true }); + $("cwdInput").value = location.pathname === "/" ? "" : ""; + renderPermissionMenu(); + renderUsage(); + renderSessionPicker(); + resizeComposer(); + updateScrollPadding(); + updateScrollToBottom($("output")); + updateIds(); + + if (embeddedInAether) { + document.body.classList.add("embed-aether"); + setAuthRequired(true); + setConversationStatus(t("正在等待云端连接"), "active"); + embedBridge.on("connect", applyEmbedConnection); + embedBridge.on("context", (message) => { + if (message.locale) i18n?.setLocale?.(message.locale, { persist: false }); + }); + embedBridge.on("disconnect", () => disconnectEmbedded("parent disconnect")); + embedBridge.on("error", (message) => { + const detail = typeof message.message === "string" && message.message ? message.message : t("云端连接已断开"); + appendOutput(detail, "error"); + setConversationStatus(detail, "warning"); + }); + } else { + // /api/health is intentionally public and only reports capability metadata. + // Probe it first so a local relay can connect automatically without a token; + // Authenticated deployments wait until a token is entered in the field. + fetch("./api/health", { cache: "no-store" }).then(async (response) => { + let health; + try { health = await response.json(); } catch { return; } + setAuthRequired(health.authRequired); + if (health.authRequired === false || (health.authRequired === true && $("tokenInput").value.trim())) connect(); + }).catch(() => undefined); + } +})(); diff --git a/aether-vscodex/public/embed-bridge.js b/aether-vscodex/public/embed-bridge.js new file mode 100644 index 000000000..0d0a7416b --- /dev/null +++ b/aether-vscodex/public/embed-bridge.js @@ -0,0 +1,112 @@ +(function (root, factory) { + "use strict"; + + const api = factory(); + if (typeof module === "object" && module.exports) module.exports = api; + if (!root || !root.document) return; + + const bridge = api.createAetherEmbedBridge(root); + root.AetherVscodexEmbed = bridge; + if (bridge.active) bridge.start(); +})(typeof window === "object" ? window : undefined, function () { + "use strict"; + + const VERSION = 1; + const PREFIX = "aether-vscodex/"; + const INBOUND_TYPES = new Set(["connect", "context", "disconnect", "error"]); + + function isAetherEmbed(locationLike) { + try { + return new URLSearchParams(locationLike?.search || "").get("embed") === "aether"; + } catch { + return false; + } + } + + function normalizeTheme(value) { + const theme = String(value || "").trim().toLowerCase(); + return theme === "dark" || theme === "light" ? theme : "system"; + } + + function createAetherEmbedBridge(windowLike) { + const active = isAetherEmbed(windowLike.location); + const listeners = new Map(); + const pending = new Map(); + let started = false; + + const emit = (name, payload) => { + for (const listener of listeners.get(name) || []) listener(payload); + }; + + const post = (type, payload = {}) => { + if (!active || windowLike.parent === windowLike) return false; + windowLike.parent.postMessage({ v: VERSION, type: `${PREFIX}${type}`, ...payload }, windowLike.location.origin); + return true; + }; + + const applyContext = (payload) => { + if (payload.locale && windowLike.VscodexI18n?.setLocale) { + windowLike.VscodexI18n.setLocale(payload.locale, { persist: false }); + } + const theme = normalizeTheme(payload.theme); + const documentElement = windowLike.document?.documentElement; + if (documentElement) { + if (theme === "system") delete documentElement.dataset.theme; + else documentElement.dataset.theme = theme; + documentElement.style.colorScheme = theme === "system" ? "" : theme; + } + }; + + const handleMessage = (event) => { + if (!active || event.origin !== windowLike.location.origin || event.source !== windowLike.parent) return; + const message = event.data; + if (!message || typeof message !== "object" || message.v !== VERSION || typeof message.type !== "string") return; + if (!message.type.startsWith(PREFIX)) return; + const name = message.type.slice(PREFIX.length); + if (!INBOUND_TYPES.has(name)) return; + if (name === "connect" || name === "context") applyContext(message); + if (!(listeners.get(name)?.size)) pending.set(name, message); + emit(name, message); + }; + + return { + active, + version: VERSION, + start() { + if (!active || started) return; + started = true; + windowLike.document.body?.classList.add("embed-aether"); + windowLike.addEventListener("message", handleMessage); + post("ready"); + }, + stop() { + if (!started) return; + started = false; + windowLike.removeEventListener("message", handleMessage); + listeners.clear(); + pending.clear(); + }, + on(name, listener) { + if (!INBOUND_TYPES.has(name) || typeof listener !== "function") return () => undefined; + if (!listeners.has(name)) listeners.set(name, new Set()); + listeners.get(name).add(listener); + if (pending.has(name)) { + const message = pending.get(name); + pending.delete(name); + listener(message); + } + return () => listeners.get(name)?.delete(listener); + }, + post, + requestTicket(payload = {}) { + return post("request-ticket", payload); + }, + reportState(state, payload = {}) { + return post("state", { state, ...payload }); + }, + _handleMessage: handleMessage, + }; + } + + return { createAetherEmbedBridge, isAetherEmbed, normalizeTheme }; +}); diff --git a/aether-vscodex/public/i18n.js b/aether-vscodex/public/i18n.js new file mode 100644 index 000000000..8c16b8f3b --- /dev/null +++ b/aether-vscodex/public/i18n.js @@ -0,0 +1,541 @@ +(function (root, factory) { + "use strict"; + + const api = factory(root); + if (typeof module === "object" && module.exports) module.exports = api; + if (root?.document) root.VscodexI18n = api; +})(typeof window === "object" ? window : undefined, function (root) { + "use strict"; + + const STORAGE_KEY = "aether-vscodex.locale"; + const SUPPORTED = new Set(["zh-CN", "en-US"]); + const EN = Object.freeze({ + "本地模式": "Local mode", + "独立模式": "Standalone mode", + "云端模式": "Cloud mode", + "控制模式": "Control mode", + "同步": "Sync", + "异步": "Async", + "同步模式跟随 VS Code 当前会话": "Sync mode follows the current VS Code conversation", + "异步模式可独立管理会话": "Async mode manages conversations independently", + "正在切换控制模式": "Switching control mode", + "控制模式已切换": "Control mode switched", + "控制模式切换失败": "Unable to switch control mode", + "当前任务或请求完成后才能切换控制模式": "The control mode can be changed after the current task or request finishes", + "同步模式下会话管理由 VS Code 控制": "VS Code controls conversation navigation in sync mode", + "当前模式不支持修改会话设置": "The current mode does not support changing conversation settings", + "本机连接(无需 token)": "Local connection (no token required)", + "本机模式无需填写;认证模式再填写": "No token is needed locally; enter one only for authenticated mode", + "访问 token(认证模式)": "Access token (authenticated mode)", + "粘贴 relay 启动时打印的 token": "Paste the token printed when the relay started", + "本地连接无需 token": "No token is needed for a local connection", + "编辑外部文件和联网时始终询问": "Always ask before editing external files or using the network", + "不限制联网或文件访问": "Allow unrestricted network and file access", + "查看请求数据": "View request data", + "查看上下文用量": "View context usage", + "创建新会话": "New conversation", + "打开会话历史": "Open conversation history", + "待处理的 Codex 请求": "Pending Codex requests", + "当前会话": "Current conversation", + "当前模型": "Current model", + "切换模型": "Change model", + "等待 VS Code 主机": "Waiting for VS Code host", + "等待连接": "Waiting for connection", + "对话内容": "Conversation", + "发送 JSON": "Send JSON", + "发送后续指令": "Send follow-up", + "发送消息": "Send message", + "返回会话列表": "Back to conversations", + "返回模型强度": "Back to model effort", + "高级": "Advanced", + "简洁": "Simple", + "更多操作": "More actions", + "更高效": "More efficient", + "更智能": "More capable", + "工作目录": "Working directory", + "工作区": "Workspace", + "工作区写入": "Workspace write", + "回到最新消息": "Jump to latest message", + "正在工作,回到最新消息": "Working, jump to latest message", + "会话历史": "Conversation history", + "会话设置": "Conversation settings", + "仅本次 turn": "This turn only", + "仅查看文件,不修改工作区": "View files without changing the workspace", + "仅对可能不安全的操作询问": "Ask only for potentially unsafe actions", + "拒绝": "Deny", + "可用会话": "Available conversations", + "连接设置": "Connection settings", + "留空使用默认模型": "Leave empty to use the default model", + "模式": "Mode", + "模型": "Model", + "模型与推理强度": "Model and reasoning effort", + "默认": "Default", + "启动新 thread": "Start new thread", + "强度": "Effort", + "切换模型与推理强度": "Change model and reasoning effort", + "清除搜索": "Clear search", + "清空当前输出": "Clear current output", + "清空对话": "Clear conversation", + "取消": "Cancel", + "权限设置": "Permission settings", + "确认": "Confirm", + "确认完全访问": "Confirm full access", + "沙箱": "Sandbox", + "上下文用量": "Context usage", + "设置": "Settings", + "审批策略": "Approval policy", + "使用 config.toml 中的权限": "Use permissions from config.toml", + "使用左右方向键调整强度": "Use the left and right arrow keys to adjust effort", + "授权范围": "Authorization scope", + "授权与输入": "Approvals and input", + "刷新会话列表": "Refresh conversations", + "搜索最近会话": "Search recent conversations", + "提交后续变更要求": "Ask for follow-up changes", + "添加工作区上下文": "Add workspace context", + "添加文件": "Add files", + "添加文件及更多内容": "Add files and more", + "添加照片": "Add photos", + "推理强度": "Reasoning effort", + "完全访问": "Full access", + "完全访问允许 Codex 执行命令、访问互联网并编辑工作区之外的文件。": "Full access lets Codex run commands, use the internet, and edit files outside the workspace.", + "网页搜索": "Web search", + "未认证": "Unauthenticated", + "显示 Codex": "Show Codex", + "修改权限": "Change permissions", + "需要时询问": "Ask when needed", + "已附着当前会话": "Attached to current conversation", + "隐藏面板": "Hide panel", + "由 Codex 审批": "Let Codex decide", + "允许": "Allow", + "允许一次": "Allow once", + "暂无待处理请求": "No pending requests", + "暂无用量数据": "No usage data", + "展开面板": "Expand panel", + "正在连接": "Connecting", + "只读": "Read only", + "中断当前 turn": "Interrupt current turn", + "重新同步": "Resync", + "子代理": "Subagent", + "自定义": "Custom", + "最近会话": "Recent conversations", + "Codex 消息": "Codex messages", + "JSON 响应": "JSON response", + "语言": "Language", + "中文": "Chinese", + "跟随浏览器": "Use browser language", + "正在连接云端会话": "Connecting to cloud conversation", + "正在等待云端连接": "Waiting for cloud connection", + "云端连接已断开": "Cloud connection disconnected", + "云端连接配置无效": "Invalid cloud connection configuration", + "云端连接地址必须与当前页面同源": "The cloud connection URL must be same-origin", + "正在获取新的连接凭证": "Requesting new connection credentials", + "父页面已断开连接": "Disconnected by the parent page", + "当前 relay 需要 token": "This relay requires a token", + "WebSocket 未连接": "WebSocket is not connected", + "连接中": "Connecting", + "同步中": "Syncing", + "已连接": "Connected", + "认证失败,准备重连": "Authentication failed; preparing to reconnect", + "准备重连": "Preparing to reconnect", + "重连中": "Reconnecting", + "收到无法解析的 relay 消息": "Received an unreadable relay message", + "等待 relay 连接": "Waiting for relay connection", + "等待 VS Code 主机连接": "Waiting for VS Code host", + "VS Code 主机未连接": "VS Code host is disconnected", + "等待 VS Code 伴随扩展连接": "Waiting for the VS Code companion extension", + "VS Code 伴随扩展未连接": "VS Code companion extension is disconnected", + "等待在 VS Code 中打开 Codex 会话": "Open a Codex conversation in VS Code to continue", + "会话已关闭": "Conversation closed", + "VS Code 会话已关闭": "VS Code conversation closed", + "会话操作失败": "Conversation operation failed", + "当前任务结束或请求处理后才能切换": "You can switch after the current task or request finishes", + "目标会话没有返回 VS Code 快照,请先在官方 Codex 面板打开它": "The target conversation did not return a VS Code snapshot. Open it in the official Codex panel first.", + "当前 relay 版本不支持此会话操作,请重启 relay": "This relay version does not support the conversation action. Restart the relay.", + "正在读取会话…": "Loading conversations...", + "正在切换会话…": "Switching conversation...", + "无法读取会话": "Unable to load conversations", + "没有匹配的会话": "No matching conversations", + "没有可附加的会话": "No attachable conversations", + "没有可控制的会话": "No controllable conversations", + "正在切换": "Switching", + "未打开": "Not open", + "当前": "Current", + "可切换": "Available", + "会话": "Conversation", + "当前角色不能创建会话": "Your current role cannot create conversations", + "正在创建新会话": "Creating a new conversation", + "无法创建新会话": "Unable to create a new conversation", + "当前任务仍在运行或等待授权,暂不能切换": "The current task is running or awaiting approval, so it cannot be switched yet", + "会话切换失败": "Conversation switch failed", + "正在确认会话": "Confirming conversation", + "正在加载会话": "Loading conversation", + "会话已切换": "Conversation switched", + "正在更新模型设置": "Updating model settings", + "模型设置已更新": "Model settings updated", + "无法更新模型设置": "Unable to update model settings", + "已停止": "Stopped", + "成功": "Succeeded", + "无输出": "No output", + "等待输出…": "Waiting for output...", + "执行步骤": "Action", + "正在读取文件": "Reading files", + "读取完成": "Finished reading", + "已读取文件运行了命令": "Read files and ran a command", + "已读取文件": "Read files", + "编辑了文件": "Edited files", + "已完成计划": "Completed plan", + "读取文件失败": "Failed to read files", + "已停止读取文件": "Stopped reading files", + "读取文件": "Read files", + "已运行命令": "Ran command", + "正在运行命令": "Running command", + "正在思考": "Thinking", + "正在制定计划": "Creating a plan", + "正在编辑文件": "Editing files", + "正在处理": "Working", + "已完成思考": "Finished thinking", + "计划完成": "Plan completed", + "文件编辑完成": "Finished editing files", + "工作说明": "Progress update", + "计划": "Plan", + "文件变更": "File changes", + "等待授权": "Waiting for approval", + "正在生成": "Generating", + "已中断": "Interrupted", + "失败": "Failed", + "已完成": "Completed", + "正在工作": "Working", + "正在等待你的回答": "Waiting for your answer", + "正在搜索网页": "Searching the web", + "执行失败": "Action failed", + "处理中": "Working", + "思考": "Reasoning", + "编辑文件": "Edit files", + "思考中": "Thinking", + "编辑中": "Editing", + "进行中": "In progress", + "异常": "Error", + "未读": "Unread", + "本地会话": "Local conversation", + "默认拒绝,请明确允许": "Denied by default; allow explicitly", + "需要远程确认或输入": "Remote confirmation or input is required", + "允许运行命令?": "Allow this command?", + "允许修改文件?": "Allow file changes?", + "需要扩大权限": "Additional permissions required", + "Codex 需要你的回答": "Codex needs your answer", + "需要外部服务确认": "External service confirmation required", + "Codex 请求确认": "Codex requests confirmation", + "高风险": "High risk", + "低风险": "Low risk", + "需确认": "Confirmation required", + "请输入": "Enter a response", + "提交回答": "Submit answer", + "发送自定义响应": "Send custom response", + "自定义响应不是有效 JSON": "The custom response is not valid JSON", + "响应不是有效 JSON": "The response is not valid JSON", + "远程参与者拒绝": "Denied by remote participant", + "状态": "Status", + "命令": "Command", + "详情": "Details", + "复制消息": "Copy message", + "复制命令": "Copy command", + "复制输出": "Copy output", + "未知": "Unknown", + "未知错误": "Unknown error", + "已附着 VS Code 当前 Codex 会话;输入、输出和授权都回到同一个会话。": "Attached to the current VS Code Codex conversation. Messages, output, and approvals all return to that conversation.", + "当前为独立 app-server 模式。": "Currently using standalone app-server mode.", + "已附着现有会话": "Attached to existing conversation", + "通用 Codex 模型": "General-purpose Codex model", + "平衡速度与推理": "Balanced speed and reasoning", + "可用模型": "Available model", + "极低": "Minimal", + "轻度": "Low", + "标准": "Medium", + "深度": "High", + "极高": "Extra high", + "最大": "Maximum", + "此模型使用默认推理强度": "This model uses its default reasoning effort", + "返回简洁模型选择": "Return to simple model selection", + "显示高级模型选项": "Show advanced model options", + "自定义权限由 config.toml 管理": "Custom permissions are managed by config.toml", + "正在等待指示": "Waiting for instructions", + "正在工作": "Working", + "命令输出": "Command output", + "工具输出": "Tool output", + "发送 Steer": "Send steer", + "会话切换失败,已恢复原会话": "Conversation switch failed; restored the previous conversation", + "(空消息)": "(empty message)", + "今天": "Today", + "昨天": "Yesterday", + "未完成": "Not completed", + "步骤": "Step", + "查看图像": "View image", + "等待输入": "Waiting for input", + "读取文件运行命令失败": "Failed to read files and run a command", + "发送输入": "Send input", + "工具": "Tool", + "工具失败": "Tool failed", + "正在搜索": "Searching", + "你停止了工作": "You stopped working", + "关闭子代理": "Close subagent", + "恢复子代理": "Resume subagent", + "启动子代理": "Start subagent", + "搜索": "Search", + "文件": "File", + "新会话已在 VS Code 中打开": "The new conversation opened in VS Code", + "事件窗口已过期,请以当前快照为准": "The event window expired; the current snapshot is authoritative", + "执行状态未知,请等待主机恢复": "Execution status is unknown; wait for the host to recover", + "文件已截断": "File truncated", + "已拒绝": "Denied", + "已开始工作": "Started working", + "已添加工作区上下文": "Added workspace context", + "已添加网页搜索": "Added web search", + "运行命令": "Run command", + "整理上下文": "Compacting context", + "正在切换会话": "Switching conversation", + "MCP 工具": "MCP tool", + " · @ 可标记代理": " · @ to mention agents", + }); + + const EN_PATTERNS = Object.freeze([ + [/^用时 1分钟(\d+)秒$/, "Worked for 1m{1}s"], + [/^用时 (\d+)分(\d+)秒$/, "Worked for {1}m{2}s"], + [/^用时 1分钟$/, "Worked for 1m"], + [/^用时 (\d+)分$/, "Worked for {1}m"], + [/^用时 (\d+)秒$/, "Worked for {1}s"], + [/^用时 (\d+)毫秒$/, "Worked for {1}ms"], + [/^用时\s+(.+)$/, "Worked for {1}"], + [/^已思考 1分钟(\d+)秒$/, "Thought for 1m{1}s"], + [/^已思考 (\d+)分(\d+)秒$/, "Thought for {1}m{2}s"], + [/^已思考 (\d+)秒$/, "Thought for {1}s"], + [/^已思考\s+(.+)$/, "Thought for {1}"], + [/^退出码\s+(.+)$/, "Exit code {1}"], + [/^正在读取\s+(.+)$/, "Reading {1}", [1]], + [/^已读取\s+(.+)$/, "Read {1}", [1]], + [/^读取失败\s*·\s*(.+)$/, "Failed to read {1}", [1]], + [/^已停止读取\s+(.+)$/, "Stopped reading {1}", [1]], + [/^读取\s+(.+)$/, "Read {1}", [1]], + [/^已读取这些内容\s*·\s*(\d+)\s*个文件(.*)$/, "Read these items · {1} files{2}"], + [/^已在\s+(.+)\s+内运行\s+(.+)$/, "Ran {2} in {1}", [2]], + [/^命令运行失败\s*·\s*(.+?)\s*·\s*((?:\d+毫秒|\d+秒|1分钟(?:\d+秒)?|\d+分(?:\d+秒)?))$/, "Command failed · {1} · {2}", [1]], + [/^命令运行失败\s*·\s*(.+)$/, "Command failed · {1}", [1]], + // Renderer-owned disclosure labels. Keep the captured command/model text + // intact; only the surrounding UI words are localized. + [/^命令\s*·\s*(.+)$/, "Command · {1}", [1]], + [/^已工具\s*·\s*(.+)$/, "Tool completed · {1}"], + [/^当前模型\s+(.+?)\s+(极低|轻度|标准|深度|极高|最大),切换模型$/, "Current model: {1} {2}. Change model", [1]], + [/^已停止\s*(.+?)\s*·\s*((?:\d+毫秒|\d+秒|1分钟(?:\d+秒)?|\d+分(?:\d+秒)?))$/, "Stopped {1} · {2}", [1]], + [/^已运行\s*(.+)$/, "Ran {1}", [1]], + [/^命令运行失败\s*(.*)$/, "Command failed{1}", [1]], + [/^命令:\s*(.+?)(执行状态未知,请等待主机恢复)$/, "Command: {1} (execution status unknown; wait for the host to recover)", [1]], + [/^命令:\s*(.+)$/, "Command: {1}", [1]], + [/^已停止\s*(.+)$/, "Stopped {1}", [1]], + [/^正在运行\s+(.+)$/, "Running {1}", [1]], + [/^(.+?)\s*·\s*失败$/, "{1} · Failed"], + [/^(.+?)\s*·\s*已中断$/, "{1} · Interrupted"], + [/^(.+)\s+失败$/, "{1} failed"], + [/^编辑了文件\s*·\s*(.+)$/, "Edited files · {1}"], + [/^已完成计划\s*·\s*(.+)$/, "Completed plan · {1}"], + [/^(\d+)\/(\d+)\s*个会话$/, "{1}/{2} conversations"], + [/^(\d+)\s*个会话$/, "{1} conversations"], + [/^会话\s+(.+)$/, "Conversation {1}", [1]], + [/^工作区\s*·\s*(.+)$/, "Workspace · {1}", [1]], + [/^昨天\s+(.+)$/, "Yesterday {1}", [1]], + [/^正在切换到「(.+)」…?$/, "Switching to “{1}”...", [1]], + [/^你在\s+(.+)\s+后停止了$/, "You stopped after {1}"], + [/^执行失败\s*·\s*(.+)$/, "Action failed · {1}"], + [/^新会话创建失败:(.+)$/, "Unable to create a new conversation: {1}", [1]], + [/^模型设置更新失败:(.+)$/, "Unable to update model settings: {1}", [1]], + [/^(.+)(执行状态未知,请等待主机恢复)$/, "{1} (execution status unknown; wait for the host to recover)", [1]], + [/^(.+) 完成$/, "{1} completed", [1]], + [/^请求 #(.+) 已提交$/, "Request #{1} submitted"], + [/^请求 #(.+) 已发送,等待 VS Code 主机确认$/, "Request #{1} sent; waiting for the VS Code host"], + [/^无法读取 (.+)$/, "Unable to read {1}", [1]], + [/^\[图片附件:(.+)\]$/, "[Image attachment: {1}]", [1]], + [/^(.+) 已开始工作$/, "{1} started working", [1]], + [/^(.+) 已完成$/, "{1} completed", [1]], + [/^(.+) 已中断$/, "{1} interrupted", [1]], + [/^…(文件已截断)$/, "... (file truncated)"], + [/^当前模型\s+(.+),切换模型$/, "Current model: {1}. Change model", [1]], + [/^切换模型(当前\s+(.+?)\s+(极低|轻度|标准|深度|极高|最大))$/, "Change model (current: {1} {2})", [1]], + [/^切换模型(当前\s+(.+))$/, "Change model (current: {1})", [1]], + [/^修改权限,当前为(.+)$/, "Change permissions. Current: {1}"], + [/^修改权限(当前:(.+))$/, "Change permissions (current: {1})"], + [/^上下文已使用\s*(\d+)%(剩余\s*(\d+)%)$/, "Context used: {1}% ({2}% remaining)"], + [/^(\d+)%\s*已使用$/, "{1}% used"], + [/^剩余\s+(.+)\s+tokens$/, "{1} tokens remaining"], + [/^当前上下文\s+(.+)\s+tokens$/, "Current context: {1} tokens"], + [/^最近请求\s+(.+)\s+tokens$/, "Latest request: {1} tokens"], + [/^累计\s+(.+)\s+tokens$/, "Total: {1} tokens"], + [/^使用\s+(.+)$/, "Using {1}", [1]], + [/^已用时\s+(.+)$/, "Elapsed: {1}"], + [/^(\d+)\s*个后台代理(.*)$/, "{1} background agents{2}"], + [/^(\d+)毫秒$/, "{1}ms"], + [/^(\d+)秒$/, "{1}s"], + [/^1分钟(\d+)秒$/, "1m{1}s"], + [/^(\d+)分(\d+)秒$/, "{1}m{2}s"], + [/^1分钟(\d+秒)?$/, "1m{1}"], + [/^(\d+)分(\d+秒)?$/, "{1}m{2}"], + ]); + const ZH = Object.freeze(Object.fromEntries(Object.entries(EN).map(([source, translated]) => [translated, source]))); + + let currentLocale = "zh-CN"; + let observer = null; + const textSources = new WeakMap(); + const textRendered = new WeakMap(); + const attributeSources = new WeakMap(); + const attributeRendered = new WeakMap(); + + function normalizeLocale(value) { + const locale = String(value || "").trim().replace("_", "-").toLowerCase(); + return locale.startsWith("zh") ? "zh-CN" : "en-US"; + } + + function embeddedMode() { + if (root?.AetherVscodexEmbed?.active) return true; + try { return new URLSearchParams(root?.location?.search || "").get("embed") === "aether"; } + catch { return false; } + } + + function interpolate(template, values) { + return String(template).replace(/\{(\d+)\}/g, (_, index) => values[Number(index)] ?? ""); + } + + function translate(value, locale = currentLocale, depth = 0) { + const source = String(value ?? ""); + if (!source) return source; + if (normalizeLocale(locale) === "zh-CN") return ZH[source] || source; + if (Object.prototype.hasOwnProperty.call(EN, source)) return EN[source]; + for (const [pattern, template, rawIndexes] of EN_PATTERNS) { + const match = source.match(pattern); + if (match) { + const translatedMatch = match.map((part, index) => index === 0 + ? part + : rawIndexes?.includes(index) ? part + : depth < 6 ? translate(part, locale, depth + 1) : (EN[part] || part)); + return interpolate(template, translatedMatch); + } + } + return source; + } + + function shouldSkipTextNode(node) { + const parent = node?.parentElement; + return Boolean(parent?.closest?.("code, pre, .message-body, .request-summary, .request-questions, .request-json, .request-command, .diff-output, .terminal-output, .session-option-title, .subagent-name, .subagent-summary-label")); + } + + function translateTextNode(node) { + if (!node || shouldSkipTextNode(node)) return; + const current = node.nodeValue; + const previousRendered = textRendered.get(node); + if (!textSources.has(node) || current !== previousRendered) textSources.set(node, current); + const source = textSources.get(node); + const leading = source.match(/^\s*/)?.[0] || ""; + const trailing = source.match(/\s*$/)?.[0] || ""; + const core = source.slice(leading.length, source.length - trailing.length); + if (!core) return; + const translated = translate(core); + const rendered = `${leading}${translated}${trailing}`; + textRendered.set(node, rendered); + if (rendered !== current) node.nodeValue = rendered; + } + + function translateAttributes(element) { + if (!element?.getAttribute || element.closest?.(".message-body, pre, code")) return; + let sources = attributeSources.get(element); + let renderedValues = attributeRendered.get(element); + if (!sources) { sources = new Map(); attributeSources.set(element, sources); } + if (!renderedValues) { renderedValues = new Map(); attributeRendered.set(element, renderedValues); } + for (const attribute of ["title", "aria-label", "placeholder", "data-placeholder"]) { + if (!element.hasAttribute(attribute)) continue; + const current = element.getAttribute(attribute); + if (!sources.has(attribute) || current !== renderedValues.get(attribute)) sources.set(attribute, current); + const source = sources.get(attribute); + const translated = translate(source); + renderedValues.set(attribute, translated); + if (translated !== current) element.setAttribute(attribute, translated); + } + } + + function translateTree(node) { + if (!root?.document || !node) return; + if (node.nodeType === 3) { + translateTextNode(node); + return; + } + if (node.nodeType !== 1 && node.nodeType !== 9 && node.nodeType !== 11) return; + if (node.nodeType === 1) translateAttributes(node); + const walker = root.document.createTreeWalker(node, root.NodeFilter.SHOW_ELEMENT | root.NodeFilter.SHOW_TEXT); + for (let current = walker.nextNode(); current; current = walker.nextNode()) { + if (current.nodeType === 3) translateTextNode(current); + else translateAttributes(current); + } + } + + function applyDocument() { + if (!root?.document) return; + root.document.documentElement.lang = currentLocale; + translateTree(root.document.body); + const selector = root.document.getElementById("localeSelect"); + if (selector && selector.value !== currentLocale) selector.value = currentLocale; + } + + function setLocale(value, options = {}) { + currentLocale = SUPPORTED.has(value) ? value : normalizeLocale(value); + if (options.persist !== false && root?.localStorage && !embeddedMode()) { + try { root.localStorage.setItem(STORAGE_KEY, currentLocale); } catch { /* storage may be disabled */ } + } + applyDocument(); + if (root?.CustomEvent) root.dispatchEvent?.(new root.CustomEvent("aether-vscodex:locale", { detail: { locale: currentLocale } })); + return currentLocale; + } + + function initialLocale() { + if (embeddedMode()) return normalizeLocale(root?.navigator?.language); + try { + const saved = root?.localStorage?.getItem(STORAGE_KEY); + if (SUPPORTED.has(saved)) return saved; + } catch { /* storage may be disabled */ } + return normalizeLocale(root?.navigator?.language); + } + + function start() { + if (!root?.document) return; + currentLocale = initialLocale(); + applyDocument(); + if (typeof root.MutationObserver === "function" && !observer) { + observer = new root.MutationObserver((records) => { + if (currentLocale === "zh-CN") return; + for (const record of records) { + if (record.type === "characterData") translateTextNode(record.target); + else if (record.type === "attributes") translateAttributes(record.target); + else for (const node of record.addedNodes) translateTree(node); + } + }); + observer.observe(root.document.documentElement, { + subtree: true, + childList: true, + characterData: true, + attributes: true, + attributeFilter: ["title", "aria-label", "placeholder", "data-placeholder"], + }); + } + } + + const api = { + locale: () => currentLocale, + normalizeLocale, + setLocale, + start, + t: (value) => translate(value), + translate, + translateTree, + messages: { "zh-CN": Object.freeze({}), "en-US": EN }, + }; + + if (root?.document) { + if (root.document.readyState === "loading") root.document.addEventListener("DOMContentLoaded", start, { once: true }); + else start(); + } + return api; +}); diff --git a/aether-vscodex/public/index.html b/aether-vscodex/public/index.html new file mode 100644 index 000000000..9137c092d --- /dev/null +++ b/aether-vscodex/public/index.html @@ -0,0 +1,303 @@ + + + + + + + Codex + + + +
+ +
+
+
+ + + 等待 VS Code 主机 +
+
+ + + + +
+
+ + + + +
+
+ +
+
+ +
+ + +
+
+ +
+
+ + + 本地模式 + +
+ + +
+
+
+
+
+ + + + + + + + + + + diff --git a/aether-vscodex/public/style.css b/aether-vscodex/public/style.css new file mode 100644 index 000000000..961bf7203 --- /dev/null +++ b/aether-vscodex/public/style.css @@ -0,0 +1,1817 @@ +:root { + color-scheme: dark light; + --color-background: var(--vscode-editor-background, #181818); + --color-background-under: var(--vscode-sideBar-background, #141414); + --color-surface: var(--vscode-sideBar-background, #1b1b1b); + --color-surface-secondary: var(--vscode-input-background, #202020); + --color-surface-tertiary: var(--vscode-titleBar-activeBackground, #1f1f1f); + --color-surface-hover: var(--vscode-list-hoverBackground, #2a2d2e); + --color-terminal-surface: var(--vscode-textCodeBlock-background, #242526); + --color-terminal-border: var(--vscode-input-border, #373738); + --color-border: var(--vscode-panel-border, #303030); + --color-border-strong: var(--vscode-input-border, #454545); + --color-text: var(--vscode-foreground, #d4d4d4); + --color-text-secondary: var(--vscode-descriptionForeground, #9d9d9d); + --color-text-tertiary: color-mix(in srgb, var(--color-text-secondary) 72%, transparent); + --color-info: var(--vscode-textLink-foreground, #75beff); + --color-warning: var(--vscode-editorWarning-foreground, #d7ba7d); + --color-danger: var(--vscode-editorError-foreground, #f48771); + --color-success: var(--vscode-testing-iconPassed, #4ec9a0); + --color-user-message: color-mix(in srgb, var(--color-text) 5%, transparent); + --font-ui: var(--vscode-font-family, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif); + --font-mono: var(--vscode-editor-font-family, ui-monospace, SFMono-Regular, Menlo, Consolas, monospace); + --font-size: var(--vscode-font-size, 13px); + --code-size: var(--vscode-editor-font-size, 12px); + --item-gap: 16px; + --panel-width: 500px; + --composer-shadow: 0 4px 16px rgba(0, 0, 0, .13); + font-family: var(--font-ui); +} + +:root[data-theme="light"] { + color-scheme: light; + --color-background: #ffffff; + --color-background-under: #f3f3f3; + --color-surface: #f8f8f8; + --color-surface-secondary: #f3f3f3; + --color-surface-tertiary: #eeeeee; + --color-surface-hover: #e8e8e8; + --color-terminal-surface: #f5f5f5; + --color-terminal-border: #d6d6d6; + --color-border: #dddddd; + --color-border-strong: #c8c8c8; + --color-text: #242424; + --color-text-secondary: #616161; + --color-info: #006ab1; + --color-warning: #8a6100; + --color-danger: #b42318; + --color-success: #267a3e; + --color-user-message: rgba(0, 0, 0, .045); + --composer-shadow: 0 4px 16px rgba(0, 0, 0, .08); +} + +:root[data-theme="dark"] { color-scheme: dark; } + +* { box-sizing: border-box; } +html, body { height: 100%; } +html { background: var(--color-background-under); } +body { + margin: 0; + min-width: 280px; + overflow: hidden; + background: var(--color-background-under); + color: var(--color-text); + font: var(--font-size)/1.45 var(--font-ui); +} + +button, input, textarea, select { font: inherit; } +button { cursor: pointer; } +button:disabled { cursor: default; opacity: .42; } +button:focus-visible, summary:focus-visible { + outline: 1px solid var(--color-info); + outline-offset: 1px; +} + +.codex-panel { + width: min(100%, var(--panel-width)); + height: 100dvh; + min-height: 0; + margin: 0 auto; + overflow: hidden; + display: flex; + flex-direction: column; + background: var(--color-background); + border-inline: 1px solid color-mix(in srgb, var(--color-border) 78%, transparent); +} + +.connection { display: none; } + +.icon-button { + align-items: center; + justify-content: center; + width: 24px; + height: 24px; + min-height: 24px; + padding: 0; + display: inline-flex; + color: var(--color-text-secondary); + background: transparent; + border: 0; + border-radius: 5px; +} +.icon-button:hover:not(:disabled) { color: var(--color-text); background: var(--color-surface-hover); } +.icon-button svg, .composer svg, .mode-icon { + width: 16px; + height: 16px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.35; +} + +.chat-shell { + position: relative; + flex: 1 1 auto; + min-height: 0; + display: flex; + flex-direction: column; + overflow: hidden; +} +.chat-header { + align-items: center; + flex: 0 0 46px; + min-height: 46px; + display: flex; + justify-content: space-between; + gap: 8px; + padding: 4px 12px; + background: var(--color-background); + border-bottom: 1px solid color-mix(in srgb, var(--color-border) 72%, transparent); +} +.thread-heading { align-items: center; display: flex; flex: 1 1 auto; min-width: 0; gap: 2px; } +.header-back-button { + flex: 0 0 24px; + width: 24px; + height: 24px; + color: var(--color-text-secondary); + opacity: .76; +} +.header-back-button:hover:not(:disabled) { background: transparent; opacity: 1; } +.header-back-button svg { width: 12px; height: 12px; } +.thread-picker-button { + min-width: 0; + max-width: min(360px, 62vw); + min-height: 28px; + padding: 2px 4px; + display: inline-flex; + align-items: center; + flex: 0 1 auto; + color: var(--color-text); + text-align: left; + background: transparent; + border: 0; + border-radius: 6px; +} +.thread-picker-button:hover, +.thread-picker-button[aria-expanded="true"] { background: transparent; opacity: .8; } +.thread-picker-button h2 { + margin: 0; + min-width: 0; + overflow: hidden; + color: var(--color-text); + font-size: 13px; + font-weight: 500; + line-height: 24px; + text-overflow: ellipsis; + white-space: nowrap; +} +.status-text { + position: absolute; + width: 1px; + height: 1px; + overflow: hidden; + clip: rect(0 0 0 0); + clip-path: inset(50%); + white-space: nowrap; +} +.thread-actions { align-items: center; display: flex; flex: 0 0 auto; gap: 2px; } +.header-history-button { margin-right: 1px; } +.new-session-button { margin-left: 1px; color: var(--color-text-secondary); } +.panel-popover { + position: absolute; + z-index: 30; + top: 53px; + right: 10px; + min-width: 150px; + padding: 5px; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 8px; + box-shadow: 0 8px 24px rgba(0, 0, 0, .28); +} +.panel-popover[hidden] { display: none; } +.panel-menu { display: grid; gap: 2px; } +.panel-menu button { + min-height: 28px; + padding: 5px 8px; + color: var(--color-text); + text-align: left; + background: transparent; + border: 0; + border-radius: 5px; + font-size: 11px; +} +.panel-menu button:hover { background: var(--color-surface-hover); } +.details-popover { width: min(320px, calc(100vw - 20px)); min-width: 220px; padding: 10px 12px; } +.popover-title { margin-bottom: 8px; font-size: 11px; font-weight: 600; } +.settings-shortcuts { display: grid; gap: 2px; } +.settings-shortcuts button { + min-height: 30px; + padding: 5px 7px; + display: flex; + align-items: center; + justify-content: space-between; + gap: 12px; + color: var(--color-text-secondary); + text-align: left; + background: transparent; + border: 0; + border-radius: 5px; + font-size: 10px; +} +.settings-shortcuts button:hover { color: var(--color-text); background: var(--color-surface-hover); } +.settings-locale { + min-height: 30px; + padding: 5px 7px; + display: flex; + align-items: center; + justify-content: space-between; + gap: 12px; + color: var(--color-text-secondary); + font-size: 10px; +} +.settings-locale select { + max-width: 124px; + padding: 2px 5px; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 4px; +} +body.embed-aether #localeSetting { display: none; } +.settings-shortcuts button span:last-child { + max-width: 110px; + overflow: hidden; + color: var(--color-text-tertiary); + text-overflow: ellipsis; + white-space: nowrap; +} +.settings-divider { height: 1px; margin: 7px 0; background: var(--color-border); } +.popover-subtitle { margin-bottom: 5px; color: var(--color-text-tertiary); font-size: 9px; } +.details-popover dl { display: grid; grid-template-columns: 48px minmax(0, 1fr); gap: 5px 8px; margin: 0; font-size: 10px; } +.details-popover dt { color: var(--color-text-secondary); } +.details-popover dd { min-width: 0; margin: 0; overflow: hidden; color: var(--color-text); text-overflow: ellipsis; white-space: nowrap; } +.session-picker { + right: auto; + left: 10px; + width: min(330px, calc(100vw - 20px)); + min-width: 0; + padding: 7px; +} +.session-picker-header { display: flex; align-items: center; justify-content: space-between; gap: 8px; padding: 1px 2px 5px 4px; } +.session-picker-header .popover-title { margin: 0; } +.session-picker-refresh { + width: 24px; + height: 24px; + min-height: 24px; + padding: 0; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-text-tertiary); + background: transparent; + border: 0; + border-radius: 5px; +} +.session-picker-refresh:hover:not(:disabled) { color: var(--color-text); background: var(--color-surface-hover); } +.session-picker-refresh svg { width: 14px; height: 14px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.25; } +.session-search { + min-height: 28px; + margin: 1px 2px 5px; + padding: 0 7px; + display: flex; + align-items: center; + gap: 6px; + color: var(--color-text-tertiary); + background: var(--color-background); + border: 1px solid var(--color-border); + border-radius: 6px; +} +.session-search:focus-within { + color: var(--color-text-secondary); + border-color: var(--color-info); +} +.session-search > svg { width: 13px; height: 13px; flex: 0 0 13px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.25; } +.session-search input { + width: 100%; + min-width: 0; + height: 26px; + padding: 0; + color: var(--color-text); + background: transparent; + border: 0; + outline: 0; + font-size: 11px; +} +.session-search input::placeholder { color: var(--color-text-tertiary); } +.session-search input::-webkit-search-cancel-button { appearance: none; } +.session-search-clear { + width: 18px; + height: 18px; + min-height: 18px; + padding: 0; + display: inline-flex; + align-items: center; + justify-content: center; + flex: 0 0 18px; + color: var(--color-text-tertiary); + background: transparent; + border: 0; + border-radius: 4px; +} +.session-search-clear:hover:not(:disabled) { color: var(--color-text); background: var(--color-surface-hover); } +.session-search-clear svg { width: 12px; height: 12px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.35; } +.session-picker-status { min-height: 16px; padding: 0 5px 4px; color: var(--color-text-tertiary); font-size: 10px; } +.session-picker-status[data-tone="warning"] { color: var(--color-warning); } +.session-list { max-height: min(58dvh, 390px); overflow-y: auto; overscroll-behavior: contain; } +.session-list:empty { display: none; } +.session-list:focus-visible { outline: 1px solid color-mix(in srgb, var(--color-info) 72%, transparent); outline-offset: -1px; } +.session-list-empty { padding: 17px 9px; color: var(--color-text-tertiary); font-size: 11px; text-align: center; } +.session-option { + position: relative; + width: 100%; + min-height: 53px; + padding: 7px 8px; + display: grid; + grid-template-columns: minmax(0, 1fr) auto; + gap: 2px 8px; + color: var(--color-text); + text-align: left; + background: transparent; + border: 0; + border-radius: 6px; +} +.session-option:hover:not(:disabled), .session-option[aria-selected="true"], .session-option[data-focused="true"] { background: var(--color-surface-hover); } +.session-option[aria-selected="true"] { box-shadow: inset 2px 0 var(--color-info); } +.session-option[data-focused="true"]::after { + position: absolute; + inset: 1px; + border: 1px solid color-mix(in srgb, var(--color-info) 55%, transparent); + border-radius: 5px; + content: ""; + pointer-events: none; +} +.session-option:disabled { opacity: .55; } +.session-option-title { min-width: 0; overflow: hidden; font-size: 11px; text-overflow: ellipsis; white-space: nowrap; } +.session-option-time { color: var(--color-text-tertiary); font-size: 10px; white-space: nowrap; } +.session-option-meta { grid-column: 1; min-width: 0; overflow: hidden; color: var(--color-text-tertiary); font-size: 9px; text-overflow: ellipsis; white-space: nowrap; } +.session-option-state { grid-column: 2; min-width: 0; display: inline-flex; align-items: center; justify-content: flex-end; gap: 4px; color: var(--color-success); font-size: 9px; white-space: nowrap; } +.session-status-dot { width: 6px; height: 6px; flex: 0 0 6px; background: var(--color-success); border-radius: 50%; } +.session-option[data-status="working"] .session-status-dot, +.session-option[data-status="thinking"] .session-status-dot, +.session-option[data-status="editing"] .session-status-dot { background: var(--color-warning); box-shadow: 0 0 0 2px color-mix(in srgb, var(--color-warning) 18%, transparent); } +.session-option[data-status="waiting"] .session-status-dot, +.session-option[data-status="approval"] .session-status-dot, +.session-option[data-status="unread"] .session-status-dot { background: var(--color-info); } +.session-option[data-status="error"] .session-status-dot { background: var(--color-danger); } +.session-option[data-unread="true"] .session-option-title { font-weight: 600; } +.session-option[data-unread="true"] .session-option-time { color: var(--color-info); } +.session-option[data-switching="true"] .session-option-state { color: var(--color-warning); } +.session-option[data-switching="true"] .session-status-dot { background: var(--color-warning); box-shadow: 0 0 0 2px color-mix(in srgb, var(--color-warning) 18%, transparent); } +.session-option[data-available="false"] .session-option-state { color: var(--color-text-tertiary); } +.session-option[data-available="false"] .session-status-dot { background: var(--color-text-tertiary); box-shadow: none; } +.session-picker[data-switching="true"] .session-option { pointer-events: none; } +.restore-panel { + position: fixed; + right: 14px; + bottom: 14px; + z-index: 50; + min-height: 30px; + padding: 5px 10px; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 7px; + box-shadow: var(--composer-shadow); +} +.restore-panel[hidden] { display: none; } +body.panel-expanded .codex-panel { width: min(100%, 900px); } +body.panel-hidden .codex-panel { display: none; } + +.chat-panel { + position: relative; + flex: 1 1 auto; + min-height: 0; + display: flex; + flex-direction: column; + overflow: hidden; +} +.chat-scroll { + flex: 1 1 auto; + min-height: 0; + overflow-x: hidden; + overflow-y: auto; + overscroll-behavior: contain; + scroll-behavior: auto; + /* JS keeps an explicit timeline anchor during disclosure/stream updates. */ + overflow-anchor: none; + scrollbar-width: thin; + scrollbar-color: var(--color-border) transparent; + scrollbar-gutter: stable; + padding: 16px 18px var(--thread-scroll-padding-bottom, 112px); + scroll-padding-bottom: var(--thread-scroll-padding-bottom, 112px); +} +.chat-scroll:hover, .chat-scroll:focus-within { scrollbar-color: var(--color-border-strong) transparent; } +.chat-scroll::-webkit-scrollbar { width: 10px; } +.chat-scroll::-webkit-scrollbar-track { background: transparent; } +.chat-scroll::-webkit-scrollbar-thumb { background: var(--color-border); border: 3px solid transparent; border-radius: 999px; background-clip: padding-box; } +.chat-scroll:hover::-webkit-scrollbar-thumb, .chat-scroll:focus-within::-webkit-scrollbar-thumb { background: var(--color-border-strong); background-clip: padding-box; } +.output { color: var(--color-text); font-size: var(--font-size); word-break: break-word; } +.output:empty::before { content: none; } + +.message { + position: relative; + width: 100%; + display: flex; + flex-direction: column; + align-items: flex-start; + margin: 0 0 var(--item-gap); +} +.message.user { align-items: flex-end; justify-content: flex-start; } +.message-content { min-width: 0; max-width: min(100%, 520px); overflow-wrap: anywhere; line-height: 1.5; } +.message.user .message-content { + max-width: min(77%, 520px); + padding: 8px 12px; + color: var(--color-text); + background: var(--color-user-message); + border: 0; + border-radius: 16px; +} +.message.assistant .message-content { max-width: 100%; } +.message.system .message-content, .message.error .message-content { color: var(--color-text-secondary); font-size: 12px; } +.message.activity .message-content { font-size: var(--font-size); } +.message.activity .message-details > summary { font-size: var(--font-size); } +.message.error .message-content { color: var(--color-danger); } +.message.tool .message-content { width: 100%; color: var(--color-text-secondary); } +.message.streaming .message-content::after { content: "▍"; margin-left: 2px; color: var(--color-text-secondary); } +.message-meta { + margin-top: 5px; + color: var(--color-text-tertiary); + font-size: 10px; + line-height: 1.35; + font-variant-numeric: tabular-nums; + opacity: 0; + transition: opacity .12s ease; +} +.message:hover .message-meta, .message:focus-within .message-meta { opacity: 1; } +.message.user .message-meta { text-align: right; } +.date-separator { + display: flex; + align-items: center; + justify-content: center; + min-height: 52px; + padding: 16px 0; + margin: 0; + color: var(--color-text-tertiary); + font-size: 13px; + font-weight: 400; + line-height: 20px; + user-select: none; + white-space: nowrap; +} +.date-separator::before, .date-separator::after { display: none; } +.date-separator time { font: inherit; color: inherit; } +.date-separator .date-label { font-weight: 500; } +.date-separator .date-time { font-weight: 400; } + +.turn-divider { + width: 100%; + display: flex; + flex-direction: column; + align-items: stretch; + gap: 4px; + margin: 4px 0 12px; + color: var(--color-text-secondary); +} +.turn-divider-toggle { + align-items: center; + align-self: flex-start; + display: inline-flex; + gap: 4px; + min-height: 22px; + padding: 1px 2px; + color: color-mix(in srgb, var(--color-text) 60%, transparent); + font-size: var(--font-size); + line-height: 20px; + text-align: left; + background: transparent; + border: 1px solid transparent; + border-radius: 4px; +} +.turn-divider-toggle:hover { color: var(--color-text); background: var(--color-surface-hover); } +.turn-divider-toggle svg { + width: 12px; + height: 12px; + flex: 0 0 12px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.35; + transition: transform .12s ease; +} +.turn-divider-toggle[aria-expanded="true"] svg { transform: rotate(90deg); } +.turn-divider-rule { width: 100%; height: 1px; background: var(--color-border); } + +.markdown-body p { margin: 0; } +.markdown-body p + p, .markdown-body ul + p, .markdown-body ol + p, +.markdown-body pre + p, .markdown-body blockquote + p, .markdown-body table + p { margin-top: 11px; } +.markdown-body h1, .markdown-body h2, .markdown-body h3 { + margin: 0 0 8px; + color: inherit; + font-weight: 600; + line-height: 1.3; +} +.markdown-body h1 { font-size: 1.45em; } +.markdown-body h2 { font-size: 1.25em; } +.markdown-body h3 { font-size: 1.1em; } +.markdown-body ul, .markdown-body ol { margin: 8px 0; padding-inline-start: 22px; } +.markdown-body li + li { margin-top: 4px; } +.markdown-body blockquote { margin: 10px 0; padding-inline-start: 11px; color: var(--color-text-secondary); border-inline-start: 2px solid var(--color-border-strong); } +.markdown-body hr { height: 1px; margin: 14px 0; background: var(--color-border); border: 0; } +.markdown-body a { color: var(--color-info); text-decoration: underline; text-underline-offset: 2px; } +.markdown-body code, .message code { padding: 1px 4px; color: inherit; font: var(--code-size)/1.35 var(--font-mono); background: color-mix(in srgb, var(--color-surface-hover) 85%, transparent); border-radius: 4px; } +.markdown-body pre, .message pre { + max-width: 100%; + margin: 10px 0; + padding: 10px 11px; + overflow: auto; + color: var(--color-text); + font: var(--code-size)/1.5 var(--font-mono); + white-space: pre; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 6px; +} +.markdown-body pre code, .message pre code { padding: 0; background: transparent; } +.markdown-body table { width: 100%; margin: 10px 0; border-collapse: collapse; font-size: .95em; } +.markdown-body th, .markdown-body td { padding: 5px 7px; text-align: left; border: 1px solid var(--color-border); } +.markdown-body th { color: var(--color-text); background: var(--color-surface-secondary); font-weight: 600; } +.markdown-body .task-list-item { list-style: none; margin-inline-start: -20px; } +.markdown-body .task-list-item input { width: 13px; height: 13px; margin: 0 6px 0 0; vertical-align: -2px; accent-color: var(--color-info); } + +.message-actions { + position: static; + min-height: 18px; + display: flex; + align-items: center; + gap: 2px; + margin-top: 1px; + opacity: 0; + transition: opacity .12s ease; +} +.message.user .message-actions { justify-content: flex-end; } +.message:hover .message-actions, .message:focus-within .message-actions { opacity: 1; } +.message-action { + min-height: 18px; + padding: 1px 4px; + color: var(--color-text-secondary); + font-size: 12px; + background: transparent; + border: 0; + border-radius: 4px; +} +.message-action:hover { color: var(--color-text); background: var(--color-surface-hover); } +.message-action svg { width: 14px; height: 14px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.25; } +.message-actions .message-meta { margin: 0 3px; opacity: 1; } + +.message-details { + min-width: min(100%, 420px); + overflow: hidden; + background: transparent; + border: 0; +} +.message-details > .details-body { + /* Managed disclosures stay mounted; app.js animates this measured box. */ + display: block; + will-change: height, opacity; +} +.message-details > summary { + min-width: 0; + padding: 2px 0; + overflow: hidden; + color: var(--color-text-secondary); + font-size: 12px; + line-height: 1.5; + list-style: none; + text-overflow: ellipsis; + white-space: nowrap; +} +.message-details > summary::-webkit-details-marker { display: none; } +.message-details > summary::after { + content: "›"; + display: inline-block; + width: 12px; + margin-left: 4px; + color: var(--color-text-secondary); + font-size: 15px; + opacity: 0; + transition: transform .18s cubic-bezier(.33,1,.68,1), opacity .12s ease; +} +.message-details > summary:hover::after, +.message-details > summary:focus-visible::after, +.message-details[data-expanded="true"] > summary::after { opacity: 1; } +.message-details[data-expanded="true"] > summary::after { transform: rotate(90deg); } +.message-details .details-body { + max-height: 260px; + margin-top: 5px; + padding: 0; + overflow: auto; + color: var(--color-text-secondary); + font: inherit; + line-height: 1.5; + white-space: normal; + background: transparent; + border: 0; + border-radius: 0; +} +.message.system .message-details > summary { color: var(--color-text-secondary); } +.message.system .message-details > summary::after { color: var(--color-text-secondary); } +.message.activity { margin-bottom: 4px; } +.message.turn-collapsed { + /* The outer worked-for disclosure removes its activity units from layout. */ + height: 0 !important; + min-height: 0 !important; + margin: 0 !important; + padding: 0 !important; + overflow: hidden !important; + visibility: hidden; + pointer-events: none; + opacity: 0; +} +.message.turn-collapsed * { margin-block: 0 !important; } +.message[data-turn-expanded="false"] .message-details > summary { pointer-events: none; } +.message.activity + .message.assistant { padding-top: 0; border-top: 0; } +.message.activity[data-kind="commentary"] .message-details > summary { display: none; } +.message.activity[data-kind="commentary"] .message-details > .details-body { margin-top: 0; } +.message.activity[data-status="inProgress"] .message-details > summary { color: var(--color-text); } +.message.activity[data-status="inProgress"][data-kind="reasoning"] .message-details > summary { + background: linear-gradient(90deg, var(--color-text-secondary), var(--color-text), var(--color-text-secondary)); + background-size: 220% 100%; + background-clip: text; + -webkit-background-clip: text; + color: transparent; + animation: thinking-shimmer 1.8s linear infinite; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .details-body { + max-height: none; + margin-top: 8px; + padding: 0; + overflow: visible; + color: var(--color-text-secondary); + font: inherit; + white-space: normal; + background: transparent; + border: 0; + border-radius: 0; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .message-details { width: 100%; min-width: 0; } +.message.activity:is([data-kind="tool"], [data-kind="read"]) .message-details > summary { + display: flex; + align-items: center; + gap: 5px; + width: 100%; + max-width: 100%; + color: var(--color-text-secondary); +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .message-details > summary::after { + flex: 0 0 12px; + margin-left: 0; +} +.message.activity[data-kind="subagent"] .message-details { width: 100%; min-width: 0; } +.message.activity[data-kind="subagent"] .message-details > summary { + display: flex; + align-items: center; + gap: 6px; + width: 100%; + max-width: 100%; + color: var(--color-text-secondary); +} +.message.activity[data-kind="subagent"] .message-details > summary::after { + flex: 0 0 12px; + margin-left: 0; +} +.subagent-summary-chip { + min-width: 0; + max-width: min(48%, 192px); + min-height: 24px; + padding: 1px 8px 1px 5px; + display: inline-flex; + align-items: center; + gap: 5px; + overflow: hidden; + color: var(--color-text-secondary); + vertical-align: middle; + background: transparent; + border: 1px solid color-mix(in srgb, var(--color-border-strong) 72%, transparent); + border-radius: 999px; +} +.subagent-summary-chip:hover { color: var(--color-text); border-color: var(--color-border-strong); background: var(--color-surface-hover); } +.subagent-summary-label { + min-width: 0; + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} +.subagent-summary-status { + min-width: 0; + overflow: hidden; + color: var(--color-text-secondary); + text-overflow: ellipsis; + white-space: nowrap; +} +.activity-summary-icon { + width: 13px; + height: 13px; + flex: 0 0 13px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.15; + opacity: .75; +} +.activity-summary-icon.subagent-summary-icon { + fill: currentColor; + stroke: none; + color: var(--color-success); + opacity: .95; +} +.subagent-summary-chip .subagent-summary-icon { width: 14px; height: 14px; flex: 0 0 14px; } +.activity-summary-label { + min-width: 0; + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} +.terminal-body { width: 100%; } +.terminal-shell { + width: 100%; + min-width: 0; + overflow: hidden; + color: var(--color-text-secondary); + background: var(--color-terminal-surface); + border: 1px solid var(--color-terminal-border); + border-radius: 10px; +} +.terminal-shell-header { + display: flex; + align-items: center; + justify-content: space-between; + gap: 8px; + min-height: 25px; + overflow: hidden; + padding: 4px 8px; + color: var(--color-text-secondary); + font: 12px/17px var(--font-ui); + user-select: none; + background: transparent; +} +.terminal-shell-label { + min-width: 0; + overflow: hidden; + padding: 0; + color: inherit; + font: inherit; + text-overflow: ellipsis; + white-space: nowrap; +} +.terminal-file-path { + min-width: 0; + overflow: hidden; + padding: 3px 8px 5px; + color: var(--color-text-secondary); + font: var(--code-size)/1.45 var(--font-mono); + text-overflow: ellipsis; + white-space: nowrap; + border-top: 1px solid color-mix(in srgb, var(--color-border) 70%, transparent); +} +.read-body { + display: grid; + gap: 2px; + padding: 2px 0 3px; + color: var(--color-text-secondary); + font-size: 11px; +} +.read-path-list { display: grid; gap: 1px; } +.read-path-row { + min-width: 0; + display: flex; + align-items: center; + gap: 6px; + padding: 3px 7px; + color: var(--color-text-secondary); + line-height: 1.45; +} +.read-path-row > span:last-child { + min-width: 0; + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} +.read-path-row:hover { color: var(--color-text); background: var(--color-surface-hover); border-radius: 4px; } +.read-path-icon { + width: 13px; + height: 13px; + flex: 0 0 13px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.1; + opacity: .72; +} +.read-output { + max-height: 140px; + margin: 4px 7px 0; + padding: 6px 7px; + overflow: auto; + color: var(--color-text-secondary); + font: var(--code-size)/1.5 var(--font-mono); + white-space: pre-wrap; + overflow-wrap: anywhere; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 5px; +} +.read-empty { padding: 3px 7px; color: var(--color-text-tertiary); } +.terminal-shell-actions { + display: flex; + align-items: center; + flex: 0 0 auto; + gap: 1px; + padding-right: 5px; +} +.terminal-action { + width: 20px; + height: 20px; + min-height: 20px; + padding: 2px; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-text-tertiary); + background: transparent; + border: 0; + border-radius: 4px; + opacity: 0; + transition: opacity .12s ease, color .12s ease, background .12s ease; +} +.terminal-action:focus-visible { opacity: 1; } +.terminal-action:hover { color: var(--color-text); background: var(--color-surface-hover); } +.terminal-action[data-copied="true"] { color: var(--color-success); opacity: 1; } +.terminal-action svg { + width: 13px; + height: 13px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.15; +} +.terminal-command-line { + position: relative; + display: flex; + align-items: baseline; + gap: 1ch; + min-width: 0; + padding: 8px 28px 0 8px; + color: var(--color-text-secondary); + font: var(--code-size)/1.5 var(--font-mono); + white-space: pre-wrap; + word-break: break-word; + cursor: pointer; + will-change: height; + scrollbar-gutter: stable; +} +.terminal-prompt { + flex: 0 0 auto; + color: var(--color-text-tertiary); + user-select: none; +} +.terminal-command-line code { + min-width: 0; + max-width: 100%; + padding: 0; + overflow-wrap: anywhere; + display: -webkit-box; + overflow: hidden; + -webkit-box-orient: vertical; + -webkit-line-clamp: 2; + color: var(--color-text-secondary); + font: inherit; + background: transparent; +} +.terminal-command-line[data-expanded="true"] code { + display: block; + overflow: visible; + -webkit-line-clamp: unset; +} +.terminal-command-chevron { + width: 12px; + height: 12px; + flex: 0 0 12px; + margin-left: auto; + color: var(--color-text-tertiary); + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.2; + transition: transform .18s cubic-bezier(.33,1,.68,1), color .12s ease; +} +.terminal-command-line[data-expanded="true"] .terminal-command-chevron { transform: rotate(180deg); color: var(--color-text-secondary); } +.terminal-command-line:hover .terminal-command-chevron, +.terminal-command-line:focus-visible .terminal-command-chevron { color: var(--color-text); } +.terminal-command-action { + position: absolute; + top: 5px; + right: 5px; +} +.terminal-command-line:hover .terminal-command-action, +.terminal-command-action:focus-visible { opacity: 1; } +.terminal-output-wrap { + position: relative; + min-height: 20px; + margin-top: 0; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output { + display: block; + width: 100%; + min-width: 0; + max-width: 100%; + max-height: 144px; + margin: 0; + padding: 0; + box-sizing: border-box; + overflow-x: auto; + overflow-y: auto; + color: var(--color-text-secondary); + font: var(--code-size)/1.5 var(--font-mono); + font-weight: 500; + white-space: pre; + background: transparent; + border: 0; + border-radius: 0; + scrollbar-width: thin; + scrollbar-color: var(--color-border) transparent; + scrollbar-gutter: stable; +} +.terminal-output-content { + display: block; + width: max-content; + min-width: 100%; + min-height: 18px; + margin: 0; + padding: 8px; + box-sizing: border-box; + color: inherit; + font: inherit; + white-space: inherit; +} +.terminal-output-empty { min-height: 34px; } +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output:hover { + scrollbar-color: var(--color-border-strong) transparent; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output::-webkit-scrollbar { width: 8px; height: 8px; } +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output::-webkit-scrollbar-track { background: transparent; } +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output::-webkit-scrollbar-thumb { + background: var(--color-border); + border: 2px solid transparent; + border-radius: 999px; + background-clip: padding-box; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output:hover::-webkit-scrollbar-thumb { + background: var(--color-border-strong); + background-clip: padding-box; +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output[data-fade-top="true"][data-fade-bottom="true"] { + -webkit-mask-image: linear-gradient(to bottom, transparent, #000 2rem, #000 calc(100% - 2rem), transparent); + mask-image: linear-gradient(to bottom, transparent, #000 2rem, #000 calc(100% - 2rem), transparent); +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output[data-fade-top="true"][data-fade-bottom="false"] { + -webkit-mask-image: linear-gradient(to bottom, transparent, #000 2rem); + mask-image: linear-gradient(to bottom, transparent, #000 2rem); +} +.message.activity:is([data-kind="tool"], [data-kind="read"]) .terminal-output[data-fade-top="false"][data-fade-bottom="true"] { + -webkit-mask-image: linear-gradient(to bottom, #000 calc(100% - 2rem), transparent); + mask-image: linear-gradient(to bottom, #000 calc(100% - 2rem), transparent); +} +.terminal-output-action { + position: absolute; + top: 0; + right: 10px; +} +.terminal-output-wrap:hover .terminal-output-action, +.terminal-output-action:focus-visible { opacity: 1; } +.terminal-no-output { + color: var(--color-text-tertiary); + font: inherit; +} +.terminal-footer { + min-height: 26px; + display: flex; + align-items: center; + justify-content: flex-end; + padding: 2px 10px 4px; + color: var(--color-text-tertiary); + font-size: var(--font-size); + line-height: 20px; +} +.terminal-status { display: inline-flex; align-items: center; gap: 4px; } +.terminal-footer[data-status="failed"], .terminal-footer[data-status="declined"] { color: var(--color-text-tertiary); } +.terminal-footer[data-status="completed"] { color: var(--color-text-tertiary); } +.terminal-status-icon { + width: 12px; + height: 12px; + flex: 0 0 12px; + fill: none; + stroke: currentColor; + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 1.2; +} +.terminal-output .ansi-bold { font-weight: 700; } +.terminal-output .ansi-dim { opacity: .5; } +.terminal-output .ansi-italic { font-style: italic; } +.terminal-output .ansi-underline { text-decoration: underline; text-underline-offset: 2px; } +.terminal-output .ansi-strikethrough { text-decoration: line-through; } +.terminal-output .ansi-black-fg { color: var(--color-codex-terminal-ansi-black, #000); } +.terminal-output .ansi-red-fg { color: var(--color-codex-terminal-ansi-red, #f66); } +.terminal-output .ansi-green-fg { color: var(--color-codex-terminal-ansi-green, #94f494); } +.terminal-output .ansi-yellow-fg { color: var(--color-codex-terminal-ansi-yellow, #f4f47b); } +.terminal-output .ansi-blue-fg { color: var(--color-codex-terminal-ansi-blue, #9e9eff); } +.terminal-output .ansi-magenta-fg { color: var(--color-codex-terminal-ansi-magenta, #db6bdb); } +.terminal-output .ansi-cyan-fg { color: var(--color-codex-terminal-ansi-cyan, #81eeee); } +.terminal-output .ansi-white-fg { color: var(--color-codex-terminal-ansi-white, #d6d6d6); } +.terminal-output .ansi-bright-black-fg { color: var(--color-codex-terminal-ansi-bright-black, #6e6e6e); } +.terminal-output .ansi-bright-red-fg { color: var(--color-codex-terminal-ansi-bright-red, #ffa8a8); } +.terminal-output .ansi-bright-green-fg { color: var(--color-codex-terminal-ansi-bright-green, #0f0); } +.terminal-output .ansi-bright-yellow-fg { color: var(--color-codex-terminal-ansi-bright-yellow, #ffffa8); } +.terminal-output .ansi-bright-blue-fg { color: var(--color-codex-terminal-ansi-bright-blue, #9494ff); } +.terminal-output .ansi-bright-magenta-fg { color: var(--color-codex-terminal-ansi-bright-magenta, #ffb3ff); } +.terminal-output .ansi-bright-cyan-fg { color: var(--color-codex-terminal-ansi-bright-cyan, #adffff); } +.terminal-output .ansi-bright-white-fg { color: var(--color-codex-terminal-ansi-bright-white, #fff); } +.terminal-output .ansi-black-bg { background-color: var(--color-codex-terminal-ansi-black, #000); } +.terminal-output .ansi-red-bg { background-color: var(--color-codex-terminal-ansi-red, #f66); } +.terminal-output .ansi-green-bg { background-color: var(--color-codex-terminal-ansi-green, #94f494); } +.terminal-output .ansi-yellow-bg { background-color: var(--color-codex-terminal-ansi-yellow, #f4f47b); } +.terminal-output .ansi-blue-bg { background-color: var(--color-codex-terminal-ansi-blue, #9e9eff); } +.terminal-output .ansi-magenta-bg { background-color: var(--color-codex-terminal-ansi-magenta, #db6bdb); } +.terminal-output .ansi-cyan-bg { background-color: var(--color-codex-terminal-ansi-cyan, #81eeee); } +.terminal-output .ansi-white-bg { background-color: var(--color-codex-terminal-ansi-white, #d6d6d6); } +.terminal-output .ansi-bright-black-bg { background-color: var(--color-codex-terminal-ansi-bright-black, #6e6e6e); } +.terminal-output .ansi-bright-red-bg { background-color: var(--color-codex-terminal-ansi-bright-red, #ffa8a8); } +.terminal-output .ansi-bright-green-bg { background-color: var(--color-codex-terminal-ansi-bright-green, #0f0); } +.terminal-output .ansi-bright-yellow-bg { background-color: var(--color-codex-terminal-ansi-bright-yellow, #ffffa8); } +.terminal-output .ansi-bright-blue-bg { background-color: var(--color-codex-terminal-ansi-bright-blue, #9494ff); } +.terminal-output .ansi-bright-magenta-bg { background-color: var(--color-codex-terminal-ansi-bright-magenta, #ffb3ff); } +.terminal-output .ansi-bright-cyan-bg { background-color: var(--color-codex-terminal-ansi-bright-cyan, #adffff); } +.terminal-output .ansi-bright-white-bg { background-color: var(--color-codex-terminal-ansi-bright-white, #fff); } +.message.activity[data-kind="reasoning"] .details-body, +.message.activity[data-kind="plan"] .details-body, +.message.activity[data-kind="edit"] .details-body { max-height: 140px; } +.diff-output { + margin: 5px 0 0; + padding: 7px 9px; + overflow: auto; + color: var(--color-text-secondary); + font: var(--code-size)/1.5 var(--font-mono); + white-space: pre; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 6px; +} +.diff-line { display: block; min-height: 1.5em; } +.diff-line.added { color: var(--color-success); background: color-mix(in srgb, var(--color-success) 10%, transparent); } +.diff-line.removed { color: var(--color-danger); background: color-mix(in srgb, var(--color-danger) 10%, transparent); } +.diff-line.context { color: var(--color-text-tertiary); } + +.scroll-to-bottom { + position: absolute; + z-index: 12; + left: 50%; + right: auto; + bottom: calc(var(--thread-scroll-padding-bottom, 112px) + 4px); + width: 32px; + height: 32px; + min-height: 32px; + padding: 0; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-text-secondary); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 50%; + box-shadow: 0 2px 8px rgba(0, 0, 0, .24); + opacity: 0; + pointer-events: none; + transform: translate(-50%, 4px); + transition: opacity .12s ease, transform .12s ease, color .12s ease; +} +.scroll-to-bottom[data-visible="true"] { opacity: 1; pointer-events: auto; transform: translate(-50%, 0); } +.scroll-to-bottom:hover { color: var(--color-text); background: var(--color-surface-hover); } +.scroll-to-bottom svg { width: 16px; height: 16px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.3; } +.scroll-working-dots { display: none; align-items: center; gap: 2px; height: 12px; } +.scroll-working-dots i { width: 3px; height: 3px; background: currentColor; border-radius: 50%; opacity: .28; animation: activity-dot 1.2s ease-in-out infinite; } +.scroll-working-dots i:nth-child(2) { animation-delay: .16s; } +.scroll-working-dots i:nth-child(3) { animation-delay: .32s; } +.scroll-to-bottom[data-working="true"] svg { display: none; } +.scroll-to-bottom[data-working="true"] .scroll-working-dots { display: inline-flex; } + +.live-activity { + display: flex; + align-items: center; + flex: 0 0 auto; + min-height: 28px; + gap: 7px; + padding: 3px 15px 6px; + color: var(--color-warning); + font-size: 11px; +} +.live-status { + position: absolute; + width: 1px; + height: 1px; + min-height: 0; + padding: 0; + overflow: hidden; + clip: rect(0, 0, 0, 0); + white-space: nowrap; + border: 0; +} +.live-activity[hidden] { display: none; } +.live-activity[data-activity="completed"] { color: var(--color-success); } +.live-activity[data-activity="failed"], .live-activity[data-activity="interrupted"] { color: var(--color-danger); } +.activity-spinner { + width: 11px; + height: 11px; + flex: 0 0 11px; + display: inline-block; + border: 1px solid currentColor; + border-right-color: transparent; + border-radius: 50%; +} +.live-activity[data-active="true"] .activity-spinner { animation: activity-spin .8s linear infinite; } +.live-activity[data-active="false"] .activity-spinner { opacity: .45; } +.activity-label { min-width: 0; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; } +.activity-dots { display: inline-flex; align-items: center; gap: 2px; height: 12px; } +.activity-dots i { + width: 2px; + height: 2px; + background: currentColor; + border-radius: 50%; + opacity: .28; + animation: activity-dot 1.2s ease-in-out infinite; +} +.activity-dots i:nth-child(2) { animation-delay: .16s; } +.activity-dots i:nth-child(3) { animation-delay: .32s; } +.activity-elapsed { margin-left: auto; color: var(--color-text-tertiary); font-variant-numeric: tabular-nums; } +@keyframes activity-spin { to { transform: rotate(360deg); } } +@keyframes activity-dot { 0%, 70%, 100% { opacity: .22; transform: translateY(0); } 35% { opacity: .9; transform: translateY(-2px); } } +@keyframes thinking-shimmer { from { background-position: 100% 0; } to { background-position: -100% 0; } } + +.inline-requests { + position: relative; + z-index: 16; + flex: 0 0 auto; + max-height: min(42vh, 360px); + display: grid; + gap: 9px; + padding: 0 12px calc(var(--thread-scroll-padding-bottom, 112px) + 10px); + background: var(--color-background); + overflow: auto; +} +.inline-requests.empty { display: none; } +.inline-requests .request { + padding: 12px 13px; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 10px; + box-shadow: 0 4px 18px rgba(0, 0, 0, .16); +} +.request-title { align-items: center; display: flex; min-width: 0; gap: 7px; } +.request-icon { + width: 18px; + height: 18px; + flex: 0 0 18px; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-warning); + font-size: 11px; + background: color-mix(in srgb, var(--color-warning) 14%, transparent); + border: 1px solid color-mix(in srgb, var(--color-warning) 42%, transparent); + border-radius: 50%; +} +.request-method { min-width: 0; overflow-wrap: anywhere; color: var(--color-text); font-size: 12px; font-weight: 600; } +.request-risk { padding: 2px 6px; color: var(--color-warning); font-size: 9px; border: 1px solid color-mix(in srgb, var(--color-warning) 42%, transparent); border-radius: 999px; white-space: nowrap; } +.request-risk[data-risk="high"] { color: var(--color-danger); border-color: color-mix(in srgb, var(--color-danger) 50%, transparent); } +.request-risk[data-risk="low"] { color: var(--color-success); border-color: color-mix(in srgb, var(--color-success) 45%, transparent); } +.request-id { margin-left: auto; color: var(--color-text-tertiary); font: 10px var(--font-mono); } +.request-summary { margin: 9px 0; color: var(--color-text-secondary); font-size: 11px; line-height: 1.5; white-space: pre-wrap; } +.request-command, .request-json { max-width: 100%; margin: 8px 0; padding: 8px 9px; overflow: auto; color: var(--color-text); font: var(--code-size)/1.45 var(--font-mono); white-space: pre-wrap; word-break: break-word; background: var(--color-background); border: 1px solid var(--color-border); border-radius: 6px; } +.request-questions { display: grid; gap: 8px; margin: 9px 0; } +.request-question { display: grid; gap: 4px; color: var(--color-text); font-size: 11px; } +.request-question input, .request-question select, .request-scope { min-height: 30px; padding: 5px 8px; color: var(--color-text); background: var(--color-background); border: 1px solid var(--color-border-strong); border-radius: 5px; } +.request-scope-wrap { max-width: 220px; display: grid; gap: 4px; margin: 9px 0; color: var(--color-text-secondary); font-size: 10px; } +.request-details { margin-top: 8px; border-top: 1px solid var(--color-border); } +.request-details > summary { padding: 7px 0 2px; color: var(--color-text-secondary); font-size: 10px; cursor: pointer; } +.request-details .request-json { margin-bottom: 7px; } +.request-response { width: 100%; min-height: 58px; margin-top: 7px; padding: 7px 8px; resize: vertical; color: var(--color-text); font: var(--code-size)/1.4 var(--font-mono); background: var(--color-background); border: 1px solid var(--color-border-strong); border-radius: 5px; } +.request-actions { display: flex; flex-wrap: wrap; gap: 6px; margin-top: 10px; padding-top: 10px; border-top: 1px solid var(--color-border); } +.request-actions button { min-height: 28px; padding: 4px 9px; font-size: 10px; } + +.composer { + position: absolute; + z-index: 15; + right: 10px; + bottom: 8px; + left: 16px; + pointer-events: none; + /* The transcript scrolls underneath this fixed composer. Keep a solid + reading band so expanded command output never bleeds through the input. */ + padding-top: 5px; + background: var(--color-background); +} +.composer > .live-activity { + pointer-events: none; + width: min(100%, max-content); + min-height: 26px; + padding: 2px 9px 5px; + color: var(--color-text-secondary); + font-size: 11px; +} +.composer > .live-activity .activity-elapsed { margin-left: 4px; } +.composer > .live-activity[data-activity="thinking"] .activity-spinner, +.composer > .live-activity[data-activity="editing"] .activity-spinner, +.composer > .live-activity[data-activity="running"] .activity-spinner { color: var(--color-success); } +.subagents-panel { + pointer-events: auto; + max-height: min(24dvh, 150px); + margin: 0 8px 4px; + overflow: hidden; + color: var(--color-text-secondary); +} +.subagents-panel[hidden] { display: none; } +.subagents-toggle { + width: 100%; + min-height: 24px; + padding: 1px 2px; + display: flex; + align-items: center; + gap: 5px; + color: var(--color-text-secondary); + text-align: left; + background: transparent; + border: 0; +} +.subagents-toggle:hover { color: var(--color-text); } +.subagents-title { font-size: 10px; } +.subagents-count { color: var(--color-text-tertiary); font-size: 10px; } +.subagents-toggle svg { width: 12px; height: 12px; margin-left: 1px; transition: transform .18s cubic-bezier(.33, 1, .68, 1); } +.subagents-toggle[aria-expanded="true"] svg { transform: rotate(90deg); } +.subagents-list { + display: grid; + gap: 1px; + max-height: 112px; + overflow: auto; + scrollbar-width: thin; + transition: grid-template-rows .2s cubic-bezier(.33, 1, .68, 1), opacity .16s ease; +} +.subagents-panel[data-collapsed="true"] .subagents-list { + max-height: 0; + overflow: hidden; + opacity: 0; + pointer-events: none; +} +.subagent-section { min-width: 0; } +.subagent-section-heading { + padding: 5px 2px 2px; + color: var(--color-text-tertiary); + font-size: 9px; + line-height: 14px; + letter-spacing: 0; +} +.subagent-empty { + padding: 4px 2px 6px; + color: var(--color-text-tertiary); + font-size: 10px; +} +.subagent-section-rows { display: grid; gap: 1px; } +.subagent-more { + min-height: 24px; + margin-top: 3px; + padding: 2px 2px; + color: var(--color-text-tertiary); + font-size: 10px; + text-align: left; + background: transparent; + border: 0; + border-radius: 4px; +} +.subagent-more:hover, +.subagent-more:focus-visible { + color: var(--color-text-secondary); + background: var(--color-surface-hover); + outline: none; +} +.subagent-row { + width: 100%; + min-width: 0; + min-height: 24px; + padding: 2px; + display: grid; + grid-template-columns: 16px minmax(0, 1fr) auto; + align-items: center; + gap: 6px; + color: inherit; + text-align: left; + background: transparent; + border: 0; + border-radius: 5px; + cursor: default; +} +.subagent-row:is(button):hover, +.subagent-row:is(button):focus-visible { + color: var(--color-text); + background: var(--color-surface-hover); + outline: none; +} +.subagent-icon { + width: 14px; + height: 14px; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-text-tertiary); +} +.subagent-icon::before { display: none; } +.subagent-icon svg { width: 14px; height: 14px; fill: none; stroke: currentColor; stroke-width: 1.15; } +.subagent-row[data-status="working"] .subagent-icon { color: var(--color-success); } +.subagent-row[data-status="failed"] .subagent-icon { color: var(--color-danger); } +.subagent-diff-stats { color: var(--color-text-tertiary); font-size: 9px; font-variant-numeric: tabular-nums; } +.subagent-copy { min-width: 0; display: block; } +.subagent-name { overflow: hidden; color: var(--color-text); font-size: 11px; text-overflow: ellipsis; white-space: nowrap; } +.subagent-objective, .subagent-elapsed { display: none; } +.subagent-status { display: inline-flex; align-items: center; gap: 4px; color: var(--color-text-tertiary); font-size: 10px; white-space: nowrap; } +.subagent-elapsed { min-width: 28px; color: var(--color-text-tertiary); font-variant-numeric: tabular-nums; text-align: right; } +.subagent-status-text { white-space: nowrap; } +.subagent-focus { animation: subagent-focus .9s ease-out; } +@keyframes subagent-focus { 0% { background: color-mix(in srgb, var(--color-info) 24%, transparent); } 100% { background: transparent; } } +.subagent-body { + display: grid; + gap: 7px; + padding: 2px 0 4px; + color: var(--color-text-secondary); +} +.subagent-prompt { min-width: 0; font-size: 11px; line-height: 1.5; } +.subagent-prompt p { margin: 0; } +.subagent-action-meta { display: flex; align-items: center; gap: 7px; color: var(--color-text-tertiary); font-size: 10px; } +.subagent-model { padding: 1px 5px; color: var(--color-text-secondary); background: var(--color-surface-hover); border-radius: 4px; font-family: var(--font-mono); } +.subagent-action-rows { display: grid; gap: 2px; padding-top: 2px; border-top: 1px solid var(--color-border); } +.subagent-action-row { + min-width: 0; + padding: 4px 5px; + display: grid; + grid-template-columns: 14px minmax(0, 1fr) auto; + align-items: center; + gap: 5px; + color: var(--color-text-secondary); + font-size: 10px; + border-radius: 4px; +} +.subagent-action-icon { width: 8px; height: 8px; border: 1px solid currentColor; border-radius: 50%; opacity: .65; } +.subagent-action-row[data-status="working"] .subagent-action-icon { color: var(--color-success); border-right-color: transparent; animation: activity-spin .8s linear infinite; } +.subagent-action-row[data-status="failed"] .subagent-action-icon { color: var(--color-danger); } +.subagent-action-label { min-width: 0; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; } +.subagent-action-status { color: var(--color-text-tertiary); white-space: nowrap; } +.subagent-action-note { grid-column: 2 / -1; overflow: hidden; color: var(--color-text-tertiary); text-overflow: ellipsis; white-space: nowrap; } +.composer-surface { + pointer-events: auto; + padding: 9px 10px 7px; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 20px; + box-shadow: var(--composer-shadow); + backdrop-filter: blur(16px); +} +.composer-surface:focus-within { border-color: color-mix(in srgb, var(--color-info) 68%, var(--color-border-strong)); } +.composer-editor { + min-height: 40px; + max-height: 25dvh; + overflow: auto; + padding: 3px 2px; + color: var(--color-text); + line-height: 1.5; + outline: none; + white-space: pre-wrap; + word-break: break-word; +} +.composer-editor:empty::before { content: attr(data-placeholder); color: var(--color-text-tertiary); pointer-events: none; } +.composer-footer { align-items: center; display: flex; justify-content: space-between; gap: 8px; padding-top: 4px; } +.composer-hint { position: relative; align-items: center; display: flex; min-width: 0; gap: 7px; color: var(--color-text-secondary); } +.composer-icon-button { + width: 28px; + height: 28px; + min-height: 28px; + padding: 0; + display: inline-flex; + align-items: center; + justify-content: center; + color: var(--color-text-secondary); + background: transparent; + border: 0; + border-radius: 6px; +} +.composer-icon-button:hover, .composer-icon-button[aria-expanded="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.composer-icon-button svg { width: 16px; height: 16px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.3; } +.permission-chip { + width: auto; + min-width: 28px; + height: 28px; + min-height: 28px; + padding: 0 4px; + align-items: center; + display: inline-flex; + justify-content: flex-start; + gap: 4px; + color: var(--color-warning); + font-size: 10px; + white-space: nowrap; + background: transparent; + border: 0; + border-radius: 6px; +} +.permission-chip:hover, .permission-chip[aria-expanded="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.permission-chip[hidden], .model-picker-button[hidden], .usage-picker[hidden] { display: none; } +.permission-chip svg { width: 14px; height: 14px; fill: none; stroke: currentColor; stroke-linecap: round; stroke-linejoin: round; stroke-width: 1.1; } +.permission-chip #permissionLabel { + min-width: 0; + max-width: 96px; + overflow: hidden; + display: inline-block; + text-overflow: ellipsis; + white-space: nowrap; +} +.permission-chip .permission-chevron { display: inline-block; flex: 0 0 14px; } +.composer-popover { + position: absolute; + z-index: 50; + min-width: 194px; + padding: 5px; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 8px; + box-shadow: 0 10px 28px rgba(0, 0, 0, .34); +} +.composer-popover[hidden] { display: none; } +.composer-popover-heading { + padding: 5px 8px 4px; + color: var(--color-text-tertiary); + font-size: 9px; + line-height: 14px; +} +.composer-plus-menu { left: 0; bottom: 31px; } +.permission-menu { left: 31px; bottom: 34px; width: min(280px, calc(100vw - 28px)); } +.permission-confirm { + position: absolute; + z-index: 60; + left: 31px; + bottom: 34px; + width: min(300px, calc(100vw - 28px)); + padding: 12px; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 9px; + box-shadow: 0 12px 32px rgba(0, 0, 0, .42); +} +.permission-confirm[hidden] { display: none; } +.permission-confirm-title { font-size: 12px; font-weight: 600; } +.permission-confirm p { margin: 7px 0 11px; color: var(--color-text-secondary); font-size: 10px; line-height: 1.5; } +.permission-confirm-actions { display: flex; justify-content: flex-end; gap: 6px; } +.permission-confirm-actions button { min-height: 27px; padding: 4px 10px; color: var(--color-text-secondary); background: transparent; border: 1px solid var(--color-border); border-radius: 5px; font-size: 10px; } +.permission-confirm-actions button:hover { color: var(--color-text); background: var(--color-surface-hover); } +.permission-confirm-actions button.primary { color: var(--color-background); background: var(--color-warning); border-color: var(--color-warning); } +.composer-popover > button { + width: 100%; + min-height: 30px; + padding: 5px 8px; + display: flex; + align-items: center; + justify-content: space-between; + gap: 8px; + color: var(--color-text-secondary); + text-align: left; + background: transparent; + border: 0; + border-radius: 5px; +} +.composer-popover > button:hover, +.composer-popover > button[aria-checked="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.permission-menu > button > span { min-width: 0; font-size: 11px; } +.permission-menu > button > small { min-width: 0; overflow: hidden; color: var(--color-text-tertiary); font-size: 9px; text-overflow: ellipsis; white-space: nowrap; } +.approval-heading { margin-top: 4px; border-top: 1px solid var(--color-border); } +.usage-picker { position: relative; } +.usage-button { + width: 28px; + height: 28px; + min-height: 28px; + padding: 0; + display: inline-grid; + place-items: center; + color: var(--color-text-tertiary); + font-size: 9px; + font-variant-numeric: tabular-nums; + background: transparent; + border: 0; + border-radius: 5px; +} +.usage-button:hover, .usage-button[aria-expanded="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.usage-ring { + --usage-percent: 0%; + position: relative; + width: 16px; + height: 16px; + display: inline-block; + background: conic-gradient(var(--color-info) var(--usage-percent), color-mix(in srgb, var(--color-text-secondary) 24%, transparent) 0); + border-radius: 50%; + transform: rotate(-90deg); +} +.usage-ring[data-level="warning"] { background: conic-gradient(var(--color-warning) var(--usage-percent), color-mix(in srgb, var(--color-text-secondary) 24%, transparent) 0); } +.usage-ring[data-level="critical"] { background: conic-gradient(var(--color-danger) var(--usage-percent), color-mix(in srgb, var(--color-text-secondary) 24%, transparent) 0); } +.usage-ring::after { + position: absolute; + inset: 3px; + content: ""; + background: var(--color-surface-secondary); + border-radius: 50%; +} +.usage-ring > span { + position: absolute; + width: 1px; + height: 1px; + overflow: hidden; + clip: rect(0, 0, 0, 0); +} +.usage-menu { right: 0; bottom: 31px; width: min(245px, calc(100vw - 28px)); } +.usage-summary { padding: 4px 8px 7px; color: var(--color-text-secondary); font-size: 11px; } +.usage-meter { height: 4px; margin: 0 8px 7px; overflow: hidden; background: var(--color-border); border-radius: 999px; } +.usage-meter span { display: block; width: 0; height: 100%; background: var(--color-info); border-radius: inherit; transition: width .2s ease; } +.usage-details { padding: 2px 8px 5px; color: var(--color-text-tertiary); font: 10px/1.45 var(--font-mono); white-space: pre-wrap; } +.composer-actions { align-items: center; display: flex; flex: 0 0 auto; gap: 5px; } +.model-picker { position: relative; } +.model-picker-button { + min-height: 26px; + padding: 2px 5px 2px 7px; + display: inline-flex; + align-items: center; + gap: 4px; + color: var(--color-text-secondary); + background: transparent; + border: 0; + border-radius: 6px; +} +.model-picker-button:hover, .model-picker-button[aria-expanded="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.model-picker-button svg { width: 12px; height: 12px; } +.model-label { max-width: 110px; overflow: hidden; color: inherit; font-size: 11px; text-overflow: ellipsis; white-space: nowrap; } +.model-effort-label { color: var(--color-text-tertiary); font-size: 10px; white-space: nowrap; } +.model-menu { + position: absolute; + z-index: 40; + right: -4px; + bottom: 34px; + width: min(280px, calc(100vw - 28px)); + max-height: min(58dvh, 440px); + padding: 7px; + overflow: auto; + color: var(--color-text); + background: var(--color-surface-secondary); + border: 1px solid var(--color-border-strong); + border-radius: 10px; + box-shadow: 0 10px 28px rgba(0, 0, 0, .32); +} +.model-menu[hidden] { display: none; } +.model-power-view { padding: 5px 4px 7px; } +.model-power-heading, +.model-advanced-toolbar { + display: flex; + align-items: center; + justify-content: space-between; + gap: 8px; + min-height: 28px; + color: var(--color-text-secondary); + font-size: 11px; +} +.model-advanced-toggle, +.model-advanced-back { + min-height: 25px; + padding: 3px 7px; + color: var(--color-text-tertiary); + font-size: 10px; + background: transparent; + border: 0; + border-radius: 5px; +} +.model-advanced-toggle:hover, +.model-advanced-toggle:focus-visible, +.model-advanced-back:hover, +.model-advanced-back:focus-visible { color: var(--color-text); background: var(--color-surface-hover); outline: none; } +.model-power-control { display: grid; grid-template-columns: auto minmax(82px, 1fr) auto; align-items: center; gap: 7px; min-height: 34px; } +.model-power-label { color: var(--color-text-tertiary); font-size: 10px; white-space: nowrap; } +.model-power-slider { + width: 100%; + height: 24px; + margin: 0; + appearance: none; + background: color-mix(in srgb, var(--color-text-secondary) 10%, transparent); + border-radius: 12px; + outline: none; +} +.model-power-slider::-webkit-slider-runnable-track { height: 24px; background: transparent; border-radius: 12px; } +.model-power-slider::-webkit-slider-thumb { + width: 28px; + height: 28px; + margin-top: -2px; + appearance: none; + background: var(--color-text); + border: 1px solid var(--color-border-strong); + border-radius: 50%; + box-shadow: 0 2px 8px rgba(0, 0, 0, .3); +} +.model-power-slider::-moz-range-track { height: 24px; background: transparent; border-radius: 12px; } +.model-power-slider::-moz-range-thumb { + width: 26px; + height: 26px; + background: var(--color-text); + border: 1px solid var(--color-border-strong); + border-radius: 50%; + box-shadow: 0 2px 8px rgba(0, 0, 0, .3); +} +.model-power-slider:focus-visible { box-shadow: 0 0 0 2px color-mix(in srgb, var(--color-info) 65%, transparent); } +.model-power-slider:disabled { opacity: .45; } +.model-power-value { min-height: 16px; color: var(--color-text-tertiary); font-size: 10px; text-align: center; } +.model-advanced-view { padding: 2px 0 0; } +.model-advanced-toolbar { margin: 0 -1px 4px; padding: 0 3px 4px; border-bottom: 1px solid var(--color-border); } +.model-advanced-back { padding-inline: 5px; font-size: 17px; line-height: 1; } +.model-menu-heading { padding: 4px 8px; color: var(--color-text-tertiary); font-size: 10px; } +.model-options { display: grid; gap: 1px; } +.model-option { + width: 100%; + min-height: 40px; + padding: 6px 8px; + display: grid; + grid-template-columns: minmax(0, 1fr) auto; + gap: 2px 8px; + color: var(--color-text); + text-align: left; + background: transparent; + border: 0; + border-radius: 5px; +} +.model-option:hover, .model-option[aria-selected="true"] { background: var(--color-surface-hover); } +.model-option-name { min-width: 0; overflow: hidden; font-size: 11px; text-overflow: ellipsis; white-space: nowrap; } +.model-option-description { grid-column: 1 / -1; overflow: hidden; color: var(--color-text-tertiary); font-size: 9px; text-overflow: ellipsis; white-space: nowrap; } +.model-option-check { color: var(--color-success); font-size: 11px; opacity: 0; } +.model-option[aria-selected="true"] .model-option-check { opacity: 1; } +.effort-heading { margin-top: 4px; border-top: 1px solid var(--color-border); padding-top: 7px; } +.effort-options { display: grid; gap: 1px; padding: 1px 3px 4px; } +.effort-option { + width: 100%; + min-width: 0; + min-height: 28px; + padding: 4px 7px; + display: flex; + align-items: center; + justify-content: space-between; + color: var(--color-text-secondary); + font-size: 10px; + background: transparent; + border: 0; + border-radius: 5px; +} +.effort-option:hover, .effort-option[aria-selected="true"] { color: var(--color-text); background: var(--color-surface-hover); } +.effort-option-check { color: var(--color-success); opacity: 0; } +.effort-option[aria-selected="true"] .effort-option-check { opacity: 1; } +.model-picker[data-pending="true"] .model-picker-button { opacity: .62; } +.compact-action { width: 28px; height: 28px; min-height: 28px; padding: 0; display: inline-flex; align-items: center; justify-content: center; border-radius: 50%; } +.compact-action svg { width: 15px; height: 15px; } +.interrupt-action { display: none; color: var(--color-danger); background: transparent; border: 1px solid color-mix(in srgb, var(--color-danger) 45%, transparent); } +.interrupt-action:hover:not(:disabled) { background: color-mix(in srgb, var(--color-danger) 12%, transparent); } +.steer-action { display: none; } +.steer-action svg { stroke: currentColor; } +.send-button { + width: 28px; + height: 28px; + min-height: 28px; + padding: 0; + display: inline-flex; + align-items: center; + justify-content: center; + color: #1a1a1a; + background: #d4d4d4; + border: 0; + border-radius: 50%; +} +.send-button:hover:not(:disabled) { background: #fff; } +.send-button svg { width: 16px; height: 16px; stroke-width: 1.45; } +.mode-row { + display: flex; + pointer-events: auto; + align-items: center; + justify-content: space-between; + min-height: 26px; + gap: 8px; + padding: 5px 8px 0; + color: var(--color-text-secondary); + font-size: 10px; + background: var(--color-background); +} +.connection-mode-label { min-width: 0; display: inline-flex; align-items: center; gap: 5px; } +.connection-mode-label > span { overflow: hidden; text-overflow: ellipsis; white-space: nowrap; } +.mode-icon { width: 15px; height: 15px; } +.control-mode-switch { + width: 82px; + height: 24px; + padding: 2px; + display: inline-grid; + grid-template-columns: repeat(2, minmax(0, 1fr)); + flex: 0 0 82px; + gap: 1px; + background: var(--color-surface-secondary); + border: 1px solid var(--color-border); + border-radius: 5px; +} +.control-mode-switch button { + min-width: 0; + min-height: 18px; + padding: 0 4px; + overflow: hidden; + color: var(--color-text-tertiary); + text-overflow: ellipsis; + white-space: nowrap; + background: transparent; + border: 0; + border-radius: 3px; + font-size: 9px; + line-height: 18px; +} +.control-mode-switch button:hover:not(:disabled) { color: var(--color-text); background: var(--color-surface-hover); } +.control-mode-switch button[aria-pressed="true"] { + color: var(--color-text); + background: var(--color-background); + box-shadow: 0 0 0 1px color-mix(in srgb, var(--color-border-strong) 72%, transparent); +} +.control-mode-switch button:focus-visible { outline: 1px solid var(--color-info); outline-offset: -1px; } +.control-mode-switch button:disabled { cursor: default; opacity: .46; } +.control-mode-switch button[aria-pressed="true"]:disabled { opacity: 1; } +.control-mode-switch[data-switching="true"] button[data-pending="true"] { color: var(--color-info); opacity: 1; } +.header-back-button[hidden], +.header-history-button[hidden], +.new-session-button[hidden], +.panel-menu button[hidden] { display: none; } +.thread-picker-button:disabled { cursor: default; opacity: 1; } +.thread-picker-button:disabled:hover { opacity: 1; } +body.turn-active .interrupt-action { display: inline-flex; } +body.turn-active .steer-action { display: inline-flex; } +body.turn-active #startTurnButton { display: none; } + +.sr-only { + position: absolute; + width: 1px; + height: 1px; + padding: 0; + margin: -1px; + overflow: hidden; + clip: rect(0, 0, 0, 0); + white-space: nowrap; + border: 0; +} +.compatibility-state { + position: absolute; + width: 1px; + height: 1px; + overflow: hidden; + clip: rect(0, 0, 0, 0); + white-space: nowrap; +} + +@media (max-width: 640px) { + .codex-panel { border-inline: 0; } + .chat-scroll { padding-inline: 12px; } + .composer { right: 8px; bottom: 6px; left: 8px; } + .composer-surface { border-radius: 15px; } + .message.user .message-content { max-width: 77%; } + .request-id { display: none; } +} +@media (prefers-reduced-motion: reduce) { + .live-activity[data-active="true"] .activity-spinner { animation: none; } + .message.activity[data-status="inProgress"][data-kind="reasoning"] .message-details > summary { animation: none; } + .message-details > summary::after { transition: none; } + .activity-dots i, .scroll-working-dots i, .subagent-row[data-status="working"] .subagent-icon::before { animation: none; } +} diff --git a/aether-vscodex/relay/server.js b/aether-vscodex/relay/server.js new file mode 100644 index 000000000..d91c26d53 --- /dev/null +++ b/aether-vscodex/relay/server.js @@ -0,0 +1,2613 @@ +"use strict"; + +const crypto = require("node:crypto"); +const fs = require("node:fs"); +const http = require("node:http"); +const net = require("node:net"); +const path = require("node:path"); +const { spawn } = require("node:child_process"); +const { URL } = require("node:url"); +const { WebSocket, WebSocketServer } = require("ws"); + +const PACKAGE_ROOT = path.resolve(__dirname, ".."); +const VUE_PUBLIC_ROOT = path.join(PACKAGE_ROOT, "web", "dist"); +const PUBLIC_ROOT = fs.existsSync(path.join(VUE_PUBLIC_ROOT, "index.html")) + ? VUE_PUBLIC_ROOT + : path.join(PACKAGE_ROOT, "public"); +const MAX_JSON_BODY = 1024 * 1024; +// Complete attached-session snapshots include the structured message/tool +// projection and can easily exceed 256 KiB. This remains bounded to protect +// the local relay while allowing long Codex conversations to hydrate. +const MAX_WS_PAYLOAD = 16 * 1024 * 1024; +const DEFAULT_EVENT_LIMIT = 2_000; +// Transcript state is represented once in the authoritative control snapshot. +// Keep replay and socket buffering smaller than an unbounded count of maximum +// sized frames so one long conversation cannot amplify into gigabytes. +const DEFAULT_EVENT_BYTE_LIMIT = 16 * 1024 * 1024; +const DEFAULT_REPLAY_BYTE_LIMIT = 2 * 1024 * 1024; +const DEFAULT_CLIENT_BUFFERED_BYTE_LIMIT = MAX_WS_PAYLOAD + 2 * 1024 * 1024; +const MAX_REPLAY_TEXT_BYTES = 64 * 1024; +const MUTATING_METHODS = new Set([ + "control/mode/set", + "thread/start", + "session/new", + "thread/settings/update", + "session/select", + "turn/start", + "turn/steer", + "turn/interrupt", +]); +const ALLOWED_METHODS = new Set(["initialize", "control/mode/get", "session/list", ...MUTATING_METHODS]); +const SERVER_REQUEST_METHODS = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", + "item/permissions/requestApproval", + "item/tool/requestUserInput", + "mcpServer/elicitation/request", + "applyPatchApproval", + "execCommandApproval", +]); +const REMOTE_RESPONSE_METHODS = new Set([ + "approval.respond", + "input.respond", + "server.request.respond", +]); +const DEFAULT_APPROVAL_TIMEOUT_MS = 5 * 60 * 1000; + +function randomToken() { + return crypto.randomBytes(24).toString("base64url"); +} + +function randomId(prefix) { + return `${prefix}_${crypto.randomBytes(9).toString("base64url")}`; +} + +// JSON-RPC treats numeric and string ids as distinct values. Keep that +// distinction in in-memory maps while retaining a small compatibility bridge +// for older browser clients that stringify numeric ids before responding. +function jsonRpcIdKey(id) { + if (typeof id === "number") { + return `number:${Object.is(id, -0) ? "-0" : String(id)}`; + } + if (typeof id === "string") return `string:${id}`; + if (id === null) return "null:"; + return `${typeof id}:${String(id)}`; +} + +function isJsonRpcId(id) { + return typeof id === "string" || typeof id === "number"; +} + +function findTypedMapKey(map, id, valueId = (value) => value?.appId, allowLegacyStringified = false) { + const exact = jsonRpcIdKey(id); + if (map.has(exact)) return exact; + // Before protocol v1, the browser console sent every request id as text. + // Allow that form only when it maps to one unambiguous pending id. If both + // `1` and `"1"` are pending, the exact typed key above wins and no cross-talk + // is possible. + if (!allowLegacyStringified || !isJsonRpcId(id)) return exact; + const text = String(id); + const candidates = []; + for (const [key, value] of map) { + const candidateId = valueId(value); + if (isJsonRpcId(candidateId) && String(candidateId) === text) candidates.push(key); + } + return candidates.length === 1 ? candidates[0] : exact; +} + +function sameJsonRpcId(left, right) { + return typeof left === typeof right && String(left) === String(right) + && (typeof left === "string" || typeof left === "number"); +} + +function responseCommandId(requestId, pendingServerRequests, pendingHostCommands) { + for (const [commandId, pending] of pendingHostCommands) { + if (pending.kind === "server-response" + && (sameJsonRpcId(pending.requestId, requestId) || sameJsonRpcId(pending.responseRequestId, requestId))) return commandId; + } + const base = `response-${String(requestId)}`; + let oppositeKey; + if (typeof requestId === "number") { + oppositeKey = jsonRpcIdKey(String(requestId)); + } else if (typeof requestId === "string") { + const numeric = Number(requestId); + if (Number.isFinite(numeric)) oppositeKey = jsonRpcIdKey(numeric); + } + if (oppositeKey && pendingServerRequests.has(oppositeKey)) return `${base}-${typeof requestId}`; + + const typed = `response-${jsonRpcIdKey(requestId)}`; + const existingBase = pendingHostCommands.get(base); + if (!existingBase || sameJsonRpcId(existingBase.requestId, requestId)) return base; + const existingTyped = pendingHostCommands.get(typed); + if (!existingTyped || sameJsonRpcId(existingTyped.requestId, requestId)) return typed; + // This is only reachable if a caller has manually occupied both stable + // names. Keep the command id deterministic and bounded while avoiding an + // accidental overwrite. + return `${typed}-${crypto.createHash("sha256").update(String(requestId)).digest("hex").slice(0, 8)}`; +} + +function secureEqual(left, right) { + const a = Buffer.from(String(left || "")); + const b = Buffer.from(String(right || "")); + return a.length === b.length && crypto.timingSafeEqual(a, b); +} + +function jsonResponse(response, statusCode, body, extraHeaders = {}) { + const payload = Buffer.from(JSON.stringify(body)); + response.writeHead(statusCode, { + "Content-Type": "application/json; charset=utf-8", + "Content-Length": payload.length, + "Cache-Control": "no-store", + ...extraHeaders, + }); + response.end(payload); +} + +function readJson(request) { + return new Promise((resolve, reject) => { + let size = 0; + const chunks = []; + request.on("data", (chunk) => { + size += chunk.length; + if (size > MAX_JSON_BODY) { + reject(Object.assign(new Error("request body is too large"), { statusCode: 413 })); + request.destroy(); + return; + } + chunks.push(chunk); + }); + request.on("end", () => { + try { + resolve(JSON.parse(Buffer.concat(chunks).toString("utf8") || "{}")); + } catch { + reject(Object.assign(new Error("invalid JSON body"), { statusCode: 400 })); + } + }); + request.on("error", reject); + }); +} + +function redactString(value) { + return String(value) + .replace(/\b(sk-[A-Za-z0-9_-]{12,})\b/g, "[REDACTED_API_KEY]") + .replace(/\b(Bearer\s+)[A-Za-z0-9._~+\/-]{12,}/gi, "$1[REDACTED]") + .replace(/\b(gh[pousr]_[A-Za-z0-9]{20,})\b/g, "[REDACTED_GITHUB_TOKEN]") + .replace(/([?&](?:token|key|secret)=)[^&\s]+/gi, "$1[REDACTED]"); +} + +function redact(value, depth = 0) { + if (depth > 8) return "[TRUNCATED]"; + if (typeof value === "string") return redactString(value); + if (Array.isArray(value)) return value.map((entry) => redact(entry, depth + 1)); + if (value && typeof value === "object") { + const result = {}; + for (const [key, entry] of Object.entries(value)) { + if (/token|authorization|cookie|private.?key|secret/i.test(key)) { + result[key] = "[REDACTED]"; + } else { + result[key] = redact(entry, depth + 1); + } + } + return result; + } + return value; +} + +function applyStructuredMessagesPatch(current, patch) { + if (!Array.isArray(current) || !patch || typeof patch !== "object" || Array.isArray(patch)) return null; + const start = Number(patch.start); + const deleteCount = Number(patch.deleteCount); + if (!Number.isInteger(start) || start < 0 || start > current.length + || !Number.isInteger(deleteCount) || deleteCount < 0 || start + deleteCount > current.length + || !Array.isArray(patch.messages)) return null; + return [ + ...current.slice(0, start), + ...patch.messages, + ...current.slice(start + deleteCount), + ]; +} + +function positiveByteLimit(value, fallback) { + const parsed = Number(value); + return Number.isFinite(parsed) && parsed > 0 ? Math.floor(parsed) : fallback; +} + +function jsonByteLength(value) { + return Buffer.byteLength(JSON.stringify(value), "utf8"); +} + +function boundedReplayText(value) { + const encoded = Buffer.from(String(value), "utf8"); + if (encoded.length <= MAX_REPLAY_TEXT_BYTES) return String(value); + // A cut through a multi-byte code point can add one replacement character, + // which is harmless for this best-effort replay hint. The following control + // snapshot carries the exact authoritative transcript. + return encoded.subarray(encoded.length - MAX_REPLAY_TEXT_BYTES).toString("utf8"); +} + +// Every subscriber receives an authoritative control snapshot after replay. +// Keep transcript-bearing live events rich, but store only their lightweight +// form in the replay ring so streaming a long session cannot retain hundreds +// of duplicate full-history projections. +function compactTranscriptEventForReplay(event) { + if (!event || !["session.snapshot", "output.snapshot", "output.chunk"].includes(event.type)) return event; + const source = event.payload && typeof event.payload === "object" && !Array.isArray(event.payload) + ? event.payload + : {}; + // Use an allow-list rather than deleting known large fields. In particular, + // current attach adapters send `messagesPatch` instead of `messages`, and a + // suffix replacement can itself be nearly as large as the full transcript. + const payload = { projectionInControlSnapshot: true }; + for (const key of [ + "threadId", + "turnId", + "requestId", + "source", + "sourceSeq", + "stream", + "encoding", + "state", + "structureChanged", + ]) { + const value = source[key]; + if (typeof value === "string" || typeof value === "number" || typeof value === "boolean" || value === null) { + payload[key] = value; + } + } + if (event.type === "output.chunk" && typeof source.text === "string" && source.text) { + payload.text = boundedReplayText(source.text); + } + return { ...event, payload }; +} + +// Token usage is telemetry, not an authentication credential. The generic +// redactor intentionally treats any key containing "token" as sensitive, so +// preserve only the numeric usage projection after redacting the rest of a +// session metadata envelope. +const SAFE_USAGE_FIELDS = [ + "totalTokens", + "total_tokens", + "inputTokens", + "input_tokens", + "cachedInputTokens", + "cached_input_tokens", + "cacheWriteInputTokens", + "cache_write_input_tokens", + "outputTokens", + "output_tokens", + "reasoningOutputTokens", + "reasoning_output_tokens", +]; + +function safeUsageNumber(value) { + if (typeof value === "number") return Number.isFinite(value) && value >= 0 ? value : undefined; + if (typeof value === "string" && /^\d+(?:\.\d+)?$/.test(value.trim())) { + const number = Number(value); + return Number.isFinite(number) && number >= 0 ? number : undefined; + } + return undefined; +} + +function safeUsageBreakdown(value) { + if (!value || typeof value !== "object" || Array.isArray(value)) return undefined; + const result = {}; + for (const field of SAFE_USAGE_FIELDS) { + const number = safeUsageNumber(value[field]); + if (number !== undefined) result[field.replace(/_([a-z])/g, (_, letter) => letter.toUpperCase())] = number; + } + return Object.keys(result).length ? result : undefined; +} + +function safeTokenUsage(value) { + if (value === null) return null; + if (!value || typeof value !== "object" || Array.isArray(value)) return undefined; + const source = value.info && typeof value.info === "object" ? value.info + : value.tokenUsage && typeof value.tokenUsage === "object" ? value.tokenUsage + : value.token_usage && typeof value.token_usage === "object" ? value.token_usage : value; + const result = {}; + const total = safeUsageBreakdown(source.total ?? source.total_token_usage ?? source.totalTokenUsage); + const last = safeUsageBreakdown(source.last ?? source.last_token_usage ?? source.lastTokenUsage); + const context = safeUsageNumber(source.modelContextWindow ?? source.model_context_window ?? source.contextWindow ?? source.context_window); + if (total) result.total = total; + if (last) result.last = last; + if (context !== undefined) result.modelContextWindow = context; + return Object.keys(result).length ? result : undefined; +} + +function redactSessionMetadata(value) { + const redacted = redact(value); + if (!redacted || typeof redacted !== "object" || Array.isArray(redacted) || !value || typeof value !== "object") return redacted; + for (const key of ["tokenUsage", "latestTokenUsageInfo"]) { + if (!Object.prototype.hasOwnProperty.call(value, key)) continue; + const usage = safeTokenUsage(value[key]); + if (usage !== undefined) redacted[key] = usage; + } + return redacted; +} + +function normalizeError(error) { + return { + code: error && error.code ? String(error.code) : "relay_error", + message: redactString(error && error.message ? error.message : String(error)), + retryable: Boolean(error && error.retryable), + }; +} + +function parsePort(value, fallback) { + const parsed = Number(value); + return Number.isInteger(parsed) && parsed >= 0 && parsed < 65536 ? parsed : fallback; +} + +/** Parse an optional auth switch without treating an invalid value as false. */ +function parseAuthRequired(value) { + if (typeof value === "boolean") return value; + if (value === 1) return true; + if (value === 0) return false; + if (typeof value !== "string" || value.trim() === "") return undefined; + const normalized = value.trim().toLowerCase(); + if (["1", "true", "yes", "on", "required", "enabled"].includes(normalized)) return true; + if (["0", "false", "no", "off", "none", "disabled", "local"].includes(normalized)) return false; + return undefined; +} + +function hasConfiguredToken(value) { + return typeof value === "string" && value.length > 0; +} + +/** Whether the configured listen address is restricted to this machine. */ +function isLoopbackIPv4(value) { + if (net.isIP(value) !== 4) return false; + const octets = value.split(".").map(Number); + return octets.length === 4 && octets[0] === 127; +} + +function isLoopbackHost(host) { + const normalized = String(host || "").trim().toLowerCase().replace(/^\[|\]$/g, ""); + if (normalized === "localhost" || normalized === "::1") return true; + if (isLoopbackIPv4(normalized)) return true; + return normalized.startsWith("::ffff:") && isLoopbackIPv4(normalized.slice("::ffff:".length)); +} + +/** Node reports IPv4 loopback peers as both 127.x and ::ffff:127.x. */ +function isLoopbackAddress(address) { + const normalized = String(address || "").trim().toLowerCase(); + if (normalized === "::1") return true; + if (isLoopbackIPv4(normalized)) return true; + return normalized.startsWith("::ffff:") && isLoopbackIPv4(normalized.slice("::ffff:".length)); +} + +function isLoopbackRequestHost(request) { + const rawHost = String(request.headers.host || "").trim(); + if (!rawHost) return false; + try { + return isLoopbackHost(new URL(`http://${rawHost}`).hostname); + } catch { + return false; + } +} + +/** Allow browser writes only from the relay's own origin; CLI requests omit Origin. */ +function isAllowedHttpOrigin(request) { + const origin = request.headers.origin; + if (!origin) return true; + const requestHost = String(request.headers.host || "").trim().toLowerCase(); + if (!requestHost) return false; + try { + return new URL(origin).host.toLowerCase() === requestHost; + } catch { + return false; + } +} + +function contentType(filePath) { + const extension = path.extname(filePath).toLowerCase(); + return ( + { + ".html": "text/html; charset=utf-8", + ".js": "text/javascript; charset=utf-8", + ".css": "text/css; charset=utf-8", + ".json": "application/json; charset=utf-8", + ".svg": "image/svg+xml", + ".png": "image/png", + }[extension] || "application/octet-stream" + ); +} + +function outputText(method, params) { + if (!params || typeof params !== "object") return ""; + const candidates = [params.delta, params.text, params.output, params.chunk, params.message]; + if (params.item && typeof params.item === "object") { + candidates.push(params.item.text, params.item.content); + } + const text = candidates.find((candidate) => typeof candidate === "string"); + if (text) return redactString(text); + if (/outputDelta|agentMessage\/delta|plan\/delta|reasoning\/.+Delta/.test(method)) { + return redactString(JSON.stringify(params)); + } + return ""; +} + +function eventTypeForAppMessage(message) { + if (message.id !== undefined && message.method) { + if (message.method === "item/tool/requestUserInput" || message.method === "mcpServer/elicitation/request") { + return "input.requested"; + } + if (SERVER_REQUEST_METHODS.has(message.method)) return "approval.requested"; + return "server.requested"; + } + if (message.id !== undefined) return "app.response"; + const method = message.method || "unknown"; + if (/outputDelta|agentMessage\/delta|plan\/delta|reasoning\/.+Delta/.test(method)) return "output.delta"; + if (method === "turn/started") return "turn.started"; + if (method === "turn/completed") return "turn.completed"; + if (method === "thread/started") return "thread.started"; + if (method === "error") return "app.error"; + return "app.notification"; +} + +function approvalDecisionForResult(result) { + if (result && typeof result === "object" && !Array.isArray(result)) { + const decision = result.decision; + if (typeof decision === "string") { + if (["acceptForSession", "accept", "approved", "approved_for_session", "approved_mcp_policy_amendment"].includes(decision)) return "allow"; + if (["cancel", "abort"].includes(decision)) return "cancel"; + if (["decline", "denied", "timed_out"].includes(decision)) return "deny"; + } + // App-server v2 encodes policy amendments as tagged objects. Require one + // known tag with its documented shape; unknown or mixed objects fail + // closed instead of being interpreted as an approval. + const decisionKind = approvalDecisionKind(decision); + if (decisionKind) return decisionKind; + } + if (result && typeof result === "object" && !Array.isArray(result)) { + if (result.action === "accept") return "allow"; + if (result.action === "cancel") return "cancel"; + if (result.action === "decline") return "deny"; + if (result.permissions && typeof result.permissions === "object") { + return Object.keys(result.permissions).length ? "allow" : "deny"; + } + } + // A custom response is still sent to the bridge; this value is only the + // local policy hint used by RelayHost when it needs a canonical decision. + return "deny"; +} + +function approvalDecisionKind(value) { + if (typeof value === "string") { + if (["allow", "accept", "acceptForSession", "approved", "approved_for_session", "approved_mcp_policy_amendment"].includes(value)) return "allow"; + if (["deny", "decline", "denied", "timed_out"].includes(value)) return "deny"; + if (["cancel", "abort"].includes(value)) return "cancel"; + return undefined; + } + const key = knownDecisionObjectKey(value); + if (key) return key === "denied" ? "deny" : "allow"; + return undefined; +} + +function responseErrorMessage(error) { + if (typeof error === "string") return redactString(error).slice(0, 1_000); + if (error && typeof error === "object" && typeof error.message === "string") return redactString(error.message).slice(0, 1_000); + return "remote response rejected"; +} + +function defaultServerResponse(method, reason = "request timed out") { + if (method === "item/permissions/requestApproval") return { permissions: {}, scope: "turn" }; + if (method === "item/tool/requestUserInput") return { answers: {} }; + if (method === "mcpServer/elicitation/request") return { action: "decline", content: null, _meta: null }; + if (method === "applyPatchApproval" || method === "execCommandApproval") { + return { decision: { denied: { rejection: reason } } }; + } + return { decision: "decline" }; +} + +function normalizeServerResponseForApp(method, response) { + response = normalizeApprovalResponseForApp(method, response); + if (method === "item/permissions/requestApproval") { + const source = isObjectPayload(response) ? response : {}; + const requested = isObjectPayload(source.permissions) ? source.permissions : {}; + const permissions = {}; + for (const [key, value] of Object.entries(requested)) { + if (value !== null && value !== undefined && isObjectPayload(value)) permissions[key] = value; + } + const normalized = { permissions, scope: source.scope === "session" ? "session" : "turn" }; + if (typeof source.strictAutoReview === "boolean") normalized.strictAutoReview = source.strictAutoReview; + return normalized; + } + if (method === "item/tool/requestUserInput") { + if (isObjectPayload(response) && Object.prototype.hasOwnProperty.call(response, "answers")) return response; + return { answers: isObjectPayload(response) ? response : {} }; + } + return response; +} + +const LEGACY_APPROVAL_METHODS = new Set(["applyPatchApproval", "execCommandApproval"]); + +function normalizeApprovalResponseForApp(method, response) { + const v2Approval = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", + ]); + if (!LEGACY_APPROVAL_METHODS.has(method) && !v2Approval.has(method)) return response; + if (!isObjectPayload(response) || !Object.prototype.hasOwnProperty.call(response, "decision")) return response; + const decision = response.decision; + const legacy = LEGACY_APPROVAL_METHODS.has(method); + let normalized = decision; + if (typeof decision === "string") { + if (legacy) { + if (["allow", "accept", "approved"].includes(decision)) normalized = "approved"; + else if (["acceptForSession", "approved_for_session"].includes(decision)) normalized = "approved_for_session"; + else if (["deny", "decline", "denied"].includes(decision)) { + normalized = { denied: { rejection: "Denied remotely" } }; + } else if (["cancel", "abort"].includes(decision)) normalized = "abort"; + else if (decision === "timed_out") normalized = "timed_out"; + } else { + if (["allow", "accept", "approved"].includes(decision)) normalized = "accept"; + else if (["acceptForSession", "approved_for_session"].includes(decision)) normalized = "acceptForSession"; + else if (["deny", "decline", "denied", "timed_out"].includes(decision)) normalized = "decline"; + else if (["cancel", "abort"].includes(decision)) normalized = "cancel"; + else if (decision === "approved_mcp_policy_amendment") normalized = "accept"; + } + } else if (knownDecisionObjectKey(decision)) { + const decisionKey = knownDecisionObjectKey(decision); + if (legacy && decisionKey === "acceptWithExecpolicyAmendment") { + const value = decision.acceptWithExecpolicyAmendment; + normalized = { approved_execpolicy_amendment: { proposed_execpolicy_amendment: value?.execpolicy_amendment ?? value } }; + } else if (legacy && decisionKey === "applyNetworkPolicyAmendment") { + const value = decision.applyNetworkPolicyAmendment; + normalized = { network_policy_amendment: { network_policy_amendment: value?.network_policy_amendment ?? value } }; + } else if (!legacy && decisionKey === "approved_execpolicy_amendment") { + const value = decision.approved_execpolicy_amendment; + normalized = { acceptWithExecpolicyAmendment: { execpolicy_amendment: value?.proposed_execpolicy_amendment ?? value } }; + } else if (!legacy && decisionKey === "network_policy_amendment") { + const value = decision.network_policy_amendment; + normalized = { applyNetworkPolicyAmendment: { network_policy_amendment: value?.network_policy_amendment ?? value } }; + } else if (!legacy && decisionKey === "denied") { + normalized = "decline"; + } + } + return normalized === decision ? response : { ...response, decision: normalized }; +} + +function isValidApprovalResponse(method, response) { + const legacy = LEGACY_APPROVAL_METHODS.has(method); + const v2 = method === "item/commandExecution/requestApproval" || method === "item/fileChange/requestApproval"; + if (!legacy && !v2) return true; + if (!isObjectPayload(response) || !Object.prototype.hasOwnProperty.call(response, "decision")) return false; + const decision = response.decision; + if (typeof decision === "string") { + return legacy + ? ["approved", "approved_for_session", "approved_mcp_policy_amendment", "timed_out", "abort"].includes(decision) + : ["accept", "acceptForSession", "decline", "cancel"].includes(decision); + } + if (!decision || typeof decision !== "object" || Array.isArray(decision)) return false; + const key = knownDecisionObjectKey(decision); + return legacy + ? key === "approved_execpolicy_amendment" || key === "network_policy_amendment" || key === "denied" + : key === "acceptWithExecpolicyAmendment" || key === "applyNetworkPolicyAmendment"; +} + +function knownDecisionObjectKey(value) { + if (!isObjectPayload(value)) return undefined; + const keys = Object.keys(value); + if (keys.length !== 1) return undefined; + const key = keys[0]; + const nested = value[key]; + if (key === "acceptWithExecpolicyAmendment") { + return isObjectPayload(nested) + && Object.keys(nested).every((field) => field === "execpolicy_amendment") + && isStringArray(nested.execpolicy_amendment) ? key : undefined; + } + if (key === "approved_execpolicy_amendment") { + return isObjectPayload(nested) + && Object.keys(nested).every((field) => field === "proposed_execpolicy_amendment") + && isStringArray(nested.proposed_execpolicy_amendment) ? key : undefined; + } + if (key === "applyNetworkPolicyAmendment") { + return isObjectPayload(nested) + && Object.keys(nested).every((field) => field === "network_policy_amendment") + && isNetworkPolicyAmendment(nested.network_policy_amendment) ? key : undefined; + } + if (key === "network_policy_amendment") { + return isObjectPayload(nested) + && Object.keys(nested).every((field) => field === "network_policy_amendment") + && isNetworkPolicyAmendment(nested.network_policy_amendment) ? key : undefined; + } + if (key === "denied") { + return isObjectPayload(nested) + && Object.keys(nested).every((field) => field === "rejection") + && typeof nested.rejection === "string" ? key : undefined; + } + return undefined; +} + +function isStringArray(value) { + return Array.isArray(value) && value.every((entry) => typeof entry === "string"); +} + +function isNetworkPolicyAmendment(value) { + return isObjectPayload(value) + && Object.keys(value).every((field) => field === "host" || field === "action") + && typeof value.host === "string" + && (value.action === "allow" || value.action === "deny"); +} + +class CodexRelay { + constructor(options = {}) { + this.host = options.host || process.env.HOST || "127.0.0.1"; + this.port = parsePort(options.port ?? process.env.PORT, 8787); + this.eventLimit = positiveByteLimit( + options.eventLimit ?? process.env.CODEX_REMOTE_EVENT_LIMIT, + DEFAULT_EVENT_LIMIT, + ); + this.eventByteLimit = positiveByteLimit( + options.eventByteLimit ?? process.env.CODEX_REMOTE_EVENT_BYTE_LIMIT, + DEFAULT_EVENT_BYTE_LIMIT, + ); + this.replayByteLimit = positiveByteLimit( + options.replayByteLimit ?? process.env.CODEX_REMOTE_REPLAY_BYTE_LIMIT, + DEFAULT_REPLAY_BYTE_LIMIT, + ); + this.clientBufferedByteLimit = positiveByteLimit( + options.clientBufferedByteLimit ?? process.env.CODEX_REMOTE_CLIENT_BUFFERED_BYTE_LIMIT, + DEFAULT_CLIENT_BUFFERED_BYTE_LIMIT, + ); + const tokenConfigured = hasConfiguredToken(options.operatorToken) + || hasConfiguredToken(options.viewerToken) + || hasConfiguredToken(options.hostToken) + || hasConfiguredToken(process.env.CODEX_REMOTE_TOKEN) + || hasConfiguredToken(process.env.CODEX_REMOTE_VIEW_TOKEN) + || hasConfiguredToken(process.env.CODEX_REMOTE_HOST_TOKEN); + const explicitAuthRequired = Object.prototype.hasOwnProperty.call(options, "authRequired") + ? parseAuthRequired(options.authRequired) + : parseAuthRequired(process.env.CODEX_REMOTE_AUTH) + ?? parseAuthRequired(process.env.CODEX_REMOTE_AUTH_REQUIRED); + // A loopback relay is a local development tool by default. The moment a + // token is configured, or the relay binds a non-loopback address, retain + // the authenticated behavior. `authRequired` can explicitly enable auth + // for a local relay; disabling it is intentionally limited to loopback. + this.authRequired = explicitAuthRequired ?? (tokenConfigured || !isLoopbackHost(this.host)); + if (!this.authRequired && !isLoopbackHost(this.host)) this.authRequired = true; + this.operatorToken = options.operatorToken || process.env.CODEX_REMOTE_TOKEN || randomToken(); + this.viewerToken = options.viewerToken || process.env.CODEX_REMOTE_VIEW_TOKEN || randomToken(); + this.hostToken = options.hostToken || process.env.CODEX_REMOTE_HOST_TOKEN || this.operatorToken; + const configuredApprovalTimeout = Number(options.approvalTimeoutMs ?? process.env.CODEX_REMOTE_APPROVAL_TIMEOUT_MS); + this.approvalTimeoutMs = Number.isFinite(configuredApprovalTimeout) && configuredApprovalTimeout >= 0 + ? configuredApprovalTimeout + : DEFAULT_APPROVAL_TIMEOUT_MS; + this.generatedOperatorToken = !options.operatorToken && !process.env.CODEX_REMOTE_TOKEN; + this.generatedViewerToken = !options.viewerToken && !process.env.CODEX_REMOTE_VIEW_TOKEN; + this.codexCommand = options.codexCommand || process.env.CODEX_BIN || "codex"; + this.codexArgs = options.codexArgs || this.readCodexArgs(); + this.codexCwd = options.codexCwd || process.env.CODEX_CWD || process.cwd(); + const spawnConfigured = options.spawnCodex === true + || process.env.CODEX_SPAWN === "true" + || process.env.CODEX_SPAWN === "1"; + // Attaching to the already-open VS Code Codex session is the safe default. + // Keep the standalone app-server path available only when it is explicit. + this.mode = options.mode + || process.env.CODEX_REMOTE_MODE + || (options.spawnCodex === false || process.env.CODEX_SPAWN === "false" + ? "host" + : spawnConfigured ? "embedded" : "host"); + this.spawnCodex = this.mode === "host" + ? false + : options.spawnCodex !== undefined + ? options.spawnCodex !== false + : process.env.CODEX_SPAWN !== "false"; + this.events = []; + this.eventSizes = []; + this.eventBytes = 0; + this.audit = []; + this.clients = new Set(); + // At most one outbound VS Code host is active for this MVP session. A + // host is optional: when absent, the relay can run its embedded stdio + // app-server. When present, browser commands are proxied to the host. + this.hostClient = null; + this.pendingHostCommands = new Map(); + this.pendingAppRequests = new Map(); + this.pendingServerRequests = new Map(); + this.commandResults = new Map(); + // Command idempotency is scoped to the embedded app or to one stable host + // session. Keep the scope metadata separate so it never crosses the wire. + this.commandResultScopes = new Map(); + this.hostCommandScope = null; + this.nextSeq = 0; + this.appRequestCounter = 0; + this.appBuffer = ""; + this.appProcess = null; + this.appGeneration = 0; + this.appTerminalGeneration = 0; + this.initializedResult = null; + this.state = { + app: this.spawnCodex ? "starting" : "waiting_for_host", + initialized: false, + activeThreadId: null, + activeTurnId: null, + cwd: this.codexCwd, + outputTail: "", + messages: [], + subagents: [], + // Keep non-sensitive session metadata available for browsers that join + // after the adapter's original session.snapshot event was replayed. + sessionMetadata: null, + // Typed turn/activity projection from the attached VS Code host. Keep + // this in the relay snapshot so a browser that connects after the last + // event still knows whether the conversation is thinking, editing, or + // waiting for approval. + executionStatus: null, + lastError: null, + mode: this.mode, + authRequired: this.authRequired, + hostConnected: false, + hostSessionId: null, + }; + } + + readCodexArgs() { + if (!process.env.CODEX_ARGS_JSON) return ["app-server", "--stdio"]; + try { + const args = JSON.parse(process.env.CODEX_ARGS_JSON); + if (!Array.isArray(args) || !args.every((entry) => typeof entry === "string")) throw new Error(); + return args; + } catch { + throw new Error("CODEX_ARGS_JSON must be a JSON array of strings"); + } + } + + async start() { + this.httpServer = http.createServer((request, response) => this.handleHttp(request, response)); + this.wsServer = new WebSocketServer({ noServer: true, maxPayload: MAX_WS_PAYLOAD }); + this.wsServer.on("connection", (socket, request) => this.handleConnection(socket, request)); + this.httpServer.on("upgrade", (request, socket, head) => this.handleUpgrade(request, socket, head)); + + await new Promise((resolve, reject) => { + const onError = (error) => reject(error); + this.httpServer.once("error", onError); + this.httpServer.listen(this.port, this.host, () => { + this.httpServer.off("error", onError); + resolve(); + }); + }); + + if (this.spawnCodex) this.startCodex(); + return this.address(); + } + + address() { + const address = this.httpServer.address(); + if (!address || typeof address === "string") return { host: this.host, port: this.port }; + return { host: address.address, port: address.port }; + } + + async stop() { + for (const client of this.clients) client.socket.close(1001, "relay shutting down"); + this.clients.clear(); + this.hostClient = null; + this.pendingHostCommands.clear(); + this.pendingAppRequests.clear(); + this.commandResults.clear(); + this.commandResultScopes.clear(); + this.hostCommandScope = null; + for (const pending of this.pendingServerRequests.values()) { + if (pending.timer) clearTimeout(pending.timer); + } + this.pendingServerRequests.clear(); + if (this.appProcess && !this.appProcess.killed) this.appProcess.kill("SIGTERM"); + if (this.wsServer) await new Promise((resolve) => this.wsServer.close(() => resolve())); + if (this.httpServer) await new Promise((resolve) => this.httpServer.close(() => resolve())); + } + + startCodex() { + if (this.appProcess && this.appProcess.exitCode === null && !this.appProcess.killed) return; + this.state.app = "starting"; + const child = spawn(this.codexCommand, this.codexArgs, { + cwd: this.codexCwd, + env: process.env, + stdio: ["pipe", "pipe", "pipe"], + }); + const generation = ++this.appGeneration; + this.appProcess = child; + + child.stdout.setEncoding("utf8"); + child.stderr.setEncoding("utf8"); + child.stdout.on("data", (chunk) => { + if (this.appProcess !== child || this.appGeneration !== generation) return; + this.consumeAppOutput(chunk); + }); + child.stderr.on("data", (chunk) => { + if (this.appProcess !== child || this.appGeneration !== generation) return; + const text = redactString(chunk).slice(0, 32_000); + this.recordEvent("app.stderr", { stream: "stderr", text }); + }); + child.stdin.on("error", (error) => { + if (this.appProcess !== child || this.appGeneration !== generation) return; + this.handleAppProcessExit(child, generation, { error }); + }); + child.on("error", (error) => { + this.handleAppProcessExit(child, generation, { error }); + }); + child.on("exit", (code, signal) => { + this.handleAppProcessExit(child, generation, { code, signal }); + }); + + this.state.app = "initializing"; + const initId = this.nextAppRequestId("initialize"); + this.pendingAppRequests.set(jsonRpcIdKey(initId), { kind: "initialize", method: "initialize" }); + try { + this.sendToApp({ + method: "initialize", + id: initId, + params: { + clientInfo: { + name: "codex-remote-collab", + title: "Codex Remote Collab", + version: "0.1.0", + }, + capabilities: { + experimentalApi: true, + requestAttestation: false, + }, + }, + }); + } catch (error) { + this.handleAppProcessExit(child, generation, { error }); + } + } + + handleAppProcessExit(child, generation, details = {}) { + if (this.appProcess !== child || this.appGeneration !== generation) return; + // ChildProcess can emit both `error` and `exit`; process one terminal + // transition so pending commands/requests are settled exactly once. + if (this.appTerminalGeneration === generation) return; + this.appTerminalGeneration = generation; + this.appProcess = null; + this.appBuffer = ""; + + const previousThreadId = this.state.activeThreadId; + const error = details.error; + const code = details.code; + const signal = details.signal; + const reason = error?.message + || `Codex app-server exited (code=${String(code)}, signal=${String(signal)})`; + const terminalError = { + code: error?.code ? String(error.code) : "app_exited", + message: redactString(reason), + retryable: true, + }; + this.state.app = "offline"; + this.state.initialized = false; + this.state.activeThreadId = null; + this.state.activeTurnId = null; + this.state.lastError = terminalError; + // Results from a dead app-server cannot safely be reused after a restart: + // a command may have applied side effects before the process crashed. + this.clearCommandResults(); + this.recordEvent("app.exited", error + ? { error: terminalError } + : { code, signal, error: terminalError }); + + const pendingCommands = [...this.pendingAppRequests.values()]; + this.pendingAppRequests.clear(); + for (const pending of pendingCommands) { + if (!pending.commandId) continue; + const payload = { + commandId: pending.commandId, + method: pending.method || null, + ok: false, + uncertain: true, + retryable: true, + error: terminalError, + }; + this.cacheCommandResult(pending.commandId, payload, "embedded"); + const event = this.recordEvent("command.result", payload); + if (pending.client) { + this.sendControl(pending.client, { type: "command.result", ...payload, seq: event.seq }); + } + } + + const pendingRequests = [...this.pendingServerRequests.values()]; + for (const pending of pendingRequests) { + this.clearPendingServerRequest(pending.appId); + const isInput = pending.method === "item/tool/requestUserInput" + || pending.method === "mcpServer/elicitation/request"; + this.recordEvent(isInput ? "input.expired" : "approval.expired", { + requestId: pending.appId, + method: pending.method, + reason: "Codex app-server exited", + error: terminalError, + }, { sessionId: previousThreadId || undefined }); + } + } + + nextAppRequestId(label) { + this.appRequestCounter += 1; + return `relay-${label}-${this.appRequestCounter}`; + } + + consumeAppOutput(chunk) { + this.appBuffer += chunk; + while (true) { + const newline = this.appBuffer.indexOf("\n"); + if (newline < 0) break; + const line = this.appBuffer.slice(0, newline).trim(); + this.appBuffer = this.appBuffer.slice(newline + 1); + if (!line) continue; + try { + this.handleAppMessage(JSON.parse(line)); + } catch (error) { + this.recordEvent("app.parse_error", { + error: normalizeError(error), + line: redactString(line).slice(0, 2_000), + }); + } + } + } + + handleAppMessage(rawMessage) { + const message = redact(rawMessage); + if (message.id !== undefined && message.method) { + const requestId = rawMessage.id; + this.clearPendingServerRequest(requestId); + const pending = { + appId: requestId, + method: message.method, + params: message.params || {}, + createdAt: Date.now(), + }; + this.scheduleServerRequestExpiry(requestId, pending); + this.pendingServerRequests.set(jsonRpcIdKey(requestId), pending); + this.recordEvent(eventTypeForAppMessage(message), { + requestId, + method: message.method, + params: message.params || {}, + }); + return; + } + + if (message.id !== undefined) { + const key = jsonRpcIdKey(rawMessage.id); + const pending = this.pendingAppRequests.get(key); + this.pendingAppRequests.delete(key); + + if (pending && pending.kind === "initialize") { + if (message.result) { + this.initializedResult = message.result; + this.state.app = "ready"; + this.state.initialized = true; + // The app-server handshake is ordered: initialized follows the + // successful initialize response. Some versions reject an early + // notification or process it before capabilities are established. + try { + this.sendToApp({ method: "initialized", params: {} }); + } catch (error) { + this.state.app = "error"; + this.state.initialized = false; + this.state.lastError = normalizeError(error); + } + this.recordEvent("app.ready", { result: message.result }); + } else { + this.state.app = "error"; + this.state.lastError = message.error || { message: "Codex initialization failed" }; + this.recordEvent("app.error", { error: this.state.lastError }); + } + return; + } + + if (pending && pending.method === "thread/start" && message.result?.thread?.id) { + this.state.activeThreadId = message.result.thread.id; + this.state.cwd = message.result.cwd || this.state.cwd; + } + if (pending && pending.method === "turn/start" && message.result?.turn?.id) { + this.state.activeTurnId = message.result.turn.id; + } + + const commandPayload = { + commandId: pending?.commandId || null, + method: pending?.method || null, + ok: !message.error, + result: message.result, + error: message.error, + }; + if (pending?.commandId) { + this.cacheCommandResult(pending.commandId, commandPayload, "embedded"); + } + const event = this.recordEvent("command.result", commandPayload); + if (pending?.client) this.sendControl(pending.client, { ...commandPayload, type: "command.result", seq: event.seq }); + return; + } + + this.updateStateFromNotification(message); + this.recordEvent(eventTypeForAppMessage(message), { + method: message.method || "unknown", + params: message.params || {}, + text: outputText(message.method || "", message.params), + emittedAtMs: message.emittedAtMs, + }); + } + + updateStateFromNotification(message) { + const params = message.params || {}; + if (message.method === "thread/started" && params.thread?.id) this.state.activeThreadId = params.thread.id; + if (message.method === "turn/started" && params.turn?.id) this.state.activeTurnId = params.turn.id; + if (message.method === "turn/completed" && (!params.turn?.id || params.turn.id === this.state.activeTurnId)) { + this.state.activeTurnId = null; + } + } + + sendToApp(message) { + if (!this.appProcess + || this.appProcess.exitCode !== null + || this.appProcess.killed + || !this.appProcess.stdin + || this.appProcess.stdin.destroyed + || !this.appProcess.stdin.writable) { + throw Object.assign(new Error("Codex app-server is offline"), { + code: "app_offline", + retryable: true, + }); + } + this.appProcess.stdin.write(`${JSON.stringify(message)}\n`); + } + + recordEvent(type, payload, options = {}) { + const event = { + v: 1, + kind: "event", + id: randomId("evt"), + seq: ++this.nextSeq, + ts: new Date().toISOString(), + type, + sessionId: options.sessionId || this.state.activeThreadId, + payload: redact(payload), + }; + const replayEvent = compactTranscriptEventForReplay(event); + const replayBytes = jsonByteLength(replayEvent); + // An individual event that exceeds the entire replay budget is still + // delivered live and represented by the authoritative state snapshot. Do + // not let it make the in-memory ring exceed its configured hard bound. + if (replayBytes <= this.eventByteLimit) { + this.events.push(replayEvent); + this.eventSizes.push(replayBytes); + this.eventBytes += replayBytes; + } + while (this.events.length > this.eventLimit || this.eventBytes > this.eventByteLimit) { + this.events.shift(); + this.eventBytes -= this.eventSizes.shift() || 0; + } + for (const client of this.clients) { + if (client.authenticated && client.subscribed && client !== options.excludeClient) this.sendControl(client, event); + } + return event; + } + + auditAction(client, action, details, outcome) { + this.audit.push({ + id: randomId("audit"), + ts: new Date().toISOString(), + actor: client?.id || "http", + role: client?.role || "unknown", + action, + details: redact(details), + outcome, + }); + if (this.audit.length > 500) this.audit.splice(0, this.audit.length - 500); + } + + handleUpgrade(request, socket, head) { + let requestUrl; + try { + requestUrl = new URL(request.url, `http://${request.headers.host || "localhost"}`); + } catch { + socket.destroy(); + return; + } + if (requestUrl.pathname !== "/ws" && requestUrl.pathname !== "/v1/connect") { + socket.destroy(); + return; + } + if (!this.authRequired && !isLoopbackRequestHost(request)) { + socket.write("HTTP/1.1 403 Forbidden\r\n\r\n"); + socket.destroy(); + return; + } + const origin = request.headers.origin; + if (origin) { + try { + if (new URL(origin).host.toLowerCase() !== String(request.headers.host || "").toLowerCase()) { + socket.write("HTTP/1.1 403 Forbidden\r\n\r\n"); + socket.destroy(); + return; + } + } catch { + socket.destroy(); + return; + } + } + this.wsServer.handleUpgrade(request, socket, head, (webSocket) => { + this.wsServer.emit("connection", webSocket, request); + }); + } + + handleConnection(socket, request) { + const client = { + id: randomId("client"), + socket, + role: null, + authenticated: false, + subscribed: false, + clientType: null, + sessionId: null, + commandScope: null, + lastSeq: 0, + remoteAddress: request.socket.remoteAddress, + }; + this.clients.add(client); + const authTimer = setTimeout(() => { + if (!client.authenticated) socket.close(1008, "authentication required"); + }, 10_000); + authTimer.unref(); + + socket.on("message", (data, isBinary) => { + if (isBinary) { + socket.close(1003, "JSON text frames only"); + return; + } + let message; + try { + message = JSON.parse(data.toString("utf8")); + } catch { + this.sendControl(client, { type: "error", code: "invalid_json", message: "Invalid JSON frame" }); + return; + } + if (!message || typeof message !== "object" || Array.isArray(message)) { + this.sendControl(client, { type: "error", code: "invalid_frame", message: "JSON frame must be an object" }); + return; + } + try { + this.handleClientMessage(client, message); + } catch (error) { + this.recordEvent("relay.error", { clientId: client.id, error: normalizeError(error) }); + this.sendControl(client, { type: "error", ...normalizeError(error) }); + } + }); + socket.on("close", () => { + clearTimeout(authTimer); + this.clients.delete(client); + if (this.hostClient === client) { + this.hostClient = null; + this.state.hostConnected = false; + this.state.hostSessionId = null; + this.state.app = "offline"; + this.state.initialized = false; + this.state.activeThreadId = null; + this.state.activeTurnId = null; + this.state.lastError = { code: "host_disconnected", message: "VS Code host disconnected" }; + this.recordEvent("host.disconnected", { clientId: client.id }, { excludeClient: client }); + // Commands waiting on a host cannot be completed after its socket is + // gone. Keep their ids reserved briefly so retries get a clear error. + for (const [commandId, pending] of this.pendingHostCommands) { + this.pendingHostCommands.delete(commandId); + this.cacheCommandResult(commandId, { + commandId, + method: pending.method, + ok: false, + uncertain: true, + error: { code: "host_disconnected", message: "VS Code host disconnected" }, + }, pending.commandScope || client.commandScope || this.hostCommandScope); + if (pending.client) { + if (pending.kind === "server-response") { + this.sendControl(pending.client, { + type: "response.rejected", + requestId: pending.responseRequestId ?? pending.requestId, + code: "host_disconnected", + message: "VS Code host disconnected", + retryable: true, + }); + } else { + this.sendControl(pending.client, { + type: "command.result", + commandId, + method: pending.method, + ok: false, + uncertain: true, + error: { code: "host_disconnected", message: "VS Code host disconnected" }, + }); + } + } + } + for (const pending of this.pendingServerRequests.values()) { + if (pending.source !== "host" || pending.hostClient !== client) continue; + this.clearPendingServerRequest(pending.appId); + this.recordEvent( + pending.method === "item/tool/requestUserInput" || pending.method === "mcpServer/elicitation/request" + ? "input.expired" + : "approval.expired", + { requestId: pending.appId, method: pending.method, reason: "VS Code host disconnected" }, + { sessionId: client.sessionId || undefined }, + ); + } + } + if (client.authenticated) this.recordEvent("presence.changed", { clientId: client.id, state: "offline" }); + }); + socket.on("error", () => {}); + } + + roleForToken(token) { + if (secureEqual(token, this.operatorToken)) return "operator"; + if (secureEqual(token, this.viewerToken)) return "viewer"; + return null; + } + + handleClientMessage(client, message) { + if (!message || typeof message !== "object" || Array.isArray(message)) { + this.sendControl(client, { type: "error", code: "invalid_frame", message: "JSON frame must be an object" }); + return; + } + if (!client.authenticated) { + // The VS Code bridge sends a hello frame before its auth frame. Keep the + // hello unauthenticated but remember the client kind and resume cursor. + const localConnection = !this.authRequired && isLoopbackAddress(client.remoteAddress); + const localNoAuthHandshake = localConnection + && (message.kind === "hello" || message.kind === "auth" || message.type === "auth"); + let token = null; + if (message.kind === "hello") { + if (message.protocol !== undefined && Number(message.protocol) !== 1) { + this.sendControl(client, { type: "error", code: "unsupported_protocol", message: "Only protocol 1 is supported" }); + client.socket.close(1002, "unsupported protocol"); + return; + } + client.clientType = message.clientType === "host" ? "host" : "web"; + client.sessionId = typeof message.sessionId === "string" ? message.sessionId : null; + client.lastSeq = Number.isFinite(Number(message.lastSeq)) ? Number(message.lastSeq) : 0; + token = typeof message.token === "string" + ? message.token + : typeof message.accessToken === "string" + ? message.accessToken + : null; + if (!token && !localConnection) return; + } else { + token = message.type === "auth" && typeof message.token === "string" + ? message.token + : message.kind === "auth" && typeof message.accessToken === "string" + ? message.accessToken + : message.kind === "auth" && typeof message.token === "string" + ? message.token + : null; + } + if (!token && !localNoAuthHandshake) { + client.socket.close(1008, "authentication required"); + return; + } + const isHost = client.clientType === "host"; + const role = localNoAuthHandshake + ? (isHost ? "host" : "operator") + : (isHost && secureEqual(token, this.hostToken) ? "operator" : this.roleForToken(token)); + if (!role || (!localNoAuthHandshake && isHost && !secureEqual(token, this.hostToken))) { + this.auditAction(client, "authenticate", {}, "denied"); + client.socket.close(1008, "invalid token"); + return; + } + client.authenticated = true; + client.clientType = client.clientType || "web"; + client.role = isHost ? "host" : role; + if (client.clientType === "host") { + if (this.mode === "embedded") { + this.sendControl(client, { type: "error", code: "host_mode_disabled", message: "Start relay with CODEX_REMOTE_MODE=host (or CODEX_SPAWN=false) for a VS Code host" }); + client.socket.close(1008, "host mode disabled"); + return; + } + if (this.hostClient && this.hostClient !== client) { + this.sendControl(client, { type: "error", code: "host_already_connected", message: "A VS Code host is already connected" }); + client.socket.close(1008, "host already connected"); + return; + } + const hostSessionId = typeof client.sessionId === "string" && client.sessionId.length > 0 + ? client.sessionId + : null; + const commandScope = hostSessionId ? `session:${hostSessionId}` : `connection:${client.id}`; + // A command id is only idempotent within the same host session. A + // reconnect with the same stable session id may reuse the cache; + // another session must never inherit old results (including uncertain + // disconnect results). + if (this.hostCommandScope !== null && this.hostCommandScope !== commandScope) { + this.clearCommandResults(); + } + client.commandScope = commandScope; + this.hostCommandScope = commandScope; + if (this.state.hostSessionId && this.state.hostSessionId !== client.sessionId) { + this.state.activeThreadId = null; + this.state.activeTurnId = null; + } + this.hostClient = client; + this.state.hostConnected = true; + this.state.hostSessionId = client.sessionId; + this.recordEvent("host.connected", { clientId: client.id, sessionId: client.sessionId }, { excludeClient: client }); + } + this.auditAction(client, "authenticate", {}, "accepted"); + this.sendControl(client, { + type: "auth.ok", + clientId: client.id, + role: client.role, + clientType: client.clientType, + protocol: 1, + authRequired: this.authRequired, + latestSeq: this.nextSeq, + }); + return; + } + + // Host frames use the versioned relay contract; browser frames use the + // compact `type` contract. Host events are ingested and re-sequenced here + // instead of echoed back to the host. + if (message.kind === "hello") { + const announcedType = message.clientType === "host" ? "host" : "web"; + if (announcedType !== client.clientType) { + this.sendControl(client, { type: "error", code: "client_type_immutable", message: "clientType cannot change after authentication" }); + return; + } + client.sessionId = typeof message.sessionId === "string" ? message.sessionId : client.sessionId; + return; + } + if (message.kind === "auth" || message.type === "auth") { + return; + } + if (message.kind === "ack") return; + if (message.kind === "event" && client.clientType === "host") { + this.ingestHostEvent(client, message); + return; + } + + if (message.type === "subscribe") { + this.subscribe(client, Number(message.fromSeq || 0)); + return; + } + if (message.type === "ping") { + this.sendControl(client, { type: "pong", ts: new Date().toISOString(), latestSeq: this.nextSeq }); + return; + } + // Browser clients historically used compact `{type:"command", method, + // params}` / `{type:"respond", requestId, result}` frames. The bridge + // contract is versioned and uses `{kind:"command", type, payload}` (and + // approval/input response command names). Normalize both forms at this + // boundary; host clients are event producers and must not issue relay + // commands back into themselves. + if (client.clientType !== "host") { + const command = normalizeBrowserCommand(message); + if (command) { + if (REMOTE_RESPONSE_METHODS.has(command.method)) { + this.dispatchServerResponse(client, normalizeBrowserResponse(message, command.method)); + } else { + this.dispatchCommand(client, command); + } + return; + } + const response = normalizeBrowserResponse(message); + if (response) { + this.dispatchServerResponse(client, response); + return; + } + } + this.sendControl(client, { type: "error", code: "unknown_frame", message: "Unknown frame type" }); + } + + subscribe(client, fromSeq) { + const firstAvailable = this.events.length ? this.events[0].seq : this.nextSeq + 1; + const replayEvents = []; + let replayBytes = 0; + let replayTooLarge = false; + if (fromSeq + 1 >= firstAvailable) { + for (const event of this.events) { + if (event.seq <= fromSeq) continue; + const size = jsonByteLength(event); + if (replayBytes + size > this.replayByteLimit) { + replayTooLarge = true; + break; + } + replayEvents.push(event); + replayBytes += size; + } + } + if (fromSeq + 1 < firstAvailable || replayTooLarge) { + this.sendControl(client, { + type: "resync.required", + requestedFromSeq: fromSeq, + firstAvailableSeq: firstAvailable, + ...(replayTooLarge ? { reason: "replay_too_large" } : {}), + }); + } else { + for (const event of replayEvents) this.sendControl(client, event); + } + client.subscribed = true; + // Other subscribers need the presence transition, while the joining + // client receives the same fact in the clients list of its snapshot. Add + // it first so `latestSeq` covers every authoritative state transition. + this.recordEvent("presence.changed", { clientId: client.id, role: client.role, state: "online" }, { excludeClient: client }); + this.sendControl(client, { type: "session.snapshot", snapshot: this.snapshot() }); + } + + ingestHostEvent(client, frame) { + const sourceType = typeof frame.type === "string" ? frame.type : "app.notification"; + const sourcePayload = frame.payload && typeof frame.payload === "object" ? frame.payload : {}; + const sourceSeq = Number.isFinite(Number(frame.seq)) ? Number(frame.seq) : undefined; + const sessionId = client.sessionId || frame.sessionId || this.state.hostSessionId || undefined; + + const executionStatus = sourcePayload.executionStatus && typeof sourcePayload.executionStatus === "object" + ? sourcePayload.executionStatus + : frame.status && typeof frame.status === "object" + ? frame.status + : sourcePayload.status && typeof sourcePayload.status === "object" + ? sourcePayload.status + : null; + if (executionStatus) this.state.executionStatus = redact(executionStatus); + + // Keep the browser snapshot useful even though the host deliberately uses + // a normalized event vocabulary instead of raw app-server notifications. + if (sourceType === "session.created" && sourcePayload.thread && typeof sourcePayload.thread === "object") { + const id = sourcePayload.thread.id; + if (typeof id === "string") this.state.activeThreadId = id; + } + if (sourceType === "session.snapshot") { + const threadId = sourcePayload.threadId || sourcePayload.thread?.id; + const turnId = sourcePayload.turnId || sourcePayload.turn?.id; + if (typeof threadId === "string") this.state.activeThreadId = threadId; + if (typeof turnId === "string") this.state.activeTurnId = turnId; + if (sourcePayload.threadId === null || sourcePayload.thread === null) this.state.activeThreadId = null; + if (sourcePayload.turnId === null || sourcePayload.turn === null) this.state.activeTurnId = null; + this.state.app = "ready"; + this.state.initialized = true; + this.state.lastError = null; + if (typeof sourcePayload.outputTail === "string") this.state.outputTail = sourcePayload.outputTail; + if (Array.isArray(sourcePayload.messages)) this.state.messages = sourcePayload.messages; + if (sourcePayload.metadata && typeof sourcePayload.metadata === "object" && !Array.isArray(sourcePayload.metadata)) { + const metadata = sourcePayload.metadata; + this.state.sessionMetadata = redactSessionMetadata({ + ...(typeof metadata.title === "string" ? { title: metadata.title } : {}), + ...(typeof metadata.name === "string" ? { name: metadata.name } : {}), + ...(typeof metadata.cwd === "string" ? { cwd: metadata.cwd } : {}), + ...(typeof metadata.mode === "string" ? { mode: metadata.mode } : {}), + ...(metadata.controlMode === "sync" || metadata.controlMode === "async" ? { controlMode: metadata.controlMode } : {}), + ...(Number.isSafeInteger(metadata.modeEpoch) && metadata.modeEpoch >= 0 ? { modeEpoch: metadata.modeEpoch } : {}), + ...(metadata.capabilities && typeof metadata.capabilities === "object" && !Array.isArray(metadata.capabilities) ? { capabilities: metadata.capabilities } : {}), + ...(typeof metadata.source === "string" ? { source: metadata.source } : {}), + ...(typeof metadata.historyComplete === "boolean" ? { historyComplete: metadata.historyComplete } : {}), + ...(typeof metadata.waitingForSession === "boolean" ? { waitingForSession: metadata.waitingForSession } : {}), + ...(typeof metadata.attachReady === "boolean" ? { attachReady: metadata.attachReady } : {}), + ...(typeof metadata.model === "string" ? { model: metadata.model } : {}), + ...(typeof metadata.latestModel === "string" ? { latestModel: metadata.latestModel } : {}), + ...(typeof metadata.effort === "string" || metadata.effort === null ? { effort: metadata.effort } : {}), + ...(typeof metadata.latestReasoningEffort === "string" || metadata.latestReasoningEffort === null ? { latestReasoningEffort: metadata.latestReasoningEffort } : {}), + ...(typeof metadata.modelName === "string" ? { modelName: metadata.modelName } : {}), + ...(typeof metadata.modelProvider === "string" ? { modelProvider: metadata.modelProvider } : {}), + ...(typeof metadata.approvalPolicy === "string" ? { approvalPolicy: metadata.approvalPolicy } : {}), + ...(typeof metadata.approvalsReviewer === "string" ? { approvalsReviewer: metadata.approvalsReviewer } : {}), + ...(typeof metadata.sandboxPolicy === "string" ? { sandboxPolicy: metadata.sandboxPolicy } : {}), + ...(metadata.approvalPolicy && typeof metadata.approvalPolicy === "object" ? { approvalPolicy: metadata.approvalPolicy } : {}), + ...(metadata.approvalsReviewer === null ? { approvalsReviewer: null } : {}), + ...(metadata.sandboxPolicy && typeof metadata.sandboxPolicy === "object" ? { sandboxPolicy: metadata.sandboxPolicy } : {}), + ...(typeof metadata.permissions === "string" || (metadata.permissions && typeof metadata.permissions === "object") || metadata.permissions === null ? { permissions: metadata.permissions } : {}), + ...(typeof metadata.currentPermissions === "string" || (metadata.currentPermissions && typeof metadata.currentPermissions === "object") || metadata.currentPermissions === null ? { currentPermissions: metadata.currentPermissions } : {}), + ...(Array.isArray(metadata.runtimeWorkspaceRoots) ? { runtimeWorkspaceRoots: metadata.runtimeWorkspaceRoots } : {}), + ...(typeof metadata.workedDurationMs === "number" ? { workedDurationMs: metadata.workedDurationMs } : {}), + ...(typeof metadata.firstTurnWorkItemStartedAtMs === "number" ? { firstTurnWorkItemStartedAtMs: metadata.firstTurnWorkItemStartedAtMs } : {}), + ...(typeof metadata.finalAssistantStartedAtMs === "number" ? { finalAssistantStartedAtMs: metadata.finalAssistantStartedAtMs } : {}), + ...(metadata.tokenUsage && typeof metadata.tokenUsage === "object" ? { tokenUsage: metadata.tokenUsage } : metadata.tokenUsage === null ? { tokenUsage: null } : {}), + ...(metadata.latestTokenUsageInfo && typeof metadata.latestTokenUsageInfo === "object" ? { latestTokenUsageInfo: metadata.latestTokenUsageInfo } : metadata.latestTokenUsageInfo === null ? { latestTokenUsageInfo: null } : {}), + ...(metadata.threadSettings && typeof metadata.threadSettings === "object" ? { threadSettings: metadata.threadSettings } : {}), + ...(Array.isArray(metadata.availableModels) ? { availableModels: metadata.availableModels } : {}), + ...(Array.isArray(metadata.models) ? { models: metadata.models } : {}), + ...(Array.isArray(metadata.subagents) ? { subagents: metadata.subagents } : {}), + ...(typeof metadata.parentThreadId === "string" ? { parentThreadId: metadata.parentThreadId } : {}), + ...(typeof metadata.agentNickname === "string" ? { agentNickname: metadata.agentNickname } : {}), + ...(typeof metadata.agentRole === "string" ? { agentRole: metadata.agentRole } : {}), + }); + } + if (Array.isArray(sourcePayload.subagents)) this.state.subagents = redact(sourcePayload.subagents); + else if (Array.isArray(sourcePayload.metadata?.subagents)) this.state.subagents = redact(sourcePayload.metadata.subagents); + for (const request of Array.isArray(sourcePayload.pendingRequests) ? sourcePayload.pendingRequests : []) { + if (!request || typeof request !== "object" || request.requestId === undefined) continue; + const requestId = request.requestId; + this.clearPendingServerRequest(requestId); + const pending = { + appId: requestId, + method: typeof request.method === "string" ? request.method : "server.request", + params: request.params && typeof request.params === "object" ? request.params : {}, + ...(typeof request.risk === "string" ? { risk: request.risk } : {}), + ...(typeof request.summary === "string" ? { summary: request.summary } : {}), + createdAt: Number.isFinite(Number(request.createdAt)) ? Number(request.createdAt) : Date.now(), + ...(Number.isFinite(Number(request.expiresAt)) ? { expiresAt: Number(request.expiresAt) } : {}), + source: "host", + hostClient: client, + commandHash: typeof request.commandHash === "string" ? request.commandHash : undefined, + }; + if (pending.expiresAt && pending.expiresAt <= Date.now()) continue; + this.pendingServerRequests.set(jsonRpcIdKey(requestId), pending); + } + } + if (sourceType === "session.switching") { + const targetThreadId = sourcePayload.targetThreadId || sourcePayload.threadId; + if (typeof targetThreadId === "string") this.state.activeThreadId = targetThreadId; + // The old transcript belongs to the previous thread. Clear it before + // the target's authoritative snapshot arrives so a remote picker never + // briefly renders messages from two sessions together. + this.state.activeTurnId = null; + this.state.outputTail = ""; + this.state.messages = []; + this.state.subagents = []; + this.state.sessionMetadata = null; + this.state.executionStatus = null; + } + if (sourceType === "session.selected") { + const selectedThreadId = sourcePayload.threadId || sourcePayload.activeThreadId; + if (typeof selectedThreadId === "string") this.state.activeThreadId = selectedThreadId; + } + if (sourceType === "output.snapshot") { + if (typeof sourcePayload.text === "string") this.state.outputTail = sourcePayload.text; + if (Array.isArray(sourcePayload.messages)) this.state.messages = sourcePayload.messages; + if (Array.isArray(sourcePayload.subagents)) this.state.subagents = redact(sourcePayload.subagents); + if (sourcePayload.metadata && typeof sourcePayload.metadata === "object" && !Array.isArray(sourcePayload.metadata)) { + // Output snapshots from older hosts occasionally carry the metadata + // projection instead of a separate session.snapshot event. Preserve + // the safe projection so model, permission, and usage controls remain + // available after reconnect. + this.state.sessionMetadata = redactSessionMetadata(sourcePayload.metadata); + } + } else if (sourceType === "output.chunk") { + // New attach adapters carry the complete role-aware projection alongside + // the append-only delta. Preserve both so reconnects do not flatten + // reasoning, tools, edits, or Markdown into one assistant transcript. + if (typeof sourcePayload.outputTail === "string") this.state.outputTail = sourcePayload.outputTail; + else if (typeof sourcePayload.text === "string") this.state.outputTail = `${this.state.outputTail || ""}${sourcePayload.text}`.slice(-32_000); + if (Array.isArray(sourcePayload.messages)) this.state.messages = sourcePayload.messages; + else { + const patchedMessages = applyStructuredMessagesPatch(this.state.messages, sourcePayload.messagesPatch); + if (patchedMessages) this.state.messages = patchedMessages; + } + if (Array.isArray(sourcePayload.subagents)) this.state.subagents = redact(sourcePayload.subagents); + if (sourcePayload.metadata && typeof sourcePayload.metadata === "object" && !Array.isArray(sourcePayload.metadata)) { + this.state.sessionMetadata = redactSessionMetadata(sourcePayload.metadata); + } + } + if (sourceType === "task.started") { + const id = sourcePayload.turnId || (sourcePayload.turn && sourcePayload.turn.id); + if (typeof id === "string") this.state.activeTurnId = id; + if (typeof sourcePayload.threadId === "string") this.state.activeThreadId = sourcePayload.threadId; + } + if (sourceType === "task.finished" || sourceType === "task.cancelled") { + this.state.activeTurnId = null; + } + + if (sourceType === "connection.opened") { + this.state.app = "ready"; + this.state.initialized = true; + this.state.hostConnected = true; + this.state.lastError = null; + } else if (sourceType === "connection.closed") { + // A replaced host socket can have one frame already queued in the + // transport. A stale close must not mark the newly connected host + // offline; acknowledge it so the old bridge does not retry forever. + if (this.hostClient !== client) { + if (sourceSeq !== undefined) { + this.sendControl(client, { v: 1, kind: "ack", sessionId: sessionId || "", seq: sourceSeq }); + } + return; + } + this.state.app = "offline"; + this.state.initialized = false; + this.state.activeTurnId = null; + this.state.outputTail = ""; + this.state.messages = []; + this.state.subagents = []; + this.state.sessionMetadata = null; + this.state.executionStatus = null; + this.state.lastError = { + code: "app_unavailable", + message: typeof sourcePayload.message === "string" + ? redactString(sourcePayload.message) + : "VS Code host app-server disconnected", + retryable: true, + }; + // This event is emitted by the authenticated host bridge when its local + // app-server exits. The relay socket remains usable, so clean only the + // app-scoped pending work here; transport close has its own handler. + this.handleHostAppUnavailable(client, sessionId, this.state.lastError.message); + } + + // RelayHost emits normalized approval/input events and keeps the original + // app-server request id in payload. Store it centrally so exactly one + // browser response can be routed back to that host. + if (sourceType === "approval.requested" + || sourceType === "input.requested" + || sourceType === "server.requested" + || sourceType === "server.request") { + const requestId = sourcePayload.requestId; + if (requestId !== undefined) { + const key = jsonRpcIdKey(requestId); + this.clearPendingServerRequest(requestId); + const pending = { + appId: requestId, + method: typeof sourcePayload.method === "string" ? sourcePayload.method : sourceType, + params: sourcePayload.params || sourcePayload, + commandHash: typeof sourcePayload.commandHash === "string" ? sourcePayload.commandHash : undefined, + risk: typeof sourcePayload.risk === "string" ? sourcePayload.risk : undefined, + summary: typeof sourcePayload.summary === "string" ? sourcePayload.summary : undefined, + expiresAt: Number.isFinite(Number(sourcePayload.expiresAt)) ? Number(sourcePayload.expiresAt) : undefined, + createdAt: Date.now(), + source: "host", + hostClient: client, + }; + // The VS Code adapter owns its local approval timer. Keeping a second + // timer in the relay would race the adapter's JSON-RPC response. + this.pendingServerRequests.set(key, pending); + } + } + if (sourceType === "approval.expired" || sourceType === "input.expired" || sourceType === "server.expired") { + const requestId = sourcePayload.requestId; + if (requestId !== undefined) { + const key = jsonRpcIdKey(requestId); + const pending = this.pendingServerRequests.get(key); + if (pending?.source === "host" && pending.hostClient === client) { + this.clearPendingServerRequest(pending.appId); + } + } + } + if (sourceType === "approval.resolved" + || sourceType === "input.resolved" + || sourceType === "server.responded" + || sourceType === "server.resolved") { + const requestId = sourcePayload.requestId; + if (requestId !== undefined) this.clearPendingServerRequest(requestId); + } + + if (sourceType === "command.accepted" || sourceType === "command.rejected" || sourceType === "command.result") { + const commandId = sourcePayload.commandId; + const pending = commandId ? this.pendingHostCommands.get(String(commandId)) : undefined; + const ok = sourceType === "command.accepted" ? sourcePayload.ok !== false : sourcePayload.ok === true; + const resultPayload = { + commandId: commandId || null, + method: sourcePayload.method || null, + ok, + result: sourcePayload.result, + error: sourcePayload.error, + sourceSeq, + }; + if (resultPayload.result && typeof resultPayload.result === "object") { + const result = resultPayload.result; + if (result.thread && typeof result.thread.id === "string") this.state.activeThreadId = result.thread.id; + if (typeof result.threadId === "string") this.state.activeThreadId = result.threadId; + if (typeof result.activeThreadId === "string") this.state.activeThreadId = result.activeThreadId; + if (typeof result.selectedThreadId === "string") this.state.activeThreadId = result.selectedThreadId; + if (result.turn && typeof result.turn.id === "string") this.state.activeTurnId = result.turn.id; + } + if (commandId) { + this.pendingHostCommands.delete(String(commandId)); + // A late terminal frame from an app that already reported + // connection.closed must not repopulate the cache we just invalidated. + if (pending || this.state.app !== "offline") { + this.cacheCommandResult( + String(commandId), + resultPayload, + pending?.commandScope || client.commandScope || this.hostCommandScope, + ); + } + } + const event = this.recordEvent("command.result", resultPayload, { sessionId }); + if (pending?.kind === "server-response") { + const requestId = pending.responseRequestId ?? pending.requestId; + if (resultPayload.ok) { + this.clearPendingServerRequest(pending.requestId); + const responseEvent = this.recordEvent("server.responded", { + requestId, + method: pending.method, + ok: true, + }, { sessionId }); + this.sendControl(pending.client, { type: "response.accepted", requestId, seq: responseEvent.seq }); + } else { + this.sendControl(pending.client, { + type: "response.rejected", + requestId, + code: "host_rejected", + message: responseErrorMessage(resultPayload.error || "VS Code host rejected the response"), + retryable: true, + }); + } + } else if (pending?.client) { + this.sendControl(pending.client, { type: "command.result", ...resultPayload, seq: event.seq }); + } + return; + } + + const payload = { + ...sourcePayload, + ...((sourceType === "approval.requested" + || sourceType === "input.requested" + || sourceType === "server.requested" + || sourceType === "server.request") && !sourcePayload.params + ? { params: sourcePayload } + : {}), + source: "vscode-host", + ...(sourceSeq !== undefined ? { sourceSeq } : {}), + ...(frame.raw !== undefined ? { raw: redact(frame.raw) } : {}), + }; + const event = this.recordEvent(sourceType, payload, { sessionId }); + // RelayHost sends event frames to its own relay transport and expects an + // ack. Acknowledge only after the frame has been accepted into our ring. + if (sourceSeq !== undefined) { + this.sendControl(client, { v: 1, kind: "ack", sessionId: sessionId || "", seq: sourceSeq }); + } + return event; + } + + handleHostAppUnavailable(client, sessionId, reason = "VS Code host app-server unavailable") { + // A delayed frame from an older host socket must never tear down the + // pending work or cache belonging to the currently authenticated host. + if (this.hostClient !== client) return; + const terminalError = { + code: "app_unavailable", + message: redactString(reason), + retryable: true, + }; + + // A local app-server crash invalidates both completed cache entries and + // in-flight host commands. Report in-flight commands as uncertain to the + // originating browser, but do not cache them: a retry after recovery must + // be explicit rather than silently replaying an unknown operation. + this.clearCommandResults(); + for (const [commandId, pending] of [...this.pendingHostCommands]) { + if (pending.hostClient && pending.hostClient !== client) continue; + if (!pending.hostClient && pending.commandScope && pending.commandScope !== client.commandScope) continue; + this.pendingHostCommands.delete(commandId); + if (pending.kind === "server-response") { + this.sendControl(pending.client, { + type: "response.rejected", + requestId: pending.responseRequestId ?? pending.requestId, + code: terminalError.code, + message: terminalError.message, + retryable: true, + }); + continue; + } + const payload = { + commandId, + method: pending.method || null, + ok: false, + uncertain: true, + retryable: true, + error: terminalError, + }; + const event = this.recordEvent("command.result", payload, { sessionId }); + if (pending.client) this.sendControl(pending.client, { type: "command.result", ...payload, seq: event.seq }); + } + + // Host approval/input requests are owned by the adapter, so the relay does + // not run a second expiry timer. Once the adapter reports its app process + // unavailable, remove every request tied to this host immediately. + for (const pending of [...this.pendingServerRequests.values()]) { + if (pending.source !== "host" || pending.hostClient !== client) continue; + this.clearPendingServerRequest(pending.appId); + const isInput = pending.method === "item/tool/requestUserInput" + || pending.method === "mcpServer/elicitation/request"; + this.recordEvent(isInput ? "input.expired" : "approval.expired", { + requestId: pending.appId, + method: pending.method, + reason: terminalError.message, + error: terminalError, + }, { sessionId }); + } + } + + dispatchCommand(client, message) { + const commandId = String(message.commandId || ""); + const method = String(message.method || ""); + if (!commandId || commandId.length > 128) { + return this.commandRejected(client, commandId, "invalid_command_id", "commandId is required"); + } + if (!ALLOWED_METHODS.has(method)) { + return this.commandRejected(client, commandId, "method_not_allowed", `Method ${method || "(empty)"} is not allowed`); + } + if (MUTATING_METHODS.has(method) && client.role !== "operator") { + this.auditAction(client, method, { commandId }, "denied"); + return this.commandRejected(client, commandId, "forbidden", "Operator token required"); + } + const cached = this.getCachedCommandResult(commandId); + if (cached) { + this.sendControl(client, { type: "command.result", ...cached, cached: true }); + return { accepted: true, cached: true }; + } + for (const pending of this.pendingHostCommands.values()) { + if (pending.commandId === commandId) { + this.sendControl(client, { type: "command.accepted", commandId, method, duplicate: true }); + return { accepted: true, duplicate: true }; + } + } + for (const pending of this.pendingAppRequests.values()) { + if (pending.commandId === commandId) { + this.sendControl(client, { type: "command.accepted", commandId, method, duplicate: true }); + return { accepted: true, duplicate: true }; + } + } + + if (method === "initialize") { + if (!this.state.initialized) { + return this.commandRejected(client, commandId, "app_initializing", "Codex is still initializing", true); + } + const result = { + commandId, + method, + ok: true, + result: this.mode === "host" + ? { protocol: 1, mode: "host", hostConnected: this.state.hostConnected, sessionId: this.state.hostSessionId } + : this.initializedResult, + cachedAt: Date.now(), + }; + this.cacheCommandResult(commandId, result, this.mode === "host" ? this.hostCommandScope : "embedded"); + this.sendControl(client, { type: "command.result", ...result, cached: true }); + return { accepted: true, cached: true }; + } + + if (!this.state.initialized) { + return this.commandRejected(client, commandId, "app_not_ready", "Codex app-server is not ready", true); + } + if (!message.params || typeof message.params !== "object" || Array.isArray(message.params)) { + return this.commandRejected(client, commandId, "invalid_params", "params must be an object"); + } + + const validationError = this.validateCommand(method, message.params); + if (validationError) return this.commandRejected(client, commandId, "invalid_params", validationError); + + // This MVP exposes one active Codex turn per relay session. Keeping the + // check at the relay boundary prevents two browser operators from racing + // a turn start or steering an outdated turn id. + const turnStartPending = [...this.pendingAppRequests.values()].some((pending) => pending.method === "turn/start") + || [...this.pendingHostCommands.values()].some((pending) => pending.method === "turn/start"); + if (method === "turn/start" && (this.state.activeTurnId || turnStartPending)) { + return this.commandRejected(client, commandId, "turn_active", "A Codex turn is already active", true); + } + if (method === "session/select" || method === "control/mode/set") { + const pendingMethod = method === "session/select" ? "session/select" : "control/mode/set"; + const sessionSwitchPending = [...this.pendingHostCommands.values()].some((pending) => pending.method === pendingMethod) + || [...this.pendingAppRequests.values()].some((pending) => pending.method === pendingMethod); + if (sessionSwitchPending) { + return this.commandRejected(client, commandId, method === "session/select" ? "session_switch_pending" : "mode_switch_pending", method === "session/select" ? "A session switch is already in progress" : "A control mode switch is already in progress", true); + } + if (this.state.activeTurnId || this.pendingServerRequests.size) { + return this.commandRejected(client, commandId, method === "session/select" ? "session_busy" : "mode_busy", "The active session has a running turn or pending request", true); + } + } + if (method === "turn/steer" && this.state.activeTurnId && message.params.expectedTurnId !== this.state.activeTurnId) { + return this.commandRejected(client, commandId, "stale_turn", "expectedTurnId does not match the active turn", true); + } + if (method === "turn/interrupt" && this.state.activeTurnId && message.params.turnId !== this.state.activeTurnId) { + return this.commandRejected(client, commandId, "stale_turn", "turnId does not match the active turn", true); + } + + // A connected VS Code bridge is the source of truth for the session. The + // relay never runs a second app-server request for the same command. + if (this.hostClient && this.hostClient.socket.readyState === WebSocket.OPEN) { + const hostFrame = { + v: 1, + kind: "command", + type: method, + commandId, + sessionId: this.hostClient.sessionId || undefined, + actor: { id: client.id || "web", role: client.role }, + payload: message.params, + }; + this.pendingHostCommands.set(commandId, { + commandId, + method, + client, + hostClient: this.hostClient, + commandScope: this.hostCommandScope || this.hostClient.commandScope || null, + createdAt: Date.now(), + }); + try { + this.hostClient.socket.send(JSON.stringify(hostFrame)); + this.auditAction(client, method, { commandId, params: message.params, target: "vscode-host" }, "forwarded"); + this.sendControl(client, { type: "command.accepted", commandId, method, target: "vscode-host" }); + return { accepted: true, commandId, method, target: "vscode-host" }; + } catch (error) { + this.pendingHostCommands.delete(commandId); + return this.commandRejected(client, commandId, "host_unavailable", error.message, true); + } + } + + const appId = this.nextAppRequestId("command"); + try { + this.pendingAppRequests.set(jsonRpcIdKey(appId), { + kind: "command", + commandId, + method, + client, + createdAt: Date.now(), + }); + this.sendToApp({ method, id: appId, params: message.params }); + this.auditAction(client, method, { commandId, params: message.params }, "forwarded"); + this.sendControl(client, { type: "command.accepted", commandId, method }); + return { accepted: true, commandId, method }; + } catch (error) { + this.pendingAppRequests.delete(jsonRpcIdKey(appId)); + return this.commandRejected(client, commandId, error.code || "app_offline", error.message, error.retryable); + } + } + + validateCommand(method, params) { + if (method === "control/mode/get") { + if (Object.keys(params).length > 0) return "control/mode/get does not accept parameters"; + } + if (method === "control/mode/set") { + if (params.mode !== "sync" && params.mode !== "async") return "mode must be sync or async"; + } + if (method === "session/list") { + if (params.limit !== undefined && (!Number.isInteger(params.limit) || params.limit < 1 || params.limit > 100)) { + return "limit must be an integer between 1 and 100"; + } + } + if (method === "session/select") { + const threadId = params.threadId ?? params.conversationId; + if (typeof threadId !== "string" || !threadId.trim()) return "threadId is required"; + if (threadId.length > 256) return "threadId is too long"; + } + if (method === "thread/start") { + const allowedSandboxes = new Set(["read-only", "workspace-write", "danger-full-access"]); + if (params.sandbox != null && !allowedSandboxes.has(params.sandbox)) { + return "sandbox must be read-only, workspace-write, or danger-full-access"; + } + if (params.cwd != null && typeof params.cwd !== "string") return "cwd must be a string"; + } + if (method === "turn/start") { + if (typeof params.threadId !== "string" || !params.threadId) return "threadId is required"; + if (!Array.isArray(params.input) || params.input.length === 0) return "input must be a non-empty array"; + } + if (method === "thread/settings/update") { + if (typeof params.threadId !== "string" || !params.threadId) return "threadId is required"; + const settings = params.threadSettings ?? params.settings; + if (!settings || typeof settings !== "object" || Array.isArray(settings)) return "threadSettings must be an object"; + if (settings.model !== undefined && typeof settings.model !== "string") return "threadSettings.model must be a string"; + // `null` is the official value for clearing a model's reasoning effort + // (some models do not expose a selectable effort). Preserve it through + // the relay instead of rejecting a valid next-turn update. + if (settings.effort !== undefined && settings.effort !== null && typeof settings.effort !== "string") return "threadSettings.effort must be a string or null"; + for (const key of ["sandboxPolicy", "approvalPolicy"]) { + const value = settings[key]; + if (value !== undefined && value !== null && typeof value !== "string" && (typeof value !== "object" || Array.isArray(value))) { + return `threadSettings.${key} must be a string, object, or null`; + } + } + if (settings.approvalsReviewer !== undefined && settings.approvalsReviewer !== null && typeof settings.approvalsReviewer !== "string") { + return "threadSettings.approvalsReviewer must be a string or null"; + } + if (settings.runtimeWorkspaceRoots !== undefined && settings.runtimeWorkspaceRoots !== null + && (!Array.isArray(settings.runtimeWorkspaceRoots) || !settings.runtimeWorkspaceRoots.every((entry) => typeof entry === "string"))) { + return "threadSettings.runtimeWorkspaceRoots must be an array of strings or null"; + } + if (settings.permissions !== undefined && settings.permissions !== null + && (typeof settings.permissions !== "string" && (typeof settings.permissions !== "object" || Array.isArray(settings.permissions)))) { + return "threadSettings.permissions must be a string, object, or null"; + } + } + if (method === "turn/steer") { + if (typeof params.threadId !== "string" || !params.threadId) return "threadId is required"; + if (typeof params.expectedTurnId !== "string" || !params.expectedTurnId) return "expectedTurnId is required"; + if (!Array.isArray(params.input) || params.input.length === 0) return "input must be a non-empty array"; + } + if (method === "turn/interrupt") { + if (typeof params.threadId !== "string" || !params.threadId) return "threadId is required"; + if (typeof params.turnId !== "string" || !params.turnId) return "turnId is required"; + } + return null; + } + + commandRejected(client, commandId, code, message, retryable = false) { + const payload = { type: "command.rejected", commandId: commandId || null, code, message, retryable: Boolean(retryable) }; + this.sendControl(client, payload); + return { accepted: false, ...payload }; + } + + scheduleServerRequestExpiry(requestId, pending) { + if (this.approvalTimeoutMs <= 0) return; + pending.timer = setTimeout(() => this.expireServerRequest(requestId), this.approvalTimeoutMs); + pending.timer.unref?.(); + } + + expireServerRequest(requestId) { + const pending = this.clearPendingServerRequest(requestId); + if (!pending) return; + const canonicalRequestId = pending.appId; + const isInput = pending.method === "item/tool/requestUserInput" || pending.method === "mcpServer/elicitation/request"; + const reason = "Remote approval timed out"; + this.recordEvent(isInput ? "input.expired" : "approval.expired", { + requestId: canonicalRequestId, + method: pending.method, + reason, + }, { sessionId: pending.hostClient?.sessionId || undefined }); + this.auditAction({ id: "relay", role: "system" }, pending.method, { requestId: canonicalRequestId }, "expired"); + + const result = defaultServerResponse(pending.method, reason); + if (pending.source === "host") { + if (!pending.hostClient || pending.hostClient.socket.readyState !== WebSocket.OPEN) return; + const commandMethod = isInput ? "server.request.respond" : "approval.respond"; + try { + pending.hostClient.socket.send(JSON.stringify({ + v: 1, + kind: "command", + type: commandMethod, + commandId: randomId("timeout"), + sessionId: pending.hostClient.sessionId || undefined, + actor: { id: "relay", role: "system" }, + payload: { + requestId: pending.appId, + decision: "deny", + response: result, + reason, + }, + })); + } catch { + // The local adapter also has its own expiry deny; no retry is needed. + } + return; + } + + try { + this.sendToApp({ id: pending.appId, result }); + } catch { + // The process may have exited while the approval was pending. + } + } + + clearPendingServerRequest(requestId) { + const key = jsonRpcIdKey(requestId); + const pending = this.pendingServerRequests.get(key); + if (pending?.timer) clearTimeout(pending.timer); + this.pendingServerRequests.delete(key); + return pending; + } + + dispatchServerResponse(client, message) { + const requestId = message.requestId ?? message.id ?? ""; + const requestKey = findTypedMapKey(this.pendingServerRequests, requestId, (value) => value?.appId, true); + if (client.role !== "operator") { + this.auditAction(client, "server-response", { requestId }, "denied"); + this.sendControl(client, { type: "response.rejected", requestId, code: "forbidden", message: "Operator token required" }); + return { accepted: false, code: "forbidden" }; + } + const pending = this.pendingServerRequests.get(requestKey); + if (!pending) { + this.sendControl(client, { type: "response.rejected", requestId, code: "unknown_request", message: "Request is no longer pending" }); + return { accepted: false, code: "unknown_request" }; + } + if (!("result" in message) && !("error" in message)) { + this.sendControl(client, { type: "response.rejected", requestId, code: "invalid_response", message: "result or error is required" }); + return { accepted: false, code: "invalid_response" }; + } + const remoteResponseAllowed = SERVER_REQUEST_METHODS.has(pending.method) + || pending.method === "item/tool/requestUserInput" + || pending.method === "mcpServer/elicitation/request"; + if (!remoteResponseAllowed) { + this.sendControl(client, { + type: "response.rejected", + requestId, + code: "unsupported_request", + message: "This app-server request must be handled by the host", + }); + return { accepted: false, code: "unsupported_request" }; + } + + const normalizedResult = Object.prototype.hasOwnProperty.call(message, "result") + ? normalizeServerResponseForApp(pending.method, message.result) + : undefined; + if (Object.prototype.hasOwnProperty.call(message, "result") + && !isValidApprovalResponse(pending.method, normalizedResult)) { + this.sendControl(client, { + type: "response.rejected", + requestId, + code: "invalid_response", + message: "Unsupported or malformed approval decision", + }); + return { accepted: false, code: "invalid_response" }; + } + if (Object.prototype.hasOwnProperty.call(message, "requestedDecision") + && (pending.method === "item/commandExecution/requestApproval" + || pending.method === "item/fileChange/requestApproval" + || pending.method === "applyPatchApproval" + || pending.method === "execCommandApproval")) { + const requested = approvalDecisionKind(message.requestedDecision); + const actual = approvalDecisionForResult(normalizedResult); + if (!requested || requested !== actual) { + this.sendControl(client, { + type: "response.rejected", + requestId, + code: "decision_mismatch", + message: "Outer approval decision does not match the response", + }); + return { accepted: false, code: "decision_mismatch" }; + } + } + + // Host-proxy mode uses the bridge's normalized command contract. Keep the + // original app-server request id in the payload, but do not forward an + // arbitrary JSON-RPC response as a relay command. + if (pending.source === "host") { + const result = normalizedResult; + const error = Object.prototype.hasOwnProperty.call(message, "error") ? message.error : undefined; + const method = pending.method || ""; + const isInput = method === "item/tool/requestUserInput" || method === "mcpServer/elicitation/request"; + const commandMethod = isInput ? "server.request.respond" : "approval.respond"; + const decision = error + ? "deny" + : isInput + ? "allow" + : approvalDecisionForResult(result); + // Generate the host command from the canonical app-server id. This + // preserves the legacy `response-77` shape when a browser merely + // stringified a numeric id, while still suffixing ids when both typed + // variants are pending concurrently. + const commandId = responseCommandId(pending.appId, this.pendingServerRequests, this.pendingHostCommands); + const hostFrame = { + v: 1, + kind: "command", + type: commandMethod, + commandId, + sessionId: pending.hostClient?.sessionId || undefined, + actor: { id: client.id || "web", role: client.role }, + payload: { + requestId: pending.appId, + decision, + ...(pending.commandHash ? { commandHash: pending.commandHash } : {}), + ...(result !== undefined ? { response: result } : {}), + ...(error !== undefined ? { reason: responseErrorMessage(error) } : {}), + }, + }; + const existing = this.pendingHostCommands.get(commandId); + if (existing?.kind === "server-response") { + this.sendControl(client, { type: "response.pending", requestId, commandId }); + return { accepted: true, pending: true, requestId, commandId }; + } + if (!pending.hostClient || pending.hostClient.socket.readyState !== WebSocket.OPEN) { + this.clearPendingServerRequest(pending.appId); + this.sendControl(client, { type: "response.rejected", requestId, code: "host_unavailable", message: "VS Code host is disconnected" }); + return { accepted: false, code: "host_unavailable" }; + } + try { + this.pendingHostCommands.set(commandId, { + kind: "server-response", + commandId, + // Keep the original app-server id for exact map cleanup. The + // browser-facing id may be a legacy stringified form of that id. + requestId: pending.appId, + responseRequestId: requestId, + method: pending.method, + client, + commandScope: this.hostCommandScope || pending.hostClient?.commandScope || null, + createdAt: Date.now(), + }); + pending.hostClient.socket.send(JSON.stringify(hostFrame)); + this.auditAction(client, pending.method, { requestId, target: "vscode-host", result, error }, "forwarded"); + this.sendControl(client, { type: "response.pending", requestId, commandId }); + return { accepted: true, pending: true, requestId, commandId }; + } catch (sendError) { + this.pendingHostCommands.delete(commandId); + this.sendControl(client, { type: "response.rejected", requestId, code: "host_unavailable", message: sendError.message, retryable: true }); + return { accepted: false, code: "host_unavailable" }; + } + } + + const appMessage = { id: pending.appId }; + if ("result" in message) appMessage.result = normalizedResult; + else appMessage.error = message.error; + try { + this.sendToApp(appMessage); + this.clearPendingServerRequest(pending.appId); + this.auditAction(client, pending.method, { requestId, result: message.result, error: message.error }, "responded"); + const event = this.recordEvent("server.responded", { + requestId, + method: pending.method, + ok: !message.error, + }); + this.sendControl(client, { type: "response.accepted", requestId, seq: event.seq }); + return { accepted: true, requestId }; + } catch (error) { + this.sendControl(client, { type: "response.rejected", requestId, ...normalizeError(error) }); + return { accepted: false, ...normalizeError(error) }; + } + } + + cacheCommandResult(commandId, payload, sessionScope) { + const key = String(commandId); + const scope = this.mode === "host" + ? (sessionScope || this.hostCommandScope || null) + : "embedded"; + this.commandResults.set(key, { ...payload, cachedAt: Date.now() }); + this.commandResultScopes.set(key, scope); + this.pruneCommandResults(); + } + + getCachedCommandResult(commandId) { + const key = String(commandId); + const cached = this.commandResults.get(key); + if (!cached) return null; + // A disconnect leaves the outcome unknown. Never replay that marker as a + // completed result; remove it so a retry can be forwarded to a reconnected + // host (or receive the normal offline error). + if (cached.uncertain) { + this.commandResults.delete(key); + this.commandResultScopes.delete(key); + return null; + } + const activeScope = this.mode === "host" + ? (this.hostClient && this.state.hostConnected ? this.hostCommandScope : null) + : "embedded"; + // Host results are never replayed while disconnected. This also prevents + // an old result from leaking across a session-id change. + if (!activeScope || this.commandResultScopes.get(key) !== activeScope) return null; + return cached; + } + + clearCommandResults() { + this.commandResults.clear(); + this.commandResultScopes.clear(); + } + + pruneCommandResults() { + const cutoff = Date.now() - 15 * 60 * 1000; + for (const [key, value] of this.commandResults) { + if (value.cachedAt < cutoff) { + this.commandResults.delete(key); + this.commandResultScopes.delete(key); + } + } + while (this.commandResults.size > 1_000) { + const [firstKey] = this.commandResults.keys(); + this.commandResults.delete(firstKey); + this.commandResultScopes.delete(firstKey); + } + } + + sendControl(client, message) { + if (client.capture) client.capture.push(message); + if (!client.socket || client.socket.readyState !== WebSocket.OPEN) return; + const serialized = JSON.stringify(message); + const frameBytes = Buffer.byteLength(serialized, "utf8"); + if (frameBytes > MAX_WS_PAYLOAD) { + client.socket.close(1009, "relay frame is too large"); + return; + } + const bufferedBytes = Number(client.socket.bufferedAmount) || 0; + if (bufferedBytes + frameBytes > this.clientBufferedByteLimit) { + client.socket.close(1013, "client is too slow"); + return; + } + client.socket.send(serialized); + } + + snapshot() { + // These projections used to be serialized both inside `state` and again + // at the top level, nearly doubling every long-history control frame. + // Keep the compact lifecycle state nested and one authoritative transcript + // projection at the stable top-level protocol fields. + const { + outputTail, + messages, + subagents, + sessionMetadata, + executionStatus, + ...state + } = this.state; + return { + protocol: 1, + latestSeq: this.nextSeq, + state, + clients: [...this.clients] + .filter((client) => client.authenticated) + .map((client) => ({ id: client.id, role: client.role })), + pendingRequests: [...this.pendingServerRequests.values()].map((request) => ({ + requestId: request.appId, + method: request.method, + params: request.params, + ...(request.commandHash ? { commandHash: request.commandHash } : {}), + ...(request.risk ? { risk: request.risk } : {}), + ...(request.summary ? { summary: request.summary } : {}), + ...(request.expiresAt ? { expiresAt: request.expiresAt } : {}), + createdAt: request.createdAt, + })), + outputTail: outputTail || "", + messages: Array.isArray(messages) ? messages : [], + subagents: Array.isArray(subagents) ? subagents : [], + ...(sessionMetadata ? { metadata: sessionMetadata } : {}), + status: executionStatus, + executionStatus, + }; + } + + tokenFromRequest(request) { + const authorization = request.headers.authorization || ""; + if (/^Bearer\s+/i.test(authorization)) return authorization.replace(/^Bearer\s+/i, ""); + return request.headers["x-codex-token"] || ""; + } + + authenticateHttp(request) { + const role = this.roleForToken(this.tokenFromRequest(request)); + if (role) return role; + if (!this.authRequired && isLoopbackAddress(request.socket?.remoteAddress)) return "operator"; + return null; + } + + async handleHttp(request, response) { + const base = `http://${request.headers.host || "localhost"}`; + let requestUrl; + try { + requestUrl = new URL(request.url, base); + } catch { + jsonResponse(response, 400, { error: "invalid_url" }); + return; + } + + if (!this.authRequired && !isLoopbackRequestHost(request)) { + jsonResponse(response, 403, { error: "loopback_host_required" }); + return; + } + + if (request.method === "GET" && requestUrl.pathname === "/api/health") { + jsonResponse(response, this.state.app === "offline" ? 503 : 200, { + ok: this.state.app !== "offline", + app: this.state.app, + initialized: this.state.initialized, + authRequired: this.authRequired, + latestSeq: this.nextSeq, + }); + return; + } + + if (requestUrl.pathname.startsWith("/api/")) { + const role = this.authenticateHttp(request); + if (!role) { + jsonResponse(response, 401, { error: "unauthorized" }, { "WWW-Authenticate": "Bearer" }); + return; + } + + if (request.method === "POST" && !isAllowedHttpOrigin(request)) { + jsonResponse(response, 403, { error: "origin_not_allowed" }); + return; + } + + if (request.method === "GET" && requestUrl.pathname === "/api/state") { + jsonResponse(response, 200, { role, ...this.snapshot(), audit: this.audit.slice(-50) }); + return; + } + if (request.method === "GET" && requestUrl.pathname === "/api/events") { + const fromSeq = Number(requestUrl.searchParams.get("fromSeq") || 0); + jsonResponse(response, 200, { + latestSeq: this.nextSeq, + events: this.events.filter((event) => event.seq > fromSeq), + }); + return; + } + if (request.method === "POST" && requestUrl.pathname === "/api/command") { + try { + const body = await readJson(request); + const client = { id: "http", role, authenticated: true, capture: [] }; + const result = this.dispatchCommand(client, { type: "command", ...body }); + jsonResponse(response, result.accepted ? 202 : 400, { ...result, messages: client.capture }); + } catch (error) { + jsonResponse(response, error.statusCode || 400, { error: normalizeError(error) }); + } + return; + } + if (request.method === "POST" && requestUrl.pathname === "/api/respond") { + try { + const body = await readJson(request); + const client = { id: "http", role, authenticated: true, capture: [] }; + const result = this.dispatchServerResponse(client, { type: "respond", ...body }); + jsonResponse(response, result.accepted ? 202 : 400, { ...result, messages: client.capture }); + } catch (error) { + jsonResponse(response, error.statusCode || 400, { error: normalizeError(error) }); + } + return; + } + jsonResponse(response, 404, { error: "not_found" }); + return; + } + + this.serveStatic(request, response, requestUrl.pathname); + } + + serveStatic(request, response, pathname) { + if (request.method !== "GET" && request.method !== "HEAD") { + response.writeHead(405, { Allow: "GET, HEAD" }); + response.end(); + return; + } + let relativePath; + try { + relativePath = pathname === "/" ? "index.html" : decodeURIComponent(pathname).replace(/^\/+/, ""); + } catch { + // A malformed percent escape must be an ordinary client error, not an + // uncaught exception from the HTTP request handler. + response.writeHead(400, { "Content-Type": "text/plain; charset=utf-8", "Cache-Control": "no-store" }); + response.end("Invalid URL"); + return; + } + const filePath = path.resolve(PUBLIC_ROOT, relativePath); + if (!filePath.startsWith(`${PUBLIC_ROOT}${path.sep}`) && filePath !== path.join(PUBLIC_ROOT, "index.html")) { + response.writeHead(403); + response.end("Forbidden"); + return; + } + fs.stat(filePath, (error, stat) => { + if (error || !stat.isFile()) { + response.writeHead(404, { "Content-Type": "text/plain; charset=utf-8" }); + response.end("Not found"); + return; + } + response.writeHead(200, { + "Content-Type": contentType(filePath), + "Content-Length": stat.size, + "Cache-Control": "no-store", + "X-Content-Type-Options": "nosniff", + "Referrer-Policy": "no-referrer", + "Content-Security-Policy": "default-src 'self'; connect-src 'self' ws: wss:; img-src 'self' data:; style-src 'self'; script-src 'self'; base-uri 'none'; frame-ancestors 'self'", + }); + if (request.method === "HEAD") response.end(); + else fs.createReadStream(filePath).pipe(response); + }); + } +} + +function isObjectPayload(value) { + return Boolean(value && typeof value === "object" && !Array.isArray(value)); +} + +function firstObject(...values) { + return values.find((value) => isObjectPayload(value)) || null; +} + +function commandMethodFromFrame(frame) { + const nested = isObjectPayload(frame.command) ? frame.command : null; + if (typeof frame.method === "string" && frame.method.trim()) return frame.method.trim(); + if (typeof nested?.type === "string" && nested.type.trim()) return nested.type.trim(); + if (typeof frame.kind === "string" && (frame.kind === "command" || frame.kind === "response") && typeof frame.type === "string") { + if (frame.type !== "command" && frame.type !== "respond" && frame.type !== "server-response") return frame.type.trim(); + } + if (typeof frame.type === "string" && frame.type !== "command" && frame.type !== "respond" && frame.type !== "server-response") { + return frame.type.trim(); + } + return ""; +} + +function normalizeWireMethod(method) { + const value = String(method || "").trim(); + const aliases = { + "control.mode.get": "control/mode/get", + "controlmode.get": "control/mode/get", + "mode.get": "control/mode/get", + "control.mode.set": "control/mode/set", + "controlmode.set": "control/mode/set", + "mode.set": "control/mode/set", + "thread.start": "thread/start", + "session.new": "session/new", + "sessionnew": "session/new", + "thread.new": "session/new", + "threadnew": "session/new", + "session/new": "session/new", + "thread/new": "session/new", + "thread.settings.update": "thread/settings/update", + "threadsettings.update": "thread/settings/update", + "session.list": "session/list", + "thread.list": "session/list", + "session.select": "session/select", + "session.switch": "session/select", + "thread.select": "session/select", + "thread.attach": "session/select", + "turn.start": "turn/start", + "turn.steer": "turn/steer", + "turn.interrupt": "turn/interrupt", + }; + return aliases[value.toLowerCase()] || value; +} + +function commandIdFromFrame(frame) { + const nested = isObjectPayload(frame.command) ? frame.command : null; + const value = frame.commandId ?? nested?.commandId ?? (typeof frame.id === "string" ? frame.id : undefined); + return value === undefined || value === null ? "" : String(value); +} + +function normalizeBrowserCommand(frame) { + const method = normalizeWireMethod(commandMethodFromFrame(frame)); + const type = frame.type; + const hasCommandEnvelope = frame.kind === "command" + || type === "command" + || (frame.kind === undefined && (ALLOWED_METHODS.has(method) || REMOTE_RESPONSE_METHODS.has(method))); + if (!hasCommandEnvelope || !method) return null; + + const nested = isObjectPayload(frame.command) ? frame.command : null; + const params = firstObject(frame.params, frame.payload, nested?.params, nested?.payload) || {}; + return { + type: "command", + commandId: commandIdFromFrame(frame), + method, + params, + }; +} + +function approvalWireDecision(decision) { + // Keep the caller's decision intact until dispatch knows the target + // app-server method. Legacy approval methods use `approved`/`abort`, while + // v2 methods use `accept`/`cancel`; method-aware normalization handles the + // conversion without discarding amendment tags. + return decision; +} + +function normalizeBrowserResponse(frame, hintedMethod) { + const type = typeof frame.type === "string" ? frame.type : ""; + const method = hintedMethod || commandMethodFromFrame(frame); + const isLegacy = type === "respond" || type === "server-response"; + const isResponse = isLegacy + || frame.kind === "response" + || REMOTE_RESPONSE_METHODS.has(method); + if (!isResponse) return null; + + const nested = isObjectPayload(frame.command) ? frame.command : null; + const payload = firstObject(frame.payload, frame.params, nested?.payload, nested?.params) || (isLegacy ? {} : frame); + const requestId = frame.requestId ?? payload.requestId ?? (typeof frame.id === "number" || typeof frame.id === "string" ? frame.id : ""); + const response = { type: "respond", requestId }; + if (payload.decision !== undefined) response.requestedDecision = payload.decision; + + if (Object.prototype.hasOwnProperty.call(frame, "result")) { + response.result = frame.result; + } else if (Object.prototype.hasOwnProperty.call(frame, "error")) { + response.error = frame.error; + } else if (Object.prototype.hasOwnProperty.call(payload, "result")) { + response.result = payload.result; + } else if (Object.prototype.hasOwnProperty.call(payload, "error")) { + response.error = payload.error; + } else if (payload.response !== undefined) { + response.result = payload.response; + } else if (method === "input.respond" || method === "server.request.respond") { + if (payload.answers !== undefined) response.result = { answers: payload.answers }; + else { + const custom = {}; + for (const [key, value] of Object.entries(payload)) { + if (!["v", "kind", "type", "method", "commandId", "id", "sessionId", "actor", "requestId", "reason", "params", "payload", "command"].includes(key)) custom[key] = value; + } + if (Object.keys(custom).length) response.result = custom; + } + } else if (payload.decision !== undefined) { + response.result = { decision: approvalWireDecision(payload.decision) }; + } + return response; +} + +async function main() { + const relay = new CodexRelay(); + const address = await relay.start(); + const displayHost = address.host === "::" || address.host === "0.0.0.0" ? "127.0.0.1" : address.host; + process.stdout.write(`Codex Remote Collab: http://${displayHost}:${address.port}\n`); + if (relay.authRequired) { + process.stdout.write(`Host token: ${relay.hostToken}\n`); + process.stdout.write(`Operator token: ${relay.operatorToken}\n`); + process.stdout.write(`Viewer token: ${relay.viewerToken}\n`); + process.stdout.write("Keep these tokens private. Use TLS before exposing this relay outside a trusted network.\n"); + } else { + process.stdout.write("Authentication: disabled for loopback connections (set CODEX_REMOTE_AUTH=required to enable tokens).\n"); + } + + const shutdown = async () => { + await relay.stop(); + process.exit(0); + }; + process.once("SIGINT", shutdown); + process.once("SIGTERM", shutdown); +} + +if (require.main === module) { + main().catch((error) => { + process.stderr.write(`${error.stack || error}\n`); + process.exitCode = 1; + }); +} + +module.exports = { + ALLOWED_METHODS, + CodexRelay, + SERVER_REQUEST_METHODS, + redact, +}; diff --git a/aether-vscodex/test/adapter-safety.test.js b/aether-vscodex/test/adapter-safety.test.js new file mode 100644 index 000000000..99e1ecece --- /dev/null +++ b/aether-vscodex/test/adapter-safety.test.js @@ -0,0 +1,745 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const test = require("node:test"); + +const { CodexAgentAdapter } = require("../vscode-extension/dist/codexAgentAdapter.js"); +const { RelayHost } = require("../vscode-extension/dist/relayHost.js"); + +class FakeRpc { + responses = []; + requests = []; + notificationListener; + requestListener; + exitListener; + overrides; + + constructor(overrides = {}) { + this.overrides = overrides; + } + + get running() { + return true; + } + + async start() {} + + async request(method, params) { + this.requests.push({ method, params }); + if (Object.prototype.hasOwnProperty.call(this.overrides, method)) { + const override = this.overrides[method]; + return typeof override === "function" ? override(params) : override; + } + if (method === "initialize") return { userAgent: "test", codexHome: "/tmp/codex" }; + if (method === "thread/start") return { thread: { id: "thread-test" }, cwd: "/tmp" }; + if (method === "turn/start") return { turn: { id: "turn-test" } }; + if (method === "turn/steer") return { turn: { id: "turn-test" } }; + if (method === "turn/interrupt") return {}; + throw new Error(`unexpected request ${method}`); + } + + notify() {} + + respond(id, result) { + this.responses.push({ id, result }); + } + + respondError(id, code, message) { + this.responses.push({ id, error: { code, message } }); + } + + onNotification(listener) { + this.notificationListener = listener; + return { dispose: () => undefined }; + } + + onServerRequest(listener) { + this.requestListener = listener; + return { dispose: () => undefined }; + } + + onExit(listener) { + this.exitListener = listener; + return { dispose: () => undefined }; + } + + close() {} + + emitRequest(request) { + this.requestListener(request); + } + + emitNotification(notification) { + this.notificationListener(notification); + } +} + +class FakeRelay { + frames = []; + listeners = new Set(); + + async connect() {} + + send(frame) { + this.frames.push(frame); + } + + onMessage(listener) { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + close() {} +} + +test("CodexAgentAdapter keeps numeric and string approval ids distinct", async () => { + const rpc = new FakeRpc(); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + + rpc.emitRequest({ + id: 1, + method: "item/commandExecution/requestApproval", + params: { threadId: "t", turnId: "u", itemId: "n", command: "echo number" }, + }); + rpc.emitRequest({ + id: "1", + method: "item/commandExecution/requestApproval", + params: { threadId: "t", turnId: "u", itemId: "s", command: "echo string" }, + }); + + const snapshot = await adapter.snapshot(); + assert.deepEqual(snapshot.pendingApprovals.map((entry) => entry.requestId), [1, "1"]); + await adapter.respondApproval(1, "deny"); + await adapter.respondApproval("1", "deny"); + assert.deepEqual(rpc.responses.map((entry) => entry.id), [1, "1"]); + assert.equal((await adapter.snapshot()).pendingApprovals.length, 0); + await adapter.dispose(); +}); + +test("commandActions are included in high-risk approval classification", async () => { + const rpc = new FakeRpc(); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + rpc.emitRequest({ + id: 2, + method: "item/commandExecution/requestApproval", + params: { + threadId: "t", + turnId: "u", + itemId: "actions", + command: null, + commandActions: [{ type: "unknown", command: "sudo rm -rf /" }], + }, + }); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.pendingApprovals[0].risk, "high"); + await adapter.respondApproval(2, "deny"); + await adapter.dispose(); +}); + +test("output snapshots stay redacted and interrupt clears the active turn", async () => { + const rpc = new FakeRpc(); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + await adapter.startThread({}); + await adapter.startTurn({ text: "hello" }); + assert.equal((await adapter.snapshot()).turnId, "turn-test"); + + rpc.emitNotification({ + method: "item/agentMessage/delta", + params: { delta: "credential Bearer abcdefghijklmnop" }, + }); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.outputTail.includes("Bearer abcdefghijklmnop"), false); + assert.match(snapshot.outputTail, /\[REDACTED\]/); + + await adapter.interruptTurn({}); + const afterInterrupt = await adapter.snapshot(); + assert.equal(afterInterrupt.turnId, null); + assert.equal(afterInterrupt.state, "idle"); + await adapter.dispose(); +}); + +test("async adapter lists app-server threads and exposes the model catalog", async () => { + const rpc = new FakeRpc({ + "model/list": { + data: [{ id: "model-1", model: "gpt-5.6-sol", displayName: "5.6 Sol", hidden: false }], + nextCursor: null, + }, + "thread/list": { + data: [ + { + id: "thread-recent", + name: null, + preview: "Inspect the workspace\nwith detail", + cwd: "/tmp/workspace", + createdAt: 1_700_000_000, + updatedAt: 1_700_000_100, + status: { type: "idle" }, + source: "vscode", + }, + ], + nextCursor: "next-page", + backwardsCursor: null, + }, + }); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + + const result = await adapter.listSessions({ limit: 500, query: "workspace", sortKey: "invalid" }); + assert.equal(result.sessions[0].threadId, "thread-recent"); + assert.equal(result.sessions[0].title, "Inspect the workspace with detail"); + assert.equal(result.sessions[0].updatedAtMs, 1_700_000_100_000); + assert.equal(result.nextCursor, "next-page"); + const listRequest = rpc.requests.find((entry) => entry.method === "thread/list"); + assert.deepEqual(listRequest.params, { + limit: 100, + sortKey: "updated_at", + sortDirection: "desc", + searchTerm: "workspace", + }); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.mode, "async"); + assert.equal(snapshot.metadata.availableModels[0].model, "gpt-5.6-sol"); + await adapter.dispose(); +}); + +test("async adapter projects live token usage notifications into metadata and snapshots", async () => { + const rpc = new FakeRpc({ + "model/list": { data: [], nextCursor: null }, + }); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + const events = []; + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + await adapter.startThread({}); + + rpc.emitNotification({ + method: "thread/tokenUsage/updated", + params: { + threadId: "thread-test", + // A usage update may arrive after the turn has completed. It must not + // make the adapter report that historical turn as active again. + turnId: "turn-finished", + tokenUsage: { + total: { + totalTokens: 1_200, + inputTokens: 800, + cachedInputTokens: 100, + cacheWriteInputTokens: 20, + outputTokens: 300, + reasoningOutputTokens: 80, + }, + last: { + totalTokens: 450, + inputTokens: 300, + cachedInputTokens: 40, + cacheWriteInputTokens: 10, + outputTokens: 100, + reasoningOutputTokens: 40, + }, + modelContextWindow: 128_000, + }, + }, + }); + + const expected = { + total: { + totalTokens: 1_200, + inputTokens: 800, + cachedInputTokens: 100, + cacheWriteInputTokens: 20, + outputTokens: 300, + reasoningOutputTokens: 80, + }, + last: { + totalTokens: 450, + inputTokens: 300, + cachedInputTokens: 40, + cacheWriteInputTokens: 10, + outputTokens: 100, + reasoningOutputTokens: 40, + }, + modelContextWindow: 128_000, + }; + const snapshot = await adapter.snapshot(); + assert.deepEqual(snapshot.metadata.tokenUsage, expected); + assert.deepEqual(snapshot.metadata.latestTokenUsageInfo, expected); + assert.equal(snapshot.turnId, null); + + const usageEvent = events.find((event) => event.raw?.method === "thread/tokenUsage/updated"); + assert.ok(usageEvent); + assert.deepEqual(usageEvent.payload.tokenUsage, expected); + assert.deepEqual(usageEvent.payload.latestTokenUsageInfo, expected); + // Keep the raw diagnostic envelope redacted while exposing only the safe + // numeric projection to the browser. + assert.equal(usageEvent.raw.params.tokenUsage, "[REDACTED]"); + + rpc.emitNotification({ + method: "thread/tokenUsage/updated", + params: { + threadId: "thread-test", + turnId: "turn-finished", + tokenUsage: { total: { inputTokens: -1 } }, + }, + }); + assert.deepEqual((await adapter.snapshot()).metadata.tokenUsage, expected); + await adapter.dispose(); +}); + +test("async adapter resumes a thread with structured history and ignores late notifications", async () => { + const thread = { + id: "thread-selected", + name: "Selected thread", + preview: "hello", + cwd: "/tmp/selected", + createdAt: 1_700_000_000, + updatedAt: 1_700_000_010, + status: { type: "idle" }, + turns: [{ + id: "turn-history", + status: "completed", + startedAt: 1_700_000_001, + completedAt: 1_700_000_004, + durationMs: 3_000, + items: [ + { type: "userMessage", id: "user-1", clientId: null, content: [{ type: "text", text: "hello", text_elements: [] }] }, + { type: "reasoning", id: "reason-1", summary: ["Checking files"], content: [] }, + { type: "commandExecution", id: "command-1", command: "pwd", cwd: "/tmp/selected", status: "completed", aggregatedOutput: "/tmp/selected\n", exitCode: 0, durationMs: 50, commandActions: [] }, + { type: "agentMessage", id: "agent-1", text: "Done", phase: "final_answer" }, + ], + }], + }; + const rpc = new FakeRpc({ + "model/list": { data: [{ id: "model-1", model: "gpt-5.6-sol" }], nextCursor: null }, + "thread/resume": { + thread, + model: "gpt-5.6-sol", + modelProvider: "openai", + serviceTier: null, + cwd: "/tmp/selected", + approvalPolicy: "on-request", + approvalsReviewer: "user", + sandbox: { type: "workspaceWrite" }, + reasoningEffort: "high", + }, + }); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + const events = []; + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + + const result = await adapter.selectSession({ threadId: "thread-selected" }); + assert.equal(result.threadId, "thread-selected"); + let snapshot = await adapter.snapshot(); + assert.equal(snapshot.messages.length, 4); + assert.deepEqual(snapshot.messages.map((message) => message.kind), ["user", "reasoning", "tool", "assistant"]); + assert.equal(snapshot.messages[2].output, "/tmp/selected\n"); + assert.equal(snapshot.metadata.title, "Selected thread"); + assert.equal(snapshot.metadata.threadSettings.effort, "high"); + assert.equal(snapshot.metadata.historyComplete, true); + assert.equal(snapshot.status.turnStatus, "completed"); + assert.match(snapshot.outputTail, /Done/); + assert.ok(events.some((event) => event.type === "output.snapshot" && event.payload.historyComplete === true)); + const authoritative = events.find((event) => event.type === "session.snapshot"); + assert.equal(authoritative.payload.threadId, "thread-selected"); + assert.equal(authoritative.payload.metadata.model, "gpt-5.6-sol"); + assert.equal(authoritative.payload.messages.length, 4); + + rpc.emitNotification({ + method: "item/completed", + params: { + threadId: "thread-old", + turnId: "turn-old", + completedAtMs: Date.now(), + item: { type: "agentMessage", id: "late-old", text: "wrong thread" }, + }, + }); + rpc.emitNotification({ + method: "item/completed", + params: { + threadId: "thread-selected", + turnId: "turn-live", + completedAtMs: Date.now(), + item: { type: "agentMessage", id: "current-item", text: "current thread" }, + }, + }); + snapshot = await adapter.snapshot(); + assert.equal(snapshot.messages.some((message) => message.itemId === "late-old"), false); + assert.equal(snapshot.messages.some((message) => message.itemId === "current-item"), true); + await adapter.dispose(); +}); + +test("async adapter hydrates paginated turns and items into chronological complete history", async () => { + const threadId = "thread-paged-history"; + const userItem = (id, text) => ({ + type: "userMessage", + id, + clientId: null, + content: [{ type: "text", text, text_elements: [] }], + }); + const assistantItem = (id, text) => ({ + type: "agentMessage", + id, + text, + phase: "final_answer", + }); + const earlyUser = userItem("early-user", "first question"); + const rpc = new FakeRpc({ + "model/list": { data: [], nextCursor: null }, + "thread/resume": { + thread: { + id: threadId, + name: "Paged history", + preview: "first question", + cwd: "/tmp/paged", + createdAt: 50, + updatedAt: 350, + historyMode: "paginated", + status: { type: "idle" }, + turns: [], + }, + model: "gpt-5.6-sol", + cwd: "/tmp/paged", + initialTurnsPage: { + data: [{ + id: "turn-late", + status: "completed", + startedAt: 300, + completedAt: 310, + itemsView: "full", + items: [userItem("late-user", "third question"), assistantItem("late-agent", "third answer")], + }], + nextCursor: "turn-page-2", + backwardsCursor: null, + }, + }, + "thread/turns/list": (params) => { + if (params.cursor === "turn-page-2") { + return { + data: [{ + id: "turn-early", + status: "completed", + startedAt: 100, + completedAt: 110, + itemsView: "summary", + items: [earlyUser], + }], + nextCursor: "turn-page-3", + backwardsCursor: null, + }; + } + assert.equal(params.cursor, "turn-page-3"); + return { + data: [{ + id: "turn-middle", + status: "completed", + startedAt: 200, + completedAt: 210, + itemsView: "full", + items: [userItem("middle-user", "second question"), assistantItem("middle-agent", "second answer")], + }], + nextCursor: null, + backwardsCursor: null, + }; + }, + "thread/items/list": (params) => { + assert.equal(params.turnId, "turn-early"); + return { + data: [ + // The summary row is repeated by the full item page; hydration must + // de-duplicate it while adding the omitted assistant response. + { turnId: "turn-early", item: earlyUser }, + { turnId: "turn-early", item: assistantItem("early-agent", "first answer") }, + ], + nextCursor: null, + backwardsCursor: null, + }; + }, + }); + const events = []; + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + await adapter.selectSession({ threadId }); + + const resume = rpc.requests.find((entry) => entry.method === "thread/resume"); + assert.deepEqual(resume.params, { + threadId, + excludeTurns: true, + initialTurnsPage: { limit: 100, sortDirection: "asc", itemsView: "full" }, + }); + const turnPages = rpc.requests.filter((entry) => entry.method === "thread/turns/list"); + assert.deepEqual(turnPages.map((entry) => entry.params.cursor), ["turn-page-2", "turn-page-3"]); + assert.ok(turnPages.every((entry) => entry.params.threadId === threadId + && entry.params.limit === 100 + && entry.params.sortDirection === "asc" + && entry.params.itemsView === "full")); + const itemPages = rpc.requests.filter((entry) => entry.method === "thread/items/list"); + assert.deepEqual(itemPages.map((entry) => entry.params), [{ + threadId, + turnId: "turn-early", + limit: 100, + sortDirection: "asc", + }]); + assert.equal(rpc.requests.some((entry) => entry.method === "thread/read"), false); + + const snapshot = await adapter.snapshot(); + assert.deepEqual(snapshot.messages.map((message) => [message.turnId, message.text]), [ + ["turn-early", "first question"], + ["turn-early", "first answer"], + ["turn-middle", "second question"], + ["turn-middle", "second answer"], + ["turn-late", "third question"], + ["turn-late", "third answer"], + ]); + assert.equal(snapshot.metadata.historyComplete, true); + const outputSnapshot = events.find((event) => event.type === "output.snapshot"); + assert.equal(outputSnapshot.payload.historyComplete, true); + assert.deepEqual(outputSnapshot.payload.messages.map((message) => message.text), [ + "first question", + "first answer", + "second question", + "second answer", + "third question", + "third answer", + ]); + await adapter.dispose(); +}); + +test("async adapter falls back to thread/read when resume omits existing history", async () => { + const metadataThread = { + id: "thread-paginated", + preview: "existing conversation", + cwd: "/tmp/project", + createdAt: 1_700_000_000, + updatedAt: 1_700_000_100, + status: { type: "idle" }, + turns: [], + }; + const rpc = new FakeRpc({ + "model/list": { data: [], nextCursor: null }, + "thread/resume": { thread: metadataThread, model: "gpt-5.6-sol", cwd: "/tmp/project" }, + "thread/read": { + thread: { + ...metadataThread, + turns: [{ + id: "turn-read", + status: "completed", + items: [{ type: "agentMessage", id: "read-agent", text: "hydrated history" }], + }], + }, + }, + }); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + await adapter.selectSession({ threadId: "thread-paginated" }); + const read = rpc.requests.find((entry) => entry.method === "thread/read"); + assert.deepEqual(read.params, { threadId: "thread-paginated", includeTurns: true }); + assert.equal((await adapter.snapshot()).messages[0].text, "hydrated history"); + await adapter.dispose(); +}); + +test("async adapter starts new sessions and sends flat durable thread settings", async () => { + const rpc = new FakeRpc({ + "model/list": { data: [], nextCursor: null }, + "thread/start": { + thread: { id: "thread-new", preview: "", cwd: "/tmp/new", status: { type: "idle" }, turns: [] }, + model: "gpt-5.6-sol", + cwd: "/tmp/new", + reasoningEffort: "medium", + }, + "thread/settings/update": { ok: true }, + }); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0, defaultCwd: "/tmp/default" }, rpc); + await adapter.start(); + await adapter.newSession({}); + await adapter.updateThreadSettings({ + threadSettings: { + model: "gpt-5.6-terra", + effort: "high", + approvalPolicy: "on-request", + approvalsReviewer: "user", + sandboxPolicy: "workspace-write", + permissions: ":workspace", + }, + }); + const start = rpc.requests.find((entry) => entry.method === "thread/start"); + assert.equal(start.params.cwd, "/tmp/default"); + const update = rpc.requests.find((entry) => entry.method === "thread/settings/update"); + assert.deepEqual(update.params, { + threadId: "thread-new", + model: "gpt-5.6-terra", + effort: "high", + approvalPolicy: "on-request", + approvalsReviewer: "user", + permissions: ":workspace", + }); + assert.equal(Object.prototype.hasOwnProperty.call(update.params, "sandboxPolicy"), false); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.model, "gpt-5.6-terra"); + assert.equal(snapshot.metadata.latestReasoningEffort, "high"); + assert.equal(snapshot.metadata.sandboxPolicy, "workspace-write"); + await adapter.dispose(); +}); + +test("thread settings updates require the send_task_input capability", async () => { + const rpc = new FakeRpc(); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + const relay = new FakeRelay(); + const host = new RelayHost({ + adapter, + relay, + capabilities: ["read_output"], + sessionId: "test-session", + }); + + await host.handleFrame({ + kind: "command", + type: "thread.settings.update", + commandId: "settings-without-capability", + actor: { role: "operator" }, + payload: { threadSettings: { model: "gpt-5.6-sol", effort: "high" } }, + }); + + const result = relay.frames.find((frame) => frame.payload?.commandId === "settings-without-capability"); + assert.equal(result.type, "command.rejected"); + assert.match(result.payload.error, /missing capability: send_task_input/); + await adapter.dispose(); +}); + +test("RelayHost exposes session list as read-only and protects session selection", async () => { + const relay = new FakeRelay(); + const calls = []; + const adapter = { + async start() {}, + async sendInput() { return {}; }, + async cancel() { return {}; }, + async respondApproval() { return {}; }, + async snapshot() { return { threadId: "thread-a", turnId: null, state: "idle", pendingApprovals: [], outputTail: "" }; }, + onEvent() { return { dispose() {} }; }, + async dispose() {}, + async listSessions(params) { + calls.push({ method: "listSessions", params }); + return { sessions: [{ threadId: "thread-a", title: "A", updatedAtMs: null, active: true, available: true }], activeThreadId: "thread-a" }; + }, + async selectSession(params) { + calls.push({ method: "selectSession", params }); + return { threadId: params.threadId, previousThreadId: "thread-a", switched: true, available: true }; + }, + async newSession(params) { + calls.push({ method: "newSession", params }); + return { opened: true, command: "chatgpt.newCodexPanel" }; + }, + getControlMode() { return "sync"; }, + async setControlMode(params) { + calls.push({ method: "setControlMode", params }); + return { changed: true, controlMode: params.mode, previousControlMode: "sync", modeEpoch: 1 }; + }, + }; + const host = new RelayHost({ + adapter, + relay, + capabilities: ["read_output", "send_task_input"], + sessionId: "test-session", + }); + + await host.handleFrame({ kind: "command", type: "session/list", commandId: "list-1", actor: { role: "viewer" }, payload: {} }); + const listed = relay.frames.find((frame) => frame.payload?.commandId === "list-1"); + assert.equal(listed.type, "command.accepted"); + assert.equal(listed.payload.result.activeThreadId, "thread-a"); + assert.equal(calls[0].method, "listSessions"); + + await host.handleFrame({ kind: "command", type: "session/select", commandId: "select-viewer", actor: { role: "viewer" }, payload: { threadId: "thread-b" } }); + const denied = relay.frames.find((frame) => frame.payload?.commandId === "select-viewer"); + assert.equal(denied.type, "command.rejected"); + + await host.handleFrame({ kind: "command", type: "session/select", commandId: "select-operator", actor: { role: "operator" }, payload: { threadId: "thread-b" } }); + const selected = relay.frames.find((frame) => frame.payload?.commandId === "select-operator"); + assert.equal(selected.type, "command.accepted"); + assert.equal(selected.payload.result.threadId, "thread-b"); + assert.equal(calls.at(-1).method, "selectSession"); + + await host.handleFrame({ kind: "command", type: "session/new", commandId: "new-viewer", actor: { role: "viewer" }, payload: {} }); + const deniedNew = relay.frames.find((frame) => frame.payload?.commandId === "new-viewer"); + assert.equal(deniedNew.type, "command.rejected"); + + await host.handleFrame({ kind: "command", type: "session/new", commandId: "new-operator", actor: { role: "operator" }, payload: {} }); + const opened = relay.frames.find((frame) => frame.payload?.commandId === "new-operator"); + assert.equal(opened.type, "command.accepted"); + assert.equal(opened.payload.result.command, "chatgpt.newCodexPanel"); + assert.equal(calls.at(-1).method, "newSession"); + + await host.handleFrame({ kind: "command", type: "control/mode/get", commandId: "mode-get-viewer", actor: { role: "viewer" }, payload: {} }); + const mode = relay.frames.find((frame) => frame.payload?.commandId === "mode-get-viewer"); + assert.equal(mode.type, "command.accepted"); + assert.equal(mode.payload.result.mode, "sync"); + + await host.handleFrame({ kind: "command", type: "control/mode/set", commandId: "mode-set-viewer", actor: { role: "viewer" }, payload: { mode: "async" } }); + const deniedMode = relay.frames.find((frame) => frame.payload?.commandId === "mode-set-viewer"); + assert.equal(deniedMode.type, "command.rejected"); + + await host.handleFrame({ kind: "command", type: "control/mode/set", commandId: "mode-set-operator", actor: { role: "operator" }, payload: { mode: "async" } }); + const changedMode = relay.frames.find((frame) => frame.payload?.commandId === "mode-set-operator"); + assert.equal(changedMode.type, "command.accepted"); + assert.equal(changedMode.payload.result.controlMode, "async"); + assert.equal(calls.at(-1).method, "setControlMode"); +}); + +test("approval decision conflicts and unknown tagged objects fail closed", async () => { + const rpc = new FakeRpc(); + const adapter = new CodexAgentAdapter({ approvalTimeoutMs: 0 }, rpc); + await adapter.start(); + const relay = new FakeRelay(); + const host = new RelayHost({ + adapter, + relay, + capabilities: ["read_output", "send_task_input", "cancel_task", "approve_low_risk"], + sessionId: "test-session", + }); + + rpc.emitRequest({ + id: 3, + method: "execCommandApproval", + params: { conversationId: "thread-test", callId: "call-3", command: ["echo", "safe"] }, + }); + await host.handleFrame({ + kind: "command", + type: "approval.respond", + commandId: "conflicting-response", + actor: { role: "operator" }, + payload: { + requestId: 3, + decision: "deny", + response: { decision: "approved_mcp_policy_amendment" }, + }, + }); + assert.deepEqual(rpc.responses[0], { + id: 3, + result: { decision: { denied: { rejection: "approval response implies allow, but decision is deny" } } }, + }); + + rpc.emitRequest({ + id: 4, + method: "item/commandExecution/requestApproval", + params: { threadId: "thread-test", turnId: "turn-test", itemId: "item-4", command: "echo safe" }, + }); + await host.handleFrame({ + kind: "command", + type: "approval.respond", + commandId: "unknown-tagged-response", + actor: { role: "operator" }, + payload: { + requestId: 4, + decision: "allow", + response: { decision: { futurePolicyGrant: { scope: "all" } } }, + }, + }); + assert.deepEqual(rpc.responses[1], { + id: 4, + result: { decision: "decline" }, + }); + await adapter.dispose(); +}); diff --git a/aether-vscodex/test/cloud-server.test.js b/aether-vscodex/test/cloud-server.test.js new file mode 100644 index 000000000..7940c3bea --- /dev/null +++ b/aether-vscodex/test/cloud-server.test.js @@ -0,0 +1,298 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const fs = require("node:fs"); +const os = require("node:os"); +const path = require("node:path"); +const test = require("node:test"); +const { WebSocket } = require("ws"); + +const { AetherVscodexCloudServer, RoomManager } = require("../cloud/server.js"); + +const internalToken = "test-internal-token-with-enough-entropy"; + +function internalFetch(base, pathname, options = {}) { + return fetch(`${base}${pathname}`, { + ...options, + headers: { + Authorization: `Bearer ${internalToken}`, + ...(options.body ? { "Content-Type": "application/json" } : {}), + ...(options.headers || {}), + }, + }); +} + +function websocketClient(base, clientType, token, sessionId) { + const socket = new WebSocket(`${base.replace(/^http/, "ws")}/api/vscodex/ws`); + const messages = []; + const waiters = []; + const wait = (predicate, timeout = 5_000, label = "websocket frame") => new Promise((resolve, reject) => { + const existing = messages.find(predicate); + if (existing) return resolve(existing); + const timer = setTimeout(() => { + const index = waiters.findIndex((entry) => entry.resolve === resolve); + if (index >= 0) waiters.splice(index, 1); + reject(new Error(`timed out waiting for ${label}; received: ${JSON.stringify(messages.map((message) => ({ type: message.type, kind: message.kind, commandId: message.commandId })))}`)); + }, timeout); + waiters.push({ + predicate, + resolve: (message) => { + clearTimeout(timer); + resolve(message); + }, + }); + }); + socket.on("message", (data) => { + const message = JSON.parse(data.toString("utf8")); + messages.push(message); + for (let index = waiters.length - 1; index >= 0; index -= 1) { + if (!waiters[index].predicate(message)) continue; + const waiter = waiters.splice(index, 1)[0]; + waiter.resolve(message); + } + }); + return new Promise((resolve, reject) => { + socket.once("open", () => { + socket.send(JSON.stringify({ v: 1, kind: "hello", clientType, protocol: 1, ...(sessionId ? { sessionId } : {}) })); + socket.send(JSON.stringify(clientType === "host" + ? { v: 1, kind: "auth", accessToken: token } + : { type: "auth", token })); + wait((message) => message.type === "auth.ok").then(() => resolve({ socket, wait, messages }), reject); + }); + socket.once("error", reject); + }); +} + +async function pairDevice(base, userId, name) { + const pairingResponse = await internalFetch(base, `/internal/v1/users/${encodeURIComponent(userId)}/pairings`, { + method: "POST", + body: JSON.stringify({ name }), + }); + assert.equal(pairingResponse.status, 201); + const pairing = await pairingResponse.json(); + const exchangeResponse = await fetch(`${base}/v1/pairings/exchange`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ code: pairing.code, name }), + }); + assert.equal(exchangeResponse.status, 201); + return exchangeResponse.json(); +} + +async function browserTicket(base, userId, deviceId) { + const response = await internalFetch(base, `/internal/v1/users/${encodeURIComponent(userId)}/ws-tickets`, { + method: "POST", + body: JSON.stringify({ device_id: deviceId }), + }); + assert.equal(response.status, 201); + return response.json(); +} + +function exchangeAttempt(base, headers = {}) { + return fetch(`${base}/v1/pairings/exchange`, { + method: "POST", + headers: { "Content-Type": "application/json", ...headers }, + body: JSON.stringify({ code: "INVALID-CODE" }), + }); +} + +test("pairing exchange trusts a gateway client IP only with valid internal authentication", async (t) => { + const dataDir = fs.mkdtempSync(path.join(os.tmpdir(), "aether-vscodex-rate-limit-")); + const server = new AetherVscodexCloudServer({ + host: "127.0.0.1", + port: 0, + internalToken, + publicWsUrl: "wss://aether.example/api/vscodex/ws", + dataDir, + }); + await server.start(); + t.after(async () => { + await server.stop(); + fs.rmSync(dataDir, { recursive: true, force: true }); + }); + const address = server.address(); + const base = `http://127.0.0.1:${address.port}`; + const trustedHeaders = (clientIp) => ({ + Authorization: `Bearer ${internalToken}`, + "X-Aether-Client-IP": clientIp, + }); + + for (let attempt = 0; attempt < 10; attempt += 1) { + assert.equal((await exchangeAttempt(base, trustedHeaders("198.51.100.10"))).status, 400); + } + assert.equal((await exchangeAttempt(base, trustedHeaders("198.51.100.10"))).status, 429); + assert.equal((await exchangeAttempt(base, trustedHeaders("198.51.100.11"))).status, 400); + assert.equal((await exchangeAttempt(base, trustedHeaders("2001:db8::10"))).status, 400); + + server.exchangeAttempts.clear(); + for (let attempt = 0; attempt < 5; attempt += 1) { + assert.equal((await exchangeAttempt(base, { "X-Aether-Client-IP": `198.51.100.${20 + attempt}` })).status, 400); + } + for (let attempt = 0; attempt < 5; attempt += 1) { + assert.equal((await exchangeAttempt(base, { + Authorization: "Bearer invalid-internal-token", + "X-Aether-Client-IP": `198.51.100.${30 + attempt}`, + })).status, 400); + } + assert.equal((await exchangeAttempt(base, { "X-Aether-Client-IP": "198.51.100.99" })).status, 429); + + server.exchangeAttempts.clear(); + const invalidForwardedAddresses = ["proxy.internal", "198.51.100.40, 198.51.100.41"]; + for (let attempt = 0; attempt < 10; attempt += 1) { + assert.equal((await exchangeAttempt(base, trustedHeaders(invalidForwardedAddresses[attempt % 2]))).status, 400); + } + assert.equal((await exchangeAttempt(base, trustedHeaders("198.51.100.42, 198.51.100.43"))).status, 429); +}); + +test("cloud sidecar pairs a device and isolates host/browser traffic by Aether user and device", async (t) => { + const dataDir = fs.mkdtempSync(path.join(os.tmpdir(), "aether-vscodex-test-")); + const server = new AetherVscodexCloudServer({ + host: "127.0.0.1", + port: 0, + internalToken, + publicWsUrl: "wss://aether.example/api/vscodex/ws", + dataDir, + pairingTtlMs: 5_000, + ticketTtlMs: 5_000, + }); + await server.start(); + t.after(async () => { + await server.stop(); + fs.rmSync(dataDir, { recursive: true, force: true }); + }); + const address = server.address(); + const base = `http://127.0.0.1:${address.port}`; + + const unauthorized = await fetch(`${base}/internal/v1/users/user-a/devices`); + assert.equal(unauthorized.status, 401); + + const paired = await pairDevice(base, "user-a", "MacBook VS Code"); + assert.match(paired.device_token, /^avx1\./); + const devicesResponse = await internalFetch(base, "/internal/v1/users/user-a/devices"); + assert.equal(devicesResponse.status, 200); + const devices = await devicesResponse.json(); + assert.deepEqual(devices.devices.map((device) => ({ id: device.id, name: device.name, connected: device.connected })), [ + { id: paired.device_id, name: "MacBook VS Code", connected: false }, + ]); + + const ticket = await browserTicket(base, "user-a", paired.device_id); + assert.equal(ticket.ws_url, "/api/vscodex/ws"); + const host = await websocketClient(base, "host", paired.device_token, "host-user-a"); + const browser = await websocketClient(base, "web", ticket.ticket); + t.after(() => host.socket.close()); + t.after(() => browser.socket.close()); + browser.socket.send(JSON.stringify({ type: "subscribe", fromSeq: 0 })); + + host.socket.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + id: "connection-a", + sessionId: "host-user-a", + seq: 1, + ts: new Date().toISOString(), + payload: {}, + })); + host.socket.send(JSON.stringify({ + v: 1, + kind: "event", + type: "session.snapshot", + id: "snapshot-a", + sessionId: "host-user-a", + seq: 2, + ts: new Date().toISOString(), + payload: { threadId: "thread-a", state: "idle", messages: [{ kind: "assistant", text: "user-a-only" }] }, + })); + const snapshot = await browser.wait((message) => message.kind === "event" && message.type === "session.snapshot", 5_000, "session snapshot"); + assert.equal(snapshot.payload.threadId, "thread-a"); + assert.equal(snapshot.payload.messages[0].text, "user-a-only"); + + browser.socket.send(JSON.stringify({ type: "command", commandId: "cmd-a", method: "session/list", params: {} })); + const command = await host.wait((message) => message.kind === "command" && message.commandId === "cmd-a", 5_000, "browser command"); + assert.equal(command.type, "session/list"); + + const secondUser = await pairDevice(base, "user-b", "Other VS Code"); + const secondTicket = await browserTicket(base, "user-b", secondUser.device_id); + const secondBrowser = await websocketClient(base, "web", secondTicket.ticket); + t.after(() => secondBrowser.socket.close()); + secondBrowser.socket.send(JSON.stringify({ type: "subscribe", fromSeq: 0 })); + await new Promise((resolve) => setTimeout(resolve, 50)); + assert.equal(secondBrowser.messages.some((message) => message.payload?.threadId === "thread-a"), false); + + const reusedTicket = new WebSocket(`${base.replace(/^http/, "ws")}/api/vscodex/ws`); + const closed = new Promise((resolve, reject) => { + reusedTicket.once("open", () => { + reusedTicket.send(JSON.stringify({ v: 1, kind: "hello", clientType: "web", protocol: 1 })); + reusedTicket.send(JSON.stringify({ type: "auth", token: ticket.ticket })); + }); + reusedTicket.once("close", (code) => resolve(code)); + reusedTicket.once("error", reject); + }); + assert.equal(await closed, 1008, "browser tickets are one-time credentials"); +}); + +test("device revocation closes its room and blocks future host authentication", async (t) => { + const dataDir = fs.mkdtempSync(path.join(os.tmpdir(), "aether-vscodex-revoke-")); + const server = new AetherVscodexCloudServer({ + host: "127.0.0.1", + port: 0, + internalToken, + publicWsUrl: "wss://aether.example/api/vscodex/ws", + dataDir, + }); + await server.start(); + t.after(async () => { + await server.stop(); + fs.rmSync(dataDir, { recursive: true, force: true }); + }); + const address = server.address(); + const base = `http://127.0.0.1:${address.port}`; + const paired = await pairDevice(base, "user-a", "Revoked device"); + const host = await websocketClient(base, "host", paired.device_token, "revoked-host"); + + const response = await internalFetch(base, `/internal/v1/users/user-a/devices/${paired.device_id}`, { method: "DELETE" }); + assert.equal(response.status, 204); + await new Promise((resolve) => host.socket.once("close", resolve)); + + const rejected = new WebSocket(`${base.replace(/^http/, "ws")}/api/vscodex/ws`); + const closed = new Promise((resolve, reject) => { + rejected.once("open", () => { + rejected.send(JSON.stringify({ v: 1, kind: "hello", clientType: "host", protocol: 1, sessionId: "retry" })); + rejected.send(JSON.stringify({ v: 1, kind: "auth", accessToken: paired.device_token })); + }); + rejected.once("close", (code) => resolve(code)); + rejected.once("error", reject); + }); + assert.equal(await closed, 1008); +}); + +test("room revocation wins a concurrent room creation", async () => { + const rooms = new RoomManager(); + let releaseCreation; + const creationGate = new Promise((resolve) => { releaseCreation = resolve; }); + let stopped = false; + const room = { + key: rooms.key("user-a", "device-a"), + userId: "user-a", + deviceId: "device-a", + relay: { stop: async () => { stopped = true; } }, + connections: 0, + lastActiveMs: Date.now(), + }; + rooms.createRoom = async (key) => { + await creationGate; + rooms.rooms.set(key, room); + return room; + }; + + const pendingGet = rooms.get("user-a", "device-a"); + await new Promise((resolve) => setImmediate(resolve)); + const pendingRevoke = rooms.revoke("user-a", "device-a"); + releaseCreation(); + + await assert.rejects(pendingGet, /device revoked/); + await pendingRevoke; + assert.equal(stopped, true); + assert.equal(rooms.rooms.has(room.key), false); + await assert.rejects(rooms.get("user-a", "device-a"), /device revoked/); +}); diff --git a/aether-vscodex/test/codex-ipc-adapter.test.js b/aether-vscodex/test/codex-ipc-adapter.test.js new file mode 100644 index 000000000..6f19ffc66 --- /dev/null +++ b/aether-vscodex/test/codex-ipc-adapter.test.js @@ -0,0 +1,2247 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const fs = require("node:fs/promises"); +const os = require("node:os"); +const path = require("node:path"); +const test = require("node:test"); + +const { CodexIpcAgentAdapter } = require("../vscode-extension/dist/codexIpcAgentAdapter.js"); + +const THREAD_ID = "11111111-1111-4111-8111-111111111111"; +const SECOND_THREAD_ID = "22222222-2222-4222-8222-222222222222"; +const STALE_THREAD_ID = "33333333-3333-4333-8333-333333333333"; +const THIRD_THREAD_ID = "44444444-4444-4444-8444-444444444444"; +const ULID_THREAD_ID = "01a0399e790373d63b6fcde4e1f97d02"; + +async function waitFor(predicate, timeoutMs = 1_000) { + const deadline = Date.now() + timeoutMs; + while (Date.now() < deadline) { + if (await predicate()) return; + await new Promise((resolve) => setTimeout(resolve, 5)); + } + assert.fail(`condition was not met within ${timeoutMs}ms`); +} + +class FakeIpcClient { + constructor(state, emitSnapshot = true) { + this.socketPath = "/tmp/fake-codex-ipc.sock"; + this.state = state; + this.emitSnapshot = emitSnapshot; + this.broadcastListeners = new Set(); + this.streamListeners = new Set(); + this.errorListeners = new Set(); + this.closeListeners = new Set(); + this.calls = []; + this.streamStates = new Map(); + } + + subscribe(set, listener) { + set.add(listener); + return { dispose: () => set.delete(listener) }; + } + + onBroadcast(listener) { return this.subscribe(this.broadcastListeners, listener); } + onStreamEvent(listener) { return this.subscribe(this.streamListeners, listener); } + onError(listener) { return this.subscribe(this.errorListeners, listener); } + onClose(listener) { return this.subscribe(this.closeListeners, listener); } + getClientId() { return "follower"; } + getConversationState(threadId) { + const state = this.streamStates.get(threadId); + if (!state) return undefined; + return { + ...state, + conversationState: JSON.parse(JSON.stringify(state.conversationState)), + }; + } + async connect() { this.calls.push({ method: "connect" }); return "follower"; } + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + return threadId === THREAD_ID ? "owner" : null; + } + async followConversation(threadId, following, options) { + this.calls.push({ method: "followConversation", threadId, following, options }); + if (!following) { + this.streamStates.delete(threadId); + return; + } + if (this.emitSnapshot) queueMicrotask(() => this.emitState(this.state, 1, "snapshot", threadId)); + } + emitState(state, revision, kind = "snapshot", threadId = THREAD_ID, ownerClientId = "owner") { + const event = { + kind, + conversationId: threadId, + hostId: "local", + ownerClientId, + revision, + conversationState: state, + raw: { type: "broadcast", method: "thread-stream-state-changed", version: 11 }, + }; + this.streamStates.set(threadId, { + conversationId: threadId, + hostId: "local", + ownerClientId, + revision, + conversationState: JSON.parse(JSON.stringify(state)), + }); + for (const listener of this.streamListeners) listener(event); + } + emitFollowing(threadId, following, options = {}) { + const frame = { + type: "broadcast", + method: "thread-stream-following-changed", + version: options.version ?? 1, + sourceClientId: options.sourceClientId ?? "owner", + ...(options.targetClientIds ? { targetClientIds: options.targetClientIds } : {}), + params: { + conversationId: threadId, + hostId: options.hostId ?? "local", + following, + }, + }; + for (const listener of this.broadcastListeners) listener(frame); + } + emitClientStatus(clientId, status, options = {}) { + const frame = { + type: "broadcast", + method: "client-status-changed", + version: options.version ?? 0, + sourceClientId: options.sourceClientId ?? clientId, + params: { + clientId, + clientType: options.clientType ?? "vscode-webview", + status, + }, + }; + for (const listener of this.broadcastListeners) listener(frame); + } + emitClose(error = new Error("fixture IPC socket closed")) { + for (const listener of this.closeListeners) listener(error); + } + async startTurn(threadId, input, options) { + this.calls.push({ method: "startTurn", threadId, input, options }); + return { turnId: "turn-new" }; + } + async steerTurn(threadId, input, options) { + this.calls.push({ method: "steerTurn", threadId, input, options }); + return { turnId: "turn-new" }; + } + async updateThreadSettings(threadId, settings, options) { + this.calls.push({ method: "updateThreadSettings", threadId, settings, options }); + return { updated: true }; + } + async interruptTurn(threadId, options) { + this.calls.push({ method: "interruptTurn", threadId, options }); + return { interrupted: true }; + } + async respondCommandApproval(threadId, requestId, decision, options) { + this.calls.push({ method: "respondCommandApproval", threadId, requestId, decision, options }); + return { accepted: true }; + } + async respondFileApproval(threadId, requestId, decision, options) { + this.calls.push({ method: "respondFileApproval", threadId, requestId, decision, options }); + return { accepted: true }; + } + async respondPermissionsApproval(threadId, requestId, response, options) { + this.calls.push({ method: "respondPermissionsApproval", threadId, requestId, response, options }); + return { accepted: true }; + } + async respondUserInput(threadId, requestId, response, options) { + this.calls.push({ method: "respondUserInput", threadId, requestId, response, options }); + return { accepted: true }; + } + async respondMcpElicitation(threadId, requestId, response, options) { + this.calls.push({ method: "respondMcpElicitation", threadId, requestId, response, options }); + return { accepted: true }; + } + async loadCompleteHistory() { this.calls.push({ method: "loadCompleteHistory" }); } + async dispose() { this.calls.push({ method: "dispose" }); } +} + +class DiscoveryIpcClient extends FakeIpcClient { + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + return "owner"; + } +} + +class MultiSessionIpcClient extends FakeIpcClient { + constructor(states, owners = {}) { + super(states.get(THREAD_ID)); + this.states = states; + this.owners = owners; + } + + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + return this.owners[threadId] || null; + } + + async followConversation(threadId, following, options) { + this.calls.push({ method: "followConversation", threadId, following, options }); + if (!following) { + this.streamStates.delete(threadId); + return; + } + if (this.states.has(threadId)) { + const owner = this.owners[threadId] || "owner"; + queueMicrotask(() => this.emitState(this.states.get(threadId), 1, "snapshot", threadId, owner)); + } + } +} + +/** Simulates a target owned by another Codex process that never answers follow. */ +class NoSnapshotSwitchIpcClient extends MultiSessionIpcClient { + constructor(states, owners = {}) { + super(states, owners); + this.targetFollowFailed = false; + } + + async followConversation(threadId, following, options) { + this.calls.push({ method: "followConversation", threadId, following, options }); + if (!following) { + this.streamStates.delete(threadId); + return; + } + if (threadId === SECOND_THREAD_ID) { + this.targetFollowFailed = true; + return; + } + // The old owner is deliberately silent after the failed target attach; + // the adapter must restore from its cached stream state instead. + if (this.targetFollowFailed && threadId === THREAD_ID) return; + if (this.states.has(threadId)) { + const owner = this.owners[threadId] || "owner"; + queueMicrotask(() => this.emitState(this.states.get(threadId), 1, "snapshot", threadId, owner)); + } + } +} + +class ThrowingTargetFollowIpcClient extends MultiSessionIpcClient { + async followConversation(threadId, following, options) { + if (threadId === SECOND_THREAD_ID && following) { + this.calls.push({ method: "followConversation", threadId, following, options }); + throw new Error("target follow failed immediately"); + } + return super.followConversation(threadId, following, options); + } +} + +class OwnerChangingIpcClient extends MultiSessionIpcClient { + constructor(states, owners = {}) { + super(states, owners); + this.targetDiscoveryCount = 0; + } + + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + if (threadId === SECOND_THREAD_ID) { + this.targetDiscoveryCount += 1; + return this.targetDiscoveryCount === 1 ? "owner-b" : "owner-c"; + } + return this.owners[threadId] || null; + } +} + +/** Delays the first target snapshot so list probing overlaps a session selection request. */ +class DelayedProbeIpcClient extends MultiSessionIpcClient { + constructor(states, owners = {}) { + super(states, owners); + this.targetFollowCount = 0; + } + + async followConversation(threadId, following, options) { + this.calls.push({ method: "followConversation", threadId, following, options }); + if (!following) { + this.streamStates.delete(threadId); + return; + } + if (!this.states.has(threadId)) return; + const owner = this.owners[threadId] || "owner"; + if (threadId === SECOND_THREAD_ID) { + this.targetFollowCount += 1; + const delay = this.targetFollowCount === 1 ? 30 : 0; + setTimeout(() => this.emitState(this.states.get(threadId), this.targetFollowCount, "snapshot", threadId, owner), delay); + return; + } + queueMicrotask(() => this.emitState(this.states.get(threadId), 1, "snapshot", threadId, owner)); + } +} + +/** Delays the first waiting attach so a newer official route can supersede it. */ +class WaitingRouteRaceIpcClient extends MultiSessionIpcClient { + async followConversation(threadId, following, options) { + this.calls.push({ method: "followConversation", threadId, following, options }); + if (!following) { + this.streamStates.delete(threadId); + return; + } + if (!this.states.has(threadId)) return; + const owner = this.owners[threadId] || "owner"; + const delay = threadId === THREAD_ID ? 30 : 0; + setTimeout(() => this.emitState(this.states.get(threadId), 1, "snapshot", threadId, owner), delay); + } +} + +/** Holds fallback discovery so an official route can arrive while it is stale. */ +class WaitingPollRouteRaceIpcClient extends WaitingRouteRaceIpcClient { + constructor(states, owners = {}) { + super(states, owners); + this.delayFallbackDiscovery = false; + this.fallbackDiscoveryStarted = false; + this.fallbackOwner = new Promise((resolve) => { + this.resolveFallbackOwner = resolve; + }); + } + + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + if (threadId === SECOND_THREAD_ID && this.delayFallbackDiscovery) { + this.fallbackDiscoveryStarted = true; + return this.fallbackOwner; + } + return this.owners[threadId] || null; + } + + releaseFallbackDiscovery() { + this.delayFallbackDiscovery = false; + const resolve = this.resolveFallbackOwner; + this.resolveFallbackOwner = null; + if (resolve) resolve(this.owners[SECOND_THREAD_ID] || null); + } +} + +/** Holds the post-snapshot owner confirmation so assertions run mid-switch. */ +class DelayedTargetOwnerConfirmationIpcClient extends MultiSessionIpcClient { + constructor(states, owners = {}) { + super(states, owners); + this.targetDiscoveryCount = 0; + this.targetConfirmationStarted = false; + this.targetOwnerConfirmation = new Promise((resolve) => { + this.resolveTargetOwnerConfirmation = resolve; + }); + } + + async findThreadOwner(threadId) { + this.calls.push({ method: "findThreadOwner", threadId }); + const owner = this.owners[threadId] || null; + if (threadId !== SECOND_THREAD_ID) return owner; + this.targetDiscoveryCount += 1; + if (this.targetDiscoveryCount === 1) return owner; + this.targetConfirmationStarted = true; + return this.targetOwnerConfirmation; + } + + releaseTargetOwnerConfirmation() { + if (!this.resolveTargetOwnerConfirmation) return; + const resolve = this.resolveTargetOwnerConfirmation; + this.resolveTargetOwnerConfirmation = null; + resolve(this.owners[SECOND_THREAD_ID] || null); + } +} + +function fixtureState(requests = []) { + return { + id: THREAD_ID, + title: "fixture session", + cwd: "/tmp/workspace", + turns: [{ + id: "turn-old", + status: "completed", + items: [ + { type: "userMessage", id: "user-1", content: [{ type: "text", text: "hello" }] }, + { type: "agentMessage", id: "agent-1", text: "hi from VS Code" }, + ], + }], + requests, + threadRuntimeStatus: { type: "idle" }, + }; +} + +test("IPC adapter follows an existing session and routes input/approval without spawning", async () => { + const client = new FakeIpcClient(fixtureState([ + { + id: 7, + method: "item/commandExecution/requestApproval", + params: { threadId: THREAD_ID, turnId: "turn-old", command: "echo safe" }, + }, + { + id: "question-1", + method: "item/tool/requestUserInput", + params: { threadId: THREAD_ID, turnId: "turn-old", questions: [{ id: "choice", question: "Pick one" }] }, + }, + ])); + const events = []; + let newSessionCalls = 0; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + openNewSession: async () => { + newSessionCalls += 1; + return { opened: true, command: "chatgpt.newCodexPanel" }; + }, + }); + adapter.onEvent((event) => events.push(event)); + + await adapter.start(); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.threadId, THREAD_ID); + assert.equal(snapshot.metadata.adapter, "codex-ipc-follower"); + assert.match(snapshot.outputTail, /hi from VS Code/); + assert.deepEqual(snapshot.pendingApprovals.map((item) => item.requestId), [7]); + assert.equal(snapshot.pendingRequests.length, 2); + assert.equal(client.calls.some((call) => call.method === "startProcess"), false); + await assert.rejects(() => adapter.startThread({}), /does not create a new thread/); + assert.deepEqual(await adapter.newSession(), { opened: true, command: "chatgpt.newCodexPanel" }); + assert.equal(newSessionCalls, 1); + + await adapter.startTurn({ text: "remote input" }); + const startCall = client.calls.find((call) => call.method === "startTurn"); + assert.deepEqual(startCall.input, "remote input"); + assert.equal(startCall.options.ownerClientId, "owner"); + + await adapter.respondApproval(7, "allow"); + const approvalCall = client.calls.find((call) => call.method === "respondCommandApproval"); + assert.equal(approvalCall.decision, "accept"); + + await adapter.respondApproval("question-1", "allow", undefined, { answers: { choice: ["yes"] } }); + const inputCall = client.calls.find((call) => call.method === "respondUserInput"); + assert.deepEqual(inputCall.response, { answers: { choice: { answers: ["yes"] } } }); + const outputSnapshot = events.find((event) => event.type === "output.snapshot"); + assert.ok(outputSnapshot); + assert.deepEqual(outputSnapshot.payload.messages.map((item) => [item.role, item.kind, item.text]), [ + ["user", "user", "hello"], + ["assistant", "assistant", "hi from VS Code"], + ]); + await adapter.dispose(); +}); + +test("IPC adapter clears stale projections on close while preserving closed session identity", async () => { + const client = new FakeIpcClient(fixtureState()); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + + const before = await adapter.snapshot(); + assert.equal(before.threadId, THREAD_ID); + assert.equal(before.metadata.ownerClientId, "owner"); + assert.equal(before.metadata.attachReady, true); + assert.ok(before.messages.length > 0); + assert.match(before.outputTail, /hi from VS Code/); + events.length = 0; + + client.emitClose(new Error("fixture IPC connection lost")); + + const closedEvent = events.find((event) => event.type === "connection.closed"); + assert.ok(closedEvent); + assert.equal(closedEvent.threadId, THREAD_ID); + assert.equal(closedEvent.payload.ownerClientId, "owner"); + assert.equal(closedEvent.payload.message, "fixture IPC connection lost"); + + const after = await adapter.snapshot(); + assert.equal(after.state, "disconnected"); + assert.equal(after.threadId, null); + assert.equal(after.turnId, null); + assert.equal(after.outputTail, ""); + assert.deepEqual(after.messages, []); + assert.deepEqual(after.subagents, []); + assert.equal(after.metadata.attachReady, false); + assert.equal(Object.hasOwn(after.metadata, "ownerClientId"), false); + assert.equal(Object.hasOwn(after.metadata, "revision"), false); + await adapter.dispose(); +}); + +test("IPC adapter persists model and effort through the official follower settings envelope", async () => { + const client = new FakeIpcClient(fixtureState()); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + + await adapter.start(); + const result = await adapter.updateThreadSettings({ + threadSettings: { model: " gpt-5.6-sol ", effort: " ultra " }, + // UI-only fields must not leak into the owner request. + commandId: "ignored", + }); + assert.deepEqual(result, { updated: true }); + const call = client.calls.find((entry) => entry.method === "updateThreadSettings"); + assert.equal(call.threadId, THREAD_ID); + assert.deepEqual(call.settings, { model: "gpt-5.6-sol", effort: "ultra" }); + assert.equal(call.options.ownerClientId, "owner"); + await assert.rejects(() => adapter.updateThreadSettings({ model: "" }), /non-empty string/); + await adapter.dispose(); +}); + +test("IPC adapter projects official latest model and reasoning effort fields", async () => { + const state = fixtureState(); + state.latestModel = "gpt-5.6-sol"; + state.latestReasoningEffort = "ultra"; + state.latestThreadSettings = { + model: "gpt-5.6-sol", + modelProvider: "aether", + effort: "ultra", + multiAgentMode: "explicitRequestOnly", + }; + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const metadata = (await adapter.snapshot()).metadata; + assert.equal(metadata.model, "gpt-5.6-sol"); + assert.equal(metadata.latestModel, "gpt-5.6-sol"); + assert.equal(metadata.effort, "ultra"); + assert.equal(metadata.latestReasoningEffort, "ultra"); + assert.equal(metadata.modelProvider, "aether"); + await adapter.dispose(); +}); + +test("IPC adapter projects a bounded model catalog from compatible state locations", async () => { + const state = fixtureState(); + // Exercise all of the state names used by different official extension + // builds. The same entries should be merged rather than duplicated. + state.availableModels = [{ model: "gpt-5.6-sol" }]; + state.models = [ + { model: "gpt-5.6-sol", description: "initial description" }, + { id: "gpt-5.6-terra", displayName: "5.6 Terra", efforts: ["low", "medium"] }, + ]; + state.modelCatalog = { + data: [{ + id: "gpt-5.6-sol", + model: "gpt-5.6-sol", + displayName: "5.6 Sol", + description: "通用 Codex 模型", + hidden: false, + isDefault: true, + defaultReasoningEffort: "medium", + supportedReasoningEfforts: [ + { reasoningEffort: "low", description: "快速" }, + { reasoningEffort: "high", description: "深入" }, + ], + apiKey: "sk-this-must-not-cross-the-relay", + capabilities: { internal: true }, + }], + nextCursor: "private-cursor", + }; + state.listModels = { data: [{ model: "gpt-5.6-terra", hidden: false }] }; + + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + const metadata = (await adapter.snapshot()).metadata; + assert.ok(Array.isArray(metadata.availableModels)); + assert.deepEqual(metadata.models, metadata.availableModels); + assert.equal(metadata.availableModels.length, 2); + const sol = metadata.availableModels.find((entry) => entry.model === "gpt-5.6-sol"); + assert.deepEqual(sol, { + model: "gpt-5.6-sol", + id: "gpt-5.6-sol", + description: "initial description", + displayName: "5.6 Sol", + hidden: false, + isDefault: true, + defaultReasoningEffort: "medium", + supportedReasoningEfforts: [ + { reasoningEffort: "low", description: "快速" }, + { reasoningEffort: "high", description: "深入" }, + ], + }); + const terra = metadata.availableModels.find((entry) => entry.model === "gpt-5.6-terra"); + assert.deepEqual(terra, { + model: "gpt-5.6-terra", + id: "gpt-5.6-terra", + displayName: "5.6 Terra", + supportedReasoningEfforts: [ + { reasoningEffort: "low" }, + { reasoningEffort: "medium" }, + ], + hidden: false, + }); + assert.equal(JSON.stringify(metadata.availableModels).includes("sk-this-must-not-cross-the-relay"), false); + assert.equal(JSON.stringify(metadata.availableModels).includes("capabilities"), false); + await adapter.dispose(); +}); + +test("IPC adapter publishes a snapshot when only thread settings metadata changes", async () => { + const state = fixtureState(); + state.latestModel = "gpt-5.6-terra"; + state.latestReasoningEffort = "medium"; + state.latestThreadSettings = { model: "gpt-5.6-terra", effort: "medium" }; + const client = new FakeIpcClient(state); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + await new Promise((resolve) => setImmediate(resolve)); + events.length = 0; + + const updated = structuredClone(state); + updated.latestModel = "gpt-5.6-sol"; + updated.latestReasoningEffort = "ultra"; + updated.latestThreadSettings = { model: "gpt-5.6-sol", effort: "ultra" }; + client.emitState(updated, 2, "patches"); + await new Promise((resolve) => setImmediate(resolve)); + + const snapshots = events.filter((event) => event.type === "session.snapshot"); + assert.equal(snapshots.length, 1); + assert.equal(snapshots[0].payload.metadata.latestModel, "gpt-5.6-sol"); + assert.equal(snapshots[0].payload.metadata.latestReasoningEffort, "ultra"); + await adapter.dispose(); +}); + +test("IPC adapter projects official activity, timestamps, command details, and turn duration", async () => { + const turnStartedAtMs = Date.now() - 20_000; + const finalAssistantStartedAtMs = turnStartedAtMs + 15_000; + const state = fixtureState(); + state.turns[0] = { + id: "turn-timed", + status: "completed", + turnStartedAtMs, + finalAssistantStartedAtMs, + durationMs: 18_000, + commandExecutionStartedAtMsById: { "command-1": turnStartedAtMs + 2_000 }, + items: [ + { type: "userMessage", id: "user-timed", content: [{ type: "text", text: "**run** it" }] }, + { + type: "commandExecution", + id: "command-1", + command: ["/bin/zsh", "-lc", "echo status-test"], + commandActions: [ + { type: "unknown", command: "/bin/zsh" }, + { type: "unknown", cmd: "echo status-test" }, + ], + cwd: "/tmp/workspace", + shellName: "zsh", + status: "completed", + aggregatedOutput: "status-test", + durationMs: 250, + exitCode: 0, + }, + { type: "agentMessage", id: "agent-timed", text: "done", phase: "final_answer" }, + ], + }; + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.status.activity, "completed"); + assert.equal(snapshot.durationMs, 18_000); + const [user, command, assistant] = snapshot.messages; + assert.equal(user.startedAtMs, turnStartedAtMs); + assert.equal(command.command, "echo status-test"); + assert.deepEqual(command.commandActions, [ + { type: "unknown", command: "/bin/zsh", cmd: "/bin/zsh" }, + { type: "unknown", cmd: "echo status-test", command: "echo status-test" }, + ]); + assert.equal(command.cwd, "/tmp/workspace"); + assert.equal(command.shellName, "zsh"); + assert.equal(command.startedAtMs, turnStartedAtMs + 2_000); + assert.equal(command.durationMs, 250); + assert.equal(command.exitCode, 0); + assert.equal(assistant.startedAtMs, finalAssistantStartedAtMs); + assert.equal(assistant.durationMs, 18_000); + + const active = structuredClone(state); + active.turns[0].status = "inProgress"; + delete active.turns[0].durationMs; + active.turns[0].items = [{ type: "reasoning", id: "reasoning-1", summary: ["checking"] }]; + client.emitState(active, 2, "patches"); + const activeSnapshot = await adapter.snapshot(); + assert.equal(activeSnapshot.activity, "thinking"); + assert.equal(activeSnapshot.turnStatus, "inprogress"); + assert.equal(activeSnapshot.turnId, "turn-timed"); + assert.equal(activeSnapshot.messages[0].turnStatus, "inProgress"); + + active.turns[0].items = [{ + type: "fileChange", + id: "edit-1", + status: "inProgress", + changes: [{ path: "src/example.ts", diff: "+const remote = true;" }], + }]; + client.emitState(active, 3, "patches"); + const editingSnapshot = await adapter.snapshot(); + assert.equal(editingSnapshot.activity, "editing"); + + active.turns[0].items = [{ + type: "commandExecution", + id: "command-active", + status: "inProgress", + command: "echo running", + }]; + client.emitState(active, 4, "patches"); + const runningSnapshot = await adapter.snapshot(); + assert.equal(runningSnapshot.activity, "running"); + + active.requests = [{ + id: "approval-active", + method: "item/commandExecution/requestApproval", + params: { threadId: THREAD_ID, turnId: "turn-timed", command: "echo approve" }, + }]; + client.emitState(active, 5, "patches"); + const waitingSnapshot = await adapter.snapshot(); + assert.equal(waitingSnapshot.activity, "waiting_approval"); + await adapter.dispose(); +}); + +test("IPC adapter keeps command-action-only items visible and suppresses shell bootstraps", async () => { + const state = fixtureState(); + state.turns[0] = { + id: "turn-command-actions", + status: "inProgress", + items: [ + { + // Older snapshots can expose a generic shell item while retaining + // the official commandActions payload. + type: "shell", + id: "command-actions-only", + status: "inProgress", + aggregatedOutput: "", + commandActions: [ + { type: "unknown", command: "/bin/zsh" }, + { type: "search", cmd: "rg --files", path: "src" }, + ], + cwd: "/tmp/workspace", + shellName: "zsh", + }, + { + type: "commandExecution", + id: "shell-bootstrap-only", + status: "inProgress", + aggregatedOutput: "", + command: "/bin/zsh", + commandActions: [{ type: "unknown", command: "/bin/zsh" }], + }, + { + type: "commandExecution", + id: "wrapped-command", + status: "inProgress", + aggregatedOutput: "", + command: "/bin/zsh -lc 'printf wrapped'", + }, + { + type: "commandExecution", + id: "wrapped-action", + status: "inProgress", + aggregatedOutput: "", + commandActions: ["/bin/zsh", "/bin/zsh -lc 'printf action-wrapped'"], + }, + ], + }; + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.messages.length, 3); + const actionsOnly = snapshot.messages.find((message) => message.itemId === "command-actions-only"); + assert.ok(actionsOnly); + assert.equal(actionsOnly.command, "rg --files"); + assert.equal(actionsOnly.text, "rg --files"); + assert.equal(actionsOnly.cwd, "/tmp/workspace"); + assert.equal(actionsOnly.shellName, "zsh"); + const wrapped = snapshot.messages.find((message) => message.itemId === "wrapped-command"); + assert.ok(wrapped); + assert.equal(wrapped.command, "'printf wrapped'"); + assert.equal(wrapped.text, "'printf wrapped'"); + const wrappedAction = snapshot.messages.find((message) => message.itemId === "wrapped-action"); + assert.ok(wrappedAction); + assert.equal(wrappedAction.command, "'printf action-wrapped'"); + assert.equal(wrappedAction.text, "'printf action-wrapped'"); + assert.doesNotMatch(snapshot.outputTail, /\/bin\/zsh/); + await adapter.dispose(); +}); + +test("IPC adapter projects collab items, metadata envelopes, and subagent lifecycle", async () => { + const childThreadId = "22222222-2222-4222-8222-222222222222"; + const state = fixtureState(); + state.id = THREAD_ID; + state.turns[0] = { + id: "turn-subagents", + status: "inProgress", + turnStartedAtMs: Date.now() - 2_000, + items: [ + { type: "userMessage", id: "user-subagents", content: [{ type: "text", text: "inspect this" }] }, + { + type: "agentMessage", + id: "message-with-collab-metadata", + text: "I am delegating this review", + metadata: { + codex_collab_agent_tool_call: { + type: "collabAgentToolCall", + id: "collab-spawn-1", + tool: "spawnAgent", + status: "inProgress", + senderThreadId: THREAD_ID, + receiverThreadIds: [childThreadId], + prompt: "Inspect the tests", + model: "gpt-5.6", + reasoningEffort: "high", + agentsStates: { [childThreadId]: { status: "running", message: null } }, + }, + }, + }, + { + type: "subAgentActivity", + id: "activity-started-1", + kind: "started", + agentThreadId: childThreadId, + agentPath: "root/reviewer_agent", + }, + ], + }; + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + let snapshot = await adapter.snapshot(); + assert.equal(snapshot.subagents.length, 1); + assert.equal(snapshot.subagents[0].threadId, childThreadId); + assert.equal(snapshot.subagents[0].displayName, "Reviewer agent"); + assert.equal(snapshot.subagents[0].prompt, "Inspect the tests"); + assert.equal(snapshot.subagents[0].objective, "Inspect the tests"); + assert.equal(snapshot.subagents[0].status, "working"); + assert.equal(snapshot.subagents[0].model, "gpt-5.6"); + assert.equal(snapshot.subagents[0].canInteract, true); + assert.ok(snapshot.messages.some((message) => message.itemType === "collabAgentToolCall" && message.uiType === "multi-agent-action")); + assert.ok(snapshot.messages.some((message) => message.itemType === "subAgentActivity" && message.uiType === "subagent-activity")); + + const completed = structuredClone(state); + completed.turns[0].status = "completed"; + completed.turns[0].items.push({ + type: "collabAgentToolCall", + id: "collab-wait-1", + tool: "wait", + status: "completed", + senderThreadId: THREAD_ID, + receiverThreadIds: [], + prompt: null, + model: null, + reasoningEffort: null, + agentsStates: {}, + }); + client.emitState(completed, 2, "patches"); + snapshot = await adapter.snapshot(); + assert.equal(snapshot.subagents[0].status, "done"); + await adapter.dispose(); +}); + +test("IPC adapter discovers pending requests retained inside official turn items", async () => { + const state = fixtureState(); + state.turns[0].status = "inProgress"; + state.turns[0].items.push({ + type: "permission-request", + id: "permission-in-turn", + threadId: THREAD_ID, + turnId: "turn-old", + summary: "需要访问工作区", + permissions: { fileSystem: { write: true } }, + }); + const client = new FakeIpcClient(state); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.pendingRequests.length, 1); + assert.equal(snapshot.pendingRequests[0].requestId, "permission-in-turn"); + assert.equal(snapshot.pendingRequests[0].method, "item/permissions/requestApproval"); + assert.deepEqual(snapshot.pendingRequests[0].params.permissions, { fileSystem: { write: true } }); + assert.equal(snapshot.messages.some((message) => message && message.itemId === "permission-in-turn"), false); + await adapter.dispose(); +}); + +test("IPC adapter expires unanswered requests and rolls back a failed follow", async () => { + const client = new FakeIpcClient(fixtureState([{ + id: 8, + method: "item/commandExecution/requestApproval", + params: { threadId: THREAD_ID, command: "echo timeout" }, + }])); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 25, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + await new Promise((resolve) => setTimeout(resolve, 60)); + assert.ok(events.some((event) => event.type === "approval.expired" && event.requestId === 8)); + assert.equal((await adapter.snapshot()).pendingApprovals.length, 0); + const expiryResponse = client.calls.find((call) => call.method === "respondCommandApproval"); + assert.equal(expiryResponse.decision, "decline"); + await adapter.dispose(); + + const noSnapshotClient = new FakeIpcClient(fixtureState(), false); + const failed = new CodexIpcAgentAdapter({ + client: noSnapshotClient, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 15, + }); + await assert.rejects(() => failed.start(), /Timed out waiting for a snapshot/); + assert.equal((await failed.snapshot()).state, "disconnected"); + assert.equal((await failed.snapshot()).threadId, null); + await failed.dispose(); +}); + +test("IPC adapter stays available while waiting for the first VS Code Codex session", async (t) => { + // An empty rollout directory represents a freshly opened VS Code window + // whose Codex panel has not created/selected a conversation yet. Starting + // the relay must remain possible so a later panel navigation can attach + // without restarting the bridge. + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-no-session-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const client = new FakeIpcClient(fixtureState()); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 25, + followTimeoutMs: 50, + }); + adapter.onEvent((event) => events.push(event)); + + await assert.doesNotReject(() => adapter.start()); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.threadId, null); + assert.equal(snapshot.state, "waiting_for_host"); + assert.equal(snapshot.metadata?.waitingForSession, true); + assert.equal(snapshot.metadata?.attachReady, false); + assert.equal(client.calls.some((call) => call.method === "followConversation" && call.following === true), false); + assert.ok(events.some((event) => event.type === "connection.opened" && event.threadId === undefined)); + + await adapter.dispose(); +}); + +test("IPC adapter attaches in-place when the official panel selects a session after waiting", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-route-after-wait-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const client = new FakeIpcClient(fixtureState()); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 100, + followTimeoutMs: 250, + }); + adapter.onEvent((event) => events.push(event)); + + await adapter.start(); + assert.equal((await adapter.snapshot()).state, "waiting_for_host"); + // The official webview broadcasts this untargeted route update when the + // user opens a conversation. The bridge should attach over the same IPC + // client rather than asking the user to restart the command. + client.emitFollowing(THREAD_ID, true, { sourceClientId: "official-vscode-panel" }); + await waitFor(async () => (await adapter.snapshot()).threadId === THREAD_ID); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.state, "idle"); + assert.equal(snapshot.metadata.waitingForSession, false); + assert.equal(snapshot.metadata.attachReady, true); + assert.equal(client.calls.filter((call) => call.method === "followConversation" && call.following === true).length, 1); + assert.equal(events.filter((event) => event.type === "connection.opened").length, 1); + assert.ok(events.some((event) => event.type === "session.snapshot" && event.threadId === THREAD_ID)); + await adapter.dispose(); +}); + +test("IPC waiting attach follows the latest official route when selection changes mid-snapshot", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-wait-route-race-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "latest route" }], + ]); + const client = new WaitingRouteRaceIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 100, + followTimeoutMs: 250, + vscodeSessionFollowDebounceMs: 0, + }); + + await adapter.start(); + client.emitFollowing(THREAD_ID, true, { sourceClientId: "official-vscode-panel" }); + await waitFor(() => client.calls.some((call) => call.method === "followConversation" + && call.threadId === THREAD_ID && call.following === true)); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "official-vscode-panel" }); + await waitFor(async () => { + const snapshot = await adapter.snapshot(); + return snapshot.threadId === SECOND_THREAD_ID && snapshot.metadata.attachReady === true; + }); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.attachReady, true); + assert.equal(snapshot.metadata.title, "latest route"); + assert.ok(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === true)); + await adapter.dispose(); +}); + +test("IPC waiting discovery never overrides a newer official VS Code route", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-wait-poll-route-race-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "stale fallback" }], + ]); + const client = new WaitingPollRouteRaceIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 100, + followTimeoutMs: 250, + vscodeSessionFollowDebounceMs: 0, + }); + + await adapter.start(); + assert.equal((await adapter.snapshot()).state, "waiting_for_host"); + + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + await fs.writeFile(path.join(sessions, `rollout-${SECOND_THREAD_ID}.jsonl`), `${JSON.stringify({ + type: "session_meta", + payload: { originator: "codex_vscode", source: "vscode", cwd: "/tmp/stale-fallback" }, + })}\n`); + + client.delayFallbackDiscovery = true; + const fallbackDiscovery = adapter.runWaitingDiscovery(); + await waitFor(() => client.fallbackDiscoveryStarted); + client.emitFollowing(THREAD_ID, true, { sourceClientId: "official-vscode-panel" }); + await waitFor(() => client.calls.some((call) => call.method === "followConversation" + && call.threadId === THREAD_ID && call.following === true)); + client.releaseFallbackDiscovery(); + await fallbackDiscovery; + await waitFor(async () => (await adapter.snapshot()).metadata.attachReady === true); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.threadId, THREAD_ID); + assert.equal(snapshot.metadata.title, "fixture session"); + assert.equal(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === true), false); + await adapter.dispose(); +}); + +test("IPC adapter waits when the configured VS Code conversation has no live owner", async () => { + const client = new FakeIpcClient(fixtureState()); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: STALE_THREAD_ID, + autoDiscoverThread: false, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 25, + followTimeoutMs: 50, + }); + + await assert.doesNotReject(() => adapter.start()); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.threadId, null); + assert.equal(snapshot.state, "waiting_for_host"); + assert.equal(snapshot.metadata?.waitingForSession, true); + assert.equal(snapshot.metadata?.attachReady, false); + assert.ok(client.calls.some((call) => call.method === "findThreadOwner" && call.threadId === STALE_THREAD_ID)); + assert.equal(client.calls.some((call) => call.method === "followConversation" && call.following === true), false); + + await adapter.dispose(); +}); + +test("IPC adapter preserves official request timestamps and streams a sliding output tail as a delta", async () => { + const startedAtMs = Date.now() - 200; + const initial = fixtureState([{ + id: 9, + method: "item/commandExecution/requestApproval", + createdAt: Date.now(), + params: { threadId: THREAD_ID, command: "echo timestamp", startedAtMs }, + }]); + initial.turns[0].items = [{ type: "agentMessage", id: "streaming", text: "x".repeat(40) }]; + const client = new FakeIpcClient(initial); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 10_000, + maxOutputTailChars: 32, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + const pending = (await adapter.snapshot()).pendingApprovals[0]; + assert.equal(pending.createdAt, startedAtMs); + + const next = structuredClone(initial); + next.turns[0].items[0].text += "y"; + client.emitState(next, 2); + const outputEvents = events.filter((event) => event.type === "output.snapshot" || event.type === "output.chunk"); + assert.equal(outputEvents.at(-1).type, "output.chunk"); + assert.equal(outputEvents.at(-1).payload.text, "y"); + assert.equal(outputEvents.at(-1).payload.messages, undefined); + assert.equal(outputEvents.at(-1).payload.messagesPatch.start, 0); + assert.equal(outputEvents.at(-1).payload.messagesPatch.deleteCount, 1); + assert.equal(outputEvents.at(-1).payload.messagesPatch.messages[0].text, `${"x".repeat(40)}y`); + + // The bounded tail can remain byte-for-byte identical when a repeated + // character arrives. It must still advance by one chunk using total length. + const repeated = structuredClone(next); + repeated.turns[0].items[0].text += "x"; + client.emitState(repeated, 3); + const repeatedEvent = events.filter((event) => event.type === "output.snapshot" || event.type === "output.chunk").at(-1); + assert.equal(repeatedEvent.type, "output.chunk"); + assert.equal(repeatedEvent.payload.text, "x"); + assert.equal(repeatedEvent.payload.messagesPatch.messages[0].text, `${"x".repeat(40)}yx`); + await adapter.dispose(); +}); + +test("IPC adapter normalizes epoch-second approval timestamps", async () => { + const nowSeconds = Math.floor(Date.now() / 1000); + const startedSeconds = nowSeconds - 2; + const expiresSeconds = nowSeconds + 60; + const initial = fixtureState([ + { + id: 10, + method: "item/commandExecution/requestApproval", + createdAt: startedSeconds, + expiresAt: expiresSeconds, + params: { threadId: THREAD_ID, command: "echo outer seconds" }, + }, + { + id: 11, + method: "item/commandExecution/requestApproval", + params: { + threadId: THREAD_ID, + command: "echo nested seconds", + startedAt: String(startedSeconds), + expiresAt: String(expiresSeconds), + }, + }, + ]); + const client = new FakeIpcClient(initial); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 10_000, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + + await adapter.start(); + const pending = (await adapter.snapshot()).pendingApprovals; + assert.equal(pending.length, 2); + for (const approval of pending) { + assert.equal(approval.createdAt, startedSeconds * 1000); + assert.equal(approval.expiresAt, expiresSeconds * 1000); + } + assert.equal(events.some((event) => event.type === "approval.expired"), false); + await adapter.dispose(); +}); + +test("IPC adapter waits for the revision returned by complete-history loading", async () => { + const initial = fixtureState(); + initial.turnsPagination = { olderCursor: "older", hasLoadedOldest: false, isLoadingOlder: false }; + const complete = fixtureState(); + complete.turnsPagination = { olderCursor: null, hasLoadedOldest: true, isLoadingOlder: false }; + class HistoryClient extends FakeIpcClient { + async loadCompleteHistory() { + this.calls.push({ method: "loadCompleteHistory" }); + // The owner can acknowledge the request before the stream broadcast. + // An unrelated owner's revision must not release our waiter. + setTimeout(() => this.emitState(initial, 2, "snapshot", THREAD_ID, "other-owner"), 0); + setTimeout(() => this.emitState(complete, 2, "snapshot", THREAD_ID, "owner"), 10); + return { revision: 2 }; + } + } + const client = new HistoryClient(initial); + const adapter = new CodexIpcAgentAdapter({ client, threadId: THREAD_ID, followTimeoutMs: 500, approvalTimeoutMs: 0 }); + await adapter.start(); + await new Promise((resolve) => setTimeout(resolve, 5)); + // The unrelated owner event at t=0 must not overwrite the attached stream. + assert.equal((await adapter.snapshot()).metadata.historyComplete, false); + await new Promise((resolve) => setTimeout(resolve, 20)); + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.historyComplete, true); + assert.equal(client.calls.filter((call) => call.method === "loadCompleteHistory").length, 1); + await adapter.dispose(); +}); + +test("IPC adapter retries complete-history loading once after a transient failure", async () => { + const initial = fixtureState(); + initial.turnsPagination = { olderCursor: "older", hasLoadedOldest: false, isLoadingOlder: false }; + const complete = fixtureState(); + complete.turnsPagination = { olderCursor: null, hasLoadedOldest: true, isLoadingOlder: false }; + class RetryingHistoryClient extends FakeIpcClient { + async loadCompleteHistory() { + this.calls.push({ method: "loadCompleteHistory" }); + const attempts = this.calls.filter((call) => call.method === "loadCompleteHistory").length; + if (attempts === 1) throw Object.assign(new Error("history request timed out"), { code: "timeout" }); + setTimeout(() => this.emitState(complete, 2, "snapshot", THREAD_ID, "owner"), 5); + return { revision: 2 }; + } + } + const client = new RetryingHistoryClient(initial); + const adapter = new CodexIpcAgentAdapter({ client, threadId: THREAD_ID, followTimeoutMs: 500, approvalTimeoutMs: 0 }); + await adapter.start(); + await new Promise((resolve) => setTimeout(resolve, 100)); + assert.equal(client.calls.filter((call) => call.method === "loadCompleteHistory").length, 2); + assert.equal((await adapter.snapshot()).metadata.historyComplete, true); + await adapter.dispose(); +}); + +test("IPC adapter cancels a pending history retry when switching sessions", async () => { + const initial = fixtureState(); + initial.turnsPagination = { olderCursor: "older", hasLoadedOldest: false, isLoadingOlder: false }; + const states = new Map([[THREAD_ID, initial], [SECOND_THREAD_ID, fixtureState()]]); + class SwitchingHistoryClient extends MultiSessionIpcClient { + async loadCompleteHistory(threadId) { + this.calls.push({ method: "loadCompleteHistory", threadId }); + throw Object.assign(new Error("history request timed out"), { code: "timeout" }); + } + } + const client = new SwitchingHistoryClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ client, threadId: THREAD_ID, followTimeoutMs: 500, approvalTimeoutMs: 0 }); + await adapter.start(); + await new Promise((resolve) => setImmediate(resolve)); + await adapter.selectSession({ threadId: SECOND_THREAD_ID }); + await new Promise((resolve) => setTimeout(resolve, 80)); + assert.deepEqual(client.calls.filter((call) => call.method === "loadCompleteHistory").map((call) => call.threadId), [THREAD_ID]); + assert.equal((await adapter.snapshot()).threadId, SECOND_THREAD_ID); + await adapter.dispose(); +}); + +test("IPC adapter cancels a pending history retry when disposed", async () => { + const initial = fixtureState(); + initial.turnsPagination = { olderCursor: "older", hasLoadedOldest: false, isLoadingOlder: false }; + class DisposedHistoryClient extends FakeIpcClient { + async loadCompleteHistory() { + this.calls.push({ method: "loadCompleteHistory" }); + throw Object.assign(new Error("history request timed out"), { code: "timeout" }); + } + } + const client = new DisposedHistoryClient(initial); + const adapter = new CodexIpcAgentAdapter({ client, threadId: THREAD_ID, followTimeoutMs: 500, approvalTimeoutMs: 0 }); + await adapter.start(); + await new Promise((resolve) => setImmediate(resolve)); + await adapter.dispose(); + await new Promise((resolve) => setTimeout(resolve, 80)); + assert.equal(client.calls.filter((call) => call.method === "loadCompleteHistory").length, 1); +}); + +test("IPC adapter reports canonical history incomplete until one complete island and all item pages", async () => { + const state = fixtureState(); + state.turns = []; + state.turnHistory = { + kind: "canonical", + history: { + isComplete: true, + islands: [{ entries: [] }, { entries: [] }], + entitiesByKey: { + "turn:1": { + id: "turn-1", + status: "completed", + items: [], + itemsPagination: { hasLoadedOldest: true }, + }, + }, + }, + }; + const client = new FakeIpcClient(state); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + assert.equal((await adapter.snapshot()).metadata.historyComplete, false); + + const itemPagePending = structuredClone(state); + itemPagePending.turnHistory.history.islands = [{ entries: [] }]; + itemPagePending.turnHistory.history.entitiesByKey["turn:1"].itemsPagination.hasLoadedOldest = false; + client.emitState(itemPagePending, 2); + assert.equal((await adapter.snapshot()).metadata.historyComplete, false); + + const complete = structuredClone(itemPagePending); + complete.turnHistory.history.entitiesByKey["turn:1"].itemsPagination.hasLoadedOldest = true; + events.length = 0; + client.emitState(complete, 3, "patches"); + assert.equal((await adapter.snapshot()).metadata.historyComplete, true); + await new Promise((resolve) => setImmediate(resolve)); + assert.ok(events.some((event) => event.type === "session.snapshot")); + await adapter.dispose(); +}); + +test("IPC auto-discovery excludes subagents and Codex Desktop tasks", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-discovery-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + + const preferredCwd = path.join(codexHome, "current-workspace"); + const subagentId = "22222222-2222-4222-8222-222222222222"; + const otherWorkspaceId = "33333333-3333-4333-8333-333333333333"; + const matchingDesktopId = "44444444-4444-4444-8444-444444444444"; + const matchingVscodeId = "55555555-5555-4555-8555-555555555555"; + const writeRollout = async (id, payload, mtimeSeconds) => { + const fileName = path.join(sessions, `rollout-2026-08-28T00-00-00-${id}.jsonl`); + await fs.writeFile(fileName, `${JSON.stringify({ type: "session_meta", payload })}\n`); + await fs.utimes(fileName, mtimeSeconds, mtimeSeconds); + }; + + await writeRollout(subagentId, { + originator: "codex_vscode", + source: { subagent: { thread_spawn: {} } }, + thread_source: "subagent", + cwd: preferredCwd, + }, 3_000); + await writeRollout(otherWorkspaceId, { + originator: "codex_vscode", + source: "vscode", + thread_source: "user", + cwd: path.join(codexHome, "other-workspace"), + }, 2_000); + await writeRollout(matchingDesktopId, { + originator: "Codex Desktop", + source: "vscode", + thread_source: "user", + cwd: preferredCwd, + }, 4_000); + await writeRollout(matchingVscodeId, { + originator: "codex_vscode", + source: "vscode", + thread_source: "user", + cwd: preferredCwd, + }, 1_000); + + const client = new DiscoveryIpcClient(fixtureState()); + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + preferredCwds: [preferredCwd], + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + + await adapter.start(); + assert.equal((await adapter.snapshot()).threadId, matchingVscodeId); + const discovered = client.calls + .filter((call) => call.method === "findThreadOwner") + .map((call) => call.threadId); + assert.equal(discovered.includes(subagentId), false); + assert.equal(discovered.includes(matchingDesktopId), false); + await adapter.dispose(); +}); + +test("IPC auto-discovery finds an older live conversation beyond the first twelve rollouts", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-deep-discovery-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + const staleIds = Array.from({ length: 20 }, (_, index) => `aaaaaaaa-aaaa-4aaa-8aaa-${String(index).padStart(12, "0")}`); + for (const [index, id] of staleIds.entries()) { + await fs.writeFile(path.join(sessions, `rollout-${id}.jsonl`), `${JSON.stringify({ + type: "session_meta", + payload: { + originator: "codex_vscode", + source: "vscode", + cwd: `/tmp/stale-${index}`, + updated_at: `2026-08-29T${String(20 - index).padStart(2, "0")}:00:00Z`, + }, + })}\n`); + } + const liveRollout = path.join(sessions, `rollout-${THREAD_ID}.jsonl`); + await fs.writeFile(liveRollout, `${JSON.stringify({ + type: "session_meta", + payload: { + originator: "codex_vscode", + source: "vscode", + cwd: "/tmp/live-old", + updated_at: "2026-08-01T00:00:00Z", + }, + })}\n`); + const oldMtime = new Date("2026-08-01T00:00:00Z"); + await fs.utimes(liveRollout, oldMtime, oldMtime); + + const client = new FakeIpcClient(fixtureState()); + const adapter = new CodexIpcAgentAdapter({ + client, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + assert.equal((await adapter.snapshot()).threadId, THREAD_ID); + const discoveryCalls = client.calls.filter((call) => call.method === "findThreadOwner"); + assert.ok(discoveryCalls.length > 12); + assert.ok(discoveryCalls.some((call) => call.threadId === THREAD_ID)); + await adapter.dispose(); +}); + +test("IPC adapter lists only indexed VS Code sessions with live attachable snapshots", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-session-list-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + const writeRollout = async (id, payload) => { + const fileName = path.join(sessions, `rollout-${id}.jsonl`); + await fs.writeFile(fileName, `${JSON.stringify({ type: "session_meta", payload })}\n`); + }; + await writeRollout(THREAD_ID, { originator: "codex_vscode", source: "vscode", cwd: "/tmp/a" }); + await writeRollout(SECOND_THREAD_ID, { originator: "codex_vscode", source: "vscode", cwd: "/tmp/b" }); + await writeRollout(STALE_THREAD_ID, { originator: "codex_vscode", source: "vscode", cwd: "/tmp/c" }); + await fs.writeFile(path.join(codexHome, "session_index.jsonl"), [ + JSON.stringify({ id: THREAD_ID, thread_name: "当前会话", updated_at: "2026-08-29T10:00:00Z" }), + JSON.stringify({ id: SECOND_THREAD_ID, thread_name: "另一个工作区", updated_at: "2026-08-29T11:00:00Z" }), + JSON.stringify({ id: STALE_THREAD_ID, thread_name: "已关闭会话", updated_at: "2026-08-29T12:00:00Z" }), + ].join("\n") + "\n"); + + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "另一个工作区" }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const result = await adapter.listSessions({ limit: 10 }); + assert.equal(result.activeThreadId, THREAD_ID); + assert.deepEqual(result.sessions.map((entry) => entry.threadId), [SECOND_THREAD_ID, THREAD_ID]); + assert.equal(result.sessions.find((entry) => entry.threadId === SECOND_THREAD_ID).available, true); + assert.equal(result.sessions.find((entry) => entry.threadId === STALE_THREAD_ID), undefined); + assert.equal(result.sessions.find((entry) => entry.threadId === SECOND_THREAD_ID).title, "另一个工作区"); + assert.ok(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === true)); + assert.ok(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === false)); + await adapter.dispose(); +}); + +test("IPC session list keeps the active attachment when newer stale history fills the limit", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-session-limit-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + const writeRollout = async (id, cwd) => { + await fs.writeFile(path.join(sessions, `rollout-${id}.jsonl`), `${JSON.stringify({ + type: "session_meta", + payload: { originator: "codex_vscode", source: "vscode", cwd }, + })}\n`); + }; + await writeRollout(THREAD_ID, "/tmp/current"); + await writeRollout(STALE_THREAD_ID, "/tmp/stale"); + await fs.writeFile(path.join(codexHome, "session_index.jsonl"), [ + JSON.stringify({ id: THREAD_ID, thread_name: "当前会话", updated_at: "2026-08-28T10:00:00Z" }), + JSON.stringify({ id: STALE_THREAD_ID, thread_name: "较新的失效记录", updated_at: "2026-08-29T12:00:00Z" }), + ].join("\n") + "\n"); + + const client = new MultiSessionIpcClient(new Map([[THREAD_ID, fixtureState()]]), { [THREAD_ID]: "owner-a" }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const result = await adapter.listSessions({ limit: 1 }); + assert.deepEqual(result.sessions.map((entry) => entry.threadId), [THREAD_ID]); + assert.equal(result.sessions[0].active, true); + await adapter.dispose(); +}); + +test("IPC session list omits an owner that does not return a matching snapshot", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-session-probe-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + const writeRollout = async (id, payload) => { + const fileName = path.join(sessions, `rollout-${id}.jsonl`); + await fs.writeFile(fileName, `${JSON.stringify({ type: "session_meta", payload })}\n`); + }; + await writeRollout(THREAD_ID, { originator: "codex_vscode", source: "vscode", cwd: "/tmp/a" }); + await writeRollout(SECOND_THREAD_ID, { originator: "codex_vscode", source: "vscode", cwd: "/tmp/b" }); + await fs.writeFile(path.join(codexHome, "session_index.jsonl"), [ + JSON.stringify({ id: THREAD_ID, thread_name: "当前会话", updated_at: "2026-08-29T10:00:00Z" }), + JSON.stringify({ id: SECOND_THREAD_ID, thread_name: "桌面会话", updated_at: "2026-08-29T11:00:00Z" }), + ].join("\n") + "\n"); + + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "桌面会话" }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "desktop-owner", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + ownerDiscoveryTimeoutMs: 250, + followTimeoutMs: 250, + }); + await adapter.start(); + const result = await adapter.listSessions({ limit: 10 }); + assert.equal(result.sessions.find((entry) => entry.threadId === THREAD_ID).available, true); + assert.equal(result.sessions.find((entry) => entry.threadId === SECOND_THREAD_ID), undefined); + const targetCalls = client.calls.filter((call) => call.method === "followConversation" && call.threadId === SECOND_THREAD_ID); + assert.equal(targetCalls.filter((call) => call.following === true).length, 1); + assert.equal(targetCalls.filter((call) => call.following === false).length, 1); + await adapter.dispose(); +}); + +test("IPC adapter switches follows only after the target owner snapshot arrives", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "目标会话", turns: [{ id: "turn-target", status: "completed", items: [{ type: "agentMessage", id: "target-message", text: "target output" }] }] }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + const result = await adapter.selectSession({ threadId: SECOND_THREAD_ID }); + assert.deepEqual(result, { + threadId: SECOND_THREAD_ID, + previousThreadId: THREAD_ID, + switched: true, + available: true, + }); + assert.equal((await adapter.snapshot()).threadId, SECOND_THREAD_ID); + assert.match((await adapter.snapshot()).outputTail, /target output/); + assert.ok(events.some((event) => event.type === "session.switching")); + assert.ok(events.some((event) => event.type === "session.selected")); + const oldUnfollow = client.calls.find((call) => call.method === "followConversation" && call.threadId === THREAD_ID && call.following === false); + assert.ok(oldUnfollow); + await adapter.dispose(); +}); + +test("IPC adapter follows paired route changes from the attached VS Code panel", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { + ...fixtureState(), + id: SECOND_THREAD_ID, + title: "面板目标会话", + turns: [{ id: "turn-target", status: "completed", items: [{ type: "agentMessage", id: "target-message", text: "panel target output" }] }], + }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "target-owner", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 5, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(async () => (await adapter.snapshot()).threadId === SECOND_THREAD_ID); + + const snapshot = await adapter.snapshot(); + assert.match(snapshot.outputTail, /panel target output/); + assert.ok(events.some((event) => event.type === "session.switching" && event.threadId === SECOND_THREAD_ID)); + assert.ok(events.some((event) => event.type === "session.selected" && event.threadId === SECOND_THREAD_ID)); + await adapter.dispose(); +}); + +test("IPC adapter ignores isolated, targeted, and other-client following broadcasts", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "不应自动切换" }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "target-owner", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 0, + }); + await adapter.start(); + + // Bind the route source with a same-thread leave/re-enter pair. + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(THREAD_ID, true, { sourceClientId: "panel-owner" }); + // A reconnect/status replay is a lone true and is not a route change. + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + // Targeted replies describe liveness to one follower, not panel navigation. + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner", targetClientIds: ["follower"] }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner", targetClientIds: ["follower"] }); + // Another Codex window shares the router but cannot take over this bridge. + client.emitFollowing(THREAD_ID, false, { sourceClientId: "other-panel" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "other-panel" }); + await new Promise((resolve) => setTimeout(resolve, 30)); + + assert.equal((await adapter.snapshot()).threadId, THREAD_ID); + assert.equal(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === true), false); + await adapter.dispose(); +}); + +test("IPC adapter does not trust an isolated false from another follower", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "真实面板目标" }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 0, + }); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "disposing-remote-follower" }); + client.emitFollowing(THREAD_ID, false, { sourceClientId: "real-vscode-panel" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "real-vscode-panel" }); + await waitFor(async () => (await adapter.snapshot()).threadId === SECOND_THREAD_ID); + await adapter.dispose(); +}); + +test("IPC adapter binds the matching route source when another unbound source emits false", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { + ...fixtureState(), + id: SECOND_THREAD_ID, + title: "真实面板目标", + turns: [{ id: "turn-target", status: "completed", items: [{ type: "agentMessage", id: "target-message", text: "real panel target" }] }], + }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 0, + }); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "real-vscode-panel" }); + client.emitFollowing(THREAD_ID, false, { sourceClientId: "other-follower" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "real-vscode-panel" }); + await waitFor(async () => (await adapter.snapshot()).threadId === SECOND_THREAD_ID); + + assert.match((await adapter.snapshot()).outputTail, /real panel target/); + await adapter.dispose(); +}); + +test("IPC adapter lets a replacement route source take over after the bound client disconnects", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { + ...fixtureState(), + id: SECOND_THREAD_ID, + title: "重连后的目标", + turns: [{ id: "turn-target", status: "completed", items: [{ type: "agentMessage", id: "target-message", text: "replacement panel target" }] }], + }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 0, + }); + await adapter.start(); + + // First bind the official panel without changing the selected conversation. + client.emitFollowing(THREAD_ID, false, { sourceClientId: "old-vscode-panel" }); + client.emitFollowing(THREAD_ID, true, { sourceClientId: "old-vscode-panel" }); + client.emitClientStatus("old-vscode-panel", "disconnected"); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "replacement-vscode-panel" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "replacement-vscode-panel" }); + await waitFor(async () => (await adapter.snapshot()).threadId === SECOND_THREAD_ID); + + assert.match((await adapter.snapshot()).outputTail, /replacement panel target/); + await adapter.dispose(); +}); + +test("IPC adapter coalesces rapid VS Code A to B to C navigation to C", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "中间会话" }], + [THIRD_THREAD_ID, { + ...fixtureState(), + id: THIRD_THREAD_ID, + title: "最终会话", + turns: [{ id: "turn-c", status: "completed", items: [{ type: "agentMessage", id: "message-c", text: "final C output" }] }], + }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "owner-b", + [THIRD_THREAD_ID]: "owner-c", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 20, + }); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(THIRD_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(async () => (await adapter.snapshot()).threadId === THIRD_THREAD_ID); + + assert.match((await adapter.snapshot()).outputTail, /final C output/); + assert.equal(client.calls.some((call) => call.method === "followConversation" + && call.threadId === SECOND_THREAD_ID && call.following === true), false); + await adapter.dispose(); +}); + +test("IPC adapter cancels an in-flight B snapshot when VS Code moves on to C", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "无快照 B" }], + [THIRD_THREAD_ID, { + ...fixtureState(), + id: THIRD_THREAD_ID, + title: "最终 C", + turns: [{ id: "turn-c", status: "completed", items: [{ type: "agentMessage", id: "message-c", text: "C arrived" }] }], + }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "owner-b", + [THIRD_THREAD_ID]: "owner-c", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 1_000, + vscodeSessionFollowDebounceMs: 0, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(() => events.some((event) => event.type === "session.switching" && event.threadId === SECOND_THREAD_ID)); + const movedOnAt = Date.now(); + client.emitFollowing(SECOND_THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(THIRD_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(async () => (await adapter.snapshot()).threadId === THIRD_THREAD_ID, 500); + + assert.ok(Date.now() - movedOnAt < 500, "C should not wait for B's 1s snapshot timeout"); + assert.equal(events.some((event) => event.type === "session.selected" + && event.threadId === SECOND_THREAD_ID && event.payload.switched === true), false); + assert.match((await adapter.snapshot()).outputTail, /C arrived/); + await adapter.dispose(); +}); + +test("IPC adapter defers the latest VS Code route while the old turn is active", async () => { + const active = fixtureState(); + active.turns[0].status = "inProgress"; + active.threadRuntimeStatus = { type: "active" }; + const completed = structuredClone(active); + completed.turns[0].status = "completed"; + completed.threadRuntimeStatus = { type: "idle" }; + const states = new Map([ + [THREAD_ID, active], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "延后目标" }], + ]); + const client = new MultiSessionIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "target-owner", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + vscodeSessionFollowDebounceMs: 0, + }); + await adapter.start(); + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await new Promise((resolve) => setTimeout(resolve, 40)); + assert.equal((await adapter.snapshot()).threadId, THREAD_ID); + + client.emitState(completed, 2, "patches", THREAD_ID, "panel-owner"); + await waitFor(async () => (await adapter.snapshot()).threadId === SECOND_THREAD_ID, 1_000); + await adapter.dispose(); +}); + +test("IPC adapter publishes an explicit rollback when a VS Code route cannot stream", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "无快照目标" }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "target-owner", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 25, + vscodeSessionFollowDebounceMs: 0, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + events.length = 0; + + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(() => events.some((event) => event.type === "session.selected" && event.payload.failed === true)); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.threadId, THREAD_ID); + assert.match(snapshot.outputTail, /hi from VS Code/); + const rollbackIndex = events.findIndex((event) => event.type === "session.selected" && event.payload.failed === true); + const restoredOutputIndex = events.findIndex((event, index) => index > rollbackIndex + && event.type === "output.snapshot" && event.threadId === THREAD_ID); + assert.ok(rollbackIndex >= 0); + assert.ok(restoredOutputIndex > rollbackIndex); + await adapter.dispose(); +}); + +test("IPC adapter keeps dispose authoritative while an automatic switch is waiting", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "等待中的目标" }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "panel-owner", + [SECOND_THREAD_ID]: "target-owner", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 1_000, + vscodeSessionFollowDebounceMs: 0, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + client.emitFollowing(THREAD_ID, false, { sourceClientId: "panel-owner" }); + client.emitFollowing(SECOND_THREAD_ID, true, { sourceClientId: "panel-owner" }); + await waitFor(() => events.some((event) => event.type === "session.switching")); + + const eventCountAtDispose = events.length; + await adapter.dispose(); + await new Promise((resolve) => setTimeout(resolve, 30)); + assert.equal((await adapter.snapshot()).state, "disconnected"); + assert.equal(events.slice(eventCountAtDispose).some((event) => [ + "session.selected", + "output.snapshot", + "session.snapshot", + ].includes(event.type)), false); +}); + +test("IPC adapter absorbs a snapshot waiter when target follow fails immediately", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "立即失败" }], + ]); + const client = new ThrowingTargetFollowIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + await assert.rejects(() => adapter.selectSession({ threadId: SECOND_THREAD_ID }), /target follow failed immediately/); + assert.equal((await adapter.snapshot()).threadId, THREAD_ID); + await adapter.dispose(); +}); + +test("IPC adapter rejects a target snapshot when owner changes before commit", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "Owner handoff" }], + ]); + const client = new OwnerChangingIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + await assert.rejects( + () => adapter.selectSession({ threadId: SECOND_THREAD_ID }), + /owner changed while switching/, + ); + assert.equal((await adapter.snapshot()).threadId, THREAD_ID); + await adapter.dispose(); +}); + +test("IPC adapter serializes list probes before selecting the same session", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-session-serialization-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + for (const [id, cwd] of [[THREAD_ID, "/tmp/current"], [SECOND_THREAD_ID, "/tmp/target"]]) { + await fs.writeFile(path.join(sessions, `rollout-${id}.jsonl`), `${JSON.stringify({ + type: "session_meta", + payload: { originator: "codex_vscode", source: "vscode", cwd }, + })}\n`); + } + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "目标会话" }], + ]); + const client = new DelayedProbeIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + const listPromise = adapter.listSessions({ limit: 10 }); + await new Promise((resolve) => setTimeout(resolve, 5)); + const selectPromise = adapter.selectSession({ threadId: SECOND_THREAD_ID }); + await Promise.all([listPromise, selectPromise]); + + const targetFollowing = client.calls + .filter((call) => call.method === "followConversation" && call.threadId === SECOND_THREAD_ID) + .map((call) => call.following); + assert.deepEqual(targetFollowing, [true, false, true]); + assert.equal((await adapter.snapshot()).threadId, SECOND_THREAD_ID); + await adapter.dispose(); +}); + +test("IPC adapter restores the cached previous projection when target follow has no snapshot", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "目标会话" }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const events = []; + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 25, + }); + adapter.onEvent((event) => events.push(event)); + await adapter.start(); + const before = await adapter.snapshot(); + assert.equal(before.threadId, THREAD_ID); + assert.match(before.outputTail, /hi from VS Code/); + + await assert.rejects( + () => adapter.selectSession({ threadId: SECOND_THREAD_ID }), + /Timed out waiting for a snapshot from VS Code conversation/, + ); + + const after = await adapter.snapshot(); + assert.equal(after.threadId, THREAD_ID); + assert.equal(after.outputTail, before.outputTail); + assert.deepEqual(after.messages, before.messages); + assert.equal(after.state, "idle"); + assert.ok(events.some((event) => event.type === "output.snapshot" && event.threadId === THREAD_ID)); + assert.ok(events.some((event) => event.type === "session.snapshot" && event.threadId === THREAD_ID)); + await adapter.dispose(); +}); + +test("IPC adapter restores its owner-validated projection after the IPC cache is polluted", async () => { + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, { ...fixtureState(), id: SECOND_THREAD_ID, title: "无快照目标" }], + ]); + const client = new NoSnapshotSwitchIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 25, + }); + await adapter.start(); + const verified = await adapter.snapshot(); + + const polluted = { + ...fixtureState(), + title: "来自旧 owner 的污染快照", + turns: [{ + id: "turn-polluted", + status: "completed", + items: [{ type: "agentMessage", id: "polluted-message", text: "poisoned stale output" }], + }], + }; + client.emitState(polluted, 99, "snapshot", THREAD_ID, "old-owner"); + + assert.equal(client.getConversationState(THREAD_ID).ownerClientId, "old-owner"); + assert.equal(client.getConversationState(THREAD_ID).conversationState.title, "来自旧 owner 的污染快照"); + assert.equal((await adapter.snapshot()).outputTail, verified.outputTail); + + await assert.rejects( + () => adapter.selectSession({ threadId: SECOND_THREAD_ID }), + /Timed out waiting for a snapshot from VS Code conversation/, + ); + + const restored = await adapter.snapshot(); + assert.equal(restored.threadId, THREAD_ID); + assert.equal(restored.outputTail, verified.outputTail); + assert.deepEqual(restored.messages, verified.messages); + assert.doesNotMatch(restored.outputTail, /poisoned stale output/); + await adapter.dispose(); +}); + +test("IPC adapter denies target approvals on dispose and blocks ordinary commands mid-switch", async () => { + const targetApprovalId = "target-approval"; + const targetState = { + ...fixtureState([{ + id: targetApprovalId, + method: "item/commandExecution/requestApproval", + params: { threadId: SECOND_THREAD_ID, turnId: "turn-target", command: "echo target" }, + }]), + id: SECOND_THREAD_ID, + title: "等待 owner 确认的目标", + }; + const states = new Map([ + [THREAD_ID, fixtureState()], + [SECOND_THREAD_ID, targetState], + ]); + const client = new DelayedTargetOwnerConfirmationIpcClient(states, { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + + const switching = adapter.selectSession({ threadId: SECOND_THREAD_ID }); + void switching.catch(() => undefined); + await waitFor(() => client.targetConfirmationStarted); + await waitFor(async () => (await adapter.snapshot()).pendingApprovals.some((entry) => entry.requestId === targetApprovalId)); + + await assert.rejects( + () => adapter.startTurn({ text: "must not be sent while switching" }), + /session switch is still in progress/, + ); + await assert.rejects( + () => adapter.respondApproval(targetApprovalId, "allow"), + /session switch is still in progress/, + ); + assert.equal(client.calls.some((call) => call.method === "startTurn"), false); + assert.equal(client.calls.some((call) => call.method === "respondCommandApproval"), false); + + await adapter.dispose(); + const denied = client.calls.find((call) => call.method === "respondCommandApproval" + && call.requestId === targetApprovalId); + assert.ok(denied); + assert.equal(denied.threadId, SECOND_THREAD_ID); + assert.equal(denied.decision, "decline"); + assert.equal(denied.options.ownerClientId, "owner-b"); + + client.releaseTargetOwnerConfirmation(); + await assert.rejects(switching, /session selection was cancelled because the IPC session closed/); +}); + +test("IPC adapter refuses session switching during an active turn or pending request", async () => { + const state = fixtureState([{ id: 12, method: "item/commandExecution/requestApproval", params: { command: "echo pending" } }]); + const client = new MultiSessionIpcClient(new Map([[THREAD_ID, state], [SECOND_THREAD_ID, fixtureState()]]), { + [THREAD_ID]: "owner-a", + [SECOND_THREAD_ID]: "owner-b", + }); + const adapter = new CodexIpcAgentAdapter({ client, threadId: THREAD_ID, loadCompleteHistory: false, approvalTimeoutMs: 0, followTimeoutMs: 500 }); + await adapter.start(); + await assert.rejects(() => adapter.selectSession({ threadId: SECOND_THREAD_ID }), /turn or approval is active/); + await adapter.dispose(); +}); + +test("IPC session discovery accepts current ULID rollout filenames", async (t) => { + const codexHome = await fs.mkdtemp(path.join(os.tmpdir(), "codex-ipc-ulid-")); + t.after(() => fs.rm(codexHome, { recursive: true, force: true })); + const sessions = path.join(codexHome, "sessions", "2026", "08"); + await fs.mkdir(sessions, { recursive: true }); + await fs.writeFile(path.join(sessions, `rollout-2026-08-29T00-00-00-${ULID_THREAD_ID}.jsonl`), `${JSON.stringify({ + type: "session_meta", + payload: { id: ULID_THREAD_ID, originator: "codex_vscode", source: "vscode", cwd: "/tmp/ulid" }, + })}\n`); + await fs.writeFile(path.join(codexHome, "session_index.jsonl"), `${JSON.stringify({ id: ULID_THREAD_ID, thread_name: "ULID 会话", updated_at: "2026-08-29T12:00:00Z" })}\n`); + const client = new DiscoveryIpcClient(fixtureState()); + const adapter = new CodexIpcAgentAdapter({ + client, + threadId: THREAD_ID, + codexHome, + loadCompleteHistory: false, + approvalTimeoutMs: 0, + followTimeoutMs: 500, + }); + await adapter.start(); + const result = await adapter.listSessions({ limit: 10 }); + const session = result.sessions.find((entry) => entry.threadId === ULID_THREAD_ID); + assert.ok(session); + assert.equal(session.title, "ULID 会话"); + assert.equal(session.available, true); + await adapter.dispose(); +}); diff --git a/aether-vscodex/test/codex-ipc.test.js b/aether-vscodex/test/codex-ipc.test.js new file mode 100644 index 000000000..7bd366b26 --- /dev/null +++ b/aether-vscodex/test/codex-ipc.test.js @@ -0,0 +1,209 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const net = require("node:net"); +const os = require("node:os"); +const path = require("node:path"); +const { mkdtempSync, rmSync } = require("node:fs"); +const test = require("node:test"); + +const { + CODEX_IPC_METHOD_VERSIONS, + CodexIpcClient, + IpcFrameDecoder, + applyIpcPatches, + encodeIpcFrame, +} = require("../vscode-extension/dist/codexIpc.js"); + +function waitFor(predicate, timeoutMs = 2_000) { + const started = Date.now(); + return new Promise((resolve, reject) => { + const poll = () => { + if (predicate()) return resolve(); + if (Date.now() - started >= timeoutMs) return reject(new Error("timed out waiting for fixture")); + setTimeout(poll, 5); + }; + poll(); + }); +} + +test("private IPC framing handles split UTF-8 frames", () => { + const message = { + type: "broadcast", + method: "thread-stream-following-changed", + sourceClientId: "client-1", + version: 1, + params: { conversationId: "thread-1", hostId: "local", following: true, text: "中文" }, + }; + const frame = encodeIpcFrame(message); + const decoder = new IpcFrameDecoder(); + const first = decoder.push(frame.subarray(0, 3)); + assert.deepEqual(first, []); + const second = decoder.push(frame.subarray(3, frame.length - 1)); + assert.deepEqual(second, []); + assert.deepEqual(decoder.push(frame.subarray(frame.length - 1)), [message]); +}); + +test("applyIpcPatches updates a conversation snapshot", () => { + const initial = { turns: [{ items: [{ text: "old" }] }], status: "idle" }; + const next = applyIpcPatches(initial, [ + { op: "replace", path: ["turns", 0, "items", 0, "text"], value: "new" }, + { op: "add", path: ["turns", 0, "items", 1], value: { text: "second" } }, + { op: "replace", path: ["status"], value: "active" }, + ]); + assert.deepEqual(next, { + turns: [{ items: [{ text: "new" }, { text: "second" }] }], + status: "active", + }); +}); + +test("fixture owner receives follow/start/steer/interrupt/approval requests", async () => { + const temp = mkdtempSync(path.join(os.tmpdir(), "codex-ipc-fixture-")); + const socketPath = path.join(temp, "ipc.sock"); + const threadId = "11111111-1111-4111-8111-111111111111"; + const ownerId = "owner-client"; + const requests = []; + const followingBroadcasts = []; + let fixtureSocket; + const server = net.createServer((socket) => { + fixtureSocket = socket; + const decoder = new IpcFrameDecoder(); + socket.on("data", (chunk) => { + for (const message of decoder.push(chunk)) { + if (message.type === "request" && message.method === "initialize") { + socket.write(encodeIpcFrame({ + type: "response", + requestId: message.requestId, + resultType: "success", + method: "initialize", + handledByClientId: "fixture-client", + result: { clientId: "fixture-client" }, + })); + continue; + } + if (message.type === "broadcast" && message.method === "thread-stream-following-changed") { + followingBroadcasts.push(message); + const target = message.sourceClientId; + socket.write(encodeIpcFrame({ + type: "broadcast", + method: "thread-stream-state-changed", + sourceClientId: ownerId, + targetClientIds: [target], + version: CODEX_IPC_METHOD_VERSIONS["thread-stream-state-changed"], + params: { + conversationId: threadId, + hostId: "local", + change: { + type: "snapshot", + revision: 1, + conversationState: { id: threadId, title: "fixture", turns: [], requests: [] }, + }, + }, + })); + continue; + } + if (message.type === "request") { + requests.push(message); + socket.write(encodeIpcFrame({ + type: "response", + requestId: message.requestId, + resultType: "success", + method: message.method, + handledByClientId: ownerId, + result: { method: message.method, ok: true }, + })); + } + } + }); + }); + + try { + await new Promise((resolve, reject) => { + server.once("error", reject); + server.listen(socketPath, resolve); + }); + const client = new CodexIpcClient({ socketPath, autoReconnect: false }); + const streamEvents = []; + client.onStreamEvent((event) => streamEvents.push(event)); + await client.connect(); + await client.followConversation(threadId); + await waitFor(() => streamEvents.some((event) => event.kind === "snapshot")); + assert.equal(client.getConversationState(threadId).ownerClientId, ownerId); + fixtureSocket.write(encodeIpcFrame({ + type: "broadcast", + method: "thread-stream-following-status-requested", + sourceClientId: ownerId, + targetClientIds: ["fixture-client"], + version: CODEX_IPC_METHOD_VERSIONS["thread-stream-following-status-requested"], + params: { conversationId: threadId, hostId: "local" }, + })); + await waitFor(() => followingBroadcasts.length >= 2); + assert.deepEqual(followingBroadcasts[1].targetClientIds, [ownerId]); + assert.deepEqual(followingBroadcasts[1].params, { + conversationId: threadId, + hostId: "local", + following: true, + }); + + await client.startTurn(threadId, "hello", { ownerClientId: ownerId }); + await client.steerTurn(threadId, "follow-up", { ownerClientId: ownerId }); + await client.updateThreadSettings(threadId, { + model: "gpt-5.6-sol", + effort: "ultra", + multiAgentMode: "explicitRequestOnly", + }, { ownerClientId: ownerId }); + await client.interruptTurn(threadId, { mode: "user-stop", expectedTurnId: "turn-1", ownerClientId: ownerId }); + await client.respondCommandApproval(threadId, 7, "decline", { ownerClientId: ownerId }); + await client.respondFileApproval(threadId, "8", "cancel", { ownerClientId: ownerId }); + await client.respondPermissionsApproval(threadId, 9, { permissions: {}, scope: "turn" }, { ownerClientId: ownerId }); + await client.respondUserInput(threadId, 10, { answers: {} }, { ownerClientId: ownerId }); + await client.respondMcpElicitation(threadId, 11, { action: "decline", content: null, _meta: null }, { ownerClientId: ownerId }); + + assert.deepEqual(requests.map((request) => request.method), [ + "thread-follower-start-turn", + "thread-follower-steer-turn", + "thread-follower-update-thread-settings", + "thread-follower-interrupt-turn", + "thread-follower-command-approval-decision", + "thread-follower-file-approval-decision", + "thread-follower-permissions-request-approval-response", + "thread-follower-submit-user-input", + "thread-follower-submit-mcp-server-elicitation-response", + ]); + assert.deepEqual(requests[0].params, { + conversationId: threadId, + turnStart: { + request: { + threadId, + input: [{ type: "text", text: "hello", text_elements: [] }], + }, + context: { inheritThreadSettings: true }, + }, + }); + assert.deepEqual(requests[2].params, { + conversationId: threadId, + threadSettings: { + model: "gpt-5.6-sol", + effort: "ultra", + multiAgentMode: "explicitRequestOnly", + }, + }); + assert.equal(requests[2].version, 1); + assert.deepEqual(requests[3].params, { + conversationId: threadId, + mode: "user-stop", + expectedTurnId: "turn-1", + }); + assert.equal(requests[3].version, 4); + assert.deepEqual(requests[4].params, { conversationId: threadId, requestId: 7, decision: "decline" }); + assert.deepEqual(requests[8].params, { + conversationId: threadId, + requestId: 11, + response: { action: "decline", content: null, _meta: null }, + }); + await client.dispose(); + } finally { + await new Promise((resolve) => server.close(resolve)); + rmSync(temp, { recursive: true, force: true }); + } +}); diff --git a/aether-vscodex/test/codex-path.test.js b/aether-vscodex/test/codex-path.test.js new file mode 100644 index 000000000..ed76ad1f6 --- /dev/null +++ b/aether-vscodex/test/codex-path.test.js @@ -0,0 +1,57 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const { chmodSync, mkdtempSync, mkdirSync, rmSync, writeFileSync } = require("node:fs"); +const os = require("node:os"); +const path = require("node:path"); +const test = require("node:test"); + +const { resolveCodexCommand } = require("../vscode-extension/dist/codexPath.js"); +const { JsonlRpcClient } = require("../vscode-extension/dist/jsonlRpc.js"); + +function temporaryDirectory() { + return mkdtempSync(path.join(os.tmpdir(), "codex-remote-path-")); +} + +test("resolveCodexCommand finds a bare command in PATH", () => { + const root = temporaryDirectory(); + try { + const bin = path.join(root, "bin"); + const executable = path.join(bin, "codex-test"); + mkdirSync(bin); + writeFileSync(executable, "#!/bin/sh\nexit 0\n"); + chmodSync(executable, 0o755); + assert.equal(resolveCodexCommand("codex-test", { env: { PATH: bin }, platform: process.platform }), executable); + } finally { + rmSync(root, { recursive: true, force: true }); + } +}); + +test("resolveCodexCommand falls back to a per-user ChatGPT.app install", () => { + const root = temporaryDirectory(); + try { + const executable = path.join(root, "Applications", "ChatGPT.app", "Contents", "Resources", "codex"); + mkdirSync(path.dirname(executable), { recursive: true }); + writeFileSync(executable, "#!/bin/sh\nexit 0\n"); + chmodSync(executable, 0o755); + assert.equal( + resolveCodexCommand("codex", { env: { PATH: "/usr/bin:/bin" }, homeDir: root, platform: "darwin" }), + executable, + ); + } finally { + rmSync(root, { recursive: true, force: true }); + } +}); + +test("a missing explicit command reports a full-path setting hint", () => { + assert.throws( + () => resolveCodexCommand("/definitely/missing/codex", { platform: process.platform }), + /Codex executable .* was not found.*codexRemoteCollab\.codexCommand.*full path/, + ); +}); + +test("JsonlRpcClient turns spawn ENOENT into an actionable error", async () => { + const client = new JsonlRpcClient({ command: "/definitely/missing/codex", args: [] }); + await assert.rejects(() => client.start(), /Codex executable .* was not found.*codexRemoteCollab\.codexCommand/); + client.close(); +}); diff --git a/aether-vscodex/test/composite-relay.test.js b/aether-vscodex/test/composite-relay.test.js new file mode 100644 index 000000000..323c5a879 --- /dev/null +++ b/aether-vscodex/test/composite-relay.test.js @@ -0,0 +1,105 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const test = require("node:test"); + +const { CompositeRelayTransport } = require("../vscode-extension/dist/compositeRelay.js"); + +class FakeRelay { + constructor({ connectError } = {}) { + this.connectError = connectError; + this.frames = []; + this.closed = false; + this.listeners = { message: new Set(), open: new Set(), close: new Set() }; + } + + async connect() { + if (this.connectError) throw this.connectError; + for (const listener of this.listeners.open) listener(); + } + + send(frame) { this.frames.push(frame); } + close() { this.closed = true; } + onMessage(listener) { return this.add("message", listener); } + onOpen(listener) { return this.add("open", listener); } + onClose(listener) { return this.add("close", listener); } + add(type, listener) { + this.listeners[type].add(listener); + return { dispose: () => this.listeners[type].delete(listener) }; + } + receive(frame) { for (const listener of this.listeners.message) listener(frame); } + disconnect(error) { for (const listener of this.listeners.close) listener(error); } +} + +test("CompositeRelayTransport keeps local control available when optional cloud connect fails", async () => { + const local = new FakeRelay(); + const cloud = new FakeRelay({ connectError: new Error("cloud offline") }); + const relay = new CompositeRelayTransport([ + { id: "local", transport: local, required: true }, + { id: "cloud", transport: cloud }, + ]); + + await relay.connect(); + assert.equal(relay.isConnected("local"), true); + assert.equal(relay.isConnected("cloud"), false); + relay.send({ kind: "event", type: "session.snapshot" }); + assert.equal(local.frames.length, 1); + assert.equal(cloud.frames.length, 1, "optional transport may queue events for reconnect"); + relay.close(); + assert.equal(local.closed, true); + assert.equal(cloud.closed, true); +}); + +test("CompositeRelayTransport forwards commands and reports offline only after every relay closes", async () => { + const local = new FakeRelay(); + const cloud = new FakeRelay(); + const relay = new CompositeRelayTransport([ + { id: "local", transport: local, required: true }, + { id: "cloud", transport: cloud }, + ]); + const messages = []; + const closes = []; + relay.onMessage((frame) => messages.push(frame)); + relay.onClose((error) => closes.push(error?.message)); + + await relay.connect(); + assert.equal(relay.isConnected("local"), true); + assert.equal(relay.isConnected("cloud"), true); + cloud.receive({ kind: "command", type: "turn.start" }); + assert.equal(messages.length, 1); + local.disconnect(new Error("local offline")); + assert.equal(relay.isConnected("local"), false); + assert.deepEqual(closes, []); + cloud.disconnect(new Error("cloud offline")); + assert.deepEqual(closes, ["cloud offline"]); + relay.close(); +}); + +test("CompositeRelayTransport surfaces each member reconnect for snapshot hydration", async () => { + const local = new FakeRelay(); + const cloud = new FakeRelay(); + const relay = new CompositeRelayTransport([ + { id: "local", transport: local, required: true }, + { id: "cloud", transport: cloud }, + ]); + let opens = 0; + relay.onOpen(() => { opens += 1; }); + await relay.connect(); + assert.equal(opens, 2); + cloud.disconnect(new Error("cloud offline")); + for (const listener of cloud.listeners.open) listener(); + assert.equal(opens, 3, "cloud recovery must prompt RelayHost to publish a fresh snapshot"); + relay.close(); +}); + +test("CompositeRelayTransport fails when the required local relay cannot connect", async () => { + const local = new FakeRelay({ connectError: new Error("local offline") }); + const cloud = new FakeRelay(); + const relay = new CompositeRelayTransport([ + { id: "local", transport: local, required: true }, + { id: "cloud", transport: cloud }, + ]); + await assert.rejects(relay.connect(), /local: local offline/); + assert.equal(local.closed, true); + assert.equal(cloud.closed, true); +}); diff --git a/aether-vscodex/test/extension-i18n-build.test.js b/aether-vscodex/test/extension-i18n-build.test.js new file mode 100644 index 000000000..ff4969dbf --- /dev/null +++ b/aether-vscodex/test/extension-i18n-build.test.js @@ -0,0 +1,34 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const fs = require("node:fs"); +const path = require("node:path"); +const test = require("node:test"); + +test("VS Code runtime strings have English and Simplified Chinese bundles", () => { + const extensionRoot = path.join(__dirname, "..", "vscode-extension"); + const source = fs.readFileSync(path.join(extensionRoot, "src", "extension.ts"), "utf8"); + const manifest = JSON.parse(fs.readFileSync(path.join(extensionRoot, "package.json"), "utf8")); + const english = JSON.parse(fs.readFileSync(path.join(extensionRoot, "l10n", "bundle.l10n.json"), "utf8")); + const chinese = JSON.parse(fs.readFileSync(path.join(extensionRoot, "l10n", "bundle.l10n.zh-cn.json"), "utf8")); + const keys = [...source.matchAll(/(? match[1]); + + assert.equal(manifest.l10n, "./l10n"); + assert.ok(keys.length > 20, "expected runtime-localized extension strings"); + for (const key of new Set(keys)) { + assert.equal(english[key], key, `missing English source string: ${key}`); + assert.equal(typeof chinese[key], "string", `missing zh-CN translation: ${key}`); + assert.ok(chinese[key].length > 0, `empty zh-CN translation: ${key}`); + } +}); + +test("production copy scripts require the Vue build instead of silently falling back", () => { + const projectRoot = path.join(__dirname, ".."); + const extensionSync = fs.readFileSync(path.join(projectRoot, "vscode-extension", "scripts", "sync-local-relay.cjs"), "utf8"); + const aetherSync = fs.readFileSync(path.join(projectRoot, "..", "frontend", "scripts", "sync-vscodex.mjs"), "utf8"); + const extensionManifest = JSON.parse(fs.readFileSync(path.join(projectRoot, "vscode-extension", "package.json"), "utf8")); + + assert.match(extensionManifest.scripts["vscode:prepublish"], /build:web/); + assert.doesNotMatch(extensionSync, /projectRoot,\s*"public"/); + assert.doesNotMatch(aetherSync, /moduleRoot,\s*'public'/); +}); diff --git a/aether-vscodex/test/local-relay.test.js b/aether-vscodex/test/local-relay.test.js new file mode 100644 index 000000000..58ef1a7c4 --- /dev/null +++ b/aether-vscodex/test/local-relay.test.js @@ -0,0 +1,120 @@ +const assert = require("node:assert/strict"); +const http = require("node:http"); +const path = require("node:path"); +const test = require("node:test"); + +const { + LocalRelayController, + localRelayTarget, + relayHealthAvailable, +} = require("../vscode-extension/dist/localRelay.js"); + +test("local relay target accepts only loopback ws URLs", () => { + assert.deepEqual(localRelayTarget("ws://localhost:8898/v1/connect"), { + host: "127.0.0.1", + port: 8898, + healthUrl: "http://127.0.0.1:8898/api/health", + webUrl: "http://127.0.0.1:8898/", + }); + assert.equal(localRelayTarget("wss://127.0.0.1:8898/v1/connect"), undefined); + assert.equal(localRelayTarget("ws://192.168.1.10:8898/v1/connect"), undefined); + assert.equal(localRelayTarget("not a url"), undefined); +}); + +test("local relay health probe recognizes a responding HTTP service", async (t) => { + const server = http.createServer((request, response) => { + if (request.url === "/api/health") { + response.writeHead(200, { "content-type": "application/json" }).end(JSON.stringify({ ok: true })); + } else if (request.url === "/aborted") { + response.writeHead(200, { "content-type": "application/json" }); + response.write('{"ok":'); + response.destroy(); + } else if (request.url === "/drip") { + response.writeHead(200, { "content-type": "application/json" }); + const interval = setInterval(() => response.write(" "), 10); + response.on("close", () => clearInterval(interval)); + } else { + response.writeHead(404).end(); + } + }); + await new Promise((resolve) => server.listen(0, "127.0.0.1", resolve)); + t.after(() => new Promise((resolve) => server.close(resolve))); + const address = server.address(); + assert.equal(await relayHealthAvailable(`http://127.0.0.1:${address.port}/api/health`), true); + assert.equal(await relayHealthAvailable(`http://127.0.0.1:${address.port}/missing`), false); + assert.equal(await relayHealthAvailable(`http://127.0.0.1:${address.port}/aborted`, 100), false); + const startedAt = Date.now(); + assert.equal(await relayHealthAvailable(`http://127.0.0.1:${address.port}/drip`, 50), false); + assert.ok(Date.now() - startedAt < 500); +}); + +test("local relay controller starts and stops a bundled loopback relay", async () => { + let starts = 0; + let stops = 0; + class FakeRelay { + async start() { starts += 1; return { host: "127.0.0.1", port: 65534 }; } + async stop() { stops += 1; } + } + const controller = new LocalRelayController({ + extensionPath: path.resolve(__dirname, "../vscode-extension"), + probeTimeoutMs: 20, + loadRelayModule: () => ({ CodexRelay: FakeRelay }), + }); + assert.equal(await controller.ensureRunning("ws://127.0.0.1:65534/v1/connect"), true); + assert.equal(starts, 1); + await controller.stop(); + assert.equal(stops, 1); +}); + +test("local relay controller does not leak a relay when stopped during startup", async () => { + let releaseStart; + const startGate = new Promise((resolve) => { releaseStart = resolve; }); + let startEntered; + const entered = new Promise((resolve) => { startEntered = resolve; }); + let stops = 0; + class SlowRelay { + async start() { + startEntered(); + await startGate; + return { host: "127.0.0.1", port: 65533 }; + } + async stop() { stops += 1; } + } + const controller = new LocalRelayController({ + extensionPath: path.resolve(__dirname, "../vscode-extension"), + probeTimeoutMs: 20, + loadRelayModule: () => ({ CodexRelay: SlowRelay }), + }); + const starting = controller.ensureRunning("ws://127.0.0.1:65533/v1/connect"); + await entered; + const stopping = controller.stop(); + releaseStart(); + await Promise.all([starting, stopping]); + assert.equal(stops, 1); +}); + +test("local relay controller does not start after stop wins an in-flight health probe", async () => { + let resolveProbe; + const probe = new Promise((resolve) => { resolveProbe = resolve; }); + let probeEntered; + const entered = new Promise((resolve) => { probeEntered = resolve; }); + let starts = 0; + class FakeRelay { + async start() { starts += 1; return { host: "127.0.0.1", port: 65532 }; } + async stop() {} + } + const controller = new LocalRelayController({ + extensionPath: path.resolve(__dirname, "../vscode-extension"), + loadRelayModule: () => ({ CodexRelay: FakeRelay }), + probeRelayHealth: async () => { + probeEntered(); + return probe; + }, + }); + const ensuring = controller.ensureRunning("ws://127.0.0.1:65532/v1/connect"); + await entered; + await controller.stop(); + resolveProbe(false); + assert.equal(await ensuring, false); + assert.equal(starts, 0); +}); diff --git a/aether-vscodex/test/public-embed-i18n.test.js b/aether-vscodex/test/public-embed-i18n.test.js new file mode 100644 index 000000000..6cffda015 --- /dev/null +++ b/aether-vscodex/test/public-embed-i18n.test.js @@ -0,0 +1,194 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const fs = require("node:fs"); +const path = require("node:path"); +const test = require("node:test"); + +const { createAetherEmbedBridge, isAetherEmbed } = require("../public/embed-bridge.js"); +const i18n = require("../public/i18n.js"); + +function embeddedWindow() { + const listeners = new Map(); + const posts = []; + const parent = { postMessage: (message, origin) => posts.push({ message, origin }) }; + const bodyClasses = new Set(); + const documentElement = { dataset: {}, style: {} }; + const windowLike = { + location: { search: "?embed=aether", origin: "https://aether.example" }, + parent, + document: { + body: { classList: { add: (value) => bodyClasses.add(value) } }, + documentElement, + }, + addEventListener: (name, listener) => listeners.set(name, listener), + removeEventListener: (name, listener) => { + if (listeners.get(name) === listener) listeners.delete(name); + }, + }; + return { bodyClasses, documentElement, listeners, parent, posts, windowLike }; +} + +test("Aether embed mode is opt-in and announces readiness only to the same-origin parent", () => { + assert.equal(isAetherEmbed({ search: "" }), false); + assert.equal(isAetherEmbed({ search: "?embed=other" }), false); + assert.equal(isAetherEmbed({ search: "?embed=aether" }), true); + + const fixture = embeddedWindow(); + const bridge = createAetherEmbedBridge(fixture.windowLike); + assert.equal(bridge.active, true); + bridge.start(); + assert.equal(fixture.bodyClasses.has("embed-aether"), true); + assert.deepEqual(fixture.posts, [{ + message: { v: 1, type: "aether-vscodex/ready" }, + origin: "https://aether.example", + }]); +}); + +test("Aether embed bridge rejects cross-origin and non-parent messages and buffers an early connect", () => { + const fixture = embeddedWindow(); + const bridge = createAetherEmbedBridge(fixture.windowLike); + bridge.start(); + const dispatch = fixture.listeners.get("message"); + const connect = { + v: 1, + type: "aether-vscodex/connect", + ticket: "one-time-ticket", + wsUrl: "/api/vscodex/ws", + locale: "en-US", + theme: "dark", + }; + + dispatch({ origin: "https://attacker.example", source: fixture.parent, data: connect }); + dispatch({ origin: "https://aether.example", source: {}, data: connect }); + let received = null; + bridge.on("connect", (message) => { received = message; }); + assert.equal(received, null); + + dispatch({ origin: "https://aether.example", source: fixture.parent, data: connect }); + assert.equal(received.ticket, "one-time-ticket"); + assert.equal(fixture.documentElement.dataset.theme, "dark"); + + const second = embeddedWindow(); + const bufferedBridge = createAetherEmbedBridge(second.windowLike); + bufferedBridge.start(); + second.listeners.get("message")({ origin: "https://aether.example", source: second.parent, data: connect }); + let buffered = null; + bufferedBridge.on("connect", (message) => { buffered = message; }); + assert.equal(buffered.ticket, "one-time-ticket"); +}); + +test("bridge ticket requests never place the ticket in a URL", () => { + const fixture = embeddedWindow(); + const bridge = createAetherEmbedBridge(fixture.windowLike); + bridge.start(); + bridge.requestTicket({ reason: "disconnected", deviceId: "device-1" }); + assert.deepEqual(fixture.posts.at(-1), { + message: { + v: 1, + type: "aether-vscodex/request-ticket", + reason: "disconnected", + deviceId: "device-1", + }, + origin: "https://aether.example", + }); +}); + +test("locale dictionary covers static shell and core dynamic status text", () => { + assert.equal(i18n.translate("设置", "en-US"), "Settings"); + assert.equal(i18n.translate("中文", "en-US"), "Chinese"); + assert.equal(i18n.translate("正在思考", "en-US"), "Thinking"); + assert.equal(i18n.translate("已读取这些内容 · 4 个文件", "en-US"), "Read these items · 4 files"); + assert.equal(i18n.translate("用时 3分45秒", "en-US"), "Worked for 3m45s"); + assert.equal(i18n.translate("修改权限,当前为需要时询问", "en-US"), "Change permissions. Current: Ask when needed"); + assert.equal(i18n.translate("模型设置更新失败:timeout", "en-US"), "Unable to update model settings: timeout"); + assert.equal(i18n.translate("请求 #17 已发送,等待 VS Code 主机确认", "en-US"), "Request #17 sent; waiting for the VS Code host"); + assert.equal(i18n.translate("无法读取 notes.md", "en-US"), "Unable to read notes.md"); + assert.equal(i18n.translate("命令: timed out", "en-US"), "Command: timed out"); + assert.equal(i18n.translate("命令: timed out(执行状态未知,请等待主机恢复)", "en-US"), "Command: timed out (execution status unknown; wait for the host to recover)"); + assert.equal(i18n.translate("子代理 失败", "en-US"), "Subagent failed"); + assert.equal(i18n.translate("已在 2秒 内运行 echo hi", "en-US"), "Ran echo hi in 2s"); + assert.equal(i18n.translate("命令运行失败 · echo hi · 2秒", "en-US"), "Command failed · echo hi · 2s"); + assert.equal(i18n.translate("命令运行失败 · echo hi", "en-US"), "Command failed · echo hi"); + assert.equal(i18n.translate("已停止 echo hi · 2秒", "en-US"), "Stopped echo hi · 2s"); + assert.equal(i18n.translate("文件变更 · 失败", "en-US"), "File changes · Failed"); + assert.equal(i18n.translate("文件变更 · 已中断", "en-US"), "File changes · Interrupted"); + assert.equal(i18n.translate("命令 · echo hi", "en-US"), "Command · echo hi"); + assert.equal(i18n.translate("命令 · 设置", "en-US"), "Command · 设置"); + assert.equal(i18n.translate("正在读取 设置", "en-US"), "Reading 设置"); + assert.equal(i18n.translate("已在 2秒 内运行 设置", "en-US"), "Ran 设置 in 2s"); + assert.equal(i18n.translate("正在切换到「设置」…", "en-US"), "Switching to “设置”..."); + assert.equal(i18n.translate("你停止了工作", "en-US"), "You stopped working"); + assert.equal(i18n.translate("工具失败", "en-US"), "Tool failed"); + assert.equal(i18n.translate("正在搜索", "en-US"), "Searching"); + assert.equal(i18n.translate("已工具 · 2秒", "en-US"), "Tool completed · 2s"); + assert.equal(i18n.translate("当前模型 5.6 Sol 标准,切换模型", "en-US"), "Current model: 5.6 Sol Medium. Change model"); + assert.equal(i18n.translate("编辑了文件", "en-US"), "Edited files"); + assert.equal(i18n.translate("编辑了文件 · 2秒", "en-US"), "Edited files · 2s"); + assert.equal(i18n.translate("已完成计划", "en-US"), "Completed plan"); + assert.equal(i18n.translate("已完成计划 · 2秒", "en-US"), "Completed plan · 2s"); + assert.equal(i18n.translate("…(文件已截断)", "en-US"), "... (file truncated)"); + assert.equal(i18n.translate("事件窗口已过期,请以当前快照为准", "en-US"), "The event window expired; the current snapshot is authoritative"); + assert.equal(i18n.translate("控制模式", "en-US"), "Control mode"); + assert.equal(i18n.translate("同步模式跟随 VS Code 当前会话", "en-US"), "Sync mode follows the current VS Code conversation"); + assert.equal(i18n.translate("异步模式可独立管理会话", "en-US"), "Async mode manages conversations independently"); + assert.equal(i18n.translate("当前任务或请求完成后才能切换控制模式", "en-US"), "The control mode can be changed after the current task or request finishes"); + assert.equal(i18n.translate("Settings", "zh-CN"), "设置"); + assert.equal(i18n.normalizeLocale("zh-Hans"), "zh-CN"); + assert.equal(i18n.normalizeLocale("en-GB"), "en-US"); +}); + +test("renderer-owned dynamic labels have English fallbacks without translating host values", () => { + assert.equal(i18n.translate("命令 · echo hi", "en-US"), "Command · echo hi"); + assert.equal(i18n.translate("你停止了工作", "en-US"), "You stopped working"); + assert.equal(i18n.translate("工具失败", "en-US"), "Tool failed"); + assert.equal(i18n.translate("正在搜索", "en-US"), "Searching"); + assert.equal(i18n.translate("当前模型 5.6 Sol 标准,切换模型", "en-US"), "Current model: 5.6 Sol Medium. Change model"); + + const app = fs.readFileSync(path.join(__dirname, "..", "public", "app.js"), "utf8"); + // Command/path/title values are appended after a locale-specific prefix; + // they are never passed through the translator as a whole. + assert.match(app, /uiWithRaw\("正在运行 ", "Running ",/); + assert.match(app, /uiWithRaw\("已读取 ", "Read ",/); + assert.match(app, /uiLocale\(\) === "en-US" \? `Switching to/); +}); + +test("public shell uses relative assets and embedded startup skips the health probe", () => { + const publicRoot = path.join(__dirname, "..", "public"); + const html = fs.readFileSync(path.join(publicRoot, "index.html"), "utf8"); + const app = fs.readFileSync(path.join(publicRoot, "app.js"), "utf8"); + assert.match(html, /href="\.\/style\.css"/); + assert.match(html, /src="\.\/embed-bridge\.js"/); + assert.match(html, /src="\.\/i18n\.js"/); + assert.match(html, /src="\.\/app\.js"/); + assert.match(app, /if \(embeddedInAether\)[\s\S]+else \{[\s\S]+fetch\("\.\/api\/health"/); + assert.doesNotMatch(app, /ticket=.*state\.embedTicket/); + assert.match(app, /empty\.textContent = t\(activity\.status === "inProgress" \? "正在读取文件" : "读取完成"\)/); + assert.match(app, /outputContent\.textContent = t\("无输出"\)/); + assert.match(app, /button\.title = t\(title\)/); + assert.match(app, /activity\.action === "spawnAgent" \? t\("启动子代理"\)/); + assert.match(app, /return t\("需要远程确认或输入"\)/); + assert.match(app, /questionPrompt === undefined \|\| questionPrompt === null \? t\("请输入"\)/); + assert.match(app, /checkbox\.setAttribute\("aria-label", t\(checkbox\.checked \? "已完成" : "未完成"\)\)/); + assert.match(app, /window\.addEventListener\("aether-vscodex:locale", \(\) => \{[\s\S]+state\.activities\.values\(\)[\s\S]+renderRequests\(\)/); +}); + +test("control mode is snapshot-authoritative and gates independent session actions", () => { + const publicRoot = path.join(__dirname, "..", "public"); + const html = fs.readFileSync(path.join(publicRoot, "index.html"), "utf8"); + const app = fs.readFileSync(path.join(publicRoot, "app.js"), "utf8"); + + assert.match(html, /id="controlModeSwitch"[\s\S]+data-control-mode="sync"[\s\S]+data-control-mode="async"/); + assert.match(app, /command\("control\/mode\/set", \{ mode \}\)/); + assert.match(app, /applyControlModeSnapshot\(payload\.metadata\)/); + assert.match(app, /const controlMetadata = \{[\s\S]+snapshot\.metadata[\s\S]+appState\.sessionMetadata[\s\S]+applyControlModeSnapshot\(controlMetadata\)/); + assert.match(app, /sessionList: source\.sessionList === true/); + assert.match(app, /Boolean\(state\.sessionListCommandId\)/); + assert.match(app, /mode_switch_pending.*return "正在切换控制模式"/); + assert.match(app, /mode_busy\|cannot switch control mode.*return "当前任务或请求完成后才能切换控制模式"/); + assert.match(app, /setConversationStatus\(sessionErrorMessage\(message, "控制模式切换失败"\), "warning"\)/); + assert.match(app, /if \(!sessionControlAllowed\("sessionList"\)\) return;/); + assert.match(app, /if \(!sessionControlAllowed\("sessionSelect"\)\)/); + assert.match(app, /if \(!sessionControlAllowed\("sessionCreate"\)\)/); + assert.match(app, /sessionPickerButton\.disabled = !listAllowed/); +}); diff --git a/aether-vscodex/test/relay-client.test.js b/aether-vscodex/test/relay-client.test.js new file mode 100644 index 000000000..f2f2adcbe --- /dev/null +++ b/aether-vscodex/test/relay-client.test.js @@ -0,0 +1,206 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const { EventEmitter } = require("node:events"); +const test = require("node:test"); + +const { RelayClient } = require("../vscode-extension/dist/relayClient.js"); + +class FakeWebSocket extends EventEmitter { + static instances = []; + + constructor(url) { + super(); + this.url = url; + this.readyState = 0; + this.sent = []; + FakeWebSocket.instances.push(this); + } + + open() { + this.readyState = 1; + this.emit("open"); + } + + receive(frame) { + this.emit("message", Buffer.from(JSON.stringify(frame))); + } + + send(data) { + this.sent.push(JSON.parse(data)); + } + + close() { + if (this.readyState === 3) return; + this.readyState = 3; + this.emit("close"); + } +} + +test("RelayClient queues application frames until auth.ok on initial connect and reconnect", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://relay.invalid/v1/connect", + accessToken: "host-token", + reconnect: false, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const firstConnect = client.connect(); + const first = FakeWebSocket.instances[0]; + first.open(); + assert.deepEqual(first.sent.map((frame) => frame.kind), ["hello", "auth"]); + + client.send({ v: 1, kind: "event", type: "output.chunk", id: "event-1", sessionId: "session-1", payload: { text: "queued" } }); + assert.equal(first.sent.length, 2, "application event must not be sent before authentication"); + first.receive({ type: "auth.ok", role: "host", clientType: "host" }); + await firstConnect; + assert.equal(first.sent.length, 3); + assert.equal(first.sent[2].id, "event-1"); + + first.close(); + const secondConnect = client.connect(); + const second = FakeWebSocket.instances[1]; + second.open(); + assert.deepEqual(second.sent.map((frame) => frame.kind), ["hello", "auth"]); + + client.send({ v: 1, kind: "event", type: "output.chunk", id: "event-2", sessionId: "session-1", payload: { text: "queued during reconnect" } }); + assert.equal(second.sent.length, 2, "reconnect window must remain auth-gated"); + second.receive({ type: "auth.ok", role: "host", clientType: "host" }); + await secondConnect; + assert.equal(second.sent.length, 3); + assert.equal(second.sent[2].id, "event-2"); +}); + +test("RelayClient coalesces queued transcript projections within a byte budget", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://relay.invalid/v1/connect", + accessToken: "host-token", + reconnect: false, + maxFrameBytes: 4_096, + maxQueuedBytes: 4_096, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const connecting = client.connect(); + const socket = FakeWebSocket.instances[0]; + socket.open(); + client.send({ v: 1, kind: "event", type: "approval.requested", id: "approval", sessionId: "session-1", payload: { text: "a".repeat(700) } }); + client.send({ v: 1, kind: "event", type: "output.snapshot", id: "old-projection", sessionId: "session-1", payload: { text: "x".repeat(1_200) } }); + client.send({ v: 1, kind: "event", type: "output.chunk", id: "new-projection", sessionId: "session-1", payload: { text: "y".repeat(1_200) } }); + client.send({ v: 1, kind: "event", type: "command.result", id: "command", sessionId: "session-1", payload: { text: "c".repeat(700) } }); + + assert.ok(client.queueBytes <= 4_096); + socket.receive({ type: "auth.ok", role: "host", clientType: "host" }); + await connecting; + const queuedIds = socket.sent.slice(2).map((frame) => frame.id); + assert.deepEqual(queuedIds, ["approval", "new-projection", "command"]); +}); + +test("RelayClient evicts reconstructible projections before queued control events", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://relay.invalid/v1/connect", + accessToken: "host-token", + reconnect: false, + maxFrameBytes: 4_096, + maxQueuedBytes: 2_500, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const connecting = client.connect(); + const socket = FakeWebSocket.instances[0]; + socket.open(); + client.send({ v: 1, kind: "event", type: "approval.requested", id: "approval", sessionId: "session-1", payload: { text: "a".repeat(850) } }); + client.send({ v: 1, kind: "event", type: "output.chunk", id: "projection", sessionId: "session-1", payload: { text: "x".repeat(900) } }); + client.send({ v: 1, kind: "event", type: "command.result", id: "command", sessionId: "session-1", payload: { text: "c".repeat(850) } }); + + assert.ok(client.queueBytes <= 2_500); + socket.receive({ type: "auth.ok", role: "host", clientType: "host" }); + await connecting; + const queuedIds = socket.sent.slice(2).map((frame) => frame.id); + assert.deepEqual(queuedIds, ["approval", "command"]); +}); + +test("RelayClient supports a tokenless local handshake", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://127.0.0.1:8787/v1/connect", + reconnect: false, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const connecting = client.connect(); + const socket = FakeWebSocket.instances[0]; + socket.open(); + assert.deepEqual(socket.sent.map((frame) => frame.kind), ["hello"]); + socket.receive({ type: "auth.ok", role: "host", clientType: "host", authRequired: false }); + await connecting; + + client.send({ v: 1, kind: "event", type: "connection.opened", id: "event-local", sessionId: "session-local", payload: {} }); + assert.equal(socket.sent.length, 2); + assert.equal(socket.sent[1].type, "connection.opened"); +}); + +test("RelayClient accepts structured history snapshots larger than the old 256 KiB limit", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://127.0.0.1:8787/v1/connect", + reconnect: false, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const connecting = client.connect(); + const socket = FakeWebSocket.instances[0]; + socket.open(); + socket.receive({ type: "auth.ok", role: "host", clientType: "host", authRequired: false }); + await connecting; + + const historyText = "x".repeat(512 * 1024); + assert.doesNotThrow(() => client.send({ + v: 1, + kind: "event", + type: "session.snapshot", + id: "large-history-snapshot", + sessionId: "session-local", + payload: { threadId: "large-thread", messages: [{ kind: "assistant", text: historyText }] }, + })); + assert.equal(socket.sent.at(-1).payload.messages[0].text.length, historyText.length); +}); + +test("RelayClient ignores late events from a replaced socket", async (t) => { + FakeWebSocket.instances.length = 0; + const client = new RelayClient({ + url: "ws://relay.invalid/v1/connect", + accessToken: "host-token", + reconnect: false, + webSocket: FakeWebSocket, + }); + t.after(() => client.close()); + + const firstConnect = client.connect(); + const first = FakeWebSocket.instances[0]; + first.open(); + client.close(); + + const secondConnect = client.connect(); + const second = FakeWebSocket.instances[1]; + second.open(); + + // Simulate a delayed event from the old socket after the replacement. + first.open(); + first.receive({ type: "auth.ok", role: "host", clientType: "host" }); + assert.equal(second.sent.length, 2, "late auth must not authenticate or flush the new socket"); + + client.send({ v: 1, kind: "event", type: "output.chunk", id: "event-after-replace", sessionId: "session-1", payload: { text: "queued" } }); + second.receive({ type: "auth.ok", role: "host", clientType: "host" }); + await secondConnect; + assert.equal(second.sent[2].id, "event-after-replace"); + await assert.rejects(firstConnect); +}); diff --git a/aether-vscodex/test/relay.test.js b/aether-vscodex/test/relay.test.js new file mode 100644 index 000000000..c248ef9ae --- /dev/null +++ b/aether-vscodex/test/relay.test.js @@ -0,0 +1,1610 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const http = require("node:http"); +const path = require("node:path"); +const test = require("node:test"); +const { WebSocket } = require("ws"); +const { CodexRelay } = require("../relay/server.js"); + +const fakeServer = path.join(__dirname, "..", "fixtures", "fake-app-server.cjs"); + +function waitFor(predicate, timeout = 5_000) { + const started = Date.now(); + return new Promise((resolve, reject) => { + const tick = () => { + try { + const result = predicate(); + if (result) return resolve(result); + } catch (error) { return reject(error); } + if (Date.now() - started > timeout) return reject(new Error("timed out waiting for condition")); + setTimeout(tick, 20); + }; + tick(); + }); +} + +function request(base, token, pathname, options = {}) { + return fetch(`${base}${pathname}`, { + ...options, + headers: { Authorization: `Bearer ${token}`, ...(options.headers || {}) }, + }); +} + +function connectWs(base, token) { + const ws = new WebSocket(base.replace(/^http/, "ws") + "/ws"); + const messages = []; + const waiters = []; + ws.on("message", (data) => { + const message = JSON.parse(data.toString()); + messages.push(message); + for (let index = waiters.length - 1; index >= 0; index -= 1) { + if (waiters[index].predicate(message)) { + const waiter = waiters.splice(index, 1)[0]; + waiter.resolve(message); + } + } + }); + const wait = (predicate, timeout = 5_000) => new Promise((resolve, reject) => { + const existing = messages.find(predicate); + if (existing) return resolve(existing); + const timer = setTimeout(() => { + const index = waiters.findIndex((waiter) => waiter.resolve === resolve); + if (index >= 0) waiters.splice(index, 1); + reject(new Error("timed out waiting for websocket message")); + }, timeout); + waiters.push({ predicate, resolve: (message) => { clearTimeout(timer); resolve(message); } }); + }); + return new Promise((resolve, reject) => { + ws.once("open", () => { + ws.send(JSON.stringify({ type: "auth", token })); + wait((message) => message.type === "auth.ok").then(() => { + ws.send(JSON.stringify({ type: "subscribe", fromSeq: 0 })); + resolve({ ws, wait, messages }); + }, reject); + }); + ws.once("error", reject); + }); +} + +function connectBrowserHello(base, token) { + const ws = new WebSocket(base.replace(/^http/, "ws") + "/ws"); + const messages = []; + const waiters = []; + ws.on("message", (data) => { + const message = JSON.parse(data.toString()); + messages.push(message); + for (let index = waiters.length - 1; index >= 0; index -= 1) { + if (waiters[index].predicate(message)) { + const waiter = waiters.splice(index, 1)[0]; + waiter.resolve(message); + } + } + }); + const wait = (predicate, timeout = 5_000) => new Promise((resolve, reject) => { + const existing = messages.find(predicate); + if (existing) return resolve(existing); + const timer = setTimeout(() => { + const index = waiters.findIndex((waiter) => waiter.resolve === resolve); + if (index >= 0) waiters.splice(index, 1); + reject(new Error("timed out waiting for websocket message")); + }, timeout); + waiters.push({ predicate, resolve: (message) => { clearTimeout(timer); resolve(message); } }); + }); + return new Promise((resolve, reject) => { + ws.once("open", () => { + ws.send(JSON.stringify({ v: 1, kind: "hello", clientType: "web", protocol: 1 })); + if (token !== undefined) ws.send(JSON.stringify({ type: "auth", token })); + wait((message) => message.type === "auth.ok").then(() => resolve({ ws, wait, messages }), reject); + }); + ws.once("error", reject); + }); +} + +function connectHost(base, token, sessionId = "host-session") { + const ws = new WebSocket(base.replace(/^http/, "ws") + "/v1/connect"); + const messages = []; + const waiters = []; + ws.on("message", (data) => { + const message = JSON.parse(data.toString()); + messages.push(message); + for (let index = waiters.length - 1; index >= 0; index -= 1) { + if (waiters[index].predicate(message)) { + const waiter = waiters.splice(index, 1)[0]; + waiter.resolve(message); + } + } + }); + const wait = (predicate, timeout = 5_000) => new Promise((resolve, reject) => { + const existing = messages.find(predicate); + if (existing) return resolve(existing); + const timer = setTimeout(() => { + const index = waiters.findIndex((waiter) => waiter.resolve === resolve); + if (index >= 0) waiters.splice(index, 1); + reject(new Error("timed out waiting for host websocket message")); + }, timeout); + waiters.push({ predicate, resolve: (message) => { clearTimeout(timer); resolve(message); } }); + }); + return new Promise((resolve, reject) => { + ws.once("open", () => { + ws.send(JSON.stringify({ v: 1, kind: "hello", clientType: "host", protocol: 1, sessionId })); + ws.send(JSON.stringify({ v: 1, kind: "auth", accessToken: token })); + wait((message) => message.type === "auth.ok" && message.clientType === "host").then(() => { + resolve({ ws, wait, messages, sessionId }); + }, reject); + }); + ws.once("error", reject); + }); +} + +test("loopback relay defaults to tokenless browser and VS Code host connections", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + }); + assert.equal(relay.authRequired, false); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const health = await fetch(`${base}/api/health`); + assert.equal(health.status, 200); + assert.equal((await health.json()).authRequired, false); + + const state = await fetch(`${base}/api/state`); + assert.equal(state.status, 200); + assert.equal((await state.json()).role, "operator"); + + const crossOriginWrite = await fetch(`${base}/api/command`, { + method: "POST", + headers: { Origin: "https://untrusted.example", "Content-Type": "application/json" }, + body: JSON.stringify({ commandId: "cross-origin", method: "turn/start", params: {} }), + }); + assert.equal(crossOriginWrite.status, 403); + + const reboundHost = await new Promise((resolve, reject) => { + const request = http.request({ + host: "127.0.0.1", + port: address.port, + path: "/api/health", + headers: { Host: `rebound.example:${address.port}` }, + }, (response) => { + response.resume(); + response.once("end", () => resolve(response)); + }); + request.once("error", reject); + request.end(); + }); + assert.equal(reboundHost.statusCode, 403); + + const deceptiveHost = await new Promise((resolve, reject) => { + const request = http.request({ + host: "127.0.0.1", + port: address.port, + path: "/api/health", + headers: { Host: `127.evil:${address.port}` }, + }, (response) => { + response.resume(); + response.once("end", () => resolve(response)); + }); + request.once("error", reject); + request.end(); + }); + assert.equal(deceptiveHost.statusCode, 403); + + const host = await connectHost(base, undefined, "tokenless-host"); + t.after(() => host.ws.close()); + const browser = await connectWs(base, undefined); + t.after(() => browser.ws.close()); + const browserHello = await connectBrowserHello(base); + t.after(() => browserHello.ws.close()); + const browserStaleToken = await connectBrowserHello(base, "stale-local-token"); + t.after(() => browserStaleToken.ws.close()); + assert.equal(host.messages.find((message) => message.type === "auth.ok")?.role, "host"); + assert.equal(browser.messages.find((message) => message.type === "auth.ok")?.role, "operator"); + assert.equal(browserHello.messages.find((message) => message.type === "auth.ok")?.authRequired, false); + assert.equal(browserStaleToken.messages.some((message) => message.type === "error"), false); +}); + +test("loopback relay hydrates attached-session snapshots larger than 256 KiB", async (t) => { + const relay = new CodexRelay({ host: "127.0.0.1", port: 0, mode: "host" }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const host = await connectHost(base, undefined, "large-snapshot-host"); + const browser = await connectWs(base, undefined); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + const historyText = "x".repeat(512 * 1024); + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "session.snapshot", + id: "large-snapshot-event", + sessionId: host.sessionId, + seq: 1, + ts: new Date().toISOString(), + payload: { + threadId: "large-thread", + state: "idle", + outputTail: historyText.slice(-32_000), + messages: [{ kind: "assistant", text: historyText }], + metadata: { title: "Large attached session", historyComplete: false }, + }, + })); + + const snapshot = await browser.wait((message) => message.kind === "event" + && message.type === "session.snapshot" + && message.payload?.sourceSeq === 1); + assert.equal(snapshot.payload.messages[0].text.length, historyText.length); + assert.equal(relay.state.messages[0].text.length, historyText.length); + assert.equal(relay.state.sessionMetadata.title, "Large attached session"); + assert.equal(relay.state.sessionMetadata.historyComplete, false); + const replayed = relay.events.find((event) => event.type === "session.snapshot" && event.payload?.sourceSeq === 1); + assert.equal(replayed.payload.messages, undefined); + assert.equal(replayed.payload.projectionInControlSnapshot, true); + const controlSnapshot = relay.snapshot(); + assert.equal(controlSnapshot.state.messages, undefined); + assert.equal(controlSnapshot.state.outputTail, undefined); + assert.equal(controlSnapshot.state.subagents, undefined); + assert.equal(controlSnapshot.state.sessionMetadata, undefined); + assert.equal(controlSnapshot.metadata.historyComplete, false); + const serializedControl = JSON.stringify(controlSnapshot); + assert.ok(Buffer.byteLength(serializedControl) < historyText.length + 100_000, "history must be serialized only once"); +}); + +test("late browser control snapshot preserves the attach waiting state", async (t) => { + const relay = new CodexRelay({ host: "127.0.0.1", port: 0, mode: "host" }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const host = await connectHost(base, undefined, "waiting-session-host"); + t.after(() => host.ws.close()); + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "session.snapshot", + id: "waiting-session-snapshot", + sessionId: host.sessionId, + seq: 1, + ts: new Date().toISOString(), + payload: { + threadId: null, + state: "waiting_for_host", + metadata: { waitingForSession: true, attachReady: false }, + }, + })); + await host.wait((message) => message.kind === "ack" && message.seq === 1); + + const browser = await connectWs(base, undefined); + t.after(() => browser.ws.close()); + const control = await browser.wait((message) => message.type === "session.snapshot" + && message.snapshot && typeof message.snapshot === "object"); + assert.equal(control.snapshot.state.activeThreadId, null); + assert.equal(control.snapshot.metadata.waitingForSession, true); + assert.equal(control.snapshot.metadata.attachReady, false); +}); + +test("relay keeps live transcript events rich while bounding the replay ring", () => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + eventByteLimit: 4_096, + }); + const liveClient = { id: "live", role: "operator", authenticated: true, subscribed: true, capture: [] }; + relay.clients.add(liveClient); + relay.recordEvent("output.chunk", { + threadId: "thread-large", + text: "delta", + outputTail: "o".repeat(32_000), + messages: [{ text: "m".repeat(32_000) }], + messagesPatch: { start: 0, deleteCount: 0, messages: [{ text: "p".repeat(32_000) }] }, + subagents: [{ output: "s".repeat(32_000) }], + raw: { transcript: "r".repeat(32_000) }, + }); + + assert.equal(liveClient.capture[0].payload.messagesPatch.messages[0].text.length, 32_000); + const replayed = relay.events[0]; + assert.equal(replayed.payload.text, "delta"); + assert.equal(replayed.payload.messages, undefined); + assert.equal(replayed.payload.messagesPatch, undefined); + assert.equal(replayed.payload.outputTail, undefined); + assert.equal(replayed.payload.subagents, undefined); + assert.equal(replayed.payload.raw, undefined); + assert.equal(replayed.payload.projectionInControlSnapshot, true); + + relay.clients.delete(liveClient); + for (let index = 0; index < 30; index += 1) { + relay.recordEvent("command.pending", { commandId: `command-${index}`, note: "n".repeat(400) }); + } + assert.ok(relay.eventBytes <= 4_096); + assert.equal(relay.eventBytes, relay.eventSizes.reduce((total, size) => total + size, 0)); + assert.ok(relay.events.length < 30, "byte budget should prune before the count limit"); +}); + +test("relay skips oversized replay windows and sends one authoritative snapshot", () => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + eventByteLimit: 64 * 1024, + replayByteLimit: 400, + }); + relay.state.messages = [{ kind: "assistant", text: "authoritative history" }]; + for (let index = 0; index < 5; index += 1) { + relay.recordEvent("command.pending", { commandId: `command-${index}`, note: "n".repeat(180) }); + } + const client = { id: "late", role: "operator", authenticated: true, subscribed: false, capture: [] }; + relay.clients.add(client); + relay.subscribe(client, 0); + + assert.equal(client.capture[0].type, "resync.required"); + assert.equal(client.capture[0].reason, "replay_too_large"); + assert.equal(client.capture.filter((message) => message.kind === "event").length, 0); + const snapshot = client.capture.find((message) => message.type === "session.snapshot"); + assert.equal(snapshot.snapshot.messages[0].text, "authoritative history"); + assert.equal(snapshot.snapshot.latestSeq, relay.nextSeq); +}); + +test("relay socket backpressure accounts for the serialized frame size", () => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + clientBufferedByteLimit: 2_048, + }); + const accepted = { + readyState: WebSocket.OPEN, + bufferedAmount: 1_000, + sent: [], + send(value) { this.sent.push(value); }, + close(code, reason) { this.closed = { code, reason }; }, + }; + relay.sendControl({ socket: accepted }, { type: "small", text: "x".repeat(500) }); + assert.equal(accepted.sent.length, 1); + assert.equal(accepted.closed, undefined); + + const slow = { + readyState: WebSocket.OPEN, + bufferedAmount: 1_900, + sent: [], + send(value) { this.sent.push(value); }, + close(code, reason) { this.closed = { code, reason }; }, + }; + relay.sendControl({ socket: slow }, { type: "small", text: "x".repeat(500) }); + assert.equal(slow.sent.length, 0); + assert.equal(slow.closed.code, 1013); +}); + +test("loopback relay can explicitly require tokens", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + authRequired: true, + operatorToken: "explicit-operator-token", + viewerToken: "explicit-viewer-token", + hostToken: "explicit-host-token", + }); + assert.equal(relay.authRequired, true); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const health = await fetch(`${base}/api/health`); + assert.equal((await health.json()).authRequired, true); + assert.equal((await fetch(`${base}/api/state`)).status, 401); + + const browser = await connectWs(base, "explicit-operator-token"); + t.after(() => browser.ws.close()); + assert.equal(browser.messages.find((message) => message.type === "auth.ok")?.role, "operator"); +}); + +test("non-loopback relay requires tokens by default", () => { + const relay = new CodexRelay({ + host: "0.0.0.0", + port: 0, + mode: "host", + }); + assert.equal(relay.authRequired, true); +}); + +test("deceptive numeric-looking hostnames never enable local no-auth", () => { + const relay = new CodexRelay({ host: "127.evil", port: 0, mode: "host" }); + assert.equal(relay.authRequired, true); +}); + +test("relay defaults to host mode so bare npm start does not spawn Codex", () => { + const relay = new CodexRelay({ host: "127.0.0.1", port: 0, authRequired: false }); + assert.equal(relay.mode, "host"); + assert.equal(relay.spawnCodex, false); +}); + +test("thread settings validation accepts null reasoning effort", () => { + const relay = new CodexRelay({ host: "127.0.0.1", port: 0, mode: "host" }); + assert.equal(relay.validateCommand("thread/settings/update", { + threadId: "thread-test", + threadSettings: { model: "gpt-5.6-sol", effort: null }, + }), null); + assert.match(relay.validateCommand("thread/settings/update", { + threadId: "thread-test", + threadSettings: { effort: 3 }, + }), /string or null/); +}); + +test("control mode validation only accepts sync and async", () => { + const relay = new CodexRelay({ host: "127.0.0.1", port: 0, mode: "host" }); + assert.equal(relay.validateCommand("control/mode/get", {}), null); + assert.equal(relay.validateCommand("control/mode/set", { mode: "sync" }), null); + assert.equal(relay.validateCommand("control/mode/set", { mode: "async" }), null); + assert.match(relay.validateCommand("control/mode/set", { mode: "attach" }), /sync or async/); + assert.match(relay.validateCommand("control/mode/get", { mode: "sync" }), /does not accept/); +}); + +test("authenticated relay forwards commands, output, approval requests, and replay", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "embedded", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + codexCommand: process.execPath, + codexArgs: [fakeServer], + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + await waitFor(() => relay.state.initialized); + + const health = await fetch(`${base}/api/health`); + assert.equal(health.status, 200); + const unauthorized = await fetch(`${base}/api/state`); + assert.equal(unauthorized.status, 401); + + const viewerCommand = await request(base, "viewer-test-token", "/api/command", { + method: "POST", + body: JSON.stringify({ commandId: "viewer-1", method: "thread/start", params: {} }), + headers: { "Content-Type": "application/json" }, + }); + assert.equal(viewerCommand.status, 400); + assert.equal((await viewerCommand.json()).code, "forbidden"); + + const client = await connectWs(base, "operator-test-token"); + t.after(() => client.ws.close()); + await client.wait((message) => message.type === "session.snapshot"); + + client.ws.send(JSON.stringify({ type: "command", commandId: "thread-1", method: "thread/start", params: { cwd: "/tmp", sandbox: "workspace-write" } })); + const threadResult = await client.wait((message) => message.type === "command.result" && message.payload?.commandId === "thread-1"); + assert.equal(threadResult.payload.ok, true); + const threadId = threadResult.payload.result.thread.id; + assert.equal(relay.state.activeThreadId, threadId); + + client.ws.send(JSON.stringify({ type: "command", commandId: "turn-1", method: "turn/start", params: { threadId, input: [{ type: "text", text: "approve this", text_elements: [] }] } })); + await client.wait((message) => message.type === "command.result" && message.payload?.commandId === "turn-1"); + const approval = await client.wait((message) => message.kind === "event" && message.type === "approval.requested"); + assert.equal(approval.payload.requestId, 9001); + assert.equal(approval.payload.method, "item/commandExecution/requestApproval"); + + client.ws.send(JSON.stringify({ type: "respond", requestId: "9001", result: { decision: "accept" } })); + await client.wait((message) => message.kind === "event" && message.type === "server.responded"); + await client.wait((message) => message.kind === "event" && message.type === "output.delta" && message.payload.text.includes("approval response")); + assert.equal(relay.pendingServerRequests.size, 0); + + const events = await request(base, "viewer-test-token", "/api/events?fromSeq=0"); + const eventBody = await events.json(); + assert.ok(eventBody.latestSeq >= 1); + assert.ok(eventBody.events.some((event) => event.type === "approval.requested")); +}); + +test("host mode proxies browser commands and preserves one-shot approval request ids", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const host = await connectHost(base, "separate-host-token", "session-from-vscode"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + await browser.wait((message) => message.type === "session.snapshot"); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + id: "host-event-1", + sessionId: host.sessionId, + seq: 1, + ts: new Date().toISOString(), + payload: {}, + })); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + assert.equal(relay.state.initialized, true); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "output.chunk", + id: "host-event-2", + sessionId: host.sessionId, + seq: 2, + ts: new Date().toISOString(), + payload: { stream: "codex", text: "output from VS Code" }, + })); + const output = await browser.wait((message) => message.kind === "event" && message.type === "output.chunk"); + assert.equal(output.payload.text, "output from VS Code"); + const ack = await host.wait((message) => message.kind === "ack" && message.seq === 2); + assert.equal(ack.sessionId, host.sessionId); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "host-thread-1", + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const hostCommand = await host.wait((message) => message.kind === "command" && message.commandId === "host-thread-1"); + assert.equal(hostCommand.type, "thread/start"); + assert.equal(hostCommand.payload.sandbox, "workspace-write"); + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "command.accepted", + id: "host-result-1", + sessionId: host.sessionId, + seq: 3, + ts: new Date().toISOString(), + payload: { commandId: "host-thread-1", method: "thread/start", ok: true, result: { thread: { id: "thread-on-host" } } }, + })); + const result = await browser.wait((message) => message.kind === "event" && message.type === "command.result" && message.payload.commandId === "host-thread-1"); + assert.equal(result.payload.ok, true); + assert.equal(relay.state.activeThreadId, "thread-on-host"); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "approval.requested", + id: "host-approval-1", + sessionId: host.sessionId, + seq: 4, + ts: new Date().toISOString(), + payload: { + requestId: 77, + method: "item/commandExecution/requestApproval", + commandHash: "approval-hash-77", + params: { command: "echo approved" }, + }, + })); + const approval = await browser.wait((message) => message.kind === "event" && message.type === "approval.requested" && message.payload.requestId === 77); + assert.equal(approval.payload.params.command, "echo approved"); + browser.ws.send(JSON.stringify({ type: "respond", requestId: "77", result: { decision: "accept" } })); + const responseCommand = await host.wait((message) => message.kind === "command" && message.type === "approval.respond"); + assert.equal(responseCommand.payload.requestId, 77); + assert.equal(responseCommand.payload.decision, "allow"); + assert.equal(responseCommand.payload.commandHash, "approval-hash-77"); + assert.deepEqual(responseCommand.payload.response, { decision: "accept" }); + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "command.accepted", + id: "host-response-result-1", + sessionId: host.sessionId, + seq: 5, + ts: new Date().toISOString(), + payload: { commandId: responseCommand.commandId, method: "approval.respond", ok: true, result: { accepted: true } }, + })); + await browser.wait((message) => message.kind === "event" && message.type === "server.responded" && String(message.payload.requestId) === "77"); + assert.equal(relay.pendingServerRequests.size, 0); + + browser.ws.send(JSON.stringify({ type: "respond", requestId: "77", result: { decision: "accept" } })); + const duplicate = await browser.wait((message) => message.type === "response.rejected" && String(message.requestId) === "77"); + assert.equal(duplicate.code, "unknown_request"); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "approval.requested", + id: "host-approval-2", + sessionId: host.sessionId, + seq: 6, + ts: new Date().toISOString(), + payload: { requestId: 78, method: "item/commandExecution/requestApproval", params: { command: "echo retry" } }, + })); + await browser.wait((message) => message.kind === "event" && message.type === "approval.requested" && message.payload.requestId === 78); + browser.ws.send(JSON.stringify({ type: "respond", requestId: "78", result: { decision: "accept" } })); + const rejectedCommand = await host.wait((message) => message.kind === "command" && message.type === "approval.respond" && message.payload.requestId === 78); + assert.equal(rejectedCommand.payload.decision, "allow"); + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "command.rejected", + id: "host-response-rejected-1", + sessionId: host.sessionId, + seq: 7, + ts: new Date().toISOString(), + payload: { commandId: rejectedCommand.commandId, method: "approval.respond", ok: false, error: { message: "local policy denied" } }, + })); + const responseRejected = await browser.wait((message) => message.type === "response.rejected" && String(message.requestId) === "78"); + assert.equal(responseRejected.code, "host_rejected"); + assert.equal([...relay.pendingServerRequests.values()].some((pending) => String(pending.appId) === "78"), true); +}); + +test("host mode forwards session list/select and publishes the selected session", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "session-picker-host"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + const sendEvent = (type, seq, payload) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type, + id: `session-picker-event-${seq}`, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload, + })); + + sendEvent("connection.opened", 1, { mode: "attach" }); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + sendEvent("session.snapshot", 2, { + threadId: "thread-one", + activeThreadId: "thread-one", + title: "当前会话", + metadata: { + controlMode: "sync", + modeEpoch: 0, + capabilities: { + followsVscodeRoute: true, + sessionList: false, + sessionSelect: false, + sessionCreate: false, + threadSettings: true, + }, + }, + }); + await browser.wait((message) => message.kind === "event" && message.type === "session.snapshot"); + assert.equal(relay.snapshot().metadata.controlMode, "sync"); + assert.equal(relay.snapshot().metadata.capabilities.sessionSelect, false); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "session-list-1", + method: "session/list", + params: { limit: 10 }, + })); + const listCommand = await host.wait((message) => message.kind === "command" && message.commandId === "session-list-1"); + assert.equal(listCommand.type, "session/list"); + assert.equal(listCommand.payload.limit, 10); + sendEvent("command.result", 3, { + commandId: "session-list-1", + method: "session/list", + ok: true, + result: { + activeThreadId: "thread-one", + sessions: [ + { threadId: "thread-one", title: "当前会话", active: true, available: true }, + { threadId: "thread-two", title: "另一个会话", active: false, available: true }, + ], + }, + }); + const listResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" && message.payload?.commandId === "session-list-1"); + assert.equal(listResult.payload.result.sessions.length, 2); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "session-select-1", + method: "session/select", + params: { threadId: "thread-two" }, + })); + const selectCommand = await host.wait((message) => message.kind === "command" && message.commandId === "session-select-1"); + assert.equal(selectCommand.type, "session/select"); + assert.equal(selectCommand.payload.threadId, "thread-two"); + sendEvent("session.switching", 4, { previousThreadId: "thread-one", targetThreadId: "thread-two" }); + sendEvent("session.snapshot", 5, { threadId: "thread-two", activeThreadId: "thread-two", title: "另一个会话" }); + sendEvent("session.selected", 6, { threadId: "thread-two", activeThreadId: "thread-two" }); + sendEvent("command.result", 7, { + commandId: "session-select-1", + method: "session/select", + ok: true, + result: { threadId: "thread-two", previousThreadId: "thread-one", switched: true, available: true }, + }); + await browser.wait((message) => message.kind === "event" && message.type === "session.switching"); + await browser.wait((message) => message.kind === "event" && message.type === "session.selected"); + const selectResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" && message.payload?.commandId === "session-select-1"); + assert.equal(selectResult.payload.result.threadId, "thread-two"); + assert.equal(relay.state.activeThreadId, "thread-two"); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "session-new-1", + method: "session/new", + params: {}, + })); + const newCommand = await host.wait((message) => message.kind === "command" && message.commandId === "session-new-1"); + assert.equal(newCommand.type, "session/new"); + sendEvent("command.result", 8, { + commandId: "session-new-1", + method: "session/new", + ok: true, + result: { opened: true, command: "chatgpt.newCodexPanel" }, + }); + const newResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" && message.payload?.commandId === "session-new-1"); + assert.equal(newResult.payload.result.command, "chatgpt.newCodexPanel"); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "mode-set-1", + method: "control/mode/set", + params: { mode: "async" }, + })); + const modeCommand = await host.wait((message) => message.kind === "command" && message.commandId === "mode-set-1"); + assert.equal(modeCommand.type, "control/mode/set"); + assert.deepEqual(modeCommand.payload, { mode: "async" }); + sendEvent("command.result", 9, { + commandId: "mode-set-1", + method: "control/mode/set", + ok: true, + result: { mode: "async", modeEpoch: 1 }, + }); + const modeResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" && message.payload?.commandId === "mode-set-1"); + assert.equal(modeResult.payload.result.mode, "async"); +}); + +test("host mode records normalized server.requested events for response routing", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "generic-request-host"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + id: "generic-ready", + sessionId: host.sessionId, + seq: 1, + ts: new Date().toISOString(), + payload: {}, + })); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "server.requested", + id: "generic-request", + sessionId: host.sessionId, + seq: 2, + ts: new Date().toISOString(), + payload: { + requestId: "generic-1", + method: "custom/request", + params: { prompt: "host-only" }, + }, + })); + const requestEvent = await browser.wait((message) => message.kind === "event" && message.type === "server.requested"); + assert.equal(requestEvent.payload.requestId, "generic-1"); + assert.equal(relay.pendingServerRequests.get("string:generic-1")?.method, "custom/request"); + + browser.ws.send(JSON.stringify({ type: "respond", requestId: "generic-1", result: { accepted: true } })); + const rejected = await browser.wait((message) => message.type === "response.rejected" && message.requestId === "generic-1"); + assert.equal(rejected.code, "unsupported_request"); + assert.equal(relay.pendingServerRequests.has("string:generic-1"), true); +}); + +test("host mode accepts versioned command, approval, and input response frames", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "versioned-host-session"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + id: "versioned-ready", + sessionId: host.sessionId, + seq: 1, + ts: new Date().toISOString(), + payload: {}, + })); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "thread.start", + commandId: "versioned-thread-1", + payload: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const threadCommand = await host.wait((message) => message.kind === "command" && message.commandId === "versioned-thread-1"); + assert.equal(threadCommand.type, "thread/start"); + assert.deepEqual(threadCommand.payload, { cwd: "/tmp", sandbox: "workspace-write" }); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "approval.requested", + id: "versioned-approval-request", + sessionId: host.sessionId, + seq: 2, + ts: new Date().toISOString(), + payload: { + requestId: 91, + method: "item/commandExecution/requestApproval", + commandHash: "versioned-hash-91", + params: { command: "echo versioned" }, + }, + })); + await browser.wait((message) => message.kind === "event" && message.type === "approval.requested" && message.payload.requestId === 91); + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "approval.respond", + commandId: "browser-approval-91", + payload: { requestId: 91, decision: "allow", response: { decision: "accept" } }, + })); + const approvalCommand = await host.wait((message) => message.kind === "command" && message.type === "approval.respond" && message.payload.requestId === 91); + assert.equal(approvalCommand.payload.commandHash, "versioned-hash-91"); + assert.equal(approvalCommand.payload.decision, "allow"); + assert.deepEqual(approvalCommand.payload.response, { decision: "accept" }); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "input.requested", + id: "versioned-input-request", + sessionId: host.sessionId, + seq: 3, + ts: new Date().toISOString(), + payload: { + requestId: 92, + method: "item/tool/requestUserInput", + params: { questions: [{ id: "choice", question: "Continue?" }] }, + }, + })); + await browser.wait((message) => message.kind === "event" && message.type === "input.requested" && message.payload.requestId === 92); + browser.ws.send(JSON.stringify({ + type: "input.respond", + payload: { requestId: 92, answers: { choice: { answers: ["yes"] } } }, + })); + const inputCommand = await host.wait((message) => message.kind === "command" && message.type === "server.request.respond" && message.payload.requestId === 92); + assert.equal(inputCommand.payload.decision, "allow"); + assert.deepEqual(inputCommand.payload.response, { answers: { choice: { answers: ["yes"] } } }); + + host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "approval.requested", + id: "versioned-tagged-approval", + sessionId: host.sessionId, + seq: 4, + ts: new Date().toISOString(), + payload: { + requestId: 93, + method: "item/commandExecution/requestApproval", + params: { command: "echo amend" }, + }, + })); + await browser.wait((message) => message.kind === "event" && message.type === "approval.requested" && message.payload.requestId === 93); + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "approval.respond", + commandId: "browser-tagged-approval-93", + payload: { + requestId: 93, + decision: "allow", + response: { decision: { acceptWithExecpolicyAmendment: { execpolicy_amendment: ["echo"] } } }, + }, + })); + const taggedCommand = await host.wait((message) => message.kind === "command" && message.type === "approval.respond" && message.payload.requestId === 93); + assert.equal(taggedCommand.payload.decision, "allow"); + assert.deepEqual(taggedCommand.payload.response, { + decision: { acceptWithExecpolicyAmendment: { execpolicy_amendment: ["echo"] } }, + }); +}); + +test("keeps numeric and string host approval ids distinct end to end", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "typed-id-session"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + const sendHostEvent = (type, seq, payload, id) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type, + id: id || `typed-${seq}`, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload, + })); + sendHostEvent("connection.opened", 1, {}); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + sendHostEvent("approval.requested", 2, { + requestId: 1, + method: "item/commandExecution/requestApproval", + params: { command: "echo numeric" }, + }); + sendHostEvent("approval.requested", 3, { + requestId: "1", + method: "item/commandExecution/requestApproval", + params: { command: "echo string" }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === 1); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === "1"); + assert.equal(relay.pendingServerRequests.size, 2); + + browser.ws.send(JSON.stringify({ type: "respond", requestId: 1, result: { decision: "accept" } })); + browser.ws.send(JSON.stringify({ type: "respond", requestId: "1", result: { decision: "accept" } })); + const firstCommand = await host.wait((message) => message.kind === "command" + && message.type === "approval.respond" + && message.payload?.requestId === 1); + const secondCommand = await host.wait((message) => message.kind === "command" + && message.type === "approval.respond" + && message.payload?.requestId === "1"); + assert.notEqual(firstCommand.commandId, secondCommand.commandId); + + sendHostEvent("command.result", 4, { + commandId: firstCommand.commandId, + method: "approval.respond", + ok: true, + result: { accepted: true }, + }); + sendHostEvent("command.result", 5, { + commandId: secondCommand.commandId, + method: "approval.respond", + ok: true, + result: { accepted: true }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "server.responded" + && message.payload?.requestId === 1); + await browser.wait((message) => message.kind === "event" + && message.type === "server.responded" + && message.payload?.requestId === "1"); + assert.equal(relay.pendingServerRequests.size, 0); +}); + +test("normalizes legacy approval decisions and rejects outer/inner conflicts", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "decision-schema-session"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + const sendHostEvent = (type, seq, payload) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type, + id: `decision-${seq}`, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload, + })); + sendHostEvent("connection.opened", 1, {}); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + sendHostEvent("approval.requested", 2, { + requestId: 201, + method: "applyPatchApproval", + params: { reason: "legacy patch" }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === 201); + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "approval.respond", + commandId: "legacy-approval-201", + payload: { requestId: 201, decision: "approved" }, + })); + const legacyCommand = await host.wait((message) => message.kind === "command" + && message.type === "approval.respond" + && message.payload?.requestId === 201); + assert.deepEqual(legacyCommand.payload.response, { decision: "approved" }); + + sendHostEvent("approval.requested", 3, { + requestId: 202, + method: "item/commandExecution/requestApproval", + params: { command: "echo conflict" }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === 202); + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "approval.respond", + commandId: "conflicting-approval-202", + payload: { + requestId: 202, + decision: "deny", + response: { decision: "accept" }, + }, + })); + const rejected = await browser.wait((message) => message.type === "response.rejected" + && message.requestId === 202); + assert.equal(rejected.code, "decision_mismatch"); + assert.equal(relay.pendingServerRequests.size, 2); + + sendHostEvent("approval.requested", 4, { + requestId: 203, + method: "item/commandExecution/requestApproval", + params: { command: "echo mixed-tag" }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === 203); + browser.ws.send(JSON.stringify({ + v: 1, + kind: "command", + type: "approval.respond", + commandId: "mixed-tag-approval-203", + payload: { + requestId: 203, + decision: "allow", + response: { + decision: { + acceptWithExecpolicyAmendment: { execpolicy_amendment: ["echo"] }, + futurePolicyGrant: { scope: "all" }, + }, + }, + }, + })); + const mixedRejected = await browser.wait((message) => message.type === "response.rejected" + && message.requestId === 203); + assert.equal(mixedRejected.code, "invalid_response"); + assert.equal(relay.pendingServerRequests.size, 3); +}); + +test("host command result cache is unavailable offline and isolated by host session", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const sendReady = (host, seq, id) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + id, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload: {}, + })); + const sendCommandResult = (host, seq, id, commandId, threadId) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "command.result", + id, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload: { + commandId, + method: "thread/start", + ok: true, + result: { thread: { id: threadId } }, + }, + })); + + const host1 = await connectHost(base, "separate-host-token", "host-session-1"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host1.ws.close()); + t.after(() => browser.ws.close()); + sendReady(host1, 1, "host1-ready"); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + const commandId = "reused-command-id"; + browser.ws.send(JSON.stringify({ + type: "command", + commandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + await host1.wait((message) => message.kind === "command" && message.commandId === commandId); + sendCommandResult(host1, 2, "host1-command-result", commandId, "thread-from-session-1"); + const firstResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" + && message.payload?.commandId === commandId + && message.payload?.result?.thread?.id === "thread-from-session-1"); + assert.equal(firstResult.payload.ok, true); + + const uncertainCommandId = "uncertain-command-id"; + browser.ws.send(JSON.stringify({ + type: "command", + commandId: uncertainCommandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + await host1.wait((message) => message.kind === "command" && message.commandId === uncertainCommandId); + + host1.ws.close(); + await waitFor(() => relay.hostClient === null); + await browser.wait((message) => message.type === "command.result" + && message.commandId === uncertainCommandId + && message.uncertain === true); + + // A stale success (or uncertain disconnect result) must not be replayed + // while no host is connected. + browser.ws.send(JSON.stringify({ + type: "command", + commandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const offline = await browser.wait((message) => message.type === "command.rejected" && message.commandId === commandId); + assert.equal(offline.code, "app_not_ready"); + browser.ws.send(JSON.stringify({ + type: "command", + commandId: uncertainCommandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const uncertainOffline = await browser.wait((message) => message.type === "command.rejected" && message.commandId === uncertainCommandId); + assert.equal(uncertainOffline.code, "app_not_ready"); + + // A reconnect carrying the same stable session id retains normal command + // idempotency and returns the cached result without forwarding a command. + const sameSessionHost = await connectHost(base, "separate-host-token", "host-session-1"); + t.after(() => sameSessionHost.ws.close()); + sendReady(sameSessionHost, 1, "same-session-ready"); + await browser.wait((message) => message.kind === "event" + && message.type === "connection.opened" + && message.payload?.source === "vscode-host"); + browser.ws.send(JSON.stringify({ + type: "command", + commandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const replay = await browser.wait((message) => message.type === "command.result" + && message.cached === true + && message.commandId === commandId); + assert.equal(replay.result.thread.id, "thread-from-session-1"); + await new Promise((resolve) => setTimeout(resolve, 50)); + assert.equal(sameSessionHost.messages.some((message) => message.kind === "command" && message.commandId === commandId), false); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: uncertainCommandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + await sameSessionHost.wait((message) => message.kind === "command" && message.commandId === uncertainCommandId); + sendCommandResult(sameSessionHost, 2, "same-session-uncertain-result", uncertainCommandId, "thread-after-uncertain"); + const uncertainRetry = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" + && message.payload?.commandId === uncertainCommandId + && message.payload?.result?.thread?.id === "thread-after-uncertain"); + assert.equal(uncertainRetry.payload.ok, true); + + sameSessionHost.ws.close(); + await waitFor(() => relay.hostClient === null); + + // A different host session cannot reuse the old command id; it must receive + // a fresh command even though the browser retries the same id. + const host2 = await connectHost(base, "separate-host-token", "host-session-2"); + t.after(() => host2.ws.close()); + sendReady(host2, 1, "host2-ready"); + await browser.wait((message) => message.kind === "event" + && message.type === "connection.opened" + && message.payload?.source === "vscode-host"); + browser.ws.send(JSON.stringify({ + type: "command", + commandId, + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const forwarded = await host2.wait((message) => message.kind === "command" && message.commandId === commandId); + assert.equal(forwarded.type, "thread/start"); + sendCommandResult(host2, 2, "host2-command-result", commandId, "thread-from-session-2"); + const secondResult = await browser.wait((message) => message.kind === "event" + && message.type === "command.result" + && message.payload?.commandId === commandId + && message.payload?.result?.thread?.id === "thread-from-session-2"); + assert.equal(secondResult.payload.ok, true); +}); + +test("rejects non-object websocket frames without taking down the relay", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "embedded", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + codexCommand: process.execPath, + codexArgs: [fakeServer], + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const ws = new WebSocket(`ws://127.0.0.1:${address.port}/ws`); + t.after(() => ws.close()); + const invalidFrame = new Promise((resolve, reject) => { + const timer = setTimeout(() => reject(new Error("timed out waiting for invalid frame response")), 5_000); + ws.once("error", reject); + ws.on("message", (data) => { + const message = JSON.parse(data.toString()); + if (message.code === "invalid_frame") { + clearTimeout(timer); + resolve(message); + } + }); + }); + await new Promise((resolve, reject) => { + ws.once("open", resolve); + ws.once("error", reject); + }); + ws.send("null"); + await invalidFrame; + const health = await fetch(`http://127.0.0.1:${address.port}/api/health`); + assert.equal(health.status, 200); +}); + +test("malformed percent escapes return 400 without crashing the relay", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + + const malformed = await fetch(`${base}/%ZZ`); + assert.equal(malformed.status, 400); + assert.equal(await malformed.text(), "Invalid URL"); + + const health = await fetch(`${base}/api/health`); + assert.equal(health.status, 200); +}); + +test("cleans embedded app pending commands and approvals when the child exits", async (t) => { + const crashServer = [ + "const readline = require('node:readline');", + "const input = readline.createInterface({ input: process.stdin });", + "const send = (message) => process.stdout.write(JSON.stringify(message) + '\\n');", + "input.on('line', (line) => {", + " let request; try { request = JSON.parse(line); } catch { return; }", + " if (request.method === 'initialize') { send({ id: request.id, result: { userAgent: 'crash-test' } }); return; }", + " if (request.method === 'thread/start') { send({ id: request.id, result: { thread: { id: 'thread-crash' }, cwd: '/tmp' } }); return; }", + " if (request.method === 'turn/start') {", + " send({ id: 4321, method: 'item/commandExecution/requestApproval', params: { command: 'echo crash' } });", + " setTimeout(() => process.exit(23), 30);", + " }", + "});", + ].join("\n"); + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "embedded", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + codexCommand: process.execPath, + codexArgs: ["-e", crashServer], + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const browser = await connectWs(base, "operator-test-token"); + t.after(() => browser.ws.close()); + await browser.wait((message) => message.type === "session.snapshot"); + await waitFor(() => relay.state.initialized === true); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "crash-thread", + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const threadResult = await browser.wait((message) => message.type === "command.result" + && message.payload?.commandId === "crash-thread"); + assert.equal(threadResult.payload.ok, true); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "crash-turn", + method: "turn/start", + params: { threadId: "thread-crash", input: [{ type: "text", text: "crash" }] }, + })); + await browser.wait((message) => message.type === "command.accepted" && message.commandId === "crash-turn"); + const approval = await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && String(message.payload?.requestId) === "4321"); + assert.equal(approval.payload.method, "item/commandExecution/requestApproval"); + assert.equal([...relay.pendingServerRequests.values()].some((pending) => String(pending.appId) === "4321"), true); + + const uncertain = await browser.wait((message) => message.type === "command.result" + && message.commandId === "crash-turn" + && message.uncertain === true); + assert.equal(uncertain.retryable, true); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.expired" + && String(message.payload?.requestId) === "4321"); + await waitFor(() => relay.state.app === "offline"); + assert.equal(relay.state.initialized, false); + assert.equal(relay.pendingAppRequests.size, 0); + assert.equal(relay.pendingServerRequests.size, 0); + assert.equal(relay.appProcess, null); + assert.throws(() => relay.sendToApp({ method: "ping" }), (error) => error.code === "app_offline"); + + // The uncertain marker is deliberately not replayed. A retry while the app + // is offline receives the normal readiness error instead of a duplicate. + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "crash-turn", + method: "turn/start", + params: { threadId: "thread-crash", input: [{ type: "text", text: "retry" }] }, + })); + const rejected = await browser.wait((message) => message.type === "command.rejected" + && message.commandId === "crash-turn"); + assert.equal(rejected.code, "app_not_ready"); +}); + +test("cleans host pending work on app connection.closed and honors expiry events", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const host = await connectHost(base, "separate-host-token", "cleanup-host-session"); + const browser = await connectWs(base, "operator-test-token"); + t.after(() => host.ws.close()); + t.after(() => browser.ws.close()); + + const sendHostEvent = (type, seq, payload, id = `cleanup-${seq}`) => host.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type, + id, + sessionId: host.sessionId, + seq, + ts: new Date().toISOString(), + payload, + })); + + sendHostEvent("connection.opened", 1, {}); + await browser.wait((message) => message.kind === "event" && message.type === "connection.opened"); + + // Complete one command first so the app-unavailable transition has a cache + // entry to invalidate as well as an in-flight command to settle. + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "host-completed-before-exit", + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + await host.wait((message) => message.kind === "command" && message.commandId === "host-completed-before-exit"); + sendHostEvent("command.result", 2, { + commandId: "host-completed-before-exit", + method: "thread/start", + ok: true, + result: { thread: { id: "host-thread-cleanup" } }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "command.result" + && message.payload?.commandId === "host-completed-before-exit"); + assert.equal(relay.commandResults.has("host-completed-before-exit"), true); + + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "host-pending-before-exit", + method: "turn/start", + params: { threadId: "host-thread-cleanup", input: [{ type: "text", text: "pending" }] }, + })); + await host.wait((message) => message.kind === "command" && message.commandId === "host-pending-before-exit"); + assert.equal(relay.pendingHostCommands.has("host-pending-before-exit"), true); + + sendHostEvent("approval.requested", 3, { + requestId: 501, + method: "item/commandExecution/requestApproval", + params: { command: "echo pending approval" }, + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.requested" + && message.payload?.requestId === 501); + assert.equal([...relay.pendingServerRequests.values()].some((pending) => String(pending.appId) === "501"), true); + + sendHostEvent("approval.expired", 4, { + requestId: 501, + method: "item/commandExecution/requestApproval", + }); + await browser.wait((message) => message.kind === "event" + && message.type === "approval.expired" + && message.payload?.requestId === 501); + assert.equal([...relay.pendingServerRequests.values()].some((pending) => String(pending.appId) === "501"), false); + + sendHostEvent("connection.closed", 5, { message: "embedded app exited" }); + await browser.wait((message) => message.kind === "event" && message.type === "connection.closed"); + const uncertain = await browser.wait((message) => message.type === "command.result" + && message.commandId === "host-pending-before-exit" + && message.uncertain === true); + assert.equal(uncertain.error.code, "app_unavailable"); + await waitFor(() => relay.state.app === "offline"); + assert.equal(relay.state.initialized, false); + assert.equal(relay.state.hostConnected, true); + assert.equal(relay.pendingHostCommands.size, 0); + assert.equal(relay.pendingServerRequests.size, 0); + assert.equal(relay.commandResults.size, 0); + + // The host transport remains connected, but commands are rejected until it + // reports a fresh connection.opened/session snapshot. + browser.ws.send(JSON.stringify({ + type: "command", + commandId: "host-completed-before-exit", + method: "thread/start", + params: { cwd: "/tmp", sandbox: "workspace-write" }, + })); + const rejected = await browser.wait((message) => message.type === "command.rejected" + && message.commandId === "host-completed-before-exit"); + assert.equal(rejected.code, "app_not_ready"); +}); + +test("ignores a stale host connection.closed frame after host replacement", async (t) => { + const relay = new CodexRelay({ + host: "127.0.0.1", + port: 0, + mode: "host", + operatorToken: "operator-test-token", + viewerToken: "viewer-test-token", + hostToken: "separate-host-token", + }); + await relay.start(); + t.after(() => relay.stop()); + const address = relay.address(); + const base = `http://127.0.0.1:${address.port}`; + const firstHost = await connectHost(base, "separate-host-token", "stale-session-1"); + t.after(() => firstHost.ws.close()); + firstHost.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + sessionId: firstHost.sessionId, + seq: 1, + payload: {}, + })); + await waitFor(() => relay.state.app === "ready"); + const staleClient = relay.hostClient; + firstHost.ws.close(); + await waitFor(() => relay.hostClient === null); + + const secondHost = await connectHost(base, "separate-host-token", "stale-session-2"); + t.after(() => secondHost.ws.close()); + secondHost.ws.send(JSON.stringify({ + v: 1, + kind: "event", + type: "connection.opened", + sessionId: secondHost.sessionId, + seq: 1, + payload: {}, + })); + await waitFor(() => relay.state.app === "ready" && relay.hostClient?.sessionId === "stale-session-2"); + const sequenceBefore = relay.nextSeq; + + relay.ingestHostEvent(staleClient, { + v: 1, + kind: "event", + type: "connection.closed", + sessionId: "stale-session-1", + seq: 99, + payload: { message: "late old app exit" }, + }); + + assert.equal(relay.state.app, "ready"); + assert.equal(relay.state.initialized, true); + assert.equal(relay.state.hostConnected, true); + assert.equal(relay.state.hostSessionId, "stale-session-2"); + assert.equal(relay.nextSeq, sequenceBefore); +}); diff --git a/aether-vscodex/test/switchable-agent-adapter.test.js b/aether-vscodex/test/switchable-agent-adapter.test.js new file mode 100644 index 000000000..643807a46 --- /dev/null +++ b/aether-vscodex/test/switchable-agent-adapter.test.js @@ -0,0 +1,387 @@ +"use strict"; + +const assert = require("node:assert/strict"); +const test = require("node:test"); + +const { SwitchableAgentAdapter } = require("../vscode-extension/dist/switchableAgentAdapter.js"); + +class FakeAdapter { + constructor(name, options = {}) { + this.name = name; + this.options = options; + this.listeners = new Set(); + this.calls = []; + this.disposed = false; + this.snapshotValue = options.snapshot ?? idleSnapshot(name); + } + + async start() { + this.calls.push(["start"]); + this.emit({ type: "candidate.starting", payload: { name: this.name } }); + if (this.options.startGate) await this.options.startGate.promise; + if (this.options.startError) throw this.options.startError; + this.emit({ type: "connection.opened", payload: { name: this.name } }); + } + + async startThread(params = {}) { return this.record("startThread", params); } + async newSession(params = {}) { return this.record("newSession", params); } + async startTurn(params) { return this.record("startTurn", params); } + async steerTurn(params) { return this.record("steerTurn", params); } + async updateThreadSettings(params) { return this.record("updateThreadSettings", params); } + async listSessions(params = {}) { return this.record("listSessions", params); } + async selectSession(params) { return this.record("selectSession", params); } + async interruptTurn(params) { return this.record("interruptTurn", params); } + async sendInput(text, params = {}) { return this.record("sendInput", { text, ...params }); } + async cancel(taskId, params = {}) { return this.record("cancel", { taskId, ...params }); } + async respondApproval(requestId, decision, reason, response) { + return this.record("respondApproval", { requestId, decision, reason, response }); + } + async denyPending(reason) { this.calls.push(["denyPending", reason]); } + + async snapshot() { + this.calls.push(["snapshot"]); + return structuredClone(this.snapshotValue); + } + + onEvent(listener) { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + emit(event) { + for (const listener of this.listeners) listener(event); + } + + async dispose() { + this.calls.push(["dispose"]); + this.disposed = true; + } + + record(method, params) { + this.calls.push([method, params]); + return { adapter: this.name, method, params }; + } +} + +function idleSnapshot(name) { + return { + threadId: `${name}-thread`, + turnId: null, + state: "idle", + pendingApprovals: [], + pendingRequests: [], + outputTail: "", + metadata: { adapter: name }, + }; +} + +function deferred() { + let resolve; + let reject; + const promise = new Promise((yes, no) => { resolve = yes; reject = no; }); + return { promise, resolve, reject }; +} + +test("sync mode decorates snapshots and enforces VS Code-owned navigation", async () => { + const sync = new FakeAdapter("sync"); + const adapter = new SwitchableAgentAdapter({ initialMode: "sync", createAdapter: () => sync }); + await adapter.start(); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.adapter, "sync"); + assert.equal(snapshot.metadata.mode, "sync"); + assert.equal(snapshot.metadata.controlMode, "sync"); + assert.equal(snapshot.metadata.modeEpoch, 0); + assert.deepEqual(snapshot.metadata.capabilities, { + followsVscodeRoute: true, + sessionList: false, + sessionSelect: false, + sessionCreate: false, + threadSettings: true, + }); + + await assert.rejects(adapter.listSessions(), /unavailable in sync mode/); + await assert.rejects(adapter.selectSession({ threadId: "other" }), /unavailable in sync mode/); + await assert.rejects(adapter.newSession(), /unavailable in sync mode/); + await assert.rejects(adapter.startThread(), /unavailable in sync mode/); + assert.equal((await adapter.sendInput("hello")).adapter, "sync"); + assert.equal((await adapter.updateThreadSettings({ model: "codex" })).adapter, "sync"); + await adapter.dispose(); +}); + +test("async mode proxies the complete AgentAdapter surface", async () => { + const independent = new FakeAdapter("async"); + const adapter = new SwitchableAgentAdapter({ initialMode: "async", createAdapter: () => independent }); + await adapter.start(); + + await adapter.startThread({ cwd: "/workspace" }); + await adapter.newSession({ model: "codex" }); + await adapter.startTurn({ text: "start" }); + await adapter.steerTurn({ text: "steer" }); + await adapter.updateThreadSettings({ effort: "high" }); + await adapter.listSessions({ limit: 10 }); + await adapter.selectSession({ threadId: "thread-2" }); + await adapter.interruptTurn({ turnId: "turn-1" }); + await adapter.sendInput("input", { source: "web" }); + await adapter.cancel("turn-2", { reason: "user" }); + await adapter.respondApproval(7, "allow", "approved", { decision: "accept" }); + await adapter.denyPending("offline"); + + assert.deepEqual( + independent.calls.map(([method]) => method).filter((method) => !["start", "snapshot", "dispose"].includes(method)), + [ + "startThread", + "newSession", + "startTurn", + "steerTurn", + "updateThreadSettings", + "listSessions", + "selectSession", + "interruptTurn", + "sendInput", + "cancel", + "respondApproval", + "denyPending", + ], + ); + await adapter.dispose(); +}); + +test("session/new falls back to thread/start for a minimal async adapter", async () => { + const independent = new FakeAdapter("async"); + independent.newSession = undefined; + const adapter = new SwitchableAgentAdapter({ initialMode: "async", createAdapter: () => independent }); + await adapter.start(); + + const result = await adapter.newSession({ cwd: "/workspace" }); + assert.equal(result.method, "startThread"); + assert.equal((await adapter.snapshot()).metadata.capabilities.sessionCreate, true); + await adapter.dispose(); +}); + +test("mode switch commits atomically, buffers candidate events, and isolates the old generation", async () => { + const sync = new FakeAdapter("sync"); + const gate = deferred(); + const asyncAdapter = new FakeAdapter("async", { startGate: gate }); + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => mode === "sync" ? sync : asyncAdapter, + }); + const events = []; + adapter.onEvent((event) => events.push(`${event.type}:${event.payload.name ?? event.payload.controlMode ?? ""}`)); + await adapter.start(); + events.length = 0; + + const switching = adapter.setControlMode({ mode: "async" }); + await Promise.resolve(); + sync.emit({ type: "old.while-current", payload: { name: "sync" } }); + assert.deepEqual(events, ["old.while-current:sync"]); + await assert.rejects(adapter.sendInput("racing input"), /mode is switching/); + gate.resolve(); + + const result = await switching; + assert.deepEqual(result, { + changed: true, + controlMode: "async", + previousControlMode: "sync", + modeEpoch: 1, + }); + assert.equal(sync.disposed, true); + assert.equal(adapter.getControlMode(), "async"); + assert.ok(events.indexOf("control.mode.changed:async") < events.indexOf("candidate.starting:async")); + assert.ok(events.includes("connection.opened:async")); + + sync.emit({ type: "old.after-commit", payload: { name: "sync" } }); + asyncAdapter.emit({ type: "new.after-commit", payload: { name: "async" } }); + assert.equal(events.includes("old.after-commit:sync"), false); + assert.equal(events.includes("new.after-commit:async"), true); + + const snapshot = await adapter.snapshot(); + assert.equal(snapshot.metadata.modeEpoch, 1); + assert.deepEqual(snapshot.metadata.capabilities, { + followsVscodeRoute: false, + sessionList: true, + sessionSelect: true, + sessionCreate: true, + threadSettings: true, + }); + assert.equal((await adapter.listSessions()).adapter, "async"); + assert.equal((await adapter.newSession()).adapter, "async"); + await adapter.dispose(); +}); + +test("delegate snapshot events always carry authoritative mode metadata", async () => { + const sync = new FakeAdapter("sync"); + const asyncAdapter = new FakeAdapter("async"); + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => mode === "sync" ? sync : asyncAdapter, + }); + const snapshots = []; + adapter.onEvent((event) => { + if (event.type === "session.snapshot") snapshots.push(event.payload); + }); + await adapter.start(); + + sync.emit({ + type: "session.snapshot", + threadId: "sync-thread-2", + payload: { threadId: "sync-thread-2", metadata: { adapter: "sync", route: "/thread/2" } }, + }); + assert.deepEqual(snapshots.at(-1).metadata, { + adapter: "sync", + route: "/thread/2", + mode: "sync", + controlMode: "sync", + modeEpoch: 0, + capabilities: { + followsVscodeRoute: true, + sessionList: false, + sessionSelect: false, + sessionCreate: false, + threadSettings: true, + }, + }); + + await adapter.setControlMode({ mode: "async" }); + snapshots.length = 0; + asyncAdapter.emit({ + type: "session.snapshot", + threadId: "async-thread-2", + payload: { threadId: "async-thread-2", metadata: { adapter: "async", title: "Second" } }, + }); + assert.equal(snapshots.length, 1); + assert.equal(snapshots[0].metadata.adapter, "async"); + assert.equal(snapshots[0].metadata.title, "Second"); + assert.equal(snapshots[0].metadata.controlMode, "async"); + assert.equal(snapshots[0].metadata.modeEpoch, 1); + assert.equal(snapshots[0].metadata.capabilities.followsVscodeRoute, false); + assert.equal(snapshots[0].metadata.capabilities.sessionSelect, true); + await adapter.dispose(); +}); + +test("active turns and pending requests prevent a mode switch", async (t) => { + const cases = [ + ["active turn", { ...idleSnapshot("sync"), turnId: "turn-1", state: "active" }], + ["active state before a turn id arrives", { ...idleSnapshot("sync"), state: "in_progress" }], + ["active runtime flag", { ...idleSnapshot("sync"), activeFlags: ["thinking"] }], + ["pending approval", { + ...idleSnapshot("sync"), + pendingApprovals: [{ requestId: 1, method: "approval", action: "run", risk: "low", summary: "run", createdAt: 1, payload: {} }], + }], + ["pending input", { + ...idleSnapshot("sync"), + pendingRequests: [{ requestId: "input-1", method: "item/tool/requestUserInput" }], + }], + ]; + + for (const [name, snapshot] of cases) { + await t.test(name, async () => { + const sync = new FakeAdapter("sync", { snapshot }); + let factoryCalls = 0; + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => { + factoryCalls += 1; + return mode === "sync" ? sync : new FakeAdapter("async"); + }, + }); + await adapter.start(); + await assert.rejects(adapter.setControlMode({ mode: "async" }), /turn or request is active/); + assert.equal(factoryCalls, 1, "busy checks happen before creating a second adapter"); + assert.equal(adapter.getControlMode(), "sync"); + await adapter.dispose(); + }); + } +}); + +test("candidate startup failure leaves the old adapter authoritative", async () => { + const sync = new FakeAdapter("sync"); + const failed = new FakeAdapter("async", { startError: new Error("candidate failed") }); + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => mode === "sync" ? sync : failed, + }); + const events = []; + adapter.onEvent((event) => events.push(event.type)); + await adapter.start(); + events.length = 0; + + await assert.rejects(adapter.setControlMode({ controlMode: "async" }), /candidate failed/); + assert.equal(adapter.getControlMode(), "sync"); + assert.equal(failed.disposed, true); + assert.equal(sync.disposed, false); + assert.equal(events.includes("candidate.starting"), false, "failed candidate events stay private"); + assert.equal((await adapter.sendInput("still attached")).adapter, "sync"); + assert.equal((await adapter.snapshot()).metadata.modeEpoch, 0); + await adapter.dispose(); +}); + +test("a mode factory cannot reuse the currently active adapter instance", async () => { + const shared = new FakeAdapter("shared"); + const adapter = new SwitchableAgentAdapter({ initialMode: "sync", createAdapter: () => shared }); + await adapter.start(); + + await assert.rejects(adapter.setControlMode({ mode: "async" }), /must return a distinct adapter/); + assert.equal(adapter.getControlMode(), "sync"); + assert.equal(shared.disposed, false); + assert.equal((await adapter.sendInput("still live")).adapter, "shared"); + await adapter.dispose(); +}); + +test("listener failures cannot turn a committed switch into a rejected command", async () => { + const sync = new FakeAdapter("sync"); + const asyncAdapter = new FakeAdapter("async"); + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => mode === "sync" ? sync : asyncAdapter, + }); + adapter.onEvent(() => { throw new Error("consumer failed"); }); + await adapter.start(); + + const result = await adapter.setControlMode({ mode: "async" }); + assert.equal(result.changed, true); + assert.equal(adapter.getControlMode(), "async"); + assert.equal(sync.disposed, true); + await adapter.dispose(); +}); + +test("a turn that appears while the candidate starts aborts before commit", async () => { + const sync = new FakeAdapter("sync"); + const gate = deferred(); + const candidate = new FakeAdapter("async", { startGate: gate }); + const adapter = new SwitchableAgentAdapter({ + initialMode: "sync", + createAdapter: (mode) => mode === "sync" ? sync : candidate, + }); + await adapter.start(); + + const switching = adapter.setControlMode({ mode: "async" }); + await Promise.resolve(); + sync.snapshotValue.turnId = "turn-race"; + sync.snapshotValue.state = "active"; + gate.resolve(); + + await assert.rejects(switching, /turn or request is active/); + assert.equal(adapter.getControlMode(), "sync"); + assert.equal(candidate.disposed, true); + assert.equal(sync.disposed, false); + sync.snapshotValue.turnId = null; + sync.snapshotValue.state = "idle"; + await adapter.dispose(); +}); + +test("control mode validation and idempotent switches are explicit", async () => { + const sync = new FakeAdapter("sync"); + const adapter = new SwitchableAgentAdapter({ initialMode: "sync", createAdapter: () => sync }); + await adapter.start(); + + await assert.rejects(adapter.setControlMode({ mode: "attach" }), /must be sync or async/); + assert.deepEqual(await adapter.setControlMode({ mode: "sync" }), { + changed: false, + controlMode: "sync", + previousControlMode: "sync", + modeEpoch: 0, + }); + await adapter.dispose(); +}); diff --git a/aether-vscodex/vscode-extension/.gitignore b/aether-vscodex/vscode-extension/.gitignore new file mode 100644 index 000000000..a08e1da2d --- /dev/null +++ b/aether-vscodex/vscode-extension/.gitignore @@ -0,0 +1,3 @@ +node_modules/ +dist/ +*.vsix diff --git a/aether-vscodex/vscode-extension/.vscodeignore b/aether-vscodex/vscode-extension/.vscodeignore new file mode 100644 index 000000000..26b7df07e --- /dev/null +++ b/aether-vscodex/vscode-extension/.vscodeignore @@ -0,0 +1,9 @@ +src/** +.gitignore +tsconfig.json +**/*.map +node_modules/@types/** +node_modules/typescript/** +node_modules/.package-lock.json +*.tsbuildinfo +*.vsix diff --git a/aether-vscodex/vscode-extension/LICENSE b/aether-vscodex/vscode-extension/LICENSE new file mode 100644 index 000000000..f0d01c030 --- /dev/null +++ b/aether-vscodex/vscode-extension/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Codex Remote Collaboration contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/aether-vscodex/vscode-extension/README.md b/aether-vscodex/vscode-extension/README.md new file mode 100644 index 000000000..a2d479a00 --- /dev/null +++ b/aether-vscodex/vscode-extension/README.md @@ -0,0 +1,237 @@ +# Codex Remote Collaboration VS Code Bridge + +This extension connects local and Aether relay channels to one switchable Codex +control host. **Synchronous mode** follows the conversation currently shown by +the official Codex VS Code extension through its private IPC protocol and does +not spawn a `codex` process. **Asynchronous mode** starts an independent +app-server and lets the Web UI list, resume, create, and select conversations. + +The attached conversation remains visible and usable in the official Codex +panel. Remote operators can observe its output, submit a new turn or steer the +active turn, interrupt it, and answer supported approval/input requests. + +The mode can be changed from the Web UI without reconnecting either relay. +Synchronous mode makes the official panel the only conversation-navigation +owner; asynchronous mode restores the browser history and new-conversation +actions. A running turn or pending request blocks mode changes. + +## Requirements + +- The official `openai.chatgpt` VS Code extension is installed and signed in. +- The target Codex conversation is open and owned by that extension. +- The bridge and official extension run as the same OS user. The default Unix + socket is `$CODEX_HOME/ipc/ipc.sock`, normally `~/.codex/ipc/ipc.sock`. +- For a loopback `ws://` URL, the extension starts and owns its bundled relay + automatically. Remote and `wss://` relay URLs remain externally hosted. + +The IPC follower protocol is private and versioned, not a public OpenAI API. +An official extension update can require a compatible bridge update. Strict +stream-version checks are enabled by default so an unknown protocol fails +closed instead of being interpreted optimistically. + +## Build and install + +```sh +npm install +npm run check +npm run build +npx --yes @vscode/vsce package +code --install-extension codex-remote-collab-0.4.0.vsix --force +``` + +Run **Developer: Reload Window** after installing or replacing the VSIX. + +## Configure control modes + +For the local default, no separate relay command is required. The extension +starts the bundled relay on the host and port from `codexRemoteCollab.localRelayUrl`. +To run the development relay manually, disable +`codexRemoteCollab.autoStartLocalRelay` and use: + +```sh +HOST=127.0.0.1 PORT=8787 CODEX_REMOTE_MODE=host npm start +``` + +To opt into authentication later, set `CODEX_REMOTE_AUTH=required` and the three +token variables before starting the relay. + +Set the extension configuration: + +```json +{ + "codexRemoteCollab.localRelayUrl": "ws://127.0.0.1:8787/v1/connect", + "codexRemoteCollab.controlMode": "sync", + "codexRemoteCollab.autoDiscoverThread": true, + "codexRemoteCollab.autoStart": true +} +``` + +Then: + +1. Open the target conversation in the official Codex panel. +2. Reload VS Code once after installing the companion extension. The local relay and bridge start automatically; no token is needed for loopback. The status item opens the Web console and is not a connect/disconnect toggle. +3. Open the relay web console; it connects automatically on localhost. The web UI uses a + Codex-style conversation stream with a bottom composer; Enter sends and Shift+Enter + inserts a newline. There is no separate connect/disconnect step for the local relay. + +If the browser says that it is waiting for the VS Code host or the recent-session list is +empty, verify that `codexRemoteCollab.localRelayUrl` uses the same port as the relay and run +**Developer: Reload Window**. Keep `codexRemoteCollab.threadId` empty unless a specific +conversation must be pinned; an old closed ID can prevent startup until it is cleared. + +When authentication is enabled, run **Codex Remote: Set Relay Token** with the +host token. It is stored in `vscode.SecretStorage`, not in settings; the browser +uses the operator or viewer token separately. + +With no configured thread ID, the bridge ranks recent VS Code rollout metadata +and shows only candidates verified by live IPC owner discovery and a matching +follower snapshot. Explicit Codex Desktop tasks, closed, stale, and other +non-attachable history entries are omitted. +In synchronous mode, switching the conversation in the official Codex panel +also switches the Web projection after the new owner snapshot is ready. The +Web UI cannot list, select, or create conversations in this mode. Switch to +asynchronous mode when the browser should own conversation navigation. +To avoid ambiguity when several Codex windows are open, run **Codex Remote: Set Existing Thread ID**. +An empty value restores automatic discovery. + +Useful commands: + +- **Codex Remote: Start Bridge** / **Stop Bridge** +- **Codex Remote: Set Existing Thread ID** +- **Codex Remote: Set Relay Token** +- **Codex Remote: Pair with Aether** +- **Codex Remote: Configure Aether Cloud Relay** +- **Codex Remote: Send Input** +- **Codex Remote: Show Snapshot** + +## Settings + +| Setting | Default | Meaning | +| --- | --- | --- | +| `codexRemoteCollab.controlMode` | `sync` | `sync` follows VS Code; `async` owns an independent app-server. | +| `codexRemoteCollab.localRelayUrl` | `ws://127.0.0.1:8787/v1/connect` | Bundled loopback relay used by the local Web control. | +| `codexRemoteCollab.aetherUrl` | empty | Aether origin remembered by the pairing command. | +| `codexRemoteCollab.cloudRelayUrl` | empty | Aether WebSocket relay URL populated by pairing. | +| `codexRemoteCollab.threadId` | empty | Exact existing conversation ID; empty enables discovery. | +| `codexRemoteCollab.autoDiscoverThread` | `true` | Discover and owner-check a local VS Code session. | +| `codexRemoteCollab.followVscodeSession` | `true` | Legacy compatibility setting; synchronous mode always follows VS Code. | +| `codexRemoteCollab.ipcSocketPath` | empty | Override the local IPC socket path. | +| `codexRemoteCollab.hostId` | `local` | Owner-discovery host identifier. | +| `codexRemoteCollab.ipcStrictVersions` | `true` | Reject unsupported stream protocol versions. | +| `codexRemoteCollab.approvalTimeoutMs` | `300000` | Deny an unanswered request locally after this delay. | +| `codexRemoteCollab.allowHighRiskApprovals` | `false` | Permit remote high-risk approvals when explicitly enabled. | + +`codexRemoteCollab.codexCommand`, `codexArgs`, and `defaultCwd` apply only to +asynchronous mode. The deprecated `mode=attach/spawn` values map to +`controlMode=sync/async` when no explicit control mode exists. + +## Pair with Aether + +The local relay stays enabled after cloud pairing. In Aether, open **Codex remote +control** and generate a one-time code. Then run **Codex Remote: Pair with Aether** +from the VS Code Command Palette, enter the Aether server URL and the code, and +the bridge will connect to both relays. The long-lived device credential is stored +only in VS Code SecretStorage. Revoke a lost or retired device from the Aether page. + +## Relay behavior + +The bridge sends a `hello` and, when a relay token is configured, a separate +bearer-auth frame over an outbound WebSocket. It publishes normalized events including: + +- `connection.opened` / `connection.closed` +- `session.snapshot` +- `output.snapshot` / `output.chunk` +- `task.started` / `task.finished` / `task.cancelled` +- `approval.requested` / `approval.resolved` / `approval.expired` +- `input.requested` / `input.resolved` / `input.expired` + +Remote commands are mapped to the existing conversation owner: + +- `control/mode/set` atomically switches between `sync` and `async`. +- `session/list`, `session/select`, and `session/new` are available only in + asynchronous mode and map to `thread/list`, `thread/resume`, and `thread/start`. +- `turn/start` starts a turn in the attached thread. +- `turn/steer` adds input to the active turn. +- `turn/interrupt` interrupts the expected active turn. +- `approval.respond`, `input.respond`, and `server.request.respond` preserve the + original request ID and use method-specific follower responses. +- `thread/start` is deliberately rejected in synchronous mode because VS Code + owns conversation navigation there. + +The browser never connects directly to the IPC socket. Relay and host both +enforce role/capability checks; high-risk command approval remains disabled +unless the local VS Code setting opts in. + +## Supported follower requests + +- `item/commandExecution/requestApproval` +- `item/fileChange/requestApproval` +- `item/permissions/requestApproval` +- `item/tool/requestUserInput` +- `mcpServer/elicitation/request` +- legacy `applyPatchApproval` and `execCommandApproval` + +Unanswered requests expire with a local deny. JSON-RPC numeric and string IDs +remain distinct, and a response can be submitted only once. + +## Legacy mode migration + +The old setting remains accepted: + +```json +{ + "codexRemoteCollab.mode": "spawn", + "codexRemoteCollab.codexCommand": "/absolute/path/to/codex", + "codexRemoteCollab.codexArgs": ["app-server", "--stdio"] +} +``` + +It maps to `controlMode=async`. Prefer the new setting directly. A +`spawn codex ENOENT` error belongs only to asynchronous mode; it is not a +synchronous-mode prerequisite or a PATH problem that needs fixing for +existing-session control. + +The standalone `npm run start:stdio` entry point and `createBridge()` helper +also retain the legacy app-server adapter for compatibility. + +## Embedding the attach adapter + +The reusable exports are in `src/index.ts`: + +```ts +import { + CodexIpcAgentAdapter, + RelayClient, + RelayHost, +} from "codex-remote-collab"; + +const adapter = new CodexIpcAgentAdapter({ + threadId: process.env.CODEX_THREAD_ID, + autoDiscoverThread: true, +}); +const relay = new RelayClient({ + url: "wss://relay.example.test/v1/connect", + accessToken: process.env.CODEX_REMOTE_HOST_TOKEN, +}); +const host = new RelayHost({ adapter, relay }); +await host.start(); +``` + +`CodexIpcClient` is exported separately for protocol fixtures and diagnostics. +Use `followConversation()` before follower mutations, and always target the +owner returned by `findThreadOwner()`. + +## Troubleshooting + +- **No existing session found:** open the target official Codex conversation, + keep that VS Code window running, then retry or set its exact thread ID. +- **Owner not found:** the rollout exists on disk but no live official client + currently owns it. Reopen the conversation in the Codex panel. +- **IPC version mismatch:** update this bridge for the installed official + extension. Disabling strict versions is diagnostic only. +- **Relay stays at waiting for host:** confirm host mode, relay URL, and that no + second host is already connected. If authentication is enabled, also check the + host token. +- **Old `spawn codex ENOENT` message:** install version `0.4.0`, reload VS Code, + and verify `codexRemoteCollab.controlMode` is `sync` unless independent + conversations are intended. diff --git a/aether-vscodex/vscode-extension/l10n/bundle.l10n.json b/aether-vscodex/vscode-extension/l10n/bundle.l10n.json new file mode 100644 index 000000000..6aa35d6db --- /dev/null +++ b/aether-vscodex/vscode-extension/l10n/bundle.l10n.json @@ -0,0 +1,56 @@ +{ + "A non-empty Aether device credential is required.": "A non-empty Aether device credential is required.", + "Aether cloud connection removed. Local control remains enabled.": "Aether cloud connection removed. Local control remains enabled.", + "Aether cloud connection saved. Restart the Codex Remote bridge to connect; local control remains available.": "Aether cloud connection saved. Restart the Codex Remote bridge to connect; local control remains available.", + "Aether cloud relay WebSocket URL": "Aether cloud relay WebSocket URL", + "Aether pairing completed. Local and cloud control are both active.": "Aether pairing completed. Local and cloud control are both active.", + "Aether pairing was saved, but the cloud connection is currently unavailable. Local control remains active and the cloud connection will retry.": "Aether pairing was saved, but the cloud connection is currently unavailable. Local control remains active and the cloud connection will retry.", + "Aether returned an invalid pairing response.": "Aether returned an invalid pairing response.", + "Aether server URL": "Aether server URL", + "Attached to the existing Codex conversation. Click to open the web control.": "Attached to the existing Codex conversation. Click to open the web control.", + "Bridge connected. Click to open the web control.": "Bridge connected. Click to open the web control.", + "Bridge paused. Click to open the web control and resume automatically.": "Bridge paused. Click to open the web control and resume automatically.", + "Codex Remote Collaboration": "Codex Remote Collaboration", + "Codex Remote will attach to {0} after the next bridge start.": "Codex Remote will attach to {0} after the next bridge start.", + "Codex Remote will auto-discover the latest VS Code Codex conversation after the next bridge start.": "Codex Remote will auto-discover the latest VS Code Codex conversation after the next bridge start.", + "Connecting to the local Codex collaboration service": "Connecting to the local Codex collaboration service", + "Device credential from the Aether pairing flow": "Device credential from the Aether pairing flow", + "Enter a valid URL.": "Enter a valid URL.", + "Enter a valid WebSocket URL.": "Enter a valid WebSocket URL.", + "Enter the 8-character pairing code.": "Enter the 8-character pairing code.", + "Enter the Aether server URL.": "Enter the Aether server URL.", + "Existing Codex conversation ID (leave blank for auto-discovery)": "Existing Codex conversation ID (leave blank for auto-discovery)", + "Independent Codex mode is connected. Click to open the web control.": "Independent Codex mode is connected. Click to open the web control.", + "One-time pairing code shown in Aether": "One-time pairing code shown in Aether", + "Relay access token (leave blank for the local relay)": "Relay access token (leave blank for the local relay)", + "Relay token stored in VS Code SecretStorage.": "Relay token stored in VS Code SecretStorage.", + "Remote Aether connections must use wss://.": "Remote Aether connections must use wss://.", + "Remote Aether servers must use https://.": "Remote Aether servers must use https://.", + "Restoring the local collaboration service": "Restoring the local collaboration service", + "Send input to the active Codex turn": "Send input to the active Codex turn", + "Set codexRemoteCollab.localRelayUrl before starting the bridge.": "Set codexRemoteCollab.localRelayUrl before starting the bridge.", + "Start the Codex remote bridge first.": "Start the Codex remote bridge first.", + "Starting the independent Codex mode.": "Starting the independent Codex mode.", + "Starting {0}": "Starting {0}", + "The Codex conversation is not connected": "The Codex conversation is not connected", + "The Codex executable is unavailable": "The Codex executable is unavailable", + "The Codex remote bridge attached to the existing VS Code Codex conversation.": "The Codex remote bridge attached to the existing VS Code Codex conversation.", + "The Codex remote bridge is already running.": "The Codex remote bridge is already running.", + "The Codex remote collaboration bridge connected.": "The Codex remote collaboration bridge connected.", + "The independent Codex mode is not connected": "The independent Codex mode is not connected", + "The independent Codex remote mode connected.": "The independent Codex remote mode connected.", + "The bridge is not connected": "The bridge is not connected", + "The local collaboration URL is invalid. Check codexRemoteCollab.localRelayUrl.": "The local collaboration URL is invalid. Check codexRemoteCollab.localRelayUrl.", + "The local collaboration service at {0} is temporarily unavailable. The extension will keep retrying.": "The local collaboration service at {0} is temporarily unavailable. The extension will keep retrying.", + "The official Codex extension new-conversation command was not found. Make sure the VS Code Codex extension is enabled.": "The official Codex extension new-conversation command was not found. Make sure the VS Code Codex extension is enabled.", + "Unable to pair with Aether: {0}": "Unable to pair with Aether: {0}", + "Unable to restore the local collaboration service": "Unable to restore the local collaboration service", + "Unable to send Codex input: {0}": "Unable to send Codex input: {0}", + "Unable to start the Codex remote bridge: {0}": "Unable to start the Codex remote bridge: {0}", + "Unable to start the local Codex collaboration service: {0}": "Unable to start the local Codex collaboration service: {0}", + "Unable to start the local collaboration service: {0}": "Unable to start the local collaboration service: {0}", + "Use a ws:// or wss:// URL.": "Use a ws:// or wss:// URL.", + "Use the Aether origin without credentials, a query, or a fragment.": "Use the Aether origin without credentials, a query, or a fragment.", + "Waiting for a Codex conversation to open in VS Code. It will connect automatically.": "Waiting for a Codex conversation to open in VS Code. It will connect automatically.", + "codexRemoteCollab.localRelayUrl must be a loopback ws:// address.": "codexRemoteCollab.localRelayUrl must be a loopback ws:// address." +} diff --git a/aether-vscodex/vscode-extension/l10n/bundle.l10n.zh-cn.json b/aether-vscodex/vscode-extension/l10n/bundle.l10n.zh-cn.json new file mode 100644 index 000000000..6a2cefdc1 --- /dev/null +++ b/aether-vscodex/vscode-extension/l10n/bundle.l10n.zh-cn.json @@ -0,0 +1,56 @@ +{ + "A non-empty Aether device credential is required.": "必须填写 Aether 设备凭据。", + "Aether cloud connection removed. Local control remains enabled.": "已移除 Aether 云端连接,本地控制仍然可用。", + "Aether cloud connection saved. Restart the Codex Remote bridge to connect; local control remains available.": "已保存 Aether 云端连接。重启 Codex Remote 桥接后即可连接,本地控制仍然可用。", + "Aether cloud relay WebSocket URL": "Aether 云端 relay WebSocket 地址", + "Aether pairing completed. Local and cloud control are both active.": "Aether 配对完成,本地与云端控制均已启用。", + "Aether pairing was saved, but the cloud connection is currently unavailable. Local control remains active and the cloud connection will retry.": "Aether 配对信息已保存,但当前无法连接云端。本地控制仍然可用,云端连接会继续重试。", + "Aether returned an invalid pairing response.": "Aether 返回了无效的配对响应。", + "Aether server URL": "Aether 服务器地址", + "Attached to the existing Codex conversation. Click to open the web control.": "已附加到现有 Codex 会话,点击打开 Web 控制页。", + "Bridge connected. Click to open the web control.": "桥接已连接,点击打开 Web 控制页。", + "Bridge paused. Click to open the web control and resume automatically.": "桥接已暂停,点击打开 Web 控制页时会自动恢复。", + "Codex Remote Collaboration": "Codex 远程协同", + "Codex Remote will attach to {0} after the next bridge start.": "Codex Remote 将在下次启动桥接后附加到 {0}。", + "Codex Remote will auto-discover the latest VS Code Codex conversation after the next bridge start.": "Codex Remote 将在下次启动桥接后自动发现最新的 VS Code Codex 会话。", + "Connecting to the local Codex collaboration service": "正在连接本地 Codex 协同服务", + "Device credential from the Aether pairing flow": "Aether 配对流程生成的设备凭据", + "Enter a valid URL.": "请输入有效的 URL。", + "Enter a valid WebSocket URL.": "请输入有效的 WebSocket URL。", + "Enter the 8-character pairing code.": "请输入 8 位配对码。", + "Enter the Aether server URL.": "请输入 Aether 服务器地址。", + "Existing Codex conversation ID (leave blank for auto-discovery)": "现有 Codex 会话 ID(留空则自动发现)", + "Independent Codex mode is connected. Click to open the web control.": "独立 Codex 模式已连接,点击打开 Web 控制页。", + "One-time pairing code shown in Aether": "Aether 中显示的一次性配对码", + "Relay access token (leave blank for the local relay)": "Relay 访问 token(本地 relay 请留空)", + "Relay token stored in VS Code SecretStorage.": "Relay token 已保存到 VS Code SecretStorage。", + "Remote Aether connections must use wss://.": "远程 Aether 连接必须使用 wss://。", + "Remote Aether servers must use https://.": "远程 Aether 服务器必须使用 https://。", + "Restoring the local collaboration service": "正在恢复本地协同服务", + "Send input to the active Codex turn": "向当前 Codex turn 发送输入", + "Set codexRemoteCollab.localRelayUrl before starting the bridge.": "请先设置 codexRemoteCollab.localRelayUrl,再启动桥接。", + "Start the Codex remote bridge first.": "请先启动 Codex 远程桥接。", + "Starting the independent Codex mode.": "正在启动独立 Codex 模式。", + "Starting {0}": "正在启动 {0}", + "The Codex conversation is not connected": "Codex 会话尚未连接", + "The Codex executable is unavailable": "Codex 可执行文件不可用", + "The Codex remote bridge attached to the existing VS Code Codex conversation.": "Codex 远程桥接已附加到现有 VS Code Codex 会话。", + "The Codex remote bridge is already running.": "Codex 远程桥接已在运行。", + "The Codex remote collaboration bridge connected.": "Codex 远程协同桥接已连接。", + "The independent Codex mode is not connected": "独立 Codex 模式尚未连接", + "The independent Codex remote mode connected.": "独立 Codex 远程模式已连接。", + "The bridge is not connected": "桥接尚未连接", + "The local collaboration URL is invalid. Check codexRemoteCollab.localRelayUrl.": "本地协同地址无效,请检查 codexRemoteCollab.localRelayUrl。", + "The local collaboration service at {0} is temporarily unavailable. The extension will keep retrying.": "本地协同服务 {0} 暂时无法连接,扩展会继续重试。", + "The official Codex extension new-conversation command was not found. Make sure the VS Code Codex extension is enabled.": "未找到官方 Codex 扩展的新会话命令,请确认 VS Code Codex 扩展已启用。", + "Unable to pair with Aether: {0}": "无法与 Aether 配对:{0}", + "Unable to restore the local collaboration service": "无法恢复本地协同服务", + "Unable to send Codex input: {0}": "无法发送 Codex 输入:{0}", + "Unable to start the Codex remote bridge: {0}": "无法启动 Codex 远程桥接:{0}", + "Unable to start the local Codex collaboration service: {0}": "无法启动本地 Codex 协同服务:{0}", + "Unable to start the local collaboration service: {0}": "无法启动本地协同服务:{0}", + "Use a ws:// or wss:// URL.": "请使用 ws:// 或 wss:// URL。", + "Use the Aether origin without credentials, a query, or a fragment.": "请填写不含凭据、查询参数或片段的 Aether 源地址。", + "Waiting for a Codex conversation to open in VS Code. It will connect automatically.": "正在等待 VS Code 中打开 Codex 会话,检测到后会自动连接。", + "codexRemoteCollab.localRelayUrl must be a loopback ws:// address.": "codexRemoteCollab.localRelayUrl 必须是回环地址上的 ws:// URL。" +} diff --git a/aether-vscodex/vscode-extension/package-lock.json b/aether-vscodex/vscode-extension/package-lock.json new file mode 100644 index 000000000..e746545ed --- /dev/null +++ b/aether-vscodex/vscode-extension/package-lock.json @@ -0,0 +1,94 @@ +{ + "name": "codex-remote-collab", + "version": "0.4.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "codex-remote-collab", + "version": "0.4.0", + "license": "MIT", + "dependencies": { + "ws": "^8.18.0" + }, + "devDependencies": { + "@types/node": "^20.14.0", + "@types/vscode": "^1.85.0", + "@types/ws": "^8.5.12", + "typescript": "^5.4.5" + }, + "engines": { + "vscode": "^1.85.0" + } + }, + "node_modules/@types/node": { + "version": "20.19.43", + "resolved": "https://registry.npmjs.org/@types/node/-/node-20.19.43.tgz", + "integrity": "sha512-6oYBAi5ikg4Pl+kGsoYtawUMBT2zZMCvPNF7pVLnHZfd1zf38DRiWn/gT01RYCdUqkv7Fhr+C9ot4/tb+2sVvA==", + "dev": true, + "license": "MIT", + "dependencies": { + "undici-types": "~6.21.0" + } + }, + "node_modules/@types/vscode": { + "version": "1.134.0", + "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.134.0.tgz", + "integrity": "sha512-NDEu0hg4sF7+vvFsADsktqUJ6f80LHSZvVK2Ovo1XiQ0/VHck1O3zst+ZZyVA/uvz6vo6LcuoqU2q48YMqOwWw==", + "dev": true, + "license": "MIT" + }, + "node_modules/@types/ws": { + "version": "8.18.1", + "resolved": "https://registry.npmjs.org/@types/ws/-/ws-8.18.1.tgz", + "integrity": "sha512-ThVF6DCVhA8kUGy+aazFQ4kXQ7E1Ty7A3ypFOe0IcJV8O/M511G99AW24irKrW56Wt44yG9+ij8FaqoBGkuBXg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/node": "*" + } + }, + "node_modules/typescript": { + "version": "5.9.3", + "resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz", + "integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==", + "dev": true, + "license": "Apache-2.0", + "bin": { + "tsc": "bin/tsc", + "tsserver": "bin/tsserver" + }, + "engines": { + "node": ">=14.17" + } + }, + "node_modules/undici-types": { + "version": "6.21.0", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", + "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/ws": { + "version": "8.21.3", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.3.tgz", + "integrity": "sha512-201TZ/kPWxoPr/OKWjquZR1SWKXcvxdH+e1xrx89b3YbmzLMFCLfnaG1HFIgWzJOEWZ7MvpK++odZufgYR50Rw==", + "license": "MIT", + "engines": { + "node": ">=10.0.0" + }, + "peerDependencies": { + "bufferutil": "^4.0.1", + "utf-8-validate": ">=5.0.2" + }, + "peerDependenciesMeta": { + "bufferutil": { + "optional": true + }, + "utf-8-validate": { + "optional": true + } + } + } + } +} diff --git a/aether-vscodex/vscode-extension/package.json b/aether-vscodex/vscode-extension/package.json new file mode 100644 index 000000000..9cf9e1e7a --- /dev/null +++ b/aether-vscodex/vscode-extension/package.json @@ -0,0 +1,211 @@ +{ + "name": "codex-remote-collab", + "displayName": "%extension.displayName%", + "description": "%extension.description%", + "version": "0.4.0", + "publisher": "local", + "license": "MIT", + "engines": { + "vscode": "^1.85.0" + }, + "categories": [ + "Other" + ], + "l10n": "./l10n", + "activationEvents": [ + "onStartupFinished", + "onCommand:codexRemoteCollab.openWeb", + "onCommand:codexRemoteCollab.start", + "onCommand:codexRemoteCollab.stop", + "onCommand:codexRemoteCollab.setThreadId", + "onCommand:codexRemoteCollab.sendInput", + "onCommand:codexRemoteCollab.setRelayToken", + "onCommand:codexRemoteCollab.configureCloud", + "onCommand:codexRemoteCollab.pairCloud", + "onCommand:codexRemoteCollab.snapshot" + ], + "main": "./dist/extension.js", + "contributes": { + "commands": [ + { + "command": "codexRemoteCollab.openWeb", + "title": "%command.openWeb%" + }, + { + "command": "codexRemoteCollab.start", + "title": "%command.start%" + }, + { + "command": "codexRemoteCollab.stop", + "title": "%command.stop%" + }, + { + "command": "codexRemoteCollab.setThreadId", + "title": "%command.setThreadId%" + }, + { + "command": "codexRemoteCollab.sendInput", + "title": "%command.sendInput%" + }, + { + "command": "codexRemoteCollab.setRelayToken", + "title": "%command.setRelayToken%" + }, + { + "command": "codexRemoteCollab.configureCloud", + "title": "%command.configureCloud%" + }, + { + "command": "codexRemoteCollab.pairCloud", + "title": "%command.pairCloud%" + }, + { + "command": "codexRemoteCollab.snapshot", + "title": "%command.snapshot%" + } + ], + "configuration": { + "title": "%configuration.title%", + "properties": { + "codexRemoteCollab.localRelayUrl": { + "type": "string", + "default": "ws://127.0.0.1:8787/v1/connect", + "description": "%configuration.localRelayUrl%" + }, + "codexRemoteCollab.relayUrl": { + "type": "string", + "default": "ws://127.0.0.1:8787/v1/connect", + "description": "%configuration.relayUrl%", + "deprecationMessage": "%configuration.relayUrl.deprecation%" + }, + "codexRemoteCollab.cloudRelayUrl": { + "type": "string", + "default": "", + "description": "%configuration.cloudRelayUrl%" + }, + "codexRemoteCollab.aetherUrl": { + "type": "string", + "default": "", + "description": "%configuration.aetherUrl%" + }, + "codexRemoteCollab.autoStart": { + "type": "boolean", + "default": true, + "description": "%configuration.autoStart%" + }, + "codexRemoteCollab.autoStartLocalRelay": { + "type": "boolean", + "default": true, + "description": "%configuration.autoStartLocalRelay%" + }, + "codexRemoteCollab.mode": { + "type": "string", + "enum": [ + "attach", + "spawn" + ], + "default": "attach", + "description": "%configuration.mode%", + "deprecationMessage": "%configuration.mode.deprecation%" + }, + "codexRemoteCollab.controlMode": { + "type": "string", + "enum": [ + "sync", + "async" + ], + "enumDescriptions": [ + "%configuration.controlMode.sync%", + "%configuration.controlMode.async%" + ], + "default": "sync", + "description": "%configuration.controlMode%" + }, + "codexRemoteCollab.threadId": { + "type": "string", + "default": "", + "description": "%configuration.threadId%" + }, + "codexRemoteCollab.autoDiscoverThread": { + "type": "boolean", + "default": true, + "description": "%configuration.autoDiscoverThread%" + }, + "codexRemoteCollab.followVscodeSession": { + "type": "boolean", + "default": true, + "description": "%configuration.followVscodeSession%" + }, + "codexRemoteCollab.ipcSocketPath": { + "type": "string", + "default": "", + "description": "%configuration.ipcSocketPath%" + }, + "codexRemoteCollab.hostId": { + "type": "string", + "default": "local", + "description": "%configuration.hostId%" + }, + "codexRemoteCollab.ipcStrictVersions": { + "type": "boolean", + "default": true, + "description": "%configuration.ipcStrictVersions%" + }, + "codexRemoteCollab.codexCommand": { + "type": "string", + "default": "codex", + "description": "%configuration.codexCommand%" + }, + "codexRemoteCollab.codexArgs": { + "type": "array", + "items": { + "type": "string" + }, + "default": [ + "app-server", + "--stdio" + ], + "description": "%configuration.codexArgs%" + }, + "codexRemoteCollab.defaultCwd": { + "type": "string", + "default": "", + "description": "%configuration.defaultCwd%" + }, + "codexRemoteCollab.approvalTimeoutMs": { + "type": "number", + "default": 300000, + "minimum": 1000, + "description": "%configuration.approvalTimeoutMs%" + }, + "codexRemoteCollab.allowHighRiskApprovals": { + "type": "boolean", + "default": false, + "description": "%configuration.allowHighRiskApprovals%" + }, + "codexRemoteCollab.relayReconnect": { + "type": "boolean", + "default": true, + "description": "%configuration.relayReconnect%" + } + } + } + }, + "scripts": { + "vscode:prepublish": "npm run build:web && npm run build", + "build:web": "npm --prefix ../web run build", + "build": "tsc -p tsconfig.json && node scripts/sync-local-relay.cjs", + "compile": "npm run build", + "check": "tsc --noEmit -p tsconfig.json", + "start:stdio": "node dist/cli.js" + }, + "dependencies": { + "ws": "^8.18.0" + }, + "devDependencies": { + "@types/node": "^20.14.0", + "@types/vscode": "^1.85.0", + "@types/ws": "^8.5.12", + "typescript": "^5.4.5" + } +} diff --git a/aether-vscodex/vscode-extension/package.nls.json b/aether-vscodex/vscode-extension/package.nls.json new file mode 100644 index 000000000..03c11081e --- /dev/null +++ b/aether-vscodex/vscode-extension/package.nls.json @@ -0,0 +1,38 @@ +{ + "extension.displayName": "Codex Remote Collaboration", + "extension.description": "Synchronize the current VS Code Codex conversation or manage independent Codex conversations from a local browser and Aether cloud.", + "command.openWeb": "Codex Remote: Open Local Web Console", + "command.start": "Codex Remote: Start Bridge", + "command.stop": "Codex Remote: Stop Bridge", + "command.setThreadId": "Codex Remote: Set Existing Thread ID", + "command.sendInput": "Codex Remote: Send Input", + "command.setRelayToken": "Codex Remote: Set Local Relay Token", + "command.configureCloud": "Codex Remote: Configure Aether Cloud Manually", + "command.pairCloud": "Codex Remote: Pair with Aether", + "command.snapshot": "Codex Remote: Show Snapshot", + "configuration.title": "Codex Remote Collaboration", + "configuration.localRelayUrl": "Loopback relay used by the local browser UI. It remains active when Aether cloud sync is enabled.", + "configuration.relayUrl": "Legacy relay setting retained for compatibility. Use localRelayUrl and cloudRelayUrl for new installations.", + "configuration.relayUrl.deprecation": "Use codexRemoteCollab.localRelayUrl for local access and codexRemoteCollab.cloudRelayUrl for Aether cloud access.", + "configuration.cloudRelayUrl": "Optional Aether cloud relay WebSocket URL. The device credential is stored separately in VS Code SecretStorage.", + "configuration.aetherUrl": "Aether server origin used by the one-time pairing flow.", + "configuration.autoStart": "Start the bridge when the extension activates.", + "configuration.autoStartLocalRelay": "Automatically host the bundled relay for loopback ws:// URLs.", + "configuration.mode": "Attach to the existing official VS Code Codex session, or spawn a separate app-server for legacy use.", + "configuration.mode.deprecation": "Use codexRemoteCollab.controlMode. attach maps to sync and spawn maps to async.", + "configuration.controlMode": "Choose whether the web console follows the current VS Code Codex conversation or manages independent conversations.", + "configuration.controlMode.sync": "Synchronize with the conversation currently shown in the official VS Code Codex panel.", + "configuration.controlMode.async": "Run an independent Codex app-server and manage its conversations from the web console.", + "configuration.threadId": "Existing VS Code Codex conversation ID to follow. Empty uses the most recent locally available session.", + "configuration.autoDiscoverThread": "Discover a recent VS Code Codex conversation when no thread ID is configured.", + "configuration.followVscodeSession": "Follow conversation changes in the attached official VS Code Codex panel.", + "configuration.ipcSocketPath": "Optional official Codex IPC socket path. Empty uses CODEX_HOME/ipc/ipc.sock.", + "configuration.hostId": "Codex host identifier used for existing-session discovery.", + "configuration.ipcStrictVersions": "Reject unknown private IPC stream versions instead of applying them optimistically.", + "configuration.codexCommand": "Asynchronous mode: Codex executable used to launch the independent app-server.", + "configuration.codexArgs": "Asynchronous mode: arguments passed to the Codex executable.", + "configuration.defaultCwd": "Asynchronous mode: working directory used when starting a conversation.", + "configuration.approvalTimeoutMs": "Milliseconds before an unanswered Codex approval or input request is denied locally.", + "configuration.allowHighRiskApprovals": "Allow the remote operator to approve high-risk commands. Keep disabled unless the relay and host are tightly controlled.", + "configuration.relayReconnect": "Reconnect outbound relay WebSockets after a disconnect." +} diff --git a/aether-vscodex/vscode-extension/package.nls.zh-cn.json b/aether-vscodex/vscode-extension/package.nls.zh-cn.json new file mode 100644 index 000000000..58314d4f0 --- /dev/null +++ b/aether-vscodex/vscode-extension/package.nls.zh-cn.json @@ -0,0 +1,38 @@ +{ + "extension.displayName": "Codex 远程协同", + "extension.description": "从本地浏览器或 Aether 云端同步 VS Code 当前 Codex 会话,或独立管理 Codex 会话。", + "command.openWeb": "Codex 远程:打开本地 Web 控制台", + "command.start": "Codex 远程:启动桥接", + "command.stop": "Codex 远程:停止桥接", + "command.setThreadId": "Codex 远程:设置现有会话 ID", + "command.sendInput": "Codex 远程:发送输入", + "command.setRelayToken": "Codex 远程:设置本地中继令牌", + "command.configureCloud": "Codex 远程:手动配置 Aether 云端", + "command.pairCloud": "Codex 远程:与 Aether 配对", + "command.snapshot": "Codex 远程:显示会话快照", + "configuration.title": "Codex 远程协同", + "configuration.localRelayUrl": "本地浏览器控制台使用的回环中继地址。启用 Aether 云同步后仍保持连接。", + "configuration.relayUrl": "为兼容旧版本保留的中继设置。新安装请使用 localRelayUrl 和 cloudRelayUrl。", + "configuration.relayUrl.deprecation": "本地访问请使用 codexRemoteCollab.localRelayUrl,Aether 云端访问请使用 codexRemoteCollab.cloudRelayUrl。", + "configuration.cloudRelayUrl": "可选的 Aether 云端 WebSocket 中继地址。设备凭据单独保存在 VS Code SecretStorage 中。", + "configuration.aetherUrl": "一次性配对流程使用的 Aether 服务地址。", + "configuration.autoStart": "扩展激活时自动启动桥接。", + "configuration.autoStartLocalRelay": "为回环 ws:// 地址自动启动扩展内置的本地中继。", + "configuration.mode": "附加到官方 VS Code Codex 现有会话,或为兼容旧版本启动独立 app-server。", + "configuration.mode.deprecation": "请改用 codexRemoteCollab.controlMode。attach 对应 sync,spawn 对应 async。", + "configuration.controlMode": "选择 Web 控制台是跟随 VS Code 当前 Codex 会话,还是独立管理会话。", + "configuration.controlMode.sync": "同步展示官方 VS Code Codex 面板当前打开的会话。", + "configuration.controlMode.async": "启动独立 Codex app-server,并从 Web 控制台管理其会话。", + "configuration.threadId": "要跟随的现有 VS Code Codex 会话 ID。留空时使用本机最近可附加的会话。", + "configuration.autoDiscoverThread": "未设置会话 ID 时自动发现最近的 VS Code Codex 会话。", + "configuration.followVscodeSession": "自动跟随官方 VS Code Codex 面板中的会话切换。", + "configuration.ipcSocketPath": "可选的官方 Codex IPC socket 路径。留空时使用 CODEX_HOME/ipc/ipc.sock。", + "configuration.hostId": "现有会话发现使用的 Codex 主机标识。", + "configuration.ipcStrictVersions": "拒绝未知的私有 IPC 流版本,不进行乐观兼容。", + "configuration.codexCommand": "异步模式:用于启动独立 app-server 的 Codex 可执行文件。", + "configuration.codexArgs": "异步模式:传给 Codex 可执行文件的参数。", + "configuration.defaultCwd": "异步模式:启动会话时使用的工作目录。", + "configuration.approvalTimeoutMs": "Codex 授权或输入请求无人处理时,在本地拒绝前等待的毫秒数。", + "configuration.allowHighRiskApprovals": "允许远程操作员批准高风险命令。仅在中继和主机均受严格控制时启用。", + "configuration.relayReconnect": "中继 WebSocket 断开后自动重连。" +} diff --git a/aether-vscodex/vscode-extension/scripts/sync-local-relay.cjs b/aether-vscodex/vscode-extension/scripts/sync-local-relay.cjs new file mode 100644 index 000000000..a4db0a17a --- /dev/null +++ b/aether-vscodex/vscode-extension/scripts/sync-local-relay.cjs @@ -0,0 +1,19 @@ +const fs = require("node:fs"); +const path = require("node:path"); + +const extensionRoot = path.resolve(__dirname, ".."); +const projectRoot = path.resolve(extensionRoot, ".."); +const outputRoot = path.join(extensionRoot, "dist", "local-relay"); +const publicRoot = path.join(extensionRoot, "dist", "public"); +const vuePublicRoot = path.join(projectRoot, "web", "dist"); + +if (!fs.existsSync(path.join(vuePublicRoot, "index.html"))) { + throw new Error("web/dist is missing; run npm run build:web before building the extension"); +} + +fs.rmSync(outputRoot, { recursive: true, force: true }); +fs.rmSync(publicRoot, { recursive: true, force: true }); +fs.mkdirSync(outputRoot, { recursive: true }); +fs.mkdirSync(publicRoot, { recursive: true }); +fs.copyFileSync(path.join(projectRoot, "relay", "server.js"), path.join(outputRoot, "server.js")); +fs.cpSync(vuePublicRoot, publicRoot, { recursive: true }); diff --git a/aether-vscodex/vscode-extension/src/bridge.ts b/aether-vscodex/vscode-extension/src/bridge.ts new file mode 100644 index 000000000..9f6050499 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/bridge.ts @@ -0,0 +1,46 @@ +import { CodexAgentAdapter, CodexAgentAdapterOptions } from "./codexAgentAdapter"; +import { RelayClient, RelayClientOptions } from "./relayClient"; +import { RelayHost, RelayHostOptions } from "./relayHost"; +import { AgentAdapter, Logger, RelayTransport } from "./protocol"; + +export interface CodexRemoteBridgeOptions { + /** Use a supplied adapter/transport when embedding or testing. */ + adapter?: AgentAdapter; + relay?: RelayTransport; + adapterOptions?: CodexAgentAdapterOptions; + relayOptions?: RelayClientOptions; + sessionId?: string; + capabilities?: Iterable; + logger?: Logger; +} +export interface CodexRemoteBridge { + adapter: AgentAdapter; + relay: RelayTransport; + host: RelayHost; + start(): Promise; + stop(): Promise; +} + +/** Construct the default outbound VS Code bridge in one call. */ +export function createBridge(options: CodexRemoteBridgeOptions): CodexRemoteBridge { + const adapter = options.adapter ?? new CodexAgentAdapter(options.adapterOptions); + const relay = options.relay ?? (() => { + if (!options.relayOptions) throw new Error("relayOptions are required when no relay transport is supplied"); + return new RelayClient(options.relayOptions); + })(); + const hostOptions: RelayHostOptions = { + adapter, + relay, + ...(options.sessionId ? { sessionId: options.sessionId } : {}), + ...(options.capabilities ? { capabilities: options.capabilities } : {}), + ...(options.logger ? { logger: options.logger } : {}), + }; + const host = new RelayHost(hostOptions); + return { + adapter, + relay, + host, + start: () => host.start(), + stop: () => host.stop(), + }; +} diff --git a/aether-vscodex/vscode-extension/src/cli.ts b/aether-vscodex/vscode-extension/src/cli.ts new file mode 100644 index 000000000..511b7defe --- /dev/null +++ b/aether-vscodex/vscode-extension/src/cli.ts @@ -0,0 +1,30 @@ +import { CodexAgentAdapter } from "./codexAgentAdapter"; +import { RelayHost } from "./relayHost"; +import { StdioRelayTransport } from "./relayClient"; + +/** Standalone bridge: relay frames in stdin, relay frames out on stdout. */ +async function main(): Promise { + const logger = { + debug: (message: string, ...args: unknown[]) => console.error(`[debug] ${message}`, ...args), + info: (message: string, ...args: unknown[]) => console.error(`[info] ${message}`, ...args), + warn: (message: string, ...args: unknown[]) => console.error(`[warn] ${message}`, ...args), + error: (message: string, ...args: unknown[]) => console.error(`[error] ${message}`, ...args), + }; + const command = process.env.CODEX_COMMAND || "codex"; + const args = process.env.CODEX_APP_SERVER_ARGS ? JSON.parse(process.env.CODEX_APP_SERVER_ARGS) as string[] : ["app-server", "--stdio"]; + const adapter = new CodexAgentAdapter({ command, args, defaultCwd: process.env.CODEX_WORKSPACE, logger }); + const relay = new StdioRelayTransport(process.stdin, process.stdout, logger); + const host = new RelayHost({ adapter, relay, sendHandshake: true, logger }); + const shutdown = async (): Promise => { + await host.stop(); + process.exit(0); + }; + process.once("SIGINT", () => void shutdown()); + process.once("SIGTERM", () => void shutdown()); + await host.start(); +} + +void main().catch((error) => { + console.error(error instanceof Error ? error.stack ?? error.message : String(error)); + process.exitCode = 1; +}); diff --git a/aether-vscodex/vscode-extension/src/codexAgentAdapter.ts b/aether-vscodex/vscode-extension/src/codexAgentAdapter.ts new file mode 100644 index 000000000..56a3fff80 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/codexAgentAdapter.ts @@ -0,0 +1,1914 @@ +import { createHash } from "node:crypto"; + +import { + AgentAdapter, + AgentEvent, + asJsonObject, + asJsonValue, + approvalDecisionKindForMethod, + Disposable, + hasApprovalDecisionField, + isRecord, + JsonObject, + JsonRpcId, + JsonRpcRequest, + JsonValue, + Logger, + PendingApproval, + SessionSnapshot, + isJsonRpcId, + jsonRpcIdKey, +} from "./protocol"; +import { JsonlRpcClient, JsonlRpcClientOptions } from "./jsonlRpc"; + +const APPROVAL_METHODS = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", + "item/permissions/requestApproval", + "applyPatchApproval", + "execCommandApproval", +]); + +const INPUT_REQUEST_METHODS = new Set([ + "item/tool/requestUserInput", + "mcpServer/elicitation/request", +]); + +// Paginated threads can contain many thousands of items. Keep hydration +// bounded so one history request cannot exhaust the extension host or exceed +// the relay frame limit, while still covering normal long-running sessions. +const HISTORY_PAGE_SIZE = 100; +const MAX_HISTORY_TURN_PAGES = 100; +const MAX_HISTORY_TURNS = HISTORY_PAGE_SIZE * MAX_HISTORY_TURN_PAGES; +const MAX_HISTORY_ITEM_PAGES = 100; + +export interface CodexAgentAdapterOptions extends JsonlRpcClientOptions { + clientName?: string; + clientTitle?: string | null; + clientVersion?: string; + initializeCapabilities?: JsonObject; + defaultCwd?: string; + maxOutputTailChars?: number; + approvalTimeoutMs?: number; + /** Handle non-approval server requests (auth refresh, tool calls, etc.). */ + onServerRequest?: (request: JsonRpcRequest) => Promise; + autoRejectUnsupportedRequests?: boolean; +} + +interface PendingRequest { + request: JsonRpcRequest; + approval?: PendingApproval; + timer?: NodeJS.Timeout; +} + +/** + * AgentAdapter implementation backed by a child `codex app-server --stdio`. + * + * It deliberately does not shell out for individual tasks. All task and + * approval operations go through the app-server JSON-RPC channel. + */ +export class CodexAgentAdapter implements AgentAdapter { + readonly rpc: JsonlRpcClient; + private readonly options: Required< + Pick + > & + Omit; + private readonly listeners = new Set<(event: AgentEvent) => void>(); + private readonly pending = new Map(); + private threadId: string | null = null; + private turnId: string | null = null; + private state = "disconnected"; + private outputTail = ""; + private messages: JsonValue[] = []; + private sessionMetadata: JsonObject = {}; + private availableModels: JsonValue[] = []; + private historyComplete = true; + private sessionSwitching = false; + private started = false; + private readonly rpcDisposables: Disposable[]; + + constructor(options: CodexAgentAdapterOptions = {}, rpc?: JsonlRpcClient) { + this.options = { + clientName: options.clientName ?? "codex-remote-collab", + clientVersion: options.clientVersion ?? "0.4.0", + maxOutputTailChars: options.maxOutputTailChars ?? 32_000, + approvalTimeoutMs: options.approvalTimeoutMs ?? 5 * 60_000, + ...options, + }; + this.rpc = rpc ?? new JsonlRpcClient(options); + this.rpcDisposables = [ + this.rpc.onNotification((message) => this.handleNotification(message.method, message.params)), + this.rpc.onServerRequest((request) => this.handleServerRequest(request)), + this.rpc.onExit((error) => { + this.started = false; + this.state = "disconnected"; + // A new app-server process cannot safely reuse ids from the dead + // process. Clear them before publishing the terminal event so a + // reconnect/restart cannot steer or interrupt a stale turn. + this.threadId = null; + this.turnId = null; + this.outputTail = ""; + this.messages = []; + this.sessionMetadata = {}; + this.historyComplete = true; + this.sessionSwitching = false; + // The child cannot receive a response after exit. Drop every pending + // approval/input and publish an explicit expiry so the relay removes + // its corresponding request instead of retaining a stale request id. + this.dropPendingRequests(error?.message ?? "app-server exited"); + this.emit({ type: "connection.closed", payload: error ? { message: error.message } : {} }); + }), + ]; + } + + async start(): Promise { + if (this.started) return; + await this.rpc.start(); + this.state = "initializing"; + const capabilities = { + experimentalApi: true, + requestAttestation: false, + ...(this.options.initializeCapabilities ?? {}), + }; + await this.rpc.request("initialize", { + clientInfo: { + name: this.options.clientName, + title: this.options.clientTitle ?? null, + version: this.options.clientVersion, + }, + capabilities, + }); + this.rpc.notify("initialized"); + this.started = true; + this.state = "idle"; + await this.refreshAvailableModels(); + this.emit({ type: "connection.opened", payload: {} }); + } + + async startThread(params: JsonObject = {}): Promise { + this.ensureStarted(); + this.ensureSessionChangeAllowed(); + const requestParams: JsonObject = { ...params }; + if (requestParams.cwd === undefined && this.options.defaultCwd) { + requestParams.cwd = this.options.defaultCwd; + } + const result = await this.rpc.request("thread/start", requestParams); + this.historyComplete = true; + this.commitThreadResult(result, true); + await this.publishHistorySnapshot(); + return result; + } + + async newSession(params: JsonObject = {}): Promise { + return this.startThread(params); + } + + async listSessions(params: JsonObject = {}): Promise { + this.ensureStarted(); + const result = await this.rpc.request("thread/list", normalizeThreadListParams(params)); + const threads = isRecord(result) && Array.isArray(result.data) + ? result.data.filter(isRecord).map((thread) => asJsonObject(thread)) + : []; + const sessions = threads.map((thread) => { + const id = typeof thread.id === "string" ? thread.id : ""; + const preview = typeof thread.preview === "string" ? thread.preview : ""; + const name = typeof thread.name === "string" ? thread.name : ""; + const cwd = typeof thread.cwd === "string" ? redactText(thread.cwd) : undefined; + const updatedAt = finiteNumber(thread.updatedAt) ?? finiteNumber(thread.createdAt); + return { + threadId: id, + title: sessionTitle(name || preview, id), + updatedAtMs: updatedAt === undefined ? null : Math.round(updatedAt * 1_000), + ...(cwd ? { cwd } : {}), + active: id === this.threadId, + available: Boolean(id), + ...(thread.status !== undefined ? { status: redactJson(thread.status) } : {}), + ...(thread.source !== undefined ? { source: redactJson(thread.source) } : {}), + }; + }).filter((session) => session.threadId); + return asJsonValue({ + sessions, + activeThreadId: this.threadId, + nextCursor: isRecord(result) && typeof result.nextCursor === "string" ? result.nextCursor : null, + backwardsCursor: isRecord(result) && typeof result.backwardsCursor === "string" ? result.backwardsCursor : null, + }); + } + + async selectSession(params: JsonObject): Promise { + this.ensureStarted(); + const target = (this.stringParam(params, "threadId") ?? this.stringParam(params, "conversationId"))?.trim(); + if (!target) throw new Error("session/select requires threadId"); + this.ensureSessionChangeAllowed(); + if (this.sessionSwitching) throw new Error("a session switch is already in progress"); + + const previousThreadId = this.threadId; + const previousState = this.state; + this.sessionSwitching = true; + this.state = "syncing"; + this.historyComplete = false; + this.emit({ + type: "session.switching", + threadId: target, + payload: { previousThreadId, targetThreadId: target }, + }); + try { + const resumeResult = await this.resumeThreadForSelection(target); + const hydrated = await this.ensureThreadHistory(resumeResult); + this.commitThreadResult(resumeResult, true, hydrated); + const switched = previousThreadId !== target; + this.emit({ + type: "session.selected", + threadId: target, + payload: { previousThreadId, threadId: target, switched, available: true }, + }); + await this.publishHistorySnapshot(); + return asJsonValue({ + threadId: target, + previousThreadId, + switched, + available: true, + result: redactJson(resumeResult), + }); + } catch (error) { + this.state = previousState; + this.historyComplete = true; + throw error; + } finally { + this.sessionSwitching = false; + } + } + + async updateThreadSettings(params: JsonObject): Promise { + this.ensureStarted(); + const threadId = this.stringParam(params, "threadId") ?? this.threadId; + if (!threadId) throw new Error("thread/settings/update requires threadId"); + if (threadId !== this.threadId) throw new Error("thread/settings/update can only target the selected thread"); + const settings = normalizeThreadSettings(params); + const result = await this.rpc.request("thread/settings/update", { threadId, ...settings.wire }); + this.mergeThreadSettings(settings.display); + await this.publishAuthoritativeSnapshot(); + return result; + } + + async startTurn(params: JsonObject): Promise { + this.ensureStarted(); + const requestParams = this.normalizeTurnParams(params); + const result = await this.rpc.request("turn/start", requestParams); + const turn = isRecord(result) && isRecord(result.turn) ? result.turn : undefined; + const nextTurnId = turn && typeof turn.id === "string" + ? turn.id + : isRecord(result) && typeof result.turnId === "string" + ? result.turnId + : undefined; + if (nextTurnId) this.turnId = nextTurnId; + if (typeof requestParams.threadId === "string") this.threadId = requestParams.threadId; + this.state = "active"; + return result; + } + + async steerTurn(params: JsonObject): Promise { + this.ensureStarted(); + const requestParams = this.normalizeTurnParams(params, true); + const result = await this.rpc.request("turn/steer", requestParams); + const nextTurnId = isRecord(result) && typeof result.turnId === "string" + ? result.turnId + : isRecord(result) && isRecord(result.turn) && typeof result.turn.id === "string" + ? result.turn.id + : undefined; + if (nextTurnId) this.turnId = nextTurnId; + if (typeof requestParams.threadId === "string") this.threadId = requestParams.threadId; + this.state = "active"; + return result; + } + + async interruptTurn(params: JsonObject): Promise { + this.ensureStarted(); + const threadId = this.stringParam(params, "threadId") ?? this.threadId; + const turnId = this.stringParam(params, "turnId") ?? this.turnId; + if (!threadId || !turnId) throw new Error("turn/interrupt requires threadId and turnId"); + const result = await this.rpc.request("turn/interrupt", { threadId, turnId }); + this.state = "idle"; + // Do this eagerly rather than waiting for the asynchronous + // `turn/completed` notification. A caller may submit the next turn as + // soon as the interrupt response resolves. + this.turnId = null; + return result; + } + + async sendInput(text: string, params: JsonObject = {}): Promise { + const body: JsonObject = { ...params, text }; + if (this.turnId) return this.steerTurn({ ...body, expectedTurnId: this.turnId }); + return this.startTurn(body); + } + + async cancel(taskId?: string, params: JsonObject = {}): Promise { + return this.interruptTurn({ ...params, ...(taskId ? { turnId: taskId } : {}) }); + } + + async respondApproval( + requestId: JsonRpcId, + decision: "allow" | "deny" | "cancel", + reason?: string, + response?: JsonValue, + ): Promise { + this.ensureStarted(); + const key = jsonRpcIdKey(requestId); + const pending = this.pending.get(key); + if (!pending) throw new Error(`unknown or already resolved approval request: ${key}`); + const rawResponse = response === undefined + ? this.defaultApprovalResponse(pending.request.method, pending.request.params, decision, reason) + : asJsonValue(response); + validateAdapterResponse(pending.request.method, rawResponse, decision); + const result = normalizeServerResponse(pending.request.method, rawResponse); + if (pending.timer) clearTimeout(pending.timer); + this.pending.delete(key); + try { + this.rpc.respond(pending.request.id, result); + } catch (error) { + // Do not leave a request retryable forever when the child exits between + // the liveness check and the JSON-RPC write. + this.emit({ + type: pending.approval ? "approval.expired" : "input.expired", + threadId: pending.approval?.threadId, + turnId: pending.approval?.turnId, + requestId, + payload: { requestId: asJsonValue(requestId), reason: "app-server unavailable" }, + }); + throw error; + } + this.emit({ + type: pending.approval ? "approval.resolved" : "input.resolved", + threadId: pending.approval?.threadId, + turnId: pending.approval?.turnId, + requestId, + payload: { + requestId: asJsonValue(requestId), + decision, + ...(reason ? { reason } : {}), + }, + }); + return result; + } + + async denyPending(reason = "relay disconnected"): Promise { + const pendingIds = [...this.pending.values()].map((entry) => entry.request.id); + for (const requestId of pendingIds) { + try { + await this.respondApproval(requestId, "deny", reason); + } catch { + // The app-server may have resolved or exited between the snapshot and + // this fail-closed cleanup pass. + } + } + } + + async snapshot(): Promise { + const status = this.statusSnapshot(); + return { + threadId: this.threadId, + turnId: this.turnId, + state: this.state, + pendingApprovals: [...this.pending.values()] + .map((entry) => entry.approval) + .filter((approval): approval is PendingApproval => Boolean(approval)), + pendingRequests: [...this.pending.values()].map((entry) => ({ + requestId: entry.request.id, + method: entry.request.method, + params: redactJson(asJsonObject(entry.request.params)), + ...(entry.approval?.commandHash ? { commandHash: entry.approval.commandHash } : {}), + ...(entry.approval?.risk ? { risk: entry.approval.risk } : {}), + ...(entry.approval?.summary ? { summary: entry.approval.summary } : {}), + ...(entry.approval?.createdAt ? { createdAt: entry.approval.createdAt } : {}), + ...(entry.approval?.expiresAt ? { expiresAt: entry.approval.expiresAt } : {}), + })), + outputTail: this.outputTail, + messages: this.messages.map((message) => asJsonValue(message)), + status, + activity: status.activity, + turnStatus: status.turnStatus, + activeFlags: [...status.activeFlags], + startedAtMs: status.startedAtMs, + durationMs: status.durationMs, + elapsedMs: status.elapsedMs, + metadata: { + adapter: "codex-app-server", + mode: "async", + started: this.started, + historyComplete: this.historyComplete, + ...this.sessionMetadata, + availableModels: this.availableModels.map((model) => asJsonValue(model)), + models: this.availableModels.map((model) => asJsonValue(model)), + }, + }; + } + + onEvent(listener: (event: AgentEvent) => void): Disposable { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + async dispose(): Promise { + for (const [key, entry] of this.pending) { + if (entry.timer) clearTimeout(entry.timer); + // A disconnected host must never leave a command approval hanging. + try { + this.rpc.respond(entry.request.id, normalizeServerResponse( + entry.request.method, + this.defaultApprovalResponse(entry.request.method, entry.request.params, "deny", "bridge stopped"), + )); + } catch { + // The child may already have exited. + } + this.pending.delete(key); + } + for (const disposable of this.rpcDisposables) disposable.dispose(); + this.rpc.close(); + this.started = false; + this.state = "disconnected"; + this.threadId = null; + this.turnId = null; + this.messages = []; + this.outputTail = ""; + this.sessionMetadata = {}; + this.historyComplete = true; + this.sessionSwitching = false; + } + + private ensureStarted(): void { + if (!this.started || !this.rpc.running) throw new Error("Codex app-server is not started"); + } + + private ensureSessionChangeAllowed(): void { + if (this.turnId || this.pending.size) { + throw new Error("cannot change sessions while a turn or approval is active"); + } + } + + private async refreshAvailableModels(): Promise { + try { + const models: JsonValue[] = []; + let cursor: string | null = null; + for (let page = 0; page < 10; page += 1) { + const result = await this.rpc.request("model/list", { + limit: 100, + includeHidden: false, + ...(cursor ? { cursor } : {}), + }); + if (!isRecord(result)) break; + if (Array.isArray(result.data)) { + for (const model of result.data) { + if (isRecord(model) && typeof model.model === "string") models.push(redactJson(model)); + } + } + cursor = typeof result.nextCursor === "string" && result.nextCursor ? result.nextCursor : null; + if (!cursor) break; + } + this.availableModels = models; + } catch (error) { + // Older app-server builds may not expose the model catalog. Thread and + // turn control should remain usable with a manually supplied model. + this.options.logger?.debug?.("Unable to load app-server model catalog", error); + } + } + + private async resumeThreadForSelection(threadId: string): Promise { + try { + // Paginated history is the stable protocol for newer app-server builds. + // Request the first page in chronological order so the renderer can use + // one consistent ordering while older pages are appended. + return await this.rpc.request("thread/resume", { + threadId, + excludeTurns: true, + initialTurnsPage: { + limit: HISTORY_PAGE_SIZE, + sortDirection: "asc", + itemsView: "full", + }, + }); + } catch (error) { + if (isPaginationUnsupportedError(error)) { + // Older app-server versions reject the pagination fields. Retry with + // the legacy full-history shape before giving up. + try { + return await this.rpc.request("thread/resume", { threadId, excludeTurns: false }); + } catch (legacyError) { + if (!isActiveWriterError(legacyError)) throw legacyError; + return this.readThreadMetadata(threadId, legacyError); + } + } + if (isActiveWriterError(error)) { + // A thread currently owned by another app-server cannot be resumed by + // this process, but its metadata and paginated history are still + // readable. Keep the conversation view available and report the + // writer limitation through metadata rather than showing an empty + // session after a successful list click. + return this.readThreadMetadata(threadId, error); + } + throw error; + } + } + + private async readThreadMetadata(threadId: string, originalError: unknown): Promise { + try { + return await this.rpc.request("thread/read", { threadId, includeTurns: false }); + } catch (error) { + this.options.logger?.debug?.("Unable to read a thread after resume failed", error); + throw originalError instanceof Error ? originalError : error; + } + } + + private async ensureThreadHistory(result: JsonValue): Promise { + const response = isRecord(result) ? result : {}; + const thread = extractThread(result); + if (!thread) throw new Error("thread/resume returned no thread"); + const turns = Array.isArray(thread.turns) ? thread.turns : undefined; + const threadId = typeof thread.id === "string" ? thread.id : undefined; + if (!threadId) throw new Error("thread/resume returned a thread without id"); + + const paginated = thread.historyMode === "paginated" || isRecord(response.initialTurnsPage); + const hasHistoryEvidence = Boolean( + (typeof thread.preview === "string" && thread.preview.trim()) + || ((finiteNumber(thread.updatedAt) ?? 0) > (finiteNumber(thread.createdAt) ?? 0)), + ); + if (!paginated && turns && (turns.length > 0 || !hasHistoryEvidence)) { + this.historyComplete = true; + return thread; + } + + if (paginated) { + try { + const initialPage = isRecord(response.initialTurnsPage) + ? response.initialTurnsPage + : undefined; + const hydratedTurns = await this.loadPaginatedTurns(threadId, initialPage); + return { ...thread, turns: hydratedTurns }; + } catch (error) { + // Some transitional server builds advertise paginated threads but do + // not implement one of the page methods. Fall back to the legacy read + // endpoint so the session remains usable instead of appearing blank. + this.options.logger?.debug?.("Paginated thread hydration failed; trying thread/read", error); + try { + const readResult = await this.rpc.request("thread/read", { threadId, includeTurns: true }); + const hydrated = extractThread(readResult); + if (hydrated) { + this.historyComplete = true; + return hydrated; + } + } catch (readError) { + this.options.logger?.debug?.("Legacy thread/read fallback failed", readError); + } + this.historyComplete = false; + return { ...thread, turns: turns ?? [] }; + } + } + + if (turns && turns.length > 0) { + this.historyComplete = true; + return thread; + } + + const readResult = await this.rpc.request("thread/read", { threadId, includeTurns: true }); + const hydrated = extractThread(readResult); + if (!hydrated) throw new Error("thread/read returned no thread"); + this.historyComplete = true; + return hydrated; + } + + private async loadPaginatedTurns(threadId: string, initialPage?: JsonObject): Promise { + const byId = new Map(); + let page: JsonObject | undefined = initialPage; + let cursor: string | null = null; + let complete = true; + + for (let index = 0; index < MAX_HISTORY_TURN_PAGES; index += 1) { + if (!page) { + const response = await this.rpc.request("thread/turns/list", { + threadId, + limit: HISTORY_PAGE_SIZE, + sortDirection: "asc", + itemsView: "full", + ...(cursor ? { cursor } : {}), + }); + page = isRecord(response) ? response : {}; + } + + const pageData = Array.isArray(page.data) ? page.data : []; + for (const value of pageData) { + if (!isRecord(value)) continue; + const id = typeof value.id === "string" ? value.id : `turn-${byId.size}`; + byId.set(id, { ...value }); + if (byId.size >= MAX_HISTORY_TURNS) { + complete = false; + break; + } + } + if (byId.size >= MAX_HISTORY_TURNS) break; + const next = typeof page.nextCursor === "string" && page.nextCursor ? page.nextCursor : null; + page = undefined; + cursor = next; + if (!cursor) break; + } + if (cursor) complete = false; + + const turns = sortHistoryTurns([...byId.values()]); + await this.hydrateTurnItems(threadId, turns, (value) => { + complete = complete && value; + }); + this.historyComplete = complete; + return turns; + } + + private async hydrateTurnItems( + threadId: string, + turns: JsonObject[], + markComplete: (complete: boolean) => void, + ): Promise { + for (const turn of turns) { + if (turn.itemsView === "full" || typeof turn.id !== "string") continue; + const items: JsonObject[] = Array.isArray(turn.items) + ? turn.items.filter(isRecord).map((item) => ({ ...(item as JsonObject) })) + : []; + const itemIds = new Set(items.map((item) => typeof item.id === "string" ? item.id : "")); + let cursor: string | null = null; + let complete = true; + try { + for (let pageIndex = 0; pageIndex < MAX_HISTORY_ITEM_PAGES; pageIndex += 1) { + const response = await this.rpc.request("thread/items/list", { + threadId, + turnId: turn.id, + limit: HISTORY_PAGE_SIZE, + sortDirection: "asc", + ...(cursor ? { cursor } : {}), + }); + const page = isRecord(response) ? response : {}; + for (const entry of Array.isArray(page.data) ? page.data : []) { + if (!isRecord(entry) || !isRecord(entry.item)) continue; + const item = { ...entry.item }; + const id = typeof item.id === "string" ? item.id : ""; + if (!id || !itemIds.has(id)) { + items.push(item); + if (id) itemIds.add(id); + } + } + const next = typeof page.nextCursor === "string" && page.nextCursor ? page.nextCursor : null; + cursor = next; + if (!cursor) break; + if (pageIndex === MAX_HISTORY_ITEM_PAGES - 1) complete = false; + } + } catch (error) { + complete = false; + this.options.logger?.debug?.(`Unable to hydrate items for turn ${turn.id}`, error); + } + turn.items = items; + turn.itemsView = "full"; + markComplete(complete); + } + } + + private commitThreadResult(result: JsonValue, replaceHistory: boolean, hydratedThread?: JsonObject): void { + const thread = hydratedThread ?? extractThread(result); + if (!thread) throw new Error("app-server thread response did not include a thread"); + const nextThreadId = typeof thread.id === "string" ? thread.id : undefined; + if (!nextThreadId) throw new Error("app-server thread response did not include thread.id"); + + this.threadId = nextThreadId; + if (replaceHistory) { + this.messages = projectThreadMessages(thread); + this.outputTail = outputTailFromMessages(this.messages, this.options.maxOutputTailChars); + } + const activeTurn = latestActiveTurn(thread); + this.turnId = activeTurn && typeof activeTurn.id === "string" ? activeTurn.id : null; + this.state = this.turnId ? "active" : statusToState(thread.status); + if (this.state === "notLoaded" || this.state === "unknown") this.state = "idle"; + + const response = isRecord(result) ? result : {}; + const title = sessionTitle( + typeof thread.name === "string" ? thread.name : typeof thread.preview === "string" ? thread.preview : "", + nextThreadId, + ); + const cwd = typeof response.cwd === "string" + ? redactText(response.cwd) + : typeof thread.cwd === "string" ? redactText(thread.cwd) : undefined; + const model = typeof response.model === "string" ? response.model : undefined; + const effort = response.reasoningEffort === null || typeof response.reasoningEffort === "string" + ? response.reasoningEffort + : undefined; + const threadSettings: JsonObject = { + ...(cwd ? { cwd } : {}), + ...(model ? { model } : {}), + ...(effort !== undefined ? { effort: asJsonValue(effort) } : {}), + ...(response.modelProvider !== undefined ? { modelProvider: redactJson(response.modelProvider) } : {}), + ...(response.serviceTier !== undefined ? { serviceTier: redactJson(response.serviceTier) } : {}), + ...(response.approvalPolicy !== undefined ? { approvalPolicy: redactJson(response.approvalPolicy) } : {}), + ...(response.approvalsReviewer !== undefined ? { approvalsReviewer: redactJson(response.approvalsReviewer) } : {}), + ...(response.sandbox !== undefined ? { sandboxPolicy: redactJson(response.sandbox) } : {}), + }; + this.sessionMetadata = { + thread: threadMetadata(thread), + title, + ...(cwd ? { cwd } : {}), + ...(model ? { model, latestModel: model } : {}), + ...(effort !== undefined ? { effort: asJsonValue(effort), latestReasoningEffort: asJsonValue(effort) } : {}), + ...(response.modelProvider !== undefined ? { modelProvider: redactJson(response.modelProvider) } : {}), + ...(response.approvalPolicy !== undefined ? { approvalPolicy: redactJson(response.approvalPolicy) } : {}), + ...(response.approvalsReviewer !== undefined ? { approvalsReviewer: redactJson(response.approvalsReviewer) } : {}), + ...(response.sandbox !== undefined ? { sandboxPolicy: redactJson(response.sandbox) } : {}), + threadSettings, + }; + } + + private mergeThreadSettings(settings: JsonObject): void { + const current = isRecord(this.sessionMetadata.threadSettings) + ? this.sessionMetadata.threadSettings + : {}; + const next = { ...current, ...redactJson(settings) as JsonObject }; + this.sessionMetadata.threadSettings = next; + for (const key of ["model", "modelProvider", "serviceTier", "approvalPolicy", "approvalsReviewer", "sandboxPolicy", "permissions", "cwd"] as const) { + if (settings[key] !== undefined) this.sessionMetadata[key] = redactJson(settings[key]); + } + if (settings.model !== undefined) this.sessionMetadata.latestModel = redactJson(settings.model); + if (settings.effort !== undefined) { + this.sessionMetadata.effort = redactJson(settings.effort); + this.sessionMetadata.latestReasoningEffort = redactJson(settings.effort); + } + } + + private async publishHistorySnapshot(): Promise { + const snapshot = await this.snapshot(); + this.emit({ + type: "output.snapshot", + threadId: this.threadId ?? undefined, + turnId: this.turnId ?? undefined, + payload: { + stream: "codex", + text: this.outputTail, + messages: this.messages.map((message) => asJsonValue(message)), + structureChanged: true, + historyComplete: this.historyComplete, + encoding: "utf8", + metadata: snapshot.metadata ?? {}, + status: snapshot.status ? asJsonValue(snapshot.status) : null, + }, + }); + await this.publishAuthoritativeSnapshot(snapshot); + } + + private async publishAuthoritativeSnapshot(existingSnapshot?: SessionSnapshot): Promise { + const snapshot = existingSnapshot ?? await this.snapshot(); + this.emit({ + type: "session.snapshot", + threadId: snapshot.threadId ?? undefined, + turnId: snapshot.turnId ?? undefined, + payload: asJsonObject(snapshot), + status: snapshot.status, + }); + } + + private statusSnapshot(): NonNullable { + const currentMessages = this.turnId + ? this.messages.filter((message) => isRecord(message) && message.turnId === this.turnId) + : []; + const startedAtValues = currentMessages + .map((message) => isRecord(message) ? finiteNumber(message.startedAtMs) : undefined) + .filter((value): value is number => value !== undefined); + const durationValues = currentMessages + .map((message) => isRecord(message) ? finiteNumber(message.durationMs) : undefined) + .filter((value): value is number => value !== undefined); + const startedAtMs = startedAtValues.length ? Math.min(...startedAtValues) : null; + const durationMs = durationValues.length ? Math.max(...durationValues) : null; + const pendingApprovals = [...this.pending.values()].some((entry) => Boolean(entry.approval)); + const pendingInput = [...this.pending.values()].some((entry) => !entry.approval); + const latest = currentMessages.length && isRecord(currentMessages[currentMessages.length - 1]) + ? currentMessages[currentMessages.length - 1] as JsonObject + : undefined; + let activity = this.turnId ? "thinking" : this.state; + if (pendingApprovals) activity = "waitingOnApproval"; + else if (pendingInput) activity = "waitingOnUserInput"; + else if (latest?.kind === "edit") activity = "editing"; + else if (latest?.itemType === "commandExecution") activity = "running"; + else if (latest?.kind === "reasoning" || latest?.kind === "plan") activity = "thinking"; + return { + activity, + turnStatus: this.turnId ? "inProgress" : this.state === "idle" ? "completed" : this.state, + activeFlags: [ + ...(pendingApprovals ? ["waitingOnApproval"] : []), + ...(pendingInput ? ["waitingOnUserInput"] : []), + ], + startedAtMs, + durationMs, + elapsedMs: this.turnId && startedAtMs !== null ? Math.max(0, Date.now() - startedAtMs) : null, + }; + } + + private upsertItem(item: JsonObject, turnId?: string, turn?: JsonObject, lifecycle: JsonObject = {}): void { + const projected = projectThreadItem(item, turnId, turn, lifecycle); + const itemId = typeof projected.itemId === "string" ? projected.itemId : undefined; + const index = itemId + ? this.messages.findIndex((message) => isRecord(message) + && message.itemId === itemId + && (turnId === undefined || message.turnId === turnId)) + : -1; + if (index >= 0) this.messages[index] = projected; + else this.messages.push(projected); + } + + private appendItemDelta(params: JsonObject, kind: "assistant" | "reasoning" | "plan" | "output", delta: string): void { + const itemId = this.extractString(params, "itemId"); + const turnId = this.extractString(params, "turnId"); + if (!itemId) return; + let index = this.messages.findIndex((message) => isRecord(message) + && message.itemId === itemId + && (turnId === undefined || message.turnId === turnId)); + if (index < 0) { + const placeholder: JsonObject = { + id: itemId, + itemId, + ...(turnId ? { turnId } : {}), + itemType: kind === "output" ? "commandExecution" : kind === "assistant" ? "agentMessage" : kind, + role: kind === "assistant" ? "assistant" : kind === "reasoning" ? "reasoning" : "tool", + kind: kind === "output" ? "tool" : kind, + text: "", + status: "inProgress", + }; + this.messages.push(placeholder); + index = this.messages.length - 1; + } + const current = isRecord(this.messages[index]) ? this.messages[index] as JsonObject : {}; + if (kind === "output") current.output = `${typeof current.output === "string" ? current.output : ""}${redactText(delta)}`; + else current.text = `${typeof current.text === "string" ? current.text : ""}${redactText(delta)}`; + this.messages[index] = current; + } + + private normalizeTurnParams(input: JsonObject, steering = false): JsonObject { + const params: JsonObject = { ...input }; + const threadId = this.stringParam(params, "threadId") ?? this.threadId; + if (!threadId) throw new Error(`${steering ? "turn/steer" : "turn/start"} requires threadId (start a thread first)`); + params.threadId = threadId; + + const suppliedInput = params.input; + if (typeof suppliedInput === "string") { + params.input = [this.textInput(suppliedInput)]; + } else if (Array.isArray(suppliedInput)) { + params.input = suppliedInput.map((item) => (typeof item === "string" ? this.textInput(item) : asJsonValue(item))); + } else { + const text = this.stringParam(params, "text") ?? this.stringParam(params, "message") ?? this.stringParam(params, "prompt"); + if (!text) throw new Error("turn request requires input or text"); + params.input = [this.textInput(text)]; + } + delete params.text; + delete params.message; + delete params.prompt; + if (steering) { + const expectedTurnId = this.stringParam(params, "expectedTurnId") ?? this.turnId; + if (!expectedTurnId) throw new Error("turn/steer requires expectedTurnId (no active turn)"); + params.expectedTurnId = expectedTurnId; + } + return params; + } + + private textInput(text: string): JsonObject { + return { type: "text", text, text_elements: [] }; + } + + private stringParam(params: JsonObject, key: string): string | undefined { + return typeof params[key] === "string" ? (params[key] as string) : undefined; + } + + private handleNotification(method: string, rawParams: JsonValue | undefined): void { + const params = asJsonObject(rawParams); + const threadId = this.extractString(params, "threadId") ?? this.extractNestedString(params, "thread", "id"); + const turnId = this.extractString(params, "turnId") ?? this.extractNestedString(params, "turn", "id"); + if (threadId && this.threadId && threadId !== this.threadId) { + this.options.logger?.debug?.(`Ignored late ${method} notification for non-selected thread ${threadId}`); + return; + } + const activeTurnId = this.turnId; + if (threadId && !this.threadId) this.threadId = threadId; + // A late completion for an earlier turn must not overwrite a newer turn + // that was started while the old completion notification was in flight. + // Token usage is thread telemetry, not a lifecycle transition. It can be + // delivered after `turn/completed`, so do not resurrect an old turn (or + // replace a newer active turn) just because the notification carries a + // turnId. + if (turnId + && method !== "thread/tokenUsage/updated" + && (method !== "turn/completed" || !activeTurnId || activeTurnId === turnId)) { + this.turnId = turnId; + } + + let type = "app-server.notification"; + let payload: JsonObject = { method, params: redactJson(params) as JsonObject }; + let outputText: string | undefined; + + switch (method) { + case "thread/started": + type = "session.created"; + payload = { thread: redactJson(params.thread ?? params) as JsonValue }; + if (isRecord(params.thread)) { + try { + this.commitThreadResult({ thread: params.thread }, true); + } catch (error) { + this.options.logger?.debug?.("Unable to hydrate thread/started notification", error); + this.state = "idle"; + } + } else { + this.state = "idle"; + } + break; + case "thread/name/updated": { + const title = this.extractString(params, "threadName")?.trim(); + if (title) this.sessionMetadata.title = redactText(title); + payload = redactJson(params) as JsonObject; + break; + } + case "thread/settings/updated": + payload = redactJson(params) as JsonObject; + if (isRecord(params.threadSettings)) this.mergeThreadSettings(params.threadSettings); + break; + case "thread/tokenUsage/updated": { + const tokenUsage = projectTokenUsage(params.tokenUsage); + // `redactJson` treats every key containing "token" as secret. Keep + // its redacted params for diagnostics, then add the numeric usage + // projection that the browser usage picker and relay snapshot need. + payload = redactJson(params) as JsonObject; + if (tokenUsage) { + this.sessionMetadata.tokenUsage = tokenUsage; + // The official extension names this field latestTokenUsageInfo; + // retain the shorter alias for existing relay/browser clients. + this.sessionMetadata.latestTokenUsageInfo = tokenUsage; + payload.tokenUsage = tokenUsage; + payload.latestTokenUsageInfo = tokenUsage; + // Persist usage through the same authoritative snapshot channel as + // thread settings so a browser reconnect does not fall back to the + // previous context-window value. + void this.publishAuthoritativeSnapshot().catch((error) => { + this.options.logger?.debug?.("Unable to publish token usage snapshot", error); + }); + } else { + // Do not let the redaction sentinel for malformed usage data look + // like a real update and clear a previously valid browser value. + delete payload.tokenUsage; + delete payload.latestTokenUsageInfo; + } + break; + } + case "thread/status/changed": + type = "session.state"; + payload = redactJson(params) as JsonObject; + this.state = statusToState(params.status); + break; + case "thread/closed": + case "thread/deleted": + type = "session.closed"; + payload = redactJson(params) as JsonObject; + this.state = "closed"; + this.threadId = null; + this.turnId = null; + this.messages = []; + this.outputTail = ""; + this.sessionMetadata = {}; + this.historyComplete = true; + break; + case "turn/started": + type = "task.started"; + payload = redactJson(params) as JsonObject; + this.state = "active"; + if (isRecord(params.turn) && Array.isArray(params.turn.items)) { + const turn = asJsonObject(params.turn); + for (const item of params.turn.items.filter(isRecord)) this.upsertItem(asJsonObject(item), turnId, turn); + } + break; + case "turn/completed": { + const status = this.extractNestedString(params, "turn", "status"); + type = status === "interrupted" ? "task.cancelled" : "task.finished"; + payload = redactJson(params) as JsonObject; + // Do not let a stale completion transition a newer active turn to + // idle. Notifications are asynchronous and can arrive after the + // caller has already started the next turn. + if (!turnId || !activeTurnId || turnId === activeTurnId) { + this.state = "idle"; + this.turnId = null; + } + if (isRecord(params.turn) && Array.isArray(params.turn.items)) { + const turn = asJsonObject(params.turn); + for (const item of params.turn.items.filter(isRecord)) this.upsertItem(asJsonObject(item), turnId, turn); + } + break; + } + case "item/agentMessage/delta": + type = "output.chunk"; + outputText = this.extractString(params, "delta"); + if (outputText) this.appendItemDelta(params, "assistant", outputText); + payload = { stream: "codex", text: redactText(outputText ?? ""), encoding: "utf8" }; + break; + case "item/plan/delta": + type = "output.chunk"; + outputText = this.extractString(params, "delta") ?? this.extractString(params, "text"); + if (outputText) this.appendItemDelta(params, "plan", outputText); + payload = { stream: "reasoning", text: redactText(outputText ?? ""), encoding: "utf8" }; + break; + case "item/reasoning/summaryTextDelta": + case "item/reasoning/textDelta": + type = "output.chunk"; + outputText = this.extractString(params, "delta") ?? this.extractString(params, "text"); + if (outputText) this.appendItemDelta(params, "reasoning", outputText); + payload = { stream: "reasoning", text: redactText(outputText ?? ""), encoding: "utf8" }; + break; + case "command/exec/outputDelta": + case "process/outputDelta": + case "item/commandExecution/outputDelta": + type = "output.chunk"; + outputText = decodeOutput(params); + if (method === "item/commandExecution/outputDelta" && outputText) { + this.appendItemDelta(params, "output", outputText); + } + payload = { + stream: outputStream(params), + text: redactText(outputText), + encoding: "utf8", + }; + break; + case "item/fileChange/outputDelta": + type = "output.chunk"; + outputText = this.extractString(params, "delta"); + if (outputText) this.appendItemDelta(params, "output", outputText); + payload = { stream: "codex", text: redactText(outputText ?? ""), encoding: "utf8" }; + break; + case "item/started": + type = "item.started"; + payload = redactJson(params) as JsonObject; + if (isRecord(params.item)) { + this.upsertItem(params.item, turnId, undefined, { + ...(finiteNumber(params.startedAtMs) !== undefined ? { startedAtMs: finiteNumber(params.startedAtMs) as number } : {}), + status: "inProgress", + }); + } + break; + case "item/completed": + type = "item.completed"; + payload = redactJson(params) as JsonObject; + if (isRecord(params.item)) { + this.upsertItem(params.item, turnId, undefined, { + ...(finiteNumber(params.completedAtMs) !== undefined ? { completedAtMs: finiteNumber(params.completedAtMs) as number } : {}), + }); + } + break; + case "serverRequest/resolved": + payload = redactJson(params) as JsonObject; + if (isJsonRpcId(params.requestId)) { + const pending = this.pending.get(jsonRpcIdKey(params.requestId)); + type = pending?.approval ? "approval.resolved" : pending ? "input.resolved" : "approval.resolved"; + if (pending?.timer) clearTimeout(pending.timer); + this.pending.delete(jsonRpcIdKey(params.requestId)); + } else { + type = "approval.resolved"; + } + break; + case "error": + type = "error"; + payload = redactJson(params) as JsonObject; + break; + case "warning": + case "guardianWarning": + type = "warning"; + payload = redactJson(params) as JsonObject; + break; + default: + break; + } + + // Events and snapshots must expose the same redacted view. Keeping raw + // text in outputTail would leak credentials through `snapshot()` even + // though the corresponding output event was redacted. + if (outputText) this.appendOutput(outputText); + this.emit({ + type, + threadId, + turnId, + payload, + raw: redactJson({ method, params }) as JsonValue, + }); + } + + private handleServerRequest(request: JsonRpcRequest): void { + const params = asJsonObject(request.params); + const requestThreadId = this.extractString(params, "threadId") ?? this.extractString(params, "conversationId"); + if (requestThreadId && this.threadId && requestThreadId !== this.threadId) { + // A resumed app-server can finish delivering an old request after the + // browser has selected another thread. Never expose or retain it as an + // approval for the selected conversation. + try { + if (APPROVAL_METHODS.has(request.method) || INPUT_REQUEST_METHODS.has(request.method)) { + this.rpc.respond(request.id, normalizeServerResponse( + request.method, + this.defaultApprovalResponse(request.method, request.params, "deny", "thread is no longer selected"), + )); + } else { + this.rpc.respondError(request.id, -32000, "thread is no longer selected"); + } + } catch { + // The child may have exited while the stale request was in flight. + } + this.options.logger?.debug?.(`Rejected stale ${request.method} request for non-selected thread ${requestThreadId}`); + return; + } + if (APPROVAL_METHODS.has(request.method)) { + const approval = this.toPendingApproval(request, params); + const entry: PendingRequest = { request, approval }; + if (this.options.approvalTimeoutMs > 0) { + entry.timer = setTimeout(() => this.expireApproval(request.id), this.options.approvalTimeoutMs); + approval.expiresAt = Date.now() + this.options.approvalTimeoutMs; + } + this.pending.set(jsonRpcIdKey(request.id), entry); + this.emit({ + type: "approval.requested", + threadId: approval.threadId, + turnId: approval.turnId, + requestId: request.id, + payload: { + ...approval.payload, + params: approval.payload, + requestId: asJsonValue(request.id), + method: request.method, + action: approval.action, + risk: approval.risk, + summary: approval.summary, + ...(approval.commandHash ? { commandHash: approval.commandHash } : {}), + ...(approval.expiresAt ? { expiresAt: approval.expiresAt } : {}), + }, + raw: redactJson(request) as JsonValue, + }); + return; + } + + if (INPUT_REQUEST_METHODS.has(request.method)) { + const entry: PendingRequest = { request }; + if (this.options.approvalTimeoutMs > 0) { + entry.timer = setTimeout(() => this.expirePendingRequest(request.id), this.options.approvalTimeoutMs); + } + this.pending.set(jsonRpcIdKey(request.id), entry); + this.emit({ + type: "input.requested", + threadId: this.extractString(params, "threadId"), + turnId: this.extractString(params, "turnId"), + requestId: request.id, + payload: { requestId: asJsonValue(request.id), method: request.method, params: redactJson(params) as JsonValue }, + raw: redactJson(request) as JsonValue, + }); + return; + } + + this.emit({ type: "server.request", requestId: request.id, payload: { method: request.method, params: redactJson(params) as JsonValue }, raw: redactJson(request) as JsonValue }); + void this.resolveServerRequest(request); + } + + private async resolveServerRequest(request: JsonRpcRequest): Promise { + try { + const result = await this.options.onServerRequest?.(request); + if (result !== undefined) { + this.rpc.respond(request.id, result); + } else if (this.options.autoRejectUnsupportedRequests !== false) { + this.rpc.respondError(request.id, -32601, `Unsupported app-server request: ${request.method}`); + } + } catch (error) { + this.rpc.respondError(request.id, -32000, error instanceof Error ? error.message : String(error)); + } + } + + private toPendingApproval(request: JsonRpcRequest, params: JsonObject): PendingApproval { + const threadId = this.extractString(params, "threadId") ?? this.extractString(params, "conversationId"); + const turnId = this.extractString(params, "turnId"); + const itemId = this.extractString(params, "itemId") ?? this.extractString(params, "callId"); + const command = this.extractString(params, "command") ?? this.extractCommand(params); + const reason = this.extractString(params, "reason"); + const action = approvalAction(request.method); + const risk = approvalRisk(request.method, command, params.commandActions); + const summary = reason || command || `${action} requested by Codex`; + return { + requestId: request.id, + method: request.method, + threadId, + turnId, + itemId, + action, + risk, + summary: redactText(summary), + commandHash: hashJson(params), + createdAt: Date.now(), + payload: redactJson(params) as JsonObject, + }; + } + + private async expireApproval(requestId: JsonRpcId): Promise { + return this.expirePendingRequest(requestId); + } + + private async expirePendingRequest(requestId: JsonRpcId): Promise { + const key = jsonRpcIdKey(requestId); + const pending = this.pending.get(key); + if (!pending) return; + this.pending.delete(key); + try { + this.rpc.respond(requestId, normalizeServerResponse( + pending.request.method, + this.expiredApprovalResponse(pending.request.method, pending.request.params), + )); + } catch { + // The app-server may have exited while the timer was pending. + } + this.emit({ + type: pending.approval ? "approval.expired" : "input.expired", + threadId: pending.approval?.threadId, + turnId: pending.approval?.turnId, + requestId, + payload: { requestId: asJsonValue(requestId), reason: "approval expired" }, + }); + } + + private dropPendingRequests(reason: string): void { + const pendingEntries = [...this.pending.values()]; + this.pending.clear(); + for (const pending of pendingEntries) { + if (pending.timer) clearTimeout(pending.timer); + this.emit({ + type: pending.approval ? "approval.expired" : "input.expired", + threadId: pending.approval?.threadId, + turnId: pending.approval?.turnId, + requestId: pending.request.id, + payload: { requestId: asJsonValue(pending.request.id), reason }, + }); + } + } + + private defaultApprovalResponse(method: string, rawParams: JsonValue | undefined, decision: "allow" | "deny" | "cancel", reason?: string): JsonValue { + const params = asJsonObject(rawParams); + if (method === "item/permissions/requestApproval") { + return { + permissions: decision === "allow" ? (params.permissions ?? {}) : {}, + scope: "turn", + }; + } + if (method === "item/tool/requestUserInput") { + return { answers: {} }; + } + if (method === "mcpServer/elicitation/request") { + return { action: decision === "allow" ? "accept" : decision === "cancel" ? "cancel" : "decline", content: null, _meta: null }; + } + if (method === "applyPatchApproval" || method === "execCommandApproval") { + if (decision === "allow") return { decision: "approved" }; + if (decision === "cancel") return { decision: "abort" }; + return { decision: { denied: { rejection: reason || "Denied remotely" } } }; + } + return { decision: decision === "allow" ? "accept" : decision === "cancel" ? "cancel" : "decline" }; + } + + private expiredApprovalResponse(method: string, rawParams: JsonValue | undefined): JsonValue { + if (method === "applyPatchApproval" || method === "execCommandApproval") { + // Preserve the legacy app-server wire decision for an actual timeout; + // `denied` is reserved for an explicit policy rejection. + return { decision: "timed_out" }; + } + return this.defaultApprovalResponse(method, rawParams, "deny", "approval expired"); + } + + private extractString(params: JsonObject, key: string): string | undefined { + return typeof params[key] === "string" ? (params[key] as string) : undefined; + } + + private extractNestedString(params: JsonObject, parent: string, key: string): string | undefined { + const nested = params[parent]; + return isRecord(nested) && typeof nested[key] === "string" ? (nested[key] as string) : undefined; + } + + private extractCommand(params: JsonObject): string | undefined { + const command = params.command; + if (Array.isArray(command)) return command.filter((item): item is string => typeof item === "string").join(" "); + // Newer command-approval requests may leave `command` null while + // providing parsed actions. Include every action command in the risk + // input so a dangerous subcommand cannot be hidden behind command:null. + if (Array.isArray(params.commandActions)) { + const commands = params.commandActions + .map((action) => isRecord(action) && typeof action.command === "string" ? action.command : undefined) + .filter((item): item is string => Boolean(item)); + if (commands.length) return commands.join(" && "); + } + return undefined; + } + + private appendOutput(text: string): void { + // Keep this invariant at the storage boundary. New notification handlers + // can append raw text later without creating a snapshot-only secret leak. + const safeText = redactText(text); + this.outputTail = `${this.outputTail}${safeText}`; + if (this.outputTail.length > this.options.maxOutputTailChars) { + this.outputTail = this.outputTail.slice(-this.options.maxOutputTailChars); + } + } + + private emit(event: AgentEvent): void { + for (const listener of this.listeners) { + try { + listener(event); + } catch (error) { + this.options.logger?.warn?.("Agent event listener failed", error); + } + } + } +} + +const THREAD_SORT_KEYS = new Set(["created_at", "updated_at", "recency_at", "section_position"]); +const THREAD_SOURCE_KINDS = new Set([ + "cli", "vscode", "exec", "appServer", "subAgent", "subAgentReview", + "subAgentCompact", "subAgentThreadSpawn", "subAgentOther", "unknown", +]); + +function normalizeThreadListParams(params: JsonObject): JsonObject { + const result: JsonObject = {}; + if (params.cursor === null || typeof params.cursor === "string") result.cursor = params.cursor; + const limit = finiteNumber(params.limit); + result.limit = Math.max(1, Math.min(100, Number.isInteger(limit) ? limit as number : 50)); + const sortKey = typeof params.sortKey === "string" ? params.sortKey : "updated_at"; + result.sortKey = THREAD_SORT_KEYS.has(sortKey) ? sortKey : "updated_at"; + result.sortDirection = params.sortDirection === "asc" ? "asc" : "desc"; + if (typeof params.archived === "boolean") result.archived = params.archived; + if (params.sectionId === null || typeof params.sectionId === "string") result.sectionId = params.sectionId; + if (typeof params.useStateDbOnly === "boolean") result.useStateDbOnly = params.useStateDbOnly; + const searchTerm = typeof params.searchTerm === "string" + ? params.searchTerm + : typeof params.query === "string" ? params.query : undefined; + if (searchTerm?.trim()) result.searchTerm = searchTerm.trim(); + if (typeof params.cwd === "string") result.cwd = params.cwd; + else if (Array.isArray(params.cwd) && params.cwd.every((value) => typeof value === "string")) { + result.cwd = asJsonValue(params.cwd); + } + if (Array.isArray(params.modelProviders) && params.modelProviders.every((value) => typeof value === "string")) { + result.modelProviders = asJsonValue(params.modelProviders); + } + if (Array.isArray(params.sourceKinds)) { + const sourceKinds = params.sourceKinds.filter((value): value is string => typeof value === "string" && THREAD_SOURCE_KINDS.has(value)); + if (sourceKinds.length) result.sourceKinds = sourceKinds; + } + return result; +} + +function normalizeThreadSettings(params: JsonObject): { wire: JsonObject; display: JsonObject } { + const source = isRecord(params.threadSettings) ? params.threadSettings : params; + const wire: JsonObject = {}; + const display: JsonObject = {}; + for (const key of ["model", "cwd", "effort", "serviceTier", "summary", "personality"] as const) { + if (!Object.prototype.hasOwnProperty.call(source, key)) continue; + const value = source[key]; + if (value !== null && (typeof value !== "string" || !value.trim())) { + throw new Error(`thread settings ${key} must be a non-empty string or null`); + } + wire[key] = typeof value === "string" ? value.trim() : null; + display[key] = wire[key]; + } + for (const key of ["collaborationMode", "multiAgentMode"] as const) { + if (!Object.prototype.hasOwnProperty.call(source, key)) continue; + const value = source[key]; + if (value !== null && typeof value !== "string" && !isRecord(value)) { + throw new Error(`thread settings ${key} must be a string, object, or null`); + } + wire[key] = asJsonValue(value); + display[key] = wire[key]; + } + for (const key of ["approvalPolicy", "approvalsReviewer"] as const) { + if (!Object.prototype.hasOwnProperty.call(source, key)) continue; + const value = source[key]; + if (value !== null && typeof value !== "string") { + throw new Error(`thread settings ${key} must be a string or null`); + } + wire[key] = asJsonValue(value); + display[key] = wire[key]; + } + + const hasPermissions = Object.prototype.hasOwnProperty.call(source, "permissions"); + const hasSandboxPolicy = Object.prototype.hasOwnProperty.call(source, "sandboxPolicy"); + if (hasPermissions) { + const value = source.permissions; + if (value !== null && (typeof value !== "string" || !value.trim())) { + throw new Error("thread settings permissions must be a non-empty string or null"); + } + wire.permissions = typeof value === "string" ? value.trim() : null; + display.permissions = wire.permissions; + } + if (hasSandboxPolicy) { + const value = source.sandboxPolicy; + if (value !== null && typeof value !== "string" && !isRecord(value)) { + throw new Error("thread settings sandboxPolicy must be a string, object, or null"); + } + display.sandboxPolicy = asJsonValue(value); + if (!hasPermissions) { + if (typeof value === "string") { + const permission = LEGACY_SANDBOX_PERMISSIONS[value]; + if (!permission) throw new Error(`unsupported legacy sandbox policy: ${value}`); + wire.permissions = permission; + display.permissions = permission; + } else { + wire.sandboxPolicy = asJsonValue(value); + } + } + } + if (!Object.keys(wire).length) throw new Error("thread settings update requires at least one setting"); + return { wire, display }; +} + +const LEGACY_SANDBOX_PERMISSIONS: Record = { + "read-only": ":read-only", + "workspace-write": ":workspace", + "danger-full-access": ":danger-full-access", +}; + +function isActiveWriterError(error: unknown): boolean { + const message = error instanceof Error ? error.message : String(error ?? ""); + return /active\s+writer|writer\s+lock|already\s+has\s+an\s+active\s+writer|thread\s+.*(?:locked|lock)|lock\s+.*thread/i.test(message); +} + +function isPaginationUnsupportedError(error: unknown): boolean { + if (isActiveWriterError(error)) return false; + const code = isRecord(error) && typeof error.code === "number" ? error.code : undefined; + if (code === -32602 || code === -32601) return true; + const message = error instanceof Error ? error.message : String(error ?? ""); + return /unknown\s+(?:field|parameter|method)|method\s+.*not\s+found|invalid\s+(?:param(?:eter)?s?|field)|unexpected\s+(?:field|property)|unsupported\s+(?:pagination|initialTurnsPage|excludeTurns|itemsView)/i.test(message); +} + +function sortHistoryTurns(turns: JsonObject[]): JsonObject[] { + return turns + .map((turn, index) => ({ turn, index })) + .sort((left, right) => { + const leftTime = [left.turn.startedAt, left.turn.createdAt, left.turn.completedAt] + .map(finiteNumber) + .find((value): value is number => value !== undefined); + const rightTime = [right.turn.startedAt, right.turn.createdAt, right.turn.completedAt] + .map(finiteNumber) + .find((value): value is number => value !== undefined); + if (leftTime !== undefined && rightTime !== undefined && leftTime !== rightTime) return leftTime - rightTime; + if (leftTime !== undefined && rightTime === undefined) return -1; + if (leftTime === undefined && rightTime !== undefined) return 1; + const leftId = typeof left.turn.id === "string" ? left.turn.id : ""; + const rightId = typeof right.turn.id === "string" ? right.turn.id : ""; + return leftId.localeCompare(rightId) || left.index - right.index; + }) + .map(({ turn }) => turn); +} + +function extractThread(result: JsonValue): JsonObject | undefined { + if (!isRecord(result)) return undefined; + if (isRecord(result.thread)) return result.thread; + return typeof result.id === "string" ? result : undefined; +} + +function threadMetadata(thread: JsonObject): JsonObject { + const result = redactJson(thread); + if (!isRecord(result)) return {}; + return { ...result, turns: [] }; +} + +const TOKEN_USAGE_FIELDS = [ + "totalTokens", + "inputTokens", + "cachedInputTokens", + "cacheWriteInputTokens", + "outputTokens", + "reasoningOutputTokens", +] as const; + +/** Keep only the numeric portion of the official token-usage projection. */ +function projectTokenUsage(value: unknown): JsonObject | undefined { + if (!isRecord(value)) return undefined; + const source = isRecord(value.info) + ? value.info + : isRecord(value.tokenUsage) + ? value.tokenUsage + : isRecord(value.token_usage) + ? value.token_usage + : value; + const total = projectTokenUsageBreakdown( + source.total + ?? source.total_token_usage + ?? source.totalTokenUsage, + ); + const last = projectTokenUsageBreakdown( + source.last + ?? source.last_token_usage + ?? source.lastTokenUsage, + ); + const modelContextWindow = tokenNumber( + source.modelContextWindow + ?? source.model_context_window + ?? source.contextWindow + ?? source.context_window, + ); + if (!total && !last && modelContextWindow === undefined) return undefined; + return { + ...(total ? { total } : {}), + ...(last ? { last } : {}), + ...(modelContextWindow !== undefined ? { modelContextWindow } : {}), + }; +} + +function projectTokenUsageBreakdown(value: unknown): JsonObject | undefined { + if (!isRecord(value)) return undefined; + const aliases: Record<(typeof TOKEN_USAGE_FIELDS)[number], string[]> = { + totalTokens: ["totalTokens", "total_tokens"], + inputTokens: ["inputTokens", "input_tokens"], + cachedInputTokens: ["cachedInputTokens", "cached_input_tokens"], + cacheWriteInputTokens: ["cacheWriteInputTokens", "cache_write_input_tokens"], + outputTokens: ["outputTokens", "output_tokens"], + reasoningOutputTokens: ["reasoningOutputTokens", "reasoning_output_tokens"], + }; + const result: JsonObject = {}; + for (const field of TOKEN_USAGE_FIELDS) { + for (const alias of aliases[field]) { + const number = tokenNumber(value[alias]); + if (number === undefined) continue; + result[field] = number; + break; + } + } + return Object.keys(result).length ? result : undefined; +} + +function tokenNumber(value: unknown): number | undefined { + if (typeof value === "number") return Number.isFinite(value) && value >= 0 ? value : undefined; + if (typeof value !== "string" || !/^\d+(?:\.\d+)?$/.test(value.trim())) return undefined; + const number = Number(value); + return Number.isFinite(number) && number >= 0 ? number : undefined; +} + +function latestActiveTurn(thread: JsonObject): JsonObject | undefined { + if (!Array.isArray(thread.turns)) return undefined; + for (let index = thread.turns.length - 1; index >= 0; index -= 1) { + const turn = thread.turns[index]; + if (isRecord(turn) && turn.status === "inProgress") return turn; + } + return undefined; +} + +function projectThreadMessages(thread: JsonObject): JsonValue[] { + if (!Array.isArray(thread.turns)) return []; + const messages: JsonValue[] = []; + for (const turn of thread.turns) { + if (!isRecord(turn) || !Array.isArray(turn.items)) continue; + const turnId = typeof turn.id === "string" ? turn.id : undefined; + for (const item of turn.items) { + if (isRecord(item)) messages.push(projectThreadItem(item, turnId, turn)); + } + } + return messages; +} + +function projectThreadItem( + item: JsonObject, + turnId?: string, + turn?: JsonObject, + lifecycle: JsonObject = {}, +): JsonObject { + const safeItem = redactJson(item); + const projected: JsonObject = isRecord(safeItem) ? { ...safeItem } : {}; + const itemType = typeof item.type === "string" ? item.type : "unknown"; + const itemId = typeof item.id === "string" ? item.id : undefined; + const turnStatus = typeof turn?.status === "string" ? turn.status : undefined; + const startedAtMs = finiteNumber(lifecycle.startedAtMs) ?? secondsToMs(turn?.startedAt); + const completedAtMs = finiteNumber(lifecycle.completedAtMs) ?? secondsToMs(turn?.completedAt); + const durationMs = finiteNumber(item.durationMs) ?? finiteNumber(turn?.durationMs); + const status = typeof lifecycle.status === "string" + ? lifecycle.status + : typeof item.status === "string" ? item.status : undefined; + + Object.assign(projected, { + ...(itemId ? { id: itemId, itemId } : {}), + ...(turnId ? { turnId } : {}), + itemType, + ...(status ? { status } : {}), + ...(turnStatus ? { turnStatus } : {}), + ...(startedAtMs !== undefined ? { startedAtMs } : {}), + ...(completedAtMs !== undefined ? { completedAtMs } : {}), + ...(durationMs !== undefined ? { durationMs } : {}), + }); + + switch (itemType) { + case "userMessage": + projected.role = "user"; + projected.kind = "user"; + projected.text = userInputText(item.content); + break; + case "agentMessage": + projected.role = "assistant"; + projected.kind = "assistant"; + projected.text = typeof item.text === "string" ? redactText(item.text) : ""; + break; + case "reasoning": + projected.role = "reasoning"; + projected.kind = "reasoning"; + projected.text = stringArrayText(item.summary) || stringArrayText(item.content); + break; + case "plan": + projected.role = "reasoning"; + projected.kind = "plan"; + projected.text = typeof item.text === "string" ? redactText(item.text) : ""; + break; + case "commandExecution": + projected.role = "tool"; + projected.kind = "tool"; + projected.command = typeof item.command === "string" ? redactText(item.command) : ""; + projected.text = typeof item.command === "string" ? redactText(item.command) : ""; + projected.output = typeof item.aggregatedOutput === "string" ? redactText(item.aggregatedOutput) : ""; + projected.label = "Command"; + projected.uiType = "commandExecution"; + break; + case "fileChange": { + projected.role = "tool"; + projected.kind = "edit"; + const paths = Array.isArray(item.changes) + ? item.changes.filter(isRecord).map((change) => { + const value = asJsonObject(change); + return typeof value.path === "string" ? redactText(value.path) : ""; + }).filter(Boolean) + : []; + projected.text = paths.join("\n"); + projected.label = "File changes"; + projected.uiType = "fileChange"; + break; + } + case "collabAgentToolCall": + projected.role = "tool"; + projected.kind = "tool"; + projected.text = typeof item.prompt === "string" ? redactText(item.prompt) : typeof item.tool === "string" ? item.tool : ""; + projected.action = typeof item.tool === "string" ? item.tool : ""; + projected.uiType = "collabAgentToolCall"; + break; + case "subAgentActivity": + projected.role = "tool"; + projected.kind = "tool"; + projected.text = typeof item.agentPath === "string" ? redactText(item.agentPath) : ""; + projected.activityKind = typeof item.kind === "string" ? item.kind : ""; + projected.uiType = "subAgentActivity"; + break; + case "webSearch": + projected.role = "tool"; + projected.kind = "tool"; + projected.text = typeof item.query === "string" ? redactText(item.query) : ""; + projected.label = "Web search"; + projected.uiType = "webSearch"; + break; + case "imageView": + projected.role = "tool"; + projected.kind = "tool"; + projected.text = typeof item.path === "string" ? redactText(item.path) : ""; + projected.label = "Image"; + break; + case "contextCompaction": + projected.role = "tool"; + projected.kind = "tool"; + projected.text = "Context compacted"; + projected.uiType = "contextCompaction"; + break; + default: + projected.role = "tool"; + projected.kind = "tool"; + projected.text = genericItemText(item); + break; + } + return projected; +} + +function userInputText(value: unknown): string { + if (!Array.isArray(value)) return ""; + return value.map((input) => { + if (!isRecord(input)) return ""; + if (input.type === "text" && typeof input.text === "string") return redactText(input.text); + if (input.type === "skill" && typeof input.name === "string") return `$${redactText(input.name)}`; + if (input.type === "mention" && typeof input.name === "string") return `@${redactText(input.name)}`; + if ((input.type === "image" || input.type === "audio") && typeof input.url === "string") return redactText(input.url); + if ((input.type === "localImage" || input.type === "localAudio") && typeof input.path === "string") return redactText(input.path); + return ""; + }).filter(Boolean).join("\n"); +} + +function stringArrayText(value: unknown): string { + return Array.isArray(value) + ? value.filter((entry): entry is string => typeof entry === "string").map(redactText).join("\n") + : ""; +} + +function genericItemText(item: JsonObject): string { + for (const key of ["text", "query", "command", "name", "tool"] as const) { + if (typeof item[key] === "string") return redactText(item[key] as string); + } + if (item.output !== undefined) return redactText(stableStringify(redactJson(item.output))); + if (item.result !== undefined) return redactText(stableStringify(redactJson(item.result))); + return ""; +} + +function outputTailFromMessages(messages: JsonValue[], maxChars: number): string { + const chunks: string[] = []; + for (const message of messages) { + if (!isRecord(message)) continue; + if (typeof message.text === "string" && message.text) chunks.push(message.text); + if (typeof message.output === "string" && message.output) chunks.push(message.output); + } + const output = redactText(chunks.join("\n\n")); + return output.length > maxChars ? output.slice(-maxChars) : output; +} + +function sessionTitle(value: string, threadId: string): string { + const title = redactText(value).replace(/\s+/g, " ").trim(); + if (title) return title.slice(0, 160); + return `Session ${threadId.slice(0, 8)}`; +} + +function secondsToMs(value: unknown): number | undefined { + const seconds = finiteNumber(value); + return seconds === undefined ? undefined : Math.round(seconds * 1_000); +} + +function finiteNumber(value: unknown): number | undefined { + return typeof value === "number" && Number.isFinite(value) ? value : undefined; +} + +/** Keep browser/embedding responses aligned with app-server response schemas. */ +function normalizeServerResponse(method: string, response: JsonValue): JsonValue { + if (method === "item/permissions/requestApproval") { + const source = isRecord(response) ? response : {}; + const requested = isRecord(source.permissions) ? source.permissions : {}; + const permissions: JsonObject = {}; + for (const [key, value] of Object.entries(requested)) { + // Request profiles use null to mean "not requested"; granted profiles + // omit those fields instead of sending an invalid explicit null. + if (value !== null && value !== undefined) permissions[key] = asJsonValue(value); + } + const normalized: JsonObject = { + permissions, + scope: source.scope === "session" ? "session" : "turn", + }; + if (typeof source.strictAutoReview === "boolean") normalized.strictAutoReview = source.strictAutoReview; + return normalized; + } + // Tool user-input responses are wrapped in an `answers` object. MCP + // elicitation has a different schema (`action`, `content`, `_meta`) and + // must be forwarded unchanged; wrapping it would make app-server reject + // an otherwise valid approval response. + if (method === "mcpServer/elicitation/request") return response; + if (method === "item/tool/requestUserInput") { + if (isRecord(response) && Object.prototype.hasOwnProperty.call(response, "answers")) return response; + return { answers: isRecord(response) ? response : {} }; + } + return response; +} + +/** Validate a response immediately before it crosses the app-server boundary. */ +function validateAdapterResponse( + method: string, + response: JsonValue, + decision: "allow" | "deny" | "cancel", +): void { + if (!isRecord(response)) throw new Error("app-server response must be a JSON object"); + + if (method === "item/permissions/requestApproval") { + if (!isRecord(response.permissions) + || (response.scope !== "turn" && response.scope !== "session") + || (response.strictAutoReview !== undefined && typeof response.strictAutoReview !== "boolean")) { + throw new Error("invalid permissions approval response"); + } + return; + } + + if (method === "item/tool/requestUserInput") { + const answers = response.answers; + if (!isRecord(answers)) throw new Error("invalid tool input response"); + return; + } + + if (method === "mcpServer/elicitation/request") { + if (!Object.prototype.hasOwnProperty.call(response, "action")) { + throw new Error("MCP elicitation response requires action"); + } + const action = approvalDecisionKindForMethod(response.action, method); + if (!action || action !== decision) throw new Error("MCP elicitation action conflicts with decision"); + return; + } + + // Approval callbacks all use a `decision` field. Unknown or mixed tagged + // objects are rejected by the method-aware classifier before write. + if (!hasApprovalDecisionField(response) || !Object.prototype.hasOwnProperty.call(response, "decision")) { + throw new Error("approval response requires decision"); + } + const responseDecision = approvalDecisionKindForMethod(response.decision, method); + if (!responseDecision) throw new Error("unsupported approval response decision"); + if (responseDecision !== decision) throw new Error(`approval response implies ${responseDecision}, but decision is ${decision}`); +} + +function statusToState(status: unknown): string { + if (typeof status === "string") return status; + if (isRecord(status) && typeof status.type === "string") return status.type; + return "unknown"; +} + +function approvalAction(method: string): string { + switch (method) { + case "item/commandExecution/requestApproval": + case "execCommandApproval": + return "command.execution"; + case "item/fileChange/requestApproval": + case "applyPatchApproval": + return "file.change"; + case "item/permissions/requestApproval": + return "permissions.grant"; + default: + return "approval"; + } +} + +function approvalRisk(method: string, command?: string, commandActions?: JsonValue): PendingApproval["risk"] { + // A permission profile can expand filesystem or network access for the + // current turn/session, so treat it like an explicit high-impact command. + if (method.includes("permissions")) return "high"; + if (method.includes("command") || method === "execCommandApproval") { + const suspicious = /(?:rm\s+-rf|sudo|curl|wget|ssh|password|token|secret)/i; + if (Array.isArray(commandActions)) { + // `commandActions` is parsed display data, not a proof of safety. Keep + // every request carrying it high-risk, scan all extracted commands, and + // treat unknown/malformed actions as high-risk as well. This prevents a + // dangerous subcommand from being hidden behind command:null. + const actions = commandActions; + const actionCommands = actions + .map((action) => isRecord(action) && typeof action.command === "string" ? action.command : undefined) + .filter((item): item is string => Boolean(item)); + const malformed = actions.some((action) => { + if (!isRecord(action) || typeof action.command !== "string") return true; + return action.type !== "read" + && action.type !== "listFiles" + && action.type !== "search" + && action.type !== "unknown"; + }); + const combined = [command, ...actionCommands].filter((item): item is string => Boolean(item)).join(" && "); + if (malformed || !actionCommands.length || suspicious.test(combined) || actions.some((action) => isRecord(action) && action.type === "unknown")) { + return "high"; + } + return "high"; + } + if (command && suspicious.test(command)) return "high"; + // An unparseable command approval is fail-closed. A missing command can + // otherwise be misclassified as medium and approved by the default host + // capability policy. + if (!command) return "high"; + return "medium"; + } + if (method.includes("fileChange") || method === "applyPatchApproval") return "medium"; + return "unknown"; +} + +function outputStream(params: JsonObject): string { + const stream = params.stream; + if (stream === "stderr" || stream === "stdout" || stream === "codex") return stream; + return "stdout"; +} + +function decodeOutput(params: JsonObject): string { + if (typeof params.delta === "string") return params.delta; + if (typeof params.deltaBase64 === "string") { + try { + return Buffer.from(params.deltaBase64, "base64").toString("utf8"); + } catch { + return "[invalid base64 output]"; + } + } + return ""; +} + +const SECRET_KEY = /(?:token|secret|password|authorization|api[_-]?key|private[_-]?key|refresh)/i; +const SECRET_VALUE = /(?:Bearer\s+)[A-Za-z0-9._~+\-/]+=*|(?:sk-[A-Za-z0-9_-]{12,}|gh[pousr]_[A-Za-z0-9_]{12,})/g; + +function redactText(text: string): string { + return text + .replace(SECRET_VALUE, "[REDACTED]") + .replace(/([?&](?:token|key|secret|password|api[_-]?key)=)[^&\s]+/gi, "$1[REDACTED]") + .replace(/((?:token|secret|password|api[_-]?key)\s*[:=]\s*)[^\s,;]+/gi, "$1[REDACTED]"); +} + +function redactJson(value: unknown): JsonValue { + if (Array.isArray(value)) return value.map((item) => redactJson(item)); + if (isRecord(value)) { + const result: JsonObject = {}; + for (const [key, child] of Object.entries(value)) { + result[key] = SECRET_KEY.test(key) ? "[REDACTED]" : redactJson(child); + } + return result; + } + if (typeof value === "string") return redactText(value); + return asJsonValue(value); +} + +function hashJson(value: JsonValue): string { + return createHash("sha256").update(stableStringify(value)).digest("hex"); +} + +function stableStringify(value: JsonValue): string { + if (Array.isArray(value)) return `[${value.map(stableStringify).join(",")}]`; + if (value !== null && typeof value === "object") { + return `{${Object.keys(value).sort().map((key) => `${JSON.stringify(key)}:${stableStringify(value[key] ?? null)}`).join(",")}}`; + } + return JSON.stringify(value); +} diff --git a/aether-vscodex/vscode-extension/src/codexIpc.ts b/aether-vscodex/vscode-extension/src/codexIpc.ts new file mode 100644 index 000000000..32842d045 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/codexIpc.ts @@ -0,0 +1,947 @@ +/** + * Minimal client for the private Codex desktop/VS Code coordination socket. + * + * This is intentionally separate from the app-server (JSONL/stdio) adapter. + * It attaches to the already running Codex UI through the local IPC router and + * therefore does not spawn another `codex` process. The wire protocol is + * private and versioned by the official extension; keep this module isolated + * so a protocol change can fail without taking down the relay bridge. + */ + +import * as crypto from "node:crypto"; +import * as net from "node:net"; +import * as os from "node:os"; +import * as path from "node:path"; + +import type { JsonObject, JsonValue } from "./protocol"; + +export const INITIALIZING_CLIENT_ID = "initializing-client"; +export const DEFAULT_IPC_REQUEST_TIMEOUT_MS = 5_000; +export const DEFAULT_MAX_IPC_FRAME_BYTES = 256 * 1024 * 1024; + +/** Versions shipped by openai.chatgpt 26.820.71523. */ +export const CODEX_IPC_METHOD_VERSIONS = Object.freeze({ + "thread-stream-state-changed": 11, + "thread-stream-following-changed": 1, + "thread-stream-following-status-requested": 1, + "ipc-connection-reset": 1, + "thread-read-state-changed": 2, + "thread-archived": 2, + "thread-unarchived": 1, + "thread-owner-discovery": 1, + "thread-follower-start-turn": 2, + "thread-follower-load-complete-history": 1, + "thread-follower-compact-thread": 1, + "thread-follower-steer-turn": 1, + "thread-follower-interrupt-turn": 4, + "thread-follower-update-thread-settings": 1, + "thread-follower-edit-last-user-turn": 2, + "thread-follower-command-approval-decision": 1, + "thread-follower-file-approval-decision": 1, + "thread-follower-permissions-request-approval-response": 1, + "thread-follower-submit-user-input": 1, + "thread-follower-submit-mcp-server-elicitation-response": 1, + "thread-follower-set-queued-follow-ups-state": 1, + "thread-queued-followups-changed": 1, +} as const); + +export type IpcMethod = keyof typeof CODEX_IPC_METHOD_VERSIONS; +export type IpcRequestId = string | number; +export type IpcPatchPathPart = string | number; + +export interface IpcRequest { + type: "request"; + requestId: IpcRequestId; + sourceClientId: string; + targetClientId?: string; + version: number; + method: string; + params?: JsonValue; + timeoutMs?: number; +} + +export interface IpcResponse { + type: "response"; + requestId: IpcRequestId; + resultType: "success" | "error"; + method?: string; + handledByClientId?: string; + result?: JsonValue; + error?: string; +} + +export interface IpcBroadcast { + type: "broadcast"; + method: string; + sourceClientId?: string; + targetClientIds?: string[]; + version: number; + params?: JsonValue; +} + +export interface IpcClientDiscoveryRequest { + type: "client-discovery-request"; + requestId: IpcRequestId; + request: IpcRequest; +} + +export interface IpcClientDiscoveryResponse { + type: "client-discovery-response"; + requestId: IpcRequestId; + response: { canHandle: boolean }; +} + +export type IpcMessage = + | IpcRequest + | IpcResponse + | IpcBroadcast + | IpcClientDiscoveryRequest + | IpcClientDiscoveryResponse; + +export interface IpcJsonPatch { + op: "add" | "remove" | "replace"; + path: IpcPatchPathPart[]; + value?: JsonValue; +} + +export interface ThreadStreamSnapshot { + type: "snapshot"; + revision: number; + conversationState: JsonObject; +} + +export interface ThreadStreamPatches { + type: "patches"; + baseRevision: number; + revision: number; + patches: IpcJsonPatch[]; +} + +export type ThreadStreamChange = ThreadStreamSnapshot | ThreadStreamPatches; + +export interface ConversationStreamState { + conversationId: string; + hostId: string; + ownerClientId: string; + revision: number; + conversationState: JsonObject; +} + +export type ConversationStreamEvent = + | (ConversationStreamState & { kind: "snapshot"; raw: IpcBroadcast }) + | (ConversationStreamState & { kind: "patches"; patches: IpcJsonPatch[]; baseRevision: number; raw: IpcBroadcast }) + | { + kind: "desync"; + conversationId: string; + hostId: string; + ownerClientId: string; + expectedRevision: number; + receivedBaseRevision: number; + receivedRevision: number; + raw: IpcBroadcast; + }; + +export interface CodexIpcClientOptions { + /** Explicit socket path; otherwise `$CODEX_HOME/ipc/ipc.sock` or `~/.codex`. */ + socketPath?: string; + codexHome?: string; + homeDir?: string; + env?: NodeJS.ProcessEnv; + platform?: NodeJS.Platform; + clientType?: string; + requestTimeoutMs?: number; + maxFrameBytes?: number; + strictVersions?: boolean; + /** Reconnect after a socket close and re-send all active following subscriptions. */ + autoReconnect?: boolean; + reconnectDelayMs?: number; + /** Optional handler for discovery requests. Default is fail-closed (`false`). */ + canHandleRequest?: (request: IpcRequest) => boolean | Promise; +} + +export interface FollowerTurnStartOptions { + request?: JsonObject; + context?: JsonObject; + clientUserMessageId?: string; + ownerClientId?: string; + timeoutMs?: number; +} + +export interface FollowerSteerOptions { + clientUserMessageId?: string; + serviceTier?: string | null; + attachments?: JsonValue[]; + additionalContext?: JsonObject | null; + restoreMessage?: JsonValue | null; + ownerClientId?: string; + timeoutMs?: number; +} + +export interface FollowerInterruptOptions { + mode?: "user-stop" | "system" | "descendant-cleanup" | string; + expectedTurnId?: string | null; + ownerClientId?: string; + timeoutMs?: number; +} + +export interface FollowOptions { + hostId?: string; + targetClientIds?: string[]; +} + +export interface RequestOptions { + targetClientId?: string; + timeoutMs?: number; + version?: number; + requestId?: IpcRequestId; +} + +export interface IpcErrorOptions { + code: string; + response?: IpcResponse; +} + +export class CodexIpcError extends Error { + readonly code: string; + readonly response?: IpcResponse; + + constructor(message: string, options: IpcErrorOptions) { + super(message); + this.name = "CodexIpcError"; + this.code = options.code; + this.response = options.response; + } +} + +export function resolveCodexIpcSocketPath(options: { + socketPath?: string; + codexHome?: string; + homeDir?: string; + env?: NodeJS.ProcessEnv; + platform?: NodeJS.Platform; +} = {}): string { + if (options.socketPath?.trim()) return options.socketPath.trim(); + const platform = options.platform ?? process.platform; + if (platform === "win32") return "\\\\.\\pipe\\codex-ipc"; + const env = options.env ?? process.env; + const homeDir = options.homeDir ?? os.homedir(); + const configuredHome = options.codexHome?.trim() || env.CODEX_HOME?.trim() || path.join(homeDir, ".codex"); + const codexHome = configuredHome === "~" + ? homeDir + : configuredHome.startsWith("~/") + ? path.join(homeDir, configuredHome.slice(2)) + : configuredHome; + return path.join(codexHome, "ipc", "ipc.sock"); +} + +/** Encode one private IPC frame: uint32 little-endian byte length + UTF-8 JSON. */ +export function encodeIpcFrame(message: IpcMessage, maxFrameBytes = DEFAULT_MAX_IPC_FRAME_BYTES): Buffer { + const json = JSON.stringify(message); + const payload = Buffer.from(json, "utf8"); + if (payload.length === 0 || payload.length > maxFrameBytes) { + throw new RangeError(`IPC frame exceeds ${maxFrameBytes} bytes`); + } + const frame = Buffer.allocUnsafe(4 + payload.length); + frame.writeUInt32LE(payload.length, 0); + payload.copy(frame, 4); + return frame; +} + +/** Incremental decoder that accepts arbitrary TCP/Unix-socket chunk boundaries. */ +export class IpcFrameDecoder { + private buffer = Buffer.alloc(0); + + constructor(private readonly maxFrameBytes = DEFAULT_MAX_IPC_FRAME_BYTES) {} + + push(chunk: Uint8Array): IpcMessage[] { + if (chunk.length === 0) return []; + this.buffer = this.buffer.length === 0 ? Buffer.from(chunk) : Buffer.concat([this.buffer, chunk]); + const messages: IpcMessage[] = []; + while (this.buffer.length >= 4) { + const payloadLength = this.buffer.readUInt32LE(0); + if (payloadLength === 0 || payloadLength > this.maxFrameBytes) { + throw new CodexIpcError(`Invalid IPC frame length (${payloadLength} bytes)`, { code: "invalid-frame-length" }); + } + if (this.buffer.length < payloadLength + 4) break; + const payload = this.buffer.subarray(4, payloadLength + 4).toString("utf8"); + this.buffer = this.buffer.subarray(payloadLength + 4); + let decoded: unknown; + try { + decoded = JSON.parse(payload); + } catch (error) { + throw new CodexIpcError(`Invalid IPC JSON: ${error instanceof Error ? error.message : String(error)}`, { + code: "invalid-json", + }); + } + if (!isRecord(decoded) || typeof decoded.type !== "string") { + throw new CodexIpcError("IPC frame must be an object with a type", { code: "invalid-message" }); + } + messages.push(decoded as unknown as IpcMessage); + } + return messages; + } + + reset(): void { + this.buffer = Buffer.alloc(0); + } +} + +type Listener = (value: T) => void; +export interface IpcSubscription { dispose(): void; } + +function subscribe(set: Set>, listener: Listener): IpcSubscription { + set.add(listener); + return { dispose: () => set.delete(listener) }; +} + +function isRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +function isJsonObject(value: unknown): value is JsonObject { + return isRecord(value); +} + +function requestIdKey(id: IpcRequestId): string { + return `${typeof id}:${String(id)}`; +} + +function cloneJson(value: T): T { + return JSON.parse(JSON.stringify(value)) as T; +} + +function versionFor(method: string, params?: JsonValue): number { + // The official client accepts interrupt v3 when expectedTurnId is absent; + // v4 is used when the active-turn precondition is present. + if (method === "thread-follower-interrupt-turn" + && (!isRecord(params) || params.expectedTurnId === undefined || params.expectedTurnId === null)) return 3; + return CODEX_IPC_METHOD_VERSIONS[method as IpcMethod] ?? 0; +} + +function textInput(text: string): JsonObject { + return { type: "text", text, text_elements: [] }; +} + +function normalizeInput(input: string | JsonValue[]): JsonValue[] { + return typeof input === "string" + ? [textInput(input)] + : input.map((entry) => typeof entry === "string" ? textInput(entry) : entry); +} + +function hasTarget(frame: IpcBroadcast, clientId: string): boolean { + return frame.targetClientIds == null || frame.targetClientIds.includes(clientId); +} + +/** Apply the JSON patch arrays generated by Immer in the official webview. */ +export function applyIpcPatches(root: JsonValue, patches: IpcJsonPatch[]): JsonValue { + let result = cloneJson(root); + for (const patch of patches) { + if (!Array.isArray(patch.path)) throw new CodexIpcError("IPC patch path must be an array", { code: "invalid-patch" }); + if (patch.path.length === 0) { + if (patch.op === "remove") throw new CodexIpcError("Removing the conversation root is unsupported", { code: "invalid-patch" }); + if (patch.value === undefined) throw new CodexIpcError("Patch value is missing", { code: "invalid-patch" }); + result = cloneJson(patch.value); + continue; + } + + const parentPath = patch.path.slice(0, -1); + const key = patch.path[patch.path.length - 1]; + assertSafePatchPart(key); + const parent = getAtPath(result, parentPath); + if (Array.isArray(parent)) { + const index = key === "-" ? parent.length : toArrayIndex(key); + if (patch.op === "add") { + if (patch.value === undefined) throw new CodexIpcError("Patch value is missing", { code: "invalid-patch" }); + parent.splice(index, 0, cloneJson(patch.value)); + } else if (patch.op === "replace") { + if (patch.value === undefined || index < 0 || index >= parent.length) throw new CodexIpcError("Invalid array replace patch", { code: "invalid-patch" }); + parent[index] = cloneJson(patch.value); + } else { + if (index < 0 || index >= parent.length) throw new CodexIpcError("Invalid array remove patch", { code: "invalid-patch" }); + parent.splice(index, 1); + } + continue; + } + if (!isRecord(parent) || typeof key !== "string") { + throw new CodexIpcError("IPC patch parent is not an object or array", { code: "invalid-patch" }); + } + if (patch.op === "remove") { + delete parent[key]; + } else { + if (patch.value === undefined) throw new CodexIpcError("Patch value is missing", { code: "invalid-patch" }); + parent[key] = cloneJson(patch.value); + } + } + return result; +} + +function getAtPath(root: JsonValue, pathParts: IpcPatchPathPart[]): JsonValue { + let current: JsonValue = root; + for (const part of pathParts) { + if (Array.isArray(current)) { + const index = toArrayIndex(part); + if (index < 0 || index >= current.length) throw new CodexIpcError("IPC patch path is out of bounds", { code: "invalid-patch" }); + current = current[index]; + } else if (isRecord(current) && typeof part === "string" && Object.prototype.hasOwnProperty.call(current, part)) { + assertSafePatchPart(part); + current = current[part]; + } else { + throw new CodexIpcError("IPC patch path does not exist", { code: "invalid-patch" }); + } + } + return current; +} + +function toArrayIndex(value: IpcPatchPathPart): number { + if (typeof value === "number" && Number.isInteger(value)) return value; + if (typeof value === "string" && /^\d+$/.test(value)) return Number(value); + throw new CodexIpcError(`Invalid array patch index: ${String(value)}`, { code: "invalid-patch" }); +} + +function assertSafePatchPart(value: IpcPatchPathPart): void { + if (value === "__proto__" || value === "prototype" || value === "constructor") { + throw new CodexIpcError("Unsafe IPC patch path", { code: "invalid-patch" }); + } +} + +export class CodexIpcClient { + readonly socketPath: string; + private readonly options: Required> & CodexIpcClientOptions; + private socket: net.Socket | undefined; + private decoder: IpcFrameDecoder; + private connectPromise: Promise | undefined; + private reconnectTimer: NodeJS.Timeout | undefined; + private disposed = false; + private clientId = INITIALIZING_CLIENT_ID; + private readonly pending = new Map void; reject: (error: Error) => void; timer: NodeJS.Timeout }>(); + private readonly followed = new Map(); + private readonly streams = new Map(); + private readonly messageListeners = new Set>(); + private readonly broadcastListeners = new Set>(); + private readonly streamListeners = new Set>(); + private readonly errorListeners = new Set>(); + private readonly closeListeners = new Set>(); + private readonly discoveryHandler?: (request: IpcRequest) => boolean | Promise; + + constructor(options: CodexIpcClientOptions = {}) { + this.options = { + ...options, + clientType: options.clientType ?? "codex-remote-collab", + requestTimeoutMs: options.requestTimeoutMs ?? DEFAULT_IPC_REQUEST_TIMEOUT_MS, + maxFrameBytes: options.maxFrameBytes ?? DEFAULT_MAX_IPC_FRAME_BYTES, + strictVersions: options.strictVersions ?? true, + autoReconnect: options.autoReconnect ?? false, + reconnectDelayMs: options.reconnectDelayMs ?? 1_000, + }; + this.socketPath = resolveCodexIpcSocketPath(options); + this.decoder = new IpcFrameDecoder(this.options.maxFrameBytes); + this.discoveryHandler = options.canHandleRequest; + } + + getClientId(): string { return this.clientId; } + + getConversationState(conversationId: string): ConversationStreamState | undefined { + const state = this.streams.get(conversationId); + return state == null ? undefined : { ...state, conversationState: cloneJson(state.conversationState) }; + } + + getFollowedConversations(): ReadonlyMap { return this.followed; } + + onMessage(listener: Listener): IpcSubscription { return subscribe(this.messageListeners, listener); } + onBroadcast(listener: Listener): IpcSubscription { return subscribe(this.broadcastListeners, listener); } + onStreamEvent(listener: Listener): IpcSubscription { return subscribe(this.streamListeners, listener); } + onError(listener: Listener): IpcSubscription { return subscribe(this.errorListeners, listener); } + onClose(listener: Listener): IpcSubscription { return subscribe(this.closeListeners, listener); } + + async connect(): Promise { + if (this.disposed) throw new CodexIpcError("IPC client is disposed", { code: "disposed" }); + if (this.reconnectTimer) { + clearTimeout(this.reconnectTimer); + this.reconnectTimer = undefined; + } + if (this.socket?.writable && this.clientId !== INITIALIZING_CLIENT_ID) return this.clientId; + if (this.connectPromise) return this.connectPromise; + this.connectPromise = new Promise((resolve, reject) => { + const socket = net.createConnection(this.socketPath); + this.socket = socket; + this.decoder.reset(); + let settled = false; + const finishError = (error: Error): void => { + if (!settled) { + settled = true; + reject(error); + } + this.emitError(error); + }; + socket.setNoDelay?.(true); + socket.on("connect", () => { + const requestId = crypto.randomUUID(); + const timer = setTimeout(() => { + this.pending.delete(requestIdKey(requestId)); + finishError(new CodexIpcError("IPC initialize timed out", { code: "timeout" })); + socket.destroy(); + }, this.options.requestTimeoutMs); + this.pending.set(requestIdKey(requestId), { + method: "initialize", + resolve: (response) => { + clearTimeout(timer); + if (response.resultType !== "success" || !isRecord(response.result) || typeof response.result.clientId !== "string") { + finishError(new CodexIpcError("IPC initialize returned an invalid response", { code: "initialize-failed", response })); + socket.destroy(); + return; + } + this.clientId = response.result.clientId; + settled = true; + resolve(this.clientId); + this.resubscribeAfterConnect().catch((error) => this.emitError(asError(error))); + }, + reject: (error) => { + clearTimeout(timer); + finishError(error); + socket.destroy(); + }, + timer, + }); + this.write({ + type: "request", + requestId, + sourceClientId: INITIALIZING_CLIENT_ID, + version: 0, + method: "initialize", + params: { clientType: this.options.clientType }, + }); + }); + socket.on("data", (chunk) => { + try { + for (const message of this.decoder.push(chunk)) this.handleMessage(message); + } catch (error) { + const normalized = asError(error); + finishError(normalized); + socket.destroy(normalized); + } + }); + socket.on("error", (error) => { + if (!settled) finishError(error); + else this.emitError(error); + }); + socket.on("close", () => { + this.handleClose(); + }); + }).finally(() => { + this.connectPromise = undefined; + }); + return this.connectPromise; + } + + async followConversation(conversationId: string, following = true, options: FollowOptions = {}): Promise { + const hostId = options.hostId ?? "local"; + await this.connect(); + if (following) this.followed.set(conversationId, hostId); + else { + this.followed.delete(conversationId); + this.streams.delete(conversationId); + } + const params: JsonObject = { conversationId, hostId, following }; + const frame: IpcBroadcast = { + type: "broadcast", + method: "thread-stream-following-changed", + sourceClientId: this.clientId, + version: CODEX_IPC_METHOD_VERSIONS["thread-stream-following-changed"], + params, + }; + if (options.targetClientIds) frame.targetClientIds = options.targetClientIds; + this.write(frame); + } + + async findThreadOwner(conversationId: string, hostId = "local", timeoutMs = this.options.requestTimeoutMs): Promise { + try { + const response = await this.request("thread-owner-discovery", { conversationId, hostId }, { timeoutMs }); + return response.handledByClientId ?? null; + } catch (error) { + if (error instanceof CodexIpcError + && (error.code === "no-client-found" || error.code.startsWith("no-client-found:"))) return null; + throw error; + } + } + + async request(method: string, params?: JsonValue, options: RequestOptions = {}): Promise { + await this.connect(); + const requestId = options.requestId ?? crypto.randomUUID(); + const timeoutMs = options.timeoutMs ?? this.options.requestTimeoutMs; + const frame: IpcRequest = { + type: "request", + requestId, + sourceClientId: this.clientId, + version: options.version ?? versionFor(method, params), + method, + params, + }; + if (options.targetClientId) frame.targetClientId = options.targetClientId; + if (timeoutMs > 0) frame.timeoutMs = timeoutMs; + return new Promise((resolve, reject) => { + const key = requestIdKey(requestId); + const timer = setTimeout(() => { + this.pending.delete(key); + reject(new CodexIpcError(`${method} timed out`, { code: "timeout" })); + }, timeoutMs > 0 ? timeoutMs : 2 ** 31 - 1); + this.pending.set(key, { method, resolve, reject, timer }); + try { + this.write(frame); + } catch (error) { + clearTimeout(timer); + this.pending.delete(key); + reject(asError(error)); + } + }).then((response) => { + if (response.resultType === "error") { + throw new CodexIpcError(response.error ?? `${method} failed`, { code: response.error ?? "ipc-error", response }); + } + if (response.method != null && response.method !== method) { + throw new CodexIpcError(`IPC response method mismatch: expected ${method}, got ${response.method}`, { + code: "response-method-mismatch", + response, + }); + } + return response; + }); + } + + async requestFollower(method: string, conversationId: string, params: JsonObject = {}, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + const ownerClientId = options.ownerClientId ?? this.streams.get(conversationId)?.ownerClientId; + if (!ownerClientId) throw new CodexIpcError(`No owner is known for conversation ${conversationId}`, { code: "owner-unknown" }); + // Do not allow a caller-provided params object to accidentally retarget a + // request after the owner has been selected from the stream snapshot. + const body: JsonObject = { ...params, conversationId }; + const { ownerClientId: _owner, ...requestOptions } = options; + return this.request(method, body, { ...requestOptions, targetClientId: ownerClientId }); + } + + /** Send the exact private `turnStart` envelope expected by the owner. */ + async startTurn(conversationId: string, input: string | JsonValue[], options: FollowerTurnStartOptions = {}): Promise { + const request: JsonObject = { + ...(options.request ?? {}), + threadId: conversationId, + input: options.request?.input ?? normalizeInput(input), + }; + const context: JsonObject = { inheritThreadSettings: true, ...(options.context ?? {}) }; + if (options.clientUserMessageId) request.clientUserMessageId = options.clientUserMessageId; + const response = await this.requestFollower("thread-follower-start-turn", conversationId, { + turnStart: { request, context }, + }, { + ownerClientId: options.ownerClientId, + timeoutMs: options.timeoutMs, + }); + return response.result; + } + + async steerTurn(conversationId: string, input: string | JsonValue[], options: FollowerSteerOptions = {}): Promise { + const params: JsonObject = { + clientUserMessageId: options.clientUserMessageId ?? crypto.randomUUID(), + input: normalizeInput(input), + attachments: options.attachments ?? [], + }; + if (options.serviceTier !== undefined) params.serviceTier = options.serviceTier; + if (options.additionalContext !== undefined) params.additionalContext = options.additionalContext; + if (options.restoreMessage !== undefined) params.restoreMessage = options.restoreMessage; + const response = await this.requestFollower("thread-follower-steer-turn", conversationId, params, { + ownerClientId: options.ownerClientId, + timeoutMs: options.timeoutMs, + }); + return response.result; + } + + /** + * Persist settings for the next turn through the official conversation + * owner. The owner-side follower handler expects the settings nested under + * `threadSettings`; `requestFollower` adds the conversation id to the + * outer envelope, yielding: + * `{ conversationId, threadSettings }`. + */ + async updateThreadSettings( + conversationId: string, + threadSettings: JsonObject, + options: RequestOptions & { ownerClientId?: string } = {}, + ): Promise { + if (!isJsonObject(threadSettings)) { + throw new CodexIpcError("thread settings must be a JSON object", { code: "invalid-thread-settings" }); + } + const response = await this.requestFollower( + "thread-follower-update-thread-settings", + conversationId, + { threadSettings: cloneJson(threadSettings) }, + options, + ); + return response.result; + } + + /** Alias matching the official app-server manager method name. */ + async updateThreadSettingsForNextTurn( + conversationId: string, + threadSettings: JsonObject, + options: RequestOptions & { ownerClientId?: string } = {}, + ): Promise { + return this.updateThreadSettings(conversationId, threadSettings, options); + } + + async interruptTurn(conversationId: string, options: FollowerInterruptOptions = {}): Promise { + const params: JsonObject = { mode: options.mode ?? "user-stop" }; + if (options.expectedTurnId !== undefined && options.expectedTurnId !== null) params.expectedTurnId = options.expectedTurnId; + const response = await this.requestFollower("thread-follower-interrupt-turn", conversationId, params, { + ownerClientId: options.ownerClientId, + timeoutMs: options.timeoutMs, + }); + return response.result; + } + + async loadCompleteHistory(conversationId: string, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + const response = await this.requestFollower("thread-follower-load-complete-history", conversationId, {}, options); + return response.result; + } + + async respondCommandApproval(conversationId: string, requestId: IpcRequestId, decision: JsonValue, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + return this.respondFollower("thread-follower-command-approval-decision", conversationId, { requestId, decision }, options); + } + + async respondFileApproval(conversationId: string, requestId: IpcRequestId, decision: JsonValue, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + return this.respondFollower("thread-follower-file-approval-decision", conversationId, { requestId, decision }, options); + } + + async respondPermissionsApproval(conversationId: string, requestId: IpcRequestId, response: JsonValue, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + return this.respondFollower("thread-follower-permissions-request-approval-response", conversationId, { requestId, response }, options); + } + + async respondUserInput(conversationId: string, requestId: IpcRequestId, response: JsonValue, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + return this.respondFollower("thread-follower-submit-user-input", conversationId, { requestId, response }, options); + } + + async respondMcpElicitation(conversationId: string, requestId: IpcRequestId, response: JsonValue, options: RequestOptions & { ownerClientId?: string } = {}): Promise { + return this.respondFollower("thread-follower-submit-mcp-server-elicitation-response", conversationId, { requestId, response }, options); + } + + private async respondFollower(method: string, conversationId: string, params: JsonObject, options: RequestOptions & { ownerClientId?: string }): Promise { + const response = await this.requestFollower(method, conversationId, params, options); + return response.result; + } + + async dispose(): Promise { + this.disposed = true; + if (this.reconnectTimer) clearTimeout(this.reconnectTimer); + this.reconnectTimer = undefined; + for (const pending of this.pending.values()) { + clearTimeout(pending.timer); + pending.reject(new CodexIpcError("IPC client disposed", { code: "disposed" })); + } + this.pending.clear(); + this.socket?.destroy(); + this.socket = undefined; + this.clientId = INITIALIZING_CLIENT_ID; + } + + private write(message: IpcMessage): void { + if (!this.socket?.writable) throw new CodexIpcError("IPC socket is not connected", { code: "not-connected" }); + this.socket.write(encodeIpcFrame(message, this.options.maxFrameBytes)); + } + + private handleMessage(message: IpcMessage): void { + for (const listener of this.messageListeners) safeCall(listener, message, (error) => this.emitError(error)); + switch (message.type) { + case "response": + this.handleResponse(message); + return; + case "broadcast": + this.handleBroadcast(message); + return; + case "client-discovery-request": + this.handleDiscoveryRequest(message).catch((error) => this.emitError(asError(error))); + return; + case "request": + this.handleUnexpectedRequest(message); + return; + case "client-discovery-response": + // Discovery responses are consumed by the router, not by clients. + return; + } + } + + private handleResponse(response: IpcResponse): void { + const key = requestIdKey(response.requestId); + const pending = this.pending.get(key); + if (!pending) return; + this.pending.delete(key); + clearTimeout(pending.timer); + pending.resolve(response); + } + + private handleBroadcast(frame: IpcBroadcast): void { + if (!hasTarget(frame, this.clientId)) return; + for (const listener of this.broadcastListeners) safeCall(listener, frame, (error) => this.emitError(error)); + if (frame.method === "thread-stream-state-changed") { + this.handleStreamStateBroadcast(frame); + } else if (frame.method === "thread-stream-following-status-requested") { + this.handleFollowingStatusRequested(frame); + } + } + + /** Re-announce active subscriptions when an owner reconnects or hands off. */ + private handleFollowingStatusRequested(frame: IpcBroadcast): void { + if (!isRecord(frame.params)) return; + if (this.options.strictVersions + && frame.version !== CODEX_IPC_METHOD_VERSIONS["thread-stream-following-status-requested"]) { + this.emitError(new CodexIpcError(`Unsupported thread following status version ${frame.version}`, { code: "version-mismatch" })); + return; + } + const conversationId = typeof frame.params.conversationId === "string" + ? frame.params.conversationId + : undefined; + const hostId = typeof frame.params.hostId === "string" ? frame.params.hostId : "local"; + const requester = frame.sourceClientId; + if (!conversationId || !requester || requester === this.clientId) return; + if (this.followed.get(conversationId) !== hostId) return; + void this.followConversation(conversationId, true, { + hostId, + targetClientIds: [requester], + }).catch((error) => this.emitError(asError(error))); + } + + private handleStreamStateBroadcast(frame: IpcBroadcast): void { + if (!isRecord(frame.params)) return; + const conversationId = typeof frame.params.conversationId === "string" ? frame.params.conversationId : undefined; + const hostId = typeof frame.params.hostId === "string" ? frame.params.hostId : "local"; + const change = frame.params.change; + if (!conversationId || !isRecord(change) || typeof change.type !== "string") return; + if (this.options.strictVersions && frame.version !== CODEX_IPC_METHOD_VERSIONS["thread-stream-state-changed"]) { + this.emitError(new CodexIpcError(`Unsupported thread stream version ${frame.version}`, { code: "version-mismatch" })); + return; + } + const ownerClientId = frame.sourceClientId ?? ""; + if (change.type === "snapshot") { + if (typeof change.revision !== "number" || !isJsonObject(change.conversationState)) return; + const state: ConversationStreamState = { + conversationId, + hostId, + ownerClientId, + revision: change.revision, + conversationState: cloneJson(change.conversationState), + }; + this.streams.set(conversationId, state); + this.emitStream({ kind: "snapshot", ...state, raw: frame }); + return; + } + if (change.type !== "patches" || typeof change.baseRevision !== "number" || typeof change.revision !== "number" || !Array.isArray(change.patches)) return; + const current = this.streams.get(conversationId); + if (!current || current.ownerClientId !== ownerClientId || current.revision !== change.baseRevision) { + const expectedRevision = current?.revision ?? 0; + this.emitStream({ + kind: "desync", + conversationId, + hostId, + ownerClientId, + expectedRevision, + receivedBaseRevision: change.baseRevision, + receivedRevision: change.revision, + raw: frame, + }); + // Re-sending `following:true` is how the official follower asks the + // owner for a fresh snapshot when a patch base revision is missed. + if (this.followed.has(conversationId)) { + this.followConversation(conversationId, true, { hostId }).catch((error) => this.emitError(asError(error))); + } + return; + } + try { + const patches = change.patches as unknown as IpcJsonPatch[]; + const nextConversationState = applyIpcPatches(current.conversationState, patches); + if (!isJsonObject(nextConversationState)) throw new CodexIpcError("Patched conversation state is not an object", { code: "invalid-patch" }); + const next: ConversationStreamState = { + ...current, + revision: change.revision, + conversationState: nextConversationState, + }; + this.streams.set(conversationId, next); + this.emitStream({ kind: "patches", ...next, patches, baseRevision: change.baseRevision, raw: frame }); + } catch (error) { + this.emitError(asError(error)); + } + } + + private async handleDiscoveryRequest(message: IpcClientDiscoveryRequest): Promise { + const request = message.request; + let canHandle = false; + try { + canHandle = this.discoveryHandler ? await this.discoveryHandler(request) : false; + } catch { + canHandle = false; + } + this.write({ + type: "client-discovery-response", + requestId: message.requestId, + response: { canHandle }, + }); + } + + private handleUnexpectedRequest(request: IpcRequest): void { + try { + this.write({ + type: "response", + requestId: request.requestId, + resultType: "error", + error: "no-handler-for-request", + }); + } catch (error) { + this.emitError(asError(error)); + } + } + + private async resubscribeAfterConnect(): Promise { + const subscriptions = [...this.followed.entries()]; + for (const [conversationId, hostId] of subscriptions) { + this.write({ + type: "broadcast", + method: "thread-stream-following-changed", + sourceClientId: this.clientId, + version: CODEX_IPC_METHOD_VERSIONS["thread-stream-following-changed"], + params: { conversationId, hostId, following: true }, + }); + } + } + + private handleClose(): void { + const socket = this.socket; + this.socket = undefined; + this.decoder.reset(); + const closeError = new CodexIpcError("IPC socket closed", { code: "connection-closed" }); + for (const pending of this.pending.values()) { + clearTimeout(pending.timer); + pending.reject(closeError); + } + this.pending.clear(); + this.clientId = INITIALIZING_CLIENT_ID; + for (const listener of this.closeListeners) safeCall(listener, closeError, (error) => this.emitError(error)); + if (!this.disposed && this.options.autoReconnect && socket) { + this.reconnectTimer = setTimeout(() => { + this.reconnectTimer = undefined; + this.connect().catch((error) => this.emitError(asError(error))); + }, this.options.reconnectDelayMs); + } + } + + private emitStream(event: ConversationStreamEvent): void { + for (const listener of this.streamListeners) safeCall(listener, event, (error) => this.emitError(error)); + } + + private emitError(error: Error): void { + for (const listener of this.errorListeners) safeCall(listener, error, () => undefined); + } +} + +function safeCall(listener: Listener, value: T, onError: (error: Error) => void): void { + try { + listener(value); + } catch (error) { + onError(asError(error)); + } +} + +function asError(error: unknown): Error { + return error instanceof Error ? error : new Error(String(error)); +} diff --git a/aether-vscodex/vscode-extension/src/codexIpcAgentAdapter.ts b/aether-vscodex/vscode-extension/src/codexIpcAgentAdapter.ts new file mode 100644 index 000000000..e7fbd3f18 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/codexIpcAgentAdapter.ts @@ -0,0 +1,4253 @@ +import { createHash, randomUUID } from "node:crypto"; +import { promises as fs } from "node:fs"; +import * as os from "node:os"; +import * as path from "node:path"; + +import { + CODEX_IPC_METHOD_VERSIONS, + CodexIpcClient, + CodexIpcClientOptions, + ConversationStreamState, + ConversationStreamEvent, + IpcBroadcast, + IpcSubscription, +} from "./codexIpc"; +import { + AgentAdapter, + AgentEvent, + AgentStatusSnapshot, + asJsonObject, + asJsonValue, + Disposable, + isJsonRpcId, + isRecord, + JsonObject, + JsonRpcId, + JsonValue, + Logger, + PendingApproval, + SessionSnapshot, + SessionListEntry, + SessionListResult, + SubagentSnapshot, + jsonRpcIdKey, +} from "./protocol"; + +const APPROVAL_METHODS = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", + "item/permissions/requestApproval", + "applyPatchApproval", + "execCommandApproval", +]); +const INPUT_METHODS = new Set([ + "item/tool/requestUserInput", + "mcpServer/elicitation/request", +]); +const TERMINAL_TURN_STATES = new Set([ + "completed", + "complete", + "failed", + "cancelled", + "canceled", + "interrupted", + "error", + "done", +]); +const HISTORY_LOAD_RETRY_DELAY_MS = 50; +const VSCODE_SESSION_FOLLOW_RETRY_DELAY_MS = 250; +const VSCODE_SESSION_FOLLOW_MAX_ATTEMPTS = 3; +/** Keep the relay alive while a user opens/selects the first official panel. */ +const WAITING_SESSION_DISCOVERY_DELAY_MS = 1_000; +/** A waiting poll is a fallback; route broadcasts handle the common fast path. */ +const WAITING_SESSION_DISCOVERY_MAX_CANDIDATES = 8; + +export interface CodexIpcAgentAdapterOptions extends Omit { + /** Existing conversation id. Empty/undefined enables local discovery. */ + threadId?: string; + hostId?: string; + autoDiscoverThread?: boolean; + /** Workspace paths used to rank auto-discovered sessions. */ + preferredCwds?: string[]; + /** Ask the owner for older paginated history after the initial snapshot. */ + loadCompleteHistory?: boolean; + /** Follow route changes in the official VS Code Codex panel. */ + followVscodeSession?: boolean; + /** Coalesce the official panel's old=false/new=true route broadcasts. */ + vscodeSessionFollowDebounceMs?: number; + ownerDiscoveryTimeoutMs?: number; + followTimeoutMs?: number; + maxOutputTailChars?: number; + approvalTimeoutMs?: number; + logger?: Logger; + /** Invoke the existing official VS Code command for a fresh Codex panel. */ + openNewSession?: () => Promise; + /** Inject a client in tests. The adapter owns it unless disabled below. */ + client?: CodexIpcClient; + disposeClient?: boolean; +} + +interface RequestEntry { + requestId: JsonRpcId; + method: string; + params: JsonObject; + threadId?: string; + turnId?: string; + createdAt: number; + expiresAt?: number; + approval?: PendingApproval; +} + +interface TurnInfo { + id?: string; + status: string; + active: boolean; + startedAt?: number; + durationMs?: number | null; + workedDurationMs?: number | null; + firstTurnWorkItemStartedAtMs?: number | null; + finalAssistantStartedAtMs?: number | null; + completedAtMs?: number | null; + error?: JsonValue; + /** The raw turn record, used to classify the currently running work item. */ + raw?: Record; +} + +/** + * AgentAdapter backed by the private IPC follower protocol used by the + * official OpenAI Codex VS Code extension. It never starts or kills a codex + * process: all work is routed to the owner of an already-open conversation. + */ +export class CodexIpcAgentAdapter implements AgentAdapter { + private readonly options: Required> & CodexIpcAgentAdapterOptions; + private readonly client: CodexIpcClient; + private readonly listeners = new Set<(event: AgentEvent) => void>(); + private readonly subscriptions: IpcSubscription[] = []; + private readonly pending = new Map(); + private readonly pendingTimers = new Map(); + private readonly pendingExpiryAt = new Map(); + private readonly optimisticallyResolved = new Set(); + private threadId: string | null = null; + private ownerClientId: string | null = null; + private revision: number | null = null; + private conversationState: JsonObject = {}; + private turnId: string | null = null; + private state = "disconnected"; + private status: AgentStatusSnapshot = { + activity: "idle", + turnStatus: "idle", + activeFlags: [], + startedAtMs: null, + durationMs: null, + workedDurationMs: null, + elapsedMs: null, + firstTurnWorkItemStartedAtMs: null, + finalAssistantStartedAtMs: null, + }; + private started = false; + private renderedOutput = ""; + private renderedOutputLength = 0; + private renderedOutputWasTruncated = false; + private renderedMessageShape = ""; + private outputTail = ""; + private outputMessages: RenderedConversationMessage[] = []; + private subagents: SubagentSnapshot[] = []; + private renderedSubagentShape = ""; + /** Fingerprint of the display-safe metadata projection last sent to peers. */ + private renderedMetadataShape = ""; + private snapshotSeen = false; + private historyComplete?: boolean; + private historyLoadRequested = false; + private historyLoadAttempts = 0; + private historyLoadGeneration = 0; + private historyLoadRetryTimer?: NodeJS.Timeout; + private disposed = false; + /** Serialize session navigation so two browser clicks cannot overlap. */ + private sessionSwitching = false; + /** Invalidates in-flight navigation when dispose/socket close begins. */ + private sessionLifecycleGeneration = 0; + /** IPC client that owns the official Codex panel route being mirrored. */ + private vscodeRouteClientId: string | null = null; + /** Untrusted old=false halves; only a matching same-source true can bind. */ + private readonly vscodeRouteCandidates = new Map(); + /** Last route that source reported as active, independent of our attachment. */ + private vscodeRouteActiveThreadId: string | null = null; + /** Official routing emits old=false before new=true, even when the picker stays open. */ + private vscodeRouteAwaitingSelection = false; + /** Latest route selected in that official panel, coalesced across rapid clicks. */ + private pendingVscodeThreadId: string | null = null; + private pendingVscodeFollowAttempts = 0; + private pendingVscodeFollowGeneration = 0; + private vscodeRouteGeneration = 0; + private activeVscodeSelection?: { target: string; generation: number }; + private vscodeSessionFollowTimer?: NodeJS.Timeout; + /** Serialize attachability probes with navigation so probe cleanup cannot unfollow a newly selected thread. */ + private sessionOperationTail: Promise = Promise.resolve(); + /** True while the IPC/relay host is usable but no conversation is attached. */ + private waitingForSession = false; + private waitingDiscoveryTimer?: NodeJS.Timeout; + private waitingDiscoveryInFlight?: Promise; + private waitingAttachPromise?: Promise; + private waitingAttachTarget: string | null = null; + private queuedWaitingAttachTarget: string | null = null; + private snapshotWaiter?: { threadId: string; ownerClientId?: string; resolve: () => void; reject: (error: Error) => void; timer: NodeJS.Timeout }; + private revisionWaiter?: { threadId: string; ownerClientId: string; revision: number; resolve: () => void; reject: (error: Error) => void; timer: NodeJS.Timeout }; + + constructor(options: CodexIpcAgentAdapterOptions = {}) { + this.options = { + ...options, + hostId: options.hostId ?? "local", + autoDiscoverThread: options.autoDiscoverThread ?? true, + followVscodeSession: options.followVscodeSession ?? true, + vscodeSessionFollowDebounceMs: Math.max(0, options.vscodeSessionFollowDebounceMs ?? 150), + // Owner discovery for an active local VS Code client normally returns + // in a few milliseconds. A short default keeps stale rollout files from + // making bridge startup look hung; callers can raise this explicitly. + ownerDiscoveryTimeoutMs: options.ownerDiscoveryTimeoutMs ?? 2_500, + followTimeoutMs: options.followTimeoutMs ?? 8_000, + maxOutputTailChars: options.maxOutputTailChars ?? 32_000, + approvalTimeoutMs: options.approvalTimeoutMs ?? 5 * 60_000, + disposeClient: options.disposeClient ?? true, + }; + this.client = options.client ?? new CodexIpcClient({ + ...options, + clientType: "codex-remote-collab-follower", + canHandleRequest: () => false, + autoReconnect: false, + }); + this.subscriptions.push(this.client.onBroadcast((frame) => this.handleBroadcast(frame))); + this.subscriptions.push(this.client.onStreamEvent((event) => this.handleStreamEvent(event))); + this.subscriptions.push(this.client.onError((error) => this.options.logger?.debug?.("Codex IPC error", error.message))); + this.subscriptions.push(this.client.onClose((error) => this.handleClose(error))); + } + + async start(): Promise { + if (this.started) return; + if (this.disposed) throw new Error("Codex IPC follower is disposed"); + this.clearWaitingDiscoveryTimer(); + this.waitingForSession = false; + this.resetHistoryLoading(); + this.renderedOutput = ""; + this.renderedOutputLength = 0; + this.renderedOutputWasTruncated = false; + this.renderedMessageShape = ""; + this.outputTail = ""; + this.outputMessages = []; + this.subagents = []; + this.renderedSubagentShape = ""; + this.renderedMetadataShape = ""; + const configured = this.options.threadId?.trim(); + await this.client.connect(); + if (configured) { + try { + await this.attachThread(configured); + return; + } catch (error) { + // A remembered thread can belong to another Codex window (or to a + // previous run). When auto-discovery is enabled, keep startup useful + // by falling back to the most recent live VS Code owner. + if (!isMissingSessionOwnerError(error)) throw error; + if (!this.options.autoDiscoverThread) { + this.enterWaitingForSession(); + return; + } + this.options.logger?.warn?.(`Configured Codex conversation ${configured} could not be followed; trying auto-discovery`, error); + const fallback = await this.discoverThreadId(new Set([configured])); + if (fallback) { + try { + await this.attachThread(fallback); + return; + } catch (fallbackError) { + if (!isMissingSessionOwnerError(fallbackError)) throw fallbackError; + } + } + this.enterWaitingForSession(); + return; + } + } + + const selectedThread = await this.discoverThreadId(); + if (!selectedThread) { + // A fresh VS Code window can have a live IPC socket before the user has + // opened a Codex conversation. Keep the follower (and therefore the + // relay/WebSocket) alive so the first later panel navigation can attach + // without restarting the bridge. + this.enterWaitingForSession(); + return; + } + await this.attachThread(selectedThread); + } + + private async attachThread(selectedThread: string, options: { fromWaiting?: boolean } = {}): Promise { + const fromWaiting = options.fromWaiting === true; + const lifecycleGeneration = this.sessionLifecycleGeneration; + const routeGeneration = this.vscodeRouteGeneration; + const wasStarted = this.started; + const wasWaiting = this.waitingForSession; + const previousRouteClientId = this.vscodeRouteClientId; + const previousRouteActiveThreadId = this.vscodeRouteActiveThreadId; + const previousRouteAwaitingSelection = this.vscodeRouteAwaitingSelection; + const owner = await this.client.findThreadOwner(selectedThread, this.options.hostId, this.options.ownerDiscoveryTimeoutMs); + if (!owner) { + throw new Error(`找不到会话 ${selectedThread} 的 VS Code Codex owner。请确认该会话已在官方 Codex 面板打开。`); + } + this.threadId = selectedThread; + this.ownerClientId = owner; + // An owner can be Codex Desktop while the visible VS Code webview follows + // it, so ownerClientId is not always the route source. Bind the route + // source on the first observed old=false/new=true navigation pair instead. + if (!this.vscodeRouteActiveThreadId) this.vscodeRouteActiveThreadId = selectedThread; + this.state = "syncing"; + if (!wasStarted) this.started = true; + // A normal startup announces the connection before waiting for the first + // stream snapshot. A waiting host already announced its IPC connection; + // emitting a second `connection.opened` would make the browser reset its + // connection indicator unnecessarily. + if (!wasStarted) this.emit({ type: "connection.opened", threadId: selectedThread, payload: { mode: "attach", ownerClientId: owner } }); + // Do not expose a provisional thread as interactive while its first + // authoritative snapshot is still in flight. + this.waitingForSession = false; + try { + const waitForSnapshot = this.waitForSnapshot(selectedThread, this.options.followTimeoutMs, owner); + void waitForSnapshot.catch(() => undefined); + await this.client.followConversation(selectedThread, true, { + hostId: this.options.hostId, + targetClientIds: [owner], + }); + await waitForSnapshot; + this.state = this.deriveSessionState(); + this.waitingForSession = false; + this.clearWaitingDiscoveryTimer(); + if (this.options.loadCompleteHistory !== false) void this.loadCompleteHistoryIfNeeded(); + this.options.logger?.info?.(`Attached to existing Codex conversation ${selectedThread}`); + } catch (error) { + // A failed follow must leave the adapter retryable and must not keep a + // stale thread/owner that could receive a later remote command. + this.clearSnapshotWaiter(error instanceof Error ? error : new Error(String(error))); + try { await this.client.followConversation(selectedThread, false, { hostId: this.options.hostId, targetClientIds: [owner] }); } catch { /* best effort */ } + const restoreWaiting = (fromWaiting || wasWaiting) + && !this.disposed + && lifecycleGeneration === this.sessionLifecycleGeneration; + // A transient attach attempt made from the waiting state must not tear + // down the IPC socket or relay. Restore the waiting projection and let + // the discovery loop try again after the official panel is ready. + this.started = restoreWaiting ? wasStarted : false; + this.state = restoreWaiting ? "waiting_for_host" : "disconnected"; + this.waitingForSession = restoreWaiting; + this.threadId = null; + this.ownerClientId = null; + if (!restoreWaiting) { + this.vscodeRouteClientId = null; + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteAwaitingSelection = false; + } else if (this.vscodeRouteGeneration === routeGeneration) { + this.vscodeRouteClientId = previousRouteClientId; + this.vscodeRouteActiveThreadId = previousRouteActiveThreadId; + this.vscodeRouteAwaitingSelection = previousRouteAwaitingSelection; + } + this.vscodeRouteCandidates.clear(); + this.resetConversationProjection(); + throw error; + } + } + + /** Attach mode deliberately has no thread creation operation. */ + async startThread(): Promise { + throw new Error("attach mode does not create a new thread; open an existing Codex conversation in VS Code"); + } + + /** + * Open a new conversation through the already-installed VS Code Codex + * extension. The callback is injected by the extension entrypoint; this + * follower never starts another Codex process. + */ + async newSession(): Promise { + // Opening the official new-session panel is also the recovery action for + // `waiting_for_host`; it does not require an existing conversation owner. + this.ensureStarted(); + if (!this.options.openNewSession) { + throw new Error("当前 VS Code Codex 扩展不支持从远程打开新会话"); + } + return asJsonValue(await this.options.openNewSession()); + } + + async startTurn(params: JsonObject): Promise { + this.ensureInteractiveReady(); + const input = extractInput(params); + const request = pickTurnRequest(params); + const result = await this.client.startTurn(this.threadId as string, input, { + request, + context: pickTurnContext(params), + clientUserMessageId: stringValue(params.clientUserMessageId), + ownerClientId: this.ownerClientId as string, + timeoutMs: this.options.followTimeoutMs, + }); + const unwrapped = unwrapFollowerResult(result); + const nextTurn = extractTurnId(unwrapped); + if (nextTurn) { + this.turnId = nextTurn; + this.state = "active"; + } + return asJsonValue(unwrapped); + } + + async steerTurn(params: JsonObject): Promise { + this.ensureInteractiveReady(); + const expected = stringValue(params.expectedTurnId) ?? this.turnId; + if (!expected) throw new Error("turn/steer requires an active turn"); + const result = await this.client.steerTurn(this.threadId as string, extractInput(params), { + clientUserMessageId: stringValue(params.clientUserMessageId) ?? randomUUID(), + serviceTier: params.serviceTier === null || typeof params.serviceTier === "string" ? params.serviceTier : undefined, + attachments: Array.isArray(params.attachments) ? params.attachments : [], + additionalContext: isRecord(params.additionalContext) ? asJsonObject(params.additionalContext) : undefined, + restoreMessage: params.restoreMessage === null || params.restoreMessage !== undefined ? asJsonValue(params.restoreMessage) : undefined, + ownerClientId: this.ownerClientId as string, + timeoutMs: this.options.followTimeoutMs, + }); + this.turnId = expected; + this.state = "active"; + return asJsonValue(unwrapFollowerResult(result)); + } + + /** Persist model/reasoning settings on the already-open official thread. */ + async updateThreadSettings(params: JsonObject): Promise { + this.ensureInteractiveReady(); + const threadSettings = pickThreadSettingsUpdate(params); + const result = await this.client.updateThreadSettings( + this.threadId as string, + threadSettings, + { + ownerClientId: this.ownerClientId as string, + timeoutMs: this.options.followTimeoutMs, + }, + ); + return asJsonValue(unwrapFollowerResult(result)); + } + + /** + * Return the local VS Code conversations that this follower can attach to. + * + * The official extension obtains this list from its app-server client via + * `thread/list`. That request is intentionally not exposed by the private + * IPC router, so the bridge uses local rollout/index metadata only to find + * candidates. A candidate is returned after live owner discovery and, for a + * non-active conversation, a matching follower snapshot. Closed, stale, or + * desktop-owned rollouts are omitted instead of being shown as selectable + * history that attach mode cannot actually open. + */ + async listSessions(params: JsonObject = {}): Promise { + // Session discovery is useful precisely while no conversation is attached + // (for example immediately after a fresh VS Code window opens). + this.ensureStarted(); + const releaseSessionOperation = await this.acquireSessionOperation(); + try { + return await this.listAttachableSessions(params); + } finally { + releaseSessionOperation(); + } + } + + private async listAttachableSessions(params: JsonObject): Promise { + const limitValue = numberValue(params.limit); + const limit = Math.max(1, Math.min(100, Number.isInteger(limitValue) ? limitValue as number : 50)); + const codexHome = resolveCodexHome(this.options); + const [candidates, index] = await Promise.all([ + recentVscodeThreadCandidates(path.join(codexHome, "sessions"), this.options.preferredCwds ?? []), + readSessionIndex(path.join(codexHome, "session_index.jsonl")), + ]); + const byId = new Map(); + for (const candidate of candidates) { + const indexed = index.get(candidate.id); + byId.set(candidate.id, { + ...candidate, + ...(indexed?.title && !candidate.title ? { title: indexed.title } : {}), + ...(indexed?.updatedAtMs !== undefined && indexed.updatedAtMs > (candidate.updatedAtMs ?? 0) + ? { updatedAtMs: indexed.updatedAtMs } : {}), + ...(indexed?.cwd && !candidate.cwd ? { cwd: indexed.cwd } : {}), + }); + } + // A configured/current thread can be valid even while its rollout has + // rotated away. Keep it in the picker so the active row is never lost. + if (this.threadId && !byId.has(this.threadId)) { + const title = stringValue(this.conversationState.title) + ?? stringValue(this.conversationState.name) + ?? stringValue(this.conversationState.threadTitle); + const cwd = stringValue(this.conversationState.cwd); + byId.set(this.threadId, { + id: this.threadId, + mtime: Date.now(), + updatedAtMs: Date.now(), + priority: 0, + ...(title ? { title } : {}), + ...(cwd ? { cwd } : {}), + }); + } + + const compareCandidates = (a: Candidate, b: Candidate) => (b.updatedAtMs ?? b.mtime) - (a.updatedAtMs ?? a.mtime) + || a.priority - b.priority; + const ordered = [...byId.values()] + .sort(compareCandidates) + .slice(0, limit); + // `limit` bounds expensive owner/snapshot probes, but a newer stale + // rollout must never consume the only slot and hide the active attachment. + const activeCandidate = this.threadId ? byId.get(this.threadId) : undefined; + if (activeCandidate && !ordered.some((candidate) => candidate.id === activeCandidate.id)) { + if (ordered.length >= limit) ordered[ordered.length - 1] = activeCandidate; + else ordered.push(activeCandidate); + ordered.sort(compareCandidates); + } + const sessions: SessionListEntry[] = []; + // Owner discovery is deliberately bounded/concurrent: stale rollout files + // are common, and one slow stale check must not block all other rows. + const ownerTimeout = Math.max(250, Math.min(this.options.ownerDiscoveryTimeoutMs, 750)); + // A discovery response only proves that *some* client knows the thread. + // Require a short, targeted snapshot probe before advertising a non-active + // row as selectable; desktop-owned/stale threads can otherwise look live + // and leave the browser blank after a failed switch. + const snapshotProbeTimeout = Math.max(250, Math.min(this.options.followTimeoutMs, 750)); + for (let offset = 0; offset < ordered.length; offset += 8) { + const batch = ordered.slice(offset, offset + 8); + const checked = await Promise.all(batch.map(async (candidate) => { + let owner: string | null = candidate.id === this.threadId ? this.ownerClientId : null; + if (!owner) { + try { + owner = await this.client.findThreadOwner(candidate.id, this.options.hostId, ownerTimeout); + } catch (error) { + this.options.logger?.debug?.(`Session owner discovery failed for ${candidate.id}`, error); + } + } + let available = false; + if (owner) { + available = candidate.id === this.threadId + || await this.probeSessionSnapshot(candidate.id, owner, snapshotProbeTimeout); + } + if (!available) return null; + const title = sanitizeSessionTitle(candidate.title) ?? `会话 ${candidate.id.slice(0, 8)}`; + return { + threadId: candidate.id, + title, + updatedAtMs: candidate.updatedAtMs ?? candidate.mtime, + ...(candidate.cwd ? { cwd: redactText(candidate.cwd) } : {}), + active: candidate.id === this.threadId, + available: true, + } satisfies SessionListEntry; + })); + for (const entry of checked) { + if (entry) sessions.push(entry); + } + // Route navigation has priority over populating more picker rows. The + // probes in this batch have already cleaned up their temporary follows; + // stop here so selectSession can acquire the shared operation lock + // instead of waiting behind dozens of stale rollout candidates. + if (this.sessionSwitching || this.pendingVscodeThreadId || this.activeVscodeSelection) break; + } + const result: SessionListResult = { + sessions, + activeThreadId: this.threadId, + }; + return asJsonValue(result); + } + + /** Attach to another already-open VS Code Codex conversation. */ + async selectSession(params: JsonObject): Promise { + this.ensureStarted(); + const origin = stringValue(params.origin) === "vscode" ? "vscode" : "web"; + const expectedRouteGeneration = origin === "vscode" + ? numberValue(params.vscodeRouteGeneration) + : undefined; + const target = (stringValue(params.threadId) ?? stringValue(params.conversationId))?.trim(); + if (!target) throw new Error("session/select requires threadId"); + if (target === this.threadId) { + return asJsonValue({ threadId: target, previousThreadId: target, switched: false, available: true }); + } + if (this.sessionSwitching) throw new Error("a session switch is already in progress"); + if (this.turnId || this.pending.size) { + throw new Error("cannot switch sessions while a turn or approval is active"); + } + const lifecycleGeneration = this.sessionLifecycleGeneration; + this.sessionSwitching = true; + const releaseSessionOperation = await this.acquireSessionOperation(); + const previousThreadId = this.threadId; + const previousOwnerClientId = this.ownerClientId; + const previousState = this.state; + // Keep a copy of the last owner-validated old-session state before changing + // the active conversation. It is a deterministic fallback if the target + // cannot produce a snapshot (for example, an already-running desktop + // writer). + let previousConversationState: ConversationStreamState | undefined; + let owner: string | null = null; + let attachmentChanged = false; + try { + this.assertSessionSelectionCurrent(target, origin, lifecycleGeneration, expectedRouteGeneration); + owner = await this.client.findThreadOwner(target, this.options.hostId, this.options.ownerDiscoveryTimeoutMs); + if (!owner) throw new Error(`找不到会话 ${target} 的 VS Code Codex owner。请确认该会话已在官方 Codex 面板打开。`); + this.assertSessionSelectionCurrent(target, origin, lifecycleGeneration, expectedRouteGeneration); + // A local turn/approval may have appeared while owner discovery was in + // flight. Re-check immediately before detaching the old projection. + if (this.turnId || this.pending.size) { + throw new Error("cannot switch sessions while a turn or approval is active"); + } + // Rollback must use the adapter's last owner-validated projection. The + // lower-level IPC cache sees every same-conversation snapshot before + // this adapter can reject a stale/unknown owner, so reading that cache + // here could resurrect content we deliberately ignored. + previousConversationState = previousThreadId && previousOwnerClientId + ? { + conversationId: previousThreadId, + hostId: this.options.hostId, + ownerClientId: previousOwnerClientId, + revision: this.revision ?? 0, + conversationState: cloneObject(this.conversationState), + } + : undefined; + this.emit({ + type: "session.switching", + threadId: target, + payload: { previousThreadId: previousThreadId ?? null, targetThreadId: target }, + }); + + // Keep the old follow alive until the new owner has supplied an + // authoritative snapshot. Events are filtered by `this.threadId`, so a + // stale old event cannot overwrite the target projection during attach. + this.threadId = target; + this.ownerClientId = owner; + this.state = "syncing"; + this.waitingForSession = false; + this.resetConversationProjection(); + attachmentChanged = true; + const waitForSnapshot = this.waitForSnapshot(target, this.options.followTimeoutMs, owner); + // If follow itself fails, the catch path rejects the waiter. Attach a + // handler immediately so that rejection can never become unhandled. + void waitForSnapshot.catch(() => undefined); + await this.client.followConversation(target, true, { + hostId: this.options.hostId, + targetClientIds: [owner], + }); + await waitForSnapshot; + this.assertSessionSelectionCurrent(target, origin, lifecycleGeneration, expectedRouteGeneration); + // Revisions are scoped to an owner. A handoff can happen after the first + // snapshot, so confirm the owner again before committing/unfollowing A. + const confirmedOwner = await this.client.findThreadOwner( + target, + this.options.hostId, + this.options.ownerDiscoveryTimeoutMs, + ); + if (confirmedOwner !== owner) { + throw new Error(`Codex conversation ${target} owner changed while switching`); + } + this.assertSessionSelectionCurrent(target, origin, lifecycleGeneration, expectedRouteGeneration); + if (previousThreadId && previousOwnerClientId) { + try { + await this.client.followConversation(previousThreadId, false, { + hostId: this.options.hostId, + targetClientIds: [previousOwnerClientId], + }); + } catch (error) { + this.options.logger?.debug?.(`Unable to unfollow previous session ${previousThreadId}`, error); + } + } + this.assertSessionSelectionCurrent(target, origin, lifecycleGeneration, expectedRouteGeneration); + this.state = this.deriveSessionState(); + this.waitingForSession = false; + this.clearWaitingDiscoveryTimer(); + if (this.options.loadCompleteHistory !== false) void this.loadCompleteHistoryIfNeeded(); + this.emit({ + type: "session.selected", + threadId: target, + payload: { + threadId: target, + activeThreadId: target, + previousThreadId: previousThreadId ?? null, + switched: true, + available: true, + }, + }); + return asJsonValue({ threadId: target, previousThreadId: previousThreadId ?? null, switched: true, available: true }); + } catch (error) { + this.clearSnapshotWaiter(error instanceof Error ? error : new Error(String(error))); + if (!attachmentChanged) throw error; + const lifecycleCurrent = lifecycleGeneration === this.sessionLifecycleGeneration + && !this.disposed + && this.started; + // Best-effort cleanup of the target subscription, then restore the old + // attachment so a failed switch does not strand the bridge disconnected. + if (lifecycleCurrent) { + try { + await this.client.followConversation(target, false, { + hostId: this.options.hostId, + ...(owner ? { targetClientIds: [owner] } : {}), + }); + } catch { /* best effort */ } + } + // dispose()/onClose owns the final disconnected state. An interrupted + // navigation must never publish rollback snapshots or reconnect after it. + if (!lifecycleCurrent) throw error; + this.threadId = previousThreadId; + this.ownerClientId = previousOwnerClientId; + this.state = previousState; + this.waitingForSession = !previousThreadId; + this.resetConversationProjection(); + // Move Relay/Web back before re-publishing the old snapshot. Browser + // command errors arrive after adapter events, so relying on only the + // command envelope would make the restored old projection look like an + // out-of-route event and leave the early target snapshot mounted. + if (previousThreadId) { + this.emit({ + type: "session.selected", + threadId: previousThreadId, + payload: { + threadId: previousThreadId, + activeThreadId: previousThreadId, + previousThreadId: target, + targetThreadId: target, + switched: false, + available: true, + failed: true, + origin, + }, + }); + } + const restoredFromCache = this.restoreCachedConversationProjection( + previousConversationState, + previousOwnerClientId, + ); + if (restoredFromCache && previousThreadId) { + // Re-assert the old follow without waiting for another snapshot. The + // cached state is already authoritative and keeps the bridge usable + // even when the owner does not answer a duplicate follow request. + try { + await this.client.followConversation(previousThreadId, true, { + hostId: this.options.hostId, + targetClientIds: this.ownerClientId ? [this.ownerClientId] : undefined, + }); + } catch (restoreError) { + this.options.logger?.debug?.("Unable to re-follow previous Codex session after switch failure", restoreError); + } + } else if (previousThreadId && previousOwnerClientId) { + try { + const restoreWaiter = this.waitForSnapshot(previousThreadId, this.options.followTimeoutMs, previousOwnerClientId); + void restoreWaiter.catch(() => undefined); + await this.client.followConversation(previousThreadId, true, { + hostId: this.options.hostId, + targetClientIds: [previousOwnerClientId], + }); + await restoreWaiter; + this.state = this.deriveSessionState(); + } catch (restoreError) { + this.options.logger?.warn?.("Unable to restore previous Codex session after switch failure", restoreError); + } + } + if (!previousThreadId && this.started && !this.disposed) this.scheduleWaitingDiscovery(); + throw error; + } finally { + releaseSessionOperation(); + this.sessionSwitching = false; + } + } + + /** Acquire a FIFO lock shared by session probes and attachment changes. */ + private async acquireSessionOperation(): Promise<() => void> { + const previous = this.sessionOperationTail; + let release!: () => void; + this.sessionOperationTail = new Promise((resolve) => { release = resolve; }); + await previous; + return release; + } + + private assertSessionSelectionCurrent( + target: string, + origin: "web" | "vscode", + lifecycleGeneration: number, + expectedRouteGeneration?: number, + ): void { + if (lifecycleGeneration !== this.sessionLifecycleGeneration || this.disposed || !this.started) { + throw new Error("session selection was cancelled because the IPC session closed"); + } + if (origin === "vscode" + && (expectedRouteGeneration === undefined + || expectedRouteGeneration !== this.vscodeRouteGeneration + || this.vscodeRouteActiveThreadId !== target)) { + throw new Error("session selection was superseded by a newer VS Code route"); + } + } + + async interruptTurn(params: JsonObject = {}): Promise { + this.ensureInteractiveReady(); + const expected = stringValue(params.turnId) ?? stringValue(params.expectedTurnId) ?? this.turnId; + if (!expected) throw new Error("turn/interrupt requires an active turn"); + const mode = stringValue(params.mode) ?? "user-stop"; + const result = await this.client.interruptTurn(this.threadId as string, { + mode, + expectedTurnId: expected, + ownerClientId: this.ownerClientId as string, + timeoutMs: this.options.followTimeoutMs, + }); + this.turnId = null; + this.state = "idle"; + const elapsed = this.status.elapsedMs ?? this.status.durationMs ?? null; + this.status = { + ...this.status, + activity: "interrupted", + turnStatus: "interrupted", + activeFlags: [], + durationMs: elapsed, + workedDurationMs: this.status.workedDurationMs ?? elapsed, + elapsedMs: elapsed, + }; + this.emit({ type: "task.cancelled", threadId: this.threadId ?? undefined, turnId: expected, payload: { mode } }); + return asJsonValue(unwrapFollowerResult(result)); + } + + async sendInput(text: string, params: JsonObject = {}): Promise { + const body = { ...params, text }; + return this.turnId ? this.steerTurn(body) : this.startTurn(body); + } + + async cancel(taskId?: string, params: JsonObject = {}): Promise { + return this.interruptTurn({ ...params, ...(taskId ? { turnId: taskId } : {}) }); + } + + async respondApproval( + requestId: JsonRpcId, + decision: "allow" | "deny" | "cancel", + reason?: string, + response?: JsonValue, + ): Promise { + return this.resolvePendingResponse(requestId, decision, reason, response, false); + } + + private async resolvePendingResponse( + requestId: JsonRpcId, + decision: "allow" | "deny" | "cancel", + reason: string | undefined, + response: JsonValue | undefined, + allowDuringSessionSwitch: boolean, + ): Promise { + this.ensureAttached(); + if (this.sessionSwitching && !allowDuringSessionSwitch) { + throw new Error("Codex session switch is still in progress"); + } + const key = jsonRpcIdKey(requestId); + const pending = this.pending.get(key); + if (!pending) throw new Error(`unknown or already resolved request: ${key}`); + const wire = this.toWireResponse(pending, decision, reason, response); + // Remove before awaiting the owner so a repeated browser click cannot send + // the same approval twice. If the IPC request fails, restore it for retry. + this.pending.delete(key); + this.clearPendingTimer(key); + this.optimisticallyResolved.add(key); + let result: JsonValue | undefined; + try { + result = await this.sendPendingResponse(pending, wire); + } catch (error) { + this.optimisticallyResolved.delete(key); + this.pending.set(key, pending); + this.schedulePendingExpiry(key, pending); + throw error; + } + this.emit({ + type: pending.approval ? "approval.resolved" : "input.resolved", + threadId: pending.threadId ?? this.threadId ?? undefined, + turnId: pending.turnId, + requestId, + payload: { requestId: asJsonValue(requestId), method: pending.method, decision }, + }); + return asJsonValue(result); + } + + async denyPending(reason = "relay disconnected"): Promise { + const entries = [...this.pending.values()]; + await Promise.all(entries.map(async (entry) => { + try { + await this.resolvePendingResponse(entry.requestId, "deny", reason, undefined, true); + } catch { /* fail closed when owner is gone */ } + })); + } + + private async sendPendingResponse(entry: RequestEntry, wire: JsonValue): Promise { + const conversationId = this.threadId as string; + const options = { ownerClientId: this.ownerClientId as string, timeoutMs: this.options.followTimeoutMs }; + if (entry.method === "item/commandExecution/requestApproval" || entry.method === "execCommandApproval") { + return this.client.respondCommandApproval(conversationId, entry.requestId, wire, options); + } + if (entry.method === "item/fileChange/requestApproval" || entry.method === "applyPatchApproval") { + return this.client.respondFileApproval(conversationId, entry.requestId, wire, options); + } + if (entry.method === "item/permissions/requestApproval") { + return this.client.respondPermissionsApproval(conversationId, entry.requestId, wire, options); + } + if (entry.method === "item/tool/requestUserInput") { + return this.client.respondUserInput(conversationId, entry.requestId, wire, options); + } + if (entry.method === "mcpServer/elicitation/request") { + return this.client.respondMcpElicitation(conversationId, entry.requestId, wire, options); + } + throw new Error(`unsupported follower request method: ${entry.method}`); + } + + private schedulePendingExpiry(key: string, entry: RequestEntry): void { + if (!entry.expiresAt || this.options.approvalTimeoutMs <= 0) { + // A later snapshot can omit an expiry that was present in an earlier + // request record. Do not leave the old timer alive in that case. + this.clearPendingTimer(key); + return; + } + const existing = this.pendingExpiryAt.get(key); + if (existing === entry.expiresAt && this.pendingTimers.has(key)) return; + this.clearPendingTimer(key); + const delay = Math.max(0, entry.expiresAt - Date.now()); + this.pendingExpiryAt.set(key, entry.expiresAt); + this.pendingTimers.set(key, setTimeout(() => { + this.pendingTimers.delete(key); + this.pendingExpiryAt.delete(key); + void this.expirePending(key); + }, delay)); + } + + private clearPendingTimer(key: string): void { + const timer = this.pendingTimers.get(key); + if (timer) clearTimeout(timer); + this.pendingTimers.delete(key); + this.pendingExpiryAt.delete(key); + } + + private async expirePending(key: string): Promise { + const entry = this.pending.get(key); + if (!entry || this.disposed || !this.started) return; + this.pending.delete(key); + this.optimisticallyResolved.add(key); + try { + const wire = this.toWireResponse(entry, "deny", "approval expired"); + await this.sendPendingResponse(entry, wire); + } catch (error) { + this.options.logger?.debug?.(`Unable to send expiry response for ${key}`, error); + } + this.emit({ + type: entry.approval ? "approval.expired" : "input.expired", + threadId: entry.threadId ?? this.threadId ?? undefined, + turnId: entry.turnId, + requestId: entry.requestId, + payload: { requestId: asJsonValue(entry.requestId), method: entry.method, reason: "approval expired" }, + }); + } + + async snapshot(): Promise { + const requests = [...this.pending.values()]; + return { + threadId: this.threadId, + turnId: this.turnId, + state: this.state, + status: { ...this.status, activeFlags: [...this.status.activeFlags] }, + activity: this.status.activity, + turnStatus: this.status.turnStatus, + activeFlags: [...this.status.activeFlags], + startedAtMs: this.status.startedAtMs ?? null, + durationMs: this.status.durationMs ?? null, + workedDurationMs: this.status.workedDurationMs ?? null, + elapsedMs: this.status.elapsedMs ?? null, + pendingApprovals: requests.map((entry) => entry.approval).filter((entry): entry is PendingApproval => Boolean(entry)), + pendingRequests: requests.map((entry) => ({ + requestId: entry.requestId, + method: entry.method, + params: redactJson(entry.params), + ...(entry.approval?.commandHash ? { commandHash: entry.approval.commandHash } : {}), + ...(entry.approval?.risk ? { risk: entry.approval.risk } : {}), + ...(entry.approval?.summary ? { summary: entry.approval.summary } : {}), + createdAt: entry.createdAt, + ...(entry.expiresAt ? { expiresAt: entry.expiresAt } : {}), + })), + outputTail: this.outputTail, + messages: asJsonValue(this.outputMessages) as JsonValue[], + subagents: this.subagents.map((subagent) => ({ ...subagent })), + metadata: { + adapter: "codex-ipc-follower", + mode: "attach", + privateProtocol: true, + socketPath: this.client.socketPath, + waitingForSession: this.waitingForSession, + attachReady: Boolean(this.threadId && this.ownerClientId && !this.waitingForSession && this.state !== "syncing"), + ...(this.ownerClientId ? { ownerClientId: this.ownerClientId } : {}), + ...(this.revision !== null ? { revision: this.revision } : {}), + ...(typeof this.conversationState.cwd === "string" ? { cwd: this.conversationState.cwd } : {}), + ...(typeof this.conversationState.title === "string" ? { title: this.conversationState.title } : {}), + ...(typeof this.conversationState.source === "string" ? { source: this.conversationState.source } : {}), + ...projectSessionMetadata(this.conversationState), + activity: this.status.activity, + turnStatus: this.status.turnStatus, + activeFlags: asJsonValue(this.status.activeFlags), + startedAtMs: this.status.startedAtMs ?? null, + durationMs: this.status.durationMs ?? null, + workedDurationMs: this.status.workedDurationMs ?? null, + elapsedMs: this.status.elapsedMs ?? null, + firstTurnWorkItemStartedAtMs: this.status.firstTurnWorkItemStartedAtMs ?? null, + finalAssistantStartedAtMs: this.status.finalAssistantStartedAtMs ?? null, + historyComplete: !hasIncompleteHistory(this.conversationState), + }, + }; + } + + onEvent(listener: (event: AgentEvent) => void): Disposable { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + async dispose(): Promise { + if (this.disposed) return; + this.sessionLifecycleGeneration += 1; + this.clearWaitingDiscoveryTimer(); + this.waitingForSession = false; + this.clearSnapshotWaiter(new Error("Codex IPC follower is stopping")); + this.clearRevisionWaiter(new Error("Codex IPC follower is stopping")); + // While the private IPC socket is still writable, close every outstanding + // owner request explicitly so stopping the bridge cannot strand an approval + // dialog in the official Codex UI. + if (this.started && this.pending.size) await this.denyPending("bridge stopped"); + this.disposed = true; + this.started = false; + this.state = "disconnected"; + this.clearPendingVscodeFollow(); + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteAwaitingSelection = false; + this.vscodeRouteClientId = null; + this.vscodeRouteCandidates.clear(); + this.activeVscodeSelection = undefined; + this.resetHistoryLoading(); + this.clearSnapshotWaiter(new Error("Codex IPC follower disposed")); + this.clearRevisionWaiter(new Error("Codex IPC follower disposed")); + for (const timer of this.pendingTimers.values()) clearTimeout(timer); + this.pendingTimers.clear(); + this.pendingExpiryAt.clear(); + if (this.threadId) { + try { + await this.client.followConversation(this.threadId, false, { + hostId: this.options.hostId, + targetClientIds: this.ownerClientId ? [this.ownerClientId] : undefined, + }); + } catch { /* socket may already be closed */ } + } + for (const entry of this.pending.values()) { + if (entry.approval) this.emit({ type: "approval.expired", threadId: entry.threadId, turnId: entry.turnId, requestId: entry.requestId, payload: { requestId: asJsonValue(entry.requestId), reason: "bridge stopped" } }); + else this.emit({ type: "input.expired", threadId: entry.threadId, turnId: entry.turnId, requestId: entry.requestId, payload: { requestId: asJsonValue(entry.requestId), reason: "bridge stopped" } }); + } + this.pending.clear(); + this.optimisticallyResolved.clear(); + for (const subscription of this.subscriptions.splice(0)) subscription.dispose(); + if (this.options.disposeClient) await this.client.dispose(); + } + + private async discoverThreadId(excluded = new Set(), maxCandidates = 64): Promise { + if (!this.options.autoDiscoverThread) return undefined; + const codexHome = resolveCodexHome(this.options); + const candidates = await recentVscodeThreadCandidates(path.join(codexHome, "sessions"), this.options.preferredCwds ?? []); + if (!candidates.length) return undefined; + // Owner discovery is the authority: a rollout file can remain on disk + // after its VS Code owner is gone. An older conversation can still be the + // one currently open in the official panel, so inspect enough candidates + // to cover its bounded recent-chat view. Eight-way batches with a short + // timeout keep the worst-case startup budget no larger than the former + // 12-candidate, four-way scan. + const limited = candidates.filter((candidate) => !excluded.has(candidate.id)).slice(0, Math.max(1, maxCandidates)); + const ownerTimeout = Math.max(250, Math.min(this.options.ownerDiscoveryTimeoutMs, 750)); + for (let offset = 0; offset < limited.length; offset += 8) { + const batch = limited.slice(offset, offset + 8); + const checks = await Promise.all(batch.map(async (candidate) => { + try { + const owner = await this.client.findThreadOwner(candidate.id, this.options.hostId, ownerTimeout); + return owner ? candidate : undefined; + } catch (error) { + this.options.logger?.debug?.(`Thread owner discovery failed for ${candidate.id}`, error); + return undefined; + } + })); + const selected = checks.find((candidate): candidate is Candidate => Boolean(candidate)); + if (selected) { + this.options.logger?.info?.(`Auto-discovered VS Code Codex conversation ${selected.id}`); + return selected.id; + } + } + return undefined; + } + + /** + * Keep the relay/IPC connection usable while the official panel has no + * selected conversation yet. The first route broadcast is handled as a + * fast path; this bounded poll covers sessions opened by older extension + * builds that do not emit the route notification. + */ + private enterWaitingForSession(): void { + if (this.disposed) return; + const wasStarted = this.started; + const wasWaiting = this.waitingForSession; + this.started = true; + this.waitingForSession = true; + this.state = "waiting_for_host"; + this.threadId = null; + this.ownerClientId = null; + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteClientId = null; + this.vscodeRouteAwaitingSelection = false; + this.vscodeRouteCandidates.clear(); + this.resetConversationProjection(); + // RelayHost normally synthesizes this event after adapter.start(). Emit it + // here as well so the adapter can be used directly and so a transition + // back to waiting clears any provisional target in existing browsers. + if (!wasWaiting) { + this.emit({ type: "connection.opened", payload: { mode: "attach", waitingForSession: true } }); + if (wasStarted) void this.emitSnapshot(); + } + this.scheduleWaitingDiscovery(); + } + + private scheduleWaitingDiscovery(delayMs = WAITING_SESSION_DISCOVERY_DELAY_MS): void { + if (this.disposed || !this.started || !this.waitingForSession || this.waitingDiscoveryTimer) return; + this.waitingDiscoveryTimer = setTimeout(() => { + this.waitingDiscoveryTimer = undefined; + void this.runWaitingDiscovery(); + }, Math.max(0, delayMs)); + } + + private clearWaitingDiscoveryTimer(): void { + if (this.waitingDiscoveryTimer) clearTimeout(this.waitingDiscoveryTimer); + this.waitingDiscoveryTimer = undefined; + } + + private async runWaitingDiscovery(): Promise { + if (this.disposed || !this.started || !this.waitingForSession || this.waitingDiscoveryInFlight) return; + const generation = this.sessionLifecycleGeneration; + const routeGeneration = this.vscodeRouteGeneration; + const officialRouteHasPriority = () => this.options.followVscodeSession + && (this.vscodeRouteGeneration !== routeGeneration + || Boolean(this.vscodeRouteClientId && this.vscodeRouteActiveThreadId)); + const task = (async () => { + // Polling is only a fallback for panels/builds that do not broadcast their + // current route. Once the official panel supplies a route, never let a + // slower filesystem/owner lookup replace it with an older recent thread. + if (officialRouteHasPriority()) return; + const configured = this.options.threadId?.trim(); + if (configured && await this.tryAttachWaitingThread(configured)) return; + if (officialRouteHasPriority()) return; + if (!this.options.autoDiscoverThread) return; + const selected = await this.discoverThreadId( + configured ? new Set([configured]) : new Set(), + WAITING_SESSION_DISCOVERY_MAX_CANDIDATES, + ); + if (selected && !officialRouteHasPriority()) await this.tryAttachWaitingThread(selected); + })(); + this.waitingDiscoveryInFlight = task; + try { + await task; + } catch (error) { + if (this.started && !this.disposed) { + this.options.logger?.debug?.("Waiting for a VS Code Codex session", error); + } + } finally { + if (this.waitingDiscoveryInFlight === task) this.waitingDiscoveryInFlight = undefined; + // A socket close/dispose may have happened while discovery was in + // flight. The generation check prevents a stale completion from + // scheduling a new timer on a dead adapter. + if (generation === this.sessionLifecycleGeneration && this.started && this.waitingForSession && !this.disposed) { + this.scheduleWaitingDiscovery(); + } + } + } + + private async tryAttachWaitingThread(target: string): Promise { + const normalizedTarget = target.trim(); + if (!normalizedTarget || this.disposed || !this.started) return false; + if (this.waitingAttachPromise) { + if (this.waitingAttachTarget !== normalizedTarget) this.queuedWaitingAttachTarget = normalizedTarget; + return this.waitingAttachTarget === normalizedTarget ? this.waitingAttachPromise : false; + } + if (!this.waitingForSession) return false; + this.waitingAttachTarget = normalizedTarget; + const task = (async () => { + const release = await this.acquireSessionOperation(); + this.sessionSwitching = true; + try { + if (this.disposed || !this.started || !this.waitingForSession) return false; + await this.attachThread(normalizedTarget, { fromWaiting: true }); + return true; + } catch (error) { + if (this.started && !this.disposed) { + this.options.logger?.debug?.(`Unable to attach waiting VS Code session ${normalizedTarget}`, error); + // attachThread restores the waiting projection for a failed + // fromWaiting attempt. Keep the retry timer alive for the next poll. + this.scheduleWaitingDiscovery(); + } + return false; + } finally { + this.sessionSwitching = false; + release(); + } + })(); + this.waitingAttachPromise = task; + let attached = false; + try { + attached = await task; + return attached; + } finally { + if (this.waitingAttachPromise === task) this.waitingAttachPromise = undefined; + this.waitingAttachTarget = null; + const queued = this.queuedWaitingAttachTarget; + this.queuedWaitingAttachTarget = null; + if (queued && !this.disposed && this.started) { + if (this.waitingForSession) { + void this.tryAttachWaitingThread(queued); + } else if (queued !== this.threadId) { + // The official panel can move A -> B while A's first snapshot is in + // flight. Preserve the latest route and reuse the normal deferred + // switching path so an active turn on A is never detached early. + this.pendingVscodeThreadId = queued; + this.pendingVscodeFollowAttempts = 0; + this.pendingVscodeFollowGeneration = this.vscodeRouteGeneration; + this.schedulePendingVscodeFollow(this.options.vscodeSessionFollowDebounceMs); + } + } + } + } + + /** + * Verify that a discovered owner can actually stream this conversation. + * Owner discovery may return a desktop client (or a stale handoff) that + * knows the id but cannot provide the follower snapshot needed for attach. + * The temporary follow is isolated to a direct stream listener and is always + * removed before returning, so probing never changes the active projection. + */ + private async probeSessionSnapshot(conversationId: string, ownerClientId: string, timeoutMs: number): Promise { + let timer: NodeJS.Timeout | undefined; + let resolveSnapshot!: (available: boolean) => void; + const snapshot = new Promise((resolve) => { + resolveSnapshot = resolve; + timer = setTimeout(() => resolve(false), timeoutMs); + }); + const subscription = this.client.onStreamEvent((event) => { + if (event.kind === "snapshot" + && event.conversationId === conversationId + && event.ownerClientId === ownerClientId) { + resolveSnapshot(true); + } + }); + try { + await this.client.followConversation(conversationId, true, { + hostId: this.options.hostId, + targetClientIds: [ownerClientId], + }); + return await snapshot; + } catch (error) { + this.options.logger?.debug?.(`Session snapshot probe failed for ${conversationId}`, error); + return false; + } finally { + if (timer) clearTimeout(timer); + subscription.dispose(); + // A user may select this row while the probe is waiting. In that case + // the temporary follow has become the active attachment; do not tear it + // down from the list request's cleanup path. + if (!(this.threadId === conversationId && this.ownerClientId === ownerClientId)) { + try { + await this.client.followConversation(conversationId, false, { + hostId: this.options.hostId, + targetClientIds: [ownerClientId], + }); + } catch (error) { + this.options.logger?.debug?.(`Unable to clean up session snapshot probe ${conversationId}`, error); + } + } + } + } + + /** + * Observe the official panel's route without reading its DOM. The official + * webview broadcasts old=false followed by new=true from one stable IPC + * client. Requiring that pair avoids treating reconnect status replays or a + * different Codex window's isolated true broadcast as a user navigation. + */ + private handleBroadcast(frame: IpcBroadcast): void { + if (!this.options.followVscodeSession || !this.started || this.disposed) return; + if (frame.method === "client-status-changed") { + this.handleRouteClientStatus(frame); + return; + } + if (frame.method !== "thread-stream-following-changed") return; + if (this.options.strictVersions !== false + && frame.version !== CODEX_IPC_METHOD_VERSIONS["thread-stream-following-changed"]) return; + // Status replies sent to a newly connected follower describe retained + // subscriptions, not a fresh route change in the VS Code panel. + if (frame.targetClientIds?.length) return; + if (!isRecord(frame.params)) return; + const conversationId = stringValue(frame.params.conversationId)?.trim(); + const hostId = stringValue(frame.params.hostId) ?? "local"; + const sourceClientId = frame.sourceClientId?.trim(); + if (!conversationId || hostId !== this.options.hostId || !sourceClientId) return; + if (sourceClientId === this.client.getClientId()) return; + + // There is no trusted old route while the first conversation is missing. + // Treat an untargeted `following:true` as a candidate hint, then verify it + // through owner discovery and an authoritative snapshot before attaching. + if (this.waitingForSession || this.waitingAttachPromise) { + if (frame.params.following !== true) return; + if (this.vscodeRouteClientId && sourceClientId !== this.vscodeRouteClientId) return; + this.vscodeRouteClientId = sourceClientId; + if (this.vscodeRouteActiveThreadId !== conversationId) { + this.vscodeRouteActiveThreadId = conversationId; + this.vscodeRouteGeneration += 1; + } + void this.tryAttachWaitingThread(conversationId); + return; + } + + // Bind only after one source completes old=false -> new=true. A different + // remote follower can emit an isolated false while disposing, and that + // must not permanently steal the trusted route source. + if (!this.vscodeRouteClientId) { + if (frame.params.following === false) { + if (conversationId === this.vscodeRouteActiveThreadId) { + this.vscodeRouteCandidates.set(sourceClientId, conversationId); + } + return; + } + if (frame.params.following !== true) return; + if (this.vscodeRouteCandidates.get(sourceClientId) !== this.vscodeRouteActiveThreadId) return; + this.vscodeRouteClientId = sourceClientId; + this.vscodeRouteCandidates.clear(); + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteAwaitingSelection = true; + this.vscodeRouteGeneration += 1; + } + if (sourceClientId !== this.vscodeRouteClientId) return; + + if (frame.params.following === false) { + if (conversationId !== this.vscodeRouteActiveThreadId) return; + this.vscodeRouteGeneration += 1; + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteAwaitingSelection = true; + if (this.activeVscodeSelection?.target === conversationId) { + this.clearSnapshotWaiter(new Error(`VS Code moved away from conversation ${conversationId}`)); + } + // A -> B -> C can happen faster than B's snapshot. Cancel the queued B + // as soon as the official route reports that B is no longer active. + if (this.pendingVscodeThreadId === conversationId) this.clearPendingVscodeFollow(); + return; + } + if (frame.params.following !== true) return; + + if (conversationId === this.vscodeRouteActiveThreadId) return; + if (!this.vscodeRouteAwaitingSelection) { + // Initial/reconnect following status. It may confirm the current route, + // but an isolated true is intentionally never treated as navigation. + if (conversationId === this.threadId) this.vscodeRouteActiveThreadId = conversationId; + return; + } + this.vscodeRouteAwaitingSelection = false; + this.vscodeRouteActiveThreadId = conversationId; + this.vscodeRouteGeneration += 1; + if (conversationId === this.threadId) { + this.clearPendingVscodeFollow(); + return; + } + this.pendingVscodeThreadId = conversationId; + this.pendingVscodeFollowAttempts = 0; + this.pendingVscodeFollowGeneration = this.vscodeRouteGeneration; + this.schedulePendingVscodeFollow(this.options.vscodeSessionFollowDebounceMs); + } + + /** + * The official webview gets a new IPC client id after its socket reconnects. + * Drop only the trusted route-source binding when that client disconnects; + * the next same-source false -> true pair can then establish the replacement. + */ + private handleRouteClientStatus(frame: IpcBroadcast): void { + if (this.options.strictVersions !== false && frame.version !== 0) return; + if (!isRecord(frame.params)) return; + const clientId = (stringValue(frame.params.clientId) ?? frame.sourceClientId)?.trim(); + const status = stringValue(frame.params.status)?.trim().toLowerCase(); + if (!clientId || status !== "disconnected") return; + this.vscodeRouteCandidates.delete(clientId); + if (clientId !== this.vscodeRouteClientId) return; + + this.options.logger?.debug?.(`VS Code Codex route client ${clientId} disconnected; awaiting a replacement source`); + this.vscodeRouteClientId = null; + this.vscodeRouteCandidates.clear(); + this.vscodeRouteAwaitingSelection = false; + // Keep the route identity, rather than forcing it back to the currently + // committed relay thread. A disconnect may occur between B's true signal + // and B's snapshot; the replacement panel will later leave B with false. + this.vscodeRouteActiveThreadId = this.activeVscodeSelection?.target + ?? this.vscodeRouteActiveThreadId + ?? this.threadId; + this.vscodeRouteGeneration += 1; + this.clearPendingVscodeFollow(); + if (this.activeVscodeSelection) { + this.clearSnapshotWaiter(new Error("VS Code Codex route client disconnected during session selection")); + } + } + + private schedulePendingVscodeFollow(delayMs: number): void { + if (!this.pendingVscodeThreadId || this.disposed || !this.started) return; + if (this.vscodeSessionFollowTimer) clearTimeout(this.vscodeSessionFollowTimer); + this.vscodeSessionFollowTimer = setTimeout(() => { + this.vscodeSessionFollowTimer = undefined; + void this.applyPendingVscodeFollow(); + }, Math.max(0, delayMs)); + } + + private async applyPendingVscodeFollow(): Promise { + const target = this.pendingVscodeThreadId; + const routeGeneration = this.pendingVscodeFollowGeneration; + if (!target || this.disposed || !this.started) return; + if (this.vscodeRouteActiveThreadId !== target || routeGeneration !== this.vscodeRouteGeneration) { + if (this.pendingVscodeThreadId === target) this.clearPendingVscodeFollow(); + return; + } + if (target === this.threadId) { + this.clearPendingVscodeFollow(); + return; + } + // Never detach a turn or approval that is still owned by the old thread. + // Keep only the latest official target and retry once that state is idle. + if (this.sessionSwitching || this.turnId || this.pending.size) { + this.schedulePendingVscodeFollow(VSCODE_SESSION_FOLLOW_RETRY_DELAY_MS); + return; + } + + this.pendingVscodeThreadId = null; + this.pendingVscodeFollowGeneration = 0; + this.activeVscodeSelection = { target, generation: routeGeneration }; + try { + await this.selectSession({ + threadId: target, + origin: "vscode", + vscodeRouteGeneration: routeGeneration, + }); + this.pendingVscodeFollowAttempts = 0; + this.options.logger?.info?.(`Followed VS Code Codex panel to conversation ${target}`); + } catch (error) { + this.options.logger?.debug?.(`Unable to follow VS Code Codex panel to ${target}`, error); + if (!this.disposed + && this.started + && !this.pendingVscodeThreadId + && this.vscodeRouteActiveThreadId === target + && this.vscodeRouteGeneration === routeGeneration + && this.pendingVscodeFollowAttempts < VSCODE_SESSION_FOLLOW_MAX_ATTEMPTS) { + this.pendingVscodeFollowAttempts += 1; + this.pendingVscodeThreadId = target; + this.pendingVscodeFollowGeneration = routeGeneration; + } + } finally { + if (this.activeVscodeSelection?.target === target + && this.activeVscodeSelection.generation === routeGeneration) { + this.activeVscodeSelection = undefined; + } + if (this.pendingVscodeThreadId) { + this.schedulePendingVscodeFollow(VSCODE_SESSION_FOLLOW_RETRY_DELAY_MS); + } + } + } + + private clearPendingVscodeFollow(): void { + if (this.vscodeSessionFollowTimer) clearTimeout(this.vscodeSessionFollowTimer); + this.vscodeSessionFollowTimer = undefined; + this.pendingVscodeThreadId = null; + this.pendingVscodeFollowAttempts = 0; + this.pendingVscodeFollowGeneration = 0; + } + + private handleStreamEvent(event: ConversationStreamEvent): void { + if (!this.threadId || event.conversationId !== this.threadId) return; + // Following is targeted at the owner discovered during attach. The IPC + // router normally filters these broadcasts, but a handoff/old socket can + // still surface another client's event; applying it would overwrite the + // active conversation and could satisfy a revision waiter incorrectly. + if (this.ownerClientId && event.ownerClientId && event.ownerClientId !== this.ownerClientId) { + this.options.logger?.debug?.(`Ignoring stream event from unexpected Codex owner ${event.ownerClientId}`); + return; + } + if (event.kind === "desync") { + this.options.logger?.warn?.(`Codex IPC stream desynchronized at revision ${event.receivedBaseRevision}; requesting snapshot`); + this.client.followConversation(this.threadId, true, { hostId: this.options.hostId, targetClientIds: this.ownerClientId ? [this.ownerClientId] : undefined }).catch((error) => this.options.logger?.warn?.("Unable to recover IPC snapshot", error)); + return; + } + this.ownerClientId = event.ownerClientId || this.ownerClientId; + this.revision = event.revision; + if (this.revisionWaiter + && this.revisionWaiter.threadId === event.conversationId + && this.revisionWaiter.ownerClientId === event.ownerClientId + && event.revision >= this.revisionWaiter.revision) { + this.clearRevisionWaiter(); + } + this.conversationState = cloneObject(event.conversationState); + this.processConversationState(event.kind === "snapshot"); + if (event.kind === "snapshot" && this.options.loadCompleteHistory !== false) void this.loadCompleteHistoryIfNeeded(); + if (event.kind === "snapshot" + && this.snapshotWaiter?.threadId === event.conversationId + && (!this.snapshotWaiter.ownerClientId || this.snapshotWaiter.ownerClientId === event.ownerClientId)) { + this.clearSnapshotWaiter(); + } + } + + private processConversationState(initial: boolean): void { + if (initial) this.snapshotSeen = true; + const previousTurnId = this.turnId; + const previousState = this.state; + const previousStatus = this.status; + const nextTurn = deriveTurn(this.conversationState); + // Request records and runtime flags are part of the same conversation + // snapshot. Derive status after extracting them so an approval/input wait + // is visible immediately with the corresponding state patch. + const nextRequests = extractRequests(this.conversationState, this.options.approvalTimeoutMs); + this.status = deriveStatusSnapshot(this.conversationState, nextTurn, nextRequests); + this.turnId = nextTurn.active ? nextTurn.id ?? null : null; + this.state = nextTurn.active ? "active" : this.deriveSessionState(); + const nextHistoryComplete = !hasIncompleteHistory(this.conversationState); + const historyChanged = this.historyComplete !== undefined && this.historyComplete !== nextHistoryComplete; + // Settings updates are broadcast as conversation-state patches and may not + // change output, turn status, or pending requests. Fingerprint only the + // redacted projection exposed by `snapshot()` so those patches still reach + // the relay without leaking opaque/private state or causing per-token + // snapshot spam. + const metadataShape = stableStringify(projectSessionMetadata(this.conversationState)); + const metadataChanged = metadataShape !== this.renderedMetadataShape; + + const previousPending = new Map(this.pending); + const nextKeys = new Set(nextRequests.map((entry) => jsonRpcIdKey(entry.requestId))); + for (const key of this.optimisticallyResolved) { + if (!nextKeys.has(key)) this.optimisticallyResolved.delete(key); + } + const previousKeys = new Set(previousPending.keys()); + for (const key of this.pendingTimers.keys()) { + if (!nextKeys.has(key)) this.clearPendingTimer(key); + } + this.pending.clear(); + for (const rawEntry of nextRequests) { + const prior = previousPending.get(jsonRpcIdKey(rawEntry.requestId)); + // Some official request records omit timestamps. Keep the first-seen + // deadline across patches instead of extending it on every output delta. + const entry = prior + ? { ...rawEntry, createdAt: prior.createdAt, expiresAt: prior.expiresAt ?? rawEntry.expiresAt } + : rawEntry; + const key = jsonRpcIdKey(entry.requestId); + if (this.optimisticallyResolved.has(key)) continue; + this.pending.set(key, entry); + this.schedulePendingExpiry(key, entry); + if (!previousKeys.has(key)) this.emitRequest(entry); + } + for (const key of previousKeys) { + if (nextKeys.has(key) || this.optimisticallyResolved.has(key)) continue; + const old = previousPending.get(key); + if (old) this.emit({ type: old.approval ? "approval.resolved" : "input.resolved", threadId: old.threadId ?? this.threadId ?? undefined, turnId: old.turnId, requestId: old.requestId, payload: { requestId: asJsonValue(old.requestId), method: old.method } }); + } + + const rendered = renderConversationOutput(this.conversationState, this.options.maxOutputTailChars); + const output = rendered.text; + const messageShape = renderedMessageShape(rendered.messages); + const messagesChanged = messageShape !== this.renderedMessageShape; + const messagesPatch = messagesChanged + ? renderedMessagesPatch(this.outputMessages, rendered.messages) + : undefined; + const subagentShape = renderedSubagentShape(rendered.subagents); + const subagentsChanged = subagentShape !== this.renderedSubagentShape; + if (output !== this.renderedOutput || rendered.totalLength !== this.renderedOutputLength || messagesChanged || subagentsChanged) { + const delta = this.renderedOutput + ? appendOnlyOutputDelta( + this.renderedOutput, + this.renderedOutputLength, + output, + rendered.totalLength, + this.renderedOutputWasTruncated, + ) + : undefined; + if (!this.renderedOutput || delta === undefined || (!delta && (messagesChanged || subagentsChanged))) { + this.emit({ type: "output.snapshot", threadId: this.threadId ?? undefined, turnId: this.turnId ?? undefined, payload: { stream: "codex", text: output, messages: asJsonValue(rendered.messages), subagents: asJsonValue(rendered.subagents), structureChanged: true, encoding: "utf8" } }); + } else { + if (delta || subagentsChanged) this.emit({ + type: "output.chunk", + threadId: this.threadId ?? undefined, + turnId: this.turnId ?? undefined, + payload: { + stream: "codex", + text: delta, + // Keep the append-only field for older relays, but include the + // authoritative projection so a browser can preserve item + // boundaries while reasoning/commands/edits stream in. + outputTail: output, + ...(messagesPatch ? { messagesPatch: asJsonValue(messagesPatch) } : {}), + subagents: asJsonValue(rendered.subagents), + structureChanged: messagesChanged || subagentsChanged, + encoding: "utf8", + }, + }); + } + this.renderedOutput = output; + this.renderedOutputLength = rendered.totalLength; + this.renderedOutputWasTruncated = rendered.truncated; + this.renderedMessageShape = messageShape; + this.renderedSubagentShape = subagentShape; + this.outputTail = output; + this.outputMessages = rendered.messages; + this.subagents = rendered.subagents; + } + + if (!initial && !previousTurnId && this.turnId) { + this.emit({ type: "task.started", threadId: this.threadId ?? undefined, turnId: this.turnId, payload: statusPayload(this.status) }); + } else if (!initial && previousTurnId && !this.turnId) { + const cancelled = new Set(["cancelled", "canceled", "interrupted"]).has(normalizeStatus(nextTurn.status)); + this.emit({ type: cancelled ? "task.cancelled" : "task.finished", threadId: this.threadId ?? undefined, turnId: previousTurnId, payload: statusPayload(this.status) }); + } else if (!initial && !sameStatus(previousStatus, this.status)) { + // A turn can remain active while moving from reasoning to a command or + // file edit, and can enter/leave an approval wait without changing its + // id. Publish a dedicated status event so remote viewers do not have to + // infer activity from output timing. + this.emit({ type: "task.status", threadId: this.threadId ?? undefined, turnId: this.turnId ?? undefined, payload: statusPayload(this.status) }); + } + + this.historyComplete = nextHistoryComplete; + this.renderedMetadataShape = metadataShape; + if (initial || metadataChanged || historyChanged || previousTurnId !== this.turnId || previousState !== this.state || !sameStatus(previousStatus, this.status) || nextRequests.length !== previousKeys.size || subagentsChanged) { + void this.emitSnapshot(); + } + } + + private async loadCompleteHistoryIfNeeded(): Promise { + if (this.historyLoadRequested || this.historyLoadRetryTimer || this.historyLoadAttempts >= 2 || !this.threadId || !this.ownerClientId || !hasIncompleteHistory(this.conversationState)) return; + this.clearHistoryLoadRetryTimer(); + // Keep the request target stable. A stream event from another owner can + // arrive while the owner is loading history; using the mutable fields + // below would otherwise wait for (or acknowledge) the wrong stream. + const threadId = this.threadId; + const ownerClientId = this.ownerClientId; + const generation = this.historyLoadGeneration; + this.historyLoadRequested = true; + this.historyLoadAttempts += 1; + try { + const result = await this.client.loadCompleteHistory(threadId, { + ownerClientId, + timeoutMs: this.options.followTimeoutMs, + }); + // The socket may close (or the owner may hand the conversation to a new + // client) while the request is in flight. Do not install a fresh waiter + // after dispose/close, where nobody could clear it. + if (!this.isCurrentHistoryLoad(threadId, ownerClientId, generation)) return; + const requestedRevision = extractRevision(result); + // The owner acknowledges the load request before broadcasting its new + // state. Wait for that revision so the relay's first snapshot contains + // the complete history rather than the old paginated tail. + if (requestedRevision !== undefined && (this.revision ?? 0) < requestedRevision) { + await this.waitForRevision(threadId, ownerClientId, requestedRevision, this.options.followTimeoutMs); + } + // A session switch can resolve the old revision waiter while installing + // a new history request. Never let that old continuation clear or retry + // the new conversation's loading state. + if (!this.isCurrentHistoryLoad(threadId, ownerClientId, generation)) return; + this.historyLoadRequested = false; + if (this.started + && this.threadId === threadId + && this.ownerClientId === ownerClientId + && hasIncompleteHistory(this.conversationState) + && this.historyLoadAttempts < 2) { + void this.loadCompleteHistoryIfNeeded(); + } + } catch (error) { + // History loading is an optional read enhancement. The live tail remains + // usable when an older extension does not implement this request. + this.options.logger?.debug?.("Unable to load complete Codex history", error); + if (!this.isCurrentHistoryLoad(threadId, ownerClientId, generation)) return; + this.historyLoadRequested = false; + if (isTransientHistoryLoadError(error) + && this.historyLoadAttempts < 2 + && hasIncompleteHistory(this.conversationState)) { + this.scheduleHistoryLoadRetry(threadId, ownerClientId, generation); + } + } + } + + private scheduleHistoryLoadRetry(threadId: string, ownerClientId: string, generation: number): void { + if (this.historyLoadRetryTimer || !this.isCurrentHistoryLoad(threadId, ownerClientId, generation)) return; + this.historyLoadRetryTimer = setTimeout(() => { + this.historyLoadRetryTimer = undefined; + if (!this.isCurrentHistoryLoad(threadId, ownerClientId, generation) + || this.historyLoadRequested + || this.historyLoadAttempts >= 2 + || !hasIncompleteHistory(this.conversationState)) return; + void this.loadCompleteHistoryIfNeeded(); + }, HISTORY_LOAD_RETRY_DELAY_MS); + } + + private isCurrentHistoryLoad(threadId: string, ownerClientId: string, generation: number): boolean { + return !this.disposed + && this.started + && this.historyLoadGeneration === generation + && this.threadId === threadId + && this.ownerClientId === ownerClientId; + } + + private clearHistoryLoadRetryTimer(): void { + if (!this.historyLoadRetryTimer) return; + clearTimeout(this.historyLoadRetryTimer); + this.historyLoadRetryTimer = undefined; + } + + private resetHistoryLoading(): void { + this.clearHistoryLoadRetryTimer(); + this.historyLoadRequested = false; + this.historyLoadAttempts = 0; + this.historyLoadGeneration += 1; + } + + private emitRequest(entry: RequestEntry): void { + const isInput = INPUT_METHODS.has(entry.method); + this.emit({ + type: isInput ? "input.requested" : APPROVAL_METHODS.has(entry.method) ? "approval.requested" : "server.requested", + threadId: entry.threadId ?? this.threadId ?? undefined, + turnId: entry.turnId, + requestId: entry.requestId, + payload: { + requestId: asJsonValue(entry.requestId), + method: entry.method, + params: redactJson(entry.params), + ...(entry.approval ? { + action: entry.approval.action, + risk: entry.approval.risk, + summary: entry.approval.summary, + ...(entry.approval.commandHash ? { commandHash: entry.approval.commandHash } : {}), + } : {}), + ...(entry.expiresAt ? { expiresAt: entry.expiresAt } : {}), + }, + raw: redactJson({ requestId: entry.requestId, method: entry.method, params: entry.params }), + }); + } + + private async emitSnapshot(): Promise { + const snapshot = await this.snapshot(); + this.emit({ type: "session.snapshot", threadId: snapshot.threadId ?? undefined, turnId: snapshot.turnId ?? undefined, payload: asJsonObject(snapshot) }); + } + + private deriveSessionState(): string { + const runtime = this.conversationState.threadRuntimeStatus; + if (isRecord(runtime) && typeof runtime.type === "string") { + const status = normalizeStatus(runtime.type); + if (!TERMINAL_TURN_STATES.has(status) && status !== "idle" && status !== "ready") return runtime.type; + } + return this.snapshotSeen ? "idle" : "syncing"; + } + + private waitForSnapshot(threadId: string, timeoutMs: number, ownerClientId?: string): Promise { + this.clearSnapshotWaiter(); + return new Promise((resolve, reject) => { + const timer = setTimeout(() => { + if (this.snapshotWaiter?.threadId === threadId + && this.snapshotWaiter.ownerClientId === ownerClientId) this.snapshotWaiter = undefined; + reject(new Error(`Timed out waiting for a snapshot from VS Code conversation ${threadId}`)); + }, timeoutMs); + this.snapshotWaiter = { threadId, ownerClientId, resolve, reject, timer }; + }); + } + + /** Rehydrate the adapter's last owner-validated state after failed navigation. */ + private restoreCachedConversationProjection( + cached: ConversationStreamState | undefined, + fallbackOwnerClientId: string | null, + ): boolean { + if (!cached || !cached.conversationState) return false; + this.ownerClientId = cached.ownerClientId || fallbackOwnerClientId; + this.revision = cached.revision; + this.conversationState = cloneObject(cached.conversationState); + this.processConversationState(true); + this.state = this.deriveSessionState(); + return true; + } + + private waitForRevision(threadId: string, ownerClientId: string, revision: number, timeoutMs: number): Promise { + this.clearRevisionWaiter(); + if (this.threadId === threadId && this.ownerClientId === ownerClientId && (this.revision ?? 0) >= revision) return Promise.resolve(); + return new Promise((resolve, reject) => { + const timer = setTimeout(() => { + if (this.revisionWaiter?.threadId === threadId + && this.revisionWaiter.ownerClientId === ownerClientId + && this.revisionWaiter.revision === revision) this.revisionWaiter = undefined; + reject(new Error(`Timed out waiting for Codex conversation revision ${revision}`)); + }, timeoutMs); + this.revisionWaiter = { threadId, ownerClientId, revision, resolve, reject, timer }; + }); + } + + private clearSnapshotWaiter(error?: Error): void { + const waiter = this.snapshotWaiter; + if (!waiter) return; + clearTimeout(waiter.timer); + this.snapshotWaiter = undefined; + if (error) waiter.reject(error); + else waiter.resolve(); + } + + private clearRevisionWaiter(error?: Error): void { + const waiter = this.revisionWaiter; + if (!waiter) return; + clearTimeout(waiter.timer); + this.revisionWaiter = undefined; + if (error) waiter.reject(error); + else waiter.resolve(); + } + + /** Clear all projections that belong to the previous conversation. */ + private resetConversationProjection(): void { + this.clearRevisionWaiter(); + this.revision = null; + this.conversationState = {}; + this.turnId = null; + this.status = { + activity: "idle", + turnStatus: "idle", + activeFlags: [], + startedAtMs: null, + durationMs: null, + workedDurationMs: null, + elapsedMs: null, + firstTurnWorkItemStartedAtMs: null, + finalAssistantStartedAtMs: null, + }; + this.renderedOutput = ""; + this.renderedOutputLength = 0; + this.renderedOutputWasTruncated = false; + this.renderedMessageShape = ""; + this.outputTail = ""; + this.outputMessages = []; + this.subagents = []; + this.renderedSubagentShape = ""; + this.renderedMetadataShape = ""; + this.snapshotSeen = false; + this.historyComplete = undefined; + this.resetHistoryLoading(); + } + + private handleClose(error?: Error): void { + if (!this.started || this.disposed) return; + const closedThreadId = this.threadId; + const closedOwnerClientId = this.ownerClientId; + this.sessionLifecycleGeneration += 1; + this.clearWaitingDiscoveryTimer(); + this.waitingForSession = false; + this.started = false; + this.state = "disconnected"; + this.clearPendingVscodeFollow(); + this.vscodeRouteActiveThreadId = null; + this.vscodeRouteAwaitingSelection = false; + this.vscodeRouteClientId = null; + this.vscodeRouteCandidates.clear(); + this.activeVscodeSelection = undefined; + this.clearSnapshotWaiter(error ?? new Error("Codex IPC socket closed")); + this.clearRevisionWaiter(error ?? new Error("Codex IPC socket closed")); + for (const timer of this.pendingTimers.values()) clearTimeout(timer); + this.pendingTimers.clear(); + this.pendingExpiryAt.clear(); + const pending = [...this.pending.values()]; + this.pending.clear(); + this.optimisticallyResolved.clear(); + // Do not retain a conversation projection after the private IPC owner has + // gone away. A later snapshot request must report a disconnected, empty + // follower rather than stale messages that can no longer be controlled. + this.resetConversationProjection(); + this.threadId = null; + this.ownerClientId = null; + this.status = { + ...this.status, + activity: "idle", + turnStatus: "disconnected", + activeFlags: [], + }; + for (const entry of pending) { + this.emit({ + type: entry.approval ? "approval.expired" : "input.expired", + threadId: entry.threadId ?? closedThreadId ?? undefined, + turnId: entry.turnId, + requestId: entry.requestId, + payload: { requestId: asJsonValue(entry.requestId), reason: error?.message ?? "IPC socket closed" }, + }); + } + this.emit({ + type: "connection.closed", + threadId: closedThreadId ?? undefined, + payload: { + message: error?.message ?? "IPC socket closed", + ...(closedOwnerClientId ? { ownerClientId: closedOwnerClientId } : {}), + }, + }); + } + + private ensureAttached(): void { + if (!this.started) throw new Error("Codex remote bridge is not attached to a VS Code Codex session"); + if (this.waitingForSession || this.state === "waiting_for_host") { + throw new Error("waiting_for_session: 请先在 VS Code 打开一个 Codex 会话"); + } + if (!this.threadId || !this.ownerClientId) throw new Error("No existing Codex conversation owner is attached"); + } + + private ensureStarted(): void { + if (!this.started) throw new Error("Codex remote bridge is not connected to the VS Code IPC host"); + } + + private ensureInteractiveReady(): void { + this.ensureAttached(); + if (this.sessionSwitching || this.state === "syncing") throw new Error("Codex session switch is still in progress"); + } + + private emit(event: AgentEvent): void { + // Every relay event carries the latest typed execution projection. Keep + // the same values in payload for older consumers that only inspect the + // untyped event envelope. + const eventStatus = event.status ?? this.status; + const normalized: AgentEvent = { + ...event, + status: { ...eventStatus, activeFlags: [...eventStatus.activeFlags] }, + payload: { ...statusPayload(eventStatus), ...event.payload }, + }; + for (const listener of this.listeners) { + try { + listener(normalized); + } catch (error) { + this.options.logger?.warn?.("Codex IPC adapter listener failed", error); + } + } + } + + private toWireResponse(entry: RequestEntry, decision: "allow" | "deny" | "cancel", reason?: string, response?: JsonValue): JsonValue { + if (entry.method === "item/permissions/requestApproval") { + const supplied = isRecord(response) ? response : {}; + const requested = isRecord(supplied.permissions) ? supplied.permissions : decision === "allow" && isRecord(entry.params.permissions) ? entry.params.permissions : {}; + const permissions: JsonObject = {}; + for (const [key, value] of Object.entries(requested)) if (value !== null && value !== undefined) permissions[key] = asJsonValue(value); + return { permissions, scope: supplied.scope === "session" ? "session" : "turn", ...(typeof supplied.strictAutoReview === "boolean" ? { strictAutoReview: supplied.strictAutoReview } : {}) }; + } + if (entry.method === "item/tool/requestUserInput") { + return normalizeUserInputResponse(response); + } + if (entry.method === "mcpServer/elicitation/request") { + if (isRecord(response) && typeof response.action === "string") return response; + return { action: decision === "allow" ? "accept" : decision === "cancel" ? "cancel" : "decline", content: null, _meta: null }; + } + const suppliedDecision = isRecord(response) && Object.prototype.hasOwnProperty.call(response, "decision") ? response.decision : undefined; + if (suppliedDecision !== undefined) return normalizeFollowerDecision(entry.method, suppliedDecision, decision, reason); + if (entry.method === "applyPatchApproval" || entry.method === "execCommandApproval") { + if (decision === "allow") return "approved"; + if (decision === "cancel") return "abort"; + return { denied: { rejection: reason || "Denied remotely" } }; + } + return decision === "allow" ? "accept" : decision === "cancel" ? "cancel" : "decline"; + } +} + +function normalizeFollowerDecision(method: string, supplied: unknown, fallback: "allow" | "deny" | "cancel", reason?: string): JsonValue { + // The relay accepts compatibility aliases, while the private follower + // methods use the app-server's method-specific wire vocabulary. + if (method === "item/commandExecution/requestApproval" || method === "item/fileChange/requestApproval") { + if (supplied === "approved" || supplied === "approved_for_session" || supplied === "approved_mcp_policy_amendment") { + return supplied === "approved" ? "accept" : supplied === "approved_for_session" ? "acceptForSession" : "accept"; + } + if (supplied === "denied" || supplied === "deny") return "decline"; + if (supplied === "abort") return "cancel"; + return asJsonValue(supplied); + } + if (method === "applyPatchApproval" || method === "execCommandApproval") { + if (supplied === "accept") return "approved"; + if (supplied === "acceptForSession") return "approved_for_session"; + if (supplied === "decline" || supplied === "deny") return { denied: { rejection: reason || "Denied remotely" } }; + if (supplied === "cancel") return "abort"; + return asJsonValue(supplied); + } + // If a caller supplied a generic wrapper with no method-specific alias, + // retain the ordinary fallback selected by the adapter. + return asJsonValue(supplied ?? (fallback === "allow" ? "accept" : fallback === "cancel" ? "cancel" : "decline")); +} + +/** Normalize the official tool-input shape: answers[id].answers is a string array. */ +function normalizeUserInputResponse(response: JsonValue | undefined): JsonObject { + const hasOuterAnswers = isRecord(response) && isRecord(response.answers); + const source = hasOuterAnswers ? response.answers as Record : isRecord(response) ? response : {}; + const answers: JsonObject = {}; + for (const [questionId, raw] of Object.entries(source)) { + if (isRecord(raw) && Array.isArray(raw.answers)) { + answers[questionId] = { answers: raw.answers.map(asJsonValue) }; + } else if (Array.isArray(raw)) { + answers[questionId] = { answers: raw.map(asJsonValue) }; + } else if (typeof raw === "string") { + answers[questionId] = { answers: [raw] }; + } else if (raw !== undefined && raw !== null) { + answers[questionId] = { answers: [asJsonValue(raw)] }; + } else { + answers[questionId] = { answers: [] }; + } + } + return { answers }; +} + +interface Candidate { + id: string; + mtime: number; + updatedAtMs?: number; + cwd?: string; + title?: string; + priority: number; +} + +async function recentVscodeThreadCandidates(root: string, preferredCwds: string[] = []): Promise { + const files: Array<{ file: string; mtime: number; id: string }> = []; + async function visit(directory: string, depth: number): Promise { + if (depth > 3) return; + let entries; + try { entries = await fs.readdir(directory, { withFileTypes: true }); } catch { return; } + await Promise.all(entries.map(async (entry) => { + const full = path.join(directory, entry.name); + if (entry.isDirectory()) return visit(full, depth + 1); + if (!entry.isFile() || !entry.name.endsWith(".jsonl")) return; + // Codex has used UUIDv7 rollouts (the current hyphenated form), compact + // UUIDs, and 26-character ULIDs across desktop/VS Code builds. Keep the + // suffix strict so an arbitrary prompt-like filename cannot become a + // selectable session. + const match = entry.name.match(/(?:^|-)((?:[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}|[0-9a-f]{32}|[0-9a-z]{26}))(?:_[^/]*)?\.jsonl$/i); + if (!match) return; + try { const stat = await fs.stat(full); files.push({ file: full, mtime: stat.mtimeMs, id: match[1] }); } catch { /* race */ } + })); + } + await visit(root, 0); + files.sort((a, b) => b.mtime - a.mtime); + const result: Candidate[] = []; + // Rollout directories can contain many old sessions. Read metadata from a + // generous bounded set so an active VS Code session is not hidden merely + // because desktop history was written more recently. + for (const file of files.slice(0, 2_000)) { + try { + const firstLine = await readFirstLine(file.file); + let payload: Record = {}; + try { + const record = JSON.parse(firstLine) as unknown; + if (isRecord(record)) { + // Current rollouts wrap metadata in `payload`; older writers put the + // same fields on the first record itself. Accept only records that + // actually contain session identity fields so an arbitrary event + // record cannot become a false discovery candidate. + const candidate = isRecord(record.payload) ? record.payload : record; + if (["originator", "source", "thread_source", "cwd"].some((key) => Object.prototype.hasOwnProperty.call(candidate, key))) { + payload = candidate; + } + } + } catch { + // Incomplete rollout records are ignored below. A regex fallback is + // intentionally omitted so text inside a long prompt cannot identify + // an unrelated session as a VS Code conversation. + } + const originator = stringValue(payload.originator); + const source = stringValue(payload.source); + const threadSource = stringValue(payload.thread_source); + const cwd = stringValue(payload.cwd); + const title = stringValue(payload.title) + ?? stringValue(payload.thread_title) + ?? stringValue(payload.threadTitle) + ?? stringValue(payload.name) + ?? stringValue(payload.thread_name); + const updatedAtMs = timestampMs(payload.updated_at_ms) + ?? timestampMs(payload.updatedAtMs) + ?? timestampMs(payload.updated_at) + ?? timestampMs(payload.updatedAt) + ?? timestampMs(payload.timestamp); + // Subagent rollouts can use the same `codex_vscode` originator as their + // user-facing parent, but they are not conversations the operator opened. + if (threadSource === "subagent") continue; + // `source: vscode` is also written by Codex Desktop tasks hosted from a + // VS Code-shaped workspace. An explicit non-VS-Code originator must win + // so attach mode never follows a desktop task merely because it shares + // the same IPC router. Keep compatibility with older official rollouts + // that omitted originator altogether. + const official = originator === "codex_vscode"; + const unattributedVscode = !originator && (source === "vscode" || threadSource === "vscode"); + const isVscode = official || unattributedVscode; + if (!isVscode) continue; + // Never select a rollout produced by this bridge's legacy spawn mode. + if (originator === "codex-remote-collab") continue; + const cwdMatch = Boolean(cwd && preferredCwds.some((candidate) => samePath(candidate, cwd))); + const id = file.id; + // The bridge runs inside a VS Code workspace, so an owner in that + // workspace is a stronger signal than the rollout writer's originator. + const priority = cwdMatch ? (official ? 0 : 1) : (official ? 2 : 3); + if (!result.some((candidate) => candidate.id === id)) { + result.push({ + id, + mtime: file.mtime, + ...(updatedAtMs !== undefined ? { updatedAtMs } : {}), + ...(cwd ? { cwd } : {}), + ...(title ? { title } : {}), + priority, + }); + } + } catch { /* ignore incomplete/rotated rollout files */ } + } + result.sort((a, b) => a.priority - b.priority || b.mtime - a.mtime); + return result; +} + +/** Read enough of a rollout file to reach its first JSONL record, bounded. */ +async function readFirstLine(fileName: string, maxBytes = 4 * 1024 * 1024): Promise { + const handle = await fs.open(fileName, "r"); + const parts: Buffer[] = []; + let offset = 0; + try { + while (offset < maxBytes) { + const size = Math.min(64 * 1024, maxBytes - offset); + const chunk = Buffer.alloc(size); + const read = await handle.read(chunk, 0, size, offset); + if (!read.bytesRead) break; + const piece = chunk.subarray(0, read.bytesRead); + const newline = piece.indexOf(0x0a); + if (newline >= 0) { + parts.push(piece.subarray(0, newline)); + return Buffer.concat(parts).toString("utf8"); + } + parts.push(piece); + offset += read.bytesRead; + if (read.bytesRead < size) break; + } + return Buffer.concat(parts).toString("utf8"); + } finally { + await handle.close(); + } +} + +interface SessionIndexEntry { + title?: string; + updatedAtMs?: number; + cwd?: string; +} + +/** Read the bounded local index used by the official Codex recent-chat list. */ +async function readSessionIndex(fileName: string, maxBytes = 8 * 1024 * 1024): Promise> { + let raw: string; + try { + raw = await readBoundedText(fileName, maxBytes); + } catch { + return new Map(); + } + const result = new Map(); + for (const line of raw.split(/\r?\n/)) { + if (!line.trim()) continue; + try { + const record = JSON.parse(line) as unknown; + if (!isRecord(record)) continue; + const id = stringValue(record.id) ?? stringValue(record.session_id) ?? stringValue(record.thread_id); + if (!id) continue; + const title = stringValue(record.thread_name) + ?? stringValue(record.title) + ?? stringValue(record.name) + ?? stringValue(record.preview); + const updatedAtMs = timestampMs(record.updated_at_ms) + ?? timestampMs(record.updatedAtMs) + ?? timestampMs(record.updated_at) + ?? timestampMs(record.updatedAt) + ?? timestampMs(record.last_updated_at); + const cwd = stringValue(record.cwd) ?? stringValue(record.workspace) ?? stringValue(record.workspacePath); + result.set(id, { + ...(title ? { title } : {}), + ...(updatedAtMs !== undefined ? { updatedAtMs } : {}), + ...(cwd ? { cwd } : {}), + }); + } catch { + // A partially-written last line must not make the rest of the index + // unavailable. + } + } + return result; +} + +async function readBoundedText(fileName: string, maxBytes: number): Promise { + const handle = await fs.open(fileName, "r"); + const parts: Buffer[] = []; + let offset = 0; + try { + while (offset < maxBytes) { + const size = Math.min(64 * 1024, maxBytes - offset); + const chunk = Buffer.alloc(size); + const read = await handle.read(chunk, 0, size, offset); + if (!read.bytesRead) break; + parts.push(chunk.subarray(0, read.bytesRead)); + offset += read.bytesRead; + if (read.bytesRead < size) break; + } + return Buffer.concat(parts).toString("utf8"); + } finally { + await handle.close(); + } +} + +function sanitizeSessionTitle(value: string | undefined): string | undefined { + if (!value) return undefined; + const title = redactText(value).replace(/\s+/g, " ").trim().slice(0, 240); + return title || undefined; +} + +function samePath(left: string, right: string): boolean { + try { return path.resolve(left) === path.resolve(right); } catch { return left === right; } +} + + +function resolveCodexHome(options: CodexIpcAgentAdapterOptions): string { + const env = options.env ?? process.env; + const configured = options.codexHome?.trim() || env.CODEX_HOME?.trim() || path.join(options.homeDir ?? os.homedir(), ".codex"); + if (configured === "~") return options.homeDir ?? os.homedir(); + if (configured.startsWith("~/")) return path.join(options.homeDir ?? os.homedir(), configured.slice(2)); + return configured; +} + +function isMissingSessionOwnerError(error: unknown): boolean { + const message = error instanceof Error ? error.message : String(error); + return /找不到会话\s+.+\s+的 VS Code Codex owner/.test(message) + || /no existing codex conversation owner is attached/i.test(message); +} + +function extractInput(params: JsonObject): string | JsonValue[] { + if (typeof params.text === "string") return params.text; + if (typeof params.message === "string") return params.message; + if (typeof params.prompt === "string") return params.prompt; + if (typeof params.input === "string") return params.input; + if (Array.isArray(params.input) && params.input.length) return params.input; + throw new Error("turn request requires text or a non-empty input array"); +} + +function pickTurnRequest(params: JsonObject): JsonObject { + const request: JsonObject = {}; + if (Array.isArray(params.attachments)) request.attachments = asJsonValue(params.attachments); + // The owner inherits all existing thread settings. Only forward fields that + // are part of the app-server turn/start request, never relay UI-only keys. + for (const key of ["model", "serviceTier", "effort", "summary", "personality", "collaborationMode", "approvalPolicy", "approvalsReviewer", "permissions", "sandboxPolicy", "runtimeWorkspaceRoots", "cwd", "outputSchema", "multiAgentMode"]) { + if (params[key] !== undefined) request[key] = asJsonValue(params[key]); + } + return request; +} + +/** + * The model picker changes durable next-turn settings, not one turn request. + * Accept the relay-friendly flat form and the official nested form while + * forwarding only the fields required by the picker. + */ +function pickThreadSettingsUpdate(params: JsonObject): JsonObject { + const source = isRecord(params.threadSettings) ? params.threadSettings : params; + const settings: JsonObject = {}; + + if (Object.prototype.hasOwnProperty.call(source, "model")) { + if (typeof source.model !== "string" || !source.model.trim()) { + throw new Error("thread settings model must be a non-empty string"); + } + settings.model = source.model.trim(); + } + + if (Object.prototype.hasOwnProperty.call(source, "effort")) { + if (source.effort !== null && (typeof source.effort !== "string" || !source.effort.trim())) { + throw new Error("thread settings effort must be a non-empty string or null"); + } + settings.effort = typeof source.effort === "string" ? source.effort.trim() : null; + } + + if (Object.prototype.hasOwnProperty.call(source, "multiAgentMode")) { + if (source.multiAgentMode !== null && typeof source.multiAgentMode !== "string") { + throw new Error("thread settings multiAgentMode must be a string or null"); + } + settings.multiAgentMode = asJsonValue(source.multiAgentMode); + } + + // The official permissions control updates the same durable thread + // settings envelope as the model picker. Keep the projection explicit so a + // browser cannot smuggle arbitrary UI state into the owner request, while + // retaining object-valued policies used by newer app-server builds. + for (const key of ["sandboxPolicy", "approvalPolicy"] as const) { + if (!Object.prototype.hasOwnProperty.call(source, key)) continue; + const value = source[key]; + if (value !== null && typeof value !== "string" && !isRecord(value)) { + throw new Error(`thread settings ${key} must be a string, object, or null`); + } + settings[key] = asJsonValue(value); + } + if (Object.prototype.hasOwnProperty.call(source, "approvalsReviewer")) { + const value = source.approvalsReviewer; + if (value !== null && typeof value !== "string") { + throw new Error("thread settings approvalsReviewer must be a string or null"); + } + settings.approvalsReviewer = asJsonValue(value); + } + if (Object.prototype.hasOwnProperty.call(source, "runtimeWorkspaceRoots")) { + const value = source.runtimeWorkspaceRoots; + if (value !== null && (!Array.isArray(value) || !value.every((entry) => typeof entry === "string"))) { + throw new Error("thread settings runtimeWorkspaceRoots must be an array of strings or null"); + } + settings.runtimeWorkspaceRoots = asJsonValue(value); + } + if (Object.prototype.hasOwnProperty.call(source, "permissions")) { + const value = source.permissions; + if (value !== null && typeof value !== "string" && !isRecord(value)) { + throw new Error("thread settings permissions must be a string, object, or null"); + } + settings.permissions = asJsonValue(value); + } + + if (Object.keys(settings).length === 0) { + throw new Error("thread settings update requires model, effort, multiAgentMode, sandboxPolicy, approvalPolicy, permissions, or approvalsReviewer"); + } + return settings; +} + +function pickTurnContext(params: JsonObject): JsonObject { + const context = isRecord(params.context) ? asJsonObject(params.context) : {}; + if (context.inheritThreadSettings === undefined) context.inheritThreadSettings = true; + if (Array.isArray(params.commentAttachments)) context.commentAttachments = asJsonValue(params.commentAttachments); + if (Array.isArray(params.mcpAppModelContextAttachments)) context.mcpAppModelContextAttachments = asJsonValue(params.mcpAppModelContextAttachments); + return context; +} + +function unwrapFollowerResult(value: JsonValue | undefined): JsonValue { + if (isRecord(value) && Object.prototype.hasOwnProperty.call(value, "result")) return asJsonValue(value.result); + return asJsonValue(value); +} + +function extractRevision(value: unknown): number | undefined { + const direct = numberValue(value); + if (direct !== undefined) return direct; + if (!isRecord(value)) return undefined; + for (const key of ["revision", "streamRevision", "stateRevision"]) { + const revision = numberValue(value[key]); + if (revision !== undefined) return revision; + } + return Object.prototype.hasOwnProperty.call(value, "result") ? extractRevision(value.result) : undefined; +} + +function extractTurnId(value: unknown): string | undefined { + if (!isRecord(value)) return undefined; + if (typeof value.turnId === "string") return value.turnId; + if (isRecord(value.turn) && typeof value.turn.id === "string") return value.turn.id; + if (isRecord(value.result)) return extractTurnId(value.result); + return undefined; +} + +function deriveTurn(state: JsonObject): TurnInfo { + const candidates: TurnInfo[] = []; + const turns = Array.isArray(state.turns) ? state.turns : []; + for (const value of turns) if (isRecord(value)) candidates.push(turnInfo(value)); + const history = isRecord(state.turnHistory) && isRecord(state.turnHistory.history) && isRecord(state.turnHistory.history.entitiesByKey) + ? Object.values(state.turnHistory.history.entitiesByKey) : []; + for (const value of history) if (isRecord(value) && (value.turnId !== undefined || value.status !== undefined || value.items !== undefined)) candidates.push(turnInfo(value)); + const runtime = isRecord(state.threadRuntimeStatus) ? state.threadRuntimeStatus : undefined; + const runtimeType = normalizeStatus(runtime?.type); + const active = candidates.filter((entry) => entry.active).sort((a, b) => (b.startedAt ?? 0) - (a.startedAt ?? 0)); + if (active[0]) return active[0]; + const selected = candidates.sort((a, b) => (b.startedAt ?? 0) - (a.startedAt ?? 0))[0]; + if (runtime && runtimeType && runtimeType !== "idle" && runtimeType !== "ready" && typeof runtime.turnId === "string") { + return { id: runtime.turnId, status: runtimeType, active: true }; + } + return selected ?? { status: runtimeType || "idle", active: false }; +} + +function turnInfo(value: Record): TurnInfo { + const id = typeof value.id === "string" ? value.id : typeof value.turnId === "string" ? value.turnId : undefined; + const status = normalizeStatus(typeof value.status === "string" ? value.status : isRecord(value.status) && typeof value.status.type === "string" ? value.status.type : "unknown"); + const startedAt = timestampMs(value.turnStartedAtMs) + ?? timestampMs(value.startedAtMs) + ?? timestampMs(value.startedAt) + ?? timestampMs(value.createdAtMs) + ?? timestampMs(value.createdAt); + const durationMs = numberValue(value.durationMs) ?? numberValue(value.duration); + const completedAtMs = timestampMs(value.completedAtMs) ?? timestampMs(value.completedAt); + const commandStarts = isRecord(value.commandExecutionStartedAtMsById) + ? value.commandExecutionStartedAtMsById + : {}; + const items = Array.isArray(value.items) ? value.items.filter(isRecord) : []; + const firstTurnWorkItemStartedAtMs = timestampMs(value.firstTurnWorkItemStartedAtMs) + ?? timestampMs(value.firstWorkItemStartedAtMs) + ?? timestampMs(value.firstTurnWorkItemStartedAt) + ?? inferFirstWorkItemStartedAtMs(items, commandStarts); + const finalAssistantStartedAtMs = timestampMs(value.finalAssistantStartedAtMs) + ?? timestampMs(value.finalAssistantStartedAt) + ?? inferFinalAssistantStartedAtMs(items, commandStarts); + const active = !TERMINAL_TURN_STATES.has(status) && status !== "idle" && status !== "unknown"; + const explicitWorkedDurationMs = numberValue(value.workedDurationMs) + ?? numberValue(value.workDurationMs) + ?? numberValue(value.workDuration); + const workedCompletedAtMs = finalAssistantStartedAtMs + ?? (!active && completedAtMs !== undefined ? completedAtMs : undefined); + const workedDurationMs = explicitWorkedDurationMs + ?? (firstTurnWorkItemStartedAtMs !== undefined && workedCompletedAtMs !== undefined + ? Math.max(0, workedCompletedAtMs - firstTurnWorkItemStartedAtMs) + : undefined); + return { + id, + status, + active, + ...(startedAt !== undefined ? { startedAt } : {}), + ...(durationMs !== undefined ? { durationMs } : {}), + ...(workedDurationMs !== undefined ? { workedDurationMs } : {}), + ...(completedAtMs !== undefined ? { completedAtMs } : {}), + ...(firstTurnWorkItemStartedAtMs !== undefined ? { firstTurnWorkItemStartedAtMs } : {}), + ...(finalAssistantStartedAtMs !== undefined ? { finalAssistantStartedAtMs } : {}), + ...(value.error !== undefined ? { error: asJsonValue(value.error) } : {}), + raw: value, + }; +} + +/** + * Infer the timestamps used by the official "worked for" row when an older + * conversation snapshot does not carry the denormalized turn fields. Persisted + * rollout records use snake_case command metadata, so both spellings are + * intentionally accepted here. + */ +function inferFirstWorkItemStartedAtMs( + items: Record[], + commandStarts: Record, +): number | undefined { + for (const item of items) { + if (isNonWorkItem(item)) continue; + const id = stringValue(item.id); + const timestamp = itemTimestamp(item) + ?? (id ? timestampMs(commandStarts[id]) : undefined); + if (timestamp !== undefined) return timestamp; + } + return undefined; +} + +function inferFinalAssistantStartedAtMs( + items: Record[], + commandStarts: Record, +): number | undefined { + let fallback: number | undefined; + for (const item of items) { + if (!isAssistantItem(item)) continue; + const id = stringValue(item.id); + const timestamp = itemTimestamp(item) + ?? (id ? timestampMs(commandStarts[id]) : undefined); + if (timestamp === undefined) continue; + fallback = timestamp; + const phase = normalizeStatus(item.phase); + if (phase === "final_answer" || phase === "finalanswer") return timestamp; + } + return fallback; +} + +function itemTimestamp(item: Record): number | undefined { + return timestampMs(item.startedAtMs) + ?? timestampMs(item.started_at_ms) + ?? timestampMs(item.startedAt) + ?? timestampMs(item.started_at) + ?? timestampMs(item.createdAtMs) + ?? timestampMs(item.created_at_ms) + ?? timestampMs(item.createdAt) + ?? timestampMs(item.created_at); +} + +function normalizedItemType(item: Record): string { + return String(item.type ?? item.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); +} + +function isAssistantItem(item: Record): boolean { + const type = normalizedItemType(item); + return type === "agentmessage" || type === "assistantmessage"; +} + +function isNonWorkItem(item: Record): boolean { + const type = normalizedItemType(item); + return type === "usermessage" + || type === "steeringusermessage" + || type === "realtimetranscript" + || type === "worktreeinit" + || type === "sleep"; +} + +/** + * Normalize the official turn/runtime records into a stable status shape for + * relay clients. This intentionally tolerates both raw webview state + * (`turnStartedAtMs`, `inProgress`) and serialized history summaries + * (`startedAt`, `completed`) used by different extension versions. + */ +function deriveStatusSnapshot(state: JsonObject, turn: TurnInfo, requests: RequestEntry[]): AgentStatusSnapshot { + const runtime = isRecord(state.threadRuntimeStatus) ? state.threadRuntimeStatus : undefined; + const activeFlags = Array.isArray(runtime?.activeFlags) + ? runtime.activeFlags.filter((flag): flag is string => typeof flag === "string") + : []; + const activity = classifyActivity(turn, activeFlags, requests, runtime); + const startedAtMs = turn.startedAt ?? null; + const durationMs = turn.durationMs + ?? (!turn.active && turn.completedAtMs !== undefined && turn.completedAtMs !== null && turn.startedAt !== undefined + ? Math.max(0, turn.completedAtMs - turn.startedAt) + : null); + // The official worked-for indicator starts at the first actual work item, + // not at the user-message/turn start, and ends when the final assistant + // message begins. Keep this separate from the broader turn duration. + const workStartedAtMs = turn.firstTurnWorkItemStartedAtMs ?? turn.startedAt ?? undefined; + const workedCompletedAtMs = turn.finalAssistantStartedAtMs + ?? (!turn.active ? turn.completedAtMs : undefined) + ?? undefined; + const workedDurationMs = turn.workedDurationMs + ?? (workStartedAtMs !== undefined && workedCompletedAtMs !== undefined + ? Math.max(0, workedCompletedAtMs - workStartedAtMs) + : turn.active && workStartedAtMs !== undefined + ? Math.max(0, Date.now() - workStartedAtMs) + : null); + const elapsedMs = turn.active && workStartedAtMs !== undefined + ? Math.max(0, Date.now() - workStartedAtMs) + : durationMs; + const status: AgentStatusSnapshot = { + activity, + turnStatus: turn.status, + activeFlags, + startedAtMs, + durationMs, + workedDurationMs, + elapsedMs, + firstTurnWorkItemStartedAtMs: turn.firstTurnWorkItemStartedAtMs ?? null, + finalAssistantStartedAtMs: turn.finalAssistantStartedAtMs ?? null, + ...(turn.error !== undefined ? { error: turn.error } : {}), + }; + return status; +} + +function classifyActivity( + turn: TurnInfo, + activeFlags: string[], + requests: RequestEntry[], + runtime?: Record, +): string { + const flags = new Set(activeFlags.map((flag) => normalizeStatus(flag))); + if (flags.has("waiting_on_approval") || flags.has("waiting_for_approval") || flags.has("waitingonapproval") || requests.some((entry) => Boolean(entry.approval))) { + return "waiting_approval"; + } + if (flags.has("waiting_on_user_input") || flags.has("waiting_for_user_input") || flags.has("waitingonuserinput") || requests.some((entry) => INPUT_METHODS.has(entry.method))) { + return "waiting_input"; + } + if (!turn.active) return terminalActivity(turn.status); + + const item = latestActiveWorkItem(turn.raw); + if (item) { + const type = normalizeStatus(item.type ?? item.kind ?? ""); + if (type === "read" || type.includes("fileread") || type.includes("readfile") || readPathsFromActions(commandActionsValue(item)).length > 0) return "reading"; + if (type.includes("filechange") || type.includes("file_change") || type.includes("patch") || type.includes("edit")) return "editing"; + if (type.includes("reasoning") || type.includes("think")) return "thinking"; + if (type.includes("command") || type.includes("exec") || type.includes("process") || type.includes("tool")) return "running"; + const phase = normalizeStatus(item.phase); + if (phase.includes("reason") || phase.includes("think")) return "thinking"; + } + + const runtimeType = normalizeStatus(runtime?.type); + if (runtimeType.includes("reason") || runtimeType.includes("think")) return "thinking"; + if (runtimeType.includes("edit") || runtimeType.includes("patch")) return "editing"; + return "running"; +} + +function latestActiveWorkItem(turn?: Record): Record | undefined { + if (!turn || !Array.isArray(turn.items)) return undefined; + for (let index = turn.items.length - 1; index >= 0; index -= 1) { + const item = turn.items[index]; + if (!isRecord(item)) continue; + const status = normalizeStatus(isRecord(item.status) ? item.status.type : item.status); + if (status && status !== "unknown" && TERMINAL_TURN_STATES.has(status)) continue; + return item; + } + return undefined; +} + +function terminalActivity(status: string): string { + const normalized = normalizeStatus(status); + if (normalized === "failed" || normalized === "error") return "failed"; + if (normalized === "cancelled" || normalized === "canceled" || normalized === "interrupted") return "interrupted"; + if (normalized === "completed" || normalized === "complete" || normalized === "done") return "completed"; + return normalized === "idle" || normalized === "ready" ? "idle" : "idle"; +} + +function statusPayload(status: AgentStatusSnapshot): JsonObject { + return { + // `status` is the legacy scalar alias; `turnStatus` retains the explicit + // name so clients can distinguish it from the coarse `activity` value. + status: status.turnStatus, + turnStatus: status.turnStatus, + activity: status.activity, + activeFlags: asJsonValue(status.activeFlags), + startedAtMs: status.startedAtMs ?? null, + durationMs: status.durationMs ?? null, + workedDurationMs: status.workedDurationMs ?? null, + elapsedMs: status.elapsedMs ?? null, + firstTurnWorkItemStartedAtMs: status.firstTurnWorkItemStartedAtMs ?? null, + finalAssistantStartedAtMs: status.finalAssistantStartedAtMs ?? null, + ...(status.error !== undefined ? { error: status.error } : {}), + }; +} + +function sameStatus(a: AgentStatusSnapshot, b: AgentStatusSnapshot): boolean { + return a.activity === b.activity + && a.turnStatus === b.turnStatus + && JSON.stringify(a.activeFlags) === JSON.stringify(b.activeFlags) + && a.startedAtMs === b.startedAtMs + && a.durationMs === b.durationMs + && a.workedDurationMs === b.workedDurationMs + && a.firstTurnWorkItemStartedAtMs === b.firstTurnWorkItemStartedAtMs + && a.finalAssistantStartedAtMs === b.finalAssistantStartedAtMs + && JSON.stringify(a.error) === JSON.stringify(b.error); +} + +function timestampMs(value: unknown): number | undefined { + if (typeof value === "string") { + const trimmed = value.trim(); + if (/^[+-]?(?:\d+(?:\.\d*)?|\.\d+)$/.test(trimmed)) { + const number = Number(trimmed); + if (Number.isFinite(number)) { + return number > 0 && number < 1_000_000_000_000 ? number * 1000 : number; + } + } + const parsed = Date.parse(value); + return Number.isFinite(parsed) ? parsed : undefined; + } + const number = numberValue(value); + if (number === undefined) return undefined; + // Serialized history in older extension builds uses epoch seconds while + // the live turn fields end in `AtMs`. Normalize both to milliseconds. + return number > 0 && number < 1_000_000_000_000 ? number * 1000 : number; +} + +function extractRequests(state: JsonObject, approvalTimeoutMs: number): RequestEntry[] { + const values: Array<{ id?: JsonRpcId; value: Record }> = []; + const seenRecords = new Set>(); + const collect = (raw: unknown, hintedId?: JsonRpcId, depth = 0): void => { + if (depth > 6 || raw === null || raw === undefined) return; + if (Array.isArray(raw)) { + for (const value of raw) collect(value, undefined, depth + 1); + return; + } + if (!isRecord(raw)) return; + if (seenRecords.has(raw)) return; + seenRecords.add(raw); + const nested = isRecord(raw.request) ? raw.request : raw; + const method = requestMethodOf(nested) ?? requestMethodOf(raw); + const requestId = requestIdOf(nested) ?? requestIdOf(raw) ?? hintedId; + if (method && isJsonRpcId(requestId)) values.push({ id: requestId, value: raw }); + // Official conversation snapshots can retain pending records inside a + // turn item (permission-request/userInput/mcp-server-elicitation), while + // older builds expose the same records under `requests`. Traverse only + // protocol containers so arbitrary message content is never interpreted as + // an approval request. + for (const key of ["requests", "pendingRequests", "pendingApprovals", "turns", "turnHistory", "history", "items", "request"]) { + const child = raw[key]; + if (child === undefined) continue; + if (isRecord(child) && !Array.isArray(child)) { + for (const [childKey, value] of Object.entries(child)) { + collect(value, isJsonRpcId(childKey) ? childKey : undefined, depth + 1); + } + } else collect(child, undefined, depth + 1); + } + }; + collect(state); + const result: RequestEntry[] = []; + const seenRequests = new Set(); + for (const item of values) { + const request = isRecord(item.value.request) ? item.value.request : item.value; + const requestId = item.id ?? requestIdOf(request); + const method = requestMethodOf(request) ?? requestMethodOf(item.value); + if (!isJsonRpcId(requestId) || !method) continue; + const dedupeKey = `${jsonRpcIdKey(requestId)}\u001f${method}`; + if (seenRequests.has(dedupeKey)) continue; + seenRequests.add(dedupeKey); + const params = asJsonObject(request.params ?? item.value.params ?? (request === item.value ? item.value : {})); + // `params.startedAtMs` is the app-server's authoritative timestamp. The + // outer fields are compatibility fallbacks for older normalized snapshots + // and can represent when a UI record was inserted rather than when the + // approval actually started. + const createdAt = timestampMs(params.startedAtMs) + ?? timestampMs(params.started_at_ms) + ?? timestampMs(params.startedAt) + ?? timestampMs(params.started_at) + ?? timestampMs(request.startedAtMs) + ?? timestampMs(request.started_at_ms) + ?? timestampMs(request.startedAt) + ?? timestampMs(request.started_at) + ?? timestampMs(item.value.startedAtMs) + ?? timestampMs(item.value.started_at_ms) + ?? timestampMs(item.value.startedAt) + ?? timestampMs(item.value.started_at) + ?? timestampMs(item.value.createdAtMs) + ?? timestampMs(item.value.created_at_ms) + ?? timestampMs(item.value.createdAt) + ?? timestampMs(item.value.created_at) + ?? Date.now(); + const expiresAt = timestampMs(item.value.expiresAtMs) + ?? timestampMs(item.value.expires_at_ms) + ?? timestampMs(item.value.expiresAt) + ?? timestampMs(item.value.expires_at) + ?? timestampMs(request.expiresAtMs) + ?? timestampMs(request.expires_at_ms) + ?? timestampMs(request.expiresAt) + ?? timestampMs(request.expires_at) + ?? timestampMs(params.expiresAtMs) + ?? timestampMs(params.expires_at_ms) + ?? timestampMs(params.expiresAt) + ?? timestampMs(params.expires_at) + ?? (approvalTimeoutMs > 0 ? createdAt + approvalTimeoutMs : undefined); + const entry: RequestEntry = { + requestId, + method, + params, + threadId: stringValue(params.threadId) ?? stringValue(params.conversationId), + turnId: stringValue(params.turnId), + createdAt, + ...(expiresAt ? { expiresAt } : {}), + }; + if (APPROVAL_METHODS.has(method)) entry.approval = toPendingApproval(entry); + result.push(entry); + } + return result; +} + +function requestMethodOf(value: Record): string | undefined { + if (typeof value.method === "string" && value.method) return value.method; + const type = String(value.type ?? value.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); + if (type.includes("permissionrequest")) return "item/permissions/requestApproval"; + if (type.includes("commandexecutionrequest") || type === "execapproval") return "item/commandExecution/requestApproval"; + if (type.includes("filechangerequest") || type === "patchapproval") return "item/fileChange/requestApproval"; + if (type === "exec" && value.approvalRequestId !== undefined && (!isRecord(value.output) || value.output.exitCode === undefined)) return "execCommandApproval"; + if (type === "patch" && value.approvalRequestId !== undefined && value.success === undefined) return "applyPatchApproval"; + if (type.includes("userinput") && value.completed !== true) return "item/tool/requestUserInput"; + if (type.includes("mcpserverelicitation") && value.completed !== true) return "mcpServer/elicitation/request"; + return undefined; +} + +function requestIdOf(value: Record): JsonRpcId | undefined { + const id = value.requestId ?? value.id; + return isJsonRpcId(id) ? id : undefined; +} + +function toPendingApproval(entry: RequestEntry): PendingApproval { + const command = typeof entry.params.command === "string" ? entry.params.command : extractCommand(entry.params); + const action = entry.method.includes("fileChange") || entry.method === "applyPatchApproval" ? "file.change" : entry.method.includes("permissions") ? "permissions.grant" : "command.execution"; + const risk: PendingApproval["risk"] = entry.method.includes("permissions") + ? "high" + : entry.method.includes("command") || entry.method === "execCommandApproval" + ? (!command || /(?:rm\s+-rf|sudo|curl|wget|ssh|password|token|secret)/i.test(command) ? "high" : "medium") + : "medium"; + const summary = stringValue(entry.params.reason) ?? command ?? `${action} requested by Codex`; + return { + requestId: entry.requestId, + method: entry.method, + threadId: entry.threadId, + turnId: entry.turnId, + itemId: stringValue(entry.params.itemId) ?? stringValue(entry.params.callId), + action, + risk, + summary: redactText(summary), + commandHash: hashJson(entry.params), + createdAt: entry.createdAt, + ...(entry.expiresAt ? { expiresAt: entry.expiresAt } : {}), + payload: redactJson(entry.params) as JsonObject, + }; +} + +function extractCommand(params: JsonObject): string | undefined { + if (Array.isArray(params.command)) return params.command.filter((value): value is string => typeof value === "string").join(" "); + const actions = params.commandActions ?? params.command_actions ?? params.parsedCmd ?? params.parsed_cmd; + const actionList = commandActionList(actions); + if (actionList.length) { + const commands = actionList.map(commandActionText).filter((value): value is string => Boolean(value)); + return commands.length ? commands.join(" && ") : undefined; + } + return undefined; +} + +interface RenderedConversationMessage { + id?: string; + turnId?: string; + itemId?: string; + role: "user" | "assistant" | "reasoning" | "tool" | "error"; + kind: "user" | "assistant" | "reasoning" | "plan" | "tool" | "edit" | "error"; + text: string; + label?: string; + itemType?: string; + status?: string; + turnStatus?: string; + startedAtMs?: number; + completedAtMs?: number; + durationMs?: number; + /** Duration of the official worked-for activity group for this turn. */ + workedDurationMs?: number; + command?: string; + /** Parsed command actions emitted by the official command renderer. */ + commandActions?: JsonValue[]; + cwd?: string | null; + shellName?: string | null; + exitCode?: number; + phase?: string; + breaksPreviousAdjacency?: boolean; + /** Official collabAgentToolCall projection. */ + action?: string; + senderThreadId?: string; + receiverThreadIds?: string[]; + /** Official webview compatibility alias. */ + receiverThreads?: string[]; + prompt?: string | null; + model?: string | null; + reasoningEffort?: string | null; + agentsStates?: JsonObject; + /** Official subAgentActivity projection. */ + agentThreadId?: string; + agentPath?: string; + displayName?: string | null; + displayStatus?: string; + activityKind?: string; + /** Semantic name used by the official webview converter. */ + uiType?: string; + /** Friendly paths extracted from a parsed `read` command action. */ + readPaths?: string[]; + /** Raw tool/file output kept separate from the compact activity summary. */ + output?: string; +} + +interface ItemDisplayProjection { + outputText: string; + projectionId?: string; + startedAtMs?: number; + completedAtMs?: number; + durationMs?: number; + workedDurationMs?: number; + role: RenderedConversationMessage["role"]; + kind: RenderedConversationMessage["kind"]; + text: string; + label?: string; + itemType?: string; + status?: string; + command?: string; + commandActions?: JsonValue[]; + cwd?: string | null; + shellName?: string | null; + action?: string; + senderThreadId?: string; + receiverThreadIds?: string[]; + receiverThreads?: string[]; + prompt?: string | null; + model?: string | null; + reasoningEffort?: string | null; + agentsStates?: JsonObject; + agentThreadId?: string; + agentPath?: string; + displayName?: string | null; + displayStatus?: string; + activityKind?: string; + uiType?: string; + readPaths?: string[]; + output?: string; +} + +function renderedMessageShape(messages: RenderedConversationMessage[]): string { + return messages.map((message, index) => [ + message.id ?? `index:${index}`, + message.turnId ?? "", + message.itemId ?? "", + message.role, + message.kind, + message.itemType ?? "", + message.text, + message.label ?? "", + message.status ?? "", + message.turnStatus ?? "", + message.startedAtMs ?? "", + message.completedAtMs ?? "", + message.durationMs ?? "", + message.workedDurationMs ?? "", + message.command ?? "", + JSON.stringify(message.commandActions ?? []), + message.cwd ?? "", + message.shellName ?? "", + message.exitCode ?? "", + message.phase ?? "", + message.action ?? "", + message.senderThreadId ?? "", + JSON.stringify(message.receiverThreadIds ?? []), + message.prompt ?? "", + message.model ?? "", + message.reasoningEffort ?? "", + JSON.stringify(message.agentsStates ?? {}), + message.agentThreadId ?? "", + message.agentPath ?? "", + message.displayName ?? "", + message.displayStatus ?? "", + message.activityKind ?? "", + message.uiType ?? "", + JSON.stringify(message.readPaths ?? []), + message.output ?? "", + message.breaksPreviousAdjacency ? "break" : "", + ].join("\u001f")).join("\u001e"); +} + +/** + * Encode a suffix replacement instead of repeating the complete structured + * history on every streaming text patch. Initial/output snapshots still carry + * the full projection, so reconnect and late-join hydration stay lossless. + */ +function renderedMessagesPatch( + previous: RenderedConversationMessage[], + next: RenderedConversationMessage[], +): JsonObject | undefined { + const sharedLength = Math.min(previous.length, next.length); + let start = 0; + while (start < sharedLength + && stableStringify(asJsonValue(previous[start])) === stableStringify(asJsonValue(next[start]))) start += 1; + if (start === previous.length && start === next.length) return undefined; + return { + start, + deleteCount: previous.length - start, + messages: asJsonValue(next.slice(start)), + }; +} + +function renderedSubagentShape(subagents: SubagentSnapshot[]): string { + return JSON.stringify(subagents); +} + +function renderConversationOutput(state: JsonObject, maxChars: number): { text: string; totalLength: number; truncated: boolean; messages: RenderedConversationMessage[]; subagents: SubagentSnapshot[] } { + const chunks: string[] = []; + const messages: RenderedConversationMessage[] = []; + const seen = new Map(); + const add = (text: string, id?: string, message?: RenderedConversationMessage): void => { + const safe = redactText(text); + if (!safe) return; + if (id && seen.has(id)) { + // The same turn is commonly present in both the canonical history and + // the active-page list. Keep the latest item text when a streaming item + // was updated, instead of dropping the active-page update entirely. + const position = seen.get(id) as number; + chunks[position] = safe; + if (message) messages[position] = { ...message, text: message.text ? redactText(message.text) : safe }; + return; + } + if (id) seen.set(id, chunks.length); + chunks.push(safe); + if (message) messages.push({ ...message, text: redactText(message.text || safe) }); + }; + const consumeTurn = (turn: Record): void => { + const turnKey = stringValue(turn.id) ?? stringValue(turn.turnId); + const turnStatus = statusValue(turn.status); + const turnStartedAtMs = timestampMs(turn.turnStartedAtMs) + ?? timestampMs(turn.startedAtMs) + ?? timestampMs(turn.startedAt) + ?? timestampMs(turn.createdAtMs) + ?? timestampMs(turn.createdAt); + const turnDurationMs = numberValue(turn.durationMs) ?? numberValue(turn.duration); + const firstTurnWorkItemStartedAtMs = timestampMs(turn.firstTurnWorkItemStartedAtMs) + ?? timestampMs(turn.firstWorkItemStartedAtMs) + ?? timestampMs(turn.firstTurnWorkItemStartedAt); + const finalAssistantStartedAtMs = timestampMs(turn.finalAssistantStartedAtMs) + ?? timestampMs(turn.finalAssistantStartedAt); + const commandStarts = isRecord(turn.commandExecutionStartedAtMsById) + ? turn.commandExecutionStartedAtMsById + : {}; + const items = Array.isArray(turn.items) ? turn.items.filter(isRecord) : []; + const inferredFirstWorkItemStartedAtMs = firstTurnWorkItemStartedAtMs + ?? inferFirstWorkItemStartedAtMs(items, commandStarts); + const inferredFinalAssistantStartedAtMs = finalAssistantStartedAtMs + ?? inferFinalAssistantStartedAtMs(items, commandStarts); + const workedCompletedAtMs = inferredFinalAssistantStartedAtMs + ?? (!TERMINAL_TURN_STATES.has(normalizeStatus(turnStatus)) + ? undefined + : timestampMs(turn.completedAtMs) ?? timestampMs(turn.completedAt)); + const workedDurationMs = numberValue(turn.workedDurationMs) + ?? numberValue(turn.workDurationMs) + ?? (inferredFirstWorkItemStartedAtMs !== undefined && workedCompletedAtMs !== undefined + ? Math.max(0, workedCompletedAtMs - inferredFirstWorkItemStartedAtMs) + : undefined); + // A turn may append bookkeeping/tool records after the final assistant + // item. Identify the final assistant from the rendered assistant records, + // with an explicit final-answer phase taking precedence over chronology. + const itemDisplays = items.map((item) => itemDisplayVariants(item)); + let finalAssistantIndex = -1; + let explicitFinalAssistantIndex = -1; + itemDisplays.forEach((displays, index) => { + if (!displays.some((display) => display.role === "assistant")) return; + finalAssistantIndex = index; + const phase = typeof items[index].phase === "string" + ? normalizeStatus(items[index].phase) + : ""; + if (phase === "final_answer" || phase === "finalanswer") explicitFinalAssistantIndex = index; + }); + if (explicitFinalAssistantIndex >= 0) finalAssistantIndex = explicitFinalAssistantIndex; + items.forEach((item, index) => { + const displays = itemDisplays[index]; + if (!displays.length) return; + const rawItemId = typeof item.id === "string" ? item.id : undefined; + const startedAtMs = timestampMs(item.startedAtMs) + ?? timestampMs(item.startedAt) + ?? (rawItemId ? timestampMs(commandStarts[rawItemId]) : undefined); + const durationMs = numberValue(item.durationMs) ?? numberValue(item.duration); + const completedAtMs = timestampMs(item.completedAtMs) + ?? timestampMs(item.finishedAtMs) + ?? timestampMs(item.completedAt) + ?? (startedAtMs !== undefined && durationMs !== undefined ? startedAtMs + durationMs : undefined); + const command = commandText(item); + const phase = typeof item.phase === "string" ? item.phase : undefined; + displays.forEach((display, displayIndex) => { + const projectionId = display.projectionId ?? rawItemId; + const itemId = projectionId + ? `id:${projectionId}` + : turnKey + ? `turn:${turnKey}:${index}:${displayIndex}` + : `raw:${JSON.stringify(item)}:${displayIndex}`; + const isFinalAssistant = display.role === "assistant" + && (phase === "final_answer" || phase === "final-answer" || index === finalAssistantIndex); + // User items do not carry their own timestamp in several official + // snapshots. Associate them with the turn start; likewise associate the + // final assistant item with the turn's final-answer start and duration. + const effectiveStartedAtMs = display.startedAtMs ?? startedAtMs + ?? (display.role === "user" ? turnStartedAtMs : undefined) + ?? (display.role === "reasoning" ? inferredFirstWorkItemStartedAtMs : undefined) + ?? (isFinalAssistant ? inferredFinalAssistantStartedAtMs : undefined); + const effectiveDurationMs = display.durationMs ?? durationMs + ?? (isFinalAssistant ? turnDurationMs : undefined); + const effectiveCompletedAtMs = display.completedAtMs ?? completedAtMs + ?? (effectiveStartedAtMs !== undefined && effectiveDurationMs !== undefined + ? effectiveStartedAtMs + effectiveDurationMs + : undefined); + const itemStatus = display.status ?? statusValue(item.status) + ?? (item.completed === true ? "completed" : item.completed === false ? "in_progress" : undefined); + const displayCommand = display.command ?? (displayIndex === 0 ? command : undefined); + add(display.outputText, itemId, { + id: itemId, + ...(turnKey ? { turnId: turnKey } : {}), + ...(projectionId ? { itemId: projectionId } : {}), + role: display.role, + kind: display.kind, + text: display.text, + ...(display.label ? { label: display.label } : {}), + ...(display.itemType ? { itemType: display.itemType } : typeof item.type === "string" ? { itemType: item.type } : typeof item.kind === "string" ? { itemType: item.kind } : {}), + ...(itemStatus ? { status: itemStatus } : {}), + ...(turnStatus ? { turnStatus } : {}), + ...(effectiveStartedAtMs !== undefined ? { startedAtMs: effectiveStartedAtMs } : {}), + ...(effectiveCompletedAtMs !== undefined ? { completedAtMs: effectiveCompletedAtMs } : {}), + ...(effectiveDurationMs !== undefined ? { durationMs: effectiveDurationMs } : {}), + ...(workedDurationMs !== undefined ? { workedDurationMs } : {}), + ...(displayCommand ? { command: displayCommand } : {}), + ...(display.commandActions?.length ? { commandActions: display.commandActions } : {}), + ...(display.cwd !== undefined ? { cwd: display.cwd } : {}), + ...(display.shellName !== undefined ? { shellName: display.shellName } : {}), + ...(numberValue(item.exitCode) !== undefined && displayIndex === 0 ? { exitCode: numberValue(item.exitCode) as number } : {}), + ...(phase && displayIndex === 0 ? { phase } : {}), + ...(display.action ? { action: display.action } : {}), + ...(display.senderThreadId ? { senderThreadId: display.senderThreadId } : {}), + ...(display.receiverThreadIds ? { receiverThreadIds: display.receiverThreadIds, receiverThreads: display.receiverThreads ?? display.receiverThreadIds } : {}), + ...(display.prompt !== undefined ? { prompt: display.prompt } : {}), + ...(display.model !== undefined ? { model: display.model } : {}), + ...(display.reasoningEffort !== undefined ? { reasoningEffort: display.reasoningEffort } : {}), + ...(display.agentsStates ? { agentsStates: display.agentsStates } : {}), + ...(display.agentThreadId ? { agentThreadId: display.agentThreadId } : {}), + ...(display.agentPath ? { agentPath: display.agentPath } : {}), + ...(display.displayName !== undefined ? { displayName: display.displayName } : {}), + ...(display.displayStatus ? { displayStatus: display.displayStatus } : {}), + ...(display.activityKind ? { activityKind: display.activityKind } : {}), + ...(display.uiType ? { uiType: display.uiType } : {}), + ...(display.readPaths?.length ? { readPaths: display.readPaths } : {}), + ...(display.output ? { output: redactText(display.output) } : {}), + ...(item.breaksPreviousAdjacency === true ? { breaksPreviousAdjacency: true } : {}), + }); + }); + }); + }; + // Canonical history islands carry the stable chronological order. The + // lightweight `turns` list is usually just the active page, so append only + // entities that are not already represented there. + for (const turn of orderedHistoryTurns(state)) consumeTurn(turn); + if (Array.isArray(state.turns)) for (const turn of state.turns) if (isRecord(turn)) consumeTurn(turn); + const rendered = chunks.join("\n\n"); + return { + text: rendered.length > maxChars ? rendered.slice(-maxChars) : rendered, + totalLength: rendered.length, + truncated: rendered.length > maxChars, + messages, + subagents: collectSubagents(state), + }; +} + +/** + * Return only newly appended text. Once the bounded output window starts + * sliding, compare the old suffix with the new prefix so a one-character + * stream update does not retransmit the entire 32 KB snapshot. + */ +function appendOnlyOutputDelta( + previous: string, + previousLength: number, + next: string, + nextLength: number, + previousWasTruncated: boolean, +): string | undefined { + // A replacement, deletion, or history prepend cannot be represented by an + // append-only chunk. Fall back to a bounded snapshot in those cases. + if (nextLength < previousLength) return undefined; + if (!previousWasTruncated) return next.startsWith(previous) ? next.slice(previous.length) : undefined; + if (!previous || !next) return undefined; + const dropped = Math.max(0, nextLength - next.length) - Math.max(0, previousLength - previous.length); + if (dropped < 0 || dropped > previous.length) return undefined; + const retained = previous.slice(dropped); + if (retained.length > next.length || next.slice(0, retained.length) !== retained) return undefined; + // If the append is larger than the retained tail, the bounded state no + // longer contains all newly appended text; a snapshot is the only lossless + // representation. + const deltaLength = nextLength - previousLength; + if (deltaLength !== next.length - retained.length) return undefined; + return next.slice(retained.length); +} + +function orderedHistoryTurns(state: JsonObject): Record[] { + const turnHistory = isRecord(state.turnHistory) ? state.turnHistory : undefined; + const history = turnHistory && isRecord(turnHistory.history) ? turnHistory.history : undefined; + const entities = history && isRecord(history.entitiesByKey) ? history.entitiesByKey : undefined; + if (!entities) return []; + const ordered: Record[] = []; + const seen = new Set(); + const add = (key: unknown): void => { + if (typeof key !== "string" || seen.has(key)) return; + const entity = entities[key]; + if (!isRecord(entity) || !looksLikeTurn(entity)) return; + seen.add(key); + ordered.push(entity); + }; + if (history && Array.isArray(history.islands)) { + for (const island of history.islands) if (isRecord(island) && Array.isArray(island.entries)) { + for (const entry of island.entries) { + if (isRecord(entry)) add(entry.key ?? entry.value); + } + } + } + // Include entities not listed by islands for forward compatibility with an + // extension that omits island metadata in a snapshot. + for (const [key, entity] of Object.entries(entities)) { + if (!seen.has(key) && isRecord(entity) && looksLikeTurn(entity)) { + seen.add(key); + ordered.push(entity); + } + } + return ordered; +} + +function looksLikeTurn(value: Record): boolean { + return Array.isArray(value.items) || value.turnId !== undefined || value.status !== undefined; +} + +function isTransientHistoryLoadError(error: unknown): boolean { + const code = isRecord(error) && typeof error.code === "string" ? error.code : ""; + if (["timeout", "connection-closed", "not-connected"].includes(code) || code.startsWith("no-client-found")) return true; + const message = error instanceof Error ? error.message : String(error ?? ""); + return /timed? out|socket (?:is )?closed|not connected|no client found/i.test(message); +} + +function hasIncompleteHistory(state: JsonObject): boolean { + const turnHistory = isRecord(state.turnHistory) ? state.turnHistory : undefined; + const history = turnHistory && isRecord(turnHistory.history) ? turnHistory.history : undefined; + if (turnHistory?.kind === "canonical" && !history) return true; + const entities = history && isRecord(history.entitiesByKey) ? Object.values(history.entitiesByKey) : []; + const turns = [ + ...entities, + ...(Array.isArray(state.turns) ? state.turns : []), + ]; + // Both canonical entities and the legacy turns list expose this per-turn + // marker. The official completeness predicate only treats an explicit + // `false` as incomplete; missing metadata is compatible with older builds. + if (turns.some((turn) => isRecord(turn) + && isRecord(turn.itemsPagination) + && turn.itemsPagination.hasLoadedOldest === false)) return true; + + // Canonical history is complete only after the owner has coalesced it into + // one island. Boundary status is deliberately not checked here: the + // official webview uses `isComplete` and island count, and some versions + // leave boundary objects in a non-exhausted transitional shape. + const canonical = Boolean(history && ( + turnHistory?.kind === "canonical" + || history.isComplete !== undefined + || Array.isArray(history.islands) + )); + if (canonical) { + return history?.isComplete !== true + || !Array.isArray(history.islands) + || history.islands.length !== 1; + } + + // Legacy snapshots carry a resume marker. Avoid requesting an unsupported + // history operation for old snapshots that expose no pagination metadata at + // all, while respecting explicit loading/unfinished states. + if (state.resumeState !== undefined && state.resumeState !== "resumed") return true; + const turnsPagination = isRecord(state.turnsPagination) ? state.turnsPagination : undefined; + return turnsPagination?.hasLoadedOldest === false; +} + +function itemDisplay(rawItem: Record): ItemDisplayProjection | undefined { + const normalizedItem = normalizeOfficialItem(rawItem); + const type = String(normalizedItem.type ?? normalizedItem.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); + + if (type === "collabagenttoolcall") { + const tool = stringValue(normalizedItem.tool) ?? "collabAgent"; + // `wait` is an internal synchronization action. The official webview + // consumes it for aggregation but intentionally omits it from the + // visible transcript. + if (tool === "wait") return undefined; + const status = statusValue(normalizedItem.status) ?? "inProgress"; + const receiverThreadIds = stringArray(normalizedItem.receiverThreadIds); + const agentsStates = collabAgentStates(normalizedItem.agentsStates); + const prompt = redactNullableString(normalizedItem.prompt); + const model = redactNullableString(normalizedItem.model); + const reasoningEffort = redactNullableString(normalizedItem.reasoningEffort); + const senderThreadId = stringValue(normalizedItem.senderThreadId); + const label = collabAgentToolLabel(tool); + const promptText = prompt?.trim() ? `: ${redactText(prompt.trim())}` : ""; + const outputText = `${label}${promptText}`; + return { + outputText, + projectionId: stringValue(normalizedItem.id), + role: "tool", + kind: "tool", + text: outputText, + label, + itemType: "collabAgentToolCall", + status, + action: tool, + ...(senderThreadId ? { senderThreadId } : {}), + receiverThreadIds, + receiverThreads: receiverThreadIds, + prompt, + model, + reasoningEffort, + agentsStates: redactJson(agentsStates) as JsonObject, + uiType: "multi-agent-action", + }; + } + + if (type === "subagentactivity") { + const activityKind = stringValue(normalizedItem.kind) ?? "started"; + const agentThreadId = stringValue(normalizedItem.agentThreadId); + if (!agentThreadId) return undefined; + const agentPath = redactNullableString(normalizedItem.agentPath); + const displayName = formatAgentPath(agentPath ?? undefined); + const displayStatus = subagentActivityDisplayStatus(activityKind); + const status = statusValue(normalizedItem.status) + ?? (activityKind === "interrupted" || activityKind === "completed" ? "completed" : "inProgress"); + const label = displayName ? `子代理 · ${displayName}` : "子代理"; + const activityText = subagentActivityText(displayName, activityKind); + return { + outputText: activityText, + projectionId: stringValue(normalizedItem.id), + role: "tool", + kind: "tool", + text: activityText, + label, + itemType: "subAgentActivity", + status, + agentThreadId, + ...(agentPath ? { agentPath } : {}), + displayName, + displayStatus, + activityKind, + uiType: "subagent-activity", + }; + } + + const item = normalizedItem; + // Pending permission/input/elicitation items are rendered by the request + // card, not as a second transcript activity. Once the owner marks one + // complete it may re-enter history and be displayed normally. + const requestItem = type.includes("permissionrequest") + || type.includes("userinput") + || type.includes("mcpserverelicitation"); + const itemStatus = normalizeStatus(statusValue(item.status)); + const requestPending = requestItem + && item.completed !== true + && !TERMINAL_TURN_STATES.has(itemStatus); + if (requestPending) return undefined; + if (["agentmessage", "assistantmessage", "usermessage"].includes(type)) { + const text = textFromValue(item.text) ?? textFromValue(item.content); + if (!text) return undefined; + if (type.startsWith("user")) return { outputText: `> ${text}`, role: "user", kind: "user", text }; + return { outputText: text, role: "assistant", kind: "assistant", text }; + } + if (type.includes("contextcompaction")) { + const text = textFromValue(item.summary) ?? textFromValue(item.content) ?? textFromValue(item.text) ?? "整理上下文"; + return { outputText: text, role: "reasoning", kind: "reasoning", text, label: "整理上下文" }; + } + if (type.includes("reasoning") || type.includes("approvalreview")) { + const text = textFromValue(item.summary) ?? textFromValue(item.content); + return text ? { outputText: text, role: "reasoning", kind: "reasoning", text, label: "思考" } : undefined; + } + if (type.includes("plan") || type.includes("todo")) { + const value = item.plan ?? item.steps ?? item.todos ?? item.content ?? item.text; + const text = planDisplayText(value); + return text ? { outputText: text, role: "reasoning", kind: "plan", text, label: "计划" } : undefined; + } + const parsedActions = commandActionsValue(item); + const readPaths = readPathsFromActions(parsedActions); + const directReadItem = type === "read" + || type.includes("fileread") + || type.includes("readfile") + || type.includes("exploration"); + if (directReadItem || readPaths.length > 0) { + const command = commandText(item); + const output = textFromValue(item.aggregatedOutput) + ?? textFromValue(item.output) + ?? textFromValue(item.stdout) + ?? textFromValue(item.content) + ?? textFromValue(item.text); + const pathSummary = readPaths.length ? readPaths.join(", ") : readPathFromItem(item); + const summary = pathSummary ? `已读取 ${pathSummary}` : "已读取文件"; + const text = summary; + const commandActions = projectCommandActions(parsedActions); + const cwd = redactNullableString(item.cwd); + const shellName = redactNullableString(item.shellName ?? item.shell); + return { + outputText: output && output.trim() ? `${summary}\n${output}` : text, + role: "tool", + kind: "tool", + text, + label: "已读取文件", + itemType: typeof item.type === "string" ? item.type : "fileRead", + ...(command ? { command } : {}), + ...(commandActions.length ? { commandActions } : {}), + ...(cwd !== undefined ? { cwd } : {}), + ...(shellName !== undefined ? { shellName } : {}), + ...(readPaths.length ? { readPaths } : {}), + ...(output && output.trim() ? { output } : {}), + activityKind: "read", + uiType: "file-read", + }; + } + if (type.includes("filechange") || type.includes("file_change") || type.includes("patch") || type.includes("edit")) { + const text = textFromValue(item.diff) + ?? textFromValue(item.patch) + ?? fileChangesText(item.changes) + ?? textFromValue(item.output) + ?? textFromValue(item.text); + return text ? { outputText: text, role: "tool", kind: "edit", text, label: "文件变更" } : undefined; + } + const hasCommandProjection = isCommandActionValue(item.commandActions) + || isCommandActionValue(item.command_actions) + || isCommandActionValue(item.parsedCmd) + || isCommandActionValue(item.parsed_cmd) + || item.command !== undefined + || item.commandLine !== undefined; + if (type.includes("command") || type.includes("exec") || type.includes("process") || hasCommandProjection) { + const command = commandText(item); + // Some official snapshots expose commandActions before output is flushed + // (and may leave aggregatedOutput as an empty string). Keep the command + // visible in that state, while avoiding a bare shell bootstrap such as + // `/bin/zsh` becoming the displayed command. + const output = textFromValue(item.aggregatedOutput) + ?? textFromValue(item.output) + ?? textFromValue(item.stdout) + ?? textFromValue(item.stderr); + const text = output || command; + const commandActions = projectCommandActions(parsedActions); + const cwd = redactNullableString(item.cwd); + const shellName = redactNullableString(item.shellName ?? item.shell); + return text ? { + outputText: text, + role: "tool", + kind: "tool", + text, + label: "命令输出", + ...(commandActions.length ? { commandActions } : {}), + ...(cwd !== undefined ? { cwd } : {}), + ...(shellName !== undefined ? { shellName } : {}), + } : undefined; + } + if (type.includes("websearch") || type.includes("mcp") || type.includes("dynamictool") + || type.includes("imageview") || type.includes("imagegeneration") || type.includes("generatedimage") + || type.includes("toolcall") || type.includes("permissionrequest") || type.includes("userinput")) { + const text = textFromValue(item.output) + ?? textFromValue(item.result) + ?? textFromValue(item.content) + ?? textFromValue(item.summary) + ?? textFromValue(item.text) + ?? textFromValue(item.query) + ?? textFromValue(item.name); + if (!text) return undefined; + const label = type.includes("websearch") ? "搜索" + : type.includes("image") ? "查看图像" + : type.includes("permissionrequest") ? "等待授权" + : type.includes("userinput") ? "等待输入" + : type.includes("mcp") ? "MCP 工具" + : "工具"; + return { outputText: text, role: "tool", kind: "tool", text, label }; + } + const text = textFromValue(item.text) ?? textFromValue(item.output); + return text ? { outputText: text, role: "assistant", kind: "assistant", text } : undefined; +} + +/** + * Official conversation messages can carry collaboration records in a + * metadata envelope instead of exposing them as a top-level `type`. Normalize + * direct item names here; metadata variants are added by + * `itemDisplayVariants` so the parent message is retained as well. + */ +function normalizeOfficialItem(item: Record): Record { + const directType = String(item.type ?? item.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); + if (directType === "collabagenttoolcall") { + return { ...item, type: "collabAgentToolCall" }; + } + if (directType === "subagentactivity") { + return { ...item, type: "subAgentActivity" }; + } + return item; +} + +/** Return the normal item plus any collaboration records attached as metadata. */ +function itemDisplayVariants(item: Record): ItemDisplayProjection[] { + const displays: ItemDisplayProjection[] = []; + const base = itemDisplay(item); + if (base) displays.push(base); + const metadata = parseRecord(item.metadata); + const candidates: Array<{ key: string; type: "collabAgentToolCall" | "subAgentActivity" }> = [ + { key: "codex_collab_agent_tool_call", type: "collabAgentToolCall" }, + { key: "codex_sub_agent_activity", type: "subAgentActivity" }, + ]; + for (const candidate of candidates) { + const value = parseRecord(metadata?.[candidate.key]); + if (!value) continue; + const normalized: Record = { ...value, type: candidate.type }; + // A direct item may carry a copy of its own metadata record. Do not render + // that record twice when the ids identify the same official item. + const normalizedId = stringValue(normalized.id); + if (normalizedId && displays.some((display) => display.projectionId === normalizedId)) continue; + const display = itemDisplay(normalized); + if (display) displays.push(display); + } + return displays; +} + +function parseRecord(value: unknown): Record | undefined { + if (isRecord(value)) return value; + if (typeof value !== "string") return undefined; + try { + const parsed: unknown = JSON.parse(value); + return isRecord(parsed) ? parsed : undefined; + } catch { + return undefined; + } +} + +function stringArray(value: unknown): string[] { + return Array.isArray(value) ? value.filter((entry): entry is string => typeof entry === "string") : []; +} + +function nullableString(value: unknown): string | null | undefined { + return typeof value === "string" ? value : value === null ? null : undefined; +} + +function redactNullableString(value: unknown): string | null | undefined { + const normalized = nullableString(value); + return normalized === undefined || normalized === null ? normalized : redactText(normalized); +} + +function collabAgentStates(value: unknown): JsonObject { + const record = parseRecord(value); + if (!record) return {}; + const states: JsonObject = {}; + for (const [threadId, rawState] of Object.entries(record)) { + const state = parseRecord(rawState); + if (!state || typeof state.status !== "string") continue; + states[threadId] = redactJson({ + status: state.status, + ...(state.message === null || typeof state.message === "string" ? { message: state.message } : {}), + }) as JsonValue; + } + return states; +} + +function collabAgentToolLabel(tool: string): string { + switch (tool) { + case "spawnAgent": return "启动子代理"; + case "sendInput": return "向子代理发送输入"; + case "resumeAgent": return "恢复子代理"; + case "wait": return "等待子代理"; + case "closeAgent": return "关闭子代理"; + default: return "子代理操作"; + } +} + +function formatAgentPath(agentPath?: string): string | null { + if (!agentPath) return null; + const leaf = agentPath.split("/").map((part) => part.trim()).filter((part) => part && part !== "root").at(-1); + if (!leaf) return null; + const normalized = leaf.replace(/[_-]+/g, " ").replace(/\s+/g, " ").trim().toLowerCase(); + return normalized ? normalized[0].toUpperCase() + normalized.slice(1) : null; +} + +function subagentActivityDisplayStatus(kind: string): string { + switch (kind) { + case "started": return "active"; + case "interacted": return "updated"; + case "interrupted": return "interrupted"; + case "completed": return "completed"; + default: return "active"; + } +} + +function subagentActivityText(displayName: string | null, kind: string): string { + const subject = displayName ? `子代理 ${displayName}` : "子代理"; + switch (kind) { + case "started": return `${subject} 已开始工作`; + case "interacted": return `${subject} 正在工作`; + case "interrupted": return `${subject} 已中断`; + case "completed": return `${subject} 已完成`; + default: return `${subject} ${kind}`; + } +} + +interface SubagentAccumulator extends SubagentSnapshot { + lastEventIndex: number; +} + +/** Rebuild the official subagent panel model from direct items and metadata. */ +function collectSubagents(state: JsonObject): SubagentSnapshot[] { + const agents = new Map(); + const parentThreadId = stringValue(state.id) + ?? stringValue(state.threadId) + ?? (isRecord(state.thread) ? stringValue(state.thread.id) : undefined) + ?? null; + let eventIndex = 0; + const ensure = (threadId: string): SubagentAccumulator => { + const current = agents.get(threadId); + if (current) { + current.lastEventIndex = eventIndex; + return current; + } + const created: SubagentAccumulator = { + threadId, + displayName: null, + prompt: null, + objective: null, + status: "working", + statusMessage: null, + canInteract: false, + parentThreadId, + lastEventIndex: eventIndex, + }; + agents.set(threadId, created); + return created; + }; + const consumeTurn = (turn: Record): void => { + const turnStartedAtMs = timestampMs(turn.turnStartedAtMs) + ?? timestampMs(turn.startedAtMs) + ?? timestampMs(turn.startedAt) + ?? timestampMs(turn.createdAtMs) + ?? timestampMs(turn.createdAt); + const turnCompletedAtMs = timestampMs(turn.completedAtMs) + ?? timestampMs(turn.completedAt) + ?? (turnStartedAtMs !== undefined && numberValue(turn.durationMs) !== undefined + ? turnStartedAtMs + (numberValue(turn.durationMs) as number) + : undefined); + for (const rawItem of Array.isArray(turn.items) ? turn.items : []) { + if (!isRecord(rawItem)) continue; + const itemStartedAtMs = timestampMs(rawItem.startedAtMs) + ?? timestampMs(rawItem.startedAt) + ?? turnStartedAtMs; + const itemCompletedAtMs = timestampMs(rawItem.completedAtMs) + ?? timestampMs(rawItem.finishedAtMs) + ?? timestampMs(rawItem.completedAt) + ?? (itemStartedAtMs !== undefined && numberValue(rawItem.durationMs) !== undefined + ? itemStartedAtMs + (numberValue(rawItem.durationMs) as number) + : turnCompletedAtMs); + for (const item of officialSubagentItems(rawItem)) { + eventIndex += 1; + const type = String(item.type ?? item.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); + if (type === "subagentactivity") { + const threadId = stringValue(item.agentThreadId); + if (!threadId) continue; + const agent = ensure(threadId); + const agentPath = stringValue(item.agentPath); + const displayName = formatAgentPath(agentPath); + const activityKind = stringValue(item.kind) ?? "started"; + if (displayName) agent.displayName = redactText(displayName); + if (agentPath) agent.agentPath = redactText(agentPath); + if (agent.startedAtMs == null && itemStartedAtMs !== undefined) agent.startedAtMs = itemStartedAtMs; + agent.statusMessage = null; + if (activityKind === "interrupted" || activityKind === "completed") { + agent.status = "done"; + if (itemCompletedAtMs !== undefined) agent.completedAtMs = itemCompletedAtMs; + } else { + agent.status = "working"; + agent.completedAtMs = null; + } + continue; + } + if (type !== "collabagenttoolcall") continue; + const tool = stringValue(item.tool) ?? ""; + const toolStatus = statusValue(item.status) ?? "inProgress"; + const receivers = stringArray(item.receiverThreadIds ?? item.receiverThreads); + const states = parseRecord(item.agentsStates) ?? {}; + const prompt = nullableString(item.prompt); + const model = nullableString(item.model); + + for (const threadId of new Set([...receivers, ...Object.keys(states)])) { + if (!threadId) continue; + const agent = ensure(threadId); + if (agent.startedAtMs == null && itemStartedAtMs !== undefined) agent.startedAtMs = itemStartedAtMs; + if (tool === "spawnAgent") { + if (prompt?.trim()) { + agent.prompt = redactText(prompt.trim()); + agent.objective = agent.prompt; + } + if (model !== undefined) agent.model = model === null ? null : redactText(model); + agent.canInteract = true; + } else if (tool === "sendInput" || tool === "resumeAgent") { + agent.canInteract = true; + } + if (toolStatus === "failed") { + agent.status = "failed"; + if (itemCompletedAtMs !== undefined) agent.completedAtMs = itemCompletedAtMs; + } else if (tool === "spawnAgent" || tool === "sendInput" || tool === "resumeAgent") { + agent.status = "working"; + agent.statusMessage = null; + agent.completedAtMs = null; + } else if (tool === "closeAgent" && toolStatus === "completed") { + agent.status = "done"; + if (itemCompletedAtMs !== undefined) agent.completedAtMs = itemCompletedAtMs; + } + } + + for (const [threadId, rawState] of Object.entries(states)) { + const agentState = parseRecord(rawState); + if (!agentState || typeof agentState.status !== "string") continue; + const agent = ensure(threadId); + agent.status = coarseSubagentStatus(agentState.status); + if (agent.status === "waiting" || agent.status === "working") { + agent.statusMessage = null; + agent.completedAtMs = null; + } else { + const statusMessage = nullableString(agentState.message); + agent.statusMessage = statusMessage?.trim() ? redactText(statusMessage.trim()) : null; + if (itemCompletedAtMs !== undefined) agent.completedAtMs = itemCompletedAtMs; + } + } + + // A completed wait means the parent observed all currently active + // receivers finishing, even if older state records were not patched. + if (tool === "wait" && toolStatus === "completed") { + for (const agent of agents.values()) { + if (agent.status !== "waiting" && agent.status !== "working") continue; + agent.status = "done"; + agent.statusMessage = null; + agent.lastEventIndex = eventIndex; + if (itemCompletedAtMs !== undefined) agent.completedAtMs = itemCompletedAtMs; + } + } + } + } + }; + + for (const turn of orderedHistoryTurns(state)) consumeTurn(turn); + if (Array.isArray(state.turns)) for (const turn of state.turns) if (isRecord(turn)) consumeTurn(turn); + + // The official UI closes stale active rows when the parent turn is no longer + // live. This prevents an old running state from lingering after reconnect. + if (!deriveTurn(state).active) { + for (const agent of agents.values()) { + if (agent.status !== "waiting" && agent.status !== "working") continue; + agent.status = "done"; + agent.statusMessage = null; + } + } + + return Array.from(agents.values(), ({ lastEventIndex: _lastEventIndex, ...agent }) => agent); +} + +function officialSubagentItems(item: Record): Record[] { + const result: Record[] = []; + const seen = new Set(); + const add = (value: Record, type: "collabAgentToolCall" | "subAgentActivity"): void => { + const normalized: Record = { ...value, type }; + const id = stringValue(normalized.id); + const key = `${type}:${id ?? JSON.stringify(normalized)}`; + if (seen.has(key)) return; + seen.add(key); + result.push(normalized); + }; + const directType = String(item.type ?? item.kind ?? "").replace(/[\s/_-]+/g, "").toLowerCase(); + if (directType === "collabagenttoolcall") add(item, "collabAgentToolCall"); + if (directType === "subagentactivity") add(item, "subAgentActivity"); + const metadata = parseRecord(item.metadata); + const collab = parseRecord(metadata?.codex_collab_agent_tool_call); + if (collab) add(collab, "collabAgentToolCall"); + const activity = parseRecord(metadata?.codex_sub_agent_activity); + if (activity) add(activity, "subAgentActivity"); + return result; +} + +function coarseSubagentStatus(status: string): "waiting" | "working" | "done" | "failed" { + switch (normalizeStatus(status)) { + case "pendinginit": return "waiting"; + case "running": return "working"; + case "completed": + case "interrupted": + case "shutdown": return "done"; + case "errored": + case "notfound": return "failed"; + default: return "working"; + } +} + +function planDisplayText(value: unknown): string | undefined { + if (typeof value === "string") return value; + if (!Array.isArray(value)) { + if (isRecord(value)) return planDisplayText(value.plan ?? value.steps ?? value.todos ?? value.text ?? value.content); + return textFromValue(value); + } + const lines = value.map((entry) => { + if (!isRecord(entry)) return textFromValue(entry); + const status = statusValue(entry.status) ?? "pending"; + const marker = ["completed", "complete", "done", "success", "succeeded"].includes(status.toLowerCase()) ? "[x]" : "[ ]"; + const label = stringValue(entry.step) ?? stringValue(entry.text) ?? stringValue(entry.title) ?? stringValue(entry.description); + return label ? `${marker} ${label}` : undefined; + }).filter((line): line is string => Boolean(line)); + return lines.length ? lines.join("\n") : undefined; +} + +function itemDisplayText(item: Record): string | undefined { + return itemDisplay(item)?.outputText; +} + +/** + * Keep the command action shape used by the official webview while making it + * safe to send over the relay. Older snapshots use `command`, whereas the + * webview's normalized action uses `cmd`; expose both aliases so either + * renderer can consume the projection without losing the original fields. + */ +function projectCommandActions(value: unknown): JsonValue[] { + const actions = commandActionList(value); + if (!actions.length) return []; + const projected: JsonValue[] = []; + for (const rawAction of actions) { + const action = redactJson(rawAction); + if (isRecord(action)) { + const command = stringValue(action.command) ?? stringValue(action.cmd); + if (command && action.command === undefined) action.command = command; + if (command && action.cmd === undefined) action.cmd = command; + if (Object.keys(action).length > 0) projected.push(action as JsonObject); + continue; + } + if (typeof action === "string" && action.trim()) projected.push(action); + } + return projected; +} + +function commandActionText(value: unknown): string | undefined { + if (typeof value === "string") return commandValue(value); + if (!isRecord(value)) return undefined; + const command = stringValue(value.command) ?? stringValue(value.cmd); + return command ? commandValue(command) : undefined; +} + +const SHELL_BOOTSTRAP = /^(?:.*[/\\])?(?:bash|cmd(?:\.exe)?|fish|powershell(?:\.exe)?|pwsh(?:\.exe)?|sh|zsh)(?:\s|$)/i; + +function isShellBootstrapCommand(value: string): boolean { + return SHELL_BOOTSTRAP.test(value.trim()); +} + +function commandValue(value: unknown): string | undefined { + if (typeof value === "string") { + const command = value.trim(); + if (!command) return undefined; + // A few IPC versions put the shell wrapper and the user command in one + // string instead of an argv array. Strip only the wrapper here; the + // frontend remains responsible for presentation-level quote cleanup. + const wrapped = command.match(/^(?:.*[/\\])?(?:bash|cmd(?:\.exe)?|fish|powershell(?:\.exe)?|pwsh(?:\.exe)?|sh|zsh)\s+-(?:l?c|c?l)\s+([\s\S]+)$/i); + return (wrapped?.[1] ?? command).trim() || undefined; + } + if (!Array.isArray(value)) return undefined; + const parts = value.filter((part): part is string => typeof part === "string").map((part) => part.trim()).filter(Boolean); + if (!parts.length) return undefined; + // The IPC snapshot may preserve the process argv (`zsh -lc `) + // rather than the command string shown by the official disclosure. + if (parts.length >= 3 && SHELL_BOOTSTRAP.test(parts[0]) && /^-(?:l?c|c?l)$/i.test(parts[1])) { + return parts.slice(2).join(" ").trim() || undefined; + } + return parts.join(" ").trim() || undefined; +} + +function commandText(item: Record): string | undefined { + // The official renderer walks actions backwards and displays the last + // non-bootstrap command. This matters when the first action is just the + // shell wrapper used to launch the process. + const actions = commandActionsValue(item); + const actionList = commandActionList(actions); + for (let index = actionList.length - 1; index >= 0; index -= 1) { + const candidate = commandActionText(actionList[index]); + if (candidate && !isShellBootstrapCommand(candidate)) return candidate; + } + const command = commandValue(item.command) ?? stringValue(item.commandLine)?.trim(); + if (!command || isShellBootstrapCommand(command)) return undefined; + return command; +} + +/** Read the parsed command-action field across live and persisted schemas. */ +function commandActionsValue(item: Record): unknown { + return item.commandActions + ?? item.command_actions + ?? item.parsedCmd + ?? item.parsed_cmd; +} + +/** Normalize live/persisted command actions, which may be a single object. */ +function commandActionList(value: unknown): unknown[] { + if (Array.isArray(value)) return value; + return isRecord(value) ? [value] : []; +} + +function isCommandActionValue(value: unknown): boolean { + return Array.isArray(value) || isRecord(value); +} + +function readPathsFromActions(value: unknown): string[] { + const actions = commandActionList(value); + if (!actions.length) return []; + const paths: string[] = []; + const seen = new Set(); + for (const raw of actions) { + if (!isRecord(raw)) continue; + const type = normalizeStatus(raw.type); + if (type !== "read") continue; + const path = stringValue(raw.path) ?? stringValue(raw.filePath) ?? stringValue(raw.file_path) ?? stringValue(raw.name); + if (!path) continue; + const safe = redactText(path.trim()); + if (!safe || seen.has(safe)) continue; + seen.add(safe); + paths.push(safe); + } + return paths.slice(0, 128); +} + +function readPathFromItem(item: Record): string | undefined { + const value = item.path ?? item.filePath ?? item.file_path ?? item.file ?? item.name; + return typeof value === "string" && value.trim() ? redactText(value.trim()) : undefined; +} + +function fileChangesText(value: unknown): string | undefined { + if (!Array.isArray(value)) return textFromValue(value); + const changes = value.map((change) => { + if (!isRecord(change)) return textFromValue(change); + const file = stringValue(change.path) ?? stringValue(change.filePath) ?? stringValue(change.file) ?? stringValue(change.name); + const kind = statusValue(change.kind ?? change.type ?? change.status); + const diff = textFromValue(change.diff) ?? textFromValue(change.patch) ?? textFromValue(change.output) ?? textFromValue(change.text); + const heading = [kind ? `[${kind}]` : undefined, file].filter(Boolean).join(" "); + return [heading, diff].filter(Boolean).join("\n"); + }).filter((part): part is string => Boolean(part)); + return changes.length ? changes.join("\n\n") : undefined; +} + +function statusValue(value: unknown): string | undefined { + if (typeof value === "string") return value; + return isRecord(value) && typeof value.type === "string" ? value.type : undefined; +} + +function textFromValue(value: unknown): string | undefined { + if (typeof value === "string") return value; + if (Array.isArray(value)) { + const parts = value.map(textFromValue).filter((part): part is string => Boolean(part)); + return parts.length ? parts.join("\n") : undefined; + } + if (isRecord(value)) { + for (const key of ["text", "value", "output", "stdout", "stderr", "delta", "summary"]) { + const text = textFromValue(value[key]); + if (text) return text; + } + } + return undefined; +} + +const MODEL_CATALOG_ROOT_KEYS = ["availableModels", "models", "modelCatalog", "listModels"] as const; +const MODEL_CATALOG_CONTAINER_KEYS = new Set([ + "data", + "items", + "models", + "availableModels", + "modelCatalog", + "listModels", +]); +const MODEL_CATALOG_META_KEYS = new Set([ + "cursor", + "nextCursor", + "next_cursor", + "hasMore", + "has_more", + "total", + "name", + "label", + "description", + "provider", + "status", + "type", + "message", + "error", +]); +const MODEL_CATALOG_MAX_ENTRIES = 256; +const MODEL_CATALOG_MAX_TEXT = 512; +const MODEL_CATALOG_MAX_MODEL = 256; +const MODEL_CATALOG_MAX_EFFORTS = 32; +const MODEL_CATALOG_MAX_SCANNED = 4_096; + +/** + * Keep the model directory useful to the browser without forwarding opaque + * provider records (which may contain credentials, URLs, or internal flags). + * The official `model/list` response is normally `{ data: Model[] }`, but + * older extension builds have exposed the same data under several state keys. + */ +function projectAvailableModels(state: JsonObject): JsonValue[] { + const sources: unknown[] = []; + const addSources = (value: unknown): void => { + if (!isRecord(value)) return; + for (const key of MODEL_CATALOG_ROOT_KEYS) { + if (value[key] !== undefined) sources.push(value[key]); + } + }; + addSources(state); + for (const key of ["thread", "metadata", "conversation", "session", "latestThreadSettings", "threadSettings", "settings"]) { + addSources(state[key]); + } + + const projected: JsonObject[] = []; + const byModel = new Map(); + const visited = new Set(); + let scanned = 0; + + const add = (value: unknown, fallbackModel?: string): void => { + const item = projectAvailableModel(value, fallbackModel); + if (!item) return; + const model = stringValue(item.model); + if (!model) return; + const key = model.toLowerCase(); + const existing = byModel.get(key); + if (!existing) { + byModel.set(key, item); + projected.push(item); + return; + } + // A state patch can first expose a bare model id and later provide the + // catalog details. Fill only absent fields so explicit false/null values + // from the first projection are not accidentally overwritten. + for (const [field, fieldValue] of Object.entries(item)) { + if (existing[field] === undefined || (Array.isArray(existing[field]) && (existing[field] as unknown[]).length === 0)) { + existing[field] = fieldValue; + } + } + }; + + const collect = (value: unknown, fallbackModel?: string, depth = 0): void => { + if (projected.length >= MODEL_CATALOG_MAX_ENTRIES || scanned >= MODEL_CATALOG_MAX_SCANNED || depth > 6 || value === undefined || value === null) return; + scanned += 1; + if (typeof value === "string") { + // Strings at the root are model ids. A string under a map key is a + // display label, so retain the key as the canonical id in that case. + if (fallbackModel && isPlausibleModelMapKey(fallbackModel)) add({ model: fallbackModel, displayName: value }); + else add(value); + return; + } + if (Array.isArray(value)) { + for (const entry of value) collect(entry, undefined, depth + 1); + return; + } + if (!isRecord(value)) return; + if (visited.has(value)) return; + visited.add(value); + + let hasContainer = false; + for (const key of MODEL_CATALOG_CONTAINER_KEYS) { + if (value[key] === undefined) continue; + hasContainer = true; + collect(value[key], undefined, depth + 1); + } + + const strongIdentity = modelCatalogText(value.model, MODEL_CATALOG_MAX_MODEL) + ?? modelCatalogText(value.id, MODEL_CATALOG_MAX_MODEL) + ?? modelCatalogText(value.slug, MODEL_CATALOG_MAX_MODEL); + const directModel = modelCatalogIdentity(value); + const isModelEntry = Boolean( + (directModel && (!hasContainer || strongIdentity)) + || (fallbackModel && isPlausibleModelMapKey(fallbackModel)), + ); + if (isModelEntry) add(value, fallbackModel); + // Once a record has an identity, its scalar fields are model properties, + // not additional map entries (for example `displayName: "Sol"`). + if (isModelEntry && !hasContainer) return; + + // A map-shaped catalog (`{ "gpt-5": { displayName: ... } }`) is used by + // a few pre-model/list extension builds. Ignore pagination metadata and + // known envelopes while walking those entries. + for (const [key, child] of Object.entries(value)) { + if (MODEL_CATALOG_CONTAINER_KEYS.has(key) || MODEL_CATALOG_META_KEYS.has(key)) continue; + if (!isPlausibleModelMapKey(key)) continue; + if (hasContainer && !isRecord(child) && !Array.isArray(child)) continue; + if (isRecord(child) || Array.isArray(child)) collect(child, key, depth + 1); + else if (typeof child === "string" && isPlausibleModelMapKey(key)) collect(child, key, depth + 1); + } + }; + + for (const source of sources) collect(source); + return projected; +} + +function isPlausibleModelMapKey(value: string): boolean { + const key = value.trim(); + return Boolean(key) + && key.length <= MODEL_CATALOG_MAX_MODEL + && !MODEL_CATALOG_CONTAINER_KEYS.has(key) + && !MODEL_CATALOG_META_KEYS.has(key) + && !/(?:token|secret|password|authorization|api[_-]?key|private[_-]?key|refresh)/i.test(key); +} + +function modelCatalogText(value: unknown, maxLength = MODEL_CATALOG_MAX_TEXT): string | undefined { + if (typeof value !== "string") return undefined; + const text = redactText(value.trim()).slice(0, maxLength).trim(); + return text || undefined; +} + +function modelCatalogStrongIdentity(value: Record): string | undefined { + return modelCatalogText(value.model, MODEL_CATALOG_MAX_MODEL) + ?? modelCatalogText(value.id, MODEL_CATALOG_MAX_MODEL) + ?? modelCatalogText(value.slug, MODEL_CATALOG_MAX_MODEL); +} + +function modelCatalogIdentity(value: Record): string | undefined { + return modelCatalogStrongIdentity(value) + ?? modelCatalogText(value.name, MODEL_CATALOG_MAX_MODEL); +} + +function projectAvailableModel(value: unknown, fallbackModel?: string): JsonObject | undefined { + if (typeof value === "string") { + const model = modelCatalogText(fallbackModel ?? value, MODEL_CATALOG_MAX_MODEL); + return model ? { model } : undefined; + } + if (!isRecord(value)) return undefined; + // For map-shaped catalogs the key is the canonical model id and `name` is + // commonly only a human-readable label. Prefer explicit model/id/slug, + // then the map key, and use `name` as a legacy fallback for standalone rows. + const model = modelCatalogStrongIdentity(value) + ?? (fallbackModel && isPlausibleModelMapKey(fallbackModel) ? modelCatalogText(fallbackModel, MODEL_CATALOG_MAX_MODEL) : undefined) + ?? modelCatalogText(value.name, MODEL_CATALOG_MAX_MODEL); + if (!model) return undefined; + + const result: JsonObject = { model }; + const id = modelCatalogText(value.id, MODEL_CATALOG_MAX_MODEL); + if (id) result.id = id; + const displayName = modelCatalogText(value.displayName ?? value.label ?? (fallbackModel ? value.name : undefined)); + if (displayName) result.displayName = displayName; + const description = modelCatalogText(value.description); + if (description) result.description = description; + const specialty = modelCatalogText(value.modelSpecialty); + if (specialty) result.modelSpecialty = specialty; + for (const key of ["hidden", "isDefault"] as const) { + if (typeof value[key] === "boolean") result[key] = value[key]; + } + const upgrade = modelCatalogText(value.upgrade, MODEL_CATALOG_MAX_MODEL); + if (upgrade) result.upgrade = upgrade; + const defaultEffort = modelCatalogText(value.defaultReasoningEffort ?? value.defaultEffort, 64); + if (defaultEffort) result.defaultReasoningEffort = defaultEffort; + else if (value.defaultReasoningEffort === null || value.defaultEffort === null) result.defaultReasoningEffort = null; + + const rawEfforts = Array.isArray(value.supportedReasoningEfforts) + ? value.supportedReasoningEfforts + : Array.isArray(value.reasoningEfforts) + ? value.reasoningEfforts + : Array.isArray(value.efforts) ? value.efforts : undefined; + const efforts = projectReasoningEfforts(rawEfforts); + if (efforts) result.supportedReasoningEfforts = efforts; + return result; +} + +function projectReasoningEfforts(value: unknown[] | undefined): JsonValue[] | undefined { + if (!value) return undefined; + const seen = new Set(); + const projected: JsonValue[] = []; + for (const entry of value.slice(0, MODEL_CATALOG_MAX_EFFORTS)) { + const effort = typeof entry === "string" + ? modelCatalogText(entry, 64) + : isRecord(entry) + ? modelCatalogText(entry.reasoningEffort ?? entry.effort, 64) + : undefined; + if (!effort) continue; + const key = effort.toLowerCase(); + if (seen.has(key)) continue; + seen.add(key); + const description = isRecord(entry) ? modelCatalogText(entry.description) : undefined; + projected.push({ + reasoningEffort: effort, + ...(description ? { description } : {}), + }); + } + return projected.length ? projected : undefined; +} + +const TOKEN_USAGE_FIELDS = [ + "totalTokens", + "inputTokens", + "cachedInputTokens", + "cacheWriteInputTokens", + "outputTokens", + "reasoningOutputTokens", +] as const; + +/** Project the official thread/tokenUsage payload without forwarding limits or account data. */ +function projectTokenUsage(value: unknown): JsonObject | undefined { + if (!isRecord(value)) return undefined; + const source = isRecord(value.info) + ? value.info + : isRecord(value.tokenUsage) + ? value.tokenUsage + : isRecord(value.token_usage) + ? value.token_usage + : value; + const total = projectTokenUsageBreakdown( + source.total + ?? source.total_token_usage + ?? source.totalTokenUsage, + ); + const last = projectTokenUsageBreakdown( + source.last + ?? source.last_token_usage + ?? source.lastTokenUsage, + ); + const contextWindow = tokenNumber( + source.modelContextWindow + ?? source.model_context_window + ?? source.contextWindow + ?? source.context_window, + ); + if (!total && !last && contextWindow === undefined) return undefined; + return { + ...(total ? { total } : {}), + ...(last ? { last } : {}), + ...(contextWindow !== undefined ? { modelContextWindow: contextWindow } : {}), + }; +} + +function projectTokenUsageBreakdown(value: unknown): JsonObject | undefined { + if (!isRecord(value)) return undefined; + const aliases: Record<(typeof TOKEN_USAGE_FIELDS)[number], string[]> = { + totalTokens: ["totalTokens", "total_tokens"], + inputTokens: ["inputTokens", "input_tokens"], + cachedInputTokens: ["cachedInputTokens", "cached_input_tokens"], + cacheWriteInputTokens: ["cacheWriteInputTokens", "cache_write_input_tokens"], + outputTokens: ["outputTokens", "output_tokens"], + reasoningOutputTokens: ["reasoningOutputTokens", "reasoning_output_tokens"], + }; + const result: JsonObject = {}; + for (const field of TOKEN_USAGE_FIELDS) { + for (const alias of aliases[field]) { + const number = tokenNumber(value[alias]); + if (number === undefined) continue; + result[field] = number; + break; + } + } + return Object.keys(result).length ? result : undefined; +} + +function tokenNumber(value: unknown): number | undefined { + if (typeof value === "number") return Number.isFinite(value) && value >= 0 ? value : undefined; + if (typeof value !== "string" || !/^\d+(?:\.\d+)?$/.test(value.trim())) return undefined; + const number = Number(value); + return Number.isFinite(number) && number >= 0 ? number : undefined; +} + +/** Project display-safe thread settings from the opaque IPC state. */ +function projectSessionMetadata(state: JsonObject): JsonObject { + const candidates: unknown[] = [ + state.latestThreadSettings, + state.threadSettings, + state.settings, + isRecord(state.thread) ? state.thread.latestThreadSettings : undefined, + isRecord(state.thread) ? state.thread.settings : undefined, + ]; + // Merge compatibility locations from oldest to newest so a partial + // `latestThreadSettings` record can still inherit provider/permission + // fields exposed by older state shapes, while the official latest record + // wins when it contains the same key. + const settings: Record = {}; + for (const candidate of candidates.slice().reverse()) { + if (isRecord(candidate)) Object.assign(settings, candidate); + } + const result: JsonObject = {}; + const latestModel = settings.model ?? state.latestModel ?? state.model; + const latestReasoningEffort = settings.effort !== undefined + ? settings.effort + : Object.prototype.hasOwnProperty.call(state, "latestReasoningEffort") + ? state.latestReasoningEffort + : state.effort; + const values: Record = { + model: latestModel, + latestModel, + modelProvider: settings.modelProvider ?? state.modelProvider, + approvalPolicy: settings.approvalPolicy ?? state.approvalPolicy, + approvalsReviewer: settings.approvalsReviewer ?? state.approvalsReviewer, + sandboxPolicy: settings.sandboxPolicy ?? settings.sandbox ?? state.sandboxPolicy ?? state.sandbox, + permissions: settings.permissions ?? state.permissions, + currentPermissions: settings.currentPermissions ?? state.currentPermissions, + runtimeWorkspaceRoots: settings.runtimeWorkspaceRoots ?? state.runtimeWorkspaceRoots, + cwd: settings.cwd ?? state.cwd, + effort: latestReasoningEffort, + latestReasoningEffort, + summary: settings.summary ?? state.summary, + }; + for (const [key, value] of Object.entries(values)) { + if (value !== undefined && value !== null) result[key] = asJsonValue(value); + } + // Preserve an explicit null effort: the official state uses null when a + // model has no selectable reasoning level, and omission would make the + // browser retain a stale prior value. + if (latestReasoningEffort === null) { + result.effort = null; + result.latestReasoningEffort = null; + } + const tokenUsageSource = [ + state.latestTokenUsageInfo, + state.tokenUsage, + state.token_usage, + settings.latestTokenUsageInfo, + settings.tokenUsage, + settings.token_usage, + ].find((candidate) => candidate !== undefined); + const tokenUsage = projectTokenUsage(tokenUsageSource); + if (tokenUsage) { + // Keep both names: `latestTokenUsageInfo` is the official state key while + // `tokenUsage` is easier for relay/browser clients to consume. + result.tokenUsage = tokenUsage; + result.latestTokenUsageInfo = tokenUsage; + } else if (tokenUsageSource === null) { + result.tokenUsage = null; + result.latestTokenUsageInfo = null; + } + if (!result.title && isRecord(state.thread)) { + const title = state.thread.name ?? state.thread.title ?? state.thread.preview; + if (title !== undefined && title !== null) result.title = asJsonValue(title); + } + const availableModels = projectAvailableModels(state); + if (availableModels.length) { + // `availableModels` is the current relay field; `models` preserves the + // name used by older browser clients and by the app-server response. + result.availableModels = availableModels; + result.models = availableModels; + } + return result; +} + +function stringValue(value: unknown): string | undefined { return typeof value === "string" ? value : undefined; } +function numberValue(value: unknown): number | undefined { return typeof value === "number" && Number.isFinite(value) ? value : undefined; } +function normalizeStatus(value: unknown): string { return typeof value === "string" ? value.replace(/[- ]/g, "_").toLowerCase() : "unknown"; } +function cloneObject(value: JsonObject): JsonObject { return JSON.parse(JSON.stringify(value)) as JsonObject; } + +const SECRET_KEY = /(?:token|secret|password|authorization|api[_-]?key|private[_-]?key|refresh)/i; +const SECRET_VALUE = /(?:Bearer\s+)[A-Za-z0-9._~+\-/]+=*|(?:sk-[A-Za-z0-9_-]{12,}|gh[pousr]_[A-Za-z0-9]{12,})/g; +function redactText(text: string): string { + return text.replace(SECRET_VALUE, "[REDACTED]").replace(/([?&](?:token|key|secret|password|api[_-]?key)=)[^&\s]+/gi, "$1[REDACTED]").replace(/((?:token|secret|password|api[_-]?key)\s*[:=]\s*)[^\s,;]+/gi, "$1[REDACTED]"); +} +function redactJson(value: unknown): JsonValue { + if (Array.isArray(value)) return value.map((item) => redactJson(item)); + if (isRecord(value)) { + const result: JsonObject = {}; + for (const [key, child] of Object.entries(value)) result[key] = SECRET_KEY.test(key) ? "[REDACTED]" : redactJson(child); + return result; + } + return typeof value === "string" ? redactText(value) : asJsonValue(value); +} +function hashJson(value: JsonValue): string { return createHash("sha256").update(stableStringify(value)).digest("hex"); } +function stableStringify(value: JsonValue): string { + if (Array.isArray(value)) return `[${value.map(stableStringify).join(",")}]`; + if (value !== null && typeof value === "object") return `{${Object.keys(value).sort().map((key) => `${JSON.stringify(key)}:${stableStringify(value[key] ?? null)}`).join(",")}}`; + return JSON.stringify(value); +} diff --git a/aether-vscodex/vscode-extension/src/codexPath.ts b/aether-vscodex/vscode-extension/src/codexPath.ts new file mode 100644 index 000000000..44873fa9c --- /dev/null +++ b/aether-vscodex/vscode-extension/src/codexPath.ts @@ -0,0 +1,152 @@ +import { accessSync, constants, Dirent, readdirSync, statSync } from "node:fs"; +import { homedir } from "node:os"; +import { delimiter, isAbsolute, join, sep } from "node:path"; + +export interface CodexPathOptions { + /** Environment used for PATH lookup. Defaults to the extension host environment. */ + env?: NodeJS.ProcessEnv; + /** Home directory used when looking for bundled installations. */ + homeDir?: string; + /** Platform override for deterministic tests. */ + platform?: NodeJS.Platform; +} + +/** + * Resolve the executable used by the VS Code bridge. + * + * VS Code launched from Finder/Dock often receives a smaller PATH than a shell. + * The default `codex` command therefore gets a few explicit installation + * fallbacks, while a user-supplied command remains authoritative. + */ +export function resolveCodexCommand(configuredCommand = "codex", options: CodexPathOptions = {}): string { + const command = configuredCommand.trim() || "codex"; + const env = options.env ?? process.env; + const platform = options.platform ?? process.platform; + const home = options.homeDir ?? homedir(); + + if (hasPathComponent(command, platform)) { + const resolved = executablePath(command, platform); + if (resolved) return resolved; + throw missingCodexError(command, platform); + } + + const fromPath = findOnPath(command, env.PATH, platform, env.PATHEXT); + if (fromPath) return fromPath; + + // Only the default command gets installation-specific fallbacks. A custom + // bare command should fail loudly instead of silently running another binary. + if (!isDefaultCommand(command, platform)) throw missingCodexError(command, platform); + + for (const candidate of bundledCandidates(home, platform)) { + const resolved = executablePath(candidate, platform); + if (resolved) return resolved; + } + + throw missingCodexError(command, platform); +} + +export function missingCodexError(command: string, platform: NodeJS.Platform = process.platform): Error { + const examples = platform === "darwin" + ? ' Set "codexRemoteCollab.codexCommand" to the full path, for example "/Applications/ChatGPT.app/Contents/Resources/codex".' + : ' Set "codexRemoteCollab.codexCommand" to the full path of the Codex executable.'; + return new Error(`Codex executable "${command}" was not found.${examples}`); +} + +function isDefaultCommand(command: string, platform: NodeJS.Platform): boolean { + return platform === "win32" ? command.toLowerCase() === "codex" || command.toLowerCase() === "codex.exe" : command === "codex"; +} + +function hasPathComponent(command: string, platform: NodeJS.Platform): boolean { + return isAbsolute(command) || command.includes(sep) || (platform === "win32" && command.includes("\\")); +} + +function executablePath(candidate: string, platform: NodeJS.Platform): string | undefined { + try { + const info = statSync(candidate); + if (!info.isFile()) return undefined; + // X_OK is meaningful on POSIX; Windows still benefits from the file check. + if (platform !== "win32") accessSync(candidate, constants.X_OK); + return candidate; + } catch { + return undefined; + } +} + +function findOnPath(command: string, pathValue: string | undefined, platform: NodeJS.Platform, pathextValue?: string): string | undefined { + if (!pathValue) return undefined; + const extensions = platform === "win32" ? windowsExtensions(command, pathextValue) : [""]; + for (const directory of pathValue.split(delimiter)) { + if (!directory) continue; + for (const extension of extensions) { + const candidate = join(directory, `${command}${extension}`); + const resolved = executablePath(candidate, platform); + if (resolved) return resolved; + } + } + return undefined; +} + +function windowsExtensions(command: string, pathextValue: string | undefined): string[] { + if (/[.][^./\\]+$/.test(command)) return [""]; + const extensions = (pathextValue ?? ".COM;.EXE;.BAT;.CMD") + .split(";") + .map((value) => value.trim()) + .filter(Boolean); + return ["", ...extensions]; +} + +function bundledCandidates(home: string, platform: NodeJS.Platform): string[] { + if (platform !== "darwin") return []; + + const candidates = [ + join(home, "Applications", "ChatGPT.app", "Contents", "Resources", "codex"), + "/Applications/ChatGPT.app/Contents/Resources/codex", + join(home, ".local", "bin", "codex"), + join(home, ".npm-global", "bin", "codex"), + ]; + + for (const extensionsRoot of [ + join(home, ".vscode", "extensions"), + join(home, ".vscode-insiders", "extensions"), + ]) { + candidates.push(...officialExtensionCandidates(extensionsRoot)); + } + return candidates; +} + +function officialExtensionCandidates(extensionsRoot: string): string[] { + let entries: Dirent[]; + try { + entries = readdirSync(extensionsRoot, { withFileTypes: true, encoding: "utf8" }); + } catch { + return []; + } + + const matches = entries + .filter((entry) => entry.isDirectory() && entry.name.startsWith("openai.chatgpt-")) + .map((entry) => { + const directory = join(extensionsRoot, entry.name); + let modified = 0; + try { + modified = statSync(directory).mtimeMs; + } catch { + // Keep an unreadable entry at the end of the deterministic sort. + } + return { directory, modified }; + }) + .sort((left, right) => right.modified - left.modified || right.directory.localeCompare(left.directory)); + + const candidates: string[] = []; + for (const match of matches) { + let architectures: Dirent[]; + try { + architectures = readdirSync(join(match.directory, "bin"), { withFileTypes: true, encoding: "utf8" }); + } catch { + continue; + } + for (const architecture of architectures) { + if (architecture.isDirectory()) candidates.push(join(match.directory, "bin", architecture.name, "codex")); + } + } + return candidates; +} diff --git a/aether-vscodex/vscode-extension/src/compositeRelay.ts b/aether-vscodex/vscode-extension/src/compositeRelay.ts new file mode 100644 index 000000000..1d83ad8ac --- /dev/null +++ b/aether-vscodex/vscode-extension/src/compositeRelay.ts @@ -0,0 +1,131 @@ +import { Disposable, RelayFrame, RelayTransport } from "./protocol"; + +export interface NamedRelayTransport { + id: string; + transport: RelayTransport; + required?: boolean; +} + +/** + * Fans host events out to local and cloud relays while presenting one + * transport lifecycle to RelayHost. A temporary cloud outage must not stop + * the local bridge (and vice versa). + */ +export class CompositeRelayTransport implements RelayTransport { + readonly handlesHandshake = true; + private readonly entries: NamedRelayTransport[]; + private readonly subscriptions: Disposable[] = []; + private readonly openEntries = new Set(); + private readonly messageListeners = new Set<(frame: RelayFrame) => void>(); + private readonly openListeners = new Set<() => void>(); + private readonly closeListeners = new Set<(error?: Error) => void>(); + private started = false; + private sessionId?: string; + + constructor(entries: NamedRelayTransport[]) { + if (entries.length === 0) throw new Error("CompositeRelayTransport requires at least one relay"); + const ids = new Set(); + for (const entry of entries) { + if (!entry.id || ids.has(entry.id)) throw new Error(`duplicate relay id: ${entry.id || "(empty)"}`); + ids.add(entry.id); + } + this.entries = [...entries]; + } + + setSessionId(sessionId: string): void { + this.sessionId = sessionId; + for (const { transport } of this.entries) { + (transport as RelayTransport & { setSessionId?: (value: string) => void }).setSessionId?.(sessionId); + } + } + + async connect(): Promise { + if (this.started) return; + this.started = true; + this.bindTransports(); + if (this.sessionId) this.setSessionId(this.sessionId); + + const results = await Promise.allSettled(this.entries.map(({ transport }) => transport.connect())); + const failures = results + .map((result, index) => ({ result, entry: this.entries[index] })) + .filter((item): item is { result: PromiseRejectedResult; entry: NamedRelayTransport } => item.result.status === "rejected"); + const requiredFailure = failures.find(({ entry }) => entry.required); + const connected = results.length - failures.length; + if (requiredFailure || connected === 0) { + this.started = false; + this.disposeSubscriptions(); + for (const { transport } of this.entries) transport.close(); + const detail = failures.map(({ entry, result }) => `${entry.id}: ${errorMessage(result.reason)}`).join("; "); + throw new Error(`unable to connect relay${failures.length === 1 ? "" : "s"}: ${detail}`); + } + } + + send(frame: RelayFrame): void { + const failures: string[] = []; + for (const { id, transport } of this.entries) { + try { + transport.send(frame); + } catch (error) { + failures.push(`${id}: ${errorMessage(error)}`); + } + } + if (failures.length === this.entries.length) { + throw new Error(`all relay sends failed: ${failures.join("; ")}`); + } + } + + onMessage(listener: (frame: RelayFrame) => void): Disposable { + this.messageListeners.add(listener); + return { dispose: () => this.messageListeners.delete(listener) }; + } + + onOpen(listener: () => void): Disposable { + this.openListeners.add(listener); + return { dispose: () => this.openListeners.delete(listener) }; + } + + onClose(listener: (error?: Error) => void): Disposable { + this.closeListeners.add(listener); + return { dispose: () => this.closeListeners.delete(listener) }; + } + + isConnected(id: string): boolean { + return this.openEntries.has(id); + } + + close(): void { + this.started = false; + this.openEntries.clear(); + this.disposeSubscriptions(); + for (const { transport } of this.entries) transport.close(); + } + + private bindTransports(): void { + for (const { id, transport } of this.entries) { + this.subscriptions.push(transport.onMessage((frame) => { + for (const listener of this.messageListeners) listener(frame); + })); + if (transport.onOpen) this.subscriptions.push(transport.onOpen(() => { + this.openEntries.add(id); + // RelayHost publishes an authoritative snapshot after an authenticated + // reconnect. Surface every member reconnect so a recovered cloud relay + // is hydrated even while the local relay remained online. + for (const listener of this.openListeners) listener(); + })); + if (transport.onClose) this.subscriptions.push(transport.onClose((error) => { + const wasOpen = this.openEntries.delete(id); + if (wasOpen && this.openEntries.size === 0) { + for (const listener of this.closeListeners) listener(error); + } + })); + } + } + + private disposeSubscriptions(): void { + for (const subscription of this.subscriptions.splice(0)) subscription.dispose(); + } +} + +function errorMessage(error: unknown): string { + return error instanceof Error ? error.message : String(error); +} diff --git a/aether-vscodex/vscode-extension/src/extension.ts b/aether-vscodex/vscode-extension/src/extension.ts new file mode 100644 index 000000000..e863ac25f --- /dev/null +++ b/aether-vscodex/vscode-extension/src/extension.ts @@ -0,0 +1,574 @@ +import { hostname } from "node:os"; + +import * as vscode from "vscode"; + +import { CodexAgentAdapter } from "./codexAgentAdapter"; +import { resolveCodexCommand } from "./codexPath"; +import { CodexIpcAgentAdapter } from "./codexIpcAgentAdapter"; +import { CompositeRelayTransport } from "./compositeRelay"; +import { LocalRelayController, localRelayTarget } from "./localRelay"; +import { AgentAdapter, ControlMode, Disposable, JsonObject, Logger } from "./protocol"; +import { RelayClient } from "./relayClient"; +import { RelayHost } from "./relayHost"; +import { SwitchableAgentAdapter } from "./switchableAgentAdapter"; + +let activeHost: RelayHost | undefined; +let activeAdapter: AgentAdapter | undefined; +let activeRelay: CompositeRelayTransport | undefined; +let activeAdapterStatusSubscription: Disposable | undefined; +let statusItem: vscode.StatusBarItem | undefined; +let autoStartRetryTimer: NodeJS.Timeout | undefined; +let autoStartRetryMs = 3_000; +let localRelayController: LocalRelayController | undefined; +const t = (message: string, ...args: Array): string => vscode.l10n.t(message, ...args); + +export async function activate(context: vscode.ExtensionContext): Promise { + const output = vscode.window.createOutputChannel(t("Codex Remote Collaboration")); + context.subscriptions.push(output); + const logger = { + debug: (message: string, ...args: unknown[]) => output.appendLine(`[debug] ${message} ${formatArgs(args)}`), + info: (message: string, ...args: unknown[]) => output.appendLine(`[info] ${message} ${formatArgs(args)}`), + warn: (message: string, ...args: unknown[]) => output.appendLine(`[warn] ${message} ${formatArgs(args)}`), + error: (message: string, ...args: unknown[]) => output.appendLine(`[error] ${message} ${formatArgs(args)}`), + }; + localRelayController = new LocalRelayController({ extensionPath: context.extensionPath, logger }); + + statusItem = vscode.window.createStatusBarItem(vscode.StatusBarAlignment.Right, 100); + statusItem.command = "codexRemoteCollab.openWeb"; + statusItem.text = "$(plug) Codex Remote"; + statusItem.tooltip = t("Connecting to the local Codex collaboration service"); + statusItem.show(); + context.subscriptions.push(statusItem); + + const start = async (automatic = false): Promise => { + if (!automatic && autoStartRetryTimer) { + clearTimeout(autoStartRetryTimer); + autoStartRetryTimer = undefined; + } + if (activeHost) { + if (!automatic) vscode.window.showInformationMessage(t("The Codex remote bridge is already running.")); + return; + } + const configuration = vscode.workspace.getConfiguration("codexRemoteCollab"); + const relayConfiguration = resolveRelayConfiguration(configuration); + const localRelayUrl = relayConfiguration.localUrl; + if (!localRelayUrl) { + vscode.window.showWarningMessage(t("Set codexRemoteCollab.localRelayUrl before starting the bridge.")); + return; + } + const localTarget = localRelayTarget(localRelayUrl); + if (!localTarget) { + vscode.window.showErrorMessage(t("codexRemoteCollab.localRelayUrl must be a loopback ws:// address.")); + return; + } + if (localTarget && configuration.get("autoStartLocalRelay", true)) { + setStatus("$(sync~spin) Codex Remote", t("Starting {0}", localTarget.webUrl), "codexRemoteCollab.openWeb"); + try { + await localRelayController?.ensureRunning(localRelayUrl); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + setStatus("$(error) Codex Remote", t("Unable to start the local collaboration service: {0}", message), "codexRemoteCollab.openWeb"); + if (automatic) { + logger.warn(`Automatic local relay start failed; retrying in ${autoStartRetryMs}ms`, error); + scheduleAutoStartRetry(start); + } else { + vscode.window.showErrorMessage(t("Unable to start the local Codex collaboration service: {0}", message)); + } + return; + } + } + const legacyToken = await context.secrets.get("codexRemoteCollab.relayToken"); + const localToken = relayConfiguration.legacyRemote ? undefined : legacyToken; + if (!localToken) logger.info("Using the loopback-only unauthenticated local relay"); + const initialControlMode = resolveInitialControlMode(configuration); + let currentControlMode: ControlMode = initialControlMode; + const createAdapter = (controlMode: ControlMode): AgentAdapter => { + if (controlMode === "sync") { + const configuredThreadId = configuration.get("threadId", "").trim(); + const socketPath = configuration.get("ipcSocketPath", "").trim(); + logger.info(`Synchronous mode enabled; following the VS Code Codex panel${configuredThreadId ? ` (initial conversation ${configuredThreadId})` : ""}`); + return new CodexIpcAgentAdapter({ + threadId: configuredThreadId || undefined, + socketPath: socketPath || undefined, + hostId: configuration.get("hostId", "local"), + autoDiscoverThread: configuration.get("autoDiscoverThread", true), + // Synchronous mode has one navigation owner: the official panel. + followVscodeSession: true, + preferredCwds: workspaceRoots(), + strictVersions: configuration.get("ipcStrictVersions", true), + logger, + approvalTimeoutMs: configuration.get("approvalTimeoutMs", 300_000), + openNewSession: () => openOfficialNewSession(logger), + }); + } + + const configuredCommand = configuration.get("codexCommand", "codex"); + const command = resolveCodexCommand(configuredCommand); + const args = configuration.get("codexArgs", ["app-server", "--stdio"]); + const defaultCwd = configuration.get("defaultCwd", "") || firstWorkspaceRoot(); + logger.info(`Asynchronous mode enabled; using independent Codex executable: ${command}`); + return new CodexAgentAdapter({ + command, + args, + defaultCwd: defaultCwd || undefined, + logger, + approvalTimeoutMs: configuration.get("approvalTimeoutMs", 300_000), + }); + }; + const adapter = new SwitchableAgentAdapter({ + initialMode: initialControlMode, + createAdapter, + logger, + onModeChanged: async (nextMode) => { + currentControlMode = nextMode; + await configuration.update("controlMode", nextMode, vscode.ConfigurationTarget.Global); + setControlModeStatus(nextMode, nextMode === "async" || Boolean((await adapter.snapshot()).threadId)); + }, + }); + const localRelay = new RelayClient({ + url: localRelayUrl, + ...(localToken ? { accessToken: localToken } : {}), + reconnect: configuration.get("relayReconnect", true), + logger, + }); + const relayEntries = [{ id: "local", transport: localRelay, required: true }]; + const cloudRelayUrl = relayConfiguration.cloudUrl; + const cloudToken = await context.secrets.get("codexRemoteCollab.cloudRelayToken") + ?? (relayConfiguration.legacyRemote ? legacyToken : undefined); + if (cloudRelayUrl && cloudToken) { + relayEntries.push({ + id: "aether-cloud", + transport: new RelayClient({ + url: cloudRelayUrl, + accessToken: cloudToken, + reconnect: configuration.get("relayReconnect", true), + logger, + }), + required: false, + }); + logger.info(`Aether cloud relay enabled: ${cloudRelayUrl}`); + } else if (cloudRelayUrl) { + logger.warn("Aether cloud relay URL is configured without a device credential; cloud sync is disabled until pairing is completed"); + } + const relay = new CompositeRelayTransport(relayEntries); + const capabilities = ["read_output", "send_task_input", "cancel_task", "approve_low_risk"]; + if (configuration.get("allowHighRiskApprovals", false)) capabilities.push("approve_high_risk"); + const host = new RelayHost({ adapter, relay, logger, capabilities }); + let controlReady = initialControlMode === "async"; + activeAdapterStatusSubscription?.dispose(); + const adapterStatusSubscription = adapter.onEvent((event) => { + if (activeAdapter !== adapter) return; + if (event.type === "control.mode.changed") { + const changedMode = event.payload.controlMode; + if (changedMode === "sync" || changedMode === "async") currentControlMode = changedMode; + } + if (event.type !== "session.snapshot") return; + const metadata = event.payload.metadata; + if (metadata !== null && typeof metadata === "object" && !Array.isArray(metadata)) { + const snapshotMode = (metadata as JsonObject).controlMode; + if (snapshotMode === "sync" || snapshotMode === "async") currentControlMode = snapshotMode; + } + const waiting = event.payload.state === "waiting_for_host" + || (metadata !== null && typeof metadata === "object" && !Array.isArray(metadata) + && (metadata as JsonObject).waitingForSession === true); + const threadId = event.threadId + ?? (typeof event.payload.threadId === "string" ? event.payload.threadId : undefined); + controlReady = currentControlMode === "async" || (Boolean(threadId) && !waiting); + setControlModeStatus(currentControlMode, controlReady); + }); + activeAdapterStatusSubscription = adapterStatusSubscription; + activeAdapter = adapter; + activeRelay = relay; + activeHost = host; + if (configuration.get("autoStartLocalRelay", true)) { + localRelay.onClose(() => { + if (activeHost !== host) return; + setStatus("$(sync~spin) Codex Remote", t("Restoring the local collaboration service"), "codexRemoteCollab.openWeb"); + void localRelayController?.ensureRunning(localRelayUrl).catch((error) => { + logger.warn("Unable to recover bundled local relay", error); + setStatus("$(error) Codex Remote", t("Unable to restore the local collaboration service"), "codexRemoteCollab.openWeb"); + }); + }); + localRelay.onOpen(() => { + if (activeHost === host) { + setControlModeStatus(currentControlMode, controlReady); + } + }); + } + try { + await host.start(); + autoStartRetryMs = 3_000; + const snapshot = await adapter.snapshot(); + const snapshotMode = snapshot.metadata?.controlMode; + if (snapshotMode === "sync" || snapshotMode === "async") currentControlMode = snapshotMode; + controlReady = currentControlMode === "async" || Boolean(snapshot.threadId); + setControlModeStatus(currentControlMode, controlReady); + if (!automatic && (currentControlMode === "async" || controlReady)) { + vscode.window.showInformationMessage(currentControlMode === "sync" + ? t("The Codex remote bridge attached to the existing VS Code Codex conversation.") + : t("The independent Codex remote mode connected.")); + } + } catch (error) { + activeHost = undefined; + activeAdapter = undefined; + activeRelay = undefined; + if (activeAdapterStatusSubscription === adapterStatusSubscription) { + activeAdapterStatusSubscription.dispose(); + activeAdapterStatusSubscription = undefined; + } + await host.stop().catch(() => undefined); + const message = error instanceof Error ? error.message : String(error); + if (initialControlMode === "sync" && isAttachSessionUnavailable(message)) { + setStatus("$(sync~spin) Codex Remote", t("Waiting for a Codex conversation to open in VS Code. It will connect automatically."), "codexRemoteCollab.openWeb"); + logger.info(`No attachable VS Code Codex session is available; retrying in ${autoStartRetryMs}ms`); + scheduleAutoStartRetry(start); + return; + } + setStatus("$(error) Codex Remote", initialControlMode === "sync" ? t("The Codex conversation is not connected") : t("The independent Codex mode is not connected"), "codexRemoteCollab.openWeb"); + if (automatic) { + logger.warn(`Automatic bridge start failed; retrying in ${autoStartRetryMs}ms`, error); + scheduleAutoStartRetry(start); + } else { + const detail = localTarget && /ECONNREFUSED|connect refused/i.test(message) + ? t("The local collaboration service at {0} is temporarily unavailable. The extension will keep retrying.", localTarget.webUrl) + : t("Unable to start the Codex remote bridge: {0}", message); + vscode.window.showErrorMessage(detail); + } + } + }; + + const stop = async (): Promise => { + if (autoStartRetryTimer) { + clearTimeout(autoStartRetryTimer); + autoStartRetryTimer = undefined; + } + autoStartRetryMs = 3_000; + const host = activeHost; + activeHost = undefined; + activeAdapter = undefined; + activeRelay = undefined; + activeAdapterStatusSubscription?.dispose(); + activeAdapterStatusSubscription = undefined; + if (host) await host.stop(); + setStatus("$(plug) Codex Remote", t("Bridge paused. Click to open the web control and resume automatically."), "codexRemoteCollab.openWeb"); + }; + + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.openWeb", async () => { + const localRelayUrl = resolveRelayConfiguration(vscode.workspace.getConfiguration("codexRemoteCollab")).localUrl; + const webUrl = localRelayController?.getWebUrl(localRelayUrl); + if (!webUrl) { + vscode.window.showErrorMessage(t("The local collaboration URL is invalid. Check codexRemoteCollab.localRelayUrl.")); + return; + } + if (!activeHost) await start(false); + if (activeHost) await vscode.env.openExternal(vscode.Uri.parse(webUrl)); + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.start", start)); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.stop", stop)); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.setThreadId", async () => { + const configuration = vscode.workspace.getConfiguration("codexRemoteCollab"); + const current = configuration.get("threadId", ""); + const value = await vscode.window.showInputBox({ + prompt: t("Existing Codex conversation ID (leave blank for auto-discovery)"), + value: current, + ignoreFocusOut: true, + }); + if (value === undefined) return; + await configuration.update("threadId", value.trim(), vscode.ConfigurationTarget.Global); + vscode.window.showInformationMessage(value.trim() + ? t("Codex Remote will attach to {0} after the next bridge start.", value.trim()) + : t("Codex Remote will auto-discover the latest VS Code Codex conversation after the next bridge start.")); + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.setRelayToken", async () => { + const token = await vscode.window.showInputBox({ prompt: t("Relay access token (leave blank for the local relay)"), password: true, ignoreFocusOut: true }); + if (token === undefined) return; + await context.secrets.store("codexRemoteCollab.relayToken", token); + vscode.window.showInformationMessage(t("Relay token stored in VS Code SecretStorage.")); + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.configureCloud", async () => { + const configuration = vscode.workspace.getConfiguration("codexRemoteCollab"); + const currentUrl = configuration.get("cloudRelayUrl", ""); + const url = await vscode.window.showInputBox({ + prompt: t("Aether cloud relay WebSocket URL"), + value: currentUrl, + placeHolder: "wss://aether.example.com/api/vscodex/ws", + ignoreFocusOut: true, + validateInput: validateCloudRelayUrl, + }); + if (url === undefined) return; + if (!url.trim()) { + await configuration.update("cloudRelayUrl", "", vscode.ConfigurationTarget.Global); + await context.secrets.delete("codexRemoteCollab.cloudRelayToken"); + vscode.window.showInformationMessage(t("Aether cloud connection removed. Local control remains enabled.")); + return; + } + const token = await vscode.window.showInputBox({ + prompt: t("Device credential from the Aether pairing flow"), + password: true, + ignoreFocusOut: true, + }); + if (token === undefined) return; + if (!token.trim()) { + vscode.window.showWarningMessage(t("A non-empty Aether device credential is required.")); + return; + } + await configuration.update("cloudRelayUrl", url.trim(), vscode.ConfigurationTarget.Global); + await context.secrets.store("codexRemoteCollab.cloudRelayToken", token.trim()); + vscode.window.showInformationMessage(t("Aether cloud connection saved. Restart the Codex Remote bridge to connect; local control remains available.")); + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.pairCloud", async () => { + const configuration = vscode.workspace.getConfiguration("codexRemoteCollab"); + const currentBaseUrl = configuration.get("aetherUrl", ""); + const baseUrl = await vscode.window.showInputBox({ + prompt: t("Aether server URL"), + value: currentBaseUrl, + placeHolder: "https://aether.example.com", + ignoreFocusOut: true, + validateInput: validateAetherBaseUrl, + }); + if (baseUrl === undefined || !baseUrl.trim()) return; + const code = await vscode.window.showInputBox({ + prompt: t("One-time pairing code shown in Aether"), + placeHolder: "ABCD-EFGH", + ignoreFocusOut: true, + validateInput: (value) => normalizePairingCode(value).length === 8 ? undefined : t("Enter the 8-character pairing code."), + }); + if (code === undefined || !code.trim()) return; + try { + const normalizedBaseUrl = baseUrl.trim().replace(/\/+$/, ""); + const response = await fetch(`${normalizedBaseUrl}/api/vscodex/pair`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ code: normalizePairingCode(code), name: hostname() || "VS Code" }), + }); + const raw = await response.text(); + let result: unknown; + try { + result = JSON.parse(raw); + } catch { + result = null; + } + if (!response.ok) { + const detail = isJsonRecord(result) && typeof result.error === "string" ? result.error : `HTTP ${response.status}`; + throw new Error(detail); + } + if (!isJsonRecord(result) || typeof result.device_token !== "string" || typeof result.ws_url !== "string") { + throw new Error(t("Aether returned an invalid pairing response.")); + } + const wsError = validateCloudRelayUrl(result.ws_url); + if (wsError) throw new Error(wsError); + await configuration.update("aetherUrl", normalizedBaseUrl, vscode.ConfigurationTarget.Global); + await configuration.update("cloudRelayUrl", result.ws_url, vscode.ConfigurationTarget.Global); + await context.secrets.store("codexRemoteCollab.cloudRelayToken", result.device_token); + if (activeHost) await stop(); + await start(false); + if (!activeHost) return; + if (activeRelay?.isConnected("aether-cloud")) { + vscode.window.showInformationMessage(t("Aether pairing completed. Local and cloud control are both active.")); + } else { + vscode.window.showWarningMessage(t("Aether pairing was saved, but the cloud connection is currently unavailable. Local control remains active and the cloud connection will retry.")); + } + } catch (error) { + vscode.window.showErrorMessage(t("Unable to pair with Aether: {0}", error instanceof Error ? error.message : String(error))); + } + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.sendInput", async () => { + if (!activeAdapter) { + vscode.window.showWarningMessage(t("Start the Codex remote bridge first.")); + return; + } + const text = await vscode.window.showInputBox({ prompt: t("Send input to the active Codex turn"), ignoreFocusOut: true }); + if (text === undefined || !text.trim()) return; + try { + await activeAdapter.sendInput(text); + } catch (error) { + vscode.window.showErrorMessage(t("Unable to send Codex input: {0}", error instanceof Error ? error.message : String(error))); + } + })); + context.subscriptions.push(vscode.commands.registerCommand("codexRemoteCollab.snapshot", async () => { + if (!activeAdapter) return vscode.window.showWarningMessage(t("Start the Codex remote bridge first.")); + const snapshot = await activeAdapter.snapshot(); + output.appendLine(JSON.stringify(snapshot)); + output.show(true); + })); + + if (vscode.workspace.getConfiguration("codexRemoteCollab").get("autoStart", true)) await start(true); +} + +export async function deactivate(): Promise { + if (autoStartRetryTimer) { + clearTimeout(autoStartRetryTimer); + autoStartRetryTimer = undefined; + } + const host = activeHost; + activeHost = undefined; + activeAdapter = undefined; + activeAdapterStatusSubscription?.dispose(); + activeAdapterStatusSubscription = undefined; + if (host) await host.stop(); + const localRelay = localRelayController; + localRelayController = undefined; + if (localRelay) await localRelay.stop(); +} + +function firstWorkspaceRoot(): string | undefined { + return vscode.workspace.workspaceFolders?.[0]?.uri.fsPath; +} + +function workspaceRoots(): string[] { + return (vscode.workspace.workspaceFolders ?? []).map((folder) => folder.uri.fsPath); +} + +function setStatus(text: string, tooltip: string, command?: string): void { + if (!statusItem) return; + statusItem.text = text; + statusItem.tooltip = tooltip; + statusItem.command = command; +} + +function setAttachStatus(ready: boolean): void { + setStatus( + ready ? "$(check) Codex Remote" : "$(sync~spin) Codex Remote", + ready ? t("Attached to the existing Codex conversation. Click to open the web control.") : t("Waiting for a Codex conversation to open in VS Code. It will connect automatically."), + "codexRemoteCollab.openWeb", + ); +} + +function setControlModeStatus(mode: ControlMode, ready: boolean): void { + if (mode === "sync") { + setAttachStatus(ready); + return; + } + setStatus( + ready ? "$(check) Codex Remote" : "$(sync~spin) Codex Remote", + ready + ? t("Independent Codex mode is connected. Click to open the web control.") + : t("Starting the independent Codex mode."), + "codexRemoteCollab.openWeb", + ); +} + +function scheduleAutoStartRetry(start: (automatic?: boolean) => Promise): void { + if (autoStartRetryTimer) return; + const delay = autoStartRetryMs; + autoStartRetryMs = Math.min(autoStartRetryMs * 2, 30_000); + autoStartRetryTimer = setTimeout(() => { + autoStartRetryTimer = undefined; + void start(true); + }, delay); +} + +function formatArgs(args: unknown[]): string { + return args.length ? args.map((arg) => (typeof arg === "string" ? arg : JSON.stringify(arg))).join(" ") : ""; +} + +function isAttachSessionUnavailable(message: string): boolean { + return message.includes("没有找到已打开的 VS Code Codex 会话") + || /找不到会话\s+.+\s+的 VS Code Codex owner/.test(message); +} + +function validateCloudRelayUrl(value: string): string | undefined { + if (!value.trim()) return undefined; + try { + const url = new URL(value.trim()); + if (url.protocol !== "wss:" && url.protocol !== "ws:") return t("Use a ws:// or wss:// URL."); + if (url.protocol === "ws:" && !isLoopbackHostname(url.hostname)) { + return t("Remote Aether connections must use wss://."); + } + return undefined; + } catch { + return t("Enter a valid WebSocket URL."); + } +} + +function resolveRelayConfiguration(configuration: vscode.WorkspaceConfiguration): { + localUrl: string; + cloudUrl: string; + legacyRemote: boolean; +} { + const defaultLocalUrl = "ws://127.0.0.1:8787/v1/connect"; + const explicitLocal = inspectedValue(configuration.inspect("localRelayUrl")); + const explicitCloud = inspectedValue(configuration.inspect("cloudRelayUrl")); + const explicitLegacy = inspectedValue(configuration.inspect("relayUrl")); + const legacyUrl = explicitLegacy?.trim() || ""; + const legacyRemote = Boolean(legacyUrl && !localRelayTarget(legacyUrl)); + const localUrl = (explicitLocal?.trim() + || (!legacyRemote ? legacyUrl : "") + || configuration.get("localRelayUrl", defaultLocalUrl).trim() + || defaultLocalUrl); + const cloudUrl = explicitCloud?.trim() + || (legacyRemote ? legacyUrl : "") + || configuration.get("cloudRelayUrl", "").trim(); + return { localUrl, cloudUrl, legacyRemote }; +} + +function inspectedValue(inspection: ReturnType | undefined): T | undefined { + if (!inspection) return undefined; + const values = inspection as { + globalLanguageValue?: T; + workspaceFolderLanguageValue?: T; + workspaceLanguageValue?: T; + workspaceFolderValue?: T; + workspaceValue?: T; + globalValue?: T; + }; + return values.workspaceFolderLanguageValue + ?? values.workspaceLanguageValue + ?? values.globalLanguageValue + ?? values.workspaceFolderValue + ?? values.workspaceValue + ?? values.globalValue; +} + +function resolveInitialControlMode(configuration: vscode.WorkspaceConfiguration): ControlMode { + const configured = inspectedValue(configuration.inspect("controlMode")); + if (configured === "sync" || configured === "async") return configured; + const legacyMode = inspectedValue<"attach" | "spawn">(configuration.inspect<"attach" | "spawn">("mode")); + return legacyMode === "spawn" ? "async" : "sync"; +} + +function validateAetherBaseUrl(value: string): string | undefined { + if (!value.trim()) return t("Enter the Aether server URL."); + try { + const url = new URL(value.trim()); + if (url.username || url.password || url.search || url.hash) return t("Use the Aether origin without credentials, a query, or a fragment."); + if (url.protocol === "https:") return undefined; + if (url.protocol === "http:" && isLoopbackHostname(url.hostname)) return undefined; + return t("Remote Aether servers must use https://."); + } catch { + return t("Enter a valid URL."); + } +} + +function normalizePairingCode(value: string): string { + return value.toUpperCase().replace(/[^A-Z2-9]/g, ""); +} + +function isLoopbackHostname(value: string): boolean { + const hostname = value.replace(/^\[|\]$/g, "").toLowerCase(); + return hostname === "127.0.0.1" || hostname === "localhost" || hostname === "::1"; +} + +function isJsonRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +/** + * Reuse the official extension's command registry for the header's new-chat + * action. This keeps the remote UI attached to the same VS Code Codex + * installation and avoids launching a second app-server process. + */ +async function openOfficialNewSession(logger: Logger): Promise { + const commands = await vscode.commands.getCommands(true); + const command = commands.includes("chatgpt.newCodexPanel") + ? "chatgpt.newCodexPanel" + : commands.includes("chatgpt.newChat") + ? "chatgpt.newChat" + : undefined; + if (!command) { + throw new Error(t("The official Codex extension new-conversation command was not found. Make sure the VS Code Codex extension is enabled.")); + } + await vscode.commands.executeCommand(command); + logger.info?.("Opened a new official Codex conversation with " + command); + return { opened: true, command }; +} diff --git a/aether-vscodex/vscode-extension/src/index.ts b/aether-vscodex/vscode-extension/src/index.ts new file mode 100644 index 000000000..4ec7c5d2a --- /dev/null +++ b/aether-vscodex/vscode-extension/src/index.ts @@ -0,0 +1,11 @@ +export * from "./protocol"; +export * from "./jsonlRpc"; +export * from "./codexAgentAdapter"; +export * from "./relayClient"; +export * from "./compositeRelay"; +export * from "./relayHost"; +export * from "./bridge"; +export * from "./codexPath"; +export * from "./codexIpc"; +export * from "./codexIpcAgentAdapter"; +export * from "./switchableAgentAdapter"; diff --git a/aether-vscodex/vscode-extension/src/jsonlRpc.ts b/aether-vscodex/vscode-extension/src/jsonlRpc.ts new file mode 100644 index 000000000..2d22d0f18 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/jsonlRpc.ts @@ -0,0 +1,268 @@ +import { ChildProcessWithoutNullStreams, spawn } from "node:child_process"; +import { createInterface, Interface as ReadLineInterface } from "node:readline"; + +import { + asJsonValue, + Disposable, + isRecord, + JsonRpcId, + JsonRpcNotification, + JsonRpcRequest, + JsonValue, + Logger, + jsonRpcIdKey, +} from "./protocol"; + +export interface JsonlRpcClientOptions { + command?: string; + args?: string[]; + cwd?: string; + env?: NodeJS.ProcessEnv; + logger?: Logger; + /** Optional request timeout. Zero disables it, which is useful for long turns. */ + requestTimeoutMs?: number; +} + +interface PendingRequest { + method: string; + resolve: (value: JsonValue) => void; + reject: (reason: Error) => void; + timer?: NodeJS.Timeout; +} + +export class JsonRpcRemoteError extends Error { + constructor( + message: string, + readonly code: number, + readonly data?: JsonValue, + ) { + super(message); + this.name = "JsonRpcRemoteError"; + } +} + +/** Minimal newline-delimited JSON-RPC client used by `codex app-server --stdio`. */ +export class JsonlRpcClient { + private readonly options: Required> & + Omit; + private child?: ChildProcessWithoutNullStreams; + private stdoutLines?: ReadLineInterface; + private nextId = 1; + private readonly pending = new Map(); + private readonly notificationListeners = new Set<(message: JsonRpcNotification) => void>(); + private readonly requestListeners = new Set<(message: JsonRpcRequest) => void>(); + private readonly exitListeners = new Set<(error?: Error) => void>(); + + constructor(options: JsonlRpcClientOptions = {}) { + this.options = { + command: options.command ?? "codex", + args: options.args ?? ["app-server", "--stdio"], + requestTimeoutMs: options.requestTimeoutMs ?? 0, + cwd: options.cwd, + env: options.env, + logger: options.logger, + }; + } + + get running(): boolean { + return Boolean(this.child && this.child.exitCode === null && !this.child.killed); + } + + async start(): Promise { + if (this.running) return; + + const child = spawn(this.options.command, this.options.args, { + cwd: this.options.cwd, + env: { ...process.env, ...(this.options.env ?? {}) }, + stdio: ["pipe", "pipe", "pipe"], + windowsHide: true, + }); + this.child = child; + + this.stdoutLines = createInterface({ input: child.stdout, crlfDelay: Infinity }); + this.stdoutLines.on("line", (line) => this.handleLine(line)); + child.stderr.on("data", (chunk: Buffer) => { + const text = redactDiagnostic(chunk.toString("utf8").trim()); + if (text) this.options.logger?.debug?.(`[app-server stderr] ${text}`); + }); + child.once("exit", (code, signal) => { + const expected = child.killed; + const error = expected + ? undefined + : new Error(`codex app-server exited (code=${String(code)}, signal=${String(signal)})`); + this.handleExit(error); + }); + + await new Promise((resolve, reject) => { + const onSpawn = (): void => { + child.off("error", onError); + resolve(); + }; + const onError = (error: Error): void => { + child.off("spawn", onSpawn); + const spawnError = error as NodeJS.ErrnoException; + if (spawnError.code === "ENOENT") { + reject(new Error(`Codex executable "${this.options.command}" was not found. Set codexRemoteCollab.codexCommand to its full path.`)); + return; + } + reject(error); + }; + child.once("spawn", onSpawn); + child.once("error", onError); + }); + } + + request(method: string, params?: JsonValue): Promise { + if (!this.running) return Promise.reject(new Error("app-server is not running")); + const id = this.nextId++; + + return new Promise((resolve, reject) => { + const pending: PendingRequest = { method, resolve, reject }; + if (this.options.requestTimeoutMs > 0) { + pending.timer = setTimeout(() => { + this.pending.delete(jsonRpcIdKey(id)); + reject(new Error(`app-server request timed out: ${method}`)); + }, this.options.requestTimeoutMs); + } + this.pending.set(jsonRpcIdKey(id), pending); + try { + this.write({ id, method, ...(params === undefined ? {} : { params }) }); + } catch (error) { + this.pending.delete(jsonRpcIdKey(id)); + if (pending.timer) clearTimeout(pending.timer); + reject(error instanceof Error ? error : new Error(String(error))); + } + }); + } + + notify(method: string, params?: JsonValue): void { + this.write({ method, ...(params === undefined ? {} : { params }) }); + } + + respond(id: JsonRpcId, result: JsonValue): void { + this.write({ id, result }); + } + + respondError(id: JsonRpcId, code: number, message: string, data?: JsonValue): void { + this.write({ id, error: { code, message, ...(data === undefined ? {} : { data }) } }); + } + + onNotification(listener: (message: JsonRpcNotification) => void): Disposable { + this.notificationListeners.add(listener); + return { dispose: () => this.notificationListeners.delete(listener) }; + } + + onServerRequest(listener: (message: JsonRpcRequest) => void): Disposable { + this.requestListeners.add(listener); + return { dispose: () => this.requestListeners.delete(listener) }; + } + + onExit(listener: (error?: Error) => void): Disposable { + this.exitListeners.add(listener); + return { dispose: () => this.exitListeners.delete(listener) }; + } + + close(): void { + const child = this.child; + this.child = undefined; + this.stdoutLines?.close(); + this.stdoutLines = undefined; + if (child && child.exitCode === null && !child.killed) child.kill(); + this.rejectAll(new Error("app-server client closed")); + } + + private write(message: unknown): void { + const child = this.child; + if (!child || child.exitCode !== null || child.killed || !child.stdin.writable) { + throw new Error("app-server is not running"); + } + child.stdin.write(`${JSON.stringify(message)}\n`, "utf8"); + } + + private handleLine(line: string): void { + const trimmed = line.trim(); + if (!trimmed) return; + + let message: unknown; + try { + message = JSON.parse(trimmed); + } catch (error) { + this.options.logger?.warn?.("Ignoring malformed app-server JSON", error, trimmed.slice(0, 500)); + return; + } + if (!isRecord(message)) return; + + const hasId = typeof message.id === "string" || typeof message.id === "number"; + const hasMethod = typeof message.method === "string"; + if (hasId && (Object.hasOwn(message, "result") || Object.hasOwn(message, "error")) && !hasMethod) { + this.handleResponse(message as Record & { id: JsonRpcId }); + return; + } + + if (hasMethod && hasId) { + const request: JsonRpcRequest = { + id: message.id as JsonRpcId, + method: message.method as string, + ...(message.params === undefined ? {} : { params: asJsonValue(message.params) }), + }; + for (const listener of this.requestListeners) listener(request); + return; + } + + if (hasMethod) { + const notification: JsonRpcNotification = { + method: message.method as string, + ...(message.params === undefined ? {} : { params: asJsonValue(message.params) }), + }; + for (const listener of this.notificationListeners) listener(notification); + return; + } + + this.options.logger?.warn?.("Ignoring unknown app-server message", message); + } + + private handleResponse(message: Record & { id: JsonRpcId }): void { + const pending = this.pending.get(jsonRpcIdKey(message.id)); + if (!pending) { + this.options.logger?.warn?.(`Received response for unknown app-server request ${String(message.id)}`); + return; + } + this.pending.delete(jsonRpcIdKey(message.id)); + if (pending.timer) clearTimeout(pending.timer); + + if (isRecord(message.error)) { + pending.reject( + new JsonRpcRemoteError( + typeof message.error.message === "string" ? message.error.message : `Request failed: ${pending.method}`, + typeof message.error.code === "number" ? message.error.code : -32000, + message.error.data === undefined ? undefined : asJsonValue(message.error.data), + ), + ); + return; + } + pending.resolve(message.result === undefined ? null : asJsonValue(message.result)); + } + + private handleExit(error?: Error): void { + this.child = undefined; + this.stdoutLines?.close(); + this.stdoutLines = undefined; + this.rejectAll(error ?? new Error("app-server exited")); + for (const listener of this.exitListeners) listener(error); + } + + private rejectAll(error: Error): void { + for (const request of this.pending.values()) { + if (request.timer) clearTimeout(request.timer); + request.reject(error); + } + this.pending.clear(); + } +} + +function redactDiagnostic(text: string): string { + return text + .replace(/Bearer\s+[A-Za-z0-9._~+\-/]+=*/gi, "Bearer [REDACTED]") + .replace(/\b(?:sk-[A-Za-z0-9_-]{12,}|gh[pousr]_[A-Za-z0-9_]{12,})\b/g, "[REDACTED]") + .replace(/((?:token|secret|password|api[_-]?key)\s*[:=]\s*)[^\s,;]+/gi, "$1[REDACTED]"); +} diff --git a/aether-vscodex/vscode-extension/src/localRelay.ts b/aether-vscodex/vscode-extension/src/localRelay.ts new file mode 100644 index 000000000..93f24bb2a --- /dev/null +++ b/aether-vscodex/vscode-extension/src/localRelay.ts @@ -0,0 +1,193 @@ +import * as http from "node:http"; +import * as path from "node:path"; + +import { Logger } from "./protocol"; + +interface BundledRelay { + start(): Promise<{ host: string; port: number }>; + stop(): Promise; +} + +interface BundledRelayModule { + CodexRelay: new (options: Record) => BundledRelay; +} + +export interface LocalRelayTarget { + host: string; + port: number; + healthUrl: string; + webUrl: string; +} + +export interface LocalRelayControllerOptions { + extensionPath: string; + logger?: Logger; + probeTimeoutMs?: number; + relayModulePath?: string; + loadRelayModule?: (modulePath: string) => BundledRelayModule; + probeRelayHealth?: (url: string, timeoutMs?: number) => Promise; +} + +/** + * Owns the loopback relay bundled with the companion extension. Remote and + * TLS relay URLs deliberately stay outside this controller. + */ +export class LocalRelayController { + private readonly options: LocalRelayControllerOptions; + private relay?: BundledRelay; + private target?: LocalRelayTarget; + private starting?: Promise; + private generation = 0; + + constructor(options: LocalRelayControllerOptions) { + this.options = options; + } + + async ensureRunning(relayUrl: string): Promise { + const target = localRelayTarget(relayUrl); + if (!target) return false; + if ((this.relay || this.starting) && this.target?.healthUrl !== target.healthUrl) await this.stop(); + const generation = this.generation; + this.target = target; + const available = await this.probeHealth(target.healthUrl); + // `stop()` may run while the health request is in flight. Do not let that + // completed probe resurrect a relay owned by a deactivated extension. + if (generation !== this.generation) return false; + if (available) return false; + if (this.relay) { + await this.relay.stop().catch(() => undefined); + this.relay = undefined; + } + if (this.starting) return this.starting; + this.starting = this.startBundledRelay(target).finally(() => { + this.starting = undefined; + }); + return this.starting; + } + + getWebUrl(relayUrl: string): string | undefined { + return localRelayTarget(relayUrl)?.webUrl; + } + + async stop(): Promise { + this.generation += 1; + const starting = this.starting; + if (starting) await starting.catch(() => undefined); + const relay = this.relay; + this.relay = undefined; + this.target = undefined; + if (relay) await relay.stop(); + } + + private async startBundledRelay(target: LocalRelayTarget): Promise { + const modulePath = this.options.relayModulePath + ?? path.join(this.options.extensionPath, "dist", "local-relay", "server.js"); + let relay: BundledRelay; + try { + const load = this.options.loadRelayModule ?? ((value: string) => require(value) as BundledRelayModule); + const module = load(modulePath); + if (typeof module?.CodexRelay !== "function") throw new Error("bundled relay module is invalid"); + relay = new module.CodexRelay({ + host: target.host, + port: target.port, + mode: "host", + spawnCodex: false, + authRequired: false, + }); + await relay.start(); + } catch (error) { + // Another VS Code window can win the listen race after our health + // probe. Treat that as success only when the expected relay responds. + if (await this.probeHealth(target.healthUrl)) { + this.options.logger?.info?.(`Using existing local relay at ${target.webUrl}`); + return false; + } + throw error; + } + this.relay = relay; + this.options.logger?.info?.(`Started bundled local relay at ${target.webUrl}`); + return true; + } + + private probeHealth(url: string): Promise { + const probe = this.options.probeRelayHealth ?? relayHealthAvailable; + return probe(url, this.options.probeTimeoutMs); + } +} + +export function localRelayTarget(relayUrl: string): LocalRelayTarget | undefined { + let url: URL; + try { + url = new URL(relayUrl); + } catch { + return undefined; + } + if (url.protocol !== "ws:" || !isLoopbackHostname(url.hostname)) return undefined; + const port = Number(url.port || 80); + if (!Number.isInteger(port) || port < 1 || port > 65_535) return undefined; + const hostname = normalizeLoopbackHostname(url.hostname); + const authorityHost = hostname.includes(":") ? `[${hostname}]` : hostname; + return { + host: hostname, + port, + healthUrl: `http://${authorityHost}:${port}/api/health`, + webUrl: `http://${authorityHost}:${port}/`, + }; +} + +function isLoopbackHostname(hostname: string): boolean { + const normalized = hostname.toLowerCase().replace(/^\[|\]$/g, ""); + return normalized === "localhost" || normalized === "127.0.0.1" || normalized === "::1"; +} + +function normalizeLoopbackHostname(hostname: string): string { + const normalized = hostname.toLowerCase().replace(/^\[|\]$/g, ""); + return normalized === "localhost" ? "127.0.0.1" : normalized; +} + +export function relayHealthAvailable(url: string, timeoutMs = 700): Promise { + return new Promise((resolve) => { + let settled = false; + let timer: NodeJS.Timeout | undefined; + const finish = (available: boolean): void => { + if (settled) return; + settled = true; + if (timer) clearTimeout(timer); + resolve(available); + }; + const request = http.get(url, (response) => { + if (response.statusCode !== 200) { + response.resume(); + finish(false); + return; + } + let body = ""; + response.setEncoding("utf8"); + response.on("data", (chunk) => { + if (body.length <= 16_384) body += chunk; + }); + response.on("end", () => { + try { + const payload = JSON.parse(body) as { ok?: unknown }; + finish(payload.ok === true); + } catch { + finish(false); + } + }); + response.on("aborted", () => finish(false)); + response.on("error", () => finish(false)); + response.on("close", () => { + if (!response.complete) finish(false); + }); + }); + request.setTimeout(timeoutMs, () => { + request.destroy(); + finish(false); + }); + request.on("error", () => finish(false)); + timer = setTimeout(() => { + request.destroy(); + finish(false); + }, timeoutMs); + }); +} diff --git a/aether-vscodex/vscode-extension/src/protocol.ts b/aether-vscodex/vscode-extension/src/protocol.ts new file mode 100644 index 000000000..19170a598 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/protocol.ts @@ -0,0 +1,496 @@ +/** + * Wire types shared by the relay host and the Codex app-server adapter. + * + * The relay intentionally treats `payload` as JSON. Keeping this boundary + * unopinionated lets the bridge continue working when app-server adds a new + * notification or request before this extension is updated. + */ + +export type JsonPrimitive = string | number | boolean | null; +export type JsonValue = JsonPrimitive | JsonValue[] | { [key: string]: JsonValue }; +export type JsonObject = { [key: string]: JsonValue }; +export type JsonRpcId = string | number; + +/** Preserve the JSON-RPC id type when using it as a map key. */ +export function jsonRpcIdKey(id: JsonRpcId): string { + return `${typeof id}:${String(id)}`; +} + +export function isJsonRpcId(value: unknown): value is JsonRpcId { + return typeof value === "string" || typeof value === "number"; +} + +export type ApprovalDecisionKind = "allow" | "deny" | "cancel"; + +const LEGACY_APPROVAL_METHODS = new Set(["applyPatchApproval", "execCommandApproval"]); +const V2_APPROVAL_METHODS = new Set([ + "item/commandExecution/requestApproval", + "item/fileChange/requestApproval", +]); + +/** + * Classify both current and legacy app-server approval decisions without + * rewriting the wire value. Unknown tagged objects intentionally return + * `undefined` so callers can fail closed instead of accidentally approving a + * newly introduced response shape. + */ +export function approvalDecisionKind(value: unknown): ApprovalDecisionKind | undefined { + if (typeof value === "string") { + if (new Set([ + "allow", + "accept", + "acceptForSession", + "approved", + "approved_for_session", + "approved_mcp_policy_amendment", + ]).has(value)) return "allow"; + if (new Set(["deny", "decline", "denied", "timed_out"]).has(value)) return "deny"; + if (new Set(["cancel", "abort"]).has(value)) return "cancel"; + return undefined; + } + if (!isRecord(value)) return undefined; + const keys = Object.keys(value); + if (keys.length !== 1) return undefined; + const key = keys[0]; + const nested = value[key]; + if (isExecpolicyAmendmentTag(key, nested) || isNetworkPolicyAmendmentTag(key, nested)) return "allow"; + if (key === "denied" && isRecord(nested) && typeof nested.rejection === "string") return "deny"; + return undefined; +} + +/** + * Classify a decision against the response schema for one app-server method. + * The generic classifier above is intentionally useful for relay envelopes; + * this method-aware variant prevents a v2 tagged object from being sent to a + * legacy callback (or vice versa), while retaining compatibility aliases that + * the relay may use in its outer `decision` field. + */ +export function approvalDecisionKindForMethod( + value: unknown, + method?: string, +): ApprovalDecisionKind | undefined { + const generic = approvalDecisionKind(value); + if (!generic || !method) return generic; + + if (LEGACY_APPROVAL_METHODS.has(method)) { + if (typeof value === "string") { + return new Set([ + "approved", + "approved_for_session", + "approved_mcp_policy_amendment", + "timed_out", + "abort", + ]).has(value) ? generic : undefined; + } + if (!isRecord(value)) return undefined; + const key = Object.keys(value)[0]; + return key === "approved_execpolicy_amendment" + || key === "network_policy_amendment" + || key === "denied" ? generic : undefined; + } + + if (V2_APPROVAL_METHODS.has(method)) { + if (typeof value === "string") { + return new Set(["accept", "acceptForSession", "decline", "cancel"]).has(value) + ? generic + : undefined; + } + if (!isRecord(value)) return undefined; + const key = Object.keys(value)[0]; + if (method === "item/fileChange/requestApproval") return undefined; + return key === "acceptWithExecpolicyAmendment" || key === "applyNetworkPolicyAmendment" + ? generic + : undefined; + } + + if (method === "mcpServer/elicitation/request") { + return typeof value === "string" && new Set(["accept", "decline", "cancel"]).has(value) + ? generic + : undefined; + } + + return generic; +} + +function isExecpolicyAmendmentTag(key: string, nested: unknown): boolean { + if (!isRecord(nested)) return false; + if (key === "acceptWithExecpolicyAmendment") { + return isStringArray(nested.execpolicy_amendment); + } + if (key === "approved_execpolicy_amendment") { + return isStringArray(nested.proposed_execpolicy_amendment); + } + return false; +} + +function isNetworkPolicyAmendmentTag(key: string, nested: unknown): boolean { + if (!isRecord(nested)) return false; + if (key === "applyNetworkPolicyAmendment") { + return isNetworkPolicyAmendment(nested.network_policy_amendment); + } + if (key === "network_policy_amendment") { + return isNetworkPolicyAmendment(nested.network_policy_amendment); + } + return false; +} + +function isStringArray(value: unknown): value is string[] { + return Array.isArray(value) && value.every((item) => typeof item === "string"); +} + +function isNetworkPolicyAmendment(value: unknown): boolean { + return isRecord(value) + && typeof value.host === "string" + && (value.action === "allow" || value.action === "deny"); +} + +/** Whether a response explicitly carries a decision/action field. */ +export function hasApprovalDecisionField(value: unknown): value is Record { + return isRecord(value) && (Object.prototype.hasOwnProperty.call(value, "decision") + || Object.prototype.hasOwnProperty.call(value, "action")); +} + +export interface Disposable { + dispose(): void; +} + +export interface JsonRpcRequest { + id: JsonRpcId; + method: string; + params?: JsonValue; +} + +export interface JsonRpcNotification { + method: string; + params?: JsonValue; +} + +export interface JsonRpcResponse { + id: JsonRpcId; + result?: JsonValue; + error?: { + code: number; + message: string; + data?: JsonValue; + }; +} + +export type JsonRpcMessage = JsonRpcRequest | JsonRpcNotification | JsonRpcResponse; + +export type RelayRole = "owner" | "operator" | "approver" | "viewer" | string; + +export interface RelayActor { + id?: string; + role?: RelayRole; +} + +/** A versioned relay event frame. `seq` is normally assigned by the relay. */ +export interface RelayEventFrame { + v: 1; + kind: "event"; + type: string; + id: string; + sessionId: string; + seq?: number; + ts: string; + actor?: RelayActor; + payload: JsonObject; + /** Optional typed execution projection attached by a VS Code host. */ + status?: AgentStatusSnapshot; +} + +export interface RelayCommandFrame { + v?: 1; + kind?: "command"; + type: string; + /** Compact relay compatibility form: `{ type: "command", method, params }`. */ + method?: string; + params?: JsonObject; + commandId?: string; + id?: string; + sessionId?: string; + actor?: RelayActor; + payload?: JsonObject; + /** Some clients put the command body under `command`. */ + command?: { + type?: string; + commandId?: string; + payload?: JsonObject; + [key: string]: JsonValue | undefined; + }; +} + +export interface RelayHelloFrame { + v: 1; + kind: "hello"; + clientType: "host" | "web" | string; + protocol?: number; + accessToken?: string; + token?: string; + lastSeq?: number; + sessionId?: string; + payload?: JsonObject; +} + +export interface RelayAckFrame { + v: 1; + kind: "ack"; + sessionId: string; + seq: number; +} + +export interface RelayErrorFrame { + v: 1; + kind: "error"; + code: string; + message: string; + retryable?: boolean; + commandId?: string; +} + +export type RelayFrame = + | RelayEventFrame + | RelayCommandFrame + | RelayHelloFrame + | RelayAckFrame + | RelayErrorFrame + | (JsonObject & { kind?: string; v?: number }); + +/** + * Live execution information projected from the official Codex conversation + * state. The private IPC protocol can add new turn statuses/flags, so the + * string fields intentionally remain open-ended for forward compatibility. + */ +export interface AgentStatusSnapshot { + /** Coarse UI activity, for example `thinking`, `editing`, or `running`. */ + activity: string; + /** Raw/normalized turn status (`inProgress`, `completed`, ...). */ + turnStatus: string; + /** Runtime flags such as `waitingOnApproval` or `waitingOnUserInput`. */ + activeFlags: string[]; + startedAtMs?: number | null; + durationMs?: number | null; + /** + * Time spent doing work in the official UI. This deliberately differs + * from `durationMs`: Codex starts the worked-for clock at the first work + * item and stops it when the final assistant response starts. + */ + workedDurationMs?: number | null; + /** Elapsed wall-clock time for an active turn. */ + elapsedMs?: number | null; + firstTurnWorkItemStartedAtMs?: number | null; + finalAssistantStartedAtMs?: number | null; + error?: JsonValue; +} + +/** Official background-agent lifecycle values emitted by Codex v2 items. */ +export type CollabAgentStatus = + | "pendingInit" + | "running" + | "interrupted" + | "completed" + | "errored" + | "shutdown" + | "notFound" + | string; + +export type CollabAgentTool = + | "spawnAgent" + | "sendInput" + | "resumeAgent" + | "wait" + | "closeAgent" + | string; + +export type CollabAgentToolCallStatus = "inProgress" | "completed" | "failed" | string; +export type SubAgentActivityKind = "started" | "interacted" | "interrupted" | "completed" | string; + +/** Last known state for one receiver in a collabAgentToolCall item. */ +export interface CollabAgentStateSnapshot { + status: CollabAgentStatus; + message?: string | null; +} + +/** + * Browser-safe projection of a background Codex subagent. The official + * webview currently uses the four coarse statuses below; the string union is + * deliberately open so a newer app-server status does not break the relay. + */ +export interface SubagentSnapshot { + threadId: string; + displayName: string | null; + prompt: string | null; + /** Alias used by the subagent side panel for the same prompt text. */ + objective?: string | null; + status: "waiting" | "working" | "done" | "failed" | string; + statusMessage: string | null; + startedAtMs?: number | null; + completedAtMs?: number | null; + canInteract?: boolean; + model?: string | null; + agentPath?: string | null; + parentThreadId?: string | null; +} + +export interface AgentEvent { + /** Normalized relay event name, for example `output.chunk`. */ + type: string; + threadId?: string; + turnId?: string; + requestId?: JsonRpcId; + payload: JsonObject; + /** Original app-server notification/request, when available. */ + raw?: JsonValue; + /** Optional typed projection of live Codex turn/runtime status. */ + status?: AgentStatusSnapshot; +} + +export interface PendingApproval { + requestId: JsonRpcId; + method: string; + threadId?: string; + turnId?: string; + itemId?: string; + action: string; + risk: "low" | "medium" | "high" | "unknown"; + summary: string; + /** SHA-256 of canonicalized, unredacted app-server request params. */ + commandHash?: string; + createdAt: number; + expiresAt?: number; + payload: JsonObject; +} + +export interface SessionSnapshot { + threadId: string | null; + turnId: string | null; + state: string; + pendingApprovals: PendingApproval[]; + pendingRequests?: Array<{ + requestId: JsonRpcId; + method: string; + params?: JsonValue; + commandHash?: string; + risk?: string; + summary?: string; + createdAt?: number; + expiresAt?: number; + }>; + outputTail: string; + /** Optional role-aware projection used by the browser renderer. */ + messages?: JsonValue[]; + /** Background/inline subagents reconstructed from official collab items. */ + subagents?: SubagentSnapshot[]; + /** Live execution projection; retained alongside the legacy `state` field. */ + status?: AgentStatusSnapshot; + /** Convenience aliases for clients that do not consume `status` yet. */ + activity?: string; + turnStatus?: string; + activeFlags?: string[]; + startedAtMs?: number | null; + durationMs?: number | null; + workedDurationMs?: number | null; + elapsedMs?: number | null; + metadata?: JsonObject; +} + +/** A live VS Code Codex conversation that the attach bridge has verified. */ +export interface SessionListEntry { + threadId: string; + title: string; + updatedAtMs: number | null; + cwd?: string | null; + active: boolean; + /** True for attach-mode results; retained for wire compatibility. */ + available: boolean; +} + +export interface SessionListResult { + sessions: SessionListEntry[]; + activeThreadId: string | null; +} + +/** Which owner controls conversation navigation for the remote surface. */ +export type ControlMode = "sync" | "async"; + +export interface AgentAdapter { + start(): Promise; + /** Switch between following VS Code and independently owned conversations. */ + setControlMode?(params: JsonObject): Promise; + /** Return the currently committed control mode without taking a snapshot. */ + getControlMode?(): ControlMode; + /** Start a new app-server thread. */ + startThread?(params?: JsonObject): Promise; + /** Ask the official VS Code Codex extension to open a fresh conversation. */ + newSession?(params?: JsonObject): Promise; + /** Start a turn; `threadId` may be supplied in params or use the active thread. */ + startTurn?(params: JsonObject): Promise; + /** Steer the active turn. */ + steerTurn?(params: JsonObject): Promise; + /** Persist model/effort and other owner-managed settings on the thread. */ + updateThreadSettings?(params: JsonObject): Promise; + /** List verified, attachable local conversations without starting another Codex process. */ + listSessions?(params?: JsonObject): Promise; + /** Attach the follower to another already-open conversation. */ + selectSession?(params: JsonObject): Promise; + /** Interrupt a turn. */ + interruptTurn?(params: JsonObject): Promise; + /** Convenience MVP aliases. */ + sendInput(text: string, params?: JsonObject): Promise; + cancel(taskId?: string, params?: JsonObject): Promise; + respondApproval( + requestId: JsonRpcId, + decision: "allow" | "deny" | "cancel", + reason?: string, + response?: JsonValue, + ): Promise; + /** Resolve all pending approvals/inputs with a deny response. */ + denyPending?(reason?: string): Promise; + snapshot(): Promise; + onEvent(listener: (event: AgentEvent) => void): Disposable; + dispose(): Promise; +} + +export interface RelayTransport { + connect(): Promise; + send(frame: RelayFrame): void; + onMessage(listener: (frame: RelayFrame) => void): Disposable; + onOpen?(listener: () => void): Disposable; + onClose?(listener: (error?: Error) => void): Disposable; + close(): void; +} + +export interface Logger { + debug?(message: string, ...args: unknown[]): void; + info?(message: string, ...args: unknown[]): void; + warn?(message: string, ...args: unknown[]): void; + error?(message: string, ...args: unknown[]): void; +} + +export const noopDisposable = (): Disposable => ({ dispose: () => undefined }); + +export function isRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +export function asJsonObject(value: unknown): JsonObject { + return isRecord(value) ? (value as JsonObject) : {}; +} + +export function asJsonValue(value: unknown): JsonValue { + if (value === undefined) return null; + if (value === null || typeof value === "string" || typeof value === "number" || typeof value === "boolean") { + return value; + } + if (Array.isArray(value)) { + return value.map(asJsonValue); + } + if (isRecord(value)) { + const output: JsonObject = {}; + for (const [key, item] of Object.entries(value)) { + if (item !== undefined) output[key] = asJsonValue(item); + } + return output; + } + return String(value); +} diff --git a/aether-vscodex/vscode-extension/src/relayClient.ts b/aether-vscodex/vscode-extension/src/relayClient.ts new file mode 100644 index 000000000..e1e0f9c60 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/relayClient.ts @@ -0,0 +1,413 @@ +import { createInterface, Interface as ReadLineInterface } from "node:readline"; +import WebSocket from "ws"; + +import { + Disposable, + isRecord, + JsonObject, + Logger, + RelayFrame, + RelayHelloFrame, + RelayTransport, +} from "./protocol"; + +export interface RelayClientOptions { + url: string; + accessToken?: string; + sessionId?: string; + lastSeq?: number; + reconnect?: boolean; + reconnectInitialMs?: number; + reconnectMaxMs?: number; + maxFrameBytes?: number; + maxQueuedBytes?: number; + logger?: Logger; + /** Injectable constructor for tests or a browser-compatible WebSocket. */ + webSocket?: new (url: string) => unknown; +} + +type SocketLike = { + readyState?: number; + send(data: string): void; + close(): void; + on?(event: string, listener: (...args: any[]) => void): void; + addEventListener?(event: string, listener: (...args: any[]) => void): void; +}; + +const OPEN = 1; +// A structured Codex history snapshot is routinely larger than 256 KiB even +// though its plain-text tail is capped. Keep a bounded limit, but leave enough +// room for the message/tool projection of a long attached conversation. +export const DEFAULT_MAX_RELAY_FRAME_BYTES = 16 * 1024 * 1024; +export const DEFAULT_MAX_RELAY_QUEUE_BYTES = DEFAULT_MAX_RELAY_FRAME_BYTES + 2 * 1024 * 1024; + +interface QueuedFrame { + serialized: string; + bytes: number; + projectionKey?: string; +} + +const QUEUED_PROJECTION_TYPES = new Set(["session.snapshot", "output.snapshot", "output.chunk"]); + +function queuedProjectionKey(frame: RelayFrame): string | undefined { + if (!isRecord(frame) || frame.kind !== "event" || typeof frame.type !== "string" + || !QUEUED_PROJECTION_TYPES.has(frame.type)) return undefined; + const sessionId = typeof frame.sessionId === "string" ? frame.sessionId : "default"; + return `${sessionId}:transcript`; +} + +/** WebSocket relay transport with bounded reconnect and frame validation. */ +export class RelayClient implements RelayTransport { + readonly handlesHandshake = true; + private readonly options: Required< + Pick + > & + Omit; + private socket?: SocketLike; + // A WebSocket can report OPEN while its relay authentication handshake is + // still in flight. Keep this separate from `socket` so events emitted by + // the adapter during reconnect are queued until the relay sends auth.ok. + private authenticatedSocket?: SocketLike; + private connecting?: Promise; + // Incremented whenever a connection attempt is replaced or explicitly + // closed. Late events from an older WebSocket must not mutate newer state. + private connectionGeneration = 0; + private reconnectTimer?: NodeJS.Timeout; + private stopped = false; + private retryMs: number; + private readonly queue: QueuedFrame[] = []; + private queueBytes = 0; + private readonly listeners = new Set<(frame: RelayFrame) => void>(); + private readonly openListeners = new Set<() => void>(); + private readonly closeListeners = new Set<(error?: Error) => void>(); + + constructor(options: RelayClientOptions) { + const maxFrameBytes = options.maxFrameBytes ?? DEFAULT_MAX_RELAY_FRAME_BYTES; + const defaultMaxQueuedBytes = Math.max( + maxFrameBytes, + Math.min(DEFAULT_MAX_RELAY_QUEUE_BYTES, maxFrameBytes * 2), + ); + this.options = { + ...options, + reconnect: options.reconnect ?? true, + reconnectInitialMs: options.reconnectInitialMs ?? 500, + reconnectMaxMs: options.reconnectMaxMs ?? 10_000, + maxFrameBytes, + maxQueuedBytes: Math.max(1, Math.floor(options.maxQueuedBytes ?? defaultMaxQueuedBytes)), + }; + this.retryMs = this.options.reconnectInitialMs; + } + + /** Let RelayHost assign its stable session id before the first hello. */ + setSessionId(sessionId: string): void { + this.options.sessionId = sessionId; + } + + async connect(): Promise { + this.stopped = false; + if (this.socket?.readyState === OPEN && this.authenticatedSocket === this.socket) return; + if (this.connecting) return this.connecting; + + const generation = ++this.connectionGeneration; + let connectionPromise: Promise; + connectionPromise = new Promise((resolve, reject) => { + let settled = false; + let authenticated = false; + const SocketCtor = this.options.webSocket ?? WebSocket; + let socket: SocketLike; + try { + socket = new SocketCtor(this.options.url) as SocketLike; + } catch (error) { + reject(error instanceof Error ? error : new Error(String(error))); + return; + } + this.socket = socket; + this.authenticatedSocket = undefined; + + const isCurrent = (): boolean => this.connectionGeneration === generation && this.socket === socket; + + const onOpen = (): void => { + if (!isCurrent() || settled || authenticated) return; + try { + // The TCP/WebSocket open event is only a transport milestone. Do + // not release queued commands until the relay has authenticated us. + this.sendHello(socket); + } catch (error) { + if (!settled) { + settled = true; + reject(error instanceof Error ? error : new Error(String(error))); + } + } + }; + const onMessage = (raw: unknown): void => { + if (!isCurrent()) return; + const data = extractMessageData(raw); + if (Buffer.byteLength(data, "utf8") > this.options.maxFrameBytes) { + this.options.logger?.warn?.("Ignoring oversized relay frame"); + return; + } + let frame: unknown; + try { + frame = JSON.parse(data); + } catch { + this.options.logger?.warn?.("Ignoring malformed relay JSON"); + return; + } + if (!isRecord(frame)) return; + if (frame.type === "auth.ok" && !authenticated && !settled) { + authenticated = true; + settled = true; + this.authenticatedSocket = socket; + this.retryMs = this.options.reconnectInitialMs; + try { + this.flush(socket); + } catch (error) { + this.options.logger?.warn?.("Unable to flush relay queue after authentication", error); + } + for (const listener of this.openListeners) listener(); + resolve(); + } else if (frame.type === "error" && !authenticated && !settled) { + settled = true; + reject(new Error(typeof frame.message === "string" ? frame.message : "relay authentication failed")); + } + if (frame.kind === "event" && typeof frame.seq === "number") { + this.options.lastSeq = Math.max(this.options.lastSeq ?? 0, frame.seq); + } + for (const listener of this.listeners) listener(frame as RelayFrame); + }; + const onError = (raw: unknown): void => { + if (!isCurrent()) return; + const error = raw instanceof Error ? raw : new Error("relay websocket error"); + this.options.logger?.warn?.(error.message); + if (!settled) { + settled = true; + reject(error); + } + }; + const onClose = (): void => { + const current = isCurrent(); + if (current) { + this.socket = undefined; + if (this.authenticatedSocket === socket) this.authenticatedSocket = undefined; + } + const error = new Error("relay websocket closed"); + // A stale socket may still need to settle the promise returned to its + // caller, but it must never notify the active host or schedule a + // second reconnect loop. + if (!current) { + if (!settled) { + settled = true; + reject(error); + } + return; + } + for (const listener of this.closeListeners) listener(error); + if (!settled) { + settled = true; + reject(error); + } + if (!this.stopped && this.options.reconnect) this.scheduleReconnect(); + }; + + bindSocket(socket, onOpen, onMessage, onError, onClose); + // A small number of test/browser WebSocket implementations can already + // be OPEN by the time listeners are attached. + if (socket.readyState === OPEN) queueMicrotask(onOpen); + }).finally(() => { + if (this.connectionGeneration === generation && this.connecting === connectionPromise) { + this.connecting = undefined; + } + }); + this.connecting = connectionPromise; + return connectionPromise; + } + + send(frame: RelayFrame): void { + const serialized = JSON.stringify(frame); + const bytes = Buffer.byteLength(serialized, "utf8"); + if (bytes > this.options.maxFrameBytes) { + throw new Error(`relay frame exceeds ${this.options.maxFrameBytes} bytes`); + } + if (this.socket?.readyState === OPEN && this.authenticatedSocket === this.socket) { + this.socket.send(serialized); + return; + } + this.enqueue({ serialized, bytes, projectionKey: queuedProjectionKey(frame) }); + } + + onMessage(listener: (frame: RelayFrame) => void): Disposable { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + onOpen(listener: () => void): Disposable { + this.openListeners.add(listener); + return { dispose: () => this.openListeners.delete(listener) }; + } + + onClose(listener: (error?: Error) => void): Disposable { + this.closeListeners.add(listener); + return { dispose: () => this.closeListeners.delete(listener) }; + } + + close(): void { + this.stopped = true; + this.connectionGeneration += 1; + if (this.reconnectTimer) clearTimeout(this.reconnectTimer); + this.reconnectTimer = undefined; + this.connecting = undefined; + const socket = this.socket; + this.socket = undefined; + this.authenticatedSocket = undefined; + if (socket && socket.readyState !== 3) socket.close(); + this.queue.length = 0; + this.queueBytes = 0; + } + + private sendHello(socket: SocketLike): void { + const hello: RelayHelloFrame = { + v: 1, + kind: "hello", + clientType: "host", + protocol: 1, + ...(this.options.sessionId ? { sessionId: this.options.sessionId } : {}), + ...(this.options.lastSeq !== undefined ? { lastSeq: this.options.lastSeq } : {}), + }; + socket.send(JSON.stringify(hello)); + if (this.options.accessToken) { + // Keep authentication separate from hello so a relay can challenge the + // host before accepting a bearer token (and so hello remains cacheable). + socket.send(JSON.stringify({ v: 1, kind: "auth", accessToken: this.options.accessToken })); + } + } + + private flush(socket: SocketLike): void { + if (socket.readyState !== OPEN || this.authenticatedSocket !== socket || this.socket !== socket) return; + while (this.queue.length > 0) { + const entry = this.queue.shift() as QueuedFrame; + this.queueBytes = Math.max(0, this.queueBytes - entry.bytes); + socket.send(entry.serialized); + } + } + + private enqueue(entry: QueuedFrame): void { + // Transcript events are reconstructible: RelayHost publishes a fresh full + // session snapshot after every authenticated reconnect. Keep only the + // newest projection per session while preserving approval/command events. + if (entry.projectionKey) { + for (let index = this.queue.length - 1; index >= 0; index -= 1) { + if (this.queue[index].projectionKey === entry.projectionKey) this.removeQueuedFrame(index); + } + } + if (entry.bytes > this.options.maxQueuedBytes) { + this.options.logger?.warn?.("Dropping relay frame that exceeds the reconnect queue byte limit"); + return; + } + while (this.queue.length >= 100 || this.queueBytes + entry.bytes > this.options.maxQueuedBytes) { + const projectionIndex = this.queue.findIndex((queued) => Boolean(queued.projectionKey)); + if (projectionIndex >= 0) { + this.removeQueuedFrame(projectionIndex); + continue; + } + // Never evict an approval/command solely to retain a transcript delta; + // the authoritative snapshot emitted after auth restores that state. + if (entry.projectionKey) { + this.options.logger?.debug?.("Dropping supersedable relay projection while reconnect queue is full"); + return; + } + this.removeQueuedFrame(0); + } + this.queue.push(entry); + this.queueBytes += entry.bytes; + } + + private removeQueuedFrame(index: number): void { + const [removed] = this.queue.splice(index, 1); + if (removed) this.queueBytes = Math.max(0, this.queueBytes - removed.bytes); + } + + private scheduleReconnect(): void { + if (this.reconnectTimer || this.stopped) return; + const delay = this.retryMs; + this.retryMs = Math.min(this.options.reconnectMaxMs, Math.max(this.retryMs * 2, this.options.reconnectInitialMs)); + this.reconnectTimer = setTimeout(() => { + this.reconnectTimer = undefined; + void this.connect().catch((error) => this.options.logger?.debug?.("relay reconnect failed", error)); + }, delay); + } +} + +/** + * Line-oriented transport for local development and CI. Pipe it to a relay + * process with `node dist/cli.js`; each line is one JSON relay frame. + */ +export class StdioRelayTransport implements RelayTransport { + private readonly listeners = new Set<(frame: RelayFrame) => void>(); + private readonly lineReader: ReadLineInterface; + private closed = false; + + constructor( + private readonly input: NodeJS.ReadableStream = process.stdin, + private readonly output: NodeJS.WritableStream = process.stdout, + private readonly logger?: Logger, + ) { + this.lineReader = createInterface({ input, crlfDelay: Infinity }); + this.lineReader.on("line", (line) => { + if (!line.trim()) return; + try { + const frame = JSON.parse(line); + if (isRecord(frame)) for (const listener of this.listeners) listener(frame as RelayFrame); + } catch (error) { + this.logger?.warn?.("Ignoring malformed relay stdin frame", error); + } + }); + } + + async connect(): Promise { + this.closed = false; + } + + send(frame: RelayFrame): void { + if (this.closed) throw new Error("stdio relay transport is closed"); + this.output.write(`${JSON.stringify(frame)}\n`); + } + + onMessage(listener: (frame: RelayFrame) => void): Disposable { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + close(): void { + this.closed = true; + this.lineReader.close(); + } +} + +function bindSocket( + socket: SocketLike, + onOpen: () => void, + onMessage: (data: unknown) => void, + onError: (error: unknown) => void, + onClose: () => void, +): void { + if (typeof socket.on === "function") { + socket.on("open", onOpen); + socket.on("message", onMessage); + socket.on("error", onError); + socket.on("close", onClose); + } else if (typeof socket.addEventListener === "function") { + socket.addEventListener("open", onOpen); + socket.addEventListener("message", onMessage); + socket.addEventListener("error", onError); + socket.addEventListener("close", onClose); + } else { + onError(new Error("WebSocket implementation has no event API")); + } +} + +function extractMessageData(raw: unknown): string { + if (typeof raw === "string") return raw; + if (Buffer.isBuffer(raw)) return raw.toString("utf8"); + if (isRecord(raw) && "data" in raw) return extractMessageData(raw.data); + return String(raw ?? ""); +} diff --git a/aether-vscodex/vscode-extension/src/relayHost.ts b/aether-vscodex/vscode-extension/src/relayHost.ts new file mode 100644 index 000000000..d9ff21df3 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/relayHost.ts @@ -0,0 +1,650 @@ +import { randomUUID } from "node:crypto"; + +import { + AgentAdapter, + AgentEvent, + approvalDecisionKind, + approvalDecisionKindForMethod, + asJsonObject, + asJsonValue, + Disposable, + hasApprovalDecisionField, + isRecord, + JsonObject, + JsonRpcId, + Logger, + JsonValue, + RelayActor, + RelayCommandFrame, + RelayEventFrame, + RelayFrame, + RelayTransport, +} from "./protocol"; + +export interface RelayHostOptions { + adapter: AgentAdapter; + relay: RelayTransport; + sessionId?: string; + actor?: RelayActor; + /** Capabilities enforced locally even when relay authorization is bypassed. */ + capabilities?: Iterable; + logger?: Logger; + /** Emit a handshake on transports that do not implement one themselves. */ + sendHandshake?: boolean; +} + +/** + * Maps relay commands to the app-server AgentAdapter and publishes normalized + * adapter events. This is the policy boundary for the VS Code host. + */ +export class RelayHost { + private readonly adapter: AgentAdapter; + private readonly relay: RelayTransport; + private readonly options: RelayHostOptions; + private readonly subscriptions: Disposable[] = []; + private readonly commandResults = new Map(); + private readonly inFlightCommands = new Set(); + private readonly capabilities: Set; + private eventSeq = 0; + private sessionId: string; + private started = false; + private adapterReady = false; + + constructor(options: RelayHostOptions); + constructor(adapter: AgentAdapter, relay: RelayTransport, options?: Omit); + constructor( + optionsOrAdapter: RelayHostOptions | AgentAdapter, + relayArg?: RelayTransport, + legacyOptions: Omit = {}, + ) { + if (isAgentAdapter(optionsOrAdapter)) { + this.adapter = optionsOrAdapter; + if (!relayArg) throw new Error("RelayHost requires a relay transport"); + this.relay = relayArg; + this.options = { ...legacyOptions, adapter: this.adapter, relay: this.relay }; + } else { + this.options = optionsOrAdapter; + this.adapter = optionsOrAdapter.adapter; + this.relay = optionsOrAdapter.relay; + } + this.capabilities = new Set(this.options.capabilities ?? [ + "read_output", + "send_task_input", + "cancel_task", + "approve_low_risk", + ]); + this.sessionId = this.options.sessionId ?? `sess_${randomUUID()}`; + } + + get id(): string { + return this.sessionId; + } + + async start(): Promise { + if (this.started) return; + this.started = true; + this.subscriptions.push(this.adapter.onEvent((event) => { + if (event.type === "connection.opened") this.adapterReady = true; + if (event.type === "connection.closed") this.adapterReady = false; + this.publishAgentEvent(event); + })); + this.subscriptions.push(this.relay.onMessage((frame) => { + void this.handleFrame(frame).catch((error) => { + this.options.logger?.warn?.("Invalid relay frame", error); + if (isRecord(frame) && typeof frame.commandId === "string") { + this.sendCommandResult(frame.commandId, false, undefined, error instanceof Error ? error.message : String(error), typeof frame.method === "string" ? frame.method : typeof frame.type === "string" ? frame.type : undefined); + } + }); + })); + if (this.relay.onClose) this.subscriptions.push(this.relay.onClose((error) => { + // A relay disconnect must not leave an app-server request waiting for a + // browser that can no longer answer. The adapter's local deny path is + // deliberately fail-closed. Do not publish `connection.closed` here: + // that event describes the app-server process, while this callback only + // describes the outbound transport and is followed by connection.opened + // on a successful reconnect. + void this.adapter.denyPending?.("relay disconnected"); + this.options.logger?.debug?.("Relay transport closed", error?.message ?? ""); + })); + if (this.relay.onOpen) this.subscriptions.push(this.relay.onOpen(() => { + // RelayClient fires onOpen only after auth.ok. On reconnect the adapter + // is already initialized, so the synthetic event restores relay state; + // during initial startup the adapter event below is authoritative. + if (this.adapterReady) { + this.publishConnectionEvent("connection.opened"); + void this.publishSnapshot(); + } + })); + + const configurableRelay = this.relay as RelayTransport & { setSessionId?: (sessionId: string) => void }; + configurableRelay.setSessionId?.(this.sessionId); + try { + await this.relay.connect(); + if (this.options.sendHandshake !== false && !transportHandlesHandshake(this.relay)) { + this.safeSend({ v: 1, kind: "hello", clientType: "host", protocol: 1, sessionId: this.sessionId }); + } + // Start app-server only after the relay handshake is queued/sent. This + // keeps standalone stdout frames protocol-ordered and prevents an early + // notification from racing the host hello. + await this.adapter.start(); + if (!this.adapterReady) { + this.adapterReady = true; + this.publishConnectionEvent("connection.opened"); + } + await this.publishSnapshot(); + } catch (error) { + this.started = false; + this.adapterReady = false; + for (const subscription of this.subscriptions.splice(0)) subscription.dispose(); + this.relay.close(); + await this.adapter.dispose().catch(() => undefined); + throw error; + } + } + + async stop(): Promise { + if (!this.started) return; + this.started = false; + this.adapterReady = false; + this.inFlightCommands.clear(); + for (const subscription of this.subscriptions.splice(0)) subscription.dispose(); + this.relay.close(); + await this.adapter.dispose(); + } + + /** Public for unit tests and local stdin bridges. */ + async handleFrame(frame: RelayFrame): Promise { + if (!isRecord(frame)) return; + if (frame.kind === "command" || isCommandLike(frame)) { + await this.handleCommand(frame as unknown as RelayCommandFrame); + return; + } + if (frame.kind === "event" && typeof frame.seq === "number") { + this.safeSend({ v: 1, kind: "ack", sessionId: frame.sessionId, seq: frame.seq }); + } + } + + private async handleCommand(frame: RelayCommandFrame): Promise { + const command = normalizeCommand(frame); + const commandId = command.commandId; + if (commandId) { + const previous = this.commandResults.get(commandId); + if (previous) { + this.safeSend(previous); + return; + } + if (this.inFlightCommands.has(commandId)) { + this.safeSend({ + v: 1, + kind: "event", + // Do not call this `command.accepted`: the relay treats that event + // as the terminal result for its pending command. A retry while the + // original operation is running is only an informational event. + type: "command.pending", + id: `evt_${randomUUID()}`, + sessionId: this.sessionId, + seq: ++this.eventSeq, + ts: new Date().toISOString(), + actor: this.options.actor ?? { id: "host", role: "host" }, + payload: { commandId, duplicate: true, pending: true }, + }); + return; + } + this.inFlightCommands.add(commandId); + } + + const role = frame.actor?.role ?? "operator"; + const denied = authorize(command.type, role, this.capabilities); + if (denied) { + // A viewer may not force a pending approval to deny (that would turn a + // read-only role into a denial-of-service primitive). Authorized roles + // can still be rejected by local capability/policy checks, in which + // case denying the app-server request is the safe terminal action. + if (role === "owner" || role === "operator" || role === "approver" || role === "host") { + await this.denyApprovalIfNeeded(command.type, command.payload, denied); + } + this.sendCommandResult(commandId, false, undefined, denied, command.type); + if (commandId) this.inFlightCommands.delete(commandId); + return; + } + + try { + const result = await this.executeCommand(command.type, command.payload); + this.sendCommandResult(commandId, true, result, undefined, command.type); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + this.options.logger?.warn?.(`Relay command ${command.type} failed`, error); + await this.denyApprovalIfNeeded(command.type, command.payload, message); + this.sendCommandResult(commandId, false, undefined, message, command.type); + } finally { + if (commandId) this.inFlightCommands.delete(commandId); + } + } + + private async denyApprovalIfNeeded(type: string, payload: JsonObject, reason: string): Promise { + const command = canonicalCommandType(type); + if (command !== "approval.respond" && command !== "input.respond" && command !== "server.request.respond") return; + const requestId = payload.requestId; + if (requestId === undefined || (typeof requestId !== "string" && typeof requestId !== "number")) return; + try { + await this.adapter.respondApproval(requestId, "deny", reason); + } catch { + // The request may already have expired or been resolved. Keep the + // original command rejection as the observable result. + } + } + + private async executeCommand(type: string, payload: JsonObject): Promise { + switch (canonicalCommandType(type)) { + case "control.mode.get": { + const mode = this.adapter.getControlMode?.(); + if (mode) return { mode }; + const snapshot = await this.adapter.snapshot(); + const controlMode = snapshot.metadata?.controlMode; + if (controlMode !== "sync" && controlMode !== "async") { + throw new Error("adapter does not expose a control mode"); + } + return { mode: controlMode }; + } + case "control.mode.set": + if (!this.adapter.setControlMode) throw new Error("adapter does not support control mode switching"); + return this.adapter.setControlMode(payload); + case "thread.start": + if (!this.adapter.startThread) throw new Error("adapter does not support thread/start"); + return this.adapter.startThread(payload); + case "session.new": + if (!this.adapter.newSession) throw new Error("adapter does not support session/new"); + return this.adapter.newSession(payload); + case "thread.settings.update": + if (!this.adapter.updateThreadSettings) throw new Error("adapter does not support thread/settings/update"); + return this.adapter.updateThreadSettings(payload); + case "session.list": + if (!this.adapter.listSessions) throw new Error("adapter does not support session/list"); + return this.adapter.listSessions(payload); + case "session.select": + if (!this.adapter.selectSession) throw new Error("adapter does not support session/select"); + return this.adapter.selectSession(payload); + case "turn.start": + if (this.adapter.startTurn) return this.adapter.startTurn(payload); + return this.adapter.sendInput(extractCommandText(payload), payload); + case "turn.steer": + if (this.adapter.steerTurn) return this.adapter.steerTurn(payload); + return this.adapter.sendInput(extractCommandText(payload), payload); + case "turn.interrupt": + if (this.adapter.interruptTurn) return this.adapter.interruptTurn(payload); + return this.adapter.cancel(typeof payload.turnId === "string" ? payload.turnId : undefined, payload); + case "task.input": { + const text = typeof payload.text === "string" ? payload.text : typeof payload.message === "string" ? payload.message : undefined; + if (!text) throw new Error("task.input requires payload.text"); + return this.adapter.sendInput(text, payload); + } + case "task.cancel": { + const taskId = typeof payload.taskId === "string" ? payload.taskId : typeof payload.turnId === "string" ? payload.turnId : undefined; + return this.adapter.cancel(taskId, payload); + } + case "approval.respond": { + const requestId = payload.requestId; + if (typeof requestId !== "string" && typeof requestId !== "number") throw new Error("approval.respond requires requestId"); + const requestedValue = payload.decision; + const decision = approvalDecisionKind(requestedValue); + if (!decision) throw new Error("decision must be a recognized allow, deny, or cancel value"); + const snapshot = await this.adapter.snapshot(); + // JSON-RPC distinguishes numeric and string ids. Keep the lookup + // type-safe so id `1` cannot accidentally authorize response `"1"`. + const approval = snapshot.pendingApprovals.find((item) => item.requestId === requestId); + const method = approval?.method ?? (typeof payload.method === "string" ? payload.method : undefined); + const response = payload.response ?? implicitApprovalResponse(requestedValue, decision, method, payload.scope); + if (decision === "allow") { + if (approval && typeof payload.commandHash === "string" && payload.commandHash !== approval.commandHash) { + throw new Error("approval commandHash does not match the pending request"); + } + if (approval?.risk === "high" && !this.capabilities.has("approve_high_risk") && !this.capabilities.has("*")) { + throw new Error("host policy requires approve_high_risk for this approval"); + } + } + validateApprovalResponse(decision, response, method); + return this.adapter.respondApproval( + requestId, + decision, + typeof payload.reason === "string" ? payload.reason : undefined, + response, + ); + } + case "input.respond": + case "server.request.respond": { + const requestId = payload.requestId; + if (typeof requestId !== "string" && typeof requestId !== "number") throw new Error(`${type} requires requestId`); + const response = payload.response ?? (payload.answers !== undefined ? payload.answers : undefined); + // Tool-input requests do not carry an allow/deny field in their wire + // response, while MCP elicitation uses `action`. Prefer an explicit + // response action when present; otherwise honor the relay decision and + // fail closed when a denial has no custom response. + if (response !== undefined && !isRecord(response)) { + throw new Error("input response must be a JSON object"); + } + const responseDecision = isRecord(response) + ? explicitResponseDecision(response) + : undefined; + const requestedDecision = payload.decision === undefined + ? undefined + : approvalDecisionKind(payload.decision); + if (payload.decision !== undefined && !requestedDecision) { + throw new Error("decision must be allow, deny, or cancel"); + } + if (responseDecision && requestedDecision + && responseDecision !== requestedDecision + // RelayHost uses `decision: "allow"` as a generic envelope for + // MCP/input responses; the nested action remains authoritative in + // that one compatibility case. + && requestedDecision !== "allow") { + throw new Error(`input response implies ${responseDecision}, but decision is ${requestedDecision}`); + } + // For MCP, `response.action` is the actual app-server decision and is + // authoritative even if a relay uses `decision: "allow"` as a generic + // input-response envelope. With no custom response, an explicit relay + // decision (or the fail-closed deny default) controls the result. + const decision = responseDecision ?? requestedDecision ?? (response === undefined ? "deny" : "allow"); + const responseForAdapter = requestedDecision && requestedDecision !== "allow" && !responseDecision + ? undefined + : response; + return this.adapter.respondApproval( + requestId, + decision, + typeof payload.reason === "string" ? payload.reason : undefined, + responseForAdapter, + ); + } + case "session.snapshot": + case "snapshot": + return this.adapter.snapshot(); + case "ping": + return { pong: true, ts: new Date().toISOString() }; + default: + throw new Error(`unsupported relay command: ${type}`); + } + } + + private sendCommandResult(commandId: string | undefined, ok: boolean, result?: unknown, error?: string, method?: string): void { + const frame: RelayEventFrame = { + v: 1, + kind: "event", + type: ok ? "command.accepted" : "command.rejected", + id: `evt_${randomUUID()}`, + sessionId: this.sessionId, + seq: ++this.eventSeq, + ts: new Date().toISOString(), + actor: this.options.actor ?? { id: "host", role: "host" }, + payload: { + ...(commandId ? { commandId } : {}), + ...(method ? { method } : {}), + ok, + ...(ok ? { result: asJsonValue(result) } : { error: error ?? "command rejected" }), + }, + }; + if (commandId) { + this.commandResults.set(commandId, frame); + if (this.commandResults.size > 1000) this.commandResults.delete(this.commandResults.keys().next().value as string); + } + this.safeSend(frame); + } + + private publishAgentEvent(event: AgentEvent): void { + if (event.threadId && this.sessionId.startsWith("sess_")) { + // Keep a stable relay session id while exposing the app-server thread id + // in the payload; a relay session may contain more than one thread. + } + const payload: JsonObject = { + ...event.payload, + ...(event.status ? { executionStatus: asJsonValue(event.status) } : {}), + ...(event.threadId ? { threadId: event.threadId } : {}), + ...(event.turnId ? { turnId: event.turnId } : {}), + ...(event.requestId !== undefined ? { requestId: asJsonValue(event.requestId) } : {}), + ...(event.raw !== undefined ? { raw: event.raw } : {}), + }; + const frame: RelayEventFrame = { + v: 1, + kind: "event", + type: event.type, + id: `evt_${randomUUID()}`, + sessionId: this.sessionId, + seq: ++this.eventSeq, + ts: new Date().toISOString(), + actor: this.options.actor ?? { id: "host", role: "host" }, + payload, + ...(event.status ? { status: { ...event.status, activeFlags: [...event.status.activeFlags] } } : {}), + }; + try { + this.safeSend(frame); + } catch (error) { + this.options.logger?.warn?.("Unable to publish relay event", error); + } + } + + private safeSend(frame: RelayFrame): void { + try { + this.relay.send(frame); + } catch (error) { + this.options.logger?.warn?.("Unable to send relay frame", error); + } + } + + private publishConnectionEvent(type: string, error?: Error): void { + if (!this.started && type === "connection.closed") return; + this.publishAgentEvent({ type, payload: error ? { message: error.message } : {} }); + } + + private async publishSnapshot(): Promise { + try { + const snapshot = await this.adapter.snapshot(); + this.publishAgentEvent({ + type: "session.snapshot", + threadId: snapshot.threadId ?? undefined, + turnId: snapshot.turnId ?? undefined, + payload: asJsonObject(snapshot), + status: snapshot.status, + }); + } catch (error) { + this.options.logger?.warn?.("Unable to publish adapter snapshot", error); + } + } +} + +interface NormalizedCommand { + type: string; + commandId?: string; + payload: JsonObject; +} + +function normalizeCommand(frame: RelayCommandFrame): NormalizedCommand { + const nested = isRecord(frame.command) ? frame.command : undefined; + const type = typeof nested?.type === "string" + ? nested.type + : typeof frame.method === "string" + ? frame.method + : frame.type === "command" + ? "" + : frame.type; + const commandId = typeof frame.commandId === "string" + ? frame.commandId + : typeof nested?.commandId === "string" + ? nested.commandId + : typeof frame.id === "string" + ? frame.id + : undefined; + if (!type) throw new Error("relay command has no type"); + + if (isRecord(nested?.payload)) return { type, commandId, payload: asJsonObject(nested.payload) }; + if (isRecord(frame.payload)) return { type, commandId, payload: asJsonObject(frame.payload) }; + if (isRecord(frame.params)) return { type, commandId, payload: asJsonObject(frame.params) }; + + const payload: JsonObject = {}; + for (const [key, value] of Object.entries(frame)) { + if (["v", "kind", "type", "method", "params", "commandId", "id", "sessionId", "actor", "command"].includes(key)) continue; + if (value !== undefined) payload[key] = asJsonValue(value); + } + return { type, commandId, payload }; +} + +function canonicalCommandType(type: string): string { + const normalized = type.trim().replace(/\//g, ".").replace(/\s+/g, ".").toLowerCase(); + if (normalized === "control.mode.get" || normalized === "controlmode.get" || normalized === "controlmodeget" || normalized === "mode.get" || normalized === "modeget") return "control.mode.get"; + if (normalized === "control.mode.set" || normalized === "controlmode.set" || normalized === "controlmodeset" || normalized === "mode.set" || normalized === "modeset") return "control.mode.set"; + if (normalized === "thread.start" || normalized === "threadstart") return "thread.start"; + if (normalized === "session.new" || normalized === "sessionnew" || normalized === "thread.new" || normalized === "threadnew") return "session.new"; + if (normalized === "thread.settings.update" || normalized === "threadsettings.update" || normalized === "threadsettingsupdate") return "thread.settings.update"; + if (normalized === "session.list" || normalized === "thread.list" || normalized === "sessionlist" || normalized === "threadlist") return "session.list"; + if (normalized === "session.select" || normalized === "session.switch" || normalized === "thread.select" || normalized === "thread.attach" || normalized === "sessionswitch" || normalized === "threadselect") return "session.select"; + if (normalized === "turn.start" || normalized === "turnstart") return "turn.start"; + if (normalized === "turn.steer" || normalized === "turnsteer") return "turn.steer"; + if (normalized === "turn.interrupt" || normalized === "turninterrupt") return "turn.interrupt"; + if (normalized === "approval.respond" || normalized === "approvalrespond") return "approval.respond"; + if (normalized === "task.input" || normalized === "taskinput") return "task.input"; + if (normalized === "task.cancel" || normalized === "taskcancel") return "task.cancel"; + if (normalized === "input.respond" || normalized === "inputrespond") return "input.respond"; + if (normalized === "server.request.respond" || normalized === "serverrequest.respond") return "server.request.respond"; + if (normalized === "session.snapshot" || normalized === "snapshot") return "session.snapshot"; + return normalized; +} + +function authorize(type: string, role: string, capabilities: Set): string | undefined { + const command = canonicalCommandType(type); + const readOnly = command === "session.snapshot" || command === "snapshot" || command === "session.list" || command === "control.mode.get" || command === "ping"; + if (readOnly) return undefined; + if (role === "viewer") return "viewer role cannot issue control commands"; + if (role !== "owner" && role !== "operator" && role !== "approver" && role !== "host") return `role ${role} is not authorized`; + if ((command === "approval.respond" || command === "input.respond" || command === "server.request.respond") && role !== "owner" && role !== "operator" && role !== "approver" && role !== "host") { + return "role is not authorized to resolve approvals"; + } + const required = command === "approval.respond" ? "approve_low_risk" : command === "task.cancel" || command === "turn.interrupt" ? "cancel_task" : command === "task.input" || command.startsWith("turn.") || command === "thread.start" || command === "thread.settings.update" || command === "session.select" || command === "session.new" || command === "control.mode.set" ? "send_task_input" : undefined; + if (required && !capabilities.has(required) && !capabilities.has("*") && role !== "owner" && role !== "host") return `missing capability: ${required}`; + return undefined; +} + +function isCommandLike(frame: Record): boolean { + if (frame.kind === "command") return true; + if (frame.type === "command" && typeof frame.method === "string") return true; + if (frame.kind !== undefined) return false; + if (typeof frame.method === "string") return true; + if (typeof frame.commandId !== "string") return false; + return KNOWN_COMMAND_TYPES.has(String(frame.type).trim().replace(/\//g, ".").toLowerCase()); +} + +const KNOWN_COMMAND_TYPES = new Set([ + "control.mode.get", + "control.mode.set", + "thread.start", + "session.new", + "thread.settings.update", + "session.list", + "session.select", + "turn.start", + "turn.steer", + "turn.interrupt", + "approval.respond", + "task.input", + "task.cancel", + "input.respond", + "server.request.respond", + "session.snapshot", + "snapshot", + "ping", +]); + +function isAgentAdapter(value: unknown): value is AgentAdapter { + return isRecord(value) && typeof value.start === "function" && typeof value.onEvent === "function" && typeof value.sendInput === "function" && typeof value.cancel === "function" && typeof value.respondApproval === "function" && typeof value.snapshot === "function"; +} + +function extractCommandText(payload: JsonObject): string { + if (typeof payload.text === "string") return payload.text; + if (typeof payload.message === "string") return payload.message; + if (typeof payload.prompt === "string") return payload.prompt; + if (Array.isArray(payload.input)) { + const first = payload.input[0]; + if (isRecord(first) && typeof first.text === "string") return first.text; + } + throw new Error("turn command requires text or input"); +} + +function validateApprovalResponse( + decision: "allow" | "deny" | "cancel", + response: JsonValue | undefined, + method?: string, +): void { + if (response === undefined) return; + if (!isRecord(response)) throw new Error("approval response must be a JSON object"); + + const hasDecision = Object.prototype.hasOwnProperty.call(response, "decision"); + const hasAction = Object.prototype.hasOwnProperty.call(response, "action"); + if (hasDecision || hasAction) { + const decisionKind = hasDecision ? approvalDecisionKindForMethod(response.decision, method) : undefined; + const actionKind = hasAction ? approvalDecisionKindForMethod(response.action, method) : undefined; + if (hasDecision && !decisionKind) throw new Error("unsupported approval response decision"); + if (hasAction && !actionKind) throw new Error("unsupported approval response action"); + if (decisionKind && actionKind && decisionKind !== actionKind) { + throw new Error("approval response decision and action conflict"); + } + const implied = decisionKind ?? actionKind; + if (implied && implied !== decision) { + throw new Error(`approval response implies ${implied}, but decision is ${decision}`); + } + return; + } + + // Permissions approvals intentionally carry a profile rather than a + // decision field. Keep the profile shape narrow; malformed/unknown objects + // must not be interpreted as an approval. + if (method === "item/permissions/requestApproval" + && isRecord(response.permissions) + && (response.scope === "turn" || response.scope === "session") + && (response.strictAutoReview === undefined || typeof response.strictAutoReview === "boolean") + && Object.keys(response).every((key) => key === "permissions" || key === "scope" || key === "strictAutoReview")) { + return; + } + throw new Error("approval response has no recognized decision or permission profile"); +} + +/** + * Convert a relay's compact outer decision into a wire response only when it + * carries a non-canonical app-server value. Canonical `allow`/`deny`/`cancel` + * remain undefined so the adapter can choose the method-specific default. + */ +function implicitApprovalResponse( + requestedValue: JsonValue, + decision: "allow" | "deny" | "cancel", + method?: string, + scope?: JsonValue, +): JsonValue | undefined { + // Permission approvals have a profile response, not a decision wrapper. + // Let the adapter construct the requested turn-scoped profile by default; + // callers that need session scope must provide the full profile explicitly. + if (method === "item/permissions/requestApproval") return undefined; + if (requestedValue === "allow" || requestedValue === "deny" || requestedValue === "cancel") { + if (requestedValue === "allow" && scope === "session") { + if (method === "applyPatchApproval" || method === "execCommandApproval") return { decision: "approved_for_session" }; + if (method === "item/commandExecution/requestApproval" || method === "item/fileChange/requestApproval") return { decision: "acceptForSession" }; + } + return undefined; + } + // The generic classifier has already rejected unknown/conflicting values. + // Preserve recognized legacy/v2 tags exactly under the app-server wrapper. + if (decision === "allow" || decision === "deny" || decision === "cancel") { + return { decision: requestedValue }; + } + return undefined; +} + +function explicitResponseDecision(response: Record): "allow" | "deny" | "cancel" | undefined { + if (!hasApprovalDecisionField(response)) return undefined; + const hasDecision = Object.prototype.hasOwnProperty.call(response, "decision"); + const hasAction = Object.prototype.hasOwnProperty.call(response, "action"); + const decision = hasDecision ? approvalDecisionKind(response.decision) : undefined; + const action = hasAction ? approvalDecisionKind(response.action) : undefined; + if (hasDecision && !decision) throw new Error("unsupported input response decision"); + if (hasAction && !action) throw new Error("unsupported input response action"); + if (decision && action && decision !== action) throw new Error("input response decision and action conflict"); + return decision ?? action; +} + +function transportHandlesHandshake(transport: RelayTransport): boolean { + return Boolean((transport as RelayTransport & { handlesHandshake?: boolean }).handlesHandshake); +} diff --git a/aether-vscodex/vscode-extension/src/switchableAgentAdapter.ts b/aether-vscodex/vscode-extension/src/switchableAgentAdapter.ts new file mode 100644 index 000000000..a78f9bd22 --- /dev/null +++ b/aether-vscodex/vscode-extension/src/switchableAgentAdapter.ts @@ -0,0 +1,425 @@ +import { + AgentAdapter, + AgentEvent, + asJsonObject, + ControlMode, + Disposable, + JsonObject, + JsonRpcId, + JsonValue, + Logger, + SessionSnapshot, +} from "./protocol"; + +export type AgentAdapterFactory = (mode: ControlMode) => AgentAdapter | Promise; + +export interface SwitchableAgentAdapterOptions { + initialMode: ControlMode; + createAdapter: AgentAdapterFactory; + /** Persist the committed mode. Persistence errors do not roll back a live adapter. */ + onModeChanged?: (mode: ControlMode, previousMode: ControlMode) => void | Promise; + logger?: Logger; +} + +interface AdapterBinding { + adapter: AgentAdapter; + mode: ControlMode; + generation: number; + committed: boolean; + bufferedEvents: AgentEvent[]; + subscription: Disposable; +} + +interface ModeCapabilities extends JsonObject { + followsVscodeRoute: boolean; + sessionList: boolean; + sessionSelect: boolean; + sessionCreate: boolean; + threadSettings: boolean; +} + +/** + * Keeps RelayHost bound to one stable AgentAdapter while atomically replacing + * the implementation behind it when the control owner changes. + */ +export class SwitchableAgentAdapter implements AgentAdapter { + private readonly options: SwitchableAgentAdapterOptions; + private readonly listeners = new Set<(event: AgentEvent) => void>(); + private binding: AdapterBinding | null = null; + private controlMode: ControlMode; + private modeEpoch = 0; + private startPromise: Promise | null = null; + private switchPromise: Promise | null = null; + private started = false; + private disposed = false; + + constructor(options: SwitchableAgentAdapterOptions) { + this.options = options; + this.controlMode = validateControlMode(options.initialMode); + } + + async start(): Promise { + if (this.started) return; + if (this.disposed) throw new Error("switchable adapter has been disposed"); + if (this.startPromise) return this.startPromise; + + const operation = this.startInitialAdapter(); + this.startPromise = operation; + try { + await operation; + } finally { + if (this.startPromise === operation) this.startPromise = null; + } + } + + getControlMode(): ControlMode { + return this.controlMode; + } + + async setControlMode(params: JsonObject): Promise { + const nextMode = controlModeFromParams(params); + this.ensureStarted(); + if (this.switchPromise) throw new Error("a control mode switch is already in progress"); + if (nextMode === this.controlMode) { + return { + changed: false, + controlMode: this.controlMode, + previousControlMode: this.controlMode, + modeEpoch: this.modeEpoch, + }; + } + + const operation = this.performModeSwitch(nextMode); + this.switchPromise = operation; + try { + return await operation; + } finally { + if (this.switchPromise === operation) this.switchPromise = null; + } + } + + async startThread(params: JsonObject = {}): Promise { + this.assertIndependentNavigation("thread/start"); + const adapter = this.activeAdapterForMutation(); + if (!adapter.startThread) throw unsupported("thread/start", this.controlMode); + return adapter.startThread(params); + } + + async newSession(params: JsonObject = {}): Promise { + this.assertIndependentNavigation("session/new"); + const adapter = this.activeAdapterForMutation(); + if (adapter.newSession) return adapter.newSession(params); + if (adapter.startThread) return adapter.startThread(params); + throw unsupported("session/new", this.controlMode); + } + + async startTurn(params: JsonObject): Promise { + const adapter = this.activeAdapterForMutation(); + if (!adapter.startTurn) throw unsupported("turn/start", this.controlMode); + return adapter.startTurn(params); + } + + async steerTurn(params: JsonObject): Promise { + const adapter = this.activeAdapterForMutation(); + if (!adapter.steerTurn) throw unsupported("turn/steer", this.controlMode); + return adapter.steerTurn(params); + } + + async updateThreadSettings(params: JsonObject): Promise { + const adapter = this.activeAdapterForMutation(); + if (!adapter.updateThreadSettings) throw unsupported("thread/settings/update", this.controlMode); + return adapter.updateThreadSettings(params); + } + + async listSessions(params: JsonObject = {}): Promise { + this.assertIndependentNavigation("session/list"); + const adapter = this.activeAdapter(); + if (!adapter.listSessions) throw unsupported("session/list", this.controlMode); + return adapter.listSessions(params); + } + + async selectSession(params: JsonObject): Promise { + this.assertIndependentNavigation("session/select"); + const adapter = this.activeAdapterForMutation(); + if (!adapter.selectSession) throw unsupported("session/select", this.controlMode); + return adapter.selectSession(params); + } + + async interruptTurn(params: JsonObject): Promise { + const adapter = this.activeAdapter(); + if (!adapter.interruptTurn) throw unsupported("turn/interrupt", this.controlMode); + return adapter.interruptTurn(params); + } + + async sendInput(text: string, params: JsonObject = {}): Promise { + return this.activeAdapterForMutation().sendInput(text, params); + } + + async cancel(taskId?: string, params: JsonObject = {}): Promise { + return this.activeAdapter().cancel(taskId, params); + } + + async respondApproval( + requestId: JsonRpcId, + decision: "allow" | "deny" | "cancel", + reason?: string, + response?: JsonValue, + ): Promise { + return this.activeAdapter().respondApproval(requestId, decision, reason, response); + } + + async denyPending(reason?: string): Promise { + await this.activeAdapter().denyPending?.(reason); + } + + async snapshot(): Promise { + const binding = this.activeBinding(); + const snapshot = await binding.adapter.snapshot(); + return this.decorateSnapshot(snapshot, binding); + } + + onEvent(listener: (event: AgentEvent) => void): Disposable { + this.listeners.add(listener); + return { dispose: () => this.listeners.delete(listener) }; + } + + async dispose(): Promise { + if (this.disposed) return; + this.disposed = true; + + const starting = this.startPromise; + const switching = this.switchPromise; + await starting?.catch(() => undefined); + await switching?.catch(() => undefined); + + const binding = this.binding; + this.binding = null; + this.started = false; + if (!binding) return; + binding.committed = false; + binding.subscription.dispose(); + await binding.adapter.dispose(); + } + + private async startInitialAdapter(): Promise { + const binding = await this.createBinding(this.controlMode, this.modeEpoch); + try { + await binding.adapter.start(); + if (this.disposed) throw new Error("switchable adapter was disposed while starting"); + binding.committed = true; + this.binding = binding; + this.started = true; + this.flushBufferedEvents(binding); + } catch (error) { + await this.releaseBinding(binding); + throw error; + } + } + + private async performModeSwitch(nextMode: ControlMode): Promise { + const previousBinding = this.activeBinding(); + const previousMode = this.controlMode; + this.assertSnapshotIdle(await previousBinding.adapter.snapshot()); + + const nextEpoch = this.modeEpoch + 1; + const candidate = await this.createBinding(nextMode, nextEpoch); + if (candidate.adapter === previousBinding.adapter) { + candidate.subscription.dispose(); + throw new Error("adapter factory must return a distinct adapter when switching control modes"); + } + try { + await candidate.adapter.start(); + if (this.disposed) throw new Error("switchable adapter was disposed while switching modes"); + + // VS Code can start a turn independently while the candidate boots. + // Recheck immediately before the synchronous commit point. + this.assertSnapshotIdle(await previousBinding.adapter.snapshot()); + const candidateSnapshot = await candidate.adapter.snapshot(); + // The candidate snapshot is an await point, so make the old adapter's + // liveness check the final operation before committing synchronously. + this.assertSnapshotIdle(await previousBinding.adapter.snapshot()); + + candidate.committed = true; + this.binding = candidate; + this.controlMode = nextMode; + this.modeEpoch = nextEpoch; + previousBinding.committed = false; + + const result: JsonObject = { + changed: true, + controlMode: nextMode, + previousControlMode: previousMode, + modeEpoch: nextEpoch, + }; + this.emit({ type: "control.mode.changed", payload: result }); + this.flushBufferedEvents(candidate); + const snapshot = this.decorateSnapshot(candidateSnapshot, candidate); + this.emit({ + type: "session.snapshot", + threadId: snapshot.threadId ?? undefined, + turnId: snapshot.turnId ?? undefined, + payload: asJsonObject(snapshot), + status: snapshot.status, + }); + + previousBinding.subscription.dispose(); + await previousBinding.adapter.dispose().catch((error) => { + this.options.logger?.warn?.("Unable to dispose the previous control mode adapter", error); + }); + await Promise.resolve(this.options.onModeChanged?.(nextMode, previousMode)).catch((error) => { + this.options.logger?.warn?.("Unable to persist the committed control mode", error); + }); + return result; + } catch (error) { + if (this.binding !== candidate) await this.releaseBinding(candidate); + throw error; + } + } + + private async createBinding(mode: ControlMode, generation: number): Promise { + const adapter = await this.options.createAdapter(mode); + if (!adapter) throw new Error(`adapter factory returned no adapter for ${mode} mode`); + const binding: AdapterBinding = { + adapter, + mode, + generation, + committed: false, + bufferedEvents: [], + subscription: { dispose: () => undefined }, + }; + binding.subscription = adapter.onEvent((event) => this.receiveAdapterEvent(binding, event)); + return binding; + } + + private receiveAdapterEvent(binding: AdapterBinding, event: AgentEvent): void { + if (!binding.committed) { + binding.bufferedEvents.push(event); + return; + } + if (this.binding !== binding || binding.generation !== this.modeEpoch) return; + this.emit(this.decorateEvent(event, binding)); + } + + private flushBufferedEvents(binding: AdapterBinding): void { + const events = binding.bufferedEvents.splice(0); + for (const event of events) { + if (this.binding !== binding || binding.generation !== this.modeEpoch) return; + this.emit(this.decorateEvent(event, binding)); + } + } + + private emit(event: AgentEvent): void { + for (const listener of this.listeners) { + try { + listener(event); + } catch (error) { + this.options.logger?.warn?.("Switchable adapter event listener failed", error); + } + } + } + + private activeBinding(): AdapterBinding { + this.ensureStarted(); + if (!this.binding) throw new Error("switchable adapter has no active adapter"); + return this.binding; + } + + private activeAdapter(): AgentAdapter { + return this.activeBinding().adapter; + } + + private activeAdapterForMutation(): AgentAdapter { + if (this.switchPromise) throw new Error("control mode is switching; retry after it completes"); + return this.activeAdapter(); + } + + private ensureStarted(): void { + if (this.disposed) throw new Error("switchable adapter has been disposed"); + if (!this.started || !this.binding) throw new Error("switchable adapter is not started"); + } + + private assertIndependentNavigation(operation: string): void { + if (this.controlMode === "sync") { + throw new Error(`${operation} is unavailable in sync mode; conversation navigation follows VS Code`); + } + } + + private assertSnapshotIdle(snapshot: SessionSnapshot): void { + const pendingApprovalCount = snapshot.pendingApprovals.length; + const pendingRequestCount = snapshot.pendingRequests?.length ?? 0; + const state = normalizeStatus(snapshot.state); + const turnStatus = normalizeStatus(snapshot.status?.turnStatus ?? snapshot.turnStatus ?? ""); + const activeFlags = snapshot.status?.activeFlags ?? snapshot.activeFlags ?? []; + const hasActiveState = ACTIVE_STATUSES.has(state) || ACTIVE_STATUSES.has(turnStatus) || activeFlags.length > 0; + if (snapshot.turnId || hasActiveState || pendingApprovalCount > 0 || pendingRequestCount > 0) { + throw new Error("cannot switch control mode while a turn or request is active"); + } + } + + private decorateSnapshot(snapshot: SessionSnapshot, binding: AdapterBinding): SessionSnapshot { + return { + ...snapshot, + metadata: { + ...(snapshot.metadata ?? {}), + mode: binding.mode, + controlMode: binding.mode, + modeEpoch: binding.generation, + capabilities: this.capabilities(binding), + }, + }; + } + + private decorateEvent(event: AgentEvent, binding: AdapterBinding): AgentEvent { + if (event.type !== "session.snapshot") return event; + return { + ...event, + payload: { + ...event.payload, + metadata: { + ...asJsonObject(event.payload.metadata), + mode: binding.mode, + controlMode: binding.mode, + modeEpoch: binding.generation, + capabilities: this.capabilities(binding), + }, + }, + }; + } + + private capabilities(binding: AdapterBinding): ModeCapabilities { + const independent = binding.mode === "async"; + return { + followsVscodeRoute: !independent, + sessionList: independent && typeof binding.adapter.listSessions === "function", + sessionSelect: independent && typeof binding.adapter.selectSession === "function", + sessionCreate: independent && (typeof binding.adapter.newSession === "function" + || typeof binding.adapter.startThread === "function"), + threadSettings: typeof binding.adapter.updateThreadSettings === "function", + }; + } + + private async releaseBinding(binding: AdapterBinding): Promise { + binding.committed = false; + binding.subscription.dispose(); + await binding.adapter.dispose().catch(() => undefined); + } +} + +function controlModeFromParams(params: JsonObject): ControlMode { + return validateControlMode(params.mode ?? params.controlMode); +} + +function validateControlMode(value: unknown): ControlMode { + if (value === "sync" || value === "async") return value; + throw new Error("control mode must be sync or async"); +} + +function unsupported(operation: string, mode: ControlMode): Error { + return new Error(`${operation} is not supported by the ${mode} adapter`); +} + +const ACTIVE_STATUSES = new Set(["active", "inprogress", "running", "starting", "thinking", "editing", "working"]); + +function normalizeStatus(value: string): string { + return value.trim().replace(/[\s_-]+/g, "").toLowerCase(); +} diff --git a/aether-vscodex/vscode-extension/tsconfig.json b/aether-vscodex/vscode-extension/tsconfig.json new file mode 100644 index 000000000..b58a0919f --- /dev/null +++ b/aether-vscodex/vscode-extension/tsconfig.json @@ -0,0 +1,28 @@ +{ + "compilerOptions": { + "target": "ES2022", + "module": "commonjs", + "lib": [ + "ES2022" + ], + "rootDir": "src", + "outDir": "dist", + "strict": true, + "esModuleInterop": true, + "forceConsistentCasingInFileNames": true, + "moduleResolution": "node", + "sourceMap": true, + "skipLibCheck": true, + "types": [ + "node", + "vscode" + ] + }, + "include": [ + "src/**/*.ts" + ], + "exclude": [ + "node_modules", + "dist" + ] +} diff --git a/aether-vscodex/web/.gitignore b/aether-vscodex/web/.gitignore new file mode 100644 index 000000000..f4e2c6d6b --- /dev/null +++ b/aether-vscodex/web/.gitignore @@ -0,0 +1,3 @@ +node_modules/ +dist/ +*.tsbuildinfo diff --git a/aether-vscodex/web/index.html b/aether-vscodex/web/index.html new file mode 100644 index 000000000..8a0d2f16d --- /dev/null +++ b/aether-vscodex/web/index.html @@ -0,0 +1,13 @@ + + + + + + + Codex + + +
+ + + diff --git a/aether-vscodex/web/package-lock.json b/aether-vscodex/web/package-lock.json new file mode 100644 index 000000000..deb51611f --- /dev/null +++ b/aether-vscodex/web/package-lock.json @@ -0,0 +1,2680 @@ +{ + "name": "@aether/vscodex-web", + "version": "0.4.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "@aether/vscodex-web", + "version": "0.4.0", + "dependencies": { + "@vitejs/plugin-vue": "^6.0.1", + "vite": "^7.1.3", + "vue": "^3.5.20" + }, + "devDependencies": { + "@types/node": "^24.3.0", + "@vue/test-utils": "^2.4.6", + "jsdom": "^26.1.0", + "typescript": "^5.9.2", + "vitest": "^3.2.4", + "vue-tsc": "^3.0.6" + }, + "engines": { + "node": ">=20" + } + }, + "node_modules/@asamuzakjp/css-color": { + "version": "3.2.0", + "resolved": "https://registry.npmjs.org/@asamuzakjp/css-color/-/css-color-3.2.0.tgz", + "integrity": "sha512-K1A6z8tS3XsmCMM86xoWdn7Fkdn9m6RSVtocUrJYIwZnFVkng/PvkEoWtOWmP+Scc6saYWHWZYbndEEXxl24jw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@csstools/css-calc": "^2.1.3", + "@csstools/css-color-parser": "^3.0.9", + "@csstools/css-parser-algorithms": "^3.0.4", + "@csstools/css-tokenizer": "^3.0.3", + "lru-cache": "^10.4.3" + } + }, + "node_modules/@babel/helper-string-parser": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-string-parser/-/helper-string-parser-7.29.7.tgz", + "integrity": "sha512-Pb5ijPrZ89GDH8223L4UP8i6QApWxs04RbPQJTeWDV0/keR2E36MeKnyr6LYmUUvqRRI+Iv87SuF1W6ErINzYw==", + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helper-validator-identifier": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-validator-identifier/-/helper-validator-identifier-7.29.7.tgz", + "integrity": "sha512-qehxGkRj55h/ff8EMaJ+cYhyaKlHIxqYDn682wQD7RNp9UujOQsHog2uS0r2vzr4pW+sXf90NeeayjcNaX3fFg==", + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/parser": { + "version": "7.29.8", + "resolved": "https://registry.npmjs.org/@babel/parser/-/parser-7.29.8.tgz", + "integrity": "sha512-E8lTAYNB1KW+FH+VGJuZM1ioAx2E6oVlvQFRrf5P8ZZmsiJXYAD9vTFV7yyEURNzgh1dFqMZuO6tUwcARbqFCA==", + "license": "MIT", + "dependencies": { + "@babel/types": "^7.29.8" + }, + "bin": { + "parser": "bin/babel-parser.js" + }, + "engines": { + "node": ">=6.0.0" + } + }, + "node_modules/@babel/types": { + "version": "7.29.8", + "resolved": "https://registry.npmjs.org/@babel/types/-/types-7.29.8.tgz", + "integrity": "sha512-Vj1jF3cPfxg7OAfoI7QnVKLoILlm2JF9pnVHrX8qx7AHMiYWT+NDAA7jChlNgRS4WTLc/fD1lXLmPixluj+3Gg==", + "license": "MIT", + "dependencies": { + "@babel/helper-string-parser": "^7.29.7", + "@babel/helper-validator-identifier": "^7.29.7" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@csstools/color-helpers": { + "version": "5.1.0", + "resolved": "https://registry.npmjs.org/@csstools/color-helpers/-/color-helpers-5.1.0.tgz", + "integrity": "sha512-S11EXWJyy0Mz5SYvRmY8nJYTFFd1LCNV+7cXyAgQtOOuzb4EsgfqDufL+9esx72/eLhsRdGZwaldu/h+E4t4BA==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT-0", + "engines": { + "node": ">=18" + } + }, + "node_modules/@csstools/css-calc": { + "version": "2.1.4", + "resolved": "https://registry.npmjs.org/@csstools/css-calc/-/css-calc-2.1.4.tgz", + "integrity": "sha512-3N8oaj+0juUw/1H3YwmDDJXCgTB1gKU6Hc/bB502u9zR0q2vd786XJH9QfrKIEgFlZmhZiq6epXl4rHqhzsIgQ==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=18" + }, + "peerDependencies": { + "@csstools/css-parser-algorithms": "^3.0.5", + "@csstools/css-tokenizer": "^3.0.4" + } + }, + "node_modules/@csstools/css-color-parser": { + "version": "3.1.0", + "resolved": "https://registry.npmjs.org/@csstools/css-color-parser/-/css-color-parser-3.1.0.tgz", + "integrity": "sha512-nbtKwh3a6xNVIp/VRuXV64yTKnb1IjTAEEh3irzS+HkKjAOYLTGNb9pmVNntZ8iVBHcWDA2Dof0QtPgFI1BaTA==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "dependencies": { + "@csstools/color-helpers": "^5.1.0", + "@csstools/css-calc": "^2.1.4" + }, + "engines": { + "node": ">=18" + }, + "peerDependencies": { + "@csstools/css-parser-algorithms": "^3.0.5", + "@csstools/css-tokenizer": "^3.0.4" + } + }, + "node_modules/@csstools/css-parser-algorithms": { + "version": "3.0.5", + "resolved": "https://registry.npmjs.org/@csstools/css-parser-algorithms/-/css-parser-algorithms-3.0.5.tgz", + "integrity": "sha512-DaDeUkXZKjdGhgYaHNJTV9pV7Y9B3b644jCLs9Upc3VeNGg6LWARAT6O+Q+/COo+2gg/bM5rhpMAtf70WqfBdQ==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=18" + }, + "peerDependencies": { + "@csstools/css-tokenizer": "^3.0.4" + } + }, + "node_modules/@csstools/css-tokenizer": { + "version": "3.0.4", + "resolved": "https://registry.npmjs.org/@csstools/css-tokenizer/-/css-tokenizer-3.0.4.tgz", + "integrity": "sha512-Vd/9EVDiu6PPJt9yAh6roZP6El1xHrdvIVGjyBsHR0RYwNHgL7FJPyIIW4fANJNG6FtyZfvlRPpFI4ZM/lubvw==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/aix-ppc64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/aix-ppc64/-/aix-ppc64-0.28.2.tgz", + "integrity": "sha512-XExcO+dvLKvVtNTibSTBej1NCAbaGhWn9Ww1ZPx80qsahhPFe/8jgWP0IchNe0F3HwkU7n8ejhH8bjonqht8mQ==", + "cpu": [ + "ppc64" + ], + "license": "MIT", + "optional": true, + "os": [ + "aix" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-arm": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-arm/-/android-arm-0.28.2.tgz", + "integrity": "sha512-kXXoiPVVGQcnIYGOeaovwOURpniDBpSq4A03qkQ+BMQqtGG6HYap3xne9C1O1yo4TR3qxlCX5IqqmX6fFo2Lqg==", + "cpu": [ + "arm" + ], + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-arm64/-/android-arm64-0.28.2.tgz", + "integrity": "sha512-5YfKeeI8qWfBZIX+u2xZC3Zlb3Os/gLS2sbEKM+I4ZOcsWmHS2WLysCcQZDAFRslDUU5Oiq44gf6PYN1vGwG5A==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-x64/-/android-x64-0.28.2.tgz", + "integrity": "sha512-O387ite7SzUyCcy3JQX4P4bLtEA7bLLkx+esve5JHnyYfNTxcVpXZo9jhdB0lTKN44gztELTdU7nS8Nr16Fs1Q==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/darwin-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.28.2.tgz", + "integrity": "sha512-n4KqkOQrraxHJcgjM1RvwbigfQKIKJVpM7xp+KsxiyUSrRdIXnt73VhrPAx0fV44hgfmIVKjxMN9J1t5jySVkw==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/darwin-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/darwin-x64/-/darwin-x64-0.28.2.tgz", + "integrity": "sha512-uq6suIWYP37qzGddBKPw5QEQPi6HiLGsO7UmkpfyaYNQ3D+rN6w6WfwH+nuqcGXWvawGwxOEroO4YGnFh95azw==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/freebsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/freebsd-arm64/-/freebsd-arm64-0.28.2.tgz", + "integrity": "sha512-n+I0BTSRIoy+d6RPKnEVwql5UwBJolytvY4mAOIEJorKlqgPII8ix6slVVrfZ5Tnj7glIZvloylbB/EJPMWEXw==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/freebsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/freebsd-x64/-/freebsd-x64-0.28.2.tgz", + "integrity": "sha512-78XJTJkvPs0kz2w61301PJjXl4g7q3JqiYMZ/M/yVI73EHBrCRTgkhu9oqG7vPqq+a/yadEW8aD+agKlk5xrmg==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-arm": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-arm/-/linux-arm-0.28.2.tgz", + "integrity": "sha512-XlDnu2q5yoqems+xay6wSAcg9DDD7K9RLKZEBOMZm3ckNpJBvOX20tSfby8KfrrhINDyv9V2YVZKY/SpoGJI8w==", + "cpu": [ + "arm" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-arm64/-/linux-arm64-0.28.2.tgz", + "integrity": "sha512-pW4AC0P3it8c7do9MVM4p51FzHzdM/TZrerurgRcHJ2WTa1VQ1CIq18xncfpBJw4ojkiZZrKW2yIBWBP92j6Ug==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-ia32": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-ia32/-/linux-ia32-0.28.2.tgz", + "integrity": "sha512-CYbnj78HsIeA+DhgUKgFCfvNsTHFhMMrinUrMZpDXJXKN8T3XViTZ/+wtHeVxEWY8ewSzTFN+nRmSwO2tZaLUQ==", + "cpu": [ + "ia32" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-loong64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-loong64/-/linux-loong64-0.28.2.tgz", + "integrity": "sha512-buwkd8nsph4R+ajRvw0qM5Hja/TXQow3ptzWO2EbG/cqcIkHloRrdlBtQlshyYGTNFvfkfJ5tpPLVkY4DtsPfQ==", + "cpu": [ + "loong64" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-mips64el": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-mips64el/-/linux-mips64el-0.28.2.tgz", + "integrity": "sha512-ZVykbDyk7519VwiNb9Lcj9m8XM6v5V9uKPvrEMkkEedVewf+0itkhahp4HDpgERXhwLRpWFypsGbG/J8s0QjJA==", + "cpu": [ + "mips64el" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-ppc64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-ppc64/-/linux-ppc64-0.28.2.tgz", + "integrity": "sha512-CAXl+Dtd9UUuJd8pKKdwh6MLm3MUMiqMPmhZ3tTSXPqfyQ3vDl6R5hZdZ/kYojK4ofXtdfSv1tFq8XzWx3heNQ==", + "cpu": [ + "ppc64" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-riscv64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-riscv64/-/linux-riscv64-0.28.2.tgz", + "integrity": "sha512-GeXCej4IQtU1B+QlDV8W/RRvbzI3O/Stss+/bCXv4lZls5WGRtu2a+3JkA3i4qIUlMXpcHebWpF8AkJhATowuA==", + "cpu": [ + "riscv64" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-s390x": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-s390x/-/linux-s390x-0.28.2.tgz", + "integrity": "sha512-3H1weTYZPxt/WOhByszQZybS9w5lKzUn1FDMsgEChbHWQwHYQQRfBxgCcZvPhjHfKyJjIievvMmEUawJrdY9Dg==", + "cpu": [ + "s390x" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-x64/-/linux-x64-0.28.2.tgz", + "integrity": "sha512-4xTZr1FUmSoQW4XIWmit3tzQrUTZM+N3P0XV8xROKYF50XfI7xeO90+1bZvNwxIufQ9hDQVRJH5YhgPVF8A/HQ==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/netbsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/netbsd-arm64/-/netbsd-arm64-0.28.2.tgz", + "integrity": "sha512-sSATRjPeDBg3pdgHoQfoYBob11Kk1FGa9lui5RIHZCoCkJa9QKlvl3/vKz2usCmYYjs7ymJR/2Nnsqe+Hjt5nw==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "netbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/netbsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/netbsd-x64/-/netbsd-x64-0.28.2.tgz", + "integrity": "sha512-lqnzCV+mM0gIADaKihiCg6ifgfU2L3h5E33rNQBN1Y4MaVGnzryzmvvf7UHxprpQdE8hpqLolJ9Rl+SkIRDpyw==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "netbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openbsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openbsd-arm64/-/openbsd-arm64-0.28.2.tgz", + "integrity": "sha512-AL2qJILH7lNjrDmCQDvdxMfAUIv8KMNZOvrwAQ8i8//ntL9FflhOyMJ8OZSMBb8/AWXe3/5v5S20y3zCoZWKoQ==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "openbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openbsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openbsd-x64/-/openbsd-x64-0.28.2.tgz", + "integrity": "sha512-QtiuPytchRyC4rwUKhexJdQKvDuZ6hWloi3igqPQNUJCS1/v9EiO3UTOXR6A3FoMo4fnAKbWJdqaIwhOzh8qEw==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "openbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openharmony-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openharmony-arm64/-/openharmony-arm64-0.28.2.tgz", + "integrity": "sha512-WkhYDmpTjLvGlScA1rwjRUmhl4k8oXR3cIbtqWmELgU/dFeHHlEllxDvdWcNJV9rbzCexB5vz8gtNewWLgCT7Q==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "openharmony" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/sunos-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/sunos-x64/-/sunos-x64-0.28.2.tgz", + "integrity": "sha512-GPMSkTOtMnv2U2F8gxe4Io6qmVs+YKyp832Etqqxr0hFngmXQ3rzwytelm3GIn7T4VviRUlf3sOgBOiTdvaf7g==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "sunos" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-arm64/-/win32-arm64-0.28.2.tgz", + "integrity": "sha512-PIhhEkE9uPBleRBrQEJpUn7MBnibZzbGzYWPmY3x+YoVg/95zbjB4CxPPOQ8l5tYYM4mMaCthF8/1DIfBQQyWQ==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-ia32": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-ia32/-/win32-ia32-0.28.2.tgz", + "integrity": "sha512-YmJbfTlvU7Sdn9BB+4PRES4oB6pxgS37MAONj+hBr/cpXS1aBPKXxNnDbu+QCWPj0o9dgyxeq79g6c5P8KeuYA==", + "cpu": [ + "ia32" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.28.2.tgz", + "integrity": "sha512-5ebpxr3nWMzrL/rnUI755Jkuee0bHL/Gq0WTF9lvcpv73wAp5eu8MfBUgWK9bhWvZjj7yX8etf/8tI8Ney695g==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@jridgewell/sourcemap-codec": { + "version": "1.6.0", + "resolved": "https://registry.npmjs.org/@jridgewell/sourcemap-codec/-/sourcemap-codec-1.6.0.tgz", + "integrity": "sha512-T7jf+5zgsZHwNJ4lvQ7/aezbyk0nNX+zJVWpmHA7VYsEx7a7qr5Rg5IbtJFqkgze5Y2sruq1RUY8Q837Od7iFw==", + "license": "MIT" + }, + "node_modules/@napi-rs/lzma-linux-x64-gnu": { + "version": "1.5.1", + "resolved": "https://registry.npmjs.org/@napi-rs/lzma-linux-x64-gnu/-/lzma-linux-x64-gnu-1.5.1.tgz", + "integrity": "sha512-oTXEIha4SsuXdTA4Iyskj0kpdx2yVXdhd75c2v3xGrHFfVMsbhTPZU/nMPL4sWKo4pBHm3aucLaqGlF696dTyQ==", + "cpu": [ + "x64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^22.20 || ^24.12 || >=25" + } + }, + "node_modules/@one-ini/wasm": { + "version": "0.2.1", + "resolved": "https://registry.npmjs.org/@one-ini/wasm/-/wasm-0.2.1.tgz", + "integrity": "sha512-TUqERXGNTifZ9y2g3wPxQrw3HpHv/02DsW3D90T9x0hhonrL1ZqpSmNrU2XkoIq0fP1N6gZfVQzy2Fw1ZvGBNg==", + "dev": true, + "license": "MIT" + }, + "node_modules/@rolldown/pluginutils": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.1.tgz", + "integrity": "sha512-2j9bGt5Jh8hj+vPtgzPtl72j0yRxHAyumoo6TNfAjsLB04UtpSvPbPcDcBMxz7n+9CYB0c1GxQFxYRg2jimqGw==", + "license": "MIT" + }, + "node_modules/@rollup/rollup-android-arm-eabi": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-android-arm-eabi/-/rollup-android-arm-eabi-4.63.1.tgz", + "integrity": "sha512-UZ8sUxPTiHWYX9QNdJedb1kDZSpS1t/VPWBWGSgqHNi9w3Cu6IXvu2mzbhiTiPvtrqgTQJ+zqiAq2iPIPilpaQ==", + "cpu": [ + "arm" + ], + "license": "MIT", + "optional": true, + "os": [ + "android" + ] + }, + "node_modules/@rollup/rollup-android-arm64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-android-arm64/-/rollup-android-arm64-4.63.1.tgz", + "integrity": "sha512-cQ4nFQABN5cDvDpbvJ7bMStCpnaVxynZrRMfUJYgxcIk9Sh54FIO1vtfkg0B69REjER77ioZ/ov+eAApx/KmLQ==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "android" + ] + }, + "node_modules/@rollup/rollup-darwin-arm64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-darwin-arm64/-/rollup-darwin-arm64-4.63.1.tgz", + "integrity": "sha512-FQNqd1lRy/0QhDk3xeRIkSBiCpXCiDnZO3YLVdcDKN1UBiKToNftCzcXYNLshmPDUMlu2TdeS8tGcsU6f3YF1Q==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@rollup/rollup-darwin-x64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-darwin-x64/-/rollup-darwin-x64-4.63.1.tgz", + "integrity": "sha512-pvD16V939D3CloK0+qikpGaxiPrDUXTe7Y5cWOMkMSy7m1cawa8EGy/kXYi/G/cKAC4HDAbSnzCIk1WmsoOKXg==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@rollup/rollup-freebsd-arm64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-freebsd-arm64/-/rollup-freebsd-arm64-4.63.1.tgz", + "integrity": "sha512-pcFGeL2345VwdTnJhA6zLbew+YgWB0qBG2+dMtXjCicf6+rm6kO6cOoh5VnTe0ZMrMRgRyuHmCJxZWrIdzYuOw==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ] + }, + "node_modules/@rollup/rollup-freebsd-x64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-freebsd-x64/-/rollup-freebsd-x64-4.63.1.tgz", + "integrity": "sha512-mRJlqSRulVzcKq/LKA6ICSIc3K/l4fzlVn/gePn2nXIHy8seRi5z/eeRE0d/XMBxcMldiXtQTSpRj0tkkC3g8Q==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ] + }, + "node_modules/@rollup/rollup-linux-arm-gnueabihf": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-arm-gnueabihf/-/rollup-linux-arm-gnueabihf-4.63.1.tgz", + "integrity": "sha512-YDUNvVM85TI3g/1OpnqKP1h4NeW/j64DfWMf+G3M809xNk1bJSnpFp4sh83NpmVE5DXnkh8ULor4LTVZKoYLHw==", + "cpu": [ + "arm" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-arm-musleabihf": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-arm-musleabihf/-/rollup-linux-arm-musleabihf-4.63.1.tgz", + "integrity": "sha512-7Mcn71p9ZuQFAj+h+dhQXy/yeLePRS2yKRnmW1DijA9thKO5qap0GNOIQK4yQ6iP3SU0Mrb/yWo8h8vgRba8lw==", + "cpu": [ + "arm" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-arm64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-arm64-gnu/-/rollup-linux-arm64-gnu-4.63.1.tgz", + "integrity": "sha512-4YiLQTX6U4CSl0L9cluep9A9W6UmTfqBDc2/CH6wlu54pl4E7Jn3cOD8oxzvBDEGk/JMKgJ47C8g+radF7mwvg==", + "cpu": [ + "arm64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-arm64-musl": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-arm64-musl/-/rollup-linux-arm64-musl-4.63.1.tgz", + "integrity": "sha512-2ra8F7w8OquwZN9z2/fKFnli69wa8PLwaVzRMIPGb13ByMJwC28Fbp8YcVGoUhlYMTt7j5j9bNgpysrN2UM+vw==", + "cpu": [ + "arm64" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-loong64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-loong64-gnu/-/rollup-linux-loong64-gnu-4.63.1.tgz", + "integrity": "sha512-Sy20ncyhjmBP0Ml+UvQbimjlk6VFgjW5uNP+qqwHB00mTE8Bl2C1TuHTlRwK2YoXeZbee5lP2XevBWVkAQAtSQ==", + "cpu": [ + "loong64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-loong64-musl": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-loong64-musl/-/rollup-linux-loong64-musl-4.63.1.tgz", + "integrity": "sha512-noITLp8oNjYliPnGWmLyelIHwULGqbHloQHGw1rtxbWhTuWooRpnZarZQJ1y9EUC4szuCusCc+HEpUtxpIwYvA==", + "cpu": [ + "loong64" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-ppc64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-ppc64-gnu/-/rollup-linux-ppc64-gnu-4.63.1.tgz", + "integrity": "sha512-hlxxXd+F1mWiAcaFR7Sv9ZQT6m6UfI8+Vy/kFJzztq2pDMU/0wZ9sish0iszNZvsQDo8Gc0i5yuFEOz5dDf6fA==", + "cpu": [ + "ppc64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-ppc64-musl": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-ppc64-musl/-/rollup-linux-ppc64-musl-4.63.1.tgz", + "integrity": "sha512-EF7OpqQTQ/BvGqLzUi4rEHuagCV9MugAUXSHemwPW5vxZ75RR+jxO/2j95Ph2dalMpFHSVECjRoioHZgA9zOYA==", + "cpu": [ + "ppc64" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-riscv64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-riscv64-gnu/-/rollup-linux-riscv64-gnu-4.63.1.tgz", + "integrity": "sha512-wQO3JesW9PRkwlabQ27y7sPfVOOTLRG73I4F2UYHG5PXun3J9U3y+b7ezVKSYbsvSKGQ1k1cq8Qlun4C9kLt3w==", + "cpu": [ + "riscv64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-riscv64-musl": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-riscv64-musl/-/rollup-linux-riscv64-musl-4.63.1.tgz", + "integrity": "sha512-ouAGwhO6wHRXdnOVCOsB0tRFkA7nhNB2Nwax6oECXN0YiN8EYUTBAOudADOB1PI+yDL61TeNx/u7MVCzksNbkQ==", + "cpu": [ + "riscv64" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-s390x-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-s390x-gnu/-/rollup-linux-s390x-gnu-4.63.1.tgz", + "integrity": "sha512-q2R38Sn+1J8RxhfJ+T54wSWmyKXWec+9jgDfqO2AtArEqHO5R2aeayp5H5OYLr5UYDVGsVaZPEFUooMhYCdz5A==", + "cpu": [ + "s390x" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-x64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-x64-gnu/-/rollup-linux-x64-gnu-4.63.1.tgz", + "integrity": "sha512-gfI5T24WLLuFfSKw7Go/zDXjAAV0fny0swTaDv+WjK7vqcw4cRhFfdsyKL1n+ukI+ooBxn3bVQnyrn06WpI50w==", + "cpu": [ + "x64" + ], + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-linux-x64-musl": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-linux-x64-musl/-/rollup-linux-x64-musl-4.63.1.tgz", + "integrity": "sha512-4h6XqthmB4Hspji84wvgk+ElodTsGj+dbZqHJHHtKxj4mYq0ANSEEPX9ys3moJueqsRjwpaJYH7874Itwnj2ow==", + "cpu": [ + "x64" + ], + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@rollup/rollup-openbsd-x64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-openbsd-x64/-/rollup-openbsd-x64-4.63.1.tgz", + "integrity": "sha512-dlfCOa87o1VAYegLQ9EKilx2JCeRofiyPGhTCmqnuXZ6bMPiycO1rq1+sKoulAp7pGLIsTIw+1x5R+zgh5LhhA==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "openbsd" + ] + }, + "node_modules/@rollup/rollup-openharmony-arm64": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-openharmony-arm64/-/rollup-openharmony-arm64-4.63.1.tgz", + "integrity": "sha512-cjkLbOlfcm3QGhMM1J5zaZjsw1GggbN6rw9UTSSRrPrR1KkcXnN7Uq9rPw34xImQ9VOY9GN+6u2Zj80B9ptkcw==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "openharmony" + ] + }, + "node_modules/@rollup/rollup-win32-arm64-msvc": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-win32-arm64-msvc/-/rollup-win32-arm64-msvc-4.63.1.tgz", + "integrity": "sha512-Li1KdUnWGE4N3e1F/B4RTB1ms+nG4WBgjByO46pkeBVX/2UBsY53xf5vK9WygVmnH3RwncIST7lkSdLSY6P9lg==", + "cpu": [ + "arm64" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/@rollup/rollup-win32-ia32-msvc": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-win32-ia32-msvc/-/rollup-win32-ia32-msvc-4.63.1.tgz", + "integrity": "sha512-t4ZYOSoLTgwhuFMrmTMLx/+i1DQVK7HYqMc6kY46EApwi8X0nIVphzdNoThU3xt6n+N5urG1/gxBdCaKDLavfg==", + "cpu": [ + "ia32" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/@rollup/rollup-win32-x64-gnu": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-win32-x64-gnu/-/rollup-win32-x64-gnu-4.63.1.tgz", + "integrity": "sha512-RgroPfMmKlD1RzSDxvwgcPiy2HNQKoYV7OmwIXDsk73uKW5t6B/V8KIy27SMv/FNXFo/oSBtWc9J0X7t91ezZg==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/@rollup/rollup-win32-x64-msvc": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/@rollup/rollup-win32-x64-msvc/-/rollup-win32-x64-msvc-4.63.1.tgz", + "integrity": "sha512-at8QVep6S3h5Y6gSbdGU06bRY5WJkf6WUduM9YtvYMbYhB1MOFfUgc6kehitQXzOtMSaT70q7f9ydPhpqu821w==", + "cpu": [ + "x64" + ], + "license": "MIT", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/@types/chai": { + "version": "5.2.3", + "resolved": "https://registry.npmjs.org/@types/chai/-/chai-5.2.3.tgz", + "integrity": "sha512-Mw558oeA9fFbv65/y4mHtXDs9bPnFMZAL/jxdPFUpOHHIXX91mcgEHbS5Lahr+pwZFR8A7GQleRWeI6cGFC2UA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/deep-eql": "*", + "assertion-error": "^2.0.1" + } + }, + "node_modules/@types/deep-eql": { + "version": "4.0.2", + "resolved": "https://registry.npmjs.org/@types/deep-eql/-/deep-eql-4.0.2.tgz", + "integrity": "sha512-c9h9dVVMigMPc4bwTvC5dxqtqJZwQPePsWjPlpSOnojbor6pGqdk541lfA7AqFQr5pB1BRdq0juY9db81BwyFw==", + "dev": true, + "license": "MIT" + }, + "node_modules/@types/estree": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/@types/estree/-/estree-1.0.9.tgz", + "integrity": "sha512-GhdPgy1el4/ImP05X05Uw4cw2/M93BCUmnEvWZNStlCzEKME4Fkk+YpoA5OiHNQmoS7Cafb8Xa3Pya8m1Qrzeg==", + "license": "MIT" + }, + "node_modules/@types/node": { + "version": "24.13.3", + "resolved": "https://registry.npmjs.org/@types/node/-/node-24.13.3.tgz", + "integrity": "sha512-Dh8vAsV36ig5wa9OX4pXvMc9D3Veibfw2wix0CUwYODLD8nkj9UsLjASr49nPg+2eKzxhBV+v7L8pXvT4e639Q==", + "devOptional": true, + "license": "MIT", + "dependencies": { + "undici-types": "~7.18.0" + } + }, + "node_modules/@vitejs/plugin-vue": { + "version": "6.0.8", + "resolved": "https://registry.npmjs.org/@vitejs/plugin-vue/-/plugin-vue-6.0.8.tgz", + "integrity": "sha512-0ZjgOg7oO6farnNGup7yvoM/YXZV84OZxHAwtflItNa/6zzQyVb5LNxyea3FEKEX2XlagIKzrlH7wwxkKgtiew==", + "license": "MIT", + "dependencies": { + "@rolldown/pluginutils": "^1.0.1" + }, + "engines": { + "node": "^20.19.0 || >=22.12.0" + }, + "peerDependencies": { + "vite": "^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0", + "vue": "^3.2.25" + } + }, + "node_modules/@vitest/expect": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-3.2.7.tgz", + "integrity": "sha512-E8eBXaKibuvH2pSZErOjdVb5vF4PbKYcrnluBTYxEk1l/VhhwZg1kZQsdtjq+CsF5CFydf2Rdkz7jDHKSisi3w==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/chai": "^5.2.2", + "@vitest/spy": "3.2.7", + "@vitest/utils": "3.2.7", + "chai": "^5.2.0", + "tinyrainbow": "^2.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/mocker": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-3.2.7.tgz", + "integrity": "sha512-Trr0hYO9CM3Wj6ksWHRhK9IZpIY6wTMO5u/MqXurMxT57sWBaOPEtP3Oq60ihZuh5JsiagKfz95OcxdEP6dBrA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/spy": "3.2.7", + "estree-walker": "^3.0.3", + "magic-string": "^0.30.17" + }, + "funding": { + "url": "https://opencollective.com/vitest" + }, + "peerDependencies": { + "msw": "^2.4.9", + "vite": "^5.0.0 || ^6.0.0 || ^7.0.0-0" + }, + "peerDependenciesMeta": { + "msw": { + "optional": true + }, + "vite": { + "optional": true + } + } + }, + "node_modules/@vitest/mocker/node_modules/estree-walker": { + "version": "3.0.3", + "resolved": "https://registry.npmjs.org/estree-walker/-/estree-walker-3.0.3.tgz", + "integrity": "sha512-7RUKfXgSMMkzt6ZuXmqapOurLGPPfgj6l9uRZ7lRGolvk0y2yocc35LdcxKC5PQZdn2DMqioAQ2NoWcrTKmm6g==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/estree": "^1.0.0" + } + }, + "node_modules/@vitest/pretty-format": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-3.2.7.tgz", + "integrity": "sha512-KUHlwqVu0sRlhCdyPdQ/wBoTfRahjUky1MubOmYw9fWfIZy1gNoHpuaaQBPAaMaVYdQYHJLurzj8ECCj5OwTqA==", + "dev": true, + "license": "MIT", + "dependencies": { + "tinyrainbow": "^2.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/runner": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-3.2.7.tgz", + "integrity": "sha512-sB9y4ovltoQP+WaUPwmSxO9WIg9Ig694Di5PalVPsYHklAdE027mehpWF2SQSVq+k6sFgaivbTjTJwZLSHbedA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/utils": "3.2.7", + "pathe": "^2.0.3", + "strip-literal": "^3.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/snapshot": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-3.2.7.tgz", + "integrity": "sha512-7C+MwShwtBSI5Buwoyg3s/iY1eHL9PKAf+O1wVh/TdnjXUtkoL/9YQtre90i4MtNXM6edP1wJ2zOBpfCyhIS7g==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/pretty-format": "3.2.7", + "magic-string": "^0.30.17", + "pathe": "^2.0.3" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/spy": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-3.2.7.tgz", + "integrity": "sha512-Q2eQGI6d2L/hBtZ0qNuKcAGid68XK6cv1xsoaIma6PaJhHPoqcEJhYpXZ/5myCMqkNgtP6UKuBhbc0nHKnrkuQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "tinyspy": "^4.0.3" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/utils": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-3.2.7.tgz", + "integrity": "sha512-x6BDOd7dyo3PFLY3I9/HJ25X/6OurhGXk2/B9gOZNPF7XDVjeBK4k01lQE5uvDpbuheErh91qYuE1E2OEjK3Rw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/pretty-format": "3.2.7", + "loupe": "^3.1.4", + "tinyrainbow": "^2.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@volar/language-core": { + "version": "2.4.28", + "resolved": "https://registry.npmjs.org/@volar/language-core/-/language-core-2.4.28.tgz", + "integrity": "sha512-w4qhIJ8ZSitgLAkVay6AbcnC7gP3glYM3fYwKV3srj8m494E3xtrCv6E+bWviiK/8hs6e6t1ij1s2Endql7vzQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@volar/source-map": "2.4.28" + } + }, + "node_modules/@volar/source-map": { + "version": "2.4.28", + "resolved": "https://registry.npmjs.org/@volar/source-map/-/source-map-2.4.28.tgz", + "integrity": "sha512-yX2BDBqJkRXfKw8my8VarTyjv48QwxdJtvRgUpNE5erCsgEUdI2DsLbpa+rOQVAJYshY99szEcRDmyHbF10ggQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/@volar/typescript": { + "version": "2.4.28", + "resolved": "https://registry.npmjs.org/@volar/typescript/-/typescript-2.4.28.tgz", + "integrity": "sha512-Ja6yvWrbis2QtN4ClAKreeUZPVYMARDYZl9LMEv1iQ1QdepB6wn0jTRxA9MftYmYa4DQ4k/DaSZpFPUfxl8giw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@volar/language-core": "2.4.28", + "path-browserify": "^1.0.1", + "vscode-uri": "^3.0.8" + } + }, + "node_modules/@vue/compiler-core": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/compiler-core/-/compiler-core-3.5.42.tgz", + "integrity": "sha512-2Ye1ilMtKXxl8qZUrQ5j0CdgenFp/HFQmta6rfRyfEsTG69L6Wk+tWuNoHYHMx9E8tF2Slvdg1FuwDvAXdy1LQ==", + "license": "MIT", + "dependencies": { + "@babel/parser": "^7.29.8", + "@vue/shared": "3.5.42", + "entities": "^7.0.1", + "estree-walker": "^2.0.2", + "source-map-js": "^1.2.1" + } + }, + "node_modules/@vue/compiler-dom": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/compiler-dom/-/compiler-dom-3.5.42.tgz", + "integrity": "sha512-qbhQZEFmycr+ni/qyuccS4sucNN7VAbDfbkvNxWOX2VfgFm90MNs3/UhRNKoPMEIVn0F8gdlYjLPvqxHwHeQOA==", + "license": "MIT", + "dependencies": { + "@vue/compiler-core": "3.5.42", + "@vue/shared": "3.5.42" + } + }, + "node_modules/@vue/compiler-sfc": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/compiler-sfc/-/compiler-sfc-3.5.42.tgz", + "integrity": "sha512-fkCAFB4okcAANGMThboWnScp/gzWjU0ZSkVnjTIiplmMDq2uq0tIB3j+xVu4rhv5rvOgBySCysudmbMd6xRRqw==", + "license": "MIT", + "dependencies": { + "@babel/parser": "^7.29.8", + "@vue/compiler-core": "3.5.42", + "@vue/compiler-dom": "3.5.42", + "@vue/compiler-ssr": "3.5.42", + "@vue/shared": "3.5.42", + "estree-walker": "^2.0.2", + "magic-string": "^0.30.21", + "postcss": "^8.5.19", + "source-map-js": "^1.2.1" + } + }, + "node_modules/@vue/compiler-ssr": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/compiler-ssr/-/compiler-ssr-3.5.42.tgz", + "integrity": "sha512-xmLk3wLkbizPAiLyomjgFFosf2ys9b5Ghb+oh/k2tnvipNz8OFrQOiTcWCzyK7MpBp9KkyGtfvgfLUivbmuGYA==", + "license": "MIT", + "dependencies": { + "@vue/compiler-dom": "3.5.42", + "@vue/shared": "3.5.42" + } + }, + "node_modules/@vue/language-core": { + "version": "3.3.11", + "resolved": "https://registry.npmjs.org/@vue/language-core/-/language-core-3.3.11.tgz", + "integrity": "sha512-QJmpliwAVpC/OxubIByPAhNzsQPRc8/gxlN2qnVzVfIMjMDz/9RnXRFoetjz5yEgXVXyp4LqhXq3V53PjmNzFw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@volar/language-core": "2.4.28", + "@vue/compiler-dom": "^3.5.0", + "@vue/shared": "^3.5.0", + "alien-signals": "^3.2.1", + "muggle-string": "^0.4.1", + "path-browserify": "^1.0.1", + "picomatch": "^4.0.4" + } + }, + "node_modules/@vue/reactivity": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/reactivity/-/reactivity-3.5.42.tgz", + "integrity": "sha512-TzNNfKpb7hDxbQltwAut8VDQA5YP+BuRlxntHUuRjyKwlMvmAPbs3unhCvieijifY6vFfVBwsS7wG/C7uq+bEQ==", + "license": "MIT", + "dependencies": { + "@vue/shared": "3.5.42" + } + }, + "node_modules/@vue/runtime-core": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/runtime-core/-/runtime-core-3.5.42.tgz", + "integrity": "sha512-9uACtuHs7vJGkm5Bp3xu4xRDLFTIYy5DgxpToVjqGIAhAEKwQfsaLvKINhM6nFVp6bZPRFGdDqd1g52MqKsotA==", + "license": "MIT", + "dependencies": { + "@vue/reactivity": "3.5.42", + "@vue/shared": "3.5.42" + } + }, + "node_modules/@vue/runtime-dom": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/runtime-dom/-/runtime-dom-3.5.42.tgz", + "integrity": "sha512-rsCmhiWLaRxGltLwhlCWyYkFn7WAbKRh0q17eZ1A6Dq6eqc2ACQ61IIryxz0LrsvCzHSilLA9JHovVwM8CNE2g==", + "license": "MIT", + "dependencies": { + "@vue/reactivity": "3.5.42", + "@vue/runtime-core": "3.5.42", + "@vue/shared": "3.5.42", + "csstype": "^3.2.3" + } + }, + "node_modules/@vue/server-renderer": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/server-renderer/-/server-renderer-3.5.42.tgz", + "integrity": "sha512-2++5dUyYS4gvo7xQXSECUDhB7TS0aOl5SeVfC5qSq1Jgfhjvegw1zqhwTIR3imZ+QYPJQw9gfcFvXGAjGZ7ajQ==", + "license": "MIT", + "dependencies": { + "@vue/compiler-ssr": "3.5.42", + "@vue/runtime-dom": "3.5.42", + "@vue/shared": "3.5.42" + } + }, + "node_modules/@vue/shared": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/@vue/shared/-/shared-3.5.42.tgz", + "integrity": "sha512-2rPxex1jQf4jvl9MOHl6YaXCPcrNqz/FstMOEh3QWY+/OME9nQTvl9WYeCwhW7AFjaR0SnngZGlp/wkR6rkI6g==", + "license": "MIT" + }, + "node_modules/@vue/test-utils": { + "version": "2.5.0", + "resolved": "https://registry.npmjs.org/@vue/test-utils/-/test-utils-2.5.0.tgz", + "integrity": "sha512-6Clu5EKR/r6cDPYrKsu+8wenciWJJ3rhS9OEGsfDlZeZIhlJeEPGIZQHxE4lHRJCzPSq3EWMsFxQUqCvrbHQuQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "js-beautify": "^2.0.0", + "vue-component-type-helpers": "^3.0.0" + }, + "peerDependencies": { + "@vue/compiler-dom": "3.x", + "@vue/server-renderer": "3.x", + "vue": "3.x" + }, + "peerDependenciesMeta": { + "@vue/server-renderer": { + "optional": true + } + } + }, + "node_modules/abbrev": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/abbrev/-/abbrev-5.0.0.tgz", + "integrity": "sha512-/XrFJgzQQQHpti1raDJC6m4ws6aNktmjBlhk8Fdlk7LwCEuDoieEJJY9OFHjfiFJFFRM2tK+Ky/IsfbbmlMu1w==", + "dev": true, + "license": "ISC", + "engines": { + "node": "^22.22.2 || ^24.15.0 || >=26.0.0" + } + }, + "node_modules/agent-base": { + "version": "7.1.4", + "resolved": "https://registry.npmjs.org/agent-base/-/agent-base-7.1.4.tgz", + "integrity": "sha512-MnA+YT8fwfJPgBx3m60MNqakm30XOkyIoH1y6huTQvC0PwZG7ki8NacLBcrPbNoo8vEZy7Jpuk7+jMO+CUovTQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 14" + } + }, + "node_modules/alien-signals": { + "version": "3.2.1", + "resolved": "https://registry.npmjs.org/alien-signals/-/alien-signals-3.2.1.tgz", + "integrity": "sha512-I8FjmltrfnDFoZedi5CG8DghVYNhzb/Ijluz7tCSJH0xpd0484Kowhbb1XDYOxfJpU1p5wnM2X54dA+IfGyD1g==", + "dev": true, + "license": "MIT" + }, + "node_modules/assertion-error": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/assertion-error/-/assertion-error-2.0.1.tgz", + "integrity": "sha512-Izi8RQcffqCeNVgFigKli1ssklIbpHnCYc6AknXGYoB6grJqyeby7jv12JUQgmTAnIDnbck1uxksT4dzN3PWBA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=12" + } + }, + "node_modules/balanced-match": { + "version": "4.0.4", + "resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-4.0.4.tgz", + "integrity": "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA==", + "dev": true, + "license": "MIT", + "engines": { + "node": "18 || 20 || >=22" + } + }, + "node_modules/brace-expansion": { + "version": "5.0.9", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.9.tgz", + "integrity": "sha512-ScQ4IuvIEF1TMlP7Zt+vjJ//9zlPb2SDcxWxM3bk8s6t6GGdJ7KO1dCcTidOPJKePW30LE/2cT7wCyPho9/Wxg==", + "dev": true, + "license": "MIT", + "dependencies": { + "balanced-match": "^4.0.2" + }, + "engines": { + "node": "20 || >=22" + } + }, + "node_modules/cac": { + "version": "6.7.14", + "resolved": "https://registry.npmjs.org/cac/-/cac-6.7.14.tgz", + "integrity": "sha512-b6Ilus+c3RrdDk+JhLKUAQfzzgLEPy6wcXqS7f/xe1EETvsDP6GORG7SFuOs6cID5YkqchW/LXZbX5bc8j7ZcQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=8" + } + }, + "node_modules/chai": { + "version": "5.3.3", + "resolved": "https://registry.npmjs.org/chai/-/chai-5.3.3.tgz", + "integrity": "sha512-4zNhdJD/iOjSH0A05ea+Ke6MU5mmpQcbQsSOkgdaUMJ9zTlDTD/GYlwohmIE2u0gaxHYiVHEn1Fw9mZ/ktJWgw==", + "dev": true, + "license": "MIT", + "dependencies": { + "assertion-error": "^2.0.1", + "check-error": "^2.1.1", + "deep-eql": "^5.0.1", + "loupe": "^3.1.0", + "pathval": "^2.0.0" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/check-error": { + "version": "2.1.3", + "resolved": "https://registry.npmjs.org/check-error/-/check-error-2.1.3.tgz", + "integrity": "sha512-PAJdDJusoxnwm1VwW07VWwUN1sl7smmC3OKggvndJFadxxDRyFJBX/ggnu/KE4kQAB7a3Dp8f/YXC1FlUprWmA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 16" + } + }, + "node_modules/commander": { + "version": "14.0.3", + "resolved": "https://registry.npmjs.org/commander/-/commander-14.0.3.tgz", + "integrity": "sha512-H+y0Jo/T1RZ9qPP4Eh1pkcQcLRglraJaSLoyOtHxu6AapkjWVCy2Sit1QQ4x3Dng8qDlSsZEet7g5Pq06MvTgw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=20" + } + }, + "node_modules/config-chain": { + "version": "1.1.13", + "resolved": "https://registry.npmjs.org/config-chain/-/config-chain-1.1.13.tgz", + "integrity": "sha512-qj+f8APARXHrM0hraqXYb2/bOVSV4PvJQlNZ/DVj0QrmNM2q2euizkeuVckQ57J+W0mRH6Hvi+k50M4Jul2VRQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "ini": "^1.3.4", + "proto-list": "~1.2.1" + } + }, + "node_modules/cssstyle": { + "version": "4.6.0", + "resolved": "https://registry.npmjs.org/cssstyle/-/cssstyle-4.6.0.tgz", + "integrity": "sha512-2z+rWdzbbSZv6/rhtvzvqeZQHrBaqgogqt85sqFNbabZOuFbCVFb8kPeEtZjiKkbrm395irpNKiYeFeLiQnFPg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@asamuzakjp/css-color": "^3.2.0", + "rrweb-cssom": "^0.8.0" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/csstype": { + "version": "3.2.3", + "resolved": "https://registry.npmjs.org/csstype/-/csstype-3.2.3.tgz", + "integrity": "sha512-z1HGKcYy2xA8AGQfwrn0PAy+PB7X/GSj3UVJW9qKyn43xWa+gl5nXmU4qqLMRzWVLFC8KusUX8T/0kCiOYpAIQ==", + "license": "MIT" + }, + "node_modules/data-urls": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/data-urls/-/data-urls-5.0.0.tgz", + "integrity": "sha512-ZYP5VBHshaDAiVZxjbRVcFJpc+4xGgT0bK3vzy1HLN8jTO975HEbuYzZJcHoQEY5K1a0z8YayJkyVETa08eNTg==", + "dev": true, + "license": "MIT", + "dependencies": { + "whatwg-mimetype": "^4.0.0", + "whatwg-url": "^14.0.0" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/debug": { + "version": "4.4.3", + "resolved": "https://registry.npmjs.org/debug/-/debug-4.4.3.tgz", + "integrity": "sha512-RGwwWnwQvkVfavKVt22FGLw+xYSdzARwm0ru6DhTVA3umU5hZc28V3kO4stgYryrTlLpuvgI9GiijltAjNbcqA==", + "dev": true, + "license": "MIT", + "dependencies": { + "ms": "^2.1.3" + }, + "engines": { + "node": ">=6.0" + }, + "peerDependenciesMeta": { + "supports-color": { + "optional": true + } + } + }, + "node_modules/decimal.js": { + "version": "10.6.0", + "resolved": "https://registry.npmjs.org/decimal.js/-/decimal.js-10.6.0.tgz", + "integrity": "sha512-YpgQiITW3JXGntzdUmyUR1V812Hn8T1YVXhCu+wO3OpS4eU9l4YdD3qjyiKdV6mvV29zapkMeD390UVEf2lkUg==", + "dev": true, + "license": "MIT" + }, + "node_modules/deep-eql": { + "version": "5.0.2", + "resolved": "https://registry.npmjs.org/deep-eql/-/deep-eql-5.0.2.tgz", + "integrity": "sha512-h5k/5U50IJJFpzfL6nO9jaaumfjO/f2NjK/oYB2Djzm4p9L+3T9qWpZqZ2hAbLPuuYq9wrU08WQyBTL5GbPk5Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6" + } + }, + "node_modules/editorconfig": { + "version": "3.0.2", + "resolved": "https://registry.npmjs.org/editorconfig/-/editorconfig-3.0.2.tgz", + "integrity": "sha512-T0ix8GhtxyKVfUFEcvdNDt3YGqlwkFHbD4/5bgFUDgFmxhI/cSRAeJ87/Sz//Cq8Eam6JX/e23RkoFO71P7aAA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@one-ini/wasm": "0.2.1", + "commander": "^14.0.3", + "minimatch": "~10.2.4", + "semver": "^7.7.4" + }, + "bin": { + "editorconfig": "bin/editorconfig" + }, + "engines": { + "node": ">=20" + } + }, + "node_modules/entities": { + "version": "7.0.1", + "resolved": "https://registry.npmjs.org/entities/-/entities-7.0.1.tgz", + "integrity": "sha512-TWrgLOFUQTH994YUyl1yT4uyavY5nNB5muff+RtWaqNVCAK408b5ZnnbNAUEWLTCpum9w6arT70i1XdQ4UeOPA==", + "license": "BSD-2-Clause", + "engines": { + "node": ">=0.12" + }, + "funding": { + "url": "https://github.com/fb55/entities?sponsor=1" + } + }, + "node_modules/es-module-lexer": { + "version": "1.7.0", + "resolved": "https://registry.npmjs.org/es-module-lexer/-/es-module-lexer-1.7.0.tgz", + "integrity": "sha512-jEQoCwk8hyb2AZziIOLhDqpm5+2ww5uIE6lkO/6jcOCusfk6LhMHpXXfBLXTZ7Ydyt0j4VoUQv6uGNYbdW+kBA==", + "dev": true, + "license": "MIT" + }, + "node_modules/esbuild": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/esbuild/-/esbuild-0.28.2.tgz", + "integrity": "sha512-HKVLS8dvII+xoKW9kmqxbRKrnWEXfJJr/FZhhJmiqIB0e053QNYFqOBouTMO/k5sID4MvCiUCvv8b9M4h32wIA==", + "hasInstallScript": true, + "license": "MIT", + "bin": { + "esbuild": "bin/esbuild" + }, + "engines": { + "node": ">=18" + }, + "optionalDependencies": { + "@esbuild/aix-ppc64": "0.28.2", + "@esbuild/android-arm": "0.28.2", + "@esbuild/android-arm64": "0.28.2", + "@esbuild/android-x64": "0.28.2", + "@esbuild/darwin-arm64": "0.28.2", + "@esbuild/darwin-x64": "0.28.2", + "@esbuild/freebsd-arm64": "0.28.2", + "@esbuild/freebsd-x64": "0.28.2", + "@esbuild/linux-arm": "0.28.2", + "@esbuild/linux-arm64": "0.28.2", + "@esbuild/linux-ia32": "0.28.2", + "@esbuild/linux-loong64": "0.28.2", + "@esbuild/linux-mips64el": "0.28.2", + "@esbuild/linux-ppc64": "0.28.2", + "@esbuild/linux-riscv64": "0.28.2", + "@esbuild/linux-s390x": "0.28.2", + "@esbuild/linux-x64": "0.28.2", + "@esbuild/netbsd-arm64": "0.28.2", + "@esbuild/netbsd-x64": "0.28.2", + "@esbuild/openbsd-arm64": "0.28.2", + "@esbuild/openbsd-x64": "0.28.2", + "@esbuild/openharmony-arm64": "0.28.2", + "@esbuild/sunos-x64": "0.28.2", + "@esbuild/win32-arm64": "0.28.2", + "@esbuild/win32-ia32": "0.28.2", + "@esbuild/win32-x64": "0.28.2" + } + }, + "node_modules/estree-walker": { + "version": "2.0.2", + "resolved": "https://registry.npmjs.org/estree-walker/-/estree-walker-2.0.2.tgz", + "integrity": "sha512-Rfkk/Mp/DL7JVje3u18FxFujQlTNR2q6QfMSMB7AvCBx91NGj/ba3kCfza0f6dVDbw7YlRf/nDrn7pQrCCyQ/w==", + "license": "MIT" + }, + "node_modules/expect-type": { + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.4.0.tgz", + "integrity": "sha512-KfYbmpRm0VbLjEvVa9yGwCi9GI34xvi7A/HXYWQO65CSD2u3MczUJSuwXKFIxlGsgBQizV9q5J9NHj4VG0n+pA==", + "dev": true, + "license": "Apache-2.0", + "engines": { + "node": ">=12.0.0" + } + }, + "node_modules/fdir": { + "version": "6.5.0", + "resolved": "https://registry.npmjs.org/fdir/-/fdir-6.5.0.tgz", + "integrity": "sha512-tIbYtZbucOs0BRGqPJkshJUYdL+SDH7dVM8gjy+ERp3WAUjLEFJE+02kanyHtwjWOnwrKYBiwAmM0p4kLJAnXg==", + "license": "MIT", + "engines": { + "node": ">=12.0.0" + }, + "peerDependencies": { + "picomatch": "^3 || ^4" + }, + "peerDependenciesMeta": { + "picomatch": { + "optional": true + } + } + }, + "node_modules/fsevents": { + "version": "2.3.3", + "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.3.tgz", + "integrity": "sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw==", + "hasInstallScript": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": "^8.16.0 || ^10.6.0 || >=11.0.0" + } + }, + "node_modules/glob": { + "version": "13.0.6", + "resolved": "https://registry.npmjs.org/glob/-/glob-13.0.6.tgz", + "integrity": "sha512-Wjlyrolmm8uDpm/ogGyXZXb1Z+Ca2B8NbJwqBVg0axK9GbBeoS7yGV6vjXnYdGm6X53iehEuxxbyiKp8QmN4Vw==", + "dev": true, + "license": "BlueOak-1.0.0", + "dependencies": { + "minimatch": "^10.2.2", + "minipass": "^7.1.3", + "path-scurry": "^2.0.2" + }, + "engines": { + "node": "18 || 20 || >=22" + }, + "funding": { + "url": "https://github.com/sponsors/isaacs" + } + }, + "node_modules/html-encoding-sniffer": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/html-encoding-sniffer/-/html-encoding-sniffer-4.0.0.tgz", + "integrity": "sha512-Y22oTqIU4uuPgEemfz7NDJz6OeKf12Lsu+QC+s3BVpda64lTiMYCyGwg5ki4vFxkMwQdeZDl2adZoqUgdFuTgQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "whatwg-encoding": "^3.1.1" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/http-proxy-agent": { + "version": "7.0.2", + "resolved": "https://registry.npmjs.org/http-proxy-agent/-/http-proxy-agent-7.0.2.tgz", + "integrity": "sha512-T1gkAiYYDWYx3V5Bmyu7HcfcvL7mUrTWiM6yOfa3PIphViJ/gFPbvidQ+veqSOHci/PxBcDabeUNCzpOODJZig==", + "dev": true, + "license": "MIT", + "dependencies": { + "agent-base": "^7.1.0", + "debug": "^4.3.4" + }, + "engines": { + "node": ">= 14" + } + }, + "node_modules/https-proxy-agent": { + "version": "7.0.6", + "resolved": "https://registry.npmjs.org/https-proxy-agent/-/https-proxy-agent-7.0.6.tgz", + "integrity": "sha512-vK9P5/iUfdl95AI+JVyUuIcVtd4ofvtrOr3HNtM2yxC9bnMbEdp3x01OhQNnjb8IJYi38VlTE3mBXwcfvywuSw==", + "dev": true, + "license": "MIT", + "dependencies": { + "agent-base": "^7.1.2", + "debug": "4" + }, + "engines": { + "node": ">= 14" + } + }, + "node_modules/iconv-lite": { + "version": "0.6.3", + "resolved": "https://registry.npmjs.org/iconv-lite/-/iconv-lite-0.6.3.tgz", + "integrity": "sha512-4fCk79wshMdzMp2rH06qWrJE4iolqLhCUH+OiuIgU++RB0+94NlDL81atO7GX55uUKueo0txHNtvEyI6D7WdMw==", + "dev": true, + "license": "MIT", + "dependencies": { + "safer-buffer": ">= 2.1.2 < 3.0.0" + }, + "engines": { + "node": ">=0.10.0" + } + }, + "node_modules/ini": { + "version": "1.3.8", + "resolved": "https://registry.npmjs.org/ini/-/ini-1.3.8.tgz", + "integrity": "sha512-JV/yugV2uzW5iMRSiZAyDtQd+nxtUnjeLt0acNdw98kKLrvuRVyB80tsREOE7yvGVgalhZ6RNXCmEHkUKBKxew==", + "dev": true, + "license": "ISC" + }, + "node_modules/is-potential-custom-element-name": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/is-potential-custom-element-name/-/is-potential-custom-element-name-1.0.1.tgz", + "integrity": "sha512-bCYeRA2rVibKZd+s2625gGnGF/t7DSqDs4dP7CrLA1m7jKWz6pps0LpYLJN8Q64HtmPKJ1hrN3nzPNKFEKOUiQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/js-beautify": { + "version": "2.0.3", + "resolved": "https://registry.npmjs.org/js-beautify/-/js-beautify-2.0.3.tgz", + "integrity": "sha512-cyFbh3tkPhknnTD/0bLf0T0yy2ZIbqL05mttzbt4y1Zfr7NxqXQZ62dkBLKs3oHH/lpjmDRAnciJiSUyOy8XwQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "config-chain": "^1.1.13", + "editorconfig": "^3.0.2", + "glob": "^13.0.6", + "js-cookie": "^3.0.8", + "nopt": "^10.0.1" + }, + "bin": { + "css-beautify": "js/bin/css-beautify.js", + "html-beautify": "js/bin/html-beautify.js", + "js-beautify": "js/bin/js-beautify.js" + }, + "engines": { + "node": ">=14" + } + }, + "node_modules/js-cookie": { + "version": "3.0.8", + "resolved": "https://registry.npmjs.org/js-cookie/-/js-cookie-3.0.8.tgz", + "integrity": "sha512-yeJd4aNAdYZQjaon2bpD/Gb0B/omw7HQOsynXXcOiWVCacbBcPlgn8S/d1X6blFSaHao7ozqtW7NZW19xpCtIw==", + "dev": true, + "license": "MIT" + }, + "node_modules/js-tokens": { + "version": "9.0.1", + "resolved": "https://registry.npmjs.org/js-tokens/-/js-tokens-9.0.1.tgz", + "integrity": "sha512-mxa9E9ITFOt0ban3j6L5MpjwegGz6lBQmM1IJkWeBZGcMxto50+eWdjC/52xDbS2vy0k7vIMK0Fe2wfL9OQSpQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/jsdom": { + "version": "26.1.0", + "resolved": "https://registry.npmjs.org/jsdom/-/jsdom-26.1.0.tgz", + "integrity": "sha512-Cvc9WUhxSMEo4McES3P7oK3QaXldCfNWp7pl2NNeiIFlCoLr3kfq9kb1fxftiwk1FLV7CvpvDfonxtzUDeSOPg==", + "dev": true, + "license": "MIT", + "dependencies": { + "cssstyle": "^4.2.1", + "data-urls": "^5.0.0", + "decimal.js": "^10.5.0", + "html-encoding-sniffer": "^4.0.0", + "http-proxy-agent": "^7.0.2", + "https-proxy-agent": "^7.0.6", + "is-potential-custom-element-name": "^1.0.1", + "nwsapi": "^2.2.16", + "parse5": "^7.2.1", + "rrweb-cssom": "^0.8.0", + "saxes": "^6.0.0", + "symbol-tree": "^3.2.4", + "tough-cookie": "^5.1.1", + "w3c-xmlserializer": "^5.0.0", + "webidl-conversions": "^7.0.0", + "whatwg-encoding": "^3.1.1", + "whatwg-mimetype": "^4.0.0", + "whatwg-url": "^14.1.1", + "ws": "^8.18.0", + "xml-name-validator": "^5.0.0" + }, + "engines": { + "node": ">=18" + }, + "peerDependencies": { + "canvas": "^3.0.0" + }, + "peerDependenciesMeta": { + "canvas": { + "optional": true + } + } + }, + "node_modules/loupe": { + "version": "3.2.1", + "resolved": "https://registry.npmjs.org/loupe/-/loupe-3.2.1.tgz", + "integrity": "sha512-CdzqowRJCeLU72bHvWqwRBBlLcMEtIvGrlvef74kMnV2AolS9Y8xUv1I0U/MNAWMhBlKIoyuEgoJ0t/bbwHbLQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/lru-cache": { + "version": "10.4.3", + "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-10.4.3.tgz", + "integrity": "sha512-JNAzZcXrCt42VGLuYz0zfAzDfAvJWW6AfYlDBQyDV5DClI2m5sAmK+OIO7s59XfsRsWHp02jAJrRadPRGTt6SQ==", + "dev": true, + "license": "ISC" + }, + "node_modules/magic-string": { + "version": "0.30.21", + "resolved": "https://registry.npmjs.org/magic-string/-/magic-string-0.30.21.tgz", + "integrity": "sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ==", + "license": "MIT", + "dependencies": { + "@jridgewell/sourcemap-codec": "^1.5.5" + } + }, + "node_modules/minimatch": { + "version": "10.2.6", + "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.2.6.tgz", + "integrity": "sha512-vpLQEs+VLCr1nU0BXS07maYoFwlDAH0gngQuuttxIwutDFEMHq2blX+8vpgxDdK3J1PwjCJiep77OitTZ4Ll1A==", + "dev": true, + "license": "BlueOak-1.0.0", + "dependencies": { + "brace-expansion": "^5.0.8" + }, + "engines": { + "node": "18 || 20 || >=22" + }, + "funding": { + "url": "https://github.com/sponsors/isaacs" + } + }, + "node_modules/minipass": { + "version": "7.1.3", + "resolved": "https://registry.npmjs.org/minipass/-/minipass-7.1.3.tgz", + "integrity": "sha512-tEBHqDnIoM/1rXME1zgka9g6Q2lcoCkxHLuc7ODJ5BxbP5d4c2Z5cGgtXAku59200Cx7diuHTOYfSBD8n6mm8A==", + "dev": true, + "license": "BlueOak-1.0.0", + "engines": { + "node": ">=16 || 14 >=14.17" + } + }, + "node_modules/ms": { + "version": "2.1.3", + "resolved": "https://registry.npmjs.org/ms/-/ms-2.1.3.tgz", + "integrity": "sha512-6FlzubTLZG3J2a/NVCAleEhjzq5oxgHyaCU9yYXvcLsvoVaHJq/s5xXI6/XXP6tz7R9xAOtHnSO/tXtF3WRTlA==", + "dev": true, + "license": "MIT" + }, + "node_modules/muggle-string": { + "version": "0.4.1", + "resolved": "https://registry.npmjs.org/muggle-string/-/muggle-string-0.4.1.tgz", + "integrity": "sha512-VNTrAak/KhO2i8dqqnqnAHOa3cYBwXEZe9h+D5h/1ZqFSTEFHdM65lR7RoIqq3tBBYavsOXV84NoHXZ0AkPyqQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/nanoid": { + "version": "3.3.18", + "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.18.tgz", + "integrity": "sha512-DTg4MJbGMWkfi6VZFdNt2/caMbQy4Ou+Op/hJQvGEWcnVfoA1QA+xzRKAzw9jD6+GVOOeYr/mIcuDSdug6F6+w==", + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/ai" + } + ], + "license": "MIT", + "bin": { + "nanoid": "bin/nanoid.cjs" + }, + "engines": { + "node": "^10 || ^12 || ^13.7 || ^14 || >=15.0.1" + } + }, + "node_modules/nopt": { + "version": "10.0.1", + "resolved": "https://registry.npmjs.org/nopt/-/nopt-10.0.1.tgz", + "integrity": "sha512-df3sBr/6ax9hSGuC3CspvLlbnX8cP5L5nZwXF8cGN8l0zSWR6BvzmQ6jPUKjvo6+/xdpkNvEcucBNUdBeeV13g==", + "dev": true, + "license": "ISC", + "dependencies": { + "abbrev": "^5.0.0" + }, + "bin": { + "nopt": "bin/nopt.js" + }, + "engines": { + "node": "^22.22.2 || ^24.15.0 || >=26.0.0" + } + }, + "node_modules/nwsapi": { + "version": "2.2.27", + "resolved": "https://registry.npmjs.org/nwsapi/-/nwsapi-2.2.27.tgz", + "integrity": "sha512-gQPNF78qebCQ6tvVFBYrvJdBNOrYZm90ZlXgpIFm06p6qHDHq/XC4TnJftN6OMbxVE0UTBAoRgcsDeJBBooITw==", + "dev": true, + "license": "MIT" + }, + "node_modules/parse5": { + "version": "7.3.0", + "resolved": "https://registry.npmjs.org/parse5/-/parse5-7.3.0.tgz", + "integrity": "sha512-IInvU7fabl34qmi9gY8XOVxhYyMyuH2xUNpb2q8/Y+7552KlejkRvqvD19nMoUW/uQGGbqNpA6Tufu5FL5BZgw==", + "dev": true, + "license": "MIT", + "dependencies": { + "entities": "^6.0.0" + }, + "funding": { + "url": "https://github.com/inikulin/parse5?sponsor=1" + } + }, + "node_modules/parse5/node_modules/entities": { + "version": "6.0.1", + "resolved": "https://registry.npmjs.org/entities/-/entities-6.0.1.tgz", + "integrity": "sha512-aN97NXWF6AWBTahfVOIrB/NShkzi5H7F9r1s9mD3cDj4Ko5f2qhhVoYMibXF7GlLveb/D2ioWay8lxI97Ven3g==", + "dev": true, + "license": "BSD-2-Clause", + "engines": { + "node": ">=0.12" + }, + "funding": { + "url": "https://github.com/fb55/entities?sponsor=1" + } + }, + "node_modules/path-browserify": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/path-browserify/-/path-browserify-1.0.1.tgz", + "integrity": "sha512-b7uo2UCUOYZcnF/3ID0lulOJi/bafxa1xPe7ZPsammBSpjSWQkjNxlt635YGS2MiR9GjvuXCtz2emr3jbsz98g==", + "dev": true, + "license": "MIT" + }, + "node_modules/path-scurry": { + "version": "2.0.2", + "resolved": "https://registry.npmjs.org/path-scurry/-/path-scurry-2.0.2.tgz", + "integrity": "sha512-3O/iVVsJAPsOnpwWIeD+d6z/7PmqApyQePUtCndjatj/9I5LylHvt5qluFaBT3I5h3r1ejfR056c+FCv+NnNXg==", + "dev": true, + "license": "BlueOak-1.0.0", + "dependencies": { + "lru-cache": "^11.0.0", + "minipass": "^7.1.2" + }, + "engines": { + "node": "18 || 20 || >=22" + }, + "funding": { + "url": "https://github.com/sponsors/isaacs" + } + }, + "node_modules/path-scurry/node_modules/lru-cache": { + "version": "11.5.2", + "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-11.5.2.tgz", + "integrity": "sha512-4pfM1Ff0x50o0tQwb5ucw/RzNyD0/YJME6IVcStalZuMWxdt3sR3huStTtxz4PUmvZfRguvDejasvQ2kifR11g==", + "dev": true, + "license": "BlueOak-1.0.0", + "engines": { + "node": "20 || >=22" + } + }, + "node_modules/pathe": { + "version": "2.0.3", + "resolved": "https://registry.npmjs.org/pathe/-/pathe-2.0.3.tgz", + "integrity": "sha512-WUjGcAqP1gQacoQe+OBJsFA7Ld4DyXuUIjZ5cc75cLHvJ7dtNsTugphxIADwspS+AraAUePCKrSVtPLFj/F88w==", + "dev": true, + "license": "MIT" + }, + "node_modules/pathval": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/pathval/-/pathval-2.0.1.tgz", + "integrity": "sha512-//nshmD55c46FuFw26xV/xFAaB5HF9Xdap7HJBBnrKdAd6/GxDBaNA1870O79+9ueg61cZLSVc+OaFlfmObYVQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 14.16" + } + }, + "node_modules/picocolors": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/picocolors/-/picocolors-1.1.1.tgz", + "integrity": "sha512-xceH2snhtb5M9liqDsmEw56le376mTZkEX/jEb/RxNFyegNul7eNslCXP9FDj/Lcu0X8KEyMceP2ntpaHrDEVA==", + "license": "ISC" + }, + "node_modules/picomatch": { + "version": "4.0.7", + "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.7.tgz", + "integrity": "sha512-qcJu88Q2IWqJsDD529JKMdwGm/dvInW4HvQnRwiH9JtihJvzGOscDtHE3x1pBKeUOTysQ8kVmLnJ2kJu7yhcGA==", + "license": "MIT", + "engines": { + "node": ">=12" + }, + "funding": { + "url": "https://github.com/sponsors/jonschlinkert" + } + }, + "node_modules/postcss": { + "version": "8.5.26", + "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.26.tgz", + "integrity": "sha512-u82N74LFzG8ca+dD8puPnplTXoGH4fTPpVGuIbt36G3qvNlkvfD0lEAZSxaly3KX8TS/L1A1gsCEmvKmBcVbkQ==", + "funding": [ + { + "type": "opencollective", + "url": "https://opencollective.com/postcss/" + }, + { + "type": "tidelift", + "url": "https://tidelift.com/funding/github/npm/postcss" + }, + { + "type": "github", + "url": "https://github.com/sponsors/ai" + } + ], + "license": "MIT", + "dependencies": { + "nanoid": "^3.3.17", + "picocolors": "^1.1.1", + "source-map-js": "^1.2.1" + }, + "engines": { + "node": "^10 || ^12 || >=14" + } + }, + "node_modules/proto-list": { + "version": "1.2.4", + "resolved": "https://registry.npmjs.org/proto-list/-/proto-list-1.2.4.tgz", + "integrity": "sha512-vtK/94akxsTMhe0/cbfpR+syPuszcuwhqVjJq26CuNDgFGj682oRBXOP5MJpv2r7JtE8MsiepGIqvvOTBwn2vA==", + "dev": true, + "license": "ISC" + }, + "node_modules/punycode": { + "version": "2.3.1", + "resolved": "https://registry.npmjs.org/punycode/-/punycode-2.3.1.tgz", + "integrity": "sha512-vYt7UD1U9Wg6138shLtLOvdAu+8DsC/ilFtEVHcH+wydcSpNE20AfSOduf6MkRFahL5FY7X1oU7nKVZFtfq8Fg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6" + } + }, + "node_modules/rollup": { + "version": "4.63.1", + "resolved": "https://registry.npmjs.org/rollup/-/rollup-4.63.1.tgz", + "integrity": "sha512-3Df9jsstwhccuEfmAMi9l8XUh/GOkVObmFTU7CCVBysEbcOZLl84jCtaAZMcPiMz2EGKsATzQcU+Xr3n/wU6cg==", + "license": "MIT", + "dependencies": { + "@types/estree": "1.0.9" + }, + "bin": { + "rollup": "dist/bin/rollup" + }, + "engines": { + "node": ">=18.0.0", + "npm": ">=8.0.0" + }, + "optionalDependencies": { + "@napi-rs/lzma-linux-x64-gnu": "1.5.1", + "@rollup/rollup-android-arm-eabi": "4.63.1", + "@rollup/rollup-android-arm64": "4.63.1", + "@rollup/rollup-darwin-arm64": "4.63.1", + "@rollup/rollup-darwin-x64": "4.63.1", + "@rollup/rollup-freebsd-arm64": "4.63.1", + "@rollup/rollup-freebsd-x64": "4.63.1", + "@rollup/rollup-linux-arm-gnueabihf": "4.63.1", + "@rollup/rollup-linux-arm-musleabihf": "4.63.1", + "@rollup/rollup-linux-arm64-gnu": "4.63.1", + "@rollup/rollup-linux-arm64-musl": "4.63.1", + "@rollup/rollup-linux-loong64-gnu": "4.63.1", + "@rollup/rollup-linux-loong64-musl": "4.63.1", + "@rollup/rollup-linux-ppc64-gnu": "4.63.1", + "@rollup/rollup-linux-ppc64-musl": "4.63.1", + "@rollup/rollup-linux-riscv64-gnu": "4.63.1", + "@rollup/rollup-linux-riscv64-musl": "4.63.1", + "@rollup/rollup-linux-s390x-gnu": "4.63.1", + "@rollup/rollup-linux-x64-gnu": "4.63.1", + "@rollup/rollup-linux-x64-musl": "4.63.1", + "@rollup/rollup-openbsd-x64": "4.63.1", + "@rollup/rollup-openharmony-arm64": "4.63.1", + "@rollup/rollup-win32-arm64-msvc": "4.63.1", + "@rollup/rollup-win32-ia32-msvc": "4.63.1", + "@rollup/rollup-win32-x64-gnu": "4.63.1", + "@rollup/rollup-win32-x64-msvc": "4.63.1", + "fsevents": "~2.3.2" + } + }, + "node_modules/rrweb-cssom": { + "version": "0.8.0", + "resolved": "https://registry.npmjs.org/rrweb-cssom/-/rrweb-cssom-0.8.0.tgz", + "integrity": "sha512-guoltQEx+9aMf2gDZ0s62EcV8lsXR+0w8915TC3ITdn2YueuNjdAYh/levpU9nFaoChh9RUS5ZdQMrKfVEN9tw==", + "dev": true, + "license": "MIT" + }, + "node_modules/safer-buffer": { + "version": "2.1.2", + "resolved": "https://registry.npmjs.org/safer-buffer/-/safer-buffer-2.1.2.tgz", + "integrity": "sha512-YZo3K82SD7Riyi0E1EQPojLz7kpepnSQI9IyPbHHg1XXXevb5dJI7tpyN2ADxGcQbHG7vcyRHk0cbwqcQriUtg==", + "dev": true, + "license": "MIT" + }, + "node_modules/saxes": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/saxes/-/saxes-6.0.0.tgz", + "integrity": "sha512-xAg7SOnEhrm5zI3puOOKyy1OMcMlIJZYNJY7xLBwSze0UjhPLnWfj2GF2EpT0jmzaJKIWKHLsaSSajf35bcYnA==", + "dev": true, + "license": "ISC", + "dependencies": { + "xmlchars": "^2.2.0" + }, + "engines": { + "node": ">=v12.22.7" + } + }, + "node_modules/semver": { + "version": "7.8.5", + "resolved": "https://registry.npmjs.org/semver/-/semver-7.8.5.tgz", + "integrity": "sha512-Y7/KDsb8LjooZpwaqGyulO6DQlksgCncchHGk+sZIY4SBvUocMBEFH5Ur1fI4dV+Jvl0w6cjvucaIi40puRioA==", + "dev": true, + "license": "ISC", + "bin": { + "semver": "bin/semver.js" + }, + "engines": { + "node": ">=10" + } + }, + "node_modules/siginfo": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/siginfo/-/siginfo-2.0.0.tgz", + "integrity": "sha512-ybx0WO1/8bSBLEWXZvEd7gMW3Sn3JFlW3TvX1nREbDLRNQNaeNN8WK0meBwPdAaOI7TtRRRJn/Es1zhrrCHu7g==", + "dev": true, + "license": "ISC" + }, + "node_modules/source-map-js": { + "version": "1.2.1", + "resolved": "https://registry.npmjs.org/source-map-js/-/source-map-js-1.2.1.tgz", + "integrity": "sha512-UXWMKhLOwVKb728IUtQPXxfYU+usdybtUrK/8uGE8CQMvrhOpwvzDBwj0QhSL7MQc7vIsISBG8VQ8+IDQxpfQA==", + "license": "BSD-3-Clause", + "engines": { + "node": ">=0.10.0" + } + }, + "node_modules/stackback": { + "version": "0.0.2", + "resolved": "https://registry.npmjs.org/stackback/-/stackback-0.0.2.tgz", + "integrity": "sha512-1XMJE5fQo1jGH6Y/7ebnwPOBEkIEnT4QF32d5R1+VXdXveM0IBMJt8zfaxX1P3QhVwrYe+576+jkANtSS2mBbw==", + "dev": true, + "license": "MIT" + }, + "node_modules/std-env": { + "version": "3.10.0", + "resolved": "https://registry.npmjs.org/std-env/-/std-env-3.10.0.tgz", + "integrity": "sha512-5GS12FdOZNliM5mAOxFRg7Ir0pWz8MdpYm6AY6VPkGpbA7ZzmbzNcBJQ0GPvvyWgcY7QAhCgf9Uy89I03faLkg==", + "dev": true, + "license": "MIT" + }, + "node_modules/strip-literal": { + "version": "3.1.0", + "resolved": "https://registry.npmjs.org/strip-literal/-/strip-literal-3.1.0.tgz", + "integrity": "sha512-8r3mkIM/2+PpjHoOtiAW8Rg3jJLHaV7xPwG+YRGrv6FP0wwk/toTpATxWYOW0BKdWwl82VT2tFYi5DlROa0Mxg==", + "dev": true, + "license": "MIT", + "dependencies": { + "js-tokens": "^9.0.1" + }, + "funding": { + "url": "https://github.com/sponsors/antfu" + } + }, + "node_modules/symbol-tree": { + "version": "3.2.4", + "resolved": "https://registry.npmjs.org/symbol-tree/-/symbol-tree-3.2.4.tgz", + "integrity": "sha512-9QNk5KwDF+Bvz+PyObkmSYjI5ksVUYtjW7AU22r2NKcfLJcXp96hkDWU3+XndOsUb+AQ9QhfzfCT2O+CNWT5Tw==", + "dev": true, + "license": "MIT" + }, + "node_modules/tinybench": { + "version": "2.9.0", + "resolved": "https://registry.npmjs.org/tinybench/-/tinybench-2.9.0.tgz", + "integrity": "sha512-0+DUvqWMValLmha6lr4kD8iAMK1HzV0/aKnCtWb9v9641TnP/MFb7Pc2bxoxQjTXAErryXVgUOfv2YqNllqGeg==", + "dev": true, + "license": "MIT" + }, + "node_modules/tinyexec": { + "version": "0.3.2", + "resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-0.3.2.tgz", + "integrity": "sha512-KQQR9yN7R5+OSwaK0XQoj22pwHoTlgYqmUscPYoknOoWCWfj/5/ABTMRi69FrKU5ffPVh5QcFikpWJI/P1ocHA==", + "dev": true, + "license": "MIT" + }, + "node_modules/tinyglobby": { + "version": "0.2.17", + "resolved": "https://registry.npmjs.org/tinyglobby/-/tinyglobby-0.2.17.tgz", + "integrity": "sha512-wXR/dYpcqKmfWpEdZjiKJOwCNFndD0DMnrW/cYjVGttEkBfVgcLFHoNrlj47mjOVic9yyNu65alsgF4NQyTa2g==", + "license": "MIT", + "dependencies": { + "fdir": "^6.5.0", + "picomatch": "^4.0.4" + }, + "engines": { + "node": ">=12.0.0" + }, + "funding": { + "url": "https://github.com/sponsors/SuperchupuDev" + } + }, + "node_modules/tinypool": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/tinypool/-/tinypool-1.1.1.tgz", + "integrity": "sha512-Zba82s87IFq9A9XmjiX5uZA/ARWDrB03OHlq+Vw1fSdt0I+4/Kutwy8BP4Y/y/aORMo61FQ0vIb5j44vSo5Pkg==", + "dev": true, + "license": "MIT", + "engines": { + "node": "^18.0.0 || >=20.0.0" + } + }, + "node_modules/tinyrainbow": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/tinyrainbow/-/tinyrainbow-2.0.0.tgz", + "integrity": "sha512-op4nsTR47R6p0vMUUoYl/a+ljLFVtlfaXkLQmqfLR1qHma1h/ysYk4hEXZ880bf2CYgTskvTa/e196Vd5dDQXw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=14.0.0" + } + }, + "node_modules/tinyspy": { + "version": "4.0.4", + "resolved": "https://registry.npmjs.org/tinyspy/-/tinyspy-4.0.4.tgz", + "integrity": "sha512-azl+t0z7pw/z958Gy9svOTuzqIk6xq+NSheJzn5MMWtWTFywIacg2wUlzKFGtt3cthx0r2SxMK0yzJOR0IES7Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=14.0.0" + } + }, + "node_modules/tldts": { + "version": "6.1.86", + "resolved": "https://registry.npmjs.org/tldts/-/tldts-6.1.86.tgz", + "integrity": "sha512-WMi/OQ2axVTf/ykqCQgXiIct+mSQDFdH2fkwhPwgEwvJ1kSzZRiinb0zF2Xb8u4+OqPChmyI6MEu4EezNJz+FQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "tldts-core": "^6.1.86" + }, + "bin": { + "tldts": "bin/cli.js" + } + }, + "node_modules/tldts-core": { + "version": "6.1.86", + "resolved": "https://registry.npmjs.org/tldts-core/-/tldts-core-6.1.86.tgz", + "integrity": "sha512-Je6p7pkk+KMzMv2XXKmAE3McmolOQFdxkKw0R8EYNr7sELW46JqnNeTX8ybPiQgvg1ymCoF8LXs5fzFaZvJPTA==", + "dev": true, + "license": "MIT" + }, + "node_modules/tough-cookie": { + "version": "5.1.2", + "resolved": "https://registry.npmjs.org/tough-cookie/-/tough-cookie-5.1.2.tgz", + "integrity": "sha512-FVDYdxtnj0G6Qm/DhNPSb8Ju59ULcup3tuJxkFb5K8Bv2pUXILbf0xZWU8PX8Ov19OXljbUyveOFwRMwkXzO+A==", + "dev": true, + "license": "BSD-3-Clause", + "dependencies": { + "tldts": "^6.1.32" + }, + "engines": { + "node": ">=16" + } + }, + "node_modules/tr46": { + "version": "5.1.1", + "resolved": "https://registry.npmjs.org/tr46/-/tr46-5.1.1.tgz", + "integrity": "sha512-hdF5ZgjTqgAntKkklYw0R03MG2x/bSzTtkxmIRw/sTNV8YXsCJ1tfLAX23lhxhHJlEf3CRCOCGGWw3vI3GaSPw==", + "dev": true, + "license": "MIT", + "dependencies": { + "punycode": "^2.3.1" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/typescript": { + "version": "5.9.3", + "resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz", + "integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==", + "devOptional": true, + "license": "Apache-2.0", + "bin": { + "tsc": "bin/tsc", + "tsserver": "bin/tsserver" + }, + "engines": { + "node": ">=14.17" + } + }, + "node_modules/undici-types": { + "version": "7.18.2", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-7.18.2.tgz", + "integrity": "sha512-AsuCzffGHJybSaRrmr5eHr81mwJU3kjw6M+uprWvCXiNeN9SOGwQ3Jn8jb8m3Z6izVgknn1R0FTCEAP2QrLY/w==", + "devOptional": true, + "license": "MIT" + }, + "node_modules/vite": { + "version": "7.3.6", + "resolved": "https://registry.npmjs.org/vite/-/vite-7.3.6.tgz", + "integrity": "sha512-4XP60spRGjSZFf1qYH+dJIkK2znL3zQfl9KkOV9MkkRR/3Dls0dxaBsQPTloEc5BLXWPL9vsOxopxyKoMmDueg==", + "license": "MIT", + "dependencies": { + "esbuild": "^0.27.0 || ^0.28.0", + "fdir": "^6.5.0", + "picomatch": "^4.0.3", + "postcss": "^8.5.6", + "rollup": "^4.43.0", + "tinyglobby": "^0.2.15" + }, + "bin": { + "vite": "bin/vite.js" + }, + "engines": { + "node": "^20.19.0 || >=22.12.0" + }, + "funding": { + "url": "https://github.com/vitejs/vite?sponsor=1" + }, + "optionalDependencies": { + "fsevents": "~2.3.3" + }, + "peerDependencies": { + "@types/node": "^20.19.0 || >=22.12.0", + "jiti": ">=1.21.0", + "less": "^4.0.0", + "lightningcss": "^1.21.0", + "sass": "^1.70.0", + "sass-embedded": "^1.70.0", + "stylus": ">=0.54.8", + "sugarss": "^5.0.0", + "terser": "^5.16.0", + "tsx": "^4.8.1", + "yaml": "^2.4.2" + }, + "peerDependenciesMeta": { + "@types/node": { + "optional": true + }, + "jiti": { + "optional": true + }, + "less": { + "optional": true + }, + "lightningcss": { + "optional": true + }, + "sass": { + "optional": true + }, + "sass-embedded": { + "optional": true + }, + "stylus": { + "optional": true + }, + "sugarss": { + "optional": true + }, + "terser": { + "optional": true + }, + "tsx": { + "optional": true + }, + "yaml": { + "optional": true + } + } + }, + "node_modules/vite-node": { + "version": "3.2.4", + "resolved": "https://registry.npmjs.org/vite-node/-/vite-node-3.2.4.tgz", + "integrity": "sha512-EbKSKh+bh1E1IFxeO0pg1n4dvoOTt0UDiXMd/qn++r98+jPO1xtJilvXldeuQ8giIB5IkpjCgMleHMNEsGH6pg==", + "dev": true, + "license": "MIT", + "dependencies": { + "cac": "^6.7.14", + "debug": "^4.4.1", + "es-module-lexer": "^1.7.0", + "pathe": "^2.0.3", + "vite": "^5.0.0 || ^6.0.0 || ^7.0.0-0" + }, + "bin": { + "vite-node": "vite-node.mjs" + }, + "engines": { + "node": "^18.0.0 || ^20.0.0 || >=22.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/vitest": { + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/vitest/-/vitest-3.2.7.tgz", + "integrity": "sha512-KrxIJ62Fd89gfysR4WotlgZABiz2dqFPgqGzX7s+CwsqLFomRH7777ZcrOD6+WVAh7khPQP41A+BKbpcJFrdEg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/chai": "^5.2.2", + "@vitest/expect": "3.2.7", + "@vitest/mocker": "3.2.7", + "@vitest/pretty-format": "^3.2.7", + "@vitest/runner": "3.2.7", + "@vitest/snapshot": "3.2.7", + "@vitest/spy": "3.2.7", + "@vitest/utils": "3.2.7", + "chai": "^5.2.0", + "debug": "^4.4.1", + "expect-type": "^1.2.1", + "magic-string": "^0.30.17", + "pathe": "^2.0.3", + "picomatch": "^4.0.2", + "std-env": "^3.9.0", + "tinybench": "^2.9.0", + "tinyexec": "^0.3.2", + "tinyglobby": "^0.2.14", + "tinypool": "^1.1.1", + "tinyrainbow": "^2.0.0", + "vite": "^5.0.0 || ^6.0.0 || ^7.0.0-0", + "vite-node": "3.2.4", + "why-is-node-running": "^2.3.0" + }, + "bin": { + "vitest": "vitest.mjs" + }, + "engines": { + "node": "^18.0.0 || ^20.0.0 || >=22.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + }, + "peerDependencies": { + "@edge-runtime/vm": "*", + "@types/debug": "^4.1.12", + "@types/node": "^18.0.0 || ^20.0.0 || >=22.0.0", + "@vitest/browser": "3.2.7", + "@vitest/ui": "3.2.7", + "happy-dom": "*", + "jsdom": "*" + }, + "peerDependenciesMeta": { + "@edge-runtime/vm": { + "optional": true + }, + "@types/debug": { + "optional": true + }, + "@types/node": { + "optional": true + }, + "@vitest/browser": { + "optional": true + }, + "@vitest/ui": { + "optional": true + }, + "happy-dom": { + "optional": true + }, + "jsdom": { + "optional": true + } + } + }, + "node_modules/vscode-uri": { + "version": "3.2.0", + "resolved": "https://registry.npmjs.org/vscode-uri/-/vscode-uri-3.2.0.tgz", + "integrity": "sha512-m2gXo3bn0G1kT9InzMf07fTbqMbGtyckj3bH5ktLO+1Ssv+yiATZ4dhwaQv9UZWxJh6E9IFGnQyjgWVDWVBDrg==", + "dev": true, + "license": "MIT" + }, + "node_modules/vue": { + "version": "3.5.42", + "resolved": "https://registry.npmjs.org/vue/-/vue-3.5.42.tgz", + "integrity": "sha512-4RyHQTbQvOPs3MfvUO1Sg0YRrKNnA0mAVtvpd12Tg1fKDN7OHBUl1IqSn8zGJjK9nI3NkNp8cgTpVrSZC5TTcA==", + "license": "MIT", + "dependencies": { + "@vue/compiler-dom": "3.5.42", + "@vue/compiler-sfc": "3.5.42", + "@vue/runtime-dom": "3.5.42", + "@vue/server-renderer": "3.5.42", + "@vue/shared": "3.5.42" + }, + "peerDependencies": { + "typescript": "*" + }, + "peerDependenciesMeta": { + "typescript": { + "optional": true + } + } + }, + "node_modules/vue-component-type-helpers": { + "version": "3.3.11", + "resolved": "https://registry.npmjs.org/vue-component-type-helpers/-/vue-component-type-helpers-3.3.11.tgz", + "integrity": "sha512-LwcxzeliO9fkQcpJG0PoX8X5kmAhKmH9wkpDLxNabwzkQ9Zeib2YVHwFV4pcWmMLfXVfjr/dSV+DaJ3cIPgSNA==", + "dev": true, + "license": "MIT" + }, + "node_modules/vue-tsc": { + "version": "3.3.11", + "resolved": "https://registry.npmjs.org/vue-tsc/-/vue-tsc-3.3.11.tgz", + "integrity": "sha512-gOb0B9rtU2+f1dszwPqSH5kAieIF9ReeLhD3kSRNHv5WZZUQz/JdVXW0RTdqhNTMlQkqKzrTTviqKr/4FYZraQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@volar/typescript": "2.4.28", + "@vue/language-core": "3.3.11" + }, + "bin": { + "vue-tsc": "bin/vue-tsc.js" + }, + "peerDependencies": { + "typescript": ">=5.0.0" + } + }, + "node_modules/w3c-xmlserializer": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/w3c-xmlserializer/-/w3c-xmlserializer-5.0.0.tgz", + "integrity": "sha512-o8qghlI8NZHU1lLPrpi2+Uq7abh4GGPpYANlalzWxyWteJOCsr/P+oPBA49TOLu5FTZO4d3F9MnWJfiMo4BkmA==", + "dev": true, + "license": "MIT", + "dependencies": { + "xml-name-validator": "^5.0.0" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/webidl-conversions": { + "version": "7.0.0", + "resolved": "https://registry.npmjs.org/webidl-conversions/-/webidl-conversions-7.0.0.tgz", + "integrity": "sha512-VwddBukDzu71offAQR975unBIGqfKZpM+8ZX6ySk8nYhVoo5CYaZyzt3YBvYtRtO+aoGlqxPg/B87NGVZ/fu6g==", + "dev": true, + "license": "BSD-2-Clause", + "engines": { + "node": ">=12" + } + }, + "node_modules/whatwg-encoding": { + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/whatwg-encoding/-/whatwg-encoding-3.1.1.tgz", + "integrity": "sha512-6qN4hJdMwfYBtE3YBTTHhoeuUrDBPZmbQaxWAqSALV/MeEnR5z1xd8UKud2RAkFoPkmB+hli1TZSnyi84xz1vQ==", + "deprecated": "Use @exodus/bytes instead for a more spec-conformant and faster implementation", + "dev": true, + "license": "MIT", + "dependencies": { + "iconv-lite": "0.6.3" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/whatwg-mimetype": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/whatwg-mimetype/-/whatwg-mimetype-4.0.0.tgz", + "integrity": "sha512-QaKxh0eNIi2mE9p2vEdzfagOKHCcj1pJ56EEHGQOVxp8r9/iszLUUV7v89x9O1p/T+NlTM5W7jW6+cz4Fq1YVg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + } + }, + "node_modules/whatwg-url": { + "version": "14.2.0", + "resolved": "https://registry.npmjs.org/whatwg-url/-/whatwg-url-14.2.0.tgz", + "integrity": "sha512-De72GdQZzNTUBBChsXueQUnPKDkg/5A5zp7pFDuQAj5UFoENpiACU0wlCvzpAGnTkj++ihpKwKyYewn/XNUbKw==", + "dev": true, + "license": "MIT", + "dependencies": { + "tr46": "^5.1.0", + "webidl-conversions": "^7.0.0" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/why-is-node-running": { + "version": "2.3.0", + "resolved": "https://registry.npmjs.org/why-is-node-running/-/why-is-node-running-2.3.0.tgz", + "integrity": "sha512-hUrmaWBdVDcxvYqnyh09zunKzROWjbZTiNy8dBEjkS7ehEDQibXJ7XvlmtbwuTclUiIyN+CyXQD4Vmko8fNm8w==", + "dev": true, + "license": "MIT", + "dependencies": { + "siginfo": "^2.0.0", + "stackback": "0.0.2" + }, + "bin": { + "why-is-node-running": "cli.js" + }, + "engines": { + "node": ">=8" + } + }, + "node_modules/ws": { + "version": "8.21.3", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.3.tgz", + "integrity": "sha512-201TZ/kPWxoPr/OKWjquZR1SWKXcvxdH+e1xrx89b3YbmzLMFCLfnaG1HFIgWzJOEWZ7MvpK++odZufgYR50Rw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=10.0.0" + }, + "peerDependencies": { + "bufferutil": "^4.0.1", + "utf-8-validate": ">=5.0.2" + }, + "peerDependenciesMeta": { + "bufferutil": { + "optional": true + }, + "utf-8-validate": { + "optional": true + } + } + }, + "node_modules/xml-name-validator": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/xml-name-validator/-/xml-name-validator-5.0.0.tgz", + "integrity": "sha512-EvGK8EJ3DhaHfbRlETOWAS5pO9MZITeauHKJyb8wyajUfQUenkIg2MvLDTZ4T/TgIcm3HU0TFBgWWboAZ30UHg==", + "dev": true, + "license": "Apache-2.0", + "engines": { + "node": ">=18" + } + }, + "node_modules/xmlchars": { + "version": "2.2.0", + "resolved": "https://registry.npmjs.org/xmlchars/-/xmlchars-2.2.0.tgz", + "integrity": "sha512-JZnDKK8B0RCDw84FNdDAIpZK+JuJw+s7Lz8nksI7SIuU3UXJJslUthsi+uWBUYOwPFwW7W7PRLRfUKpxjtjFCw==", + "dev": true, + "license": "MIT" + } + } +} diff --git a/aether-vscodex/web/package.json b/aether-vscodex/web/package.json new file mode 100644 index 000000000..abd573349 --- /dev/null +++ b/aether-vscodex/web/package.json @@ -0,0 +1,28 @@ +{ + "name": "@aether/vscodex-web", + "version": "0.4.0", + "private": true, + "type": "module", + "scripts": { + "dev": "vite", + "build": "vue-tsc -b && vite build", + "typecheck": "vue-tsc -b --pretty false", + "test": "vitest run" + }, + "dependencies": { + "@vitejs/plugin-vue": "^6.0.1", + "vite": "^7.1.3", + "vue": "^3.5.20" + }, + "devDependencies": { + "@types/node": "^24.3.0", + "@vue/test-utils": "^2.4.6", + "jsdom": "^26.1.0", + "typescript": "^5.9.2", + "vitest": "^3.2.4", + "vue-tsc": "^3.0.6" + }, + "engines": { + "node": ">=20" + } +} diff --git a/aether-vscodex/web/src/App.vue b/aether-vscodex/web/src/App.vue new file mode 100644 index 000000000..bcafc50d0 --- /dev/null +++ b/aether-vscodex/web/src/App.vue @@ -0,0 +1,7 @@ + + + diff --git a/aether-vscodex/web/src/components/CodexSurface.vue b/aether-vscodex/web/src/components/CodexSurface.vue new file mode 100644 index 000000000..76e4106f6 --- /dev/null +++ b/aether-vscodex/web/src/components/CodexSurface.vue @@ -0,0 +1,269 @@ + diff --git a/aether-vscodex/web/src/main.ts b/aether-vscodex/web/src/main.ts new file mode 100644 index 000000000..62e2f50f8 --- /dev/null +++ b/aether-vscodex/web/src/main.ts @@ -0,0 +1,52 @@ +import { createApp } from "vue"; + +import appRuntimeUrl from "../../public/app.js?url"; +import embedBridgeUrl from "../../public/embed-bridge.js?url"; +import i18nRuntimeUrl from "../../public/i18n.js?url"; +import "../../public/style.css"; +import App from "./App.vue"; +import { installRequestTemplate } from "./runtime/request-template"; + +type RuntimeAsset = { + id: string; + url: string; +}; + +const runtimeAssets: RuntimeAsset[] = [ + { id: "vscodex-i18n-runtime", url: i18nRuntimeUrl }, + { id: "vscodex-embed-bridge", url: embedBridgeUrl }, + { id: "vscodex-compat-runtime", url: appRuntimeUrl }, +]; + +function loadRuntimeAsset(asset: RuntimeAsset): Promise { + const existing = document.getElementById(asset.id) as HTMLScriptElement | null; + if (existing?.dataset.loaded === "true") return Promise.resolve(); + + return new Promise((resolve, reject) => { + const script = existing ?? document.createElement("script"); + script.id = asset.id; + script.async = false; + script.src = asset.url; + script.addEventListener("load", () => { + script.dataset.loaded = "true"; + resolve(); + }, { once: true }); + script.addEventListener("error", () => reject(new Error(`Unable to load ${asset.id}`)), { once: true }); + if (!existing) document.body.append(script); + }); +} + +async function startCompatibilityRuntime(): Promise { + for (const asset of runtimeAssets) await loadRuntimeAsset(asset); +} + +createApp(App).mount("#app"); +installRequestTemplate(); + +void startCompatibilityRuntime().catch((error: unknown) => { + const message = error instanceof Error ? error.message : String(error); + const status = document.getElementById("appState"); + if (status) status.textContent = message; + document.body.dataset.runtimeError = "true"; + console.error("Failed to start the Codex compatibility runtime", error); +}); diff --git a/aether-vscodex/web/src/runtime/request-template.ts b/aether-vscodex/web/src/runtime/request-template.ts new file mode 100644 index 000000000..92153b4b8 --- /dev/null +++ b/aether-vscodex/web/src/runtime/request-template.ts @@ -0,0 +1,34 @@ +export function installRequestTemplate(): HTMLTemplateElement { + const existing = document.getElementById("requestTemplate"); + if (existing instanceof HTMLTemplateElement) return existing; + + const template = document.createElement("template"); + template.id = "requestTemplate"; + template.innerHTML = ` +
+
+

+

+      
+ +
+ 查看请求数据 +

+      
+ +
+ + + +
+
+ `; + document.body.append(template); + return template; +} diff --git a/aether-vscodex/web/src/vite-env.d.ts b/aether-vscodex/web/src/vite-env.d.ts new file mode 100644 index 000000000..be2c61784 --- /dev/null +++ b/aether-vscodex/web/src/vite-env.d.ts @@ -0,0 +1,12 @@ +/// + +interface Window { + AetherVscodexEmbed?: { + active: boolean; + stop?: () => void; + }; + VscodexI18n?: { + locale: () => string; + setLocale: (locale: string, options?: { persist?: boolean }) => string; + }; +} diff --git a/aether-vscodex/web/tests/CodexSurface.test.ts b/aether-vscodex/web/tests/CodexSurface.test.ts new file mode 100644 index 000000000..d8a352778 --- /dev/null +++ b/aether-vscodex/web/tests/CodexSurface.test.ts @@ -0,0 +1,49 @@ +import { readFileSync } from "node:fs"; +import { resolve } from "node:path"; + +import { mount } from "@vue/test-utils"; +import { afterEach, describe, expect, it } from "vitest"; + +import CodexSurface from "../src/components/CodexSurface.vue"; +import { installRequestTemplate } from "../src/runtime/request-template"; + +afterEach(() => { + document.body.innerHTML = ""; +}); + +describe("CodexSurface", () => { + it("mounts the compatibility shell expected by the existing runtime", () => { + const wrapper = mount(CodexSurface, { attachTo: document.body }); + + expect(wrapper.find("#output").exists()).toBe(true); + expect(wrapper.find("#messageInput").attributes("contenteditable")).toBe("true"); + expect(wrapper.find("#sessionPicker").exists()).toBe(true); + expect(wrapper.find("#modelMenu").exists()).toBe(true); + expect(wrapper.find("#permissionMenu").exists()).toBe(true); + expect(wrapper.find("#requests").exists()).toBe(true); + expect(wrapper.find("#controlModeSwitch").attributes("data-mode")).toBe("sync"); + const controlModes = wrapper.findAll("#controlModeSwitch [data-control-mode]"); + expect(controlModes).toHaveLength(2); + expect(controlModes[0].attributes("aria-pressed")).toBe("true"); + expect(controlModes.every((button) => button.attributes("disabled") !== undefined)).toBe(true); + + wrapper.unmount(); + }); + + it("keeps every compatibility element from the legacy shell", () => { + mount(CodexSurface, { attachTo: document.body }); + installRequestTemplate(); + + const legacyHtml = readFileSync(resolve(process.cwd(), "../public/index.html"), "utf8"); + const legacyDocument = new DOMParser().parseFromString(legacyHtml, "text/html"); + const expected = [...legacyDocument.querySelectorAll("[id]")] + .map((element) => ({ id: element.id, tag: element.tagName, className: element.className })) + .sort((left, right) => left.id.localeCompare(right.id)); + const actual = [...document.querySelectorAll("[id]")] + .filter((element) => element.id !== "app") + .map((element) => ({ id: element.id, tag: element.tagName, className: element.className })) + .sort((left, right) => left.id.localeCompare(right.id)); + + expect(actual).toEqual(expected); + }); +}); diff --git a/aether-vscodex/web/tsconfig.app.json b/aether-vscodex/web/tsconfig.app.json new file mode 100644 index 000000000..568840ae7 --- /dev/null +++ b/aether-vscodex/web/tsconfig.app.json @@ -0,0 +1,19 @@ +{ + "compilerOptions": { + "tsBuildInfoFile": "./node_modules/.tmp/tsconfig.app.tsbuildinfo", + "target": "ES2022", + "useDefineForClassFields": true, + "module": "ESNext", + "lib": ["ES2022", "DOM", "DOM.Iterable"], + "skipLibCheck": true, + "moduleResolution": "Bundler", + "allowImportingTsExtensions": true, + "verbatimModuleSyntax": true, + "moduleDetection": "force", + "noEmit": true, + "strict": true, + "jsx": "preserve", + "types": ["vite/client", "vitest/globals"] + }, + "include": ["src/**/*.ts", "src/**/*.vue", "tests/**/*.ts"] +} diff --git a/aether-vscodex/web/tsconfig.json b/aether-vscodex/web/tsconfig.json new file mode 100644 index 000000000..1ffef600d --- /dev/null +++ b/aether-vscodex/web/tsconfig.json @@ -0,0 +1,7 @@ +{ + "files": [], + "references": [ + { "path": "./tsconfig.app.json" }, + { "path": "./tsconfig.node.json" } + ] +} diff --git a/aether-vscodex/web/tsconfig.node.json b/aether-vscodex/web/tsconfig.node.json new file mode 100644 index 000000000..506d2e826 --- /dev/null +++ b/aether-vscodex/web/tsconfig.node.json @@ -0,0 +1,17 @@ +{ + "compilerOptions": { + "tsBuildInfoFile": "./node_modules/.tmp/tsconfig.node.tsbuildinfo", + "target": "ES2023", + "lib": ["ES2023"], + "module": "ESNext", + "skipLibCheck": true, + "moduleResolution": "Bundler", + "allowImportingTsExtensions": true, + "verbatimModuleSyntax": true, + "moduleDetection": "force", + "noEmit": true, + "strict": true, + "types": ["node"] + }, + "include": ["vite.config.ts"] +} diff --git a/aether-vscodex/web/vite.config.ts b/aether-vscodex/web/vite.config.ts new file mode 100644 index 000000000..ade8df114 --- /dev/null +++ b/aether-vscodex/web/vite.config.ts @@ -0,0 +1,45 @@ +/// + +import { fileURLToPath, URL } from "node:url"; + +import vue from "@vitejs/plugin-vue"; +import { defineConfig } from "vite"; + +const relayTarget = "http://127.0.0.1:8787"; + +export default defineConfig({ + // Relative assets let the same build run at the local relay root and under + // Aether's /aether-vscodex/ static subpath. + base: "./", + plugins: [vue()], + resolve: { + alias: { + "@": fileURLToPath(new URL("./src", import.meta.url)), + }, + }, + server: { + fs: { + allow: [fileURLToPath(new URL("..", import.meta.url))], + }, + proxy: { + "/api": { + target: relayTarget, + changeOrigin: true, + }, + "/ws": { + target: relayTarget.replace("http", "ws"), + changeOrigin: true, + ws: true, + }, + }, + }, + build: { + outDir: "dist", + emptyOutDir: true, + assetsInlineLimit: 0, + }, + test: { + environment: "jsdom", + include: ["tests/**/*.test.ts"], + }, +}); diff --git a/apps/aether-gateway/Cargo.toml b/apps/aether-gateway/Cargo.toml index 2488d6e5f..ce44803eb 100644 --- a/apps/aether-gateway/Cargo.toml +++ b/apps/aether-gateway/Cargo.toml @@ -65,14 +65,14 @@ http.workspace = true http-body-util = "0.1" hyper = { version = "1", features = ["client", "server", "http1", "http2"] } hyper-util = { version = "0.1", features = ["client-legacy", "client-pool", "server-auto", "service", "tokio"] } -ldap3 = { version = "0.11", default-features = false, features = ["sync", "tls-rustls"] } +ldap3 = { version = "0.12.1", default-features = false, features = ["sync", "tls-rustls-ring"] } libc = "0.2" md-5 = "0.10" object_store.workspace = true parking_lot = "0.12" +percent-encoding.workspace = true regex.workspace = true reqwest.workspace = true -rsa = "0.9.10" rustls.workspace = true serde.workspace = true serde_json.workspace = true @@ -86,6 +86,7 @@ sysinfo = "0.32" thiserror.workspace = true tokio.workspace = true tokio-util.workspace = true +tokio-tungstenite = { version = "0.28", features = ["rustls-tls-webpki-roots"] } tower = { version = "0.5", features = ["util"] } tower-http = { version = "0.6", features = ["fs", "compression-gzip", "set-header"] } tracing.workspace = true @@ -102,4 +103,5 @@ tikv-jemalloc-sys = { version = "0.6", optional = true } [dev-dependencies] aether-test-support.workspace = true +aws-lc-rs.workspace = true tracing-subscriber.workspace = true diff --git a/apps/aether-gateway/examples/execution_runtime_harness.rs b/apps/aether-gateway/examples/execution_runtime_harness.rs index 994951c5c..7cc050093 100644 --- a/apps/aether-gateway/examples/execution_runtime_harness.rs +++ b/apps/aether-gateway/examples/execution_runtime_harness.rs @@ -30,7 +30,7 @@ struct Args { #[arg( long, env = "AETHER_EXECUTION_RUNTIME_UNIX_SOCKET", - default_value = "/tmp/aether-execution-runtime.sock" + default_value = "/tmp/aether-execution-runtime/aether-execution-runtime.sock" )] unix_socket: PathBuf, diff --git a/apps/aether-gateway/src/admin_api.rs b/apps/aether-gateway/src/admin_api.rs index 547fbaef3..f53d96c9e 100644 --- a/apps/aether-gateway/src/admin_api.rs +++ b/apps/aether-gateway/src/admin_api.rs @@ -1,17 +1,18 @@ pub(crate) use crate::handlers::admin::{ admin_provider_ops_local_action_response, admin_provider_pool_config, build_internal_control_error_response, create_provider_oauth_catalog_key, - find_duplicate_provider_oauth_key, maybe_build_local_admin_pool_response, - maybe_build_local_admin_response, persist_provider_quota_refresh_state, - provider_oauth_maintenance_endpoint_for_provider, provider_oauth_runtime_endpoint_for_provider, - provider_quota_refresh_endpoint_for_provider, provider_type_supports_quota_refresh, - reconcile_admin_fixed_provider_template_endpoints, + execute_admin_system_import_exclusively, find_duplicate_provider_oauth_key, + maybe_build_local_admin_pool_response, maybe_build_local_admin_response, + persist_provider_quota_refresh_state, provider_oauth_maintenance_endpoint_for_provider, + provider_oauth_runtime_endpoint_for_provider, provider_quota_refresh_endpoint_for_provider, + provider_type_supports_quota_refresh, reconcile_admin_fixed_provider_template_endpoints, refresh_provider_oauth_account_state_after_update, refresh_provider_pool_quota_locally, - store_admin_provider_ops_balance_cache, update_existing_provider_oauth_catalog_key, + release_admin_system_import_lease, store_admin_provider_ops_balance_cache, + try_acquire_admin_system_import_lease, update_existing_provider_oauth_catalog_key, AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError, AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult, - AdminStatsTimeRange, AdminStatsUsageFilter, OAUTH_ACCOUNT_BLOCK_PREFIX, - OAUTH_REQUEST_FAILED_PREFIX, + AdminStatsTimeRange, AdminStatsUsageFilter, AdminSystemImportLockError, SystemExportMode, + OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REQUEST_FAILED_PREFIX, }; use crate::handlers::admin::{ diff --git a/apps/aether-gateway/src/ai_serving/adaptation/private_envelope/sync.rs b/apps/aether-gateway/src/ai_serving/adaptation/private_envelope/sync.rs index ec2887590..3baae3caa 100644 --- a/apps/aether-gateway/src/ai_serving/adaptation/private_envelope/sync.rs +++ b/apps/aether-gateway/src/ai_serving/adaptation/private_envelope/sync.rs @@ -1,3 +1,4 @@ +use aether_usage_runtime::decode_internal_report_body_base64; use base64::Engine as _; use serde_json::Value; @@ -43,9 +44,8 @@ pub(crate) fn maybe_normalize_provider_private_sync_report_payload( } if let Some(body_base64) = payload.body_base64.as_deref() { - let body_bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let body_bytes = + decode_internal_report_body_base64(body_base64).map_err(GatewayError::Internal)?; let Some(normalized_bytes) = normalize_provider_private_stream_bytes(report_context, &body_bytes)? else { diff --git a/apps/aether-gateway/src/ai_serving/finalize/tests_sync.rs b/apps/aether-gateway/src/ai_serving/finalize/tests_sync.rs index 8b346e5a6..cd34e1417 100644 --- a/apps/aether-gateway/src/ai_serving/finalize/tests_sync.rs +++ b/apps/aether-gateway/src/ai_serving/finalize/tests_sync.rs @@ -11,10 +11,10 @@ use super::{ convert_gemini_chat_response_to_openai_chat, convert_gemini_response_to_openai_responses, maybe_build_local_core_sync_finalize_response, }; -use crate::ai_serving::GatewayControlDecision; use crate::ai_serving::{ convert_openai_chat_response_to_openai_responses, - convert_openai_responses_response_to_openai_chat, + convert_openai_responses_response_to_openai_chat, openai_responses_message_item_id, + GatewayControlDecision, }; use crate::usage::GatewaySyncReportRequest; @@ -192,7 +192,7 @@ fn aggregates_openai_responses_stream_completed_event_to_final_response() { "output_text": "Hello", "output": [{ "type": "message", - "id": "resp_123_msg", + "id": openai_responses_message_item_id("resp_123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -843,7 +843,7 @@ fn converts_claude_cli_response_to_openai_responses_response() { "output_text": "Hello Claude CLI", "output": [{ "type": "message", - "id": "msg_cli_123_msg", + "id": openai_responses_message_item_id("msg_cli_123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -907,7 +907,7 @@ fn converts_claude_cli_tool_use_to_openai_responses_function_call() { "output": [ { "type": "message", - "id": "msg_cli_tool_123_msg", + "id": openai_responses_message_item_id("msg_cli_tool_123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -977,7 +977,7 @@ fn converts_gemini_cli_response_to_openai_responses_response() { "output_text": "Hello Gemini CLI", "output": [{ "type": "message", - "id": "resp_cli_123_msg", + "id": openai_responses_message_item_id("resp_cli_123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -1046,7 +1046,7 @@ fn converts_gemini_cli_function_call_to_openai_responses_function_call() { "output": [ { "type": "message", - "id": "resp_cli_tool_123_msg", + "id": openai_responses_message_item_id("resp_cli_tool_123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -1252,7 +1252,7 @@ fn local_finalize_handles_openai_responses_openai_family_sync_response_even_when "model": "gpt-5", "output": [{ "type": "message", - "id": "resp_cli_family_123_msg", + "id": openai_responses_message_item_id("resp_cli_family_123", 0), "role": "assistant", "status": "completed", "content": [{ diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs index 3aff505cf..8e1b740f0 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs @@ -9,7 +9,7 @@ use aether_ai_serving::{ use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem}; use aether_routing_core::{ rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts, - RoutingCandidateTrace, RoutingDecisionTrace, + RoutingCandidateTrace, RoutingDecisionTrace, RoutingExecutionPolicy, }; use aether_scheduler_core::{ ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate, SchedulerRankingOutcome, @@ -47,13 +47,12 @@ use crate::cache::{ use crate::clock::current_unix_ms; use crate::dispatch::refs::dispatch_ref_for_local_candidate; use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value; -use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity}; +use crate::orchestration::{ExecutionAttemptIdentity, POOL_KEY_RETRY_INDEX_STRIDE}; use crate::scheduler::candidate::is_auth_api_key_concurrency_limit_skip_reason; use crate::scheduler::config::SchedulerSchedulingMode; use crate::stage_metrics::observe_gateway_stage_ms; use crate::{AppState, GatewayError}; -const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100; const AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET: Duration = Duration::from_millis(100); const AUTH_API_KEY_CONCURRENCY_RETRY_DELAY: Duration = Duration::from_millis(10); @@ -80,6 +79,13 @@ type DecorateSkippedCandidateFn<'a> = Arc< pub(crate) trait LocalExecutionAttemptSource: Send { async fn next_execution_attempt(&mut self) -> Result, GatewayError>; + /// Returns the request-scoped execution behaviour selected by routing. + /// Execution wrappers use this snapshot before consuming the first + /// attempt, avoiding a second lookup against mutable system settings. + fn routing_execution_policy(&self) -> Option { + None + } + async fn drain_execution_attempts(&mut self) -> Result, GatewayError>; async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError>; @@ -481,10 +487,6 @@ where type ExtraData = Value; type Error = Infallible; - fn attempt_slot_count(&self, candidate: &Self::Candidate) -> u32 { - local_attempt_slot_count(&candidate.transport) - } - fn build_extra_data(&self, candidate: &Self::Candidate) -> Option { available_candidate_extra_data_with_dispatch_ref(candidate, &self.build_extra_data) } @@ -1242,9 +1244,7 @@ async fn scheduler_cache_affinity_enabled( state: PlannerAppState<'_>, routing_policy: Option<&ResolvedRoutingPolicy>, ) -> bool { - scheduler_ordering_config_for_routing_policy(state, routing_policy) - .await - .scheduling_mode + scheduler_ordering_config_for_routing_policy(routing_policy).scheduling_mode == SchedulerSchedulingMode::CacheAffinity } @@ -1610,7 +1610,8 @@ async fn persist_available_local_execution_candidate_at_index( where F: Fn(&EligibleLocalExecutionCandidate) -> Option + Send + Sync, { - let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1); + // Exactly one attempt is materialized per candidate; same-key retries are + // derived lazily by the attempt loop after a failure. let extra_data = ai_candidate_extra_data_with_ranking( available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data), candidate.ranking.as_ref(), @@ -1625,53 +1626,34 @@ where Some(candidate_index), extra_data, ); - let should_persist = should_persist_available_local_candidate(&candidate); - let mut attempts = Vec::with_capacity(attempt_slots as usize); - let mut owned_candidate = Some(candidate); + let retry_index = effective_retry_index(0, candidate.orchestration.pool_key_index); + let generated_candidate_id = Uuid::new_v4().to_string(); + let candidate_id = if should_persist_available_local_candidate(&candidate) { + state + .persist_available_local_candidate( + trace_id, + context.user_id, + context.api_key_id, + &candidate.candidate, + candidate_index, + retry_index, + generated_candidate_id.as_str(), + context.required_capabilities, + extra_data, + current_unix_ms(), + context.error_context, + ) + .await + } else { + generated_candidate_id + }; - for retry_index in 0..attempt_slots { - let candidate_ref = owned_candidate - .as_ref() - .expect("candidate should remain available until final retry"); - let generated_candidate_id = Uuid::new_v4().to_string(); - let candidate_id = if should_persist { - state - .persist_available_local_candidate( - trace_id, - context.user_id, - context.api_key_id, - &candidate_ref.candidate, - candidate_index, - effective_retry_index(retry_index, candidate_ref.orchestration.pool_key_index), - generated_candidate_id.as_str(), - context.required_capabilities, - extra_data.clone(), - current_unix_ms(), - context.error_context, - ) - .await - } else { - generated_candidate_id - }; - - let candidate = if retry_index + 1 == attempt_slots { - owned_candidate - .take() - .expect("final retry should consume owned candidate") - } else { - candidate_ref.clone() - }; - let retry_index = - effective_retry_index(retry_index, candidate.orchestration.pool_key_index); - attempts.push(LocalExecutionCandidateAttempt { - eligible: candidate, - candidate_index, - retry_index, - candidate_id, - }); - } - - attempts + vec![LocalExecutionCandidateAttempt { + eligible: candidate, + candidate_index, + retry_index, + candidate_id, + }] } fn available_candidate_extra_data_with_dispatch_ref( @@ -1840,6 +1822,7 @@ fn routing_trace_for_candidate( CandidateKind::Provider => Some(candidate.key_id.clone()), CandidateKind::PoolGroup => None, }, + api_format: Some(candidate.endpoint_api_format.clone()), provider_priority: candidate.provider_priority, key_priority: candidate .key_global_priority_for_format @@ -1923,32 +1906,15 @@ fn build_unpersisted_local_execution_candidate_attempts( candidate: EligibleLocalExecutionCandidate, candidate_index: u32, ) -> VecDeque { - let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1); - let mut attempts = VecDeque::with_capacity(attempt_slots as usize); - let mut owned_candidate = Some(candidate); - - for retry_index in 0..attempt_slots { - let candidate = if retry_index + 1 == attempt_slots { - owned_candidate - .take() - .expect("final retry should consume owned candidate") - } else { - owned_candidate - .as_ref() - .expect("candidate should remain available until final retry") - .clone() - }; - let retry_index = - effective_retry_index(retry_index, candidate.orchestration.pool_key_index); - attempts.push_back(LocalExecutionCandidateAttempt { - eligible: candidate, - candidate_index, - retry_index, - candidate_id: Uuid::new_v4().to_string(), - }); - } - - attempts + // One attempt per candidate; same-key retries are derived lazily by the + // attempt loop after a failure. + let retry_index = effective_retry_index(0, candidate.orchestration.pool_key_index); + VecDeque::from([LocalExecutionCandidateAttempt { + eligible: candidate, + candidate_index, + retry_index, + candidate_id: Uuid::new_v4().to_string(), + }]) } async fn persist_pool_group_exhaustion_skipped_candidate( @@ -2276,6 +2242,8 @@ mod tests { pool_key_index, pool_key_lease: None, scheduler_affinity_epoch: None, + // These tests cover persistence shape, not same-key retries. + sticky_key_attempts: Some(1), }, ranking: None, } @@ -2357,16 +2325,11 @@ mod tests { assert_eq!(stored.len(), 1); assert_eq!(stored[0].key_id.as_deref(), Some("normal-key")); assert_eq!(stored[0].candidate_index, 2); - assert_eq!( - stored[0] - .extra_data - .as_ref() - .and_then(|value| value.get("dispatch_ref")) - .and_then(|value| value.get("SingleKey")) - .and_then(|value| value.get("key")) - .and_then(|value| value.get("key_id")), - Some(&json!("normal-key")) - ); + assert!(stored[0] + .extra_data + .as_ref() + .and_then(|value| value.get("dispatch_ref")) + .is_none()); } #[test] @@ -2514,14 +2477,23 @@ mod tests { assert!(should_cache_resolved_candidate_page(&cursor)); - let fixed_order_app = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::disabled().with_system_config_values_for_tests([( - "scheduling_mode".to_string(), - json!("fixed_order"), - )]), - ); + let fixed_order_app = AppState::new().expect("state should build"); + let fixed_order_policy = ResolvedRoutingPolicy { + group_id: Some("routing-group-fixed-order".to_string()), + group_version: Some(1), + selection_source: "test".to_string(), + requested_model: "gpt-5".to_string(), + resolved_model: "gpt-5".to_string(), + priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, + scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, + keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), + ranking_overlay: Default::default(), + mutation_plan: Default::default(), + pool_policy_overrides: Default::default(), + matched_rules: Vec::new(), + }; let mut page_cursor = LocalCandidatePreselectionPageCursor::new( PlannerAppState::new(&fixed_order_app), &model_directive_policy, @@ -2531,7 +2503,7 @@ mod tests { true, None, &auth_snapshot, - None, + Some(&fixed_order_policy), None, None, false, @@ -2549,7 +2521,7 @@ mod tests { auth_snapshot, client_session_affinity: None, required_capabilities: None, - routing_policy: None, + routing_policy: Some(fixed_order_policy), sticky_session_token: None, request_auth_channel: None, skipped_user_id: "user-1".to_string(), @@ -2647,16 +2619,11 @@ mod tests { ); assert_eq!(stored[1].key_id.as_deref(), Some("normal-key")); assert_eq!(stored[1].candidate_index, 1); - assert_eq!( - stored[1] - .extra_data - .as_ref() - .and_then(|value| value.get("dispatch_ref")) - .and_then(|value| value.get("SingleKey")) - .and_then(|value| value.get("key")) - .and_then(|value| value.get("key_id")), - Some(&json!("normal-key")) - ); + assert!(stored[1] + .extra_data + .as_ref() + .and_then(|value| value.get("dispatch_ref")) + .is_none()); } #[test] @@ -2726,7 +2693,7 @@ mod tests { .as_ref() .and_then(serde_json::Value::as_object) .expect("ranking metadata should persist as object extra data"); - assert_eq!(extra_data.get("existing"), Some(&json!("value"))); + assert!(extra_data.get("existing").is_none()); assert_eq!( extra_data.get("ranking_mode"), Some(&json!("CacheAffinity")) @@ -2739,14 +2706,7 @@ mod tests { Some(&json!("cached_affinity")) ); assert_eq!(extra_data.get("demoted_by"), Some(&json!("cross_format"))); - assert_eq!( - extra_data - .get("dispatch_ref") - .and_then(|value| value.get("SingleKey")) - .and_then(|value| value.get("key")) - .and_then(|value| value.get("key_id")), - Some(&json!("ranked-key")) - ); + assert!(extra_data.get("dispatch_ref").is_none()); } #[tokio::test] @@ -3084,7 +3044,7 @@ mod tests { .as_ref() .and_then(serde_json::Value::as_object) .expect("skipped ranking metadata should persist"); - assert_eq!(extra_data.get("existing"), Some(&json!("value"))); + assert!(extra_data.get("existing").is_none()); assert_eq!( extra_data.get("ranking_mode"), Some(&json!("CacheAffinity")) diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_metadata.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_metadata.rs index 699941e8a..3a0482013 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_metadata.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_metadata.rs @@ -278,13 +278,21 @@ mod tests { assert_eq!(metadata["transport_diagnostics"]["provider_type"], "codex"); assert_eq!( - metadata["transport_diagnostics"]["fingerprint"]["transport_profile"]["profile_id"], - "chrome_136" + metadata["transport_diagnostics"]["key_fingerprint_configured"], + Value::Bool(true) + ); + assert_eq!( + metadata["transport_diagnostics"]["key_transport_profile_configured"], + Value::Bool(true) ); assert_eq!( metadata["transport_diagnostics"]["resolved_transport_profile_id"], "chrome_136" ); + assert_eq!( + metadata["transport_diagnostics"]["resolved_transport_profile"]["profile_id"], + "chrome_136" + ); assert_eq!( metadata["transport_diagnostics"]["request_pair"]["conversion_enabled"], Value::Bool(true) diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs index dadf50839..a3840edc8 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs @@ -3,17 +3,14 @@ use aether_ai_serving::{ AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig, AiRankingSchedulingMode, }; -use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode}; +use aether_routing_core::ResolvedRoutingPolicy; use async_trait::async_trait; use tokio::sync::Mutex; -use tracing::warn; use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState}; use crate::clock::current_unix_ms; use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value; -use crate::scheduler::config::{ - read_scheduler_ordering_config, SchedulerOrderingConfig, SchedulerSchedulingMode, -}; +use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode}; use aether_scheduler_core::{ matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate, SchedulerPriorityMode, SchedulerRankableCandidate, @@ -133,7 +130,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates( required_capabilities: Option<&serde_json::Value>, routing_policy: Option<&ResolvedRoutingPolicy>, ) -> Vec { - let ordering_config = scheduler_ordering_config_for_routing_policy(state, routing_policy).await; + let ordering_config = scheduler_ordering_config_for_routing_policy(routing_policy); let port = GatewayLocalCandidateRankingPort { state, requested_model, @@ -184,35 +181,24 @@ fn ai_ranking_scheduling_mode(mode: SchedulerSchedulingMode) -> AiRankingSchedul } } -pub(crate) async fn scheduler_ordering_config_for_routing_policy( - state: PlannerAppState<'_>, +/// Return the immutable scheduler snapshot carried by a resolved routing +/// policy. A missing policy is a programming error in production request +/// paths; unit tests may use the scheduler default for isolated ranking tests. +pub(crate) fn scheduler_ordering_config_for_routing_policy( routing_policy: Option<&ResolvedRoutingPolicy>, ) -> SchedulerOrderingConfig { - let system_config = read_scheduler_ordering_config_or_default(state).await; match routing_policy { - Some(policy) => { - let mut config = scheduler_ordering_config_from_routing_policy(policy); - config.keep_priority_on_conversion |= system_config.keep_priority_on_conversion; - config + Some(policy) => SchedulerOrderingConfig::from_routing_policy(policy), + None => { + #[cfg(test)] + { + SchedulerOrderingConfig::default() + } + #[cfg(not(test))] + { + panic!("resolved routing policy is required before candidate scheduling") + } } - None => system_config, - } -} - -fn scheduler_ordering_config_from_routing_policy( - policy: &ResolvedRoutingPolicy, -) -> SchedulerOrderingConfig { - SchedulerOrderingConfig { - priority_mode: match policy.priority_mode { - RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider, - RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey, - }, - scheduling_mode: match policy.scheduling_mode { - RoutingSchedulingMode::FixedOrder => SchedulerSchedulingMode::FixedOrder, - RoutingSchedulingMode::CacheAffinity => SchedulerSchedulingMode::CacheAffinity, - RoutingSchedulingMode::LoadBalance => SchedulerSchedulingMode::LoadBalance, - }, - keep_priority_on_conversion: policy.keep_priority_on_conversion, } } @@ -231,37 +217,32 @@ fn routing_overlaid_candidate( let overlaid_key_priority = match kind { LocalExecutionCandidateKind::SingleKey => policy .ranking_overlay - .key_priority_overrides - .get(candidate.key_id.as_str()), + .key_priority_override_matching_format(candidate.key_id.as_str(), |format| { + crate::ai_serving::api_format_alias_matches( + format, + candidate.endpoint_api_format.as_str(), + ) + }) + .or_else(|| { + policy + .ranking_overlay + .key_priority_overrides + .get(candidate.key_id.as_str()) + .copied() + }), LocalExecutionCandidateKind::PoolGroup => policy .ranking_overlay .pool_priority_overrides - .get(candidate.provider_id.as_str()), + .get(candidate.provider_id.as_str()) + .copied(), }; - if let Some(overlaid_key_priority) = overlaid_key_priority.copied() { + if let Some(overlaid_key_priority) = overlaid_key_priority { overlaid.key_internal_priority = overlaid_key_priority; overlaid.key_global_priority_for_format = Some(overlaid_key_priority); } overlaid } -async fn read_scheduler_ordering_config_or_default( - state: PlannerAppState<'_>, -) -> SchedulerOrderingConfig { - match read_scheduler_ordering_config(state.app()).await { - Ok(config) => config, - Err(error) => { - warn!( - event_name = "planner_scheduler_ordering_config_load_failed", - log_type = "event", - error = ?error, - "failed to load scheduler ordering config while ranking local execution candidates" - ); - SchedulerOrderingConfig::default() - } - } -} - #[cfg(test)] mod tests { use std::collections::BTreeMap; @@ -270,10 +251,17 @@ mod tests { use aether_ai_serving::{ ai_ranking_context, build_ai_rankable_candidate, AiRankableCandidateParts, }; - use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::{ + provider_catalog::InMemoryProviderCatalogReadRepository, + routing_profiles::InMemoryRoutingGroupRepository, + }; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; + use aether_data_contracts::repository::routing_profiles::{ + CreateRoutingGroupRecord, RoutingGroupWriteRepository, + }; use aether_scheduler_core::{ apply_scheduler_candidate_ranking, build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, @@ -303,7 +291,11 @@ mod tests { required_capabilities: Option<&serde_json::Value>, ) -> Vec { let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase(); - let ordering_config = super::read_scheduler_ordering_config_or_default(state).await; + let ordering_config = + crate::scheduler::config::read_system_default_routing_ordering_config(state.app()) + .await + .expect("routing strategy should load") + .unwrap_or_default(); let mut candidates = candidates; let mut rankables = Vec::with_capacity(candidates.len()); let mut ordering_cache = CandidateTransportRankingFactsCache::default(); @@ -378,6 +370,8 @@ mod tests { priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: aether_routing_core::RankingOverlay::default(), mutation_plan: Default::default(), pool_policy_overrides: BTreeMap::new(), @@ -396,7 +390,7 @@ mod tests { } #[tokio::test] - async fn routing_policy_inherits_global_conversion_priority_override() { + async fn routing_policy_ignores_legacy_global_conversion_priority_override() { let data_state = GatewayDataState::default().with_system_config_values_for_tests([( "keep_priority_on_conversion".to_string(), json!(true), @@ -413,23 +407,24 @@ mod tests { priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: Default::default(), mutation_plan: Default::default(), pool_policy_overrides: Default::default(), matched_rules: Vec::new(), }; - let ordering = super::scheduler_ordering_config_for_routing_policy( - PlannerAppState::new(&state), - Some(&policy), - ) - .await; + let ordering = super::scheduler_ordering_config_for_routing_policy(Some(&policy)); assert_eq!( ordering.scheduling_mode, crate::scheduler::config::SchedulerSchedulingMode::FixedOrder ); - assert!(ordering.keep_priority_on_conversion); + assert!( + !ordering.keep_priority_on_conversion, + "a resolved routing policy must not inherit the legacy system-config flag" + ); } #[test] @@ -447,6 +442,8 @@ mod tests { priority_mode: aether_routing_core::RoutingSetPriorityMode::GlobalKey, scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: aether_routing_core::RankingOverlay { pool_priority_overrides: BTreeMap::from([("provider-1".to_string(), 4)]), key_priority_overrides: BTreeMap::from([("representative-key".to_string(), 1)]), @@ -570,6 +567,15 @@ mod tests { api_formats: Option, allowed_models: Option, ) -> StoredProviderCatalogKey { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key(provider_id, id, "plain-upstream-key") + .expect("api key should encrypt"); StoredProviderCatalogKey::new( id.to_string(), provider_id.to_string(), @@ -581,7 +587,7 @@ mod tests { .expect("key should build") .with_transport_fields( api_formats, - "plain-upstream-key".to_string(), + encrypted_api_key, None, None, Some(json!({"openai:chat": 1})), @@ -695,7 +701,7 @@ mod tests { let observed_at_unix_secs = current_unix_secs(); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ) .with_system_config_values_for_tests(vec![ ("provider_priority_mode".to_string(), json!("provider")), @@ -704,6 +710,7 @@ mod tests { serde_json::to_value(TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: "http://gateway-b:8080".to_string(), + tunnel_generation: "test-generation-remote".to_string(), conn_count: 1, observed_at_unix_secs, }) @@ -714,6 +721,7 @@ mod tests { serde_json::to_value(TunnelAttachmentRecord { gateway_instance_id: "gateway-a".to_string(), relay_base_url: "http://gateway-a:8080".to_string(), + tunnel_generation: "test-generation-local".to_string(), conn_count: 1, observed_at_unix_secs, }) @@ -772,7 +780,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -825,7 +833,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ) .with_system_config_values_for_tests(vec![( "scheduling_mode".to_string(), @@ -882,7 +890,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -918,7 +926,8 @@ mod tests { } #[tokio::test] - async fn local_execution_ranking_keeps_cross_format_priority_when_global_override_is_enabled() { + async fn local_execution_ranking_keeps_cross_format_priority_when_strategy_override_is_enabled() + { let provider_catalog = InMemoryProviderCatalogReadRepository::seed( vec![ sample_provider_with_options("provider-same", false, 10), @@ -933,14 +942,32 @@ mod tests { sample_key_for_provider("provider-cross", "key-cross", ""), ], ); + let routing_repository = std::sync::Arc::new(InMemoryRoutingGroupRepository::default()); + routing_repository + .create_routing_group(CreateRoutingGroupRecord { + id: "strategy-default".to_string(), + name: "strategy-default".to_string(), + description: None, + enabled: true, + is_system_default: true, + sort_order: 0, + config_json: json!({ + "default_policy": { + "keep_priority_on_conversion": true + } + }), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .expect("routing strategy should be created"); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ) - .with_system_config_values_for_tests(vec![( - "keep_priority_on_conversion".to_string(), - json!(true), - )]); + .with_routing_group_repository_for_tests(routing_repository); let state = AppState::new() .expect("state should build") .with_data_state_for_tests(data_state); @@ -1000,7 +1027,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ) .with_system_config_values_for_tests(vec![( "provider_priority_mode".to_string(), @@ -1066,7 +1093,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1119,7 +1146,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1193,7 +1220,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1273,7 +1300,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1349,7 +1376,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1416,7 +1443,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1499,7 +1526,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1564,7 +1591,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1653,7 +1680,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1739,7 +1766,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1836,7 +1863,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -1941,7 +1968,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -2035,7 +2062,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") 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 fdf7b00cf..abac4c3e6 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs @@ -23,7 +23,9 @@ use crate::ai_serving::{ use crate::orchestration::LocalExecutionCandidateMetadata; use crate::stage_metrics::observe_gateway_stage_ms; -use super::candidate_ranking::rank_eligible_local_execution_candidates; +use super::candidate_ranking::{ + rank_eligible_local_execution_candidates, scheduler_ordering_config_for_routing_policy, +}; #[derive(Debug, Clone, PartialEq)] pub(crate) struct EligibleLocalExecutionCandidate { @@ -378,8 +380,17 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion( "candidate_resolution_core", started_at.elapsed().as_millis() as u64, ); + let sticky_key_attempts = if outcome.eligible_candidates.is_empty() { + None + } else { + Some( + scheduler_ordering_config_for_routing_policy(routing_policy) + .sticky_key_attempts, + ) + }; for candidate in &mut outcome.eligible_candidates { candidate.orchestration.scheduler_affinity_epoch = Some(scheduler_affinity_epoch); + candidate.orchestration.sticky_key_attempts = sticky_key_attempts; } (outcome.eligible_candidates, outcome.skipped_candidates) } diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs index 391b7e36d..57a12e259 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs @@ -174,6 +174,9 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> { self.ranking_seed, false, self.request_operation, + super::candidate_ranking::scheduler_ordering_config_for_routing_policy( + self.routing_policy, + ), ) .await?; @@ -425,11 +428,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { ); let ordering_config = - super::candidate_ranking::scheduler_ordering_config_for_routing_policy( - state, - routing_policy, - ) - .await; + super::candidate_ranking::scheduler_ordering_config_for_routing_policy(routing_policy); Self { state, @@ -1291,6 +1290,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { .then_some(self.client_session_affinity.as_ref()) .flatten(), self.ranking_seed, + self.ordering_config, ) .await?; let skipped_candidates = skipped_candidates @@ -1473,6 +1473,7 @@ mod tests { use super::*; use crate::data::GatewayDataState; use crate::AppState; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::DataLayerError; @@ -1884,6 +1885,8 @@ mod tests { priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: Default::default(), mutation_plan: Default::default(), pool_policy_overrides: Default::default(), @@ -1947,6 +1950,8 @@ mod tests { priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: Default::default(), mutation_plan: Default::default(), pool_policy_overrides: Default::default(), @@ -2170,6 +2175,19 @@ mod tests { None, ) .expect("endpoint transport should build"); + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key( + row.provider_id.as_str(), + row.key_id.as_str(), + "plain-upstream-key", + ) + .expect("api key should encrypt"); let key = StoredProviderCatalogKey::new( row.key_id.clone(), row.provider_id.clone(), @@ -2181,7 +2199,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(serde_json::json!([row.endpoint_api_format.clone()])), - "plain-upstream-key".to_string(), + encrypted_api_key, None, None, None, @@ -2536,7 +2554,7 @@ mod tests { provider_repository, candidate_repository, ) - .with_encryption_key_for_tests("development-key"); + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); let app = AppState::new() .expect("gateway state should build") .with_data_state_for_tests(data_state); @@ -2656,15 +2674,17 @@ mod tests { provider_repository, candidate_repository, ) - .with_encryption_key_for_tests("development-key") + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + // Legacy keys deliberately disagree with the routing policy: the + // resolved policy must be the only source of scheduler ordering. .with_system_config_values_for_tests([ ( "scheduling_mode".to_string(), - serde_json::json!("fixed_order"), + serde_json::json!("cache_affinity"), ), ( "keep_priority_on_conversion".to_string(), - serde_json::json!(true), + serde_json::json!(false), ), ]); let app = AppState::new() @@ -2681,7 +2701,9 @@ mod tests { resolved_model: "gpt-5.4-mini".to_string(), priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, - keep_priority_on_conversion: false, + keep_priority_on_conversion: true, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: Default::default(), mutation_plan: Default::default(), pool_policy_overrides: Default::default(), diff --git a/apps/aether-gateway/src/ai_serving/planner/decision/stream.rs b/apps/aether-gateway/src/ai_serving/planner/decision/stream.rs index 4e1bd8f39..9c75fee92 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision/stream.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision/stream.rs @@ -11,6 +11,7 @@ use crate::ai_serving::planner::route::{ is_matching_stream_request, resolve_execution_runtime_stream_plan_kind, }; use crate::ai_serving::{resolve_decision_execution_runtime_auth_context, GatewayControlDecision}; +use crate::state::VideoTaskRouteAccess; use crate::{AiExecutionDecision, AppState, GatewayError}; pub(crate) async fn maybe_build_stream_decision_payload( @@ -155,16 +156,37 @@ async fn maybe_build_local_video_task_content_stream_decision_payload( return Ok(None); } - let _ = state - .hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path()) - .await?; + let Some(user_id) = decision + .auth_context + .as_ref() + .filter(|auth_context| auth_context.access_allowed) + .map(|auth_context| auth_context.user_id.trim()) + .filter(|value| !value.is_empty()) + else { + return Err(crate::video_tasks::not_found_error()); + }; + if state + .hydrate_video_task_for_route_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + user_id, + ) + .await? + != VideoTaskRouteAccess::Allowed + { + return Err(crate::video_tasks::not_found_error()); + } - let Some(action) = state.video_tasks.prepare_openai_content_stream_action( - parts.uri.path(), - parts.uri.query(), - trace_id, - ) else { - return Ok(None); + let Some(action) = state + .video_tasks + .prepare_openai_content_stream_action_for_user( + parts.uri.path(), + parts.uri.query(), + trace_id, + user_id, + ) + else { + return Err(crate::video_tasks::not_found_error()); }; let crate::video_tasks::LocalVideoTaskContentAction::StreamPlan(plan) = action else { diff --git a/apps/aether-gateway/src/ai_serving/planner/decision/sync.rs b/apps/aether-gateway/src/ai_serving/planner/decision/sync.rs index 1255f4e24..cb2422430 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision/sync.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision/sync.rs @@ -16,6 +16,7 @@ use crate::ai_serving::{ build_execution_runtime_auth_context, resolve_execution_runtime_auth_context, GatewayControlDecision, }; +use crate::state::VideoTaskRouteAccess; use crate::{AiExecutionDecision, AppState, GatewayError}; pub(crate) async fn maybe_build_sync_decision_payload( @@ -191,10 +192,6 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload( return Ok(None); } - let _ = state - .hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path()) - .await?; - let auth_context = resolve_execution_runtime_auth_context( state, decision, @@ -204,16 +201,30 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload( ) .await?; let Some(auth_context) = auth_context else { - return Ok(None); + return Err(crate::video_tasks::not_found_error()); }; - let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan( + if !auth_context.access_allowed || auth_context.user_id.trim().is_empty() { + return Err(crate::video_tasks::not_found_error()); + } + if state + .hydrate_video_task_for_route_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + &auth_context.user_id, + ) + .await? + != VideoTaskRouteAccess::Allowed + { + return Err(crate::video_tasks::not_found_error()); + } + let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan_for_user( plan_kind, parts.uri.path(), Some(body_json), Some(&auth_context), trace_id, ) else { - return Ok(None); + return Err(crate::video_tasks::not_found_error()); }; let aether_video_tasks_core::LocalVideoTaskFollowUpPlan { @@ -236,8 +247,7 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload( downstream_path = %parts.uri.path(), provider_api_format = %plan.provider_api_format, client_api_format = %plan.client_api_format, - upstream_base_url = ?upstream_base_url, - upstream_url = %plan.url, + upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), "gateway built local video follow-up sync decision payload" ); diff --git a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs index a7d5a13e4..efa915eb0 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs @@ -37,6 +37,10 @@ const ROUTING_GROUP_SELECTION_CACHE_TTL: Duration = Duration::from_secs(30); const ROUTING_GROUP_SELECTION_CACHE_STALE_TTL: Duration = Duration::from_secs(120); const CODEX_ACCOUNT_ID_HEADER: &str = "chatgpt-account-id"; const CODEX_FEDRAMP_HEADER: &str = "x-openai-fedramp"; +const INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE: &str = + "routing provider request violates provider contract"; +const INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE: &str = + "invalid provider request headers in routing mutation"; #[derive(Debug, Clone)] pub(crate) struct ResolvedLocalDecisionAuthInput { @@ -292,6 +296,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m crate::ai_serving::openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + provider_model.as_str(), ) }) .unwrap_or_default(); @@ -311,10 +316,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m ) } } - .map_err(|violation| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: format!("routing provider_request violates provider contract: {violation:?}"), - })?; + .map_err(|_| invalid_routing_provider_contract())?; } let provider_model = provider_request_body .get("model") @@ -649,21 +651,17 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()), )); } - None + return Err(routing_selection_error( + GatewayRoutingSelectionError::NoDefault, + )); } }; let Some((group_id, group_version, group_config_json, selection_source)) = selected_group else { - input.client_session_affinity = client_session_affinity_from_api_request( - client_api_format, - &parts.headers, - Some(body_json), - ); - input.routing_policy = None; - input.routing_trace_seed = None; - input.routing_context = None; - return Ok(()); + return Err(routing_selection_error( + GatewayRoutingSelectionError::NoDefault, + )); }; if try_attach_static_default_routing_policy_to_input( @@ -887,10 +885,36 @@ fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError GatewayRoutingSelectionError::Repository(message) => { GatewayError::Internal(format!("routing group repository lookup failed: {message}")) } - error => GatewayError::Client { - status: StatusCode::FORBIDDEN, - message: error.to_string(), + GatewayRoutingSelectionError::NoDefault => GatewayError::Client { + status: StatusCode::SERVICE_UNAVAILABLE, + message: "no enabled routing strategy is configured for this request".to_string(), }, + GatewayRoutingSelectionError::NotFound(_) => GatewayError::Client { + status: StatusCode::FORBIDDEN, + message: "requested routing group was not found".to_string(), + }, + GatewayRoutingSelectionError::Disabled(_) => GatewayError::Client { + status: StatusCode::FORBIDDEN, + message: "requested routing group is not enabled".to_string(), + }, + GatewayRoutingSelectionError::Forbidden(_) => GatewayError::Client { + status: StatusCode::FORBIDDEN, + message: "requested routing group is not allowed for this principal".to_string(), + }, + } +} + +fn invalid_routing_provider_contract() -> GatewayError { + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE.to_string(), + } +} + +fn invalid_routing_provider_headers() -> GatewayError { + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE.to_string(), } } @@ -945,14 +969,9 @@ fn btree_headers_to_header_map( ) -> Result { let mut output = HeaderMap::new(); for (name, value) in headers { - let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: format!("invalid provider request header name in routing mutation: {err}"), - })?; - let value = HeaderValue::from_str(value).map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: format!("invalid provider request header value in routing mutation: {err}"), - })?; + let name = HeaderName::from_bytes(name.as_bytes()) + .map_err(|_| invalid_routing_provider_headers())?; + let value = HeaderValue::from_str(value).map_err(|_| invalid_routing_provider_headers())?; output.insert(name, value); } Ok(output) @@ -1097,6 +1116,7 @@ fn ensure_report_context_routing_trace( endpoint_id: decision.endpoint_id.clone().unwrap_or_default(), model_id, key_id, + api_format: decision.provider_api_format.clone(), provider_priority, key_priority, }, @@ -1174,6 +1194,50 @@ mod tests { } } + #[test] + fn routing_selection_errors_do_not_echo_explicit_group() { + let secret = "private-group?token=Bearer-secret"; + + for error in [ + GatewayRoutingSelectionError::NotFound(secret.to_string()), + GatewayRoutingSelectionError::Disabled(secret.to_string()), + GatewayRoutingSelectionError::Forbidden(secret.to_string()), + ] { + let error = routing_selection_error(error); + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::FORBIDDEN, + ref message, + } if !message.contains(secret) + )); + } + } + + #[test] + fn routing_provider_errors_do_not_echo_dynamic_details() { + let secret = "https://internal.example/?token=Bearer-secret"; + let contract_error = invalid_routing_provider_contract(); + let header_error = btree_headers_to_header_map(&BTreeMap::from([( + format!("Authorization: {secret}"), + secret.to_string(), + )])) + .expect_err("invalid header should fail"); + + for (error, expected_message) in [ + (contract_error, INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE), + (header_error, INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE), + ] { + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + ref message, + } if message == expected_message && !message.contains(secret) + )); + } + } + #[tokio::test] async fn explicit_auth_snapshot_override_does_not_fall_back_to_the_planner_cache() { // AppState::new has no auth snapshot repository. Without the explicit @@ -1230,6 +1294,7 @@ mod tests { description: None, enabled: true, is_system_default: false, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/candidates.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/candidates.rs index f357405ab..99e2c7291 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/candidates.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/candidates.rs @@ -140,6 +140,9 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts( current_unix_secs(), false, spec.operation.map(|operation| operation.as_str()), + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await?; let outcome = materialize_local_execution_candidates_with_serving( @@ -246,6 +249,9 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a current_unix_secs(), false, spec.operation.map(|operation| operation.as_str()), + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await?; diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs index ed3797c25..e65590f58 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs @@ -203,6 +203,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_ client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: body_json .get("stream") .and_then(serde_json::Value::as_bool) 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 fbedcef7b..41008896e 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 @@ -183,6 +183,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts( let reasoning_replay_policy = openai_responses_reasoning_replay_policy( prepared.transport.provider.provider_type.as_str(), prepared.transport.endpoint.base_url.as_str(), + prepared.mapped_model.as_str(), ); let redaction = resolve_provider_chat_pii_redaction( state, diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs index fe4070b74..e4117cd37 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs @@ -26,6 +26,7 @@ use super::{ LocalSameFormatProviderCandidateAttemptSource, LocalSameFormatProviderDecisionInput, LocalSameFormatProviderSpec, }; +use aether_routing_core::RoutingExecutionPolicy; pub(crate) struct LocalSameFormatProviderSyncAttemptSource<'a> { state: &'a AppState, @@ -189,6 +190,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalSameFormatProviderSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { @@ -234,6 +242,13 @@ impl LocalExecutionAttemptSource for LocalSameFormatProviderSyncA impl LocalExecutionAttemptSource for LocalSameFormatProviderStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_stream_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/report_context.rs b/apps/aether-gateway/src/ai_serving/planner/report_context.rs index 7e3a29343..190ebd8fc 100644 --- a/apps/aether-gateway/src/ai_serving/planner/report_context.rs +++ b/apps/aether-gateway/src/ai_serving/planner/report_context.rs @@ -4,7 +4,7 @@ use aether_ai_serving::{ build_ai_execution_report_context, insert_provider_stream_event_api_format as insert_ai_provider_stream_event_api_format, provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type, - AiExecutionReportContextParts, AiRequestOrigin, + AiExecutionReportContextParts, AiRequestOrigin, STICKY_KEY_ATTEMPTS_REPORT_FIELD, }; use aether_routing_core::ResolvedRoutingPolicy; use aether_runtime_state::RuntimeLockLease; @@ -21,7 +21,8 @@ use crate::client_session_affinity::{ }; use crate::orchestration::{ insert_pool_key_lease_report_context_fields, ExecutionAttemptIdentity, - ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, + ROUTING_EXECUTION_POLICY_REPORT_FIELD, ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, + SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, }; use crate::scheduler::affinity::insert_scheduler_affinity_policy_report_context_field; @@ -59,6 +60,9 @@ pub(crate) struct LocalExecutionReportContextParts<'a> { pub(crate) client_session_affinity: Option<&'a ClientSessionAffinity>, pub(crate) routing_policy: Option<&'a ResolvedRoutingPolicy>, pub(crate) scheduler_affinity_epoch: Option, + /// Routing policy sticky-key attempt budget; read back by the attempt + /// loop to derive same-key retries lazily. + pub(crate) sticky_key_attempts: Option, pub(crate) client_requested_stream: bool, pub(crate) upstream_is_stream: bool, pub(crate) has_envelope: bool, @@ -72,10 +76,12 @@ pub(crate) fn build_local_execution_report_context( let RequestOrigin { client_ip, user_agent, + forwarded_headers_trusted, } = parts .request_origin .unwrap_or_else(|| request_origin_from_headers(parts.original_headers)); - let original_headers = crate::ai_serving::collect_control_headers(parts.original_headers); + let original_headers = + collect_report_context_original_headers(parts.original_headers, forwarded_headers_trusted); let original_request_body = crate::ai_serving::build_report_context_original_request_echo( parts.original_request_body_json, parts.original_request_body_base64, @@ -102,13 +108,20 @@ pub(crate) fn build_local_execution_report_context( value, ); } - if let Some(incoming_tls) = - crate::ai_serving::tls_fingerprint_from_headers(parts.original_headers) - { - merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls); + if forwarded_headers_trusted { + if let Some(incoming_tls) = + crate::ai_serving::tls_fingerprint_from_headers(parts.original_headers) + { + merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls); + } } insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease); insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy); + if let Some(policy) = parts.routing_policy { + if let Ok(value) = serde_json::to_value(policy.execution_policy) { + extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value); + } + } if let Some(override_policy) = parts .routing_policy .and_then(|policy| policy.pool_policy_overrides.get(parts.provider_id)) @@ -124,6 +137,12 @@ pub(crate) fn build_local_execution_report_context( Value::Number(epoch.into()), ); } + if let Some(sticky_key_attempts) = parts.sticky_key_attempts { + extra_fields.insert( + STICKY_KEY_ATTEMPTS_REPORT_FIELD.to_string(), + Value::Number(sticky_key_attempts.into()), + ); + } insert_request_path_fields( &mut extra_fields, parts.request_path, @@ -174,6 +193,17 @@ pub(crate) fn build_local_execution_report_context( }) } +fn collect_report_context_original_headers( + headers: &http::HeaderMap, + forwarded_headers_trusted: bool, +) -> BTreeMap { + let mut collected = crate::ai_serving::collect_control_headers(headers); + if !forwarded_headers_trusted { + collected.retain(|name, _| !name.starts_with("x-aether-tls-")); + } + collected +} + fn insert_request_path_fields( extra_fields: &mut Map, request_path: Option<&str>, @@ -243,8 +273,8 @@ mod tests { use serde_json::{json, Map, Value}; use super::{ - build_local_execution_report_context, provider_stream_event_api_format_for_provider_type, - LocalExecutionReportContextParts, + build_local_execution_report_context, collect_report_context_original_headers, + provider_stream_event_api_format_for_provider_type, LocalExecutionReportContextParts, }; use crate::ai_serving::ExecutionRuntimeAuthContext; use crate::ai_serving::RequestOrigin; @@ -274,6 +304,26 @@ mod tests { ); } + #[test] + fn untrusted_tls_forwarding_headers_are_excluded_from_report_context() { + let mut headers = http::HeaderMap::new(); + headers.insert("x-aether-tls-ja3", "spoofed-ja3".parse().unwrap()); + headers.insert(http::header::USER_AGENT, "test-client".parse().unwrap()); + + let untrusted = collect_report_context_original_headers(&headers, false); + assert!(!untrusted.contains_key("x-aether-tls-ja3")); + assert_eq!( + untrusted.get("user-agent").map(String::as_str), + Some("test-client") + ); + + let trusted = collect_report_context_original_headers(&headers, true); + assert_eq!( + trusted.get("x-aether-tls-ja3").map(String::as_str), + Some("spoofed-ja3") + ); + } + #[test] fn local_execution_report_context_records_request_origin_and_session_affinity() { let auth_context = ExecutionRuntimeAuthContext { @@ -324,12 +374,14 @@ mod tests { request_origin: Some(RequestOrigin { client_ip: Some("203.0.113.8".to_string()), user_agent: Some("Claude-Code/1.0".to_string()), + forwarded_headers_trusted: false, }), original_request_body_json: Some(&json!({"model": "gpt-5"})), original_request_body_base64: None, client_session_affinity: Some(&client_session_affinity), routing_policy: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, client_requested_stream: false, upstream_is_stream: false, has_envelope: false, @@ -413,6 +465,7 @@ mod tests { client_session_affinity: None, routing_policy: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, client_requested_stream: false, upstream_is_stream: true, has_envelope: false, @@ -474,12 +527,17 @@ mod tests { original_headers: &original_headers, request_path: None, request_query_string: None, - request_origin: None, + request_origin: Some(RequestOrigin { + client_ip: None, + user_agent: None, + forwarded_headers_trusted: true, + }), original_request_body_json: Some(&json!({"model": "gpt-5"})), original_request_body_base64: None, client_session_affinity: None, routing_policy: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, client_requested_stream: false, upstream_is_stream: false, has_envelope: false, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs index 2d8ec9a62..99524fda7 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs @@ -17,6 +17,7 @@ use crate::ai_serving::{ resolve_gemini_files_sync_spec as resolve_sync_spec, LocalGeminiFilesSpec, }; use crate::{AiExecutionDecision, AppState, GatewayError}; +use aether_routing_core::RoutingExecutionPolicy; use self::decision::maybe_build_local_gemini_files_decision_payload_for_candidate; use self::support::{ @@ -174,6 +175,13 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalGeminiFilesSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { @@ -212,6 +220,13 @@ impl LocalExecutionAttemptSource for LocalGeminiFilesSyncAttemptS #[async_trait] impl LocalExecutionAttemptSource for LocalGeminiFilesStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_stream_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs index 8f329cf8b..9dffd02df 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs @@ -109,6 +109,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: spec_metadata.require_streaming, upstream_is_stream: spec_metadata.require_streaming, has_envelope: false, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files/request.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files/request.rs index a2873c3af..9bcdcef71 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files/request.rs @@ -8,7 +8,10 @@ use crate::ai_serving::transport::{ GeminiFilesRequestBodyError, }; use crate::ai_serving::GEMINI_FILES_UPLOAD_PLAN_KIND; -use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot}; +use crate::ai_serving::{ + CandidateFailureDiagnostic, GatewayProviderTransportSnapshot, GEMINI_FILES_DELETE_PLAN_KIND, + GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, +}; use crate::AppState; use super::support::{ @@ -47,6 +50,26 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts( let transport = &attempt.eligible.transport; let effective_headers = input.effective_headers(&parts.headers); + if matches!( + spec_metadata.decision_kind, + GEMINI_FILES_GET_PLAN_KIND + | GEMINI_FILES_DELETE_PLAN_KIND + | GEMINI_FILES_DOWNLOAD_PLAN_KIND + ) && !candidate_matches_owned_gemini_file_mapping(state, parts, input, attempt).await + { + mark_skipped_local_gemini_files_candidate( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "gemini_file_mapping_mismatch", + ) + .await; + return None; + } + if let Some(skip_reason) = gemini_files_transport_unsupported_reason(transport, GEMINI_FILES_CANDIDATE_API_FORMAT) { @@ -191,3 +214,64 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts( file_name, }) } + +async fn candidate_matches_owned_gemini_file_mapping( + state: &AppState, + parts: &http::request::Parts, + input: &LocalGeminiFilesDecisionInput, + attempt: &LocalGeminiFilesCandidateAttempt, +) -> bool { + let Some(file_name) = normalize_gemini_file_name_from_path(parts.uri.path()) else { + return false; + }; + let user_id = input.auth_context.user_id.trim(); + if user_id.is_empty() || !state.has_gemini_file_mapping_data_reader() { + return false; + } + let Ok(Some(mapping)) = state + .find_active_gemini_file_mapping_for_owner( + file_name.as_str(), + &attempt.eligible.transport.key.id, + user_id, + crate::clock::current_unix_secs(), + ) + .await + else { + return false; + }; + + mapping.user_id.as_deref().map(str::trim) == Some(user_id) + && mapping.key_id == attempt.eligible.transport.key.id +} + +pub(crate) fn normalize_gemini_file_name_from_path(path: &str) -> Option { + let suffix = path.strip_prefix("/v1beta/files/")?.trim_matches('/'); + let suffix = suffix.strip_suffix(":download").unwrap_or(suffix).trim(); + let suffix = suffix.strip_prefix("files/").unwrap_or(suffix).trim(); + if suffix.is_empty() || suffix.contains('/') { + return None; + } + Some(format!("files/{suffix}")) +} + +#[cfg(test)] +mod tests { + use super::normalize_gemini_file_name_from_path; + + #[test] + fn normalizes_supported_gemini_file_object_paths() { + assert_eq!( + normalize_gemini_file_name_from_path("/v1beta/files/file-123"), + Some("files/file-123".to_string()) + ); + assert_eq!( + normalize_gemini_file_name_from_path("/v1beta/files/file-123:download"), + Some("files/file-123".to_string()) + ); + assert_eq!( + normalize_gemini_file_name_from_path("/v1beta/files/files/abc-123"), + Some("files/abc-123".to_string()) + ); + assert_eq!(normalize_gemini_file_name_from_path("/v1beta/files"), None); + } +} diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files/support.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files/support.rs index c2d736d61..4c90f7a72 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files/support.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files/support.rs @@ -108,6 +108,9 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts( Some(&input.auth_snapshot), input.client_session_affinity.as_ref(), current_unix_secs(), + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await?; let outcome = materialize_local_execution_candidates_with_serving( @@ -181,6 +184,9 @@ pub(super) async fn build_local_gemini_files_candidate_attempt_source<'a>( Some(&input.auth_snapshot), input.client_session_affinity.as_ref(), current_unix_secs(), + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await?; Ok(build_local_execution_candidate_attempt_source_with_serving( diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs index b02027337..c68a40fde 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs @@ -19,6 +19,7 @@ use crate::ai_serving::{ resolve_local_image_sync_spec as resolve_sync_spec, }; use crate::{AiExecutionDecision, AppState, GatewayError}; +use aether_routing_core::RoutingExecutionPolicy; use self::decision::maybe_build_local_openai_image_decision_payload_for_candidate; use self::support::{ @@ -252,6 +253,13 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiImageSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { @@ -290,6 +298,13 @@ impl LocalExecutionAttemptSource for LocalOpenAiImageSyncAttemptS #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiImageStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_stream_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs index ea13663db..c45cf08c8 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs @@ -122,6 +122,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: spec_metadata.require_streaming, upstream_is_stream, has_envelope: false, 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 37cc057cb..3ca11f0d2 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 @@ -9,6 +9,7 @@ use crate::ai_serving::planner::candidate_preparation::{ }; use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata; use crate::ai_serving::pure::normalize_openai_image_request_with_options; +use crate::ai_serving::transport::antigravity::is_antigravity_provider_transport; use crate::ai_serving::transport::{ build_grok_browser_headers, build_grok_upstream_url, build_openai_image_headers, build_openai_image_upstream_url, build_standard_provider_request_headers, @@ -338,6 +339,25 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts( let candidate = &attempt.eligible.candidate; let transport = &attempt.eligible.transport; let provider_api_format = "gemini:generate_content"; + + // The gemini:generate_content URL hook rewrites an Antigravity endpoint to + // /v1internal:, and this image path has no v1internal envelope to match it. + // Skip the candidate instead of posting a bare Gemini body that upstream + // would only reject. + if is_antigravity_provider_transport(transport) { + mark_skipped_local_openai_image_candidate( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "transport_unsupported", + ) + .await; + return None; + } + let effective_headers = input.effective_headers(&parts.headers); let prepared_candidate = match prepare_header_authenticated_candidate( diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image/support.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image/support.rs index 3a59fcf56..dc79c525d 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image/support.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image/support.rs @@ -127,6 +127,9 @@ pub(super) async fn list_local_openai_image_candidate_attempts( input.client_session_affinity.as_ref(), current_unix_secs(), false, + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await { @@ -201,6 +204,9 @@ pub(super) async fn build_local_openai_image_candidate_attempt_source<'a>( input.client_session_affinity.as_ref(), current_unix_secs(), false, + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await { diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs index 172bd0a69..daeee1e2b 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs @@ -16,6 +16,7 @@ use crate::ai_serving::{ LocalVideoCreateSpec, }; use crate::{AiExecutionDecision, AppState, GatewayError}; +use aether_routing_core::RoutingExecutionPolicy; use self::decision::maybe_build_local_video_create_decision_payload_for_candidate; use self::support::{ @@ -104,6 +105,13 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalVideoCreateSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { 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 f252354db..c1207443e 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 @@ -90,6 +90,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: false, upstream_is_stream: false, has_envelope: false, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/support.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/support.rs index ffaf3eedf..ddcc9d9c0 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/support.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/support.rs @@ -133,6 +133,9 @@ pub(super) async fn list_local_video_create_candidate_attempts( input.client_session_affinity.as_ref(), current_unix_secs(), false, + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await { @@ -190,6 +193,9 @@ pub(super) async fn build_local_video_create_candidate_attempt_source<'a>( input.client_session_affinity.as_ref(), current_unix_secs(), false, + crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy( + input.routing_policy.as_ref(), + ), ) .await { 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 405c8b138..860da8b16 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,7 +505,7 @@ 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"), - Some(&"codex_cli_rs/0.144.1".to_string()) + Some(&"codex_cli_rs/0.153.3".to_string()) ); assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); @@ -615,7 +615,7 @@ fn injects_only_codex_client_headers_for_images_requests() { ); assert_eq!( headers.get("user-agent"), - Some(&"codex_cli_rs/0.144.1".to_string()) + Some(&"codex_cli_rs/0.153.3".to_string()) ); assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); @@ -699,7 +699,7 @@ fn preserves_client_context_headers_and_enforces_codex_provider_identity() { ); assert_eq!( headers.get("user-agent"), - Some(&"codex_cli_rs/0.144.1".to_string()) + Some(&"codex_cli_rs/0.153.3".to_string()) ); assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert_eq!( @@ -763,7 +763,7 @@ 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"), - Some(&"codex_cli_rs/0.144.1".to_string()) + Some(&"codex_cli_rs/0.153.3".to_string()) ); assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); 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 14389089a..38fae7c71 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs @@ -15,11 +15,25 @@ pub(crate) fn is_deepseek_provider(provider_type: &str, base_url: &str) -> bool host == "deepseek.com" || host.ends_with(".deepseek.com") } +fn is_deepseek_model(provider_model: &str) -> bool { + let provider_model = provider_model.trim().to_ascii_lowercase(); + let leaf = provider_model + .rsplit(['/', ':']) + .next() + .unwrap_or(provider_model.as_str()); + leaf == "deepseek" || leaf.starts_with("deepseek-") || leaf.starts_with("deepseek_") +} + +fn is_deepseek_upstream(provider_type: &str, base_url: &str, provider_model: &str) -> bool { + is_deepseek_provider(provider_type, base_url) || is_deepseek_model(provider_model) +} + pub(crate) fn openai_responses_reasoning_replay_policy( provider_type: &str, base_url: &str, + provider_model: &str, ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { - if is_deepseek_provider(provider_type, base_url) { + if is_deepseek_upstream(provider_type, base_url, provider_model) { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque } else { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds @@ -33,7 +47,11 @@ pub(crate) fn apply_deepseek_tool_call_thinking_compat( provider_api_format: &str, original_request_body: Option<&Value>, ) { - if !is_deepseek_provider(provider_type, base_url) { + let provider_model = provider_request_body + .get("model") + .and_then(Value::as_str) + .unwrap_or_default(); + if !is_deepseek_upstream(provider_type, base_url, provider_model) { return; } @@ -302,11 +320,35 @@ mod tests { )); assert!(!is_deepseek_provider("custom", "ftp://api.deepseek.com/v1")); assert_eq!( - openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"), + openai_responses_reasoning_replay_policy( + "custom", + "https://api.deepseek.com/v1", + "deepseek-v4-flash", + ), crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque ); assert_eq!( - openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"), + openai_responses_reasoning_replay_policy( + "openai", + "https://api.openai.com/v1", + "gpt-5.6-sol", + ), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds + ); + assert_eq!( + openai_responses_reasoning_replay_policy( + "custom", + "https://api.b.ai/v1", + "deepseek-v4-flash", + ), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque + ); + assert_eq!( + openai_responses_reasoning_replay_policy( + "custom", + "https://api.b.ai/v1", + "not-deepseek-compatible", + ), crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds ); } @@ -330,8 +372,11 @@ mod tests { "input": reasoning_items.clone(), "future_request_field": {"preserve": true} }); - let replay_policy = - openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"); + let replay_policy = openai_responses_reasoning_replay_policy( + "custom", + "https://api.deepseek.com/v1", + "deepseek-v4-flash", + ); let mut provider_body = crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers_and_reasoning_replay_policy( &request, "openai:responses", @@ -373,7 +418,11 @@ mod tests { crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy( &mut deepseek, "openai:responses", - openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"), + openai_responses_reasoning_replay_policy( + "custom", + "https://api.deepseek.com/v1", + "deepseek-v4-flash", + ), ), 0 ); @@ -383,7 +432,11 @@ mod tests { crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy( &mut openai, "openai:responses", - openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"), + openai_responses_reasoning_replay_policy( + "openai", + "https://api.openai.com/v1", + "gpt-5.6-sol", + ), ), 66 ); @@ -417,6 +470,52 @@ mod tests { assert_eq!(body["messages"][1]["reasoning_content"], ""); } + #[test] + fn custom_relay_deepseek_model_adds_chat_thinking_compat() { + let mut body = json!({ + "model": "deepseek-v4-flash", + "messages": [ + {"role": "user", "content": "inspect the repository"}, + {"role": "assistant", "content": null, "tool_calls": [{ + "id": "call_1", + "type": "function", + "function": {"name": "inspect", "arguments": "{}"} + }]}, + {"role": "tool", "tool_call_id": "call_1", "content": "done"} + ] + }); + + apply_deepseek_tool_call_thinking_compat( + &mut body, + "custom", + "https://api.b.ai/v1", + "openai:chat", + None, + ); + + assert_eq!(body["thinking"]["type"], "enabled"); + assert_eq!(body["messages"][1]["reasoning_content"], ""); + } + + #[test] + fn custom_relay_non_deepseek_model_is_not_rewritten() { + let original = json!({ + "model": "not-deepseek-compatible", + "messages": [{"role": "assistant", "content": "done"}] + }); + let mut body = original.clone(); + + apply_deepseek_tool_call_thinking_compat( + &mut body, + "custom", + "https://api.b.ai/v1", + "openai:chat", + None, + ); + + assert_eq!(body, original); + } + #[test] fn openai_chat_deepseek_honors_disabled_thinking() { let original = json!({"reasoning_effort": "none"}); diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs b/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs index d2d599b49..384f8c928 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs @@ -18,6 +18,7 @@ use crate::ai_serving::planner::spec_metadata::{ }; use crate::ai_serving::GatewayControlDecision; use crate::{AiExecutionDecision, AppState, GatewayError}; +use aether_routing_core::RoutingExecutionPolicy; use super::candidates::{ build_local_standard_candidate_attempt_source, resolve_local_standard_decision_input, @@ -177,6 +178,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalStandardSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { @@ -220,6 +228,13 @@ impl LocalExecutionAttemptSource for LocalStandardSyncAttemptSour #[async_trait] impl LocalExecutionAttemptSource for LocalStandardStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_stream_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs index aabce02e3..bd0877b10 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs @@ -142,6 +142,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate( client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: body_json .get("stream") .and_then(serde_json::Value::as_bool) diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs b/apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs index 7b3f128c9..6769d0f97 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs @@ -4,6 +4,10 @@ use std::sync::Arc; use aether_contracts::ResolvedTransportProfile; use serde_json::Value; +use crate::ai_serving::planner::antigravity::{ + build_antigravity_v1internal_provider_request, AntigravityV1InternalRequestError, + AntigravityV1InternalRequestInput, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, +}; use crate::ai_serving::planner::candidate_preparation::{ prepare_header_authenticated_candidate, prepare_header_authenticated_candidate_from_auth, OauthPreparationContext, @@ -26,6 +30,7 @@ use crate::ai_serving::planner::standard::{ openai_provider_request_contract_failure_extra_data, openai_responses_reasoning_replay_policy, request_body_build_failure_extra_data, request_conversion_failure_extra_data, }; +use crate::ai_serving::transport::antigravity::is_antigravity_provider_transport; use crate::ai_serving::transport::kiro::{ build_kiro_provider_headers, build_kiro_provider_request_body, is_kiro_claude_messages_transport, KiroProviderHeadersInput, KiroRequestAuth, @@ -587,7 +592,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts( } }; crate::ai_serving::hydrate_openai_response_history( - state.runtime_state(), + state, body_json, spec_metadata.api_format, provider_api_format, @@ -597,6 +602,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts( let reasoning_replay_policy = openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + prepared_candidate.mapped_model.as_str(), ); let redaction = resolve_provider_chat_pii_redaction( state, @@ -836,6 +842,29 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts( .await); } + if normalized_provider_api_format == "gemini:generate_content" + && is_antigravity_provider_transport(transport) + { + return Ok(build_antigravity_cross_format_payload_parts( + state, + parts, + trace_id, + body_json, + input, + attempt, + transport, + spec_metadata.api_format, + provider_api_format, + prepared_candidate.mapped_model, + prepared_candidate.auth_header, + prepared_candidate.auth_value, + provider_request_body, + upstream_is_stream, + redaction.redacted, + ) + .await); + } + if normalized_provider_api_format == "gemini:generate_content" && is_gemini_cli_provider_transport(transport) { @@ -962,6 +991,145 @@ fn apply_transport_request_body_semantics( ) } +#[allow(clippy::too_many_arguments)] +async fn build_antigravity_cross_format_payload_parts( + state: &AppState, + parts: &http::request::Parts, + trace_id: &str, + original_body_json: &serde_json::Value, + input: &LocalStandardDecisionInput, + attempt: &LocalStandardCandidateAttempt, + transport: &Arc, + client_api_format: &str, + provider_api_format: &str, + mapped_model: String, + auth_header: String, + auth_value: String, + gemini_request_body: Value, + upstream_is_stream: bool, + request_redacted: bool, +) -> Option { + let candidate = &attempt.eligible.candidate; + let effective_headers = input.effective_headers(&parts.headers); + let resolved = + match build_antigravity_v1internal_provider_request(AntigravityV1InternalRequestInput { + state, + parts, + transport, + trace_id, + mapped_model: &mapped_model, + provider_api_format, + auth_header: &auth_header, + auth_value: &auth_value, + request_headers: effective_headers, + original_request_body: original_body_json, + gemini_request_body: &gemini_request_body, + upstream_is_stream, + same_format: false, + }) + .await + { + Ok(resolved) => resolved, + Err(AntigravityV1InternalRequestError::TransportUnsupported) => { + mark_skipped_local_standard_candidate( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "transport_unsupported", + ) + .await; + return None; + } + Err(AntigravityV1InternalRequestError::EnvelopeUnsupported) => { + mark_skipped_local_standard_candidate_with_extra_data( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "provider_request_body_build_failed", + request_body_build_failure_extra_data( + original_body_json, + client_api_format, + provider_api_format, + ), + ) + .await; + return None; + } + Err(AntigravityV1InternalRequestError::UpstreamUrlUnavailable) => { + mark_skipped_local_standard_candidate_with_failure_diagnostic( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "upstream_url_missing", + CandidateFailureDiagnostic::upstream_url_missing( + client_api_format, + provider_api_format, + "standard_family_antigravity_url", + ), + ) + .await; + return None; + } + Err(AntigravityV1InternalRequestError::HeaderRulesApplyFailed) => { + mark_skipped_local_standard_candidate_with_failure_diagnostic( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + "transport_header_rules_apply_failed", + CandidateFailureDiagnostic::header_rules_apply_failed( + client_api_format, + provider_api_format, + "standard_family_antigravity_headers", + ), + ) + .await; + return None; + } + }; + + let mut provider_request_headers = resolved.headers.headers; + apply_codex_openai_special_headers( + &mut provider_request_headers, + &resolved.body, + effective_headers, + resolved.transport.provider.provider_type.as_str(), + provider_api_format, + Some(trace_id), + resolved.transport.key.decrypted_auth_config.as_deref(), + ); + request_identity_response_encoding_when_redacted( + &mut provider_request_headers, + request_redacted, + ); + + Some(LocalStandardCandidatePayloadParts { + auth_header: resolved.headers.auth_header, + auth_value: resolved.headers.auth_value, + mapped_model, + provider_api_format: provider_api_format.to_string(), + provider_request_body: resolved.body, + provider_request_headers, + upstream_url: resolved.upstream_url, + upstream_is_stream, + envelope_name: Some(ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME), + transport: resolved.transport, + transport_profile: None, + request_redacted, + }) +} + #[allow(clippy::too_many_arguments)] async fn build_gemini_cli_cross_format_payload_parts( state: &AppState, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs index 4b6644845..496e57153 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs @@ -195,6 +195,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: body_json .get("stream") .and_then(serde_json::Value::as_bool) 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 920aa36b3..4bfa07dea 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 @@ -159,6 +159,7 @@ fn finalize_openai_chat_provider_request_body( openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + mapped_model, ), ) .err() @@ -2740,7 +2741,7 @@ mod tests { .provider_request_headers .get("x-client-version") .map(String::as_str), - Some("1.2.3") + Some("4.3.0") ); assert_eq!( payload @@ -2760,7 +2761,7 @@ mod tests { assert_eq!(payload.provider_request_body["model"], "gemini-2.5-pro"); assert_eq!( payload.provider_request_body["userAgent"], - "antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)" + "vscode/1.X.X (Antigravity/4.3.0)" ); assert_eq!(payload.provider_request_body["requestType"], "agent"); assert!(payload.provider_request_body.get("contents").is_none()); 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 b25a1be84..9873bcf9d 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,3 +1,4 @@ +use aether_routing_core::RoutingExecutionPolicy; use async_trait::async_trait; use std::collections::VecDeque; use tracing::warn; @@ -119,6 +120,13 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiChatStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { let select_started_at = std::time::Instant::now(); let selected = self.next_execution_attempt_with_target_select().await?; diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs index 812def7c6..6c54303d8 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs @@ -1,3 +1,4 @@ +use aether_routing_core::RoutingExecutionPolicy; use async_trait::async_trait; use tracing::warn; @@ -92,6 +93,13 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiChatSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/stream.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/stream.rs index bbda7d62a..3abb46dae 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/stream.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/stream.rs @@ -1,7 +1,6 @@ use std::collections::BTreeMap; use aether_contracts::RequestBody; -use tracing::debug; use super::super::{ augment_sync_report_context, build_ai_execution_plan_from_decision, @@ -10,7 +9,6 @@ use super::super::{ AiStreamAttempt, }; use crate::ai_serving::planner::common::enforce_provider_body_stream_policy; -use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log; use crate::ai_serving::provider_adaptation_requires_eventstream_accept; use crate::ai_serving::transport::{ build_standard_plan_fallback_headers, build_standard_plan_fallback_openai_chat_url, @@ -157,21 +155,16 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision( let Some(auth_pair) = take_ai_upstream_auth_pair(&mut payload) else { return Ok(None); }; - let (url, url_source) = if let Some(upstream_url) = - take_non_empty_string(&mut payload.upstream_url) - { - (upstream_url, "upstream_url") + let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) { + upstream_url } else { let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else { return Ok(None); }; - ( - build_standard_plan_fallback_openai_responses_url( - &upstream_base_url, - parts.uri.query(), - compact, - ), - "upstream_base_url", + build_standard_plan_fallback_openai_responses_url( + &upstream_base_url, + parts.uri.query(), + compact, ) }; let Some(provider_request_body_value) = payload.provider_request_body.take() else { @@ -238,16 +231,7 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision( .uri .query() .and_then(crate::ai_serving::api::sanitize_request_query_string); - let log_decision_upstream_base_url = payload - .upstream_base_url - .as_deref() - .map(sanitize_upstream_url_for_log); - let log_decision_upstream_url = payload - .upstream_url - .as_deref() - .map(sanitize_upstream_url_for_log); - let log_plan_url = sanitize_upstream_url_for_log(plan.url.as_str()); - debug!( + tracing::debug!( event_name = "local_openai_responses_stream_plan_built", log_type = "debug", request_id = %plan.request_id, @@ -255,12 +239,13 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision( provider_id = %plan.provider_id, endpoint_id = %plan.endpoint_id, key_id = %plan.key_id, + downstream_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query( + parts.uri.path(), + parts.uri.query(), + ).unwrap_or_else(|| "/".to_string()), + upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), downstream_path = %parts.uri.path(), downstream_query = ?log_downstream_query, - url_source, - decision_upstream_base_url = ?log_decision_upstream_base_url, - decision_upstream_url = ?log_decision_upstream_url, - plan_url = %log_plan_url, client_api_format = %plan.client_api_format, provider_api_format = %plan.provider_api_format, upstream_is_stream = effective_upstream_is_stream, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/sync.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/sync.rs index 44882aa11..e759144dc 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/sync.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/plan_builders/sync.rs @@ -1,7 +1,6 @@ use std::collections::BTreeMap; use aether_contracts::RequestBody; -use tracing::debug; use super::super::{ augment_sync_report_context, build_ai_execution_plan_from_decision, @@ -10,7 +9,6 @@ use super::super::{ AiSyncAttempt, }; use crate::ai_serving::planner::common::enforce_provider_body_stream_policy; -use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log; use crate::ai_serving::transport::{ build_standard_plan_fallback_headers, build_standard_plan_fallback_openai_chat_url, build_standard_plan_fallback_openai_responses_url, StandardPlanFallbackAcceptPolicy, @@ -142,21 +140,16 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision( let Some(auth_pair) = take_ai_upstream_auth_pair(&mut payload) else { return Ok(None); }; - let (url, url_source) = if let Some(upstream_url) = - take_non_empty_string(&mut payload.upstream_url) - { - (upstream_url, "upstream_url") + let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) { + upstream_url } else { let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else { return Ok(None); }; - ( - build_standard_plan_fallback_openai_responses_url( - &upstream_base_url, - parts.uri.query(), - compact, - ), - "upstream_base_url", + build_standard_plan_fallback_openai_responses_url( + &upstream_base_url, + parts.uri.query(), + compact, ) }; let Some(provider_request_body_value) = payload.provider_request_body.take() else { @@ -205,16 +198,7 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision( .uri .query() .and_then(crate::ai_serving::api::sanitize_request_query_string); - let log_decision_upstream_base_url = payload - .upstream_base_url - .as_deref() - .map(sanitize_upstream_url_for_log); - let log_decision_upstream_url = payload - .upstream_url - .as_deref() - .map(sanitize_upstream_url_for_log); - let log_plan_url = sanitize_upstream_url_for_log(plan.url.as_str()); - debug!( + tracing::debug!( event_name = "local_openai_responses_sync_plan_built", log_type = "debug", request_id = %plan.request_id, @@ -222,12 +206,13 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision( provider_id = %plan.provider_id, endpoint_id = %plan.endpoint_id, key_id = %plan.key_id, + downstream_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query( + parts.uri.path(), + parts.uri.query(), + ).unwrap_or_else(|| "/".to_string()), + upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), downstream_path = %parts.uri.path(), downstream_query = ?log_downstream_query, - url_source, - decision_upstream_base_url = ?log_decision_upstream_base_url, - decision_upstream_url = ?log_decision_upstream_url, - plan_url = %log_plan_url, client_api_format = %plan.client_api_format, provider_api_format = %plan.provider_api_format, upstream_is_stream = payload.upstream_is_stream, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs index dfc950f2d..6a8711140 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs @@ -3,7 +3,6 @@ use tracing::debug; use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision_with_websocket_mode; -use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log; use crate::ai_serving::planner::report_context::{ build_local_execution_report_context, insert_native_client_envelope_name, insert_provider_stream_event_api_format, LocalExecutionReportContextParts, @@ -184,6 +183,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand client_session_affinity: input.client_session_affinity.as_ref(), routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, + sticky_key_attempts: eligible.orchestration.sticky_key_attempts, client_requested_stream: body_json .get("stream") .and_then(serde_json::Value::as_bool) @@ -204,12 +204,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand &resolved.transport, ); - let log_base_url = sanitize_upstream_url_for_log(resolved.transport.endpoint.base_url.as_str()); let log_request_query = parts .uri .query() .and_then(crate::ai_serving::api::sanitize_request_query_string); - let log_upstream_url = sanitize_upstream_url_for_log(resolved.upstream_url.as_str()); debug!( event_name = "local_openai_responses_decision_payload_built", log_type = "debug", @@ -226,9 +224,12 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand client_api_format = spec_metadata.api_format, provider_api_format = %resolved.provider_api_format, request_path = %parts.uri.path(), + request_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query( + parts.uri.path(), + parts.uri.query(), + ).unwrap_or_else(|| "/".to_string()), + upstream_origin = %crate::handlers::shared::security_log_url_origin(&resolved.upstream_url), request_query = ?log_request_query, - upstream_base_url = %log_base_url, - upstream_url = %log_upstream_url, upstream_is_stream = resolved.upstream_is_stream, has_envelope = resolved.envelope_name.is_some(), "gateway built local openai responses decision payload" 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 8d09af2ae..a9170ce59 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 @@ -24,7 +24,6 @@ use crate::ai_serving::planner::gemini_cli::{ }; use crate::ai_serving::planner::redaction::{ request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction, - sanitize_upstream_url_for_log, }; use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata; use crate::ai_serving::planner::standard::{ @@ -428,7 +427,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_ } }; crate::ai_serving::hydrate_openai_response_history( - state.runtime_state(), + state, body_json, spec_metadata.api_format, provider_api_format, @@ -438,6 +437,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_ let reasoning_replay_policy = openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + mapped_model.as_str(), ); let redaction = resolve_provider_chat_pii_redaction( state, @@ -866,17 +866,10 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_ let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format); - let log_base_url = sanitize_upstream_url_for_log(transport.endpoint.base_url.as_str()); - let log_custom_path = transport - .endpoint - .custom_path - .as_deref() - .map(sanitize_upstream_url_for_log); let log_request_query = parts .uri .query() .and_then(crate::ai_serving::api::sanitize_request_query_string); - let log_upstream_url = sanitize_upstream_url_for_log(upstream_url.as_str()); debug!( event_name = "local_openai_responses_upstream_url_resolved", @@ -892,12 +885,14 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_ provider_api_format = %provider_api_format, execution_strategy = execution_strategy.as_str(), conversion_mode = conversion_mode.as_str(), - base_url = %log_base_url, - custom_path = ?log_custom_path, + request_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query( + parts.uri.path(), + parts.uri.query(), + ).unwrap_or_else(|| "/".to_string()), + upstream_origin = %crate::handlers::shared::security_log_url_origin(&upstream_url), request_path = %parts.uri.path(), request_query = ?log_request_query, mapped_model = %mapped_model, - upstream_url = %log_upstream_url, upstream_is_stream, "gateway resolved local openai responses upstream url" ); @@ -2010,8 +2005,6 @@ async fn build_kiro_openai_responses_payload_parts( }; let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(client_api_format, provider_api_format); - let log_upstream_url = sanitize_upstream_url_for_log(upstream_url.as_str()); - debug!( event_name = "local_openai_responses_kiro_upstream_url_resolved", log_type = "debug", @@ -2026,7 +2019,7 @@ async fn build_kiro_openai_responses_payload_parts( provider_api_format = %provider_api_format, execution_strategy = execution_strategy.as_str(), conversion_mode = conversion_mode.as_str(), - upstream_url = %log_upstream_url, + upstream_origin = %crate::handlers::shared::security_log_url_origin(&upstream_url), upstream_is_stream, "gateway resolved local openai responses kiro upstream url" ); 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 7b4b3f3d2..4d27a4b29 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 @@ -1005,6 +1005,7 @@ pub(crate) async fn maybe_build_responses_websocket_decision( reasoning_replay_policy: openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + mapped_model.as_str(), ), model_directive_patch: input .model_directive_policy diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs index 649f173b7..53cfd5d77 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs @@ -1,3 +1,4 @@ +use aether_routing_core::RoutingExecutionPolicy; use async_trait::async_trait; use tracing::warn; @@ -161,6 +162,13 @@ pub(super) async fn build_local_stream_attempt_source<'a>( #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiResponsesSyncAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_sync_attempt(attempt).await? { @@ -204,6 +212,13 @@ impl LocalExecutionAttemptSource for LocalOpenAiResponsesSyncAtte #[async_trait] impl LocalExecutionAttemptSource for LocalOpenAiResponsesStreamAttemptSource<'_> { + fn routing_execution_policy(&self) -> Option { + self.input + .routing_policy + .as_ref() + .map(|policy| policy.execution_policy) + } + async fn next_execution_attempt(&mut self) -> Result, GatewayError> { while let Some(attempt) = self.candidates.next_attempt().await? { match self.build_stream_attempt(attempt).await? { diff --git a/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs b/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs index 916307790..724e77a6e 100644 --- a/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs +++ b/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs @@ -8,9 +8,13 @@ use crate::constants::{ API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS, API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS, }; use crate::scheduler::candidate::SchedulerSkippedCandidate; +use crate::scheduler::config::SchedulerOrderingConfig; use crate::GatewayError; impl<'a> PlannerAppState<'a> { + /// `ordering_config` is the immutable scheduler snapshot derived from the + /// request's resolved routing policy. + #[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates( self, api_format: &str, @@ -21,6 +25,7 @@ impl<'a> PlannerAppState<'a> { client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { crate::scheduler::candidate::list_selectable_candidates( self.app().data.as_ref(), @@ -33,10 +38,12 @@ impl<'a> PlannerAppState<'a> { client_session_affinity, now_unix_secs, enable_model_directives, + ordering_config, ) .await } + #[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_with_skip_reasons( self, api_format: &str, @@ -47,6 +54,7 @@ impl<'a> PlannerAppState<'a> { client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -64,10 +72,12 @@ impl<'a> PlannerAppState<'a> { now_unix_secs, enable_model_directives, None, + ordering_config, ) .await } + #[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_operation( self, api_format: &str, @@ -79,6 +89,7 @@ impl<'a> PlannerAppState<'a> { now_unix_secs: u64, enable_model_directives: bool, request_operation: Option<&str>, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -103,6 +114,7 @@ impl<'a> PlannerAppState<'a> { attempt_now_unix_secs, enable_model_directives, request_operation, + ordering_config, ) .await?; @@ -123,6 +135,7 @@ impl<'a> PlannerAppState<'a> { } } + #[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons( self, api_format: &str, @@ -132,6 +145,7 @@ impl<'a> PlannerAppState<'a> { auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -148,10 +162,12 @@ impl<'a> PlannerAppState<'a> { auth_snapshot, client_session_affinity, now_unix_secs, + ordering_config, ) .await } + #[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_for_required_capability_without_requested_model( self, candidate_api_format: &str, @@ -160,6 +176,7 @@ impl<'a> PlannerAppState<'a> { auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS); let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1)); @@ -176,6 +193,7 @@ impl<'a> PlannerAppState<'a> { auth_snapshot, client_session_affinity, attempt_now_unix_secs, + ordering_config, ) .await?; diff --git a/apps/aether-gateway/src/ai_serving/pure/mod.rs b/apps/aether-gateway/src/ai_serving/pure/mod.rs index faaa16c2f..91115566b 100644 --- a/apps/aether-gateway/src/ai_serving/pure/mod.rs +++ b/apps/aether-gateway/src/ai_serving/pure/mod.rs @@ -177,6 +177,7 @@ pub(crate) use aether_ai_formats::{ api_format_defaults_to_client_error_failover, api_format_defaults_to_non_stream, api_format_permission_covers, codex_responses_lite_tool_is_client_executed, intersect_api_format_allowed_lists, is_embedding_api_format, is_rerank_api_format, + 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, ApiOperation, ClientSurface, diff --git a/apps/aether-gateway/src/ai_serving/response_history.rs b/apps/aether-gateway/src/ai_serving/response_history.rs index 2578d0b5e..5de43fd05 100644 --- a/apps/aether-gateway/src/ai_serving/response_history.rs +++ b/apps/aether-gateway/src/ai_serving/response_history.rs @@ -2,14 +2,15 @@ use crate::ai_serving::{ hydrate_response_history, normalize_api_format_alias, record_converted_response_history, response_history_is_loaded, response_history_storage_key, ResponseHistoryRecord, }; -use aether_runtime_state::RuntimeState; use serde_json::Value; use tracing::warn; -use crate::GatewayError; +use crate::{AppState, GatewayError}; + +const RESPONSE_HISTORY_SECRET_PURPOSE: &str = "openai-response-history"; pub(crate) async fn hydrate_openai_response_history( - runtime_state: &RuntimeState, + state: &AppState, request: &Value, client_api_format: &str, provider_api_format: &str, @@ -33,6 +34,7 @@ pub(crate) async fn hydrate_openai_response_history( } let storage_key = response_history_storage_key(previous_response_id, Some(history_scope)); + let runtime_state = state.runtime_state(); let payload = runtime_state.kv_get(&storage_key).await.map_err(|error| { warn!( event_name = "openai_response_history_read_failed", @@ -46,8 +48,24 @@ pub(crate) async fn hydrate_openai_response_history( let Some(payload) = payload else { return Ok(()); }; + let Some(payload) = crate::handlers::shared::open_runtime_secret_payload( + state, + RESPONSE_HISTORY_SECRET_PURPOSE, + &payload, + ) else { + let _ = runtime_state.kv_delete(&storage_key).await; + warn!( + event_name = "openai_response_history_decryption_failed", + log_type = "ops", + backend = runtime_state.backend_kind().as_str(), + "gateway rejected undecryptable shared OpenAI response history" + ); + return Err(GatewayError::Internal( + "OpenAI response history decryption failed".to_string(), + )); + }; if let Err(error) = - hydrate_response_history(previous_response_id, Some(history_scope), &payload) + hydrate_response_history(previous_response_id, Some(history_scope), payload.as_str()) { let _ = runtime_state.kv_delete(&storage_key).await; warn!( @@ -65,11 +83,25 @@ pub(crate) async fn hydrate_openai_response_history( } pub(crate) async fn persist_response_history_record( - runtime_state: &RuntimeState, + state: &AppState, record: ResponseHistoryRecord, ) { + let runtime_state = state.runtime_state(); + let Some(sealed_payload) = crate::handlers::shared::seal_runtime_secret_payload( + state, + RESPONSE_HISTORY_SECRET_PURPOSE, + &record.payload, + ) else { + warn!( + event_name = "openai_response_history_encryption_unavailable", + log_type = "ops", + backend = runtime_state.backend_kind().as_str(), + "gateway refused to persist unencrypted OpenAI response history" + ); + return; + }; if let Err(error) = runtime_state - .kv_set(&record.storage_key, record.payload, Some(record.ttl)) + .kv_set(&record.storage_key, sealed_payload, Some(record.ttl)) .await { warn!( @@ -83,7 +115,7 @@ pub(crate) async fn persist_response_history_record( } pub(crate) async fn persist_converted_response_history( - runtime_state: &RuntimeState, + state: &AppState, report_context: &Value, response: Option<&Value>, ) { @@ -91,6 +123,126 @@ pub(crate) async fn persist_converted_response_history( return; }; if let Some(record) = record_converted_response_history(report_context, response) { - persist_response_history_record(runtime_state, record).await; + persist_response_history_record(state, record).await; + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::time::{Duration, SystemTime, UNIX_EPOCH}; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; + use serde_json::json; + use sha2::{Digest, Sha256}; + + use super::{ + hydrate_openai_response_history, persist_response_history_record, ResponseHistoryRecord, + }; + use crate::{ai_serving::response_history_storage_key, data::GatewayDataState, AppState}; + + fn response_history_test_state() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_runtime_state(Arc::new(RuntimeState::memory( + MemoryRuntimeStateConfig::default(), + ))) + } + + fn response_history_payload(response_id: &str, scope: &str, marker: &str) -> String { + let expires_at_unix_secs = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() + .saturating_add(3600); + json!({ + "version": 1, + "response_id": response_id, + "scope_fingerprint": format!("{:x}", Sha256::digest(scope.trim().as_bytes())), + "expires_at_unix_secs": expires_at_unix_secs, + "transcript": [{"type": "message", "content": marker}], + }) + .to_string() + } + + #[tokio::test] + async fn response_history_is_encrypted_at_rest_and_hydrates() { + let state = response_history_test_state(); + let response_id = "resp_gateway_encrypted_history_v1"; + let scope = "response-history-encrypted-scope"; + let marker = "private-response-history-marker"; + let storage_key = response_history_storage_key(response_id, Some(scope)); + let payload = response_history_payload(response_id, scope, marker); + + persist_response_history_record( + &state, + ResponseHistoryRecord { + storage_key: storage_key.clone(), + payload, + ttl: Duration::from_secs(6 * 60 * 60), + }, + ) + .await; + + let stored = state + .runtime_kv_get(&storage_key) + .await + .expect("history lookup should succeed") + .expect("history should be persisted"); + assert!(crate::handlers::shared::runtime_secret_payload_is_sealed( + &stored + )); + assert!(!stored.contains(marker)); + + hydrate_openai_response_history( + &state, + &json!({"previous_response_id": response_id}), + "openai:responses", + "openai:chat", + scope, + ) + .await + .expect("encrypted history should hydrate"); + assert!(crate::ai_serving::response_history_is_loaded( + response_id, + Some(scope) + )); + } + + #[tokio::test] + async fn response_history_reader_rejects_and_deletes_legacy_plaintext() { + let state = response_history_test_state(); + let response_id = "resp_gateway_legacy_history_v1"; + let scope = "response-history-legacy-scope"; + let storage_key = response_history_storage_key(response_id, Some(scope)); + let payload = response_history_payload(response_id, scope, "legacy-private-history"); + state + .runtime_kv_setex(&storage_key, &payload, 6 * 60 * 60) + .await + .expect("legacy history should store"); + + let result = hydrate_openai_response_history( + &state, + &json!({"previous_response_id": response_id}), + "openai:responses", + "openai:chat", + scope, + ) + .await; + assert!(result.is_err()); + assert!(!crate::ai_serving::response_history_is_loaded( + response_id, + Some(scope) + )); + assert!(state + .runtime_kv_get(&storage_key) + .await + .expect("history lookup should succeed") + .is_none()); } } diff --git a/apps/aether-gateway/src/api/backend/public.rs b/apps/aether-gateway/src/api/backend/public.rs index 628a9badf..a2f334f7a 100644 --- a/apps/aether-gateway/src/api/backend/public.rs +++ b/apps/aether-gateway/src/api/backend/public.rs @@ -1,7 +1,10 @@ use axum::routing::get; use axum::Router; -use crate::{handlers::proxy::proxy_request, state::AppState}; +use crate::{ + handlers::{proxy::proxy_request, public::vscodex_ws_proxy}, + state::AppState, +}; pub(crate) fn mount_public_support_routes(router: Router) -> Router { router @@ -26,6 +29,7 @@ pub(crate) fn mount_public_support_routes(router: Router) -> Router) -> impl In "internal_gateway": { "route_groups": INTERNAL_GATEWAY_ROUTE_GROUPS, "path_prefixes": INTERNAL_GATEWAY_PATH_PREFIXES, - "status": "rust_native_control_plane", + "status": state.internal_gateway_auth_status(), }, }, "features": { diff --git a/apps/aether-gateway/src/api/ops.rs b/apps/aether-gateway/src/api/ops.rs index 7c03c1813..6fb341455 100644 --- a/apps/aether-gateway/src/api/ops.rs +++ b/apps/aether-gateway/src/api/ops.rs @@ -1,5 +1,14 @@ +use std::net::SocketAddr; + +use axum::body::Body; +use axum::extract::{ConnectInfo, Request, State}; +use axum::http::{self, HeaderValue, StatusCode}; +use axum::middleware::{self, Next}; +use axum::response::{IntoResponse, Response}; use axum::routing::{get, post}; -use axum::Router; +use axum::{Json, Router}; +use serde_json::json; +use tracing::warn; use crate::async_task::{ cancel_video_task, get_video_task_detail, get_video_task_stats, get_video_task_video, @@ -10,8 +19,18 @@ use crate::hooks::{get_request_audit_bundle, get_request_usage_audit}; use crate::router::metrics; use crate::state::AppState; -pub(crate) fn mount_operational_routes(router: Router) -> Router { - router +#[derive(Clone, Copy)] +struct OperationalPermission { + required_permissions: &'static [&'static str], + write: bool, + requires_full_admin_role: bool, +} + +pub(crate) fn mount_operational_routes( + router: Router, + state: AppState, +) -> Router { + let operational = Router::::new() .route("/_gateway/metrics", get(metrics)) .route("/_gateway/async-tasks/video-tasks", get(list_video_tasks)) .route( @@ -50,4 +69,236 @@ pub(crate) fn mount_operational_routes(router: Router) -> Router, + request: Request, + next: Next, +) -> Response { + let Some(permission) = operational_permission(request.method(), request.uri().path()) else { + return operational_error_response( + StatusCode::FORBIDDEN, + "operational route permission is not configured", + None, + ); + }; + let Some(remote_addr) = request + .extensions() + .get::>() + .map(|value| value.0) + else { + return operational_error_response( + StatusCode::SERVICE_UNAVAILABLE, + "operational authentication unavailable", + None, + ); + }; + let headers = request.headers().clone(); + let uri = request.uri().clone(); + if headers.get_all(http::header::AUTHORIZATION).iter().count() > 1 { + return operational_auth_required_response(); + } + + match crate::control::resolve_local_admin_session_principal(&state, &headers, &uri).await { + Ok(Some(principal)) => { + if permission.requires_full_admin_role + && !crate::roles::is_full_admin_role(&principal.user_role) + { + return operational_permission_denied_response(permission.required_permissions[0]); + } + if permission.write && !crate::roles::can_write_admin_console(&principal.user_role) { + return operational_permission_denied_response(permission.required_permissions[0]); + } + } + Ok(None) => { + let client_ip = crate::headers::effective_client_ip(&headers, &remote_addr); + let authenticated = match crate::management_token_auth::authenticate_management_token( + &state, &headers, client_ip, + ) + .await + { + Ok(authenticated) => authenticated, + Err( + crate::management_token_auth::ManagementTokenAuthError::Missing + | crate::management_token_auth::ManagementTokenAuthError::Invalid, + ) => return operational_auth_required_response(), + Err(crate::management_token_auth::ManagementTokenAuthError::Unavailable) => { + return operational_error_response( + StatusCode::SERVICE_UNAVAILABLE, + "operational authentication unavailable", + None, + ) + } + }; + + if permission.requires_full_admin_role + && !crate::roles::is_full_admin_role(&authenticated.user.role) + { + return operational_permission_denied_response(permission.required_permissions[0]); + } + if permission.write && !crate::roles::can_write_admin_console(&authenticated.user.role) + { + return operational_permission_denied_response(permission.required_permissions[0]); + } + let missing_permission = + permission + .required_permissions + .iter() + .copied() + .find(|required| { + !management_token_has_operational_permission( + &authenticated.permissions, + required, + ) + }); + if let Some(required_permission) = missing_permission { + return operational_permission_denied_response(required_permission); + } + + let client_ip = client_ip.to_string(); + if let Err(err) = state + .record_management_token_usage(&authenticated.token.id, Some(client_ip.as_str())) + .await + { + warn!( + token_id = %authenticated.token.id, + error = ?err, + "gateway failed to record operational management token usage" + ); + } + } + Err(err) => { + warn!(error = ?err, "operational admin session authentication failed"); + return operational_error_response( + StatusCode::SERVICE_UNAVAILABLE, + "operational authentication unavailable", + None, + ); + } + } + + let mut response = next.run(request).await; + response.headers_mut().insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("no-store"), + ); + response +} + +fn operational_permission(method: &http::Method, path: &str) -> Option { + if path == "/_gateway/metrics" { + return Some(OperationalPermission { + required_permissions: &["admin:monitoring:read"], + write: false, + requires_full_admin_role: false, + }); + } + if path.starts_with("/_gateway/async-tasks/video-tasks") { + let write = *method == http::Method::POST && path.ends_with("/cancel"); + return Some(OperationalPermission { + required_permissions: if write { + &["admin:video_tasks:write"] + } else { + &["admin:video_tasks:read"] + }, + write, + requires_full_admin_role: false, + }); + } + if path.starts_with("/_gateway/audit/auth/users/") { + return Some(OperationalPermission { + required_permissions: &["admin:api_keys:read"], + write: false, + requires_full_admin_role: false, + }); + } + if path.starts_with("/_gateway/audit/request-audit/") { + return Some(OperationalPermission { + required_permissions: &[ + "admin:monitoring:admin", + "admin:usage:read", + "admin:api_keys:read", + ], + write: false, + requires_full_admin_role: true, + }); + } + if path.starts_with("/_gateway/audit/request-candidates/") + || path.starts_with("/_gateway/audit/decision-trace/") + { + return Some(OperationalPermission { + required_permissions: &["admin:monitoring:admin"], + write: false, + requires_full_admin_role: true, + }); + } + if path.starts_with("/_gateway/audit/") { + return Some(OperationalPermission { + required_permissions: &["admin:usage:read"], + write: false, + requires_full_admin_role: false, + }); + } + None +} + +fn management_token_has_operational_permission( + permissions: &[String], + required_permission: &str, +) -> bool { + let scope = required_permission + .rsplit_once(':') + .map(|(scope, _)| scope) + .unwrap_or(required_permission); + let admin_permission = format!("{scope}:admin"); + permissions + .iter() + .any(|permission| permission == required_permission || permission == &admin_permission) +} + +fn operational_auth_required_response() -> Response { + let mut response = operational_error_response( + StatusCode::UNAUTHORIZED, + "admin authentication required", + None, + ); + response.headers_mut().insert( + http::header::WWW_AUTHENTICATE, + HeaderValue::from_static("Bearer"), + ); + response +} + +fn operational_permission_denied_response(required_permission: &'static str) -> Response { + operational_error_response( + StatusCode::FORBIDDEN, + "operational permission denied", + Some(required_permission), + ) +} + +fn operational_error_response( + status: StatusCode, + detail: &'static str, + required_permission: Option<&'static str>, +) -> Response { + let mut response = ( + status, + Json(json!({ + "detail": detail, + "required_permission": required_permission, + })), + ) + .into_response(); + response.headers_mut().insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("no-store"), + ); + response } diff --git a/apps/aether-gateway/src/api/response.rs b/apps/aether-gateway/src/api/response.rs index 403037c7a..8b49717e1 100644 --- a/apps/aether-gateway/src/api/response.rs +++ b/apps/aether-gateway/src/api/response.rs @@ -11,6 +11,7 @@ use crate::constants::*; use crate::control::GatewayControlDecision; use crate::control::GatewayLocalAuthRejection; use crate::headers::should_skip_response_header; +use crate::plan_usage_policy::PlanUsagePolicyRejection; use crate::rate_limit::FrontdoorUserRpmRejection; use crate::{insert_header_if_missing, GatewayError}; @@ -52,22 +53,34 @@ pub(crate) fn apply_streaming_response_headers(headers: &mut http::HeaderMap) { ); } +fn apply_gateway_browser_security_headers(headers: &mut http::HeaderMap) { + // Provider responses are API data, even when an untrusted provider labels + // them as HTML or SVG. Keep a direct navigation to a gateway API route + // from becoming same-origin active content, and prevent referrer leakage + // if a user follows a link rendered from such a response. + headers.insert( + http::header::X_CONTENT_TYPE_OPTIONS, + HeaderValue::from_static("nosniff"), + ); + headers.insert( + HeaderName::from_static("content-security-policy"), + HeaderValue::from_static( + "default-src 'none'; base-uri 'none'; form-action 'none'; frame-ancestors 'none'; sandbox", + ), + ); + headers.insert( + HeaderName::from_static("referrer-policy"), + HeaderValue::from_static("no-referrer"), + ); +} + pub(crate) fn build_client_response( upstream_response: reqwest::Response, trace_id: &str, control_decision: Option<&GatewayControlDecision>, ) -> Result, GatewayError> { let status = upstream_response.status(); - let upstream_headers = upstream_response - .headers() - .iter() - .map(|(name, value)| { - ( - name.as_str().to_string(), - value.to_str().unwrap_or_default().to_string(), - ) - }) - .collect::>(); + let upstream_headers = collect_safe_response_headers(upstream_response.headers()); let upstream_stream = upstream_response.bytes_stream(); build_client_response_from_parts( status.as_u16(), @@ -78,6 +91,40 @@ pub(crate) fn build_client_response( ) } +fn collect_safe_response_headers(headers: &http::HeaderMap) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); + headers + .iter() + .filter_map(|(name, value)| { + let normalized = name.as_str().to_ascii_lowercase(); + if should_skip_client_response_header(&normalized) + || connection_declared.contains(&normalized) + { + return None; + } + value + .to_str() + .ok() + .map(|value| (normalized, value.to_string())) + }) + .collect() +} + +fn should_skip_client_response_header(name: &str) -> bool { + should_skip_response_header(name) + // A provider Location is relative to the provider, not to the gateway. + // Forwarding it lets redirect-following clients bypass the gateway and + // can disclose their gateway Authorization header to another origin. + // Keep Location available inside execution reports, but never expose + // it on the client-facing response boundary. + || name.eq_ignore_ascii_case(http::header::LOCATION.as_str()) +} + pub(crate) fn build_client_response_from_parts( status_code: u16, upstream_headers: &BTreeMap, @@ -111,8 +158,17 @@ where .body(body) .map_err(|err| GatewayError::Internal(err.to_string()))?; + let connection_declared = aether_http::connection_declared_header_names( + upstream_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); + for (name, value) in upstream_headers { - if should_skip_response_header(name.as_str()) { + if should_skip_client_response_header(name.as_str()) + || connection_declared.contains(&name.to_ascii_lowercase()) + { continue; } let header_name = HeaderName::from_bytes(name.as_bytes()) @@ -123,6 +179,7 @@ where } mutate_headers(response.headers_mut())?; apply_streaming_response_headers(response.headers_mut()); + apply_gateway_browser_security_headers(response.headers_mut()); insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, trace_id)?; insert_header_if_missing(response.headers_mut(), GATEWAY_HEADER, "rust-phase3b")?; if let Some(decision) = control_decision { @@ -258,6 +315,57 @@ pub(crate) fn build_local_user_rpm_limited_response( ) } +pub(crate) fn build_local_plan_usage_limited_response( + trace_id: &str, + control_decision: Option<&GatewayControlDecision>, + rejection: &PlanUsagePolicyRejection, +) -> Result, GatewayError> { + let message = "套餐使用限制已达到上限,请稍后重试"; + let fallback_payload = json!({ + "error": { + "type": "plan_usage_limit_exceeded", + "message": message, + "details": { + "metric": rejection.metric, + "window": rejection.window, + "limit": rejection.limit, + "retry_after": rejection.retry_after, + } + } + }); + let payload = build_local_error_payload( + control_decision, + None, + message, + LocalCoreSyncErrorKind::RateLimit, + fallback_payload, + ); + let body = + serde_json::to_vec(&payload).map_err(|err| GatewayError::Internal(err.to_string()))?; + let headers = BTreeMap::from([ + ("content-type".to_string(), "application/json".to_string()), + ("Retry-After".to_string(), rejection.retry_after.to_string()), + ("X-RateLimit-Limit".to_string(), rejection.limit.to_string()), + ("X-RateLimit-Remaining".to_string(), "0".to_string()), + ("X-RateLimit-Scope".to_string(), "plan".to_string()), + ( + "X-RateLimit-Metric".to_string(), + rejection.metric.to_string(), + ), + ( + "X-RateLimit-Window".to_string(), + rejection.window.to_string(), + ), + ]); + build_client_response_from_parts( + StatusCode::TOO_MANY_REQUESTS.as_u16(), + &headers, + Body::from(body), + trace_id, + control_decision, + ) +} + pub(crate) fn build_local_http_error_response( trace_id: &str, control_decision: Option<&GatewayControlDecision>, @@ -454,11 +562,13 @@ fn local_error_kind_for_status(status: StatusCode) -> LocalCoreSyncErrorKind { #[cfg(test)] mod tests { use super::{ - build_client_response_from_parts, build_local_auth_rejection_response, + build_client_response, build_client_response_from_parts, + build_client_response_from_parts_with_mutator, build_local_auth_rejection_response, build_local_http_error_response_with_request_path, build_local_overloaded_response, - build_local_user_rpm_limited_response, + build_local_plan_usage_limited_response, build_local_user_rpm_limited_response, }; use crate::control::{GatewayControlDecision, GatewayLocalAuthRejection}; + use crate::plan_usage_policy::PlanUsagePolicyRejection; use crate::rate_limit::FrontdoorUserRpmRejection; use axum::body::{to_bytes, Body}; use std::collections::BTreeMap; @@ -490,6 +600,163 @@ mod tests { ); } + #[test] + fn upstream_security_headers_are_stripped_before_gateway_headers_are_added() { + let response = build_client_response_from_parts_with_mutator( + 200, + &BTreeMap::from([ + ("set-cookie".to_string(), "session=attacker".to_string()), + ( + "x-aether-gateway".to_string(), + "attacker-gateway".to_string(), + ), + ( + "x-aether-control-action".to_string(), + "attacker-action".to_string(), + ), + ( + "x-aether-future-control".to_string(), + "attacker-future".to_string(), + ), + ( + "x-accel-redirect".to_string(), + "/internal/private-file".to_string(), + ), + ("x-sendfile".to_string(), "/etc/passwd".to_string()), + ( + "x-reproxy-url".to_string(), + "http://127.0.0.1:9000/private".to_string(), + ), + ( + "access-control-allow-origin".to_string(), + "https://attacker.example".to_string(), + ), + ( + "access-control-allow-credentials".to_string(), + "true".to_string(), + ), + ("content-length".to_string(), "999999".to_string()), + ( + "content-security-policy".to_string(), + "default-src * 'unsafe-inline' 'unsafe-eval'".to_string(), + ), + ( + "content-security-policy-report-only".to_string(), + "default-src 'none'; report-uri https://attacker.example/csp".to_string(), + ), + ( + "reporting-endpoints".to_string(), + "attacker=\"https://attacker.example/reports\"".to_string(), + ), + ("report-to".to_string(), "attacker".to_string()), + ( + "nel".to_string(), + "{\"report_to\":\"attacker\"}".to_string(), + ), + ( + "refresh".to_string(), + "0; url=https://attacker.example".to_string(), + ), + ("referrer-policy".to_string(), "unsafe-url".to_string()), + ("x-content-type-options".to_string(), "invalid".to_string()), + ( + "location".to_string(), + "https://provider.example/direct".to_string(), + ), + ("x-upstream-visible".to_string(), "ok".to_string()), + ]), + Body::empty(), + "trace-upstream-header-filter", + None, + |headers| { + headers.insert( + http::HeaderName::from_static("x-aether-control-action"), + http::HeaderValue::from_static("gateway-action"), + ); + Ok(()) + }, + ) + .expect("response should build"); + + assert!(response.headers().get(http::header::SET_COOKIE).is_none()); + assert!(response.headers().get("x-aether-future-control").is_none()); + assert!(response.headers().get("x-accel-redirect").is_none()); + assert!(response.headers().get("x-sendfile").is_none()); + assert!(response.headers().get("x-reproxy-url").is_none()); + assert!(response + .headers() + .get("access-control-allow-origin") + .is_none()); + assert!(response + .headers() + .get("access-control-allow-credentials") + .is_none()); + assert!(response + .headers() + .get(http::header::CONTENT_LENGTH) + .is_none()); + assert!(response + .headers() + .get("content-security-policy-report-only") + .is_none()); + assert!(response.headers().get("reporting-endpoints").is_none()); + assert!(response.headers().get("report-to").is_none()); + assert!(response.headers().get("nel").is_none()); + assert!(response.headers().get("refresh").is_none()); + assert!(response.headers().get(http::header::LOCATION).is_none()); + assert_eq!( + response.headers()["content-security-policy"], + "default-src 'none'; base-uri 'none'; form-action 'none'; frame-ancestors 'none'; sandbox" + ); + assert_eq!(response.headers()["referrer-policy"], "no-referrer"); + assert_eq!( + response.headers()[http::header::X_CONTENT_TYPE_OPTIONS], + "nosniff" + ); + assert_eq!(response.headers()["x-aether-gateway"], "rust-phase3b"); + assert_eq!( + response.headers()["x-aether-control-action"], + "gateway-action" + ); + assert_eq!(response.headers()["x-upstream-visible"], "ok"); + } + + #[tokio::test] + async fn raw_response_collector_honors_all_connection_header_lines() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener"); + let addr = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("connection"); + use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; + let mut request = [0_u8; 1024]; + let _ = stream.read(&mut request).await.expect("request read"); + stream + .write_all( + b"HTTP/1.1 200 OK\r\nConnection: x-first-hop\r\nConnection: x-second-hop\r\nX-First-Hop: first-secret\r\nX-Second-Hop: second-secret\r\nContent-Length: 2\r\n\r\nok", + ) + .await + .expect("response write"); + }); + let upstream = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("client") + .get(format!("http://{addr}/")) + .send() + .await + .expect("upstream response"); + + let response = build_client_response(upstream, "trace-connection-lines", None) + .expect("client response"); + server.await.expect("server"); + + assert!(response.headers().get("connection").is_none()); + assert!(response.headers().get("x-first-hop").is_none()); + assert!(response.headers().get("x-second-hop").is_none()); + } + fn claude_decision() -> GatewayControlDecision { GatewayControlDecision::synthetic( "/v1/messages", @@ -581,4 +848,25 @@ mod tests { ); } } + + #[tokio::test] + async fn plan_usage_rejection_exposes_machine_readable_limit_headers() { + let response = build_local_plan_usage_limited_response( + "trace-plan-limit", + None, + &PlanUsagePolicyRejection { + metric: "request_count", + limit: 100.0, + retry_after: 42, + window: "calendar_week", + }, + ) + .expect("response"); + assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS); + assert_eq!(response.headers()["retry-after"], "42"); + assert_eq!(response.headers()["x-ratelimit-scope"], "plan"); + assert_eq!(response.headers()["x-ratelimit-window"], "calendar_week"); + let payload = response_json(response).await; + assert_eq!(payload["error"]["type"], "plan_usage_limit_exceeded"); + } } diff --git a/apps/aether-gateway/src/async_task/http.rs b/apps/aether-gateway/src/async_task/http.rs index d9d452947..28ce43eee 100644 --- a/apps/aether-gateway/src/async_task/http.rs +++ b/apps/aether-gateway/src/async_task/http.rs @@ -1,3 +1,4 @@ +use std::net::{IpAddr, SocketAddr}; use std::time::{SystemTime, UNIX_EPOCH}; use aether_contracts::ExecutionResult; @@ -21,7 +22,9 @@ use super::{ }; use crate::{AppState, GatewayError}; -pub(crate) use self::cancel::{cancel_video_task_record, CancelVideoTaskError}; +pub(crate) use self::cancel::{ + cancel_video_task_record, cancel_video_task_record_for_user, CancelVideoTaskError, +}; #[derive(Debug, Deserialize)] pub(crate) struct ListVideoTasksQuery { @@ -146,18 +149,21 @@ pub(crate) async fn get_video_task_video( } pub(crate) async fn build_video_task_video_response( - state: &AppState, + _state: &AppState, task_id: &str, source: VideoTaskVideoSource, ) -> Result { match source { - VideoTaskVideoSource::Redirect { url } => Ok(Redirect::temporary(&url).into_response()), + VideoTaskVideoSource::Redirect { url } => { + resolve_public_video_target(&url).await?; + Ok(Redirect::temporary(url.as_str()).into_response()) + } VideoTaskVideoSource::Proxy { url, header_name, header_value, filename, - } => proxy_video_stream(state, task_id, &url, &header_name, &header_value, &filename).await, + } => proxy_video_stream(task_id, &url, &header_name, &header_value, &filename).await, } } @@ -209,25 +215,34 @@ fn video_task_status_name(status: VideoTaskStatus) -> &'static str { } async fn proxy_video_stream( - state: &AppState, task_id: &str, - url: &str, + url: &url::Url, header_name: &str, header_value: &str, filename: &str, ) -> Result { - let response = state - .client - .get(url) + let target = resolve_public_video_target(url).await?; + let client = build_pinned_video_client(&target)?; + let response = client + .get(url.clone()) .header(header_name, header_value) .send() .await .map_err(|err| GatewayError::UpstreamUnavailable { trace_id: task_id.to_string(), - message: err.to_string(), + message: video_request_failure_message(&err).to_string(), })?; - if response.status().is_client_error() || response.status().is_server_error() { + if response.status().is_redirection() { + return Err(GatewayError::UpstreamUnavailable { + trace_id: task_id.to_string(), + message: format!( + "video upstream redirect was rejected with HTTP {}", + response.status() + ), + }); + } + if !response.status().is_success() { return Err(GatewayError::UpstreamUnavailable { trace_id: task_id.to_string(), message: format!("video upstream returned HTTP {}", response.status()), @@ -235,47 +250,321 @@ async fn proxy_video_stream( } let status = response.status(); - let content_type = response - .headers() - .get(axum::http::header::CONTENT_TYPE) - .cloned() - .unwrap_or_else(|| axum::http::HeaderValue::from_static("video/mp4")); - let content_length = response - .headers() - .get(axum::http::header::CONTENT_LENGTH) - .cloned(); - let cache_control = response - .headers() - .get(axum::http::header::CACHE_CONTROL) - .cloned(); + // Do not copy the provider's Content-Length onto a newly wrapped stream. + // Reqwest may decode transfer/content encodings and the provider controls + // the declaration; forwarding a stale value would make the client-facing + // HTTP framing disagree with the bytes produced by this Body. Axum/Hyper + // will select safe framing for the actual stream. + let upstream_headers = response.headers().clone(); let body = Body::from_stream(response.bytes_stream()); let mut outbound = axum::http::Response::builder() .status(status) .body(body) .map_err(|err| GatewayError::Internal(err.to_string()))?; - outbound - .headers_mut() - .insert(axum::http::header::CONTENT_TYPE, content_type); - outbound.headers_mut().insert( - axum::http::header::CONTENT_DISPOSITION, - axum::http::HeaderValue::from_str(&format!("inline; filename=\"{filename}\"")) - .map_err(|err| GatewayError::Internal(err.to_string()))?, - ); - if let Some(content_length) = content_length { - outbound - .headers_mut() - .insert(axum::http::header::CONTENT_LENGTH, content_length); - } - if let Some(cache_control) = cache_control { - outbound - .headers_mut() - .insert(axum::http::header::CACHE_CONTROL, cache_control); - } else { - outbound.headers_mut().insert( - axum::http::header::CACHE_CONTROL, - axum::http::HeaderValue::from_static("private, max-age=3600"), - ); - } + apply_safe_video_response_metadata(outbound.headers_mut(), &upstream_headers, filename)?; Ok(outbound) } + +fn apply_safe_video_response_metadata( + outbound: &mut axum::http::HeaderMap, + upstream: &axum::http::HeaderMap, + filename: &str, +) -> Result<(), GatewayError> { + let content_type = upstream + .get(axum::http::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .and_then(safe_video_content_type) + .unwrap_or_else(|| axum::http::HeaderValue::from_static("application/octet-stream")); + outbound.insert(axum::http::header::CONTENT_TYPE, content_type); + outbound.insert( + axum::http::header::CONTENT_DISPOSITION, + axum::http::HeaderValue::from_str(&format!( + "inline; filename=\"{}\"", + safe_video_filename(filename) + )) + .map_err(|err| GatewayError::Internal(err.to_string()))?, + ); + outbound.remove(axum::http::header::CONTENT_LENGTH); + outbound.insert( + axum::http::header::CACHE_CONTROL, + axum::http::HeaderValue::from_static("private, no-store"), + ); + outbound.insert( + axum::http::header::X_CONTENT_TYPE_OPTIONS, + axum::http::HeaderValue::from_static("nosniff"), + ); + Ok(()) +} + +fn safe_video_content_type(raw_value: &str) -> Option { + let media_type = raw_value.split(';').next()?.trim().to_ascii_lowercase(); + let subtype = media_type.strip_prefix("video/")?; + if subtype.is_empty() + || !subtype.bytes().all(|byte| { + byte.is_ascii_alphanumeric() + || matches!( + byte, + b'!' | b'#' | b'$' | b'&' | b'-' | b'^' | b'_' | b'.' | b'+' + ) + }) + { + return None; + } + axum::http::HeaderValue::from_str(raw_value).ok() +} + +fn safe_video_filename(filename: &str) -> String { + let filename = filename + .chars() + .take(255) + .map(|character| { + if character.is_ascii_alphanumeric() || matches!(character, '.' | '-' | '_') { + character + } else { + '_' + } + }) + .collect::(); + if filename.is_empty() { + "video.mp4".to_string() + } else { + filename + } +} + +struct ResolvedVideoTarget { + host: String, + addrs: Vec, +} + +async fn resolve_public_video_target(url: &url::Url) -> Result { + if !matches!(url.scheme(), "http" | "https") + || !url.username().is_empty() + || url.password().is_some() + { + return Err(video_target_rejected( + "video URL must be an absolute HTTP(S) URL without credentials", + )); + } + let port = url + .port_or_known_default() + .ok_or_else(|| video_target_rejected("video URL is missing a port"))?; + let (host, addrs) = match url.host() { + Some(url::Host::Ipv4(ip)) => (ip.to_string(), vec![SocketAddr::new(IpAddr::V4(ip), port)]), + Some(url::Host::Ipv6(ip)) => (ip.to_string(), vec![SocketAddr::new(IpAddr::V6(ip), port)]), + Some(url::Host::Domain(host)) if !host.is_empty() => { + let addrs = aether_http::lookup_host_with_limits( + host, + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| video_target_rejected("video URL DNS resolution failed"))?; + (host.to_string(), addrs) + } + _ => return Err(video_target_rejected("video URL is missing a host")), + }; + if addrs.is_empty() + || addrs + .iter() + .any(|addr| aether_http::is_private_or_reserved_ip(addr.ip())) + { + return Err(video_target_rejected( + "video URL resolves to a private or reserved address", + )); + } + Ok(ResolvedVideoTarget { host, addrs }) +} + +fn build_pinned_video_client( + target: &ResolvedVideoTarget, +) -> Result { + let mut builder = aether_http::apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &aether_http::HttpClientConfig { + connect_timeout_ms: Some(10_000), + request_timeout_ms: Some(300_000), + http2_adaptive_window: true, + ..aether_http::HttpClientConfig::default() + }, + ); + if target.host.parse::().is_err() { + builder = builder.resolve_to_addrs(&target.host, &target.addrs); + } + builder + .build() + .map_err(|_| GatewayError::Internal("video HTTP client initialization failed".to_string())) +} + +fn video_target_rejected(message: &str) -> GatewayError { + GatewayError::Client { + status: axum::http::StatusCode::BAD_GATEWAY, + message: message.to_string(), + } +} + +fn video_request_failure_message(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "video upstream request timed out" + } else if error.is_connect() { + "video upstream connection failed" + } else if error.is_body() || error.is_decode() { + "video upstream response failed" + } else { + "video upstream request failed" + } +} + +#[cfg(test)] +mod tests { + use axum::response::IntoResponse; + + use super::{ + apply_safe_video_response_metadata, build_video_task_video_response, + resolve_public_video_target, safe_video_content_type, safe_video_filename, + VideoTaskVideoSource, + }; + use crate::AppState; + + #[tokio::test] + async fn video_redirect_response_accepts_public_target() { + let state = AppState::new().expect("gateway state should build"); + let target = "https://8.8.8.8/video.mp4"; + + let response = build_video_task_video_response( + &state, + "task-public-redirect", + VideoTaskVideoSource::Redirect { + url: url::Url::parse(target).expect("public target should parse"), + }, + ) + .await + .expect("public redirect should build"); + + assert_eq!( + response.status(), + axum::http::StatusCode::TEMPORARY_REDIRECT + ); + assert_eq!( + response + .headers() + .get(axum::http::header::LOCATION) + .and_then(|value| value.to_str().ok()), + Some(target) + ); + } + + #[tokio::test] + async fn video_redirect_response_rejects_private_and_reserved_targets() { + let state = AppState::new().expect("gateway state should build"); + + for raw_url in [ + "http://127.0.0.1/video.mp4", + "http://169.254.169.254/latest/meta-data", + "http://10.0.0.1/video.mp4", + "http://[::1]/video.mp4", + ] { + let error = build_video_task_video_response( + &state, + "task-rejected-redirect", + VideoTaskVideoSource::Redirect { + url: url::Url::parse(raw_url).expect("target should parse"), + }, + ) + .await + .expect_err("private or reserved redirect target should be rejected"); + + assert_eq!( + error.into_response().status(), + axum::http::StatusCode::BAD_GATEWAY, + "unexpected status for {raw_url}" + ); + } + } + + #[tokio::test] + async fn video_target_resolution_rejects_private_and_reserved_ip_literals() { + for raw_url in [ + "http://127.0.0.1/video.mp4", + "http://169.254.169.254/latest/meta-data", + "http://10.0.0.1/video.mp4", + "http://[::1]/video.mp4", + "http://[::ffff:127.0.0.1]/video.mp4", + ] { + let url = url::Url::parse(raw_url).unwrap(); + assert!( + resolve_public_video_target(&url).await.is_err(), + "target should be rejected: {raw_url}" + ); + } + } + + #[tokio::test] + async fn video_target_resolution_accepts_public_ip_literals() { + for raw_url in [ + "https://8.8.8.8/video.mp4", + "https://[2606:4700:4700::1111]/video.mp4", + ] { + let url = url::Url::parse(raw_url).unwrap(); + assert!( + resolve_public_video_target(&url).await.is_ok(), + "target should be accepted: {raw_url}" + ); + } + } + + #[test] + fn video_response_metadata_rejects_active_content_and_sanitizes_filename() { + assert!(safe_video_content_type("video/mp4").is_some()); + assert!(safe_video_content_type("video/webm; charset=binary").is_some()); + assert!(safe_video_content_type("video/").is_none()); + assert!(safe_video_content_type("video/; charset=binary").is_none()); + assert!(safe_video_content_type("text/html").is_none()); + assert!(safe_video_content_type("video/mp4\r\nx-test: injected").is_none()); + assert_eq!( + safe_video_filename("video_123.mp4\"; filename=\"attack.html"), + "video_123.mp4___filename__attack.html" + ); + assert_eq!(safe_video_filename(&"x".repeat(1024)).len(), 255); + + let mut upstream = axum::http::HeaderMap::new(); + upstream.insert( + axum::http::header::CONTENT_TYPE, + axum::http::HeaderValue::from_static("text/html"), + ); + upstream.insert( + axum::http::header::CONTENT_LENGTH, + axum::http::HeaderValue::from_static("999999"), + ); + let mut outbound = upstream.clone(); + apply_safe_video_response_metadata( + &mut outbound, + &upstream, + "video.mp4\"; filename=\"attack.html", + ) + .expect("video metadata should build"); + + assert_eq!( + outbound + .get(axum::http::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()), + Some("application/octet-stream") + ); + assert!(outbound.get(axum::http::header::CONTENT_LENGTH).is_none()); + assert_eq!( + outbound + .get(axum::http::header::X_CONTENT_TYPE_OPTIONS) + .and_then(|value| value.to_str().ok()), + Some("nosniff") + ); + assert_eq!( + outbound + .get(axum::http::header::CONTENT_DISPOSITION) + .and_then(|value| value.to_str().ok()), + Some("inline; filename=\"video.mp4___filename__attack.html\"") + ); + } +} diff --git a/apps/aether-gateway/src/async_task/http/cancel.rs b/apps/aether-gateway/src/async_task/http/cancel.rs index e5671d6cf..ca489f465 100644 --- a/apps/aether-gateway/src/async_task/http/cancel.rs +++ b/apps/aether-gateway/src/async_task/http/cancel.rs @@ -3,12 +3,13 @@ use aether_data_contracts::repository::video_tasks::{ }; use axum::response::IntoResponse; use axum::Json; -use serde_json::{json, Map, Value}; +use serde_json::json; +use crate::state::VideoTaskRouteAccess; use crate::{AppState, GatewayError}; use super::super::finalize_video_task_if_terminal; -use super::super::read_video_task_detail; +use super::super::{read_video_task_detail, read_video_task_detail_for_user}; use super::current_unix_secs; #[derive(Debug)] @@ -29,7 +30,31 @@ pub(crate) async fn cancel_video_task_record( state: &AppState, task_id: &str, ) -> Result { - let Some(task) = read_video_task_detail(state, task_id).await? else { + cancel_video_task_record_inner(state, task_id, None).await +} + +pub(crate) async fn cancel_video_task_record_for_user( + state: &AppState, + task_id: &str, + user_id: &str, +) -> Result { + let user_id = user_id.trim(); + if user_id.is_empty() { + return Err(CancelVideoTaskError::NotFound); + } + cancel_video_task_record_inner(state, task_id, Some(user_id)).await +} + +async fn cancel_video_task_record_inner( + state: &AppState, + task_id: &str, + expected_user_id: Option<&str>, +) -> Result { + let task = match expected_user_id { + Some(user_id) => read_video_task_detail_for_user(state, task_id, user_id).await?, + None => read_video_task_detail(state, task_id).await?, + }; + let Some(task) = task else { return Err(CancelVideoTaskError::NotFound); }; @@ -45,39 +70,84 @@ pub(crate) async fn cancel_video_task_record( } let trace_id = format!("async-task-admin-cancel-{task_id}"); + let mut finalize_mutation = None; if let Some(cancel_plan) = build_video_task_cancel_plan(&task) { - state - .hydrate_video_task_for_route(Some(cancel_plan.route_family), &cancel_plan.request_path) - .await?; - let body_json = json!({}); - let follow_up = state.video_tasks.prepare_follow_up_sync_plan( - cancel_plan.plan_kind, - &cancel_plan.request_path, - Some(&body_json), - None, - &trace_id, - ); + let follow_up = if let Some(user_id) = expected_user_id { + if state + .hydrate_video_task_for_route_for_user( + Some(cancel_plan.route_family), + &cancel_plan.request_path, + user_id, + ) + .await? + != VideoTaskRouteAccess::Allowed + { + return Err(CancelVideoTaskError::NotFound); + } + state.video_tasks.prepare_follow_up_sync_plan_for_user_id( + cancel_plan.plan_kind, + &cancel_plan.request_path, + Some(&body_json), + user_id, + task.api_key_id.as_deref(), + &trace_id, + ) + } else { + state + .hydrate_video_task_for_route( + Some(cancel_plan.route_family), + &cancel_plan.request_path, + ) + .await?; + state.video_tasks.prepare_follow_up_sync_plan( + cancel_plan.plan_kind, + &cancel_plan.request_path, + Some(&body_json), + None, + &trace_id, + ) + }; if let Some(follow_up) = follow_up { execute_video_task_cancel_plan(state, &trace_id, follow_up.plan) .await .map_err(CancelVideoTaskError::Response)?; + finalize_mutation = Some(( + cancel_plan.request_path, + cancel_plan.report_kind.to_string(), + )); + } else if expected_user_id.is_none() { + finalize_mutation = Some(( + cancel_plan.request_path, + cancel_plan.report_kind.to_string(), + )); } - - state - .video_tasks - .apply_finalize_mutation(&cancel_plan.request_path, cancel_plan.report_kind); } - let request_metadata = build_cancelled_request_metadata(state, &task).await?; - let stored = persist_cancelled_video_task(state, &task, request_metadata) - .await? - .ok_or_else(|| { - CancelVideoTaskError::Gateway(GatewayError::Internal( + let stored = match persist_cancelled_video_task(state, &task).await? { + Some(stored) => stored, + None => { + let current = match expected_user_id { + Some(user_id) => read_video_task_detail_for_user(state, task_id, user_id).await?, + None => read_video_task_detail(state, task_id).await?, + }; + let Some(current) = current else { + return Err(CancelVideoTaskError::NotFound); + }; + if !current.status.is_active() { + return Err(CancelVideoTaskError::InvalidStatus(current.status)); + } + return Err(CancelVideoTaskError::Gateway(GatewayError::Internal( "video task repository is unavailable".to_string(), - )) - })?; + ))); + } + }; + if let Some((request_path, report_kind)) = finalize_mutation { + state + .video_tasks + .apply_finalize_mutation(&request_path, &report_kind); + } finalize_video_task_if_terminal(state, &stored).await; Ok(stored) } @@ -91,12 +161,7 @@ struct VideoTaskCancelPlan<'a> { } fn build_video_task_cancel_plan(task: &StoredVideoTask) -> Option> { - let provider_api_format = task - .provider_api_format - .as_deref() - .or(task.client_api_format.as_deref()) - .map(str::trim) - .filter(|value| !value.is_empty())?; + let provider_api_format = task.effective_api_format()?; match provider_api_format { "openai:video" => Some(VideoTaskCancelPlan { @@ -131,99 +196,52 @@ async fn execute_video_task_cancel_plan( let result = crate::execution_runtime::execute_execution_runtime_sync_plan(state, Some(trace_id), &plan) .await - .map_err(|err| { + .map_err(|_| { GatewayError::UpstreamUnavailable { trace_id: trace_id.to_string(), - message: format!("{err:?}"), + message: "video cancellation request failed".to_string(), } .into_response() })?; if result.status_code >= 400 { - let status = axum::http::StatusCode::from_u16(result.status_code) - .unwrap_or(axum::http::StatusCode::BAD_GATEWAY); - let body_json = result - .body - .and_then(|body| body.json_body) - .unwrap_or_else(|| { - json!({ - "error": { - "message": result - .error - .as_ref() - .map(|error| error.message.clone()) - .unwrap_or_else(|| { - format!("execution runtime returned {}", result.status_code) - }), - } - }) - }); - return Err((status, Json(body_json)).into_response()); + return Err(build_video_task_cancel_upstream_error_response(&result)); } Ok(()) } -async fn build_cancelled_request_metadata( - state: &AppState, - task: &StoredVideoTask, -) -> Result, GatewayError> { - let mut metadata = match task.request_metadata.clone() { - Some(Value::Object(object)) => object, - _ => Map::new(), - }; - let mut snapshot_value = metadata.get("rust_local_snapshot").cloned(); - if snapshot_value.is_none() { - snapshot_value = state - .reconstruct_video_task_snapshot(task) - .await? - .map(|snapshot| { - serde_json::to_value(snapshot) - .map_err(|err| GatewayError::Internal(err.to_string())) - }) - .transpose()?; - } - if let Some(snapshot_value_ref) = snapshot_value.as_mut() { - mark_snapshot_value_cancelled(snapshot_value_ref); - metadata.insert( - "rust_owner".to_string(), - Value::String("async_task".to_string()), - ); - metadata.insert( - "rust_local_snapshot".to_string(), - snapshot_value_ref.clone(), - ); - return Ok(Some(Value::Object(metadata))); - } - - Ok(task.request_metadata.clone()) -} - -fn mark_snapshot_value_cancelled(snapshot_value: &mut Value) { - if let Some(object) = snapshot_value - .get_mut("OpenAi") - .and_then(Value::as_object_mut) - { - object.insert("status".to_string(), Value::String("Cancelled".to_string())); - return; - } - if let Some(object) = snapshot_value - .get_mut("Gemini") - .and_then(Value::as_object_mut) - { - object.insert("status".to_string(), Value::String("Cancelled".to_string())); - } +fn build_video_task_cancel_upstream_error_response( + result: &aether_contracts::ExecutionResult, +) -> axum::response::Response { + let status = axum::http::StatusCode::from_u16(result.status_code) + .unwrap_or(axum::http::StatusCode::BAD_GATEWAY); + tracing::warn!( + event_name = "video_task_cancel_upstream_error", + upstream_status = result.status_code, + "video cancellation upstream response body discarded" + ); + ( + status, + Json(json!({ + "error": { + "message": format!( + "video cancellation upstream returned HTTP {}", + result.status_code + ), + } + })), + ) + .into_response() } async fn persist_cancelled_video_task( state: &AppState, task: &StoredVideoTask, - request_metadata: Option, ) -> Result, GatewayError> { let now_unix_secs = current_unix_secs(); state - .data - .upsert_video_task(UpsertVideoTask { + .update_active_video_task(UpsertVideoTask { id: task.id.clone(), short_id: task.short_id.clone(), request_id: task.request_id.clone(), @@ -240,14 +258,14 @@ async fn persist_cancelled_video_task( format_converted: task.format_converted, model: task.model.clone(), prompt: task.prompt.clone(), - original_request_body: task.original_request_body.clone(), + original_request_body: None, duration_seconds: task.duration_seconds, resolution: task.resolution.clone(), aspect_ratio: task.aspect_ratio.clone(), size: task.size.clone(), status: VideoTaskStatus::Cancelled, progress_percent: task.progress_percent, - progress_message: task.progress_message.clone(), + progress_message: None, retry_count: task.retry_count, poll_interval_seconds: task.poll_interval_seconds, next_poll_at_unix_secs: None, @@ -258,10 +276,73 @@ async fn persist_cancelled_video_task( completed_at_unix_secs: Some(now_unix_secs), updated_at_unix_secs: now_unix_secs, error_code: task.error_code.clone(), - error_message: task.error_message.clone(), + error_message: None, video_url: task.video_url.clone(), - request_metadata, + request_metadata: None, }) .await - .map_err(|err| GatewayError::Internal(err.to_string())) +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use aether_contracts::{ + ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResult, ResponseBody, + }; + use axum::body::to_bytes; + use serde_json::json; + + use super::build_video_task_cancel_upstream_error_response; + + #[tokio::test] + async fn cancellation_upstream_errors_do_not_expose_runtime_payloads() { + let result = ExecutionResult { + request_id: "cancel-secret-request-id".to_string(), + candidate_id: Some("cancel-secret-candidate-id".to_string()), + status_code: 502, + headers: BTreeMap::from([( + "x-internal-secret".to_string(), + "cancel-secret-header".to_string(), + )]), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({ + "error": { + "message": "cancel-secret-upstream-body", + } + })), + body_bytes_b64: None, + }), + telemetry: None, + error: Some(ExecutionError { + kind: ExecutionErrorKind::Upstream5xx, + phase: ExecutionPhase::FirstByte, + message: "cancel-secret-runtime-error".to_string(), + upstream_status: Some(502), + retryable: true, + failover_recommended: false, + }), + }; + + let response = build_video_task_cancel_upstream_error_response(&result); + assert_eq!(response.status(), axum::http::StatusCode::BAD_GATEWAY); + assert!(response.headers().get("x-internal-secret").is_none()); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body should read"); + let payload: serde_json::Value = + serde_json::from_slice(&body).expect("response body should parse"); + + assert_eq!( + payload, + json!({ + "error": { + "message": "video cancellation upstream returned HTTP 502", + } + }) + ); + let body = String::from_utf8(body.to_vec()).expect("response body should be utf-8"); + assert!(!body.contains("cancel-secret")); + } } diff --git a/apps/aether-gateway/src/async_task/mod.rs b/apps/aether-gateway/src/async_task/mod.rs index 5953c82f0..bb91d8a3a 100644 --- a/apps/aether-gateway/src/async_task/mod.rs +++ b/apps/aether-gateway/src/async_task/mod.rs @@ -6,13 +6,14 @@ pub(crate) use crate::video_tasks::VideoTaskService; pub use crate::video_tasks::VideoTaskTruthSourceMode; pub(crate) use http::{ build_video_task_video_response, cancel_video_task, cancel_video_task_record, - get_video_task_detail, get_video_task_stats, get_video_task_video, list_video_tasks, - CancelVideoTaskError, + cancel_video_task_record_for_user, get_video_task_detail, get_video_task_stats, + get_video_task_video, list_video_tasks, CancelVideoTaskError, }; pub(crate) use query::{ - read_video_task_detail, read_video_task_page, read_video_task_page_summary, - read_video_task_stats, read_video_task_video_source, VideoTaskPageResponse, - VideoTaskStatsResponse, VideoTaskVideoSource, + read_video_task_detail, read_video_task_detail_for_user, read_video_task_page, + read_video_task_page_summary, read_video_task_stats, read_video_task_video_source, + video_task_video_source_from_task, VideoTaskPageResponse, VideoTaskStatsResponse, + VideoTaskVideoSource, }; pub(crate) use runtime::{ execute_video_task_refresh_plan, finalize_video_task_if_terminal, spawn_video_task_poller, diff --git a/apps/aether-gateway/src/async_task/query.rs b/apps/aether-gateway/src/async_task/query.rs index 287d148e0..ccccc3d8f 100644 --- a/apps/aether-gateway/src/async_task/query.rs +++ b/apps/aether-gateway/src/async_task/query.rs @@ -26,13 +26,12 @@ pub(crate) struct VideoTaskStatsResponse { pub(crate) processing_count: u64, } -#[derive(Debug, Clone)] pub(crate) enum VideoTaskVideoSource { Redirect { - url: String, + url: url::Url, }, Proxy { - url: String, + url: url::Url, header_name: String, header_value: String, filename: String, @@ -102,6 +101,14 @@ pub(crate) async fn read_video_task_detail( state.find_video_task_by_id(task_id).await } +pub(crate) async fn read_video_task_detail_for_user( + state: &AppState, + task_id: &str, + user_id: &str, +) -> Result, GatewayError> { + state.find_video_task_by_id_for_user(task_id, user_id).await +} + pub(crate) async fn read_video_task_video_source( state: &AppState, task_id: &str, @@ -109,6 +116,13 @@ pub(crate) async fn read_video_task_video_source( let Some(task) = read_video_task_detail(state, task_id).await? else { return Ok(None); }; + video_task_video_source_from_task(state, &task).await +} + +pub(crate) async fn video_task_video_source_from_task( + state: &AppState, + task: &StoredVideoTask, +) -> Result, GatewayError> { let Some(video_url) = task .video_url .as_deref() @@ -119,7 +133,9 @@ pub(crate) async fn read_video_task_video_source( return Ok(None); }; - if !video_url.contains("generativelanguage.googleapis.com") { + let video_url = parse_video_url(&video_url)?; + + if task.effective_api_format() != Some("gemini:video") { return Ok(Some(VideoTaskVideoSource::Redirect { url: video_url })); } @@ -148,6 +164,15 @@ pub(crate) async fn read_video_task_video_source( )); }; + let endpoint_url = parse_video_url(transport.endpoint.base_url.trim()).map_err(|_| { + GatewayError::Internal("provider endpoint URL is invalid for proxied video".to_string()) + })?; + if !video_urls_share_origin(&endpoint_url, &video_url) { + return Err(GatewayError::Client { + status: axum::http::StatusCode::BAD_GATEWAY, + message: "video URL origin does not match its provider endpoint".to_string(), + }); + } let api_key = transport.key.decrypted_api_key.trim(); if api_key.is_empty() { return Err(GatewayError::Internal( @@ -159,10 +184,34 @@ pub(crate) async fn read_video_task_video_source( url: video_url, header_name: "x-goog-api-key".to_string(), header_value: api_key.to_string(), - filename: format!("video_{task_id}.mp4"), + filename: format!("video_{}.mp4", task.id), })) } +fn parse_video_url(raw_url: &str) -> Result { + let url = url::Url::parse(raw_url.trim()).map_err(|_| GatewayError::Client { + status: axum::http::StatusCode::BAD_GATEWAY, + message: "video URL is invalid".to_string(), + })?; + if !matches!(url.scheme(), "http" | "https") + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + { + return Err(GatewayError::Client { + status: axum::http::StatusCode::BAD_GATEWAY, + message: "video URL must be an absolute HTTP(S) URL without credentials".to_string(), + }); + } + Ok(url) +} + +fn video_urls_share_origin(left: &url::Url, right: &url::Url) -> bool { + left.scheme() == right.scheme() + && left.host() == right.host() + && left.port_or_known_default() == right.port_or_known_default() +} + pub(crate) async fn read_video_task_stats( state: &AppState, filter: &VideoTaskQueryFilter, @@ -226,3 +275,212 @@ fn status_key(status: VideoTaskStatus) -> String { fn start_of_utc_day(now_unix_secs: u64) -> u64 { now_unix_secs - (now_unix_secs % 86_400) } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; + use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogReadRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogProvider, + }; + use aether_data_contracts::repository::video_tasks::{UpsertVideoTask, VideoTaskStatus}; + use serde_json::json; + + use super::{ + parse_video_url, video_task_video_source_from_task, video_urls_share_origin, + VideoTaskVideoSource, + }; + use crate::{data::GatewayDataState, AppState}; + + fn legacy_gemini_video_task() -> aether_data_contracts::repository::video_tasks::StoredVideoTask + { + UpsertVideoTask { + id: "legacy-gemini-task".to_string(), + short_id: Some("legacy-short".to_string()), + request_id: "legacy-request".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("client-key-1".to_string()), + username: None, + api_key_name: None, + external_task_id: Some("operations/upstream-1".to_string()), + provider_id: Some("provider-1".to_string()), + endpoint_id: Some("endpoint-1".to_string()), + key_id: Some("provider-key-1".to_string()), + client_api_format: Some("gemini:video".to_string()), + provider_api_format: None, + format_converted: false, + model: Some("veo-3".to_string()), + prompt: None, + original_request_body: None, + duration_seconds: Some(8), + resolution: Some("720p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("1280x720".to_string()), + status: VideoTaskStatus::Completed, + progress_percent: 100, + progress_message: None, + retry_count: 0, + poll_interval_seconds: 10, + next_poll_at_unix_secs: None, + poll_count: 1, + max_poll_count: 360, + created_at_unix_ms: 1, + submitted_at_unix_secs: Some(1), + completed_at_unix_secs: Some(2), + updated_at_unix_secs: 2, + error_code: None, + error_message: None, + video_url: Some( + "https://generativelanguage.googleapis.com/v1beta/files/video-1:download?alt=media" + .to_string(), + ), + request_metadata: None, + } + .into_stored() + } + + fn state_with_gemini_transport() -> AppState { + let state = AppState::new().expect("gateway state should build"); + let provider = StoredProviderCatalogProvider::new( + "provider-1".to_string(), + "Gemini".to_string(), + Some("https://ai.google.dev".to_string()), + "gemini".to_string(), + ) + .expect("provider should build"); + let endpoint = StoredProviderCatalogEndpoint::new( + "endpoint-1".to_string(), + "provider-1".to_string(), + "gemini:video".to_string(), + None, + None, + true, + ) + .expect("endpoint should build") + .with_transport_fields( + "https://generativelanguage.googleapis.com".to_string(), + None, + None, + None, + None, + None, + None, + None, + ) + .expect("endpoint transport should build"); + let encrypted_api_key = state + .seal_provider_catalog_key_api_key( + "provider-1", + "provider-key-1", + "gemini-provider-secret", + ) + .expect("provider key should encrypt"); + let key = StoredProviderCatalogKey::new( + "provider-key-1".to_string(), + "provider-1".to_string(), + "default".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("provider key should build") + .with_transport_fields( + Some(json!(["gemini:video"])), + encrypted_api_key, + None, + None, + None, + None, + None, + None, + None, + ) + .expect("provider key transport should build"); + let provider_catalog: Arc = Arc::new( + InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key]), + ); + let video_tasks = Arc::new(InMemoryVideoTaskRepository::default()); + let data = GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + video_tasks, + provider_catalog, + DEVELOPMENT_ENCRYPTION_KEY, + ); + state.with_data_state_for_tests(data) + } + + #[test] + fn video_url_parser_rejects_non_http_and_embedded_credentials() { + for raw_url in [ + "file:///etc/passwd", + "data:video/mp4;base64,AAAA", + "https://user@example.com/video.mp4", + "https://user:password@example.com/video.mp4", + "/relative/video.mp4", + ] { + assert!( + parse_video_url(raw_url).is_err(), + "URL should be rejected: {raw_url}" + ); + } + } + + #[test] + fn video_origin_comparison_uses_scheme_host_and_effective_port() { + let base = parse_video_url("https://generativelanguage.googleapis.com/v1beta").unwrap(); + for same_origin in [ + "https://generativelanguage.googleapis.com/file", + "https://generativelanguage.googleapis.com:443/file", + ] { + assert!(video_urls_share_origin( + &base, + &parse_video_url(same_origin).unwrap() + )); + } + for different_origin in [ + "http://generativelanguage.googleapis.com/file", + "https://generativelanguage.googleapis.com:444/file", + "https://generativelanguage.googleapis.com.evil.test/file", + "https://evil.test/generativelanguage.googleapis.com/file", + ] { + assert!(!video_urls_share_origin( + &base, + &parse_video_url(different_origin).unwrap() + )); + } + } + + #[tokio::test] + async fn legacy_gemini_client_format_uses_authenticated_proxy_source() { + let source = video_task_video_source_from_task( + &state_with_gemini_transport(), + &legacy_gemini_video_task(), + ) + .await + .expect("video source should resolve") + .expect("video source should exist"); + + match source { + VideoTaskVideoSource::Proxy { + url, + header_name, + header_value, + filename, + } => { + assert_eq!( + url.as_str(), + "https://generativelanguage.googleapis.com/v1beta/files/video-1:download?alt=media" + ); + assert_eq!(header_name, "x-goog-api-key"); + assert_eq!(header_value, "gemini-provider-secret"); + assert_eq!(filename, "video_legacy-gemini-task.mp4"); + } + VideoTaskVideoSource::Redirect { .. } => { + panic!("legacy Gemini video must not bypass the authenticated proxy") + } + } + } +} diff --git a/apps/aether-gateway/src/async_task/runtime.rs b/apps/aether-gateway/src/async_task/runtime.rs index 0bb8782b1..d1e4a64b5 100644 --- a/apps/aether-gateway/src/async_task/runtime.rs +++ b/apps/aether-gateway/src/async_task/runtime.rs @@ -20,7 +20,7 @@ const VIDEO_TASK_POLL_CLAIM_SECONDS: u64 = 30; #[derive(Debug, Clone)] struct VideoTaskRefreshError { - message: String, + category: &'static str, permanent: bool, } @@ -55,7 +55,7 @@ pub(crate) async fn execute_video_task_refresh_plan( warn!( event_name = "video_task_refresh_failed", log_type = "event", - error = %err.message, + error_category = err.category, permanent = err.permanent, "gateway video task refresh failed" ); @@ -79,23 +79,32 @@ async fn poll_video_tasks_once(state: &AppState, batch_size: usize) -> Result { - let Some(updated) = - build_successful_poll_update(&task, &provider_body, now_unix_secs)? + let Some(updated) = build_successful_poll_update( + &task, + snapshot.clone(), + &provider_body, + now_unix_secs, + )? else { continue; }; match state.update_active_video_task(updated).await? { Some(stored) => { - if let Some(snapshot) = LocalVideoTaskSnapshot::from_stored_task(&stored) { + if let Some(snapshot) = + state.reconstruct_video_task_snapshot(&stored).await? + { state.video_tasks.record_snapshot(snapshot); } info!( @@ -116,7 +125,9 @@ async fn poll_video_tasks_once(state: &AppState, batch_size: usize) -> Result { - if let Some(snapshot) = LocalVideoTaskSnapshot::from_stored_task(&stored) { + if let Some(snapshot) = + state.reconstruct_video_task_snapshot(&stored).await? + { state.video_tasks.record_snapshot(snapshot); } info!( @@ -190,9 +201,9 @@ async fn fetch_video_task_refresh_attempt( .await { Ok(result) => result, - Err(err) => { + Err(_) => { return Ok(VideoTaskRefreshAttempt::Error(VideoTaskRefreshError { - message: format!("{err:?}"), + category: "transport_error", permanent: false, })); } @@ -209,7 +220,7 @@ async fn fetch_video_task_refresh_attempt( .and_then(|body| body.as_object().cloned()) else { return Ok(VideoTaskRefreshAttempt::Error(VideoTaskRefreshError { - message: "video task refresh missing json provider body".to_string(), + category: "invalid_provider_response", permanent: false, })); }; @@ -223,20 +234,19 @@ fn classify_refresh_result_error(result: &ExecutionResult) -> VideoTaskRefreshEr .as_ref() .and_then(|error| error.upstream_status) .unwrap_or(result.status_code); - let message = result - .error - .as_ref() - .map(|error| error.message.clone()) - .or_else(|| { - result - .body - .as_ref() - .and_then(|body| body.json_body.as_ref()) - .and_then(|value| value.get("error")) - .and_then(Value::as_str) - .map(str::to_string) - }) - .unwrap_or_else(|| format!("upstream returned {status_code}")); + let category = if status_code == 401 { + "authentication_error" + } else if status_code == 403 { + "permission_denied" + } else if status_code == 404 { + "not_found" + } else if status_code == 429 { + "rate_limit" + } else if status_code >= 500 { + "server_error" + } else { + "provider_error" + }; let permanent = result.error.as_ref().map_or( matches!(status_code, 400 | 401 | 403 | 404 | 422), |error| match error.kind { @@ -253,17 +263,18 @@ fn classify_refresh_result_error(result: &ExecutionResult) -> VideoTaskRefreshEr }, ); - VideoTaskRefreshError { message, permanent } + VideoTaskRefreshError { + category, + permanent, + } } fn build_successful_poll_update( task: &StoredVideoTask, + mut snapshot: LocalVideoTaskSnapshot, provider_body: &Map, now_unix_secs: u64, ) -> Result, GatewayError> { - let Some(mut snapshot) = LocalVideoTaskSnapshot::from_stored_task(task) else { - return Ok(None); - }; snapshot.apply_provider_body(provider_body); let mut record = snapshot.to_upsert_record(); @@ -283,10 +294,7 @@ fn build_successful_poll_update( record.format_converted = task.format_converted; record.model = task.model.clone().or(record.model); record.prompt = task.prompt.clone().or(record.prompt); - record.original_request_body = task - .original_request_body - .clone() - .or(record.original_request_body); + record.original_request_body = None; record.duration_seconds = task.duration_seconds.or(record.duration_seconds); record.resolution = task.resolution.clone().or(record.resolution); record.aspect_ratio = task.aspect_ratio.clone().or(record.aspect_ratio); @@ -309,17 +317,11 @@ fn build_successful_poll_update( if record.status.is_active() && record.poll_count >= record.max_poll_count { record.status = VideoTaskStatus::Failed; record.error_code = Some("poll_timeout".to_string()); - record.error_message = Some(format!("Task timed out after {} polls", record.poll_count)); + record.error_message = None; record.completed_at_unix_secs = Some(now_unix_secs); record.next_poll_at_unix_secs = None; } - record.request_metadata = merge_video_task_request_metadata( - task.request_metadata.clone(), - &snapshot, - Some(provider_body), - None, - ) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + record.request_metadata = None; Ok(Some(record)) } @@ -332,11 +334,11 @@ fn build_failed_poll_update( let mut record = stored_task_to_upsert(task); record.updated_at_unix_secs = now_unix_secs; record.poll_count = task.poll_count.saturating_add(1); - record.progress_message = Some(format!("Poll error: {}", err.message)); + record.progress_message = None; if err.permanent { record.status = VideoTaskStatus::Failed; record.error_code = Some("poll_permanent_error".to_string()); - record.error_message = Some(err.message.clone()); + record.error_message = None; record.completed_at_unix_secs = Some(now_unix_secs); record.next_poll_at_unix_secs = None; } else { @@ -348,28 +350,15 @@ fn build_failed_poll_update( if record.status.is_active() && record.poll_count >= record.max_poll_count { record.status = VideoTaskStatus::Failed; record.error_code = Some("poll_timeout".to_string()); - record.error_message = Some(format!("Task timed out after {} polls", record.poll_count)); + record.error_message = None; record.completed_at_unix_secs = Some(now_unix_secs); record.next_poll_at_unix_secs = None; } - record.request_metadata = LocalVideoTaskSnapshot::from_stored_task(task) - .and_then(|snapshot| { - merge_video_task_request_metadata( - task.request_metadata.clone(), - &snapshot, - None, - Some(err), - ) - .ok() - .flatten() - }) - .or(task.request_metadata.clone()); + record.request_metadata = None; record } fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask { - let snapshot_record = - LocalVideoTaskSnapshot::from_stored_task(task).map(|snapshot| snapshot.to_upsert_record()); UpsertVideoTask { id: task.id.clone(), short_id: task.short_id.clone(), @@ -386,39 +375,15 @@ fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask { provider_api_format: task.provider_api_format.clone(), format_converted: task.format_converted, model: task.model.clone(), - prompt: task.prompt.clone().or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.prompt.clone()) - }), - original_request_body: task.original_request_body.clone().or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.original_request_body.clone()) - }), - duration_seconds: task.duration_seconds.or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.duration_seconds) - }), - resolution: task.resolution.clone().or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.resolution.clone()) - }), - aspect_ratio: task.aspect_ratio.clone().or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.aspect_ratio.clone()) - }), - size: task.size.clone().or_else(|| { - snapshot_record - .as_ref() - .and_then(|record| record.size.clone()) - }), + prompt: task.prompt.clone(), + original_request_body: None, + duration_seconds: task.duration_seconds, + resolution: task.resolution.clone(), + aspect_ratio: task.aspect_ratio.clone(), + size: task.size.clone(), status: task.status, progress_percent: task.progress_percent, - progress_message: task.progress_message.clone(), + progress_message: None, retry_count: task.retry_count, poll_interval_seconds: task.poll_interval_seconds.max(1), next_poll_at_unix_secs: task.next_poll_at_unix_secs, @@ -429,9 +394,9 @@ fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask { completed_at_unix_secs: task.completed_at_unix_secs, updated_at_unix_secs: task.updated_at_unix_secs, error_code: task.error_code.clone(), - error_message: task.error_message.clone(), + error_message: None, video_url: task.video_url.clone(), - request_metadata: task.request_metadata.clone(), + request_metadata: None, } } @@ -443,44 +408,6 @@ fn compute_poll_backoff_seconds(poll_interval_seconds: u32, retry_count: u32) -> .min(MAX_VIDEO_TASK_POLL_BACKOFF_SECONDS) } -fn merge_video_task_request_metadata( - existing: Option, - snapshot: &LocalVideoTaskSnapshot, - provider_body: Option<&Map>, - poll_error: Option<&VideoTaskRefreshError>, -) -> Result, serde_json::Error> { - let mut metadata = match existing { - Some(Value::Object(object)) => object, - _ => Map::new(), - }; - metadata.insert( - "rust_owner".to_string(), - Value::String("async_task".to_string()), - ); - metadata.insert( - "rust_local_snapshot".to_string(), - serde_json::to_value(snapshot)?, - ); - if let Some(provider_body) = provider_body { - metadata.insert( - "poll_raw_response".to_string(), - Value::Object(provider_body.clone()), - ); - metadata.remove("poll_error"); - } - if let Some(poll_error) = poll_error { - metadata.insert( - "poll_error".to_string(), - serde_json::json!({ - "message": poll_error.message, - "permanent": poll_error.permanent, - "observed_at_unix_secs": now_unix_secs(), - }), - ); - } - Ok(Some(Value::Object(metadata))) -} - pub(crate) async fn finalize_video_task_if_terminal(state: &AppState, task: &StoredVideoTask) { let Some(event) = build_video_task_terminal_usage_event(task) else { return; @@ -543,9 +470,9 @@ fn build_video_task_terminal_usage_event(task: &StoredVideoTask) -> Option Option, ) { + let sanitized_path_and_query = sanitize_admin_audit_path(path_and_query); let Some(decision) = control_decision else { return; }; @@ -64,9 +65,10 @@ pub(crate) fn emit_admin_audit( }, route_kind, default_target_type(route_family), - path_and_query.to_string(), + sanitized_path_and_query.clone(), ) }; + let target_id = sanitize_admin_audit_target_id(target_id); let (audit_status, log_level) = classify_admin_audit_response(method, response.status()); if log_level == AdminAuditLogLevel::Info { @@ -83,7 +85,7 @@ pub(crate) fn emit_admin_audit( route_family, route_kind, method = %method, - path = %path_and_query, + path = %sanitized_path_and_query, action, target_type, target_id = %target_id, @@ -103,7 +105,7 @@ pub(crate) fn emit_admin_audit( route_family, route_kind, method = %method, - path = %path_and_query, + path = %sanitized_path_and_query, action, target_type, target_id = %target_id, @@ -112,6 +114,17 @@ pub(crate) fn emit_admin_audit( } } +fn sanitize_admin_audit_path(path_and_query: &str) -> String { + crate::middleware::sanitize_access_log_path(path_and_query) +} + +fn sanitize_admin_audit_target_id(target_id: String) -> String { + if target_id.trim_start().starts_with('/') { + return sanitize_admin_audit_path(&target_id); + } + target_id +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum AdminAuditLogLevel { Info, @@ -151,7 +164,10 @@ fn is_admin_read_method(method: &http::Method) -> bool { #[cfg(test)] mod tests { - use super::{classify_admin_audit_response, AdminAuditLogLevel}; + use super::{ + classify_admin_audit_response, sanitize_admin_audit_path, sanitize_admin_audit_target_id, + AdminAuditLogLevel, + }; use axum::http::{Method, StatusCode}; #[test] @@ -169,4 +185,32 @@ mod tests { ("failed", AdminAuditLogLevel::Warn) ); } + + #[test] + fn audit_paths_drop_sensitive_query_values() { + assert_eq!( + sanitize_admin_audit_path( + "/api/admin/providers?token=secret&api_key=live-key&limit=25" + ), + "/api/admin/providers?limit=25" + ); + assert_eq!( + sanitize_admin_audit_path("/install/one-time-secret?view=raw"), + "/install/[redacted]?view=raw" + ); + } + + #[test] + fn path_shaped_audit_targets_drop_sensitive_query_values() { + assert_eq!( + sanitize_admin_audit_target_id( + "/api/admin/monitoring/trace/request-1?token=secret&limit=25".to_string(), + ), + "/api/admin/monitoring/trace/request-1?limit=25" + ); + assert_eq!( + sanitize_admin_audit_target_id("resource-id?literal".to_string()), + "resource-id?literal" + ); + } } diff --git a/apps/aether-gateway/src/audit/http.rs b/apps/aether-gateway/src/audit/http.rs index 8691f7be9..08d30be62 100644 --- a/apps/aether-gateway/src/audit/http.rs +++ b/apps/aether-gateway/src/audit/http.rs @@ -26,7 +26,10 @@ pub(crate) async fn get_request_candidate_trace( .map_err(|err| GatewayError::Internal(err.to_string()).into_response())?; match trace { - Some(trace) => Ok(Json(trace)), + Some(mut trace) => { + trace.sanitize_sensitive_diagnostics(); + Ok(Json(trace)) + } None => Err(( axum::http::StatusCode::NOT_FOUND, Json(json!({ @@ -52,7 +55,10 @@ pub(crate) async fn get_decision_trace( .map_err(|err| GatewayError::Internal(err.to_string()).into_response())?; match trace { - Some(trace) => Ok(Json(trace)), + Some(mut trace) => { + trace.sanitize_sensitive_diagnostics(); + Ok(Json(trace)) + } None => Err(( axum::http::StatusCode::NOT_FOUND, Json(json!({ diff --git a/apps/aether-gateway/src/backup/config.rs b/apps/aether-gateway/src/backup/config.rs index 860df0ecc..a06d75039 100644 --- a/apps/aether-gateway/src/backup/config.rs +++ b/apps/aether-gateway/src/backup/config.rs @@ -5,7 +5,7 @@ use serde_json::{Map, Value}; use super::schedule::{BackupSchedule, BackupScheduleUnit}; use super::scopes::BackupScope; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub(crate) struct S3BackupConfig { pub(crate) enabled: bool, pub(crate) scope: BackupScope, @@ -22,6 +22,28 @@ pub(crate) struct S3BackupConfig { pub(crate) retention_count: u32, } +impl fmt::Debug for S3BackupConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let endpoint_origin = sanitized_endpoint_origin(&self.endpoint); + formatter + .debug_struct("S3BackupConfig") + .field("enabled", &self.enabled) + .field("scope", &self.scope) + .field("endpoint_origin", &endpoint_origin) + .field("region", &self.region) + .field("user_agent", &self.user_agent) + .field("bucket", &self.bucket) + .field("prefix", &self.prefix) + .field("has_access_key_id", &!self.access_key_id.is_empty()) + .field("has_secret_access_key", &!self.secret_access_key.is_empty()) + .field("path_style", &self.path_style) + .field("compression", &self.compression) + .field("schedule", &self.schedule) + .field("retention_count", &self.retention_count) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct BackupConfigError { message: String, @@ -84,6 +106,9 @@ impl S3BackupConfig { "Endpoint(S3 地址)", enabled, )?; + if enabled { + validate_s3_endpoint(&endpoint)?; + } let bucket = required_or_disabled_string(entries, "backup_s3_bucket", "Bucket(存储桶)", enabled)?; let access_key_id = required_or_disabled_string( @@ -99,6 +124,11 @@ impl S3BackupConfig { enabled, )?; + let prefix = normalize_s3_prefix( + &optional_string(entries, "backup_s3_prefix")? + .unwrap_or_else(|| "aether/backups/".to_string()), + )?; + Ok(Self { enabled, scope, @@ -108,8 +138,7 @@ impl S3BackupConfig { user_agent: optional_string(entries, "backup_s3_user_agent")? .unwrap_or_else(|| "rclone/v1.68.0".to_string()), bucket, - prefix: optional_string(entries, "backup_s3_prefix")? - .unwrap_or_else(|| "aether/backups/".to_string()), + prefix, access_key_id, secret_access_key, path_style: optional_bool(entries, "backup_s3_path_style")?.unwrap_or(true), @@ -121,6 +150,48 @@ impl S3BackupConfig { } } +fn normalize_s3_prefix(prefix: &str) -> Result { + let prefix = prefix.trim().trim_matches('/'); + if prefix.is_empty() { + return Ok(String::new()); + } + if prefix + .split('/') + .any(|segment| segment.is_empty() || segment == "." || segment == "..") + || prefix.contains('\\') + { + return Err(BackupConfigError::new( + "Prefix(备份前缀)不能包含空路径段、相对路径段或反斜杠", + )); + } + + Ok(format!("{prefix}/")) +} + +fn validate_s3_endpoint(endpoint: &str) -> Result<(), BackupConfigError> { + let parsed = url::Url::parse(endpoint) + .map_err(|_| BackupConfigError::new("Endpoint(S3 地址)必须是有效的 HTTPS URL"))?; + if parsed.scheme() != "https" + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return Err(BackupConfigError::new( + "Endpoint(S3 地址)必须使用 HTTPS,且不能包含用户凭据、查询参数或片段", + )); + } + Ok(()) +} + +fn sanitized_endpoint_origin(endpoint: &str) -> String { + url::Url::parse(endpoint) + .ok() + .map(|parsed| parsed.origin().ascii_serialization()) + .unwrap_or_else(|| "".to_string()) +} + fn validate_range(label: &str, value: u32, min: u32, max: u32) -> Result<(), BackupConfigError> { if (min..=max).contains(&value) { Ok(()) @@ -374,6 +445,69 @@ mod tests { assert!(err.to_string().contains("Endpoint")); } + #[test] + fn rejects_insecure_or_credential_bearing_endpoints() { + for endpoint in [ + "http://s3.example.com", + "https://user:password@s3.example.com", + "https://s3.example.com?token=secret", + "https://s3.example.com/#fragment", + ] { + let entries = serde_json::json!({ + "backup_s3_enabled": true, + "backup_s3_endpoint": endpoint, + "backup_s3_bucket": "aether-backups", + "backup_s3_access_key_id": "access", + "backup_s3_secret_access_key": "secret" + }); + + let error = S3BackupConfig::from_json_map(entries.as_object().unwrap()) + .expect_err("unsafe endpoint should fail closed"); + assert!(error.to_string().contains("Endpoint")); + } + } + + #[test] + fn debug_output_does_not_expose_s3_credentials() { + let entries = serde_json::json!({ + "backup_s3_enabled": true, + "backup_s3_endpoint": "https://s3.example.com/path", + "backup_s3_bucket": "aether-backups", + "backup_s3_access_key_id": "access-key-value", + "backup_s3_secret_access_key": "secret-key-value" + }); + let config = S3BackupConfig::from_json_map(entries.as_object().unwrap()) + .expect("config should parse"); + + let debug = format!("{config:?}"); + assert!(debug.contains("https://s3.example.com")); + assert!(!debug.contains("/path")); + assert!(!debug.contains("access-key-value")); + assert!(!debug.contains("secret-key-value")); + } + + #[test] + fn canonicalizes_s3_backup_prefix_once() { + let entries = serde_json::json!({ + "backup_s3_enabled": true, + "backup_s3_endpoint": "https://s3.example.com", + "backup_s3_bucket": "aether-backups", + "backup_s3_prefix": "/prod/backups//", + "backup_s3_access_key_id": "access", + "backup_s3_secret_access_key": "secret" + }); + let config = S3BackupConfig::from_json_map(entries.as_object().unwrap()) + .expect("prefix should be canonicalized"); + + assert_eq!(config.prefix, "prod/backups/"); + + for invalid_prefix in ["prod//backups", "prod/../backups", "prod\\backups"] { + let mut entries = entries.clone(); + entries["backup_s3_prefix"] = serde_json::json!(invalid_prefix); + assert!(S3BackupConfig::from_json_map(entries.as_object().unwrap()).is_err()); + } + } + #[test] fn applies_default_values_from_system_config_contract() { let entries = serde_json::json!({ diff --git a/apps/aether-gateway/src/backup/executor.rs b/apps/aether-gateway/src/backup/executor.rs index ee6df856e..a7bc20baf 100644 --- a/apps/aether-gateway/src/backup/executor.rs +++ b/apps/aether-gateway/src/backup/executor.rs @@ -1,11 +1,189 @@ +use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload}; +use aes_gcm::Aes256Gcm; +use aether_crypto::derive_python_fernet_key; use bytes::Bytes; use chrono::{DateTime, SecondsFormat, Utc}; +use hmac::{Hmac, Mac}; +use serde::Deserialize; use serde_json::Value; use sha2::{Digest, Sha256}; +use std::io::Read; use super::config::S3BackupConfig; use super::scopes::BackupScope; -use super::store::{BackupObjectStore, BackupStoreError}; +use super::store::{BackupObjectCreateResult, BackupObjectStore, BackupStoreError}; +use super::BackupRestoreScope; + +// V1: magic || version || nonce || ciphertext/tag. +// V2: magic || version || fixed-length key ID || nonce || ciphertext/tag. +const BACKUP_ENVELOPE_MAGIC: &[u8; 8] = b"AETHERBK"; +const BACKUP_ENVELOPE_VERSION_V1: u8 = 1; +const BACKUP_ENVELOPE_VERSION_V2: u8 = 2; +const BACKUP_ENVELOPE_KEY_ID_LEN: usize = 16; +const BACKUP_ENVELOPE_NONCE_LEN: usize = 12; +const BACKUP_ENCRYPTION_CONTEXT_V1: &[u8] = b"aether-s3-backup-aes-256-gcm-v1"; +const BACKUP_ENCRYPTION_CONTEXT_V2: &[u8] = b"aether-s3-backup-aes-256-gcm-v2"; +const BACKUP_KEY_ID_CONTEXT: &[u8] = b"aether-s3-backup-key-id-v1"; +const BACKUP_ENCRYPTION_NAME: &str = "aes-256-gcm-v2"; +const MAX_LEGACY_V1_CANDIDATES: usize = 16; +const MAX_V2_CANDIDATES: usize = 256; +const MAX_BACKUP_OBJECTS_PER_PREFIX: usize = 10_000; +const MAX_ENCRYPTED_VARIANTS_PER_LEGACY_OBJECT: usize = 32; +const ZSTD_MAX_WINDOW_LOG: u32 = 27; +const MIN_NEW_BACKUP_ENCRYPTION_SECRET_BYTES: usize = 32; +const INSECURE_NEW_BACKUP_ENCRYPTION_SECRETS: &[&str] = &[ + "change-this-to-another-secure-random-string", + "change-this-to-a-secure-random-string", + "dev-encryption-key-do-not-use-in-production", +]; +pub const DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES: usize = 512 * 1024 * 1024; +pub const DEFAULT_BACKUP_MAX_JSON_BYTES: usize = 1024 * 1024 * 1024; +type HmacSha256 = Hmac; + +pub struct BackupDecryptionKey { + secret: String, + v2_key_id: [u8; BACKUP_ENVELOPE_KEY_ID_LEN], + allow_legacy_v1: bool, +} + +impl BackupDecryptionKey { + pub fn current(secret: impl Into) -> Result { + Self::new(secret, false) + } + + pub fn historical(secret: impl Into) -> Result { + Self::new(secret, true) + } + + pub fn v2_only(secret: impl Into) -> Result { + Self::new(secret, false) + } + + fn new(secret: impl Into, allow_legacy_v1: bool) -> Result { + let secret = secret.into(); + if secret.trim().is_empty() { + return Err(BackupRestoreError::InvalidKeyMaterial); + } + let v2_key = derive_backup_encryption_key(&secret, BACKUP_ENCRYPTION_CONTEXT_V2) + .map_err(|_| BackupRestoreError::InvalidKeyMaterial)?; + let v2_key_id = + derive_backup_key_id(&v2_key).map_err(|_| BackupRestoreError::InvalidKeyMaterial)?; + Ok(Self { + secret, + v2_key_id, + allow_legacy_v1, + }) + } +} + +impl std::fmt::Debug for BackupDecryptionKey { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("BackupDecryptionKey") + .field("secret", &"[REDACTED]") + .field("v2_key_id", &encode_key_id(&self.v2_key_id)) + .field("allow_legacy_v1", &self.allow_legacy_v1) + .finish() + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct BackupRestoreLimits { + pub max_encrypted_bytes: usize, + pub max_json_bytes: usize, +} + +impl Default for BackupRestoreLimits { + fn default() -> Self { + Self { + max_encrypted_bytes: DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES, + max_json_bytes: DEFAULT_BACKUP_MAX_JSON_BYTES, + } + } +} + +#[derive(Debug)] +pub struct RestoredBackupJson { + json_bytes: Vec, + pub envelope_version: u8, + pub key_id: Option, + pub export_version: Option, + pub exported_at: Option, + restore_authority: BackupRestoreAuthority, +} + +#[derive(Debug)] +pub(crate) struct BackupRestoreAuthority { + scope: BackupRestoreScope, +} + +impl BackupRestoreAuthority { + fn new(scope: BackupRestoreScope) -> Self { + Self { scope } + } + + pub(crate) fn scope(&self) -> BackupRestoreScope { + self.scope + } +} + +impl RestoredBackupJson { + pub fn json_bytes(&self) -> &[u8] { + &self.json_bytes + } + + pub fn scope(&self) -> BackupRestoreScope { + self.restore_authority.scope() + } + + pub(crate) fn into_authenticated_parts(self) -> (Vec, BackupRestoreAuthority) { + (self.json_bytes, self.restore_authority) + } +} + +#[derive(Debug, Deserialize)] +struct BackupJsonMetadata { + #[serde(default)] + version: Option, + #[serde(default)] + exported_at: Option, +} + +#[derive(Debug, thiserror::Error)] +pub enum BackupRestoreError { + #[error("backup envelope is invalid or unsupported")] + InvalidEnvelope, + + #[error("backup object key is invalid")] + InvalidObjectKey, + + #[error("backup encryption key is empty or invalid")] + InvalidKeyMaterial, + + #[error("no configured backup key matches key ID {0}")] + UnknownKeyId(String), + + #[error("backup authentication failed")] + AuthenticationFailed, + + #[error("legacy v1 backup restore allows at most 16 candidate keys")] + TooManyLegacyKeys, + + #[error("backup restore allows at most 256 v2 candidate keys")] + TooManyV2Keys, + + #[error("encrypted backup exceeds the configured {limit} byte limit")] + EncryptedSizeLimit { limit: usize }, + + #[error("decompressed backup JSON exceeds the configured {limit} byte limit")] + JsonSizeLimit { limit: usize }, + + #[error("backup zstd decompression failed: {0}")] + Decompression(#[source] std::io::Error), + + #[error("decompressed backup is not valid JSON: {0}")] + InvalidJson(#[source] serde_json::Error), +} #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct BackupRunResult { @@ -17,7 +195,22 @@ pub(crate) struct BackupRunResult { pub(crate) export_version: String, pub(crate) exported_at: String, pub(crate) compression: String, - pub(crate) deleted_old_objects: usize, + pub(crate) encryption: String, + pub(crate) key_id: String, + pub(crate) legacy_encrypted_copies_created: usize, + pub(crate) legacy_encrypted_copies_verified: usize, + pub(crate) legacy_plaintext_objects_deleted: usize, + pub(crate) legacy_plaintext_objects_retained: usize, + pub(crate) retention_cleanup_candidates: usize, + pub(crate) versioned_storage_cleanup_required: bool, +} + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +struct LegacyMigrationResult { + encrypted_copies_created: usize, + encrypted_copies_verified: usize, + plaintext_objects_deleted: usize, + plaintext_objects_retained: usize, } #[derive(Debug, thiserror::Error)] @@ -28,11 +221,22 @@ pub(crate) enum BackupExecutionError { #[error("S3 backup compression failed: {0}")] Compression(#[from] std::io::Error), + #[error("S3 backup encryption failed")] + Encryption, + + #[error("S3 backup encryption key is unsafe for creating new backups: {0}")] + UnsafeEncryptionKey(&'static str), + #[error("{0}")] Store(#[from] BackupStoreError), #[error("S3 backup compression `{0}` is not supported; expected `zstd`")] UnknownCompression(String), + + #[error( + "legacy backup `{key}` has more than {limit} encrypted variants; refusing unbounded verification" + )] + TooManyEncryptedVariants { key: String, limit: usize }, } pub(crate) async fn run_backup_with_store( @@ -40,28 +244,38 @@ pub(crate) async fn run_backup_with_store( store: &S, payload: Value, now_utc: DateTime, + encryption_secret: &str, ) -> Result where S: BackupObjectStore + ?Sized, { + validate_new_backup_encryption_secret(encryption_secret)?; let export_version = payload_string_field(&payload, "version").unwrap_or_default(); let exported_at = payload_string_field(&payload, "exported_at") .unwrap_or_else(|| now_utc.to_rfc3339_opts(SecondsFormat::Secs, true)); let json_bytes = serde_json::to_vec(&payload)?; let compression = config.compression.trim().to_string(); - let upload_bytes = match compression.as_str() { + let compressed_bytes = match compression.as_str() { "zstd" => zstd::stream::encode_all(json_bytes.as_slice(), 0)?, other => return Err(BackupExecutionError::UnknownCompression(other.to_string())), }; + let timestamp = now_utc.format("%Y%m%d-%H%M%S").to_string(); + let preferred_object_key = config.scope.object_key(&config.prefix, ×tamp); + let (upload_bytes, key_id) = + encrypt_backup_bytes(encryption_secret, &preferred_object_key, &compressed_bytes)?; let bytes = upload_bytes.len(); let sha256 = format!("{:x}", Sha256::digest(&upload_bytes)); - let timestamp = now_utc.format("%Y%m%d-%H%M%S").to_string(); - let object_key = config.scope.object_key(&config.prefix, ×tamp); - store - .put_object(&object_key, Bytes::from(upload_bytes)) - .await?; - let deleted_old_objects = prune_old_backups(config, store, &object_key).await?; + let object_key = put_encrypted_object_without_overwrite( + store, + &preferred_object_key, + Bytes::from(upload_bytes), + ) + .await?; + let legacy_migration = + migrate_legacy_plaintext_backups(config, store, encryption_secret).await?; + let retention_cleanup_candidates = + count_retention_cleanup_candidates(config, store, &object_key).await?; Ok(BackupRunResult { scope: config.scope, @@ -72,11 +286,462 @@ where export_version, exported_at, compression, - deleted_old_objects, + encryption: BACKUP_ENCRYPTION_NAME.to_string(), + key_id: encode_key_id(&key_id), + legacy_encrypted_copies_created: legacy_migration.encrypted_copies_created, + legacy_encrypted_copies_verified: legacy_migration.encrypted_copies_verified, + legacy_plaintext_objects_deleted: legacy_migration.plaintext_objects_deleted, + legacy_plaintext_objects_retained: legacy_migration.plaintext_objects_retained, + retention_cleanup_candidates, + versioned_storage_cleanup_required: legacy_migration.plaintext_objects_deleted > 0 + || legacy_migration.plaintext_objects_retained > 0 + || retention_cleanup_candidates > 0, }) } -async fn prune_old_backups( +fn validate_new_backup_encryption_secret( + encryption_secret: &str, +) -> Result<(), BackupExecutionError> { + let encryption_secret = encryption_secret.trim(); + if encryption_secret.as_bytes().len() < MIN_NEW_BACKUP_ENCRYPTION_SECRET_BYTES { + return Err(BackupExecutionError::UnsafeEncryptionKey( + "must contain at least 32 bytes", + )); + } + if INSECURE_NEW_BACKUP_ENCRYPTION_SECRETS.contains(&encryption_secret) { + return Err(BackupExecutionError::UnsafeEncryptionKey( + "must not use a published example or development value", + )); + } + Ok(()) +} + +pub(crate) fn encrypt_backup_bytes( + encryption_secret: &str, + object_key: &str, + plaintext: &[u8], +) -> Result<(Vec, [u8; BACKUP_ENVELOPE_KEY_ID_LEN]), BackupExecutionError> { + let key = derive_backup_encryption_key(encryption_secret, BACKUP_ENCRYPTION_CONTEXT_V2)?; + let key_id = derive_backup_key_id(&key)?; + let cipher = Aes256Gcm::new_from_slice(&key).map_err(|_| BackupExecutionError::Encryption)?; + let nonce = Aes256Gcm::generate_nonce(&mut OsRng); + let aad = backup_envelope_aad(BACKUP_ENVELOPE_VERSION_V2, Some(&key_id), object_key)?; + let ciphertext = cipher + .encrypt( + &nonce, + Payload { + msg: plaintext, + aad: &aad, + }, + ) + .map_err(|_| BackupExecutionError::Encryption)?; + + let envelope_header_len = BACKUP_ENVELOPE_MAGIC.len() + 1 + BACKUP_ENVELOPE_KEY_ID_LEN; + let mut envelope = Vec::with_capacity(envelope_header_len + nonce.len() + ciphertext.len()); + envelope.extend_from_slice(&aad[..envelope_header_len]); + envelope.extend_from_slice(&nonce); + envelope.extend_from_slice(&ciphertext); + Ok((envelope, key_id)) +} + +fn derive_backup_encryption_key( + encryption_secret: &str, + context: &[u8], +) -> Result<[u8; 32], BackupExecutionError> { + let encryption_secret = encryption_secret.trim(); + if encryption_secret.is_empty() { + return Err(BackupExecutionError::Encryption); + } + let root_key = derive_python_fernet_key(encryption_secret); + let mut mac = ::new_from_slice(root_key.as_bytes()) + .map_err(|_| BackupExecutionError::Encryption)?; + mac.update(context); + Ok(mac.finalize().into_bytes().into()) +} + +fn derive_backup_key_id( + encryption_key: &[u8; 32], +) -> Result<[u8; BACKUP_ENVELOPE_KEY_ID_LEN], BackupExecutionError> { + let mut mac = ::new_from_slice(encryption_key) + .map_err(|_| BackupExecutionError::Encryption)?; + mac.update(BACKUP_KEY_ID_CONTEXT); + let digest = mac.finalize().into_bytes(); + let mut key_id = [0_u8; BACKUP_ENVELOPE_KEY_ID_LEN]; + key_id.copy_from_slice(&digest[..BACKUP_ENVELOPE_KEY_ID_LEN]); + Ok(key_id) +} + +fn backup_envelope_aad( + version: u8, + key_id: Option<&[u8; BACKUP_ENVELOPE_KEY_ID_LEN]>, + object_key: &str, +) -> Result, BackupExecutionError> { + let object_key = canonical_encrypted_object_key(object_key)?; + let key_id_len = key_id.map_or(0, |_| BACKUP_ENVELOPE_KEY_ID_LEN); + let mut aad = + Vec::with_capacity(BACKUP_ENVELOPE_MAGIC.len() + 1 + key_id_len + object_key.len()); + aad.extend_from_slice(BACKUP_ENVELOPE_MAGIC); + aad.push(version); + if let Some(key_id) = key_id { + aad.extend_from_slice(key_id); + } + aad.extend_from_slice(object_key.as_bytes()); + Ok(aad) +} + +fn encode_key_id(key_id: &[u8; BACKUP_ENVELOPE_KEY_ID_LEN]) -> String { + key_id.iter().map(|byte| format!("{byte:02x}")).collect() +} + +async fn put_encrypted_object_without_overwrite( + store: &S, + preferred_key: &str, + bytes: Bytes, +) -> Result +where + S: BackupObjectStore + ?Sized, +{ + match store + .put_object_if_absent(preferred_key, bytes.clone()) + .await + { + Ok(BackupObjectCreateResult::Created) => return Ok(preferred_key.to_string()), + Ok(BackupObjectCreateResult::AlreadyExists) => {} + Err(error) => return Err(error.into()), + } + + let collision_key = collision_safe_encrypted_key(preferred_key, &bytes)?; + match store + .put_object_if_absent(&collision_key, bytes.clone()) + .await + { + Ok(BackupObjectCreateResult::Created) => Ok(collision_key), + Ok(BackupObjectCreateResult::AlreadyExists) => { + let existing = store + .get_object_limited(&collision_key, bytes.len()) + .await?; + if existing == bytes { + Ok(collision_key) + } else { + Err(BackupExecutionError::Encryption) + } + } + Err(error) => Err(error.into()), + } +} + +fn collision_safe_encrypted_key( + preferred_key: &str, + encrypted_bytes: &[u8], +) -> Result { + const SUFFIX: &str = ".json.zst.aes256gcm"; + let Some(stem) = preferred_key.strip_suffix(SUFFIX) else { + return Err(BackupExecutionError::Encryption); + }; + let digest = format!("{:x}", Sha256::digest(encrypted_bytes)); + Ok(format!("{stem}-{digest}{SUFFIX}")) +} + +fn canonical_encrypted_object_key(object_key: &str) -> Result { + const SUFFIX: &str = ".json.zst.aes256gcm"; + if BackupScope::from_encrypted_object_key(object_key).is_none() { + return Err(BackupExecutionError::Encryption); + } + let Some(stem) = object_key.strip_suffix(SUFFIX) else { + return Err(BackupExecutionError::Encryption); + }; + let canonical_stem = stem + .rsplit_once('-') + .filter(|(_, digest)| { + digest.len() == 64 + && digest + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) + }) + .map(|(canonical_stem, _)| canonical_stem) + .unwrap_or(stem); + Ok(format!("{canonical_stem}{SUFFIX}")) +} + +fn decrypt_backup_bytes( + encryption_secret: &str, + object_key: &str, + envelope: &[u8], +) -> Result, BackupExecutionError> { + let version = envelope + .get(BACKUP_ENVELOPE_MAGIC.len()) + .copied() + .ok_or(BackupExecutionError::Encryption)?; + let (context, key_id_len) = match version { + BACKUP_ENVELOPE_VERSION_V1 => (BACKUP_ENCRYPTION_CONTEXT_V1, 0), + BACKUP_ENVELOPE_VERSION_V2 => (BACKUP_ENCRYPTION_CONTEXT_V2, BACKUP_ENVELOPE_KEY_ID_LEN), + _ => return Err(BackupExecutionError::Encryption), + }; + let envelope_header_len = BACKUP_ENVELOPE_MAGIC.len() + 1 + key_id_len; + let header_len = envelope_header_len + BACKUP_ENVELOPE_NONCE_LEN; + if envelope.len() <= header_len || !envelope.starts_with(BACKUP_ENVELOPE_MAGIC) { + return Err(BackupExecutionError::Encryption); + } + let key_id = if version == BACKUP_ENVELOPE_VERSION_V2 { + let mut key_id = [0_u8; BACKUP_ENVELOPE_KEY_ID_LEN]; + key_id.copy_from_slice(&envelope[BACKUP_ENVELOPE_MAGIC.len() + 1..envelope_header_len]); + Some(key_id) + } else { + None + }; + let aad = backup_envelope_aad(version, key_id.as_ref(), object_key)?; + if envelope.get(..envelope_header_len) != Some(&aad[..envelope_header_len]) { + return Err(BackupExecutionError::Encryption); + } + let key = derive_backup_encryption_key(encryption_secret, context)?; + if key_id.is_some_and(|expected| derive_backup_key_id(&key).ok() != Some(expected)) { + return Err(BackupExecutionError::Encryption); + } + let cipher = Aes256Gcm::new_from_slice(&key).map_err(|_| BackupExecutionError::Encryption)?; + let nonce = aes_gcm::Nonce::from_slice(&envelope[envelope_header_len..header_len]); + cipher + .decrypt( + nonce, + Payload { + msg: &envelope[header_len..], + aad: &aad, + }, + ) + .map_err(|_| BackupExecutionError::Encryption) +} + +pub fn restore_backup_json( + object_key: &str, + envelope: &[u8], + candidates: &[BackupDecryptionKey], + limits: BackupRestoreLimits, +) -> Result { + if envelope.len() > limits.max_encrypted_bytes { + return Err(BackupRestoreError::EncryptedSizeLimit { + limit: limits.max_encrypted_bytes, + }); + } + if candidates.len() > MAX_V2_CANDIDATES { + return Err(BackupRestoreError::TooManyV2Keys); + } + let scope = BackupScope::from_encrypted_object_key(object_key) + .ok_or(BackupRestoreError::InvalidObjectKey)?; + canonical_encrypted_object_key(object_key).map_err(|_| BackupRestoreError::InvalidObjectKey)?; + if envelope.len() <= BACKUP_ENVELOPE_MAGIC.len() || !envelope.starts_with(BACKUP_ENVELOPE_MAGIC) + { + return Err(BackupRestoreError::InvalidEnvelope); + } + + let version = envelope[BACKUP_ENVELOPE_MAGIC.len()]; + let (compressed, key_id) = match version { + BACKUP_ENVELOPE_VERSION_V2 => { + let key_id_start = BACKUP_ENVELOPE_MAGIC.len() + 1; + let key_id_end = key_id_start + BACKUP_ENVELOPE_KEY_ID_LEN; + let key_id_bytes = envelope + .get(key_id_start..key_id_end) + .ok_or(BackupRestoreError::InvalidEnvelope)?; + let mut key_id = [0_u8; BACKUP_ENVELOPE_KEY_ID_LEN]; + key_id.copy_from_slice(key_id_bytes); + let encoded_key_id = encode_key_id(&key_id); + let candidate = candidates + .iter() + .find(|candidate| candidate.v2_key_id == key_id) + .ok_or_else(|| BackupRestoreError::UnknownKeyId(encoded_key_id.clone()))?; + let compressed = decrypt_backup_bytes(&candidate.secret, object_key, envelope) + .map_err(|_| BackupRestoreError::AuthenticationFailed)?; + (compressed, Some(encoded_key_id)) + } + BACKUP_ENVELOPE_VERSION_V1 => { + let legacy_candidates: Vec<_> = candidates + .iter() + .filter(|candidate| candidate.allow_legacy_v1) + .collect(); + if legacy_candidates.len() > MAX_LEGACY_V1_CANDIDATES { + return Err(BackupRestoreError::TooManyLegacyKeys); + } + let compressed = legacy_candidates + .into_iter() + .find_map(|candidate| { + decrypt_backup_bytes(&candidate.secret, object_key, envelope).ok() + }) + .ok_or(BackupRestoreError::AuthenticationFailed)?; + (compressed, None) + } + _ => return Err(BackupRestoreError::InvalidEnvelope), + }; + + let json_bytes = decompress_backup_json(&compressed, limits.max_json_bytes)?; + let mut deserializer = serde_json::Deserializer::from_slice(&json_bytes); + let metadata = BackupJsonMetadata::deserialize(&mut deserializer) + .map_err(BackupRestoreError::InvalidJson)?; + deserializer + .end() + .map_err(BackupRestoreError::InvalidJson)?; + if metadata + .version + .as_deref() + .map(str::trim) + .is_none_or(str::is_empty) + || metadata + .exported_at + .as_deref() + .map(str::trim) + .is_none_or(str::is_empty) + { + return Err(BackupRestoreError::InvalidJson(serde_json::Error::io( + std::io::Error::new( + std::io::ErrorKind::InvalidData, + "backup JSON must include non-empty version and exported_at fields", + ), + ))); + } + Ok(RestoredBackupJson { + json_bytes, + envelope_version: version, + key_id, + export_version: metadata.version, + exported_at: metadata.exported_at, + restore_authority: BackupRestoreAuthority::new(match scope { + BackupScope::Config => BackupRestoreScope::Config, + BackupScope::Users => BackupRestoreScope::Users, + BackupScope::Data => BackupRestoreScope::Data, + }), + }) +} + +fn decompress_backup_json( + compressed: &[u8], + max_json_bytes: usize, +) -> Result, BackupRestoreError> { + let mut decoder = + zstd::stream::read::Decoder::new(compressed).map_err(BackupRestoreError::Decompression)?; + decoder + .window_log_max(ZSTD_MAX_WINDOW_LOG) + .map_err(BackupRestoreError::Decompression)?; + let read_limit = u64::try_from(max_json_bytes) + .unwrap_or(u64::MAX) + .saturating_add(1); + let mut limited = decoder.take(read_limit); + let mut json_bytes = Vec::with_capacity(max_json_bytes.min(8 * 1024 * 1024)); + limited + .read_to_end(&mut json_bytes) + .map_err(BackupRestoreError::Decompression)?; + if json_bytes.len() > max_json_bytes { + return Err(BackupRestoreError::JsonSizeLimit { + limit: max_json_bytes, + }); + } + Ok(json_bytes) +} + +async fn migrate_legacy_plaintext_backups( + config: &S3BackupConfig, + store: &S, + encryption_secret: &str, +) -> Result +where + S: BackupObjectStore + ?Sized, +{ + let keys = store + .list_keys_limited(&config.prefix, MAX_BACKUP_OBJECTS_PER_PREFIX) + .await?; + let mut legacy_plaintext_keys = Vec::new(); + let mut encrypted_keys = Vec::new(); + for scope in [BackupScope::Config, BackupScope::Users, BackupScope::Data] { + legacy_plaintext_keys.extend( + scope.matching_legacy_plaintext_backup_keys(&config.prefix, keys.iter().cloned()), + ); + encrypted_keys + .extend(scope.matching_encrypted_backup_keys(&config.prefix, keys.iter().cloned())); + } + legacy_plaintext_keys.sort(); + legacy_plaintext_keys.dedup(); + encrypted_keys.sort(); + encrypted_keys.dedup(); + + let mut result = LegacyMigrationResult::default(); + for key in legacy_plaintext_keys { + let plaintext = store + .get_object_limited(&key, DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES) + .await?; + let candidates: Vec<_> = encrypted_keys + .iter() + .filter(|encrypted_key| is_encrypted_variant_of_legacy_key(&key, encrypted_key)) + .collect(); + if candidates.len() > MAX_ENCRYPTED_VARIANTS_PER_LEGACY_OBJECT { + return Err(BackupExecutionError::TooManyEncryptedVariants { + key, + limit: MAX_ENCRYPTED_VARIANTS_PER_LEGACY_OBJECT, + }); + } + let mut matching_encrypted_copy_exists = false; + for candidate in candidates { + let existing = store + .get_object_limited(candidate, DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES) + .await?; + if decrypt_backup_bytes(encryption_secret, candidate, &existing) + .is_ok_and(|decrypted| decrypted.as_slice() == &plaintext[..]) + { + matching_encrypted_copy_exists = true; + break; + } + } + + if matching_encrypted_copy_exists { + result.encrypted_copies_verified += 1; + store.delete_object(&key).await?; + result.plaintext_objects_deleted += 1; + continue; + } + + let encrypted_key = format!("{key}.aes256gcm"); + let (encrypted, _) = encrypt_backup_bytes(encryption_secret, &encrypted_key, &plaintext)?; + let created_key = + put_encrypted_object_without_overwrite(store, &encrypted_key, Bytes::from(encrypted)) + .await?; + let stored_encrypted = store + .get_object_limited(&created_key, DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES) + .await?; + let verified_plaintext = + decrypt_backup_bytes(encryption_secret, &created_key, stored_encrypted.as_ref()) + .map_err(|_| BackupExecutionError::Encryption)?; + if verified_plaintext.as_slice() != plaintext.as_ref() { + return Err(BackupExecutionError::Encryption); + } + encrypted_keys.push(created_key); + result.encrypted_copies_created += 1; + store.delete_object(&key).await?; + result.plaintext_objects_deleted += 1; + } + Ok(result) +} + +fn is_encrypted_variant_of_legacy_key(legacy_key: &str, encrypted_key: &str) -> bool { + const LEGACY_SUFFIX: &str = ".json.zst"; + const ENCRYPTED_SUFFIX: &str = ".json.zst.aes256gcm"; + + if encrypted_key == format!("{legacy_key}.aes256gcm") { + return true; + } + + let Some(legacy_stem) = legacy_key.strip_suffix(LEGACY_SUFFIX) else { + return false; + }; + let Some(digest) = encrypted_key + .strip_prefix(legacy_stem) + .and_then(|rest| rest.strip_prefix('-')) + .and_then(|rest| rest.strip_suffix(ENCRYPTED_SUFFIX)) + else { + return false; + }; + + digest.len() == 64 + && digest + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) +} + +async fn count_retention_cleanup_candidates( config: &S3BackupConfig, store: &S, current_object_key: &str, @@ -84,11 +749,15 @@ async fn prune_old_backups( where S: BackupObjectStore + ?Sized, { - let keys = store.list_keys(&config.prefix).await?; - let mut matching_keys = config.scope.matching_backup_keys(&config.prefix, keys); + let keys = store + .list_keys_limited(&config.prefix, MAX_BACKUP_OBJECTS_PER_PREFIX) + .await?; + let mut matching_keys = config + .scope + .matching_encrypted_backup_keys(&config.prefix, keys); matching_keys.sort_by(|left, right| right.cmp(left)); - let mut deleted = 0; + let mut cleanup_candidates = 0; let mut retained = usize::from( config.retention_count > 0 && matching_keys.iter().any(|key| key == current_object_key), ); @@ -101,11 +770,10 @@ where continue; } - store.delete_object(&key).await?; - deleted += 1; + cleanup_candidates += 1; } - Ok(deleted) + Ok(cleanup_candidates) } fn payload_string_field(payload: &Value, field: &str) -> Option { @@ -121,11 +789,50 @@ mod tests { use super::super::schedule::BackupSchedule; use super::super::scopes::BackupScope; use super::super::store::{BackupObjectStore, FakeBackupObjectStore}; - use super::run_backup_with_store; + use super::{ + backup_envelope_aad, decrypt_backup_bytes, derive_backup_encryption_key, + encrypt_backup_bytes, restore_backup_json, run_backup_with_store, + validate_new_backup_encryption_secret, BackupDecryptionKey, BackupRestoreError, + BackupRestoreLimits, BACKUP_ENCRYPTION_CONTEXT_V1, BACKUP_ENVELOPE_MAGIC, + BACKUP_ENVELOPE_NONCE_LEN, BACKUP_ENVELOPE_VERSION_V1, BACKUP_ENVELOPE_VERSION_V2, + }; + use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload}; + use aes_gcm::Aes256Gcm; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use bytes::Bytes; use chrono::{DateTime, Utc}; use serde_json::json; + const TEST_NEW_BACKUP_ENCRYPTION_SECRET: &str = "test-only-backup-encryption-secret-2026-08-30"; + + fn compressed_json(value: serde_json::Value) -> Vec { + zstd::stream::encode_all(serde_json::to_vec(&value).unwrap().as_slice(), 0).unwrap() + } + + fn encrypt_v1_for_test(secret: &str, object_key: &str, plaintext: &[u8]) -> Vec { + let key = derive_backup_encryption_key(secret, BACKUP_ENCRYPTION_CONTEXT_V1).unwrap(); + let cipher = Aes256Gcm::new_from_slice(&key).unwrap(); + let nonce = Aes256Gcm::generate_nonce(&mut OsRng); + let aad = backup_envelope_aad(BACKUP_ENVELOPE_VERSION_V1, None, object_key).unwrap(); + let ciphertext = cipher + .encrypt( + &nonce, + Payload { + msg: plaintext, + aad: &aad, + }, + ) + .unwrap(); + let mut envelope = Vec::with_capacity( + BACKUP_ENVELOPE_MAGIC.len() + 1 + BACKUP_ENVELOPE_NONCE_LEN + ciphertext.len(), + ); + envelope.extend_from_slice(BACKUP_ENVELOPE_MAGIC); + envelope.push(BACKUP_ENVELOPE_VERSION_V1); + envelope.extend_from_slice(&nonce); + envelope.extend_from_slice(&ciphertext); + envelope + } + fn sample_backup_config(scope: BackupScope, retention_count: u32) -> S3BackupConfig { S3BackupConfig { enabled: true, @@ -145,7 +852,7 @@ mod tests { } #[tokio::test] - async fn backup_executor_uploads_payload_and_prunes_same_scope_only() { + async fn backup_executor_encrypts_then_deletes_legacy_plaintext_objects() { let store = FakeBackupObjectStore::default(); store .put_object( @@ -173,33 +880,173 @@ mod tests { .unwrap() .with_timezone(&Utc); - let result = run_backup_with_store(&config, &store, payload, now_utc) - .await - .expect("backup should succeed"); + let result = run_backup_with_store( + &config, + &store, + payload, + now_utc, + TEST_NEW_BACKUP_ENCRYPTION_SECRET, + ) + .await + .expect("backup should succeed"); assert_eq!(result.scope, BackupScope::Data); assert_eq!(result.bucket, "aether-backups"); assert_eq!( result.object_key, - "prod/aether-data-backup-20260523-191500.json.zst" + "prod/aether-data-backup-20260523-191500.json.zst.aes256gcm" ); assert!(result.bytes > 0); assert_eq!(result.sha256.len(), 64); assert_eq!(result.export_version, "1.0"); assert_eq!(result.exported_at, "2026-05-24T03:15:00Z"); assert_eq!(result.compression, "zstd"); - assert_eq!(result.deleted_old_objects, 1); + assert_eq!(result.encryption, "aes-256-gcm-v2"); + assert_eq!(result.key_id.len(), 32); + assert!(result.key_id.bytes().all(|byte| byte.is_ascii_hexdigit())); + assert_eq!(result.legacy_encrypted_copies_created, 2); + assert_eq!(result.legacy_encrypted_copies_verified, 0); + assert_eq!(result.legacy_plaintext_objects_deleted, 2); + assert_eq!(result.legacy_plaintext_objects_retained, 0); + assert_eq!(result.retention_cleanup_candidates, 1); + assert!(result.versioned_storage_cleanup_required); - let keys = store.list_keys("prod/").await.unwrap(); + let keys = store.list_keys_limited("prod/", 10).await.unwrap(); assert!(keys + .iter() + .any(|key| key == "prod/aether-config-backup-20260524-010000.json.zst.aes256gcm")); + assert!(!keys .iter() .any(|key| key == "prod/aether-config-backup-20260524-010000.json.zst")); assert!(keys .iter() - .any(|key| key == "prod/aether-data-backup-20260523-191500.json.zst")); + .any(|key| key == "prod/aether-data-backup-20260523-191500.json.zst.aes256gcm")); assert!(!keys .iter() .any(|key| key == "prod/aether-data-backup-20260524-010000.json.zst")); + + let uploaded = store + .object_bytes(&result.object_key) + .await + .expect("uploaded backup should be readable in the fake store"); + assert!(uploaded.starts_with(BACKUP_ENVELOPE_MAGIC)); + assert!(!uploaded + .windows("config_data".len()) + .any(|window| window == b"config_data")); + let compressed = decrypt_backup_bytes( + TEST_NEW_BACKUP_ENCRYPTION_SECRET, + &result.object_key, + &uploaded, + ) + .expect("encrypted backup should decrypt"); + let decoded = zstd::stream::decode_all(compressed.as_slice()) + .expect("decrypted backup should decompress"); + let decoded: serde_json::Value = + serde_json::from_slice(&decoded).expect("decrypted backup should contain JSON"); + assert_eq!(decoded["version"], "1.0"); + + let migrated_copy = store + .object_bytes("prod/aether-config-backup-20260524-010000.json.zst.aes256gcm") + .await + .expect("legacy backup should be migrated"); + assert_eq!( + decrypt_backup_bytes( + TEST_NEW_BACKUP_ENCRYPTION_SECRET, + "prod/aether-config-backup-20260524-010000.json.zst.aes256gcm", + &migrated_copy, + ) + .expect("migrated backup should decrypt") + .as_slice(), + b"keep-config".as_slice() + ); + } + + #[tokio::test] + async fn legacy_migration_never_overwrites_an_existing_encrypted_object() { + let store = FakeBackupObjectStore::default(); + let legacy_key = "prod/aether-data-backup-20260524-010000.json.zst"; + let preferred_encrypted_key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + store + .put_object(legacy_key, Bytes::from_static(b"legacy payload")) + .await + .unwrap(); + store + .put_object( + preferred_encrypted_key, + Bytes::from_static(b"existing encrypted object"), + ) + .await + .unwrap(); + + let config = sample_backup_config(BackupScope::Data, 10); + let migration = + super::migrate_legacy_plaintext_backups(&config, &store, DEVELOPMENT_ENCRYPTION_KEY) + .await + .expect("legacy backup should migrate without overwriting"); + + assert_eq!(migration.encrypted_copies_created, 1); + assert_eq!(migration.encrypted_copies_verified, 0); + assert_eq!(migration.plaintext_objects_deleted, 1); + assert_eq!(migration.plaintext_objects_retained, 0); + assert!(store.object_bytes(legacy_key).await.is_none()); + assert_eq!( + store.object_bytes(preferred_encrypted_key).await.as_deref(), + Some(b"existing encrypted object".as_slice()) + ); + + let keys = store.list_keys_limited("prod/", 10).await.unwrap(); + let collision_key = keys + .iter() + .find(|key| { + key.starts_with("prod/aether-data-backup-20260524-010000-") + && key.ends_with(".json.zst.aes256gcm") + }) + .expect("migration should create a collision-safe encrypted object"); + let migrated_bytes = store.object_bytes(collision_key).await.unwrap(); + assert_eq!( + decrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, collision_key, &migrated_bytes) + .expect("collision-safe backup should decrypt") + .as_slice(), + b"legacy payload".as_slice() + ); + } + + #[tokio::test] + async fn legacy_migration_retry_reuses_matching_encrypted_copy() { + let store = FakeBackupObjectStore::default(); + let legacy_key = "prod/aether-users-backup-20260524-010000.json.zst"; + let encrypted_key = "prod/aether-users-backup-20260524-010000.json.zst.aes256gcm"; + let (encrypted, _) = super::encrypt_backup_bytes( + DEVELOPMENT_ENCRYPTION_KEY, + encrypted_key, + b"legacy payload", + ) + .expect("test payload should encrypt"); + store + .put_object(legacy_key, Bytes::from_static(b"legacy payload")) + .await + .unwrap(); + store + .put_object(encrypted_key, Bytes::from(encrypted.clone())) + .await + .unwrap(); + + let config = sample_backup_config(BackupScope::Data, 10); + let migration = + super::migrate_legacy_plaintext_backups(&config, &store, DEVELOPMENT_ENCRYPTION_KEY) + .await + .expect("retry should recognize the existing encrypted copy"); + + assert_eq!(migration.encrypted_copies_created, 0); + assert_eq!(migration.encrypted_copies_verified, 1); + assert_eq!(migration.plaintext_objects_deleted, 1); + assert_eq!(migration.plaintext_objects_retained, 0); + assert!(store.object_bytes(legacy_key).await.is_none()); + assert_eq!( + store.object_bytes(encrypted_key).await.as_deref(), + Some(encrypted.as_slice()) + ); + assert_eq!(store.list_keys_limited("prod/", 10).await.unwrap().len(), 1); } #[tokio::test] @@ -216,10 +1063,220 @@ mod tests { .unwrap() .with_timezone(&Utc); - let error = run_backup_with_store(&config, &store, payload, now_utc) - .await - .expect_err("unknown compression should fail"); + let error = run_backup_with_store( + &config, + &store, + payload, + now_utc, + TEST_NEW_BACKUP_ENCRYPTION_SECRET, + ) + .await + .expect_err("unknown compression should fail"); assert!(error.to_string().contains("brotli")); } + + #[test] + fn new_backup_key_policy_rejects_weak_creation_keys_but_keeps_restore_compatibility() { + for insecure in [ + "short-backup-key", + "change-this-to-another-secure-random-string", + "change-this-to-a-secure-random-string", + DEVELOPMENT_ENCRYPTION_KEY, + ] { + let error = validate_new_backup_encryption_secret(insecure) + .expect_err("new backup creation must reject weak or published keys"); + assert!(error.to_string().contains("unsafe")); + + BackupDecryptionKey::historical(insecure) + .expect("historical restore must continue accepting legacy weak keys"); + } + + validate_new_backup_encryption_secret(TEST_NEW_BACKUP_ENCRYPTION_SECRET) + .expect("strong new backup key should be accepted"); + } + + #[test] + fn backup_envelope_rejects_wrong_keys_and_tampering() { + let object_key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + let (encrypted, _) = + super::encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, b"secret payload") + .expect("backup should encrypt"); + assert!(decrypt_backup_bytes("wrong-key", object_key, &encrypted).is_err()); + assert!(decrypt_backup_bytes( + DEVELOPMENT_ENCRYPTION_KEY, + "prod/aether-users-backup-20260524-010000.json.zst.aes256gcm", + &encrypted, + ) + .is_err()); + + let mut tampered = encrypted; + let last = tampered.last_mut().expect("envelope should not be empty"); + *last ^= 1; + assert!(decrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &tampered).is_err()); + } + + #[test] + fn v2_envelope_has_stable_non_secret_key_id_and_restores_json() { + let object_key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + let compressed = compressed_json(json!({ + "version": "1.6", + "exported_at": "2026-05-24T01:00:00Z", + "value": 42 + })); + let (first, first_id) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed).unwrap(); + let (second, second_id) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed).unwrap(); + + assert_eq!( + first[BACKUP_ENVELOPE_MAGIC.len()], + BACKUP_ENVELOPE_VERSION_V2 + ); + assert_eq!(first_id, second_id); + assert_ne!( + first, second, + "fresh nonces must produce different envelopes" + ); + let (_, other_id) = + encrypt_backup_bytes("different-secret", object_key, &compressed).unwrap(); + assert_ne!(first_id, other_id); + + let restored = restore_backup_json( + object_key, + &first, + &[ + BackupDecryptionKey::current("wrong-secret").unwrap(), + BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY).unwrap(), + ], + BackupRestoreLimits::default(), + ) + .unwrap(); + assert_eq!(restored.envelope_version, BACKUP_ENVELOPE_VERSION_V2); + assert_eq!(restored.scope(), super::BackupRestoreScope::Data); + let expected_key_id = super::encode_key_id(&first_id); + assert_eq!(restored.key_id.as_deref(), Some(expected_key_id.as_str())); + let json: serde_json::Value = serde_json::from_slice(restored.json_bytes()).unwrap(); + assert_eq!(json["value"], 42); + } + + #[test] + fn v2_unknown_key_id_fails_before_authenticated_decryption() { + let object_key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + let compressed = compressed_json(json!({ + "version": "1.6", + "exported_at": "2026-05-24T01:00:00Z" + })); + let (encrypted, _) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed).unwrap(); + let mut forged_key_id = encrypted.clone(); + forged_key_id[BACKUP_ENVELOPE_MAGIC.len() + 1] ^= 1; + + let error = restore_backup_json( + object_key, + &forged_key_id, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY).unwrap()], + BackupRestoreLimits::default(), + ) + .unwrap_err(); + assert!(matches!(error, BackupRestoreError::UnknownKeyId(_))); + } + + #[test] + fn v2_aad_rejects_object_key_swap() { + let object_key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + let compressed = compressed_json(json!({ + "version": "1.6", + "exported_at": "2026-05-24T01:00:00Z" + })); + let (encrypted, _) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed).unwrap(); + let error = restore_backup_json( + "prod/aether-users-backup-20260524-010000.json.zst.aes256gcm", + &encrypted, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY).unwrap()], + BackupRestoreLimits::default(), + ) + .unwrap_err(); + assert!(matches!(error, BackupRestoreError::AuthenticationFailed)); + } + + #[test] + fn restore_rejects_traversal_and_non_backup_object_keys_before_decryption() { + for object_key in [ + "../aether-data-backup-20260524-010000.json.zst.aes256gcm", + "/aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod//aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod/aether-data-backup-invalid.json.zst.aes256gcm", + "prod/not-an-aether-backup-20260524-010000.json.zst.aes256gcm", + ] { + let error = restore_backup_json(object_key, b"", &[], BackupRestoreLimits::default()) + .expect_err("unsafe object key must fail before envelope processing"); + assert!( + matches!(error, BackupRestoreError::InvalidObjectKey), + "unexpected error for {object_key}: {error}" + ); + } + } + + #[test] + fn restore_keeps_v1_compatibility_and_tries_only_legacy_candidates() { + let object_key = "prod/aether-config-backup-20260524-010000.json.zst.aes256gcm"; + let compressed = compressed_json(json!({ + "version": "2.3", + "exported_at": "2026-05-24T01:00:00Z", + "config_data": {} + })); + let encrypted = encrypt_v1_for_test(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed); + let restored = restore_backup_json( + object_key, + &encrypted, + &[ + BackupDecryptionKey::v2_only(DEVELOPMENT_ENCRYPTION_KEY).unwrap(), + BackupDecryptionKey::historical("wrong-legacy-secret").unwrap(), + BackupDecryptionKey::historical(DEVELOPMENT_ENCRYPTION_KEY).unwrap(), + ], + BackupRestoreLimits::default(), + ) + .unwrap(); + assert_eq!(restored.envelope_version, BACKUP_ENVELOPE_VERSION_V1); + assert_eq!(restored.key_id, None); + assert_eq!(restored.export_version.as_deref(), Some("2.3")); + + let too_many: Vec<_> = (0..17) + .map(|index| BackupDecryptionKey::historical(format!("legacy-{index}")).unwrap()) + .collect(); + assert!(matches!( + restore_backup_json( + object_key, + &encrypted, + &too_many, + BackupRestoreLimits::default() + ), + Err(BackupRestoreError::TooManyLegacyKeys) + )); + } + + #[test] + fn restore_rejects_zstd_output_over_limit() { + let object_key = "prod/aether-users-backup-20260524-010000.json.zst.aes256gcm"; + let compressed = compressed_json(json!({ + "version": "1.6", + "exported_at": "2026-05-24T01:00:00Z", + "padding": "x".repeat(4096) + })); + let (encrypted, _) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed).unwrap(); + let error = restore_backup_json( + object_key, + &encrypted, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY).unwrap()], + BackupRestoreLimits { + max_encrypted_bytes: encrypted.len(), + max_json_bytes: 128, + }, + ) + .unwrap_err(); + assert!(matches!(error, BackupRestoreError::JsonSizeLimit { .. })); + } } diff --git a/apps/aether-gateway/src/backup/mod.rs b/apps/aether-gateway/src/backup/mod.rs index 4aafc24db..63bc07938 100644 --- a/apps/aether-gateway/src/backup/mod.rs +++ b/apps/aether-gateway/src/backup/mod.rs @@ -6,5 +6,167 @@ pub(crate) mod store; pub(crate) mod task; pub(crate) mod worker; +pub use executor::{ + restore_backup_json, BackupDecryptionKey, BackupRestoreError, BackupRestoreLimits, + RestoredBackupJson, DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES, DEFAULT_BACKUP_MAX_JSON_BYTES, +}; + +use axum::body::Bytes; +use serde_json::Value; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BackupRestoreScope { + Config, + Users, + Data, +} + +impl BackupRestoreScope { + pub const fn as_str(self) -> &'static str { + match self { + Self::Config => "config", + Self::Users => "users", + Self::Data => "data", + } + } +} + +#[derive(Debug, thiserror::Error)] +#[error("backup database apply failed: {0}")] +pub struct BackupApplyError(String); + +pub async fn apply_restored_backup( + app: &crate::AppState, + restored: RestoredBackupJson, + scope: BackupRestoreScope, + operator_id: Option<&str>, +) -> Result, BackupApplyError> { + let (json_bytes, authority) = restored.into_authenticated_parts(); + if authority.scope() != scope { + return Err(BackupApplyError(format!( + "authenticated {} backup cannot be applied to {} scope", + authority.scope().as_str(), + scope.as_str(), + ))); + } + let request_body = Bytes::from(json_bytes); + let state = crate::admin_api::AdminAppState::new(app); + let result = crate::admin_api::execute_admin_system_import_exclusively(app, async { + match scope { + BackupRestoreScope::Config => { + state + .restore_admin_system_config_backup(&request_body, authority) + .await + } + BackupRestoreScope::Users => { + state + .restore_admin_system_users_backup(&request_body, operator_id, authority) + .await + } + BackupRestoreScope::Data => { + state + .restore_admin_system_data_backup(&request_body, operator_id, authority) + .await + } + } + }) + .await + .map_err(|error| { + let message = match error { + crate::admin_api::AdminSystemImportLockError::Conflict => { + "another system import or restore is already running" + } + crate::admin_api::AdminSystemImportLockError::Unavailable => { + "system import coordination is unavailable" + } + crate::admin_api::AdminSystemImportLockError::Lost => { + "system import coordination lease was lost; restore was cancelled and may have partially applied changes" + } + }; + BackupApplyError(message.to_string()) + })?; + result.map_err(|error| BackupApplyError(error.into_message())) +} + pub(crate) const S3_BACKUP_ENABLED_KEY: &str = "backup_s3_enabled"; pub(crate) const S3_BACKUP_LAST_SLOT_KEY: &str = "backup_s3_last_slot"; + +#[cfg(test)] +mod tests { + use super::{ + apply_restored_backup, BackupDecryptionKey, BackupRestoreLimits, BackupRestoreScope, + RestoredBackupJson, + }; + use crate::backup::executor::encrypt_backup_bytes; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use serde_json::json; + + fn authenticated_users_backup() -> RestoredBackupJson { + let object_key = "prod/aether-users-backup-20260830-120000.json.zst.aes256gcm"; + let compressed = zstd::stream::encode_all( + serde_json::to_vec(&json!({ + "version": "1.5", + "exported_at": "2026-08-30T12:00:00Z", + "users": [], + "standalone_keys": [], + })) + .expect("test backup should serialize") + .as_slice(), + 0, + ) + .expect("test backup should compress"); + let (envelope, _) = + encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed) + .expect("test backup should encrypt"); + super::restore_backup_json( + object_key, + &envelope, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY) + .expect("test restore key should build")], + BackupRestoreLimits::default(), + ) + .expect("test backup should authenticate") + } + + #[tokio::test] + async fn authenticated_backup_cannot_be_applied_to_a_different_scope() { + let restored = authenticated_users_backup(); + + let error = apply_restored_backup( + &crate::AppState::new().expect("test state should build"), + restored, + BackupRestoreScope::Config, + None, + ) + .await + .expect_err("scope mismatch must fail before database access"); + + assert_eq!( + error.to_string(), + "backup database apply failed: authenticated users backup cannot be applied to config scope" + ); + } + + #[tokio::test] + async fn authenticated_backup_apply_uses_the_shared_system_import_lock() { + let app = crate::AppState::new().expect("test state should build"); + let lock = crate::admin_api::try_acquire_admin_system_import_lease(&app) + .await + .expect("test should acquire the shared import lease"); + + let error = apply_restored_backup( + &app, + authenticated_users_backup(), + BackupRestoreScope::Users, + None, + ) + .await + .expect_err("restore must not interleave with another system import"); + + assert_eq!( + error.to_string(), + "backup database apply failed: another system import or restore is already running" + ); + crate::admin_api::release_admin_system_import_lease(&app, &lock).await; + } +} diff --git a/apps/aether-gateway/src/backup/scopes.rs b/apps/aether-gateway/src/backup/scopes.rs index f8993d55f..2f6b5ac9e 100644 --- a/apps/aether-gateway/src/backup/scopes.rs +++ b/apps/aether-gateway/src/backup/scopes.rs @@ -1,5 +1,8 @@ use std::fmt; +const ENCRYPTED_BACKUP_FILE_SUFFIX: &str = ".json.zst.aes256gcm"; +const LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX: &str = ".json.zst"; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum BackupScope { Config, @@ -52,10 +55,81 @@ impl BackupScope { } } + pub(crate) fn from_encrypted_object_key(object_key: &str) -> Option { + if object_key.is_empty() + || object_key.starts_with('/') + || object_key.contains('\0') + || object_key.contains('\\') + { + return None; + } + let mut segments = object_key.split('/').peekable(); + let mut file_name = None; + while let Some(segment) = segments.next() { + if segment.is_empty() + || segment == "." + || segment == ".." + || segment.chars().any(char::is_control) + { + return None; + } + if segments.peek().is_none() { + file_name = Some(segment); + } + } + let file_name = file_name?; + + [Self::Config, Self::Users, Self::Data] + .into_iter() + .find(|scope| { + file_name + .strip_prefix(&format!("{}-", scope.file_stem())) + .and_then(|rest| rest.strip_suffix(ENCRYPTED_BACKUP_FILE_SUFFIX)) + .is_some_and(is_aether_backup_object_id) + }) + } + + #[cfg(test)] pub(crate) fn matching_backup_keys( self, prefix: &str, keys: impl IntoIterator, + ) -> Vec { + self.matching_backup_keys_with_suffixes( + prefix, + keys, + &[ + ENCRYPTED_BACKUP_FILE_SUFFIX, + LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX, + ], + ) + } + + pub(crate) fn matching_encrypted_backup_keys( + self, + prefix: &str, + keys: impl IntoIterator, + ) -> Vec { + self.matching_backup_keys_with_suffixes(prefix, keys, &[ENCRYPTED_BACKUP_FILE_SUFFIX]) + } + + pub(crate) fn matching_legacy_plaintext_backup_keys( + self, + prefix: &str, + keys: impl IntoIterator, + ) -> Vec { + self.matching_backup_keys_with_suffixes( + prefix, + keys, + &[LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX], + ) + } + + fn matching_backup_keys_with_suffixes( + self, + prefix: &str, + keys: impl IntoIterator, + file_suffixes: &[&str], ) -> Vec { let normalized_prefix = normalized_prefix(prefix); let expected_prefix = if normalized_prefix.is_empty() { @@ -64,7 +138,6 @@ impl BackupScope { format!("{normalized_prefix}/") }; let file_prefix = format!("{}-", self.file_stem()); - let file_suffix = ".json.zst"; keys.into_iter() .filter(|key| { @@ -74,20 +147,24 @@ impl BackupScope { if file_name.contains('/') { return false; } - let Some(timestamp) = file_name - .strip_prefix(&file_prefix) - .and_then(|rest| rest.strip_suffix(file_suffix)) - else { + let Some(timestamp) = file_name.strip_prefix(&file_prefix).and_then(|rest| { + file_suffixes + .iter() + .find_map(|suffix| rest.strip_suffix(suffix)) + }) else { return false; }; - is_aether_backup_timestamp(timestamp) + is_aether_backup_object_id(timestamp) }) .collect() } fn file_name(self, timestamp: &str) -> String { - format!("{}-{timestamp}.json.zst", self.file_stem()) + format!( + "{}-{timestamp}{ENCRYPTED_BACKUP_FILE_SUFFIX}", + self.file_stem() + ) } } @@ -110,6 +187,25 @@ fn is_aether_backup_timestamp(timestamp: &str) -> bool { && bytes[9..].iter().all(|byte| byte.is_ascii_digit()) } +fn is_aether_backup_object_id(value: &str) -> bool { + if is_aether_backup_timestamp(value) { + return true; + } + + let Some((timestamp, collision_digest)) = value.split_once('-').and_then(|(date, rest)| { + let (time, digest) = rest.split_once('-')?; + Some((format!("{date}-{time}"), digest)) + }) else { + return false; + }; + + is_aether_backup_timestamp(×tamp) + && collision_digest.len() == 64 + && collision_digest + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) +} + #[cfg(test)] mod tests { use super::BackupScope; @@ -130,15 +226,15 @@ mod tests { assert_eq!( BackupScope::Config.object_key("prod/", "20260524-031500"), - "prod/aether-config-backup-20260524-031500.json.zst" + "prod/aether-config-backup-20260524-031500.json.zst.aes256gcm" ); assert_eq!( BackupScope::Users.object_key("prod/", "20260524-031500"), - "prod/aether-users-backup-20260524-031500.json.zst" + "prod/aether-users-backup-20260524-031500.json.zst.aes256gcm" ); assert_eq!( BackupScope::Data.object_key("prod/", "20260524-031500"), - "prod/aether-data-backup-20260524-031500.json.zst" + "prod/aether-data-backup-20260524-031500.json.zst.aes256gcm" ); } @@ -146,7 +242,7 @@ mod tests { fn retention_filter_only_matches_same_scope() { let keys = vec![ "prod/aether-config-backup-20260524-010000.json.zst".to_string(), - "prod/aether-users-backup-20260524-010000.json.zst".to_string(), + "prod/aether-users-backup-20260524-010000.json.zst.aes256gcm".to_string(), "prod/aether-data-backup-20260524-010000.json.zst".to_string(), "prod/random.json.zst".to_string(), ]; @@ -155,14 +251,18 @@ mod tests { assert_eq!( matched, - vec!["prod/aether-users-backup-20260524-010000.json.zst"] + vec!["prod/aether-users-backup-20260524-010000.json.zst.aes256gcm"] ); } #[test] fn retention_filter_requires_aether_timestamp_format() { + let collision_digest = "a".repeat(64); let keys = vec![ "prod/aether-users-backup-20260524-010000.json.zst".to_string(), + format!( + "prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm" + ), "prod/aether-users-backup-foo.json.zst".to_string(), "prod/aether-users-backup-2026052-010000.json.zst".to_string(), "prod/aether-users-backup-202605240-010000.json.zst".to_string(), @@ -171,13 +271,19 @@ mod tests { "prod/aether-users-backup-20260524010000.json.zst".to_string(), "prod/aether-users-backup-2026052a-010000.json.zst".to_string(), "prod/aether-users-backup-20260524-01000x.json.zst".to_string(), + "prod/aether-users-backup-20260524-010000-short.json.zst.aes256gcm".to_string(), ]; let matched = BackupScope::Users.matching_backup_keys("prod/", keys); assert_eq!( matched, - vec!["prod/aether-users-backup-20260524-010000.json.zst"] + vec![ + "prod/aether-users-backup-20260524-010000.json.zst".to_string(), + format!( + "prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm" + ), + ] ); } @@ -185,11 +291,11 @@ mod tests { fn backup_key_prefix_boundaries_are_exact() { assert_eq!( BackupScope::Config.object_key("", "20260524-031500"), - "aether-config-backup-20260524-031500.json.zst" + "aether-config-backup-20260524-031500.json.zst.aes256gcm" ); assert_eq!( BackupScope::Config.object_key("prod", "20260524-031500"), - "prod/aether-config-backup-20260524-031500.json.zst" + "prod/aether-config-backup-20260524-031500.json.zst.aes256gcm" ); let keys = vec![ @@ -208,4 +314,36 @@ mod tests { vec!["prod/aether-config-backup-20260524-010000.json.zst"] ); } + + #[test] + fn encrypted_object_key_parser_binds_scope_and_rejects_path_traversal() { + let collision_digest = "a".repeat(64); + assert_eq!( + BackupScope::from_encrypted_object_key( + "prod/aether-config-backup-20260524-010000.json.zst.aes256gcm" + ), + Some(BackupScope::Config) + ); + assert_eq!( + BackupScope::from_encrypted_object_key(&format!( + "prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm" + )), + Some(BackupScope::Users) + ); + for key in [ + "../aether-data-backup-20260524-010000.json.zst.aes256gcm", + "/aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod//aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod/./aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod\\aether-data-backup-20260524-010000.json.zst.aes256gcm", + "prod/aether-data-backup-invalid.json.zst.aes256gcm", + "prod/unrelated-20260524-010000.json.zst.aes256gcm", + ] { + assert_eq!( + BackupScope::from_encrypted_object_key(key), + None, + "unsafe or unrelated key: {key}" + ); + } + } } diff --git a/apps/aether-gateway/src/backup/store.rs b/apps/aether-gateway/src/backup/store.rs index 74cdaf1e2..0c76dabe3 100644 --- a/apps/aether-gateway/src/backup/store.rs +++ b/apps/aether-gateway/src/backup/store.rs @@ -2,23 +2,45 @@ use std::collections::BTreeMap; use std::fmt; use std::sync::Arc; -use bytes::Bytes; +use bytes::{Bytes, BytesMut}; use futures_util::TryStreamExt; use object_store::aws::AmazonS3Builder; use object_store::path::Path; -use object_store::{ClientOptions, ObjectStore}; +use object_store::{ClientOptions, ObjectStore, ObjectStoreExt, PutMode, PutOptions}; use reqwest::header::HeaderValue; use tokio::sync::RwLock; use super::config::S3BackupConfig; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum BackupObjectCreateResult { + Created, + AlreadyExists, +} + #[async_trait::async_trait] pub(crate) trait BackupObjectStore: Send + Sync { async fn put_object(&self, key: &str, bytes: Bytes) -> Result<(), BackupStoreError>; - async fn list_keys(&self, prefix: &str) -> Result, BackupStoreError>; + async fn put_object_if_absent( + &self, + key: &str, + bytes: Bytes, + ) -> Result; + + async fn get_object_limited( + &self, + key: &str, + max_bytes: usize, + ) -> Result; async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError>; + + async fn list_keys_limited( + &self, + prefix: &str, + max_objects: usize, + ) -> Result, BackupStoreError>; } #[derive(Debug, Clone, PartialEq, Eq)] @@ -60,21 +82,72 @@ impl BackupObjectStore for FakeBackupObjectStore { Ok(()) } - async fn list_keys(&self, prefix: &str) -> Result, BackupStoreError> { + async fn put_object_if_absent( + &self, + key: &str, + bytes: Bytes, + ) -> Result { + let mut objects = self.objects.write().await; + if objects.contains_key(key) { + Ok(BackupObjectCreateResult::AlreadyExists) + } else { + objects.insert(key.to_string(), bytes); + Ok(BackupObjectCreateResult::Created) + } + } + + async fn get_object_limited( + &self, + key: &str, + max_bytes: usize, + ) -> Result { + let bytes = self + .objects + .read() + .await + .get(key) + .cloned() + .ok_or_else(|| BackupStoreError::new(format!("backup object `{key}` not found")))?; + if bytes.len() > max_bytes { + return Err(BackupStoreError::new(format!( + "backup object `{key}` exceeds the configured {max_bytes} byte read limit" + ))); + } + Ok(bytes) + } + + async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> { + self.objects.write().await.remove(key); + Ok(()) + } + + async fn list_keys_limited( + &self, + prefix: &str, + max_objects: usize, + ) -> Result, BackupStoreError> { let prefix = directory_list_prefix(prefix); - Ok(self + let keys: Vec<_> = self .objects .read() .await .keys() .filter(|key| key.starts_with(&prefix)) .cloned() - .collect()) + .collect(); + if keys.len() > max_objects { + return Err(BackupStoreError::new(format!( + "backup object listing exceeds the configured {max_objects} object limit" + ))); + } + Ok(keys) } +} - async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> { - self.objects.write().await.remove(key); - Ok(()) +#[cfg(test)] +impl FakeBackupObjectStore { + pub(crate) async fn object_bytes(&self, key: &str) -> Option { + self.objects.read().await.get(key).cloned() } } @@ -125,17 +198,68 @@ impl BackupObjectStore for ObjectStoreS3BackupStore { .map_err(|error| BackupStoreError::object_store("put", key, error)) } - async fn list_keys(&self, prefix: &str) -> Result, BackupStoreError> { - let prefix_path = list_prefix_path(prefix); - let mut keys = self + async fn put_object_if_absent( + &self, + key: &str, + bytes: Bytes, + ) -> Result { + let options = PutOptions { + mode: PutMode::Create, + ..PutOptions::default() + }; + match self .store - .list(prefix_path.as_ref()) - .map_ok(|meta| meta.location.to_string()) - .try_collect::>() + .put_opts(&Path::from(key), bytes.into(), options) .await - .map_err(|error| BackupStoreError::object_store("list", prefix, error))?; - keys.sort(); - Ok(keys) + { + Ok(_) => Ok(BackupObjectCreateResult::Created), + Err(object_store::Error::AlreadyExists { .. }) => { + Ok(BackupObjectCreateResult::AlreadyExists) + } + Err(error) => Err(BackupStoreError::object_store( + "conditional put", + key, + error, + )), + } + } + + async fn get_object_limited( + &self, + key: &str, + max_bytes: usize, + ) -> Result { + let result = self + .store + .get(&Path::from(key)) + .await + .map_err(|error| BackupStoreError::object_store("get", key, error))?; + if result.meta.size > u64::try_from(max_bytes).unwrap_or(u64::MAX) { + return Err(BackupStoreError::new(format!( + "backup object `{key}` exceeds the configured {max_bytes} byte read limit" + ))); + } + let object_size = result.meta.size; + let mut stream = result.into_stream(); + let mut bytes = BytesMut::with_capacity( + usize::try_from(object_size) + .unwrap_or(max_bytes) + .min(max_bytes) + .min(8 * 1024 * 1024), + ); + while let Some(chunk) = stream + .try_next() + .await + .map_err(|error| BackupStoreError::object_store("read", key, error))? + { + if bytes.len().saturating_add(chunk.len()) > max_bytes { + return Err(BackupStoreError::new(format!( + "backup object `{key}` exceeds the configured {max_bytes} byte read limit" + ))); + } + bytes.extend_from_slice(&chunk); + } + Ok(bytes.freeze()) } async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> { @@ -144,6 +268,30 @@ impl BackupObjectStore for ObjectStoreS3BackupStore { .await .map_err(|error| BackupStoreError::object_store("delete", key, error)) } + + async fn list_keys_limited( + &self, + prefix: &str, + max_objects: usize, + ) -> Result, BackupStoreError> { + let prefix_path = list_prefix_path(prefix); + let mut objects = self.store.list(prefix_path.as_ref()); + let mut keys = Vec::new(); + while let Some(meta) = objects + .try_next() + .await + .map_err(|error| BackupStoreError::object_store("list", prefix, error))? + { + if keys.len() >= max_objects { + return Err(BackupStoreError::new(format!( + "backup object listing exceeds the configured {max_objects} object limit" + ))); + } + keys.push(meta.location.to_string()); + } + keys.sort(); + Ok(keys) + } } fn directory_list_prefix(prefix: &str) -> String { @@ -166,10 +314,12 @@ fn list_prefix_path(prefix: &str) -> Option { #[cfg(test)] mod tests { - use super::{list_prefix_path, BackupObjectStore, FakeBackupObjectStore}; + use super::{ + list_prefix_path, BackupObjectCreateResult, BackupObjectStore, FakeBackupObjectStore, + }; #[tokio::test] - async fn fake_backup_object_store_puts_lists_and_deletes() { + async fn fake_backup_object_store_puts_and_lists() { let store = FakeBackupObjectStore::default(); store .put_object( @@ -186,17 +336,59 @@ mod tests { .await .unwrap(); - let keys = store.list_keys("prod/").await.unwrap(); - assert_eq!(keys.len(), 2); - - store - .delete_object("prod/aether-data-backup-20260524-010000.json.zst") - .await - .unwrap(); - let keys = store.list_keys("prod/").await.unwrap(); + let keys = store.list_keys_limited("prod/", 2).await.unwrap(); assert_eq!( keys, - vec!["prod/aether-data-backup-20260524-020000.json.zst"] + vec![ + "prod/aether-data-backup-20260524-010000.json.zst", + "prod/aether-data-backup-20260524-020000.json.zst", + ] + ); + } + + #[tokio::test] + async fn fake_backup_object_store_enforces_read_and_listing_limits() { + let store = FakeBackupObjectStore::default(); + store + .put_object("prod/one", bytes::Bytes::from_static(b"1234")) + .await + .unwrap(); + store + .put_object("prod/two", bytes::Bytes::from_static(b"5678")) + .await + .unwrap(); + + assert!(store.get_object_limited("prod/one", 3).await.is_err()); + assert_eq!( + store.get_object_limited("prod/one", 4).await.unwrap(), + bytes::Bytes::from_static(b"1234") + ); + assert!(store.list_keys_limited("prod/", 1).await.is_err()); + assert_eq!(store.list_keys_limited("prod/", 2).await.unwrap().len(), 2); + } + + #[tokio::test] + async fn fake_backup_object_store_conditional_put_never_overwrites() { + let store = FakeBackupObjectStore::default(); + let key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm"; + + assert_eq!( + store + .put_object_if_absent(key, bytes::Bytes::from_static(b"first")) + .await + .unwrap(), + BackupObjectCreateResult::Created + ); + assert_eq!( + store + .put_object_if_absent(key, bytes::Bytes::from_static(b"second")) + .await + .unwrap(), + BackupObjectCreateResult::AlreadyExists + ); + assert_eq!( + store.object_bytes(key).await.as_deref(), + Some(b"first".as_slice()) ); } @@ -218,7 +410,7 @@ mod tests { .await .unwrap(); - let keys = store.list_keys("prod").await.unwrap(); + let keys = store.list_keys_limited("prod", 10).await.unwrap(); assert_eq!( keys, diff --git a/apps/aether-gateway/src/backup/task.rs b/apps/aether-gateway/src/backup/task.rs index 4a754a052..754cbd622 100644 --- a/apps/aether-gateway/src/backup/task.rs +++ b/apps/aether-gateway/src/backup/task.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::future::Future; use std::time::Duration; use aether_admin::system::admin_system_config_default_value; @@ -12,14 +13,15 @@ use chrono::Utc; use futures_util::FutureExt; use serde::Serialize; use serde_json::{json, Map, Value}; +use tokio::task::{JoinError, JoinHandle}; use tracing::warn; use super::config::S3BackupConfig; use super::executor::{run_backup_with_store, BackupRunResult}; use super::scopes::BackupScope; use super::store::ObjectStoreS3BackupStore; -use crate::admin_api::AdminAppState; -use crate::handlers::shared::decrypt_catalog_secret_with_fallbacks; +use crate::admin_api::{AdminAppState, SystemExportMode}; +use crate::handlers::shared::decrypt_or_migrate_system_config_secret; use crate::task_runtime::{ append_event_with_logging, build_task_run_id, now_unix_secs, spawn_fire_and_forget, task_definition, update_run_status, upsert_run_with_logging, TASK_KEY_SYSTEM_S3_BACKUP, @@ -48,6 +50,9 @@ const S3_BACKUP_CONFIG_KEYS: &[&str] = &[ ]; const S3_BACKUP_QUEUED_MESSAGE: &str = "S3 备份任务已提交"; +const S3_BACKUP_INTERNAL_ERROR_DETAIL: &str = "S3 备份服务暂时不可用"; +const S3_BACKUP_TASK_FAILURE_CODE: &str = "s3_backup_failed"; +const S3_BACKUP_SLOT_RECORD_FAILURE_CODE: &str = "s3_backup_slot_record_failed"; const S3_BACKUP_TASK_LOCK_KEY: &str = "task_runtime:lock:system.s3.backup"; const S3_BACKUP_TASK_LOCK_TTL: Duration = Duration::from_secs(60 * 60 * 6); const S3_BACKUP_TASK_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(60 * 5); @@ -67,6 +72,16 @@ pub(crate) struct S3BackupTaskError { detail: String, } +enum BackupLockRenewalFailure { + Lost, + Backend(E), +} + +enum BackupLockRaceOutcome { + BackupCompleted(T), + LeaseLost(Result<(), JoinError>), +} + impl S3BackupTaskError { fn bad_request(detail: impl Into) -> Self { Self { @@ -114,8 +129,12 @@ impl fmt::Display for S3BackupTaskError { impl std::error::Error for S3BackupTaskError {} impl From for S3BackupTaskError { - fn from(error: GatewayError) -> Self { - Self::internal(format!("{error:?}")) + fn from(_error: GatewayError) -> Self { + warn!( + error_category = "dependency_failed", + "S3 backup dependency failed" + ); + Self::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL) } } @@ -220,8 +239,6 @@ fn s3_backup_task_payload_json( ) -> Value { let mut payload = json!({ "scope": config.scope.as_config_value(), - "bucket": config.bucket.clone(), - "prefix": config.prefix.clone(), "compression": config.compression.clone(), "trigger": trigger, }); @@ -259,7 +276,7 @@ fn spawn_s3_backup_worker( Some(100), Some("S3 备份任务异常退出".to_string()), None, - Some("S3 backup task panicked".to_string()), + Some("background_task_panicked".to_string()), None, Some(now_unix_secs()), ) @@ -294,16 +311,67 @@ async fn run_s3_backup_worker_inner( .await; append_event_with_logging(&app, &run_id, "running", "S3 backup task started", None).await; - let heartbeat = spawn_s3_backup_task_heartbeat(app.clone(), run_id.clone(), lock); - let result = run_s3_backup_once(&app, &config).await; - heartbeat.abort(); - let _ = heartbeat.await; + let heartbeat = spawn_s3_backup_task_heartbeat(app.clone(), run_id.clone(), lock.clone()); + let result = match race_backup_with_lock_heartbeat(run_s3_backup_once(&app, &config), heartbeat) + .await + { + BackupLockRaceOutcome::BackupCompleted(result) => { + match require_successful_backup_lock_renewal( + app.runtime_state + .lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL) + .await, + ) { + Ok(()) => result, + Err(BackupLockRenewalFailure::Lost) => { + warn!( + run_id = %run_id, + lock_key = %lock.key, + "S3 backup task lost its distributed lock before publishing completion" + ); + Err(S3BackupTaskError::service_unavailable( + "S3 备份任务锁已失效,任务完成状态未发布", + )) + } + Err(BackupLockRenewalFailure::Backend(error)) => { + warn!( + run_id = %run_id, + lock_key = %lock.key, + error = %error, + "S3 backup task could not verify its distributed lock before publishing completion" + ); + Err(S3BackupTaskError::service_unavailable( + "无法确认 S3 备份任务锁所有权,任务完成状态未发布", + )) + } + } + } + BackupLockRaceOutcome::LeaseLost(heartbeat_result) => { + match heartbeat_result { + Ok(()) => warn!( + run_id = %run_id, + "S3 backup task stopped after losing its distributed lock" + ), + Err(error) => warn!( + run_id = %run_id, + error = %error, + "S3 backup lock heartbeat task failed" + ), + } + Err(S3BackupTaskError::service_unavailable( + "S3 备份任务锁已失效,任务已停止", + )) + } + }; match result { Ok(result) => { if let Some(slot) = scheduled_backup_slot_to_record(scheduled_slot.as_deref(), true) { - if let Err(error) = record_scheduled_backup_slot(&app, &slot).await { - warn!(error = ?error, run_id = %run_id, "S3 backup slot record failed"); + if record_scheduled_backup_slot(&app, &slot).await.is_err() { + warn!( + error_category = "slot_record_failed", + run_id = %run_id, + "S3 backup slot record failed" + ); let _ = update_run_status( &app, &run_id, @@ -311,7 +379,7 @@ async fn run_s3_backup_worker_inner( Some(100), Some("S3 备份任务完成,但记录调度时间失败".to_string()), None, - Some(format!("S3 backup slot record failed: {error:?}")), + Some(S3_BACKUP_SLOT_RECORD_FAILURE_CODE.to_string()), None, Some(now_unix_secs()), ) @@ -321,7 +389,7 @@ async fn run_s3_backup_worker_inner( &run_id, "failed", "S3 backup slot record failed", - Some(json!({ "error": format!("{error:?}") })), + Some(json!({ "error_code": S3_BACKUP_SLOT_RECORD_FAILURE_CODE })), ) .await; return; @@ -349,8 +417,12 @@ async fn run_s3_backup_worker_inner( ) .await; } - Err(error) => { - warn!(error = %error, run_id = %run_id, "S3 backup task failed"); + Err(_) => { + warn!( + error_category = "backup_execution_failed", + run_id = %run_id, + "S3 backup task failed" + ); let _ = update_run_status( &app, &run_id, @@ -358,7 +430,7 @@ async fn run_s3_backup_worker_inner( Some(100), Some("S3 备份任务失败".to_string()), None, - Some(error.to_string()), + Some(S3_BACKUP_TASK_FAILURE_CODE.to_string()), None, Some(now_unix_secs()), ) @@ -368,13 +440,44 @@ async fn run_s3_backup_worker_inner( &run_id, "failed", "S3 backup task failed", - Some(json!({ "error": error.to_string() })), + Some(json!({ "error_code": S3_BACKUP_TASK_FAILURE_CODE })), ) .await; } } } +async fn race_backup_with_lock_heartbeat( + backup: F, + mut heartbeat: JoinHandle<()>, +) -> BackupLockRaceOutcome +where + F: Future, +{ + tokio::pin!(backup); + tokio::select! { + biased; + heartbeat_result = &mut heartbeat => { + BackupLockRaceOutcome::LeaseLost(heartbeat_result) + } + result = &mut backup => { + heartbeat.abort(); + let _ = heartbeat.await; + BackupLockRaceOutcome::BackupCompleted(result) + } + } +} + +fn require_successful_backup_lock_renewal( + result: Result, +) -> Result<(), BackupLockRenewalFailure> { + match result { + Ok(true) => Ok(()), + Ok(false) => Err(BackupLockRenewalFailure::Lost), + Err(error) => Err(BackupLockRenewalFailure::Backend(error)), + } +} + fn spawn_s3_backup_task_heartbeat( app: AppState, run_id: String, @@ -386,10 +489,30 @@ fn spawn_s3_backup_task_heartbeat( interval.tick().await; loop { interval.tick().await; - let _ = app - .runtime_state - .lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL) - .await; + match require_successful_backup_lock_renewal( + app.runtime_state + .lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL) + .await, + ) { + Ok(()) => {} + Err(BackupLockRenewalFailure::Lost) => { + warn!( + run_id = %run_id, + lock_key = %lock.key, + "S3 backup task distributed lock is no longer owned" + ); + return; + } + Err(BackupLockRenewalFailure::Backend(error)) => { + warn!( + run_id = %run_id, + lock_key = %lock.key, + error = %error, + "S3 backup task distributed lock renewal failed" + ); + return; + } + } let _ = update_run_status( &app, &run_id, @@ -458,9 +581,15 @@ async fn acquire_s3_backup_task_lock( Ok(None) => Err(S3BackupTaskError::conflict( "已有 S3 备份任务正在执行,请等待当前任务完成后再试", )), - Err(error) => Err(S3BackupTaskError::service_unavailable(format!( - "无法获取 S3 备份任务锁:{error}" - ))), + Err(_) => { + warn!( + error_category = "lock_acquisition_failed", + "S3 backup task lock acquisition failed" + ); + Err(S3BackupTaskError::service_unavailable( + "无法获取 S3 备份任务锁,请稍后重试", + )) + } } } @@ -509,25 +638,77 @@ async fn run_s3_backup_once( app: &AppState, config: &S3BackupConfig, ) -> Result { - let admin_state = AdminAppState::new(app); - let payload = match config.scope { - BackupScope::Config => { - admin_state - .build_admin_system_config_export_payload() - .await? - } - BackupScope::Users => { - admin_state - .build_admin_system_users_export_payload() - .await? - } - BackupScope::Data => admin_state.build_admin_system_data_export_payload().await?, + let Some(encryption_secret) = effective_backup_encryption_secret(app) else { + return Err(S3BackupTaskError::service_unavailable( + "S3 备份需要 AETHER_BACKUP_ENCRYPTION_KEY 或可用的数据加密密钥", + )); }; - let store = ObjectStoreS3BackupStore::from_config(config) - .map_err(|error| S3BackupTaskError::internal(error.to_string()))?; - run_backup_with_store(config, &store, payload, Utc::now()) + let payload = build_s3_backup_payload_exclusively(app, config.scope).await?; + let store = ObjectStoreS3BackupStore::from_config(config).map_err(|_| { + warn!( + error_category = "object_store_initialization_failed", + "S3 backup object store initialization failed" + ); + S3BackupTaskError::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL) + })?; + run_backup_with_store(config, &store, payload, Utc::now(), &encryption_secret) .await - .map_err(|error| S3BackupTaskError::internal(error.to_string())) + .map_err(|_| { + warn!( + error_category = "backup_execution_failed", + "S3 backup execution failed" + ); + S3BackupTaskError::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL) + }) +} + +async fn build_s3_backup_payload_exclusively( + app: &AppState, + scope: BackupScope, +) -> Result { + let admin_state = AdminAppState::new(app); + crate::admin_api::execute_admin_system_import_exclusively(app, async { + match scope { + BackupScope::Config => { + admin_state + .build_admin_system_config_export_payload(SystemExportMode::RecoveryBackup) + .await + } + BackupScope::Users => { + admin_state + .build_admin_system_users_export_payload(SystemExportMode::RecoveryBackup) + .await + } + BackupScope::Data => { + admin_state + .build_admin_system_data_export_payload(SystemExportMode::RecoveryBackup) + .await + } + } + }) + .await + .map_err(|error| { + warn!( + error_category = "system_import_coordination_failed", + lock_error = ?error, + "S3 backup snapshot could not acquire or retain the system import lock" + ); + S3BackupTaskError::service_unavailable(S3_BACKUP_INTERNAL_ERROR_DETAIL) + })? + .map_err(S3BackupTaskError::from) +} + +fn effective_backup_encryption_secret(app: &AppState) -> Option { + std::env::var("AETHER_BACKUP_ENCRYPTION_KEY") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .or_else(|| { + app.encryption_key() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) } async fn load_s3_backup_config_for_run( @@ -551,7 +732,7 @@ pub(crate) async fn load_s3_backup_config_values( .or_else(|| admin_system_config_default_value(key)); if let Some(value) = value { let value = if *key == "backup_s3_secret_access_key" { - decrypt_s3_secret_access_key(app, value)? + decrypt_s3_secret_access_key(app, value).await? } else { value }; @@ -561,39 +742,56 @@ pub(crate) async fn load_s3_backup_config_values( Ok(values) } -fn decrypt_s3_secret_access_key(app: &AppState, value: Value) -> Result { - let Some(ciphertext) = value +async fn decrypt_s3_secret_access_key( + app: &AppState, + value: Value, +) -> Result { + let Some(stored_value) = value .as_str() .map(str::trim) .filter(|value| !value.is_empty()) else { return Ok(value); }; - let Some(plaintext) = decrypt_catalog_secret_with_fallbacks(app.encryption_key(), ciphertext) - else { - return Err(S3BackupTaskError::bad_request( + let plaintext = decrypt_or_migrate_system_config_secret( + app, + "backup_s3_secret_access_key", + stored_value.to_string(), + ) + .await + .map_err(|_| { + S3BackupTaskError::bad_request( "S3 备份配置无效:Secret Access Key(访问密钥)无法解密,请重新填写", - )); - }; + ) + })?; Ok(Value::String(plaintext)) } fn backup_run_result_json(result: &BackupRunResult) -> Value { json!({ "scope": result.scope.as_config_value(), - "bucket": result.bucket, - "object_key": result.object_key, "bytes": result.bytes, "sha256": result.sha256, "export_version": result.export_version, "exported_at": result.exported_at, "compression": result.compression, - "deleted_old_objects": result.deleted_old_objects, + "encryption": result.encryption, + "legacy_encrypted_copies_created": result.legacy_encrypted_copies_created, + "legacy_encrypted_copies_verified": result.legacy_encrypted_copies_verified, + "legacy_plaintext_objects_deleted": result.legacy_plaintext_objects_deleted, + "legacy_plaintext_objects_retained": result.legacy_plaintext_objects_retained, + "retention_cleanup_candidates": result.retention_cleanup_candidates, + "automatic_deletions": result.legacy_plaintext_objects_deleted, + "object_cleanup_mode": "legacy_plaintext_deleted_after_verified_encryption", + "versioned_storage_cleanup_required": result.versioned_storage_cleanup_required, + "versioned_storage_cleanup_notice": "legacy_plaintext_versions_require_external_cleanup", }) } #[cfg(test)] mod tests { + use std::convert::Infallible; + use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; @@ -603,9 +801,78 @@ mod tests { }; use crate::data::GatewayDataState; + use crate::handlers::shared::decrypt_system_config_secret; use crate::state::AppState; use crate::task_runtime::{now_unix_secs, TASK_KEY_SYSTEM_S3_BACKUP}; + #[test] + fn backup_lock_renewal_requires_ownership_and_preserves_backend_errors() { + assert!(matches!( + super::require_successful_backup_lock_renewal::(Ok(true)), + Ok(()) + )); + assert!(matches!( + super::require_successful_backup_lock_renewal::(Ok(false)), + Err(super::BackupLockRenewalFailure::Lost) + )); + assert!(matches!( + super::require_successful_backup_lock_renewal(Err("redis unavailable")), + Err(super::BackupLockRenewalFailure::Backend( + "redis unavailable" + )) + )); + } + + #[tokio::test] + async fn lost_backup_lock_stops_race_without_publishing_backup_result() { + let destructive_stage_reached = Arc::new(AtomicBool::new(false)); + let destructive_stage_for_backup = Arc::clone(&destructive_stage_reached); + let backup = async move { + std::future::pending::<()>().await; + destructive_stage_for_backup.store(true, Ordering::Release); + Ok::<(), super::S3BackupTaskError>(()) + }; + let heartbeat = tokio::spawn(async {}); + let outcome = super::race_backup_with_lock_heartbeat(backup, heartbeat).await; + assert!(matches!( + outcome, + super::BackupLockRaceOutcome::LeaseLost(Ok(())) + )); + assert!(!destructive_stage_reached.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn completed_heartbeat_wins_when_backup_completion_is_also_ready() { + let heartbeat = tokio::spawn(async {}); + tokio::task::yield_now().await; + + let outcome = super::race_backup_with_lock_heartbeat(async { 42_u8 }, heartbeat).await; + + assert!(matches!( + outcome, + super::BackupLockRaceOutcome::LeaseLost(Ok(())) + )); + } + + #[tokio::test] + async fn s3_backup_snapshot_refuses_to_overlap_system_import() { + let app = AppState::new().expect("app state should build"); + let lease = crate::admin_api::try_acquire_admin_system_import_lease(&app) + .await + .expect("test should acquire the system import lease"); + + let error = super::build_s3_backup_payload_exclusively( + &app, + crate::backup::scopes::BackupScope::Config, + ) + .await + .expect_err("backup snapshot must not overlap a system import"); + + crate::admin_api::release_admin_system_import_lease(&app, &lease).await; + assert_eq!(error.status(), axum::http::StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(error.detail(), super::S3_BACKUP_INTERNAL_ERROR_DETAIL); + } + fn valid_s3_backup_config_values() -> Vec<(String, serde_json::Value)> { vec![ ( @@ -632,6 +899,92 @@ mod tests { ] } + #[tokio::test] + async fn legacy_plaintext_s3_secret_is_migrated_when_config_loads() { + let plaintext = "legacy-s3-secret-access-key"; + let mut entries = valid_s3_backup_config_values(); + entries + .iter_mut() + .find(|(key, _)| key == "backup_s3_secret_access_key") + .expect("secret config fixture should exist") + .1 = serde_json::json!(plaintext); + let app = AppState::new() + .expect("app state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests(entries), + ); + + let values = super::load_s3_backup_config_values(&app) + .await + .expect("legacy S3 config should load"); + assert_eq!( + values.get("backup_s3_secret_access_key"), + Some(&serde_json::json!(plaintext)) + ); + + let stored = app + .read_system_config_json_value_strong("backup_s3_secret_access_key") + .await + .expect("stored S3 secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored S3 secret should remain a string"); + assert_ne!(stored, plaintext); + assert_eq!( + decrypt_system_config_secret(&app, "backup_s3_secret_access_key", &stored) + .expect("migrated S3 secret should decrypt"), + plaintext + ); + } + + #[tokio::test] + async fn undecryptable_s3_fernet_secret_fails_closed() { + let plaintext = "s3-secret-from-unavailable-key"; + let ciphertext = encrypt_python_fernet_plaintext("unavailable-s3-key", plaintext) + .expect("unknown-key fixture should encrypt"); + let mut entries = valid_s3_backup_config_values(); + entries + .iter_mut() + .find(|(key, _)| key == "backup_s3_secret_access_key") + .expect("secret config fixture should exist") + .1 = serde_json::json!(ciphertext.clone()); + let app = AppState::new() + .expect("app state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests(entries), + ); + + let error = super::load_s3_backup_config_values(&app) + .await + .expect_err("unknown-key S3 ciphertext must fail closed"); + let error_text = error.to_string(); + assert!(!error_text.contains(plaintext)); + assert!(!error_text.contains(&ciphertext)); + assert_eq!( + app.read_system_config_json_value_strong("backup_s3_secret_access_key") + .await + .expect("stored S3 secret should read"), + Some(serde_json::json!(ciphertext)) + ); + } + + #[test] + fn gateway_dependency_errors_are_not_exposed_to_backup_clients() { + let error = super::S3BackupTaskError::from(crate::GatewayError::Internal( + "postgresql://admin:database-secret@db.internal/aether".to_string(), + )); + + assert_eq!( + error.status(), + axum::http::StatusCode::INTERNAL_SERVER_ERROR + ); + assert_eq!(error.detail(), super::S3_BACKUP_INTERNAL_ERROR_DETAIL); + assert!(!error.detail().contains("database-secret")); + } + fn stored_s3_backup_run(status: BackgroundTaskStatus) -> StoredBackgroundTaskRun { let now = now_unix_secs(); StoredBackgroundTaskRun { @@ -774,7 +1127,9 @@ mod tests { let payload = super::s3_backup_task_payload_json(&config, "manual", None); - assert!(payload["bucket"].is_string()); + assert!(payload.get("bucket").is_none()); + assert!(payload.get("prefix").is_none()); + assert_eq!(payload["scope"], serde_json::json!("data")); assert_eq!(payload["trigger"], serde_json::json!("manual")); assert!(!payload.to_string().contains("secret")); } diff --git a/apps/aether-gateway/src/backup/worker.rs b/apps/aether-gateway/src/backup/worker.rs index 9fe12d156..e43d942ca 100644 --- a/apps/aether-gateway/src/backup/worker.rs +++ b/apps/aether-gateway/src/backup/worker.rs @@ -29,8 +29,11 @@ pub(crate) fn spawn_s3_backup_worker(app: AppState) -> Option> { interval.tick().await; loop { interval.tick().await; - if let Err(error) = run_s3_backup_schedule_tick(&app, Utc::now()).await { - warn!(error = ?error, "S3 backup schedule tick failed"); + if run_s3_backup_schedule_tick(&app, Utc::now()).await.is_err() { + warn!( + error_category = "schedule_tick_failed", + "S3 backup schedule tick failed" + ); } } }, @@ -43,15 +46,21 @@ async fn run_s3_backup_schedule_tick( ) -> Result<(), GatewayError> { let values = match super::task::load_s3_backup_config_values(app).await { Ok(values) => values, - Err(error) => { - warn!(error = %error, "S3 backup schedule config load failed"); + Err(_) => { + warn!( + error_category = "config_load_failed", + "S3 backup schedule config load failed" + ); return Ok(()); } }; let config = match S3BackupConfig::from_json_map(&values) { Ok(config) => config, - Err(error) => { - warn!(error = %error, "S3 backup schedule config is invalid"); + Err(_) => { + warn!( + error_category = "config_invalid", + "S3 backup schedule config is invalid" + ); return Ok(()); } }; @@ -68,8 +77,11 @@ async fn run_s3_backup_schedule_tick( match super::task::start_s3_backup_task_for_schedule(app.clone(), slot).await { Ok(_) => {} - Err(error) => { - warn!(error = %error, "S3 backup scheduled task submission failed"); + Err(_) => { + warn!( + error_category = "task_submission_failed", + "S3 backup scheduled task submission failed" + ); } } Ok(()) diff --git a/apps/aether-gateway/src/bark_push.rs b/apps/aether-gateway/src/bark_push.rs index cbf6e09c7..5b3de17dc 100644 --- a/apps/aether-gateway/src/bark_push.rs +++ b/apps/aether-gateway/src/bark_push.rs @@ -1,8 +1,10 @@ use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, system_config_bool, system_config_string, + bark_device_key_binding, canonical_bark_server_url, decrypt_or_migrate_bark_device_key, + system_config_bool, system_config_string, }; use crate::{AppState, GatewayError}; use serde_json::{json, Value}; +use std::net::{IpAddr, SocketAddr}; pub(crate) const BARK_PUSH_ENABLED_KEY: &str = "module.bark_push.enabled"; pub(crate) const BARK_PUSH_DEVICE_KEY_KEY: &str = "module.bark_push.device_key"; @@ -10,8 +12,20 @@ pub(crate) const BARK_PUSH_SERVER_URL_KEY: &str = "module.bark_push.server_url"; pub(crate) const BARK_PUSH_TEMPLATE_KEY: &str = "module.bark_push.template"; const DEFAULT_BARK_API_BASE: &str = "https://api.day.app"; +const BARK_ALLOW_HTTP_ENV: &str = "AETHER_BARK_ALLOW_HTTP"; +const BARK_ALLOW_PRIVATE_TARGETS_ENV: &str = "AETHER_BARK_ALLOW_PRIVATE_TARGETS"; +const MAX_BARK_RESPONSE_BYTES: usize = 64 * 1024; +const BARK_CONNECT_TIMEOUT_MS: u64 = 10_000; +const BARK_REQUEST_TIMEOUT_MS: u64 = 300_000; +const MAX_BARK_SERVER_URL_BYTES: usize = 2 * 1024; +const MAX_BARK_DEVICE_KEY_BYTES: usize = 512; +const MAX_BARK_TEMPLATE_BYTES: usize = 256 * 1024; +const MAX_BARK_TITLE_BYTES: usize = 512; +const MAX_BARK_BODY_BYTES: usize = 2 * 1024 * 1024; +const MAX_BARK_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024; +const MAX_BARK_RESOLVED_ADDRESSES: usize = 32; -#[derive(Debug, Clone)] +#[derive(Clone)] pub(crate) struct BarkPushConfig { pub(crate) enabled: bool, pub(crate) device_key: Option, @@ -19,6 +33,21 @@ pub(crate) struct BarkPushConfig { pub(crate) template: Option, } +impl std::fmt::Debug for BarkPushConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("BarkPushConfig") + .field("enabled", &self.enabled) + .field( + "device_key", + &self.device_key.as_ref().map(|_| "[REDACTED]"), + ) + .field("server_url", &self.server_url) + .field("template", &self.template.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + pub(crate) async fn bark_push_module_enabled(state: &AppState) -> Result { let value = state .read_system_config_json_value(BARK_PUSH_ENABLED_KEY) @@ -35,23 +64,34 @@ pub(crate) async fn read_bark_push_config( state: &AppState, ) -> Result { let enabled = bark_push_module_enabled(state).await?; - let device_key = state - .read_system_config_json_value(BARK_PUSH_DEVICE_KEY_KEY) - .await? - .and_then(|value| system_config_string(Some(&value))) - .map(|value| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value) - }); let server_url = state .read_system_config_json_value(BARK_PUSH_SERVER_URL_KEY) .await? .and_then(|value| system_config_string(Some(&value))) .filter(|value| !value.trim().is_empty()) .unwrap_or_else(|| DEFAULT_BARK_API_BASE.to_string()); + validate_bark_config_field("server_url", &server_url, MAX_BARK_SERVER_URL_BYTES)?; + let server_url = normalized_bark_server_url(&server_url)?; + let binding = bark_device_key_binding(&server_url) + .ok_or_else(|| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string()))?; + let device_key = state + .read_system_config_json_value(BARK_PUSH_DEVICE_KEY_KEY) + .await? + .and_then(|value| system_config_string(Some(&value))); + let device_key = match device_key { + Some(value) => Some(decrypt_or_migrate_bark_device_key(state, &binding, value).await?), + None => None, + }; + if let Some(device_key) = device_key.as_deref() { + validate_bark_config_field("device_key", device_key, MAX_BARK_DEVICE_KEY_BYTES)?; + } let template = state .read_system_config_json_value(BARK_PUSH_TEMPLATE_KEY) .await? .and_then(|value| system_config_string(Some(&value))); + if let Some(template) = template.as_deref() { + validate_bark_config_field("template", template, MAX_BARK_TEMPLATE_BYTES)?; + } Ok(BarkPushConfig { enabled, @@ -62,7 +102,7 @@ pub(crate) async fn read_bark_push_config( } pub(crate) async fn send_bark_push( - state: &AppState, + _state: &AppState, config: &BarkPushConfig, title: &str, markdown_body: &str, @@ -76,11 +116,13 @@ pub(crate) async fn send_bark_push( "Bark Device Key 不能为空".to_string(), )); } - let server_url = normalized_bark_server_url(&config.server_url)?; - let body = render_bark_body(config.template.as_deref(), title, markdown_body); - let response = state - .client - .post(format!("{server_url}/push")) + validate_bark_config_field("device_key", device_key, MAX_BARK_DEVICE_KEY_BYTES)?; + validate_bark_content_field("title", title, MAX_BARK_TITLE_BYTES)?; + validate_bark_content_field("body", markdown_body, MAX_BARK_BODY_BYTES)?; + let (client, push_url) = build_bark_push_client_and_url(&config.server_url).await?; + let body = render_bark_body(config.template.as_deref(), title, markdown_body)?; + let response = client + .post(push_url) .json(&json!({ "device_key": device_key, "title": title, @@ -88,16 +130,14 @@ pub(crate) async fn send_bark_push( })) .send() .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + .map_err(|err| GatewayError::Internal(bark_request_error_message(&err)))?; let status = response.status(); - let text = response - .text() + let body = aether_http::read_response_bytes_with_limit(response, MAX_BARK_RESPONSE_BYTES) .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + .map_err(|err| GatewayError::Internal(bark_response_body_error_message(&err)))?; + let text = String::from_utf8_lossy(&body); if !status.is_success() { - return Err(GatewayError::Internal(format!( - "Bark 返回 HTTP {status}: {text}" - ))); + return Err(GatewayError::Internal(format!("Bark 返回 HTTP {status}"))); } if let Ok(payload) = serde_json::from_str::(&text) { let code_is_ok = payload @@ -114,53 +154,259 @@ pub(crate) async fn send_bark_push( }) .unwrap_or(true); if !code_is_ok { - return Err(GatewayError::Internal(format!("Bark 返回失败: {payload}"))); + return Err(GatewayError::Internal("Bark 返回失败".to_string())); } } Ok(()) } -fn normalized_bark_server_url(server_url: &str) -> Result { - let server_url = server_url.trim().trim_end_matches('/'); - if server_url.is_empty() { - return Err(GatewayError::Internal( - "Bark 服务器地址不能为空".to_string(), - )); - } - if !server_url.starts_with("https://") && !server_url.starts_with("http://") { - return Err(GatewayError::Internal( - "Bark 服务器地址必须以 http:// 或 https:// 开头".to_string(), - )); - } - Ok(server_url.to_string()) +fn bark_request_error_message(error: &reqwest::Error) -> String { + format!("Bark 请求失败 ({})", bark_reqwest_error_kind(error)) } -fn render_bark_body(template: Option<&str>, title: &str, markdown_body: &str) -> String { - match template { - Some(template) if !template.trim().is_empty() => template - .replace("{title}", title) - .replace("{body}", markdown_body), - _ => markdown_body.to_string(), +fn bark_response_body_error_message(error: &aether_http::ResponseBodyReadError) -> String { + match error { + aether_http::ResponseBodyReadError::TooLarge { max_bytes } => { + format!("Bark 响应超过 {max_bytes} 字节") + } + aether_http::ResponseBodyReadError::Read(error) => { + format!("Bark 响应读取失败 ({})", bark_reqwest_error_kind(error)) + } } } +fn bark_reqwest_error_kind(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else { + "transport" + } +} + +fn normalized_bark_server_url(server_url: &str) -> Result { + canonical_bark_server_url(server_url) + .ok_or_else(|| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string())) +} + +async fn build_bark_push_client_and_url( + server_url: &str, +) -> Result<(reqwest::Client, url::Url), GatewayError> { + validate_bark_config_field("server_url", server_url, MAX_BARK_SERVER_URL_BYTES)?; + let normalized = normalized_bark_server_url(server_url)?; + let mut push_url = url::Url::parse(&normalized) + .map_err(|_| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string()))?; + validate_bark_transport_policy(&push_url, env_flag_enabled(BARK_ALLOW_HTTP_ENV))?; + + let host = push_url + .host_str() + .ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少主机名".to_string()))? + .to_string(); + let port = push_url + .port_or_known_default() + .ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?; + let addresses = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + tokio::time::timeout( + std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS), + tokio::net::lookup_host((host.as_str(), port)), + ) + .await + .map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析超时".to_string()))? + .map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析失败".to_string()))? + .take(MAX_BARK_RESOLVED_ADDRESSES) + .collect::>() + }; + let allow_benchmarking_ip = push_url.scheme() == "https" + && push_url.port_or_known_default() == Some(443) + && host.eq_ignore_ascii_case("api.day.app"); + validate_bark_resolved_addresses( + &addresses, + env_flag_enabled(BARK_ALLOW_PRIVATE_TARGETS_ENV), + allow_benchmarking_ip, + )?; + + push_url + .path_segments_mut() + .map_err(|_| GatewayError::Internal("Bark 服务器地址不能作为基础 URL".to_string()))? + .pop_if_empty() + .push("push"); + + let mut builder = aether_http::apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &aether_http::HttpClientConfig { + connect_timeout_ms: Some(BARK_CONNECT_TIMEOUT_MS), + request_timeout_ms: Some(BARK_REQUEST_TIMEOUT_MS), + http2_adaptive_window: true, + ..aether_http::HttpClientConfig::default() + }, + ); + if host.parse::().is_err() { + builder = builder.resolve_to_addrs(&host, &addresses); + } + let client = builder + .build() + .map_err(|_| GatewayError::Internal("Bark HTTP 客户端初始化失败".to_string()))?; + Ok((client, push_url)) +} + +fn validate_bark_transport_policy(url: &url::Url, allow_http: bool) -> Result<(), GatewayError> { + if url.scheme() == "http" && !allow_http { + return Err(GatewayError::Internal(format!( + "Bark 服务器必须使用 HTTPS;如确需明文 HTTP,请显式设置 {BARK_ALLOW_HTTP_ENV}=true" + ))); + } + Ok(()) +} + +fn validate_bark_resolved_addresses( + addresses: &[SocketAddr], + allow_private: bool, + allow_benchmarking_ip: bool, +) -> Result<(), GatewayError> { + if addresses.is_empty() { + return Err(GatewayError::Internal( + "Bark 服务器 DNS 解析未返回地址".to_string(), + )); + } + if !allow_private + && addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(allow_benchmarking_ip + && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) + { + return Err(GatewayError::Internal(format!( + "Bark 服务器解析到私有或保留地址;如确需内网自建服务,请显式设置 {BARK_ALLOW_PRIVATE_TARGETS_ENV}=true" + ))); + } + Ok(()) +} + +fn env_flag_enabled(key: &str) -> bool { + std::env::var(key).ok().is_some_and(|value| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "1" | "true" | "yes" | "on" + ) + }) +} + +fn validate_bark_config_field( + field: &str, + value: &str, + max_bytes: usize, +) -> Result<(), GatewayError> { + if value.len() > max_bytes || value.bytes().any(|byte| byte == 0) { + return Err(GatewayError::Internal(format!( + "Bark {field} exceeds the allowed size or contains a NUL byte" + ))); + } + Ok(()) +} + +fn validate_bark_content_field( + field: &str, + value: &str, + max_bytes: usize, +) -> Result<(), GatewayError> { + validate_bark_config_field(field, value, max_bytes) +} + +fn render_bark_body( + template: Option<&str>, + title: &str, + markdown_body: &str, +) -> Result { + validate_bark_content_field("title", title, MAX_BARK_TITLE_BYTES)?; + validate_bark_content_field("body", markdown_body, MAX_BARK_BODY_BYTES)?; + let template = template + .filter(|value| !value.trim().is_empty()) + .unwrap_or("{body}"); + validate_bark_config_field("template", template, MAX_BARK_TEMPLATE_BYTES)?; + + let mut rendered = String::with_capacity(template.len().min(MAX_BARK_RENDERED_BODY_BYTES)); + let mut cursor = 0usize; + while cursor < template.len() { + let remaining = &template[cursor..]; + let title_match = remaining.find("{title}"); + let body_match = remaining.find("{body}"); + let next = match (title_match, body_match) { + (None, None) => { + append_bark_rendered_part(&mut rendered, remaining)?; + cursor = template.len(); + continue; + } + (Some(index), None) => (index, "{title}", title), + (None, Some(index)) => (index, "{body}", markdown_body), + (Some(title_index), Some(body_index)) if title_index <= body_index => { + (title_index, "{title}", title) + } + (Some(_), Some(body_index)) => (body_index, "{body}", markdown_body), + }; + append_bark_rendered_part(&mut rendered, &remaining[..next.0])?; + append_bark_rendered_part(&mut rendered, next.2)?; + cursor += next.0 + next.1.len(); + } + if rendered.is_empty() && template.is_empty() { + return Ok(String::new()); + } + Ok(rendered) +} + +fn append_bark_rendered_part(output: &mut String, part: &str) -> Result<(), GatewayError> { + let next_len = output + .len() + .checked_add(part.len()) + .ok_or_else(|| GatewayError::Internal("Bark rendered body is too large".to_string()))?; + if next_len > MAX_BARK_RENDERED_BODY_BYTES { + return Err(GatewayError::Internal( + "Bark rendered body exceeds the allowed size".to_string(), + )); + } + output.push_str(part); + Ok(()) +} + #[cfg(test)] mod tests { - use super::{normalized_bark_server_url, render_bark_body}; + use super::{ + bark_request_error_message, bark_response_body_error_message, normalized_bark_server_url, + render_bark_body, validate_bark_resolved_addresses, validate_bark_transport_policy, + }; + use std::net::SocketAddr; #[test] fn bark_body_uses_template_when_provided() { - let rendered = render_bark_body(Some("{title}\n\n{body}"), "告警", "原始正文"); + let rendered = render_bark_body(Some("{title}\n\n{body}"), "告警", "原始正文") + .expect("template should render"); assert_eq!(rendered, "告警\n\n原始正文"); } #[test] fn bark_body_falls_back_to_markdown_body_for_empty_template() { - assert_eq!(render_bark_body(None, "告警", "原始正文"), "原始正文"); assert_eq!( - render_bark_body(Some(" "), "告警", "原始正文"), + render_bark_body(None, "告警", "原始正文").expect("fallback should render"), "原始正文" ); + assert_eq!( + render_bark_body(Some(" "), "告警", "原始正文").expect("fallback should render"), + "原始正文" + ); + } + + #[test] + fn bark_body_rejects_template_expansion_bombs_and_oversized_content() { + let template = "x".repeat(super::MAX_BARK_TEMPLATE_BYTES + 1); + assert!(render_bark_body(Some(&template), "告警", "正文").is_err()); + let body = "x".repeat(super::MAX_BARK_BODY_BYTES + 1); + assert!(render_bark_body(None, "告警", &body).is_err()); } #[test] @@ -170,4 +416,61 @@ mod tests { "https://api.day.app" ); } + + #[test] + fn bark_server_url_rejects_credentials_query_and_fragments() { + for invalid in [ + "https://user@example.com", + "https://example.com?target=internal", + "https://example.com/#fragment", + ] { + assert!(normalized_bark_server_url(invalid).is_err(), "{invalid}"); + } + } + + #[test] + fn bark_http_transport_requires_explicit_opt_in() { + let url = url::Url::parse("http://bark.example.com").unwrap(); + assert!(validate_bark_transport_policy(&url, false).is_err()); + assert!(validate_bark_transport_policy(&url, true).is_ok()); + } + + #[test] + fn bark_private_targets_require_explicit_opt_in() { + let private = [SocketAddr::from(([127, 0, 0, 1], 443))]; + assert!(validate_bark_resolved_addresses(&private, false, false).is_err()); + assert!(validate_bark_resolved_addresses(&private, true, false).is_ok()); + } + + #[test] + fn bark_builtin_server_allows_benchmarking_ip_only_with_https_default_port() { + let fake = [SocketAddr::from(([198, 18, 75, 234], 443))]; + assert!(validate_bark_resolved_addresses(&fake, false, true).is_ok()); + assert!(validate_bark_resolved_addresses( + &[fake[0], SocketAddr::from(([127, 0, 0, 1], 443))], + false, + true, + ) + .is_err()); + assert!(validate_bark_resolved_addresses(&fake, false, false).is_err()); + } + + #[tokio::test] + async fn bark_transport_errors_do_not_expose_server_url_or_response_body() { + let secret = "bark-secret-query"; + let error = reqwest::Client::new() + .post(format!("ftp://bark.example.test/push?token={secret}")) + .send() + .await + .expect_err("unsupported URL scheme should fail before network I/O"); + + let message = bark_request_error_message(&error); + assert!(!message.contains(secret)); + assert!(!message.contains("bark.example.test")); + + let body_error = aether_http::ResponseBodyReadError::Read(error); + let message = bark_response_body_error_message(&body_error); + assert!(!message.contains(secret)); + assert!(!message.contains("bark.example.test")); + } } diff --git a/apps/aether-gateway/src/bin/aether-backup-restore.rs b/apps/aether-gateway/src/bin/aether-backup-restore.rs new file mode 100644 index 000000000..2c20deb47 --- /dev/null +++ b/apps/aether-gateway/src/bin/aether-backup-restore.rs @@ -0,0 +1,1047 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::fs::{self, File, OpenOptions}; +use std::io::{Read, Write}; +use std::path::{Path, PathBuf}; + +use aether_gateway::{ + restore_backup_json, BackupDecryptionKey, BackupRestoreLimits, + DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES, DEFAULT_BACKUP_MAX_JSON_BYTES, +}; +use clap::Parser; +use serde::Deserialize; +use sha2::{Digest, Sha256}; +use uuid::Uuid; + +const HISTORICAL_KEYS_ENV: &str = "AETHER_BACKUP_HISTORICAL_KEYS_JSON"; +const MAX_SECRET_FILE_BYTES: usize = 1024 * 1024; +const MAX_KEY_FILES: usize = 16; +const MAX_V2_KEY_CANDIDATES: usize = 256; +const MAX_LEGACY_V1_KEY_CANDIDATES: usize = 16; +const AUTOMATIC_KEY_ENV_VARS: [(&str, bool); 3] = [ + ("AETHER_BACKUP_ENCRYPTION_KEY", false), + ("AETHER_GATEWAY_DATA_ENCRYPTION_KEY", true), + ("ENCRYPTION_KEY", true), +]; + +#[derive(Debug, Parser)] +#[command( + name = "aether-backup-restore", + about = "Decrypt and verify an Aether S3 backup into a local JSON file" +)] +struct Args { + /// Local encrypted .json.zst.aes256gcm file. + #[arg(long)] + input: PathBuf, + + /// Complete canonical S3 object key used when the backup was encrypted. + #[arg(long)] + object_key: String, + + /// Destination for verified JSON. Existing files are rejected by default. + #[arg(long)] + output: PathBuf, + + /// Plaintext secret file; may be repeated. The secret itself is never accepted as an argument. + #[arg(long = "key-file")] + key_files: Vec, + + /// Structured JSON keyring file. May also be set with AETHER_BACKUP_KEYRING_FILE. + #[arg(long, env = "AETHER_BACKUP_KEYRING_FILE")] + keyring_file: Option, + + /// Replace an existing output file atomically. + #[arg(long)] + overwrite: bool, + + /// Maximum encrypted input size in MiB. + #[arg(long, default_value_t = mib(DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES), value_parser = clap::value_parser!(u64).range(1..=4096))] + max_encrypted_mib: u64, + + /// Maximum decompressed JSON size in MiB. + #[arg(long, default_value_t = mib(DEFAULT_BACKUP_MAX_JSON_BYTES), value_parser = clap::value_parser!(u64).range(1..=8192))] + max_json_mib: u64, +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct KeyringDocument { + version: u8, + #[serde(default)] + keys: Vec, + #[serde(default)] + legacy_v1: Vec, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum KeyringSecret { + Direct(String), + Named { + #[serde(alias = "key")] + secret: String, + }, +} + +impl KeyringSecret { + fn into_secret(self) -> String { + match self { + Self::Direct(secret) | Self::Named { secret } => secret, + } + } +} + +#[derive(Debug, thiserror::Error)] +enum CliError { + #[error("{0}")] + Message(String), + + #[error("file operation failed: {0}")] + Io(#[from] std::io::Error), + + #[error("backup restore failed: {0}")] + Restore(#[from] aether_gateway::BackupRestoreError), +} + +fn main() { + if let Err(error) = run(Args::parse()) { + eprintln!("{error}"); + std::process::exit(1); + } +} + +fn run(args: Args) -> Result<(), CliError> { + reject_output_aliases(&args)?; + let limits = BackupRestoreLimits { + max_encrypted_bytes: checked_mib(args.max_encrypted_mib)?, + max_json_bytes: checked_mib(args.max_json_mib)?, + }; + let encrypted = read_limited_file(&args.input, limits.max_encrypted_bytes, true, false)?; + let candidates = load_key_candidates(&args)?; + if candidates.is_empty() { + return Err(CliError::Message(format!( + "no backup keys configured; set AETHER_BACKUP_ENCRYPTION_KEY, AETHER_GATEWAY_DATA_ENCRYPTION_KEY, or ENCRYPTION_KEY, or use a protected key/keyring file ({HISTORICAL_KEYS_ENV} is also supported)" + ))); + } + + let cipher_sha256 = format!("{:x}", Sha256::digest(&encrypted)); + let restored = restore_backup_json(&args.object_key, &encrypted, &candidates, limits)?; + write_atomic_private(&args.output, restored.json_bytes(), args.overwrite)?; + let restored_scope = restored.scope().as_str(); + + let summary = serde_json::json!({ + "status": "verified_json_written", + "object_key": args.object_key, + "output": args.output.display().to_string(), + "cipher_sha256": cipher_sha256, + "envelope_version": restored.envelope_version, + "key_id": restored.key_id, + "export_version": restored.export_version, + "exported_at": restored.exported_at, + "scope": restored_scope, + "database_applied": false, + }); + println!( + "{}", + serde_json::to_string(&summary).map_err(|error| { + CliError::Message(format!("could not serialize restore summary: {error}")) + })? + ); + Ok(()) +} + +fn reject_output_aliases(args: &Args) -> Result<(), CliError> { + let Ok(output) = fs::canonicalize(&args.output) else { + return Ok(()); + }; + let mut protected_inputs = Vec::with_capacity(args.key_files.len() + 2); + protected_inputs.push(&args.input); + protected_inputs.extend(args.key_files.iter()); + if let Some(keyring_file) = &args.keyring_file { + protected_inputs.push(keyring_file); + } + for input in protected_inputs { + if fs::canonicalize(input).is_ok_and(|canonical| canonical == output) { + return Err(CliError::Message(format!( + "output {} must not replace the encrypted input or a key file", + args.output.display() + ))); + } + } + Ok(()) +} + +const fn mib(bytes: usize) -> u64 { + (bytes / (1024 * 1024)) as u64 +} + +fn checked_mib(value: u64) -> Result { + value + .checked_mul(1024 * 1024) + .and_then(|bytes| usize::try_from(bytes).ok()) + .ok_or_else(|| CliError::Message("configured size limit is too large".to_string())) +} + +fn load_key_candidates(args: &Args) -> Result, CliError> { + if args.key_files.len() > MAX_KEY_FILES { + return Err(CliError::Message(format!( + "at most {MAX_KEY_FILES} --key-file values are allowed" + ))); + } + let mut values = Vec::<(String, bool)>::new(); + append_automatic_environment_keys(&mut values, |env_name| std::env::var(env_name).ok())?; + if let Ok(value) = std::env::var(HISTORICAL_KEYS_ENV) { + append_keyring_document( + &mut values, + parse_keyring(value.as_bytes(), HISTORICAL_KEYS_ENV)?, + )?; + } + if let Some(path) = &args.keyring_file { + let bytes = read_limited_file(path, MAX_SECRET_FILE_BYTES, true, true)?; + append_keyring_document( + &mut values, + parse_keyring(&bytes, &path.display().to_string())?, + )?; + } + for path in &args.key_files { + let bytes = read_limited_file(path, MAX_SECRET_FILE_BYTES, true, true)?; + let value = String::from_utf8(bytes).map_err(|_| { + CliError::Message(format!("key file {} is not valid UTF-8", path.display())) + })?; + push_secret(&mut values, value, true)?; + } + + let unique = deduplicate_and_validate_key_values(values)?; + unique + .into_iter() + .map(|(secret, allow_v1)| { + if allow_v1 { + BackupDecryptionKey::historical(secret) + } else { + BackupDecryptionKey::v2_only(secret) + } + .map_err(CliError::from) + }) + .collect() +} + +fn append_automatic_environment_keys( + values: &mut Vec<(String, bool)>, + mut get_env: impl FnMut(&str) -> Option, +) -> Result<(), CliError> { + for (env_name, allow_v1) in AUTOMATIC_KEY_ENV_VARS { + if let Some(value) = get_env(env_name) { + push_secret(values, value, allow_v1)?; + } + } + Ok(()) +} + +fn deduplicate_and_validate_key_values( + values: Vec<(String, bool)>, +) -> Result, CliError> { + let mut indexes = HashMap::::new(); + let mut unique = Vec::<(String, bool)>::new(); + let mut legacy_v1_count = 0_usize; + for (secret, allow_v1) in values { + if let Some(index) = indexes.get(&secret).copied() { + if allow_v1 && !unique[index].1 { + unique[index].1 = true; + legacy_v1_count += 1; + } + } else { + let index = unique.len(); + indexes.insert(secret.clone(), index); + unique.push((secret, allow_v1)); + legacy_v1_count += usize::from(allow_v1); + + if unique.len() > MAX_V2_KEY_CANDIDATES { + return Err(CliError::Message(format!( + "backup restore allows at most {MAX_V2_KEY_CANDIDATES} v2 candidate keys" + ))); + } + } + + if legacy_v1_count > MAX_LEGACY_V1_KEY_CANDIDATES { + return Err(CliError::Message(format!( + "legacy v1 backup restore allows at most {MAX_LEGACY_V1_KEY_CANDIDATES} candidate keys" + ))); + } + } + Ok(unique) +} + +fn parse_keyring(bytes: &[u8], source: &str) -> Result { + serde_json::from_slice(bytes).map_err(|error| { + CliError::Message(format!( + "keyring {source} is not valid structured JSON: {error}" + )) + }) +} + +fn append_keyring_document( + values: &mut Vec<(String, bool)>, + keyring: KeyringDocument, +) -> Result<(), CliError> { + if keyring.version != 1 { + return Err(CliError::Message(format!( + "unsupported backup keyring version {}", + keyring.version + ))); + } + for secret in keyring.keys { + push_secret(values, secret.into_secret(), false)?; + } + for secret in keyring.legacy_v1 { + push_secret(values, secret.into_secret(), true)?; + } + Ok(()) +} + +fn push_secret( + values: &mut Vec<(String, bool)>, + value: String, + allow_v1: bool, +) -> Result<(), CliError> { + let value = value.trim().to_string(); + if value.is_empty() || value.contains('\0') { + return Err(CliError::Message( + "backup key material must not be empty or contain NUL".to_string(), + )); + } + values.push((value, allow_v1)); + Ok(()) +} + +fn read_limited_file( + path: &Path, + limit: usize, + reject_symlink: bool, + require_private_permissions: bool, +) -> Result, CliError> { + let mut file = open_file_without_following_symlinks(path, reject_symlink)?; + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(CliError::Message(format!( + "{} is not a regular file", + path.display() + ))); + } + if require_private_permissions { + validate_secret_file_permissions(path, &metadata)?; + } + #[cfg(not(unix))] + if reject_symlink && fs::symlink_metadata(path)?.file_type().is_symlink() { + return Err(CliError::Message(format!( + "input file {} changed to a symbolic link while being opened", + path.display() + ))); + } + if metadata.len() > u64::try_from(limit).unwrap_or(u64::MAX) { + return Err(CliError::Message(format!( + "{} exceeds the configured {} byte limit", + path.display(), + limit + ))); + } + let read_limit = u64::try_from(limit).unwrap_or(u64::MAX).saturating_add(1); + let mut bytes = Vec::with_capacity((metadata.len() as usize).min(limit)); + Read::by_ref(&mut file) + .take(read_limit) + .read_to_end(&mut bytes)?; + if bytes.len() > limit { + return Err(CliError::Message(format!( + "{} exceeds the configured {} byte limit", + path.display(), + limit + ))); + } + Ok(bytes) +} + +fn open_file_without_following_symlinks( + path: &Path, + reject_symlink: bool, +) -> Result { + #[cfg(unix)] + if reject_symlink { + return open_file_beneath_real_directories(path); + } + + #[cfg(not(unix))] + if reject_symlink && fs::symlink_metadata(path)?.file_type().is_symlink() { + return Err(CliError::Message(format!( + "input file {} must not be a symbolic link", + path.display() + ))); + } + let mut options = OpenOptions::new(); + options.read(true); + #[cfg(windows)] + if reject_symlink { + use std::os::windows::fs::OpenOptionsExt; + const FILE_FLAG_OPEN_REPARSE_POINT: u32 = 0x0020_0000; + options.custom_flags(FILE_FLAG_OPEN_REPARSE_POINT); + } + options.open(path).map_err(CliError::from) +} + +#[cfg(unix)] +fn open_file_beneath_real_directories(path: &Path) -> Result { + use std::os::fd::{AsRawFd, FromRawFd}; + + let (parent, file_name) = open_real_parent_directory(path)?; + let file_name = unix_path_component(&file_name, "input file name")?; + // SAFETY: `parent` is a valid open directory descriptor, `file_name` is NUL-terminated, + // and ownership of a successful descriptor is transferred immediately to `File`. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + file_name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_NOFOLLOW, + ) + }; + if descriptor < 0 { + let error = std::io::Error::last_os_error(); + if matches!( + error.raw_os_error(), + Some(libc::ELOOP) | Some(libc::ENOTDIR) + ) && unix_file_mode_at(&parent, &file_name)? + .is_some_and(|mode| mode & libc::S_IFMT == libc::S_IFLNK) + { + return Err(CliError::Message(format!( + "input file {} must not be a symbolic link", + path.display() + ))); + } + return Err(CliError::Io(error)); + } + // SAFETY: `openat` returned a new owned descriptor and no other owner exists. + Ok(unsafe { File::from_raw_fd(descriptor) }) +} + +#[cfg(unix)] +fn open_real_parent_directory(path: &Path) -> Result<(File, OsString), CliError> { + let file_name = path.file_name().map(OsString::from).ok_or_else(|| { + CliError::Message(format!( + "path {} must include a regular file name", + path.display() + )) + })?; + let parent = path + .parent() + .filter(|path| !path.as_os_str().is_empty()) + .unwrap_or(Path::new(".")); + Ok((open_real_directory(parent)?, file_name)) +} + +#[cfg(unix)] +fn open_real_directory(path: &Path) -> Result { + use std::os::fd::{AsRawFd, FromRawFd}; + use std::path::Component; + + let mut directory = File::open(if path.is_absolute() { "/" } else { "." })?; + for component in path.components() { + let name = match component { + Component::RootDir | Component::CurDir => continue, + Component::Normal(name) => name, + Component::ParentDir => { + return Err(CliError::Message(format!( + "path {} must not contain '..' components", + path.display() + ))) + } + Component::Prefix(_) => { + return Err(CliError::Message(format!( + "path {} uses an unsupported prefix", + path.display() + ))) + } + }; + let name = unix_path_component(name, "directory component")?; + // SAFETY: `directory` is a valid open directory descriptor and `name` is a valid + // NUL-terminated component. A successful descriptor is immediately owned by `File`. + let descriptor = unsafe { + libc::openat( + directory.as_raw_fd(), + name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_DIRECTORY, + ) + }; + if descriptor < 0 { + let error = std::io::Error::last_os_error(); + if matches!( + error.raw_os_error(), + Some(libc::ELOOP) | Some(libc::ENOTDIR) + ) && unix_file_mode_at(&directory, &name)? + .is_some_and(|mode| mode & libc::S_IFMT == libc::S_IFLNK) + { + return Err(CliError::Message(format!( + "path {} must not contain symbolic-link directory components", + path.display() + ))); + } + return Err(CliError::Io(error)); + } + // SAFETY: `openat` returned a new owned descriptor and no other owner exists. + directory = unsafe { File::from_raw_fd(descriptor) }; + } + Ok(directory) +} + +#[cfg(unix)] +fn unix_path_component( + component: &std::ffi::OsStr, + description: &str, +) -> Result { + use std::os::unix::ffi::OsStrExt; + + std::ffi::CString::new(component.as_bytes()).map_err(|_| { + CliError::Message(format!( + "{description} must not contain an embedded NUL byte" + )) + }) +} + +#[cfg(unix)] +fn validate_secret_file_permissions(path: &Path, metadata: &fs::Metadata) -> Result<(), CliError> { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + if metadata.permissions().mode() & 0o077 != 0 { + return Err(CliError::Message(format!( + "secret file {} is accessible by group or other users; require mode 0600 or stricter", + path.display() + ))); + } + // SAFETY: `geteuid` has no preconditions and does not retain pointers or borrowed state. + let effective_uid = unsafe { libc::geteuid() }; + if metadata.uid() != effective_uid { + return Err(CliError::Message(format!( + "secret file {} must be owned by the current effective user (uid {effective_uid})", + path.display() + ))); + } + Ok(()) +} + +#[cfg(not(unix))] +fn validate_secret_file_permissions( + _path: &Path, + _metadata: &fs::Metadata, +) -> Result<(), CliError> { + Ok(()) +} + +#[cfg(unix)] +fn write_atomic_private(path: &Path, bytes: &[u8], overwrite: bool) -> Result<(), CliError> { + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::fs::PermissionsExt; + + let (parent, output_file_name) = open_real_parent_directory(path)?; + let output_name = unix_path_component(&output_file_name, "output file name")?; + if let Some(mode) = unix_file_mode_at(&parent, &output_name)? { + if mode & libc::S_IFMT != libc::S_IFREG { + return Err(CliError::Message(format!( + "output {} must be a regular file path, not a symbolic link or special file", + path.display() + ))); + } + if !overwrite { + return Err(CliError::Message(format!( + "output {} already exists; pass --overwrite to replace it", + path.display() + ))); + } + } + + let safe_file_name = safe_temp_file_component(&output_file_name); + let temp_file_name = OsString::from(format!( + ".{}.aether-restore-{}-{}.tmp", + safe_file_name.to_string_lossy(), + std::process::id(), + Uuid::new_v4() + )); + let temp_name = unix_path_component(&temp_file_name, "temporary output file name")?; + // SAFETY: `parent` is a valid directory descriptor, `temp_name` is NUL-terminated, and a + // successful descriptor is transferred immediately to `File`. O_EXCL prevents name reuse. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + temp_name.as_ptr(), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_CLOEXEC | libc::O_NOFOLLOW, + 0o600, + ) + }; + if descriptor < 0 { + return Err(CliError::Io(std::io::Error::last_os_error())); + } + // SAFETY: `openat` returned a new owned descriptor and no other owner exists. + let mut temp = unsafe { File::from_raw_fd(descriptor) }; + + let result = (|| -> Result<(), CliError> { + temp.set_permissions(fs::Permissions::from_mode(0o600))?; + temp.write_all(bytes)?; + temp.sync_all()?; + drop(temp); + if overwrite { + unix_rename_at(&parent, &temp_name, &output_name)?; + } else { + unix_link_at(&parent, &temp_name, &output_name).map_err(|error| { + if error.kind() == std::io::ErrorKind::AlreadyExists { + CliError::Message(format!( + "output {} already exists; pass --overwrite to replace it", + path.display() + )) + } else { + CliError::Io(error) + } + })?; + unix_unlink_at(&parent, &temp_name)?; + } + parent.sync_all()?; + Ok(()) + })(); + if result.is_err() { + let _ = unix_unlink_at(&parent, &temp_name); + } + result +} + +#[cfg(unix)] +fn unix_file_mode_at( + parent: &File, + file_name: &std::ffi::CStr, +) -> Result, CliError> { + use std::mem::MaybeUninit; + use std::os::fd::AsRawFd; + + let mut stat = MaybeUninit::::uninit(); + // SAFETY: `parent` and `file_name` remain valid for the call, and `stat` points to writable + // storage. The value is initialized only when fstatat reports success. + let result = unsafe { + libc::fstatat( + parent.as_raw_fd(), + file_name.as_ptr(), + stat.as_mut_ptr(), + libc::AT_SYMLINK_NOFOLLOW, + ) + }; + if result == 0 { + // SAFETY: successful fstatat initialized the complete stat structure. + return Ok(Some(unsafe { stat.assume_init() }.st_mode)); + } + let error = std::io::Error::last_os_error(); + if error.raw_os_error() == Some(libc::ENOENT) { + Ok(None) + } else { + Err(CliError::Io(error)) + } +} + +#[cfg(unix)] +fn unix_link_at( + parent: &File, + source: &std::ffi::CStr, + destination: &std::ffi::CStr, +) -> Result<(), std::io::Error> { + use std::os::fd::AsRawFd; + + // SAFETY: both names are valid NUL-terminated components and both directory descriptors are + // the same live `parent` descriptor. No pointers are retained after the call. + let result = unsafe { + libc::linkat( + parent.as_raw_fd(), + source.as_ptr(), + parent.as_raw_fd(), + destination.as_ptr(), + 0, + ) + }; + if result == 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error()) + } +} + +#[cfg(unix)] +fn unix_rename_at( + parent: &File, + source: &std::ffi::CStr, + destination: &std::ffi::CStr, +) -> Result<(), CliError> { + use std::os::fd::AsRawFd; + + // SAFETY: both names are valid NUL-terminated components and `parent` remains open for the + // duration of the call. renameat does not retain either pointer. + let result = unsafe { + libc::renameat( + parent.as_raw_fd(), + source.as_ptr(), + parent.as_raw_fd(), + destination.as_ptr(), + ) + }; + if result == 0 { + Ok(()) + } else { + Err(CliError::Io(std::io::Error::last_os_error())) + } +} + +#[cfg(unix)] +fn unix_unlink_at(parent: &File, file_name: &std::ffi::CStr) -> Result<(), CliError> { + use std::os::fd::AsRawFd; + + // SAFETY: `parent` is a valid directory descriptor and `file_name` is a NUL-terminated + // component. unlinkat does not retain either argument. + let result = unsafe { libc::unlinkat(parent.as_raw_fd(), file_name.as_ptr(), 0) }; + if result == 0 { + Ok(()) + } else { + Err(CliError::Io(std::io::Error::last_os_error())) + } +} + +#[cfg(not(unix))] +fn write_atomic_private(path: &Path, bytes: &[u8], overwrite: bool) -> Result<(), CliError> { + let parent = path + .parent() + .filter(|path| !path.as_os_str().is_empty()) + .unwrap_or(Path::new(".")); + let file_name = + safe_temp_file_component(path.file_name().ok_or_else(|| { + CliError::Message("output path must include a file name".to_string()) + })?); + if !parent.is_dir() { + return Err(CliError::Message(format!( + "output directory {} does not exist", + parent.display() + ))); + } + if fs::symlink_metadata(parent)?.file_type().is_symlink() { + return Err(CliError::Message(format!( + "output directory {} must not be a symbolic link", + parent.display() + ))); + } + if let Ok(metadata) = fs::symlink_metadata(path) { + if metadata.file_type().is_symlink() || metadata.is_dir() || !metadata.is_file() { + return Err(CliError::Message(format!( + "output {} must be a regular file path, not a symbolic link or special file", + path.display() + ))); + } + } + if !overwrite && fs::symlink_metadata(path).is_ok() { + return Err(CliError::Message(format!( + "output {} already exists; pass --overwrite to replace it", + path.display() + ))); + } + + let temp_name = format!( + ".{}.aether-restore-{}-{}.tmp", + file_name.to_string_lossy(), + std::process::id(), + Uuid::new_v4() + ); + let temp_path = parent.join(temp_name); + let mut options = OpenOptions::new(); + options.write(true).create_new(true); + let mut temp = options.open(&temp_path)?; + let result = (|| -> Result<(), CliError> { + temp.write_all(bytes)?; + temp.sync_all()?; + drop(temp); + if overwrite { + replace_output_file(&temp_path, path)?; + } else { + fs::hard_link(&temp_path, path).map_err(|error| { + if error.kind() == std::io::ErrorKind::AlreadyExists { + CliError::Message(format!( + "output {} already exists; pass --overwrite to replace it", + path.display() + )) + } else { + CliError::Io(error) + } + })?; + fs::remove_file(&temp_path)?; + } + Ok(()) + })(); + if result.is_err() { + let _ = fs::remove_file(&temp_path); + } + result +} + +fn safe_temp_file_component(file_name: &std::ffi::OsStr) -> OsString { + let sanitized: String = file_name + .to_string_lossy() + .chars() + .map(|character| match character { + '/' | '\\' | ':' | '\0' => '_', + character => character, + }) + .collect(); + OsString::from(sanitized) +} + +#[cfg(test)] +mod tests { + use super::{ + append_automatic_environment_keys, deduplicate_and_validate_key_values, read_limited_file, + write_atomic_private, MAX_LEGACY_V1_KEY_CANDIDATES, MAX_V2_KEY_CANDIDATES, + }; + + #[cfg(unix)] + fn unix_test_directory(prefix: &str) -> std::path::PathBuf { + std::fs::canonicalize(std::env::temp_dir()) + .expect("system temporary directory should canonicalize") + .join(format!("{prefix}-{}", uuid::Uuid::new_v4())) + } + + #[test] + fn automatic_key_environment_priority_matches_backup_generation() { + let mut requested = Vec::new(); + let mut values = Vec::new(); + append_automatic_environment_keys(&mut values, |env_name| { + requested.push(env_name.to_string()); + Some(format!("secret-for-{env_name}")) + }) + .expect("automatic keys should load"); + + assert_eq!( + requested, + vec![ + "AETHER_BACKUP_ENCRYPTION_KEY", + "AETHER_GATEWAY_DATA_ENCRYPTION_KEY", + "ENCRYPTION_KEY", + ] + ); + assert_eq!( + values, + vec![ + ("secret-for-AETHER_BACKUP_ENCRYPTION_KEY".to_string(), false,), + ( + "secret-for-AETHER_GATEWAY_DATA_ENCRYPTION_KEY".to_string(), + true, + ), + ("secret-for-ENCRYPTION_KEY".to_string(), true), + ] + ); + } + + #[test] + fn candidate_deduplication_preserves_priority_and_promotes_legacy_access() { + let candidates = deduplicate_and_validate_key_values(vec![ + ("backup-current".to_string(), false), + ("gateway-fallback".to_string(), true), + ("default-fallback".to_string(), true), + ("backup-current".to_string(), true), + ]) + .expect("candidate set should be valid"); + + assert_eq!( + candidates, + vec![ + ("backup-current".to_string(), true), + ("gateway-fallback".to_string(), true), + ("default-fallback".to_string(), true), + ] + ); + } + + #[test] + fn rejects_too_many_v2_candidates_before_key_construction() { + let values = (0..=MAX_V2_KEY_CANDIDATES) + .map(|index| (format!("v2-key-{index}"), false)) + .collect(); + + let error = deduplicate_and_validate_key_values(values) + .expect_err("candidate count above the restore limit must fail"); + + assert_eq!( + error.to_string(), + format!("backup restore allows at most {MAX_V2_KEY_CANDIDATES} v2 candidate keys") + ); + } + + #[test] + fn accepts_candidate_counts_at_both_restore_limits() { + let values = (0..MAX_LEGACY_V1_KEY_CANDIDATES) + .map(|index| (format!("legacy-key-{index}"), true)) + .chain( + (MAX_LEGACY_V1_KEY_CANDIDATES..MAX_V2_KEY_CANDIDATES) + .map(|index| (format!("v2-key-{index}"), false)), + ) + .collect(); + + let candidates = deduplicate_and_validate_key_values(values) + .expect("candidate counts at the restore limits must be accepted"); + + assert_eq!(candidates.len(), MAX_V2_KEY_CANDIDATES); + assert_eq!( + candidates.iter().filter(|(_, allow_v1)| *allow_v1).count(), + MAX_LEGACY_V1_KEY_CANDIDATES + ); + } + + #[test] + fn rejects_too_many_legacy_candidates_before_key_construction() { + let values = (0..=MAX_LEGACY_V1_KEY_CANDIDATES) + .map(|index| (format!("legacy-key-{index}"), true)) + .collect(); + + let error = deduplicate_and_validate_key_values(values) + .expect_err("legacy candidate count above the restore limit must fail"); + + assert_eq!( + error.to_string(), + format!( + "legacy v1 backup restore allows at most {MAX_LEGACY_V1_KEY_CANDIDATES} candidate keys" + ) + ); + } + + #[cfg(unix)] + #[test] + fn encrypted_input_symbolic_links_are_rejected() { + use std::os::unix::fs::symlink; + + let directory = unix_test_directory("aether-backup-restore-symlink-test"); + std::fs::create_dir(&directory).expect("test directory should be created"); + let target = directory.join("backup.bin"); + let link = directory.join("backup-link.bin"); + std::fs::write(&target, b"encrypted-backup").expect("test target should be written"); + symlink(&target, &link).expect("test symlink should be created"); + + let error = read_limited_file(&link, 1024, true, false) + .expect_err("encrypted backup symlink must be rejected"); + assert!(error.to_string().contains("must not be a symbolic link")); + + std::fs::remove_dir_all(directory).expect("test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn encrypted_input_symbolic_link_ancestors_are_rejected() { + use std::os::unix::fs::symlink; + + let directory = unix_test_directory("aether-backup-restore-ancestor-test"); + let real_parent = directory.join("real-parent"); + let linked_parent = directory.join("linked-parent"); + std::fs::create_dir_all(&real_parent).expect("real parent should be created"); + std::fs::write(real_parent.join("backup.bin"), b"encrypted-backup") + .expect("test input should be written"); + symlink(&real_parent, &linked_parent).expect("parent symlink should be created"); + + read_limited_file(&linked_parent.join("backup.bin"), 1024, true, false) + .expect_err("a symbolic-link ancestor must be rejected"); + + std::fs::remove_dir_all(directory).expect("test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn atomic_private_output_rejects_symbolic_link_ancestors() { + use std::os::unix::fs::symlink; + + let directory = unix_test_directory("aether-backup-output-ancestor-test"); + let real_parent = directory.join("real-parent"); + let linked_parent = directory.join("linked-parent"); + std::fs::create_dir_all(&real_parent).expect("real parent should be created"); + symlink(&real_parent, &linked_parent).expect("parent symlink should be created"); + + write_atomic_private(&linked_parent.join("restored.json"), b"{}", false) + .expect_err("output through a symbolic-link ancestor must be rejected"); + assert!(!real_parent.join("restored.json").exists()); + + std::fs::remove_dir_all(directory).expect("test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn atomic_private_output_is_mode_0600_and_preserves_no_overwrite() { + use std::os::unix::fs::PermissionsExt; + + let directory = unix_test_directory("aether-backup-output-mode-test"); + std::fs::create_dir(&directory).expect("test directory should be created"); + let output = directory.join("restored.json"); + + write_atomic_private(&output, b"first", false).expect("first output should be written"); + assert_eq!( + std::fs::metadata(&output) + .expect("output metadata should load") + .permissions() + .mode() + & 0o777, + 0o600 + ); + write_atomic_private(&output, b"second", false) + .expect_err("no-overwrite mode must preserve an existing output"); + assert_eq!( + std::fs::read(&output).expect("output should remain readable"), + b"first" + ); + + write_atomic_private(&output, b"second", true) + .expect("overwrite mode should atomically replace a regular output"); + assert_eq!( + std::fs::read(&output).expect("replaced output should be readable"), + b"second" + ); + + std::fs::remove_dir_all(directory).expect("test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn secret_files_require_current_owner_and_private_permissions() { + use std::os::unix::fs::PermissionsExt; + + let directory = unix_test_directory("aether-backup-secret-permissions-test"); + std::fs::create_dir(&directory).expect("test directory should be created"); + let secret_file = directory.join("backup.key"); + std::fs::write(&secret_file, b"private-backup-key").expect("test secret should be written"); + std::fs::set_permissions(&secret_file, std::fs::Permissions::from_mode(0o600)) + .expect("test secret permissions should be private"); + + assert_eq!( + read_limited_file(&secret_file, 1024, true, true) + .expect("current-user-owned private secret should load"), + b"private-backup-key" + ); + + std::fs::set_permissions(&secret_file, std::fs::Permissions::from_mode(0o640)) + .expect("test secret permissions should become group-readable"); + assert!(read_limited_file(&secret_file, 1024, true, true) + .expect_err("group-readable secret must be rejected") + .to_string() + .contains("accessible by group or other users")); + + std::fs::remove_dir_all(directory).expect("test directory should be removed"); + } +} + +#[cfg(windows)] +fn replace_output_file(temp_path: &Path, path: &Path) -> Result<(), CliError> { + if fs::symlink_metadata(path).is_ok() { + return Err(CliError::Message( + "--overwrite cannot atomically replace an existing file on Windows; choose a new output path" + .to_string(), + )); + } + fs::rename(temp_path, path).map_err(CliError::from) +} + +#[cfg(not(any(unix, windows)))] +fn replace_output_file(temp_path: &Path, path: &Path) -> Result<(), CliError> { + if fs::symlink_metadata(path).is_ok() { + return Err(CliError::Message( + "--overwrite is unsupported on this platform; choose a new output path".to_string(), + )); + } + fs::rename(temp_path, path).map_err(CliError::from) +} diff --git a/apps/aether-gateway/src/bin/support/responses_ws_probe.rs b/apps/aether-gateway/src/bin/support/responses_ws_probe.rs index 061a97cf6..78231fac7 100644 --- a/apps/aether-gateway/src/bin/support/responses_ws_probe.rs +++ b/apps/aether-gateway/src/bin/support/responses_ws_probe.rs @@ -193,6 +193,7 @@ pub(crate) fn parse_probe_url(raw: &str) -> Result { || url.password().is_some() || url.query().is_some() || url.fragment().is_some() + || (url.scheme() == "ws" && !aether_http::url_has_literal_loopback_host(&url)) { return Err(ProbeFailure::InvalidEndpoint); } @@ -210,6 +211,7 @@ pub(crate) const fn turn_timeout(args: &ProbeArgs) -> Duration { async fn run_probe(config: &ProbeConfig, started_at: Instant) -> Result { let client = wreq::Client::builder() + .no_proxy() .connect_timeout(config.turn_timeout) .timeout(config.turn_timeout) .build() @@ -434,7 +436,14 @@ mod tests { #[test] fn probe_url_rejects_credentials_and_query_strings() { assert!(parse_probe_url("wss://example.test/v1/responses").is_ok()); + assert!(parse_probe_url("ws://localhost:8080/v1/responses").is_ok()); + assert!(parse_probe_url("ws://127.42.0.1:8080/v1/responses").is_ok()); + assert!(parse_probe_url("ws://[::1]:8080/v1/responses").is_ok()); assert!(parse_probe_url("https://example.test/v1/responses").is_err()); + assert!(parse_probe_url("ws://example.test/v1/responses").is_err()); + assert!(parse_probe_url("ws://10.0.0.1/v1/responses").is_err()); + assert!(parse_probe_url("ws://0.0.0.0:8080/v1/responses").is_err()); + assert!(parse_probe_url("ws://[::ffff:127.0.0.1]:8080/v1/responses").is_err()); assert!(parse_probe_url("wss://token@example.test/v1/responses").is_err()); assert!(parse_probe_url("wss://example.test/v1/responses?token=secret").is_err()); } diff --git a/apps/aether-gateway/src/cache/auth_context.rs b/apps/aether-gateway/src/cache/auth_context.rs index cc2c3e9dd..c626247bd 100644 --- a/apps/aether-gateway/src/cache/auth_context.rs +++ b/apps/aether-gateway/src/cache/auth_context.rs @@ -422,6 +422,7 @@ mod tests { local_rejection: None, allowed_models: None, ip_rules: None, + verified_api_key_hash: None, } } diff --git a/apps/aether-gateway/src/cache/system_config.rs b/apps/aether-gateway/src/cache/system_config.rs index bef43c4fe..5aa315b58 100644 --- a/apps/aether-gateway/src/cache/system_config.rs +++ b/apps/aether-gateway/src/cache/system_config.rs @@ -173,6 +173,15 @@ impl SystemConfigCache { self.detach_all_loads(); } + pub(crate) fn invalidate(&self, key: &str) { + let Ok(_mutation) = self.mutation.lock() else { + return; + }; + self.generation.fetch_add(1, Ordering::AcqRel); + self.entries.remove(&key.to_string()); + self.detach_all_loads(); + } + pub(crate) fn insert_if_generation( &self, key: String, diff --git a/apps/aether-gateway/src/constants.rs b/apps/aether-gateway/src/constants.rs index a99de0675..e1edeb5d9 100644 --- a/apps/aether-gateway/src/constants.rs +++ b/apps/aether-gateway/src/constants.rs @@ -18,6 +18,7 @@ pub(crate) const TUNNEL_AFFINITY_FORWARDED_BY_HEADER: &str = "x-aether-tunnel-affinity-forwarded-by"; pub(crate) const TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER: &str = "x-aether-tunnel-affinity-owner-instance-id"; +pub(crate) const TUNNEL_AFFINITY_NODE_ID_HEADER: &str = "x-aether-tunnel-affinity-node-id"; pub(crate) const EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH: &str = "public_proxy_passthrough"; pub(crate) const EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED: &str = "local_proxy_passthrough_removed"; diff --git a/apps/aether-gateway/src/control/auth/credentials.rs b/apps/aether-gateway/src/control/auth/credentials.rs index f9c0a1825..ae4953d2b 100644 --- a/apps/aether-gateway/src/control/auth/credentials.rs +++ b/apps/aether-gateway/src/control/auth/credentials.rs @@ -23,7 +23,7 @@ pub(crate) fn extract_requested_model( body: &Bytes, ) -> Option { if decision.route_family.as_deref() == Some("gemini") { - if let Some(model) = extract_gemini_model_from_path(uri.path()) { + if let Some(model) = extract_gemini_requested_model_from_path(uri.path()) { return Some(model); } } @@ -43,23 +43,46 @@ pub(crate) fn extract_requested_model( .filter(|value| !value.is_empty()) } +fn extract_gemini_requested_model_from_path(path: &str) -> Option { + let model = extract_gemini_model_from_path(path)?; + Some( + model + .split_once("/operations/") + .map(|(model, _)| model) + .unwrap_or(model.as_str()) + .to_string(), + ) +} + pub(super) fn extract_request_credentials( headers: &http::HeaderMap, uri: &Uri, auth_endpoint_signature: &str, +) -> GatewayExtractedCredentials { + extract_request_credentials_with_trusted_auth(headers, uri, auth_endpoint_signature, cfg!(test)) +} + +pub(super) fn extract_request_credentials_with_trusted_auth( + headers: &http::HeaderMap, + uri: &Uri, + auth_endpoint_signature: &str, + trusted_auth_verified: bool, ) -> GatewayExtractedCredentials { let bundle = GatewayCredentialBundle { - authorization_bearer: header_value_str(headers, http::header::AUTHORIZATION.as_str()) - .as_deref() - .and_then(extract_bearer_token) - .map(ToOwned::to_owned), + authorization_bearer: unique_header_value_str( + headers, + http::header::AUTHORIZATION.as_str(), + ) + .as_deref() + .and_then(extract_bearer_token) + .map(ToOwned::to_owned), x_api_key: header_value_str(headers, "x-api-key"), api_key: header_value_str(headers, "api-key"), x_goog_api_key: header_value_str(headers, "x-goog-api-key"), query_key: extract_query_api_key(uri), cookie_header: header_value_str(headers, http::header::COOKIE.as_str()), }; - let trusted_headers = extract_trusted_auth_headers(headers); + let trusted_headers = extract_trusted_auth_headers(headers, trusted_auth_verified); let trusted_admin_headers = extract_trusted_admin_headers(headers); let primary = select_primary_credential(auth_endpoint_signature, &bundle); @@ -71,6 +94,20 @@ pub(super) fn extract_request_credentials( } } +fn unique_header_value_str(headers: &http::HeaderMap, key: &str) -> Option { + let mut values = headers.get_all(key).iter(); + let value = values.next()?; + if values.next().is_some() { + return None; + } + value + .to_str() + .ok() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + pub(in crate::control) fn resolve_gateway_credential_carrier( headers: &http::HeaderMap, uri: &Uri, @@ -85,6 +122,7 @@ pub(in crate::control) fn resolve_gateway_credential_carrier( }) } +#[cfg(test)] fn has_trusted_gateway_marker(headers: &http::HeaderMap) -> bool { header_value_str(headers, crate::constants::GATEWAY_HEADER) .unwrap_or_default() @@ -97,13 +135,32 @@ pub(super) fn build_auth_context_cache_key( headers: &http::HeaderMap, uri: &Uri, auth_endpoint_signature: &str, +) -> Option { + build_auth_context_cache_key_with_trusted_auth( + headers, + uri, + auth_endpoint_signature, + cfg!(test), + ) +} + +pub(super) fn build_auth_context_cache_key_with_trusted_auth( + headers: &http::HeaderMap, + uri: &Uri, + auth_endpoint_signature: &str, + trusted_auth_verified: bool, ) -> Option { let signature = auth_endpoint_signature.trim(); if signature.is_empty() { return None; } - let extracted = extract_request_credentials(headers, uri, signature); + let extracted = extract_request_credentials_with_trusted_auth( + headers, + uri, + signature, + trusted_auth_verified, + ); let trusted_headers = extracted.trusted_headers; let bundle = extracted.bundle; if bundle.authorization_bearer.is_none() @@ -135,7 +192,7 @@ pub(super) fn build_auth_context_cache_key( }) .unwrap_or_default(); - Some(format!( + let raw_cache_identity = format!( "{signature}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}", bundle.authorization_bearer.unwrap_or_default(), bundle.x_api_key.unwrap_or_default(), @@ -147,11 +204,26 @@ pub(super) fn build_auth_context_cache_key( trusted_api_key_id, trusted_balance_remaining, trusted_access_allowed, - )) + ); + let mut hasher = Sha256::new(); + hasher.update(raw_cache_identity.as_bytes()); + Some(format!("auth-context:sha256:{:x}", hasher.finalize())) } -fn extract_trusted_auth_headers(headers: &http::HeaderMap) -> Option { - if !has_trusted_gateway_marker(headers) { +fn extract_trusted_auth_headers( + headers: &http::HeaderMap, + trusted_auth_verified: bool, +) -> Option { + if !trusted_auth_verified { + return None; + } + #[cfg(test)] + if !header_value_str(headers, crate::constants::GATEWAY_HEADER) + .unwrap_or_default() + .trim() + .to_ascii_lowercase() + .starts_with("rust-phase3") + { return None; } let user_id = header_value_str(headers, crate::constants::TRUSTED_AUTH_USER_ID_HEADER) @@ -387,7 +459,7 @@ fn extract_bearer_token(value: &str) -> Option<&str> { return None; } let token = token.trim(); - if token.is_empty() { + if token.is_empty() || token.chars().any(char::is_whitespace) { None } else { Some(token) @@ -472,6 +544,46 @@ mod tests { assert_eq!(requested_model.as_deref(), Some("gpt-5.4")); } + #[test] + fn extract_requested_model_handles_gemini_generation_and_operation_paths() { + let generation_decision = GatewayControlDecision::synthetic( + "/v1beta/models/gemini-2.5-pro:generateContent", + Some("ai_public".to_string()), + Some("gemini".to_string()), + Some("generate_content".to_string()), + Some("gemini:generate_content".to_string()), + ); + let operation_decision = GatewayControlDecision::synthetic( + "/v1beta/models/veo-3/operations/task-123:cancel", + Some("ai_public".to_string()), + Some("gemini".to_string()), + Some("video".to_string()), + Some("gemini:video".to_string()), + ); + let headers = http::HeaderMap::new(); + + assert_eq!( + extract_requested_model( + &generation_decision, + &uri("/v1beta/models/gemini-2.5-pro:generateContent"), + &headers, + &Bytes::new(), + ) + .as_deref(), + Some("gemini-2.5-pro") + ); + assert_eq!( + extract_requested_model( + &operation_decision, + &uri("/v1beta/models/veo-3/operations/task-123:cancel"), + &headers, + &Bytes::new(), + ) + .as_deref(), + Some("veo-3") + ); + } + #[test] fn selects_openai_bearer_as_provider_api_key() { let mut headers = http::HeaderMap::new(); @@ -491,6 +603,33 @@ mod tests { ); } + #[test] + fn rejects_duplicate_or_combined_authorization_credentials() { + let mut duplicate = http::HeaderMap::new(); + duplicate.append( + http::header::AUTHORIZATION, + "Bearer first-token".parse().unwrap(), + ); + duplicate.append( + http::header::AUTHORIZATION, + "Bearer second-token".parse().unwrap(), + ); + let extracted = + extract_request_credentials(&duplicate, &uri("/api/admin/system"), "admin:operational"); + assert!(extracted.bundle.authorization_bearer.is_none()); + assert!(extracted.primary.is_none()); + + let mut combined = http::HeaderMap::new(); + combined.insert( + http::header::AUTHORIZATION, + "Bearer first-token, Bearer second-token".parse().unwrap(), + ); + let extracted = + extract_request_credentials(&combined, &uri("/api/admin/system"), "admin:operational"); + assert!(extracted.bundle.authorization_bearer.is_none()); + assert!(extracted.primary.is_none()); + } + #[test] fn selects_codex_live_bearer_as_provider_api_key() { let mut headers = http::HeaderMap::new(); @@ -608,7 +747,7 @@ mod tests { } #[test] - fn cache_key_includes_cookie_header() { + fn cache_key_hashes_cookie_header_instead_of_retaining_session_secret() { let mut headers = http::HeaderMap::new(); headers.insert(http::header::COOKIE, "session=abc123".parse().unwrap()); @@ -618,7 +757,8 @@ mod tests { "internal:session", ) .expect("cache key should exist"); - assert!(cache_key.contains("session=abc123")); + assert!(cache_key.starts_with("auth-context:sha256:")); + assert!(!cache_key.contains("session=abc123")); } #[test] @@ -669,12 +809,10 @@ mod tests { .expect("trusted cache key should exist"); assert_ne!(first, second); - assert!(first.contains("user-1")); - assert!(first.contains("key-1")); - assert!(first.contains("1.5")); - assert!(first.contains("true")); - assert!(second.contains("user-2")); - assert!(second.contains("false")); + for raw_identity in ["user-1", "key-1", "1.5", "user-2"] { + assert!(!first.contains(raw_identity)); + assert!(!second.contains(raw_identity)); + } } #[test] diff --git a/apps/aether-gateway/src/control/auth/gate.rs b/apps/aether-gateway/src/control/auth/gate.rs index 0b5ce4d33..58f0c8d72 100644 --- a/apps/aether-gateway/src/control/auth/gate.rs +++ b/apps/aether-gateway/src/control/auth/gate.rs @@ -224,7 +224,7 @@ fn wallet_finite_available_usd( Some(wallet.balance.max(0.0) + wallet.gift_balance.max(0.0)) } -async fn estimate_execution_plan_cost_upper_bound_usd( +pub(crate) async fn estimate_execution_plan_cost_upper_bound_usd( state: &AppState, plan: &aether_contracts::ExecutionPlan, report_context: Option<&serde_json::Value>, @@ -925,6 +925,7 @@ mod tests { local_rejection: None, allowed_models: Some(allowed_models), ip_rules: None, + verified_api_key_hash: None, }); decision } diff --git a/apps/aether-gateway/src/control/auth/mod.rs b/apps/aether-gateway/src/control/auth/mod.rs index 591e96ce0..9de07e7bf 100644 --- a/apps/aether-gateway/src/control/auth/mod.rs +++ b/apps/aether-gateway/src/control/auth/mod.rs @@ -7,13 +7,16 @@ mod types; pub(crate) use credentials::extract_requested_model; pub(super) use credentials::resolve_gateway_credential_carrier; pub(crate) use gate::{ - execution_plan_balance_capacity_rejection, request_model_local_rejection, - should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayLocalAuthRejection, + estimate_execution_plan_cost_upper_bound_usd, execution_plan_balance_capacity_rejection, + request_model_local_rejection, should_buffer_request_for_local_auth, + trusted_auth_local_rejection, GatewayLocalAuthRejection, }; pub(crate) use resolution::{ refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot, - resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext, - GatewayControlAuthContext, + resolve_execution_runtime_auth_context, resolve_local_admin_session_principal, + GatewayAdminPrincipalContext, GatewayControlAuthContext, +}; +pub(super) use resolution::{ + resolve_control_decision_auth_with_trusted_auth, ControlDecisionAuthResolution, }; -pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution}; pub(crate) use types::GatewayCredentialCarrier; diff --git a/apps/aether-gateway/src/control/auth/resolution.rs b/apps/aether-gateway/src/control/auth/resolution.rs index bd572db6e..c16ac7a8a 100644 --- a/apps/aether-gateway/src/control/auth/resolution.rs +++ b/apps/aether-gateway/src/control/auth/resolution.rs @@ -4,11 +4,9 @@ use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogProvider, }; use axum::http::Uri; -use base64::Engine as _; -use hmac::Mac; use serde::{Deserialize, Serialize}; use serde_json::Value; -use tracing::{debug, info}; +use tracing::{debug, info, warn}; use crate::wallet_runtime::{ local_rejection_from_wallet_access, resolve_wallet_auth_gate_uncached, @@ -17,7 +15,8 @@ use crate::{AppState, GatewayError}; use super::super::GatewayControlDecision; use super::credentials::{ - build_auth_context_cache_key, current_unix_secs, extract_request_credentials, + build_auth_context_cache_key, build_auth_context_cache_key_with_trusted_auth, + current_unix_secs, extract_request_credentials, extract_request_credentials_with_trusted_auth, extract_trusted_admin_headers, hash_api_key, }; use super::gate::GatewayLocalAuthRejection; @@ -27,6 +26,9 @@ use super::types::{ }; use crate::cache::{AuthContextCacheGeneration, AuthContextInflightRegistration}; use crate::headers::header_value_str; +use crate::local_auth_token::{ + decode_local_auth_token, local_auth_token_identity_matches_user, LocalAuthTokenType, +}; const AUTH_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(60); const AUTH_CONTEXT_CACHE_REFRESH_INTERVAL: Duration = Duration::from_secs(10); @@ -93,6 +95,30 @@ pub(crate) struct GatewayControlAuthContext { pub(crate) allowed_models: Option>, #[serde(skip)] pub(crate) ip_rules: Option>, + /// Credential verifier that established this API-key identity. Long-lived + /// executions use it to prove that a later row with the same IDs is still + /// the record authenticated by the original request. + #[serde(skip)] + pub(crate) verified_api_key_hash: Option, +} + +#[derive(Clone)] +pub(crate) struct VerifiedApiKeyHash(String); + +impl VerifiedApiKeyHash { + fn new(value: String) -> Self { + Self(value) + } + + fn as_str(&self) -> &str { + self.0.as_str() + } +} + +impl std::fmt::Debug for VerifiedApiKeyHash { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("VerifiedApiKeyHash([REDACTED])") + } } #[derive(Debug, Clone, PartialEq, Eq)] @@ -109,11 +135,30 @@ pub(in super::super) enum ControlDecisionAuthResolution { } pub(in super::super) async fn resolve_control_decision_auth( + state: &AppState, + headers: &http::HeaderMap, + uri: &Uri, + trace_id: &str, + decision: GatewayControlDecision, +) -> Result { + resolve_control_decision_auth_with_trusted_auth( + state, + headers, + uri, + trace_id, + decision, + cfg!(test), + ) + .await +} + +pub(in super::super) async fn resolve_control_decision_auth_with_trusted_auth( state: &AppState, headers: &http::HeaderMap, uri: &Uri, trace_id: &str, mut decision: GatewayControlDecision, + trusted_auth_verified: bool, ) -> Result { if let Some(admin_principal) = resolve_trusted_admin_principal(headers, decision.auth_endpoint_signature.as_deref()) @@ -132,10 +177,18 @@ pub(in super::super) async fn resolve_control_decision_auth( decision.admin_principal = Some(admin_principal); } - let auth_context_cache_key = decision - .auth_endpoint_signature - .as_deref() - .and_then(|signature| build_auth_context_cache_key(headers, uri, signature)); + let auth_context_cache_key = + decision + .auth_endpoint_signature + .as_deref() + .and_then(|signature| { + build_auth_context_cache_key_with_trusted_auth( + headers, + uri, + signature, + trusted_auth_verified, + ) + }); let mut resolved_auth_context = None; if let Some(cache_key) = auth_context_cache_key.as_deref() { @@ -149,6 +202,7 @@ pub(in super::super) async fn resolve_control_decision_auth( decision.auth_endpoint_signature.as_deref(), headers, uri, + trusted_auth_verified, ) .await?, ); @@ -168,6 +222,7 @@ pub(in super::super) async fn resolve_control_decision_auth( uri, decision.auth_endpoint_signature.as_deref(), true, + trusted_auth_verified, ) .await?; } @@ -336,7 +391,7 @@ async fn resolve_local_admin_principal( let Some(access_token) = extracted.bundle.authorization_bearer.as_deref() else { return Ok(None); }; - let claims = match decode_local_auth_token(access_token, "access") { + let claims = match decode_local_auth_token(access_token, LocalAuthTokenType::Access) { Ok(claims) => claims, Err(_) => return Ok(None), }; @@ -351,6 +406,14 @@ async fn resolve_local_admin_principal( resolve_local_admin_principal_from_claims(state, headers, uri, &claims).await } +pub(crate) async fn resolve_local_admin_session_principal( + state: &AppState, + headers: &http::HeaderMap, + uri: &Uri, +) -> Result, GatewayError> { + resolve_local_admin_principal(state, headers, uri, Some("admin:operational")).await +} + async fn resolve_local_admin_principal_from_claims( state: &AppState, headers: &http::HeaderMap, @@ -373,6 +436,9 @@ async fn resolve_local_admin_principal_from_claims( if !user.is_active || user.is_deleted || !crate::roles::can_access_admin_console(&user.role) { return Ok(None); } + if !local_auth_token_identity_matches_user(claims, &user) { + return Ok(None); + } let now = chrono::Utc::now(); let Some(session) = state.find_user_session(user_id, session_id).await? else { @@ -380,6 +446,7 @@ async fn resolve_local_admin_principal_from_claims( }; if session.is_revoked() || session.is_expired(now) + || session.security_version != user.security_version || session.client_device_id != client_device_id { return Ok(None); @@ -431,68 +498,6 @@ fn local_admin_user_agent(headers: &http::HeaderMap) -> Option { .map(|value| value.chars().take(1000).collect()) } -fn local_auth_secret() -> String { - std::env::var("JWT_SECRET_KEY") - .ok() - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) - .unwrap_or_else(|| "aether-rust-dev-jwt-secret".to_string()) -} - -fn decode_local_auth_token( - token: &str, - expected_type: &str, -) -> Result, String> { - let mut parts = token.split('.'); - let Some(header_segment) = parts.next() else { - return Err("invalid token".to_string()); - }; - let Some(payload_segment) = parts.next() else { - return Err("invalid token".to_string()); - }; - let Some(signature_segment) = parts.next() else { - return Err("invalid token".to_string()); - }; - if parts.next().is_some() { - return Err("invalid token".to_string()); - } - - let signing_input = format!("{header_segment}.{payload_segment}"); - let signature = base64::engine::general_purpose::URL_SAFE_NO_PAD - .decode(signature_segment) - .map_err(|_| "invalid token".to_string())?; - let mut mac = hmac::Hmac::::new_from_slice(local_auth_secret().as_bytes()) - .map_err(|_| "invalid token".to_string())?; - mac.update(signing_input.as_bytes()); - mac.verify_slice(&signature) - .map_err(|_| "invalid token".to_string())?; - - let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD - .decode(payload_segment) - .map_err(|_| "invalid token".to_string())?; - let payload = - serde_json::from_slice::(&payload_bytes).map_err(|_| "invalid token".to_string())?; - let payload = payload - .as_object() - .cloned() - .ok_or_else(|| "invalid token".to_string())?; - let actual_type = payload - .get("type") - .and_then(Value::as_str) - .unwrap_or_default(); - if actual_type != expected_type { - return Err("invalid token".to_string()); - } - let exp = payload - .get("exp") - .and_then(Value::as_i64) - .ok_or_else(|| "invalid token".to_string())?; - if exp <= chrono::Utc::now().timestamp() { - return Err("expired token".to_string()); - } - Ok(payload) -} - pub(crate) async fn resolve_execution_runtime_auth_context( state: &AppState, decision: &GatewayControlDecision, @@ -525,6 +530,7 @@ pub(crate) async fn resolve_execution_runtime_auth_context( Some(auth_endpoint_signature), headers, uri, + cfg!(test), ) .await .map(Some); @@ -539,6 +545,7 @@ pub(crate) async fn resolve_execution_runtime_auth_context( uri, Some(auth_endpoint_signature), true, + cfg!(test), ) .await? { @@ -558,6 +565,7 @@ async fn revalidate_cached_auth_context( auth_endpoint_signature: Option<&str>, headers: &http::HeaderMap, uri: &Uri, + trusted_auth_verified: bool, ) -> Result { if is_negative_auth_context(&auth_context) || !auth_context.access_allowed @@ -581,6 +589,7 @@ async fn revalidate_cached_auth_context( uri, auth_context.clone(), auth_endpoint_signature, + trusted_auth_verified, ) .await { @@ -622,6 +631,7 @@ async fn revalidate_cached_auth_context( uri, auth_context, auth_endpoint_signature, + trusted_auth_verified, ) .await; if refreshed.is_err() { @@ -639,9 +649,16 @@ async fn resolve_security_fresh_auth_context( uri: &Uri, stale: GatewayControlAuthContext, auth_endpoint_signature: Option<&str>, + trusted_auth_verified: bool, ) -> Result { - if let Some(refreshed) = - resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature).await? + if let Some(refreshed) = resolve_data_backed_auth_context_with_trusted_auth( + state, + headers, + uri, + auth_endpoint_signature, + trusted_auth_verified, + ) + .await? { return Ok(refreshed); } @@ -660,19 +677,27 @@ async fn resolve_data_backed_auth_context_cached( uri: &Uri, auth_endpoint_signature: Option<&str>, cache_negative: bool, + trusted_auth_verified: bool, ) -> Result, GatewayError> { let Some(cache_key) = cache_key else { - return resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature) - .await; + return resolve_data_backed_auth_context_with_trusted_auth( + state, + headers, + uri, + auth_endpoint_signature, + trusted_auth_verified, + ) + .await; }; loop { match state.auth_context_cache.register_inflight(cache_key) { AuthContextInflightRegistration::Leader(guard) => { - let resolved = match resolve_data_backed_auth_context( + let resolved = match resolve_data_backed_auth_context_with_trusted_auth( state, headers, uri, auth_endpoint_signature, + trusted_auth_verified, ) .await { @@ -708,11 +733,12 @@ async fn resolve_data_backed_auth_context_cached( } } AuthContextInflightRegistration::Bypass => { - return resolve_data_backed_auth_context( + return resolve_data_backed_auth_context_with_trusted_auth( state, headers, uri, auth_endpoint_signature, + trusted_auth_verified, ) .await; } @@ -768,28 +794,39 @@ pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot( return Ok((auth_context, None)); } + let verified_api_key_hash = auth_context.verified_api_key_hash.clone(); let snapshot = { let _permit = state.acquire_auth_snapshot_load_gate().await?; - state - .data - .read_auth_api_key_snapshot_strong( - &auth_context.user_id, - &auth_context.api_key_id, - current_unix_secs(), - ) - .await - .map_err(|err| GatewayError::Internal(err.to_string()))? + if let Some(key_hash) = verified_api_key_hash.as_ref() { + state + .data + .read_auth_api_key_snapshot_by_key_hash_strong( + key_hash.as_str(), + current_unix_secs(), + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + } else { + state + .data + .read_auth_api_key_snapshot_strong( + &auth_context.user_id, + &auth_context.api_key_id, + current_unix_secs(), + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + } }; let Some(snapshot) = snapshot else { - let mut denied = auth_context; - denied.access_allowed = false; - denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey); - denied.balance_remaining = None; - return Ok((denied, None)); + return Ok((deny_refreshed_auth_context(auth_context), None)); + }; + if snapshot.user_id != auth_context.user_id || snapshot.api_key_id != auth_context.api_key_id { + return Ok((deny_refreshed_auth_context(auth_context), None)); }; let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?; - let refreshed = build_data_backed_auth_context( + let mut refreshed = build_data_backed_auth_context( state, snapshot.clone(), auth_endpoint_signature, @@ -798,9 +835,19 @@ pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot( wallet_access, ) .await; + refreshed.verified_api_key_hash = verified_api_key_hash; Ok((refreshed, Some(snapshot))) } +fn deny_refreshed_auth_context( + mut auth_context: GatewayControlAuthContext, +) -> GatewayControlAuthContext { + auth_context.access_allowed = false; + auth_context.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey); + auth_context.balance_remaining = None; + auth_context +} + fn put_cached_auth_context( state: &AppState, cache_key: String, @@ -913,6 +960,23 @@ pub(super) async fn resolve_data_backed_auth_context( headers: &http::HeaderMap, uri: &Uri, auth_endpoint_signature: Option<&str>, +) -> Result, GatewayError> { + resolve_data_backed_auth_context_with_trusted_auth( + state, + headers, + uri, + auth_endpoint_signature, + cfg!(test), + ) + .await +} + +async fn resolve_data_backed_auth_context_with_trusted_auth( + state: &AppState, + headers: &http::HeaderMap, + uri: &Uri, + auth_endpoint_signature: Option<&str>, + trusted_auth_verified: bool, ) -> Result, GatewayError> { let Some(signature) = auth_endpoint_signature .map(str::trim) @@ -923,7 +987,12 @@ pub(super) async fn resolve_data_backed_auth_context( if !state.has_auth_api_key_reader() { return Ok(None); } - let extracted = extract_request_credentials(headers, uri, signature); + let extracted = extract_request_credentials_with_trusted_auth( + headers, + uri, + signature, + trusted_auth_verified, + ); let principal = derive_principal_candidate(&extracted); let now_unix_secs = current_unix_secs(); @@ -955,6 +1024,7 @@ pub(super) async fn resolve_data_backed_auth_context( local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey), allowed_models: None, ip_rules: None, + verified_api_key_hash: None, })); }; @@ -963,17 +1033,17 @@ pub(super) async fn resolve_data_backed_auth_context( .await; let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?; - Ok(Some( - build_data_backed_auth_context( - state, - snapshot, - signature, - None, - None, - wallet_access, - ) - .await, - )) + let mut auth_context = build_data_backed_auth_context( + state, + snapshot, + signature, + None, + None, + wallet_access, + ) + .await; + auth_context.verified_api_key_hash = Some(VerifiedApiKeyHash::new(key_hash)); + Ok(Some(auth_context)) } Some(GatewayPrincipalCandidate::DeferredBearerToken { raw, carrier }) => { if let Some(auth_context) = resolve_antigravity_bearer_bridge_auth_context( @@ -1068,6 +1138,7 @@ async fn resolve_antigravity_bearer_bridge_auth_context( local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey), allowed_models: None, ip_rules: None, + verified_api_key_hash: None, })); }; @@ -1127,6 +1198,7 @@ async fn resolve_trusted_auth_context( local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey), allowed_models: None, ip_rules: None, + verified_api_key_hash: None, })); }; @@ -1158,9 +1230,7 @@ async fn build_data_backed_auth_context( let invalid_api_key = !snapshot.user_is_active || snapshot.user_is_deleted || !snapshot.api_key_is_active - || snapshot - .api_key_expires_at_unix_secs - .is_some_and(|expires_at| expires_at < current_unix_secs()); + || api_key_is_expired(snapshot.api_key_expires_at_unix_secs, current_unix_secs()); let locked_api_key = snapshot.api_key_is_locked && !snapshot.api_key_is_standalone; let key_access_allowed = header_access_allowed .map(|value| value && snapshot.currently_usable) @@ -1225,9 +1295,14 @@ async fn build_data_backed_auth_context( local_rejection, allowed_models, ip_rules: snapshot.api_key_ip_rules, + verified_api_key_hash: None, } } +fn api_key_is_expired(expires_at_unix_secs: Option, now_unix_secs: u64) -> bool { + expires_at_unix_secs.is_some_and(|expires_at| expires_at <= now_unix_secs) +} + fn contains_api_format_or_alias(items: &[String], target: &str) -> bool { items.iter().any(|item| api_format_matches(item, target)) } @@ -1282,18 +1357,21 @@ async fn auth_snapshot_allows_requested_provider( return true; } if !state.has_provider_catalog_data_reader() { - return true; + debug!( + "deny requested provider {}: provider catalog is unavailable for allowlist resolution", + requested_provider + ); + return false; } let providers = match state.list_provider_catalog_providers(true).await { Ok(value) => value, Err(err) => { - debug!( - "skip local provider auth gate for requested provider {}: provider catalog lookup failed: {:?}", - requested_provider, - err + warn!( + "deny requested provider {}: provider catalog lookup failed: {:?}", + requested_provider, err ); - return true; + return false; } }; @@ -1331,11 +1409,11 @@ async fn auth_snapshot_allows_requested_provider( { Ok(value) => value, Err(err) => { - debug!( - "skip local provider auth gate for requested provider {}: provider endpoint lookup failed: {:?}", + warn!( + "deny requested provider {}: provider endpoint lookup failed: {:?}", requested_provider, err ); - return true; + return false; } }; @@ -1426,7 +1504,8 @@ mod tests { use std::time::Duration; use aether_data::repository::auth::{ - AuthApiKeyWriteRepository, InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, + AuthApiKeyWriteRepository, CreateUserApiKeyRecord, InMemoryAuthApiKeySnapshotRepository, + StoredAuthApiKeySnapshot, }; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::wallet::{ @@ -1441,9 +1520,10 @@ mod tests { use futures_util::future::join_all; use super::{ - get_cached_auth_context, resolve_control_decision_auth, resolve_data_backed_auth_context, - resolve_execution_runtime_auth_context, ControlDecisionAuthResolution, - GatewayLocalAuthRejection, + api_key_is_expired, get_cached_auth_context, + refresh_execution_runtime_auth_context_with_snapshot, resolve_control_decision_auth, + resolve_data_backed_auth_context, resolve_execution_runtime_auth_context, + ControlDecisionAuthResolution, GatewayLocalAuthRejection, }; use crate::control::auth::credentials::{build_auth_context_cache_key, hash_api_key}; use crate::control::GatewayControlDecision; @@ -1481,6 +1561,14 @@ mod tests { path.parse().expect("uri should parse") } + #[test] + fn api_key_expiry_is_inclusive_at_the_declared_second() { + assert!(!api_key_is_expired(None, 100)); + assert!(!api_key_is_expired(Some(101), 100)); + assert!(api_key_is_expired(Some(100), 100)); + assert!(api_key_is_expired(Some(99), 100)); + } + fn sample_provider(id: &str, name: &str, provider_type: &str) -> StoredProviderCatalogProvider { StoredProviderCatalogProvider::new( id.to_string(), @@ -1769,6 +1857,97 @@ mod tests { assert_eq!(repository.touch_count("key-1"), 1); } + #[tokio::test] + async fn long_lived_refresh_rejects_same_ids_recreated_with_a_different_credential() { + let old_api_key = "sk-old-websocket-credential"; + let new_api_key = "sk-new-websocket-credential"; + let old_key_hash = hash_api_key(old_api_key); + let new_key_hash = hash_api_key(new_api_key); + let mut old_snapshot = sample_snapshot("key-stable-id", "user-stable-id"); + old_snapshot.user_allowed_api_formats = Some(vec!["openai:responses".to_string()]); + old_snapshot.api_key_allowed_api_formats = Some(vec!["openai:responses".to_string()]); + let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(old_key_hash.clone()), + old_snapshot, + )])); + let data = GatewayDataState::with_auth_api_key_repository_for_tests(repository.clone()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data); + let mut headers = HeaderMap::new(); + headers.insert( + http::header::AUTHORIZATION, + format!("Bearer {old_api_key}").parse().unwrap(), + ); + + let original = resolve_data_backed_auth_context( + &state, + &headers, + &uri("/v1/responses"), + Some("openai:responses"), + ) + .await + .expect("initial auth resolution should succeed") + .expect("the old API key should authenticate"); + assert!(original.access_allowed); + assert!(original.verified_api_key_hash.is_some()); + assert!( + !format!("{original:?}").contains(&old_key_hash), + "the credential verifier must stay redacted from Debug output" + ); + + assert!(repository + .delete_user_api_key("user-stable-id", "key-stable-id") + .await + .expect("old API key deletion should succeed")); + repository + .create_user_api_key(CreateUserApiKeyRecord { + user_id: "user-stable-id".to_string(), + api_key_id: "key-stable-id".to_string(), + key_hash: new_key_hash, + key_encrypted: None, + name: Some("restored-with-new-secret".to_string()), + allowed_providers: Some(vec!["openai".to_string()]), + allowed_api_formats: Some(vec!["openai:responses".to_string()]), + allowed_models: Some(vec!["gpt-4.1".to_string()]), + ip_rules: None, + rate_limit: 60, + concurrent_limit: Some(5), + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: Some(4_102_444_800), + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + }) + .await + .expect("same-ID API key recreation should resolve") + .expect("same-ID API key recreation should persist"); + + let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot( + &state, + original, + Some("openai:responses"), + ) + .await + .expect("long-lived auth refresh should resolve"); + + assert!(!refreshed.access_allowed); + assert_eq!( + refreshed.local_rejection, + Some(GatewayLocalAuthRejection::InvalidApiKey) + ); + assert!(snapshot.is_none()); + assert_eq!(repository.key_hash_lookup_count(&old_key_hash), 1); + assert_eq!( + repository.snapshot_lookup_count("key-stable-id"), + 0, + "a bound long-lived credential must not fall back to identity-only lookup" + ); + } + #[tokio::test] async fn control_auth_context_singleflights_concurrent_cache_misses() { let api_key = "sk-test-concurrent-auth-miss"; @@ -2396,6 +2575,44 @@ mod tests { assert_eq!(auth_context.local_rejection, None); } + #[tokio::test] + async fn data_backed_auth_context_denies_unresolved_provider_id_without_catalog_reader() { + let api_key = "sk-test-provider-no-catalog"; + let mut snapshot = sample_snapshot("key-no-catalog", "user-no-catalog"); + snapshot.user_allowed_providers = Some(vec!["provider-custom-claude".to_string()]); + snapshot.api_key_allowed_providers = Some(vec!["provider-custom-claude".to_string()]); + snapshot.user_allowed_api_formats = None; + snapshot.api_key_allowed_api_formats = None; + let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key(api_key)), + snapshot, + )])); + let data = GatewayDataState::with_auth_api_key_reader_for_tests(repository); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data); + let mut headers = HeaderMap::new(); + headers.insert("x-api-key", api_key.parse().unwrap()); + + let auth_context = resolve_data_backed_auth_context( + &state, + &headers, + &uri("/v1/messages"), + Some("claude:messages"), + ) + .await + .expect("resolution should succeed") + .expect("auth context should exist"); + + assert!(!auth_context.access_allowed); + assert_eq!( + auth_context.local_rejection, + Some(GatewayLocalAuthRejection::ProviderNotAllowed { + provider: "claude".to_string(), + }) + ); + } + #[tokio::test] async fn due_antigravity_bearer_refresh_observes_cross_node_allowlist_revocation() { let raw_bearer = "google-oauth-access-token-revoked-cross-node"; diff --git a/apps/aether-gateway/src/control/auth/types.rs b/apps/aether-gateway/src/control/auth/types.rs index 5773edfaa..f3cf21b03 100644 --- a/apps/aether-gateway/src/control/auth/types.rs +++ b/apps/aether-gateway/src/control/auth/types.rs @@ -44,7 +44,7 @@ pub(super) struct GatewayTrustedAdminHeaders { pub(super) management_token_id: Option, } -#[derive(Debug, Clone, Default, PartialEq, Eq)] +#[derive(Clone, Default, PartialEq, Eq)] pub(super) struct GatewayCredentialBundle { pub(super) authorization_bearer: Option, pub(super) x_api_key: Option, @@ -54,7 +54,25 @@ pub(super) struct GatewayCredentialBundle { pub(super) cookie_header: Option, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for GatewayCredentialBundle { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let redacted = |value: &Option| value.as_ref().map(|_| "[REDACTED]"); + formatter + .debug_struct("GatewayCredentialBundle") + .field( + "authorization_bearer", + &redacted(&self.authorization_bearer), + ) + .field("x_api_key", &redacted(&self.x_api_key)) + .field("api_key", &redacted(&self.api_key)) + .field("x_goog_api_key", &redacted(&self.x_goog_api_key)) + .field("query_key", &redacted(&self.query_key)) + .field("cookie_header", &redacted(&self.cookie_header)) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub(super) enum GatewayPrimaryCredential { ProviderApiKey { raw: String, @@ -70,6 +88,21 @@ pub(super) enum GatewayPrimaryCredential { }, } +impl std::fmt::Debug for GatewayPrimaryCredential { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let (variant, carrier) = match self { + Self::ProviderApiKey { carrier, .. } => ("ProviderApiKey", carrier), + Self::BearerToken { carrier, .. } => ("BearerToken", carrier), + Self::CookieHeader { carrier, .. } => ("CookieHeader", carrier), + }; + formatter + .debug_struct(variant) + .field("raw", &"[REDACTED]") + .field("carrier", carrier) + .finish() + } +} + #[derive(Debug, Clone, PartialEq)] pub(super) struct GatewayExtractedCredentials { pub(super) trusted_headers: Option, @@ -78,7 +111,7 @@ pub(super) struct GatewayExtractedCredentials { pub(super) primary: Option, } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub(super) enum GatewayPrincipalCandidate { TrustedHeaders(GatewayTrustedAuthHeaders), ApiKeyHash { @@ -94,3 +127,58 @@ pub(super) enum GatewayPrincipalCandidate { carrier: GatewayCredentialCarrier, }, } + +impl std::fmt::Debug for GatewayPrincipalCandidate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::TrustedHeaders(headers) => formatter + .debug_tuple("TrustedHeaders") + .field(headers) + .finish(), + Self::ApiKeyHash { carrier, .. } => formatter + .debug_struct("ApiKeyHash") + .field("key_hash", &"[REDACTED]") + .field("carrier", carrier) + .finish(), + Self::DeferredBearerToken { carrier, .. } => formatter + .debug_struct("DeferredBearerToken") + .field("raw", &"[REDACTED]") + .field("carrier", carrier) + .finish(), + Self::DeferredCookieHeader { carrier, .. } => formatter + .debug_struct("DeferredCookieHeader") + .field("raw", &"[REDACTED]") + .field("carrier", carrier) + .finish(), + } + } +} + +#[cfg(test)] +mod debug_redaction_tests { + use super::{GatewayCredentialBundle, GatewayCredentialCarrier, GatewayPrimaryCredential}; + + #[test] + fn gateway_credential_debug_output_redacts_raw_authorization_values() { + let bundle = GatewayCredentialBundle { + authorization_bearer: Some("bundle-bearer-canary".to_string()), + api_key: Some("bundle-api-key-canary".to_string()), + cookie_header: Some("bundle-cookie-canary".to_string()), + ..GatewayCredentialBundle::default() + }; + let primary = GatewayPrimaryCredential::ProviderApiKey { + raw: "primary-api-key-canary".to_string(), + carrier: GatewayCredentialCarrier::ApiKey, + }; + let debug = format!("{bundle:?} {primary:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "bundle-bearer-canary", + "bundle-api-key-canary", + "bundle-cookie-canary", + "primary-api-key-canary", + ] { + assert!(!debug.contains(secret), "debug output leaked {secret}"); + } + } +} diff --git a/apps/aether-gateway/src/control/management_token_permissions.rs b/apps/aether-gateway/src/control/management_token_permissions.rs index b4fbedd19..9052f680c 100644 --- a/apps/aether-gateway/src/control/management_token_permissions.rs +++ b/apps/aether-gateway/src/control/management_token_permissions.rs @@ -186,6 +186,52 @@ const PERMISSION_GROUPS: &[PermissionGroup] = &[ const ACCESS_LEVELS: &[(&str, &str)] = &[("read", "读取"), ("write", "写入"), ("admin", "管理")]; +// Freeze the implicit permissions granted before per-token permissions were +// introduced. New scopes must never expand legacy NULL-permission tokens. The +// management_tokens scope is deliberately withheld so a legacy token cannot +// mint an unrestricted replacement that outlives its own IP or expiry bounds. +const LEGACY_FULL_PERMISSION_SCOPES: &[&str] = &[ + "adaptive", + "announcements", + "api_keys", + "billing", + "endpoints_health", + "endpoints_manage", + "endpoints_rpm", + "gemini_files", + "ldap", + "models", + "modules", + "monitoring", + "oauth", + "payments", + "pool", + "provider_ops", + "provider_oauth", + "provider_query", + "provider_strategy", + "providers", + "proxy_nodes", + "security", + "stats", + "system", + "usage", + "users", + "video_tasks", + "wallets", +]; + +pub(crate) fn legacy_full_management_token_permissions() -> Vec { + LEGACY_FULL_PERMISSION_SCOPES + .iter() + .flat_map(|scope| { + ACCESS_LEVELS + .iter() + .map(move |(access, _)| permission_key(scope, access).to_string()) + }) + .collect() +} + pub(crate) fn management_token_permission_catalog_items( ) -> Vec { PERMISSION_GROUPS @@ -249,7 +295,7 @@ pub(crate) fn normalize_assignable_management_token_permissions( return Ok(json!(all_assignable_management_token_permissions())); }; if value.is_null() { - return Ok(json!(all_assignable_management_token_permissions())); + return Err("permissions 必须是非空字符串数组;省略该字段可使用默认权限".to_string()); } let Some(items) = value.as_array() else { return Err("permissions 必须是字符串数组".to_string()); @@ -283,7 +329,7 @@ pub(crate) fn management_token_permission_keys_from_value( return Ok(None); }; if value.is_null() { - return Ok(None); + return Err("management token permissions JSON null is invalid".to_string()); } let Some(items) = value.as_array() else { return Err("management token permissions must be an array".to_string()); @@ -338,7 +384,10 @@ pub(crate) fn management_token_required_permission( if scope.is_empty() { return None; } - Some(format!("admin:{scope}:{}", access_for_method(method))) + Some(format!( + "admin:{scope}:{}", + access_for_route(method, decision) + )) } pub(crate) fn validate_management_token_admin_route_permission( @@ -346,23 +395,26 @@ pub(crate) fn validate_management_token_admin_route_permission( decision: &GatewayControlDecision, token_permissions: Option<&[String]>, ) -> Result<(), ManagementTokenPermissionDenied> { - let Some(required_permission) = management_token_required_permission(method, decision) else { - return Ok(()); - }; let Some(token_permissions) = token_permissions else { return Ok(()); }; + let Some(required_permission) = management_token_required_permission(method, decision) else { + return if decision.route_class.as_deref() == Some("admin_proxy") { + Err(ManagementTokenPermissionDenied { + required_permission: "admin:unknown:admin".to_string(), + }) + } else { + Ok(()) + }; + }; let scope = required_permission .strip_prefix("admin:") .and_then(|value| value.rsplit_once(':').map(|(scope, _)| scope)) .unwrap_or_default(); let admin_permission = format!("admin:{scope}:admin"); - let has_full_assignable_access = - management_token_permissions_cover_all_assignable_permissions(token_permissions); if token_permissions .iter() .any(|permission| permission == &required_permission || permission == &admin_permission) - || (scope == "management_tokens" && has_full_assignable_access) { Ok(()) } else { @@ -372,6 +424,30 @@ pub(crate) fn validate_management_token_admin_route_permission( } } +pub(crate) fn management_token_principal_has_permission( + decision: &GatewayControlDecision, + required_permission: &str, +) -> bool { + let Some(principal) = decision.admin_principal.as_ref() else { + return false; + }; + if principal.management_token_id.is_none() { + return true; + } + + let legacy_permissions; + let permissions = match principal.management_token_permissions.as_deref() { + Some(permissions) => permissions, + None => { + legacy_permissions = legacy_full_management_token_permissions(); + legacy_permissions.as_slice() + } + }; + permissions + .iter() + .any(|permission| permission == required_permission) +} + fn access_for_method(method: &http::Method) -> &'static str { if matches!( *method, @@ -383,6 +459,147 @@ fn access_for_method(method: &http::Method) -> &'static str { } } +fn access_for_route(method: &http::Method, decision: &GatewayControlDecision) -> &'static str { + let signature = decision.auth_endpoint_signature.as_deref(); + let route_kind = decision.route_kind.as_deref(); + let requires_admin_permission = matches!( + (signature, route_kind), + ( + Some("admin:system"), + Some( + "prepare_update" + | "apply_update" + | "rollback" + | "config_export" + | "users_export" + | "data_export" + | "s3_backup_run" + | "config_import" + | "users_import" + | "data_import" + | "smtp_test" + | "important_notification_test" + | "cleanup" + | "cleanup_usage_manual" + | "purge_config" + | "purge_users" + | "purge_usage" + | "purge_audit_logs" + | "purge_request_bodies" + | "purge_request_bodies_task" + | "purge_stats" + | "settings_set" + | "config_set" + | "config_delete" + ) + ) | ( + Some("admin:endpoints_manage"), + Some( + "reveal_key" + | "export_key" + | "create_provider_key" + | "update_key" + | "create_endpoint" + | "update_endpoint" + | "refresh_quota" + | "codex_reset_credit_consume" + ) + ) | (Some("admin:providers"), Some("update_provider")) + | ( + Some("admin:provider_query"), + Some("query_models" | "test_model" | "test_model_failover") + ) + | ( + Some("admin:provider_oauth"), + Some( + "complete_key_oauth" + | "refresh_key_oauth" + | "complete_provider_oauth" + | "import_refresh_token" + | "cookie_authorize" + | "start_cookie_authorize_task" + | "start_agent_identity_import_task" + | "batch_import_oauth" + | "start_batch_import_oauth_task" + | "device_poll" + ) + ) + | ( + Some("admin:management_tokens"), + Some("create_token" | "regenerate_token") + ) + | ( + Some("admin:provider_ops"), + Some( + "connect_provider" + | "verify_provider" + | "get_provider_balance" + | "refresh_provider_balance" + | "provider_checkin" + | "execute_provider_action" + | "batch_balance" + ) + ) + | ( + Some("admin:security"), + Some("blacklist_add" | "blacklist_remove" | "whitelist_add" | "whitelist_remove") + ) + | (Some("admin:ldap"), Some("set_config" | "test_connection")) + | (Some("admin:modules"), Some("set_enabled")) + | ( + Some("admin:users"), + Some("reveal_user_api_key" | "create_user_api_key") + ) + | ( + Some("admin:oauth"), + Some("upsert_provider" | "delete_provider") + ) + | ( + Some("admin:payments"), + Some("update_epay_gateway" | "update_payment_gateway") + ) + | ( + Some("admin:api_keys"), + Some("create_api_key" | "create_api_key_install_session") + ) + | ( + Some("admin:proxy_nodes"), + Some("create_proxy_node_install_session") + ) + | (Some("admin:usage"), Some("detail" | "curl" | "replay")) + | (Some("admin:monitoring"), Some("trace_request")) + | (Some("admin:tasks"), Some("detail" | "events")) + | ( + Some("admin:pool"), + Some("batch_action_keys" | "batch_update_keys") + ) + ) || (*method == http::Method::GET + && signature == Some("admin:api_keys") + && route_kind == Some("api_key_detail") + && query_param_enabled(decision.public_query_string.as_deref(), "include_key")); + + if requires_admin_permission { + // These routes reveal usable credentials or replace security-critical + // identity and system configuration. They are not delegated writes. + "admin" + } else { + access_for_method(method) + } +} + +fn query_param_enabled(query: Option<&str>, key: &str) -> bool { + query + .into_iter() + .flat_map(|query| url::form_urlencoded::parse(query.as_bytes())) + .find(|(entry_key, _)| entry_key == key) + .is_some_and(|(_, value)| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "1" | "true" | "yes" | "on" + ) + }) +} + fn permission_key(scope: &str, access: &str) -> &'static str { match (scope, access) { ("adaptive", "read") => "admin:adaptive:read", @@ -492,18 +709,6 @@ fn is_assignable_management_token_permission(key: &str) -> bool { .any(|item| item.key == key) } -pub(crate) fn management_token_permissions_cover_all_assignable_permissions( - token_permissions: &[String], -) -> bool { - let permission_set = token_permissions - .iter() - .map(String::as_str) - .collect::>(); - all_assignable_management_token_permissions() - .iter() - .all(|permission| permission_set.contains(permission.as_str())) -} - #[cfg(test)] mod tests { use super::*; @@ -618,7 +823,7 @@ mod tests { } #[test] - fn full_assignable_token_permissions_can_cover_management_tokens_scope() { + fn full_assignable_token_permissions_cannot_cover_management_tokens_scope() { let decision = GatewayControlDecision::synthetic( "/api/admin/management-tokens".to_string(), Some("admin_proxy".to_string()), @@ -628,12 +833,42 @@ mod tests { ); let permissions = all_assignable_management_token_permissions(); - assert!(validate_management_token_admin_route_permission( - &http::Method::GET, - &decision, - Some(&permissions), - ) - .is_ok()); + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + Some(&permissions), + ) + .expect_err("management-token administration is not assignable") + .required_permission, + "admin:management_tokens:read" + ); + } + + #[test] + fn legacy_tokens_cannot_administer_management_tokens() { + let decision = GatewayControlDecision::synthetic( + "/api/admin/management-tokens".to_string(), + Some("admin_proxy".to_string()), + Some("management_tokens_manage".to_string()), + Some("create_token".to_string()), + Some("admin:management_tokens".to_string()), + ); + let permissions = legacy_full_management_token_permissions(); + + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&permissions), + ) + .expect_err("legacy tokens must not mint replacement credentials") + .required_permission, + "admin:management_tokens:admin" + ); + assert!(!permissions + .iter() + .any(|permission| permission.starts_with("admin:management_tokens:"))); } #[test] @@ -665,6 +900,953 @@ mod tests { ); } + #[test] + fn scoped_tokens_fail_closed_for_unclassified_admin_permissions() { + let decision = GatewayControlDecision::synthetic( + "/api/admin/unclassified".to_string(), + Some("admin_proxy".to_string()), + Some("unclassified_manage".to_string()), + Some("unclassified".to_string()), + None, + ); + let permissions = vec!["admin:system:admin".to_string()]; + + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + Some(&permissions), + ) + .expect_err("scoped tokens must not bypass an unclassified admin route") + .required_permission, + "admin:unknown:admin" + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + None, + ) + .is_ok()); + } + + #[test] + fn legacy_full_permissions_are_frozen_to_the_original_scope_catalog() { + let permissions = legacy_full_management_token_permissions(); + + for permission in [ + "admin:system:admin", + "admin:oauth:admin", + "admin:proxy_nodes:admin", + ] { + assert!(permissions.iter().any(|item| item == permission)); + } + for permission in [ + "admin:routing_profiles:read", + "admin:routing_profiles:admin", + "admin:tasks:read", + "admin:tasks:admin", + ] { + assert!( + !permissions.iter().any(|item| item == permission), + "legacy token unexpectedly gained {permission}" + ); + } + } + + #[test] + fn json_null_permissions_never_inherit_legacy_full_access() { + assert!(normalize_assignable_management_token_permissions(Some(&Value::Null)).is_err()); + assert!(management_token_permission_keys_from_value(Some(&Value::Null)).is_err()); + + assert_eq!( + management_token_permission_keys_from_value(None) + .expect("SQL NULL remains the explicit legacy representation"), + None + ); + } + + #[test] + fn sensitive_action_permission_distinguishes_sessions_and_scoped_tokens() { + let mut decision = GatewayControlDecision::synthetic( + "/api/admin/users/user-1".to_string(), + Some("admin_proxy".to_string()), + Some("users_manage".to_string()), + Some("update_user".to_string()), + Some("admin:users".to_string()), + ); + + assert!(!management_token_principal_has_permission( + &decision, + "admin:users:admin" + )); + + decision.admin_principal = Some(crate::control::GatewayAdminPrincipalContext { + user_id: "admin-1".to_string(), + user_role: "admin".to_string(), + session_id: Some("session-1".to_string()), + management_token_id: None, + management_token_permissions: None, + }); + assert!(management_token_principal_has_permission( + &decision, + "admin:users:admin" + )); + + let principal = decision + .admin_principal + .as_mut() + .expect("principal should exist"); + principal.session_id = None; + principal.management_token_id = Some("management-token-1".to_string()); + principal.management_token_permissions = Some(vec!["admin:users:write".to_string()]); + assert!(!management_token_principal_has_permission( + &decision, + "admin:users:admin" + )); + + decision + .admin_principal + .as_mut() + .expect("principal should exist") + .management_token_permissions = Some(vec!["admin:users:admin".to_string()]); + assert!(management_token_principal_has_permission( + &decision, + "admin:users:admin" + )); + } + + #[test] + fn sensitive_system_data_transfers_require_admin_permission() { + let delegated_permissions = vec![ + "admin:system:read".to_string(), + "admin:system:write".to_string(), + ]; + let admin_permissions = vec!["admin:system:admin".to_string()]; + + for (method, route_kind) in [ + (http::Method::GET, "config_export"), + (http::Method::GET, "users_export"), + (http::Method::GET, "data_export"), + (http::Method::POST, "s3_backup_run"), + (http::Method::POST, "config_import"), + (http::Method::POST, "users_import"), + (http::Method::POST, "data_import"), + ] { + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/system/{route_kind}"), + Some("admin_proxy".to_string()), + Some("system_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:system".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:system:admin") + ); + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(&delegated_permissions), + ) + .expect_err("delegated system access must not transfer sensitive data") + .required_permission, + "admin:system:admin" + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn oauth_provider_upsert_requires_oauth_admin_permission() { + let decision = GatewayControlDecision::synthetic( + "/api/admin/oauth/providers/custom".to_string(), + Some("admin_proxy".to_string()), + Some("oauth_manage".to_string()), + Some("upsert_provider".to_string()), + Some("admin:oauth".to_string()), + ); + let write_permissions = vec!["admin:oauth:write".to_string()]; + let admin_permissions = vec!["admin:oauth:admin".to_string()]; + + assert_eq!( + management_token_required_permission(&http::Method::PUT, &decision).as_deref(), + Some("admin:oauth:admin") + ); + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::PUT, + &decision, + Some(&write_permissions), + ) + .expect_err("oauth write must not replace security-sensitive provider configuration") + .required_permission, + "admin:oauth:admin" + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::PUT, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + + #[test] + fn oauth_provider_delete_requires_oauth_admin_permission() { + let decision = GatewayControlDecision::synthetic( + "/api/admin/oauth/providers/custom".to_string(), + Some("admin_proxy".to_string()), + Some("oauth_manage".to_string()), + Some("delete_provider".to_string()), + Some("admin:oauth".to_string()), + ); + let write_permissions = vec!["admin:oauth:write".to_string()]; + let admin_permissions = vec!["admin:oauth:admin".to_string()]; + + assert_eq!( + management_token_required_permission(&http::Method::DELETE, &decision).as_deref(), + Some("admin:oauth:admin") + ); + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::DELETE, + &decision, + Some(&write_permissions), + ) + .expect_err("oauth write must not delete an authentication provider") + .required_permission, + "admin:oauth:admin" + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::DELETE, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + + #[test] + fn destructive_system_actions_require_system_admin_permission() { + let write_permissions = vec!["admin:system:write".to_string()]; + let admin_permissions = vec!["admin:system:admin".to_string()]; + + for (method, route_kind) in [ + (http::Method::POST, "prepare_update"), + (http::Method::POST, "apply_update"), + (http::Method::POST, "rollback"), + (http::Method::POST, "cleanup"), + (http::Method::POST, "cleanup_usage_manual"), + (http::Method::POST, "purge_config"), + (http::Method::POST, "purge_users"), + (http::Method::POST, "purge_usage"), + (http::Method::POST, "purge_audit_logs"), + (http::Method::POST, "purge_request_bodies"), + (http::Method::POST, "purge_request_bodies_task"), + (http::Method::POST, "purge_stats"), + (http::Method::PUT, "settings_set"), + (http::Method::PUT, "config_set"), + (http::Method::DELETE, "config_delete"), + ] { + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/system/{route_kind}"), + Some("admin_proxy".to_string()), + Some("system_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:system".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:system:admin"), + "unexpected permission for {method} {route_kind}" + ); + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(&write_permissions), + ) + .expect_err("system write must not perform destructive or security-critical action") + .required_permission, + "admin:system:admin" + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn provider_operations_that_use_stored_credentials_require_provider_ops_admin_permission() { + let read_permissions = vec!["admin:provider_ops:read".to_string()]; + let write_permissions = vec!["admin:provider_ops:write".to_string()]; + let admin_permissions = vec!["admin:provider_ops:admin".to_string()]; + + for (method, route_kind) in [ + (http::Method::POST, "connect_provider"), + (http::Method::POST, "verify_provider"), + (http::Method::GET, "get_provider_balance"), + (http::Method::POST, "refresh_provider_balance"), + (http::Method::POST, "provider_checkin"), + (http::Method::POST, "execute_provider_action"), + (http::Method::POST, "batch_balance"), + ] { + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/provider-ops/{route_kind}"), + Some("admin_proxy".to_string()), + Some("provider_ops_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:provider_ops".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:provider_ops:admin") + ); + for delegated_permissions in [&read_permissions, &write_permissions] { + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(delegated_permissions), + ) + .expect_err("delegated access must not execute stored provider credentials") + .required_permission, + "admin:provider_ops:admin" + ); + } + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn provider_queries_that_execute_stored_credentials_require_admin_permission() { + let write_permissions = vec!["admin:provider_query:write".to_string()]; + let admin_permissions = vec!["admin:provider_query:admin".to_string()]; + + for route_kind in ["query_models", "test_model", "test_model_failover"] { + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/provider-query/{route_kind}"), + Some("admin_proxy".to_string()), + Some("provider_query_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:provider_query".to_string()), + ); + + assert_eq!( + management_token_required_permission(&http::Method::POST, &decision).as_deref(), + Some("admin:provider_query:admin") + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&write_permissions), + ) + .is_err()); + assert!(validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn endpoint_actions_that_execute_stored_credentials_require_admin_permission() { + let write_permissions = vec!["admin:endpoints_manage:write".to_string()]; + let admin_permissions = vec!["admin:endpoints_manage:admin".to_string()]; + + for (method, route_kind) in [ + (http::Method::PUT, "update_key"), + (http::Method::POST, "create_endpoint"), + (http::Method::PUT, "update_endpoint"), + (http::Method::POST, "refresh_quota"), + (http::Method::POST, "codex_reset_credit_consume"), + ] { + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/endpoints/{route_kind}"), + Some("admin_proxy".to_string()), + Some("endpoints_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:endpoints_manage".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:endpoints_manage:admin") + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&write_permissions), + ) + .is_err()); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn stored_credential_egress_configuration_requires_admin_permission() { + for (method, signature, route_kind) in [ + (http::Method::PATCH, "admin:providers", "update_provider"), + (http::Method::POST, "admin:pool", "batch_action_keys"), + (http::Method::PATCH, "admin:pool", "batch_update_keys"), + ] { + let scope = signature.trim_start_matches("admin:"); + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/{scope}/{route_kind}"), + Some("admin_proxy".to_string()), + Some(format!("{scope}_manage")), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + let write_permissions = vec![format!("{signature}:write")]; + let admin_permissions = vec![format!("{signature}:admin")]; + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some(format!("{signature}:admin").as_str()) + ); + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(&write_permissions), + ) + .expect_err("delegated writes must not redirect stored provider credentials") + .required_permission, + format!("{signature}:admin") + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn connection_tests_that_transmit_stored_secrets_require_admin_permission() { + for (signature, route_kind, read_permission, write_permission, admin_permission) in [ + ( + "admin:system", + "smtp_test", + "admin:system:read", + "admin:system:write", + "admin:system:admin", + ), + ( + "admin:system", + "important_notification_test", + "admin:system:read", + "admin:system:write", + "admin:system:admin", + ), + ( + "admin:ldap", + "test_connection", + "admin:ldap:read", + "admin:ldap:write", + "admin:ldap:admin", + ), + ] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/security-sensitive-test".to_string(), + Some("admin_proxy".to_string()), + Some("security_test".to_string()), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + let read_permissions = vec![read_permission.to_string()]; + let write_permissions = vec![write_permission.to_string()]; + let admin_permissions = vec![admin_permission.to_string()]; + + assert_eq!( + management_token_required_permission(&http::Method::POST, &decision).as_deref(), + Some(admin_permission) + ); + for delegated_permissions in [&read_permissions, &write_permissions] { + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(delegated_permissions), + ) + .expect_err("delegated access must not transmit stored service secrets") + .required_permission, + admin_permission + ); + } + assert!(validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn security_control_plane_mutations_require_admin_permission() { + for (method, signature, route_kind) in [ + (http::Method::POST, "admin:security", "blacklist_add"), + (http::Method::DELETE, "admin:security", "blacklist_remove"), + (http::Method::POST, "admin:security", "whitelist_add"), + (http::Method::DELETE, "admin:security", "whitelist_remove"), + (http::Method::PUT, "admin:ldap", "set_config"), + (http::Method::PUT, "admin:modules", "set_enabled"), + ] { + let scope = signature.trim_start_matches("admin:"); + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/{scope}/{route_kind}"), + Some("admin_proxy".to_string()), + Some(format!("{scope}_manage")), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + let read_permissions = vec![format!("{signature}:read")]; + let write_permissions = vec![format!("{signature}:write")]; + let admin_permissions = vec![format!("{signature}:admin")]; + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some(format!("{signature}:admin").as_str()) + ); + for delegated_permissions in [&read_permissions, &write_permissions] { + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(delegated_permissions), + ) + .expect_err("delegated tokens must not change security controls") + .required_permission, + format!("{signature}:admin") + ); + } + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn payment_gateway_updates_require_payments_admin_permission() { + for route_kind in ["update_epay_gateway", "update_payment_gateway"] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/payments/gateways/epay".to_string(), + Some("admin_proxy".to_string()), + Some("payments_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:payments".to_string()), + ); + let write_permissions = vec!["admin:payments:write".to_string()]; + let admin_permissions = vec!["admin:payments:admin".to_string()]; + + assert_eq!( + management_token_required_permission(&http::Method::PUT, &decision).as_deref(), + Some("admin:payments:admin") + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::PUT, + &decision, + Some(&write_permissions), + ) + .is_err()); + assert!(validate_management_token_admin_route_permission( + &http::Method::PUT, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn plaintext_credential_reads_require_admin_permission() { + let read_only_permissions = read_only_management_token_permissions(); + let cases = [ + ( + "admin:endpoints_manage", + "reveal_key", + None, + "admin:endpoints_manage:admin", + ), + ( + "admin:endpoints_manage", + "export_key", + None, + "admin:endpoints_manage:admin", + ), + ( + "admin:users", + "reveal_user_api_key", + None, + "admin:users:admin", + ), + ( + "admin:api_keys", + "api_key_detail", + Some("include_key=true"), + "admin:api_keys:admin", + ), + ( + "admin:api_keys", + "create_api_key_install_session", + None, + "admin:api_keys:admin", + ), + ( + "admin:proxy_nodes", + "create_proxy_node_install_session", + None, + "admin:proxy_nodes:admin", + ), + ]; + + for (signature, route_kind, query, expected) in cases { + let mut decision = GatewayControlDecision::synthetic( + "/api/admin/sensitive-read".to_string(), + Some("admin_proxy".to_string()), + Some("security_test".to_string()), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + decision.public_query_string = query.map(str::to_string); + + let method = if route_kind.ends_with("install_session") { + http::Method::POST + } else { + http::Method::GET + }; + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some(expected) + ); + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(&read_only_permissions), + ) + .expect_err("read-only access must not reveal plaintext credentials") + .required_permission, + expected + ); + } + } + + #[test] + fn raw_usage_and_trace_reads_require_admin_permission() { + let cases = [ + ("admin:usage", "detail", http::Method::GET), + ("admin:usage", "curl", http::Method::GET), + ("admin:usage", "replay", http::Method::POST), + ("admin:monitoring", "trace_request", http::Method::GET), + ]; + + for (signature, route_kind, method) in cases { + let scope = signature.trim_start_matches("admin:"); + let read_permissions = vec![format!("{signature}:read")]; + let admin_permissions = vec![format!("{signature}:admin")]; + let decision = GatewayControlDecision::synthetic( + format!("/api/admin/{scope}/sensitive-record"), + Some("admin_proxy".to_string()), + Some(format!("{scope}_manage")), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some(format!("{signature}:admin").as_str()) + ); + assert_eq!( + validate_management_token_admin_route_permission( + &method, + &decision, + Some(&read_permissions), + ) + .expect_err("delegated read access must not expose raw request diagnostics") + .required_permission, + format!("{signature}:admin") + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + + for (signature, route_kind) in [ + ("admin:usage", "records"), + ("admin:monitoring", "trace_provider_stats"), + ] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/summary".to_string(), + Some("admin_proxy".to_string()), + Some("summary".to_string()), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + assert_eq!( + management_token_required_permission(&http::Method::GET, &decision).as_deref(), + Some(format!("{signature}:read").as_str()) + ); + } + } + + #[test] + fn task_read_permission_excludes_raw_detail_and_events() { + let read_permissions = vec!["admin:tasks:read".to_string()]; + let admin_permissions = vec!["admin:tasks:admin".to_string()]; + + for route_kind in ["list_tasks", "stats"] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/tasks".to_string(), + Some("admin_proxy".to_string()), + Some("tasks_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:tasks".to_string()), + ); + assert_eq!( + management_token_required_permission(&http::Method::GET, &decision).as_deref(), + Some("admin:tasks:read") + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + Some(&read_permissions), + ) + .is_ok()); + } + + for route_kind in ["detail", "events"] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/tasks/run-1".to_string(), + Some("admin_proxy".to_string()), + Some("tasks_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:tasks".to_string()), + ); + assert_eq!( + management_token_required_permission(&http::Method::GET, &decision).as_deref(), + Some("admin:tasks:admin") + ); + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + Some(&read_permissions), + ) + .expect_err("task read access must not expose raw task diagnostics") + .required_permission, + "admin:tasks:admin" + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + assert!(validate_management_token_admin_route_permission( + &http::Method::GET, + &decision, + None, + ) + .is_ok()); + } + } + + #[test] + fn user_lifecycle_routes_preserve_write_permission_for_field_level_checks() { + let write_permissions = vec!["admin:users:write".to_string()]; + let cases = [ + (http::Method::POST, "/api/admin/users", "create_user"), + (http::Method::PUT, "/api/admin/users/user-1", "update_user"), + ( + http::Method::POST, + "/api/admin/users/batch-action", + "batch_action_users", + ), + ]; + + for (method, path, route_kind) in cases { + let decision = GatewayControlDecision::synthetic( + path.to_string(), + Some("admin_proxy".to_string()), + Some("users_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:users".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:users:write"), + "unexpected permission for {method} {path}" + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&write_permissions), + ) + .is_ok()); + } + } + + #[test] + fn ordinary_user_resource_mutations_still_require_users_write_permission() { + let write_permissions = vec!["admin:users:write".to_string()]; + + for (method, route_kind) in [ + (http::Method::PUT, "update_user_api_key"), + (http::Method::DELETE, "delete_user_session"), + ] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/users/user-1/resource".to_string(), + Some("admin_proxy".to_string()), + Some("users_manage".to_string()), + Some(route_kind.to_string()), + Some("admin:users".to_string()), + ); + + assert_eq!( + management_token_required_permission(&method, &decision).as_deref(), + Some("admin:users:write") + ); + assert!(validate_management_token_admin_route_permission( + &method, + &decision, + Some(&write_permissions), + ) + .is_ok()); + } + } + + #[test] + fn usable_credential_creation_requires_admin_permission() { + for (signature, route_kind, write_permission, admin_permission) in [ + ( + "admin:management_tokens", + "create_token", + "admin:management_tokens:write", + "admin:management_tokens:admin", + ), + ( + "admin:management_tokens", + "regenerate_token", + "admin:management_tokens:write", + "admin:management_tokens:admin", + ), + ( + "admin:endpoints_manage", + "create_provider_key", + "admin:endpoints_manage:write", + "admin:endpoints_manage:admin", + ), + ( + "admin:users", + "create_user_api_key", + "admin:users:write", + "admin:users:admin", + ), + ( + "admin:api_keys", + "create_api_key", + "admin:api_keys:write", + "admin:api_keys:admin", + ), + ] { + let decision = GatewayControlDecision::synthetic( + "/api/admin/credential".to_string(), + Some("admin_proxy".to_string()), + Some("credential_security_test".to_string()), + Some(route_kind.to_string()), + Some(signature.to_string()), + ); + let write_permissions = vec![write_permission.to_string()]; + let admin_permissions = vec![admin_permission.to_string()]; + + assert_eq!( + management_token_required_permission(&http::Method::POST, &decision).as_deref(), + Some(admin_permission) + ); + assert_eq!( + validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&write_permissions), + ) + .expect_err("delegated write access must not issue usable credentials") + .required_permission, + admin_permission + ); + assert!(validate_management_token_admin_route_permission( + &http::Method::POST, + &decision, + Some(&admin_permissions), + ) + .is_ok()); + } + } + + #[test] + fn ordinary_system_reads_still_require_read_permission() { + let decision = GatewayControlDecision::synthetic( + "/api/admin/system/settings".to_string(), + Some("admin_proxy".to_string()), + Some("system_manage".to_string()), + Some("settings_get".to_string()), + Some("admin:system".to_string()), + ); + + assert_eq!( + management_token_required_permission(&http::Method::GET, &decision).as_deref(), + Some("admin:system:read") + ); + + let mut api_key_detail = GatewayControlDecision::synthetic( + "/api/admin/api-keys/key-1".to_string(), + Some("admin_proxy".to_string()), + Some("api_keys_manage".to_string()), + Some("api_key_detail".to_string()), + Some("admin:api_keys".to_string()), + ); + for query in [None, Some("include_key=false"), Some("include_key=invalid")] { + api_key_detail.public_query_string = query.map(str::to_string); + assert_eq!( + management_token_required_permission(&http::Method::GET, &api_key_detail) + .as_deref(), + Some("admin:api_keys:read") + ); + } + } + #[test] fn audit_admin_read_only_permissions_allow_management_tokens_reads_and_reject_writes() { let decision = GatewayControlDecision::synthetic( diff --git a/apps/aether-gateway/src/control/mod.rs b/apps/aether-gateway/src/control/mod.rs index ca9531692..40ef186ac 100644 --- a/apps/aether-gateway/src/control/mod.rs +++ b/apps/aether-gateway/src/control/mod.rs @@ -8,9 +8,10 @@ mod public; mod route; pub(crate) use auth::{ - execution_plan_balance_capacity_rejection, extract_requested_model, - refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot, - request_model_local_rejection, resolve_execution_runtime_auth_context, + estimate_execution_plan_cost_upper_bound_usd, execution_plan_balance_capacity_rejection, + extract_requested_model, refresh_execution_runtime_auth_context, + refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection, + resolve_execution_runtime_auth_context, resolve_local_admin_session_principal, should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayCredentialCarrier, GatewayLocalAuthRejection, @@ -18,14 +19,16 @@ pub(crate) use auth::{ pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control}; pub(crate) use management_token_permissions::{ all_assignable_management_token_permissions, - audit_admin_read_only_management_token_permissions, + audit_admin_read_only_management_token_permissions, legacy_full_management_token_permissions, management_token_permission_catalog_payload, management_token_permission_keys_from_value, - management_token_permission_mode_and_summary, - management_token_permissions_cover_all_assignable_permissions, + management_token_permission_mode_and_summary, management_token_principal_has_permission, management_token_required_permission, normalize_assignable_management_token_permissions, read_only_management_token_permissions, validate_management_token_admin_route_permission, }; -pub(crate) use public::{resolve_public_request_context, GatewayPublicRequestContext}; +pub(crate) use public::{ + resolve_public_request_context, resolve_public_request_context_with_trusted_auth, + resolve_public_request_context_without_trusted_auth, GatewayPublicRequestContext, +}; #[cfg(test)] pub(crate) use route::classify_control_route; pub(crate) use route::{resolve_control_route, GatewayControlDecision}; diff --git a/apps/aether-gateway/src/control/public.rs b/apps/aether-gateway/src/control/public.rs index 3c12672a1..1e20fe5a9 100644 --- a/apps/aether-gateway/src/control/public.rs +++ b/apps/aether-gateway/src/control/public.rs @@ -2,7 +2,9 @@ use axum::http::Uri; use crate::{AppState, GatewayError}; -use super::{resolve_control_route, GatewayControlDecision}; +use super::{ + resolve_control_route, route::resolve_control_route_with_trusted_auth, GatewayControlDecision, +}; pub(crate) type GatewayPublicRequestContext = aether_gateway_control::PublicRequestContext; @@ -23,3 +25,41 @@ pub(crate) async fn resolve_public_request_context( control_decision, )) } + +pub(crate) async fn resolve_public_request_context_with_trusted_auth( + state: &AppState, + method: &http::Method, + uri: &Uri, + headers: &http::HeaderMap, + trace_id: &str, +) -> Result { + let control_decision = + resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, true) + .await?; + Ok(GatewayPublicRequestContext::from_request_parts( + trace_id, + method, + uri, + headers, + control_decision, + )) +} + +pub(crate) async fn resolve_public_request_context_without_trusted_auth( + state: &AppState, + method: &http::Method, + uri: &Uri, + headers: &http::HeaderMap, + trace_id: &str, +) -> Result { + let control_decision = + resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, false) + .await?; + Ok(GatewayPublicRequestContext::from_request_parts( + trace_id, + method, + uri, + headers, + control_decision, + )) +} diff --git a/apps/aether-gateway/src/control/route/admin/operations_families.rs b/apps/aether-gateway/src/control/route/admin/operations_families.rs index 7ceca03cd..69cf18354 100644 --- a/apps/aether-gateway/src/control/route/admin/operations_families.rs +++ b/apps/aether-gateway/src/control/route/admin/operations_families.rs @@ -797,6 +797,18 @@ pub(super) fn classify_admin_operations_family_route( "admin:users", false, )) + } else if method == http::Method::DELETE + && normalized_path_no_trailing.starts_with("/api/admin/users/") + && normalized_path_no_trailing.contains("/billing/entitlements/") + && normalized_path_no_trailing.matches('/').count() == 7 + { + Some(classified( + "admin_proxy", + "users_manage", + "revoke_user_billing_entitlement", + "admin:users", + false, + )) } else if method == http::Method::GET && normalized_path.starts_with("/api/admin/users/") && normalized_path.ends_with("/sessions") diff --git a/apps/aether-gateway/src/control/route/internal.rs b/apps/aether-gateway/src/control/route/internal.rs index 040c4270f..463b8d845 100644 --- a/apps/aether-gateway/src/control/route/internal.rs +++ b/apps/aether-gateway/src/control/route/internal.rs @@ -5,19 +5,21 @@ pub(super) fn classify_internal_route( method: &http::Method, normalized_path: &str, ) -> Option { - if method == http::Method::POST && normalized_path.starts_with("/api/internal/gateway/") { - let route_kind = match normalized_path { - "/api/internal/gateway/resolve" => "resolve", - "/api/internal/gateway/auth-context" => "auth_context", - "/api/internal/gateway/decision-sync" => "decision_sync", - "/api/internal/gateway/decision-stream" => "decision_stream", - "/api/internal/gateway/plan-sync" => "plan_sync", - "/api/internal/gateway/plan-stream" => "plan_stream", - "/api/internal/gateway/report-sync" => "report_sync", - "/api/internal/gateway/report-stream" => "report_stream", - "/api/internal/gateway/finalize-sync" => "finalize_sync", - "/api/internal/gateway/execute-sync" => "execute_sync", - "/api/internal/gateway/execute-stream" => "execute_stream", + if normalized_path == "/api/internal/gateway" + || normalized_path.starts_with("/api/internal/gateway/") + { + let route_kind = match (method, normalized_path) { + (&http::Method::POST, "/api/internal/gateway/resolve") => "resolve", + (&http::Method::POST, "/api/internal/gateway/auth-context") => "auth_context", + (&http::Method::POST, "/api/internal/gateway/decision-sync") => "decision_sync", + (&http::Method::POST, "/api/internal/gateway/decision-stream") => "decision_stream", + (&http::Method::POST, "/api/internal/gateway/plan-sync") => "plan_sync", + (&http::Method::POST, "/api/internal/gateway/plan-stream") => "plan_stream", + (&http::Method::POST, "/api/internal/gateway/report-sync") => "report_sync", + (&http::Method::POST, "/api/internal/gateway/report-stream") => "report_stream", + (&http::Method::POST, "/api/internal/gateway/finalize-sync") => "finalize_sync", + (&http::Method::POST, "/api/internal/gateway/execute-sync") => "execute_sync", + (&http::Method::POST, "/api/internal/gateway/execute-stream") => "execute_stream", _ => "unhandled", }; Some(classified( diff --git a/apps/aether-gateway/src/control/route/mod.rs b/apps/aether-gateway/src/control/route/mod.rs index ee76b2765..8e27588bf 100644 --- a/apps/aether-gateway/src/control/route/mod.rs +++ b/apps/aether-gateway/src/control/route/mod.rs @@ -11,7 +11,7 @@ mod oauth; mod public_support; use super::auth::{ - resolve_control_decision_auth, resolve_gateway_credential_carrier, + resolve_control_decision_auth_with_trusted_auth, resolve_gateway_credential_carrier, ControlDecisionAuthResolution, GatewayCredentialCarrier, }; use super::{GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayLocalAuthRejection}; @@ -175,6 +175,17 @@ pub(crate) async fn resolve_control_route( uri: &Uri, headers: &http::HeaderMap, trace_id: &str, +) -> Result, GatewayError> { + resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, cfg!(test)).await +} + +pub(crate) async fn resolve_control_route_with_trusted_auth( + state: &AppState, + method: &http::Method, + uri: &Uri, + headers: &http::HeaderMap, + trace_id: &str, + trusted_auth_verified: bool, ) -> Result, GatewayError> { let Some(mut decision) = classify_control_route(method, uri, headers) else { return Ok(None); @@ -185,7 +196,16 @@ pub(crate) async fn resolve_control_route( crate::system_features::ModelDirectivePolicySnapshot::load(state).await; } - match resolve_control_decision_auth(state, headers, uri, trace_id, decision).await? { + match resolve_control_decision_auth_with_trusted_auth( + state, + headers, + uri, + trace_id, + decision, + trusted_auth_verified, + ) + .await? + { ControlDecisionAuthResolution::Resolved(decision) => Ok(Some(decision)), } } diff --git a/apps/aether-gateway/src/control/route/oauth.rs b/apps/aether-gateway/src/control/route/oauth.rs index ab6279f74..6fd9e77ff 100644 --- a/apps/aether-gateway/src/control/route/oauth.rs +++ b/apps/aether-gateway/src/control/route/oauth.rs @@ -285,7 +285,7 @@ pub(super) fn classify_oauth_route( "admin_proxy", "provider_oauth_manage", "batch_import_oauth", - "admin:pool", + "admin:provider_oauth", false, )) } else if method == http::Method::POST @@ -296,7 +296,7 @@ pub(super) fn classify_oauth_route( "admin_proxy", "provider_oauth_manage", "start_batch_import_oauth_task", - "admin:pool", + "admin:provider_oauth", false, )) } else if method == http::Method::GET @@ -307,7 +307,7 @@ pub(super) fn classify_oauth_route( "admin_proxy", "provider_oauth_manage", "get_batch_import_task_status", - "admin:pool", + "admin:provider_oauth", false, )) } else if method == http::Method::POST diff --git a/apps/aether-gateway/src/control/route/public_support.rs b/apps/aether-gateway/src/control/route/public_support.rs index 634ff1afd..77e662919 100644 --- a/apps/aether-gateway/src/control/route/public_support.rs +++ b/apps/aether-gateway/src/control/route/public_support.rs @@ -197,18 +197,22 @@ pub(super) fn classify_public_support_route( "public:auth", false, )) - } else if matches!(method, &http::Method::GET | &http::Method::POST) + // Authentication state-changing endpoints must never be dispatched for GET. + // Besides violating HTTP method semantics, accepting GET here would allow + // browser prefetches/cross-site requests to trigger login, refresh, logout, + // registration, or verification side effects. `/me` is the sole read route. + } else if (method == http::Method::POST && matches!( normalized_path, "/api/auth/login" | "/api/auth/refresh" | "/api/auth/register" - | "/api/auth/me" | "/api/auth/logout" | "/api/auth/send-verification-code" | "/api/auth/verify-email" | "/api/auth/verification-status" - ) + )) + || (method == http::Method::GET && normalized_path == "/api/auth/me") { let route_kind = match normalized_path { "/api/auth/login" => "login", @@ -520,6 +524,65 @@ pub(super) fn classify_public_support_route( "aether:ccswitch_usage", false, )) + } else if method == http::Method::POST + && matches!(normalized_path, "/api/vscodex/pair" | "/api/vscodex/pair/") + { + Some(classified( + "public_support", + "vscodex", + "pairing_exchange", + "public:vscodex", + false, + )) + } else if method == http::Method::GET + && matches!( + normalized_path, + "/api/users/me/vscodex/devices" | "/api/users/me/vscodex/devices/" + ) + { + Some(classified( + "public_support", + "users_me", + "vscodex_devices_list", + "user:self", + false, + )) + } else if method == http::Method::POST + && matches!( + normalized_path, + "/api/users/me/vscodex/pairings" | "/api/users/me/vscodex/pairings/" + ) + { + Some(classified( + "public_support", + "users_me", + "vscodex_pairing_create", + "user:self", + false, + )) + } else if method == http::Method::POST + && matches!( + normalized_path, + "/api/users/me/vscodex/ws-tickets" | "/api/users/me/vscodex/ws-tickets/" + ) + { + Some(classified( + "public_support", + "users_me", + "vscodex_ws_ticket_create", + "user:self", + false, + )) + } else if method == http::Method::DELETE + && has_single_segment_after_prefix(normalized_path, "/api/users/me/vscodex/devices/") + { + Some(classified( + "public_support", + "users_me", + "vscodex_device_delete", + "user:self", + false, + )) } else if method == http::Method::GET && matches!( normalized_path, diff --git a/apps/aether-gateway/src/control/tests/admin_core.rs b/apps/aether-gateway/src/control/tests/admin_core.rs index bd772a1e9..1b8808419 100644 --- a/apps/aether-gateway/src/control/tests/admin_core.rs +++ b/apps/aether-gateway/src/control/tests/admin_core.rs @@ -628,7 +628,7 @@ fn classifies_admin_management_token_write_routes_and_permission_catalog() { http::Method::POST, "/api/admin/management-tokens", "create_token", - "admin:management_tokens:write", + "admin:management_tokens:admin", ), ( http::Method::PUT, @@ -640,7 +640,7 @@ fn classifies_admin_management_token_write_routes_and_permission_catalog() { http::Method::POST, "/api/admin/management-tokens/token-123/regenerate", "regenerate_token", - "admin:management_tokens:write", + "admin:management_tokens:admin", ), ]; diff --git a/apps/aether-gateway/src/control/tests/admin_oauth.rs b/apps/aether-gateway/src/control/tests/admin_oauth.rs index 53d406a80..7131059dd 100644 --- a/apps/aether-gateway/src/control/tests/admin_oauth.rs +++ b/apps/aether-gateway/src/control/tests/admin_oauth.rs @@ -68,11 +68,11 @@ fn classifies_admin_provider_oauth_batch_import_task_status_as_admin_proxy_route ); assert_eq!( decision.auth_endpoint_signature.as_deref(), - Some("admin:pool") + Some("admin:provider_oauth") ); assert_eq!( management_token_required_permission(&http::Method::GET, &decision).as_deref(), - Some("admin:pool:read") + Some("admin:provider_oauth:read") ); assert!(!decision.is_execution_runtime_candidate()); } @@ -86,42 +86,42 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() { "/api/admin/provider-oauth/keys/key-123/complete", "complete_key_oauth", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/keys/key-123/refresh", "refresh_key_oauth", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/complete", "complete_provider_oauth", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/import-refresh-token", "import_refresh_token", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/cookie-authorize", "cookie_authorize", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/cookie-authorize/tasks", "start_cookie_authorize_task", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::GET, @@ -135,7 +135,7 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() { "/api/admin/provider-oauth/providers/provider-123/agent-identity-import/tasks", "start_agent_identity_import_task", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ( http::Method::GET, @@ -148,22 +148,22 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() { http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/batch-import", "batch_import_oauth", - "admin:pool", - "admin:pool:write", + "admin:provider_oauth", + "admin:provider_oauth:admin", ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/batch-import/tasks", "start_batch_import_oauth_task", - "admin:pool", - "admin:pool:write", + "admin:provider_oauth", + "admin:provider_oauth:admin", ), ( http::Method::GET, "/api/admin/provider-oauth/providers/provider-123/batch-import/tasks/task-123", "get_batch_import_task_status", - "admin:pool", - "admin:pool:read", + "admin:provider_oauth", + "admin:provider_oauth:read", ), ( http::Method::POST, @@ -177,7 +177,7 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() { "/api/admin/provider-oauth/providers/provider-123/device-poll", "device_poll", "admin:provider_oauth", - "admin:provider_oauth:write", + "admin:provider_oauth:admin", ), ] { let uri: Uri = path.parse().expect("uri should parse"); diff --git a/apps/aether-gateway/src/control/tests/admin_users.rs b/apps/aether-gateway/src/control/tests/admin_users.rs index 17831adfc..4dcac4598 100644 --- a/apps/aether-gateway/src/control/tests/admin_users.rs +++ b/apps/aether-gateway/src/control/tests/admin_users.rs @@ -103,6 +103,22 @@ fn classifies_admin_user_billing_routes_as_admin_proxy_route() { Some("admin:users") ); + let revoke_uri: Uri = "/api/admin/users/user-1/billing/entitlements/entitlement-1" + .parse() + .expect("uri should parse"); + let revoke = classify_control_route(&http::Method::DELETE, &revoke_uri, &headers) + .expect("route should classify"); + assert_eq!(revoke.route_class.as_deref(), Some("admin_proxy")); + assert_eq!(revoke.route_family.as_deref(), Some("users_manage")); + assert_eq!( + revoke.route_kind.as_deref(), + Some("revoke_user_billing_entitlement") + ); + assert_eq!( + revoke.auth_endpoint_signature.as_deref(), + Some("admin:users") + ); + let context = GatewayPublicRequestContext::from_request_parts( "trace-user-billing-grant", &http::Method::POST, diff --git a/apps/aether-gateway/src/control/tests/public_support.rs b/apps/aether-gateway/src/control/tests/public_support.rs index 078f47a2b..a5c1efa19 100644 --- a/apps/aether-gateway/src/control/tests/public_support.rs +++ b/apps/aether-gateway/src/control/tests/public_support.rs @@ -440,6 +440,26 @@ fn classifies_users_me_routes_as_public_support_route() { "/api/users/me/available-models", "available_models", ), + ( + http::Method::GET, + "/api/users/me/vscodex/devices", + "vscodex_devices_list", + ), + ( + http::Method::POST, + "/api/users/me/vscodex/pairings", + "vscodex_pairing_create", + ), + ( + http::Method::DELETE, + "/api/users/me/vscodex/devices/device-1", + "vscodex_device_delete", + ), + ( + http::Method::POST, + "/api/users/me/vscodex/ws-tickets", + "vscodex_ws_ticket_create", + ), ( http::Method::PUT, "/api/users/me/model-capabilities", @@ -496,6 +516,49 @@ fn classifies_users_me_routes_as_public_support_route() { } } +#[test] +fn vscodex_post_routes_buffer_request_body() { + let headers = headers(&[]); + for path in [ + "/api/vscodex/pair", + "/api/users/me/vscodex/pairings", + "/api/users/me/vscodex/ws-tickets", + ] { + let uri: Uri = path.parse().expect("uri should parse"); + let decision = classify_control_route(&http::Method::POST, &uri, &headers) + .expect("route should classify"); + let context = GatewayPublicRequestContext::from_request_parts( + "trace-vscodex", + &http::Method::POST, + &uri, + &headers, + Some(decision), + ); + + assert!( + local_proxy_route_requires_buffered_body(&context), + "{path} should buffer its JSON body" + ); + } +} + +#[test] +fn classifies_public_vscodex_pairing_exchange() { + let headers = headers(&[]); + let uri: Uri = "/api/vscodex/pair".parse().expect("uri should parse"); + let decision = + classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify"); + + assert_eq!(decision.route_class.as_deref(), Some("public_support")); + assert_eq!(decision.route_family.as_deref(), Some("vscodex")); + assert_eq!(decision.route_kind.as_deref(), Some("pairing_exchange")); + assert_eq!( + decision.auth_endpoint_signature.as_deref(), + Some("public:vscodex") + ); + assert!(!decision.is_execution_runtime_candidate()); +} + #[test] fn classifies_ccswitch_usage_as_api_key_public_support_route() { let headers = headers(&[]); @@ -884,6 +947,34 @@ fn classifies_auth_routes_as_public_support_route() { } } +#[test] +fn does_not_classify_state_changing_auth_routes_for_get() { + for path in [ + "/api/auth/login", + "/api/auth/refresh", + "/api/auth/register", + "/api/auth/logout", + "/api/auth/send-verification-code", + "/api/auth/verify-email", + "/api/auth/verification-status", + ] { + let headers = headers(&[]); + let uri: Uri = path.parse().expect("uri should parse"); + assert!( + classify_control_route(&http::Method::GET, &uri, &headers).is_none(), + "state-changing auth route {path} must not accept GET" + ); + } + + let headers = headers(&[]); + let uri: Uri = "/api/auth/me".parse().expect("uri should parse"); + assert_eq!( + classify_control_route(&http::Method::GET, &uri, &headers) + .and_then(|decision| decision.route_kind), + Some("me".to_string()) + ); +} + #[test] fn classifies_oauth_public_providers_route() { let headers = headers(&[]); diff --git a/apps/aether-gateway/src/data/decision_trace.rs b/apps/aether-gateway/src/data/decision_trace.rs index 84f91ae37..96a1b5ec7 100644 --- a/apps/aether-gateway/src/data/decision_trace.rs +++ b/apps/aether-gateway/src/data/decision_trace.rs @@ -172,7 +172,7 @@ mod tests { candidates: vec![DecisionTraceCandidate { candidate: sample_candidate("req-1"), provider_name: Some("OpenAI".to_string()), - provider_website: Some("https://openai.com".to_string()), + provider_website: Some("https://openai.com/".to_string()), provider_type: Some("custom".to_string()), provider_priority: Some(0), provider_keep_priority_on_conversion: Some(false), diff --git a/apps/aether-gateway/src/data/state/auth.rs b/apps/aether-gateway/src/data/state/auth.rs index be7308c0d..69280d5aa 100644 --- a/apps/aether-gateway/src/data/state/auth.rs +++ b/apps/aether-gateway/src/data/state/auth.rs @@ -1,19 +1,38 @@ use super::{ - AuthApiKeyLookupKey, CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot, - GatewayDataState, ManagementTokenCounterDelta, ManagementTokenListQuery, ProxyNodeCounterDelta, - ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, - ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, - ProxyNodeTunnelStatusMutation, RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord, - StoredAuthApiKeySnapshot, StoredLdapModuleConfig, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenWithUser, StoredOAuthProviderConfig, - StoredOAuthProviderModuleConfig, StoredProxyFleetMetricsBucket, StoredProxyNode, - StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, StoredUserAuthRecord, - StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, - StoredWalletSnapshot, UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord, + ActivateManagementTokenIfMatches, AuthApiKeyLookupKey, CompareAndSwapLdapConfigResult, + CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot, GatewayDataState, + InitializeAuthWalletOutcome, LdapBindPasswordUpdate, ManagementTokenCounterDelta, + ManagementTokenListQuery, ProxyNodeCounterDelta, ProxyNodeHeartbeatMutation, + ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeRegistrationMutation, + ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation, + RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, + StoredLdapModuleConfig, StoredManagementToken, StoredManagementTokenListPage, + StoredManagementTokenWithUser, StoredOAuthProviderConfig, StoredOAuthProviderModuleConfig, + StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent, + StoredProxyNodeMetricsBucket, StoredUserAuthRecord, StoredUserOAuthLinkSummary, + StoredUserPreferenceRecord, StoredUserSessionRecord, StoredWalletSnapshot, + UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord, }; use crate::LocalMutationOutcome; use aether_data::repository::auth::ResolvedAuthApiKeySnapshotReader; +/// Result of synchronizing an LDAP identity and, when applicable, creating its +/// first wallet. The wallet id is only exposed when this invocation created +/// the row, which gives callers a safe compensation token without exposing an +/// existing wallet as their own work. +pub(crate) struct LdapAuthProvisioningResult { + pub(crate) user: StoredUserAuthRecord, + pub(crate) owned_wallet_id: Option, +} + +fn auth_user_wallet_matches(wallet: &StoredWalletSnapshot, user_id: &str) -> bool { + wallet.user_id.as_deref() == Some(user_id) && wallet.api_key_id.is_none() +} +use aether_data::repository::users::ResolveOAuthLinkedUserOutcome; +use aether_data::repository::users::{ + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, +}; + #[derive(Debug, Clone, Default)] pub(crate) struct GatewayUserEffectiveListPolicies { pub(crate) allowed_providers: Option>, @@ -140,6 +159,21 @@ impl GatewayDataState { } } + pub(crate) async fn restore_user_group_if_matches( + &self, + expected: &aether_data::repository::users::StoredUserGroup, + restored: &aether_data::repository::users::StoredUserGroup, + ) -> Result { + match &self.user_reader { + Some(repository) => { + repository + .restore_user_group_if_matches(expected, restored) + .await + } + None => Ok(false), + } + } + pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result { match &self.user_reader { Some(repository) => repository.delete_user_group(group_id).await, @@ -212,6 +246,22 @@ impl GatewayDataState { } } + pub(crate) async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + match &self.user_reader { + Some(repository) => { + repository + .restore_user_groups_if_matches(user_id, expected_group_ids, restored_group_ids) + .await + } + None => Ok(false), + } + } + pub(crate) async fn add_user_to_group( &self, group_id: &str, @@ -246,6 +296,35 @@ impl GatewayDataState { .await } + #[allow(clippy::too_many_arguments)] + pub(crate) async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: chrono::DateTime, + provider_enabled_snapshot: bool, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(ResolveOAuthLinkedUserOutcome::ProviderUnavailable); + }; + repository + .resolve_enabled_oauth_linked_user( + provider_type, + provider_user_id, + provider_username, + provider_email, + extra_data, + verified_email, + touched_at, + provider_enabled_snapshot, + ) + .await + } + pub(crate) async fn touch_oauth_link( &self, provider_type: &str, @@ -273,6 +352,7 @@ impl GatewayDataState { pub(crate) async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: chrono::DateTime, ) -> Result, DataLayerError> { @@ -280,7 +360,7 @@ impl GatewayDataState { return Ok(None); }; repository - .create_oauth_auth_user(email, username, created_at) + .create_oauth_auth_user(email, email_verified, username, created_at) .await } @@ -320,8 +400,18 @@ impl GatewayDataState { repository.count_user_oauth_links(user_id).await } + pub(crate) async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository.has_oauth_links_for_provider(provider_type).await + } + #[allow(clippy::too_many_arguments)] - pub(crate) async fn upsert_user_oauth_link( + pub(crate) async fn bind_user_oauth_link( &self, user_id: &str, provider_type: &str, @@ -330,12 +420,14 @@ impl GatewayDataState { provider_email: Option<&str>, extra_data: Option, linked_at: chrono::DateTime, - ) -> Result<(), DataLayerError> { + provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result { let Some(repository) = self.user_reader.as_ref() else { - return Ok(()); + return Ok(BindUserOAuthLinkOutcome::UserNotFound); }; repository - .upsert_user_oauth_link( + .bind_user_oauth_link_if_provider_enabled( user_id, provider_type, provider_user_id, @@ -343,20 +435,43 @@ impl GatewayDataState { provider_email, extra_data, linked_at, + provider_enabled_snapshot, + session_expectation, ) .await } + pub(crate) async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: chrono::DateTime, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .upgrade_oauth_email_verification_if_matches(user_id, verified_email, verified_at) + .await + } + pub(crate) async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { + local_password_login_allowed: bool, + enabled_provider_types_snapshot: &[String], + ) -> Result { let Some(repository) = self.user_reader.as_ref() else { - return Ok(false); + return Ok(DeleteUserOAuthLinkOutcome::NotFound); }; repository - .delete_user_oauth_link(user_id, provider_type) + .delete_user_oauth_link( + user_id, + provider_type, + local_password_login_allowed, + enabled_provider_types_snapshot, + ) .await } @@ -438,6 +553,19 @@ impl GatewayDataState { repository.create_user_session(session).await } + pub(crate) async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(None); + }; + repository + .create_user_session_if_password_matches(session, expected_password_hash) + .await + } + pub(crate) async fn update_user_model_capability_settings( &self, user_id: &str, @@ -467,14 +595,45 @@ impl GatewayDataState { pub(crate) async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { let Some(repository) = self.user_reader.as_ref() else { return Ok(None); }; repository - .update_local_auth_user_profile(user_id, email, username) + .update_local_auth_user_profile(user_id, email_present, email, email_verified, username) + .await + } + + #[allow(clippy::too_many_arguments)] + pub(crate) async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &aether_data::repository::users::StoredUserAuthRecord, + restored_auth: &aether_data::repository::users::StoredUserAuthRecord, + expected_export: &aether_data::repository::users::StoredUserExportRow, + restored_export: &aether_data::repository::users::StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .restore_local_auth_user_state_if_matches( + expected_auth, + restored_auth, + expected_export, + restored_export, + expected_model_capability_settings, + restored_model_capability_settings, + expected_feature_settings, + restored_feature_settings, + ) .await } @@ -492,6 +651,62 @@ impl GatewayDataState { .await } + pub(crate) async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: chrono::DateTime, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .restore_local_auth_user_password_hash_if_matches( + user_id, + expected_password_hash, + password_hash, + updated_at, + ) + .await + } + + pub(crate) async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .reset_local_auth_user_password_and_revoke_sessions(user_id, password_hash, changed_at) + .await + } + + pub(crate) async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .change_local_auth_password_and_revoke_sessions( + user_id, + current_session_id, + expected_password_hash, + next_password_hash, + changed_at, + ) + .await + } + #[allow(dead_code)] pub(crate) async fn create_local_auth_user( &self, @@ -620,6 +835,30 @@ impl GatewayDataState { initial_gift_usd: f64, unlimited: bool, ) -> Result, DataLayerError> { + Ok(self + .get_or_create_ldap_auth_user_with_wallet_outcome( + email, + username, + ldap_dn, + ldap_username, + logged_in_at, + initial_gift_usd, + unlimited, + ) + .await? + .map(|result| result.user)) + } + + pub(crate) async fn get_or_create_ldap_auth_user_with_wallet_outcome( + &self, + email: String, + username: String, + ldap_dn: Option, + ldap_username: Option, + logged_in_at: chrono::DateTime, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { let Some(repository) = self.user_reader.as_ref() else { return Ok(None); }; @@ -629,23 +868,54 @@ impl GatewayDataState { else { return Ok(None); }; - if outcome.created { - match self - .initialize_auth_user_wallet(&outcome.user.id, initial_gift_usd, unlimited) - .await - { - Ok(Some(_wallet)) => {} - Ok(None) => { - let _ = self.delete_local_auth_user(&outcome.user.id).await; - return Ok(None); - } - Err(err) => { - let _ = self.delete_local_auth_user(&outcome.user.id).await; - return Err(err); - } - } + if !outcome.created { + return Ok(Some(LdapAuthProvisioningResult { + user: outcome.user, + owned_wallet_id: None, + })); } - Ok(Some(outcome.user)) + + let initialized = match self + .initialize_auth_user_wallet_with_outcome(&outcome.user.id, initial_gift_usd, unlimited) + .await + { + Ok(Some(initialized)) => initialized, + Ok(None) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&outcome.user.id, None) + .await; + return Ok(None); + } + Err(err) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&outcome.user.id, None) + .await; + return Err(err); + } + }; + + // A user wallet initializer must return a user-owned wallet. If a + // custom/legacy backend violates that contract, only remove the wallet + // when this invocation actually created it; an existing wallet must be + // preserved and the user rollback must fail closed. + let wallet_is_user_owned = auth_user_wallet_matches(&initialized.wallet, &outcome.user.id); + if !wallet_is_user_owned { + let owned_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let _ = self + .rollback_provisional_auth_user_with_wallet( + &outcome.user.id, + owned_wallet_id.as_deref(), + ) + .await; + return Err(DataLayerError::UnexpectedValue( + "LDAP user wallet owner does not match the provisioned user".to_string(), + )); + } + + Ok(Some(LdapAuthProvisioningResult { + user: outcome.user, + owned_wallet_id: initialized.created.then_some(initialized.wallet.id), + })) } #[allow(dead_code)] @@ -663,6 +933,20 @@ impl GatewayDataState { .await } + pub(crate) async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + let Some(repository) = self.wallet_reader.as_ref() else { + return Ok(None); + }; + repository + .initialize_auth_user_wallet_with_outcome(user_id, initial_gift_usd, unlimited) + .await + } + pub(crate) async fn initialize_auth_api_key_wallet( &self, api_key_id: &str, @@ -677,6 +961,81 @@ impl GatewayDataState { .await } + pub(crate) async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + let Some(repository) = self.wallet_reader.as_ref() else { + return Ok(None); + }; + repository + .initialize_auth_api_key_wallet_with_outcome(api_key_id, initial_gift_usd, unlimited) + .await + } + + pub(crate) async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result { + match &self.wallet_writer { + Some(repository) => { + repository + .delete_provisional_auth_user_wallet(wallet_id, user_id) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + match &self.wallet_writer { + Some(repository) => { + repository + .delete_wallet_if_unreferenced(wallet_id, owner) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + match &self.wallet_writer { + Some(repository) => { + repository + .delete_wallet_if_snapshot_matches_and_unreferenced(expected, owner) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + match &self.wallet_writer { + Some(repository) => { + repository + .restore_wallet_if_snapshot_matches(before, after, owner) + .await + } + None => Ok(false), + } + } + pub(crate) async fn update_auth_user_wallet_limit_mode( &self, user_id: &str, @@ -823,6 +1182,222 @@ impl GatewayDataState { repository.delete_local_auth_user(user_id).await } + pub(crate) async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + let Some(repository) = self.user_reader.as_ref() else { + return Ok(false); + }; + repository + .delete_local_auth_user_if_wallet_absent(user_id) + .await + } + + pub(crate) async fn rollback_provisional_auth_user( + &self, + user_id: &str, + ) -> Result { + self.rollback_provisional_auth_user_with_wallet(user_id, None) + .await + } + + /// Check for every wallet that can still keep a user-owned account alive. + /// + /// API-key wallets do not carry the owning user id themselves, so a direct + /// `find(UserId(..))` is insufficient. When an auth-key reader is + /// available we resolve the user's non-standalone key ids first. A + /// reader-less in-memory repository cannot distinguish those keys from + /// standalone keys; in that case any API-key wallet is treated as a + /// blocking reference (fail closed). + async fn wallet_exists_for_user_or_api_key( + &self, + user_id: &str, + ) -> Result { + let Some(repository) = self.wallet_reader.as_ref() else { + // A writer without a reader cannot prove wallet absence. Preserve + // the existing conservative behavior used by rollback callers. + return Ok(self.wallet_writer.is_some()); + }; + + if repository + .find(aether_data::repository::wallet::WalletLookupKey::UserId( + user_id, + )) + .await? + .is_some() + { + return Ok(true); + } + + if self.has_auth_api_key_reader() { + let user_ids = vec![user_id.to_string()]; + let api_key_ids = self + .list_auth_api_key_export_records_by_user_ids(&user_ids) + .await? + .into_iter() + .map(|record| record.api_key_id) + .collect::>(); + if api_key_ids.is_empty() { + return Ok(false); + } + return Ok(!repository + .list_wallets_by_api_key_ids(&api_key_ids) + .await? + .is_empty()); + } + + // Without an auth-key reader, inspect the wallet owner type directly. + // A single-row page is enough to establish that an API-key wallet + // exists while avoiding an unbounded read during error compensation. + let page = repository + .list_admin_wallets(&aether_data::repository::wallet::AdminWalletListQuery { + status: None, + owner_type: Some("api_key".to_string()), + limit: 1, + offset: 0, + }) + .await?; + Ok(page.total > 0) + } + + /// Compensate a user provision only when the caller can prove ownership of the wallet it + /// created. A missing wallet id is intentionally fail-closed: a concurrent initializer may + /// have created a valid wallet after this operation started, and an owner-only structural + /// lookup would then be able to delete that wallet. + pub(crate) async fn rollback_provisional_auth_user_with_wallet( + &self, + user_id: &str, + wallet_id: Option<&str>, + ) -> Result { + if user_id.trim().is_empty() { + return Ok(false); + } + if wallet_id.is_some_and(|wallet_id| wallet_id.trim().is_empty()) { + return Err(DataLayerError::InvalidInput( + "wallet compensation wallet id cannot be empty".to_string(), + )); + } + + let database_backend_has_atomic_guard = self.backends.is_some(); + let wallet_removed = if let Some(wallet_id) = wallet_id { + let removed = self + .delete_provisional_auth_user_wallet(wallet_id, user_id) + .await?; + if !removed { + // The caller supplied an ownership token for a specific + // wallet. If that row still exists, an owner-only lookup is + // not enough to justify deleting the user: the row may be + // funded, attached to another owner, or simply ineligible + // for provisional cleanup. + let exact_wallet_exists = match &self.wallet_reader { + Some(repository) => repository + .find(aether_data::repository::wallet::WalletLookupKey::WalletId( + wallet_id, + )) + .await? + .is_some(), + None => self.wallet_writer.is_some(), + }; + if exact_wallet_exists { + return Err(DataLayerError::UnexpectedValue(format!( + "refusing to delete provisional auth user {user_id}: supplied wallet still exists" + ))); + } + let wallet_exists = self.wallet_exists_for_user_or_api_key(user_id).await?; + if wallet_exists { + return Err(DataLayerError::UnexpectedValue(format!( + "refusing to delete provisional auth user {user_id}: wallet is not eligible for rollback" + ))); + } + } + // Removing the supplied wallet does not prove that it was the + // user's only financial reference. Check again before falling + // through to the user deletion path when no SQL atomic guard is + // available (notably in-memory/test repositories). + if removed && !database_backend_has_atomic_guard { + if self.wallet_exists_for_user_or_api_key(user_id).await? { + return Err(DataLayerError::UnexpectedValue(format!( + "refusing to delete provisional auth user {user_id}: another wallet reference exists" + ))); + } + } + removed + } else if database_backend_has_atomic_guard { + // The SQL user repository performs the wallet-absence check in the + // same transaction as deletion. A separate read here would + // reintroduce the provisioning TOCTOU race. + false + } else { + let wallet_exists = self.wallet_exists_for_user_or_api_key(user_id).await?; + if wallet_exists { + return Err(DataLayerError::UnexpectedValue(format!( + "refusing to delete provisional auth user {user_id}: wallet ownership is unknown" + ))); + } + false + }; + let _ = wallet_removed; + if database_backend_has_atomic_guard { + self.delete_local_auth_user_if_wallet_absent(user_id).await + } else { + self.delete_local_auth_user(user_id).await + } + } + + pub(crate) async fn register_local_auth_user_with_wallet_outcome( + &self, + email: Option, + email_verified: bool, + username: String, + password_hash: String, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + let Some(user) = self + .create_local_auth_user(email, email_verified, username, password_hash) + .await? + else { + return Ok(None); + }; + + let initialized = match self + .initialize_auth_user_wallet_with_outcome(&user.id, initial_gift_usd, unlimited) + .await + { + Ok(Some(initialized)) => initialized, + Ok(None) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; + return Ok(None); + } + Err(err) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; + return Err(err); + } + }; + + // A local account may only be paired with a user-owned wallet. Keep + // this contract check at the data boundary as well as in the LDAP + // path: custom or legacy repositories must not be able to hand a + // caller an API-key wallet (or another user's wallet) as its new + // account balance. + let wallet_is_user_owned = auth_user_wallet_matches(&initialized.wallet, &user.id); + if !wallet_is_user_owned { + let owned_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let _ = self + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await; + return Err(DataLayerError::UnexpectedValue( + "local user wallet owner does not match the provisioned user".to_string(), + )); + } + Ok(Some((user, initialized.wallet, initialized.created))) + } + pub(crate) async fn register_local_auth_user( &self, email: Option, @@ -832,27 +1407,17 @@ impl GatewayDataState { initial_gift_usd: f64, unlimited: bool, ) -> Result, DataLayerError> { - let Some(user) = self - .create_local_auth_user(email, email_verified, username, password_hash) + Ok(self + .register_local_auth_user_with_wallet_outcome( + email, + email_verified, + username, + password_hash, + initial_gift_usd, + unlimited, + ) .await? - else { - return Ok(None); - }; - - match self - .initialize_auth_user_wallet(&user.id, initial_gift_usd, unlimited) - .await - { - Ok(Some(wallet)) => Ok(Some((user, wallet))), - Ok(None) => { - let _ = self.delete_local_auth_user(&user.id).await; - Ok(None) - } - Err(err) => { - let _ = self.delete_local_auth_user(&user.id).await; - Err(err) - } - } + .map(|(user, wallet, _created)| (user, wallet))) } pub(crate) async fn touch_user_session( @@ -891,7 +1456,7 @@ impl GatewayDataState { &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: chrono::DateTime, expires_at: chrono::DateTime, @@ -905,7 +1470,7 @@ impl GatewayDataState { .rotate_user_session_refresh_token( user_id, session_id, - previous_refresh_token_hash, + expected_refresh_token_hash, next_refresh_token_hash, rotated_at, expires_at, @@ -962,16 +1527,46 @@ impl GatewayDataState { } } - pub(crate) async fn upsert_ldap_module_config( + pub(crate) async fn compare_and_swap_ldap_module_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, DataLayerError> { + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result, DataLayerError> { match &self.auth_module_writer { - Some(repository) => repository.upsert_ldap_config(config).await, + Some(repository) => repository + .compare_and_swap_ldap_config(expected, replacement, bind_password_update) + .await + .map(Some), None => Ok(None), } } + pub(crate) async fn delete_ldap_module_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result { + match &self.auth_module_writer { + Some(repository) => repository.delete_ldap_config_if_matches(expected).await, + None => Ok(false), + } + } + + pub(crate) async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + match &self.auth_module_writer { + Some(repository) => { + repository + .compare_and_swap_ldap_bind_password(expected, replacement) + .await + } + None => Ok(false), + } + } + pub(crate) async fn list_oauth_provider_configs( &self, ) -> Result, DataLayerError> { @@ -991,40 +1586,99 @@ impl GatewayDataState { } } + pub(crate) async fn compare_and_swap_oauth_provider_client_secret( + &self, + provider_type: &str, + expected: &str, + replacement: &str, + ) -> Result { + match &self.oauth_provider_writer { + Some(repository) => { + repository + .compare_and_swap_oauth_provider_client_secret( + provider_type, + expected, + replacement, + ) + .await + } + None => Ok(false), + } + } + pub(crate) async fn count_locked_users_if_oauth_provider_disabled( &self, provider_type: &str, ldap_exclusive: bool, ) -> Result { - match &self.oauth_provider_reader { + let repository_count = match &self.oauth_provider_reader { Some(repository) => { repository .count_locked_users_if_provider_disabled(provider_type, ldap_exclusive) .await } None => Ok(0), - } + }?; + let enabled_provider_types = match &self.oauth_provider_reader { + Some(repository) => repository + .list_oauth_provider_configs() + .await? + .into_iter() + .filter(|provider| provider.is_enabled) + .map(|provider| provider.provider_type) + .collect::>(), + None => Vec::new(), + }; + let user_count = match &self.user_reader { + Some(repository) => { + repository + .count_locked_users_if_oauth_provider_disabled( + provider_type, + &enabled_provider_types, + ldap_exclusive, + ) + .await? + } + None => 0, + }; + Ok(repository_count.max(user_count)) } pub(crate) async fn upsert_oauth_provider_config( &self, record: &UpsertOAuthProviderConfigRecord, - ) -> Result, DataLayerError> { + ldap_exclusive: bool, + force_disable: bool, + locked_users_snapshot: usize, + ) -> Result< + Option, + DataLayerError, + > { match &self.oauth_provider_writer { Some(repository) => repository - .upsert_oauth_provider_config(record) + .upsert_oauth_provider_config_guarded( + record, + ldap_exclusive, + force_disable, + locked_users_snapshot, + ) .await .map(Some), None => Ok(None), } } - pub(crate) async fn delete_oauth_provider_config( + pub(crate) async fn delete_oauth_provider_config_if_unlinked( &self, provider_type: &str, ) -> Result { + let has_links_snapshot = self.has_oauth_links_for_provider(provider_type).await?; match &self.oauth_provider_writer { - Some(repository) => repository.delete_oauth_provider_config(provider_type).await, + Some(repository) => { + repository + .delete_oauth_provider_config_if_unlinked(provider_type, has_links_snapshot) + .await + } None => Ok(false), } } @@ -1099,6 +1753,27 @@ impl GatewayDataState { } } + pub(crate) async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, DataLayerError> { + match &self.management_token_writer { + Some(repository) => match repository + .update_management_token_for_user(record, user_id) + .await + { + Ok(Some(token)) => Ok(LocalMutationOutcome::Applied(token)), + Ok(None) => Ok(LocalMutationOutcome::NotFound), + Err(DataLayerError::InvalidInput(detail)) => { + Ok(LocalMutationOutcome::Invalid(detail)) + } + Err(err) => Err(err), + }, + None => Ok(LocalMutationOutcome::Unavailable), + } + } + pub(crate) async fn delete_management_token( &self, token_id: &str, @@ -1109,6 +1784,21 @@ impl GatewayDataState { } } + pub(crate) async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + match &self.management_token_writer { + Some(repository) => { + repository + .delete_management_token_for_user(token_id, user_id) + .await + } + None => Ok(false), + } + } + pub(crate) async fn record_management_token_usage( &self, token_id: &str, @@ -1226,6 +1916,38 @@ impl GatewayDataState { } } + pub(crate) async fn compare_and_set_proxy_node_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result { + match &self.proxy_node_writer { + Some(repository) => { + repository + .compare_and_set_proxy_password(node_id, expected, replacement) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn compare_and_set_proxy_node_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result { + match &self.proxy_node_writer { + Some(repository) => { + repository + .compare_and_set_proxy_metadata(node_id, expected, replacement) + .await + } + None => Ok(false), + } + } + pub(crate) async fn create_manual_proxy_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -1293,6 +2015,7 @@ impl GatewayDataState { let enqueued = repository .enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta { node_id: mutation.node_id.clone(), + expected_tunnel_generation: mutation.expected_tunnel_generation.clone(), total_requests_delta: mutation.total_requests_delta, failed_requests_delta: mutation.failed_requests_delta, dns_failures_delta: mutation.dns_failures_delta, @@ -1365,6 +2088,50 @@ impl GatewayDataState { } } + pub(crate) async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + match &self.management_token_writer { + Some(repository) => { + repository + .set_management_token_active_for_user(token_id, user_id, is_active) + .await + } + None => Ok(None), + } + } + + pub(crate) async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + match &self.management_token_writer { + Some(repository) => { + repository + .activate_management_token_if_matches(mutation) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + match &self.management_token_writer { + Some(repository) => { + repository + .delete_inactive_management_token_if_matches(mutation) + .await + } + None => Ok(false), + } + } + pub(crate) async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, @@ -1385,6 +2152,27 @@ impl GatewayDataState { } } + pub(crate) async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, DataLayerError> { + match &self.management_token_writer { + Some(repository) => match repository + .regenerate_management_token_secret_for_user(mutation, user_id) + .await + { + Ok(Some(token)) => Ok(LocalMutationOutcome::Applied(token)), + Ok(None) => Ok(LocalMutationOutcome::NotFound), + Err(DataLayerError::InvalidInput(detail)) => { + Ok(LocalMutationOutcome::Invalid(detail)) + } + Err(err) => Err(err), + }, + None => Ok(LocalMutationOutcome::Unavailable), + } + } + pub(in crate::data) async fn find_auth_api_key_snapshot( &self, key: AuthApiKeyLookupKey<'_>, @@ -1556,6 +2344,21 @@ impl GatewayDataState { } } + #[cfg(test)] + pub(crate) async fn synchronize_user_api_key_owner_for_tests( + &self, + user: &StoredUserAuthRecord, + ) -> Result<(), DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .synchronize_user_api_key_owner_for_tests(user) + .await + } + None => Ok(()), + } + } + pub(crate) async fn create_standalone_api_key( &self, record: aether_data::repository::auth::CreateStandaloneApiKeyRecord, @@ -1576,6 +2379,34 @@ impl GatewayDataState { } } + pub(crate) async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &aether_data::repository::auth::CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .compare_and_swap_api_key_ciphertext(mutation) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn update_user_api_key_basic_if_unlocked( + &self, + record: aether_data::repository::auth::UpdateUserApiKeyBasicRecord, + ) -> Result, DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .update_user_api_key_basic_if_unlocked(record) + .await + } + None => Ok(None), + } + } + pub(crate) async fn update_standalone_api_key_basic( &self, record: aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord, @@ -1586,6 +2417,21 @@ impl GatewayDataState { } } + pub(crate) async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .restore_api_key_if_matches(expected, restored) + .await + } + None => Ok(false), + } + } + pub(crate) async fn set_user_api_key_active( &self, user_id: &str, @@ -1602,6 +2448,22 @@ impl GatewayDataState { } } + pub(crate) async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .set_user_api_key_active_if_unlocked(user_id, api_key_id, is_active) + .await + } + None => Ok(None), + } + } + pub(crate) async fn set_standalone_api_key_active( &self, api_key_id: &str, @@ -1649,6 +2511,26 @@ impl GatewayDataState { } } + pub(crate) async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .set_user_api_key_allowed_providers_if_unlocked( + user_id, + api_key_id, + allowed_providers, + ) + .await + } + None => Ok(None), + } + } + pub(crate) async fn set_user_api_key_force_capabilities( &self, user_id: &str, @@ -1665,6 +2547,26 @@ impl GatewayDataState { } } + pub(crate) async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .set_user_api_key_force_capabilities_if_unlocked( + user_id, + api_key_id, + force_capabilities, + ) + .await + } + None => Ok(None), + } + } + pub(crate) async fn set_user_api_key_feature_settings( &self, user_id: &str, @@ -1681,6 +2583,26 @@ impl GatewayDataState { } } + pub(crate) async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, DataLayerError> { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .set_user_api_key_feature_settings_if_unlocked( + user_id, + api_key_id, + feature_settings, + ) + .await + } + None => Ok(None), + } + } + pub(crate) async fn set_api_key_usage_totals( &self, api_key_id: &str, @@ -1729,6 +2651,21 @@ impl GatewayDataState { } } + pub(crate) async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + match &self.auth_api_key_writer { + Some(repository) => { + repository + .delete_user_api_key_if_unlocked(user_id, api_key_id) + .await + } + None => Ok(false), + } + } + pub(crate) async fn delete_standalone_api_key( &self, api_key_id: &str, @@ -2130,6 +3067,9 @@ mod tests { InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserGroup, UpsertUserGroupRecord, UserReadRepository, }; + use aether_data::repository::wallet::{ + InMemoryWalletRepository, WalletLookupKey, WalletReadRepository, + }; use crate::data::GatewayDataState; @@ -2137,6 +3077,332 @@ mod tests { sample_snapshot_with_role(api_key_id, user_id, "user") } + #[test] + fn auth_user_wallet_owner_contract_rejects_cross_owner_rows() { + let user_wallet = StoredWalletSnapshot::new( + "user-wallet".to_string(), + Some("user-1".to_string()), + None, + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + 1, + ) + .expect("wallet should build"); + let api_key_wallet = StoredWalletSnapshot::new( + "api-key-wallet".to_string(), + None, + Some("key-1".to_string()), + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + 1, + ) + .expect("wallet should build"); + let other_user_wallet = StoredWalletSnapshot::new( + "other-user-wallet".to_string(), + Some("user-2".to_string()), + None, + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + 1, + ) + .expect("wallet should build"); + + assert!(auth_user_wallet_matches(&user_wallet, "user-1")); + assert!(!auth_user_wallet_matches(&user_wallet, "user-2")); + assert!(!auth_user_wallet_matches(&api_key_wallet, "user-1")); + assert!(!auth_user_wallet_matches(&other_user_wallet, "user-1")); + } + + #[tokio::test] + async fn provisioning_rollback_removes_initial_wallet_before_user() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "provisional-user".to_string(), + Some("provisional@example.com".to_string()), + true, + "provisional".to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user])); + let wallet_repository = Arc::new(InMemoryWalletRepository::default()); + let wallet = wallet_repository + .initialize_auth_user_wallet("provisional-user", 10.0, false) + .await + .expect("wallet should initialize") + .expect("wallet should exist"); + let state = GatewayDataState::with_user_and_wallet_for_tests( + user_repository.clone(), + wallet_repository.clone(), + ); + + assert!(state + .rollback_provisional_auth_user_with_wallet( + "provisional-user", + Some(wallet.id.as_str()), + ) + .await + .expect("rollback should succeed")); + assert!(user_repository + .find_user_auth_by_id("provisional-user") + .await + .expect("user lookup should succeed") + .is_none()); + assert!(wallet_repository + .find(WalletLookupKey::UserId("provisional-user")) + .await + .expect("wallet lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn provisioning_rollback_preserves_user_when_wallet_has_financial_activity() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "active-user".to_string(), + Some("active@example.com".to_string()), + true, + "active-user".to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let wallet = StoredWalletSnapshot::new( + "active-wallet".to_string(), + Some("active-user".to_string()), + None, + 1.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 1.0, + 0.0, + 0.0, + 0.0, + now.timestamp(), + ) + .expect("wallet should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user])); + let wallet_repository = Arc::new(InMemoryWalletRepository::seed([wallet])); + let state = GatewayDataState::with_user_and_wallet_for_tests( + user_repository.clone(), + wallet_repository.clone(), + ); + + let error = state + .rollback_provisional_auth_user_with_wallet("active-user", Some("active-wallet")) + .await + .expect_err("financial activity must block provisioning rollback"); + assert!(matches!(error, DataLayerError::UnexpectedValue(_))); + assert!(user_repository + .find_user_auth_by_id("active-user") + .await + .expect("user lookup should succeed") + .is_some()); + assert!(wallet_repository + .find(WalletLookupKey::UserId("active-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn provisioning_rollback_removes_user_when_wallet_is_confirmed_absent() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "no-wallet-user".to_string(), + Some("no-wallet@example.com".to_string()), + true, + "no-wallet-user".to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user])); + let wallet_repository = Arc::new(InMemoryWalletRepository::default()); + let state = GatewayDataState::with_user_and_wallet_for_tests( + user_repository.clone(), + wallet_repository, + ); + + assert!(state + .rollback_provisional_auth_user("no-wallet-user") + .await + .expect("confirmed wallet absence should allow rollback")); + assert!(user_repository + .find_user_auth_by_id("no-wallet-user") + .await + .expect("user lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn provisioning_rollback_preserves_user_when_api_key_wallet_exists() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "api-key-wallet-user".to_string(), + Some("api-key-wallet@example.com".to_string()), + true, + "api-key-wallet-user".to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let api_key_wallet = StoredWalletSnapshot::new( + "api-key-wallet-row".to_string(), + None, + Some("api-key-wallet-id".to_string()), + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + now.timestamp(), + ) + .expect("wallet should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user])); + let wallet_repository = Arc::new(InMemoryWalletRepository::seed([api_key_wallet])); + let state = GatewayDataState::with_user_and_wallet_for_tests( + user_repository.clone(), + wallet_repository.clone(), + ); + + let error = state + .rollback_provisional_auth_user_with_wallet("api-key-wallet-user", None) + .await + .expect_err("an API-key wallet must block user rollback"); + assert!(matches!(error, DataLayerError::UnexpectedValue(_))); + assert!(user_repository + .find_user_auth_by_id("api-key-wallet-user") + .await + .expect("user lookup should succeed") + .is_some()); + assert!(wallet_repository + .find(WalletLookupKey::ApiKeyId("api-key-wallet-id")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn provisioning_rollback_rejects_a_wallet_id_owned_by_another_user() { + let now = chrono::Utc::now(); + let target_user = StoredUserAuthRecord::new( + "target-user".to_string(), + Some("target@example.com".to_string()), + true, + "target-user".to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("target user should build"); + let other_user_wallet = StoredWalletSnapshot::new( + "other-user-wallet".to_string(), + Some("other-user".to_string()), + None, + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + now.timestamp(), + ) + .expect("wallet should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([target_user])); + let wallet_repository = Arc::new(InMemoryWalletRepository::seed([other_user_wallet])); + let state = GatewayDataState::with_user_and_wallet_for_tests( + user_repository.clone(), + wallet_repository.clone(), + ); + + let error = state + .rollback_provisional_auth_user_with_wallet("target-user", Some("other-user-wallet")) + .await + .expect_err("a live supplied wallet id must block user deletion"); + assert!(matches!(error, DataLayerError::UnexpectedValue(_))); + assert!(user_repository + .find_user_auth_by_id("target-user") + .await + .expect("user lookup should succeed") + .is_some()); + assert!(wallet_repository + .find(WalletLookupKey::WalletId("other-user-wallet")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + fn sample_snapshot_with_role( api_key_id: &str, user_id: &str, diff --git a/apps/aether-gateway/src/data/state/catalog.rs b/apps/aether-gateway/src/data/state/catalog.rs index 6e1aa6dbf..ca1f0ccc6 100644 --- a/apps/aether-gateway/src/data/state/catalog.rs +++ b/apps/aether-gateway/src/data/state/catalog.rs @@ -1,23 +1,42 @@ use super::{ ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery, GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate, - ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate, - ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, - ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, - ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket, - StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, - StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, - StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, - StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord, + ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate, + ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery, + ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, + ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, + ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, PublicHealthStatusCount, + PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage, + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, + StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate, + UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord, }; +fn sanitize_request_candidate_rows( + mut candidates: Vec, +) -> Vec { + for candidate in &mut candidates { + candidate.sanitize_sensitive_diagnostics(); + } + candidates +} + +fn sanitize_request_candidate_row(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate { + candidate.sanitize_sensitive_diagnostics(); + candidate +} + impl GatewayDataState { pub(crate) async fn list_request_candidates_by_request_id( &self, request_id: &str, ) -> Result, DataLayerError> { match &self.request_candidate_reader { - Some(repository) => repository.list_by_request_id(request_id).await, + Some(repository) => repository + .list_by_request_id(request_id) + .await + .map(sanitize_request_candidate_rows), None => Ok(Vec::new()), } } @@ -27,7 +46,10 @@ impl GatewayDataState { request_id: &str, ) -> Result, DataLayerError> { match &self.request_candidate_reader { - Some(repository) => repository.list_attempted_by_request_id(request_id).await, + Some(repository) => repository + .list_attempted_by_request_id(request_id) + .await + .map(sanitize_request_candidate_rows), None => Ok(Vec::new()), } } @@ -38,7 +60,10 @@ impl GatewayDataState { limit: usize, ) -> Result, DataLayerError> { match &self.request_candidate_reader { - Some(repository) => repository.list_by_provider_id(provider_id, limit).await, + Some(repository) => repository + .list_by_provider_id(provider_id, limit) + .await + .map(sanitize_request_candidate_rows), None => Ok(Vec::new()), } } @@ -48,7 +73,10 @@ impl GatewayDataState { limit: usize, ) -> Result, DataLayerError> { match &self.request_candidate_reader { - Some(repository) => repository.list_recent(limit).await, + Some(repository) => repository + .list_recent(limit) + .await + .map(sanitize_request_candidate_rows), None => Ok(Vec::new()), } } @@ -60,11 +88,10 @@ impl GatewayDataState { limit: usize, ) -> Result, DataLayerError> { match &self.request_candidate_reader { - Some(repository) => { - repository - .list_finalized_by_endpoint_ids_since(endpoint_ids, since_unix_secs, limit) - .await - } + Some(repository) => repository + .list_finalized_by_endpoint_ids_since(endpoint_ids, since_unix_secs, limit) + .await + .map(sanitize_request_candidate_rows), None => Ok(Vec::new()), } } @@ -108,14 +135,19 @@ impl GatewayDataState { pub(crate) async fn upsert_request_candidate( &self, - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, ) -> Result, DataLayerError> { + candidate.sanitize_for_persistence(); crate::request_diagnostics::observe_db_operation( "request_candidate_upsert", self.database_pool_summary(), async { match &self.request_candidate_writer { - Some(repository) => repository.upsert(candidate).await.map(Some), + Some(repository) => repository + .upsert(candidate) + .await + .map(sanitize_request_candidate_row) + .map(Some), None => Ok(None), } }, @@ -170,6 +202,16 @@ impl GatewayDataState { } } + pub(crate) async fn upsert_gemini_file_mapping_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, DataLayerError> { + match &self.gemini_file_mapping_writer { + Some(repository) => repository.upsert_if_owner_matches(record).await, + None => Ok(None), + } + } + pub(crate) async fn list_gemini_file_mappings( &self, query: &GeminiFileMappingListQuery, @@ -183,6 +225,49 @@ impl GatewayDataState { } } + pub(crate) async fn find_gemini_file_mapping_by_file_name( + &self, + file_name: &str, + ) -> Result, DataLayerError> { + match &self.gemini_file_mapping_reader { + Some(repository) => repository.find_by_file_name(file_name).await, + None => Ok(None), + } + } + + pub(crate) async fn find_active_gemini_file_mapping_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + match &self.gemini_file_mapping_reader { + Some(repository) => { + repository + .find_active_by_file_name_for_user(file_name, user_id, now_unix_secs) + .await + } + None => Ok(None), + } + } + + pub(crate) async fn find_active_gemini_file_mapping_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + match &self.gemini_file_mapping_reader { + Some(repository) => { + repository + .find_active_by_file_name_for_owner(file_name, key_id, user_id, now_unix_secs) + .await + } + None => Ok(None), + } + } + pub(crate) async fn summarize_gemini_file_mappings( &self, now_unix_secs: u64, @@ -208,6 +293,37 @@ impl GatewayDataState { } } + pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + match &self.gemini_file_mapping_writer { + Some(repository) => { + repository + .delete_by_file_name_for_user(file_name, user_id) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + match &self.gemini_file_mapping_writer { + Some(repository) => { + repository + .delete_by_file_name_for_owner(file_name, key_id, user_id) + .await + } + None => Ok(false), + } + } + pub(crate) async fn delete_gemini_file_mapping_by_id( &self, mapping_id: &str, @@ -357,38 +473,11 @@ impl GatewayDataState { } } - pub(crate) async fn update_provider_catalog_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - let updated = match &self.provider_catalog_writer { - Some(repository) => { - repository - .update_key_oauth_credentials( - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } - None => Ok(false), - }?; - if updated { - self.clear_provider_catalog_cache(); - } - Ok(updated) - } - pub(crate) async fn update_provider_catalog_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { let updated = match &self.provider_catalog_writer { @@ -398,7 +487,6 @@ impl GatewayDataState { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -475,6 +563,30 @@ impl GatewayDataState { Ok(updated) } + pub(crate) async fn compare_and_swap_provider_catalog_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + let updated = match &self.provider_catalog_writer { + Some(repository) => repository.compare_and_swap_provider_config(update).await, + None => Ok(false), + }?; + self.clear_provider_catalog_cache(); + Ok(updated) + } + + pub(crate) async fn compare_and_swap_provider_catalog_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let updated = match &self.provider_catalog_writer { + Some(repository) => repository.compare_and_swap_provider_proxy(update).await, + None => Ok(false), + }?; + self.clear_provider_catalog_cache(); + Ok(updated) + } + pub(crate) async fn delete_provider_catalog_provider( &self, provider_id: &str, @@ -543,6 +655,18 @@ impl GatewayDataState { Ok(updated) } + pub(crate) async fn compare_and_swap_provider_catalog_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let updated = match &self.provider_catalog_writer { + Some(repository) => repository.compare_and_swap_endpoint_proxy(update).await, + None => Ok(false), + }?; + self.clear_provider_catalog_cache(); + Ok(updated) + } + pub(crate) async fn delete_provider_catalog_endpoint( &self, endpoint_id: &str, @@ -571,6 +695,32 @@ impl GatewayDataState { Ok(updated) } + pub(crate) async fn compare_and_swap_provider_catalog_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let updated = match &self.provider_catalog_writer { + Some(repository) => repository.compare_and_swap_key_proxy(update).await, + None => Ok(false), + }?; + self.clear_provider_catalog_cache(); + Ok(updated) + } + + pub(crate) async fn compare_and_swap_provider_catalog_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + let updated = match &self.provider_catalog_writer { + Some(repository) => repository.compare_and_swap_key_credentials(update).await, + None => Ok(false), + }?; + // Clear on both outcomes: a CAS miss proves the cached credential + // generation was stale and the retry must observe the winning record. + self.clear_provider_catalog_cache(); + Ok(updated) + } + pub(crate) async fn compare_and_update_provider_catalog_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -843,3 +993,81 @@ impl GatewayDataState { Ok(updated) } } + +#[cfg(test)] +mod request_candidate_security_tests { + use serde_json::json; + + use super::{ + sanitize_request_candidate_row, sanitize_request_candidate_rows, StoredRequestCandidate, + }; + use aether_data_contracts::repository::candidates::RequestCandidateStatus; + + fn untrusted_candidate() -> StoredRequestCandidate { + let mut candidate = StoredRequestCandidate::new( + "candidate-untrusted".to_string(), + "request-1".to_string(), + None, + None, + None, + None, + 0, + 0, + Some("provider-1".to_string()), + Some("endpoint-1".to_string()), + Some("key-1".to_string()), + RequestCandidateStatus::Failed, + None, + false, + Some(500), + None, + None, + None, + None, + None, + None, + 1, + None, + Some(2), + ) + .expect("candidate should build"); + candidate.skip_reason = Some("Bearer candidate-secret".to_string()); + candidate.error_type = Some("candidate-secret".to_string()); + candidate.error_message = Some("Bearer candidate-secret".to_string()); + candidate.extra_data = Some(json!({ + "gateway_execution_runtime": true, + "request_body": {"token": "candidate-secret"} + })); + candidate.required_capabilities = Some(json!({ + "vision": 1, + "tenant_secret": "candidate-secret" + })); + candidate + } + + fn assert_candidate_is_sanitized(candidate: &StoredRequestCandidate) { + assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip")); + assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error")); + assert!(candidate.error_message.is_none()); + assert_eq!( + candidate.extra_data, + Some(json!({"gateway_execution_runtime": true})) + ); + assert_eq!( + candidate.required_capabilities, + Some(json!({"vision": true})) + ); + assert!(!serde_json::to_string(candidate) + .expect("candidate should serialize") + .contains("candidate-secret")); + } + + #[test] + fn gateway_candidate_boundary_sanitizes_repository_rows_and_write_results() { + let candidate = sanitize_request_candidate_row(untrusted_candidate()); + assert_candidate_is_sanitized(&candidate); + + let candidates = sanitize_request_candidate_rows(vec![untrusted_candidate()]); + assert_candidate_is_sanitized(&candidates[0]); + } +} diff --git a/apps/aether-gateway/src/data/state/core.rs b/apps/aether-gateway/src/data/state/core.rs index 992fc6a02..c41e4593c 100644 --- a/apps/aether-gateway/src/data/state/core.rs +++ b/apps/aether-gateway/src/data/state/core.rs @@ -895,6 +895,44 @@ impl GatewayDataState { .await } + pub(crate) async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + if let Some(values) = &self.system_config_values { + let updated = { + let mut values = values.write().expect("system config values lock"); + match values.get_mut(key) { + Some(entry) if entry.value.as_str() == Some(expected) => { + entry.value = serde_json::Value::String(replacement.to_string()); + entry.updated_at_unix_secs = + Some(current_system_config_updated_at_unix_secs()); + true + } + _ => false, + } + }; + self.clear_cached_system_config_value(key); + return Ok(updated); + } + + let result = match self.backends.as_ref() { + Some(backends) => { + crate::request_diagnostics::observe_db_operation( + "system_config_compare_and_set", + self.database_pool_summary(), + backends.compare_and_set_system_config_string_value(key, expected, replacement), + ) + .await + } + None => Ok(false), + }; + self.clear_cached_system_config_value(key); + result + } + pub(crate) async fn upsert_system_config_value( &self, key: &str, diff --git a/apps/aether-gateway/src/data/state/integrations.rs b/apps/aether-gateway/src/data/state/integrations.rs index 5ad389012..b453afa3d 100644 --- a/apps/aether-gateway/src/data/state/integrations.rs +++ b/apps/aether-gateway/src/data/state/integrations.rs @@ -11,7 +11,11 @@ use aether_data_contracts::repository::candidates::DecisionTrace; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; -use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput}; +use aether_data_contracts::repository::proxy_nodes::ProxyNodeTrafficMutation; +use aether_data_contracts::repository::settlement::{ + ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement, + UsageSettlementInput, +}; use aether_data_contracts::repository::usage::{ ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord, UsageWriteRepository, }; @@ -34,7 +38,7 @@ const LEGACY_REQUEST_LOG_LEVEL_KEY: &str = "request_log_level"; fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestRecordLevel { let Some(value) = value.and_then(Value::as_str).map(str::trim) else { - return UsageRequestRecordLevel::Full; + return UsageRequestRecordLevel::Basic; }; if value.eq_ignore_ascii_case("basic") @@ -45,7 +49,9 @@ fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestR { UsageRequestRecordLevel::Basic } else { - UsageRequestRecordLevel::Full + // Raw HTTP payload capture is disabled at the runtime boundary. The setting remains + // accepted for compatibility, but no longer authorizes collecting request/response data. + UsageRequestRecordLevel::Basic } } @@ -95,6 +101,14 @@ impl StoredVideoTaskReadSide for GatewayDataState { ) -> Result, DataLayerError> { GatewayDataState::find_video_task(self, key).await } + + async fn find_stored_video_task_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + GatewayDataState::find_video_task_for_user(self, key, user_id).await + } } #[async_trait] @@ -219,6 +233,13 @@ impl UsageSettlementWriter for GatewayDataState { GatewayDataState::has_settlement_writer(self) } + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + GatewayDataState::reconcile_usage_policy_cost(self, input).await + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -285,10 +306,26 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState { failed_delta: i64, latency_ms: Option, ) -> Result<(), DataLayerError> { + // This API predates incarnation fences and only receives a node id. Read + // the selected node first so every durable path can bind the delta to + // the observed generation. If the node disappeared, fail closed rather + // than allowing a bare id to target a replacement node. + let Some(node) = self.find_proxy_node(node_id).await? else { + return Ok(()); + }; + if !node.is_manual { + return Ok(()); + } + let expected_tunnel_generation = node.tunnel_generation.trim().to_string(); + if expected_tunnel_generation.is_empty() { + return Ok(()); + } + if let Some(repository) = &self.usage_writer { let enqueued = repository .enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta { node_id: node_id.to_string(), + expected_tunnel_generation: Some(expected_tunnel_generation.clone()), total_requests_delta: total_delta, failed_requests_delta: failed_delta, dns_failures_delta: 0, @@ -300,14 +337,25 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState { } } - match &self.proxy_node_writer { - Some(repository) => { - repository - .increment_manual_node_requests(node_id, total_delta, failed_delta, latency_ms) - .await - } - None => Ok(()), + // The legacy increment method has no generation argument and would + // re-read the current row, which is vulnerable to an id reuse between + // the read above and the write. Use the fenced traffic mutation as the + // only fallback. It intentionally omits the legacy latency-only field; + // preserving counters safely is more important than an unfenced write. + if let Some(repository) = &self.proxy_node_writer { + let _ = repository + .record_traffic(&ProxyNodeTrafficMutation { + node_id: node_id.to_string(), + expected_tunnel_generation: Some(expected_tunnel_generation), + total_requests_delta: total_delta, + failed_requests_delta: failed_delta, + dns_failures_delta: 0, + stream_errors_delta: 0, + }) + .await?; } + let _ = latency_ms; + Ok(()) } } @@ -452,6 +500,20 @@ mod tests { assert_eq!(level, UsageRequestRecordLevel::Basic); } + #[tokio::test] + async fn usage_runtime_access_disables_full_http_capture() { + let state = GatewayDataState::disabled().with_system_config_values_for_tests([( + "request_record_level".to_string(), + json!("full"), + )]); + + let level = UsageRuntimeAccess::request_record_level(&state) + .await + .expect("request record level should read"); + + assert_eq!(level, UsageRequestRecordLevel::Basic); + } + #[tokio::test] async fn usage_runtime_access_falls_back_to_legacy_request_log_level_alias() { let state = GatewayDataState::disabled().with_system_config_values_for_tests([( @@ -467,14 +529,14 @@ mod tests { } #[tokio::test] - async fn usage_runtime_access_defaults_missing_request_record_level_to_full() { + async fn usage_runtime_access_defaults_missing_request_record_level_to_basic() { let state = GatewayDataState::disabled(); let level = UsageRuntimeAccess::request_record_level(&state) .await .expect("missing request record level should fall back"); - assert_eq!(level, UsageRequestRecordLevel::Full); + assert_eq!(level, UsageRequestRecordLevel::Basic); } #[tokio::test] @@ -488,6 +550,20 @@ mod tests { .await .expect("body capture policy should read"); - assert_eq!(policy.record_level, UsageRequestRecordLevel::Full); + assert_eq!(policy.record_level, UsageRequestRecordLevel::Basic); + } + + #[tokio::test] + async fn usage_runtime_access_fails_closed_for_unknown_record_level() { + let state = GatewayDataState::disabled().with_system_config_values_for_tests([( + "request_record_level".to_string(), + json!("everything"), + )]); + + let level = UsageRuntimeAccess::request_record_level(&state) + .await + .expect("request record level should read"); + + assert_eq!(level, UsageRequestRecordLevel::Basic); } } diff --git a/apps/aether-gateway/src/data/state/mod.rs b/apps/aether-gateway/src/data/state/mod.rs index fcffb2a1e..0a455684b 100644 --- a/apps/aether-gateway/src/data/state/mod.rs +++ b/apps/aether-gateway/src/data/state/mod.rs @@ -27,8 +27,8 @@ use aether_data::repository::auth::{ StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, }; use aether_data::repository::auth_modules::{ - AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig, - StoredOAuthProviderModuleConfig, + AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult, + LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig, }; use aether_data::repository::gemini_file_mappings::{ GeminiFileMappingListQuery, GeminiFileMappingReadRepository, GeminiFileMappingStats, @@ -36,9 +36,10 @@ use aether_data::repository::gemini_file_mappings::{ UpsertGeminiFileMappingRecord, }; use aether_data::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenWithUser, + UpdateManagementTokenRecord, }; use aether_data::repository::oauth_providers::{ OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig, @@ -63,21 +64,24 @@ pub(crate) use aether_data::repository::users::{ use aether_data::repository::wallet::{ AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, - AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, + AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, - FailAdminWalletRefundInput, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, - ProcessPaymentCallbackOutcome, RedeemWalletCodeInput, RedeemWalletCodeOutcome, + FailAdminWalletRefundInput, FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, + ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage, StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use aether_data::{ @@ -92,8 +96,9 @@ use aether_data_contracts::repository::background_tasks::{ use aether_data_contracts::repository::billing::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; use aether_data_contracts::repository::candidate_selection::{ @@ -122,12 +127,14 @@ use aether_data_contracts::repository::pool_scores::{ }; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate, - ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery, - ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, - ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, - ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, - StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, - StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate, + ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, + ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, + ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate, + ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogWriteRepository, + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, + StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, }; use aether_data_contracts::repository::quota::{ ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot, @@ -136,7 +143,10 @@ use aether_data_contracts::repository::routing_profiles::{ RoutingGroupReadRepository, RoutingGroupWriteRepository, }; use aether_data_contracts::repository::settlement::{ - SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation, + StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsageSettlementInput, }; use aether_data_contracts::repository::usage::{ ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary, @@ -150,9 +160,9 @@ use aether_data_contracts::repository::video_tasks::{ use aether_runtime_state::RuntimeQueueStore; pub(crate) use self::referrals::{ - ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery, - ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery, - ReferralRewardRecord, ReferralUserDashboard, + ReferralAdminStats, ReferralMutationStatus, ReferralReconciliationSummary, + ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig, + ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard, }; #[derive(Clone, Default)] diff --git a/apps/aether-gateway/src/data/state/referrals.rs b/apps/aether-gateway/src/data/state/referrals.rs index 7ebb8bf25..5c5719d66 100644 --- a/apps/aether-gateway/src/data/state/referrals.rs +++ b/apps/aether-gateway/src/data/state/referrals.rs @@ -4,9 +4,9 @@ use aether_data::DataLayerError; use super::GatewayDataState; pub(crate) use aether_data::backend::{ - ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery, - ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery, - ReferralRewardRecord, ReferralUserDashboard, + ReferralAdminStats, ReferralMutationStatus, ReferralReconciliationSummary, + ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig, + ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard, }; impl GatewayDataState { @@ -115,4 +115,13 @@ impl GatewayDataState { .reverse_referral_rewards_for_order(order_id, amount_usd) .await } + + pub(crate) async fn reconcile_referral_rewards_once( + &self, + reward_config: Option, + ) -> Result { + self.referrals() + .reconcile_referral_rewards_once(reward_config) + .await + } } diff --git a/apps/aether-gateway/src/data/state/routing_group_cache.rs b/apps/aether-gateway/src/data/state/routing_group_cache.rs index 47eda5871..b03ac9011 100644 --- a/apps/aether-gateway/src/data/state/routing_group_cache.rs +++ b/apps/aether-gateway/src/data/state/routing_group_cache.rs @@ -684,7 +684,7 @@ mod tests { .unwrap_or_default(); // One Arc is retained by the map and every active request // owns one through its leader guard or follower state. - if participant_count >= participants + 1 { + if participant_count > participants { break; } tokio::task::yield_now().await; diff --git a/apps/aether-gateway/src/data/state/runtime.rs b/apps/aether-gateway/src/data/state/runtime.rs index 96cb79121..aadf9215f 100644 --- a/apps/aether-gateway/src/data/state/runtime.rs +++ b/apps/aether-gateway/src/data/state/runtime.rs @@ -7,33 +7,39 @@ use super::{ AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery, BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey, BillingModelContextCacheState, BillingModelContextInflightState, BillingPlanRecord, - BillingPlanWriteInput, CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + BillingPlanWriteInput, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, CreateAnnouncementRecord, CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, DataLayerError, DatabaseMaintenanceSummary, DecisionTrace, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse, - PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, ProcessAdminWalletRefundInput, - ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, RedeemWalletCodeInput, - RedeemWalletCodeOutcome, RequestAuditBundle, RequestCandidateTrace, StoredAdminAuditLogPage, - StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, - StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, - StoredAdminWalletLedgerPage, StoredAdminWalletListPage, StoredAdminWalletRefund, - StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, + 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, StoredUsageSettlement, - StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, - StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, - StoredWalletSnapshot, UpdateAnnouncementRecord, UpsertBackgroundTaskEvent, - UpsertBackgroundTaskRun, UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, - UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, VideoTaskLookupKey, - VideoTaskModelCount, VideoTaskQueryFilter, VideoTaskStatusCount, - WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult, WalletLookupKey, - WalletMutationOutcome, + 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, + VideoTaskStatusCount, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult, + WalletLookupKey, WalletMutationOutcome, }; use aether_data_contracts::repository::usage::{ PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest, @@ -43,7 +49,9 @@ use aether_data_contracts::repository::usage::{ UsageDailyHeatmapQuery, }; use aether_runtime_state::RuntimeQueueStore; -use aether_video_tasks_core::read_data_backed_video_task_response; +use aether_video_tasks_core::{ + read_data_backed_video_task_response, read_data_backed_video_task_response_for_user, +}; use std::time::{Duration, Instant}; use tokio::time::timeout; @@ -558,6 +566,17 @@ impl GatewayDataState { } } + pub(crate) async fn find_video_task_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + match &self.video_task_reader { + Some(repository) => repository.find_for_user(key, user_id).await, + None => Ok(None), + } + } + pub(crate) async fn list_video_task_page( &self, filter: &VideoTaskQueryFilter, @@ -894,6 +913,21 @@ impl GatewayDataState { } } + pub(crate) async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, DataLayerError> { + match &self.wallet_reader { + Some(repository) => { + repository + .find_wallet_recharge_order_by_order_no(user_id, order_no) + .await + } + None => Ok(None), + } + } + pub(crate) async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, @@ -909,6 +943,16 @@ impl GatewayDataState { } } + pub(crate) async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, DataLayerError> { + match &self.wallet_reader { + Some(repository) => repository.find_payment_order_by_order_no(order_no).await, + None => Ok(None), + } + } + pub(crate) async fn find_wallet_refund( &self, wallet_id: &str, @@ -934,6 +978,58 @@ impl GatewayDataState { } } + pub(crate) async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result>, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .update_wallet_recharge_checkout(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .compare_and_swap_payment_order_stripe_client_secret(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result>, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .fail_wallet_recharge_checkout(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result>, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .reclaim_wallet_recharge_checkout(input) + .await + .map(Some), + None => Ok(None), + } + } + pub(crate) async fn create_plan_purchase_order( &self, input: CreatePlanPurchaseOrderInput, @@ -1009,6 +1105,19 @@ impl GatewayDataState { } } + pub(crate) async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result>, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .update_admin_wallet_refund_gateway(input) + .await + .map(Some), + None => Ok(None), + } + } + pub(crate) async fn complete_admin_wallet_refund( &self, input: CompleteAdminWalletRefundInput, @@ -1151,6 +1260,83 @@ impl GatewayDataState { } } + pub(crate) async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result, DataLayerError> { + match &self.settlement_writer { + Some(repository) => repository.reserve_usage_policy_cost(input).await.map(Some), + None => Ok(None), + } + } + + pub(crate) async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result, DataLayerError> { + match &self.settlement_writer { + Some(repository) => repository + .reserve_usage_policy_request(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, DataLayerError> { + match &self.settlement_writer { + Some(repository) => { + repository + .release_usage_policy_request_admission(input) + .await + } + None => Ok(None), + } + } + + pub(crate) async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + match &self.settlement_writer { + Some(repository) => { + repository + .cleanup_usage_policy_request_admissions(now_unix_secs, batch_size) + .await + } + None => Ok(0), + } + } + + pub(crate) async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + match &self.settlement_writer { + Some(repository) => repository.reconcile_usage_policy_cost(input).await, + None => Ok(None), + } + } + + pub(crate) async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + match &self.settlement_writer { + Some(repository) => { + repository + .cleanup_usage_policy_cost_reservations(now_unix_secs, batch_size) + .await + } + None => Ok(0), + } + } + pub(crate) async fn reset_due_provider_quotas( &self, now_unix_secs: u64, @@ -2496,6 +2682,48 @@ impl GatewayDataState { } } + pub(crate) async fn find_payment_gateway_config_strong( + &self, + provider: &str, + ) -> Result, DataLayerError> { + match &self.billing_reader { + Some(repository) => { + repository + .find_payment_gateway_config_strong(provider) + .await + } + None => Ok(None), + } + } + + pub(crate) async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + match &self.billing_reader { + Some(repository) => { + repository + .compare_and_swap_payment_gateway_secret(update) + .await + } + None => Ok(false), + } + } + + pub(crate) async fn compare_and_swap_payment_gateway_config( + &self, + input: &PaymentGatewayConfigCasWriteInput, + ) -> Result, DataLayerError> { + match &self.billing_reader { + Some(repository) => { + repository + .compare_and_swap_payment_gateway_config(input) + .await + } + None => Ok(AdminBillingMutationOutcome::Unavailable), + } + } + pub(crate) async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -2578,6 +2806,21 @@ impl GatewayDataState { } } + pub(crate) async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, DataLayerError> { + match &self.billing_reader { + Some(repository) => { + repository + .revoke_user_plan_entitlement(user_id, entitlement_id) + .await + } + None => Ok(AdminBillingMutationOutcome::Unavailable), + } + } + pub(crate) async fn find_user_daily_quota_availability( &self, user_id: &str, @@ -2652,6 +2895,16 @@ impl GatewayDataState { read_data_backed_video_task_response(self, route_family, request_path).await } + pub(crate) async fn read_video_task_response_for_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + ) -> Result, DataLayerError> { + read_data_backed_video_task_response_for_user(self, route_family, request_path, user_id) + .await + } + pub(crate) async fn find_background_task_run( &self, run_id: &str, diff --git a/apps/aether-gateway/src/data/state/testing/mod.rs b/apps/aether-gateway/src/data/state/testing/mod.rs index 486d6f435..fd7595dec 100644 --- a/apps/aether-gateway/src/data/state/testing/mod.rs +++ b/apps/aether-gateway/src/data/state/testing/mod.rs @@ -1,10 +1,15 @@ use std::collections::BTreeMap; use std::sync::{Arc, RwLock}; +use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; use aether_data_contracts::repository::candidates::RequestCandidateRepository; use aether_data_contracts::repository::pool_scores::PoolMemberScoreRepository; use aether_data_contracts::repository::quota::ProviderQuotaRepository; +use aether_data_contracts::repository::routing_profiles::{ + StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion, +}; use aether_data_contracts::repository::usage::UsageRepository; +use aether_routing_core::RoutingGroupConfig; use super::{ AnnouncementReadRepository, AnnouncementWriteRepository, AuthApiKeyReadRepository, @@ -213,6 +218,21 @@ impl GatewayDataState { self } + #[cfg(test)] + pub(crate) fn with_cached_provider_catalog_reader_for_tests( + mut self, + repository: Arc, + ) -> Self + where + T: ProviderCatalogReadRepository + 'static, + { + let inner: Arc = repository; + self.provider_catalog_reader = Some(Arc::new( + super::provider_catalog_cache::CachedProviderCatalogReadRepository::new(inner), + )); + self + } + #[cfg(test)] pub(crate) fn with_request_candidate_reader( mut self, @@ -877,6 +897,30 @@ impl GatewayDataState { self } + #[cfg(test)] + pub(crate) fn with_system_default_routing_group_for_tests(self) -> Self { + let now = 1; + let repository = Arc::new(InMemoryRoutingGroupRepository::seed( + [StoredRoutingGroup { + id: "system-default".to_string(), + name: "system-default".to_string(), + description: Some("test system default routing strategy".to_string()), + enabled: true, + is_system_default: true, + sort_order: 0, + config_json: serde_json::to_value(RoutingGroupConfig::default()) + .expect("default routing config should serialize"), + version: 1, + created_at: now, + updated_at: now, + published_at: Some(now), + }], + std::iter::empty::(), + std::iter::empty::(), + )); + self.with_routing_group_repository_for_tests(repository) + } + #[cfg(test)] pub(crate) fn with_auth_api_key_reader( mut self, @@ -1746,6 +1790,18 @@ impl GatewayDataState { } } + #[cfg(test)] + pub(crate) fn attach_auth_api_key_repository_for_tests(mut self, repository: Arc) -> Self + where + T: aether_data::repository::auth::AuthRepository + 'static, + { + let auth_api_key_reader: Arc = repository.clone(); + let auth_api_key_writer: Arc = repository; + self.auth_api_key_reader = Some(auth_api_key_reader); + self.auth_api_key_writer = Some(auth_api_key_writer); + self + } + #[cfg(test)] pub(crate) fn with_decision_trace_readers_for_tests( request_candidate_repository: Arc, @@ -2335,6 +2391,36 @@ impl GatewayDataState { } } + #[cfg(test)] + pub(crate) fn with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests< + T, + U, + V, + >( + auth_api_key_repository: Arc, + candidate_selection_repository: Arc, + provider_catalog_repository: Arc, + request_candidate_repository: Arc, + gemini_file_mapping_repository: Arc, + encryption_key: impl Into, + ) -> Self + where + T: RequestCandidateRepository + 'static, + U: ProviderCatalogReadRepository + ProviderCatalogWriteRepository + 'static, + V: aether_data::repository::gemini_file_mappings::GeminiFileMappingRepository + 'static, + { + let mut state = Self::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests( + auth_api_key_repository, + candidate_selection_repository, + provider_catalog_repository, + request_candidate_repository, + encryption_key, + ); + state.gemini_file_mapping_reader = Some(gemini_file_mapping_repository.clone()); + state.gemini_file_mapping_writer = Some(gemini_file_mapping_repository); + state + } + #[cfg(test)] pub(crate) fn with_auth_candidate_selection_provider_catalog_request_candidates_for_tests< T, diff --git a/apps/aether-gateway/src/data/state/testkit.rs b/apps/aether-gateway/src/data/state/testkit.rs index 0e652582d..f1cf3856f 100644 --- a/apps/aether-gateway/src/data/state/testkit.rs +++ b/apps/aether-gateway/src/data/state/testkit.rs @@ -12,10 +12,113 @@ use aether_data_contracts::repository::usage::{ }; use aether_data::repository::auth::AuthApiKeyReadRepository; +use aether_data::repository::management_tokens::{ + InMemoryManagementTokenRepository, ManagementTokenReadRepository, + ManagementTokenWriteRepository, StoredManagementToken, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, +}; +use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; +use aether_data::repository::proxy_nodes::{ProxyNodeReadRepository, ProxyNodeWriteRepository}; +use aether_data::repository::users::{ + InMemoryUserReadRepository, StoredUserAuthRecord, UserReadRepository, +}; +use sha2::{Digest, Sha256}; use super::{GatewayDataConfig, GatewayDataState}; impl GatewayDataState { + pub(crate) fn with_tunnel_management_auth_for_testkit( + node_id: &str, + tunnel_generation: &str, + raw_token: &str, + encryption_key: impl Into, + ) -> Result { + const TOKEN_ID: &str = "token-tunnel-harness"; + const USER_ID: &str = "user-tunnel-harness"; + + let node = StoredProxyNode::new( + node_id.to_string(), + "tunnel harness node".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "offline".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + true, + false, + 0, + )? + .with_tunnel_generation(tunnel_generation.to_string()); + let proxy_repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + + let user_summary = StoredManagementTokenUserSummary::new( + USER_ID.to_string(), + Some("tunnel-harness@example.com".to_string()), + "tunnel_harness_admin".to_string(), + "admin".to_string(), + )?; + let token = StoredManagementToken::new( + TOKEN_ID.to_string(), + USER_ID.to_string(), + "tunnel harness token".to_string(), + )? + .with_permissions(Some(serde_json::json!(["admin:proxy_nodes:admin"]))); + let token_hash = format!("{:x}", Sha256::digest(raw_token.as_bytes())); + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + [StoredManagementTokenWithUser::new(token, user_summary)], + [(token_hash, TOKEN_ID.to_string())], + )); + let token_reader: Arc = token_repository.clone(); + let token_writer: Arc = token_repository; + + let user = StoredUserAuthRecord::new( + USER_ID.to_string(), + Some("tunnel-harness@example.com".to_string()), + true, + "tunnel_harness_admin".to_string(), + None, + "admin".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + None, + None, + )?; + let user_reader: Arc = + Arc::new(InMemoryUserReadRepository::seed_auth_users([user])); + + let mut state = + Self::with_proxy_node_repository_for_testkit(proxy_repository, encryption_key); + state.management_token_reader = Some(token_reader); + state.management_token_writer = Some(token_writer); + state.user_reader = Some(user_reader); + Ok(state) + } + + pub(crate) fn with_proxy_node_repository_for_testkit( + repository: Arc, + encryption_key: impl Into, + ) -> Self + where + T: ProxyNodeReadRepository + ProxyNodeWriteRepository + 'static, + { + let proxy_node_reader: Arc = repository.clone(); + let proxy_node_writer: Arc = repository; + let mut state = Self::disabled(); + state.config = GatewayDataConfig::disabled().with_encryption_key(encryption_key); + state.proxy_node_reader = Some(proxy_node_reader); + state.proxy_node_writer = Some(proxy_node_writer); + state + } + pub(crate) fn with_openai_chat_pressure_repositories_for_testkit( auth_api_key_repository: Arc, candidate_selection_repository: Arc, diff --git a/apps/aether-gateway/src/data/tests.rs b/apps/aether-gateway/src/data/tests.rs index 25c1154a0..de857c64a 100644 --- a/apps/aether-gateway/src/data/tests.rs +++ b/apps/aether-gateway/src/data/tests.rs @@ -312,7 +312,7 @@ async fn data_state_checks_user_uniqueness_through_user_reader() { Some("admin@example.com".to_string()), true, "admin".to_string(), - Some(format!("$2b$12${}", "a".repeat(53))), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), "admin".to_string(), "local".to_string(), None, diff --git a/apps/aether-gateway/src/dispatch/pool_scheduler.rs b/apps/aether-gateway/src/dispatch/pool_scheduler.rs index eff085bd1..1b35ef641 100644 --- a/apps/aether-gateway/src/dispatch/pool_scheduler.rs +++ b/apps/aether-gateway/src/dispatch/pool_scheduler.rs @@ -318,7 +318,9 @@ fn active_probe_member_is_unschedulable_for_request( }) { return true; } - key_context.is_some_and(|context| context.account_blocked || context.quota_exhausted) + key_context.is_some_and(|context| { + context.account_blocked || context.quota_exhausted || context.quota_hard_blocked + }) } async fn expand_pool_group_candidate( @@ -598,11 +600,17 @@ impl<'a> PoolKeyCursor<'a> { return; }; self.exhaustion_skip_recorded = true; - record_local_runtime_candidate_skip_reason( - self.state.app(), - trace_id, - self.runtime_miss_pool_exhaustion_skip_reason(), - ); + if self.skip_reason_counts.is_empty() { + record_local_runtime_candidate_skip_reason( + self.state.app(), + trace_id, + "pool_group_exhausted", + ); + return; + } + for reason in self.skip_reason_counts.keys() { + record_local_runtime_candidate_skip_reason(self.state.app(), trace_id, reason); + } } fn runtime_miss_pool_exhaustion_skip_reason(&self) -> &'static str { @@ -748,19 +756,28 @@ impl<'a> PoolKeyCursor<'a> { return None; } }; + let api_format = self.group.candidate.endpoint_api_format.as_str(); rows.sort_by(|left, right| { - let left_priority = self - .routing_overlay - .as_ref() - .map_or(left.key_internal_priority, |overlay| { - overlay.key_priority(&left.key_id, left.key_internal_priority) - }); - let right_priority = self - .routing_overlay - .as_ref() - .map_or(right.key_internal_priority, |overlay| { - overlay.key_priority(&right.key_id, right.key_internal_priority) - }); + let left_priority = self.routing_overlay.as_ref().map_or( + left.key_internal_priority, + |overlay| { + overlay.key_priority_for_format( + &left.key_id, + api_format, + left.key_internal_priority, + ) + }, + ); + let right_priority = self.routing_overlay.as_ref().map_or( + right.key_internal_priority, + |overlay| { + overlay.key_priority_for_format( + &right.key_id, + api_format, + right.key_internal_priority, + ) + }, + ); left_priority .cmp(&right_priority) .then(left.key_id.cmp(&right.key_id)) @@ -1441,13 +1458,30 @@ async fn read_pool_catalog_key_contexts_by_id( key_count = key_ids.len(), "gateway pool scheduler: failed to read catalog key metadata" ); - return BTreeMap::new(); + // Do not fail open when the quota metadata read is unavailable. A + // missing context must never turn an exhausted account into an + // eligible candidate and produce another upstream 429. The caller + // treats this marker as a pool quota skip and the next request will + // retry the metadata read. + return key_ids + .into_iter() + .map(|key_id| { + ( + key_id, + PoolCatalogKeyContext { + quota_hard_blocked: true, + ..PoolCatalogKeyContext::default() + }, + ) + }) + .collect(); } }; let provider_pool_service = ProviderPoolService::with_builtin_adapters(); - keys.into_iter() + let mut contexts = keys + .into_iter() .map(|key| { let provider_type = provider_type_by_key_id .get(&key.id) @@ -1464,7 +1498,19 @@ async fn read_pool_catalog_key_contexts_by_id( ), ) }) - .collect() + .collect::>(); + // A key can disappear between the candidate-row and catalog reads. Keep + // the snapshot non-empty and fail closed for those IDs so the caller does + // not interpret an incomplete read as "all accounts are healthy". + for key_id in key_ids { + contexts + .entry(key_id) + .or_insert_with(|| PoolCatalogKeyContext { + quota_hard_blocked: true, + ..PoolCatalogKeyContext::default() + }); + } + contexts } fn build_pool_catalog_key_context( @@ -1709,7 +1755,21 @@ fn run_local_execution_pool_scheduler_with_runtime_map( let key_context = key_context_by_id .get(&candidate.candidate.key_id) .cloned() - .unwrap_or_default(); + .unwrap_or_else(|| { + // An explicitly non-empty metadata snapshot should contain + // every catalog key in this page. If one disappeared between + // reads, fail closed for that key instead of sending traffic + // with an unknown quota state. Empty maps are retained for + // callers/tests that intentionally provide no runtime context. + if key_context_by_id.is_empty() { + PoolCatalogKeyContext::default() + } else { + PoolCatalogKeyContext { + quota_hard_blocked: true, + ..PoolCatalogKeyContext::default() + } + } + }); let admin_pool_config = effective_pool_config_by_provider .get(&candidate.candidate.provider_id) .cloned() @@ -1936,11 +1996,13 @@ fn apply_pool_orchestration( orchestration: PoolCandidateOrchestration, ) -> EligibleLocalExecutionCandidate { let scheduler_affinity_epoch = candidate.orchestration.scheduler_affinity_epoch; + let sticky_key_attempts = candidate.orchestration.sticky_key_attempts; candidate.orchestration = LocalExecutionCandidateMetadata { candidate_group_id: orchestration.candidate_group_id, pool_key_index: orchestration.pool_key_index, pool_key_lease: None, scheduler_affinity_epoch, + sticky_key_attempts, }; candidate } @@ -1983,7 +2045,7 @@ mod tests { use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; - use aether_pool_core::PoolSchedulingPreset; + use aether_pool_core::{PoolSchedulingPreset, POOL_ACCOUNT_EXHAUSTED_SKIP_REASON}; use aether_provider_pool::ProviderPoolService; use aether_provider_transport::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, @@ -2095,6 +2157,55 @@ mod tests { ); } + #[test] + fn pool_scheduler_skips_quota_exhausted_key_when_flag_is_false() { + let ready = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "key-ready", + 10, + Some(json!({ "pool_advanced": {} })), + ); + let exhausted = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "key-exhausted", + 10, + Some(json!({ "pool_advanced": { "skip_exhausted_accounts": false } })), + ); + let key_context_by_id = BTreeMap::from([ + ("key-ready".to_string(), PoolCatalogKeyContext::default()), + ( + "key-exhausted".to_string(), + PoolCatalogKeyContext { + quota_exhausted: true, + ..PoolCatalogKeyContext::default() + }, + ), + ]); + + let (scheduled, skipped) = apply_local_execution_pool_scheduler_with_runtime_map( + vec![ready, exhausted], + &BTreeMap::new(), + &key_context_by_id, + ); + + assert_eq!( + scheduled + .iter() + .map(|item| item.candidate.key_id.as_str()) + .collect::>(), + vec!["key-ready"] + ); + assert_eq!( + skipped + .iter() + .map(|item| (item.candidate.key_id.as_str(), item.skip_reason)) + .collect::>(), + vec![("key-exhausted", POOL_ACCOUNT_EXHAUSTED_SKIP_REASON)] + ); + } + #[test] fn pool_scheduler_attaches_group_and_pool_metadata_to_ranked_candidates() { let pool_first = sample_eligible_candidate( @@ -2144,6 +2255,7 @@ mod tests { pool_key_index: Some(0), pool_key_lease: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, } ); assert_eq!(reordered[1].orchestration.pool_key_index, Some(1)); @@ -2161,6 +2273,7 @@ mod tests { pool_key_index: None, pool_key_lease: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, } ); } @@ -3245,8 +3358,12 @@ mod tests { .take_local_execution_runtime_miss_diagnostic(trace_id) .expect("runtime miss diagnostic should exist"); assert_eq!(diagnostic.reason, "all_candidates_skipped"); - assert_eq!(diagnostic.skipped_candidate_count, Some(1)); + assert_eq!(diagnostic.skipped_candidate_count, Some(2)); assert_eq!(diagnostic.skip_reasons.get("pool_cooldown"), Some(&1)); + assert_eq!( + diagnostic.skip_reasons.get("transport_snapshot_missing"), + Some(&1) + ); } #[tokio::test] @@ -4689,6 +4806,15 @@ 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, @@ -4739,10 +4865,18 @@ mod tests { ) .expect("endpoint transport should build"); + let credential_state = provider_catalog_credential_state(); 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(), @@ -4754,7 +4888,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:chat"])), - Some(format!("secret-{index}")), + encrypted_api_key, None, None, None, @@ -4890,6 +5024,9 @@ mod tests { } 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(), @@ -4901,7 +5038,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:responses"])), - Some(format!("secret-{key_id}")), + encrypted_api_key, None, None, Some(json!({"openai:responses": 1})), @@ -5053,6 +5190,8 @@ mod tests { priority_mode: RoutingSetPriorityMode::Provider, scheduling_mode: RoutingSchedulingMode::CacheAffinity, keep_priority_on_conversion: false, + sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: Default::default(), ranking_overlay: RankingOverlay { allowed_keys: key_ids.into_iter().map(str::to_string).collect(), ..RankingOverlay::default() diff --git a/apps/aether-gateway/src/dispatch/refs.rs b/apps/aether-gateway/src/dispatch/refs.rs index e8acb0a1d..4c689cfb3 100644 --- a/apps/aether-gateway/src/dispatch/refs.rs +++ b/apps/aether-gateway/src/dispatch/refs.rs @@ -208,6 +208,7 @@ mod tests { pool_key_index: None, pool_key_lease: None, scheduler_affinity_epoch: None, + sticky_key_attempts: None, }, ranking: None, } diff --git a/apps/aether-gateway/src/email_delivery.rs b/apps/aether-gateway/src/email_delivery.rs index dc98e21a7..f6d5a6e96 100644 --- a/apps/aether-gateway/src/email_delivery.rs +++ b/apps/aether-gateway/src/email_delivery.rs @@ -1,13 +1,28 @@ use base64::Engine; use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, system_config_bool, system_config_string, + decrypt_or_migrate_smtp_password, smtp_password_binding, system_config_bool, }; use crate::{AppState, GatewayError}; const SMTP_TIMEOUT_SECS: u64 = 30; +const SMTP_MAX_HOST_BYTES: usize = 255; +const SMTP_MAX_ADDRESS_BYTES: usize = 320; +const SMTP_MAX_HEADER_VALUE_BYTES: usize = 512; +const SMTP_MAX_USERNAME_BYTES: usize = 320; +const SMTP_MAX_PASSWORD_BYTES: usize = 16 * 1024; +const SMTP_MAX_STORED_PASSWORD_BYTES: usize = 64 * 1024; +const SMTP_MAX_BODY_BYTES: usize = 2 * 1024 * 1024; +const SMTP_MAX_MESSAGE_BYTES: usize = 8 * 1024 * 1024; +const SMTP_MAX_DIAGNOSTIC_BYTES: usize = 4096; +// SMTP servers normally emit short ASCII status lines. Keep parser buffers +// bounded even when the peer is untrusted or compromised; these limits apply +// only to control responses, not to the message body being submitted. +const SMTP_MAX_RESPONSE_LINE_BYTES: usize = 16 * 1024; +const SMTP_MAX_RESPONSE_BYTES: usize = 256 * 1024; +const SMTP_MAX_RESPONSE_LINES: usize = 128; -#[derive(Debug, Clone)] +#[derive(Clone)] pub(crate) struct SmtpDeliveryConfig { pub(crate) host: String, pub(crate) port: u16, @@ -19,7 +34,7 @@ pub(crate) struct SmtpDeliveryConfig { pub(crate) from_name: String, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub(crate) struct ComposedEmail { pub(crate) to_email: String, pub(crate) subject: String, @@ -27,6 +42,26 @@ pub(crate) struct ComposedEmail { pub(crate) text_body: String, } +fn bounded_system_config_string( + field: &str, + value: Option<&serde_json::Value>, + max_bytes: usize, +) -> Result, GatewayError> { + let Some(serde_json::Value::String(raw)) = value else { + return Ok(None); + }; + let value = raw.trim(); + if value.is_empty() { + return Ok(None); + } + if value.len() > max_bytes { + return Err(GatewayError::Internal(format!( + "smtp {field} exceeds the allowed size" + ))); + } + Ok(Some(value.to_string())) +} + pub(crate) async fn read_smtp_delivery_config( state: &AppState, ) -> Result, GatewayError> { @@ -34,10 +69,16 @@ pub(crate) async fn read_smtp_delivery_config( let smtp_from_email = state .read_system_config_json_value("smtp_from_email") .await?; - let Some(host) = system_config_string(smtp_host.as_ref()) else { + let Some(host) = bounded_system_config_string("host", smtp_host.as_ref(), SMTP_MAX_HOST_BYTES)? + else { return Ok(None); }; - let Some(from_email) = system_config_string(smtp_from_email.as_ref()) else { + let Some(from_email) = bounded_system_config_string( + "from_email", + smtp_from_email.as_ref(), + SMTP_MAX_ADDRESS_BYTES, + )? + else { return Ok(None); }; @@ -50,20 +91,43 @@ pub(crate) async fn read_smtp_delivery_config( .read_system_config_json_value("smtp_from_name") .await?; - let password = system_config_string(smtp_password.as_ref()).map(|value| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value) - }); + let port = system_config_u16(smtp_port.as_ref(), 587); + let user = bounded_system_config_string("user", smtp_user.as_ref(), SMTP_MAX_USERNAME_BYTES)?; + let use_tls = system_config_bool(smtp_use_tls.as_ref(), true); + let use_ssl = system_config_bool(smtp_use_ssl.as_ref(), false); + let password = match ( + bounded_system_config_string( + "stored_password", + smtp_password.as_ref(), + SMTP_MAX_STORED_PASSWORD_BYTES, + )?, + smtp_password_binding(&host, port, user.as_deref(), use_tls, use_ssl), + ) { + (Some(value), Some(binding)) => { + Some(decrypt_or_migrate_smtp_password(state, &binding, value).await?) + } + (Some(_), None) => { + return Err(GatewayError::Internal( + "SMTP password binding is invalid".to_string(), + )); + } + (None, _) => None, + }; Ok(Some(SmtpDeliveryConfig { host, - port: system_config_u16(smtp_port.as_ref(), 587), - user: system_config_string(smtp_user.as_ref()), + port, + user, password, - use_tls: system_config_bool(smtp_use_tls.as_ref(), true), - use_ssl: system_config_bool(smtp_use_ssl.as_ref(), false), + use_tls, + use_ssl, from_email, - from_name: system_config_string(smtp_from_name.as_ref()) - .unwrap_or_else(|| "Aether".to_string()), + from_name: bounded_system_config_string( + "from_name", + smtp_from_name.as_ref(), + SMTP_MAX_HEADER_VALUE_BYTES, + )? + .unwrap_or_else(|| "Aether".to_string()), })) } @@ -71,17 +135,150 @@ pub(crate) async fn send_smtp_email( config: SmtpDeliveryConfig, email: ComposedEmail, ) -> Result<(), GatewayError> { + validate_smtp_delivery_inputs(&config, &email)?; tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email)) .await .map_err(|err| GatewayError::Internal(err.to_string()))? } pub(crate) async fn probe_smtp_connection(config: SmtpDeliveryConfig) -> Result<(), GatewayError> { + validate_smtp_config(&config)?; tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config)) .await .map_err(|err| GatewayError::Internal(err.to_string()))? } +fn validate_smtp_control_field(field: &str, value: &str) -> Result<(), GatewayError> { + if value.bytes().any(|byte| byte < 0x20 || byte == 0x7f) { + return Err(GatewayError::Internal(format!( + "smtp {field} contains forbidden control characters" + ))); + } + Ok(()) +} + +fn validate_smtp_bounded_field( + field: &str, + value: &str, + max_bytes: usize, +) -> Result<(), GatewayError> { + validate_smtp_control_field(field, value)?; + if value.len() > max_bytes { + return Err(GatewayError::Internal(format!( + "smtp {field} exceeds the allowed size" + ))); + } + Ok(()) +} + +fn validate_smtp_body_field(field: &str, value: &str) -> Result<(), GatewayError> { + // Bodies are base64 encoded before DATA is written, so line breaks and + // tabs are valid content. NUL is still rejected because it is not valid + // textual mail content and can confuse downstream gateways. + if value.bytes().any(|byte| byte == 0) { + return Err(GatewayError::Internal(format!( + "smtp {field} contains a forbidden NUL byte" + ))); + } + if value.len() > SMTP_MAX_BODY_BYTES { + return Err(GatewayError::Internal(format!( + "smtp {field} exceeds the allowed size" + ))); + } + Ok(()) +} + +fn validate_smtp_address(field: &str, value: &str) -> Result<(), GatewayError> { + validate_smtp_bounded_field(field, value, SMTP_MAX_ADDRESS_BYTES)?; + if value.is_empty() || value.trim() != value || value.chars().any(char::is_whitespace) { + return Err(GatewayError::Internal(format!( + "smtp {field} must be a single mailbox address" + ))); + } + // Addresses are inserted inside SMTP angle brackets. Reject delimiters + // that could turn one envelope/header value into multiple fields. + if value + .bytes() + .any(|byte| matches!(byte, b'<' | b'>' | b',' | b';' | b'"' | b'\\')) + { + return Err(GatewayError::Internal(format!( + "smtp {field} contains invalid mailbox delimiters" + ))); + } + let Some((local, domain)) = value.split_once('@') else { + return Err(GatewayError::Internal(format!( + "smtp {field} must contain a mailbox domain" + ))); + }; + if local.is_empty() || domain.is_empty() || domain.contains('@') { + return Err(GatewayError::Internal(format!( + "smtp {field} must contain a valid mailbox domain" + ))); + } + Ok(()) +} + +fn validate_smtp_auth_config(config: &SmtpDeliveryConfig) -> Result<(), GatewayError> { + let username = config.user.as_deref().map(str::trim); + let has_username = username.is_some_and(|value| !value.is_empty()); + let has_password = config.password.is_some(); + if has_username != has_password { + return Err(GatewayError::Internal( + "smtp username and password must be configured together".to_string(), + )); + } + if let Some(username) = username.filter(|value| !value.is_empty()) { + validate_smtp_bounded_field("user", username, SMTP_MAX_USERNAME_BYTES)?; + if !config.use_tls && !config.use_ssl { + return Err(GatewayError::Internal( + "smtp authentication requires TLS or SSL encryption".to_string(), + )); + } + } + if let Some(password) = config.password.as_deref() { + validate_smtp_bounded_field("password", password, SMTP_MAX_PASSWORD_BYTES)?; + } + Ok(()) +} + +fn validate_smtp_config(config: &SmtpDeliveryConfig) -> Result<(), GatewayError> { + validate_smtp_bounded_field("host", &config.host, SMTP_MAX_HOST_BYTES)?; + if config.host.is_empty() || config.host.trim() != config.host { + return Err(GatewayError::Internal( + "smtp host must not be empty or padded".to_string(), + )); + } + if config.host.chars().any(char::is_whitespace) { + return Err(GatewayError::Internal( + "smtp host contains invalid whitespace".to_string(), + )); + } + if config.port == 0 { + return Err(GatewayError::Internal( + "smtp port must be non-zero".to_string(), + )); + } + if config.use_tls && config.use_ssl { + return Err(GatewayError::Internal( + "smtp TLS and SSL modes cannot both be enabled".to_string(), + )); + } + validate_smtp_address("from_email", &config.from_email)?; + validate_smtp_bounded_field("from_name", &config.from_name, SMTP_MAX_HEADER_VALUE_BYTES)?; + validate_smtp_auth_config(config) +} + +fn validate_smtp_delivery_inputs( + config: &SmtpDeliveryConfig, + email: &ComposedEmail, +) -> Result<(), GatewayError> { + validate_smtp_config(config)?; + validate_smtp_address("to_email", &email.to_email)?; + validate_smtp_bounded_field("subject", &email.subject, SMTP_MAX_HEADER_VALUE_BYTES)?; + validate_smtp_body_field("html_body", &email.html_body)?; + validate_smtp_body_field("text_body", &email.text_body) +} + pub(crate) fn system_config_u16(value: Option<&serde_json::Value>, default: u16) -> u16 { match value { Some(serde_json::Value::Number(value)) => value @@ -132,8 +329,42 @@ fn resolve_server_name(host: &str) -> Result Result { - let stream = std::net::TcpStream::connect((config.host.as_str(), config.port)) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + use std::net::ToSocketAddrs; + let addresses = (config.host.as_str(), config.port) + .to_socket_addrs() + .map_err(|err| GatewayError::Internal(err.to_string()))? + .take(16) + .collect::>(); + if addresses.is_empty() { + return Err(GatewayError::Internal( + "smtp host did not resolve to an address".to_string(), + )); + } + let deadline = std::time::Instant::now() + .checked_add(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS)) + .unwrap_or_else(std::time::Instant::now); + let mut last_error = None; + let mut stream = None; + for address in addresses { + let remaining = deadline.saturating_duration_since(std::time::Instant::now()); + if remaining.is_zero() { + break; + } + match std::net::TcpStream::connect_timeout(&address, remaining) { + Ok(candidate) => { + stream = Some(candidate); + break; + } + Err(err) => last_error = Some(err), + } + } + let stream = stream.ok_or_else(|| { + GatewayError::Internal( + last_error + .map(|err| err.to_string()) + .unwrap_or_else(|| "smtp connection timed out".to_string()), + ) + })?; stream .set_read_timeout(Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))) .map_err(|err| GatewayError::Internal(err.to_string()))?; @@ -155,38 +386,109 @@ fn wrap_tls_stream( fn smtp_read_response(reader: &mut T) -> Result<(u16, String), GatewayError> { let mut message = String::new(); - let code = loop { - let parsed_code; - let continuation; - let trimmed; - { - let mut line = String::new(); - let bytes = reader - .read_line(&mut line) - .map_err(|err| GatewayError::Internal(err.to_string()))?; - if bytes == 0 { + let mut expected_code = None; + for line_number in 0..SMTP_MAX_RESPONSE_LINES { + let mut line = Vec::new(); + let bytes = read_smtp_response_line(reader, &mut line) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if bytes == 0 { + return Err(GatewayError::Internal( + "smtp connection closed unexpectedly".to_string(), + )); + } + if line.len() > SMTP_MAX_RESPONSE_LINE_BYTES { + return Err(GatewayError::Internal( + "smtp response line exceeds the allowed size".to_string(), + )); + } + + // Strip only the protocol line ending. The remaining bytes are kept + // for diagnostics after validating that they are UTF-8. + while matches!(line.last(), Some(b'\r' | b'\n')) { + line.pop(); + } + if line.len() < 3 || !line[..3].iter().all(|byte| byte.is_ascii_digit()) { + return Err(GatewayError::Internal("invalid smtp response".to_string())); + } + let parsed_code = u16::from(line[0] - b'0') * 100 + + u16::from(line[1] - b'0') * 10 + + u16::from(line[2] - b'0'); + if let Some(expected_code) = expected_code { + if parsed_code != expected_code { return Err(GatewayError::Internal( - "smtp connection closed unexpectedly".to_string(), + "smtp response continuation code changed".to_string(), )); } - trimmed = line.trim_end_matches(['\r', '\n']).to_string(); - if trimmed.len() < 3 { - return Err(GatewayError::Internal("invalid smtp response".to_string())); - } - parsed_code = trimmed[..3] - .parse::() - .map_err(|err| GatewayError::Internal(err.to_string()))?; - continuation = trimmed.as_bytes().get(3).copied() == Some(b'-'); + } else { + expected_code = Some(parsed_code); + } + let separator = line.get(3).copied().unwrap_or(b' '); + if separator != b'-' && separator != b' ' { + return Err(GatewayError::Internal("invalid smtp response".to_string())); + } + + let trimmed = std::str::from_utf8(&line) + .map_err(|_| GatewayError::Internal("smtp response is not valid UTF-8".to_string()))?; + let additional = trimmed.len() + usize::from(!message.is_empty()); + if message + .len() + .checked_add(additional) + .is_none_or(|length| length > SMTP_MAX_RESPONSE_BYTES) + { + return Err(GatewayError::Internal( + "smtp response exceeds the allowed size".to_string(), + )); } if !message.is_empty() { message.push('\n'); } - message.push_str(&trimmed); - if !continuation { - break parsed_code; + message.push_str(trimmed); + + if separator != b'-' { + return Ok((parsed_code, message)); } - }; - Ok((code, message)) + + if line_number + 1 == SMTP_MAX_RESPONSE_LINES { + return Err(GatewayError::Internal( + "smtp response has too many continuation lines".to_string(), + )); + } + } + + Err(GatewayError::Internal( + "smtp response has too many continuation lines".to_string(), + )) +} + +/// Read one SMTP response line without allowing `BufRead::read_until` to +/// allocate an attacker-controlled amount of memory before a size check. +fn read_smtp_response_line( + reader: &mut T, + line: &mut Vec, +) -> std::io::Result { + loop { + let buffered = reader.fill_buf()?; + if buffered.is_empty() { + return Ok(line.len()); + } + let newline = buffered.iter().position(|byte| *byte == b'\n'); + let take = newline.map_or(buffered.len(), |index| index + 1); + if line + .len() + .checked_add(take) + .is_none_or(|length| length > SMTP_MAX_RESPONSE_LINE_BYTES) + { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "smtp response line exceeds the allowed size", + )); + } + line.extend_from_slice(&buffered[..take]); + reader.consume(take); + if newline.is_some() { + return Ok(line.len()); + } + } } fn smtp_expect( @@ -197,11 +499,42 @@ fn smtp_expect( if allowed_codes.contains(&code) { return Ok(message); } + let message = sanitize_smtp_diagnostic(&message); Err(GatewayError::Internal(format!( "unexpected smtp response {code}: {message}" ))) } +/// SMTP responses are controlled by a remote server. Keep diagnostics useful +/// while preventing terminal escapes, log/UI line injection, and oversized +/// error payloads from crossing the API boundary. +fn sanitize_smtp_diagnostic(message: &str) -> String { + let mut sanitized = String::new(); + let mut previous_space = false; + for character in message.chars() { + if character == '\u{1b}' || character.is_control() { + if !previous_space { + sanitized.push(' '); + previous_space = true; + } + continue; + } + if sanitized.len() + character.len_utf8() > SMTP_MAX_DIAGNOSTIC_BYTES { + break; + } + if character.is_whitespace() { + if !previous_space { + sanitized.push(' '); + previous_space = true; + } + } else { + sanitized.push(character); + previous_space = false; + } + } + sanitized.trim().to_string() +} + fn smtp_write_line(writer: &mut T, line: &str) -> Result<(), GatewayError> { writer .write_all(line.as_bytes()) @@ -223,7 +556,11 @@ fn smtp_send_command( smtp_expect(reader, allowed_codes) } -fn build_email_message(config: &SmtpDeliveryConfig, email: &ComposedEmail) -> String { +fn build_email_message( + config: &SmtpDeliveryConfig, + email: &ComposedEmail, +) -> Result { + validate_smtp_delivery_inputs(config, email)?; let boundary = format!("aether-{}", uuid::Uuid::new_v4().simple()); let text_body = wrap_base64(&base64::engine::general_purpose::STANDARD.encode(email.text_body.as_bytes())); @@ -238,11 +575,17 @@ fn build_email_message(config: &SmtpDeliveryConfig, email: &ComposedEmail) -> St config.from_email ) }; - format!( + let message = format!( "From: {from_header}\r\nTo: <{to_email}>\r\nSubject: {subject}\r\nMIME-Version: 1.0\r\nContent-Type: multipart/alternative; boundary=\"{boundary}\"\r\n\r\n--{boundary}\r\nContent-Type: text/plain; charset=\"utf-8\"\r\nContent-Transfer-Encoding: base64\r\n\r\n{text_body}--{boundary}\r\nContent-Type: text/html; charset=\"utf-8\"\r\nContent-Transfer-Encoding: base64\r\n\r\n{html_body}--{boundary}--\r\n", to_email = email.to_email, subject = encode_mime_header(&email.subject), - ) + ); + if message.len() > SMTP_MAX_MESSAGE_BYTES { + return Err(GatewayError::Internal( + "smtp message exceeds the allowed size".to_string(), + )); + } + Ok(message) } fn smtp_authenticate( @@ -257,6 +600,11 @@ fn smtp_authenticate( else { return Ok(()); }; + if !config.use_tls && !config.use_ssl { + return Err(GatewayError::Internal( + "smtp authentication requires TLS or SSL encryption".to_string(), + )); + } let password = config.password.as_deref().unwrap_or(""); smtp_send_command(reader, "AUTH LOGIN", &[334])?; smtp_send_command( @@ -277,6 +625,9 @@ fn smtp_deliver_message( config: &SmtpDeliveryConfig, email: &ComposedEmail, ) -> Result<(), GatewayError> { + // Keep this check next to command construction for callers that bypass + // the async delivery wrapper. + validate_smtp_delivery_inputs(config, email)?; smtp_send_command( reader, &format!("MAIL FROM:<{}>", config.from_email), @@ -288,7 +639,7 @@ fn smtp_deliver_message( &[250, 251], )?; smtp_send_command(reader, "DATA", &[354])?; - let message = build_email_message(config, email); + let message = build_email_message(config, email)?; reader .get_mut() .write_all(message.as_bytes()) @@ -379,3 +730,135 @@ fn probe_smtp_connection_blocking(config: SmtpDeliveryConfig) -> Result<(), Gate let _ = smtp_send_command(&mut reader, "QUIT", &[221]); Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + fn config() -> SmtpDeliveryConfig { + SmtpDeliveryConfig { + host: "smtp.example.com".to_string(), + port: 587, + user: Some("user@example.com".to_string()), + password: Some("password".to_string()), + use_tls: true, + use_ssl: false, + from_email: "sender@example.com".to_string(), + from_name: "Aether".to_string(), + } + } + + fn email() -> ComposedEmail { + ComposedEmail { + to_email: "recipient@example.com".to_string(), + subject: "Subject".to_string(), + html_body: "

hello

".to_string(), + text_body: "hello".to_string(), + } + } + + #[test] + fn rejects_crlf_in_smtp_envelope_and_header_fields() { + let mut malicious_config = config(); + malicious_config.from_email = + "sender@example.com\r\nRCPT TO:".to_string(); + let error = validate_smtp_delivery_inputs(&malicious_config, &email()) + .expect_err("CRLF in an envelope address must be rejected"); + assert!(format!("{error:?}").contains("from_email")); + + let mut malicious_email = email(); + malicious_email.subject = "Subject\nX-Injected: yes".to_string(); + let error = validate_smtp_delivery_inputs(&config(), &malicious_email) + .expect_err("CRLF in a header value must be rejected"); + assert!(format!("{error:?}").contains("subject")); + } + + #[test] + fn allows_normal_smtp_values() { + assert!(validate_smtp_delivery_inputs(&config(), &email()).is_ok()); + } + + #[test] + fn rejects_authentication_over_plaintext_smtp() { + let mut insecure = config(); + insecure.use_tls = false; + insecure.use_ssl = false; + let error = validate_smtp_delivery_inputs(&insecure, &email()) + .expect_err("SMTP credentials must never be sent over plaintext"); + assert!(format!("{error:?}").contains("requires TLS or SSL")); + } + + #[test] + fn rejects_malformed_mailboxes_and_non_protocol_delimiters() { + let mut malicious = email(); + malicious.to_email = "recipient@example.com>\x01RCPT TO:".to_string(); + let error = validate_smtp_delivery_inputs(&config(), &malicious) + .expect_err("control characters and envelope delimiters must be rejected"); + assert!(format!("{error:?}").contains("to_email")); + + let mut malformed = email(); + malformed.to_email = "not-an-email".to_string(); + assert!(validate_smtp_delivery_inputs(&config(), &malformed).is_err()); + } + + #[test] + fn bounds_message_bodies_before_smtp_submission() { + let mut oversized = email(); + oversized.html_body = "x".repeat(SMTP_MAX_BODY_BYTES + 1); + let error = validate_smtp_delivery_inputs(&config(), &oversized) + .expect_err("oversized message bodies must be rejected"); + assert!(format!("{error:?}").contains("html_body")); + + let mut textual = email(); + textual.text_body = "line one\nline two\t✓".to_string(); + assert!(validate_smtp_delivery_inputs(&config(), &textual).is_ok()); + + textual.text_body.push('\0'); + assert!(validate_smtp_delivery_inputs(&config(), &textual).is_err()); + } + + #[test] + fn sanitizes_remote_response_diagnostics() { + let mut reader = std::io::BufReader::new("550 bad\u{1b}[31m\r\n".as_bytes()); + let error = smtp_expect(&mut reader, &[250]).expect_err("unexpected response must fail"); + let GatewayError::Internal(message) = error else { + panic!("expected internal SMTP error"); + }; + assert!(!message.contains('\u{1b}')); + assert!(!message.contains('\n')); + assert!(message.len() < SMTP_MAX_DIAGNOSTIC_BYTES); + } + + #[test] + fn smtp_response_rejects_non_ascii_status_prefix_without_panicking() { + let mut reader = std::io::BufReader::new("é00 greeting\r\n".as_bytes()); + let error = smtp_read_response(&mut reader) + .expect_err("a non-ASCII status prefix must be rejected"); + assert!(format!("{error:?}").contains("invalid smtp response")); + } + + #[test] + fn smtp_response_rejects_oversized_lines_before_allocating_unbounded_memory() { + let mut input = vec![b'2'; SMTP_MAX_RESPONSE_LINE_BYTES + 1]; + input.push(b'\n'); + let mut reader = std::io::BufReader::new(input.as_slice()); + let error = smtp_read_response(&mut reader).expect_err("oversized line must be rejected"); + assert!(format!("{error:?}").contains("response line exceeds")); + } + + #[test] + fn smtp_response_bounds_continuation_lines_and_code_changes() { + let repeated = (0..=SMTP_MAX_RESPONSE_LINES) + .map(|_| "250-more\r\n") + .collect::(); + let mut reader = std::io::BufReader::new(repeated.as_bytes()); + let error = smtp_read_response(&mut reader) + .expect_err("too many continuation lines must be rejected"); + assert!(format!("{error:?}").contains("too many continuation lines")); + + let mut reader = std::io::BufReader::new("250-more\r\n550 done\r\n".as_bytes()); + let error = smtp_read_response(&mut reader) + .expect_err("continuation response code changes must be rejected"); + assert!(format!("{error:?}").contains("continuation code changed")); + } +} diff --git a/apps/aether-gateway/src/error.rs b/apps/aether-gateway/src/error.rs index c0f76ff5f..92c62078a 100644 --- a/apps/aether-gateway/src/error.rs +++ b/apps/aether-gateway/src/error.rs @@ -3,6 +3,7 @@ use axum::http::{Response, StatusCode}; use axum::response::IntoResponse; use axum::Json; use serde_json::json; +use sha2::{Digest, Sha256}; use tracing::warn; use crate::ai_serving::AiSurfaceFinalizeError; @@ -33,6 +34,9 @@ pub(crate) enum GatewayError { status: StatusCode, message: String, }, + PlanUsageLimited(crate::plan_usage_policy::PlanUsagePolicyRejection), + LastActiveAdminUpdateDenied, + LastActiveAdminDeleteDenied, Internal(String), } @@ -43,6 +47,10 @@ impl GatewayError { | Self::ControlUnavailable { message, .. } | Self::Client { message, .. } | Self::Internal(message) => message, + Self::PlanUsageLimited(rejection) => format!( + "subscription plan {} limit {} reached for {} window; retry after {} seconds", + rejection.metric, rejection.limit, rejection.window, rejection.retry_after + ), Self::LocalExecutionPlanningTimeout { phase, timeout_ms, .. } => { @@ -55,6 +63,8 @@ impl GatewayError { } => { format!("gateway admission gate {gate} timed out after {queue_budget_ms}ms") } + Self::LastActiveAdminUpdateDenied => "不能降级或停用最后一个管理员账户".to_string(), + Self::LastActiveAdminDeleteDenied => "不能删除最后一个管理员账户".to_string(), } } } @@ -63,7 +73,13 @@ impl IntoResponse for GatewayError { fn into_response(self) -> Response { match self { Self::UpstreamUnavailable { trace_id, message } => { - warn!(trace_id = %trace_id, error = %message, "gateway proxy unavailable"); + let error_fingerprint = gateway_error_fingerprint(&message); + warn!( + trace_id = %trace_id, + error_fingerprint, + error_length = message.len(), + "gateway proxy unavailable" + ); let body = Json(json!({ "error": { "message": "gateway proxy unavailable", @@ -81,7 +97,13 @@ impl IntoResponse for GatewayError { response } Self::ControlUnavailable { trace_id, message } => { - warn!(trace_id = %trace_id, error = %message, "gateway control unavailable"); + let error_fingerprint = gateway_error_fingerprint(&message); + warn!( + trace_id = %trace_id, + error_fingerprint, + error_length = message.len(), + "gateway control unavailable" + ); let body = Json(json!({ "error": { "message": "gateway control unavailable", @@ -162,19 +184,59 @@ impl IntoResponse for GatewayError { })), ) .into_response(), - Self::Internal(message) => ( - StatusCode::INTERNAL_SERVER_ERROR, + Self::PlanUsageLimited(rejection) => ( + StatusCode::TOO_MANY_REQUESTS, Json(json!({ "error": { - "message": message, + "type": "plan_usage_limit_exceeded", + "message": "套餐使用限制已达到上限,请稍后重试", + "details": { + "metric": rejection.metric, + "window": rejection.window, + "limit": rejection.limit, + "retry_after": rejection.retry_after, + } } })), ) .into_response(), + Self::LastActiveAdminUpdateDenied => ( + StatusCode::BAD_REQUEST, + Json(json!({ "detail": "不能降级或停用最后一个管理员账户" })), + ) + .into_response(), + Self::LastActiveAdminDeleteDenied => ( + StatusCode::BAD_REQUEST, + Json(json!({ "detail": "不能删除最后一个管理员账户" })), + ) + .into_response(), + Self::Internal(message) => { + let error_fingerprint = gateway_error_fingerprint(&message); + tracing::error!( + event_name = "gateway_internal_error", + error_fingerprint, + error_length = message.len(), + "internal gateway error hidden from client" + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ + "error": { + "message": "internal server error", + } + })), + ) + .into_response() + } } } } +fn gateway_error_fingerprint(message: &str) -> String { + let digest = Sha256::digest(message.as_bytes()); + format!("{:x}", digest)[..16].to_string() +} + impl From for GatewayError { fn from(error: AiSurfaceFinalizeError) -> Self { GatewayError::Internal(error.0) @@ -183,12 +245,41 @@ impl From for GatewayError { #[cfg(test)] mod tests { + use axum::body::to_bytes; use axum::http::{header::RETRY_AFTER, StatusCode}; use axum::response::IntoResponse; use crate::constants::TRACE_ID_HEADER; - use super::GatewayError; + use super::{gateway_error_fingerprint, GatewayError}; + + #[tokio::test] + async fn internal_errors_do_not_expose_internal_details() { + let response = GatewayError::Internal( + "database connection failed: password=internal-secret".to_string(), + ) + .into_response(); + + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("internal error response body should read"); + let payload: serde_json::Value = + serde_json::from_slice(&body).expect("internal error response should be JSON"); + assert_eq!(payload["error"]["message"], "internal server error"); + assert!(!String::from_utf8_lossy(&body).contains("internal-secret")); + } + + #[test] + fn error_fingerprints_are_stable_without_containing_source_text() { + let secret_error = "postgresql://admin:database-secret@db.internal/aether"; + let fingerprint = gateway_error_fingerprint(secret_error); + + assert_eq!(fingerprint, gateway_error_fingerprint(secret_error)); + assert_eq!(fingerprint.len(), 16); + assert!(fingerprint.bytes().all(|byte| byte.is_ascii_hexdigit())); + assert!(!fingerprint.contains("database-secret")); + } #[test] fn admission_timeout_returns_429_with_retry_after_without_panicking() { diff --git a/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs b/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs new file mode 100644 index 000000000..7bd968827 --- /dev/null +++ b/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs @@ -0,0 +1,567 @@ +//! Terminal settlement for a local stream attempt whose future is dropped +//! mid-flight. +//! +//! A local stream attempt writes its `usage` row and its `request_candidates` +//! slot as `pending` before it dispatches to the provider, then keeps running +//! inside the downstream request future. When the client disconnects, axum drops +//! that future: the remaining `.await`s never resume and nothing settles either +//! row. They stay `pending` until the maintenance sweeper rewrites them as a 504 +//! timeout roughly ten minutes later, which loses the real outcome and the real +//! latency. +//! +//! The stream transport therefore keeps a guard alive across the window between +//! the `pending` write and terminal settlement, and settles the attempt from +//! `Drop` when that window is left by cancellation instead of by a terminal +//! state. + +use std::sync::Arc; +use std::time::Instant; + +use aether_contracts::ExecutionPlan; +use aether_data_contracts::repository::candidates::RequestCandidateStatus; +use aether_scheduler_core::SchedulerRequestCandidateStatusUpdate; +use aether_usage_runtime::{ + build_usage_event_data_seed_describing_request_bodies, UsageEvent, UsageEventData, + UsageEventType, +}; +use serde_json::{json, Value}; +use tracing::warn; + +use crate::clock::current_unix_ms as current_request_candidate_unix_ms; +use crate::execution_runtime::attempt_lifecycle::CLIENT_CANCELLED_STATUS_CODE; +use crate::execution_runtime::transport_failure::StreamCandidateWatchdogProgress; +use crate::log_ids::short_request_id; +use crate::request_candidate_runtime::{ + record_local_request_candidate_status_snapshot, LocalRequestCandidateStatusSnapshot, +}; +use crate::request_diagnostics::{ + attach_request_diagnostics_to_report_context, current_request_diagnostics, RequestDiagnostics, +}; +use crate::AppState; + +fn elapsed_ms_since(started_at: Instant) -> u64 { + started_at.elapsed().as_millis().min(u128::from(u64::MAX)) as u64 +} + +/// The facts the guard needs to settle the attempt it is watching. +/// +/// This is held for the whole attempt, so it is deliberately free of request +/// bodies. A request body can be megabytes, and holding one per in-flight +/// attempt would cost far more than the row it settles: the usage seed is built +/// with [`build_usage_event_data_seed_describing_request_bodies`], which derives +/// every capture state, body reference and derived request fact from the real +/// plan and report context but keeps neither body. The terminal write it +/// produces therefore preserves the capture the `pending` write recorded instead +/// of clearing it. +struct ArmedAttempt { + request_id: String, + candidate_id: Option, + candidate: Option, + // Boxed: the guard lives inside the stream request future, which is already + // very large, and `UsageEventData` is a wide struct. + usage_seed: Option>, + request_diagnostics: Option>, + candidate_started_unix_ms: u64, + candidate_started_at: Instant, +} + +/// Settles an attempt as cancelled when its future is dropped before the +/// transport reaches a terminal state. +/// +/// The guard is created disarmed and stays inert until [`Self::arm`] is called, +/// so an attempt that is dropped before it owns any `pending` row does not grow +/// a settlement row it never had. The owner disarms it as soon as the attempt +/// completes, whichever way it completes: from that point terminal settlement +/// belongs to the transport (for streams, to the stream finalizer that lives in +/// the response body), and the guard must not write a second terminal state. +/// +/// A stream candidate also runs under a first-byte watchdog that drops the +/// attempt future when it gives up. That drop is not a client disconnect and the +/// watchdog settles the attempt itself, so the guard stands down for it. +pub(crate) struct AttemptCancellationGuard { + state: AppState, + error_type: &'static str, + error_message: &'static str, + watchdog: Option>, + armed: Option, +} + +impl AttemptCancellationGuard { + pub(crate) fn disarmed( + state: &AppState, + error_type: &'static str, + error_message: &'static str, + ) -> Self { + Self { + state: state.clone(), + error_type, + error_message, + watchdog: StreamCandidateWatchdogProgress::current(), + armed: None, + } + } + + /// Takes ownership of the attempt's settlement until it is disarmed. + pub(crate) fn arm( + &mut self, + plan: &ExecutionPlan, + report_context: Option<&Value>, + candidate: Option<&LocalRequestCandidateStatusSnapshot>, + candidate_started_unix_ms: u64, + candidate_started_at: Instant, + ) { + let usage_seed = self.state.usage_runtime.is_enabled().then(|| { + Box::new(build_usage_event_data_seed_describing_request_bodies( + plan, + report_context, + )) + }); + self.armed = Some(ArmedAttempt { + request_id: plan.request_id.clone(), + candidate_id: plan.candidate_id.clone(), + candidate: candidate.cloned(), + usage_seed, + request_diagnostics: current_request_diagnostics(), + candidate_started_unix_ms, + candidate_started_at, + }); + } + + pub(crate) fn disarm(&mut self) { + self.armed = None; + } +} + +/// Writes the candidate terminal row and the terminal usage event for an attempt +/// that never reached its own terminal path. +async fn settle_cancelled_attempt( + state: AppState, + armed: ArmedAttempt, + error_type: &'static str, + error_message: &'static str, +) { + let ArmedAttempt { + request_id, + candidate_id: _, + candidate, + usage_seed, + request_diagnostics, + candidate_started_unix_ms, + candidate_started_at, + } = armed; + let terminal_unix_ms = current_request_candidate_unix_ms(); + let latency_ms = elapsed_ms_since(candidate_started_at); + + if let Some(candidate) = candidate.as_ref() { + record_local_request_candidate_status_snapshot( + &state, + candidate, + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Cancelled, + status_code: Some(CLIENT_CANCELLED_STATUS_CODE), + error_type: Some(error_type.to_string()), + error_message: Some(error_message.to_string()), + latency_ms: Some(latency_ms), + started_at_unix_ms: Some(candidate_started_unix_ms), + finished_at_unix_ms: Some(terminal_unix_ms), + }, + ) + .await; + } + + let Some(usage_data) = usage_seed else { + return; + }; + let mut usage_data = *usage_data; + // The seed was built when the attempt was armed, so it predates the + // diagnostics it should carry. Attaching them to the seed's metadata is the + // same write the report context would have carried into a seed built here: + // both land the same keys in the same object. + usage_data.request_metadata = attach_request_diagnostics_to_report_context( + usage_data.request_metadata.take(), + request_diagnostics.as_ref(), + ); + usage_data.status_code = Some(CLIENT_CANCELLED_STATUS_CODE); + usage_data.error_message = Some(error_message.to_string()); + usage_data.error_category = Some("cancelled".to_string()); + usage_data.response_time_ms = Some(latency_ms); + let error_body = json!({ + "error": { + "type": error_type, + "message": error_message, + "code": CLIENT_CANCELLED_STATUS_CODE + } + }); + usage_data.response_headers = Some(json!({"content-type": "application/json"})); + usage_data.response_body = Some(error_body.clone()); + usage_data.client_response_headers = Some(json!({"content-type": "application/json"})); + usage_data.client_response_body = Some(error_body); + + state + .usage_runtime + .record_terminal_event_direct( + state.usage_lifecycle_data_state().as_ref(), + UsageEvent::new(UsageEventType::Cancelled, request_id, usage_data), + ) + .await; +} + +impl Drop for AttemptCancellationGuard { + fn drop(&mut self) { + let Some(armed) = self.armed.take() else { + return; + }; + if self + .watchdog + .as_ref() + .is_some_and(|watchdog| watchdog.abandoned()) + { + return; + } + let state = self.state.clone(); + let error_type = self.error_type; + let error_message = self.error_message; + // `Drop` cannot await, and the settlement writes touch the database. + // Hand them to the runtime so they survive the dropped request future. + let Ok(handle) = tokio::runtime::Handle::try_current() else { + warn!( + event_name = "local_attempt_cancellation_guard_no_runtime", + log_type = "ops", + request_id = %short_request_id(armed.request_id.as_str()), + candidate_id = ?armed.candidate_id, + error_type, + "gateway could not settle dropped local attempt because no Tokio runtime is available" + ); + return; + }; + handle.spawn(async move { + settle_cancelled_attempt(state, armed, error_type, error_message).await; + }); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use aether_contracts::RequestBody; + use aether_data::repository::candidates::InMemoryRequestCandidateRepository; + use aether_data::repository::usage::InMemoryUsageReadRepository; + use aether_data_contracts::repository::candidates::RequestCandidateReadRepository; + use aether_data_contracts::repository::usage::{ + StoredRequestUsageAudit, UsageBodyCaptureState, UsageReadRepository, UsageWriteRepository, + }; + use aether_usage_runtime::{ + build_lifecycle_usage_seed, build_pending_usage_record, UsageRuntimeConfig, + }; + use std::collections::BTreeMap; + use std::time::Duration; + + use crate::request_candidate_runtime::{ + ensure_execution_request_candidate_slot, snapshot_local_request_candidate_status, + }; + + const TEST_ERROR_TYPE: &str = "local_stream_attempt_cancelled"; + const TEST_ERROR_MESSAGE: &str = + "Local stream attempt was dropped before terminal finalization."; + + fn test_stream_plan(request_id: &str) -> ExecutionPlan { + ExecutionPlan { + request_id: request_id.to_string(), + candidate_id: None, + provider_name: Some("Anthropic".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-1".to_string(), + method: "POST".to_string(), + url: "https://example.test/v1/messages".to_string(), + headers: BTreeMap::new(), + content_type: Some("application/json".to_string()), + content_encoding: None, + body: RequestBody::from_json(json!({"stream": true, "service_tier": "priority"})), + stream: true, + client_api_format: "claude:messages".to_string(), + provider_api_format: "claude:messages".to_string(), + model_name: Some("claude-sonnet-4-5".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + } + } + + fn test_report_context() -> Option { + Some(json!({ + "candidate_index": 0, + "retry_index": 0, + "user_id": "user-cancel", + "api_key_id": "api-key-cancel", + "client_api_format": "claude:messages", + "provider_api_format": "claude:messages", + "request_path": "/v1/messages", + "request_path_and_query": "/v1/messages?beta=true", + "upstream_url": "https://example.test/v1/messages", + "mapped_model": "claude-sonnet-4-5", + "original_request_body": {"stream": true, "messages": []}, + })) + } + + fn test_state( + usage_repository: &Arc, + request_candidate_repository: &Arc, + ) -> AppState { + AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests( + Arc::clone(request_candidate_repository), + Arc::clone(usage_repository), + ), + ) + .with_usage_runtime_for_tests(UsageRuntimeConfig { + enabled: true, + ..UsageRuntimeConfig::default() + }) + } + + /// Writes the `pending` rows the same way a stream attempt does before it + /// dispatches to the provider, and returns the candidate slot snapshot the + /// attempt owns from that point on. + async fn record_pending_attempt( + state: &AppState, + plan: &mut ExecutionPlan, + report_context: &mut Option, + candidate_started_unix_ms: u64, + ) -> LocalRequestCandidateStatusSnapshot { + ensure_execution_request_candidate_slot(state, plan, report_context).await; + state.usage_runtime.record_pending( + state.usage_lifecycle_data_state().as_ref(), + build_lifecycle_usage_seed(plan, report_context.as_ref()), + ); + let snapshot = snapshot_local_request_candidate_status(plan, report_context.as_ref()) + .expect("attempt should own a candidate slot"); + record_local_request_candidate_status_snapshot( + state, + &snapshot, + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Pending, + status_code: None, + error_type: None, + error_message: None, + latency_ms: None, + started_at_unix_ms: Some(candidate_started_unix_ms), + finished_at_unix_ms: None, + }, + ) + .await; + snapshot + } + + async fn wait_for_usage_status( + usage_repository: &InMemoryUsageReadRepository, + request_id: &str, + status: &str, + ) -> Option { + for _ in 0..50 { + if let Some(usage) = usage_repository + .find_by_request_id(request_id) + .await + .expect("usage should read") + { + if usage.status == status { + return Some(usage); + } + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + None + } + + #[tokio::test] + async fn armed_guard_settles_a_dropped_attempt_as_cancelled() { + let usage_repository = Arc::new(InMemoryUsageReadRepository::default()); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = test_state(&usage_repository, &request_candidate_repository); + let mut plan = test_stream_plan("stream-cancel-guard-request"); + let mut report_context = test_report_context(); + let candidate_started_unix_ms = current_request_candidate_unix_ms(); + let snapshot = record_pending_attempt( + &state, + &mut plan, + &mut report_context, + candidate_started_unix_ms, + ) + .await; + + { + let mut guard = + AttemptCancellationGuard::disarmed(&state, TEST_ERROR_TYPE, TEST_ERROR_MESSAGE); + guard.arm( + &plan, + report_context.as_ref(), + Some(&snapshot), + candidate_started_unix_ms, + Instant::now(), + ); + } + + let usage = wait_for_usage_status( + usage_repository.as_ref(), + "stream-cancel-guard-request", + "cancelled", + ) + .await + .expect("cancelled usage should be recorded"); + assert_eq!(usage.billing_status, "void"); + assert_eq!(usage.status_code, Some(CLIENT_CANCELLED_STATUS_CODE)); + assert_eq!(usage.error_category.as_deref(), Some("cancelled")); + assert!(usage.response_time_ms.is_some()); + + let candidates = request_candidate_repository + .list_by_request_id("stream-cancel-guard-request") + .await + .expect("candidates should read"); + let candidate = candidates.first().expect("candidate row should exist"); + assert_eq!(candidate.status, RequestCandidateStatus::Cancelled); + assert_eq!(candidate.status_code, Some(CLIENT_CANCELLED_STATUS_CODE)); + assert_eq!(candidate.error_type.as_deref(), Some(TEST_ERROR_TYPE)); + assert!(candidate.finished_at_unix_ms.is_some()); + } + + /// The guard holds no request body, and the persistence boundary intentionally + /// rejects request/response capture material. A dropped-attempt settlement + /// must not re-introduce an inline body or a caller-controlled body reference. + #[tokio::test] + async fn settling_a_dropped_attempt_does_not_reintroduce_request_body_capture() { + let usage_repository = Arc::new(InMemoryUsageReadRepository::default()); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = test_state(&usage_repository, &request_candidate_repository); + let mut plan = test_stream_plan("stream-cancel-guard-capture"); + let mut report_context = test_report_context(); + let candidate_started_unix_ms = current_request_candidate_unix_ms(); + let snapshot = record_pending_attempt( + &state, + &mut plan, + &mut report_context, + candidate_started_unix_ms, + ) + .await; + // This deliberately supplies capture material to prove that the usage + // persistence boundary strips it before either lifecycle write stores it. + let captured_body = json!({"stream": true, "service_tier": "priority"}); + let mut capture = build_pending_usage_record( + &plan, + report_context.as_ref(), + current_request_candidate_unix_ms() / 1_000, + ) + .expect("pending usage record should build"); + capture.provider_request_body = Some(captured_body.clone()); + capture.provider_request_body_state = Some(UsageBodyCaptureState::Inline); + usage_repository + .upsert(capture) + .await + .expect("captured request body should upsert"); + + { + let mut guard = + AttemptCancellationGuard::disarmed(&state, TEST_ERROR_TYPE, TEST_ERROR_MESSAGE); + guard.arm( + &plan, + report_context.as_ref(), + Some(&snapshot), + candidate_started_unix_ms, + Instant::now(), + ); + } + + let usage = wait_for_usage_status( + usage_repository.as_ref(), + "stream-cancel-guard-capture", + "cancelled", + ) + .await + .expect("cancelled usage should be recorded"); + assert_eq!(usage.provider_request_body, None); + assert_eq!(usage.provider_request_body_ref, None); + assert_eq!(usage.provider_request_body_state, None); + } + + #[tokio::test] + async fn guard_stands_down_when_the_watchdog_abandons_the_attempt() { + let usage_repository = Arc::new(InMemoryUsageReadRepository::default()); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = test_state(&usage_repository, &request_candidate_repository); + let mut plan = test_stream_plan("stream-watchdog-guard-request"); + let mut report_context = test_report_context(); + let candidate_started_unix_ms = current_request_candidate_unix_ms(); + let snapshot = record_pending_attempt( + &state, + &mut plan, + &mut report_context, + candidate_started_unix_ms, + ) + .await; + + let watchdog = StreamCandidateWatchdogProgress::shared(); + Arc::clone(&watchdog) + .scope(async { + let mut guard = + AttemptCancellationGuard::disarmed(&state, TEST_ERROR_TYPE, TEST_ERROR_MESSAGE); + guard.arm( + &plan, + report_context.as_ref(), + Some(&snapshot), + candidate_started_unix_ms, + Instant::now(), + ); + // The watchdog gives up and takes over settlement before the + // abandoned attempt is dropped. + watchdog.mark_abandoned(); + }) + .await; + + assert!(wait_for_usage_status( + usage_repository.as_ref(), + "stream-watchdog-guard-request", + "cancelled", + ) + .await + .is_none()); + } + + #[tokio::test] + async fn disarmed_guard_leaves_the_attempt_pending() { + let usage_repository = Arc::new(InMemoryUsageReadRepository::default()); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = test_state(&usage_repository, &request_candidate_repository); + let mut plan = test_stream_plan("stream-disarmed-guard-request"); + let mut report_context = test_report_context(); + let candidate_started_unix_ms = current_request_candidate_unix_ms(); + let snapshot = record_pending_attempt( + &state, + &mut plan, + &mut report_context, + candidate_started_unix_ms, + ) + .await; + + { + let mut guard = + AttemptCancellationGuard::disarmed(&state, TEST_ERROR_TYPE, TEST_ERROR_MESSAGE); + guard.arm( + &plan, + report_context.as_ref(), + Some(&snapshot), + candidate_started_unix_ms, + Instant::now(), + ); + guard.disarm(); + } + + assert!(wait_for_usage_status( + usage_repository.as_ref(), + "stream-disarmed-guard-request", + "cancelled", + ) + .await + .is_none()); + } +} diff --git a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs index 12a7485b0..46ff852bd 100644 --- a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs +++ b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs @@ -36,10 +36,11 @@ use aether_contracts::{ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionT use aether_data_contracts::repository::candidates::RequestCandidateStatus; use aether_data_contracts::repository::usage::UsageBodyCaptureState; use aether_scheduler_core::SchedulerRequestCandidateStatusUpdate; +#[cfg(test)] +use aether_usage_runtime::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES; use aether_usage_runtime::{ build_lifecycle_usage_seed, build_stream_terminal_usage_payload_seed, build_terminal_usage_context_seed, stream_report_represents_failure, - DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, }; use base64::Engine as _; use serde_json::Value; @@ -433,7 +434,7 @@ impl AttemptBodyCapture { if bytes.is_empty() || self.truncated { return; } - let max_bytes = DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES; + let max_bytes = crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES; if self.buffer.len() >= max_bytes { self.truncated = true; return; @@ -447,11 +448,19 @@ impl AttemptBodyCapture { } pub(crate) fn encode(&self) -> (Option, Option) { - let body = (!self.buffer.is_empty()) - .then(|| base64::engine::general_purpose::STANDARD.encode(&self.buffer)); - let state = if self.truncated { + self.encode_with_limit(crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES) + } + + fn encode_with_limit( + &self, + max_bytes: usize, + ) -> (Option, Option) { + let captured = &self.buffer[..self.buffer.len().min(max_bytes)]; + let body = (!captured.is_empty()) + .then(|| base64::engine::general_purpose::STANDARD.encode(captured)); + let state = if self.truncated || captured.len() < self.buffer.len() { UsageBodyCaptureState::Truncated - } else if self.buffer.is_empty() { + } else if captured.is_empty() { UsageBodyCaptureState::None } else { UsageBodyCaptureState::Inline @@ -1418,15 +1427,13 @@ mod stage_tests { ); } - /// body capture 的编码状态。截断分支这里到不了:共享的 - /// `DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES` 是 `usize::MAX`, - /// 也就是默认不限长;截断只在 usage 侧把上限调低后才可能发生。 + /// body capture 的编码状态。Full 记录级别仍受 gateway 的硬上限约束, + /// 这样长连接不会把审计副本无限累积。 #[test] fn body_capture_encodes_inline_and_empty_states() { assert_eq!( - super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, - usize::MAX, - "the default capture limit is unbounded; truncation is not reachable here" + crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES, + crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES ); let mut capture = AttemptBodyCapture::default(); @@ -1454,6 +1461,20 @@ mod stage_tests { state, Some(aether_data_contracts::repository::usage::UsageBodyCaptureState::None) ); + + let defensive = AttemptBodyCapture { + buffer: b"abcdef".to_vec(), + truncated: false, + }; + let (body, state) = defensive.encode_with_limit(3); + let decoded = base64::engine::general_purpose::STANDARD + .decode(body.expect("bounded capture should be encoded")) + .expect("capture is valid base64"); + assert_eq!(decoded, b"abc"); + assert_eq!( + state, + Some(aether_data_contracts::repository::usage::UsageBodyCaptureState::Truncated) + ); } /// candidate 行的 error_type 映射:投递失败与供应商侧失败必须各有名字。 diff --git a/apps/aether-gateway/src/execution_runtime/chatgpt_web_image.rs b/apps/aether-gateway/src/execution_runtime/chatgpt_web_image.rs index ac336ba7b..5d09cd58a 100644 --- a/apps/aether-gateway/src/execution_runtime/chatgpt_web_image.rs +++ b/apps/aether-gateway/src/execution_runtime/chatgpt_web_image.rs @@ -1,16 +1,17 @@ use std::collections::{BTreeMap, BTreeSet}; use std::io::Error as IoError; +use std::net::{IpAddr, SocketAddr}; use std::time::{Duration, Instant}; use aether_admin::provider::quota::{ parse_chatgpt_web_conversation_init_response, quota_refresh_success_invalid_state, }; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; use aether_contracts::{ ExecutionPlan, ExecutionResult, ExecutionStreamTerminalSummary, ExecutionTelemetry, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile, ResponseBody, - StreamFrame, StreamFramePayload, StreamFrameType, - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, - TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY, + StreamFrame, StreamFramePayload, StreamFrameType, TRANSPORT_BACKEND_BROWSER_WREQ, + TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY, }; use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate; use aether_provider_pool::{ @@ -30,7 +31,10 @@ use crate::ai_serving::api::StreamingStandardTerminalObserver; use crate::clock::current_unix_secs; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::transport::{ - with_non_stream_total_timeout, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, + decode_base64_body_with_limit, format_upstream_request_error, json_value_fits_serialized_limit, + maximum_base64_len_for_decoded_limit, safe_transport_error_message, + serialize_json_body_with_limit, with_non_stream_total_timeout, DirectSyncExecutionRuntime, + ExecutionRuntimeTransportError, }; use crate::handlers::shared::{ sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot, @@ -53,6 +57,34 @@ const GPT_IMAGE2_TOKEN_MAX_PIXELS: u64 = 8_294_400; const GPT_IMAGE2_TOKEN_MAX_EDGE: u64 = 3_840; const GPT_IMAGE2_TOKEN_MAX_ASPECT_RATIO: u64 = 3; const GPT_IMAGE2_PARTIAL_IMAGE_OUTPUT_TOKENS: u64 = 100; +const CHATGPT_WEB_IMAGE_DOWNLOAD_MAX_REDIRECTS: usize = 10; +// A generated SSE response contains the image's base64 text plus JSON/event +// framing. Keep its decoded envelope bounded independently from the raw image +// limit, while retaining support for a raw image up to the default 64 MiB cap. +const CHATGPT_WEB_IMAGE_SSE_WRAPPER_OVERHEAD_BYTES: usize = 256 * 1024; +const CHATGPT_WEB_IMAGE_SSE_HARD_MAX_BYTES: usize = 128 * 1024 * 1024; +const CHATGPT_WEB_IMAGE_STREAM_CHUNK_BYTES: usize = 1024 * 1024; +const CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS: u64 = 10_000; +const CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS: u64 = 30_000; +const CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS: u64 = 300_000; +const CHATGPT_WEB_OPAQUE_ID_MAX_BYTES: usize = 256; +const CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES: usize = 32; +const CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES: usize = 64 * 1024; +const CHATGPT_WEB_IMAGE_UPLOAD_RESPONSE_LIMIT_BYTES: usize = 64 * 1024; +const CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES: usize = 32 * 1024; +const CHATGPT_WEB_IMAGE_MAX_MODEL_BYTES: usize = 256; +const CHATGPT_WEB_IMAGE_MAX_OPTION_BYTES: usize = 128; +const CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES: usize = 64 * 1024; +const CHATGPT_WEB_IMAGE_MAX_INPUT_IMAGES: usize = 16; +const CHATGPT_WEB_IMAGE_MAX_DIMENSION: u32 = 16_384; +// Values in a provider SSE/poll response are merged across the initial +// response and up to 24 follow-up polls. Keep the retained candidate set +// bounded independently of the per-response body limit so a peer cannot make +// the gateway grow memory over the lifetime of one request. +const CHATGPT_WEB_IMAGE_SUMMARY_MAX_ITEMS: usize = 256; +const CHATGPT_WEB_IMAGE_SUMMARY_MAX_DIRECT_URLS: usize = 16; +const CHATGPT_WEB_IMAGE_SUMMARY_MAX_ID_BYTES: usize = 1024 * 1024; +const CHATGPT_WEB_IMAGE_SUMMARY_MAX_TEXT_BYTES: usize = 64 * 1024; pub(crate) struct ChatGptWebImageStream { pub(crate) frame_stream: BoxStream<'static, Result>, @@ -94,6 +126,102 @@ struct WebImageSseSummary { last_text: Option, } +#[derive(Debug, Clone, Copy)] +enum WebImageSummaryCollection { + FileId, + SedimentId, + DirectUrl, +} + +impl WebImageSseSummary { + fn retained_item_count(&self) -> usize { + self.file_ids + .len() + .saturating_add(self.sediment_ids.len()) + .saturating_add(self.direct_urls.len()) + } + + fn retained_value_bytes(&self) -> usize { + saturating_string_bytes(&self.file_ids) + .saturating_add(saturating_string_bytes(&self.sediment_ids)) + .saturating_add(saturating_string_bytes(&self.direct_urls)) + } + + fn add_values(&mut self, collection: WebImageSummaryCollection, incoming: I) + where + I: IntoIterator, + { + for value in incoming { + self.add_value(collection, value); + } + } + + fn add_value(&mut self, collection: WebImageSummaryCollection, value: String) { + if value.is_empty() { + return; + } + let (max_items, collection_budget) = match collection { + WebImageSummaryCollection::FileId | WebImageSummaryCollection::SedimentId => ( + CHATGPT_WEB_IMAGE_SUMMARY_MAX_ITEMS, + CHATGPT_WEB_IMAGE_SUMMARY_MAX_ID_BYTES, + ), + WebImageSummaryCollection::DirectUrl if is_data_image_reference(&value) => { + (4, chatgpt_web_image_sse_envelope_limit_bytes()) + } + WebImageSummaryCollection::DirectUrl => ( + CHATGPT_WEB_IMAGE_SUMMARY_MAX_DIRECT_URLS, + CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES, + ), + }; + // Reject an over-sized value before it can become part of the retained + // set. The caller may have had to materialize it to parse a JSON + // field, but this prevents repeated poll responses from accumulating + // it and bounds the final synthetic SSE envelope. + if value.len() > collection_budget + || self.retained_item_count() >= CHATGPT_WEB_IMAGE_SUMMARY_MAX_ITEMS + { + return; + } + let values = match collection { + WebImageSummaryCollection::FileId => &self.file_ids, + WebImageSummaryCollection::SedimentId => &self.sediment_ids, + WebImageSummaryCollection::DirectUrl => &self.direct_urls, + }; + if values.len() >= max_items || values.iter().any(|existing| existing == &value) { + return; + } + let collection_bytes = match collection { + WebImageSummaryCollection::FileId => saturating_string_bytes(&self.file_ids), + WebImageSummaryCollection::SedimentId => saturating_string_bytes(&self.sediment_ids), + WebImageSummaryCollection::DirectUrl => saturating_string_bytes(&self.direct_urls), + }; + let total_budget = chatgpt_web_image_sse_envelope_limit_bytes() + .saturating_add(CHATGPT_WEB_IMAGE_SUMMARY_MAX_ID_BYTES); + if value.len() > collection_budget.saturating_sub(collection_bytes) + || value.len() > total_budget.saturating_sub(self.retained_value_bytes()) + { + return; + } + match collection { + WebImageSummaryCollection::FileId => self.file_ids.push(value), + WebImageSummaryCollection::SedimentId => self.sediment_ids.push(value), + WebImageSummaryCollection::DirectUrl => self.direct_urls.push(value), + } + } +} + +fn saturating_string_bytes(values: &[String]) -> usize { + values + .iter() + .fold(0usize, |total, value| total.saturating_add(value.len())) +} + +fn is_data_image_reference(value: &str) -> bool { + value + .get(..11) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case("data:image/")) +} + #[derive(Debug, Clone)] struct DownloadedImage { b64_json: String, @@ -102,6 +230,17 @@ struct DownloadedImage { height: Option, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum WebImageDownloadTrust { + UntrustedInput, + ProviderOutput, +} + +struct WebImageHttpPayload { + data: Vec, + content_type: Option, +} + pub(crate) async fn maybe_execute_chatgpt_web_image_sync( state: &AppState, plan: &ExecutionPlan, @@ -151,7 +290,7 @@ pub(crate) async fn maybe_execute_chatgpt_web_image_stream( Err(error) => return Err(error), }; Ok(Some(ChatGptWebImageStream { - frame_stream: execution_result_frame_stream(plan, &result, report_context), + frame_stream: execution_result_frame_stream(plan, &result, report_context)?, report_context: report_context.cloned(), })) } @@ -203,7 +342,7 @@ async fn execute_chatgpt_web_image( log_type = "debug", request_id = %plan.request_id, candidate_id = ?plan.candidate_id, - base_url = %base_url, + upstream_origin = %crate::handlers::shared::security_log_url_origin(&base_url), operation = %request.operation, image_count = request.images.len(), size = %request.size, @@ -315,7 +454,7 @@ async fn execute_chatgpt_web_image( ) }; - Ok(bytes_execution_result( + bytes_execution_result( plan, 200, BTreeMap::from([ @@ -324,7 +463,7 @@ async fn execute_chatgpt_web_image( ]), body.into_bytes(), started_at, - )) + ) } #[derive(Debug, Clone)] @@ -343,38 +482,105 @@ struct ChatGptWebImageRequest { impl ChatGptWebImageRequest { fn from_body(body: &Value) -> Result { - let text = |key: &str| { - body.get(key) - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - }; - let images = body - .get("images") - .and_then(Value::as_array) - .into_iter() - .flatten() - .filter_map(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .collect::>(); + let model = + bounded_chatgpt_web_text_field(body, "model", CHATGPT_WEB_IMAGE_MAX_MODEL_BYTES)? + .unwrap_or_else(|| "gpt-image-2".to_string()); + let web_model = + bounded_chatgpt_web_text_field(body, "web_model", CHATGPT_WEB_IMAGE_MAX_MODEL_BYTES)? + .unwrap_or_else(|| "gpt-5-5-thinking".to_string()); + let prompt = + bounded_chatgpt_web_text_field(body, "prompt", CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES)? + .unwrap_or_else(|| "Generate a high quality image.".to_string()); + let size = + bounded_chatgpt_web_text_field(body, "size", CHATGPT_WEB_IMAGE_MAX_OPTION_BYTES)? + .unwrap_or_else(|| "1024x1024".to_string()); + let ratio = + bounded_chatgpt_web_text_field(body, "ratio", CHATGPT_WEB_IMAGE_MAX_OPTION_BYTES)? + .unwrap_or_else(|| "1:1".to_string()); + let output_format = bounded_chatgpt_web_text_field( + body, + "output_format", + CHATGPT_WEB_IMAGE_MAX_OPTION_BYTES, + )? + .unwrap_or_else(|| "png".to_string()); + let quality = + bounded_chatgpt_web_text_field(body, "quality", CHATGPT_WEB_IMAGE_MAX_OPTION_BYTES)?; + + let mut images = Vec::new(); + if let Some(values) = body.get("images").and_then(Value::as_array) { + if values.len() > CHATGPT_WEB_IMAGE_MAX_INPUT_IMAGES { + return Err(chatgpt_web_image_request_field_too_large("images")); + } + for value in values { + let Some(value) = value.as_str() else { + continue; + }; + let value = value.trim(); + if value.is_empty() { + continue; + } + let max_bytes = if value + .get(..5) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case("data:")) + { + maximum_base64_len_for_decoded_limit( + chatgpt_web_image_raw_payload_limit_bytes(), + ) + .saturating_add(128) + } else { + CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES + }; + if value.len() > max_bytes { + return Err(chatgpt_web_image_request_field_too_large("image reference")); + } + images.push(value.to_string()); + } + } + let partial_images = json_u64(body.get("partial_images")).unwrap_or(0); + if partial_images > 3 { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image partial_images must be between 0 and 3".to_string(), + )); + } Ok(Self { operation: chatgpt_web_image_operation(body.get("operation")), - model: text("model").unwrap_or_else(|| "gpt-image-2".to_string()), - web_model: text("web_model").unwrap_or_else(|| "gpt-5-5-thinking".to_string()), - prompt: text("prompt").unwrap_or_else(|| "Generate a high quality image.".to_string()), - size: text("size").unwrap_or_else(|| "1024x1024".to_string()), - ratio: text("ratio").unwrap_or_else(|| "1:1".to_string()), - output_format: text("output_format").unwrap_or_else(|| "png".to_string()), - quality: text("quality"), - partial_images: json_u64(body.get("partial_images")).unwrap_or(0), + model, + web_model, + prompt, + size, + ratio, + output_format, + quality, + partial_images, images, }) } } +fn bounded_chatgpt_web_text_field( + body: &Value, + key: &str, + max_bytes: usize, +) -> Result, ExecutionRuntimeTransportError> { + let Some(value) = body.get(key).and_then(Value::as_str) else { + return Ok(None); + }; + let value = value.trim(); + if value.is_empty() { + return Ok(None); + } + if value.len() > max_bytes { + return Err(chatgpt_web_image_request_field_too_large(key)); + } + Ok(Some(value.to_string())) +} + +fn chatgpt_web_image_request_field_too_large(field: &str) -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image request {field} exceeds the supported size" + )) +} + impl WebFingerprint { fn new() -> Self { Self { @@ -594,6 +800,7 @@ async fn web_poll_conversation( conversation_id: &str, uploads: &[WebUploadMeta], ) -> Result { + let conversation_id = validated_web_opaque_id(conversation_id, "conversation ID")?; let path = format!("/backend-api/conversation/{conversation_id}"); let mut headers = web_base_headers(fp, token, path.as_str()); headers.insert("accept".to_string(), "application/json".to_string()); @@ -628,7 +835,17 @@ async fn resolve_and_download_images( add_unique_values(&mut urls, resolved); let mut downloaded = Vec::new(); for url in urls { - match web_download_image(state, plan, base_url, fp, token, url.as_str()).await { + match web_download_image( + state, + plan, + base_url, + fp, + token, + url.as_str(), + WebImageDownloadTrust::ProviderOutput, + ) + .await + { Ok(image) => { downloaded.push(image); break; @@ -639,7 +856,7 @@ async fn resolve_and_download_images( log_type = "debug", request_id = %plan.request_id, candidate_id = ?plan.candidate_id, - error = %err, + error = %safe_transport_error_message(&err), "gateway failed to download one ChatGPT-Web image URL" ); } @@ -658,12 +875,18 @@ async fn web_resolve_image_urls( ) -> Result, ExecutionRuntimeTransportError> { let mut urls = Vec::new(); let uploaded_ids = uploaded_file_ids(uploads); - for file_id in &summary.file_ids { + let conversation_id = summary + .conversation_id + .as_deref() + .map(|value| validated_web_opaque_id(value, "conversation ID")) + .transpose()?; + for raw_file_id in &summary.file_ids { + let file_id = validated_web_file_id(raw_file_id)?; if uploaded_ids.contains(file_id) || file_id == "file_upload" { continue; } let mut path = format!("/backend-api/files/download/{file_id}"); - if let Some(conversation_id) = summary.conversation_id.as_deref() { + if let Some(conversation_id) = conversation_id { path.push_str("?conversation_id="); path.push_str(conversation_id); path.push_str("&inline=false"); @@ -672,8 +895,9 @@ async fn web_resolve_image_urls( add_unique_values(&mut urls, [url]); } } - if let Some(conversation_id) = summary.conversation_id.as_deref() { - for sediment_id in &summary.sediment_ids { + if let Some(conversation_id) = conversation_id { + for raw_sediment_id in &summary.sediment_ids { + let sediment_id = validated_web_opaque_id(raw_sediment_id, "sediment ID")?; if uploaded_ids.contains(sediment_id) { continue; } @@ -715,8 +939,27 @@ async fn web_download_url( .or_else(|| body.get("url")) .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned)) + .and_then(|value| { + if value.is_empty() || value.len() > CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES { + return None; + } + let base = url::Url::parse(base_url).ok()?; + let absolute = url::Url::parse(value).is_ok(); + let url = if absolute { + url::Url::parse(value).ok()? + } else { + base.join(value).ok()? + }; + validate_web_image_http_url(&url).ok()?; + // Relative download paths are expected from the authenticated + // ChatGPT API. Do not let a provider response turn one into an + // arbitrary cross-origin target through URL joining. + if !absolute && !web_download_url_is_same_origin(&base, &url) { + return None; + } + let serialized = url.to_string(); + (serialized.len() <= CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES).then_some(serialized) + })) } async fn web_download_image( @@ -726,47 +969,22 @@ async fn web_download_image( fp: &WebFingerprint, token: &str, raw_url: &str, + trust: WebImageDownloadTrust, ) -> Result { if let Some(data) = parse_data_url(raw_url) { return Ok(data); } - let download_url = if raw_url.starts_with('/') { - format!("{base_url}{raw_url}") - } else { - raw_url.to_string() + let payload = match trust { + WebImageDownloadTrust::UntrustedInput => { + let url = parse_absolute_web_image_url(raw_url)?; + download_public_web_image(url, plan.timeouts.as_ref(), false).await? + } + WebImageDownloadTrust::ProviderOutput => { + download_provider_web_image(plan, base_url, fp, token, raw_url).await? + } }; - let mut headers = BTreeMap::from([( - EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), - "true".to_string(), - )]); - if should_use_web_download_headers(base_url, download_url.as_str()) { - let path = url::Url::parse(download_url.as_str()) - .ok() - .map(|url| url.path().to_string()) - .filter(|path| !path.is_empty()) - .unwrap_or_else(|| "/".to_string()); - headers.extend(web_base_headers(fp, token, path.as_str())); - headers.insert( - "accept".to_string(), - "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8".to_string(), - ); - } - let result = execute_subrequest(plan, "GET", download_url, headers, None, false).await?; - ensure_success(&result, "ChatGPT-Web image download")?; - let data = execution_result_bytes(&result)?; - if data.is_empty() { - return Err(ExecutionRuntimeTransportError::UpstreamRequest( - "ChatGPT-Web image download returned empty body".to_string(), - )); - } - let mime = result - .headers - .get("content-type") - .and_then(|value| value.split(';').next()) - .map(str::trim) - .filter(|value| value.starts_with("image/")) - .unwrap_or("image/png") - .to_string(); + let data = payload.data; + let mime = validate_web_image_payload(&data, payload.content_type.as_deref())?.to_string(); let (width, height) = image_dimensions(&data); Ok(DownloadedImage { b64_json: base64::engine::general_purpose::STANDARD.encode(data), @@ -776,6 +994,554 @@ async fn web_download_image( }) } +fn parse_absolute_web_image_url(raw_url: &str) -> Result { + let url = url::Url::parse(raw_url.trim()).map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image URL is invalid: {err}" + )) + })?; + validate_web_image_http_url(&url)?; + Ok(url) +} + +fn validate_web_image_http_url(url: &url::Url) -> Result<(), ExecutionRuntimeTransportError> { + if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL must be an absolute http or https URL".to_string(), + )); + } + if !url.username().is_empty() || url.password().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL must not contain credentials".to_string(), + )); + } + Ok(()) +} + +/// Return the canonical MIME type for a supported, non-active image payload. +/// +/// Content-Type is metadata supplied by an untrusted upstream and must not be +/// used as the sole type check: an HTML/SVG response can be labelled as +/// `image/png`. Require a real PNG/JPEG/WebP signature and, when a concrete +/// content type is supplied, require it to agree with the signature. +fn validate_web_image_payload( + data: &[u8], + content_type: Option<&str>, +) -> Result<&'static str, ExecutionRuntimeTransportError> { + if data.is_empty() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image download returned empty body".to_string(), + )); + } + let detected = detected_web_image_mime(data).ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image download returned an unsupported image payload".to_string(), + ) + })?; + let declared = declared_web_image_mime(content_type)?; + if let Some(declared) = declared { + if declared != detected { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image content type does not match its payload".to_string(), + )); + } + } + Ok(detected) +} + +/// Parse a response Content-Type into one of the formats we can safely pass +/// through the OpenAI image response surface. Generic octet-stream is +/// allowed only because the payload signature is checked independently. +fn declared_web_image_mime( + content_type: Option<&str>, +) -> Result, ExecutionRuntimeTransportError> { + let Some(content_type) = content_type else { + return Ok(None); + }; + let token = content_type + .split(';') + .next() + .map(str::trim) + .unwrap_or_default(); + if token.is_empty() + || token + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte.is_ascii_control() || byte == b',') + { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image response has an invalid content type".to_string(), + )); + } + if token.eq_ignore_ascii_case("application/octet-stream") + || token.eq_ignore_ascii_case("binary/octet-stream") + { + return Ok(None); + } + let mime = + if token.eq_ignore_ascii_case("image/png") || token.eq_ignore_ascii_case("image/x-png") { + Some("image/png") + } else if token.eq_ignore_ascii_case("image/jpeg") + || token.eq_ignore_ascii_case("image/jpg") + || token.eq_ignore_ascii_case("image/pjpeg") + { + Some("image/jpeg") + } else if token.eq_ignore_ascii_case("image/webp") { + Some("image/webp") + } else { + // This intentionally rejects image/svg+xml, image/avif, generic + // image/*, text/html, and all other active/unsupported types. + None + }; + mime.ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image response has an unsupported content type".to_string(), + ) + }) + .map(Some) +} + +fn detected_web_image_mime(data: &[u8]) -> Option<&'static str> { + // Require the PNG signature and an IHDR chunk before treating the body as + // PNG. This also gives image_dimensions a safe minimum length. + if data.len() >= 24 && data.starts_with(b"\x89PNG\r\n\x1a\n") && &data[12..16] == b"IHDR" { + return Some("image/png"); + } + // JPEG's SOI marker must be followed by a marker prefix. This rejects a + // bare/truncated `ff d8` body while leaving full structural validation to + // the image decoder downstream. + if data.len() >= 3 && data.starts_with(&[0xff, 0xd8, 0xff]) { + return Some("image/jpeg"); + } + // WebP is a RIFF container with a WEBP form type. + if data.len() >= 12 && &data[..4] == b"RIFF" && &data[8..12] == b"WEBP" { + return Some("image/webp"); + } + None +} + +async fn download_provider_web_image( + plan: &ExecutionPlan, + base_url: &str, + fp: &WebFingerprint, + token: &str, + raw_url: &str, +) -> Result { + let base = parse_absolute_web_image_url(base_url)?; + let mut current = base.join(raw_url.trim()).map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image URL is invalid: {err}" + )) + })?; + validate_web_image_http_url(¤t)?; + let mut redirects = 0usize; + + loop { + if !web_download_url_is_same_origin(&base, ¤t) { + // Provider-generated storage URLs may be resolved to RFC 2544 + // synthetic addresses by a local DNS interception tool. The + // public downloader still decides whether the exact storage + // origin is eligible; this flag is never enabled for untrusted + // request input. + return download_public_web_image(current, plan.timeouts.as_ref(), true).await; + } + let path = match current.query() { + Some(query) => format!("{}?{query}", current.path()), + None => current.path().to_string(), + }; + let mut headers = BTreeMap::new(); + if is_authenticated_web_download_url(&base, ¤t) { + headers.extend(web_base_headers(fp, token, path.as_str())); + } + headers.insert( + "accept".to_string(), + "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8".to_string(), + ); + let result = execute_subrequest( + plan, + "GET", + current.as_str().to_string(), + headers, + None, + false, + ) + .await?; + + if (300..400).contains(&result.status_code) { + if redirects >= CHATGPT_WEB_IMAGE_DOWNLOAD_MAX_REDIRECTS { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image download exceeded redirect limit".to_string(), + )); + } + let location = result.headers.get("location").ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image redirect is missing Location header".to_string(), + ) + })?; + current = current.join(location).map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image redirect URL is invalid: {err}" + )) + })?; + validate_web_image_http_url(¤t)?; + redirects += 1; + continue; + } + + ensure_success(&result, "ChatGPT-Web image download")?; + return Ok(WebImageHttpPayload { + data: execution_result_bytes_with_limit( + &result, + chatgpt_web_image_raw_payload_limit_bytes(), + )?, + content_type: result.headers.get("content-type").cloned(), + }); + } +} + +async fn download_public_web_image( + mut current: url::Url, + timeouts: Option<&ExecutionTimeouts>, + allow_benchmarking_fake_ip: bool, +) -> Result { + let mut redirects = 0usize; + let total_timeout = bounded_chatgpt_web_image_timeout( + timeouts.and_then(|value| value.total_ms), + CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS, + ); + let deadline = Instant::now() + total_timeout; + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web public image download timed out".to_string(), + )); + } + validate_web_image_http_url(¤t)?; + let connect_timeout = bounded_chatgpt_web_image_timeout( + timeouts.and_then(|value| value.connect_ms), + CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS, + ); + let (host, resolved) = resolve_public_web_image_addrs( + ¤t, + connect_timeout.min(remaining), + allow_benchmarking_fake_ip, + ) + .await?; + let mut builder = reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()); + builder = builder + .connect_timeout( + bounded_chatgpt_web_image_timeout( + timeouts.and_then(|value| value.connect_ms), + CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS, + ) + .min(remaining), + ) + .read_timeout( + bounded_chatgpt_web_image_timeout( + timeouts.and_then(|value| value.read_ms), + CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS, + ) + .min(remaining), + ) + .timeout(remaining); + if host.parse::().is_err() { + builder = builder.resolve_to_addrs(host.as_str(), &resolved); + } + let client = builder + .build() + .map_err(ExecutionRuntimeTransportError::ClientBuild)?; + let response = client + .get(current.clone()) + .header( + reqwest::header::ACCEPT, + "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8", + ) + .send() + .await + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web public image download failed: {}", + format_upstream_request_error(&err) + )) + })?; + + if response.status().is_redirection() { + if redirects >= CHATGPT_WEB_IMAGE_DOWNLOAD_MAX_REDIRECTS { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image download exceeded redirect limit".to_string(), + )); + } + let location = response + .headers() + .get(reqwest::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image redirect is missing Location header".to_string(), + ) + })?; + current = current.join(location).map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image redirect URL is invalid: {err}" + )) + })?; + redirects += 1; + continue; + } + if !response.status().is_success() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image download returned {}", + response.status().as_u16() + ))); + } + let content_type = response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + let data = aether_http::read_response_bytes_with_limit( + response, + chatgpt_web_image_raw_payload_limit_bytes(), + ) + .await + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web public image body read failed: {err}" + )) + })?; + return Ok(WebImageHttpPayload { data, content_type }); + } +} + +fn bounded_chatgpt_web_image_timeout(configured_ms: Option, default_ms: u64) -> Duration { + Duration::from_millis(configured_ms.unwrap_or(default_ms).clamp(1, 1_200_000)) +} + +async fn resolve_public_web_image_addrs( + url: &url::Url, + lookup_timeout: Duration, + allow_benchmarking_fake_ip: bool, +) -> Result<(String, Vec), ExecutionRuntimeTransportError> { + let host = url.host_str().ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL is missing a host".to_string(), + ) + })?; + let port = url.port_or_known_default().ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL is missing a port".to_string(), + ) + })?; + let resolved = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + tokio::time::timeout(lookup_timeout, tokio::net::lookup_host((host, port))) + .await + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL DNS resolution timed out".to_string(), + ) + })? + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web image URL DNS resolution failed: {err}" + )) + })? + .take(CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES) + .collect::>() + }; + if resolved.is_empty() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL DNS resolution returned no addresses".to_string(), + )); + } + validate_public_web_image_addresses(url, &resolved, allow_benchmarking_fake_ip)?; + Ok((host.to_string(), resolved)) +} + +fn validate_public_web_image_addresses( + url: &url::Url, + addresses: &[SocketAddr], + allow_benchmarking_fake_ip: bool, +) -> Result<(), ExecutionRuntimeTransportError> { + let allows_benchmarking_fake_ip = + allow_benchmarking_fake_ip && web_image_storage_origin_allows_benchmarking_fake_ip(url); + if addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(allows_benchmarking_fake_ip + && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web image URL resolves to a private or reserved address".to_string(), + )); + } + Ok(()) +} + +/// Synthetic DNS is accepted only for the storage origins that ChatGPT uses +/// for generated assets and upload blobs. In particular, a user-supplied URL +/// on an arbitrary host cannot opt into this exception merely by resolving to +/// the RFC 2544 benchmark range. +fn web_image_storage_origin_allows_benchmarking_fake_ip(url: &url::Url) -> bool { + url.scheme().eq_ignore_ascii_case("https") + && url.port_or_known_default() == Some(443) + && url.username().is_empty() + && url.password().is_none() + && url + .host_str() + .is_some_and(chatgpt_web_upload_host_is_allowed) +} + +/// Validate the destination returned by ChatGPT's upload-metadata endpoint. +/// +/// The upload URL is provider-controlled data, not a trusted request target. +/// Keep this boundary narrower than the generic execution URL policy: uploads +/// must go to the storage origins used by ChatGPT, over HTTPS, without +/// credentials or fragments. Azure SAS query parameters are intentionally +/// retained because they carry the upload authorization. +fn validate_chatgpt_web_upload_url( + raw_url: &str, +) -> Result { + let raw_url = raw_url.trim(); + if raw_url.is_empty() || raw_url.len() > CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL is invalid or too large".to_string(), + )); + } + let url = url::Url::parse(raw_url).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL is invalid".to_string(), + ) + })?; + if url.scheme() != "https" + || url.host_str().is_none() + || url.port().is_some_and(|port| port != 443) + { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL must use HTTPS on the default port".to_string(), + )); + } + if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL must not contain credentials or a fragment".to_string(), + )); + } + let host = url.host_str().unwrap_or_default(); + if host.parse::().is_ok() || !chatgpt_web_upload_host_is_allowed(host) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL host is not an allowed storage origin".to_string(), + )); + } + if url.path().len() > CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES + || url + .query() + .is_some_and(|query| query.len() > CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES) + { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload URL is too large".to_string(), + )); + } + Ok(url) +} + +fn chatgpt_web_upload_host_is_allowed(host: &str) -> bool { + if !web_dns_host_is_valid(host) { + return false; + } + if web_host_is_domain_or_subdomain(host, "files.oaiusercontent.com") { + return true; + } + if web_host_is_domain_or_subdomain(host, "oaidalleapiprodscus.blob.core.windows.net") { + return true; + } + web_host_is_strict_subdomain(host, "blob.core.windows.net") + && !web_host_is_domain_or_subdomain(host, "openaiassets.blob.core.windows.net") +} + +/// PUT image bytes to a validated ChatGPT storage URL using a DNS-pinned, +/// proxy-free client. The generic execution runtime intentionally supports +/// configured proxies and broad public HTTPS targets; that is inappropriate +/// for a provider-supplied upload destination. +async fn upload_chatgpt_web_blob( + plan: &ExecutionPlan, + upload_url: &url::Url, + content_type: &str, + user_agent: &str, + base_url: &str, + body: Vec, +) -> Result<(), ExecutionRuntimeTransportError> { + let total_timeout = bounded_chatgpt_web_image_timeout( + plan.timeouts + .as_ref() + .and_then(|timeouts| timeouts.total_ms), + CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS, + ); + let connect_timeout = bounded_chatgpt_web_image_timeout( + plan.timeouts + .as_ref() + .and_then(|timeouts| timeouts.connect_ms), + CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS, + ) + .min(total_timeout); + let read_timeout = bounded_chatgpt_web_image_timeout( + plan.timeouts.as_ref().and_then(|timeouts| timeouts.read_ms), + CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS, + ) + .min(total_timeout); + let (host, resolved) = + resolve_public_web_image_addrs(upload_url, connect_timeout, true).await?; + let client = reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(connect_timeout) + .read_timeout(read_timeout) + .timeout(total_timeout) + .resolve_to_addrs(host.as_str(), &resolved) + .build() + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload client initialization failed".to_string(), + ) + })?; + + let response = client + .put(upload_url.clone()) + .header(reqwest::header::CONTENT_TYPE, content_type) + .header("x-ms-blob-type", "BlockBlob") + .header("x-ms-version", "2020-04-08") + .header(reqwest::header::ORIGIN, base_url) + .header(reqwest::header::REFERER, format!("{base_url}/")) + .header(reqwest::header::USER_AGENT, user_agent) + .body(body) + .send() + .await + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload request failed".to_string(), + ) + })?; + let status_code = response.status().as_u16(); + if !(200..300).contains(&status_code) { + return Err(ExecutionRuntimeTransportError::UpstreamHttpStatus { + status_code, + message: chatgpt_web_stage_http_error_message("ChatGPT-Web upload blob", status_code), + }); + } + aether_http::read_response_bytes_with_limit( + response, + CHATGPT_WEB_IMAGE_UPLOAD_RESPONSE_LIMIT_BYTES, + ) + .await + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web upload response body is invalid or too large".to_string(), + ) + })?; + Ok(()) +} + async fn web_upload_image( state: &AppState, plan: &ExecutionPlan, @@ -785,10 +1551,21 @@ async fn web_upload_image( ref_url: &str, file_name: String, ) -> Result { - let image = web_download_image(state, plan, base_url, fp, token, ref_url).await?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(image.b64_json.as_bytes()) - .map_err(ExecutionRuntimeTransportError::BodyDecode)?; + let image = web_download_image( + state, + plan, + base_url, + fp, + token, + ref_url, + WebImageDownloadTrust::UntrustedInput, + ) + .await?; + let bytes = decode_base64_body_with_limit( + image.b64_json.as_str(), + chatgpt_web_image_raw_payload_limit_bytes(), + )?; + let file_size = bytes.len(); let path = "/backend-api/files"; let mut headers = web_base_headers(fp, token, path); headers.insert("content-type".to_string(), "application/json".to_string()); @@ -811,7 +1588,7 @@ async fn web_upload_image( .await?; ensure_success(&result, "ChatGPT-Web upload metadata")?; let upload_payload = execution_result_json(&result)?; - let file_id = upload_payload + let raw_file_id = upload_payload .get("file_id") .and_then(Value::as_str) .map(str::trim) @@ -820,8 +1597,8 @@ async fn web_upload_image( ExecutionRuntimeTransportError::UpstreamRequest( "ChatGPT-Web upload response missing file_id".to_string(), ) - })? - .to_string(); + })?; + let file_id = validated_web_file_id(raw_file_id)?.to_string(); let upload_url = upload_payload .get("upload_url") .and_then(Value::as_str) @@ -832,29 +1609,16 @@ async fn web_upload_image( "ChatGPT-Web upload response missing upload_url".to_string(), ) })?; - - let put_headers = BTreeMap::from([ - ("content-type".to_string(), image.mime.clone()), - ("x-ms-blob-type".to_string(), "BlockBlob".to_string()), - ("x-ms-version".to_string(), "2020-04-08".to_string()), - ("origin".to_string(), base_url.to_string()), - ("referer".to_string(), format!("{base_url}/")), - ("user-agent".to_string(), fp.user_agent.to_string()), - ]); - let put_result = execute_subrequest( + let upload_url = validate_chatgpt_web_upload_url(upload_url)?; + upload_chatgpt_web_blob( plan, - "PUT", - upload_url.to_string(), - put_headers, - Some(RequestBody { - json_body: None, - body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(&bytes)), - body_ref: None, - }), - false, + &upload_url, + image.mime.as_str(), + fp.user_agent, + base_url, + bytes, ) .await?; - ensure_success(&put_result, "ChatGPT-Web upload blob")?; let uploaded_path = format!("/backend-api/files/{file_id}/uploaded"); let mut uploaded_headers = web_base_headers(fp, token, uploaded_path.as_str()); @@ -883,7 +1647,7 @@ async fn web_upload_image( file_id, library_file_id, file_name, - file_size: bytes.len(), + file_size, mime: image.mime, width: image.width, height: image.height, @@ -931,7 +1695,7 @@ async fn web_process_upload_stream( .and_then(|extra| extra.get("metadata_object_id")) .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| web_opaque_id_is_safe(value)) .map(ToOwned::to_owned) }) })) @@ -941,14 +1705,10 @@ async fn execute_subrequest( plan: &ExecutionPlan, method: &str, url: String, - mut headers: BTreeMap, + headers: BTreeMap, body: Option, stream: bool, ) -> Result { - headers.insert( - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER.to_string(), - "true".to_string(), - ); let subplan = ExecutionPlan { request_id: plan.request_id.clone(), candidate_id: plan.candidate_id.clone(), @@ -1079,7 +1839,7 @@ async fn apply_chatgpt_web_image_quota_request_delta( let Some(mut latest_key) = state .read_provider_catalog_keys_by_ids(&[key_id.to_string()]) .await - .map_err(|err| err.into_message())? + .map_err(|_| "ChatGPT-Web quota state read failed".to_string())? .into_iter() .find(|key| key.id == key_id && key.provider_id == provider_id) else { @@ -1106,7 +1866,8 @@ async fn apply_chatgpt_web_image_quota_request_delta( return Ok(false); } - let namespace_value = Value::Object(metadata); + let namespace_value = + admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(&Value::Object(metadata))); let updated_upstream_metadata = merge_provider_metadata_object( latest_key.upstream_metadata.as_ref(), "chatgpt_web", @@ -1139,7 +1900,7 @@ async fn apply_chatgpt_web_image_quota_request_delta( }, ) .await - .map_err(|err| err.into_message())?; + .map_err(|_| "ChatGPT-Web quota state update failed".to_string())?; if persisted { return Ok(true); } @@ -1422,7 +2183,7 @@ async fn refresh_chatgpt_web_image_quota_after_success( let key_available = state .read_provider_catalog_keys_by_ids(&key_ids) .await - .map_err(|err| err.into_message())? + .map_err(|_| "ChatGPT-Web quota key read failed".to_string())? .into_iter() .any(|key| key.id == key_id && key.provider_id == provider_id); if !key_available { @@ -1431,7 +2192,7 @@ async fn refresh_chatgpt_web_image_quota_after_success( let Some(provider) = state .read_provider_catalog_providers_by_ids(&provider_ids) .await - .map_err(|err| err.into_message())? + .map_err(|_| "ChatGPT-Web quota provider read failed".to_string())? .into_iter() .find(|provider| provider.id == provider_id) else { @@ -1454,19 +2215,16 @@ async fn refresh_chatgpt_web_image_quota_after_success( let result = DirectSyncExecutionRuntime::new() .execute_sync("a_plan) .await - .map_err(|err| err.to_string())?; + .map_err(|_| "ChatGPT-Web quota refresh request failed".to_string())?; if result.status_code != 200 { - let body_excerpt = String::from_utf8_lossy(&execution_result_body_bytes_lossy(&result)) - .chars() - .take(320) - .collect::(); return Err(format!( - "conversation/init returned {}: {}", - result.status_code, body_excerpt + "ChatGPT-Web quota refresh returned HTTP {}", + result.status_code )); } - let body_json = execution_result_json(&result).map_err(|err| err.to_string())?; + let body_json = execution_result_json(&result) + .map_err(|_| "ChatGPT-Web quota refresh response was invalid".to_string())?; let now_unix_secs = current_unix_secs(); let Some(metadata) = parse_chatgpt_web_conversation_init_response(&body_json, now_unix_secs) else { @@ -1475,7 +2233,7 @@ async fn refresh_chatgpt_web_image_quota_after_success( let Some(latest_key) = state .read_provider_catalog_keys_by_ids(&key_ids) .await - .map_err(|err| err.into_message())? + .map_err(|_| "ChatGPT-Web quota key read failed".to_string())? .into_iter() .find(|key| key.id == key_id && key.provider_id == provider_id) else { @@ -1489,6 +2247,7 @@ async fn refresh_chatgpt_web_image_quota_after_success( .cloned(); let mut metadata = metadata.clone(); normalize_chatgpt_web_image_quota_limit(&mut metadata, latest_key.upstream_metadata.as_ref()); + metadata = admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(&metadata)); let mut updated_key = latest_key; let namespace_value = metadata.clone(); @@ -1524,18 +2283,17 @@ async fn refresh_chatgpt_web_image_quota_after_success( updated_at_unix_secs: updated_key.updated_at_unix_secs, }) .await - .map_err(|err| err.into_message())?; + .map_err(|_| "ChatGPT-Web quota state update failed".to_string())?; if persisted { return state .update_provider_catalog_key_oauth_runtime_state( &updated_key.id, updated_key.oauth_invalid_at_unix_secs, updated_key.oauth_invalid_reason.as_deref(), - None, updated_key.updated_at_unix_secs, ) .await - .map_err(|err| err.into_message()); + .map_err(|_| "ChatGPT-Web OAuth state update failed".to_string()); } // The conversation/init response is an authoritative snapshot. A // conflict means a newer local delta won; do not overwrite it with @@ -1553,20 +2311,13 @@ fn build_chatgpt_web_image_quota_refresh_plan( quota_kind: _, method, url, - mut headers, + headers, content_type, json_body, client_api_format, provider_api_format, model_name, - accept_invalid_certs, } = spec; - if accept_invalid_certs { - headers.insert( - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER.to_string(), - "true".to_string(), - ); - } let body = json_body .map(RequestBody::from_json) .unwrap_or(RequestBody { @@ -1912,28 +2663,37 @@ fn flush_sse_data(data_lines: &mut Vec, summary: &mut WebImageSseSummary value.get("type").and_then(Value::as_str), Some("error" | "response.failed") ) { - summary.failure = Some(value.clone()); + summary.failure = Some(bounded_web_failure_value(&value)); } if let Some(text) = extract_assistant_text(&value) { summary.last_text = Some(text); } - if let Some(result) = value - .get("item") - .filter(|item| { - item.get("type").and_then(Value::as_str) == Some("image_generation_call") - }) - .and_then(|item| item.get("result")) - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - { - add_unique_values( - &mut summary.direct_urls, - [format!("data:image/png;base64,{result}")], - ); + if let Some(item) = value.get("item").filter(|item| { + item.get("type").and_then(Value::as_str) == Some("image_generation_call") + }) { + // Keep the provider's declared output format when constructing a + // data URL. The bytes are still verified by `parse_data_url` + // before download, but labelling every output as PNG would create + // an avoidable MIME/signature mismatch (and an extra failed + // download attempt) for JPEG/WebP results. + if let Some(result) = item + .get("result") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + { + let mime = mime_for_web_output_format( + item.get("output_format") + .and_then(Value::as_str) + .unwrap_or_default(), + ); + if let Some(url) = bounded_web_image_data_url(mime, result) { + summary.add_values(WebImageSummaryCollection::DirectUrl, [url]); + } + } } - add_unique_values( - &mut summary.direct_urls, + summary.add_values( + WebImageSummaryCollection::DirectUrl, extract_web_image_payload_urls(&value), ); extract_web_image_values(&value, summary); @@ -1973,7 +2733,9 @@ fn extract_web_image_payload_urls(value: &Value) -> Vec { .and_then(Value::as_str) .unwrap_or_default(), ); - add_unique_values(&mut urls, [format!("data:{mime};base64,{partial_b64}")]); + if let Some(url) = bounded_web_image_data_url(mime, partial_b64) { + add_unique_values(&mut urls, [url]); + } } } _ => { @@ -2019,6 +2781,9 @@ fn image_payload_url_from_object(value: &Value) -> Option { .map(str::trim) .filter(|value| !value.is_empty()) { + if url.len() > CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES { + return None; + } return Some(url.to_string()); } let b64 = value @@ -2034,14 +2799,32 @@ fn image_payload_url_from_object(value: &Value) -> Option { .and_then(Value::as_str) .unwrap_or_default(), ); + bounded_web_image_data_url(mime, b64) +} + +fn bounded_web_image_data_url(mime: &str, b64: &str) -> Option { + let b64 = b64.trim(); + if b64.is_empty() + || b64.len() + > maximum_base64_len_for_decoded_limit(chatgpt_web_image_raw_payload_limit_bytes()) + { + return None; + } + let prefix_len = "data:;base64,".len().saturating_add(mime.len()); + if prefix_len.saturating_add(b64.len()) > chatgpt_web_image_sse_envelope_limit_bytes() { + return None; + } Some(format!("data:{mime};base64,{b64}")) } fn mime_for_web_output_format(format: &str) -> &'static str { - match format.trim().to_ascii_lowercase().as_str() { - "jpeg" | "jpg" => "image/jpeg", - "webp" => "image/webp", - _ => "image/png", + let format = format.trim(); + if format.eq_ignore_ascii_case("jpeg") || format.eq_ignore_ascii_case("jpg") { + "image/jpeg" + } else if format.eq_ignore_ascii_case("webp") { + "image/webp" + } else { + "image/png" } } @@ -2053,7 +2836,7 @@ fn extract_web_image_values(value: &Value, summary: &mut WebImageSseSummary) { if let Some(conversation_id) = value .as_str() .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| web_opaque_id_is_safe(value)) { summary .conversation_id @@ -2070,15 +2853,21 @@ fn extract_web_image_values(value: &Value, summary: &mut WebImageSseSummary) { } Value::String(text) => { let text = text.trim(); - if text.starts_with("sediment://") { - add_unique_values( - &mut summary.sediment_ids, - [text.trim_start_matches("sediment://").to_string()], - ); + if let Some(sediment_id) = text.strip_prefix("sediment://") { + if web_opaque_id_is_safe(sediment_id) { + summary.add_values( + WebImageSummaryCollection::SedimentId, + [sediment_id.to_string()], + ); + } } else if is_web_file_id(text) { - add_unique_values(&mut summary.file_ids, [text.to_string()]); - } else if is_generated_web_asset_url(text) || text.starts_with("data:image/") { - add_unique_values(&mut summary.direct_urls, [text.to_string()]); + summary.add_values(WebImageSummaryCollection::FileId, [text.to_string()]); + } else if (text.len() <= CHATGPT_WEB_IMAGE_MAX_EXTERNAL_URL_BYTES + && is_generated_web_asset_url(text)) + || (text.len() <= chatgpt_web_image_sse_envelope_limit_bytes() + && is_data_image_reference(text)) + { + summary.add_values(WebImageSummaryCollection::DirectUrl, [text.to_string()]); } } _ => {} @@ -2093,17 +2882,63 @@ fn extract_assistant_text(value: &Value) -> Option { .and_then(Value::as_array) .and_then(|parts| parts.iter().filter_map(Value::as_str).next()) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| { + !value.is_empty() && value.len() <= CHATGPT_WEB_IMAGE_SUMMARY_MAX_TEXT_BYTES + }) .map(ToOwned::to_owned) } +fn bounded_web_failure_value(value: &Value) -> Value { + if json_value_fits_serialized_limit(value, CHATGPT_WEB_IMAGE_SUMMARY_MAX_TEXT_BYTES) { + return value.clone(); + } + if value.get("type").and_then(Value::as_str) == Some("response.failed") { + json!({ + "type": "response.failed", + "response": { + "status": "failed", + "error": { + "code": "chatgpt_web_image_failed", + "message": "ChatGPT-Web image provider returned an oversized failure" + } + } + }) + } else { + json!({ + "type": "error", + "error": { + "code": "chatgpt_web_image_failed", + "message": "ChatGPT-Web image provider returned an oversized failure" + } + }) + } +} + fn merge_web_summary(target: &mut WebImageSseSummary, source: &mut WebImageSseSummary) { if target.conversation_id.is_none() { - target.conversation_id = source.conversation_id.take(); + target.conversation_id = source + .conversation_id + .take() + .filter(|value| web_opaque_id_is_safe(value)); } - add_unique_values(&mut target.file_ids, source.file_ids.drain(..)); - add_unique_values(&mut target.sediment_ids, source.sediment_ids.drain(..)); - add_unique_values(&mut target.direct_urls, source.direct_urls.drain(..)); + target.add_values( + WebImageSummaryCollection::FileId, + source + .file_ids + .drain(..) + .filter(|value| is_web_file_id(value)), + ); + target.add_values( + WebImageSummaryCollection::SedimentId, + source + .sediment_ids + .drain(..) + .filter(|value| web_opaque_id_is_safe(value)), + ); + target.add_values( + WebImageSummaryCollection::DirectUrl, + source.direct_urls.drain(..), + ); if target.failure.is_none() { target.failure = source.failure.take(); } @@ -2130,10 +2965,19 @@ fn uploaded_file_ids(uploads: &[WebUploadMeta]) -> BTreeSet { } fn add_unique_values(values: &mut Vec, incoming: impl IntoIterator) { + let budget = chatgpt_web_image_sse_envelope_limit_bytes(); + let mut retained_bytes = saturating_string_bytes(values); for value in incoming { - if !value.is_empty() && !values.iter().any(|existing| existing == &value) { - values.push(value); + if value.is_empty() + || value.len() > budget + || values.len() >= CHATGPT_WEB_IMAGE_SUMMARY_MAX_DIRECT_URLS + || value.len() > budget.saturating_sub(retained_bytes) + || values.iter().any(|existing| existing == &value) + { + continue; } + retained_bytes = retained_bytes.saturating_add(value.len()); + values.push(value); } } @@ -2521,12 +3365,14 @@ fn json_u64(value: Option<&Value>) -> Option { } fn chatgpt_web_image_operation(value: Option<&Value>) -> String { - value - .and_then(Value::as_str) - .map(str::trim) - .map(str::to_ascii_lowercase) - .filter(|value| matches!(value.as_str(), "generate" | "edit")) - .unwrap_or_else(|| "generate".to_string()) + let Some(value) = value.and_then(Value::as_str).map(str::trim) else { + return "generate".to_string(); + }; + if value.eq_ignore_ascii_case("edit") { + "edit".to_string() + } else { + "generate".to_string() + } } fn build_failed_sse(request: &ChatGptWebImageRequest, failure: &Value) -> String { @@ -2609,9 +3455,15 @@ fn bytes_execution_result( headers: BTreeMap, body: Vec, started_at: Instant, -) -> ExecutionResult { +) -> Result { + let envelope_limit = chatgpt_web_image_sse_envelope_limit_bytes(); + if body.len() > envelope_limit { + return Err(ExecutionRuntimeTransportError::BodyTooLarge { + limit_bytes: envelope_limit, + }); + } let body_len = body.len() as u64; - ExecutionResult { + Ok(ExecutionResult { request_id: plan.request_id.clone(), candidate_id: plan.candidate_id.clone(), status_code, @@ -2623,15 +3475,19 @@ fn bytes_execution_result( }), telemetry: Some(telemetry(started_at, body_len)), error: None, - } + }) } fn execution_result_frame_stream( plan: &ExecutionPlan, result: &ExecutionResult, report_context: Option<&Value>, -) -> BoxStream<'static, Result> { - let body = execution_result_body_bytes_lossy(result); +) -> Result>, ExecutionRuntimeTransportError> { + // The synthetic ChatGPT-Web SSE body embeds an image as base64, so its + // envelope is larger than the decoded image/body limit. Use the bounded + // envelope budget here instead of rejecting valid images near 64 MiB. + let body = + execution_result_bytes_with_limit(result, chatgpt_web_image_sse_envelope_limit_bytes())?; let terminal_summary = chatgpt_web_stream_terminal_summary(plan, result, report_context, &body); let mut frames = vec![ StreamFrame { @@ -2653,11 +3509,11 @@ fn execution_result_frame_stream( }, }, ]; - if !body.is_empty() { + for chunk in body.chunks(CHATGPT_WEB_IMAGE_STREAM_CHUNK_BYTES) { frames.push(StreamFrame { frame_type: StreamFrameType::Data, payload: StreamFramePayload::Data { - chunk_b64: Some(base64::engine::general_purpose::STANDARD.encode(body.as_slice())), + chunk_b64: Some(base64::engine::general_purpose::STANDARD.encode(chunk)), text: None, }, }); @@ -2673,12 +3529,12 @@ fn execution_result_frame_stream( }, }); frames.push(StreamFrame::eof_with_summary(terminal_summary)); - stream::iter( + Ok(stream::iter( frames .into_iter() .map(|frame| encode_stream_frame_ndjson(&frame)), ) - .boxed() + .boxed()) } fn chatgpt_web_stream_terminal_summary( @@ -2793,20 +3649,49 @@ fn execution_result_json( fn execution_result_bytes( result: &ExecutionResult, ) -> Result, ExecutionRuntimeTransportError> { - Ok(execution_result_body_bytes_lossy(result)) + execution_result_bytes_with_limit(result, crate::headers::max_internal_buffered_body_bytes()) } -fn execution_result_body_bytes_lossy(result: &ExecutionResult) -> Vec { +fn execution_result_bytes_with_limit( + result: &ExecutionResult, + body_limit: usize, +) -> Result, ExecutionRuntimeTransportError> { let Some(body) = result.body.as_ref() else { - return Vec::new(); + return Ok(Vec::new()); }; if let Some(json_body) = body.json_body.as_ref() { - return serde_json::to_vec(json_body).unwrap_or_default(); + return serialize_json_body_with_limit(json_body, body_limit); } body.body_bytes_b64 .as_deref() - .and_then(|value| base64::engine::general_purpose::STANDARD.decode(value).ok()) - .unwrap_or_default() + .map(|value| decode_base64_body_with_limit(value, body_limit)) + .unwrap_or_else(|| Ok(Vec::new())) +} + +fn execution_result_body_bytes_lossy(result: &ExecutionResult) -> Vec { + execution_result_bytes(result).unwrap_or_default() +} + +pub(super) fn chatgpt_web_image_sse_envelope_limit_bytes() -> usize { + let image_limit = crate::headers::max_internal_buffered_body_bytes(); + maximum_base64_len_for_decoded_limit(image_limit) + .saturating_add(CHATGPT_WEB_IMAGE_SSE_WRAPPER_OVERHEAD_BYTES) + .min(CHATGPT_WEB_IMAGE_SSE_HARD_MAX_BYTES) +} + +fn chatgpt_web_image_raw_payload_limit_bytes() -> usize { + let configured_limit = crate::headers::max_internal_buffered_body_bytes(); + let envelope_limit = chatgpt_web_image_sse_envelope_limit_bytes(); + let available_for_base64 = + envelope_limit.saturating_sub(CHATGPT_WEB_IMAGE_SSE_WRAPPER_OVERHEAD_BYTES); + // Standard base64 expands three bytes into four. Use the floor of the + // inverse expansion so a raw image can always be represented by the + // synthetic SSE envelope without first allocating an over-sized body. + let representable_raw_limit = available_for_base64 + .saturating_div(4) + .saturating_mul(3) + .max(1); + configured_limit.min(representable_raw_limit) } fn ensure_success( @@ -2816,17 +3701,16 @@ fn ensure_success( if (200..300).contains(&result.status_code) { return Ok(()); } - let body = String::from_utf8_lossy(&execution_result_body_bytes_lossy(result)).to_string(); Err(ExecutionRuntimeTransportError::UpstreamHttpStatus { status_code: result.status_code, - message: format!( - "{stage} returned {}: {}", - result.status_code, - body.chars().take(320).collect::() - ), + message: chatgpt_web_stage_http_error_message(stage, result.status_code), }) } +fn chatgpt_web_stage_http_error_message(stage: &str, status_code: u16) -> String { + format!("{stage} returned HTTP {status_code}") +} + fn chatgpt_web_base_url_from_plan(plan: &ExecutionPlan) -> String { let Ok(url) = url::Url::parse(&plan.url) else { return CHATGPT_WEB_DEFAULT_BASE_URL.to_string(); @@ -2908,7 +3792,10 @@ fn pow_generate(seed: &str, difficulty: &str, config: Vec) -> (String, bo let Some(diff_bytes) = hex_to_bytes(difficulty) else { return (encode_pow_seed(seed), false); }; - if diff_bytes.is_empty() { + // `sha3_512` yields exactly 64 bytes. Difficulty comes from the + // upstream sentinel response, so reject an overlong value before the + // comparison below could slice the digest out of bounds. + if diff_bytes.is_empty() || diff_bytes.len() > 64 { return (encode_pow_seed(seed), false); } @@ -2943,7 +3830,13 @@ fn encode_pow_seed(seed: &str) -> String { } fn hex_to_bytes(value: &str) -> Option> { - let mut hex = value.trim().to_string(); + // Avoid copying/allocating an unbounded upstream difficulty string. The + // proof comparison cannot consume more than the 64-byte SHA-3 digest. + let trimmed = value.trim(); + if trimmed.len() > 128 { + return None; + } + let mut hex = trimmed.to_string(); if hex.len() % 2 == 1 { hex.insert(0, '0'); } @@ -3062,20 +3955,57 @@ fn keccak_f1600(state: &mut [u64; 25]) { } fn parse_data_url(value: &str) -> Option { + parse_data_url_with_limit(value, chatgpt_web_image_raw_payload_limit_bytes()) +} + +fn parse_data_url_with_limit(value: &str, decoded_limit: usize) -> Option { let (header, data) = value.trim().split_once(',')?; - let mime = header - .strip_prefix("data:") - .and_then(|value| value.split(';').next()) - .filter(|value| value.starts_with("image/")) - .unwrap_or("image/png") - .to_string(); - let bytes = base64::engine::general_purpose::STANDARD - .decode(data) - .ok()?; + if header.is_empty() + || data.is_empty() + || header + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte.is_ascii_control()) + || data + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte.is_ascii_control()) + { + return None; + } + + // Only pass image formats that the OpenAI image surface can represent safely. + // In particular, accepting arbitrary `image/*` values would allow SVG/XML + // payloads to cross a JSON image boundary and be interpreted as active markup. + let (scheme, metadata) = header.split_once(':')?; + if !scheme.eq_ignore_ascii_case("data") { + return None; + } + let (mime, encoding) = metadata.rsplit_once(';')?; + if !encoding.eq_ignore_ascii_case("base64") || mime.contains(';') { + return None; + } + let mime = if mime.eq_ignore_ascii_case("image/png") { + "image/png" + } else if mime.eq_ignore_ascii_case("image/jpeg") || mime.eq_ignore_ascii_case("image/jpg") { + "image/jpeg" + } else if mime.eq_ignore_ascii_case("image/webp") { + "image/webp" + } else { + return None; + }; + + // Check the encoded length before invoking the decoder. The base64 engine + // allocates from the input length, so a decoded-size check performed after + // decoding would still leave an allocation DoS. Keep the operator-configured + // 64 MiB default so valid large image responses remain supported. + let bytes = decode_base64_body_with_limit(data, decoded_limit).ok()?; + let detected_mime = validate_web_image_payload(&bytes, Some(mime)).ok()?; let (width, height) = image_dimensions(&bytes); Some(DownloadedImage { - b64_json: base64::engine::general_purpose::STANDARD.encode(bytes), - mime, + // `decode_base64_body_with_limit` has already validated the canonical + // alphabet and padding. Preserve the source text to avoid a second + // 64 MiB-scale allocation when handling large images. + b64_json: data.to_string(), + mime: detected_mime.to_string(), width, height, }) @@ -3085,6 +4015,13 @@ fn image_dimensions(bytes: &[u8]) -> (Option, Option) { if bytes.starts_with(b"\x89PNG\r\n\x1a\n") && bytes.len() >= 24 { let width = u32::from_be_bytes([bytes[16], bytes[17], bytes[18], bytes[19]]); let height = u32::from_be_bytes([bytes[20], bytes[21], bytes[22], bytes[23]]); + if width == 0 + || height == 0 + || width > CHATGPT_WEB_IMAGE_MAX_DIMENSION + || height > CHATGPT_WEB_IMAGE_MAX_DIMENSION + { + return (None, None); + } return (Some(width), Some(height)); } if bytes.starts_with(&[0xff, 0xd8]) { @@ -3114,6 +4051,13 @@ fn image_dimensions(bytes: &[u8]) -> (Option, Option) { { let height = u16::from_be_bytes([bytes[cursor + 5], bytes[cursor + 6]]) as u32; let width = u16::from_be_bytes([bytes[cursor + 7], bytes[cursor + 8]]) as u32; + if width == 0 + || height == 0 + || width > CHATGPT_WEB_IMAGE_MAX_DIMENSION + || height > CHATGPT_WEB_IMAGE_MAX_DIMENSION + { + return (None, None); + } return (Some(width), Some(height)); } if segment_len < 2 { @@ -3127,39 +4071,118 @@ fn image_dimensions(bytes: &[u8]) -> (Option, Option) { fn is_web_file_id(value: &str) -> bool { let value = value.trim(); - (value.starts_with("file-") || value.starts_with("file_")) && value.len() >= 10 + (value.starts_with("file-") || value.starts_with("file_")) + && value.len() >= 10 + && web_opaque_id_is_safe(value) +} + +fn web_opaque_id_is_safe(value: &str) -> bool { + !value.is_empty() + && value.len() <= CHATGPT_WEB_OPAQUE_ID_MAX_BYTES + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) +} + +fn validated_web_opaque_id<'a>( + value: &'a str, + field: &str, +) -> Result<&'a str, ExecutionRuntimeTransportError> { + let value = value.trim(); + if !web_opaque_id_is_safe(value) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( + "ChatGPT-Web response contains an invalid {field}" + ))); + } + Ok(value) +} + +fn validated_web_file_id(value: &str) -> Result<&str, ExecutionRuntimeTransportError> { + let value = value.trim(); + if !is_web_file_id(value) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "ChatGPT-Web response contains an invalid file ID".to_string(), + )); + } + Ok(value) +} + +fn web_dns_host_is_valid(host: &str) -> bool { + let host = host.strip_suffix('.').unwrap_or(host); + !host.is_empty() + && !host.ends_with('.') + && host.len() <= 253 + && host.split('.').all(|label| { + !label.is_empty() + && label.len() <= 63 + && label + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') + && label + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && label + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric) + }) +} + +fn web_host_is_domain_or_subdomain(host: &str, domain: &str) -> bool { + let host = host.strip_suffix('.').unwrap_or(host); + if !web_dns_host_is_valid(host) { + return false; + } + if host.eq_ignore_ascii_case(domain) { + return true; + } + host.len() > domain.len() + && host.as_bytes()[host.len() - domain.len() - 1] == b'.' + && host[host.len() - domain.len()..].eq_ignore_ascii_case(domain) +} + +fn web_host_is_strict_subdomain(host: &str, domain: &str) -> bool { + web_host_is_domain_or_subdomain(host, domain) + && !host.trim_end_matches('.').eq_ignore_ascii_case(domain) } fn is_generated_web_asset_url(raw_url: &str) -> bool { let Ok(url) = url::Url::parse(raw_url.trim()) else { return false; }; - let Some(host) = url.host_str().map(str::to_ascii_lowercase) else { + if validate_web_image_http_url(&url).is_err() { + return false; + } + let Some(host) = url.host_str() else { return false; }; let path = url.path().to_ascii_lowercase(); - if host.contains("openaiassets.blob.core.windows.net") { + if web_host_is_domain_or_subdomain(host, "openaiassets.blob.core.windows.net") { return false; } if path.contains("/$web/chatgpt/") { return false; } - host.contains("files.oaiusercontent.com") - || host.contains("oaidalleapiprodscus.blob.core.windows.net") - || (host.ends_with(".blob.core.windows.net") && !path.contains("/$web/")) + web_host_is_domain_or_subdomain(host, "files.oaiusercontent.com") + || web_host_is_domain_or_subdomain(host, "oaidalleapiprodscus.blob.core.windows.net") + || (web_host_is_strict_subdomain(host, "blob.core.windows.net") && !path.contains("/$web/")) } -fn should_use_web_download_headers(base_url: &str, raw_url: &str) -> bool { - let Ok(url) = url::Url::parse(raw_url) else { - return raw_url.starts_with("/backend-api/"); - }; - if url.path().starts_with("/backend-api/") { - return true; - } - let Ok(base) = url::Url::parse(base_url) else { - return false; - }; - url.domain() == base.domain() +fn is_authenticated_web_download_url(base: &url::Url, target: &url::Url) -> bool { + target.path().starts_with("/backend-api/") + && web_download_url_is_same_origin(base, target) + && target.username().is_empty() + && target.password().is_none() +} + +fn web_download_url_is_same_origin(base: &url::Url, target: &url::Url) -> bool { + target.scheme().eq_ignore_ascii_case(base.scheme()) + && target + .host_str() + .zip(base.host_str()) + .is_some_and(|(target, base)| target.eq_ignore_ascii_case(base)) + && target.port_or_known_default() == base.port_or_known_default() } #[cfg(test)] @@ -3256,7 +4279,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:image"])), - Some("test-access-token".to_string()), + None, None, None, None, @@ -3326,6 +4349,37 @@ mod tests { assert_eq!(gpt_image2_output_tokens(1024, 1024, "high"), 7024); } + #[test] + fn chatgpt_web_http_status_error_omits_upstream_body() { + let result = ExecutionResult { + request_id: "req-chatgpt-web-image-test".to_string(), + candidate_id: None, + status_code: 502, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({"error": "Bearer secret-chatgpt-web-body"})), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }; + + let error = ensure_success(&result, "ChatGPT-Web bootstrap") + .expect_err("non-success result should be rejected"); + let message = error.to_string(); + + assert!(matches!( + error, + ExecutionRuntimeTransportError::UpstreamHttpStatus { + status_code: 502, + .. + } + )); + assert_eq!(message, "ChatGPT-Web bootstrap returned HTTP 502"); + assert!(!message.contains("secret-chatgpt-web-body")); + } + #[test] fn chatgpt_web_success_sse_includes_estimated_image_usage() { let request = ChatGptWebImageRequest { @@ -3499,13 +4553,6 @@ mod tests { quota_plan.headers.get("authorization").map(String::as_str), Some("Bearer test-access-token") ); - assert_eq!( - quota_plan - .headers - .get(EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER) - .map(String::as_str), - Some("true") - ); assert_eq!( quota_plan .transport_profile @@ -3824,6 +4871,22 @@ data: [DONE] ); } + #[test] + fn parse_web_image_sse_preserves_inline_output_format() { + let jpeg_payload = + base64::engine::general_purpose::STANDARD.encode([0xff, 0xd8, 0xff, 0xd9]); + let event = format!( + "data: {{\"type\":\"response.output_item.done\",\"item\":{{\"type\":\"image_generation_call\",\"result\":\"{jpeg_payload}\",\"output_format\":\"jpeg\"}}}}\n\n" + ); + + let summary = parse_web_image_sse(event.as_bytes()); + + assert_eq!( + summary.direct_urls, + vec![format!("data:image/jpeg;base64,{jpeg_payload}")] + ); + } + #[test] fn parse_web_image_sse_preserves_response_failed_event() { let summary = parse_web_image_sse( @@ -3846,12 +4909,414 @@ data: [DONE] #[test] fn generated_asset_filter_does_not_drop_icon_or_logo_outputs() { - assert!(is_generated_web_asset_url( - "https://files.oaiusercontent.com/generated/icon-logo-output.png" - )); + for accepted in [ + "https://files.oaiusercontent.com/generated/icon-logo-output.png", + "https://cdn.files.oaiusercontent.com/generated/image.png", + "https://oaidalleapiprodscus.blob.core.windows.net/generated/image.png", + "https://tenant.blob.core.windows.net/generated/image.png", + "https://files.oaiusercontent.com./generated/image.png", + ] { + assert!( + is_generated_web_asset_url(accepted), + "generated image host should be accepted: {accepted}" + ); + } assert!(!is_generated_web_asset_url( "https://openaiassets.blob.core.windows.net/$web/chatgpt/filled-plus-icon.svg" )); + + for rejected in [ + "https://notfiles.oaiusercontent.com/generated/image.png", + "https://files.oaiusercontent.com.attacker.invalid/generated/image.png", + "https://not-oaidalleapiprodscus.blob.core.windows.net.attacker.invalid/image.png", + "https://blob.core.windows.net/generated/image.png", + "https://openaiassets.blob.core.windows.net/generated/image.png", + "https://sub.openaiassets.blob.core.windows.net/generated/image.png", + "ftp://files.oaiusercontent.com/generated/image.png", + "https://user@files.oaiusercontent.com/generated/image.png", + ] { + assert!( + !is_generated_web_asset_url(rejected), + "lookalike or static asset host should be rejected: {rejected}" + ); + } + } + + #[test] + fn web_opaque_ids_reject_path_and_query_injection() { + for accepted in ["conv-test_123", "file-generated-123456", "sediment_123"] { + assert!(web_opaque_id_is_safe(accepted)); + } + assert!(is_web_file_id("file-generated-123456")); + + for rejected in [ + "../admin", + "conv/other", + "conv?inline=true", + "conv#fragment", + "conv%2fadmin", + "conv&inline=true", + "conv=value", + "conv value", + "\r\nX-Injected: true", + ] { + assert!( + !web_opaque_id_is_safe(rejected), + "unsafe opaque ID should be rejected: {rejected:?}" + ); + } + assert!(!web_opaque_id_is_safe( + "a".repeat(CHATGPT_WEB_OPAQUE_ID_MAX_BYTES + 1).as_str() + )); + assert!(!is_web_file_id("file-generated-123456/../../admin")); + assert!(validated_web_file_id("file-generated-123456?download=1").is_err()); + } + + #[test] + fn web_image_value_extraction_keeps_only_safe_opaque_ids() { + let mut summary = WebImageSseSummary::default(); + extract_web_image_values( + &json!({ + "conversation_id": "conv-test_123", + "file": "file-generated-123456", + "sediment": "sediment://sediment_123" + }), + &mut summary, + ); + assert_eq!(summary.conversation_id.as_deref(), Some("conv-test_123")); + assert_eq!(summary.file_ids, vec!["file-generated-123456"]); + assert_eq!(summary.sediment_ids, vec!["sediment_123"]); + + let mut malicious = WebImageSseSummary::default(); + extract_web_image_values( + &json!({ + "conversation_id": "conv-test?inline=true", + "file": "file-generated-123456/../../admin", + "sediment": "sediment://sediment_123?download=1" + }), + &mut malicious, + ); + assert!(malicious.conversation_id.is_none()); + assert!(malicious.file_ids.is_empty()); + assert!(malicious.sediment_ids.is_empty()); + } + + #[test] + fn chatgpt_web_image_url_validation_requires_absolute_http_without_credentials() { + assert!(parse_absolute_web_image_url("https://cdn.example/image.png").is_ok()); + + for rejected in [ + "/relative.png", + "file:///etc/passwd", + "ftp://cdn.example/image.png", + "https://user:password@cdn.example/image.png", + ] { + assert!( + parse_absolute_web_image_url(rejected).is_err(), + "URL should be rejected: {rejected}" + ); + } + } + + #[test] + fn chatgpt_web_upload_url_is_restricted_to_signed_storage_origins() { + for accepted in [ + "https://files.oaiusercontent.com/upload/blob?sig=abc&se=123", + "https://cdn.files.oaiusercontent.com/upload/blob?sig=abc", + "https://oaidalleapiprodscus.blob.core.windows.net/container/blob?sig=abc", + "https://tenant.blob.core.windows.net/container/blob?sig=abc", + "https://tenant.blob.core.windows.net:443/container/blob?sig=abc", + ] { + assert!( + validate_chatgpt_web_upload_url(accepted).is_ok(), + "valid storage URL should be accepted: {accepted}" + ); + } + + for rejected in [ + "http://files.oaiusercontent.com/upload/blob?sig=abc", + "https://127.0.0.1/upload/blob?sig=abc", + "https://user:pass@files.oaiusercontent.com/upload/blob?sig=abc", + "https://files.oaiusercontent.com.attacker.invalid/upload/blob?sig=abc", + "https://attacker.invalid/upload/blob?sig=abc", + "https://blob.core.windows.net/upload/blob?sig=abc", + "https://openaiassets.blob.core.windows.net/upload/blob?sig=abc", + "https://tenant.blob.core.windows.net:8443/upload/blob?sig=abc", + "https://tenant.blob.core.windows.net/upload/blob?sig=abc#fragment", + ] { + assert!( + validate_chatgpt_web_upload_url(rejected).is_err(), + "unsafe storage URL should be rejected: {rejected}" + ); + } + + let oversized = format!( + "https://files.oaiusercontent.com/upload/blob?sig={}", + "a".repeat(CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES) + ); + assert!(validate_chatgpt_web_upload_url(&oversized).is_err()); + } + + #[test] + fn chatgpt_web_image_request_fields_are_bounded() { + let oversized_prompt = json!({ + "prompt": "x".repeat(CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES + 1) + }); + assert!(ChatGptWebImageRequest::from_body(&oversized_prompt).is_err()); + + let oversized_images = json!({ + "images": vec!["data:image/png;base64,AA=="; CHATGPT_WEB_IMAGE_MAX_INPUT_IMAGES + 1] + }); + assert!(ChatGptWebImageRequest::from_body(&oversized_images).is_err()); + + let too_many_partial_images = json!({"partial_images": 4}); + assert!(ChatGptWebImageRequest::from_body(&too_many_partial_images).is_err()); + } + + #[test] + fn chatgpt_web_image_summary_bounds_assets_across_merges() { + let mut summary = WebImageSseSummary::default(); + for round in 0..32 { + let mut poll = WebImageSseSummary::default(); + poll.add_values( + WebImageSummaryCollection::FileId, + (0..16).map(|index| format!("file-{round}-{index}")), + ); + poll.add_values( + WebImageSummaryCollection::SedimentId, + (0..16).map(|index| format!("sediment-{round}-{index}")), + ); + poll.add_values( + WebImageSummaryCollection::DirectUrl, + (0..16).map(|index| format!("https://files.oaiusercontent.com/{round}/{index}")), + ); + merge_web_summary(&mut summary, &mut poll); + } + assert!(summary.retained_item_count() <= CHATGPT_WEB_IMAGE_SUMMARY_MAX_ITEMS); + assert!(summary.retained_value_bytes() <= chatgpt_web_image_sse_envelope_limit_bytes()); + } + + #[test] + fn chatgpt_web_image_data_url_is_bounded_before_formatting() { + let oversized = "A".repeat( + maximum_base64_len_for_decoded_limit(chatgpt_web_image_raw_payload_limit_bytes()) + .saturating_add(1), + ); + assert!(bounded_web_image_data_url("image/png", &oversized).is_none()); + assert!(bounded_web_image_data_url("image/png", "AAAA").is_some()); + } + + #[test] + fn chatgpt_web_image_same_origin_requires_scheme_host_and_effective_port() { + let base = url::Url::parse("https://chatgpt.example").expect("base URL should parse"); + + for same_origin in [ + "https://chatgpt.example/backend-api/files/download/file-1", + "https://CHATGPT.example:443/backend-api/files/download/file-1", + ] { + let target = url::Url::parse(same_origin).expect("target URL should parse"); + assert!(web_download_url_is_same_origin(&base, &target)); + } + + for cross_origin in [ + "http://chatgpt.example/backend-api/files/download/file-1", + "https://chatgpt.example:444/backend-api/files/download/file-1", + "https://cdn.chatgpt.example/backend-api/files/download/file-1", + ] { + let target = url::Url::parse(cross_origin).expect("target URL should parse"); + assert!(!web_download_url_is_same_origin(&base, &target)); + } + } + + #[test] + fn chatgpt_web_image_authentication_is_limited_to_same_origin_backend_api_paths() { + let base = url::Url::parse("https://chatgpt.example").expect("base URL should parse"); + let authenticated = + url::Url::parse("https://chatgpt.example/backend-api/files/download/file-1") + .expect("authenticated URL should parse"); + assert!(is_authenticated_web_download_url(&base, &authenticated)); + + for unauthenticated in [ + "https://chatgpt.example/generated.png", + "https://chatgpt.example/backend-api-impersonator/image.png", + "https://cdn.example/backend-api/files/download/file-1", + "http://chatgpt.example/backend-api/files/download/file-1", + "https://chatgpt.example:444/backend-api/files/download/file-1", + "https://user@chatgpt.example/backend-api/files/download/file-1", + ] { + let target = url::Url::parse(unauthenticated).expect("target URL should parse"); + assert!( + !is_authenticated_web_download_url(&base, &target), + "provider credentials must not be sent to {unauthenticated}" + ); + } + } + + #[test] + fn chatgpt_web_data_url_parser_accepts_only_bounded_supported_image_types() { + let payload = base64::engine::general_purpose::STANDARD.encode(png_header_bytes(2, 3)); + let png = + parse_data_url_with_limit(format!("data:image/png;base64,{payload}").as_str(), 64) + .expect("png data URL should parse"); + assert_eq!(png.mime, "image/png"); + assert_eq!(png.b64_json, payload); + + let jpeg_payload = + base64::engine::general_purpose::STANDARD.encode([0xff, 0xd8, 0xff, 0xd9]); + let jpeg = parse_data_url_with_limit( + format!("DATA:IMAGE/JPEG;BASE64,{jpeg_payload}").as_str(), + 64, + ) + .expect("jpeg data URL should parse"); + assert_eq!(jpeg.mime, "image/jpeg"); + + for rejected in [ + "data:text/html;base64,PGh0bWw+", + "data:image/svg+xml;base64,PHN2Zz4=", + "data:image/gif;base64,R0lGODlh", + "data:image/png;base64,", + "data:image/png;base64,!!!!", + "data:image/png;charset=utf-8;base64,aW1hZ2U=", + "data:image/png;base64,aW1h\nZ2U=", + "data:image/png;base64,PHN2Zz4=", + ] { + assert!( + parse_data_url_with_limit(rejected, 64).is_none(), + "unsafe data URL should be rejected: {rejected}" + ); + } + } + + #[test] + fn chatgpt_web_data_url_parser_enforces_decoded_limit_before_allocation() { + let exact_bytes = png_header_bytes(2, 3); + let exact_payload = base64::engine::general_purpose::STANDARD.encode(&exact_bytes); + let exact = parse_data_url_with_limit( + format!("data:image/png;base64,{exact_payload}").as_str(), + exact_bytes.len(), + ) + .expect("payload at the decoded limit should parse"); + assert_eq!(exact.b64_json, exact_payload); + + let exact_len = exact_bytes.len(); + let mut over_bytes = exact_bytes.clone(); + over_bytes.push(0); + let over_payload = base64::engine::general_purpose::STANDARD.encode(over_bytes); + assert!( + parse_data_url_with_limit( + format!("data:image/png;base64,{over_payload}").as_str(), + exact_len, + ) + .is_none(), + "payload over the decoded limit must be rejected" + ); + } + + #[test] + fn chatgpt_web_image_payload_requires_supported_magic_and_matching_mime() { + let png = png_header_bytes(2, 3); + assert_eq!( + validate_web_image_payload(&png, Some("image/png; charset=binary")) + .expect("valid png should pass"), + "image/png" + ); + assert_eq!( + validate_web_image_payload(&png, Some("application/octet-stream")) + .expect("octet-stream with a valid signature should pass"), + "image/png" + ); + assert!(validate_web_image_payload(&png, Some("image/jpeg")).is_err()); + assert!(validate_web_image_payload(&png, Some("image/svg+xml")).is_err()); + assert!(validate_web_image_payload(b"", None).is_err()); + assert!( + validate_web_image_payload(b"not an image", Some("image/png")).is_err() + ); + assert_eq!( + validate_web_image_payload(&[0xff, 0xd8, 0xff, 0xd9], Some("image/jpg")) + .expect("jpeg signature should pass"), + "image/jpeg" + ); + assert_eq!( + validate_web_image_payload(b"RIFF\x04\0\0\0WEBP", None) + .expect("webp signature should pass"), + "image/webp" + ); + assert!(validate_web_image_payload(&[0xff, 0xd8], Some("image/jpeg")).is_err()); + } + + #[test] + fn chatgpt_web_image_sse_envelope_budget_covers_base64_expansion() { + let raw_limit = crate::headers::max_internal_buffered_body_bytes(); + let expected_minimum = maximum_base64_len_for_decoded_limit(raw_limit) + .saturating_add(CHATGPT_WEB_IMAGE_SSE_WRAPPER_OVERHEAD_BYTES) + .min(CHATGPT_WEB_IMAGE_SSE_HARD_MAX_BYTES); + assert!(chatgpt_web_image_sse_envelope_limit_bytes() >= expected_minimum); + assert!(chatgpt_web_image_sse_envelope_limit_bytes() >= raw_limit.min(64 * 1024 * 1024)); + } + + #[test] + fn chatgpt_web_execution_result_body_decode_is_bounded() { + let result = ExecutionResult { + request_id: "req-chatgpt-web-image-test".to_string(), + candidate_id: None, + status_code: 200, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: None, + body_bytes_b64: Some("!!!!".to_string()), + }), + telemetry: None, + error: None, + }; + assert!(execution_result_bytes(&result).is_err()); + assert!(execution_result_body_bytes_lossy(&result).is_empty()); + } + + #[tokio::test] + async fn chatgpt_web_public_image_resolution_rejects_private_ip_literals() { + for private_url in [ + "http://127.0.0.1/image.png", + "http://169.254.169.254/latest/meta-data", + ] { + let url = url::Url::parse(private_url).expect("private URL should parse"); + let error = resolve_public_web_image_addrs(&url, Duration::from_secs(5), false) + .await + .expect_err("private address must be rejected"); + assert!( + error.to_string().contains("private or reserved"), + "unexpected error for {private_url}: {error}" + ); + } + + let public_url = + url::Url::parse("https://8.8.8.8/image.png").expect("public URL should parse"); + let (host, addresses) = + resolve_public_web_image_addrs(&public_url, Duration::from_secs(5), false) + .await + .expect("public IP literal should be accepted"); + assert_eq!(host, "8.8.8.8"); + assert_eq!(addresses, vec!["8.8.8.8:443".parse().unwrap()]); + } + + #[test] + fn chatgpt_web_image_fake_ip_exception_is_limited_to_storage_origins() { + let storage = + url::Url::parse("https://files.oaiusercontent.com/generated/image.png?sig=test") + .expect("storage URL should parse"); + let fake = vec!["198.18.75.234:443".parse().unwrap()]; + assert!(validate_public_web_image_addresses(&storage, &fake, true).is_ok()); + assert!(validate_public_web_image_addresses(&storage, &fake, false).is_err()); + + let arbitrary = url::Url::parse("https://cdn.example/generated/image.png") + .expect("arbitrary URL should parse"); + assert!(validate_public_web_image_addresses(&arbitrary, &fake, true).is_err()); + + let mixed = vec![ + "198.18.75.234:443".parse().unwrap(), + "10.0.0.1:443".parse().unwrap(), + ]; + assert!(validate_public_web_image_addresses(&storage, &mixed, true).is_err()); } #[test] @@ -3872,6 +5337,18 @@ data: [DONE] assert!(!answer.is_empty()); } + #[test] + fn pow_generate_rejects_difficulty_larger_than_digest() { + let seed = "seed"; + let (answer, solved) = pow_generate( + seed, + "f".repeat(130).as_str(), + pow_config(CHATGPT_WEB_USER_AGENT), + ); + assert!(!solved); + assert_eq!(answer, encode_pow_seed(seed)); + } + #[tokio::test] async fn chatgpt_web_image_executor_downloads_file_id_result_as_openai_image_sse() { let (base_url, handle) = start_mock_chatgpt_web().await; diff --git a/apps/aether-gateway/src/execution_runtime/constants.rs b/apps/aether-gateway/src/execution_runtime/constants.rs index e2a98f672..525e59c32 100644 --- a/apps/aether-gateway/src/execution_runtime/constants.rs +++ b/apps/aether-gateway/src/execution_runtime/constants.rs @@ -3,3 +3,13 @@ pub(crate) use aether_gateway_execution::{ MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES, }; + +// Usage/audit captures are secondary copies of the stream. Keep a hard +// ceiling even when the configurable "full" record level is otherwise +// unbounded; this does not limit bytes forwarded to the client. +pub(crate) const MAX_STREAM_BODY_CAPTURE_BYTES: usize = 64 * 1024 * 1024; + +// Stream frames are newline-delimited JSON. Binary response chunks are base64 +// encoded before framing, so this must be larger than the normal 64 MiB raw +// response limit while still bounding an attacker-controlled unterminated line. +pub(crate) const MAX_EXECUTION_STREAM_FRAME_LINE_BYTES: usize = 128 * 1024 * 1024; diff --git a/apps/aether-gateway/src/execution_runtime/fallback.rs b/apps/aether-gateway/src/execution_runtime/fallback.rs index 3c7287b6d..907932abc 100644 --- a/apps/aether-gateway/src/execution_runtime/fallback.rs +++ b/apps/aether-gateway/src/execution_runtime/fallback.rs @@ -156,7 +156,22 @@ pub(crate) fn should_fallback_to_control_sync( return true; }; - body_json.get("error").is_some() + sync_body_has_embedded_error(Some(body_json)) +} + +/// Mirrors the error-like body markers used by the formats layer. Successful OpenAI Responses +/// bodies contain `"error": null`, which must not route them through error finalization. +fn sync_body_has_embedded_error(body_json: Option<&serde_json::Value>) -> bool { + let Some(object) = body_json.and_then(serde_json::Value::as_object) else { + return false; + }; + + object.get("error").is_some_and(|error| !error.is_null()) + || object.get("status").and_then(serde_json::Value::as_str) == Some("failed") + || object + .get("type") + .and_then(serde_json::Value::as_str) + .is_some_and(|value| value == "error") } pub(crate) fn should_finalize_sync_response(report_kind: Option<&str>) -> bool { @@ -168,7 +183,7 @@ pub(crate) fn resolve_core_sync_error_finalize_report_kind( result: &ExecutionResult, body_json: Option<&serde_json::Value>, ) -> Option { - let has_embedded_error = body_json.is_some_and(|value| value.get("error").is_some()); + let has_embedded_error = sync_body_has_embedded_error(body_json); if result.status_code < 400 && !has_embedded_error { return None; } @@ -355,6 +370,7 @@ pub(crate) fn resolve_core_stream_direct_finalize_report_kind(plan_kind: &str) - #[cfg(test)] mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use std::collections::BTreeSet; use aether_contracts::{ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResult}; @@ -435,6 +451,15 @@ mod tests { } fn sample_key() -> StoredProviderCatalogKey { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key("provider-1", "key-1", "plain-upstream-key") + .expect("api key should encrypt"); StoredProviderCatalogKey::new( "key-1".to_string(), "provider-1".to_string(), @@ -446,7 +471,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(serde_json::json!(["openai:chat"])), - "plain-upstream-key".to_string(), + encrypted_api_key, None, None, Some(serde_json::json!({"openai:chat": 1})), @@ -466,7 +491,7 @@ mod tests { ); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( std::sync::Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); AppState::new() .expect("state should build") @@ -500,6 +525,74 @@ mod tests { ); } + #[test] + fn successful_responses_body_with_null_error_stays_on_success_path() { + let result = ExecutionResult { + request_id: "req-1".to_string(), + candidate_id: None, + status_code: 200, + headers: Default::default(), + response_observation: None, + body: None, + telemetry: None, + error: None, + }; + let body_json = serde_json::json!({ + "id": "resp_1", + "object": "response", + "status": "completed", + "error": null, + "output": [], + }); + + assert_eq!( + resolve_core_sync_error_finalize_report_kind( + "openai_responses_sync", + &result, + Some(&body_json) + ), + None + ); + assert!(!should_fallback_to_control_sync( + "openai_responses_sync", + &result, + Some(&body_json), + true, + false, + false, + )); + } + + #[test] + fn error_like_success_status_bodies_still_map_to_error_finalize() { + let result = ExecutionResult { + request_id: "req-1".to_string(), + candidate_id: None, + status_code: 200, + headers: Default::default(), + response_observation: None, + body: None, + telemetry: None, + error: None, + }; + + for body_json in [ + serde_json::json!({"status": "failed", "error": null}), + serde_json::json!({"type": "error"}), + serde_json::json!({"error": {"message": "boom"}}), + ] { + assert_eq!( + resolve_core_sync_error_finalize_report_kind( + "openai_responses_sync", + &result, + Some(&body_json) + ), + Some("openai_responses_sync_finalize".to_string()), + "error-like body must not escape through the success path: {body_json}" + ); + } + } + #[test] fn stream_failover_marks_chat_errors() { assert!(should_fallback_to_control_stream( diff --git a/apps/aether-gateway/src/execution_runtime/grok.rs b/apps/aether-gateway/src/execution_runtime/grok.rs index af27d5750..4aec7e745 100644 --- a/apps/aether-gateway/src/execution_runtime/grok.rs +++ b/apps/aether-gateway/src/execution_runtime/grok.rs @@ -26,15 +26,18 @@ use crate::ai_serving::api::{ CanonicalContentPart, CanonicalStreamEvent, CanonicalStreamFrame, ClaudeClientEmitter, OpenAIChatClientEmitter, OpenAIResponsesClientEmitter, StreamingCanonicalUsage, }; -use crate::ai_serving::openai_responses_synthetic_reasoning_item_id; +use crate::ai_serving::{ + openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, +}; use crate::clock::current_unix_secs; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::transport::{ build_browser_wreq_client, build_request_body, build_request_headers, - decode_response_body_bytes, format_hyper_error_chain, format_upstream_request_error, - format_wreq_upstream_request_error, resolve_stream_first_byte_timeout, send_request, - stream_first_byte_timeout_message, with_non_stream_total_timeout, DirectHttpResponse, - ExecutionRuntimeTransportError, ExecutionTransportControls, + execution_plan_response_body_limit_bytes, format_hyper_error_chain, + format_upstream_request_error, format_wreq_upstream_request_error, + resolve_stream_first_byte_timeout, send_request, stream_first_byte_timeout_message, + with_non_stream_total_timeout, DirectHttpResponse, ExecutionRuntimeTransportError, + ExecutionTransportControls, UpstreamResponseBodyPhase, }; const GROK_INTERNAL_HEADER: &str = "x-aether-grok-runtime"; @@ -44,6 +47,33 @@ const GROK_MEDIA_POST_PATH: &str = "/rest/media/post/create"; const GROK_IMAGINE_WS_URL: &str = "wss://grok.com/ws/imagine/listen"; const GROK_STANDARD_PROVIDER_API_FORMAT: &str = "openai:responses"; const GROK_PROMPT_OVERHEAD_TOKENS: u64 = 4; +const GROK_MAX_ATTACHMENT_BYTES: usize = 64 * 1024 * 1024; +// Attachment uploads each require a fetch/decode plus a provider upload. Cap +// the number independently of the request-body byte limit so a compact JSON +// array cannot turn into an unbounded sequence of outbound requests. +const GROK_MAX_ATTACHMENT_COUNT: usize = 16; +const GROK_MAX_ATTACHMENT_URL_BYTES: usize = 64 * 1024; +const GROK_MAX_ATTACHMENT_FILENAME_BYTES: usize = 1024; +const GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES: usize = 256; +const GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES: usize = 4 * 1024; +// A websocket image response may contain several slots. Bound the aggregate +// encoded blob storage before retaining provider text in the slot map; the +// per-message 64 MiB websocket limit alone would otherwise permit roughly +// 1 GiB across the 16 transient slots. +const GROK_MAX_IMAGINE_BLOB_TOTAL_DECODED_BYTES: usize = 128 * 1024 * 1024; +// Collected image responses are rendered into one client-facing SSE payload +// before being wrapped in a StreamFrame. Keep this compatibility envelope +// bounded independently of the normal streaming path; regular token streams +// remain incremental and are not subject to this aggregate cap. +const GROK_SYNTHETIC_STREAM_ENVELOPE_MAX_BYTES: usize = 256 * 1024 * 1024; +const GROK_SYNTHETIC_STREAM_ENVELOPE_OVERHEAD_BYTES: usize = 64 * 1024; +const GROK_SYNTHETIC_STREAM_BODY_MAX_BYTES: usize = + (GROK_SYNTHETIC_STREAM_ENVELOPE_MAX_BYTES - GROK_SYNTHETIC_STREAM_ENVELOPE_OVERHEAD_BYTES) / 4 + * 3; +const GROK_MAX_IMAGE_COUNT: usize = 4; +// A provider can emit progress frames for more IDs than the client requested. +// Keep transient IDs bounded so a malformed stream cannot grow the map forever. +const GROK_MAX_IMAGINE_SLOTS: usize = 16; const GROK_MAX_ATTACHMENT_REDIRECTS: usize = 5; const GROK_IMAGINE_STREAM_TIMEOUT_MS: u64 = 10_000; const GROK_IMAGINE_ROUND_TIMEOUT_MS: u64 = 120_000; @@ -209,12 +239,14 @@ async fn execute_grok_app_chat( let mut upstream_bytes = 0u64; let mut raw_body = Vec::new(); let mut adapter = GrokStreamAdapter::default(); + let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan); collect_grok_response_stream( response, status_code, &mut upstream_bytes, &mut raw_body, &mut adapter, + response_body_limit_bytes, ) .await?; if (200..300).contains(&status_code) { @@ -223,12 +255,10 @@ async fn execute_grok_app_chat( let elapsed_ms = started_at.elapsed().as_millis() as u64; if !(200..300).contains(&status_code) { - let decoded = decode_response_body_bytes(&headers, &raw_body)?; - let text = String::from_utf8_lossy(decoded.as_ref()).to_string(); return Ok(GrokCollected { status_code, headers, - text, + text: grok_upstream_http_error_message(status_code), telemetry: ExecutionTelemetry { ttfb_ms: Some(ttfb_ms), elapsed_ms: Some(elapsed_ms), @@ -266,21 +296,21 @@ async fn execute_grok_app_chat_stream( let mut upstream_bytes = 0u64; let mut raw_body = Vec::new(); let mut adapter = GrokStreamAdapter::default(); + let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan); collect_grok_response_stream( response, status_code, &mut upstream_bytes, &mut raw_body, &mut adapter, + response_body_limit_bytes, ) .await?; - let decoded = decode_response_body_bytes(&headers, &raw_body)?; - let text = String::from_utf8_lossy(decoded.as_ref()).to_string(); let elapsed_ms = started_at.elapsed().as_millis() as u64; let collected = GrokCollected { status_code, headers, - text, + text: grok_upstream_http_error_message(status_code), telemetry: ExecutionTelemetry { ttfb_ms: Some(elapsed_ms), elapsed_ms: Some(elapsed_ms), @@ -349,7 +379,7 @@ async fn execute_grok_imagine_websocket( "Grok Imagine requires a non-empty prompt".to_string(), ) })?; - let requested = grok_image_count_from_provider_body(body).clamp(1, 4); + let requested = grok_image_count_from_provider_body(body); let enable_pro = grok_upstream_model_name(report_context)? .to_ascii_lowercase() .contains("pro"); @@ -366,7 +396,7 @@ async fn execute_grok_imagine_websocket( .filter_map(|image| { image .url - .or_else(|| image.blob_b64.map(grok_data_image_url)) + .or_else(|| image.blob_b64.and_then(grok_data_image_url)) }) .collect(), telemetry: ExecutionTelemetry { @@ -491,6 +521,7 @@ async fn collect_grok_response_stream( upstream_bytes: &mut u64, raw_body: &mut Vec, adapter: &mut GrokStreamAdapter, + response_body_limit_bytes: usize, ) -> Result<(), ExecutionRuntimeTransportError> { match response { DirectHttpResponse::Reqwest(response) => { @@ -501,7 +532,14 @@ async fn collect_grok_response_stream( &err, )) })?; - collect_grok_response_chunk(status_code, upstream_bytes, raw_body, adapter, &chunk); + collect_grok_response_chunk_with_limit( + status_code, + upstream_bytes, + raw_body, + adapter, + &chunk, + response_body_limit_bytes, + )?; } } DirectHttpResponse::HyperH2c(response) => { @@ -510,7 +548,14 @@ async fn collect_grok_response_stream( let chunk = chunk.map_err(|err| { ExecutionRuntimeTransportError::UpstreamRequest(format_hyper_error_chain(&err)) })?; - collect_grok_response_chunk(status_code, upstream_bytes, raw_body, adapter, &chunk); + collect_grok_response_chunk_with_limit( + status_code, + upstream_bytes, + raw_body, + adapter, + &chunk, + response_body_limit_bytes, + )?; } } DirectHttpResponse::BrowserWreq(response) => { @@ -521,25 +566,44 @@ async fn collect_grok_response_stream( format_wreq_upstream_request_error(&err), ) })?; - collect_grok_response_chunk(status_code, upstream_bytes, raw_body, adapter, &chunk); + collect_grok_response_chunk_with_limit( + status_code, + upstream_bytes, + raw_body, + adapter, + &chunk, + response_body_limit_bytes, + )?; } } } Ok(()) } -fn collect_grok_response_chunk( +fn collect_grok_response_chunk_with_limit( status_code: u16, upstream_bytes: &mut u64, raw_body: &mut Vec, adapter: &mut GrokStreamAdapter, chunk: &[u8], -) { + response_body_limit_bytes: usize, +) -> Result<(), ExecutionRuntimeTransportError> { + if chunk.len() + > response_body_limit_bytes + .saturating_sub(usize::try_from(*upstream_bytes).unwrap_or(usize::MAX)) + { + return Err(ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Wire, + limit_bytes: response_body_limit_bytes, + }); + } *upstream_bytes += chunk.len() as u64; - raw_body.extend_from_slice(chunk); if (200..300).contains(&status_code) { adapter.push_chunk(chunk); + } else { + raw_body.extend_from_slice(chunk); } + Ok(()) } type GrokUpstreamBodyStream = BoxStream<'static, Result>; @@ -589,6 +653,7 @@ fn grok_success_frame_stream( mut body_stream: GrokUpstreamBodyStream, ) -> BoxStream<'static, Result> { let stream_first_byte_timeout = resolve_stream_first_byte_timeout(&plan); + let response_body_limit_bytes = execution_plan_response_body_limit_bytes(&plan); async_stream::stream! { match encode_grok_headers_frame( status_code, @@ -653,6 +718,27 @@ fn grok_success_frame_stream( break; } }; + if chunk.len() + > response_body_limit_bytes.saturating_sub( + usize::try_from(upstream_bytes).unwrap_or(usize::MAX), + ) + { + match encode_grok_error_frame( + ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Wire, + limit_bytes: response_body_limit_bytes, + } + .to_string(), + ) { + Ok(frame) => yield Ok(frame), + Err(err) => { + yield Err(err); + return; + } + } + terminal_error_emitted = true; + break; + } if ttfb_ms.is_none() { ttfb_ms = Some(started_at.elapsed().as_millis() as u64); } @@ -1106,7 +1192,7 @@ async fn attach_grok_uploaded_files( original_body: &Value, upstream_body: &mut Value, ) -> Result<(), ExecutionRuntimeTransportError> { - let inputs = extract_grok_attachment_inputs(plan.client_api_format.as_str(), original_body); + let inputs = extract_grok_attachment_inputs(plan.client_api_format.as_str(), original_body)?; if inputs.is_empty() { return Ok(()); } @@ -1136,7 +1222,7 @@ async fn attach_grok_image_edit_references( original_body: &Value, upstream_body: &mut Value, ) -> Result<(), ExecutionRuntimeTransportError> { - let inputs = extract_grok_attachment_inputs(plan.client_api_format.as_str(), original_body); + let inputs = extract_grok_attachment_inputs(plan.client_api_format.as_str(), original_body)?; if inputs.is_empty() { return Err(ExecutionRuntimeTransportError::UpstreamRequest( "Grok image edit requires at least one reference image".to_string(), @@ -1158,7 +1244,7 @@ async fn attach_grok_image_edit_references( fn extract_grok_attachment_inputs( client_api_format: &str, body: &Value, -) -> Vec { +) -> Result, ExecutionRuntimeTransportError> { match client_api_format.trim().to_ascii_lowercase().as_str() { "openai:responses" | "openai:responses:compact" => { extract_responses_attachment_inputs(body) @@ -1168,103 +1254,173 @@ fn extract_grok_attachment_inputs( } } -fn extract_openai_chat_attachment_inputs(body: &Value) -> Vec { +fn extract_openai_chat_attachment_inputs( + body: &Value, +) -> Result, ExecutionRuntimeTransportError> { let mut out = Vec::new(); if let Some(messages) = body.get("messages").and_then(Value::as_array) { for message in messages { - collect_content_attachment_inputs(message.get("content"), &mut out); + collect_content_attachment_inputs(message.get("content"), &mut out)?; } } - out + Ok(out) } -fn extract_responses_attachment_inputs(body: &Value) -> Vec { +fn extract_responses_attachment_inputs( + body: &Value, +) -> Result, ExecutionRuntimeTransportError> { let mut out = Vec::new(); - collect_responses_input_attachment_inputs(body.get("input"), &mut out); - out + collect_responses_input_attachment_inputs(body.get("input"), &mut out)?; + Ok(out) } fn collect_responses_input_attachment_inputs( value: Option<&Value>, out: &mut Vec, -) { +) -> Result<(), ExecutionRuntimeTransportError> { let Some(value) = value else { - return; + return Ok(()); }; match value { Value::Array(items) => { for item in items { if item.get("type").and_then(Value::as_str) == Some("message") { - collect_content_attachment_inputs(item.get("content"), out); + collect_content_attachment_inputs(item.get("content"), out)?; } else { - collect_attachment_input_from_object(item, out); + collect_attachment_input_from_object(item, out)?; } } } - Value::Object(_) => collect_attachment_input_from_object(value, out), + Value::Object(_) => collect_attachment_input_from_object(value, out)?, _ => {} } + Ok(()) } -fn extract_claude_attachment_inputs(body: &Value) -> Vec { +fn extract_claude_attachment_inputs( + body: &Value, +) -> Result, ExecutionRuntimeTransportError> { let mut out = Vec::new(); if let Some(messages) = body.get("messages").and_then(Value::as_array) { for message in messages { - collect_content_attachment_inputs(message.get("content"), &mut out); + collect_content_attachment_inputs(message.get("content"), &mut out)?; } } - out + Ok(out) } -fn collect_content_attachment_inputs(value: Option<&Value>, out: &mut Vec) { +fn collect_content_attachment_inputs( + value: Option<&Value>, + out: &mut Vec, +) -> Result<(), ExecutionRuntimeTransportError> { let Some(value) = value else { - return; + return Ok(()); }; match value { Value::Array(items) => { for item in items { - collect_attachment_input_from_object(item, out); + collect_attachment_input_from_object(item, out)?; } } - Value::Object(_) => collect_attachment_input_from_object(value, out), + Value::Object(_) => collect_attachment_input_from_object(value, out)?, _ => {} } + Ok(()) } -fn collect_attachment_input_from_object(value: &Value, out: &mut Vec) { +fn collect_attachment_input_from_object( + value: &Value, + out: &mut Vec, +) -> Result<(), ExecutionRuntimeTransportError> { let Some(object) = value.as_object() else { - return; + return Ok(()); }; - if let Some(input) = claude_source_attachment(object) { - out.push(input); - return; + // Do this cheap shape check before parsing/copying any source field. In + // particular, a seventeenth Claude base64 block must not be materialized + // into another multi-megabyte data URI merely to reject it below. + if out.len() >= GROK_MAX_ATTACHMENT_COUNT && object_may_contain_grok_attachment(object) { + return Err(grok_attachment_count_error()); } - if let Some(source) = image_url_source(object) { - out.push(GrokAttachmentInput { - source, - filename: None, - mime_type: None, - }); - return; + if let Some(input) = claude_source_attachment(object)? { + return push_grok_attachment_input(out, input); } - if let Some(input) = file_source(object) { - out.push(input); + if let Some(source) = image_url_source(object)? { + return push_grok_attachment_input( + out, + GrokAttachmentInput { + source, + filename: None, + mime_type: None, + }, + ); } + if let Some(input) = file_source(object)? { + push_grok_attachment_input(out, input)?; + } + Ok(()) } -fn image_url_source(object: &Map) -> Option { - if let Some(source) = object.get("image_url").and_then(string_or_url_value) { - return Some(source); +fn push_grok_attachment_input( + out: &mut Vec, + input: GrokAttachmentInput, +) -> Result<(), ExecutionRuntimeTransportError> { + if out.len() >= GROK_MAX_ATTACHMENT_COUNT { + return Err(grok_attachment_count_error()); + } + out.push(input); + Ok(()) +} + +fn grok_attachment_count_error() -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok requests support at most {GROK_MAX_ATTACHMENT_COUNT} attachments" + )) +} + +fn object_may_contain_grok_attachment(object: &Map) -> bool { + if object.contains_key("image_url") + || object.contains_key("file_data") + || object.contains_key("file_url") + || object.contains_key("file") + { + return true; + } + object + .get("type") + .and_then(Value::as_str) + .map(str::trim) + .is_some_and(|value| { + [ + "image_url", + "input_image", + "input_file", + "file", + "image", + "document", + ] + .iter() + .any(|candidate| value.eq_ignore_ascii_case(candidate)) + }) +} + +fn image_url_source( + object: &Map, +) -> Result, ExecutionRuntimeTransportError> { + if let Some(source) = object + .get("image_url") + .map(string_or_url_value) + .transpose()? + .flatten() + { + return Ok(Some(source)); } if object .get("type") .and_then(Value::as_str) .is_some_and(|value| value.eq_ignore_ascii_case("image_url")) { - return object - .get("url") - .and_then(Value::as_str) - .map(trimmed_string); + return bounded_grok_attachment_source(object.get("url").and_then(Value::as_str)) + .map(|value| value.map(ToOwned::to_owned)); } if object .get("type") @@ -1274,14 +1430,18 @@ fn image_url_source(object: &Map) -> Option { return object .get("image_url") .or_else(|| object.get("source")) - .and_then(string_or_url_value); + .map(string_or_url_value) + .transpose() + .map(Option::flatten); } - None + Ok(None) } -fn file_source(object: &Map) -> Option { +fn file_source( + object: &Map, +) -> Result, ExecutionRuntimeTransportError> { let file_object = object.get("file").and_then(Value::as_object); - let source = file_object + let raw_source = file_object .and_then(|file| { file.get("file_data") .or_else(|| file.get("data")) @@ -1291,96 +1451,171 @@ fn file_source(object: &Map) -> Option { }) .or_else(|| object.get("file_data").and_then(Value::as_str)) .or_else(|| object.get("file_url").and_then(Value::as_str)) - .or_else(|| object.get("data").and_then(Value::as_str)) - .map(trimmed_string) - .filter(|value| !value.is_empty())?; - let filename = file_object + .or_else(|| object.get("data").and_then(Value::as_str)); + let Some(source) = bounded_grok_attachment_source(raw_source)? else { + return Ok(None); + }; + let raw_filename = file_object .and_then(|file| file.get("filename").or_else(|| file.get("name"))) .or_else(|| object.get("filename")) .or_else(|| object.get("name")) - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty()); - let mime_type = file_object + .and_then(Value::as_str); + let raw_mime_type = file_object .and_then(|file| file.get("mime_type").or_else(|| file.get("mimeType"))) .or_else(|| object.get("mime_type")) .or_else(|| object.get("mimeType")) - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty()); - Some(GrokAttachmentInput { - source, + .and_then(Value::as_str); + let filename = bounded_grok_attachment_field( + raw_filename, + "filename", + GROK_MAX_ATTACHMENT_FILENAME_BYTES, + )? + .map(ToOwned::to_owned); + let mime_type = bounded_grok_attachment_field( + raw_mime_type, + "MIME type", + GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES, + )? + .map(ToOwned::to_owned); + Ok(Some(GrokAttachmentInput { + source: source.to_owned(), filename, mime_type, - }) + })) } -fn claude_source_attachment(object: &Map) -> Option { +fn claude_source_attachment( + object: &Map, +) -> Result, ExecutionRuntimeTransportError> { let block_type = object .get("type") .and_then(Value::as_str) .map(str::trim) .unwrap_or_default(); if !matches!(block_type, "image" | "document") { - return None; + return Ok(None); } - let source = object.get("source").and_then(Value::as_object)?; + let Some(source) = object.get("source").and_then(Value::as_object) else { + return Ok(None); + }; let source_type = source .get("type") .and_then(Value::as_str) .map(str::trim) .unwrap_or_default(); - let mime_type = source + let raw_mime_type = source .get("media_type") .or_else(|| source.get("mediaType")) - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty()); - let filename = object + .and_then(Value::as_str); + let raw_filename = object .get("filename") .or_else(|| object.get("name")) - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty()); + .and_then(Value::as_str); + let mime_type = bounded_grok_attachment_field( + raw_mime_type, + "MIME type", + GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES, + )? + .map(ToOwned::to_owned); + let filename = bounded_grok_attachment_field( + raw_filename, + "filename", + GROK_MAX_ATTACHMENT_FILENAME_BYTES, + )? + .map(ToOwned::to_owned); match source_type { "base64" => { - let data = source - .get("data") - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty())?; - let mime = mime_type - .clone() - .unwrap_or_else(|| "application/octet-stream".to_string()); - Some(GrokAttachmentInput { - source: format!("data:{mime};base64,{data}"), + let Some(data) = bounded_grok_attachment_field( + source.get("data").and_then(Value::as_str), + "base64 data", + maximum_base64_len_for_decoded_limit(GROK_MAX_ATTACHMENT_BYTES), + )? + else { + return Ok(None); + }; + let mime = mime_type.as_deref().unwrap_or("application/octet-stream"); + let mut data_uri = String::with_capacity( + mime.len() + .saturating_add(data.len()) + .saturating_add("data:;base64,".len()), + ); + data_uri.push_str("data:"); + data_uri.push_str(mime); + data_uri.push_str(";base64,"); + data_uri.push_str(data); + Ok(Some(GrokAttachmentInput { + source: data_uri, filename, mime_type, - }) + })) } "url" => { - let url = source - .get("url") - .and_then(Value::as_str) - .map(trimmed_string) - .filter(|value| !value.is_empty())?; - Some(GrokAttachmentInput { - source: url, + let Some(url) = + bounded_grok_attachment_source(source.get("url").and_then(Value::as_str))? + else { + return Ok(None); + }; + Ok(Some(GrokAttachmentInput { + source: url.to_owned(), filename, mime_type, - }) + })) } - _ => None, + _ => Ok(None), } } -fn string_or_url_value(value: &Value) -> Option { - value +fn string_or_url_value(value: &Value) -> Result, ExecutionRuntimeTransportError> { + let raw = value .as_str() - .map(trimmed_string) - .or_else(|| value.get("url").and_then(Value::as_str).map(trimmed_string)) - .filter(|value| !value.is_empty()) + .or_else(|| value.get("url").and_then(Value::as_str)); + bounded_grok_attachment_source(raw).map(|value| value.map(ToOwned::to_owned)) +} + +fn bounded_grok_attachment_source( + raw: Option<&str>, +) -> Result, ExecutionRuntimeTransportError> { + let Some(value) = raw.map(str::trim).filter(|value| !value.is_empty()) else { + return Ok(None); + }; + let max_bytes = if grok_is_data_uri(value) { + let metadata_bytes = value.find(',').unwrap_or(value.len()); + if metadata_bytes > GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES { + return Err(grok_attachment_field_too_large( + "data URI metadata", + GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES, + )); + } + maximum_base64_len_for_decoded_limit(GROK_MAX_ATTACHMENT_BYTES) + .saturating_add(GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES) + } else { + GROK_MAX_ATTACHMENT_URL_BYTES + }; + bounded_grok_attachment_field(Some(value), "source", max_bytes) +} + +fn bounded_grok_attachment_field<'a>( + raw: Option<&'a str>, + field: &str, + max_bytes: usize, +) -> Result, ExecutionRuntimeTransportError> { + let Some(value) = raw.map(str::trim).filter(|value| !value.is_empty()) else { + return Ok(None); + }; + if value.len() > max_bytes { + return Err(grok_attachment_field_too_large(field, max_bytes)); + } + Ok(Some(value)) +} + +fn grok_attachment_field_too_large( + field: &str, + max_bytes: usize, +) -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok attachment {field} exceeds {max_bytes} byte limit" + )) } fn trimmed_string(value: &str) -> String { @@ -1391,47 +1626,123 @@ async fn resolve_grok_attachment_payload( input: &GrokAttachmentInput, index: usize, ) -> Result { - if input.source.starts_with("data:") { + validate_grok_attachment_input_fields(input)?; + if grok_is_data_uri(input.source.as_str()) { return grok_attachment_payload_from_data_uri(input, index); } grok_attachment_payload_from_url(input, index).await } +fn validate_grok_attachment_input_fields( + input: &GrokAttachmentInput, +) -> Result<(), ExecutionRuntimeTransportError> { + let source = input.source.trim(); + if source.is_empty() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment source is empty".to_string(), + )); + } + let max_source_bytes = if grok_is_data_uri(source) { + let metadata_bytes = source.find(',').unwrap_or(source.len()); + if metadata_bytes > GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES { + return Err(grok_attachment_field_too_large( + "data URI metadata", + GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES, + )); + } + maximum_base64_len_for_decoded_limit(GROK_MAX_ATTACHMENT_BYTES) + .saturating_add(GROK_MAX_ATTACHMENT_DATA_URI_METADATA_BYTES) + } else { + GROK_MAX_ATTACHMENT_URL_BYTES + }; + if source.len() > max_source_bytes { + return Err(grok_attachment_field_too_large("source", max_source_bytes)); + } + if input + .filename + .as_deref() + .is_some_and(|filename| filename.trim().len() > GROK_MAX_ATTACHMENT_FILENAME_BYTES) + { + return Err(grok_attachment_field_too_large( + "filename", + GROK_MAX_ATTACHMENT_FILENAME_BYTES, + )); + } + if input + .mime_type + .as_deref() + .is_some_and(|mime_type| mime_type.trim().len() > GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES) + { + return Err(grok_attachment_field_too_large( + "MIME type", + GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES, + )); + } + Ok(()) +} + fn grok_attachment_payload_from_data_uri( input: &GrokAttachmentInput, index: usize, ) -> Result { - let (header, content_b64) = input.source.split_once(',').ok_or_else(|| { + grok_attachment_payload_from_data_uri_with_limit(input, index, GROK_MAX_ATTACHMENT_BYTES) +} + +fn grok_attachment_payload_from_data_uri_with_limit( + input: &GrokAttachmentInput, + index: usize, + limit_bytes: usize, +) -> Result { + let source = input.source.trim(); + let (header, content_b64) = source.split_once(',').ok_or_else(|| { ExecutionRuntimeTransportError::UpstreamRequest( "Grok attachment data URI is missing comma separator".to_string(), ) })?; - if !header.contains(";base64") { + if !header + .split(';') + .skip(1) + .any(|parameter| parameter.trim().eq_ignore_ascii_case("base64")) + { return Err(ExecutionRuntimeTransportError::UpstreamRequest( "Grok attachment data URI must be base64 encoded".to_string(), )); } + let header_mime = header + .get(..5) + .filter(|prefix| prefix.eq_ignore_ascii_case("data:")) + .and_then(|_| header.get(5..)) + .and_then(|value| value.split(';').next()) + .map(str::trim) + .filter(|value| !value.is_empty()); + if header_mime.is_some_and(|value| value.len() > GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES) { + return Err(grok_attachment_field_too_large( + "MIME type", + GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES, + )); + } let mime_type = input .mime_type - .clone() - .or_else(|| { - header - .strip_prefix("data:") - .and_then(|value| value.split(';').next()) - .map(trimmed_string) - .filter(|value| !value.is_empty()) - }) + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| header_mime.map(ToOwned::to_owned)) .unwrap_or_else(|| "application/octet-stream".to_string()); - let normalized_b64 = content_b64.split_whitespace().collect::(); - drop( - base64::engine::general_purpose::STANDARD - .decode(&normalized_b64) - .map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok attachment data URI base64 is invalid: {err}" - )) - })?, - ); + let max_base64_len = maximum_base64_len_for_decoded_limit(limit_bytes); + let normalized_b64 = normalize_base64_with_limit(content_b64, max_base64_len) + .map_err(|_| grok_attachment_too_large(limit_bytes))?; + let decoded_len = base64::engine::general_purpose::STANDARD + .decode(&normalized_b64) + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok attachment data URI base64 is invalid: {err}" + )) + })? + .len(); + if decoded_len > limit_bytes { + return Err(grok_attachment_too_large(limit_bytes)); + } Ok(GrokAttachmentPayload { filename: input .filename @@ -1442,6 +1753,30 @@ fn grok_attachment_payload_from_data_uri( }) } +fn normalize_base64_with_limit(input: &str, limit_bytes: usize) -> Result { + // Do not reserve based on the untrusted URI length. Whitespace is ignored, + // so an attacker could otherwise force a large allocation before validation. + let mut normalized = String::with_capacity(limit_bytes.min(4096)); + for character in input.chars() { + if character.is_whitespace() { + continue; + } + if normalized.len().saturating_add(character.len_utf8()) > limit_bytes { + return Err(()); + } + normalized.push(character); + } + Ok(normalized) +} + +fn maximum_base64_len_for_decoded_limit(limit_bytes: usize) -> usize { + limit_bytes + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4) +} + async fn grok_attachment_payload_from_url( input: &GrokAttachmentInput, index: usize, @@ -1451,25 +1786,24 @@ async fn grok_attachment_payload_from_url( "Grok attachment URL is invalid: {err}" )) })?; - if !matches!(url.scheme(), "http" | "https") { - return Err(ExecutionRuntimeTransportError::UpstreamRequest( - "Grok attachment URL must use http or https".to_string(), - )); - } + validate_grok_attachment_url(&url)?; let response = fetch_grok_attachment_url(url.clone(), 0).await?; let final_url = response.url().clone(); + let response_mime_type = response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.split(';').next()); + let response_mime_type = bounded_grok_attachment_field( + response_mime_type, + "MIME type", + GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES, + )? + .map(ToOwned::to_owned); let mime_type = input .mime_type .clone() - .or_else(|| { - response - .headers() - .get(reqwest::header::CONTENT_TYPE) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.split(';').next()) - .map(trimmed_string) - .filter(|value| !value.is_empty()) - }) + .or(response_mime_type) .unwrap_or_else(|| "application/octet-stream".to_string()); let bytes = collect_grok_attachment_url_bytes(response).await?; Ok(GrokAttachmentPayload { @@ -1488,14 +1822,16 @@ async fn fetch_grok_attachment_url( mut redirects: usize, ) -> Result { loop { - validate_grok_attachment_public_url(&url).await?; + // Validate every hop. A relative Location may inherit credentials or + // a fragment from the previous URL, while an absolute Location can + // introduce either explicitly. + validate_grok_attachment_url(&url)?; + let public_addr = public_socket_addr_for_url(&url).await?; let response = reqwest::Client::builder() + .no_proxy() .timeout(Duration::from_secs(30)) .redirect(reqwest::redirect::Policy::none()) - .resolve_to_addrs( - url.host_str().unwrap_or_default(), - &[public_socket_addr_for_url(&url).await?], - ) + .resolve_to_addrs(url.host_str().unwrap_or_default(), &[public_addr]) .build() .map_err(ExecutionRuntimeTransportError::ClientBuild)? .get(url.clone()) @@ -1537,22 +1873,34 @@ async fn fetch_grok_attachment_url( } } -async fn validate_grok_attachment_public_url( - url: &reqwest::Url, -) -> Result<(), ExecutionRuntimeTransportError> { +fn validate_grok_attachment_url(url: &reqwest::Url) -> Result<(), ExecutionRuntimeTransportError> { if !matches!(url.scheme(), "http" | "https") { - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok attachment URL scheme is unsupported: {}", - url.scheme() - ))); + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL must use http or https".to_string(), + )); } - public_socket_addr_for_url(url).await.map(|_| ()) + if url.host_str().is_none() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL is missing a host".to_string(), + )); + } + if !url.username().is_empty() || url.password().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL must not contain credentials".to_string(), + )); + } + if url.fragment().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL must not contain a fragment".to_string(), + )); + } + Ok(()) } async fn public_socket_addr_for_url( url: &reqwest::Url, ) -> Result { - let host = url.host_str().ok_or_else(|| { + let host = url.host().ok_or_else(|| { ExecutionRuntimeTransportError::UpstreamRequest( "Grok attachment URL is missing a host".to_string(), ) @@ -1562,6 +1910,27 @@ async fn public_socket_addr_for_url( "Grok attachment URL is missing a port".to_string(), ) })?; + let host = match host { + url::Host::Ipv4(ip) => { + let ip = IpAddr::V4(ip); + if !grok_attachment_ip_is_public(ip) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL resolves to a non-public address".to_string(), + )); + } + return Ok(std::net::SocketAddr::new(ip, port)); + } + url::Host::Ipv6(ip) => { + let ip = IpAddr::V6(ip); + if !grok_attachment_ip_is_public(ip) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok attachment URL resolves to a non-public address".to_string(), + )); + } + return Ok(std::net::SocketAddr::new(ip, port)); + } + url::Host::Domain(host) => host, + }; if let Ok(ip) = host.parse::() { if !grok_attachment_ip_is_public(ip) { return Err(ExecutionRuntimeTransportError::UpstreamRequest( @@ -1572,11 +1941,15 @@ async fn public_socket_addr_for_url( } let mut public_addr = None; let mut resolved_any = false; - for addr in tokio::net::lookup_host((host, port)).await.map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok attachment URL DNS resolution failed: {err}" - )) - })? { + for addr in + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok attachment URL DNS resolution failed: {err}" + )) + })? + { resolved_any = true; if !grok_attachment_ip_is_public(addr.ip()) { return Err(ExecutionRuntimeTransportError::UpstreamRequest( @@ -1598,44 +1971,38 @@ async fn public_socket_addr_for_url( } fn grok_attachment_ip_is_public(ip: IpAddr) -> bool { - match ip { - IpAddr::V4(ip) => { - !(ip.is_private() - || ip.is_loopback() - || ip.is_link_local() - || ip.is_broadcast() - || ip.is_documentation() - || ip.is_unspecified()) - } - IpAddr::V6(ip) => { - !(ip.is_loopback() - || ip.is_unspecified() - || ip.is_unique_local() - || ip.is_unicast_link_local() - || is_ipv6_documentation_addr(ip)) - } - } -} - -fn is_ipv6_documentation_addr(ip: std::net::Ipv6Addr) -> bool { - let segments = ip.segments(); - segments[0] == 0x2001 && segments[1] == 0x0db8 + !aether_http::is_private_or_reserved_ip(ip) } async fn collect_grok_attachment_url_bytes( response: reqwest::Response, ) -> Result, ExecutionRuntimeTransportError> { + if response + .content_length() + .is_some_and(|length| length > GROK_MAX_ATTACHMENT_BYTES as u64) + { + return Err(grok_attachment_too_large(GROK_MAX_ATTACHMENT_BYTES)); + } let mut bytes = Vec::new(); let mut stream = response.bytes_stream(); while let Some(chunk) = stream.next().await { let chunk = chunk.map_err(|err| { ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err)) })?; + if chunk.len() > GROK_MAX_ATTACHMENT_BYTES.saturating_sub(bytes.len()) { + return Err(grok_attachment_too_large(GROK_MAX_ATTACHMENT_BYTES)); + } bytes.extend_from_slice(&chunk); } Ok(bytes) } +fn grok_attachment_too_large(limit_bytes: usize) -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok attachment exceeds {limit_bytes} byte limit" + )) +} + async fn upload_grok_attachment( plan: &ExecutionPlan, payload: GrokAttachmentPayload, @@ -1655,12 +2022,11 @@ async fn upload_grok_attachment( let request_body = build_request_body(&upload_plan)?; let response = send_request(&upload_plan, request_body).await?; let status_code = response.status_code(); - let bytes = response.bytes().await?; + let bytes = response + .bytes_with_limit(execution_plan_response_body_limit_bytes(&upload_plan)) + .await?; if !(200..300).contains(&status_code) { - let text = String::from_utf8_lossy(&bytes).to_string(); - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok attachment upload returned {status_code}: {text}" - ))); + return Err(grok_auxiliary_http_error("attachment upload", status_code)); } let value = serde_json::from_slice::(&bytes) .map_err(ExecutionRuntimeTransportError::InvalidJson)?; @@ -1705,12 +2071,11 @@ async fn create_grok_media_post( let request_body = build_request_body(&media_plan)?; let response = send_request(&media_plan, request_body).await?; let status_code = response.status_code(); - let bytes = response.bytes().await?; + let bytes = response + .bytes_with_limit(execution_plan_response_body_limit_bytes(&media_plan)) + .await?; if !(200..300).contains(&status_code) { - let text = String::from_utf8_lossy(&bytes).to_string(); - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok media post create returned {status_code}: {text}" - ))); + return Err(grok_auxiliary_http_error("media post create", status_code)); } let value = serde_json::from_slice::(&bytes) .map_err(ExecutionRuntimeTransportError::InvalidJson)?; @@ -1858,6 +2223,7 @@ fn grok_image_count_from_provider_body(body: &Value) -> usize { }) .and_then(|value| usize::try_from(value).ok()) .unwrap_or(1) + .clamp(1, GROK_MAX_IMAGE_COUNT) } fn grok_aspect_ratio_from_provider_body(body: &Value) -> String { @@ -1978,28 +2344,42 @@ fn grok_imagine_request_message(prompt: &str, aspect_ratio: &str, enable_pro: bo fn grok_handle_imagine_ws_message( value: &Value, slots: &mut BTreeMap, +) -> Result<(), ExecutionRuntimeTransportError> { + grok_handle_imagine_ws_message_with_slot_limit(value, slots, GROK_MAX_IMAGINE_SLOTS) +} + +fn grok_handle_imagine_ws_message_with_slot_limit( + value: &Value, + slots: &mut BTreeMap, + max_slots: usize, ) -> Result<(), ExecutionRuntimeTransportError> { match value.get("type").and_then(Value::as_str) { - Some("json") => grok_handle_imagine_json_frame(value, slots), + Some("json") => grok_handle_imagine_json_frame(value, slots, max_slots), Some("image") => { - grok_handle_imagine_image_frame(value, slots); + grok_handle_imagine_image_frame(value, slots, max_slots); Ok(()) } - Some("error") => Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok Imagine websocket error: {}", - value - .get("err_msg") - .or_else(|| value.get("error")) - .and_then(Value::as_str) - .unwrap_or("unknown") - ))), + Some("error") => Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Grok Imagine websocket returned an error".to_string(), + )), _ => Ok(()), } } +fn grok_auxiliary_http_error(stage: &str, status_code: u16) -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok {stage} returned HTTP {status_code}" + )) +} + +fn grok_upstream_http_error_message(status_code: u16) -> String { + format!("Grok upstream request returned HTTP {status_code}") +} + fn grok_handle_imagine_json_frame( value: &Value, slots: &mut BTreeMap, + max_slots: usize, ) -> Result<(), ExecutionRuntimeTransportError> { let status = value.get("current_status").and_then(Value::as_str); let Some(image_id) = value @@ -2016,6 +2396,9 @@ fn grok_handle_imagine_json_frame( .and_then(Value::as_u64) .and_then(|value| usize::try_from(value).ok()) .unwrap_or_default(); + if !slots.contains_key(&image_id) && slots.len() >= max_slots { + return Ok(()); + } match status { Some("start_stage") => { slots.entry(image_id.clone()).or_insert(GrokImagineImage { @@ -2048,12 +2431,57 @@ fn grok_handle_imagine_json_frame( Ok(()) } -fn grok_handle_imagine_image_frame(value: &Value, slots: &mut BTreeMap) { - let Some(url) = value.get("url").and_then(Value::as_str).map(grok_asset_url) else { +fn grok_handle_imagine_image_frame( + value: &Value, + slots: &mut BTreeMap, + max_slots: usize, +) { + let Some(raw_url) = value + .get("url") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { return; }; + if raw_url.len() > GROK_MAX_ATTACHMENT_URL_BYTES { + return; + } + let url = grok_asset_url(raw_url); + if url.len() > GROK_MAX_ATTACHMENT_URL_BYTES { + return; + } let image_id = grok_imagine_image_id_from_url(&url).unwrap_or_else(|| Uuid::new_v4().to_string()); + if !slots.contains_key(&image_id) && slots.len() >= max_slots { + return; + } + let blob = value + .get("blob") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + let Some(blob) = blob else { + let fallback_order = slots.len(); + let slot = slots.entry(image_id.clone()).or_insert(GrokImagineImage { + image_id, + order: fallback_order, + url: None, + blob_b64: None, + done: false, + moderated: false, + }); + slot.url = Some(url); + return; + }; + let existing_blob_len = slots + .get(&image_id) + .and_then(|slot| slot.blob_b64.as_ref()) + .map(String::len) + .unwrap_or(0); + if !grok_imagine_blob_can_be_retained(slots, existing_blob_len, blob.len()) { + return; + } let fallback_order = slots.len(); let slot = slots.entry(image_id.clone()).or_insert(GrokImagineImage { image_id, @@ -2064,11 +2492,39 @@ fn grok_handle_imagine_image_frame(value: &Value, slots: &mut BTreeMap, + existing_blob_len: usize, + incoming_blob_len: usize, +) -> bool { + let retained_blob_bytes = slots.values().fold(0usize, |total, slot| { + total.saturating_add(slot.blob_b64.as_ref().map(String::len).unwrap_or(0)) + }); + grok_imagine_blob_lengths_can_be_retained( + retained_blob_bytes, + existing_blob_len, + incoming_blob_len, + ) +} + +fn grok_imagine_blob_lengths_can_be_retained( + retained_blob_bytes: usize, + existing_blob_len: usize, + incoming_blob_len: usize, +) -> bool { + let per_blob_limit = maximum_base64_len_for_decoded_limit(GROK_MAX_ATTACHMENT_BYTES); + if incoming_blob_len > per_blob_limit { + return false; + } + let total_blob_limit = + maximum_base64_len_for_decoded_limit(GROK_MAX_IMAGINE_BLOB_TOTAL_DECODED_BYTES); + retained_blob_bytes + .saturating_sub(existing_blob_len) + .saturating_add(incoming_blob_len) + <= total_blob_limit } fn grok_imagine_image_id_from_url(url: &str) -> Option { @@ -2090,8 +2546,22 @@ fn grok_imagine_completed_count(slots: &BTreeMap) -> u .count() } -fn grok_data_image_url(blob_b64: String) -> String { - format!("data:image/png;base64,{blob_b64}") +fn grok_data_image_url(blob_b64: String) -> Option { + // A websocket `blob` is untrusted provider data. Do not pass it through + // as an arbitrary data URL: malformed base64 and active formats such as + // SVG/HTML could otherwise cross the public image boundary. Validate the + // encoded and decoded sizes plus the raster magic before retaining it. + // Bound the interpolation itself before constructing the candidate URL; + // otherwise formatting would allocate from an attacker-controlled blob + // before the parser has a chance to reject it. + if blob_b64.len() > maximum_base64_len_for_decoded_limit(GROK_MAX_ATTACHMENT_BYTES) { + return None; + } + let mut candidate = String::with_capacity(22usize.saturating_add(blob_b64.len())); + candidate.push_str("data:image/png;base64,"); + candidate.push_str(&blob_b64); + grok_data_image_parts(&candidate)?; + Some(candidate) } fn set_grok_image_edit_config( @@ -2119,10 +2589,11 @@ fn set_grok_image_edit_config( } fn filename_from_url_path(path: &str) -> Option { - path.rsplit('/') - .next() - .map(trimmed_string) - .filter(|value| !value.is_empty()) + let filename = path.rsplit('/').next()?.trim(); + if filename.is_empty() || filename.len() > GROK_MAX_ATTACHMENT_FILENAME_BYTES { + return None; + } + Some(filename.to_owned()) } fn default_attachment_filename(index: usize, mime_type: &str) -> String { @@ -2130,9 +2601,10 @@ fn default_attachment_filename(index: usize, mime_type: &str) -> String { .rsplit('/') .next() .map(|value| value.split('+').next().unwrap_or(value)) - .map(trimmed_string) + .map(str::trim) .filter(|value| !value.is_empty()) - .unwrap_or_else(|| "bin".to_string()); + .filter(|value| value.len() <= GROK_MAX_ATTACHMENT_FILENAME_BYTES) + .unwrap_or("bin"); format!("file-{}.{}", index + 1, ext) } @@ -2147,7 +2619,7 @@ fn grok_execution_result( } else { json!({ "error": { - "message": collected.text, + "message": grok_upstream_http_error_message(status_code), "type": "grok_upstream_error", "code": status_code, } @@ -2206,48 +2678,74 @@ fn grok_collected_frame_stream( collected: GrokCollected, report_context: Option<&Value>, ) -> BoxStream<'static, Result> { - let body = grok_client_stream_body(&plan, &collected, report_context); let telemetry = collected.telemetry.clone(); let status_code = collected.status_code; - let frames = vec![ - StreamFrame { - frame_type: StreamFrameType::Headers, - payload: StreamFramePayload::Headers { - status_code, - headers: BTreeMap::from([( - "content-type".to_string(), - if (200..300).contains(&status_code) { - "text/event-stream".to_string() - } else { - "application/json".to_string() - }, - )]), - response_observation: None, + let headers_frame = || StreamFrame { + frame_type: StreamFrameType::Headers, + payload: StreamFramePayload::Headers { + status_code, + headers: BTreeMap::from([( + "content-type".to_string(), + if (200..300).contains(&status_code) { + "text/event-stream".to_string() + } else { + "application/json".to_string() + }, + )]), + response_observation: None, + }, + }; + let initial_telemetry_frame = || StreamFrame { + frame_type: StreamFrameType::Telemetry, + payload: StreamFramePayload::Telemetry { + telemetry: ExecutionTelemetry { + ttfb_ms: telemetry.ttfb_ms, + elapsed_ms: telemetry.ttfb_ms, + upstream_bytes: Some(0), }, }, - StreamFrame { - frame_type: StreamFrameType::Telemetry, - payload: StreamFramePayload::Telemetry { - telemetry: ExecutionTelemetry { - ttfb_ms: telemetry.ttfb_ms, - elapsed_ms: telemetry.ttfb_ms, - upstream_bytes: Some(0), + }; + let frames = match bounded_grok_client_stream_body(&plan, &collected, report_context) { + Ok(body) => vec![ + headers_frame(), + initial_telemetry_frame(), + StreamFrame { + frame_type: StreamFrameType::Data, + payload: StreamFramePayload::Data { + chunk_b64: Some( + base64::engine::general_purpose::STANDARD.encode(body.as_bytes()), + ), + text: None, }, }, - }, - StreamFrame { - frame_type: StreamFrameType::Data, - payload: StreamFramePayload::Data { - chunk_b64: Some(base64::engine::general_purpose::STANDARD.encode(body.as_bytes())), - text: None, + StreamFrame { + frame_type: StreamFrameType::Telemetry, + payload: StreamFramePayload::Telemetry { telemetry }, }, - }, - StreamFrame { - frame_type: StreamFrameType::Telemetry, - payload: StreamFramePayload::Telemetry { telemetry }, - }, - StreamFrame::eof(), - ]; + StreamFrame::eof(), + ], + Err(error) => vec![ + headers_frame(), + StreamFrame { + frame_type: StreamFrameType::Error, + payload: StreamFramePayload::Error { + error: aether_contracts::ExecutionError { + kind: aether_contracts::ExecutionErrorKind::ProtocolError, + phase: aether_contracts::ExecutionPhase::Finalize, + message: error.to_string(), + upstream_status: None, + retryable: true, + failover_recommended: true, + }, + }, + }, + StreamFrame { + frame_type: StreamFrameType::Telemetry, + payload: StreamFramePayload::Telemetry { telemetry }, + }, + StreamFrame::eof_with_summary(None), + ], + }; stream::iter( frames .into_iter() @@ -2256,6 +2754,21 @@ fn grok_collected_frame_stream( .boxed() } +fn bounded_grok_client_stream_body( + plan: &ExecutionPlan, + collected: &GrokCollected, + report_context: Option<&Value>, +) -> Result { + let body = grok_client_stream_body(plan, collected, report_context); + if body.len() > GROK_SYNTHETIC_STREAM_BODY_MAX_BYTES { + return Err(ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Decoded, + limit_bytes: GROK_SYNTHETIC_STREAM_ENVELOPE_MAX_BYTES, + }); + } + Ok(body) +} + fn grok_client_stream_body( plan: &ExecutionPlan, collected: &GrokCollected, @@ -2264,7 +2777,7 @@ fn grok_client_stream_body( if !(200..300).contains(&collected.status_code) { return serde_json::to_string(&json!({ "error": { - "message": collected.text, + "message": grok_upstream_http_error_message(collected.status_code), "type": "grok_upstream_error", "code": collected.status_code, } @@ -2563,8 +3076,15 @@ impl GrokStreamAdapter { } fn push_image_url(&mut self, url: String) { - if !url.trim().is_empty() && !self.images.iter().any(|item| item == &url) { - self.images.push(url); + let url = url.trim(); + if url.is_empty() + || url.len() > GROK_MAX_ATTACHMENT_URL_BYTES + || self.images.len() >= GROK_MAX_IMAGE_COUNT + { + return; + } + if !self.images.iter().any(|item| item == url) { + self.images.push(url.to_owned()); } } @@ -2725,7 +3245,7 @@ fn openai_responses_body( }; if !message_text.trim().is_empty() { output.push(json!({ - "id": format!("{response_id}_msg"), + "id": openai_responses_message_item_id(response_id.as_str(), output.len()), "type": "message", "role": "assistant", "content": [{"type": "output_text", "text": message_text, "annotations": []}], @@ -2734,11 +3254,13 @@ fn openai_responses_body( } if images_as_generation_calls { for (index, image) in collected.images.iter().enumerate() { - output.push(grok_openai_responses_image_generation_item( + if let Some(item) = grok_openai_responses_image_generation_item( response_id.as_str(), index, image.as_str(), - )); + ) { + output.push(item); + } } } json!({ @@ -2756,7 +3278,11 @@ fn grok_openai_responses_image_generation_item( response_id: &str, index: usize, image: &str, -) -> Value { +) -> Option { + let image = image.trim(); + if image.is_empty() { + return None; + } let mut item = Map::new(); item.insert( "id".to_string(), @@ -2768,7 +3294,8 @@ fn grok_openai_responses_image_generation_item( ); item.insert("status".to_string(), Value::String("completed".to_string())); item.insert("action".to_string(), Value::String("generate".to_string())); - if let Some((mime_type, b64_json)) = grok_data_image_parts(image) { + if grok_is_data_uri(image) { + let (mime_type, b64_json) = grok_data_image_parts(image)?; item.insert("result".to_string(), Value::String(b64_json)); item.insert( "output_format".to_string(), @@ -2782,7 +3309,7 @@ fn grok_openai_responses_image_generation_item( Value::String("png".to_string()), ); } - Value::Object(item) + Some(Value::Object(item)) } fn grok_output_format_from_mime_type(mime_type: &str) -> String { @@ -2985,19 +3512,24 @@ fn openai_image_body(collected: &GrokCollected) -> Value { "data": collected .images .iter() - .map(|url| grok_openai_image_item(url.as_str())) + .filter_map(|url| grok_openai_image_item(url.as_str())) .collect::>(), }) } -fn grok_openai_image_item(url: &str) -> Value { - if let Some((mime_type, b64_json)) = grok_data_image_parts(url) { - return json!({ +fn grok_openai_image_item(url: &str) -> Option { + let url = url.trim(); + if url.is_empty() { + return None; + } + if grok_is_data_uri(url) { + let (mime_type, b64_json) = grok_data_image_parts(url)?; + return Some(json!({ "b64_json": b64_json, "mime_type": mime_type, - }); + })); } - json!({ "url": url }) + Some(json!({ "url": url })) } async fn materialize_grok_image_assets(plan: &ExecutionPlan, collected: &mut GrokCollected) { @@ -3050,27 +3582,71 @@ async fn grok_download_image_asset( return Ok(None); } let headers = response.headers(); - let content_type = headers + let declared_content_type = headers .get("content-type") - .map(String::as_str) + .and_then(|value| value.split(';').next()) .map(str::trim) - .filter(|value| value.starts_with("image/")) - .unwrap_or("image/png") - .to_string(); - let bytes = response.bytes().await.map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Grok image asset download failed: {err}" - )) - })?; + .filter(|value| !value.is_empty()); + let bytes = response + .bytes_with_limit(GROK_MAX_ATTACHMENT_BYTES) + .await + .map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Grok image asset download failed: {err}" + )) + })?; if bytes.is_empty() { return Ok(None); } + let content_type = + grok_image_mime_for_payload(&bytes, declared_content_type).ok_or_else(|| { + ExecutionRuntimeTransportError::UpstreamRequest( + "Grok image asset response is not a supported image".to_string(), + ) + })?; Ok(Some(format!( "data:{content_type};base64,{}", base64::engine::general_purpose::STANDARD.encode(bytes) ))) } +fn grok_image_mime_for_payload( + bytes: &[u8], + declared_content_type: Option<&str>, +) -> Option<&'static str> { + let detected = if bytes.starts_with(b"\x89PNG\r\n\x1a\n") { + "image/png" + } else if bytes.starts_with(&[0xff, 0xd8, 0xff]) { + "image/jpeg" + } else if bytes.len() >= 12 && bytes.starts_with(b"RIFF") && &bytes[8..12] == b"WEBP" { + "image/webp" + } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") { + "image/gif" + } else if bytes.len() >= 12 + && &bytes[4..8] == b"ftyp" + && matches!(&bytes[8..12], b"avif" | b"avis") + { + "image/avif" + } else { + return None; + }; + + let declared = declared_content_type + .and_then(|value| value.split(';').next()) + .map(|value| value.trim().to_ascii_lowercase()); + let declared = declared.as_deref().map(|value| { + if value == "image/jpg" { + "image/jpeg" + } else { + value + } + }); + if declared.is_some_and(|value| value != detected) { + return None; + } + Some(detected) +} + fn grok_image_asset_url_is_supported(raw_url: &str) -> bool { if raw_url.starts_with("data:image/") { return true; @@ -3091,28 +3667,98 @@ fn grok_image_asset_url_is_supported(raw_url: &str) -> bool { } fn grok_data_image_parts(raw_url: &str) -> Option<(String, String)> { + grok_data_image_parts_with_limit(raw_url, GROK_MAX_ATTACHMENT_BYTES) +} + +fn grok_data_image_parts_with_limit( + raw_url: &str, + decoded_limit: usize, +) -> Option<(String, String)> { let Some((header, data)) = raw_url.trim().split_once(',') else { return None; }; - if !header.starts_with("data:image/") || !header.contains(";base64") { + let metadata = header + .get(..5) + .filter(|prefix| prefix.eq_ignore_ascii_case("data:")) + .and_then(|_| header.get(5..))?; + let mut parameters = metadata.split(';'); + let declared_mime = parameters.next()?.trim(); + let declared_mime = grok_inline_image_mime(declared_mime)?; + if !parameters.any(|parameter| parameter.trim().eq_ignore_ascii_case("base64")) { return None; } - let mime = header - .strip_prefix("data:") - .and_then(|value| value.split(';').next()) - .map(str::trim) - .filter(|value| value.starts_with("image/"))? - .to_string(); - let normalized = data.split_whitespace().collect::(); + + // Check the encoded length before decoding. The base64 engine allocates + // based on that length, so a post-decode check alone would still permit an + // allocation denial of service. + let max_base64_len = maximum_base64_len_for_decoded_limit(decoded_limit); + let normalized = normalize_base64_with_limit(data, max_base64_len).ok()?; if normalized.is_empty() { return None; } - Some((mime, normalized)) + let decoded = base64::engine::general_purpose::STANDARD + .decode(&normalized) + .ok()?; + if decoded.len() > decoded_limit { + return None; + } + let detected = grok_image_mime_for_payload(&decoded, Some(declared_mime))?; + Some((detected.to_string(), normalized)) +} + +fn grok_inline_image_mime(value: &str) -> Option<&'static str> { + // Keep the declaration bounded and restricted to passive raster formats. + // In particular, never permit image/svg+xml or a generic image/* token to + // cross a client-facing data URL boundary. + if value.is_empty() + || value.len() > 64 + || value + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte.is_ascii_control() || byte == b',') + { + return None; + } + if value.eq_ignore_ascii_case("image/png") { + Some("image/png") + } else if value.eq_ignore_ascii_case("image/jpeg") || value.eq_ignore_ascii_case("image/jpg") { + Some("image/jpeg") + } else if value.eq_ignore_ascii_case("image/webp") { + Some("image/webp") + } else if value.eq_ignore_ascii_case("image/gif") { + Some("image/gif") + } else if value.eq_ignore_ascii_case("image/avif") { + Some("image/avif") + } else { + None + } +} + +fn grok_is_data_uri(value: &str) -> bool { + value + .trim() + .get(..5) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case("data:")) +} + +fn grok_public_image_reference(value: &str) -> Option<&str> { + let value = value.trim(); + if value.is_empty() { + return None; + } + if grok_is_data_uri(value) { + // Validate before exposing any data URL. Keep the original borrowed + // text to avoid an additional large allocation for stream responses. + grok_data_image_parts(value)?; + } + Some(value) } fn openai_image_sse(collected: &GrokCollected) -> String { let mut body = String::new(); for (index, url) in collected.images.iter().enumerate() { + let Some(url) = grok_public_image_reference(url) else { + continue; + }; push_sse_event( &mut body, "image_generation.completed", @@ -3130,6 +3776,9 @@ fn openai_image_sse(collected: &GrokCollected) -> String { fn chat_text_with_images(collected: &GrokCollected) -> String { let mut text = collected.text.clone(); for image in &collected.images { + let Some(image) = grok_public_image_reference(image) else { + continue; + }; if !text.is_empty() { text.push_str("\n\n"); } @@ -3235,18 +3884,23 @@ mod tests { use http::{Method, StatusCode}; use super::{ - encode_grok_error_frame, encode_grok_first_byte_timeout_frame, + chat_text_with_images, encode_grok_error_frame, encode_grok_first_byte_timeout_frame, extract_grok_attachment_inputs, grok_aspect_ratio_from_provider_body, grok_asset_url, - grok_attachment_ip_is_public, grok_attachment_payload_from_data_uri, grok_client_json_body, - grok_client_stream_body, grok_handle_imagine_ws_message, - grok_image_count_from_provider_body, grok_image_prompt_from_provider_body, - grok_imagine_request_message, grok_imagine_reset_message, grok_media_post_url, + grok_attachment_ip_is_public, grok_attachment_payload_from_data_uri, + grok_attachment_payload_from_data_uri_with_limit, grok_auxiliary_http_error, + grok_client_json_body, grok_client_stream_body, grok_data_image_parts_with_limit, + grok_execution_result, grok_handle_imagine_ws_message, + grok_handle_imagine_ws_message_with_slot_limit, grok_image_count_from_provider_body, + grok_image_mime_for_payload, grok_image_prompt_from_provider_body, + grok_imagine_blob_lengths_can_be_retained, grok_imagine_request_message, + grok_imagine_reset_message, grok_media_post_url, grok_plan_uses_structured_image_generation, grok_should_collect_image_stream, grok_should_use_imagine_websocket, grok_success_frame_stream, grok_upload_url, grok_upstream_model_name, grok_usage_estimate, grok_user_id_from_cookie_header, - materialize_grok_image_assets, openai_chat_body, openai_image_body, openai_responses_body, - set_grok_image_edit_config, GrokAttachmentInput, GrokCollected, GrokImagineImage, - GrokStreamAdapter, + materialize_grok_image_assets, maximum_base64_len_for_decoded_limit, openai_chat_body, + openai_image_body, openai_image_sse, openai_responses_body, public_socket_addr_for_url, + set_grok_image_edit_config, validate_grok_attachment_url, GrokAttachmentInput, + GrokCollected, GrokImagineImage, GrokStreamAdapter, }; fn sample_plan(body: serde_json::Value, client_api_format: &str) -> ExecutionPlan { @@ -3454,6 +4108,56 @@ mod tests { ); } + #[test] + fn grok_response_chunk_limit_rejects_wire_overflow_without_growing_buffers() { + let mut upstream_bytes = 4; + let mut raw_body = b"1234".to_vec(); + let mut adapter = GrokStreamAdapter::default(); + + let error = super::collect_grok_response_chunk_with_limit( + 502, + &mut upstream_bytes, + &mut raw_body, + &mut adapter, + b"56", + 5, + ) + .expect_err("wire body above the plan limit must fail"); + + assert_eq!(upstream_bytes, 4); + assert_eq!(raw_body, b"1234"); + assert!(matches!( + error, + super::ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: super::UpstreamResponseBodyPhase::Wire, + limit_bytes: 5, + } + )); + } + + #[test] + fn grok_success_response_does_not_duplicate_raw_body_storage() { + let mut upstream_bytes = 0; + let mut raw_body = Vec::new(); + let mut adapter = GrokStreamAdapter::default(); + let line = + b"data: {\"result\":{\"response\":{\"token\":\"ok\",\"messageTag\":\"final\"}}}\n"; + + super::collect_grok_response_chunk_with_limit( + 200, + &mut upstream_bytes, + &mut raw_body, + &mut adapter, + line, + line.len(), + ) + .expect("body exactly at the limit should pass"); + + assert_eq!(upstream_bytes, line.len() as u64); + assert!(raw_body.is_empty()); + assert_eq!(adapter.text, "ok"); + } + #[test] fn adapter_extracts_grok_image_edit_streaming_response() { let line = format!( @@ -3505,6 +4209,16 @@ mod tests { ); } + #[test] + fn adapter_bounds_collected_image_urls() { + let mut adapter = GrokStreamAdapter::default(); + for index in 0..(super::GROK_MAX_IMAGE_COUNT + 3) { + adapter.push_image_url(format!("https://assets.grok.com/image-{index}.png")); + } + + assert_eq!(adapter.images.len(), super::GROK_MAX_IMAGE_COUNT); + } + #[test] fn grok_image_edit_config_sets_references_and_parent_post() { let mut body = serde_json::json!({ @@ -3594,6 +4308,53 @@ mod tests { ); } + #[test] + fn grok_auxiliary_errors_omit_response_and_websocket_error_bodies() { + let upstream_body = "Bearer secret-grok-error-body"; + let upload_message = grok_auxiliary_http_error("attachment upload", 502).to_string(); + + assert!(upload_message.contains("Grok attachment upload returned HTTP 502")); + assert!(!upload_message.contains(upstream_body)); + + let mut slots = BTreeMap::new(); + let error = grok_handle_imagine_ws_message( + &serde_json::json!({ + "type": "error", + "err_msg": upstream_body, + }), + &mut slots, + ) + .expect_err("websocket error frame should fail"); + let message = error.to_string(); + + assert!(message.contains("Grok Imagine websocket returned an error")); + assert!(!message.contains(upstream_body)); + } + + #[test] + fn grok_http_errors_do_not_copy_upstream_response_bodies() { + let secret = "authorization=Bearer secret-grok-error-body"; + let plan = sample_plan(serde_json::json!({"message": "test"}), "openai:chat"); + let collected = GrokCollected { + status_code: 502, + text: secret.to_string(), + ..GrokCollected::default() + }; + + let stream_body = grok_client_stream_body(&plan, &collected, None); + let result = grok_execution_result(&plan, collected, None); + let sync_body = result + .body + .and_then(|body| body.json_body) + .expect("sync error body should exist") + .to_string(); + + assert!(stream_body.contains("Grok upstream request returned HTTP 502")); + assert!(sync_body.contains("Grok upstream request returned HTTP 502")); + assert!(!stream_body.contains("secret-grok-error-body")); + assert!(!sync_body.contains("secret-grok-error-body")); + } + #[test] fn grok_imagine_ws_parser_collects_completed_image() { let mut slots = BTreeMap::::new(); @@ -3653,6 +4414,14 @@ mod tests { Some("a chair".to_string()) ); assert_eq!(grok_image_count_from_provider_body(&body), 3); + assert_eq!( + grok_image_count_from_provider_body(&serde_json::json!({"n": 0})), + 1 + ); + assert_eq!( + grok_image_count_from_provider_body(&serde_json::json!({"n": u64::MAX})), + 4 + ); assert_eq!(grok_aspect_ratio_from_provider_body(&body), "16:9"); let plan = sample_plan(body, "openai:image"); @@ -3668,6 +4437,49 @@ mod tests { .expect("route should resolve")); } + #[test] + fn grok_imagine_ws_parser_bounds_transient_image_slots() { + let mut slots = BTreeMap::::new(); + for index in 0..8 { + grok_handle_imagine_ws_message_with_slot_limit( + &serde_json::json!({ + "type": "json", + "current_status": "start_stage", + "image_id": format!("image-{index}"), + "order": index, + }), + &mut slots, + 4, + ) + .expect("progress frame should parse"); + } + + assert_eq!(slots.len(), 4); + } + + #[test] + fn grok_imagine_ws_parser_bounds_aggregate_blob_storage() { + let total_limit = + maximum_base64_len_for_decoded_limit(super::GROK_MAX_IMAGINE_BLOB_TOTAL_DECODED_BYTES); + let existing_len = total_limit.saturating_sub(8); + assert!(!grok_imagine_blob_lengths_can_be_retained( + existing_len, + 0, + 16, + )); + assert!(grok_imagine_blob_lengths_can_be_retained( + existing_len, + existing_len, + maximum_base64_len_for_decoded_limit(super::GROK_MAX_ATTACHMENT_BYTES), + )); + assert!(!grok_imagine_blob_lengths_can_be_retained( + 0, + 0, + maximum_base64_len_for_decoded_limit(super::GROK_MAX_ATTACHMENT_BYTES) + .saturating_add(1), + )); + } + #[test] fn grok_attachment_public_ip_guard_rejects_private_ranges() { for ip in [ @@ -3679,6 +4491,12 @@ mod tests { "::1", "fc00::1", "fe80::1", + "100.64.0.1", + "198.18.0.1", + "224.0.0.1", + "64:ff9b::10.0.0.1", + "2002:0a00:0001::1", + "2001:0000:4136:e378:8000:63bf:3fff:fdd2", ] { assert!( !grok_attachment_ip_is_public(ip.parse().expect("ip should parse")), @@ -3693,6 +4511,54 @@ mod tests { )); } + #[tokio::test] + async fn grok_attachment_public_target_rejects_bracketed_private_ipv6_literals() { + for raw_url in [ + "http://[::1]/attachment", + "http://[fc00::1]/attachment", + "http://[fe80::1]/attachment", + ] { + let url = reqwest::Url::parse(raw_url).expect("URL should parse"); + assert!( + public_socket_addr_for_url(&url).await.is_err(), + "private IPv6 literal should be rejected: {raw_url}" + ); + } + + let url = reqwest::Url::parse("https://[2606:4700:4700::1111]/attachment") + .expect("URL should parse"); + assert_eq!( + public_socket_addr_for_url(&url) + .await + .expect("public IPv6 literal should pass"), + "[2606:4700:4700::1111]:443".parse().unwrap() + ); + } + + #[test] + fn grok_attachment_url_rejects_credentials_and_fragments_on_every_hop() { + for raw_url in [ + "https://user@example.com/attachment.png", + "https://:password@example.com/attachment.png", + "https://user:password@example.com/attachment.png", + "https://example.com/attachment.png#private-fragment", + ] { + let url = reqwest::Url::parse(raw_url).expect("URL should parse"); + assert!( + validate_grok_attachment_url(&url).is_err(), + "unsafe attachment URL should be rejected: {raw_url}" + ); + } + + let initial = reqwest::Url::parse("https://example.com/attachment.png?signature=abc") + .expect("URL should parse"); + validate_grok_attachment_url(&initial).expect("signed query URL should remain supported"); + let redirected = initial + .join("https://redirect-user:redirect-pass@example.net/next.png") + .expect("redirect URL should parse"); + assert!(validate_grok_attachment_url(&redirected).is_err()); + } + #[test] fn adapter_cleans_inline_citation_render_tags() { let card_json = serde_json::json!({ @@ -3774,6 +4640,9 @@ mod tests { ); assert_eq!(body["output"][0]["type"], serde_json::json!("reasoning")); assert_eq!(body["output"][1]["type"], serde_json::json!("message")); + assert!(body["output"][1]["id"] + .as_str() + .is_some_and(|id| id.starts_with("msg_"))); } #[test] @@ -3781,7 +4650,7 @@ mod tests { let plan = sample_plan(serde_json::json!({"input": "draw"}), "openai:responses"); let collected = GrokCollected { text: "done".to_string(), - images: vec!["data:image/png;base64,aW1hZ2U=".to_string()], + images: vec!["data:image/png;base64,iVBORw0KGgo=".to_string()], ..GrokCollected::default() }; let usage = grok_usage_estimate(&plan, &collected); @@ -3792,7 +4661,10 @@ mod tests { body["output"][1]["type"], serde_json::json!("image_generation_call") ); - assert_eq!(body["output"][1]["result"], serde_json::json!("aW1hZ2U=")); + assert_eq!( + body["output"][1]["result"], + serde_json::json!("iVBORw0KGgo=") + ); assert_eq!(body["output"][1]["output_format"], serde_json::json!("png")); } @@ -3894,7 +4766,7 @@ mod tests { plan.model_name = Some("grok-imagine-image-lite".to_string()); let collected = GrokCollected { status_code: 200, - images: vec!["data:image/png;base64,aW1hZ2U=".to_string()], + images: vec!["data:image/png;base64,iVBORw0KGgo=".to_string()], ..GrokCollected::default() }; let body = grok_client_json_body(&plan, &collected, None); @@ -3906,7 +4778,7 @@ mod tests { ); assert_eq!( body["choices"][0]["message"]["content"][0]["image_url"]["url"], - serde_json::json!("data:image/png;base64,aW1hZ2U=") + serde_json::json!("data:image/png;base64,iVBORw0KGgo=") ); } @@ -3922,7 +4794,7 @@ mod tests { let report_context = report_context_with_mapped_model("grok-imagine-image-lite"); let collected = GrokCollected { status_code: 200, - images: vec!["data:image/png;base64,aW1hZ2U=".to_string()], + images: vec!["data:image/png;base64,iVBORw0KGgo=".to_string()], ..GrokCollected::default() }; let body = grok_client_json_body(&plan, &collected, Some(&report_context)); @@ -3970,14 +4842,14 @@ mod tests { ); let collected = GrokCollected { status_code: 200, - images: vec!["data:image/png;base64,aGVsbG8=".to_string()], + images: vec!["data:image/png;base64,iVBORw0KGgo=".to_string()], ..GrokCollected::default() }; let body = grok_client_stream_body(&plan, &collected, None); assert!(body.contains("event: response.output_item.done")); assert!(body.contains("\"type\":\"image_generation_call\"")); - assert!(body.contains("\"result\":\"aGVsbG8=\"")); + assert!(body.contains("\"result\":\"iVBORw0KGgo=\"")); assert!(body.contains("event: response.completed")); } @@ -3994,7 +4866,7 @@ mod tests { let report_context = report_context_with_mapped_model("grok-imagine-image-lite"); let collected = GrokCollected { status_code: 200, - images: vec!["data:image/png;base64,aGVsbG8=".to_string()], + images: vec!["data:image/png;base64,iVBORw0KGgo=".to_string()], ..GrokCollected::default() }; let body = grok_client_stream_body(&plan, &collected, Some(&report_context)); @@ -4005,7 +4877,7 @@ mod tests { )); assert!(body.contains("event: response.output_item.done")); assert!(body.contains("\"type\":\"image_generation_call\"")); - assert!(body.contains("\"result\":\"aGVsbG8=\"")); + assert!(body.contains("\"result\":\"iVBORw0KGgo=\"")); } #[test] @@ -4084,7 +4956,8 @@ mod tests { ] }] }), - ); + ) + .expect("attachment inputs should be within limits"); assert_eq!(inputs.len(), 2); assert_eq!(inputs[0].source.as_str(), "data:image/png;base64,aGVsbG8="); @@ -4093,9 +4966,88 @@ mod tests { } #[test] - fn grok_data_uri_attachment_accepts_content_above_previous_size_cap() { - const PREVIOUS_CAP_BYTES: usize = 25 * 1024 * 1024; - let base64_blocks = PREVIOUS_CAP_BYTES / 3 + 1; + fn grok_attachment_inputs_reject_excess_count_before_uploads() { + let content = (0..=super::GROK_MAX_ATTACHMENT_COUNT) + .map(|index| { + serde_json::json!({ + "type": "image_url", + "image_url": format!("https://example.com/image-{index}.png") + }) + }) + .collect::>(); + let error = extract_grok_attachment_inputs( + "openai:chat", + &serde_json::json!({ + "messages": [{"role": "user", "content": content}] + }), + ) + .expect_err("more than the bounded attachment count must be rejected"); + + assert!(error.to_string().contains("at most 16 attachments")); + } + + #[test] + fn grok_attachment_inputs_bound_source_and_metadata_fields() { + let long_url = format!( + "https://example.com/{}", + "u".repeat(super::GROK_MAX_ATTACHMENT_URL_BYTES) + ); + let source_error = extract_grok_attachment_inputs( + "openai:chat", + &serde_json::json!({ + "messages": [{ + "role": "user", + "content": [{"type": "image_url", "image_url": long_url}] + }] + }), + ) + .expect_err("an oversized URL field must be rejected before DNS/fetch"); + assert!(source_error.to_string().contains("source exceeds")); + + let long_filename = "f".repeat(super::GROK_MAX_ATTACHMENT_FILENAME_BYTES + 1); + let filename_error = extract_grok_attachment_inputs( + "openai:chat", + &serde_json::json!({ + "messages": [{ + "role": "user", + "content": [{ + "type": "file", + "file": { + "filename": long_filename, + "file_data": "data:text/plain;base64,Zm9v" + } + }] + }] + }), + ) + .expect_err("an oversized filename field must be rejected"); + assert!(filename_error.to_string().contains("filename exceeds")); + + let long_mime = "a".repeat(super::GROK_MAX_ATTACHMENT_MIME_TYPE_BYTES + 1); + let mime_error = extract_grok_attachment_inputs( + "claude:messages", + &serde_json::json!({ + "messages": [{ + "role": "user", + "content": [{ + "type": "document", + "source": { + "type": "base64", + "media_type": long_mime, + "data": "Zm9v" + } + }] + }] + }), + ) + .expect_err("an oversized MIME field must be rejected"); + assert!(mime_error.to_string().contains("MIME type exceeds")); + } + + #[test] + fn grok_data_uri_attachment_accepts_content_within_hard_size_cap() { + const PAYLOAD_BYTES: usize = 25 * 1024 * 1024 + 1; + let base64_blocks = PAYLOAD_BYTES / 3 + 1; let mut source = String::from("data:application/octet-stream;base64,"); source.extend(std::iter::repeat_n('A', base64_blocks * 4)); let input = GrokAttachmentInput { @@ -4105,13 +5057,88 @@ mod tests { }; let payload = grok_attachment_payload_from_data_uri(&input, 0) - .expect("attachment above the previous size cap should be accepted"); + .expect("attachment within the hard cap should be accepted"); assert_eq!(payload.filename, "large.bin"); assert_eq!(payload.mime_type, "application/octet-stream"); assert_eq!(payload.content_b64.len(), base64_blocks * 4); } + #[test] + fn grok_data_uri_attachment_rejects_content_above_hard_size_cap() { + let source = "data:application/octet-stream;base64,MTIzNDU2Nzg5".to_string(); + let input = GrokAttachmentInput { + source, + filename: Some("oversized.bin".to_string()), + mime_type: None, + }; + + let error = grok_attachment_payload_from_data_uri_with_limit(&input, 0, 8) + .expect_err("attachment above the hard cap must be rejected"); + + assert!(error.to_string().contains("exceeds")); + } + + #[test] + fn grok_public_data_image_requires_bounded_base64_and_matching_raster_magic() { + let png = [ + 0x89, b'P', b'N', b'G', 0x0d, 0x0a, 0x1a, 0x0a, 0, 0, 0, 13, b'I', b'H', b'D', b'R', + ]; + let encoded_png = base64::engine::general_purpose::STANDARD.encode(png); + let valid = format!("data:image/png;base64,{encoded_png}"); + let (mime, normalized) = grok_data_image_parts_with_limit(&valid, png.len()) + .expect("valid PNG data URL should pass"); + assert_eq!(mime, "image/png"); + assert_eq!(normalized, encoded_png); + + let html = base64::engine::general_purpose::STANDARD.encode(b"not an image"); + assert!(grok_data_image_parts_with_limit( + format!("data:image/png;base64,{html}").as_str(), + 64, + ) + .is_none()); + assert!(grok_data_image_parts_with_limit( + format!("data:image/svg+xml;base64,{html}").as_str(), + 64, + ) + .is_none()); + assert!(grok_data_image_parts_with_limit( + format!("data:image/jpeg;base64,{encoded_png}").as_str(), + png.len(), + ) + .is_none()); + assert!( + grok_data_image_parts_with_limit("data:image/png;base64,not-valid-***", 64,).is_none() + ); + assert!(grok_data_image_parts_with_limit(&valid, png.len() - 1).is_none()); + } + + #[test] + fn grok_public_image_outputs_drop_invalid_data_urls_but_keep_http_urls() { + let svg = + base64::engine::general_purpose::STANDARD.encode(b""); + let invalid = format!("data:image/svg+xml;base64,{svg}"); + let ordinary_url = "https://assets.grok.com/generated/example.png"; + let collected = GrokCollected { + status_code: 200, + text: "done".to_string(), + images: vec![invalid.clone(), ordinary_url.to_string()], + ..GrokCollected::default() + }; + + let body = openai_image_body(&collected); + assert_eq!(body["data"].as_array().map(Vec::len), Some(1)); + assert_eq!(body["data"][0]["url"], serde_json::json!(ordinary_url)); + + let text = chat_text_with_images(&collected); + assert!(text.contains(ordinary_url)); + assert!(!text.contains(&invalid)); + + let sse = openai_image_sse(&collected); + assert!(sse.contains(ordinary_url)); + assert!(!sse.contains(&invalid)); + } + #[test] fn extracts_responses_and_claude_attachment_inputs() { let responses = extract_grok_attachment_inputs( @@ -4127,7 +5154,8 @@ mod tests { ] }] }), - ); + ) + .expect("responses attachments should be within limits"); let claude = extract_grok_attachment_inputs( "claude:messages", &serde_json::json!({ @@ -4140,7 +5168,8 @@ mod tests { ] }] }), - ); + ) + .expect("Claude attachments should be within limits"); assert_eq!(responses.len(), 2); assert_eq!(responses[0].source.as_str(), "https://example.com/a.png"); @@ -4179,7 +5208,7 @@ mod tests { (Method::GET, "/generated.png") => response( StatusCode::OK, "image/png", - Body::from(vec![0x89, b'P', b'N', b'G']), + Body::from(vec![0x89, b'P', b'N', b'G', 0x0d, 0x0a, 0x1a, 0x0a]), ), _ => response(StatusCode::NOT_FOUND, "text/plain", "not found"), } @@ -4210,6 +5239,21 @@ mod tests { assert!(body["data"][0].get("url").is_none()); } + #[test] + fn grok_downloaded_image_requires_matching_magic_and_mime() { + let png = [0x89, b'P', b'N', b'G', 0x0d, 0x0a, 0x1a, 0x0a]; + assert_eq!( + grok_image_mime_for_payload(&png, Some("image/png; charset=binary")), + Some("image/png") + ); + assert_eq!(grok_image_mime_for_payload(&png, Some("text/html")), None); + assert_eq!( + grok_image_mime_for_payload(b"not an image", Some("image/png")), + None + ); + assert_eq!(grok_image_mime_for_payload(&png, None), Some("image/png")); + } + #[test] fn grok_openai_image_body_preserves_plain_urls_when_asset_is_not_materialized() { let body = openai_image_body(&GrokCollected { diff --git a/apps/aether-gateway/src/execution_runtime/kiro_web_search.rs b/apps/aether-gateway/src/execution_runtime/kiro_web_search.rs index 1afbf8ab0..6987788e6 100644 --- a/apps/aether-gateway/src/execution_runtime/kiro_web_search.rs +++ b/apps/aether-gateway/src/execution_runtime/kiro_web_search.rs @@ -21,11 +21,13 @@ use crate::execution_runtime::kiro_cache::{ }; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::transport::{ - DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, + apply_upstream_response_body_limit, decode_base64_body_with_limit, + json_value_fits_serialized_limit, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, }; use crate::AppState; const WEB_SEARCH_TOOL_NAME: &str = "web_search"; +const MAX_KIRO_MCP_JSON_BYTES: usize = 8 * 1024 * 1024; const WEB_SEARCH_TOOL_TYPE_PREFIX: &str = "web_search"; const WEB_SEARCH_QUERY_PREFIX: &str = "Perform a web search for the query: "; @@ -73,8 +75,6 @@ struct McpResponse { struct McpError { #[serde(default)] code: Option, - #[serde(default)] - message: Option, } #[derive(Debug, Deserialize)] @@ -146,14 +146,14 @@ pub(crate) async fn maybe_execute_kiro_web_search_stream( request_id = %plan.request_id, candidate_id = ?plan.candidate_id, status_code = mcp_execution.result.status_code, - mcp_url = %mcp_execution.url, + mcp_origin = %crate::handlers::shared::security_log_url_origin(&mcp_execution.url), profile_arn_present = mcp_execution.profile_arn_present, "gateway executed Kiro web_search through MCP endpoint" ); if !(200..300).contains(&mcp_execution.result.status_code) { return Ok(Some(KiroWebSearchStream { - frame_stream: execution_result_frame_stream(&mcp_execution.result), + frame_stream: execution_result_failure_frame_stream(&mcp_execution.result), report_context: report_context.cloned(), })); } @@ -183,7 +183,7 @@ pub(crate) async fn maybe_execute_kiro_web_search_stream( cache_usage, ) .map_err(ExecutionRuntimeTransportError::BodyEncode)?; - let mut synthetic_context = synthetic_report_context(report_context, mcp_execution.url); + let mut synthetic_context = synthetic_report_context(report_context); if let Some(context) = synthetic_context.as_mut().and_then(Value::as_object_mut) { context.insert("kiro_web_search_mcp".to_string(), Value::Bool(true)); } @@ -214,13 +214,13 @@ async fn kiro_simulated_cache_enabled(state: &AppState, plan: &ExecutionPlan) -> .is_some_and(|provider| { kiro_simulated_cache_enabled_from_provider_config(provider.config.as_ref()) }), - Err(err) => { + Err(_) => { warn!( event_name = "kiro_simulated_cache_config_read_failed", log_type = "event", request_id = %plan.request_id, provider_id = %plan.provider_id, - error = ?err, + error_category = "provider_catalog_read_failed", "failed to read Kiro simulated cache provider config; defaulting disabled" ); false @@ -228,30 +228,28 @@ async fn kiro_simulated_cache_enabled(state: &AppState, plan: &ExecutionPlan) -> } } -fn execute_result_body_bytes(result: &ExecutionResult) -> Vec { - let Some(body) = result.body.as_ref() else { - return Vec::new(); - }; - if let Some(json_body) = body.json_body.as_ref() { - return serde_json::to_vec(json_body).unwrap_or_default(); - } - body.body_bytes_b64 - .as_deref() - .and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok()) - .unwrap_or_default() -} - -fn execution_result_frame_stream( +fn execution_result_failure_frame_stream( result: &ExecutionResult, ) -> BoxStream<'static, Result> { + let body = serde_json::to_vec(&kiro_mcp_failure_body(result.status_code)).unwrap_or_default(); raw_response_frame_stream( result.status_code, - result.headers.clone(), - Bytes::from(execute_result_body_bytes(result)), + BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), + Bytes::from(body), result.telemetry.clone(), ) } +fn kiro_mcp_failure_body(status_code: u16) -> Value { + json!({ + "error": { + "type": "kiro_web_search_error", + "message": "Kiro web search request failed", + "code": status_code, + } + }) +} + fn sse_frame_stream(body: Bytes) -> BoxStream<'static, Result> { raw_response_frame_stream( 200, @@ -328,12 +326,7 @@ async fn execute_mcp_request( request: &McpRequest, ) -> Result { let mcp_url = aether_provider_transport::kiro::build_kiro_mcp_url_from_resolved_url(&plan.url) - .ok_or_else(|| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to build Kiro MCP url from {}", - plan.url - )) - })?; + .ok_or_else(kiro_mcp_url_build_error)?; let mut request_context = build_mcp_request_context(state, plan).await; if !request_context.profile_arn_present { if let Some(profile_arn) = discover_kiro_profile_arn(state, plan, &request_context).await? { @@ -346,7 +339,7 @@ async fn execute_mcp_request( } let body_json = serde_json::to_value(request).map_err(ExecutionRuntimeTransportError::BodyEncode)?; - let mcp_plan = ExecutionPlan { + let mut mcp_plan = ExecutionPlan { request_id: plan.request_id.clone(), candidate_id: plan.candidate_id.clone(), provider_name: plan.provider_name.clone(), @@ -367,6 +360,7 @@ async fn execute_mcp_request( transport_profile: plan.transport_profile.clone(), timeouts: plan.timeouts.clone(), }; + apply_upstream_response_body_limit(&mut mcp_plan, MAX_KIRO_MCP_JSON_BYTES); let result = DirectSyncExecutionRuntime::new() .execute_sync(&mcp_plan) .await?; @@ -390,7 +384,7 @@ async fn build_mcp_request_context( { Ok(Some(transport)) => transport, Ok(None) => return fallback(), - Err(err) => { + Err(_) => { warn!( event_name = "kiro_web_search_transport_snapshot_unavailable", log_type = "ops", @@ -399,7 +393,7 @@ async fn build_mcp_request_context( provider_id = %plan.provider_id, endpoint_id = %plan.endpoint_id, key_id = %plan.key_id, - error = ?err, + error_category = "transport_snapshot_read_failed", "gateway could not read Kiro transport snapshot for web_search MCP" ); return fallback(); @@ -574,7 +568,7 @@ async fn discover_kiro_profile_arn_in_region( if let Some(token) = next_token.as_deref() { body.insert("nextToken".to_string(), Value::String(token.to_string())); } - let list_plan = ExecutionPlan { + let mut list_plan = ExecutionPlan { request_id: plan.request_id.clone(), candidate_id: plan.candidate_id.clone(), provider_name: plan.provider_name.clone(), @@ -598,6 +592,7 @@ async fn discover_kiro_profile_arn_in_region( transport_profile: plan.transport_profile.clone(), timeouts: plan.timeouts.clone(), }; + apply_upstream_response_body_limit(&mut list_plan, MAX_KIRO_MCP_JSON_BYTES); let result = DirectSyncExecutionRuntime::new() .execute_sync(&list_plan) .await?; @@ -676,6 +671,10 @@ fn kiro_runtime_base_url_for_region(region: &str) -> String { } fn kiro_runtime_host_for_region(region: &str) -> String { + // Region values can originate in persisted OAuth metadata. Normalize + // before interpolation so a crafted value cannot turn the host header or + // URL into an attacker-controlled origin. + let region = aether_provider_transport::kiro::normalize_kiro_region(region); match region { "us-gov-east-1" | "us-gov-west-1" => format!("q-fips.{region}.amazonaws.com"), "us-iso-east-1" => "q.us-iso-east-1.c2s.ic.gov".to_string(), @@ -959,7 +958,6 @@ fn parse_mcp_search_results(result: &ExecutionResult) -> Option Option Option { let body = result.body.as_ref()?; if let Some(json_body) = body.json_body.as_ref() { - return Some(json_body.clone()); + return json_value_fits_serialized_limit(json_body, MAX_KIRO_MCP_JSON_BYTES) + .then(|| json_body.clone()); } let body = body.body_bytes_b64.as_deref()?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(body) - .ok()?; + // MCP responses are parsed into an owned JSON tree. Keep a fixed ceiling + // here even if the operator disables the general internal body cap, so a + // forged execution result cannot trigger an unbounded base64 allocation. + let bytes = decode_base64_body_with_limit( + body, + crate::headers::max_internal_buffered_body_bytes().min(MAX_KIRO_MCP_JSON_BYTES), + ) + .ok()?; serde_json::from_slice(&bytes).ok() } @@ -1198,13 +1202,14 @@ fn generate_search_summary(query: &str, results: Option<&WebSearchResults>) -> S } summary.push_str(&format!(" Source: {}\n\n", result.url)); } - if let Some(error) = results + if results .error .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) + .is_some() { - summary.push_str(&format!("Search warning: {error}\n\n")); + summary.push_str("Search warning: the provider reported an incomplete result.\n\n"); } } _ => summary.push_str("No results found.\n"), @@ -1237,12 +1242,15 @@ fn estimate_text_tokens(text: &str) -> u64 { ((text.len() as u64 + 3) / 4).max(1) } -fn synthetic_report_context(report_context: Option<&Value>, mcp_url: String) -> Option { +fn kiro_mcp_url_build_error() -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest("failed to build Kiro MCP URL".to_string()) +} + +fn synthetic_report_context(report_context: Option<&Value>) -> Option { let mut context = report_context.cloned()?; if let Some(object) = context.as_object_mut() { object.insert("has_envelope".to_string(), Value::Bool(false)); object.insert("needs_conversion".to_string(), Value::Bool(false)); - object.insert("upstream_url".to_string(), Value::String(mcp_url)); object.remove("envelope_name"); } Some(context) @@ -1257,7 +1265,9 @@ mod tests { use super::{ build_mcp_headers_from_plan, build_web_search_sse_body, detect_kiro_web_search_request, - parse_mcp_search_results, strip_search_query_prefix, KiroPromptCacheUsage, + execution_result_body_json, generate_search_summary, kiro_mcp_failure_body, + kiro_mcp_url_build_error, parse_mcp_search_results, strip_search_query_prefix, + KiroPromptCacheUsage, WebSearchResult, WebSearchResults, }; fn sample_plan(body: serde_json::Value) -> ExecutionPlan { @@ -1443,6 +1453,43 @@ mod tests { ); } + #[test] + fn kiro_mcp_url_build_error_omits_resolved_url() { + let sensitive_url = "https://token:secret-kiro-url@internal.example.invalid/path"; + let message = kiro_mcp_url_build_error().to_string(); + + assert_eq!( + message, + "failed to execute upstream request: failed to build Kiro MCP URL" + ); + assert!(!message.contains(sensitive_url)); + } + + #[test] + fn kiro_mcp_errors_and_search_warnings_do_not_copy_upstream_secrets() { + let secret = "authorization=Bearer upstream-secret"; + let failure = kiro_mcp_failure_body(502).to_string(); + let summary = generate_search_summary( + "test", + Some(&WebSearchResults { + results: vec![WebSearchResult { + title: "Safe result".to_string(), + url: "https://example.test/result".to_string(), + snippet: None, + published_date: None, + }], + total_results: Some(1), + query: None, + error: Some(secret.to_string()), + }), + ); + + assert!(failure.contains("Kiro web search request failed")); + assert!(summary.contains("provider reported an incomplete result")); + assert!(!failure.contains("upstream-secret")); + assert!(!summary.contains("upstream-secret")); + } + #[test] fn parses_mcp_search_result_text_payload() { let result = aether_contracts::ExecutionResult { @@ -1474,6 +1521,29 @@ mod tests { assert_eq!(parsed.results[0].title, "Example"); } + #[test] + fn kiro_mcp_result_rejects_oversized_base64_before_decode() { + let encoded_limit = + crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + super::MAX_KIRO_MCP_JSON_BYTES, + ); + let result = aether_contracts::ExecutionResult { + request_id: "req-oversized".to_string(), + candidate_id: None, + status_code: 200, + headers: BTreeMap::new(), + response_observation: None, + body: Some(aether_contracts::ResponseBody { + json_body: None, + body_bytes_b64: Some("A".repeat(encoded_limit + 1)), + }), + telemetry: None, + error: None, + }; + + assert!(execution_result_body_json(&result).is_none()); + } + #[test] fn builds_anthropic_web_search_sse() { let sse = build_web_search_sse_body( diff --git a/apps/aether-gateway/src/execution_runtime/mod.rs b/apps/aether-gateway/src/execution_runtime/mod.rs index 89fbf1653..b9e48c3a8 100644 --- a/apps/aether-gateway/src/execution_runtime/mod.rs +++ b/apps/aether-gateway/src/execution_runtime/mod.rs @@ -4,6 +4,7 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; pub(crate) mod admission; +pub(crate) mod attempt_cancellation; pub(crate) mod attempt_lifecycle; mod chatgpt_web_image; mod constants; @@ -30,7 +31,8 @@ pub(crate) use self::admission::{ }; pub(crate) use self::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync; pub(crate) use self::constants::{ - MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES, + MAX_ERROR_BODY_BYTES, MAX_EXECUTION_STREAM_FRAME_LINE_BYTES, MAX_STREAM_BODY_CAPTURE_BYTES, + MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES, }; pub(crate) use self::fallback::{ analyze_local_candidate_failover_sync, local_failover_response_text, diff --git a/apps/aether-gateway/src/execution_runtime/oauth_retry.rs b/apps/aether-gateway/src/execution_runtime/oauth_retry.rs index 293193ebd..86c1c4533 100644 --- a/apps/aether-gateway/src/execution_runtime/oauth_retry.rs +++ b/apps/aether-gateway/src/execution_runtime/oauth_retry.rs @@ -190,7 +190,12 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry( endpoint_id = %plan.endpoint_id, key_id = %plan.key_id, status_code, - error = %err, + // `LocalOAuthRefreshError` carries provider-generated details + // in a few variants. Its `Debug` implementation deliberately + // redacts those details; using `%` here would invoke + // `Display` and could expose a raw transport URL/query or an + // upstream error body in the ops log. + error = ?err, "gateway oauth retry refresh failed" ); false @@ -336,6 +341,23 @@ mod tests { assert!(!status_may_be_oauth_invalid(429, Some("token bucket"))); } + #[test] + fn oauth_retry_failure_log_uses_redacted_error_debug() { + let error = crate::provider_transport::LocalOAuthRefreshError::TransportMessage { + provider_type: "kiro", + message: "request failed for https://user:password@example.test/token?refresh_token=oauth-retry-canary" + .to_string(), + }; + + // The retry path logs this error with `?error` (Debug), whose contract + // is to omit dynamic transport details. Keep the assertion here so a + // future change back to `%error` cannot silently reintroduce leakage. + let debug = format!("{error:?}"); + assert!(!debug.contains("oauth-retry-canary")); + assert!(!debug.contains("password")); + assert!(debug.contains("TransportMessage")); + } + #[test] fn separates_retry_candidate_from_access_token_invalid_proof() { assert!(status_proves_access_token_invalid(401, None)); diff --git a/apps/aether-gateway/src/execution_runtime/remote_compat.rs b/apps/aether-gateway/src/execution_runtime/remote_compat.rs index 30156b3ff..1ec63b903 100644 --- a/apps/aether-gateway/src/execution_runtime/remote_compat.rs +++ b/apps/aether-gateway/src/execution_runtime/remote_compat.rs @@ -3,21 +3,56 @@ use aether_contracts::{ExecutionPlan, ExecutionResult}; use crate::constants::TRACE_ID_HEADER; use crate::{AppState, GatewayError}; +fn remote_runtime_request_error_kind(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else if error.is_body() { + "body" + } else if error.is_decode() { + "decode" + } else { + "transport" + } +} + fn build_remote_execution_runtime_request( state: &AppState, remote_execution_runtime_base_url: &str, path: &str, trace_id: Option<&str>, plan: &ExecutionPlan, -) -> reqwest::RequestBuilder { +) -> Result { + let envelope_limit = crate::execution_runtime::transport::execution_result_envelope_limit_bytes( + crate::headers::max_internal_buffered_body_bytes(), + ); + let body = crate::execution_runtime::transport::serialize_serializable_with_limit( + plan, + envelope_limit, + ) + .map_err(|error| { + let kind = match error { + crate::execution_runtime::transport::ExecutionRuntimeTransportError::BodyTooLarge { + .. + } => "too_large", + _ => "encode", + }; + GatewayError::Internal(format!( + "remote execution runtime request body failed ({kind})" + )) + })?; let mut request = state .client .post(format!("{remote_execution_runtime_base_url}{path}")) - .json(plan); + .header(reqwest::header::CONTENT_TYPE, "application/json") + .body(body); if let Some(trace_id) = trace_id.map(str::trim).filter(|value| !value.is_empty()) { request = request.header(TRACE_ID_HEADER, trace_id); } - request + Ok(request) } pub(crate) async fn post_sync_plan_to_remote_execution_runtime( @@ -32,10 +67,15 @@ pub(crate) async fn post_sync_plan_to_remote_execution_runtime( "/v1/execute/sync", trace_id, plan, - ) + )? .send() .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| { + GatewayError::Internal(format!( + "remote execution runtime request failed ({})", + remote_runtime_request_error_kind(&err) + )) + }) } pub(crate) async fn post_stream_plan_to_remote_execution_runtime( @@ -50,10 +90,15 @@ pub(crate) async fn post_stream_plan_to_remote_execution_runtime( "/v1/execute/stream", trace_id, plan, - ) + )? .send() .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| { + GatewayError::Internal(format!( + "remote execution runtime request failed ({})", + remote_runtime_request_error_kind(&err) + )) + }) } pub(crate) async fn execute_sync_plan_via_remote_execution_runtime( @@ -76,8 +121,27 @@ pub(crate) async fn execute_sync_plan_via_remote_execution_runtime( ))); } - response - .json::() - .await - .map_err(|err| GatewayError::Internal(err.to_string())) + let body = aether_http::read_response_bytes_with_limit( + response, + crate::execution_runtime::transport::execution_result_envelope_limit_bytes( + crate::headers::max_internal_buffered_body_bytes(), + ), + ) + .await + .map_err(|err| { + GatewayError::Internal(format!( + "remote execution runtime response body failed ({})", + match err { + aether_http::ResponseBodyReadError::TooLarge { .. } => "too_large", + aether_http::ResponseBodyReadError::Read(error) => { + remote_runtime_request_error_kind(&error) + } + } + )) + })?; + serde_json::from_slice::(&body).map_err(|_| { + GatewayError::Internal( + "remote execution runtime returned invalid execution JSON".to_string(), + ) + }) } diff --git a/apps/aether-gateway/src/execution_runtime/server.rs b/apps/aether-gateway/src/execution_runtime/server.rs index 8308e58ea..a2852f839 100644 --- a/apps/aether-gateway/src/execution_runtime/server.rs +++ b/apps/aether-gateway/src/execution_runtime/server.rs @@ -1,5 +1,11 @@ -use std::path::Path; +use std::io; +use std::net::SocketAddr; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; +#[cfg(unix)] +use std::sync::{Mutex, OnceLock}; +use std::time::Duration; use aether_contracts::ExecutionPlan; use aether_runtime::{ @@ -13,8 +19,13 @@ use axum::http::StatusCode; use axum::response::{IntoResponse, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; +use hyper::body::Incoming; +use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; +use hyper_util::server::conn::auto::Builder as HyperServerBuilder; +use hyper_util::service::TowerToHyperService; use serde_json::json; use thiserror::Error; +use tower::{Service as _, ServiceExt as _}; use crate::execution_runtime::{ build_direct_execution_frame_stream, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, @@ -25,6 +36,135 @@ const EXECUTION_RUNTIME_COMPONENT: &str = "aether-gateway-execution-runtime"; const REQUEST_GATE_NAME: &str = "execution_runtime_requests"; const DISTRIBUTED_REQUEST_GATE_NAME: &str = "execution_runtime_requests_distributed"; +// These limits protect only connection metadata. Once the first complete +// request reaches the service, request and response bodies remain fully +// streaming (including long-lived streaming responses). Keep the HTTP/2 +// stream default high enough for high-concurrency hosts; this is not a body +// admission limit. +const EXECUTION_RUNTIME_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 16_384; +const EXECUTION_RUNTIME_HTTP_HEADER_READ_TIMEOUT: Duration = Duration::from_secs(30); +const EXECUTION_RUNTIME_HTTP_HEADER_MAX_BYTES: usize = 64 * 1024; +const EXECUTION_RUNTIME_HTTP_MAX_HEADERS: usize = 256; +const EXECUTION_RUNTIME_REQUEST_BODY_HARD_LIMIT_BYTES: usize = 256 * 1024 * 1024; + +/// Coordinates the connection-level deadline that covers protocol detection +/// and the first request header block. Hyper's HTTP/1 timer starts only after +/// the auto protocol detector has finished, while HTTP/2 has no equivalent +/// header timer. Keeping this gate outside the parser closes that initial gap +/// without imposing a deadline on request or response bodies. +#[derive(Clone)] +struct ExecutionRuntimeFirstRequestGate { + seen: Arc, + notify: Arc, +} + +impl ExecutionRuntimeFirstRequestGate { + fn new() -> Self { + Self { + seen: Arc::new(AtomicBool::new(false)), + notify: Arc::new(tokio::sync::Notify::new()), + } + } + + fn mark_seen(&self) { + if !self.seen.swap(true, Ordering::Release) { + self.notify.notify_one(); + } + } + + fn is_seen(&self) -> bool { + self.seen.load(Ordering::Acquire) + } +} + +#[derive(Clone)] +struct ExecutionRuntimeFirstRequestService { + inner: S, + gate: ExecutionRuntimeFirstRequestGate, +} + +impl tower::Service for ExecutionRuntimeFirstRequestService +where + S: tower::Service, +{ + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, request: Req) -> Self::Future { + self.gate.mark_seen(); + self.inner.call(request) + } +} + +/// Drive one Hyper connection while enforcing the deadline for its first +/// request. The timeout is deliberately limited to protocol detection and +/// initial headers; after `gate.mark_seen()` body and response streaming are +/// not interrupted by this helper. +async fn drive_execution_runtime_connection( + connection: F, + gate: ExecutionRuntimeFirstRequestGate, + timeout: Duration, +) -> Result<(), E> +where + F: std::future::Future>, +{ + if gate.is_seen() { + return connection.await; + } + + let mut connection = Box::pin(connection); + let timeout = tokio::time::sleep(timeout); + tokio::pin!(timeout); + let notified = gate.notify.notified(); + tokio::pin!(notified); + + tokio::select! { + result = &mut connection => result, + _ = &mut timeout => { + if gate.is_seen() { + (&mut connection).await + } else { + tracing::debug!( + "execution runtime connection closed before the first request header completed" + ); + Ok(()) + } + } + _ = &mut notified => (&mut connection).await, + } +} + +fn execution_runtime_http_builder() -> HyperServerBuilder { + let mut builder = HyperServerBuilder::new(TokioExecutor::new()); + + // Hyper's HTTP/1 header timer is opt-in when using the custom connection + // builder. Configure both protocol parsers explicitly: HTTP/1 gets a + // slow-header deadline and bounded parser buffer; HTTP/2 gets a + // decompressed header-list limit. These settings apply to metadata only. + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(EXECUTION_RUNTIME_HTTP_HEADER_READ_TIMEOUT) + .max_buf_size(EXECUTION_RUNTIME_HTTP_HEADER_MAX_BYTES) + .max_headers(EXECUTION_RUNTIME_HTTP_MAX_HEADERS); + builder + .http2() + .timer(TokioTimer::new()) + .enable_connect_protocol() + .max_concurrent_streams(EXECUTION_RUNTIME_HTTP2_MAX_CONCURRENT_STREAMS) + .max_header_list_size(EXECUTION_RUNTIME_HTTP_HEADER_MAX_BYTES as u32); + + builder +} + #[derive(Debug, Clone, Default)] struct ExecutionRuntimeAppState { execution_runtime: DirectSyncExecutionRuntime, @@ -143,16 +283,60 @@ pub async fn serve_execution_runtime_tcp( max_in_flight_requests: Option, distributed_request_gate: Option, ) -> Result<(), Box> { - let listener = tokio::net::TcpListener::bind(bind).await?; - axum::serve( - listener, - build_execution_runtime_router_with_request_gates( - max_in_flight_requests, - distributed_request_gate, - ), - ) - .await?; - Ok(()) + // The execution runtime accepts plans containing upstream credentials and + // can issue arbitrary provider requests. It has no network + // authentication layer, so a TCP listener must remain local-only. + let bind_addr = validate_execution_runtime_tcp_bind(bind)?; + let listener = tokio::net::TcpListener::bind(bind_addr).await?; + let router = build_execution_runtime_router_with_request_gates( + max_in_flight_requests, + distributed_request_gate, + ); + let mut make_service = router.into_make_service(); + + loop { + let (io, _remote_addr) = listener.accept().await?; + let tower_service = make_service + .call(()) + .await + .unwrap_or_else(|error| match error {}) + .map_request(|request: http::Request| request.map(Body::new)); + let first_request_gate = ExecutionRuntimeFirstRequestGate::new(); + let hyper_service = TowerToHyperService::new(ExecutionRuntimeFirstRequestService { + inner: tower_service, + gate: first_request_gate.clone(), + }); + let io = TokioIo::new(io); + + tokio::spawn(async move { + let builder = execution_runtime_http_builder(); + let result = drive_execution_runtime_connection( + builder.serve_connection_with_upgrades(io, hyper_service), + first_request_gate, + EXECUTION_RUNTIME_HTTP_HEADER_READ_TIMEOUT, + ) + .await; + if let Err(error) = result { + tracing::trace!(error = ?error, "execution runtime TCP connection closed with error"); + } + }); + } +} + +fn validate_execution_runtime_tcp_bind(bind: &str) -> io::Result { + let address = bind.parse::().map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "execution runtime TCP bind must be a literal loopback socket address", + ) + })?; + if !address.ip().is_loopback() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime TCP bind must use a loopback address", + )); + } + Ok(address) } #[cfg(unix)] @@ -161,25 +345,323 @@ pub async fn serve_execution_runtime_unix( max_in_flight_requests: Option, distributed_request_gate: Option, ) -> Result<(), Box> { - if let Some(parent) = socket_path.parent() { - std::fs::create_dir_all(parent)?; + let listener = bind_secure_execution_runtime_socket(socket_path).await?; + let router = build_execution_runtime_router_with_request_gates( + max_in_flight_requests, + distributed_request_gate, + ); + let mut make_service = router.into_make_service(); + + loop { + let (io, _peer_addr) = listener.accept().await?; + let tower_service = make_service + .call(()) + .await + .unwrap_or_else(|error| match error {}) + .map_request(|request: http::Request| request.map(Body::new)); + let first_request_gate = ExecutionRuntimeFirstRequestGate::new(); + let hyper_service = TowerToHyperService::new(ExecutionRuntimeFirstRequestService { + inner: tower_service, + gate: first_request_gate.clone(), + }); + let io = TokioIo::new(io); + + tokio::spawn(async move { + let builder = execution_runtime_http_builder(); + let result = drive_execution_runtime_connection( + builder.serve_connection_with_upgrades(io, hyper_service), + first_request_gate, + EXECUTION_RUNTIME_HTTP_HEADER_READ_TIMEOUT, + ) + .await; + if let Err(error) = result { + tracing::trace!(error = ?error, "execution runtime Unix connection closed with error"); + } + }); } - if socket_path.exists() { - std::fs::remove_file(socket_path)?; +} + +/// Bind the execution-runtime UDS without exposing an unauthenticated socket +/// to other local users. In particular, do not unlink an arbitrary path before +/// binding: a symlink/regular file could otherwise be replaced during that +/// gap. An occupied path is removed only after it is proven to be a stale +/// current-user socket. +#[cfg(unix)] +async fn bind_secure_execution_runtime_socket( + requested_path: &Path, +) -> io::Result { + let socket_path = prepare_execution_runtime_socket_path(requested_path)?; + if let Ok(metadata) = std::fs::symlink_metadata(&socket_path) { + validate_existing_execution_runtime_socket(&metadata)?; } - let listener = tokio::net::UnixListener::bind(socket_path)?; - axum::serve( - listener, - build_execution_runtime_router_with_request_gates( - max_in_flight_requests, - distributed_request_gate, - ), - ) - .await?; + // Try the requested path first. If it is occupied, only a current-user + // socket that is demonstrably stale may be removed; every other path + // fails closed. This removes the old unlink-before-bind TOCTOU window. + match bind_execution_runtime_listener(&socket_path) { + Ok(listener) => { + harden_execution_runtime_socket(&listener, &socket_path)?; + Ok(listener) + } + Err(error) if error.kind() == io::ErrorKind::AddrInUse => { + remove_stale_execution_runtime_socket(&socket_path).await?; + let listener = bind_execution_runtime_listener(&socket_path)?; + harden_execution_runtime_socket(&listener, &socket_path)?; + Ok(listener) + } + Err(error) => Err(error), + } +} + +#[cfg(unix)] +fn bind_execution_runtime_listener(socket_path: &Path) -> io::Result { + // Unix bind derives the socket mode from the process umask. Darwin does + // not support fchmod on a Unix-socket fd, so make the inode private at + // creation time instead of briefly publishing a world-accessible socket. + // Serialize the temporary process-wide umask change with other binds made + // by this component and restore it before returning. + static BIND_LOCK: OnceLock> = OnceLock::new(); + let _lock = BIND_LOCK + .get_or_init(|| Mutex::new(())) + .lock() + .map_err(|_| io::Error::other("execution runtime socket bind lock poisoned"))?; + let previous_umask = unsafe { libc::umask(0o177) }; + let result = tokio::net::UnixListener::bind(socket_path); + unsafe { + libc::umask(previous_umask); + } + result +} + +#[cfg(unix)] +fn prepare_execution_runtime_socket_path(requested_path: &Path) -> io::Result { + use std::ffi::OsStr; + use std::os::unix::fs::{DirBuilderExt, MetadataExt, PermissionsExt}; + + let file_name = requested_path.file_name().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "execution runtime socket path must name a file", + ) + })?; + if file_name == OsStr::new(".") || file_name == OsStr::new("..") { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "execution runtime socket path has an invalid file name", + )); + } + + let requested_parent = requested_path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + + // Do not let a caller-controlled symlink redirect directory creation (or + // the eventual socket) into an unrelated tree. Root-owned compatibility + // links such as macOS `/tmp -> /private/tmp` and Linux `/var/run -> /run` + // remain allowed; links owned by an unprivileged user fail closed. + validate_requested_execution_runtime_parent_components(requested_parent)?; + let mut directory_builder = std::fs::DirBuilder::new(); + directory_builder.recursive(true).mode(0o700); + directory_builder.create(requested_parent)?; + validate_requested_execution_runtime_parent_components(requested_parent)?; + + let parent = std::fs::canonicalize(requested_parent)?; + validate_execution_runtime_socket_parent(&parent)?; + let metadata = std::fs::symlink_metadata(&parent)?; + if metadata.mode() & 0o022 != 0 + && metadata.mode() & 0o1000 == 0 + && metadata.uid() == unsafe { libc::geteuid() } + { + std::fs::set_permissions(&parent, std::fs::Permissions::from_mode(0o700))?; + } + + Ok(parent.join(file_name)) +} + +#[cfg(unix)] +fn validate_requested_execution_runtime_parent_components(parent: &Path) -> io::Result<()> { + use std::os::unix::fs::MetadataExt; + use std::path::Component; + + let mut prefix = PathBuf::new(); + for component in parent.components() { + match component { + Component::Prefix(prefix_component) => prefix.push(prefix_component.as_os_str()), + Component::RootDir => prefix.push(Path::new("/")), + Component::CurDir => {} + Component::ParentDir => prefix.push(".."), + Component::Normal(name) => { + prefix.push(name); + let metadata = match std::fs::symlink_metadata(&prefix) { + Ok(metadata) => metadata, + Err(error) if error.kind() == io::ErrorKind::NotFound => break, + Err(error) => return Err(error), + }; + if metadata.file_type().is_symlink() { + if metadata.uid() != 0 { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket parent contains an untrusted symlink", + )); + } + } else if !metadata.is_dir() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket parent contains a non-directory component", + )); + } + } + } + } Ok(()) } +#[cfg(unix)] +fn validate_execution_runtime_socket_parent(parent: &Path) -> io::Result<()> { + use std::os::unix::fs::MetadataExt; + + let effective_uid = unsafe { libc::geteuid() }; + let mut current = Some(parent); + let mut is_immediate_parent = true; + while let Some(path) = current { + let metadata = std::fs::symlink_metadata(path)?; + if metadata.file_type().is_symlink() || !metadata.is_dir() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket parent contains an unsafe path component", + )); + } + let mode = metadata.mode(); + if metadata.uid() != effective_uid && metadata.uid() != 0 { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket parent has untrusted ownership", + )); + } + if mode & 0o022 != 0 && mode & 0o1000 == 0 && metadata.uid() != effective_uid { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket parent must not be writable without sticky protection", + )); + } + if !is_immediate_parent && mode & 0o022 != 0 && mode & 0o1000 == 0 { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket ancestor is writable without sticky protection", + )); + } + is_immediate_parent = false; + current = path.parent(); + } + Ok(()) +} + +#[cfg(unix)] +fn harden_execution_runtime_socket( + listener: &tokio::net::UnixListener, + socket_path: &Path, +) -> io::Result<()> { + use std::mem::MaybeUninit; + use std::os::unix::fs::{FileTypeExt, MetadataExt}; + use std::os::unix::io::AsRawFd; + + let metadata = std::fs::symlink_metadata(socket_path)?; + let effective_uid = unsafe { libc::geteuid() }; + let mut stat = MaybeUninit::::uninit(); + if unsafe { libc::fstat(listener.as_raw_fd(), stat.as_mut_ptr()) } != 0 { + return Err(io::Error::last_os_error()); + } + let stat = unsafe { stat.assume_init() }; + let fd_mode = stat.st_mode as libc::mode_t; + if metadata.file_type().is_symlink() + || !metadata.file_type().is_socket() + || metadata.uid() != effective_uid + || metadata.nlink() != 1 + || metadata.mode() & 0o777 != 0o600 + // A pathname Unix socket and its connected sockfs inode do not have + // portable pathname identity. In particular, Linux can report a + // different inode number (and always reports a different device) for + // the descriptor than for the dentry. Validate the descriptor's own + // type and owner instead; the canonical, owner-checked + // parent and the path checks above prevent an untrusted user from + // replacing this private socket. + || (fd_mode & libc::S_IFMT) != libc::S_IFSOCK + || stat.st_uid != effective_uid + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "execution runtime socket path changed or has unsafe ownership", + )); + } + Ok(()) +} + +#[cfg(unix)] +fn validate_existing_execution_runtime_socket(metadata: &std::fs::Metadata) -> io::Result<()> { + use std::os::unix::fs::{FileTypeExt, MetadataExt}; + let effective_uid = unsafe { libc::geteuid() }; + if metadata.file_type().is_symlink() + || !metadata.file_type().is_socket() + || metadata.uid() != effective_uid + || metadata.nlink() != 1 + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "occupied execution runtime socket path is not a current-user socket", + )); + } + Ok(()) +} + +#[cfg(unix)] +async fn remove_stale_execution_runtime_socket(socket_path: &Path) -> io::Result<()> { + use std::os::unix::fs::{FileTypeExt, MetadataExt}; + use std::time::Duration; + + let metadata = match std::fs::symlink_metadata(socket_path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(error), + }; + validate_existing_execution_runtime_socket(&metadata)?; + match tokio::time::timeout( + Duration::from_millis(100), + tokio::net::UnixStream::connect(socket_path), + ) + .await + { + Ok(Ok(_stream)) => { + return Err(io::Error::new( + io::ErrorKind::AddrInUse, + "execution runtime socket is already serving requests", + )); + } + Ok(Err(error)) + if matches!( + error.kind(), + io::ErrorKind::ConnectionRefused | io::ErrorKind::NotFound + ) => {} + Ok(Err(error)) => return Err(error), + Err(_) => { + return Err(io::Error::new( + io::ErrorKind::TimedOut, + "could not determine whether execution runtime socket is active", + )); + } + } + + let latest = std::fs::symlink_metadata(socket_path)?; + validate_existing_execution_runtime_socket(&latest)?; + if latest.ino() != metadata.ino() || latest.dev() != metadata.dev() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "occupied execution runtime socket path changed during cleanup", + )); + } + std::fs::remove_file(socket_path) +} + #[cfg(not(unix))] pub async fn serve_execution_runtime_unix( _socket_path: &Path, @@ -309,19 +791,24 @@ async fn parse_request_json(request: Request) -> Result usize { + usize::try_from(configured_limit) + .unwrap_or(EXECUTION_RUNTIME_REQUEST_BODY_HARD_LIMIT_BYTES) + .min(EXECUTION_RUNTIME_REQUEST_BODY_HARD_LIMIT_BYTES) +} + fn build_overloaded_response(message: &str) -> Response { ( StatusCode::SERVICE_UNAVAILABLE, @@ -352,7 +839,7 @@ struct ExecutionRuntimeAppError(ExecutionRuntimeServerError); impl IntoResponse for ExecutionRuntimeAppError { fn into_response(self) -> Response { - let status_code = match self.0 { + let status_code = match &self.0 { ExecutionRuntimeServerError::RequestRead(_) | ExecutionRuntimeServerError::InvalidRequestJson(_) => StatusCode::BAD_REQUEST, ExecutionRuntimeServerError::Overloaded { .. } => { @@ -360,7 +847,9 @@ impl IntoResponse for ExecutionRuntimeAppError { } ExecutionRuntimeServerError::Transport( ExecutionRuntimeTransportError::RequestBodyRequired + | ExecutionRuntimeTransportError::RequestBodyAmbiguous | ExecutionRuntimeTransportError::BodyDecode(_) + | ExecutionRuntimeTransportError::BodyTooLarge { .. } | ExecutionRuntimeTransportError::UnsupportedContentEncoding(_) | ExecutionRuntimeTransportError::ProxyUnsupported | ExecutionRuntimeTransportError::InvalidMethod(_) @@ -372,7 +861,7 @@ impl IntoResponse for ExecutionRuntimeAppError { ) => StatusCode::BAD_REQUEST, ExecutionRuntimeServerError::Transport( ExecutionRuntimeTransportError::UpstreamHttpStatus { status_code, .. }, - ) => StatusCode::from_u16(status_code).unwrap_or(StatusCode::BAD_GATEWAY), + ) => StatusCode::from_u16(*status_code).unwrap_or(StatusCode::BAD_GATEWAY), ExecutionRuntimeServerError::Transport( ExecutionRuntimeTransportError::ClientBuild(_) | ExecutionRuntimeTransportError::BrowserClientBuild(_) @@ -384,11 +873,43 @@ impl IntoResponse for ExecutionRuntimeAppError { | ExecutionRuntimeTransportError::InvalidJson(_), ) => StatusCode::BAD_GATEWAY, }; + let message = match &self.0 { + ExecutionRuntimeServerError::RequestRead(_) + | ExecutionRuntimeServerError::InvalidRequestJson(_) + | ExecutionRuntimeServerError::Transport( + ExecutionRuntimeTransportError::RequestBodyRequired + | ExecutionRuntimeTransportError::RequestBodyAmbiguous + | ExecutionRuntimeTransportError::BodyDecode(_) + | ExecutionRuntimeTransportError::BodyTooLarge { .. } + | ExecutionRuntimeTransportError::UnsupportedContentEncoding(_) + | ExecutionRuntimeTransportError::ProxyUnsupported + | ExecutionRuntimeTransportError::InvalidMethod(_) + | ExecutionRuntimeTransportError::InvalidHeaderName(_) + | ExecutionRuntimeTransportError::InvalidHeaderValue(_) + | ExecutionRuntimeTransportError::InvalidProxy(_) + | ExecutionRuntimeTransportError::UnsupportedTransportProfile(_) + | ExecutionRuntimeTransportError::BodyEncode(_), + ) => "Invalid execution runtime request".to_string(), + ExecutionRuntimeServerError::Transport( + ExecutionRuntimeTransportError::UpstreamHttpStatus { status_code, .. }, + ) => format!("Upstream request returned HTTP {status_code}"), + ExecutionRuntimeServerError::Transport( + ExecutionRuntimeTransportError::ClientBuild(_) + | ExecutionRuntimeTransportError::BrowserClientBuild(_) + | ExecutionRuntimeTransportError::BrowserBody(_) + | ExecutionRuntimeTransportError::UpstreamRequest(_) + | ExecutionRuntimeTransportError::UpstreamResponseTooLarge { .. } + | ExecutionRuntimeTransportError::UpstreamResponseDecode { .. } + | ExecutionRuntimeTransportError::RelayError(_) + | ExecutionRuntimeTransportError::InvalidJson(_), + ) => "Upstream request failed".to_string(), + ExecutionRuntimeServerError::Overloaded { .. } => unreachable!(), + }; ( status_code, Json(json!({ - "error": self.0.to_string(), + "error": message, })), ) .into_response() @@ -399,7 +920,9 @@ impl IntoResponse for ExecutionRuntimeAppError { mod tests { use super::{ build_execution_runtime_router_with_request_concurrency_limit, - build_execution_runtime_router_with_request_gates, DISTRIBUTED_REQUEST_GATE_NAME, + build_execution_runtime_router_with_request_gates, + execution_runtime_request_body_limit_bytes, validate_execution_runtime_tcp_bind, + ExecutionRuntimeAppError, ExecutionRuntimeServerError, DISTRIBUTED_REQUEST_GATE_NAME, }; use aether_contracts::{ ExecutionPlan, ExecutionTimeouts, RequestBody, StreamFrame, StreamFrameType, @@ -408,13 +931,27 @@ mod tests { MemoryRuntimeStateConfig, RuntimeSemaphore, RuntimeSemaphoreConfig, RuntimeState, }; use axum::body::{Body, Bytes}; - use axum::response::Response; + use axum::response::{IntoResponse, Response}; use axum::routing::any; use axum::{extract::Request, Router}; use http::StatusCode; + use http_body_util::{BodyExt, Full}; + use hyper::body::Incoming as HyperIncoming; + use hyper::{Request as HyperRequest, Response as HyperResponse}; + use hyper_util::rt::{TokioExecutor, TokioIo}; use std::convert::Infallible; + use std::fs; + #[cfg(unix)] + use std::os::unix::fs::{FileTypeExt, MetadataExt}; + #[cfg(unix)] + use std::os::unix::net::UnixListener as StdUnixListener; + use std::path::PathBuf; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tower::service_fn; + + use crate::execution_runtime::ExecutionRuntimeTransportError; fn distributed_gate(gate: &'static str, limit: usize) -> RuntimeSemaphore { RuntimeState::memory(MemoryRuntimeStateConfig::default()) @@ -433,6 +970,183 @@ mod tests { (format!("http://{addr}"), handle) } + #[test] + fn execution_runtime_tcp_bind_accepts_only_literal_loopback_addresses() { + for bind in ["127.0.0.1:0", "127.42.17.9:5219", "[::1]:0"] { + let address = validate_execution_runtime_tcp_bind(bind) + .unwrap_or_else(|error| panic!("{bind} should be accepted: {error}")); + assert!(address.ip().is_loopback()); + } + } + + #[test] + fn execution_runtime_tcp_bind_rejects_wildcard_non_loopback_and_unparseable_addresses() { + for bind in [ + "0.0.0.0:5219", + "[::]:5219", + "10.0.0.1:5219", + "192.168.1.10:5219", + "localhost:5219", + "not-a-socket-address", + "127.0.0.1", + ] { + assert!( + validate_execution_runtime_tcp_bind(bind).is_err(), + "{bind} must be rejected" + ); + } + } + + #[test] + fn execution_runtime_parser_defaults_preserve_high_http2_concurrency() { + assert_eq!( + super::EXECUTION_RUNTIME_HTTP2_MAX_CONCURRENT_STREAMS, + 16_384 + ); + assert_eq!( + super::EXECUTION_RUNTIME_HTTP_HEADER_READ_TIMEOUT, + std::time::Duration::from_secs(30) + ); + assert_eq!(super::EXECUTION_RUNTIME_HTTP_HEADER_MAX_BYTES, 64 * 1024); + assert_eq!(super::EXECUTION_RUNTIME_HTTP_MAX_HEADERS, 256); + } + + #[test] + fn execution_runtime_request_body_limit_never_accepts_unbounded_sentinel() { + assert_eq!( + execution_runtime_request_body_limit_bytes(u64::MAX), + super::EXECUTION_RUNTIME_REQUEST_BODY_HARD_LIMIT_BYTES + ); + assert_eq!( + execution_runtime_request_body_limit_bytes(512 * 1024 * 1024), + super::EXECUTION_RUNTIME_REQUEST_BODY_HARD_LIMIT_BYTES + ); + assert_eq!(execution_runtime_request_body_limit_bytes(1024), 1024); + } + + #[tokio::test] + async fn execution_runtime_first_request_deadline_covers_partial_protocol_input() { + let prefixes: &[&[u8]] = &[ + b"G", + b"GET /health HTTP/1.1\r\nHost: localhost\r\n", + b"PRI * HTTP/2.0\r\n\r\nSM\r\n", + ]; + + for prefix in prefixes { + let (mut client, server) = tokio::io::duplex(16 * 1024); + client + .write_all(prefix) + .await + .expect("fixture prefix should be writable"); + + let gate = super::ExecutionRuntimeFirstRequestGate::new(); + let service = service_fn(|_request: HyperRequest| async { + Ok::<_, Infallible>(HyperResponse::new(Full::new(bytes::Bytes::from_static( + b"ok", + )))) + }); + let builder = super::execution_runtime_http_builder(); + let result = tokio::time::timeout( + std::time::Duration::from_secs(1), + super::drive_execution_runtime_connection( + builder.serve_connection_with_upgrades( + TokioIo::new(server), + super::TowerToHyperService::new( + super::ExecutionRuntimeFirstRequestService { + inner: service, + gate: gate.clone(), + }, + ), + ), + gate, + std::time::Duration::from_millis(5), + ), + ) + .await + .expect("partial protocol input should hit the first-request deadline"); + assert!(result.is_ok()); + + let mut byte = [0u8; 1]; + let read = + tokio::time::timeout(std::time::Duration::from_secs(1), client.read(&mut byte)) + .await + .expect("timed-out connection should close its peer"); + assert!(matches!(read, Ok(0) | Err(_))); + } + } + + #[tokio::test] + async fn execution_runtime_first_request_deadline_does_not_cut_off_body_streaming() { + let (mut client, server) = tokio::io::duplex(16 * 1024); + let gate = super::ExecutionRuntimeFirstRequestGate::new(); + let (headers_seen_tx, mut headers_seen_rx) = tokio::sync::mpsc::unbounded_channel(); + let service = service_fn(move |request: HyperRequest| { + let _ = headers_seen_tx.send(()); + async move { + let body = request + .into_body() + .collect() + .await + .expect("test body should decode") + .to_bytes(); + Ok::<_, Infallible>(HyperResponse::new(Full::new(body))) + } + }); + let builder = super::execution_runtime_http_builder(); + let server_task = tokio::spawn(async move { + super::drive_execution_runtime_connection( + builder.serve_connection_with_upgrades( + TokioIo::new(server), + super::TowerToHyperService::new(super::ExecutionRuntimeFirstRequestService { + inner: service, + gate: gate.clone(), + }), + ), + gate, + std::time::Duration::from_millis(20), + ) + .await + }); + + let large_body = vec![b'x'; 128 * 1024]; + let request_headers = format!( + "POST /health HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\nContent-Length: {}\r\n\r\n", + large_body.len() + ); + client + .write_all(request_headers.as_bytes()) + .await + .expect("request headers should be writable"); + tokio::time::timeout(std::time::Duration::from_secs(1), headers_seen_rx.recv()) + .await + .expect("request headers should reach the service") + .expect("service notification should remain available"); + + tokio::time::sleep(std::time::Duration::from_millis(40)).await; + // The parser's 64 KiB metadata buffer must not become a body-size or + // body-throughput limit. Send a body larger than that buffer after the + // first-request gate has opened and verify it is echoed intact. + client + .write_all(&large_body) + .await + .expect("body should remain writable after the header deadline"); + let mut response = Vec::new(); + tokio::time::timeout( + std::time::Duration::from_secs(1), + client.read_to_end(&mut response), + ) + .await + .expect("streaming response should complete") + .expect("response should be readable"); + let result = server_task + .await + .expect("server connection task should join"); + assert!(result.is_ok()); + assert!(response + .windows(large_body.len()) + .any(|window| window == large_body.as_slice())); + } + fn stream_plan(url: String) -> ExecutionPlan { ExecutionPlan { request_id: "req-1".into(), @@ -465,6 +1179,159 @@ mod tests { } } + #[cfg(unix)] + fn test_socket_path() -> PathBuf { + std::env::temp_dir() + .join(format!("ar{}", uuid::Uuid::new_v4().simple())) + .join("n") + .join("s.sock") + } + + #[cfg(unix)] + fn remove_test_socket_path(socket_path: &std::path::Path) { + if let Some(parent) = socket_path.parent() { + if let Some(root) = parent.parent() { + let _ = fs::remove_dir_all(root); + } + } + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_and_new_parent_are_private() { + let socket_path = test_socket_path(); + let listener = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect("socket should bind"); + + let socket_metadata = fs::symlink_metadata(&socket_path).expect("socket should exist"); + assert!(socket_metadata.file_type().is_socket()); + assert_eq!(socket_metadata.uid(), unsafe { libc::geteuid() }); + assert_eq!(socket_metadata.mode() & 0o777, 0o600); + + let parent = socket_path.parent().expect("socket should have a parent"); + let parent_mode = fs::symlink_metadata(parent) + .expect("parent should exist") + .mode() + & 0o777; + assert_eq!(parent_mode, 0o700); + + drop(listener); + remove_test_socket_path(&socket_path); + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_rejects_regular_file_path() { + let socket_path = test_socket_path(); + let parent = socket_path.parent().expect("socket should have a parent"); + fs::create_dir_all(parent).expect("parent should be created"); + fs::write(&socket_path, b"do not replace this file").expect("file should be created"); + + let error = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect_err("regular file must not be replaced"); + assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied); + assert_eq!( + fs::read(&socket_path).expect("file should remain"), + b"do not replace this file" + ); + + remove_test_socket_path(&socket_path); + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_rejects_symlink_path() { + use std::os::unix::fs::symlink; + + let socket_path = test_socket_path(); + let parent = socket_path.parent().expect("socket should have a parent"); + fs::create_dir_all(parent).expect("parent should be created"); + let target = parent.join("target"); + fs::write(&target, b"target must not be followed").expect("target should be created"); + symlink(&target, &socket_path).expect("symlink should be created"); + + let error = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect_err("symlink must not be replaced or followed"); + assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied); + assert_eq!( + fs::read(&target).expect("target should remain"), + b"target must not be followed" + ); + + remove_test_socket_path(&socket_path); + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_rejects_untrusted_parent_symlink() { + use std::os::unix::fs::symlink; + + let socket_path = test_socket_path(); + let parent_root = socket_path + .parent() + .and_then(std::path::Path::parent) + .expect("socket should have a test root"); + fs::create_dir_all(parent_root).expect("test root should be created"); + let target = parent_root.join("target"); + fs::create_dir_all(&target).expect("symlink target should be created"); + let link = parent_root.join("link"); + symlink(&target, &link).expect("parent symlink should be created"); + let linked_socket_path = link.join("runtime.sock"); + + let error = super::bind_secure_execution_runtime_socket(&linked_socket_path) + .await + .expect_err("untrusted parent symlink must not be followed"); + assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied); + assert!(!target.join("runtime.sock").exists()); + + remove_test_socket_path(&socket_path); + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_does_not_replace_active_listener() { + let socket_path = test_socket_path(); + let first_listener = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect("first socket should bind"); + + let error = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect_err("active socket must not be replaced"); + assert_eq!(error.kind(), std::io::ErrorKind::AddrInUse); + assert!(fs::symlink_metadata(&socket_path) + .expect("active socket should remain") + .file_type() + .is_socket()); + + drop(first_listener); + remove_test_socket_path(&socket_path); + } + + #[cfg(unix)] + #[tokio::test] + async fn execution_runtime_unix_socket_rebinds_stale_current_user_socket() { + let socket_path = test_socket_path(); + let parent = socket_path.parent().expect("socket should have a parent"); + fs::create_dir_all(parent).expect("parent should be created"); + let stale_listener = StdUnixListener::bind(&socket_path).expect("stale socket should bind"); + drop(stale_listener); + + let listener = super::bind_secure_execution_runtime_socket(&socket_path) + .await + .expect("stale socket should be replaced"); + let rebound_metadata = + fs::symlink_metadata(&socket_path).expect("rebound socket should exist"); + assert!(rebound_metadata.file_type().is_socket()); + assert_eq!(rebound_metadata.mode() & 0o777, 0o600); + + drop(listener); + remove_test_socket_path(&socket_path); + } + #[tokio::test] async fn execution_runtime_stream_endpoint_carries_non_stream_upstream_plan() { let upstream = Router::new().route( @@ -667,4 +1534,41 @@ mod tests { runtime_handle.abort(); } + + #[tokio::test] + async fn execution_runtime_transport_errors_do_not_expose_internal_details() { + let secret = "Bearer upstream-secret https://user:password@example.test?q=token"; + let response = ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport( + ExecutionRuntimeTransportError::UpstreamRequest(secret.to_string()), + )) + .into_response(); + + assert_eq!(response.status(), StatusCode::BAD_GATEWAY); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("error body should read"); + let body = String::from_utf8(body.to_vec()).expect("error body should be utf8"); + assert_eq!(body, r#"{"error":"Upstream request failed"}"#); + assert!(!body.contains(secret)); + assert!(!body.contains("upstream-secret")); + } + + #[tokio::test] + async fn execution_runtime_upstream_status_keeps_only_status_diagnostics() { + let response = ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport( + ExecutionRuntimeTransportError::UpstreamHttpStatus { + status_code: 429, + message: "authorization=Bearer upstream-secret".to_string(), + }, + )) + .into_response(); + + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("error body should read"); + let body = String::from_utf8(body.to_vec()).expect("error body should be utf8"); + assert_eq!(body, r#"{"error":"Upstream request returned HTTP 429"}"#); + assert!(!body.contains("upstream-secret")); + } } diff --git a/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs b/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs index 7645c67e2..9cfd0cae1 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs @@ -5,6 +5,7 @@ use serde_json::Value; use crate::execution_runtime::MAX_STREAM_PREFETCH_BYTES; const ANTHROPIC_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750); +const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750); #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(super) enum StreamCommitPolicy { @@ -14,6 +15,10 @@ pub(super) enum StreamCommitPolicy { max_bytes: usize, max_wait: Duration, }, + FirstGeminiSemanticEvent { + max_bytes: usize, + max_wait: Duration, + }, } impl StreamCommitPolicy { @@ -51,6 +56,12 @@ impl StreamCommitPolicy { max_wait: ANTHROPIC_PRECOMMIT_MAX_WAIT, }; } + if provider_api_format.eq_ignore_ascii_case("gemini:generate_content") { + return Self::FirstGeminiSemanticEvent { + max_bytes: MAX_STREAM_PREFETCH_BYTES, + max_wait: GEMINI_PRECOMMIT_MAX_WAIT, + }; + } return Self::ResponseHeaders; } @@ -78,12 +89,16 @@ impl StreamCommitPolicy { } pub(super) const fn requires_bounded_frame_wait(self) -> bool { - matches!(self, Self::FirstAnthropicSemanticEvent { .. }) + matches!( + self, + Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. } + ) } pub(super) const fn max_precommit_wait(self) -> Option { match self { - Self::FirstAnthropicSemanticEvent { max_wait, .. } => Some(max_wait), + Self::FirstAnthropicSemanticEvent { max_wait, .. } + | Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait), Self::ResponseHeaders | Self::FirstClassifiedBody => None, } } @@ -91,6 +106,10 @@ impl StreamCommitPolicy { pub(super) const fn is_native_anthropic(self) -> bool { matches!(self, Self::FirstAnthropicSemanticEvent { .. }) } + + pub(super) const fn is_gemini(self) -> bool { + matches!(self, Self::FirstGeminiSemanticEvent { .. }) + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -113,6 +132,7 @@ pub(super) struct StreamCommitGate { state: StreamCommitState, observed_bytes: usize, anthropic: AnthropicSsePrecommitInspector, + gemini: GeminiSsePrecommitInspector, } impl StreamCommitGate { @@ -127,6 +147,7 @@ impl StreamCommitGate { state, observed_bytes: 0, anthropic: AnthropicSsePrecommitInspector::default(), + gemini: GeminiSsePrecommitInspector::default(), } } @@ -143,21 +164,32 @@ impl StreamCommitGate { return StreamPrecommitObservation::Commit; } - let StreamCommitPolicy::FirstAnthropicSemanticEvent { max_bytes, .. } = self.policy else { - return StreamPrecommitObservation::Pending; + let (max_bytes, observation) = match self.policy { + StreamCommitPolicy::FirstAnthropicSemanticEvent { max_bytes, .. } => { + (max_bytes, self.anthropic.observe(chunk, max_bytes)) + } + StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => { + (max_bytes, self.gemini.observe(chunk, max_bytes)) + } + StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => { + return StreamPrecommitObservation::Pending; + } }; self.observed_bytes = self.observed_bytes.saturating_add(chunk.len()); - match self.anthropic.observe(chunk, max_bytes) { - AnthropicSseObservation::Pending => {} - AnthropicSseObservation::SemanticEvent => { + match observation { + SemanticSseObservation::Pending => {} + SemanticSseObservation::SemanticEvent => { self.state = StreamCommitState::Committed; return StreamPrecommitObservation::Commit; } - AnthropicSseObservation::Error(body_json) => { + SemanticSseObservation::Error { + status_code, + body_json, + } => { self.state = StreamCommitState::Terminal; return StreamPrecommitObservation::UpstreamError { - status_code: anthropic_error_status_code(&body_json), + status_code, body_json, }; } @@ -179,10 +211,10 @@ impl StreamCommitGate { } #[derive(Debug)] -enum AnthropicSseObservation { +enum SemanticSseObservation { Pending, SemanticEvent, - Error(Value), + Error { status_code: u16, body_json: Value }, } #[derive(Debug, Default)] @@ -191,7 +223,7 @@ struct AnthropicSsePrecommitInspector { } impl AnthropicSsePrecommitInspector { - fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> AnthropicSseObservation { + fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation { let remaining = max_bytes.saturating_sub(self.buffered.len()); let truncated = chunk.len() > remaining; self.buffered @@ -201,15 +233,44 @@ impl AnthropicSsePrecommitInspector { let record = self.buffered[..record_end].to_vec(); self.buffered.drain(..record_end + separator_len); match classify_anthropic_sse_record(&record) { - AnthropicSseObservation::Pending => {} + SemanticSseObservation::Pending => {} decision => return decision, } } if truncated { - AnthropicSseObservation::SemanticEvent + SemanticSseObservation::SemanticEvent } else { - AnthropicSseObservation::Pending + SemanticSseObservation::Pending + } + } +} + +#[derive(Debug, Default)] +struct GeminiSsePrecommitInspector { + buffered: Vec, +} + +impl GeminiSsePrecommitInspector { + fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation { + let remaining = max_bytes.saturating_sub(self.buffered.len()); + let truncated = chunk.len() > remaining; + self.buffered + .extend_from_slice(&chunk[..chunk.len().min(remaining)]); + + while let Some((record_end, separator_len)) = find_sse_record_boundary(&self.buffered) { + let record = self.buffered[..record_end].to_vec(); + self.buffered.drain(..record_end + separator_len); + match classify_gemini_sse_record(&record) { + SemanticSseObservation::Pending => {} + decision => return decision, + } + } + + if truncated { + SemanticSseObservation::SemanticEvent + } else { + SemanticSseObservation::Pending } } } @@ -249,9 +310,9 @@ fn next_sse_line_ending(buffer: &[u8], start: usize) -> Option<(usize, usize)> { Some((index, ending_len)) } -fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation { +fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation { let Ok(record) = std::str::from_utf8(record) else { - return AnthropicSseObservation::Pending; + return SemanticSseObservation::Pending; }; let normalized_record = record.replace("\r\n", "\n").replace('\r', "\n"); let mut event_type = None; @@ -275,15 +336,18 @@ fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation { } } if data.trim().is_empty() { - return AnthropicSseObservation::Pending; + return SemanticSseObservation::Pending; } let Ok(body_json) = serde_json::from_str::(data.trim()) else { - return AnthropicSseObservation::Pending; + return SemanticSseObservation::Pending; }; let payload_type = body_json.get("type").and_then(Value::as_str).map(str::trim); if event_type == Some("error") || payload_type == Some("error") { - return AnthropicSseObservation::Error(body_json); + return SemanticSseObservation::Error { + status_code: anthropic_error_status_code(&body_json), + body_json, + }; } let semantic_type = match (event_type, payload_type) { @@ -292,12 +356,123 @@ fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation { _ => None, }; if semantic_type.is_some_and(is_anthropic_semantic_event_type) { - AnthropicSseObservation::SemanticEvent + SemanticSseObservation::SemanticEvent } else { - AnthropicSseObservation::Pending + SemanticSseObservation::Pending } } +fn classify_gemini_sse_record(record: &[u8]) -> SemanticSseObservation { + let Ok(record) = std::str::from_utf8(record) else { + return SemanticSseObservation::Pending; + }; + let data = record + .replace("\r\n", "\n") + .replace('\r', "\n") + .lines() + .filter_map(|line| line.strip_prefix("data:").map(str::trim_start)) + .collect::>() + .join("\n"); + if data.trim().is_empty() { + return SemanticSseObservation::Pending; + } + if data.trim() == "[DONE]" { + return SemanticSseObservation::SemanticEvent; + } + + let Ok(body_json) = serde_json::from_str::(data.trim()) else { + return SemanticSseObservation::Pending; + }; + let response = body_json.get("response").unwrap_or(&body_json); + let Some(candidates) = response.get("candidates").and_then(Value::as_array) else { + return SemanticSseObservation::Pending; + }; + + for candidate in candidates { + let finish_reason = candidate + .get("finishReason") + .or_else(|| candidate.get("finish_reason")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + if let Some(finish_reason) = finish_reason.filter(|reason| { + matches!( + *reason, + "MALFORMED_FUNCTION_CALL" + | "UNEXPECTED_TOOL_CALL" + | "TOO_MANY_TOOL_CALLS" + | "MISSING_THOUGHT_SIGNATURE" + | "MALFORMED_RESPONSE" + ) + }) { + let message = candidate + .get("finishMessage") + .or_else(|| candidate.get("finish_message")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| format!("Gemini stream ended with {finish_reason}")); + return SemanticSseObservation::Error { + status_code: 502, + body_json: serde_json::json!({ + "error": { + "type": "upstream_gemini_finish_error", + "code": finish_reason, + "message": message, + "upstream_status": 200 + } + }), + }; + } + + if finish_reason.is_some() { + return SemanticSseObservation::SemanticEvent; + } + let Some(parts) = candidate + .get("content") + .and_then(|content| content.get("parts")) + .and_then(Value::as_array) + else { + continue; + }; + if parts.iter().any(gemini_part_is_client_semantic) { + return SemanticSseObservation::SemanticEvent; + } + } + + SemanticSseObservation::Pending +} + +fn gemini_part_is_client_semantic(part: &Value) -> bool { + let Some(part) = part.as_object() else { + return true; + }; + if part + .keys() + .any(|key| !matches!(key.as_str(), "text" | "thought" | "thoughtSignature")) + { + return true; + } + if part.get("thought").and_then(Value::as_bool) == Some(true) { + return part + .get("text") + .and_then(Value::as_str) + .is_some_and(|text| !text.is_empty()); + } + if part.keys().all(|key| key == "thoughtSignature") { + return false; + } + if part + .get("text") + .and_then(Value::as_str) + .is_some_and(|text| !text.is_empty()) + { + return true; + } + false +} + fn is_anthropic_semantic_event_type(event_type: &str) -> bool { matches!( event_type, @@ -346,6 +521,13 @@ mod tests { } } + fn gemini_policy() -> StreamCommitPolicy { + StreamCommitPolicy::FirstGeminiSemanticEvent { + max_bytes: 16_384, + max_wait: Duration::from_millis(750), + } + } + #[test] fn policy_selects_bounded_anthropic_gate_only_for_native_same_format_sse() { let native = StreamCommitPolicy::for_response( @@ -384,6 +566,104 @@ mod tests { .commits_on_response_headers()); } + #[test] + fn policy_selects_bounded_gemini_gate_for_event_streams() { + let policy = StreamCommitPolicy::for_response( + true, + Some("text/event-stream"), + "gemini:generate_content", + "openai:responses", + false, + true, + false, + ); + + assert!(policy.is_gemini()); + assert!(policy.requires_bounded_frame_wait()); + assert_eq!( + policy.max_precommit_wait(), + Some(Duration::from_millis(750)) + ); + } + + #[test] + fn gemini_gate_commits_on_first_nonempty_thought() { + let mut gate = StreamCommitGate::new(gemini_policy()); + let thought = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thought\":true,\"text\":\"checking\"}]}}]}}\n\n"; + + assert_eq!( + gate.observe_provider_bytes(thought), + StreamPrecommitObservation::Commit + ); + assert_eq!(gate.state(), StreamCommitState::Committed); + } + + #[test] + fn gemini_gate_commits_on_function_call_even_with_thought_marker() { + let mut gate = StreamCommitGate::new(gemini_policy()); + let tool_call = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thought\":true,\"functionCall\":{\"name\":\"validate\",\"args\":{}}}]}}]}}\n\n"; + + assert_eq!( + gate.observe_provider_bytes(tool_call), + StreamPrecommitObservation::Commit + ); + assert_eq!(gate.state(), StreamCommitState::Committed); + } + + #[test] + fn gemini_gate_rejects_malformed_function_call_before_commit() { + let mut gate = StreamCommitGate::new(gemini_policy()); + let thought = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thoughtSignature\":\"signature\",\"text\":\"\"}]}}]}}\n\n"; + let malformed = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thoughtSignature\":\"signature\",\"text\":\"\"}]},\"finishReason\":\"MALFORMED_FUNCTION_CALL\",\"finishMessage\":\"Malformed function call: Function call is empty - no input to parse.\"}]}}\n\n"; + + assert_eq!( + gate.observe_provider_bytes(thought), + StreamPrecommitObservation::Pending + ); + let StreamPrecommitObservation::UpstreamError { + status_code, + body_json, + } = gate.observe_provider_bytes(malformed) + else { + panic!("malformed Gemini function call should fail before stream commit"); + }; + + assert_eq!(status_code, 502); + assert_eq!(body_json["error"]["code"], "MALFORMED_FUNCTION_CALL"); + assert_eq!( + body_json["error"]["message"], + "Malformed function call: Function call is empty - no input to parse." + ); + assert_eq!(gate.state(), StreamCommitState::Terminal); + } + + #[test] + fn gemini_gate_detects_malformed_function_call_across_chunk_boundaries() { + let malformed = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thoughtSignature\":\"signature\",\"text\":\"\"}]},\"finishReason\":\"MALFORMED_FUNCTION_CALL\",\"finishMessage\":\"empty call\"}]}}\r\n\r\n"; + + for split in 1..malformed.len() { + let mut gate = StreamCommitGate::new(gemini_policy()); + let first_observation = gate.observe_provider_bytes(&malformed[..split]); + if !matches!( + first_observation, + StreamPrecommitObservation::UpstreamError { + status_code: 502, + .. + } + ) { + assert_eq!(first_observation, StreamPrecommitObservation::Pending); + assert!(matches!( + gate.observe_provider_bytes(&malformed[split..]), + StreamPrecommitObservation::UpstreamError { + status_code: 502, + .. + } + )); + } + assert_eq!(gate.state(), StreamCommitState::Terminal); + } + } + #[test] fn gate_detects_anthropic_error_across_every_chunk_boundary() { let event = b"event: error\r\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"busy\"}}\r\n\r\n"; diff --git a/apps/aether-gateway/src/execution_runtime/stream/error.rs b/apps/aether-gateway/src/execution_runtime/stream/error.rs index a9e983c6e..1c7f4e2c6 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/error.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/error.rs @@ -191,6 +191,46 @@ pub(super) fn inspect_prefetched_stream_body( } } +fn append_error_frame_payload( + body: &mut Vec, + chunk_b64: Option<&str>, + text: Option<&str>, +) -> Result { + let remaining = MAX_ERROR_BODY_BYTES.saturating_sub(body.len()); + if remaining == 0 { + return Ok(false); + } + + if let Some(chunk_b64) = chunk_b64 { + // Do not decode an attacker-controlled megabyte-scale base64 value + // merely to retain the first 16 KiB of an error body. A standard + // base64 value representing at most `remaining` bytes cannot exceed + // this length. + let max_encoded_len = remaining + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4); + if chunk_b64.len() > max_encoded_len { + warn!( + encoded_bytes = chunk_b64.len(), + max_encoded_bytes = max_encoded_len, + "execution runtime error frame body exceeded capture limit" + ); + return Ok(false); + } + let chunk = base64::engine::general_purpose::STANDARD + .decode(chunk_b64) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + body.extend_from_slice(&chunk[..chunk.len().min(remaining)]); + } else if let Some(text) = text { + let text_bytes = text.as_bytes(); + body.extend_from_slice(&text_bytes[..text_bytes.len().min(remaining)]); + } + + Ok(body.len() < MAX_ERROR_BODY_BYTES) +} + pub(super) async fn collect_error_body( lines: &mut FramedRead, ) -> Result, GatewayError> @@ -201,23 +241,19 @@ where while let Some(frame) = read_next_frame(lines).await? { match frame.payload { StreamFramePayload::Data { chunk_b64, text } => { - let chunk = if let Some(chunk_b64) = chunk_b64 { - base64::engine::general_purpose::STANDARD - .decode(chunk_b64) - .map_err(|err| GatewayError::Internal(err.to_string()))? - } else { - text.unwrap_or_default().into_bytes() - }; - body.extend_from_slice(&chunk); - if body.len() >= MAX_ERROR_BODY_BYTES { - body.truncate(MAX_ERROR_BODY_BYTES); + if !append_error_frame_payload(&mut body, chunk_b64.as_deref(), text.as_deref())? { break; } } StreamFramePayload::Telemetry { .. } => {} StreamFramePayload::Eof { .. } => break, StreamFramePayload::Error { error } => { - warn!(error = %error.message, "execution runtime stream emitted error frame while collecting error body"); + warn!( + error_kind = ?error.kind, + error_phase = ?error.phase, + upstream_status = ?error.upstream_status, + "execution runtime stream emitted error frame while collecting error body" + ); break; } StreamFramePayload::Headers { .. } => {} @@ -242,3 +278,32 @@ where } Ok(None) } + +#[cfg(test)] +mod tests { + use super::{append_error_frame_payload, MAX_ERROR_BODY_BYTES}; + + #[test] + fn oversized_base64_error_frame_is_rejected_before_decode() { + let mut body = Vec::new(); + let encoded = "x".repeat((MAX_ERROR_BODY_BYTES + 2) / 3 * 4 + 1); + + let keep_reading = append_error_frame_payload(&mut body, Some(&encoded), None) + .expect("oversized frame should be handled without a decode error"); + + assert!(!keep_reading); + assert!(body.is_empty()); + } + + #[test] + fn error_frame_payload_is_capped_to_remaining_capture_budget() { + let mut body = vec![b'a'; MAX_ERROR_BODY_BYTES - 2]; + + let keep_reading = append_error_frame_payload(&mut body, None, Some("hello")) + .expect("text payload should append"); + + assert!(!keep_reading); + assert_eq!(body.len(), MAX_ERROR_BODY_BYTES); + assert_eq!(&body[MAX_ERROR_BODY_BYTES - 2..], b"he"); + } +} diff --git a/apps/aether-gateway/src/execution_runtime/stream/execution.rs b/apps/aether-gateway/src/execution_runtime/stream/execution.rs index 80830455e..db8853a44 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/execution.rs @@ -21,11 +21,12 @@ use aether_data_contracts::repository::usage::UsageBodyCaptureState; use aether_scheduler_core::{ parse_request_candidate_report_context, SchedulerRequestCandidateStatusUpdate, }; +#[cfg(test)] +use aether_usage_runtime::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES; use aether_usage_runtime::{ build_lifecycle_usage_seed, build_stream_terminal_usage_payload_seed, build_sync_terminal_usage_payload_seed, build_terminal_usage_context_seed, LifecycleUsageSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed, UsageRequestRecordLevel, - DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, }; use async_stream::stream; use axum::body::{Body, Bytes}; @@ -71,12 +72,15 @@ use crate::ai_serving::api::{ UPSTREAM_IS_STREAM_KEY, }; use crate::ai_serving::is_openai_responses_family_format; +use crate::ai_serving::record_local_runtime_candidate_skip_reason; use crate::api::response::{ attach_control_metadata_headers, build_client_response, build_client_response_from_parts, + build_client_response_from_parts_with_mutator, }; use crate::clock::current_unix_ms as current_request_candidate_unix_ms; use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER}; use crate::control::GatewayControlDecision; +use crate::execution_runtime::attempt_cancellation::AttemptCancellationGuard; use crate::execution_runtime::build_direct_execution_frame_stream; use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_stream; use crate::execution_runtime::grok::maybe_execute_grok_stream; @@ -96,11 +100,12 @@ use crate::execution_runtime::submission::{ strip_utf8_bom_and_ws, submit_local_core_error_or_sync_finalize, }; use crate::execution_runtime::transport::{ - execute_stream_plan_via_local_tunnel, format_hyper_error_chain, format_upstream_request_error, - format_wreq_upstream_request_error, record_manual_proxy_request_failure, - record_manual_proxy_request_success, record_manual_proxy_stream_error, - stream_first_byte_timeout_message, DirectSyncExecutionRuntime, DirectUpstreamResponse, - DirectUpstreamStreamExecution, ExecutionRuntimeTransportError, + decode_base64_body_with_limit, execute_stream_plan_via_local_tunnel, format_hyper_error_chain, + format_upstream_request_error, format_wreq_upstream_request_error, + record_manual_proxy_request_failure, record_manual_proxy_request_success, + record_manual_proxy_stream_error, stream_first_byte_timeout_message, + DirectSyncExecutionRuntime, DirectUpstreamResponse, DirectUpstreamStreamExecution, + ExecutionRuntimeTransportError, }; use crate::execution_runtime::windsurf::maybe_execute_windsurf_stream; use crate::execution_runtime::{ @@ -117,8 +122,7 @@ use crate::execution_runtime::{ use crate::log_ids::short_request_id; use crate::orchestration::{ apply_local_execution_effect, build_local_error_flow_metadata, classify_failure_disposition, - cyber_continue_failover_enabled, spawn_local_oauth_success_effect, - trace_upstream_response_body, with_error_flow_report_context, + spawn_local_oauth_success_effect, trace_upstream_response_body, with_error_flow_report_context, with_upstream_response_report_context, FailureDisposition, FailureTokenAction, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext, LocalFailoverAnalysis, @@ -149,6 +153,11 @@ use crate::{ AppState, GatewayError, GEMINI_FILES_DOWNLOAD_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND, }; +/// Settlement labels for a stream attempt whose future is dropped before the +/// transport reaches a terminal state. +const STREAM_ATTEMPT_CANCELLED_ERROR_TYPE: &str = "local_stream_attempt_cancelled"; +const STREAM_ATTEMPT_CANCELLED_ERROR_MESSAGE: &str = "Local stream attempt was dropped before terminal finalization, usually because the client disconnected or the request task was cancelled."; + const SSE_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(15); const SSE_KEEPALIVE_BYTES: &[u8] = b": aether-keepalive\n\n"; const SSE_CONTROL_FILTER_MAX_BUFFER_BYTES: usize = 1024 * 1024; @@ -156,6 +165,7 @@ const SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES: usize = 1024 * 1024; const SSE_TERMINAL_DETECTOR_MAX_RECORD_BYTES: usize = SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES; const PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES: usize = SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES; const BASIC_STREAM_BODY_ANALYSIS_LIMIT_BYTES: usize = 5 * 1024 * 1024; +const MAX_EXECUTION_STREAM_DATA_CHUNK_BYTES: usize = 64 * 1024 * 1024; const STREAM_IDLE_LOG_INTERVAL: Duration = Duration::from_secs(60); const STREAM_IDLE_LOG_INTERVAL_MS: u64 = 60_000; const REWRITTEN_STREAM_PREFETCH_TIMEOUT: Duration = Duration::from_millis(750); @@ -183,7 +193,47 @@ impl ProviderStreamErrorInspection { if chunk.is_empty() { return None; } - if let Some(error_body) = extract_provider_private_stream_error_body(report_context, chunk) + + // A transport implementation may deliver a very large chunk. Keep + // every parser invocation bounded: inspect the prefix (including the + // previous rolling tail for events split across chunks) and suffix, + // while retaining only the bounded suffix for the next observation. + // The middle of an oversized chunk is deliberately skipped because + // this observer is best-effort and must never duplicate the client + // stream or turn a single upstream read into an unbounded parse. + if chunk.len() > PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES { + let prefix_len = PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES + .saturating_sub(self.buffered.len()) + .min(chunk.len()); + let mut boundary = Vec::with_capacity( + self.buffered + .len() + .saturating_add(prefix_len) + .min(PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES), + ); + boundary.extend_from_slice(&self.buffered); + boundary.extend_from_slice(&chunk[..prefix_len]); + if let Some(error_body) = + extract_provider_private_stream_error_body(report_context, &boundary) + { + return Some(error_body); + } + + let prefix = &chunk[..PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES]; + if let Some(error_body) = + extract_provider_private_stream_error_body(report_context, prefix) + { + return Some(error_body); + } + + let suffix_start = chunk.len() - PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES; + if let Some(error_body) = + extract_provider_private_stream_error_body(report_context, &chunk[suffix_start..]) + { + return Some(error_body); + } + } else if let Some(error_body) = + extract_provider_private_stream_error_body(report_context, chunk) { return Some(error_body); } @@ -350,7 +400,7 @@ fn direct_passthrough_mode() -> DirectPassthroughMode { fn stream_body_buffer_limit_for_record_level(record_level: UsageRequestRecordLevel) -> usize { match record_level { UsageRequestRecordLevel::Basic => BASIC_STREAM_BODY_ANALYSIS_LIMIT_BYTES, - UsageRequestRecordLevel::Full => DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, + UsageRequestRecordLevel::Full => crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES, } } @@ -365,15 +415,15 @@ async fn resolve_stream_body_buffer_limit(state: &AppState) -> usize { .await { Ok(policy) => stream_body_buffer_limit_for_record_level(policy.record_level), - Err(error) => { + Err(_error) => { warn!( event_name = "stream_body_capture_policy_read_failed", log_type = "ops", - error = %error, - fallback = "full", + error_category = "capture_policy_read_failed", + fallback = "basic", "gateway could not resolve stream body capture policy" ); - DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES + BASIC_STREAM_BODY_ANALYSIS_LIMIT_BYTES } } } @@ -442,11 +492,11 @@ async fn record_sync_terminal_usage_with_handoff_after_spawn( ) .await; }); - if let Err(err) = task.await { + if let Err(_err) = task.await { warn!( event_name = "sync_terminal_usage_handoff_failed", log_type = "ops", - error = %err, + error_category = "terminal_usage_handoff_failed", "gateway sync terminal usage handoff task failed" ); } @@ -582,11 +632,24 @@ fn build_stream_body_capture( body: &[u8], truncated: bool, ) -> (Option, Option) { + build_stream_body_capture_with_limit( + body, + truncated, + crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES, + ) +} + +fn build_stream_body_capture_with_limit( + body: &[u8], + truncated: bool, + max_bytes: usize, +) -> (Option, Option) { + let captured = &body[..body.len().min(max_bytes)]; let body_base64 = - (!body.is_empty()).then(|| base64::engine::general_purpose::STANDARD.encode(body)); - let body_state = Some(if truncated { + (!captured.is_empty()).then(|| base64::engine::general_purpose::STANDARD.encode(captured)); + let body_state = Some(if truncated || captured.len() < body.len() { UsageBodyCaptureState::Truncated - } else if body.is_empty() { + } else if captured.is_empty() { UsageBodyCaptureState::None } else { UsageBodyCaptureState::Inline @@ -597,7 +660,7 @@ fn build_stream_body_capture( fn wrap_non_json_binary_stream_error_for_client( plan_kind: &str, headers: &BTreeMap, - error_body: &[u8], + _error_body: &[u8], ) -> Result, GatewayError> { let content_type = headers .get("content-type") @@ -609,7 +672,7 @@ fn wrap_non_json_binary_stream_error_for_client( let body = match plan_kind { GEMINI_FILES_DOWNLOAD_PLAN_KIND => json!({ - "error": String::from_utf8_lossy(error_body).to_string(), + "error": "File download failed", }), OPENAI_VIDEO_CONTENT_PLAN_KIND => json!({ "error": { @@ -617,7 +680,12 @@ fn wrap_non_json_binary_stream_error_for_client( "message": "Video not available", } }), - _ => return Ok(None), + _ => json!({ + "error": { + "type": "upstream_error", + "message": "Upstream request failed", + } + }), }; Ok(Some(body)) } @@ -733,13 +801,13 @@ async fn seed_kiro_simulated_cache_enabled( .is_some_and(|provider| { kiro_simulated_cache_enabled_from_provider_config(provider.config.as_ref()) }), - Err(err) => { + Err(_err) => { warn!( event_name = "kiro_simulated_cache_config_read_failed", log_type = "event", request_id = %plan.request_id, provider_id = %plan.provider_id, - error = ?err, + error_category = "provider_catalog_read_failed", "failed to read Kiro simulated cache provider config; defaulting disabled" ); false @@ -1004,8 +1072,8 @@ fn observe_stream_usage_bytes( remaining = &remaining[line_part_len..]; if buffered.last() == Some(&b'\n') { let line = std::mem::take(buffered); - if let Err(err) = observer.push_line(report_context, line) { - observer.disable_with_error(err.to_string()); + if let Err(_err) = observer.push_line(report_context, line) { + observer.disable_with_error("stream usage parsing failed"); buffered.clear(); return; } @@ -1024,15 +1092,15 @@ fn finalize_stream_usage_observer( if !buffered.is_empty() { let line = std::mem::take(buffered); - if let Err(err) = observer.push_line(report_context, line) { - observer.disable_with_error(err.to_string()); + if let Err(_err) = observer.push_line(report_context, line) { + observer.disable_with_error("stream usage parsing failed"); } } match observer.finish(report_context) { Ok(summary) => summary, - Err(err) => { - observer.disable_with_error(err.to_string()); + Err(_err) => { + observer.disable_with_error("stream usage parsing failed"); observer.latest_summary().cloned() } } @@ -1864,7 +1932,7 @@ impl DirectPassthroughFinalizer { && core.terminal_failure.is_none() } - fn log_terminal_error_event_encode_failed(&self, err: impl std::fmt::Debug) { + fn log_terminal_error_event_encode_failed(&self, _err: impl std::fmt::Debug) { let core = self.core(); warn!( event_name = "direct_passthrough_terminal_error_event_encode_failed", @@ -1872,7 +1940,7 @@ impl DirectPassthroughFinalizer { trace_id = %core.trace_id, request_id = %core.request_id_for_log, candidate_id = ?core.candidate_id.as_deref(), - error = ?err, + error_category = "terminal_error_event_encode_failed", "gateway direct passthrough failed to encode terminal SSE error event" ); } @@ -2027,11 +2095,11 @@ impl DirectPassthroughFinalizer { let task = tokio::spawn(async move { core.finalize(downstream_dropped).await; }); - if let Err(err) = task.await { + if let Err(_err) = task.await { warn!( event_name = "direct_passthrough_terminal_handoff_failed", log_type = "ops", - error = %err, + error_category = "terminal_handoff_failed", "gateway direct passthrough terminal handoff task failed" ); } @@ -2432,7 +2500,7 @@ impl DirectPassthroughFinalizerCore { .await; if should_submit_report { - if let Err(err) = submit_stream_report(&state, usage_payload).await { + if let Err(_err) = submit_stream_report(&state, usage_payload).await { warn!( event_name = "execution_report_submit_failed", log_type = "ops", @@ -2440,7 +2508,7 @@ impl DirectPassthroughFinalizerCore { request_id = %request_id_for_log, candidate_id = ?candidate_id.as_deref(), report_scope = "direct_passthrough_stream", - error = ?err, + error_category = "stream_report_submit_failed", "gateway failed to submit direct passthrough stream execution report" ); } @@ -2680,7 +2748,7 @@ impl DirectPassthroughInlineBodyState { trace_id = %core.trace_id, request_id = %core.request_id_for_log, candidate_id = ?core.candidate_id.as_deref(), - error = %message, + error_category = "upstream_body_read_failed", "gateway ignored direct passthrough teardown error after Anthropic message_stop" ); return; @@ -2693,7 +2761,7 @@ impl DirectPassthroughInlineBodyState { request_id = %core.request_id_for_log, candidate_id = ?core.candidate_id.as_deref(), upstream_bytes = core.provider_stream_bytes, - error = %message, + error_category = "upstream_body_read_failed", "gateway direct passthrough upstream body read failed" ); finalizer.set_terminal_failure(build_stream_transport_failure_report( @@ -2820,6 +2888,7 @@ async fn execute_stream_from_direct_passthrough( candidate_id: _, status_code, mut headers, + upstream_content_length: _, provider_api_format: _, stream_summary_report_context: _, prefetched_body, @@ -3124,7 +3193,7 @@ async fn execute_stream_from_direct_passthrough( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = %message, + error_category = "upstream_body_read_failed", "gateway ignored direct passthrough teardown error after Anthropic message_stop" ); break; @@ -3136,7 +3205,7 @@ async fn execute_stream_from_direct_passthrough( request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), upstream_bytes = provider_stream_bytes, - error = %message, + error_category = "upstream_body_read_failed", "gateway direct passthrough upstream body read failed" ); terminal_failure = Some(build_stream_transport_failure_report( @@ -3335,14 +3404,14 @@ async fn execute_stream_from_direct_passthrough( ) .await; } - Err(err) => { + Err(_err) => { warn!( event_name = "direct_passthrough_terminal_error_event_encode_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 = ?err, + error_category = "terminal_error_event_encode_failed", "gateway direct passthrough failed to encode terminal SSE error event" ); } @@ -3613,7 +3682,7 @@ async fn execute_stream_from_direct_passthrough( .await; if should_submit_report { - if let Err(err) = submit_stream_report(&state_for_report, usage_payload).await { + if let Err(_err) = submit_stream_report(&state_for_report, usage_payload).await { warn!( event_name = "execution_report_submit_failed", log_type = "ops", @@ -3621,7 +3690,7 @@ async fn execute_stream_from_direct_passthrough( request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), report_scope = "direct_passthrough_stream", - error = ?err, + error_category = "stream_report_submit_failed", "gateway failed to submit direct passthrough stream execution report" ); } @@ -3655,17 +3724,30 @@ pub(crate) fn execute_execution_runtime_stream<'a>( report_kind: Option, report_context: Option, ) -> Pin>, GatewayError>> + Send + 'a>> { - Box::pin(execute_execution_runtime_stream_inner( - state, - plan, - trace_id, - decision, - plan_kind, - report_kind, - report_context, - None, - None, - )) + Box::pin(async move { + let mut cancellation_guard = AttemptCancellationGuard::disarmed( + state, + STREAM_ATTEMPT_CANCELLED_ERROR_TYPE, + STREAM_ATTEMPT_CANCELLED_ERROR_MESSAGE, + ); + let result = execute_execution_runtime_stream_inner( + state, + plan, + trace_id, + decision, + plan_kind, + report_kind, + report_context, + None, + None, + &mut cancellation_guard, + ) + .await; + // The attempt reached its own terminal path, or handed settlement to the + // stream finalizer that now lives in the response body. + cancellation_guard.disarm(); + result + }) } #[allow(clippy::too_many_arguments)] @@ -3687,7 +3769,12 @@ pub(crate) fn execute_execution_runtime_stream_with_retry_scope<'a>( Box::pin(async move { let mut retry_scope = AiAttemptRetryScope::Candidate; let mut fallback_response = None; - let response = execute_execution_runtime_stream_inner( + let mut cancellation_guard = AttemptCancellationGuard::disarmed( + state, + STREAM_ATTEMPT_CANCELLED_ERROR_TYPE, + STREAM_ATTEMPT_CANCELLED_ERROR_MESSAGE, + ); + let result = execute_execution_runtime_stream_inner( state, plan, trace_id, @@ -3697,8 +3784,13 @@ pub(crate) fn execute_execution_runtime_stream_with_retry_scope<'a>( report_context, Some(&mut retry_scope), Some(&mut fallback_response), + &mut cancellation_guard, ) - .await?; + .await; + // The attempt reached its own terminal path, or handed settlement to the + // stream finalizer that now lives in the response body. + cancellation_guard.disarm(); + let response = result?; Ok(match response { Some(response) => AiAttemptExecutionOutcome::Responded(response), None => AiAttemptExecutionOutcome::Retry { @@ -3744,6 +3836,7 @@ async fn maybe_build_stream_transport_error_stop_response( .map(Some) } +#[allow(clippy::too_many_arguments)] // internal function, grouping would add unnecessary indirection async fn execute_execution_runtime_stream_inner( state: &AppState, mut plan: ExecutionPlan, @@ -3754,6 +3847,7 @@ async fn execute_execution_runtime_stream_inner( mut report_context: Option, mut retry_scope_out: Option<&mut AiAttemptRetryScope>, mut retry_fallback_out: Option<&mut Option>>, + cancellation_guard: &mut AttemptCancellationGuard, ) -> Result>, GatewayError> { let stream_started_at = Instant::now(); let mut stage_trace = RequestStageTrace::from_env(); @@ -3778,6 +3872,11 @@ async fn execute_execution_runtime_stream_inner( match acquire_provider_pool_execution_guard(state, &plan).await? { ProviderPoolInFlightAdmission::Acquired(guard) => guard, ProviderPoolInFlightAdmission::Saturated { limit } => { + record_local_runtime_candidate_skip_reason( + state, + trace_id, + "provider_key_concurrency_limit_reached", + ); if let Some(retry_scope) = retry_scope_out.as_deref_mut() { *retry_scope = AiAttemptRetryScope::Candidate; } @@ -3832,6 +3931,16 @@ async fn execute_execution_runtime_stream_inner( ) .await; } + // From here the attempt owns non-terminal rows, and everything that could + // settle them runs inside the downstream request future. Arm the guard so a + // client disconnect before the stream finalizer exists still settles them. + cancellation_guard.arm( + &plan, + report_context.as_ref(), + request_candidate_status_snapshot.as_ref(), + candidate_started_unix_secs, + stream_started_at, + ); let plan_request_id_for_log = short_request_id(plan.request_id.as_str()); let provider_name = plan .provider_name @@ -3868,8 +3977,8 @@ async fn execute_execution_runtime_stream_inner( .await; } Ok(None) => {} - Err(err) => { - let transport_error_message = err.to_string(); + Err(_err) => { + let transport_error_message = "Grok stream execution unavailable".to_string(); info!( event_name = "grok_execution_unavailable", log_type = "ops", @@ -3881,7 +3990,7 @@ async fn execute_execution_runtime_stream_inner( key_id = %key_id, model_name = model_name.as_str(), candidate_index = candidate_index.as_str(), - error = %err, + error_category = "grok_execution_unavailable", "gateway Grok stream execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -3941,8 +4050,8 @@ async fn execute_execution_runtime_stream_inner( .await; } Ok(None) => {} - Err(err) => { - let transport_error_message = err.to_string(); + Err(_err) => { + let transport_error_message = "Windsurf stream execution unavailable".to_string(); info!( event_name = "windsurf_native_execution_unavailable", log_type = "ops", @@ -3954,7 +4063,7 @@ async fn execute_execution_runtime_stream_inner( key_id = %key_id, model_name = model_name.as_str(), candidate_index = candidate_index.as_str(), - error = %err, + error_category = "windsurf_execution_unavailable", "gateway native Windsurf stream execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4014,8 +4123,8 @@ async fn execute_execution_runtime_stream_inner( .await; } Ok(None) => {} - Err(err) => { - let transport_error_message = err.to_string(); + Err(_err) => { + let transport_error_message = "Kiro web search execution unavailable".to_string(); info!( event_name = "kiro_web_search_mcp_unavailable", log_type = "ops", @@ -4027,7 +4136,7 @@ async fn execute_execution_runtime_stream_inner( key_id = %key_id, model_name = model_name.as_str(), candidate_index = candidate_index.as_str(), - error = %err, + error_category = "kiro_web_search_unavailable", "gateway Kiro web_search MCP execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4087,8 +4196,8 @@ async fn execute_execution_runtime_stream_inner( .await; } Ok(None) => {} - Err(err) => { - let transport_error_message = err.to_string(); + Err(_err) => { + let transport_error_message = "ChatGPT-Web image execution unavailable".to_string(); info!( event_name = "chatgpt_web_image_execution_unavailable", log_type = "ops", @@ -4100,7 +4209,7 @@ async fn execute_execution_runtime_stream_inner( key_id = %key_id, model_name = model_name.as_str(), candidate_index = candidate_index.as_str(), - error = %err, + error_category = "chatgpt_web_image_execution_unavailable", "gateway ChatGPT-Web image stream execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4161,8 +4270,8 @@ async fn execute_execution_runtime_stream_inner( } return Err(err); } - Err(InProcessStreamExecutionError::Transport(err)) => { - let transport_error_message = err.to_string(); + Err(InProcessStreamExecutionError::Transport(_err)) => { + let transport_error_message = "Execution runtime unavailable".to_string(); info!( event_name = "stream_execution_runtime_unavailable", log_type = "ops", @@ -4174,7 +4283,7 @@ async fn execute_execution_runtime_stream_inner( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error_category = "execution_runtime_unavailable", "gateway in-process stream execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4301,8 +4410,8 @@ async fn execute_execution_runtime_stream_inner( } return Err(err); } - Err(InProcessStreamExecutionError::Transport(err)) => { - let transport_error_message = err.to_string(); + Err(InProcessStreamExecutionError::Transport(_err)) => { + let transport_error_message = "Execution runtime unavailable".to_string(); info!( event_name = "stream_execution_runtime_unavailable", log_type = "ops", @@ -4314,7 +4423,7 @@ async fn execute_execution_runtime_stream_inner( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error_category = "execution_runtime_unavailable", "gateway in-process stream execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4429,15 +4538,15 @@ async fn execute_execution_runtime_stream_inner( .await { Ok(response) => response, - Err(err) => { - let transport_error_message = format!("{err:?}"); + Err(_err) => { + let transport_error_message = "Remote execution runtime unavailable".to_string(); warn!( event_name = "stream_execution_runtime_remote_unavailable", log_type = "ops", trace_id = %trace_id, request_id = %plan_request_id_for_log, candidate_id = ?plan.candidate_id, - error = ?err, + error_category = "execution_runtime_unavailable", "gateway remote execution runtime stream unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -4537,13 +4646,26 @@ async fn execute_execution_runtime_stream_inner( fn decode_stream_data_chunk( chunk_b64: Option<&str>, text: Option<&str>, +) -> Result, GatewayError> { + decode_stream_data_chunk_with_limit(chunk_b64, text, MAX_EXECUTION_STREAM_DATA_CHUNK_BYTES) +} + +fn decode_stream_data_chunk_with_limit( + chunk_b64: Option<&str>, + text: Option<&str>, + max_bytes: usize, ) -> Result, GatewayError> { if let Some(chunk_b64) = chunk_b64 { - return base64::engine::general_purpose::STANDARD - .decode(chunk_b64) + return decode_base64_body_with_limit(chunk_b64, max_bytes) .map_err(|err| GatewayError::Internal(err.to_string())); } - Ok(text.unwrap_or_default().as_bytes().to_vec()) + let text = text.unwrap_or_default().as_bytes(); + if text.len() > max_bytes { + return Err(GatewayError::Internal(format!( + "execution runtime stream data chunk exceeds {max_bytes} bytes" + ))); + } + Ok(text.to_vec()) } fn response_headers_indicate_sse(headers: &BTreeMap) -> bool { @@ -5387,6 +5509,10 @@ fn serialized_stream_frame_len(frame: &StreamFrame) -> usize { serde_json::to_vec(frame).map_or(usize::MAX, |encoded| encoded.len()) } +fn execution_stream_frame_codec() -> LinesCodec { + LinesCodec::new_with_max_length(crate::execution_runtime::MAX_EXECUTION_STREAM_FRAME_LINE_BYTES) +} + fn should_refresh_stream_usage_telemetry( previous: Option<&ExecutionTelemetry>, next: &ExecutionTelemetry, @@ -5673,7 +5799,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .map(|value| value.to_string()) .unwrap_or_else(|| "-".to_string()); let reader = PostStopLimitedStreamReader::new(frame_stream, PostStopFrameReadBudget::new()); - let mut lines = FramedRead::new(reader, LinesCodec::new()); + let mut lines = FramedRead::new(reader, execution_stream_frame_codec()); let first_frame_started_at = Instant::now(); let first_frame = read_next_frame(&mut lines).await?.ok_or_else(|| { @@ -6112,14 +6238,33 @@ async fn execute_stream_from_frame_stream_with_retry_scope( candidate_id, )?)); } - return Ok(Some(attach_control_metadata_headers( + let response = if (300..400).contains(&status_code) { + build_client_response_from_parts_with_mutator( + client_status_code, + &client_response_headers, + Body::from(client_error_body), + trace_id, + Some(decision), + |headers| { + headers.insert( + http::HeaderName::from_static("x-aether-upstream-status"), + http::HeaderValue::from_str(&status_code.to_string()) + .map_err(|error| GatewayError::Internal(error.to_string()))?, + ); + Ok(()) + }, + )? + } else { build_client_response_from_parts( client_status_code, &client_response_headers, Body::from(client_error_body), trace_id, Some(decision), - )?, + )? + }; + return Ok(Some(attach_control_metadata_headers( + response, Some(request_id), candidate_id, )?)); @@ -6167,7 +6312,10 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } let prefetch_for_cyber_failover = is_openai_responses_family_format(plan.provider_api_format.as_str()) - && cyber_continue_failover_enabled(state).await; + && crate::orchestration::routing_execution_policy_from_report_context( + report_context.as_ref(), + ) + .is_some_and(|policy| policy.cyber_continue_failover); let stream_commit_policy = StreamCommitPolicy::for_response( direct_stream_finalize_kind.is_some(), upstream_content_type, @@ -6490,7 +6638,9 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } } - let inspection = if stream_commit_policy.is_native_anthropic() { + let inspection = if stream_commit_policy.is_native_anthropic() + || stream_commit_policy.is_gemini() + { StreamPrefetchInspection::NeedMore } else { inspect_prefetched_stream_body( @@ -6569,8 +6719,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( Ok(Some(outcome)) => { if let Some(record) = outcome.response_history_record { crate::ai_serving::persist_response_history_record( - state.runtime_state(), - record, + state, record, ) .await; } @@ -6749,7 +6898,9 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id, request_id, candidate_id = ?candidate_id, - error = %error.message, + error_kind = ?error.kind, + error_phase = ?error.phase, + upstream_status = ?error.upstream_status, "execution runtime stream emitted error frame during prefetch" ); return handle_prefetch_stream_failure( @@ -6782,7 +6933,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .as_mut() .and_then(|rewriter| rewriter.take_response_history_record()) { - crate::ai_serving::persist_response_history_record(state.runtime_state(), record).await; + crate::ai_serving::persist_response_history_record(state, record).await; true } else { false @@ -7058,7 +7209,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_normalization_restore_failed", "gateway failed to restore private stream normalization state after prefetch" ); terminal_failure = Some(build_stream_failure_report( @@ -7111,7 +7262,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_rewrite_restore_failed", "gateway failed to restore local stream rewrite state after prefetch" ); terminal_failure = Some(build_stream_failure_report( @@ -7177,7 +7328,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_frame_decode_failed", "gateway ignored execution runtime teardown error after Anthropic message_stop" ); break; @@ -7188,7 +7339,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_frame_decode_failed", "gateway failed to decode execution runtime stream frame" ); terminal_failure = Some(build_stream_failure_report( @@ -7276,7 +7427,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_chunk_decode_failed", "gateway failed to decode execution runtime chunk" ); terminal_failure = Some(build_stream_failure_report( @@ -7329,7 +7480,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_chunk_normalize_failed", "gateway failed to normalize execution runtime stream chunk" ); terminal_failure = Some(build_stream_failure_report( @@ -7367,7 +7518,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_chunk_rewrite_failed", "gateway failed to rewrite execution runtime stream chunk" ); terminal_failure = Some(build_stream_failure_report( @@ -7388,7 +7539,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .and_then(|rewriter| rewriter.take_response_history_record()) { crate::ai_serving::persist_response_history_record( - state_for_report.runtime_state(), + &state_for_report, record, ) .await; @@ -7502,7 +7653,9 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = %error.message, + error_kind = ?error.kind, + error_phase = ?error.phase, + upstream_status = ?error.upstream_status, "gateway ignored execution runtime error frame after Anthropic message_stop" ); continue; @@ -7513,7 +7666,9 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = %error.message, + error_kind = ?error.kind, + error_phase = ?error.phase, + upstream_status = ?error.upstream_status, "execution runtime stream emitted error frame" ); terminal_failure = Some(build_stream_failure_from_execution_error(&error)); @@ -7571,7 +7726,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_flush_rewrite_failed", "gateway failed to rewrite normalized private stream chunk during flush" ); let failure = build_stream_failure_report( @@ -7592,7 +7747,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .and_then(|rewriter| rewriter.take_response_history_record()) { crate::ai_serving::persist_response_history_record( - state_for_report.runtime_state(), + &state_for_report, record, ) .await; @@ -7656,7 +7811,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_normalization_flush_failed", "gateway failed to flush private stream normalization" ); terminal_failure.get_or_insert_with(|| { @@ -7673,11 +7828,8 @@ async fn execute_stream_from_frame_stream_with_retry_scope( if let Some(rewriter) = local_stream_rewriter.as_mut() { let finish_result = rewriter.finish(); if let Some(record) = rewriter.take_response_history_record() { - crate::ai_serving::persist_response_history_record( - state_for_report.runtime_state(), - record, - ) - .await; + crate::ai_serving::persist_response_history_record(&state_for_report, record) + .await; } match finish_result { Ok(flushed_chunk) if !flushed_chunk.is_empty() => { @@ -7722,7 +7874,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( trace_id = %trace_id_owned, request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), - error = ?err, + error_category = "stream_rewrite_flush_failed", "gateway failed to flush local stream rewrite" ); terminal_failure.get_or_insert_with(|| { @@ -7742,11 +7894,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .as_mut() .and_then(|rewriter| rewriter.take_response_history_record()) { - crate::ai_serving::persist_response_history_record( - state_for_report.runtime_state(), - record, - ) - .await; + crate::ai_serving::persist_response_history_record(&state_for_report, record).await; } } @@ -7799,14 +7947,14 @@ async fn execute_stream_from_frame_stream_with_retry_scope( ); } } - Err(err) => { + Err(_err) => { warn!( event_name = "stream_execution_terminal_error_event_encode_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 = ?err, + error_category = "terminal_error_event_encode_failed", "gateway failed to encode terminal SSE error event" ); } @@ -8095,7 +8243,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .await; if should_submit_report { - if let Err(err) = submit_stream_report(&state_for_report, usage_payload).await { + if let Err(_err) = submit_stream_report(&state_for_report, usage_payload).await { warn!( event_name = "execution_report_submit_failed", log_type = "ops", @@ -8103,7 +8251,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( request_id = %request_id_for_report_log, candidate_id = ?candidate_id_for_report.as_deref(), report_scope = "stream", - error = ?err, + error_category = "stream_report_submit_failed", "gateway failed to submit stream execution report" ); } @@ -8164,6 +8312,7 @@ mod tests { use std::time::{Duration, Instant}; use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope}; + use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; use aether_contracts::{ ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTelemetry, ExecutionTimeouts, RequestBody, @@ -8171,6 +8320,7 @@ mod tests { }; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data_contracts::repository::candidates::{ PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository, @@ -8215,9 +8365,10 @@ mod tests { ensure_stream_terminal_summary_for_missing_observed_finish, execute_execution_runtime_stream, execute_in_process_stream_with_oauth_retry, execute_stream_from_frame_stream, execute_stream_from_frame_stream_with_retry_scope, - maybe_apply_kiro_prompt_cache_usage_to_stream_summary, merge_stream_terminal_summary, - normalize_declared_stream_response_headers, parse_direct_passthrough_mode, - prefetch_direct_stream_error_body, prefetched_openai_responses_body_has_output_boundary, + execution_stream_frame_codec, maybe_apply_kiro_prompt_cache_usage_to_stream_summary, + merge_stream_terminal_summary, normalize_declared_stream_response_headers, + parse_direct_passthrough_mode, prefetch_direct_stream_error_body, + prefetched_openai_responses_body_has_output_boundary, record_sync_terminal_usage_with_handoff, record_sync_terminal_usage_with_handoff_after_spawn, resolve_provider_stream_error_status_code, select_direct_anthropic_prefetch_wait, @@ -8227,11 +8378,13 @@ mod tests { stream_terminal_summary_missing_observed_finish, stream_terminal_summary_missing_observed_finish_with_requirement, stream_terminal_summary_represents_failure_with_requirement, - ClientVisibleStreamCompletionTracker, DirectPassthroughFinalizer, - DirectPassthroughFinalizerCore, DirectPassthroughInlineBodyState, DirectPassthroughMode, - PostStopFrameReadBudget, PostStopLimitedStreamReader, ProviderStreamErrorInspection, + wrap_non_json_binary_stream_error_for_client, ClientVisibleStreamCompletionTracker, + DirectPassthroughFinalizer, DirectPassthroughFinalizerCore, + DirectPassthroughInlineBodyState, DirectPassthroughMode, PostStopFrameReadBudget, + PostStopLimitedStreamReader, ProviderStreamErrorInspection, ANTHROPIC_POST_STOP_DRAIN_MAX_BYTES, GEMINI_FILES_DOWNLOAD_PLAN_KIND, - OPENAI_CHAT_STREAM_PLAN_KIND, POST_STOP_MAX_EMPTY_CHUNKS_PER_POLL, + OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND, + POST_STOP_MAX_EMPTY_CHUNKS_PER_POLL, PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES, }; use crate::control::GatewayControlDecision; use crate::stage_metrics::RequestStageTrace; @@ -8256,6 +8409,19 @@ mod tests { plan: &ExecutionPlan, provider_config: Option, ) -> InMemoryProviderCatalogReadRepository { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key( + &plan.provider_id, + &plan.key_id, + "plain-upstream-key", + ) + .expect("api key should encrypt"); let provider_type = plan.provider_name.as_deref().unwrap_or("custom"); let provider = StoredProviderCatalogProvider::new( plan.provider_id.clone(), @@ -8306,7 +8472,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!([plan.provider_api_format.clone()])), - "plain-upstream-key".to_string(), + encrypted_api_key, None, None, Some(json!({ "openai:chat": 1 })), @@ -8320,6 +8486,22 @@ mod tests { InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key]) } + #[test] + fn non_json_upstream_error_body_is_not_projected_to_clients() { + let secret = b"Bearer upstream-secret https://user:password@example.test/private"; + let body = wrap_non_json_binary_stream_error_for_client( + "openai_chat_stream", + &BTreeMap::from([("content-type".to_string(), "text/plain".to_string())]), + secret, + ) + .expect("error body projection should succeed") + .expect("non-JSON errors should receive a client projection"); + + assert_eq!(body["error"]["message"], "Upstream request failed"); + assert!(!body.to_string().contains("upstream-secret")); + assert!(!body.to_string().contains("password")); + } + fn provider_catalog_for_stream_auth_plan( plan: &ExecutionPlan, provider_type: &str, @@ -8490,16 +8672,8 @@ mod tests { let provider_catalog = provider_catalog_for_plan(&plan, None); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); - let data_state = if continue_failover { - data_state.with_system_config_values_for_tests([( - crate::orchestration::CYBER_CONTINUE_FAILOVER_CONFIG_KEY.to_string(), - json!(true), - )]) - } else { - data_state - }; let state = AppState::new() .expect("app state should build") .with_data_state_for_tests(data_state); @@ -8548,7 +8722,10 @@ mod tests { "candidate_index": 0, "retry_index": 0, "provider_api_format": "openai:responses", - "client_api_format": "openai:responses" + "client_api_format": "openai:responses", + "routing_execution_policy": { + "cyber_continue_failover": continue_failover + } })), crate::clock::current_unix_ms(), Instant::now(), @@ -8580,7 +8757,7 @@ mod tests { let provider_catalog = provider_catalog_for_plan(&plan, provider_config); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("app state should build") @@ -8673,7 +8850,7 @@ mod tests { ); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("app state should build") @@ -8776,6 +8953,36 @@ mod tests { } } + fn antigravity_gemini_stream_plan(request_id: &str) -> ExecutionPlan { + ExecutionPlan { + request_id: request_id.to_string(), + candidate_id: Some(format!("candidate-{request_id}")), + provider_name: Some("antigravity".to_string()), + provider_id: format!("provider-{request_id}"), + endpoint_id: format!("endpoint-{request_id}"), + key_id: format!("key-{request_id}"), + method: "POST".to_string(), + url: "https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent".to_string(), + headers: BTreeMap::from([ + ("content-type".to_string(), "application/json".to_string()), + ("accept".to_string(), "text/event-stream".to_string()), + ]), + content_type: Some("application/json".to_string()), + content_encoding: None, + body: RequestBody::from_json(json!({ + "model": "gemini-3.7-flash-tiered", + "contents": [{"role": "user", "parts": [{"text": "validate"}]}] + })), + stream: true, + client_api_format: "openai:responses".to_string(), + provider_api_format: "gemini:generate_content".to_string(), + model_name: Some("gemini-3.7-flash-tiered".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + } + } + struct StreamDropFlag(Arc); impl Drop for StreamDropFlag { @@ -8875,7 +9082,7 @@ mod tests { let provider_catalog = provider_catalog_for_plan(&plan, None); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( Arc::new(provider_catalog), - "development-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("app state should build") @@ -10349,12 +10556,37 @@ mod tests { } #[test] - fn stream_capture_policy_keeps_full_unbounded_and_caps_basic_analysis_buffer() { + fn stream_capture_encoding_defensively_caps_an_oversized_slice() { + let (body, state) = super::build_stream_body_capture_with_limit(b"abcdef", false, 3); + let decoded = base64::engine::general_purpose::STANDARD + .decode(body.expect("bounded capture should be encoded")) + .expect("capture should be valid base64"); + + assert_eq!(decoded, b"abc"); + assert_eq!(state, Some(UsageBodyCaptureState::Truncated)); + } + + #[test] + fn execution_stream_data_chunk_decode_is_bounded_before_allocation() { + assert_eq!( + super::decode_stream_data_chunk_with_limit(Some("YWJj"), None, 3) + .expect("three decoded bytes"), + b"abc" + ); + assert!(super::decode_stream_data_chunk_with_limit(Some("YWJjZA=="), None, 3).is_err()); + assert!(super::decode_stream_data_chunk_with_limit(None, Some("abcd"), 3).is_err()); + } + + #[test] + fn stream_capture_policy_hard_caps_full_and_basic_analysis_buffers() { let oversized_chunk = vec![b'x'; super::BASIC_STREAM_BODY_ANALYSIS_LIMIT_BYTES + 1]; let full_limit = super::stream_body_buffer_limit_for_record_level(UsageRequestRecordLevel::Full); - assert_eq!(full_limit, usize::MAX); + assert_eq!( + full_limit, + crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES + ); let mut full_buffer = Vec::new(); let mut full_truncated = false; super::append_stream_capture_bytes( @@ -10466,6 +10698,54 @@ mod tests { assert_eq!(detected.pointer("/error/param"), Some(&json!("input"))); } + #[test] + fn provider_error_inspection_bounds_oversized_chunks_and_keeps_boundary_detection() { + let error_event = concat!( + "event: response.failed\n", + "data: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"code\":\"cyber_policy_violation\"}}}\n\n", + ) + .as_bytes(); + + let mut prefix_chunk = vec![b'x'; PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES + 64]; + prefix_chunk[..error_event.len()].copy_from_slice(error_event); + let mut inspection = ProviderStreamErrorInspection::default(); + let detected = inspection + .observe(None, &prefix_chunk) + .expect("error at the bounded chunk prefix should be detected"); + assert_eq!( + detected.pointer("/error/code"), + Some(&json!("cyber_policy_violation")) + ); + + let mut suffix_chunk = vec![b'x'; PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES + 64]; + let suffix_start = suffix_chunk.len() - error_event.len(); + suffix_chunk[suffix_start..].copy_from_slice(error_event); + let mut inspection = ProviderStreamErrorInspection::default(); + let detected = inspection + .observe(None, &suffix_chunk) + .expect("error at the bounded chunk suffix should be detected"); + assert_eq!( + detected.pointer("/error/code"), + Some(&json!("cyber_policy_violation")) + ); + + // The JSON payload is split across chunks. The previous rolling tail + // must still be combined with the prefix of the oversized chunk. + let split = b"event: response.failed\ndata: {".len(); + let mut inspection = ProviderStreamErrorInspection::default(); + assert!(inspection.observe(None, &error_event[..split]).is_none()); + let mut continuation = vec![b'x'; PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES + 64]; + let continuation_len = error_event.len() - split; + continuation[..continuation_len].copy_from_slice(&error_event[split..]); + let detected = inspection + .observe(None, &continuation) + .expect("error split across an oversized chunk boundary should be detected"); + assert_eq!( + detected.pointer("/error/code"), + Some(&json!("cyber_policy_violation")) + ); + } + #[tokio::test] async fn prefetched_codex_cyber_policy_violation_stops_failover_by_default() { let response = execute_prefetched_codex_cyber_policy_failure(false) @@ -10476,7 +10756,7 @@ mod tests { } #[tokio::test] - async fn prefetched_codex_cyber_policy_violation_retries_when_system_setting_is_enabled() { + async fn prefetched_codex_cyber_policy_violation_retries_when_routing_strategy_is_enabled() { assert!( execute_prefetched_codex_cyber_policy_failure(true) .await @@ -10524,6 +10804,117 @@ mod tests { assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); } + #[tokio::test] + async fn malformed_antigravity_function_call_streams_thought_then_fails_in_band() { + let request_id = "req-antigravity-malformed-function-call"; + let plan = antigravity_gemini_stream_plan(request_id); + let provider_catalog = provider_catalog_for_plan( + &plan, + Some(json!({ + "failover_rules": { + "continue_status_codes": [502] + } + })), + ); + let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( + Arc::new(provider_catalog), + "development-key", + ); + let state = AppState::new() + .expect("app state should build") + .with_data_state_for_tests(data_state); + let frame_stream = stream! { + yield Ok::(ndjson_frame(StreamFrame { + frame_type: StreamFrameType::Headers, + payload: StreamFramePayload::Headers { + status_code: 200, + headers: BTreeMap::from([( + "content-type".to_string(), + "text/event-stream".to_string(), + )]), + response_observation: None, + }, + })); + for chunk in [ + r#"data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"thought":true,"text":"Validating the document."}]} }],"modelVersion":"gemini-3.7-flash-tiered"}} + +"#, + r#"data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"thoughtSignature":"signature","text":""}]},"finishReason":"MALFORMED_FUNCTION_CALL","finishMessage":"Malformed function call: Function call is empty - no input to parse."}],"modelVersion":"gemini-3.7-flash-tiered"}} + +"#, + ] { + yield Ok::(ndjson_frame(StreamFrame { + frame_type: StreamFrameType::Data, + payload: StreamFramePayload::Data { + chunk_b64: None, + text: Some(chunk.to_string()), + }, + })); + } + yield Ok::(ndjson_frame(StreamFrame::eof())); + } + .boxed(); + let mut retry_scope = AiAttemptRetryScope::Provider; + + let response = execute_stream_from_frame_stream_with_retry_scope( + &state, + plan, + "trace-antigravity-malformed-function-call", + &test_decision(), + OPENAI_RESPONSES_STREAM_PLAN_KIND, + Some("openai_responses_stream_success".to_string()), + Some(json!({ + "request_id": request_id, + "candidate_id": format!("candidate-{request_id}"), + "candidate_index": 0, + "retry_index": 0, + "provider_api_format": "gemini:generate_content", + "client_api_format": "openai:responses", + "needs_conversion": true + })), + crate::clock::current_unix_ms(), + Instant::now(), + RequestStageTrace::from_env(), + true, + frame_stream, + false, + None, + Some(&mut retry_scope), + None, + None, + ) + .await + .expect("malformed Antigravity stream should return a client stream") + .expect("the first reasoning delta should commit the selected candidate"); + + assert_eq!(response.status(), StatusCode::OK); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .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}" + ); + assert!( + body.contains("\"delta\":\"Validating the document.\""), + "{body}" + ); + assert!(body.contains("event: response.failed\n"), "{body}"); + assert!( + body.contains("\"code\":\"MALFORMED_FUNCTION_CALL\""), + "{body}" + ); + assert!( + body.contains( + "\"message\":\"Malformed function call: Function call is empty - no input to parse.\"" + ), + "{body}" + ); + assert!(!body.contains("unsupported_finish_reason"), "{body}"); + assert_eq!(retry_scope, AiAttemptRetryScope::Provider); + } + fn tunnel_proxy_snapshot(base_url: String) -> aether_contracts::ProxySnapshot { aether_contracts::ProxySnapshot { enabled: Some(true), @@ -10535,6 +10926,66 @@ mod tests { } } + const LOCAL_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const LOCAL_TUNNEL_TEST_GENERATION: &str = "stream-test-generation-1"; + + fn authenticated_local_tunnel_test_state() -> AppState { + let node = StoredProxyNode::new( + "node-1".to_string(), + "Node 1".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "online".to_string(), + 30, + 1, + 0, + 0, + 0, + 0, + true, + true, + 1, + ) + .expect("tunnel node should build") + .with_runtime_fields( + None, + None, + None, + None, + Some(json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": LOCAL_TUNNEL_TEST_PSK, + } + })), + None, + None, + None, + None, + None, + None, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()); + let data = crate::data::GatewayDataState::with_proxy_node_repository_for_tests(Arc::new( + InMemoryProxyNodeRepository::seed([node]), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + AppState::new() + .expect("app state should build") + .with_data_state_for_tests(data) + } + + async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, + ) -> Message { + tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) + } + fn connect_json_frame(flags: u8, payload: &[u8]) -> Vec { let mut out = Vec::with_capacity(5 + payload.len()); out.push(flags); @@ -10549,6 +11000,14 @@ mod tests { Bytes::from(bytes) } + #[test] + fn execution_stream_frame_codec_has_a_bounded_line_length() { + assert_eq!( + execution_stream_frame_codec().max_length(), + crate::execution_runtime::MAX_EXECUTION_STREAM_FRAME_LINE_BYTES + ); + } + #[test] fn post_stop_reader_yields_after_bounded_empty_chunks() { let polls = Arc::new(AtomicUsize::new(0)); @@ -10597,8 +11056,7 @@ mod tests { let frame_stream = futures_util::stream::iter([Ok::(Bytes::from(combined))]); let reader = PostStopLimitedStreamReader::new(frame_stream, PostStopFrameReadBudget::new()); - let mut lines = - tokio_util::codec::FramedRead::new(reader, tokio_util::codec::LinesCodec::new()); + let mut lines = tokio_util::codec::FramedRead::new(reader, execution_stream_frame_codec()); super::read_next_frame(&mut lines) .await @@ -10632,8 +11090,7 @@ mod tests { fn post_stop_activation_releases_over_limit_framed_buffer() { let frame_stream = futures_util::stream::empty::>(); let reader = PostStopLimitedStreamReader::new(frame_stream, PostStopFrameReadBudget::new()); - let mut lines = - tokio_util::codec::FramedRead::new(reader, tokio_util::codec::LinesCodec::new()); + let mut lines = tokio_util::codec::FramedRead::new(reader, execution_stream_frame_codec()); lines .read_buffer_mut() .resize(ANTHROPIC_POST_STOP_DRAIN_MAX_BYTES + 1, b'x'); @@ -11531,7 +11988,8 @@ mod tests { assert!(event.starts_with("event: error\ndata: ")); assert!(event.contains("\"type\":\"error\"")); assert!(event.contains("\"type\":\"api_error\"")); - assert!(event.contains("upstream disconnected")); + assert!(event.contains("Upstream response stream failed")); + assert!(!event.contains("upstream disconnected")); assert!(!event.contains("[DONE]")); } @@ -11655,7 +12113,8 @@ mod tests { assert!(body.starts_with(message_start)); assert!(body.contains("event: error\ndata: {\"type\":\"error\"")); assert!(body.contains("\"type\":\"api_error\"")); - assert!(body.contains(original_error)); + assert!(body.contains("Execution runtime stream protocol failed")); + assert!(!body.contains(original_error)); assert!(!body.contains("[DONE]")); } @@ -11723,7 +12182,7 @@ mod tests { &usage_repository, )) .with_provider_catalog_reader(Arc::new(provider_catalog)) - .with_encryption_key_for_tests("development-key"), + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_usage_runtime_for_tests(UsageRuntimeConfig { enabled: true, @@ -11848,7 +12307,7 @@ mod tests { &usage_repository, )) .with_provider_catalog_reader(Arc::new(provider_catalog)) - .with_encryption_key_for_tests("development-key"), + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_usage_runtime_for_tests(UsageRuntimeConfig { enabled: true, @@ -12768,7 +13227,8 @@ mod tests { assert!(!body_text.contains("event: message_stop")); assert!(!body_text.contains("\"stop_reason\":\"tool_use\"")); assert!(body_text.contains("\"error\"")); - assert!(body_text.contains("unexpected internal error encountered")); + assert!(body_text.contains("Execution runtime stream failed")); + assert!(!body_text.contains("unexpected internal error encountered")); assert!(body_text.contains("data: [DONE]")); let candidates = tokio::time::timeout(Duration::from_secs(1), async { @@ -12925,7 +13385,7 @@ mod tests { Arc::clone(&usage_repository), ) .with_provider_catalog_reader(Arc::new(provider_catalog_stop_429_for_plan(&plan))) - .with_encryption_key_for_tests("development-key"), + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let trailer_error = connect_json_frame( 2, @@ -13064,7 +13524,7 @@ mod tests { Arc::clone(&usage_repository), ) .with_provider_catalog_reader(Arc::new(provider_catalog_stop_429_for_plan(&plan))) - .with_encryption_key_for_tests("development-key"), + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let connect_error = connect_json_frame( 2, @@ -13140,15 +13600,10 @@ mod tests { .await .expect("usage should be written"); assert_eq!(record.status_code, Some(429)); - assert_eq!( - record - .response_body - .as_ref() - .and_then(|body| body.get("error")) - .and_then(|error| error.get("code")), - Some(&json!("resource_exhausted")) - ); + assert!(record.response_body.is_none()); assert!(record.response_body_ref.is_none()); + assert!(record.client_response_body.is_none()); + assert!(record.client_response_body_ref.is_none()); } #[tokio::test] @@ -14080,7 +14535,11 @@ mod tests { crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests( Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), - ), + ) + .with_system_config_values_for_tests([( + "request_record_level".to_string(), + json!("full"), + )]), ) .with_usage_runtime_for_tests(UsageRuntimeConfig { enabled: true, @@ -14201,35 +14660,16 @@ mod tests { assert_eq!(usage.status_code, Some(302)); assert_eq!(usage.error_category.as_deref(), Some("redirect")); - assert!(usage - .error_message - .as_deref() - .is_some_and(|value| value.contains("non-success status 302"))); - assert_eq!( - usage - .client_response_headers - .as_ref() - .and_then(|headers| headers.get("x-aether-upstream-status")), - Some(&json!("302")) - ); - assert_eq!( - usage - .response_headers - .as_ref() - .and_then(|headers| headers.get("location")), - Some(&json!("/")) - ); + assert!(usage.error_message.is_none()); + // HTTP capture is intentionally disabled at the persistence boundary. Keep the + // protocol facts above, but do not turn provider/client headers into an audit store. + assert!(usage.client_response_headers.is_none()); + assert!(usage.response_headers.is_none()); assert!( usage.response_body.is_none(), "upstream redirect did not include a body" ); - assert_eq!( - usage - .client_response_body - .as_ref() - .and_then(|body| body.pointer("/error/upstream_status")), - Some(&json!(302)) - ); + assert!(usage.client_response_body.is_none()); let candidates = request_candidate_repository .list_by_request_id("req-remote-runtime-stream-redirect") .await @@ -14242,10 +14682,9 @@ mod tests { candidate_extra["upstream_response"]["status_code"], json!(302) ); - assert_eq!( - candidate_extra["upstream_response"]["headers"]["location"], - json!("/") - ); + assert!(candidate_extra["upstream_response"] + .get("headers") + .is_none()); assert!(candidate_extra["upstream_response"].get("body").is_none()); assert!(candidate_extra.get("client_response").is_none()); @@ -14366,21 +14805,24 @@ mod tests { } #[tokio::test] - async fn execute_execution_runtime_stream_returns_client_error_with_local_tunnel_message_before_first_data( - ) { - let state = AppState::new().expect("app state should build"); + async fn execute_execution_runtime_stream_sanitizes_local_tunnel_error_before_first_data() { + let state = authenticated_local_tunnel_test_state(); let tunnel_app = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_app.hub.register_proxy(Arc::new(TunnelProxyConn::new( - 901, - "node-1".to_string(), - "Node 1".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + tunnel_app.hub.register_proxy(Arc::new( + TunnelProxyConn::new( + 901, + "node-1".to_string(), + "Node 1".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let plan = ExecutionPlan { request_id: "req-client-stream-error-1".into(), @@ -14428,7 +14870,7 @@ mod tests { .await }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -14436,7 +14878,7 @@ mod tests { .expect("request header frame should parse"); assert_eq!(request_header.msg_type, tunnel_protocol::REQUEST_HEADERS); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -14489,28 +14931,29 @@ mod tests { .and_then(Value::as_str) .expect("response body should contain error.message"); - assert_eq!(error_message, original_error); - assert!( - !error_message.contains("unexpected EOF during chunk size line"), - "client-facing response should preserve the original local tunnel error" - ); + assert_eq!(error_message, "Upstream response stream failed"); + assert!(!error_message.contains(original_error)); } #[tokio::test] async fn execute_execution_runtime_stream_emits_terminal_sse_error_event_after_body_started() { - let state = AppState::new().expect("app state should build"); + let state = authenticated_local_tunnel_test_state(); let tunnel_app = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_app.hub.register_proxy(Arc::new(TunnelProxyConn::new( - 902, - "node-1".to_string(), - "Node 1".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + tunnel_app.hub.register_proxy(Arc::new( + TunnelProxyConn::new( + 902, + "node-1".to_string(), + "Node 1".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let plan = ExecutionPlan { request_id: "req-client-stream-sse-error-1".into(), @@ -14558,7 +15001,7 @@ mod tests { .await }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -14566,7 +15009,7 @@ mod tests { .expect("request header frame should parse"); assert_eq!(request_header.msg_type, tunnel_protocol::REQUEST_HEADERS); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -14633,11 +15076,8 @@ mod tests { let body = body_task.await.expect("body task should complete"); assert!(body.contains("data: hello\n\n")); assert!(body.contains("data: {\"error\":")); - assert!(body.contains(original_error)); + assert!(body.contains("Upstream response stream failed")); + assert!(!body.contains(original_error)); assert!(body.contains("data: [DONE]\n\n")); - assert!( - !body.contains("unexpected EOF during chunk size line"), - "same-format SSE path should surface the original terminal error event" - ); } } diff --git a/apps/aether-gateway/src/execution_runtime/stream/execution_failures.rs b/apps/aether-gateway/src/execution_runtime/stream/execution_failures.rs index 0361035fd..2a7224a93 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/execution_failures.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/execution_failures.rs @@ -68,6 +68,60 @@ enum StreamFailureHandling { HonorLocalFailover, } +const UPSTREAM_STREAM_FAILURE_MESSAGE: &str = "Upstream response stream failed"; +const EXECUTION_STREAM_PROTOCOL_FAILURE_MESSAGE: &str = "Execution runtime stream protocol failed"; +const EXECUTION_STREAM_PROCESSING_FAILURE_MESSAGE: &str = + "Execution runtime stream processing failed"; + +fn encode_bounded_stream_capture(body: &[u8]) -> Option { + encode_stream_capture_with_limit( + body, + crate::execution_runtime::MAX_STREAM_BODY_CAPTURE_BYTES, + ) +} + +fn encode_stream_capture_with_limit(body: &[u8], max_bytes: usize) -> Option { + let captured = &body[..body.len().min(max_bytes)]; + (!captured.is_empty()).then(|| base64::engine::general_purpose::STANDARD.encode(captured)) +} + +fn stable_stream_failure_message(error_type: &str) -> &'static str { + match error_type { + "first_byte_timeout" => "Upstream response timed out before the first byte", + "read_timeout" => "Upstream response stream timed out", + "execution_runtime_stream_read_error" => UPSTREAM_STREAM_FAILURE_MESSAGE, + "execution_runtime_stream_frame_decode_error" + | "execution_runtime_stream_chunk_decode_error" => { + EXECUTION_STREAM_PROTOCOL_FAILURE_MESSAGE + } + "execution_runtime_sync_json_stream_bridge_error" + | "execution_runtime_stream_rewrite_error" + | "execution_runtime_stream_rewrite_flush_error" => { + EXECUTION_STREAM_PROCESSING_FAILURE_MESSAGE + } + _ => "Execution runtime stream failed", + } +} + +fn public_execution_error_message(error: &ExecutionError) -> String { + match &error.kind { + ExecutionErrorKind::ConnectTimeout => "Upstream connection timed out".to_string(), + ExecutionErrorKind::FirstByteTimeout => { + "Upstream response timed out before the first byte".to_string() + } + ExecutionErrorKind::ReadTimeout => "Upstream response stream timed out".to_string(), + ExecutionErrorKind::Upstream4xx | ExecutionErrorKind::Upstream5xx => error + .upstream_status + .map(|status| format!("Upstream request returned HTTP {status}")) + .unwrap_or_else(|| "Upstream request failed".to_string()), + ExecutionErrorKind::TlsError => "Upstream TLS connection failed".to_string(), + ExecutionErrorKind::ProxyError => "Upstream proxy request failed".to_string(), + ExecutionErrorKind::ProtocolError => UPSTREAM_STREAM_FAILURE_MESSAGE.to_string(), + ExecutionErrorKind::Cancelled => "Request was cancelled".to_string(), + ExecutionErrorKind::Internal => "Execution runtime stream failed".to_string(), + } +} + impl StreamFailureReport { fn into_body_jsons(self) -> (Value, Option) { let Self { @@ -110,11 +164,11 @@ impl StreamFailureReport { pub(super) fn build_stream_failure_report( error_type: impl Into, - error_message: impl Into, + _error_message: impl Into, status_code: u16, ) -> StreamFailureReport { let error_type = error_type.into(); - let error_message = error_message.into(); + let error_message = stable_stream_failure_message(error_type.as_str()).to_string(); StreamFailureReport { status_code, error_type, @@ -129,13 +183,14 @@ pub(super) fn build_stream_failure_report( pub(super) fn build_stream_transport_failure_report( error_type: impl Into, - error_message: impl Into, + _error_message: impl Into, status_code: u16, ) -> StreamFailureReport { + let error_type = error_type.into(); StreamFailureReport { status_code, - error_type: error_type.into(), - error_message: error_message.into(), + error_message: stable_stream_failure_message(error_type.as_str()).to_string(), + error_type, upstream_status_code: None, transport_error: true, honor_http_failover: false, @@ -163,7 +218,7 @@ pub(super) fn build_stream_failure_from_execution_error( .ok() .and_then(|value| value.as_str().map(ToOwned::to_owned)) .unwrap_or_else(|| "internal".to_string()); - let error_message = error.message.trim().to_string(); + let error_message = public_execution_error_message(error); let phase = serde_json::to_value(&error.phase).unwrap_or(Value::Null); let mut error_object = Map::from_iter([ ("phase".to_string(), phase), @@ -323,8 +378,7 @@ fn build_stream_failure_sync_payload( headers, body_json: Some(body), client_body_json: client_body, - body_base64: (!provider_buffered_body.is_empty()) - .then(|| base64::engine::general_purpose::STANDARD.encode(provider_buffered_body)), + body_base64: encode_bounded_stream_capture(provider_buffered_body), telemetry, } } @@ -524,8 +578,7 @@ pub(super) async fn handle_prefetch_provider_private_stream_error( headers, body_json: Some(body_json), client_body_json: None, - body_base64: (!buffered_body.is_empty()) - .then(|| base64::engine::general_purpose::STANDARD.encode(buffered_body)), + body_base64: encode_bounded_stream_capture(buffered_body), telemetry, }; let failure_analysis = record_stream_sync_failure( @@ -813,11 +866,11 @@ pub(super) async fn submit_midstream_stream_failure( started_at_unix_ms: u64, failure: StreamFailureReport, ) { - let Some(report_kind) = - direct_stream_finalize_kind.and_then(resolve_core_error_background_report_kind) - else { - return; - }; + let background_report_kind = + direct_stream_finalize_kind.and_then(resolve_core_error_background_report_kind); + let submit_background_report = background_report_kind.is_some(); + let report_kind = + background_report_kind.unwrap_or_else(|| "execution_runtime_stream_error".to_string()); let candidate_status_code = failure.upstream_status_code; let payload = build_stream_failure_sync_payload( @@ -839,7 +892,10 @@ pub(super) async fn submit_midstream_stream_failure( StreamFailureHandling::Terminal, ) .await; - if let Err(err) = submit_sync_report(state, payload).await { + if !submit_background_report { + return; + } + if let Err(_err) = submit_sync_report(state, payload).await { let request_id = short_request_id(plan.request_id.as_str()); warn!( event_name = "execution_report_submit_failed", @@ -848,7 +904,7 @@ pub(super) async fn submit_midstream_stream_failure( request_id = %request_id, candidate_id = ?plan.candidate_id, report_scope = "stream_failure", - error = ?err, + error_category = "stream_report_submit_failed", "gateway failed to submit sync execution report for terminal stream failure" ); } @@ -864,9 +920,22 @@ mod tests { use super::{ build_stream_failure_from_execution_error, build_stream_failure_from_provider_error_body, - build_stream_failure_sync_payload, build_stream_transport_failure_report, + build_stream_failure_report, build_stream_failure_sync_payload, + build_stream_transport_failure_report, encode_stream_capture_with_limit, }; + #[test] + fn failure_capture_encoding_defensively_caps_an_oversized_slice() { + let encoded = encode_stream_capture_with_limit(b"abcdef", 3) + .expect("bounded failure capture should be encoded"); + let decoded = base64::engine::general_purpose::STANDARD + .decode(encoded) + .expect("capture should be valid base64"); + + assert_eq!(decoded, b"abc"); + assert!(encode_stream_capture_with_limit(b"abc", 0).is_none()); + } + #[test] fn committed_transport_failure_has_no_upstream_status() { for status_code in [502, 504] { @@ -913,6 +982,61 @@ mod tests { assert!(!failure.transport_error); } + #[test] + fn execution_error_details_are_not_projected_to_stream_clients() { + let secret = "Bearer stream-secret https://user:password@example.test/private"; + let failure = build_stream_failure_from_execution_error(&ExecutionError { + kind: ExecutionErrorKind::ProtocolError, + phase: ExecutionPhase::StreamRead, + message: secret.to_string(), + upstream_status: None, + retryable: true, + failover_recommended: true, + }); + + assert_eq!(failure.error_message, "Upstream response stream failed"); + let client_body = failure + .to_json_string() + .expect("stream failure should serialize"); + assert!(!client_body.contains(secret)); + assert!(!client_body.contains("stream-secret")); + } + + #[test] + fn upstream_execution_error_keeps_only_http_status_diagnostics() { + let secret = "authorization=Bearer upstream-secret"; + let failure = build_stream_failure_from_execution_error(&ExecutionError { + kind: ExecutionErrorKind::Upstream4xx, + phase: ExecutionPhase::FirstByte, + message: secret.to_string(), + upstream_status: Some(429), + retryable: true, + failover_recommended: true, + }); + + assert_eq!(failure.error_message, "Upstream request returned HTTP 429"); + assert!(!failure + .to_json_string() + .expect("stream failure should serialize") + .contains(secret)); + } + + #[test] + fn internal_stream_failure_details_are_replaced_with_stable_text() { + let secret = "failed near /Users/admin/.config with token=stream-secret"; + let failure = + build_stream_failure_report("execution_runtime_stream_rewrite_error", secret, 502); + + assert_eq!( + failure.error_message, + "Execution runtime stream processing failed" + ); + assert!(!failure + .to_json_string() + .expect("stream failure should serialize") + .contains(secret)); + } + #[test] fn midstream_failure_trace_uses_terminal_error_instead_of_buffered_sse() { let provider_buffered_body = concat!( diff --git a/apps/aether-gateway/src/execution_runtime/stream_pump.rs b/apps/aether-gateway/src/execution_runtime/stream_pump.rs index dbe07ec8f..ab4e2e408 100644 --- a/apps/aether-gateway/src/execution_runtime/stream_pump.rs +++ b/apps/aether-gateway/src/execution_runtime/stream_pump.rs @@ -22,13 +22,14 @@ use crate::ai_serving::api::{ }; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::transport::{ - append_upstream_response_body_chunk, decode_response_body_bytes, format_hyper_error_chain, - format_wreq_upstream_request_error, stream_first_byte_timeout_message, DirectUpstreamResponse, + append_upstream_response_body_chunk, decode_response_body_bytes, + stream_first_byte_timeout_message, DirectUpstreamResponse, }; use crate::execution_runtime::DirectUpstreamStreamExecution; use crate::GatewayError; const STREAM_USAGE_OBSERVER_MAX_LINE_BYTES: usize = 1024 * 1024; +const UPSTREAM_STREAM_READ_ERROR_MESSAGE: &str = "Upstream response stream failed"; pub(crate) fn build_direct_execution_frame_stream( execution: DirectUpstreamStreamExecution, @@ -39,6 +40,7 @@ pub(crate) fn build_direct_execution_frame_stream( candidate_id: _, status_code, headers, + upstream_content_length, provider_api_format, stream_summary_report_context, prefetched_body, @@ -74,7 +76,11 @@ pub(crate) fn build_direct_execution_frame_stream( let mut stream_terminal_observer = StreamingStandardTerminalObserver::default(); let mut observer_buffered = Vec::new(); - if should_buffer_non_stream_response(&headers, &observer_context) { + if should_buffer_non_stream_response( + &headers, + upstream_content_length, + &observer_context, + ) { let original_headers = headers.clone(); match buffer_non_sse_upstream_body( prefetched_body, @@ -104,8 +110,10 @@ pub(crate) fn build_direct_execution_frame_stream( summary = outcome.terminal_summary; } Ok(None) => {} - Err(err) => { - yield Err(IoError::other(format!("{err:?}"))); + Err(_err) => { + yield Err(IoError::other( + "Execution runtime stream conversion failed", + )); return; } } @@ -250,16 +258,16 @@ pub(crate) fn build_direct_execution_frame_stream( } } } - Err(message) => { + Err(_message) => { warn!( event_name = "stream_pump_body_read_error", log_type = "ops", status_code, upstream_bytes, - error = %message, + error_category = "prefetched_body_read_failed", "upstream body stream read error" ); - match encode_error_frame(message) { + match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) { Ok(frame) => yield Ok(frame), Err(encode_err) => { yield Err(encode_err); @@ -333,14 +341,14 @@ pub(crate) fn build_direct_execution_frame_stream( } } } - Err(err) => { - let message = format_error_chain(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", status_code, upstream_bytes, - error = %message, + error_category = "reqwest_body_read_failed", "upstream body stream read error" ); match encode_error_frame(message) { @@ -415,14 +423,14 @@ pub(crate) fn build_direct_execution_frame_stream( } } } - Err(err) => { - let message = format_hyper_error_chain(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", status_code, upstream_bytes, - error = %message, + error_category = "hyper_body_read_failed", "upstream body stream read error" ); match encode_error_frame(message) { @@ -497,14 +505,14 @@ pub(crate) fn build_direct_execution_frame_stream( } } } - Err(err) => { - let message = format_wreq_upstream_request_error(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", status_code, upstream_bytes, - error = %message, + error_category = "browser_body_read_failed", "upstream body stream read error" ); match encode_error_frame(message) { @@ -575,16 +583,16 @@ pub(crate) fn build_direct_execution_frame_stream( } } Ok(None) => break, - Err(message) => { + Err(_message) => { warn!( event_name = "stream_pump_body_read_error", log_type = "ops", status_code, upstream_bytes, - error = %message, + error_category = "tunnel_body_read_failed", "upstream body stream read error" ); - match encode_error_frame(message) { + match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) { Ok(frame) => yield Ok(frame), Err(encode_err) => { yield Err(encode_err); @@ -664,14 +672,14 @@ fn encode_data_frame(chunk: &Bytes) -> Result { }) } -fn encode_error_frame(message: String) -> Result { +fn encode_error_frame(_message: String) -> Result { encode_stream_frame_ndjson(&StreamFrame { frame_type: StreamFrameType::Error, payload: StreamFramePayload::Error { error: ExecutionError { kind: ExecutionErrorKind::ProtocolError, phase: ExecutionPhase::StreamRead, - message, + message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(), upstream_status: None, retryable: true, failover_recommended: true, @@ -770,6 +778,7 @@ fn should_treat_upstream_response_as_stream( fn should_buffer_non_stream_response( headers: &BTreeMap, + upstream_content_length: Option, report_context: &Value, ) -> bool { if should_treat_upstream_response_as_stream(headers, report_context) { @@ -784,6 +793,22 @@ fn should_buffer_non_stream_response( return true; } + // `content-length` is intentionally removed from the response header map + // before it reaches the execution stream. Retain its parsed value as + // internal metadata so only a declared fixed-length JSON response is + // converted to the client's SSE contract. + if report_context + .get("upstream_is_stream") + .and_then(Value::as_bool) + == Some(true) + && upstream_content_length.is_some() + && headers + .get("content-type") + .is_some_and(|value| value.to_ascii_lowercase().contains("json")) + { + return true; + } + headers .get("content-length") .and_then(|value| value.trim().parse::().ok()) @@ -813,9 +838,9 @@ async fn buffer_non_sse_upstream_body( &mut upstream_bytes, )?; } - Err(message) => { + Err(_message) => { return Err(BufferedUpstreamBodyError { - message, + message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(), ttfb_ms, upstream_bytes, first_byte_timeout: None, @@ -864,13 +889,13 @@ async fn buffer_non_sse_upstream_body( &mut upstream_bytes, )?; } - Err(err) => { - let message = format_error_chain(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", upstream_bytes, - error = %message, + error_category = "reqwest_body_read_failed", "upstream body stream read error" ); return Err(BufferedUpstreamBodyError { @@ -922,13 +947,13 @@ async fn buffer_non_sse_upstream_body( &mut upstream_bytes, )?; } - Err(err) => { - let message = format_hyper_error_chain(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", upstream_bytes, - error = %message, + error_category = "hyper_body_read_failed", "upstream body stream read error" ); return Err(BufferedUpstreamBodyError { @@ -980,13 +1005,13 @@ async fn buffer_non_sse_upstream_body( &mut upstream_bytes, )?; } - Err(err) => { - let message = format_wreq_upstream_request_error(&err); + Err(_err) => { + let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); warn!( event_name = "stream_pump_body_read_error", log_type = "ops", upstream_bytes, - error = %message, + error_category = "browser_body_read_failed", "upstream body stream read error" ); return Err(BufferedUpstreamBodyError { @@ -1034,16 +1059,16 @@ async fn buffer_non_sse_upstream_body( )?; } Ok(None) => break, - Err(message) => { + Err(_message) => { warn!( event_name = "stream_pump_body_read_error", log_type = "ops", upstream_bytes, - error = %message, + error_category = "tunnel_body_read_failed", "upstream body stream read error" ); return Err(BufferedUpstreamBodyError { - message, + message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(), ttfb_ms, upstream_bytes, first_byte_timeout: None, @@ -1071,14 +1096,16 @@ fn maybe_bridge_non_sse_sync_json_to_stream( return Ok(None); } - let decoded_body_bytes = decode_response_body_bytes(headers, body_bytes) - .map_err(|error| GatewayError::Internal(error.to_string()))?; + let decoded_body_bytes = decode_response_body_bytes(headers, body_bytes).map_err(|_error| { + GatewayError::Internal("execution runtime response decode failed".to_string()) + })?; if !response_body_is_json(headers, decoded_body_bytes.as_ref()) { return Ok(None); } - let body_json: Value = serde_json::from_slice(decoded_body_bytes.as_ref()) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let body_json: Value = serde_json::from_slice(decoded_body_bytes.as_ref()).map_err(|_err| { + GatewayError::Internal("execution runtime response JSON decode failed".to_string()) + })?; let client_api_format = report_context .get("client_api_format") .and_then(Value::as_str) @@ -1114,17 +1141,6 @@ fn response_body_is_json(headers: &BTreeMap, body_bytes: &[u8]) serde_json::from_slice::(body_bytes).is_ok() } -fn format_error_chain(err: &(dyn std::error::Error + 'static)) -> String { - let mut message = err.to_string(); - let mut source = err.source(); - while let Some(cause) = source { - message.push_str(": "); - message.push_str(&cause.to_string()); - source = cause.source(); - } - message -} - fn observe_stream_chunk( observer: &mut StreamingStandardTerminalObserver, report_context: &Value, @@ -1135,10 +1151,8 @@ fn observe_stream_chunk( let normalized = if let Some(normalizer) = private_stream_normalizer { match normalizer.push_chunk(chunk) { Ok(normalized) => normalized, - Err(err) => { - observer.disable_with_error(format!( - "failed to normalize provider private stream chunk: {err:?}" - )); + Err(_err) => { + observer.disable_with_error("provider stream normalization failed"); return; } } @@ -1160,23 +1174,21 @@ fn finalize_stream_terminal_summary( Ok(flushed) => { observe_normalized_bytes(observer, report_context, observer_buffered, &flushed) } - Err(err) => observer.disable_with_error(format!( - "failed to flush provider private stream normalization: {err:?}" - )), + Err(_err) => observer.disable_with_error("provider stream normalization failed"), } } if !observer_buffered.is_empty() { let line = std::mem::take(observer_buffered); - if let Err(err) = observer.push_line(report_context, line) { - observer.disable_with_error(err.to_string()); + if let Err(_err) = observer.push_line(report_context, line) { + observer.disable_with_error("stream usage parsing failed"); } } match observer.finish(report_context) { Ok(summary) => summary, - Err(err) => { - observer.disable_with_error(err.to_string()); + Err(_err) => { + observer.disable_with_error("stream usage parsing failed"); observer.latest_summary().cloned() } } @@ -1216,8 +1228,8 @@ fn observe_normalized_bytes( remaining = &remaining[line_part_len..]; if observer_buffered.last() == Some(&b'\n') { let line = std::mem::take(observer_buffered); - if let Err(err) = observer.push_line(report_context, line) { - observer.disable_with_error(err.to_string()); + if let Err(_err) = observer.push_line(report_context, line) { + observer.disable_with_error("stream usage parsing failed"); observer_buffered.clear(); return; } @@ -1232,7 +1244,10 @@ mod tests { use std::sync::Arc; use std::time::Duration; + use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody}; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use async_stream::stream; use axum::body::{Body, Bytes}; use axum::extract::ws::Message; @@ -1245,9 +1260,9 @@ mod tests { use tokio::sync::watch; use super::{ - build_direct_execution_frame_stream, observe_normalized_bytes, + build_direct_execution_frame_stream, encode_error_frame, observe_normalized_bytes, should_buffer_non_stream_response, should_treat_upstream_response_as_stream, - STREAM_USAGE_OBSERVER_MAX_LINE_BYTES, + STREAM_USAGE_OBSERVER_MAX_LINE_BYTES, UPSTREAM_STREAM_READ_ERROR_MESSAGE, }; use crate::ai_serving::api::StreamingStandardTerminalObserver; use crate::execution_runtime::transport::{ @@ -1267,6 +1282,66 @@ mod tests { } } + const LOCAL_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const LOCAL_TUNNEL_TEST_GENERATION: &str = "stream-pump-test-generation-1"; + + fn authenticated_local_tunnel_test_state() -> AppState { + let node = StoredProxyNode::new( + "node-1".to_string(), + "Node 1".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "online".to_string(), + 30, + 1, + 0, + 0, + 0, + 0, + true, + true, + 1, + ) + .expect("tunnel node should build") + .with_runtime_fields( + None, + None, + None, + None, + Some(serde_json::json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": LOCAL_TUNNEL_TEST_PSK, + } + })), + None, + None, + None, + None, + None, + None, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()); + let data = crate::data::GatewayDataState::with_proxy_node_repository_for_tests(Arc::new( + InMemoryProxyNodeRepository::seed([node]), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + AppState::new() + .expect("app state should build") + .with_data_state_for_tests(data) + } + + async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, + ) -> Message { + tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) + } + #[test] fn treats_kiro_eventstream_envelope_as_stream_even_when_content_type_is_json() { let headers = BTreeMap::from([("content-type".into(), "application/json".into())]); @@ -1295,10 +1370,12 @@ mod tests { assert!(!should_buffer_non_stream_response( &BTreeMap::from([("content-type".into(), "application/json".into())]), + None, &streaming_context )); assert!(should_buffer_non_stream_response( &BTreeMap::from([("content-type".into(), "application/json".into())]), + None, &non_stream_context )); assert!(should_buffer_non_stream_response( @@ -1306,14 +1383,27 @@ mod tests { ("content-type".into(), "application/json".into()), ("content-length".into(), "128".into()), ]), + Some(128), &streaming_context )); assert!(!should_buffer_non_stream_response( &BTreeMap::from([("content-type".into(), "text/event-stream".into())]), + None, &non_stream_context )); } + #[test] + fn error_frames_do_not_include_transport_details() { + let secret = "Bearer stream-secret https://user:password@example.test/private"; + let frame = encode_error_frame(secret.to_string()).expect("error frame should encode"); + let frame = String::from_utf8(frame.to_vec()).expect("error frame should be utf8"); + + assert!(frame.contains(UPSTREAM_STREAM_READ_ERROR_MESSAGE)); + assert!(!frame.contains(secret)); + assert!(!frame.contains("stream-secret")); + } + #[test] fn oversized_usage_line_disables_observation_without_retaining_the_line() { let mut observer = StreamingStandardTerminalObserver::default(); @@ -1920,20 +2010,24 @@ mod tests { } #[tokio::test] - async fn direct_execution_frame_stream_preserves_local_tunnel_stream_error_message() { - let state = AppState::new().expect("app state should build"); + async fn direct_execution_frame_stream_sanitizes_local_tunnel_stream_error_message() { + let state = authenticated_local_tunnel_test_state(); let tunnel_app = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_app.hub.register_proxy(Arc::new(TunnelProxyConn::new( - 801, - "node-1".to_string(), - "Node 1".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + tunnel_app.hub.register_proxy(Arc::new( + TunnelProxyConn::new( + 801, + "node-1".to_string(), + "Node 1".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let plan = ExecutionPlan { request_id: "req-local-stream-error-1".into(), @@ -1967,7 +2061,7 @@ mod tests { execute_stream_plan_via_local_tunnel(&state_for_task, &plan_for_task).await }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -1975,7 +2069,7 @@ mod tests { .expect("request header frame should parse"); assert_eq!(request_header.msg_type, tunnel_protocol::REQUEST_HEADERS); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -2057,28 +2151,29 @@ mod tests { .and_then(Value::as_str) .expect("error frame should include a message"); - assert_eq!(error_message, original_error); - assert!( - !error_message.contains("unexpected EOF during chunk size line"), - "local tunnel path should preserve the original proxy error text" - ); + assert_eq!(error_message, UPSTREAM_STREAM_READ_ERROR_MESSAGE); + assert!(!error_message.contains(original_error)); } #[tokio::test] async fn second_local_tunnel_request_works_after_first_completes() { - let state = AppState::new().expect("app state should build"); + let state = authenticated_local_tunnel_test_state(); let tunnel_app = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_app.hub.register_proxy(Arc::new(TunnelProxyConn::new( - 900, - "node-1".to_string(), - "Node 1".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + tunnel_app.hub.register_proxy(Arc::new( + TunnelProxyConn::new( + 900, + "node-1".to_string(), + "Node 1".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let plan = ExecutionPlan { request_id: "req-reuse-1".into(), @@ -2115,13 +2210,13 @@ mod tests { ); // Read request frames from proxy side - let req1_headers = match proxy_rx.recv().await.expect("req1 headers") { + let req1_headers = match recv_tunnel_test_frame(&mut proxy_rx, "req1 headers").await { Message::Binary(data) => data, other => panic!("unexpected: {other:?}"), }; let req1_header = tunnel_protocol::FrameHeader::parse(&req1_headers).expect("req1 header parse"); - let _req1_body = proxy_rx.recv().await.expect("req1 body"); + let _req1_body = recv_tunnel_test_frame(&mut proxy_rx, "req1 body").await; // Simulate proxy response let resp_meta = serde_json::to_vec(&tunnel_protocol::ResponseMeta { @@ -2188,10 +2283,7 @@ mod tests { ); // Read second request's frames - let req2_headers = tokio::time::timeout(Duration::from_secs(2), proxy_rx.recv()) - .await - .expect("second request should arrive within 2s") - .expect("req2 headers"); + let req2_headers = recv_tunnel_test_frame(&mut proxy_rx, "req2 headers").await; let req2_data = match req2_headers { Message::Binary(data) => data, other => panic!("unexpected: {other:?}"), diff --git a/apps/aether-gateway/src/execution_runtime/submission.rs b/apps/aether-gateway/src/execution_runtime/submission.rs index 861b5792b..9e7186058 100644 --- a/apps/aether-gateway/src/execution_runtime/submission.rs +++ b/apps/aether-gateway/src/execution_runtime/submission.rs @@ -9,9 +9,9 @@ use crate::api::response::build_client_response_from_parts; use crate::control::GatewayControlDecision; use crate::usage::spawn_sync_report; use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError}; +use aether_usage_runtime::decode_internal_report_body_base64; use axum::body::Body; use axum::http::{Response, StatusCode}; -use base64::Engine as _; use tracing::warn; #[derive(Clone, Debug)] @@ -144,9 +144,8 @@ fn build_local_core_sync_finalize_fallback_response( } if let Some(body_base64) = payload.body_base64.as_ref() { - let body_bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let body_bytes = + decode_internal_report_body_base64(body_base64).map_err(GatewayError::Internal)?; return build_local_sync_response_from_bytes(trace_id, decision, payload, body_bytes); } @@ -299,9 +298,8 @@ fn resolve_local_sync_source_body_json( let body_json = if let Some(body_json) = payload.body_json.clone() { body_json } else if let Some(body_base64) = payload.body_base64.as_deref() { - let body_bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let body_bytes = + decode_internal_report_body_base64(body_base64).map_err(GatewayError::Internal)?; let stripped = strip_utf8_bom_and_ws(&body_bytes); let Ok(body_json) = serde_json::from_slice::(stripped) else { return Ok(None); @@ -328,9 +326,8 @@ fn decode_local_sync_body_text( let Some(body_base64) = payload.body_base64.as_deref() else { return Ok(None); }; - let body_bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let body_bytes = + decode_internal_report_body_base64(body_base64).map_err(GatewayError::Internal)?; let stripped = strip_utf8_bom_and_ws(&body_bytes); let body_text = String::from_utf8_lossy(stripped).trim().to_string(); if body_text.is_empty() { diff --git a/apps/aether-gateway/src/execution_runtime/sync/execution.rs b/apps/aether-gateway/src/execution_runtime/sync/execution.rs index 84d67a25a..fe72956be 100644 --- a/apps/aether-gateway/src/execution_runtime/sync/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/sync/execution.rs @@ -34,6 +34,7 @@ use crate::ai_serving::api::{ implicit_sync_finalize_report_kind, maybe_build_sync_finalize_outcome, LocalCoreSyncErrorKind, LocalCoreSyncFinalizeOutcome, }; +use crate::ai_serving::record_local_runtime_candidate_skip_reason; use crate::api::response::{ attach_control_metadata_headers, build_client_response, build_client_response_from_parts, build_client_response_from_parts_with_mutator, @@ -59,8 +60,8 @@ use crate::execution_runtime::transport::{ build_request_body, collect_response_headers, decode_response_body_bytes_with_limit, execution_plan_response_body_limit_bytes, execution_response_body_mode, format_hyper_error_chain, format_upstream_request_error, format_wreq_upstream_request_error, - response_body_is_json, send_request, DirectHttpResponse, DirectSyncExecutionRuntime, - ExecutionRuntimeTransportError, + response_body_is_json, safe_transport_error_message, send_request, DirectHttpResponse, + DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, }; use crate::execution_runtime::windsurf::maybe_execute_windsurf_sync; use crate::execution_runtime::{ @@ -100,7 +101,7 @@ mod policy; #[path = "execution/response.rs"] mod response; -use policy::decode_execution_result_body; +use policy::{decode_execution_result_body, decode_execution_result_body_with_limit}; pub(crate) use response::{ maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response, maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind, @@ -114,6 +115,11 @@ const SYNC_EXECUTION_IDLE_LOG_INTERVAL: Duration = Duration::from_secs(60); const OPENAI_IMAGE_SYNC_JSON_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(15); const OPENAI_IMAGE_SYNC_JSON_HEARTBEAT_BYTES: &[u8] = b"\n"; const OPENAI_IMAGE_SYNC_PROGRESS_WRITE_INTERVAL: Duration = Duration::from_secs(5); +// Progress parsing must retain only the incomplete SSE record. This is not a +// response-body limit: the full upstream body is still handled by the normal +// execution body policy, while malformed streams cannot grow telemetry state +// forever by withholding a record separator. +const OPENAI_IMAGE_SYNC_PROGRESS_MAX_BUFFER_BYTES: usize = 16 * 1024 * 1024; const INVALID_GEMINI_PROVIDER_SUCCESS_MESSAGE: &str = "Provider returned HTTP 200 but the Gemini response did not contain visible model output; refusing to finalize it as a successful response."; fn elapsed_ms_since(started_at: Instant) -> u64 { @@ -173,6 +179,39 @@ struct SyncAttemptTerminalGuard { armed: bool, } +/// Keep forced-terminal records useful for operations without copying an +/// arbitrary `GatewayError` into the candidate/usage stores. Gateway errors +/// can wrap provider URLs, credentials, query strings, or database details; +/// those values belong in the internal logging path only. +fn persisted_sync_abort_message(error: &GatewayError) -> &'static str { + match error { + GatewayError::UpstreamUnavailable { .. } => { + "local sync attempt aborted before terminal finalization: upstream unavailable" + } + GatewayError::ControlUnavailable { .. } => { + "local sync attempt aborted before terminal finalization: control unavailable" + } + GatewayError::LocalExecutionPlanningTimeout { .. } => { + "local sync attempt aborted before terminal finalization: planning timeout" + } + GatewayError::AdmissionTimeout { .. } => { + "local sync attempt aborted before terminal finalization: admission timeout" + } + GatewayError::Client { .. } => { + "local sync attempt aborted before terminal finalization: client error" + } + GatewayError::PlanUsageLimited(_) => { + "local sync attempt aborted before terminal finalization: usage limit" + } + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + "local sync attempt aborted before terminal finalization: policy denied" + } + GatewayError::Internal(_) => { + "local sync attempt aborted before terminal finalization: internal error" + } + } +} + impl SyncAttemptTerminalGuard { fn new( state: &AppState, @@ -212,7 +251,7 @@ impl SyncAttemptTerminalGuard { RequestCandidateStatus::Failed, StatusCode::INTERNAL_SERVER_ERROR.as_u16(), "local_sync_attempt_aborted", - format!("Local sync attempt failed before terminal finalization: {error:?}"), + persisted_sync_abort_message(error), ) .await; } @@ -345,7 +384,7 @@ impl SyncExecutionFailure { error_type: fallback_kind .map(SyncExecutionFailureFallbackKind::error_type) .unwrap_or("execution_runtime_unavailable"), - message: err.to_string(), + message: safe_transport_error_message(&err), status_code: fallback_kind.map(|_| StatusCode::BAD_GATEWAY.as_u16()), latency_ms: None, fallback_kind, @@ -1131,10 +1170,42 @@ impl<'a> OpenAiImageSyncProgressRecorder<'a> { if chunk.is_empty() { return; } - self.buffer.extend_from_slice(chunk); + let mut blocks = Vec::new(); + let mut remaining = chunk; + let mut parser_overflowed = false; + loop { + while let Some(block_end) = find_sse_block_end(&self.buffer) { + blocks.push(self.buffer.drain(..block_end).collect::>()); + } + if remaining.is_empty() { + break; + } + + let capacity = + OPENAI_IMAGE_SYNC_PROGRESS_MAX_BUFFER_BYTES.saturating_sub(self.buffer.len()); + if capacity == 0 { + // The progress recorder is observational. Drop an incomplete + // oversized record and keep the client-facing response alive. + self.buffer.clear(); + parser_overflowed = true; + break; + } + let take = find_sse_block_end(remaining) + .map_or(remaining.len(), |block_end| block_end) + .min(capacity); + if take == 0 { + self.buffer.clear(); + parser_overflowed = true; + break; + } + self.buffer.extend_from_slice(&remaining[..take]); + remaining = &remaining[take..]; + } + if parser_overflowed { + debug!("openai image sync progress parser dropped an oversized incomplete SSE record"); + } let mut force_persist = false; - while let Some(block_end) = find_sse_block_end(&self.buffer) { - let block = self.buffer.drain(..block_end).collect::>(); + for block in blocks { let Some(frame) = parse_openai_image_sync_sse_frame(&block) else { continue; }; @@ -1812,12 +1883,16 @@ async fn openai_image_sync_json_heartbeat_final_bytes( { Ok(bytes) if !bytes.is_empty() => bytes.to_vec(), Ok(_) => openai_image_sync_json_heartbeat_error_body("empty sync image response"), - Err(err) => openai_image_sync_json_heartbeat_error_body(&err.to_string()), + Err(_err) => { + openai_image_sync_json_heartbeat_error_body("sync image response read failed") + } }, Ok(None) => openai_image_sync_json_heartbeat_error_body( "sync image execution ended without a local response", ), - Err(err) => openai_image_sync_json_heartbeat_error_body(&format!("{err:?}")), + // Do not serialize the internal error: Debug output can contain + // upstream URLs, credentials, or other request-specific details. + Err(_err) => openai_image_sync_json_heartbeat_error_body("sync image execution failed"), } } @@ -1839,7 +1914,7 @@ async fn apply_sync_success_effects( ) { if let Some(report_context) = report_context { crate::ai_serving::persist_converted_response_history( - state.runtime_state(), + state, report_context, payload .client_body_json @@ -2002,6 +2077,11 @@ async fn execute_execution_runtime_sync_impl( { ProviderPoolInFlightAdmission::Acquired(guard) => guard, ProviderPoolInFlightAdmission::Saturated { limit } => { + record_local_runtime_candidate_skip_reason( + state, + trace_id, + "provider_key_concurrency_limit_reached", + ); if let Some(retry_scope) = retry_scope_out.as_deref_mut() { *retry_scope = AiAttemptRetryScope::Candidate; } @@ -2149,7 +2229,7 @@ async fn execute_execution_runtime_sync_impl( } }, Err(err) => { - let transport_error_message = err.to_string(); + let transport_error_message = safe_transport_error_message(&err); warn!( event_name = "chatgpt_web_image_execution_unavailable", log_type = "ops", @@ -2161,7 +2241,7 @@ async fn execute_execution_runtime_sync_impl( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error = %transport_error_message, "gateway ChatGPT-Web image execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -2201,7 +2281,7 @@ async fn execute_execution_runtime_sync_impl( } } Err(err) => { - let transport_error_message = err.to_string(); + let transport_error_message = safe_transport_error_message(&err); warn!( event_name = "grok_execution_unavailable", log_type = "ops", @@ -2213,7 +2293,7 @@ async fn execute_execution_runtime_sync_impl( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error = %transport_error_message, "gateway Grok execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -2256,7 +2336,7 @@ async fn execute_execution_runtime_sync_impl( match (override_fn.0)(&plan) { Ok(result) => result, Err(err) => { - let transport_error_message = format!("{err:?}"); + let transport_error_message = persisted_sync_abort_message(&err).to_string(); warn!( event_name = "sync_execution_runtime_test_override_failed", log_type = "ops", @@ -2402,7 +2482,7 @@ async fn execute_execution_runtime_sync_impl( } }, Err(err) => { - let transport_error_message = err.to_string(); + let transport_error_message = safe_transport_error_message(&err); warn!( event_name = "chatgpt_web_image_execution_unavailable", log_type = "ops", @@ -2414,7 +2494,7 @@ async fn execute_execution_runtime_sync_impl( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error = %transport_error_message, "gateway ChatGPT-Web image execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -2453,7 +2533,7 @@ async fn execute_execution_runtime_sync_impl( } }, Err(err) => { - let transport_error_message = err.to_string(); + let transport_error_message = safe_transport_error_message(&err); warn!( event_name = "grok_execution_unavailable", log_type = "ops", @@ -2465,7 +2545,7 @@ async fn execute_execution_runtime_sync_impl( key_id, model_name, candidate_index = candidate_index.as_str(), - error = %err, + error = %transport_error_message, "gateway Grok execution unavailable" ); let terminal_unix_secs = current_request_candidate_unix_ms(); @@ -2566,8 +2646,26 @@ async fn execute_execution_runtime_sync_impl( .as_ref() .and_then(|telemetry| telemetry.elapsed_ms); let mut headers = std::mem::take(&mut result.headers); - let (body_bytes, mut body_json, body_base64) = - decode_execution_result_body(result.body.take(), &mut headers)?; + let result_body = result.body.take(); + let chatgpt_web_image_result = plan + .provider_api_format + .eq_ignore_ascii_case("openai:image") + && (plan.headers.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("x-aether-chatgpt-web-image") && value == "1" + }) || report_context + .as_ref() + .and_then(|value| value.get("chatgpt_web_image")) + .and_then(Value::as_bool) + .unwrap_or(false)); + let (body_bytes, mut body_json, body_base64) = if chatgpt_web_image_result { + decode_execution_result_body_with_limit( + result_body, + &mut headers, + crate::execution_runtime::chatgpt_web_image::chatgpt_web_image_sse_envelope_limit_bytes(), + )? + } else { + decode_execution_result_body(result_body, &mut headers)? + }; if let Some(message) = invalid_gemini_provider_success_message( &plan, report_context.as_ref(), @@ -3376,7 +3474,7 @@ async fn execute_sync_via_remote_execution_runtime( status: RequestCandidateStatus::Failed, status_code: None, error_type: Some("execution_runtime_unavailable".to_string()), - error_message: Some(format!("{err:?}")), + error_message: Some(persisted_sync_abort_message(&err).to_string()), latency_ms: Some(elapsed_ms_since(candidate_started_at)), started_at_unix_ms: Some(candidate_started_unix_secs), finished_at_unix_ms: Some(terminal_unix_secs), @@ -3417,9 +3515,15 @@ async fn execute_sync_via_remote_execution_runtime( } let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms(); - let mut result = response - .json::() - .await + let response_body = aether_http::read_response_bytes_with_limit( + response, + crate::execution_runtime::transport::execution_result_envelope_limit_bytes( + crate::headers::max_internal_buffered_body_bytes(), + ), + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut result = serde_json::from_slice::(&response_body) .map_err(|err| GatewayError::Internal(err.to_string()))?; result .response_observation @@ -3739,6 +3843,62 @@ mod tests { assert!(message.contains("visible model output")); } + #[test] + fn invalid_gemini_provider_success_accepts_thought_only_max_tokens() { + let plan = test_gemini_chat_plan(); + let body = json!({ + "candidates": [{ + "content": { + "role": "model", + "parts": [{"text": "hidden plan", "thought": true}] + }, + "finishReason": "MAX_TOKENS" + }], + "usageMetadata": { + "promptTokenCount": 8, + "candidatesTokenCount": 0, + "thoughtsTokenCount": 24, + "totalTokenCount": 32 + } + }); + + let message = invalid_gemini_provider_success_message( + &plan, + None, + StatusCode::OK.as_u16(), + Some(&body), + ); + + assert!(message.is_none()); + } + + #[test] + fn invalid_gemini_provider_stream_success_accepts_signature_only_reasoning_exhaustion() { + let plan = test_gemini_chat_plan(); + let report_context = json!({ + "has_envelope": true, + "envelope_name": "antigravity:v1internal", + "provider_api_format": "gemini:generate_content", + }); + let body = concat!( + "data: {\"response\":{\"responseId\":\"resp_signature_only_123\",\"modelVersion\":\"gemini-3.7-flash-tiered\",", + "\"candidates\":[{\"index\":0,\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"\",\"thoughtSignature\":\"opaque-thought-signature\"}]},\"finishReason\":\"MAX_TOKENS\"}],", + "\"usageMetadata\":{\"promptTokenCount\":22,\"thoughtsTokenCount\":29,\"totalTokenCount\":51}},", + "\"traceId\":\"trace-signature-only\"}\n\n", + ); + + let message = invalid_gemini_provider_stream_success_message( + &plan, + Some(&report_context), + StatusCode::OK.as_u16(), + None, + body.as_bytes(), + true, + ); + + assert!(message.is_none()); + } + #[test] fn invalid_gemini_provider_success_error_is_retryable_candidate_failure() { let error = invalid_gemini_provider_success_execution_error( @@ -3831,6 +3991,20 @@ mod tests { assert!(message.is_none()); } + #[test] + fn forced_sync_abort_message_does_not_copy_internal_error_details() { + let secret = "https://user:password@example.test/v1?api_key=should-not-persist"; + let error = GatewayError::Internal(secret.to_string()); + + let message = persisted_sync_abort_message(&error); + assert_eq!( + message, + "local sync attempt aborted before terminal finalization: internal error" + ); + assert!(!message.contains("password")); + assert!(!message.contains("api_key")); + } + #[tokio::test] async fn sync_attempt_terminal_guard_marks_dropped_pending_attempt_cancelled() { let usage_repository = Arc::new(InMemoryUsageReadRepository::default()); @@ -4420,4 +4594,22 @@ mod tests { Some(&report_context), )); } + + #[tokio::test] + async fn openai_image_sync_json_heartbeat_hides_internal_error_details() { + let secret = "https://user:password@example.test/private?api_key=top-secret"; + let bytes = openai_image_sync_json_heartbeat_final_bytes(Err(GatewayError::Internal( + secret.to_string(), + ))) + .await; + + let body: Value = serde_json::from_slice(&bytes).expect("heartbeat error body is JSON"); + assert_eq!( + body.pointer("/error/message").and_then(Value::as_str), + Some("sync image execution failed") + ); + let body_text = String::from_utf8(bytes).expect("heartbeat body is UTF-8"); + assert!(!body_text.contains(secret)); + assert!(!body_text.contains("top-secret")); + } } diff --git a/apps/aether-gateway/src/execution_runtime/sync/execution/policy.rs b/apps/aether-gateway/src/execution_runtime/sync/execution/policy.rs index f865ba025..5a7adc95a 100644 --- a/apps/aether-gateway/src/execution_runtime/sync/execution/policy.rs +++ b/apps/aether-gateway/src/execution_runtime/sync/execution/policy.rs @@ -1,15 +1,35 @@ +use aether_contracts::ResponseBody; +use serde_json::Value; use std::collections::BTreeMap; -use aether_contracts::ResponseBody; -use base64::Engine as _; - +use crate::execution_runtime::transport::{ + decode_base64_body_with_limit, + serialize_json_body_with_limit as serialize_transport_json_body_with_limit, +}; use crate::GatewayError; type DecodedBody = (Vec, Option, Option); +fn serialize_json_body_with_limit(body: &Value, limit: usize) -> Result, GatewayError> { + serialize_transport_json_body_with_limit(body, limit) + .map_err(|error| GatewayError::Internal(error.to_string())) +} + pub(super) fn decode_execution_result_body( body: Option, headers: &mut BTreeMap, +) -> Result { + decode_execution_result_body_with_limit( + body, + headers, + crate::headers::max_internal_buffered_body_bytes(), + ) +} + +pub(super) fn decode_execution_result_body_with_limit( + body: Option, + headers: &mut BTreeMap, + body_limit: usize, ) -> Result { let Some(body) = body else { return Ok((Vec::new(), None, None)); @@ -18,10 +38,8 @@ pub(super) fn decode_execution_result_body( json_body, body_bytes_b64, } = body; - if let Some(body_bytes_b64) = body_bytes_b64 { - let bytes = base64::engine::general_purpose::STANDARD - .decode(&body_bytes_b64) + let bytes = decode_base64_body_with_limit(&body_bytes_b64, body_limit) .map_err(|err| GatewayError::Internal(err.to_string()))?; return Ok((bytes, json_body, Some(body_bytes_b64))); } @@ -32,8 +50,7 @@ pub(super) fn decode_execution_result_body( headers .entry("content-type".to_string()) .or_insert_with(|| "application/json".to_string()); - let bytes = serde_json::to_vec(&json_body) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let bytes = serialize_json_body_with_limit(&json_body, body_limit)?; headers.insert("content-length".to_string(), bytes.len().to_string()); return Ok((bytes, Some(json_body), None)); } @@ -118,4 +135,34 @@ mod tests { Some(raw_len.as_str()) ); } + + #[test] + fn scoped_decode_limit_can_cover_a_bounded_synthetic_envelope() { + let raw = vec![b'x'; 65 * 1024]; + let encoded = base64::engine::general_purpose::STANDARD.encode(&raw); + let mut headers = BTreeMap::new(); + + let (decoded, json, retained) = super::decode_execution_result_body_with_limit( + Some(ResponseBody { + json_body: None, + body_bytes_b64: Some(encoded.clone()), + }), + &mut headers, + raw.len(), + ) + .expect("body at the scoped limit should decode"); + assert_eq!(decoded, raw); + assert_eq!(json, None); + assert_eq!(retained.as_deref(), Some(encoded.as_str())); + + assert!(super::decode_execution_result_body_with_limit( + Some(ResponseBody { + json_body: None, + body_bytes_b64: Some(encoded), + }), + &mut headers, + raw.len() - 1, + ) + .is_err()); + } } diff --git a/apps/aether-gateway/src/execution_runtime/sync/execution/response.rs b/apps/aether-gateway/src/execution_runtime/sync/execution/response.rs index 34f360abe..9f5f9ccb9 100644 --- a/apps/aether-gateway/src/execution_runtime/sync/execution/response.rs +++ b/apps/aether-gateway/src/execution_runtime/sync/execution/response.rs @@ -2,13 +2,10 @@ use std::collections::BTreeMap; use aether_contracts::ExecutionPlan; use axum::body::Body; -use axum::http::header::HeaderValue; use axum::http::Response; use serde_json::json; -use crate::api::response::{ - build_client_response_from_parts, build_client_response_from_parts_with_mutator, -}; +use crate::api::response::build_client_response_from_parts; use crate::async_task::VideoTaskService; use crate::control::GatewayControlDecision; use crate::video_tasks::{ @@ -156,35 +153,55 @@ pub(crate) fn maybe_build_local_video_error_response( return Ok(None); } - let empty_body = json!({}); - let response_body = payload.body_json.as_ref().unwrap_or(&empty_body); - let body_bytes = - serde_json::to_vec(response_body).map_err(|err| GatewayError::Internal(err.to_string()))?; - let body_len = body_bytes.len().to_string(); + let response_body = + local_video_error_response_body(payload.report_kind.as_str(), payload.status_code); + let body_bytes = serde_json::to_vec(&response_body) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let headers = BTreeMap::from([ + ("content-type".to_string(), "application/json".to_string()), + ("content-length".to_string(), body_bytes.len().to_string()), + ]); - Ok(Some(build_client_response_from_parts_with_mutator( + Ok(Some(build_client_response_from_parts( payload.status_code, - &payload.headers, + &headers, Body::from(body_bytes), trace_id, Some(decision), - |headers| { - headers.remove(http::header::CONTENT_ENCODING); - headers.remove(http::header::CONTENT_LENGTH); - headers.insert( - http::header::CONTENT_TYPE, - HeaderValue::from_static("application/json"), - ); - headers.insert( - http::header::CONTENT_LENGTH, - HeaderValue::from_str(body_len.as_str()) - .map_err(|err| GatewayError::Internal(err.to_string()))?, - ); - Ok(()) - }, )?)) } +fn local_video_error_response_body(report_kind: &str, status_code: u16) -> serde_json::Value { + let (code, gemini_status) = match status_code { + 400 => ("invalid_request", "INVALID_ARGUMENT"), + 401 => ("authentication_error", "UNAUTHENTICATED"), + 403 => ("permission_denied", "PERMISSION_DENIED"), + 404 => ("not_found", "NOT_FOUND"), + 429 => ("rate_limit_exceeded", "RESOURCE_EXHAUSTED"), + 503 => ("server_error", "UNAVAILABLE"), + 500..=599 => ("server_error", "INTERNAL"), + _ => ("provider_error", "UNKNOWN"), + }; + + if report_kind.starts_with("gemini_video_") { + json!({ + "error": { + "code": status_code, + "message": "Video generation failed", + "status": gemini_status, + } + }) + } else { + json!({ + "error": { + "message": "Video generation failed", + "type": code, + "code": code, + } + }) + } +} + #[cfg(test)] mod tests { use super::*; @@ -193,7 +210,7 @@ mod tests { use serde_json::json; #[tokio::test] - async fn local_video_error_response_rewrites_headers_without_mutating_payload() { + async fn local_video_error_response_does_not_expose_upstream_payload_or_headers() { let decision = GatewayControlDecision::synthetic( "/v1/videos", Some("ai_public".to_string()), @@ -212,12 +229,15 @@ mod tests { headers: BTreeMap::from([ ("content-encoding".to_string(), "gzip".to_string()), ("content-length".to_string(), "999".to_string()), - ("x-upstream-id".to_string(), "video-123".to_string()), + ( + "x-upstream-debug".to_string(), + "Authorization: Bearer header-secret".to_string(), + ), ]), body_json: Some(json!({ "error": { "type": "video_backend_error", - "message": "backend failed", + "message": "Authorization: Bearer body-secret at https://internal.test/?key=secret", } })), client_body_json: None, @@ -239,13 +259,7 @@ mod tests { Some("application/json") ); assert_eq!(response.headers().get(http::header::CONTENT_ENCODING), None); - assert_eq!( - response - .headers() - .get("x-upstream-id") - .and_then(|value| value.to_str().ok()), - Some("video-123") - ); + assert_eq!(response.headers().get("x-upstream-debug"), None); assert_eq!( payload.headers.get("content-encoding").map(String::as_str), Some("gzip") @@ -258,12 +272,36 @@ mod tests { let body = to_bytes(response.into_body(), usize::MAX) .await .expect("response body should read"); + let response_body = + serde_json::from_slice::(&body).expect("response body should parse"); assert_eq!( - serde_json::from_slice::(&body).expect("response body should parse"), - payload - .body_json - .clone() - .expect("payload body should exist") + response_body, + json!({ + "error": { + "message": "Video generation failed", + "type": "server_error", + "code": "server_error", + } + }) + ); + let encoded = response_body.to_string(); + for sensitive in ["Bearer", "body-secret", "header-secret", "internal.test"] { + assert!(!encoded.contains(sensitive)); + } + assert!(payload.body_json.is_some()); + } + + #[test] + fn gemini_video_error_response_uses_fixed_schema() { + assert_eq!( + local_video_error_response_body("gemini_video_create_sync_finalize", 429), + json!({ + "error": { + "code": 429, + "message": "Video generation failed", + "status": "RESOURCE_EXHAUSTED", + } + }) ); } } diff --git a/apps/aether-gateway/src/execution_runtime/transport.rs b/apps/aether-gateway/src/execution_runtime/transport.rs index 2eeb8c38d..411c69fcf 100644 --- a/apps/aether-gateway/src/execution_runtime/transport.rs +++ b/apps/aether-gateway/src/execution_runtime/transport.rs @@ -1,23 +1,30 @@ use std::borrow::Cow; -use std::collections::{BTreeMap, HashMap, HashSet, VecDeque}; +use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet, VecDeque}; use std::error::Error as _; use std::future::Future; use std::io::Read; use std::io::Write; +use std::net::{IpAddr, SocketAddr}; +use std::pin::Pin; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, LazyLock, Mutex as StdMutex, OnceLock, RwLock as StdRwLock}; +use std::task::{Context, Poll}; use std::time::{Duration, Instant}; +use aether_contracts::tunnel::MAX_TUNNEL_RELAY_META_LEN; use aether_contracts::{ ExecutionPlan, ExecutionResponseBodyMode, ExecutionResponseObservation, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile, ResponseBody, - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, - EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER, + EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, + EXECUTION_RESPONSE_BODY_MODE_HEADER, PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY, }; use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation; -use aether_http::{apply_http_client_config, HttpClientConfig}; +use aether_http::{ + apply_http_client_config, is_https_or_loopback_http_url, is_ipv4_benchmarking_fake_ip, + is_private_or_reserved_ip, HttpClientConfig, +}; use aether_runtime::{MetricKind, MetricSample}; use axum::body::Bytes; use base64::Engine as _; @@ -56,17 +63,27 @@ use crate::upstream_admission::UpstreamTargetAdmissionPermit; use crate::{AppState, GatewayError}; const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope"; +pub(crate) const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str = + aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY; const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error"; -const TUNNEL_RELAY_PATH_PREFIX: &str = "/api/internal/tunnel/relay"; +const MAX_SAFE_REDIRECTS: usize = 10; +const MAX_UPSTREAM_ERROR_DETAIL_BYTES: usize = 2_048; const DEFAULT_TUNNEL_TIMEOUT_MS: u64 = 60_000; const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS: u64 = 30_000; const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000; const DEFAULT_CODEX_COMPACT_TOTAL_TIMEOUT_MS: u64 = 1_200_000; const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1; const EXECUTION_RESPONSE_BODY_LIMIT_HEADER: &str = "x-aether-execution-response-body-limit-bytes"; +const LEGACY_EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER: &str = + "x-aether-execution-accept-invalid-certs"; const DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 8 * 1024 * 1024; const MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024; const MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024 * 1024; +// A remote execution result is JSON and may carry both the parsed JSON body +// and the original bytes. Keep room for base64 expansion, the second body +// representation, and bounded response metadata while retaining a hard cap. +const MAX_EXECUTION_RESULT_ENVELOPE_BYTES: usize = 256 * 1024 * 1024; +const EXECUTION_RESULT_ENVELOPE_METADATA_BYTES: usize = 8 * 1024 * 1024; const DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_H2_CLIENT_SHARDS"; const DIRECT_REQWEST_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_CLIENT_SHARDS"; const DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT_ENV: &str = @@ -75,6 +92,8 @@ const DIRECT_REQWEST_HTTP1_TARGET_STREAMS_PER_CLIENT_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_HTTP1_TARGET_STREAMS_PER_CLIENT"; const DIRECT_REQWEST_STREAM_HTTP_MODE_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_STREAM_HTTP_MODE"; const DIRECT_REQWEST_CACHE_PER_ORIGIN_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_CACHE_PER_ORIGIN"; +const DIRECT_REQWEST_CACHE_MAX_ENTRIES_ENV: &str = + "AETHER_GATEWAY_DIRECT_REQWEST_CACHE_MAX_ENTRIES"; const DIRECT_H2C_FAST_PATH_ENV: &str = "AETHER_GATEWAY_DIRECT_H2C_FAST_PATH"; const DIRECT_H2C_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_H2C_CLIENT_SHARDS"; const DIRECT_H2C_POOL_MAX_IDLE_PER_HOST_ENV: &str = @@ -106,9 +125,14 @@ const DEFAULT_DIRECT_REQWEST_SYNC_WARM_CLIENTS: usize = 4; const MAX_DIRECT_REQWEST_SYNC_WARM_CLIENTS: usize = 16; const MAX_DIRECT_H2C_CLIENT_SHARDS: usize = 512; const MAX_DIRECT_REQWEST_H2_CLIENT_SHARDS: usize = 2048; +// This bounds distinct cached transport configurations, not request concurrency, +// HTTP/2 streams, or the number of clients/shards within an active entry. +const DEFAULT_DIRECT_REQWEST_CACHE_MAX_ENTRIES: usize = 1024; +const MAX_DIRECT_REQWEST_CACHE_MAX_ENTRIES: usize = 16_384; type DirectHyperH2cRequestBody = Full; -type DirectHyperH2cClient = HyperLegacyClient; +type DirectHyperH2cClient = + HyperLegacyClient, DirectHyperH2cRequestBody>; type DirectHyperH2cSender = HyperH2cSendRequest; type DirectHyperH2cSenderCacheCell = TokioOnceCell>; @@ -117,10 +141,9 @@ struct DirectReqwestClientCacheKey { upstream_origin: Option, pool_partition: Option, connect_timeout_ms: Option, - proxy_url: Option, + proxy_digest: Option, follow_redirects: bool, http1_only: bool, - accept_invalid_certs: bool, transport_profile: Option, } @@ -146,6 +169,7 @@ struct DirectReqwestClientCacheEntry { next: AtomicU64, target_len: usize, warming: bool, + last_used: u64, } impl DirectReqwestClientCacheEntry { @@ -155,6 +179,7 @@ impl DirectReqwestClientCacheEntry { next: AtomicU64::new(0), target_len: target_len.max(1), warming, + last_used: next_direct_reqwest_client_cache_clock(), } } @@ -177,6 +202,10 @@ impl DirectReqwestClientCacheEntry { fn should_warm(&self) -> bool { self.clients.len() < self.target_len && !self.warming } + + fn touch(&mut self) { + self.last_used = next_direct_reqwest_client_cache_clock(); + } } struct DirectHyperH2cClientCacheEntry { @@ -345,6 +374,8 @@ static DIRECT_REQWEST_CLIENT_CACHE: LazyLock< StdMutex>, > = LazyLock::new(|| StdMutex::new(HashMap::new())); +static DIRECT_REQWEST_CLIENT_CACHE_CLOCK: AtomicU64 = AtomicU64::new(0); + static DIRECT_H2C_CLIENT_CACHE: LazyLock< StdMutex>, > = LazyLock::new(|| StdMutex::new(HashMap::new())); @@ -382,6 +413,7 @@ struct DirectReqwestClientCacheMetrics { http1_selections: AtomicU64, h2c_selections: AtomicU64, auto_selections: AtomicU64, + evictions: AtomicU64, } static DIRECT_REQWEST_CLIENT_CACHE_METRICS: LazyLock = @@ -410,6 +442,325 @@ struct DirectHyperH2cSenderCacheMetrics { static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock = LazyLock::new(DirectHyperH2cSenderCacheMetrics::default); +/// DNS resolver used for direct provider connections. +/// +/// Provider endpoint URLs are frequently user/configuration supplied. The +/// platform resolver may return a different answer on every lookup, so merely +/// checking a URL's host (or resolving it once before constructing a client) +/// is not sufficient to prevent DNS rebinding. This resolver validates every +/// answer at the point reqwest/wreq asks for it. Explicit loopback targets are +/// retained for the supported local-provider workflow, but a hostname that is +/// not itself `localhost` can never resolve to a loopback/private address. +#[derive(Debug, Clone, Copy, Default)] +struct ExecutionSafeDnsResolver; + +/// Resolver adapter for the legacy Hyper client retained for compatibility +/// with the non-fast-path H2C cache. Keep this path subject to the same +/// private-address and rebinding checks as reqwest/wreq clients. +#[derive(Debug, Clone, Copy, Default)] +struct ExecutionSafeHyperDnsResolver; + +// Local DNS interception tools may use RFC 2544's 198.18.0.0/15 range for +// synthetic answers. This exception is deliberately an allowlist rather +// than a property of the address range itself: a custom provider hostname +// must not be able to turn a local synthetic mapping into an SSRF primitive. +// Keep this list limited to origins that Aether constructs as built-in +// provider/model-fetch targets. In particular, do not use a +// suffix match for ordinary hosts (for example, `evil.chatgpt.com`). +const TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS: &[&str] = &[ + "aiplatform.googleapis.com", + "antigravity.googleapis.com", + "api.openai.com", + "api.anthropic.com", + "api.deepseek.com", + "chatgpt.com", + "cloudcode-pa.googleapis.com", + "daily-cloudcode-pa.googleapis.com", + "daily-cloudcode-pa.sandbox.googleapis.com", + "dashscope.aliyuncs.com", + "generativelanguage.googleapis.com", + "grok.com", + "open.bigmodel.cn", + "q.us-iso-east-1.c2s.ic.gov", + "q.us-isob-east-1.sc2s.sgov.gov", + "q.us-isof-east-1.csp.hci.ic.gov", + "q.us-isof-south-1.csp.hci.ic.gov", + "server.codeium.com", +]; + +const TRUSTED_EXECUTION_VERTEX_DNS_REGIONS: &[&str] = &[ + "africa-south1", + "asia-east1", + "asia-east2", + "asia-northeast1", + "asia-northeast2", + "asia-northeast3", + "asia-south1", + "asia-south2", + "asia-southeast1", + "asia-southeast2", + "australia-southeast1", + "australia-southeast2", + "europe-central2", + "europe-north1", + "europe-southwest1", + "europe-west1", + "europe-west2", + "europe-west3", + "europe-west4", + "europe-west6", + "europe-west8", + "europe-west9", + "europe-west10", + "europe-west12", + "me-central1", + "me-central2", + "me-west1", + "northamerica-northeast1", + "northamerica-northeast2", + "southamerica-east1", + "southamerica-west1", + "us-central1", + "us-east1", + "us-east4", + "us-east5", + "us-south1", + "us-west1", + "us-west2", + "us-west3", + "us-west4", +]; + +const TRUSTED_EXECUTION_AWS_DNS_REGIONS: &[&str] = &[ + "af-south-1", + "ap-east-1", + "ap-northeast-1", + "ap-northeast-2", + "ap-northeast-3", + "ap-south-1", + "ap-south-2", + "ap-southeast-1", + "ap-southeast-2", + "ap-southeast-3", + "ap-southeast-4", + "ca-central-1", + "ca-west-1", + "eu-central-1", + "eu-central-2", + "eu-north-1", + "eu-south-1", + "eu-south-2", + "eu-west-1", + "eu-west-2", + "eu-west-3", + "il-central-1", + "me-central-1", + "me-south-1", + "mx-central-1", + "sa-east-1", + "us-east-1", + "us-east-2", + "us-gov-east-1", + "us-gov-west-1", + "us-west-1", + "us-west-2", +]; + +static EXECUTION_EXTRA_TRUSTED_DNS_HOSTS: LazyLock>> = + LazyLock::new(|| StdRwLock::new(BTreeSet::new())); + +pub(crate) fn refresh_execution_extra_trusted_dns_hosts(value: Option<&Value>) { + let hosts = value + .cloned() + .and_then(|value| { + aether_admin::system::normalize_execution_extra_trusted_dns_hosts_config_value(value) + .ok() + }) + .and_then(|value| { + value.as_array().map(|hosts| { + hosts + .iter() + .filter_map(Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }) + }) + .unwrap_or_default(); + + if let Ok(mut current) = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS.write() { + *current = hosts; + } +} + +/// Return whether `host` is one of the fixed provider origins for which a +/// local RFC-2544 synthetic answer can be accepted. The resolver receives only +/// a hostname (not the URL scheme/path), so all policy that can be expressed +/// here is intentionally host based. URL validation still requires HTTPS for +/// non-loopback upstreams before this resolver is used. +fn execution_host_allows_benchmarking_dns_answer(host: &str) -> bool { + let extra_hosts = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS + .read() + .map(|hosts| hosts.clone()) + .unwrap_or_default(); + execution_host_allows_benchmarking_dns_answer_with_extra_hosts(host, &extra_hosts) +} + +fn execution_host_allows_benchmarking_dns_answer_with_extra_hosts( + host: &str, + extra_hosts: &BTreeSet, +) -> bool { + let host = host.trim().trim_end_matches('.').to_ascii_lowercase(); + if extra_hosts.contains(&host) + || TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS + .iter() + .any(|trusted| *trusted == host) + { + return true; + } + + // Vertex service-account requests use `-aiplatform.googleapis.com`. + // Keep this compatibility exception limited to known provider regions. + if let Some(region) = host.strip_suffix("-aiplatform.googleapis.com") { + return TRUSTED_EXECUTION_VERTEX_DNS_REGIONS.contains(®ion); + } + + // Kiro uses a small, fixed set of regional service origins. Match each + // supported AWS partition explicitly; never use a broad suffix check that + // could accept an attacker-controlled subdomain. + matches_regional_service_host(&host, "q", ".amazonaws.com") + || matches_regional_service_host(&host, "q-fips", ".amazonaws.com") + || matches_regional_service_host(&host, "codewhisperer", ".amazonaws.com") + || matches_regional_service_host(&host, "oidc", ".amazonaws.com") + || matches_regional_service_host(&host, "prod", ".auth.desktop.kiro.dev") +} + +fn matches_regional_service_host(host: &str, service: &str, suffix: &str) -> bool { + let Some(region) = host + .strip_prefix(service) + .and_then(|value| value.strip_prefix('.')) + .and_then(|value| value.strip_suffix(suffix)) + else { + return false; + }; + TRUSTED_EXECUTION_AWS_DNS_REGIONS.contains(®ion) +} + +fn dns_host_explicitly_allows_loopback(host: &str) -> bool { + let host = host.trim_end_matches('.'); + host.eq_ignore_ascii_case("localhost") + || host + .parse::() + .map(|ip| ip.is_loopback()) + .unwrap_or(false) +} + +fn validate_execution_dns_answers( + host: &str, + addresses: Vec, +) -> Result, std::io::Error> { + validate_execution_dns_answers_with_policy(host, addresses, true) +} + +fn validate_execution_dns_answers_with_policy( + host: &str, + addresses: Vec, + allow_trusted_benchmarking_dns_answer: bool, +) -> Result, std::io::Error> { + if addresses.is_empty() { + return Err(std::io::Error::new( + std::io::ErrorKind::NotFound, + "upstream DNS resolution returned no addresses", + )); + } + + let allows_loopback = dns_host_explicitly_allows_loopback(host); + let allows_benchmarking_dns_answer = allow_trusted_benchmarking_dns_answer + && execution_host_allows_benchmarking_dns_answer(host); + let unsafe_answer = addresses.iter().any(|address| { + if allows_loopback { + !address.ip().is_loopback() + } else { + is_private_or_reserved_ip(address.ip()) + && !(allows_benchmarking_dns_answer && is_ipv4_benchmarking_fake_ip(address.ip())) + } + }); + if unsafe_answer { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + "upstream DNS resolution returned a private or reserved address", + )); + } + + Ok(addresses) +} + +async fn resolve_execution_dns_addresses(host: &str) -> Result, std::io::Error> { + resolve_execution_target_addresses_with_policy(host, 0, true).await +} + +async fn resolve_execution_target_addresses_with_policy( + host: &str, + port: u16, + allow_trusted_benchmarking_dns_answer: bool, +) -> Result, std::io::Error> { + let addresses = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await? + }; + validate_execution_dns_answers_with_policy( + host, + addresses, + allow_trusted_benchmarking_dns_answer, + ) +} + +impl reqwest::dns::Resolve for ExecutionSafeDnsResolver { + fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving { + let host = name.as_str().to_owned(); + Box::pin(async move { + let addresses = resolve_execution_dns_addresses(&host) + .await + .map_err(|error| Box::new(error) as Box)?; + Ok(Box::new(addresses.into_iter()) as reqwest::dns::Addrs) + }) + } +} + +impl wreq::dns::Resolve for ExecutionSafeDnsResolver { + fn resolve(&self, name: wreq::dns::Name) -> wreq::dns::Resolving { + let host = name.as_str().to_owned(); + Box::pin(async move { + let addresses = resolve_execution_dns_addresses(&host) + .await + .map_err(|error| Box::new(error) as Box)?; + Ok(Box::new(addresses.into_iter()) as wreq::dns::Addrs) + }) + } +} + +impl tower::Service + for ExecutionSafeHyperDnsResolver +{ + type Response = std::vec::IntoIter; + type Error = std::io::Error; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, name: hyper_util::client::legacy::connect::dns::Name) -> Self::Future { + let host = name.as_str().to_owned(); + Box::pin(async move { + resolve_execution_dns_addresses(&host) + .await + .map(|addrs| addrs.into_iter()) + }) + } +} + #[derive(Debug, Clone, Default)] pub struct DirectH2cSenderPrewarmReport { pub requested_urls: u64, @@ -466,7 +817,7 @@ pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String { detail.push(']'); } - detail + sanitize_error_detail(&detail) } fn sanitize_upstream_request_error_detail(detail: &str, upstream_url: &str) -> (String, String) { @@ -476,8 +827,21 @@ fn sanitize_upstream_request_error_detail(detail: &str, upstream_url: &str) -> ( fn sanitize_upstream_url_text(upstream_url: &str) -> String { if let Ok(mut parsed_url) = reqwest::Url::parse(upstream_url) { + // URL userinfo can contain proxy or upstream credentials. reqwest's + // error chain may include the original URL, so remove it alongside + // query and fragment data before the error crosses a trust boundary. + let _ = parsed_url.set_username(""); + let _ = parsed_url.set_password(None); parsed_url.set_query(None); parsed_url.set_fragment(None); + let private_literal = match parsed_url.host() { + Some(url::Host::Ipv4(address)) => is_private_or_reserved_ip(IpAddr::V4(address)), + Some(url::Host::Ipv6(address)) => is_private_or_reserved_ip(IpAddr::V6(address)), + _ => false, + }; + if private_literal { + let _ = parsed_url.set_host(Some("redacted.invalid")); + } return parsed_url.to_string(); } @@ -485,7 +849,89 @@ fn sanitize_upstream_url_text(upstream_url: &str) -> String { .char_indices() .find_map(|(offset, character)| matches!(character, '?' | '#').then_some(offset)) .unwrap_or(upstream_url.len()); - upstream_url[..suffix_offset].to_string() + let mut sanitized = upstream_url[..suffix_offset].to_string(); + // Keep malformed URL diagnostics useful without carrying userinfo across + // the boundary. All indices here are ASCII delimiters discovered in a + // UTF-8 string, so the range boundaries remain valid. + if let Some(scheme_end) = sanitized.find("://") { + let authority_end = sanitized[scheme_end + 3..] + .find('/') + .map(|offset| scheme_end + 3 + offset) + .unwrap_or(sanitized.len()); + if let Some(at) = sanitized[scheme_end + 3..authority_end].rfind('@') { + let at = scheme_end + 3 + at; + sanitized.replace_range(scheme_end + 3..=at, ""); + } + } + sanitized +} + +fn sanitize_error_detail(detail: &str) -> String { + let mut sanitized = String::with_capacity(detail.len().min(MAX_UPSTREAM_ERROR_DETAIL_BYTES)); + for (index, token) in detail.split_whitespace().enumerate() { + if index > 0 { + sanitized.push(' '); + } + sanitized.push_str(&sanitize_error_token(token)); + } + if sanitized.len() > MAX_UPSTREAM_ERROR_DETAIL_BYTES { + let mut end = MAX_UPSTREAM_ERROR_DETAIL_BYTES; + while !sanitized.is_char_boundary(end) { + end = end.saturating_sub(1); + } + sanitized.truncate(end); + sanitized.push_str("..."); + } + sanitized +} + +fn sanitize_error_token(token: &str) -> String { + let Some(scheme_offset) = token.find("://") else { + return token.to_string(); + }; + let mut start = scheme_offset; + while start > 0 { + let previous = token[..start] + .chars() + .next_back() + .expect("non-empty URL prefix should contain a character"); + if matches!( + previous, + '(' | '[' | '{' | '"' | '\'' | '=' | ';' | ',' | ':' + ) { + break; + } + start -= previous.len_utf8(); + } + let mut end = token.len(); + while end > start { + let last = token.as_bytes()[end - 1] as char; + if matches!(last, ')' | ']' | '}' | '"' | '\'' | ',' | ';') { + end -= 1; + } else { + break; + } + } + let candidate = &token[start..end]; + let Ok(parsed) = reqwest::Url::parse(candidate) else { + return token.to_string(); + }; + let sanitized = sanitize_upstream_url_text(parsed.as_str()); + let mut result = String::with_capacity(token.len()); + result.push_str(&token[..start]); + result.push_str(&sanitized); + result.push_str(&token[end..]); + result +} + +/// Return a bounded diagnostic suitable for scheduler/usage records and +/// structured logs. `ExecutionRuntimeTransportError` keeps rich dynamic +/// details for local control flow, but its `Display` implementation is also +/// used by older call sites that persist the message. Route those boundaries +/// through the same URL/query/credential sanitizer as the custom `Debug` +/// implementation so a future error constructor cannot leak request secrets. +pub(crate) fn safe_transport_error_message(error: &ExecutionRuntimeTransportError) -> String { + sanitize_error_detail(&error.to_string()) } pub(crate) fn format_wreq_upstream_request_error(err: &wreq::Error) -> String { @@ -535,7 +981,7 @@ pub(crate) fn format_wreq_upstream_request_error(err: &wreq::Error) -> String { detail.push(']'); } - detail + sanitize_error_detail(&detail) } pub(crate) fn format_hyper_error_chain(err: &dyn std::error::Error) -> String { @@ -549,54 +995,157 @@ pub(crate) fn format_hyper_error_chain(err: &dyn std::error::Error) -> String { } source = cause.source(); } - detail + sanitize_error_detail(&detail) } -#[derive(Debug, Error)] +#[derive(Error)] pub(crate) enum ExecutionRuntimeTransportError { #[error("request body must contain json_body or body_bytes_b64")] RequestBodyRequired, + #[error("request body must not contain both json_body and body_bytes_b64")] + RequestBodyAmbiguous, #[error("request body base64 is invalid: {0}")] BodyDecode(base64::DecodeError), - #[error("request content-encoding is not supported: {0}")] + #[error("request body exceeds {limit_bytes} decoded bytes")] + BodyTooLarge { limit_bytes: usize }, + #[error("request content-encoding is not supported: {}", sanitize_error_detail(.0))] UnsupportedContentEncoding(String), #[error("proxy execution is not supported")] ProxyUnsupported, - #[error("invalid method: {0}")] + #[error("invalid method: {}", sanitize_error_detail(&.0.to_string()))] InvalidMethod(#[from] http::method::InvalidMethod), - #[error("invalid upstream header name: {0}")] + #[error("invalid upstream header name: {}", sanitize_error_detail(.0))] InvalidHeaderName(String), - #[error("invalid upstream header value for {0}")] + #[error("invalid upstream header value for {}", sanitize_error_detail(.0))] InvalidHeaderValue(String), - #[error("invalid proxy configuration: {0}")] - InvalidProxy(reqwest::Error), - #[error("unsupported transport profile backend: {0}")] + #[error("invalid proxy configuration")] + InvalidProxy(#[source] reqwest::Error), + #[error("unsupported transport profile backend: {}", sanitize_error_detail(.0))] UnsupportedTransportProfile(String), - #[error("failed to encode request body: {0}")] - BodyEncode(serde_json::Error), - #[error("failed to build HTTP client: {0}")] - ClientBuild(reqwest::Error), - #[error("failed to build browser impersonation HTTP client: {0}")] - BrowserClientBuild(wreq::Error), - #[error("browser impersonation response body failed: {0}")] + #[error("failed to encode request body")] + BodyEncode(#[source] serde_json::Error), + #[error("failed to build HTTP client")] + ClientBuild(#[source] reqwest::Error), + #[error("failed to build browser impersonation HTTP client")] + BrowserClientBuild(#[source] wreq::Error), + #[error("browser impersonation response body failed: {}", sanitize_error_detail(.0))] BrowserBody(String), - #[error("{message}")] + #[error("{}", sanitize_error_detail(message))] UpstreamHttpStatus { status_code: u16, message: String }, - #[error("failed to execute upstream request: {0}")] + #[error("failed to execute upstream request: {}", sanitize_error_detail(.0))] UpstreamRequest(String), #[error("upstream response {phase} body exceeds {limit_bytes} bytes")] UpstreamResponseTooLarge { phase: UpstreamResponseBodyPhase, limit_bytes: usize, }, - #[error("failed to decode upstream response body with content-encoding {encoding}: {message}")] + #[error( + "failed to decode upstream response body with content-encoding {}: {}", + sanitize_error_detail(encoding), + sanitize_error_detail(message) + )] UpstreamResponseDecode { encoding: String, message: String }, - #[error("hub relay request failed: {0}")] + #[error("hub relay request failed: {}", sanitize_error_detail(.0))] RelayError(String), #[error("upstream response is not valid JSON: {0}")] InvalidJson(serde_json::Error), } +// `reqwest::Error` and `wreq::Error` retain the URL associated with a failed +// request. Their derived `Debug` implementations therefore may include +// proxy credentials or query-string tokens. This error is logged with +// structured `?error` fields in a few execution paths, so both `Debug` and +// `Display` must be safe if a caller accidentally crosses that boundary. +// Dynamic details are passed through the same URL-aware, bounded sanitizer +// used by the upstream request formatters. +impl std::fmt::Debug for ExecutionRuntimeTransportError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::RequestBodyRequired => formatter.write_str("RequestBodyRequired"), + Self::RequestBodyAmbiguous => formatter.write_str("RequestBodyAmbiguous"), + Self::BodyDecode(error) => formatter + .debug_tuple("BodyDecode") + .field(&sanitize_error_detail(&error.to_string())) + .finish(), + Self::BodyTooLarge { limit_bytes } => formatter + .debug_struct("BodyTooLarge") + .field("limit_bytes", limit_bytes) + .finish(), + Self::UnsupportedContentEncoding(encoding) => formatter + .debug_tuple("UnsupportedContentEncoding") + .field(&sanitize_error_detail(encoding)) + .finish(), + Self::ProxyUnsupported => formatter.write_str("ProxyUnsupported"), + Self::InvalidMethod(error) => formatter + .debug_tuple("InvalidMethod") + .field(&sanitize_error_detail(&error.to_string())) + .finish(), + Self::InvalidHeaderName(name) => formatter + .debug_tuple("InvalidHeaderName") + .field(&sanitize_error_detail(name)) + .finish(), + Self::InvalidHeaderValue(name) => formatter + .debug_tuple("InvalidHeaderValue") + .field(&sanitize_error_detail(name)) + .finish(), + Self::InvalidProxy(error) => formatter + .debug_tuple("InvalidProxy") + .field(&format_upstream_request_error(error)) + .finish(), + Self::UnsupportedTransportProfile(profile) => formatter + .debug_tuple("UnsupportedTransportProfile") + .field(&sanitize_error_detail(profile)) + .finish(), + Self::BodyEncode(error) => formatter + .debug_tuple("BodyEncode") + .field(&sanitize_error_detail(&error.to_string())) + .finish(), + Self::ClientBuild(error) => formatter + .debug_tuple("ClientBuild") + .field(&format_upstream_request_error(error)) + .finish(), + Self::BrowserClientBuild(error) => formatter + .debug_tuple("BrowserClientBuild") + .field(&format_wreq_upstream_request_error(error)) + .finish(), + Self::BrowserBody(detail) => formatter + .debug_tuple("BrowserBody") + .field(&sanitize_error_detail(detail)) + .finish(), + Self::UpstreamHttpStatus { + status_code, + message, + } => formatter + .debug_struct("UpstreamHttpStatus") + .field("status_code", status_code) + .field("message", &sanitize_error_detail(message)) + .finish(), + Self::UpstreamRequest(detail) => formatter + .debug_tuple("UpstreamRequest") + .field(&sanitize_error_detail(detail)) + .finish(), + Self::UpstreamResponseTooLarge { phase, limit_bytes } => formatter + .debug_struct("UpstreamResponseTooLarge") + .field("phase", phase) + .field("limit_bytes", limit_bytes) + .finish(), + Self::UpstreamResponseDecode { encoding, message } => formatter + .debug_struct("UpstreamResponseDecode") + .field("encoding", &sanitize_error_detail(encoding)) + .field("message", &sanitize_error_detail(message)) + .finish(), + Self::RelayError(detail) => formatter + .debug_tuple("RelayError") + .field(&sanitize_error_detail(detail)) + .finish(), + Self::InvalidJson(error) => formatter + .debug_tuple("InvalidJson") + .field(&sanitize_error_detail(&error.to_string())) + .finish(), + } + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum UpstreamResponseBodyPhase { Wire, @@ -617,16 +1166,19 @@ pub(crate) fn with_upstream_response_body_limit( limit_bytes: usize, ) -> ExecutionPlan { let mut bounded_plan = plan.clone(); + apply_upstream_response_body_limit(&mut bounded_plan, limit_bytes); bounded_plan - .headers +} + +pub(crate) fn apply_upstream_response_body_limit(plan: &mut ExecutionPlan, limit_bytes: usize) { + plan.headers .retain(|name, _| !name.eq_ignore_ascii_case(EXECUTION_RESPONSE_BODY_LIMIT_HEADER)); - bounded_plan.headers.insert( + plan.headers.insert( EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_string(), normalize_scoped_response_body_limit(limit_bytes) .unwrap_or(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES) .to_string(), ); - bounded_plan } pub(crate) fn execution_plan_response_body_limit_bytes(plan: &ExecutionPlan) -> usize { @@ -688,6 +1240,140 @@ pub(crate) fn append_upstream_response_body_chunk_with_limit( Ok(()) } +/// Return the maximum base64 text length that can decode to at most +/// `decoded_limit` bytes. This is intentionally checked before invoking the +/// base64 decoder, whose allocation is based on the input text length. +pub(crate) fn maximum_base64_len_for_decoded_limit(decoded_limit: usize) -> usize { + decoded_limit + .checked_add(2) + .and_then(|value| value.checked_div(3)) + .and_then(|value| value.checked_mul(4)) + .unwrap_or(usize::MAX) +} + +/// Decode a body carried in an execution plan/result only after enforcing a +/// decoded-size bound. Both representations are checked: the encoded check +/// prevents an attacker-controlled allocation, while the decoded check covers +/// padding and decoder edge cases. +pub(crate) fn decode_base64_body_with_limit( + body_base64: &str, + decoded_limit: usize, +) -> Result, ExecutionRuntimeTransportError> { + if body_base64.len() > maximum_base64_len_for_decoded_limit(decoded_limit) { + return Err(ExecutionRuntimeTransportError::BodyTooLarge { + limit_bytes: decoded_limit, + }); + } + + let bytes = base64::engine::general_purpose::STANDARD + .decode(body_base64) + .map_err(ExecutionRuntimeTransportError::BodyDecode)?; + if bytes.len() > decoded_limit { + return Err(ExecutionRuntimeTransportError::BodyTooLarge { + limit_bytes: decoded_limit, + }); + } + Ok(bytes) +} + +struct JsonSerializedSizeLimiter { + remaining: usize, +} + +impl Write for JsonSerializedSizeLimiter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.remaining { + return Err(std::io::Error::other("serialized JSON exceeds limit")); + } + self.remaining -= bytes.len(); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +pub(crate) fn json_value_fits_serialized_limit(value: &Value, limit_bytes: usize) -> bool { + serde_json::to_writer( + JsonSerializedSizeLimiter { + remaining: limit_bytes, + }, + value, + ) + .is_ok() +} + +/// Bound the JSON envelope used by the test/compatibility remote execution +/// runtime. A result can contain a raw JSON representation and a base64 wire +/// representation at the same time, so the limit is larger than either body +/// limit. It remains capped even when the raw-body cap is explicitly +/// disabled (`usize::MAX`). +pub(crate) fn execution_result_envelope_limit_bytes(decoded_body_limit: usize) -> usize { + maximum_base64_len_for_decoded_limit(decoded_body_limit) + .saturating_add(decoded_body_limit) + .saturating_add(EXECUTION_RESULT_ENVELOPE_METADATA_BYTES) + .min(MAX_EXECUTION_RESULT_ENVELOPE_BYTES) +} + +/// Serialize a JSON body without allowing serde_json to grow an unbounded +/// temporary `Vec`. The value itself is already owned by the execution plan; +/// this bound covers the wire representation that will be sent upstream. +pub(crate) fn serialize_json_body_with_limit( + body: &Value, + limit_bytes: usize, +) -> Result, ExecutionRuntimeTransportError> { + serialize_serializable_with_limit(body, limit_bytes) +} + +pub(crate) fn serialize_serializable_with_limit( + value: &T, + limit_bytes: usize, +) -> Result, ExecutionRuntimeTransportError> { + let mut writer = LimitedJsonWriter::new(limit_bytes); + match serde_json::to_writer(&mut writer, value) { + Ok(()) => Ok(writer.bytes), + Err(_error) if writer.exceeded => { + Err(ExecutionRuntimeTransportError::BodyTooLarge { limit_bytes }) + } + Err(error) => Err(ExecutionRuntimeTransportError::BodyEncode(error)), + } +} + +struct LimitedJsonWriter { + bytes: Vec, + limit: usize, + exceeded: bool, +} + +impl LimitedJsonWriter { + fn new(limit: usize) -> Self { + Self { + bytes: Vec::with_capacity(limit.min(16 * 1024)), + limit, + exceeded: false, + } + } +} + +impl Write for LimitedJsonWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.limit.saturating_sub(self.bytes.len()) { + self.exceeded = true; + return Err(std::io::Error::new( + std::io::ErrorKind::WriteZero, + "json body exceeds configured limit", + )); + } + self.bytes.extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + #[derive(Debug, Serialize)] struct RelayRequestMeta { provider_id: String, @@ -718,7 +1404,6 @@ pub(crate) struct DirectSyncExecutionRuntime; pub(crate) struct ExecutionTransportControls { follow_redirects: Option, http1_only: bool, - accept_invalid_certs: bool, } #[derive(Debug, Clone, Copy)] @@ -740,6 +1425,9 @@ pub(crate) struct DirectUpstreamStreamExecution { pub(crate) candidate_id: Option, pub(crate) status_code: u16, pub(crate) headers: BTreeMap, + /// The upstream length is retained for stream classification only. The + /// hop-by-hop header itself remains filtered from the client-facing map. + pub(crate) upstream_content_length: Option, pub(crate) provider_api_format: String, pub(crate) stream_summary_report_context: Value, pub(crate) prefetched_body: VecDeque>, @@ -857,6 +1545,7 @@ impl DirectSyncExecutionRuntime { started_at.elapsed().as_millis() as u64, ); let status_code = response.status_code(); + let upstream_content_length = response.content_length(); let headers = response.headers(); let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms(); @@ -867,6 +1556,7 @@ impl DirectSyncExecutionRuntime { candidate_id: plan.candidate_id.clone(), status_code, headers, + upstream_content_length, provider_api_format: plan.provider_api_format.clone(), stream_summary_report_context, prefetched_body: VecDeque::new(), @@ -917,7 +1607,7 @@ pub(crate) async fn execute_sync_plan_with_report_context( if resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).is_some() { return execute_sync_plan_via_local_tunnel(state, plan, report_context) .await - .map_err(|err| GatewayError::Internal(err.to_string())); + .map_err(|err| GatewayError::Internal(safe_transport_error_message(&err))); } match super::grok::maybe_execute_grok_sync(plan, report_context).await { @@ -928,7 +1618,7 @@ pub(crate) async fn execute_sync_plan_with_report_context( Ok(None) => {} Err(err) => { record_manual_proxy_request_failure(state, plan).await; - return Err(GatewayError::Internal(err.to_string())); + return Err(GatewayError::Internal(safe_transport_error_message(&err))); } } @@ -936,7 +1626,7 @@ pub(crate) async fn execute_sync_plan_with_report_context( match maybe_execute_windsurf_sync(state, plan, None).await { Ok(Some(result)) => return Ok(result), Ok(None) => {} - Err(err) => return Err(GatewayError::Internal(err.to_string())), + Err(err) => return Err(GatewayError::Internal(safe_transport_error_message(&err))), } let state_for_response_started = state.clone(); match DirectSyncExecutionRuntime::new() @@ -962,7 +1652,7 @@ pub(crate) async fn execute_sync_plan_with_report_context( } Err(err) => { record_manual_proxy_request_failure(state, plan).await; - Err(GatewayError::Internal(err.to_string())) + Err(GatewayError::Internal(safe_transport_error_message(&err))) } } } @@ -975,6 +1665,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel( return Ok(None); }; + validate_execution_upstream_url(plan.url.as_str())?; if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) { return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail)); } @@ -999,6 +1690,11 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel( .await .map_err(ExecutionRuntimeTransportError::RelayError)?; let status_code = response.status(); + let upstream_content_length = response + .headers() + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("content-length")) + .and_then(|(_, value)| value.trim().parse::().ok()); let headers = collect_tunnel_response_headers(response.headers()); let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms(); @@ -1007,6 +1703,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel( candidate_id: plan.candidate_id.clone(), status_code, headers, + upstream_content_length, provider_api_format: plan.provider_api_format.clone(), stream_summary_report_context: build_stream_summary_report_context(plan), prefetched_body: VecDeque::new(), @@ -1061,11 +1758,14 @@ async fn record_manual_proxy_traffic( dns_failures_delta: i64, stream_errors_delta: i64, ) { - let Some(node_id) = manual_proxy_node_id(plan.proxy.as_ref()) else { + let Some((node_id, expected_tunnel_generation)) = + manual_proxy_node_binding(plan.proxy.as_ref()) + else { return; }; let mutation = ProxyNodeTrafficMutation { node_id: node_id.clone(), + expected_tunnel_generation: Some(expected_tunnel_generation), total_requests_delta, failed_requests_delta, dns_failures_delta, @@ -1081,17 +1781,26 @@ async fn record_manual_proxy_traffic( } } -fn manual_proxy_node_id(proxy: Option<&ProxySnapshot>) -> Option { +fn manual_proxy_node_binding(proxy: Option<&ProxySnapshot>) -> Option<(String, String)> { let proxy = proxy?; if proxy.enabled == Some(false) || resolve_tunnel_node_id(Some(proxy)).is_some() { return None; } - proxy + let node_id = proxy .node_id .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .map(ToOwned::to_owned)?; + let expected_tunnel_generation = proxy + .extra + .as_ref() + .and_then(|extra| extra.get(PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned)?; + Some((node_id, expected_tunnel_generation)) } async fn execute_sync_plan_via_local_tunnel( @@ -1114,6 +1823,7 @@ async fn execute_sync_plan_via_local_tunnel_inner( let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| { ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string()) })?; + validate_execution_upstream_url(plan.url.as_str())?; if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) { return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail)); } @@ -1311,6 +2021,7 @@ async fn send_request_inner( body_bytes: Vec, apply_request_total_timeout: bool, ) -> Result { + validate_execution_upstream_url(plan.url.as_str())?; if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) { return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail)); } @@ -1430,12 +2141,21 @@ impl DirectHttpResponse { } } + pub(crate) fn content_length(&self) -> Option { + let value = match self { + DirectHttpResponse::Reqwest(response) => response.headers().get("content-length"), + DirectHttpResponse::HyperH2c(response) => response.headers().get("content-length"), + DirectHttpResponse::BrowserWreq(response) => response.headers().get("content-length"), + }?; + value.to_str().ok()?.trim().parse::().ok() + } + pub(crate) async fn bytes(self) -> Result { self.bytes_with_limit(crate::headers::max_internal_buffered_body_bytes()) .await } - async fn bytes_with_limit( + pub(crate) async fn bytes_with_limit( self, response_body_limit_bytes: usize, ) -> Result { @@ -1653,7 +2373,6 @@ fn direct_h2c_fast_path_applies( if !direct_h2c_fast_path_enabled() || !plan.stream || transport_controls.http1_only - || transport_controls.accept_invalid_certs || plan.proxy.is_some() || !transport_profile_h2c_prior_knowledge(plan.transport_profile.as_ref()) { @@ -1741,7 +2460,7 @@ async fn prewarm_direct_h2c_sender_cache_urls( .prewarm_failed .fetch_add(1, Ordering::Relaxed); if first_error.is_none() { - first_error = Some(err.to_string()); + first_error = Some(safe_transport_error_message(&err)); } } } @@ -1899,10 +2618,14 @@ fn direct_h2c_client_cache_key( request_url: &str, timeouts: Option<&aether_contracts::ExecutionTimeouts>, ) -> Result { + if reqwest::Url::parse(request_url).is_err() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "invalid h2c upstream origin".to_string(), + )); + } + validate_execution_upstream_url(request_url)?; let upstream_origin = direct_reqwest_upstream_origin(request_url).ok_or_else(|| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "invalid h2c upstream origin: {request_url}" - )) + ExecutionRuntimeTransportError::UpstreamRequest("invalid h2c upstream origin".to_string()) })?; Ok(DirectHyperH2cClientCacheKey { upstream_origin, @@ -1961,52 +2684,56 @@ async fn connect_direct_h2c_sender_on_current_runtime( cache_key: &DirectHyperH2cClientCacheKey, ) -> Result { let upstream = reqwest::Url::parse(&cache_key.upstream_origin).map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "invalid h2c upstream origin {}: {err}", - cache_key.upstream_origin - )) + tracing::debug!(error = %err, "invalid direct h2c upstream origin"); + ExecutionRuntimeTransportError::UpstreamRequest("invalid h2c upstream origin".to_string()) })?; let host = upstream.host_str().ok_or_else(|| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "missing h2c upstream host: {}", - cache_key.upstream_origin - )) + ExecutionRuntimeTransportError::UpstreamRequest("missing h2c upstream host".to_string()) })?; let port = upstream.port_or_known_default().ok_or_else(|| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "missing h2c upstream port: {}", - cache_key.upstream_origin - )) + ExecutionRuntimeTransportError::UpstreamRequest("missing h2c upstream port".to_string()) })?; - let connect = TcpStream::connect((host, port)); + let addresses = resolve_execution_target_addresses_with_policy(host, port, true) + .await + .map_err(|error| { + let message = if error.kind() == std::io::ErrorKind::PermissionDenied { + "h2c upstream DNS resolution returned a private or reserved address" + } else { + "h2c upstream DNS resolution failed" + }; + ExecutionRuntimeTransportError::UpstreamRequest(message.to_string()) + })?; + // Passing concrete socket addresses prevents TcpStream from performing a + // second hostname lookup after the validated DNS answer. + let connect = TcpStream::connect(addresses.as_slice()); let stream = if let Some(timeout_ms) = cache_key.connect_timeout_ms { let timeout = Duration::from_millis(timeout_ms); tokio::time::timeout(timeout, connect) .await .map_err(|_| { - ExecutionRuntimeTransportError::UpstreamRequest(stream_first_byte_timeout_message( + ExecutionRuntimeTransportError::UpstreamRequest(direct_h2c_connect_timeout_message( timeout, )) })? .map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to connect h2c upstream {}: {err}", - cache_key.upstream_origin - )) + tracing::debug!(error = %err, "failed to connect direct h2c upstream"); + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to connect h2c upstream".to_string(), + ) })? } else { connect.await.map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to connect h2c upstream {}: {err}", - cache_key.upstream_origin - )) + tracing::debug!(error = %err, "failed to connect direct h2c upstream"); + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to connect h2c upstream".to_string(), + ) })? }; stream.set_nodelay(true).map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to configure h2c upstream socket {}: {err}", - cache_key.upstream_origin - )) + tracing::debug!(error = %err, "failed to configure direct h2c upstream socket"); + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to configure h2c upstream socket".to_string(), + ) })?; let io = TokioIo::new(stream); let mut builder = hyper::client::conn::http2::Builder::new(TokioExecutor::new()); @@ -2079,7 +2806,7 @@ fn cached_direct_h2c_client( fn build_direct_h2c_client_from_cache_key( cache_key: &DirectHyperH2cClientCacheKey, ) -> DirectHyperH2cClient { - let mut connector = HttpConnector::new(); + let mut connector = HttpConnector::new_with_resolver(ExecutionSafeHyperDnsResolver); connector.enforce_http(true); connector.set_nodelay(true); connector.set_connect_timeout(cache_key.connect_timeout_ms.map(Duration::from_millis)); @@ -2205,8 +2932,8 @@ async fn send_via_direct_h2c_fast_path( ); let request_build_started_at = Instant::now(); - let uri = plan.url.parse::().map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!("invalid h2c upstream uri: {err}")) + let uri = plan.url.parse::().map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest("invalid h2c upstream uri".to_string()) })?; let authority = uri .authority() @@ -2329,6 +3056,13 @@ fn direct_h2c_remaining_timeout(deadline: Instant) -> Option { deadline.checked_duration_since(Instant::now()) } +fn direct_h2c_connect_timeout_message(timeout: Duration) -> String { + format!( + "direct h2c upstream connect timeout after {} ms", + timeout.as_millis() + ) +} + async fn send_via_browser_wreq_transport( plan: &ExecutionPlan, method: reqwest::Method, @@ -2373,8 +3107,12 @@ async fn send_via_tunnel_relay( stream_first_byte_timeout: Option, transport_controls: ExecutionTransportControls, ) -> Result { - let client = build_relay_client(plan.timeouts.as_ref())?; - let relay_url = build_relay_url(plan.proxy.as_ref(), node_id); + let relay_url = build_relay_url(plan.proxy.as_ref(), node_id)?; + let (relay_host, relay_addresses) = resolve_relay_target_addresses(&relay_url).await?; + let client = build_relay_client_with_pinned_target( + plan.timeouts.as_ref(), + Some((&relay_host, &relay_addresses)), + )?; let timeout_metadata = resolve_tunnel_timeout_metadata(plan); let timeout_secs = timeout_metadata.legacy_timeout_secs; let envelope = build_relay_envelope( @@ -2406,17 +3144,29 @@ async fn send_via_tunnel_relay( node_id, path = "tunnel_relay", body_bytes_len = body_bytes.len(), - envelope_bytes_len = envelope.len(), + envelope_bytes_len = envelope.body.len(), timeout_secs, follow_redirects = ?transport_controls.follow_redirects, http1_only = transport_controls.http1_only, "gateway execution runtime tunnel relay request prepared" ); - let mut request = client - .request(reqwest::Method::POST, relay_url) - .header(reqwest::header::CONTENT_TYPE, HUB_RELAY_CONTENT_TYPE) - .body(envelope); + let relay_auth = resolve_tunnel_owner_instance_id(plan.proxy.as_ref()) + .map(ToOwned::to_owned) + .unwrap_or_else(tunnel::resolve_tunnel_instance_id); + let relay_auth = tunnel::build_relay_auth_headers_from_environment( + &relay_auth, + node_id, + envelope.metadata_envelope(), + envelope.request_body(), + ) + .map_err(ExecutionRuntimeTransportError::RelayError)?; + let mut request = relay_auth.apply( + client + .request(reqwest::Method::POST, relay_url) + .header(reqwest::header::CONTENT_TYPE, HUB_RELAY_CONTENT_TYPE) + .body(envelope.body), + ); if !plan.stream { if let Some(timeout) = total_timeout { request = request.timeout(timeout); @@ -2472,12 +3222,13 @@ async fn send_via_tunnel_relay( ); } - if let Some(kind) = response + if let Some(raw_kind) = response .headers() .get(HUB_RELAY_ERROR_HEADER) .and_then(|value| value.to_str().ok()) .map(str::to_owned) { + let kind = sanitize_relay_error_kind(&raw_kind); tracing::warn!( request_id = %plan.request_id, provider_id = %plan.provider_id, @@ -2492,37 +3243,61 @@ async fn send_via_tunnel_relay( error_kind = %kind, "gateway execution runtime tunnel relay returned relay error" ); - let response_headers = collect_response_headers(response.headers()); let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan); - let (wire_body, _) = - collect_reqwest_stream_body(response, Instant::now(), None, response_body_limit_bytes) - .await - .map_err(|error| { - ExecutionRuntimeTransportError::RelayError(format!( - "hub relay error: {kind}: bounded error body read failed: {error}" - )) - })?; - let decoded_body = decode_response_body_bytes_with_limit( - &response_headers, - &wire_body, - response_body_limit_bytes, - ) - .map_err(|error| { - ExecutionRuntimeTransportError::RelayError(format!( - "hub relay error: {kind}: bounded error body decode failed: {error}" - )) - })?; - let message = if decoded_body.is_empty() { - format!("hub relay error: {kind}") + let drain_timeout = stream_first_byte_timeout.or(total_timeout); + let drain = + collect_reqwest_stream_body(response, Instant::now(), None, response_body_limit_bytes); + let drain_result = if let Some(timeout) = drain_timeout { + match tokio::time::timeout(timeout, drain).await { + Ok(result) => result, + Err(_) => Err(ExecutionRuntimeTransportError::UpstreamRequest( + "tunnel relay error body drain timeout".to_string(), + )), + } } else { - String::from_utf8_lossy(decoded_body.as_ref()).into_owned() + // There is no configured budget to drain an untrusted error body; + // dropping the response cancels the body stream instead of waiting + // indefinitely for a peer that never terminates it. + drop(drain); + Err(ExecutionRuntimeTransportError::UpstreamRequest( + "tunnel relay error body drain skipped".to_string(), + )) }; - return Err(ExecutionRuntimeTransportError::RelayError(message)); + match drain_result { + Ok((body, _)) => { + // Consume the bounded body so the connection can be reused, + // but never propagate relay/upstream text across the error + // boundary. The body length is sufficient for diagnostics. + tracing::debug!( + error_kind = %kind, + error_body_bytes = body.len(), + "discarded tunnel relay error body" + ); + } + Err(error) => { + tracing::debug!( + error_kind = %kind, + drain_error_kind = %relay_body_drain_error_kind(&error), + "failed to consume tunnel relay error body" + ); + } + } + return Err(ExecutionRuntimeTransportError::RelayError(format!( + "hub relay error: {kind}" + ))); } Ok(response) } +fn relay_body_drain_error_kind(error: &ExecutionRuntimeTransportError) -> &'static str { + match error { + ExecutionRuntimeTransportError::UpstreamResponseTooLarge { .. } => "too_large", + ExecutionRuntimeTransportError::UpstreamRequest(_) => "request", + _ => "unknown", + } +} + async fn send_relay_request( request: reqwest::RequestBuilder, first_byte_timeout: Option, @@ -2530,23 +3305,60 @@ async fn send_relay_request( if let Some(timeout) = first_byte_timeout { return match tokio::time::timeout(timeout, request.send()).await { Ok(Ok(response)) => Ok(response), - Ok(Err(error)) => Err(error.to_string()), + Ok(Err(error)) => Err(format_relay_request_error(&error)), Err(_) => Err("tunnel relay first byte timeout".to_string()), }; } - request.send().await.map_err(|err| err.to_string()) + request + .send() + .await + .map_err(|error| format_relay_request_error(&error)) +} + +fn sanitize_relay_error_kind(raw_kind: &str) -> String { + // The relay is outside this process' trust boundary. Do not merely strip + // punctuation: a value such as `https://user:secret@example.invalid` + // would still carry the secret after character filtering. Relay errors + // are a small protocol enum, so only accept the categories emitted by the + // embedded relay and collapse everything else to a stable value. + let normalized = raw_kind.trim().to_ascii_lowercase(); + match normalized.as_str() { + "overloaded" | "forbidden" | "connect" | "relay" | "timeout" => normalized, + _ => "unknown".to_string(), + } +} + +fn format_relay_request_error(error: &reqwest::Error) -> String { + let kind = if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else if error.is_redirect() { + "redirect" + } else if error.is_body() { + "body" + } else if error.is_decode() { + "decode" + } else { + "unknown" + }; + format!("tunnel relay request failed [kind={kind}]") } pub(crate) fn build_request_body( plan: &ExecutionPlan, ) -> Result, ExecutionRuntimeTransportError> { - let mut body_bytes = if let Some(json_body) = plan.body.json_body.clone() { - serde_json::to_vec(&json_body).map_err(ExecutionRuntimeTransportError::BodyEncode)? + if plan.body.json_body.is_some() && plan.body.body_bytes_b64.is_some() { + return Err(ExecutionRuntimeTransportError::RequestBodyAmbiguous); + } + let body_limit = crate::headers::max_internal_buffered_body_bytes(); + let mut body_bytes = if let Some(json_body) = plan.body.json_body.as_ref() { + serialize_json_body_with_limit(json_body, body_limit)? } else if let Some(body_b64) = plan.body.body_bytes_b64.as_deref() { - base64::engine::general_purpose::STANDARD - .decode(body_b64) - .map_err(ExecutionRuntimeTransportError::BodyDecode)? + decode_base64_body_with_limit(body_b64, body_limit)? } else { Vec::new() }; @@ -2586,44 +3398,169 @@ fn zstd_bytes(body_bytes: &[u8]) -> Result, ExecutionRuntimeTransportErr fn build_relay_client( timeouts: Option<&aether_contracts::ExecutionTimeouts>, +) -> Result { + build_relay_client_with_pinned_target(timeouts, None) +} + +fn build_relay_client_with_pinned_target( + timeouts: Option<&aether_contracts::ExecutionTimeouts>, + pinned_target: Option<(&str, &[SocketAddr])>, ) -> Result { let builder = apply_http_client_config( - reqwest::Client::builder(), + reqwest::Client::builder() + .no_proxy() + .redirect(Policy::none()), &HttpClientConfig { connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms), use_rustls_tls: false, ..HttpClientConfig::default() }, ); + let builder = if let Some((host, addresses)) = pinned_target { + builder.resolve_to_addrs(host, addresses) + } else { + builder + }; builder .build() .map_err(ExecutionRuntimeTransportError::ClientBuild) } +async fn resolve_relay_target_addresses( + relay_url: &str, +) -> Result<(String, Vec), ExecutionRuntimeTransportError> { + let url = reqwest::Url::parse(relay_url).map_err(|_| { + ExecutionRuntimeTransportError::RelayError("invalid tunnel relay URL".to_string()) + })?; + validate_relay_target_url(&url)?; + let host = url.host_str().ok_or_else(|| { + ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no host".to_string()) + })?; + let port = url.port_or_known_default().ok_or_else(|| { + ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no port".to_string()) + })?; + // Relay destinations remain strict even when their hostname happens to be + // an official provider origin. The RFC-2544 compatibility exception is + // only for direct provider execution; allowing it here would weaken the + // relay SSRF guard. + let addresses = resolve_execution_target_addresses_with_policy(host, port, false) + .await + .map_err(|error| match error.kind() { + std::io::ErrorKind::PermissionDenied => ExecutionRuntimeTransportError::RelayError( + "tunnel relay DNS resolution returned a private or reserved address".to_string(), + ), + std::io::ErrorKind::NotFound => ExecutionRuntimeTransportError::RelayError( + "tunnel relay DNS resolution returned no addresses".to_string(), + ), + _ => ExecutionRuntimeTransportError::RelayError( + "tunnel relay DNS resolution failed".to_string(), + ), + })?; + Ok((host.to_string(), addresses)) +} + +fn validate_relay_target_url(url: &url::Url) -> Result<(), ExecutionRuntimeTransportError> { + tunnel::validate_tunnel_relay_transport_url(url) + .map_err(ExecutionRuntimeTransportError::RelayError)?; + if url.query().is_some() || url.fragment().is_some() { + return Err(ExecutionRuntimeTransportError::RelayError( + "tunnel relay URL must not include a query or fragment".to_string(), + )); + } + + let host = url.host_str().ok_or_else(|| { + ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no host".to_string()) + })?; + let explicit_loopback = dns_host_explicitly_allows_loopback(host); + // Relay credentials and envelopes must never be sent to a private target. + // Local relays are an explicit exception, and are intentionally plain HTTP + // so an operator cannot mistake a loopback TLS endpoint for a trusted peer. + if explicit_loopback && url.scheme() != "http" { + return Err(ExecutionRuntimeTransportError::RelayError( + "loopback tunnel relay URL must use HTTP".to_string(), + )); + } + if let Some(ip) = match url.host() { + Some(url::Host::Ipv4(address)) => Some(IpAddr::V4(address)), + Some(url::Host::Ipv6(address)) => Some(IpAddr::V6(address)), + _ => None, + } { + if is_private_or_reserved_ip(ip) && !(url.scheme() == "http" && ip.is_loopback()) { + return Err(ExecutionRuntimeTransportError::RelayError( + "tunnel relay URL must not target a private or reserved address".to_string(), + )); + } + } + Ok(()) +} + +struct RelayEnvelope { + body: Vec, + metadata_len: usize, +} + +impl RelayEnvelope { + fn metadata_envelope(&self) -> &[u8] { + &self.body[..self.metadata_len] + } + + fn request_body(&self) -> &[u8] { + &self.body[self.metadata_len..] + } +} + fn build_relay_envelope( meta: RelayRequestMeta, body_bytes: &[u8], -) -> Result, ExecutionRuntimeTransportError> { +) -> Result { let meta_bytes = - serde_json::to_vec(&meta).map_err(ExecutionRuntimeTransportError::BodyEncode)?; - let mut envelope = Vec::with_capacity(4 + meta_bytes.len() + body_bytes.len()); - envelope.extend_from_slice(&(meta_bytes.len() as u32).to_be_bytes()); + serialize_serializable_with_limit(&meta, MAX_TUNNEL_RELAY_META_LEN).map_err(|error| { + match error { + ExecutionRuntimeTransportError::BodyTooLarge { .. } => { + ExecutionRuntimeTransportError::RelayError( + "tunnel relay metadata exceeds configured limit".to_string(), + ) + } + other => other, + } + })?; + let metadata_len_u32 = u32::try_from(meta_bytes.len()).map_err(|_| { + ExecutionRuntimeTransportError::RelayError("tunnel relay metadata too large".to_string()) + })?; + let envelope_capacity = 4usize + .checked_add(meta_bytes.len()) + .and_then(|value| value.checked_add(body_bytes.len())) + .ok_or_else(|| { + ExecutionRuntimeTransportError::RelayError( + "tunnel relay envelope too large".to_string(), + ) + })?; + let mut envelope = Vec::with_capacity(envelope_capacity); + envelope.extend_from_slice(&metadata_len_u32.to_be_bytes()); envelope.extend_from_slice(&meta_bytes); + let metadata_len = envelope.len(); envelope.extend_from_slice(body_bytes); - Ok(envelope) + Ok(RelayEnvelope { + body: envelope, + metadata_len, + }) } -fn build_relay_url(proxy: Option<&ProxySnapshot>, node_id: &str) -> String { +fn build_relay_url( + proxy: Option<&ProxySnapshot>, + node_id: &str, +) -> Result { let base_url = proxy .and_then(resolve_tunnel_base_url_from_proxy) .or_else(|| std::env::var("AETHER_TUNNEL_BASE_URL").ok()) .unwrap_or_else(configured_gateway_frontdoor_base_url); - format!( - "{}{}/{}", - base_url.trim_end_matches('/'), - TUNNEL_RELAY_PATH_PREFIX, - node_id - ) + let relay_url = tunnel::build_tunnel_owner_relay_url(&base_url, node_id) + .map_err(ExecutionRuntimeTransportError::RelayError)?; + let parsed = reqwest::Url::parse(&relay_url).map_err(|_| { + ExecutionRuntimeTransportError::RelayError("invalid tunnel relay URL".to_string()) + })?; + validate_relay_target_url(&parsed)?; + Ok(relay_url) } fn resolve_tunnel_base_url_from_proxy(proxy: &ProxySnapshot) -> Option { @@ -2635,6 +3572,16 @@ fn resolve_tunnel_base_url_from_proxy(proxy: &ProxySnapshot) -> Option { None } +fn resolve_tunnel_owner_instance_id(proxy: Option<&ProxySnapshot>) -> Option<&str> { + proxy? + .extra + .as_ref()? + .get("tunnel_owner_instance_id")? + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) +} + fn resolve_relay_timeout_seconds(plan: &ExecutionPlan) -> u64 { resolve_tunnel_timeout_metadata(plan).legacy_timeout_secs } @@ -2860,11 +3807,11 @@ fn build_client( request_url, key_id, timeouts, - resolved_proxy_url, + resolved_proxy_url.clone(), transport_profile, transport_controls, ); - cached_direct_reqwest_client(cache_key) + cached_direct_reqwest_client(cache_key, resolved_proxy_url) } fn direct_reqwest_effective_transport_controls( @@ -2900,7 +3847,7 @@ pub(crate) fn prewarm_direct_reqwest_client_cache_for_plan(plan: &ExecutionPlan) Ok(false) => {} Err(err) => { tracing::debug!( - error = ?err, + error = %sanitize_error_detail(&err.to_string()), request_id = %plan.request_id, candidate_id = ?plan.candidate_id, provider_id = %plan.provider_id, @@ -2938,17 +3885,19 @@ fn try_prewarm_direct_reqwest_client_cache_for_plan( &plan.url, &plan.key_id, plan.timeouts.as_ref(), - resolved_proxy_url, + resolved_proxy_url.clone(), plan.transport_profile.as_ref(), transport_controls, ); - prewarm_direct_reqwest_client_cache(cache_key)?; + prewarm_direct_reqwest_client_cache(cache_key, resolved_proxy_url)?; Ok(true) } fn prewarm_direct_reqwest_client_cache( cache_key: DirectReqwestClientCacheKey, + proxy_url: Option, ) -> Result<(), ExecutionRuntimeTransportError> { + validate_direct_reqwest_proxy_material(&cache_key, proxy_url.as_deref())?; let mut warm_after_unlock = None; let cache_lock_started_at = Instant::now(); if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() { @@ -2957,14 +3906,21 @@ fn prewarm_direct_reqwest_client_cache( cache_lock_started_at.elapsed().as_millis() as u64, ); if let Some(entry) = cache.get_mut(&cache_key) { + entry.touch(); if entry.should_warm() { entry.warming = true; - warm_after_unlock = Some((cache_key.clone(), entry.len(), entry.target_len)); + warm_after_unlock = Some(( + cache_key.clone(), + proxy_url.clone(), + entry.len(), + entry.target_len, + )); } drop(cache); - if let Some((cache_key, existing_len, target_len)) = warm_after_unlock { + if let Some((cache_key, proxy_url, existing_len, target_len)) = warm_after_unlock { let spawned = spawn_direct_reqwest_client_cache_warm( cache_key.clone(), + proxy_url, existing_len, target_len, ); @@ -2979,7 +3935,10 @@ fn prewarm_direct_reqwest_client_cache( let initial_len = direct_reqwest_prewarm_client_shard_count(target_len); let mut clients = Vec::with_capacity(initial_len); for _ in 0..initial_len { - clients.push(build_direct_reqwest_client_from_cache_key(&cache_key)?); + clients.push(build_direct_reqwest_client_from_cache_key( + &cache_key, + proxy_url.as_deref(), + )?); DIRECT_REQWEST_CLIENT_CACHE_METRICS .builds .fetch_add(1, Ordering::Relaxed); @@ -2987,14 +3946,19 @@ fn prewarm_direct_reqwest_client_cache( let entry = DirectReqwestClientCacheEntry::new(clients, target_len, target_len > initial_len); let warm_key = (target_len > initial_len).then(|| cache_key.clone()); + evict_direct_reqwest_client_cache_for_insert(&mut cache, &cache_key); cache.insert(cache_key, entry); if let Some(warm_key) = warm_key { - warm_after_unlock = Some((warm_key, initial_len, target_len)); + warm_after_unlock = Some((warm_key, proxy_url, initial_len, target_len)); } drop(cache); - if let Some((cache_key, existing_len, target_len)) = warm_after_unlock { - let spawned = - spawn_direct_reqwest_client_cache_warm(cache_key.clone(), existing_len, target_len); + if let Some((cache_key, proxy_url, existing_len, target_len)) = warm_after_unlock { + let spawned = spawn_direct_reqwest_client_cache_warm( + cache_key.clone(), + proxy_url, + existing_len, + target_len, + ); if !spawned { mark_direct_reqwest_client_cache_not_warming(&cache_key); } @@ -3010,7 +3974,9 @@ fn prewarm_direct_reqwest_client_cache( fn cached_direct_reqwest_client( cache_key: DirectReqwestClientCacheKey, + proxy_url: Option, ) -> Result { + validate_direct_reqwest_proxy_material(&cache_key, proxy_url.as_deref())?; let mut warm_after_unlock = None; let cache_lock_started_at = Instant::now(); if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() { @@ -3019,6 +3985,7 @@ fn cached_direct_reqwest_client( cache_lock_started_at.elapsed().as_millis() as u64, ); if let Some(entry) = cache.get_mut(&cache_key) { + entry.touch(); DIRECT_REQWEST_CLIENT_CACHE_METRICS .hits .fetch_add(1, Ordering::Relaxed); @@ -3026,12 +3993,18 @@ fn cached_direct_reqwest_client( let client = entry.select(); if entry.should_warm() { entry.warming = true; - warm_after_unlock = Some((cache_key.clone(), entry.len(), entry.target_len)); + warm_after_unlock = Some(( + cache_key.clone(), + proxy_url.clone(), + entry.len(), + entry.target_len, + )); } drop(cache); - if let Some((cache_key, existing_len, target_len)) = warm_after_unlock { + if let Some((cache_key, proxy_url, existing_len, target_len)) = warm_after_unlock { let spawned = spawn_direct_reqwest_client_cache_warm( cache_key.clone(), + proxy_url, existing_len, target_len, ); @@ -3048,7 +4021,10 @@ fn cached_direct_reqwest_client( let initial_len = direct_reqwest_initial_client_shard_count(target_len); let mut clients = Vec::with_capacity(initial_len); for _ in 0..initial_len { - clients.push(build_direct_reqwest_client_from_cache_key(&cache_key)?); + clients.push(build_direct_reqwest_client_from_cache_key( + &cache_key, + proxy_url.as_deref(), + )?); DIRECT_REQWEST_CLIENT_CACHE_METRICS .builds .fetch_add(1, Ordering::Relaxed); @@ -3058,14 +4034,19 @@ fn cached_direct_reqwest_client( record_direct_reqwest_client_protocol_selection(&cache_key); let client = entry.select(); let warm_key = (target_len > initial_len).then(|| cache_key.clone()); + evict_direct_reqwest_client_cache_for_insert(&mut cache, &cache_key); cache.insert(cache_key, entry); if let Some(warm_key) = warm_key { - warm_after_unlock = Some((warm_key, initial_len, target_len)); + warm_after_unlock = Some((warm_key, proxy_url, initial_len, target_len)); } drop(cache); - if let Some((cache_key, existing_len, target_len)) = warm_after_unlock { - let spawned = - spawn_direct_reqwest_client_cache_warm(cache_key.clone(), existing_len, target_len); + if let Some((cache_key, proxy_url, existing_len, target_len)) = warm_after_unlock { + let spawned = spawn_direct_reqwest_client_cache_warm( + cache_key.clone(), + proxy_url, + existing_len, + target_len, + ); if !spawned { mark_direct_reqwest_client_cache_not_warming(&cache_key); } @@ -3081,7 +4062,7 @@ fn cached_direct_reqwest_client( .misses .fetch_add(1, Ordering::Relaxed); record_direct_reqwest_client_protocol_selection(&cache_key); - let client = build_direct_reqwest_client_from_cache_key(&cache_key)?; + let client = build_direct_reqwest_client_from_cache_key(&cache_key, proxy_url.as_deref())?; DIRECT_REQWEST_CLIENT_CACHE_METRICS .builds .fetch_add(1, Ordering::Relaxed); @@ -3090,6 +4071,7 @@ fn cached_direct_reqwest_client( fn spawn_direct_reqwest_client_cache_warm( cache_key: DirectReqwestClientCacheKey, + proxy_url: Option, existing_len: usize, target_len: usize, ) -> bool { @@ -3111,7 +4093,7 @@ fn spawn_direct_reqwest_client_cache_warm( let enqueue_started_at = Instant::now(); handle.spawn_blocking(move || { for _ in existing_len..target_len { - match build_direct_reqwest_client_from_cache_key(&cache_key) { + match build_direct_reqwest_client_from_cache_key(&cache_key, proxy_url.as_deref()) { Ok(client) => { DIRECT_REQWEST_CLIENT_CACHE_METRICS .builds @@ -3134,7 +4116,7 @@ fn spawn_direct_reqwest_client_cache_warm( } Err(err) => { tracing::debug!( - error = ?err, + error = %sanitize_error_detail(&err.to_string()), "gateway direct reqwest client cache warm failed" ); mark_direct_reqwest_client_cache_not_warming(&cache_key); @@ -3174,6 +4156,41 @@ fn mark_direct_reqwest_client_cache_not_warming(cache_key: &DirectReqwestClientC } } +fn next_direct_reqwest_client_cache_clock() -> u64 { + DIRECT_REQWEST_CLIENT_CACHE_CLOCK + .fetch_add(1, Ordering::Relaxed) + .wrapping_add(1) +} + +fn direct_reqwest_client_cache_max_entries() -> usize { + env_positive_usize(DIRECT_REQWEST_CACHE_MAX_ENTRIES_ENV) + .unwrap_or(DEFAULT_DIRECT_REQWEST_CACHE_MAX_ENTRIES) + .clamp(1, MAX_DIRECT_REQWEST_CACHE_MAX_ENTRIES) +} + +fn evict_direct_reqwest_client_cache_for_insert( + cache: &mut HashMap, + incoming: &DirectReqwestClientCacheKey, +) { + if cache.contains_key(incoming) { + return; + } + let max_entries = direct_reqwest_client_cache_max_entries(); + while cache.len() >= max_entries { + let Some(oldest_key) = cache + .iter() + .min_by_key(|(_, entry)| entry.last_used) + .map(|(key, _)| key.clone()) + else { + break; + }; + cache.remove(&oldest_key); + DIRECT_REQWEST_CLIENT_CACHE_METRICS + .evictions + .fetch_add(1, Ordering::Relaxed); + } +} + fn direct_reqwest_client_cache_key( request_url: &str, key_id: &str, @@ -3188,14 +4205,28 @@ fn direct_reqwest_client_cache_key( .flatten(), pool_partition: direct_reqwest_pool_partition(transport_profile, key_id), connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms), - proxy_url, + proxy_digest: proxy_url.map(|proxy_url| direct_reqwest_proxy_digest(&proxy_url)), follow_redirects: transport_controls.follow_redirects == Some(true), http1_only: transport_controls.http1_only, - accept_invalid_certs: transport_controls.accept_invalid_certs, transport_profile: transport_profile.map(direct_reqwest_transport_profile_cache_key), } } +fn direct_reqwest_proxy_digest(proxy_url: &str) -> String { + format!("{:x}", sha2::Sha256::digest(proxy_url.as_bytes())) +} + +fn validate_direct_reqwest_proxy_material( + cache_key: &DirectReqwestClientCacheKey, + proxy_url: Option<&str>, +) -> Result<(), ExecutionRuntimeTransportError> { + if cache_key.proxy_digest.as_deref() == proxy_url.map(direct_reqwest_proxy_digest).as_deref() { + Ok(()) + } else { + Err(ExecutionRuntimeTransportError::ProxyUnsupported) + } +} + fn direct_reqwest_pool_partition( transport_profile: Option<&ResolvedTransportProfile>, key_id: &str, @@ -3228,7 +4259,11 @@ fn direct_reqwest_upstream_origin(request_url: &str) -> Option { } let host = url.host_str()?; let port = url.port_or_known_default()?; - Some(format!("{scheme}://{host}:{port}")) + let authority_host = match url.host() { + Some(url::Host::Ipv6(_)) if !host.starts_with('[') => format!("[{host}]"), + _ => host.to_string(), + }; + Some(format!("{scheme}://{authority_host}:{port}")) } fn direct_reqwest_transport_profile_cache_key( @@ -3250,11 +4285,14 @@ fn stable_json_cache_key(value: Option<&Value>) -> Option { fn build_direct_reqwest_client_cache_entry_from_cache_key( cache_key: &DirectReqwestClientCacheKey, + proxy_url: Option<&str>, ) -> Result { let shard_count = direct_reqwest_client_shard_count(cache_key); let mut clients = Vec::with_capacity(shard_count); for _ in 0..shard_count { - clients.push(build_direct_reqwest_client_from_cache_key(cache_key)?); + clients.push(build_direct_reqwest_client_from_cache_key( + cache_key, proxy_url, + )?); } Ok(DirectReqwestClientCacheEntry::new( clients, @@ -3368,11 +4406,21 @@ fn env_positive_usize(name: &str) -> Option { fn build_direct_reqwest_client_from_cache_key( cache_key: &DirectReqwestClientCacheKey, + proxy_url: Option<&str>, ) -> Result { - let mut builder = reqwest::Client::builder(); - if !cache_key.follow_redirects { - builder = builder.redirect(Policy::none()); + validate_direct_reqwest_proxy_material(cache_key, proxy_url)?; + if let Some(proxy_url) = proxy_url { + validate_execution_proxy_url(proxy_url)?; } + let mut builder = reqwest::Client::builder().no_proxy(); + if proxy_url.is_none() { + builder = builder.dns_resolver(Arc::new(ExecutionSafeDnsResolver)); + } + builder = builder.redirect(if cache_key.follow_redirects { + same_origin_reqwest_redirect_policy() + } else { + Policy::none() + }); if cache_key.http1_only || cache_key .transport_profile @@ -3400,10 +4448,7 @@ fn build_direct_reqwest_client_from_cache_key( cache_key.transport_profile.as_ref(), cache_key.http1_only, ); - if cache_key.accept_invalid_certs { - builder = builder.danger_accept_invalid_certs(true); - } - if let Some(proxy_url) = cache_key.proxy_url.as_deref() { + if let Some(proxy_url) = proxy_url { let proxy = reqwest::Proxy::all(proxy_url).map_err(ExecutionRuntimeTransportError::InvalidProxy)?; builder = builder.proxy(proxy); @@ -3591,6 +4636,14 @@ pub(crate) fn direct_reqwest_client_cache_metric_samples() -> Vec .warm_skipped_total .load(Ordering::Relaxed), ), + MetricSample::new( + "direct_reqwest_client_cache_evictions_total", + "Number of least-recently-used direct reqwest client cache entries evicted at the configured capacity.", + MetricKind::Counter, + DIRECT_REQWEST_CLIENT_CACHE_METRICS + .evictions + .load(Ordering::Relaxed), + ), MetricSample::new( "direct_reqwest_client_http1_select_total", "Number of direct reqwest client selections using forced HTTP/1.", @@ -3757,16 +4810,19 @@ pub(crate) fn build_browser_wreq_client( apply_total_timeout: bool, ) -> Result { let emulation = browser_wreq_emulation_from_profile(transport_profile)?; - let mut builder = wreq::Client::builder().emulation(emulation); - if transport_controls.follow_redirects == Some(true) { - builder = builder.redirect(wreq::redirect::Policy::limited(10)); + let proxy_url = resolve_proxy_url(proxy)?; + let mut builder = wreq::Client::builder().no_proxy().emulation(emulation); + if proxy_url.is_none() { + builder = builder.dns_resolver(ExecutionSafeDnsResolver); } + builder = builder.redirect(if transport_controls.follow_redirects == Some(true) { + same_origin_wreq_redirect_policy() + } else { + wreq::redirect::Policy::none() + }); if transport_controls.http1_only || transport_profile_http1_only(Some(transport_profile)) { builder = builder.http1_only(); } - if transport_controls.accept_invalid_certs { - builder = builder.cert_verification(false).verify_hostname(false); - } if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) { builder = builder.connect_timeout(Duration::from_millis(connect_ms)); } @@ -3778,7 +4834,7 @@ pub(crate) fn build_browser_wreq_client( if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) { builder = builder.read_timeout(Duration::from_millis(read_ms)); } - if let Some(proxy_url) = resolve_proxy_url(proxy)? { + if let Some(proxy_url) = proxy_url { let proxy = wreq::Proxy::all(proxy_url.as_str()) .map_err(ExecutionRuntimeTransportError::BrowserClientBuild)?; builder = builder.proxy(proxy); @@ -3788,6 +4844,85 @@ pub(crate) fn build_browser_wreq_client( .map_err(ExecutionRuntimeTransportError::BrowserClientBuild) } +fn same_origin_reqwest_redirect_policy() -> Policy { + Policy::custom(|attempt| { + let same_origin = attempt + .previous() + .last() + .is_some_and(|previous| reqwest_urls_have_same_origin(previous, attempt.url())); + match safe_redirect_decision(attempt.previous().len(), same_origin) { + SafeRedirectDecision::Follow => attempt.follow(), + SafeRedirectDecision::Stop => attempt.stop(), + SafeRedirectDecision::TooMany => attempt.error("too many redirects"), + } + }) +} + +fn same_origin_wreq_redirect_policy() -> wreq::redirect::Policy { + wreq::redirect::Policy::custom(|attempt| { + let same_origin = attempt + .previous + .last() + .is_some_and(|previous| http_uris_have_same_origin(previous, &attempt.uri)); + match safe_redirect_decision(attempt.previous.len(), same_origin) { + SafeRedirectDecision::Follow => attempt.follow(), + SafeRedirectDecision::Stop => attempt.stop(), + SafeRedirectDecision::TooMany => attempt.error("too many redirects"), + } + }) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SafeRedirectDecision { + Follow, + Stop, + TooMany, +} + +fn safe_redirect_decision(previous_len: usize, same_origin: bool) -> SafeRedirectDecision { + if previous_len > MAX_SAFE_REDIRECTS { + SafeRedirectDecision::TooMany + } else if same_origin { + SafeRedirectDecision::Follow + } else { + SafeRedirectDecision::Stop + } +} + +fn reqwest_urls_have_same_origin(previous: &reqwest::Url, next: &reqwest::Url) -> bool { + previous.scheme().eq_ignore_ascii_case(next.scheme()) + && previous + .host_str() + .zip(next.host_str()) + .is_some_and(|(previous, next)| previous.eq_ignore_ascii_case(next)) + && previous.port_or_known_default() == next.port_or_known_default() +} + +fn http_uris_have_same_origin(previous: &http::Uri, next: &http::Uri) -> bool { + previous + .scheme_str() + .zip(next.scheme_str()) + .is_some_and(|(previous, next)| previous.eq_ignore_ascii_case(next)) + && previous + .host() + .zip(next.host()) + .is_some_and(|(previous, next)| previous.eq_ignore_ascii_case(next)) + && http_uri_effective_port(previous) == http_uri_effective_port(next) +} + +fn http_uri_effective_port(uri: &http::Uri) -> Option { + uri.port_u16().or_else(|| { + let scheme = uri.scheme_str()?; + if scheme.eq_ignore_ascii_case("http") { + Some(80) + } else if scheme.eq_ignore_ascii_case("https") { + Some(443) + } else { + None + } + }) +} + fn browser_wreq_emulation_from_profile( profile: &ResolvedTransportProfile, ) -> Result { @@ -4005,14 +5140,49 @@ fn resolve_proxy_url( .map(|url| url.trim()) .filter(|url| !url.is_empty()) { - return Ok(Some(proxy_url.to_string())); + return normalize_execution_proxy_url(proxy_url).map(Some); } - if proxy.node_id.is_some() || proxy.mode.as_deref() == Some("tunnel") { + Err(ExecutionRuntimeTransportError::ProxyUnsupported) +} + +fn validate_execution_proxy_url(raw_url: &str) -> Result<(), ExecutionRuntimeTransportError> { + parse_execution_proxy_url(raw_url).map(|_| ()) +} + +/// Normalize a configured proxy URL before handing it to reqwest/wreq. +/// +/// `socks5://` has a particularly dangerous ambiguity in a gateway: reqwest +/// and wreq interpret it as *local* target-name resolution, while +/// `socks5h://` delegates target resolution to the proxy. Local resolution +/// would bypass the execution DNS guard (and could turn a rebinding hostname +/// into a private address). Keep accepting the established `socks5` config +/// syntax for compatibility, but make its runtime semantics the safe remote +/// DNS variant. HTTP/HTTPS and already-remote `socks5h` URLs are unchanged. +pub(crate) fn normalize_execution_proxy_url( + raw_url: &str, +) -> Result { + let mut parsed = parse_execution_proxy_url(raw_url)?; + if parsed.scheme().eq_ignore_ascii_case("socks5") { + parsed + .set_scheme("socks5h") + .map_err(|_| ExecutionRuntimeTransportError::ProxyUnsupported)?; + } + Ok(parsed.to_string()) +} + +fn parse_execution_proxy_url(raw_url: &str) -> Result { + let parsed = + url::Url::parse(raw_url).map_err(|_| ExecutionRuntimeTransportError::ProxyUnsupported)?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + || !matches!(parsed.path(), "" | "/") + || parsed.query().is_some() + || parsed.fragment().is_some() + { return Err(ExecutionRuntimeTransportError::ProxyUnsupported); } - - Ok(None) + Ok(parsed) } pub(crate) fn build_request_headers( @@ -4021,6 +5191,12 @@ pub(crate) fn build_request_headers( allow_passthrough_content_encoding: bool, ) -> Result { let mut out = HeaderMap::new(); + let connection_declared = aether_http::connection_declared_header_names( + headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); let normalized_content_encoding = normalize_content_encoding(content_encoding); if let Some(encoding) = normalized_content_encoding.as_deref() { if !matches!(encoding, "gzip" | "zstd") && !allow_passthrough_content_encoding { @@ -4033,10 +5209,11 @@ pub(crate) fn build_request_headers( let normalized_key = key.trim().to_ascii_lowercase(); if crate::headers::should_skip_request_header(&normalized_key) || is_hop_by_hop_header(&normalized_key) + || connection_declared.contains(&normalized_key) || normalized_key == "content-encoding" || normalized_key == EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER || normalized_key == EXECUTION_REQUEST_HTTP1_ONLY_HEADER - || normalized_key == EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER + || normalized_key == LEGACY_EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER || normalized_key == EXECUTION_RESPONSE_BODY_MODE_HEADER || normalized_key == EXECUTION_RESPONSE_BODY_LIMIT_HEADER { @@ -4072,12 +5249,6 @@ fn resolve_execution_transport_controls( http1_only: execution_transport_header_value(headers, EXECUTION_REQUEST_HTTP1_ONLY_HEADER) .and_then(|value| parse_execution_transport_bool(value)) .unwrap_or(false), - accept_invalid_certs: execution_transport_header_value( - headers, - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, - ) - .and_then(|value| parse_execution_transport_bool(value)) - .unwrap_or(false), } } @@ -4149,12 +5320,34 @@ fn is_hop_by_hop_header(name: &str) -> bool { } pub(crate) fn collect_response_headers(headers: &HeaderMap) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); header_map_to_string_map(headers) + .into_iter() + .filter(|(name, _)| { + !crate::headers::should_skip_response_header(name) + && !connection_declared.contains(&name.to_ascii_lowercase()) + }) + .collect() } fn collect_tunnel_response_headers(headers: &[(String, String)]) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); headers .iter() + .filter(|(name, _)| { + !crate::headers::should_skip_response_header(name) + && !connection_declared.contains(&name.to_ascii_lowercase()) + }) .map(|(name, value)| (name.to_ascii_lowercase(), value.clone())) .collect() } @@ -4176,6 +5369,47 @@ fn execution_log_url_host(url: &str) -> String { .unwrap_or_else(|| "-".to_string()) } +fn validate_execution_upstream_url( + raw_url: &str, +) -> Result { + let url = url::Url::parse(raw_url).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest("invalid upstream URL".to_string()) + })?; + if url.host().is_none() || !matches!(url.scheme(), "http" | "https") { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "upstream URL must use HTTP or HTTPS and include a host".to_string(), + )); + } + if !url.username().is_empty() || url.password().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "upstream URL must not include credentials".to_string(), + )); + } + if url.fragment().is_some() { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "upstream URL must not include a fragment".to_string(), + )); + } + if !is_https_or_loopback_http_url(&url) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "remote upstream URL must use HTTPS".to_string(), + )); + } + let literal_ip = match url.host() { + Some(url::Host::Ipv4(address)) => Some(IpAddr::V4(address)), + Some(url::Host::Ipv6(address)) => Some(IpAddr::V6(address)), + _ => None, + }; + if literal_ip.is_some_and(|ip| { + is_private_or_reserved_ip(ip) && !(url.scheme() == "http" && ip.is_loopback()) + }) { + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "upstream URL must not target a private or reserved address".to_string(), + )); + } + Ok(url) +} + pub(crate) fn decode_response_body_bytes<'a>( headers: &BTreeMap, body_bytes: &'a [u8], @@ -4309,13 +5543,21 @@ mod tests { use std::io::{Read, Write}; use std::sync::{Arc, Mutex, MutexGuard, OnceLock}; + use aether_contracts::tunnel::{ + 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_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; use aether_contracts::{ ExecutionPlan, ExecutionResponseBodyMode, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER, - TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, + PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY, TRANSPORT_BACKEND_BROWSER_WREQ, + TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY, }; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::proxy_nodes::{ InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode, }; @@ -4332,16 +5574,22 @@ mod tests { use super::{ append_upstream_response_body_chunk_with_limit, build_browser_wreq_client, build_client, - build_direct_tunnel_request_meta, build_execution_response_body, build_request_headers, + build_direct_tunnel_request_meta, build_execution_response_body, build_relay_client, + build_relay_url, build_request_headers, collect_response_headers, + collect_tunnel_response_headers, decode_base64_body_with_limit, decode_response_body_bytes_with_limit, effective_response_body_limit_bytes, execute_sync_plan, execution_plan_response_body_limit_bytes, execution_response_body_mode, + execution_result_envelope_limit_bytes, http_uris_have_same_origin, + json_value_fits_serialized_limit, maximum_base64_len_for_decoded_limit, record_manual_proxy_request_failure, record_manual_proxy_request_outcome, record_manual_proxy_request_success, record_manual_proxy_stream_error, - resolve_execution_transport_controls, resolve_non_stream_total_timeout, - resolve_stream_first_byte_timeout, response_body_is_json, - with_upstream_response_body_limit, DirectSyncExecutionRuntime, - ExecutionRuntimeTransportError, ExecutionTransportControls, UpstreamResponseBodyPhase, - DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES, EXECUTION_RESPONSE_BODY_LIMIT_HEADER, + reqwest_urls_have_same_origin, resolve_execution_transport_controls, + resolve_non_stream_total_timeout, resolve_proxy_url, resolve_stream_first_byte_timeout, + response_body_is_json, safe_redirect_decision, validate_execution_upstream_url, + validate_relay_target_url, with_upstream_response_body_limit, DirectSyncExecutionRuntime, + ExecutionRuntimeTransportError, ExecutionTransportControls, RelayRequestMeta, + SafeRedirectDecision, UpstreamResponseBodyPhase, DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES, + EXECUTION_RESPONSE_BODY_LIMIT_HEADER, MAX_SAFE_REDIRECTS, MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES, MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES, }; use crate::constants::{ @@ -4355,11 +5603,355 @@ mod tests { use crate::AppState; const LOCAL_HTTP_SUCCESS_TIMEOUT_MS: u64 = 15_000; + const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes"; + + #[test] + fn execution_upstream_url_requires_https_or_literal_loopback_http() { + for allowed in [ + "https://api.example.test/v1/responses?api-version=1", + "http://localhost:8080/v1/responses", + "http://127.42.0.1:8080/v1/responses", + "http://[::1]:8080/v1/responses", + ] { + assert!( + validate_execution_upstream_url(allowed).is_ok(), + "URL should be accepted: {allowed}" + ); + } + + for rejected in [ + "http://api.example.test/v1/responses", + "http://10.0.0.1/v1/responses", + "http://0.0.0.0:8080/v1/responses", + "http://[::ffff:127.0.0.1]:8080/v1/responses", + "https://127.0.0.1:8443/v1/responses", + "https://10.0.0.1:8443/v1/responses", + "https://token@example.test/v1/responses", + "https://example.test/v1/responses#secret", + "ftp://localhost/resource", + ] { + assert!( + validate_execution_upstream_url(rejected).is_err(), + "URL should be rejected: {rejected}" + ); + } + } + + #[test] + fn invalid_execution_upstream_url_error_does_not_echo_credentials() { + let error = validate_execution_upstream_url( + "https://sensitive-user:sensitive-password@example.test/v1/responses", + ) + .expect_err("URL userinfo should be rejected") + .to_string(); + + assert!(!error.contains("sensitive-user")); + assert!(!error.contains("sensitive-password")); + } + + #[test] + fn execution_dns_answers_reject_private_addresses_and_allow_explicit_loopback() { + let public = "93.184.216.34:443".parse().unwrap(); + let private = "10.0.0.8:443".parse().unwrap(); + let loopback_v4 = "127.0.0.1:8080".parse().unwrap(); + let loopback_v6 = "[::1]:8080".parse().unwrap(); + + assert!(super::validate_execution_dns_answers("api.example.test", vec![public]).is_ok()); + assert!(super::validate_execution_dns_answers("api.example.test", vec![private]).is_err()); + assert!( + super::validate_execution_dns_answers("localhost", vec![loopback_v4, loopback_v6]) + .is_ok() + ); + assert!(super::validate_execution_dns_answers("localhost", vec![private]).is_err()); + assert!(super::validate_execution_dns_answers("api.example.test", Vec::new()).is_err()); + } + + #[test] + fn execution_dns_answers_allow_benchmarking_range_only_for_fixed_provider_hosts() { + let fake = "198.18.75.234:443".parse().unwrap(); + for host in [ + "api.openai.com", + "CHATGPT.COM.", + "us-central1-aiplatform.googleapis.com", + "me-central2-aiplatform.googleapis.com", + "q.us-east-1.amazonaws.com", + "q-fips.us-gov-west-1.amazonaws.com", + "codewhisperer.us-west-2.amazonaws.com", + "oidc.us-east-1.amazonaws.com", + "prod.us-east-1.auth.desktop.kiro.dev", + "q.us-iso-east-1.c2s.ic.gov", + "q.us-isob-east-1.sc2s.sgov.gov", + "q.us-isof-east-1.csp.hci.ic.gov", + ] { + assert!( + super::validate_execution_dns_answers(host, vec![fake]).is_ok(), + "fixed provider host should accept a benchmarking DNS answer: {host}" + ); + } + + for host in [ + "api.example.test", + "evil.chatgpt.com", + "api.openai.com.evil.test", + "q.us-east-1.evil.amazonaws.com", + "q.us-east-1.amazonaws.com.attacker.test", + "q.localhost.amazonaws.com", + "evil-1-aiplatform.googleapis.com", + "q.evil-1.amazonaws.com", + "q-fips.evil-1.amazonaws.com", + "codewhisperer.evil-1.amazonaws.com", + "prod.evil-1.auth.desktop.kiro.dev", + "oidc.evil-1.amazonaws.com", + "q.us-central1.amazonaws.com", + "us-east-1-aiplatform.googleapis.com", + "q.us-east-1.c2s.ic.gov", + "q.us-iso-east-1.sc2s.sgov.gov", + "q-fips.us-gov-west-1.evil.amazonaws.com", + "codewhisperer.us-west-2.evil.amazonaws.com", + "oidc.us-east-1.evil.amazonaws.com", + "prod.us-east-1.auth.desktop.kiro.dev.attacker.test", + "prod.us-east-1.evil.auth.desktop.kiro.dev", + "q.us-iso-east-1.evil.c2s.ic.gov", + "q.us-iso-east-1.c2s.ic.gov.attacker.test", + "q.us-iso-east-1.c2s.ic.gov.evil", + "198.18.75.234", + ] { + assert!( + super::validate_execution_dns_answers(host, vec![fake]).is_err(), + "untrusted or lookalike host must reject a benchmarking DNS answer: {host}" + ); + } + } + + #[test] + fn execution_dns_answers_allow_benchmarking_range_for_configured_exact_hosts() { + let fake = "198.18.75.234:443".parse().unwrap(); + super::refresh_execution_extra_trusted_dns_hosts(Some(&json!(["custom.example.com",]))); + + assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_ok()); + assert!( + super::validate_execution_dns_answers("api.custom.example.com", vec![fake]).is_err() + ); + + super::refresh_execution_extra_trusted_dns_hosts(None); + assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_err()); + } + + #[test] + fn execution_dns_answers_reject_mixed_private_results_and_strict_relay_policy() { + let fake = "198.18.75.234:443".parse().unwrap(); + let public = "93.184.216.34:443".parse().unwrap(); + let private = "10.0.0.8:443".parse().unwrap(); + + // A trusted host may have a synthetic answer alongside a genuine public + // answer, but any real private answer still fails closed. + assert!( + super::validate_execution_dns_answers("api.openai.com", vec![fake, public]).is_ok() + ); + assert!( + super::validate_execution_dns_answers("api.openai.com", vec![fake, private]).is_err() + ); + + // Tunnel relay resolution opts out of the compatibility exception. + assert!(super::validate_execution_dns_answers_with_policy( + "api.openai.com", + vec![fake], + false, + ) + .is_err()); + } + + #[test] + fn execution_proxy_url_policy_rejects_non_origin_components() { + for rejected in [ + "mailto:proxy@example.test", + "http://proxy.example.test/path", + "http://proxy.example.test?token=secret", + "http://proxy.example.test#fragment", + ] { + assert!( + super::validate_execution_proxy_url(rejected).is_err(), + "proxy URL should be rejected: {rejected}" + ); + } + assert!(super::validate_execution_proxy_url( + "http://alice:password@proxy.example.test:8080" + ) + .is_ok()); + } + + #[test] + fn execution_proxy_url_normalizes_local_socks_dns_to_remote_dns() { + assert_eq!( + super::normalize_execution_proxy_url("socks5://alice:password@proxy.example.test:1080") + .expect("socks5 URL should normalize"), + "socks5h://alice:password@proxy.example.test:1080" + ); + assert_eq!( + super::normalize_execution_proxy_url("socks5h://proxy.example.test:1080") + .expect("socks5h URL should remain valid"), + "socks5h://proxy.example.test:1080" + ); + assert_eq!( + super::normalize_execution_proxy_url("https://proxy.example.test:8443") + .expect("https URL should remain valid"), + "https://proxy.example.test:8443/" + ); + } + + #[test] + fn relay_error_kind_accepts_only_protocol_categories() { + assert_eq!(super::sanitize_relay_error_kind("TIMEOUT"), "timeout"); + assert_eq!(super::sanitize_relay_error_kind("upstream"), "unknown"); + assert_eq!( + super::sanitize_relay_error_kind("https://relay-user:secret@example.test"), + "unknown" + ); + } + + #[test] + fn enabled_proxy_without_a_usable_target_is_rejected() { + let proxy = ProxySnapshot { + enabled: Some(true), + mode: Some("unavailable".to_string()), + ..ProxySnapshot::default() + }; + + assert!(matches!( + resolve_proxy_url(Some(&proxy)), + Err(ExecutionRuntimeTransportError::ProxyUnsupported) + )); + assert_eq!( + resolve_proxy_url(Some(&ProxySnapshot { + enabled: Some(false), + ..ProxySnapshot::default() + })) + .expect("disabled proxy should be accepted"), + None + ); + } + + #[test] + fn redirect_origin_checks_scheme_host_and_effective_port() { + let reqwest_base = + reqwest::Url::parse("https://api.example.com/v1").expect("base URL should parse"); + for same_origin in [ + "https://api.example.com/v2", + "https://API.EXAMPLE.COM:443/v2", + ] { + let next = reqwest::Url::parse(same_origin).expect("same-origin URL should parse"); + assert!(reqwest_urls_have_same_origin(&reqwest_base, &next)); + } + for cross_origin in [ + "http://api.example.com/v2", + "https://other.example.com/v2", + "https://api.example.com:444/v2", + ] { + let next = reqwest::Url::parse(cross_origin).expect("cross-origin URL should parse"); + assert!(!reqwest_urls_have_same_origin(&reqwest_base, &next)); + } + + let wreq_base: http::Uri = "https://api.example.com/v1" + .parse() + .expect("base URI should parse"); + for same_origin in [ + "https://api.example.com/v2", + "https://API.EXAMPLE.COM:443/v2", + ] { + let next: http::Uri = same_origin.parse().expect("same-origin URI should parse"); + assert!(http_uris_have_same_origin(&wreq_base, &next)); + } + for cross_origin in [ + "http://api.example.com/v2", + "https://other.example.com/v2", + "https://api.example.com:444/v2", + ] { + let next: http::Uri = cross_origin.parse().expect("cross-origin URI should parse"); + assert!(!http_uris_have_same_origin(&wreq_base, &next)); + } + } + + #[test] + fn safe_redirect_decision_preserves_the_ten_hop_limit() { + assert_eq!( + safe_redirect_decision(MAX_SAFE_REDIRECTS, true), + SafeRedirectDecision::Follow + ); + assert_eq!( + safe_redirect_decision(MAX_SAFE_REDIRECTS + 1, true), + SafeRedirectDecision::TooMany + ); + assert_eq!(safe_redirect_decision(1, false), SafeRedirectDecision::Stop); + } + + #[test] + fn direct_and_tunnel_response_collectors_strip_upstream_security_headers() { + let mut direct = reqwest::header::HeaderMap::new(); + direct.insert( + reqwest::header::SET_COOKIE, + reqwest::header::HeaderValue::from_static("session=attacker"), + ); + direct.insert( + "x-aether-future-control", + reqwest::header::HeaderValue::from_static("attacker"), + ); + direct.insert( + reqwest::header::CONTENT_TYPE, + reqwest::header::HeaderValue::from_static("application/json"), + ); + direct.append( + reqwest::header::CONNECTION, + reqwest::header::HeaderValue::from_static("x-first-hop"), + ); + direct.append( + reqwest::header::CONNECTION, + reqwest::header::HeaderValue::from_static("x-second-hop"), + ); + direct.insert( + "x-first-hop", + reqwest::header::HeaderValue::from_static("first-secret"), + ); + direct.insert( + "x-second-hop", + reqwest::header::HeaderValue::from_static("second-secret"), + ); + + let direct = collect_response_headers(&direct); + assert!(!direct.contains_key("set-cookie")); + assert!(!direct.contains_key("x-aether-future-control")); + assert!(!direct.contains_key("x-first-hop")); + assert!(!direct.contains_key("x-second-hop")); + assert_eq!( + direct.get("content-type").map(String::as_str), + Some("application/json") + ); + + let tunnel = collect_tunnel_response_headers(&[ + ("Set-Cookie".to_string(), "session=attacker".to_string()), + ( + "X-Aether-Future-Control".to_string(), + "attacker".to_string(), + ), + ("Content-Type".to_string(), "application/json".to_string()), + ("Connection".to_string(), "x-first-hop".to_string()), + ("connection".to_string(), "x-second-hop".to_string()), + ("x-first-hop".to_string(), "first-secret".to_string()), + ("x-second-hop".to_string(), "second-secret".to_string()), + ]); + assert!(!tunnel.contains_key("set-cookie")); + assert!(!tunnel.contains_key("x-aether-future-control")); + assert!(!tunnel.contains_key("x-first-hop")); + assert!(!tunnel.contains_key("x-second-hop")); + assert_eq!( + tunnel.get("content-type").map(String::as_str), + Some("application/json") + ); + } #[test] fn upstream_error_url_sanitization_removes_secrets_everywhere() { let upstream_url = - "https://api.example.test/v1/messages?key=query-secret&alt=sse#fragment-secret"; + "https://upstream-user:upstream-password@api.example.test/v1/messages?key=query-secret&alt=sse#fragment-secret"; let detail = format!( "error sending request for url ({upstream_url}); source repeated {upstream_url}" ); @@ -4374,6 +5966,100 @@ mod tests { ); assert!(!sanitized_detail.contains("query-secret")); assert!(!sanitized_detail.contains("fragment-secret")); + assert!(!sanitized_detail.contains("upstream-user")); + assert!(!sanitized_detail.contains("upstream-password")); + } + + #[test] + fn upstream_error_detail_redacts_embedded_proxy_urls_and_is_bounded() { + let detail = format!( + "proxy=https://proxy-user:proxy-password@10.0.0.8:8443/connect?token=secret#fragment {}", + "diagnostic ".repeat(400) + ); + let sanitized = super::sanitize_error_detail(&detail); + + assert!(!sanitized.contains("proxy-password")); + assert!(!sanitized.contains("token=secret")); + assert!(!sanitized.contains("10.0.0.8")); + assert!(sanitized.len() <= super::MAX_UPSTREAM_ERROR_DETAIL_BYTES + 3); + } + + #[test] + fn transport_error_debug_sanitizes_dynamic_url_details() { + let secret_url = "https://upstream-user:upstream-password@127.0.0.1:8443/path?token=query-secret#fragment-secret"; + let upstream = ExecutionRuntimeTransportError::UpstreamRequest(format!( + "request failed for url={secret_url}" + )); + let upstream_debug = format!("{upstream:?}"); + assert!(!upstream_debug.contains("upstream-user")); + assert!(!upstream_debug.contains("upstream-password")); + assert!(!upstream_debug.contains("query-secret")); + assert!(!upstream_debug.contains("fragment-secret")); + assert!(!upstream_debug.contains("127.0.0.1")); + assert!(upstream_debug.contains("redacted.invalid")); + let upstream_message = super::safe_transport_error_message(&upstream); + assert!(!upstream_message.contains("upstream-user")); + assert!(!upstream_message.contains("upstream-password")); + assert!(!upstream_message.contains("query-secret")); + assert!(!upstream_message.contains("fragment-secret")); + assert!(!upstream_message.contains("127.0.0.1")); + assert!(upstream_message.contains("redacted.invalid")); + // Display is used by a few legacy error/logging boundaries. Keep it + // safe as well, so a missed `?error`/`safe_transport_error_message` + // conversion cannot reintroduce URL credential leakage. + let upstream_display = format!("{upstream}"); + assert!(!upstream_display.contains("upstream-user")); + assert!(!upstream_display.contains("upstream-password")); + assert!(!upstream_display.contains("query-secret")); + assert!(!upstream_display.contains("fragment-secret")); + assert!(!upstream_display.contains("127.0.0.1")); + assert!(upstream_display.contains("redacted.invalid")); + + let status = ExecutionRuntimeTransportError::UpstreamHttpStatus { + status_code: 502, + message: secret_url.to_string(), + }; + let status_debug = format!("{status:?}"); + assert!(!status_debug.contains("upstream-password")); + assert!(!status_debug.contains("query-secret")); + assert!(!status_debug.contains("127.0.0.1")); + let status_display = format!("{status}"); + assert!(!status_display.contains("upstream-user")); + assert!(!status_display.contains("upstream-password")); + assert!(!status_display.contains("query-secret")); + assert!(!status_display.contains("fragment-secret")); + assert!(!status_display.contains("127.0.0.1")); + + let decode = ExecutionRuntimeTransportError::UpstreamResponseDecode { + encoding: "gzip".to_string(), + message: format!("decode failed for {secret_url}"), + }; + let decode_display = format!("{decode}"); + assert!(!decode_display.contains("upstream-user")); + assert!(!decode_display.contains("upstream-password")); + assert!(!decode_display.contains("query-secret")); + assert!(!decode_display.contains("fragment-secret")); + assert!(!decode_display.contains("127.0.0.1")); + + let relay = ExecutionRuntimeTransportError::RelayError(format!( + "relay failed while contacting {secret_url}" + )); + let relay_display = format!("{relay}"); + assert!(!relay_display.contains("upstream-user")); + assert!(!relay_display.contains("upstream-password")); + assert!(!relay_display.contains("query-secret")); + assert!(!relay_display.contains("fragment-secret")); + assert!(!relay_display.contains("127.0.0.1")); + + let source = reqwest::Proxy::all("http://[") + .expect_err("malformed proxy should produce a reqwest error") + .with_url(reqwest::Url::parse(secret_url).expect("test URL should parse")); + let invalid_proxy = ExecutionRuntimeTransportError::InvalidProxy(source); + let invalid_proxy_debug = format!("{invalid_proxy:?}"); + assert!(!invalid_proxy_debug.contains("upstream-user")); + assert!(!invalid_proxy_debug.contains("upstream-password")); + assert!(!invalid_proxy_debug.contains("query-secret")); + assert!(!invalid_proxy_debug.contains("127.0.0.1")); } #[test] @@ -4499,6 +6185,70 @@ mod tests { ); } + #[test] + fn bounded_execution_body_base64_checks_encoded_and_decoded_sizes() { + let exact = base64::engine::general_purpose::STANDARD.encode([1_u8, 2, 3]); + assert_eq!( + decode_base64_body_with_limit(&exact, 3).expect("exact decoded limit should pass"), + vec![1, 2, 3] + ); + + let encoded_too_large = base64::engine::general_purpose::STANDARD.encode([1_u8, 2, 3, 4]); + assert!(matches!( + decode_base64_body_with_limit(&encoded_too_large, 2), + Err(ExecutionRuntimeTransportError::BodyTooLarge { limit_bytes: 2 }) + )); + + // Eight encoded bytes are within the limit for a four-byte bound, but + // this valid value decodes to six bytes and must fail the second check. + let decoded_too_large = "YWJjZGVm"; + assert!(matches!( + decode_base64_body_with_limit(decoded_too_large, 4), + Err(ExecutionRuntimeTransportError::BodyTooLarge { limit_bytes: 4 }) + )); + + assert!(matches!( + decode_base64_body_with_limit("!!!!", 3), + Err(ExecutionRuntimeTransportError::BodyDecode(_)) + )); + } + + #[test] + fn serialized_json_limit_is_inclusive_and_does_not_allocate_an_encoded_copy() { + let value = json!({"value": "abc"}); + let encoded_len = serde_json::to_vec(&value).unwrap().len(); + + assert!(json_value_fits_serialized_limit(&value, encoded_len)); + assert!(!json_value_fits_serialized_limit(&value, encoded_len - 1)); + } + + #[test] + fn request_body_rejects_ambiguous_json_and_base64_representations() { + let mut plan = tunnel_timeout_plan(false); + plan.body = RequestBody { + json_body: Some(json!({"json": true})), + body_bytes_b64: Some("e30=".to_string()), + body_ref: None, + }; + + assert!(matches!( + super::build_request_body(&plan), + Err(ExecutionRuntimeTransportError::RequestBodyAmbiguous) + )); + } + + #[test] + fn execution_result_envelope_limit_accounts_for_base64_expansion() { + let raw_limit = 64 * 1024 * 1024; + let envelope_limit = execution_result_envelope_limit_bytes(raw_limit); + assert!(envelope_limit > maximum_base64_len_for_decoded_limit(raw_limit)); + assert!(envelope_limit <= 256 * 1024 * 1024); + assert_eq!( + execution_result_envelope_limit_bytes(usize::MAX), + 256 * 1024 * 1024 + ); + } + #[test] fn scoped_response_body_wire_limit_rejects_overflow() { let bounded_plan = with_upstream_response_body_limit( @@ -4618,10 +6368,22 @@ mod tests { 8084, "http://127.0.0.1:8084/v1/messages" )); + assert!(gateway_frontdoor_self_loop_guard_matches_with_port( + 8084, + "http://127.42.0.1:8084/v1/messages" + )); assert!(gateway_frontdoor_self_loop_guard_matches_with_port( 8084, "http://localhost:8084/v1/responses" )); + assert!(gateway_frontdoor_self_loop_guard_matches_with_port( + 8084, + "http://[::ffff:127.0.0.1]:8084/v1/responses" + )); + assert!(gateway_frontdoor_self_loop_guard_matches_with_port( + 8084, + "http://0.0.0.0:8084/v1/responses" + )); assert!(gateway_frontdoor_self_loop_guard_matches_with_port( 8084, "http://localhost:8084/v1internal:streamGenerateContent?alt=sse" @@ -4653,12 +6415,25 @@ mod tests { "http://localhost:8084/v1/responses" ), Some( - "upstream execution target resolves back to the local aether-gateway frontdoor: http://localhost:8084/v1/responses" + "upstream execution target resolves back to the local aether-gateway frontdoor" .to_string() ) ); } + #[test] + fn gateway_frontdoor_self_loop_guard_does_not_echo_target_secrets() { + let error = gateway_frontdoor_self_loop_guard_error_with_port( + 8084, + "http://user:password@localhost:8084/v1/responses?api_key=query-secret#fragment-secret", + ) + .expect("frontdoor self-loop should be rejected"); + assert!(!error.contains("password")); + assert!(!error.contains("query-secret")); + assert!(!error.contains("fragment-secret")); + assert!(!error.contains("localhost:8084")); + } + #[test] fn direct_sync_execution_runtime_builds_clients_for_socks_proxy_urls() { let timeouts = ExecutionTimeouts { @@ -4707,6 +6482,12 @@ mod tests { TestEnvVarGuard { key, previous } } + fn unset_test_env_var(key: &'static str) -> TestEnvVarGuard { + let previous = std::env::var(key).ok(); + std::env::remove_var(key); + TestEnvVarGuard { key, previous } + } + fn direct_reqwest_env_lock() -> MutexGuard<'static, ()> { static LOCK: OnceLock> = OnceLock::new(); LOCK.get_or_init(|| Mutex::new(())) @@ -4779,6 +6560,77 @@ mod tests { )); } + #[test] + fn direct_reqwest_proxy_cache_identity_is_digest_only() { + let proxy_url = "http://alice:proxy-password@proxy.example.test:8080"; + let rotated_proxy_url = "http://alice:rotated-password@proxy.example.test:8080"; + let cache_key = super::direct_reqwest_client_cache_key( + "https://api.example.test/v1/messages", + "key-1", + None, + Some(proxy_url.to_string()), + None, + ExecutionTransportControls::default(), + ); + let rotated = super::direct_reqwest_client_cache_key( + "https://api.example.test/v1/messages", + "key-1", + None, + Some(rotated_proxy_url.to_string()), + None, + ExecutionTransportControls::default(), + ); + + assert_ne!(cache_key, rotated); + assert_eq!(cache_key.proxy_digest.as_deref().map(str::len), Some(64)); + let debug = format!("{cache_key:?}"); + assert!(!debug.contains("alice")); + assert!(!debug.contains("proxy-password")); + assert!(!debug.contains("proxy.example.test")); + super::build_direct_reqwest_client_from_cache_key(&cache_key, Some(proxy_url)) + .expect("authenticated proxy client should build from transient URL material"); + } + + #[test] + fn direct_upstream_origin_brackets_ipv6_literals() { + assert_eq!( + super::direct_reqwest_upstream_origin("https://[::1]:8443/v1/messages").as_deref(), + Some("https://[::1]:8443") + ); + } + + #[test] + fn direct_reqwest_client_cache_evicts_least_recently_used_entry_at_capacity() { + let _guard = direct_reqwest_env_lock(); + let _capacity = set_test_env_var(super::DIRECT_REQWEST_CACHE_MAX_ENTRIES_ENV, "2"); + let cache_key = |suffix| { + super::direct_reqwest_client_cache_key( + "https://api.example.test/v1/messages", + "key-1", + None, + Some(format!("http://proxy-{suffix}.example.test:8080")), + None, + ExecutionTransportControls::default(), + ) + }; + let oldest = cache_key("oldest"); + let recent = cache_key("recent"); + let incoming = cache_key("incoming"); + let mut cache = std::collections::HashMap::new(); + let mut oldest_entry = super::DirectReqwestClientCacheEntry::new(Vec::new(), 1, false); + oldest_entry.last_used = 1; + let mut recent_entry = super::DirectReqwestClientCacheEntry::new(Vec::new(), 1, false); + recent_entry.last_used = 2; + cache.insert(oldest.clone(), oldest_entry); + cache.insert(recent.clone(), recent_entry); + + super::evict_direct_reqwest_client_cache_for_insert(&mut cache, &incoming); + + assert_eq!(cache.len(), 1); + assert!(!cache.contains_key(&oldest)); + assert!(cache.contains_key(&recent)); + } + #[test] fn direct_reqwest_client_cache_key_partitions_key_scoped_pools_by_hashed_key_id() { let profile = ResolvedTransportProfile { @@ -5488,7 +7340,7 @@ mod tests { ]); let controls = resolve_execution_transport_controls(&headers); - assert!(controls.accept_invalid_certs); + assert!(!controls.http1_only); let forwarded = build_request_headers(&headers, None, false) .expect("headers should build after stripping internal controls"); @@ -5736,6 +7588,163 @@ mod tests { } } + const LOCAL_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const LOCAL_TUNNEL_TEST_GENERATION: &str = "transport-test-generation-1"; + + fn authenticated_local_tunnel_test_state() -> AppState { + let node = StoredProxyNode::new( + "node-1".to_string(), + "Node 1".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "online".to_string(), + 30, + 1, + 0, + 0, + 0, + 0, + true, + true, + 1, + ) + .expect("tunnel node should build") + .with_runtime_fields( + None, + None, + None, + None, + Some(json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": LOCAL_TUNNEL_TEST_PSK, + } + })), + None, + None, + None, + None, + None, + None, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()); + let data = crate::data::GatewayDataState::with_proxy_node_repository_for_tests(Arc::new( + InMemoryProxyNodeRepository::seed([node]), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + AppState::new() + .expect("app state should build") + .with_data_state_for_tests(data) + } + + async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, + ) -> Message { + tokio::time::timeout(std::time::Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) + } + + #[test] + fn execution_tunnel_relay_url_policy_allows_https_and_loopback_http() { + let https_proxy = tunnel_proxy_snapshot("https://gateway.example.com/base".to_string()); + let https_url = build_relay_url(Some(&https_proxy), "node-1") + .expect("remote HTTPS relay should be allowed"); + assert_eq!( + https_url, + "https://gateway.example.com/base/api/internal/tunnel/relay/node-1" + ); + + let loopback_proxy = tunnel_proxy_snapshot("http://127.0.0.1:8084".to_string()); + let loopback_url = build_relay_url(Some(&loopback_proxy), "node-1") + .expect("loopback HTTP relay should be allowed"); + assert_eq!( + loopback_url, + "http://127.0.0.1:8084/api/internal/tunnel/relay/node-1" + ); + } + + #[test] + fn execution_tunnel_relay_url_policy_rejects_remote_http() { + let proxy = tunnel_proxy_snapshot("http://gateway.example.com".to_string()); + let error = build_relay_url(Some(&proxy), "node-1") + .expect_err("remote HTTP relay must be rejected"); + + assert!(matches!( + error, + ExecutionRuntimeTransportError::RelayError(message) + if message.contains("HTTPS") && message.contains("loopback") + )); + } + + #[test] + fn relay_target_url_policy_rejects_ambiguous_or_private_targets() { + for rejected in [ + "http://relay.example.test", + "https://10.0.0.8:8443", + "https://127.0.0.1:8443", + "http://localhost:8084?token=secret", + "http://localhost:8084#fragment", + "https://relay-user:relay-password@relay.example.test", + "file:///tmp/relay", + ] { + let url = reqwest::Url::parse(rejected).expect("test URL should parse"); + assert!( + validate_relay_target_url(&url).is_err(), + "relay target should be rejected: {rejected}" + ); + } + + for accepted in [ + "https://relay.example.test:8443/api/internal/tunnel/relay/node-1", + "http://localhost:8084/api/internal/tunnel/relay/node-1", + "http://127.0.0.1:8084/api/internal/tunnel/relay/node-1", + ] { + let url = reqwest::Url::parse(accepted).expect("test URL should parse"); + validate_relay_target_url(&url) + .unwrap_or_else(|error| panic!("relay target should pass: {accepted}: {error}")); + } + } + + #[test] + fn limited_json_body_serialization_rejects_before_growing_to_full_body() { + let body = json!({"payload": "x".repeat(1024)}); + assert!(matches!( + super::serialize_json_body_with_limit(&body, 32), + Err(ExecutionRuntimeTransportError::BodyTooLarge { limit_bytes: 32 }) + )); + let encoded = super::serialize_json_body_with_limit(&json!({"ok": true}), 32) + .expect("small JSON body should pass"); + assert_eq!(encoded, br#"{"ok":true}"#); + } + + #[test] + fn relay_envelope_rejects_oversized_metadata_before_length_cast() { + let meta = RelayRequestMeta { + provider_id: "provider".to_string(), + endpoint_id: "endpoint".to_string(), + key_id: "key".to_string(), + method: "POST".to_string(), + url: "https://relay.example.test".to_string(), + headers: BTreeMap::from([("x-large".to_string(), "x".repeat(300 * 1024))]), + stream: false, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 60, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }; + assert!(matches!( + super::build_relay_envelope(meta, &[]), + Err(ExecutionRuntimeTransportError::RelayError(message)) + if message.contains("metadata") + )); + } + fn manual_proxy_snapshot(node_id: &str) -> ProxySnapshot { ProxySnapshot { enabled: Some(true), @@ -5743,7 +7752,9 @@ mod tests { node_id: Some(node_id.to_string()), label: Some("manual-proxy".into()), url: Some("http://127.0.0.1:1".into()), - extra: None, + extra: Some(json!({ + PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY: "test-manual-generation" + })), } } @@ -5766,9 +7777,23 @@ mod tests { 0, ) .expect("manual proxy node should build") + .with_tunnel_generation("test-manual-generation".into()) .with_manual_proxy_fields(Some("http://127.0.0.1:1".into()), None, None) } + #[test] + fn manual_proxy_binding_requires_incarnation_generation() { + let proxy = ProxySnapshot { + enabled: Some(true), + mode: Some("http".into()), + node_id: Some("manual-node".into()), + label: None, + url: Some("http://127.0.0.1:1".into()), + extra: None, + }; + assert!(super::manual_proxy_node_binding(Some(&proxy)).is_none()); + } + fn decode_relay_envelope(body: &[u8]) -> (serde_json::Value, Vec) { assert!( body.len() >= 4, @@ -6804,41 +8829,58 @@ mod tests { #[tokio::test] async fn direct_sync_execution_runtime_supports_tunnel_relay() { + let _env_lock = direct_reqwest_env_lock(); + let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET); let listener = crate::test_support::bind_loopback_listener() .await .expect("listener should bind"); let addr = listener.local_addr().expect("local addr should resolve"); let app = Router::new().route( "/api/internal/tunnel/relay/{node_id}", - post(|Path(node_id): Path, body: Bytes| async move { - let (meta, request_body) = decode_relay_envelope(&body); - assert_eq!(node_id, "node-1"); - assert_eq!(meta["method"], "POST"); - assert_eq!(meta["url"], "https://example.com/chat"); - let headers = meta["headers"] - .as_object() - .expect("relay meta headers should be an object"); - assert!( - !headers.contains_key(EXECUTION_RUNTIME_LOOP_GUARD_HEADER), - "tunnel relay metadata must not leak internal execution loop guard headers" - ); - let via = headers - .get("via") - .and_then(|value| value.as_str()) - .unwrap_or_default(); - assert!( + post( + |Path(node_id): Path, headers: AxumHeaderMap, body: Bytes| async move { + let (meta, request_body) = decode_relay_envelope(&body); + assert_eq!(node_id, "node-1"); + for name in [ + TUNNEL_RELAY_AUTH_SENDER_HEADER, + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + TUNNEL_RELAY_AUTH_NONCE_HEADER, + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + ] { + assert!( + headers.contains_key(name), + "missing relay auth header {name}" + ); + } + assert_eq!(meta["method"], "POST"); + assert_eq!(meta["url"], "https://example.com/chat"); + let headers = meta["headers"] + .as_object() + .expect("relay meta headers should be an object"); + assert!( + !headers.contains_key(EXECUTION_RUNTIME_LOOP_GUARD_HEADER), + "tunnel relay metadata must not leak internal execution loop guard headers" + ); + let via = headers + .get("via") + .and_then(|value| value.as_str()) + .unwrap_or_default(); + assert!( !via.to_ascii_lowercase() .contains(EXECUTION_RUNTIME_LOOP_GUARD_VIA_TOKEN), "tunnel relay metadata must not leak internal execution runtime Via markers" ); - let request_json: serde_json::Value = - serde_json::from_slice(&request_body).expect("request body should be json"); - assert_eq!(request_json["model"], "gpt-4.1"); - ( - axum::http::StatusCode::OK, - Json(json!({"tunnel": true, "node_id": node_id})), - ) - }), + let request_json: serde_json::Value = + serde_json::from_slice(&request_body).expect("request body should be json"); + assert_eq!(request_json["model"], "gpt-4.1"); + ( + axum::http::StatusCode::OK, + Json(json!({"tunnel": true, "node_id": node_id})), + ) + }, + ), ); let server = tokio::spawn(async move { axum::serve(listener, app) @@ -6885,21 +8927,175 @@ mod tests { ); } + #[tokio::test] + async fn tunnel_relay_client_never_forwards_signed_envelopes_across_redirects() { + let redirected_hits = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let redirected_hits_clone = Arc::clone(&redirected_hits); + let redirected_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect target listener should bind"); + let redirected_addr = redirected_listener + .local_addr() + .expect("redirect target address should resolve"); + let redirected_app = Router::new().route( + "/captured", + post(move || { + let hits = Arc::clone(&redirected_hits_clone); + async move { + hits.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + axum::http::StatusCode::OK + } + }), + ); + let redirected_server = tokio::spawn(async move { + axum::serve(redirected_listener, redirected_app) + .await + .expect("redirect target server should run"); + }); + + let redirect_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect source listener should bind"); + let redirect_addr = redirect_listener + .local_addr() + .expect("redirect source address should resolve"); + let location = format!("http://{redirected_addr}/captured"); + let redirect_app = Router::new().route( + "/relay", + post(move || { + let location = location.clone(); + async move { + ( + axum::http::StatusCode::TEMPORARY_REDIRECT, + [(axum::http::header::LOCATION, location)], + ) + } + }), + ); + let redirect_server = tokio::spawn(async move { + axum::serve(redirect_listener, redirect_app) + .await + .expect("redirect source server should run"); + }); + + let response = build_relay_client(None) + .expect("relay client should build") + .post(format!("http://{redirect_addr}/relay")) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, "sensitive-signature") + .body("sensitive-relay-envelope") + .send() + .await + .expect("relay client should return the redirect response"); + + assert_eq!(response.status(), reqwest::StatusCode::TEMPORARY_REDIRECT); + assert_eq!(redirected_hits.load(std::sync::atomic::Ordering::SeqCst), 0); + + redirect_server.abort(); + redirected_server.abort(); + } + + #[tokio::test] + async fn direct_sync_execution_runtime_rejects_short_tunnel_relay_secret_before_send() { + let _env_lock = direct_reqwest_env_lock(); + let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", &"x".repeat(31)); + let execution_runtime = DirectSyncExecutionRuntime::new(); + let error = execution_runtime + .execute_sync(&ExecutionPlan { + request_id: "req-short-relay-secret".into(), + candidate_id: None, + provider_name: None, + provider_id: "prov-1".into(), + endpoint_id: "ep-1".into(), + key_id: "key-1".into(), + method: "POST".into(), + url: "https://example.com/chat".into(), + headers: BTreeMap::from([("content-type".into(), "application/json".into())]), + content_type: Some("application/json".into()), + content_encoding: None, + body: RequestBody::from_json(json!({"model": "gpt-4.1"})), + stream: false, + client_api_format: "openai:chat".into(), + provider_api_format: "openai:chat".into(), + model_name: Some("gpt-4.1".into()), + proxy: Some(tunnel_proxy_snapshot("http://127.0.0.1:1".to_string())), + transport_profile: None, + timeouts: Some(ExecutionTimeouts { + connect_ms: Some(5_000), + total_ms: Some(LOCAL_HTTP_SUCCESS_TIMEOUT_MS), + ..ExecutionTimeouts::default() + }), + }) + .await + .expect_err("short relay secret must fail closed"); + + assert!(matches!( + error, + ExecutionRuntimeTransportError::RelayError(message) + if message.contains("at least 32 bytes") + )); + } + + #[tokio::test] + async fn direct_sync_execution_runtime_requires_tunnel_relay_secret_before_send() { + let _env_lock = direct_reqwest_env_lock(); + let _relay_secret = unset_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET"); + let execution_runtime = DirectSyncExecutionRuntime::new(); + let error = execution_runtime + .execute_sync(&ExecutionPlan { + request_id: "req-missing-relay-secret".into(), + candidate_id: None, + provider_name: None, + provider_id: "prov-1".into(), + endpoint_id: "ep-1".into(), + key_id: "key-1".into(), + method: "POST".into(), + url: "https://example.com/chat".into(), + headers: BTreeMap::from([("content-type".into(), "application/json".into())]), + content_type: Some("application/json".into()), + content_encoding: None, + body: RequestBody::from_json(json!({"model": "gpt-4.1"})), + stream: false, + client_api_format: "openai:chat".into(), + provider_api_format: "openai:chat".into(), + model_name: Some("gpt-4.1".into()), + proxy: Some(tunnel_proxy_snapshot("http://127.0.0.1:1".to_string())), + transport_profile: None, + timeouts: Some(ExecutionTimeouts { + connect_ms: Some(5_000), + total_ms: Some(LOCAL_HTTP_SUCCESS_TIMEOUT_MS), + ..ExecutionTimeouts::default() + }), + }) + .await + .expect_err("missing relay secret must fail closed"); + + assert!(matches!( + error, + ExecutionRuntimeTransportError::RelayError(message) + if message.contains("AETHER_TUNNEL_RELAY_AUTH_SECRET") + && message.contains("required") + )); + } + #[tokio::test] async fn execute_sync_plan_prefers_local_tunnel_stream_over_http_relay_loopback() { - let state = AppState::new().expect("app state should build"); + let state = authenticated_local_tunnel_test_state(); let tunnel_app = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_app.hub.register_proxy(Arc::new(TunnelProxyConn::new( - 701, - "node-1".to_string(), - "Node 1".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + tunnel_app.hub.register_proxy(Arc::new( + TunnelProxyConn::new( + 701, + "node-1".to_string(), + "Node 1".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let plan = ExecutionPlan { request_id: "req-local-tunnel-1".into(), @@ -6933,7 +9129,7 @@ mod tests { execute_sync_plan(&state_for_task, Some("trace-local-tunnel"), &plan_for_task).await }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -6949,7 +9145,7 @@ mod tests { assert_eq!(request_meta.method, "POST"); assert_eq!(request_meta.url, "https://example.com/chat"); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -7168,26 +9364,31 @@ mod tests { #[tokio::test] async fn direct_sync_execution_runtime_forwards_http1_only_control_to_tunnel_relay() { + let _env_lock = direct_reqwest_env_lock(); + let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET); let listener = crate::test_support::bind_loopback_listener() .await .expect("listener should bind"); let addr = listener.local_addr().expect("local addr should resolve"); let app = Router::new().route( "/api/internal/tunnel/relay/{node_id}", - post(|Path(node_id): Path, body: Bytes| async move { - let (meta, request_body) = decode_relay_envelope(&body); - assert_eq!(node_id, "node-1"); - assert_eq!(meta["provider_id"], "prov-1"); - assert_eq!(meta["endpoint_id"], "ep-1"); - assert_eq!(meta["key_id"], "key-1"); - assert_eq!(meta["http1_only"], true); - assert_eq!(meta["follow_redirects"], json!(false)); - assert_eq!(meta["transport_profile"]["profile_id"], "relay-profile"); - let request_json: serde_json::Value = - serde_json::from_slice(&request_body).expect("request body should be json"); - assert_eq!(request_json["model"], "gpt-4.1"); - (axum::http::StatusCode::OK, Json(json!({"ok": true}))) - }), + post( + |Path(node_id): Path, headers: AxumHeaderMap, body: Bytes| async move { + let (meta, request_body) = decode_relay_envelope(&body); + assert_eq!(node_id, "node-1"); + assert!(headers.contains_key(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER)); + assert_eq!(meta["provider_id"], "prov-1"); + assert_eq!(meta["endpoint_id"], "ep-1"); + assert_eq!(meta["key_id"], "key-1"); + assert_eq!(meta["http1_only"], true); + assert_eq!(meta["follow_redirects"], json!(false)); + assert_eq!(meta["transport_profile"]["profile_id"], "relay-profile"); + let request_json: serde_json::Value = + serde_json::from_slice(&request_body).expect("request body should be json"); + assert_eq!(request_json["model"], "gpt-4.1"); + (axum::http::StatusCode::OK, Json(json!({"ok": true}))) + }, + ), ); let server = tokio::spawn(async move { axum::serve(listener, app) @@ -7373,7 +9574,7 @@ mod tests { .map(|profile| profile.http_mode.as_str()), Some(TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE) ); - super::build_direct_reqwest_client_from_cache_key(&cache_key) + super::build_direct_reqwest_client_from_cache_key(&cache_key, None) .expect("h2c prior-knowledge client should build"); } diff --git a/apps/aether-gateway/src/execution_runtime/transport_failure.rs b/apps/aether-gateway/src/execution_runtime/transport_failure.rs index ccd532adf..c3f1885fb 100644 --- a/apps/aether-gateway/src/execution_runtime/transport_failure.rs +++ b/apps/aether-gateway/src/execution_runtime/transport_failure.rs @@ -22,6 +22,7 @@ const TRANSPORT_ERROR_CLIENT_MESSAGE: &str = #[derive(Debug, Default)] pub(crate) struct StreamCandidateWatchdogProgress { terminal_started: AtomicBool, + abandoned: AtomicBool, } tokio::task_local! { @@ -37,6 +38,24 @@ impl StreamCandidateWatchdogProgress { self.terminal_started.load(Ordering::Acquire) } + /// The watchdog gave up waiting and settles this attempt itself. + /// + /// The attempt future is dropped once the watchdog returns, so its own + /// cancellation guard must stay out of the way instead of racing the + /// watchdog's terminal rows with a cancellation. + pub(crate) fn mark_abandoned(&self) { + self.abandoned.store(true, Ordering::Release); + } + + pub(crate) fn abandoned(&self) -> bool { + self.abandoned.load(Ordering::Acquire) + } + + /// The watchdog watching the attempt on this task, if it runs under one. + pub(crate) fn current() -> Option> { + STREAM_CANDIDATE_WATCHDOG_PROGRESS.try_with(Arc::clone).ok() + } + pub(crate) async fn scope(self: Arc, future: F) -> F::Output where F: Future, diff --git a/apps/aether-gateway/src/execution_runtime/windsurf.rs b/apps/aether-gateway/src/execution_runtime/windsurf.rs index 98b695337..2810eb229 100644 --- a/apps/aether-gateway/src/execution_runtime/windsurf.rs +++ b/apps/aether-gateway/src/execution_runtime/windsurf.rs @@ -1,6 +1,6 @@ use std::collections::{BTreeMap, HashMap, HashSet}; use std::fs; -use std::io::{Error as IoError, Read, Seek, SeekFrom}; +use std::io::{self, Error as IoError, Write}; use std::net::{SocketAddr, TcpListener, TcpStream}; use std::path::{Path, PathBuf}; use std::process::{Child, Command, Stdio}; @@ -36,12 +36,13 @@ use tracing::{debug, info, warn}; use uuid::Uuid; use super::ndjson::encode_stream_frame_ndjson; -use super::transport::{with_non_stream_total_timeout, ExecutionRuntimeTransportError}; +use super::transport::{ + with_non_stream_total_timeout, ExecutionRuntimeTransportError, UpstreamResponseBodyPhase, +}; use crate::AppState; const LS_SERVICE: &str = "/exa.language_server_pb.LanguageServerService"; const DEFAULT_LS_PORT: u16 = 42100; -const DEFAULT_CSRF_TOKEN: &str = "windsurf-api-csrf-fixed-token"; const DEFAULT_CODEIUM_API_URL: &str = "https://server.self-serve.windsurf.com"; const DEFAULT_REGISTER_USER_URL: &str = "https://api.codeium.com/register_user/"; const POLL_INTERVAL: Duration = Duration::from_millis(500); @@ -55,6 +56,8 @@ const WINDOWS_LS_READY_TIMEOUT: Duration = Duration::from_secs(90); const GRPC_SHORT_TIMEOUT: Duration = Duration::from_secs(5); const GRPC_STATUS_TIMEOUT: Duration = Duration::from_secs(10); const GRPC_REQUEST_TIMEOUT: Duration = Duration::from_secs(45); +const WINDSURF_GRPC_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024 * 1024; +const WINDSURF_STREAM_FRAME_CHANNEL_CAPACITY: usize = 64; const SEND_CASCADE_MAX_RETRIES: usize = 3; const WARMUP_TRANSPORT_MAX_RESTARTS: usize = 2; const WORKSPACE_PATH_HINT: &str = "Workspace path hidden; \"\" is a redaction marker, NOT a path. Use tool calls to inspect real files or execute commands."; @@ -65,7 +68,7 @@ pub(crate) struct WindsurfNativeStream { pub(crate) report_context: Option, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] struct WindsurfRequestInput { api_key: String, model: String, @@ -77,6 +80,28 @@ struct WindsurfRequestInput { native_bridge: Option, } +impl std::fmt::Debug for WindsurfRequestInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("WindsurfRequestInput") + .field("api_key", &"[REDACTED]") + .field("model", &self.model) + .field("message", &"[REDACTED]") + .field("image_count", &self.images.len()) + .field("tool_count", &self.tools.len()) + .field( + "tool_preamble", + &self.tool_preamble.as_ref().map(|_| "[REDACTED]"), + ) + .field("tool_dialect", &self.tool_dialect) + .field( + "native_bridge", + &self.native_bridge.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq)] struct WindsurfToolDefinition { name: String, @@ -148,18 +173,31 @@ enum WindsurfPollEvent { Heartbeat, } -#[derive(Debug)] struct LsProcessEntry { port: u16, csrf_token: String, session_id: String, workspace_path: PathBuf, proxy_url: Option, - stderr_log_path: Option, + stderr_capture_enabled: bool, _child: Child, } -#[derive(Debug, Clone)] +impl std::fmt::Debug for LsProcessEntry { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("LsProcessEntry") + .field("port", &self.port) + .field("csrf_token", &"[REDACTED]") + .field("session_id", &"[REDACTED]") + .field("workspace_path", &"[REDACTED]") + .field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]")) + .field("stderr_capture_enabled", &self.stderr_capture_enabled) + .finish_non_exhaustive() + } +} + +#[derive(Clone)] struct LsHandle { pool_key: String, port: u16, @@ -168,6 +206,19 @@ struct LsHandle { workspace_path: PathBuf, } +impl std::fmt::Debug for LsHandle { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("LsHandle") + .field("pool_key", &self.pool_key) + .field("port", &self.port) + .field("csrf_token", &"[REDACTED]") + .field("session_id", &"[REDACTED]") + .field("workspace_path", &"[REDACTED]") + .finish() + } +} + #[derive(Clone)] struct PreparedCascade { plan: ExecutionPlan, @@ -220,16 +271,17 @@ pub(crate) async fn maybe_execute_windsurf_sync( let key_upstream_metadata = read_windsurf_key_upstream_metadata(state, plan).await; let prepared = prepare_windsurf_cascade(plan, input, key_upstream_metadata).await?; let started_at = Instant::now(); - let mut deltas = Vec::new(); + let mut content = String::new(); + let buffered_body_limit = crate::headers::max_internal_buffered_body_bytes(); let poll_result = poll_windsurf_cascade_with_transport_recovery(&prepared, |event| { if let WindsurfPollEvent::TextDelta(delta) = event { - deltas.push(sanitize_windsurf_text(&delta)); + let delta = sanitize_windsurf_text(&delta); + append_windsurf_delta_with_limit(&mut content, &delta, buffered_body_limit)?; } Ok(()) }) .await?; let elapsed_ms = started_at.elapsed().as_millis() as u64; - let content = deltas.concat(); let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input); let mut tool_calls = poll_result.native_tool_calls; tool_calls.extend(parsed_tool_calls.tool_calls); @@ -289,10 +341,9 @@ async fn prepare_windsurf_cascade( ) -> Result { let model = resolve_windsurf_execution_model(&input.model, key_upstream_metadata.as_ref()) .ok_or_else(|| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "unsupported Windsurf model {}", - input.model - )) + ExecutionRuntimeTransportError::UpstreamRequest( + "unsupported Windsurf model".to_string(), + ) })?; let mut ls = ensure_windsurf_language_server(plan).await?; ls = warmup_windsurf_cascade_with_transport_recovery(plan, ls, &input.api_key).await?; @@ -304,7 +355,7 @@ async fn prepare_windsurf_cascade( event_name = "windsurf_panel_state_missing_on_start", log_type = "ops", request_id = %plan.request_id, - error = %err, + error_category = "panel_state_missing", "gateway rewarming Windsurf language server after missing panel state on StartCascade" ); ls = force_rewarm_windsurf_cascade(plan, &ls, &input.api_key).await?; @@ -360,7 +411,11 @@ async fn prepare_windsurf_cascade( &ls.session_id, &send_options, ) - .map_err(|err| ExecutionRuntimeTransportError::UpstreamRequest(err.to_string()))?; + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to build Windsurf cascade message".to_string(), + ) + })?; match windsurf_grpc_unary( ls.port, &ls.csrf_token, @@ -374,17 +429,16 @@ async fn prepare_windsurf_cascade( Err(err) if is_windsurf_send_retryable_error(&err) => { send_retry += 1; if send_retry > SEND_CASCADE_MAX_RETRIES { - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf SendUserCascadeMessage retry limit exceeded after {} rewarm attempts: {err}", - SEND_CASCADE_MAX_RETRIES - ))); + return Err(ExecutionRuntimeTransportError::UpstreamRequest( + "Windsurf SendUserCascadeMessage retry limit exceeded".to_string(), + )); } warn!( event_name = "windsurf_send_retryable_error", log_type = "ops", request_id = %plan.request_id, retry = send_retry, - error = %err, + error_category = "retryable_send_failure", "gateway rewarming Windsurf cascade after retryable SendUserCascadeMessage error" ); ls = force_rewarm_windsurf_cascade(plan, &ls, &input.api_key).await?; @@ -425,14 +479,14 @@ async fn read_windsurf_key_upstream_metadata( .into_iter() .find(|key| key.id == plan.key_id && key.provider_id == plan.provider_id) .and_then(|key| key.upstream_metadata), - Err(err) => { + Err(_) => { warn!( event_name = "windsurf_key_upstream_metadata_unavailable", log_type = "ops", request_id = %plan.request_id, provider_id = %plan.provider_id, key_id = %plan.key_id, - error = ?err, + error_category = "provider_catalog_read_failed", "gateway could not read Windsurf key upstream metadata; falling back to static model catalog" ); None @@ -532,9 +586,12 @@ fn build_windsurf_stream_frame_stream( }, }); - let (tx, mut rx) = mpsc::unbounded_channel::>(); + let (tx, mut rx) = mpsc::channel::>( + WINDSURF_STREAM_FRAME_CHANNEL_CAPACITY, + ); tokio::spawn(async move { - let mut deltas = Vec::new(); + let mut content = String::new(); + let buffered_body_limit = crate::headers::max_internal_buffered_body_bytes(); let mut streamed_native_call_ids = HashSet::new(); let mut next_tool_index = 0usize; let buffer_for_tool_calls = should_parse_windsurf_tool_calls(&prepared.input); @@ -542,7 +599,11 @@ fn build_windsurf_stream_frame_stream( match event { WindsurfPollEvent::TextDelta(delta) => { let delta = sanitize_windsurf_text(&delta); - deltas.push(delta.clone()); + append_windsurf_delta_with_limit( + &mut content, + &delta, + buffered_body_limit, + )?; if !buffer_for_tool_calls { send_stream_frame(&tx, sse_data_frame(&prepared.request_id, &prepared.model, &delta))?; } @@ -570,7 +631,6 @@ fn build_windsurf_stream_frame_stream( match poll_result { Ok(poll_result) => { - let content = deltas.concat(); let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input); let mut tool_calls = poll_result .native_tool_calls @@ -590,30 +650,46 @@ fn build_windsurf_stream_frame_stream( next_tool_index, &tool_calls, ) { - let _ = tx.send(encode_stream_frame_ndjson(&frame)); + if tx.send(encode_stream_frame_ndjson(&frame)).await.is_err() { + return; + } } - let _ = tx.send(encode_stream_frame_ndjson(&sse_finish_frame_with_reason( + if tx.send(encode_stream_frame_ndjson(&sse_finish_frame_with_reason( &prepared.request_id, &prepared.model, "tool_calls", - ))); + ))).await.is_err() { + return; + } } else { if buffer_for_tool_calls && !content.is_empty() { - let _ = tx.send(encode_stream_frame_ndjson(&sse_data_frame( + if tx.send(encode_stream_frame_ndjson(&sse_data_frame( &prepared.request_id, &prepared.model, &content, - ))); + ))).await.is_err() { + return; + } } - let _ = tx.send(encode_stream_frame_ndjson(&sse_finish_frame_with_reason( + if tx.send(encode_stream_frame_ndjson(&sse_finish_frame_with_reason( &prepared.request_id, &prepared.model, finish_reason, - ))); + ))).await.is_err() { + return; + } + } + if tx + .send(encode_stream_frame_ndjson(&raw_sse_data_frame( + b"data: [DONE]\n\n", + ))) + .await + .is_err() + { + return; } - let _ = tx.send(encode_stream_frame_ndjson(&raw_sse_data_frame(b"data: [DONE]\n\n"))); let elapsed_ms = started_at.elapsed().as_millis() as u64; - let _ = tx.send(encode_stream_frame_ndjson(&StreamFrame { + if tx.send(encode_stream_frame_ndjson(&StreamFrame { frame_type: StreamFrameType::Telemetry, payload: StreamFramePayload::Telemetry { telemetry: ExecutionTelemetry { @@ -622,25 +698,29 @@ fn build_windsurf_stream_frame_stream( upstream_bytes: Some(content.len() as u64), }, }, - })); + })).await.is_err() { + return; + } let _ = tx.send(encode_stream_frame_ndjson(&StreamFrame::eof_with_summary( windsurf_terminal_summary( poll_result.usage, Some(prepared.model.as_str()), Some(finish_reason), ), - ))); + ))).await; } Err(err) => { let execution_error = windsurf_execution_error_from_transport_error(&err, ExecutionPhase::StreamRead); - let _ = tx.send(encode_stream_frame_ndjson(&StreamFrame { + if tx.send(encode_stream_frame_ndjson(&StreamFrame { frame_type: StreamFrameType::Error, payload: StreamFramePayload::Error { error: execution_error, }, - })); - let _ = tx.send(encode_stream_frame_ndjson(&StreamFrame::eof())); + })).await.is_err() { + return; + } + let _ = tx.send(encode_stream_frame_ndjson(&StreamFrame::eof())).await; } } }); @@ -651,23 +731,43 @@ fn build_windsurf_stream_frame_stream( } } +fn append_windsurf_delta_with_limit( + content: &mut String, + delta: &str, + limit_bytes: usize, +) -> Result<(), ExecutionRuntimeTransportError> { + if delta.len() > limit_bytes.saturating_sub(content.len()) { + return Err(ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Decoded, + limit_bytes, + }); + } + content.push_str(delta); + Ok(()) +} + fn send_stream_frame( - tx: &mpsc::UnboundedSender>, + tx: &mpsc::Sender>, frame: StreamFrame, ) -> Result<(), ExecutionRuntimeTransportError> { - tx.send(encode_stream_frame_ndjson(&frame)).map_err(|_| { - ExecutionRuntimeTransportError::UpstreamRequest( - "Windsurf stream cancelled by downstream client".to_string(), - ) - }) + tx.try_send(encode_stream_frame_ndjson(&frame)) + .map_err(|err| { + let detail = match err { + mpsc::error::TrySendError::Full(_) => "Windsurf stream buffer is full", + mpsc::error::TrySendError::Closed(_) => { + "Windsurf stream cancelled by downstream client" + } + }; + ExecutionRuntimeTransportError::UpstreamRequest(detail.to_string()) + }) } fn windsurf_execution_error_from_transport_error( err: &ExecutionRuntimeTransportError, phase: ExecutionPhase, ) -> ExecutionError { - let message = err.to_string(); - let lower = message.to_ascii_lowercase(); + let lower = err.to_string().to_ascii_lowercase(); + let message = windsurf_public_error_message(err, &lower); if lower.contains("stream cancelled by downstream client") { return ExecutionError { kind: ExecutionErrorKind::Cancelled, @@ -678,6 +778,20 @@ fn windsurf_execution_error_from_transport_error( failover_recommended: false, }; } + if matches!( + err, + ExecutionRuntimeTransportError::UpstreamResponseTooLarge { .. } + ) || lower.contains("stream buffer is full") + { + return ExecutionError { + kind: ExecutionErrorKind::ProtocolError, + phase, + message, + upstream_status: None, + retryable: false, + failover_recommended: false, + }; + } if lower.contains("reached message rate limit") || lower.contains("resource_exhausted") || lower.contains("rate limit") @@ -713,6 +827,76 @@ fn windsurf_execution_error_from_transport_error( } } +fn windsurf_public_error_message(err: &ExecutionRuntimeTransportError, normalized: &str) -> String { + if normalized.contains("stream cancelled by downstream client") { + return "Windsurf request was cancelled".to_string(); + } + if matches!( + err, + ExecutionRuntimeTransportError::UpstreamResponseTooLarge { .. } + ) || normalized.contains("stream buffer is full") + { + return "Windsurf response exceeded local runtime limits".to_string(); + } + if normalized.contains("reached message rate limit") + || normalized.contains("resource_exhausted") + || normalized.contains("rate limit") + || normalized.contains("rate_limit") + { + return "Windsurf upstream rate limit reached".to_string(); + } + if normalized.contains("unsupported windsurf model") { + return "Unsupported Windsurf model".to_string(); + } + if normalized.contains("proxy") { + return "Windsurf proxy request failed".to_string(); + } + if ["tls", "ssl", "certificate", "handshake"] + .iter() + .any(|needle| normalized.contains(needle)) + { + return "Windsurf TLS request failed".to_string(); + } + if ["timed out", "timeout", "deadline exceeded"] + .iter() + .any(|needle| normalized.contains(needle)) + { + return "Windsurf request timed out".to_string(); + } + "Windsurf request failed".to_string() +} + +fn windsurf_trajectory_error(detail: &str) -> ExecutionRuntimeTransportError { + let normalized = detail.trim().to_ascii_lowercase(); + let message = if normalized.contains("reached message rate limit") + || normalized.contains("resource_exhausted") + || normalized.contains("rate limit") + || normalized.contains("rate_limit") + { + "Windsurf trajectory reported a rate limit" + } else if normalized.contains("panel state") + && (normalized.contains("not found") || normalized.contains("not_found")) + { + "Windsurf panel state not found" + } else if normalized.contains("untrusted workspace") + || (normalized.contains("workspace") + && normalized.contains("not") + && normalized.contains("trusted")) + { + "Windsurf workspace is not trusted" + } else if (normalized.contains("cascade") || normalized.contains("trajectory")) + && (normalized.contains("not found") + || normalized.contains("not_found") + || normalized.contains("expired") + || normalized.contains("unknown")) + { + "Windsurf cascade not found or expired" + } else { + "Windsurf trajectory reported an upstream error" + }; + ExecutionRuntimeTransportError::UpstreamRequest(message.to_string()) +} + fn classify_windsurf_transport_execution_error( message: &str, phase: ExecutionPhase, @@ -794,7 +978,7 @@ where request_id = %prepared.request_id, cascade_id = %prepared.cascade_id, port = prepared.ls.port, - error = %first_err, + error_category = "poll_transport_failure", "gateway restarting Windsurf language server after pre-output polling transport failure" ); invalidate_windsurf_language_server_handle( @@ -846,17 +1030,15 @@ where GRPC_REQUEST_TIMEOUT, ) .await?; - let steps = parse_trajectory_steps(&steps_response).map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to parse Windsurf trajectory steps: {err}" - )) + let steps = parse_trajectory_steps(&steps_response).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to parse Windsurf trajectory steps".to_string(), + ) })?; for step in &steps { if step.step_type == 17 && !step.error_text.trim().is_empty() { - return Err(ExecutionRuntimeTransportError::UpstreamRequest( - step.error_text.trim().to_string(), - )); + return Err(windsurf_trajectory_error(&step.error_text)); } } for (index, step) in steps.iter().enumerate() { @@ -925,10 +1107,10 @@ where GRPC_REQUEST_TIMEOUT, ) .await?; - let final_steps = parse_trajectory_steps(&final_steps_response).map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to parse final Windsurf trajectory steps: {err}" - )) + let final_steps = parse_trajectory_steps(&final_steps_response).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to parse final Windsurf trajectory steps".to_string(), + ) })?; for (index, step) in final_steps.iter().enumerate() { if let Some(usage) = step.usage { @@ -995,13 +1177,13 @@ async fn fetch_windsurf_generator_usage(prepared: &PreparedCascade) -> Option response, - Err(err) => { + Err(_) => { debug!( event_name = "windsurf_generator_metadata_fetch_failed", log_type = "debug", request_id = %prepared.request_id, cascade_id = %prepared.cascade_id, - error = %err, + error_category = "generator_metadata_fetch_failed", "gateway could not fetch Windsurf generator token usage" ); return None; @@ -1010,13 +1192,13 @@ async fn fetch_windsurf_generator_usage(prepared: &PreparedCascade) -> Option usage, - Err(err) => { + Err(_) => { debug!( event_name = "windsurf_generator_metadata_parse_failed", log_type = "debug", request_id = %prepared.request_id, cascade_id = %prepared.cascade_id, - error = %err, + error_category = "generator_metadata_parse_failed", "gateway could not parse Windsurf generator token usage" ); None @@ -1194,7 +1376,7 @@ async fn ensure_windsurf_language_server( .and_then(windsurf_language_server_stale_reason) { if let Some(entry) = guard.remove(&key) { - terminate_windsurf_language_server_entry(&key, entry, &reason); + terminate_windsurf_language_server_entry(&key, entry, reason); } } if let Some(entry) = guard.get(&key) { @@ -1203,29 +1385,39 @@ async fn ensure_windsurf_language_server( } let binary_path = resolve_language_server_binary_path()?; - repair_executable_mode(&binary_path); let port = find_free_language_server_port()?; - let data_dir = language_server_data_dir(&key); - let workspace_path = language_server_workspace_path(plan); - fs::create_dir_all(data_dir.join("db")).map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to create Windsurf LS data dir {}: {err}", - data_dir.display() - )) + let data_dir = absolute_runtime_path(language_server_data_dir(&key)).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to resolve Windsurf language server data directory".to_string(), + ) })?; - ensure_workspace_dir(&workspace_path); - let (stderr, stderr_log_path) = language_server_stderr(&data_dir); + let database_dir = data_dir.join("db"); + ensure_private_directory(&data_dir) + .and_then(|_| ensure_private_directory(&database_dir)) + .map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to secure Windsurf language server data directory".to_string(), + ) + })?; + let workspace_path = data_dir.join("workspace"); + ensure_workspace_dir(&workspace_path).map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to secure Windsurf placeholder workspace".to_string(), + ) + })?; + let (stderr, stderr_capture_enabled) = language_server_stderr(&data_dir); let proxy_url = language_server_proxy_url(plan); + let csrf_token = new_windsurf_csrf_token(); let mut command = Command::new(&binary_path); command .arg(format!("--api_server_url={}", codeium_api_url())) .arg("--run_child") .arg(format!("--server_port={port}")) - .arg(format!("--csrf_token={DEFAULT_CSRF_TOKEN}")) + .arg(format!("--csrf_token={csrf_token}")) .arg(format!("--register_user_url={DEFAULT_REGISTER_USER_URL}")) .arg(format!("--codeium_dir={}", data_dir.display())) - .arg(format!("--database_dir={}", data_dir.join("db").display())) + .arg(format!("--database_dir={}", database_dir.display())) .arg("--detect_proxy=false"); if !cfg!(target_os = "windows") { @@ -1238,39 +1430,34 @@ async fn ensure_windsurf_language_server( .stdout(Stdio::null()) .stderr(stderr); - let mut child = command.spawn().map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "failed to start Windsurf language server {}: {err}", - binary_path.display() - )) + let mut child = command.spawn().map_err(|_| { + ExecutionRuntimeTransportError::UpstreamRequest( + "failed to start Windsurf language server".to_string(), + ) })?; - if let Err(err) = wait_language_server_ready(port, &mut child, stderr_log_path.as_deref()).await - { + if let Err(err) = wait_language_server_ready(port, &mut child).await { let _ = child.kill(); return Err(err); } - let stderr_log_display = stderr_log_path - .as_ref() - .map(|path| path.display().to_string()); info!( event_name = "windsurf_language_server_ready", log_type = "ops", port, pool_key = %key, proxy_configured = proxy_url.is_some(), - stderr_log_path = stderr_log_display.as_deref(), + stderr_capture_enabled, "gateway native Windsurf language server ready" ); let entry = LsProcessEntry { port, - csrf_token: DEFAULT_CSRF_TOKEN.to_string(), + csrf_token, session_id: Uuid::new_v4().to_string(), workspace_path, proxy_url, - stderr_log_path, + stderr_capture_enabled, _child: child, }; let mut guard = pool.lock().map_err(|_| { @@ -1283,7 +1470,7 @@ async fn ensure_windsurf_language_server( .and_then(windsurf_language_server_stale_reason) { if let Some(existing) = guard.remove(&key) { - terminate_windsurf_language_server_entry(&key, existing, &reason); + terminate_windsurf_language_server_entry(&key, existing, reason); } } if let Some(existing) = guard.get(&key) { @@ -1301,6 +1488,14 @@ async fn ensure_windsurf_language_server( Ok(ls_handle_from_entry(&key, entry)) } +fn new_windsurf_csrf_token() -> String { + format!( + "aether-{}{}", + Uuid::new_v4().simple(), + Uuid::new_v4().simple() + ) +} + async fn warmup_windsurf_cascade( ls: &LsHandle, api_key: &str, @@ -1361,7 +1556,7 @@ async fn warmup_windsurf_cascade_with_transport_recovery( port = ls.port, attempt = attempt + 1, max_restarts = WARMUP_TRANSPORT_MAX_RESTARTS, - error = %err, + error_category = "warmup_transport_failure", "gateway restarting Windsurf language server after warmup transport failure" ); invalidate_windsurf_language_server_handle(&ls, "warmup transport failure")?; @@ -1389,7 +1584,7 @@ async fn force_rewarm_windsurf_cascade( event_name = "windsurf_rewarm_transport_restart", log_type = "ops", port = refreshed.port, - error = %err, + error_category = "rewarm_transport_failure", "gateway restarting Windsurf language server after rewarm transport failure" ); invalidate_windsurf_language_server_handle(&refreshed, "rewarm transport failure")?; @@ -1438,7 +1633,7 @@ async fn windsurf_warmup_unary( event_name = "windsurf_workspace_trust_update_failed", log_type = "ops", port, - error = %err, + error_category = "workspace_trust_update_failed", "gateway Windsurf workspace trust update failed; continuing to match WindsurfAPI warmup behavior" ); } else { @@ -1447,7 +1642,7 @@ async fn windsurf_warmup_unary( log_type = "ops", port, stage, - error = %err, + error_category = "cascade_warmup_stage_failed", "gateway Windsurf cascade warmup stage failed; continuing to match WindsurfAPI warmup behavior" ); } @@ -1515,12 +1710,12 @@ async fn sync_windsurf_user_status_with_panel(ls: &LsHandle, api_key: &str) { .await { Ok(response) => response, - Err(err) => { + Err(_) => { warn!( event_name = "windsurf_user_status_sync_failed", log_type = "ops", port = ls.port, - error = %err, + error_category = "user_status_sync_failed", "gateway failed to fetch Windsurf user status for panel sync" ); return; @@ -1535,7 +1730,7 @@ async fn sync_windsurf_user_status_with_panel(ls: &LsHandle, api_key: &str) { ); return; }; - if let Err(err) = windsurf_grpc_unary( + if windsurf_grpc_unary( ls.port, &ls.csrf_token, "UpdatePanelStateWithUserStatus", @@ -1547,12 +1742,13 @@ async fn sync_windsurf_user_status_with_panel(ls: &LsHandle, api_key: &str) { GRPC_SHORT_TIMEOUT, ) .await + .is_err() { warn!( event_name = "windsurf_panel_user_status_update_failed", log_type = "ops", port = ls.port, - error = %err, + error_category = "panel_user_status_update_failed", "gateway failed to update Windsurf panel state with user status" ); } @@ -1568,6 +1764,8 @@ async fn windsurf_grpc_unary( let url = format!("http://127.0.0.1:{port}{LS_SERVICE}/{method}"); let client = reqwest::Client::builder() .http1_only() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) .timeout(timeout) .build() .map_err(ExecutionRuntimeTransportError::ClientBuild)?; @@ -1587,19 +1785,44 @@ async fn windsurf_grpc_unary( )) })?; let status = response.status(); - let body = response.bytes().await.map_err(|err| { - ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf Connect {method} response read failed: {}", - super::transport::format_upstream_request_error(&err) - )) - })?; + let body = collect_windsurf_grpc_response_body(response, method).await?; if !status.is_success() { - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf Connect {method} returned HTTP {status}: {}", - String::from_utf8_lossy(&body) - ))); + return Err(windsurf_connect_http_error(method, status.as_u16())); } - Ok(body.to_vec()) + Ok(body) +} + +async fn collect_windsurf_grpc_response_body( + response: reqwest::Response, + method: &str, +) -> Result, ExecutionRuntimeTransportError> { + if response.content_length().is_some_and(|length| { + length > u64::try_from(WINDSURF_GRPC_RESPONSE_BODY_LIMIT_BYTES).unwrap_or(u64::MAX) + }) { + return Err(ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Wire, + limit_bytes: WINDSURF_GRPC_RESPONSE_BODY_LIMIT_BYTES, + }); + } + + let mut body = Vec::new(); + let mut stream = response.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(|err| { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Windsurf Connect {method} response read failed: {}", + super::transport::format_upstream_request_error(&err) + )) + })?; + if chunk.len() > WINDSURF_GRPC_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) { + return Err(ExecutionRuntimeTransportError::UpstreamResponseTooLarge { + phase: UpstreamResponseBodyPhase::Wire, + limit_bytes: WINDSURF_GRPC_RESPONSE_BODY_LIMIT_BYTES, + }); + } + body.extend_from_slice(&chunk); + } + Ok(body) } fn detect_windsurf_request( @@ -3757,21 +3980,21 @@ fn ls_handle_from_entry(pool_key: &str, entry: &LsProcessEntry) -> LsHandle { } } -fn windsurf_language_server_stale_reason(entry: &mut LsProcessEntry) -> Option { +fn windsurf_language_server_stale_reason(entry: &mut LsProcessEntry) -> Option<&'static str> { match entry._child.try_wait() { - Ok(Some(status)) => { - return Some(format!("process exited with status {status}")); + Ok(Some(_)) => { + return Some("process_exited"); } Ok(None) => {} - Err(err) => { - return Some(format!("failed to inspect process status: {err}")); + Err(_) => { + return Some("process_status_unavailable"); } } if language_server_port_accepts(entry.port) { None } else { - Some(format!("port {} is not accepting connections", entry.port)) + Some("port_unavailable") } } @@ -3808,18 +4031,14 @@ fn terminate_windsurf_language_server_entry( ) { let _ = entry._child.kill(); let _ = entry._child.wait(); - let stderr_log_display = entry - .stderr_log_path - .as_ref() - .map(|path| path.display().to_string()); warn!( event_name = "windsurf_language_server_removed", log_type = "ops", pool_key, port = entry.port, proxy_configured = entry.proxy_url.is_some(), - stderr_log_path = stderr_log_display.as_deref(), - reason, + stderr_capture_enabled = entry.stderr_capture_enabled, + reason_category = reason, "gateway removed Windsurf language server from pool" ); } @@ -3849,25 +4068,11 @@ fn language_server_proxy_url(plan: &ExecutionPlan) -> Option { fn resolve_language_server_binary_path() -> Result { for key in ["WINDSURF_LS_BINARY_PATH", "LS_BINARY_PATH"] { if let Some(path) = std::env::var_os(key).filter(|value| !value.is_empty()) { - let path = PathBuf::from(path); - if path.exists() { - return Ok(path); - } - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "{key} points to missing Windsurf language server binary: {}", - path.display() - ))); + return validate_language_server_binary_path(Path::new(&path)); } } let default = default_language_server_binary_path(); - if default.exists() { - Ok(default) - } else { - Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf language server binary not found at {}", - default.display() - ))) - } + validate_language_server_binary_path(&default) } fn default_language_server_binary_path() -> PathBuf { @@ -3892,44 +4097,75 @@ fn default_language_server_binary_path() -> PathBuf { } } -fn language_server_stderr(data_dir: &Path) -> (Stdio, Option) { +fn language_server_stderr(data_dir: &Path) -> (Stdio, bool) { let path = data_dir.join("language-server.stderr.log"); - match fs::OpenOptions::new().create(true).append(true).open(&path) { - Ok(file) => (Stdio::from(file), Some(path)), - Err(err) => { + match open_private_append_file(&path) { + Ok(file) => (Stdio::from(file), true), + Err(_) => { warn!( event_name = "windsurf_language_server_stderr_open_failed", log_type = "ops", - path = %path.display(), - error = %err, + error_category = "stderr_log_open_failed", "gateway could not open Windsurf language server stderr log" ); - (Stdio::null(), None) + (Stdio::null(), false) } } } -fn repair_executable_mode(path: &Path) { +fn validate_language_server_binary_path( + path: &Path, +) -> Result { + let invalid = || { + ExecutionRuntimeTransportError::UpstreamRequest( + "Windsurf language server binary path is missing or unsafe".to_string(), + ) + }; + + let requested_metadata = fs::symlink_metadata(path).map_err(|_| invalid())?; + if requested_metadata.file_type().is_symlink() { + return Err(invalid()); + } + + let canonical = fs::canonicalize(path).map_err(|_| invalid())?; + let metadata = fs::symlink_metadata(&canonical).map_err(|_| invalid())?; + if !metadata.is_file() || metadata.file_type().is_symlink() { + return Err(invalid()); + } + #[cfg(unix)] { - use std::os::unix::fs::PermissionsExt; - if let Ok(metadata) = fs::metadata(path) { - let mode = metadata.permissions().mode(); - if mode & 0o111 == 0 { - let mut permissions = metadata.permissions(); - permissions.set_mode(mode | 0o111); - if let Err(err) = fs::set_permissions(path, permissions) { - warn!( - event_name = "windsurf_language_server_chmod_failed", - log_type = "ops", - path = %path.display(), - error = %err, - "gateway failed to repair Windsurf language server executable bit" - ); - } + use std::os::unix::fs::MetadataExt; + + let effective_uid = unsafe { libc::geteuid() }; + let trusted_owner = |uid| uid == 0 || (effective_uid != 0 && uid == effective_uid); + if requested_metadata.dev() != metadata.dev() + || requested_metadata.ino() != metadata.ino() + || metadata.nlink() != 1 + || !trusted_owner(metadata.uid()) + || metadata.mode() & 0o022 != 0 + || metadata.mode() & 0o6000 != 0 + || metadata.mode() & 0o111 == 0 + { + return Err(invalid()); + } + + let Some(parent) = canonical.parent() else { + return Err(invalid()); + }; + for directory in parent.ancestors() { + let directory_metadata = fs::symlink_metadata(directory).map_err(|_| invalid())?; + if !directory_metadata.is_dir() + || directory_metadata.file_type().is_symlink() + || !trusted_owner(directory_metadata.uid()) + || directory_metadata.mode() & 0o022 != 0 + { + return Err(invalid()); } } } + + Ok(canonical) } fn find_free_language_server_port() -> Result { @@ -3950,7 +4186,6 @@ fn port_is_free(port: u16) -> bool { async fn wait_language_server_ready( port: u16, child: &mut Child, - stderr_log_path: Option<&Path>, ) -> Result<(), ExecutionRuntimeTransportError> { let timeout = if cfg!(target_os = "windows") { WINDOWS_LS_READY_TIMEOUT @@ -3960,13 +4195,8 @@ async fn wait_language_server_ready( let started = Instant::now(); let addr = SocketAddr::from(([127, 0, 0, 1], port)); while started.elapsed() < timeout { - if let Ok(Some(status)) = child.try_wait() { - let stderr_tail = stderr_log_path - .and_then(|path| read_log_tail(path, 8 * 1024)) - .unwrap_or_else(|| "".to_string()); - return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf language server exited before port {port} became ready with status {status}; stderr tail: {stderr_tail}" - ))); + if let Ok(Some(_)) = child.try_wait() { + return Err(windsurf_language_server_exited_before_ready_error()); } if TcpStream::connect_timeout(&addr, Duration::from_millis(200)).is_ok() { debug!( @@ -3979,28 +4209,25 @@ async fn wait_language_server_ready( } tokio::time::sleep(Duration::from_millis(250)).await; } - let child_status = child.try_wait().ok().flatten(); - let stderr_tail = stderr_log_path - .and_then(|path| read_log_tail(path, 8 * 1024)) - .unwrap_or_else(|| "".to_string()); - Err(ExecutionRuntimeTransportError::UpstreamRequest(format!( - "Windsurf language server port {port} was not ready after {}ms{}; stderr tail: {stderr_tail}", - timeout.as_millis(), - child_status - .map(|status| format!(" (child status: {status})")) - .unwrap_or_default() - ))) + Err(windsurf_language_server_ready_timeout_error()) } -fn read_log_tail(path: &Path, max_bytes: usize) -> Option { - let mut file = fs::File::open(path).ok()?; - let len = file.metadata().ok()?.len() as i64; - let max_bytes = max_bytes.max(1) as i64; - let start = len.saturating_sub(max_bytes); - file.seek(SeekFrom::Start(start as u64)).ok()?; - let mut buf = Vec::with_capacity((len - start) as usize); - file.read_to_end(&mut buf).ok()?; - Some(String::from_utf8_lossy(&buf).to_string()) +fn windsurf_connect_http_error(method: &str, status_code: u16) -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest(format!( + "Windsurf Connect {method} returned HTTP {status_code}" + )) +} + +fn windsurf_language_server_exited_before_ready_error() -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest( + "Windsurf language server exited before becoming ready".to_string(), + ) +} + +fn windsurf_language_server_ready_timeout_error() -> ExecutionRuntimeTransportError { + ExecutionRuntimeTransportError::UpstreamRequest( + "Windsurf language server was not ready before timeout".to_string(), + ) } fn language_server_data_dir(key: &str) -> PathBuf { @@ -4016,32 +4243,164 @@ fn language_server_data_dir(key: &str) -> PathBuf { .join(key) } -fn language_server_workspace_path(plan: &ExecutionPlan) -> PathBuf { - std::env::temp_dir() - .join("aether-windsurf") - .join(format!("workspace-{}", hash_hex(&plan.key_id, 16))) +fn absolute_runtime_path(path: PathBuf) -> io::Result { + if path.is_absolute() { + Ok(path) + } else { + Ok(std::env::current_dir()?.join(path)) + } } -fn ensure_workspace_dir(path: &Path) { - if let Err(err) = fs::create_dir_all(path) { - warn!( - event_name = "windsurf_workspace_create_failed", - log_type = "ops", - path = %path.display(), - error = %err, - "gateway failed to create Windsurf placeholder workspace" - ); - return; +fn ensure_private_directory(path: &Path) -> io::Result<()> { + validate_runtime_path_components(path)?; + fs::create_dir_all(path)?; + validate_runtime_path_components(path)?; + + let metadata = fs::symlink_metadata(path)?; + if !metadata.is_dir() || metadata.file_type().is_symlink() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "private runtime path is not a regular directory", + )); } - let _ = fs::write( - path.join("package.json"), - "{\n \"name\": \"aether-windsurf-workspace-stub\",\n \"private\": true,\n \"version\": \"0.0.0\"\n}\n", - ); - let _ = fs::write( - path.join("README.md"), - "# Aether Windsurf workspace placeholder\n\nThis directory is only registered so the Windsurf language server has a trusted workspace.\n", - ); - let _ = fs::write(path.join(".gitignore"), "# placeholder\n"); + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + + let effective_uid = unsafe { libc::geteuid() }; + if metadata.uid() != effective_uid { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime directory is owned by another user", + )); + } + fs::set_permissions(path, fs::Permissions::from_mode(0o700))?; + } + Ok(()) +} + +fn validate_runtime_path_components(path: &Path) -> io::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + let absolute = if path.is_absolute() { + path.to_path_buf() + } else { + std::env::current_dir()?.join(path) + }; + let effective_uid = unsafe { libc::geteuid() }; + for component in absolute.ancestors().collect::>().into_iter().rev() { + let metadata = match fs::symlink_metadata(component) { + Ok(metadata) => metadata, + Err(error) if error.kind() == io::ErrorKind::NotFound => continue, + Err(error) => return Err(error), + }; + if metadata.file_type().is_symlink() { + if component == absolute || (metadata.uid() != 0 && metadata.uid() != effective_uid) + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime path contains an untrusted symbolic link", + )); + } + validate_runtime_path_components(&fs::canonicalize(component)?)?; + continue; + } + if !metadata.is_dir() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "private runtime path contains a non-directory component", + )); + } + let mode = metadata.mode(); + if (metadata.uid() != 0 && metadata.uid() != effective_uid) + || (mode & 0o022 != 0 && mode & 0o1000 == 0) + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime path contains an unsafe writable directory", + )); + } + } + } + + #[cfg(not(unix))] + if let Ok(metadata) = fs::symlink_metadata(path) { + if metadata.file_type().is_symlink() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime path must not be a symbolic link", + )); + } + } + Ok(()) +} + +fn open_private_file(path: &Path, append: bool) -> io::Result { + if let Ok(metadata) = fs::symlink_metadata(path) { + if metadata.file_type().is_symlink() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime file must not be a symbolic link", + )); + } + } + + let mut options = fs::OpenOptions::new(); + options.write(true).create(true).append(append); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options + .mode(0o600) + .custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); + } + let file = options.open(path)?; + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "private runtime file is not a regular file", + )); + } + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + + if metadata.uid() != unsafe { libc::geteuid() } || metadata.nlink() != 1 { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "private runtime file has unsafe ownership or hard links", + )); + } + file.set_permissions(fs::Permissions::from_mode(0o600))?; + } + Ok(file) +} + +fn open_private_append_file(path: &Path) -> io::Result { + open_private_file(path, true) +} + +fn write_private_file(path: &Path, content: &[u8]) -> io::Result<()> { + let mut file = open_private_file(path, false)?; + file.set_len(0)?; + file.write_all(content)?; + file.sync_all() +} + +fn ensure_workspace_dir(path: &Path) -> io::Result<()> { + ensure_private_directory(path)?; + write_private_file( + &path.join("package.json"), + b"{\n \"name\": \"aether-windsurf-workspace-stub\",\n \"private\": true,\n \"version\": \"0.0.0\"\n}\n", + )?; + write_private_file( + &path.join("README.md"), + b"# Aether Windsurf workspace placeholder\n\nThis directory is only registered so the Windsurf language server has a trusted workspace.\n", + )?; + write_private_file(&path.join(".gitignore"), b"# placeholder\n") } fn language_server_env(proxy_url: Option<&str>) -> BTreeMap { @@ -4183,10 +4542,197 @@ mod tests { use super::ExecutionRuntimeTransportError; use super::{ build_openai_chat_sse_body, detect_windsurf_request, emit_windsurf_step_text_deltas, - is_windsurf_cascade_transport_error, is_windsurf_send_retryable_error, WindsurfToolCall, - WindsurfToolDefinition, + is_windsurf_cascade_transport_error, is_windsurf_send_retryable_error, + new_windsurf_csrf_token, windsurf_connect_http_error, + windsurf_language_server_exited_before_ready_error, + windsurf_language_server_ready_timeout_error, WindsurfToolCall, WindsurfToolDefinition, }; + #[test] + fn windsurf_language_server_csrf_tokens_are_per_process_secrets() { + let first = new_windsurf_csrf_token(); + let second = new_windsurf_csrf_token(); + + assert!(first.starts_with("aether-")); + assert_eq!(first.len(), "aether-".len() + 64); + assert_ne!(first, second); + assert_ne!(first, "windsurf-api-csrf-fixed-token"); + } + + #[cfg(unix)] + #[test] + fn windsurf_language_server_binary_must_be_trusted_and_already_executable() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let root = std::env::current_dir() + .expect("current directory") + .join(format!( + ".aether-windsurf-binary-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&root).expect("test directory should be created"); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700)) + .expect("test directory mode should be private"); + + let binary = root.join("language-server"); + std::fs::write(&binary, b"test binary").expect("binary fixture should be written"); + std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o700)) + .expect("binary fixture should be executable"); + assert_eq!( + super::validate_language_server_binary_path(&binary) + .expect("trusted executable should validate"), + std::fs::canonicalize(&binary).expect("binary fixture should canonicalize") + ); + + let linked = root.join("language-server-link"); + symlink(&binary, &linked).expect("binary symlink should be created"); + assert!(super::validate_language_server_binary_path(&linked).is_err()); + + let hard_linked = root.join("language-server-hard-link"); + std::fs::hard_link(&binary, &hard_linked).expect("binary hard link should be created"); + assert!(super::validate_language_server_binary_path(&binary).is_err()); + assert!(super::validate_language_server_binary_path(&hard_linked).is_err()); + std::fs::remove_file(&hard_linked).expect("binary hard link should be removed"); + + assert!(super::validate_language_server_binary_path(&root).is_err()); + + std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o720)) + .expect("binary fixture should become group writable"); + assert!(super::validate_language_server_binary_path(&binary).is_err()); + + std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o600)) + .expect("binary fixture should become non-executable"); + assert!(super::validate_language_server_binary_path(&binary).is_err()); + assert_eq!( + std::fs::metadata(&binary) + .expect("binary fixture metadata") + .permissions() + .mode() + & 0o777, + 0o600, + "validation must not grant execute permission" + ); + + let unsafe_parent = root.join("group-writable"); + std::fs::create_dir(&unsafe_parent).expect("unsafe parent fixture should be created"); + std::fs::set_permissions(&unsafe_parent, std::fs::Permissions::from_mode(0o770)) + .expect("unsafe parent should be group writable"); + let nested_binary = unsafe_parent.join("language-server"); + std::fs::write(&nested_binary, b"test binary") + .expect("nested binary fixture should be written"); + std::fs::set_permissions(&nested_binary, std::fs::Permissions::from_mode(0o700)) + .expect("nested binary fixture should be executable"); + assert!(super::validate_language_server_binary_path(&nested_binary).is_err()); + + std::fs::remove_dir_all(root).expect("test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn windsurf_runtime_files_are_private_and_reject_links() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-windsurf-private-test-{}", + uuid::Uuid::new_v4() + )); + super::ensure_private_directory(&root).expect("private root should be created"); + let workspace = root.join("workspace"); + super::ensure_workspace_dir(&workspace).expect("private workspace should be created"); + + assert_eq!( + std::fs::metadata(&root) + .expect("private root metadata") + .permissions() + .mode() + & 0o777, + 0o700 + ); + assert_eq!( + std::fs::metadata(workspace.join("package.json")) + .expect("private workspace file metadata") + .permissions() + .mode() + & 0o777, + 0o600 + ); + + let log_path = root.join("language-server.stderr.log"); + drop(super::open_private_append_file(&log_path).expect("private log should open")); + assert_eq!( + std::fs::metadata(&log_path) + .expect("private log metadata") + .permissions() + .mode() + & 0o777, + 0o600 + ); + + let target = root.join("target.log"); + super::write_private_file(&target, b"target").expect("target fixture should be written"); + let link = root.join("linked.log"); + symlink(&target, &link).expect("symlink fixture should be created"); + assert!(super::open_private_append_file(&link).is_err()); + + let hard_link = root.join("hard-linked.log"); + std::fs::hard_link(&target, &hard_link).expect("hard-link fixture should be created"); + assert!(super::open_private_append_file(&target).is_err()); + + std::fs::remove_dir_all(root).expect("private runtime fixture should be removed"); + } + + #[test] + fn windsurf_delta_buffer_rejects_overflow_without_mutation() { + let mut content = "1234".to_string(); + + let error = super::append_windsurf_delta_with_limit(&mut content, "56", 5) + .expect_err("decoded content above the limit must fail"); + + assert_eq!(content, "1234"); + assert!(matches!( + error, + ExecutionRuntimeTransportError::UpstreamResponseTooLarge { limit_bytes: 5, .. } + )); + } + + #[test] + fn windsurf_auxiliary_errors_omit_response_and_local_diagnostics() { + let sensitive_detail = "Bearer secret-windsurf-error-body stderr tail /private/path"; + let connect_message = windsurf_connect_http_error("GetUserStatus", 502).to_string(); + + assert!(connect_message.contains("Windsurf Connect GetUserStatus returned HTTP 502")); + assert!(!connect_message.contains(sensitive_detail)); + + let exited_message = windsurf_language_server_exited_before_ready_error().to_string(); + let timeout_message = windsurf_language_server_ready_timeout_error().to_string(); + assert_eq!( + exited_message, + "failed to execute upstream request: Windsurf language server exited before becoming ready" + ); + assert_eq!( + timeout_message, + "failed to execute upstream request: Windsurf language server was not ready before timeout" + ); + assert!(!exited_message.contains(sensitive_detail)); + assert!(!timeout_message.contains(sensitive_detail)); + + let raw_error = ExecutionRuntimeTransportError::UpstreamRequest(format!( + "request failed: {sensitive_detail}" + )); + let public_error = super::windsurf_execution_error_from_transport_error( + &raw_error, + ExecutionPhase::StreamRead, + ); + let trajectory_error = + super::windsurf_trajectory_error(&format!("rate limit: {sensitive_detail}")); + assert_eq!(public_error.message, "Windsurf request failed"); + assert!(trajectory_error.to_string().contains("rate limit")); + assert!(!public_error.message.contains("secret-windsurf-error-body")); + assert!(!trajectory_error + .to_string() + .contains("secret-windsurf-error-body")); + } + fn windsurf_plan() -> ExecutionPlan { ExecutionPlan { request_id: "req-windsurf".to_string(), @@ -4683,6 +5229,10 @@ mod tests { native_bridge: None, }; + let input_debug = format!("{input:?}"); + assert!(input_debug.contains("[REDACTED]")); + assert!(!input_debug.contains("windsurf-api-key")); + let parsed = super::parse_and_filter_windsurf_tool_calls( r#"I will use WebSearch(query="today tech", domain="example.com")."#, &input, diff --git a/apps/aether-gateway/src/executor/candidate_loop.rs b/apps/aether-gateway/src/executor/candidate_loop.rs index 5b0e5409f..8806e66e3 100644 --- a/apps/aether-gateway/src/executor/candidate_loop.rs +++ b/apps/aether-gateway/src/executor/candidate_loop.rs @@ -1,11 +1,12 @@ use std::collections::{BTreeMap, BTreeSet}; +use std::sync::atomic::{AtomicBool, Ordering}; use aether_ai_serving::{ run_ai_attempt_loop, AiAttemptExecutionOutcome, AiAttemptLoopOutcome, AiAttemptLoopPort, AiAttemptRetryScope, AiExecutionAttempt, }; use aether_data_contracts::repository::candidates::RequestCandidateStatus; -use aether_runtime::ConcurrencyPermit; +use aether_runtime::{AdmissionPermit, ConcurrencyPermit}; use aether_scheduler_core::{ parse_request_candidate_report_context, SchedulerRequestCandidateStatusUpdate, }; @@ -27,7 +28,8 @@ use crate::execution_runtime::{ UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME, }; use crate::executor::{ - build_local_execution_exhaustion, mark_deferred_upstream_response, LocalExecutionRequestOutcome, + attach_deferred_usage_context, build_local_execution_exhaustion, + mark_deferred_upstream_response, LocalExecutionRequestOutcome, }; use crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease; use crate::log_ids::short_request_id; @@ -51,6 +53,17 @@ const UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE_ENV: &str = const UPSTREAM_EXECUTION_GATE_STREAM_HOLD_MODE_ENV: &str = "AETHER_GATEWAY_UPSTREAM_EXECUTION_GATE_STREAM_HOLD_MODE"; +#[derive(Clone, Debug)] +pub(crate) struct BackgroundAdmissionPermit { + _permit: AdmissionPermit, +} + +impl BackgroundAdmissionPermit { + pub(crate) fn new(permit: AdmissionPermit) -> Self { + Self { _permit: permit } + } +} + fn attach_redaction_execution_candidate(response: &mut Response, candidate_id: Option<&str>) { if let Some(candidate_id) = candidate_id .map(str::trim) @@ -73,7 +86,7 @@ pub(crate) async fn execute_sync_plan_and_reports( where T: AiExecutionAttempt + Send + Sync + 'static, { - let transfer_tracker = ProviderTransferTracker::default(); + let transfer_tracker = ProviderTransferTracker::for_request(parts); execute_sync_plan_and_reports_with_transfer_tracker( state, parts, @@ -130,7 +143,17 @@ where plan_kind, transfer_tracker, }; - match run_ai_attempt_loop(&port, plan_and_reports).await? { + let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await; + if loop_result.is_err() { + release_active_plan_usage_policy_cost_best_effort( + state, + decision, + transfer_tracker, + "candidate_loop_error", + ) + .await; + } + match loop_result? { AiAttemptLoopOutcome::Responded(response) => { Ok(LocalExecutionRequestOutcome::responded(response)) } @@ -159,7 +182,7 @@ where T: AiExecutionAttempt + Send + Sync + 'static, S: LocalExecutionAttemptSource, { - let transfer_tracker = ProviderTransferTracker::default(); + let transfer_tracker = ProviderTransferTracker::for_request(parts); execute_sync_attempt_source_with_transfer_tracker( state, parts, @@ -204,7 +227,7 @@ where plan_kind, transfer_tracker, }; - run_dynamic_attempt_loop( + let loop_result = run_dynamic_attempt_loop( &port, &mut source, trace_id, @@ -213,7 +236,17 @@ where .frontdoor_runtime_guards .local_execution_planning_timeout, ) - .await + .await; + if loop_result.is_err() { + release_active_plan_usage_policy_cost_best_effort( + state, + decision, + transfer_tracker, + "dynamic_candidate_loop_error", + ) + .await; + } + loop_result } .instrument(span) .await @@ -252,6 +285,10 @@ where Ok(()) } + async fn next_same_key_retry(&self, attempt: &T) -> Result, Self::Error> { + Ok(crate::orchestration::next_same_key_retry_attempt(attempt)) + } + async fn record_attempt_failed(&self, attempt: &T) -> Result<(), Self::Error> { record_provider_transfer_attempt_failed( self.state, @@ -269,22 +306,73 @@ where attempt: &T, ) -> Result, Self::Error> { let plan = attempt.execution_plan(); - let report_context = attempt.report_context(); - if let Some(response) = execution_plan_balance_capacity_response( + let report_context = attach_plan_usage_reservation_token( + attempt.report_context(), + self.transfer_tracker.usage_policy_reservation_token(), + ); + let balance_response = execution_plan_balance_capacity_response( self.state, self.trace_id, self.decision, plan, report_context.as_ref(), ) + .await; + let balance_response = match balance_response { + Ok(response) => response, + Err(error) => { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "balance_capacity_error", + ) + .await; + return Err(error); + } + }; + if let Some(response) = balance_response { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "balance_capacity_rejected", + ) + .await; + return Ok(AiAttemptExecutionOutcome::Responded(response)); + } + prewarm_direct_reqwest_candidate_client(plan); + let _permit = match acquire_upstream_execution_gate(self.state, self.trace_id).await { + Ok(permit) => permit, + Err(error) => { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "upstream_admission_error", + ) + .await; + return Err(error); + } + }; + if let Some(response) = execution_plan_cost_capacity_response( + self.state, + self.trace_id, + self.decision, + plan, + report_context.as_ref(), + self.transfer_tracker, + ) .await? { return Ok(AiAttemptExecutionOutcome::Responded(response)); } - prewarm_direct_reqwest_candidate_client(plan); - let _permit = acquire_upstream_execution_gate(self.state, self.trace_id).await?; let upstream_execution_gate_held_started_at = std::time::Instant::now(); - let mut execution = execute_execution_runtime_sync_with_retry_scope( + let deferred_report_context = report_context.clone(); + let execution = execute_execution_runtime_sync_with_retry_scope( self.state, self.parts.uri.path(), plan.clone(), @@ -294,19 +382,39 @@ where attempt.report_kind(), report_context, ) - .await?; + .await; + let execution = match execution { + Ok(execution) => execution, + Err(error) => { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "sync_execution_error", + ) + .await; + return Err(error); + } + }; observe_gateway_stage_ms( "upstream_execution_gate_held", upstream_execution_gate_held_started_at .elapsed() .as_millis() as u64, ); + let mut execution = execution; match &mut execution { - AiAttemptExecutionOutcome::Responded(response) - | AiAttemptExecutionOutcome::Retry { + AiAttemptExecutionOutcome::Responded(response) => { + attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()); + } + AiAttemptExecutionOutcome::Retry { fallback_response: Some(response), .. - } => attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()), + } => { + attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()); + attach_deferred_usage_context(response, plan, deferred_report_context.as_ref()); + } AiAttemptExecutionOutcome::Retry { fallback_response: None, .. @@ -338,10 +446,16 @@ where model_name = last_plan.model_name.as_deref().unwrap_or("-"), "candidate loop exhausted local sync candidates" ); - Ok( - build_local_execution_exhaustion(self.state, &last_plan, last_report_context.as_ref()) - .await, + Ok(build_local_execution_exhaustion( + self.state, + &last_plan, + attach_plan_usage_reservation_token( + last_report_context, + self.transfer_tracker.usage_policy_reservation_token(), + ) + .as_ref(), ) + .await) } } @@ -408,8 +522,19 @@ where decision, plan_kind, transfer_tracker, + request_first_byte_started_at: Instant::now(), }; - match run_ai_attempt_loop(&port, plan_and_reports).await? { + let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await; + if loop_result.is_err() { + release_active_plan_usage_policy_cost_best_effort( + state, + decision, + transfer_tracker, + "candidate_loop_error", + ) + .await; + } + match loop_result? { AiAttemptLoopOutcome::Responded(response) => { Ok(LocalExecutionRequestOutcome::responded(response)) } @@ -478,8 +603,9 @@ where decision, plan_kind, transfer_tracker, + request_first_byte_started_at: Instant::now(), }; - run_dynamic_attempt_loop( + let loop_result = run_dynamic_attempt_loop( &port, &mut source, trace_id, @@ -488,7 +614,17 @@ where .frontdoor_runtime_guards .local_execution_planning_timeout, ) - .await + .await; + if loop_result.is_err() { + release_active_plan_usage_policy_cost_best_effort( + state, + decision, + transfer_tracker, + "dynamic_candidate_loop_error", + ) + .await; + } + loop_result } .instrument(span) .await @@ -526,6 +662,49 @@ struct ProviderTransferStateTracker { #[derive(Clone, Debug, Default)] pub(crate) struct ProviderTransferTracker { state: std::sync::Arc>, + usage_policy_reservation: Option, + usage_policy_reservation_plan: + std::sync::Arc>>, + usage_policy_cost_reserved: std::sync::Arc, + _background_admission_permit: Option, +} + +impl ProviderTransferTracker { + pub(crate) fn for_request(parts: &http::request::Parts) -> Self { + let background_admission_permit = + parts.extensions.get::().cloned(); + let usage_policy_reservation = parts + .extensions + .get::() + .cloned(); + Self { + state: Default::default(), + usage_policy_reservation, + usage_policy_reservation_plan: Default::default(), + usage_policy_cost_reserved: Default::default(), + _background_admission_permit: background_admission_permit, + } + } + + fn usage_policy_reservation_token(&self) -> Option<&str> { + self.usage_policy_reservation + .as_ref() + .map(crate::plan_usage_policy::PlanUsageReservationContext::token) + } + + fn record_usage_policy_reservation_plan(&self, plan: &aether_contracts::ExecutionPlan) { + *self + .usage_policy_reservation_plan + .lock() + .expect("usage policy reservation plan lock") = Some(plan.clone()); + } + + fn usage_policy_reservation_plan(&self) -> Option { + self.usage_policy_reservation_plan + .lock() + .expect("usage policy reservation plan lock") + .clone() + } } #[derive(Debug, Clone, PartialEq, Eq)] @@ -788,18 +967,31 @@ where { let mut last_attempted = None; let mut fallback_response = None; + // A same-key retry derived after a candidate-scoped failure runs before + // the source is asked for the next candidate. + let mut pending_same_key_retry: Option = None; loop { - let next_started_at = std::time::Instant::now(); - let next_attempt = - next_execution_attempt_with_timeout(source, trace_id, plan_kind, planning_timeout) + let attempt = match pending_same_key_retry.take() { + Some(attempt) => attempt, + None => { + let next_started_at = std::time::Instant::now(); + let next_attempt = next_execution_attempt_with_timeout( + source, + trace_id, + plan_kind, + planning_timeout, + ) .await?; - observe_gateway_stage_ms( - "stream_candidate_next", - next_started_at.elapsed().as_millis() as u64, - ); - let Some(attempt) = next_attempt else { - break; + observe_gateway_stage_ms( + "stream_candidate_next", + next_started_at.elapsed().as_millis() as u64, + ); + let Some(attempt) = next_attempt else { + break; + }; + attempt + } }; if port.should_skip_attempt(&attempt).await? { let provider_id = attempt.execution_plan().provider_id.clone(); @@ -839,6 +1031,9 @@ where if attempt_fallback_response.is_some() { fallback_response = attempt_fallback_response; } + if scope == AiAttemptRetryScope::Candidate { + pending_same_key_retry = port.next_same_key_retry(&attempt).await?; + } apply_attempt_retry_scope(source, &attempt, scope).await?; } } @@ -926,6 +1121,10 @@ struct StreamAttemptLoopPort<'a> { decision: &'a GatewayControlDecision, plan_kind: &'a str, transfer_tracker: &'a ProviderTransferTracker, + /// All candidates in one downstream stream request share this origin. + /// Without it every retry receives a fresh full first-byte timeout and a + /// 30-second provider timeout can accumulate into a 60-120 second stall. + request_first_byte_started_at: Instant, } #[async_trait] @@ -952,6 +1151,10 @@ where Ok(()) } + async fn next_same_key_retry(&self, attempt: &T) -> Result, Self::Error> { + Ok(crate::orchestration::next_same_key_retry_attempt(attempt)) + } + async fn record_attempt_failed(&self, attempt: &T) -> Result<(), Self::Error> { record_provider_transfer_attempt_failed( self.state, @@ -969,7 +1172,10 @@ where attempt: &T, ) -> Result, Self::Error> { let plan = attempt.execution_plan(); - let report_context = attempt.report_context(); + let report_context = attach_plan_usage_reservation_token( + attempt.report_context(), + self.transfer_tracker.usage_policy_reservation_token(), + ); let candidate_index = parse_request_candidate_report_context(report_context.as_ref()) .and_then(|context| context.candidate_index) .map(|value| value.to_string()) @@ -988,35 +1194,49 @@ where candidate_index = candidate_index.as_str(), "candidate loop attempting stream execution candidate" ); - if let Some(response) = execution_plan_balance_capacity_response( + let balance_response = execution_plan_balance_capacity_response( self.state, self.trace_id, self.decision, plan, report_context.as_ref(), ) - .await? - { + .await; + let balance_response = match balance_response { + Ok(response) => response, + Err(error) => { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "balance_capacity_error", + ) + .await; + return Err(error); + } + }; + if let Some(response) = balance_response { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "balance_capacity_rejected", + ) + .await; return Ok(AiAttemptExecutionOutcome::Responded(response)); } prewarm_direct_reqwest_candidate_client(plan); - // The attempt owns the canonical report context. Borrow it for the - // watchdog; only third-party/synthesized attempts using the default - // trait implementation need an owned fallback clone. - let watchdog_report_context_owned = if attempt.report_context_ref().is_none() { - report_context.clone() - } else { - None - }; - let watchdog_report_context = attempt - .report_context_ref() - .or(watchdog_report_context_owned.as_ref()); + let watchdog_report_context_owned = report_context.clone(); + let watchdog_report_context = watchdog_report_context_owned.as_ref(); let execution_state = self.state.clone(); let execution_trace_id = self.trace_id.to_string(); let execution_plan_kind = self.plan_kind.to_string(); let execution_decision = self.decision.clone(); let execution_report_kind = attempt.report_kind(); let execution_plan = plan.clone(); + let execution_transfer_tracker = self.transfer_tracker.clone(); let stop_on_transport_errors = matches!( resolve_local_transport_failover_analysis_for_attempt( self.state, @@ -1034,8 +1254,21 @@ where self.plan_kind, plan, watchdog_report_context, + self.request_first_byte_started_at, stop_on_transport_errors, move || async move { + if let Some(response) = execution_plan_cost_capacity_response( + &execution_state, + execution_trace_id.as_str(), + &execution_decision, + &execution_plan, + report_context.as_ref(), + &execution_transfer_tracker, + ) + .await? + { + return Ok(AiAttemptExecutionOutcome::Responded(response)); + } execute_execution_runtime_stream_with_retry_scope( &execution_state, execution_plan, @@ -1048,7 +1281,21 @@ where .await }, ) - .await?; + .await; + let execution = match execution { + Ok(execution) => execution, + Err(error) => { + release_plan_usage_policy_cost_best_effort( + self.state, + self.decision, + plan, + self.transfer_tracker, + "stream_execution_error", + ) + .await; + return Err(error); + } + }; let mut execution = match execution { StreamCandidateWatchdogOutcome::TransportTimeout => { AiAttemptExecutionOutcome::Responded( @@ -1061,7 +1308,7 @@ where http::StatusCode::GATEWAY_TIMEOUT.as_u16(), "local_stream_candidate_watchdog_timeout", stream_candidate_watchdog_timeout_message(), - watchdog_started_at.elapsed().as_millis() as u64, + self.request_first_byte_started_at.elapsed().as_millis() as u64, ) .await?, ) @@ -1069,11 +1316,16 @@ where StreamCandidateWatchdogOutcome::Executed(execution) => execution, }; match &mut execution { - AiAttemptExecutionOutcome::Responded(response) - | AiAttemptExecutionOutcome::Retry { + AiAttemptExecutionOutcome::Responded(response) => { + attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()); + } + AiAttemptExecutionOutcome::Retry { fallback_response: Some(response), .. - } => attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()), + } => { + attach_redaction_execution_candidate(response, plan.candidate_id.as_deref()); + attach_deferred_usage_context(response, plan, watchdog_report_context); + } AiAttemptExecutionOutcome::Retry { fallback_response: None, .. @@ -1105,10 +1357,16 @@ where model_name = last_plan.model_name.as_deref().unwrap_or("-"), "candidate loop exhausted local stream candidates" ); - Ok( - build_local_execution_exhaustion(self.state, &last_plan, last_report_context.as_ref()) - .await, + Ok(build_local_execution_exhaustion( + self.state, + &last_plan, + attach_plan_usage_reservation_token( + last_report_context, + self.transfer_tracker.usage_policy_reservation_token(), + ) + .as_ref(), ) + .await) } } @@ -1121,6 +1379,47 @@ fn prewarm_direct_reqwest_candidate_client(plan: &aether_contracts::ExecutionPla ); } +fn attach_plan_usage_reservation_token( + report_context: Option, + reservation_token: Option<&str>, +) -> Option { + match (report_context, reservation_token) { + (Some(serde_json::Value::Object(mut context)), Some(reservation_token)) => { + context.insert( + "plan_usage_reservation_token".to_string(), + serde_json::Value::String(reservation_token.to_string()), + ); + context.insert( + aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY + .to_string(), + serde_json::Value::Bool(false), + ); + Some(serde_json::Value::Object(context)) + } + (Some(serde_json::Value::Object(mut context)), None) => { + context.remove("plan_usage_reservation_token"); + context.remove( + aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, + ); + Some(serde_json::Value::Object(context)) + } + (_, Some(reservation_token)) => { + let mut context = serde_json::Map::new(); + context.insert( + "plan_usage_reservation_token".to_string(), + serde_json::Value::String(reservation_token.to_string()), + ); + context.insert( + aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY + .to_string(), + serde_json::Value::Bool(false), + ); + Some(serde_json::Value::Object(context)) + } + (report_context, None) => report_context, + } +} + async fn execution_plan_balance_capacity_response( state: &AppState, trace_id: &str, @@ -1155,6 +1454,121 @@ async fn execution_plan_balance_capacity_response( Ok(Some(response)) } +async fn execution_plan_cost_capacity_response( + state: &AppState, + trace_id: &str, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + report_context: Option<&serde_json::Value>, + transfer_tracker: &ProviderTransferTracker, +) -> Result>, GatewayError> { + let outcome = match crate::plan_usage_policy::reserve_admitted_http_plan_usage_policy_cost( + state, + decision, + plan, + report_context, + transfer_tracker.usage_policy_reservation.as_ref(), + ) + .await + { + Ok(outcome) => outcome, + Err(err) => { + release_plan_usage_policy_cost_best_effort( + state, + decision, + plan, + transfer_tracker, + "reservation_error", + ) + .await; + mark_unused_local_candidate(state, plan, report_context).await; + return Err(err); + } + }; + let rejection = match outcome { + crate::plan_usage_policy::PlanUsageCostReservationOutcome::NotRequired => return Ok(None), + crate::plan_usage_policy::PlanUsageCostReservationOutcome::Reserved => { + transfer_tracker.record_usage_policy_reservation_plan(plan); + transfer_tracker + .usage_policy_cost_reserved + .store(true, Ordering::Release); + return Ok(None); + } + crate::plan_usage_policy::PlanUsageCostReservationOutcome::Rejected(rejection) => rejection, + }; + release_plan_usage_policy_cost_best_effort( + state, + decision, + plan, + transfer_tracker, + "reservation_rejected", + ) + .await; + mark_unused_local_candidate(state, plan, report_context).await; + let mut response = crate::api::response::build_local_plan_usage_limited_response( + trace_id, + Some(decision), + &rejection, + )?; + attach_redaction_execution_candidate(&mut response, plan.candidate_id.as_deref()); + Ok(Some(response)) +} + +async fn release_plan_usage_policy_cost_best_effort( + state: &AppState, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + transfer_tracker: &ProviderTransferTracker, + reason: &'static str, +) { + if !transfer_tracker + .usage_policy_cost_reserved + .swap(false, Ordering::AcqRel) + { + return; + } + if let Err(error) = crate::plan_usage_policy::release_plan_usage_policy_cost( + state, + decision, + plan, + transfer_tracker + .usage_policy_reservation_token() + .expect("a reserved plan cost always has a reservation token"), + current_unix_ms() / 1_000, + ) + .await + { + warn!( + event_name = "plan_usage_cost_reservation_release_failed", + log_type = "ops", + request_id = %plan.request_id, + candidate_id = ?plan.candidate_id, + reservation_token = transfer_tracker + .usage_policy_reservation_token() + .unwrap_or("-"), + reason, + error = ?error, + "gateway failed to release plan usage cost reservation after terminal capacity failure" + ); + transfer_tracker + .usage_policy_cost_reserved + .store(true, Ordering::Release); + } +} + +async fn release_active_plan_usage_policy_cost_best_effort( + state: &AppState, + decision: &GatewayControlDecision, + transfer_tracker: &ProviderTransferTracker, + reason: &'static str, +) { + let Some(plan) = transfer_tracker.usage_policy_reservation_plan() else { + return; + }; + release_plan_usage_policy_cost_best_effort(state, decision, &plan, transfer_tracker, reason) + .await; +} + pub(crate) async fn mark_unused_local_candidates(state: &AppState, remaining: Vec) where T: AiExecutionAttempt, @@ -1344,6 +1758,7 @@ async fn execute_stream_candidate_with_watchdog( plan_kind: &str, plan: &aether_contracts::ExecutionPlan, report_context: Option<&serde_json::Value>, + request_first_byte_started_at: Instant, stop_on_transport_errors: bool, execute: impl FnOnce() -> Fut, ) -> Result @@ -1353,6 +1768,7 @@ where > + Send, { let timeout_duration = resolve_stream_candidate_watchdog_timeout(plan, report_context); + let request_first_byte_deadline = request_first_byte_started_at + timeout_duration; let candidate_started_at = std::time::Instant::now(); let candidate_started_unix_ms = current_unix_ms(); let permit = match acquire_upstream_execution_gate(state, trace_id).await { @@ -1378,7 +1794,14 @@ where let watchdog_progress = StreamCandidateWatchdogProgress::shared(); let execution = watchdog_progress.clone().scope(execute()); tokio::pin!(execution); - let deadline = tokio::time::sleep(timeout_duration); + // This is an absolute request-level deadline, not a new timeout for this + // candidate. Retries therefore consume only the budget left by earlier + // candidates instead of resetting the full provider timeout. + let candidate_budget_ms = request_first_byte_deadline + .saturating_duration_since(Instant::now()) + .as_millis() + .min(u128::from(u64::MAX)) as u64; + let deadline = tokio::time::sleep_until(request_first_byte_deadline); tokio::pin!(deadline); let execution_result = tokio::select! { biased; @@ -1394,6 +1817,10 @@ where let outcome = match execution_result { Some(result) => result.map(StreamCandidateWatchdogOutcome::Executed), None => { + // The abandoned attempt is dropped when this function returns. + // Claim its settlement before that so its cancellation guard does + // not race the watchdog rows written just below. + watchdog_progress.mark_abandoned(); let finished_at_unix_ms = current_unix_ms(); let request_id = short_request_id(plan.request_id.as_str()); let provider_name = plan.provider_name.as_deref().unwrap_or("-"); @@ -1403,6 +1830,10 @@ where .map(|value| value.to_string()) .unwrap_or_else(|| "-".to_string()); let timeout_ms = u64::try_from(timeout_duration.as_millis()).unwrap_or(u64::MAX); + let request_elapsed_ms = request_first_byte_started_at + .elapsed() + .as_millis() + .min(u128::from(u64::MAX)) as u64; record_local_request_candidate_status( state, plan, @@ -1431,6 +1862,8 @@ where model_name, candidate_index = candidate_index.as_str(), timeout_ms, + candidate_budget_ms, + request_elapsed_ms, "gateway local stream candidate watchdog timed out" ); if stop_on_transport_errors { @@ -1649,10 +2082,20 @@ pub(crate) async fn mark_unused_local_candidate_items( mod tests { use std::sync::{Arc, Mutex as StdMutex}; - use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody}; + use aether_contracts::{ + ExecutionPlan, ExecutionResult, ExecutionTimeouts, RequestBody, ResponseBody, + }; + use aether_data::repository::settlement::{ + InMemorySettlementRepository, ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, + SettlementWriteRepository, UsagePolicyCostReservationState, UsagePolicyCostWindow, + }; + use aether_data_contracts::repository::billing::{ + BillingReadRepository, StoredBillingModelContext, UserPlanEntitlementRecord, + }; use aether_data_contracts::repository::candidates::{ RequestCandidateStatus, UpsertRequestCandidateRecord, }; + use aether_data_contracts::DataLayerError; use async_trait::async_trait; use serde_json::json; use tokio::sync::Mutex; @@ -1742,6 +2185,83 @@ mod tests { } } + #[derive(Debug)] + struct CostReservationBillingRepository { + model_context: StoredBillingModelContext, + entitlement: UserPlanEntitlementRecord, + } + + #[async_trait] + impl BillingReadRepository for CostReservationBillingRepository { + async fn find_model_context( + &self, + _provider_id: &str, + _provider_api_key_id: Option<&str>, + _global_model_name: &str, + ) -> Result, DataLayerError> { + Ok(Some(self.model_context.clone())) + } + + async fn find_model_context_by_model_id( + &self, + _provider_id: &str, + _provider_api_key_id: Option<&str>, + _model_id: &str, + ) -> Result, DataLayerError> { + Ok(Some(self.model_context.clone())) + } + + async fn list_user_plan_entitlements( + &self, + user_id: &str, + ) -> Result>, DataLayerError> { + Ok(Some( + (self.entitlement.user_id == user_id) + .then(|| self.entitlement.clone()) + .into_iter() + .collect(), + )) + } + } + + struct PlanningErrorAfterFirstSyncAttemptSource { + first: Option, + } + + #[async_trait] + impl LocalExecutionAttemptSource + for PlanningErrorAfterFirstSyncAttemptSource + { + async fn next_execution_attempt( + &mut self, + ) -> Result, GatewayError> { + match self.first.take() { + Some(first) => Ok(Some(first)), + None => Err(GatewayError::Internal( + "synthetic next-candidate planning failure".to_string(), + )), + } + } + + async fn drain_execution_attempts( + &mut self, + ) -> Result, GatewayError> { + Ok(self.first.take().into_iter().collect()) + } + + async fn skip_credential(&mut self, _key_id: &str) -> Result<(), GatewayError> { + Ok(()) + } + + async fn skip_endpoint(&mut self, _endpoint_id: &str) -> Result<(), GatewayError> { + Ok(()) + } + + async fn skip_provider(&mut self, _provider_id: &str) -> Result<(), GatewayError> { + Ok(()) + } + } + #[derive(Clone)] struct TransferTestAttempt { label: &'static str, @@ -2206,6 +2726,190 @@ mod tests { } } + #[tokio::test] + async fn dynamic_sync_planning_error_releases_reserved_http_plan_cost() { + let now_unix_secs = current_unix_ms() / 1_000; + let request_id = "req-http-reservation-planning-error"; + let subject_id = "user-http-reservation-planning-error"; + let settlement = Arc::new(InMemorySettlementRepository::default()); + let billing = Arc::new(CostReservationBillingRepository { + model_context: StoredBillingModelContext::new( + "provider-1".to_string(), + None, + Some("key-1".to_string()), + None, + None, + "global-model-1".to_string(), + "gpt-5".to_string(), + None, + Some(0.25), + None, + Some("model-1".to_string()), + Some("gpt-5".to_string()), + None, + None, + None, + ) + .expect("billing context should build"), + entitlement: UserPlanEntitlementRecord { + id: "ent-http-reservation-planning-error".to_string(), + user_id: subject_id.to_string(), + plan_id: "plan-http-reservation-planning-error".to_string(), + payment_order_id: "order-http-reservation-planning-error".to_string(), + status: "active".to_string(), + starts_at_unix_secs: now_unix_secs.saturating_sub(60), + expires_at_unix_secs: now_unix_secs.saturating_add(3_600), + entitlements_snapshot: json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "actual_cost_usd", + "window": {"kind": "rolling", "seconds": 3_600}, + "limit": 10.0 + }] + }]), + created_at_unix_secs: now_unix_secs.saturating_sub(60), + updated_at_unix_secs: now_unix_secs.saturating_sub(60), + }, + }); + let data = crate::data::GatewayDataState::with_billing_reader_for_tests(billing) + .with_settlement_writer_for_tests(settlement.clone()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data) + .with_usage_runtime_for_tests(crate::usage::UsageRuntimeConfig { + enabled: true, + ..crate::usage::UsageRuntimeConfig::default() + }) + .with_execution_runtime_sync_override_for_tests(|plan| { + Ok(ExecutionResult { + request_id: plan.request_id.clone(), + candidate_id: plan.candidate_id.clone(), + status_code: http::StatusCode::TOO_MANY_REQUESTS.as_u16(), + headers: Default::default(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({"error": {"message": "retry elsewhere"}})), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }) + }); + + let mut decision = GatewayControlDecision::synthetic( + "/v1/chat/completions", + Some("ai_public".to_string()), + Some("openai".to_string()), + Some("chat".to_string()), + Some("openai:chat".to_string()), + ) + .with_execution_runtime_candidate(true); + decision.auth_context = Some(crate::control::GatewayControlAuthContext { + user_id: subject_id.to_string(), + api_key_id: "api-key-http-reservation-planning-error".to_string(), + username: None, + api_key_name: None, + balance_remaining: None, + access_allowed: true, + user_rate_limit: None, + api_key_rate_limit: None, + api_key_is_standalone: false, + admin_bypass_limits: false, + local_rejection: None, + allowed_models: None, + ip_rules: None, + verified_api_key_hash: None, + }); + let admission = crate::plan_usage_policy::check_and_acquire_http_plan_usage_policy( + &state, + Some(&decision), + request_id, + now_unix_secs.saturating_mul(1_000), + ) + .await + .expect("HTTP plan admission should succeed"); + let reservation = admission + .reservation_context + .expect("cost policy should freeze a reservation context"); + let reservation_token = reservation.token().to_string(); + let admitted_at_unix_secs = reservation.admitted_at_unix_secs(); + + let mut request = http::Request::builder() + .uri("/v1/chat/completions") + .body(()) + .expect("request should build"); + request.extensions_mut().insert(reservation); + let (parts, _) = request.into_parts(); + let mut plan = test_plan(None); + plan.request_id = request_id.to_string(); + plan.candidate_id = Some("cand-http-reservation-planning-error".to_string()); + plan.provider_id = "provider-1".to_string(); + plan.endpoint_id = "endpoint-1".to_string(); + plan.key_id = "key-1".to_string(); + plan.client_api_format = "openai:chat".to_string(); + plan.provider_api_format = "openai:chat".to_string(); + plan.model_name = Some("gpt-5".to_string()); + plan.stream = false; + plan.body = RequestBody::from_json(json!({ + "model": "gpt-5", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 16 + })); + let source = PlanningErrorAfterFirstSyncAttemptSource { + first: Some(crate::ai_serving::AiSyncAttempt { + plan, + report_kind: None, + report_context: Some(json!({ + "candidate_index": 0, + "retry_index": 0, + "model_id": "model-1", + "global_model_name": "gpt-5" + })), + }), + }; + + let error = execute_sync_attempt_source::( + &state, + &parts, + "trace-http-reservation-planning-error", + &decision, + "openai_chat_sync", + source, + ) + .await + .expect_err("second candidate planning should fail"); + assert!(matches!( + error, + GatewayError::Internal(message) + if message == "synthetic next-candidate planning failure" + )); + + let probe = settlement + .reserve_usage_policy_cost(ReserveUsagePolicyCostInput { + request_id: request_id.to_string(), + subject_id: subject_id.to_string(), + reservation_token, + admitted_at_unix_secs, + reserved_cost_units: 1, + reservation_expires_at_unix_secs: admitted_at_unix_secs.saturating_add(86_400), + retain_until_unix_secs: admitted_at_unix_secs.saturating_add(32 * 86_400), + windows: vec![UsagePolicyCostWindow { + window_id: "release-state-probe".to_string(), + starts_at_unix_secs: admitted_at_unix_secs.saturating_sub(1), + ends_at_unix_secs: admitted_at_unix_secs.saturating_add(1), + limit_cost_units: 1_000_000_000, + }], + }) + .await + .expect("reservation state probe should succeed"); + assert_eq!( + probe, + ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: UsagePolicyCostReservationState::Released, + } + ); + } + fn test_report_context() -> serde_json::Value { json!({ "request_id": "req_watchdog", @@ -2217,6 +2921,91 @@ mod tests { }) } + #[test] + fn request_tracker_preserves_server_reservation_identity_across_clones() { + let mut request = http::Request::new(()); + let gate = aether_runtime::ConcurrencyGate::new("image_heartbeat_request", 1); + let admission = AdmissionPermit::from(gate.try_acquire().expect("request admission")); + request + .extensions_mut() + .insert(BackgroundAdmissionPermit::new(admission.clone())); + request.extensions_mut().insert( + crate::plan_usage_policy::PlanUsageReservationContext::for_test( + "user-1", + "server-token", + 12_345, + crate::plan_usage_policy::EffectivePlanUsagePolicy::default(), + ), + ); + let (parts, _) = request.into_parts(); + + let tracker = ProviderTransferTracker::for_request(&parts); + let cloned = tracker.clone(); + + let reservation = tracker + .usage_policy_reservation + .as_ref() + .expect("request reservation snapshot"); + let cloned_reservation = cloned + .usage_policy_reservation + .as_ref() + .expect("cloned reservation snapshot"); + assert_eq!(reservation.admitted_at_unix_secs(), 12_345); + assert_eq!(reservation.subject_id(), "user-1"); + assert_eq!(reservation.token(), "server-token"); + assert_eq!(cloned_reservation.token(), "server-token"); + assert!(std::ptr::eq( + reservation.policy(), + cloned_reservation.policy() + )); + + drop((parts, admission, tracker)); + assert_eq!(gate.snapshot().in_flight, 1); + assert!( + gate.try_acquire().is_err(), + "image heartbeat tracker clone should retain request admission" + ); + drop(cloned); + assert_eq!(gate.snapshot().in_flight, 0); + } + + #[test] + fn server_reservation_token_overrides_report_context_value() { + let context = attach_plan_usage_reservation_token( + Some(json!({ + "candidate_index": 2, + "plan_usage_reservation_token": "client-controlled-token" + })), + Some("server-token"), + ) + .expect("reservation context"); + + assert_eq!(context["candidate_index"], 2); + assert_eq!(context["plan_usage_reservation_token"], "server-token"); + assert_eq!(context["plan_usage_reservation_deferred"], false); + } + + #[test] + fn missing_reservation_does_not_create_or_propagate_a_token() { + let request = http::Request::new(()); + let (parts, _) = request.into_parts(); + let tracker = ProviderTransferTracker::for_request(&parts); + assert!(tracker.usage_policy_reservation.is_none()); + assert!(tracker.usage_policy_reservation_token().is_none()); + + let original = json!({ + "candidate_index": 2, + "plan_usage_reservation_token": "untrusted-seed-token", + "plan_usage_reservation_deferred": true + }); + let expected = json!({"candidate_index": 2}); + assert_eq!( + attach_plan_usage_reservation_token(Some(original), None), + Some(expected) + ); + assert_eq!(attach_plan_usage_reservation_token(None, None), None); + } + #[test] fn stream_candidate_watchdog_prefers_first_byte_timeout() { let report_context = json!({"upstream_is_stream": true}); @@ -2361,6 +3150,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), false, || { std::future::pending::< @@ -2399,6 +3189,56 @@ mod tests { assert_eq!(record.candidate_index, 2); } + #[tokio::test] + async fn stream_candidate_retry_does_not_reset_an_expired_request_first_byte_budget() { + let writer = Arc::new(TestRequestCandidateWriter::default()); + let plan = test_plan(Some(ExecutionTimeouts { + first_byte_ms: Some(250), + ..ExecutionTimeouts::default() + })); + let report_context = test_report_context(); + // Stand in for earlier candidates having already consumed the request's + // complete first-byte budget. A per-candidate watchdog would wait a new + // 250 ms here; the shared absolute deadline must settle immediately. + let request_first_byte_started_at = Instant::now() - Duration::from_millis(300); + + let result = tokio::time::timeout( + Duration::from_millis(100), + execute_stream_candidate_with_watchdog( + writer.as_ref(), + "trace_watchdog_shared_budget", + "claude_cli_stream", + &plan, + Some(&report_context), + request_first_byte_started_at, + false, + || { + std::future::pending::< + Result>, GatewayError>, + >() + }, + ), + ) + .await + .expect("an expired request-level first-byte budget must not restart per candidate"); + + assert!(matches!( + result, + Ok(StreamCandidateWatchdogOutcome::Executed( + AiAttemptExecutionOutcome::Retry { + scope: AiAttemptRetryScope::Candidate, + fallback_response: None, + } + )) + )); + let records = writer.records.lock().await; + assert_eq!(records.len(), 1); + assert_eq!( + records[0].error_type.as_deref(), + Some("local_stream_candidate_watchdog_timeout") + ); + } + #[tokio::test] async fn stream_candidate_watchdog_can_stop_on_transport_error() { let writer = Arc::new(TestRequestCandidateWriter::default()); @@ -2414,6 +3254,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), true, || { std::future::pending::< @@ -2451,6 +3292,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), true, || async { mark_stream_candidate_watchdog_terminal_started(); @@ -2483,6 +3325,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), true, || async { Err(GatewayError::UpstreamUnavailable { @@ -2522,6 +3365,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), false, || async { panic!("execute future should not run while upstream execution gate is saturated") @@ -2569,6 +3413,7 @@ mod tests { "claude_cli_stream", &plan, Some(&report_context), + Instant::now(), false, || async { Err(GatewayError::AdmissionTimeout { diff --git a/apps/aether-gateway/src/executor/mod.rs b/apps/aether-gateway/src/executor/mod.rs index fe2fc9e6a..3031806d8 100644 --- a/apps/aether-gateway/src/executor/mod.rs +++ b/apps/aether-gateway/src/executor/mod.rs @@ -17,10 +17,11 @@ pub(crate) use candidate_loop::{ }; pub(crate) use orchestration::*; pub(crate) use outcome::{ - beautify_local_execution_client_error_message, build_fast_local_execution_exhaustion, - build_fast_local_execution_runtime_miss_context, build_local_execution_exhaustion, - build_local_execution_runtime_miss_context, is_deferred_upstream_response, - mark_deferred_upstream_response, record_failed_usage_for_exhausted_request, + attach_deferred_usage_context, beautify_local_execution_client_error_message, + build_fast_local_execution_exhaustion, build_fast_local_execution_runtime_miss_context, + build_local_execution_exhaustion, build_local_execution_runtime_miss_context, + is_deferred_upstream_response, mark_deferred_upstream_response, + record_failed_usage_for_deferred_response, record_failed_usage_for_exhausted_request, record_failed_usage_for_runtime_miss_request, LocalExecutionExhaustion, LocalExecutionRequestOutcome, LocalExecutionRuntimeMissContext, }; diff --git a/apps/aether-gateway/src/executor/orchestration.rs b/apps/aether-gateway/src/executor/orchestration.rs index 366690283..7fd6ebca6 100644 --- a/apps/aether-gateway/src/executor/orchestration.rs +++ b/apps/aether-gateway/src/executor/orchestration.rs @@ -54,13 +54,10 @@ use crate::executor::{ record_failed_usage_for_exhausted_request, LocalExecutionExhaustion, LocalExecutionRequestOutcome, }; -use crate::handlers::shared::system_config_bool; use crate::request_diagnostics::{current_request_diagnostics, scope_request_diagnostics_with}; use crate::stage_metrics::observe_gateway_stage_ms; use crate::{AiExecutionDecision, AppState, GatewayError}; -const ENABLE_OPENAI_IMAGE_SYNC_HEARTBEAT_CONFIG_KEY: &str = "enable_openai_image_sync_heartbeat"; -const ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY: &str = "enable_standard_text_sync_heartbeat"; const OPENAI_IMAGE_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS: u16 = 502; const OPENAI_IMAGE_SYNC_HEARTBEAT_EXHAUSTED_STATUS: u16 = 503; const OPENAI_IMAGE_SYNC_HEARTBEAT_ERROR_MESSAGE_LIMIT: usize = 4096; @@ -107,7 +104,10 @@ pub(crate) async fn maybe_execute_sync_via_local_decision( return Ok(LocalExecutionRequestOutcome::NoPath); }; - if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await { + if standard_text_sync_heartbeat_should_wrap( + plan_kind, + attempt_source.routing_execution_policy(), + ) { let parts_for_task = parts.clone(); let body_json_for_task = body_json.clone(); let transfer_tracker_for_task = transfer_tracker.clone(); @@ -264,7 +264,10 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision( return Ok(LocalExecutionRequestOutcome::NoPath); }; - if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await { + if standard_text_sync_heartbeat_should_wrap( + plan_kind, + attempt_source.routing_execution_policy(), + ) { let parts_for_task = parts.clone(); let body_json_for_task = body_json.clone(); let transfer_tracker_for_task = transfer_tracker.clone(); @@ -381,7 +384,10 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision( return Ok(LocalExecutionRequestOutcome::NoPath); }; - if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await { + if standard_text_sync_heartbeat_should_wrap( + plan_kind, + attempt_source.routing_execution_policy(), + ) { let parts_for_task = parts.clone(); let body_json_for_task = body_json.clone(); let transfer_tracker_for_task = transfer_tracker.clone(); @@ -609,7 +615,10 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision( return Ok(LocalExecutionRequestOutcome::NoPath); }; - if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await { + if standard_text_sync_heartbeat_should_wrap( + plan_kind, + attempt_source.routing_execution_policy(), + ) { let parts_for_task = parts.clone(); let body_json_for_task = body_json.clone(); let transfer_tracker_for_task = transfer_tracker.clone(); @@ -746,42 +755,6 @@ pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision( .await } -async fn openai_image_sync_heartbeat_enabled(state: &AppState) -> bool { - match state - .read_system_config_json_value(ENABLE_OPENAI_IMAGE_SYNC_HEARTBEAT_CONFIG_KEY) - .await - { - Ok(value) => system_config_bool(value.as_ref(), false), - Err(err) => { - tracing::warn!( - event_name = "openai_image_sync_heartbeat_config_read_failed", - log_type = "ops", - error = ?err, - "gateway failed to read sync image heartbeat config; defaulting disabled" - ); - false - } - } -} - -async fn standard_text_sync_heartbeat_enabled(state: &AppState) -> bool { - match state - .read_system_config_json_value(ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY) - .await - { - Ok(value) => system_config_bool(value.as_ref(), false), - Err(err) => { - tracing::warn!( - event_name = "standard_text_sync_heartbeat_config_read_failed", - log_type = "ops", - error = ?err, - "gateway failed to read standard text sync heartbeat config; defaulting disabled" - ); - false - } - } -} - fn standard_text_sync_heartbeat_applies_to_plan_kind(plan_kind: &str) -> bool { matches!( plan_kind, @@ -795,9 +768,12 @@ fn standard_text_sync_heartbeat_applies_to_plan_kind(plan_kind: &str) -> bool { ) } -async fn standard_text_sync_heartbeat_should_wrap(state: &AppState, plan_kind: &str) -> bool { +fn standard_text_sync_heartbeat_should_wrap( + plan_kind: &str, + execution_policy: Option, +) -> bool { standard_text_sync_heartbeat_applies_to_plan_kind(plan_kind) - && standard_text_sync_heartbeat_enabled(state).await + && execution_policy.is_some_and(|policy| policy.enable_cf_heartbeat) } fn standard_text_sync_heartbeat_client_api_format_for_plan_kind(plan_kind: &str) -> &'static str { @@ -923,10 +899,10 @@ async fn standard_text_sync_heartbeat_final_bytes( STANDARD_TEXT_SYNC_HEARTBEAT_EXHAUSTED_STATUS, "standard text sync exhausted all local candidates", ), - Err(err) => standard_text_sync_heartbeat_error_body( + Err(_err) => standard_text_sync_heartbeat_error_body( client_api_format, STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS, - &format!("{err:?}"), + "internal gateway error while executing request", ), } } @@ -946,11 +922,11 @@ async fn standard_text_sync_heartbeat_response_body_bytes( bytes.as_ref(), ) { Ok(body) => body, - Err(err) => { + Err(_err) => { return standard_text_sync_heartbeat_error_body( client_api_format, STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS, - &format!("{err:?}"), + "internal gateway error while restoring response", ); } }; @@ -970,10 +946,10 @@ async fn standard_text_sync_heartbeat_response_body_bytes( "empty standard text sync response", ) } - Err(err) => standard_text_sync_heartbeat_error_body( + Err(_err) => standard_text_sync_heartbeat_error_body( client_api_format, STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS, - &err.to_string(), + "internal gateway error while reading response", ), } } @@ -1224,9 +1200,9 @@ async fn openai_image_sync_heartbeat_final_bytes( OPENAI_IMAGE_SYNC_HEARTBEAT_EXHAUSTED_STATUS, "OpenAI image sync exhausted all local candidates", ), - Err(err) => openai_image_sync_heartbeat_error_body( + Err(_err) => openai_image_sync_heartbeat_error_body( OPENAI_IMAGE_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS, - &format!("{err:?}"), + "internal gateway error while executing image request", ), } } @@ -1247,9 +1223,9 @@ async fn openai_image_sync_heartbeat_response_body_bytes(response: Response openai_image_sync_heartbeat_error_body( + Err(_err) => openai_image_sync_heartbeat_error_body( OPENAI_IMAGE_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS, - &err.to_string(), + "internal gateway error while reading image response", ), } } @@ -1329,7 +1305,10 @@ pub(crate) async fn maybe_execute_sync_via_local_image_decision( return Ok(LocalExecutionRequestOutcome::NoPath); }; - if openai_image_sync_heartbeat_enabled(state).await { + if attempt_source + .routing_execution_policy() + .is_some_and(|policy| policy.enable_cf_heartbeat) + { let mut attempts = Vec::new(); while let Some(attempt) = attempt_source.next_execution_attempt().await? { attempts.push(attempt); @@ -1663,6 +1642,22 @@ mod tests { candidate_index: u32, endpoint_id: &str, candidate_id: &str, + ) -> AiSyncAttempt { + test_openai_image_heartbeat_attempt_with_sticky_key_attempts( + candidate_index, + endpoint_id, + candidate_id, + 1, + ) + } + + /// `sticky_key_attempts` is pinned so these tests exercise candidate + /// failover; the default same-key retry is covered separately. + fn test_openai_image_heartbeat_attempt_with_sticky_key_attempts( + candidate_index: u32, + endpoint_id: &str, + candidate_id: &str, + sticky_key_attempts: u32, ) -> AiSyncAttempt { AiSyncAttempt { plan: test_openai_image_heartbeat_plan(endpoint_id, candidate_id), @@ -1670,6 +1665,7 @@ mod tests { report_context: Some(json!({ "candidate_index": candidate_index, "retry_index": 0, + "sticky_key_attempts": sticky_key_attempts, })), } } @@ -1828,6 +1824,9 @@ mod tests { report_context: Some(json!({ "candidate_index": candidate_index, "retry_index": 0, + // Pin to a single attempt so this helper exercises candidate + // failover rather than the default same-key retry. + "sticky_key_attempts": 1, "client_api_format": client_api_format, "provider_api_format": client_api_format, })), @@ -1847,11 +1846,10 @@ mod tests { assert_eq!(body, json!({"data": [{"b64_json": "x"}]})); } - #[tokio::test] - async fn openai_image_sync_heartbeat_missing_config_defaults_disabled() { - let state = AppState::new().expect("state should build"); - - assert!(!openai_image_sync_heartbeat_enabled(&state).await); + #[test] + fn openai_image_sync_heartbeat_missing_routing_policy_defaults_disabled() { + assert!(!Option::::None + .is_some_and(|policy| policy.enable_cf_heartbeat)); } #[tokio::test] @@ -1980,6 +1978,90 @@ mod tests { assert_eq!(body, json!({"data": [{"b64_json": "second-candidate"}]})); } + #[tokio::test] + async fn openai_image_sync_heartbeat_retries_sticky_key_lazily_before_failover() { + let seen_plans = Arc::new(std::sync::Mutex::new(Vec::<(String, Option)>::new())); + let seen_plans_for_override = Arc::clone(&seen_plans); + let state = AppState::new() + .expect("state should build") + .with_execution_runtime_sync_override_for_tests(move |plan| { + seen_plans_for_override + .lock() + .expect("mutex should lock") + .push((plan.endpoint_id.clone(), plan.candidate_id.clone())); + if plan.endpoint_id == "endpoint-retry" { + Ok(test_openai_image_execution_result( + plan, + StatusCode::TOO_MANY_REQUESTS.as_u16(), + json!({"error": {"message": "retry this candidate"}}), + )) + } else { + Ok(test_openai_image_execution_result( + plan, + StatusCode::OK.as_u16(), + json!({"data": [{"b64_json": "second-candidate"}]}), + )) + } + }); + // Three total attempts on the sticky key; only one attempt is + // materialized up front, the other two are derived after each failure. + let attempts = vec![ + test_openai_image_heartbeat_attempt_with_sticky_key_attempts( + 0, + "endpoint-retry", + "candidate-retry", + 3, + ), + test_openai_image_heartbeat_attempt_with_sticky_key_attempts( + 1, + "endpoint-success", + "candidate-success", + 3, + ), + ]; + let outcome = execute_openai_image_sync_heartbeat_attempts( + state, + "/v1/images/generations".to_string(), + "trace-image-heartbeat-sticky-retry".to_string(), + test_openai_image_heartbeat_decision(), + TEST_OPENAI_IMAGE_SYNC_PLAN_KIND.to_string(), + attempts, + ProviderTransferTracker::default(), + Instant::now(), + ) + .await + .expect("heartbeat attempts should execute"); + let LocalExecutionRequestOutcome::Responded(response) = outcome else { + panic!("second candidate should return a response"); + }; + let bytes = openai_image_sync_heartbeat_response_body_bytes(response).await; + let body: Value = serde_json::from_slice(&bytes).expect("body should decode"); + + let seen_plans = seen_plans.lock().expect("mutex should lock").clone(); + assert_eq!( + seen_plans + .iter() + .map(|(endpoint_id, _)| endpoint_id.as_str()) + .collect::>(), + [ + "endpoint-retry", + "endpoint-retry", + "endpoint-retry", + "endpoint-success" + ] + ); + let sticky_candidate_ids = seen_plans[..3] + .iter() + .map(|(_, candidate_id)| candidate_id.clone()) + .collect::>(); + assert_eq!( + sticky_candidate_ids.len(), + 3, + "each derived same-key retry must carry a fresh candidate id" + ); + assert_eq!(body, json!({"data": [{"b64_json": "second-candidate"}]})); + } + #[tokio::test] async fn openai_image_sync_heartbeat_honors_provider_transfer_limit() { let call_count = Arc::new(AtomicUsize::new(0)); @@ -2013,6 +2095,7 @@ mod tests { attempt.report_context = Some(json!({ "candidate_index": index, "retry_index": 0, + "sticky_key_attempts": 1, "local_failover_policy": { "max_transfer_count": 1, "max_transfer_timeout_seconds": 0 @@ -2044,23 +2127,9 @@ mod tests { assert_eq!(body, json!({"data": [{"b64_json": "fallback-provider"}]})); } - #[tokio::test] - async fn standard_text_sync_heartbeat_missing_config_defaults_disabled() { - let state = AppState::new().expect("state should build"); - - assert!(!standard_text_sync_heartbeat_enabled(&state).await); - } - #[tokio::test] async fn standard_text_sync_heartbeat_no_local_candidates_preserves_no_path() { - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - crate::data::GatewayDataState::disabled().with_system_config_values_for_tests([( - ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY.to_string(), - json!(true), - )]), - ); + let state = AppState::new().expect("state should build"); let (parts, _) = http::Request::builder() .method(http::Method::POST) .uri("/v1/responses") @@ -2198,6 +2267,70 @@ mod tests { let _ = release_tx.send(()); } + #[tokio::test] + async fn standard_text_sync_heartbeat_background_holds_request_admission_after_disconnect() { + let state = AppState::new().expect("state should build"); + let gate = aether_runtime::ConcurrencyGate::new("heartbeat_request", 1); + let admission = aether_runtime::AdmissionPermit::from( + gate.try_acquire().expect("request admission permit"), + ); + let (mut parts, _) = http::Request::builder() + .method(http::Method::POST) + .uri("/v1/responses") + .body(()) + .expect("request should build") + .into_parts(); + parts.extensions.insert( + crate::executor::candidate_loop::BackgroundAdmissionPermit::new(admission.clone()), + ); + drop(admission); + + let (started_tx, started_rx) = tokio::sync::oneshot::channel::<()>(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>(); + let response = build_standard_text_sync_heartbeat_shell_response( + state, + parts, + "trace-standard-text-heartbeat-admission".to_string(), + test_standard_text_heartbeat_decision(), + TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(), + move |_state, parts, _trace_id, _decision, _plan_kind, _started_at| async move { + assert!( + parts + .extensions + .get::() + .is_some(), + "background request parts should carry admission" + ); + let _ = started_tx.send(()); + let _ = release_rx.await; + Ok(LocalExecutionRequestOutcome::responded( + Response::builder() + .status(StatusCode::OK) + .body(Body::from(r#"{"id":"resp_done","output":[]}"#)) + .expect("response should build"), + )) + }, + ) + .expect("heartbeat shell should build"); + + started_rx.await.expect("background execution should start"); + drop(response); + assert_eq!(gate.snapshot().in_flight, 1); + assert!( + gate.try_acquire().is_err(), + "disconnect must not release background admission" + ); + + let _ = release_tx.send(()); + tokio::time::timeout(Duration::from_secs(1), async { + while gate.snapshot().in_flight != 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("background completion should release admission"); + } + #[tokio::test] async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() { let (state, usage_repository) = heartbeat_usage_test_state(json!({ diff --git a/apps/aether-gateway/src/executor/outcome.rs b/apps/aether-gateway/src/executor/outcome.rs index a9a255743..4d06f48b9 100644 --- a/apps/aether-gateway/src/executor/outcome.rs +++ b/apps/aether-gateway/src/executor/outcome.rs @@ -3,14 +3,13 @@ use std::time::{Duration, Instant}; use aether_contracts::ExecutionPlan; use aether_data_contracts::repository::candidates::{ + sanitize_request_candidate_error_type, sanitize_request_candidate_skip_reason, RequestCandidateStatus, StoredRequestCandidate, }; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; -use aether_data_contracts::repository::usage::{ - ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, -}; +use aether_data_contracts::repository::usage::ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY; use aether_usage_runtime::{ build_usage_event_data_seed, UsageEvent, UsageEventData, UsageEventType, }; @@ -39,6 +38,12 @@ pub(crate) enum LocalExecutionRequestOutcome { #[derive(Debug, Clone, Copy)] pub(crate) struct DeferredUpstreamResponse; +#[derive(Debug, Clone)] +pub(crate) struct DeferredUsageContext { + plan: ExecutionPlan, + report_context: Option, +} + #[derive(Debug, Clone)] pub(crate) struct LocalExecutionExhaustion { request_id: String, @@ -47,7 +52,6 @@ pub(crate) struct LocalExecutionExhaustion { candidate_index: Option, upstream_status_code: Option, upstream_error_type: Option, - upstream_error_message: Option, } #[derive(Debug, Clone, Default)] @@ -82,6 +86,53 @@ pub(crate) fn mark_deferred_upstream_response(mut response: Response) -> R response } +pub(crate) fn attach_deferred_usage_context( + response: &mut Response, + plan: &ExecutionPlan, + report_context: Option<&Value>, +) { + response.extensions_mut().insert(DeferredUsageContext { + plan: plan.clone(), + report_context: report_context.cloned(), + }); +} + +pub(crate) fn record_failed_usage_for_deferred_response<'a>( + state: &'a AppState, + response: &Response, +) -> impl std::future::Future + Send + 'a { + let context = response.extensions().get::().cloned(); + let status_code = response.status().as_u16(); + async move { + if !state.usage_runtime.is_enabled() { + return; + } + let Some(context) = context else { + return; + }; + let mut data = build_usage_event_data_seed(&context.plan, context.report_context.as_ref()); + data.status_code = Some(status_code); + data.error_message = + Some("all local candidates failed; returning preserved upstream error".to_string()); + data.error_category = error_category_for_failed_status(status_code) + .or_else(|| Some("upstream_error".to_string())); + data.response_headers = Some(json_header_map()); + data.client_response_headers = Some(json_header_map()); + + state + .usage_runtime + .record_terminal_event_direct( + state.usage_lifecycle_data_state().as_ref(), + UsageEvent::new( + UsageEventType::Failed, + context.plan.request_id.clone(), + data, + ), + ) + .await; + } +} + pub(crate) fn is_deferred_upstream_response(response: &Response) -> bool { response .extensions() @@ -111,6 +162,22 @@ impl LocalExecutionRuntimeMissContext { }) } + pub(crate) fn all_candidates_skipped_for_reasons(&self, reasons: &[&str]) -> bool { + if reasons.is_empty() || self.candidate_contexts.is_empty() { + return false; + } + + self.candidate_contexts.iter().all(|candidate| { + candidate.candidate.status == RequestCandidateStatus::Skipped + && candidate + .candidate + .skip_reason + .as_deref() + .map(str::trim) + .is_some_and(|value| reasons.contains(&value)) + }) + } + pub(crate) fn candidate_summary(&self) -> Option { const MAX_ITEMS: usize = 5; @@ -146,24 +213,10 @@ impl LocalExecutionRuntimeMissContext { return None; } - let diagnostic = self - .candidate_contexts - .iter() - .find_map(runtime_miss_candidate_failure_diagnostic)?; - let mut detail = format!("上游请求体转换失败:{}", diagnostic.message); - if diagnostic.path != "$" { - detail.push_str(&format!(";字段路径:{}", diagnostic.path)); - } - detail.push_str("(原因代码: provider_request_body_build_failed)"); - Some(detail) + Some("上游请求体转换失败(原因代码: provider_request_body_build_failed)".to_string()) } } -struct RuntimeMissFailureDiagnostic { - path: String, - message: String, -} - pub(crate) async fn build_local_execution_exhaustion( state: &AppState, plan: &ExecutionPlan, @@ -211,16 +264,11 @@ pub(crate) async fn build_local_execution_exhaustion( exhaustion.upstream_status_code = last_failed_candidate .as_ref() .and_then(|candidate| candidate.status_code); - exhaustion.upstream_error_type = last_failed_candidate - .as_ref() - .and_then(|candidate| candidate.error_type.clone()) - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - exhaustion.upstream_error_message = last_failed_candidate - .as_ref() - .and_then(|candidate| candidate.error_message.clone()) - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); + exhaustion.upstream_error_type = sanitize_request_candidate_error_type( + last_failed_candidate + .as_ref() + .and_then(|candidate| candidate.error_type.clone()), + ); exhaustion } @@ -240,7 +288,6 @@ pub(crate) fn build_fast_local_execution_exhaustion( data, upstream_status_code: None, upstream_error_type: None, - upstream_error_message: None, } } @@ -296,15 +343,12 @@ pub(crate) async fn record_failed_usage_for_exhausted_request( candidate_index, upstream_status_code, upstream_error_type, - upstream_error_message, } = exhaustion; let status_code = http::StatusCode::SERVICE_UNAVAILABLE.as_u16(); let candidate_status_code = upstream_status_code.unwrap_or(status_code); data.status_code = Some(status_code); - data.error_message = upstream_error_message - .clone() - .or_else(|| Some(local_execution_runtime_miss_detail.to_string())); + data.error_message = Some(local_execution_runtime_miss_detail.to_string()); data.error_category = error_category_for_failed_status(status_code); data.response_time_ms = Some(started_at.elapsed().as_millis() as u64); data.response_headers = Some(json_header_map()); @@ -314,10 +358,7 @@ pub(crate) async fn record_failed_usage_for_exhausted_request( .as_deref() .filter(|value| !value.trim().is_empty()) .unwrap_or("upstream_error"), - "message": upstream_error_message - .as_deref() - .filter(|value| !value.trim().is_empty()) - .unwrap_or(local_execution_runtime_miss_detail), + "message": local_execution_runtime_miss_detail, "code": candidate_status_code, } })); @@ -977,82 +1018,13 @@ fn insert_runtime_miss_candidate_usage_metadata( metadata: &mut Map, candidate: &StoredRequestCandidate, ) { - if let Some(skip_reason) = candidate - .skip_reason - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) + if let Some(skip_reason) = sanitize_request_candidate_skip_reason(candidate.skip_reason.clone()) { metadata.insert( ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY.to_string(), - Value::String(skip_reason.to_string()), + Value::String(skip_reason), ); } - - let diagnostic = candidate - .extra_data - .as_ref() - .and_then(Value::as_object) - .and_then(|extra_data| { - extra_data - .get("failure_diagnostic") - .filter(|value| { - value.as_object().is_some_and(|diagnostic| { - diagnostic.get("safe_to_show") != Some(&Value::Bool(false)) - }) - }) - .or_else(|| { - extra_data - .get("request_conversion_error") - .filter(|v| v.is_object()) - }) - .or_else(|| { - extra_data - .get("request_body_build_error") - .filter(|v| v.is_object()) - }) - }); - if let Some(diagnostic) = diagnostic { - metadata.insert( - ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY.to_string(), - diagnostic.clone(), - ); - } -} - -fn runtime_miss_candidate_failure_diagnostic( - candidate: &RuntimeMissCandidateContext, -) -> Option { - let extra_data = candidate.candidate.extra_data.as_ref()?.as_object()?; - let diagnostic = extra_data - .get("failure_diagnostic") - .and_then(Value::as_object) - .filter(|diagnostic| diagnostic.get("safe_to_show") != Some(&Value::Bool(false))) - .or_else(|| { - extra_data - .get("request_conversion_error") - .and_then(Value::as_object) - }) - .or_else(|| { - extra_data - .get("request_body_build_error") - .and_then(Value::as_object) - })?; - let message = diagnostic - .get("message") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty())?; - let path = diagnostic - .get("path") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or("$"); - Some(RuntimeMissFailureDiagnostic { - path: path.to_string(), - message: message.to_string(), - }) } fn build_runtime_miss_candidate_endpoint_url( @@ -1120,7 +1092,11 @@ fn format_runtime_miss_candidate_summary(candidate: &RuntimeMissCandidateContext .map(str::trim) .filter(|value| !value.is_empty()) { - parts.push(format!("url={endpoint_url}")); + // Candidate URLs can contain provider API keys or other query + // credentials. Runtime-miss summaries are emitted to logs and may + // cross an operator/client boundary, so retain only the safe origin. + let origin = crate::handlers::shared::security_log_url_origin(endpoint_url); + parts.push(format!("url={origin}")); } if let Some(key_label) = format_name_with_id( candidate.key_name.as_deref(), @@ -1392,10 +1368,10 @@ mod tests { request_metadata["routing_candidate_skip_reason"], "provider_request_body_build_failed" ); - assert_eq!( - request_metadata["routing_failure_diagnostic"]["path"], - "$.reasoning.summary" - ); + assert!(request_metadata.get("routing_failure_diagnostic").is_none()); + assert!(!Value::Object(request_metadata.clone()) + .to_string() + .contains("invalid reasoning summary")); assert!(!request_candidate_represents_provider_execution( &skipped_candidate @@ -1421,7 +1397,7 @@ mod tests { } #[test] - fn runtime_miss_context_surfaces_request_conversion_field_diagnostic() { + fn runtime_miss_context_uses_fixed_request_body_build_failure_detail() { let skipped_candidate = StoredRequestCandidate::new( "cand-skipped".to_string(), "req-1".to_string(), @@ -1477,10 +1453,11 @@ mod tests { let detail = context .all_provider_request_body_build_failures_detail() - .expect("detail should include conversion diagnostic"); + .expect("detail should identify the fixed failure category"); - assert!(detail.contains("字段 n")); - assert!(detail.contains("字段路径:$.n")); + assert!(!detail.contains("字段 n")); + assert!(!detail.contains("字段路径")); + assert!(!detail.contains("OpenAI Responses")); assert!(detail.contains("provider_request_body_build_failed")); } } diff --git a/apps/aether-gateway/src/executor/stream_path.rs b/apps/aether-gateway/src/executor/stream_path.rs index 921122e16..0976961b5 100644 --- a/apps/aether-gateway/src/executor/stream_path.rs +++ b/apps/aether-gateway/src/executor/stream_path.rs @@ -18,6 +18,7 @@ use crate::ai_serving::api::{ use crate::api::response::build_client_response_from_parts; use crate::control::GatewayControlDecision; use crate::stage_metrics::observe_gateway_stage_ms; +use crate::state::VideoTaskRouteAccess; use crate::{AppState, GatewayError, GatewayFallbackReason}; use super::{ @@ -101,7 +102,7 @@ pub(crate) async fn maybe_execute_via_stream_decision_path( if skip_direct_plan { return Ok(LocalExecutionRequestOutcome::NoPath); } - let transfer_tracker = ProviderTransferTracker::default(); + let transfer_tracker = ProviderTransferTracker::for_request(parts); if plan_kind == OPENAI_CHAT_STREAM_PLAN_KIND && supports_stream_execution_decision_kind(plan_kind) @@ -490,29 +491,73 @@ async fn maybe_execute_local_video_task_content_stream( return Ok(LocalExecutionRequestOutcome::NoPath); } - let _ = state - .hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path()) - .await?; + let Some(user_id) = decision + .auth_context + .as_ref() + .filter(|auth_context| auth_context.access_allowed) + .map(|auth_context| auth_context.user_id.trim()) + .filter(|value| !value.is_empty()) + else { + return Ok(LocalExecutionRequestOutcome::Responded( + build_json_response( + trace_id, + decision, + 404, + &crate::video_tasks::not_found_body(), + )?, + )); + }; + + if state + .hydrate_video_task_for_route_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + user_id, + ) + .await? + != VideoTaskRouteAccess::Allowed + { + return Ok(LocalExecutionRequestOutcome::Responded( + build_json_response( + trace_id, + decision, + 404, + &crate::video_tasks::not_found_body(), + )?, + )); + } if let Some(task_id) = crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path()) { let refresh_path = format!("/v1/videos/{task_id}"); - if let Some(refresh_plan) = state.video_tasks.prepare_read_refresh_sync_plan( + if let Some(refresh_plan) = state.video_tasks.prepare_read_refresh_sync_plan_for_user( Some("openai"), &refresh_path, + user_id, trace_id, ) { state.execute_video_task_refresh_plan(&refresh_plan).await?; } } - let Some(action) = state.video_tasks.prepare_openai_content_stream_action( - parts.uri.path(), - parts.uri.query(), - trace_id, - ) else { - return Ok(LocalExecutionRequestOutcome::NoPath); + let Some(action) = state + .video_tasks + .prepare_openai_content_stream_action_for_user( + parts.uri.path(), + parts.uri.query(), + trace_id, + user_id, + ) + else { + return Ok(LocalExecutionRequestOutcome::Responded( + build_json_response( + trace_id, + decision, + 404, + &crate::video_tasks::not_found_body(), + )?, + )); }; match action { diff --git a/apps/aether-gateway/src/executor/sync_path.rs b/apps/aether-gateway/src/executor/sync_path.rs index 3af16aa78..8e5475ee9 100644 --- a/apps/aether-gateway/src/executor/sync_path.rs +++ b/apps/aether-gateway/src/executor/sync_path.rs @@ -16,6 +16,7 @@ use crate::ai_serving::api::{ use crate::api::response::build_client_response_from_parts; use crate::control::resolve_execution_runtime_auth_context; use crate::control::GatewayControlDecision; +use crate::state::VideoTaskRouteAccess; use crate::{AppState, GatewayError, GatewayFallbackReason}; use super::{ @@ -90,7 +91,7 @@ pub(crate) async fn maybe_execute_via_sync_decision_path( plan_kind, bypass_cache_key, scheduler_supported: supports_sync_execution_decision_kind(plan_kind), - transfer_tracker: ProviderTransferTracker::default(), + transfer_tracker: ProviderTransferTracker::for_request(parts), }; Ok(from_ai_serving_outcome( @@ -321,13 +322,41 @@ async fn maybe_build_local_video_task_read_response( return Ok(LocalExecutionRequestOutcome::NoPath); } - let _ = state - .hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path()) - .await?; - - let refresh_plan = state.video_tasks.prepare_read_refresh_sync_plan( + if crate::video_tasks::resolve_video_task_read_lookup_key( decision.route_family.as_deref(), parts.uri.path(), + ) + .is_none() + { + return Ok(LocalExecutionRequestOutcome::NoPath); + } + + let Some(user_id) = decision + .auth_context + .as_ref() + .filter(|auth_context| auth_context.access_allowed) + .map(|auth_context| auth_context.user_id.trim()) + .filter(|value| !value.is_empty()) + else { + return build_video_task_not_found_outcome(trace_id, decision); + }; + + if state + .hydrate_video_task_for_route_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + user_id, + ) + .await? + == VideoTaskRouteAccess::Denied + { + return build_video_task_not_found_outcome(trace_id, decision); + } + + let refresh_plan = state.video_tasks.prepare_read_refresh_sync_plan_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + user_id, trace_id, ); @@ -335,22 +364,25 @@ async fn maybe_build_local_video_task_read_response( state.execute_video_task_refresh_plan(&refresh_plan).await?; } - let read_response = state - .video_tasks - .read_response(decision.route_family.as_deref(), parts.uri.path()); + let read_response = state.video_tasks.read_response_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + user_id, + ); let read_response = match read_response { Some(read_response) => Some(read_response), None => { state - .read_data_backed_video_task_response( + .read_data_backed_video_task_response_for_user( decision.route_family.as_deref(), parts.uri.path(), + user_id, ) .await? } }; let Some(read_response) = read_response else { - return Ok(LocalExecutionRequestOutcome::NoPath); + return build_video_task_not_found_outcome(trace_id, decision); }; let body_bytes = serde_json::to_vec(&read_response.body_json) @@ -389,10 +421,6 @@ async fn maybe_execute_local_video_task_follow_up_sync( return Ok(LocalExecutionRequestOutcome::NoPath); } - let _ = state - .hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path()) - .await?; - let auth_context = resolve_execution_runtime_auth_context( state, decision, @@ -401,14 +429,30 @@ async fn maybe_execute_local_video_task_follow_up_sync( trace_id, ) .await?; - let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan( + let Some(auth_context) = auth_context.filter(|auth_context| { + auth_context.access_allowed && !auth_context.user_id.trim().is_empty() + }) else { + return build_video_task_not_found_outcome(trace_id, decision); + }; + if state + .hydrate_video_task_for_route_for_user( + decision.route_family.as_deref(), + parts.uri.path(), + &auth_context.user_id, + ) + .await? + != VideoTaskRouteAccess::Allowed + { + return build_video_task_not_found_outcome(trace_id, decision); + } + let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan_for_user( plan_kind, parts.uri.path(), Some(body_json), - auth_context.as_ref(), + Some(&auth_context), trace_id, ) else { - return Ok(LocalExecutionRequestOutcome::NoPath); + return build_video_task_not_found_outcome(trace_id, decision); }; execute_sync_plan_and_reports_with_transfer_tracker( @@ -426,3 +470,23 @@ async fn maybe_execute_local_video_task_follow_up_sync( ) .await } + +fn build_video_task_not_found_outcome( + trace_id: &str, + decision: &GatewayControlDecision, +) -> Result { + let body_bytes = serde_json::to_vec(&crate::video_tasks::not_found_body()) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut headers = BTreeMap::new(); + headers.insert("content-type".to_string(), "application/json".to_string()); + headers.insert("content-length".to_string(), body_bytes.len().to_string()); + Ok(LocalExecutionRequestOutcome::Responded( + build_client_response_from_parts( + 404, + &headers, + Body::from(body_bytes), + trace_id, + Some(decision), + )?, + )) +} diff --git a/apps/aether-gateway/src/frontdoor_loop_guard.rs b/apps/aether-gateway/src/frontdoor_loop_guard.rs index 19cfd7e7c..52c50feff 100644 --- a/apps/aether-gateway/src/frontdoor_loop_guard.rs +++ b/apps/aether-gateway/src/frontdoor_loop_guard.rs @@ -77,9 +77,9 @@ pub(crate) fn gateway_frontdoor_self_loop_guard_error_with_port( url: &str, ) -> Option { gateway_frontdoor_self_loop_guard_matches_with_port(app_port, url).then(|| { - format!( - "upstream execution target resolves back to the local aether-gateway frontdoor: {url}" - ) + // Do not echo the configured target: provider URLs can carry API keys, + // signed query parameters, or proxy credentials. + "upstream execution target resolves back to the local aether-gateway frontdoor".to_string() }) } @@ -143,7 +143,28 @@ fn normalize_host_for_frontdoor_loop_guard(host: &str) -> String { } fn is_loopbackish_host(host: &str) -> bool { - matches!(host, "localhost" | "127.0.0.1" | "::1" | "0.0.0.0" | "::") + if host.eq_ignore_ascii_case("localhost") { + return true; + } + + // The execution URL validator deliberately permits literal loopback HTTP + // targets for local providers. The frontdoor loop guard must therefore + // use IP semantics instead of a short allowlist: every address in + // 127.0.0.0/8 can reach a listener bound to 0.0.0.0, and IPv4-mapped + // IPv6 forms can represent the same destinations. + let Ok(address) = host.parse::() else { + return false; + }; + match address { + std::net::IpAddr::V4(address) => address.is_loopback() || address.is_unspecified(), + std::net::IpAddr::V6(address) => { + address.is_loopback() + || address.is_unspecified() + || address + .to_ipv4() + .is_some_and(|mapped| mapped.is_loopback() || mapped.is_unspecified()) + } + } } #[cfg(test)] diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/install_routes.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/install_routes.rs index b2e59e180..09ce6977b 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/install_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/install_routes.rs @@ -8,16 +8,13 @@ use crate::handlers::public::{ build_api_key_install_session_response, CreateApiKeyInstallSessionRequest, }; use crate::GatewayError; -use axum::{ - body::Body, - http, - response::{IntoResponse, Response}, -}; +use axum::{body::Body, http, response::Response}; pub(super) async fn build_admin_create_api_key_install_session_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, request_headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&axum::body::Bytes>, ) -> Result, GatewayError> { if !state.has_auth_api_key_data_reader() { @@ -48,40 +45,26 @@ pub(super) async fn build_admin_create_api_key_install_session_response( else { return Ok(build_admin_api_keys_not_found_response()); }; - let Some(ciphertext) = record - .key_encrypted - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - else { - return Ok(build_admin_api_keys_bad_request_response( - "该密钥没有存储完整密钥信息", - )); - }; - let Some(api_key) = state.decrypt_catalog_secret_with_fallbacks(ciphertext) else { - return Ok(( - http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(serde_json::json!({ "detail": "解密密钥失败" })), - ) - .into_response()); - }; let response = build_api_key_install_session_response( state.app(), request_context.public(), request_headers, - record.api_key_id.clone(), - record.name.unwrap_or_else(|| "API Key".to_string()), - api_key, + remote_addr, + &record, payload, ) .await; - Ok(attach_admin_audit_response( - response, - "admin_standalone_api_key_install_session_created", - "create_standalone_api_key_install_session", - "api_key", - &api_key_id, - )) + if response.status().is_success() { + Ok(attach_admin_audit_response( + response, + "admin_standalone_api_key_install_session_created", + "create_standalone_api_key_install_session", + "api_key", + &api_key_id, + )) + } else { + Ok(response) + } } diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/mod.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/mod.rs index f278f177c..4ebaf8fd3 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/mod.rs @@ -6,8 +6,7 @@ use super::super::users::{ }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{ - decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, query_param_bool, - query_param_optional_bool, query_param_value, + query_param_bool, query_param_optional_bool, query_param_value, }; use crate::GatewayError; use axum::{ @@ -42,12 +41,14 @@ pub(crate) async fn maybe_build_local_admin_api_keys_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, request_headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&axum::body::Bytes>, ) -> Result>, GatewayError> { routes::maybe_build_local_admin_api_keys_routes_response( state, request_context, request_headers, + remote_addr, request_body, ) .await diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/mutation_routes.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/mutation_routes.rs index 23af1b9f4..fc232ed73 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/mutation_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/mutation_routes.rs @@ -5,7 +5,9 @@ use super::shared::{ AdminStandaloneApiKeyToggleRequest, AdminStandaloneApiKeyUpdatePatch, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; -use crate::handlers::admin::shared::attach_admin_audit_response; +use crate::handlers::admin::shared::{ + attach_admin_audit_response, mark_sensitive_admin_response_no_store, +}; use crate::handlers::admin::users::{ default_admin_user_api_key_name, format_optional_unix_secs_iso8601, generate_admin_user_api_key_plaintext, hash_admin_user_api_key, masked_user_api_key_display, @@ -13,7 +15,9 @@ use crate::handlers::admin::users::{ normalize_admin_user_api_formats, normalize_admin_user_ip_rules, normalize_admin_user_string_list, }; -use crate::handlers::shared::normalize_optional_api_key_concurrent_limit; +use crate::handlers::shared::{ + normalize_optional_api_key_concurrent_limit, seal_auth_api_key_secret, +}; use crate::GatewayError; use aether_admin::system::serialize_admin_system_users_export_wallet; use axum::{ @@ -60,6 +64,48 @@ fn normalize_standalone_initial_balance( Ok((initial_balance_usd, false)) } +/// Compensate the rows created by a standalone API-key request when wallet +/// provisioning cannot complete. The wallet is deleted first, and only when +/// its exact API-key owner and untouched state still match. If a wallet has +/// become funded or otherwise referenced, preserve both rows rather than +/// deleting the key and leaving an orphaned financial account. +async fn compensate_standalone_api_key_creation( + state: &AdminAppState<'_>, + api_key_id: &str, +) -> Result<(), GatewayError> { + let wallet = state + .find_wallet(aether_data::repository::wallet::WalletLookupKey::ApiKeyId( + api_key_id, + )) + .await?; + if let Some(wallet) = wallet { + let removed = state + .delete_wallet_if_unreferenced( + &wallet.id, + aether_data::repository::wallet::WalletLookupKey::ApiKeyId(api_key_id), + ) + .await?; + if !removed + && state + .find_wallet(aether_data::repository::wallet::WalletLookupKey::WalletId( + &wallet.id, + )) + .await? + .is_some() + { + return Err(GatewayError::Internal(format!( + "refusing to delete standalone API key {api_key_id}: wallet is still referenced" + ))); + } + } + + // Treat an already-removed key as an idempotent successful cleanup. The + // wallet owner check above prevents deleting a key while its funded wallet + // remains attached. + let _ = state.delete_standalone_api_key(api_key_id).await?; + Ok(()) +} + pub(super) async fn build_admin_create_api_key_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -148,7 +194,16 @@ pub(super) async fn build_admin_create_api_key_response( }; let plaintext_key = generate_admin_user_api_key_plaintext(); - let Some(key_encrypted) = state.encrypt_catalog_secret_with_fallbacks(&plaintext_key) else { + let api_key_id = uuid::Uuid::new_v4().to_string(); + let key_hash = hash_admin_user_api_key(&plaintext_key); + let Ok(key_encrypted) = seal_auth_api_key_secret( + state.app(), + &operator_id, + &api_key_id, + &key_hash, + true, + &plaintext_key, + ) else { return Ok(( http::StatusCode::INTERNAL_SERVER_ERROR, Json(json!({ "detail": "API密钥加密失败" })), @@ -160,8 +215,8 @@ pub(super) async fn build_admin_create_api_key_response( .create_standalone_api_key( aether_data::repository::auth::CreateStandaloneApiKeyRecord { user_id: operator_id, - api_key_id: uuid::Uuid::new_v4().to_string(), - key_hash: hash_admin_user_api_key(&plaintext_key), + api_key_id, + key_hash, key_encrypted: Some(key_encrypted), name: Some(name), allowed_providers, @@ -184,11 +239,39 @@ pub(super) async fn build_admin_create_api_key_response( return Ok(build_admin_api_keys_data_unavailable_response()); }; let wallet = match state - .initialize_auth_api_key_wallet(&created.api_key_id, initial_balance_usd, unlimited_balance) - .await? + .initialize_auth_api_key_wallet_with_outcome( + &created.api_key_id, + initial_balance_usd, + unlimited_balance, + ) + .await { - Some(wallet) => wallet, - None => return Ok(build_admin_api_keys_data_unavailable_response()), + Ok(Some(initialized)) => initialized.wallet, + Ok(None) => { + if let Err(error) = + compensate_standalone_api_key_creation(state, &created.api_key_id).await + { + tracing::error!( + api_key_id = %created.api_key_id, + error = ?error, + "standalone API key wallet provisioning cleanup failed" + ); + return Err(error); + } + return Ok(build_admin_api_keys_data_unavailable_response()); + } + Err(error) => { + if let Err(cleanup_error) = + compensate_standalone_api_key_creation(state, &created.api_key_id).await + { + tracing::error!( + api_key_id = %created.api_key_id, + error = ?cleanup_error, + "standalone API key wallet provisioning cleanup failed" + ); + } + return Err(error); + } }; let created = if feature_settings.is_some() { state @@ -199,30 +282,32 @@ pub(super) async fn build_admin_create_api_key_response( created }; - Ok(attach_admin_audit_response( - Json(json!({ - "id": created.api_key_id, - "key": plaintext_key, - "name": created.name, - "key_display": masked_user_api_key_display(state, created.key_encrypted.as_deref()), - "is_standalone": true, - "is_active": created.is_active, - "rate_limit": created.rate_limit, - "concurrent_limit": created.concurrent_limit, - "allowed_providers": created.allowed_providers, - "allowed_api_formats": created.allowed_api_formats, - "allowed_models": created.allowed_models, - "expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs), - "auto_delete_on_expiry": created.auto_delete_on_expiry, - "feature_settings": created.feature_settings, - "wallet": serialize_admin_system_users_export_wallet(Some(&wallet)), - "message": "独立余额Key创建成功,请妥善保存完整密钥,后续将无法查看", - })) - .into_response(), - "admin_standalone_api_key_created", - "create_standalone_api_key", - "api_key", - &created.api_key_id, + Ok(mark_sensitive_admin_response_no_store( + attach_admin_audit_response( + Json(json!({ + "id": created.api_key_id, + "key": plaintext_key, + "name": created.name, + "key_display": masked_user_api_key_display(state, &created), + "is_standalone": true, + "is_active": created.is_active, + "rate_limit": created.rate_limit, + "concurrent_limit": created.concurrent_limit, + "allowed_providers": created.allowed_providers, + "allowed_api_formats": created.allowed_api_formats, + "allowed_models": created.allowed_models, + "expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs), + "auto_delete_on_expiry": created.auto_delete_on_expiry, + "feature_settings": created.feature_settings, + "wallet": serialize_admin_system_users_export_wallet(Some(&wallet)), + "message": "独立余额Key创建成功,请妥善保存完整密钥,后续将无法查看", + })) + .into_response(), + "admin_standalone_api_key_created", + "create_standalone_api_key", + "api_key", + &created.api_key_id, + ), )) } @@ -405,7 +490,11 @@ pub(super) async fn build_admin_update_api_key_response( .update_standalone_api_key_basic( aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord { api_key_id: api_key_id.clone(), + key_encrypted: None, + key_encrypted_present: false, name, + name_present: field_presence.contains("name"), + force_capabilities: None, rate_limit_present: field_presence.contains("rate_limit"), rate_limit: payload.rate_limit, concurrent_limit_present: field_presence.contains("concurrent_limit"), @@ -540,3 +629,127 @@ pub(super) async fn build_admin_delete_api_key_response( false => Ok(build_admin_api_keys_not_found_response()), } } + +#[cfg(test)] +mod tests { + use super::compensate_standalone_api_key_creation; + use crate::data::GatewayDataState; + use crate::handlers::admin::request::AdminAppState; + use crate::state::AppState; + use aether_data::repository::auth::{ + AuthApiKeyReadRepository, AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, + InMemoryAuthApiKeySnapshotRepository, + }; + use aether_data::repository::wallet::{StoredWalletSnapshot, WalletLookupKey}; + use std::sync::Arc; + + async fn seed_standalone_key( + repository: &InMemoryAuthApiKeySnapshotRepository, + api_key_id: &str, + ) { + repository + .create_standalone_api_key(CreateStandaloneApiKeyRecord { + user_id: "admin-user".to_string(), + api_key_id: api_key_id.to_string(), + key_hash: format!("hash-{api_key_id}"), + key_encrypted: None, + name: Some("test-key".to_string()), + allowed_providers: None, + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + rate_limit: None, + concurrent_limit: None, + force_capabilities: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + }) + .await + .expect("key creation should succeed") + .expect("key should be returned"); + } + + fn wallet_for_key(api_key_id: &str, balance: f64) -> StoredWalletSnapshot { + StoredWalletSnapshot::new( + format!("wallet-{api_key_id}"), + None, + Some(api_key_id.to_string()), + balance, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + if balance > 0.0 { balance } else { 0.0 }, + 0.0, + 0.0, + 0.0, + 1, + ) + .expect("wallet should build") + } + + #[tokio::test] + async fn standalone_key_compensation_removes_unreferenced_wallet_and_key() { + let api_key_id = "compensate-key"; + let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + seed_standalone_key(&repository, api_key_id).await; + let wallet = wallet_for_key(api_key_id, 0.0); + let state = AppState::new() + .expect("state should build") + .with_auth_wallets_for_tests([wallet]) + .with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests( + Arc::clone(&repository), + )); + let admin_state = AdminAppState::new(&state); + + compensate_standalone_api_key_creation(&admin_state, api_key_id) + .await + .expect("compensation should succeed"); + + assert!(state + .find_wallet(WalletLookupKey::ApiKeyId(api_key_id)) + .await + .expect("wallet lookup should succeed") + .is_none()); + assert!(repository + .find_export_standalone_api_key_by_id(api_key_id) + .await + .expect("key lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn standalone_key_compensation_preserves_funded_wallet_and_key() { + let api_key_id = "funded-compensate-key"; + let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + seed_standalone_key(&repository, api_key_id).await; + let wallet = wallet_for_key(api_key_id, 5.0); + let state = AppState::new() + .expect("state should build") + .with_auth_wallets_for_tests([wallet]) + .with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests( + Arc::clone(&repository), + )); + let admin_state = AdminAppState::new(&state); + + assert!( + compensate_standalone_api_key_creation(&admin_state, api_key_id) + .await + .is_err() + ); + assert!(state + .find_wallet(WalletLookupKey::ApiKeyId(api_key_id)) + .await + .expect("wallet lookup should succeed") + .is_some()); + assert!(repository + .find_export_standalone_api_key_by_id(api_key_id) + .await + .expect("key lookup should succeed") + .is_some()); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/read_routes.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/read_routes.rs index b61494cc0..d528d925f 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/read_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/read_routes.rs @@ -4,9 +4,12 @@ use super::shared::{ build_admin_api_keys_bad_request_response, build_admin_api_keys_data_unavailable_response, build_admin_api_keys_not_found_response, }; -use super::{decrypt_catalog_secret_with_fallbacks, query_param_bool, query_param_optional_bool}; +use super::{query_param_bool, query_param_optional_bool}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; -use crate::handlers::admin::shared::attach_admin_audit_response; +use crate::handlers::admin::shared::{ + attach_admin_audit_response, mark_sensitive_admin_response_no_store, +}; +use crate::handlers::shared::decrypt_or_migrate_auth_api_key_secret; use crate::GatewayError; use axum::{ body::Body, @@ -132,20 +135,24 @@ pub(super) async fn build_admin_api_key_detail_response( "该密钥没有存储完整密钥信息", )); }; - let Some(key) = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - else { - return Ok(( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": "解密密钥失败" })), - ) - .into_response()); + let key = match decrypt_or_migrate_auth_api_key_secret(state.app(), &record).await { + Ok(value) => value, + Err(_) => { + return Ok(( + http::StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ "detail": "解密或校验密钥失败" })), + ) + .into_response()) + } }; - return Ok(attach_admin_audit_response( - Json(json!({ "key": key })).into_response(), - "admin_standalone_api_key_revealed", - "reveal_standalone_api_key", - "api_key", - &api_key_id, + return Ok(mark_sensitive_admin_response_no_store( + attach_admin_audit_response( + Json(json!({ "key": key })).into_response(), + "admin_standalone_api_key_revealed", + "reveal_standalone_api_key", + "api_key", + &api_key_id, + ), )); } diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/routes.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/routes.rs index 72e0ebaa6..8991d3208 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/routes.rs @@ -13,6 +13,7 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, request_headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&axum::body::Bytes>, ) -> Result>, GatewayError> { let Some(decision) = request_context.decision() else { @@ -68,6 +69,7 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response( state, request_context, request_headers, + remote_addr, request_body, ) .await?, diff --git a/apps/aether-gateway/src/handlers/admin/auth/api_keys/shared.rs b/apps/aether-gateway/src/handlers/admin/auth/api_keys/shared.rs index 3adb5ce64..1df71c26c 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/api_keys/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/api_keys/shared.rs @@ -146,8 +146,11 @@ pub(super) fn admin_api_keys_parse_limit(query: Option<&str>) -> Result, ciphertext: Option<&str>) -> String { - masked_user_api_key_display(state, ciphertext) +fn masked_admin_api_key_display( + state: &AdminAppState<'_>, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, +) -> String { + masked_user_api_key_display(state, record) } pub(super) fn build_admin_api_key_list_item_payload( @@ -159,7 +162,7 @@ pub(super) fn build_admin_api_key_list_item_payload( "id": record.api_key_id, "user_id": record.user_id, "name": record.name, - "key_display": masked_admin_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": masked_admin_api_key_display(state, record), "is_active": record.is_active, "is_standalone": true, "total_requests": record.total_requests, @@ -190,7 +193,7 @@ pub(super) fn build_admin_api_key_detail_payload( "id": record.api_key_id, "user_id": record.user_id, "name": record.name, - "key_display": masked_admin_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": masked_admin_api_key_display(state, record), "is_active": record.is_active, "is_standalone": true, "total_requests": record.total_requests, diff --git a/apps/aether-gateway/src/handlers/admin/auth/ldap/builders.rs b/apps/aether-gateway/src/handlers/admin/auth/ldap/builders.rs index 47e130508..5739d125e 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/ldap/builders.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/ldap/builders.rs @@ -1,9 +1,10 @@ use super::shared::*; use crate::handlers::admin::request::AdminAppState; use crate::GatewayError; +use aether_data::repository::auth_modules::{LdapBindPasswordUpdate, StoredLdapModuleConfig}; use serde::Deserialize; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(super) struct AdminLdapConfigUpdateRequest { server_url: String, bind_dn: String, @@ -28,7 +29,7 @@ pub(super) struct AdminLdapConfigUpdateRequest { connect_timeout: i32, } -#[derive(Debug, Default, Deserialize)] +#[derive(Default, Deserialize)] pub(super) struct AdminLdapConfigTestRequest { #[serde(default)] server_url: Option, @@ -56,7 +57,7 @@ pub(super) struct AdminLdapConfigTestRequest { connect_timeout: Option, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub(super) struct AdminLdapConnectionTestConfig { server_url: String, bind_dn: String, @@ -66,13 +67,26 @@ pub(super) struct AdminLdapConnectionTestConfig { connect_timeout: i32, } +pub(super) struct AdminLdapConfigUpdate { + pub(super) expected: Option, + pub(super) replacement: StoredLdapModuleConfig, + pub(super) bind_password_update: LdapBindPasswordUpdate, +} + pub(super) async fn build_admin_ldap_update_config( state: &AdminAppState<'_>, payload: AdminLdapConfigUpdateRequest, -) -> Result { +) -> Result { let server_url = admin_ldap_trim_required(payload.server_url, "LDAP 服务器地址不能为空")?; + let server_url = admin_ldap_normalize_server_url(&server_url, payload.use_starttls) + .ok_or_else(|| { + "LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段" + .to_string() + })?; let bind_dn = admin_ldap_trim_required(payload.bind_dn, "绑定 DN 不能为空")?; let base_dn = admin_ldap_trim_required(payload.base_dn, "Base DN 不能为空")?; + admin_ldap_validate_distinguished_name(&bind_dn, "绑定 DN")?; + admin_ldap_validate_distinguished_name(&base_dn, "Base DN")?; let user_search_filter = admin_ldap_trim_required(payload.user_search_filter, "搜索过滤器不能为空")?; admin_ldap_validate_search_filter(&user_search_filter)?; @@ -80,39 +94,42 @@ pub(super) async fn build_admin_ldap_update_config( let email_attr = admin_ldap_trim_required(payload.email_attr, "邮箱属性不能为空")?; let display_name_attr = admin_ldap_trim_required(payload.display_name_attr, "显示名称属性不能为空")?; + admin_ldap_validate_attribute_description(&username_attr, "用户名属性")?; + admin_ldap_validate_attribute_description(&email_attr, "邮箱属性")?; + admin_ldap_validate_attribute_description(&display_name_attr, "显示名称属性")?; if !(1..=60).contains(&payload.connect_timeout) { return Err("连接超时时间必须在 1 到 60 秒之间".to_string()); } - let existing = state + let mut existing = state .get_ldap_module_config() .await .map_err(|err| format!("{err:?}"))?; - let bind_password_update_requested = payload - .bind_password - .as_ref() - .is_some_and(|value| !value.is_empty()); - let bind_password = match payload.bind_password { - Some(value) if value.is_empty() => Some(String::new()), - Some(value) => Some(admin_ldap_trim_required(value, "绑定密码不能为空")?), - None => None, - }; + if payload.bind_password.is_none() { + if let Some(config) = existing.as_ref() { + crate::handlers::shared::decrypt_or_migrate_ldap_bind_password(state.app(), config) + .await + .map_err(|_| "已保存的 LDAP 绑定密码无法解密".to_string())?; + existing = state + .get_ldap_module_config() + .await + .map_err(|err| format!("{err:?}"))?; + } + } + let requested_bind_password = payload.bind_password; let is_new_config = existing.is_none(); - if is_new_config && bind_password.as_deref().unwrap_or("").is_empty() { + let will_have_password = match requested_bind_password.as_deref() { + Some(value) => !value.trim().is_empty(), + None => existing + .as_ref() + .and_then(|config| config.bind_password_encrypted.as_deref()) + .map(str::trim) + .is_some_and(|value: &str| !value.is_empty()), + }; + if is_new_config && !will_have_password { return Err("首次配置 LDAP 时必须设置绑定密码".to_string()); } - let will_have_password = bind_password - .as_ref() - .map(|value| !value.is_empty()) - .unwrap_or_else(|| { - existing - .as_ref() - .and_then(|config| config.bind_password_encrypted.as_deref()) - .map(str::trim) - .is_some_and(|value: &str| !value.is_empty()) - }); - if payload.is_exclusive && !payload.is_enabled { return Err("仅允许 LDAP 登录 需要先启用 LDAP 认证".to_string()); } @@ -135,31 +152,58 @@ pub(super) async fn build_admin_ldap_update_config( } } - let bind_password_encrypted = match bind_password { - Some(value) if value.is_empty() => None, - Some(value) => state.encrypt_catalog_secret_with_fallbacks(&value), - None => existing.and_then(|config| config.bind_password_encrypted), + let replacement = StoredLdapModuleConfig { + server_url, + bind_dn, + // The repository ignores this field for config mutations. Password changes are + // carried exclusively by `bind_password_update`, so Preserve never copies an old + // ciphertext into the replacement record. + bind_password_encrypted: None, + base_dn, + user_search_filter: Some(user_search_filter), + username_attr: Some(username_attr), + email_attr: Some(email_attr), + display_name_attr: Some(display_name_attr), + is_enabled: payload.is_enabled, + is_exclusive: payload.is_exclusive, + use_starttls: payload.use_starttls, + connect_timeout: Some(payload.connect_timeout), }; - if bind_password_update_requested && bind_password_encrypted.is_none() { - return Err("LDAP 绑定密码加密失败,请检查 Rust 数据加密配置".to_string()); - } - - Ok( - aether_data::repository::auth_modules::StoredLdapModuleConfig { - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter: Some(user_search_filter), - username_attr: Some(username_attr), - email_attr: Some(email_attr), - display_name_attr: Some(display_name_attr), - is_enabled: payload.is_enabled, - is_exclusive: payload.is_exclusive, - use_starttls: payload.use_starttls, - connect_timeout: Some(payload.connect_timeout), - }, - ) + let bind_password_update = match requested_bind_password { + Some(value) if value.is_empty() => LdapBindPasswordUpdate::Clear, + Some(value) => { + let password = admin_ldap_trim_required(value, "绑定密码不能为空")?; + let ciphertext = state + .encrypt_ldap_bind_password(&replacement, &password) + .ok_or_else(|| "LDAP 绑定密码加密失败,请检查 Rust 数据加密配置".to_string())?; + LdapBindPasswordUpdate::Set(ciphertext) + } + None => { + if let Some(existing_config) = existing.as_ref() { + if existing_config + .bind_password_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + && !crate::handlers::shared::ldap_bind_password_binding_matches( + existing_config, + &replacement, + ) + .unwrap_or(false) + { + return Err( + "修改 LDAP 服务器、StartTLS、bind DN 或 Base DN 时必须重新提供绑定密码" + .to_string(), + ); + } + } + LdapBindPasswordUpdate::Preserve + } + }; + Ok(AdminLdapConfigUpdate { + expected: existing, + replacement, + bind_password_update, + }) } pub(super) async fn build_admin_ldap_test_config( @@ -167,7 +211,16 @@ pub(super) async fn build_admin_ldap_test_config( payload: AdminLdapConfigTestRequest, ) -> Result, String> { if let Some(value) = payload.user_search_filter.as_deref() { - admin_ldap_validate_search_filter(value.trim())?; + admin_ldap_validate_search_filter(value)?; + } + for (value, label) in [ + (payload.username_attr.as_deref(), "用户名属性"), + (payload.email_attr.as_deref(), "邮箱属性"), + (payload.display_name_attr.as_deref(), "显示名称属性"), + ] { + if let Some(value) = value { + admin_ldap_validate_attribute_description(value.trim(), label)?; + } } if let Some(connect_timeout) = payload.connect_timeout { if !(1..=60).contains(&connect_timeout) { @@ -199,9 +252,14 @@ pub(super) async fn build_admin_ldap_test_config( .as_ref() .and_then(|config| config.connect_timeout) .unwrap_or_else(admin_ldap_default_connect_timeout); - let mut bind_password = saved - .as_ref() - .and_then(|config| admin_ldap_read_saved_bind_password(state, config)); + let mut bind_password = match saved.as_ref() { + Some(config) => { + crate::handlers::shared::decrypt_or_migrate_ldap_bind_password(state.app(), config) + .await + .map_err(|_| "已保存的 LDAP 绑定密码无法解密".to_string())? + } + None => None, + }; if let Some(value) = payload.server_url { server_url = Some(admin_ldap_trim_required(value, "LDAP 服务器地址不能为空")?); @@ -239,11 +297,21 @@ pub(super) async fn build_admin_ldap_test_config( return Ok(None); } + let server_url = server_url.expect("server_url already checked"); + let server_url = admin_ldap_normalize_server_url(&server_url, use_starttls).ok_or_else(|| { + "LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段" + .to_string() + })?; + let bind_dn = bind_dn.expect("bind_dn already checked"); + let base_dn = base_dn.expect("base_dn already checked"); + admin_ldap_validate_distinguished_name(&bind_dn, "绑定 DN")?; + admin_ldap_validate_distinguished_name(&base_dn, "Base DN")?; + Ok(Some(AdminLdapConnectionTestConfig { - server_url: server_url.expect("server_url already checked"), - bind_dn: bind_dn.expect("bind_dn already checked"), + server_url, + bind_dn, bind_password: bind_password.expect("bind_password already checked"), - base_dn: base_dn.expect("base_dn already checked"), + base_dn, use_starttls, connect_timeout, })) @@ -270,7 +338,8 @@ pub(super) async fn admin_ldap_test_connection( } fn admin_ldap_test_connection_blocking(config: AdminLdapConnectionTestConfig) -> (bool, String) { - let Some(server_url): Option = admin_ldap_normalize_server_url(&config.server_url) + let Some(server_url): Option = + admin_ldap_normalize_server_url(&config.server_url, config.use_starttls) else { return (false, ADMIN_LDAP_TEST_FAILURE_MESSAGE.to_string()); }; @@ -302,51 +371,21 @@ fn admin_ldap_trim_required(value: String, detail: &str) -> Result Result<(), String> { - if value.is_empty() { - return Err("搜索过滤器不能为空".to_string()); - } - if !value.contains("{username}") { - return Err("搜索过滤器必须包含 {username} 占位符".to_string()); - } - - let mut depth = 0i32; - let mut max_depth = 0i32; - for ch in value.chars() { - if ch == '(' { - depth += 1; - max_depth = max_depth.max(depth); - } else if ch == ')' { - depth -= 1; - if depth < 0 { - return Err("搜索过滤器括号不匹配".to_string()); - } - } - } - if depth != 0 { - return Err("搜索过滤器括号不匹配".to_string()); - } - if max_depth > 5 { - return Err("搜索过滤器嵌套层数过深(最多5层)".to_string()); - } - if value.len() > 200 { - return Err("搜索过滤器过长(最多200字符)".to_string()); - } - Ok(()) -} - -fn admin_ldap_read_saved_bind_password( - state: &AdminAppState<'_>, - config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, -) -> Option { - config - .bind_password_encrypted - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(|value| { - state - .decrypt_catalog_secret_with_fallbacks(value) - .or_else(|| Some(value.to_string())) + crate::handlers::shared::ldap_search_filter_is_valid(value) + .then_some(()) + .ok_or_else(|| { + "搜索过滤器格式无效,必须包含 {username} 且使用唯一、有限的外层括号结构".to_string() }) - .filter(|value| !value.trim().is_empty()) +} + +fn admin_ldap_validate_distinguished_name(value: &str, label: &str) -> Result<(), String> { + crate::handlers::shared::ldap_distinguished_name_is_valid(value) + .then_some(()) + .ok_or_else(|| format!("{label}格式无效或过长")) +} + +fn admin_ldap_validate_attribute_description(value: &str, label: &str) -> Result<(), String> { + crate::handlers::shared::ldap_attribute_description_is_valid(value) + .then_some(()) + .ok_or_else(|| format!("{label}必须是有效的 LDAP 属性名称")) } diff --git a/apps/aether-gateway/src/handlers/admin/auth/ldap/routes.rs b/apps/aether-gateway/src/handlers/admin/auth/ldap/routes.rs index 962f52a33..dbc6233b9 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/ldap/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/ldap/routes.rs @@ -6,6 +6,7 @@ use super::shared::*; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::attach_admin_audit_response; use crate::GatewayError; +use aether_data::repository::auth_modules::CompareAndSwapLdapConfigResult; use axum::{ body::{Body, Bytes}, http, @@ -61,9 +62,20 @@ pub(super) async fn maybe_build_local_admin_ldap_response( Ok(config) => config, Err(detail) => return Ok(Some(admin_ldap_bad_request_response(detail))), }; - let saved = state.upsert_ldap_module_config(&update).await?; - if saved.is_none() { + let saved = state + .compare_and_swap_ldap_module_config( + update.expected.as_ref(), + &update.replacement, + &update.bind_password_update, + ) + .await?; + let Some(saved) = saved else { return Ok(Some(admin_ldap_unavailable_response())); + }; + if saved == CompareAndSwapLdapConfigResult::Conflict { + return Ok(Some(admin_ldap_conflict_response( + "LDAP 配置已被其他请求更新,请重新加载后重试", + ))); } return Ok(Some( Json(json!({ "message": "LDAP配置更新成功" })).into_response(), diff --git a/apps/aether-gateway/src/handlers/admin/auth/ldap/shared.rs b/apps/aether-gateway/src/handlers/admin/auth/ldap/shared.rs index a5c6e5c0b..50073583c 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/ldap/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/ldap/shared.rs @@ -102,6 +102,14 @@ pub(super) fn admin_ldap_bad_request_response(detail: impl Into) -> Resp .into_response() } +pub(super) fn admin_ldap_conflict_response(detail: impl Into) -> Response { + ( + http::StatusCode::CONFLICT, + Json(json!({ "detail": detail.into() })), + ) + .into_response() +} + pub(super) fn admin_ldap_unavailable_response() -> Response { ( http::StatusCode::SERVICE_UNAVAILABLE, @@ -110,13 +118,9 @@ pub(super) fn admin_ldap_unavailable_response() -> Response { .into_response() } -pub(super) fn admin_ldap_normalize_server_url(server_url: &str) -> Option { - let server_url = server_url.trim(); - if server_url.is_empty() { - return None; - } - if server_url.contains("://") { - return Some(server_url.to_string()); - } - Some(format!("ldap://{server_url}")) +pub(super) fn admin_ldap_normalize_server_url( + server_url: &str, + use_starttls: bool, +) -> Option { + crate::handlers::shared::normalize_ldap_transport_server_url(server_url, use_starttls) } diff --git a/apps/aether-gateway/src/handlers/admin/auth/oauth_config.rs b/apps/aether-gateway/src/handlers/admin/auth/oauth_config.rs index b1a78a060..f1cb17fd4 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/oauth_config.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/oauth_config.rs @@ -1,13 +1,14 @@ use crate::handlers::admin::request::AdminAppState; use aether_data::repository::oauth_providers::{ - EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord, + validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config, + validate_oauth_redirect_uri, EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord, }; use axum::http; use serde::Deserialize; use serde_json::json; -use url::Url; +use url::{Host, Url}; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(crate) struct AdminOAuthProviderUpsertRequest { pub(super) display_name: String, pub(super) client_id: String, @@ -122,7 +123,9 @@ pub(super) fn admin_oauth_is_supported_provider(provider_type: &str) -> bool { }) } -fn admin_oauth_builtin_allowed_domains(provider_type: &str) -> Option<&'static [&'static str]> { +pub(super) fn admin_oauth_builtin_allowed_domains( + provider_type: &str, +) -> Option<&'static [&'static str]> { if provider_type.eq_ignore_ascii_case("linuxdo") { Some(&["linux.do", "connect.linux.do", "connect.linuxdo.org"]) } else { @@ -130,7 +133,9 @@ fn admin_oauth_builtin_allowed_domains(provider_type: &str) -> Option<&'static [ } } -fn admin_oauth_custom_allowed_domains(extra_config: Option<&serde_json::Value>) -> Vec { +pub(super) fn admin_oauth_custom_allowed_domains( + extra_config: Option<&serde_json::Value>, +) -> Vec { extra_config .and_then(serde_json::Value::as_object) .and_then(|object| { @@ -146,42 +151,57 @@ fn admin_oauth_custom_allowed_domains(extra_config: Option<&serde_json::Value>) .map(str::trim) .filter(|value| !value.is_empty()) .map(|value| value.trim_end_matches('.').to_ascii_lowercase()) + .filter(|value| value.parse::().is_err()) .collect::>() }) .unwrap_or_default() } -fn validate_admin_oauth_frontend_callback_url(url: &str) -> Result<(), String> { - let parsed = Url::parse(url).map_err(|_| "frontend_callback_url 必须是绝对 URL".to_string())?; - if !matches!(parsed.scheme(), "http" | "https") { - return Err("frontend_callback_url scheme 必须是 http/https".to_string()); - } - if parsed.host_str().is_none() { - return Err("frontend_callback_url 必须是绝对 URL".to_string()); - } - let path = parsed.path().trim_end_matches('/'); - if !path.ends_with("/auth/callback") { - return Err("frontend_callback_url 路径必须以 /auth/callback 结尾".to_string()); - } - Ok(()) +pub(super) fn validate_admin_oauth_url_override( + url: &str, + allowed_domains: &[&str], +) -> Result<(), String> { + validate_admin_oauth_url_override_with_options(url, allowed_domains, false) } -fn validate_admin_oauth_redirect_uri(url: &str) -> Result<(), String> { - let parsed = Url::parse(url).map_err(|_| "redirect_uri 必须是绝对 URL".to_string())?; - if !matches!(parsed.scheme(), "http" | "https") { - return Err("redirect_uri scheme 必须是 http/https".to_string()); - } - if parsed.host_str().is_none() { - return Err("redirect_uri 必须是绝对 URL".to_string()); - } - Ok(()) +fn validate_admin_oauth_authorization_url_override( + url: &str, + allowed_domains: &[&str], +) -> Result<(), String> { + validate_admin_oauth_url_override_with_options(url, allowed_domains, true) } -fn validate_admin_oauth_url_override(url: &str, allowed_domains: &[&str]) -> Result<(), String> { +fn validate_admin_oauth_url_override_with_options( + url: &str, + allowed_domains: &[&str], + reject_authorization_parameters: bool, +) -> Result<(), String> { let parsed = Url::parse(url).map_err(|_| "端点覆盖必须是 https 绝对 URL".to_string())?; if parsed.scheme() != "https" || parsed.host_str().is_none() { return Err("端点覆盖必须是 https 绝对 URL".to_string()); } + if matches!(parsed.host(), Some(Host::Ipv4(_)) | Some(Host::Ipv6(_))) { + return Err("端点覆盖必须使用 DNS 主机名,不能使用 IP 字面量".to_string()); + } + if !parsed.username().is_empty() || parsed.password().is_some() || parsed.fragment().is_some() { + return Err("端点覆盖不得包含 URL 凭据或 fragment".to_string()); + } + if reject_authorization_parameters + && parsed.query_pairs().any(|(name, _)| { + matches!( + name.to_ascii_lowercase().as_str(), + "response_type" + | "client_id" + | "redirect_uri" + | "state" + | "scope" + | "code_challenge" + | "code_challenge_method" + ) + }) + { + return Err("authorization endpoint 不得预置 OAuth authorization 参数".to_string()); + } let host = parsed .host_str() .map(|value| value.trim().trim_end_matches('.').to_ascii_lowercase()) @@ -207,6 +227,17 @@ fn validate_admin_oauth_url_override_for_domains( validate_admin_oauth_url_override(url, &allowed) } +fn validate_admin_oauth_authorization_url_override_for_domains( + url: &str, + allowed_domains: &[String], +) -> Result<(), String> { + let allowed = allowed_domains + .iter() + .map(String::as_str) + .collect::>(); + validate_admin_oauth_authorization_url_override(url, &allowed) +} + pub(super) fn build_admin_oauth_upsert_record( state: &AdminAppState<'_>, provider_type: &str, @@ -236,8 +267,8 @@ pub(super) fn build_admin_oauth_upsert_record( return Err("frontend_callback_url 不能为空".to_string()); } - validate_admin_oauth_frontend_callback_url(frontend_callback_url)?; - validate_admin_oauth_redirect_uri(redirect_uri)?; + validate_oauth_frontend_callback_url(frontend_callback_url)?; + validate_oauth_redirect_uri(redirect_uri)?; let is_custom_oidc = admin_oauth_is_custom_provider_type(&provider_type); let custom_allowed_domains = if is_custom_oidc { @@ -268,16 +299,26 @@ pub(super) fn build_admin_oauth_upsert_record( let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else { return Err(format!("custom_oidc 必须配置 {field_name}")); }; - validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?; + if field_name == "authorization_url_override" { + validate_admin_oauth_authorization_url_override_for_domains( + value, + &custom_allowed_domains, + )?; + } else { + validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?; + } } } if let Some(value) = payload.authorization_url_override.as_deref().map(str::trim) { if !value.is_empty() { if let Some(allowed_domains) = builtin_allowed_domains { - validate_admin_oauth_url_override(value, allowed_domains)?; + validate_admin_oauth_authorization_url_override(value, allowed_domains)?; } else { - validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?; + validate_admin_oauth_authorization_url_override_for_domains( + value, + &custom_allowed_domains, + )?; } } } @@ -322,40 +363,34 @@ pub(super) fn build_admin_oauth_upsert_record( return Err("scopes 不能为空".to_string()); } - let client_secret_encrypted = match payload.client_secret.as_deref() { - None => EncryptedSecretUpdate::Preserve, - Some(raw) => { - let secret = raw.trim(); - if secret == "__CLEAR__" { - EncryptedSecretUpdate::Clear - } else if secret.is_empty() { - EncryptedSecretUpdate::Preserve - } else { - let encrypted = state - .encrypt_catalog_secret_with_fallbacks(secret) - .ok_or_else(|| "gateway 未配置 OAuth provider 加密密钥".to_string())?; - EncryptedSecretUpdate::Set(encrypted) - } - } - }; + validate_oauth_provider_endpoint_config( + &provider_type, + payload.authorization_url_override.as_deref(), + payload.token_url_override.as_deref(), + payload.userinfo_url_override.as_deref(), + payload.extra_config.as_ref(), + )?; - Ok(UpsertOAuthProviderConfigRecord { + let authorization_url_override = payload.authorization_url_override.and_then(|value| { + let value = value.trim().to_string(); + (!value.is_empty()).then_some(value) + }); + let token_url_override = payload.token_url_override.and_then(|value| { + let value = value.trim().to_string(); + (!value.is_empty()).then_some(value) + }); + let userinfo_url_override = payload.userinfo_url_override.and_then(|value| { + let value = value.trim().to_string(); + (!value.is_empty()).then_some(value) + }); + let mut record = UpsertOAuthProviderConfigRecord { provider_type, display_name: display_name.to_string(), client_id: client_id.to_string(), - client_secret_encrypted, - authorization_url_override: payload.authorization_url_override.and_then(|value| { - let value = value.trim().to_string(); - (!value.is_empty()).then_some(value) - }), - token_url_override: payload.token_url_override.and_then(|value| { - let value = value.trim().to_string(); - (!value.is_empty()).then_some(value) - }), - userinfo_url_override: payload.userinfo_url_override.and_then(|value| { - let value = value.trim().to_string(); - (!value.is_empty()).then_some(value) - }), + client_secret_encrypted: EncryptedSecretUpdate::Preserve, + authorization_url_override, + token_url_override, + userinfo_url_override, scopes: payload.scopes.map(|items| { items .into_iter() @@ -372,5 +407,27 @@ pub(super) fn build_admin_oauth_upsert_record( (!value.is_empty()).then_some(value) }), is_enabled: payload.is_enabled, - }) + }; + record.client_secret_encrypted = match payload.client_secret.as_deref() { + None => EncryptedSecretUpdate::Preserve, + Some(raw) => { + let secret = raw.trim(); + if secret == "__CLEAR__" { + EncryptedSecretUpdate::Clear + } else if secret.is_empty() { + EncryptedSecretUpdate::Preserve + } else { + let encrypted = + crate::handlers::shared::seal_identity_oauth_provider_client_secret( + state.as_ref(), + &record, + secret, + ) + .map_err(str::to_string)?; + EncryptedSecretUpdate::Set(encrypted) + } + } + }; + + Ok(record) } diff --git a/apps/aether-gateway/src/handlers/admin/auth/oauth_routes.rs b/apps/aether-gateway/src/handlers/admin/auth/oauth_routes.rs index f612ec2bc..759f54f9f 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/oauth_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/oauth_routes.rs @@ -1,8 +1,9 @@ use super::oauth_config::{ + admin_oauth_builtin_allowed_domains, admin_oauth_custom_allowed_domains, admin_oauth_is_supported_provider, admin_oauth_provider_type_from_path, admin_oauth_test_provider_type_from_path, build_admin_oauth_provider_payload, build_admin_oauth_supported_types_payload, build_admin_oauth_upsert_record, - AdminOAuthProviderUpsertRequest, + validate_admin_oauth_url_override, AdminOAuthProviderUpsertRequest, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, build_proxy_error_response}; @@ -14,9 +15,11 @@ use axum::{ Json, }; use serde_json::json; +use std::net::{IpAddr, SocketAddr}; use std::time::Duration; const ADMIN_OAUTH_TEST_TIMEOUT_SECS: u64 = 10; +const ADMIN_OAUTH_TEST_MAX_REDIRECTS: usize = 3; const LINUXDO_AUTHORIZATION_URL: &str = "https://connect.linux.do/oauth2/authorize"; const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token"; @@ -36,27 +39,204 @@ fn admin_oauth_secret_status(has_secret: bool) -> &'static str { } } -async fn admin_oauth_endpoint_reachable(client: &reqwest::Client, url: &str) -> bool { - let Ok(parsed) = reqwest::Url::parse(url) else { +async fn admin_oauth_endpoint_reachable( + url: &str, + allowed_domains: &[&str], + allow_benchmarking_ip: bool, +) -> bool { + let Ok(mut current) = reqwest::Url::parse(url) else { return false; }; - if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() { + for redirects in 0..=ADMIN_OAUTH_TEST_MAX_REDIRECTS { + if validate_admin_oauth_url_override(current.as_str(), allowed_domains).is_err() { + return false; + } + let Ok((host, addrs)) = + resolve_public_admin_oauth_endpoint_with_policy(¤t, allow_benchmarking_ip).await + else { + return false; + }; + let mut builder = reqwest::Client::builder() + .timeout(Duration::from_secs(ADMIN_OAUTH_TEST_TIMEOUT_SECS)) + .redirect(reqwest::redirect::Policy::none()) + .no_proxy(); + if host.parse::().is_err() { + builder = builder.resolve_to_addrs(host.as_str(), &addrs); + } + let Ok(client) = builder.build() else { + return false; + }; + let Ok(response) = client + .get(current.clone()) + .header(reqwest::header::ACCEPT, "*/*") + .header( + reqwest::header::USER_AGENT, + "Aether OAuth configuration tester", + ) + .send() + .await + else { + return false; + }; + if !response.status().is_redirection() { + return response.status().as_u16() < 500; + } + if redirects == ADMIN_OAUTH_TEST_MAX_REDIRECTS { + return false; + } + let Some(location) = response + .headers() + .get(reqwest::header::LOCATION) + .and_then(|value| value.to_str().ok()) + else { + return false; + }; + let Ok(next) = current.join(location) else { + return false; + }; + current = next; + } + false +} + +async fn resolve_public_admin_oauth_endpoint( + url: &reqwest::Url, +) -> Result<(String, Vec), ()> { + resolve_public_admin_oauth_endpoint_with_policy(url, false).await +} + +async fn resolve_public_admin_oauth_endpoint_with_policy( + url: &reqwest::Url, + allow_benchmarking_ip: bool, +) -> Result<(String, Vec), ()> { + if url.scheme() != "https" + || !url.username().is_empty() + || url.password().is_some() + || url.host_str().is_none() + { + return Err(()); + } + let host = url.host_str().ok_or(())?; + let port = url.port_or_known_default().ok_or(())?; + let addrs = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .map_err(|_| ())? + }; + if validate_public_admin_oauth_resolved_addrs(url, &addrs, allow_benchmarking_ip).is_err() { + return Err(()); + } + Ok((host.to_string(), addrs)) +} + +fn validate_public_admin_oauth_resolved_addrs( + url: &reqwest::Url, + addrs: &[SocketAddr], + allow_benchmarking_ip: bool, +) -> Result<(), ()> { + if addrs.is_empty() + || addrs.iter().any(|addr| { + aether_http::is_private_or_reserved_ip(addr.ip()) + && !(allow_benchmarking_ip + && is_fixed_linuxdo_oauth_origin(url) + && aether_http::is_ipv4_benchmarking_fake_ip(addr.ip())) + }) + { + return Err(()); + } + Ok(()) +} + +fn is_fixed_linuxdo_oauth_origin(url: &reqwest::Url) -> bool { + url.scheme() == "https" + && url.host_str().is_some_and(|host| { + host.trim_end_matches('.') + .eq_ignore_ascii_case("connect.linux.do") + }) + && url.port_or_known_default() == Some(443) + && url.username().is_empty() + && url.password().is_none() + && url.query().is_none() + && url.fragment().is_none() +} + +fn admin_oauth_test_allowed_domains( + provider_type: &str, + payload: &serde_json::Value, + persisted_config: Option<&aether_data::repository::oauth_providers::StoredOAuthProviderConfig>, +) -> Vec { + if let Some(domains) = admin_oauth_builtin_allowed_domains(provider_type) { + return domains.iter().map(|domain| (*domain).to_string()).collect(); + } + let payload_extra = payload.get("extra_config"); + let domains = admin_oauth_custom_allowed_domains(payload_extra); + if domains.is_empty() { + admin_oauth_custom_allowed_domains( + persisted_config.and_then(|provider| provider.extra_config.as_ref()), + ) + } else { + domains + } +} + +fn management_token_may_configure_frontend_callback( + request_context: &AdminRequestContext<'_>, + existing: Option<&aether_data::repository::oauth_providers::StoredOAuthProviderConfig>, + requested_callback: &str, +) -> bool { + let Some(principal) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + else { return false; + }; + if principal.management_token_id.is_none() { + return true; + } + let callback_changed = + existing.is_none_or(|provider| provider.frontend_callback_url != requested_callback.trim()); + if !callback_changed { + return true; } - match client - .get(parsed) - .header(reqwest::header::ACCEPT, "*/*") - .header( - reqwest::header::USER_AGENT, - "Aether OAuth configuration tester", + // A missing permission list is the legacy full-access token representation. + principal + .management_token_permissions + .as_ref() + .is_none_or(|permissions| { + permissions + .iter() + .any(|permission| permission == "admin:oauth:admin") + }) +} + +fn oauth_frontend_callback_permission_denied_response( + request_context: &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"); + attach_admin_audit_response( + ( + http::StatusCode::FORBIDDEN, + Json(json!({ + "detail": "management token permission denied", + "required_permission": "admin:oauth:admin", + "route_family": "oauth_manage", + "route_kind": "upsert_provider", + "request_path": request_context.path(), + })), ) - .send() - .await - { - Ok(response) => response.status().as_u16() < 500, - Err(_) => false, - } + .into_response(), + "admin_oauth_frontend_callback_permission_denied", + "permission_denied", + "oauth_frontend_callback", + actor_id, + ) } async fn build_admin_oauth_test_payload( @@ -117,34 +297,38 @@ async fn build_admin_oauth_test_payload( })); }; - let proxy_snapshot = state.app().resolve_system_proxy_snapshot().await; - let mut client_builder = reqwest::Client::builder() - .timeout(Duration::from_secs(ADMIN_OAUTH_TEST_TIMEOUT_SECS)) - .redirect(reqwest::redirect::Policy::limited(3)); - if let Some(proxy_url) = proxy_snapshot.as_ref().and_then(|p| p.url.as_deref()) { - if let Ok(proxy) = reqwest::Proxy::all(proxy_url) { - client_builder = client_builder.proxy(proxy); - } - } - let client = client_builder.build(); - let Ok(client) = client else { + let allowed_domains = + admin_oauth_test_allowed_domains(provider_type, payload, persisted_config.as_ref()); + let allowed_domain_refs = allowed_domains + .iter() + .map(String::as_str) + .collect::>(); + if allowed_domain_refs.is_empty() + || validate_admin_oauth_url_override(&authorization_url, &allowed_domain_refs).is_err() + || validate_admin_oauth_url_override(&token_url, &allowed_domain_refs).is_err() + { return Ok(json!({ "authorization_url_reachable": false, "token_url_reachable": false, "secret_status": admin_oauth_secret_status(has_secret), - "details": "OAuth 配置测试 HTTP client 初始化失败", + "details": "OAuth 端点必须使用 https 且位于 provider 域名白名单中", })); - }; + } + let allow_benchmarking_ip = provider_type.eq_ignore_ascii_case("linuxdo"); let (authorization_url_reachable, token_url_reachable) = tokio::join!( - admin_oauth_endpoint_reachable(&client, &authorization_url), - admin_oauth_endpoint_reachable(&client, &token_url), + admin_oauth_endpoint_reachable( + &authorization_url, + &allowed_domain_refs, + allow_benchmarking_ip, + ), + admin_oauth_endpoint_reachable(&token_url, &allowed_domain_refs, allow_benchmarking_ip), ); let details = if authorization_url_reachable && token_url_reachable { "OAuth 端点可达;client_secret 仅在授权回调兑换 code 时校验" } else { - "OAuth 端点不可达或返回不可用状态;请检查端点 URL、网络和代理配置" + "OAuth 端点不可达或返回不可用状态;请检查端点 URL 和网络配置" }; Ok(json!({ @@ -155,6 +339,57 @@ async fn build_admin_oauth_test_payload( })) } +#[cfg(test)] +mod tests { + use super::{ + is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint, + validate_public_admin_oauth_resolved_addrs, + }; + use std::net::SocketAddr; + + #[tokio::test] + async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() { + let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL"); + + assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err()); + } + + #[test] + fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() { + let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token") + .expect("LinuxDo URL should parse"); + let fake = SocketAddr::from(([198, 18, 75, 234], 443)); + assert!(is_fixed_linuxdo_oauth_origin(&fixed)); + assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok()); + assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err()); + assert!(validate_public_admin_oauth_resolved_addrs( + &fixed, + &[fake, SocketAddr::from(([127, 0, 0, 1], 443))], + true, + ) + .is_err()); + } + + #[test] + fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() { + let fake = SocketAddr::from(([198, 18, 75, 234], 443)); + for raw_url in [ + "https://oauth.example.test/token", + "https://connect.linux.do:8443/oauth2/token", + "https://connect.linuxdo.org/oauth2/token", + "https://connect.linux.do.evil.test/oauth2/token", + "https://connect.linux.do/oauth2/token?tenant=unexpected", + ] { + let url = reqwest::Url::parse(raw_url).expect("test URL should parse"); + assert!( + !is_fixed_linuxdo_oauth_origin(&url), + "must not trust {raw_url}" + ); + assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err()); + } + } +} + pub(crate) async fn maybe_build_local_admin_oauth_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -263,34 +498,16 @@ pub(crate) async fn maybe_build_local_admin_oauth_response( } }; let existing = state.get_oauth_provider_config(&provider_type).await?; - let ldap_exclusive = state.get_ldap_module_config().await?.is_some_and(|config| { - config.is_enabled - && config.is_exclusive - && config - .bind_password_encrypted - .as_deref() - .map(str::trim) - .is_some_and(|value| !value.is_empty()) - }); - if existing - .as_ref() - .is_some_and(|provider| provider.is_enabled && !payload.is_enabled) - { - let affected_count = state - .count_locked_users_if_oauth_provider_disabled(&provider_type, ldap_exclusive) - .await?; - if affected_count > 0 && !payload.force { - return Ok(Some(build_proxy_error_response( - http::StatusCode::CONFLICT, - "confirmation_required", - format!("禁用该 Provider 会导致 {affected_count} 个用户无法登录"), - Some(json!({ - "affected_count": affected_count, - "action": "disable_oauth_provider", - })), - ))); - } + if !management_token_may_configure_frontend_callback( + request_context, + existing.as_ref(), + &payload.frontend_callback_url, + ) { + return Ok(Some(oauth_frontend_callback_permission_denied_response( + request_context, + ))); } + let force_disable = payload.force; let record = match build_admin_oauth_upsert_record(state, &provider_type, payload) { Ok(record) => record, Err(message) => { @@ -302,9 +519,63 @@ pub(crate) async fn maybe_build_local_admin_oauth_response( ))); } }; - let Some(provider) = state.upsert_oauth_provider_config(&record).await? else { + if let Some(existing) = existing.as_ref() { + if existing + .client_secret_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + && matches!( + record.client_secret_encrypted, + aether_data::repository::oauth_providers::EncryptedSecretUpdate::Preserve + ) + { + match crate::handlers::shared::identity_oauth_provider_secret_binding_matches( + existing, &record, + ) { + Ok(true) => {} + Ok(false) => { + return Ok(Some(build_proxy_error_response( + http::StatusCode::BAD_REQUEST, + "invalid_request", + "修改 OAuth Provider 的 Client ID、端点或 redirect_uri 时必须重新提供 client_secret", + None, + ))); + } + Err(_) => { + return Ok(Some(build_proxy_error_response( + http::StatusCode::BAD_REQUEST, + "invalid_request", + "OAuth Provider 密钥绑定校验失败,请重新提供 client_secret", + None, + ))); + } + } + } + } + let Some(outcome) = state + .upsert_oauth_provider_config_with_force_disable(&record, force_disable) + .await? + else { return Ok(None); }; + let provider = match outcome { + aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::Upserted( + provider, + ) => provider, + aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count, + } => { + return Ok(Some(build_proxy_error_response( + http::StatusCode::CONFLICT, + "confirmation_required", + format!("禁用该 Provider 会导致 {affected_count} 个用户无法登录"), + Some(json!({ + "affected_count": affected_count, + "action": "disable_oauth_provider", + })), + ))); + } + }; return Ok(Some( Json(build_admin_oauth_provider_payload(&provider)).into_response(), )); @@ -322,7 +593,7 @@ pub(crate) async fn maybe_build_local_admin_oauth_response( None, ))); }; - let Some(existing) = state.get_oauth_provider_config(&provider_type).await? else { + let Some(_existing) = state.get_oauth_provider_config(&provider_type).await? else { return Ok(Some(build_proxy_error_response( http::StatusCode::BAD_REQUEST, "invalid_request", @@ -330,32 +601,27 @@ pub(crate) async fn maybe_build_local_admin_oauth_response( None, ))); }; - if existing.is_enabled { - let ldap_exclusive = state.get_ldap_module_config().await?.is_some_and(|config| { - config.is_enabled - && config.is_exclusive - && config - .bind_password_encrypted - .as_deref() - .map(str::trim) - .is_some_and(|value| !value.is_empty()) - }); - let affected_count = state - .count_locked_users_if_oauth_provider_disabled(&provider_type, ldap_exclusive) - .await?; - if affected_count > 0 { + let _mutation_guard = crate::oauth::lock_identity_oauth_mutation().await; + if state.has_oauth_links_for_provider(&provider_type).await? { + return Ok(Some(build_proxy_error_response( + http::StatusCode::CONFLICT, + "provider_has_bindings", + "Provider 仍有用户绑定,必须先解除全部绑定", + None, + ))); + } + let deleted = state + .delete_oauth_provider_config_if_unlinked(&provider_type) + .await?; + if !deleted { + if state.has_oauth_links_for_provider(&provider_type).await? { return Ok(Some(build_proxy_error_response( - http::StatusCode::BAD_REQUEST, - "invalid_request", - format!( - "删除该 Provider 会导致部分用户无法登录(数量: {affected_count}),已阻止操作" - ), + http::StatusCode::CONFLICT, + "provider_has_bindings", + "Provider 仍有用户绑定,必须先解除全部绑定", None, ))); } - } - let deleted = state.delete_oauth_provider_config(&provider_type).await?; - if !deleted { return Ok(Some(build_proxy_error_response( http::StatusCode::BAD_REQUEST, "invalid_request", diff --git a/apps/aether-gateway/src/handlers/admin/auth/routes.rs b/apps/aether-gateway/src/handlers/admin/auth/routes.rs index 8dec1148d..3eb75e60c 100644 --- a/apps/aether-gateway/src/handlers/admin/auth/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/auth/routes.rs @@ -18,6 +18,7 @@ pub(crate) async fn maybe_build_local_admin_auth_response( &request.state(), &request.request_context(), request.request_headers(), + request.remote_addr(), request.request_body(), ) .await? diff --git a/apps/aether-gateway/src/handlers/admin/billing/collectors/support.rs b/apps/aether-gateway/src/handlers/admin/billing/collectors/support.rs index d744cabf0..f2dfc7a58 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/collectors/support.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/collectors/support.rs @@ -5,6 +5,7 @@ use super::super::{ use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::unix_secs_to_rfc3339; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::{Body, Bytes}, http, @@ -53,7 +54,7 @@ pub(super) fn build_admin_billing_collector_payload_from_record( "default_value": record.default_value, "priority": record.priority, "is_enabled": record.is_enabled, - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(record.updated_at_unix_secs), }) } @@ -174,17 +175,7 @@ pub(super) async fn parse_admin_billing_collector_request( )); } Ok(false) => {} - Err(err) => { - let detail = match err { - GatewayError::Internal(message) => message, - other => format!("{other:?}"), - }; - return Err(( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": detail })), - ) - .into_response()); - } + Err(_err) => return Err(build_admin_billing_internal_error_response()), } } @@ -202,6 +193,16 @@ pub(super) async fn parse_admin_billing_collector_request( }) } +fn build_admin_billing_internal_error_response() -> Response { + ( + http::StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ + "detail": "计费采集器服务暂不可用,请稍后重试" + })), + ) + .into_response() +} + pub(in super::super) fn admin_billing_parse_page(query: Option<&str>) -> Result { super::super::admin_billing_parse_page(query) } diff --git a/apps/aether-gateway/src/handlers/admin/billing/mod.rs b/apps/aether-gateway/src/handlers/admin/billing/mod.rs index 8df84d47e..ad2d74c76 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/mod.rs @@ -20,6 +20,7 @@ mod routes; mod rules; mod wallets; +pub(in crate::handlers::admin) use self::payments::admin_payment_gateway_response_projection; pub(super) use self::payments::maybe_build_local_admin_payments_response; pub(super) use self::routes::maybe_build_local_admin_billing_routes_response; pub(super) use self::wallets::maybe_build_local_admin_wallets_response; diff --git a/apps/aether-gateway/src/handlers/admin/billing/payments/gateways.rs b/apps/aether-gateway/src/handlers/admin/billing/payments/gateways.rs index 49ca001d9..851a38448 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/payments/gateways.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/payments/gateways.rs @@ -3,12 +3,16 @@ use super::{ }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::shared::{ + normalize_payment_callback_base_url, normalize_payment_currency, normalize_payment_https_url, payment_gateway_allow_user_refund, payment_gateway_channels_config_json, payment_gateway_channels_json, payment_gateway_config_json, payment_gateway_refund_enabled, - payment_gateway_secret_keys_json, + payment_gateway_secret_is_legacy_unbound, payment_gateway_secret_keys_json, + PaymentGatewaySecretBinding, }; use crate::{GatewayError, LocalMutationOutcome}; -use aether_data_contracts::repository::billing::PaymentGatewayConfigWriteInput; +use aether_data_contracts::repository::billing::{ + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigWriteInput, +}; use axum::{ body::Body, http, @@ -18,7 +22,9 @@ use axum::{ use serde::Deserialize; use serde_json::{json, Value}; -#[derive(Debug, Deserialize)] +const PAYMENT_GATEWAY_CONFIG_CAS_MAX_ATTEMPTS: usize = 8; + +#[derive(Deserialize)] struct PaymentGatewayConfigRequest { #[serde(default)] enabled: bool, @@ -60,6 +66,14 @@ fn default_min_recharge_usd() -> f64 { 1.0 } +fn build_payment_gateway_conflict_response(detail: impl Into) -> Response { + ( + http::StatusCode::CONFLICT, + Json(json!({ "detail": detail.into() })), + ) + .into_response() +} + fn default_channels() -> Value { json!([ {"channel": "alipay", "display_name": "支付宝", "fee_rate": 0.0}, @@ -121,6 +135,18 @@ fn admin_payment_gateway_provider_from_path(path: &str) -> Option { Some(provider) } +fn resolve_admin_payment_gateway_provider(path: &str, route_kind: &str) -> Option { + match route_kind { + "get_epay_gateway" | "update_epay_gateway" | "test_epay_gateway" => { + Some("epay".to_string()) + } + "get_payment_gateway" | "update_payment_gateway" | "test_payment_gateway" => { + admin_payment_gateway_provider_from_path(path) + } + _ => None, + } +} + fn default_provider_channels(provider: &str) -> Value { match provider { "epay" => default_channels(), @@ -263,29 +289,97 @@ fn normalize_config_object(config: Value) -> Result { Err("config must be an object".to_string()) } +fn merge_gateway_secret_maps( + existing_plaintext: Option<&str>, + updates: serde_json::Map, +) -> Result, &'static str> { + let mut merged = match existing_plaintext { + Some(plaintext) => serde_json::from_str::(plaintext) + .ok() + .and_then(|value| value.as_object().cloned()) + .ok_or("existing gateway secrets have invalid format")?, + None => serde_json::Map::new(), + }; + merged.extend(updates); + Ok(merged) +} + +/// A legacy gateway ciphertext has no authenticated destination (or only the +/// provider in v2). Reusing it while changing endpoint/merchant would carry +/// an unknown credential into a different payment account. Require the +/// administrator to provide a replacement secret in that case. +fn legacy_secret_reuse_requires_reentry( + existing: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>, + requested_binding: &PaymentGatewaySecretBinding, +) -> bool { + let Some(record) = existing else { + return false; + }; + let Some(ciphertext) = record.merchant_key_encrypted.as_deref() else { + return false; + }; + if !payment_gateway_secret_is_legacy_unbound(ciphertext) { + return false; + } + + // An invalid historical binding cannot establish that the legacy value + // belongs to the requested destination, so fail closed as well. + PaymentGatewaySecretBinding::from_record(record) + .map(|stored_binding| stored_binding != requested_binding.clone()) + .unwrap_or(true) +} + fn encrypted_gateway_secret( state: &AdminAppState<'_>, - provider: &str, + binding: &PaymentGatewaySecretBinding, payload: &PaymentGatewayConfigRequest, -) -> Result, Response> { + existing: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>, +) -> Result<(Option, Vec), Response> { + let provider = binding.provider.as_str(); + let decrypt_existing = || { + if legacy_secret_reuse_requires_reentry(existing, binding) { + return Err(build_admin_payments_bad_request_response( + "endpoint_url or merchant_id changed; re-enter the gateway secret", + )); + } + existing + .and_then(|record| record.merchant_key_encrypted.as_deref()) + .map(|ciphertext| { + crate::handlers::shared::open_payment_gateway_secret( + state.app(), binding, ciphertext, + ) + .map(|projection| projection.plaintext) + .map_err(|_| { + build_admin_payments_backend_unavailable_response( + "existing gateway secrets are not valid for the requested destination; re-enter the secret", + ) + }) + }) + .transpose() + }; let secret_plaintext = if provider == "epay" { - payload + let supplied = payload .merchant_key .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .map(ToOwned::to_owned); + if supplied.is_none() { + decrypt_existing()?; + } + supplied } else { let Some(secrets) = payload.secrets.as_object() else { return if payload.secrets.is_null() { - Ok(None) + decrypt_existing()?; + Ok((None, existing_gateway_secret_keys(existing))) } else { Err(build_admin_payments_bad_request_response( "secrets must be an object", )) }; }; - let filtered = secrets + let updates = secrets .iter() .filter_map(|(key, value)| { let value = value.as_str()?.trim(); @@ -293,39 +387,60 @@ fn encrypted_gateway_secret( .then(|| (key.trim().to_string(), Value::String(value.to_string()))) }) .collect::>(); - if filtered.is_empty() { - None - } else { - Some(Value::Object(filtered).to_string()) + if updates.is_empty() { + decrypt_existing()?; + return Ok((None, existing_gateway_secret_keys(existing))); } + + let existing_plaintext = decrypt_existing()?; + let merged = match merge_gateway_secret_maps(existing_plaintext.as_deref(), updates) { + Ok(value) => value, + Err(detail) => { + return Err(build_admin_payments_backend_unavailable_response(detail)); + } + }; + Some(Value::Object(merged).to_string()) }; let Some(secret_plaintext) = secret_plaintext else { - return Ok(None); + return Ok((None, existing_gateway_secret_keys(existing))); }; - state - .encrypt_catalog_secret_with_fallbacks(&secret_plaintext) - .ok_or_else(|| { - build_admin_payments_backend_unavailable_response("encryption key is not configured") - }) - .map(Some) + let encrypted = crate::handlers::shared::seal_payment_gateway_secret( + state.app(), + binding, + &secret_plaintext, + ) + .map_err(build_admin_payments_backend_unavailable_response)?; + let secret_keys = if provider == "epay" { + Vec::new() + } else { + let mut keys = serde_json::from_str::(&secret_plaintext) + .ok() + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default() + .into_iter() + .map(|(key, _)| Value::String(key)) + .collect::>(); + keys.sort_by(|left, right| left.as_str().cmp(&right.as_str())); + keys + }; + Ok((Some(encrypted), secret_keys)) } -async fn existing_gateway_secret_keys( - state: &AdminAppState<'_>, - provider: &str, -) -> Result, GatewayError> { - let Some(record) = state.app().find_payment_gateway_config(provider).await? else { - return Ok(Vec::new()); +fn existing_gateway_secret_keys( + record: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>, +) -> Vec { + let Some(record) = record else { + return Vec::new(); }; - let (_, _, secret_keys, _, _) = split_gateway_channels_config(&record); - Ok(secret_keys + let (_, _, secret_keys, _, _) = split_gateway_channels_config(record); + secret_keys .as_array() .cloned() .unwrap_or_default() .into_iter() .filter(|value| value.as_str().is_some_and(|item| !item.trim().is_empty())) - .collect()) + .collect() } pub(super) async fn maybe_build_local_admin_payment_gateways_response( @@ -336,8 +451,14 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response( ) -> Result>, GatewayError> { match route_kind { Some("get_epay_gateway") | Some("get_payment_gateway") => { - let provider = admin_payment_gateway_provider_from_path(request_context.path()) - .unwrap_or_else(|| "epay".to_string()); + let Some(provider) = resolve_admin_payment_gateway_provider( + request_context.path(), + route_kind.expect("matched payment gateway route kind"), + ) else { + return Ok(Some(build_admin_payments_bad_request_response( + "unsupported payment gateway provider", + ))); + }; let record = state.app().find_payment_gateway_config(&provider).await?; let payload = record .map(gateway_config_payload) @@ -345,8 +466,14 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response( Ok(Some(Json(payload).into_response())) } Some("update_epay_gateway") | Some("update_payment_gateway") => { - let provider = admin_payment_gateway_provider_from_path(request_context.path()) - .unwrap_or_else(|| "epay".to_string()); + let Some(provider) = resolve_admin_payment_gateway_provider( + request_context.path(), + route_kind.expect("matched payment gateway route kind"), + ) else { + return Ok(Some(build_admin_payments_bad_request_response( + "unsupported payment gateway provider", + ))); + }; let Some(body) = request_body else { return Ok(Some(build_admin_payments_bad_request_response( "缺少请求体", @@ -371,109 +498,167 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response( ))); } - let merchant_key_encrypted = match encrypted_gateway_secret(state, &provider, &payload) - { - Ok(value) => value, - Err(response) => return Ok(Some(response)), - }; let endpoint_url = if provider == "epay" { - match normalize_text(payload.endpoint_url, "endpoint_url", 512) { - Ok(value) => value, + match normalize_text(payload.endpoint_url.clone(), "endpoint_url", 512) { + Ok(value) => match normalize_payment_https_url(&value, "endpoint_url") { + Ok(value) => value, + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }, Err(detail) => { return Ok(Some(build_admin_payments_bad_request_response(detail))) } } } else { - match normalize_optional_text(Some(payload.endpoint_url), 512) { - Ok(value) => value.unwrap_or_default(), + match normalize_optional_text(Some(payload.endpoint_url.clone()), 512) { + Ok(Some(value)) => match normalize_payment_https_url(&value, "endpoint_url") { + Ok(value) => value, + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }, + Ok(None) => String::new(), Err(detail) => { return Ok(Some(build_admin_payments_bad_request_response(detail))) } } }; - let callback_base_url = match normalize_optional_text(payload.callback_base_url, 512) { - Ok(value) => value, - Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))), - }; + let callback_base_url = + match normalize_optional_text(payload.callback_base_url.clone(), 512) { + Ok(Some(value)) => match normalize_payment_callback_base_url(&value) { + Ok(value) => Some(value), + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }, + Ok(None) => None, + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }; let merchant_id = if provider == "epay" { - match normalize_text(payload.merchant_id, "merchant_id", 128) { + match normalize_text(payload.merchant_id.clone(), "merchant_id", 128) { Ok(value) => value, Err(detail) => { return Ok(Some(build_admin_payments_bad_request_response(detail))) } } } else { - match normalize_optional_text(Some(payload.merchant_id), 128) { + match normalize_optional_text(Some(payload.merchant_id.clone()), 128) { Ok(value) => value.unwrap_or_default(), Err(detail) => { return Ok(Some(build_admin_payments_bad_request_response(detail))) } } }; - let pay_currency = match normalize_text(payload.pay_currency, "pay_currency", 16) { + let pay_currency = + match normalize_payment_currency(&payload.pay_currency, "pay_currency") { + Ok(value) => value, + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }; + let config = match normalize_config_object(payload.config.clone()) { Ok(value) => value, Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))), }; - let config = match normalize_config_object(payload.config) { - Ok(value) => value, - Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))), - }; - let submitted_secret_keys = payload - .secrets - .as_object() - .map(|secrets| { - secrets - .iter() - .filter(|(_, value)| { - value.as_str().is_some_and(|value| !value.trim().is_empty()) - }) - .map(|(key, _)| Value::String(key.clone())) - .collect::>() - }) - .unwrap_or_default(); - let secret_keys = if provider == "epay" || !submitted_secret_keys.is_empty() { - submitted_secret_keys - } else { - existing_gateway_secret_keys(state, &provider).await? - }; - let channels = match normalize_gateway_channels(&provider, payload.channels) { + let channels = match normalize_gateway_channels(&provider, payload.channels.clone()) { Ok(value) => value, Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))), }; let refund_enabled = payload.refund_enabled; let allow_user_refund = refund_enabled && payload.allow_user_refund; - let channels_json = payment_gateway_channels_config_json( - channels, - config, - Value::Array(secret_keys), - refund_enabled, - allow_user_refund, - ); - let input = PaymentGatewayConfigWriteInput { - provider: provider.clone(), - enabled: payload.enabled, - endpoint_url, - callback_base_url, - merchant_id, - preserve_existing_secret: merchant_key_encrypted.is_none(), - merchant_key_encrypted, - pay_currency, - usd_exchange_rate: payload.usd_exchange_rate, - min_recharge_usd: payload.min_recharge_usd, - channels_json, - }; - match state.app().upsert_payment_gateway_config(&input).await? { - LocalMutationOutcome::Applied(record) => { - Ok(Some(Json(gateway_config_payload(record)).into_response())) + let binding = + match PaymentGatewaySecretBinding::new(&provider, &endpoint_url, &merchant_id) { + Ok(value) => value, + Err(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))) + } + }; + let mut existing_record = state.app().find_payment_gateway_config(&provider).await?; + let expected_existing = existing_record.is_some(); + for _ in 0..PAYMENT_GATEWAY_CONFIG_CAS_MAX_ATTEMPTS { + if expected_existing && existing_record.is_none() { + return Ok(Some(build_payment_gateway_conflict_response( + "payment gateway config was removed concurrently", + ))); + } + let (merchant_key_encrypted, secret_keys) = match encrypted_gateway_secret( + state, + &binding, + &payload, + existing_record.as_ref(), + ) { + Ok(value) => value, + Err(response) => return Ok(Some(response)), + }; + let channels_json = payment_gateway_channels_config_json( + channels.clone(), + config.clone(), + Value::Array(secret_keys), + refund_enabled, + allow_user_refund, + ); + let mutation = PaymentGatewayConfigCasWriteInput { + input: PaymentGatewayConfigWriteInput { + provider: provider.clone(), + enabled: payload.enabled, + endpoint_url: endpoint_url.clone(), + callback_base_url: callback_base_url.clone(), + merchant_id: merchant_id.clone(), + preserve_existing_secret: merchant_key_encrypted.is_none(), + merchant_key_encrypted, + pay_currency: pay_currency.clone(), + usd_exchange_rate: payload.usd_exchange_rate, + min_recharge_usd: payload.min_recharge_usd, + channels_json, + }, + expected_existing, + expected_merchant_key_encrypted: existing_record + .as_ref() + .and_then(|record| record.merchant_key_encrypted.clone()), + }; + match state + .app() + .compare_and_swap_payment_gateway_config(&mutation) + .await? + { + LocalMutationOutcome::Applied(record) => { + return Ok(Some(Json(gateway_config_payload(record)).into_response())); + } + LocalMutationOutcome::NotFound if !expected_existing => { + return Ok(Some(build_payment_gateway_conflict_response( + "payment gateway config was created concurrently", + ))); + } + LocalMutationOutcome::NotFound => { + existing_record = + state.app().find_payment_gateway_config(&provider).await?; + } + LocalMutationOutcome::Invalid(detail) => { + return Ok(Some(build_admin_payments_bad_request_response(detail))); + } + LocalMutationOutcome::Unavailable => { + return Ok(Some(build_admin_payments_backend_unavailable_response( + "payment gateway config backend unavailable", + ))); + } } - _ => Ok(Some(build_admin_payments_backend_unavailable_response( - "payment gateway config backend unavailable", - ))), } + Ok(Some(build_payment_gateway_conflict_response( + "payment gateway config changed too frequently; retry the request", + ))) } Some("test_epay_gateway") | Some("test_payment_gateway") => { - let provider = admin_payment_gateway_provider_from_path(request_context.path()) - .unwrap_or_else(|| "epay".to_string()); + let Some(provider) = resolve_admin_payment_gateway_provider( + request_context.path(), + route_kind.expect("matched payment gateway route kind"), + ) else { + return Ok(Some(build_admin_payments_bad_request_response( + "unsupported payment gateway provider", + ))); + }; let status = state.app().find_payment_gateway_config(&provider).await?; let ok = status .as_ref() @@ -493,3 +678,164 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response( _ => Ok(None), } } + +#[cfg(test)] +mod tests { + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; + use aether_data_contracts::repository::billing::PaymentGatewayConfigRecord; + use serde_json::{json, Value}; + + use super::{ + legacy_secret_reuse_requires_reentry, merge_gateway_secret_maps, + resolve_admin_payment_gateway_provider, + }; + use crate::handlers::shared::PaymentGatewaySecretBinding; + + fn gateway_record( + endpoint_url: &str, + merchant_id: &str, + merchant_key_encrypted: Option, + ) -> PaymentGatewayConfigRecord { + PaymentGatewayConfigRecord { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: endpoint_url.to_string(), + callback_base_url: None, + merchant_id: merchant_id.to_string(), + merchant_key_encrypted, + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: json!({}), + created_at_unix_secs: 1, + updated_at_unix_secs: 1, + } + } + + #[test] + fn legacy_secret_reuse_requires_reentry_after_binding_change() { + let legacy = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-secret") + .expect("legacy secret should encrypt"); + let old_record = gateway_record("https://api.stripe.com", "merchant-old", Some(legacy)); + let changed_binding = + PaymentGatewaySecretBinding::new("stripe", "https://api.stripe.com", "merchant-new") + .expect("changed binding should be valid"); + assert!(legacy_secret_reuse_requires_reentry( + Some(&old_record), + &changed_binding, + )); + + let v2_record = gateway_record( + "https://api.stripe.com", + "merchant-old", + Some("aether-payment-gateway-secret-v2:legacy".to_string()), + ); + assert!(legacy_secret_reuse_requires_reentry( + Some(&v2_record), + &changed_binding, + )); + } + + #[test] + fn legacy_secret_reuse_is_allowed_only_for_same_binding_or_bound_v3() { + let legacy = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-secret") + .expect("legacy secret should encrypt"); + let old_record = + gateway_record("https://API.STRIPE.COM:443/", "merchant-old", Some(legacy)); + let same_binding = PaymentGatewaySecretBinding::new( + "stripe", + "https://api.stripe.com:443/", + " merchant-old ", + ) + .expect("same binding should be valid"); + assert!(!legacy_secret_reuse_requires_reentry( + Some(&old_record), + &same_binding, + )); + + let v3_record = gateway_record( + "https://api.stripe.com", + "merchant-old", + Some("aether-payment-gateway-secret-v3:bound".to_string()), + ); + let changed_binding = + PaymentGatewaySecretBinding::new("stripe", "https://api.stripe.com", "merchant-new") + .expect("changed binding should be valid"); + assert!(!legacy_secret_reuse_requires_reentry( + Some(&v3_record), + &changed_binding, + )); + } + + #[test] + fn stripe_secret_rotation_preserves_omitted_secret_fields() { + let existing = json!({ + "secret_key": "old-secret-key", + "webhook_secret": "old-webhook" + }) + .to_string(); + let updates = json!({"webhook_secret": "new-webhook"}) + .as_object() + .cloned() + .expect("updates should be an object"); + + let merged = Value::Object( + merge_gateway_secret_maps(Some(&existing), updates) + .expect("valid secret maps should merge"), + ); + assert_eq!(merged["secret_key"], "old-secret-key"); + assert_eq!(merged["webhook_secret"], "new-webhook"); + } + + #[test] + fn wxpay_secret_rotation_preserves_omitted_secret_fields() { + let existing = json!({ + "private_key": "old-private", + "api_v3_key": "old-api-v3-key", + "public_key": "old-public" + }) + .to_string(); + let updates = json!({"api_v3_key": "new-api-v3-key"}) + .as_object() + .cloned() + .expect("updates should be an object"); + + let merged = Value::Object( + merge_gateway_secret_maps(Some(&existing), updates) + .expect("valid secret maps should merge"), + ); + assert_eq!(merged["private_key"], "old-private"); + assert_eq!(merged["api_v3_key"], "new-api-v3-key"); + assert_eq!(merged["public_key"], "old-public"); + } + + #[test] + fn generic_gateway_routes_never_fall_back_to_epay() { + assert_eq!( + resolve_admin_payment_gateway_provider( + "/api/admin/payments/gateways/stripe", + "update_payment_gateway", + ) + .as_deref(), + Some("stripe") + ); + assert!(resolve_admin_payment_gateway_provider( + "/api/admin/payments/gateways/unsupported", + "update_payment_gateway", + ) + .is_none()); + assert!(resolve_admin_payment_gateway_provider( + "/api/admin/payments/gateways/stripe/extra", + "get_payment_gateway", + ) + .is_none()); + assert_eq!( + resolve_admin_payment_gateway_provider( + "/api/admin/payments/epay", + "update_epay_gateway", + ) + .as_deref(), + Some("epay") + ); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/billing/payments/mod.rs b/apps/aether-gateway/src/handlers/admin/billing/payments/mod.rs index 951a7b26e..9af3b42ef 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/payments/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/payments/mod.rs @@ -11,6 +11,7 @@ mod redeem_codes; mod routes; mod shared; +pub(in crate::handlers::admin) use self::shared::admin_payment_gateway_response_projection; use self::shared::{ admin_payment_operator_id, admin_payment_order_id_from_detail_path, admin_payment_order_id_from_suffix_path, build_admin_payment_callback_payload_from_record, @@ -19,7 +20,8 @@ use self::shared::{ build_admin_payments_bad_request_response, build_admin_payments_data_unavailable_response, normalize_admin_payment_currency, normalize_admin_payment_optional_string, normalize_admin_payment_positive_number, parse_admin_payments_limit, - parse_admin_payments_offset, AdminPaymentOrderCreditRequest, + parse_admin_payments_offset, prepare_admin_payment_gateway_response_for_storage, + AdminPaymentOrderCreditRequest, }; pub(crate) async fn maybe_build_local_admin_payments_response( diff --git a/apps/aether-gateway/src/handlers/admin/billing/payments/orders.rs b/apps/aether-gateway/src/handlers/admin/billing/payments/orders.rs index 860631452..cb2fcbca1 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/payments/orders.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/payments/orders.rs @@ -5,7 +5,8 @@ use super::{ build_admin_payments_backend_unavailable_response, build_admin_payments_bad_request_response, normalize_admin_payment_currency, normalize_admin_payment_optional_string, normalize_admin_payment_positive_number, parse_admin_payments_limit, - parse_admin_payments_offset, AdminPaymentOrderCreditRequest, + parse_admin_payments_offset, prepare_admin_payment_gateway_response_for_storage, + AdminPaymentOrderCreditRequest, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value}; @@ -124,7 +125,9 @@ async fn close_direct_gateway_order_before_terminal_mark( ))) } }; - if order.status != "pending" || !matches!(order.payment_method.as_str(), "alipay" | "wxpay") { + if order.status != "pending" + || !matches!(order.payment_method.as_str(), "alipay" | "wxpay" | "stripe") + { return Ok(None); } crate::handlers::shared::close_direct_gateway_order(state.app(), &order) @@ -231,6 +234,8 @@ async fn build_admin_payment_credit_order_response( "gateway_response 必须为对象", )); } + let gateway_response = + prepare_admin_payment_gateway_response_for_storage(payload.gateway_response); let operator_id = admin_payment_operator_id(request_context); match state .admin_credit_payment_order( @@ -239,7 +244,7 @@ async fn build_admin_payment_credit_order_response( pay_amount, pay_currency.as_deref(), exchange_rate, - payload.gateway_response, + gateway_response, operator_id.as_deref(), ) .await? diff --git a/apps/aether-gateway/src/handlers/admin/billing/payments/redeem_codes.rs b/apps/aether-gateway/src/handlers/admin/billing/payments/redeem_codes.rs index cc980a437..860090b8f 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/payments/redeem_codes.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/payments/redeem_codes.rs @@ -8,6 +8,7 @@ use crate::handlers::admin::shared::{ attach_admin_audit_response, query_param_value, unix_secs_to_rfc3339, }; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, http, @@ -122,7 +123,7 @@ fn build_batch_payload( "description": batch.description, "created_by": batch.created_by, "expires_at": batch.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), - "created_at": unix_secs_to_rfc3339(batch.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(batch.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(batch.updated_at_unix_secs), }) } @@ -146,7 +147,7 @@ fn build_code_payload( "redeemed_at": code.redeemed_at_unix_secs.and_then(unix_secs_to_rfc3339), "disabled_by": code.disabled_by, "expires_at": code.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), - "created_at": unix_secs_to_rfc3339(code.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(code.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(code.updated_at_unix_secs), }) } diff --git a/apps/aether-gateway/src/handlers/admin/billing/payments/shared.rs b/apps/aether-gateway/src/handlers/admin/billing/payments/shared.rs index 85cdc0437..8795df8a7 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/payments/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/payments/shared.rs @@ -1,17 +1,19 @@ use crate::handlers::admin::request::AdminRequestContext; use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339}; +use crate::handlers::shared::normalize_payment_currency; use crate::GatewayAdminPaymentCallbackView; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, http, response::{IntoResponse, Response}, Json, }; -use serde_json::json; +use serde_json::{json, Value}; const ADMIN_PAYMENTS_DATA_UNAVAILABLE_DETAIL: &str = "Admin payments data unavailable"; -#[derive(Debug, Default, serde::Deserialize)] +#[derive(Default, serde::Deserialize)] pub(super) struct AdminPaymentOrderCreditRequest { #[serde(default)] pub(super) gateway_order_id: Option, @@ -25,6 +27,22 @@ pub(super) struct AdminPaymentOrderCreditRequest { pub(super) gateway_response: Option, } +impl std::fmt::Debug for AdminPaymentOrderCreditRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPaymentOrderCreditRequest") + .field("gateway_order_id", &self.gateway_order_id) + .field("pay_amount", &self.pay_amount) + .field("pay_currency", &self.pay_currency) + .field("exchange_rate", &self.exchange_rate) + .field( + "gateway_response", + &self.gateway_response.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + pub(super) fn build_admin_payments_data_unavailable_response() -> Response { ( http::StatusCode::SERVICE_UNAVAILABLE, @@ -153,11 +171,9 @@ pub(super) fn normalize_admin_payment_currency( let Some(value) = normalize_admin_payment_optional_string(value, "pay_currency", 3)? else { return Ok(None); }; - let normalized = value.to_ascii_uppercase(); - if normalized.len() != 3 { - return Err("pay_currency 必须是 3 位货币代码".to_string()); - } - Ok(Some(normalized)) + normalize_payment_currency(&value, "pay_currency") + .map(Some) + .map_err(|_| "pay_currency 必须是 3 位 ASCII 货币代码".to_string()) } pub(super) fn normalize_admin_payment_positive_number( @@ -187,13 +203,184 @@ pub(super) fn admin_payment_effective_status( expires_at_unix_secs: Option, ) -> String { let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64; - if status == "pending" && expires_at_unix_secs.is_some_and(|value| value < now_unix_secs) { + if status == "pending" && expires_at_unix_secs.is_some_and(|value| value <= now_unix_secs) { "expired".to_string() } else { status.to_string() } } +fn admin_payment_bounded_string(value: &Value, max_chars: usize) -> Option { + let value = value.as_str()?.trim(); + (!value.is_empty() && value.chars().count() <= max_chars) + .then(|| Value::String(value.to_string())) +} + +fn admin_payment_identifier(value: &Value, max_chars: usize) -> Option { + let value = value.as_str()?.trim(); + (!value.is_empty() + && value.chars().count() <= max_chars + && value + .chars() + .all(|character| character.is_ascii_alphanumeric() || matches!(character, '_' | '-'))) + .then(|| Value::String(value.to_string())) +} + +fn admin_payment_gateway_response_field(key: &str, value: &Value) -> Option { + match key { + "gateway" | "submit_method" | "payment_channel" => admin_payment_identifier(value, 64), + "pay_currency" => admin_payment_identifier(value, 16), + "display_name" | "provider_label" => admin_payment_bounded_string(value, 128), + "gateway_order_id" | "intent_id" => admin_payment_bounded_string(value, 256), + "expires_at" => admin_payment_bounded_string(value, 64), + "pay_amount" | "base_pay_amount" | "fee_rate" | "fee_amount" => { + value.is_number().then(|| value.clone()) + } + "manual_credit" => value.as_bool().map(Value::Bool), + "payment_method_types" => { + let values = value.as_array()?; + if values.len() > 16 { + return None; + } + values + .iter() + .map(|value| admin_payment_identifier(value, 64)) + .collect::>>() + .map(Value::Array) + } + _ => None, + } +} + +pub(in crate::handlers::admin) fn admin_payment_gateway_response_projection( + value: Option<&Value>, +) -> Value { + let Some(object) = value.and_then(Value::as_object) else { + return Value::Null; + }; + Value::Object( + object + .iter() + .filter_map(|(key, value)| { + admin_payment_gateway_response_field(key, value).map(|value| (key.clone(), value)) + }) + .collect(), + ) +} + +pub(super) fn prepare_admin_payment_gateway_response_for_storage( + value: Option, +) -> Option { + value.map(|value| admin_payment_gateway_response_projection(Some(&value))) +} + +#[derive(Default)] +struct AdminPaymentJsonShape { + objects: u64, + arrays: u64, + strings: u64, + numbers: u64, + booleans: u64, + nulls: u64, + object_fields: u64, + array_items: u64, + max_depth: u64, +} + +impl AdminPaymentJsonShape { + fn observe(&mut self, value: &Value, depth: u64) { + self.max_depth = self.max_depth.max(depth); + match value { + Value::Object(object) => { + self.objects = self.objects.saturating_add(1); + self.object_fields = self + .object_fields + .saturating_add(u64::try_from(object.len()).unwrap_or(u64::MAX)); + for value in object.values() { + self.observe(value, depth.saturating_add(1)); + } + } + Value::Array(values) => { + self.arrays = self.arrays.saturating_add(1); + self.array_items = self + .array_items + .saturating_add(u64::try_from(values.len()).unwrap_or(u64::MAX)); + for value in values { + self.observe(value, depth.saturating_add(1)); + } + } + Value::String(_) => self.strings = self.strings.saturating_add(1), + Value::Number(_) => self.numbers = self.numbers.saturating_add(1), + Value::Bool(_) => self.booleans = self.booleans.saturating_add(1), + Value::Null => self.nulls = self.nulls.saturating_add(1), + } + } +} + +fn admin_payment_json_kind(value: &Value) -> &'static str { + match value { + Value::Null => "null", + Value::Bool(_) => "boolean", + Value::Number(_) => "number", + Value::String(_) => "string", + Value::Array(_) => "array", + Value::Object(_) => "object", + } +} + +fn admin_payment_payload_summary(value: Option<&Value>) -> Value { + let Some(value) = value else { + return Value::Null; + }; + let mut shape = AdminPaymentJsonShape::default(); + shape.observe(value, 1); + json!({ + "kind": admin_payment_json_kind(value), + "serialized_bytes": serde_json::to_vec(value).map_or(0, |encoded| encoded.len()), + "objects": shape.objects, + "arrays": shape.arrays, + "strings": shape.strings, + "numbers": shape.numbers, + "booleans": shape.booleans, + "nulls": shape.nulls, + "object_fields": shape.object_fields, + "array_items": shape.array_items, + "max_depth": shape.max_depth, + }) +} + +fn admin_payment_callback_error_projection(value: Option<&str>) -> Option { + const SAFE_ERRORS: &[&str] = &[ + "callback amount mismatch", + "callback key reused with different payment payload", + "invalid callback signature", + "invalid payment callback numeric or identity fields", + "payment channel mismatch", + "payment currency mismatch", + "payment gateway order belongs to another payment order", + "payment gateway order identifier mismatch", + "payment gateway order mismatch", + "payment method mismatch", + "payment order expired", + "payment order not found", + "payment order number mismatch", + "payment order user missing", + "payment provider mismatch", + "plan purchase limit reached", + "wallet is not active", + "wallet not found", + ]; + + let value = value?.trim(); + if SAFE_ERRORS.contains(&value) { + return Some(value.to_string()); + } + if value.starts_with("payment order is not creditable:") { + return Some("payment order is not creditable".to_string()); + } + Some("payment callback processing failed".to_string()) +} + pub(super) fn build_admin_payment_order_payload( record: &crate::AdminWalletPaymentOrderRecord, ) -> serde_json::Value { @@ -210,9 +397,10 @@ pub(super) fn build_admin_payment_order_payload( "refundable_amount_usd": record.refundable_amount_usd, "payment_method": record.payment_method, "gateway_order_id": record.gateway_order_id, - "gateway_response": record.gateway_response, + "gateway_response": admin_payment_gateway_response_projection(record.gateway_response.as_ref()), + "has_gateway_response": record.gateway_response.is_some(), "status": admin_payment_effective_status(&record.status, record.expires_at_unix_secs), - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "paid_at": record.paid_at_unix_secs.and_then(unix_secs_to_rfc3339), "credited_at": record.credited_at_unix_secs.and_then(unix_secs_to_rfc3339), "expires_at": record.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), @@ -232,9 +420,196 @@ pub(super) fn build_admin_payment_callback_payload_from_record( "payload_hash": record.payload_hash, "signature_valid": record.signature_valid, "status": record.status, - "payload": record.payload, - "error_message": record.error_message, - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "payload": Value::Null, + "has_payload": record.payload.is_some(), + "payload_summary": admin_payment_payload_summary(record.payload.as_ref()), + "error_message": admin_payment_callback_error_projection(record.error_message.as_deref()), + "has_error_message": record.error_message.is_some(), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "processed_at": record.processed_at_unix_secs.and_then(unix_secs_to_rfc3339), }) } + +#[cfg(test)] +mod tests { + use super::{ + build_admin_payment_callback_payload_from_record, build_admin_payment_order_payload, + prepare_admin_payment_gateway_response_for_storage, + }; + use crate::{AdminWalletPaymentOrderRecord, GatewayAdminPaymentCallbackView}; + use serde_json::json; + + #[test] + fn admin_payment_order_projection_excludes_replayable_gateway_fields() { + let record = AdminWalletPaymentOrderRecord { + id: "order-1".to_string(), + order_no: "merchant-order-1".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "stripe".to_string(), + gateway_order_id: Some("pi_1".to_string()), + status: "pending".to_string(), + gateway_response: Some(json!({ + "gateway": "stripe", + "intent_id": "pi_1", + "client_secret": "pi_1_secret_replayable", + "payment_url": "https://pay.example/checkout?token=secret", + "payment_params": {"sign": "signed-secret"}, + "customer_email": "payer@example.com" + })), + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: None, + }; + + let payload = build_admin_payment_order_payload(&record); + assert_eq!(payload["has_gateway_response"], true); + assert_eq!( + payload.pointer("/gateway_response/gateway"), + Some(&json!("stripe")) + ); + assert_eq!( + payload.pointer("/gateway_response/intent_id"), + Some(&json!("pi_1")) + ); + for key in [ + "client_secret", + "payment_url", + "payment_params", + "customer_email", + ] { + assert!(payload + .pointer(&format!("/gateway_response/{key}")) + .is_none()); + } + } + + #[test] + fn admin_payment_order_projection_rejects_nested_or_mistyped_safe_fields() { + let mut record = AdminWalletPaymentOrderRecord { + id: "order-1".to_string(), + order_no: "merchant-order-1".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "stripe".to_string(), + gateway_order_id: Some("pi_1".to_string()), + status: "pending".to_string(), + gateway_response: None, + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: None, + }; + record.gateway_response = Some(json!({ + "gateway": {"client_secret": "secret-in-nested-object"}, + "intent_id": ["pi_1", "secret-in-array"], + "payment_method_types": ["card", {"secret": "nested"}], + "manual_credit": "secret-in-string", + })); + + let encoded = build_admin_payment_order_payload(&record).to_string(); + assert!(!encoded.contains("secret-in-nested-object")); + assert!(!encoded.contains("secret-in-array")); + assert!(!encoded.contains("nested")); + assert!(!encoded.contains("secret-in-string")); + } + + #[test] + fn admin_payment_gateway_response_is_projected_before_storage() { + let projected = prepare_admin_payment_gateway_response_for_storage(Some(json!({ + "gateway": "stripe", + "intent_id": "pi_1", + "client_secret": "pi_1_secret_replayable", + "customer": {"email": "payer@example.com"}, + "payment_params": {"authorization": "Bearer secret"}, + }))) + .expect("provided gateway response should remain present"); + + assert_eq!(projected, json!({"gateway": "stripe", "intent_id": "pi_1"})); + let encoded = projected.to_string(); + for forbidden in [ + "client_secret", + "replayable", + "customer", + "payer@example.com", + "authorization", + "Bearer secret", + ] { + assert!(!encoded.contains(forbidden), "persisted {forbidden}"); + } + } + + #[test] + fn admin_payment_callback_projection_does_not_return_raw_payload() { + let record = GatewayAdminPaymentCallbackView { + id: "callback-1".to_string(), + payment_order_id: Some("order-1".to_string()), + payment_method: "stripe".to_string(), + callback_key: "stripe:event-1".to_string(), + order_no: Some("merchant-order-1".to_string()), + gateway_order_id: Some("pi_1".to_string()), + payload_hash: Some("hash-1".to_string()), + signature_valid: true, + status: "processed".to_string(), + payload: Some(json!({ + "data": {"object": {"client_secret": "secret", "customer_email": "payer@example.com"}} + })), + error_message: None, + created_at_unix_ms: 1, + processed_at_unix_secs: Some(1), + }; + + let payload = build_admin_payment_callback_payload_from_record(&record); + assert_eq!(payload["has_payload"], true); + assert!(payload["payload"].is_null()); + assert_eq!(payload["payload_summary"]["kind"], "object"); + assert_eq!(payload["payload_summary"]["objects"], 3); + assert_eq!(payload["payload_summary"]["strings"], 2); + assert_eq!(payload["payload_summary"]["max_depth"], 4); + let encoded = payload.to_string(); + assert!(!encoded.contains("customer_email")); + assert!(!encoded.contains("payer@example.com")); + assert!(!encoded.contains("client_secret")); + assert!(!encoded.contains("secret")); + } + + #[test] + fn admin_payment_callback_projection_does_not_return_unknown_historical_errors() { + let record = GatewayAdminPaymentCallbackView { + id: "callback-1".to_string(), + payment_order_id: None, + payment_method: "stripe".to_string(), + callback_key: "stripe:event-1".to_string(), + order_no: None, + gateway_order_id: None, + payload_hash: None, + signature_valid: false, + status: "failed".to_string(), + payload: None, + error_message: Some("upstream rejected sk_live_secret_value".to_string()), + created_at_unix_ms: 1, + processed_at_unix_secs: Some(1), + }; + + let payload = build_admin_payment_callback_payload_from_record(&record); + assert_eq!(payload["has_error_message"], true); + assert_eq!( + payload["error_message"], + "payment callback processing failed" + ); + assert!(!payload.to_string().contains("sk_live_secret_value")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/billing/plans.rs b/apps/aether-gateway/src/handlers/admin/billing/plans.rs index 710bec985..696186f04 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/plans.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/plans.rs @@ -3,8 +3,12 @@ use super::{ build_admin_billing_data_unavailable_response, build_admin_billing_not_found_response, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::shared::normalize_payment_currency; use crate::{GatewayError, LocalMutationOutcome}; -use aether_data_contracts::repository::billing::{BillingPlanRecord, BillingPlanWriteInput}; +use aether_data_contracts::repository::billing::{ + checked_plan_duration_days, parse_usage_policy_entitlements, + validate_entitlement_replacement_groups, BillingPlanRecord, BillingPlanWriteInput, +}; use axum::{ body::{Body, Bytes}, http, @@ -172,11 +176,16 @@ fn validate_entitlements(value: &serde_json::Value) -> Result<(), String> { } } } + "usage_policy" => {} _ => return Err(format!("unsupported entitlement type: {kind}")), } } + validate_entitlement_replacement_groups(value).map_err(|error| error.to_string())?; + parse_usage_policy_entitlements(value).map_err(|error| error.to_string())?; if !entitlements_include_package_rights(items) { - return Err("套餐至少需要包含每日额度或会员分组;钱包充值请使用充值功能".to_string()); + return Err( + "套餐至少需要包含每日额度、会员分组或使用限制;钱包充值请使用充值功能".to_string(), + ); } Ok(()) } @@ -185,7 +194,7 @@ fn entitlements_include_package_rights(items: &[serde_json::Value]) -> bool { items.iter().any(|item| { matches!( item.get("type").and_then(|value| value.as_str()), - Some("daily_quota" | "membership_group") + Some("daily_quota" | "membership_group" | "usage_policy") ) }) } @@ -204,6 +213,7 @@ fn normalize_plan_input(payload: BillingPlanRequest) -> Result Result bool { + !value.is_empty() + && value.len() <= 128 + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) +} + +fn gateway_refund_mode_allowed(refund_mode: &str) -> bool { + refund_mode.trim().eq_ignore_ascii_case("original_channel") +} + +fn stored_refund_to_gateway( + refund: aether_data::repository::wallet::StoredAdminWalletRefund, +) -> crate::AdminWalletRefundRecord { + crate::AdminWalletRefundRecord { + id: refund.id, + refund_no: refund.refund_no, + wallet_id: refund.wallet_id, + user_id: refund.user_id, + payment_order_id: refund.payment_order_id, + source_type: refund.source_type, + source_id: refund.source_id, + refund_mode: refund.refund_mode, + amount_usd: refund.amount_usd, + status: refund.status, + reason: refund.reason, + failure_reason: refund.failure_reason, + gateway_refund_id: refund.gateway_refund_id, + payout_method: refund.payout_method, + payout_reference: refund.payout_reference, + payout_proof: refund.payout_proof, + requested_by: refund.requested_by, + approved_by: refund.approved_by, + processed_by: refund.processed_by, + created_at_unix_ms: refund.created_at_unix_ms, + updated_at_unix_secs: refund.updated_at_unix_secs, + processed_at_unix_secs: refund.processed_at_unix_secs, + completed_at_unix_secs: refund.completed_at_unix_secs, + } +} + fn merge_gateway_refund_proof( proof: Option, gateway_refund: Option<&crate::handlers::shared::DirectGatewayRefundResult>, @@ -29,14 +72,7 @@ fn merge_gateway_refund_proof( let mut object = proof .and_then(|value| value.as_object().cloned()) .unwrap_or_default(); - object.insert( - "gateway_refund".to_string(), - json!({ - "id": gateway_refund.gateway_refund_id, - "status": gateway_refund.status, - "payload": gateway_refund.payload, - }), - ); + object.insert("gateway_refund".to_string(), gateway_refund.proof.clone()); Some(Value::Object(object)) } @@ -64,7 +100,12 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( "gateway_refund_id", 128, ) { - Ok(value) => value, + Ok(value) if value.as_deref().is_none_or(is_safe_gateway_refund_id) => value, + Ok(_) => { + return Ok(build_admin_wallets_bad_request_response( + "gateway_refund_id 格式无效", + )) + } Err(detail) => return Ok(build_admin_wallets_bad_request_response(detail)), }; let payout_reference = match normalize_admin_wallet_optional_text( @@ -107,9 +148,55 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( else { return Ok(build_admin_wallet_refund_not_found_response()); }; + let refund_before_complete = stored_refund_to_gateway(refund_before_complete); + if !refund_before_complete.amount_usd.is_finite() || refund_before_complete.amount_usd <= 0.0 { + return Ok(build_admin_wallets_bad_request_response("退款金额无效")); + } + if refund_before_complete.status == "succeeded" { + if let Some(order_id) = refund_before_complete.payment_order_id.as_deref() { + if let Err(err) = state + .app() + .reverse_referral_rewards_for_order(order_id, refund_before_complete.amount_usd) + .await + { + warn!( + error = ?err, + order_id = %order_id, + refund_id = %refund_before_complete.id, + "failed to reconcile referral rewards for completed refund" + ); + return Ok(build_admin_wallets_data_unavailable_response()); + } + } + let response = Json(json!({ + "refund": build_admin_wallet_refund_payload( + &wallet, + &owner, + &refund_before_complete, + ), + })) + .into_response(); + return Ok(attach_admin_audit_response( + response, + "admin_wallet_refund_completed", + "complete_wallet_refund", + "wallet_refund", + &refund_id, + )); + } let mut gateway_refund_id = gateway_refund_id; let mut payout_proof = payload.payout_proof; if payload.gateway_refund { + // A line-item refund in `offline_payout` mode has no provider-side + // settlement contract. Calling a gateway before recording evidence + // would let `/fail` concurrently release the local reservation and + // leave an external refund with no durable proof. Keep the mode + // constraint at the boundary, before any network request. + if !gateway_refund_mode_allowed(&refund_before_complete.refund_mode) { + return Ok(build_admin_wallets_bad_request_response( + "只有原支付渠道退款可以调用支付网关", + )); + } let Some(payment_order_id) = refund_before_complete.payment_order_id.as_deref() else { return Ok(build_admin_wallets_bad_request_response( "网关原路退款需要退款申请关联支付订单", @@ -120,8 +207,10 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( crate::AdminWalletMutationOutcome::NotFound => { return Ok(build_admin_wallets_bad_request_response("支付订单不存在")) } - crate::AdminWalletMutationOutcome::Invalid(detail) => { - return Ok(build_admin_wallets_bad_request_response(detail)) + crate::AdminWalletMutationOutcome::Invalid(_) => { + return Ok(build_admin_wallets_bad_request_response( + "支付订单状态或数据无效", + )) } crate::AdminWalletMutationOutcome::Unavailable => { return Ok(build_admin_wallets_data_unavailable_response()) @@ -150,14 +239,96 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( { Ok(Some(result)) => { gateway_refund_id = Some(result.gateway_refund_id.clone()); + if result.is_pending() { + let persisted = match state + .app() + .update_admin_wallet_refund_gateway( + aether_data::repository::wallet::UpdateAdminWalletRefundGatewayInput { + wallet_id: wallet_id.clone(), + refund_id: refund_id.clone(), + gateway_refund_id: result.gateway_refund_id.clone(), + payout_proof: Some(result.proof.clone()), + }, + ) + .await? + { + Some(aether_data::repository::wallet::WalletMutationOutcome::Applied( + refund, + )) => refund, + Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound) => { + return Ok(build_admin_wallet_refund_not_found_response()) + } + Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid( + detail, + )) => return Ok(build_admin_wallets_bad_request_response(detail)), + None => return Ok(build_admin_wallets_data_unavailable_response()), + }; + let persisted = stored_refund_to_gateway(persisted); + let response = ( + http::StatusCode::ACCEPTED, + Json(json!({ + "refund": build_admin_wallet_refund_payload(&wallet, &owner, &persisted), + "gateway_refund": { + "id": result.gateway_refund_id, + "status": result.status, + }, + })), + ) + .into_response(); + return Ok(attach_admin_audit_response( + response, + "admin_wallet_refund_pending", + "complete_wallet_refund", + "wallet_refund", + &refund_id, + )); + } + if !result.is_succeeded() { + return Ok(build_admin_wallets_bad_request_response("上游退款未成功")); + } payout_proof = merge_gateway_refund_proof(payout_proof, Some(&result)); + + // Persist the provider evidence before releasing the local refund reservation. + // If the local completion transaction fails after a successful gateway call, + // a retry can reuse the idempotent gateway identifier instead of issuing a + // second refund with no durable proof of the first one. + match state + .app() + .update_admin_wallet_refund_gateway( + aether_data::repository::wallet::UpdateAdminWalletRefundGatewayInput { + wallet_id: wallet_id.clone(), + refund_id: refund_id.clone(), + gateway_refund_id: result.gateway_refund_id.clone(), + payout_proof: payout_proof.clone(), + }, + ) + .await? + { + Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(_)) => {} + Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound) => { + return Ok(build_admin_wallet_refund_not_found_response()) + } + Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid( + detail, + )) => return Ok(build_admin_wallets_bad_request_response(detail)), + None => return Ok(build_admin_wallets_data_unavailable_response()), + } } Ok(None) => { return Ok(build_admin_wallets_bad_request_response( "该支付方式不支持官方直连退款,请使用线下完成", )) } - Err(detail) => return Ok(build_admin_wallets_bad_request_response(detail)), + Err(detail) => { + warn!( + error = %detail, + refund_id = %refund_id, + "direct payment gateway refund failed" + ); + return Ok(build_admin_wallets_bad_request_response( + "支付网关退款请求失败", + )); + } } } match state @@ -183,6 +354,7 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( refund_id = %refund.id, "failed to reverse referral rewards for completed refund" ); + return Ok(build_admin_wallets_data_unavailable_response()); } } let response = Json(json!({ @@ -204,7 +376,7 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( let detail = if detail == "refund status must be processing before completion" { "只有 processing 状态的退款可以标记完成".to_string() } else { - detail + "退款状态或参数无效".to_string() }; Ok(build_admin_wallets_bad_request_response(detail)) } @@ -213,3 +385,70 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response( } } } + +#[cfg(test)] +mod tests { + use super::{ + gateway_refund_mode_allowed, is_safe_gateway_refund_id, merge_gateway_refund_proof, + }; + use crate::handlers::shared::DirectGatewayRefundResult; + use serde_json::json; + + #[test] + fn gateway_refund_merge_replaces_legacy_raw_payload() { + let existing = json!({ + "channel": "manual", + "gateway_refund": { + "payload": { + "authorization": "Bearer legacy-secret", + "payer": {"openid": "openid-secret"} + } + } + }); + let result = DirectGatewayRefundResult { + gateway_refund_id: "refund-1".to_string(), + status: "success".to_string(), + proof: json!({ + "gateway": "wxpay", + "id": "refund-1", + "status": "success", + "order_no": "order-1", + "refund_no": "request-1", + "amount": 8.5, + "currency": "CNY", + "processed_at": "2026-08-27T12:00:00Z" + }), + }; + + let merged = merge_gateway_refund_proof(Some(existing), Some(&result)) + .expect("gateway proof should be merged"); + assert_eq!(merged["channel"], "manual"); + assert_eq!(merged["gateway_refund"], result.proof); + let encoded = merged.to_string(); + assert!(!encoded.contains("legacy-secret")); + assert!(!encoded.contains("openid-secret")); + assert!(!encoded.contains("payload")); + } + + #[test] + fn manual_gateway_refund_ids_use_the_same_strict_identifier_policy() { + assert!(is_safe_gateway_refund_id("refund_123-ABC")); + for value in [ + "Authorization: Bearer secret", + "https://internal.example/refund?token=secret", + "refund id", + "payer/openid", + ] { + assert!(!is_safe_gateway_refund_id(value)); + } + assert!(!is_safe_gateway_refund_id(&"a".repeat(129))); + } + + #[test] + fn gateway_refunds_are_limited_to_original_channel_mode() { + assert!(gateway_refund_mode_allowed("original_channel")); + assert!(gateway_refund_mode_allowed(" Original_Channel ")); + assert!(!gateway_refund_mode_allowed("offline_payout")); + assert!(!gateway_refund_mode_allowed("")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/fail_refund.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/fail_refund.rs index 8a14687d0..bf9c37fa6 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/fail_refund.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/fail_refund.rs @@ -10,6 +10,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -86,7 +87,7 @@ pub(in super::super) async fn build_admin_wallet_fail_refund_response( transaction.link_id.as_deref(), transaction.operator_id.as_deref(), transaction.description.as_deref(), - unix_secs_to_rfc3339(transaction.created_at_unix_ms), + unix_secs_to_rfc3339(stored_timestamp_unix_secs(transaction.created_at_unix_ms)), ) }) .unwrap_or(serde_json::Value::Null), diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/process_refund.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/process_refund.rs index 9e1eb4111..60e6cf5a2 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/process_refund.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/process_refund.rs @@ -9,6 +9,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -71,7 +72,7 @@ pub(in super::super) async fn build_admin_wallet_process_refund_response( transaction.link_id.as_deref(), transaction.operator_id.as_deref(), transaction.description.as_deref(), - unix_secs_to_rfc3339(transaction.created_at_unix_ms), + unix_secs_to_rfc3339(stored_timestamp_unix_secs(transaction.created_at_unix_ms)), ), })) .into_response(); diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/recharge.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/recharge.rs index 140d7f750..1e2bac14a 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/recharge.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/recharge.rs @@ -10,6 +10,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -89,7 +90,7 @@ pub(in super::super) async fn build_admin_wallet_recharge_response( payment_order.amount_usd, payment_order.payment_method, payment_order.status, - unix_secs_to_rfc3339(payment_order.created_at_unix_ms), + unix_secs_to_rfc3339(stored_timestamp_unix_secs(payment_order.created_at_unix_ms)), payment_order .credited_at_unix_secs .and_then(unix_secs_to_rfc3339), diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/ledger.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/ledger.rs index dfaae88ec..5d856c377 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/ledger.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/ledger.rs @@ -6,6 +6,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -69,7 +70,10 @@ pub(in super::super) async fn build_admin_wallet_ledger_response( "operator_name": entry.operator_name, "operator_email": entry.operator_email, "description": entry.description, - "created_at": entry.created_at_unix_ms.and_then(unix_secs_to_rfc3339), + "created_at": entry + .created_at_unix_ms + .map(stored_timestamp_unix_secs) + .and_then(unix_secs_to_rfc3339), }) }) .collect::>(); diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/list.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/list.rs index 200b5e038..c8cda09a2 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/list.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/list.rs @@ -6,6 +6,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -61,7 +62,10 @@ pub(in super::super) async fn build_admin_wallet_list_response( "total_consumed": wallet.total_consumed, "total_refunded": wallet.total_refunded, "total_adjusted": wallet.total_adjusted, - "created_at": wallet.created_at_unix_ms.and_then(unix_secs_to_rfc3339), + "created_at": wallet + .created_at_unix_ms + .map(stored_timestamp_unix_secs) + .and_then(unix_secs_to_rfc3339), "updated_at": wallet.updated_at_unix_secs.and_then(unix_secs_to_rfc3339), }); enrich_admin_wallet_package_summary( diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/refund_requests.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/refund_requests.rs index 6ad1adc9b..43a5a2ddc 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/refund_requests.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/refund_requests.rs @@ -1,12 +1,13 @@ use super::super::shared::{ - build_admin_wallets_bad_request_response, parse_admin_wallets_limit, - parse_admin_wallets_offset, parse_admin_wallets_owner_type_filter, + admin_wallet_payout_proof_projection, build_admin_wallets_bad_request_response, + parse_admin_wallets_limit, parse_admin_wallets_offset, parse_admin_wallets_owner_type_filter, resolve_admin_wallet_owner_summary, wallet_owner_summary_from_fields, ADMIN_WALLETS_API_KEY_REFUND_DETAIL, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339}; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -40,6 +41,7 @@ pub(in super::super) async fn build_admin_wallet_refund_requests_response( .await?; let mut items = Vec::with_capacity(refunds.len()); for refund in refunds { + let payout_proof = admin_wallet_payout_proof_projection(refund.payout_proof.as_ref()); let mut owner = wallet_owner_summary_from_fields( refund.wallet_user_id.as_deref(), refund.wallet_user_name.clone(), @@ -75,11 +77,14 @@ pub(in super::super) async fn build_admin_wallet_refund_requests_response( "gateway_refund_id": refund.gateway_refund_id, "payout_method": refund.payout_method, "payout_reference": refund.payout_reference, - "payout_proof": refund.payout_proof, + "payout_proof": payout_proof, "requested_by": refund.requested_by, "approved_by": refund.approved_by, "processed_by": refund.processed_by, - "created_at": refund.created_at_unix_ms.and_then(unix_secs_to_rfc3339), + "created_at": refund + .created_at_unix_ms + .map(stored_timestamp_unix_secs) + .and_then(unix_secs_to_rfc3339), "updated_at": refund.updated_at_unix_secs.and_then(unix_secs_to_rfc3339), "processed_at": refund.processed_at_unix_secs.and_then(unix_secs_to_rfc3339), "completed_at": refund.completed_at_unix_secs.and_then(unix_secs_to_rfc3339), diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/transactions.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/transactions.rs index 02fce43c5..87444a54c 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/transactions.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/reads/transactions.rs @@ -6,6 +6,7 @@ use super::super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::unix_secs_to_rfc3339; use crate::GatewayError; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use axum::{ body::Body, response::{IntoResponse, Response}, @@ -81,7 +82,10 @@ pub(in super::super) async fn build_admin_wallet_transactions_response( "operator_name": operator_name, "operator_email": operator_email, "description": transaction.description, - "created_at": transaction.created_at_unix_ms.and_then(unix_secs_to_rfc3339), + "created_at": transaction + .created_at_unix_ms + .map(stored_timestamp_unix_secs) + .and_then(unix_secs_to_rfc3339), })); } diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/normalizers.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/normalizers.rs index b0f6d4ffd..e49343cac 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/normalizers.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/normalizers.rs @@ -51,11 +51,12 @@ pub(in super::super) fn normalize_admin_wallet_optional_text( pub(in super::super) fn normalize_admin_wallet_payment_method( value: String, ) -> Result { - let normalized = value.trim(); - if normalized.is_empty() { - return Err("payment_method 不能为空".to_string()); + let normalized = aether_data::repository::wallet::canonicalize_payment_method(&value) + .map_err(|detail| format!("payment_method 无效: {detail}"))?; + if normalized.chars().count() > 30 { + return Err("payment_method 长度不能超过 30".to_string()); } - Ok(normalized.chars().take(30).collect()) + Ok(normalized) } pub(in super::super) fn normalize_admin_wallet_balance_type( diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/payloads.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/payloads.rs index 50201299e..a280ff86a 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/shared/payloads.rs @@ -2,7 +2,8 @@ use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::unix_secs_to_rfc3339; use crate::handlers::shared::round_to; use crate::GatewayError; -use serde_json::json; +use aether_data::repository::wallet::stored_timestamp_unix_secs; +use serde_json::{json, Map, Value}; #[derive(Clone)] pub(in super::super) struct AdminWalletOwnerSummary { @@ -10,6 +11,10 @@ pub(in super::super) struct AdminWalletOwnerSummary { pub(in super::super) owner_name: Option, } +fn api_key_display_prefix(api_key_id: &str) -> String { + api_key_id.chars().take(8).collect() +} + pub(in super::super) fn build_admin_wallet_payment_order_payload( order_id: String, order_no: String, @@ -92,7 +97,7 @@ pub(in super::super) fn wallet_owner_summary_from_fields( owner_type: "api_key", owner_name: api_key_name .filter(|value| !value.trim().is_empty()) - .or_else(|| Some(format!("Key-{}", &api_key_id[..api_key_id.len().min(8)]))), + .or_else(|| Some(format!("Key-{}", api_key_display_prefix(api_key_id)))), }; } AdminWalletOwnerSummary { @@ -121,7 +126,7 @@ pub(in super::super) async fn resolve_admin_wallet_owner_summary( .find(|snapshot| snapshot.api_key_id == api_key_id) .and_then(|snapshot| snapshot.api_key_name) .filter(|value| !value.trim().is_empty()) - .or_else(|| Some(format!("Key-{}", &api_key_id[..api_key_id.len().min(8)]))); + .or_else(|| Some(format!("Key-{}", api_key_display_prefix(api_key_id)))); Ok(AdminWalletOwnerSummary { owner_type: "api_key", owner_name, @@ -234,6 +239,104 @@ pub(in super::super) async fn enrich_admin_wallet_package_summary( Ok(()) } +fn admin_refund_proof_identifier(value: Option<&Value>, max_bytes: usize) -> Option { + let value = value?.as_str()?.trim(); + if value.is_empty() + || value.len() > max_bytes + || !value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) + { + return None; + } + Some(value.to_string()) +} + +fn admin_gateway_refund_proof_projection(value: &Value) -> Option { + let source = value.as_object()?; + let mut projected = Map::new(); + + if let Some(gateway) = source + .get("gateway") + .and_then(Value::as_str) + .map(str::trim) + .map(str::to_ascii_lowercase) + .filter(|value| matches!(value.as_str(), "alipay" | "wxpay")) + { + projected.insert("gateway".to_string(), json!(gateway)); + } + for (key, max_bytes) in [ + ("id", 128usize), + ("order_no", 64usize), + ("refund_no", 64usize), + ] { + if let Some(value) = admin_refund_proof_identifier(source.get(key), max_bytes) { + projected.insert(key.to_string(), json!(value)); + } + } + if let Some(status) = source + .get("status") + .and_then(Value::as_str) + .map(str::trim) + .map(str::to_ascii_lowercase) + .and_then(|value| match value.as_str() { + "success" | "succeeded" => Some("success"), + "pending" | "processing" => Some("processing"), + "failed" | "closed" | "abnormal" => Some("failed"), + _ => None, + }) + { + projected.insert("status".to_string(), json!(status)); + } + if let Some(amount) = source + .get("amount") + .and_then(Value::as_f64) + .filter(|value| value.is_finite() && *value > 0.0) + { + projected.insert("amount".to_string(), json!(amount)); + } + if let Some(currency) = source + .get("currency") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| value.len() == 3 && value.bytes().all(|byte| byte.is_ascii_alphabetic())) + .map(str::to_ascii_uppercase) + { + projected.insert("currency".to_string(), json!(currency)); + } + if let Some(processed_at) = source + .get("processed_at") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| chrono::DateTime::parse_from_rfc3339(value).is_ok()) + { + projected.insert("processed_at".to_string(), json!(processed_at)); + } + + (!projected.is_empty()).then_some(Value::Object(projected)) +} + +pub(in super::super) fn admin_wallet_payout_proof_projection( + payout_proof: Option<&Value>, +) -> Option { + let source = payout_proof?.as_object()?; + let mut projected = source.clone(); + if source.contains_key("gateway_refund") { + match source + .get("gateway_refund") + .and_then(admin_gateway_refund_proof_projection) + { + Some(gateway_refund) => { + projected.insert("gateway_refund".to_string(), gateway_refund); + } + None => { + projected.remove("gateway_refund"); + } + } + } + Some(Value::Object(projected)) +} + pub(in super::super) fn build_admin_wallet_refund_payload( wallet: &aether_data::repository::wallet::StoredWalletSnapshot, owner: &AdminWalletOwnerSummary, @@ -258,13 +361,94 @@ pub(in super::super) fn build_admin_wallet_refund_payload( "gateway_refund_id": refund.gateway_refund_id.clone(), "payout_method": refund.payout_method.clone(), "payout_reference": refund.payout_reference.clone(), - "payout_proof": refund.payout_proof.clone(), + "payout_proof": admin_wallet_payout_proof_projection(refund.payout_proof.as_ref()), "requested_by": refund.requested_by.clone(), "approved_by": refund.approved_by.clone(), "processed_by": refund.processed_by.clone(), - "created_at": unix_secs_to_rfc3339(refund.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(refund.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(refund.updated_at_unix_secs), "processed_at": refund.processed_at_unix_secs.and_then(unix_secs_to_rfc3339), "completed_at": refund.completed_at_unix_secs.and_then(unix_secs_to_rfc3339), }) } + +#[cfg(test)] +mod tests { + use super::admin_wallet_payout_proof_projection; + use serde_json::json; + + #[test] + fn admin_refund_proof_projection_removes_historical_gateway_payloads() { + let proof = json!({ + "channel": "manual", + "gateway_refund": { + "gateway": "WXPAY", + "id": "refund-1", + "status": "SUCCESS", + "order_no": "order-1", + "refund_no": "request-1", + "amount": 8.5, + "currency": "cny", + "processed_at": "2026-08-27T12:00:00Z", + "payload": { + "authorization": "Bearer payment-secret", + "url": "https://internal.example/refund?token=secret", + "payer": {"openid": "openid-secret"}, + "credential": "gateway-credential" + }, + "message": "upstream secret message" + } + }); + + let projection = admin_wallet_payout_proof_projection(Some(&proof)) + .expect("object payout proof should be projected"); + assert_eq!(projection["channel"], "manual"); + assert_eq!(projection["gateway_refund"]["gateway"], "wxpay"); + assert_eq!(projection["gateway_refund"]["status"], "success"); + assert_eq!( + projection["gateway_refund"] + .as_object() + .expect("gateway proof should be an object") + .len(), + 8 + ); + let encoded = projection.to_string(); + for sensitive in [ + "payment-secret", + "?token=secret", + "openid-secret", + "gateway-credential", + "upstream secret message", + "payload", + "authorization", + "payer", + "openid", + "credential", + "message", + ] { + assert!(!encoded.contains(sensitive)); + } + } + + #[test] + fn admin_refund_proof_projection_drops_mistyped_gateway_fields() { + let proof = json!({ + "operator": "finance", + "gateway_refund": { + "gateway": {"credential": "secret"}, + "id": ["refund-1"], + "status": "unknown-secret-status", + "order_no": "https://example.test/?token=secret", + "refund_no": 123, + "amount": "8.5", + "currency": "CNY?token=secret", + "processed_at": "Bearer secret" + } + }); + + let projection = admin_wallet_payout_proof_projection(Some(&proof)) + .expect("manual payout proof should remain available"); + assert_eq!(projection, json!({"operator": "finance"})); + assert!(!projection.to_string().contains("secret")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/features/background_tasks/routes.rs b/apps/aether-gateway/src/handlers/admin/features/background_tasks/routes.rs index 71c771e86..7fa1771af 100644 --- a/apps/aether-gateway/src/handlers/admin/features/background_tasks/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/features/background_tasks/routes.rs @@ -1,13 +1,14 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{ - attach_admin_audit_response, query_param_value, unix_secs_to_rfc3339, + attach_admin_audit_response, mark_sensitive_admin_response_no_store, query_param_value, + unix_secs_to_rfc3339, }; use crate::task_runtime::{ self, set_cancel_signal, TASK_KEY_PROVIDER_DELETE, TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT, }; use crate::GatewayError; use aether_data_contracts::repository::background_tasks::{ - BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskStatus, + BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskStatus, StoredBackgroundTaskRun, }; use axum::{ body::{Body, Bytes}, @@ -21,6 +22,30 @@ const DEFAULT_PAGE_SIZE: usize = 20; const MAX_PAGE_SIZE: usize = 100; const DEFAULT_EVENTS_PAGE_SIZE: usize = 50; +fn build_background_task_list_item(run: &StoredBackgroundTaskRun) -> serde_json::Value { + json!({ + "id": run.id, + "task_key": run.task_key, + "kind": run.kind.as_database(), + "trigger": run.trigger, + "status": run.status.as_database(), + "attempt": run.attempt, + "max_attempts": run.max_attempts, + "owner_instance": run.owner_instance, + "progress_percent": run.progress_percent, + "progress_message": run.progress_message, + "has_payload": run.payload_json.is_some(), + "has_result": run.result_json.is_some(), + "has_error": run.error_message.is_some(), + "cancel_requested": run.cancel_requested, + "created_by": run.created_by, + "created_at": unix_secs_to_rfc3339(run.created_at_unix_secs), + "started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339), + "finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339), + "updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs), + }) +} + pub(super) async fn maybe_build_local_admin_background_tasks_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -71,29 +96,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response( let items = response .items .iter() - .map(|run| { - json!({ - "id": run.id, - "task_key": run.task_key, - "kind": run.kind.as_database(), - "trigger": run.trigger, - "status": run.status.as_database(), - "attempt": run.attempt, - "max_attempts": run.max_attempts, - "owner_instance": run.owner_instance, - "progress_percent": run.progress_percent, - "progress_message": run.progress_message, - "payload": run.payload_json, - "result": run.result_json, - "error_message": run.error_message, - "cancel_requested": run.cancel_requested, - "created_by": run.created_by, - "created_at": unix_secs_to_rfc3339(run.created_at_unix_secs), - "started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339), - "finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339), - "updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs), - }) - }) + .map(build_background_task_list_item) .collect::>(); let definitions = task_runtime::task_definitions() .iter() @@ -153,8 +156,9 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response( .into_response(), )); }; - return Ok(Some(attach_admin_audit_response( - Json(json!({ + return Ok(Some(mark_sensitive_admin_response_no_store( + attach_admin_audit_response( + Json(json!({ "id": run.id, "task_key": run.task_key, "kind": run.kind.as_database(), @@ -174,12 +178,13 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response( "started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339), "finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339), "updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs), - })) - .into_response(), - "admin_task_detail_viewed", - "view_task_detail", - "background_task", - run_id, + })) + .into_response(), + "admin_task_detail_viewed", + "view_task_detail", + "background_task", + run_id, + ), ))); } Some("events") if request_context.method() == http::Method::GET => { @@ -205,7 +210,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response( let events = state .list_background_task_events(run_id, offset, page_size) .await?; - return Ok(Some( + return Ok(Some(mark_sensitive_admin_response_no_store( Json(json!({ "items": events.into_iter().map(|event| { json!({ @@ -221,7 +226,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response( "page_size": page_size, })) .into_response(), - )); + ))); } Some("cancel") if request_context.method() == http::Method::POST => { let Some(run_id) = nested_task_id_from_path(request_context.path(), "/cancel") else { @@ -371,3 +376,55 @@ fn parse_json_payload(request_body: Option<&Bytes>) -> Result(body) .map_err(|err| GatewayError::Internal(format!("invalid json body: {err}"))) } + +#[cfg(test)] +mod tests { + use super::build_background_task_list_item; + use aether_data_contracts::repository::background_tasks::{ + BackgroundTaskKind, BackgroundTaskStatus, StoredBackgroundTaskRun, + }; + use serde_json::json; + + #[test] + fn task_list_item_exposes_only_safe_diagnostic_presence_flags() { + let run = StoredBackgroundTaskRun { + id: "run-1".to_string(), + task_key: "provider.oauth.import".to_string(), + kind: BackgroundTaskKind::OnDemand, + trigger: "manual".to_string(), + status: BackgroundTaskStatus::Failed, + attempt: 1, + max_attempts: 3, + owner_instance: Some("gateway-1".to_string()), + progress_percent: 100, + progress_message: Some("task failed".to_string()), + payload_json: Some(json!({"refresh_token": "secret-refresh-token"})), + result_json: Some(json!({"access_token": "secret-access-token"})), + error_message: Some("upstream error containing secret-api-key".to_string()), + cancel_requested: false, + created_by: Some("admin".to_string()), + created_at_unix_secs: 1, + started_at_unix_secs: Some(2), + finished_at_unix_secs: Some(3), + updated_at_unix_secs: 3, + }; + + let item = build_background_task_list_item(&run); + assert_eq!(item["status"], "failed"); + assert_eq!(item["has_payload"], true); + assert_eq!(item["has_result"], true); + assert_eq!(item["has_error"], true); + assert!(item.get("payload").is_none()); + assert!(item.get("result").is_none()); + assert!(item.get("error_message").is_none()); + + let serialized = item.to_string(); + for secret in [ + "secret-refresh-token", + "secret-access-token", + "secret-api-key", + ] { + assert!(!serialized.contains(secret)); + } + } +} diff --git a/apps/aether-gateway/src/handlers/admin/features/gemini_files/read_routes.rs b/apps/aether-gateway/src/handlers/admin/features/gemini_files/read_routes.rs index 1fcee2fac..71a183af7 100644 --- a/apps/aether-gateway/src/handlers/admin/features/gemini_files/read_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/features/gemini_files/read_routes.rs @@ -173,6 +173,7 @@ pub(super) async fn maybe_build_local_admin_gemini_files_read_response( let mappings = state .list_gemini_file_mappings( &aether_data::repository::gemini_file_mappings::GeminiFileMappingListQuery { + user_id: None, include_expired: page.include_expired, search: page.search.clone(), offset: (page.page - 1).saturating_mul(page.page_size), diff --git a/apps/aether-gateway/src/handlers/admin/features/gemini_files/upload/request.rs b/apps/aether-gateway/src/handlers/admin/features/gemini_files/upload/request.rs index d2cf2d1a5..c9ed970e2 100644 --- a/apps/aether-gateway/src/handlers/admin/features/gemini_files/upload/request.rs +++ b/apps/aether-gateway/src/handlers/admin/features/gemini_files/upload/request.rs @@ -1,4 +1,11 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::shared::{ + find_multipart_boundary, find_multipart_boundary_after_crlf, parse_multipart_boundary, + MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES, +}; +use aether_data_contracts::repository::gemini_file_mappings::{ + GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS, +}; use axum::body::Bytes; use base64::Engine as _; @@ -21,7 +28,8 @@ pub(super) fn admin_gemini_files_parse_upload_request( .map(str::trim) .filter(|value| !value.is_empty()) .ok_or_else(|| "Content-Type 缺失".to_string())?; - let boundary = admin_gemini_files_multipart_boundary(content_type)?; + let boundary = parse_multipart_boundary(content_type) + .ok_or_else(|| "multipart boundary 缺失或无效".to_string())?; let body = request_body .filter(|body| !body.is_empty()) .ok_or_else(|| "上传文件不能为空".to_string())?; @@ -35,47 +43,32 @@ pub(super) fn admin_gemini_files_parse_upload_request( }) } -fn admin_gemini_files_multipart_boundary(content_type: &str) -> Result { - let normalized = content_type.trim(); - if !normalized - .to_ascii_lowercase() - .starts_with("multipart/form-data") - { - return Err("Content-Type 必须是 multipart/form-data".to_string()); - } - for part in normalized.split(';').skip(1) { - let Some((key, value)) = part.trim().split_once('=') else { - continue; - }; - if !key.trim().eq_ignore_ascii_case("boundary") { - continue; - } - let boundary = value.trim().trim_matches('"').trim(); - if !boundary.is_empty() { - return Ok(boundary.to_string()); - } - } - Err("multipart boundary 缺失".to_string()) -} - fn admin_gemini_files_extract_file_part( body: &[u8], boundary: &str, ) -> Result<(String, String, Vec), String> { let boundary_marker = format!("--{boundary}"); - let next_boundary_marker = format!("\r\n--{boundary}"); let boundary_bytes = boundary_marker.as_bytes(); - let next_boundary_bytes = next_boundary_marker.as_bytes(); let mut cursor = 0usize; + let mut part_count = 0usize; + let mut file_part = None; while cursor < body.len() { - if !body[cursor..].starts_with(boundary_bytes) { + if find_multipart_boundary(&body[cursor..], boundary_bytes) != Some(0) { return Err("multipart body 格式无效".to_string()); } cursor += boundary_bytes.len(); if body[cursor..].starts_with(b"--") { + let closing_suffix = body.get(cursor + 2..).unwrap_or_default(); + if !(closing_suffix.is_empty() || closing_suffix.starts_with(b"\r\n")) { + return Err("multipart 结束边界格式无效".to_string()); + } break; } + part_count = part_count.saturating_add(1); + if part_count > MAX_MULTIPART_PARTS { + return Err("multipart part 数量超过上限".to_string()); + } if !body[cursor..].starts_with(b"\r\n") { return Err("multipart body 缺少头部分隔符".to_string()); } @@ -85,34 +78,60 @@ fn admin_gemini_files_extract_file_part( return Err("multipart part 缺少头部".to_string()); }; let headers_end = cursor + headers_end_rel; + if headers_end_rel > MAX_MULTIPART_PART_HEADER_BYTES { + return Err("multipart part 头部超过大小上限".to_string()); + } let headers_text = std::str::from_utf8(&body[cursor..headers_end]) .map_err(|_| "multipart part 头部编码无效".to_string())?; cursor = headers_end + 4; let Some(next_boundary_rel) = - admin_gemini_files_find_subslice(&body[cursor..], next_boundary_bytes) + find_multipart_boundary_after_crlf(&body[cursor..], boundary_bytes) else { return Err("multipart body 缺少结束边界".to_string()); }; let content_end = cursor + next_boundary_rel; - let content = &body[cursor..content_end]; - cursor = content_end + 2; + // The CRLF immediately before the delimiter belongs to the + // multipart framing, not to the uploaded file bytes. + let content = body[cursor..content_end] + .strip_suffix(b"\r\n") + .unwrap_or(&body[cursor..content_end]); + cursor = content_end; let Some((field_name, file_name, mime_type)) = admin_gemini_files_parse_part_headers(headers_text) else { - continue; + return Err("multipart part 头部无效".to_string()); }; if field_name != "file" { continue; } - return Ok(( + if file_part.is_some() { + return Err("multipart body 包含多个 file 字段".to_string()); + } + file_part = Some(( file_name.unwrap_or_else(|| "uploaded-file".to_string()), mime_type.unwrap_or_else(|| "application/octet-stream".to_string()), content.to_vec(), )); } - Err("multipart body 中缺少 file 字段".to_string()) + let (display_name, mime_type, content) = + file_part.ok_or_else(|| "multipart body 中缺少 file 字段".to_string())?; + if display_name + .chars() + .nth(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS) + .is_some() + { + return Err("上传文件名超过长度上限".to_string()); + } + if mime_type + .chars() + .nth(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS) + .is_some() + { + return Err("上传文件 Content-Type 超过长度上限".to_string()); + } + Ok((display_name, mime_type, content)) } fn admin_gemini_files_parse_part_headers( @@ -121,27 +140,29 @@ fn admin_gemini_files_parse_part_headers( let mut field_name = None; let mut file_name = None; let mut mime_type = None; + let mut disposition_seen = false; + let mut content_type_seen = false; for line in headers_text.split("\r\n") { - let Some((header_name, header_value)) = line.split_once(':') else { - continue; - }; + let (header_name, header_value) = line.split_once(':')?; let header_name = header_name.trim(); let header_value = header_value.trim(); if header_name.eq_ignore_ascii_case("content-disposition") { - for part in header_value.split(';').skip(1) { - let Some((key, value)) = part.trim().split_once('=') else { - continue; - }; - let key = key.trim(); - let value = value.trim().trim_matches('"').trim(); - if key.eq_ignore_ascii_case("name") && !value.is_empty() { - field_name = Some(value.to_string()); - } else if key.eq_ignore_ascii_case("filename") && !value.is_empty() { - file_name = Some(value.to_string()); - } + if disposition_seen { + return None; } - } else if header_name.eq_ignore_ascii_case("content-type") && !header_value.is_empty() { + disposition_seen = true; + let (name, filename) = admin_gemini_files_parse_content_disposition(header_value)?; + field_name = Some(name); + file_name = filename; + } else if header_name.eq_ignore_ascii_case("content-type") { + if content_type_seen + || header_value.is_empty() + || header_value.chars().any(char::is_control) + { + return None; + } + content_type_seen = true; mime_type = Some(header_value.to_string()); } } @@ -149,6 +170,150 @@ fn admin_gemini_files_parse_part_headers( field_name.map(|field_name| (field_name, file_name, mime_type)) } +fn admin_gemini_files_parse_content_disposition(value: &str) -> Option<(String, Option)> { + let segments = admin_gemini_files_split_header_parameters(value)?; + if !segments.first()?.trim().eq_ignore_ascii_case("form-data") { + return None; + } + + let mut seen_keys = Vec::new(); + let mut name = None; + let mut filename = None; + for segment in segments.into_iter().skip(1) { + let segment = segment.trim(); + if segment.is_empty() { + return None; + } + let (raw_key, raw_value) = segment.split_once('=')?; + let key = raw_key.trim(); + if key.is_empty() + || !key + .as_bytes() + .iter() + .copied() + .all(admin_gemini_files_is_token_byte) + { + return None; + } + if seen_keys + .iter() + .any(|seen: &String| seen.eq_ignore_ascii_case(key)) + { + return None; + } + seen_keys.push(key.to_ascii_lowercase()); + + let parsed_value = admin_gemini_files_parse_parameter_value(raw_value.trim())?; + if key.eq_ignore_ascii_case("name") { + if parsed_value.is_empty() { + return None; + } + name = Some(parsed_value); + } else if key.eq_ignore_ascii_case("filename") { + if !parsed_value.is_empty() { + filename = Some(parsed_value); + } + } + } + + Some((name?, filename)) +} + +fn admin_gemini_files_split_header_parameters(value: &str) -> Option> { + let mut segments = Vec::new(); + let mut start = 0usize; + let mut in_quotes = false; + let mut escaped = false; + + for (index, byte) in value.as_bytes().iter().copied().enumerate() { + if in_quotes { + if escaped { + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + in_quotes = false; + } + } else if byte == b'"' { + in_quotes = true; + } else if byte == b';' { + segments.push(&value[start..index]); + start = index + 1; + } + } + + if in_quotes || escaped { + return None; + } + segments.push(&value[start..]); + Some(segments) +} + +fn admin_gemini_files_parse_parameter_value(value: &str) -> Option { + if value.is_empty() { + return None; + } + if value.starts_with('"') { + if value.len() < 2 || !value.ends_with('"') { + return None; + } + let inner = &value[1..value.len() - 1]; + let mut parsed = String::with_capacity(inner.len()); + let mut escaped = false; + for character in inner.chars() { + if escaped { + if character.is_control() { + return None; + } + parsed.push(character); + escaped = false; + } else if character == '\\' { + escaped = true; + } else { + if character == '"' || character.is_control() { + return None; + } + parsed.push(character); + } + } + if escaped { + return None; + } + return Some(parsed); + } + + value + .as_bytes() + .iter() + .copied() + .all(admin_gemini_files_is_token_byte) + .then(|| value.to_string()) +} + +fn admin_gemini_files_is_token_byte(byte: u8) -> bool { + matches!( + byte, + b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + | b'!' + | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) +} + fn admin_gemini_files_find_subslice(haystack: &[u8], needle: &[u8]) -> Option { if haystack.is_empty() || needle.is_empty() || haystack.len() < needle.len() { return None; @@ -157,3 +322,224 @@ fn admin_gemini_files_find_subslice(haystack: &[u8], needle: &[u8]) -> Option= 400 { return Err(admin_gemini_files_execution_error_message(&result)); } @@ -241,7 +250,7 @@ async fn admin_gemini_files_upload_single_key( .or(Some(upload.mime_type.as_str())), ) .await - .map_err(|err| format!("上传成功但本地映射写入失败: {err:?}"))?; + .map_err(|_| "上传成功但本地映射写入失败".to_string())?; Ok(success) } @@ -261,7 +270,8 @@ fn admin_gemini_files_execution_json_body(result: &ExecutionResult) -> Option Option Strin .map(str::trim) .filter(|value| !value.is_empty()) { - return message.to_string(); + return bound_gemini_files_error_message(message); } if let Some(message) = body_json .get("message") @@ -338,7 +364,7 @@ fn admin_gemini_files_execution_error_message(result: &ExecutionResult) -> Strin .map(str::trim) .filter(|value| !value.is_empty()) { - return message.to_string(); + return bound_gemini_files_error_message(message); } } if let Some(error) = result @@ -347,7 +373,112 @@ fn admin_gemini_files_execution_error_message(result: &ExecutionResult) -> Strin .map(|error| error.message.trim()) .filter(|value| !value.is_empty()) { - return error.to_string(); + return bound_gemini_files_error_message(error); } format!("上传失败,状态码 {}", result.status_code) } + +fn bound_gemini_files_error_message(value: &str) -> String { + let value = value.trim(); + let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES)); + value[..end].to_string() +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use aether_contracts::{ExecutionResult, ResponseBody}; + use aether_data_contracts::repository::gemini_file_mappings::{ + GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS, + GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS, + }; + + use super::super::request::AdminGeminiFilesUploadRequest; + use super::{ + admin_gemini_files_execution_error_message, admin_gemini_files_execution_json_body, + admin_gemini_files_upload_success_from_body, + }; + + fn sample_upload() -> AdminGeminiFilesUploadRequest { + AdminGeminiFilesUploadRequest { + display_name: "fallback.bin".to_string(), + mime_type: "application/octet-stream".to_string(), + body_bytes: vec![1], + body_bytes_b64: "AQ==".to_string(), + } + } + + #[test] + fn gemini_upload_result_rejects_file_name_beyond_storage_limit() { + let body = serde_json::json!({ + "file": {"name": "n".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS + 1)} + }); + + assert!(admin_gemini_files_upload_success_from_body(&body, &sample_upload()).is_none()); + } + + #[test] + fn gemini_upload_result_ignores_oversized_optional_metadata() { + let body = serde_json::json!({ + "file": { + "name": "files/safe", + "displayName": "d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS + 1), + "mimeType": "m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS + 1), + } + }); + + let success = admin_gemini_files_upload_success_from_body(&body, &sample_upload()) + .expect("valid file name should remain usable"); + assert_eq!(success.display_name.as_deref(), Some("fallback.bin")); + assert_eq!( + success.mime_type.as_deref(), + Some("application/octet-stream") + ); + } + + #[test] + fn gemini_upload_result_rejects_oversized_base64_before_decode() { + let encoded_limit = + crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + super::MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES, + ); + let result = ExecutionResult { + request_id: "gemini-upload-oversized".to_string(), + candidate_id: None, + status_code: 200, + headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), + response_observation: None, + body: Some(ResponseBody { + json_body: None, + body_bytes_b64: Some("A".repeat(encoded_limit + 1)), + }), + telemetry: None, + error: None, + }; + + assert!(admin_gemini_files_execution_json_body(&result).is_none()); + } + + #[test] + fn gemini_upload_error_message_is_bounded_without_splitting_utf8() { + let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES)); + let result = ExecutionResult { + request_id: "gemini-upload-oversized-error".to_string(), + candidate_id: None, + status_code: 500, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(serde_json::json!({"error": {"message": message}})), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }; + + let detail = admin_gemini_files_execution_error_message(&result); + assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES); + assert!(detail.bytes().all(|byte| byte == b'x')); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/features/video_tasks/builders.rs b/apps/aether-gateway/src/handlers/admin/features/video_tasks/builders.rs index d1ab06b63..e830787ee 100644 --- a/apps/aether-gateway/src/handlers/admin/features/video_tasks/builders.rs +++ b/apps/aether-gateway/src/handlers/admin/features/video_tasks/builders.rs @@ -22,6 +22,20 @@ pub(super) fn admin_video_task_status_name(status: VideoTaskStatus) -> &'static } } +pub(super) fn admin_video_task_error_projection(task: &StoredVideoTask) -> Option { + if task.error_message.as_deref().is_none_or(str::is_empty) + && task.error_code.as_deref().is_none_or(str::is_empty) + { + return None; + } + task.error_code + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| Some("provider_error".to_string())) +} + pub(super) fn admin_video_task_timestamp(unix_secs: Option) -> Option { unix_secs.and_then(|value| { chrono::DateTime::::from_timestamp(value as i64, 0) @@ -91,7 +105,7 @@ pub(super) fn build_admin_video_task_list_item( "aspect_ratio": task.aspect_ratio, "video_url": task.video_url, "error_code": task.error_code, - "error_message": task.error_message, + "error_message": admin_video_task_error_projection(task), "poll_count": task.poll_count, "max_poll_count": task.max_poll_count, "created_at": admin_video_task_timestamp(Some(task.created_at_unix_ms)), @@ -127,3 +141,65 @@ pub(super) fn current_admin_video_task_unix_secs() -> u64 { .unwrap_or_default() .as_secs() } + +#[cfg(test)] +mod tests { + use super::{admin_video_task_error_projection, build_admin_video_task_list_item}; + use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus}; + use std::collections::BTreeMap; + + fn failed_task() -> StoredVideoTask { + StoredVideoTask::new( + "task-1".to_string(), + None, + "request-1".to_string(), + Some("user-1".to_string()), + None, + Some("alice".to_string()), + None, + None, + Some("provider-1".to_string()), + Some("endpoint-1".to_string()), + Some("key-1".to_string()), + Some("openai:video".to_string()), + Some("openai:video".to_string()), + false, + Some("video-model".to_string()), + None, + None, + None, + None, + None, + None, + VideoTaskStatus::Failed, + 100, + None, + 0, + 10, + None, + 1, + 10, + 1, + None, + Some(2), + 2, + Some("authentication_error".to_string()), + Some("Authorization: Bearer live-secret at https://api.example?key=secret".to_string()), + None, + None, + ) + .expect("task") + } + + #[test] + fn video_task_payload_does_not_return_historical_raw_error_text() { + let task = failed_task(); + assert_eq!( + admin_video_task_error_projection(&task).as_deref(), + Some("authentication_error") + ); + let payload = build_admin_video_task_list_item(&task, &BTreeMap::new()); + assert_eq!(payload["error_message"], "authentication_error"); + assert!(!payload.to_string().contains("live-secret")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/features/video_tasks/routes.rs b/apps/aether-gateway/src/handlers/admin/features/video_tasks/routes.rs index ba8338b3b..e28335495 100644 --- a/apps/aether-gateway/src/handlers/admin/features/video_tasks/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/features/video_tasks/routes.rs @@ -15,9 +15,10 @@ use axum::{ use serde_json::json; use super::builders::{ - admin_video_task_detail_id_from_path, admin_video_task_nested_id_from_path, - admin_video_task_status_name, admin_video_task_timestamp, build_admin_video_task_list_item, - build_admin_video_task_provider_names, current_admin_video_task_unix_secs, + admin_video_task_detail_id_from_path, admin_video_task_error_projection, + admin_video_task_nested_id_from_path, admin_video_task_status_name, admin_video_task_timestamp, + build_admin_video_task_list_item, build_admin_video_task_provider_names, + current_admin_video_task_unix_secs, }; pub(super) async fn maybe_build_local_admin_video_tasks_response( @@ -259,7 +260,10 @@ pub(super) async fn maybe_build_local_admin_video_tasks_response( payload.insert("stored_video_path".to_string(), serde_json::Value::Null); payload.insert("storage_provider".to_string(), serde_json::Value::Null); payload.insert("error_code".to_string(), json!(task.error_code)); - payload.insert("error_message".to_string(), json!(task.error_message)); + payload.insert( + "error_message".to_string(), + json!(admin_video_task_error_projection(&task)), + ); payload.insert("retry_count".to_string(), json!(task.retry_count)); payload.insert("max_retries".to_string(), serde_json::Value::Null); payload.insert( diff --git a/apps/aether-gateway/src/handlers/admin/mod.rs b/apps/aether-gateway/src/handlers/admin/mod.rs index 0626429e8..5c118a2cd 100644 --- a/apps/aether-gateway/src/handlers/admin/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/mod.rs @@ -41,6 +41,7 @@ pub(crate) use self::provider::oauth::runtime::{ refresh_provider_oauth_account_state_after_update, }; pub(crate) use self::provider::ops::providers::actions::admin_provider_ops_local_action_response; +pub(crate) use self::provider::ops::providers::admin_provider_ops_credential_snapshot; pub(crate) use self::provider::ops::providers::store_admin_provider_ops_balance_cache; pub(crate) use self::provider::pool::config::admin_provider_pool_config; pub(crate) use self::provider::pool_admin::maybe_build_local_admin_pool_response; @@ -53,7 +54,7 @@ pub(crate) use self::provider::{ }; pub(crate) use self::request::{ AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError, - AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult, + AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult, SystemExportMode, }; pub(crate) use self::routes::maybe_build_local_admin_response; #[cfg(test)] @@ -61,3 +62,7 @@ pub(crate) use self::system::{ clear_proxy_node_references_with_cache_failure_for_tests, override_proxy_connectivity_probe_url_for_tests, }; +pub(crate) use self::system::{ + execute_admin_system_import_exclusively, release_admin_system_import_lease, + try_acquire_admin_system_import_lease, AdminSystemImportLockError, +}; diff --git a/apps/aether-gateway/src/handlers/admin/model/external_cache.rs b/apps/aether-gateway/src/handlers/admin/model/external_cache.rs index 92a16c846..44f14e435 100644 --- a/apps/aether-gateway/src/handlers/admin/model/external_cache.rs +++ b/apps/aether-gateway/src/handlers/admin/model/external_cache.rs @@ -10,6 +10,7 @@ use axum::http; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use std::collections::BTreeMap; +use std::net::{IpAddr, SocketAddr}; use std::time::Duration; use tracing::warn; @@ -19,10 +20,18 @@ const ADMIN_EXTERNAL_MODELS_CACHE_VERSION: u8 = 2; const ADMIN_EXTERNAL_MODELS_CACHE_TTL_SECS: u64 = 15 * 60; const ADMIN_EXTERNAL_MODELS_SOURCE_URL_ENV: &str = "AETHER_GATEWAY_EXTERNAL_MODELS_URL"; const ADMIN_EXTERNAL_MODELS_SOURCE_URL_DEFAULT: &str = "https://models.dev/api.json"; +const ADMIN_EXTERNAL_MODELS_OFFICIAL_HOST: &str = "models.dev"; +const ADMIN_EXTERNAL_MODELS_OFFICIAL_PATH: &str = "/api.json"; pub(in crate::handlers::admin) const ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY: &str = "external_models_proxy_node_id"; const ADMIN_EXTERNAL_MODELS_CONNECT_TIMEOUT_MS: u64 = 10_000; -const ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS: u64 = 300_000; +const ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS: u64 = 30_000; +const ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES: usize = 8 * 1024 * 1024; +// Keep the cache envelope bounded independently of the upstream body limit. +// Normalization adds a small amount of metadata, while a corrupted/shared +// runtime KV value must never be allowed to drive an unbounded serde +// allocation during cache reads. +const ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES: usize = 16 * 1024 * 1024; pub(crate) const ADMIN_EXTERNAL_MODELS_CONFIG_MUTATION_LOCK_KEY: &str = "admin:external_models_proxy_node_config:mutation"; const ADMIN_EXTERNAL_MODELS_CONFIG_MUTATION_LOCK_TTL: Duration = Duration::from_secs(10 * 60); @@ -34,6 +43,13 @@ struct AdminExternalModelsCacheEnvelope { payload: Value, } +#[derive(Debug)] +struct ResolvedAdminExternalModelsSource { + url: url::Url, + host: String, + addresses: Vec, +} + #[cfg(test)] pub(crate) struct AdminExternalModelsSourceUrlEnvGuard { previous: Option, @@ -76,13 +92,139 @@ fn admin_external_models_source_url() -> String { .unwrap_or_else(|| ADMIN_EXTERNAL_MODELS_SOURCE_URL_DEFAULT.to_string()) } +fn parse_admin_external_models_source_url( + raw_url: &str, + allow_insecure_test_target: bool, +) -> Result<(url::Url, String, u16), GatewayError> { + let url = url::Url::parse(raw_url) + .map_err(|_| GatewayError::Internal("external models source URL is invalid".to_string()))?; + let allowed_scheme = + url.scheme() == "https" || (allow_insecure_test_target && url.scheme() == "http"); + if !allowed_scheme + || !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + { + return Err(GatewayError::Internal( + "external models source must be an HTTPS URL without credentials, query, or fragment" + .to_string(), + )); + } + let host = url.host_str().map(ToOwned::to_owned).ok_or_else(|| { + GatewayError::Internal("external models source is missing a host".to_string()) + })?; + let port = url.port_or_known_default().ok_or_else(|| { + GatewayError::Internal("external models source is missing a port".to_string()) + })?; + Ok((url, host, port)) +} + +fn validate_admin_external_models_source_addresses( + url: &url::Url, + addresses: &[SocketAddr], + allow_insecure_test_target: bool, +) -> Result<(), GatewayError> { + if addresses.is_empty() { + return Err(GatewayError::Internal( + "external models source DNS resolution returned no addresses".to_string(), + )); + } + if !allow_insecure_test_target + && addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(is_official_external_models_catalog_url(url) + && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) + { + return Err(GatewayError::Internal( + "external models source resolves to a private or reserved address".to_string(), + )); + } + Ok(()) +} + +fn is_official_external_models_catalog_url(url: &url::Url) -> bool { + url.scheme() == "https" + && url + .host_str() + .is_some_and(|host| host.eq_ignore_ascii_case(ADMIN_EXTERNAL_MODELS_OFFICIAL_HOST)) + && url.port_or_known_default() == Some(443) + && url.path() == ADMIN_EXTERNAL_MODELS_OFFICIAL_PATH + && url.username().is_empty() + && url.password().is_none() + && url.query().is_none() + && url.fragment().is_none() +} + +async fn resolve_admin_external_models_source( + raw_url: &str, + allow_insecure_test_target: bool, +) -> Result { + let (url, host, port) = + parse_admin_external_models_source_url(raw_url, allow_insecure_test_target)?; + let addresses = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits( + host.as_str(), + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| { + GatewayError::Internal("external models source DNS resolution failed".to_string()) + })? + }; + validate_admin_external_models_source_addresses(&url, &addresses, allow_insecure_test_target)?; + Ok(ResolvedAdminExternalModelsSource { + url, + host, + addresses, + }) +} + +fn build_admin_external_models_direct_client( + source: &ResolvedAdminExternalModelsSource, +) -> Result { + let mut builder = aether_http::apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &aether_http::HttpClientConfig { + connect_timeout_ms: Some(ADMIN_EXTERNAL_MODELS_CONNECT_TIMEOUT_MS), + request_timeout_ms: Some(ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS), + http2_adaptive_window: true, + ..aether_http::HttpClientConfig::default() + }, + ); + if source.host.parse::().is_err() { + builder = builder.resolve_to_addrs(&source.host, &source.addresses); + } + builder.build().map_err(|_| { + GatewayError::Internal("external models HTTP client initialization failed".to_string()) + }) +} + fn normalize_admin_external_models_payload(payload: serde_json::Value) -> serde_json::Value { mark_external_models_official_providers(&payload).unwrap_or(payload) } fn classify_admin_external_models_transport_error(message: &str) -> &'static str { let message = message.to_ascii_lowercase(); - if message.contains("timed out") || message.contains("timeout") { + if message.contains("dns resolution") + || message.contains("dns lookup") + || (message.contains("resolve") && message.contains("host")) + { + "dns_resolution" + } else if message.contains("private or reserved") + || message.contains("ssrf") + || message.contains("address policy") + { + "ssrf_blocked" + } else if message.contains("invalid") && message.contains("url") { + "invalid_url" + } else if message.contains("timed out") || message.contains("timeout") { "timeout" } else if message.contains("relay") || message.contains("tunnel") { "relay" @@ -96,6 +238,8 @@ fn classify_admin_external_models_transport_error(message: &str) -> &'static str "response_decode" } else if message.contains("connect") || message.contains("dns") || message.contains("tcp") { "connect" + } else if message.contains("source returned http") || message.contains("status ") { + "upstream_http" } else if message.contains("header") || message.contains("method") || message.contains("build") { "request_build" @@ -116,6 +260,11 @@ async fn store_admin_external_models_cache( }; let serialized = serde_json::to_string(&envelope).map_err(|err| GatewayError::Internal(err.to_string()))?; + if serialized.len() > ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES { + return Err(GatewayError::Internal( + "external models cache envelope exceeds the allowed size".to_string(), + )); + } state .as_ref() .runtime_kv_setex( @@ -127,6 +276,13 @@ async fn store_admin_external_models_cache( Ok(()) } +fn parse_admin_external_models_cache(raw: &str) -> Option { + if raw.len() > ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES { + return None; + } + serde_json::from_str::(raw).ok() +} + fn normalize_admin_external_models_proxy_node_id( value: Option<&Value>, ) -> Result, GatewayError> { @@ -345,8 +501,15 @@ async fn fetch_admin_external_models_from_source( request_id: &str, proxy_node_id: Option<&str>, ) -> Result { - let url = admin_external_models_source_url(); + let source_url = admin_external_models_source_url(); + // A proxy/tunnel resolves the target in its own network namespace. Do not + // resolve it locally first: local DNS may be unavailable, may intentionally + // return synthetic addresses, or may not be able to see an internal target + // that the configured proxy can reach. URL shape is still validated below, + // and the direct path retains the local DNS validation/pinning guard. + let parsed_source = parse_admin_external_models_source_url(&source_url, cfg!(test))?; if let Some(node_id) = proxy_node_id { + let (source_url, _host, _port) = parsed_source; let Some(proxy) = state.resolve_admin_proxy_node_snapshot(Some(node_id)).await else { warn!( request_id = %request_id, @@ -374,7 +537,7 @@ async fn fetch_admin_external_models_from_source( ), ( EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), - "true".to_string(), + "false".to_string(), ), ]); let plan = ExecutionPlan { @@ -385,7 +548,7 @@ async fn fetch_admin_external_models_from_source( endpoint_id: String::new(), key_id: String::new(), method: http::Method::GET.as_str().to_string(), - url, + url: source_url.to_string(), headers, content_type: None, content_encoding: None, @@ -409,8 +572,12 @@ async fn fetch_admin_external_models_from_source( ..ExecutionTimeouts::default() }), }; + let bounded_plan = crate::execution_runtime::transport::with_upstream_response_body_limit( + &plan, + ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES, + ); let result = match state - .execute_execution_runtime_sync_plan(Some(request_id), &plan) + .execute_execution_runtime_sync_plan(Some(request_id), &bounded_plan) .await { Ok(result) => result, @@ -444,19 +611,65 @@ async fn fetch_admin_external_models_from_source( return Ok(normalize_admin_external_models_payload(payload)); } - let response = state - .http_client() - .get(&url) + let (url, host, port) = parsed_source; + let addresses = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits( + host.as_str(), + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| { + GatewayError::Internal("external models source DNS resolution failed".to_string()) + })? + }; + validate_admin_external_models_source_addresses(&url, &addresses, cfg!(test))?; + let source = ResolvedAdminExternalModelsSource { + url, + host, + addresses, + }; + + let client = build_admin_external_models_direct_client(&source)?; + let response = client + .get(source.url) + .header(reqwest::header::ACCEPT, "application/json") + .header( + reqwest::header::USER_AGENT, + "aether-gateway/external-models", + ) .send() .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; - let response = response - .error_for_status() - .map_err(|err| GatewayError::Internal(err.to_string()))?; - let payload = response - .json::() - .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + .map_err(|err| { + let error_message = err.to_string(); + let transport_error_kind = + classify_admin_external_models_transport_error(&error_message); + warn!( + request_id = %request_id, + transport_error_kind, + "external models direct request failed" + ); + GatewayError::Internal("external models source request failed".to_string()) + })?; + if !response.status().is_success() { + return Err(GatewayError::Internal(format!( + "external models source returned HTTP {}", + response.status().as_u16() + ))); + } + let body = aether_http::read_response_bytes_with_limit( + response, + ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES, + ) + .await + .map_err(|_| { + GatewayError::Internal("external models source response read failed".to_string()) + })?; + let payload = serde_json::from_slice::(&body).map_err(|_| { + GatewayError::Internal("external models source returned invalid JSON".to_string()) + })?; Ok(normalize_admin_external_models_payload(payload)) } @@ -470,8 +683,8 @@ pub(crate) async fn read_admin_external_models_cache( .runtime_kv_get(ADMIN_EXTERNAL_MODELS_CACHE_KEY) .await? { - match serde_json::from_str::(&raw) { - Ok(envelope) + match parse_admin_external_models_cache(&raw) { + Some(envelope) if envelope.schema_version == ADMIN_EXTERNAL_MODELS_CACHE_VERSION && envelope.proxy_node_id == proxy_node_id => { @@ -479,25 +692,36 @@ pub(crate) async fn read_admin_external_models_cache( envelope.payload, ))); } - Ok(_) => {} - Err(err) => { - warn!(error = %err, "failed to parse cached external models payload"); - } + Some(_) => {} + None => warn!("failed to parse cached external models payload"), } } match fetch_admin_external_models_from_source(state, request_id, proxy_node_id.as_deref()).await { Ok(payload) => { - if let Err(err) = - store_admin_external_models_cache(state, proxy_node_id.as_deref(), &payload).await + if store_admin_external_models_cache(state, proxy_node_id.as_deref(), &payload) + .await + .is_err() { - warn!(error = ?err, "failed to store fetched external models cache"); + warn!("failed to store fetched external models cache"); } Ok(Some(payload)) } - Err(err) => { - warn!(error = ?err, "failed to fetch external models catalog"); + Err(error) => { + // Keep the client-facing response generic, but leave an actionable, + // low-cardinality diagnostic for operators. The underlying error + // is intentionally not logged here because a future transport + // implementation could include URL, proxy, or credential details. + let error_message = error.into_message(); + let transport_error_kind = + classify_admin_external_models_transport_error(&error_message); + warn!( + request_id = %request_id, + proxy_mode = if proxy_node_id.is_some() { "configured_node" } else { "direct" }, + transport_error_kind, + "failed to fetch external models catalog" + ); Ok(None) } } @@ -518,7 +742,10 @@ mod tests { use super::{ admin_external_models_source_url, classify_admin_external_models_transport_error, normalize_admin_external_models_payload, normalize_admin_external_models_proxy_node_id, - read_admin_external_models_cache, set_admin_external_models_source_url_for_tests, + parse_admin_external_models_cache, parse_admin_external_models_source_url, + read_admin_external_models_cache, resolve_admin_external_models_source, + set_admin_external_models_source_url_for_tests, + validate_admin_external_models_source_addresses, ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES, }; use crate::handlers::admin::request::AdminAppState; use crate::tests::{start_server, AppState}; @@ -564,12 +791,22 @@ mod tests { fn classifies_external_models_transport_errors_without_exposing_details() { for (message, expected) in [ ("request timeout after 300000ms", "timeout"), + ( + "external models source DNS resolution failed", + "dns_resolution", + ), + ( + "external models source resolves to a private or reserved address", + "ssrf_blocked", + ), + ("external models source URL is invalid", "invalid_url"), ("hub relay request failed", "relay"), ("invalid proxy configuration", "proxy_config"), ("upstream response body exceeds limit", "response_too_large"), ("upstream response is not valid JSON", "invalid_json"), ("failed to decode content-encoding gzip", "response_decode"), ("tcp connect error", "connect"), + ("external models source returned HTTP 503", "upstream_http"), ("invalid upstream header value", "request_build"), ("opaque execution failure", "unknown_transport"), ] { @@ -590,6 +827,102 @@ mod tests { ); } + #[test] + fn production_external_models_source_requires_safe_https_url_shape() { + assert!( + parse_admin_external_models_source_url("https://models.dev/api.json", false).is_ok() + ); + + for source_url in [ + "http://models.dev/api.json", + "file:///etc/passwd", + "https://user:secret@models.dev/api.json", + "https://models.dev/api.json?next=http://169.254.169.254", + "https://models.dev/api.json#fragment", + ] { + assert!( + parse_admin_external_models_source_url(source_url, false).is_err(), + "source URL should be rejected: {source_url}" + ); + } + } + + #[test] + fn external_models_cache_parser_rejects_oversized_runtime_values() { + let oversized = "x".repeat(ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES + 1); + assert!(parse_admin_external_models_cache(&oversized).is_none()); + + let valid = r#"{"schema_version":2,"proxy_node_id":null,"payload":{}}"#; + assert!(parse_admin_external_models_cache(valid).is_some()); + } + + #[tokio::test] + async fn production_external_models_source_rejects_private_ip_literals() { + for source_url in [ + "https://127.0.0.1/api.json", + "https://169.254.169.254/latest/meta-data", + "https://[::1]/api.json", + ] { + assert!( + resolve_admin_external_models_source(source_url, false) + .await + .is_err(), + "private source URL should be rejected: {source_url}" + ); + } + } + + #[test] + fn official_external_models_catalog_allows_benchmarking_fake_ip_addresses() { + let url = url::Url::parse("https://models.dev/api.json") + .expect("official catalog URL should parse"); + + for address in ["198.18.75.234:443", "198.19.255.254:443"] { + let address = address + .parse() + .expect("fake IP socket address should parse"); + assert!( + validate_admin_external_models_source_addresses(&url, &[address], false).is_ok(), + "official catalog should accept local proxy fake IP {address}" + ); + } + } + + #[test] + fn custom_external_models_sources_reject_benchmarking_fake_ip_addresses() { + let fake_ip = "198.18.75.234:443" + .parse() + .expect("fake IP socket address should parse"); + + for source_url in [ + "https://catalog.example/api.json", + "https://models.dev/other.json", + "https://models.dev:444/api.json", + ] { + let url = url::Url::parse(source_url).expect("custom catalog URL should parse"); + assert!( + validate_admin_external_models_source_addresses(&url, &[fake_ip], false).is_err(), + "custom source must reject the benchmarking fake IP: {source_url}" + ); + } + } + + #[test] + fn official_external_models_catalog_still_rejects_private_addresses() { + let url = url::Url::parse("https://models.dev/api.json") + .expect("official catalog URL should parse"); + + for address in ["127.0.0.1:443", "10.0.0.1:443", "169.254.169.254:443"] { + let address = address + .parse() + .expect("private socket address should parse"); + assert!( + validate_admin_external_models_source_addresses(&url, &[address], false).is_err(), + "official catalog must reject private address {address}" + ); + } + } + #[tokio::test] async fn read_external_models_fetches_remote_payload_when_cache_missing() { let upstream = Router::new().route( @@ -623,4 +956,38 @@ mod tests { upstream_handle.abort(); } + + #[tokio::test] + async fn direct_external_models_fetch_does_not_follow_redirects() { + let upstream = Router::new() + .route( + "/redirect", + get(|| async { axum::response::Redirect::temporary("/api.json") }), + ) + .route( + "/api.json", + get(|| async { + Json(json!({ + "openai": { + "name": "redirected payload", + "models": {} + } + })) + }), + ); + let (upstream_url, upstream_handle) = start_server(upstream).await; + let _guard = + set_admin_external_models_source_url_for_tests(&format!("{upstream_url}/redirect")); + + let state = AppState::new().expect("gateway should build"); + let payload = read_admin_external_models_cache( + &AdminAppState::new(&state), + "external-models-redirect", + ) + .await + .expect("external models read should not fail"); + + assert!(payload.is_none(), "redirected payload must not be accepted"); + upstream_handle.abort(); + } } diff --git a/apps/aether-gateway/src/handlers/admin/model/global_models/routes/core/writes.rs b/apps/aether-gateway/src/handlers/admin/model/global_models/routes/core/writes.rs index 0aaf6d9ce..b1cf71749 100644 --- a/apps/aether-gateway/src/handlers/admin/model/global_models/routes/core/writes.rs +++ b/apps/aether-gateway/src/handlers/admin/model/global_models/routes/core/writes.rs @@ -27,6 +27,36 @@ use axum::{ Json, }; use serde_json::json; +use std::collections::HashSet; + +const MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS: usize = 100; + +fn normalize_admin_global_model_batch_ids( + ids: Vec, + field_name: &str, +) -> Result, String> { + if ids.len() > MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS { + return Err(format!( + "{field_name} 最多 {MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS} 个" + )); + } + + let mut seen = HashSet::with_capacity(ids.len()); + let mut normalized = Vec::with_capacity(ids.len()); + for id in ids { + let trimmed = id.trim(); + if trimmed.is_empty() { + // Keep the original value so batch-delete retains its existing per-item failure. + normalized.push(id); + continue; + } + let trimmed = trimmed.to_string(); + if seen.insert(trimmed.clone()) { + normalized.push(trimmed); + } + } + Ok(normalized) +} pub(super) async fn maybe_build_local_admin_global_models_write_response( state: &AdminAppState<'_>, @@ -212,17 +242,20 @@ async fn build_batch_delete_global_models_response( Ok(payload) => payload, Err(response) => return Ok(response), }; + let ids = match normalize_admin_global_model_batch_ids(payload.ids, "ids") { + Ok(ids) => ids, + Err(detail) => return Ok(bad_request_response(detail)), + }; let mut success_count = 0usize; let mut failed = Vec::new(); - for id in payload.ids { - let trimmed = id.trim(); - if trimmed.is_empty() { + for id in ids { + if id.trim().is_empty() { failed.push(json!({"id": id, "error": "not found"})); continue; } - let Some(existing) = state.get_admin_global_model_by_id(trimmed).await? else { - failed.push(json!({"id": trimmed, "error": "not found"})); + let Some(existing) = state.get_admin_global_model_by_id(&id).await? else { + failed.push(json!({"id": id, "error": "not found"})); continue; }; if state.delete_admin_global_model(&existing.id).await? { @@ -245,6 +278,46 @@ async fn build_batch_delete_global_models_response( )) } +#[cfg(test)] +mod batch_boundary_tests { + use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS}; + + #[test] + fn global_model_batch_ids_are_bounded_and_deduplicated() { + assert_eq!( + normalize_admin_global_model_batch_ids( + vec![ + "model-2".to_string(), + "model-1".to_string(), + " model-2 ".to_string(), + " ".to_string(), + ], + "ids", + ) + .expect("valid ids"), + vec![ + "model-2".to_string(), + "model-1".to_string(), + " ".to_string(), + ] + ); + assert!(normalize_admin_global_model_batch_ids( + (0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS) + .map(|index| format!("model-{index}")) + .collect(), + "ids", + ) + .is_err()); + assert!(normalize_admin_global_model_batch_ids( + (0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS) + .map(|index| format!("provider-{index}")) + .collect(), + "provider_ids", + ) + .is_ok()); + } +} + async fn build_assign_to_providers_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -259,10 +332,15 @@ async fn build_assign_to_providers_response( Ok(payload) => payload, Err(response) => return Ok(response), }; + let provider_ids = + match normalize_admin_global_model_batch_ids(payload.provider_ids, "provider_ids") { + Ok(provider_ids) => provider_ids, + Err(detail) => return Ok(bad_request_response(detail)), + }; let payload: serde_json::Value = match build_admin_assign_global_model_to_providers_payload( state, &global_model_id, - payload.provider_ids, + provider_ids, payload.create_models.unwrap_or(false), ) .await diff --git a/apps/aether-gateway/src/handlers/admin/model/routing.rs b/apps/aether-gateway/src/handlers/admin/model/routing.rs index 8d602b187..10863f1fd 100644 --- a/apps/aether-gateway/src/handlers/admin/model/routing.rs +++ b/apps/aether-gateway/src/handlers/admin/model/routing.rs @@ -67,27 +67,17 @@ pub(crate) async fn build_admin_global_model_routing_payload( .push(key); } - let scheduling_mode = state - .read_system_config_json_value("scheduling_mode") - .await - .ok() - .flatten() - .and_then(|value| value.as_str().map(ToOwned::to_owned)) - .unwrap_or_else(|| "cache_affinity".to_string()); - let priority_mode = state - .read_system_config_json_value("provider_priority_mode") - .await - .ok() - .flatten() - .and_then(|value| value.as_str().map(ToOwned::to_owned)) - .unwrap_or_else(|| "provider".to_string()); - let keep_priority_on_conversion = state - .read_system_config_json_value("keep_priority_on_conversion") - .await - .ok() - .flatten() - .and_then(|value| value.as_bool()) - .unwrap_or(false); + // The admin view reports the system-default routing strategy. + let ordering_config = + match crate::scheduler::config::read_system_default_routing_ordering_config(state.app()) + .await + { + Ok(Some(config)) => config, + Ok(None) | Err(_) => crate::scheduler::config::SchedulerOrderingConfig::default(), + }; + let scheduling_mode = ordering_config.scheduling_mode_str().to_string(); + let priority_mode = ordering_config.priority_mode_str().to_string(); + let keep_priority_on_conversion = ordering_config.keep_priority_on_conversion; let now_unix_secs = SystemTime::now() .duration_since(UNIX_EPOCH) .map(|duration| duration.as_secs()) diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_affinity_reads.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_affinity_reads.rs index db5a9bb89..bfa72dce4 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_affinity_reads.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_affinity_reads.rs @@ -141,9 +141,8 @@ pub(super) async fn build_admin_monitoring_cache_affinities_response( let key = affinity.key_id.as_ref().and_then(|id| key_by_id.get(id)); let user_api_key_name = user_api_key.and_then(|item| item.name.clone()); - let user_api_key_prefix = user_api_key.and_then(|item| { - admin_monitoring_masked_user_api_key_prefix(state, item.key_encrypted.as_deref()) - }); + let user_api_key_prefix = + user_api_key.and_then(|item| admin_monitoring_masked_user_api_key_prefix(state, item)); let provider_name = provider.map(|item| item.name.clone()); let endpoint_url = endpoint .map(|item| item.base_url.clone()) diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_payloads.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_payloads.rs index 89dc478c0..165eff953 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_payloads.rs @@ -1,28 +1,16 @@ use crate::handlers::admin::request::AdminAppState; +use crate::handlers::shared::{masked_secret_display, open_auth_api_key_secret}; use crate::provider_key_auth::{ provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization, }; -use aether_crypto::decrypt_python_fernet_ciphertext; -#[cfg(test)] -use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; pub(super) fn admin_monitoring_masked_user_api_key_prefix( state: &AdminAppState<'_>, - ciphertext: Option<&str>, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, ) -> Option { - let Some(ciphertext) = ciphertext.map(str::trim).filter(|value| !value.is_empty()) else { - return None; - }; - let full_key = admin_monitoring_try_decrypt_secret(state, ciphertext)?; - let prefix_len = full_key.len().min(10); - let prefix = &full_key[..prefix_len]; - let suffix = if full_key.len() >= 4 { - &full_key[full_key.len().saturating_sub(4)..] - } else { - "" - }; - Some(format!("{prefix}...{suffix}")) + let projection = open_auth_api_key_secret(state.app(), record).ok()?; + Some(masked_secret_display(&projection.plaintext, 10, 4, "...")) } pub(super) fn admin_monitoring_masked_provider_key_prefix( @@ -43,59 +31,16 @@ pub(super) fn admin_monitoring_masked_provider_key_prefix( } } _ => { - let full_key = key - .encrypted_api_key - .as_deref() - .and_then(|ciphertext| admin_monitoring_try_decrypt_secret(state, ciphertext))?; - if full_key.len() <= 12 { - Some(format!("{full_key}***")) - } else { - Some(format!( - "{}***{}", - &full_key[..8], - &full_key[full_key.len().saturating_sub(4)..] - )) - } + let full_key = state + .app() + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten()?; + Some(masked_secret_display(&full_key, 8, 4, "***")) } } } -fn admin_monitoring_try_decrypt_secret( - state: &AdminAppState<'_>, - ciphertext: &str, -) -> Option { - let ciphertext = ciphertext.trim(); - if ciphertext.is_empty() { - return None; - } - let encryption_key = state.encryption_key().map(str::trim).unwrap_or(""); - if !encryption_key.is_empty() { - if let Ok(value) = decrypt_python_fernet_ciphertext(encryption_key, ciphertext) { - return Some(value); - } - } - for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] { - let Ok(candidate) = std::env::var(env_key) else { - continue; - }; - let candidate = candidate.trim(); - if candidate.is_empty() || candidate == encryption_key { - continue; - } - if let Ok(value) = decrypt_python_fernet_ciphertext(candidate, ciphertext) { - return Some(value); - } - } - #[cfg(test)] - if encryption_key != DEVELOPMENT_ENCRYPTION_KEY { - if let Ok(value) = decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, ciphertext) - { - return Some(value); - } - } - None -} - pub(super) fn admin_monitoring_cache_affinity_sort_value(value: Option<&serde_json::Value>) -> f64 { let Some(value) = value else { return 0.0; @@ -122,11 +67,14 @@ pub(super) fn admin_monitoring_cache_affinity_sort_value(value: Option<&serde_js #[cfg(test)] mod tests { - use super::admin_monitoring_masked_provider_key_prefix; + use super::{ + admin_monitoring_masked_provider_key_prefix, admin_monitoring_masked_user_api_key_prefix, + }; use crate::handlers::admin::request::AdminAppState; use crate::AppState; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; + use sha2::{Digest, Sha256}; #[test] fn monitoring_labels_agent_identity_instead_of_oauth_token() { @@ -167,4 +115,41 @@ mod tests { Some("[Agent Identity]") ); } + + #[test] + fn monitoring_never_exposes_complete_short_credentials() { + let app = AppState::new().expect("gateway should build"); + let state = AdminAppState::new(&app); + let plaintext = "short-key"; + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, plaintext) + .expect("secret should encrypt"); + let mut hasher = Sha256::new(); + hasher.update(plaintext.as_bytes()); + let record = aether_data::repository::auth::StoredAuthApiKeyExportRecord::new( + "owner-1".to_string(), + "key-1".to_string(), + format!("{:x}", hasher.finalize()), + Some(ciphertext), + None, + None, + None, + None, + None, + None, + None, + true, + None, + false, + 0, + 0, + 0.0, + false, + ) + .expect("API-key record should build"); + + let masked = admin_monitoring_masked_user_api_key_prefix(&state, &record) + .expect("secret should decrypt"); + assert_ne!(masked, plaintext); + assert!(!masked.contains(plaintext)); + } } diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_store.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_store.rs index d5a852cf5..37c66fce3 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_store.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/cache_store.rs @@ -264,16 +264,12 @@ async fn list_admin_monitoring_cache_affinity_records_matching( pub(super) async fn build_admin_monitoring_cache_snapshot( state: &AdminAppState<'_>, ) -> Result { - let scheduling_mode = state - .read_system_config_json_value("scheduling_mode") - .await? - .and_then(|value| value.as_str().map(ToOwned::to_owned)) - .unwrap_or_else(|| "cache_affinity".to_string()); - let provider_priority_mode = state - .read_system_config_json_value("provider_priority_mode") - .await? - .and_then(|value| value.as_str().map(ToOwned::to_owned)) - .unwrap_or_else(|| "provider".to_string()); + let ordering_config = + crate::scheduler::config::read_system_default_routing_ordering_config(state.app()) + .await? + .unwrap_or_default(); + let scheduling_mode = ordering_config.scheduling_mode_str().to_string(); + let provider_priority_mode = ordering_config.priority_mode_str().to_string(); let now = chrono::Utc::now(); let usage_summary = if state.has_usage_data_reader() { diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/resilience/snapshot.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/resilience/snapshot.rs index 2e8129c08..88c73a83f 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/resilience/snapshot.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/resilience/snapshot.rs @@ -218,7 +218,12 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot( "model": item.model, "api_format": item.api_format, "status_code": item.status_code, - "error_message": item.error_message, + "error_message": item + .error_category + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("request_failed"), } }) }) diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/mod.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/mod.rs index 0ae03812b..03f27de15 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/mod.rs @@ -3,6 +3,7 @@ use super::test_support::*; use crate::control::GatewayPublicRequestContext; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::AppState; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data_contracts::repository::{ candidates::{RequestCandidateStatus, StoredRequestCandidate}, provider_catalog::{ @@ -119,9 +120,12 @@ async fn admin_monitoring_cache_affinities_and_affinity_return_local_payload_fro let state = AppState::new() .expect("state should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog) - .with_user_reader(user_repository) - .with_auth_api_key_reader(auth_repository), + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog, + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_user_reader(user_repository) + .with_auth_api_key_reader(auth_repository), ) .with_admin_monitoring_cache_affinity_entry_for_tests( "cache_affinity:user-key-1:openai:model-alpha", @@ -228,9 +232,12 @@ async fn admin_monitoring_cache_affinities_and_delete_use_runtime_scheduler_affi let state = AppState::new() .expect("state should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog) - .with_user_reader(user_repository) - .with_auth_api_key_reader(auth_repository), + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog, + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_user_reader(user_repository) + .with_auth_api_key_reader(auth_repository), ); let affinity_cache_key = aether_scheduler_core::build_scheduler_affinity_cache_key_for_api_key_id( @@ -358,9 +365,12 @@ async fn admin_monitoring_cache_affinities_parse_session_scoped_scheduler_affini let state = AppState::new() .expect("state should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog) - .with_user_reader(user_repository) - .with_auth_api_key_reader(auth_repository), + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog, + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_user_reader(user_repository) + .with_auth_api_key_reader(auth_repository), ); let client_session = aether_scheduler_core::ClientSessionAffinity::new( Some("Codex".to_string()), diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/trace.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/trace.rs index f5f183391..5c83c25eb 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/trace.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/tests/trace.rs @@ -69,7 +69,7 @@ async fn admin_monitoring_trace_request_returns_local_payload() { assert_eq!(payload["candidates"][0]["provider_name"], json!("OpenAI")); assert_eq!( payload["candidates"][0]["provider_website"], - json!("https://openai.com") + json!("https://openai.com/") ); assert_eq!( payload["candidates"][0]["endpoint_name"], @@ -292,10 +292,10 @@ async fn admin_monitoring_trace_request_falls_back_to_usage_routing_snapshot() { payload["candidates"][0]["extra_data"]["execution_path"], json!("local_execution_runtime_miss") ); - assert_eq!( - payload["candidates"][0]["extra_data"]["failure_diagnostic"]["path"], - json!("$.reasoning.summary") - ); + assert!(payload["candidates"][0]["error_message"].is_null()); + assert!(payload["candidates"][0]["extra_data"] + .get("failure_diagnostic") + .is_none()); } #[tokio::test] @@ -340,11 +340,9 @@ async fn admin_monitoring_trace_request_returns_oauth_account_label_from_auth_co vec![sample_endpoint()], vec![oauth_key], )); - let data_state = GatewayDataState::with_decision_trace_readers_for_tests( - request_candidates, - provider_catalog, - ) - .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let data_state = GatewayDataState::with_request_candidate_reader_for_tests(request_candidates) + .attach_provider_catalog_repository_for_tests(provider_catalog) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); let state = AppState::new() .expect("state should build") .with_data_state_for_tests(data_state); @@ -615,7 +613,7 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit() usage.candidate_id = Some("cand-used".to_string()); usage.request_metadata = Some(json!({ "request_path": "/v1beta/models/gemini-2.5-pro:generateContent", - "request_query_string": "alt=sse" + "request_query_string": "alt=sse&key=gemini-secret&access_token=oauth-secret" })); let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage])); let data_state = @@ -649,14 +647,19 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit() payload["request_path_and_query"], json!("/v1beta/models/gemini-2.5-pro:generateContent?alt=sse") ); - assert_eq!( - payload["candidates"][0]["extra_data"]["request_path_and_query"], - json!("/v1beta/models/gemini-2.5-pro:generateContent?alt=sse") - ); + assert!(payload["candidates"][0]["extra_data"] + .get("request_path") + .is_none()); + assert!(payload["candidates"][0]["extra_data"] + .get("request_query_string") + .is_none()); + assert!(payload["candidates"][0]["extra_data"] + .get("request_path_and_query") + .is_none()); } #[tokio::test] -async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_response_boundary() { +async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloads() { let mut candidate = sample_candidate( "cand-used", "request-1", @@ -739,19 +742,18 @@ async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_respon let extra = &payload["candidates"][0]["extra_data"]; assert_eq!(extra["upstream_response"]["status_code"], json!(302)); assert_eq!( - extra["upstream_response"]["headers"]["location"], - json!("/") - ); - assert_eq!( - extra["upstream_response"]["body"]["error"]["message"], - json!("redirect blocked") + extra["upstream_response"]["source"], + json!("upstream_response") ); + assert!(extra["upstream_response"].get("headers").is_none()); + assert!(extra["upstream_response"].get("body").is_none()); + assert!(extra["upstream_response"].get("body_ref").is_none()); assert!(extra.get("client_response").is_none()); assert!(extra.get("provider_response").is_none()); } #[tokio::test] -async fn admin_monitoring_trace_request_prefers_ref_backed_usage_response_body() { +async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_response_body() { let mut candidate = sample_candidate( "cand-used", "request-ref-body", @@ -832,30 +834,16 @@ async fn admin_monitoring_trace_request_prefers_ref_backed_usage_response_body() .expect("body should read"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"]; - assert_eq!( - upstream_response["headers"], - json!({ - "content-type": "application/json", - "x-request-id": "req_usage-cyber-risk-demo" - }) - ); - assert_eq!( - upstream_response["body"]["error"], - json!({ - "type": "invalid_request", - "message": "This content was flagged for possible cybersecurity risk.", - "code": 400 - }) - ); - assert!(upstream_response["body"].get("input").is_none()); - assert_eq!( - upstream_response["body_ref"], - json!("usage://request/request-ref-body/response_body") - ); + assert_eq!(upstream_response["status_code"], json!(400)); + assert_eq!(upstream_response["source"], json!("upstream_response")); + assert_eq!(upstream_response["body_state"], json!("reference")); + assert!(upstream_response.get("headers").is_none()); + assert!(upstream_response.get("body").is_none()); + assert!(upstream_response.get("body_ref").is_none()); } #[tokio::test] -async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs() { +async fn admin_monitoring_trace_request_does_not_expose_inline_connect_json_response_body() { let mut candidate = sample_candidate( "cand-used", "request-connect", @@ -922,19 +910,11 @@ async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs( let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"]; assert_eq!(upstream_response["status_code"], json!(429)); - assert_eq!( - upstream_response["body"]["error"]["code"], - json!("resource_exhausted") - ); - assert_eq!( - upstream_response["body"]["error"]["message"], - json!("quota exhausted") - ); - assert_eq!( - upstream_response["body_ref"], - json!("usage://request/request-connect/response_body") - ); + assert_eq!(upstream_response["source"], json!("upstream_response")); assert_eq!(upstream_response["body_state"], json!("inline")); + assert!(upstream_response.get("headers").is_none()); + assert!(upstream_response.get("body").is_none()); + assert!(upstream_response.get("body_ref").is_none()); } #[tokio::test] diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/trace.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/trace.rs index c3e5fd80c..28abec2c0 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/trace.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/trace.rs @@ -11,8 +11,9 @@ use aether_admin::observability::monitoring::{ }; use aether_data_contracts::repository::{ candidates::{ - DecisionTrace, DecisionTraceCandidate, RequestCandidateFinalStatus, RequestCandidateStatus, - StoredRequestCandidate, + sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data, + sanitize_request_candidate_skip_reason, DecisionTrace, DecisionTraceCandidate, + RequestCandidateFinalStatus, RequestCandidateStatus, StoredRequestCandidate, }, provider_catalog::StoredProviderCatalogKey, usage::StoredRequestUsageAudit, @@ -30,25 +31,6 @@ struct ResolvedAdminMonitoringTrace { usage: Option, } -async fn hydrate_admin_monitoring_trace_response_body( - state: &AdminAppState<'_>, - mut usage: StoredRequestUsageAudit, -) -> Result { - let is_error_node = !usage.status.eq_ignore_ascii_case("completed") - || usage - .status_code - .is_some_and(|status| !(200..300).contains(&status)); - let response_body_ref = if is_error_node && usage.response_body.is_none() { - usage.response_body_ref.clone() - } else { - None - }; - if let Some(body_ref) = response_body_ref.as_deref() { - usage.response_body = state.resolve_request_usage_body_ref(body_ref).await?; - } - Ok(usage) -} - pub(super) async fn build_admin_monitoring_trace_request_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -111,10 +93,6 @@ async fn resolve_admin_monitoring_trace( .read_request_usage_audit_shallow(request_id) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; - let usage = match usage { - Some(usage) => Some(hydrate_admin_monitoring_trace_response_body(state, usage).await?), - None => None, - }; return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage })); } @@ -125,10 +103,9 @@ async fn resolve_admin_monitoring_trace( .await .map_err(|err| GatewayError::Internal(err.to_string()))? { - usage_candidates.push(hydrate_admin_monitoring_trace_response_body(state, usage).await?); + usage_candidates.push(usage); } if let Some(usage) = state.find_request_usage_by_id(request_id).await? { - let usage = hydrate_admin_monitoring_trace_response_body(state, usage).await?; if !usage_candidates.iter().any(|item| item.id == usage.id) { usage_candidates.push(usage); } @@ -203,17 +180,23 @@ fn build_admin_monitoring_usage_routing_snapshot_trace( endpoint_id: usage.provider_endpoint_id.clone(), key_id: usage.provider_api_key_id.clone(), status, - skip_reason: usage.routing_candidate_skip_reason().map(ToOwned::to_owned), + skip_reason: sanitize_request_candidate_skip_reason( + usage.routing_candidate_skip_reason().map(ToOwned::to_owned), + ), is_cached: false, status_code: usage.status_code, - error_type: usage - .routing_local_execution_runtime_miss_reason() - .or(usage.error_category.as_deref()) - .map(ToOwned::to_owned), - error_message: usage.error_message.clone(), + error_type: sanitize_request_candidate_error_type( + usage + .routing_local_execution_runtime_miss_reason() + .or(usage.error_category.as_deref()) + .map(ToOwned::to_owned), + ), + error_message: None, latency_ms: usage.response_time_ms, concurrent_requests: None, - extra_data: build_admin_monitoring_usage_routing_snapshot_extra_data(usage), + extra_data: sanitize_request_candidate_extra_data( + build_admin_monitoring_usage_routing_snapshot_extra_data(usage), + ), required_capabilities: None, created_at_unix_ms: usage.created_at_unix_ms, started_at_unix_ms: Some(usage.created_at_unix_ms), @@ -485,8 +468,12 @@ fn parse_admin_monitoring_key_auth_config( state: &AdminAppState<'_>, key: &StoredProviderCatalogKey, ) -> Option> { - let ciphertext = key.encrypted_auth_config.as_deref()?; - let plaintext = state.decrypt_catalog_secret_with_fallbacks(ciphertext)?; + let _ciphertext = key.encrypted_auth_config.as_deref()?; + let plaintext = state + .app() + .decrypt_provider_catalog_key_auth_config(key) + .ok() + .flatten()?; serde_json::from_str::(&plaintext) .ok()? .as_object() diff --git a/apps/aether-gateway/src/handlers/admin/observability/usage/replay.rs b/apps/aether-gateway/src/handlers/admin/observability/usage/replay.rs index 3b6733b21..93446b47e 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/usage/replay.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/usage/replay.rs @@ -6,7 +6,10 @@ use aether_admin::observability::usage::{ }; use aether_data_contracts::repository::{ provider_catalog::StoredProviderCatalogEndpoint, - usage::{StoredRequestUsageAudit, UsageBodyCaptureState, UsageBodyField}, + usage::{ + canonical_usage_body_ref_for, StoredRequestUsageAudit, UsageBodyCaptureState, + UsageBodyField, + }, }; use axum::{ body::Body, @@ -64,7 +67,10 @@ pub(super) async fn admin_usage_resolve_body_value( } Some(UsageBodyCaptureState::Reference) | None => {} } - let resolved_ref_body = match item.body_ref(field) { + let body_ref = item + .body_ref(field) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, &item.request_id, field)); + let resolved_ref_body = match body_ref.as_deref() { Some(body_ref) => state.resolve_request_usage_body_ref(body_ref).await?, None => None, }; 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 be83243e4..14a3f0d95 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 @@ -14,7 +14,9 @@ use aether_admin::observability::usage::{ }; use aether_data::repository::users::StoredUserSummary; use aether_data_contracts::repository::{ - candidates::{RequestCandidateStatus, StoredRequestCandidate}, + candidates::{ + sanitize_request_candidate_extra_data, RequestCandidateStatus, StoredRequestCandidate, + }, usage::{ StoredRequestUsageAudit, UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery, @@ -255,10 +257,8 @@ fn latest_admin_usage_image_progress( candidates .iter() .filter_map(|candidate| { - let progress = candidate - .extra_data - .as_ref() - .and_then(|value| value.get("image_progress"))? + let progress = sanitize_request_candidate_extra_data(candidate.extra_data.clone())? + .get("image_progress")? .clone(); Some(( candidate @@ -348,9 +348,6 @@ pub(super) fn admin_usage_terminal_candidate_state_override( if let Some(status_code) = candidate.status_code { payload["status_code"] = json!(status_code); } - if let Some(error_message) = candidate.error_message.as_ref() { - payload["error_message"] = json!(error_message); - } Some(payload) } @@ -1038,7 +1035,8 @@ mod tests { use super::{ admin_usage_terminal_candidate_state_override, build_admin_usage_keyword_search_query, - build_admin_usage_records_query, AdminUsageSearchContext, + build_admin_usage_records_query, latest_admin_usage_image_progress, + AdminUsageSearchContext, }; fn sample_candidate( @@ -1140,6 +1138,32 @@ mod tests { assert!(payload.is_none()); } + #[test] + fn admin_usage_image_progress_sanitizes_untrusted_candidate_data() { + let mut candidate = + sample_candidate(0, RequestCandidateStatus::Streaming, None, None, None); + candidate.extra_data = Some(json!({ + "image_progress": { + "phase": "upstream_streaming", + "upstream_sse_frame_count": 3, + "message": "Bearer candidate-secret", + "request_body": {"token": "candidate-secret"} + } + })); + + let progress = latest_admin_usage_image_progress(&[candidate]) + .expect("safe progress summary should remain"); + + assert_eq!( + progress, + json!({ + "phase": "upstream_streaming", + "upstream_sse_frame_count": 3 + }) + ); + assert!(!progress.to_string().contains("candidate-secret")); + } + #[test] fn admin_usage_transport_statuses_are_disjoint_in_list_and_keyword_queries() { for status in ["websocket", "ws", "WS"] { diff --git a/apps/aether-gateway/src/handlers/admin/provider/delete_task.rs b/apps/aether-gateway/src/handlers/admin/provider/delete_task.rs index 4a4ae441b..b07cc3190 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/delete_task.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/delete_task.rs @@ -3,11 +3,9 @@ use crate::handlers::admin::provider::shared::support::{ ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS, ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS, }; use crate::handlers::admin::request::AdminAppState; -use crate::handlers::admin::shared::{ - decrypt_catalog_secret_with_fallbacks, json_string_list, parse_catalog_auth_config_json, - take_secret_prefix, take_secret_suffix, -}; +use crate::handlers::admin::shared::{json_string_list, parse_catalog_auth_config_json}; use crate::handlers::public::matches_model_mapping_for_models; +use crate::handlers::shared::masked_secret_display; use crate::provider_key_auth::provider_key_auth_config_is_agent_identity; use crate::{GatewayError, LocalProviderDeleteTaskState}; use aether_data_contracts::repository::global_models::{ @@ -182,21 +180,12 @@ pub(crate) fn mapping_preview_masked_catalog_api_key( return "***".to_string(); } - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - .map(|value| { - let char_count = value.chars().count(); - if char_count > 8 { - format!( - "{}***{}", - take_secret_prefix(&value, 4), - take_secret_suffix(&value, 4) - ) - } else if char_count >= 2 { - format!("{}***", take_secret_prefix(&value, 2)) - } else { - "***".to_string() - } - }) + state + .as_ref() + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten() + .map(|value| masked_secret_display(&value, 4, 4, "***")) .unwrap_or_else(|| "***".to_string()) } diff --git a/apps/aether-gateway/src/handlers/admin/provider/endpoint_keys/reads.rs b/apps/aether-gateway/src/handlers/admin/provider/endpoint_keys/reads.rs index f031d924e..f73bdb602 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/endpoint_keys/reads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/endpoint_keys/reads.rs @@ -2,7 +2,9 @@ use crate::handlers::admin::provider::shared::paths::{ admin_export_key_id, admin_provider_id_for_keys, admin_reveal_key_id, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; -use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value}; +use crate::handlers::admin::shared::{ + attach_admin_audit_response, mark_sensitive_admin_response_no_store, query_param_value, +}; use crate::GatewayError; use axum::{ body::{Body, Bytes}, @@ -94,13 +96,13 @@ pub(super) async fn maybe_handle( )); }; return Ok(Some(match state.build_admin_reveal_key_payload(&key) { - Ok(payload) => attach_admin_audit_response( + Ok(payload) => mark_sensitive_admin_response_no_store(attach_admin_audit_response( Json(payload).into_response(), "admin_provider_key_revealed", "reveal_provider_key", "provider_key", &key_id, - ), + )), Err(detail) => ( http::StatusCode::BAD_REQUEST, Json(json!({ "detail": detail })), @@ -141,13 +143,13 @@ pub(super) async fn maybe_handle( }; return Ok(Some( match state.build_admin_export_key_payload(&key).await { - Ok(payload) => attach_admin_audit_response( + Ok(payload) => mark_sensitive_admin_response_no_store(attach_admin_audit_response( Json(payload).into_response(), "admin_provider_key_exported", "export_provider_key", "provider_key_export", &key_id, - ), + )), Err(detail) => ( http::StatusCode::BAD_REQUEST, Json(json!({ "detail": detail })), diff --git a/apps/aether-gateway/src/handlers/admin/provider/models/batch.rs b/apps/aether-gateway/src/handlers/admin/provider/models/batch.rs index 620a84b56..292223d60 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/models/batch.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/models/batch.rs @@ -13,6 +13,33 @@ use serde_json::json; use std::collections::BTreeSet; use std::time::{SystemTime, UNIX_EPOCH}; +const MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS: usize = 100; + +fn validate_admin_provider_model_batch( + payloads: Vec, +) -> Result, String> { + if payloads.len() > MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS { + return Err(format!( + "批量创建模型最多支持 {MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS} 条" + )); + } + + let mut normalized = Vec::with_capacity(payloads.len()); + let mut seen = BTreeSet::new(); + for mut payload in payloads { + let normalized_name = payload.provider_model_name.trim().to_string(); + if normalized_name.is_empty() { + return Err("provider_model_name 不能为空".to_string()); + } + if !seen.insert(normalized_name.clone()) { + return Err(format!("批量请求中包含重复模型 {normalized_name}")); + } + payload.provider_model_name = normalized_name.clone(); + normalized.push((normalized_name, payload)); + } + Ok(normalized) +} + pub(super) async fn maybe_handle( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -68,28 +95,23 @@ pub(super) async fn maybe_handle( )); } }; - let mut created = Vec::new(); - let mut seen = BTreeSet::new(); - for payload in payloads { - let normalized_name = payload.provider_model_name.trim().to_string(); - if normalized_name.is_empty() { + let payloads = match validate_admin_provider_model_batch(payloads) { + Ok(payloads) => payloads, + Err(detail) => { return Ok(Some( ( http::StatusCode::BAD_REQUEST, - Json(json!({ "detail": "provider_model_name 不能为空" })), - ) - .into_response(), - )); - } - if !seen.insert(normalized_name.clone()) { - return Ok(Some( - ( - http::StatusCode::BAD_REQUEST, - Json(json!({ "detail": format!("批量请求中包含重复模型 {normalized_name}") })), + Json(json!({ "detail": detail })), ) .into_response(), )); } + }; + + // Complete request validation before the first write so a bad later item cannot leave + // an earlier subset committed while the endpoint returns a validation error. + let mut staged = Vec::new(); + for (normalized_name, payload) in payloads { if admin_provider_model_name_exists(state, &provider_id, &normalized_name, None).await? { continue; @@ -109,6 +131,11 @@ pub(super) async fn maybe_handle( )); } }; + staged.push(record); + } + + let mut created = Vec::with_capacity(staged.len()); + for record in staged { let Some(model) = state.create_admin_provider_model(&record).await? else { return Ok(Some( ( @@ -140,3 +167,36 @@ pub(super) async fn maybe_handle( Ok(None) } + +#[cfg(test)] +mod tests { + use super::{ + validate_admin_provider_model_batch, AdminProviderModelCreateRequest, + MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS, + }; + + fn payload(name: &str) -> AdminProviderModelCreateRequest { + serde_json::from_value(serde_json::json!({ + "provider_model_name": name, + "global_model_id": "global-1" + })) + .expect("payload should deserialize") + } + + #[test] + fn provider_model_batch_is_bounded_and_prevalidates_duplicates() { + let at_limit = (0..MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS) + .map(|index| payload(&format!("model-{index}"))) + .collect(); + assert!(validate_admin_provider_model_batch(at_limit).is_ok()); + + let oversized = (0..=MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS) + .map(|index| payload(&format!("model-{index}"))) + .collect(); + assert!(validate_admin_provider_model_batch(oversized).is_err()); + + let duplicate = vec![payload("model-1"), payload(" model-1 ")]; + assert!(validate_admin_provider_model_batch(duplicate).is_err()); + assert!(validate_admin_provider_model_batch(vec![payload(" ")]).is_err()); + } +} 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 9aefb9a54..42d9c0c15 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 @@ -394,7 +394,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens( }); } } - return Err(format!("Token 验证失败: {detail}")); + return Err("Token 验证失败".to_string()); } }; diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/kiro_import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/kiro_import.rs index 19b8de37f..ba1442f7d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/kiro_import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/kiro_import.rs @@ -138,12 +138,12 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import( .await { Ok(config) => config, - Err(err) => { + Err(_) => { failed += 1; results.push(json!({ "index": index, "status": "error", - "error": format!("Token 验证失败: {err}"), + "error": "Token 验证失败", "replaced": false, })); maybe_report_admin_provider_oauth_batch_import_progress( 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 a8cb9de1f..7d05e18f0 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 @@ -16,13 +16,13 @@ use serde_json::json; use std::collections::BTreeMap; use std::time::{SystemTime, UNIX_EPOCH}; -#[derive(Debug, Clone, Deserialize)] +#[derive(Clone, Deserialize)] pub(super) struct AdminProviderOAuthBatchImportRequest { pub credentials: String, pub proxy_node_id: Option, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub(super) struct AdminProviderOAuthBatchImportEntry { pub parse_error: Option, pub refresh_token: Option, @@ -52,7 +52,7 @@ pub(super) struct AdminProviderOAuthBatchImportEntry { pub rate_limit_tier: Option, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub(super) struct AdminProviderOAuthBatchImportOutcome { pub total: usize, pub success: usize, @@ -663,7 +663,7 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries( .collect(); } Ok(_) => {} - Err(error) => return vec![parse_error_entry(format!("JSON 数组解析失败: {error}"))], + Err(_) => return vec![parse_error_entry("JSON 数组解析失败".to_string())], } } @@ -697,8 +697,8 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries( "JSON 行必须是账号对象,不能作为 raw token 导入".to_string(), )); } - Err(error) => { - return Some(parse_error_entry(format!("JSON 行解析失败: {error}"))); + Err(_) => { + return Some(parse_error_entry("JSON 行解析失败".to_string())); } } } @@ -719,7 +719,7 @@ pub(super) fn parse_admin_provider_oauth_agent_identity_import_entries( return Err("Agent Identity 凭据不能为空".to_string()); } let value = serde_json::from_str::(raw) - .map_err(|error| format!("Agent Identity JSON 解析失败: {error}"))?; + .map_err(|_| "Agent Identity JSON 解析失败".to_string())?; let entries = match &value { serde_json::Value::Array(items) => items .iter() @@ -808,6 +808,14 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints( return; } if provider_type == "antigravity" { + // The Google refresh-token response does not include the account email. + // Preserve the identity supplied by the imported Antigravity credentials so + // account naming and duplicate detection can use it after token exchange. + if let Some(email) = entry.email.as_ref() { + auth_config + .entry("email".to_string()) + .or_insert_with(|| json!(email)); + } if let Some(project_id) = entry.project_id.as_ref() { auth_config .entry("project_id".to_string()) @@ -932,23 +940,7 @@ pub(super) async fn extract_admin_provider_oauth_batch_error_detail( response: Response, ) -> String { let status = response.status(); - let raw_body = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES) - .await - .ok(); - if let Some(raw_body) = raw_body { - if let Ok(value) = serde_json::from_slice::(&raw_body) { - if let Some(detail) = value.get("detail").and_then(serde_json::Value::as_str) { - let normalized = detail.trim(); - if !normalized.is_empty() { - return normalized.to_string(); - } - } - } - let normalized = String::from_utf8_lossy(&raw_body).trim().to_string(); - if !normalized.is_empty() { - return normalized; - } - } + let _ = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES).await; format!("HTTP {}", status.as_u16()) } @@ -1433,7 +1425,7 @@ mod tests { fn applies_antigravity_project_and_user_agent_hints_to_auth_config() { let entries = parse_admin_provider_oauth_batch_import_entries( "antigravity", - r#"{"refreshToken":"rt-1","cloudaicompanionProject":{"id":"project-antigravity-2"},"userAgent":"antigravity"}"#, + r#"{"refreshToken":"rt-1","email":"anti@example.com","cloudaicompanionProject":{"id":"project-antigravity-2"},"userAgent":"antigravity"}"#, ); let mut auth_config = serde_json::Map::new(); @@ -1444,6 +1436,35 @@ mod tests { Some(&json!("project-antigravity-2")) ); assert_eq!(auth_config.get("user_agent"), Some(&json!("antigravity"))); + assert_eq!(auth_config.get("email"), Some(&json!("anti@example.com"))); + } + + #[test] + fn antigravity_batch_import_keeps_json_email_for_key_naming() { + let entries = parse_admin_provider_oauth_batch_import_entries( + "antigravity", + r#"{"access_token":"at-1","refresh_token":"rt-1","email":"anti@example.com","project_id":"project-antigravity-3","type":"antigravity"}"#, + ); + // Simulate Google's refresh-token response, which carries no email. + let mut auth_config = json!({ + "provider_type": "antigravity", + "refresh_token": "rt-1", + }) + .as_object() + .cloned() + .expect("auth config should be an object"); + + apply_admin_provider_oauth_batch_import_hints("antigravity", &entries[0], &mut auth_config); + + assert_eq!(auth_config.get("email"), Some(&json!("anti@example.com"))); + assert_eq!( + super::super::super::helpers::admin_provider_oauth_key_name_from_auth_config( + "antigravity", + &auth_config, + Some(0), + ), + "antigravity_anti@example.com" + ); } #[test] diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/task.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/task.rs index d3559ee2d..11c20e2ea 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/task.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/task.rs @@ -65,13 +65,13 @@ fn codex_agent_identity_import_auth_configs( .iter() .enumerate() .map(|(index, entry)| { - if let Some(error) = entry.parse_error.as_deref() { - return Err(format!("第 {} 个条目无效: {error}", index + 1)); + if entry.parse_error.is_some() { + return Err(format!("第 {} 个条目无效", index + 1)); } match codex_agent_identity_auth_config_from_import(entry) { Ok(Some(auth_config)) => Ok(auth_config), Ok(None) => Err(format!("第 {} 个条目不是 Agent Identity", index + 1)), - Err(error) => Err(format!("第 {} 个条目无效: {error}", index + 1)), + Err(_) => Err(format!("第 {} 个条目无效", index + 1)), } }) .collect() @@ -135,11 +135,10 @@ async fn acquire_provider_agent_identity_import_locks( "其中一个 Agent Identity 正在导入或创建,请稍后重试", )); } - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, lock_key = %lock_key, - error = ?error, "gateway Agent Identity import lock unavailable" ); release_provider_agent_identity_import_locks(state, leases).await; @@ -164,9 +163,8 @@ async fn release_provider_agent_identity_import_locks( lock_key = %lease.key, "gateway Agent Identity import lock was not owned during release" ), - Err(error) => tracing::warn!( + Err(_) => tracing::warn!( lock_key = %lease.key, - error = ?error, "gateway Agent Identity import lock release failed" ), } @@ -329,10 +327,10 @@ async fn handle_admin_provider_oauth_start_import_task( let agent_identity_auth_configs = if agent_identity_only { match codex_agent_identity_import_auth_configs(&payload.credentials) { Ok(auth_configs) => Some(auth_configs), - Err(detail) => { + Err(_) => { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - format!("该接口仅接受有效的 Agent Identity JSON: {detail}"), + "该接口仅接受有效的 Agent Identity JSON", )); } } @@ -652,13 +650,12 @@ async fn handle_admin_provider_oauth_start_import_task( ) .await; } - Err(err) => { + Err(_) => { let finished_at = SystemTime::now() .duration_since(UNIX_EPOCH) .ok() .map(|duration| duration.as_secs()) .unwrap_or(started_at); - let error_message = format!("{err:?}"); let failed_state = build_admin_provider_oauth_batch_task_state( &task_id_for_worker, &provider_id_for_worker, @@ -672,7 +669,7 @@ async fn handle_admin_provider_oauth_start_import_task( 0, 0, Some("导入任务执行失败"), - Some(error_message.as_str()), + Some("provider_oauth_batch_import_failed"), Vec::new(), created_at, Some(started_at), @@ -688,7 +685,7 @@ async fn handle_admin_provider_oauth_start_import_task( Some(100), Some("provider oauth batch import failed".to_string()), None, - Some(error_message.clone()), + Some("provider_oauth_batch_import_failed".to_string()), None, Some(finished_at), ) @@ -698,13 +695,15 @@ async fn handle_admin_provider_oauth_start_import_task( &task_id_for_worker, "failed", "provider oauth batch import failed", - Some(json!({ "error": error_message.clone() })), + Some(json!({ + "error_code": "provider_oauth_batch_import_failed" + })), ) .await; tracing::warn!( task_id = %task_id_for_worker, provider_id = %provider_id_for_worker, - error = %error_message, + error_category = "provider_oauth_batch_import_failed", "provider oauth batch import task failed" ); } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/key.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/key.rs index dd3a73d56..dd251ea3b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/key.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/key.rs @@ -15,7 +15,8 @@ use super::super::super::state::{ is_fixed_provider_type_for_provider_oauth, json_non_empty_string, }; use super::shared::{ - parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body, + admin_provider_oauth_state_matches_principal, parse_admin_provider_oauth_complete_callback, + parse_admin_provider_oauth_complete_request_body, }; use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_key_id; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; @@ -23,7 +24,7 @@ use crate::handlers::shared::sync_provider_key_oauth_status_snapshot; use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::GatewayError; use aether_data_contracts::repository::provider_catalog::{ - ProviderCatalogKeyOAuthRuntimeStateCasUpdate, + ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogUpstreamMetadataNamespaceExpectation, }; use axum::{ @@ -96,10 +97,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( Err(response) => return Ok(response), }; - let state_data = match state - .consume_provider_oauth_state(&callback.state_nonce) - .await - { + let preview = match state.load_provider_oauth_state(&callback.state_nonce).await { Ok(Some(state_data)) => state_data, Ok(None) => { return Ok(build_internal_control_error_response( @@ -107,6 +105,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( "state 无效或已过期", )); } + Err(GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + .. + }) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } Err(_) => { return Ok(build_internal_control_error_response( http::StatusCode::SERVICE_UNAVAILABLE, @@ -114,12 +121,41 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( )); } }; - if state_data.key_id != key_id { + if preview.key_id != key_id + || !admin_provider_oauth_state_matches_principal(&preview, request_context) + { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, "state 无效或已过期", )); } + let state_data = match state + .consume_provider_oauth_state(&callback.state_nonce) + .await + { + Ok(Some(state_data)) if state_data == preview => state_data, + Ok(Some(_)) | Ok(None) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } + Err(GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + .. + }) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } + Err(_) => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth redis unavailable", + )); + } + }; let key = state .read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id)) @@ -248,7 +284,11 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( } enrich_admin_provider_oauth_auth_config(&provider_type, &mut auth_config, &token_payload); - let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(&access_token) else { + let Ok(encrypted_api_key) = + state + .app() + .seal_provider_catalog_key_api_key(&provider_id, &key_id, &access_token) + else { return Ok(build_internal_control_error_response( http::StatusCode::SERVICE_UNAVAILABLE, "provider oauth encryption unavailable", @@ -256,8 +296,10 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( }; let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone())) .map_err(|err| GatewayError::Internal(err.to_string()))?; - let Some(encrypted_auth_config) = - state.encrypt_catalog_secret_with_fallbacks(&auth_config_json) + let Ok(encrypted_auth_config) = + state + .app() + .seal_provider_catalog_key_auth_config(&provider_id, &key_id, &auth_config_json) else { return Ok(build_internal_control_error_response( http::StatusCode::SERVICE_UNAVAILABLE, @@ -340,6 +382,12 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( CODEX_CREDENTIAL_GENERATION_KEY: uuid::Uuid::now_v7().to_string() }); let expected_encrypted_auth_config = state_data.expected_encrypted_auth_config.clone(); + let expected_credential = ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: key.encrypted_api_key.clone(), + auth_type: key.auth_type.clone(), + provider_id: key.provider_id.clone(), + provider_type: provider.provider_type.clone(), + }; let updated_result: Result = async { let max_namespace_retries = if provider_type == "codex" { CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES @@ -353,7 +401,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key( &ProviderCatalogKeyOAuthRuntimeStateCasUpdate { key_id: key_id.clone(), expected_encrypted_auth_config: expected_encrypted_auth_config.clone(), - expected_credential: None, + expected_credential: Some(expected_credential.clone()), expected_upstream_metadata_namespace: (provider_type == "codex").then( || ProviderCatalogUpstreamMetadataNamespaceExpectation { namespace: "codex".to_string(), diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs index 7b7760e38..8217c81df 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs @@ -8,7 +8,8 @@ use super::super::super::state::{ is_fixed_provider_type_for_provider_oauth, }; use super::shared::{ - parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body, + admin_provider_oauth_state_matches_principal, parse_admin_provider_oauth_complete_callback, + parse_admin_provider_oauth_complete_request_body, }; use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_provider_id; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; @@ -43,10 +44,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider( Err(response) => return Ok(response), }; - let state_data = match state - .consume_provider_oauth_state(&callback.state_nonce) - .await - { + let preview = match state.load_provider_oauth_state(&callback.state_nonce).await { Ok(Some(state_data)) => state_data, Ok(None) => { return Ok(build_internal_control_error_response( @@ -54,6 +52,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider( "state 无效或已过期", )); } + Err(GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + .. + }) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } Err(_) => { return Ok(build_internal_control_error_response( http::StatusCode::SERVICE_UNAVAILABLE, @@ -61,12 +68,42 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider( )); } }; - if !state_data.key_id.trim().is_empty() || state_data.provider_id != provider_id { + if !preview.key_id.trim().is_empty() + || preview.provider_id != provider_id + || !admin_provider_oauth_state_matches_principal(&preview, request_context) + { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, "state 无效或已过期", )); } + let state_data = match state + .consume_provider_oauth_state(&callback.state_nonce) + .await + { + Ok(Some(state_data)) if state_data == preview => state_data, + Ok(Some(_)) | Ok(None) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } + Err(GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + .. + }) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "state 无效或已过期", + )); + } + Err(_) => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth redis unavailable", + )); + } + }; let Some(provider) = state .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/shared.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/shared.rs index 811d52df5..3ccf349a7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/shared.rs @@ -1,5 +1,8 @@ use super::super::super::errors::build_internal_control_error_response; use super::super::super::state::parse_provider_oauth_callback_params; +use crate::control::GatewayAdminPrincipalContext; +use crate::handlers::admin::request::AdminRequestContext; +use aether_data::repository::provider_oauth::StoredAdminProviderOAuthState; use axum::{ body::{Body, Bytes}, http, @@ -17,6 +20,31 @@ pub(super) struct AdminProviderOAuthCompleteCallback { pub(super) state_nonce: String, } +pub(super) fn admin_provider_oauth_state_matches_principal( + state: &StoredAdminProviderOAuthState, + request_context: &AdminRequestContext<'_>, +) -> bool { + admin_provider_oauth_state_matches_resolved_principal( + state, + request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()), + ) +} + +fn admin_provider_oauth_state_matches_resolved_principal( + state: &StoredAdminProviderOAuthState, + principal: Option<&GatewayAdminPrincipalContext>, +) -> bool { + let Some(principal) = principal else { + return false; + }; + state.initiated_by_user_id == principal.user_id + && state.initiated_by_session_id == principal.session_id + && state.initiated_by_management_token_id == principal.management_token_id + && (principal.session_id.is_some() || principal.management_token_id.is_some()) +} + pub(super) fn parse_admin_provider_oauth_callback_url( raw_payload: &serde_json::Map, ) -> Result> { @@ -114,3 +142,59 @@ pub(super) fn parse_admin_provider_oauth_complete_callback( Ok(AdminProviderOAuthCompleteCallback { code, state_nonce }) } + +#[cfg(test)] +mod tests { + use super::admin_provider_oauth_state_matches_resolved_principal; + use crate::control::GatewayAdminPrincipalContext; + use aether_data::repository::provider_oauth::StoredAdminProviderOAuthState; + + fn state() -> StoredAdminProviderOAuthState { + StoredAdminProviderOAuthState { + nonce: "a".repeat(64), + key_id: "key-1".to_string(), + provider_id: "provider-1".to_string(), + provider_type: "codex".to_string(), + pkce_verifier: Some("verifier".to_string()), + expected_encrypted_auth_config: None, + initiated_by_user_id: "admin-1".to_string(), + initiated_by_session_id: Some("session-1".to_string()), + initiated_by_management_token_id: None, + created_at: 1, + } + } + + fn principal(user_id: &str, session_id: Option<&str>) -> GatewayAdminPrincipalContext { + GatewayAdminPrincipalContext { + user_id: user_id.to_string(), + user_role: "admin".to_string(), + session_id: session_id.map(ToOwned::to_owned), + management_token_id: None, + management_token_permissions: None, + } + } + + #[test] + fn provider_oauth_state_is_bound_to_exact_admin_session() { + let state = state(); + let matching = principal("admin-1", Some("session-1")); + let wrong_user = principal("admin-2", Some("session-1")); + let wrong_session = principal("admin-1", Some("session-2")); + + assert!(admin_provider_oauth_state_matches_resolved_principal( + &state, + Some(&matching) + )); + assert!(!admin_provider_oauth_state_matches_resolved_principal( + &state, + Some(&wrong_user) + )); + assert!(!admin_provider_oauth_state_matches_resolved_principal( + &state, + Some(&wrong_session) + )); + assert!(!admin_provider_oauth_state_matches_resolved_principal( + &state, None + )); + } +} 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 19d32d9b7..b4ce3c5e5 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 @@ -192,6 +192,18 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( "设备授权仅支持 Kiro / Windsurf provider", )); } + let Some(principal) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + .filter(|principal| { + principal.session_id.is_some() || principal.management_token_id.is_some() + }) + else { + return Ok(build_internal_control_error_response( + http::StatusCode::UNAUTHORIZED, + "管理员身份不可用", + )); + }; let endpoint_resolution = resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?; let runtime_endpoint = endpoint_resolution.runtime_endpoint; @@ -240,10 +252,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( .build_authorize_url(&ctx, &session_id, None) { Ok(authorization) => authorization, - Err(error) => { + Err(_) => { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - format!("Windsurf 授权 URL 构建失败: {error}"), + "Windsurf 授权 URL 构建失败", )); } }; @@ -251,7 +263,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( build_windsurf_authorization_url(&authorization.authorize_url, &login_option); let now_unix_secs = current_unix_secs(); let session = StoredAdminProviderOAuthDeviceSession { + session_id: session_id.clone(), provider_id: provider_id.clone(), + 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: String::new(), client_secret: String::new(), @@ -330,7 +346,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( ); let now_unix_secs = current_unix_secs(); let session = StoredAdminProviderOAuthDeviceSession { + session_id: session_id.clone(), provider_id: provider_id.clone(), + 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: "us-east-1".to_string(), client_id: String::new(), client_secret: String::new(), @@ -468,7 +488,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( 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.clone(), + 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, client_id, client_secret, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/lease.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/lease.rs new file mode 100644 index 000000000..3c2273fa0 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/lease.rs @@ -0,0 +1,277 @@ +use crate::handlers::admin::request::AdminAppState; +use aether_runtime_state::{RuntimeLockLease, RuntimeState}; +use sha2::{Digest, Sha256}; +use std::future::Future; +use std::time::Duration; + +const DEVICE_POLL_LEASE_TTL: Duration = Duration::from_secs(300); +const DEVICE_POLL_LEASE_RENEW_INTERVAL: Duration = Duration::from_secs(60); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum AdminProviderOAuthDevicePollLeaseFailure { + Lost, + Unavailable, +} + +pub(super) enum AdminProviderOAuthDevicePollLeaseAcquire { + Acquired(AdminProviderOAuthDevicePollLease), + Contended, + Unavailable, +} + +pub(super) struct AdminProviderOAuthDevicePollLease { + runtime: RuntimeState, + lease: Option, +} + +impl AdminProviderOAuthDevicePollLease { + pub(super) async fn try_acquire( + state: &AdminAppState<'_>, + session_id: &str, + ) -> AdminProviderOAuthDevicePollLeaseAcquire { + let runtime = state.runtime_state().clone(); + let lock_key = admin_provider_oauth_device_poll_lock_key(session_id); + let owner = format!( + "aether-gateway-admin-provider-oauth-device-poll:{}", + uuid::Uuid::new_v4() + ); + match runtime + .lock_try_acquire(&lock_key, &owner, DEVICE_POLL_LEASE_TTL) + .await + { + Ok(Some(lease)) => AdminProviderOAuthDevicePollLeaseAcquire::Acquired(Self { + runtime, + lease: Some(lease), + }), + Ok(None) => AdminProviderOAuthDevicePollLeaseAcquire::Contended, + Err(error) => { + tracing::warn!( + lock_key = %lock_key, + error = ?error, + "gateway provider OAuth device poll lease acquisition failed" + ); + AdminProviderOAuthDevicePollLeaseAcquire::Unavailable + } + } + } + + pub(super) async fn run( + &self, + operation: F, + ) -> Result + where + F: Future, + { + let output = prefer_lease_loss( + operation, + wait_for_admin_provider_oauth_device_poll_lease_loss( + self.runtime.clone(), + self.lease + .as_ref() + .expect("an acquired device poll lease must contain its runtime lease") + .clone(), + ), + ) + .await?; + self.confirm_ownership().await?; + Ok(output) + } + + async fn confirm_ownership(&self) -> Result<(), AdminProviderOAuthDevicePollLeaseFailure> { + let lease = self + .lease + .as_ref() + .expect("an acquired device poll lease must contain its runtime lease"); + match self.runtime.lock_renew(lease, DEVICE_POLL_LEASE_TTL).await { + Ok(renewed) => ensure_admin_provider_oauth_device_poll_lease_renewed(renewed), + Err(error) => { + tracing::error!( + lock_key = %lease.key, + error = ?error, + "gateway provider OAuth device poll final lease renewal failed" + ); + Err(AdminProviderOAuthDevicePollLeaseFailure::Unavailable) + } + } + } + + pub(super) async fn release(mut self) { + let Some(lease) = self.lease.as_ref().cloned() else { + return; + }; + match self.runtime.lock_release(&lease).await { + Ok(_) => { + self.lease.take(); + } + Err(error) => { + tracing::warn!( + lock_key = %lease.key, + error = ?error, + "gateway provider OAuth device poll lease release failed" + ); + // Keep the lease in the guard so Drop can make one best-effort retry. + } + } + } +} + +impl Drop for AdminProviderOAuthDevicePollLease { + fn drop(&mut self) { + let Some(lease) = self.lease.take() else { + return; + }; + let runtime = self.runtime.clone(); + let Ok(handle) = tokio::runtime::Handle::try_current() else { + return; + }; + handle.spawn(async move { + if let Err(error) = runtime.lock_release(&lease).await { + tracing::warn!( + lock_key = %lease.key, + error = ?error, + "gateway provider OAuth device poll lease Drop release failed" + ); + } + }); + } +} + +fn admin_provider_oauth_device_poll_lock_key(session_id: &str) -> String { + format!( + "admin-provider-oauth-device-poll:sha256:{:x}", + Sha256::digest(session_id.as_bytes()) + ) +} + +fn ensure_admin_provider_oauth_device_poll_lease_renewed( + renewed: bool, +) -> Result<(), AdminProviderOAuthDevicePollLeaseFailure> { + if renewed { + Ok(()) + } else { + Err(AdminProviderOAuthDevicePollLeaseFailure::Lost) + } +} + +async fn wait_for_admin_provider_oauth_device_poll_lease_loss( + runtime: RuntimeState, + lease: RuntimeLockLease, +) -> AdminProviderOAuthDevicePollLeaseFailure { + let first_renewal = tokio::time::Instant::now() + DEVICE_POLL_LEASE_RENEW_INTERVAL; + let mut renewal_timer = + tokio::time::interval_at(first_renewal, DEVICE_POLL_LEASE_RENEW_INTERVAL); + renewal_timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + loop { + renewal_timer.tick().await; + match runtime.lock_renew(&lease, DEVICE_POLL_LEASE_TTL).await { + Ok(true) => {} + Ok(false) => { + tracing::error!( + lock_key = %lease.key, + "gateway provider OAuth device poll lease was lost" + ); + return AdminProviderOAuthDevicePollLeaseFailure::Lost; + } + Err(error) => { + tracing::error!( + lock_key = %lease.key, + error = ?error, + "gateway provider OAuth device poll lease renewal failed" + ); + return AdminProviderOAuthDevicePollLeaseFailure::Unavailable; + } + } + } +} + +async fn prefer_lease_loss( + operation: Operation, + lease_loss: LeaseLoss, +) -> Result +where + Operation: Future, + LeaseLoss: Future, +{ + tokio::pin!(operation); + tokio::pin!(lease_loss); + tokio::select! { + biased; + failure = &mut lease_loss => Err(failure), + output = &mut operation => Ok(output), + } +} + +#[cfg(test)] +mod tests { + use super::{ + admin_provider_oauth_device_poll_lock_key, + ensure_admin_provider_oauth_device_poll_lease_renewed, prefer_lease_loss, + AdminProviderOAuthDevicePollLeaseFailure, + }; + use std::future::Future; + use std::pin::Pin; + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; + use std::task::{Context, Poll}; + + struct ReadyOperation { + polled: Arc, + dropped: Arc, + } + + impl Future for ReadyOperation { + type Output = &'static str; + + fn poll(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll { + self.polled.store(true, Ordering::Release); + Poll::Ready("must not be published") + } + } + + impl Drop for ReadyOperation { + fn drop(&mut self) { + self.dropped.store(true, Ordering::Release); + } + } + + #[test] + fn device_poll_lease_lock_key_hashes_session_id() { + let session_id = "secret-device-session"; + let first = admin_provider_oauth_device_poll_lock_key(session_id); + let second = admin_provider_oauth_device_poll_lock_key(session_id); + + assert_eq!(first, second); + assert!(!first.contains(session_id)); + } + + #[test] + fn device_poll_lease_renew_false_fails_closed() { + assert_eq!( + ensure_admin_provider_oauth_device_poll_lease_renewed(false), + Err(AdminProviderOAuthDevicePollLeaseFailure::Lost) + ); + assert!(ensure_admin_provider_oauth_device_poll_lease_renewed(true).is_ok()); + } + + #[tokio::test] + async fn device_poll_lease_loss_future_has_priority_and_cancels_operation() { + let polled = Arc::new(AtomicBool::new(false)); + let dropped = Arc::new(AtomicBool::new(false)); + let operation = ReadyOperation { + polled: Arc::clone(&polled), + dropped: Arc::clone(&dropped), + }; + + let result = prefer_lease_loss( + operation, + std::future::ready(AdminProviderOAuthDevicePollLeaseFailure::Lost), + ) + .await; + + assert_eq!(result, Err(AdminProviderOAuthDevicePollLeaseFailure::Lost)); + assert!(!polled.load(Ordering::Acquire)); + assert!(dropped.load(Ordering::Acquire)); + } +} 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 0b4932323..295a93336 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 @@ -1,4 +1,5 @@ mod authorize; +mod lease; mod poll; mod session; 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 144303403..d30f52598 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 @@ -2,9 +2,11 @@ use super::super::kiro::{ admin_provider_oauth_kiro_refresh_base_url_override, fetch_admin_provider_oauth_kiro_email, refresh_admin_provider_oauth_kiro_auth_config, }; +use super::lease::{AdminProviderOAuthDevicePollLease, AdminProviderOAuthDevicePollLeaseAcquire}; use super::session::{ attach_admin_provider_oauth_device_poll_terminal_response, AdminProviderOAuthDevicePollPayload, }; +use crate::control::GatewayAdminPrincipalContext; 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, @@ -103,11 +105,7 @@ fn sanitize_windsurf_browser_poll_detail(detail: impl AsRef) -> String { if detail.is_empty() { return "-".to_string(); } - if contains_windsurf_sensitive_marker(detail) { - "[REDACTED upstream error body]".to_string() - } else { - detail.chars().take(500).collect() - } + "[REDACTED upstream error body]".to_string() } fn sanitize_windsurf_browser_poll_callback_error(error: &str, description: &str) -> String { @@ -141,25 +139,6 @@ fn sanitize_windsurf_browser_poll_oauth_error(error: &aether_oauth::core::OAuthE } } -fn contains_windsurf_sensitive_marker(value: &str) -> bool { - let lowered = value.to_ascii_lowercase(); - [ - "token", - "api_key", - "apikey", - "sessiontoken", - "firebase_id_token", - "idtoken", - "authorization", - "password", - "secret", - "devin-session-token$", - ] - .iter() - .any(|marker| lowered.contains(marker)) - || value.contains("sk-") -} - fn secret_fingerprint(value: &str) -> Option { let value = value.trim(); if value.is_empty() { @@ -301,14 +280,9 @@ async fn exchange_admin_provider_oauth_kiro_social_code( proxy, ) .await - .map_err(|err| format!("Kiro social token 请求失败: {err}"))?; + .map_err(|_| "Kiro social token 请求失败".to_string())?; if !response.status.is_success() { - let detail = response.body_text.trim(); - return Err(if detail.is_empty() { - format!("HTTP {}", response.status.as_u16()) - } else { - detail.to_string() - }); + return Err(format!("HTTP {}", response.status.as_u16())); } response .json_body @@ -316,11 +290,64 @@ async fn exchange_admin_provider_oauth_kiro_social_code( .ok_or_else(|| "Kiro social token 返回了非 JSON 响应".to_string()) } +fn admin_provider_oauth_device_poll_principal_has_authenticator( + principal: &GatewayAdminPrincipalContext, +) -> bool { + let valid_identity = + |value: Option<&str>| value.is_some_and(|value| !value.is_empty() && value == value.trim()); + valid_identity(principal.session_id.as_deref()) + || valid_identity(principal.management_token_id.as_deref()) +} + +fn admin_provider_oauth_device_session_matches_resolved_principal( + session: &StoredAdminProviderOAuthDeviceSession, + principal: Option<&GatewayAdminPrincipalContext>, +) -> bool { + let Some(principal) = principal else { + return false; + }; + admin_provider_oauth_device_poll_principal_has_authenticator(principal) + && session.initiated_by_user_id == principal.user_id + && session.initiated_by_session_id == principal.session_id + && session.initiated_by_management_token_id == principal.management_token_id +} + +fn admin_provider_oauth_device_session_unavailable_response() -> Response { + Json(json!({ + "status": "expired", + "error": "会话不存在或已过期", + "replaced": false, + })) + .into_response() +} + +fn admin_provider_oauth_device_poll_busy_response(status: http::StatusCode) -> Response { + let mut response = Json(json!({ + "status": "pending", + "busy": true, + "replaced": false, + })) + .into_response(); + *response.status_mut() = status; + response +} + pub(super) async fn handle_admin_provider_oauth_device_poll( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, request_body: Option<&Bytes>, ) -> Result, GatewayError> { + let Some(principal) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + .filter(|principal| admin_provider_oauth_device_poll_principal_has_authenticator(principal)) + .cloned() + else { + return Ok(build_internal_control_error_response( + http::StatusCode::UNAUTHORIZED, + "管理员身份不可用", + )); + }; if !state.has_provider_catalog_data_reader() { return Ok(build_admin_provider_oauth_backend_unavailable_response()); } @@ -355,140 +382,57 @@ pub(super) async fn handle_admin_provider_oauth_device_poll( )); } - let Some(mut session) = state.read_provider_oauth_device_session(session_id).await? else { - return Ok(Json(json!({ - "status": "expired", - "error": "会话不存在或已过期", - "replaced": false, - })) - .into_response()); + let lease = match AdminProviderOAuthDevicePollLease::try_acquire(state, session_id).await { + AdminProviderOAuthDevicePollLeaseAcquire::Acquired(lease) => lease, + AdminProviderOAuthDevicePollLeaseAcquire::Contended => { + return Ok(admin_provider_oauth_device_poll_busy_response( + http::StatusCode::CONFLICT, + )); + } + AdminProviderOAuthDevicePollLeaseAcquire::Unavailable => { + return Ok(admin_provider_oauth_device_poll_busy_response( + http::StatusCode::SERVICE_UNAVAILABLE, + )); + } }; - if session.provider_id != provider_id { - return Ok(Json(json!({ - "status": "error", - "error": "会话与 Provider 不匹配", - "replaced": false, - })) - .into_response()); - } - if session.status == "authorized" { - return Ok(Json(json!({ - "status": "authorized", - "key_id": session.key_id, - "email": session.email, - "replaced": session.replaced, - })) - .into_response()); - } - if matches!(session.status.as_str(), "expired" | "error") { - return Ok(Json(json!({ - "status": session.status, - "error": session.error_msg, - "replaced": session.replaced, - })) - .into_response()); - } - if current_unix_secs() > session.expires_at_unix_secs { - session.status = "expired".to_string(); - session.error_msg = Some("设备码已过期".to_string()); - let _ = state - .save_provider_oauth_device_session(session_id, &session, 30) - .await; - return Ok(attach_admin_provider_oauth_device_poll_terminal_response( - session_id, - "expired", - Json(json!({ - "status": "expired", - "error": "设备码已过期", + let operation = async { + let Some(mut session) = state.read_provider_oauth_device_session(session_id).await? else { + return Ok(admin_provider_oauth_device_session_unavailable_response()); + }; + if !admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&principal), + ) { + return Ok(admin_provider_oauth_device_session_unavailable_response()); + } + if session.provider_id != provider_id { + return Ok(Json(json!({ + "status": "error", + "error": "会话与 Provider 不匹配", "replaced": false, })) - .into_response(), - )); - } - - let Some(provider) = state - .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) - .await? - .into_iter() - .next() - else { - return Ok(build_internal_control_error_response( - http::StatusCode::NOT_FOUND, - "Provider 不存在", - )); - }; - let provider_type = provider.provider_type.trim().to_ascii_lowercase(); - let endpoint_resolution = - resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?; - let endpoints = endpoint_resolution.endpoints; - let runtime_endpoint = endpoint_resolution.runtime_endpoint; - let request_proxy = state - .resolve_admin_provider_oauth_operation_proxy_snapshot( - session.proxy_node_id.as_deref(), - &[ - runtime_endpoint - .as_ref() - .and_then(|endpoint| endpoint.proxy.as_ref()), - provider.proxy.as_ref(), - ], - ) - .await; - - if provider_type == "windsurf" { - return handle_admin_provider_oauth_windsurf_browser_device_poll( - state, - &provider, - &endpoints, - request_proxy, - session_id, - session, - payload.callback_url.as_deref(), - payload.token.as_deref(), - ) - .await; - } - - if kiro_device_session_is_social(&session) { - return handle_admin_provider_oauth_kiro_social_device_poll( - state, - &provider, - &endpoints, - request_proxy, - session_id, - session, - payload.callback_url.as_deref(), - ) - .await; - } - - let token_result = match state - .poll_admin_kiro_device_token( - &session.region, - &session.client_id, - &session.client_secret, - &session.device_code, - request_proxy.clone(), - ) - .await - { - Ok(payload) => payload, - Err(response) => return Ok(response), - }; - - if token_result - .get("_error") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false) - { - let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default(); - if error_code == "authorization_pending" { - return Ok(Json(json!({"status": "pending", "replaced": false})).into_response()); + .into_response()); } - if error_code == "slow_down" { - return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response()); + if session.status == "authorized" { + return Ok(Json(json!({ + "status": "authorized", + "key_id": session.key_id, + "email": session.email, + "replaced": session.replaced, + })) + .into_response()); } - if error_code == "expired_token" { + if matches!(session.status.as_str(), "expired" | "error") { + return Ok(Json(json!({ + "status": session.status, + "error": session.error_msg, + "replaced": session.replaced, + })) + .into_response()); + } + + if current_unix_secs() > session.expires_at_unix_secs { session.status = "expired".to_string(); session.error_msg = Some("设备码已过期".to_string()); let _ = state @@ -505,240 +449,351 @@ pub(super) async fn handle_admin_provider_oauth_device_poll( .into_response(), )); } - if error_code == "access_denied" { - session.status = "error".to_string(); - session.error_msg = Some("用户拒绝授权".to_string()); - let _ = state - .save_provider_oauth_device_session(session_id, &session, 30) - .await; - return Ok(attach_admin_provider_oauth_device_poll_terminal_response( + + let Some(provider) = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) + .await? + .into_iter() + .next() + else { + return Ok(build_internal_control_error_response( + http::StatusCode::NOT_FOUND, + "Provider 不存在", + )); + }; + let provider_type = provider.provider_type.trim().to_ascii_lowercase(); + let endpoint_resolution = + resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?; + let endpoints = endpoint_resolution.endpoints; + let runtime_endpoint = endpoint_resolution.runtime_endpoint; + let request_proxy = state + .resolve_admin_provider_oauth_operation_proxy_snapshot( + session.proxy_node_id.as_deref(), + &[ + runtime_endpoint + .as_ref() + .and_then(|endpoint| endpoint.proxy.as_ref()), + provider.proxy.as_ref(), + ], + ) + .await; + + if provider_type == "windsurf" { + return handle_admin_provider_oauth_windsurf_browser_device_poll( + state, + &provider, + &endpoints, + request_proxy, session_id, - "error", - Json(json!({ + session, + payload.callback_url.as_deref(), + payload.token.as_deref(), + ) + .await; + } + + if kiro_device_session_is_social(&session) { + return handle_admin_provider_oauth_kiro_social_device_poll( + state, + &provider, + &endpoints, + request_proxy, + session_id, + session, + payload.callback_url.as_deref(), + ) + .await; + } + + let token_result = match state + .poll_admin_kiro_device_token( + &session.region, + &session.client_id, + &session.client_secret, + &session.device_code, + request_proxy.clone(), + ) + .await + { + Ok(payload) => payload, + Err(response) => return Ok(response), + }; + + if token_result + .get("_error") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + { + let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default(); + if error_code == "authorization_pending" { + return Ok(Json(json!({"status": "pending", "replaced": false})).into_response()); + } + if error_code == "slow_down" { + return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response()); + } + if error_code == "expired_token" { + session.status = "expired".to_string(); + session.error_msg = Some("设备码已过期".to_string()); + let _ = state + .save_provider_oauth_device_session(session_id, &session, 30) + .await; + return Ok(attach_admin_provider_oauth_device_poll_terminal_response( + session_id, + "expired", + Json(json!({ + "status": "expired", + "error": "设备码已过期", + "replaced": false, + })) + .into_response(), + )); + } + if error_code == "access_denied" { + session.status = "error".to_string(); + session.error_msg = Some("用户拒绝授权".to_string()); + let _ = state + .save_provider_oauth_device_session(session_id, &session, 30) + .await; + return Ok(attach_admin_provider_oauth_device_poll_terminal_response( + session_id, + "error", + Json(json!({ + "status": "error", + "error": "用户拒绝授权", + "replaced": false, + })) + .into_response(), + )); + } + let error_message = if error_code.is_empty() { + "授权失败".to_string() + } else { + sanitize_windsurf_browser_poll_error_code(&error_code) + }; + return Ok(Json(json!({ + "status": "error", + "error": error_message, + "replaced": false, + })) + .into_response()); + } + + let Some(access_token) = json_non_empty_string( + token_result + .get("accessToken") + .or_else(|| token_result.get("access_token")), + ) else { + return Ok(Json(json!({ + "status": "error", + "error": "token 响应缺少 accessToken 或 refreshToken", + "replaced": false, + })) + .into_response()); + }; + let Some(refresh_token) = json_non_empty_string( + token_result + .get("refreshToken") + .or_else(|| token_result.get("refresh_token")), + ) else { + return Ok(Json(json!({ + "status": "error", + "error": "token 响应缺少 accessToken 或 refreshToken", + "replaced": false, + })) + .into_response()); + }; + let initial_expires_at = json_u64_value( + token_result + .get("expiresIn") + .or_else(|| token_result.get("expires_in")), + ) + .map(|expires_in| current_unix_secs().saturating_add(expires_in)) + .unwrap_or_else(|| current_unix_secs().saturating_add(3600)); + let social_refresh_base_url = + admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh"); + let idc_refresh_base_url = + admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh"); + let mut refreshed_auth_config = match refresh_admin_provider_oauth_kiro_auth_config( + state, + &AdminKiroAuthConfig { + auth_method: Some("idc".to_string()), + refresh_token: Some(refresh_token.clone()), + expires_at: Some(initial_expires_at), + profile_arn: None, + region: Some(session.region.clone()), + auth_region: Some(session.region.clone()), + api_region: None, + client_id: Some(session.client_id.clone()), + client_secret: Some(session.client_secret.clone()), + machine_id: None, + kiro_version: None, + system_version: None, + node_version: None, + access_token: Some(access_token.clone()), + }, + request_proxy.clone(), + social_refresh_base_url.as_deref(), + idc_refresh_base_url.as_deref(), + ) + .await + { + Ok(config) => config, + Err(_) => { + return Ok(Json(json!({ "status": "error", - "error": "用户拒绝授权", + "error": "token 验证失败", "replaced": false, })) - .into_response(), - )); + .into_response()); + } + }; + if refreshed_auth_config.auth_method.is_none() { + refreshed_auth_config.auth_method = Some("idc".to_string()); } - let error_message = json_non_empty_string(token_result.get("error_description")) - .or_else(|| (!error_code.is_empty()).then_some(error_code.clone())) - .unwrap_or_else(|| "未知错误".to_string()); - return Ok(Json(json!({ - "status": "error", - "error": error_message, - "replaced": false, - })) - .into_response()); - } - - let Some(access_token) = json_non_empty_string( - token_result - .get("accessToken") - .or_else(|| token_result.get("access_token")), - ) else { - return Ok(Json(json!({ - "status": "error", - "error": "token 响应缺少 accessToken 或 refreshToken", - "replaced": false, - })) - .into_response()); - }; - let Some(refresh_token) = json_non_empty_string( - token_result - .get("refreshToken") - .or_else(|| token_result.get("refresh_token")), - ) else { - return Ok(Json(json!({ - "status": "error", - "error": "token 响应缺少 accessToken 或 refreshToken", - "replaced": false, - })) - .into_response()); - }; - let initial_expires_at = json_u64_value( - token_result - .get("expiresIn") - .or_else(|| token_result.get("expires_in")), - ) - .map(|expires_in| current_unix_secs().saturating_add(expires_in)) - .unwrap_or_else(|| current_unix_secs().saturating_add(3600)); - let social_refresh_base_url = - admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh"); - let idc_refresh_base_url = - admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh"); - let mut refreshed_auth_config = match refresh_admin_provider_oauth_kiro_auth_config( - state, - &AdminKiroAuthConfig { - auth_method: Some("idc".to_string()), - refresh_token: Some(refresh_token.clone()), - expires_at: Some(initial_expires_at), - profile_arn: None, - region: Some(session.region.clone()), - auth_region: Some(session.region.clone()), - api_region: None, - client_id: Some(session.client_id.clone()), - client_secret: Some(session.client_secret.clone()), - machine_id: None, - kiro_version: None, - system_version: None, - node_version: None, - access_token: Some(access_token.clone()), - }, - request_proxy.clone(), - social_refresh_base_url.as_deref(), - idc_refresh_base_url.as_deref(), - ) - .await - { - Ok(config) => config, - Err(detail) => { + let Some(access_token) = refreshed_auth_config + .access_token + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { return Ok(Json(json!({ "status": "error", - "error": format!("token 验证失败: {detail}"), + "error": "token 验证失败: accessToken 为空", "replaced": false, })) .into_response()); + }; + let expires_at = refreshed_auth_config + .expires_at + .unwrap_or_else(|| current_unix_secs().saturating_add(3600)); + let mut email = decode_jwt_claims(&access_token) + .and_then(|claims| claims.get("email").cloned()) + .and_then(|value| value.as_str().map(ToOwned::to_owned)); + if email.is_none() { + email = fetch_admin_provider_oauth_kiro_email( + state, + &refreshed_auth_config, + request_proxy.clone(), + ) + .await; } - }; - if refreshed_auth_config.auth_method.is_none() { - refreshed_auth_config.auth_method = Some("idc".to_string()); - } - let Some(access_token) = refreshed_auth_config - .access_token - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - else { - return Ok(Json(json!({ - "status": "error", - "error": "token 验证失败: accessToken 为空", - "replaced": false, - })) - .into_response()); - }; - let expires_at = refreshed_auth_config - .expires_at - .unwrap_or_else(|| current_unix_secs().saturating_add(3600)); - let mut email = decode_jwt_claims(&access_token) - .and_then(|claims| claims.get("email").cloned()) - .and_then(|value| value.as_str().map(ToOwned::to_owned)); - if email.is_none() { - email = fetch_admin_provider_oauth_kiro_email( - state, - &refreshed_auth_config, + + let mut auth_config = refreshed_auth_config + .to_json_value() + .as_object() + .cloned() + .unwrap_or_default(); + auth_config.insert("provider_type".to_string(), json!("kiro")); + if let Some(email) = email.as_ref() { + auth_config.insert("email".to_string(), json!(email)); + } + + 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 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(), + Some(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 = build_kiro_device_key_name( + email.as_deref(), + refreshed_auth_config.refresh_token.as_deref(), + ); + match state + .create_provider_oauth_catalog_key( + &provider_id, + &provider.provider_type, + &key_name, + &access_token, + &auth_config, + &api_formats, + key_proxy, + Some(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(), - ) - .await; - } - - let mut auth_config = refreshed_auth_config - .to_json_value() - .as_object() - .cloned() - .unwrap_or_default(); - auth_config.insert("provider_type".to_string(), json!("kiro")); - if let Some(email) = email.as_ref() { - auth_config.insert("email".to_string(), json!(email)); - } - - 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 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(), - Some(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 = build_kiro_device_key_name( - email.as_deref(), - refreshed_auth_config.refresh_token.as_deref(), ); - match state - .create_provider_oauth_catalog_key( - &provider_id, - &provider.provider_type, - &key_name, - &access_token, - &auth_config, - &api_formats, - key_proxy, - Some(expires_at), - ) - .await? - { - Some(key) => key, - None => { - return Ok(build_internal_control_error_response( - http::StatusCode::SERVICE_UNAVAILABLE, - "provider oauth write unavailable", - )); - } - } + + 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(), + )) }; - 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(), - )) + let result = lease.run(operation).await; + lease.release().await; + match result { + Ok(result) => result, + Err(_) => Ok(admin_provider_oauth_device_poll_busy_response( + http::StatusCode::SERVICE_UNAVAILABLE, + )), + } } fn windsurf_raw_api_key(value: &str) -> Option<&str> { @@ -1023,19 +1078,17 @@ async fn handle_admin_provider_oauth_kiro_social_device_poll( let callback_params = parse_provider_oauth_callback_params(callback_url); if let Some(error) = callback_params.get("error").map(String::as_str) { - let error_description = callback_params - .get("error_description") - .map(String::as_str) - .unwrap_or("用户拒绝授权"); + let error = sanitize_windsurf_browser_poll_error_code(error); + let error_message = format!("{error}: 授权失败"); session.status = "error".to_string(); - session.error_msg = Some(format!("{error}: {error_description}")); + session.error_msg = Some(error_message.clone()); let _ = state .save_provider_oauth_device_session(session_id, &session, 30) .await; return Ok(attach_admin_provider_oauth_device_poll_terminal_response( session_id, "error", - kiro_social_poll_error_response(format!("{error}: {error_description}")), + kiro_social_poll_error_response(error_message), )); } @@ -1094,16 +1147,16 @@ async fn handle_admin_provider_oauth_kiro_social_device_poll( .await { Ok(payload) => payload, - Err(detail) => { + Err(_) => { session.status = "error".to_string(); - session.error_msg = Some(format!("token exchange 失败: {detail}")); + session.error_msg = Some("token exchange 失败".to_string()); let _ = state .save_provider_oauth_device_session(session_id, &session, 30) .await; return Ok(attach_admin_provider_oauth_device_poll_terminal_response( session_id, "error", - kiro_social_poll_error_response(format!("token exchange 失败: {detail}")), + kiro_social_poll_error_response("token exchange 失败"), )); } }; @@ -1299,6 +1352,105 @@ async fn handle_admin_provider_oauth_kiro_social_device_poll( #[cfg(test)] mod tests { + use super::admin_provider_oauth_device_session_matches_resolved_principal; + use crate::control::GatewayAdminPrincipalContext; + use aether_data::repository::provider_oauth::StoredAdminProviderOAuthDeviceSession; + + fn device_session() -> StoredAdminProviderOAuthDeviceSession { + StoredAdminProviderOAuthDeviceSession { + session_id: "device-session-1".to_string(), + provider_id: "provider-1".to_string(), + initiated_by_user_id: "admin-1".to_string(), + initiated_by_session_id: Some("admin-session-1".to_string()), + initiated_by_management_token_id: None, + region: "us-east-1".to_string(), + client_id: "client-1".to_string(), + client_secret: "client-secret".to_string(), + device_code: "device-code".to_string(), + auth_type: None, + social_provider: None, + code_verifier: None, + redirect_uri: None, + machine_id: None, + interval: 5, + expires_at_unix_secs: 2, + status: "pending".to_string(), + proxy_node_id: None, + created_at_unix_ms: 1, + key_id: None, + email: None, + replaced: false, + error_msg: None, + } + } + + fn principal( + user_id: &str, + session_id: Option<&str>, + management_token_id: Option<&str>, + ) -> GatewayAdminPrincipalContext { + GatewayAdminPrincipalContext { + user_id: user_id.to_string(), + user_role: "admin".to_string(), + session_id: session_id.map(ToOwned::to_owned), + management_token_id: management_token_id.map(ToOwned::to_owned), + management_token_permissions: None, + } + } + + #[test] + fn device_poll_session_is_bound_to_exact_admin_principal() { + let session = device_session(); + let matching = principal("admin-1", Some("admin-session-1"), None); + let wrong_user = principal("admin-2", Some("admin-session-1"), None); + let wrong_session = principal("admin-1", Some("admin-session-2"), None); + let wrong_authenticator = principal("admin-1", None, Some("token-1")); + let unbound = principal("admin-1", None, None); + + assert!( + admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&matching) + ) + ); + assert!( + !admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&wrong_user) + ) + ); + assert!( + !admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&wrong_session) + ) + ); + assert!( + !admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&wrong_authenticator) + ) + ); + assert!( + !admin_provider_oauth_device_session_matches_resolved_principal( + &session, + Some(&unbound) + ) + ); + assert!(!admin_provider_oauth_device_session_matches_resolved_principal(&session, None)); + + let mut token_session = session; + token_session.initiated_by_session_id = None; + token_session.initiated_by_management_token_id = Some("token-1".to_string()); + let matching_token = principal("admin-1", None, Some("token-1")); + assert!( + admin_provider_oauth_device_session_matches_resolved_principal( + &token_session, + Some(&matching_token) + ) + ); + } + #[test] fn windsurf_browser_poll_callback_error_redacts_sensitive_values() { let detail = super::sanitize_windsurf_browser_poll_callback_error( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/session.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/session.rs index c53203abc..470631b23 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/session.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/session.rs @@ -17,7 +17,7 @@ pub(super) struct AdminProviderOAuthDeviceAuthorizePayload { pub(super) proxy_node_id: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(super) struct AdminProviderOAuthDevicePollPayload { pub(super) session_id: String, pub(super) callback_url: Option, 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 ad48fce19..5ef0ea8b9 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 @@ -225,10 +225,9 @@ async fn prepare_codex_agent_identity_enrollment( "该 ChatGPT 账号正在创建 Agent Identity,请稍后重试", )); } - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, - error = ?error, "gateway Agent Identity enrollment lock unavailable" ); release_codex_agent_identity_leases(state, leases).await; @@ -296,11 +295,10 @@ fn spawn_codex_agent_identity_enrollment_heartbeat( ); return; } - Err(error) => { + Err(_) => { lease_lost.store(true, Ordering::Release); tracing::error!( lock_key = %lease.key, - error = ?error, "gateway Agent Identity enrollment lock renewal failed" ); return; @@ -316,10 +314,9 @@ async fn release_codex_agent_identity_leases( leases: Vec, ) { for lease in leases { - if let Err(error) = state.runtime_state().lock_release(&lease).await { + if state.runtime_state().lock_release(&lease).await.is_err() { tracing::warn!( lock_key = %lease.key, - error = ?error, "gateway Agent Identity enrollment lock release failed" ); } @@ -460,6 +457,9 @@ fn apply_single_import_hints( .or_insert_with(|| json!(project_id)); } for (target, keys) in [ + // Antigravity token responses omit the account email, so retain the + // identity carried by the imported credential payload. + ("email", &["email", "oauth_email"][..]), ( "client_version", &[ @@ -1198,11 +1198,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( match state.resolve_local_oauth_request_auth(&transport).await { Ok(Some(_)) => true, Ok(None) => false, - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, key_id = %persisted_key.id, - error = ?error, "gateway Agent Identity initial task registration failed" ); false @@ -1210,11 +1209,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( } } Ok(None) => false, - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, key_id = %persisted_key.id, - error = ?error, "gateway Agent Identity pending transport reload failed" ); false @@ -1461,7 +1459,8 @@ mod tests { }, "clientVersion": "1.99.0", "sessionId": "session-antigravity-1", - "userAgent": "antigravity" + "userAgent": "antigravity", + "email": "anti@example.com" }) .as_object() .cloned() @@ -1470,6 +1469,7 @@ mod tests { apply_single_import_hints("antigravity", &payload, &mut auth_config); + assert_eq!(auth_config.get("email"), Some(&json!("anti@example.com"))); assert_eq!( auth_config.get("project_id"), Some(&json!("project-antigravity-1")) diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/kiro.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/kiro.rs index 7c9926628..e76d82bea 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/kiro.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/kiro.rs @@ -56,21 +56,50 @@ fn admin_provider_oauth_kiro_refresh_error( "social refresh" }; match error { - OAuthError::HttpStatus { - status_code, - body_excerpt, - } => { - let detail = body_excerpt.trim(); - if detail.is_empty() { - format!("{prefix} 失败: HTTP {status_code}") - } else { - format!("{prefix} 失败: {detail}") - } + OAuthError::HttpStatus { status_code, .. } => { + format!("{prefix} 失败: HTTP {status_code}") } - OAuthError::Transport(message) => format!("{prefix} 请求失败: {message}"), - OAuthError::InvalidRequest(message) => format!("{prefix} 参数无效: {message}"), - OAuthError::InvalidResponse(message) => format!("{prefix} 返回无效响应: {message}"), - error => format!("{prefix} 失败: {error}"), + OAuthError::InvalidRequest(_) => format!("{prefix} 参数无效"), + OAuthError::InvalidResponse(_) => format!("{prefix} 返回无效响应"), + OAuthError::Transport(_) => format!("{prefix} 请求失败"), + _ => format!("{prefix} 失败"), + } +} + +#[cfg(test)] +mod refresh_error_tests { + use super::admin_provider_oauth_kiro_refresh_error; + use crate::handlers::admin::request::AdminKiroAuthConfig; + use aether_oauth::core::OAuthError; + + #[test] + fn kiro_refresh_error_does_not_reflect_upstream_body() { + let auth_config = AdminKiroAuthConfig { + auth_method: None, + refresh_token: None, + expires_at: None, + profile_arn: None, + region: None, + auth_region: None, + api_region: None, + client_id: None, + client_secret: None, + machine_id: None, + kiro_version: None, + system_version: None, + node_version: None, + access_token: None, + }; + let detail = admin_provider_oauth_kiro_refresh_error( + &auth_config, + OAuthError::HttpStatus { + status_code: 502, + body_excerpt: "authorization=Bearer upstream-secret".to_string(), + }, + ); + + assert_eq!(detail, "social refresh 失败: HTTP 502"); + assert!(!detail.contains("upstream-secret")); } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/execution.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/execution.rs index 936b3cb1a..fc4b96f84 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/execution.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/execution.rs @@ -88,43 +88,44 @@ pub(super) async fn execute_admin_provider_oauth_refresh( response::oauth_refresh_failed_bad_request_response(&error_reason), )); } - Err(AdminLocalOAuthRefreshError::Transport { source, .. }) => { + Err(AdminLocalOAuthRefreshError::Transport { .. }) => { tracing::warn!( trace_id = %trace_id, key_id = %key_id, provider_id = %provider.id, provider_type = %provider_type, - error = %source, "gateway manual provider oauth refresh transport failed" ); return Ok(RefreshDispatch::Respond( - response::oauth_refresh_failed_service_unavailable_response(source.to_string()), + response::oauth_refresh_failed_service_unavailable_response( + "Token 刷新网络请求失败", + ), )); } - Err(AdminLocalOAuthRefreshError::TransportMessage { message, .. }) => { + Err(AdminLocalOAuthRefreshError::TransportMessage { .. }) => { tracing::warn!( trace_id = %trace_id, key_id = %key_id, provider_id = %provider.id, provider_type = %provider_type, - error = %message, "gateway manual provider oauth refresh transport failed" ); return Ok(RefreshDispatch::Respond( - response::oauth_refresh_failed_service_unavailable_response(message), + response::oauth_refresh_failed_service_unavailable_response( + "Token 刷新网络请求失败", + ), )); } - Err(AdminLocalOAuthRefreshError::InvalidResponse { message, .. }) => { + Err(AdminLocalOAuthRefreshError::InvalidResponse { .. }) => { tracing::warn!( trace_id = %trace_id, key_id = %key_id, provider_id = %provider.id, provider_type = %provider_type, - reason = %message, "gateway manual provider oauth refresh returned invalid response" ); return Ok(RefreshDispatch::Respond( - response::oauth_refresh_failed_bad_request_response(&message), + response::oauth_refresh_failed_bad_request_response("Token 刷新响应无效"), )); } }; @@ -140,12 +141,7 @@ pub(super) async fn execute_admin_provider_oauth_refresh( .and_then(|entry| entry.metadata.as_ref()) .and_then(serde_json::Value::as_object) .cloned() - .unwrap_or_else(|| { - helpers::refreshed_auth_config_object( - state, - refreshed_key.encrypted_auth_config.as_deref(), - ) - }); + .unwrap_or_else(|| helpers::refreshed_auth_config_object(state, &refreshed_key)); let refreshed_expires_at_unix_secs = refreshed_entry .as_ref() .and_then(|entry| entry.expires_at_unix_secs) diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/helpers.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/helpers.rs index a07166bfa..8914ca0d6 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/helpers.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/helpers.rs @@ -1,5 +1,4 @@ use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; -use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogKey, StoredProviderCatalogProvider, }; @@ -31,9 +30,13 @@ pub(super) struct RefreshSuccessContext { pub(super) fn decrypt_auth_config( state: &AdminAppState<'_>, - encrypted_auth_config: &str, + key: &StoredProviderCatalogKey, ) -> Option { - state.decrypt_catalog_secret_with_fallbacks(encrypted_auth_config) + state + .app() + .decrypt_provider_catalog_key_auth_config(key) + .ok() + .flatten() } pub(super) fn parse_auth_config_object(plaintext: &str) -> Map { @@ -45,10 +48,9 @@ pub(super) fn parse_auth_config_object(plaintext: &str) -> Map { pub(super) fn refreshed_auth_config_object( state: &AdminAppState<'_>, - encrypted_auth_config: Option<&str>, + key: &StoredProviderCatalogKey, ) -> Map { - encrypted_auth_config - .and_then(|ciphertext| decrypt_auth_config(state, ciphertext)) + decrypt_auth_config(state, key) .map(|plaintext| parse_auth_config_object(&plaintext)) .unwrap_or_default() } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/request.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/request.rs index 1d507d339..486ae591e 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/request.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/refresh/request.rs @@ -29,14 +29,13 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request( "Key 不存在", ))); }; - let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else { + let Some(_encrypted_auth_config) = key.encrypted_auth_config.as_deref() else { return Ok(RefreshDispatch::Respond(response::control_error_response( http::StatusCode::BAD_REQUEST, "缺少 auth_config,无法 refresh", ))); }; - let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, encrypted_auth_config) - else { + let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, &key) else { return Ok(RefreshDispatch::Respond(response::control_error_response( http::StatusCode::SERVICE_UNAVAILABLE, "provider oauth encryption unavailable", 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 adde11074..d909b0320 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 @@ -8,6 +8,7 @@ use crate::handlers::admin::provider::shared::paths::{ admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::admin::shared::mark_sensitive_admin_response_no_store; use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::GatewayError; use axum::{ @@ -75,6 +76,15 @@ pub(super) async fn handle_admin_provider_oauth_start_key( "该 Provider 不支持 OAuth 授权", )); }; + let Some(principal) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + else { + return Ok(build_internal_control_error_response( + http::StatusCode::UNAUTHORIZED, + "管理员身份不可用", + )); + }; let pkce_verifier = template .use_pkce @@ -87,6 +97,9 @@ pub(super) async fn handle_admin_provider_oauth_start_key( &provider_type, pkce_verifier.as_deref(), key.encrypted_auth_config.as_deref(), + &principal.user_id, + principal.session_id.as_deref(), + principal.management_token_id.as_deref(), ) .await { @@ -99,12 +112,19 @@ pub(super) async fn handle_admin_provider_oauth_start_key( } }; - Ok(Json(build_provider_oauth_start_response( - template, - &nonce, - code_challenge.as_deref(), + let payload = + match build_provider_oauth_start_response(template, &nonce, code_challenge.as_deref()) { + Ok(payload) => payload, + Err(_) => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "OAuth 客户端配置不可用", + )); + } + }; + Ok(mark_sensitive_admin_response_no_store( + Json(payload).into_response(), )) - .into_response()) } pub(super) async fn handle_admin_provider_oauth_start_provider( @@ -153,6 +173,15 @@ pub(super) async fn handle_admin_provider_oauth_start_provider( "该 Provider 不支持 OAuth 授权", )); }; + let Some(principal) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + else { + return Ok(build_internal_control_error_response( + http::StatusCode::UNAUTHORIZED, + "管理员身份不可用", + )); + }; let pkce_verifier = template .use_pkce @@ -165,6 +194,9 @@ pub(super) async fn handle_admin_provider_oauth_start_provider( &provider_type, pkce_verifier.as_deref(), None, + &principal.user_id, + principal.session_id.as_deref(), + principal.management_token_id.as_deref(), ) .await { @@ -177,10 +209,17 @@ pub(super) async fn handle_admin_provider_oauth_start_provider( } }; - Ok(Json(build_provider_oauth_start_response( - template, - &nonce, - code_challenge.as_deref(), + let payload = + match build_provider_oauth_start_response(template, &nonce, code_challenge.as_deref()) { + Ok(payload) => payload, + Err(_) => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "OAuth 客户端配置不可用", + )); + } + }; + Ok(mark_sensitive_admin_response_no_store( + Json(payload).into_response(), )) - .into_response()) } 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 0cc4794f2..f446d4104 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 @@ -7,8 +7,17 @@ use base64::{ use serde_json::{json, Map, Value}; use std::collections::BTreeMap; +const MAX_UNVERIFIED_JWT_PART_BYTES: usize = 64 * 1024; + fn decode_base64_url_part(value: &str) -> Option> { - URL_SAFE_NO_PAD + if value.len() + > crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + MAX_UNVERIFIED_JWT_PART_BYTES, + ) + { + return None; + } + let bytes = URL_SAFE_NO_PAD .decode(value.as_bytes()) .or_else(|_| URL_SAFE.decode(value.as_bytes())) .or_else(|_| { @@ -19,7 +28,8 @@ fn decode_base64_url_part(value: &str) -> Option> { } URL_SAFE.decode(padded.as_bytes()) }) - .ok() + .ok()?; + (bytes.len() <= MAX_UNVERIFIED_JWT_PART_BYTES).then_some(bytes) } fn decode_unverified_jwt_json_part(part: &str) -> Option> { @@ -31,15 +41,26 @@ fn decode_unverified_jwt_json_part(part: &str) -> Option> { } pub(super) fn looks_like_access_token(token: &str) -> bool { - let parts = token.trim().split('.').collect::>(); - if parts.len() != 3 || parts.iter().any(|part| part.is_empty()) { + let mut parts = token.trim().split('.'); + let Some(header_part) = parts.next().filter(|part| !part.is_empty()) else { + return false; + }; + let Some(payload_part) = parts.next().filter(|part| !part.is_empty()) else { + return false; + }; + let Some(_signature_part) = parts.next().filter(|part| !part.is_empty()) else { + return false; + }; + // A JWT has exactly three dot-separated parts. Avoid collecting all + // attacker-controlled segments into a temporary Vec just to reject extras. + if parts.next().is_some() { return false; } - let Some(header) = decode_unverified_jwt_json_part(parts[0]) else { + let Some(header) = decode_unverified_jwt_json_part(header_part) else { return false; }; - let Some(payload) = decode_unverified_jwt_json_part(parts[1]) else { + let Some(payload) = decode_unverified_jwt_json_part(payload_part) else { return false; }; @@ -367,6 +388,13 @@ mod tests { assert_eq!(access_token.as_deref(), Some(token.as_str())); } + #[test] + fn rejects_jwt_with_many_extra_segments_without_collecting() { + let token = unsigned_jwt(json!({"exp": 2_000_000_000u64})); + let oversized = format!("{token}.{}", "x.".repeat(4096)); + assert!(!looks_like_access_token(&oversized)); + } + #[test] fn builds_codex_temporary_auth_config_from_access_token() { let token = unsigned_jwt(json!({ diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs index 40cd83e7b..e40c79462 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs @@ -211,12 +211,11 @@ pub(crate) async fn acquire_codex_oauth_account_locks( release_codex_oauth_account_locks(state, leases).await; return Err(CodexOAuthAccountLockError::Contended); } - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, lock_key = %lock_key, operation, - error = ?error, "gateway Codex OAuth account lock unavailable" ); release_codex_oauth_account_locks(state, leases).await; @@ -249,9 +248,8 @@ pub(crate) async fn release_provider_oauth_account_locks( lock_key = %lease.key, "gateway provider OAuth account lock was not owned during release" ), - Err(error) => tracing::warn!( + Err(_) => tracing::warn!( lock_key = %lease.key, - error = ?error, "gateway provider OAuth account lock release failed" ), } @@ -308,12 +306,11 @@ pub(crate) async fn acquire_claude_oauth_account_lock( { Ok(Some(lease)) => lease, Ok(None) => return Err(ClaudeOAuthAccountLockError::Contended), - Err(error) => { + Err(_) => { tracing::warn!( provider_id = %provider_id, lock_key = %lock_key, operation, - error = ?error, "gateway Claude OAuth account lock unavailable" ); return Err(ClaudeOAuthAccountLockError::Unavailable); diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/errors.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/errors.rs index a76588d6d..71695568d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/errors.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/errors.rs @@ -41,7 +41,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message( ) -> String { let mut message = None::; let mut error_code = None::; - let mut error_type = None::; if let Some(body_excerpt) = body_excerpt { if let Ok(value) = serde_json::from_str::(body_excerpt) { @@ -62,12 +61,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message( .map(str::trim) .filter(|value| !value.is_empty()) .map(|value| value.to_ascii_lowercase()); - error_type = error_object - .get("type") - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); } if message.is_none() { message = object @@ -86,14 +79,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message( .filter(|value| !value.is_empty()) .map(|value| value.to_ascii_lowercase()); } - if error_type.is_none() { - error_type = object - .get("type") - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); - } } } } @@ -108,7 +93,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message( .unwrap_or_default(); let lowered = message.to_ascii_lowercase(); let error_code = error_code.unwrap_or_default(); - let error_type = error_type.unwrap_or_default(); if error_code == "refresh_token_reused" || lowered.contains("already been used to generate a new access token") @@ -126,15 +110,9 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message( { return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string(); } - if error_type == "invalid_request_error" && !message.is_empty() { - return message; - } - if !message.is_empty() { - return message; - } status_code .map(|status_code| format!("HTTP {status_code}")) - .unwrap_or_else(|| "未知错误".to_string()) + .unwrap_or_else(|| "Token 刷新失败".to_string()) } pub(crate) fn merge_provider_oauth_refresh_failure_reason( @@ -182,6 +160,17 @@ mod tests { ); } + #[test] + fn refresh_error_does_not_reflect_unknown_upstream_text_or_credentials() { + let body = r#"{"error":{"message":"authorization=Bearer upstream-secret https://user:pass@example.test?q=secret","type":"invalid_request_error","code":"unexpected"}}"#; + + let normalized = normalize_provider_oauth_refresh_error_message(Some(502), Some(body)); + assert_eq!(normalized, "HTTP 502"); + for secret in ["upstream-secret", "user:pass", "q=secret"] { + assert!(!normalized.contains(secret), "leaked {secret}"); + } + } + #[test] fn refresh_failure_does_not_replace_account_level_block() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs index c27160e9b..37accf7ee 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs @@ -353,13 +353,20 @@ pub(crate) async fn create_provider_oauth_catalog_key( proxy: Option, expires_at_unix_secs: Option, ) -> Result, GatewayError> { - let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else { + let key_id = Uuid::new_v4().to_string(); + let Ok(encrypted_api_key) = + state + .app() + .seal_provider_catalog_key_api_key(provider_id, &key_id, access_token) + else { return Ok(None); }; let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone())) .map_err(|err| GatewayError::Internal(err.to_string()))?; - let Some(encrypted_auth_config) = - state.encrypt_catalog_secret_with_fallbacks(&auth_config_json) + let Ok(encrypted_auth_config) = + state + .app() + .seal_provider_catalog_key_auth_config(provider_id, &key_id, &auth_config_json) else { return Ok(None); }; @@ -369,7 +376,7 @@ pub(crate) async fn create_provider_oauth_catalog_key( .map(|duration| duration.as_secs()) .unwrap_or(0); let mut record = StoredProviderCatalogKey::new( - Uuid::new_v4().to_string(), + key_id, provider_id.to_string(), name.to_string(), "oauth".to_string(), @@ -422,14 +429,20 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key( proxy: Option, expires_at_unix_secs: Option, ) -> Result, GatewayError> { - let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else { + let Ok(encrypted_api_key) = state.app().seal_provider_catalog_key_api_key( + &existing_key.provider_id, + &existing_key.id, + access_token, + ) else { return Ok(None); }; let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone())) .map_err(|err| GatewayError::Internal(err.to_string()))?; - let Some(encrypted_auth_config) = - state.encrypt_catalog_secret_with_fallbacks(&auth_config_json) - else { + let Ok(encrypted_auth_config) = state.app().seal_provider_catalog_key_auth_config( + &existing_key.provider_id, + &existing_key.id, + &auth_config_json, + ) else { return Ok(None); }; let now_unix_secs = SystemTime::now() @@ -492,11 +505,10 @@ pub(super) async fn seed_provider_oauth_pool_score( .await { Ok(mut providers) => providers.pop(), - Err(err) => { + Err(_) => { tracing::debug!( provider_id = %provider_id, key_id = %key.id, - error = ?err, "gateway provider oauth provisioning: failed to read provider for pool score seed" ); return; @@ -524,11 +536,10 @@ pub(super) async fn seed_provider_oauth_pool_score( .await { Ok(mut scores) => scores.pop(), - Err(err) => { + Err(_) => { tracing::debug!( provider_id = %provider_id, key_id = %key.id, - error = ?err, "gateway provider oauth provisioning: failed to read existing pool score" ); return; @@ -541,11 +552,16 @@ pub(super) async fn seed_provider_oauth_pool_score( now_unix_secs, pool_config.score_rules, ); - if let Err(err) = state.app().data.upsert_pool_member_score(upsert).await { + if state + .app() + .data + .upsert_pool_member_score(upsert) + .await + .is_err() + { tracing::debug!( provider_id = %provider_id, key_id = %key.id, - error = ?err, "gateway provider oauth provisioning: failed to refresh pool score row" ); } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/antigravity.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/antigravity.rs index 3e60995f1..ee24a7161 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/antigravity.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/antigravity.rs @@ -5,17 +5,81 @@ use super::shared::{ quota_key_auto_removed, quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome, }; +use crate::handlers::admin::provider::shared::payloads::AdminImportProviderModelsRequest; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; use crate::GatewayError; -use aether_admin::provider::quota::parse_antigravity_usage_response; +use aether_admin::provider::quota::{ + parse_antigravity_quota_summary_response, parse_antigravity_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_antigravity_pool_quota_request; +use aether_provider_pool::{ + build_antigravity_pool_quota_request, build_antigravity_pool_quota_summary_request, +}; use serde_json::json; use std::collections::BTreeMap; use std::time::{SystemTime, UNIX_EPOCH}; +use tracing::warn; + +fn antigravity_discovered_model_ids(metadata_update: Option<&serde_json::Value>) -> Vec { + metadata_update + .and_then(|value| value.pointer("/antigravity/quota_by_model")) + .and_then(serde_json::Value::as_object) + .into_iter() + .flat_map(|models| models.keys()) + .map(String::as_str) + .filter(|model_id| aether_model_fetch::antigravity_model_id_is_routable(model_id)) + .map(ToOwned::to_owned) + .collect() +} + +async fn sync_antigravity_discovered_models( + state: &AdminAppState<'_>, + provider_id: &str, + metadata_update: Option<&serde_json::Value>, +) { + if !state.has_global_model_data_reader() || !state.has_global_model_data_writer() { + return; + } + let model_ids = antigravity_discovered_model_ids(metadata_update); + if model_ids.is_empty() { + return; + } + + let result = state + .build_admin_import_provider_models_payload( + provider_id, + AdminImportProviderModelsRequest { + model_ids, + tiered_pricing: None, + price_per_request: None, + }, + ) + .await; + match result { + Ok(payload) => { + let errors = payload + .get("errors") + .and_then(serde_json::Value::as_array) + .map(Vec::len) + .unwrap_or(0); + if errors > 0 { + warn!( + provider_id, + errors, "Antigravity discovered-model catalog sync completed with item errors" + ); + } + } + Err(error) => warn!( + provider_id, + error = %error, + "Antigravity discovered-model catalog sync failed" + ), + } +} async fn execute_antigravity_quota_plan( state: &AdminAppState<'_>, @@ -55,6 +119,79 @@ async fn execute_antigravity_quota_plan( execute_provider_quota_plan(state, transport, plan, "antigravity").await } +async fn fetch_antigravity_quota_summary_best_effort( + state: &AdminAppState<'_>, + transport: &AdminGatewayProviderTransportSnapshot, + authorization: (String, String), + project_id: &str, + identity_headers: BTreeMap, + proxy_override: Option<&ProxySnapshot>, +) -> Option { + let mut request_project_id = Some(project_id); + + loop { + let proxy = match proxy_override { + Some(proxy) => Some(proxy.clone()), + None => { + state + .resolve_transport_proxy_snapshot_with_tunnel_affinity(transport) + .await + } + }; + let timeouts = Some(resolve_provider_quota_execution_timeouts( + state.resolve_transport_execution_timeouts(transport), + proxy.as_ref(), + )); + let spec = build_antigravity_pool_quota_summary_request( + &transport.key.id, + &transport.endpoint.base_url, + authorization.clone(), + request_project_id, + identity_headers.clone(), + ); + let plan = build_provider_quota_execution_plan( + transport, + spec, + proxy, + state.resolve_transport_profile(transport), + timeouts, + ); + let outcome = match execute_provider_quota_plan(state, transport, plan, "antigravity").await + { + Ok(outcome) => outcome, + Err(error) => { + warn!(error = ?error, "Antigravity grouped quota request failed"); + return None; + } + }; + let result = match outcome { + ProviderQuotaExecutionOutcome::Response(result) => result, + ProviderQuotaExecutionOutcome::Failure(detail) => { + warn!(detail = %detail, "Antigravity grouped quota execution failed"); + return None; + } + }; + + if result.status_code == 200 { + return result + .body + .as_ref() + .and_then(|body| body.json_body.as_ref()) + .and_then(parse_antigravity_quota_summary_response); + } + if result.status_code == 403 && request_project_id.is_some() { + request_project_id = None; + continue; + } + + warn!( + status_code = result.status_code, + "Antigravity grouped quota request returned a non-success status" + ); + return None; + } +} + pub(crate) async fn refresh_antigravity_provider_quota_locally( state: &AdminAppState<'_>, provider: &StoredProviderCatalogProvider, @@ -130,21 +267,21 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally( let result = match execute_antigravity_quota_plan( state, &transport, - authorization, + authorization.clone(), &project_id, - identity_headers, + identity_headers.clone(), proxy_override.as_ref(), ) .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("fetchAvailableModels 请求执行失败: {detail}"), + "message": "fetchAvailableModels 请求执行失败", "status_code": 502, })); continue; @@ -168,9 +305,31 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally( .as_ref() .and_then(|body| body.json_body.as_ref()) { - metadata_update = parse_antigravity_usage_response(body_json, now_unix_secs) - .map(|metadata| json!({ "antigravity": metadata })); - if metadata_update.is_some() { + if let Some(mut metadata) = + parse_antigravity_usage_response(body_json, now_unix_secs) + { + if let Some(metadata) = metadata.as_object_mut() { + metadata.insert("project_id".to_string(), json!(project_id)); + } + if let Some(quota_groups) = fetch_antigravity_quota_summary_best_effort( + state, + &transport, + authorization, + &project_id, + identity_headers, + proxy_override.as_ref(), + ) + .await + { + if let Some(metadata) = metadata.as_object_mut() { + metadata.insert("quota_groups".to_string(), quota_groups); + metadata.insert( + "quota_groups_updated_at".to_string(), + json!(now_unix_secs), + ); + } + } + metadata_update = Some(json!({ "antigravity": metadata })); status = "success".to_string(); } else { status = "no_metadata".to_string(); @@ -181,21 +340,12 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally( message = Some("响应中未包含配额信息".to_string()); } } else { - let err_msg = extract_execution_error_message(&result); - message = Some(match err_msg.as_deref() { - Some(detail) if !detail.is_empty() => { - format!( - "fetchAvailableModels 返回状态码 {}: {}", - result.status_code, detail - ) - } - _ => format!("fetchAvailableModels 返回状态码 {}", result.status_code), - }); + message = Some(format!( + "fetchAvailableModels 返回状态码 {}", + result.status_code + )); if result.status_code == 403 { - let reason = err_msg - .clone() - .filter(|value| !value.trim().is_empty()) - .unwrap_or_else(|| "账户访问被禁止".to_string()); + let reason = "账户访问被禁止".to_string(); oauth_invalid_at_unix_secs = Some(now_unix_secs); oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}")); metadata_update = Some(json!({ @@ -230,6 +380,10 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally( continue; } + if status == "success" { + sync_antigravity_discovered_models(state, &provider.id, metadata_update.as_ref()).await; + } + if status == "success" { success_count += 1; } else { @@ -246,9 +400,11 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally( if let Some(metadata) = metadata_update .as_ref() .and_then(|value| value.get("antigravity")) - .cloned() { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("antigravity", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "antigravity", diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/chatgpt_web.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/chatgpt_web.rs index 868ffc20f..e689ad793 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/chatgpt_web.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/chatgpt_web.rs @@ -1,15 +1,17 @@ use super::shared::{ - build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message, + build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message_ref, oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state, quota_key_auto_removed, quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome, }; +use crate::execution_runtime::transport::decode_base64_body_with_limit; use crate::handlers::admin::provider::shared::payloads::{ OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX, }; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; use crate::GatewayError; use aether_admin::provider::quota::parse_chatgpt_web_conversation_init_response; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; use aether_contracts::{ ExecutionResult, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY, @@ -21,7 +23,6 @@ use aether_provider_pool::{ build_chatgpt_web_pool_quota_request, enrich_chatgpt_web_quota_metadata, normalize_chatgpt_web_image_quota_limit, }; -use base64::Engine as _; use serde_json::json; use std::time::{SystemTime, UNIX_EPOCH}; @@ -125,14 +126,21 @@ fn default_chatgpt_web_quota_transport_profile() -> ResolvedTransportProfile { } fn chatgpt_web_quota_error_detail(result: &ExecutionResult) -> Option { - extract_execution_error_message(result).or_else(|| { - let body = result.body.as_ref()?.body_bytes_b64.as_deref()?; - let decoded = base64::engine::general_purpose::STANDARD - .decode(body) - .ok()?; - let text = String::from_utf8_lossy(&decoded).trim().to_string(); - (!text.is_empty()).then_some(text) - }) + extract_execution_error_message_ref(result) + .map(bound_quota_error_detail) + .or_else(|| { + let body = result.body.as_ref()?.body_bytes_b64.as_deref()?; + let decoded = decode_base64_body_with_limit(body, crate::MAX_ERROR_BODY_BYTES).ok()?; + let text = String::from_utf8_lossy(&decoded); + let text = text.trim(); + (!text.is_empty()).then(|| bound_quota_error_detail(text)) + }) +} + +fn bound_quota_error_detail(value: &str) -> String { + let value = value.trim(); + let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES)); + value[..end].to_string() } fn chatgpt_web_is_structured_account_block(message: &str) -> bool { @@ -162,11 +170,8 @@ fn chatgpt_web_is_structured_account_block(message: &str) -> bool { } fn chatgpt_web_quota_403_refresh_failed_reason(message: Option<&str>) -> String { - let detail = message - .map(str::trim) - .filter(|value| !value.is_empty()) - .filter(|value| !value.contains('<')) - .unwrap_or("ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制"); + let _ = message; + let detail = "ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制"; format!("{OAUTH_REFRESH_FAILED_PREFIX}{detail}") } @@ -175,14 +180,10 @@ fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<& if status_code == 403 && !chatgpt_web_is_structured_account_block(message) { return chatgpt_web_quota_403_refresh_failed_reason(upstream_message); } - let detail = if message.is_empty() { - match status_code { - 401 => "ChatGPT Web Token 无效或已过期", - 403 => "ChatGPT Web 账户访问受限", - _ => "ChatGPT Web 请求失败", - } - } else { - message + let detail = match status_code { + 401 => "ChatGPT Web Token 无效或已过期", + 403 => "ChatGPT Web 账户访问受限", + _ => "ChatGPT Web 请求失败", }; match status_code { 401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"), @@ -263,13 +264,13 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("conversation/init 请求执行失败: {detail}"), + "message": "conversation/init 请求执行失败", "status_code": 502, })); continue; @@ -304,6 +305,8 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally( &mut metadata, key.upstream_metadata.as_ref(), ); + metadata = + admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(&metadata)); metadata_update = Some(json!({ "chatgpt_web": metadata })); (oauth_invalid_at_unix_secs, oauth_invalid_reason) = quota_refresh_success_invalid_state(&key); @@ -328,8 +331,7 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally( }; let display_detail = invalid_reason .as_deref() - .map(chatgpt_web_quota_result_message) - .or_else(|| err_msg.clone()); + .map(chatgpt_web_quota_result_message); message = Some(match display_detail.as_deref() { Some(detail) if !detail.is_empty() => { format!( @@ -395,9 +397,11 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally( if let Some(metadata) = metadata_update .as_ref() .and_then(|value| value.get("chatgpt_web")) - .cloned() { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "chatgpt_web", @@ -496,4 +500,65 @@ mod tests { assert!(reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX)); } + + #[test] + fn quota_error_detail_is_bounded_before_copying_json_message() { + let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES)); + let result = ExecutionResult { + request_id: "chatgpt-web-quota:oversized-json".to_string(), + candidate_id: None, + status_code: 500, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({"error": {"message": message}})), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }; + + let detail = chatgpt_web_quota_error_detail(&result).expect("JSON error detail"); + + assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES); + assert!(detail.bytes().all(|byte| byte == b'x')); + } + + #[test] + fn quota_error_detail_rejects_oversized_base64_before_decode() { + let encoded_limit = + crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + crate::MAX_ERROR_BODY_BYTES, + ); + let result = ExecutionResult { + request_id: "chatgpt-web-quota:oversized-base64".to_string(), + candidate_id: None, + status_code: 500, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: None, + body_bytes_b64: Some("A".repeat(encoded_limit + 1)), + }), + telemetry: None, + error: None, + }; + + assert_eq!(chatgpt_web_quota_error_detail(&result), None); + } + + #[test] + fn quota_invalid_reason_does_not_persist_upstream_credentials() { + let reason = chatgpt_web_quota_invalid_reason( + 401, + Some("authorization=Bearer upstream-secret https://user:pass@example.test?q=secret"), + ); + + assert_eq!( + reason, + format!("{OAUTH_EXPIRED_PREFIX}ChatGPT Web Token 无效或已过期") + ); + assert!(!reason.contains("upstream-secret")); + assert!(!reason.contains("user:pass")); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/codex/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/codex/mod.rs index b5c490fea..1b59cc306 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/codex/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/codex/mod.rs @@ -30,6 +30,7 @@ use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTranspo use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::state::ProviderTransportCredentialFence; use crate::GatewayError; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; use aether_contracts::ProxySnapshot; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyOAuthCredentialCasDelete, @@ -282,12 +283,58 @@ fn truncate_codex_reset_credit_detail_error(message: impl Into) -> Strin let message = message.into(); let mut sanitized = message.replace('\n', " "); if sanitized.len() > 240 { - sanitized.truncate(240); + let mut truncate_at = 240; + while !sanitized.is_char_boundary(truncate_at) { + truncate_at -= 1; + } + sanitized.truncate(truncate_at); sanitized.push('…'); } sanitized } +fn safe_codex_quota_refresh_error(error: &GatewayError) -> &'static str { + if matches!( + error, + GatewayError::LocalExecutionPlanningTimeout { .. } | GatewayError::AdmissionTimeout { .. } + ) { + return "Quota refresh timed out"; + } + + let message = match error { + GatewayError::UpstreamUnavailable { message, .. } + | GatewayError::ControlUnavailable { message, .. } + | GatewayError::Client { message, .. } + | GatewayError::Internal(message) => message.as_str(), + GatewayError::PlanUsageLimited(_) + | GatewayError::LastActiveAdminUpdateDenied + | GatewayError::LastActiveAdminDeleteDenied => return "Quota refresh failed", + GatewayError::LocalExecutionPlanningTimeout { .. } + | GatewayError::AdmissionTimeout { .. } => unreachable!(), + }; + let lower = message.to_ascii_lowercase(); + if lower.contains("timeout") || lower.contains("timed out") { + "Quota refresh timed out" + } else if [ + "connection", + "connect", + "dns", + "network", + "proxy", + "socket", + "tls", + "certificate", + "transport", + ] + .iter() + .any(|marker| lower.contains(marker)) + { + "Quota refresh connection failed" + } else { + "Quota refresh failed" + } +} + fn merge_codex_reset_credit_detail_metadata( codex_metadata: &mut Map, detail_metadata: &Value, @@ -355,26 +402,21 @@ async fn enrich_codex_reset_credit_details( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { mark_codex_reset_credit_detail_failed( codex_metadata, now_unix_secs, - format!("reset credit detail 请求执行失败: {detail}"), + "reset credit detail 请求执行失败".to_string(), ); return Ok(()); } }; if result.status_code != 200 { - let detail = extract_execution_error_message(&result) - .unwrap_or_else(|| format!("HTTP {}", result.status_code)); mark_codex_reset_credit_detail_failed( codex_metadata, now_unix_secs, - format!( - "reset credit detail 返回状态码 {}: {detail}", - result.status_code - ), + format!("reset credit detail 返回状态码 {}", result.status_code), ); return Ok(()); } @@ -506,8 +548,7 @@ async fn finish_codex_reset_replay( } Err(err) => { refresh_status = "failed".to_string(); - refresh_error = - Some(truncate_codex_reset_credit_detail_error(err.into_message())); + refresh_error = Some(safe_codex_quota_refresh_error(&err).to_string()); } } } @@ -529,7 +570,10 @@ async fn finish_codex_reset_replay( payload.insert("refresh_error".to_string(), json!(refresh_error)); } if let Some(metadata) = metadata { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("codex", Some(&metadata)), + ); } if let Some(quota_snapshot) = quota_snapshot { payload.insert("quota_snapshot".to_string(), quota_snapshot); @@ -702,14 +746,14 @@ pub(crate) async fn consume_codex_reset_credit_locally( let result = match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { return Ok(( StatusCode::BAD_GATEWAY, json!({ "key_id": key.id, "status": "error", "outcome": "error", - "message": format!("reset credit consume 请求执行失败: {detail}"), + "message": "reset credit consume 请求执行失败", }), )); } @@ -726,8 +770,6 @@ pub(crate) async fn consume_codex_reset_credit_locally( "reset" | "already_redeemed" | "nothing_to_reset" | "no_credit" ); if !known_terminal_outcome { - let detail = extract_execution_error_message(&result) - .unwrap_or_else(|| format!("HTTP {}", result.status_code)); return Ok(( StatusCode::BAD_GATEWAY, json!({ @@ -735,7 +777,7 @@ pub(crate) async fn consume_codex_reset_credit_locally( "status": "error", "outcome": "error", "idempotency_key": idempotency_key, - "message": format!("reset credit consume outcome is ambiguous: {detail}"), + "message": "reset credit consume outcome is ambiguous", "status_code": result.status_code, }), )); @@ -802,7 +844,7 @@ pub(crate) async fn consume_codex_reset_credit_locally( } Err(err) => ( "failed".to_string(), - Some(truncate_codex_reset_credit_detail_error(err.into_message())), + Some(safe_codex_quota_refresh_error(&err).to_string()), None, None, ), @@ -825,7 +867,7 @@ pub(crate) async fn consume_codex_reset_credit_locally( } Err(err) => ( "failed".to_string(), - Some(truncate_codex_reset_credit_detail_error(err.into_message())), + Some(safe_codex_quota_refresh_error(&err).to_string()), None, None, ), @@ -845,7 +887,10 @@ pub(crate) async fn consume_codex_reset_credit_locally( payload.insert("refresh_error".to_string(), json!(refresh_error)); } if let Some(metadata) = metadata { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("codex", Some(&metadata)), + ); } if let Some(quota_snapshot) = quota_snapshot { payload.insert("quota_snapshot".to_string(), quota_snapshot); @@ -1003,13 +1048,13 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("wham/usage 请求执行失败: {detail}"), + "message": "wham/usage 请求执行失败", "status_code": 502, })); continue; @@ -1080,15 +1125,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( } } else { let err_msg = extract_execution_error_message(&result); - message = Some(match err_msg.as_deref() { - Some(detail) if !detail.is_empty() => { - format!( - "wham/usage API 返回状态码 {}: {}", - result.status_code, detail - ) - } - _ => format!("wham/usage API 返回状态码 {}", result.status_code), - }); + message = Some(format!("wham/usage API 返回状态码 {}", result.status_code)); match result.status_code { 401 => { @@ -1114,12 +1151,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( codex_meta.insert("updated_at".to_string(), json!(now_unix_secs)); codex_meta.insert("account_disabled".to_string(), json!(true)); codex_meta.insert("reason".to_string(), json!("deactivated_workspace")); - codex_meta.insert( - "message".to_string(), - json!(err_msg - .clone() - .unwrap_or_else(|| "deactivated_workspace".to_string())), - ); + codex_meta.insert("message".to_string(), json!("deactivated_workspace")); let plan_type = transport .key .decrypted_auth_config @@ -1206,7 +1238,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( account_reset_fence_id, coverage: quota_window_coverage, }, - Some(&expected_credential.credential), + &expected_credential.credential, ) .await? } else { @@ -1214,9 +1246,6 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( state, &key.id, metadata_update.as_ref(), - oauth_invalid_at_unix_secs, - oauth_invalid_reason.clone(), - None, aether_admin::provider::quota::CodexQuotaMergeContext { observed_at_unix_secs: now_unix_secs, request_started_at_unix_ms: Some(quota_request_started_at_unix_ms), @@ -1366,9 +1395,11 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence( if let Some(metadata_update) = metadata_update .as_ref() .and_then(|value| value.get("codex")) - .cloned() { - payload.insert("metadata".to_string(), metadata_update); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("codex", Some(metadata_update)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "codex", @@ -1470,6 +1501,53 @@ mod tests { ); } + #[test] + fn codex_reset_credit_detail_truncation_preserves_utf8_boundaries() { + let detail = format!("{}密钥", "a".repeat(239)); + + let truncated = truncate_codex_reset_credit_detail_error(detail); + + assert_eq!(truncated, format!("{}…", "a".repeat(239))); + } + + #[test] + fn codex_quota_refresh_error_projection_discards_internal_details() { + let secrets = [ + "https://user:password@internal.test/v1/quota?q=secret", + "Authorization: Bearer upstream-secret", + "user:password", + "upstream-secret", + ]; + let connection_error = GatewayError::Internal(format!( + "connection failed for {}; {}", + secrets[0], secrets[1] + )); + let generic_error = GatewayError::Internal(format!( + "repository failure while processing {}; {}", + secrets[0], secrets[1] + )); + let timeout_error = GatewayError::LocalExecutionPlanningTimeout { + trace_id: secrets[1].to_string(), + phase: "quota_refresh", + timeout_ms: 5_000, + }; + + let projected = [ + safe_codex_quota_refresh_error(&connection_error), + safe_codex_quota_refresh_error(&generic_error), + safe_codex_quota_refresh_error(&timeout_error), + ]; + + assert_eq!(projected[0], "Quota refresh connection failed"); + assert_eq!(projected[1], "Quota refresh failed"); + assert_eq!(projected[2], "Quota refresh timed out"); + for safe_error in projected { + for secret in secrets { + assert!(!safe_error.contains(secret)); + } + } + } + #[test] fn codex_quota_coverage_only_replaces_observed_window_families() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/gemini_cli.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/gemini_cli.rs index 184c44ec9..536fb7e7b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/gemini_cli.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/gemini_cli.rs @@ -8,6 +8,7 @@ use super::shared::{ use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; use crate::GatewayError; use aether_admin::provider::quota::parse_gemini_cli_retrieve_user_quota_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, @@ -137,13 +138,13 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("retrieveUserQuota 请求执行失败: {detail}"), + "message": "retrieveUserQuota 请求执行失败", "status_code": 502, })); continue; @@ -181,21 +182,12 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally( message = Some("响应中未包含配额信息".to_string()); } } else { - let err_msg = extract_execution_error_message(&result); - message = Some(match err_msg.as_deref() { - Some(detail) if !detail.is_empty() => { - format!( - "retrieveUserQuota 返回状态码 {}: {}", - result.status_code, detail - ) - } - _ => format!("retrieveUserQuota 返回状态码 {}", result.status_code), - }); + message = Some(format!( + "retrieveUserQuota 返回状态码 {}", + result.status_code + )); if result.status_code == 403 { - let reason = err_msg - .clone() - .filter(|value| !value.trim().is_empty()) - .unwrap_or_else(|| "账户访问被禁止".to_string()); + let reason = "账户访问被禁止".to_string(); oauth_invalid_at_unix_secs = Some(now_unix_secs); oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}")); metadata_update = Some(json!({ @@ -246,9 +238,11 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally( if let Some(metadata) = metadata_update .as_ref() .and_then(|value| value.get("gemini_cli")) - .cloned() { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("gemini_cli", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "gemini_cli", diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/grok.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/grok.rs index 2a9b707c3..51c0e2de2 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/grok.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/grok.rs @@ -1,13 +1,15 @@ use super::shared::{ - build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message, + build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message_ref, persist_provider_quota_refresh_state, quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome, }; +use crate::execution_runtime::transport::decode_base64_body_with_limit; use crate::handlers::admin::provider::shared::payloads::{ OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX, }; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; use crate::GatewayError; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; use aether_contracts::{ ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody, ResolvedTransportProfile, }; @@ -18,7 +20,6 @@ use aether_provider_pool::{ grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier, }; use aether_provider_transport::grok_browser_profile_metadata_from_resolved_transport_profile; -use base64::Engine as _; use serde_json::json; use std::collections::BTreeMap; use std::time::{SystemTime, UNIX_EPOCH}; @@ -288,14 +289,21 @@ async fn execute_grok_quota_plan( } fn grok_quota_error_detail(result: &ExecutionResult) -> Option { - extract_execution_error_message(result).or_else(|| { - let body = result.body.as_ref()?.body_bytes_b64.as_deref()?; - let decoded = base64::engine::general_purpose::STANDARD - .decode(body) - .ok()?; - let text = String::from_utf8_lossy(&decoded).trim().to_string(); - (!text.is_empty()).then_some(text) - }) + extract_execution_error_message_ref(result) + .map(bound_quota_error_detail) + .or_else(|| { + let body = result.body.as_ref()?.body_bytes_b64.as_deref()?; + let decoded = decode_base64_body_with_limit(body, crate::MAX_ERROR_BODY_BYTES).ok()?; + let text = String::from_utf8_lossy(&decoded); + let text = text.trim(); + (!text.is_empty()).then(|| bound_quota_error_detail(text)) + }) +} + +fn bound_quota_error_detail(value: &str) -> String { + let value = value.trim(); + let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES)); + value[..end].to_string() } fn grok_is_cloudflare_challenge(message: &str) -> bool { @@ -313,14 +321,10 @@ fn grok_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) - "{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败,请重新从同一浏览器复制最新 Cookie 和 User-Agent,或配置可通过 Cloudflare 的代理运行时" ); } - let detail = if message.is_empty() { - match status_code { - 401 => "Grok Token 无效或已过期", - 403 => "Grok 账户访问受限", - _ => "Grok 请求失败", - } - } else { - message + let detail = match status_code { + 401 => "Grok Token 无效或已过期", + 403 => "Grok 账户访问受限", + _ => "Grok 请求失败", }; match status_code { 401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"), @@ -406,8 +410,8 @@ pub(crate) async fn refresh_grok_provider_quota_locally( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { - last_error_message = Some(format!("rate-limits 请求执行失败: {detail}")); + ProviderQuotaExecutionOutcome::Failure(_) => { + last_error_message = Some("rate-limits 请求执行失败".to_string()); continue; } }; @@ -467,10 +471,8 @@ pub(crate) async fn refresh_grok_provider_quota_locally( )); last_error_message = invalid_reason.as_deref().map(grok_quota_result_message); } else { - let error_detail = - grok_quota_error_detail(&result).unwrap_or_else(|| "Grok 请求失败".to_string()); last_error_message = Some(format!( - "Grok rate-limits 请求失败({}): {error_detail}", + "Grok rate-limits 请求失败 ({})", result.status_code )); } @@ -544,8 +546,11 @@ pub(crate) async fn refresh_grok_provider_quota_locally( "status".to_string(), json!(if refreshed { "success" } else { "error" }), ); - if let Some(metadata) = metadata_update.get("quota_by_model").cloned() { - payload.insert("metadata".to_string(), metadata); + if let Some(metadata) = metadata_update.get("quota_by_model") { + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("grok", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "grok", @@ -813,6 +818,52 @@ mod tests { assert!(reason.contains("Cloudflare")); } + #[test] + fn quota_error_detail_is_bounded_before_copying_json_message() { + let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES)); + let result = ExecutionResult { + request_id: "grok-quota:oversized-json".to_string(), + candidate_id: None, + status_code: 500, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({"error": {"message": message}})), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }; + + let detail = grok_quota_error_detail(&result).expect("JSON error detail"); + + assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES); + assert!(detail.bytes().all(|byte| byte == b'x')); + } + + #[test] + fn quota_error_detail_rejects_oversized_base64_before_decode() { + let encoded_limit = + crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + crate::MAX_ERROR_BODY_BYTES, + ); + let result = ExecutionResult { + request_id: "grok-quota:oversized-base64".to_string(), + candidate_id: None, + status_code: 500, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: None, + body_bytes_b64: Some("A".repeat(encoded_limit + 1)), + }), + telemetry: None, + error: None, + }; + + assert_eq!(grok_quota_error_detail(&result), None); + } + #[test] fn quota_result_message_removes_status_prefix() { let reason = format!("{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败"); @@ -822,4 +873,16 @@ mod tests { "Grok Cloudflare 验证失败" ); } + + #[test] + fn quota_invalid_reason_does_not_persist_upstream_credentials() { + let reason = grok_quota_invalid_reason( + 401, + Some("authorization=Bearer upstream-secret https://user:pass@example.test?q=secret"), + ); + + assert_eq!(reason, "[OAUTH_EXPIRED] Grok Token 无效或已过期"); + assert!(!reason.contains("upstream-secret")); + assert!(!reason.contains("user:pass")); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/kiro/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/kiro/mod.rs index 673246064..79cf31dca 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/kiro/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/kiro/mod.rs @@ -5,18 +5,17 @@ use self::parse::parse_kiro_usage_response; use self::plan::execute_kiro_quota_plan; use super::shared::{ build_quota_snapshot_payload, extract_execution_error_message, - oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state, + oauth_refresh_auto_removed_result, persist_credential_fenced_provider_quota_refresh_state, persist_quota_oauth_refresh_failure_state, provider_auto_remove_quota_exhausted_keys, quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome, }; use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError}; use crate::GatewayError; +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_transport::kiro::build_kiro_request_auth_from_config; -use aether_provider_transport::{CachedOAuthEntry, LocalResolvedOAuthRequestAuth}; use serde_json::json; use std::time::{SystemTime, UNIX_EPOCH}; @@ -54,20 +53,6 @@ fn kiro_quota_error_is_account_banned(detail: Option<&str>) -> bool { .any(|keyword| normalized.contains(keyword)) } -fn kiro_auth_from_refreshed_entry( - entry: &CachedOAuthEntry, -) -> Option { - if !entry.provider_type.trim().eq_ignore_ascii_case("kiro") { - return None; - } - let auth_config = entry - .metadata - .as_ref() - .and_then(aether_provider_transport::kiro::KiroAuthConfig::from_json_value)?; - let auth = build_kiro_request_auth_from_config(auth_config, None)?; - Some(LocalResolvedOAuthRequestAuth::Kiro(auth)) -} - fn kiro_quota_refresh_failure_status(err: &AdminLocalOAuthRefreshError) -> Option { match err { AdminLocalOAuthRefreshError::HttpStatus { status_code, .. } => Some(*status_code), @@ -77,12 +62,13 @@ fn kiro_quota_refresh_failure_status(err: &AdminLocalOAuthRefreshError) -> Optio fn kiro_quota_refresh_failure_message(err: &AdminLocalOAuthRefreshError) -> String { match err { - AdminLocalOAuthRefreshError::HttpStatus { - status_code, - body_excerpt, - .. - } => format!("Kiro Token 刷新失败 ({status_code}): {body_excerpt}"), - _ => format!("Kiro Token 刷新失败: {err}"), + AdminLocalOAuthRefreshError::HttpStatus { status_code, .. } => { + format!("Kiro Token 刷新失败: HTTP {status_code}") + } + AdminLocalOAuthRefreshError::InvalidResponse { .. } => { + "Kiro Token 刷新失败: 无效响应".to_string() + } + _ => "Kiro Token 刷新失败: 网络错误".to_string(), } } @@ -118,36 +104,8 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( } }; - let auth = match state.force_local_oauth_refresh_entry(&transport).await { - Ok(Some(entry)) => match kiro_auth_from_refreshed_entry(&entry) { - Some(LocalResolvedOAuthRequestAuth::Kiro(auth)) => auth, - _ => { - failed_count += 1; - results.push(json!({ - "key_id": key.id, - "key_name": key.name, - "status": "error", - "message": "Kiro Token 刷新成功但认证信息解析失败", - })); - continue; - } - }, - Ok(None) => match state - .resolve_local_oauth_kiro_request_auth(&transport) - .await? - { - Some(auth) => auth, - None => { - failed_count += 1; - results.push(json!({ - "key_id": key.id, - "key_name": key.name, - "status": "error", - "message": "缺少 Kiro 认证配置 (auth_config)", - })); - continue; - } - }, + match state.force_local_oauth_refresh_entry(&transport).await { + Ok(_) => {} Err(err) => { if persist_quota_oauth_refresh_failure_state(state, &transport, &err).await? || super::shared::quota_key_auto_removed(state, &key.id).await? @@ -171,19 +129,60 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( results.push(serde_json::Value::Object(payload)); continue; } + } + + let Some(transport) = state + .read_provider_transport_snapshot_uncached(&provider.id, &endpoint.id, &key.id) + .await? + else { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Kiro Token 刷新后凭据快照不可用", + })); + continue; + }; + let Some(auth) = state + .resolve_local_oauth_kiro_request_auth(&transport) + .await? + else { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "缺少 Kiro 认证配置 (auth_config)", + })); + continue; + }; + let Some(credential_fence) = state + .app() + .capture_provider_transport_credential_fence(&transport) + .await? + else { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Kiro 凭据在请求前已变化", + })); + continue; }; let result = match execute_kiro_quota_plan(state, &transport, &auth, proxy_override.as_ref()).await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("getUsageLimits 请求执行失败: {detail}"), + "message": "getUsageLimits 请求执行失败", "status_code": 502, })); continue; @@ -230,11 +229,12 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( .or_insert_with(|| json!("kiro")); let auth_config_json = serde_json::Value::Object(auth_config_object).to_string(); - if let Some(auth_config_json) = - state.encrypt_catalog_secret_with_fallbacks(auth_config_json.as_str()) - { - encrypted_auth_config = Some(auth_config_json); - } + encrypted_auth_config = + Some(state.app().seal_provider_catalog_key_auth_config( + &transport.provider.id, + &transport.key.id, + auth_config_json.as_str(), + )?); status = "success".to_string(); } else { status = "no_metadata".to_string(); @@ -246,25 +246,14 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( } } else { let err_msg = extract_execution_error_message(&result); - message = Some(match err_msg.as_deref() { - Some(detail) if !detail.is_empty() => { - format!( - "getUsageLimits 返回状态码 {}: {}", - result.status_code, detail - ) - } - _ => format!("getUsageLimits 返回状态码 {}", result.status_code), - }); + message = Some(format!("getUsageLimits 返回状态码 {}", result.status_code)); match result.status_code { 401 => { oauth_invalid_at_unix_secs = Some(now_unix_secs); oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string()); } 403 | 423 => { - let reason = err_msg - .clone() - .filter(|value| !value.trim().is_empty()) - .unwrap_or_else(|| format!("HTTP {}", result.status_code)); + let reason = format!("HTTP {}", result.status_code); if kiro_quota_error_is_token_invalid(err_msg.as_deref()) { oauth_invalid_at_unix_secs = Some(now_unix_secs); oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string()); @@ -286,13 +275,14 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( } } - if !persist_provider_quota_refresh_state( + if !persist_credential_fenced_provider_quota_refresh_state( state, &key.id, metadata_update.as_ref(), oauth_invalid_at_unix_secs, oauth_invalid_reason, encrypted_auth_config, + &credential_fence, ) .await? { @@ -336,12 +326,11 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( 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("kiro")) - .cloned() - { - payload.insert("metadata".to_string(), metadata); + if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("kiro")) { + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("kiro", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "kiro", @@ -369,7 +358,11 @@ pub(crate) async fn refresh_kiro_provider_quota_locally( #[cfg(test)] mod tests { - use super::{kiro_quota_error_is_account_banned, kiro_quota_error_is_token_invalid}; + use super::{ + kiro_quota_error_is_account_banned, kiro_quota_error_is_token_invalid, + kiro_quota_refresh_failure_message, + }; + use crate::handlers::admin::request::AdminLocalOAuthRefreshError; #[test] fn bearer_token_invalid_is_not_classified_as_banned() { @@ -394,4 +387,20 @@ mod tests { assert!(!kiro_quota_error_is_token_invalid(detail)); assert!(kiro_quota_error_is_account_banned(detail)); } + + #[test] + fn quota_refresh_failure_does_not_reflect_upstream_body() { + let message = + kiro_quota_refresh_failure_message(&AdminLocalOAuthRefreshError::HttpStatus { + provider_type: "kiro", + status_code: 502, + body_excerpt: + "authorization=Bearer upstream-secret https://user:pass@example.test?q=secret" + .to_string(), + }); + + assert_eq!(message, "Kiro Token 刷新失败: HTTP 502"); + assert!(!message.contains("upstream-secret")); + assert!(!message.contains("user:pass")); + } } 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 6c5e55937..7f1129faf 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 @@ -7,11 +7,13 @@ use crate::handlers::admin::request::{ use crate::handlers::shared::{ sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot, }; +use crate::state::ProviderTransportCredentialFence; use crate::GatewayError; use aether_admin::provider::quota as admin_provider_quota_pure; +use aether_admin::provider::redaction::admin_provider_upstream_metadata_safe_json; use aether_contracts::{ ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody, - ResolvedTransportProfile, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, + ResolvedTransportProfile, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, }; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, @@ -240,7 +242,16 @@ fn merge_upstream_metadata( .unwrap_or_default(); if let Some(update_object) = updates.as_object() { for (key, value) in update_object { - merged.insert(key.clone(), value.clone()); + let mut next = value.clone(); + if let (Some(current_namespace), Some(next_namespace)) = ( + merged.get(key).and_then(serde_json::Value::as_object), + next.as_object_mut(), + ) { + let mut combined = current_namespace.clone(); + combined.extend(next_namespace.clone()); + next = serde_json::Value::Object(combined); + } + merged.insert(key.clone(), next); } } serde_json::Value::Object(merged) @@ -250,6 +261,40 @@ pub(super) fn extract_execution_error_message(result: &ExecutionResult) -> Optio admin_provider_quota_pure::extract_execution_error_message(result) } +pub(super) fn extract_execution_error_message_ref(result: &ExecutionResult) -> Option<&str> { + if let Some(body_json) = result + .body + .as_ref() + .and_then(|body| body.json_body.as_ref()) + .and_then(serde_json::Value::as_object) + { + if let Some(message) = body_json + .get("error") + .and_then(serde_json::Value::as_object) + .and_then(|error| error.get("message")) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|message| !message.is_empty()) + { + return Some(message); + } + if let Some(message) = body_json + .get("message") + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|message| !message.is_empty()) + { + return Some(message); + } + } + + result + .error + .as_ref() + .map(|error| error.message.trim()) + .filter(|message| !message.is_empty()) +} + fn extract_execution_error_detail(result: &ExecutionResult) -> Option { admin_provider_quota_pure::extract_execution_error_detail(result) } @@ -561,6 +606,50 @@ pub(crate) async fn reserve_codex_account_reset( Ok(None) } +fn record_locally_consumed_codex_reset_credit( + codex: &mut serde_json::Map, + observed_at_unix_secs: u64, +) { + let Some(reset_credits) = codex + .get_mut("reset_credits") + .and_then(serde_json::Value::as_object_mut) + else { + return; + }; + let Some(available_count) = reset_credits + .get("available_count") + .and_then(admin_provider_quota_pure::coerce_json_u64) + else { + return; + }; + + reset_credits.insert( + "available_count".to_string(), + serde_json::json!(available_count.saturating_sub(1)), + ); + reset_credits.insert( + "updated_at".to_string(), + serde_json::json!(observed_at_unix_secs), + ); + reset_credits.insert( + "detail_source".to_string(), + serde_json::json!("local_consume"), + ); + reset_credits.insert( + "detail_status".to_string(), + serde_json::json!("pending_refresh"), + ); + reset_credits.remove("detail_error"); + if let Some(credits) = reset_credits + .get_mut("credits") + .and_then(serde_json::Value::as_array_mut) + { + if !credits.is_empty() { + credits.remove(0); + } + } +} + pub(crate) async fn complete_codex_account_reset( state: &AdminAppState<'_>, key_id: &str, @@ -626,6 +715,9 @@ pub(crate) async fn complete_codex_account_reset( generation: reservation.generation, outcome: outcome.to_string(), }; + if outcome == "reset" { + record_locally_consumed_codex_reset_credit(&mut codex, fence_unix_ms / 1_000); + } codex_reset_write_bounded_history(&mut codex, &terminal); if codex_reset_reservation_from_object(&codex).as_ref() == Some(reservation) { codex.remove(admin_provider_quota_pure::CODEX_QUOTA_ACCOUNT_RESET_RESERVATION_KEY); @@ -696,6 +788,12 @@ pub(crate) async fn persist_codex_account_reset_fence( fence_id: &str, idempotency_key: &str, ) -> Result, GatewayError> { + if expected_encrypted_auth_config.is_some() != expected_credential.is_some() { + return Err(GatewayError::Internal( + "Codex reset credential fence must include auth_config and credential identity" + .to_string(), + )); + } let fence_id = fence_id.trim(); let idempotency_key = idempotency_key.trim(); if fence_unix_ms == 0 || fence_id.is_empty() || idempotency_key.is_empty() { @@ -893,14 +991,8 @@ pub(super) fn build_provider_quota_execution_plan( client_api_format, provider_api_format, model_name, - accept_invalid_certs, } = spec; - if accept_invalid_certs { - headers.insert( - EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER.to_string(), - "true".to_string(), - ); - } + force_provider_quota_redirects_disabled(&mut headers); let body = json_body .map(RequestBody::from_json) .unwrap_or(RequestBody { @@ -931,6 +1023,16 @@ pub(super) fn build_provider_quota_execution_plan( } } +fn force_provider_quota_redirects_disabled( + headers: &mut std::collections::BTreeMap, +) { + headers.retain(|name, _| !name.eq_ignore_ascii_case(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)); + headers.insert( + EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), + "false".to_string(), + ); +} + fn codex_reset_refresh_is_superseded( current: Option<&serde_json::Value>, context: admin_provider_quota_pure::CodexQuotaMergeContext<'_>, @@ -968,6 +1070,29 @@ pub(crate) async fn persist_provider_quota_refresh_state( oauth_invalid_at_unix_secs, oauth_invalid_reason, encrypted_auth_config, + None, + std::future::ready(()), + ) + .await +} + +pub(crate) async fn persist_credential_fenced_provider_quota_refresh_state( + state: &AdminAppState<'_>, + key_id: &str, + metadata_update: Option<&serde_json::Value>, + oauth_invalid_at_unix_secs: Option, + oauth_invalid_reason: Option, + encrypted_auth_config: Option, + expected_credential_fence: &ProviderTransportCredentialFence, +) -> Result { + persist_provider_quota_refresh_state_after_read( + state, + key_id, + metadata_update, + oauth_invalid_at_unix_secs, + oauth_invalid_reason, + encrypted_auth_config, + Some(expected_credential_fence), std::future::ready(()), ) .await @@ -977,9 +1102,6 @@ pub(crate) async fn persist_codex_provider_quota_refresh_state( state: &AdminAppState<'_>, key_id: &str, metadata_update: Option<&serde_json::Value>, - oauth_invalid_at_unix_secs: Option, - oauth_invalid_reason: Option, - encrypted_auth_config: Option, merge_context: admin_provider_quota_pure::CodexQuotaMergeContext<'_>, ) -> Result { let Some(incoming_codex) = metadata_update.and_then(|value| value.get("codex")) else { @@ -987,9 +1109,9 @@ pub(crate) async fn persist_codex_provider_quota_refresh_state( state, key_id, metadata_update, - oauth_invalid_at_unix_secs, - oauth_invalid_reason, - encrypted_auth_config, + None, + None, + None, ) .await; }; @@ -1030,86 +1152,32 @@ pub(crate) async fn persist_codex_provider_quota_refresh_state( latest_key.upstream_metadata.as_ref(), &merged_update, )); - let current_encrypted_auth_config = latest_key.encrypted_auth_config.clone(); - if let Some(encrypted_auth_config) = encrypted_auth_config.as_ref() { - latest_key.encrypted_auth_config = Some(encrypted_auth_config.clone()); - } - if encrypted_auth_config.is_some() { - ( - latest_key.oauth_invalid_at_unix_secs, - latest_key.oauth_invalid_reason, - ) = merge_codex_oauth_response_state( - &latest_key, - oauth_invalid_at_unix_secs, - oauth_invalid_reason.as_deref(), - merge_context.observed_at_unix_secs, - ); - } latest_key.status_snapshot = sync_provider_key_quota_status_snapshot( latest_key.status_snapshot.as_ref(), "codex", latest_key.upstream_metadata.as_ref(), "refresh_api", ); - if encrypted_auth_config.is_some() { - latest_key.status_snapshot = sync_provider_key_oauth_status_snapshot( - latest_key.status_snapshot.as_ref(), - &latest_key, - ); - } latest_key.updated_at_unix_secs = SystemTime::now() .duration_since(UNIX_EPOCH) .ok() .map(|duration| duration.as_secs()); - let persisted = if let Some(encrypted_auth_config) = encrypted_auth_config.as_ref() { - state - .app() - .compare_and_update_provider_catalog_key_oauth_runtime_state( - &ProviderCatalogKeyOAuthRuntimeStateCasUpdate { - key_id: key_id.to_string(), - expected_encrypted_auth_config: current_encrypted_auth_config, - expected_credential: None, - expected_upstream_metadata_namespace: Some( - ProviderCatalogUpstreamMetadataNamespaceExpectation { - namespace: "codex".to_string(), - expected_value: expected_codex, - }, - ), - encrypted_auth_config: encrypted_auth_config.clone(), - encrypted_api_key_update: None, - expires_at_unix_secs_update: None, - oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs, - oauth_invalid_reason: latest_key.oauth_invalid_reason.clone(), - upstream_metadata_patch: Some(serde_json::json!({ - "codex": outcome.metadata - })), - upstream_metadata_namespace_to_remove: None, - status_snapshot_patch: provider_quota_refresh_status_patch( - latest_key.status_snapshot.as_ref(), - ), - reset_error_count: false, - updated_at_unix_secs: latest_key.updated_at_unix_secs, - }, - ) - .await? - } else { - state - .app() - .update_provider_catalog_key_runtime_metadata( - &ProviderCatalogKeyRuntimeMetadataUpdate { - key_id: key_id.to_string(), - namespace: "codex".to_string(), - expected_upstream_metadata_value: expected_codex, - upstream_metadata_value: outcome.metadata, - status_snapshot_patch: provider_quota_refresh_status_patch( - latest_key.status_snapshot.as_ref(), - ), - updated_at_unix_secs: latest_key.updated_at_unix_secs, - }, - ) - .await? - }; + let persisted = state + .app() + .update_provider_catalog_key_runtime_metadata( + &ProviderCatalogKeyRuntimeMetadataUpdate { + key_id: key_id.to_string(), + namespace: "codex".to_string(), + expected_upstream_metadata_value: expected_codex, + upstream_metadata_value: outcome.metadata, + status_snapshot_patch: provider_quota_refresh_status_patch( + latest_key.status_snapshot.as_ref(), + ), + updated_at_unix_secs: latest_key.updated_at_unix_secs, + }, + ) + .await?; if persisted { return Ok(true); } @@ -1133,7 +1201,7 @@ pub(crate) async fn persist_fenced_provider_quota_refresh_state( oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option, merge_context: admin_provider_quota_pure::CodexQuotaMergeContext<'_>, - expected_credential: Option<&ProviderCatalogKeyOAuthCredentialFence>, + expected_credential: &ProviderCatalogKeyOAuthCredentialFence, ) -> Result { let expected_encrypted_auth_config = expected_encrypted_auth_config.trim(); if expected_encrypted_auth_config.is_empty() { @@ -1260,7 +1328,7 @@ pub(crate) async fn persist_fenced_provider_quota_refresh_state( expected_encrypted_auth_config: Some( expected_encrypted_auth_config.to_string(), ), - expected_credential: expected_credential.cloned(), + expected_credential: Some(expected_credential.clone()), expected_upstream_metadata_namespace: Some( ProviderCatalogUpstreamMetadataNamespaceExpectation { namespace: "codex".to_string(), @@ -1300,13 +1368,18 @@ async fn persist_provider_quota_refresh_state_after_read( oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option, encrypted_auth_config: Option, + expected_credential_fence: Option<&ProviderTransportCredentialFence>, after_read: F, ) -> Result where F: std::future::Future, { + let safe_metadata_update = + metadata_update.map(|value| admin_provider_upstream_metadata_safe_json(Some(value))); + let metadata_update = safe_metadata_update.as_ref(); let Some(mut latest_key) = state - .read_provider_catalog_keys_by_ids(&[key_id.to_string()]) + .app() + .list_provider_catalog_keys_by_ids_strong(&[key_id.to_string()]) .await? .into_iter() .next() @@ -1317,6 +1390,7 @@ where // Keep the namespace values observed before applying the refresh response; // each runtime metadata write uses them as its CAS expectation. + let observed_encrypted_auth_config = latest_key.encrypted_auth_config.clone(); let observed_upstream_metadata = latest_key.upstream_metadata.clone(); let mut quota_snapshot_provider_type = None::; if let Some(metadata_update) = metadata_update { @@ -1350,19 +1424,85 @@ where let metadata_updates = metadata_update .and_then(serde_json::Value::as_object) .map(|updates| { + let merged = latest_key + .upstream_metadata + .as_ref() + .and_then(serde_json::Value::as_object); updates - .iter() - .map(|(namespace, value)| (namespace.clone(), value.clone())) + .keys() + .filter_map(|namespace| { + merged + .and_then(|metadata| metadata.get(namespace)) + .cloned() + .map(|value| (namespace.clone(), value)) + }) .collect::>() }) .unwrap_or_default(); + if let Some(expected_credential_fence) = expected_credential_fence { + if observed_encrypted_auth_config.as_deref() + != Some(expected_credential_fence.encrypted_auth_config.as_str()) + || latest_key.encrypted_api_key + != expected_credential_fence.credential.encrypted_api_key + || latest_key.auth_type != expected_credential_fence.credential.auth_type + || latest_key.provider_id != expected_credential_fence.credential.provider_id + { + return Ok(false); + } + if metadata_updates.len() > 1 { + return Err(GatewayError::Internal( + "credential-fenced quota refresh may update at most one metadata namespace" + .to_string(), + )); + } + let expected_upstream_metadata_namespace = + metadata_updates.first().map(|(namespace, _)| { + ProviderCatalogUpstreamMetadataNamespaceExpectation { + namespace: namespace.clone(), + expected_value: observed_upstream_metadata + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|metadata| metadata.get(namespace)) + .cloned(), + } + }); + return state + .app() + .compare_and_update_provider_catalog_key_oauth_runtime_state( + &ProviderCatalogKeyOAuthRuntimeStateCasUpdate { + key_id: key_id.to_string(), + expected_encrypted_auth_config: Some( + expected_credential_fence.encrypted_auth_config.clone(), + ), + expected_credential: Some(expected_credential_fence.credential.clone()), + expected_upstream_metadata_namespace, + encrypted_auth_config: encrypted_auth_config + .clone() + .unwrap_or_else(|| expected_credential_fence.encrypted_auth_config.clone()), + encrypted_api_key_update: None, + expires_at_unix_secs_update: None, + oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs, + oauth_invalid_reason: latest_key.oauth_invalid_reason.clone(), + upstream_metadata_patch: metadata_update.cloned(), + upstream_metadata_namespace_to_remove: None, + status_snapshot_patch: status_patch, + reset_error_count: false, + updated_at_unix_secs: latest_key.updated_at_unix_secs, + }, + ) + .await; + } + if encrypted_auth_config.is_some() { + return Err(GatewayError::Internal( + "provider quota credential update requires a pre-request credential fence".to_string(), + )); + } if metadata_updates.is_empty() { if !state .update_provider_catalog_key_oauth_runtime_state( key_id, latest_key.oauth_invalid_at_unix_secs, latest_key.oauth_invalid_reason.as_deref(), - encrypted_auth_config.as_deref(), latest_key.updated_at_unix_secs, ) .await? @@ -1384,7 +1524,7 @@ where } else { serde_json::json!({}) }; - let mut expected = observed_upstream_metadata + let expected = observed_upstream_metadata .as_ref() .and_then(serde_json::Value::as_object) .and_then(|metadata| metadata.get(namespace)) @@ -1413,7 +1553,6 @@ where key_id, latest_key.oauth_invalid_at_unix_secs, latest_key.oauth_invalid_reason.as_deref(), - encrypted_auth_config.as_deref(), latest_key.updated_at_unix_secs, ) .await @@ -1439,6 +1578,21 @@ pub(super) async fn execute_provider_quota_plan( plan: ExecutionPlan, quota_kind: &str, ) -> Result { + let provider_name = plan.provider_name.as_deref().unwrap_or_default(); + if !provider_quota_url_has_allowed_origin(provider_name, &plan.url) { + warn!( + key_id = %transport.key.id, + endpoint_id = %transport.endpoint.id, + provider_name, + quota_kind, + upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), + "gateway provider quota request blocked by origin policy" + ); + return Ok(ProviderQuotaExecutionOutcome::Failure( + "Provider quota request origin is not allowed".to_string(), + )); + } + match state.execute_execution_runtime_sync_plan(None, &plan).await { Ok(result) => { if !crate::provider_transport::is_codex_agent_identity_transport(transport) @@ -1457,17 +1611,16 @@ pub(super) async fn execute_provider_quota_plan( "Agent Identity 任务重注册未返回认证信息".to_string(), )); } - Err(error) => { + Err(_) => { warn!( key_id = %transport.key.id, endpoint_id = %transport.endpoint.id, quota_kind = %quota_kind, - error = %error, "gateway Agent Identity quota task recovery failed" ); - return Ok(ProviderQuotaExecutionOutcome::Failure(format!( - "Agent Identity 任务重注册失败: {error}" - ))); + return Ok(ProviderQuotaExecutionOutcome::Failure( + "Agent Identity 任务重注册失败".to_string(), + )); } }; let header_name = refreshed_entry.auth_header_name.trim().to_ascii_lowercase(); @@ -1490,21 +1643,20 @@ pub(super) async fn execute_provider_quota_plan( .await { Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)), - Err(error) => { - let error = error.into_message(); + Err(_) => { warn!( key_id = %transport.key.id, endpoint_id = %transport.endpoint.id, quota_kind = %quota_kind, - error = %error, "gateway Agent Identity quota task recovery retry failed" ); - Ok(ProviderQuotaExecutionOutcome::Failure(error)) + Ok(ProviderQuotaExecutionOutcome::Failure( + "Provider quota request failed".to_string(), + )) } } } - Err(err) => { - let error = err.into_message(); + Err(_) => { let proxy_node_id = plan .proxy .as_ref() @@ -1524,24 +1676,78 @@ pub(super) async fn execute_provider_quota_plan( warn!( key_id = %transport.key.id, endpoint_id = %transport.endpoint.id, - url = %plan.url, + upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), proxy_source = ?proxy_source, proxy_node_id = ?proxy_node_id, proxy_url_present, - error = %error, quota_kind = %quota_kind, "gateway provider quota execution runtime request failed" ); - Ok(ProviderQuotaExecutionOutcome::Failure(error)) + Ok(ProviderQuotaExecutionOutcome::Failure( + "Provider quota request failed".to_string(), + )) } } } +fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bool { + let Ok(url) = url::Url::parse(value) else { + return false; + }; + if url.scheme() != "https" + || !url.username().is_empty() + || url.password().is_some() + || url.port_or_known_default() != Some(443) + { + return false; + } + let Some(host) = url.host_str() else { + return false; + }; + + match provider_name.trim().to_ascii_lowercase().as_str() { + "antigravity" => matches!( + host, + "cloudcode-pa.googleapis.com" + | "daily-cloudcode-pa.googleapis.com" + | "daily-cloudcode-pa.sandbox.googleapis.com" + ), + "gemini_cli" => host == "cloudcode-pa.googleapis.com", + "chatgpt_web" | "codex" => host == "chatgpt.com", + "grok" => host == "grok.com", + "windsurf" => host == "server.codeium.com", + "kiro" => kiro_quota_host_is_allowed(host), + _ => false, + } +} + +fn kiro_quota_host_is_allowed(host: &str) -> bool { + let Some(region) = host + .strip_prefix("q.") + .and_then(|host| host.strip_suffix(".amazonaws.com")) + else { + return false; + }; + !region.is_empty() + && region + .bytes() + .all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'-') + && region + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && region + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric) +} + #[cfg(test)] mod tests { use super::*; use crate::data::GatewayDataState; use crate::AppState; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogKey, @@ -1550,6 +1756,124 @@ mod tests { use serde_json::json; use std::sync::Arc; + #[test] + fn quota_plans_cannot_enable_redirects_through_header_overrides() { + let mut headers = std::collections::BTreeMap::from([ + ( + "X-Aether-Execution-Follow-Redirects".to_string(), + "true".to_string(), + ), + ("authorization".to_string(), "Bearer secret".to_string()), + ]); + + force_provider_quota_redirects_disabled(&mut headers); + + assert_eq!( + headers + .get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) + .map(String::as_str), + Some("false") + ); + assert_eq!( + headers + .keys() + .filter(|name| { + name.eq_ignore_ascii_case(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) + }) + .count(), + 1 + ); + assert_eq!( + headers.get("authorization").map(String::as_str), + Some("Bearer secret") + ); + } + + #[test] + fn quota_origin_policy_accepts_only_provider_owned_https_origins() { + for (provider_name, url) in [ + ( + "antigravity", + "https://cloudcode-pa.googleapis.com/v1internal:fetchAvailableModels", + ), + ( + "antigravity", + "https://daily-cloudcode-pa.googleapis.com/v1internal:fetchAvailableModels", + ), + ( + "antigravity", + "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:fetchAvailableModels", + ), + ( + "gemini_cli", + "https://cloudcode-pa.googleapis.com/v1internal:retrieveUserQuota", + ), + ( + "chatgpt_web", + "https://chatgpt.com/backend-api/conversation/init", + ), + ("codex", "https://chatgpt.com/backend-api/wham/usage"), + ("grok", "https://grok.com/rest/rate-limits"), + ( + "windsurf", + "https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus", + ), + ( + "kiro", + "https://q.us-east-1.amazonaws.com/getUsageLimits?origin=AI_EDITOR", + ), + ] { + assert!( + provider_quota_url_has_allowed_origin(provider_name, url), + "expected {provider_name} quota URL to be allowed: {url}" + ); + } + } + + #[test] + fn quota_origin_policy_rejects_ssrf_and_credential_redirect_origins() { + for (provider_name, url) in [ + ( + "chatgpt_web", + "http://chatgpt.com/backend-api/conversation/init", + ), + ("codex", "https://chatgpt.com:444/backend-api/wham/usage"), + ( + "codex", + "https://user:secret@chatgpt.com/backend-api/wham/usage", + ), + ( + "codex", + "https://chatgpt.com.attacker.test/backend-api/wham/usage", + ), + ("grok", "https://grok.com.attacker.test/rest/rate-limits"), + ("windsurf", "https://server.codeium.com.attacker.test/quota"), + ( + "gemini_cli", + "https://quota-proxy.internal/retrieveUserQuota", + ), + ( + "antigravity", + "https://cloudcode-pa.googleapis.com.attacker.test/quota", + ), + ("kiro", "https://q.localhost:8443/getUsageLimits"), + ( + "kiro", + "https://q.us-east-1.evil.amazonaws.com/getUsageLimits", + ), + ( + "kiro", + "https://q.us-east-1.amazonaws.com.attacker.test/getUsageLimits", + ), + ("unknown", "https://chatgpt.com/backend-api/wham/usage"), + ] { + assert!( + !provider_quota_url_has_allowed_origin(provider_name, url), + "expected {provider_name} quota URL to be rejected: {url}" + ); + } + } + fn codex_merge_context( request_started_at_unix_ms: u64, ) -> admin_provider_quota_pure::CodexQuotaMergeContext<'static> { @@ -1590,18 +1914,41 @@ mod tests { fn codex_refresh_test_state( key_id: &str, - encrypted_auth_config: Option<&str>, - ) -> (AppState, Arc) { + auth_config: Option<&str>, + ) -> ( + AppState, + Arc, + Option, + ) { + let provider = StoredProviderCatalogProvider::new( + "provider-codex-refresh".to_string(), + "Codex Refresh".to_string(), + None, + "codex".to_string(), + ) + .expect("provider should build"); + let bootstrap = AppState::new() + .expect("bootstrap app should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_auth_config = auth_config + .map(|plaintext| { + bootstrap.seal_provider_catalog_key_auth_config(&provider.id, key_id, plaintext) + }) + .transpose() + .expect("auth config should seal"); let mut key = StoredProviderCatalogKey::new( key_id.to_string(), - "provider-codex-refresh".to_string(), + provider.id.clone(), "Codex Refresh".to_string(), "oauth".to_string(), None, true, ) .expect("key should build"); - key.encrypted_auth_config = encrypted_auth_config.map(ToOwned::to_owned); + key.encrypted_auth_config = encrypted_auth_config.clone(); key.upstream_metadata = Some(json!({ "codex": { "plan_type": "plus", @@ -1612,8 +1959,18 @@ mod tests { "updated_at": 200u64 } })); + let credential_fence = + encrypted_auth_config.map(|encrypted_auth_config| ProviderTransportCredentialFence { + encrypted_auth_config, + credential: ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: key.encrypted_api_key.clone(), + auth_type: key.auth_type.clone(), + provider_id: key.provider_id.clone(), + provider_type: provider.provider_type.clone(), + }, + }); let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![], + vec![provider], vec![], vec![key], )); @@ -1622,9 +1979,10 @@ mod tests { .with_data_state_for_tests( GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone( &repository, - )), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); - (app, repository) + (app, repository, credential_fence) } fn codex_reset_state_machine_test_state( @@ -1632,7 +1990,7 @@ mod tests { ) -> ( AppState, Arc, - ProviderCatalogKeyOAuthCredentialFence, + ProviderTransportCredentialFence, ) { let provider = StoredProviderCatalogProvider::new( "provider-codex-reset-state".to_string(), @@ -1641,6 +1999,15 @@ mod tests { "codex".to_string(), ) .expect("provider should build"); + let bootstrap = AppState::new() + .expect("bootstrap app should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_auth_config = bootstrap + .seal_provider_catalog_key_auth_config(&provider.id, key_id, "auth-v1") + .expect("auth config should seal"); let mut key = StoredProviderCatalogKey::new( key_id.to_string(), provider.id.clone(), @@ -1650,15 +2017,30 @@ mod tests { true, ) .expect("key should build"); - key.encrypted_auth_config = Some("auth-v1".to_string()); + key.encrypted_auth_config = Some(encrypted_auth_config.clone()); key.upstream_metadata = Some(json!({ - "codex": {"credential_generation": "credential-v1"} + "codex": { + "credential_generation": "credential-v1", + "reset_credits": { + "available_count": 2, + "updated_at": 100u64, + "detail_source": "wham_readonly", + "detail_status": "available", + "credits": [ + {"id": "credit-1", "expires_at": 20_000u64}, + {"id": "credit-2", "expires_at": 30_000u64} + ] + } + } })); - let credential = ProviderCatalogKeyOAuthCredentialFence { - encrypted_api_key: None, - auth_type: key.auth_type.clone(), - provider_id: provider.id.clone(), - provider_type: provider.provider_type.clone(), + let credential_fence = ProviderTransportCredentialFence { + encrypted_auth_config, + credential: ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: None, + auth_type: key.auth_type.clone(), + provider_id: provider.id.clone(), + provider_type: provider.provider_type.clone(), + }, }; let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -1670,9 +2052,10 @@ mod tests { .with_data_state_for_tests( GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone( &repository, - )), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); - (app, repository, credential) + (app, repository, credential_fence) } #[tokio::test] @@ -1684,8 +2067,8 @@ mod tests { let first = reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-a", ) @@ -1699,8 +2082,8 @@ mod tests { let same = reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-a", ) @@ -1714,8 +2097,8 @@ mod tests { let other = reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-b", ) @@ -1746,12 +2129,19 @@ mod tests { let key_id = "key-codex-reset-credential-generation"; let (app, repository, credential) = codex_reset_state_machine_test_state(key_id); let admin_state = AdminAppState::new(&app); + let original_metadata = repository + .list_keys_by_ids(&[key_id.to_string()]) + .await + .expect("key should load before reservation") + .pop() + .expect("key should exist before reservation") + .upstream_metadata; let result = reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-before-rebind"), "reset-from-old-account", ) @@ -1769,10 +2159,7 @@ mod tests { .expect("key should reload") .pop() .expect("key should exist"); - assert_eq!( - stored.upstream_metadata.unwrap()["codex"], - json!({"credential_generation":"credential-v1"}) - ); + assert_eq!(stored.upstream_metadata, original_metadata); } #[tokio::test] @@ -1783,8 +2170,8 @@ mod tests { let reservation = match reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-noop", ) @@ -1799,8 +2186,8 @@ mod tests { complete_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &reservation, "nothing_to_reset", 200_000, @@ -1812,8 +2199,8 @@ mod tests { let next = reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-next", ) @@ -1847,8 +2234,8 @@ mod tests { let reservation = match reserve_codex_account_reset( &admin_state, &key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "same-id", ) @@ -1862,8 +2249,8 @@ mod tests { complete_codex_account_reset( &admin_state, &key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &reservation, first_outcome, 200_000, @@ -1874,8 +2261,8 @@ mod tests { complete_codex_account_reset( &admin_state, &key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &reservation, second_outcome, 210_000, @@ -1907,6 +2294,11 @@ mod tests { codex["account_quota_reset_history"][0]["outcome"], json!("reset") ); + assert_eq!(codex["reset_credits"]["available_count"], json!(1u64)); + assert_eq!( + codex["reset_credits"]["credits"], + json!([{"id": "credit-2", "expires_at": 30_000u64}]) + ); } } @@ -1918,8 +2310,8 @@ mod tests { let first = match reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-first", ) @@ -1933,8 +2325,8 @@ mod tests { complete_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &first, "nothing_to_reset", 200_000, @@ -1945,8 +2337,8 @@ mod tests { let second = match reserve_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, Some("credential-v1"), "reset-second", ) @@ -1963,8 +2355,8 @@ mod tests { complete_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &first, "reset", 210_000, @@ -1994,8 +2386,8 @@ mod tests { complete_codex_account_reset( &admin_state, key_id, - "auth-v1", - &credential, + credential.encrypted_auth_config.as_str(), + &credential.credential, &second, "reset", 220_000, @@ -2020,7 +2412,7 @@ mod tests { #[tokio::test] async fn stale_codex_refresh_cannot_lower_realtime_usage() { let key_id = "key-codex-refresh-monotonic"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); let stale_refresh = json!({"codex": { "plan_type": "plus", @@ -2034,9 +2426,6 @@ mod tests { &admin_state, key_id, Some(&stale_refresh), - None, - None, - None, codex_merge_context(100_000), ) .await @@ -2059,7 +2448,7 @@ mod tests { #[tokio::test] async fn codex_reset_fence_is_idempotent_and_rejects_pre_reset_response() { let key_id = "key-codex-reset-fence"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); let initial_fence = persist_codex_account_reset_fence( @@ -2109,9 +2498,6 @@ mod tests { &admin_state, key_id, Some(&stale), - None, - None, - None, codex_merge_context(200_000), ) .await @@ -2126,9 +2512,6 @@ mod tests { &admin_state, key_id, Some(&baseline), - None, - None, - None, codex_reset_merge_context(260_000, "fence-a"), ) .await @@ -2149,7 +2532,7 @@ mod tests { #[tokio::test] async fn codex_reset_fence_barrier_never_moves_backward_and_remembers_processed_ids() { let key_id = "key-codex-reset-fence-order"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); let newer = persist_codex_account_reset_fence( @@ -2209,7 +2592,7 @@ mod tests { #[tokio::test] async fn concurrent_codex_reset_fences_converge_on_newest_barrier() { let key_id = "key-codex-reset-fence-concurrent"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); let (older, newer) = tokio::join!( @@ -2273,7 +2656,7 @@ mod tests { #[tokio::test] async fn superseded_codex_reset_refresh_cannot_confirm_newer_fence() { let key_id = "key-codex-reset-fence-stale-refresh"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); for (fence_unix_ms, fence_id, redeem_id) in [ @@ -2304,9 +2687,6 @@ mod tests { &admin_state, key_id, Some(&stale_baseline), - None, - None, - None, admin_provider_quota_pure::CodexQuotaMergeContext { observed_at_unix_secs: 310, request_started_at_unix_ms: Some(310_000), @@ -2336,7 +2716,7 @@ mod tests { #[tokio::test] async fn replaying_historical_codex_reset_does_not_reopen_pending() { let key_id = "key-codex-reset-fence-replay"; - let (app, repository) = codex_refresh_test_state(key_id, None); + let (app, repository, _) = codex_refresh_test_state(key_id, None); let admin_state = AdminAppState::new(&app); for (fence_unix_ms, fence_id, redeem_id, request_started_at_unix_ms, usage) in [ @@ -2365,9 +2745,6 @@ mod tests { &admin_state, key_id, Some(&baseline), - None, - None, - None, admin_provider_quota_pure::CodexQuotaMergeContext { observed_at_unix_secs: request_started_at_unix_ms / 1_000, request_started_at_unix_ms: Some(request_started_at_unix_ms), @@ -2416,7 +2793,8 @@ mod tests { #[tokio::test] async fn fenced_stale_codex_refresh_keeps_usage_and_oauth_state() { let key_id = "key-codex-fenced-refresh-monotonic"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let stale_refresh = json!({"codex": { "primary_used_percent": 50.0, @@ -2427,12 +2805,12 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some(&stale_refresh), Some(300), Some("refresh-state".to_string()), codex_merge_context(100_000), - None, + &credential_fence.credential, ) .await .expect("fenced refresh persistence should complete")); @@ -2454,7 +2832,8 @@ mod tests { #[tokio::test] async fn fenced_older_refresh_cannot_overwrite_newer_oauth_state() { let key_id = "key-codex-fenced-oauth-watermark"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let quota = |used_percent| { json!({"codex": { @@ -2467,24 +2846,24 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a(70.0)), None, None, codex_merge_context(300_000), - None, + &credential_fence.credential, ) .await .expect("newer refresh should persist")); assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a(65.0)), Some(250), Some("stale-invalid".to_string()), codex_merge_context(250_000), - None, + &credential_fence.credential, ) .await .expect("older refresh should merge without replacing OAuth state")); @@ -2507,7 +2886,8 @@ mod tests { #[tokio::test] async fn fenced_same_millisecond_refresh_uses_request_id_for_oauth_order() { let key_id = "key-codex-fenced-oauth-id-watermark"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let quota = json!({"codex": { "primary_used_percent": 70.0, @@ -2518,24 +2898,24 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), None, None, codex_merge_context_with_id(300_000, Some("request-b")), - None, + &credential_fence.credential, ) .await .expect("newer same-millisecond refresh should persist")); assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), Some(300), Some("stale-invalid".to_string()), codex_merge_context_with_id(300_000, Some("request-a")), - None, + &credential_fence.credential, ) .await .expect("older same-millisecond refresh should merge without replacing OAuth state")); @@ -2562,7 +2942,8 @@ mod tests { #[tokio::test] async fn fenced_same_millisecond_newer_request_id_can_replace_oauth_state() { let key_id = "key-codex-fenced-oauth-id-watermark-newer-invalid"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let quota = json!({"codex": { "primary_used_percent": 70.0, @@ -2573,24 +2954,24 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), None, None, codex_merge_context_with_id(300_000, Some("request-a")), - None, + &credential_fence.credential, ) .await .expect("older same-millisecond refresh should persist")); assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), Some(300), Some("newer-invalid".to_string()), codex_merge_context_with_id(300_000, Some("request-b")), - None, + &credential_fence.credential, ) .await .expect("newer same-millisecond refresh should replace OAuth state")); @@ -2620,7 +3001,8 @@ mod tests { #[tokio::test] async fn fenced_older_success_cannot_clear_newer_oauth_invalid_state() { let key_id = "key-codex-newer-invalid-older-success"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let quota = json!({"codex": { "primary_used_percent": 70.0, @@ -2631,24 +3013,24 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), Some(300), Some("newer-invalid".to_string()), codex_merge_context_with_id(300_000, Some("request-newer")), - None, + &credential_fence.credential, ) .await .expect("newer invalid response should persist")); assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), None, None, codex_merge_context_with_id(250_000, Some("request-older")), - None, + &credential_fence.credential, ) .await .expect("older success should be harmlessly acknowledged")); @@ -2678,7 +3060,8 @@ mod tests { #[tokio::test] async fn fenced_same_millisecond_older_success_cannot_clear_newer_invalid_state() { let key_id = "key-codex-same-ms-newer-invalid-older-success"; - let (app, repository) = codex_refresh_test_state(key_id, Some("auth-v1")); + let (app, repository, credential_fence) = codex_refresh_test_state(key_id, Some("auth-v1")); + let credential_fence = credential_fence.expect("credential fence should exist"); let admin_state = AdminAppState::new(&app); let quota = json!({"codex": { "primary_used_percent": 70.0, @@ -2689,24 +3072,24 @@ mod tests { assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), Some(300), Some("newer-invalid".to_string()), codex_merge_context_with_id(300_000, Some("request-b")), - None, + &credential_fence.credential, ) .await .expect("newer same-millisecond invalid response should persist")); assert!(persist_fenced_provider_quota_refresh_state( &admin_state, key_id, - "auth-v1", + credential_fence.encrypted_auth_config.as_str(), Some("a), None, None, codex_merge_context_with_id(300_000, Some("request-a")), - None, + &credential_fence.credential, ) .await .expect("older same-millisecond success should be acknowledged")); @@ -2731,23 +3114,43 @@ mod tests { #[tokio::test] async fn metadata_cas_conflict_does_not_persist_stale_oauth_runtime_state() { + let provider_id = "provider-codex-cas"; + let key_id = "key-codex-cas"; + let bootstrap = AppState::new() + .expect("bootstrap app should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let old_auth_config = bootstrap + .seal_provider_catalog_key_auth_config(provider_id, key_id, "old-auth-config") + .expect("old auth config should seal"); + let new_auth_config = bootstrap + .seal_provider_catalog_key_auth_config(provider_id, key_id, "new-auth-config") + .expect("new auth config should seal"); let mut key = StoredProviderCatalogKey::new( - "key-codex-cas".to_string(), - "provider-codex-cas".to_string(), + key_id.to_string(), + provider_id.to_string(), "Codex CAS".to_string(), "oauth".to_string(), None, true, ) .expect("key should build"); - key.encrypted_auth_config = Some("old-auth-config".to_string()); + key.encrypted_auth_config = Some(old_auth_config.clone()); key.oauth_invalid_at_unix_secs = Some(100); key.oauth_invalid_reason = Some("old-invalid-reason".to_string()); key.upstream_metadata = Some(json!({"codex":{"remaining":5}})); key.status_snapshot = Some(json!({"oauth":{"invalid":true}})); let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![], + vec![StoredProviderCatalogProvider::new( + provider_id.to_string(), + "Codex CAS".to_string(), + None, + "codex".to_string(), + ) + .expect("provider should build")], vec![], vec![key], )); @@ -2756,19 +3159,30 @@ mod tests { .with_data_state_for_tests( GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone( &repository, - )), + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let admin_state = AdminAppState::new(&app); + let credential_fence = ProviderTransportCredentialFence { + encrypted_auth_config: old_auth_config.clone(), + credential: ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: None, + auth_type: "oauth".to_string(), + provider_id: provider_id.to_string(), + provider_type: "codex".to_string(), + }, + }; let concurrent_repository = Arc::clone(&repository); let metadata_update = json!({"codex":{"remaining":3}}); let persisted = persist_provider_quota_refresh_state_after_read( &admin_state, - "key-codex-cas", + key_id, Some(&metadata_update), Some(200), Some("new-invalid-reason".to_string()), - Some("new-auth-config".to_string()), + Some(new_auth_config), + Some(&credential_fence), async move { assert!(concurrent_repository .update_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate { @@ -2795,7 +3209,7 @@ mod tests { .expect("key should remain"); assert_eq!( stored.encrypted_auth_config.as_deref(), - Some("old-auth-config") + Some(old_auth_config.as_str()) ); assert_eq!(stored.oauth_invalid_at_unix_secs, Some(100)); assert_eq!( @@ -2807,4 +3221,121 @@ mod tests { json!({"remaining":4}) ); } + + #[tokio::test] + async fn quota_refresh_strong_read_bypasses_stale_provider_catalog_cache() { + let mut key = StoredProviderCatalogKey::new( + "key-antigravity-stale-cache".to_string(), + "provider-antigravity-stale-cache".to_string(), + "Antigravity stale cache".to_string(), + "oauth".to_string(), + None, + true, + ) + .expect("key should build"); + key.upstream_metadata = Some(json!({ + "antigravity": { + "project_id": "project-1", + "quota_by_model": { + "gemini-3.7-flash-tiered": {"remaining_fraction": 0.9} + } + } + })); + + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![], + vec![], + vec![key], + )); + let data = + GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(&repository)) + .with_cached_provider_catalog_reader_for_tests(Arc::clone(&repository)); + let app = AppState::new() + .expect("app should build") + .with_data_state_for_tests(data); + let admin_state = AdminAppState::new(&app); + let key_ids = ["key-antigravity-stale-cache".to_string()]; + + let cached = app + .read_provider_catalog_keys_by_ids(&key_ids) + .await + .expect("initial cached read should succeed"); + assert_eq!( + cached[0].upstream_metadata.as_ref().unwrap()["antigravity"]["quota_by_model"] + ["gemini-3.7-flash-tiered"]["remaining_fraction"], + json!(0.9) + ); + + let current_namespace = json!({ + "project_id": "project-1", + "model_fetch_revision": 2, + "quota_by_model": { + "gemini-3.7-flash-tiered": {"remaining_fraction": 0.7} + } + }); + assert!(repository + .upsert_key_upstream_metadata_namespace( + "key-antigravity-stale-cache", + "antigravity", + ¤t_namespace, + None, + ) + .await + .expect("out-of-band metadata update should succeed")); + let still_cached = app + .read_provider_catalog_keys_by_ids(&key_ids) + .await + .expect("stale cached read should succeed"); + assert_eq!( + still_cached[0].upstream_metadata.as_ref().unwrap()["antigravity"]["quota_by_model"] + ["gemini-3.7-flash-tiered"]["remaining_fraction"], + json!(0.9), + "regression setup must keep the ordinary read stale" + ); + + let metadata_update = json!({ + "antigravity": { + "project_id": "project-1", + "quota_by_model": { + "gemini-3.7-flash-tiered": {"remaining_fraction": 0.6} + }, + "quota_groups": [{ + "display_name": "Gemini models", + "buckets": [{"bucket_id": "gemini-weekly", "window": "weekly"}] + }] + } + }); + assert!(persist_provider_quota_refresh_state( + &admin_state, + "key-antigravity-stale-cache", + Some(&metadata_update), + None, + None, + None, + ) + .await + .expect("quota refresh persistence should not error")); + + let stored = repository + .list_keys_by_ids(&key_ids) + .await + .expect("key should reload") + .pop() + .expect("key should exist"); + assert_eq!( + stored.upstream_metadata.as_ref().unwrap()["antigravity"]["quota_groups"][0]["buckets"] + [0]["bucket_id"], + json!("gemini-weekly") + ); + assert_eq!( + stored.upstream_metadata.as_ref().unwrap()["antigravity"]["model_fetch_revision"], + json!(2), + "quota refresh must preserve fields written by another Antigravity metadata producer" + ); + assert_eq!( + stored.upstream_metadata.as_ref().unwrap()["antigravity"]["quota_by_model"] + ["gemini-3.7-flash-tiered"]["remaining_fraction"], + json!(0.6) + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/windsurf.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/windsurf.rs index f506feb21..7e6eed480 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/windsurf.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/windsurf.rs @@ -6,6 +6,7 @@ use super::shared::{ }; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; use crate::GatewayError; +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, @@ -132,7 +133,7 @@ fn merge_windsurf_probe_metadata( user_status_metadata } -fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, message: String) { +fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, code: &str) { let Some(target) = metadata.as_object_mut() else { return; }; @@ -142,7 +143,7 @@ fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, if let Some(items) = warnings.as_array_mut() { items.push(json!({ "probe": probe, - "message": message, + "code": code, })); } } @@ -162,7 +163,13 @@ fn build_windsurf_metadata_update( for (key, value) in patch_object { merged_bucket.insert(key.clone(), value.clone()); } - json!({ "windsurf": merged_bucket }) + let merged_bucket = serde_json::Value::Object(merged_bucket); + json!({ + "windsurf": admin_provider_metadata_bucket_safe_json( + "windsurf", + Some(&merged_bucket), + ) + }) } fn sanitize_windsurf_probe_detail(detail: impl AsRef) -> String { @@ -170,77 +177,7 @@ fn sanitize_windsurf_probe_detail(detail: impl AsRef) -> String { if detail.is_empty() { return "-".to_string(); } - if let Ok(mut value) = serde_json::from_str::(detail) { - redact_windsurf_sensitive_json(&mut value); - return value.to_string().chars().take(500).collect(); - } - if contains_windsurf_sensitive_marker(detail) { - "[REDACTED upstream error body]".to_string() - } else { - detail.chars().take(500).collect() - } -} - -fn redact_windsurf_sensitive_json(value: &mut serde_json::Value) { - match value { - serde_json::Value::Object(object) => { - for (key, value) in object { - if is_windsurf_sensitive_key(key) { - *value = json!("[REDACTED]"); - } else { - redact_windsurf_sensitive_json(value); - } - } - } - serde_json::Value::Array(items) => { - for item in items { - redact_windsurf_sensitive_json(item); - } - } - serde_json::Value::String(text) if looks_like_windsurf_secret(text) => { - *text = "[REDACTED]".to_string(); - } - _ => {} - } -} - -fn is_windsurf_sensitive_key(key: &str) -> bool { - let normalized = key - .chars() - .filter(|ch| ch.is_ascii_alphanumeric()) - .collect::() - .to_ascii_lowercase(); - normalized.contains("token") - || normalized.contains("apikey") - || normalized.contains("password") - || normalized.contains("authorization") - || normalized.contains("secret") -} - -fn looks_like_windsurf_secret(value: &str) -> bool { - let value = value.trim(); - value.starts_with("devin-session-token$") - || value.starts_with("sk-") - || (value.len() > 80 && value.split('.').count() == 3) -} - -fn contains_windsurf_sensitive_marker(value: &str) -> bool { - let lowered = value.to_ascii_lowercase(); - [ - "token", - "api_key", - "apikey", - "sessiontoken", - "firebase_id_token", - "idtoken", - "authorization", - "password", - "secret", - "devin-session-token$", - ] - .iter() - .any(|marker| lowered.contains(marker)) - || value.contains("sk-") + "[REDACTED upstream error body]".to_string() } pub(crate) async fn refresh_windsurf_provider_quota_locally( @@ -293,14 +230,13 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( .await? { ProviderQuotaExecutionOutcome::Response(result) => result, - ProviderQuotaExecutionOutcome::Failure(detail) => { + ProviderQuotaExecutionOutcome::Failure(_) => { failed_count += 1; - let detail = sanitize_windsurf_probe_detail(detail); results.push(json!({ "key_id": key.id, "key_name": key.name, "status": "error", - "message": format!("GetUserStatus 请求执行失败: {detail}"), + "message": "GetUserStatus 请求执行失败", "status_code": 502, })); continue; @@ -353,22 +289,18 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( }) } ProviderQuotaExecutionOutcome::Response(model_result) => { - let detail = extract_execution_error_message(&model_result) - .unwrap_or_else(|| format!("HTTP {}", model_result.status_code)); - let detail = sanitize_windsurf_probe_detail(detail); append_windsurf_probe_warning( &mut metadata, "model_configs", - format!("GetCascadeModelConfigs 返回: {detail}"), + "response_failed", ); None } - ProviderQuotaExecutionOutcome::Failure(detail) => { - let detail = sanitize_windsurf_probe_detail(detail); + ProviderQuotaExecutionOutcome::Failure(_) => { append_windsurf_probe_warning( &mut metadata, "model_configs", - format!("GetCascadeModelConfigs 执行失败: {detail}"), + "execution_failed", ); None } @@ -396,22 +328,18 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( }) } ProviderQuotaExecutionOutcome::Response(rate_limit_result) => { - let detail = extract_execution_error_message(&rate_limit_result) - .unwrap_or_else(|| format!("HTTP {}", rate_limit_result.status_code)); - let detail = sanitize_windsurf_probe_detail(detail); append_windsurf_probe_warning( &mut metadata, "rate_limit", - format!("CheckUserMessageRateLimit 返回: {detail}"), + "response_failed", ); None } - ProviderQuotaExecutionOutcome::Failure(detail) => { - let detail = sanitize_windsurf_probe_detail(detail); + ProviderQuotaExecutionOutcome::Failure(_) => { append_windsurf_probe_warning( &mut metadata, "rate_limit", - format!("CheckUserMessageRateLimit 执行失败: {detail}"), + "execution_failed", ); None } @@ -457,7 +385,7 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( 401 | 403 => { oauth_invalid_at_unix_secs = Some(now_unix_secs); oauth_invalid_reason = - Some(format!("Windsurf token 无效或已被拒绝: {}", detail)); + Some("Windsurf token is invalid or rejected".to_string()); metadata.insert("banned".to_string(), json!(result.status_code == 403)); status = if result.status_code == 401 { "auth_invalid".to_string() @@ -470,10 +398,6 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( "rate_limit".to_string(), json!({ "limited": true, - "message": metadata - .get("last_error") - .cloned() - .unwrap_or_else(|| json!("rate limited")), }), ); status = "rate_limited".to_string(); @@ -525,9 +449,11 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally( if let Some(metadata) = metadata_update .as_ref() .and_then(|value| value.get("windsurf")) - .cloned() { - payload.insert("metadata".to_string(), metadata); + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("windsurf", Some(metadata)), + ); } if let Some(quota_snapshot) = build_quota_snapshot_payload( "windsurf", @@ -592,7 +518,7 @@ mod tests { r#"{"error":{"message":"bad"},"apiKey":"sk-secret","sessionToken":"devin-session-token$secret"}"#, ); - assert!(detail.contains("[REDACTED]")); + assert_eq!(detail, "[REDACTED upstream error body]"); assert!(!detail.contains("sk-secret")); assert!(!detail.contains("devin-session-token$secret")); } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs index bfda465fb..789b7b8e7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs @@ -4,8 +4,9 @@ use super::super::errors::{ use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate}; use aether_contracts::ProxySnapshot; use aether_oauth::provider::providers::{ - ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter, CLAUDE_CODE_PROVIDER_TYPE, - CLAUDE_CODE_TOKEN_URL, CLAUDE_CODE_WEB_BASE_URL, + AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter, + ANTIGRAVITY_USER_INFO_URL, CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_TOKEN_URL, + CLAUDE_CODE_WEB_BASE_URL, }; use aether_oauth::provider::{ ProviderOAuthCookieAuthorizationInput, ProviderOAuthService, ProviderOAuthTransportContext, @@ -13,12 +14,8 @@ use aether_oauth::provider::{ use axum::{body::Body, http, response::Response}; use std::sync::Arc; -fn provider_oauth_transport_error_detail(prefix: &str, error: &str) -> String { - let error = error.trim(); - if error.is_empty() { - return prefix.to_string(); - } - format!("{prefix}: {error}") +fn provider_oauth_transport_error_detail(prefix: &str, _error: &str) -> String { + prefix.to_string() } fn provider_oauth_exchange_context( @@ -43,7 +40,19 @@ fn provider_oauth_exchange_context( fn provider_oauth_service_for_template( template: AdminProviderOAuthTemplate, token_url: String, + antigravity_user_info_url: String, ) -> Result> { + if template.provider_type.eq_ignore_ascii_case("antigravity") { + let adapter = AntigravityProviderOAuthAdapter::default() + .with_token_url_override(token_url) + .with_user_info_url_override(antigravity_user_info_url); + #[cfg(test)] + let adapter = adapter.with_oauth_credentials_for_tests( + "gateway-test-antigravity-client-id", + "gateway-test-antigravity-client-secret", + ); + return Ok(ProviderOAuthService::new().with_adapter(Arc::new(adapter))); + } GenericProviderOAuthAdapter::for_provider_type(template.provider_type) .map(|adapter| adapter.with_token_url_override(token_url)) .map(|adapter| ProviderOAuthService::new().with_adapter(Arc::new(adapter))) @@ -75,7 +84,10 @@ pub(crate) async fn exchange_admin_provider_oauth_code( proxy: Option, ) -> Result> { let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url); - let service = provider_oauth_service_for_template(template, token_url)?; + let antigravity_user_info_url = + state.provider_oauth_token_url("antigravity_user_info", ANTIGRAVITY_USER_INFO_URL); + let service = + provider_oauth_service_for_template(template, token_url, antigravity_user_info_url)?; let ctx = provider_oauth_exchange_context(template.provider_type, proxy); let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state); let result = service @@ -103,7 +115,10 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token( proxy: Option, ) -> Result> { let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url); - let service = provider_oauth_service_for_template(template, token_url)?; + let antigravity_user_info_url = + state.provider_oauth_token_url("antigravity_user_info", ANTIGRAVITY_USER_INFO_URL); + let service = + provider_oauth_service_for_template(template, token_url, antigravity_user_info_url)?; let ctx = provider_oauth_exchange_context(template.provider_type, proxy); let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state); let input = aether_oauth::provider::ProviderOAuthImportInput { @@ -182,3 +197,20 @@ pub(crate) async fn authorize_admin_provider_oauth_with_cookie( ) }) } + +#[cfg(test)] +mod tests { + use super::provider_oauth_transport_error_detail; + + #[test] + fn provider_oauth_transport_error_does_not_reflect_network_details() { + let detail = provider_oauth_transport_error_detail( + "token exchange 失败", + "request failed for https://user:pass@example.test/token?secret=value authorization=Bearer upstream-secret", + ); + + assert_eq!(detail, "token exchange 失败"); + assert!(!detail.contains("upstream-secret")); + assert!(!detail.contains("user:pass")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs index aa2c8f384..3fa74dee7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs @@ -1,31 +1,29 @@ use crate::handlers::admin::request::AdminProviderOAuthTemplate; +use aether_oauth::core::OAuthError; use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext}; use serde_json::json; -use url::form_urlencoded; pub(crate) fn build_provider_oauth_start_response( template: AdminProviderOAuthTemplate, nonce: &str, code_challenge: Option<&str>, -) -> serde_json::Value { - let authorization_url = build_provider_oauth_authorization_url(template, nonce, code_challenge) - .unwrap_or_else(|| { - build_provider_oauth_authorization_url_legacy(template, nonce, code_challenge) - }); +) -> Result { + let authorization_url = + build_provider_oauth_authorization_url(template, nonce, code_challenge)?; - json!({ + Ok(json!({ "authorization_url": authorization_url, "redirect_uri": template.redirect_uri, "provider_type": template.provider_type, "instructions": "1) 打开 authorization_url 完成授权\n2) 复制授权页面显示的授权码或浏览器中的完整回调 URL\n3) 调用 complete 接口粘贴 callback_url", - }) + })) } fn build_provider_oauth_authorization_url( template: AdminProviderOAuthTemplate, nonce: &str, code_challenge: Option<&str>, -) -> Option { +) -> Result { let ctx = ProviderOAuthTransportContext { provider_id: String::new(), provider_type: template.provider_type.to_string(), @@ -41,32 +39,5 @@ fn build_provider_oauth_authorization_url( }; ProviderOAuthService::with_builtin_adapters() .build_authorize_url(&ctx, nonce, code_challenge) - .ok() .map(|response| response.authorize_url) } - -fn build_provider_oauth_authorization_url_legacy( - template: AdminProviderOAuthTemplate, - nonce: &str, - code_challenge: Option<&str>, -) -> String { - let mut serializer = form_urlencoded::Serializer::new(String::new()); - serializer.append_pair("client_id", template.client_id); - serializer.append_pair("response_type", "code"); - serializer.append_pair("redirect_uri", template.redirect_uri); - serializer.append_pair("scope", &template.scopes.join(" ")); - serializer.append_pair("state", nonce); - if template.provider_type == "codex" { - serializer.append_pair("prompt", "login"); - serializer.append_pair("id_token_add_organizations", "true"); - serializer.append_pair("codex_cli_simplified_flow", "true"); - } - if template.use_pkce { - if let Some(code_challenge) = code_challenge { - serializer.append_pair("code_challenge", code_challenge); - serializer.append_pair("code_challenge_method", "S256"); - } - } - - format!("{}?{}", template.authorize_url, serializer.finish()) -} diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/probe.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/probe.rs index aa003d8a4..d015ea8a3 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/probe.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/probe.rs @@ -24,11 +24,13 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin( .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or("/api/user/checkin"); - let url = admin_provider_ops_request_url( + let Ok(url) = admin_provider_ops_request_url( base_url, &admin_provider_ops_json_object_map(json!({ "endpoint": endpoint })), endpoint, - ); + ) else { + return None; + }; let (status, response_json) = match admin_provider_ops_execute_json_request( state, "provider-ops-action:probe_checkin", @@ -58,8 +60,15 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin( cookie_expired: true, }); } + if status != http::StatusCode::OK { + return Some(AdminProviderOpsCheckinOutcome { + success: Some(false), + message: "签到失败".to_string(), + cookie_expired: false, + }); + } - let message = response_json + let upstream_message = response_json .get("message") .and_then(serde_json::Value::as_str) .unwrap_or_default() @@ -72,43 +81,27 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin( { return Some(AdminProviderOpsCheckinOutcome { success: Some(true), - message: if message.is_empty() { - "签到成功".to_string() - } else { - message - }, + message: "签到成功".to_string(), cookie_expired: false, }); } - if admin_provider_ops_checkin_already_done(&message) { + if admin_provider_ops_checkin_already_done(&upstream_message) { return Some(AdminProviderOpsCheckinOutcome { success: None, - message: if message.is_empty() { - "今日已签到".to_string() - } else { - message - }, + message: "今日已签到".to_string(), cookie_expired: false, }); } - if admin_provider_ops_checkin_auth_failure(&message) { + if admin_provider_ops_checkin_auth_failure(&upstream_message) { return has_cookie.then(|| AdminProviderOpsCheckinOutcome { success: None, - message: if message.is_empty() { - "Cookie 已失效".to_string() - } else { - message - }, + message: "Cookie 已失效".to_string(), cookie_expired: true, }); } Some(AdminProviderOpsCheckinOutcome { success: Some(false), - message: if message.is_empty() { - "签到失败".to_string() - } else { - message - }, + message: "签到失败".to_string(), cookie_expired: false, }) } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/run.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/run.rs index be38d1948..b97235f85 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/run.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/run.rs @@ -32,7 +32,12 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action( ); } - let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin"); + let url = match admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin") { + Ok(url) => url, + Err(message) => { + return admin_provider_ops_action_error("not_configured", "checkin", message, None) + } + }; let method = admin_provider_ops_request_method(action_config, "POST"); let (status, response_json) = match admin_provider_ops_execute_json_request( state, @@ -129,7 +134,7 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action( ); } - let message = response_json + let upstream_message = response_json .get("message") .and_then(serde_json::Value::as_str) .unwrap_or_default() @@ -143,23 +148,23 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action( return admin_provider_ops_action_response( "success", "checkin", - admin_provider_ops_checkin_payload(&response_json, Some(message)), + admin_provider_ops_checkin_payload(&response_json, Some("签到成功".to_string())), None, response_time_ms, 3600, ); } - if admin_provider_ops_checkin_already_done(&message) { + if admin_provider_ops_checkin_already_done(&upstream_message) { return admin_provider_ops_action_response( "already_done", "checkin", - admin_provider_ops_checkin_payload(&response_json, Some(message)), + admin_provider_ops_checkin_payload(&response_json, Some("今日已签到".to_string())), None, response_time_ms, 3600, ); } - if admin_provider_ops_checkin_auth_failure(&message) { + if admin_provider_ops_checkin_auth_failure(&upstream_message) { return admin_provider_ops_action_error( if has_cookie { "auth_expired" @@ -167,28 +172,15 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action( "auth_failed" }, "checkin", - if message.is_empty() { - if has_cookie { - "Cookie 已失效" - } else { - "认证失败" - } + if has_cookie { + "Cookie 已失效" } else { - message.as_str() + "认证失败" }, response_time_ms, ); } - admin_provider_ops_action_error( - "unknown_error", - "checkin", - if message.is_empty() { - "签到失败" - } else { - message.as_str() - }, - response_time_ms, - ) + admin_provider_ops_action_error("unknown_error", "checkin", "签到失败", response_time_ms) } fn admin_provider_ops_network_error_message(error: &str) -> String { @@ -197,5 +189,5 @@ fn admin_provider_ops_network_error_message(error: &str) -> String { if lower.contains("timeout") || normalized.contains("超时") { return "请求超时".to_string(); } - format!("网络错误: {normalized}") + "网络错误".to_string() } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/shared.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/shared.rs index 273cefd5d..aa1bd12c4 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/checkin/shared.rs @@ -34,7 +34,7 @@ pub(super) fn admin_provider_ops_checkin_auth_failure(message: &str) -> bool { pub(super) fn admin_provider_ops_checkin_payload( response_json: &serde_json::Value, - fallback_message: Option, + message: Option, ) -> serde_json::Value { let details = response_json .get("data") @@ -54,30 +54,45 @@ pub(super) fn admin_provider_ops_checkin_payload( let next_reward = details.and_then(|value| { admin_provider_ops_value_as_f64(value.get("next_reward").or_else(|| value.get("next"))) }); - let message = fallback_message.or_else(|| { - response_json - .get("message") - .and_then(serde_json::Value::as_str) - .map(ToOwned::to_owned) - }); - let mut extra = serde_json::Map::new(); - if let Some(details) = details { - for (key, value) in details { - if matches!( - key.as_str(), - "reward" - | "quota" - | "amount" - | "streak_days" - | "streak" - | "next_reward" - | "next" - | "message" - ) { - continue; - } - extra.insert(key.clone(), value.clone()); - } + admin_provider_ops_checkin_data( + reward, + streak_days, + next_reward, + message, + serde_json::Map::new(), + ) +} + +#[cfg(test)] +mod tests { + use super::admin_provider_ops_checkin_payload; + use serde_json::json; + + #[test] + fn checkin_payload_keeps_metrics_without_copying_upstream_secrets() { + let payload = admin_provider_ops_checkin_payload( + &json!({ + "success": true, + "message": "authorization=Bearer upstream-secret", + "data": { + "reward": 1.5, + "streak_days": 3, + "next_reward": 2.0, + "api_key": "secret-api-key", + "profile": {"access_token": "secret-token"} + } + }), + Some("签到成功".to_string()), + ); + + assert_eq!(payload["reward"], json!(1.5)); + assert_eq!(payload["streak_days"], json!(3)); + assert_eq!(payload["next_reward"], json!(2.0)); + assert_eq!(payload["message"], json!("签到成功")); + assert_eq!(payload["extra"], json!({})); + let serialized = payload.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("secret-api-key")); + assert!(!serialized.contains("secret-token")); } - admin_provider_ops_checkin_data(reward, streak_days, next_reward, message, extra) } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/mod.rs index 0f2ecb056..6f9e62ff5 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/mod.rs @@ -5,7 +5,7 @@ mod support; use super::config::{ admin_provider_ops_config_object, admin_provider_ops_connector_object, - admin_provider_ops_decrypted_credentials, resolve_admin_provider_ops_base_url, + admin_provider_ops_credential_snapshot, }; use super::support::ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE; use super::verify::{ @@ -36,13 +36,23 @@ pub(crate) async fn admin_provider_ops_local_action_response( state: &AdminAppState<'_>, provider_id: &str, provider: Option<&StoredProviderCatalogProvider>, - endpoints: &[StoredProviderCatalogEndpoint], + _endpoints: &[StoredProviderCatalogEndpoint], action_type: &str, request_config: Option<&serde_json::Map>, ) -> serde_json::Value { let Some(provider) = provider else { return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置"); }; + let credential_snapshot = match admin_provider_ops_credential_snapshot(state, provider).await { + Ok(snapshot) => snapshot, + Err(_) => { + return responses::admin_provider_ops_action_not_configured( + action_type, + "已保存的 Provider Ops 凭据无法解密或迁移", + ) + } + }; + let provider = &credential_snapshot.provider; let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else { return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置"); }; @@ -57,14 +67,11 @@ pub(crate) async fn admin_provider_ops_local_action_response( ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE, ); }; - let Some(base_url) = - resolve_admin_provider_ops_base_url(provider, endpoints, Some(provider_ops_config)) - else { - return responses::admin_provider_ops_action_not_configured( - action_type, - "Provider 未配置 base_url", - ); - }; + let base_url = credential_snapshot + .binding + .destination + .base_url() + .to_string(); let mut connector_config = admin_provider_ops_connector_object(provider_ops_config) .and_then(|connector| connector.get("config")) @@ -84,12 +91,7 @@ pub(crate) async fn admin_provider_ops_local_action_response( let proxy_snapshot = admin_provider_ops_resolve_proxy_snapshot(state, Some(&connector_config)).await; - let credentials = admin_provider_ops_decrypted_credentials( - state, - admin_provider_ops_config_object(provider) - .and_then(admin_provider_ops_connector_object) - .and_then(|connector| connector.get("credentials")), - ); + let credentials = credential_snapshot.credentials; let headers = match build_headers( architecture.architecture_id, &connector_config, diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/mod.rs index 36c312c0f..25a2494a7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/mod.rs @@ -73,7 +73,17 @@ pub(super) async fn admin_provider_ops_run_query_balance_action( } let start = std::time::Instant::now(); - let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/balance"); + let url = match admin_provider_ops_request_url(base_url, action_config, "/api/user/balance") { + Ok(url) => url, + Err(message) => { + return admin_provider_ops_action_error( + "not_configured", + "query_balance", + message, + None, + ) + } + }; let method = admin_provider_ops_request_method(action_config, "GET"); let (status, response_json) = match admin_provider_ops_execute_json_request( state, @@ -200,5 +210,5 @@ fn admin_provider_ops_network_error_message(error: &str) -> String { if lower.contains("timeout") || normalized.contains("超时") { return "请求超时".to_string(); } - format!("网络错误: {normalized}") + "网络错误".to_string() } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/sub2api.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/sub2api.rs index 6e03eb8b0..f1adffd9f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/sub2api.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/sub2api.rs @@ -63,7 +63,17 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload( .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or("/api/v1/auth/me?timezone=Asia/Shanghai"); - let me_url = admin_provider_ops_sub2api_request_url(base_url, me_endpoint); + let me_url = match admin_provider_ops_sub2api_request_url(base_url, me_endpoint) { + Ok(url) => url, + Err(message) => { + return admin_provider_ops_action_error( + "not_configured", + "query_balance", + message, + None, + ) + } + }; let subscription_endpoint = admin_provider_ops_json_object_map(json!({ "endpoint": action_config .get("subscription_endpoint") @@ -77,7 +87,17 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload( .unwrap_or("/api/v1/subscriptions/summary") .to_string(); let subscription_url = - admin_provider_ops_sub2api_request_url(base_url, subscription_endpoint.as_str()); + match admin_provider_ops_sub2api_request_url(base_url, subscription_endpoint.as_str()) { + Ok(url) => url, + Err(message) => { + return admin_provider_ops_action_error( + "not_configured", + "query_balance", + message, + None, + ) + } + }; let auth_value = match reqwest::header::HeaderValue::from_str(&format!("Bearer {access_token}")) { @@ -197,8 +217,5 @@ fn network_error_message(error: &str) -> String { if lower.contains("timeout") || normalized.contains("超时") { return "请求超时".to_string(); } - if normalized.starts_with("网络错误:") { - return normalized.to_string(); - } - format!("网络错误: {normalized}") + "网络错误".to_string() } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/yescode.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/yescode.rs index 7406c91ac..d683c5a01 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/yescode.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/query_balance/yescode.rs @@ -5,6 +5,9 @@ use super::super::responses::{ admin_provider_ops_action_error, admin_provider_ops_action_response, }; use crate::handlers::admin::request::AdminAppState; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url, +}; use aether_admin::provider::ops::parse_yescode_combined_balance_payload; use aether_contracts::ProxySnapshot; use serde_json::json; @@ -17,8 +20,44 @@ pub(super) async fn admin_provider_ops_yescode_balance_payload( proxy_snapshot: Option<&ProxySnapshot>, ) -> serde_json::Value { let start = std::time::Instant::now(); - let balance_url = format!("{}/api/v1/user/balance", base_url.trim_end_matches('/')); - let profile_url = format!("{}/api/v1/auth/profile", base_url.trim_end_matches('/')); + // Resolve fixed action paths against a canonical origin. Avoid string + // concatenation so malformed bases (or path-like inputs) cannot redirect + // credentials to another host. + let destination = match canonicalize_provider_ops_base_url(base_url) { + Ok(destination) => destination, + Err(_) => { + return admin_provider_ops_action_error( + "auth_failed", + "query_balance", + "Cookie 已失效,请重新配置", + Some(start.elapsed().as_millis() as u64), + ); + } + }; + let balance_url = + match resolve_provider_ops_same_origin_url(&destination, "/api/v1/user/balance") { + Ok(url) => url, + Err(_) => { + return admin_provider_ops_action_error( + "auth_failed", + "query_balance", + "Cookie 已失效,请重新配置", + Some(start.elapsed().as_millis() as u64), + ); + } + }; + let profile_url = + match resolve_provider_ops_same_origin_url(&destination, "/api/v1/auth/profile") { + Ok(url) => url, + Err(_) => { + return admin_provider_ops_action_error( + "auth_failed", + "query_balance", + "Cookie 已失效,请重新配置", + Some(start.elapsed().as_millis() as u64), + ); + } + }; let (balance_result, profile_result) = tokio::join!( admin_provider_ops_execute_json_request( state, diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/support.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/support.rs index 468c63f82..669e2d220 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/support.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/actions/support.rs @@ -1,5 +1,9 @@ use serde_json::json; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url, +}; + pub(super) fn admin_provider_ops_checkin_data( reward: Option, streak_days: Option, @@ -26,18 +30,15 @@ pub(super) fn admin_provider_ops_request_url( base_url: &str, action_config: &serde_json::Map, default_endpoint: &str, -) -> String { +) -> Result { let endpoint = action_config .get("endpoint") .and_then(serde_json::Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or(default_endpoint); - if endpoint.starts_with("http://") || endpoint.starts_with("https://") { - endpoint.to_string() - } else { - format!("{}{}", base_url.trim_end_matches('/'), endpoint) - } + let destination = canonicalize_provider_ops_base_url(base_url).map_err(ToString::to_string)?; + resolve_provider_ops_same_origin_url(&destination, endpoint).map_err(ToString::to_string) } pub(super) fn admin_provider_ops_request_method( diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/balance_cache.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/balance_cache.rs index 836fd1d1d..2e64b08a8 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/balance_cache.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/balance_cache.rs @@ -1,7 +1,7 @@ use super::actions::admin_provider_ops_local_action_response; use crate::handlers::admin::request::AdminAppState; use crate::task_runtime::{spawn_fire_and_forget, TASK_KEY_PROVIDER_BALANCE_REFRESH}; -use serde_json::{json, Value}; +use serde_json::{json, Map, Value}; use std::collections::HashSet; use std::time::Duration; use tokio::sync::{Mutex, Semaphore}; @@ -75,10 +75,13 @@ pub(crate) async fn store_admin_provider_ops_balance_cache( provider_id: &str, payload: &Value, ) { - let Some(ttl_seconds) = balance_cache_ttl_seconds(payload) else { + let Some(projected) = project_admin_provider_ops_balance_cache_payload(payload) else { return; }; - let serialized = match serde_json::to_string(payload) { + let Some(ttl_seconds) = balance_cache_ttl_seconds(&projected) else { + return; + }; + let serialized = match serde_json::to_string(&projected) { Ok(serialized) => serialized, Err(err) => { warn!( @@ -222,6 +225,284 @@ fn balance_cache_ttl_seconds(payload: &Value) -> Option { } } +const BALANCE_CACHE_EXTRA_NUMERIC_FIELDS: &[&str] = &[ + "balance", + "points", + "active_subscriptions", + "total_used_usd", + "normal_balance", + "subscription_balance", + "charity_balance", + "pay_as_you_go_balance", + "daily_limit", + "weekly_limit", + "weekly_spent", + "daily_spent", + "daily_used_quota", + "daily_quota_limit", + "daily_remaining_quota", +]; + +const BALANCE_CACHE_EXTRA_STRING_FIELDS: &[&str] = &[ + "plan_name", + "subscription_status", + "status", + "group_name", + "effective_start_date", + "effective_end_date", +]; + +const BALANCE_CACHE_EXTRA_BOOL_FIELDS: &[&str] = &["checkin_success", "cookie_expired"]; + +const BALANCE_CACHE_EXTRA_NESTED_FIELDS: &[&str] = &[ + "five_hour_limit", + "weekly_limit", + "month_stats", + "subscriptions", +]; +const BALANCE_CACHE_LIMIT_FIELDS: &[&str] = &["limit", "used", "remaining", "resets_at"]; +const BALANCE_CACHE_MONTH_STATS_FIELDS: &[&str] = &[ + "total_input_tokens", + "total_output_tokens", + "total_quota", + "total_requests", +]; + +fn project_admin_provider_ops_balance_cache_payload(payload: &Value) -> Option { + let source = payload.as_object()?; + let status = source.get("status").and_then(Value::as_str)?.trim(); + if !matches!(status, "success" | "auth_expired" | "auth_failed") { + return None; + } + if source.get("action_type").and_then(Value::as_str) != Some("query_balance") { + return None; + } + + let mut projected = Map::new(); + projected.insert("status".to_string(), Value::String(status.to_string())); + projected.insert( + "action_type".to_string(), + Value::String("query_balance".to_string()), + ); + + let data = match source.get("data") { + Some(Value::Null) | None => Value::Null, + Some(value) => project_admin_provider_ops_balance_data(value)?, + }; + projected.insert("data".to_string(), data); + projected.insert( + "message".to_string(), + match status { + "auth_failed" => Value::String("认证失败".to_string()), + "auth_expired" => Value::String("认证已过期".to_string()), + _ => Value::Null, + }, + ); + if let Some(value) = source + .get("executed_at") + .and_then(project_admin_provider_ops_safe_string) + { + projected.insert("executed_at".to_string(), Value::String(value)); + } + if let Some(value) = source + .get("response_time_ms") + .and_then(project_admin_provider_ops_finite_number) + { + projected.insert("response_time_ms".to_string(), value); + } + projected.insert( + "cache_ttl_seconds".to_string(), + Value::from(if status == "auth_failed" { + ADMIN_PROVIDER_OPS_BALANCE_AUTH_FAILED_CACHE_TTL_SECS + } else { + ADMIN_PROVIDER_OPS_BALANCE_CACHE_TTL_SECS + }), + ); + Some(Value::Object(projected)) +} + +fn project_admin_provider_ops_balance_data(value: &Value) -> Option { + let source = value.as_object()?; + let mut projected = Map::new(); + for field in ["total_granted", "total_used", "total_available"] { + if let Some(value) = source.get(field) { + projected.insert( + field.to_string(), + project_admin_provider_ops_finite_number_or_null(value)?, + ); + } + } + if let Some(value) = source.get("expires_at") { + projected.insert( + "expires_at".to_string(), + project_admin_provider_ops_finite_number_or_null(value)?, + ); + } + if let Some(value) = source.get("currency") { + let currency = project_admin_provider_ops_safe_string(value)?; + if currency.len() > 32 + || !currency + .chars() + .all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.' | '/')) + { + return None; + } + projected.insert("currency".to_string(), Value::String(currency)); + } + if let Some(extra) = source.get("extra") { + projected.insert( + "extra".to_string(), + project_admin_provider_ops_balance_extra(extra)?, + ); + } + Some(Value::Object(projected)) +} + +fn project_admin_provider_ops_balance_extra(value: &Value) -> Option { + let source = value.as_object()?; + let mut projected = Map::new(); + for (field, value) in source { + let projected_value = if BALANCE_CACHE_EXTRA_NUMERIC_FIELDS.contains(&field.as_str()) + && project_admin_provider_ops_finite_number(value).is_some() + { + project_admin_provider_ops_finite_number(value) + } else if BALANCE_CACHE_EXTRA_STRING_FIELDS.contains(&field.as_str()) { + project_admin_provider_ops_safe_string(value).map(Value::String) + } else if BALANCE_CACHE_EXTRA_BOOL_FIELDS.contains(&field.as_str()) { + value.as_bool().map(Value::Bool) + } else if BALANCE_CACHE_EXTRA_NESTED_FIELDS.contains(&field.as_str()) { + project_admin_provider_ops_balance_extra_nested(field, value) + } else if matches!( + field.as_str(), + "weekly_resets_at" | "daily_resets_at" | "resets_at" + ) { + project_admin_provider_ops_finite_number_or_safe_string(value) + } else if matches!(field.as_str(), "checkin_message" | "cookie_expired_message") { + project_admin_provider_ops_safe_string(value).map(Value::String) + } else { + None + }; + if let Some(projected_value) = projected_value { + projected.insert(field.clone(), projected_value); + } + } + Some(Value::Object(projected)) +} + +fn project_admin_provider_ops_balance_extra_nested(field: &str, value: &Value) -> Option { + if field == "subscriptions" { + let items = value.as_array()?; + return Some(Value::Array( + items + .iter() + .take(128) + .filter_map(project_admin_provider_ops_subscription) + .collect(), + )); + } + let source = value.as_object()?; + let mut projected = Map::new(); + let allowed = if field == "month_stats" { + BALANCE_CACHE_MONTH_STATS_FIELDS + } else { + BALANCE_CACHE_LIMIT_FIELDS + }; + for key in allowed { + if let Some(value) = source.get(*key) { + let projected_value = if *key == "resets_at" { + project_admin_provider_ops_finite_number_or_safe_string(value) + } else { + project_admin_provider_ops_finite_number(value) + }; + if let Some(projected_value) = projected_value { + projected.insert((*key).to_string(), projected_value); + } + } + } + Some(Value::Object(projected)) +} + +fn project_admin_provider_ops_subscription(value: &Value) -> Option { + let source = value.as_object()?; + let mut projected = Map::new(); + for field in ["group_name", "status"] { + if let Some(value) = source.get(field) { + projected.insert( + field.to_string(), + Value::String(project_admin_provider_ops_safe_string(value)?), + ); + } + } + for field in [ + "daily_used_usd", + "daily_limit_usd", + "weekly_used_usd", + "weekly_limit_usd", + "monthly_used_usd", + "monthly_limit_usd", + ] { + if let Some(value) = source.get(field) { + if let Some(value) = project_admin_provider_ops_finite_number(value) { + projected.insert(field.to_string(), value); + } + } + } + if let Some(value) = source.get("expires_at") { + if let Some(value) = project_admin_provider_ops_finite_number_or_safe_string(value) { + projected.insert("expires_at".to_string(), value); + } + } + Some(Value::Object(projected)) +} + +fn project_admin_provider_ops_finite_number(value: &Value) -> Option { + if let Some(number) = value.as_f64() { + return number.is_finite().then(|| value.clone()); + } + let number = value.as_str()?.trim().parse::().ok()?; + number.is_finite().then(|| Value::from(number)) +} + +fn project_admin_provider_ops_finite_number_or_null(value: &Value) -> Option { + if value.is_null() { + Some(Value::Null) + } else { + project_admin_provider_ops_finite_number(value) + } +} + +fn project_admin_provider_ops_finite_number_or_safe_string(value: &Value) -> Option { + project_admin_provider_ops_finite_number(value) + .or_else(|| project_admin_provider_ops_safe_string(value).map(Value::String)) +} + +fn project_admin_provider_ops_safe_string(value: &Value) -> Option { + let value = value.as_str()?.trim(); + if value.is_empty() || value.len() > 256 || value.chars().any(char::is_control) { + return None; + } + let lower = value.to_ascii_lowercase(); + if [ + "authorization", + "bearer ", + "api_key", + "apikey", + "access_token", + "refresh_token", + "password", + "cookie", + "session", + "secret", + "token=", + ] + .iter() + .any(|needle| lower.contains(needle)) + { + return None; + } + Some(value.to_string()) +} + fn admin_provider_ops_balance_refresh_key(state: &AdminAppState<'_>, provider_id: &str) -> String { let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_REFRESH_PREFIX}{provider_id}"); format!( @@ -263,7 +544,10 @@ fn admin_provider_ops_action_response( #[cfg(test)] mod tests { - use super::{admin_provider_ops_pending_balance_response, balance_cache_ttl_seconds}; + use super::{ + admin_provider_ops_pending_balance_response, balance_cache_ttl_seconds, + project_admin_provider_ops_balance_cache_payload, + }; use serde_json::json; #[test] @@ -292,4 +576,65 @@ mod tests { None ); } + + #[test] + fn balance_cache_projection_drops_untrusted_messages_and_fields() { + let payload = json!({ + "status": "auth_failed", + "action_type": "query_balance", + "message": "authorization=Bearer upstream-secret", + "data": { + "total_available": 1.25, + "currency": "USD", + "extra": { + "balance": 1.0, + "access_token": "upstream-secret", + "today_stats": {"private_note": "upstream-secret"}, + "checkin_message": "签到失败" + } + }, + "cache_ttl_seconds": 999999 + }); + let projected = project_admin_provider_ops_balance_cache_payload(&payload) + .expect("known balance payload should project"); + assert_eq!(projected["message"], json!("认证失败")); + assert_eq!(projected["data"]["extra"]["balance"], json!(1.0)); + assert!(projected.to_string().find("upstream-secret").is_none()); + assert!(projected["data"]["extra"].get("access_token").is_none()); + assert!(projected["data"]["extra"].get("today_stats").is_none()); + assert_eq!(projected["cache_ttl_seconds"], json!(60)); + } + + #[test] + fn balance_cache_projection_keeps_sub2api_subscription_allowlist() { + let payload = json!({ + "status": "success", + "action_type": "query_balance", + "data": { + "total_available": 8.5, + "currency": "USD", + "extra": { + "subscriptions": [{ + "group_name": "default", + "status": "active", + "monthly_used_usd": 1.2, + "private_token": "must-drop" + }] + } + } + }); + let projected = project_admin_provider_ops_balance_cache_payload(&payload) + .expect("known balance payload should project"); + assert_eq!( + projected["data"]["extra"]["subscriptions"][0]["group_name"], + json!("default") + ); + assert_eq!( + projected["data"]["extra"]["subscriptions"][0]["monthly_used_usd"], + json!(1.2) + ); + assert!(projected["data"]["extra"]["subscriptions"][0] + .get("private_token") + .is_none()); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/config.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/config.rs index 934e11294..0ed706938 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/config.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/config.rs @@ -1,19 +1,58 @@ -use super::support::{ - AdminProviderOpsQuotaAlertConfigRequest, AdminProviderOpsSaveConfigRequest, - ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS, -}; +use super::support::{AdminProviderOpsQuotaAlertConfigRequest, AdminProviderOpsSaveConfigRequest}; use crate::handlers::admin::request::AdminAppState; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, masked_secret_display, open_provider_ops_credential, + provider_ops_credential_binding_from_config, provider_ops_credential_field_is_secret, + seal_provider_ops_credential, ProviderOpsCredentialBinding, + PROVIDER_OPS_PERSISTENT_SECRET_FIELDS, PROVIDER_OPS_TRANSIENT_METADATA_FIELDS, +}; use crate::GatewayError; use aether_admin::provider::ops as admin_provider_ops_pure; use aether_data_contracts::repository::provider_catalog::{ - StoredProviderCatalogEndpoint, StoredProviderCatalogProvider, + ProviderCatalogProviderConfigCasUpdate, StoredProviderCatalogEndpoint, + StoredProviderCatalogProvider, }; use serde_json::json; -use std::time::{SystemTime, UNIX_EPOCH}; const PROVIDER_OPS_QUOTA_ALERT_DEFAULT_FETCH_INTERVAL_SECS: u64 = 30; const PROVIDER_OPS_QUOTA_ALERT_MIN_FETCH_INTERVAL_SECS: u64 = 30; const PROVIDER_OPS_QUOTA_ALERT_MAX_FETCH_INTERVAL_SECS: u64 = 86_400; +const PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES: usize = 8; + +struct AdminProviderOpsDecodedCredentials { + values: serde_json::Map, + protected_values: serde_json::Map, + migration_required: bool, +} + +pub(crate) struct AdminProviderOpsCredentialSnapshot { + pub(crate) provider: StoredProviderCatalogProvider, + pub(crate) credentials: serde_json::Map, + pub(crate) binding: ProviderOpsCredentialBinding, +} + +impl std::fmt::Debug for AdminProviderOpsCredentialSnapshot { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminProviderOpsCredentialSnapshot") + .field("provider_id", &self.provider.id) + .field("credentials", &"[REDACTED]") + .field("binding", &"[REDACTED]") + .finish_non_exhaustive() + } +} + +pub(super) struct AdminProviderOpsMergedCredentialSnapshot { + pub(super) provider: StoredProviderCatalogProvider, + pub(super) credentials: serde_json::Map, + pub(super) saved_binding: ProviderOpsCredentialBinding, + pub(super) reused_saved_secret: bool, +} + +pub(super) struct AdminProviderOpsSavedConfigSnapshot { + pub(super) provider: StoredProviderCatalogProvider, + pub(super) provider_ops_config: serde_json::Value, +} pub(super) fn admin_provider_ops_config_object( provider: &StoredProviderCatalogProvider, @@ -27,57 +66,77 @@ pub(super) fn admin_provider_ops_connector_object( admin_provider_ops_pure::admin_provider_ops_connector_object(provider_ops_config) } -fn admin_provider_ops_masked_secret( +pub(super) fn admin_provider_ops_binding_from_config( + provider_id: &str, + provider_ops_config: &serde_json::Map, + effective_base_url: &str, +) -> Result { + provider_ops_credential_binding_from_config( + provider_id, + provider_ops_config, + effective_base_url, + ) + .map_err(ToString::to_string) +} + +async fn admin_provider_ops_binding_for_provider( state: &AdminAppState<'_>, - field: &str, - ciphertext: &str, -) -> serde_json::Value { - let plaintext = state - .decrypt_catalog_secret_with_fallbacks(ciphertext) - .unwrap_or_else(|| ciphertext.to_string()); + provider: &StoredProviderCatalogProvider, +) -> Result<(ProviderOpsCredentialBinding, bool), GatewayError> { + let provider_ops_config = admin_provider_ops_config_object(provider) + .ok_or_else(|| GatewayError::Internal("Provider Ops 配置格式无效".to_string()))?; + let explicit_base_url = provider_ops_config + .get("base_url") + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + let endpoints = if explicit_base_url.is_some() { + Vec::new() + } else { + state + .list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id)) + .await? + }; + let effective_base_url = + resolve_admin_provider_ops_base_url(provider, &endpoints, Some(provider_ops_config)) + .ok_or_else(|| GatewayError::Internal("Provider Ops 未配置 base_url".to_string()))?; + let binding = admin_provider_ops_binding_from_config( + &provider.id, + provider_ops_config, + &effective_base_url, + ) + .map_err(GatewayError::Internal)?; + let needs_materialized_base_url = explicit_base_url != Some(binding.destination.base_url()); + Ok((binding, needs_materialized_base_url)) +} + +fn admin_provider_ops_masked_secret(field: &str, plaintext: &str) -> serde_json::Value { if plaintext.is_empty() { return serde_json::Value::String(String::new()); } let masked = if field == "password" { "********".to_string() - } else if plaintext.len() > 12 { - format!( - "{}****{}", - &plaintext[..4], - &plaintext[plaintext.len().saturating_sub(4)..] - ) - } else if plaintext.len() > 8 { - format!( - "{}****{}", - &plaintext[..2], - &plaintext[plaintext.len().saturating_sub(2)..] - ) } else { - "*".repeat(plaintext.len()) + masked_secret_display(plaintext, 4, 4, "****") }; serde_json::Value::String(masked) } fn admin_provider_ops_masked_credentials( - state: &AdminAppState<'_>, - raw_credentials: Option<&serde_json::Value>, + credentials: &serde_json::Map, ) -> serde_json::Value { - let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else { - return json!({}); - }; - let mut masked = serde_json::Map::new(); for (key, value) in credentials { if key.starts_with('_') { continue; } - if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) { + if provider_ops_credential_field_is_secret(key) { if let Some(ciphertext) = value.as_str().filter(|value| !value.is_empty()) { masked.insert( key.clone(), - admin_provider_ops_masked_secret(state, key, ciphertext), + admin_provider_ops_masked_secret(key, ciphertext), ); continue; } @@ -91,53 +150,71 @@ fn admin_provider_ops_is_supported_auth_type(auth_type: &str) -> bool { admin_provider_ops_pure::admin_provider_ops_is_supported_auth_type(auth_type) } -pub(super) fn admin_provider_ops_decrypted_credentials( +fn admin_provider_ops_decode_credentials( state: &AdminAppState<'_>, + binding: &ProviderOpsCredentialBinding, raw_credentials: Option<&serde_json::Value>, -) -> serde_json::Map { +) -> Result { let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else { - return serde_json::Map::new(); + return Ok(AdminProviderOpsDecodedCredentials { + values: serde_json::Map::new(), + protected_values: serde_json::Map::new(), + migration_required: false, + }); }; - let mut decrypted = serde_json::Map::new(); + let mut values = serde_json::Map::new(); + let mut protected_values = credentials.clone(); + let mut migration_required = false; for (key, value) in credentials { - if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) { - if let Some(ciphertext) = value.as_str() { - let plaintext = state - .decrypt_catalog_secret_with_fallbacks(ciphertext) - .unwrap_or_else(|| ciphertext.to_string()); - decrypted.insert(key.clone(), serde_json::Value::String(plaintext)); + if provider_ops_credential_field_is_secret(key) { + if let Some(stored_value) = value.as_str() { + if stored_value.is_empty() { + values.insert(key.clone(), value.clone()); + continue; + } + let projection = + open_provider_ops_credential(state.app(), binding, key, stored_value).map_err( + |message| format!("已保存的 Provider Ops 凭据无法解密: {message}"), + )?; + migration_required |= projection.migration_required; + protected_values + .insert(key.clone(), serde_json::Value::String(projection.protected)); + values.insert(key.clone(), serde_json::Value::String(projection.plaintext)); continue; } } - decrypted.insert(key.clone(), value.clone()); + values.insert(key.clone(), value.clone()); } - decrypted + Ok(AdminProviderOpsDecodedCredentials { + values, + protected_values, + migration_required, + }) } fn admin_provider_ops_sensitive_placeholder_or_empty(value: Option<&serde_json::Value>) -> bool { admin_provider_ops_pure::admin_provider_ops_sensitive_placeholder_or_empty(value) } -pub(super) fn admin_provider_ops_merge_credentials( +pub(super) async fn admin_provider_ops_merge_credentials( state: &AdminAppState<'_>, architecture_id: &str, provider: &StoredProviderCatalogProvider, mut request_credentials: serde_json::Map, -) -> serde_json::Map { - let mut saved_credentials = admin_provider_ops_decrypted_credentials( - state, - admin_provider_ops_config_object(provider) - .and_then(admin_provider_ops_connector_object) - .and_then(|connector| connector.get("credentials")), - ); +) -> Result { + let snapshot = admin_provider_ops_credential_snapshot(state, provider) + .await + .map_err(|_| "已保存的 Provider Ops 凭据无法解密或迁移".to_string())?; + let mut saved_credentials = snapshot.credentials; let preserve_internal_runtime_fields = admin_provider_ops_pure::normalize_architecture_id(architecture_id) == "sub2api"; if !preserve_internal_runtime_fields { saved_credentials.retain(|key, _| !key.starts_with('_')); } - for field in ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS { + let mut reused_saved_secret = false; + for field in PROVIDER_OPS_PERSISTENT_SECRET_FIELDS { if field.starts_with('_') { continue; } @@ -146,6 +223,7 @@ pub(super) fn admin_provider_ops_merge_credentials( { if let Some(saved_value) = saved_credentials.get(*field) { request_credentials.insert((*field).to_string(), saved_value.clone()); + reused_saved_secret = true; } } } @@ -158,23 +236,29 @@ pub(super) fn admin_provider_ops_merge_credentials( } } - request_credentials + Ok(AdminProviderOpsMergedCredentialSnapshot { + provider: snapshot.provider, + credentials: request_credentials, + saved_binding: snapshot.binding, + reused_saved_secret, + }) } fn admin_provider_ops_encrypt_credentials( state: &AdminAppState<'_>, + binding: &ProviderOpsCredentialBinding, credentials: serde_json::Map, ) -> Result, String> { let mut encrypted = serde_json::Map::new(); for (key, value) in credentials { - if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) { + if provider_ops_credential_field_is_secret(&key) { if let Some(plaintext) = value.as_str() { if plaintext.is_empty() { encrypted.insert(key, value); } else { - let ciphertext = state - .encrypt_catalog_secret_with_fallbacks(plaintext) - .ok_or_else(|| "gateway 未配置 Provider Ops 加密密钥".to_string())?; + let ciphertext = + seal_provider_ops_credential(state.app(), binding, &key, plaintext) + .map_err(ToString::to_string)?; encrypted.insert(key, serde_json::Value::String(ciphertext)); } continue; @@ -185,6 +269,104 @@ fn admin_provider_ops_encrypt_credentials( Ok(encrypted) } +fn admin_provider_ops_config_with_credentials( + provider: &StoredProviderCatalogProvider, + credentials: serde_json::Map, + binding: &ProviderOpsCredentialBinding, +) -> Result, String> { + let mut provider_config = provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .cloned() + .ok_or_else(|| "Provider Ops 配置格式无效".to_string())?; + let mut provider_ops_config = provider_config + .get("provider_ops") + .and_then(serde_json::Value::as_object) + .cloned() + .ok_or_else(|| "Provider Ops 配置格式无效".to_string())?; + let mut connector_config = provider_ops_config + .get("connector") + .and_then(serde_json::Value::as_object) + .cloned() + .ok_or_else(|| "Provider Ops connector 配置格式无效".to_string())?; + + connector_config.insert( + "credentials".to_string(), + serde_json::Value::Object(credentials), + ); + provider_ops_config.insert( + "connector".to_string(), + serde_json::Value::Object(connector_config), + ); + provider_ops_config.insert( + "base_url".to_string(), + serde_json::Value::String(binding.destination.base_url().to_string()), + ); + provider_config.insert( + "provider_ops".to_string(), + serde_json::Value::Object(provider_ops_config), + ); + Ok(Some(serde_json::Value::Object(provider_config))) +} + +pub(crate) async fn admin_provider_ops_credential_snapshot( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, +) -> Result { + let mut current = provider.clone(); + for _ in 0..PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES { + let (binding, needs_materialized_base_url) = + admin_provider_ops_binding_for_provider(state, ¤t).await?; + let raw_credentials = admin_provider_ops_config_object(¤t) + .and_then(admin_provider_ops_connector_object) + .and_then(|connector| connector.get("credentials")); + let decoded = admin_provider_ops_decode_credentials(state, &binding, raw_credentials) + .map_err(GatewayError::Internal)?; + if !decoded.migration_required && !needs_materialized_base_url { + return Ok(AdminProviderOpsCredentialSnapshot { + provider: current, + credentials: decoded.values, + binding, + }); + } + + let migrated_config = admin_provider_ops_config_with_credentials( + ¤t, + decoded.protected_values, + &binding, + ) + .map_err(GatewayError::Internal)?; + let update = ProviderCatalogProviderConfigCasUpdate { + provider_id: current.id.clone(), + expected_config: current.config.clone(), + config: migrated_config.clone(), + }; + if state + .compare_and_swap_provider_catalog_provider_config(&update) + .await? + { + current.config = migrated_config; + return Ok(AdminProviderOpsCredentialSnapshot { + provider: current, + credentials: decoded.values, + binding, + }); + } + + current = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(¤t.id)) + .await? + .into_iter() + .next() + .ok_or_else(|| GatewayError::Internal("Provider Ops Provider 不存在".to_string()))?; + } + + Err(GatewayError::Internal( + "Provider Ops 凭据迁移未能稳定完成".to_string(), + )) +} + pub(super) async fn persist_admin_provider_ops_runtime_credentials( state: &AdminAppState<'_>, provider: &StoredProviderCatalogProvider, @@ -193,82 +375,91 @@ pub(super) async fn persist_admin_provider_ops_runtime_credentials( if updated_credentials.is_empty() || !state.has_provider_catalog_data_writer() { return Ok(None); } - - let mut updated_provider = provider.clone(); - let mut provider_config = updated_provider - .config - .as_ref() - .and_then(serde_json::Value::as_object) - .cloned() - .unwrap_or_default(); - let Some(provider_ops_config) = provider_config - .get("provider_ops") - .and_then(serde_json::Value::as_object) - .cloned() - else { - return Ok(None); - }; - let Some(connector_config) = provider_ops_config - .get("connector") - .and_then(serde_json::Value::as_object) - .cloned() - else { - return Ok(None); - }; - - let mut decrypted_credentials = - admin_provider_ops_decrypted_credentials(state, connector_config.get("credentials")); - for (key, value) in updated_credentials { - decrypted_credentials.insert(key.clone(), value.clone()); + for key in updated_credentials.keys() { + if key != "refresh_token" + && key != "_cached_access_token" + && !PROVIDER_OPS_TRANSIENT_METADATA_FIELDS.contains(&key.as_str()) + { + return Err(GatewayError::Internal(format!( + "不允许持久化未知的 Provider Ops runtime credential 字段 '{key}'" + ))); + } } - let encrypted_credentials = - admin_provider_ops_encrypt_credentials(state, decrypted_credentials) - .map_err(GatewayError::Internal)?; - let mut updated_connector = connector_config.clone(); - updated_connector.insert( - "credentials".to_string(), - serde_json::Value::Object(encrypted_credentials), - ); - - let mut updated_provider_ops = provider_ops_config.clone(); - updated_provider_ops.insert( - "connector".to_string(), - serde_json::Value::Object(updated_connector), - ); - - provider_config.insert( - "provider_ops".to_string(), - serde_json::Value::Object(updated_provider_ops), - ); - updated_provider.config = Some(serde_json::Value::Object(provider_config)); - updated_provider.updated_at_unix_secs = SystemTime::now() - .duration_since(UNIX_EPOCH) - .ok() - .map(|duration| duration.as_secs()); - - state - .update_provider_catalog_provider(&updated_provider) - .await + let mut current = provider.clone(); + for _ in 0..PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES { + let snapshot = admin_provider_ops_credential_snapshot(state, ¤t).await?; + let mut decrypted_credentials = snapshot.credentials; + for (key, value) in updated_credentials { + decrypted_credentials.insert(key.clone(), value.clone()); + } + let encrypted_credentials = + admin_provider_ops_encrypt_credentials(state, &snapshot.binding, decrypted_credentials) + .map_err(GatewayError::Internal)?; + let config = admin_provider_ops_config_with_credentials( + &snapshot.provider, + encrypted_credentials, + &snapshot.binding, + ) + .map_err(GatewayError::Internal)?; + let update = ProviderCatalogProviderConfigCasUpdate { + provider_id: snapshot.provider.id.clone(), + expected_config: snapshot.provider.config.clone(), + config, + }; + if state + .compare_and_swap_provider_catalog_provider_config(&update) + .await? + { + return Ok(state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider.id)) + .await? + .into_iter() + .next()); + } + current = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider.id)) + .await? + .into_iter() + .next() + .ok_or_else(|| GatewayError::Internal("Provider Ops Provider 不存在".to_string()))?; + } + Err(GatewayError::Internal( + "Provider Ops runtime credential 并发更新未能稳定完成".to_string(), + )) } -pub(super) fn build_admin_provider_ops_saved_config_value( +pub(super) async fn build_admin_provider_ops_saved_config_value( state: &AdminAppState<'_>, provider: &StoredProviderCatalogProvider, payload: AdminProviderOpsSaveConfigRequest, -) -> Result { +) -> Result { + let architecture_id = payload.architecture_id.trim(); + let normalized_architecture_id = + admin_provider_ops_pure::normalize_architecture_id(architecture_id); + if architecture_id.is_empty() || architecture_id != normalized_architecture_id { + return Err("architecture_id 必须是合法的 Provider Ops 架构".to_string()); + } let auth_type = payload.connector.auth_type.trim().to_string(); if auth_type.is_empty() || !admin_provider_ops_is_supported_auth_type(auth_type.as_str()) { return Err("connector.auth_type 必须是合法的认证类型".to_string()); } - let merged_credentials = admin_provider_ops_merge_credentials( + let merged = admin_provider_ops_merge_credentials( state, - payload.architecture_id.as_str(), + normalized_architecture_id, provider, payload.connector.credentials, - ); - let encrypted_credentials = admin_provider_ops_encrypt_credentials(state, merged_credentials)?; + ) + .await?; + let canonical_base_url = payload + .base_url + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| merged.saved_binding.destination.base_url()); + let canonical_destination = + canonicalize_provider_ops_base_url(canonical_base_url).map_err(ToString::to_string)?; let actions = payload .actions @@ -284,19 +475,48 @@ pub(super) fn build_admin_provider_ops_saved_config_value( }) .collect::>(); let quota_alert = normalize_admin_provider_ops_quota_alert(payload.quota_alert)?; - - Ok(json!({ - "architecture_id": payload.architecture_id, - "base_url": payload.base_url, + let mut provider_ops_config = json!({ + "architecture_id": normalized_architecture_id, + "base_url": canonical_destination.base_url(), "connector": { "auth_type": auth_type, "config": payload.connector.config, - "credentials": encrypted_credentials, + "credentials": {}, }, "actions": actions, "schedule": payload.schedule, "quota_alert": quota_alert, - })) + }); + let new_binding = admin_provider_ops_binding_from_config( + &merged.provider.id, + provider_ops_config + .as_object() + .ok_or_else(|| "Provider Ops 配置格式无效".to_string())?, + canonical_destination.base_url(), + )?; + let same_secret_destination = merged.saved_binding.provider_id == new_binding.provider_id + && merged.saved_binding.architecture_id == new_binding.architecture_id + && merged.saved_binding.auth_type == new_binding.auth_type + && merged.saved_binding.destination == new_binding.destination; + if merged.reused_saved_secret && !same_secret_destination { + return Err("修改 Provider Ops 架构、认证类型或目标地址时必须重新填写凭据".to_string()); + } + let mut merged_credentials = merged.credentials; + if merged.saved_binding != new_binding { + for field in PROVIDER_OPS_TRANSIENT_METADATA_FIELDS { + merged_credentials.remove(*field); + } + merged_credentials.retain(|field, _| !field.starts_with("_cached_")); + } + let encrypted_credentials = + admin_provider_ops_encrypt_credentials(state, &new_binding, merged_credentials)?; + provider_ops_config["connector"]["credentials"] = + serde_json::Value::Object(encrypted_credentials); + + Ok(AdminProviderOpsSavedConfigSnapshot { + provider: merged.provider, + provider_ops_config, + }) } fn normalize_admin_provider_ops_quota_alert( @@ -356,27 +576,35 @@ pub(super) fn build_admin_provider_ops_status_payload( admin_provider_ops_pure::build_admin_provider_ops_status_payload(provider_id, provider) } -pub(super) fn build_admin_provider_ops_config_payload( +pub(super) async fn build_admin_provider_ops_config_payload( state: &AdminAppState<'_>, provider_id: &str, provider: Option<&StoredProviderCatalogProvider>, endpoints: &[StoredProviderCatalogEndpoint], -) -> serde_json::Value { +) -> Result { let Some(provider) = provider else { - return json!({ + return Ok(json!({ "provider_id": provider_id, "is_configured": false, - }); + })); }; - let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else { - return json!({ + if admin_provider_ops_config_object(provider).is_none() { + return Ok(json!({ "provider_id": provider_id, "is_configured": false, - }); + })); + } + let snapshot = admin_provider_ops_credential_snapshot(state, provider).await?; + let provider = &snapshot.provider; + let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else { + return Ok(json!({ + "provider_id": provider_id, + "is_configured": false, + })); }; let connector = admin_provider_ops_connector_object(provider_ops_config); - json!({ + Ok(json!({ "provider_id": provider_id, "is_configured": true, "architecture_id": provider_ops_config @@ -398,15 +626,159 @@ pub(super) fn build_admin_provider_ops_config_payload( .filter(|value| value.is_object()) .cloned() .unwrap_or_else(|| json!({})), - "credentials": admin_provider_ops_masked_credentials( - state, - connector.and_then(|connector| connector.get("credentials")), - ), + "credentials": admin_provider_ops_masked_credentials(&snapshot.credentials), }, "quota_alert": provider_ops_config .get("quota_alert") .filter(|value| value.is_object()) .cloned() .unwrap_or_else(default_admin_provider_ops_quota_alert), - }) + })) +} + +#[cfg(test)] +mod tests { + use super::{admin_provider_ops_credential_snapshot, open_provider_ops_credential}; + use crate::data::GatewayDataState; + use crate::handlers::admin::request::AdminAppState; + use crate::AppState; + use aether_crypto::{ + decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, + looks_like_python_fernet_ciphertext, DEVELOPMENT_ENCRYPTION_KEY, + }; + use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogReadRepository, StoredProviderCatalogProvider, + }; + use serde_json::json; + use std::sync::Arc; + + const TEST_PROVIDER_ID: &str = "provider-ops-secret-test"; + const TEST_API_KEY: &str = "legacy-provider-ops-api-key"; + + fn provider_with_api_key(api_key: &str) -> StoredProviderCatalogProvider { + StoredProviderCatalogProvider::new( + TEST_PROVIDER_ID.to_string(), + "Provider Ops Secret Test".to_string(), + None, + "openai".to_string(), + ) + .expect("provider should build") + .with_transport_fields( + true, + false, + false, + None, + None, + None, + None, + None, + Some(json!({ + "provider_ops": { + "architecture_id": "generic_api", + "base_url": "https://provider.example.com", + "connector": { + "auth_type": "api_key", + "config": {}, + "credentials": { + "api_key": api_key, + "account_id": "account-1" + } + }, + "actions": {}, + "schedule": {} + } + })), + ) + } + + fn state_with_provider( + provider: StoredProviderCatalogProvider, + ) -> (AppState, Arc) { + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + Vec::new(), + )); + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests(repository.clone()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + (state, repository) + } + + async fn stored_provider( + repository: &InMemoryProviderCatalogReadRepository, + ) -> StoredProviderCatalogProvider { + repository + .list_providers_by_ids(&[TEST_PROVIDER_ID.to_string()]) + .await + .expect("provider should read") + .into_iter() + .next() + .expect("provider should exist") + } + + #[tokio::test] + async fn legacy_provider_ops_credentials_are_lazily_migrated() { + let provider = provider_with_api_key(TEST_API_KEY); + let (state, repository) = state_with_provider(provider.clone()); + let admin_state = AdminAppState::new(&state); + + let snapshot = admin_provider_ops_credential_snapshot(&admin_state, &provider) + .await + .expect("legacy Provider Ops credential should migrate"); + assert_eq!(snapshot.credentials["api_key"], TEST_API_KEY); + assert_eq!(snapshot.credentials["account_id"], "account-1"); + + let stored = stored_provider(repository.as_ref()).await; + let ciphertext = stored + .config + .as_ref() + .and_then(|config| config.pointer("/provider_ops/connector/credentials/api_key")) + .and_then(serde_json::Value::as_str) + .expect("stored API key should exist"); + assert_ne!(ciphertext, TEST_API_KEY); + // New migrations use a binding-aware runtime-secret envelope. Keep + // the legacy Fernet assertion below only for the tamper fixture; a + // migrated value must no longer be treated as an unbound Fernet blob. + assert!(ciphertext.starts_with("aether-provider-ops-credential-v2:")); + assert_eq!( + open_provider_ops_credential(&state, &snapshot.binding, "api_key", ciphertext) + .expect("migrated Provider Ops API key should decrypt") + .plaintext, + TEST_API_KEY + ); + } + + #[tokio::test] + async fn tampered_provider_ops_ciphertext_fails_closed() { + let mut tampered = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_API_KEY) + .expect("Provider Ops API key should encrypt"); + tampered.replace_range(tampered.len() - 2.., "AA"); + assert!(looks_like_python_fernet_ciphertext(&tampered)); + let provider = provider_with_api_key(&tampered); + let (state, repository) = state_with_provider(provider.clone()); + let admin_state = AdminAppState::new(&state); + + let error = admin_provider_ops_credential_snapshot(&admin_state, &provider) + .await + .expect_err("tampered Provider Ops ciphertext must not be used as plaintext"); + assert!(format!("{error:?}").contains("无法解密")); + + let stored = stored_provider(repository.as_ref()).await; + assert_eq!( + stored + .config + .as_ref() + .and_then(|config| { + config.pointer("/provider_ops/connector/credentials/api_key") + }) + .and_then(serde_json::Value::as_str), + Some(tampered.as_str()) + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/mod.rs index 5ad8ff84e..71e722b78 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/mod.rs @@ -5,4 +5,5 @@ mod routes; mod support; mod verify; pub(crate) use self::balance_cache::store_admin_provider_ops_balance_cache; +pub(crate) use self::config::admin_provider_ops_credential_snapshot; pub(super) use self::routes::maybe_build_local_admin_provider_ops_providers_response; diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/batch.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/batch.rs index 71f5f600b..c8314d609 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/batch.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/batch.rs @@ -15,7 +15,7 @@ use axum::{ }; use futures_util::stream::{self, StreamExt}; use serde_json::json; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; pub(super) async fn handle_admin_provider_ops_batch_balance( state: &AdminAppState<'_>, @@ -166,12 +166,48 @@ fn parse_provider_ids(body: &Bytes) -> Result, Response> { ) .into_response() })?; + let mut seen = HashSet::with_capacity(items.len()); + let mut provider_ids = Vec::with_capacity(items.len()); + for item in items { + let Some(provider_id) = item.as_str().map(str::trim) else { + return Err(( + http::StatusCode::BAD_REQUEST, + Json(json!({ "detail": "provider_ids 必须是非空字符串数组" })), + ) + .into_response()); + }; + if provider_id.is_empty() { + return Err(( + http::StatusCode::BAD_REQUEST, + Json(json!({ "detail": "provider_ids 必须是非空字符串数组" })), + ) + .into_response()); + } + if seen.insert(provider_id) { + provider_ids.push(provider_id.to_string()); + } + } + Ok(provider_ids) +} - Ok(items - .iter() - .filter_map(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .collect()) +#[cfg(test)] +mod tests { + use super::parse_provider_ids; + use axum::body::Bytes; + + #[test] + fn provider_ops_batch_ids_are_deduplicated_without_reordering() { + let body = + Bytes::from_static(br#"{"provider_ids":["provider-1"," provider-1 ","provider-2"]}"#); + assert_eq!( + parse_provider_ids(&body).expect("valid ids"), + vec!["provider-1".to_string(), "provider-2".to_string()] + ); + } + + #[test] + fn provider_ops_batch_ids_reject_non_string_entries() { + let body = Bytes::from_static(br#"{"provider_ids":["provider-1",42]}"#); + assert!(parse_provider_ids(&body).is_err()); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/config.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/config.rs index 70116f727..7d0422741 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/config.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/config.rs @@ -3,6 +3,7 @@ use super::super::config::build_admin_provider_ops_saved_config_value; use super::super::support::AdminProviderOpsSaveConfigRequest; use crate::handlers::admin::request::AdminAppState; use crate::GatewayError; +use aether_data_contracts::repository::provider_catalog::ProviderCatalogProviderConfigCasUpdate; use axum::{ body::{Body, Bytes}, http, @@ -10,7 +11,8 @@ use axum::{ Json, }; use serde_json::json; -use std::time::{SystemTime, UNIX_EPOCH}; + +const ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES: usize = 8; pub(super) async fn handle_admin_provider_ops_save_config( state: &AdminAppState<'_>, @@ -23,7 +25,7 @@ pub(super) async fn handle_admin_provider_ops_save_config( Err(response) => return Ok(Some(response)), }; let provider_ids = [provider_id.to_string()]; - let Some(existing_provider) = state + let Some(mut existing_provider) = state .read_provider_catalog_providers_by_ids(&provider_ids) .await? .into_iter() @@ -32,31 +34,57 @@ pub(super) async fn handle_admin_provider_ops_save_config( return Ok(Some(provider_not_found_response())); }; - let provider_ops_config = - match build_admin_provider_ops_saved_config_value(state, &existing_provider, payload) { - Ok(config) => config, + let mut saved = false; + for _ in 0..ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES { + let snapshot = match build_admin_provider_ops_saved_config_value( + state, + &existing_provider, + payload.clone(), + ) + .await + { + Ok(snapshot) => snapshot, Err(detail) => return Ok(Some(bad_request_detail_response(&detail))), }; - - let mut updated_provider = existing_provider.clone(); - let mut provider_config = updated_provider - .config - .as_ref() - .and_then(serde_json::Value::as_object) - .cloned() - .unwrap_or_default(); - provider_config.insert("provider_ops".to_string(), provider_ops_config); - updated_provider.config = Some(serde_json::Value::Object(provider_config)); - updated_provider.updated_at_unix_secs = SystemTime::now() - .duration_since(UNIX_EPOCH) - .ok() - .map(|duration| duration.as_secs()); - let Some(_updated) = state - .update_provider_catalog_provider(&updated_provider) - .await? - else { - return Ok(None); - }; + let mut provider_config = snapshot + .provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .cloned() + .unwrap_or_default(); + provider_config.insert("provider_ops".to_string(), snapshot.provider_ops_config); + let update = ProviderCatalogProviderConfigCasUpdate { + provider_id: snapshot.provider.id.clone(), + expected_config: snapshot.provider.config.clone(), + config: Some(serde_json::Value::Object(provider_config)), + }; + if state + .compare_and_swap_provider_catalog_provider_config(&update) + .await? + { + saved = true; + break; + } + let Some(current) = state + .read_provider_catalog_providers_by_ids(&provider_ids) + .await? + .into_iter() + .next() + else { + return Ok(Some(provider_not_found_response())); + }; + existing_provider = current; + } + if !saved { + return Ok(Some( + ( + http::StatusCode::CONFLICT, + Json(json!({ "detail": "Provider Ops 配置并发更新冲突,请重试" })), + ) + .into_response(), + )); + } clear_admin_provider_ops_balance_cache(state, provider_id).await; Ok(Some( @@ -73,7 +101,7 @@ pub(super) async fn handle_admin_provider_ops_delete_config( provider_id: &str, ) -> Result>, GatewayError> { let provider_ids = [provider_id.to_string()]; - let Some(existing_provider) = state + let Some(mut existing_provider) = state .read_provider_catalog_providers_by_ids(&provider_ids) .await? .into_iter() @@ -82,25 +110,49 @@ pub(super) async fn handle_admin_provider_ops_delete_config( return Ok(Some(provider_not_found_response())); }; - let mut updated_provider = existing_provider.clone(); - let mut provider_config = updated_provider - .config - .as_ref() - .and_then(serde_json::Value::as_object) - .cloned() - .unwrap_or_default(); - if provider_config.remove("provider_ops").is_some() { - updated_provider.config = Some(serde_json::Value::Object(provider_config)); - updated_provider.updated_at_unix_secs = SystemTime::now() - .duration_since(UNIX_EPOCH) - .ok() - .map(|duration| duration.as_secs()); - let Some(_updated) = state - .update_provider_catalog_provider(&updated_provider) - .await? - else { - return Ok(None); + let mut removed = false; + for _ in 0..ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES { + let mut provider_config = existing_provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .cloned() + .unwrap_or_default(); + if provider_config.remove("provider_ops").is_none() { + break; + } + let update = ProviderCatalogProviderConfigCasUpdate { + provider_id: existing_provider.id.clone(), + expected_config: existing_provider.config.clone(), + config: Some(serde_json::Value::Object(provider_config)), }; + if state + .compare_and_swap_provider_catalog_provider_config(&update) + .await? + { + removed = true; + break; + } + let Some(current) = state + .read_provider_catalog_providers_by_ids(&provider_ids) + .await? + .into_iter() + .next() + else { + return Ok(Some(provider_not_found_response())); + }; + existing_provider = current; + } + if admin_provider_ops_config_still_exists(&existing_provider) && !removed { + return Ok(Some( + ( + http::StatusCode::CONFLICT, + Json(json!({ "detail": "Provider Ops 配置并发更新冲突,请重试" })), + ) + .into_response(), + )); + } + if removed { clear_admin_provider_ops_balance_cache(state, provider_id).await; } @@ -113,6 +165,16 @@ pub(super) async fn handle_admin_provider_ops_delete_config( )) } +fn admin_provider_ops_config_still_exists( + provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider, +) -> bool { + provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .is_some_and(|config| config.contains_key("provider_ops")) +} + fn parse_json_object_payload(request_body: Option<&Bytes>) -> Result> where T: serde::de::DeserializeOwned, diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/connect.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/connect.rs index 6c37608b4..3be36330d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/connect.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/connect.rs @@ -1,6 +1,6 @@ use super::super::config::{ - admin_provider_ops_config_object, admin_provider_ops_connector_object, - admin_provider_ops_decrypted_credentials, resolve_admin_provider_ops_base_url, + admin_provider_ops_config_object, admin_provider_ops_credential_snapshot, + resolve_admin_provider_ops_base_url, }; use super::super::support::{ AdminProviderOpsConnectRequest, ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE, @@ -33,6 +33,9 @@ pub(super) async fn handle_admin_provider_ops_connect( else { return Ok(bad_request_detail_response("Provider 不存在")); }; + let credential_snapshot = + admin_provider_ops_credential_snapshot(state, &existing_provider).await?; + let existing_provider = credential_snapshot.provider; let Some(provider_ops_config) = admin_provider_ops_config_object(&existing_provider) else { return Ok(bad_request_detail_response("未配置操作设置")); }; @@ -52,14 +55,7 @@ pub(super) async fn handle_admin_provider_ops_connect( let actual_credentials = payload .credentials .filter(|value| !value.is_empty()) - .unwrap_or_else(|| { - admin_provider_ops_decrypted_credentials( - state, - admin_provider_ops_config_object(&existing_provider) - .and_then(admin_provider_ops_connector_object) - .and_then(|connector| connector.get("credentials")), - ) - }); + .unwrap_or(credential_snapshot.credentials); if actual_credentials.is_empty() { return Ok(bad_request_detail_response("未提供凭据")); } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/read.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/read.rs index c78ed1b95..ae433308a 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/read.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/read.rs @@ -31,7 +31,7 @@ pub(super) async fn handle_admin_provider_ops_read( let payload = if route_kind == "get_provider_status" { build_admin_provider_ops_status_payload(provider_id, provider) } else { - build_admin_provider_ops_config_payload(state, provider_id, provider, &endpoints) + build_admin_provider_ops_config_payload(state, provider_id, provider, &endpoints).await? }; let response = Json(payload).into_response(); diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/verify.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/verify.rs index 21332063a..e73b084c8 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/verify.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/routes/verify.rs @@ -1,13 +1,13 @@ use super::super::config::{ - admin_provider_ops_config_object, admin_provider_ops_merge_credentials, - resolve_admin_provider_ops_base_url, + admin_provider_ops_binding_from_config, admin_provider_ops_config_object, + admin_provider_ops_merge_credentials, resolve_admin_provider_ops_base_url, }; use super::super::support::AdminProviderOpsSaveConfigRequest; use super::super::verify::admin_provider_ops_local_verify_response; use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::attach_admin_audit_response; use crate::GatewayError; -use aether_admin::provider::ops::{admin_provider_ops_verify_failure, normalize_architecture_id}; +use aether_admin::provider::ops::admin_provider_ops_verify_failure; use axum::{ body::{Body, Bytes}, http, @@ -39,14 +39,41 @@ pub(super) async fn handle_admin_provider_ops_verify( } else { Vec::new() }; - let base_url = payload - .base_url - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + let (effective_provider, mut credentials, saved_binding, reused_saved_secret) = + match existing_provider.as_ref() { + Some(provider) if admin_provider_ops_config_object(provider).is_some() => { + match admin_provider_ops_merge_credentials( + state, + &payload.architecture_id, + provider, + payload.connector.credentials.clone(), + ) + .await + { + Ok(snapshot) => ( + Some(snapshot.provider), + snapshot.credentials, + Some(snapshot.saved_binding), + snapshot.reused_saved_secret, + ), + Err(detail) => { + return Ok(Json(admin_provider_ops_verify_failure(detail)).into_response()) + } + } + } + Some(provider) => ( + Some(provider.clone()), + payload.connector.credentials.clone(), + None, + false, + ), + None => (None, payload.connector.credentials.clone(), None, false), + }; + let fallback_base_url = saved_binding + .as_ref() + .map(|binding| binding.destination.base_url().to_string()) .or_else(|| { - existing_provider.as_ref().and_then(|provider| { + effective_provider.as_ref().and_then(|provider| { resolve_admin_provider_ops_base_url( provider, &endpoints, @@ -54,27 +81,70 @@ pub(super) async fn handle_admin_provider_ops_verify( ) }) }); + let base_url = payload + .base_url + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or(fallback_base_url); let Some(base_url) = base_url else { return Ok(Json(admin_provider_ops_verify_failure("请提供 API 地址")).into_response()); }; - - let architecture_id = normalize_architecture_id(&payload.architecture_id); - let credentials = existing_provider.as_ref().map_or_else( - || payload.connector.credentials.clone(), - |provider| { - admin_provider_ops_merge_credentials( - state, - architecture_id, - provider, - payload.connector.credentials.clone(), + let actions = payload + .actions + .iter() + .map(|(action_type, action)| { + ( + action_type.clone(), + serde_json::json!({ + "enabled": action.enabled, + "config": action.config, + }), ) + }) + .collect::>(); + let requested_provider_ops_config = serde_json::json!({ + "architecture_id": payload.architecture_id, + "base_url": base_url, + "connector": { + "auth_type": payload.connector.auth_type, + "config": payload.connector.config, + "credentials": {}, }, - ); + "actions": actions, + }); + let requested_binding = match admin_provider_ops_binding_from_config( + provider_id, + requested_provider_ops_config + .as_object() + .expect("Provider Ops verify config should be an object"), + &base_url, + ) { + Ok(binding) => binding, + Err(detail) => return Ok(Json(admin_provider_ops_verify_failure(detail)).into_response()), + }; + if let Some(saved_binding) = saved_binding.as_ref() { + let same_secret_destination = saved_binding.provider_id == requested_binding.provider_id + && saved_binding.architecture_id == requested_binding.architecture_id + && saved_binding.auth_type == requested_binding.auth_type + && saved_binding.destination == requested_binding.destination; + if reused_saved_secret && !same_secret_destination { + return Ok(Json(admin_provider_ops_verify_failure( + "验证不同的 Provider Ops 架构、认证类型或目标地址时必须重新填写凭据", + )) + .into_response()); + } + if saved_binding != &requested_binding { + credentials.retain(|field, _| !field.starts_with("_cached_")); + } + }; + let payload = admin_provider_ops_local_verify_response( state, - existing_provider.as_ref(), - &base_url, - architecture_id, + effective_provider.as_ref(), + requested_binding.destination.base_url(), + &requested_binding.architecture_id, &payload.connector.config, &credentials, ) diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/support.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/support.rs index 90c43c557..b9677300b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/support.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/support.rs @@ -3,18 +3,6 @@ use std::collections::BTreeMap; pub(super) use aether_admin::provider::ops::ProviderOpsCheckinOutcome as AdminProviderOpsCheckinOutcome; -pub(super) const ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS: &[&str] = &[ - "api_key", - "password", - "refresh_token", - "_cached_access_token", - "session_token", - "session_cookie", - "token_cookie", - "auth_cookie", - "cookie_string", - "cookie", -]; pub(super) const ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE: &str = "Provider 连接仅支持 Rust execution runtime"; pub(super) const ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE: &str = @@ -22,7 +10,7 @@ pub(super) const ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE: &str = pub(super) const ADMIN_PROVIDER_OPS_VERIFY_RUST_ONLY_MESSAGE: &str = "认证验证仅支持 Rust execution runtime"; -#[derive(Debug, Deserialize)] +#[derive(Clone, Deserialize)] pub(super) struct AdminProviderOpsSaveConfigRequest { #[serde(default = "default_admin_provider_ops_architecture_id")] pub(crate) architecture_id: String, @@ -37,7 +25,7 @@ pub(super) struct AdminProviderOpsSaveConfigRequest { pub(crate) quota_alert: Option, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Deserialize)] pub(super) struct AdminProviderOpsConnectorConfigRequest { pub(crate) auth_type: String, #[serde(default)] @@ -46,7 +34,7 @@ pub(super) struct AdminProviderOpsConnectorConfigRequest { pub(crate) credentials: serde_json::Map, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Deserialize)] pub(super) struct AdminProviderOpsActionConfigRequest { #[serde(default = "default_admin_provider_ops_action_enabled")] pub(crate) enabled: bool, @@ -54,7 +42,7 @@ pub(super) struct AdminProviderOpsActionConfigRequest { pub(crate) config: serde_json::Map, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Deserialize)] pub(super) struct AdminProviderOpsQuotaAlertConfigRequest { #[serde(default)] pub(crate) enabled: bool, @@ -64,13 +52,13 @@ pub(super) struct AdminProviderOpsQuotaAlertConfigRequest { pub(crate) fetch_interval_seconds: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(super) struct AdminProviderOpsConnectRequest { #[serde(default)] pub(crate) credentials: Option>, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(super) struct AdminProviderOpsExecuteActionRequest { #[serde(default)] pub(crate) config: Option>, diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/mod.rs index 77006b0e4..ea6c1ca0f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/mod.rs @@ -3,6 +3,9 @@ mod request; mod sub2api; use crate::handlers::admin::request::AdminAppState; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url, +}; use aether_admin::provider::ops::{ admin_provider_ops_verify_failure, build_headers, get_architecture, normalize_architecture_id, parse_verify_payload, ProviderOpsVerifyMode, @@ -68,7 +71,15 @@ pub(super) async fn admin_provider_ops_local_verify_response( Ok(headers) => headers, Err(message) => return admin_provider_ops_verify_failure(message), }; - let verify_url = format!("{base_url}{}", architecture.verify_endpoint); + let destination = match canonicalize_provider_ops_base_url(base_url) { + Ok(destination) => destination, + Err(message) => return admin_provider_ops_verify_failure(message), + }; + let verify_url = + match resolve_provider_ops_same_origin_url(&destination, architecture.verify_endpoint) { + Ok(url) => url, + Err(message) => return admin_provider_ops_verify_failure(message), + }; let (status, response_json) = match request::admin_provider_ops_execute_get_json( state, &format!("provider-ops-verify:{}", architecture.architecture_id), diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/request.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/request.rs index ef041b919..bfc93fe9f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/request.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/request.rs @@ -4,13 +4,11 @@ use aether_contracts::{ ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, }; -use base64::{engine::general_purpose::STANDARD, Engine as _}; -use flate2::read::{DeflateDecoder, GzDecoder}; use serde_json::{json, Value}; use std::collections::BTreeMap; -use std::io::Read; const ADMIN_PROVIDER_OPS_VERIFY_TIMEOUT_MS: u64 = 30_000; +const ADMIN_PROVIDER_OPS_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; pub(super) struct AdminProviderOpsTextResponse { pub(super) body: String, @@ -170,7 +168,9 @@ async fn admin_provider_ops_execute_request( key_id: String::new(), method: method.as_str().to_string(), url: url.to_string(), - headers: admin_provider_ops_execution_headers(headers), + headers: admin_provider_ops_execution_headers( + &admin_provider_ops_headers_with_transport_controls(headers, Some(false), false), + ), content_type: has_json_body.then(|| "application/json".to_string()), content_encoding: None, body, @@ -189,8 +189,12 @@ async fn admin_provider_ops_execute_request( ..ExecutionTimeouts::default() }), }; + let bounded_plan = crate::execution_runtime::transport::with_upstream_response_body_limit( + &plan, + ADMIN_PROVIDER_OPS_RESPONSE_BODY_LIMIT_BYTES, + ); state - .execute_execution_runtime_sync_plan(None, &plan) + .execute_execution_runtime_sync_plan(None, &bounded_plan) .await .map_err(admin_provider_ops_gateway_error_message) } @@ -254,9 +258,9 @@ fn admin_provider_ops_execution_json_response( match serde_json::from_slice::(&bytes) { Ok(value) => Ok((status, value)), Err(_) if status != http::StatusCode::OK => Ok((status, json!({}))), - Err(err) => Err(AdminProviderOpsExecuteJsonError::InvalidJson(format!( - "upstream response is not valid JSON: {err}" - ))), + Err(_) => Err(AdminProviderOpsExecuteJsonError::InvalidJson( + "upstream response is not valid JSON".to_string(), + )), } } @@ -264,49 +268,28 @@ fn admin_provider_ops_execution_body_bytes( headers: &BTreeMap, body: &aether_contracts::ResponseBody, ) -> Option> { - let bytes = body - .body_bytes_b64 - .as_deref() - .and_then(|value| STANDARD.decode(value).ok())?; - admin_provider_ops_decode_response_bytes( + let bytes = body.body_bytes_b64.as_deref().and_then(|value| { + crate::execution_runtime::transport::decode_base64_body_with_limit( + value, + ADMIN_PROVIDER_OPS_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + })?; + crate::execution_runtime::transport::decode_response_body_bytes_with_limit( + headers, &bytes, - headers.get("content-encoding").map(String::as_str), + ADMIN_PROVIDER_OPS_RESPONSE_BODY_LIMIT_BYTES, ) - .or(Some(bytes)) -} - -fn admin_provider_ops_decode_response_bytes( - bytes: &[u8], - content_encoding: Option<&str>, -) -> Option> { - let encoding = content_encoding - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); - match encoding.as_deref() { - Some("gzip") => { - let mut decoder = GzDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) - } - Some("deflate") => { - let mut decoder = DeflateDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) - } - _ => None, - } + .ok() + .map(std::borrow::Cow::into_owned) } fn admin_provider_ops_gateway_error_message(error: GatewayError) -> String { - error.into_message() + admin_provider_ops_verify_execution_error_message(&error.into_message()) } pub(super) fn admin_provider_ops_verify_execution_error_message(error: &str) -> String { - let normalized = error.trim(); - let lower = normalized.to_ascii_lowercase(); + let lower = error.trim().to_ascii_lowercase(); if lower.contains("timeout") || lower.contains("timed out") { return "连接超时".to_string(); } @@ -316,7 +299,22 @@ pub(super) fn admin_provider_ops_verify_execution_error_message(error: &str) -> || lower.contains("proxy") || lower.contains("relay") { - return format!("连接失败: {normalized}"); + return "连接失败".to_string(); + } + "验证失败".to_string() +} + +#[cfg(test)] +mod tests { + use super::admin_provider_ops_verify_execution_error_message; + + #[test] + fn provider_ops_transport_errors_do_not_echo_urls_or_credentials() { + let raw = "connection failed for https://user:password@api.example.test/path?token=secret"; + let projected = admin_provider_ops_verify_execution_error_message(raw); + assert_eq!(projected, "连接失败"); + for secret in ["user", "password", "token", "secret", "api.example.test"] { + assert!(!projected.contains(secret), "leaked {secret}"); + } } - format!("验证失败: {normalized}") } diff --git a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/sub2api.rs b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/sub2api.rs index b1026e5d1..858cdca90 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/sub2api.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/ops/providers/verify/sub2api.rs @@ -4,6 +4,9 @@ use super::request::{ }; use crate::handlers::admin::provider::ops::providers::config::persist_admin_provider_ops_runtime_credentials; use crate::handlers::admin::request::AdminAppState; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url, +}; use aether_admin::provider::ops::{ admin_provider_ops_frontend_updated_credentials, admin_provider_ops_verify_failure, parse_verify_payload, ADMIN_PROVIDER_OPS_USER_AGENT, @@ -46,7 +49,10 @@ pub(super) async fn admin_provider_ops_local_sub2api_verify_response( } } - let verify_url = admin_provider_ops_sub2api_request_url(base_url, verify_endpoint); + let verify_url = match admin_provider_ops_sub2api_request_url(base_url, verify_endpoint) { + Ok(url) => url, + Err(message) => return admin_provider_ops_verify_failure(message), + }; let auth_value = match reqwest::header::HeaderValue::from_str(&format!("Bearer {access_token}")) { Ok(value) => value, @@ -98,20 +104,9 @@ pub(super) async fn admin_provider_ops_local_sub2api_verify_response( pub(in super::super) fn admin_provider_ops_sub2api_request_url( base_url: &str, endpoint: &str, -) -> String { - let trimmed_base_url = base_url.trim().trim_end_matches('/'); - let trimmed_endpoint = endpoint.trim(); - if trimmed_endpoint.is_empty() { - return trimmed_base_url.to_string(); - } - if trimmed_endpoint.starts_with("http://") || trimmed_endpoint.starts_with("https://") { - return trimmed_endpoint.to_string(); - } - - reqwest::Url::parse(trimmed_base_url) - .and_then(|base| base.join(trimmed_endpoint)) - .map(|url| url.to_string()) - .unwrap_or_else(|_| format!("{trimmed_base_url}{trimmed_endpoint}")) +) -> Result { + let destination = canonicalize_provider_ops_base_url(base_url).map_err(ToString::to_string)?; + resolve_provider_ops_same_origin_url(&destination, endpoint).map_err(ToString::to_string) } fn admin_provider_ops_sub2api_updated_credentials( @@ -212,7 +207,7 @@ async fn admin_provider_ops_sub2api_token_request( default_error: &str, proxy_snapshot: Option<&ProxySnapshot>, ) -> Result, String> { - let url = admin_provider_ops_sub2api_request_url(base_url, path); + let url = admin_provider_ops_sub2api_request_url(base_url, path)?; let default_headers = reqwest::header::HeaderMap::from_iter([ ( reqwest::header::USER_AGENT, @@ -246,11 +241,10 @@ async fn admin_provider_ops_sub2api_token_request( if status != http::StatusCode::OK || payload.get("code").and_then(Value::as_i64).unwrap_or(-1) != 0 { - let message = payload - .get("message") - .and_then(Value::as_str) - .unwrap_or(default_error); - return Err(message.to_string()); + // Upstream messages are untrusted response data and may echo credentials or + // request headers. Callers only need a stable local classification here; + // never reflect the upstream text into verify responses, logs, or caches. + return Err(default_error.to_string()); } payload .get("data") diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs index 9b2fbb96c..71ac5890c 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs @@ -1,9 +1,19 @@ +use sha2::{Digest, Sha256}; + pub(super) fn pool_sticky_pattern(provider_id: &str) -> String { format!("ap:{provider_id}:sticky:*") } pub(super) fn pool_sticky_key(provider_id: &str, session_token: &str) -> String { - format!("ap:{provider_id}:sticky:{session_token}") + let digest = Sha256::digest( + format!( + "aether-provider-pool-sticky-v1\0{}\0{}", + provider_id.trim(), + session_token.trim() + ) + .as_bytes(), + ); + format!("ap:{provider_id}:sticky:v1:{digest:x}") } pub(super) fn pool_lru_key(provider_id: &str) -> String { @@ -64,3 +74,19 @@ pub(super) fn pool_latency_keys(provider_id: &str, key_ids: &[String]) -> Vec) -> bool { + let body = error_body.unwrap_or_default().to_ascii_lowercase(); + [ + "quota exhausted", + "quota_exhausted", + "quota exceeded", + "quota_exceeded", + "insufficient_quota", + "resource exhausted", + "resource has been exhausted", + "resource_exhausted", + "usage_limit_reached", + "limit_reached", + "quota limit reached", + "credits exhausted", + "insufficient credits", + ] + .iter() + .any(|marker| body.contains(marker)) +} + pub(crate) async fn record_admin_provider_pool_stream_timeout( runtime: &RuntimeState, provider_id: &str, @@ -589,9 +608,10 @@ pub(crate) async fn record_admin_provider_pool_stream_timeout( #[cfg(test)] mod tests { use super::{ - admin_provider_pool_key_terminal_error_reason, parse_google_quota_cooldown_seconds_at, - record_admin_provider_pool_error, record_admin_provider_pool_stream_timeout, - record_admin_provider_pool_success, resolve_transient_cooldown_ttl, + admin_provider_pool_key_terminal_error_reason, error_body_indicates_quota_exhaustion, + parse_google_quota_cooldown_seconds_at, record_admin_provider_pool_error, + record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success, + resolve_transient_cooldown_ttl, }; use crate::handlers::admin::provider::pool::runtime::reads::read_admin_provider_pool_runtime_state; use crate::handlers::admin::provider::shared::support::{ @@ -803,6 +823,16 @@ mod tests { assert_eq!(resolve_transient_cooldown_ttl(500, None, &pool_config), 0); } + #[test] + fn detects_quota_exhaustion_markers_in_429_bodies() { + assert!(error_body_indicates_quota_exhaustion(Some( + r#"{"error":{"status":"RESOURCE_EXHAUSTED","message":"quota exceeded"}}"#, + ))); + assert!(!error_body_indicates_quota_exhaustion(Some( + r#"{"error":{"message":"temporary rate limit"}}"#, + ))); + } + #[tokio::test] async fn success_feedback_writes_sticky_lru_cost_and_latency() { let Some(redis) = start_managed_redis_or_skip().await else { @@ -1110,7 +1140,7 @@ mod tests { .cooldown_reason_by_key .get("key-google-429") .map(String::as_str), - Some("rate_limited_429") + Some("quota_exhausted_429") ); assert!(runtime .cooldown_ttl_by_key diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/action.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/action.rs index 2085ecd81..ca23de9ba 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/action.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/action.rs @@ -11,6 +11,15 @@ use axum::{ response::Response, }; +const MAX_ADMIN_POOL_BATCH_ITEMS: usize = 100; + +pub(super) fn validate_admin_pool_batch_item_count(item_count: usize) -> Result<(), String> { + if item_count > MAX_ADMIN_POOL_BATCH_ITEMS { + return Err(format!("key_ids 最多 {MAX_ADMIN_POOL_BATCH_ITEMS} 个")); + } + Ok(()) +} + pub(super) async fn build_admin_pool_batch_action_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -54,8 +63,25 @@ pub(super) async fn build_admin_pool_batch_action_response( )); } }; + if let Err(detail) = validate_admin_pool_batch_item_count(payload.key_ids.len()) { + return Ok(build_admin_pool_error_response( + http::StatusCode::BAD_REQUEST, + detail, + )); + } state .build_admin_pool_batch_action_response(&provider_id, payload) .await } + +#[cfg(test)] +mod tests { + use super::{validate_admin_pool_batch_item_count, MAX_ADMIN_POOL_BATCH_ITEMS}; + + #[test] + fn admin_pool_batch_item_count_has_an_inclusive_boundary() { + assert!(validate_admin_pool_batch_item_count(MAX_ADMIN_POOL_BATCH_ITEMS).is_ok()); + assert!(validate_admin_pool_batch_item_count(MAX_ADMIN_POOL_BATCH_ITEMS + 1).is_err()); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/update.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/update.rs index f85a0bc3a..ad0b8a7da 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/update.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/batch_routes/update.rs @@ -1,6 +1,6 @@ use super::{ - admin_pool_provider_id_from_path, build_admin_pool_error_response, - ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL, + admin_pool_provider_id_from_path, batch_action::validate_admin_pool_batch_item_count, + build_admin_pool_error_response, ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL, ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL, }; use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyBatchUpdateRequest; @@ -53,6 +53,12 @@ pub(super) async fn build_admin_pool_batch_update_response( )); } }; + if let Err(detail) = validate_admin_pool_batch_item_count(payload.key_ids.len()) { + return Ok(build_admin_pool_error_response( + http::StatusCode::BAD_REQUEST, + detail, + )); + } state .build_admin_pool_batch_update_response(&provider_id, payload) 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 98076f0df..5dec0f9d5 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 @@ -10,6 +10,11 @@ use crate::provider_key_auth::{ }; use aether_admin::provider::pool as admin_provider_pool_pure; use aether_admin::provider::quota as admin_provider_quota_pure; +use aether_admin::provider::redaction::{ + admin_provider_metadata_bucket_safe_json, admin_provider_oauth_invalid_reason_safe_text, + admin_provider_status_snapshot_safe_json, admin_provider_upstream_metadata_safe_json, + admin_secret_safe_json, admin_secret_safe_proxy, +}; use aether_data_contracts::repository::pool_scores::StoredPoolMemberScore; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, @@ -46,6 +51,14 @@ fn admin_pool_json_object( .filter(|value| !value.is_empty()) } +fn admin_pool_secret_safe_json_object(value: Option<&serde_json::Value>) -> serde_json::Value { + admin_pool_json_object(value) + .map(serde_json::Value::Object) + .as_ref() + .map(|value| admin_secret_safe_json(Some(value))) + .unwrap_or(serde_json::Value::Null) +} + fn admin_pool_json_to_f64(value: Option<&serde_json::Value>) -> Option { let parsed = match value { Some(serde_json::Value::Number(number)) => number.as_f64(), @@ -1087,7 +1100,10 @@ pub(super) fn build_admin_pool_key_payload( codex_cycle_usage_by_code: Option<&BTreeMap>, now_unix_secs: u64, ) -> serde_json::Value { - let cooldown_reason = runtime.cooldown_reason_by_key.get(&key.id).cloned(); + let cooldown_reason = runtime + .cooldown_reason_by_key + .get(&key.id) + .map(|_| "Provider key is cooling down".to_string()); let cooldown_ttl_seconds = cooldown_reason .as_ref() .and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied()); @@ -1251,7 +1267,9 @@ pub(super) fn build_admin_pool_key_payload( "oauth_invalid_reason".to_string(), json!(auth_semantics .can_show_oauth_metadata() - .then_some(key.oauth_invalid_reason.clone()) + .then(|| { + admin_provider_oauth_invalid_reason_safe_text(key.oauth_invalid_reason.as_deref()) + }) .flatten()), ); payload.insert("oauth_plan_type".to_string(), json!(oauth_plan_type)); @@ -1290,7 +1308,10 @@ pub(super) fn build_admin_pool_key_payload( "account_status_source".to_string(), json!(account_status_source), ); - payload.insert("status_snapshot".to_string(), status_snapshot); + payload.insert( + "status_snapshot".to_string(), + admin_provider_status_snapshot_safe_json(Some(&status_snapshot)), + ); payload.insert("quota_updated_at".to_string(), json!(quota_updated_at)); payload.insert("health_score".to_string(), json!(health_score)); payload.insert( @@ -1305,7 +1326,10 @@ pub(super) fn build_admin_pool_key_payload( "score": score.score, "hard_state": score.hard_state.as_database(), "score_version": score.score_version, - "score_reason": score.score_reason.clone(), + "score_reason": admin_provider_metadata_bucket_safe_json( + "pool_score", + Some(&score.score_reason), + ), "last_ranked_at": score.last_ranked_at, "last_scheduled_at": score.last_scheduled_at, "last_success_at": score.last_success_at, @@ -1335,7 +1359,7 @@ pub(super) fn build_admin_pool_key_payload( ); payload.insert( "rate_multipliers".to_string(), - json!(admin_pool_json_object(key.rate_multipliers.as_ref())), + admin_pool_secret_safe_json_object(key.rate_multipliers.as_ref()), ); payload.insert( "internal_priority".to_string(), @@ -1358,7 +1382,7 @@ pub(super) fn build_admin_pool_key_payload( ); payload.insert( "capabilities".to_string(), - json!(admin_pool_json_object(key.capabilities.as_ref())), + admin_pool_secret_safe_json_object(key.capabilities.as_ref()), ); payload.insert( "auto_fetch_models".to_string(), @@ -1378,10 +1402,16 @@ pub(super) fn build_admin_pool_key_payload( ); payload.insert( "upstream_metadata".to_string(), - json!(key.upstream_metadata.clone()), + admin_provider_upstream_metadata_safe_json(key.upstream_metadata.as_ref()), + ); + payload.insert( + "proxy".to_string(), + admin_secret_safe_proxy(key.proxy.as_ref()), + ); + payload.insert( + "fingerprint".to_string(), + admin_secret_safe_json(key.fingerprint.as_ref()), ); - payload.insert("proxy".to_string(), json!(key.proxy.clone())); - payload.insert("fingerprint".to_string(), json!(key.fingerprint.clone())); payload.insert("account_quota".to_string(), json!(account_quota)); payload.insert("cooldown_reason".to_string(), json!(cooldown_reason)); payload.insert( diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/scores.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/scores.rs index 1dc3dcf75..bac3159af 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/scores.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/scores.rs @@ -5,6 +5,7 @@ use super::{ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::shared::query_param_value; use crate::GatewayError; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; use aether_data_contracts::repository::pool_scores::{ ListPoolMemberScoresQuery, PoolMemberHardState, PoolMemberProbeStatus, POOL_KIND_PROVIDER_KEY_POOL, POOL_SCORE_CAPABILITY_ACCOUNT, POOL_SCORE_SCOPE_KIND_ACCOUNT, @@ -114,7 +115,10 @@ pub(super) async fn build_admin_pool_scores_response( "score": score.score, "hard_state": score.hard_state.as_database(), "score_version": score.score_version, - "score_reason": score.score_reason, + "score_reason": admin_provider_metadata_bucket_safe_json( + "pool_score", + Some(&score.score_reason), + ), "last_ranked_at": score.last_ranked_at, "last_scheduled_at": score.last_scheduled_at, "last_success_at": score.last_success_at, diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/selection.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/selection.rs index 17db13e63..43aa305b7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/selection.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/selection.rs @@ -15,7 +15,11 @@ fn admin_pool_parse_auth_config_json( if ciphertext.is_empty() { return None; } - let plaintext = state.decrypt_catalog_secret_with_fallbacks(ciphertext)?; + let plaintext = state + .app() + .decrypt_provider_catalog_key_auth_config(key) + .ok() + .flatten()?; serde_json::from_str::(&plaintext) .ok()? .as_object() 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 e88bc5356..cb0cc44e5 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 @@ -29,7 +29,7 @@ use crate::handlers::shared::{ parse_catalog_auth_config_json, provider_key_health_summary, provider_key_status_snapshot_payload, }; -use crate::model_fetch::ModelFetchRuntimeState; +use crate::model_fetch::{safe_model_fetch_error, ModelFetchRuntimeState}; use crate::provider_key_auth::{ provider_key_auth_semantics, provider_key_configured_api_formats, provider_key_inherits_provider_api_formats, @@ -63,7 +63,7 @@ use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use aether_model_fetch::{ - aggregate_models_for_cache, fetch_models_from_transports, json_string_list, + aggregate_models_for_cache, fetch_models_from_transports_for_management, json_string_list, model_catalog_upstream_metadata, preset_models_for_provider, selected_models_fetch_endpoints, upstream_metadata_namespace_updates, }; @@ -319,6 +319,15 @@ fn provider_query_codex_preset_fallback( }) } +fn provider_query_project_model_fetch_errors( + errors: impl IntoIterator, +) -> Vec { + errors + .into_iter() + .map(|error| safe_model_fetch_error(&error)) + .collect() +} + async fn provider_query_persist_preset_models( state: &AdminAppState<'_>, provider: &StoredProviderCatalogProvider, @@ -496,10 +505,34 @@ async fn provider_query_fetch_models_for_key( key: &StoredProviderCatalogKey, force_refresh: bool, ) -> Result { + let is_codex = provider.provider_type.trim().eq_ignore_ascii_case("codex"); + let codex_catalog = if is_codex { + crate::model_fetch::read_codex_management_catalog(state.app(), &provider.id, &key.id).await + } else { + None + }; if !force_refresh { - if let Some(cached_models) = + // Read through the live, credential-scoped directory before the admin cache. + // Never reuse the old versionless cache for Codex (including cached presets). + let cached_models = if is_codex { + codex_catalog + .as_ref() + .and_then(|catalog| catalog.models.as_ref()) + .filter(|_| { + selected_models_fetch_endpoints(endpoints, key) + .iter() + .any(|endpoint| endpoint.api_format == "openai:responses") + }) + .map(|models| { + aether_model_fetch::project_codex_models_for_legacy_cache([( + "openai:responses", + models.as_slice(), + )]) + }) + } else { provider_query_read_cached_models(state, &provider.id, &key.id).await - { + }; + if let Some(cached_models) = cached_models { let models = provider_query_filter_models_for_key(provider, key, cached_models); return Ok(ProviderQueryKeyFetchResult { models, @@ -561,30 +594,51 @@ async fn provider_query_fetch_models_for_key( }); } - let outcome = match fetch_models_from_transports(state.app(), &transports).await { - Ok(outcome) => outcome, - Err(err) => { - all_errors.push(err); - if let Some(fallback) = - provider_query_codex_preset_fallback(provider, &all_errors.join("; ")) - { - provider_query_persist_preset_models(state, provider, key, &fallback.models) - .await?; - return Ok(fallback); + 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) + }); + let outcome = + match fetch_models_from_transports_for_management(state.app(), &transports, client_version) + .await + { + Ok(outcome) => outcome, + Err(err) => { + all_errors.extend(provider_query_project_model_fetch_errors([err])); + if let Some(fallback) = + provider_query_codex_preset_fallback(provider, &all_errors.join("; ")) + { + provider_query_persist_preset_models(state, provider, key, &fallback.models) + .await?; + return Ok(fallback); + } + return Ok(ProviderQueryKeyFetchResult { + models: Vec::new(), + error: Some(all_errors.join("; ")), + warning: None, + from_cache: false, + has_success: false, + }); } - return Ok(ProviderQueryKeyFetchResult { - models: Vec::new(), - error: Some(all_errors.join("; ")), - warning: None, - from_cache: false, - has_success: false, - }); - } - }; + }; - all_errors.extend(outcome.errors); + all_errors.extend(provider_query_project_model_fetch_errors(outcome.errors)); let unique_models = outcome.legacy_models; if outcome.has_success && !unique_models.is_empty() { + if all_errors.is_empty() && outcome.native_codex_catalog { + if let Some(catalog) = codex_catalog.as_ref() { + crate::model_fetch::store_codex_management_catalog( + state.app(), + catalog, + &transports, + outcome.cached_models, + outcome.etag.as_deref(), + ) + .await; + } + } ::write_upstream_models_cache( state.app(), &provider.id, @@ -1026,4 +1080,32 @@ mod tests { json!(true) ); } + + #[test] + fn provider_query_projects_transport_errors_before_exposing_them() { + let projected = provider_query_project_model_fetch_errors(vec![ + "HTTP 401 response body: Authorization: Bearer upstream-secret".to_string(), + "connection failed for https://user:password@example.test/v1/models?key=secret" + .to_string(), + ]); + + assert_eq!( + projected, + [ + "Upstream models fetch authentication failed (status 401)", + "Upstream models fetch connection failed", + ] + ); + let exposed = projected.join("; "); + for secret in [ + "upstream-secret", + "Bearer", + "user", + "password", + "example.test", + "?key=secret", + ] { + assert!(!exposed.contains(secret)); + } + } } 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 ea6ff99d5..8319b1409 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 @@ -74,7 +74,6 @@ use axum::{ response::{IntoResponse, Response}, Json, }; -use base64::Engine as _; use serde_json::{json, Map, Value}; use std::collections::{BTreeMap, BTreeSet}; use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering}; @@ -106,7 +105,8 @@ use self::model_mapping::{ provider_query_resolve_global_effective_model, }; use self::summary::{ - provider_query_candidate_summary_payload, provider_query_test_attempt_payload, + provider_query_candidate_summary_payload, provider_query_error_projection, + provider_query_success_response_body, provider_query_test_attempt_payload, }; pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str = @@ -209,7 +209,9 @@ fn provider_query_test_candidate_trace_extra_data( "admin_model_test": { "provider_type": provider.provider_type, "endpoint_api_format": candidate.endpoint.api_format, - "endpoint_base_url": candidate.endpoint.base_url, + "endpoint_base_url": aether_admin::provider::redaction::admin_secret_safe_url( + Some(&candidate.endpoint.base_url) + ), "effective_model": candidate.effective_model, } }) @@ -386,6 +388,7 @@ async fn provider_query_finish_test_candidate_trace( "skipped" => RequestCandidateStatus::Skipped, _ => RequestCandidateStatus::Failed, }; + let projected_error_message = provider_query_error_projection(execution); provider_query_persist_test_candidate_trace( state, trace_id, @@ -395,7 +398,7 @@ async fn provider_query_finish_test_candidate_trace( status, ProviderQueryTestTraceUpdate { skip_reason: execution.skip_reason.as_deref(), - error_message: execution.error_message.as_deref(), + error_message: projected_error_message.as_deref(), status_code: execution.status_code, latency_ms: execution.latency_ms, finished_at_unix_ms: Some(current_unix_ms()), @@ -1646,11 +1649,19 @@ async fn provider_query_reconcile_fixed_provider_endpoints_for_test_model( fn provider_query_decode_execution_body( result: &aether_contracts::ExecutionResult, ) -> Option> { + const MAX_PROVIDER_QUERY_RESULT_BODY_BYTES: usize = 64 * 1024 * 1024; result .body .as_ref() .and_then(|body| body.body_bytes_b64.as_deref()) - .and_then(|value| base64::engine::general_purpose::STANDARD.decode(value).ok()) + .and_then(|value| { + crate::execution_runtime::transport::decode_base64_body_with_limit( + value, + crate::headers::max_internal_buffered_body_bytes() + .min(MAX_PROVIDER_QUERY_RESULT_BODY_BYTES), + ) + .ok() + }) } fn provider_query_execution_json_body(result: &aether_contracts::ExecutionResult) -> Option { @@ -3177,6 +3188,7 @@ async fn provider_query_execute_standard_test_candidate( crate::ai_serving::openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + request_model, ), ) else { @@ -3238,6 +3250,7 @@ async fn provider_query_execute_standard_test_candidate( crate::ai_serving::openai_responses_reasoning_replay_policy( transport.provider.provider_type.as_str(), transport.endpoint.base_url.as_str(), + request_model, ), ) .is_err() @@ -3936,7 +3949,7 @@ async fn build_admin_provider_query_kiro_failover_response( total_attempts += 1; } let is_success = execution.status == "success"; - let response_body = execution.response_body.clone(); + let response_body = provider_query_success_response_body(&execution); attempts.push(provider_query_test_attempt_payload( candidate_index, candidate, diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs index 72618af5c..7c98303af 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs @@ -1,7 +1,8 @@ -use std::collections::BTreeMap; +use std::collections::{BTreeMap, BTreeSet}; use super::super::provider_query_key_display_name; use super::{ProviderQueryExecutionOutcome, ProviderQueryTestCandidate}; +use aether_admin::provider::redaction::admin_secret_safe_url; use serde_json::{json, Value}; pub(super) fn provider_query_test_attempt_payload( @@ -10,6 +11,7 @@ pub(super) fn provider_query_test_attempt_payload( execution: &ProviderQueryExecutionOutcome, ) -> Value { let endpoint_route = provider_query_endpoint_route_payload(candidate, execution); + let response_body = provider_query_success_response_body(execution); let endpoint_product = endpoint_route .get("product") .cloned() @@ -27,7 +29,7 @@ pub(super) fn provider_query_test_attempt_payload( "candidate_index": candidate_index, "retry_index": 0, "endpoint_api_format": candidate.endpoint.api_format, - "endpoint_base_url": candidate.endpoint.base_url, + "endpoint_base_url": admin_secret_safe_url(Some(&candidate.endpoint.base_url)), "endpoint_product": endpoint_product, "endpoint_variant": endpoint_variant, "endpoint_action": endpoint_action, @@ -38,14 +40,25 @@ pub(super) fn provider_query_test_attempt_payload( "effective_model": candidate.effective_model, "status": execution.status, "skip_reason": execution.skip_reason, - "error_message": execution.error_message, + "error_message": provider_query_error_projection(execution), "status_code": execution.status_code, "latency_ms": execution.latency_ms, - "request_url": execution.request_url, + "request_url": admin_secret_safe_url(Some(&execution.request_url)), "request_headers": redacted_provider_query_headers(&execution.request_headers), "request_body": redacted_provider_query_value(&execution.request_body), "response_headers": redacted_provider_query_headers(&execution.response_headers), - "response_body": execution.response_body, + "response_body": response_body, + }) +} + +pub(super) fn provider_query_error_projection( + execution: &ProviderQueryExecutionOutcome, +) -> Option { + execution.error_message.as_ref().map(|_| { + execution + .status_code + .map(|status| format!("HTTP {status}")) + .unwrap_or_else(|| "Provider request failed".to_string()) }) } @@ -195,7 +208,28 @@ fn redacted_provider_query_headers(headers: &BTreeMap) -> BTreeM .collect() } -fn redacted_provider_query_value(value: &Value) -> Value { +pub(super) fn redacted_provider_query_value(value: &Value) -> Value { + redacted_provider_query_value_with_sensitive_values(value, &BTreeSet::new()) +} + +pub(super) fn provider_query_success_response_body( + execution: &ProviderQueryExecutionOutcome, +) -> Option { + if execution.status != "success" { + return None; + } + + let sensitive_values = provider_query_execution_sensitive_values(execution); + execution + .response_body + .as_ref() + .map(|value| redacted_provider_query_value_with_sensitive_values(value, &sensitive_values)) +} + +fn redacted_provider_query_value_with_sensitive_values( + value: &Value, + sensitive_values: &BTreeSet, +) -> Value { match value { Value::Object(object) => Value::Object( object @@ -204,7 +238,13 @@ fn redacted_provider_query_value(value: &Value) -> Value { if provider_query_field_is_sensitive(key) { (key.clone(), Value::String("[REDACTED]".to_string())) } else { - (key.clone(), redacted_provider_query_value(value)) + ( + key.clone(), + redacted_provider_query_value_with_sensitive_values( + value, + sensitive_values, + ), + ) } }) .collect(), @@ -212,19 +252,238 @@ fn redacted_provider_query_value(value: &Value) -> Value { Value::Array(items) => Value::Array( items .iter() - .map(redacted_provider_query_value) + .map(|value| { + redacted_provider_query_value_with_sensitive_values(value, sensitive_values) + }) .collect::>(), ), + Value::String(value) + if provider_query_string_contains_sensitive_material(value, sensitive_values) => + { + Value::String("[REDACTED]".to_string()) + } other => other.clone(), } } +fn provider_query_execution_sensitive_values( + execution: &ProviderQueryExecutionOutcome, +) -> BTreeSet { + let mut sensitive_values = BTreeSet::new(); + for (key, value) in &execution.request_headers { + if provider_query_field_is_sensitive(key) { + provider_query_insert_sensitive_value(&mut sensitive_values, value); + } + } + provider_query_collect_sensitive_field_values(&execution.request_body, &mut sensitive_values); + provider_query_collect_sensitive_url_values(&execution.request_url, &mut sensitive_values); + sensitive_values +} + +fn provider_query_collect_sensitive_field_values( + value: &Value, + sensitive_values: &mut BTreeSet, +) { + match value { + Value::Object(object) => { + for (key, value) in object { + if provider_query_field_is_sensitive(key) { + provider_query_collect_string_values(value, sensitive_values); + } else { + provider_query_collect_sensitive_field_values(value, sensitive_values); + } + } + } + Value::Array(items) => { + for value in items { + provider_query_collect_sensitive_field_values(value, sensitive_values); + } + } + _ => {} + } +} + +fn provider_query_collect_string_values(value: &Value, sensitive_values: &mut BTreeSet) { + match value { + Value::String(value) => provider_query_insert_sensitive_value(sensitive_values, value), + Value::Object(object) => { + for value in object.values() { + provider_query_collect_string_values(value, sensitive_values); + } + } + Value::Array(items) => { + for value in items { + provider_query_collect_string_values(value, sensitive_values); + } + } + _ => {} + } +} + +fn provider_query_collect_sensitive_url_values( + value: &str, + sensitive_values: &mut BTreeSet, +) { + let Ok(parsed) = url::Url::parse(value) else { + return; + }; + if !parsed.username().is_empty() { + provider_query_insert_sensitive_value(sensitive_values, parsed.username()); + } + if let Some(password) = parsed.password() { + provider_query_insert_sensitive_value(sensitive_values, password); + } + for (key, value) in parsed.query_pairs() { + if provider_query_url_query_field_is_sensitive(&key) { + provider_query_insert_sensitive_value(sensitive_values, &value); + } + } +} + +fn provider_query_insert_sensitive_value(sensitive_values: &mut BTreeSet, value: &str) { + let value = value.trim(); + if value.is_empty() || value == "[REDACTED]" { + return; + } + sensitive_values.insert(value.to_string()); + + let lower = value.to_ascii_lowercase(); + for scheme in ["bearer ", "basic ", "token "] { + if lower.starts_with(scheme) { + let credential = value[scheme.len()..].trim(); + if !credential.is_empty() { + sensitive_values.insert(credential.to_string()); + } + } + } +} + +fn provider_query_string_contains_sensitive_material( + value: &str, + sensitive_values: &BTreeSet, +) -> bool { + let value = value.trim(); + if value.is_empty() { + return false; + } + if sensitive_values + .iter() + .any(|secret| value == secret || (secret.len() >= 8 && value.contains(secret.as_str()))) + { + return true; + } + + let lower = value.to_ascii_lowercase(); + if provider_query_contains_credential_scheme(&lower, "bearer") + || provider_query_contains_credential_scheme(&lower, "basic") + { + return true; + } + + [ + "authorization", + "proxy-authorization", + "api_key", + "api-key", + "api key", + "apikey", + "x-api-key", + "x-goog-api-key", + "access_token", + "access token", + "refresh_token", + "refresh token", + "id_token", + "id token", + "client_secret", + "client secret", + "password", + "passwd", + "secret", + ] + .iter() + .any(|marker| provider_query_contains_secret_assignment(&lower, marker)) + || provider_query_contains_known_token_prefix(&lower) +} + +fn provider_query_contains_credential_scheme(value: &str, scheme: &str) -> bool { + value.match_indices(scheme).any(|(index, _)| { + let before_is_boundary = + index == 0 || !value.as_bytes()[index.saturating_sub(1)].is_ascii_alphanumeric(); + if !before_is_boundary { + return false; + } + let credential = value[index + scheme.len()..].trim_start(); + let credential = credential + .strip_prefix(':') + .or_else(|| credential.strip_prefix('=')) + .unwrap_or(credential) + .trim_start(); + credential + .split(|ch: char| ch.is_whitespace() || matches!(ch, '"' | '\'' | ',' | ';')) + .next() + .is_some_and(|token| token.len() >= 6) + }) +} + +fn provider_query_contains_secret_assignment(value: &str, marker: &str) -> bool { + value.match_indices(marker).any(|(index, _)| { + let before_is_boundary = + index == 0 || !value.as_bytes()[index.saturating_sub(1)].is_ascii_alphanumeric(); + if !before_is_boundary { + return false; + } + let remainder = value[index + marker.len()..] + .trim_start_matches(|ch: char| ch.is_whitespace() || matches!(ch, '"' | '\'')); + let Some(remainder) = remainder + .strip_prefix(':') + .or_else(|| remainder.strip_prefix('=')) + else { + return false; + }; + let assigned = + remainder.trim_start_matches(|ch: char| ch.is_whitespace() || matches!(ch, '"' | '\'')); + !assigned.is_empty() + && !assigned.starts_with("null") + && !assigned.starts_with("none") + && !assigned.starts_with("[redacted]") + }) +} + +fn provider_query_contains_known_token_prefix(value: &str) -> bool { + [ + "sk-", + "sk_", + "ghp_", + "github_pat_", + "xoxb-", + "xoxp-", + "aiza", + ] + .iter() + .any(|prefix| { + value.match_indices(*prefix).any(|(index, _)| { + value[index..] + .split(|ch: char| ch.is_whitespace() || matches!(ch, '"' | '\'' | ',' | ';')) + .next() + .is_some_and(|token| token.len() >= 12) + }) + }) +} + +fn provider_query_url_query_field_is_sensitive(key: &str) -> bool { + if provider_query_field_is_sensitive(key) { + return true; + } + matches!( + provider_query_normalized_field_key(key).as_str(), + "key" | "credential" | "signature" | "sig" + ) +} + fn provider_query_field_is_sensitive(key: &str) -> bool { let key = key.trim().to_ascii_lowercase(); - let normalized = key - .chars() - .filter(|ch| ch.is_ascii_alphanumeric()) - .collect::(); + let normalized = provider_query_normalized_field_key(&key); if matches!( normalized.as_str(), "maxtokens" @@ -262,6 +521,13 @@ fn provider_query_field_is_sensitive(key: &str) -> bool { || normalized.contains("authorization") } +fn provider_query_normalized_field_key(key: &str) -> String { + key.chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .map(|ch| ch.to_ascii_lowercase()) + .collect() +} + pub(super) fn provider_query_candidate_summary_payload( total_candidates: usize, total_attempts: usize, @@ -369,10 +635,234 @@ pub(super) fn provider_query_candidate_summary_payload( #[cfg(test)] mod tests { - use super::{redacted_provider_query_headers, redacted_provider_query_value}; + use super::{ + provider_query_test_attempt_payload, redacted_provider_query_headers, + redacted_provider_query_value, ProviderQueryExecutionOutcome, ProviderQueryTestCandidate, + }; + use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + }; use serde_json::json; use std::collections::BTreeMap; + #[test] + fn attempt_payload_strips_credentials_and_queries_from_urls() { + let mut endpoint = StoredProviderCatalogEndpoint::new( + "endpoint-1".to_string(), + "provider-1".to_string(), + "openai:chat".to_string(), + Some("openai".to_string()), + Some("chat".to_string()), + true, + ) + .expect("endpoint"); + endpoint.base_url = + "https://base-user:base-password@api.example.test/v1?base-secret=1#fragment" + .to_string(); + let candidate = ProviderQueryTestCandidate { + endpoint, + key: StoredProviderCatalogKey::new( + "key-1".to_string(), + "provider-1".to_string(), + "key".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key"), + effective_model: "gpt-test".to_string(), + scheduler_skip_reason: None, + }; + let execution = ProviderQueryExecutionOutcome { + status: "success", + skip_reason: None, + error_message: None, + status_code: Some(200), + latency_ms: Some(1), + request_url: + "https://request-user:request-password@api.example.test/v1/chat?key=request-secret#fragment" + .to_string(), + request_headers: BTreeMap::new(), + request_body: json!({}), + response_headers: BTreeMap::new(), + response_body: None, + }; + + let payload = provider_query_test_attempt_payload(0, &candidate, &execution); + assert_eq!(payload["endpoint_base_url"], "https://api.example.test/v1"); + assert_eq!(payload["request_url"], "https://api.example.test/v1/chat"); + let serialized = payload.to_string(); + for secret in [ + "base-user", + "base-password", + "base-secret", + "request-user", + "request-password", + "request-secret", + "fragment", + ] { + assert!(!serialized.contains(secret), "leaked {secret}"); + } + } + + #[test] + fn attempt_payload_does_not_reflect_upstream_error_text_or_response_secrets() { + let candidate = ProviderQueryTestCandidate { + endpoint: StoredProviderCatalogEndpoint::new( + "endpoint-1".to_string(), + "provider-1".to_string(), + "openai:chat".to_string(), + Some("openai".to_string()), + Some("chat".to_string()), + true, + ) + .expect("endpoint"), + key: StoredProviderCatalogKey::new( + "key-1".to_string(), + "provider-1".to_string(), + "key".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key"), + effective_model: "gpt-test".to_string(), + scheduler_skip_reason: None, + }; + let mut execution = ProviderQueryExecutionOutcome { + status: "failed", + skip_reason: None, + error_message: Some( + "authorization=Bearer upstream-secret https://user:pass@example.test?q=secret" + .to_string(), + ), + status_code: Some(502), + latency_ms: Some(1), + request_url: "https://api.example.test/v1/chat".to_string(), + request_headers: BTreeMap::new(), + request_body: json!({}), + response_headers: BTreeMap::new(), + response_body: Some(json!({ + "error": {"message": "authorization=Bearer response-message-secret"}, + "content": "password=response-content-secret", + "access_token": "response-secret", + "nested": {"password": "private-password"} + })), + }; + + let payload = provider_query_test_attempt_payload(0, &candidate, &execution); + assert_eq!(payload["error_message"], json!("HTTP 502")); + assert!(payload["response_body"].is_null()); + let serialized = payload.to_string(); + for secret in [ + "upstream-secret", + "response-message-secret", + "response-content-secret", + "response-secret", + "private-password", + "user:pass", + ] { + assert!(!serialized.contains(secret), "leaked {secret}"); + } + + execution.status_code = None; + let network_failure_payload = + provider_query_test_attempt_payload(0, &candidate, &execution); + assert_eq!( + network_failure_payload["error_message"], + json!("Provider request failed") + ); + assert!(network_failure_payload["response_body"].is_null()); + } + + #[test] + fn successful_attempt_keeps_diagnostics_but_redacts_embedded_credentials() { + let candidate = ProviderQueryTestCandidate { + endpoint: StoredProviderCatalogEndpoint::new( + "endpoint-1".to_string(), + "provider-1".to_string(), + "openai:chat".to_string(), + Some("openai".to_string()), + Some("chat".to_string()), + true, + ) + .expect("endpoint"), + key: StoredProviderCatalogKey::new( + "key-1".to_string(), + "provider-1".to_string(), + "key".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key"), + effective_model: "gpt-test".to_string(), + scheduler_skip_reason: None, + }; + let execution = ProviderQueryExecutionOutcome { + status: "success", + skip_reason: None, + error_message: None, + status_code: Some(200), + latency_ms: Some(1), + request_url: "https://api.example.test/v1/chat?key=query-credential-token-9012" + .to_string(), + request_headers: BTreeMap::from([( + "authorization".to_string(), + "Bearer request-credential-token-1234".to_string(), + )]), + request_body: json!({ + "metadata": {"apiKey": "body-credential-token-5678"} + }), + response_headers: BTreeMap::new(), + response_body: Some(json!({ + "choices": [ + {"message": {"content": "request-credential-token-1234"}}, + {"message": {"content": "Normal model response"}} + ], + "echoed_body": "body-credential-token-5678", + "echoed_query": "query-credential-token-9012", + "warning": {"message": "password=upstream-private-password"}, + "notice": {"content": "authorization: Bearer upstream-auth-secret"}, + "api_key_warning": {"message": "api_key=upstream-api-key-secret"}, + "token_warning": {"content": "access_token=upstream-access-token"}, + "api_key": "direct-response-secret", + "usage": {"input_tokens": 3, "output_tokens": 2} + })), + }; + + let payload = provider_query_test_attempt_payload(0, &candidate, &execution); + assert_eq!( + payload.pointer("/response_body/choices/0/message/content"), + Some(&json!("[REDACTED]")) + ); + assert_eq!( + payload.pointer("/response_body/choices/1/message/content"), + Some(&json!("Normal model response")) + ); + assert_eq!( + payload.pointer("/response_body/usage/input_tokens"), + Some(&json!(3)) + ); + assert_eq!( + payload.pointer("/response_body/usage/output_tokens"), + Some(&json!(2)) + ); + let serialized = payload.to_string(); + for secret in [ + "request-credential-token-1234", + "body-credential-token-5678", + "query-credential-token-9012", + "upstream-private-password", + "upstream-auth-secret", + "upstream-api-key-secret", + "upstream-access-token", + "direct-response-secret", + ] { + assert!(!serialized.contains(secret), "leaked {secret}"); + } + } + #[test] fn redacts_sensitive_provider_query_headers() { let headers = BTreeMap::from([ diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs index 2130651ea..8c700a3e5 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs @@ -1,5 +1,6 @@ use super::*; use crate::handlers::admin::request::AdminGatewayProviderTransportSnapshot; +use base64::Engine as _; use serde_json::json; fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTransportSnapshot { diff --git a/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs b/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs index 3b75b15d7..0b74b8778 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs @@ -3,7 +3,7 @@ use crate::handlers::admin::shared::{ }; use serde::Deserialize; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(crate) struct AdminProviderKeyCreateRequest { #[serde(default)] pub(crate) api_formats: Option>, @@ -48,7 +48,31 @@ pub(crate) struct AdminProviderKeyCreateRequest { pub(crate) fingerprint: Option, } -#[derive(Debug, Deserialize)] +impl std::fmt::Debug for AdminProviderKeyCreateRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminProviderKeyCreateRequest") + .field("api_formats", &self.api_formats) + .field("api_key", &self.api_key.as_ref().map(|_| "[REDACTED]")) + .field("auth_type", &self.auth_type) + .field( + "auth_config", + &self.auth_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("name", &self.name) + .field("internal_priority", &self.internal_priority) + .field("rpm_limit", &self.rpm_limit) + .field("concurrent_limit", &self.concurrent_limit) + .field("auto_fetch_models", &self.auto_fetch_models) + .field( + "fingerprint", + &self.fingerprint.as_ref().map(|_| "[REDACTED]"), + ) + .finish_non_exhaustive() + } +} + +#[derive(Deserialize)] pub(crate) struct AdminProviderKeyUpdateRequest { #[serde(default)] pub(crate) api_formats: Option>, @@ -100,6 +124,32 @@ pub(crate) struct AdminProviderKeyUpdateRequest { pub(crate) fingerprint: Option, } +impl std::fmt::Debug for AdminProviderKeyUpdateRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminProviderKeyUpdateRequest") + .field("api_formats", &self.api_formats) + .field("api_key", &self.api_key.as_ref().map(|_| "[REDACTED]")) + .field("auth_type", &self.auth_type) + .field( + "auth_config", + &self.auth_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("name", &self.name) + .field("internal_priority", &self.internal_priority) + .field("rpm_limit", &self.rpm_limit) + .field("concurrent_limit", &self.concurrent_limit) + .field("is_active", &self.is_active) + .field("auto_fetch_models", &self.auto_fetch_models) + .field("proxy", &self.proxy.as_ref().map(|_| "[REDACTED]")) + .field( + "fingerprint", + &self.fingerprint.as_ref().map(|_| "[REDACTED]"), + ) + .finish_non_exhaustive() + } +} + pub(crate) type AdminProviderKeyUpdatePatch = AdminTypedObjectPatch; #[derive(Debug, Deserialize)] @@ -347,7 +397,54 @@ pub(crate) struct AdminImportProviderModelsRequest { #[cfg(test)] mod tests { - use super::AdminCodexResetCreditConsumeRequest; + use super::{ + AdminCodexResetCreditConsumeRequest, AdminProviderKeyCreateRequest, + AdminProviderKeyUpdateRequest, + }; + + #[test] + fn provider_key_request_debug_output_redacts_authorization_material() { + let create = serde_json::from_value::(serde_json::json!({ + "name": "key", + "api_key": "create-api-key-canary", + "auth_config": {"refresh_token": "create-refresh-token-canary"}, + "fingerprint": {"device_id": "create-device-canary"} + })) + .expect("create request should deserialize"); + let update = serde_json::from_value::(serde_json::json!({ + "api_key": "update-api-key-canary", + "auth_config": {"refresh_token": "update-refresh-token-canary"}, + "proxy": {"password": "update-proxy-canary"}, + "fingerprint": {"device_id": "update-device-canary"} + })) + .expect("update request should deserialize"); + + let create_debug = format!("{create:?}"); + let update_debug = format!("{update:?}"); + assert!(create_debug.contains("[REDACTED]")); + assert!(update_debug.contains("[REDACTED]")); + for secret in [ + "create-api-key-canary", + "create-refresh-token-canary", + "create-device-canary", + ] { + assert!( + !create_debug.contains(secret), + "create debug leaked {secret}" + ); + } + for secret in [ + "update-api-key-canary", + "update-refresh-token-canary", + "update-proxy-canary", + "update-device-canary", + ] { + assert!( + !update_debug.contains(secret), + "update debug leaked {secret}" + ); + } + } #[test] fn codex_reset_credit_consume_requires_an_explicit_credential_generation() { diff --git a/apps/aether-gateway/src/handlers/admin/provider/summary/health.rs b/apps/aether-gateway/src/handlers/admin/provider/summary/health.rs index f3bec7ec7..ecdea68d3 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/summary/health.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/summary/health.rs @@ -3,7 +3,7 @@ use crate::handlers::admin::shared::unix_secs_to_rfc3339; use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label}; use crate::handlers::shared::unix_ms_to_rfc3339; use aether_data_contracts::repository::candidates::{ - RequestCandidateStatus, StoredRequestCandidate, + sanitize_request_candidate_error_type, RequestCandidateStatus, StoredRequestCandidate, }; use serde_json::json; use std::collections::BTreeMap; @@ -128,8 +128,8 @@ pub(crate) async fn build_admin_provider_health_monitor_payload( "status": request_candidate_status_label(candidate.status), "status_code": candidate.status_code, "latency_ms": candidate.latency_ms, - "error_type": candidate.error_type, - "error_message": candidate.error_message, + "error_type": sanitize_request_candidate_error_type(candidate.error_type), + "error_message": serde_json::Value::Null, })) }) .collect::>(); diff --git a/apps/aether-gateway/src/handlers/admin/provider/summary/list.rs b/apps/aether-gateway/src/handlers/admin/provider/summary/list.rs index bdc81e1e1..fdaa6b319 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/summary/list.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/summary/list.rs @@ -1,5 +1,6 @@ use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::unix_secs_to_rfc3339; +use aether_admin::provider::redaction::admin_secret_safe_url; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint; use serde_json::json; use std::collections::{BTreeMap, BTreeSet}; @@ -87,7 +88,9 @@ pub(crate) async fn build_admin_providers_payload( "id": provider_id.clone(), "name": provider.name, "api_format": endpoint.map(|item| item.api_format.clone()), - "base_url": endpoint.map(|item| item.base_url.clone()), + "base_url": endpoint + .map(|item| admin_secret_safe_url(Some(&item.base_url))) + .unwrap_or(serde_json::Value::Null), "api_key": has_any_key_by_provider.contains(&provider_id).then_some("***"), "priority": provider.provider_priority, "is_active": provider.is_active, diff --git a/apps/aether-gateway/src/handlers/admin/provider/summary/value.rs b/apps/aether-gateway/src/handlers/admin/provider/summary/value.rs index d57954030..27aad0d91 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/summary/value.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/summary/value.rs @@ -6,6 +6,7 @@ use crate::handlers::admin::shared::unix_secs_to_rfc3339; use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label}; use crate::orchestration::{codex_cyber_flag_passthrough_enabled, responses_websocket_adapter}; use crate::provider_key_auth::provider_key_effective_api_formats; +use aether_admin::provider::redaction::{admin_secret_safe_json, admin_secret_safe_proxy}; use aether_data_contracts::repository::candidates::{ RequestCandidateStatus, StoredRequestCandidate, }; @@ -196,13 +197,13 @@ pub(crate) fn build_admin_provider_summary_value( "max_retries": provider.max_retries, "max_transfer_count": max_transfer_count, "max_transfer_timeout_seconds": max_transfer_timeout_seconds, - "proxy": provider.proxy.clone(), + "proxy": admin_secret_safe_proxy(provider.proxy.as_ref()), "stream_first_byte_timeout": provider.stream_first_byte_timeout_secs, "request_timeout": provider.request_timeout_secs, - "claude_code_advanced": config.and_then(|cfg| cfg.get("claude_code_advanced")).cloned(), - "pool_advanced": config.and_then(|cfg| cfg.get("pool_advanced")).cloned(), - "failover_rules": config.and_then(|cfg| cfg.get("failover_rules")).cloned(), - "chat_pii_redaction": config.and_then(|cfg| cfg.get("chat_pii_redaction")).cloned(), + "claude_code_advanced": admin_secret_safe_json(config.and_then(|cfg| cfg.get("claude_code_advanced"))), + "pool_advanced": admin_secret_safe_json(config.and_then(|cfg| cfg.get("pool_advanced"))), + "failover_rules": admin_secret_safe_json(config.and_then(|cfg| cfg.get("failover_rules"))), + "chat_pii_redaction": admin_secret_safe_json(config.and_then(|cfg| cfg.get("chat_pii_redaction"))), "total_endpoints": total_endpoints, "active_endpoints": active_endpoints, "total_keys": total_keys, diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/keys/create.rs b/apps/aether-gateway/src/handlers/admin/provider/write/keys/create.rs index 535b09f39..c6aa2d7fe 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/keys/create.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/keys/create.rs @@ -7,7 +7,6 @@ use crate::handlers::admin::provider::write::normalize::{ }; use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::{ - decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, normalize_json_object, normalize_string_list, parse_catalog_auth_config_json, }; use crate::handlers::shared::normalize_optional_api_key_concurrent_limit; @@ -73,6 +72,8 @@ pub(crate) async fn build_admin_create_provider_key_record( _ => {} } + let key_id = Uuid::new_v4().to_string(); + let existing_keys = state .list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id)) .await @@ -83,12 +84,10 @@ pub(crate) async fn build_admin_create_provider_key_record( .iter() .filter(|existing| raw_secret_auth_type(&existing.auth_type)) { - let Some(decrypted) = existing - .encrypted_api_key - .as_deref() - .and_then(|ciphertext| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - }) + let Some(decrypted) = state + .decrypt_provider_catalog_key_api_key(existing) + .ok() + .flatten() else { continue; }; @@ -139,8 +138,9 @@ pub(crate) async fn build_admin_create_provider_key_record( let encrypted_api_key = match auth_type.as_str() { "api_key" | "bearer" if !api_key.is_empty() => Some( - encrypt_catalog_secret_with_fallbacks(state, &api_key) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?, + state + .seal_provider_catalog_key_api_key(&provider.id, &key_id, &api_key) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string())?, ), _ => None, }; @@ -150,7 +150,12 @@ pub(crate) async fn build_admin_create_provider_key_record( .map(serde_json::to_string) .transpose() .map_err(|err| err.to_string())? - .and_then(|plaintext| encrypt_catalog_secret_with_fallbacks(state, &plaintext)); + .map(|plaintext| { + state + .seal_provider_catalog_key_auth_config(&provider.id, &key_id, &plaintext) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string()) + }) + .transpose()?; let now_unix_secs = SystemTime::now() .duration_since(UNIX_EPOCH) @@ -160,7 +165,7 @@ pub(crate) async fn build_admin_create_provider_key_record( let inherits_provider_api_formats = auth_type == "oauth" && provider_type_is_fixed(&provider.provider_type); let mut key = StoredProviderCatalogKey::new( - Uuid::new_v4().to_string(), + key_id, provider.id.clone(), name.to_string(), auth_type, diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/keys/update.rs b/apps/aether-gateway/src/handlers/admin/provider/write/keys/update.rs index 5c114e2b6..ca6f2b9c2 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/keys/update.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/keys/update.rs @@ -8,11 +8,13 @@ use crate::handlers::admin::provider::write::normalize::{ }; use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::{ - decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, json_string_list, - normalize_json_object, normalize_string_list, parse_catalog_auth_config_json, + json_string_list, normalize_json_object, normalize_string_list, parse_catalog_auth_config_json, }; use crate::handlers::shared::normalize_optional_api_key_concurrent_limit; use crate::provider_key_auth::provider_key_is_oauth_managed; +use aether_admin::provider::redaction::{ + admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, +}; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence, StoredProviderCatalogKey, StoredProviderCatalogProvider, @@ -50,6 +52,27 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( ) -> Result { let state = state.as_ref(); let mut updated = existing.clone(); + if provider.id != existing.provider_id { + updated.encrypted_api_key = state + .decrypt_provider_catalog_key_api_key(existing) + .map_err(|_| "无法验证现有 provider API Key".to_string())? + .map(|plaintext| { + state + .seal_provider_catalog_key_api_key(&provider.id, &existing.id, &plaintext) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string()) + }) + .transpose()?; + updated.encrypted_auth_config = state + .decrypt_provider_catalog_key_auth_config(existing) + .map_err(|_| "无法验证现有 provider auth_config".to_string())? + .map(|plaintext| { + state + .seal_provider_catalog_key_auth_config(&provider.id, &existing.id, &plaintext) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string()) + }) + .transpose()?; + updated.provider_id = provider.id.clone(); + } let (fields, payload) = patch.into_parts(); let auto_fetch_disabled = existing.auto_fetch_models && matches!(payload.auto_fetch_models, Some(false)); @@ -101,16 +124,10 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( .iter() .filter(|key| key.id != existing.id && raw_secret_auth_type(&key.auth_type)) { - let Some(decrypted) = - existing_key - .encrypted_api_key - .as_deref() - .and_then(|ciphertext| { - decrypt_catalog_secret_with_fallbacks( - state.encryption_key(), - ciphertext, - ) - }) + let Some(decrypted) = state + .decrypt_provider_catalog_key_api_key(existing_key) + .ok() + .flatten() else { continue; }; @@ -122,8 +139,9 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( } } updated.encrypted_api_key = Some( - encrypt_catalog_secret_with_fallbacks(state, api_key) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?, + state + .seal_provider_catalog_key_api_key(&provider.id, &existing.id, api_key) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string())?, ); } else if api_key_present { updated.encrypted_api_key = None; @@ -188,8 +206,13 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( .transpose() .map_err(|err| err.to_string())? .map(|plaintext| { - encrypt_catalog_secret_with_fallbacks(state, &plaintext) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string()) + state + .seal_provider_catalog_key_auth_config( + &provider.id, + &existing.id, + &plaintext, + ) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string()) }) .transpose()?; } @@ -348,10 +371,12 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( normalize_string_list(payload.model_exclude_patterns).map(|value| json!(value)); } if fields.contains("proxy") { - updated.proxy = normalize_json_object(payload.proxy, "proxy")?; + updated.proxy = normalize_json_object(payload.proxy, "proxy")? + .map(|value| admin_restore_secret_safe_proxy(existing.proxy.as_ref(), &value)); } if fields.contains("fingerprint") { - updated.fingerprint = normalize_json_object(payload.fingerprint, "fingerprint")?; + updated.fingerprint = normalize_json_object(payload.fingerprint, "fingerprint")? + .map(|value| admin_restore_secret_safe_json(existing.fingerprint.as_ref(), &value)); } if auth_config_present && !auth_type_switch && !raw_secret_auth_type(&updated.auth_type) { updated.encrypted_auth_config = auth_config @@ -360,8 +385,9 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys( .transpose() .map_err(|err| err.to_string())? .map(|plaintext| { - encrypt_catalog_secret_with_fallbacks(state, &plaintext) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string()) + state + .seal_provider_catalog_key_auth_config(&provider.id, &existing.id, &plaintext) + .map_err(|_| "gateway 未配置 provider key 加密密钥".to_string()) }) .transpose()?; } diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/provider/update.rs b/apps/aether-gateway/src/handlers/admin/provider/write/provider/update.rs index 778bfcc36..29715765f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/provider/update.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/provider/update.rs @@ -11,6 +11,9 @@ use crate::handlers::admin::provider::write::normalize::set_responses_websocket_ use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config; use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::shared::normalize_json_object; +use aether_admin::provider::redaction::{ + admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, +}; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider; use serde_json::json; use std::time::{SystemTime, UNIX_EPOCH}; @@ -202,7 +205,8 @@ pub(crate) async fn build_admin_update_provider_record( } if fields.contains("proxy") { - updated.proxy = normalize_json_object(payload.proxy, "proxy")?; + updated.proxy = normalize_json_object(payload.proxy, "proxy")? + .map(|value| admin_restore_secret_safe_proxy(existing.proxy.as_ref(), &value)); } if fields.contains("stream_first_byte_timeout") { @@ -233,6 +237,7 @@ pub(crate) async fn build_admin_update_provider_record( } else { let value = normalize_json_object(payload.config, "config")? .ok_or_else(|| "config 必须是 JSON 对象".to_string())?; + let value = admin_restore_secret_safe_json(existing.config.as_ref(), &value); let serde_json::Value::Object(patch_map) = value else { return Err("config 必须是 JSON 对象".to_string()); }; @@ -306,6 +311,13 @@ pub(crate) async fn build_admin_update_provider_record( let value = normalize_json_object(payload.claude_code_advanced, "claude_code_advanced")? .ok_or_else(|| "claude_code_advanced 必须是 JSON 对象".to_string())?; + let value = admin_restore_secret_safe_json( + existing + .config + .as_ref() + .and_then(|config| config.get("claude_code_advanced")), + &value, + ); config_map.insert("claude_code_advanced".to_string(), value); } } else if target_provider_type != "claude_code" { @@ -318,6 +330,13 @@ pub(crate) async fn build_admin_update_provider_record( } else { let value = normalize_pool_advanced_config(payload.pool_advanced)? .ok_or_else(|| "pool_advanced 必须是 JSON 对象".to_string())?; + let value = admin_restore_secret_safe_json( + existing + .config + .as_ref() + .and_then(|config| config.get("pool_advanced")), + &value, + ); config_map.insert("pool_advanced".to_string(), value); } } @@ -328,6 +347,13 @@ pub(crate) async fn build_admin_update_provider_record( } else { let value = normalize_json_object(payload.failover_rules, "failover_rules")? .ok_or_else(|| "failover_rules 必须是 JSON 对象".to_string())?; + let value = admin_restore_secret_safe_json( + existing + .config + .as_ref() + .and_then(|config| config.get("failover_rules")), + &value, + ); config_map.insert("failover_rules".to_string(), value); } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/reveal.rs b/apps/aether-gateway/src/handlers/admin/provider/write/reveal.rs index 51b1c280d..c215bb06d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/reveal.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/reveal.rs @@ -55,10 +55,11 @@ pub(crate) fn build_admin_reveal_key_payload( "auth_config": auth_config, })); } - let decrypted = key - .encrypted_api_key - .as_deref() - .and_then(|ciphertext| state.decrypt_catalog_secret_with_fallbacks(ciphertext)) + let decrypted = state + .app() + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten() .ok_or_else(|| { "无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。".to_string() })?; @@ -73,7 +74,9 @@ pub(crate) fn build_admin_reveal_key_payload( let decrypted = match key.encrypted_api_key.as_deref().map(str::trim) { Some(ciphertext) if !ciphertext.is_empty() => state - .decrypt_catalog_secret_with_fallbacks(ciphertext) + .app() + .decrypt_provider_catalog_key_api_key(key) + .map_err(|_| "无法解密 API Key,可能是加密密钥已更改。请重新添加该密钥。".to_string())? .ok_or_else(|| { "无法解密 API Key,可能是加密密钥已更改。请重新添加该密钥。".to_string() })?, @@ -179,14 +182,16 @@ pub(crate) async fn build_admin_export_key_payload( state: &AdminAppState<'_>, key: &StoredProviderCatalogKey, ) -> Result { - let ciphertext = key + let _ciphertext = key .encrypted_auth_config .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) .ok_or_else(|| "缺少认证配置,无法导出".to_string())?; let plaintext = state - .decrypt_catalog_secret_with_fallbacks(ciphertext) + .app() + .decrypt_provider_catalog_key_auth_config(key) + .map_err(|_| "无法解密认证配置".to_string())? .ok_or_else(|| "无法解密认证配置".to_string())?; let auth_config = serde_json::from_str::(&plaintext) .ok() @@ -229,7 +234,13 @@ pub(crate) async fn build_admin_export_key_payload( .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) - .and_then(|ciphertext| state.decrypt_catalog_secret_with_fallbacks(ciphertext)); + .and_then(|_| { + state + .app() + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten() + }); let mut payload = provider_oauth_export_payload( &provider_type, &auth_config, diff --git a/apps/aether-gateway/src/handlers/admin/request/auth.rs b/apps/aether-gateway/src/handlers/admin/request/auth.rs index c99075137..6f8912244 100644 --- a/apps/aether-gateway/src/handlers/admin/request/auth.rs +++ b/apps/aether-gateway/src/handlers/admin/request/auth.rs @@ -38,11 +38,33 @@ impl<'a> AdminAppState<'a> { self.app.upsert_oauth_provider_config(record).await } - pub(crate) async fn delete_oauth_provider_config( + pub(crate) async fn upsert_oauth_provider_config_with_force_disable( + &self, + record: &aether_data::repository::oauth_providers::UpsertOAuthProviderConfigRecord, + force_disable: bool, + ) -> Result< + Option, + GatewayError, + > { + self.app + .upsert_oauth_provider_config_with_force_disable(record, force_disable) + .await + } + + pub(crate) async fn delete_oauth_provider_config_if_unlinked( &self, provider_type: &str, ) -> Result { - self.app.delete_oauth_provider_config(provider_type).await + self.app + .delete_oauth_provider_config_if_unlinked(provider_type) + .await + } + + pub(crate) async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + self.app.has_oauth_links_for_provider(provider_type).await } pub(crate) async fn count_locked_users_if_oauth_provider_disabled( @@ -179,12 +201,27 @@ impl<'a> AdminAppState<'a> { self.app.list_admin_security_whitelist().await } - pub(crate) async fn upsert_ldap_module_config( + pub(crate) async fn compare_and_swap_ldap_module_config( &self, - config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, - ) -> Result, GatewayError> - { - self.app.upsert_ldap_module_config(config).await + expected: Option<&aether_data::repository::auth_modules::StoredLdapModuleConfig>, + replacement: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + bind_password_update: &aether_data::repository::auth_modules::LdapBindPasswordUpdate, + ) -> Result< + Option, + GatewayError, + > { + self.app + .compare_and_swap_ldap_module_config(expected, replacement, bind_password_update) + .await + } + + pub(crate) async fn delete_ldap_module_config_if_matches( + &self, + expected: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + ) -> Result { + self.app + .delete_ldap_module_config_if_matches(expected) + .await } pub(crate) async fn count_active_local_admin_users_with_valid_password( diff --git a/apps/aether-gateway/src/handlers/admin/request/capabilities.rs b/apps/aether-gateway/src/handlers/admin/request/capabilities.rs index afa3f05db..d5a488615 100644 --- a/apps/aether-gateway/src/handlers/admin/request/capabilities.rs +++ b/apps/aether-gateway/src/handlers/admin/request/capabilities.rs @@ -22,6 +22,18 @@ impl<'a> AdminAppState<'a> { crate::handlers::admin::shared::encrypt_catalog_secret_with_fallbacks(self.app, secret) } + pub(crate) fn encrypt_system_config_secret(&self, key: &str, secret: &str) -> Option { + crate::handlers::shared::encrypt_system_config_secret(self.app, key, secret) + } + + pub(crate) fn encrypt_ldap_bind_password( + &self, + config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + secret: &str, + ) -> Option { + crate::handlers::shared::encrypt_ldap_bind_password(self.app, config, secret) + } + pub(crate) fn decrypt_catalog_secret_with_fallbacks(&self, ciphertext: &str) -> Option { crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks( self.app.encryption_key(), diff --git a/apps/aether-gateway/src/handlers/admin/request/mod.rs b/apps/aether-gateway/src/handlers/admin/request/mod.rs index 8fa32e6de..70e5ca84b 100644 --- a/apps/aether-gateway/src/handlers/admin/request/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/request/mod.rs @@ -22,3 +22,4 @@ pub(crate) use self::provider_oauth::{ }; pub(crate) use self::route_request::{AdminCancelVideoTaskError, AdminRouteRequest}; pub(crate) use self::state::{AdminAppState, AdminRouteResponse, AdminRouteResult}; +pub(crate) use self::system::{SystemExportMode, ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED}; diff --git a/apps/aether-gateway/src/handlers/admin/request/models.rs b/apps/aether-gateway/src/handlers/admin/request/models.rs index 0b63f592d..452e9942c 100644 --- a/apps/aether-gateway/src/handlers/admin/request/models.rs +++ b/apps/aether-gateway/src/handlers/admin/request/models.rs @@ -16,6 +16,18 @@ use serde_json::json; use std::collections::{BTreeMap, BTreeSet}; use uuid::Uuid; +const ADMIN_MODEL_DATA_UNAVAILABLE_DETAIL: &str = "Model data temporarily unavailable"; + +fn admin_model_repository_error(operation: &'static str) -> String { + tracing::error!( + event_name = "admin_model_repository_error", + operation, + error_category = "repository_unavailable", + "admin model repository operation failed" + ); + ADMIN_MODEL_DATA_UNAVAILABLE_DETAIL.to_string() +} + fn normalize_provider_model_mapping_scopes( value: Option, ) -> Option { @@ -103,7 +115,7 @@ impl<'a> AdminAppState<'a> { { self.get_admin_global_model_by_id(global_model_id) .await - .map_err(|err| format!("{err:?}"))? + .map_err(|_| admin_model_repository_error("resolve_global_model_by_id"))? .ok_or_else(|| format!("GlobalModel {global_model_id} 不存在")) } @@ -144,7 +156,7 @@ impl<'a> AdminAppState<'a> { if self .admin_provider_model_name_exists(provider_id, &provider_model_name, None) .await - .map_err(|err| format!("{err:?}"))? + .map_err(|_| admin_model_repository_error("check_provider_model_name"))? { return Err(format!("模型 '{provider_model_name}' 已存在")); } @@ -202,7 +214,7 @@ impl<'a> AdminAppState<'a> { if self .admin_provider_model_name_exists(&existing.provider_id, &name, Some(&existing.id)) .await - .map_err(|err| format!("{err:?}"))? + .map_err(|_| admin_model_repository_error("check_provider_model_name"))? { return Err(format!("模型 '{name}' 已存在")); } @@ -311,7 +323,7 @@ impl<'a> AdminAppState<'a> { limit: 10_000, }) .await - .map_err(|err| format!("{err:?}"))?; + .map_err(|_| admin_model_repository_error("list_provider_models_for_import"))?; let mut existing_by_name = existing_models .iter() .map(|model| (model.provider_model_name.clone(), model.clone())) @@ -350,7 +362,7 @@ impl<'a> AdminAppState<'a> { let global_model = if let Some(existing) = self .get_admin_global_model_by_name(&trimmed) .await - .map_err(|err| format!("{err:?}"))? + .map_err(|_| admin_model_repository_error("lookup_global_model_for_import"))? { existing } else { @@ -365,7 +377,7 @@ impl<'a> AdminAppState<'a> { .map_err(|err| err.to_string())?, ) .await - .map_err(|err| format!("{err:?}"))?; + .map_err(|_| admin_model_repository_error("create_global_model_for_import"))?; let Some(created) = created else { errors.push(json!({"model_id": trimmed, "error": "Create GlobalModel failed"})); continue; @@ -399,9 +411,9 @@ impl<'a> AdminAppState<'a> { "model_id": trimmed, "error": "Create provider model failed", })), - Err(err) => errors.push(json!({ + Err(_) => errors.push(json!({ "model_id": trimmed, - "error": format!("{err:?}"), + "error": admin_model_repository_error("create_provider_model_for_import"), })), } } @@ -425,7 +437,7 @@ impl<'a> AdminAppState<'a> { limit: 10_000, }) .await - .map_err(|err| format!("{err:?}"))?; + .map_err(|_| admin_model_repository_error("list_provider_models_for_assignment"))?; let existing_global_model_ids = existing_models .into_iter() .map(|model| model.global_model_id) @@ -475,9 +487,11 @@ impl<'a> AdminAppState<'a> { "global_model_id": global_model.id, "error": "Create provider model failed", })), - Err(err) => errors.push(json!({ + Err(_) => errors.push(json!({ "global_model_id": global_model.id, - "error": format!("{err:?}"), + "error": admin_model_repository_error( + "create_provider_model_for_assignment" + ), })), } } diff --git a/apps/aether-gateway/src/handlers/admin/request/provider/catalog.rs b/apps/aether-gateway/src/handlers/admin/request/provider/catalog.rs index cd57b5f37..0b73678a5 100644 --- a/apps/aether-gateway/src/handlers/admin/request/provider/catalog.rs +++ b/apps/aether-gateway/src/handlers/admin/request/provider/catalog.rs @@ -430,6 +430,15 @@ impl<'a> AdminAppState<'a> { self.app.update_provider_catalog_provider(provider).await } + pub(crate) async fn compare_and_swap_provider_catalog_provider_config( + &self, + update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + self.app + .compare_and_swap_provider_catalog_provider_config(update) + .await + } + pub(crate) async fn cleanup_deleted_provider_catalog_refs( &self, provider_id: &str, diff --git a/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs b/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs index 3306f8764..c3517200b 100644 --- a/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs +++ b/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs @@ -5,15 +5,15 @@ use aether_contracts::{ EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, }; use aether_data::repository::provider_oauth::{ - build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key, + build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_secret_purpose, + provider_oauth_batch_task_storage_key, provider_oauth_device_session_secret_purpose, provider_oauth_device_session_storage_key, provider_oauth_state_storage_key, StoredAdminProviderOAuthDeviceSession, StoredAdminProviderOAuthState, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, PROVIDER_OAUTH_STATE_TTL_SECS, }; use axum::http; -use base64::{engine::general_purpose::STANDARD, Engine as _}; use flate2::read::{DeflateDecoder, GzDecoder}; -use serde_json::json; +use serde_json::{json, Value}; use std::collections::BTreeMap; use std::io::Read; use url::Url; @@ -21,6 +21,9 @@ use url::Url; const KIRO_IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE"; const ADMIN_PROVIDER_OAUTH_TIMEOUT_MS: u64 = 30_000; const ADMIN_PROVIDER_OAUTH_PROXY_TIMEOUT_MS: u64 = 60_000; +const ADMIN_PROVIDER_OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; +const ADMIN_PROVIDER_OAUTH_STATE_SECRET_PURPOSE: &str = "provider-oauth-state"; +const ADMIN_PROVIDER_OAUTH_STATE_MAX_CLOCK_SKEW_SECS: u64 = 60; pub(crate) struct AdminProviderOAuthHttpResponse { pub(crate) status: http::StatusCode, @@ -28,30 +31,177 @@ pub(crate) struct AdminProviderOAuthHttpResponse { pub(crate) json_body: Option, } -impl<'a> AdminAppState<'a> { - pub(crate) async fn update_provider_catalog_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - crate::oauth::ProviderOAuthRepository::update_provider_catalog_key_oauth_credentials( - self, - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } +fn admin_provider_oauth_state_secret_purpose(nonce: &str) -> String { + format!( + "{ADMIN_PROVIDER_OAUTH_STATE_SECRET_PURPOSE}:{}", + provider_oauth_state_storage_key(nonce.trim()) + ) +} +fn is_generated_admin_provider_oauth_nonce(nonce: &str) -> bool { + let nonce = nonce.trim(); + nonce.len() == 64 + && nonce + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) +} + +fn decode_admin_provider_oauth_state( + state: &AdminAppState<'_>, + expected_nonce: &str, + stored: &str, +) -> Result { + let expected_nonce = expected_nonce.trim(); + let purpose = admin_provider_oauth_state_secret_purpose(expected_nonce); + let plaintext = + crate::handlers::shared::open_runtime_secret_payload(state.as_ref(), &purpose, stored) + .ok_or_else(invalid_admin_provider_oauth_state_error)?; + let record = serde_json::from_str::(&plaintext) + .map_err(|_| invalid_admin_provider_oauth_state_error())?; + validate_admin_provider_oauth_state(expected_nonce, &record)?; + Ok(record) +} + +fn invalid_admin_provider_oauth_state_error() -> GatewayError { + GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + message: "provider OAuth state is invalid".to_string(), + } +} + +fn invalid_admin_provider_oauth_device_session_error() -> GatewayError { + GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + message: "provider OAuth device session is invalid".to_string(), + } +} + +fn is_valid_admin_provider_oauth_ephemeral_id(value: &str) -> bool { + !value.is_empty() + && value.len() <= 128 + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b':')) +} + +fn validate_admin_provider_oauth_device_session( + expected_session_id: &str, + record: &StoredAdminProviderOAuthDeviceSession, +) -> Result<(), GatewayError> { + let raw_expected_session_id = expected_session_id; + let expected_session_id = expected_session_id.trim(); + let optional_identity_is_invalid = |value: Option<&str>| { + value.is_some_and(|value| value.trim().is_empty() || value != value.trim()) + }; + if raw_expected_session_id != expected_session_id + || !is_valid_admin_provider_oauth_ephemeral_id(expected_session_id) + || record.session_id != expected_session_id + || !is_valid_admin_provider_oauth_ephemeral_id(&record.session_id) + || record.provider_id.trim().is_empty() + || record.provider_id != record.provider_id.trim() + || record.initiated_by_user_id.trim().is_empty() + || record.initiated_by_user_id != record.initiated_by_user_id.trim() + || optional_identity_is_invalid(record.initiated_by_session_id.as_deref()) + || optional_identity_is_invalid(record.initiated_by_management_token_id.as_deref()) + || (record.initiated_by_session_id.is_none() + && record.initiated_by_management_token_id.is_none()) + || !matches!( + record.status.as_str(), + "pending" | "authorized" | "expired" | "error" + ) + || record.created_at_unix_ms == 0 + || record.expires_at_unix_secs < record.created_at_unix_ms + || (record.status == "authorized" + && record + .key_id + .as_deref() + .map(str::trim) + .is_none_or(str::is_empty)) + { + return Err(invalid_admin_provider_oauth_device_session_error()); + } + Ok(()) +} + +fn decode_admin_provider_oauth_device_session( + state: &AdminAppState<'_>, + expected_session_id: &str, + stored: &str, +) -> Result { + let purpose = provider_oauth_device_session_secret_purpose(expected_session_id); + let plaintext = + crate::handlers::shared::open_runtime_secret_payload(state.as_ref(), &purpose, stored) + .ok_or_else(invalid_admin_provider_oauth_device_session_error)?; + let record = serde_json::from_str::(&plaintext) + .map_err(|_| invalid_admin_provider_oauth_device_session_error())?; + validate_admin_provider_oauth_device_session(expected_session_id, &record)?; + Ok(record) +} + +fn decode_admin_provider_oauth_batch_task( + state: &AdminAppState<'_>, + expected_task_id: &str, + stored: &str, +) -> Result { + let invalid = || GatewayError::Internal("provider OAuth batch task is invalid".to_string()); + let purpose = provider_oauth_batch_task_secret_purpose(expected_task_id); + let plaintext = + crate::handlers::shared::open_runtime_secret_payload(state.as_ref(), &purpose, stored) + .ok_or_else(invalid)?; + let parsed = serde_json::from_str::(&plaintext).map_err(|_| invalid())?; + let state = parsed.as_object().ok_or_else(invalid)?; + if state.get("task_id").and_then(serde_json::Value::as_str) != Some(expected_task_id) { + return Err(invalid()); + } + Ok(parsed) +} + +fn validate_admin_provider_oauth_state( + expected_nonce: &str, + record: &StoredAdminProviderOAuthState, +) -> Result<(), GatewayError> { + let invalid = invalid_admin_provider_oauth_state_error; + if !is_generated_admin_provider_oauth_nonce(expected_nonce) + || record.nonce != expected_nonce + || !is_generated_admin_provider_oauth_nonce(&record.nonce) + || record.provider_id.trim().is_empty() + || record.provider_id != record.provider_id.trim() + || record.provider_type.trim().is_empty() + || record.provider_type != record.provider_type.trim().to_ascii_lowercase() + || record.initiated_by_user_id.trim().is_empty() + || record.initiated_by_user_id != record.initiated_by_user_id.trim() + || record + .pkce_verifier + .as_deref() + .is_some_and(|value| value.trim().is_empty()) + || record + .initiated_by_session_id + .as_deref() + .is_some_and(|value| value.trim().is_empty()) + || record + .initiated_by_management_token_id + .as_deref() + .is_some_and(|value| value.trim().is_empty()) + || (record.initiated_by_session_id.is_none() + && record.initiated_by_management_token_id.is_none()) + { + return Err(invalid()); + } + let now = aether_admin::provider::state::current_unix_secs(); + if record.created_at > now.saturating_add(ADMIN_PROVIDER_OAUTH_STATE_MAX_CLOCK_SKEW_SECS) + || now.saturating_sub(record.created_at) > PROVIDER_OAUTH_STATE_TTL_SECS + { + return Err(invalid()); + } + Ok(()) +} + +impl<'a> AdminAppState<'a> { pub(crate) async fn update_provider_catalog_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { self.app @@ -59,7 +209,6 @@ impl<'a> AdminAppState<'a> { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -91,21 +240,36 @@ impl<'a> AdminAppState<'a> { provider_type: &str, pkce_verifier: Option<&str>, expected_encrypted_auth_config: Option<&str>, + initiated_by_user_id: &str, + initiated_by_session_id: Option<&str>, + initiated_by_management_token_id: Option<&str>, ) -> Result { let nonce = aether_admin::provider::state::generate_provider_oauth_nonce(); - let payload = json!({ - "nonce": nonce, - "key_id": key_id, - "provider_id": provider_id, - "provider_type": provider_type, - "pkce_verifier": pkce_verifier, - "expected_encrypted_auth_config": expected_encrypted_auth_config, - "created_at": aether_admin::provider::state::current_unix_secs(), - }); + let record = StoredAdminProviderOAuthState { + nonce: nonce.clone(), + key_id: key_id.to_string(), + provider_id: provider_id.to_string(), + provider_type: provider_type.to_string(), + pkce_verifier: pkce_verifier.map(ToOwned::to_owned), + expected_encrypted_auth_config: expected_encrypted_auth_config.map(ToOwned::to_owned), + initiated_by_user_id: initiated_by_user_id.to_string(), + initiated_by_session_id: initiated_by_session_id.map(ToOwned::to_owned), + initiated_by_management_token_id: initiated_by_management_token_id + .map(ToOwned::to_owned), + created_at: aether_admin::provider::state::current_unix_secs(), + }; + validate_admin_provider_oauth_state(&nonce, &record)?; let key = provider_oauth_state_storage_key(&nonce); - let value = payload.to_string(); + let value = serde_json::to_string(&record) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let purpose = admin_provider_oauth_state_secret_purpose(&nonce); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(self.as_ref(), &purpose, &value) + .ok_or_else(|| { + GatewayError::Internal("provider OAuth state encryption unavailable".to_string()) + })?; self.as_ref() - .runtime_kv_setex(&key, &value, PROVIDER_OAUTH_STATE_TTL_SECS) + .runtime_kv_setex(&key, &sealed, PROVIDER_OAUTH_STATE_TTL_SECS) .await?; self.as_ref() .save_provider_oauth_state_for_tests(&key, &value); @@ -118,11 +282,18 @@ impl<'a> AdminAppState<'a> { ) -> Result, GatewayError> { let key = provider_oauth_state_storage_key(nonce); let raw = self.as_ref().runtime_kv_getdel(&key).await?; - raw.map(|value| { - serde_json::from_str::(&value) - .map_err(|err| GatewayError::Internal(err.to_string())) - }) - .transpose() + raw.map(|value| decode_admin_provider_oauth_state(self, nonce, &value)) + .transpose() + } + + pub(crate) async fn load_provider_oauth_state( + &self, + nonce: &str, + ) -> Result, GatewayError> { + let key = provider_oauth_state_storage_key(nonce); + let raw = self.as_ref().runtime_kv_get(&key).await?; + raw.map(|value| decode_admin_provider_oauth_state(self, nonce, &value)) + .transpose() } pub(crate) async fn exchange_admin_provider_oauth_code( @@ -164,12 +335,37 @@ impl<'a> AdminAppState<'a> { task_id: &str, task_state: &serde_json::Value, ) -> Result<(), GatewayError> { + let Some(state) = task_state.as_object() else { + return Err(GatewayError::Internal( + "provider OAuth batch task is invalid".to_string(), + )); + }; + if !is_valid_admin_provider_oauth_ephemeral_id(task_id) + || state.get("task_id").and_then(serde_json::Value::as_str) != Some(task_id) + || state + .get("provider_id") + .and_then(serde_json::Value::as_str) + .is_none_or(|value| value.trim().is_empty() || value != value.trim()) + { + return Err(GatewayError::Internal( + "provider OAuth batch task is invalid".to_string(), + )); + } let key = provider_oauth_batch_task_storage_key(task_id); let serialized = serde_json::to_string(task_state) .map_err(|err| GatewayError::Internal(err.to_string()))?; + let purpose = provider_oauth_batch_task_secret_purpose(task_id); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + self.as_ref(), + &purpose, + &serialized, + ) + .ok_or_else(|| { + GatewayError::Internal("provider OAuth batch task encryption unavailable".to_string()) + })?; self.as_ref() - .runtime_kv_setex(&key, &serialized, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS) + .runtime_kv_setex(&key, &sealed, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS) .await?; self.as_ref() .save_provider_oauth_batch_task_for_tests(&key, &serialized); @@ -181,17 +377,20 @@ impl<'a> AdminAppState<'a> { provider_id: &str, task_id: &str, ) -> Result, GatewayError> { + if provider_id.trim().is_empty() + || provider_id != provider_id.trim() + || !is_valid_admin_provider_oauth_ephemeral_id(task_id) + { + return Ok(None); + } let key = provider_oauth_batch_task_storage_key(task_id); let raw = self.as_ref().runtime_kv_get(&key).await?; let Some(raw) = raw else { return Ok(None); }; - let parsed = match serde_json::from_str::(&raw) { - Ok(value) => value, - Err(_) => return Ok(None), - }; + let parsed = decode_admin_provider_oauth_batch_task(self, task_id, &raw)?; let Some(state) = parsed.as_object() else { - return Ok(None); + unreachable!("validated provider OAuth batch task must be an object"); }; if state .get("provider_id") @@ -213,6 +412,12 @@ impl<'a> AdminAppState<'a> { session: &StoredAdminProviderOAuthDeviceSession, ttl_seconds: u64, ) -> Result<(), Response> { + if validate_admin_provider_oauth_device_session(session_id, session).is_err() { + return Err(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "provider oauth device session is invalid", + )); + } let key = provider_oauth_device_session_storage_key(session_id); let value = serde_json::to_string(session).map_err(|_| { build_internal_control_error_response( @@ -220,8 +425,17 @@ impl<'a> AdminAppState<'a> { "provider oauth redis unavailable", ) })?; + let purpose = provider_oauth_device_session_secret_purpose(session_id); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(self.as_ref(), &purpose, &value) + .ok_or_else(|| { + build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth session encryption unavailable", + ) + })?; self.as_ref() - .runtime_kv_setex(&key, &value, ttl_seconds) + .runtime_kv_setex(&key, &sealed, ttl_seconds) .await .map_err(|_| { build_internal_control_error_response( @@ -238,13 +452,15 @@ impl<'a> AdminAppState<'a> { &self, session_id: &str, ) -> Result, GatewayError> { + if session_id != session_id.trim() + || !is_valid_admin_provider_oauth_ephemeral_id(session_id) + { + return Err(invalid_admin_provider_oauth_device_session_error()); + } let key = provider_oauth_device_session_storage_key(session_id); let raw = self.as_ref().runtime_kv_get(&key).await?; - raw.map(|value| { - serde_json::from_str::(&value) - .map_err(|err| GatewayError::Internal(err.to_string())) - }) - .transpose() + raw.map(|value| decode_admin_provider_oauth_device_session(self, session_id, &value)) + .transpose() } pub(crate) async fn register_admin_kiro_device_oidc_client( @@ -253,6 +469,7 @@ impl<'a> AdminAppState<'a> { start_url: &str, proxy: Option, ) -> Result> { + let region = aether_provider_transport::kiro::normalize_kiro_region(region); let payload = post_kiro_device_oidc_json( self, "kiro_device_register", @@ -281,14 +498,10 @@ impl<'a> AdminAppState<'a> { .and_then(serde_json::Value::as_bool) .unwrap_or(false) { - let error_desc = aether_admin::provider::state::json_non_empty_string( - payload.get("error_description"), - ) - .or_else(|| aether_admin::provider::state::json_non_empty_string(payload.get("error"))) - .unwrap_or_else(|| "unknown".to_string()); + let error_code = kiro_device_oidc_error_code(&payload); return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - format!("注册 OIDC 客户端失败: {error_desc}"), + format!("注册 OIDC 客户端失败: {error_code}"), )); } Ok(payload) @@ -302,6 +515,7 @@ impl<'a> AdminAppState<'a> { start_url: &str, proxy: Option, ) -> Result> { + let region = aether_provider_transport::kiro::normalize_kiro_region(region); let payload = post_kiro_device_oidc_json( self, "kiro_device_authorize", @@ -319,14 +533,10 @@ impl<'a> AdminAppState<'a> { .and_then(serde_json::Value::as_bool) .unwrap_or(false) { - let error_desc = aether_admin::provider::state::json_non_empty_string( - payload.get("error_description"), - ) - .or_else(|| aether_admin::provider::state::json_non_empty_string(payload.get("error"))) - .unwrap_or_else(|| "unknown".to_string()); + let error_code = kiro_device_oidc_error_code(&payload); return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - format!("发起设备授权失败: {error_desc}"), + format!("发起设备授权失败: {error_code}"), )); } Ok(payload) @@ -340,6 +550,7 @@ impl<'a> AdminAppState<'a> { device_code: &str, proxy: Option, ) -> Result> { + let region = aether_provider_transport::kiro::normalize_kiro_region(region); post_kiro_device_oidc_json( self, "kiro_device_poll", @@ -505,29 +716,48 @@ async fn post_kiro_device_oidc_json( "发起设备授权失败: unknown", ) })?; - let status = response.status; - let body_text = response.body_text; - Ok( - match serde_json::from_str::(&body_text) { - Ok(mut payload) => { - if !status.is_success() { - if let Some(object) = payload.as_object_mut() { - object.insert("_error".to_string(), json!(true)); - } else { - payload = json!({ - "_error": true, - "data": payload, - }); - } - } - payload - } - Err(_) => json!({ - "_error": !status.is_success(), - "error": body_text.trim(), - }), - }, - ) + Ok(project_kiro_device_oidc_response( + response.status, + &response.body_text, + )) +} + +fn project_kiro_device_oidc_response(status: http::StatusCode, body_text: &str) -> Value { + match serde_json::from_str::(body_text) { + Ok(payload) if status.is_success() => payload, + Ok(payload) => json!({ + "_error": true, + "error": kiro_device_oidc_error_code(&payload), + }), + Err(_) => json!({ + "_error": true, + "error": if status.is_success() { + "invalid_response" + } else { + "upstream_error" + }, + }), + } +} + +fn kiro_device_oidc_error_code(payload: &Value) -> &'static str { + let Some(error_code) = payload.get("error").and_then(Value::as_str).map(str::trim) else { + return "upstream_error"; + }; + match error_code { + "access_denied" => "access_denied", + "authorization_pending" => "authorization_pending", + "expired_token" => "expired_token", + "invalid_client" => "invalid_client", + "invalid_client_metadata" => "invalid_client_metadata", + "invalid_grant" => "invalid_grant", + "invalid_redirect_uri" => "invalid_redirect_uri", + "invalid_request" => "invalid_request", + "invalid_scope" => "invalid_scope", + "slow_down" => "slow_down", + "unauthorized_client" => "unauthorized_client", + _ => "upstream_error", + } } impl<'a> AdminAppState<'a> { @@ -647,10 +877,13 @@ fn admin_provider_oauth_execution_body_bytes( headers: &BTreeMap, body: &aether_contracts::ResponseBody, ) -> Option> { - let bytes = body - .body_bytes_b64 - .as_deref() - .and_then(|value| STANDARD.decode(value).ok())?; + let bytes = body.body_bytes_b64.as_deref().and_then(|value| { + crate::execution_runtime::transport::decode_base64_body_with_limit( + value, + ADMIN_PROVIDER_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + })?; admin_provider_oauth_decode_response_bytes( &bytes, headers.get("content-encoding").map(String::as_str), @@ -669,20 +902,287 @@ fn admin_provider_oauth_decode_response_bytes( match encoding.as_deref() { Some("gzip") => { let mut decoder = GzDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) + read_admin_provider_oauth_decoder_with_limit( + &mut decoder, + ADMIN_PROVIDER_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) } Some("deflate") => { let mut decoder = DeflateDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) + read_admin_provider_oauth_decoder_with_limit( + &mut decoder, + ADMIN_PROVIDER_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) } _ => None, } } +fn read_admin_provider_oauth_decoder_with_limit( + decoder: &mut impl Read, + limit_bytes: usize, +) -> Option> { + let read_limit = u64::try_from(limit_bytes) + .unwrap_or(u64::MAX) + .saturating_add(1); + let mut limited = decoder.take(read_limit); + let mut out = Vec::new(); + limited.read_to_end(&mut out).ok()?; + (out.len() <= limit_bytes).then_some(out) +} + fn admin_provider_oauth_gateway_error_message(error: GatewayError) -> String { error.into_message() } + +#[cfg(test)] +mod response_decode_tests { + use std::io::Cursor; + + use super::{ + admin_provider_oauth_state_secret_purpose, decode_admin_provider_oauth_batch_task, + decode_admin_provider_oauth_device_session, decode_admin_provider_oauth_state, + project_kiro_device_oidc_response, read_admin_provider_oauth_decoder_with_limit, + }; + use crate::{data::GatewayDataState, AppState}; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::provider_oauth::{ + provider_oauth_batch_task_secret_purpose, provider_oauth_device_session_secret_purpose, + StoredAdminProviderOAuthDeviceSession, StoredAdminProviderOAuthState, + }; + use axum::http::StatusCode; + use serde_json::json; + + #[test] + fn admin_oauth_decoder_accepts_exact_limit_and_rejects_limit_plus_one() { + let mut exact = Cursor::new(vec![b'x'; 8]); + assert_eq!( + read_admin_provider_oauth_decoder_with_limit(&mut exact, 8), + Some(vec![b'x'; 8]) + ); + + let mut oversized = Cursor::new(vec![b'x'; 9]); + assert!(read_admin_provider_oauth_decoder_with_limit(&mut oversized, 8).is_none()); + } + + #[test] + fn admin_provider_oauth_state_ciphertext_is_bound_to_its_nonce() { + let app = AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let admin = super::AdminAppState::new(&app); + let record = StoredAdminProviderOAuthState { + nonce: "a".repeat(64), + key_id: "key-1".to_string(), + provider_id: "provider-1".to_string(), + provider_type: "codex".to_string(), + pkce_verifier: Some("verifier".to_string()), + expected_encrypted_auth_config: None, + initiated_by_user_id: "admin-1".to_string(), + initiated_by_session_id: Some("session-1".to_string()), + initiated_by_management_token_id: None, + created_at: aether_admin::provider::state::current_unix_secs(), + }; + let plaintext = serde_json::to_string(&record).expect("state should serialize"); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &app, + &admin_provider_oauth_state_secret_purpose(&record.nonce), + &plaintext, + ) + .expect("state should seal"); + + assert_eq!( + decode_admin_provider_oauth_state(&admin, &record.nonce, &sealed) + .expect("matching state should open"), + record + ); + assert!(matches!( + decode_admin_provider_oauth_state(&admin, &"b".repeat(64), &sealed), + Err(crate::GatewayError::Client { + status: StatusCode::BAD_REQUEST, + .. + }) + )); + assert!(matches!( + decode_admin_provider_oauth_state(&admin, &record.nonce, "damaged-ciphertext"), + Err(crate::GatewayError::Client { + status: StatusCode::BAD_REQUEST, + .. + }) + )); + } + + #[test] + fn admin_provider_oauth_device_session_ciphertext_is_bound_to_session_and_principal() { + let app = AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let admin = super::AdminAppState::new(&app); + let now = aether_admin::provider::state::current_unix_secs(); + let record = StoredAdminProviderOAuthDeviceSession { + session_id: "session-123".to_string(), + provider_id: "provider-1".to_string(), + initiated_by_user_id: "admin-1".to_string(), + initiated_by_session_id: Some("admin-session-1".to_string()), + initiated_by_management_token_id: None, + region: "us-east-1".to_string(), + client_id: "client-1".to_string(), + client_secret: "secret-1".to_string(), + device_code: "device-code-1".to_string(), + auth_type: Some("idc".to_string()), + social_provider: None, + code_verifier: None, + redirect_uri: None, + machine_id: None, + interval: 5, + expires_at_unix_secs: now.saturating_add(600), + status: "pending".to_string(), + proxy_node_id: None, + created_at_unix_ms: now, + key_id: None, + email: None, + replaced: false, + error_msg: None, + }; + let plaintext = serde_json::to_string(&record).expect("device session should serialize"); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &app, + &provider_oauth_device_session_secret_purpose(&record.session_id), + &plaintext, + ) + .expect("device session should seal"); + + let decoded = + decode_admin_provider_oauth_device_session(&admin, &record.session_id, &sealed) + .expect("matching device session should open"); + assert_eq!(decoded.session_id, record.session_id); + assert_eq!(decoded.initiated_by_user_id, record.initiated_by_user_id); + assert_eq!( + decoded.initiated_by_session_id, + record.initiated_by_session_id + ); + assert!(matches!( + decode_admin_provider_oauth_device_session(&admin, "session-456", &sealed), + Err(crate::GatewayError::Client { + status: StatusCode::BAD_REQUEST, + .. + }) + )); + assert!(matches!( + decode_admin_provider_oauth_device_session( + &admin, + &record.session_id, + "damaged-ciphertext" + ), + Err(crate::GatewayError::Client { + status: StatusCode::BAD_REQUEST, + .. + }) + )); + + let mut mismatched = record.clone(); + mismatched.session_id = "session-456".to_string(); + let mismatched_plaintext = + serde_json::to_string(&mismatched).expect("mismatched session should serialize"); + let mismatched_sealed = crate::handlers::shared::seal_runtime_secret_payload( + &app, + &provider_oauth_device_session_secret_purpose("session-123"), + &mismatched_plaintext, + ) + .expect("mismatched session should seal"); + assert!(matches!( + decode_admin_provider_oauth_device_session(&admin, "session-123", &mismatched_sealed), + Err(crate::GatewayError::Client { + status: StatusCode::BAD_REQUEST, + .. + }) + )); + } + + #[test] + fn admin_provider_oauth_batch_task_ciphertext_is_bound_to_task_id() { + let app = AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let admin = super::AdminAppState::new(&app); + let task = json!({ + "task_id": "task-123", + "provider_id": "provider-1", + "provider_type": "codex", + "status": "processing", + }); + let plaintext = task.to_string(); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &app, + &provider_oauth_batch_task_secret_purpose("task-123"), + &plaintext, + ) + .expect("batch task should seal"); + + assert_eq!( + decode_admin_provider_oauth_batch_task(&admin, "task-123", &sealed) + .expect("matching batch task should open"), + task + ); + assert!(decode_admin_provider_oauth_batch_task(&admin, "task-456", &sealed).is_err()); + + let mismatched = json!({ + "task_id": "task-456", + "provider_id": "provider-1", + "status": "processing", + }); + let mismatched_sealed = crate::handlers::shared::seal_runtime_secret_payload( + &app, + &provider_oauth_batch_task_secret_purpose("task-123"), + &mismatched.to_string(), + ) + .expect("mismatched batch task should seal"); + assert!( + decode_admin_provider_oauth_batch_task(&admin, "task-123", &mismatched_sealed).is_err() + ); + } + + #[test] + fn kiro_oidc_error_projection_discards_upstream_free_text() { + let known = project_kiro_device_oidc_response( + StatusCode::BAD_REQUEST, + r#"{ + "error": "authorization_pending", + "error_description": "Bearer upstream-secret at https://internal.test" + }"#, + ); + assert_eq!( + known, + json!({"_error": true, "error": "authorization_pending"}) + ); + + let unknown = project_kiro_device_oidc_response( + StatusCode::BAD_REQUEST, + r#"{ + "error": "Bearer-upstream-secret", + "error_description": "https://user:password@internal.test" + }"#, + ); + assert_eq!(unknown, json!({"_error": true, "error": "upstream_error"})); + + let non_json = project_kiro_device_oidc_response( + StatusCode::BAD_GATEWAY, + "Bearer upstream-secret at https://internal.test", + ); + assert_eq!(non_json, json!({"_error": true, "error": "upstream_error"})); + + let exposed = format!("{known}{unknown}{non_json}"); + for secret in ["upstream-secret", "password", "internal.test", "Bearer"] { + assert!(!exposed.contains(secret)); + } + } +} diff --git a/apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs b/apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs index 4b3b7dab7..fece72492 100644 --- a/apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs +++ b/apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs @@ -20,6 +20,17 @@ fn provider_requires_credential_cas_cleanup(provider: &StoredProviderCatalogProv provider.provider_type.trim().eq_ignore_ascii_case("codex") } +fn build_model_sync_failure_payload(requested: usize) -> Value { + json!({ + "requested": requested, + "attempted": 0, + "succeeded": 0, + "failed": requested, + "skipped": 0, + "error": "Model synchronization failed", + }) +} + impl<'a> AdminAppState<'a> { pub(crate) async fn clear_admin_provider_pool_cooldown(&self, provider_id: &str, key_id: &str) { crate::handlers::admin::provider::pool::runtime::clear_admin_provider_pool_cooldown( @@ -185,8 +196,12 @@ impl<'a> AdminAppState<'a> { .collect::>(); let mut known_api_keys = existing_keys .iter() - .filter_map(|key| key.encrypted_api_key.as_deref()) - .filter_map(|ciphertext| self.decrypt_catalog_secret_with_fallbacks(ciphertext)) + .filter_map(|key| { + self.app() + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten() + }) .filter(|value| value != "__placeholder__") .collect::>(); let mut imported = 0usize; @@ -281,7 +296,10 @@ impl<'a> AdminAppState<'a> { continue; } - let Some(encrypted_api_key) = self.encrypt_catalog_secret_with_fallbacks(api_key) + let key_id = uuid::Uuid::new_v4().to_string(); + let Ok(encrypted_api_key) = + self.app() + .seal_provider_catalog_key_api_key(&provider.id, &key_id, api_key) else { errors.push(json!({ "index": index, @@ -290,7 +308,7 @@ impl<'a> AdminAppState<'a> { continue; }; let record = match admin_provider_pool_pure::build_admin_pool_batch_import_key_record( - uuid::Uuid::new_v4().to_string(), + key_id, provider.id.clone(), name.to_string(), auth_type, @@ -819,7 +837,7 @@ impl<'a> AdminAppState<'a> { .unwrap_or(0); let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize).clamp(1, 50_000); - if let Err(err) = ensure_provider_key_pool_scores_for_keys( + if ensure_provider_key_pool_scores_for_keys( self.as_ref(), &provider, &pool_config, @@ -829,11 +847,12 @@ impl<'a> AdminAppState<'a> { score_ensure_budget, ) .await + .is_err() { tracing::debug!( provider_id = %provider.id, updated_keys = updated_keys.len(), - error = ?err, + error_category = "pool_score_repository_unavailable", "gateway admin provider key batch update: failed to seed pool score rows" ); } @@ -853,14 +872,7 @@ impl<'a> AdminAppState<'a> { "failed": summary.failed, "skipped": summary.skipped, }), - Err(err) => json!({ - "requested": requested, - "attempted": 0, - "succeeded": 0, - "failed": requested, - "skipped": 0, - "error": err.into_message(), - }), + Err(_) => build_model_sync_failure_payload(requested), } }; @@ -876,7 +888,13 @@ impl<'a> AdminAppState<'a> { #[cfg(test)] mod automatic_cleanup_tests { - use super::{provider_requires_credential_cas_cleanup, AdminAppState}; + use super::{ + build_model_sync_failure_payload, provider_requires_credential_cas_cleanup, AdminAppState, + }; + use crate::handlers::shared::{ + seal_provider_catalog_credential, ProviderCatalogCredentialField, + }; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider, @@ -900,11 +918,49 @@ mod automatic_cleanup_tests { assert!(!provider_requires_credential_cas_cleanup(&provider("kiro"))); } + #[test] + fn model_sync_failure_payload_does_not_accept_internal_error_details() { + assert_eq!( + build_model_sync_failure_payload(3), + serde_json::json!({ + "requested": 3, + "attempted": 0, + "succeeded": 0, + "failed": 3, + "skipped": 0, + "error": "Model synchronization failed", + }) + ); + } + #[tokio::test] async fn cleanup_removes_existing_hard_invalid_codex_key() { let provider = provider("codex"); + let credential_state = crate::AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let key_id = "key-codex-invalid"; + let encrypted_api_key = seal_provider_catalog_credential( + &credential_state, + &provider.id, + key_id, + ProviderCatalogCredentialField::ApiKey, + "test-api-key", + ) + .expect("api key should seal"); + let encrypted_auth_config = seal_provider_catalog_credential( + &credential_state, + &provider.id, + key_id, + ProviderCatalogCredentialField::AuthConfig, + r#"{"provider_type":"codex"}"#, + ) + .expect("auth config should seal"); let mut key = StoredProviderCatalogKey::new( - "key-codex-invalid".to_string(), + key_id.to_string(), provider.id.clone(), "Invalid Codex OAuth".to_string(), "oauth".to_string(), @@ -914,8 +970,8 @@ mod automatic_cleanup_tests { .expect("key should build") .with_transport_fields( None, - "encrypted-api-key".to_string(), - Some("encrypted-auth-config".to_string()), + encrypted_api_key, + Some(encrypted_auth_config), None, None, None, @@ -939,7 +995,8 @@ mod automatic_cleanup_tests { .with_data_state_for_tests( crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( repository.clone(), - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let admin_state = AdminAppState::new(&state); diff --git a/apps/aether-gateway/src/handlers/admin/request/provider/transport.rs b/apps/aether-gateway/src/handlers/admin/request/provider/transport.rs index 1d983d563..759ecba5a 100644 --- a/apps/aether-gateway/src/handlers/admin/request/provider/transport.rs +++ b/apps/aether-gateway/src/handlers/admin/request/provider/transport.rs @@ -126,20 +126,25 @@ impl<'a> AdminAppState<'a> { &self, connector_config: Option<&Map>, ) -> Option { - let explicit_node_id = connector_config - .and_then(|config| admin_provider_transport_string_field(config, "proxy_node_id")); - if let Some(snapshot) = self - .resolve_admin_proxy_node_snapshot(explicit_node_id.as_deref()) - .await - { - return Some(snapshot); + let explicit_node_id_value = + connector_config.and_then(|config| config.get("proxy_node_id")); + if explicit_node_id_value.is_some_and(|value| !value.is_null()) { + let explicit_node_id = connector_config + .and_then(|config| admin_provider_transport_string_field(config, "proxy_node_id")); + return self + .resolve_admin_proxy_node_snapshot(explicit_node_id.as_deref()) + .await + .or_else(|| { + Some(crate::state::unavailable_proxy_snapshot( + "admin_connector_proxy_node_unavailable", + )) + }); } - let proxy = connector_config - .and_then(|config| config.get("proxy")) - .and_then(admin_provider_transport_proxy_snapshot); - if proxy.is_some() { - return proxy; + if let Some(raw_proxy) = connector_config.and_then(|config| config.get("proxy")) { + if let Some(proxy) = admin_provider_transport_proxy_snapshot(raw_proxy) { + return Some(proxy); + } } self.app.resolve_system_proxy_snapshot().await @@ -276,15 +281,24 @@ impl<'a> AdminAppState<'a> { } fn admin_provider_transport_proxy_snapshot(value: &Value) -> Option { - let object = value.as_object()?; + let Some(object) = value.as_object() else { + return Some(crate::state::unavailable_proxy_snapshot( + "admin_connector_proxy_invalid", + )); + }; if object.get("enabled").and_then(Value::as_bool) == Some(false) { return None; } - let proxy_url = object + let Some(proxy_url) = object .get("url") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty())?; + .filter(|value| !value.is_empty()) + else { + return Some(crate::state::unavailable_proxy_snapshot( + "admin_connector_proxy_url_unavailable", + )); + }; let username = object .get("username") .and_then(Value::as_str) @@ -295,6 +309,14 @@ fn admin_provider_transport_proxy_snapshot(value: &Value) -> Option url, + None => { + return Some(crate::state::unavailable_proxy_snapshot( + "admin_connector_proxy_auth_unavailable", + )); + } + }; Some(ProxySnapshot { enabled: Some(true), mode: object @@ -311,8 +333,7 @@ fn admin_provider_transport_proxy_snapshot(value: &Value) -> Option, password: Option<&str>, ) -> Option { - let username = username.filter(|value| !value.is_empty())?; + let username = username.filter(|value| !value.is_empty()); + let password = password.filter(|value| !value.is_empty()); let mut parsed = Url::parse(proxy_url).ok()?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + { + return None; + } + if username.is_none() && password.is_none() { + return Some(parsed.to_string()); + } + let username = username?; parsed.set_username(username).ok()?; parsed.set_password(password).ok()?; Some(parsed.to_string()) @@ -363,21 +394,21 @@ mod tests { use super::admin_provider_transport_proxy_snapshot; #[test] - fn connector_proxy_snapshot_requires_object_value() { - assert_eq!( - admin_provider_transport_proxy_snapshot(&json!("http://proxy.example:8080")), - None - ); + fn connector_proxy_snapshot_keeps_invalid_explicit_value_fail_closed() { + let snapshot = admin_provider_transport_proxy_snapshot(&json!("http://proxy.example:8080")) + .expect("invalid explicit proxy should remain represented"); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); } #[test] - fn connector_proxy_snapshot_requires_url_field() { - assert_eq!( - admin_provider_transport_proxy_snapshot(&json!({ - "proxy_url": "http://proxy.example:8080" - })), - None - ); + fn connector_proxy_snapshot_keeps_missing_url_fail_closed() { + let snapshot = admin_provider_transport_proxy_snapshot(&json!({ + "proxy_url": "http://proxy.example:8080" + })) + .expect("invalid explicit proxy should remain represented"); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); } #[test] @@ -400,4 +431,39 @@ mod tests { }) ); } + + #[test] + fn connector_proxy_auth_injection_failure_is_not_unauthenticated_fallback() { + for value in [ + json!({ + "url": "http://proxy.example:8080", + "password": "secret" + }), + json!({ + "url": "not a proxy url", + "username": "alice", + "password": "secret" + }), + json!({ + "url": "mailto:proxy@example.com", + "username": "alice" + }), + ] { + let snapshot = admin_provider_transport_proxy_snapshot(&value) + .expect("invalid explicit proxy should remain represented"); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); + } + } + + #[test] + fn disabled_connector_proxy_allows_lower_priority_resolution() { + assert_eq!( + admin_provider_transport_proxy_snapshot(&json!({ + "enabled": false, + "url": "http://proxy.example:8080" + })), + None + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/request/route_request.rs b/apps/aether-gateway/src/handlers/admin/request/route_request.rs index 42c39b822..17c3f6194 100644 --- a/apps/aether-gateway/src/handlers/admin/request/route_request.rs +++ b/apps/aether-gateway/src/handlers/admin/request/route_request.rs @@ -2,6 +2,7 @@ use super::{AdminAppState, AdminRequestContext}; use crate::{AppState, GatewayError}; use axum::body::{Body, Bytes}; use axum::http::{HeaderMap, Response}; +use std::net::SocketAddr; pub(crate) enum AdminCancelVideoTaskError { NotFound, @@ -14,6 +15,7 @@ pub(crate) enum AdminCancelVideoTaskError { pub(crate) struct AdminRouteRequest<'a> { state: AdminAppState<'a>, request_context: AdminRequestContext<'a>, + remote_addr: &'a SocketAddr, request_headers: &'a HeaderMap, request_body: Option<&'a Bytes>, } @@ -22,12 +24,14 @@ impl<'a> AdminRouteRequest<'a> { pub(crate) fn new( state: &'a AppState, request_context: &'a crate::control::GatewayPublicRequestContext, + remote_addr: &'a SocketAddr, request_headers: &'a HeaderMap, request_body: Option<&'a Bytes>, ) -> Self { Self { state: AdminAppState::new(state), request_context: AdminRequestContext::new(request_context), + remote_addr, request_headers, request_body, } @@ -41,6 +45,10 @@ impl<'a> AdminRouteRequest<'a> { self.request_context } + pub(crate) fn remote_addr(self) -> &'a SocketAddr { + self.remote_addr + } + pub(crate) fn request_headers(self) -> &'a HeaderMap { self.request_headers } diff --git a/apps/aether-gateway/src/handlers/admin/request/system/export.rs b/apps/aether-gateway/src/handlers/admin/request/system/export.rs index 96cc28227..7e736d611 100644 --- a/apps/aether-gateway/src/handlers/admin/request/system/export.rs +++ b/apps/aether-gateway/src/handlers/admin/request/system/export.rs @@ -3,9 +3,16 @@ use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::system::shared::configs::is_sensitive_admin_system_config_key; use crate::handlers::admin::system::shared::export::{ build_admin_system_export_providers_payload, decrypt_admin_system_export_secret, - ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, + project_admin_system_export_json, project_admin_system_export_optional_url, + project_admin_system_export_url, ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, +}; +use crate::handlers::shared::{ + decrypt_or_migrate_auth_api_key_secret, + decrypt_or_migrate_identity_oauth_provider_client_secret, + decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_smtp_password, + decrypt_or_migrate_system_config_secret, smtp_password_binding, system_config_bool, + system_config_string, unix_secs_to_rfc3339, }; -use crate::handlers::shared::{system_config_string, unix_secs_to_rfc3339}; use crate::GatewayError; use aether_admin::system::{ serialize_admin_system_users_export_wallet, AdminSystemConfigDocument, AdminSystemConfigEntry, @@ -13,24 +20,220 @@ use aether_admin::system::{ AdminSystemConfigProxyNode, ADMIN_SYSTEM_CONFIG_EXPORT_VERSION, ADMIN_SYSTEM_USERS_EXPORT_VERSION, }; -use aether_data_contracts::repository::global_models::AdminGlobalModelListQuery; +use aether_data_contracts::repository::global_models::{ + AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModel, + StoredAdminProviderModel, +}; use chrono::Utc; use serde_json::json; -use std::collections::BTreeMap; +use std::collections::{BTreeMap, BTreeSet}; + +pub(crate) const ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED: &str = "not_exported"; +const ADMIN_SYSTEM_USERS_RECOVERY_EXPORT_VERSION: &str = "1.5"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum SystemExportMode { + InteractiveDownload, + RecoveryBackup, + /// Internal checkpoint used by aggregate imports. It keeps operational enabled/disabled + /// flags while retaining the interactive export's credential redaction guarantees. + RollbackCheckpoint, +} + +impl SystemExportMode { + pub(crate) fn credentials_are_exported(self) -> bool { + self == Self::RecoveryBackup + } + + pub(crate) fn preserves_active_state(self) -> bool { + matches!(self, Self::RecoveryBackup | Self::RollbackCheckpoint) + } + + pub(crate) fn credential_state(self) -> Option { + (!self.credentials_are_exported()) + .then(|| ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED.to_string()) + } + + fn users_export_version(self) -> &'static str { + if self.credentials_are_exported() { + ADMIN_SYSTEM_USERS_RECOVERY_EXPORT_VERSION + } else { + ADMIN_SYSTEM_USERS_EXPORT_VERSION + } + } +} impl<'a> AdminAppState<'a> { + pub(crate) async fn list_all_admin_global_models_for_system_transfer( + &self, + ) -> Result, GatewayError> { + self.list_all_admin_global_models_for_system_transfer_with_page_limit( + ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, + ) + .await + } + + async fn list_all_admin_global_models_for_system_transfer_with_page_limit( + &self, + page_limit: usize, + ) -> Result, GatewayError> { + if page_limit == 0 { + return Err(GatewayError::Internal( + "system transfer global-model page size must be positive".to_string(), + )); + } + + let first = self + .scan_admin_global_models_for_system_transfer(page_limit) + .await?; + let second = self + .scan_admin_global_models_for_system_transfer(page_limit) + .await?; + if first != second { + return Err(GatewayError::Internal( + "global-model catalog changed while building system transfer; retry".to_string(), + )); + } + Ok(second) + } + + async fn scan_admin_global_models_for_system_transfer( + &self, + page_limit: usize, + ) -> Result, GatewayError> { + let mut models = Vec::new(); + let mut seen_ids = BTreeSet::new(); + let mut expected_total = None; + let mut offset = 0_usize; + loop { + let page = self + .list_admin_global_models(&AdminGlobalModelListQuery { + offset, + limit: page_limit, + is_active: None, + search: None, + }) + .await?; + let total = *expected_total.get_or_insert(page.total); + if page.total != total { + return Err(GatewayError::Internal( + "global-model catalog changed while building system transfer; retry" + .to_string(), + )); + } + let page_len = page.items.len(); + if page_len == 0 { + if offset == total { + break; + } + return Err(GatewayError::Internal( + "global-model catalog pagination ended before the advertised total".to_string(), + )); + } + for model in page.items { + if !seen_ids.insert(model.id.clone()) { + return Err(GatewayError::Internal( + "global-model catalog changed while building system transfer; retry" + .to_string(), + )); + } + models.push(model); + } + offset = offset.checked_add(page_len).ok_or_else(|| { + GatewayError::Internal("global-model catalog pagination overflow".to_string()) + })?; + if offset >= total { + if offset != total { + return Err(GatewayError::Internal( + "global-model catalog returned more rows than its advertised total" + .to_string(), + )); + } + break; + } + } + Ok(models) + } + + pub(crate) async fn list_all_admin_provider_models_for_system_transfer( + &self, + provider_id: &str, + ) -> Result, GatewayError> { + self.list_all_admin_provider_models_for_system_transfer_with_page_limit( + provider_id, + ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, + ) + .await + } + + async fn list_all_admin_provider_models_for_system_transfer_with_page_limit( + &self, + provider_id: &str, + page_limit: usize, + ) -> Result, GatewayError> { + if page_limit == 0 { + return Err(GatewayError::Internal( + "system transfer provider-model page size must be positive".to_string(), + )); + } + + let first = self + .scan_admin_provider_models_for_system_transfer(provider_id, page_limit) + .await?; + let second = self + .scan_admin_provider_models_for_system_transfer(provider_id, page_limit) + .await?; + if first != second { + return Err(GatewayError::Internal(format!( + "provider-model catalog for '{provider_id}' changed while building system transfer; retry" + ))); + } + Ok(second) + } + + async fn scan_admin_provider_models_for_system_transfer( + &self, + provider_id: &str, + page_limit: usize, + ) -> Result, GatewayError> { + let mut models = Vec::new(); + let mut seen_ids = BTreeSet::new(); + let mut offset = 0_usize; + loop { + let page = self + .list_admin_provider_models(&AdminProviderModelListQuery { + provider_id: provider_id.to_string(), + is_active: None, + offset, + limit: page_limit, + }) + .await?; + let page_len = page.len(); + for model in page { + if !seen_ids.insert(model.id.clone()) { + return Err(GatewayError::Internal(format!( + "provider-model catalog for '{provider_id}' changed while building system transfer; retry" + ))); + } + models.push(model); + } + if page_len < page_limit { + break; + } + offset = offset.checked_add(page_len).ok_or_else(|| { + GatewayError::Internal("provider-model catalog pagination overflow".to_string()) + })?; + } + Ok(models) + } + pub(crate) async fn build_admin_system_config_export_payload( &self, + mode: SystemExportMode, ) -> Result { let global_models = self - .list_admin_global_models(&AdminGlobalModelListQuery { - offset: 0, - limit: ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, - is_active: None, - search: None, - }) - .await? - .items; + .list_all_admin_global_models_for_system_transfer() + .await?; let global_model_name_by_id = global_models .iter() .map(|model| (model.id.clone(), model.name.clone())) @@ -42,7 +245,10 @@ impl<'a> AdminAppState<'a> { display_name: model.display_name.clone(), usage_count: Some(model.usage_count), default_price_per_request: model.default_price_per_request, - default_tiered_pricing: model.default_tiered_pricing.clone(), + default_tiered_pricing: project_admin_system_export_json( + mode, + model.default_tiered_pricing.as_ref(), + ), supported_capabilities: model.supported_capabilities.as_ref().and_then(|value| { value.as_array().map(|items| { items @@ -52,86 +258,173 @@ impl<'a> AdminAppState<'a> { .collect::>() }) }), - config: model.config.clone(), + config: project_admin_system_export_json(mode, model.config.as_ref()), is_active: model.is_active, }) .collect::>(); let providers_data = - build_admin_system_export_providers_payload(self, &global_model_name_by_id).await?; + build_admin_system_export_providers_payload(self, &global_model_name_by_id, mode) + .await?; - let ldap_data = self - .get_ldap_module_config() - .await? - .map(|config| AdminSystemConfigLdap { - server_url: config.server_url, - bind_dn: config.bind_dn, - bind_password: Some( - config - .bind_password_encrypted - .as_deref() - .and_then(|ciphertext| decrypt_admin_system_export_secret(self, ciphertext)) - .unwrap_or_default(), - ), - base_dn: config.base_dn, - user_search_filter: config.user_search_filter, - username_attr: config.username_attr, - email_attr: config.email_attr, - display_name_attr: config.display_name_attr, - is_enabled: config.is_enabled, - is_exclusive: config.is_exclusive, - use_starttls: config.use_starttls, - connect_timeout: config.connect_timeout, - }); + let ldap_config = self.get_ldap_module_config().await?; + let ldap_bind_password = if mode.credentials_are_exported() { + match ldap_config.as_ref() { + Some(config) => decrypt_or_migrate_ldap_bind_password(self.app(), config).await?, + None => None, + } + } else { + None + }; + let ldap_data = ldap_config.map(|config| AdminSystemConfigLdap { + server_url: config.server_url, + bind_dn: config.bind_dn, + bind_password: ldap_bind_password, + base_dn: config.base_dn, + user_search_filter: config.user_search_filter, + username_attr: config.username_attr, + email_attr: config.email_attr, + display_name_attr: config.display_name_attr, + is_enabled: mode.preserves_active_state() && config.is_enabled, + is_exclusive: mode.preserves_active_state() && config.is_exclusive, + use_starttls: config.use_starttls, + connect_timeout: config.connect_timeout, + }); let system_configs = self.list_system_config_entries().await?; - let system_configs_data = system_configs - .iter() - .map(|entry| { - let value = if is_sensitive_admin_system_config_key(&entry.key) { - entry - .value - .as_str() - .and_then(|ciphertext| decrypt_admin_system_export_secret(self, ciphertext)) - .map(serde_json::Value::String) - .unwrap_or_else(|| entry.value.clone()) - } else { - entry.value.clone() - }; - AdminSystemConfigEntry { - key: entry.key.clone(), - value, - description: entry.description.clone(), - } + let smtp_password_binding = if mode.credentials_are_exported() { + let host = self + .read_system_config_json_value("smtp_host") + .await? + .and_then(|value| system_config_string(Some(&value))); + let port = self + .read_system_config_json_value("smtp_port") + .await? + .map(|value| crate::email_delivery::system_config_u16(Some(&value), 587)) + .unwrap_or(587); + let user = self + .read_system_config_json_value("smtp_user") + .await? + .and_then(|value| system_config_string(Some(&value))); + let use_tls = self + .read_system_config_json_value("smtp_use_tls") + .await? + .map(|value| system_config_bool(Some(&value), true)) + .unwrap_or(true); + let use_ssl = self + .read_system_config_json_value("smtp_use_ssl") + .await? + .map(|value| system_config_bool(Some(&value), false)) + .unwrap_or(false); + host.and_then(|host| { + smtp_password_binding(&host, port, user.as_deref(), use_tls, use_ssl) }) - .collect::>(); + } else { + None + }; + let mut system_configs_data = Vec::new(); + for entry in system_configs.iter().filter(|entry| { + mode.credentials_are_exported() + || !is_sensitive_admin_system_config_key(&entry.key) + && !is_interactive_export_private_system_config_key(&entry.key) + }) { + let value = if mode.credentials_are_exported() + && is_sensitive_admin_system_config_key(&entry.key) + { + match entry.value.as_str() { + Some(stored) if !stored.trim().is_empty() => { + let plaintext = if entry.key.eq_ignore_ascii_case("smtp_password") { + let Some(binding) = smtp_password_binding.as_ref() else { + return Err(GatewayError::Internal( + "RecoveryBackup SMTP password binding is unavailable" + .to_string(), + )); + }; + decrypt_or_migrate_smtp_password( + self.as_ref(), + binding, + stored.to_string(), + ) + .await? + } else { + decrypt_or_migrate_system_config_secret( + self.as_ref(), + &entry.key, + stored.to_string(), + ) + .await? + }; + serde_json::Value::String(plaintext) + } + Some(_) => serde_json::Value::Null, + None if entry.value.is_null() => serde_json::Value::Null, + None => { + return Err(GatewayError::Internal(format!( + "RecoveryBackup 敏感系统配置 '{}' 不是密文字符串或 null", + entry.key, + ))) + } + } + } else { + project_admin_system_export_json(mode, Some(&entry.value)) + .unwrap_or(serde_json::Value::Null) + }; + system_configs_data.push(AdminSystemConfigEntry { + key: entry.key.clone(), + value, + description: entry.description.clone(), + }); + } let oauth_providers = self.list_oauth_provider_configs().await?; - let oauth_data = oauth_providers - .iter() - .map(|provider| AdminSystemConfigOAuthProvider { + let mut oauth_data = Vec::with_capacity(oauth_providers.len()); + for provider in &oauth_providers { + let client_secret = if mode.credentials_are_exported() { + decrypt_or_migrate_identity_oauth_provider_client_secret(self.as_ref(), provider) + .await? + } else { + None + }; + oauth_data.push(AdminSystemConfigOAuthProvider { provider_type: provider.provider_type.clone(), display_name: provider.display_name.clone(), client_id: provider.client_id.clone(), - client_secret: Some( - provider - .client_secret_encrypted - .as_deref() - .and_then(|ciphertext| decrypt_admin_system_export_secret(self, ciphertext)) - .unwrap_or_default(), + client_secret, + authorization_url_override: project_admin_system_export_optional_url( + mode, + provider.authorization_url_override.as_deref(), + ), + token_url_override: project_admin_system_export_optional_url( + mode, + provider.token_url_override.as_deref(), + ), + userinfo_url_override: project_admin_system_export_optional_url( + mode, + provider.userinfo_url_override.as_deref(), ), - authorization_url_override: provider.authorization_url_override.clone(), - token_url_override: provider.token_url_override.clone(), - userinfo_url_override: provider.userinfo_url_override.clone(), scopes: provider.scopes.clone(), - redirect_uri: provider.redirect_uri.clone(), - frontend_callback_url: provider.frontend_callback_url.clone(), - attribute_mapping: provider.attribute_mapping.clone(), - extra_config: provider.extra_config.clone(), - is_enabled: provider.is_enabled, - }) - .collect::>(); + redirect_uri: project_admin_system_export_url(mode, &provider.redirect_uri), + frontend_callback_url: project_admin_system_export_url( + mode, + &provider.frontend_callback_url, + ), + attribute_mapping: project_admin_system_export_json( + mode, + provider.attribute_mapping.as_ref(), + ), + extra_config: project_admin_system_export_json( + mode, + provider.extra_config.as_ref(), + ), + is_enabled: mode.preserves_active_state() && provider.is_enabled, + }); + } - let proxy_nodes = self.list_proxy_nodes().await?; + let mut proxy_nodes = self.list_proxy_nodes().await?; + if mode.credentials_are_exported() { + for node in &mut proxy_nodes { + node.proxy_password = self.app().decrypt_proxy_node_password(&node.id).await?; + } + } let proxy_nodes_data = proxy_nodes .iter() .map(|node| AdminSystemConfigProxyNode { @@ -141,12 +434,21 @@ impl<'a> AdminAppState<'a> { port: Some(node.port), region: node.region.clone(), is_manual: Some(node.is_manual), - proxy_url: node.proxy_url.clone(), - proxy_username: node.proxy_username.clone(), - proxy_password: node.proxy_password.clone(), + proxy_url: project_admin_system_export_optional_url( + mode, + node.proxy_url.as_deref(), + ), + proxy_username: mode + .credentials_are_exported() + .then(|| node.proxy_username.clone()) + .flatten(), + proxy_password: mode + .credentials_are_exported() + .then(|| node.proxy_password.clone()) + .flatten(), tunnel_mode: Some(node.tunnel_mode), heartbeat_interval: Some(node.heartbeat_interval), - remote_config: node.remote_config.clone(), + remote_config: project_admin_system_export_json(mode, node.remote_config.as_ref()), config_version: Some(node.config_version), }) .collect::>(); @@ -154,6 +456,7 @@ impl<'a> AdminAppState<'a> { let document = AdminSystemConfigDocument { version: ADMIN_SYSTEM_CONFIG_EXPORT_VERSION.to_string(), exported_at: Utc::now().to_rfc3339(), + credential_state: mode.credential_state(), global_models: global_models_data, providers: providers_data, proxy_nodes: proxy_nodes_data, @@ -167,6 +470,7 @@ impl<'a> AdminAppState<'a> { pub(crate) async fn build_admin_system_users_export_payload( &self, + mode: SystemExportMode, ) -> Result { let users = self.list_non_admin_export_users().await?; let user_ids = users.iter().map(|user| user.id.clone()).collect::>(); @@ -195,6 +499,29 @@ impl<'a> AdminAppState<'a> { .await?; let usage_aggregates = self.export_admin_system_usage_aggregates().await?; + let mut recovery_api_key_plaintext_by_id = BTreeMap::::new(); + if mode.credentials_are_exported() { + for key in user_api_keys + .iter() + .filter(|key| !key.is_standalone) + .chain(standalone_api_keys.iter()) + { + if key.key_encrypted.is_none() { + continue; + } + let plaintext = decrypt_or_migrate_auth_api_key_secret(self.app(), key).await?; + if recovery_api_key_plaintext_by_id + .insert(key.api_key_id.clone(), plaintext) + .is_some() + { + return Err(GatewayError::Internal(format!( + "RecoveryBackup API Key ID '{}' is not unique", + key.api_key_id, + ))); + } + } + } + let wallets_by_user_id = user_wallets .into_iter() .filter_map(|wallet| wallet.user_id.clone().map(|user_id| (user_id, wallet))) @@ -266,17 +593,24 @@ impl<'a> AdminAppState<'a> { let api_keys_payload = api_keys .iter() .map(|key| { - self.build_admin_system_users_export_api_key_payload(key, None, true) + self.build_admin_system_users_export_api_key_payload( + key, + None, + true, + mode, + recovery_api_key_plaintext_by_id + .get(&key.api_key_id) + .map(String::as_str), + ) }) - .collect::>(); + .collect::, GatewayError>>()?; let usage_totals = user_usage_totals.get(&user.id); - json!({ + let mut payload = json!({ "id": user.id.clone(), "email": user.email.clone(), "email_verified": user.email_verified, "username": user.username.clone(), - "password_hash": user.password_hash.clone(), "role": user.role.clone(), "allowed_providers": user.allowed_providers.clone(), "allowed_providers_mode": user.allowed_providers_mode.clone(), @@ -302,9 +636,13 @@ impl<'a> AdminAppState<'a> { .map(|totals| totals.total_tokens) .unwrap_or(0), "api_keys": api_keys_payload, - }) + }); + if mode.credentials_are_exported() { + payload["password_hash"] = json!(user.password_hash.clone()); + } + Ok::<_, GatewayError>(payload) }) - .collect::>(); + .collect::, GatewayError>>()?; let standalone_keys_data = standalone_api_keys .iter() @@ -313,12 +651,16 @@ impl<'a> AdminAppState<'a> { key, wallets_by_api_key_id.get(&key.api_key_id), false, + mode, + recovery_api_key_plaintext_by_id + .get(&key.api_key_id) + .map(String::as_str), ) }) - .collect::>(); + .collect::, GatewayError>>()?; Ok(json!({ - "version": ADMIN_SYSTEM_USERS_EXPORT_VERSION, + "version": mode.users_export_version(), "exported_at": Utc::now().to_rfc3339(), "user_groups": user_groups_data, "users": users_data, @@ -329,9 +671,10 @@ impl<'a> AdminAppState<'a> { pub(crate) async fn build_admin_system_data_export_payload( &self, + mode: SystemExportMode, ) -> Result { - let config_data = self.build_admin_system_config_export_payload().await?; - let user_data = self.build_admin_system_users_export_payload().await?; + let config_data = self.build_admin_system_config_export_payload(mode).await?; + let user_data = self.build_admin_system_users_export_payload(mode).await?; Ok(json!({ "version": ADMIN_SYSTEM_DATA_EXPORT_VERSION, @@ -346,10 +689,11 @@ impl<'a> AdminAppState<'a> { key: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, wallet: Option<&aether_data::repository::wallet::StoredWalletSnapshot>, include_is_standalone: bool, - ) -> serde_json::Value { + mode: SystemExportMode, + recovery_plaintext: Option<&str>, + ) -> Result { let mut payload = serde_json::Map::from_iter([ ("api_key_id".to_string(), json!(key.api_key_id.clone())), - ("key_hash".to_string(), json!(key.key_hash.clone())), ("name".to_string(), json!(key.name.clone())), ( "allowed_providers".to_string(), @@ -374,7 +718,10 @@ impl<'a> AdminAppState<'a> { "feature_settings".to_string(), json!(key.feature_settings.clone()), ), - ("is_active".to_string(), json!(key.is_active)), + ( + "is_active".to_string(), + json!(mode.preserves_active_state() && key.is_active), + ), ( "expires_at".to_string(), json!(key.expires_at_unix_secs.and_then(unix_secs_to_rfc3339)), @@ -393,21 +740,169 @@ impl<'a> AdminAppState<'a> { ), ]); - if let Some(ciphertext) = key.key_encrypted.as_deref() { - if let Some(plaintext) = decrypt_admin_system_export_secret(self, ciphertext) { - payload.insert("key".to_string(), serde_json::Value::String(plaintext)); - } else { + if mode.credentials_are_exported() { + payload.insert("key_hash".to_string(), json!(key.key_hash.clone())); + if key.key_encrypted.is_some() { + let plaintext = recovery_plaintext.ok_or_else(|| { + GatewayError::Internal(format!( + "RecoveryBackup 无法解密或校验用户 API Key '{}'", + key.api_key_id, + )) + })?; payload.insert( - "key_encrypted".to_string(), - serde_json::Value::String(ciphertext.to_string()), + "key".to_string(), + serde_json::Value::String(plaintext.to_string()), ); } + } else { + payload.insert( + "credential_state".to_string(), + json!(ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED), + ); } if include_is_standalone { payload.insert("is_standalone".to_string(), json!(key.is_standalone)); } - serde_json::Value::Object(payload) + Ok(serde_json::Value::Object(payload)) + } +} + +pub(crate) fn is_interactive_export_private_system_config_key(key: &str) -> bool { + matches!( + key.trim().to_ascii_lowercase().as_str(), + "backup_s3_access_key_id" | "smtp_user" | "turnstile_site_key" | "backup_s3_last_slot" + ) +} + +#[cfg(test)] +mod tests { + use super::{is_interactive_export_private_system_config_key, AdminAppState, SystemExportMode}; + use aether_data::repository::global_models::InMemoryGlobalModelReadRepository; + use aether_data_contracts::repository::global_models::{ + StoredAdminGlobalModel, StoredAdminProviderModel, + }; + use std::sync::Arc; + + #[test] + fn export_modes_keep_interactive_and_recovery_credentials_separate() { + assert!(!SystemExportMode::InteractiveDownload.credentials_are_exported()); + assert_eq!( + SystemExportMode::InteractiveDownload.users_export_version(), + "1.6" + ); + assert_eq!( + SystemExportMode::RecoveryBackup.users_export_version(), + "1.5" + ); + assert!(SystemExportMode::RecoveryBackup.credentials_are_exported()); + assert!(!SystemExportMode::InteractiveDownload.preserves_active_state()); + assert!(SystemExportMode::RollbackCheckpoint.preserves_active_state()); + assert!(!SystemExportMode::RollbackCheckpoint.credentials_are_exported()); + } + + #[test] + fn interactive_system_export_omits_credential_companion_fields() { + for key in [ + "backup_s3_access_key_id", + "smtp_user", + "turnstile_site_key", + "backup_s3_last_slot", + ] { + assert!(is_interactive_export_private_system_config_key(key)); + } + assert!(!is_interactive_export_private_system_config_key( + "site_name" + )); + } + + #[tokio::test] + async fn system_transfer_model_queries_read_every_page() { + let global_models = (0..3) + .map(|index| { + StoredAdminGlobalModel::new( + format!("global-{index}"), + format!("model-{index}"), + format!("Model {index}"), + true, + None, + None, + None, + None, + 0, + 0, + 0, + Some(index), + None, + ) + .expect("test global model should be valid") + }) + .collect::>(); + let provider_models = (0..3) + .map(|index| StoredAdminProviderModel { + id: format!("provider-model-{index}"), + provider_id: "provider-1".to_string(), + global_model_id: format!("global-{index}"), + provider_model_name: format!("upstream-model-{index}"), + provider_model_mappings: None, + price_per_request: None, + tiered_pricing: None, + supports_vision: None, + supports_function_calling: None, + supports_streaming: None, + supports_extended_thinking: None, + supports_image_generation: None, + is_active: true, + is_available: true, + config: None, + created_at_unix_ms: Some(index), + updated_at_unix_secs: None, + global_model_name: Some(format!("model-{index}")), + global_model_display_name: Some(format!("Model {index}")), + global_model_default_price_per_request: None, + global_model_default_tiered_pricing: None, + global_model_supported_capabilities: None, + global_model_config: None, + }) + .collect::>(); + let repository = Arc::new( + InMemoryGlobalModelReadRepository::seed(Vec::new()) + .with_admin_global_models(global_models) + .with_admin_provider_models(provider_models), + ); + let app = crate::AppState::new() + .expect("test app state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_global_model_repository_for_tests(repository), + ); + let state = AdminAppState::new(&app); + + let global_models = state + .list_all_admin_global_models_for_system_transfer_with_page_limit(2) + .await + .expect("all global-model pages should load"); + let provider_models = state + .list_all_admin_provider_models_for_system_transfer_with_page_limit("provider-1", 2) + .await + .expect("all provider-model pages should load"); + + assert_eq!(global_models.len(), 3); + assert_eq!(provider_models.len(), 3); + assert_eq!( + global_models + .iter() + .map(|model| model.id.as_str()) + .collect::>(), + vec!["global-0", "global-1", "global-2"] + ); + assert_eq!( + provider_models + .iter() + .map(|model| model.id.as_str()) + .collect::>(), + vec!["provider-model-2", "provider-model-1", "provider-model-0"] + ); } } diff --git a/apps/aether-gateway/src/handlers/admin/request/system/import.rs b/apps/aether-gateway/src/handlers/admin/request/system/import.rs index 312d41ab1..0a9bf2759 100644 --- a/apps/aether-gateway/src/handlers/admin/request/system/import.rs +++ b/apps/aether-gateway/src/handlers/admin/request/system/import.rs @@ -1,4 +1,7 @@ -use super::{AdminAppState, ADMIN_SYSTEM_DATA_EXPORT_VERSION}; +use super::{ + is_interactive_export_private_system_config_key, AdminAppState, SystemExportMode, + ADMIN_SYSTEM_DATA_EXPORT_VERSION, ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED, +}; use crate::ai_serving::build_provider_key_pool_score_upsert; use crate::api::ai::admin_endpoint_signature_parts; use crate::handlers::admin::admin_provider_pool_config; @@ -9,35 +12,53 @@ use crate::handlers::admin::provider::shared::payloads::{ AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch, AdminProviderUpdatePatch, }; -use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update; +use crate::handlers::admin::provider::write::keys::{ + build_admin_update_provider_key_record_with_existing_keys, + build_provider_catalog_key_admin_cas_update, +}; use crate::handlers::admin::shared::{ normalize_json_array, normalize_json_object, normalize_string_list, }; -use crate::handlers::admin::system::shared::configs::apply_admin_system_config_update; +use crate::handlers::admin::system::shared::configs::{ + apply_admin_system_config_update, is_sensitive_admin_system_config_key, +}; use crate::handlers::admin::users::{ hash_admin_user_api_key, normalize_admin_feature_settings, normalize_admin_list_policy_mode, normalize_admin_rate_limit_policy_mode, normalize_admin_user_api_formats, normalize_admin_user_ip_rules, normalize_admin_user_string_list, }; use crate::handlers::public::normalize_admin_base_url; +use crate::handlers::shared::{ + canonicalize_provider_ops_base_url, ldap_attribute_description_is_valid, + ldap_distinguished_name_is_valid, ldap_search_filter_is_valid, + normalize_ldap_transport_server_url, provider_ops_credential_binding_from_config, + seal_auth_api_key_secret, seal_provider_ops_credential, PROVIDER_OPS_PERSISTENT_SECRET_FIELDS, + PROVIDER_OPS_TRANSIENT_METADATA_FIELDS, PROVIDER_OPS_TRANSIENT_SECRET_FIELDS, +}; use crate::GatewayError; use aether_admin::provider::endpoints as admin_provider_endpoints_pure; use aether_admin::provider::models_write as admin_provider_models_write_pure; +use aether_admin::provider::redaction::{ + admin_restore_secret_safe_body_rules, admin_restore_secret_safe_header_rules, + admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, +}; use aether_admin::system::{ normalize_admin_system_config_key, parse_admin_system_config_array, parse_admin_system_config_import_request, parse_admin_system_config_nested_array, - parse_admin_system_config_optional_object, AdminImportMergeMode, - AdminSystemConfigEndpoint as ImportedEndpoint, AdminSystemConfigEntry as ImportedSystemConfig, + parse_admin_system_config_optional_object, parse_admin_system_config_update, + AdminImportMergeMode, AdminSystemConfigEndpoint as ImportedEndpoint, + AdminSystemConfigEntry as ImportedSystemConfig, AdminSystemConfigGlobalModel as ImportedGlobalModel, AdminSystemConfigImportCounter, AdminSystemConfigImportStats, AdminSystemConfigLdap as ImportedLdapConfig, AdminSystemConfigOAuthProvider as ImportedOAuthProvider, AdminSystemConfigProvider as ImportedProvider, AdminSystemConfigProviderKey as ImportedProviderKey, AdminSystemConfigProviderModel as ImportedProviderModel, - AdminSystemConfigProxyNode as ImportedProxyNode, - ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS, ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS, + AdminSystemConfigProxyNode as ImportedProxyNode, ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS, +}; +use aether_data::repository::auth_modules::{ + CompareAndSwapLdapConfigResult, LdapBindPasswordUpdate, StoredLdapModuleConfig, }; -use aether_data::repository::auth_modules::StoredLdapModuleConfig; use aether_data::repository::oauth_providers::{ EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord, }; @@ -45,18 +66,172 @@ use aether_data::repository::system::{ AdminSystemStatsUserDailyAggregate, AdminSystemUsageAggregateImportMode, AdminSystemUsageAggregateImportSummary, AdminSystemUsageAggregateSnapshot, }; -use aether_data::repository::wallet::WalletLookupKey; +use aether_data::repository::wallet::{StoredWalletSnapshot, WalletLookupKey}; use aether_data_contracts::repository::global_models::{ - AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord, - UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord, + CreateAdminGlobalModelRecord, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord, }; use aether_data_contracts::repository::pool_scores::PoolMemberScoreUpsertMode; use axum::{body::Bytes, http}; use serde_json::{json, Map, Value}; +use sha2::{Digest, Sha256}; use std::collections::{BTreeMap, BTreeSet}; use std::time::{SystemTime, UNIX_EPOCH}; use uuid::Uuid; +const ADMIN_SYSTEM_USERS_RECOVERY_IMPORT_VERSION: (u32, u32) = (1, 5); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SystemImportMode { + InteractiveUpload, + RecoveryBackup, + /// Internal aggregate rollback. Credentials remain redacted, but operational active flags + /// from the checkpoint are restored exactly like a recovery backup. + RollbackCheckpoint, + /// Internal aggregate rollback for a recovery backup. Unlike the interactive checkpoint, + /// this mode carries the encrypted/decrypted credential fields needed to restore values that + /// the failed recovery import may already have overwritten. + RecoveryRollbackCheckpoint, +} + +impl SystemImportMode { + fn restores_credentials(self) -> bool { + matches!( + self, + Self::RecoveryBackup | Self::RecoveryRollbackCheckpoint + ) + } + + fn preserves_active_state(self) -> bool { + matches!( + self, + Self::RecoveryBackup | Self::RollbackCheckpoint | Self::RecoveryRollbackCheckpoint + ) + } + + fn is_rollback_checkpoint(self) -> bool { + matches!( + self, + Self::RollbackCheckpoint | Self::RecoveryRollbackCheckpoint + ) + } + + fn allows_audit_admin_restore(self) -> bool { + matches!( + self, + Self::RecoveryBackup | Self::RecoveryRollbackCheckpoint + ) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ImportedApiKeyMaterial { + key_hash: String, + key_plaintext: Option, +} + +#[derive(Debug, Clone)] +struct ExistingWalletMutation { + before: StoredWalletSnapshot, + // `None` means the import observed an existing wallet but did not receive a + // verifiable post-write snapshot. Rollback must report that as a failure + // rather than guessing or applying an owner-blind overwrite. + after: Option, +} + +#[derive(Debug, Clone)] +struct ExistingUserMutation { + before_auth: aether_data::repository::users::StoredUserAuthRecord, + after_auth: aether_data::repository::users::StoredUserAuthRecord, + before_export: Option, + after_export: Option, + before_model_capability_settings: Option, + after_model_capability_settings: Option, + before_feature_settings: Option, + after_feature_settings: Option, + before_group_ids: Vec, + after_group_ids: Vec, +} + +#[derive(Debug, Clone)] +struct ExistingUserGroupMutation { + before: aether_data::repository::users::StoredUserGroup, + after: aether_data::repository::users::StoredUserGroup, +} + +#[derive(Debug, Clone)] +struct ExistingApiKeyMutation { + before: aether_data::repository::auth::StoredAuthApiKeyExportRecord, + after: aether_data::repository::auth::StoredAuthApiKeyExportRecord, +} + +#[cfg(test)] +fn synthetic_rollback_export_row( + auth: &aether_data::repository::users::StoredUserAuthRecord, + model_capability_settings: Option, + feature_settings: Option, +) -> Result { + aether_data::repository::users::StoredUserExportRow::new( + auth.id.clone(), + auth.email.clone(), + auth.email_verified, + auth.username.clone(), + auth.password_hash.clone(), + auth.role.clone(), + auth.auth_source.clone(), + auth.allowed_providers.clone().map(Value::from), + auth.allowed_api_formats.clone().map(Value::from), + auth.allowed_models.clone().map(Value::from), + None, + model_capability_settings, + auth.is_active, + ) + .map(|row| row.with_feature_settings(feature_settings)) + .and_then(|row| { + row.with_policy_modes( + auth.allowed_providers_mode.clone(), + auth.allowed_api_formats_mode.clone(), + auth.allowed_models_mode.clone(), + "system".to_string(), + ) + }) + .map_err(|err| GatewayError::Internal(err.to_string())) +} + +/// Records rows created by one aggregate import invocation. A post-failure full-table diff is +/// unsafe because ordinary admin mutations may run concurrently with the aggregate operation; +/// only these IDs are eligible for compensation. +#[derive(Debug, Default)] +struct AggregateMutationJournal { + global_model_ids: BTreeSet, + provider_ids: BTreeSet, + provider_endpoint_ids: BTreeSet<(String, String)>, + provider_key_ids: BTreeSet<(String, String)>, + provider_model_ids: BTreeSet<(String, String)>, + oauth_provider_types: BTreeSet, + system_config_keys: BTreeSet, + created_ldap_config: Option, + user_group_ids: BTreeSet, + user_ids: BTreeSet, + user_wallet_snapshots: BTreeMap<(String, String), StoredWalletSnapshot>, + api_key_wallet_snapshots: BTreeMap<(String, String), StoredWalletSnapshot>, + existing_user_wallets: BTreeMap<(String, String), ExistingWalletMutation>, + existing_api_key_wallets: BTreeMap<(String, String), ExistingWalletMutation>, + existing_users: BTreeMap, + existing_user_groups: BTreeMap, + existing_user_api_keys: BTreeMap<(String, String), ExistingApiKeyMutation>, + existing_standalone_api_keys: BTreeMap, + user_api_key_ids: BTreeSet<(String, String)>, + standalone_api_key_ids: BTreeSet, +} + +/// Result of compensating config rows created by one aggregate import. LDAP is tracked +/// separately because its checkpoint restore can overwrite a configuration written by another +/// admin while the import was running. +struct ConfigCleanupOutcome { + result: Result<(), GatewayError>, + skip_ldap_restore: bool, +} + fn invalid_request(detail: impl Into) -> (http::StatusCode, Value) { ( http::StatusCode::BAD_REQUEST, @@ -93,6 +268,186 @@ fn build_admin_system_data_import_part_body( .map_err(|err| invalid_request(format!("{field_name} 序列化失败: {err}"))) } +fn build_aggregate_rollback_body( + checkpoint: &Value, + is_config: bool, +) -> Result { + build_aggregate_rollback_body_with_options(checkpoint, is_config, false) +} + +/// Build the user half of an aggregate rollback without carrying wallet data. +/// Wallet rows are compensated through the journal's owner-checked CAS path; +/// feeding the checkpoint wallet fields back through the regular importer +/// would otherwise perform an unconditional overwrite and could erase a +/// concurrent recharge or adjustment. +fn build_aggregate_users_rollback_body(checkpoint: &Value) -> Result { + let mut object = checkpoint.as_object().cloned().ok_or_else(|| { + GatewayError::Internal("aggregate rollback checkpoint must be a JSON object".to_string()) + })?; + object.insert("merge_mode".to_string(), json!("overwrite")); + + // Usage aggregates and denormalized counters are runtime state. Replaying them during + // compensation could erase requests completed while the failed import was running. + object.remove("usage_aggregates"); + + if let Some(users) = object.get_mut("users") { + let Value::Array(users) = users else { + return Err(GatewayError::Internal( + "aggregate users rollback checkpoint users must be an array".to_string(), + )); + }; + for (index, user) in users.iter_mut().enumerate() { + let Some(user) = user.as_object_mut() else { + return Err(GatewayError::Internal(format!( + "aggregate users rollback checkpoint users[{index}] must be an object" + ))); + }; + user.remove("request_count"); + user.remove("total_tokens"); + user.remove("wallet"); + if let Some(api_keys) = user.get_mut("api_keys") { + let Value::Array(api_keys) = api_keys else { + return Err(GatewayError::Internal(format!( + "aggregate users rollback checkpoint users[{index}].api_keys must be an array" + ))); + }; + for (key_index, api_key) in api_keys.iter_mut().enumerate() { + let Some(api_key) = api_key.as_object_mut() else { + return Err(GatewayError::Internal(format!( + "aggregate users rollback checkpoint users[{index}].api_keys[{key_index}] must be an object" + ))); + }; + api_key.remove("total_requests"); + api_key.remove("total_tokens"); + api_key.remove("total_cost_usd"); + api_key.remove("wallet"); + } + } + } + } + + if let Some(standalone_keys) = object.get_mut("standalone_keys") { + let Value::Array(standalone_keys) = standalone_keys else { + return Err(GatewayError::Internal( + "aggregate users rollback checkpoint standalone_keys must be an array".to_string(), + )); + }; + for (index, key) in standalone_keys.iter_mut().enumerate() { + let Some(key) = key.as_object_mut() else { + return Err(GatewayError::Internal(format!( + "aggregate users rollback checkpoint standalone_keys[{index}] must be an object" + ))); + }; + key.remove("total_requests"); + key.remove("total_tokens"); + key.remove("total_cost_usd"); + key.remove("wallet"); + } + } + + serde_json::to_vec(&Value::Object(object)) + .map(Bytes::from) + .map_err(|err| { + GatewayError::Internal(format!( + "serialize aggregate users rollback checkpoint: {err}" + )) + }) +} + +fn build_aggregate_rollback_body_with_options( + checkpoint: &Value, + is_config: bool, + skip_ldap_config: bool, +) -> Result { + let mut object = checkpoint.as_object().cloned().ok_or_else(|| { + GatewayError::Internal("aggregate rollback checkpoint must be a JSON object".to_string()) + })?; + object.insert("merge_mode".to_string(), json!("overwrite")); + if is_config { + // Proxy nodes are deployment-local and are deliberately not restored by the admin + // config importer. Excluding them also prevents a rollback from changing local routing + // resources while it is restoring the portable catalog. + object.insert("proxy_nodes".to_string(), Value::Array(Vec::new())); + if skip_ldap_config { + // A failed owner-checked LDAP delete means another writer may have changed or + // recreated the row. Missing the field makes the config importer leave that row + // untouched while still restoring every unrelated config section. + object.remove("ldap_config"); + } + } + serde_json::to_vec(&Value::Object(object)) + .map(Bytes::from) + .map_err(|err| { + GatewayError::Internal(format!("serialize aggregate rollback checkpoint: {err}")) + }) +} + +fn aggregate_rollback_failure( + phase: &str, + original_kind: &'static str, + rollback: GatewayError, +) -> GatewayError { + let rollback_kind = gateway_error_kind(&rollback); + tracing::error!( + phase, + original_kind, + rollback_kind, + "aggregate system import failed and compensation failed" + ); + GatewayError::Internal(format!( + "aggregate system import compensation failed in {phase}" + )) +} + +fn gateway_error_kind(error: &GatewayError) -> &'static str { + match error { + GatewayError::UpstreamUnavailable { .. } => "upstream_unavailable", + GatewayError::ControlUnavailable { .. } => "control_unavailable", + GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timeout", + GatewayError::AdmissionTimeout { .. } => "admission_timeout", + GatewayError::Client { .. } => "client", + GatewayError::PlanUsageLimited(_) => "plan_usage_limited", + GatewayError::LastActiveAdminUpdateDenied => "last_admin_update_denied", + GatewayError::LastActiveAdminDeleteDenied => "last_admin_delete_denied", + GatewayError::Internal(_) => "internal", + } +} + +fn aggregate_rollback_error( + phase: &str, + original: GatewayError, + rollback: GatewayError, +) -> GatewayError { + aggregate_rollback_failure(phase, gateway_error_kind(&original), rollback) +} + +fn aggregate_rollback_http_error( + phase: &str, + original: &(http::StatusCode, Value), + rollback: GatewayError, +) -> GatewayError { + let original_kind = if original.0.is_client_error() { + "http_client_error" + } else { + "http_server_error" + }; + aggregate_rollback_failure(phase, original_kind, rollback) +} + +fn combine_rollback_results( + first: Result<(), GatewayError>, + second: Result<(), GatewayError>, + phase: &str, +) -> Result<(), GatewayError> { + match (first, second) { + (Ok(()), Ok(())) => Ok(()), + (Err(error), Ok(())) | (Ok(()), Err(error)) => Err(error), + (Err(_), Err(_)) => Err(GatewayError::Internal(format!( + "multiple aggregate rollback operations failed in {phase}" + ))), + } +} + fn trim_required(value: &str, field_name: &str) -> Result { let trimmed = value.trim(); if trimmed.is_empty() { @@ -102,7 +457,11 @@ fn trim_required(value: &str, field_name: &str) -> Result { } fn normalize_optional_price(value: Option, field_name: &str) -> Result, String> { - admin_provider_models_write_pure::normalize_optional_price(value, field_name) + let value = admin_provider_models_write_pure::normalize_optional_price(value, field_name)?; + if let Some(value) = value { + validate_imported_decimal_storage(value, field_name)?; + } + Ok(value) } fn normalize_supported_capabilities(value: Option>) -> Option { @@ -128,17 +487,284 @@ fn normalize_import_auth_config(value: Option) -> Result, S } } +fn is_imported_redacted_secret(value: &str) -> bool { + matches!(value.trim(), "***" | "********") +} + +fn imported_value_contains_redacted_secret(value: &Value) -> bool { + match value { + Value::String(value) => is_imported_redacted_secret(value), + Value::Array(items) => items.iter().any(imported_value_contains_redacted_secret), + Value::Object(object) => object.values().any(imported_value_contains_redacted_secret), + _ => false, + } +} + +fn imported_config_credentials_not_exported(root: &Map) -> Result { + let Some(value) = root.get("credential_state") else { + return Ok(false); + }; + match value { + Value::String(value) if value.trim() == ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED => { + Ok(true) + } + _ => Err("配置导出 credential_state 无效".to_string()), + } +} + +fn contains_rule_redaction_marker(value: &Value) -> bool { + match value { + Value::Array(items) => items.iter().any(contains_rule_redaction_marker), + Value::Object(object) => { + object.iter().any(|(key, value)| { + matches!( + key.as_str(), + "has_value" | "has_pattern" | "has_replacement" + ) && value.as_bool() == Some(true) + }) || object.values().any(contains_rule_redaction_marker) + } + _ => false, + } +} + +fn strip_imported_redaction_placeholders(value: Value) -> Option { + match value { + Value::String(value) if is_imported_redacted_secret(&value) => None, + Value::Array(items) => Some(Value::Array( + items + .into_iter() + .filter(|item| !contains_rule_redaction_marker(item)) + .filter_map(strip_imported_redaction_placeholders) + .collect(), + )), + Value::Object(object) => Some(Value::Object( + object + .into_iter() + .filter(|(key, _)| { + !matches!( + key.as_str(), + "has_credentials" | "has_value" | "has_pattern" | "has_replacement" + ) + }) + .filter_map(|(key, value)| { + strip_imported_redaction_placeholders(value).map(|value| (key, value)) + }) + .collect(), + )), + value => Some(value), + } +} + +fn prepare_imported_secret_safe_json( + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, +) -> Option { + let incoming = incoming?; + if !credentials_not_exported { + return Some(incoming); + } + if let Some(existing) = existing { + return strip_imported_redaction_placeholders(admin_restore_secret_safe_json( + Some(existing), + &incoming, + )); + } + strip_imported_redaction_placeholders(incoming) +} + +fn prepare_imported_secret_safe_rules( + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, + restore: fn(Option<&Value>, &Value) -> Value, +) -> Option { + let incoming = incoming?; + if !credentials_not_exported { + return Some(incoming); + } + + let restored = existing + .map(|existing| restore(Some(existing), &incoming)) + .unwrap_or_else(|| incoming.clone()); + let (Some(incoming_rules), Some(restored_rules)) = (incoming.as_array(), restored.as_array()) + else { + return strip_imported_redaction_placeholders(restored); + }; + + Some(Value::Array( + incoming_rules + .iter() + .zip(restored_rules) + .filter_map(|(incoming_rule, restored_rule)| { + let contains_placeholder = contains_rule_redaction_marker(incoming_rule) + || imported_value_contains_redacted_secret(incoming_rule); + if contains_placeholder + && (existing.is_none() + || imported_value_contains_redacted_secret(restored_rule)) + { + return None; + } + strip_imported_redaction_placeholders(restored_rule.clone()) + }) + .collect(), + )) +} + +fn prepare_imported_secret_safe_header_rules( + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, +) -> Option { + prepare_imported_secret_safe_rules( + existing, + incoming, + credentials_not_exported, + admin_restore_secret_safe_header_rules, + ) +} + +fn prepare_imported_secret_safe_body_rules( + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, +) -> Option { + prepare_imported_secret_safe_rules( + existing, + incoming, + credentials_not_exported, + admin_restore_secret_safe_body_rules, + ) +} + +fn prepare_imported_secret_safe_proxy( + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, + node_id_map: &BTreeMap, +) -> Option { + let incoming = remap_import_proxy(incoming, node_id_map)?; + let incoming = if credentials_not_exported { + if existing.is_some() { + admin_restore_secret_safe_proxy(existing, &incoming) + } else { + strip_imported_redaction_placeholders(incoming)? + } + } else { + incoming + }; + Some(incoming) +} + +fn prepare_imported_provider_config( + state: &AdminAppState<'_>, + provider_id: &str, + fallback_base_url: Option<&str>, + existing: Option<&Value>, + incoming: Option, + credentials_not_exported: bool, +) -> Result, String> { + if credentials_not_exported { + return normalize_json_object( + prepare_imported_secret_safe_json(existing, incoming, true), + "config", + ); + } + encrypt_imported_provider_config(state, provider_id, fallback_base_url, incoming) +} + +fn imported_provider_ops_fallback_base_url(raw_provider: &Map) -> Option { + raw_provider + .get("endpoints") + .and_then(Value::as_array) + .and_then(|endpoints| endpoints.first()) + .and_then(Value::as_object) + .and_then(|endpoint| endpoint.get("base_url")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| { + raw_provider + .get("website") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) +} + +fn imported_provider_key_credentials_not_exported(item: &ImportedProviderKey) -> bool { + item.credential_state.as_deref().map(str::trim) + == Some(ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED) +} + +fn validate_imported_provider_key_credential_state( + item: &ImportedProviderKey, +) -> Result { + let credentials_not_exported = imported_provider_key_credentials_not_exported(item); + if item.credential_state.is_some() && !credentials_not_exported { + return Err("Provider Key credential_state 无效".to_string()); + } + if credentials_not_exported + && (item + .api_key + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + || item.auth_config.is_some()) + { + return Err("credential_state=not_exported 的 Provider Key 不允许包含凭据字段".to_string()); + } + if item + .api_key + .as_deref() + .is_some_and(is_imported_redacted_secret) + || item + .auth_config + .as_ref() + .is_some_and(imported_value_contains_redacted_secret) + { + return Err("Provider Key 脱敏占位符不能作为凭据导入".to_string()); + } + Ok(credentials_not_exported) +} + fn encrypt_imported_provider_config( state: &AdminAppState<'_>, + provider_id: &str, + fallback_base_url: Option<&str>, config: Option, ) -> Result, String> { let Some(mut config) = normalize_json_object(config, "config")? else { return Ok(None); }; - let Some(credentials) = config + let Some(provider_ops) = config .get_mut("provider_ops") .and_then(Value::as_object_mut) - .and_then(|provider_ops| provider_ops.get_mut("connector")) + else { + return Ok(Some(config)); + }; + let raw_base_url = provider_ops + .get("base_url") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .or(fallback_base_url) + .ok_or_else(|| "RecoveryBackup Provider Ops 缺少 base_url".to_string())?; + let destination = + canonicalize_provider_ops_base_url(raw_base_url).map_err(ToString::to_string)?; + provider_ops.insert( + "base_url".to_string(), + Value::String(destination.base_url().to_string()), + ); + let binding = provider_ops_credential_binding_from_config( + provider_id, + provider_ops, + destination.base_url(), + ) + .map_err(ToString::to_string)?; + let Some(credentials) = provider_ops + .get_mut("connector") .and_then(Value::as_object_mut) .and_then(|connector| connector.get_mut("credentials")) .and_then(Value::as_object_mut) @@ -146,16 +772,29 @@ fn encrypt_imported_provider_config( return Ok(Some(config)); }; - for field in ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS { + for field in PROVIDER_OPS_TRANSIENT_SECRET_FIELDS + .iter() + .chain(PROVIDER_OPS_TRANSIENT_METADATA_FIELDS) + { + credentials.remove(*field); + } + for field in PROVIDER_OPS_PERSISTENT_SECRET_FIELDS { let Some(Value::String(raw)) = credentials.get_mut(*field) else { continue; }; if raw.is_empty() { continue; } - let encrypted = state - .encrypt_catalog_secret_with_fallbacks(raw) - .ok_or_else(|| "gateway 未配置 Provider Ops 加密密钥".to_string())?; + if is_imported_redacted_secret(raw) { + return Err("Provider Ops 脱敏占位符不能作为凭据导入".to_string()); + } + if raw.starts_with("aether-") { + return Err( + "RecoveryBackup Provider Ops 凭据必须是明文,不能包含密文 envelope".to_string(), + ); + } + let encrypted = seal_provider_ops_credential(state.app(), &binding, field, raw) + .map_err(ToString::to_string)?; *raw = encrypted; } @@ -273,6 +912,26 @@ fn imported_service_account_email(config: Option<&Value>) -> Option { } } +fn imported_provider_credential_identity( + imported_key: &ImportedProviderKey, + auth_type: &str, + normalized_auth_config: Option<&Value>, +) -> Option { + if matches!(auth_type, "api_key" | "bearer") { + return imported_key + .api_key + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(|value| format!("secret:{value}")); + } + if matches!(auth_type, "service_account" | "vertex_ai") { + return imported_service_account_email(normalized_auth_config) + .map(|email| format!("service_account:{email}")); + } + None +} + fn build_import_key_match_name(item: &ImportedProviderKey) -> Option { item.name .as_deref() @@ -348,8 +1007,14 @@ fn normalize_import_key_raw_payload( auth_type: &str, normalized_api_formats: &[String], normalized_auth_config: Option, + credentials_not_exported: bool, ) -> Map { let mut payload = raw_key.clone(); + payload.remove("credential_state"); + if credentials_not_exported { + payload.remove("api_key"); + payload.remove("auth_config"); + } if auth_type == "oauth" { payload.remove("api_key"); } @@ -369,10 +1034,12 @@ fn normalize_import_key_raw_payload( allow_auth_channel_mismatch_formats, ); } - if let Some(auth_config) = normalized_auth_config { - payload.insert("auth_config".to_string(), auth_config); - } else if raw_key.contains_key("auth_config") { - payload.insert("auth_config".to_string(), Value::Null); + if !credentials_not_exported { + if let Some(auth_config) = normalized_auth_config { + payload.insert("auth_config".to_string(), auth_config); + } else if raw_key.contains_key("auth_config") { + payload.insert("auth_config".to_string(), Value::Null); + } } payload } @@ -394,14 +1061,22 @@ fn apply_imported_oauth_key_credentials( .as_str() .map(str::trim) .filter(|value| !value.is_empty()); + if plaintext.is_some_and(is_imported_redacted_secret) { + return Err("Provider Key 脱敏占位符不能作为凭据导入".to_string()); + } record.encrypted_api_key = match plaintext { Some(plaintext) => { credentials_supplied = true; api_key_supplied = true; Some( state - .encrypt_catalog_secret_with_fallbacks(plaintext) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?, + .app() + .seal_provider_catalog_key_api_key( + &record.provider_id, + &record.id, + plaintext, + ) + .map_err(GatewayError::into_message)?, ) } None => None, @@ -416,8 +1091,13 @@ fn apply_imported_oauth_key_credentials( serde_json::to_string(auth_config).map_err(|err| err.to_string())?; Some( state - .encrypt_catalog_secret_with_fallbacks(&plaintext) - .ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?, + .app() + .seal_provider_catalog_key_auth_config( + &record.provider_id, + &record.id, + &plaintext, + ) + .map_err(GatewayError::into_message)?, ) } None => None, @@ -577,17 +1257,37 @@ async fn seed_imported_oauth_pool_score( fn build_import_provider_model_record( provider_id: &str, existing_id: Option<&str>, + existing: Option<&aether_data_contracts::repository::global_models::StoredAdminProviderModel>, global_model_id: &str, item: &ImportedProviderModel, + credentials_not_exported: bool, ) -> Result { let provider_model_name = trim_required(&item.provider_model_name, "provider_model_name")?; let provider_model_mappings = normalize_json_array( - item.provider_model_mappings.clone(), + prepare_imported_secret_safe_json( + existing.and_then(|model| model.provider_model_mappings.as_ref()), + item.provider_model_mappings.clone(), + credentials_not_exported, + ), "provider_model_mappings", )?; let price_per_request = normalize_optional_price(item.price_per_request, "price_per_request")?; - let tiered_pricing = normalize_json_object(item.tiered_pricing.clone(), "tiered_pricing")?; - let config = normalize_json_object(item.config.clone(), "config")?; + let tiered_pricing = normalize_json_object( + prepare_imported_secret_safe_json( + existing.and_then(|model| model.tiered_pricing.as_ref()), + item.tiered_pricing.clone(), + credentials_not_exported, + ), + "tiered_pricing", + )?; + let config = normalize_json_object( + prepare_imported_secret_safe_json( + existing.and_then(|model| model.config.as_ref()), + item.config.clone(), + credentials_not_exported, + ), + "config", + )?; UpsertAdminProviderModelRecord::new( existing_id @@ -611,6 +1311,324 @@ fn build_import_provider_model_record( .map_err(|err| err.to_string()) } +fn build_imported_oauth_provider_record( + oauth_provider: &ImportedOAuthProvider, + client_secret_encrypted: EncryptedSecretUpdate, +) -> Result { + let record = UpsertOAuthProviderConfigRecord { + provider_type: trim_required(&oauth_provider.provider_type, "provider_type")?, + display_name: trim_required(&oauth_provider.display_name, "display_name")?, + client_id: trim_required(&oauth_provider.client_id, "client_id")?, + client_secret_encrypted, + authorization_url_override: oauth_provider + .authorization_url_override + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), + token_url_override: oauth_provider + .token_url_override + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), + userinfo_url_override: oauth_provider + .userinfo_url_override + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), + scopes: normalize_string_list(oauth_provider.scopes.clone()), + redirect_uri: trim_required(&oauth_provider.redirect_uri, "redirect_uri")?, + frontend_callback_url: trim_required( + &oauth_provider.frontend_callback_url, + "frontend_callback_url", + )?, + attribute_mapping: normalize_json_object( + oauth_provider.attribute_mapping.clone(), + "attribute_mapping", + )?, + extra_config: normalize_json_object(oauth_provider.extra_config.clone(), "extra_config")?, + icon_url: None, + is_enabled: oauth_provider.is_enabled, + }; + record.validate().map_err(|err| err.to_string())?; + Ok(record) +} + +fn is_custom_identity_oauth_provider_type(provider_type: &str) -> bool { + let provider_type = provider_type.trim().to_ascii_lowercase(); + provider_type == "custom_oidc" + || provider_type.starts_with("custom_oidc_") + || provider_type.starts_with("custom_") + || provider_type.starts_with("oidc_") +} + +fn legacy_custom_oauth_provider_type(provider_type: &str) -> String { + let normalized = provider_type.trim().to_ascii_lowercase(); + let mut suffix = String::with_capacity(normalized.len()); + let mut previous_was_separator = false; + for character in normalized.chars() { + if character.is_ascii_lowercase() || character.is_ascii_digit() || character == '-' { + suffix.push(character); + previous_was_separator = false; + } else if !previous_was_separator { + suffix.push('_'); + previous_was_separator = true; + } + } + let suffix = suffix.trim_matches(['_', '-']); + let candidate = format!("custom_{suffix}"); + if !suffix.is_empty() && candidate.len() <= 64 { + candidate + } else { + let digest = format!("{:x}", Sha256::digest(normalized.as_bytes())); + format!("custom_legacy_{}", &digest[..16]) + } +} + +fn legacy_oauth_endpoint_domains( + oauth_provider: &ImportedOAuthProvider, +) -> Result, String> { + let mut domains = BTreeSet::new(); + for (field, value) in [ + ( + "authorization_url_override", + oauth_provider.authorization_url_override.as_deref(), + ), + ( + "token_url_override", + oauth_provider.token_url_override.as_deref(), + ), + ( + "userinfo_url_override", + oauth_provider.userinfo_url_override.as_deref(), + ), + ] { + let value = value + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| format!("legacy custom OAuth provider is missing {field}"))?; + let parsed = url::Url::parse(value) + .map_err(|_| format!("legacy custom OAuth provider has invalid {field}"))?; + let host = parsed + .host_str() + .map(|host| host.trim_end_matches('.').to_ascii_lowercase()) + .filter(|host| !host.is_empty()) + .ok_or_else(|| format!("legacy custom OAuth provider has invalid {field}"))?; + domains.insert(host); + } + Ok(domains.into_iter().collect()) +} + +fn normalize_legacy_imported_oauth_provider( + mut oauth_provider: ImportedOAuthProvider, + source_version: &str, +) -> Result { + if !matches!(source_version.trim(), "2.0" | "2.1" | "2.2") { + return Ok(oauth_provider); + } + + let original_provider_type = + trim_required(&oauth_provider.provider_type, "provider_type")?.to_ascii_lowercase(); + if original_provider_type == "linuxdo" { + oauth_provider.provider_type = original_provider_type; + return Ok(oauth_provider); + } + + let mut requires_review = false; + if is_custom_identity_oauth_provider_type(&original_provider_type) { + oauth_provider.provider_type = original_provider_type; + } else { + oauth_provider.provider_type = legacy_custom_oauth_provider_type(&original_provider_type); + requires_review = true; + } + + let mut extra_config = match oauth_provider.extra_config.take() { + Some(Value::Object(config)) => config, + Some(_) => return Err("extra_config must be an object".to_string()), + None => Map::new(), + }; + let has_allowed_domains = extra_config + .get("allowed_domains") + .or_else(|| extra_config.get("oauth_allowed_domains")) + .and_then(Value::as_array) + .is_some_and(|domains| !domains.is_empty()); + if !has_allowed_domains { + extra_config.insert( + "allowed_domains".to_string(), + serde_json::to_value(legacy_oauth_endpoint_domains(&oauth_provider)?) + .map_err(|err| err.to_string())?, + ); + requires_review = true; + } + oauth_provider.extra_config = Some(Value::Object(extra_config)); + if requires_review { + oauth_provider.is_enabled = false; + } + Ok(oauth_provider) +} + +fn find_imported_provider_key_index( + state: &AdminAppState<'_>, + imported_key: &ImportedProviderKey, + auth_type: &str, + normalized_auth_config: Option<&Value>, + existing_keys: &[aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey], +) -> Result, String> { + if auth_type == "api_key" { + let target_key = imported_key + .api_key + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + for (index, existing_key) in existing_keys.iter().enumerate() { + let decrypted_existing = state + .app() + .decrypt_provider_catalog_key_api_key(existing_key) + .map_err(GatewayError::into_message)?; + if target_key + .zip(decrypted_existing.as_deref()) + .is_some_and(|(target, decrypted)| decrypted == target) + { + return Ok(Some(index)); + } + } + Ok(None) + } else if matches!(auth_type, "service_account" | "vertex_ai") { + let target_email = imported_service_account_email(normalized_auth_config); + for (index, existing_key) in existing_keys.iter().enumerate() { + let existing_email = imported_existing_provider_auth_config(state, existing_key)? + .and_then(|config| { + config + .get("client_email") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }); + if target_email + .as_deref() + .zip(existing_email.as_deref()) + .is_some_and(|(target, existing)| target == existing) + { + return Ok(Some(index)); + } + } + Ok(None) + } else { + Ok( + build_import_key_match_name(imported_key).and_then(|target_name| { + existing_keys.iter().position(|existing_key| { + existing_key + .auth_type + .trim() + .eq_ignore_ascii_case(auth_type) + && existing_key.name == target_name + }) + }), + ) + } +} + +fn imported_existing_provider_auth_config( + state: &AdminAppState<'_>, + key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey, +) -> Result>, String> { + let Some(plaintext) = state + .app() + .decrypt_provider_catalog_key_auth_config(key) + .map_err(GatewayError::into_message)? + else { + return Ok(None); + }; + let value = serde_json::from_str::(&plaintext).map_err(|_| { + format!( + "Provider Key '{}' 已保存的 auth_config 不是有效 JSON", + key.name + ) + })?; + value.as_object().cloned().map(Some).ok_or_else(|| { + format!( + "Provider Key '{}' 已保存的 auth_config 不是 JSON 对象", + key.name + ) + }) +} + +fn prevalidate_imported_provider_key_uniqueness( + state: &AdminAppState<'_>, + imported_key: &ImportedProviderKey, + auth_type: &str, + normalized_auth_config: Option<&Value>, + existing_keys: &[aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey], +) -> Result<(), String> { + if matches!(auth_type, "api_key" | "bearer") { + let Some(target_key) = imported_key + .api_key + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Ok(()); + }; + let target_key = target_key.to_string(); + for existing in existing_keys.iter().filter(|existing| { + matches!( + existing.auth_type.trim().to_ascii_lowercase().as_str(), + "api_key" | "bearer" + ) + }) { + let Some(decrypted) = state + .app() + .decrypt_provider_catalog_key_api_key(existing) + .map_err(GatewayError::into_message)? + else { + continue; + }; + if decrypted != "__placeholder__" && decrypted == target_key { + return Err(format!( + "该 API Key 已存在于当前 Provider 中(名称: {})", + existing.name + )); + } + } + } + + if auth_type == "service_account" { + let Some(target_email) = imported_service_account_email(normalized_auth_config) else { + return Ok(()); + }; + let target_email = target_email.to_string(); + for existing in existing_keys.iter().filter(|existing| { + matches!( + existing.auth_type.trim().to_ascii_lowercase().as_str(), + "service_account" | "vertex_ai" + ) + }) { + let Some(existing_email) = imported_existing_provider_auth_config(state, existing)? + .and_then(|config| { + config + .get("client_email") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) + else { + continue; + }; + if existing_email == target_email { + return Err(format!( + "该 Service Account ({target_email}) 已存在于当前 Provider 中(名称: {})", + existing.name + )); + } + } + } + Ok(()) +} + #[derive(Debug, Clone, Default, serde::Serialize)] struct AdminSystemUsersImportStats { user_groups: AdminSystemConfigImportCounter, @@ -636,6 +1654,82 @@ struct ImportedWalletTarget { updated_at_unix_secs: Option, } +#[derive(Debug, Clone)] +struct SimulatedImportedUser { + id: String, + email: Option, + username: String, + role: String, + existed_before_import: bool, +} + +#[derive(Debug, Clone)] +struct SimulatedImportedApiKey { + owner_id: String, + is_standalone: bool, + target_id: String, + existed_before_import: bool, +} + +fn replace_simulated_imported_user( + users_by_id: &mut BTreeMap, + email_owners: &mut BTreeMap, + username_owners: &mut BTreeMap, + released_emails: &mut BTreeSet, + released_usernames: &mut BTreeSet, + user: SimulatedImportedUser, +) { + if let Some(previous) = users_by_id.remove(&user.id) { + if previous.email != user.email { + if previous + .email + .as_ref() + .is_some_and(|email| email_owners.get(email) == Some(&previous.id)) + { + let previous_email = previous.email.as_deref().unwrap(); + email_owners.remove(previous_email); + released_emails.insert(previous_email.to_string()); + } + } + if previous.username != user.username + && username_owners.get(&previous.username) == Some(&previous.id) + { + username_owners.remove(&previous.username); + released_usernames.insert(previous.username); + } + } + if let Some(email) = user.email.as_ref() { + released_emails.remove(email); + email_owners.insert(email.clone(), user.id.clone()); + } + released_usernames.remove(&user.username); + username_owners.insert(user.username.clone(), user.id.clone()); + users_by_id.insert(user.id.clone(), user); +} + +fn simulated_imported_user_id_by_identifier( + email_owners: &BTreeMap, + username_owners: &BTreeMap, + identifier: &str, +) -> Option { + email_owners + .get(identifier) + .or_else(|| username_owners.get(identifier)) + .cloned() +} + +fn simulated_imported_user_from_auth_record( + user: &aether_data::repository::users::StoredUserAuthRecord, +) -> SimulatedImportedUser { + SimulatedImportedUser { + id: user.id.clone(), + email: user.email.clone(), + username: user.username.clone(), + role: user.role.clone(), + existed_before_import: true, + } +} + fn imported_system_export_version(version: Option<&Value>) -> Result<(u32, u32), String> { let Some(Value::String(version)) = version else { return Err("version 必须是 x.y 字符串".to_string()); @@ -654,7 +1748,9 @@ fn imported_system_export_version(version: Option<&Value>) -> Result<(u32, u32), Ok((major, minor)) } -fn validate_imported_system_users_export_version(version: Option<&Value>) -> Result<(), String> { +fn validate_imported_system_users_export_version( + version: Option<&Value>, +) -> Result<(u32, u32), String> { let Some(Value::String(raw_version)) = version else { return Err("version 必须是 x.y 字符串".to_string()); }; @@ -662,14 +1758,29 @@ fn validate_imported_system_users_export_version(version: Option<&Value>) -> Res if normalized.is_empty() { return Err("version 必须是 x.y 字符串".to_string()); } - let _ = imported_system_export_version(version)?; + let parsed = imported_system_export_version(version)?; if !ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS.contains(&normalized) { return Err(format!( "不支持的用户数据版本: {normalized},支持的版本: {}", ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS.join(", ") )); } - Ok(()) + Ok(parsed) +} + +fn validate_imported_system_users_export_version_for_mode( + version: Option<&Value>, + mode: SystemImportMode, +) -> Result<(u32, u32), String> { + let parsed = validate_imported_system_users_export_version(version)?; + if mode.restores_credentials() && parsed != ADMIN_SYSTEM_USERS_RECOVERY_IMPORT_VERSION { + return Err(format!( + "恢复备份仅支持用户数据版本 {}.{}", + ADMIN_SYSTEM_USERS_RECOVERY_IMPORT_VERSION.0, + ADMIN_SYSTEM_USERS_RECOVERY_IMPORT_VERSION.1, + )); + } + Ok(parsed) } fn usage_aggregate_import_mode( @@ -706,6 +1817,90 @@ fn imported_optional_string(value: Option<&Value>) -> Result, Str } } +fn normalize_imported_system_user_role( + value: Option<&Value>, + mode: SystemImportMode, +) -> Result, String> { + let raw_role = imported_optional_string(value)?.unwrap_or_else(|| "user".to_string()); + let role = crate::roles::normalize_assignable_user_role(&raw_role) + .ok_or_else(|| format!("不支持的用户角色: {raw_role}"))?; + + if crate::roles::is_full_admin_role(role) + || (crate::roles::is_audit_admin_role(role) && !mode.allows_audit_admin_restore()) + { + return Ok(None); + } + + Ok(Some(role.to_string())) +} + +fn imported_existing_user_is_protected(role: &str, mode: SystemImportMode) -> bool { + crate::roles::is_full_admin_role(role) + || (crate::roles::is_audit_admin_role(role) && !mode.allows_audit_admin_restore()) +} + +fn validate_rollback_user_source_id( + mode: SystemImportMode, + source_user_id: Option<&str>, +) -> Result<(), String> { + if mode.is_rollback_checkpoint() && source_user_id.is_none() { + return Err( + "回滚检查点中的用户必须包含稳定的 users[].id;拒绝按 email/username 猜测用户" + .to_string(), + ); + } + Ok(()) +} + +const IMPORTED_CREDENTIAL_TOMBSTONE_PREFIX: &str = "$aether-import-revoked$"; + +fn imported_credential_tombstone(identity: &str) -> String { + let digest = format!("{:x}", Sha256::digest(identity.as_bytes())); + let digest_length = 64usize.saturating_sub(IMPORTED_CREDENTIAL_TOMBSTONE_PREFIX.len()); + format!( + "{IMPORTED_CREDENTIAL_TOMBSTONE_PREFIX}{}", + &digest[..digest_length] + ) +} + +fn imported_api_key_tombstone(api_key_id: &str) -> String { + imported_credential_tombstone(&format!("api-key-id:{api_key_id}")) +} + +fn imported_api_key_id_for_mode(source_api_key_id: Option<&str>, mode: SystemImportMode) -> String { + if mode.is_rollback_checkpoint() { + if let Some(source_api_key_id) = source_api_key_id { + return source_api_key_id.to_string(); + } + } + Uuid::new_v4().to_string() +} + +fn imported_password_tombstone() -> String { + imported_credential_tombstone(&format!("password:{}", Uuid::new_v4())) +} + +fn resolve_imported_password_hash( + user: &Map, + users_export_version: (u32, u32), + mode: SystemImportMode, +) -> Result, String> { + if users_export_version >= (1, 6) && user.contains_key("password_hash") { + return Err("用户数据 1.6+ 不允许包含 password_hash 凭据字段".to_string()); + } + let password_hash = imported_optional_string(user.get("password_hash"))?; + if !mode.restores_credentials() { + return Ok(password_hash.map(|_| imported_password_tombstone())); + } + if password_hash + .as_deref() + .is_some_and(|value| !aether_data::repository::users::is_valid_bcrypt_hash(value)) + { + return Err("恢复备份中的 password_hash 不是有效的 bcrypt 哈希".to_string()); + } + Ok(password_hash) +} + fn imported_optional_bool(value: Option<&Value>) -> Result, String> { match value { None | Some(Value::Null) => Ok(None), @@ -737,8 +1932,38 @@ fn imported_optional_u64(value: Option<&Value>, field_name: &str) -> Result Result<(), String> { + i64::try_from(value) + .map(|_| ()) + .map_err(|_| format!("{field_name} 超出数据库整数范围")) +} + +fn validate_imported_request_count(value: u64, field_name: &str) -> Result<(), String> { + i32::try_from(value) + .map(|_| ()) + .map_err(|_| format!("{field_name} 超出数据库请求计数范围")) +} + +fn validate_imported_timestamp(value: u64, field_name: &str) -> Result<(), String> { + validate_imported_u64_storage(value, field_name)?; + chrono::DateTime::::from_timestamp(value as i64, 0) + .map(|_| ()) + .ok_or_else(|| format!("{field_name} 超出数据库时间范围")) +} + +fn validate_imported_decimal_storage(value: f64, field_name: &str) -> Result<(), String> { + if !value.is_finite() { + return Err(format!("{field_name} 必须是有限数值")); + } + // PostgreSQL persists imported monetary values as NUMERIC(20,8). + if value.abs() >= 1_000_000_000_000.0 { + return Err(format!("{field_name} 超出数据库金额范围")); + } + Ok(()) +} + fn imported_optional_f64(value: Option<&Value>, field_name: &str) -> Result, String> { - match value { + let parsed = match value { None | Some(Value::Null) => Ok(None), Some(Value::Number(number)) => number .as_f64() @@ -753,7 +1978,11 @@ fn imported_optional_f64(value: Option<&Value>, field_name: &str) -> Result Err(format!("{field_name} 必须是有限数值")), + }?; + if let Some(value) = parsed { + validate_imported_decimal_storage(value, field_name)?; } + Ok(parsed) } fn imported_optional_json_object( @@ -881,6 +2110,200 @@ fn build_imported_user_usage_total_aggregates( Ok(rows) } +fn build_imported_usage_aggregate_snapshot( + value: Option<&Value>, + supplemental_user_daily: &[AdminSystemStatsUserDailyAggregate], +) -> Result { + let mut snapshot = match value { + Some(value) if !value.is_null() => { + serde_json::from_value::(value.clone()) + .map_err(|err| format!("usage_aggregates 格式无效: {err}"))? + } + _ => AdminSystemUsageAggregateSnapshot::default(), + }; + let mut existing_user_totals = BTreeMap::::new(); + for row in &snapshot.stats_user_daily { + let total_tokens = row + .input_tokens + .saturating_add(row.output_tokens) + .saturating_add(row.cache_creation_tokens) + .saturating_add(row.cache_read_tokens); + let entry = existing_user_totals + .entry(row.user_id.clone()) + .or_insert((0, 0)); + entry.0 = entry.0.saturating_add(row.total_requests); + entry.1 = entry.1.saturating_add(total_tokens); + } + for row in supplemental_user_daily { + let existing = existing_user_totals + .get(&row.user_id) + .copied() + .unwrap_or_default(); + let request_delta = row.total_requests.saturating_sub(existing.0); + let token_delta = row.input_tokens.saturating_sub(existing.1); + if request_delta == 0 && token_delta == 0 { + continue; + } + if let Some(existing_row) = snapshot + .stats_user_daily + .iter_mut() + .rev() + .find(|existing_row| existing_row.user_id == row.user_id) + { + existing_row.total_requests = existing_row.total_requests.saturating_add(request_delta); + existing_row.success_requests = + existing_row.success_requests.saturating_add(request_delta); + existing_row.input_tokens = existing_row.input_tokens.saturating_add(token_delta); + } else { + let mut row = row.clone(); + row.total_requests = request_delta; + row.success_requests = request_delta; + row.input_tokens = token_delta; + snapshot.stats_user_daily.push(row); + } + } + Ok(snapshot) +} + +fn validate_imported_usage_aggregate_storage( + snapshot: &AdminSystemUsageAggregateSnapshot, +) -> Result<(), String> { + macro_rules! validate_fields { + ($row:expr, $prefix:expr, [$($field:ident),+ $(,)?]) => { + $(validate_imported_u64_storage( + $row.$field, + &format!("{}.{}", $prefix, stringify!($field)), + )?;)+ + }; + } + + for (index, row) in snapshot.stats_daily.iter().enumerate() { + let prefix = format!("usage_aggregates.stats_daily[{index}]"); + validate_imported_timestamp(row.date_unix_secs, &format!("{prefix}.date_unix_secs"))?; + validate_imported_request_count(row.total_requests, &format!("{prefix}.total_requests"))?; + validate_imported_request_count( + row.success_requests, + &format!("{prefix}.success_requests"), + )?; + validate_imported_request_count(row.error_requests, &format!("{prefix}.error_requests"))?; + validate_fields!( + row, + prefix, + [ + input_tokens, + output_tokens, + cache_creation_tokens, + cache_read_tokens, + ] + ); + if let Some(value) = row.aggregated_at_unix_secs { + validate_imported_timestamp(value, &format!("{prefix}.aggregated_at_unix_secs"))?; + } + validate_imported_decimal_storage(row.total_cost, &format!("{prefix}.total_cost"))?; + validate_imported_decimal_storage( + row.actual_total_cost, + &format!("{prefix}.actual_total_cost"), + )?; + } + for (index, row) in snapshot.stats_user_daily.iter().enumerate() { + let prefix = format!("usage_aggregates.stats_user_daily[{index}]"); + validate_imported_timestamp(row.date_unix_secs, &format!("{prefix}.date_unix_secs"))?; + validate_imported_request_count(row.total_requests, &format!("{prefix}.total_requests"))?; + validate_imported_request_count( + row.success_requests, + &format!("{prefix}.success_requests"), + )?; + validate_imported_request_count(row.error_requests, &format!("{prefix}.error_requests"))?; + validate_fields!( + row, + prefix, + [ + input_tokens, + output_tokens, + cache_creation_tokens, + cache_read_tokens, + ] + ); + validate_imported_decimal_storage(row.total_cost, &format!("{prefix}.total_cost"))?; + } + for (index, row) in snapshot.stats_daily_api_key.iter().enumerate() { + let prefix = format!("usage_aggregates.stats_daily_api_key[{index}]"); + validate_imported_timestamp(row.date_unix_secs, &format!("{prefix}.date_unix_secs"))?; + validate_imported_request_count(row.total_requests, &format!("{prefix}.total_requests"))?; + validate_imported_request_count( + row.success_requests, + &format!("{prefix}.success_requests"), + )?; + validate_imported_request_count(row.error_requests, &format!("{prefix}.error_requests"))?; + validate_fields!( + row, + prefix, + [ + input_tokens, + output_tokens, + cache_creation_tokens, + cache_read_tokens, + ] + ); + validate_imported_decimal_storage(row.total_cost, &format!("{prefix}.total_cost"))?; + } + Ok(()) +} + +fn validate_imported_usage_aggregate_dimensions( + snapshot: &AdminSystemUsageAggregateSnapshot, + user_id_map: &BTreeMap, + api_key_id_map: &BTreeMap, +) -> Result<(), String> { + let mut seen_daily = BTreeSet::new(); + for row in &snapshot.stats_daily { + if !seen_daily.insert(row.date_unix_secs) { + return Err(format!( + "stats_daily aggregate already exists for date_unix_secs={}", + row.date_unix_secs + )); + } + } + let mut seen_user_daily = BTreeSet::new(); + for row in &snapshot.stats_user_daily { + let Some(target_user_id) = user_id_map.get(&row.user_id) else { + continue; + }; + if !seen_user_daily.insert((target_user_id.clone(), row.date_unix_secs)) { + return Err(format!( + "stats_user_daily aggregate already exists for date_unix_secs={}", + row.date_unix_secs + )); + } + } + let mut seen_api_key_daily = BTreeSet::new(); + for row in &snapshot.stats_daily_api_key { + let Some(target_api_key_id) = api_key_id_map.get(&row.api_key_id) else { + continue; + }; + if !seen_api_key_daily.insert((target_api_key_id.clone(), row.date_unix_secs)) { + return Err(format!( + "stats_daily_api_key aggregate already exists for date_unix_secs={}", + row.date_unix_secs + )); + } + } + Ok(()) +} + +fn insert_imported_id_mapping( + mappings: &mut BTreeMap, + source_id: String, + target_id: String, + field_name: &str, +) -> Result<(), String> { + if mappings.contains_key(&source_id) { + return Err(format!("{field_name} 在导入文档中重复: {source_id}")); + } + mappings.insert(source_id, target_id); + Ok(()) +} + fn imported_export_day_unix_secs(exported_at: Option<&Value>) -> u64 { imported_optional_string(exported_at) .ok() @@ -1145,21 +2568,36 @@ fn normalize_imported_wallet_target( .unwrap_or_else(|| "USD".to_string()); let status = imported_optional_string(wallet.and_then(|map| map.get("status")))? .unwrap_or_else(|| "active".to_string()); - let total_recharged = imported_optional_f64( + if currency.chars().count() > 3 { + return Err("wallet.currency 最多允许 3 个字符".to_string()); + } + if status.chars().count() > 20 { + return Err("wallet.status 最多允许 20 个字符".to_string()); + } + let imported_total_recharged = imported_optional_f64( wallet.and_then(|map| map.get("total_recharged")), "wallet.total_recharged", - )? - .unwrap_or(recharge_balance); - let total_consumed = imported_optional_f64( + )?; + let total_recharged = imported_total_recharged.unwrap_or_else(|| recharge_balance.max(0.0)); + if total_recharged < 0.0 { + return Err("wallet.total_recharged 必须是非负有限数值".to_string()); + } + let imported_total_consumed = imported_optional_f64( wallet.and_then(|map| map.get("total_consumed")), "wallet.total_consumed", - )? - .unwrap_or(0.0); - let total_refunded = imported_optional_f64( + )?; + let total_consumed = imported_total_consumed.unwrap_or(0.0); + if total_consumed < 0.0 { + return Err("wallet.total_consumed 必须是非负有限数值".to_string()); + } + let imported_total_refunded = imported_optional_f64( wallet.and_then(|map| map.get("total_refunded")), "wallet.total_refunded", - )? - .unwrap_or(0.0); + )?; + let total_refunded = imported_total_refunded.unwrap_or(0.0); + if total_refunded < 0.0 { + return Err("wallet.total_refunded 必须是非负有限数值".to_string()); + } let total_adjusted = imported_optional_f64( wallet.and_then(|map| map.get("total_adjusted")), "wallet.total_adjusted", @@ -1170,6 +2608,17 @@ fn normalize_imported_wallet_target( "wallet.updated_at", )?; + for (field_name, value) in [ + ("wallet.recharge_balance", recharge_balance), + ("wallet.gift_balance", gift_balance), + ("wallet.total_recharged", total_recharged), + ("wallet.total_consumed", total_consumed), + ("wallet.total_refunded", total_refunded), + ("wallet.total_adjusted", total_adjusted), + ] { + validate_imported_decimal_storage(value, field_name)?; + } + Ok(ImportedWalletTarget { recharge_balance, gift_balance, @@ -1185,10 +2634,1777 @@ fn normalize_imported_wallet_target( } impl<'a> AdminAppState<'a> { + async fn prevalidate_admin_system_config_import( + &self, + request_body: &[u8], + mode: SystemImportMode, + ) -> Result, GatewayError> { + macro_rules! invalid { + ($expr:expr) => { + match $expr { + Ok(value) => value, + Err(detail) => return Ok(Err(invalid_request(detail))), + } + }; + } + macro_rules! routed { + ($expr:expr) => { + match $expr { + Ok(value) => value, + Err(err) => return Ok(Err(err)), + } + }; + } + + let parsed = routed!(parse_admin_system_config_import_request(request_body)); + let source_version = parsed.request.document.version.clone(); + let root = parsed.root; + let credentials_not_exported = invalid!(imported_config_credentials_not_exported(&root)); + let merge_mode = parsed.request.merge_mode; + let imported_global_models = routed!( + parse_admin_system_config_array::(&root, "global_models") + ); + let imported_providers = routed!(parse_admin_system_config_array::( + &root, + "providers", + )); + let imported_proxy_nodes = routed!(parse_admin_system_config_array::( + &root, + "proxy_nodes", + )); + if mode.restores_credentials() && !imported_proxy_nodes.is_empty() { + return Ok(Err(invalid_request( + "恢复备份包含 proxy_nodes,但当前恢复入口不支持安全恢复代理节点", + ))); + } + let imported_ldap = routed!(parse_admin_system_config_optional_object::< + ImportedLdapConfig, + >(&root, "ldap_config")); + let imported_oauth_providers = routed!(parse_admin_system_config_array::< + ImportedOAuthProvider, + >(&root, "oauth_providers")); + let imported_system_configs = routed!(parse_admin_system_config_array::< + ImportedSystemConfig, + >(&root, "system_configs")); + let (imported_external_models_configs, mut imported_system_configs): (Vec<_>, Vec<_>) = + imported_system_configs.into_iter().partition(|item| { + normalize_imported_system_config_key(&item.value.key) + == ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY + }); + // Destination-bound secrets must be applied after their destination fields. Recovery + // documents are user-controlled JSON and do not guarantee any ordering. + imported_system_configs.sort_by_key(|item| { + matches!( + normalize_imported_system_config_key(&item.value.key).as_str(), + "smtp_password" | "module.bark_push.device_key" + ) + }); + let mut existing_system_config_keys = self + .list_system_config_entries() + .await? + .into_iter() + .map(|entry| normalize_imported_system_config_key(&entry.key)) + .collect::>(); + for imported_config_item in imported_external_models_configs { + let config = imported_config_item.value; + let exists = + existing_system_config_keys.contains(ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY); + match (exists, merge_mode) { + (true, AdminImportMergeMode::Skip) => continue, + (true, AdminImportMergeMode::Error) => { + return Ok(Err(invalid_request(format!( + "SystemConfig '{ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY}' 已存在" + )))); + } + _ => {} + } + match config.value { + Value::Null => {} + Value::String(value) if !value.trim().is_empty() => {} + Value::String(_) => { + return Ok(Err(invalid_request( + "external_models_proxy_node_id 不能为空", + ))); + } + _ => { + return Ok(Err(invalid_request( + "external_models_proxy_node_id 必须是字符串或 null", + ))); + } + } + existing_system_config_keys + .insert(ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY.to_string()); + } + + let mut global_models_by_name = self + .list_all_admin_global_models_for_system_transfer() + .await? + .into_iter() + .map(|model| (model.name.clone(), (model.id.clone(), Some(model)))) + .collect::>(); + for imported_model_item in &imported_global_models { + let model = &imported_model_item.value; + let name = invalid!(trim_required(&model.name, "name")); + let display_name = invalid!(trim_required(&model.display_name, "display_name")); + let default_price_per_request = invalid!(normalize_optional_price( + model.default_price_per_request, + "default_price_per_request", + )); + let existing_model = global_models_by_name + .get(&name) + .and_then(|(_, model)| model.as_ref()); + let default_tiered_pricing = invalid!(normalize_json_object( + prepare_imported_secret_safe_json( + existing_model.and_then(|model| model.default_tiered_pricing.as_ref()), + model.default_tiered_pricing.clone(), + credentials_not_exported, + ), + "default_tiered_pricing", + )); + let supported_capabilities = + normalize_supported_capabilities(model.supported_capabilities.clone()); + let config = invalid!(normalize_json_object( + prepare_imported_secret_safe_json( + existing_model.and_then(|model| model.config.as_ref()), + model.config.clone(), + credentials_not_exported, + ), + "config", + )); + if let Some((existing_id, _)) = global_models_by_name.get(&name).cloned() { + match merge_mode { + AdminImportMergeMode::Skip => continue, + AdminImportMergeMode::Error => { + return Ok(Err(invalid_request(format!("GlobalModel '{name}' 已存在")))); + } + AdminImportMergeMode::Overwrite => { + invalid!(UpdateAdminGlobalModelRecord::new( + existing_id, + display_name, + model.is_active, + default_price_per_request, + default_tiered_pricing, + supported_capabilities, + config, + ) + .map_err(|err| err.to_string())); + } + } + } else { + let id = Uuid::new_v4().to_string(); + invalid!(CreateAdminGlobalModelRecord::new( + id.clone(), + name.clone(), + display_name, + model.is_active, + default_price_per_request, + default_tiered_pricing, + supported_capabilities, + config, + ) + .map_err(|err| err.to_string())); + global_models_by_name.insert(name, (id, None)); + } + } + + let mut providers_by_name = self + .list_provider_catalog_providers(false) + .await? + .into_iter() + .map(|provider| (provider.name.clone(), provider)) + .collect::>(); + let mut endpoints_by_provider = BTreeMap::< + String, + BTreeMap< + String, + aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint, + >, + >::new(); + let mut keys_by_provider = BTreeMap::< + String, + Vec, + >::new(); + let mut models_by_provider = BTreeMap::< + String, + BTreeMap< + String, + ( + String, + Option< + aether_data_contracts::repository::global_models::StoredAdminProviderModel, + >, + ), + >, + >::new(); + let node_id_map = BTreeMap::::new(); + + for imported_provider_item in &imported_providers { + let raw_provider = &imported_provider_item.raw; + let imported_provider = &imported_provider_item.value; + let provider_name = invalid!(trim_required(&imported_provider.name, "name")); + invalid!( + crate::provider_transport::validate_anthropic_compatibility_profile_config( + imported_provider.config.as_ref(), + ) + .map_err(|_| "无效的 Anthropic compatibility profile".to_string()) + ); + + let existing_provider = providers_by_name.get(&provider_name).cloned(); + if existing_provider.is_some() && merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "Provider '{provider_name}' 已存在" + )))); + } + let provider = if let Some(existing) = existing_provider { + if merge_mode == AdminImportMergeMode::Skip { + existing + } else { + let patch = AdminProviderUpdatePatch::from_object(raw_provider.clone()) + .map_err(|_| "Provider 配置格式无效".to_string()); + let patch = invalid!(patch); + let mut updated = invalid!( + self.build_admin_update_provider_record(&existing, patch) + .await + ); + updated.proxy = prepare_imported_secret_safe_proxy( + existing.proxy.as_ref(), + imported_provider.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); + let provider_ops_fallback_base_url = + imported_provider_ops_fallback_base_url(raw_provider); + updated.config = invalid!(prepare_imported_provider_config( + self, + &updated.id, + provider_ops_fallback_base_url.as_deref(), + existing.config.as_ref(), + imported_provider.config.clone(), + credentials_not_exported, + )); + updated + } + } else { + let payload = serde_json::from_value::(Value::Object( + raw_provider.clone(), + )) + .map_err(|_| format!("Provider '{provider_name}' 配置格式无效")); + let payload = invalid!(payload); + let (mut record, _) = + invalid!(self.build_admin_create_provider_record(payload).await); + record.name = provider_name.clone(); + if let Some(enable_format_conversion) = imported_provider.enable_format_conversion { + record.enable_format_conversion = enable_format_conversion; + } + record.proxy = prepare_imported_secret_safe_proxy( + None, + imported_provider.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); + let provider_ops_fallback_base_url = + imported_provider_ops_fallback_base_url(raw_provider); + record.config = invalid!(prepare_imported_provider_config( + self, + &record.id, + provider_ops_fallback_base_url.as_deref(), + None, + imported_provider.config.clone(), + credentials_not_exported, + )); + record + }; + + let mut existing_endpoints_by_format = + match endpoints_by_provider.remove(&provider_name) { + Some(endpoints) => endpoints, + None => self + .list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref( + &provider.id, + )) + .await? + .into_iter() + .map(|endpoint| (endpoint.api_format.clone(), endpoint)) + .collect(), + }; + let imported_endpoints = routed!(parse_admin_system_config_nested_array::< + ImportedEndpoint, + >(raw_provider, "endpoints")); + for imported_endpoint_item in imported_endpoints { + let (raw_endpoint, imported_endpoint) = imported_endpoint_item.into_parts(); + let normalized_api_format = invalid!(normalize_import_endpoint_format( + &imported_endpoint.api_format, + )); + invalid!( + crate::provider_transport::validate_anthropic_compatibility_profile_config( + imported_endpoint.config.as_ref(), + ) + .map_err(|_| "无效的 Anthropic compatibility profile".to_string()) + ); + if !fixed_provider_import_endpoint_supported( + &provider.provider_type, + &normalized_api_format, + ) { + existing_endpoints_by_format.remove(&normalized_api_format); + continue; + } + let existing_endpoint = existing_endpoints_by_format + .get(&normalized_api_format) + .cloned(); + if existing_endpoint.is_some() { + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "Endpoint '{normalized_api_format}' 已存在于 Provider '{provider_name}'" + )))); + } + if merge_mode == AdminImportMergeMode::Skip { + continue; + } + } + let Some((signature, api_family, endpoint_kind)) = + admin_endpoint_signature_parts(&normalized_api_format) + else { + return Ok(Err(invalid_request(format!( + "无效的 api_format: {}", + imported_endpoint.api_format + )))); + }; + let endpoint = if let Some(existing) = existing_endpoint.as_ref() { + let patch = AdminProviderEndpointUpdatePatch::from_object(raw_endpoint) + .map_err(|_| "Provider Endpoint 配置格式无效".to_string()); + let patch = invalid!(patch); + let (fields, payload) = patch.into_parts(); + let normalized_base_url = payload + .base_url + .as_deref() + .map(normalize_admin_base_url) + .transpose(); + let normalized_base_url = invalid!(normalized_base_url); + let update_fields = + admin_provider_endpoints_pure::AdminProviderEndpointUpdateFields { + base_url: normalized_base_url, + custom_path: payload.custom_path, + header_rules: prepare_imported_secret_safe_header_rules( + existing.header_rules.as_ref(), + payload.header_rules, + credentials_not_exported, + ), + body_rules: prepare_imported_secret_safe_body_rules( + existing.body_rules.as_ref(), + payload.body_rules, + credentials_not_exported, + ), + max_retries: payload.max_retries, + is_active: payload.is_active, + config: prepare_imported_secret_safe_json( + existing.config.as_ref(), + payload.config, + credentials_not_exported, + ), + proxy: payload.proxy, + format_acceptance_config: prepare_imported_secret_safe_json( + existing.format_acceptance_config.as_ref(), + payload.format_acceptance_config, + credentials_not_exported, + ), + }; + let mut updated = invalid!( + admin_provider_endpoints_pure::apply_admin_provider_endpoint_update_fields( + existing, + |field| fields.contains(field), + |field| fields.is_null(field), + &update_fields, + ) + ); + updated.api_format = signature.to_string(); + updated.api_family = Some(api_family.to_string()); + updated.endpoint_kind = Some(endpoint_kind.to_string()); + updated + } else { + invalid!( + admin_provider_endpoints_pure::build_admin_provider_endpoint_record( + Uuid::new_v4().to_string(), + provider.id.clone(), + signature.to_string(), + api_family.to_string(), + endpoint_kind.to_string(), + invalid!(normalize_admin_base_url(&imported_endpoint.base_url)), + imported_endpoint.custom_path, + prepare_imported_secret_safe_header_rules( + None, + imported_endpoint.header_rules, + credentials_not_exported, + ), + prepare_imported_secret_safe_body_rules( + None, + imported_endpoint.body_rules, + credentials_not_exported, + ), + imported_endpoint.max_retries.unwrap_or(2), + prepare_imported_secret_safe_json( + None, + imported_endpoint.config, + credentials_not_exported, + ), + prepare_imported_secret_safe_proxy( + None, + imported_endpoint.proxy, + credentials_not_exported, + &node_id_map, + ), + prepare_imported_secret_safe_json( + None, + imported_endpoint.format_acceptance_config, + credentials_not_exported, + ), + 0, + ) + ) + }; + existing_endpoints_by_format.insert(normalized_api_format, endpoint); + } + + let endpoint_formats = existing_endpoints_by_format + .keys() + .cloned() + .collect::>(); + let mut existing_keys = match keys_by_provider.remove(&provider_name) { + Some(keys) => keys, + None => { + self.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref( + &provider.id, + )) + .await? + } + }; + let imported_keys = routed!(parse_admin_system_config_nested_array::< + ImportedProviderKey, + >(raw_provider, "api_keys")); + let mut imported_credential_identities = BTreeSet::new(); + for imported_key_item in imported_keys { + let (raw_key, imported_key) = imported_key_item.into_parts(); + let (normalized_api_formats, _) = + normalize_import_key_formats(&imported_key, &endpoint_formats); + if normalized_api_formats.is_empty() { + continue; + } + let normalized_auth_config = invalid!(normalize_import_auth_config( + imported_key.auth_config.clone() + )); + let auth_type = imported_key_auth_type(&imported_key); + let credentials_not_exported = invalid!( + validate_imported_provider_key_credential_state(&imported_key) + ); + if imported_provider_credential_identity( + &imported_key, + &auth_type, + normalized_auth_config.as_ref(), + ) + .is_some_and(|identity| !imported_credential_identities.insert(identity)) + { + return Ok(Err(invalid_request(format!( + "Provider '{provider_name}' 的凭据在导入文档中重复" + )))); + } + let normalized_raw_key = normalize_import_key_raw_payload( + &raw_key, + &auth_type, + &normalized_api_formats, + normalized_auth_config.clone(), + credentials_not_exported, + ); + let existing_key_index = if credentials_not_exported { + build_import_key_match_name(&imported_key).and_then(|target_name| { + existing_keys.iter().position(|existing_key| { + existing_key + .auth_type + .trim() + .eq_ignore_ascii_case(&auth_type) + && existing_key.name == target_name + }) + }) + } else { + invalid!(find_imported_provider_key_index( + self, + &imported_key, + &auth_type, + normalized_auth_config.as_ref(), + &existing_keys, + )) + }; + if credentials_not_exported && existing_key_index.is_none() { + continue; + } + if existing_key_index.is_some() && merge_mode == AdminImportMergeMode::Skip { + continue; + } + let mut record = if let Some(existing_index) = existing_key_index { + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "Provider '{provider_name}' 中存在重复 Key" + )))); + } + let patch = AdminProviderKeyUpdatePatch::from_object(normalized_raw_key) + .map_err(|_| "Provider Key 配置格式无效".to_string()); + let patch = invalid!(patch); + let mut record = + invalid!(build_admin_update_provider_key_record_with_existing_keys( + self, + &provider, + &existing_keys[existing_index], + &existing_keys, + patch, + )); + if credentials_not_exported { + record.encrypted_api_key = + existing_keys[existing_index].encrypted_api_key.clone(); + record.encrypted_auth_config = + existing_keys[existing_index].encrypted_auth_config.clone(); + } + record + } else { + invalid!(prevalidate_imported_provider_key_uniqueness( + self, + &imported_key, + &auth_type, + normalized_auth_config.as_ref(), + &existing_keys, + )); + let payload = serde_json::from_value::( + Value::Object(normalized_raw_key), + ) + .map_err(|_| "Provider Key 配置格式无效".to_string()); + invalid!( + self.build_admin_create_provider_key_record(&provider, invalid!(payload)) + .await + ) + }; + if auth_type == "oauth" { + invalid!(apply_imported_oauth_key_credentials( + self, + &provider.provider_type, + None, + &raw_key, + normalized_auth_config.as_ref(), + &mut record, + )); + } + if existing_key_index.is_none() || merge_mode != AdminImportMergeMode::Skip { + invalid!(normalize_json_object( + imported_key.global_priority_by_format.clone(), + "global_priority_by_format", + )); + invalid!(normalize_json_object( + imported_key.fingerprint.clone(), + "fingerprint", + )); + } + if let Some(index) = existing_key_index { + existing_keys[index] = record; + } else { + existing_keys.push(record); + } + } + + let imported_models = routed!(parse_admin_system_config_nested_array::< + ImportedProviderModel, + >(raw_provider, "models")); + let mut existing_models_by_name = match models_by_provider.remove(&provider_name) { + Some(models) => models, + None => self + .list_all_admin_provider_models_for_system_transfer(&provider.id) + .await? + .into_iter() + .map(|model| { + ( + model.provider_model_name.clone(), + (model.id.clone(), Some(model)), + ) + }) + .collect::>(), + }; + for imported_model_item in imported_models { + let imported_model = imported_model_item.value; + let Some(global_model_name) = imported_model + .global_model_name + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + continue; + }; + let Some((global_model_id, _)) = global_models_by_name.get(global_model_name) + else { + continue; + }; + let provider_model_name = invalid!(trim_required( + &imported_model.provider_model_name, + "provider_model_name", + )); + if existing_models_by_name.contains_key(&provider_model_name) { + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "Model '{provider_model_name}' 已存在于 Provider '{provider_name}'" + )))); + } + if merge_mode == AdminImportMergeMode::Skip { + continue; + } + } + let existing_model = existing_models_by_name + .get(&provider_model_name) + .and_then(|(_, model)| model.as_ref()); + invalid!(build_import_provider_model_record( + &provider.id, + existing_models_by_name + .get(&provider_model_name) + .map(|(id, _)| id.as_str()), + existing_model, + global_model_id, + &imported_model, + credentials_not_exported, + )); + existing_models_by_name + .entry(provider_model_name) + .or_insert_with(|| (Uuid::new_v4().to_string(), None)); + } + + providers_by_name.insert(provider_name.clone(), provider); + endpoints_by_provider.insert(provider_name.clone(), existing_endpoints_by_format); + keys_by_provider.insert(provider_name.clone(), existing_keys); + models_by_provider.insert(provider_name, existing_models_by_name); + } + + if let Some(imported_ldap) = imported_ldap.filter(|_| self.has_auth_module_writer()) { + let ldap_config = imported_ldap.value; + let existing = self.get_ldap_module_config().await?; + let server_url = invalid!(trim_required(&ldap_config.server_url, "LDAP 服务器地址")); + let server_url = invalid!(normalize_ldap_transport_server_url( + &server_url, + ldap_config.use_starttls, + ) + .ok_or_else(|| { + "LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段" + .to_string() + })); + let bind_dn = invalid!(trim_required(&ldap_config.bind_dn, "绑定 DN")); + let base_dn = invalid!(trim_required(&ldap_config.base_dn, "Base DN")); + if !ldap_distinguished_name_is_valid(&bind_dn) + || !ldap_distinguished_name_is_valid(&base_dn) + { + return Ok(Err(invalid_request( + "LDAP 绑定 DN 或 Base DN 格式无效或过长", + ))); + } + let user_search_filter = invalid!(trim_required( + ldap_config + .user_search_filter + .as_deref() + .unwrap_or("(uid={username})"), + "搜索过滤器", + )); + if !ldap_search_filter_is_valid(&user_search_filter) { + return Ok(Err(invalid_request( + "LDAP 搜索过滤器格式无效,必须包含 {username} 且使用有限的括号结构", + ))); + } + let username_attr = invalid!(trim_required( + ldap_config.username_attr.as_deref().unwrap_or("uid"), + "用户名属性", + )); + let email_attr = invalid!(trim_required( + ldap_config.email_attr.as_deref().unwrap_or("mail"), + "邮箱属性", + )); + let display_name_attr = invalid!(trim_required( + ldap_config.display_name_attr.as_deref().unwrap_or("cn"), + "显示名称属性", + )); + if [ + username_attr.as_str(), + email_attr.as_str(), + display_name_attr.as_str(), + ] + .into_iter() + .any(|attribute| !ldap_attribute_description_is_valid(attribute)) + { + return Ok(Err(invalid_request( + "LDAP 用户名、邮箱或显示名称属性格式无效", + ))); + } + let connect_timeout = ldap_config.connect_timeout.unwrap_or(10); + if !(1..=60).contains(&connect_timeout) { + return Ok(Err(invalid_request( + "LDAP connect_timeout 必须在 1 到 60 秒之间", + ))); + } + let config = StoredLdapModuleConfig { + server_url, + bind_dn, + // Password mutation is explicit and separate from the replacement snapshot. + // In particular, Preserve never copies a previously read ciphertext here. + bind_password_encrypted: None, + base_dn, + user_search_filter: Some(user_search_filter), + username_attr: Some(username_attr), + email_attr: Some(email_attr), + display_name_attr: Some(display_name_attr), + is_enabled: ldap_config.is_enabled, + is_exclusive: ldap_config.is_exclusive, + use_starttls: ldap_config.use_starttls, + connect_timeout: Some(connect_timeout), + }; + let bind_password = ldap_config + .bind_password + .as_deref() + .map(str::trim) + .map(ToOwned::to_owned); + if bind_password + .as_deref() + .is_some_and(is_imported_redacted_secret) + { + return Ok(Err(invalid_request("LDAP 脱敏占位符不能作为绑定密码导入"))); + } + let bind_password_update = match bind_password { + Some(password) if password.is_empty() => LdapBindPasswordUpdate::Clear, + Some(password) => LdapBindPasswordUpdate::Set(invalid!(self + .encrypt_ldap_bind_password(&config, &password) + .ok_or_else(|| { + "LDAP 绑定密码加密失败,请检查 Rust 数据加密配置".to_string() + }))), + None => LdapBindPasswordUpdate::Preserve, + }; + if matches!(&bind_password_update, LdapBindPasswordUpdate::Preserve) { + if let Some(existing) = existing.as_ref() { + if existing + .bind_password_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + let binding_matches = + invalid!(crate::handlers::shared::ldap_bind_password_binding_matches( + existing, &config, + )); + if !binding_matches { + return Ok(Err(invalid_request( + "导入 LDAP 时修改了服务器、StartTLS、bind DN 或 Base DN,必须提供绑定密码", + ))); + } + } + } + } + let will_have_password = match &bind_password_update { + LdapBindPasswordUpdate::Set(ciphertext) => !ciphertext.trim().is_empty(), + LdapBindPasswordUpdate::Clear => false, + LdapBindPasswordUpdate::Preserve => existing + .as_ref() + .and_then(|config| config.bind_password_encrypted.as_deref()) + .map(str::trim) + .is_some_and(|value| !value.is_empty()), + }; + if existing.is_none() + && !matches!(&bind_password_update, LdapBindPasswordUpdate::Set(_)) + { + return Ok(Err(invalid_request("首次配置 LDAP 时必须设置绑定密码"))); + } + if ldap_config.is_exclusive && !ldap_config.is_enabled { + return Ok(Err(invalid_request( + "仅允许 LDAP 登录 需要先启用 LDAP 认证", + ))); + } + if ldap_config.is_enabled && !will_have_password { + return Ok(Err(invalid_request("启用 LDAP 认证 需要先设置绑定密码"))); + } + if ldap_config.is_enabled && ldap_config.is_exclusive { + let admin_count = self + .count_active_local_admin_users_with_valid_password() + .await?; + if admin_count < 1 { + return Ok(Err(invalid_request( + "启用 LDAP 独占模式前,必须至少保留 1 个有效的本地管理员账户(含有效密码)作为紧急恢复通道", + ))); + } + } + if existing.is_some() && merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request("LDAP 配置已存在"))); + } + } + + let existing_oauth_providers = self.list_oauth_provider_configs().await?; + let existing_oauth_by_type = existing_oauth_providers + .iter() + .map(|provider| (provider.provider_type.clone(), provider)) + .collect::>(); + let mut oauth_provider_types = existing_oauth_by_type + .keys() + .cloned() + .collect::>(); + for imported_oauth_item in imported_oauth_providers { + let oauth_provider = invalid!(normalize_legacy_imported_oauth_provider( + imported_oauth_item.value, + &source_version, + )); + let provider_type = invalid!(trim_required( + &oauth_provider.provider_type, + "provider_type", + )); + if oauth_provider_types.contains(&provider_type) { + match merge_mode { + AdminImportMergeMode::Skip => continue, + AdminImportMergeMode::Error => { + return Ok(Err(invalid_request(format!( + "OAuth Provider '{provider_type}' 已存在" + )))); + } + AdminImportMergeMode::Overwrite => {} + } + } + // Construct and validate the complete record before sealing the secret. The + // envelope binding includes client_id, redirect URI, and endpoint overrides; + // sealing against provider_type alone would allow a secret to be replayed after + // those fields change. + let mut record = invalid!(build_imported_oauth_provider_record( + &oauth_provider, + EncryptedSecretUpdate::Preserve, + )); + let client_secret_update = match oauth_provider.client_secret.as_deref().map(str::trim) + { + Some(secret) if is_imported_redacted_secret(secret) => { + EncryptedSecretUpdate::Preserve + } + Some(secret) if !secret.is_empty() => EncryptedSecretUpdate::Set(invalid!( + crate::handlers::shared::seal_identity_oauth_provider_client_secret( + self.as_ref(), + &record, + secret, + ) + .map_err(str::to_string) + )), + _ => EncryptedSecretUpdate::Preserve, + }; + record.client_secret_encrypted = client_secret_update; + if matches!( + &record.client_secret_encrypted, + EncryptedSecretUpdate::Preserve + ) { + if let Some(existing) = existing_oauth_by_type.get(&provider_type) { + if existing + .client_secret_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + let binding_matches = invalid!( + crate::handlers::shared::identity_oauth_provider_secret_binding_matches( + existing, &record, + ) + ); + if !binding_matches { + return Ok(Err(invalid_request( + "导入 OAuth Provider 时修改了 Client ID、端点或 redirect_uri,必须提供 client_secret", + ))); + } + } + } + } + oauth_provider_types.insert(provider_type); + } + + for imported_config_item in imported_system_configs { + let config = imported_config_item.value; + let normalized_key = normalize_imported_system_config_key(&config.key); + if credentials_not_exported + && (is_sensitive_admin_system_config_key(&normalized_key) + || is_interactive_export_private_system_config_key(&normalized_key)) + { + continue; + } + let exists = existing_system_config_keys.contains(&normalized_key); + match (exists, merge_mode) { + (true, AdminImportMergeMode::Skip) => continue, + (true, AdminImportMergeMode::Error) => { + return Ok(Err(invalid_request(format!( + "SystemConfig '{normalized_key}' 已存在" + )))); + } + _ => {} + } + let request_body = serde_json::to_vec(&json!({ + "value": config.value, + "description": config.description, + })) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let update = routed!(parse_admin_system_config_update(&config.key, &request_body)); + if is_sensitive_admin_system_config_key(&update.normalized_key) + && update.value.as_str().is_some_and(|raw| !raw.is_empty()) + { + let Some(_) = self.encrypt_system_config_secret( + &update.normalized_key, + update + .value + .as_str() + .expect("sensitive imported config value was a string"), + ) else { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "系统配置写入需要可用的加密密钥" }), + ))); + }; + } + existing_system_config_keys.insert(normalized_key); + } + + Ok(Ok(())) + } + + async fn prevalidate_admin_system_users_import( + &self, + request_body: &[u8], + operator_id: Option<&str>, + mode: SystemImportMode, + ) -> Result, GatewayError> { + macro_rules! invalid { + ($expr:expr) => { + match $expr { + Ok(value) => value, + Err(detail) => return Ok(Err(invalid_request(detail))), + } + }; + } + + let root = match serde_json::from_slice::(request_body) { + Ok(Value::Object(map)) => map, + _ => return Ok(Err(invalid_request("请求数据验证失败"))), + }; + let merge_mode = match serde_json::from_value::( + root.get("merge_mode").cloned().unwrap_or(Value::Null), + ) { + Ok(value) => value, + Err(_) => { + return Ok(Err(invalid_request( + "merge_mode 仅支持 skip / overwrite / error", + ))) + } + }; + let users_export_version = invalid!( + validate_imported_system_users_export_version_for_mode(root.get("version"), mode) + ); + let empty = Vec::new(); + let users = match root.get("users") { + Some(Value::Array(items)) => items, + Some(_) => return Ok(Err(invalid_request("users 必须是数组"))), + None => &empty, + }; + let standalone_keys = match root.get("standalone_keys") { + Some(Value::Array(items)) => items, + Some(_) => return Ok(Err(invalid_request("standalone_keys 必须是数组"))), + None => &empty, + }; + let imported_user_groups = match root.get("user_groups") { + Some(Value::Array(items)) => items, + Some(_) => return Ok(Err(invalid_request("user_groups 必须是数组"))), + None => &empty, + }; + // Aggregate rollback restores identity/configuration only. Runtime usage rows are left + // untouched so a request that completes concurrently cannot be overwritten by the + // checkpoint. Skip both parsing and conflict validation in this mode; the rollback body + // also strips these fields before it reaches the regular importer. + let usage_aggregate_snapshot = if mode.is_rollback_checkpoint() { + AdminSystemUsageAggregateSnapshot::default() + } else { + let supplemental = invalid!(build_imported_user_usage_total_aggregates( + users, + root.get("exported_at") + )); + let snapshot = invalid!(build_imported_usage_aggregate_snapshot( + root.get("usage_aggregates"), + &supplemental, + )); + invalid!(validate_imported_usage_aggregate_storage(&snapshot)); + snapshot + }; + + let default_group_id = self.effective_default_user_group_id().await?; + let mut groups_by_name = self + .list_user_groups() + .await? + .into_iter() + .map(|group| { + ( + aether_data::repository::users::normalize_user_group_name(&group.name) + .to_ascii_lowercase(), + group, + ) + }) + .collect::>(); + let mut imported_group_id_map = BTreeMap::::new(); + let mut imported_group_name_map = BTreeMap::::new(); + for (index, raw_group) in imported_user_groups.iter().enumerate() { + let group = invalid!(imported_object_field( + raw_group, + &format!("user_groups[{index}]"), + )); + let (export_id, normalized_name, record) = invalid!(build_imported_user_group_record( + group, + &format!("user_groups[{index}]") + )); + if default_group_id + .as_deref() + .is_some_and(|group_id| export_id.as_deref() == Some(group_id)) + || normalized_name == "default" + { + if let Some(default_group_id) = default_group_id.as_ref() { + if let Some(export_id) = export_id { + imported_group_id_map.insert(export_id, default_group_id.clone()); + } + imported_group_name_map.insert(normalized_name, default_group_id.clone()); + } + continue; + } + let existing_by_id = mode + .is_rollback_checkpoint() + .then(|| { + export_id.as_deref().and_then(|export_id| { + groups_by_name.values().find(|group| group.id == export_id) + }) + }) + .flatten(); + if mode.is_rollback_checkpoint() && export_id.is_some() && existing_by_id.is_none() { + return Ok(Err(invalid_request(format!( + "回滚检查点用户组 '{}' 不存在;拒绝按名称匹配", + export_id.as_deref().unwrap_or_default() + )))); + } + if let Some(existing) = existing_by_id.or_else(|| groups_by_name.get(&normalized_name)) + { + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "用户组 '{}' 已存在", + existing.name + )))); + } + if let Some(export_id) = export_id { + imported_group_id_map.insert(export_id, existing.id.clone()); + } + imported_group_name_map.insert(normalized_name, existing.id.clone()); + } else { + let synthetic_id = format!("prevalidated-group-{index}"); + let stored = aether_data::repository::users::StoredUserGroup::new( + synthetic_id.clone(), + record.name.clone(), + normalized_name.clone(), + record.description.clone(), + record.priority, + record.allowed_providers.clone().map(Value::from), + record.allowed_providers_mode.clone(), + record.allowed_api_formats.clone().map(Value::from), + record.allowed_api_formats_mode.clone(), + record.allowed_models.clone().map(Value::from), + record.allowed_models_mode.clone(), + record.rate_limit, + record.rate_limit_mode.clone(), + None, + None, + ) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if let Some(export_id) = export_id { + imported_group_id_map.insert(export_id, synthetic_id.clone()); + } + imported_group_name_map.insert(normalized_name.clone(), synthetic_id); + groups_by_name.insert(normalized_name, stored); + } + } + + let standalone_owner_id = match operator_id { + Some(candidate) => match self.find_user_auth_by_id(candidate).await? { + Some(user) if crate::roles::is_full_admin_role(&user.role) => Some(user.id), + _ => None, + }, + None => None, + }; + let mut simulated_users_by_id = BTreeMap::::new(); + let mut simulated_email_owners = BTreeMap::::new(); + let mut simulated_username_owners = BTreeMap::::new(); + let mut released_emails = BTreeSet::::new(); + let mut released_usernames = BTreeSet::::new(); + let mut api_keys_by_hash = BTreeMap::::new(); + let mut imported_api_key_hashes = BTreeSet::::new(); + let mut imported_user_id_map = BTreeMap::::new(); + let mut imported_api_key_id_map = BTreeMap::::new(); + for user in self.list_export_users().await? { + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + SimulatedImportedUser { + id: user.id, + email: user.email, + username: user.username, + role: user.role, + existed_before_import: true, + }, + ); + } + #[cfg(test)] + if let Some(store) = self.app().auth_user_store.as_ref() { + for user in store.lock().expect("auth user store should lock").values() { + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + simulated_imported_user_from_auth_record(user), + ); + } + } + for record in self.list_auth_api_key_export_standalone_records().await? { + let simulated = SimulatedImportedApiKey { + owner_id: record.user_id, + is_standalone: true, + target_id: record.api_key_id.clone(), + existed_before_import: true, + }; + api_keys_by_hash.insert(record.key_hash, simulated.clone()); + if mode.is_rollback_checkpoint() { + api_keys_by_hash + .entry(imported_api_key_tombstone(&record.api_key_id)) + .or_insert(simulated); + } + } + let now_unix_secs = SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0); + + for (index, raw_user) in users.iter().enumerate() { + let user = invalid!(imported_object_field(raw_user, &format!("users[{index}]"))); + let source_user_id = invalid!(imported_optional_string(user.get("id"))); + invalid!(validate_rollback_user_source_id( + mode, + source_user_id.as_deref(), + )); + let Some(role) = invalid!(normalize_imported_system_user_role(user.get("role"), mode)) + else { + invalid!(imported_optional_string(user.get("email"))); + invalid!(imported_optional_string(user.get("username"))); + continue; + }; + let email = invalid!(imported_optional_string(user.get("email"))) + .map(|value| value.to_ascii_lowercase()); + let username = invalid!(imported_optional_string(user.get("username"))) + .or_else(|| { + email.as_ref().map(|value| { + value + .split('@') + .next() + .unwrap_or(value.as_str()) + .to_string() + }) + }) + .unwrap_or_else(|| format!("imported-user-{index}")); + invalid!(imported_optional_bool(user.get("email_verified"))); + invalid!(resolve_imported_password_hash( + user, + users_export_version, + mode, + )); + let allowed_providers = invalid!(normalize_imported_user_string_list( + user, + "allowed_providers" + )); + let allowed_api_formats = invalid!(normalize_imported_user_api_formats( + user, + "allowed_api_formats" + )); + let allowed_models = + invalid!(normalize_imported_user_string_list(user, "allowed_models")); + let rate_limit = invalid!(imported_optional_i32(user.get("rate_limit"), "rate_limit")); + invalid!(imported_user_list_policy_mode( + user, + "allowed_providers_mode", + "allowed_providers", + &allowed_providers, + )); + invalid!(imported_user_list_policy_mode( + user, + "allowed_api_formats_mode", + "allowed_api_formats", + &allowed_api_formats, + )); + invalid!(imported_user_list_policy_mode( + user, + "allowed_models_mode", + "allowed_models", + &allowed_models, + )); + invalid!(imported_user_rate_limit_policy_mode( + user, + "rate_limit_mode", + "rate_limit", + rate_limit, + )); + let group_ids = invalid!(resolve_imported_user_group_ids( + user, + &imported_group_id_map, + &imported_group_name_map, + &groups_by_name, + )); + if user.contains_key("group_ids") || user.contains_key("group_names") { + let group_ids = self.include_default_user_group_ids(&group_ids).await?; + let known_group_ids = groups_by_name + .values() + .map(|group| group.id.as_str()) + .collect::>(); + if group_ids + .iter() + .any(|group_id| !known_group_ids.contains(group_id.as_str())) + { + return Ok(Err(invalid_request(format!( + "用户 '{}' 的用户组不存在", + email.clone().unwrap_or(username.clone()) + )))); + } + } + invalid!(imported_optional_bool(user.get("is_active"))); + invalid!(imported_optional_json_object( + user.get("model_capability_settings"), + "model_capability_settings", + )); + invalid!(imported_optional_json_object( + user.get("feature_settings"), + "feature_settings", + ) + .and_then(normalize_admin_feature_settings)); + let wallet = match user.get("wallet") { + Some(Value::Object(map)) => Some(map), + Some(Value::Null) | None => None, + Some(_) => return Ok(Err(invalid_request("wallet 必须是对象"))), + }; + if let Some(wallet) = wallet { + invalid!(normalize_imported_wallet_target(Some(wallet), false)); + } + + // Rollback checkpoints originate from this deployment and carry stable user IDs. + // Never fall back to mutable email/username fields: if the checkpoint ID disappeared + // concurrently, guessing could overwrite an unrelated user. + let existing_user = if mode.is_rollback_checkpoint() { + let mut existing = source_user_id + .as_deref() + .and_then(|user_id| simulated_users_by_id.get(user_id).cloned()); + if existing.is_none() { + if let Some(source_user_id) = source_user_id.as_deref() { + if let Some(record) = self.find_user_auth_by_id(source_user_id).await? { + let simulated = simulated_imported_user_from_auth_record(&record); + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + simulated.clone(), + ); + existing = Some(simulated); + } + } + } + if existing.is_none() { + let source_user_id = source_user_id.as_deref().unwrap_or_default(); + return Ok(Err(invalid_request(format!( + "回滚检查点用户 '{source_user_id}' 不存在;拒绝按 email/username 匹配" + )))); + } + existing + } else { + let mut existing_user = None; + if let Some(email) = email.as_deref() { + if let Some(user_id) = simulated_imported_user_id_by_identifier( + &simulated_email_owners, + &simulated_username_owners, + email, + ) { + existing_user = simulated_users_by_id.get(&user_id).cloned(); + } else if !released_emails.contains(email) + && !released_usernames.contains(email) + { + if let Some(record) = self.find_user_auth_by_identifier(email).await? { + let simulated = simulated_imported_user_from_auth_record(&record); + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + simulated.clone(), + ); + existing_user = Some(simulated); + } + } + } + if existing_user.is_none() { + if let Some(user_id) = simulated_imported_user_id_by_identifier( + &simulated_email_owners, + &simulated_username_owners, + &username, + ) { + existing_user = simulated_users_by_id.get(&user_id).cloned(); + } else if !released_emails.contains(&username) + && !released_usernames.contains(&username) + { + if let Some(record) = self.find_user_auth_by_identifier(&username).await? { + let simulated = simulated_imported_user_from_auth_record(&record); + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + simulated.clone(), + ); + existing_user = Some(simulated); + } + } + } + existing_user + }; + + let label = email.clone().unwrap_or(username.clone()); + let simulated_user = if let Some(existing) = existing_user { + if imported_existing_user_is_protected(&existing.role, mode) { + continue; + } + match merge_mode { + AdminImportMergeMode::Skip => continue, + AdminImportMergeMode::Error => { + return Ok(Err(invalid_request(format!("用户 '{label}' 已存在")))); + } + AdminImportMergeMode::Overwrite => {} + } + if let Some(email) = email.as_deref() { + let taken_in_simulation = simulated_email_owners + .get(email) + .is_some_and(|owner_id| owner_id != &existing.id); + let taken_in_database = !released_emails.contains(email) + && self + .is_other_user_auth_email_taken(email, &existing.id) + .await?; + if taken_in_simulation || taken_in_database { + return Ok(Err(invalid_request(format!("邮箱已存在: {email}")))); + } + } + let username_taken_in_simulation = simulated_username_owners + .get(&username) + .is_some_and(|owner_id| owner_id != &existing.id); + let username_taken_in_database = !released_usernames.contains(&username) + && self + .is_other_user_auth_username_taken(&username, &existing.id) + .await?; + if username_taken_in_simulation || username_taken_in_database { + return Ok(Err(invalid_request(format!("用户名已存在: {username}")))); + } + SimulatedImportedUser { + id: existing.id, + email: email.clone().or(existing.email), + username, + role, + existed_before_import: existing.existed_before_import, + } + } else { + if email.as_ref().is_some_and(|email| { + simulated_email_owners.contains_key(email) + || (!released_emails.contains(email) + && simulated_username_owners.contains_key(email)) + }) || simulated_username_owners.contains_key(&username) + { + return Ok(Err(invalid_request(format!("用户 '{label}' 已存在")))); + } + SimulatedImportedUser { + id: format!("prevalidated-user-{index}"), + email, + username, + role, + existed_before_import: false, + } + }; + let user_id = simulated_user.id.clone(); + let existed_before_import = simulated_user.existed_before_import; + replace_simulated_imported_user( + &mut simulated_users_by_id, + &mut simulated_email_owners, + &mut simulated_username_owners, + &mut released_emails, + &mut released_usernames, + simulated_user, + ); + if let Some(source_user_id) = source_user_id { + invalid!(insert_imported_id_mapping( + &mut imported_user_id_map, + source_user_id, + user_id.clone(), + "users[].id", + )); + } + + if existed_before_import { + for record in self + .list_auth_api_key_export_records_by_user_ids(std::slice::from_ref(&user_id)) + .await? + .into_iter() + .filter(|record| !record.is_standalone) + { + let simulated = SimulatedImportedApiKey { + owner_id: user_id.clone(), + is_standalone: false, + target_id: record.api_key_id.clone(), + existed_before_import: true, + }; + api_keys_by_hash + .entry(record.key_hash) + .or_insert_with(|| simulated.clone()); + if mode.is_rollback_checkpoint() { + api_keys_by_hash + .entry(imported_api_key_tombstone(&record.api_key_id)) + .or_insert(simulated); + } + } + } + + let imported_api_keys = match user.get("api_keys") { + Some(Value::Array(items)) => items, + Some(_) => return Ok(Err(invalid_request("api_keys 必须是数组"))), + None => &empty, + }; + for (key_index, raw_key) in imported_api_keys.iter().enumerate() { + let key = invalid!(imported_object_field( + raw_key, + &format!("users[{index}].api_keys[{key_index}]"), + )); + invalid!(self.prevalidate_imported_auth_api_key(key, users_export_version, mode,)); + let source_api_key_id = invalid!(imported_optional_string(key.get("api_key_id"))); + let Some(key_material) = invalid!(self + .resolve_imported_system_user_api_key_material( + key, + users_export_version, + mode, + )) + else { + continue; + }; + let key_hash = key_material.key_hash; + if !imported_api_key_hashes.insert(key_hash.clone()) { + return Ok(Err(invalid_request("API Key 在导入文档中重复"))); + } + if !api_keys_by_hash.contains_key(&key_hash) { + if let Some(snapshot) = self + .app() + .data + .read_auth_api_key_snapshot_by_key_hash_strong(&key_hash, now_unix_secs) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + api_keys_by_hash.insert( + key_hash.clone(), + SimulatedImportedApiKey { + owner_id: snapshot.user_id, + is_standalone: snapshot.api_key_is_standalone, + target_id: snapshot.api_key_id, + existed_before_import: true, + }, + ); + } + } + if let Some(existing_key) = api_keys_by_hash.get(&key_hash) { + if existing_key.owner_id != user_id || existing_key.is_standalone { + return Ok(Err(invalid_request( + "API Key 已存在且属于其他用户或独立余额 Key", + ))); + } + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request(format!( + "用户 '{label}' 的 API Key 已存在" + )))); + } + if merge_mode == AdminImportMergeMode::Overwrite { + if let Some(source_api_key_id) = source_api_key_id { + invalid!(insert_imported_id_mapping( + &mut imported_api_key_id_map, + source_api_key_id, + existing_key.target_id.clone(), + "api_key_id", + )); + } + } + } else { + let target_id = format!("prevalidated-api-key-{index}-{key_index}"); + if let Some(source_api_key_id) = source_api_key_id { + invalid!(insert_imported_id_mapping( + &mut imported_api_key_id_map, + source_api_key_id, + target_id.clone(), + "api_key_id", + )); + } + api_keys_by_hash.insert( + key_hash, + SimulatedImportedApiKey { + owner_id: user_id.clone(), + is_standalone: false, + target_id, + existed_before_import: false, + }, + ); + } + } + } + + if let Some(standalone_owner_id) = standalone_owner_id { + for (index, raw_key) in standalone_keys.iter().enumerate() { + let key = invalid!(imported_object_field( + raw_key, + &format!("standalone_keys[{index}]"), + )); + invalid!(self.prevalidate_imported_auth_api_key(key, users_export_version, mode,)); + let source_api_key_id = invalid!(imported_optional_string(key.get("api_key_id"))); + let Some(key_material) = invalid!(self + .resolve_imported_system_user_api_key_material( + key, + users_export_version, + mode, + )) + else { + continue; + }; + let key_hash = key_material.key_hash; + if !imported_api_key_hashes.insert(key_hash.clone()) { + return Ok(Err(invalid_request("API Key 在导入文档中重复"))); + } + let wallet = match key.get("wallet") { + Some(Value::Object(map)) => Some(map), + Some(Value::Null) | None => None, + Some(_) => return Ok(Err(invalid_request("wallet 必须是对象"))), + }; + let unlimited = + invalid!(imported_optional_bool(key.get("unlimited"))).unwrap_or(false); + if let Some(wallet) = wallet { + invalid!(normalize_imported_wallet_target(Some(wallet), unlimited)); + } + if !api_keys_by_hash.contains_key(&key_hash) { + if let Some(snapshot) = self + .app() + .data + .read_auth_api_key_snapshot_by_key_hash_strong(&key_hash, now_unix_secs) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + api_keys_by_hash.insert( + key_hash.clone(), + SimulatedImportedApiKey { + owner_id: snapshot.user_id, + is_standalone: snapshot.api_key_is_standalone, + target_id: snapshot.api_key_id, + existed_before_import: true, + }, + ); + } + } + if let Some(existing_key) = api_keys_by_hash.get(&key_hash) { + if !existing_key.is_standalone { + return Ok(Err(invalid_request("独立余额 Key 已存在且属于普通用户"))); + } + if merge_mode == AdminImportMergeMode::Error { + return Ok(Err(invalid_request("独立余额 Key 已存在"))); + } + if merge_mode == AdminImportMergeMode::Overwrite { + if let Some(source_api_key_id) = source_api_key_id { + invalid!(insert_imported_id_mapping( + &mut imported_api_key_id_map, + source_api_key_id, + existing_key.target_id.clone(), + "api_key_id", + )); + } + } + } else { + let target_id = format!("prevalidated-standalone-key-{index}"); + if let Some(source_api_key_id) = source_api_key_id { + invalid!(insert_imported_id_mapping( + &mut imported_api_key_id_map, + source_api_key_id, + target_id.clone(), + "api_key_id", + )); + } + api_keys_by_hash.insert( + key_hash, + SimulatedImportedApiKey { + owner_id: standalone_owner_id.clone(), + is_standalone: true, + target_id, + existed_before_import: false, + }, + ); + } + } + } + + invalid!(validate_imported_usage_aggregate_dimensions( + &usage_aggregate_snapshot, + &imported_user_id_map, + &imported_api_key_id_map, + )); + + if merge_mode == AdminImportMergeMode::Error { + let persisted_user_target_ids = simulated_users_by_id + .values() + .filter(|user| user.existed_before_import) + .map(|user| user.id.clone()) + .collect::>(); + let persisted_api_key_target_ids = api_keys_by_hash + .values() + .filter(|key| key.existed_before_import) + .map(|key| key.target_id.clone()) + .collect::>(); + invalid!( + self.prevalidate_imported_usage_aggregate_conflicts( + &usage_aggregate_snapshot, + &imported_user_id_map, + &imported_api_key_id_map, + &persisted_user_target_ids, + &persisted_api_key_target_ids, + ) + .await + ); + } + + Ok(Ok(())) + } + + fn prevalidate_imported_auth_api_key( + &self, + key: &Map, + users_export_version: (u32, u32), + mode: SystemImportMode, + ) -> Result<(), String> { + if self + .resolve_imported_system_user_api_key_material(key, users_export_version, mode)? + .is_none() + { + return Ok(()); + } + imported_optional_string(key.get("api_key_id"))?; + imported_optional_string(key.get("name"))?; + normalize_imported_user_string_list(key, "allowed_providers")?; + normalize_imported_user_api_formats(key, "allowed_api_formats")?; + normalize_imported_user_string_list(key, "allowed_models")?; + normalize_imported_user_ip_rules(key)?; + imported_optional_i32(key.get("rate_limit"), "rate_limit")?; + let concurrent_limit = + imported_optional_i32(key.get("concurrent_limit"), "concurrent_limit")?; + if concurrent_limit.is_some_and(|value| value < 0) { + return Err("concurrent_limit 必须是非负整数".to_string()); + } + imported_optional_bool(key.get("is_active"))?; + imported_rfc3339_to_unix_secs(key.get("expires_at"), "expires_at")?; + imported_optional_bool(key.get("auto_delete_on_expiry"))?; + if let Some(value) = imported_optional_u64(key.get("total_requests"), "total_requests")? { + validate_imported_u64_storage(value, "total_requests")?; + } + if let Some(value) = imported_optional_u64(key.get("total_tokens"), "total_tokens")? { + validate_imported_u64_storage(value, "total_tokens")?; + } + imported_optional_f64(key.get("total_cost_usd"), "total_cost_usd")?; + imported_optional_json_object(key.get("feature_settings"), "feature_settings") + .and_then(normalize_admin_feature_settings)?; + Ok(()) + } + + async fn prevalidate_imported_usage_aggregate_conflicts( + &self, + snapshot: &AdminSystemUsageAggregateSnapshot, + user_id_map: &BTreeMap, + api_key_id_map: &BTreeMap, + persisted_user_target_ids: &BTreeSet, + persisted_api_key_target_ids: &BTreeSet, + ) -> Result<(), String> { + if snapshot.stats_daily.is_empty() + && snapshot.stats_user_daily.is_empty() + && snapshot.stats_daily_api_key.is_empty() + { + return Ok(()); + } + + if !self.app().data.has_backends() { + return Ok(()); + } + + let persisted_user_ids = user_id_map + .iter() + .filter(|(_, target_id)| persisted_user_target_ids.contains(*target_id)) + .map(|(source_id, target_id)| (source_id.clone(), target_id.clone())) + .collect::>(); + let persisted_api_key_ids = api_key_id_map + .iter() + .filter(|(_, target_id)| persisted_api_key_target_ids.contains(*target_id)) + .map(|(source_id, target_id)| (source_id.clone(), target_id.clone())) + .collect::>(); + + match self + .import_admin_system_usage_aggregates( + snapshot, + &persisted_user_ids, + &persisted_api_key_ids, + AdminSystemUsageAggregateImportMode::ValidateError, + ) + .await + { + Err(GatewayError::Client { message, .. }) => Err(message), + Err(_) => { + tracing::error!( + event_name = "admin_system_import_usage_prevalidation_error", + operation = "prevalidate_usage_aggregate_conflicts", + error_category = "repository_unavailable", + "admin system import usage prevalidation failed" + ); + Err("Usage aggregate data temporarily unavailable".to_string()) + } + Ok(_) => Ok(()), + } + } + pub(crate) async fn import_admin_system_data( &self, request_body: &Bytes, operator_id: Option<&str>, + ) -> Result, GatewayError> { + self.import_admin_system_data_with_mode( + request_body, + operator_id, + SystemImportMode::InteractiveUpload, + ) + .await + } + + pub(crate) async fn restore_admin_system_data_backup( + &self, + request_body: &Bytes, + operator_id: Option<&str>, + _authority: crate::backup::executor::BackupRestoreAuthority, + ) -> Result, GatewayError> { + self.import_admin_system_data_with_mode( + request_body, + operator_id, + SystemImportMode::RecoveryBackup, + ) + .await + } + + async fn import_admin_system_data_with_mode( + &self, + request_body: &Bytes, + operator_id: Option<&str>, + mode: SystemImportMode, ) -> Result, GatewayError> { if !self.has_global_model_data_reader() || !self.has_global_model_data_writer() @@ -1247,16 +4463,127 @@ impl<'a> AdminAppState<'a> { Err(err) => return Ok(Err(err)), }; - let config_result = match self.import_admin_system_config(&config_body).await? { - Ok(payload) => payload, - Err(err) => return Ok(Err(err)), - }; - let users_result = match self - .import_admin_system_users(&users_body, operator_id) + match self + .prevalidate_admin_system_config_import(&config_body, mode) .await? { - Ok(payload) => payload, + Ok(()) => {} Err(err) => return Ok(Err(err)), + } + match self + .prevalidate_admin_system_users_import(&users_body, operator_id, mode) + .await? + { + Ok(()) => {} + Err(err) => return Ok(Err(err)), + } + + // The config and users repositories are intentionally exposed as independent write + // handles, so this aggregate operation cannot share a database transaction across all + // supported drivers. Capture both sides immediately before the first write. Interactive + // imports use a redacted checkpoint. Recovery restores are authorized to + // hold a credential-bearing checkpoint briefly in memory; otherwise a failed restore + // could not put an overwritten secret back. The import lock held by the route serializes + // other aggregate imports while this checkpoint is being applied. + let checkpoint_export_mode = if mode == SystemImportMode::RecoveryBackup { + SystemExportMode::RecoveryBackup + } else { + SystemExportMode::RollbackCheckpoint + }; + let rollback_mode = if mode == SystemImportMode::RecoveryBackup { + SystemImportMode::RecoveryRollbackCheckpoint + } else { + SystemImportMode::RollbackCheckpoint + }; + let config_checkpoint = self + .build_admin_system_config_export_payload(checkpoint_export_mode) + .await?; + let users_checkpoint = self + .build_admin_system_users_export_payload(checkpoint_export_mode) + .await?; + + let mut mutation_journal = AggregateMutationJournal::default(); + let config_result = match self + .import_admin_system_config_with_mode(&config_body, mode, Some(&mut mutation_journal)) + .await + { + Ok(Ok(payload)) => payload, + Ok(Err(original)) => { + match self + .rollback_aggregate_config(&config_checkpoint, rollback_mode, &mutation_journal) + .await + { + Ok(()) => return Ok(Err(original)), + Err(rollback_error) => { + return Err(aggregate_rollback_http_error( + "配置阶段", + &original, + rollback_error, + )); + } + } + } + Err(error) => { + let original_error = error.clone(); + self.rollback_aggregate_config( + &config_checkpoint, + rollback_mode, + &mutation_journal, + ) + .await + .map_err(|rollback_error| { + aggregate_rollback_error("配置阶段", original_error, rollback_error) + })?; + return Err(error); + } + }; + + let users_result = match self + .import_admin_system_users_with_mode( + &users_body, + operator_id, + mode, + Some(&mut mutation_journal), + ) + .await + { + Ok(Ok(payload)) => payload, + Ok(Err(original)) => { + match self + .rollback_aggregate_import( + &config_checkpoint, + &users_checkpoint, + operator_id, + rollback_mode, + &mutation_journal, + ) + .await + { + Ok(()) => return Ok(Err(original)), + Err(rollback_error) => { + return Err(aggregate_rollback_http_error( + "用户阶段", + &original, + rollback_error, + )); + } + } + } + Err(error) => { + let original_error = error.clone(); + self.rollback_aggregate_import( + &config_checkpoint, + &users_checkpoint, + operator_id, + rollback_mode, + &mutation_journal, + ) + .await + .map_err(|rollback_error| { + aggregate_rollback_error("用户阶段", original_error, rollback_error) + })?; + return Err(error); + } }; Ok(Ok(json!({ @@ -1266,9 +4593,772 @@ impl<'a> AdminAppState<'a> { }))) } + /// Restore a config checkpoint using overwrite semantics. A redacted checkpoint preserves + /// existing encrypted values; a recovery checkpoint carries the original credentials. + async fn rollback_aggregate_config( + &self, + checkpoint: &Value, + rollback_mode: SystemImportMode, + mutation_journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let cleanup_result = self.rollback_created_config(mutation_journal).await; + let body = build_aggregate_rollback_body_with_options( + checkpoint, + true, + cleanup_result.skip_ldap_restore, + )?; + let restore_result = match self + .import_admin_system_config_with_mode(&body, rollback_mode, None) + .await + { + Ok(Ok(_)) => Ok(()), + Ok(Err(_)) => Err(GatewayError::Internal( + "aggregate config rollback rejected".to_string(), + )), + Err(error) => Err(error), + }; + combine_rollback_results(cleanup_result.result, restore_result, "config") + } + + async fn rollback_created_config( + &self, + journal: &AggregateMutationJournal, + ) -> ConfigCleanupOutcome { + let mut failures = Vec::new(); + let mut skip_ldap_restore = false; + + if let Some(expected) = journal.created_ldap_config.as_ref() { + match self.delete_ldap_module_config_if_matches(expected).await { + Ok(true) => {} + Ok(false) => { + // Even if the follow-up read observes no row, another writer can create one + // between that read and the checkpoint restore. Skip LDAP restoration for + // every non-successful compare/delete result to avoid a TOCTOU overwrite. + skip_ldap_restore = true; + // A missing row means another cleanup attempt already removed it. If a row + // remains, however, it was changed concurrently and must not be deleted by + // an owner-blind rollback. + match self.get_ldap_module_config().await { + Ok(None) => {} + Ok(Some(_)) => failures.push(GatewayError::Internal( + "LDAP configuration changed during aggregate rollback".to_string(), + )), + Err(error) => failures.push(error), + } + } + Err(error) => { + skip_ldap_restore = true; + failures.push(error); + } + } + } + + for (provider_id, model_id) in &journal.provider_model_ids { + if let Err(error) = self + .delete_admin_provider_model(provider_id, model_id) + .await + { + failures.push(error); + } + } + for (_, key_id) in &journal.provider_key_ids { + if let Err(error) = self.delete_provider_catalog_key(key_id).await { + failures.push(error); + } + } + for (_, endpoint_id) in &journal.provider_endpoint_ids { + if let Err(error) = self.delete_provider_catalog_endpoint(endpoint_id).await { + failures.push(error); + } + } + for provider_id in &journal.provider_ids { + let endpoint_ids = journal + .provider_endpoint_ids + .iter() + .filter(|(owner_id, _)| owner_id == provider_id) + .map(|(_, id)| id.clone()) + .collect::>(); + let key_ids = journal + .provider_key_ids + .iter() + .filter(|(owner_id, _)| owner_id == provider_id) + .map(|(_, id)| id.clone()) + .collect::>(); + if let Err(error) = self + .cleanup_deleted_provider_catalog_refs(provider_id, true, &endpoint_ids, &key_ids) + .await + { + failures.push(error); + } + if let Err(error) = self + .app() + .delete_provider_catalog_provider(provider_id) + .await + { + failures.push(error); + } + } + for global_model_id in &journal.global_model_ids { + if let Err(error) = self.delete_admin_global_model(global_model_id).await { + failures.push(error); + } + } + for provider_type in &journal.oauth_provider_types { + match self + .delete_oauth_provider_config_if_unlinked(provider_type) + .await + { + Ok(_) => {} + Err(error) => failures.push(error), + } + } + for key in &journal.system_config_keys { + if let Err(error) = self.delete_system_config_value(key).await { + failures.push(error); + } + } + + let result = if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "aggregate config mutation cleanup failed for {} object(s)", + failures.len() + ))) + }; + ConfigCleanupOutcome { + result, + skip_ldap_restore, + } + } + + async fn rollback_created_users( + &self, + journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let mut failures = Vec::new(); + let mut blocked_user_ids = BTreeSet::new(); + let mut blocked_api_key_ids = BTreeSet::new(); + + // Remove wallets before their owning API keys/users. The owner predicate is part of the + // delete operation, so a wallet cannot be detached and accidentally reclaimed by another + // object between the journal lookup and compensation. + for ((user_id, wallet_id), expected) in &journal.user_wallet_snapshots { + match self + .delete_wallet_if_snapshot_matches_and_unreferenced( + expected, + WalletLookupKey::UserId(user_id.as_str()), + ) + .await + { + Ok(true) => {} + Ok(false) => { + blocked_user_ids.insert(user_id.clone()); + failures.push(GatewayError::Internal(format!( + "import rollback could not delete wallet {wallet_id} for user {user_id}" + ))); + } + Err(error) => { + blocked_user_ids.insert(user_id.clone()); + failures.push(error); + } + } + } + for ((api_key_id, wallet_id), expected) in &journal.api_key_wallet_snapshots { + match self + .delete_wallet_if_snapshot_matches_and_unreferenced( + expected, + WalletLookupKey::ApiKeyId(api_key_id.as_str()), + ) + .await + { + Ok(true) => {} + Ok(false) => { + blocked_api_key_ids.insert(api_key_id.clone()); + if let Some((user_id, _)) = journal + .user_api_key_ids + .iter() + .find(|(_, candidate_api_key_id)| candidate_api_key_id == api_key_id) + { + // A failed API-key-wallet compensation also blocks + // deleting the owning user. Otherwise the later + // user rollback can remove the key while leaving its + // funded wallet orphaned. + blocked_user_ids.insert(user_id.clone()); + } + failures.push(GatewayError::Internal(format!( + "import rollback could not delete wallet {wallet_id} for API key {api_key_id}" + ))); + } + Err(error) => { + blocked_api_key_ids.insert(api_key_id.clone()); + if let Some((user_id, _)) = journal + .user_api_key_ids + .iter() + .find(|(_, candidate_api_key_id)| candidate_api_key_id == api_key_id) + { + blocked_user_ids.insert(user_id.clone()); + } + failures.push(error); + } + } + } + for (user_id, api_key_id) in &journal.user_api_key_ids { + if blocked_user_ids.contains(user_id) || blocked_api_key_ids.contains(api_key_id) { + continue; + } + if let Err(error) = self.delete_user_api_key(user_id, api_key_id).await { + failures.push(error); + } + } + for api_key_id in &journal.standalone_api_key_ids { + if blocked_api_key_ids.contains(api_key_id) { + continue; + } + if let Err(error) = self.delete_standalone_api_key(api_key_id).await { + failures.push(error); + } + } + for user_id in &journal.user_ids { + if blocked_user_ids.contains(user_id) { + continue; + } + if let Err(error) = self + .app() + .rollback_provisional_auth_user_with_wallet(user_id, None) + .await + { + failures.push(error); + } + } + for group_id in &journal.user_group_ids { + if let Err(error) = self.delete_user_group(group_id).await { + failures.push(error); + } + } + + if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "aggregate user mutation cleanup failed for {} object(s)", + failures.len() + ))) + } + } + + async fn rollback_existing_wallets( + &self, + journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let mut failures = Vec::new(); + + for ((user_id, wallet_id), mutation) in &journal.existing_user_wallets { + let Some(after) = mutation.after.as_ref() else { + failures.push(GatewayError::Internal(format!( + "import rollback has no verified post-state for existing user wallet {wallet_id} ({user_id})" + ))); + continue; + }; + match self + .restore_wallet_if_snapshot_matches( + &mutation.before, + after, + WalletLookupKey::UserId(user_id.as_str()), + ) + .await + { + Ok(true) => {} + Ok(false) => failures.push(GatewayError::Internal(format!( + "import rollback wallet CAS conflict for user {user_id}, wallet {wallet_id}" + ))), + Err(error) => failures.push(error), + } + } + + for ((api_key_id, wallet_id), mutation) in &journal.existing_api_key_wallets { + let Some(after) = mutation.after.as_ref() else { + failures.push(GatewayError::Internal(format!( + "import rollback has no verified post-state for existing API-key wallet {wallet_id} ({api_key_id})" + ))); + continue; + }; + match self + .restore_wallet_if_snapshot_matches( + &mutation.before, + after, + WalletLookupKey::ApiKeyId(api_key_id.as_str()), + ) + .await + { + Ok(true) => {} + Ok(false) => failures.push(GatewayError::Internal(format!( + "import rollback wallet CAS conflict for API key {api_key_id}, wallet {wallet_id}" + ))), + Err(error) => failures.push(error), + } + } + + if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "aggregate existing wallet rollback failed for {} object(s)", + failures.len() + ))) + } + } + + async fn capture_existing_user_mutation( + &self, + journal: &mut AggregateMutationJournal, + user: &aether_data::repository::users::StoredUserAuthRecord, + ) -> Result<(), GatewayError> { + if journal.existing_users.contains_key(&user.id) { + return Ok(()); + } + let mut before_export = self.find_export_user_by_id(&user.id).await?; + let before_model_capability_settings = self + .app() + .read_user_model_capability_settings(&user.id) + .await?; + let before_feature_settings = before_export + .as_ref() + .and_then(|record| record.feature_settings.clone()) + .or(self.app().read_user_feature_settings(&user.id).await?); + #[cfg(test)] + if before_export.is_none() { + before_export = Some(synthetic_rollback_export_row( + user, + before_model_capability_settings.clone(), + before_feature_settings.clone(), + )?); + } + let mut before_group_ids = self + .list_user_groups_for_user(&user.id) + .await? + .into_iter() + .map(|group| group.id) + .collect::>(); + before_group_ids.sort(); + before_group_ids.dedup(); + // A role/active-state update revokes every key owned by the user. Capture those keys before + // the first user write so a later import failure can restore them through their own CAS + // path instead of leaving an unrelated key permanently disabled. + let existing_api_keys = self + .list_auth_api_key_export_records_by_user_ids(std::slice::from_ref(&user.id)) + .await?; + for record in existing_api_keys { + if record.is_standalone || record.user_id != user.id { + continue; + } + let key = (user.id.clone(), record.api_key_id.clone()); + journal + .existing_user_api_keys + .entry(key) + .or_insert_with(|| ExistingApiKeyMutation { + before: record.clone(), + after: record, + }); + } + journal.existing_users.insert( + user.id.clone(), + ExistingUserMutation { + before_auth: user.clone(), + after_auth: user.clone(), + before_export: before_export.clone(), + after_export: before_export, + before_model_capability_settings: before_model_capability_settings.clone(), + after_model_capability_settings: before_model_capability_settings, + before_feature_settings: before_feature_settings.clone(), + after_feature_settings: before_feature_settings, + before_group_ids: before_group_ids.clone(), + after_group_ids: before_group_ids, + }, + ); + Ok(()) + } + + /// Refresh a journal entry after each successful user mutation. Keeping the latest + /// post-state lets compensation recover even when a later setter in the same user fails. + async fn refresh_existing_user_mutation( + &self, + mutation_journal: Option<&mut AggregateMutationJournal>, + user_id: &str, + ) -> Result<(), GatewayError> { + let Some(journal) = mutation_journal else { + return Ok(()); + }; + if !journal.existing_users.contains_key(user_id) { + return Ok(()); + } + let Some(auth) = self.find_user_auth_by_id(user_id).await? else { + return Ok(()); + }; + let security_state_changed = journal.existing_users.get(user_id).is_some_and(|mutation| { + mutation.after_auth.role != auth.role || mutation.after_auth.is_active != auth.is_active + }); + let model_capability_settings = self + .app() + .read_user_model_capability_settings(user_id) + .await?; + let mut export = self.find_export_user_by_id(user_id).await?; + let feature_settings = export + .as_ref() + .and_then(|record| record.feature_settings.clone()) + .or(self.app().read_user_feature_settings(user_id).await?); + #[cfg(test)] + if export.is_none() { + export = Some(synthetic_rollback_export_row( + &auth, + model_capability_settings.clone(), + feature_settings.clone(), + )?); + } + let mut group_ids = self + .list_user_groups_for_user(user_id) + .await? + .into_iter() + .map(|group| group.id) + .collect::>(); + group_ids.sort(); + group_ids.dedup(); + if let Some(mutation) = journal.existing_users.get_mut(user_id) { + mutation.after_auth = auth; + mutation.after_export = export; + mutation.after_model_capability_settings = model_capability_settings; + mutation.after_feature_settings = feature_settings; + mutation.after_group_ids = group_ids; + } + if security_state_changed { + // The user CAS revokes all active API keys in the same database transaction. Refresh + // the post-state of every pre-captured key so the later key CAS can undo that exact + // revocation while still refusing any key changed by another writer. + for record in self + .list_auth_api_key_export_records_by_user_ids(&[user_id.to_string()]) + .await? + { + if record.is_standalone || record.user_id != user_id { + continue; + } + if let Some(key_mutation) = journal + .existing_user_api_keys + .get_mut(&(user_id.to_string(), record.api_key_id.clone())) + { + key_mutation.after = record; + } + } + } + Ok(()) + } + + async fn refresh_existing_api_key_mutation( + &self, + mutation_journal: Option<&mut AggregateMutationJournal>, + user_id: Option<&str>, + api_key_id: &str, + standalone: bool, + ) -> Result<(), GatewayError> { + let Some(journal) = mutation_journal else { + return Ok(()); + }; + let record = if standalone { + self.find_auth_api_key_export_standalone_record_by_id(api_key_id) + .await? + } else { + self.list_auth_api_key_export_records_by_ids(&[api_key_id.to_string()]) + .await? + .into_iter() + .find(|record| { + !record.is_standalone + && user_id.is_none_or(|expected| record.user_id == expected) + }) + }; + let Some(record) = record else { + return Ok(()); + }; + if standalone { + if let Some(mutation) = journal.existing_standalone_api_keys.get_mut(api_key_id) { + mutation.after = record; + } + } else if let Some(user_id) = user_id { + if let Some(mutation) = journal + .existing_user_api_keys + .get_mut(&(user_id.to_string(), api_key_id.to_string())) + { + mutation.after = record; + } + } + Ok(()) + } + + async fn rollback_existing_user_groups( + &self, + journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let mut failures = Vec::new(); + for (group_id, mutation) in &journal.existing_user_groups { + if mutation.before == mutation.after { + continue; + } + match self + .restore_user_group_if_matches(&mutation.after, &mutation.before) + .await + { + Ok(true) => {} + Ok(false) => { + tracing::warn!( + group_id, + "skipping missing or concurrently changed user group rollback" + ); + } + Err(error) => failures.push(error), + } + } + if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "existing user group rollback failed for {} object(s)", + failures.len() + ))) + } + } + + async fn rollback_existing_users( + &self, + journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let mut failures = Vec::new(); + for (user_id, mutation) in &journal.existing_users { + let before = &mutation.before_auth; + let after = &mutation.after_auth; + // Passwords are intentionally excluded from the aggregate CAS because a nullable + // password hash has its own compare-and-write operation below. + let auth_state_changed = !before.matches_restore_state(after); + let export_state_changed = match (&mutation.before_export, &mutation.after_export) { + (Some(before_export), Some(after_export)) => { + !before_export.matches_restore_state(after_export) + || before_export.rate_limit != after_export.rate_limit + || before_export.rate_limit_mode != after_export.rate_limit_mode + } + (None, None) => false, + _ => true, + }; + let model_settings_changed = mutation.before_model_capability_settings + != mutation.after_model_capability_settings; + let feature_settings_changed = + mutation.before_feature_settings != mutation.after_feature_settings; + + if auth_state_changed + || export_state_changed + || model_settings_changed + || feature_settings_changed + { + let restore_result = match (&mutation.after_export, &mutation.before_export) { + (Some(expected_export), Some(restored_export)) => { + self.restore_local_auth_user_state_if_matches( + after, + before, + expected_export, + restored_export, + mutation.after_model_capability_settings.as_ref(), + mutation.before_model_capability_settings.clone(), + mutation.after_feature_settings.as_ref(), + mutation.before_feature_settings.clone(), + ) + .await + } + _ => { + tracing::warn!( + user_id, + "skipping user state rollback because export snapshot is unavailable" + ); + Ok(false) + } + }; + match restore_result { + Ok(true) => {} + Ok(false) => { + tracing::warn!( + user_id, + "skipping concurrently changed user state rollback" + ); + } + Err(error) => failures.push(error), + } + } + + if before.password_hash != after.password_hash { + match self + .restore_local_auth_user_password_hash_if_matches( + user_id, + after.password_hash.as_deref(), + before.password_hash.clone(), + chrono::Utc::now(), + ) + .await + { + Ok(true) => {} + Ok(false) => { + tracing::warn!(user_id, "skipping concurrently changed user password"); + } + Err(error) => failures.push(error), + } + } + if mutation.before_group_ids != mutation.after_group_ids { + match self + .restore_user_groups_if_matches( + user_id, + &mutation.after_group_ids, + &mutation.before_group_ids, + ) + .await + { + Ok(true) => {} + Ok(false) => { + tracing::warn!( + user_id, + "skipping concurrently changed user groups rollback" + ); + } + Err(error) => failures.push(error), + } + } + } + if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "existing user rollback failed for {} object(s)", + failures.len() + ))) + } + } + + async fn rollback_existing_api_keys( + &self, + journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let mut failures = Vec::new(); + for ((user_id, api_key_id), mutation) in &journal.existing_user_api_keys { + self.rollback_one_existing_api_key( + Some(user_id), + api_key_id, + false, + mutation, + &mut failures, + ) + .await?; + } + for (api_key_id, mutation) in &journal.existing_standalone_api_keys { + self.rollback_one_existing_api_key(None, api_key_id, true, mutation, &mut failures) + .await?; + } + if failures.is_empty() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "existing API key rollback failed for {} object(s)", + failures.len() + ))) + } + } + + async fn rollback_one_existing_api_key( + &self, + user_id: Option<&str>, + api_key_id: &str, + standalone: bool, + mutation: &ExistingApiKeyMutation, + failures: &mut Vec, + ) -> Result<(), GatewayError> { + let _ = (user_id, standalone); + if mutation.before == mutation.after { + return Ok(()); + } + + match self + .restore_api_key_if_matches(&mutation.after, &mutation.before) + .await + { + Ok(true) => {} + Ok(false) => { + tracing::warn!( + api_key_id, + "skipping API key rollback after deletion, identity, or concurrent-state conflict" + ); + } + Err(error) => failures.push(error), + } + Ok(()) + } + + /// Compensate only objects touched by this import, then restore config. Replaying the whole + /// user checkpoint is intentionally avoided because it could overwrite a concurrent admin + /// change made after the import wrote a row. + async fn rollback_aggregate_import( + &self, + config_checkpoint: &Value, + _users_checkpoint: &Value, + _operator_id: Option<&str>, + rollback_mode: SystemImportMode, + mutation_journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + let cleanup_result = self.rollback_created_users(mutation_journal).await; + let existing_wallet_result = self.rollback_existing_wallets(mutation_journal).await; + let existing_users_result = self.rollback_existing_users(mutation_journal).await; + let existing_groups_result = self.rollback_existing_user_groups(mutation_journal).await; + let existing_api_keys_result = self.rollback_existing_api_keys(mutation_journal).await; + let config_result = self + .rollback_aggregate_config(config_checkpoint, rollback_mode, mutation_journal) + .await; + let restore_result = combine_rollback_results( + existing_users_result, + existing_groups_result, + "aggregate existing users/groups", + ); + let restore_result = combine_rollback_results( + existing_api_keys_result, + restore_result, + "aggregate existing API keys", + ); + let restore_result = combine_rollback_results(restore_result, config_result, "aggregate"); + let restore_result = + combine_rollback_results(existing_wallet_result, restore_result, "aggregate wallets"); + combine_rollback_results(cleanup_result, restore_result, "aggregate") + } + pub(crate) async fn import_admin_system_config( &self, request_body: &Bytes, + ) -> Result, GatewayError> { + self.import_admin_system_config_with_mode( + request_body, + SystemImportMode::InteractiveUpload, + None, + ) + .await + } + + pub(crate) async fn restore_admin_system_config_backup( + &self, + request_body: &Bytes, + _authority: crate::backup::executor::BackupRestoreAuthority, + ) -> Result, GatewayError> { + self.import_admin_system_config_with_mode( + request_body, + SystemImportMode::RecoveryBackup, + None, + ) + .await + } + + async fn import_admin_system_config_with_mode( + &self, + request_body: &Bytes, + mode: SystemImportMode, + mut mutation_journal: Option<&mut AggregateMutationJournal>, ) -> Result, GatewayError> { macro_rules! invalid { ($expr:expr) => { @@ -1297,8 +5387,17 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + match self + .prevalidate_admin_system_config_import(request_body, mode) + .await? + { + Ok(()) => {} + Err(err) => return Ok(Err(err)), + } let parsed = routed!(parse_admin_system_config_import_request(request_body)); + let source_version = parsed.request.document.version.clone(); let root = parsed.root; + let credentials_not_exported = invalid!(imported_config_credentials_not_exported(&root)); let merge_mode = parsed.request.merge_mode; let imported_global_models = routed!( @@ -1329,11 +5428,19 @@ impl<'a> AdminAppState<'a> { // object, and turn a non-empty exported node reference into direct mode. This keeps a // clean-environment restore portable and prevents a late selector validation failure from // leaving the rest of the document partially imported. - let (imported_external_models_configs, imported_system_configs): (Vec<_>, Vec<_>) = + let (imported_external_models_configs, mut imported_system_configs): (Vec<_>, Vec<_>) = imported_system_configs.into_iter().partition(|item| { normalize_imported_system_config_key(&item.value.key) == ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY }); + // Destination-bound secrets must be applied after their destination fields. Recovery + // documents are user-controlled JSON and do not guarantee any ordering. + imported_system_configs.sort_by_key(|item| { + matches!( + normalize_imported_system_config_key(&item.value.key).as_str(), + "smtp_password" | "module.bark_push.device_key" + ) + }); let mut existing_system_config_keys = self .list_system_config_entries() .await? @@ -1374,8 +5481,22 @@ impl<'a> AdminAppState<'a> { ))) } }; + // A portable import must not retain a deployment-local node reference. Rollback is + // different: it runs on the same deployment and should restore the selector when + // that node still exists, so a failed aggregate operation does not silently switch + // the catalog to direct mode. + let selector = if mode.is_rollback_checkpoint() { + match imported_proxy_node_id.as_deref() { + Some(node_id) if self.find_proxy_node(node_id).await?.is_some() => { + Some(node_id) + } + _ => None, + } + } else { + None + }; let request_bytes = Bytes::from( - serde_json::to_vec(&json!({ "proxy_node_id": null })) + serde_json::to_vec(&json!({ "proxy_node_id": selector })) .map_err(|err| GatewayError::Internal(err.to_string()))?, ); match self @@ -1389,11 +5510,18 @@ impl<'a> AdminAppState<'a> { stats.system_configs.created += 1; existing_system_config_keys .insert(ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY.to_string()); + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .system_config_keys + .insert(ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY.to_string()); + } } - if let Some(node_id) = imported_proxy_node_id { - stats.errors.push(format!( + if !mode.is_rollback_checkpoint() { + if let Some(node_id) = imported_proxy_node_id { + stats.errors.push(format!( "外部模型目录代理节点 '{node_id}' 是当前部署的本地引用;代理节点未导入,已切换为直连" )); + } } } Err((status, payload)) => return Ok(Err((status, payload))), @@ -1401,14 +5529,8 @@ impl<'a> AdminAppState<'a> { } let mut global_models_by_name = self - .list_admin_global_models(&AdminGlobalModelListQuery { - offset: 0, - limit: 10_000, - is_active: None, - search: None, - }) + .list_all_admin_global_models_for_system_transfer() .await? - .items .into_iter() .map(|model| (model.name.clone(), model)) .collect::>(); @@ -1446,13 +5568,25 @@ impl<'a> AdminAppState<'a> { model.default_price_per_request, "default_price_per_request", )); + let existing_model = global_models_by_name.get(&name); let default_tiered_pricing = invalid!(normalize_json_object( - model.default_tiered_pricing, + prepare_imported_secret_safe_json( + existing_model.and_then(|model| model.default_tiered_pricing.as_ref()), + model.default_tiered_pricing, + credentials_not_exported, + ), "default_tiered_pricing", )); let supported_capabilities = normalize_supported_capabilities(model.supported_capabilities); - let config = invalid!(normalize_json_object(model.config, "config")); + let config = invalid!(normalize_json_object( + prepare_imported_secret_safe_json( + existing_model.and_then(|model| model.config.as_ref()), + model.config, + credentials_not_exported, + ), + "config", + )); if let Some(existing) = global_models_by_name.get(&name).cloned() { match merge_mode { @@ -1503,8 +5637,12 @@ impl<'a> AdminAppState<'a> { "创建 GlobalModel '{name}' 失败" )))); }; + let created_global_model_id = created.id.clone(); global_models_by_name.insert(name, created); stats.global_models.created += 1; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.global_model_ids.insert(created_global_model_id); + } } let mut providers_by_name = self @@ -1550,11 +5688,21 @@ impl<'a> AdminAppState<'a> { self.build_admin_update_provider_record(&existing, patch) .await ); - updated.proxy = - remap_import_proxy(imported_provider.proxy.clone(), &node_id_map); - updated.config = invalid!(encrypt_imported_provider_config( + updated.proxy = prepare_imported_secret_safe_proxy( + existing.proxy.as_ref(), + imported_provider.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); + let provider_ops_fallback_base_url = + imported_provider_ops_fallback_base_url(&raw_provider); + updated.config = invalid!(prepare_imported_provider_config( self, + &updated.id, + provider_ops_fallback_base_url.as_deref(), + existing.config.as_ref(), imported_provider.config.clone(), + credentials_not_exported, )); let Some(persisted) = self.update_provider_catalog_provider(&updated).await? @@ -1584,10 +5732,21 @@ impl<'a> AdminAppState<'a> { if let Some(enable_format_conversion) = imported_provider.enable_format_conversion { record.enable_format_conversion = enable_format_conversion; } - record.proxy = remap_import_proxy(imported_provider.proxy.clone(), &node_id_map); - record.config = invalid!(encrypt_imported_provider_config( + record.proxy = prepare_imported_secret_safe_proxy( + None, + imported_provider.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); + let provider_ops_fallback_base_url = + imported_provider_ops_fallback_base_url(&raw_provider); + record.config = invalid!(prepare_imported_provider_config( self, + &record.id, + provider_ops_fallback_base_url.as_deref(), + None, imported_provider.config.clone(), + credentials_not_exported, )); let Some(created) = self .create_provider_catalog_provider(&record, shift_existing_priorities_from) @@ -1599,6 +5758,9 @@ impl<'a> AdminAppState<'a> { }; providers_by_name.insert(provider_name.clone(), created.clone()); stats.providers.created += 1; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.provider_ids.insert(created.id.clone()); + } created }; @@ -1697,13 +5859,29 @@ impl<'a> AdminAppState<'a> { admin_provider_endpoints_pure::AdminProviderEndpointUpdateFields { base_url: normalized_base_url, custom_path: payload.custom_path, - header_rules: payload.header_rules, - body_rules: payload.body_rules, + header_rules: prepare_imported_secret_safe_header_rules( + existing_endpoint.header_rules.as_ref(), + payload.header_rules, + credentials_not_exported, + ), + body_rules: prepare_imported_secret_safe_body_rules( + existing_endpoint.body_rules.as_ref(), + payload.body_rules, + credentials_not_exported, + ), max_retries: payload.max_retries, is_active: payload.is_active, - config: payload.config, + config: prepare_imported_secret_safe_json( + existing_endpoint.config.as_ref(), + payload.config, + credentials_not_exported, + ), proxy: payload.proxy, - format_acceptance_config: payload.format_acceptance_config, + format_acceptance_config: prepare_imported_secret_safe_json( + existing_endpoint.format_acceptance_config.as_ref(), + payload.format_acceptance_config, + credentials_not_exported, + ), }; let mut updated = invalid!( admin_provider_endpoints_pure::apply_admin_provider_endpoint_update_fields( @@ -1714,8 +5892,10 @@ impl<'a> AdminAppState<'a> { ) ); if fields.contains("proxy") { - updated.proxy = remap_import_proxy( + updated.proxy = prepare_imported_secret_safe_proxy( + existing_endpoint.proxy.as_ref(), imported_endpoint.proxy.clone(), + credentials_not_exported, &node_id_map, ); } @@ -1763,12 +5943,33 @@ impl<'a> AdminAppState<'a> { endpoint_kind.to_string(), invalid!(normalize_admin_base_url(&imported_endpoint.base_url)), imported_endpoint.custom_path.clone(), - imported_endpoint.header_rules.clone(), - imported_endpoint.body_rules.clone(), + prepare_imported_secret_safe_header_rules( + None, + imported_endpoint.header_rules.clone(), + credentials_not_exported, + ), + prepare_imported_secret_safe_body_rules( + None, + imported_endpoint.body_rules.clone(), + credentials_not_exported, + ), imported_endpoint.max_retries.unwrap_or(2), - imported_endpoint.config.clone(), - remap_import_proxy(imported_endpoint.proxy.clone(), &node_id_map), - imported_endpoint.format_acceptance_config.clone(), + prepare_imported_secret_safe_json( + None, + imported_endpoint.config.clone(), + credentials_not_exported, + ), + prepare_imported_secret_safe_proxy( + None, + imported_endpoint.proxy.clone(), + credentials_not_exported, + &node_id_map, + ), + prepare_imported_secret_safe_json( + None, + imported_endpoint.format_acceptance_config.clone(), + credentials_not_exported, + ), now_unix_secs, ) ); @@ -1779,8 +5980,14 @@ impl<'a> AdminAppState<'a> { "创建 Endpoint '{normalized_api_format}' 失败" )))); }; - existing_endpoints_by_format.insert(normalized_api_format, created); + let created_endpoint_id = created.id.clone(); + existing_endpoints_by_format.insert(normalized_api_format.clone(), created); stats.endpoints.created += 1; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .provider_endpoint_ids + .insert((provider.id.clone(), created_endpoint_id)); + } } let provider_endpoint_formats = existing_endpoints_by_format @@ -1819,50 +6026,17 @@ impl<'a> AdminAppState<'a> { imported_key.auth_config.clone() )); let auth_type = imported_key_auth_type(&imported_key); + let credentials_not_exported = invalid!( + validate_imported_provider_key_credential_state(&imported_key) + ); let normalized_raw_key = normalize_import_key_raw_payload( &raw_key, &auth_type, &normalized_api_formats, normalized_auth_config.clone(), + credentials_not_exported, ); - let existing_key_index = if auth_type == "api_key" { - let target_key = imported_key - .api_key - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned); - existing_keys.iter().position(|existing_key| { - let decrypted_existing = existing_key - .encrypted_api_key - .as_deref() - .and_then(|ciphertext| { - self.decrypt_catalog_secret_with_fallbacks(ciphertext) - }); - target_key - .as_deref() - .zip(decrypted_existing.as_deref()) - .is_some_and(|(target, decrypted)| decrypted == target) - }) - } else if matches!(auth_type.as_str(), "service_account" | "vertex_ai") { - let target_email = - imported_service_account_email(normalized_auth_config.as_ref()); - existing_keys.iter().position(|existing_key| { - target_email.as_deref().is_some_and(|target_email| { - self.parse_catalog_auth_config_json(existing_key) - .and_then(|config| { - config - .get("client_email") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - }) - .as_deref() - == Some(target_email) - }) - }) - } else { + let existing_key_index = if credentials_not_exported { build_import_key_match_name(&imported_key).and_then(|target_name| { existing_keys.iter().position(|existing_key| { existing_key @@ -1872,8 +6046,21 @@ impl<'a> AdminAppState<'a> { && existing_key.name == target_name }) }) + } else { + invalid!(find_imported_provider_key_index( + self, + &imported_key, + &auth_type, + normalized_auth_config.as_ref(), + &existing_keys, + )) }; + if credentials_not_exported && existing_key_index.is_none() { + stats.keys.skipped += 1; + continue; + } + if let Some(existing_index) = existing_key_index { let existing_key = existing_keys[existing_index].clone(); let previous_codex_credential_generation = existing_key @@ -1911,6 +6098,11 @@ impl<'a> AdminAppState<'a> { ) .await ); + if credentials_not_exported { + updated.encrypted_api_key = existing_key.encrypted_api_key.clone(); + updated.encrypted_auth_config = + existing_key.encrypted_auth_config.clone(); + } let oauth_credentials_supplied = if auth_type == "oauth" { invalid!(apply_imported_oauth_key_credentials( self, @@ -1923,10 +6115,18 @@ impl<'a> AdminAppState<'a> { } else { false }; - updated.proxy = - remap_import_proxy(imported_key.proxy.clone(), &node_id_map); + updated.proxy = prepare_imported_secret_safe_proxy( + existing_key.proxy.as_ref(), + imported_key.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); updated.fingerprint = invalid!(normalize_json_object( - imported_key.fingerprint.clone(), + prepare_imported_secret_safe_json( + existing_key.fingerprint.as_ref(), + imported_key.fingerprint.clone(), + credentials_not_exported, + ), "fingerprint", )); let admin_update = build_provider_catalog_key_admin_cas_update( @@ -2037,9 +6237,18 @@ impl<'a> AdminAppState<'a> { imported_key.global_priority_by_format.clone(), "global_priority_by_format", )); - record.proxy = remap_import_proxy(imported_key.proxy.clone(), &node_id_map); + record.proxy = prepare_imported_secret_safe_proxy( + None, + imported_key.proxy.clone(), + credentials_not_exported, + &node_id_map, + ); record.fingerprint = invalid!(normalize_json_object( - imported_key.fingerprint.clone(), + prepare_imported_secret_safe_json( + None, + imported_key.fingerprint.clone(), + credentials_not_exported, + ), "fingerprint", )); let Some(created) = self.create_provider_catalog_key(&record).await? else { @@ -2047,6 +6256,14 @@ impl<'a> AdminAppState<'a> { "创建 Provider '{provider_name}' 的 Key 失败" )))); }; + // Journal the row immediately after creation. The pool-score seed below is a + // separate write and may fail; recording first ensures aggregate compensation + // can still remove this key when that follow-up operation aborts the import. + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .provider_key_ids + .insert((provider.id.clone(), created.id.clone())); + } if oauth_credentials_supplied { seed_imported_oauth_pool_score(self, &provider.id, &created, now_unix_secs) .await?; @@ -2059,12 +6276,7 @@ impl<'a> AdminAppState<'a> { ImportedProviderModel, >(&raw_provider, "models")); let mut existing_models_by_name = self - .list_admin_provider_models(&AdminProviderModelListQuery { - provider_id: provider.id.clone(), - is_active: None, - offset: 0, - limit: 10_000, - }) + .list_all_admin_provider_models_for_system_transfer(&provider.id) .await? .into_iter() .map(|model| (model.provider_model_name.clone(), model)) @@ -2113,8 +6325,10 @@ impl<'a> AdminAppState<'a> { let record = invalid!(build_import_provider_model_record( &provider.id, Some(&existing_model.id), + Some(&existing_model), &global_model_id, &imported_model, + credentials_not_exported, )); let Some(updated) = self.update_admin_provider_model(&record).await? else { @@ -2132,14 +6346,21 @@ impl<'a> AdminAppState<'a> { let record = invalid!(build_import_provider_model_record( &provider.id, None, + None, &global_model_id, &imported_model, + credentials_not_exported, )); let Some(created) = self.create_admin_provider_model(&record).await? else { return Ok(Err(invalid_request(format!( "创建 Provider '{provider_name}' 的模型 '{provider_model_name}' 失败" )))); }; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .provider_model_ids + .insert((provider.id.clone(), created.id.clone())); + } existing_models_by_name.insert(provider_model_name, created); stats.models.created += 1; } @@ -2156,8 +6377,23 @@ impl<'a> AdminAppState<'a> { let existing = self.get_ldap_module_config().await?; let server_url = invalid!(trim_required(&ldap_config.server_url, "LDAP 服务器地址")); + let server_url = invalid!(normalize_ldap_transport_server_url( + &server_url, + ldap_config.use_starttls, + ) + .ok_or_else(|| { + "LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段" + .to_string() + })); let bind_dn = invalid!(trim_required(&ldap_config.bind_dn, "绑定 DN")); let base_dn = invalid!(trim_required(&ldap_config.base_dn, "Base DN")); + if !ldap_distinguished_name_is_valid(&bind_dn) + || !ldap_distinguished_name_is_valid(&base_dn) + { + return Ok(Err(invalid_request( + "LDAP 绑定 DN 或 Base DN 格式无效或过长", + ))); + } let user_search_filter = invalid!(trim_required( ldap_config .user_search_filter @@ -2165,6 +6401,11 @@ impl<'a> AdminAppState<'a> { .unwrap_or("(uid={username})"), "搜索过滤器", )); + if !ldap_search_filter_is_valid(&user_search_filter) { + return Ok(Err(invalid_request( + "LDAP 搜索过滤器格式无效,必须包含 {username} 且使用有限的括号结构", + ))); + } let username_attr = invalid!(trim_required( ldap_config.username_attr.as_deref().unwrap_or("uid"), "用户名属性", @@ -2177,28 +6418,92 @@ impl<'a> AdminAppState<'a> { ldap_config.display_name_attr.as_deref().unwrap_or("cn"), "显示名称属性", )); + if [ + username_attr.as_str(), + email_attr.as_str(), + display_name_attr.as_str(), + ] + .into_iter() + .any(|attribute| !ldap_attribute_description_is_valid(attribute)) + { + return Ok(Err(invalid_request( + "LDAP 用户名、邮箱或显示名称属性格式无效", + ))); + } let connect_timeout = ldap_config.connect_timeout.unwrap_or(10); if !(1..=60).contains(&connect_timeout) { return Ok(Err(invalid_request( "LDAP connect_timeout 必须在 1 到 60 秒之间", ))); } + let config = StoredLdapModuleConfig { + server_url, + bind_dn, + // Password mutation is explicit and separate from the replacement snapshot. + // In particular, Preserve never copies a previously read ciphertext here. + bind_password_encrypted: None, + base_dn, + user_search_filter: Some(user_search_filter), + username_attr: Some(username_attr), + email_attr: Some(email_attr), + display_name_attr: Some(display_name_attr), + is_enabled: ldap_config.is_enabled, + is_exclusive: ldap_config.is_exclusive, + use_starttls: ldap_config.use_starttls, + connect_timeout: Some(connect_timeout), + }; let bind_password = ldap_config .bind_password .as_deref() .map(str::trim) .map(ToOwned::to_owned); - let will_have_password = bind_password + if bind_password .as_deref() - .map(|value| !value.is_empty()) - .unwrap_or_else(|| { - existing - .as_ref() - .and_then(|config| config.bind_password_encrypted.as_deref()) - .map(str::trim) - .is_some_and(|value| !value.is_empty()) - }); - if existing.is_none() && !will_have_password { + .is_some_and(is_imported_redacted_secret) + { + return Ok(Err(invalid_request("LDAP 脱敏占位符不能作为绑定密码导入"))); + } + let bind_password_update = match bind_password { + Some(password) if password.is_empty() => LdapBindPasswordUpdate::Clear, + Some(password) => LdapBindPasswordUpdate::Set(routed!(self + .encrypt_ldap_bind_password(&config, &password) + .ok_or_else(|| { + invalid_request("LDAP 绑定密码加密失败,请检查 Rust 数据加密配置") + }))), + None => LdapBindPasswordUpdate::Preserve, + }; + if matches!(&bind_password_update, LdapBindPasswordUpdate::Preserve) { + if let Some(existing) = existing.as_ref() { + if existing + .bind_password_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + let binding_matches = invalid!( + crate::handlers::shared::ldap_bind_password_binding_matches( + existing, &config, + ) + ); + if !binding_matches { + return Ok(Err(invalid_request( + "导入 LDAP 时修改了服务器、StartTLS、bind DN 或 Base DN,必须提供绑定密码", + ))); + } + } + } + } + let will_have_password = match &bind_password_update { + LdapBindPasswordUpdate::Set(ciphertext) => !ciphertext.trim().is_empty(), + LdapBindPasswordUpdate::Clear => false, + LdapBindPasswordUpdate::Preserve => existing + .as_ref() + .and_then(|config| config.bind_password_encrypted.as_deref()) + .map(str::trim) + .is_some_and(|value| !value.is_empty()), + }; + if existing.is_none() + && !matches!(&bind_password_update, LdapBindPasswordUpdate::Set(_)) + { return Ok(Err(invalid_request("首次配置 LDAP 时必须设置绑定密码"))); } if ldap_config.is_exclusive && !ldap_config.is_enabled { @@ -2219,47 +6524,56 @@ impl<'a> AdminAppState<'a> { ))); } } - let bind_password_encrypted = match bind_password { - Some(password) if password.is_empty() => None, - Some(password) => Some(routed!(self - .encrypt_catalog_secret_with_fallbacks(&password) - .ok_or_else(|| { - invalid_request("LDAP 绑定密码加密失败,请检查 Rust 数据加密配置") - }))), - None => existing - .as_ref() - .and_then(|config| config.bind_password_encrypted.clone()), - }; - let config = StoredLdapModuleConfig { - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter: Some(user_search_filter), - username_attr: Some(username_attr), - email_attr: Some(email_attr), - display_name_attr: Some(display_name_attr), - is_enabled: ldap_config.is_enabled, - is_exclusive: ldap_config.is_exclusive, - use_starttls: ldap_config.use_starttls, - connect_timeout: Some(connect_timeout), - }; - match (existing.is_some(), merge_mode) { (true, AdminImportMergeMode::Skip) => stats.ldap.skipped += 1, (true, AdminImportMergeMode::Error) => { return Ok(Err(invalid_request("LDAP 配置已存在"))); } (true, AdminImportMergeMode::Overwrite) => { - let Some(_) = self.upsert_ldap_module_config(&config).await? else { + let Some(result) = self + .compare_and_swap_ldap_module_config( + existing.as_ref(), + &config, + &bind_password_update, + ) + .await? + else { return Ok(Err(invalid_request("更新 LDAP 配置失败"))); }; + if result == CompareAndSwapLdapConfigResult::Conflict { + return Ok(Err(( + http::StatusCode::CONFLICT, + json!({ + "detail": "LDAP 配置已被其他请求更新,请重试" + }), + ))); + } stats.ldap.updated += 1; } (false, _) => { - let Some(_) = self.upsert_ldap_module_config(&config).await? else { + let Some(result) = self + .compare_and_swap_ldap_module_config( + None, + &config, + &bind_password_update, + ) + .await? + else { return Ok(Err(invalid_request("创建 LDAP 配置失败"))); }; + let CompareAndSwapLdapConfigResult::Applied(created) = result else { + return Ok(Err(( + http::StatusCode::CONFLICT, + json!({ + "detail": "LDAP 配置已被其他请求创建,请重试" + }), + ))); + }; + // Record the exact persisted snapshot before any later phase can fail. + // Compensation will delete it only if no concurrent write changed it. + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.created_ldap_config = Some(created); + } stats.ldap.created += 1; } } @@ -2277,6 +6591,12 @@ impl<'a> AdminAppState<'a> { for (index, imported_oauth_item) in imported_oauth_providers.into_iter().enumerate() { let (_, oauth_provider) = imported_oauth_item.into_parts(); + let original_provider_type = oauth_provider.provider_type.clone(); + let original_enabled = oauth_provider.is_enabled; + let oauth_provider = invalid!(normalize_legacy_imported_oauth_provider( + oauth_provider, + &source_version, + )); let provider_type = invalid!(trim_required( &oauth_provider.provider_type, "provider_type", @@ -2306,49 +6626,61 @@ impl<'a> AdminAppState<'a> { &oauth_provider.frontend_callback_url, "frontend_callback_url", )); - let client_secret_encrypted = + let mut record = invalid!(build_imported_oauth_provider_record( + &oauth_provider, + EncryptedSecretUpdate::Preserve, + )); + record.provider_type = provider_type.clone(); + record.display_name = display_name; + record.client_id = client_id; + record.redirect_uri = redirect_uri; + record.frontend_callback_url = frontend_callback_url; + + // Bind imported plaintext to the final normalized record, including all + // endpoint and redirect fields. Never seal using provider_type alone. + record.client_secret_encrypted = match oauth_provider.client_secret.as_deref().map(str::trim) { - Some(secret) if !secret.is_empty() => { - EncryptedSecretUpdate::Set(routed!(self - .encrypt_catalog_secret_with_fallbacks(secret) - .ok_or_else(|| { - invalid_request("gateway 未配置 OAuth provider 加密密钥") - }))) + Some(secret) if is_imported_redacted_secret(secret) => { + EncryptedSecretUpdate::Preserve } + Some(secret) if !secret.is_empty() => EncryptedSecretUpdate::Set(routed!( + crate::handlers::shared::seal_identity_oauth_provider_client_secret( + self.as_ref(), + &record, + secret, + ) + .map_err(|message| invalid_request(message)) + )), _ => EncryptedSecretUpdate::Preserve, }; - let record = UpsertOAuthProviderConfigRecord { - provider_type: provider_type.clone(), - display_name, - client_id, - client_secret_encrypted, - authorization_url_override: oauth_provider - .authorization_url_override - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()), - token_url_override: oauth_provider - .token_url_override - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()), - userinfo_url_override: oauth_provider - .userinfo_url_override - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()), - scopes: normalize_string_list(oauth_provider.scopes), - redirect_uri, - frontend_callback_url, - attribute_mapping: invalid!(normalize_json_object( - oauth_provider.attribute_mapping, - "attribute_mapping", - )), - extra_config: invalid!(normalize_json_object( - oauth_provider.extra_config, - "extra_config", - )), - icon_url: None, - is_enabled: oauth_provider.is_enabled, - }; - invalid!(record.validate().map_err(|err| err.to_string())); + + // A redacted/omitted secret may preserve an existing value only when the + // complete OAuth binding is unchanged. Otherwise the old secret would be + // replayed against a different client or endpoint after an overwrite import. + if matches!( + &record.client_secret_encrypted, + EncryptedSecretUpdate::Preserve + ) { + if let Some(existing) = oauth_by_type.get(&provider_type) { + if existing + .client_secret_encrypted + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + let binding_matches = invalid!( + crate::handlers::shared::identity_oauth_provider_secret_binding_matches( + existing, + &record, + ) + ); + if !binding_matches { + return Ok(Err(invalid_request( + "导入 OAuth Provider 时修改了 Client ID、端点或 redirect_uri,必须提供 client_secret", + ))); + } + } + } + } let Some(persisted) = self.upsert_oauth_provider_config(&record).await? else { stats.oauth.skipped += (imported_oauth_provider_count - index) as u64; @@ -2358,11 +6690,19 @@ impl<'a> AdminAppState<'a> { ); break; }; - oauth_by_type.insert(provider_type, persisted); + oauth_by_type.insert(provider_type.clone(), persisted); + if original_enabled && !oauth_provider.is_enabled { + stats.errors.push(format!( + "旧版 OAuth Provider '{original_provider_type}' 已安全迁移并停用,请复核域名白名单后重新启用" + )); + } if existed { stats.oauth.updated += 1; } else { stats.oauth.created += 1; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.oauth_provider_types.insert(provider_type.clone()); + } } } } @@ -2375,6 +6715,13 @@ impl<'a> AdminAppState<'a> { description, } = system_config; let normalized_key = normalize_imported_system_config_key(&key); + if credentials_not_exported + && (is_sensitive_admin_system_config_key(&normalized_key) + || is_interactive_export_private_system_config_key(&normalized_key)) + { + stats.system_configs.skipped += 1; + continue; + } let exists = existing_system_config_keys.contains(&normalized_key); match (exists, merge_mode) { (true, AdminImportMergeMode::Skip) => { @@ -2404,7 +6751,10 @@ impl<'a> AdminAppState<'a> { stats.system_configs.updated += 1; } else { stats.system_configs.created += 1; - existing_system_config_keys.insert(normalized_key); + existing_system_config_keys.insert(normalized_key.clone()); + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.system_config_keys.insert(normalized_key); + } } } Err((status, payload)) => return Ok(Err((status, payload))), @@ -2421,6 +6771,110 @@ impl<'a> AdminAppState<'a> { &self, request_body: &Bytes, operator_id: Option<&str>, + ) -> Result, GatewayError> { + let mut mutation_journal = AggregateMutationJournal::default(); + let result = self + .import_admin_system_users_with_mode( + request_body, + operator_id, + SystemImportMode::InteractiveUpload, + Some(&mut mutation_journal), + ) + .await; + self.finish_standalone_users_import(result, &mutation_journal) + .await + } + + pub(crate) async fn restore_admin_system_users_backup( + &self, + request_body: &Bytes, + operator_id: Option<&str>, + _authority: crate::backup::executor::BackupRestoreAuthority, + ) -> Result, GatewayError> { + let mut mutation_journal = AggregateMutationJournal::default(); + let result = self + .import_admin_system_users_with_mode( + request_body, + operator_id, + SystemImportMode::RecoveryBackup, + Some(&mut mutation_journal), + ) + .await; + self.finish_standalone_users_import(result, &mutation_journal) + .await + } + + /// Interactive and standalone recovery imports do not have the aggregate config checkpoint + /// available to their caller. Compensate rows created by this invocation and restore existing + /// rows only when their mutable fields still match the recorded post-state. + async fn finish_standalone_users_import( + &self, + result: Result, GatewayError>, + mutation_journal: &AggregateMutationJournal, + ) -> Result, GatewayError> { + match result { + Ok(Ok(payload)) => Ok(Ok(payload)), + Ok(Err(original)) => match self + .rollback_standalone_users_mutations(mutation_journal) + .await + { + Ok(()) => Ok(Err(original)), + Err(rollback_error) => Err(aggregate_rollback_http_error( + "用户阶段", + &original, + rollback_error, + )), + }, + Err(original) => match self + .rollback_standalone_users_mutations(mutation_journal) + .await + { + Ok(()) => Err(original), + Err(rollback_error) => Err(aggregate_rollback_error( + "用户阶段", + original, + rollback_error, + )), + }, + } + } + + async fn rollback_standalone_users_mutations( + &self, + mutation_journal: &AggregateMutationJournal, + ) -> Result<(), GatewayError> { + // Run both compensations even when one fails. Newly-created rows are removed first, while + // pre-existing wallets are restored through an owner- and snapshot-checked CAS so a + // concurrent recharge cannot be overwritten by an import failure. + let created_result = self.rollback_created_users(mutation_journal).await; + let existing_wallet_result = self.rollback_existing_wallets(mutation_journal).await; + let existing_users_result = self.rollback_existing_users(mutation_journal).await; + let existing_groups_result = self.rollback_existing_user_groups(mutation_journal).await; + let existing_api_keys_result = self.rollback_existing_api_keys(mutation_journal).await; + let existing_result = combine_rollback_results( + existing_users_result, + existing_groups_result, + "standalone existing users/groups", + ); + let existing_result = combine_rollback_results( + existing_api_keys_result, + existing_result, + "standalone existing API keys", + ); + let restore_result = combine_rollback_results( + created_result, + existing_wallet_result, + "standalone users wallets", + ); + combine_rollback_results(restore_result, existing_result, "standalone users") + } + + async fn import_admin_system_users_with_mode( + &self, + request_body: &Bytes, + operator_id: Option<&str>, + mode: SystemImportMode, + mut mutation_journal: Option<&mut AggregateMutationJournal>, ) -> Result, GatewayError> { if !self.has_auth_user_write_capability() || !self.has_auth_wallet_write_capability() @@ -2431,6 +6885,13 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + match self + .prevalidate_admin_system_users_import(request_body, operator_id, mode) + .await? + { + Ok(()) => {} + Err(err) => return Ok(Err(err)), + } let root = match serde_json::from_slice::(request_body) { Ok(Value::Object(map)) => map, _ => return Ok(Err(invalid_request("请求数据验证失败"))), @@ -2461,10 +6922,9 @@ impl<'a> AdminAppState<'a> { Some(_) => return Ok(Err(invalid_request("user_groups 必须是数组"))), None => &empty, }; - let standalone_owner_id = match operator_id { Some(candidate) => match self.find_user_auth_by_id(candidate).await? { - Some(user) if user.role.eq_ignore_ascii_case("admin") => Some(user.id), + Some(user) if crate::roles::is_full_admin_role(&user.role) => Some(user.id), _ => None, }, None => None, @@ -2479,13 +6939,38 @@ impl<'a> AdminAppState<'a> { }; } - invalid_value!(validate_imported_system_users_export_version( - root.get("version") - )); + // Repository write helpers return `None` when the corresponding writer is unavailable + // or the target row disappeared. Treat that as a failed import step so the caller's + // mutation journal can compensate newly-created rows instead of returning a half-imported + // success. + macro_rules! require_persisted { + ($expr:expr) => { + match $expr { + Some(value) => value, + None => { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "Admin system data unavailable" }), + ))) + } + } + }; + } - let supplemental_user_usage_aggregates = invalid_value!( - build_imported_user_usage_total_aggregates(users, root.get("exported_at")) + let users_export_version = invalid_value!( + validate_imported_system_users_export_version_for_mode(root.get("version"), mode) ); + + let supplemental_user_usage_aggregates = if mode.is_rollback_checkpoint() { + // Rollback checkpoints must not mutate runtime usage state. In particular, do not + // synthesize daily rows from the denormalized counters on the user records. + Vec::new() + } else { + invalid_value!(build_imported_user_usage_total_aggregates( + users, + root.get("exported_at") + )) + }; let mut stats = AdminSystemUsersImportStats::default(); let mut imported_user_id_map = BTreeMap::::new(); let mut imported_api_key_id_map = BTreeMap::::new(); @@ -2526,7 +7011,26 @@ impl<'a> AdminAppState<'a> { stats.user_groups.skipped += 1; continue; } - if let Some(existing) = groups_by_name.get(&normalized_name).cloned() { + let existing_by_id = mode + .is_rollback_checkpoint() + .then(|| { + export_id.as_deref().and_then(|export_id| { + groups_by_name + .values() + .find(|group| group.id == export_id) + .cloned() + }) + }) + .flatten(); + if mode.is_rollback_checkpoint() && export_id.is_some() && existing_by_id.is_none() { + return Ok(Err(invalid_request(format!( + "回滚检查点用户组 '{}' 不存在;拒绝按名称匹配", + export_id.as_deref().unwrap_or_default() + )))); + } + if let Some(existing) = + existing_by_id.or_else(|| groups_by_name.get(&normalized_name).cloned()) + { if let Some(export_id) = export_id { imported_group_id_map.insert(export_id, existing.id.clone()); } @@ -2542,6 +7046,15 @@ impl<'a> AdminAppState<'a> { )))); } AdminImportMergeMode::Overwrite => { + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .existing_user_groups + .entry(existing.id.clone()) + .or_insert_with(|| ExistingUserGroupMutation { + before: existing.clone(), + after: existing.clone(), + }); + } let Some(updated) = self.update_user_group(&existing.id, record).await? else { return Ok(Err(( @@ -2549,6 +7062,14 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); }; + if let Some(journal) = mutation_journal.as_deref_mut() { + if let Some(mutation) = + journal.existing_user_groups.get_mut(&existing.id) + { + mutation.after = updated.clone(); + } + } + groups_by_name.retain(|_, group| group.id != existing.id); groups_by_name.insert(normalized_name, updated); stats.user_groups.updated += 1; } @@ -2562,6 +7083,9 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); }; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.user_group_ids.insert(created.id.clone()); + } if let Some(export_id) = export_id { imported_group_id_map.insert(export_id, created.id.clone()); } @@ -2576,22 +7100,25 @@ impl<'a> AdminAppState<'a> { Err(detail) => return Ok(Err(invalid_request(detail))), }; let source_user_id = invalid_value!(imported_optional_string(user.get("id"))); - let role = invalid_value!(imported_optional_string(user.get("role"))) - .unwrap_or_else(|| "user".to_string()) - .to_ascii_lowercase(); - if role == "admin" { + invalid_value!(validate_rollback_user_source_id( + mode, + source_user_id.as_deref(), + )); + let Some(role) = + invalid_value!(normalize_imported_system_user_role(user.get("role"), mode,)) + else { let skipped_email = invalid_value!(imported_optional_string(user.get("email"))); let skipped_username = invalid_value!(imported_optional_string(user.get("username"))); stats.users.skipped += 1; stats.errors.push(format!( - "跳过管理员用户: {}", + "跳过受保护的管理员用户: {}", skipped_email .or(skipped_username) .unwrap_or_else(|| format!("users[{index}]")) )); continue; - } + }; let email = invalid_value!(imported_optional_string(user.get("email"))) .map(|value| value.to_ascii_lowercase()); @@ -2608,7 +7135,11 @@ impl<'a> AdminAppState<'a> { }) }) .unwrap_or_else(|| format!("imported-user-{index}")); - let password_hash = invalid_value!(imported_optional_string(user.get("password_hash"))); + let password_hash = invalid_value!(resolve_imported_password_hash( + user, + users_export_version, + mode, + )); let allowed_providers = invalid_value!(normalize_imported_user_string_list( user, "allowed_providers" @@ -2684,23 +7215,43 @@ impl<'a> AdminAppState<'a> { Some(Value::Null) | None => None, Some(_) => return Ok(Err(invalid_request("wallet 必须是对象"))), }; - let wallet_target = - invalid_value!(normalize_imported_wallet_target(wallet_payload, false)); + let wallet_target = match wallet_payload { + Some(wallet) => Some(invalid_value!(normalize_imported_wallet_target( + Some(wallet), + false, + ))), + None => None, + }; - let mut existing_user = if let Some(email) = email.as_deref() { - self.find_user_auth_by_identifier(email).await? + // Checkpoint IDs are the only safe identity during compensation. Email and username + // are mutable fields and may have been changed by the failed import itself. Refuse to + // guess if the stable row disappeared instead of overwriting an unrelated account. + let existing_user = if mode.is_rollback_checkpoint() { + let source_user_id = source_user_id.as_deref().unwrap_or_default(); + let existing = self.find_user_auth_by_id(source_user_id).await?; + if existing.is_none() { + return Ok(Err(invalid_request(format!( + "回滚检查点用户 '{source_user_id}' 不存在;拒绝按 email/username 匹配" + )))); + } + existing } else { - None + let mut existing = if let Some(email) = email.as_deref() { + self.find_user_auth_by_identifier(email).await? + } else { + None + }; + if existing.is_none() { + existing = self.find_user_auth_by_identifier(&username).await?; + } + existing }; - if existing_user.is_none() { - existing_user = self.find_user_auth_by_identifier(&username).await?; - } let user_id = if let Some(existing) = existing_user { - if existing.role.eq_ignore_ascii_case("admin") { + if imported_existing_user_is_protected(&existing.role, mode) { stats.users.skipped += 1; stats.errors.push(format!( - "跳过管理员用户记录: {}", + "跳过受保护的管理员用户记录: {}", email.clone().unwrap_or(username.clone()) )); continue; @@ -2717,6 +7268,10 @@ impl<'a> AdminAppState<'a> { )))); } AdminImportMergeMode::Overwrite => { + if let Some(journal) = mutation_journal.as_deref_mut() { + self.capture_existing_user_mutation(journal, &existing) + .await?; + } if let Some(email) = email.as_deref() { if self .is_other_user_auth_email_taken(email, &existing.id) @@ -2731,10 +7286,27 @@ impl<'a> AdminAppState<'a> { { return Ok(Err(invalid_request(format!("用户名已存在: {username}")))); } + let email_present = email.is_some() || mode.is_rollback_checkpoint(); + let email_verified_update = if mode.is_rollback_checkpoint() { + user.contains_key("email_verified") + .then_some(email_verified) + } else { + email.as_deref().and_then(|email| { + existing + .email + .as_deref() + .is_none_or(|current| { + !current.trim().eq_ignore_ascii_case(email.trim()) + }) + .then_some(email_verified) + }) + }; let updated_profile = self .update_local_auth_user_profile( &existing.id, + email_present, email.clone(), + email_verified_update, Some(username.clone()), ) .await?; @@ -2744,22 +7316,32 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; if let Some(password_hash) = password_hash.as_deref().filter(|value| !value.is_empty()) { let updated_password = self - .update_local_auth_user_password_hash( + .reset_local_auth_user_password_and_revoke_sessions( &existing.id, password_hash.to_string(), chrono::Utc::now(), ) .await?; - if updated_password.is_none() { + if !updated_password { return Ok(Err(( http::StatusCode::SERVICE_UNAVAILABLE, json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; } let updated_admin_fields = self .update_local_auth_user_admin_fields( @@ -2782,6 +7364,11 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; if user.contains_key("email_verified") { stats.errors.push(format!( "用户 '{}' 的 email_verified 当前不会覆盖已有值", @@ -2789,20 +7376,42 @@ impl<'a> AdminAppState<'a> { )); } if user.contains_key("model_capability_settings") { - let _ = self + let updated = self .update_user_model_capability_settings( &existing.id, model_capability_settings.clone(), ) .await?; + if model_capability_settings.is_some() && updated.is_none() { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "Admin system data unavailable" }), + ))); + } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; } if user.contains_key("feature_settings") { - let _ = self + let updated = self .update_user_feature_settings( &existing.id, feature_settings.clone(), ) .await?; + if feature_settings.is_some() && updated.is_none() { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "Admin system data unavailable" }), + ))); + } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; } if allowed_providers_mode.is_some() || allowed_api_formats_mode.is_some() @@ -2824,17 +7433,37 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; } if let Some(group_ids) = group_ids.as_ref() { - self.replace_user_groups_for_user(&existing.id, group_ids) + let persisted_groups = self + .replace_user_groups_for_user(&existing.id, group_ids) .await?; + if persisted_groups.len() != group_ids.len() { + return Ok(Err(invalid_request(format!( + "用户 '{}' 的用户组未能完整写入", + email.clone().unwrap_or(username.clone()) + )))); + } + self.refresh_existing_user_mutation( + mutation_journal.as_deref_mut(), + &existing.id, + ) + .await?; + } + if let Some(wallet_target) = wallet_target.as_ref() { + self.sync_imported_user_wallet( + &existing.id, + wallet_target, + &email.clone().unwrap_or(username.clone()), + mutation_journal.as_deref_mut(), + ) + .await?; } - self.sync_imported_user_wallet( - &existing.id, - &wallet_target, - &email.clone().unwrap_or(username.clone()), - ) - .await?; stats.users.updated += 1; existing.id } @@ -2845,7 +7474,7 @@ impl<'a> AdminAppState<'a> { email.clone(), email_verified, username.clone(), - password_hash.unwrap_or_default(), + password_hash.unwrap_or_else(imported_password_tombstone), role.clone(), allowed_providers.clone(), allowed_api_formats.clone(), @@ -2859,18 +7488,33 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); }; + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.user_ids.insert(created.id.clone()); + } if user.contains_key("model_capability_settings") { - let _ = self + let updated = self .update_user_model_capability_settings( &created.id, model_capability_settings.clone(), ) .await?; + if model_capability_settings.is_some() && updated.is_none() { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "Admin system data unavailable" }), + ))); + } } if user.contains_key("feature_settings") { - let _ = self + let updated = self .update_user_feature_settings(&created.id, feature_settings.clone()) .await?; + if feature_settings.is_some() && updated.is_none() { + return Ok(Err(( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": "Admin system data unavailable" }), + ))); + } } let created = if allowed_providers_mode.is_some() || allowed_api_formats_mode.is_some() @@ -2897,15 +7541,25 @@ impl<'a> AdminAppState<'a> { created }; if let Some(group_ids) = group_ids.as_ref() { - self.replace_user_groups_for_user(&created.id, group_ids) + let persisted_groups = self + .replace_user_groups_for_user(&created.id, group_ids) .await?; + if persisted_groups.len() != group_ids.len() { + return Ok(Err(invalid_request(format!( + "用户 '{}' 的用户组未能完整写入", + email.clone().unwrap_or(username.clone()) + )))); + } + } + if let Some(wallet_target) = wallet_target.as_ref() { + self.sync_imported_user_wallet( + &created.id, + wallet_target, + &email.clone().unwrap_or(username.clone()), + mutation_journal.as_deref_mut(), + ) + .await?; } - self.sync_imported_user_wallet( - &created.id, - &wallet_target, - &email.clone().unwrap_or(username.clone()), - ) - .await?; stats.users.created += 1; created.id }; @@ -2924,10 +7578,16 @@ impl<'a> AdminAppState<'a> { Some(_) => return Ok(Err(invalid_request("api_keys 必须是数组"))), None => &empty, }; - let mut existing_api_keys_by_hash = existing_api_keys - .into_iter() - .map(|record| (record.key_hash.clone(), record)) - .collect::>(); + let mut existing_api_keys_by_hash = BTreeMap::new(); + for record in existing_api_keys { + let api_key_id = record.api_key_id.clone(); + existing_api_keys_by_hash.insert(record.key_hash.clone(), record.clone()); + if mode.is_rollback_checkpoint() { + existing_api_keys_by_hash + .entry(imported_api_key_tombstone(&api_key_id)) + .or_insert(record); + } + } for (key_index, raw_key) in imported_api_keys.iter().enumerate() { let key = match imported_object_field( @@ -2937,8 +7597,12 @@ impl<'a> AdminAppState<'a> { Ok(value) => value, Err(detail) => return Ok(Err(invalid_request(detail))), }; - let Some((key_hash, key_encrypted)) = - invalid_value!(self.resolve_imported_system_user_api_key_material(key)) + let Some(key_material) = invalid_value!(self + .resolve_imported_system_user_api_key_material( + key, + users_export_version, + mode, + )) else { stats.api_keys.skipped += 1; stats.errors.push(format!( @@ -2947,6 +7611,8 @@ impl<'a> AdminAppState<'a> { )); continue; }; + let key_hash = key_material.key_hash; + let key_plaintext = key_material.key_plaintext; let source_api_key_id = invalid_value!(imported_optional_string(key.get("api_key_id"))); let name = invalid_value!(imported_optional_string(key.get("name"))); @@ -2961,9 +7627,16 @@ impl<'a> AdminAppState<'a> { let allowed_models = invalid_value!(normalize_imported_user_string_list(key, "allowed_models")); let ip_rules = invalid_value!(normalize_imported_user_ip_rules(key)); - let rate_limit = - invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit")) - .unwrap_or(0); + let imported_rate_limit = + invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit")); + // Legacy uploads historically normalize an omitted rate limit to zero. Rollback + // checkpoints instead preserve the nullable database value exactly. + let rate_limit = imported_rate_limit.unwrap_or(0); + let rate_limit_value = if mode.is_rollback_checkpoint() { + imported_rate_limit + } else { + Some(rate_limit) + }; let concurrent_limit = invalid_value!(imported_optional_i32( key.get("concurrent_limit"), "concurrent_limit" @@ -2973,7 +7646,7 @@ impl<'a> AdminAppState<'a> { } let force_capabilities = imported_optional_value(key.get("force_capabilities")); let is_active = - invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(true); + invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(false); let expires_at_unix_secs = invalid_value!(imported_rfc3339_to_unix_secs( key.get("expires_at"), "expires_at" @@ -3014,20 +7687,52 @@ impl<'a> AdminAppState<'a> { )))); } AdminImportMergeMode::Overwrite => { + let key_encrypted = invalid_value!(self + .seal_imported_auth_api_key_secret( + key_plaintext.as_deref(), + &user_id, + &existing_key.api_key_id, + &key_hash, + false, + )); + if let Some(journal) = mutation_journal.as_deref_mut() { + let key = (user_id.clone(), existing_key.api_key_id.clone()); + journal + .existing_user_api_keys + .entry(key) + .or_insert_with(|| ExistingApiKeyMutation { + before: existing_key.clone(), + after: existing_key.clone(), + }); + } let updated = self .update_user_api_key_basic( aether_data::repository::auth::UpdateUserApiKeyBasicRecord { user_id: user_id.clone(), api_key_id: existing_key.api_key_id.clone(), + key_encrypted: key_encrypted.clone(), + key_encrypted_present: key_encrypted.is_some() + || mode == SystemImportMode::RecoveryRollbackCheckpoint, name: name.clone(), - rate_limit: Some(rate_limit), - concurrent_limit: if key.contains_key("concurrent_limit") { + name_present: name.is_some() + || mode.is_rollback_checkpoint(), + rate_limit: rate_limit_value, + rate_limit_present: true, + concurrent_limit: if key.contains_key("concurrent_limit") + || mode.is_rollback_checkpoint() + { concurrent_limit } else { None }, + concurrent_limit_present: key + .contains_key("concurrent_limit") + || mode.is_rollback_checkpoint(), ip_rules: imported_ip_rules_present(key) .then(|| ip_rules.clone()), + feature_settings: key + .contains_key("feature_settings") + .then(|| feature_settings.clone()), }, ) .await?; @@ -3037,36 +7742,58 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } - let _ = self - .set_user_api_key_allowed_providers( + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + Some(&user_id), + &existing_key.api_key_id, + false, + ) + .await?; + let _ = require_persisted!( + self.set_user_api_key_allowed_providers( &user_id, &existing_key.api_key_id, allowed_providers.clone(), ) - .await?; - let _ = self - .set_user_api_key_force_capabilities( + .await? + ); + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + Some(&user_id), + &existing_key.api_key_id, + false, + ) + .await?; + let _ = require_persisted!( + self.set_user_api_key_force_capabilities( &user_id, &existing_key.api_key_id, force_capabilities.clone(), ) - .await?; - if key.contains_key("feature_settings") { - let _ = self - .set_user_api_key_feature_settings( - &user_id, - &existing_key.api_key_id, - feature_settings.clone(), - ) - .await?; - } - let _ = self - .set_user_api_key_active( + .await? + ); + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + Some(&user_id), + &existing_key.api_key_id, + false, + ) + .await?; + let _ = require_persisted!( + self.set_user_api_key_active( &user_id, &existing_key.api_key_id, - is_active, + mode.preserves_active_state() && is_active, ) - .await?; + .await? + ); + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + Some(&user_id), + &existing_key.api_key_id, + false, + ) + .await?; if imported_total_requests.is_some() || imported_total_tokens.is_some() || imported_total_cost_usd.is_some() @@ -3087,6 +7814,13 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + Some(&user_id), + &existing_key.api_key_id, + false, + ) + .await?; } if key.contains_key("allowed_api_formats") || key.contains_key("allowed_models") @@ -3108,10 +7842,20 @@ impl<'a> AdminAppState<'a> { continue; } + let api_key_id = imported_api_key_id_for_mode(source_api_key_id.as_deref(), mode); + let key_encrypted = invalid_value!(self.seal_imported_auth_api_key_secret( + key_plaintext.as_deref(), + &user_id, + &api_key_id, + &key_hash, + false, + )); let created = self .create_user_api_key(aether_data::repository::auth::CreateUserApiKeyRecord { user_id: user_id.clone(), - api_key_id: Uuid::new_v4().to_string(), + // Preserve a checkpoint API-key ID when a missing row must be recreated; + // ordinary imports continue to receive fresh IDs. + api_key_id, key_hash: key_hash.clone(), key_encrypted, name, @@ -3122,7 +7866,8 @@ impl<'a> AdminAppState<'a> { rate_limit, concurrent_limit, force_capabilities, - is_active, + feature_settings: feature_settings.clone(), + is_active: mode.preserves_active_state() && is_active, expires_at_unix_secs, auto_delete_on_expiry, total_requests, @@ -3136,16 +7881,12 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); }; - if key.contains_key("feature_settings") { - let _ = self - .set_user_api_key_feature_settings( - &user_id, - &created.api_key_id, - feature_settings.clone(), - ) - .await?; - } let created_api_key_id = created.api_key_id.clone(); + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .user_api_key_ids + .insert((user_id.clone(), created_api_key_id.clone())); + } existing_api_keys_by_hash.insert(key_hash, created); if let Some(source_api_key_id) = source_api_key_id { imported_api_key_id_map.insert(source_api_key_id, created_api_key_id); @@ -3160,17 +7901,19 @@ impl<'a> AdminAppState<'a> { stats .errors .push("无法导入独立余额 Key: 当前管理员用户记录不存在".to_string()); - if let Some(summary) = self - .import_admin_system_user_usage_aggregates( - root.get("usage_aggregates"), - &supplemental_user_usage_aggregates, - &imported_user_id_map, - &imported_api_key_id_map, - merge_mode, - ) - .await? - { - stats.usage_aggregates = Some(summary); + if !mode.is_rollback_checkpoint() { + if let Some(summary) = self + .import_admin_system_user_usage_aggregates( + root.get("usage_aggregates"), + &supplemental_user_usage_aggregates, + &imported_user_id_map, + &imported_api_key_id_map, + merge_mode, + ) + .await? + { + stats.usage_aggregates = Some(summary); + } } return Ok(Ok(json!({ "message": "用户数据导入成功", @@ -3183,10 +7926,16 @@ impl<'a> AdminAppState<'a> { .await? .into_iter() .collect::>(); - let mut existing_standalone_by_hash = existing_standalone_keys - .into_iter() - .map(|record| (record.key_hash.clone(), record)) - .collect::>(); + let mut existing_standalone_by_hash = BTreeMap::new(); + for record in existing_standalone_keys { + let api_key_id = record.api_key_id.clone(); + existing_standalone_by_hash.insert(record.key_hash.clone(), record.clone()); + if mode.is_rollback_checkpoint() { + existing_standalone_by_hash + .entry(imported_api_key_tombstone(&api_key_id)) + .or_insert(record); + } + } for (index, raw_key) in standalone_keys.iter().enumerate() { let key = match imported_object_field(raw_key, &format!("standalone_keys[{index}]")) @@ -3194,8 +7943,12 @@ impl<'a> AdminAppState<'a> { Ok(value) => value, Err(detail) => return Ok(Err(invalid_request(detail))), }; - let Some((key_hash, key_encrypted)) = - invalid_value!(self.resolve_imported_system_user_api_key_material(key)) + let Some(key_material) = invalid_value!(self + .resolve_imported_system_user_api_key_material( + key, + users_export_version, + mode, + )) else { stats.standalone_keys.skipped += 1; stats @@ -3203,6 +7956,8 @@ impl<'a> AdminAppState<'a> { .push(format!("跳过无效独立余额 Key: standalone_keys[{index}]")); continue; }; + let key_hash = key_material.key_hash; + let key_plaintext = key_material.key_plaintext; let source_api_key_id = invalid_value!(imported_optional_string(key.get("api_key_id"))); let name = invalid_value!(imported_optional_string(key.get("name"))); @@ -3229,7 +7984,7 @@ impl<'a> AdminAppState<'a> { } let force_capabilities = imported_optional_value(key.get("force_capabilities")); let is_active = - invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(true); + invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(false); let expires_at_unix_secs = invalid_value!(imported_rfc3339_to_unix_secs( key.get("expires_at"), "expires_at" @@ -3264,8 +8019,13 @@ impl<'a> AdminAppState<'a> { }; let unlimited = invalid_value!(imported_optional_bool(key.get("unlimited"))).unwrap_or(false); - let wallet_target = - invalid_value!(normalize_imported_wallet_target(wallet_payload, unlimited)); + let wallet_target = match wallet_payload { + Some(wallet) => Some(invalid_value!(normalize_imported_wallet_target( + Some(wallet), + unlimited, + ))), + None => None, + }; if let Some(existing_key) = existing_standalone_by_hash.get(&key_hash).cloned() { match merge_mode { @@ -3276,11 +8036,33 @@ impl<'a> AdminAppState<'a> { return Ok(Err(invalid_request("独立余额 Key 已存在"))); } AdminImportMergeMode::Overwrite => { + let key_encrypted = invalid_value!(self + .seal_imported_auth_api_key_secret( + key_plaintext.as_deref(), + &existing_key.user_id, + &existing_key.api_key_id, + &key_hash, + true, + )); + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .existing_standalone_api_keys + .entry(existing_key.api_key_id.clone()) + .or_insert_with(|| ExistingApiKeyMutation { + before: existing_key.clone(), + after: existing_key.clone(), + }); + } let updated = self .update_standalone_api_key_basic( aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord { api_key_id: existing_key.api_key_id.clone(), + key_encrypted: key_encrypted.clone(), + key_encrypted_present: key_encrypted.is_some() + || mode == SystemImportMode::RecoveryRollbackCheckpoint, name: name.clone(), + name_present: name.is_some(), + force_capabilities: None, rate_limit_present: true, rate_limit: Some(rate_limit), concurrent_limit_present: key.contains_key("concurrent_limit"), @@ -3303,16 +8085,42 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } - let _ = self - .set_standalone_api_key_active(&existing_key.api_key_id, is_active) - .await?; + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + None, + &existing_key.api_key_id, + true, + ) + .await?; + let _ = require_persisted!( + self.set_standalone_api_key_active( + &existing_key.api_key_id, + mode.preserves_active_state() && is_active, + ) + .await? + ); + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + None, + &existing_key.api_key_id, + true, + ) + .await?; if key.contains_key("feature_settings") { - let _ = self - .set_standalone_api_key_feature_settings( + let _ = require_persisted!( + self.set_standalone_api_key_feature_settings( &existing_key.api_key_id, feature_settings.clone(), ) - .await?; + .await? + ); + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + None, + &existing_key.api_key_id, + true, + ) + .await?; } if imported_total_requests.is_some() || imported_total_tokens.is_some() @@ -3334,6 +8142,13 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); } + self.refresh_existing_api_key_mutation( + mutation_journal.as_deref_mut(), + None, + &existing_key.api_key_id, + true, + ) + .await?; } if key.contains_key("expires_at") || key.contains_key("auto_delete_on_expiry") @@ -3344,14 +8159,17 @@ impl<'a> AdminAppState<'a> { .to_string(), ); } - self.sync_imported_api_key_wallet( - &existing_key.api_key_id, - &wallet_target, - key.get("name") - .and_then(Value::as_str) - .unwrap_or("独立余额 Key"), - ) - .await?; + if let Some(wallet_target) = wallet_target.as_ref() { + self.sync_imported_api_key_wallet( + &existing_key.api_key_id, + wallet_target, + key.get("name") + .and_then(Value::as_str) + .unwrap_or("独立余额 Key"), + mutation_journal.as_deref_mut(), + ) + .await?; + } stats.standalone_keys.updated += 1; if let Some(source_api_key_id) = source_api_key_id.clone() { imported_api_key_id_map @@ -3362,11 +8180,19 @@ impl<'a> AdminAppState<'a> { continue; } + let api_key_id = imported_api_key_id_for_mode(source_api_key_id.as_deref(), mode); + let key_encrypted = invalid_value!(self.seal_imported_auth_api_key_secret( + key_plaintext.as_deref(), + &standalone_owner_id, + &api_key_id, + &key_hash, + true, + )); let created = self .create_standalone_api_key( aether_data::repository::auth::CreateStandaloneApiKeyRecord { user_id: standalone_owner_id.clone(), - api_key_id: Uuid::new_v4().to_string(), + api_key_id, key_hash: key_hash.clone(), key_encrypted, name, @@ -3377,7 +8203,7 @@ impl<'a> AdminAppState<'a> { rate_limit: Some(rate_limit), concurrent_limit, force_capabilities, - is_active, + is_active: mode.preserves_active_state() && is_active, expires_at_unix_secs, auto_delete_on_expiry, total_requests, @@ -3392,21 +8218,30 @@ impl<'a> AdminAppState<'a> { json!({ "detail": "Admin system data unavailable" }), ))); }; + let created_api_key_id = created.api_key_id.clone(); + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .standalone_api_key_ids + .insert(created_api_key_id.clone()); + } if key.contains_key("feature_settings") { - let _ = self - .set_standalone_api_key_feature_settings( + let _ = require_persisted!( + self.set_standalone_api_key_feature_settings( &created.api_key_id, feature_settings.clone(), ) - .await?; + .await? + ); + } + if let Some(wallet_target) = wallet_target.as_ref() { + self.sync_imported_api_key_wallet( + &created.api_key_id, + wallet_target, + created.name.as_deref().unwrap_or("独立余额 Key"), + mutation_journal.as_deref_mut(), + ) + .await?; } - self.sync_imported_api_key_wallet( - &created.api_key_id, - &wallet_target, - created.name.as_deref().unwrap_or("独立余额 Key"), - ) - .await?; - let created_api_key_id = created.api_key_id.clone(); existing_standalone_by_hash.insert(key_hash, created); if let Some(source_api_key_id) = source_api_key_id { imported_api_key_id_map.insert(source_api_key_id, created_api_key_id); @@ -3415,17 +8250,19 @@ impl<'a> AdminAppState<'a> { } } - if let Some(summary) = self - .import_admin_system_user_usage_aggregates( - root.get("usage_aggregates"), - &supplemental_user_usage_aggregates, - &imported_user_id_map, - &imported_api_key_id_map, - merge_mode, - ) - .await? - { - stats.usage_aggregates = Some(summary); + if !mode.is_rollback_checkpoint() { + if let Some(summary) = self + .import_admin_system_user_usage_aggregates( + root.get("usage_aggregates"), + &supplemental_user_usage_aggregates, + &imported_user_id_map, + &imported_api_key_id_map, + merge_mode, + ) + .await? + { + stats.usage_aggregates = Some(summary); + } } Ok(Ok(json!({ @@ -3442,58 +8279,11 @@ impl<'a> AdminAppState<'a> { api_key_id_map: &BTreeMap, merge_mode: AdminImportMergeMode, ) -> Result, GatewayError> { - let mut snapshot = match value { - Some(value) if !value.is_null() => serde_json::from_value::< - AdminSystemUsageAggregateSnapshot, - >(value.clone()) - .map_err(|err| GatewayError::Client { + let snapshot = build_imported_usage_aggregate_snapshot(value, supplemental_user_daily) + .map_err(|message| GatewayError::Client { status: http::StatusCode::BAD_REQUEST, - message: format!("usage_aggregates 格式无效: {err}"), - })?, - _ => AdminSystemUsageAggregateSnapshot::default(), - }; - let mut existing_user_totals = BTreeMap::::new(); - for row in &snapshot.stats_user_daily { - let total_tokens = row - .input_tokens - .saturating_add(row.output_tokens) - .saturating_add(row.cache_creation_tokens) - .saturating_add(row.cache_read_tokens); - let entry = existing_user_totals - .entry(row.user_id.clone()) - .or_insert((0, 0)); - entry.0 = entry.0.saturating_add(row.total_requests); - entry.1 = entry.1.saturating_add(total_tokens); - } - for row in supplemental_user_daily { - let existing = existing_user_totals - .get(&row.user_id) - .copied() - .unwrap_or_default(); - let request_delta = row.total_requests.saturating_sub(existing.0); - let token_delta = row.input_tokens.saturating_sub(existing.1); - if request_delta == 0 && token_delta == 0 { - continue; - } - if let Some(existing_row) = snapshot - .stats_user_daily - .iter_mut() - .rev() - .find(|existing_row| existing_row.user_id == row.user_id) - { - existing_row.total_requests = - existing_row.total_requests.saturating_add(request_delta); - existing_row.success_requests = - existing_row.success_requests.saturating_add(request_delta); - existing_row.input_tokens = existing_row.input_tokens.saturating_add(token_delta); - } else { - let mut row = row.clone(); - row.total_requests = request_delta; - row.success_requests = request_delta; - row.input_tokens = token_delta; - snapshot.stats_user_daily.push(row); - } - } + message, + })?; if snapshot.stats_daily.is_empty() && snapshot.stats_user_daily.is_empty() && snapshot.stats_daily_api_key.is_empty() @@ -3515,23 +8305,73 @@ impl<'a> AdminAppState<'a> { user_id: &str, wallet_target: &ImportedWalletTarget, label: &str, + mut mutation_journal: Option<&mut AggregateMutationJournal>, ) -> Result<(), GatewayError> { - if self - .find_wallet(WalletLookupKey::UserId(user_id)) - .await? - .is_none() + let initialized = self + .initialize_auth_user_wallet_with_outcome(user_id, 0.0, false) + .await?; + let Some(initialized) = initialized else { + return Err(GatewayError::Internal(format!( + "failed to initialize imported wallet for {label}" + ))); + }; + if initialized.wallet.user_id.as_deref() != Some(user_id) + || initialized.wallet.api_key_id.is_some() { - let created = self - .initialize_auth_user_wallet(user_id, 0.0, false) - .await?; - if created.is_none() { - return Err(GatewayError::Internal(format!( - "failed to initialize imported wallet for {label}" - ))); + return Err(GatewayError::Internal(format!( + "imported user wallet owner does not match {label}" + ))); + } + let created_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let existing_wallet_key = + (!initialized.created).then(|| (user_id.to_string(), initialized.wallet.id.clone())); + // Record only a row this invocation actually created. The repository returns this bit + // from the same atomic operation, so a concurrent initializer's wallet is never treated + // as import-owned and later deleted during compensation. The initial snapshot is a + // fallback for a failure before the imported values are persisted; a later successful + // sync replaces it with the complete post-import snapshot. + if let Some(wallet_id) = created_wallet_id.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.user_wallet_snapshots.insert( + (user_id.to_string(), wallet_id.clone()), + initialized.wallet.clone(), + ); } } - self.sync_wallet_snapshot(WalletOwner::User(user_id), wallet_target, label) - .await + if let Some(key) = existing_wallet_key.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .existing_user_wallets + .entry(key.clone()) + .or_insert_with(|| ExistingWalletMutation { + before: initialized.wallet.clone(), + after: None, + }); + } + } + let synced = self + .sync_wallet_snapshot(WalletOwner::User(user_id), wallet_target, label) + .await?; + if let Some(key) = existing_wallet_key.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + if let Some(mutation) = journal.existing_user_wallets.get_mut(key) { + mutation.after = Some(synced.clone()); + } + } + } + if !Self::imported_wallet_snapshot_matches_target(&synced, wallet_target) { + return Err(GatewayError::Internal(format!( + "persisted imported wallet snapshot changed during sync for {label}" + ))); + } + if let Some(wallet_id) = created_wallet_id { + if let Some(journal) = mutation_journal { + journal + .user_wallet_snapshots + .insert((user_id.to_string(), wallet_id), synced); + } + } + Ok(()) } async fn sync_imported_api_key_wallet( @@ -3539,23 +8379,68 @@ impl<'a> AdminAppState<'a> { api_key_id: &str, wallet_target: &ImportedWalletTarget, label: &str, + mut mutation_journal: Option<&mut AggregateMutationJournal>, ) -> Result<(), GatewayError> { - if self - .find_wallet(WalletLookupKey::ApiKeyId(api_key_id)) - .await? - .is_none() + let initialized = self + .initialize_auth_api_key_wallet_with_outcome(api_key_id, 0.0, false) + .await?; + let Some(initialized) = initialized else { + return Err(GatewayError::Internal(format!( + "failed to initialize imported wallet for {label}" + ))); + }; + if initialized.wallet.api_key_id.as_deref() != Some(api_key_id) + || initialized.wallet.user_id.is_some() { - let created = self - .initialize_auth_api_key_wallet(api_key_id, 0.0, false) - .await?; - if created.is_none() { - return Err(GatewayError::Internal(format!( - "failed to initialize imported wallet for {label}" - ))); + return Err(GatewayError::Internal(format!( + "imported API-key wallet owner does not match {label}" + ))); + } + let created_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let existing_wallet_key = + (!initialized.created).then(|| (api_key_id.to_string(), initialized.wallet.id.clone())); + if let Some(wallet_id) = created_wallet_id.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + journal.api_key_wallet_snapshots.insert( + (api_key_id.to_string(), wallet_id.clone()), + initialized.wallet.clone(), + ); } } - self.sync_wallet_snapshot(WalletOwner::ApiKey(api_key_id), wallet_target, label) - .await + if let Some(key) = existing_wallet_key.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + journal + .existing_api_key_wallets + .entry(key.clone()) + .or_insert_with(|| ExistingWalletMutation { + before: initialized.wallet.clone(), + after: None, + }); + } + } + let synced = self + .sync_wallet_snapshot(WalletOwner::ApiKey(api_key_id), wallet_target, label) + .await?; + if let Some(key) = existing_wallet_key.as_ref() { + if let Some(journal) = mutation_journal.as_deref_mut() { + if let Some(mutation) = journal.existing_api_key_wallets.get_mut(key) { + mutation.after = Some(synced.clone()); + } + } + } + if !Self::imported_wallet_snapshot_matches_target(&synced, wallet_target) { + return Err(GatewayError::Internal(format!( + "persisted imported wallet snapshot changed during sync for {label}" + ))); + } + if let Some(wallet_id) = created_wallet_id { + if let Some(journal) = mutation_journal { + journal + .api_key_wallet_snapshots + .insert((api_key_id.to_string(), wallet_id), synced); + } + } + Ok(()) } async fn sync_wallet_snapshot( @@ -3563,7 +8448,7 @@ impl<'a> AdminAppState<'a> { owner: WalletOwner<'_>, wallet_target: &ImportedWalletTarget, label: &str, - ) -> Result<(), GatewayError> { + ) -> Result { let updated = match owner { WalletOwner::User(user_id) => { self.update_auth_user_wallet_snapshot( @@ -3598,28 +8483,141 @@ impl<'a> AdminAppState<'a> { .await? } }; - if updated.is_none() { + let Some(updated) = updated else { return Err(GatewayError::Internal(format!( "failed to persist imported wallet snapshot for {label}" ))); - } - Ok(()) + }; + Ok(updated) + } + + fn imported_wallet_snapshot_matches_target( + snapshot: &StoredWalletSnapshot, + target: &ImportedWalletTarget, + ) -> bool { + const AMOUNT_EPSILON_USD: f64 = 0.00000001; + let amount_matches = |actual: f64, expected: f64| { + actual.is_finite() + && expected.is_finite() + && (actual - expected).abs() <= AMOUNT_EPSILON_USD + }; + amount_matches(snapshot.balance, target.recharge_balance) + && amount_matches(snapshot.gift_balance, target.gift_balance) + && snapshot.limit_mode == target.limit_mode + && snapshot.currency == target.currency + && snapshot.status == target.status + && amount_matches(snapshot.total_recharged, target.total_recharged) + && amount_matches(snapshot.total_consumed, target.total_consumed) + && amount_matches(snapshot.total_refunded, target.total_refunded) + && amount_matches(snapshot.total_adjusted, target.total_adjusted) + && target + .updated_at_unix_secs + .is_none_or(|expected| snapshot.updated_at_unix_secs == expected) } fn resolve_imported_system_user_api_key_material( &self, key: &Map, - ) -> Result)>, String> { + users_export_version: (u32, u32), + mode: SystemImportMode, + ) -> Result, String> { + let source_api_key_id = imported_optional_string(key.get("api_key_id"))?; let plaintext_key = imported_optional_string(key.get("key"))?; - if let Some(plaintext_key) = plaintext_key.filter(|value| !value.is_empty()) { - return Ok(Some(( - hash_admin_user_api_key(&plaintext_key), - self.encrypt_catalog_secret_with_fallbacks(&plaintext_key), - ))); - } let key_hash = imported_optional_string(key.get("key_hash"))?; let key_encrypted = imported_optional_string(key.get("key_encrypted"))?; - Ok(key_hash.map(|key_hash| (key_hash, key_encrypted))) + + if users_export_version >= (1, 6) { + if key.contains_key("key") + || key.contains_key("key_hash") + || key.contains_key("key_encrypted") + { + return Err( + "用户数据 1.6+ 不允许包含 key、key_hash 或 key_encrypted 凭据字段".to_string(), + ); + } + let credential_state = imported_optional_string(key.get("credential_state"))?; + if credential_state.as_deref() != Some("not_exported") { + return Err( + "用户数据 1.6+ API Key 必须标记 credential_state=not_exported".to_string(), + ); + } + let source_api_key_id = source_api_key_id + .filter(|value| !value.is_empty()) + .ok_or_else(|| "用户数据 1.6+ API Key 必须包含 api_key_id".to_string())?; + return Ok(Some(ImportedApiKeyMaterial { + key_hash: imported_credential_tombstone(&format!("api-key-id:{source_api_key_id}")), + key_plaintext: None, + })); + } + + if mode.restores_credentials() { + let key_hash = key_hash + .filter(|value| !value.is_empty()) + .ok_or_else(|| "恢复备份中的 API Key 必须包含 key_hash".to_string())?; + if key_hash.len() != 64 + || !key_hash + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) + { + return Err("恢复备份中的 key_hash 必须是规范的小写 SHA-256 十六进制".to_string()); + } + if plaintext_key.is_some() && key_encrypted.is_some() { + return Err("恢复备份中的 API Key 不能同时包含 key 和 key_encrypted".to_string()); + } + let plaintext = match (plaintext_key, key_encrypted) { + (Some(plaintext), None) => Some(plaintext), + (None, Some(ciphertext)) => Some( + self.decrypt_catalog_secret_with_fallbacks(&ciphertext) + .ok_or_else(|| { + "恢复备份中的 API Key 旧密文无法使用当前或历史数据密钥解密".to_string() + })?, + ), + (None, None) => None, + (Some(_), Some(_)) => unreachable!(), + }; + if let Some(plaintext) = plaintext.as_deref() { + if hash_admin_user_api_key(plaintext) != key_hash { + return Err("恢复备份中的 API Key 明文与 key_hash 不匹配".to_string()); + } + } + return Ok(Some(ImportedApiKeyMaterial { + key_hash, + key_plaintext: plaintext, + })); + } + + let identity = source_api_key_id + .map(|value| format!("api-key-id:{value}")) + .or_else(|| key_hash.map(|value| format!("legacy-key-hash:{value}"))) + .or_else(|| plaintext_key.map(|value| format!("legacy-key:{value}"))) + .or_else(|| key_encrypted.map(|value| format!("legacy-key-encrypted:{value}"))); + Ok(identity.map(|identity| ImportedApiKeyMaterial { + key_hash: imported_credential_tombstone(&identity), + key_plaintext: None, + })) + } + + fn seal_imported_auth_api_key_secret( + &self, + plaintext: Option<&str>, + user_id: &str, + api_key_id: &str, + key_hash: &str, + is_standalone: bool, + ) -> Result, String> { + plaintext + .map(|plaintext| { + seal_auth_api_key_secret( + self.app(), + user_id, + api_key_id, + key_hash, + is_standalone, + plaintext, + ) + .map_err(|_| "gateway 无法为目的 API Key 记录加密恢复凭据".to_string()) + }) + .transpose() } } @@ -3641,13 +8639,20 @@ mod tests { use serde_json::json; use super::{ - build_imported_user_usage_total_aggregates, imported_oauth_auth_config_has_credentials, - imported_oauth_expiry_after_import, imported_optional_bool, imported_optional_f64, - imported_optional_i32, imported_optional_u64, imported_rfc3339_to_unix_secs, - imported_string_list_from_value, normalize_import_endpoint_format, - normalize_import_key_formats, normalize_import_key_raw_payload, - normalize_imported_wallet_target, seed_imported_oauth_pool_score, - validate_imported_system_users_export_version, ImportedProviderKey, + build_imported_user_usage_total_aggregates, imported_api_key_id_for_mode, + imported_credential_tombstone, imported_existing_user_is_protected, + imported_oauth_auth_config_has_credentials, imported_oauth_expiry_after_import, + imported_optional_bool, imported_optional_f64, imported_optional_i32, + imported_optional_u64, imported_rfc3339_to_unix_secs, imported_string_list_from_value, + normalize_import_endpoint_format, normalize_import_key_formats, + normalize_import_key_raw_payload, normalize_imported_system_user_role, + normalize_imported_wallet_target, prepare_imported_secret_safe_body_rules, + prepare_imported_secret_safe_header_rules, prepare_imported_secret_safe_json, + resolve_imported_password_hash, seed_imported_oauth_pool_score, + validate_imported_provider_key_credential_state, + validate_imported_system_users_export_version, + validate_imported_system_users_export_version_for_mode, ImportedProviderKey, + SystemImportMode, }; use crate::admin_api::AdminAppState; use crate::data::GatewayDataState; @@ -3658,9 +8663,10 @@ mod tests { assert!(validate_imported_system_users_export_version(Some(&json!("1.3"))).is_ok()); assert!(validate_imported_system_users_export_version(Some(&json!("1.4"))).is_ok()); assert!(validate_imported_system_users_export_version(Some(&json!("1.5"))).is_ok()); + assert!(validate_imported_system_users_export_version(Some(&json!("1.6"))).is_ok()); assert_eq!( validate_imported_system_users_export_version(Some(&json!("2.2"))).unwrap_err(), - "不支持的用户数据版本: 2.2,支持的版本: 1.3, 1.4, 1.5" + "不支持的用户数据版本: 2.2,支持的版本: 1.3, 1.4, 1.5, 1.6" ); assert_eq!( validate_imported_system_users_export_version(Some(&json!(null))).unwrap_err(), @@ -3668,6 +8674,278 @@ mod tests { ); } + #[test] + fn recovery_users_import_requires_v15_and_valid_bcrypt() { + assert_eq!( + validate_imported_system_users_export_version_for_mode( + Some(&json!("1.5")), + SystemImportMode::RecoveryBackup, + ), + Ok((1, 5)), + ); + assert!(validate_imported_system_users_export_version_for_mode( + Some(&json!("1.4")), + SystemImportMode::RecoveryBackup, + ) + .is_err()); + assert!(resolve_imported_password_hash( + json!({ "password_hash": "attacker-controlled" }) + .as_object() + .expect("fixture should be an object"), + (1, 5), + SystemImportMode::RecoveryBackup, + ) + .is_err()); + } + + #[test] + fn system_user_role_interactive_import_cannot_assign_admin_console_roles() { + assert_eq!( + normalize_imported_system_user_role(None, SystemImportMode::InteractiveUpload), + Ok(Some("user".to_string())) + ); + assert_eq!( + normalize_imported_system_user_role( + Some(&json!(" ADMIN ")), + SystemImportMode::InteractiveUpload, + ), + Ok(None) + ); + assert_eq!( + normalize_imported_system_user_role( + Some(&json!("audit_admin")), + SystemImportMode::InteractiveUpload, + ), + Ok(None) + ); + assert_eq!( + normalize_imported_system_user_role( + Some(&json!("audit_admin")), + SystemImportMode::RollbackCheckpoint, + ), + Ok(None) + ); + assert!(normalize_imported_system_user_role( + Some(&json!("owner")), + SystemImportMode::InteractiveUpload, + ) + .expect_err("unknown roles must be rejected before import writes") + .contains("不支持的用户角色")); + } + + #[test] + fn system_user_role_authenticated_recovery_restores_only_audit_admin() { + for mode in [ + SystemImportMode::RecoveryBackup, + SystemImportMode::RecoveryRollbackCheckpoint, + ] { + assert_eq!( + normalize_imported_system_user_role(Some(&json!("audit_admin")), mode), + Ok(Some("audit_admin".to_string())) + ); + assert_eq!( + normalize_imported_system_user_role(Some(&json!("admin")), mode), + Ok(None) + ); + } + } + + #[test] + fn system_user_role_ordinary_import_protects_existing_admin_console_users() { + for mode in [ + SystemImportMode::InteractiveUpload, + SystemImportMode::RollbackCheckpoint, + ] { + assert!(imported_existing_user_is_protected("admin", mode)); + assert!(imported_existing_user_is_protected("audit_admin", mode)); + assert!(!imported_existing_user_is_protected("user", mode)); + } + + assert!(imported_existing_user_is_protected( + "admin", + SystemImportMode::RecoveryBackup, + )); + assert!(!imported_existing_user_is_protected( + "audit_admin", + SystemImportMode::RecoveryBackup, + )); + assert!(!imported_existing_user_is_protected( + "audit_admin", + SystemImportMode::RecoveryRollbackCheckpoint, + )); + } + + #[test] + fn rollback_checkpoint_body_forces_overwrite_and_omits_local_proxy_nodes() { + let checkpoint = json!({ + "version": "2.3", + "credential_state": "not_exported", + "proxy_nodes": [{"id": "local-node"}], + "system_configs": [] + }); + let body = super::build_aggregate_rollback_body(&checkpoint, true) + .expect("rollback body should serialize"); + let parsed: serde_json::Value = serde_json::from_slice(&body).expect("body is JSON"); + + assert_eq!(parsed["merge_mode"], json!("overwrite")); + assert_eq!(parsed["proxy_nodes"], json!([])); + assert_eq!(parsed["credential_state"], json!("not_exported")); + } + + #[test] + fn rollback_checkpoint_can_skip_ldap_without_dropping_other_config_sections() { + let checkpoint = json!({ + "version": "2.3", + "ldap_config": { + "server_url": "ldaps://checkpoint.example.test", + "bind_dn": "cn=admin,dc=example,dc=test", + "base_dn": "dc=example,dc=test" + }, + "system_configs": [{"key": "module.example.enabled", "value": true}], + }); + let body = super::build_aggregate_rollback_body_with_options(&checkpoint, true, true) + .expect("rollback body should serialize"); + let parsed: serde_json::Value = serde_json::from_slice(&body).expect("body is JSON"); + + assert!(parsed.get("ldap_config").is_none()); + assert_eq!(parsed["system_configs"], checkpoint["system_configs"]); + assert_eq!(parsed["merge_mode"], json!("overwrite")); + assert_eq!(parsed["proxy_nodes"], json!([])); + } + + #[test] + fn rollback_checkpoint_body_rejects_non_object() { + let error = super::build_aggregate_rollback_body(&json!([1, 2, 3]), false) + .expect_err("non-object checkpoint must be rejected"); + assert!(error + .into_message() + .contains("checkpoint must be a JSON object")); + } + + #[test] + fn aggregate_users_rollback_body_excludes_all_wallet_snapshots() { + let checkpoint = json!({ + "version": "1.5", + "users": [{ + "id": "user-1", + "username": "checkpoint-user", + "request_count": 100, + "total_tokens": 200, + "wallet": {"balance": 10.0}, + "api_keys": [{ + "api_key_id": "key-1", + "name": "user key", + "total_requests": 101, + "total_tokens": 201, + "total_cost_usd": 1.25, + "wallet": {"balance": 20.0} + }] + }], + "standalone_keys": [{ + "api_key_id": "standalone-1", + "name": "standalone key", + "total_requests": 102, + "total_tokens": 202, + "total_cost_usd": 2.5, + "wallet": {"balance": 30.0} + }], + "usage_aggregates": { + "stats_daily": [{"date_unix_secs": 1}] + } + }); + + let body = super::build_aggregate_users_rollback_body(&checkpoint) + .expect("users rollback body should serialize"); + let parsed: serde_json::Value = serde_json::from_slice(&body).expect("body is JSON"); + + assert_eq!(parsed["merge_mode"], json!("overwrite")); + assert_eq!(parsed["users"][0]["username"], json!("checkpoint-user")); + assert!(parsed["users"][0].get("request_count").is_none()); + assert!(parsed["users"][0].get("total_tokens").is_none()); + assert!(parsed["users"][0].get("wallet").is_none()); + assert!(parsed["users"][0]["api_keys"][0].get("wallet").is_none()); + assert!(parsed["users"][0]["api_keys"][0] + .get("total_requests") + .is_none()); + assert!(parsed["users"][0]["api_keys"][0] + .get("total_tokens") + .is_none()); + assert!(parsed["users"][0]["api_keys"][0] + .get("total_cost_usd") + .is_none()); + assert!(parsed["standalone_keys"][0].get("wallet").is_none()); + assert!(parsed["standalone_keys"][0].get("total_requests").is_none()); + assert!(parsed["standalone_keys"][0].get("total_tokens").is_none()); + assert!(parsed["standalone_keys"][0].get("total_cost_usd").is_none()); + assert!(parsed.get("usage_aggregates").is_none()); + } + + #[test] + fn imported_api_key_tombstones_fit_legacy_columns_and_cannot_authenticate() { + use sha2::{Digest, Sha256}; + + let identity = "api-key-id:public-source-key-id"; + let tombstone = imported_credential_tombstone(identity); + let normal_auth_hash = format!("{:x}", Sha256::digest(identity.as_bytes())); + + assert_eq!(tombstone.len(), 64); + assert!(tombstone.starts_with("$aether-import-revoked$")); + assert!(!tombstone + .chars() + .all(|character| character.is_ascii_hexdigit())); + assert_ne!(tombstone, normal_auth_hash); + } + + #[test] + fn rollback_checkpoint_reuses_api_key_id_while_interactive_imports_rotate_it() { + let source_id = Some("checkpoint-api-key-id"); + let rollback_id = + imported_api_key_id_for_mode(source_id, SystemImportMode::RollbackCheckpoint); + assert_eq!(rollback_id, "checkpoint-api-key-id"); + + let interactive_id = + imported_api_key_id_for_mode(source_id, SystemImportMode::InteractiveUpload); + assert_ne!(interactive_id, "checkpoint-api-key-id"); + assert!(!interactive_id.is_empty()); + } + + #[test] + fn recovery_rollback_checkpoint_keeps_credentials_and_stable_ids() { + let mode = SystemImportMode::RecoveryRollbackCheckpoint; + assert!(mode.restores_credentials()); + assert!(mode.preserves_active_state()); + assert!(mode.is_rollback_checkpoint()); + assert_eq!( + imported_api_key_id_for_mode(Some("recovery-key-id"), mode), + "recovery-key-id" + ); + assert!(!SystemImportMode::InteractiveUpload.restores_credentials()); + } + + #[test] + fn rollback_checkpoint_requires_stable_user_id_instead_of_identifier_guessing() { + assert!(super::validate_rollback_user_source_id( + SystemImportMode::RollbackCheckpoint, + None, + ) + .expect_err("rollback without a source ID must be rejected") + .contains("拒绝按 email/username 猜测用户")); + assert!(super::validate_rollback_user_source_id( + SystemImportMode::RecoveryRollbackCheckpoint, + None, + ) + .is_err()); + assert!(super::validate_rollback_user_source_id( + SystemImportMode::RollbackCheckpoint, + Some("stable-user-id"), + ) + .is_ok()); + assert!( + super::validate_rollback_user_source_id(SystemImportMode::InteractiveUpload, None,) + .is_ok() + ); + } + #[test] fn users_import_builds_supplemental_usage_aggregates_from_summary_fields() { let users = vec![ @@ -3754,6 +9032,7 @@ mod tests { is_active: true, proxy: None, fingerprint: None, + credential_state: None, }; let (formats, missing) = normalize_import_key_formats(&item, &endpoint_formats); @@ -3784,6 +9063,7 @@ mod tests { "api_key", &["openai:responses".to_string()], None, + false, ); assert_eq!(payload["api_formats"], json!(["openai:responses"])); @@ -3812,11 +9092,169 @@ mod tests { "api_key", &["openai:responses".to_string()], None, + false, ); assert_eq!(payload["allow_auth_channel_mismatch_formats"], json!([])); } + #[test] + fn config_import_never_turns_redaction_markers_into_credentials() { + let mut item = ImportedProviderKey { + api_key: None, + auth_type: Some("api_key".to_string()), + auth_config: None, + name: Some("primary".to_string()), + note: None, + api_formats: Some(vec!["openai:chat".to_string()]), + supported_endpoints: None, + rate_multipliers: None, + internal_priority: None, + global_priority_by_format: None, + auth_type_by_format: None, + allow_auth_channel_mismatch_formats: None, + rpm_limit: None, + allowed_models: None, + capabilities: None, + cache_ttl_minutes: None, + max_probe_interval_minutes: None, + auto_fetch_models: None, + locked_models: None, + model_include_patterns: None, + model_exclude_patterns: None, + is_active: false, + proxy: None, + fingerprint: None, + credential_state: Some("not_exported".to_string()), + }; + + assert!(validate_imported_provider_key_credential_state(&item).unwrap()); + item.api_key = Some("***".to_string()); + assert!(validate_imported_provider_key_credential_state(&item).is_err()); + item.api_key = None; + item.credential_state = Some("unknown".to_string()); + assert!(validate_imported_provider_key_credential_state(&item).is_err()); + } + + #[test] + fn config_import_restores_existing_secret_safe_json_and_strips_new_placeholders() { + let existing = json!({ + "credentials": {"refresh_token": "old-refresh"}, + "region": "us-east-1" + }); + let incoming = json!({ + "credentials": "***", + "region": "eu-west-1" + }); + + let restored = + prepare_imported_secret_safe_json(Some(&existing), Some(incoming.clone()), true) + .expect("existing config should restore"); + assert_eq!( + restored["credentials"], + json!({"refresh_token": "old-refresh"}) + ); + assert_eq!(restored["region"], "eu-west-1"); + + let created = prepare_imported_secret_safe_json(None, Some(incoming), true) + .expect("safe config should remain"); + assert!(created.get("credentials").is_none()); + assert_eq!(created["region"], "eu-west-1"); + } + + #[test] + fn config_import_restores_matching_endpoint_rules_without_persisting_markers() { + let existing_headers = json!([{ + "action": "set", + "key": "Authorization", + "value": "Bearer old-secret" + }]); + let incoming_headers = json!([{ + "action": "set", + "key": "Authorization", + "value": "***", + "has_value": true + }]); + let restored_headers = prepare_imported_secret_safe_header_rules( + Some(&existing_headers), + Some(incoming_headers), + true, + ) + .expect("matching header rule should remain"); + assert_eq!(restored_headers[0]["value"], "Bearer old-secret"); + assert!(restored_headers[0].get("has_value").is_none()); + + let existing_body = json!([{ + "action": "set", + "path": "$.credentials.token", + "value": "old-body-secret" + }]); + let incoming_body = json!([{ + "action": "set", + "path": "$.credentials.token", + "value": "***", + "has_value": true + }]); + let restored_body = prepare_imported_secret_safe_body_rules( + Some(&existing_body), + Some(incoming_body), + true, + ) + .expect("matching body rule should remain"); + assert_eq!(restored_body[0]["value"], "old-body-secret"); + assert!(restored_body[0].get("has_value").is_none()); + } + + #[test] + fn config_import_drops_unrecoverable_endpoint_rule_placeholders() { + let incoming_headers = json!([{ + "action": "set", + "key": "Authorization", + "value": "***", + "has_value": true + }]); + assert_eq!( + prepare_imported_secret_safe_header_rules(None, Some(incoming_headers.clone()), true), + Some(json!([])) + ); + let markerless_placeholder = json!([{ + "action": "set", + "key": "Authorization", + "value": "***" + }]); + assert_eq!( + prepare_imported_secret_safe_header_rules(None, Some(markerless_placeholder), true,), + Some(json!([])) + ); + + let different_existing = json!([{ + "action": "set", + "key": "X-Different", + "value": "must-not-move" + }]); + assert_eq!( + prepare_imported_secret_safe_header_rules( + Some(&different_existing), + Some(incoming_headers), + true, + ), + Some(json!([])) + ); + + let incoming_body = json!([{ + "action": "regex_replace", + "path": "$.credentials.token", + "pattern": "***", + "replacement": "***", + "has_pattern": true, + "has_replacement": true + }]); + assert_eq!( + prepare_imported_secret_safe_body_rules(None, Some(incoming_body), true), + Some(json!([])) + ); + } + #[test] fn oauth_import_only_treats_non_empty_secret_fields_as_credentials() { assert!(!imported_oauth_auth_config_has_credentials(&json!({}))); @@ -3989,7 +9427,7 @@ mod tests { let target = normalize_imported_wallet_target(Some(wallet), false).unwrap(); assert_eq!(target.recharge_balance, -5.25); assert_eq!(target.gift_balance, 0.75); - assert_eq!(target.total_recharged, -5.25); + assert_eq!(target.total_recharged, 0.0); } #[test] @@ -4006,6 +9444,34 @@ mod tests { assert_eq!(target.gift_balance, 0.75); } + #[test] + fn import_rejects_negative_wallet_history_totals() { + for (field, value) in [ + ("total_recharged", -1.0), + ("total_consumed", -1.0), + ("total_refunded", -1.0), + ] { + let mut wallet = serde_json::Map::new(); + wallet.insert(field.to_string(), json!(value)); + let error = normalize_imported_wallet_target(Some(&wallet), false) + .expect_err("negative wallet history total must be rejected"); + assert!(error.contains(field), "error should identify {field}"); + assert!( + error.contains("非负"), + "error should require non-negative {field}" + ); + } + } + + #[test] + fn import_allows_signed_wallet_adjustment_total() { + let wallet = json!({"total_adjusted": -3.5}); + let wallet = wallet.as_object().expect("wallet should be object"); + let target = normalize_imported_wallet_target(Some(wallet), false) + .expect("signed adjustment history should remain valid"); + assert_eq!(target.total_adjusted, -3.5); + } + #[test] fn import_rejects_legacy_string_lists() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/admin/request/system/mod.rs b/apps/aether-gateway/src/handlers/admin/request/system/mod.rs index 708a62ad5..b43241a3f 100644 --- a/apps/aether-gateway/src/handlers/admin/request/system/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/request/system/mod.rs @@ -10,6 +10,11 @@ mod templates; const ADMIN_SYSTEM_DATA_EXPORT_VERSION: &str = "1.0"; +pub(crate) use self::export::{ + is_interactive_export_private_system_config_key, SystemExportMode, + ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED, +}; + impl<'a> AdminAppState<'a> { pub(crate) async fn upsert_system_config_json_value( &self, diff --git a/apps/aether-gateway/src/handlers/admin/request/users.rs b/apps/aether-gateway/src/handlers/admin/request/users.rs index 12c802676..39d8554a3 100644 --- a/apps/aether-gateway/src/handlers/admin/request/users.rs +++ b/apps/aether-gateway/src/handlers/admin/request/users.rs @@ -105,6 +105,16 @@ impl<'a> AdminAppState<'a> { self.app.update_user_group(group_id, record).await } + pub(crate) async fn restore_user_group_if_matches( + &self, + expected: &aether_data::repository::users::StoredUserGroup, + restored: &aether_data::repository::users::StoredUserGroup, + ) -> Result { + self.app + .restore_user_group_if_matches(expected, restored) + .await + } + pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result { self.app.delete_user_group(group_id).await } @@ -152,6 +162,17 @@ impl<'a> AdminAppState<'a> { .await } + pub(crate) async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + self.app + .restore_user_groups_if_matches(user_id, expected_group_ids, restored_group_ids) + .await + } + pub(crate) async fn include_default_user_group_ids( &self, group_ids: &[String], @@ -246,6 +267,18 @@ impl<'a> AdminAppState<'a> { .await } + pub(crate) async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, GatewayError> + { + self.app + .initialize_auth_user_wallet_with_outcome(user_id, initial_gift_usd, unlimited) + .await + } + pub(crate) async fn initialize_auth_api_key_wallet( &self, api_key_id: &str, @@ -257,14 +290,85 @@ impl<'a> AdminAppState<'a> { .await } + pub(crate) async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, GatewayError> + { + self.app + .initialize_auth_api_key_wallet_with_outcome(api_key_id, initial_gift_usd, unlimited) + .await + } + + pub(crate) async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + self.app + .delete_wallet_if_unreferenced(wallet_id, owner) + .await + } + + pub(crate) async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &aether_data::repository::wallet::StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + self.app + .delete_wallet_if_snapshot_matches_and_unreferenced(expected, owner) + .await + } + + pub(crate) async fn restore_wallet_if_snapshot_matches( + &self, + before: &aether_data::repository::wallet::StoredWalletSnapshot, + after: &aether_data::repository::wallet::StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + self.app + .restore_wallet_if_snapshot_matches(before, after, owner) + .await + } + pub(crate) async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, GatewayError> { self.app - .update_local_auth_user_profile(user_id, email, username) + .update_local_auth_user_profile(user_id, email_present, email, email_verified, username) + .await + } + + #[allow(clippy::too_many_arguments)] + pub(crate) async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &aether_data::repository::users::StoredUserAuthRecord, + restored_auth: &aether_data::repository::users::StoredUserAuthRecord, + expected_export: &aether_data::repository::users::StoredUserExportRow, + restored_export: &aether_data::repository::users::StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + self.app + .restore_local_auth_user_state_if_matches( + expected_auth, + restored_auth, + expected_export, + restored_export, + expected_model_capability_settings, + restored_model_capability_settings, + expected_feature_settings, + restored_feature_settings, + ) .await } @@ -279,6 +383,34 @@ impl<'a> AdminAppState<'a> { .await } + pub(crate) async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: chrono::DateTime, + ) -> Result { + self.app + .restore_local_auth_user_password_hash_if_matches( + user_id, + expected_password_hash, + password_hash, + updated_at, + ) + .await + } + + pub(crate) async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + self.app + .reset_local_auth_user_password_and_revoke_sessions(user_id, password_hash, changed_at) + .await + } + #[allow(clippy::too_many_arguments)] pub(crate) async fn update_local_auth_user_admin_fields( &self, @@ -456,6 +588,16 @@ impl<'a> AdminAppState<'a> { self.app.delete_local_auth_user(user_id).await } + pub(crate) async fn rollback_provisional_auth_user_with_wallet( + &self, + user_id: &str, + wallet_id: Option<&str>, + ) -> Result<(), GatewayError> { + self.app + .rollback_provisional_auth_user_with_wallet(user_id, wallet_id) + .await + } + pub(crate) async fn list_user_sessions( &self, user_id: &str, @@ -625,6 +767,16 @@ impl<'a> AdminAppState<'a> { self.app.update_standalone_api_key_basic(record).await } + pub(crate) async fn restore_api_key_if_matches( + &self, + expected: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + restored: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + ) -> Result { + self.app + .restore_api_key_if_matches(expected, restored) + .await + } + pub(crate) async fn set_standalone_api_key_active( &self, api_key_id: &str, diff --git a/apps/aether-gateway/src/handlers/admin/routing/mod.rs b/apps/aether-gateway/src/handlers/admin/routing/mod.rs index 942df8a36..49d5e904a 100644 --- a/apps/aether-gateway/src/handlers/admin/routing/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/routing/mod.rs @@ -44,6 +44,8 @@ struct AdminRoutingGroupCreateRequest { #[serde(default)] is_system_default: bool, #[serde(default)] + sort_order: i64, + #[serde(default)] config_json: Option, } @@ -137,6 +139,7 @@ async fn maybe_build_routing_groups_response( description: payload.description, enabled: payload.enabled, is_system_default: payload.is_system_default, + sort_order: payload.sort_order, config_json, version: 1, created_at: now, @@ -464,6 +467,9 @@ fn build_routing_group_update_patch( if let Some(value) = object.get("is_system_default") { patch.is_system_default = Some(required_bool(value, "is_system_default")?); } + if let Some(value) = object.get("sort_order") { + patch.sort_order = Some(required_i64(value, "sort_order")?.max(0)); + } if let Some(value) = object.get("config_json") { validate_config_json(value)?; patch.config_json = Some(value.clone()); @@ -604,6 +610,7 @@ fn routing_group_payload(group: &StoredRoutingGroup) -> Value { "description": group.description, "enabled": group.enabled, "is_system_default": group.is_system_default, + "sort_order": group.sort_order, "config_json": group.config_json, "version": group.version, "created_at": group.created_at, diff --git a/apps/aether-gateway/src/handlers/admin/shared/mod.rs b/apps/aether-gateway/src/handlers/admin/shared/mod.rs index 9d13f5c47..28f1262d0 100644 --- a/apps/aether-gateway/src/handlers/admin/shared/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/shared/mod.rs @@ -11,10 +11,11 @@ pub(crate) use crate::handlers::shared::{ attach_admin_audit_response, build_admin_provider_key_response, decrypt_catalog_secret_with_fallbacks, default_provider_key_status_snapshot, effective_catalog_encryption_key, encrypt_catalog_secret_with_fallbacks, json_string_list, - masked_catalog_api_key, masked_catalog_api_key_for_provider, normalize_json_array, - normalize_json_object, normalize_string_list, parse_catalog_auth_config_json, - provider_catalog_key_supports_format, provider_key_health_summary, - provider_key_health_summary_at, provider_key_status_snapshot_payload, query_param_bool, - query_param_optional_bool, query_param_value, take_secret_prefix, take_secret_suffix, - unix_secs_to_rfc3339, OFFICIAL_EXTERNAL_MODEL_PROVIDERS, + mark_sensitive_admin_response_no_store, masked_catalog_api_key, + masked_catalog_api_key_for_provider, normalize_json_array, normalize_json_object, + normalize_string_list, parse_catalog_auth_config_json, provider_catalog_key_supports_format, + provider_key_health_summary, provider_key_health_summary_at, + provider_key_status_snapshot_payload, query_param_bool, query_param_optional_bool, + query_param_value, take_secret_prefix, take_secret_suffix, unix_secs_to_rfc3339, + OFFICIAL_EXTERNAL_MODEL_PROVIDERS, }; diff --git a/apps/aether-gateway/src/handlers/admin/system/core/system_routes.rs b/apps/aether-gateway/src/handlers/admin/system/core/system_routes.rs index 819bc2410..ad76409ea 100644 --- a/apps/aether-gateway/src/handlers/admin/system/core/system_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/system/core/system_routes.rs @@ -1,5 +1,5 @@ use super::ADMIN_AWS_REGIONS; -use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::admin::request::{AdminAppState, AdminRequestContext, SystemExportMode}; use crate::handlers::admin::shared::attach_admin_audit_response; use crate::handlers::admin::shared::build_proxy_error_response; use crate::handlers::admin::system::shared::configs::{ @@ -23,10 +23,15 @@ use crate::handlers::admin::system::shared::update::{ prepare_admin_system_update_task, read_update_history, read_update_task_status, self_update_supported, start_admin_system_rollback_task, start_admin_system_update_task, }; +use crate::handlers::admin::system::{ + execute_admin_system_import_exclusively, release_admin_system_import_lease, + try_acquire_admin_system_import_lease, AdminSystemImportLockError, +}; use crate::important_notification::build_important_notification_test_payload; use crate::maintenance::{ManualUsageCleanupMode, ManualUsageCleanupOptions}; use crate::GatewayError; use aether_data_contracts::repository::usage::UsageCleanupTargets; +use aether_runtime_state::RuntimeLockLease; use axum::{ body::{Body, Bytes}, http, @@ -34,6 +39,7 @@ use axum::{ Json, }; use serde_json::json; +use std::future::Future; use std::time::Instant; use url::form_urlencoded; @@ -233,12 +239,21 @@ pub(super) async fn maybe_build_local_admin_core_system_response( && request_method == http::Method::GET && request_path == "/api/admin/system/config/export" { - return Ok(Some(attach_admin_audit_response( - Json(state.build_admin_system_config_export_payload().await?).into_response(), - "admin_system_config_exported", - "export_system_config", - "system_config_export", - "global", + return Ok(Some(sensitive_system_export_response( + attach_admin_audit_response( + Json( + state + .build_admin_system_config_export_payload( + SystemExportMode::InteractiveDownload, + ) + .await?, + ) + .into_response(), + "admin_system_config_exported", + "export_system_config", + "system_config_export", + "global", + ), ))); } @@ -255,30 +270,46 @@ pub(super) async fn maybe_build_local_admin_core_system_response( .into_response(), )); }; - return Ok(Some( - match state.import_admin_system_config(request_body).await? { - Ok(payload) => attach_admin_audit_response( - Json(payload).into_response(), - "admin_system_config_imported", - "import_system_config", - "system_config_import", - "global", - ), - Err((status, payload)) => (status, Json(payload)).into_response(), - }, - )); + let import_result = match execute_admin_system_import_with_lock( + state, + state.import_admin_system_config(request_body), + ) + .await + { + Ok(result) => result, + Err(response) => return Ok(Some(response)), + }; + return Ok(Some(match import_result? { + Ok(payload) => attach_admin_audit_response( + Json(payload).into_response(), + "admin_system_config_imported", + "import_system_config", + "system_config_import", + "global", + ), + Err((status, payload)) => (status, Json(payload)).into_response(), + })); } if decision.route_kind.as_deref() == Some("users_export") && request_method == http::Method::GET && request_path == "/api/admin/system/users/export" { - return Ok(Some(attach_admin_audit_response( - Json(state.build_admin_system_users_export_payload().await?).into_response(), - "admin_system_users_exported", - "export_system_users", - "user_export", - "all_users", + return Ok(Some(sensitive_system_export_response( + attach_admin_audit_response( + Json( + state + .build_admin_system_users_export_payload( + SystemExportMode::InteractiveDownload, + ) + .await?, + ) + .into_response(), + "admin_system_users_exported", + "export_system_users", + "user_export", + "all_users", + ), ))); } @@ -295,39 +326,52 @@ pub(super) async fn maybe_build_local_admin_core_system_response( .into_response(), )); }; - return Ok(Some( - match state - .import_admin_system_users( - request_body, - decision - .admin_principal - .as_ref() - .map(|principal| principal.user_id.as_str()), - ) - .await? - { - Ok(payload) => attach_admin_audit_response( - Json(payload).into_response(), - "admin_system_users_imported", - "import_system_users", - "system_users_import", - "global", - ), - Err((status, payload)) => (status, Json(payload)).into_response(), - }, - )); + let import_result = match execute_admin_system_import_with_lock( + state, + state.import_admin_system_users( + request_body, + decision + .admin_principal + .as_ref() + .map(|principal| principal.user_id.as_str()), + ), + ) + .await + { + Ok(result) => result, + Err(response) => return Ok(Some(response)), + }; + return Ok(Some(match import_result? { + Ok(payload) => attach_admin_audit_response( + Json(payload).into_response(), + "admin_system_users_imported", + "import_system_users", + "system_users_import", + "global", + ), + Err((status, payload)) => (status, Json(payload)).into_response(), + })); } if decision.route_kind.as_deref() == Some("data_export") && request_method == http::Method::GET && request_path == "/api/admin/system/data/export" { - return Ok(Some(attach_admin_audit_response( - Json(state.build_admin_system_data_export_payload().await?).into_response(), - "admin_system_data_exported", - "export_system_data", - "system_data_export", - "global", + return Ok(Some(sensitive_system_export_response( + attach_admin_audit_response( + Json( + state + .build_admin_system_data_export_payload( + SystemExportMode::InteractiveDownload, + ) + .await?, + ) + .into_response(), + "admin_system_data_exported", + "export_system_data", + "system_data_export", + "global", + ), ))); } @@ -344,27 +388,31 @@ pub(super) async fn maybe_build_local_admin_core_system_response( .into_response(), )); }; - return Ok(Some( - match state - .import_admin_system_data( - request_body, - decision - .admin_principal - .as_ref() - .map(|principal| principal.user_id.as_str()), - ) - .await? - { - Ok(payload) => attach_admin_audit_response( - Json(payload).into_response(), - "admin_system_data_imported", - "import_system_data", - "system_data_import", - "global", - ), - Err((status, payload)) => (status, Json(payload)).into_response(), - }, - )); + let import_result = match execute_admin_system_import_with_lock( + state, + state.import_admin_system_data( + request_body, + decision + .admin_principal + .as_ref() + .map(|principal| principal.user_id.as_str()), + ), + ) + .await + { + Ok(result) => result, + Err(response) => return Ok(Some(response)), + }; + return Ok(Some(match import_result? { + Ok(payload) => attach_admin_audit_response( + Json(payload).into_response(), + "admin_system_data_imported", + "import_system_data", + "system_data_import", + "global", + ), + Err((status, payload)) => (status, Json(payload)).into_response(), + })); } if decision.route_kind.as_deref() == Some("s3_backup_run") @@ -1103,6 +1151,62 @@ fn bad_manual_cleanup_request(detail: impl Into) -> Response { .into_response() } +fn sensitive_system_export_response(mut response: Response) -> Response { + response.headers_mut().insert( + http::header::CACHE_CONTROL, + http::HeaderValue::from_static("no-store"), + ); + response +} + +async fn acquire_admin_system_import_lock( + state: &AdminAppState<'_>, +) -> Result> { + try_acquire_admin_system_import_lease(state.app()) + .await + .map_err(admin_system_import_lock_error_response) +} + +async fn execute_admin_system_import_with_lock( + state: &AdminAppState<'_>, + operation: F, +) -> Result> +where + F: Future, +{ + execute_admin_system_import_exclusively(state.app(), operation) + .await + .map_err(admin_system_import_lock_error_response) +} + +fn admin_system_import_lock_error_response(error: AdminSystemImportLockError) -> Response { + match error { + AdminSystemImportLockError::Conflict => ( + http::StatusCode::CONFLICT, + Json(json!({ + "detail": "已有系统数据导入正在执行,请等待当前导入完成后再试" + })), + ) + .into_response(), + AdminSystemImportLockError::Unavailable => ( + http::StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ "detail": "系统数据导入服务暂时不可用" })), + ) + .into_response(), + AdminSystemImportLockError::Lost => ( + http::StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ + "detail": "系统数据导入锁已丢失,操作已停止;可能存在部分已提交变更,请检查运行状态后重试" + })), + ) + .into_response(), + } +} + +async fn release_admin_system_import_lock(state: &AdminAppState<'_>, lock: &RuntimeLockLease) { + release_admin_system_import_lease(state.app(), lock).await; +} + fn query_param(query_string: Option<&str>, name: &str) -> Option { let query = query_string.filter(|value| !value.is_empty())?; form_urlencoded::parse(query.as_bytes()) @@ -1142,6 +1246,26 @@ fn parse_older_than_days_query(query_string: Option<&str>) -> Result mod tests { use super::*; + #[tokio::test] + async fn admin_system_import_lock_serializes_and_releases_imports() { + let app = crate::AppState::new().expect("app state should build"); + let state = AdminAppState::new(&app); + + let first = acquire_admin_system_import_lock(&state) + .await + .expect("first import should acquire the lock"); + let conflict = acquire_admin_system_import_lock(&state) + .await + .expect_err("concurrent import must be rejected"); + assert_eq!(conflict.status(), http::StatusCode::CONFLICT); + + release_admin_system_import_lock(&state, &first).await; + let second = acquire_admin_system_import_lock(&state) + .await + .expect("lock should be reusable after release"); + release_admin_system_import_lock(&state, &second).await; + } + #[test] fn manual_cleanup_request_defaults_to_policy_targets() { let options = parse_manual_usage_cleanup_request(None).expect("default request is valid"); @@ -1188,4 +1312,14 @@ mod tests { assert!(parse_manual_usage_cleanup_request(Some(&body)).is_err()); } + + #[test] + fn sensitive_system_exports_are_never_cacheable() { + let response = sensitive_system_export_response(Json(json!({})).into_response()); + + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&http::HeaderValue::from_static("no-store")) + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/system/import_lock.rs b/apps/aether-gateway/src/handlers/admin/system/import_lock.rs new file mode 100644 index 000000000..96423975c --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/system/import_lock.rs @@ -0,0 +1,325 @@ +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use aether_runtime_state::{RuntimeLockLease, RuntimeState}; +use tokio::time::MissedTickBehavior; + +use crate::AppState; + +const ADMIN_SYSTEM_IMPORT_LOCK_KEY: &str = "admin:system:import"; +const ADMIN_SYSTEM_IMPORT_LOCK_TTL: Duration = Duration::from_secs(60 * 60 * 6); +const ADMIN_SYSTEM_IMPORT_LOCK_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(60 * 5); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AdminSystemImportLockError { + Conflict, + Unavailable, + Lost, +} + +#[derive(Debug, PartialEq, Eq)] +enum AdminSystemImportRaceOutcome { + OperationCompleted(T), + LeaseLost, +} + +#[derive(Debug, PartialEq, Eq)] +enum AdminSystemImportLockRenewalFailure { + Lost, + Backend(E), +} + +struct AdminSystemImportLeaseGuard { + runtime_state: Arc, + lease: Option, +} + +impl AdminSystemImportLeaseGuard { + fn new(app: &AppState, lease: RuntimeLockLease) -> Self { + Self { + runtime_state: app.runtime_state.clone(), + lease: Some(lease), + } + } + + fn lease(&self) -> &RuntimeLockLease { + self.lease + .as_ref() + .expect("admin system import lease guard must own a lease") + } + + async fn release(&mut self) -> Result<(), AdminSystemImportLockError> { + let Some(lease) = self.lease.clone() else { + return Ok(()); + }; + match self.runtime_state.lock_release(&lease).await { + Ok(true) => { + self.lease.take(); + Ok(()) + } + Ok(false) => { + tracing::warn!( + lock_key = %lease.key, + "admin system import lock was no longer owned during release" + ); + self.lease.take(); + Err(AdminSystemImportLockError::Lost) + } + Err(error) => { + tracing::warn!(error = %error, "admin system import lock release failed"); + // Keep the lease in the guard so Drop can make one best-effort retry. + Err(AdminSystemImportLockError::Lost) + } + } + } +} + +impl Drop for AdminSystemImportLeaseGuard { + fn drop(&mut self) { + let Some(lease) = self.lease.take() else { + return; + }; + let runtime_state = self.runtime_state.clone(); + let Ok(handle) = tokio::runtime::Handle::try_current() else { + return; + }; + drop(handle.spawn(async move { + match runtime_state.lock_release(&lease).await { + Ok(true) => {} + Ok(false) => tracing::warn!( + lock_key = %lease.key, + "admin system import lock was no longer owned during asynchronous release" + ), + Err(error) => { + tracing::warn!( + error = %error, + lock_key = %lease.key, + "admin system import lock asynchronous release failed" + ); + } + } + })); + } +} + +pub(crate) async fn try_acquire_admin_system_import_lease( + app: &AppState, +) -> Result { + match app + .runtime_state() + .lock_try_acquire( + ADMIN_SYSTEM_IMPORT_LOCK_KEY, + app.tunnel.local_instance_id(), + ADMIN_SYSTEM_IMPORT_LOCK_TTL, + ) + .await + { + Ok(Some(lock)) => Ok(lock), + Ok(None) => Err(AdminSystemImportLockError::Conflict), + Err(error) => { + tracing::warn!(error = %error, "admin system import lock acquisition failed"); + Err(AdminSystemImportLockError::Unavailable) + } + } +} + +pub(crate) async fn release_admin_system_import_lease(app: &AppState, lock: &RuntimeLockLease) { + match app.runtime_state().lock_release(lock).await { + Ok(true) => {} + Ok(false) => tracing::warn!( + lock_key = %lock.key, + "admin system import lock was no longer owned during release" + ), + Err(error) => { + tracing::warn!(error = %error, "admin system import lock release failed"); + } + } +} + +fn require_successful_admin_system_import_lock_renewal( + result: Result, +) -> Result<(), AdminSystemImportLockRenewalFailure> { + match result { + Ok(true) => Ok(()), + Ok(false) => Err(AdminSystemImportLockRenewalFailure::Lost), + Err(error) => Err(AdminSystemImportLockRenewalFailure::Backend(error)), + } +} + +async fn wait_for_admin_system_import_lease_loss( + runtime_state: Arc, + lease: RuntimeLockLease, +) { + let mut heartbeat = tokio::time::interval(ADMIN_SYSTEM_IMPORT_LOCK_HEARTBEAT_INTERVAL); + heartbeat.set_missed_tick_behavior(MissedTickBehavior::Delay); + heartbeat.tick().await; + loop { + heartbeat.tick().await; + match require_successful_admin_system_import_lock_renewal( + runtime_state + .lock_renew(&lease, ADMIN_SYSTEM_IMPORT_LOCK_TTL) + .await, + ) { + Ok(()) => {} + Err(AdminSystemImportLockRenewalFailure::Lost) => { + tracing::warn!( + lock_key = %lease.key, + fencing_token = lease.fencing_token, + "admin system import lock is no longer owned; cancelling the import" + ); + return; + } + Err(AdminSystemImportLockRenewalFailure::Backend(error)) => { + tracing::warn!( + error = %error, + lock_key = %lease.key, + fencing_token = lease.fencing_token, + "admin system import lock renewal failed; cancelling the import" + ); + return; + } + } + } +} + +async fn race_admin_system_import_with_lease_loss( + operation: F, + lease_loss: L, +) -> AdminSystemImportRaceOutcome +where + F: Future, + L: Future, +{ + tokio::pin!(operation); + tokio::pin!(lease_loss); + tokio::select! { + biased; + _ = &mut lease_loss => AdminSystemImportRaceOutcome::LeaseLost, + result = &mut operation => AdminSystemImportRaceOutcome::OperationCompleted(result), + } +} + +pub(crate) async fn execute_admin_system_import_exclusively( + app: &AppState, + operation: F, +) -> Result +where + F: Future, +{ + let lease = try_acquire_admin_system_import_lease(app).await?; + let mut guard = AdminSystemImportLeaseGuard::new(app, lease); + let lease_loss = + wait_for_admin_system_import_lease_loss(guard.runtime_state.clone(), guard.lease().clone()); + + match race_admin_system_import_with_lease_loss(operation, lease_loss).await { + AdminSystemImportRaceOutcome::OperationCompleted(result) => { + guard.release().await?; + Ok(result) + } + AdminSystemImportRaceOutcome::LeaseLost => Err(AdminSystemImportLockError::Lost), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::convert::Infallible; + use std::sync::atomic::{AtomicBool, Ordering}; + + struct DropSignal(Arc); + + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + #[test] + fn admin_system_import_lock_renewal_requires_current_ownership() { + assert_eq!( + require_successful_admin_system_import_lock_renewal::(Ok(true)), + Ok(()) + ); + assert_eq!( + require_successful_admin_system_import_lock_renewal::(Ok(false)), + Err(AdminSystemImportLockRenewalFailure::Lost) + ); + assert_eq!( + require_successful_admin_system_import_lock_renewal(Err("redis unavailable")), + Err(AdminSystemImportLockRenewalFailure::Backend( + "redis unavailable" + )) + ); + } + + #[tokio::test] + async fn admin_system_import_operation_is_cancelled_when_lease_is_lost() { + let dropped = Arc::new(AtomicBool::new(false)); + let operation_dropped = dropped.clone(); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let operation = async move { + let _drop_signal = DropSignal(operation_dropped); + let _ = started_tx.send(()); + std::future::pending::<()>().await; + }; + let lease_loss = async move { + started_rx + .await + .expect("operation should be polled before reporting lease loss"); + }; + + let outcome = race_admin_system_import_with_lease_loss(operation, lease_loss).await; + + assert_eq!(outcome, AdminSystemImportRaceOutcome::LeaseLost); + assert!(dropped.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn completed_admin_system_import_releases_lease_for_reuse() { + let app = AppState::new().expect("app state should build"); + + let result = execute_admin_system_import_exclusively(&app, async { 42_u8 }) + .await + .expect("import should complete"); + assert_eq!(result, 42); + + let lease = try_acquire_admin_system_import_lease(&app) + .await + .expect("completed import should release its lease"); + release_admin_system_import_lease(&app, &lease).await; + } + + #[tokio::test] + async fn cancelled_admin_system_import_releases_lease_without_waiting_for_ttl() { + let app = AppState::new().expect("app state should build"); + let task_app = app.clone(); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let task = tokio::spawn(async move { + execute_admin_system_import_exclusively(&task_app, async move { + let _ = started_tx.send(()); + std::future::pending::<()>().await; + }) + .await + }); + started_rx + .await + .expect("operation should start after acquiring the lease"); + + task.abort(); + let _ = task.await; + + let lease = tokio::time::timeout(Duration::from_secs(1), async { + loop { + match try_acquire_admin_system_import_lease(&app).await { + Ok(lease) => break lease, + Err(AdminSystemImportLockError::Conflict) => tokio::task::yield_now().await, + Err(error) => panic!("unexpected lock acquisition error: {error:?}"), + } + } + }) + .await + .expect("cancelled import should release its lease promptly"); + release_admin_system_import_lease(&app, &lease).await; + } +} diff --git a/apps/aether-gateway/src/handlers/admin/system/management_tokens.rs b/apps/aether-gateway/src/handlers/admin/system/management_tokens.rs index 4a105eaa7..ec8fde00a 100644 --- a/apps/aether-gateway/src/handlers/admin/system/management_tokens.rs +++ b/apps/aether-gateway/src/handlers/admin/system/management_tokens.rs @@ -1,7 +1,5 @@ use crate::control::{ - management_token_permission_catalog_payload, - management_token_permissions_cover_all_assignable_permissions, - normalize_assignable_management_token_permissions, + management_token_permission_catalog_payload, normalize_assignable_management_token_permissions, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value}; @@ -58,6 +56,35 @@ fn admin_management_token_read_only_response() -> Response { .into_response() } +fn admin_management_token_secret_response(mut response: Response) -> Response { + response.headers_mut().insert( + http::header::CACHE_CONTROL, + http::HeaderValue::from_static("no-store"), + ); + response +} + +fn admin_management_token_internal_error_response( + trace_id: &str, + event_name: &'static str, + error: impl std::fmt::Debug, +) -> Response { + tracing::error!( + event_name, + trace_id, + error = ?error, + "management token operation failed" + ); + ( + http::StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ + "detail": "Management Token 服务暂不可用,请稍后重试", + "trace_id": trace_id, + })), + ) + .into_response() +} + #[derive(Debug, Clone)] struct AdminManagementTokenCreateInput { name: String, @@ -93,8 +120,12 @@ fn hash_admin_management_token(value: &str) -> String { } fn admin_management_token_prefix(value: &str) -> Option { - (!value.is_empty()) - .then(|| value[..value.len().min(ADMIN_MANAGEMENT_TOKEN_DISPLAY_PREFIX_LEN)].to_string()) + (!value.is_empty()).then(|| { + value + .chars() + .take(ADMIN_MANAGEMENT_TOKEN_DISPLAY_PREFIX_LEN) + .collect() + }) } fn admin_parse_management_token_allowed_ips( @@ -228,13 +259,14 @@ fn admin_parse_management_token_update_input( async fn admin_management_token_user_summary( state: &AdminAppState<'_>, user_id: &str, + trace_id: &str, ) -> Result> { let Some(user) = state.find_user_auth_by_id(user_id).await.map_err(|err| { - ( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": format!("management token user lookup failed: {err:?}") })), + admin_management_token_internal_error_response( + trace_id, + "admin_management_token_user_lookup_failed", + err, ) - .into_response() })? else { return Err(( @@ -243,19 +275,15 @@ async fn admin_management_token_user_summary( ) .into_response()); }; - StoredManagementTokenUserSummary::new( - user.id, - user.email, - user.username, - user.role, + StoredManagementTokenUserSummary::new(user.id, user.email, user.username, user.role).map_err( + |err| { + admin_management_token_internal_error_response( + trace_id, + "admin_management_token_user_summary_build_failed", + err, + ) + }, ) - .map_err(|err| { - ( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": format!("management token user summary build failed: {err:?}") })), - ) - .into_response() - }) } pub(crate) async fn maybe_build_local_admin_management_tokens_response( @@ -275,17 +303,12 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response( .as_ref() .and_then(|principal| principal.management_token_id.as_deref()) .is_some(); - let management_token_is_full = decision - .admin_principal - .as_ref() - .and_then(|principal| principal.management_token_permissions.as_deref()) - .is_none_or(management_token_permissions_cover_all_assignable_permissions); - if is_management_token && !management_token_is_full { + if is_management_token { return Ok(Some( ( http::StatusCode::FORBIDDEN, Json(json!({ - "detail": "不允许使用 Management Token 管理其他 Token,请使用 Web 界面或 JWT 认证" + "detail": "不允许使用 Management Token 管理其他 Token,请使用管理员会话认证" })), ) .into_response(), @@ -359,7 +382,12 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response( let Some(admin_principal) = decision.admin_principal.as_ref() else { return Ok(None); }; - let user = match admin_management_token_user_summary(state, &admin_principal.user_id).await + let user = match admin_management_token_user_summary( + state, + &admin_principal.user_id, + request_context.trace_id(), + ) + .await { Ok(value) => value, Err(response) => return Ok(Some(response)), @@ -380,15 +408,17 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response( }; return Ok(Some(match state.create_management_token(&record).await? { - LocalMutationOutcome::Applied(token) => ( - http::StatusCode::CREATED, - Json(json!({ - "message": "Management Token 创建成功", - "token": raw_token, - "data": build_management_token_payload(&token, Some(&record.user)), - })), - ) - .into_response(), + LocalMutationOutcome::Applied(token) => admin_management_token_secret_response( + ( + http::StatusCode::CREATED, + Json(json!({ + "message": "Management Token 创建成功", + "token": raw_token, + "data": build_management_token_payload(&token, Some(&record.user)), + })), + ) + .into_response(), + ), LocalMutationOutcome::Invalid(detail) => { admin_management_token_bad_request_response(detail) } @@ -546,12 +576,14 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response( return Ok(Some( match state.regenerate_management_token_secret(&mutation).await? { - LocalMutationOutcome::Applied(token) => Json(json!({ - "message": "Token 已重新生成", - "token": raw_token, - "data": build_management_token_payload(&token, Some(&existing.user)), - })) - .into_response(), + LocalMutationOutcome::Applied(token) => admin_management_token_secret_response( + Json(json!({ + "message": "Token 已重新生成", + "token": raw_token, + "data": build_management_token_payload(&token, Some(&existing.user)), + })) + .into_response(), + ), LocalMutationOutcome::NotFound => admin_management_token_not_found_response(), LocalMutationOutcome::Invalid(detail) => { admin_management_token_bad_request_response(detail) @@ -563,3 +595,20 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response( Ok(None) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn plaintext_management_token_responses_are_never_cacheable() { + let response = admin_management_token_secret_response( + Json(json!({ "token": "ae-secret" })).into_response(), + ); + + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&http::HeaderValue::from_static("no-store")) + ); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/system/mod.rs b/apps/aether-gateway/src/handlers/admin/system/mod.rs index e8004221c..45188a7b6 100644 --- a/apps/aether-gateway/src/handlers/admin/system/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/system/mod.rs @@ -1,11 +1,16 @@ mod adaptive; mod core; +mod import_lock; mod management_tokens; mod modules; mod proxy_nodes; mod routes; pub(super) mod shared; +pub(crate) use self::import_lock::{ + execute_admin_system_import_exclusively, release_admin_system_import_lease, + try_acquire_admin_system_import_lease, AdminSystemImportLockError, +}; #[cfg(test)] pub(crate) use self::proxy_nodes::{ clear_proxy_node_references_with_cache_failure_for_tests, diff --git a/apps/aether-gateway/src/handlers/admin/system/proxy_nodes.rs b/apps/aether-gateway/src/handlers/admin/system/proxy_nodes.rs index 783ccad34..134b51952 100644 --- a/apps/aether-gateway/src/handlers/admin/system/proxy_nodes.rs +++ b/apps/aether-gateway/src/handlers/admin/system/proxy_nodes.rs @@ -19,11 +19,9 @@ use aether_admin::system::{ build_admin_proxy_node_payload, build_admin_proxy_nodes_data_unavailable_response, build_admin_proxy_nodes_not_found_response, }; -use aether_contracts::tunnel::{ - TUNNEL_RELAY_FORWARDED_BY_HEADER, TUNNEL_RELAY_OWNER_INSTANCE_HEADER, -}; +use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER; use aether_data::repository::management_tokens::{ - CreateManagementTokenRecord, StoredManagementTokenUserSummary, + CreateManagementTokenRecord, StoredManagementToken, StoredManagementTokenUserSummary, }; use aether_data::repository::proxy_nodes::{ProxyNodeEventQuery, ProxyNodeMetricsStep}; use axum::{ @@ -103,7 +101,7 @@ struct ProxyNodeUnregisterRequest { node_id: String, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct ManualProxyNodeCreateRequest { name: String, proxy_url: String, @@ -115,7 +113,7 @@ struct ManualProxyNodeCreateRequest { region: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct ManualProxyNodeUpdateRequest { #[serde(default)] name: Option, @@ -129,7 +127,7 @@ struct ManualProxyNodeUpdateRequest { region: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct ProxyNodeTestUrlRequest { proxy_url: String, #[serde(default)] @@ -177,6 +175,7 @@ const MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES: usize = 64 * 1024; const PROXY_NODE_METRICS_MAX_POINTS: usize = 50_000; const PROXY_NODE_METRICS_1M_MAX_WINDOW_SECS: u64 = 30 * 24 * 60 * 60; const PROXY_NODE_METRICS_1H_MAX_WINDOW_SECS: u64 = 365 * 24 * 60 * 60; +const PROXY_INSTALL_INTERNAL_ERROR_DETAIL: &str = "Service temporarily unavailable"; #[cfg(test)] fn manual_proxy_connectivity_probe_url_override() -> &'static std::sync::RwLock> { @@ -233,6 +232,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&Bytes>, ) -> Result>, GatewayError> { let Some(decision) = request_context.decision() else { @@ -374,9 +374,11 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( .tunnel .register_secure_tunnel_key(node.id.clone(), key); } + state.app().tunnel.request_close_proxies_for_node(&node.id); return Ok(Some( Json(json!({ "node_id": node.id, + "tunnel_generation": node.tunnel_generation, "node": build_admin_proxy_node_payload(&node), })) .into_response(), @@ -434,6 +436,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( let Some(node) = state.unregister_proxy_node(&node_id).await? else { return Ok(Some(build_admin_proxy_nodes_not_found_response())); }; + state.app().tunnel.request_close_proxies_for_node(&node.id); return Ok(Some( Json(json!({ "message": "unregistered", @@ -483,7 +486,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( Ok(node_name) => node_name, Err(response) => return Ok(Some(response)), }; - let raw_token = + let (token_record, raw_token) = match create_proxy_install_management_token(state, request_context, &node_name).await { Ok(token) => token, Err(response) => return Ok(Some(response)), @@ -493,7 +496,9 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( state.app(), request_context.public(), headers, + remote_addr, node_name, + &token_record, raw_token, ) .await, @@ -555,6 +560,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response( let Some(_deleted_node) = state.delete_proxy_node(&node_id).await? else { return Ok(build_admin_proxy_nodes_not_found_response()); }; + state.app().tunnel.request_close_proxies_for_node(&node_id); Ok(Json(json!({ "message": build_delete_proxy_node_message(&cleanup), "node_id": node_id, @@ -901,13 +907,7 @@ struct DeletedProxyNodeCleanup { fn build_admin_proxy_node_detail_payload( node: &aether_data::repository::proxy_nodes::StoredProxyNode, ) -> Value { - let mut payload = build_admin_proxy_node_payload(node); - if node.is_manual { - if let Value::Object(object) = &mut payload { - object.insert("proxy_password".to_string(), json!(node.proxy_password)); - } - } - payload + build_admin_proxy_node_payload(node) } #[derive(Debug, Clone)] @@ -1147,12 +1147,33 @@ async fn test_proxy_node_connectivity( ); } }; - let proxy_url = proxy_url_with_auth( + let proxy_password = match state.app().decrypt_proxy_node_password(&node.id).await { + Ok(password) => password, + Err(_) => { + return build_proxy_connectivity_result( + &probe_url, + PROXY_CONNECTIVITY_TIMEOUT_SECS, + false, + None, + None, + Some("手动节点密码不可用".to_string()), + ); + } + }; + let Some(proxy_url) = proxy_url_with_auth( &endpoint.proxy_url, node.proxy_username.as_deref(), - node.proxy_password.as_deref(), - ) - .unwrap_or(endpoint.proxy_url); + proxy_password.as_deref(), + ) else { + return build_proxy_connectivity_result( + &probe_url, + PROXY_CONNECTIVITY_TIMEOUT_SECS, + false, + None, + None, + Some("手动节点认证配置不可用".to_string()), + ); + }; return test_manual_proxy_connectivity(&proxy_url).await; } @@ -1251,7 +1272,25 @@ async fn test_manual_proxy_connectivity_with_probe_url( timeout_secs: u64, ) -> Value { let started_at = Instant::now(); - let proxy = match reqwest::Proxy::all(proxy_url) { + // reqwest resolves the destination locally for `socks5://`, which would + // bypass the gateway's private-address/DNS-rebinding guard. Normalize + // legacy SOCKS URLs to the remote-DNS form before probing; HTTP/HTTPS and + // already-normalized `socks5h://` URLs are unchanged. + let proxy_url = + match crate::execution_runtime::transport::normalize_execution_proxy_url(proxy_url) { + Ok(proxy_url) => proxy_url, + Err(_) => { + return build_proxy_connectivity_result( + probe_url, + timeout_secs, + false, + None, + None, + Some("代理 URL 无效".to_string()), + ); + } + }; + let proxy = match reqwest::Proxy::all(&proxy_url) { Ok(proxy) => proxy, Err(error) => { return build_proxy_connectivity_result( @@ -1260,24 +1299,17 @@ async fn test_manual_proxy_connectivity_with_probe_url( false, None, None, - Some(sanitize_proxy_error(&format_upstream_request_error(&error))), + Some(sanitize_proxy_error(&error.to_string())), ); } }; - let mut builder = reqwest::Client::builder() + let builder = reqwest::Client::builder() .no_proxy() .redirect(reqwest::redirect::Policy::none()) .connect_timeout(Duration::from_secs(5)) .timeout(Duration::from_secs(timeout_secs)) .proxy(proxy) .user_agent("aether-gateway/proxy-connectivity"); - if proxy_url - .trim() - .to_ascii_lowercase() - .starts_with("https://") - { - builder = builder.danger_accept_invalid_certs(true); - } let client = match builder.build() { Ok(client) => client, Err(error) => { @@ -1306,8 +1338,13 @@ async fn test_manual_proxy_connectivity_with_probe_url( } }; let status = response.status(); - let body = match response.text().await { - Ok(body) => body, + let body = match aether_http::read_response_bytes_with_limit( + response, + MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES, + ) + .await + { + Ok(body) => String::from_utf8_lossy(&body).into_owned(), Err(error) => { return build_proxy_connectivity_result( probe_url, @@ -1315,7 +1352,7 @@ async fn test_manual_proxy_connectivity_with_probe_url( false, None, None, - Some(sanitize_proxy_error(&format_upstream_request_error(&error))), + Some(sanitize_proxy_error(&error.to_string())), ); } }; @@ -1424,10 +1461,21 @@ async fn probe_tunnel_proxy_connectivity_via_owner( relay_base_url: &str, owner_instance_id: &str, ) -> Result { - let owner_url = build_tunnel_owner_relay_url(relay_base_url, node_id)?; + let owner_url = crate::tunnel::build_tunnel_owner_relay_url(relay_base_url, node_id)?; + let payload = build_tunnel_probe_relay_envelope(probe_url, timeout_secs)?; + let relay_auth = state.tunnel.build_relay_auth_headers( + owner_instance_id, + node_id, + true, + false, + &payload, + &[], + )?; let started_at = Instant::now(); - let response = state - .client + let owner_client = + crate::tunnel::owner_forward_client_for_url(&state.owner_forward_client, &owner_url) + .await?; + let request = owner_client .post(owner_url) .header( http::header::CONTENT_TYPE, @@ -1437,23 +1485,20 @@ async fn probe_tunnel_proxy_connectivity_via_owner( TUNNEL_RELAY_FORWARDED_BY_HEADER, state.tunnel.local_instance_id(), ) - .header(TUNNEL_RELAY_OWNER_INSTANCE_HEADER, owner_instance_id) .timeout(Duration::from_secs(timeout_secs)) - .body(build_tunnel_probe_relay_envelope(probe_url, timeout_secs)?) + .body(payload); + let response = relay_auth + .apply(request) .send() .await - .map_err(|error| format!("owner tunnel relay probe failed: {error}"))?; + .map_err(|error| crate::tunnel::owner_forward_request_error(&error))?; let status = response.status(); - let body = response - .bytes() - .await - .map_err(|error| format!("failed to read owner tunnel relay probe body: {error}"))?; - if body.len() > MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES { - return Err(format!( - "owner tunnel relay probe body exceeds {} bytes", - MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES - )); - } + let body = aether_http::read_response_bytes_with_limit( + response, + MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES, + ) + .await + .map_err(|_| "failed to read owner tunnel relay probe body".to_string())?; Ok(TunnelConnectivityProbeResult { status: status.as_u16(), @@ -1489,23 +1534,6 @@ fn build_tunnel_probe_relay_envelope( Ok(envelope) } -fn build_tunnel_owner_relay_url(relay_base_url: &str, node_id: &str) -> Result { - let mut url = url::Url::parse(relay_base_url) - .map_err(|error| format!("invalid owner relay base url: {error}"))?; - { - let mut segments = url - .path_segments_mut() - .map_err(|_| "owner relay base url cannot be a base-less URL".to_string())?; - segments.pop_if_empty(); - segments.push("api"); - segments.push("internal"); - segments.push("tunnel"); - segments.push("relay"); - segments.push(node_id.trim()); - } - Ok(url.to_string()) -} - fn validate_register_request( input: ProxyNodeRegisterRequest, request_context: &AdminRequestContext<'_>, @@ -1577,6 +1605,7 @@ fn validate_register_request( Ok( aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation { + node_id: None, name, ip, port: i32::from(input.port.unwrap_or_default()), @@ -1611,6 +1640,7 @@ fn validate_manual_create_request( Ok( aether_data::repository::proxy_nodes::ProxyNodeManualCreateMutation { + node_id: None, name: normalize_required_string(&input.name, "name", 100)?, ip: endpoint.node_ip, port: endpoint.node_port, @@ -1665,12 +1695,12 @@ fn validate_proxy_test_url_request( let username = normalize_optional_string(input.username.as_deref(), "username", 255)?; let password = normalize_optional_string(input.password.as_deref(), "password", 500)?; let endpoint = normalize_manual_proxy_endpoint(&input.proxy_url)?; - Ok(proxy_url_with_auth( + proxy_url_with_auth( &endpoint.proxy_url, username.as_deref(), password.as_deref(), ) - .unwrap_or(endpoint.proxy_url)) + .ok_or_else(|| bad_request_response("password 需要非空 username,且代理 URL 必须支持认证")) } fn admin_proxy_node_upgrade_action_node_id_from_path(path: &str, suffix: &str) -> Option { @@ -1760,6 +1790,7 @@ async fn dispatch_proxy_node_upgrade_targets( .update_proxy_node_remote_config( &aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation { node_id: node.id.clone(), + expected_tunnel_generation: None, node_name: None, allowed_ports: None, log_level: None, @@ -1851,6 +1882,7 @@ fn validate_heartbeat_request( Ok( aether_data::repository::proxy_nodes::ProxyNodeHeartbeatMutation { node_id, + expected_tunnel_generation: None, heartbeat_interval: input.heartbeat_interval, active_connections: input.active_connections, total_requests_delta: input.total_requests, @@ -1964,6 +1996,7 @@ fn validate_remote_config_request( Ok( aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation { node_id, + expected_tunnel_generation: None, node_name, allowed_ports, log_level, @@ -2036,6 +2069,12 @@ fn parse_manual_proxy_endpoint( if !parsed.username().is_empty() || parsed.password().is_some() { return Err(format!("{field} 不应包含用户名或密码,请使用独立字段")); } + if !matches!(parsed.path(), "" | "/") || parsed.query().is_some() || parsed.fragment().is_some() + { + return Err(format!( + "{field} 必须是代理 origin,不能包含 path、query 或 fragment" + )); + } let host = parsed .host_str() .map(str::trim) @@ -2119,13 +2158,74 @@ fn normalize_ip_address(value: &str) -> Result> { } fn sanitize_proxy_error(detail: &str) -> String { - match detail.split_once("://") { - Some((scheme, rest)) => match rest.split_once('@') { - Some((_, tail)) => format!("{scheme}://***@{tail}"), - None => detail.to_string(), - }, - None => detail.to_string(), + const HTTP_STATUS_PREFIX: &str = "代理探测返回 HTTP "; + const CLASSIFICATION_PREFIX_BYTES: usize = 4 * 1024; + + if let Some(status) = detail + .strip_prefix(HTTP_STATUS_PREFIX) + .and_then(|rest| { + rest.split(|character: char| !character.is_ascii_digit()) + .next() + }) + .and_then(|status| status.parse::().ok()) + .filter(|status| (100..=599).contains(status)) + { + return format!("{HTTP_STATUS_PREFIX}{status}"); } + + let detail = if detail.len() <= CLASSIFICATION_PREFIX_BYTES { + detail + } else { + let mut end = CLASSIFICATION_PREFIX_BYTES; + while !detail.is_char_boundary(end) { + end = end.saturating_sub(1); + } + &detail[..end] + }; + let normalized = detail.to_ascii_lowercase(); + if normalized.contains("timed out") || normalized.contains("timeout") { + return "代理探测超时".to_string(); + } + if normalized.contains("overloaded") + || normalized.contains("backpressure") + || normalized.contains("congested") + || normalized.contains("busy") + { + return "代理探测服务繁忙".to_string(); + } + if normalized.contains("unauthorized") + || normalized.contains("forbidden") + || normalized.contains("authentication") + || normalized.contains("credential") + { + return "代理认证失败".to_string(); + } + if normalized.contains("too large") + || normalized.contains("body exceeds") + || normalized.contains("response exceeds") + { + return "代理探测响应过大".to_string(); + } + if normalized.contains("response body") + || normalized.contains("body read") + || normalized.contains("decode") + { + return "代理探测响应读取失败".to_string(); + } + if normalized.contains("connect") + || normalized.contains("dns") + || normalized.contains("socket") + || normalized.contains("not connected") + || normalized.contains("unavailable") + || normalized.contains("offline") + { + return "代理连接失败".to_string(); + } + + // Error strings can originate in reqwest, a remote gateway, or a tunnel + // peer. Keep arbitrary URLs, credentials, paths, and control characters + // out of the admin response by projecting unknown details to one category. + "代理探测失败".to_string() } fn proxy_url_with_auth( @@ -2133,13 +2233,22 @@ fn proxy_url_with_auth( username: Option<&str>, password: Option<&str>, ) -> Option { - let username = username.map(str::trim).filter(|value| !value.is_empty())?; + let username = username.filter(|value| !value.is_empty()); + let password = password.filter(|value| !value.is_empty()); let mut parsed = url::Url::parse(proxy_url).ok()?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + { + return None; + } + if username.is_none() && password.is_none() { + return Some(parsed.to_string()); + } + let username = username.unwrap_or(""); if parsed.set_username(username).is_err() { return None; } - let password = password.map(str::trim).filter(|value| !value.is_empty()); if parsed.set_password(password).is_err() { return None; } @@ -2152,27 +2261,18 @@ fn parse_proxy_probe_exit_ip(body: &str) -> Option { if key.trim() != "ip" { return None; } - let value = value.trim(); - if value.is_empty() { - return None; - } - Some(value.to_string()) + value + .trim() + .parse::() + .ok() + .map(|ip| ip.to_string()) }) } -fn format_proxy_probe_status_error(status: reqwest::StatusCode, body: &str) -> String { - let body = body.trim(); - if body.is_empty() { - return format!("代理探测返回 HTTP {}", status.as_u16()); - } - - let truncated = if body.chars().count() > 200 { - let shortened: String = body.chars().take(200).collect(); - format!("{shortened}...") - } else { - body.to_string() - }; - format!("代理探测返回 HTTP {}: {truncated}", status.as_u16()) +fn format_proxy_probe_status_error(status: reqwest::StatusCode, _body: &str) -> String { + // The body is controlled by the probe target and may echo proxy + // credentials or contain private upstream diagnostics. + format!("代理探测返回 HTTP {}", status.as_u16()) } fn validate_optional_counter(value: Option, field: &str) -> Result<(), Response> { @@ -2238,14 +2338,14 @@ fn hash_proxy_install_management_token(value: &str) -> String { } fn proxy_install_management_token_prefix(value: &str) -> Option { - (!value.is_empty()).then(|| value[..value.len().min(12)].to_string()) + (!value.is_empty()).then(|| value.chars().take(12).collect()) } async fn create_proxy_install_management_token( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, node_name: &str, -) -> Result> { +) -> Result<(StoredManagementToken, String), Response> { let Some(principal) = request_context .decision() .and_then(|decision| decision.admin_principal.as_ref()) @@ -2257,39 +2357,70 @@ async fn create_proxy_install_management_token( .into_response()); }; + let (allowed_ips, expires_at_unix_secs) = + if let Some(parent_token_id) = principal.management_token_id.as_deref() { + let parent = match state.get_management_token_with_user(parent_token_id).await { + Ok(Some(parent)) => parent, + Ok(None) => return Err(proxy_install_parent_token_denied_response()), + Err(_) => { + return Err(proxy_install_internal_error_response( + "proxy_install_parent_management_token_lookup", + "repository_lookup_failed", + )) + } + }; + let now = chrono::Utc::now().timestamp().max(0) as u64; + if parent.token.user_id != principal.user_id + || !parent.token.is_active + || parent + .token + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at <= now) + { + return Err(proxy_install_parent_token_denied_response()); + } + let permissions = crate::control::management_token_permission_keys_from_value( + parent.token.permissions.as_ref(), + ) + .map_err(|_| proxy_install_parent_token_denied_response())?; + let Some(decision) = request_context.decision() else { + return Err(proxy_install_parent_token_denied_response()); + }; + crate::control::validate_management_token_admin_route_permission( + request_context.method(), + decision, + permissions.as_deref(), + ) + .map_err(|_| proxy_install_parent_token_denied_response())?; + ( + parent.token.allowed_ips.clone(), + parent.token.expires_at_unix_secs, + ) + } else { + (None, None) + }; + let user = match state.app().find_user_auth_by_id(&principal.user_id).await { Ok(value) => value, - Err(err) => { - return Err(( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": format!("admin user lookup failed: {err:?}") })), - ) - .into_response()) + Err(_) => { + return Err(proxy_install_internal_error_response( + "proxy_install_admin_user_lookup", + "repository_lookup_failed", + )) } }; - let user = user - .map(|user| { - StoredManagementTokenUserSummary::new( - user.id, - user.email, - user.username, - user.role, + let Some(user) = user else { + return Err(proxy_install_parent_token_denied_response()); + }; + if !user.is_active || user.is_deleted || !user.role.eq_ignore_ascii_case("admin") { + return Err(proxy_install_parent_token_denied_response()); + } + let user = StoredManagementTokenUserSummary::new(user.id, user.email, user.username, user.role) + .map_err(|_| { + proxy_install_internal_error_response( + "proxy_install_management_token_user_summary_build", + "invalid_user_summary", ) - }) - .unwrap_or_else(|| { - StoredManagementTokenUserSummary::new( - principal.user_id.clone(), - None, - principal.user_id.clone(), - principal.user_role.clone(), - ) - }) - .map_err(|err| { - ( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": format!("management token user summary build failed: {err:?}") })), - ) - .into_response() })?; let raw_token = generate_proxy_install_management_token_plaintext(); @@ -2307,14 +2438,15 @@ async fn create_proxy_install_management_token( token_prefix: proxy_install_management_token_prefix(&raw_token), name: format!("aether-tunnel {node_name} {short_id}"), description: Some("Created by proxy node one-click installer".to_string()), - allowed_ips: None, + allowed_ips, permissions: Some(json!(["admin:proxy_nodes:write"])), - expires_at_unix_secs: None, - is_active: true, + expires_at_unix_secs, + // The bearer secret is not usable until its one-time install session is consumed. + is_active: false, }; match state.app().create_management_token(&record).await { - Ok(LocalMutationOutcome::Applied(_)) => Ok(raw_token), + Ok(LocalMutationOutcome::Applied(stored)) => Ok((stored, raw_token)), Ok(LocalMutationOutcome::Invalid(detail)) => Err(bad_request_response(detail)), Ok(LocalMutationOutcome::Unavailable) => { Err(build_admin_proxy_nodes_data_unavailable_response()) @@ -2324,14 +2456,38 @@ async fn create_proxy_install_management_token( Json(json!({ "detail": "管理员不存在" })), ) .into_response()), - Err(err) => Err(( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": format!("management token create failed: {err:?}") })), - ) - .into_response()), + Err(_) => Err(proxy_install_internal_error_response( + "proxy_install_management_token_create", + "repository_write_failed", + )), } } +fn proxy_install_internal_error_response( + operation: &'static str, + error_category: &'static str, +) -> Response { + warn!( + event_name = "proxy_install_internal_error", + operation, error_category, "proxy install operation failed" + ); + ( + http::StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ "detail": PROXY_INSTALL_INTERNAL_ERROR_DETAIL })), + ) + .into_response() +} + +fn proxy_install_parent_token_denied_response() -> Response { + ( + http::StatusCode::FORBIDDEN, + Json(json!({ + "detail": "parent management token is no longer authorized to create install sessions" + })), + ) + .into_response() +} + fn parse_proxy_node_event_query( query: Option<&str>, ) -> Result> { @@ -2419,3 +2575,132 @@ fn bad_request_response(detail: impl Into) -> Response { ) .into_response() } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn proxy_auth_url_construction_never_falls_back_to_unauthenticated() { + assert_eq!( + proxy_url_with_auth("http://proxy.example:8080", None, None).as_deref(), + Some("http://proxy.example:8080/") + ); + assert_eq!( + proxy_url_with_auth("http://proxy.example:8080", None, Some("secret")).as_deref(), + Some("http://:secret@proxy.example:8080/") + ); + assert!(proxy_url_with_auth("not a proxy url", Some("alice"), Some("secret")).is_none()); + assert!(proxy_url_with_auth("mailto:proxy@example.com", Some("alice"), None).is_none()); + } + + #[test] + fn manual_proxy_endpoint_accepts_only_an_origin_without_ambiguous_components() { + for value in [ + "http://proxy.example:8080/path", + "http://proxy.example:8080?token=secret", + "http://proxy.example:8080#fragment", + "http://alice:password@proxy.example:8080", + "file:///tmp/proxy", + ] { + assert!( + parse_manual_proxy_endpoint(value, "proxy_url").is_err(), + "proxy URL should be rejected: {value}" + ); + } + + for value in [ + "http://proxy.example:8080", + "https://proxy.example:8443/", + "socks5://proxy.example:1080", + "socks5h://proxy.example:1080", + ] { + assert!( + parse_manual_proxy_endpoint(value, "proxy_url").is_ok(), + "proxy origin should be accepted: {value}" + ); + } + } + + #[test] + fn proxy_connectivity_errors_are_projected_without_sensitive_details() { + let details = [ + "request failed for https://alice:proxy-secret@10.0.0.8/probe?access_token=query-secret", + "Bearer bearer-secret\r\nx-injected: true", + "failed to open /private/var/proxy-secret.pem", + ]; + + for detail in details { + let projected = sanitize_proxy_error(detail); + for secret in [ + "alice", + "proxy-secret", + "10.0.0.8", + "query-secret", + "bearer-secret", + "x-injected", + "/private/var", + ] { + assert!(!projected.contains(secret), "leaked {secret}: {projected}"); + } + assert!(!projected.contains(['\r', '\n'])); + } + + assert_eq!( + sanitize_proxy_error("connection timed out for https://secret.internal"), + "代理探测超时" + ); + assert_eq!( + sanitize_proxy_error("DNS connect error for http://10.0.0.2"), + "代理连接失败" + ); + } + + #[test] + fn proxy_probe_status_error_does_not_echo_untrusted_body() { + let detail = format_proxy_probe_status_error( + reqwest::StatusCode::BAD_GATEWAY, + "Bearer upstream-secret at http://10.0.0.9/private?token=query-secret", + ); + + assert_eq!(detail, "代理探测返回 HTTP 502"); + assert_eq!(sanitize_proxy_error(&detail), detail); + assert!(!detail.contains("upstream-secret")); + assert!(!detail.contains("10.0.0.9")); + } + + #[test] + fn proxy_probe_exit_ip_only_accepts_ip_addresses() { + assert_eq!( + parse_proxy_probe_exit_ip("fl=1\nip=2001:db8::1\nts=2"), + Some("2001:db8::1".to_string()) + ); + assert_eq!( + parse_proxy_probe_exit_ip("ip=Bearer upstream-secret\nx=1"), + None + ); + } + + #[tokio::test] + async fn proxy_install_internal_error_response_hides_internal_details() { + let response = proxy_install_internal_error_response( + "secret-bearing operation https://internal.example/token", + "Bearer super-secret", + ); + + assert_eq!(response.status(), http::StatusCode::INTERNAL_SERVER_ERROR); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("internal error response body should be readable"); + let payload: Value = + serde_json::from_slice(&body).expect("internal error response should be JSON"); + + assert_eq!( + payload, + json!({ "detail": PROXY_INSTALL_INTERNAL_ERROR_DETAIL }) + ); + let body = String::from_utf8_lossy(&body); + assert!(!body.contains("internal.example")); + assert!(!body.contains("super-secret")); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/system/routes.rs b/apps/aether-gateway/src/handlers/admin/system/routes.rs index b0151b12a..e95ee8a11 100644 --- a/apps/aether-gateway/src/handlers/admin/system/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/system/routes.rs @@ -39,6 +39,7 @@ pub(crate) async fn maybe_build_local_admin_system_response( &request.state(), &request.request_context(), request.request_headers(), + request.remote_addr(), request.request_body(), ) .await? diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/configs.rs b/apps/aether-gateway/src/handlers/admin/system/shared/configs.rs index 8cd5ad72a..ebb94dd1f 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/configs.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/configs.rs @@ -1,6 +1,9 @@ use crate::handlers::admin::model::ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY; use crate::handlers::admin::request::AdminAppState; -use crate::handlers::shared::unix_secs_to_rfc3339; +use crate::handlers::shared::{ + bark_device_key_binding, encrypt_bark_device_key, encrypt_smtp_password, smtp_password_binding, + system_config_bool, system_config_string, unix_secs_to_rfc3339, +}; use crate::GatewayError; use aether_admin::system::{ admin_system_config_default_value as admin_system_config_default_value_pure, @@ -13,7 +16,6 @@ use aether_admin::system::{ normalize_admin_system_config_key as normalize_admin_system_config_key_pure, parse_admin_system_config_update, }; -use aether_crypto::encrypt_python_fernet_plaintext; use axum::body::Bytes; use axum::http; use serde_json::json; @@ -125,18 +127,70 @@ pub(crate) async fn apply_admin_system_config_update( if is_sensitive_admin_system_config_key(&normalized_key) && value.as_str().is_some_and(|raw| !raw.is_empty()) { - let Some(encryption_key) = state - .encryption_key() - .filter(|value| !value.trim().is_empty()) - else { + let plaintext = value + .as_str() + .expect("sensitive config value was a non-empty string"); + let encrypted = if normalized_key.eq_ignore_ascii_case("smtp_password") { + let host = state + .read_system_config_json_value("smtp_host") + .await? + .and_then(|value| system_config_string(Some(&value))); + let port = state + .read_system_config_json_value("smtp_port") + .await? + .map(|value| crate::email_delivery::system_config_u16(Some(&value), 587)) + .unwrap_or(587); + let user = state + .read_system_config_json_value("smtp_user") + .await? + .and_then(|value| system_config_string(Some(&value))); + let use_tls = state + .read_system_config_json_value("smtp_use_tls") + .await? + .map(|value| system_config_bool(Some(&value), true)) + .unwrap_or(true); + let use_ssl = state + .read_system_config_json_value("smtp_use_ssl") + .await? + .map(|value| system_config_bool(Some(&value), false)) + .unwrap_or(false); + let Some(binding) = host.as_deref().and_then(|host| { + smtp_password_binding(host, port, user.as_deref(), use_tls, use_ssl) + }) else { + return Ok(Err(( + http::StatusCode::BAD_REQUEST, + json!({ + "detail": "保存 SMTP 密码前必须先配置有效的 smtp_host 和 smtp_user" + }), + ))); + }; + encrypt_smtp_password(state.app(), &binding, plaintext) + } else if normalized_key.eq_ignore_ascii_case("module.bark_push.device_key") { + let server_url = state + .read_system_config_json_value("module.bark_push.server_url") + .await? + .and_then(|value| system_config_string(Some(&value))) + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| "https://api.day.app".to_string()); + let Some(binding) = bark_device_key_binding(&server_url) else { + return Ok(Err(( + http::StatusCode::BAD_REQUEST, + json!({ + "detail": "保存 Bark Device Key 前必须先配置有效的 module.bark_push.server_url" + }), + ))); + }; + encrypt_bark_device_key(state.app(), &binding, plaintext) + } else { + state.encrypt_system_config_secret(&normalized_key, plaintext) + }; + let Some(encrypted) = encrypted else { return Ok(Err(( http::StatusCode::SERVICE_UNAVAILABLE, json!({ "detail": "系统配置写入需要可用的加密密钥" }), ))); }; - let plaintext = value.as_str().unwrap(); - value = json!(encrypt_python_fernet_plaintext(encryption_key, plaintext) - .map_err(|err| GatewayError::Internal(err.to_string()))?); + value = json!(encrypted); } let updated = state diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/export/mod.rs b/apps/aether-gateway/src/handlers/admin/system/shared/export/mod.rs index 44dd3c8a9..a81b2d3c9 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/export/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/export/mod.rs @@ -3,6 +3,7 @@ mod support; pub(crate) use self::providers::build_admin_system_export_providers_payload; pub(crate) use self::support::{ - decrypt_admin_system_export_secret, ADMIN_SYSTEM_CONFIG_EXPORT_VERSION, - ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, + decrypt_admin_system_export_secret, project_admin_system_export_json, + project_admin_system_export_optional_url, project_admin_system_export_url, + ADMIN_SYSTEM_CONFIG_EXPORT_VERSION, ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, }; diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/export/providers.rs b/apps/aether-gateway/src/handlers/admin/system/shared/export/providers.rs index 6d77ce643..f932ed84d 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/export/providers.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/export/providers.rs @@ -1,22 +1,44 @@ use super::support::{ - collect_admin_system_export_provider_endpoint_formats, - decrypt_admin_system_export_provider_config, decrypt_admin_system_export_secret, - resolve_admin_system_export_key_api_formats, ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, + collect_admin_system_export_provider_endpoint_formats, project_admin_system_export_body_rules, + project_admin_system_export_header_rules, project_admin_system_export_json, + project_admin_system_export_optional_url, project_admin_system_export_provider_config, + project_admin_system_export_proxy, project_admin_system_export_url, + resolve_admin_system_export_key_api_formats, +}; +use crate::handlers::admin::admin_provider_ops_credential_snapshot; +use crate::handlers::admin::request::{ + AdminAppState, SystemExportMode, ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED, }; -use crate::handlers::admin::request::AdminAppState; use crate::GatewayError; use aether_admin::system::{ AdminSystemConfigEndpoint, AdminSystemConfigProvider, AdminSystemConfigProviderKey, AdminSystemConfigProviderModel, }; -use aether_data_contracts::repository::global_models::AdminProviderModelListQuery; use std::collections::BTreeMap; pub(crate) async fn build_admin_system_export_providers_payload( state: &AdminAppState<'_>, global_model_name_by_id: &BTreeMap, + mode: SystemExportMode, ) -> Result, GatewayError> { - let providers = state.list_provider_catalog_providers(false).await?; + let mut providers = state.list_provider_catalog_providers(false).await?; + let mut provider_ops_credentials = BTreeMap::new(); + if mode.credentials_are_exported() { + for provider in &mut providers { + let has_provider_ops = provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|config| config.get("provider_ops")) + .is_some(); + if !has_provider_ops { + continue; + } + let snapshot = admin_provider_ops_credential_snapshot(state, provider).await?; + provider_ops_credentials.insert(provider.id.clone(), snapshot.credentials); + *provider = snapshot.provider; + } + } let provider_ids = providers .iter() .map(|provider| provider.id.clone()) @@ -46,188 +68,232 @@ pub(crate) async fn build_admin_system_export_providers_payload( let mut provider_models_by_provider = BTreeMap::>::new(); for provider in &providers { let models = state - .list_admin_provider_models(&AdminProviderModelListQuery { - provider_id: provider.id.clone(), - is_active: None, - offset: 0, - limit: ADMIN_SYSTEM_EXPORT_PAGE_LIMIT, - }) + .list_all_admin_provider_models_for_system_transfer(&provider.id) .await?; provider_models_by_provider.insert(provider.id.clone(), models); } - Ok(providers + providers .iter() - .map(|provider| { - let endpoints = endpoints_by_provider - .remove(&provider.id) - .unwrap_or_default(); - let provider_endpoint_formats = - collect_admin_system_export_provider_endpoint_formats(&endpoints); - let endpoints_data = endpoints - .iter() - .map(|endpoint| AdminSystemConfigEndpoint { - api_format: endpoint.api_format.clone(), - base_url: endpoint.base_url.clone(), - header_rules: endpoint.header_rules.clone(), - body_rules: endpoint.body_rules.clone(), - max_retries: endpoint.max_retries, - is_active: endpoint.is_active, - custom_path: endpoint.custom_path.clone(), - config: endpoint.config.clone(), - format_acceptance_config: endpoint.format_acceptance_config.clone(), - proxy: endpoint.proxy.clone(), - }) - .collect::>(); + .map( + |provider| -> Result { + let endpoints = endpoints_by_provider + .remove(&provider.id) + .unwrap_or_default(); + let provider_endpoint_formats = + collect_admin_system_export_provider_endpoint_formats(&endpoints); + let endpoints_data = endpoints + .iter() + .map(|endpoint| AdminSystemConfigEndpoint { + api_format: endpoint.api_format.clone(), + base_url: project_admin_system_export_url(mode, &endpoint.base_url), + header_rules: project_admin_system_export_header_rules( + mode, + endpoint.header_rules.as_ref(), + ), + body_rules: project_admin_system_export_body_rules( + mode, + endpoint.body_rules.as_ref(), + ), + max_retries: endpoint.max_retries, + is_active: endpoint.is_active, + custom_path: endpoint.custom_path.clone(), + config: project_admin_system_export_json(mode, endpoint.config.as_ref()), + format_acceptance_config: project_admin_system_export_json( + mode, + endpoint.format_acceptance_config.as_ref(), + ), + proxy: project_admin_system_export_proxy(mode, endpoint.proxy.as_ref()), + }) + .collect::>(); - let mut keys = keys_by_provider.remove(&provider.id).unwrap_or_default(); - keys.sort_by(|left, right| { - left.internal_priority - .cmp(&right.internal_priority) - .then( - left.created_at_unix_ms - .unwrap_or(0) - .cmp(&right.created_at_unix_ms.unwrap_or(0)), + let mut keys = keys_by_provider.remove(&provider.id).unwrap_or_default(); + keys.sort_by(|left, right| { + left.internal_priority + .cmp(&right.internal_priority) + .then( + left.created_at_unix_ms + .unwrap_or(0) + .cmp(&right.created_at_unix_ms.unwrap_or(0)), + ) + .then(left.id.cmp(&right.id)) + }); + let keys_data = keys + .iter() + .map( + |key| -> Result { + let api_formats = resolve_admin_system_export_key_api_formats( + key.api_formats.as_ref(), + &provider_endpoint_formats, + ); + let auth_config = if mode.credentials_are_exported() { + state + .app() + .decrypt_provider_catalog_key_auth_config(key)? + .map(serde_json::Value::String) + } else { + None + }; + let api_key = if mode.credentials_are_exported() { + state.app().decrypt_provider_catalog_key_api_key(key)? + } else { + None + }; + Ok(AdminSystemConfigProviderKey { + api_key, + auth_type: Some(key.auth_type.clone()), + auth_config, + name: Some(key.name.clone()), + note: key.note.clone(), + api_formats: Some(api_formats.clone()), + supported_endpoints: Some(api_formats), + rate_multipliers: project_admin_system_export_json( + mode, + key.rate_multipliers.as_ref(), + ), + internal_priority: Some(key.internal_priority), + global_priority_by_format: project_admin_system_export_json( + mode, + key.global_priority_by_format.as_ref(), + ), + auth_type_by_format: project_admin_system_export_json( + mode, + key.auth_type_by_format.as_ref(), + ), + allow_auth_channel_mismatch_formats: key + .allow_auth_channel_mismatch_formats + .as_ref() + .and_then(serde_json::Value::as_array) + .map(|items| { + items + .iter() + .filter_map(serde_json::Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }), + rpm_limit: key.rpm_limit, + allowed_models: key.allowed_models.as_ref().and_then(|value| { + value.as_array().map(|items| { + items + .iter() + .filter_map(serde_json::Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }) + }), + capabilities: project_admin_system_export_json( + mode, + key.capabilities.as_ref(), + ), + cache_ttl_minutes: Some(key.cache_ttl_minutes), + max_probe_interval_minutes: Some(key.max_probe_interval_minutes), + auto_fetch_models: Some(key.auto_fetch_models), + locked_models: key.locked_models.as_ref().and_then(|value| { + value.as_array().map(|items| { + items + .iter() + .filter_map(serde_json::Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }) + }), + model_include_patterns: key + .model_include_patterns + .as_ref() + .and_then(|value| { + value.as_array().map(|items| { + items + .iter() + .filter_map(serde_json::Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }) + }), + model_exclude_patterns: key + .model_exclude_patterns + .as_ref() + .and_then(|value| { + value.as_array().map(|items| { + items + .iter() + .filter_map(serde_json::Value::as_str) + .map(ToOwned::to_owned) + .collect::>() + }) + }), + is_active: mode.preserves_active_state() && key.is_active, + proxy: project_admin_system_export_proxy(mode, key.proxy.as_ref()), + fingerprint: project_admin_system_export_json( + mode, + key.fingerprint.as_ref(), + ), + credential_state: (!mode.credentials_are_exported()).then(|| { + ADMIN_SYSTEM_EXPORT_CREDENTIALS_NOT_EXPORTED.to_string() + }), + }) + }, ) - .then(left.id.cmp(&right.id)) - }); - let keys_data = keys - .iter() - .map(|key| { - let api_formats = resolve_admin_system_export_key_api_formats( - key.api_formats.as_ref(), - &provider_endpoint_formats, - ); - let auth_config = key - .encrypted_auth_config - .as_deref() - .and_then(|ciphertext| { - decrypt_admin_system_export_secret(state, ciphertext) - }) - .map(serde_json::Value::String); - AdminSystemConfigProviderKey { - api_key: key.encrypted_api_key.as_deref().map(|ciphertext| { - decrypt_admin_system_export_secret(state, ciphertext) - .unwrap_or_default() - }), - auth_type: Some(key.auth_type.clone()), - auth_config, - name: Some(key.name.clone()), - note: key.note.clone(), - api_formats: Some(api_formats.clone()), - supported_endpoints: Some(api_formats), - rate_multipliers: key.rate_multipliers.clone(), - internal_priority: Some(key.internal_priority), - global_priority_by_format: key.global_priority_by_format.clone(), - auth_type_by_format: key.auth_type_by_format.clone(), - allow_auth_channel_mismatch_formats: key - .allow_auth_channel_mismatch_formats - .as_ref() - .and_then(serde_json::Value::as_array) - .map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .collect::>() - }), - rpm_limit: key.rpm_limit, - allowed_models: key.allowed_models.as_ref().and_then(|value| { - value.as_array().map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .collect::>() - }) - }), - capabilities: key.capabilities.clone(), - cache_ttl_minutes: Some(key.cache_ttl_minutes), - max_probe_interval_minutes: Some(key.max_probe_interval_minutes), - auto_fetch_models: Some(key.auto_fetch_models), - locked_models: key.locked_models.as_ref().and_then(|value| { - value.as_array().map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .collect::>() - }) - }), - model_include_patterns: key.model_include_patterns.as_ref().and_then( - |value| { - value.as_array().map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .collect::>() - }) - }, - ), - model_exclude_patterns: key.model_exclude_patterns.as_ref().and_then( - |value| { - value.as_array().map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .collect::>() - }) - }, - ), - is_active: key.is_active, - proxy: key.proxy.clone(), - fingerprint: key.fingerprint.clone(), - } - }) - .collect::>(); + .collect::, GatewayError>>()?; - let models_data = provider_models_by_provider - .remove(&provider.id) - .unwrap_or_default() - .into_iter() - .map(|model| AdminSystemConfigProviderModel { - global_model_name: global_model_name_by_id.get(&model.global_model_id).cloned(), - provider_model_name: model.provider_model_name, - provider_model_mappings: model.provider_model_mappings, - price_per_request: model.price_per_request, - tiered_pricing: model.tiered_pricing, - supports_vision: model.supports_vision, - supports_function_calling: model.supports_function_calling, - supports_streaming: model.supports_streaming, - supports_extended_thinking: model.supports_extended_thinking, - supports_image_generation: model.supports_image_generation, - is_active: model.is_active, - config: model.config, - }) - .collect::>(); + let models_data = provider_models_by_provider + .remove(&provider.id) + .unwrap_or_default() + .into_iter() + .map(|model| AdminSystemConfigProviderModel { + global_model_name: global_model_name_by_id + .get(&model.global_model_id) + .cloned(), + provider_model_name: model.provider_model_name, + provider_model_mappings: project_admin_system_export_json( + mode, + model.provider_model_mappings.as_ref(), + ), + price_per_request: model.price_per_request, + tiered_pricing: project_admin_system_export_json( + mode, + model.tiered_pricing.as_ref(), + ), + supports_vision: model.supports_vision, + supports_function_calling: model.supports_function_calling, + supports_streaming: model.supports_streaming, + supports_extended_thinking: model.supports_extended_thinking, + supports_image_generation: model.supports_image_generation, + is_active: model.is_active, + config: project_admin_system_export_json(mode, model.config.as_ref()), + }) + .collect::>(); - AdminSystemConfigProvider { - name: provider.name.clone(), - description: provider.description.clone(), - website: provider.website.clone(), - provider_type: Some(provider.provider_type.clone()), - billing_type: provider.billing_type.clone(), - monthly_quota_usd: provider.monthly_quota_usd, - quota_reset_day: provider.quota_reset_day, - provider_priority: Some(provider.provider_priority), - keep_priority_on_conversion: Some(provider.keep_priority_on_conversion), - enable_format_conversion: Some(provider.enable_format_conversion), - is_active: provider.is_active, - concurrent_limit: provider.concurrent_limit, - max_retries: provider.max_retries, - stream_first_byte_timeout: provider.stream_first_byte_timeout_secs, - request_timeout: provider.request_timeout_secs, - proxy: provider.proxy.clone(), - config: decrypt_admin_system_export_provider_config( - state, - provider.config.as_ref(), - ), - endpoints: endpoints_data, - api_keys: keys_data, - models: models_data, - } - }) - .collect::>()) + Ok(AdminSystemConfigProvider { + name: provider.name.clone(), + description: provider.description.clone(), + website: project_admin_system_export_optional_url( + mode, + provider.website.as_deref(), + ), + provider_type: Some(provider.provider_type.clone()), + billing_type: provider.billing_type.clone(), + monthly_quota_usd: provider.monthly_quota_usd, + quota_reset_day: provider.quota_reset_day, + provider_priority: Some(provider.provider_priority), + keep_priority_on_conversion: Some(provider.keep_priority_on_conversion), + enable_format_conversion: Some(provider.enable_format_conversion), + is_active: provider.is_active, + concurrent_limit: provider.concurrent_limit, + max_retries: provider.max_retries, + stream_first_byte_timeout: provider.stream_first_byte_timeout_secs, + request_timeout: provider.request_timeout_secs, + proxy: project_admin_system_export_proxy(mode, provider.proxy.as_ref()), + config: project_admin_system_export_provider_config( + state, + mode, + provider.config.as_ref(), + provider_ops_credentials.get(&provider.id), + )?, + endpoints: endpoints_data, + api_keys: keys_data, + models: models_data, + }) + }, + ) + .collect::, GatewayError>>() } diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/export/support.rs b/apps/aether-gateway/src/handlers/admin/system/shared/export/support.rs index e3cc96889..d711d086d 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/export/support.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/export/support.rs @@ -1,9 +1,15 @@ use super::super::configs::is_sensitive_admin_system_config_key; use crate::api::ai::admin_endpoint_signature_parts; -use crate::handlers::admin::request::AdminAppState; -use crate::handlers::shared::decrypt_catalog_secret_with_fallbacks; +use crate::handlers::admin::request::{AdminAppState, SystemExportMode}; +use crate::handlers::shared::{ + decrypt_catalog_secret_with_fallbacks, PROVIDER_OPS_PERSISTENT_SECRET_FIELDS, + PROVIDER_OPS_TRANSIENT_METADATA_FIELDS, PROVIDER_OPS_TRANSIENT_SECRET_FIELDS, +}; +use aether_admin::provider::redaction::{ + admin_secret_safe_body_rules, admin_secret_safe_header_rules, admin_secret_safe_json, + admin_secret_safe_proxy, admin_secret_safe_url, +}; pub(crate) use aether_admin::system::ADMIN_SYSTEM_CONFIG_EXPORT_VERSION; -use aether_admin::system::ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint; pub(crate) const ADMIN_SYSTEM_EXPORT_PAGE_LIMIT: usize = 10_000; @@ -47,11 +53,18 @@ pub(super) fn collect_admin_system_export_provider_endpoint_formats( ) } -pub(super) fn decrypt_admin_system_export_provider_config( - state: &AdminAppState<'_>, +pub(super) fn project_admin_system_export_provider_config( + _state: &AdminAppState<'_>, + mode: SystemExportMode, config: Option<&serde_json::Value>, -) -> Option { - let mut decrypted = config.cloned()?; + provider_ops_plaintext_credentials: Option<&serde_json::Map>, +) -> Result, crate::GatewayError> { + if !mode.credentials_are_exported() { + return Ok(project_admin_system_export_json(mode, config)); + } + let Some(mut decrypted) = config.cloned() else { + return Ok(None); + }; let Some(credentials) = decrypted .get_mut("provider_ops") .and_then(serde_json::Value::as_object_mut) @@ -60,17 +73,96 @@ pub(super) fn decrypt_admin_system_export_provider_config( .and_then(|connector| connector.get_mut("credentials")) .and_then(serde_json::Value::as_object_mut) else { - return Some(decrypted); + return Ok(Some(decrypted)); }; - for field in ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS { - let Some(serde_json::Value::String(ciphertext)) = credentials.get(*field).cloned() else { - continue; - }; - if let Some(plaintext) = decrypt_admin_system_export_secret(state, &ciphertext) { - credentials.insert((*field).to_string(), serde_json::Value::String(plaintext)); + for field in PROVIDER_OPS_TRANSIENT_SECRET_FIELDS + .iter() + .chain(PROVIDER_OPS_TRANSIENT_METADATA_FIELDS) + { + credentials.remove(*field); + } + let plaintext_credentials = provider_ops_plaintext_credentials.ok_or_else(|| { + crate::GatewayError::Internal( + "RecoveryBackup 缺少已认证的 Provider Ops 凭据快照".to_string(), + ) + })?; + for field in PROVIDER_OPS_PERSISTENT_SECRET_FIELDS { + if let Some(plaintext) = plaintext_credentials.get(*field) { + credentials.insert((*field).to_string(), plaintext.clone()); } } - Some(decrypted) + Ok(Some(decrypted)) +} + +pub(crate) fn project_admin_system_export_json( + mode: SystemExportMode, + value: Option<&serde_json::Value>, +) -> Option { + value.map(|value| { + if mode.credentials_are_exported() { + value.clone() + } else { + admin_secret_safe_json(Some(value)) + } + }) +} + +pub(super) fn project_admin_system_export_header_rules( + mode: SystemExportMode, + value: Option<&serde_json::Value>, +) -> Option { + value.map(|value| { + if mode.credentials_are_exported() { + value.clone() + } else { + admin_secret_safe_header_rules(Some(value)) + } + }) +} + +pub(super) fn project_admin_system_export_body_rules( + mode: SystemExportMode, + value: Option<&serde_json::Value>, +) -> Option { + value.map(|value| { + if mode.credentials_are_exported() { + value.clone() + } else { + admin_secret_safe_body_rules(Some(value)) + } + }) +} + +pub(super) fn project_admin_system_export_proxy( + mode: SystemExportMode, + value: Option<&serde_json::Value>, +) -> Option { + value.map(|value| { + if mode.credentials_are_exported() { + value.clone() + } else { + admin_secret_safe_proxy(Some(value)) + } + }) +} + +pub(crate) fn project_admin_system_export_optional_url( + mode: SystemExportMode, + value: Option<&str>, +) -> Option { + value.and_then(|value| { + if mode.credentials_are_exported() { + Some(value.to_string()) + } else { + admin_secret_safe_url(Some(value)) + .as_str() + .map(ToOwned::to_owned) + } + }) +} + +pub(crate) fn project_admin_system_export_url(mode: SystemExportMode, value: &str) -> String { + project_admin_system_export_optional_url(mode, Some(value)).unwrap_or_default() } diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/modules.rs b/apps/aether-gateway/src/handlers/admin/system/shared/modules.rs index a925a1765..a5eb4b28b 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/modules.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/modules.rs @@ -248,7 +248,7 @@ pub(crate) fn oauth_module_config_is_valid( pub(crate) fn ldap_module_config_is_valid( config: Option<&aether_data::repository::auth_modules::StoredLdapModuleConfig>, ) -> bool { - admin_system_kernel::ldap_module_config_is_valid(config) + crate::handlers::shared::ldap_module_config_is_valid(config) } pub(crate) async fn build_admin_module_runtime_state( diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/settings.rs b/apps/aether-gateway/src/handlers/admin/system/shared/settings.rs index b1845ea91..7484d6b20 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/settings.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/settings.rs @@ -28,6 +28,7 @@ use std::time::Duration; const AETHER_RELEASES_API_URL: &str = "https://api.github.com/repos/fawney19/Aether/releases?per_page=20"; const AETHER_RELEASE_TAG_URL_BASE: &str = "https://github.com/fawney19/Aether/releases/tag"; +const MAX_GITHUB_RELEASES_RESPONSE_BYTES: usize = 4 * 1024 * 1024; const SOURCE_BUILD_UPDATE_BLOCKER: &str = "当前为源码构建,请使用 git pull 后重新编译。"; const SOURCE_BUILD_RELEASE_BLOCKER: &str = "当前为源码构建,请手动切换到对应标签后重新编译。"; @@ -348,18 +349,22 @@ async fn fetch_github_releases_with_client( rate_limited: false, })?; let status = response.status(); + let body = + aether_http::read_response_bytes_with_limit(response, MAX_GITHUB_RELEASES_RESPONSE_BYTES) + .await + .map_err(|err| GitHubReleaseFetchError { + message: format!("读取 GitHub Releases 响应失败: {err}"), + rate_limited: false, + })?; if !status.is_success() { - let body = response.text().await.unwrap_or_default(); + let body = String::from_utf8_lossy(&body); return Err(github_release_response_error(status, &body)); } - response - .json() - .await - .map_err(|err| GitHubReleaseFetchError { - message: format!("解析 GitHub Releases 失败: {err}"), - rate_limited: false, - }) + serde_json::from_slice(&body).map_err(|err| GitHubReleaseFetchError { + message: format!("解析 GitHub Releases 失败: {err}"), + rate_limited: false, + }) } fn github_release_response_error( diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/smtp.rs b/apps/aether-gateway/src/handlers/admin/system/shared/smtp.rs index bb0256d90..602cfb0e5 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/smtp.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/smtp.rs @@ -1,12 +1,15 @@ use crate::email_delivery::{probe_smtp_connection, system_config_u16, SmtpDeliveryConfig}; use crate::handlers::admin::request::AdminAppState; -use crate::handlers::shared::{system_config_bool, system_config_string}; +use crate::handlers::shared::{ + decrypt_or_migrate_smtp_password, smtp_password_binding, system_config_bool, + system_config_string, +}; use crate::GatewayError; use axum::body::Bytes; use serde::Deserialize; use serde_json::json; -#[derive(Debug, Default, Deserialize)] +#[derive(Default, Deserialize)] struct AdminSmtpTestRequest { smtp_host: Option, smtp_port: Option, @@ -18,7 +21,7 @@ struct AdminSmtpTestRequest { smtp_from_name: Option, } -#[derive(Debug, Clone)] +#[derive(Clone)] struct ResolvedSmtpConfig { host: Option, port: u16, @@ -75,43 +78,70 @@ async fn resolve_admin_smtp_config( .read_system_config_json_value("smtp_from_name") .await?; - let stored_password = system_config_string(smtp_password.as_ref()).map(|value| { - state - .decrypt_catalog_secret_with_fallbacks(&value) - .unwrap_or(value) - }); + let host = request + .smtp_host + .as_ref() + .and_then(|value| system_config_string(Some(value))) + .or_else(|| system_config_string(smtp_host.as_ref())); + let port = request + .smtp_port + .as_ref() + .map(|value| system_config_u16(Some(value), 587)) + .unwrap_or_else(|| system_config_u16(smtp_port.as_ref(), 587)); + let user = request + .smtp_user + .as_ref() + .and_then(|value| system_config_string(Some(value))) + .or_else(|| system_config_string(smtp_user.as_ref())); + let use_tls = request + .smtp_use_tls + .as_ref() + .map(|value| system_config_bool(Some(value), true)) + .unwrap_or_else(|| system_config_bool(smtp_use_tls.as_ref(), true)); + let use_ssl = request + .smtp_use_ssl + .as_ref() + .map(|value| system_config_bool(Some(value), false)) + .unwrap_or_else(|| system_config_bool(smtp_use_ssl.as_ref(), false)); + let requested_password = request + .smtp_password + .as_ref() + .and_then(|value| system_config_string(Some(value))); + let password = if requested_password.is_some() { + requested_password + } else { + let saved_binding = system_config_string(smtp_host.as_ref()).and_then(|saved_host| { + smtp_password_binding( + &saved_host, + system_config_u16(smtp_port.as_ref(), 587), + system_config_string(smtp_user.as_ref()).as_deref(), + system_config_bool(smtp_use_tls.as_ref(), true), + system_config_bool(smtp_use_ssl.as_ref(), false), + ) + }); + let effective_binding = host.as_deref().and_then(|effective_host| { + smtp_password_binding(effective_host, port, user.as_deref(), use_tls, use_ssl) + }); + match (saved_binding.as_ref(), effective_binding.as_ref()) { + (Some(saved), Some(effective)) if saved == effective => { + match system_config_string(smtp_password.as_ref()) { + Some(value) => Some( + decrypt_or_migrate_smtp_password(state.as_ref(), effective, value).await?, + ), + None => None, + } + } + _ => None, + } + }; Ok(ResolvedSmtpConfig { - host: request - .smtp_host - .as_ref() - .and_then(|value| system_config_string(Some(value))) - .or_else(|| system_config_string(smtp_host.as_ref())), - port: request - .smtp_port - .as_ref() - .map(|value| system_config_u16(Some(value), 587)) - .unwrap_or_else(|| system_config_u16(smtp_port.as_ref(), 587)), - user: request - .smtp_user - .as_ref() - .and_then(|value| system_config_string(Some(value))) - .or_else(|| system_config_string(smtp_user.as_ref())), - password: request - .smtp_password - .as_ref() - .and_then(|value| system_config_string(Some(value))) - .or(stored_password), - use_tls: request - .smtp_use_tls - .as_ref() - .map(|value| system_config_bool(Some(value), true)) - .unwrap_or_else(|| system_config_bool(smtp_use_tls.as_ref(), true)), - use_ssl: request - .smtp_use_ssl - .as_ref() - .map(|value| system_config_bool(Some(value), false)) - .unwrap_or_else(|| system_config_bool(smtp_use_ssl.as_ref(), false)), + host, + port, + user, + password, + use_tls, + use_ssl, from_email: request .smtp_from_email .as_ref() diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/update.rs b/apps/aether-gateway/src/handlers/admin/system/shared/update.rs index b13f6270b..8adb56f4f 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/update.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/update.rs @@ -1,9 +1,14 @@ -use crate::handlers::admin::system::shared::update_client::build_update_http_client; +use crate::handlers::admin::system::shared::update_client::{ + build_update_http_client, is_trusted_update_url, +}; use crate::GatewayError; use axum::http; use futures_util::StreamExt; use serde_json::json; use sha2::{Digest, Sha256}; +use std::collections::HashSet; +#[cfg(unix)] +use std::io::{Read, Write}; use std::path::Component; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicBool, Ordering}; @@ -63,18 +68,18 @@ fn set_update_task_download_progress(label: &str, downloaded_bytes: u64, total_b fn set_update_task_failed(error: String) { if let Ok(mut guard) = update_task_status_lock().lock() { guard.phase = "failed"; - guard.error = Some(error); + guard.error = Some(safe_system_update_error(&error).to_string()); } } -fn set_update_task_output(output: String) { +fn set_update_task_output(_output: String) { if let Ok(mut guard) = update_task_status_lock().lock() { - guard.output = Some(output); + guard.output = Some("Update package prepared".to_string()); } } pub(crate) fn read_update_task_status() -> SystemUpdateTaskStatus { - update_task_status_lock() + let mut status = update_task_status_lock() .lock() .map(|guard| guard.clone()) .unwrap_or(SystemUpdateTaskStatus { @@ -85,10 +90,20 @@ pub(crate) fn read_update_task_status() -> SystemUpdateTaskStatus { downloaded_bytes: None, total_bytes: None, progress_percent: None, - }) + }); + status.error = status + .error + .as_deref() + .map(|error| safe_system_update_error(error).to_string()); + status.output = status + .output + .as_ref() + .map(|_| "System update step completed".to_string()); + status } static PREPARED_VERSION: std::sync::OnceLock>> = std::sync::OnceLock::new(); +static UPDATE_HISTORY_LOCK: std::sync::OnceLock> = std::sync::OnceLock::new(); fn prepared_version_lock() -> &'static Mutex> { PREPARED_VERSION.get_or_init(|| Mutex::new(None)) @@ -100,6 +115,12 @@ fn set_prepared_version(version: String) { } } +fn clear_prepared_version() { + if let Ok(mut guard) = prepared_version_lock().lock() { + *guard = None; + } +} + pub(crate) fn get_prepared_version() -> Option { prepared_version_lock().lock().ok()?.clone() } @@ -107,10 +128,14 @@ pub(crate) fn get_prepared_version() -> Option { const UPDATE_HISTORY_FILENAME: &str = ".aether-update-history.json"; const PREVIOUS_RELEASE_FILENAME: &str = ".aether-previous-release"; const MAX_HISTORY_ENTRIES: usize = 50; +const MAX_UPDATE_HISTORY_BYTES: usize = 256 * 1024; +const MAX_PREVIOUS_RELEASE_BYTES: usize = 256; const RESTART_EXIT_CODE: i32 = 75; const MAX_RELEASE_DOWNLOAD_BYTES: u64 = 512 * 1024 * 1024; const MAX_SHA256SUMS_DOWNLOAD_BYTES: u64 = 1024 * 1024; const MAX_EXTRACTED_RELEASE_BYTES: u64 = 1024 * 1024 * 1024; +const MAX_RELEASE_ARCHIVE_ENTRIES: usize = 100_000; +const MAX_RELEASE_ARCHIVE_PATH_DEPTH: usize = 64; const DEFAULT_UPDATE_DOWNLOAD_TIMEOUT_SECS: u64 = 600; const DEFAULT_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS: u64 = 30; const SOURCE_BUILD_UPDATE_BLOCKER: &str = "当前为源码构建,请使用 git pull 后重新编译。"; @@ -120,6 +145,8 @@ const MANUAL_UPDATE_BLOCKER: &str = "当前部署策略不支持在线自更新,请手动下载 Release 或使用安装脚本更新。"; const MULTI_NODE_UPDATE_BLOCKER: &str = "多节点部署不支持在管理后台更新单个节点,请使用镜像滚动更新或外部发布编排。"; +const STORAGE_UPDATE_BLOCKER: &str = + "当前安装目录未提供安全的自更新写权限,请使用安装脚本或受限的系统更新服务。"; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum UpdateStrategy { @@ -221,22 +248,41 @@ fn release_dir_for_version(version: &str) -> Result { Ok(releases_base_dir().join(safe_release_name(version)?)) } -fn current_symlink_path() -> PathBuf { - aether_base_dir().join("current") +fn current_release_name() -> Option { + current_release_name_at(&aether_base_dir()) } -fn current_release_name() -> Option { - std::fs::read_link(current_symlink_path()) +fn current_release_name_at(base_dir: &Path) -> Option { + let releases = std::fs::canonicalize(base_dir.join("releases")).ok()?; + let current = base_dir.join("current"); + if !std::fs::symlink_metadata(¤t) .ok()? - .file_name() - .and_then(|name| name.to_str()) - .map(str::to_string) + .file_type() + .is_symlink() + { + return None; + } + let target = std::fs::canonicalize(current).ok()?; + let relative = target.strip_prefix(releases).ok()?; + let mut components = relative.components(); + let Component::Normal(name) = components.next()? else { + return None; + }; + if components.next().is_some() { + return None; + } + let name = name.to_str()?; + safe_release_name(name).ok() } fn update_history_path() -> PathBuf { aether_base_dir().join(UPDATE_HISTORY_FILENAME) } +fn update_history_lock() -> &'static Mutex<()> { + UPDATE_HISTORY_LOCK.get_or_init(|| Mutex::new(())) +} + fn append_update_history( operation: &str, success: bool, @@ -244,40 +290,505 @@ fn append_update_history( output: Option<&str>, ) { let path = update_history_path(); + let _guard = update_history_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + append_update_history_at_path(&path, operation, success, error, output); +} +fn append_update_history_at_path( + path: &Path, + operation: &str, + success: bool, + error: Option<&str>, + output: Option<&str>, +) { let entry = UpdateHistoryEntry { timestamp: chrono::Utc::now().to_rfc3339(), - operation: operation.to_string(), + operation: safe_system_update_operation(operation).to_string(), success, - error: error.map(|s| s.to_string()), - output_tail: output.map(|s| { - let lines: Vec<&str> = s.lines().collect(); - let start = lines.len().saturating_sub(20); - lines[start..].join("\n") - }), + error: error.map(|error| safe_system_update_error(error).to_string()), + output_tail: output.map(|_| safe_system_update_output(operation).to_string()), }; - let mut entries: Vec = std::fs::read_to_string(&path) - .ok() - .and_then(|content| serde_json::from_str(&content).ok()) - .unwrap_or_default(); + let (mut entries, _) = load_and_sanitize_update_history(path); entries.push(entry); - if entries.len() > MAX_HISTORY_ENTRIES { - entries.drain(..entries.len() - MAX_HISTORY_ENTRIES); - } - - if let Ok(json) = serde_json::to_string_pretty(&entries) { - let _ = std::fs::write(&path, json); - } + sanitize_update_history_entries(&mut entries); + persist_update_history(path, &entries); } pub(crate) fn read_update_history() -> Vec { let path = update_history_path(); - std::fs::read_to_string(&path) - .ok() - .and_then(|content| serde_json::from_str(&content).ok()) - .unwrap_or_default() + let _guard = update_history_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + read_update_history_at_path(&path) +} + +fn read_update_history_at_path(path: &Path) -> Vec { + let (entries, changed) = load_and_sanitize_update_history(path); + if changed { + persist_update_history(path, &entries); + } + entries +} + +fn load_and_sanitize_update_history(path: &Path) -> (Vec, bool) { + let Ok(Some(content)) = read_update_metadata_file(path, MAX_UPDATE_HISTORY_BYTES) else { + return (Vec::new(), false); + }; + let Ok(content) = String::from_utf8(content) else { + return (Vec::new(), true); + }; + let Ok(mut entries) = serde_json::from_str::>(&content) else { + return (Vec::new(), !content.trim().is_empty()); + }; + let changed = sanitize_update_history_entries(&mut entries); + (entries, changed) +} + +fn sanitize_update_history_entries(entries: &mut Vec) -> bool { + // Deliberately use a non-short-circuiting fold: every historical entry + // must be sanitized even after one entry changes. + #[allow(clippy::unnecessary_fold)] + let mut changed = entries.iter_mut().fold(false, |changed, entry| { + sanitize_update_history_entry(entry) || changed + }); + if entries.len() > MAX_HISTORY_ENTRIES { + entries.drain(..entries.len() - MAX_HISTORY_ENTRIES); + changed = true; + } + changed +} + +fn sanitize_update_history_entry(entry: &mut UpdateHistoryEntry) -> bool { + let mut changed = false; + + if chrono::DateTime::parse_from_rfc3339(&entry.timestamp).is_err() { + entry.timestamp = "1970-01-01T00:00:00Z".to_string(); + changed = true; + } + + let operation = safe_system_update_operation(&entry.operation).to_string(); + if entry.operation != operation { + entry.operation = operation; + changed = true; + } + + let error = entry + .error + .as_deref() + .map(|error| safe_system_update_error(error).to_string()); + if entry.error != error { + entry.error = error; + changed = true; + } + + let output_tail = entry + .output_tail + .as_ref() + .map(|_| safe_system_update_output(&entry.operation).to_string()); + if entry.output_tail != output_tail { + entry.output_tail = output_tail; + changed = true; + } + + changed +} + +fn persist_update_history(path: &Path, entries: &[UpdateHistoryEntry]) { + let result = serde_json::to_vec_pretty(entries) + .map_err(|err| format!("序列化更新历史失败: {err}")) + .and_then(|json| write_update_metadata_atomic(path, &json)); + if let Err(err) = result { + tracing::warn!(error = %safe_system_update_error(&err), "failed to persist update history"); + } +} + +fn read_update_metadata_file(path: &Path, max_bytes: usize) -> Result>, String> { + #[cfg(unix)] + { + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::fs::MetadataExt; + + let (parent, file_name) = open_real_update_parent(path)?; + let file_name = unix_update_path_component(&file_name, "更新元数据文件名")?; + // SAFETY: parent is a live directory descriptor, file_name is NUL-terminated, and a + // successful descriptor is transferred immediately into File. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + file_name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_NOFOLLOW, + ) + }; + if descriptor < 0 { + let error = std::io::Error::last_os_error(); + if error.raw_os_error() == Some(libc::ENOENT) { + return Ok(None); + } + return Err(format!("安全打开更新元数据失败: {error}")); + } + // SAFETY: openat returned a new owned descriptor. + let mut file = unsafe { std::fs::File::from_raw_fd(descriptor) }; + let metadata = file + .metadata() + .map_err(|err| format!("读取更新元数据属性失败: {err}"))?; + if !metadata.is_file() { + return Err("更新元数据必须是普通文件".to_string()); + } + // SAFETY: geteuid has no preconditions and retains no pointers. + let effective_uid = unsafe { libc::geteuid() }; + if metadata.uid() != effective_uid || metadata.mode() & 0o022 != 0 { + return Err("更新元数据所有者或写权限不安全".to_string()); + } + if metadata.len() > u64::try_from(max_bytes).unwrap_or(u64::MAX) { + return Err("更新元数据超过大小限制".to_string()); + } + let mut bytes = Vec::with_capacity((metadata.len() as usize).min(max_bytes)); + Read::by_ref(&mut file) + .take( + u64::try_from(max_bytes) + .unwrap_or(u64::MAX) + .saturating_add(1), + ) + .read_to_end(&mut bytes) + .map_err(|err| format!("读取更新元数据失败: {err}"))?; + if bytes.len() > max_bytes { + return Err("更新元数据超过大小限制".to_string()); + } + Ok(Some(bytes)) + } + + #[cfg(not(unix))] + { + let _ = (path, max_bytes); + Err("当前平台不支持安全读取自更新元数据".to_string()) + } +} + +fn write_update_metadata_atomic(path: &Path, bytes: &[u8]) -> Result<(), String> { + #[cfg(unix)] + { + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + + let (parent, output_file_name) = open_real_update_parent(path)?; + let parent_metadata = parent + .metadata() + .map_err(|err| format!("读取更新元数据目录属性失败: {err}"))?; + // SAFETY: geteuid has no preconditions and retains no pointers. + let effective_uid = unsafe { libc::geteuid() }; + let parent_mode = parent_metadata.mode(); + if parent_metadata.uid() != effective_uid + || parent_mode & 0o022 != 0 + || ((parent_mode >> 6) & 0o3) != 0o3 + { + return Err("更新元数据目录所有者或写权限不安全".to_string()); + } + + let output_name = unix_update_path_component(&output_file_name, "更新元数据文件名")?; + if let Some(stat) = unix_update_file_stat_at(&parent, &output_name)? { + if stat.st_mode & libc::S_IFMT != libc::S_IFREG + || stat.st_uid != effective_uid + || stat.st_mode & 0o022 != 0 + { + return Err("已有更新元数据不是安全的普通文件".to_string()); + } + } + + let temp_file_name = std::ffi::OsString::from(format!( + ".aether-update-metadata-{}-{}.tmp", + std::process::id(), + uuid::Uuid::new_v4() + )); + let temp_name = unix_update_path_component(&temp_file_name, "临时更新元数据文件名")?; + // SAFETY: parent and temp_name remain valid for the call. O_EXCL prevents collisions and + // the successful descriptor is transferred immediately into File. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + temp_name.as_ptr(), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_CLOEXEC | libc::O_NOFOLLOW, + 0o600, + ) + }; + if descriptor < 0 { + return Err(format!( + "创建临时更新元数据失败: {}", + std::io::Error::last_os_error() + )); + } + // SAFETY: openat returned a new owned descriptor. + let mut temp = unsafe { std::fs::File::from_raw_fd(descriptor) }; + let result = (|| -> Result<(), String> { + temp.set_permissions(std::fs::Permissions::from_mode(0o600)) + .map_err(|err| format!("设置更新元数据权限失败: {err}"))?; + temp.write_all(bytes) + .map_err(|err| format!("写入更新元数据失败: {err}"))?; + temp.sync_all() + .map_err(|err| format!("同步更新元数据失败: {err}"))?; + drop(temp); + unix_update_rename_at(&parent, &temp_name, &output_name)?; + parent + .sync_all() + .map_err(|err| format!("同步更新元数据目录失败: {err}"))?; + Ok(()) + })(); + if result.is_err() { + let _ = unix_update_unlink_at(&parent, &temp_name); + } + result + } + + #[cfg(not(unix))] + { + let _ = (path, bytes); + Err("当前平台不支持安全写入自更新元数据".to_string()) + } +} + +#[cfg(unix)] +fn open_real_update_parent(path: &Path) -> Result<(std::fs::File, std::ffi::OsString), String> { + let file_name = path + .file_name() + .map(std::ffi::OsString::from) + .ok_or_else(|| "更新元数据路径缺少文件名".to_string())?; + let parent = path + .parent() + .filter(|value| !value.as_os_str().is_empty()) + .unwrap_or(Path::new(".")); + let canonical_parent = + std::fs::canonicalize(parent).map_err(|err| format!("解析更新元数据目录失败: {err}"))?; + Ok((open_real_update_directory(&canonical_parent)?, file_name)) +} + +#[cfg(unix)] +fn open_real_update_directory(path: &Path) -> Result { + use std::os::fd::{AsRawFd, FromRawFd}; + + let mut directory = std::fs::File::open(if path.is_absolute() { "/" } else { "." }) + .map_err(|err| format!("打开更新元数据根目录失败: {err}"))?; + for component in path.components() { + let name = match component { + Component::RootDir | Component::CurDir => continue, + Component::Normal(name) => name, + Component::ParentDir | Component::Prefix(_) => { + return Err("更新元数据目录包含不安全路径组件".to_string()) + } + }; + let name = unix_update_path_component(name, "更新元数据目录组件")?; + // SAFETY: directory is a live directory descriptor and name is NUL-terminated. A + // successful descriptor is transferred immediately into File. + let descriptor = unsafe { + libc::openat( + directory.as_raw_fd(), + name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_DIRECTORY, + ) + }; + if descriptor < 0 { + return Err(format!( + "更新元数据目录必须存在且不能经过符号链接: {}", + std::io::Error::last_os_error() + )); + } + // SAFETY: openat returned a new owned descriptor. + directory = unsafe { std::fs::File::from_raw_fd(descriptor) }; + } + Ok(directory) +} + +#[cfg(unix)] +fn unix_update_path_component( + component: &std::ffi::OsStr, + description: &str, +) -> Result { + use std::os::unix::ffi::OsStrExt; + + std::ffi::CString::new(component.as_bytes()).map_err(|_| format!("{description} 包含 NUL 字节")) +} + +#[cfg(unix)] +fn unix_update_file_stat_at( + parent: &std::fs::File, + file_name: &std::ffi::CStr, +) -> Result, String> { + use std::mem::MaybeUninit; + use std::os::fd::AsRawFd; + + let mut stat = MaybeUninit::::uninit(); + // SAFETY: parent and file_name remain valid, and stat points to writable storage initialized + // only after a successful fstatat call. + let status = unsafe { + libc::fstatat( + parent.as_raw_fd(), + file_name.as_ptr(), + stat.as_mut_ptr(), + libc::AT_SYMLINK_NOFOLLOW, + ) + }; + if status == 0 { + // SAFETY: successful fstatat initialized the complete stat value. + return Ok(Some(unsafe { stat.assume_init() })); + } + let error = std::io::Error::last_os_error(); + if error.raw_os_error() == Some(libc::ENOENT) { + Ok(None) + } else { + Err(format!("检查更新元数据目标失败: {error}")) + } +} + +#[cfg(unix)] +fn unix_update_rename_at( + parent: &std::fs::File, + source: &std::ffi::CStr, + destination: &std::ffi::CStr, +) -> Result<(), String> { + use std::os::fd::AsRawFd; + + // SAFETY: parent is live and both names are NUL-terminated components retained for the call. + let status = unsafe { + libc::renameat( + parent.as_raw_fd(), + source.as_ptr(), + parent.as_raw_fd(), + destination.as_ptr(), + ) + }; + if status == 0 { + Ok(()) + } else { + Err(format!( + "原子替换更新元数据失败: {}", + std::io::Error::last_os_error() + )) + } +} + +#[cfg(unix)] +fn unix_update_unlink_at(parent: &std::fs::File, file_name: &std::ffi::CStr) -> Result<(), String> { + use std::os::fd::AsRawFd; + + // SAFETY: parent is live and file_name is a NUL-terminated component retained for the call. + let status = unsafe { libc::unlinkat(parent.as_raw_fd(), file_name.as_ptr(), 0) }; + if status == 0 { + Ok(()) + } else { + Err(format!( + "删除临时更新元数据失败: {}", + std::io::Error::last_os_error() + )) + } +} + +fn remove_update_metadata_file(path: &Path) -> Result<(), String> { + #[cfg(unix)] + { + use std::os::fd::AsRawFd; + + let (parent, file_name) = open_real_update_parent(path)?; + let file_name = unix_update_path_component(&file_name, "更新元数据文件名")?; + let Some(stat) = unix_update_file_stat_at(&parent, &file_name)? else { + return Ok(()); + }; + // SAFETY: geteuid has no preconditions and retains no pointers. + let effective_uid = unsafe { libc::geteuid() }; + if stat.st_mode & libc::S_IFMT != libc::S_IFREG + || stat.st_uid != effective_uid + || stat.st_mode & 0o022 != 0 + { + return Err("拒绝删除不安全的更新元数据".to_string()); + } + // SAFETY: parent is live and file_name is a NUL-terminated component retained for the + // call. unlinkat removes the directory entry itself and never follows it. + let status = unsafe { libc::unlinkat(parent.as_raw_fd(), file_name.as_ptr(), 0) }; + if status != 0 { + return Err(format!( + "删除更新元数据失败: {}", + std::io::Error::last_os_error() + )); + } + parent + .sync_all() + .map_err(|err| format!("同步更新元数据目录失败: {err}"))?; + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = path; + Err("当前平台不支持安全删除自更新元数据".to_string()) + } +} + +fn safe_system_update_operation(operation: &str) -> &'static str { + match operation.trim() { + "prepare" => "prepare", + "apply" => "apply", + "rollback" => "rollback", + _ => "unknown", + } +} + +fn safe_system_update_output(operation: &str) -> &'static str { + match safe_system_update_operation(operation) { + "prepare" => "Update package prepared", + "apply" => "System update applied", + "rollback" => "System rollback applied", + _ => "System update step completed", + } +} + +fn safe_system_update_error(error: &str) -> &'static str { + let lowered = error.trim().to_ascii_lowercase(); + if lowered.contains("timeout") || lowered.contains("timed out") || lowered.contains("超时") { + "System update timed out" + } else if lowered.contains("sha256") || lowered.contains("checksum") || lowered.contains("校验") + { + "Update package checksum verification failed" + } else if [ + "download", + "http", + "connect", + "connection", + "dns", + "tls", + "certificate", + "下载", + ] + .iter() + .any(|marker| lowered.contains(marker)) + { + "Update package download failed" + } else if ["archive", "extract", "unpack", "gzip", "tar", "解压"] + .iter() + .any(|marker| lowered.contains(marker)) + { + "Update package extraction failed" + } else if [ + "symlink", + "rename", + "permission", + "directory", + "filesystem", + "install", + "符号链接", + "目录", + "切换", + ] + .iter() + .any(|marker| lowered.contains(marker)) + { + "Update installation failed" + } else if lowered.contains("utf-8") || lowered.contains("invalid") || lowered.contains("无效") + { + "Update package is invalid" + } else { + "System update failed" + } } static SYSTEM_UPDATE_RUNNING: AtomicBool = AtomicBool::new(false); @@ -347,7 +858,7 @@ pub(crate) fn self_update_supported() -> bool { is_release_build(), current_update_strategy(), current_deployment_topology(), - ) + ) && self_update_storage_ready() } pub(crate) fn current_self_update_blocker() -> &'static str { @@ -359,25 +870,97 @@ pub(crate) fn current_self_update_blocker() -> &'static str { } match current_update_strategy() { - UpdateStrategy::SelfManaged => "一键更新可用", + UpdateStrategy::SelfManaged if self_update_storage_ready() => "一键更新可用", + UpdateStrategy::SelfManaged => STORAGE_UPDATE_BLOCKER, UpdateStrategy::Docker => DOCKER_UPDATE_BLOCKER, UpdateStrategy::Manual => MANUAL_UPDATE_BLOCKER, } } -fn update_logs_dir() -> PathBuf { - std::env::var("AETHER_LOG_DIR") - .ok() - .filter(|value| !value.trim().is_empty()) - .map(PathBuf::from) - .unwrap_or_else(|| aether_base_dir().join("logs")) +fn self_update_storage_ready() -> bool { + self_update_storage_ready_at(&aether_base_dir()) } -fn docker_update_command() -> String { - std::env::var("AETHER_DOCKER_UPDATE_COMMAND") - .ok() - .filter(|value| !value.trim().is_empty()) - .unwrap_or_else(|| "./update.sh".to_string()) +fn self_update_storage_ready_at(base_dir: &Path) -> bool { + if !base_dir.is_absolute() + || base_dir.components().any(|component| { + matches!( + component, + Component::CurDir | Component::ParentDir | Component::Prefix(_) + ) + }) + { + return false; + } + + let releases_dir = base_dir.join("releases"); + if !is_safe_writable_update_directory(base_dir) + || !is_safe_writable_update_directory(&releases_dir) + { + return false; + } + + let Ok(base_canonical) = std::fs::canonicalize(base_dir) else { + return false; + }; + if base_canonical != base_dir { + return false; + } + let Ok(releases_canonical) = std::fs::canonicalize(&releases_dir) else { + return false; + }; + if releases_canonical.parent() != Some(base_canonical.as_path()) { + return false; + } + + let current = base_dir.join("current"); + let Ok(current_metadata) = std::fs::symlink_metadata(¤t) else { + return false; + }; + if !current_metadata.file_type().is_symlink() { + return false; + } + let Ok(current_canonical) = std::fs::canonicalize(¤t) else { + return false; + }; + let Ok(relative) = current_canonical.strip_prefix(&releases_canonical) else { + return false; + }; + let mut components = relative.components(); + let Some(Component::Normal(release_name)) = components.next() else { + return false; + }; + components.next().is_none() + && release_name + .to_str() + .is_some_and(|name| safe_release_name(name).is_ok()) +} + +fn is_safe_writable_update_directory(path: &Path) -> bool { + let Ok(metadata) = std::fs::symlink_metadata(path) else { + return false; + }; + if metadata.file_type().is_symlink() || !metadata.is_dir() { + return false; + } + + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + let mode = metadata.mode(); + if mode & 0o022 != 0 { + return false; + } + // SAFETY: these libc calls take no pointers and have no preconditions. + let effective_uid = unsafe { libc::geteuid() }; + metadata.uid() == effective_uid && ((mode >> 6) & 0o3 == 0o3) + } + + #[cfg(not(unix))] + { + false + } } pub(crate) fn build_admin_system_update_capability_payload() -> serde_json::Value { @@ -385,16 +968,15 @@ pub(crate) fn build_admin_system_update_capability_payload() -> serde_json::Valu let update_strategy = current_update_strategy(); let deployment_topology = current_deployment_topology(); let supported = - self_update_supported_for(is_release_build(), update_strategy, deployment_topology); + self_update_supported_for(is_release_build(), update_strategy, deployment_topology) + && self_update_storage_ready(); let rollback_available = supported && find_rollback_target().is_some(); let task_status = read_update_task_status(); - let base_dir = aether_base_dir(); let docker_command = if update_strategy == UpdateStrategy::Docker { - Some(docker_update_command()) + Some("./update.sh") } else { None }; - let data_dir = base_dir.join("data"); json!({ "supported": supported, "enabled": supported, @@ -406,10 +988,6 @@ pub(crate) fn build_admin_system_update_capability_payload() -> serde_json::Valu "strategy": update_strategy.as_str(), "deployment_topology": deployment_topology.as_str(), "topology": deployment_topology.as_str(), - "install_root": base_dir.clone(), - "base_dir": base_dir, - "data_dir": data_dir, - "logs_dir": update_logs_dir(), "docker_update_command": docker_command, "message": if supported { "一键更新可用" @@ -420,18 +998,17 @@ pub(crate) fn build_admin_system_update_capability_payload() -> serde_json::Valu } fn find_rollback_target() -> Option { - let previous_path = aether_base_dir().join(PREVIOUS_RELEASE_FILENAME); - let previous = std::fs::read_to_string(previous_path).ok()?; - let previous = previous.trim().to_string(); - if previous.is_empty() { - return None; - } - let target_dir = release_dir_for_version(&previous).ok()?; - if target_dir.is_dir() { - Some(previous) - } else { - None - } + find_rollback_target_at(&aether_base_dir()) +} + +fn find_rollback_target_at(base_dir: &Path) -> Option { + let previous_path = base_dir.join(PREVIOUS_RELEASE_FILENAME); + let previous = read_update_metadata_file(&previous_path, MAX_PREVIOUS_RELEASE_BYTES).ok()??; + let previous = std::str::from_utf8(&previous).ok()?.trim(); + let previous = safe_release_name(previous).ok()?; + let target_dir = base_dir.join("releases").join(&previous); + validate_release_payload_dir(&target_dir).ok()?; + Some(previous) } pub(crate) async fn prepare_admin_system_update_task( @@ -448,27 +1025,22 @@ pub(crate) async fn prepare_admin_system_update_task( json!({ "detail": "缺少 SHA256SUMS 校验文件,已拒绝在线更新" }), ))); }; + if validate_update_release_urls(&version, &tarball_url, &sha256sums_url).is_err() { + return Ok(Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "更新资产 URL 与官方版本或当前平台不匹配" }), + ))); + } let Some(guard) = SystemUpdateGuard::try_acquire() else { return Ok(Err(update_already_running_response())); }; + clear_prepared_version(); set_update_task_phase("preparing"); tokio::spawn(async move { let _guard = guard; - let total_timeout = update_download_total_timeout(); - let result = match tokio::time::timeout( - total_timeout, - download_and_extract_release(&version, &tarball_url, &sha256sums_url), - ) - .await - { - Ok(result) => result, - Err(_) => Err(format!( - "下载更新包超时: 超过 {} 秒", - total_timeout.as_secs() - )), - }; + let result = download_and_extract_release(&version, &tarball_url, &sha256sums_url).await; match result { Ok(output) => { @@ -496,23 +1068,30 @@ async fn download_and_extract_release( tarball_url: &str, sha256sums_url: &str, ) -> Result { - let client = build_update_http_client(update_download_total_timeout(), "更新下载")?; + validate_update_release_urls(version, tarball_url, sha256sums_url)?; + let total_timeout = update_download_total_timeout(); + let client = build_update_http_client(total_timeout, "更新下载")?; + let (tarball_bytes, sha256_text) = tokio::time::timeout(total_timeout, async { + set_update_task_phase("downloading"); + let tarball_bytes = + download_update_bytes(&client, tarball_url, MAX_RELEASE_DOWNLOAD_BYTES, "更新包") + .await?; - set_update_task_phase("downloading"); - let tarball_bytes = - download_update_bytes(&client, tarball_url, MAX_RELEASE_DOWNLOAD_BYTES, "更新包").await?; - - set_update_task_phase("downloading_checksum"); - let sha256_text = String::from_utf8( - download_update_bytes( - &client, - sha256sums_url, - MAX_SHA256SUMS_DOWNLOAD_BYTES, - "校验文件", + set_update_task_phase("downloading_checksum"); + let sha256_text = String::from_utf8( + download_update_bytes( + &client, + sha256sums_url, + MAX_SHA256SUMS_DOWNLOAD_BYTES, + "校验文件", + ) + .await?, ) - .await?, - ) - .map_err(|err| format!("校验文件不是有效 UTF-8: {err}"))?; + .map_err(|_| "校验文件不是有效 UTF-8".to_string())?; + Ok::<_, String>((tarball_bytes, sha256_text)) + }) + .await + .map_err(|_| format!("下载更新包超时: 超过 {} 秒", total_timeout.as_secs()))??; let tarball_url_owned = tarball_url.to_string(); let version_owned = version.to_string(); @@ -523,7 +1102,7 @@ async fn download_and_extract_release( extract_release(&version_owned, &tarball_bytes) }) .await - .map_err(|err| format!("\u{89e3}\u{538b}\u{4efb}\u{52a1}\u{5f02}\u{5e38}: {err}"))? + .map_err(|_| "\u{89e3}\u{538b}\u{4efb}\u{52a1}\u{5f02}\u{5e38}".to_string())? } async fn download_update_bytes( @@ -549,9 +1128,13 @@ async fn download_update_bytes( idle_timeout.as_secs() ) })? - .map_err(|err| format!("下载{label}失败: {err}"))? - .error_for_status() - .map_err(|err| format!("下载{label}返回错误: {err}"))?; + .map_err(|_| format!("下载{label}失败: 网络连接错误"))?; + if !response.status().is_success() { + return Err(format!( + "下载{label}返回错误状态: {}", + response.status().as_u16() + )); + } if let Some(content_length) = response.content_length() { if content_length > max_bytes { @@ -574,7 +1157,7 @@ async fn download_update_bytes( ) })? { - let chunk = chunk.map_err(|err| format!("读取{label}数据失败: {err}"))?; + let chunk = chunk.map_err(|_| format!("读取{label}数据失败"))?; let next_len = data.len() as u64 + chunk.len() as u64; if next_len > max_bytes { return Err(format!("{label}超过大小限制: 最大允许 {max_bytes} bytes")); @@ -610,21 +1193,56 @@ fn update_timeout_from_env(key: &str, default_secs: u64) -> std::time::Duration } fn validate_update_download_url(raw_url: &str) -> Result<(), String> { - let parsed = url::Url::parse(raw_url).map_err(|err| format!("下载 URL 无效: {err}"))?; - if parsed.scheme() != "https" { - return Err("下载 URL 必须使用 HTTPS".to_string()); - } - let Some(host) = parsed.host_str() else { - return Err("下载 URL 缺少主机名".to_string()); - }; - if host == "github.com" - || host.ends_with(".github.com") - || host == "objects.githubusercontent.com" - || host.ends_with(".objects.githubusercontent.com") - { + let parsed = url::Url::parse(raw_url).map_err(|_| "下载 URL 无效".to_string())?; + if is_trusted_update_url(&parsed) { return Ok(()); } - Err(format!("下载 URL 主机不受信任: {host}")) + Err("下载 URL 必须使用无凭据的 HTTPS GitHub 发布主机".to_string()) +} + +fn validate_update_release_urls( + version: &str, + tarball_url: &str, + sha256sums_url: &str, +) -> Result<(), String> { + let safe_version = safe_release_name(version)?; + let platform = if cfg!(target_os = "macos") { + "macos" + } else if cfg!(target_os = "linux") { + "linux" + } else { + return Err("当前平台没有受支持的官方更新资产".to_string()); + }; + let arch = if cfg!(target_arch = "aarch64") { + "arm64" + } else if cfg!(target_arch = "x86_64") { + "amd64" + } else { + return Err("当前架构没有受支持的官方更新资产".to_string()); + }; + let asset_name = format!("aether-{safe_version}-{platform}-{arch}.tar.gz"); + let release_prefix = format!("/fawney19/Aether/releases/download/{safe_version}/"); + + for (raw_url, expected_name) in [ + (tarball_url, asset_name.as_str()), + (sha256sums_url, "SHA256SUMS"), + ] { + let parsed = url::Url::parse(raw_url).map_err(|_| "更新资产 URL 无效".to_string())?; + if parsed.scheme() != "https" + || !parsed.username().is_empty() + || parsed.password().is_some() + || !parsed + .host_str() + .is_some_and(|host| host.eq_ignore_ascii_case("github.com")) + || parsed.port().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + || parsed.path() != format!("{release_prefix}{expected_name}") + { + return Err("更新资产 URL 必须精确指向官方同版本发布资产".to_string()); + } + } + Ok(()) } fn verify_sha256(data: &[u8], sums_text: &str, tarball_url: &str) -> Result<(), String> { @@ -632,27 +1250,43 @@ fn verify_sha256(data: &[u8], sums_text: &str, tarball_url: &str) -> Result<(), "\u{65e0}\u{6cd5}\u{4ece} URL \u{63d0}\u{53d6}\u{6587}\u{4ef6}\u{540d}".to_string() })?; - let expected_hash = sums_text + let expected_hashes = sums_text .lines() - .find_map(|line| { - let (hash, name) = line.split_once(char::is_whitespace)?; - let name = name.trim().trim_start_matches('*'); - if name == tarball_filename { - Some(hash.to_string()) - } else { - None + .filter_map(|line| { + let mut fields = line.split_ascii_whitespace(); + let (Some(hash), Some(name)) = (fields.next(), fields.next()) else { + return None; + }; + if fields.next().is_some() { + return None; } + let name = name.strip_prefix('*').unwrap_or(name); + (name == tarball_filename + && hash.len() == 64 + && hash.bytes().all(|byte| byte.is_ascii_hexdigit())) + .then(|| hash.to_ascii_lowercase()) }) - .ok_or_else(|| { - format!("SHA256SUMS \u{4e2d}\u{672a}\u{627e}\u{5230} {tarball_filename} \u{7684}\u{6821}\u{9a8c}\u{503c}") - })?; + .collect::>(); + let expected_hash = match expected_hashes.as_slice() { + [hash] => hash, + [] => { + return Err(format!( + "SHA256SUMS \u{4e2d}\u{672a}\u{627e}\u{5230} {tarball_filename} \u{7684}\u{552f}\u{4e00}\u{6709}\u{6548}\u{6821}\u{9a8c}\u{503c}" + )); + } + _ => { + return Err(format!( + "SHA256SUMS \u{4e2d} {tarball_filename} \u{5b58}\u{5728}\u{591a}\u{4e2a}\u{6709}\u{6548}\u{6821}\u{9a8c}\u{503c}" + )); + } + }; let mut hasher = Sha256::new(); hasher.update(data); let hash = hasher.finalize(); let actual_hash: String = hash.iter().map(|b| format!("{b:02x}")).collect(); - if actual_hash != expected_hash { + if actual_hash != *expected_hash { return Err(format!( "SHA256 \u{6821}\u{9a8c}\u{5931}\u{8d25}: \u{671f}\u{671b} {expected_hash}, \u{5b9e}\u{9645} {actual_hash}" )); @@ -672,9 +1306,7 @@ fn extract_release(version: &str, tarball_bytes: &[u8]) -> Result Result Result Result { + for _ in 0..16 { + let path = base_dir.join(format!( + ".prepare-{}-{}-{}", + safe_version, + std::process::id(), + uuid::Uuid::new_v4() + )); + + #[cfg(unix)] + let result = { + use std::os::unix::fs::DirBuilderExt; + let mut builder = std::fs::DirBuilder::new(); + builder.mode(0o700).create(&path) + }; + #[cfg(not(unix))] + let result = std::fs::create_dir(&path); + + match result { + Ok(()) => return Ok(path), + Err(err) if err.kind() == std::io::ErrorKind::AlreadyExists => continue, + Err(err) => return Err(format!("创建临时版本目录失败: {err}")), + } + } + Err("无法分配唯一的临时版本目录".to_string()) +} + +fn rename_release_dir_noreplace(source: &Path, destination: &Path) -> Result<(), String> { + #[cfg(any(target_os = "linux", target_os = "macos"))] + { + use std::ffi::CString; + use std::os::unix::ffi::OsStrExt; + + let source = CString::new(source.as_os_str().as_bytes()) + .map_err(|_| "临时版本目录包含无效字符".to_string())?; + let destination = CString::new(destination.as_os_str().as_bytes()) + .map_err(|_| "版本目录包含无效字符".to_string())?; + + #[cfg(target_os = "linux")] + // SAFETY: both paths are valid NUL-terminated strings and remain alive for the call. + let status = unsafe { + libc::renameat2( + libc::AT_FDCWD, + source.as_ptr(), + libc::AT_FDCWD, + destination.as_ptr(), + libc::RENAME_NOREPLACE, + ) + }; + #[cfg(target_os = "macos")] + // SAFETY: both paths are valid NUL-terminated strings and remain alive for the call. + let status = + unsafe { libc::renamex_np(source.as_ptr(), destination.as_ptr(), libc::RENAME_EXCL) }; + + if status == 0 { + return Ok(()); + } + let error = std::io::Error::last_os_error(); + if error.kind() == std::io::ErrorKind::AlreadyExists { + return Err("版本目录已存在,拒绝覆盖以保留回滚完整性".to_string()); + } + Err(format!("安装版本目录失败: {error}")) + } + + #[cfg(not(any(target_os = "linux", target_os = "macos")))] + { + let _ = (source, destination); + Err("当前平台不支持安全的原子版本安装".to_string()) + } +} + +fn sync_release_tree(path: &Path) -> Result<(), String> { + let metadata = std::fs::symlink_metadata(path) + .map_err(|err| format!("读取待安装版本元数据失败: {err}"))?; + if metadata.file_type().is_symlink() { + return Err(format!("待安装版本包含符号链接: {}", path.display())); + } + if metadata.is_file() { + return std::fs::File::open(path) + .and_then(|file| file.sync_all()) + .map_err(|err| format!("同步待安装版本文件失败: {err}")); + } + if !metadata.is_dir() { + return Err(format!("待安装版本包含特殊文件: {}", path.display())); + } + for entry in std::fs::read_dir(path).map_err(|err| format!("读取待安装版本目录失败: {err}"))? + { + let entry = entry.map_err(|err| format!("读取待安装版本条目失败: {err}"))?; + sync_release_tree(&entry.path())?; + } + sync_update_directory(path).map_err(|err| format!("同步待安装版本目录失败: {err}")) +} + +fn sync_update_directory(path: &Path) -> std::io::Result<()> { + std::fs::File::open(path)?.sync_all() +} + fn unpack_release_archive(tarball_bytes: &[u8], staging_dir: &Path) -> Result<(), String> { + unpack_release_archive_with_limits( + tarball_bytes, + staging_dir, + MAX_RELEASE_ARCHIVE_ENTRIES, + MAX_EXTRACTED_RELEASE_BYTES, + ) +} + +fn unpack_release_archive_with_limits( + tarball_bytes: &[u8], + staging_dir: &Path, + max_entries: usize, + max_extracted_bytes: u64, +) -> Result<(), String> { let decoder = flate2::read::GzDecoder::new(std::io::Cursor::new(tarball_bytes)); let mut archive = tar::Archive::new(decoder); let entries = archive .entries() .map_err(|err| format!("读取更新包失败: {err}"))?; let mut extracted_bytes = 0u64; + let mut entry_count = 0usize; + let mut entry_paths = HashSet::new(); for entry in entries { + entry_count = entry_count.saturating_add(1); + if entry_count > max_entries { + return Err(format!("更新包条目过多: 最大允许 {max_entries} 个条目")); + } let mut entry = entry.map_err(|err| format!("读取更新包条目失败: {err}"))?; let path = entry .path() .map_err(|err| format!("读取更新包路径失败: {err}"))? .to_path_buf(); - validate_archive_entry_path(&path)?; + let normalized_path = validate_archive_entry_path(&path)?; + if !entry_paths.insert(normalized_path) { + return Err(format!("更新包包含重复路径: {}", path.display())); + } let entry_type = entry.header().entry_type(); if entry_type.is_file() { @@ -737,14 +1501,21 @@ fn unpack_release_archive(tarball_bytes: &[u8], staging_dir: &Path) -> Result<() .size() .map_err(|err| format!("读取更新包文件大小失败: {err}"))?; extracted_bytes = extracted_bytes.saturating_add(size); - if extracted_bytes > MAX_EXTRACTED_RELEASE_BYTES { + if extracted_bytes > max_extracted_bytes { return Err(format!( - "更新包解压后过大: 最大允许 {MAX_EXTRACTED_RELEASE_BYTES} bytes" + "更新包解压后过大: 最大允许 {max_extracted_bytes} bytes" )); } } else if !entry_type.is_dir() { return Err(format!("更新包包含不支持的条目: {}", path.display())); } + let mode = entry + .header() + .mode() + .map_err(|err| format!("读取更新包权限失败: {err}"))?; + if mode & 0o7022 != 0 { + return Err(format!("更新包包含不安全权限: {}", path.display())); + } let unpacked = entry .unpack_in(staging_dir) @@ -757,21 +1528,30 @@ fn unpack_release_archive(tarball_bytes: &[u8], staging_dir: &Path) -> Result<() Ok(()) } -fn validate_archive_entry_path(path: &Path) -> Result<(), String> { - let mut has_normal_component = false; +fn validate_archive_entry_path(path: &Path) -> Result { + let mut normalized = PathBuf::new(); + let mut depth = 0usize; for component in path.components() { match component { - Component::Normal(_) => has_normal_component = true, - Component::CurDir => {} - Component::ParentDir | Component::RootDir | Component::Prefix(_) => { + Component::Normal(value) => { + depth = depth.saturating_add(1); + if depth > MAX_RELEASE_ARCHIVE_PATH_DEPTH { + return Err(format!("更新包路径层级过深: {}", path.display())); + } + normalized.push(value); + } + Component::CurDir + | Component::ParentDir + | Component::RootDir + | Component::Prefix(_) => { return Err(format!("更新包包含非法路径: {}", path.display())); } } } - if has_normal_component { - Ok(()) - } else { + if normalized.as_os_str().is_empty() { Err("更新包包含空路径".to_string()) + } else { + Ok(normalized) } } @@ -781,10 +1561,12 @@ fn find_release_payload_dir(staging_dir: &Path) -> Result { } let mut candidates = Vec::new(); + let mut top_level_entries = 0usize; let entries = std::fs::read_dir(staging_dir).map_err(|err| format!("读取更新包目录失败: {err}"))?; for entry in entries { let entry = entry.map_err(|err| format!("读取更新包条目失败: {err}"))?; + top_level_entries = top_level_entries.saturating_add(1); let path = entry.path(); if path.is_dir() && looks_like_release_payload(&path) { candidates.push(path); @@ -792,7 +1574,8 @@ fn find_release_payload_dir(staging_dir: &Path) -> Result { } match candidates.len() { - 1 => Ok(candidates.remove(0)), + 1 if top_level_entries == 1 => Ok(candidates.remove(0)), + 1 => Err("更新包必须只包含一个顶层版本目录".to_string()), 0 => Err( "\u{66f4}\u{65b0}\u{5305}\u{4e2d}\u{672a}\u{627e}\u{5230} bin/aether-gateway" .to_string(), @@ -802,28 +1585,98 @@ fn find_release_payload_dir(staging_dir: &Path) -> Result { } fn looks_like_release_payload(path: &Path) -> bool { - path.join("bin/aether-gateway").is_file() && path.join("frontend").is_dir() + is_nonsymlink_regular_file(&path.join("bin/aether-gateway")) + && is_nonsymlink_directory(&path.join("frontend")) } fn validate_release_payload_dir(path: &Path) -> Result<(), String> { - if !path.join("bin/aether-gateway").is_file() { + let mut entry_count = 0usize; + validate_release_tree_entry(path, 0, &mut entry_count)?; + + if !is_nonsymlink_regular_file(&path.join("bin/aether-gateway")) { return Err( "\u{66f4}\u{65b0}\u{5305}\u{4e2d}\u{672a}\u{627e}\u{5230} bin/aether-gateway" .to_string(), ); } - if !path.join("frontend/index.html").is_file() { + if !is_nonsymlink_regular_file(&path.join("frontend/index.html")) { return Err("更新包中未找到 frontend/index.html".to_string()); } Ok(()) } -fn ensure_release_binary_permissions(binary_path: &Path) { +fn is_nonsymlink_regular_file(path: &Path) -> bool { + std::fs::symlink_metadata(path) + .is_ok_and(|metadata| !metadata.file_type().is_symlink() && metadata.is_file()) +} + +fn is_nonsymlink_directory(path: &Path) -> bool { + std::fs::symlink_metadata(path) + .is_ok_and(|metadata| !metadata.file_type().is_symlink() && metadata.is_dir()) +} + +fn validate_release_tree_entry( + path: &Path, + depth: usize, + entry_count: &mut usize, +) -> Result<(), String> { + if depth > MAX_RELEASE_ARCHIVE_PATH_DEPTH { + return Err(format!("版本目录路径层级过深: {}", path.display())); + } + *entry_count = entry_count.saturating_add(1); + if *entry_count > MAX_RELEASE_ARCHIVE_ENTRIES.saturating_add(1) { + return Err(format!( + "版本目录条目过多: 最大允许 {MAX_RELEASE_ARCHIVE_ENTRIES} 个条目" + )); + } + + let metadata = + std::fs::symlink_metadata(path).map_err(|err| format!("读取版本目录属性失败: {err}"))?; + if metadata.file_type().is_symlink() { + return Err(format!("版本目录包含符号链接: {}", path.display())); + } + if !metadata.is_file() && !metadata.is_dir() { + return Err(format!("版本目录包含特殊文件: {}", path.display())); + } + + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + // SAFETY: geteuid has no preconditions and retains no pointers. + let effective_uid = unsafe { libc::geteuid() }; + if metadata.uid() != effective_uid { + return Err(format!( + "版本目录包含非当前用户所有的条目: {}", + path.display() + )); + } + if metadata.mode() & 0o7022 != 0 { + return Err(format!("版本目录包含不安全权限: {}", path.display())); + } + if metadata.is_file() && metadata.nlink() != 1 { + return Err(format!("版本目录包含硬链接文件: {}", path.display())); + } + } + + if metadata.is_dir() { + for entry in std::fs::read_dir(path).map_err(|err| format!("读取版本目录失败: {err}"))? + { + let entry = entry.map_err(|err| format!("读取版本目录条目失败: {err}"))?; + validate_release_tree_entry(&entry.path(), depth.saturating_add(1), entry_count)?; + } + } + Ok(()) +} + +fn ensure_release_binary_permissions(binary_path: &Path) -> Result<(), String> { #[cfg(unix)] { use std::os::unix::fs::PermissionsExt; - let _ = std::fs::set_permissions(binary_path, std::fs::Permissions::from_mode(0o755)); + std::fs::set_permissions(binary_path, std::fs::Permissions::from_mode(0o755)) + .map_err(|err| format!("设置更新程序权限失败: {err}"))?; } + Ok(()) } fn remove_path_if_exists(path: &Path) -> std::io::Result<()> { @@ -856,17 +1709,17 @@ pub(crate) async fn start_admin_system_update_task( let release_dir = match release_dir_for_version(&version) { Ok(dir) => dir, - Err(err) => { + Err(_) => { return Ok(Err(( http::StatusCode::BAD_REQUEST, - json!({ "detail": err }), + json!({ "detail": "版本号无效" }), ))); } }; if !release_dir.join("bin/aether-gateway").is_file() { return Ok(Err(( http::StatusCode::PRECONDITION_REQUIRED, - json!({ "detail": format!("\u{7248}\u{672c} {version} \u{5c1a}\u{672a}\u{51c6}\u{5907}\u{597d}\u{ff0c}\u{8bf7}\u{5148}\u{6267}\u{884c} prepare-update") }), + json!({ "detail": "指定版本尚未准备好,请先执行 prepare-update" }), ))); } @@ -874,7 +1727,14 @@ pub(crate) async fn start_admin_system_update_task( return Ok(Err(update_already_running_response())); }; - save_previous_release(); + if let Err(err) = save_previous_release() { + set_update_task_failed(err.clone()); + append_update_history("apply", false, Some(&err), None); + return Ok(Err(( + http::StatusCode::INTERNAL_SERVER_ERROR, + json!({ "detail": "无法安全保存回滚状态,已取消更新" }), + ))); + } set_update_task_phase("restarting"); tokio::spawn(async move { @@ -893,7 +1753,15 @@ pub(crate) async fn start_admin_system_update_task( request_process_restart(); } Err(err) => { - tracing::error!(error = %err, "admin system update apply failed"); + let prev_path = aether_base_dir().join(PREVIOUS_RELEASE_FILENAME); + if let Err(cleanup_err) = remove_update_metadata_file(&prev_path) { + tracing::warn!( + error = %safe_system_update_error(&cleanup_err), + "failed to clear rollback metadata after update failure" + ); + } + let safe_error = safe_system_update_error(&err); + tracing::error!(error = %safe_error, "admin system update apply failed"); append_update_history("apply", false, Some(&err), None); set_update_task_failed(err); } @@ -908,19 +1776,31 @@ pub(crate) async fn start_admin_system_update_task( }))) } -fn save_previous_release() { - let current = current_symlink_path(); - if let Ok(target) = std::fs::read_link(¤t) { - if let Some(name) = target.file_name().and_then(|n| n.to_str()) { - let prev_path = aether_base_dir().join(PREVIOUS_RELEASE_FILENAME); - let _ = std::fs::write(prev_path, name); - } - } +fn save_previous_release() -> Result<(), String> { + save_previous_release_at(&aether_base_dir()) +} + +fn save_previous_release_at(base_dir: &Path) -> Result<(), String> { + let name = + current_release_name_at(base_dir).ok_or_else(|| "当前版本符号链接不安全".to_string())?; + let prev_path = base_dir.join(PREVIOUS_RELEASE_FILENAME); + write_update_metadata_atomic(&prev_path, name.as_bytes()) } fn switch_current_symlink(version: &str) -> Result<(), String> { - let target = release_dir_for_version(version)?; - if !target.is_dir() { + switch_current_symlink_at(&aether_base_dir(), version) +} + +fn switch_current_symlink_at(base_dir: &Path, version: &str) -> Result<(), String> { + let safe_version = safe_release_name(version)?; + let target = base_dir.join("releases").join(&safe_version); + let target_metadata = std::fs::symlink_metadata(&target).map_err(|err| { + format!( + "\u{7248}\u{672c}\u{76ee}\u{5f55}\u{4e0d}\u{5b58}\u{5728}: {} ({err})", + target.display() + ) + })?; + if target_metadata.file_type().is_symlink() || !target_metadata.is_dir() { return Err(format!( "\u{7248}\u{672c}\u{76ee}\u{5f55}\u{4e0d}\u{5b58}\u{5728}: {}", target.display() @@ -928,22 +1808,76 @@ fn switch_current_symlink(version: &str) -> Result<(), String> { } validate_release_payload_dir(&target)?; - let current = current_symlink_path(); - let current_new = current.with_file_name("current.new"); - - let _ = remove_path_if_exists(¤t_new); - #[cfg(unix)] - std::os::unix::fs::symlink(&target, ¤t_new) - .map_err(|err| format!("\u{521b}\u{5efa}\u{4e34}\u{65f6}\u{7b26}\u{53f7}\u{94fe}\u{63a5}\u{5931}\u{8d25}: {err}"))?; - #[cfg(windows)] - std::os::windows::fs::symlink_dir(&target, ¤t_new) - .map_err(|err| format!("\u{521b}\u{5efa}\u{4e34}\u{65f6}\u{7b26}\u{53f7}\u{94fe}\u{63a5}\u{5931}\u{8d25}: {err}"))?; + { + use std::os::fd::AsRawFd; - std::fs::rename(¤t_new, ¤t) - .map_err(|err| format!("\u{539f}\u{5b50}\u{5207}\u{6362}\u{7b26}\u{53f7}\u{94fe}\u{63a5}\u{5931}\u{8d25}: {err}"))?; + let canonical_base = + std::fs::canonicalize(base_dir).map_err(|err| format!("解析安装目录失败: {err}"))?; + let parent = open_real_update_directory(&canonical_base)?; + let current_name = + unix_update_path_component(std::ffi::OsStr::new("current"), "当前版本符号链接名")?; + let Some(current_stat) = unix_update_file_stat_at(&parent, ¤t_name)? else { + return Err("当前版本符号链接不存在".to_string()); + }; + // SAFETY: geteuid has no preconditions and retains no pointers. + let effective_uid = unsafe { libc::geteuid() }; + if current_stat.st_mode & libc::S_IFMT != libc::S_IFLNK + || current_stat.st_uid != effective_uid + { + return Err("当前版本入口不是受当前进程管理的符号链接".to_string()); + } - Ok(()) + let temp_file_name = std::ffi::OsString::from(format!( + ".current-{}-{}.new", + std::process::id(), + uuid::Uuid::new_v4() + )); + let temp_name = unix_update_path_component(&temp_file_name, "临时版本符号链接名")?; + let relative_target = std::ffi::CString::new(format!("releases/{safe_version}")) + .map_err(|_| "版本符号链接目标包含 NUL 字节".to_string())?; + + // SAFETY: parent is live and both target and temp_name are NUL-terminated for the call. + let symlink_status = unsafe { + libc::symlinkat( + relative_target.as_ptr(), + parent.as_raw_fd(), + temp_name.as_ptr(), + ) + }; + if symlink_status != 0 { + return Err(format!( + "创建唯一临时版本符号链接失败: {}", + std::io::Error::last_os_error() + )); + } + + // SAFETY: parent is live and both names are NUL-terminated for the call. renameat + // atomically replaces the current directory entry without following either symlink. + let rename_status = unsafe { + libc::renameat( + parent.as_raw_fd(), + temp_name.as_ptr(), + parent.as_raw_fd(), + current_name.as_ptr(), + ) + }; + if rename_status != 0 { + let error = std::io::Error::last_os_error(); + let _ = unix_update_unlink_at(&parent, &temp_name); + return Err(format!("原子切换版本符号链接失败: {error}")); + } + parent + .sync_all() + .map_err(|err| format!("同步版本入口目录失败: {err}"))?; + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = (base_dir, safe_version); + Err("当前平台不支持安全的原子版本切换".to_string()) + } } pub(crate) async fn start_admin_system_rollback_task( @@ -969,7 +1903,9 @@ pub(crate) async fn start_admin_system_rollback_task( match switch_current_symlink(&previous) { Ok(_) => { let prev_path = aether_base_dir().join(PREVIOUS_RELEASE_FILENAME); - let _ = std::fs::remove_file(prev_path); + if let Err(err) = remove_update_metadata_file(&prev_path) { + tracing::warn!(error = %safe_system_update_error(&err), "failed to clear rollback metadata"); + } append_update_history( "rollback", @@ -983,7 +1919,8 @@ pub(crate) async fn start_admin_system_rollback_task( request_process_restart(); } Err(err) => { - tracing::error!(error = %err, "admin system rollback failed"); + let safe_error = safe_system_update_error(&err); + tracing::error!(error = %safe_error, "admin system rollback failed"); append_update_history("rollback", false, Some(&err), None); set_update_task_failed(err); } @@ -1041,6 +1978,212 @@ mod tests { .expect("frontend index should be written"); } + #[test] + fn reading_update_history_rewrites_legacy_sensitive_fields() { + let dir = temp_test_dir("history-read-redaction"); + let path = dir.join(UPDATE_HISTORY_FILENAME); + std::fs::create_dir_all(&dir).expect("history directory should be created"); + let legacy = json!([{ + "timestamp": "Bearer timestamp-secret", + "operation": "prepare?token=operation-secret", + "success": false, + "error": "download failed for https://user:password@internal.test/file?token=query-secret; Authorization: Bearer error-secret", + "output_tail": "installed /opt/private; access_token=output-secret" + }]); + std::fs::write( + &path, + serde_json::to_vec_pretty(&legacy).expect("legacy history should serialize"), + ) + .expect("legacy history should be written"); + + let entries = read_update_history_at_path(&path); + + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].timestamp, "1970-01-01T00:00:00Z"); + assert_eq!(entries[0].operation, "unknown"); + assert_eq!( + entries[0].error.as_deref(), + Some("Update package download failed") + ); + assert_eq!( + entries[0].output_tail.as_deref(), + Some("System update step completed") + ); + + let persisted = std::fs::read_to_string(&path).expect("history should remain readable"); + for secret in [ + "timestamp-secret", + "operation-secret", + "user:password", + "query-secret", + "error-secret", + "/opt/private", + "output-secret", + ] { + assert!( + !persisted.contains(secret), + "persisted history leaked {secret}" + ); + } + std::fs::remove_dir_all(dir).ok(); + } + + #[test] + fn appending_update_history_sanitizes_existing_entries_before_write() { + let dir = temp_test_dir("history-append-redaction"); + let path = dir.join(UPDATE_HISTORY_FILENAME); + std::fs::create_dir_all(&dir).expect("history directory should be created"); + let legacy = json!([{ + "timestamp": "2026-08-27T00:00:00Z", + "operation": "apply", + "success": false, + "error": "Authorization: Bearer legacy-secret", + "output_tail": "legacy output token=old-secret" + }]); + std::fs::write( + &path, + serde_json::to_vec_pretty(&legacy).expect("legacy history should serialize"), + ) + .expect("legacy history should be written"); + + append_update_history_at_path( + &path, + "rollback", + false, + Some("failed with Bearer new-secret"), + Some("private output new-output-secret"), + ); + + let persisted = std::fs::read_to_string(&path).expect("history should remain readable"); + assert!(!persisted.contains("legacy-secret")); + assert!(!persisted.contains("old-secret")); + assert!(!persisted.contains("new-secret")); + assert!(!persisted.contains("new-output-secret")); + let entries: Vec = + serde_json::from_str(&persisted).expect("sanitized history should parse"); + assert_eq!(entries.len(), 2); + assert_eq!(entries[0].error.as_deref(), Some("System update failed")); + assert_eq!( + entries[1].output_tail.as_deref(), + Some("System rollback applied") + ); + std::fs::remove_dir_all(dir).ok(); + } + + #[cfg(unix)] + #[test] + fn update_metadata_atomic_write_is_private_and_rejects_symlink_destination() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let dir = temp_test_dir("metadata-atomic"); + std::fs::create_dir_all(&dir).expect("metadata directory should be created"); + let path = dir.join(UPDATE_HISTORY_FILENAME); + + write_update_metadata_atomic(&path, b"first").expect("metadata should be written"); + write_update_metadata_atomic(&path, b"second").expect("metadata should be replaced"); + assert_eq!(std::fs::read(&path).unwrap(), b"second"); + assert_eq!( + std::fs::metadata(&path).unwrap().permissions().mode() & 0o077, + 0, + "update metadata must remain private" + ); + + std::fs::remove_file(&path).unwrap(); + let victim = dir.join("victim"); + std::fs::write(&victim, b"known-good").unwrap(); + symlink(&victim, &path).expect("metadata symlink should be created"); + + let err = write_update_metadata_atomic(&path, b"attacker-controlled") + .expect_err("symlink destination must be rejected"); + + assert!(err.contains("安全的普通文件")); + assert!(remove_update_metadata_file(&path).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + assert!( + std::fs::symlink_metadata(&path) + .unwrap() + .file_type() + .is_symlink(), + "rejected metadata symlink must not be followed or replaced" + ); + assert!( + std::fs::read_dir(&dir).unwrap().all(|entry| { + !entry + .unwrap() + .file_name() + .to_string_lossy() + .starts_with(".aether-update-metadata-") + }), + "atomic write must not leave temporary files" + ); + std::fs::remove_dir_all(dir).ok(); + } + + #[cfg(unix)] + #[test] + fn update_history_does_not_read_or_rewrite_symlink_target() { + use std::os::unix::fs::symlink; + + let dir = temp_test_dir("history-symlink"); + std::fs::create_dir_all(&dir).expect("history directory should be created"); + let victim = dir.join("victim.json"); + let victim_contents = br#"[{"timestamp":"2026-01-01T00:00:00Z","operation":"apply","success":true,"error":null,"output_tail":null}]"#; + std::fs::write(&victim, victim_contents).unwrap(); + let path = dir.join(UPDATE_HISTORY_FILENAME); + symlink(&victim, &path).expect("history symlink should be created"); + + assert!(read_update_history_at_path(&path).is_empty()); + assert_eq!(std::fs::read(&victim).unwrap(), victim_contents); + assert!(std::fs::symlink_metadata(&path) + .unwrap() + .file_type() + .is_symlink()); + std::fs::remove_dir_all(dir).ok(); + } + + #[cfg(unix)] + #[test] + fn previous_release_metadata_is_atomic_and_rejects_symlink_input() { + use std::os::unix::fs::symlink; + + let base = temp_test_dir("previous-release"); + let release = base.join("releases/v1.2.3"); + write_release_payload(&release); + symlink(&release, base.join("current")).expect("current link should be created"); + + save_previous_release_at(&base).expect("previous release should be persisted"); + assert_eq!(find_rollback_target_at(&base).as_deref(), Some("v1.2.3")); + + let previous = base.join(PREVIOUS_RELEASE_FILENAME); + std::fs::remove_file(&previous).unwrap(); + let victim = base.join("victim"); + std::fs::write(&victim, b"v1.2.3").unwrap(); + symlink(&victim, &previous).expect("previous-release symlink should be created"); + + assert!(find_rollback_target_at(&base).is_none()); + assert!(save_previous_release_at(&base).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"v1.2.3"); + std::fs::remove_dir_all(base).ok(); + } + + #[cfg(unix)] + #[test] + fn previous_release_rejects_current_link_outside_managed_releases() { + use std::os::unix::fs::symlink; + + let base = temp_test_dir("outside-current"); + std::fs::create_dir_all(base.join("releases")).unwrap(); + let outside = temp_test_dir("outside-current-target"); + write_release_payload(&outside); + symlink(&outside, base.join("current")).expect("outside current link should be created"); + + assert!(current_release_name_at(&base).is_none()); + assert!(save_previous_release_at(&base).is_err()); + assert!(!base.join(PREVIOUS_RELEASE_FILENAME).exists()); + std::fs::remove_dir_all(base).ok(); + std::fs::remove_dir_all(outside).ok(); + } + #[test] fn update_strategy_defaults_to_self_only_for_release_builds() { assert_eq!( @@ -1095,6 +2238,21 @@ mod tests { )); } + #[test] + fn update_capability_omits_internal_paths_and_dynamic_commands() { + let payload = build_admin_system_update_capability_payload(); + + for field in ["install_root", "base_dir", "data_dir", "logs_dir"] { + assert!(payload.get(field).is_none(), "capability exposed {field}"); + } + if let Some(command) = payload + .get("docker_update_command") + .and_then(serde_json::Value::as_str) + { + assert_eq!(command, "./update.sh"); + } + } + #[test] fn update_finds_nested_release_payload_dir() { let staging = temp_test_dir("nested"); @@ -1107,6 +2265,78 @@ mod tests { std::fs::remove_dir_all(staging).ok(); } + #[test] + fn update_rejects_nested_payload_with_extra_top_level_entries() { + let staging = temp_test_dir("nested-extra"); + let bundle = staging.join("aether-v1.2.3-linux-amd64"); + write_release_payload(&bundle); + std::fs::write(staging.join("unexpected"), b"extra") + .expect("extra entry should be written"); + + let err = find_release_payload_dir(&staging) + .expect_err("extra top-level entries must be rejected"); + + assert!(err.contains("一个顶层版本目录")); + std::fs::remove_dir_all(staging).ok(); + } + + #[cfg(unix)] + #[test] + fn update_revalidates_release_tree_before_activation() { + use std::os::unix::fs::symlink; + + let release = temp_test_dir("payload-revalidation"); + write_release_payload(&release); + let victim = release.join("victim"); + std::fs::write(&victim, b"outside-index").unwrap(); + std::fs::remove_file(release.join("frontend/index.html")).unwrap(); + symlink(&victim, release.join("frontend/index.html")).unwrap(); + + let err = validate_release_payload_dir(&release) + .expect_err("post-extraction symlink must be rejected"); + assert!(err.contains("符号链接")); + + std::fs::remove_file(release.join("frontend/index.html")).unwrap(); + std::fs::hard_link(&victim, release.join("frontend/index.html")).unwrap(); + let err = validate_release_payload_dir(&release) + .expect_err("post-extraction hard link must be rejected"); + assert!(err.contains("硬链接")); + std::fs::remove_dir_all(release).ok(); + } + + #[cfg(unix)] + #[test] + fn update_switch_uses_unique_relative_symlink_without_deleting_predictable_paths() { + use std::os::unix::fs::symlink; + + let base = temp_test_dir("atomic-current-switch"); + write_release_payload(&base.join("releases/v1.0.0")); + write_release_payload(&base.join("releases/v2.0.0")); + symlink("releases/v1.0.0", base.join("current")).unwrap(); + let predictable = base.join("current.new"); + std::fs::create_dir(&predictable).unwrap(); + std::fs::write(predictable.join("keep"), b"known-good").unwrap(); + + switch_current_symlink_at(&base, "v2.0.0").expect("version switch should succeed"); + + assert_eq!( + std::fs::read_link(base.join("current")).unwrap(), + PathBuf::from("releases/v2.0.0") + ); + assert_eq!( + std::fs::read(predictable.join("keep")).unwrap(), + b"known-good" + ); + assert!(std::fs::read_dir(&base).unwrap().all(|entry| { + !entry + .unwrap() + .file_name() + .to_string_lossy() + .starts_with(".current-") + })); + std::fs::remove_dir_all(base).ok(); + } + #[test] fn update_finds_flat_release_payload_dir() { let staging = temp_test_dir("flat"); @@ -1118,6 +2348,38 @@ mod tests { std::fs::remove_dir_all(found).ok(); } + #[test] + fn update_accepts_release_workflow_archive_layout() { + let fixture = temp_test_dir("workflow-layout-source"); + let payload = fixture.join("payload"); + write_release_payload(&payload); + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), Compression::default()); + { + let mut builder = tar::Builder::new(&mut encoder); + builder + .append_dir_all("aether-v1.2.3-linux-amd64", &payload) + .expect("workflow-style bundle should be appended"); + builder.finish().expect("tar builder should finish"); + } + let tarball = encoder.finish().expect("gzip encoder should finish"); + let staging = temp_test_dir("workflow-layout-destination"); + std::fs::create_dir_all(&staging).expect("staging should be created"); + + unpack_release_archive(&tarball, &staging) + .expect("workflow-style release archive should unpack"); + let found = find_release_payload_dir(&staging) + .expect("workflow-style release payload should be found"); + validate_release_payload_dir(&found) + .expect("workflow-style release payload should validate"); + + assert_eq!( + found.file_name().and_then(|name| name.to_str()), + Some("aether-v1.2.3-linux-amd64") + ); + std::fs::remove_dir_all(fixture).ok(); + std::fs::remove_dir_all(staging).ok(); + } + #[test] fn update_rejects_unsafe_release_names() { assert!(safe_release_name("v1.2.3").is_ok()); @@ -1136,16 +2398,95 @@ mod tests { "https://objects.githubusercontent.com/github-production-release-asset/test" ) .is_ok()); + assert!(validate_update_download_url( + "https://release-assets.githubusercontent.com/github-production-release-asset/test" + ) + .is_ok()); assert!(validate_update_download_url( "http://github.com/fawney19/Aether/releases/download/v1/aether.tar.gz" ) .is_err()); + assert!(validate_update_download_url( + "https://user@github.com/fawney19/Aether/releases/download/v1/aether.tar.gz" + ) + .is_err()); + assert!( + validate_update_download_url("https://github.com.evil.test/aether.tar.gz").is_err() + ); assert!(validate_update_download_url("https://example.com/aether.tar.gz").is_err()); } + #[test] + fn update_binds_tarball_and_checksum_to_official_same_version_release() { + let platform = if cfg!(target_os = "macos") { + "macos" + } else { + "linux" + }; + let arch = if cfg!(target_arch = "aarch64") { + "arm64" + } else { + "amd64" + }; + let tarball = format!( + "https://github.com/fawney19/Aether/releases/download/v1.2.3/aether-v1.2.3-{platform}-{arch}.tar.gz" + ); + let checksums = "https://github.com/fawney19/Aether/releases/download/v1.2.3/SHA256SUMS"; + + validate_update_release_urls("v1.2.3", &tarball, checksums) + .expect("matching official release assets should pass"); + + for (bad_tarball, bad_checksums) in [ + ( + tarball.replace("fawney19/Aether", "attacker/project"), + checksums.to_string(), + ), + (tarball.clone(), checksums.replace("/v1.2.3/", "/v1.2.2/")), + ( + tarball.replace(&format!("-{arch}.tar.gz"), "-wrong.tar.gz"), + checksums.to_string(), + ), + (format!("{tarball}?token=unexpected"), checksums.to_string()), + ] { + assert!( + validate_update_release_urls("v1.2.3", &bad_tarball, &bad_checksums).is_err(), + "mismatched update assets must be rejected" + ); + } + } + + #[test] + fn update_url_validation_and_error_projection_do_not_echo_sensitive_details() { + let rejected = validate_update_download_url( + "https://user:password@github.com/release.tar.gz?token=query-secret", + ) + .expect_err("credential-bearing update URL must be rejected"); + assert!(!rejected.contains("user:password")); + assert!(!rejected.contains("query-secret")); + + let projected = safe_system_update_error( + "request failed for https://user:password@internal.test?q=secret; Authorization: Bearer upstream-secret", + ); + assert_eq!(projected, "Update package download failed"); + assert!(!projected.contains("upstream-secret")); + assert!(!projected.contains("user:password")); + + assert_eq!( + safe_system_update_error( + "download failed: https://user:password@internal.test?q=secret" + ), + "Update package download failed" + ); + assert_eq!( + safe_system_update_error("version directory missing: /opt/aether/releases/secret"), + "Update installation failed" + ); + } + #[test] fn update_rejects_archive_path_traversal() { assert!(validate_archive_entry_path(Path::new("bundle/bin/aether-gateway")).is_ok()); + assert!(validate_archive_entry_path(Path::new("./bundle/bin/aether-gateway")).is_err()); assert!(validate_archive_entry_path(Path::new("../escape")).is_err()); assert!(validate_archive_entry_path(Path::new("/tmp/escape")).is_err()); } @@ -1178,6 +2519,200 @@ mod tests { std::fs::remove_dir_all(staging).ok(); } + #[test] + fn update_rejects_duplicate_archive_paths() { + let staging = temp_test_dir("duplicate-path"); + std::fs::create_dir_all(&staging).expect("staging dir should be created"); + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), Compression::default()); + { + let mut builder = tar::Builder::new(&mut encoder); + for contents in [b"first".as_slice(), b"second".as_slice()] { + let mut header = tar::Header::new_gnu(); + header.set_entry_type(tar::EntryType::Regular); + header.set_size(contents.len() as u64); + header.set_mode(0o644); + header.set_cksum(); + builder + .append_data(&mut header, "bundle/frontend/index.html", contents) + .expect("duplicate fixture entry should be appended"); + } + builder.finish().expect("tar builder should finish"); + } + let tarball = encoder.finish().expect("gzip encoder should finish"); + + let err = unpack_release_archive(&tarball, &staging) + .expect_err("duplicate archive path should fail"); + + assert!(err.contains("重复路径")); + std::fs::remove_dir_all(staging).ok(); + } + + #[test] + fn update_enforces_archive_entry_limit() { + let staging = temp_test_dir("entry-limit"); + std::fs::create_dir_all(&staging).expect("staging dir should be created"); + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), Compression::default()); + { + let mut builder = tar::Builder::new(&mut encoder); + for (path, contents) in [("bundle/one", b"one"), ("bundle/two", b"two")] { + let mut header = tar::Header::new_gnu(); + header.set_entry_type(tar::EntryType::Regular); + header.set_size(contents.len() as u64); + header.set_mode(0o644); + header.set_cksum(); + builder + .append_data(&mut header, path, contents.as_slice()) + .expect("entry-limit fixture should be appended"); + } + builder.finish().expect("tar builder should finish"); + } + let tarball = encoder.finish().expect("gzip encoder should finish"); + + let err = unpack_release_archive_with_limits(&tarball, &staging, 1, u64::MAX) + .expect_err("archive entry limit should be enforced"); + + assert!(err.contains("条目过多")); + std::fs::remove_dir_all(staging).ok(); + } + + #[test] + fn update_rejects_group_or_world_writable_archive_members() { + let staging = temp_test_dir("writable-member"); + std::fs::create_dir_all(&staging).expect("staging dir should be created"); + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), Compression::default()); + { + let mut builder = tar::Builder::new(&mut encoder); + let contents = b"test-binary"; + let mut header = tar::Header::new_gnu(); + header.set_entry_type(tar::EntryType::Regular); + header.set_size(contents.len() as u64); + header.set_mode(0o777); + header.set_cksum(); + builder + .append_data( + &mut header, + "bundle/bin/aether-gateway", + contents.as_slice(), + ) + .expect("writable fixture entry should be appended"); + builder.finish().expect("tar builder should finish"); + } + let tarball = encoder.finish().expect("gzip encoder should finish"); + + let err = unpack_release_archive(&tarball, &staging) + .expect_err("writable archive member should fail"); + + assert!(err.contains("不安全权限")); + std::fs::remove_dir_all(staging).ok(); + } + + #[test] + fn update_preserves_existing_release_destination() { + let dir = temp_test_dir("existing-release"); + let source = dir.join("prepared-v1.2.3"); + let release = dir.join("releases/v1.2.3"); + std::fs::create_dir_all(&source).expect("prepared release should be created"); + std::fs::create_dir_all(&release).expect("existing release should be created"); + let sentinel = release.join("keep"); + std::fs::write(&sentinel, b"known-good").expect("sentinel should be written"); + + let err = rename_release_dir_noreplace(&source, &release) + .expect_err("existing release must not be replaceable"); + + assert!(err.contains("拒绝覆盖")); + assert_eq!(std::fs::read(&sentinel).unwrap(), b"known-good"); + assert!(source.is_dir(), "failed install must retain its source"); + std::fs::remove_dir_all(dir).ok(); + } + + #[cfg(any(target_os = "linux", target_os = "macos"))] + #[test] + fn update_atomically_installs_absent_release_destination() { + let dir = temp_test_dir("new-release"); + let source = dir.join("prepared-v1.2.3"); + let release = dir.join("releases/v1.2.3"); + std::fs::create_dir_all(&source).expect("prepared release should be created"); + std::fs::create_dir_all(release.parent().unwrap()) + .expect("releases parent should be created"); + std::fs::write(source.join("sentinel"), b"prepared") + .expect("prepared sentinel should be written"); + + rename_release_dir_noreplace(&source, &release) + .expect("absent release should install atomically"); + + assert!(!source.exists()); + assert_eq!( + std::fs::read(release.join("sentinel")).unwrap(), + b"prepared" + ); + std::fs::remove_dir_all(dir).ok(); + } + + #[cfg(unix)] + #[test] + fn update_staging_directory_is_unique_private_and_non_destructive() { + use std::os::unix::fs::PermissionsExt; + + let base = temp_test_dir("private-staging"); + std::fs::create_dir_all(&base).expect("base should be created"); + let legacy = base.join(format!(".prepare-v1.2.3-{}", std::process::id())); + std::fs::create_dir(&legacy).expect("legacy staging should be created"); + std::fs::write(legacy.join("keep"), b"existing") + .expect("legacy sentinel should be written"); + + let staging = + create_release_staging_dir(&base, "v1.2.3").expect("unique staging should be created"); + + assert_ne!(staging, legacy); + assert_eq!( + std::fs::metadata(&staging).unwrap().permissions().mode() & 0o777, + 0o700 + ); + assert_eq!(std::fs::read(legacy.join("keep")).unwrap(), b"existing"); + std::fs::remove_dir_all(base).ok(); + } + + #[cfg(unix)] + #[test] + fn self_update_capability_requires_managed_writable_directories() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let base = temp_test_dir("storage-capability"); + let current_release = base.join("releases/v1.2.3"); + write_release_payload(¤t_release); + symlink(¤t_release, base.join("current")).expect("current link should be created"); + let base = std::fs::canonicalize(base).expect("base should canonicalize"); + std::fs::set_permissions(&base, std::fs::Permissions::from_mode(0o750)).unwrap(); + std::fs::set_permissions( + base.join("releases"), + std::fs::Permissions::from_mode(0o750), + ) + .unwrap(); + + assert!(self_update_storage_ready_at(&base)); + + std::fs::set_permissions( + base.join("releases"), + std::fs::Permissions::from_mode(0o777), + ) + .unwrap(); + assert!(!self_update_storage_ready_at(&base)); + + std::fs::set_permissions( + base.join("releases"), + std::fs::Permissions::from_mode(0o750), + ) + .unwrap(); + std::fs::remove_file(base.join("current")).unwrap(); + let outside = temp_test_dir("storage-capability-outside"); + std::fs::create_dir_all(&outside).unwrap(); + symlink(&outside, base.join("current")).expect("outside current link should be created"); + assert!(!self_update_storage_ready_at(&base)); + + std::fs::remove_dir_all(base).ok(); + std::fs::remove_dir_all(outside).ok(); + } + #[test] fn update_verifies_sha256sum_for_asset_name() { let data = b"release-bytes"; @@ -1196,5 +2731,21 @@ mod tests { "https://example.test/aether-v1.2.3-linux-amd64.tar.gz", ) .expect("sha256 should match"); + + let duplicate = format!( + "{expected} aether-v1.2.3-linux-amd64.tar.gz\n{expected} *aether-v1.2.3-linux-amd64.tar.gz\n" + ); + assert!(verify_sha256( + data, + &duplicate, + "https://example.test/aether-v1.2.3-linux-amd64.tar.gz", + ) + .is_err()); + assert!(verify_sha256( + data, + "not-a-hash aether-v1.2.3-linux-amd64.tar.gz\n", + "https://example.test/aether-v1.2.3-linux-amd64.tar.gz", + ) + .is_err()); } } diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/update_client.rs b/apps/aether-gateway/src/handlers/admin/system/shared/update_client.rs index 6658f8ef6..ae1bfe7d6 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/update_client.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/update_client.rs @@ -1,5 +1,9 @@ +use std::net::SocketAddr; +use std::sync::Arc; use std::time::Duration; +use reqwest::dns::{Addrs, Name, Resolve, Resolving}; + const EXPLICIT_UPDATE_PROXY_ENV_KEYS: &[&str] = &["AETHER_UPDATE_PROXY_URL", "UPDATE_PROXY_URL"]; const UPDATE_PROXY_ENV_KEYS: &[&str] = &[ @@ -20,8 +24,15 @@ pub(crate) fn build_update_http_client( timeout: Duration, label: &str, ) -> Result { - let mut builder = base_update_http_client_builder(timeout); - if let Some(proxy_url) = update_proxy_url_from_env() { + let proxy_url = update_proxy_url_from_env(); + let proxy_host = proxy_url + .as_deref() + .and_then(update_proxy_host) + .map(|host| host.trim_end_matches('.').to_ascii_lowercase()); + let mut builder = base_update_http_client_builder(timeout) + .no_proxy() + .dns_resolver(Arc::new(SafeUpdateDnsResolver::new(proxy_host))); + if let Some(proxy_url) = proxy_url { let proxy = reqwest::Proxy::all(proxy_url) .map_err(|_| format!("创建{label}代理失败,请检查更新代理环境变量"))? .no_proxy(reqwest::NoProxy::from_env()); @@ -38,6 +49,7 @@ pub(crate) fn build_direct_update_http_client( ) -> Result { base_update_http_client_builder(timeout) .no_proxy() + .dns_resolver(Arc::new(SafeUpdateDnsResolver::new(None))) .build() .map_err(|err| format!("创建{label}客户端失败: {err}")) } @@ -47,13 +59,112 @@ pub(crate) fn has_explicit_update_proxy_env() -> bool { } fn base_update_http_client_builder(timeout: Duration) -> reqwest::ClientBuilder { - reqwest::Client::builder().timeout(timeout) + reqwest::Client::builder() + .timeout(timeout) + .redirect(reqwest::redirect::Policy::custom(|attempt| { + if attempt.previous().len() >= 10 { + return attempt.error("too many update download redirects"); + } + if is_trusted_update_url(attempt.url()) { + attempt.follow() + } else { + attempt.error("update download redirected to an untrusted URL") + } + })) +} + +#[derive(Debug)] +struct SafeUpdateDnsResolver { + private_allowed_host: Option, +} + +impl SafeUpdateDnsResolver { + fn new(private_allowed_host: Option) -> Self { + Self { + private_allowed_host, + } + } +} + +impl Resolve for SafeUpdateDnsResolver { + fn resolve(&self, name: Name) -> Resolving { + let host = name.as_str().trim_end_matches('.').to_ascii_lowercase(); + let allow_private = self.private_allowed_host.as_deref() == Some(host.as_str()); + // Transparent DNS proxies may use RFC 2544's 198.18.0.0/15 benchmark + // range as a synthetic address. Update destinations are compiled-in + // GitHub hosts, so accepting that range + // for those exact hosts preserves proxy compatibility without opening + // the resolver to arbitrary custom destinations. + let allow_benchmarking_ip = is_trusted_update_host(&host); + Box::pin(async move { + let addresses = aether_http::lookup_host_with_limits( + host.as_str(), + 0, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|error| -> Box { Box::new(error) })?; + validate_update_resolved_addrs(&addresses, allow_private, allow_benchmarking_ip) + .map_err(|message| { + Box::new(std::io::Error::other(message)) + as Box + })?; + Ok(Box::new(addresses.into_iter()) as Addrs) + }) + } +} + +fn validate_update_resolved_addrs( + addresses: &[SocketAddr], + allow_private: bool, + allow_benchmarking_ip: bool, +) -> Result<(), &'static str> { + if addresses.is_empty() { + return Err("update DNS resolution returned no addresses"); + } + if !allow_private + && addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(allow_benchmarking_ip + && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) + { + return Err("update DNS resolution returned a private or reserved address"); + } + Ok(()) +} + +fn is_trusted_update_host(host: &str) -> bool { + host.eq_ignore_ascii_case("github.com") + || host.eq_ignore_ascii_case("api.github.com") + || host.eq_ignore_ascii_case("objects.githubusercontent.com") + || host.ends_with(".objects.githubusercontent.com") + || host.eq_ignore_ascii_case("release-assets.githubusercontent.com") + || host.ends_with(".release-assets.githubusercontent.com") +} + +pub(crate) fn is_trusted_update_url(url: &url::Url) -> bool { + if url.scheme() != "https" || !url.username().is_empty() || url.password().is_some() { + return false; + } + let Some(host) = url.host_str() else { + return false; + }; + is_trusted_update_host(host) } fn update_proxy_url_from_env() -> Option { read_nonempty_env_value(UPDATE_PROXY_ENV_KEYS) } +fn update_proxy_host(proxy_url: &str) -> Option { + let parsed = url::Url::parse(proxy_url) + .ok() + .filter(url::Url::has_host) + .or_else(|| url::Url::parse(&format!("http://{proxy_url}")).ok())?; + parsed.host_str().map(ToOwned::to_owned) +} + pub(crate) fn update_github_token_from_env() -> Option { read_nonempty_env_value(UPDATE_GITHUB_TOKEN_ENV_KEYS) } @@ -66,3 +177,72 @@ fn read_nonempty_env_value(keys: &[&str]) -> Option { .filter(|value| !value.is_empty()) }) } + +#[cfg(test)] +mod tests { + use super::{ + is_trusted_update_host, is_trusted_update_url, update_proxy_host, + validate_update_resolved_addrs, + }; + use std::net::SocketAddr; + + #[test] + fn update_url_trust_rejects_credentials_and_untrusted_hosts() { + for trusted in [ + "https://github.com/fawney19/Aether/releases/download/v1/aether.tar.gz", + "https://api.github.com/repos/fawney19/Aether/releases", + "https://objects.githubusercontent.com/github-production-release-asset/test", + "https://release-assets.githubusercontent.com/github-production-release-asset/test", + ] { + assert!(is_trusted_update_url(&url::Url::parse(trusted).unwrap())); + } + for untrusted in [ + "http://github.com/fawney19/Aether/releases/download/v1/aether.tar.gz", + "https://github.com.evil.example/aether.tar.gz", + "https://user@github.com/aether.tar.gz", + "https://example.com/aether.tar.gz", + ] { + assert!(!is_trusted_update_url(&url::Url::parse(untrusted).unwrap())); + } + } + + #[test] + fn update_dns_rejects_private_or_mixed_target_answers() { + let public = "8.8.8.8:443".parse::().unwrap(); + let private = "127.0.0.1:443".parse::().unwrap(); + + assert!(validate_update_resolved_addrs(&[public], false, false).is_ok()); + assert!(validate_update_resolved_addrs(&[private], false, false).is_err()); + assert!(validate_update_resolved_addrs(&[public, private], false, false).is_err()); + assert!(validate_update_resolved_addrs(&[private], true, false).is_ok()); + assert!(validate_update_resolved_addrs(&[], false, false).is_err()); + } + + #[test] + fn update_dns_allows_benchmarking_ip_only_for_trusted_github_hosts() { + let fake = "198.18.75.234:443".parse::().unwrap(); + assert!(validate_update_resolved_addrs(&[fake], false, true).is_ok()); + assert!(validate_update_resolved_addrs( + &[fake, "127.0.0.1:443".parse().unwrap()], + false, + true, + ) + .is_err()); + assert!(validate_update_resolved_addrs(&[fake], false, false).is_err()); + assert!(is_trusted_update_host("api.github.com")); + assert!(is_trusted_update_host("foo.objects.githubusercontent.com")); + assert!(!is_trusted_update_host("github.com.evil.example")); + } + + #[test] + fn update_proxy_host_supports_explicit_and_legacy_proxy_urls() { + assert_eq!( + update_proxy_host("http://user:secret@proxy.example.test:8080").as_deref(), + Some("proxy.example.test") + ); + assert_eq!( + update_proxy_host("127.0.0.1:7890").as_deref(), + Some("127.0.0.1") + ); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/helpers.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/helpers.rs index 98f9e8d3a..fd63a2433 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/helpers.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/helpers.rs @@ -1,9 +1,8 @@ use crate::handlers::admin::request::AdminAppState; -use crate::handlers::admin::shared::{ - attach_admin_audit_response, decrypt_catalog_secret_with_fallbacks, -}; +use crate::handlers::admin::shared::attach_admin_audit_response; use crate::handlers::shared::{ - api_key_placeholder_display, generate_gateway_api_key_plaintext, masked_gateway_api_key_display, + api_key_placeholder_display, generate_gateway_api_key_plaintext, + masked_gateway_api_key_display, open_auth_api_key_secret, }; use axum::{body::Body, response::Response}; use serde_json::json; @@ -17,16 +16,12 @@ pub(crate) fn format_optional_unix_secs_iso8601(value: Option) -> Option, - ciphertext: Option<&str>, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, ) -> String { - let Some(ciphertext) = ciphertext.map(str::trim).filter(|value| !value.is_empty()) else { + let Ok(projection) = open_auth_api_key_secret(state.app(), record) else { return api_key_placeholder_display(); }; - let Some(full_key) = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - else { - return api_key_placeholder_display(); - }; - masked_gateway_api_key_display(Some(full_key.as_str())) + masked_gateway_api_key_display(Some(projection.plaintext.as_str())) } pub(super) fn build_admin_user_api_key_detail_payload( @@ -37,7 +32,7 @@ pub(super) fn build_admin_user_api_key_detail_payload( json!({ "id": record.api_key_id, "name": record.name, - "key_display": masked_user_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": masked_user_api_key_display(state, record), "is_active": record.is_active, "is_locked": is_locked, "total_requests": record.total_requests, diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/create.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/create.rs index dfe0cb37d..f5e00d487 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/create.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/create.rs @@ -11,7 +11,9 @@ use super::super::helpers::{ use super::super::paths::admin_user_id_from_api_keys_path; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::admin::shared::mark_sensitive_admin_response_no_store; use crate::handlers::shared::normalize_optional_api_key_concurrent_limit; +use crate::handlers::shared::seal_auth_api_key_secret; use crate::GatewayError; use axum::{ body::Body, @@ -141,7 +143,16 @@ pub(crate) async fn build_admin_create_user_api_key_response( }; let plaintext_key = generate_admin_user_api_key_plaintext(); - let Some(key_encrypted) = state.encrypt_catalog_secret_with_fallbacks(&plaintext_key) else { + let api_key_id = uuid::Uuid::new_v4().to_string(); + let key_hash = hash_admin_user_api_key(&plaintext_key); + let Ok(key_encrypted) = seal_auth_api_key_secret( + state.app(), + &target_user_id, + &api_key_id, + &key_hash, + false, + &plaintext_key, + ) else { return Ok(( http::StatusCode::INTERNAL_SERVER_ERROR, Json(json!({ "detail": "API密钥加密失败" })), @@ -152,17 +163,18 @@ pub(crate) async fn build_admin_create_user_api_key_response( let Some(created) = state .create_user_api_key(aether_data::repository::auth::CreateUserApiKeyRecord { user_id: target_user_id.clone(), - api_key_id: uuid::Uuid::new_v4().to_string(), - key_hash: hash_admin_user_api_key(&plaintext_key), + api_key_id, + key_hash, key_encrypted: Some(key_encrypted), name: Some(name.clone()), - allowed_providers: None, + allowed_providers, allowed_api_formats: None, allowed_models: None, ip_rules, rate_limit, concurrent_limit, force_capabilities: None, + feature_settings, is_active: true, expires_at_unix_secs: None, auto_delete_on_expiry: false, @@ -175,56 +187,27 @@ pub(crate) async fn build_admin_create_user_api_key_response( return Ok(build_admin_users_data_unavailable_response()); }; - let created = if allowed_providers.is_some() { - match state - .set_user_api_key_allowed_providers( - &target_user_id, - &created.api_key_id, - allowed_providers, - ) - .await? - { - Some(updated) => updated, - None => created, - } - } else { - created - }; - let created = if feature_settings.is_some() { - match state - .set_user_api_key_feature_settings( - &target_user_id, - &created.api_key_id, - feature_settings.clone(), - ) - .await? - { - Some(updated) => updated, - None => created, - } - } else { - created - }; - - Ok(attach_audit_response( - Json(json!({ - "id": created.api_key_id, - "key": plaintext_key, - "name": created.name, - "key_display": masked_user_api_key_display(state, created.key_encrypted.as_deref()), - "rate_limit": created.rate_limit, - "concurrent_limit": created.concurrent_limit, - "ip_rules": created.ip_rules, - "expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs), - "last_used_at": format_optional_unix_secs_iso8601(created.last_used_at_unix_secs), - "created_at": format_optional_unix_secs_iso8601(created.created_at_unix_secs), - "feature_settings": created.feature_settings, - "message": "API Key创建成功,请妥善保存完整密钥", - })) - .into_response(), - "admin_user_api_key_created", - "create_user_api_key", - "user_api_key", - &created.api_key_id, + Ok(mark_sensitive_admin_response_no_store( + attach_audit_response( + Json(json!({ + "id": created.api_key_id, + "key": plaintext_key, + "name": created.name, + "key_display": masked_user_api_key_display(state, &created), + "rate_limit": created.rate_limit, + "concurrent_limit": created.concurrent_limit, + "ip_rules": created.ip_rules, + "expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs), + "last_used_at": format_optional_unix_secs_iso8601(created.last_used_at_unix_secs), + "created_at": format_optional_unix_secs_iso8601(created.created_at_unix_secs), + "feature_settings": created.feature_settings, + "message": "API Key创建成功,请妥善保存完整密钥", + })) + .into_response(), + "admin_user_api_key_created", + "create_user_api_key", + "user_api_key", + &created.api_key_id, + ), )) } 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 50f3ee248..dc328761c 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 @@ -56,7 +56,7 @@ pub(crate) async fn build_admin_list_user_api_keys_response( json!({ "id": record.api_key_id, "name": record.name, - "key_display": masked_user_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": masked_user_api_key_display(state, &record), "is_active": record.is_active, "is_locked": is_locked, "total_requests": record.total_requests, diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/reveal.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/reveal.rs index 704ccaaf4..ddbb172a6 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/reveal.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/reveal.rs @@ -3,6 +3,8 @@ use super::super::helpers::attach_audit_response; use super::super::paths::admin_user_api_key_full_key_parts; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::handlers::admin::shared::mark_sensitive_admin_response_no_store; +use crate::handlers::shared::decrypt_or_migrate_auth_api_key_secret; use crate::GatewayError; use axum::{ body::Body, @@ -51,19 +53,24 @@ pub(crate) async fn build_admin_reveal_user_api_key_response( .into_response()); } - let Some(full_key) = state.decrypt_catalog_secret_with_fallbacks(ciphertext) else { - return Ok(( - http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": "解密密钥失败" })), - ) - .into_response()); + let full_key = match decrypt_or_migrate_auth_api_key_secret(state.app(), &record).await { + Ok(value) => value, + Err(_) => { + return Ok(( + http::StatusCode::INTERNAL_SERVER_ERROR, + Json(json!({ "detail": "解密或校验密钥失败" })), + ) + .into_response()) + } }; - Ok(attach_audit_response( - Json(json!({ "key": full_key })).into_response(), - "admin_user_api_key_revealed", - "reveal_user_api_key", - "user_api_key", - &key_id, + Ok(mark_sensitive_admin_response_no_store( + attach_audit_response( + Json(json!({ "key": full_key })).into_response(), + "admin_user_api_key_revealed", + "reveal_user_api_key", + "user_api_key", + &key_id, + ), )) } diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs index 29956628f..08668a6c8 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs @@ -107,15 +107,24 @@ pub(crate) async fn build_admin_update_user_api_key_response( }, None => None, }; + let name_present = name.is_some(); + let rate_limit_present = payload.rate_limit.is_some(); + let concurrent_limit_present = concurrent_limit.is_some(); let Some(updated) = state .update_user_api_key_basic(aether_data::repository::auth::UpdateUserApiKeyBasicRecord { user_id: user_id.clone(), api_key_id: api_key_id.clone(), + key_encrypted: None, + key_encrypted_present: false, name, + name_present, rate_limit: payload.rate_limit, + rate_limit_present, concurrent_limit, + concurrent_limit_present, ip_rules, + feature_settings, }) .await? else { @@ -125,14 +134,6 @@ pub(crate) async fn build_admin_update_user_api_key_response( ) .into_response()); }; - let updated = if let Some(feature_settings) = feature_settings { - state - .set_user_api_key_feature_settings(&user_id, &api_key_id, feature_settings) - .await? - .unwrap_or(updated) - } else { - updated - }; let is_locked = state .list_auth_api_key_snapshots_by_ids(std::slice::from_ref(&api_key_id)) diff --git a/apps/aether-gateway/src/handlers/admin/users/batch.rs b/apps/aether-gateway/src/handlers/admin/users/batch.rs index 72274bdcf..75d945001 100644 --- a/apps/aether-gateway/src/handlers/admin/users/batch.rs +++ b/apps/aether-gateway/src/handlers/admin/users/batch.rs @@ -1,6 +1,7 @@ use super::{ - build_admin_users_bad_request_response, build_admin_users_read_only_response, - disabled_user_policy_detail, disabled_user_policy_field, normalize_admin_user_role, + 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, + 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; @@ -131,6 +132,24 @@ 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 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)), + }; + let mutates_admin_account = mutation + .role + .as_deref() + .is_some_and(crate::roles::can_access_admin_console) + || (mutation.has_auth_user_fields() + && resolved + .items + .iter() + .any(|item| crate::roles::can_access_admin_console(&item.role))); + if mutates_admin_account && !management_token_may_administer_user_accounts(request_context) { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); + } if mutation.has_auth_user_fields() && !state.has_auth_user_write_capability() { return Ok(build_admin_users_read_only_response( "当前为只读模式,无法批量更新用户", @@ -141,10 +160,6 @@ pub(in super::super) async fn build_admin_user_batch_action_response( "当前为只读模式,无法批量更新用户钱包", )); } - 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)), - }; let active_admin_demotions = count_active_admin_demotions(&mutation, &resolved.items); let active_admin_count = if active_admin_demotions > 0 { state.count_active_admin_users().await? diff --git a/apps/aether-gateway/src/handlers/admin/users/billing.rs b/apps/aether-gateway/src/handlers/admin/users/billing.rs index b12c67593..b30cbcd19 100644 --- a/apps/aether-gateway/src/handlers/admin/users/billing.rs +++ b/apps/aether-gateway/src/handlers/admin/users/billing.rs @@ -1,4 +1,5 @@ use super::{build_admin_users_bad_request_response, build_admin_users_data_unavailable_response}; +use crate::handlers::admin::billing::admin_payment_gateway_response_projection; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339}; use crate::handlers::shared::unix_ms_to_rfc3339; @@ -34,6 +35,22 @@ fn admin_user_id_from_billing_path(request_path: &str, suffix: &str) -> Option Option<(String, String)> { + let rest = request_path + .trim_end_matches('/') + .strip_prefix("/api/admin/users/")?; + let mut parts = rest.split('/'); + let user_id = parts.next()?.trim(); + if parts.next()? != "billing" || parts.next()? != "entitlements" { + return None; + } + let entitlement_id = parts.next()?.trim(); + if user_id.is_empty() || entitlement_id.is_empty() || parts.next().is_some() { + return None; + } + Some((user_id.to_string(), entitlement_id.to_string())) +} + fn admin_user_billing_operator_id(request_context: &AdminRequestContext<'_>) -> Option { request_context .decision() @@ -102,7 +119,7 @@ fn plan_has_package_rights(record: &BillingPlanRecord) -> bool { items.iter().any(|item| { matches!( item.get("type").and_then(|value| value.as_str()), - Some("daily_quota" | "membership_group") + Some("daily_quota" | "membership_group" | "usage_policy") ) }) }) @@ -122,7 +139,8 @@ fn admin_payment_order_payload(record: &crate::AdminWalletPaymentOrderRecord) -> "refundable_amount_usd": record.refundable_amount_usd, "payment_method": record.payment_method, "gateway_order_id": record.gateway_order_id, - "gateway_response": record.gateway_response, + "gateway_response": admin_payment_gateway_response_projection(record.gateway_response.as_ref()), + "has_gateway_response": record.gateway_response.is_some(), "status": record.status, "order_kind": "plan_purchase", "created_at": unix_ms_to_rfc3339(record.created_at_unix_ms), @@ -202,6 +220,60 @@ pub(in super::super) async fn build_admin_list_user_billing_entitlements_respons } } +pub(in super::super) async fn build_admin_revoke_user_billing_entitlement_response( + state: &AdminAppState<'_>, + request_context: &AdminRequestContext<'_>, +) -> Result, GatewayError> { + let Some((user_id, entitlement_id)) = + admin_user_entitlement_ids_from_path(request_context.path()) + else { + return Ok(build_admin_users_bad_request_response("缺少套餐权益 ID")); + }; + if state.find_user_auth_by_id(&user_id).await?.is_none() { + return Ok(( + http::StatusCode::NOT_FOUND, + Json(json!({ "detail": "用户不存在" })), + ) + .into_response()); + } + match state + .app() + .revoke_user_plan_entitlement(&user_id, &entitlement_id) + .await? + { + crate::LocalMutationOutcome::Applied(()) => {} + crate::LocalMutationOutcome::NotFound => { + return Ok(( + http::StatusCode::NOT_FOUND, + Json(json!({ "detail": "套餐权益不存在或已失效" })), + ) + .into_response()); + } + crate::LocalMutationOutcome::Invalid(detail) => { + return Ok(build_admin_users_bad_request_response(detail)); + } + crate::LocalMutationOutcome::Unavailable => { + return Ok(build_admin_users_data_unavailable_response()); + } + } + let entitlements = match load_admin_user_entitlements_payload(state, &user_id).await? { + Some(value) => value, + None => return Ok(build_admin_users_data_unavailable_response()), + }; + Ok(attach_admin_audit_response( + Json(json!({ + "items": entitlements["items"].clone(), + "entitlements": entitlements["items"].clone(), + "total": entitlements["total"].clone(), + })) + .into_response(), + "admin_user_plan_revoked", + "revoke_user_billing_entitlement", + "user_plan_entitlement", + &entitlement_id, + )) +} + pub(in super::super) async fn build_admin_grant_user_billing_plan_response( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, @@ -264,7 +336,10 @@ pub(in super::super) async fn build_admin_grant_user_billing_plan_response( user_id: user_id.clone(), amount_usd: 0.0, pay_amount: 0.0, - pay_currency: plan.price_currency.clone(), + // Manual grants have no provider settlement. Keep the order + // currency stable and independent of the plan's display + // currency (the amount is always zero). + pay_currency: "USD".to_string(), exchange_rate: 1.0, payment_method: "admin_grant".to_string(), payment_provider: Some("admin".to_string()), @@ -304,7 +379,7 @@ pub(in super::super) async fn build_admin_grant_user_billing_plan_response( &order.id, Some(&order_no), Some(0.0), - Some(&plan.price_currency), + Some("USD"), Some(1.0), Some(json!({ "admin_grant": true })), operator_id.as_deref(), diff --git a/apps/aether-gateway/src/handlers/admin/users/lifecycle/create.rs b/apps/aether-gateway/src/handlers/admin/users/lifecycle/create.rs index ad5ea3f30..61b55ac7c 100644 --- a/apps/aether-gateway/src/handlers/admin/users/lifecycle/create.rs +++ b/apps/aether-gateway/src/handlers/admin/users/lifecycle/create.rs @@ -1,6 +1,7 @@ use super::super::{ - admin_default_user_initial_gift, build_admin_users_read_only_response, - disabled_user_policy_detail, disabled_user_policy_field, normalize_admin_feature_settings, + admin_default_user_initial_gift, build_admin_users_permission_denied_response, + build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field, + management_token_may_administer_user_accounts, normalize_admin_feature_settings, normalize_admin_optional_user_email, normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_username, validate_admin_user_password, AdminCreateUserRequest, }; @@ -18,7 +19,7 @@ use serde_json::{json, Value}; pub(in super::super) async fn build_admin_create_user_response( state: &AdminAppState<'_>, - _request_context: &AdminRequestContext<'_>, + request_context: &AdminRequestContext<'_>, request_body: Option<&axum::body::Bytes>, ) -> Result, GatewayError> { if !state.has_auth_user_write_capability() { @@ -107,6 +108,13 @@ pub(in super::super) async fn build_admin_create_user_response( .into_response()) } }; + if crate::roles::can_access_admin_console(&role) + && !management_token_may_administer_user_accounts(request_context) + { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); + } let password_policy = admin_user_password_policy(state).await?; if let Err(detail) = validate_admin_user_password(&payload.password, &password_policy) { return Ok(( @@ -206,25 +214,152 @@ pub(in super::super) async fn build_admin_create_user_response( )); }; - if state - .initialize_auth_user_wallet(&user.id, initial_gift_usd, payload.unlimited) - .await? - .is_none() + let initialized = match state + .initialize_auth_user_wallet_with_outcome(&user.id, initial_gift_usd, payload.unlimited) + .await { - return Ok(build_admin_users_read_only_response( - "当前为只读模式,无法初始化用户钱包", + Ok(Some(initialized)) => initialized, + Ok(None) => { + // The user row was created before wallet provisioning. Remove it + // when the backend reports that wallet initialization is unavailable; + // the guarded rollback refuses to delete a concurrently-created or + // funded wallet. + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + "admin user wallet-unavailable cleanup failed" + ); + } + return Ok(build_admin_users_read_only_response( + "当前为只读模式,无法初始化用户钱包", + )); + } + Err(err) => { + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + "admin user wallet initialization cleanup failed" + ); + } + return Err(err); + } + }; + let owned_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let wallet_is_user_owned = initialized.wallet.user_id.as_deref() == Some(user.id.as_str()) + && initialized.wallet.api_key_id.is_none(); + if !wallet_is_user_owned { + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + wallet_id = ?owned_wallet_id, + "admin user wallet owner validation cleanup failed" + ); + } + return Err(GatewayError::Internal( + "admin user wallet owner does not match the provisioned user".to_string(), )); } if !group_ids.is_empty() { - state + let replaced_groups = match state .replace_user_groups_for_user(&user.id, &group_ids) - .await?; + .await + { + Ok(groups) if groups.len() == group_ids.len() => groups, + Ok(_) => { + let error = + GatewayError::Internal("user groups could not be persisted".to_string()); + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + wallet_id = ?owned_wallet_id, + "admin user group provisioning cleanup failed" + ); + } + return Err(error); + } + Err(error) => { + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + wallet_id = ?owned_wallet_id, + "admin user group provisioning cleanup failed" + ); + } + return Err(error); + } + }; + let _ = replaced_groups; } - let feature_settings = if feature_settings.is_some() { - state - .update_user_feature_settings(&user.id, feature_settings.clone()) - .await? - .or(feature_settings) + let feature_settings = if let Some(requested_feature_settings) = feature_settings { + match state + .update_user_feature_settings(&user.id, Some(requested_feature_settings)) + .await + { + Ok(Some(updated_feature_settings)) => Some(updated_feature_settings), + Ok(None) => { + let error = GatewayError::Internal( + "user feature settings could not be persisted".to_string(), + ); + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + wallet_id = ?owned_wallet_id, + "admin user feature settings cleanup failed" + ); + } + return Err(error); + } + Err(error) => { + if let Err(cleanup_error) = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await + { + tracing::error!( + error = ?cleanup_error, + user_id = %user.id, + wallet_id = ?owned_wallet_id, + "admin user feature settings cleanup failed" + ); + } + return Err(error); + } + } } else { None }; diff --git a/apps/aether-gateway/src/handlers/admin/users/lifecycle/delete.rs b/apps/aether-gateway/src/handlers/admin/users/lifecycle/delete.rs index 33ee805de..8e79d5119 100644 --- a/apps/aether-gateway/src/handlers/admin/users/lifecycle/delete.rs +++ b/apps/aether-gateway/src/handlers/admin/users/lifecycle/delete.rs @@ -1,4 +1,7 @@ -use super::super::build_admin_users_bad_request_response; +use super::super::{ + build_admin_users_bad_request_response, build_admin_users_permission_denied_response, + management_token_may_administer_user_accounts, +}; use super::support::admin_user_id_from_detail_path; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::attach_admin_audit_response; @@ -26,7 +29,19 @@ pub(in super::super) async fn build_admin_delete_user_response( .into_response()); }; - if user.role.eq_ignore_ascii_case("admin") && state.count_active_admin_users().await? <= 1 { + if crate::roles::can_access_admin_console(&user.role) + && !management_token_may_administer_user_accounts(request_context) + { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); + } + + if user.is_active + && !user.is_deleted + && user.role.eq_ignore_ascii_case("admin") + && state.count_active_admin_users().await? <= 1 + { return Ok(( http::StatusCode::BAD_REQUEST, Json(json!({ "detail": "不能删除最后一个管理员账户" })), diff --git a/apps/aether-gateway/src/handlers/admin/users/lifecycle/update.rs b/apps/aether-gateway/src/handlers/admin/users/lifecycle/update.rs index 0cb174166..dcbac4efd 100644 --- a/apps/aether-gateway/src/handlers/admin/users/lifecycle/update.rs +++ b/apps/aether-gateway/src/handlers/admin/users/lifecycle/update.rs @@ -1,9 +1,10 @@ use super::super::{ build_admin_users_bad_request_response, build_admin_users_data_unavailable_response, - build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field, - normalize_admin_feature_settings, normalize_admin_optional_user_email, - normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_username, - validate_admin_user_password, AdminUpdateUserPatch, + build_admin_users_permission_denied_response, build_admin_users_read_only_response, + disabled_user_policy_detail, disabled_user_policy_field, + management_token_may_administer_user_accounts, normalize_admin_feature_settings, + normalize_admin_optional_user_email, normalize_admin_user_group_ids, normalize_admin_user_role, + normalize_admin_username, validate_admin_user_password, AdminUpdateUserPatch, }; use super::support::{ admin_user_id_from_detail_path, admin_user_password_policy, @@ -70,6 +71,7 @@ pub(in super::super) async fn build_admin_update_user_response( } }; let (field_presence, payload) = patch.into_parts(); + let email_present = field_presence.contains("email"); let feature_settings = if field_presence.contains("feature_settings") { match normalize_admin_feature_settings(payload.feature_settings.flatten()) { Ok(value) => Some(value), @@ -150,16 +152,47 @@ pub(in super::super) async fn build_admin_update_user_response( }, None => None, }; + let changes_admin_role = role.as_deref().is_some_and(|role| { + !role.eq_ignore_ascii_case(&existing_user.role) + && (crate::roles::can_access_admin_console(&existing_user.role) + || crate::roles::can_access_admin_console(role)) + }); + let resets_admin_password = + payload.password.is_some() && crate::roles::can_access_admin_console(&existing_user.role); + let changes_admin_active_state = payload + .is_active + .is_some_and(|is_active| is_active != existing_user.is_active) + && crate::roles::can_access_admin_console(&existing_user.role); + let mutates_existing_admin = crate::roles::can_access_admin_console(&existing_user.role) + && (email.is_some() + || username.is_some() + || payload.password.is_some() + || role.is_some() + || payload.is_active.is_some() + || field_presence.contains("group_ids") + || field_presence.contains("feature_settings") + || payload.unlimited.is_some()); + if (changes_admin_role + || resets_admin_password + || changes_admin_active_state + || mutates_existing_admin) + && !management_token_may_administer_user_accounts(request_context) + { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); + } if existing_user.is_active && crate::roles::is_full_admin_role(&existing_user.role) - && role + && (role .as_deref() .is_some_and(|role| !crate::roles::is_full_admin_role(role)) + || payload.is_active == Some(false)) && state.count_active_admin_users().await? <= 1 { return Ok(( http::StatusCode::BAD_REQUEST, - Json(json!({ "detail": "不能降级最后一个管理员账户" })), + Json(json!({ "detail": "不能降级或停用最后一个管理员账户" })), ) .into_response()); } @@ -193,7 +226,7 @@ pub(in super::super) async fn build_admin_update_user_response( } } } - let needs_auth_user_write = email.is_some() + let needs_auth_user_write = email_present || username.is_some() || payload.password.is_some() || role.is_some() @@ -211,9 +244,21 @@ pub(in super::super) async fn build_admin_update_user_response( )); } - if email.is_some() || username.is_some() { + if email_present || username.is_some() { + let resets_email_verification = email.as_deref().is_some_and(|email| { + existing_user + .email + .as_deref() + .is_none_or(|existing| !existing.trim().eq_ignore_ascii_case(email.trim())) + }); if state - .update_local_auth_user_profile(&user_id, email.clone(), username.clone()) + .update_local_auth_user_profile( + &user_id, + email_present, + email.clone(), + resets_email_verification.then_some(false), + username.clone(), + ) .await? .is_none() { @@ -249,10 +294,13 @@ pub(in super::super) async fn build_admin_update_user_response( .into_response()) } }; - if state - .update_local_auth_user_password_hash(&user_id, password_hash, chrono::Utc::now()) + if !state + .reset_local_auth_user_password_and_revoke_sessions( + &user_id, + password_hash, + chrono::Utc::now(), + ) .await? - .is_none() { return Ok(( http::StatusCode::NOT_FOUND, diff --git a/apps/aether-gateway/src/handlers/admin/users/mod.rs b/apps/aether-gateway/src/handlers/admin/users/mod.rs index 957cb7cf0..e5c854748 100644 --- a/apps/aether-gateway/src/handlers/admin/users/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/users/mod.rs @@ -28,6 +28,7 @@ use self::batch::{ use self::billing::{ build_admin_grant_user_billing_plan_response, build_admin_list_user_billing_entitlements_response, + build_admin_revoke_user_billing_entitlement_response, }; use self::groups::{ build_admin_create_user_group_response, build_admin_delete_user_group_response, @@ -47,9 +48,10 @@ use self::sessions::{ 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_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, + 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, diff --git a/apps/aether-gateway/src/handlers/admin/users/routes.rs b/apps/aether-gateway/src/handlers/admin/users/routes.rs index ed0c2e4ca..9ec0a0186 100644 --- a/apps/aether-gateway/src/handlers/admin/users/routes.rs +++ b/apps/aether-gateway/src/handlers/admin/users/routes.rs @@ -8,10 +8,11 @@ use super::{ build_admin_list_user_group_members_response, build_admin_list_user_groups_response, build_admin_list_user_sessions_response, build_admin_list_users_response, build_admin_replace_user_group_members_response, build_admin_resolve_user_selection_response, - build_admin_reveal_user_api_key_response, build_admin_set_default_user_group_response, - build_admin_toggle_user_api_key_lock_response, build_admin_update_user_api_key_response, - build_admin_update_user_group_response, build_admin_update_user_response, - build_admin_user_batch_action_response, build_admin_users_data_unavailable_response, + build_admin_reveal_user_api_key_response, build_admin_revoke_user_billing_entitlement_response, + build_admin_set_default_user_group_response, build_admin_toggle_user_api_key_lock_response, + build_admin_update_user_api_key_response, build_admin_update_user_group_response, + build_admin_update_user_response, build_admin_user_batch_action_response, + build_admin_users_data_unavailable_response, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::GatewayError; @@ -58,6 +59,10 @@ fn is_admin_users_route(request_context: &AdminRequestContext<'_>) -> bool { && path.starts_with("/api/admin/users/") && path.ends_with("/billing/grant-plan") && path.matches('/').count() == 6) + || (request_context.method() == http::Method::DELETE + && path.starts_with("/api/admin/users/") + && path.contains("/billing/entitlements/") + && path.matches('/').count() == 7) || ((request_context.method() == http::Method::GET || request_context.method() == http::Method::PUT || request_context.method() == http::Method::DELETE) @@ -155,6 +160,9 @@ pub(super) async fn maybe_build_local_admin_users_routes_response( build_admin_grant_user_billing_plan_response(state, request_context, request_body) .await?, )), + Some("revoke_user_billing_entitlement") => Ok(Some( + build_admin_revoke_user_billing_entitlement_response(state, request_context).await?, + )), Some("get_user") => Ok(Some( build_admin_get_user_response(state, request_context).await?, )), diff --git a/apps/aether-gateway/src/handlers/admin/users/sessions.rs b/apps/aether-gateway/src/handlers/admin/users/sessions.rs index 9904ef602..c2cbfdb10 100644 --- a/apps/aether-gateway/src/handlers/admin/users/sessions.rs +++ b/apps/aether-gateway/src/handlers/admin/users/sessions.rs @@ -1,4 +1,7 @@ -use super::{build_admin_users_bad_request_response, format_optional_datetime_iso8601}; +use super::{ + build_admin_users_bad_request_response, build_admin_users_permission_denied_response, + format_optional_datetime_iso8601, management_token_may_administer_user_accounts, +}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::attach_admin_audit_response; use crate::{GatewayError, GatewayUserSessionView}; @@ -98,12 +101,19 @@ pub(super) async fn build_admin_delete_user_session_response( )); }; - if state.find_user_auth_by_id(&user_id).await?.is_none() { + let Some(user) = state.find_user_auth_by_id(&user_id).await? else { return Ok(( http::StatusCode::NOT_FOUND, Json(json!({ "detail": "用户不存在" })), ) .into_response()); + }; + if crate::roles::can_access_admin_console(&user.role) + && !management_token_may_administer_user_accounts(request_context) + { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); } if state @@ -144,12 +154,19 @@ pub(super) async fn build_admin_delete_user_sessions_response( return Ok(build_admin_users_bad_request_response("缺少 user_id")); }; - if state.find_user_auth_by_id(&user_id).await?.is_none() { + let Some(user) = state.find_user_auth_by_id(&user_id).await? else { return Ok(( http::StatusCode::NOT_FOUND, Json(json!({ "detail": "用户不存在" })), ) .into_response()); + }; + if crate::roles::can_access_admin_console(&user.role) + && !management_token_may_administer_user_accounts(request_context) + { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); } let revoked_count = state diff --git a/apps/aether-gateway/src/handlers/admin/users/shared.rs b/apps/aether-gateway/src/handlers/admin/users/shared.rs index dbae70da8..b304932c9 100644 --- a/apps/aether-gateway/src/handlers/admin/users/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/users/shared.rs @@ -66,7 +66,7 @@ pub(super) struct AdminToggleUserApiKeyLockRequest { pub(super) locked: Option, } -#[derive(Debug, serde::Deserialize)] +#[derive(serde::Deserialize)] pub(super) struct AdminCreateUserRequest { pub(super) username: String, pub(super) password: String, @@ -84,7 +84,7 @@ pub(super) struct AdminCreateUserRequest { pub(super) feature_settings: Option, } -#[derive(Debug, serde::Deserialize)] +#[derive(serde::Deserialize)] pub(super) struct AdminUpdateUserRequest { #[serde(default)] pub(super) email: Option, @@ -157,6 +157,41 @@ pub(super) fn build_admin_users_bad_request_response(detail: impl Into) .into_response() } +pub(super) fn management_token_may_administer_user_accounts( + request_context: &crate::handlers::admin::request::AdminRequestContext<'_>, +) -> bool { + request_context.decision().is_some_and(|decision| { + crate::control::management_token_principal_has_permission(decision, "admin:users:admin") + }) +} + +pub(super) fn build_admin_users_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_permission": "admin:users:admin", + "route_family": request_context.route_family(), + "route_kind": request_context.route_kind(), + "request_path": request_context.path(), + })), + ) + .into_response(), + "admin_user_account_permission_denied", + "permission_denied", + "admin_user_account", + actor_id, + ) +} + pub(super) fn normalize_admin_optional_user_email( value: Option<&str>, ) -> Result, String> { diff --git a/apps/aether-gateway/src/handlers/internal/gateway.rs b/apps/aether-gateway/src/handlers/internal/gateway.rs index 05560c13f..03d115acc 100644 --- a/apps/aether-gateway/src/handlers/internal/gateway.rs +++ b/apps/aether-gateway/src/handlers/internal/gateway.rs @@ -4,9 +4,9 @@ use super::{ build_internal_gateway_header_map, build_internal_gateway_passthrough_payload, build_internal_gateway_proxy_public_response, build_internal_gateway_request_parts, build_internal_gateway_resolve_payload, build_internal_gateway_uri, - build_internal_tunnel_heartbeat_ack, build_management_token_payload, gateway_error_message, - maybe_build_internal_finalize_video_response, parse_internal_tunnel_heartbeat_request, - parse_internal_tunnel_node_status_request, + build_internal_tunnel_heartbeat_ack, build_management_token_payload, + internal_finalize_report_kind_is_supported, maybe_build_internal_finalize_video_response, + parse_internal_tunnel_heartbeat_request, parse_internal_tunnel_node_status_request, }; use crate::ai_serving::api; use crate::constants::{ @@ -19,6 +19,7 @@ use crate::execution_runtime::{execute_execution_runtime_stream, execute_executi use crate::handlers::shared::{ InternalGatewayAuthContextRequest, InternalGatewayExecuteRequest, InternalGatewayResolveRequest, }; +use crate::tunnel::{claim_tunnel_heartbeat, finish_tunnel_heartbeat_claim}; use crate::tunnel::{is_tunnel_heartbeat_path, is_tunnel_node_status_path, TUNNEL_ROUTE_FAMILY}; use crate::{AppState, GatewayError}; use aether_data::repository::proxy_nodes::{ @@ -28,31 +29,70 @@ use axum::body::{Body, Bytes}; use axum::http::{self, HeaderName, HeaderValue, Response}; use axum::response::IntoResponse; use axum::Json; -use serde_json::json; +use serde_json::{json, Value}; -async fn apply_supplied_auth_context( - state: &AppState, - decision: &mut GatewayControlDecision, - auth_context: Option, -) -> Result { - let Some(auth_context) = auth_context else { - return Ok(false); - }; - let refreshed = crate::control::refresh_execution_runtime_auth_context( - state, - auth_context, - decision.auth_endpoint_signature.as_deref(), +fn reject_supplied_auth_context( + auth_context: Option<&crate::control::GatewayControlAuthContext>, +) -> Result<(), Response> { + if auth_context.is_some() { + return Err(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "supplied auth_context is not accepted; authenticate through request headers", + )); + } + Ok(()) +} + +fn internal_gateway_data_error_response(operation: &'static str) -> Response { + tracing::error!( + event_name = "internal_gateway_data_error", + operation, + error_category = "repository_unavailable", + "internal gateway data operation failed" + ); + build_internal_control_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "internal gateway data unavailable", ) - .await?; - decision.local_auth_rejection = refreshed.local_rejection.clone(); - decision.auth_context = Some(refreshed); - Ok(true) +} + +async fn resolve_bound_internal_report_context( + state: &AppState, + trace_id: &str, + report_kind: &str, + report_context: Option<&Value>, + operation: &'static str, +) -> Result> { + match crate::usage::resolve_bound_internal_gateway_report_context( + state, + trace_id, + report_kind, + report_context, + ) + .await + { + Ok(Some(report_context)) => Ok(report_context), + Ok(None) => { + tracing::warn!( + event_name = "internal_gateway_report_context_rejected", + operation, + error_category = "unbound_report_context", + "internal gateway report context did not carry a valid planner capability" + ); + Err(build_internal_control_error_response( + http::StatusCode::CONFLICT, + "internal gateway report context does not carry a valid planner capability", + )) + } + Err(_) => Err(internal_gateway_data_error_response(operation)), + } } pub(crate) async fn maybe_build_local_internal_proxy_response_impl( state: &AppState, request_context: &GatewayPublicRequestContext, remote_addr: &std::net::SocketAddr, + request_headers: &http::HeaderMap, request_body: Option<&Bytes>, ) -> Result>, GatewayError> { let Some(decision) = request_context.control_decision.as_ref() else { @@ -62,12 +102,40 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( return Ok(None); } if decision.route_family.as_deref() == Some("internal_gateway") { - if !remote_addr.ip().is_loopback() { + if request_context.request_method != http::Method::POST + || decision.route_kind.as_deref() == Some("unhandled") + { return Ok(Some(build_internal_control_error_response( - http::StatusCode::FORBIDDEN, - "loopback access only", + http::StatusCode::NOT_FOUND, + "route not found", ))); } + let authenticated_body = request_body.map_or(&[][..], Bytes::as_ref); + if let Err(error) = crate::internal_gateway_auth::authenticate_internal_gateway_request( + state, + remote_addr, + &request_context.request_method, + &request_context.request_path_and_query(), + request_headers, + authenticated_body, + ) + .await + { + let (status, message) = match error { + crate::internal_gateway_auth::InternalGatewayAuthError::Disabled => { + (http::StatusCode::NOT_FOUND, "route not found") + } + crate::internal_gateway_auth::InternalGatewayAuthError::Invalid => ( + http::StatusCode::FORBIDDEN, + "invalid internal gateway authentication", + ), + crate::internal_gateway_auth::InternalGatewayAuthError::Unavailable => ( + http::StatusCode::SERVICE_UNAVAILABLE, + "internal gateway authentication unavailable", + ), + }; + return Ok(Some(build_internal_control_error_response(status, message))); + } match decision.route_kind.as_deref() { Some("resolve") if request_context.request_path == "/api/internal/gateway/resolve" => { let Some(request_body) = request_body else { @@ -192,6 +260,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -225,23 +296,14 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( Json(build_internal_gateway_fallback_plan_payload(None)).into_response(), )); }; - let provided_auth_context = - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; let auth_context = resolved.auth_context.as_ref(); if auth_context .map(|value| !value.access_allowed) .unwrap_or(true) { - let fallback_auth_context = if !provided_auth_context { - auth_context - } else { - None - }; return Ok(Some( - Json(build_internal_gateway_fallback_plan_payload( - fallback_auth_context, - )) - .into_response(), + Json(build_internal_gateway_fallback_plan_payload(auth_context)) + .into_response(), )); } let Some(mut local_payload) = api::maybe_build_sync_decision_payload( @@ -255,21 +317,20 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ) .await? else { - let fallback_auth_context = if !provided_auth_context { - auth_context - } else { - None - }; return Ok(Some( - Json(build_internal_gateway_fallback_plan_payload( - fallback_auth_context, - )) - .into_response(), + Json(build_internal_gateway_fallback_plan_payload(auth_context)) + .into_response(), )); }; - if provided_auth_context { - local_payload.auth_context = None; - } + let report_kind = local_payload.report_kind.clone(); + crate::usage::attach_internal_gateway_report_capability( + state, + trace_id.as_str(), + report_kind.as_deref(), + &local_payload.provider_request_headers, + &mut local_payload.report_context, + ) + .await?; return Ok(Some(Json(local_payload).into_response())); } Some("decision_stream") @@ -291,6 +352,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -324,23 +388,14 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( Json(build_internal_gateway_fallback_plan_payload(None)).into_response(), )); }; - let provided_auth_context = - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; let auth_context = resolved.auth_context.as_ref(); if auth_context .map(|value| !value.access_allowed) .unwrap_or(true) { - let fallback_auth_context = if !provided_auth_context { - auth_context - } else { - None - }; return Ok(Some( - Json(build_internal_gateway_fallback_plan_payload( - fallback_auth_context, - )) - .into_response(), + Json(build_internal_gateway_fallback_plan_payload(auth_context)) + .into_response(), )); } let Some(mut local_payload) = api::maybe_build_stream_decision_payload( @@ -353,21 +408,20 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ) .await? else { - let fallback_auth_context = if !provided_auth_context { - auth_context - } else { - None - }; return Ok(Some( - Json(build_internal_gateway_fallback_plan_payload( - fallback_auth_context, - )) - .into_response(), + Json(build_internal_gateway_fallback_plan_payload(auth_context)) + .into_response(), )); }; - if provided_auth_context { - local_payload.auth_context = None; - } + let report_kind = local_payload.report_kind.clone(); + crate::usage::attach_internal_gateway_report_capability( + state, + trace_id.as_str(), + report_kind.as_deref(), + &local_payload.provider_request_headers, + &mut local_payload.report_context, + ) + .await?; return Ok(Some(Json(local_payload).into_response())); } Some("plan_sync") @@ -389,6 +443,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -420,8 +477,6 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( else { return Ok(Some(build_internal_gateway_proxy_public_response())); }; - let provided_auth_context = - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; if let Some(mut planned) = api::maybe_build_sync_plan_payload( state, &parts, @@ -433,9 +488,24 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ) .await? { - if provided_auth_context { - planned.auth_context = None; - } + let report_kind = planned.report_kind.clone(); + let provider_request_headers = planned + .plan + .as_ref() + .map(|plan| &plan.headers) + .ok_or_else(|| { + crate::GatewayError::Internal( + "internal gateway sync plan omitted its execution plan".to_string(), + ) + })?; + crate::usage::attach_internal_gateway_report_capability( + state, + trace_id.as_str(), + report_kind.as_deref(), + provider_request_headers, + &mut planned.report_context, + ) + .await?; return Ok(Some(Json(planned).into_response())); } return Ok(Some(build_internal_gateway_proxy_public_response())); @@ -459,6 +529,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -484,8 +557,6 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( else { return Ok(Some(build_internal_gateway_proxy_public_response())); }; - let provided_auth_context = - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; if let Some(mut planned) = api::maybe_build_stream_plan_payload( state, &parts, @@ -496,9 +567,25 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ) .await? { - if provided_auth_context { - planned.auth_context = None; - } + let report_kind = planned.report_kind.clone(); + let provider_request_headers = planned + .plan + .as_ref() + .map(|plan| &plan.headers) + .ok_or_else(|| { + crate::GatewayError::Internal( + "internal gateway stream plan omitted its execution plan" + .to_string(), + ) + })?; + crate::usage::attach_internal_gateway_report_capability( + state, + trace_id.as_str(), + report_kind.as_deref(), + provider_request_headers, + &mut planned.report_context, + ) + .await?; return Ok(Some(Json(planned).into_response())); } return Ok(Some(build_internal_gateway_proxy_public_response())); @@ -522,6 +609,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -553,7 +643,6 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( else { return Ok(None); }; - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; if let Some(plan_payload) = api::maybe_build_sync_plan_payload( state, &parts, @@ -609,6 +698,9 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if let Err(response) = reject_supplied_auth_context(payload.auth_context.as_ref()) { + return Ok(Some(response)); + } let parts = match build_internal_gateway_request_parts( &payload.method, &payload.path, @@ -634,7 +726,6 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( else { return Ok(None); }; - apply_supplied_auth_context(state, &mut resolved, payload.auth_context).await?; if let Some(plan_payload) = api::maybe_build_stream_plan_payload( state, &parts, @@ -678,9 +769,10 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( "invalid internal gateway report-sync payload", ))); }; - let payload = match serde_json::from_slice::( - request_body, - ) { + let mut payload = match serde_json::from_slice::< + crate::usage::GatewaySyncReportRequest, + >(request_body) + { Ok(payload) => payload, Err(_) => { return Ok(Some(build_internal_control_error_response( @@ -689,6 +781,20 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + payload.report_context = Some( + match resolve_bound_internal_report_context( + state, + payload.trace_id.as_str(), + payload.report_kind.as_str(), + payload.report_context.as_ref(), + "report_sync", + ) + .await + { + Ok(report_context) => report_context, + Err(response) => return Ok(Some(response)), + }, + ); crate::usage::submit_sync_report(state, payload).await?; return Ok(Some(Json(json!({ "ok": true })).into_response())); } @@ -701,9 +807,10 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( "invalid internal gateway report-stream payload", ))); }; - let payload = match serde_json::from_slice::( - request_body, - ) { + let mut payload = match serde_json::from_slice::< + crate::usage::GatewayStreamReportRequest, + >(request_body) + { Ok(payload) => payload, Err(_) => { return Ok(Some(build_internal_control_error_response( @@ -712,6 +819,20 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + payload.report_context = Some( + match resolve_bound_internal_report_context( + state, + payload.trace_id.as_str(), + payload.report_kind.as_str(), + payload.report_context.as_ref(), + "report_stream", + ) + .await + { + Ok(report_context) => report_context, + Err(response) => return Ok(Some(response)), + }, + ); crate::usage::submit_stream_report(state, payload).await?; return Ok(Some(Json(json!({ "ok": true })).into_response())); } @@ -724,9 +845,10 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( "invalid internal gateway finalize-sync payload", ))); }; - let payload = match serde_json::from_slice::( - request_body, - ) { + let mut payload = match serde_json::from_slice::< + crate::usage::GatewaySyncReportRequest, + >(request_body) + { Ok(payload) => payload, Err(_) => { return Ok(Some(build_internal_control_error_response( @@ -735,6 +857,26 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( ))); } }; + if !internal_finalize_report_kind_is_supported(payload.report_kind.as_str()) { + return Ok(Some(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "Unsupported gateway sync finalize kind", + ))); + } + payload.report_context = Some( + match resolve_bound_internal_report_context( + state, + payload.trace_id.as_str(), + payload.report_kind.as_str(), + payload.report_context.as_ref(), + "finalize_sync", + ) + .await + { + Ok(report_context) => report_context, + Err(response) => return Ok(Some(response)), + }, + ); let Some(synthetic_decision) = build_internal_finalize_decision(&payload) else { return Ok(Some(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, @@ -806,8 +948,66 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( Err(response) => return Ok(Some(response)), }; let node_id = payload.node_id.trim().to_string(); + let authenticated_generation = match state + .tunnel + .authenticate_control_plane_request( + request_headers, + request_context.request_method.as_str(), + request_context.request_path.as_str(), + &node_id, + request_body, + ) + .await + { + Ok(generation) => generation, + Err(error) => { + let (status, message) = match error { + crate::tunnel::ControlPlaneAuthError::Unavailable => ( + http::StatusCode::SERVICE_UNAVAILABLE, + "tunnel control-plane authentication unavailable", + ), + crate::tunnel::ControlPlaneAuthError::Invalid => ( + http::StatusCode::FORBIDDEN, + "invalid tunnel control-plane authentication", + ), + }; + return Ok(Some(build_internal_control_error_response(status, message))); + } + }; + let claim = match claim_tunnel_heartbeat( + state.runtime_state.as_ref(), + &node_id, + &payload.heartbeat_session_id, + payload.heartbeat_id, + ) + .await + { + Ok(claim) => claim, + Err(error) => { + return Ok(Some(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + error, + ))); + } + }; + if claim.is_none() { + let response = match state.find_proxy_node(&node_id).await { + Ok(Some(node)) if node.tunnel_generation == authenticated_generation => Json( + build_internal_tunnel_heartbeat_ack(&node, payload.heartbeat_id), + ) + .into_response(), + Ok(Some(_)) | Ok(None) => build_internal_control_error_response( + http::StatusCode::FORBIDDEN, + "proxy tunnel credential was revoked", + ), + Err(_) => internal_gateway_data_error_response("heartbeat_duplicate_lookup"), + }; + return Ok(Some(response)); + } + let claim = claim.expect("fresh heartbeat claim should be present"); let mutation = ProxyNodeHeartbeatMutation { node_id: node_id.clone(), + expected_tunnel_generation: Some(authenticated_generation.clone()), heartbeat_interval: payload.heartbeat_interval, active_connections: payload.active_connections, total_requests_delta: payload.window_total_requests.or(payload.total_requests), @@ -820,19 +1020,25 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( }; let response = match state.apply_proxy_node_heartbeat(&mutation).await { - Ok(Some(node)) => Json(build_internal_tunnel_heartbeat_ack( - &node, - payload.heartbeat_id, - )) - .into_response(), - Ok(None) => build_internal_control_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("heartbeat sync failed: ProxyNode {node_id} 不存在"), - ), - Err(err) => build_internal_control_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("heartbeat sync failed: {}", gateway_error_message(err)), - ), + Ok(Some(node)) => { + finish_tunnel_heartbeat_claim(state.runtime_state.as_ref(), claim).await; + Json(build_internal_tunnel_heartbeat_ack( + &node, + payload.heartbeat_id, + )) + .into_response() + } + Ok(None) => { + finish_tunnel_heartbeat_claim(state.runtime_state.as_ref(), claim).await; + build_internal_control_error_response( + http::StatusCode::FORBIDDEN, + "proxy tunnel credential was revoked", + ) + } + Err(_) => { + finish_tunnel_heartbeat_claim(state.runtime_state.as_ref(), claim).await; + internal_gateway_data_error_response("heartbeat_sync") + } }; return Ok(Some(response)); } @@ -849,8 +1055,36 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( Ok(payload) => payload, Err(response) => return Ok(Some(response)), }; + let node_id = payload.node_id.trim().to_string(); + let authenticated_generation = match state + .tunnel + .authenticate_control_plane_request( + request_headers, + request_context.request_method.as_str(), + request_context.request_path.as_str(), + &node_id, + request_body, + ) + .await + { + Ok(generation) => generation, + Err(error) => { + let (status, message) = match error { + crate::tunnel::ControlPlaneAuthError::Unavailable => ( + http::StatusCode::SERVICE_UNAVAILABLE, + "tunnel control-plane authentication unavailable", + ), + crate::tunnel::ControlPlaneAuthError::Invalid => ( + http::StatusCode::FORBIDDEN, + "invalid tunnel control-plane authentication", + ), + }; + return Ok(Some(build_internal_control_error_response(status, message))); + } + }; let mutation = ProxyNodeTunnelStatusMutation { - node_id: payload.node_id.trim().to_string(), + node_id, + expected_tunnel_generation: Some(authenticated_generation), connected: payload.connected, conn_count: payload.conn_count, detail: None, @@ -858,11 +1092,12 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl( }; let response = match state.update_proxy_node_tunnel_status(&mutation).await { - Ok(node) => Json(json!({ "updated": node.is_some() })).into_response(), - Err(err) => build_internal_control_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("node status sync failed: {}", gateway_error_message(err)), + Ok(Some(_)) => Json(json!({ "updated": true })).into_response(), + Ok(None) => build_internal_control_error_response( + http::StatusCode::FORBIDDEN, + "proxy tunnel credential was revoked", ), + Err(_) => internal_gateway_data_error_response("node_status_sync"), }; return Ok(Some(response)); } diff --git a/apps/aether-gateway/src/handlers/internal/gateway_helpers.rs b/apps/aether-gateway/src/handlers/internal/gateway_helpers.rs index ae87f655d..afcac6986 100644 --- a/apps/aether-gateway/src/handlers/internal/gateway_helpers.rs +++ b/apps/aether-gateway/src/handlers/internal/gateway_helpers.rs @@ -204,6 +204,11 @@ pub(crate) fn build_internal_gateway_request_parts( }; let (mut parts, _) = request.into_parts(); parts.headers = mapped_headers; + parts + .extensions + .insert(crate::headers::request_origin_from_trusted_headers( + &parts.headers, + )); Ok(parts) } @@ -246,6 +251,30 @@ pub(crate) fn build_internal_finalize_decision( ) } +pub(crate) fn internal_finalize_report_kind_is_supported(report_kind: &str) -> bool { + matches!( + report_kind.trim().to_ascii_lowercase().as_str(), + "openai_chat_sync_finalize" + | "openai_responses_sync_finalize" + | "openai_responses_compact_sync_finalize" + | "openai_compact_sync_finalize" + | "openai_cli_sync_finalize" + | "openai_embedding_sync_finalize" + | "openai_image_sync_finalize" + | "claude_chat_sync_finalize" + | "claude_cli_sync_finalize" + | "gemini_chat_sync_finalize" + | "gemini_interactions_sync_finalize" + | "gemini_cli_sync_finalize" + | "openai_video_create_sync_finalize" + | "openai_video_remix_sync_finalize" + | "openai_video_delete_sync_finalize" + | "openai_video_cancel_sync_finalize" + | "gemini_video_create_sync_finalize" + | "gemini_video_cancel_sync_finalize" + ) +} + pub(crate) async fn maybe_build_internal_finalize_video_response( state: &AppState, trace_id: &str, @@ -353,10 +382,6 @@ pub(crate) async fn maybe_build_internal_finalize_video_response( Ok(None) } -pub(crate) fn gateway_error_message(error: GatewayError) -> String { - error.into_message() -} - pub(crate) fn build_internal_tunnel_heartbeat_ack( node: &StoredProxyNode, heartbeat_id: u64, @@ -391,7 +416,12 @@ pub(crate) fn parse_internal_tunnel_heartbeat_request( })?; let node_id = payload.node_id.trim(); - if node_id.is_empty() || node_id.len() > 36 || payload.heartbeat_id == 0 { + if node_id.is_empty() + || node_id.len() > 36 + || payload.heartbeat_id == 0 + || crate::tunnel::validate_tunnel_heartbeat_session_id(&payload.heartbeat_session_id) + .is_err() + { return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, "invalid heartbeat payload", diff --git a/apps/aether-gateway/src/handlers/proxy/body_buffer.rs b/apps/aether-gateway/src/handlers/proxy/body_buffer.rs index c27ded5bb..94dc3fefc 100644 --- a/apps/aether-gateway/src/handlers/proxy/body_buffer.rs +++ b/apps/aether-gateway/src/handlers/proxy/body_buffer.rs @@ -3,6 +3,7 @@ use crate::control::GatewayPublicRequestContext; use crate::headers::RequestBodyNormalizationError; use crate::{AppState, GatewayError}; use aether_gateway_frontdoor::{BodyBufferError, BodyBufferPolicy as FrontdoorBodyBufferPolicy}; +use aether_usage_runtime::MAX_INTERNAL_REPORT_BODY_BYTES; use axum::body::{Body, Bytes}; use axum::http::{self, Response}; use std::sync::Arc; @@ -13,6 +14,11 @@ use tracing::{info, warn}; const REQUEST_BODY_READ_TIMEOUT_DETAIL: &str = "Request body read timed out before the gateway could route the request"; const REQUEST_BODY_READ_FAILED_DETAIL: &str = "Failed to read request body"; +// A report can carry two 64 MiB decoded provider/client bodies. Base64 expands +// each body by roughly one third, so the request envelope needs about 180 MiB +// plus JSON metadata. Keep a bounded 192 MiB control-plane ceiling rather +// than inheriting the generic 256 MiB public request limit. +const MAX_INTERNAL_REPORT_REQUEST_BODY_BYTES: u64 = 192 * 1024 * 1024; #[derive(Debug, Clone)] pub(super) struct RequestBodyBufferPolicy { @@ -21,9 +27,44 @@ pub(super) struct RequestBodyBufferPolicy { impl RequestBodyBufferPolicy { pub(super) fn from_state(state: &AppState) -> Self { + Self::from_state_with_max_bytes(state, crate::headers::max_request_body_bytes()) + } + + pub(super) fn for_internal_report(state: &AppState) -> Self { + // Keep the envelope limit tied to the per-field decoded limit. The + // explicit constant leaves room for base64 expansion and metadata. + let envelope_limit = MAX_INTERNAL_REPORT_REQUEST_BODY_BYTES + .min((MAX_INTERNAL_REPORT_BODY_BYTES as u64).saturating_mul(3)); + Self::from_state_with_max_bytes(state, envelope_limit) + } + + pub(super) fn for_request_context( + state: &AppState, + request_context: &GatewayPublicRequestContext, + ) -> Self { + let is_internal_report = + request_context + .control_decision + .as_ref() + .is_some_and(|decision| { + decision.route_class.as_deref() == Some("internal_proxy") + && decision.route_family.as_deref() == Some("internal_gateway") + && matches!( + decision.route_kind.as_deref(), + Some("report_sync" | "report_stream" | "finalize_sync") + ) + }); + if is_internal_report { + Self::for_internal_report(state) + } else { + Self::from_state(state) + } + } + + fn from_state_with_max_bytes(state: &AppState, max_bytes: u64) -> Self { Self { - inner: FrontdoorBodyBufferPolicy::with_permit_bytes( - crate::headers::max_request_body_bytes(), + inner: FrontdoorBodyBufferPolicy::with_optional_read_timeout_and_permit_bytes( + max_bytes, state.frontdoor_runtime_guards.request_body_read_timeout, state.frontdoor_runtime_guards.internal_gate_queue_budget, state @@ -55,6 +96,26 @@ impl RequestBodyBufferPolicy { } } + #[cfg(test)] + pub(super) fn for_tests_without_read_timeout(max_bytes: u64) -> Self { + let budget_bytes = usize::try_from(max_bytes) + .unwrap_or(usize::MAX) + .max(crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES); + Self { + inner: FrontdoorBodyBufferPolicy::with_optional_read_timeout_and_permit_bytes( + max_bytes, + None, + Duration::from_secs(1), + budget_bytes, + crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES, + Arc::new(Semaphore::new( + budget_bytes.saturating_add(crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES - 1) + / crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES, + )), + ), + } + } + #[cfg(test)] pub(super) fn for_tests_with_budget( max_bytes: u64, @@ -79,12 +140,16 @@ impl RequestBodyBufferPolicy { self.inner.max_bytes() } + fn effective_max_bytes(&self) -> u64 { + self.inner.effective_max_bytes() + } + fn budget_bytes(&self) -> usize { self.inner.budget_bytes() } - fn read_timeout(&self) -> Duration { - self.inner.read_timeout() + fn read_timeout(&self) -> Option { + self.inner.optional_read_timeout() } async fn reserve( @@ -97,6 +162,9 @@ impl RequestBodyBufferPolicy { #[derive(Debug)] pub(super) enum RequestBodyBufferError { + InvalidHeaders { + message: String, + }, Normalization(RequestBodyNormalizationError), TooLarge { limit_bytes: u64, @@ -117,6 +185,7 @@ pub(super) enum RequestBodyBufferError { impl RequestBodyBufferError { pub(super) fn http_status(&self) -> http::StatusCode { match self { + Self::InvalidHeaders { .. } => http::StatusCode::BAD_REQUEST, Self::Normalization(error) => error.http_status(), Self::TooLarge { .. } => http::StatusCode::PAYLOAD_TOO_LARGE, Self::Overloaded { .. } => http::StatusCode::SERVICE_UNAVAILABLE, @@ -125,8 +194,9 @@ impl RequestBodyBufferError { } } - fn client_message(&self) -> String { + pub(super) fn client_message(&self) -> String { match self { + Self::InvalidHeaders { .. } => "Invalid request body headers".to_string(), Self::Normalization(error) => error.client_message(), Self::TooLarge { limit_bytes } => format!("Request body exceeds {limit_bytes} bytes"), Self::Overloaded { .. } => { @@ -139,7 +209,12 @@ impl RequestBodyBufferError { fn reason(&self) -> &'static str { match self { + Self::InvalidHeaders { .. } => "invalid_request_body_headers", Self::Normalization(error) => match error { + RequestBodyNormalizationError::InvalidBodyFraming + | RequestBodyNormalizationError::AmbiguousBodyFraming => { + "invalid_request_body_headers" + } RequestBodyNormalizationError::UnsupportedContentEncoding(_) => { "unsupported_content_encoding" } @@ -162,6 +237,7 @@ impl RequestBodyBufferError { impl From for RequestBodyBufferError { fn from(error: BodyBufferError) -> Self { match error { + BodyBufferError::InvalidHeaders { message } => Self::InvalidHeaders { message }, BodyBufferError::TooLarge { limit_bytes } => Self::TooLarge { limit_bytes }, BodyBufferError::Overloaded { requested_bytes, @@ -188,23 +264,26 @@ pub(super) async fn buffer_and_normalize_request_body( phase: &'static str, policy: RequestBodyBufferPolicy, ) -> Result { + let sanitized_path_and_query = crate::middleware::sanitize_access_log_path(path_and_query); let reservation = policy .reserve(headers) .await .map_err(RequestBodyBufferError::from)?; let reservation_bytes = reservation.requested_bytes(); + let read_timeout = policy.read_timeout(); info!( event_name = "frontdoor_request_body_buffer_started", log_type = "event", trace_id, method = %method, - path = %path_and_query, + path = %sanitized_path_and_query, phase, max_body_bytes = policy.max_bytes(), reserved_body_bytes = reservation_bytes, body_buffer_budget_bytes = policy.budget_bytes(), - timeout_ms = policy.read_timeout().as_millis() as u64, + timeout_enabled = read_timeout.is_some(), + timeout_ms = read_timeout.map(|timeout| timeout.as_millis() as u64).unwrap_or(0), "gateway started buffering request body" ); @@ -218,7 +297,7 @@ pub(super) async fn buffer_and_normalize_request_body( crate::headers::normalize_request_body_headers_and_bytes_with_limit( headers, body, - policy.max_bytes(), + policy.effective_max_bytes(), ) }) .map_err(RequestBodyBufferError::Normalization)?; @@ -227,7 +306,7 @@ pub(super) async fn buffer_and_normalize_request_body( log_type = "event", trace_id, method = %method, - path = %path_and_query, + path = %sanitized_path_and_query, phase, body_bytes = normalized.len(), elapsed_ms, @@ -241,12 +320,14 @@ pub(super) fn build_request_body_buffer_error_response( request_context: &GatewayPublicRequestContext, error: &RequestBodyBufferError, ) -> Result, GatewayError> { + let sanitized_path_and_query = + crate::middleware::sanitize_access_log_path(&request_context.request_path_and_query()); warn!( event_name = "frontdoor_request_body_buffer_failed", log_type = "ops", trace_id, method = %request_context.request_method, - path = %request_context.request_path_and_query(), + path = %sanitized_path_and_query, status_code = error.http_status().as_u16(), reason = error.reason(), detail = %error.client_message(), diff --git a/apps/aether-gateway/src/handlers/proxy/finalize.rs b/apps/aether-gateway/src/handlers/proxy/finalize.rs index 9044d9db5..01225bed3 100644 --- a/apps/aether-gateway/src/handlers/proxy/finalize.rs +++ b/apps/aether-gateway/src/handlers/proxy/finalize.rs @@ -50,6 +50,7 @@ pub(super) fn finalize_gateway_response( mut response: Response, trace_id: &str, remote_addr: &std::net::SocketAddr, + client_ip: std::net::IpAddr, method: &http::Method, path_and_query: &str, control_decision: Option<&GatewayControlDecision>, @@ -57,6 +58,7 @@ pub(super) fn finalize_gateway_response( started_at: &Instant, request_permit: Option, ) -> Response { + apply_sensitive_route_cache_policy(response.headers_mut(), path_and_query, control_decision); attach_control_decision_headers(&mut response, control_decision); if !response.headers().contains_key(TRACE_ID_HEADER) { response.headers_mut().insert( @@ -117,6 +119,7 @@ pub(super) fn finalize_gateway_response( trace_id = %trace_id, request_id, remote_addr = %remote_addr, + client_ip = %client_ip, method = %method, path = %sanitized_path_and_query, user_id, @@ -137,6 +140,7 @@ pub(super) fn finalize_gateway_response( trace_id = %trace_id, request_id, remote_addr = %remote_addr, + client_ip = %client_ip, method = %method, path = %sanitized_path_and_query, user_id, @@ -157,6 +161,7 @@ pub(super) fn finalize_gateway_response( trace_id = %trace_id, request_id, remote_addr = %remote_addr, + client_ip = %client_ip, method = %method, path = %sanitized_path_and_query, user_id, @@ -174,6 +179,45 @@ pub(super) fn finalize_gateway_response( maybe_hold_axum_response_permit(response, request_permit) } +fn apply_sensitive_route_cache_policy( + headers: &mut http::HeaderMap, + path_and_query: &str, + control_decision: Option<&GatewayControlDecision>, +) { + let is_admin_route = path_and_query + .split_once('?') + .map_or(path_and_query, |(path, _)| path) + .starts_with("/api/admin/") + || control_decision + .is_some_and(|decision| decision.route_class.as_deref() == Some("admin_proxy")); + let has_authenticated_principal = control_decision.is_some_and(|decision| { + decision.auth_context.is_some() || decision.admin_principal.is_some() + }); + let is_sensitive_public_support = control_decision.is_some_and(|decision| { + if decision.route_class.as_deref() != Some("public_support") { + return false; + } + match decision.route_family.as_deref() { + Some( + "auth" | "dashboard" | "monitoring_user" | "announcement_user" | "wallet" + | "ccswitch" | "users_me" | "models" | "oauth" | "install", + ) => true, + Some("billing") => decision.route_kind.as_deref() != Some("plans"), + Some("system_catalog") => decision.route_kind.as_deref() == Some("test_connection"), + _ => false, + } + }); + if !(is_admin_route || has_authenticated_principal || is_sensitive_public_support) { + return; + } + + headers.insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("no-store"), + ); + headers.insert(http::header::PRAGMA, HeaderValue::from_static("no-cache")); +} + fn attach_control_decision_headers( response: &mut Response, control_decision: Option<&GatewayControlDecision>, @@ -247,11 +291,17 @@ pub(super) fn finalize_gateway_response_with_context( started_at: &Instant, request_permit: Option, ) -> Response { + let client_ip = request_context + .client_ip + .as_deref() + .and_then(|value| value.parse().ok()) + .unwrap_or_else(|| remote_addr.ip()); finalize_gateway_response( state, response, &request_context.trace_id, remote_addr, + client_ip, &request_context.request_method, &request_context.request_path_and_query(), request_context.control_decision.as_ref(), @@ -263,7 +313,9 @@ pub(super) fn finalize_gateway_response_with_context( #[cfg(test)] mod tests { - use super::{finalize_gateway_response, request_wants_stream}; + use super::{ + apply_sensitive_route_cache_policy, finalize_gateway_response, request_wants_stream, + }; use crate::control::{GatewayControlDecision, GatewayPublicRequestContext}; use crate::AppState; use axum::body::{Body, Bytes}; @@ -278,6 +330,102 @@ mod tests { struct SharedBufferWriter(Arc>>); + #[test] + fn admin_responses_are_never_cacheable() { + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("public, max-age=3600"), + ); + + apply_sensitive_route_cache_policy( + &mut headers, + "/api/admin/endpoints/keys/key-1/reveal?include_key=true", + None, + ); + + assert_eq!( + headers.get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("no-store")) + ); + assert_eq!( + headers.get(http::header::PRAGMA), + Some(&HeaderValue::from_static("no-cache")) + ); + } + + #[test] + fn authenticated_user_data_responses_are_never_cacheable() { + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("public, max-age=3600"), + ); + let decision = GatewayControlDecision::synthetic( + "/api/wallet/transactions", + Some("public_support".to_string()), + Some("wallet".to_string()), + Some("transactions".to_string()), + None, + ); + + apply_sensitive_route_cache_policy( + &mut headers, + "/api/wallet/transactions", + Some(&decision), + ); + + assert_eq!( + headers.get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("no-store")) + ); + assert_eq!( + headers.get(http::header::PRAGMA), + Some(&HeaderValue::from_static("no-cache")) + ); + } + + #[test] + fn public_catalog_responses_keep_their_existing_cache_policy() { + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("public, max-age=3600"), + ); + let decision = GatewayControlDecision::synthetic( + "/api/public/models", + Some("public_support".to_string()), + Some("public_catalog".to_string()), + Some("models".to_string()), + None, + ); + + apply_sensitive_route_cache_policy(&mut headers, "/api/public/models", Some(&decision)); + + assert_eq!( + headers.get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("public, max-age=3600")) + ); + assert!(headers.get(http::header::PRAGMA).is_none()); + } + + #[test] + fn non_admin_responses_keep_their_existing_cache_policy() { + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CACHE_CONTROL, + HeaderValue::from_static("public, max-age=3600"), + ); + + apply_sensitive_route_cache_policy(&mut headers, "/api/models", None); + + assert_eq!( + headers.get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("public, max-age=3600")) + ); + assert!(headers.get(http::header::PRAGMA).is_none()); + } + impl SharedBuffer { fn lines(&self) -> Vec { String::from_utf8(self.0.lock().expect("buffer should lock").clone()) @@ -320,6 +468,7 @@ mod tests { request_query_string: None, request_content_type: Some("application/json".to_string()), host_header: None, + client_ip: None, control_decision: None, }; let mut headers = HeaderMap::new(); @@ -373,6 +522,7 @@ mod tests { response, "trace-finalize", &remote_addr, + remote_addr.ip(), &Method::GET, "/v1beta/models/gemini-3-flash-preview:generateContent?key=secret&alt=sse", Some(&control_decision), diff --git a/apps/aether-gateway/src/handlers/proxy/local.rs b/apps/aether-gateway/src/handlers/proxy/local.rs index 9794438b3..0d22a0a7e 100644 --- a/apps/aether-gateway/src/handlers/proxy/local.rs +++ b/apps/aether-gateway/src/handlers/proxy/local.rs @@ -17,12 +17,14 @@ pub(super) async fn maybe_build_local_internal_proxy_response( state: &AppState, request_context: &GatewayPublicRequestContext, remote_addr: &std::net::SocketAddr, + request_headers: &http::HeaderMap, request_body: Option<&Bytes>, ) -> Result>, GatewayError> { internal::maybe_build_local_internal_proxy_response_impl( state, request_context, remote_addr, + request_headers, request_body, ) .await @@ -31,6 +33,7 @@ pub(super) async fn maybe_build_local_internal_proxy_response( pub(super) async fn maybe_build_local_admin_proxy_response( state: &AppState, request_context: &GatewayPublicRequestContext, + remote_addr: &std::net::SocketAddr, request_headers: &http::HeaderMap, request_body: Option<&Bytes>, ) -> Result>, GatewayError> { @@ -51,6 +54,7 @@ pub(super) async fn maybe_build_local_admin_proxy_response( admin_api::maybe_build_local_admin_response(admin_api::AdminRouteRequest::new( state, request_context, + remote_addr, request_headers, request_body, )) diff --git a/apps/aether-gateway/src/handlers/proxy/mod.rs b/apps/aether-gateway/src/handlers/proxy/mod.rs index 00b17feac..5d8886ce4 100644 --- a/apps/aether-gateway/src/handlers/proxy/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/mod.rs @@ -22,7 +22,8 @@ use crate::ai_serving::api::{ use crate::api::response::{ build_client_response, build_client_response_from_parts, build_local_auth_rejection_response, build_local_http_error_response, build_local_http_error_response_with_request_path, - build_local_overloaded_response, build_local_user_rpm_limited_response, + build_local_overloaded_response, build_local_plan_usage_limited_response, + build_local_user_rpm_limited_response, }; use crate::constants::{ CONTROL_CANDIDATE_ID_HEADER, DEPENDENCY_REASON_HEADER, EXECUTION_PATH_CONTROL_EXECUTE_STREAM, @@ -38,12 +39,12 @@ use crate::constants::{ FORWARDED_PROTO_HEADER, GATEWAY_HEADER, LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER, TRUSTED_AUTH_BALANCE_HEADER, TRUSTED_AUTH_USER_ID_HEADER, TUNNEL_AFFINITY_FORWARDED_BY_HEADER, - TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + TUNNEL_AFFINITY_NODE_ID_HEADER, TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, }; use crate::control::{ - allows_control_execute_emergency, management_token_permission_keys_from_value, - maybe_execute_via_control, request_model_local_rejection, should_buffer_request_for_local_auth, - trusted_auth_local_rejection, GatewayControlDecision, GatewayPublicRequestContext, + allows_control_execute_emergency, maybe_execute_via_control, request_model_local_rejection, + should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayControlDecision, + GatewayPublicRequestContext, }; use crate::executor::{ beautify_local_execution_client_error_message, build_local_execution_runtime_miss_context, @@ -56,8 +57,9 @@ use crate::frontdoor_loop_guard::{ }; use crate::handlers::shared::{ build_admin_proxy_auth_required_response, build_unhandled_admin_proxy_response, ip_rules_allow, - json_ip_rules_allow, local_proxy_route_requires_buffered_body, request_enables_control_execute, - should_strip_forwarded_provider_credential_header, should_strip_forwarded_trusted_admin_header, + local_proxy_route_requires_buffered_body, request_enables_control_execute, + sanitize_upstream_path_and_query, should_strip_forwarded_provider_credential_header, + should_strip_forwarded_trusted_admin_header, }; use crate::headers::{ effective_client_ip, extract_or_generate_trace_id, request_origin_from_headers_and_remote_addr, @@ -68,7 +70,6 @@ use crate::scheduler::candidate::{ is_auth_api_key_concurrency_limit_skip_reason, AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, }; -use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode}; use crate::stage_metrics::observe_gateway_stage_ms; use crate::{ AppState, FrontdoorUserRpmOutcome, GatewayError, GatewayFallbackMetricKind, @@ -78,7 +79,6 @@ use axum::body::{to_bytes, Body, Bytes}; use axum::extract::{ConnectInfo, Request, State}; use axum::http::{self, header::HeaderName, header::HeaderValue, Response}; use futures_util::StreamExt; -use sha2::{Digest, Sha256}; use std::{collections::BTreeMap, time::Instant}; use tracing::{debug, info, warn}; @@ -105,6 +105,12 @@ const LOCAL_EXECUTION_LOOP_DETECTED_DETAIL: &str = "Gateway detected an execution runtime request loop back into the local frontdoor"; const AUTH_API_KEY_CONCURRENCY_LIMIT_REACHED_DETAIL: &str = "当前调用方 API Key 并发请求数已达上限,请稍后重试"; +const PROVIDER_KEY_CAPACITY_LIMIT_REACHED_DETAIL: &str = + "所有可用上游账号当前均已达到并发或 RPM 上限,请稍后重试"; +const PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS: &[&str] = &[ + "provider_key_concurrency_limit_reached", + "key_rpm_exhausted", +]; const LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL: &str = "当前 AI 请求在本地执行规划阶段超时,请稍后重试"; const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward"; @@ -157,12 +163,14 @@ fn finalize_local_execution_planning_timeout( phase: &'static str, timeout_ms: u64, ) -> Result, GatewayError> { + let sanitized_path_and_query = + crate::middleware::sanitize_access_log_path(&request_context.request_path_and_query()); warn!( event_name = "frontdoor_local_execution_planning_timeout", log_type = "ops", trace_id, method = %request_context.request_method, - path = %request_context.request_path_and_query(), + path = %sanitized_path_and_query, route_family = control_decision .and_then(|decision| decision.route_family.as_deref()) .unwrap_or("-"), @@ -211,29 +219,6 @@ fn execution_runtime_candidate_header_value(decision: &GatewayControlDecision) - } } -fn extract_management_token_bearer(headers: &http::HeaderMap) -> Option { - let header = crate::headers::header_value_str(headers, http::header::AUTHORIZATION.as_str())?; - let token = header - .strip_prefix("Bearer ") - .or_else(|| header.strip_prefix("bearer "))? - .trim() - .to_string(); - (!token.is_empty() - && (token.starts_with(MANAGEMENT_TOKEN_PREFIX) - || token.starts_with(LEGACY_MANAGEMENT_TOKEN_PREFIX))) - .then_some(token) -} - -fn hash_management_token(value: &str) -> String { - let mut hasher = Sha256::new(); - hasher.update(value.as_bytes()); - format!("{:x}", hasher.finalize()) -} - -fn remote_ip_allowed(allowed_ips: Option<&serde_json::Value>, remote_ip: std::net::IpAddr) -> bool { - json_ip_rules_allow(allowed_ips, remote_ip) -} - fn api_key_remote_ip_allowed(ip_rules: Option<&[String]>, remote_ip: std::net::IpAddr) -> bool { ip_rules_allow(ip_rules, remote_ip) } @@ -253,67 +238,39 @@ async fn maybe_promote_management_token_admin_principal( return Ok(()); } - let Some(token) = extract_management_token_bearer(headers) else { - return Ok(()); - }; - let token_hash = hash_management_token(&token); - let Some(token_with_user) = state - .get_management_token_with_user_by_hash(&token_hash) - .await? - else { - return Ok(()); - }; - - if !token_with_user.token.is_active { - return Ok(()); - } - if token_with_user - .token - .expires_at_unix_secs - .is_some_and(|value| value <= chrono::Utc::now().timestamp().max(0) as u64) + let authenticated = match crate::management_token_auth::authenticate_management_token( + state, headers, client_ip, + ) + .await { - return Ok(()); - } - if !remote_ip_allowed(token_with_user.token.allowed_ips.as_ref(), client_ip) { - return Ok(()); - } - let Some(user) = state.find_user_auth_by_id(&token_with_user.user.id).await? else { - return Ok(()); - }; - if !user.is_active || user.is_deleted || !crate::roles::can_access_admin_console(&user.role) { - return Ok(()); - } - let management_token_permissions = match management_token_permission_keys_from_value( - token_with_user.token.permissions.as_ref(), - ) { - Ok(value) => value, - Err(err) => { - warn!( - trace_id = %trace_id, - token_id = %token_with_user.token.id, - error = %err, - "gateway rejected management token with invalid permissions" - ); - return Ok(()); + Ok(authenticated) => authenticated, + Err(crate::management_token_auth::ManagementTokenAuthError::Unavailable) => { + return Err(GatewayError::Internal( + "management token authentication unavailable".to_string(), + )) } + Err( + crate::management_token_auth::ManagementTokenAuthError::Missing + | crate::management_token_auth::ManagementTokenAuthError::Invalid, + ) => return Ok(()), }; decision.admin_principal = Some(crate::control::GatewayAdminPrincipalContext { - user_id: user.id.clone(), - user_role: user.role.clone(), + user_id: authenticated.user.id.clone(), + user_role: authenticated.user.role.clone(), session_id: None, - management_token_id: Some(token_with_user.token.id.clone()), - management_token_permissions, + management_token_id: Some(authenticated.token.id.clone()), + management_token_permissions: Some(authenticated.permissions), }); let remote_ip = client_ip.to_string(); if let Err(err) = state - .record_management_token_usage(&token_with_user.token.id, Some(remote_ip.as_str())) + .record_management_token_usage(&authenticated.token.id, Some(remote_ip.as_str())) .await { warn!( trace_id = %trace_id, - token_id = %token_with_user.token.id, + token_id = %authenticated.token.id, error = ?err, "gateway failed to record management token usage" ); @@ -405,27 +362,7 @@ async fn maybe_forward_public_request_to_tunnel_owner( policy_context, ) } else { - let cache_affinity_enabled = match read_scheduler_ordering_config(state).await { - Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, - Err(err) => { - warn!( - trace_id = %request_context.trace_id, - error = ?err, - "gateway failed to load scheduler config while checking tunnel affinity forwarding mode" - ); - SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity - } - }; - if !cache_affinity_enabled { - return Ok(None); - } - crate::scheduler::affinity::read_cached_scheduler_affinity_target( - state, - &auth_context.api_key_id, - affinity_context.client_session_affinity.as_ref(), - api_format, - &affinity_context.requested_model, - ) + return Ok(None); }; let Some(target) = target else { return Ok(None); @@ -486,11 +423,22 @@ async fn maybe_forward_public_request_to_tunnel_owner( return Ok(None); } - let owner_url = format!( - "{}{}", - owner.relay_base_url.trim_end_matches('/'), - request_context.request_path_and_query() + let sanitized_path_and_query = sanitize_upstream_path_and_query( + Some(decision), + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), ); + let affinity_uri = sanitized_path_and_query + .parse::() + .map_err(|error| { + GatewayError::Internal(format!("invalid sanitized tunnel affinity URI: {error}")) + })?; + let owner_url = + crate::tunnel::build_tunnel_affinity_forward_url(&owner.relay_base_url, &affinity_uri) + .map_err(GatewayError::Internal)?; let is_stream = owner_forward_request_is_stream(parts, decision, buffered_body.unwrap_or(&empty_body)); let transport_timeouts = @@ -506,55 +454,130 @@ async fn maybe_forward_public_request_to_tunnel_owner( is_stream, transport_timeouts.as_ref(), ); - let mut upstream_request = state - .owner_forward_client - .request(parts.method.clone(), owner_url); - if let Some(timeout) = non_stream_timeout { - upstream_request = upstream_request.timeout(timeout); - } + let mut forwarded_headers = http::HeaderMap::new(); + let connection_declared = aether_http::connection_declared_header_names( + parts + .headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); for (name, value) in &parts.headers { - if should_skip_request_header(name.as_str()) || name == http::header::HOST { + if should_skip_request_header(name.as_str()) + || name == http::header::HOST + || connection_declared.contains(&name.as_str().to_ascii_lowercase()) + { continue; } if should_strip_forwarded_provider_credential_header(Some(decision), name) { continue; } + if matches!(name.as_str(), "cookie" | "cookie2") { + continue; + } + if name.as_str() == "x-real-ip" { + continue; + } if should_strip_forwarded_trusted_admin_header(Some(decision), name) { continue; } - upstream_request = upstream_request.header(name, value); + if should_strip_tunnel_affinity_forward_header(name) { + continue; + } + forwarded_headers.append(name.clone(), value.clone()); } if let Some(host) = request_context.host_header.as_deref() { - if !parts.headers.contains_key(FORWARDED_HOST_HEADER) { - upstream_request = upstream_request.header(FORWARDED_HOST_HEADER, host); - } + insert_tunnel_affinity_forward_header(&mut forwarded_headers, FORWARDED_HOST_HEADER, host)?; } - if !parts.headers.contains_key(FORWARDED_FOR_HEADER) { - upstream_request = - upstream_request.header(FORWARDED_FOR_HEADER, remote_addr.ip().to_string()); - } - if !parts.headers.contains_key(FORWARDED_PROTO_HEADER) { - upstream_request = upstream_request.header(FORWARDED_PROTO_HEADER, "http"); - } - if !parts.headers.contains_key(TRACE_ID_HEADER) { - upstream_request = upstream_request.header(TRACE_ID_HEADER, &request_context.trace_id); - } - upstream_request = upstream_request - .header(GATEWAY_HEADER, "rust-phase3b-affinity") - .header( - TUNNEL_AFFINITY_FORWARDED_BY_HEADER, - state.tunnel.local_instance_id(), - ) - .header( - TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, - owner.gateway_instance_id.as_str(), - ) - .header(TRUSTED_AUTH_USER_ID_HEADER, &auth_context.user_id) - .header(TRUSTED_AUTH_API_KEY_ID_HEADER, &auth_context.api_key_id) - .header(TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, "true"); + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + FORWARDED_FOR_HEADER, + &effective_client_ip(&parts.headers, remote_addr).to_string(), + )?; + insert_tunnel_affinity_forward_header(&mut forwarded_headers, FORWARDED_PROTO_HEADER, "http")?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TRACE_ID_HEADER, + &request_context.trace_id, + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + GATEWAY_HEADER, + "rust-phase3b-affinity", + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER, + state.tunnel.local_instance_id(), + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + state.tunnel.local_instance_id(), + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + owner.gateway_instance_id.as_str(), + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TUNNEL_AFFINITY_NODE_ID_HEADER, + node_id, + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TRUSTED_AUTH_USER_ID_HEADER, + &auth_context.user_id, + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TRUSTED_AUTH_API_KEY_ID_HEADER, + &auth_context.api_key_id, + )?; + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + "true", + )?; if let Some(balance_remaining) = auth_context.balance_remaining { - upstream_request = - upstream_request.header(TRUSTED_AUTH_BALANCE_HEADER, balance_remaining.to_string()); + insert_tunnel_affinity_forward_header( + &mut forwarded_headers, + TRUSTED_AUTH_BALANCE_HEADER, + &balance_remaining.to_string(), + )?; + } + + let auth_metadata = crate::tunnel::build_tunnel_affinity_auth_metadata( + &parts.method, + &affinity_uri, + &forwarded_headers, + ) + .map_err(GatewayError::Internal)?; + let relay_auth = state + .tunnel + .build_relay_auth_headers( + owner.gateway_instance_id.as_str(), + node_id, + true, + false, + &auth_metadata, + buffered_body.unwrap_or(&empty_body), + ) + .map_err(GatewayError::Internal)?; + relay_auth + .apply_to_headers(&mut forwarded_headers) + .map_err(GatewayError::Internal)?; + + let owner_client = + crate::tunnel::owner_forward_client_for_url(&state.owner_forward_client, &owner_url) + .await + .map_err(GatewayError::Internal)?; + let mut upstream_request = owner_client + .request(parts.method.clone(), owner_url) + .headers(forwarded_headers); + if let Some(timeout) = non_stream_timeout { + upstream_request = upstream_request.timeout(timeout); } let upstream_response = crate::tunnel::send_owner_forward_request( @@ -583,6 +606,37 @@ async fn maybe_forward_public_request_to_tunnel_owner( Ok(Some(response)) } +fn should_strip_tunnel_affinity_forward_header(name: &http::HeaderName) -> bool { + crate::tunnel::is_tunnel_relay_auth_header(name.as_str()) + || matches!( + name.as_str(), + TUNNEL_AFFINITY_FORWARDED_BY_HEADER + | TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER + | TUNNEL_AFFINITY_NODE_ID_HEADER + | TRUSTED_AUTH_USER_ID_HEADER + | TRUSTED_AUTH_API_KEY_ID_HEADER + | TRUSTED_AUTH_BALANCE_HEADER + | TRUSTED_AUTH_ACCESS_ALLOWED_HEADER + | FORWARDED_HOST_HEADER + | FORWARDED_FOR_HEADER + | FORWARDED_PROTO_HEADER + ) +} + +fn insert_tunnel_affinity_forward_header( + headers: &mut http::HeaderMap, + name: &'static str, + value: &str, +) -> Result<(), GatewayError> { + let value = HeaderValue::from_str(value).map_err(|error| { + GatewayError::Internal(format!( + "invalid tunnel affinity forward header {name}: {error}" + )) + })?; + headers.insert(HeaderName::from_static(name), value); + Ok(()) +} + fn routing_overlay_allows_affinity_target( routing_overlay: Option<&aether_routing_core::RankingOverlay>, target: &aether_scheduler_core::SchedulerAffinityTarget, @@ -626,25 +680,49 @@ fn upstream_response_is_sse(headers: &reqwest::header::HeaderMap) -> bool { fn collect_upstream_response_headers( headers: &reqwest::header::HeaderMap, ) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); headers .iter() - .map(|(name, value)| { - ( - name.as_str().to_string(), - value.to_str().unwrap_or_default().to_string(), - ) + .filter_map(|(name, value)| { + let normalized = name.as_str().to_ascii_lowercase(); + if crate::headers::should_skip_response_header(&normalized) + || connection_declared.contains(&normalized) + { + return None; + } + value + .to_str() + .ok() + .map(|value| (normalized, value.to_string())) }) .collect() } fn collect_response_headers(headers: &http::HeaderMap) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); headers .iter() - .map(|(name, value)| { - ( - name.as_str().to_string(), - value.to_str().unwrap_or_default().to_string(), - ) + .filter_map(|(name, value)| { + let normalized = name.as_str().to_ascii_lowercase(); + if crate::headers::should_skip_response_header(&normalized) + || connection_declared.contains(&normalized) + { + return None; + } + value + .to_str() + .ok() + .map(|value| (normalized, value.to_string())) }) .collect() } @@ -712,6 +790,7 @@ fn restore_redacted_stream_execution_response( }; let headers = collect_response_headers(&parts.headers); let _ = crate::privacy::StreamingResponseRestorer::new(&headers, &session)?; + replace_response_headers(&mut parts.headers, &headers)?; parts.headers.remove(http::header::CONTENT_LENGTH); let stream_headers = headers; let stream = async_stream::stream! { @@ -886,10 +965,24 @@ async fn build_sync_aware_affinity_forward_response( let status_code = upstream_response.status().as_u16(); let headers = collect_upstream_response_headers(upstream_response.headers()); - let body_bytes = upstream_response - .bytes() - .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + // A response can only be buffered here when the owner-forward path has + // to translate between sync JSON and SSE. Keep the ordinary passthrough + // and streaming paths untouched, but never let this protocol bridge trust + // an upstream `Content-Length` or allocate without a bound. + let body_bytes = aether_http::read_response_bytes_with_limit( + upstream_response, + crate::headers::max_internal_buffered_body_bytes(), + ) + .await + .map_err(|error| { + let kind = match error { + aether_http::ResponseBodyReadError::TooLarge { .. } => "too_large", + aether_http::ResponseBodyReadError::Read(_) => "read", + }; + GatewayError::Internal(format!( + "tunnel affinity response body buffering failed ({kind})" + )) + })?; if stream_request { if (200..300).contains(&status_code) { if let Some(client_api_format) = resolve_affinity_forward_client_api_format( @@ -956,7 +1049,7 @@ async fn proxy_request_inner( request: Request, ) -> Result, GatewayError> { let started_at = Instant::now(); - let client_ip = effective_client_ip(request.headers(), &remote_addr); + let direct_client_ip = effective_client_ip(request.headers(), &remote_addr); let trace_id = extract_or_generate_trace_id(request.headers()); let accepted_at = request .extensions() @@ -993,6 +1086,7 @@ async fn proxy_request_inner( response, &trace_id, &remote_addr, + direct_client_ip, request.method(), request .uri() @@ -1029,6 +1123,7 @@ async fn proxy_request_inner( response, &trace_id, &remote_addr, + direct_client_ip, request.method(), request .uri() @@ -1047,28 +1142,26 @@ async fn proxy_request_inner( }; let request_admission_ms = started_at.elapsed().as_millis() as u64; observe_gateway_stage_ms("frontdoor_admission", request_admission_ms); - match state.admin_security_ip_blacklisted(client_ip).await { - Ok(true) => { - warn!( - event_name = "frontdoor_ip_blacklist_rejected", - log_type = "event", - trace_id = %trace_id, - client_ip = %client_ip, - path = %request.uri().path(), - "gateway rejected blacklisted client IP" - ); + let pending_affinity_auth = match state + .tunnel + .prepare_tunnel_affinity_auth_request(request.method(), request.uri(), request.headers()) + .await + { + Ok(context) => context, + Err(crate::tunnel::RelayAuthError::Invalid) => { let response = build_local_http_error_response_with_request_path( &trace_id, None, Some(request.uri().path()), http::StatusCode::FORBIDDEN, - "当前 IP 已被禁止访问", + "invalid tunnel affinity authentication", )?; return Ok(finalize_gateway_response( &state, response, &trace_id, &remote_addr, + direct_client_ip, request.method(), request .uri() @@ -1081,26 +1174,228 @@ async fn proxy_request_inner( request_permit.take(), )); } - Ok(false) => {} - Err(err) => warn!( - event_name = "frontdoor_ip_blacklist_check_failed", - log_type = "ops", - trace_id = %trace_id, - client_ip = %client_ip, - error = ?err, - "gateway failed open after IP blacklist check error" - ), - } + Err(crate::tunnel::RelayAuthError::Unavailable) => { + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(request.uri().path()), + http::StatusCode::SERVICE_UNAVAILABLE, + "tunnel affinity authentication is unavailable", + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + direct_client_ip, + request.method(), + request + .uri() + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + }; let (mut parts, body) = request.into_parts(); + let mut request_body = Some(body); + let mut authenticated_affinity_body = None; + let affinity_auth = if let Some(pending) = pending_affinity_auth { + let body = match buffer_and_normalize_request_body( + &mut request_body, + &mut parts.headers, + "tunnel affinity authentication should own request body", + &trace_id, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + "tunnel_affinity_auth", + RequestBodyBufferPolicy::from_state(&state), + ) + .await + { + Ok(body) => body, + Err(error) => { + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(parts.uri.path()), + error.http_status(), + &error.client_message(), + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + direct_client_ip, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + }; + match state + .tunnel + .commit_tunnel_affinity_auth_request(pending, &body) + .await + { + Ok(context) => { + authenticated_affinity_body = Some(body); + Some(context) + } + Err(crate::tunnel::RelayAuthError::Invalid) => { + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(parts.uri.path()), + http::StatusCode::FORBIDDEN, + "invalid tunnel affinity authentication", + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + direct_client_ip, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + Err(crate::tunnel::RelayAuthError::Unavailable) => { + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(parts.uri.path()), + http::StatusCode::SERVICE_UNAVAILABLE, + "tunnel affinity authentication is unavailable", + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + direct_client_ip, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + } + } else { + None + }; + let trusted_affinity_auth = affinity_auth.is_some(); + let client_ip = affinity_auth + .map(|context| context.client_ip) + .unwrap_or(direct_client_ip); + match state.admin_security_ip_blacklisted(client_ip).await { + Ok(true) => { + warn!( + event_name = "frontdoor_ip_blacklist_rejected", + log_type = "event", + trace_id = %trace_id, + client_ip = %client_ip, + path = %parts.uri.path(), + "gateway rejected blacklisted client IP" + ); + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(parts.uri.path()), + http::StatusCode::FORBIDDEN, + "当前 IP 已被禁止访问", + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + client_ip, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + Ok(false) => {} + Err(err) => { + warn!( + event_name = "frontdoor_ip_blacklist_check_failed", + log_type = "ops", + trace_id = %trace_id, + client_ip = %client_ip, + error = ?err, + "gateway rejected request because IP blacklist state is unavailable" + ); + let response = build_local_http_error_response_with_request_path( + &trace_id, + None, + Some(parts.uri.path()), + http::StatusCode::SERVICE_UNAVAILABLE, + "IP 访问控制暂时不可用", + )?; + return Ok(finalize_gateway_response( + &state, + response, + &trace_id, + &remote_addr, + client_ip, + &parts.method, + parts + .uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/"), + None, + EXECUTION_PATH_LOCAL_AUTH_DENIED, + &started_at, + request_permit.take(), + )); + } + } crate::ai_serving::codex_context::install_codex_fingerprint_context_slot(&mut parts); let redaction_slot = crate::privacy::RedactionSessionSlot::default(); parts.extensions.insert(redaction_slot.clone()); - parts - .extensions - .insert(request_origin_from_headers_and_remote_addr( - &parts.headers, - &remote_addr, - )); + let mut request_origin = + request_origin_from_headers_and_remote_addr(&parts.headers, &remote_addr); + request_origin.client_ip = Some(client_ip.to_string()); + parts.extensions.insert(request_origin); let trace_id = extract_or_generate_trace_id(&parts.headers); state.clear_local_execution_runtime_miss_diagnostic(&trace_id); if request_hits_execution_loop_guard(&parts) { @@ -1109,11 +1404,7 @@ async fn proxy_request_inner( log_type = "ops", trace_id = %trace_id, method = %parts.method, - path = %parts - .uri - .path_and_query() - .map(|value| value.as_str()) - .unwrap_or("/"), + path = %parts.uri.path(), loop_guard_header = EXECUTION_RUNTIME_LOOP_GUARD_HEADER, "gateway rejected execution runtime request loop into frontdoor" ); @@ -1129,6 +1420,7 @@ async fn proxy_request_inner( response, &trace_id, &remote_addr, + client_ip, &parts.method, parts .uri @@ -1142,14 +1434,26 @@ async fn proxy_request_inner( )); } let request_context_started_at = Instant::now(); - let mut request_context = crate::control::resolve_public_request_context( - &state, - &parts.method, - &parts.uri, - &parts.headers, - &trace_id, - ) - .await?; + let mut request_context = if trusted_affinity_auth { + crate::control::resolve_public_request_context_with_trusted_auth( + &state, + &parts.method, + &parts.uri, + &parts.headers, + &trace_id, + ) + .await? + } else { + crate::control::resolve_public_request_context_without_trusted_auth( + &state, + &parts.method, + &parts.uri, + &parts.headers, + &trace_id, + ) + .await? + }; + request_context.client_ip = Some(client_ip.to_string()); maybe_promote_management_token_admin_principal( &state, client_ip, @@ -1198,47 +1502,47 @@ async fn proxy_request_inner( log_type = "event", trace_id = %trace_id, method = %parts.method, - path = %parts - .uri - .path_and_query() - .map(|value| value.as_str()) - .unwrap_or("/"), + path = %parts.uri.path(), request_admission_ms, request_context_ms, "measured admin api keys route pre-handler timing" ); } - let mut request_body = Some(body); let local_proxy_body = if local_proxy_route_requires_buffered_body(&request_context) { - let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state); - let stage_started_at = Instant::now(); - let body = buffer_and_normalize_request_body( - &mut request_body, - &mut parts.headers, - "local proxy body buffering should own request body", - &trace_id, - &parts.method, - &request_context.request_path_and_query(), - "local_proxy", - body_buffer_policy, - ) - .await; - observe_gateway_stage_ms( - "frontdoor_body_buffer", - stage_started_at.elapsed().as_millis() as u64, - ); - match body { - Ok(body) => Some(body), - Err(err) => { - return finalize_request_body_buffer_rejection( - &state, - &request_context, - &remote_addr, - &started_at, - &trace_id, - request_permit.take(), - &err, - ); + if let Some(body) = authenticated_affinity_body.as_ref() { + Some(body.clone()) + } else { + let body_buffer_policy = + RequestBodyBufferPolicy::for_request_context(&state, &request_context); + let stage_started_at = Instant::now(); + let body = buffer_and_normalize_request_body( + &mut request_body, + &mut parts.headers, + "local proxy body buffering should own request body", + &trace_id, + &parts.method, + &request_context.request_path_and_query(), + "local_proxy", + body_buffer_policy, + ) + .await; + observe_gateway_stage_ms( + "frontdoor_body_buffer", + stage_started_at.elapsed().as_millis() as u64, + ); + match body { + Ok(body) => Some(body), + Err(err) => { + return finalize_request_body_buffer_rejection( + &state, + &request_context, + &remote_addr, + &started_at, + &trace_id, + request_permit.take(), + &err, + ); + } } } } else { @@ -1252,6 +1556,7 @@ async fn proxy_request_inner( &state, &request_context, &remote_addr, + &parts.headers, local_proxy_body.as_ref(), ) .await? @@ -1271,6 +1576,7 @@ async fn proxy_request_inner( if let Some(response) = maybe_build_local_admin_proxy_response( &state, &request_context, + &remote_addr, &parts.headers, local_proxy_body.as_ref(), ) @@ -1342,10 +1648,8 @@ async fn proxy_request_inner( &state, &request_context, &parts.headers, - parts - .extensions - .get::() - .map(|value| value.0.as_str()), + &remote_addr, + client_ip, local_proxy_body.as_ref(), ) .await @@ -1401,41 +1705,48 @@ async fn proxy_request_inner( && request_enables_control_execute(&parts.headers); let buffered_body = if should_buffer_body { - let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state); - let stage_started_at = Instant::now(); - let body = buffer_and_normalize_request_body( - &mut request_body, - &mut parts.headers, - "buffered auth/execution runtime path should own request body", - &trace_id, - &parts.method, - &request_context.request_path_and_query(), - "auth_execution", - body_buffer_policy, - ) - .await; - observe_gateway_stage_ms( - "frontdoor_body_buffer", - stage_started_at.elapsed().as_millis() as u64, - ); - match body { - Ok(body) => Some(body), - Err(err) => { - return finalize_request_body_buffer_rejection( - &state, - &request_context, - &remote_addr, - &started_at, - &trace_id, - request_permit.take(), - &err, - ); + if let Some(body) = authenticated_affinity_body.as_ref() { + Some(body.clone()) + } else { + let body_buffer_policy = + RequestBodyBufferPolicy::for_request_context(&state, &request_context); + let stage_started_at = Instant::now(); + let body = buffer_and_normalize_request_body( + &mut request_body, + &mut parts.headers, + "buffered auth/execution runtime path should own request body", + &trace_id, + &parts.method, + &request_context.request_path_and_query(), + "auth_execution", + body_buffer_policy, + ) + .await; + observe_gateway_stage_ms( + "frontdoor_body_buffer", + stage_started_at.elapsed().as_millis() as u64, + ); + match body { + Ok(body) => Some(body), + Err(err) => { + return finalize_request_body_buffer_rejection( + &state, + &request_context, + &remote_addr, + &started_at, + &trace_id, + request_permit.take(), + &err, + ); + } } } } else { None }; + // The first gateway forwards before plan admission. The owner sees the loop guard above, + // skips forwarding, and performs admission exactly once before local execution. let owner_forward_started_at = Instant::now(); let owner_forward_response = maybe_forward_public_request_to_tunnel_owner( &state, @@ -1544,7 +1855,8 @@ async fn proxy_request_inner( let api_key_id = auth_context .map(|auth_context| auth_context.api_key_id.as_str()) .unwrap_or("-"); - let path_and_query = request_context.request_path_and_query(); + let path_and_query = + crate::middleware::sanitize_access_log_path(&request_context.request_path_and_query()); info!( event_name = "frontdoor_user_rpm_rejected", log_type = "event", @@ -1591,6 +1903,106 @@ async fn proxy_request_inner( )); } + let plan_usage_started_at = Instant::now(); + let plan_usage_event_id = uuid::Uuid::new_v4().to_string(); + let now_unix_ms = chrono::Utc::now().timestamp_millis().max(0) as u64; + let plan_usage_admission = + match crate::plan_usage_policy::check_and_acquire_http_plan_usage_policy( + &state, + control_decision, + &plan_usage_event_id, + now_unix_ms, + ) + .await + { + Ok(admission) => admission, + Err(crate::plan_usage_policy::PlanUsageAdmissionError::Rejected(rejection)) => { + let response = build_local_plan_usage_limited_response( + &trace_id, + control_decision, + &rejection, + )?; + return Ok(finalize_gateway_response_with_context( + &state, + response, + &remote_addr, + &request_context, + EXECUTION_PATH_LOCAL_RATE_LIMITED, + &started_at, + request_permit.take(), + )); + } + Err(crate::plan_usage_policy::PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::Saturated { limit, .. }, + )) => { + let rejection = crate::plan_usage_policy::PlanUsagePolicyRejection { + metric: "concurrency", + limit: limit as f64, + retry_after: 1, + window: "concurrent", + }; + let response = build_local_plan_usage_limited_response( + &trace_id, + control_decision, + &rejection, + )?; + return Ok(finalize_gateway_response_with_context( + &state, + response, + &remote_addr, + &request_context, + EXECUTION_PATH_LOCAL_RATE_LIMITED, + &started_at, + request_permit.take(), + )); + } + Err(crate::plan_usage_policy::PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. }, + )) => { + let response = build_local_overloaded_response( + &trace_id, + control_decision, + Some(request_context.request_path.as_str()), + gate, + limit, + )?; + return Ok(finalize_gateway_response_with_context( + &state, + response, + &remote_addr, + &request_context, + EXECUTION_PATH_DISTRIBUTED_OVERLOADED, + &started_at, + request_permit.take(), + )); + } + Err(crate::plan_usage_policy::PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message), + )) => return Err(GatewayError::Internal(message)), + Err(crate::plan_usage_policy::PlanUsageAdmissionError::Gateway(error)) => { + return Err(error) + } + }; + let crate::plan_usage_policy::HttpPlanUsageAdmission { + permit: plan_usage_permit, + reservation_context, + } = plan_usage_admission; + if let Some(reservation_context) = reservation_context { + parts.extensions.insert(reservation_context); + } + observe_gateway_stage_ms( + "frontdoor_plan_usage", + plan_usage_started_at.elapsed().as_millis() as u64, + ); + request_permit = aether_runtime::AdmissionPermit::combine( + request_permit.into_iter().chain(plan_usage_permit), + ); + if let Some(request_permit) = request_permit.as_ref() { + parts.extensions.insert( + crate::executor::candidate_loop::BackgroundAdmissionPermit::new(request_permit.clone()), + ); + } + let local_ai_public_started_at = Instant::now(); let local_ai_public_response = super::public::maybe_build_local_ai_public_response( &state, @@ -1690,11 +2102,18 @@ async fn proxy_request_inner( route_kind = control_decision .and_then(|decision| decision.route_kind.as_deref()) .unwrap_or("-"), - request_path = %request_context.request_path_and_query(), + request_path = %crate::middleware::sanitize_access_log_path( + &request_context.request_path_and_query() + ), "gateway local stream execution returned to proxy" ); match stream_outcome { LocalExecutionRequestOutcome::Responded(execution_runtime_response) => { + crate::executor::record_failed_usage_for_deferred_response( + &state, + &execution_runtime_response, + ) + .await; let execution_runtime_response = restore_redacted_stream_execution_response( execution_runtime_response, &redaction_slot, @@ -1754,6 +2173,11 @@ async fn proxy_request_inner( ); match sync_outcome { LocalExecutionRequestOutcome::Responded(execution_runtime_response) => { + crate::executor::record_failed_usage_for_deferred_response( + &state, + &execution_runtime_response, + ) + .await; let execution_runtime_response = restore_redacted_sync_execution_response( execution_runtime_response, &redaction_slot, @@ -1815,6 +2239,11 @@ async fn proxy_request_inner( ); match stream_outcome { LocalExecutionRequestOutcome::Responded(execution_runtime_response) => { + crate::executor::record_failed_usage_for_deferred_response( + &state, + &execution_runtime_response, + ) + .await; let execution_runtime_response = restore_redacted_stream_execution_response( execution_runtime_response, &redaction_slot, @@ -1848,6 +2277,11 @@ async fn proxy_request_inner( .await? { LocalExecutionRequestOutcome::Responded(control_response) => { + crate::executor::record_failed_usage_for_deferred_response( + &state, + &control_response, + ) + .await; let reason = GatewayFallbackReason::ControlExecuteEmergency; let control_execution_path = if stream_request { EXECUTION_PATH_CONTROL_EXECUTE_STREAM @@ -1909,12 +2343,23 @@ async fn proxy_request_inner( .all_candidates_skipped_for_reason(AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON) || local_execution_runtime_miss_context .all_candidates_skipped_for_reason(LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON); - let local_execution_runtime_miss_detail = (!auth_api_key_concurrency_limited) - .then(|| { + let provider_key_capacity_limited = local_execution_runtime_miss_diagnostic + .as_ref() + .map(|diagnostic| diagnostic_is_provider_key_capacity_limited(Some(diagnostic))) + .unwrap_or_else(|| { local_execution_runtime_miss_context - .all_provider_request_body_build_failures_detail() + .all_candidates_skipped_for_reasons(PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS) + }); + let local_execution_runtime_miss_detail = provider_key_capacity_limited + .then_some(PROVIDER_KEY_CAPACITY_LIMIT_REACHED_DETAIL.to_string()) + .or_else(|| { + (!auth_api_key_concurrency_limited) + .then(|| { + local_execution_runtime_miss_context + .all_provider_request_body_build_failures_detail() + }) + .flatten() }) - .flatten() .or_else(|| { local_execution_runtime_miss_detail( control_decision, @@ -2030,7 +2475,7 @@ async fn proxy_request_inner( let mut response = build_local_http_error_response( &trace_id, control_decision, - http::StatusCode::SERVICE_UNAVAILABLE, + local_execution_runtime_miss_status(provider_key_capacity_limited), local_execution_runtime_miss_client_message( local_execution_runtime_miss_detail.as_str(), ) @@ -2356,6 +2801,30 @@ fn diagnostic_is_auth_api_key_concurrency_limited( })) } +fn diagnostic_is_provider_key_capacity_limited( + diagnostic: Option<&LocalExecutionRuntimeMissDiagnostic>, +) -> bool { + let Some(diagnostic) = diagnostic else { + return false; + }; + PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS.contains(&diagnostic.reason.as_str()) + || (diagnostic.candidate_count.is_some_and(|candidate_count| { + candidate_count > 0 + && diagnostic.skipped_candidate_count.unwrap_or(0) >= candidate_count + }) && !diagnostic.skip_reasons.is_empty() + && diagnostic.skip_reasons.iter().all(|(reason, count)| { + PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS.contains(&reason.as_str()) && *count > 0 + })) +} + +fn local_execution_runtime_miss_status(provider_key_capacity_limited: bool) -> http::StatusCode { + if provider_key_capacity_limited { + http::StatusCode::TOO_MANY_REQUESTS + } else { + http::StatusCode::SERVICE_UNAVAILABLE + } +} + fn local_execution_runtime_miss_route_detail( decision: Option<&GatewayControlDecision>, ) -> Option<&'static str> { @@ -2394,15 +2863,18 @@ mod tests { use super::{ api_key_remote_ip_allowed, buffer_and_normalize_request_body, - diagnostic_is_auth_api_key_concurrency_limited, local_execution_runtime_miss_detail, - owner_forward_request_is_stream, restore_redacted_stream_execution_response, - restore_redacted_sync_execution_response, routing_overlay_allows_affinity_target, - GatewayControlDecision, LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, - RequestBodyBufferPolicy, + diagnostic_is_auth_api_key_concurrency_limited, + diagnostic_is_provider_key_capacity_limited, local_execution_runtime_miss_detail, + local_execution_runtime_miss_status, owner_forward_request_is_stream, + restore_redacted_stream_execution_response, restore_redacted_sync_execution_response, + routing_overlay_allows_affinity_target, GatewayControlDecision, + LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, RequestBodyBufferPolicy, }; use axum::body::{to_bytes, Body, Bytes}; - use axum::http::{header, HeaderMap, HeaderValue, Method, Response}; + use axum::http::{header, HeaderMap, HeaderValue, Method, Response, StatusCode}; + use flate2::{write::GzEncoder, Compression}; use serde_json::json; + use std::io::Write; use tokio::sync::Semaphore; #[test] @@ -2635,20 +3107,68 @@ mod tests { ); } + #[tokio::test] + async fn proxy_pii_redaction_sync_wrapper_strips_all_connection_declared_headers() { + let (slot, sentinel) = redaction_slot_for_email(); + let body = serde_json::to_vec(&json!({ + "choices": [{"message": {"role": "assistant", "content": sentinel}}] + })) + .expect("response should serialize"); + let mut response = Response::builder() + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(body)) + .expect("response should build"); + response + .headers_mut() + .append(header::CONNECTION, HeaderValue::from_static("x-first-hop")); + response + .headers_mut() + .append(header::CONNECTION, HeaderValue::from_static("x-second-hop")); + response + .headers_mut() + .insert("x-first-hop", HeaderValue::from_static("first-secret")); + response + .headers_mut() + .insert("x-second-hop", HeaderValue::from_static("second-secret")); + + let restored = restore_redacted_sync_execution_response(response, &slot) + .await + .expect("sync wrapper should restore"); + + assert!(restored.headers().get(header::CONNECTION).is_none()); + assert!(restored.headers().get("x-first-hop").is_none()); + assert!(restored.headers().get("x-second-hop").is_none()); + } + #[tokio::test] async fn proxy_pii_redaction_stream_response_wrapper_restores_current_request_sentinel() { let (slot, sentinel) = redaction_slot_for_email(); - let response = Response::builder() + let mut response = Response::builder() .header(header::CONTENT_TYPE, "text/event-stream") .header(header::CONTENT_LENGTH, "999") .body(Body::from(format!( "data: {{\"choices\":[{{\"delta\":{{\"content\":\"hello {sentinel}\"}}}}]}}\n\n" ))) .expect("response should build"); + response + .headers_mut() + .append(header::CONNECTION, HeaderValue::from_static("x-first-hop")); + response + .headers_mut() + .append(header::CONNECTION, HeaderValue::from_static("x-second-hop")); + response + .headers_mut() + .insert("x-first-hop", HeaderValue::from_static("first-secret")); + response + .headers_mut() + .insert("x-second-hop", HeaderValue::from_static("second-secret")); let restored = restore_redacted_stream_execution_response(response, &slot) .expect("stream wrapper should restore"); assert!(restored.headers().get(header::CONTENT_LENGTH).is_none()); + assert!(restored.headers().get(header::CONNECTION).is_none()); + assert!(restored.headers().get("x-first-hop").is_none()); + assert!(restored.headers().get("x-second-hop").is_none()); let body = to_bytes(restored.into_body(), usize::MAX) .await .expect("body should read"); @@ -2704,6 +3224,51 @@ mod tests { )); } + #[tokio::test] + async fn request_body_buffer_caps_decompressed_body_at_shared_budget() { + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder + .write_all(&vec![b'a'; 128]) + .expect("test gzip body should encode"); + let encoded = encoder.finish().expect("test gzip body should finish"); + assert!( + encoded.len() <= 64, + "fixture must fit the compressed budget" + ); + + let mut body = Some(Body::from(encoded)); + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + let err = buffer_and_normalize_request_body( + &mut body, + &mut headers, + "test owns body", + "trace-body-decompression-budget", + &Method::POST, + "/v1/responses", + "test", + RequestBodyBufferPolicy::for_tests_with_budget( + 1024, + Duration::from_secs(1), + Duration::from_secs(1), + 64, + Arc::new(Semaphore::new(1)), + ), + ) + .await + .expect_err("decoded body above the shared budget must fail closed"); + + assert!(matches!( + err, + RequestBodyBufferError::Normalization( + crate::headers::RequestBodyNormalizationError::DecompressedBodyTooLarge { + limit_bytes: 64, + .. + } + ) + )); + } + #[tokio::test] async fn request_body_buffer_times_out_instead_of_waiting_forever() { let stream = async_stream::stream! { @@ -2732,6 +3297,36 @@ mod tests { )); } + #[tokio::test] + async fn request_body_buffer_without_read_timeout_allows_slow_body_to_complete() { + let stream = async_stream::stream! { + yield Ok::(Bytes::from_static(b"{")); + tokio::time::sleep(Duration::from_millis(20)).await; + yield Ok::(Bytes::from_static(b"}")); + }; + let mut body = Some(Body::from_stream(stream)); + let mut headers = HeaderMap::new(); + + let buffered = tokio::time::timeout( + Duration::from_secs(1), + buffer_and_normalize_request_body( + &mut body, + &mut headers, + "test owns body", + "trace-body-no-timeout", + &Method::POST, + "/v1/responses", + "test", + RequestBodyBufferPolicy::for_tests_without_read_timeout(1024), + ), + ) + .await + .expect("test body should finish") + .expect("disabled read timeout should allow a slow body"); + + assert_eq!(buffered.as_ref(), b"{}"); + } + #[tokio::test] async fn request_body_buffer_rejects_when_weighted_budget_is_exhausted() { let budget = Arc::new(Semaphore::new(1)); @@ -2881,6 +3476,45 @@ mod tests { Some("当前调用方 API Key 并发请求数已达上限,请稍后重试") ); } + + #[test] + fn provider_key_capacity_requires_every_skip_reason_to_be_capacity_related() { + let capacity_limited = LocalExecutionRuntimeMissDiagnostic { + reason: "candidate_evaluation_incomplete".to_string(), + candidate_count: Some(2), + skipped_candidate_count: Some(2), + skip_reasons: std::collections::BTreeMap::from([ + ("provider_key_concurrency_limit_reached".to_string(), 1), + ("key_rpm_exhausted".to_string(), 1), + ]), + ..LocalExecutionRuntimeMissDiagnostic::default() + }; + let mixed_failure = LocalExecutionRuntimeMissDiagnostic { + reason: "all_candidates_skipped".to_string(), + candidate_count: Some(2), + skipped_candidate_count: Some(2), + skip_reasons: std::collections::BTreeMap::from([ + ("provider_key_concurrency_limit_reached".to_string(), 1), + ("account_quota_exhausted".to_string(), 1), + ]), + ..LocalExecutionRuntimeMissDiagnostic::default() + }; + + assert!(diagnostic_is_provider_key_capacity_limited(Some( + &capacity_limited + ))); + assert!(!diagnostic_is_provider_key_capacity_limited(Some( + &mixed_failure + ))); + assert_eq!( + local_execution_runtime_miss_status(true), + StatusCode::TOO_MANY_REQUESTS + ); + assert_eq!( + local_execution_runtime_miss_status(false), + StatusCode::SERVICE_UNAVAILABLE + ); + } } #[path = "finalize.rs"] diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs b/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs index 39a9ad888..2c724408f 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs @@ -1027,6 +1027,7 @@ mod tests { local_rejection: None, allowed_models: None, ip_rules: None, + verified_api_key_hash: None, }); let diagnostic = LocalExecutionRuntimeMissDiagnostic { reason: "candidate_list_empty".to_string(), @@ -1207,6 +1208,7 @@ mod tests { local_rejection: None, allowed_models: None, ip_rules: None, + verified_api_key_hash: None, }); record_live_websocket_preflight_failure( diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/live/http.rs b/apps/aether-gateway/src/handlers/proxy/websocket/live/http.rs index 0e2bd589c..45b71b063 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/live/http.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/live/http.rs @@ -19,6 +19,9 @@ use crate::api::response::{ }; use crate::control::{execution_plan_balance_capacity_rejection, GatewayPublicRequestContext}; use crate::execution_runtime::execute_execution_runtime_sync_plan_with_report_context; +use crate::execution_runtime::transport::{ + decode_base64_body_with_limit, serialize_json_body_with_limit, +}; use crate::handlers::proxy::websocket::responses::ResponsesWebSocketTurnAdmission; use crate::{AppState, GatewayError}; @@ -609,6 +612,10 @@ fn gateway_error_status(error: &GatewayError) -> StatusCode { } GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT, GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS, + GatewayError::PlanUsageLimited(_) => StatusCode::TOO_MANY_REQUESTS, + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + StatusCode::BAD_REQUEST + } GatewayError::Client { status, .. } => *status, GatewayError::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR, } @@ -748,10 +755,18 @@ fn execution_result_body(result: &ExecutionResult) -> Result Result { StatusCode::TOO_MANY_REQUESTS } + Self::Gateway(GatewayError::PlanUsageLimited(_)) => StatusCode::TOO_MANY_REQUESTS, + Self::Gateway(GatewayError::LastActiveAdminUpdateDenied) + | Self::Gateway(GatewayError::LastActiveAdminDeleteDenied) => StatusCode::BAD_REQUEST, Self::Gateway(GatewayError::Client { status, .. }) => *status, Self::Gateway(GatewayError::LocalExecutionPlanningTimeout { .. }) => { StatusCode::GATEWAY_TIMEOUT @@ -114,6 +117,13 @@ impl LiveRelayAdmissionError { Self::Gateway(GatewayError::AdmissionTimeout { .. }) => { "Gateway capacity is busy; retry this Live connection" } + Self::Gateway(GatewayError::PlanUsageLimited(_)) => { + "Subscription plan usage limit reached" + } + Self::Gateway(GatewayError::LastActiveAdminUpdateDenied) + | Self::Gateway(GatewayError::LastActiveAdminDeleteDenied) => { + "Codex Live request was not allowed" + } Self::Gateway(GatewayError::Client { .. }) => "Codex Live request was not allowed", Self::Gateway(GatewayError::LocalExecutionPlanningTimeout { .. }) => { "Codex Live admission planning timed out" @@ -127,6 +137,9 @@ impl LiveRelayAdmissionError { Self::PlanUnavailable => "admission_plan_unavailable", Self::BalanceRejected => "balance_rejected", Self::Gateway(GatewayError::AdmissionTimeout { .. }) => "admission_timeout", + Self::Gateway(GatewayError::PlanUsageLimited(_)) => "plan_usage_limited", + Self::Gateway(GatewayError::LastActiveAdminUpdateDenied) => "last_admin_update_denied", + Self::Gateway(GatewayError::LastActiveAdminDeleteDenied) => "last_admin_delete_denied", Self::Gateway(GatewayError::Client { .. }) => "request_rejected", Self::Gateway(GatewayError::LocalExecutionPlanningTimeout { .. }) => { "admission_planning_timeout" @@ -1226,6 +1239,9 @@ fn gateway_error_kind(error: &GatewayError) -> &'static str { GatewayError::ControlUnavailable { .. } => "control_unavailable", GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timeout", GatewayError::AdmissionTimeout { .. } => "admission_timeout", + GatewayError::PlanUsageLimited(_) => "plan_usage_limited", + GatewayError::LastActiveAdminUpdateDenied => "last_admin_update_denied", + GatewayError::LastActiveAdminDeleteDenied => "last_admin_delete_denied", GatewayError::Client { .. } => "client_error", GatewayError::Internal(_) => "internal_error", } @@ -1238,6 +1254,10 @@ fn gateway_error_status(error: &GatewayError) -> StatusCode { } GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT, GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS, + GatewayError::PlanUsageLimited(_) => StatusCode::TOO_MANY_REQUESTS, + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + StatusCode::BAD_REQUEST + } GatewayError::Client { status, .. } => *status, GatewayError::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR, } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/realtime/session.rs b/apps/aether-gateway/src/handlers/proxy/websocket/realtime/session.rs index 4845244c9..db40505ce 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/realtime/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/realtime/session.rs @@ -552,6 +552,10 @@ fn rejection( fn admission_error_status(error: &GatewayError) -> StatusCode { match error { GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS, + GatewayError::PlanUsageLimited(_) => StatusCode::TOO_MANY_REQUESTS, + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + StatusCode::BAD_REQUEST + } GatewayError::Client { status, .. } => *status, GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT, _ => StatusCode::INTERNAL_SERVER_ERROR, @@ -564,6 +568,9 @@ fn gateway_error_kind(error: &GatewayError) -> &'static str { GatewayError::ControlUnavailable { .. } => "control_unavailable", GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timeout", GatewayError::AdmissionTimeout { .. } => "admission_timeout", + GatewayError::PlanUsageLimited(_) => "plan_usage_limited", + GatewayError::LastActiveAdminUpdateDenied => "last_admin_update_denied", + GatewayError::LastActiveAdminDeleteDenied => "last_admin_delete_denied", GatewayError::Client { .. } => "client_error", GatewayError::Internal(_) => "internal_error", } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs index 96d11c344..e1a200af8 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs @@ -17,6 +17,9 @@ use super::ownership::{ await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease, spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision, }; +use super::plan_admission::{ + acquire_responses_websocket_plan_admission, send_responses_websocket_plan_admission_error, +}; use super::quota::mark_active_response_retry_unsafe; use super::redaction::redact_responses_websocket_client_event_with_reasoning_replay_policy; use super::request::{ @@ -35,6 +38,8 @@ use super::turn::{ use super::turn_state::LogicalTurn; use super::upstream::{ bind_responses_upstream, decision_bound_upstream_change_fields, decision_reuses_bound_upstream, + send_responses_websocket_upstream_message, ResponsesWebSocketUpstreamBindError, + ResponsesWebSocketUpstreamSendError, }; use crate::ai_serving::ResponsesWebSocketPinnedCandidate; use crate::clock::current_unix_secs; @@ -44,8 +49,9 @@ use crate::handlers::proxy::websocket::session::{CLOSE_INTERNAL_ERROR, WEBSOCKET use crate::handlers::proxy::websocket::transport::{ close_client_socket, close_upstream_socket, send_client_message, send_gateway_error, send_gateway_error_with_status, send_gateway_error_with_stream_id, - send_responses_websocket_error_with_param, send_upstream_message, + send_responses_websocket_error_with_param, }; +use crate::plan_usage_policy::PlanUsagePolicySnapshot; use crate::privacy::RedactionSession; use crate::rate_limit::FrontdoorUserRpmOutcome; use crate::AppState; @@ -69,6 +75,7 @@ pub(super) enum RelayDisposition { Continue, Close, UpstreamError(&'static str), + PlanUsagePermitLost, } pub(super) fn adapter_drain_ready( @@ -317,6 +324,37 @@ pub(super) async fn forward_client_message( } } + // Subscription-plan concurrency is admitted per logical + // response.create. Keep the permit and policy snapshot attached + // to this turn through any transparent provider rebind. + let plan_usage_admission = match acquire_responses_websocket_plan_admission( + state, + &turn_control.decision, + &logical_turn_id, + ) + .await + { + Ok(admission) => admission, + Err(error) => { + warn!( + event_name = "responses_websocket_followup_plan_usage_rejected", + log_type = "event", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + logical_turn_id = %logical_turn_id, + error = ?error, + "gateway rejected a Responses WebSocket follow-up at its subscription plan limit" + ); + send_responses_websocket_plan_admission_error(client_socket, &error).await; + return RelayDisposition::Continue; + } + }; + let crate::plan_usage_policy::PlanUsageAdmission { + permit: plan_usage_permit, + policy_snapshot: plan_usage_policy_snapshot, + } = plan_usage_admission; + let response_id_owned_by_connection = client_event .get("previous_response_id") .and_then(Value::as_str) @@ -465,6 +503,8 @@ pub(super) async fn forward_client_message( codex_fingerprint_context, turn_control, turn_redaction_session, + plan_usage_permit, + plan_usage_policy_snapshot, ) .await; } @@ -482,6 +522,8 @@ pub(super) async fn forward_client_message( raw_responses_lite_static_config .expect("independent turns always retain their raw static config"), turn_redaction_session, + plan_usage_permit, + plan_usage_policy_snapshot, ) .await } @@ -529,6 +571,8 @@ async fn forward_pinned_continuation( codex_fingerprint_context: CodexFingerprintConvergenceContext, turn_control: ResponsesWebSocketTurnControl, turn_redaction_session: Option, + plan_usage_permit: Option, + plan_usage_policy_snapshot: Option, ) -> RelayDisposition { let Some(pinned_candidate) = ResponsesWebSocketPinnedCandidate::from_decision(&bound.decision_template) @@ -725,6 +769,7 @@ async fn forward_pinned_continuation( &turn_control.decision, turn_decision, &client_event, + plan_usage_policy_snapshot.clone(), planned_lease, ) .await @@ -755,21 +800,42 @@ async fn forward_pinned_continuation( .await; return RelayDisposition::UpstreamError("responses_websocket_send_failed"); }; - if send_upstream_message(upstream, WreqWsMessage::text(outbound)) - .await - .is_err() + match send_responses_websocket_upstream_message( + upstream, + WreqWsMessage::text(outbound), + plan_usage_permit.as_ref(), + |request_state| turn.record_upstream_request_state(request_state), + ) + .await { - queue_turn_finalization( - bound, - state, - turn, - ResponsesWebSocketTurnOutcome::upstream_send_failed(), - ) - .await; - return RelayDisposition::UpstreamError("responses_websocket_send_failed"); + Ok(()) => {} + Err(ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost) => { + turn.release_plan_usage_cost_before_upstream_send( + state, + "continuation_plan_usage_permit_lost", + ) + .await; + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ) + .await; + return RelayDisposition::PlanUsagePermitLost; + } + Err(ResponsesWebSocketUpstreamSendError::Transport(_)) => { + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::upstream_send_failed(), + ) + .await; + return RelayDisposition::UpstreamError("responses_websocket_send_failed"); + } } - turn.mark_upstream_request_sent(); turn.set_provider_response_headers(bound.upstream_response_headers.clone()); if let Some(session) = turn_redaction_session { bound.redaction_restorer.register(session); @@ -781,7 +847,9 @@ async fn forward_pinned_continuation( LogicalTurn::new(client_event, turn_index, logical_turn_id) .with_codex_fingerprint_context(codex_fingerprint_context) .with_provider_store(provider_event.get("store") == Some(&Value::Bool(true))) - .with_turn_control(turn_control), + .with_turn_control(turn_control) + .with_plan_usage_permit(plan_usage_permit) + .with_plan_usage_policy_snapshot(plan_usage_policy_snapshot), turn, ); bound.next_turn_index = bound.next_turn_index.saturating_add(1); @@ -819,6 +887,8 @@ async fn forward_replanned_response_create( turn_control: ResponsesWebSocketTurnControl, raw_responses_lite_static_config: ResponsesLiteStaticConfig, turn_redaction_session: Option, + plan_usage_permit: Option, + plan_usage_policy_snapshot: Option, ) -> RelayDisposition { let turn_request_id = Uuid::new_v4().to_string(); let now_unix_secs = current_unix_secs(); @@ -919,6 +989,7 @@ async fn forward_replanned_response_create( &turn_control.decision, turn_decision, &client_event, + plan_usage_policy_snapshot.clone(), planned_lease, ) .await @@ -970,18 +1041,40 @@ async fn forward_replanned_response_create( .await; return RelayDisposition::UpstreamError("responses_websocket_send_failed"); }; - if send_upstream_message(upstream, WreqWsMessage::text(outbound)) - .await - .is_err() + match send_responses_websocket_upstream_message( + upstream, + WreqWsMessage::text(outbound), + plan_usage_permit.as_ref(), + |request_state| turn.record_upstream_request_state(request_state), + ) + .await { - queue_turn_finalization( - bound, - state, - turn, - ResponsesWebSocketTurnOutcome::upstream_send_failed(), - ) - .await; - return RelayDisposition::UpstreamError("responses_websocket_send_failed"); + Ok(()) => {} + Err(ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost) => { + turn.release_plan_usage_cost_before_upstream_send( + state, + "replanned_reuse_plan_usage_permit_lost", + ) + .await; + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ) + .await; + return RelayDisposition::PlanUsagePermitLost; + } + Err(ResponsesWebSocketUpstreamSendError::Transport(_)) => { + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::upstream_send_failed(), + ) + .await; + return RelayDisposition::UpstreamError("responses_websocket_send_failed"); + } } // A response.create without previous_response_id starts a new chain. @@ -991,7 +1084,6 @@ async fn forward_replanned_response_create( bound .redaction_restorer .start_new_chain(turn_redaction_session); - turn.mark_upstream_request_sent(); turn.set_provider_response_headers(bound.upstream_response_headers.clone()); let provider_model = provider_model_from_decision(&decision).unwrap_or_else(|| bound.provider_model.clone()); @@ -1009,7 +1101,9 @@ async fn forward_replanned_response_create( LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone()) .with_codex_fingerprint_context(codex_fingerprint_context.clone()) .with_provider_store(provider_event.get("store") == Some(&Value::Bool(true))) - .with_turn_control(turn_control), + .with_turn_control(turn_control) + .with_plan_usage_permit(plan_usage_permit) + .with_plan_usage_policy_snapshot(plan_usage_policy_snapshot), turn, ); bound.next_turn_index = bound.next_turn_index.saturating_add(1); @@ -1031,42 +1125,64 @@ async fn forward_replanned_response_create( return RelayDisposition::Continue; } - let mut replacement = - match bind_responses_upstream(&decision, normalization, &client_event, adapter).await { - Ok(connection) => connection, - Err(code) => { - queue_turn_finalization( - bound, - state, - turn, - ResponsesWebSocketTurnOutcome::upstream_connect_failed(code), - ) - .await; - warn!( - event_name = "responses_websocket_followup_model_rebind_failed", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %context.trace_id, - requested_model = %requested_model, - error_code = code, - "gateway failed to rebind Responses WebSocket follow-up model" - ); - send_gateway_error_with_status( - client_socket, - 502, - code, - "Gateway could not establish the requested model", - ) - .await; - return RelayDisposition::Continue; - } - }; + let mut replacement = match bind_responses_upstream( + &decision, + normalization, + &client_event, + adapter, + plan_usage_permit.as_ref(), + |request_state| turn.record_upstream_request_state(request_state), + ) + .await + { + Ok(connection) => connection, + Err(ResponsesWebSocketUpstreamBindError::PlanUsagePermitLost) => { + turn.release_plan_usage_cost_before_upstream_send( + state, + "replanned_rebind_plan_usage_permit_lost", + ) + .await; + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ) + .await; + return RelayDisposition::PlanUsagePermitLost; + } + Err(ResponsesWebSocketUpstreamBindError::Transport(code)) => { + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::upstream_connect_failed(code), + ) + .await; + warn!( + event_name = "responses_websocket_followup_model_rebind_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + requested_model = %requested_model, + error_code = code, + "gateway failed to rebind Responses WebSocket follow-up model" + ); + send_gateway_error_with_status( + client_socket, + 502, + code, + "Gateway could not establish the requested model", + ) + .await; + return RelayDisposition::Continue; + } + }; if replacement.responses_lite_static_config.is_some() { replacement.responses_lite_static_config = Some(raw_responses_lite_static_config); } - turn.mark_upstream_request_sent(); turn.set_provider_response_headers(replacement.upstream_response_headers.clone()); let previous_client_model = bound.client_model.clone(); let previous_provider_model = bound.provider_model.clone(); @@ -1094,7 +1210,9 @@ async fn forward_replanned_response_create( LogicalTurn::new(client_event, turn_index, logical_turn_id) .with_codex_fingerprint_context(codex_fingerprint_context) .with_provider_store(provider_event.get("store") == Some(&Value::Bool(true))) - .with_turn_control(turn_control), + .with_turn_control(turn_control) + .with_plan_usage_permit(plan_usage_permit) + .with_plan_usage_policy_snapshot(plan_usage_policy_snapshot), turn, ); bound.next_turn_index = bound.next_turn_index.saturating_add(1); diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs index 64ca00a2f..5749b714c 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs @@ -17,6 +17,7 @@ use super::lifecycle::{ await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization, settle_turn_finalization, spawn_bounded_adapter_observation, PreviousAttemptSettled, }; +use super::plan_admission::terminate_responses_websocket_for_plan_permit_loss; use super::quota::{ detach_exhausted_upstream, is_usage_limit_error_event, mark_active_response_retry_unsafe, observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion, @@ -144,6 +145,19 @@ pub(super) async fn relay_bound_connection( ).await; break; } + RelayDisposition::PlanUsagePermitLost => { + warn!( + event_name = "responses_websocket_plan_usage_concurrency_lost_before_send", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a Responses WebSocket turn before its upstream send after the subscription plan concurrency lease became unhealthy" + ); + close_bound_upstream(bound).await; + terminate_responses_websocket_for_plan_permit_loss(client_socket).await; + break; + } RelayDisposition::UpstreamError(code) => { warn!( event_name = "responses_websocket_upstream_send_failed", @@ -469,18 +483,33 @@ pub(super) async fn relay_bound_connection( ) .await } - None => PreviousAttemptSettled::nothing_to_settle(), + None => Some(PreviousAttemptSettled::nothing_to_settle()), }; // Planning and binding a replacement carries the complete // scheduler/provider state machine. Keep that large future // off the relay task's stack; the default Tokio/test worker // stack is otherwise easy to exhaust on this rare branch. - if Box::pin(retry_active_turn_after_quota_exhaustion( - bound, state, context, settled, - )) - .await - { - continue; + if let Some(settled) = settled { + match Box::pin(retry_active_turn_after_quota_exhaustion( + bound, state, context, settled, + )) + .await + { + Ok(true) => continue, + Ok(false) => {} + Err(()) => { + finalize_active_turn( + bound, + state, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ) + .await; + close_bound_upstream(bound).await; + terminate_responses_websocket_for_plan_permit_loss(client_socket) + .await; + break; + } + } } // 重试失败。旧 attempt 已经结算,logical turn 仍停在 // Replanning,所以后面分支里的 end() / finalize_active_turn @@ -531,27 +560,32 @@ pub(super) async fn relay_bound_connection( let mut relay_serialization_failed = false; match relay_directive { Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => { - let restored = parsed_upstream_frame - .as_ref() - .and_then(|frame| { - bound - .redaction_restorer - .restore_provider_frame_text(frame.event()) - }); - let client_frame = match restored { - Some(text) => AxumWsMessage::Text(text.into()), - None => upstream_message_to_client(upstream_message.clone()), + let client_frame = match parsed_upstream_frame.as_ref().map(|frame| { + bound + .redaction_restorer + .restore_provider_frame_text(frame.event()) + }) { + Some(Ok(Some(text))) => Some(AxumWsMessage::Text(text.into())), + Some(Ok(None)) | None => { + Some(upstream_message_to_client(upstream_message.clone())) + } + Some(Err(_)) => { + relay_serialization_failed = true; + None + } }; - match send_client_message(client_socket, client_frame).await { - Ok(()) => { - if let (Some(turn), Some(frame)) = ( - bound.turn_state.attempt_mut(), - parsed_upstream_frame.as_ref(), - ) { - turn.capture_client_frame(frame.event()); + if let Some(client_frame) = client_frame { + match send_client_message(client_socket, client_frame).await { + Ok(()) => { + if let (Some(turn), Some(frame)) = ( + bound.turn_state.attempt_mut(), + parsed_upstream_frame.as_ref(), + ) { + turn.capture_client_frame(frame.event()); + } } + Err(error) => relay_send_error = Some(error), } - Err(error) => relay_send_error = Some(error), } } Some(ResponsesWebSocketRelayDirective::ForwardEvents(events)) => { @@ -560,14 +594,18 @@ pub(super) async fn relay_bound_connection( .redaction_restorer .restore_provider_frame_text(event) { - Some(restored) => restored, - None => match encode_opaque_websocket_event(event) { + Ok(Some(restored)) => restored, + Ok(None) => match encode_opaque_websocket_event(event) { Ok(encoded) => encoded, Err(_) => { relay_serialization_failed = true; break; } }, + Err(_) => { + relay_serialization_failed = true; + break; + } }; match send_client_message( client_socket, diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs index 708517b66..d0db88cc5 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs @@ -18,6 +18,7 @@ use crate::handlers::proxy::websocket::session::{ CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT, }; use crate::handlers::proxy::websocket::transport::send_responses_websocket_error; +use crate::plan_usage_policy::PlanUsagePolicySnapshot; use crate::{AppState, GatewayError}; const RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT: Duration = Duration::from_secs(5); @@ -85,6 +86,7 @@ pub(super) async fn begin_responses_websocket_turn( control_decision: &crate::control::GatewayControlDecision, decision: crate::ai_serving::AiExecutionDecision, client_event: &serde_json::Value, + plan_usage_policy_snapshot: Option, ) -> Result { let state = state.clone(); let trace_id = trace_id.to_string(); @@ -108,6 +110,7 @@ pub(super) async fn begin_responses_websocket_turn( &control_decision, decision, &client_event, + plan_usage_policy_snapshot, ) .await?; Ok(ActiveProviderAttempt::new(&state, turn)) @@ -231,12 +234,29 @@ impl PreviousAttemptSettled { pub(super) async fn settle_turn_finalization( bound: &mut BoundResponsesConnection, state: &AppState, - turn: ActiveProviderAttempt, + mut turn: ActiveProviderAttempt, outcome: ResponsesWebSocketTurnOutcome, -) -> PreviousAttemptSettled { +) -> Option { + // This awaited finalization entry is used only by transparent quota retry. + // Free the old attempt's durable cost reservation before planning can + // reserve the replacement, otherwise two per-attempt estimates briefly + // count against the same user and can manufacture a false local 429. + if let Err(error) = turn.release_plan_usage_cost_for_retry(state).await { + warn!( + event_name = "responses_websocket_plan_usage_cost_retry_release_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + error = ?error, + "gateway stopped a transparent Responses WebSocket retry because the old plan cost reservation could not be released" + ); + queue_turn_finalization(bound, state, turn, outcome).await; + await_pending_turn_finalization(bound).await; + return None; + } queue_turn_finalization(bound, state, turn, outcome).await; await_pending_turn_finalization(bound).await; - PreviousAttemptSettled(()) + Some(PreviousAttemptSettled(())) } pub(super) fn spawn_bounded_adapter_observation( @@ -334,6 +354,16 @@ pub(super) async fn send_responses_websocket_turn_start_error( ) { let status_code = responses_websocket_turn_start_http_status(error); match error { + GatewayError::PlanUsageLimited(_) => { + send_responses_websocket_error( + client_socket, + status_code, + "rate_limit_error", + "plan_usage_limit_exceeded", + "Subscription plan usage limit exceeded; retry later", + ) + .await; + } GatewayError::Client { status, message } => { let (error_type, code) = if status.as_u16() == 429 { ("rate_limit_error", "gateway_request_capacity_exceeded") @@ -378,6 +408,7 @@ pub(super) async fn send_responses_websocket_turn_start_error( fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 { match error { + GatewayError::PlanUsageLimited(_) => StatusCode::TOO_MANY_REQUESTS.as_u16(), GatewayError::Client { status, .. } => status.as_u16(), GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS.as_u16(), GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT.as_u16(), @@ -387,6 +418,7 @@ fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 { pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) { match error { + GatewayError::PlanUsageLimited(_) => (CLOSE_TRY_AGAIN, "plan_usage_limit_exceeded"), GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"), GatewayError::AdmissionTimeout { .. } | GatewayError::LocalExecutionPlanningTimeout { .. } => (CLOSE_TRY_AGAIN, "gateway_busy"), @@ -405,6 +437,7 @@ mod tests { await_turn_finalization_handle, responses_websocket_turn_start_close, responses_websocket_turn_start_http_status, spawn_bounded_adapter_observation_with_timeout, }; + use crate::plan_usage_policy::PlanUsagePolicyRejection; use crate::GatewayError; #[test] @@ -422,6 +455,22 @@ mod tests { ); } + #[test] + fn plan_usage_limit_uses_http_429_and_retry_later_close_code() { + let error = GatewayError::PlanUsageLimited(PlanUsagePolicyRejection { + metric: "actual_cost_usd", + limit: 10.0, + retry_after: 60, + window: "calendar_month", + }); + + assert_eq!(responses_websocket_turn_start_http_status(&error), 429); + assert_eq!( + responses_websocket_turn_start_close(&error), + (1013, "plan_usage_limit_exceeded") + ); + } + /// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。 /// /// 透明重试在这之后立刻按 health / adaptive / pool 状态规划下一个 attempt, diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs index b262513fe..2ea47ea6b 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs @@ -18,6 +18,7 @@ mod frame; mod lifecycle; mod observation; mod ownership; +mod plan_admission; mod quota; mod redaction; mod relay_policy; diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs index 361c264ee..ce27af50b 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs @@ -18,6 +18,7 @@ use crate::ai_serving::{ }; use crate::control::GatewayControlDecision; use crate::orchestration::release_pool_key_lease_from_report_context; +use crate::plan_usage_policy::PlanUsagePolicySnapshot; use crate::{AppState, GatewayError}; /// Owns a selected pool-key lease until the attempt lifecycle has taken over @@ -149,6 +150,7 @@ pub(super) async fn begin_responses_websocket_turn_with_planned_lease( control_decision: &GatewayControlDecision, decision: AiExecutionDecision, client_event: &Value, + plan_usage_policy_snapshot: Option, mut planned_lease: PlannedPoolKeyLeaseGuard, ) -> Result { let state = state.clone(); @@ -163,6 +165,7 @@ pub(super) async fn begin_responses_websocket_turn_with_planned_lease( &control_decision, decision, &client_event, + plan_usage_policy_snapshot, ) .await?; // ActiveProviderAttempt now owns the report context containing the diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/plan_admission.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/plan_admission.rs new file mode 100644 index 000000000..a1ff2717e --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/plan_admission.rs @@ -0,0 +1,232 @@ +//! Subscription-plan request and concurrency admission for logical WS turns. +//! +//! The HTTP Upgrade is only a transport connection. Each accepted +//! `response.create` is admitted here as one logical request, while transparent +//! provider retries reuse the permit stored on that logical turn. + +use axum::extract::ws::WebSocket; +use std::time::Duration; + +use crate::handlers::proxy::websocket::session::{CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN}; +use crate::handlers::proxy::websocket::transport::{ + close_client_socket, send_responses_websocket_error, +}; +use crate::plan_usage_policy::PlanUsageAdmissionError; +use crate::{AppState, GatewayError}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct ResponsesWebSocketPlanAdmissionError { + status: u16, + error_type: &'static str, + code: &'static str, + message: &'static str, + close_code: u16, + close_reason: &'static str, +} + +pub(super) async fn acquire_responses_websocket_plan_admission( + state: &AppState, + decision: &crate::control::GatewayControlDecision, + logical_turn_id: &str, +) -> Result { + crate::plan_usage_policy::check_and_acquire_plan_usage_policy_admission( + state, + Some(decision), + logical_turn_id, + crate::clock::current_unix_ms(), + ) + .await +} + +pub(super) fn responses_websocket_plan_permit_is_healthy( + permit: Option<&aether_runtime::AdmissionPermit>, +) -> bool { + permit.is_none_or(aether_runtime::AdmissionPermit::is_healthy) +} + +pub(super) async fn send_responses_websocket_plan_admission_error( + client_socket: &mut WebSocket, + error: &PlanUsageAdmissionError, +) { + let mapped = map_plan_admission_error(error); + send_responses_websocket_error( + client_socket, + mapped.status, + mapped.error_type, + mapped.code, + mapped.message, + ) + .await; +} + +pub(super) fn responses_websocket_plan_admission_close( + error: &PlanUsageAdmissionError, +) -> (u16, &'static str) { + let mapped = map_plan_admission_error(error); + (mapped.close_code, mapped.close_reason) +} + +pub(super) async fn wait_for_admission_permit_loss( + permit: Option<&aether_runtime::AdmissionPermit>, +) { + wait_for_admission_permit_loss_with_interval(permit, Duration::from_secs(1)).await; +} + +async fn wait_for_admission_permit_loss_with_interval( + permit: Option<&aether_runtime::AdmissionPermit>, + health_poll_interval: Duration, +) { + let Some(permit) = permit else { + std::future::pending::<()>().await; + return; + }; + let mut health = tokio::time::interval(health_poll_interval); + health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + loop { + health.tick().await; + if !permit.is_healthy() { + return; + } + } +} + +pub(super) async fn terminate_responses_websocket_for_plan_permit_loss( + client_socket: &mut WebSocket, +) { + let mapped = plan_permit_loss_error(); + send_responses_websocket_error( + client_socket, + mapped.status, + mapped.error_type, + mapped.code, + mapped.message, + ) + .await; + close_client_socket(client_socket, mapped.close_code, mapped.close_reason).await; +} + +fn map_plan_admission_error( + error: &PlanUsageAdmissionError, +) -> ResponsesWebSocketPlanAdmissionError { + match error { + PlanUsageAdmissionError::Rejected(_) => plan_limit_exceeded_error(), + PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::Saturated { .. }, + ) => plan_limit_exceeded_error(), + PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::Unavailable { .. }, + ) => ResponsesWebSocketPlanAdmissionError { + status: 503, + error_type: "server_error", + code: "plan_usage_policy_unavailable", + message: "Gateway could not evaluate the subscription plan limit", + close_code: CLOSE_TRY_AGAIN, + close_reason: "plan_usage_policy_unavailable", + }, + PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(_), + ) + | PlanUsageAdmissionError::Gateway(GatewayError::Internal(_)) => { + ResponsesWebSocketPlanAdmissionError { + status: 500, + error_type: "server_error", + code: "plan_usage_policy_unavailable", + message: "Gateway could not evaluate the subscription plan limit", + close_code: CLOSE_INTERNAL_ERROR, + close_reason: "plan_usage_policy_unavailable", + } + } + PlanUsageAdmissionError::Gateway(GatewayError::Client { status, .. }) => { + ResponsesWebSocketPlanAdmissionError { + status: status.as_u16(), + error_type: "invalid_request_error", + code: "plan_usage_policy_unavailable", + message: "Gateway could not evaluate the subscription plan limit", + close_code: CLOSE_INTERNAL_ERROR, + close_reason: "plan_usage_policy_unavailable", + } + } + PlanUsageAdmissionError::Gateway(_) => ResponsesWebSocketPlanAdmissionError { + status: 500, + error_type: "server_error", + code: "plan_usage_policy_unavailable", + message: "Gateway could not evaluate the subscription plan limit", + close_code: CLOSE_INTERNAL_ERROR, + close_reason: "plan_usage_policy_unavailable", + }, + } +} + +const fn plan_permit_loss_error() -> ResponsesWebSocketPlanAdmissionError { + ResponsesWebSocketPlanAdmissionError { + status: 503, + error_type: "server_error", + code: "plan_usage_concurrency_unavailable", + message: "Subscription plan concurrency admission was lost; retry this response", + close_code: CLOSE_TRY_AGAIN, + close_reason: "plan_usage_concurrency_unavailable", + } +} + +const fn plan_limit_exceeded_error() -> ResponsesWebSocketPlanAdmissionError { + ResponsesWebSocketPlanAdmissionError { + status: 429, + error_type: "rate_limit_error", + code: "plan_usage_limit_exceeded", + message: "Subscription plan usage limit exceeded; retry later", + close_code: CLOSE_TRY_AGAIN, + close_reason: "plan_usage_limit_exceeded", + } +} + +#[cfg(test)] +mod tests { + use super::{ + map_plan_admission_error, plan_permit_loss_error, responses_websocket_plan_admission_close, + }; + use crate::plan_usage_policy::{PlanUsageAdmissionError, PlanUsagePolicyRejection}; + + #[test] + fn plan_limit_rejection_is_a_machine_readable_429() { + let error = PlanUsageAdmissionError::Rejected(PlanUsagePolicyRejection { + metric: "request_count", + limit: 10.0, + retry_after: 60, + window: "rolling", + }); + + let mapped = map_plan_admission_error(&error); + assert_eq!(mapped.status, 429); + assert_eq!(mapped.error_type, "rate_limit_error"); + assert_eq!(mapped.code, "plan_usage_limit_exceeded"); + assert_eq!( + responses_websocket_plan_admission_close(&error), + (1013, "plan_usage_limit_exceeded") + ); + } + + #[test] + fn saturated_plan_concurrency_uses_the_same_limit_error() { + let error = PlanUsageAdmissionError::Runtime( + aether_runtime_state::RuntimeSemaphoreError::Saturated { + gate: "plan_usage_concurrency", + limit: 2, + }, + ); + + let mapped = map_plan_admission_error(&error); + assert_eq!(mapped.status, 429); + assert_eq!(mapped.code, "plan_usage_limit_exceeded"); + } + + #[test] + fn lost_plan_permit_is_a_retryable_client_visible_503() { + let mapped = plan_permit_loss_error(); + + assert_eq!(mapped.status, 503); + assert_eq!(mapped.error_type, "server_error"); + assert_eq!(mapped.code, "plan_usage_concurrency_unavailable"); + assert_eq!(mapped.close_code, 1013); + assert_eq!(mapped.close_reason, "plan_usage_concurrency_unavailable"); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs index 116058015..fbb475db8 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs @@ -18,7 +18,9 @@ use super::request::{ }; use super::state::BoundResponsesConnection; use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome}; -use super::upstream::{bind_responses_upstream, close_bound_upstream}; +use super::upstream::{ + bind_responses_upstream, close_bound_upstream, ResponsesWebSocketUpstreamBindError, +}; use crate::clock::current_unix_secs; use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT; @@ -103,7 +105,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( state: &AppState, context: &WebSocketRequestContext, _previous_settled: PreviousAttemptSettled, -) -> bool { +) -> Result { // `LogicalTurn::client_event` is intentionally redacted before it is // retained for replay. The binding, however, keeps the hash of the raw // client-side Responses Lite tools/instructions so a later continuation @@ -113,7 +115,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( // redacted replay event. let responses_lite_static_config = bound.responses_lite_static_config.clone(); let Some(active) = bound.turn_state.logical_mut() else { - return false; + return Ok(false); }; if let Some(reason) = active.quota_retry_block_reason() { debug!( @@ -128,7 +130,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( reason, "gateway will not transparently replay an unsafe Responses WebSocket turn" ); - return false; + return Ok(false); } active.retry_attempted = true; active.turn_attempt = active.turn_attempt.saturating_add(1); @@ -142,12 +144,14 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( trace_id = %context.trace_id, "gateway refused to retry a WebSocket turn without its live authorization snapshot" ); - return false; + return Ok(false); }; let turn_index = active.turn_index; let logical_turn_id = active.logical_turn_id.clone(); let codex_fingerprint_context = active.codex_fingerprint_context.clone(); let turn_attempt = active.turn_attempt; + let plan_usage_permit = active.plan_usage_permit.clone(); + let plan_usage_policy_snapshot = active.plan_usage_policy_snapshot.clone(); let retry_exclusion_until_unix_secs = bound .pending_adapter_drain @@ -193,7 +197,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( exhausted_key_id = ?exhausted_key_id, "gateway could not find an alternate Responses WebSocket provider after quota exhaustion" ); - return false; + return Ok(false); } Err(error) => { warn!( @@ -206,7 +210,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( error = ?error, "gateway could not plan an alternate Responses WebSocket provider after quota exhaustion" ); - return false; + return Ok(false); } }; let OwnedResponsesWebSocketDecision { @@ -228,7 +232,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( key_id = ?decision.key_id, "gateway rejected an alternate Responses WebSocket plan that reused the exhausted key" ); - return false; + return Ok(false); } let provider_event = match planned_response_create_event( &decision, @@ -250,7 +254,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( error_code = code, "gateway could not rebuild a Responses response.create for transparent quota retry" ); - return false; + return Ok(false); } }; let replacement_provider_store = provider_event.get("store") == Some(&Value::Bool(true)); @@ -272,6 +276,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( &turn_control.decision, turn_decision, &client_event, + plan_usage_policy_snapshot, planned_lease, ) .await @@ -287,7 +292,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( error = ?error, "gateway could not start usage and audit tracking for transparent quota retry" ); - return false; + return Ok(false); } }; let mut replacement = match bind_responses_upstream( @@ -295,11 +300,41 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( normalization, &client_event, adapter, + plan_usage_permit.as_ref(), + |state| turn.record_upstream_request_state(state), ) .await { Ok(connection) => connection, - Err(code) => { + Err(ResponsesWebSocketUpstreamBindError::PlanUsagePermitLost) => { + turn.release_plan_usage_cost_before_upstream_send( + state, + "transparent_retry_plan_usage_permit_lost", + ) + .await; + queue_turn_finalization( + bound, + state, + turn, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ) + .await; + warn!( + event_name = "responses_websocket_quota_retry_plan_usage_concurrency_lost", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a transparent Responses WebSocket retry before its upstream send after the subscription plan concurrency lease became unhealthy" + ); + return Err(()); + } + Err(ResponsesWebSocketUpstreamBindError::Transport(code)) => { + turn.release_plan_usage_cost_before_upstream_send( + state, + "transparent_retry_rebind_failed", + ) + .await; queue_turn_finalization( bound, state, @@ -316,11 +351,10 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( error_code = code, "gateway could not bind an alternate Responses WebSocket provider after quota exhaustion" ); - return false; + return Ok(false); } }; - turn.mark_upstream_request_sent(); turn.set_provider_response_headers(replacement.upstream_response_headers.clone()); let replacement_upstream = replacement .upstream @@ -348,8 +382,15 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( // drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了 // pending usage 行、占着 candidate 和 pool key lease 的 attempt。 if let Err(orphan) = bound.turn_state.resume(turn) { + let mut orphan = orphan; + orphan + .release_plan_usage_cost_before_upstream_send( + state, + "transparent_retry_state_handoff_failed", + ) + .await; drop(orphan); - return false; + return Ok(false); } if let Some(logical) = bound.turn_state.logical_mut() { logical.provider_store = replacement_provider_store; @@ -369,7 +410,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( key_id = ?bound.decision_template.key_id, "gateway transparently rebound a Responses WebSocket turn after quota exhaustion" ); - true + Ok(true) } fn responses_lite_static_config_after_rebind( diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs index 6312c81a0..118b89fbc 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs @@ -15,7 +15,7 @@ //! (`privacy::restore_sync_response_body` / `privacy::StreamingResponseRestorer`), //! WS 少了这一步,客户端就会直接看到 ``。 //! [`ResponsesWebSocketRedactionRestorer`] 补上这一跳,语义与 HTTP 完全一致: -//! 复用 `privacy::restore_json_strings`,只还原本连接自己 mask 出来的映射, +//! 复用 `privacy::restore_json_strings_with_budget`,只还原本连接自己 mask 出来的映射, //! 未映射的占位符原样透传。 //! //! ## session 为什么活在连接上而不是活在这一轮里 @@ -48,7 +48,11 @@ use crate::ai_serving::{ resolve_local_decision_execution_runtime_auth_context, resolve_provider_chat_pii_redaction, }; use crate::control::GatewayControlDecision; -use crate::privacy::{restore_json_strings, RedactionSession, RedactionSessionSlot}; +use crate::privacy::{ + restore_json_strings_with_budget, serialize_json_value_with_limit, RedactionSession, + RedactionSessionSlot, RestoreExpansionBudget, MAX_STREAM_RESTORE_EXPANSION_BYTES, + MAX_STREAM_RESTORE_OUTPUT_BYTES, +}; use crate::{AppState, GatewayError}; /// Responses WebSocket 只承载 `openai:responses`,脱敏规则按这个客户端格式选取。 @@ -191,29 +195,48 @@ impl ResponsesWebSocketRedactionRestorer { /// 把一帧 provider 事件里的占位符换回真实值,返回要发给客户端的帧文本。 /// - /// `None` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节:未启用 - /// 脱敏(没有任何 session)时连 clone 都不做。 + /// `Ok(None)` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节: + /// 未启用脱敏(没有任何 session)时连 clone 都不做。`Err` 表示恢复或有界 + /// 序列化失败,调用方必须停止转发,不能把它降级成“未命中”而泄漏占位符。 /// /// 入参只读:审计与终态观测继续消费脱敏态的事件,还原只作用于发往客户端的 /// 那一份拷贝,和 HTTP 侧「审计存脱敏体、线上还原」保持一致。 - pub(super) fn restore_provider_frame_text(&self, event: &Value) -> Option { + pub(super) fn restore_provider_frame_text( + &self, + event: &Value, + ) -> Result, GatewayError> { if self.sessions.is_empty() { - return None; + return Ok(None); } let mut restored_event = event.clone(); let mut restored = false; + // Fresh per provider frame: this bounds one allocation amplification + // without imposing a cumulative byte or duration cap on the socket. + let mut budget = RestoreExpansionBudget::new(MAX_STREAM_RESTORE_EXPANSION_BYTES); for session in &self.sessions { // 逐 session 还原而不是合并映射:每个 session 只认自己 mask 过的 // sentinel(`RedactionSession::restore_text`),跨 session 合并会绕开 // 这条边界。同一个值在不同轮派生出的 sentinel 相同,所以顺序无关。 - restored |= restore_json_strings(&mut restored_event, session); + restored |= + restore_json_strings_with_budget(&mut restored_event, session, &mut budget)?; } if !restored { - return None; + return Ok(None); } - // 刚从 JSON 解析出来的 Value 再序列化不会失败;真失败时宁可让客户端看到 - // 占位符,也不能丢掉这一帧——丢帧会让客户端的协议状态机卡死。 - serde_json::to_string(&restored_event).ok() + // The parsed input frame is bounded at ingress, but use the same bounded + // serializer here so calculating the restored output limit cannot create + // an unchecked temporary allocation. + let original_len = + serialize_json_value_with_limit(event, MAX_STREAM_RESTORE_OUTPUT_BYTES)?.len(); + let output_limit = original_len + .checked_add(MAX_STREAM_RESTORE_EXPANSION_BYTES) + .ok_or_else(|| { + GatewayError::Internal( + "WebSocket redaction restored frame length overflow".to_string(), + ) + })? + .min(MAX_STREAM_RESTORE_OUTPUT_BYTES); + serialize_json_value_with_limit(&restored_event, output_limit).map(Some) } } @@ -334,6 +357,7 @@ mod tests { local_rejection: None, allowed_models: None, ip_rules: None, + verified_api_key_hash: None, }); decision } @@ -691,6 +715,7 @@ mod tests { let frame = provider_delta_frame(&format!("your mail is {sentinel}")); let restored = restorer .restore_provider_frame_text(&frame) + .expect("frame restoration must not fail") .expect("a frame echoing this turn's sentinel must be restored"); assert!( @@ -725,6 +750,7 @@ mod tests { }); let restored = restorer .restore_provider_frame_text(&frame) + .expect("frame restoration must not fail") .expect("a batched sentinel must be restored"); assert!(restored.contains(TEST_EMAIL), "{restored}"); @@ -744,6 +770,7 @@ mod tests { let frame = provider_delta_frame(&format!("{FOREIGN_SENTINEL} and {sentinel}")); let restored = restorer .restore_provider_frame_text(&frame) + .expect("frame restoration must not fail") .expect("the mapped sentinel is still restored"); assert!(restored.contains(TEST_EMAIL), "{restored}"); @@ -764,17 +791,36 @@ mod tests { assert!( restorer .restore_provider_frame_text(&provider_delta_frame("nothing to restore")) + .expect("frame restoration must not fail") .is_none(), "a frame with no mapped sentinel must be relayed byte-for-byte" ); assert!( restorer .restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL)) + .expect("frame restoration must not fail") .is_none(), "a frame that only carries unmapped placeholders must not be rewritten" ); } + #[tokio::test] + async fn frame_restore_budget_does_not_accumulate_across_frames() { + let state = redaction_enabled_state(); + let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await; + let sentinel = sentinel_for(&redaction, TEST_EMAIL); + let mut restorer = ResponsesWebSocketRedactionRestorer::default(); + restorer.register(redaction.session); + + for _ in 0..3 { + let restored = restorer + .restore_provider_frame_text(&provider_delta_frame(&sentinel)) + .expect("frame restoration must not fail") + .expect("each provider frame gets an independent expansion budget"); + assert!(restored.contains(TEST_EMAIL)); + } + } + /// 未启用脱敏(或这条连接从没 mask 到东西)时,还原器必须完全不介入: /// 连 clone 都不做,输出就是上游原字节。 #[tokio::test] @@ -783,9 +829,11 @@ mod tests { assert!(restorer .restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL)) + .expect("frame restoration must not fail") .is_none()); assert!(restorer .restore_provider_frame_text(&provider_delta_frame(TEST_EMAIL)) + .expect("frame restoration must not fail") .is_none()); } @@ -806,6 +854,7 @@ mod tests { assert!(restorer .restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL)) + .expect("frame restoration must not fail") .is_none()); } @@ -822,6 +871,7 @@ mod tests { let before = frame.clone(); let _ = restorer .restore_provider_frame_text(&frame) + .expect("frame restoration must not fail") .expect("the frame is restored for the client"); assert_eq!( @@ -849,6 +899,7 @@ mod tests { let frame = provider_delta_frame(&format!("{first_sentinel} then {second_sentinel}")); let restored = restorer .restore_provider_frame_text(&frame) + .expect("frame restoration must not fail") .expect("both turns' sentinels are restorable on this connection"); assert!(restored.contains(TEST_EMAIL), "{restored}"); @@ -876,11 +927,13 @@ mod tests { assert!( restorer .restore_provider_frame_text(&provider_delta_frame(&prior_sentinel)) + .expect("prior-chain restoration check must not fail") .is_none(), "an independent chain must not restore PII from its predecessor" ); let restored = restorer .restore_provider_frame_text(&provider_delta_frame(¤t_sentinel)) + .expect("the new chain's first-turn restoration must not fail") .expect("the new chain's first-turn mapping must remain available"); assert!(restored.contains(OTHER_TEST_EMAIL), "{restored}"); assert!(!restored.contains(¤t_sentinel), "{restored}"); @@ -901,6 +954,7 @@ mod tests { assert!(!restorer.has_sessions()); assert!(restorer .restore_provider_frame_text(&provider_delta_frame(&prior_sentinel)) + .expect("cleared-chain restoration check must not fail") .is_none()); } @@ -927,12 +981,14 @@ mod tests { assert!( restorer .restore_provider_frame_text(&provider_delta_frame(&oldest_sentinel)) + .expect("frame restoration must not fail") .is_none(), "the evicted turn's sentinel is relayed verbatim, never mis-restored" ); assert!( restorer .restore_provider_frame_text(&provider_delta_frame(&newest_sentinel)) + .expect("frame restoration must not fail") .is_some(), "the most recent turns stay restorable" ); diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs index d8aa948f1..4393e1e99 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -30,6 +30,11 @@ use super::ownership::{ await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease, spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision, }; +use super::plan_admission::{ + acquire_responses_websocket_plan_admission, responses_websocket_plan_admission_close, + send_responses_websocket_plan_admission_error, + terminate_responses_websocket_for_plan_permit_loss, +}; use super::redaction::redact_responses_websocket_client_event_with_reasoning_replay_policy; use super::relay_policy::{fatal_relay_policy, FatalRelaySignal}; use super::request::{ @@ -41,7 +46,9 @@ use super::request::{ use super::state::BoundResponsesConnection; use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome}; use super::turn_state::LogicalTurn; -use super::upstream::{bind_responses_upstream, close_bound_upstream}; +use super::upstream::{ + bind_responses_upstream, close_bound_upstream, ResponsesWebSocketUpstreamBindError, +}; use crate::ai_serving::ResponsesWebSocketPinnedCandidate; use crate::handlers::proxy::websocket::ingress::{ @@ -522,6 +529,35 @@ async fn bootstrap_responses_websocket( } } + let plan_usage_admission = match acquire_responses_websocket_plan_admission( + &state, + &turn_control.decision, + &first_logical_turn_id, + ) + .await + { + Ok(admission) => admission, + Err(error) => { + warn!( + event_name = "responses_websocket_initial_plan_usage_rejected", + log_type = "event", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + logical_turn_id = %first_logical_turn_id, + error = ?error, + "gateway rejected the initial Responses WebSocket turn at its subscription plan limit" + ); + send_responses_websocket_plan_admission_error(client_socket, &error).await; + let (close_code, close_reason) = responses_websocket_plan_admission_close(&error); + close_client_socket(client_socket, close_code, close_reason).await; + return None; + } + }; + let crate::plan_usage_policy::PlanUsageAdmission { + permit: plan_usage_permit, + policy_snapshot: plan_usage_policy_snapshot, + } = plan_usage_admission; // A cross-socket response chain must be owned by this exact live // authenticated principal. Missing, expired, corrupt or unavailable state // fails closed; allowing the normal scheduler to choose a provider/key @@ -885,6 +921,7 @@ async fn bootstrap_responses_websocket( &turn_control.decision, first_turn_decision, &first_event, + plan_usage_policy_snapshot.clone(), planned_lease, ) .await @@ -907,36 +944,68 @@ async fn bootstrap_responses_websocket( } }; - let mut bound = - match bind_responses_upstream(&decision, normalization, &first_event, adapter).await { - Ok(connection) => connection, - Err(code) => { - let finalizer = finalize_unbound_turn( - state.clone(), - first_turn, - ResponsesWebSocketTurnOutcome::upstream_connect_failed(code), - ); - warn!( - event_name = "responses_websocket_upstream_connect_failed", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %context.trace_id, - error_code = code, - "gateway failed to establish Responses WebSocket upstream" - ); - send_gateway_error_with_status( - client_socket, - 502, - code, - "Gateway could not establish the Provider connection", + let mut bound = match bind_responses_upstream( + &decision, + normalization, + &first_event, + adapter, + plan_usage_permit.as_ref(), + |state| first_turn.record_upstream_request_state(state), + ) + .await + { + Ok(connection) => connection, + Err(ResponsesWebSocketUpstreamBindError::PlanUsagePermitLost) => { + first_turn + .release_plan_usage_cost_before_upstream_send( + &state, + "initial_plan_usage_permit_lost", ) .await; - close_client_socket(client_socket, CLOSE_TRY_AGAIN, code).await; - await_turn_finalization_handle(finalizer).await; - return None; - } - }; + let finalizer = finalize_unbound_turn( + state.clone(), + first_turn, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ); + warn!( + event_name = "responses_websocket_initial_plan_usage_concurrency_lost", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped the initial Responses WebSocket turn before its upstream send after the subscription plan concurrency lease became unhealthy" + ); + terminate_responses_websocket_for_plan_permit_loss(client_socket).await; + await_turn_finalization_handle(finalizer).await; + return None; + } + Err(ResponsesWebSocketUpstreamBindError::Transport(code)) => { + let finalizer = finalize_unbound_turn( + state.clone(), + first_turn, + ResponsesWebSocketTurnOutcome::upstream_connect_failed(code), + ); + warn!( + event_name = "responses_websocket_upstream_connect_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error_code = code, + "gateway failed to establish Responses WebSocket upstream" + ); + send_gateway_error_with_status( + client_socket, + 502, + code, + "Gateway could not establish the Provider connection", + ) + .await; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, code).await; + await_turn_finalization_handle(finalizer).await; + return None; + } + }; if bound.responses_lite_static_config.is_some() { bound.responses_lite_static_config = continuation_record .as_ref() @@ -953,7 +1022,6 @@ async fn bootstrap_responses_websocket( .continuation_response_ids .remember_persisted(previous_response_id); } - first_turn.mark_upstream_request_sent(); first_turn.set_provider_response_headers(bound.upstream_response_headers.clone()); if let Some(session) = first_turn_redaction_session { register_initial_redaction_session(&mut bound, session); @@ -962,7 +1030,9 @@ async fn bootstrap_responses_websocket( LogicalTurn::new(first_event, 1, first_logical_turn_id) .with_codex_fingerprint_context(first_codex_fingerprint_context) .with_provider_store(first_provider_event.get("store") == Some(&Value::Bool(true))) - .with_turn_control(turn_control), + .with_turn_control(turn_control) + .with_plan_usage_permit(plan_usage_permit) + .with_plan_usage_policy_snapshot(plan_usage_policy_snapshot), first_turn, ); @@ -2067,6 +2137,8 @@ mod tests { resolve_responses_websocket_adapter( crate::orchestration::ResponsesWebSocketAdapter::Standard, ), + None, + |_| {}, ) .await .expect("upstream binding should succeed"); @@ -2278,6 +2350,7 @@ mod tests { let restored = ordinary_bound .redaction_restorer .restore_provider_frame_text(&provider_event) + .expect("ordinary OpenAI replay policy restoration must not fail") .expect("ordinary OpenAI replay policy should restore response text"); let restored: serde_json::Value = serde_json::from_str(&restored).expect("restored provider event should remain JSON"); @@ -2297,6 +2370,7 @@ mod tests { deepseek_bound .redaction_restorer .restore_provider_frame_text(&provider_event) + .expect("DeepSeek opaque restoration check must not fail") .is_none(), "the authenticated DeepSeek binding must keep opaque reasoning state byte-identical" ); diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs index fab1cac80..db2551efe 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs @@ -283,4 +283,28 @@ mod tests { // 投递失败不是供应商的错误,摘要不该因此补 parser_error。 assert_eq!(facts.forced_error(), None); } + + #[test] + fn plan_permit_loss_is_gateway_cancellation_not_provider_failure() { + let facts = attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::connection_admission_lost(), + ); + + assert_eq!( + facts.provider, + aborted( + 499, + "gateway WebSocket connection admission became unhealthy" + ) + ); + assert_eq!( + facts.delivery, + AttemptClientDelivery::Aborted { + reason: "gateway WebSocket connection admission became unhealthy" + } + ); + assert_eq!(facts.forced_error(), None); + } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs index 79f9bc2f3..4fa9d7b07 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs @@ -17,7 +17,8 @@ use aether_contracts::{ }; use aether_data_contracts::repository::candidates::RequestCandidateStatus; use aether_data_contracts::repository::usage::{ - UsageBodyCaptureState, WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, + UsageBodyCaptureState, PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, + WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, }; use aether_scheduler_core::SchedulerRequestCandidateStatusUpdate; use aether_usage_runtime::{ @@ -51,6 +52,7 @@ use crate::orchestration::{ release_local_pool_key_lease, release_pool_key_lease_from_report_context, LocalExecutionEffectContext, LocalStreamFailureEffect, }; +use crate::plan_usage_policy::PlanUsagePolicySnapshot; use crate::request_candidate_runtime::{ ensure_execution_request_candidate_slot, record_local_request_candidate_status, }; @@ -146,6 +148,10 @@ impl ResponsesWebSocketTurnOutcome { } pub(super) const fn connection_admission_lost() -> Self { + // This gateway-owned cancellation is also used when a logical turn's + // subscription-plan permit becomes unhealthy. Keeping it Cancelled + // prevents the settlement layer from projecting a provider failure; + // the client-visible 503 is emitted separately by plan_admission. Self::Cancelled { reason: "gateway WebSocket connection admission became unhealthy", } @@ -250,16 +256,226 @@ pub(super) struct ResponsesProviderAttempt { first_event_timeout: Duration, terminal_timeout: Duration, admission: Option, + /// A provider attempt gets its own server-issued cost reservation. The + /// trusted token also lives in the lifecycle report context so the normal + /// terminal usage path can reconcile actual cost. We retain the control + /// snapshot here only for the quota-retry path, which must release the old + /// attempt synchronously before reserving the replacement. + plan_usage_cost_reservation: Option, terminal_error_body: Option, /// 观察到的 provider 终态事实,与「为什么现在结算」这个信号分开保存。 /// 客户端投递失败不会把它擦掉。 provider_outcome: Option, /// 这一个 attempt 的内容是否完整交付给了客户端。与 provider 终态正交。 client_delivery: AttemptClientDelivery, - /// True only after `response.create` has been accepted by the upstream - /// socket writer. Cancellation before this point must not be projected as - /// provider failure or billed usage. - upstream_request_sent: bool, + /// Tracks the irreversible transport handoff separately from flush + /// confirmation. Once start_send succeeds, a later cancellation cannot + /// prove the provider did not receive or execute response.create. + upstream_request_state: UpstreamRequestState, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum UpstreamRequestState { + NotStarted, + PossiblySent, + Sent, +} + +impl UpstreamRequestState { + const fn may_have_reached_provider(self) -> bool { + !matches!(self, Self::NotStarted) + } +} + +struct ResponsesPlanUsageCostReservation { + control_decision: GatewayControlDecision, + reservation_token: String, +} + +struct OwnedResponsesPlanUsageCostReservation { + control_decision: GatewayControlDecision, + plan: ExecutionPlan, + reservation_token: String, +} + +/// Owns a successful durable reservation until an attempt lifecycle owns the +/// trusted report context. If startup is cancelled anywhere in between, Drop +/// releases the reservation instead of leaving a false 429 until its TTL. +struct ResponsesPlanUsageCostReservationGuard { + state: AppState, + reservation: Option, +} + +impl ResponsesPlanUsageCostReservationGuard { + fn new( + state: &AppState, + control_decision: GatewayControlDecision, + plan: ExecutionPlan, + reservation_token: String, + ) -> Self { + Self { + state: state.clone(), + reservation: Some(OwnedResponsesPlanUsageCostReservation { + control_decision, + plan, + reservation_token, + }), + } + } + + fn reservation_token(&self) -> &str { + self.reservation + .as_ref() + .expect("an armed plan cost reservation guard owns its reservation") + .reservation_token + .as_str() + } + + fn disarm(mut self) -> ResponsesPlanUsageCostReservation { + let reservation = self + .reservation + .take() + .expect("an armed plan cost reservation guard owns its reservation"); + ResponsesPlanUsageCostReservation { + control_decision: reservation.control_decision, + reservation_token: reservation.reservation_token, + } + } +} + +impl Drop for ResponsesPlanUsageCostReservationGuard { + fn drop(&mut self) { + let Some(reservation) = self.reservation.take() else { + return; + }; + let state = self.state.clone(); + if let Ok(handle) = tokio::runtime::Handle::try_current() { + handle.spawn(async move { + release_responses_plan_usage_cost_best_effort( + &state, + &reservation.control_decision, + &reservation.plan, + reservation.reservation_token.as_str(), + "turn_start_cancelled", + ) + .await; + }); + } + } +} + +enum ResponsesPlanUsageCostReservationStart { + NotRequired, + Reserved(ResponsesPlanUsageCostReservationGuard), + Rejected(crate::plan_usage_policy::PlanUsagePolicyRejection), +} + +async fn reserve_responses_plan_usage_cost_owned( + state: &AppState, + control_decision: &GatewayControlDecision, + plan: &ExecutionPlan, + report_context: Option<&Value>, + plan_usage_policy_snapshot: Option<&PlanUsagePolicySnapshot>, + reservation_token: String, +) -> Result { + let outcome = crate::plan_usage_policy::reserve_admitted_plan_usage_policy_cost( + state, + control_decision, + plan, + report_context, + plan_usage_policy_snapshot, + reservation_token.as_str(), + ) + .await?; + match outcome { + crate::plan_usage_policy::PlanUsageCostReservationOutcome::NotRequired => { + Ok(ResponsesPlanUsageCostReservationStart::NotRequired) + } + crate::plan_usage_policy::PlanUsageCostReservationOutcome::Reserved => { + Ok(ResponsesPlanUsageCostReservationStart::Reserved( + ResponsesPlanUsageCostReservationGuard::new( + state, + control_decision.clone(), + plan.clone(), + reservation_token, + ), + )) + } + crate::plan_usage_policy::PlanUsageCostReservationOutcome::Rejected(rejection) => { + Ok(ResponsesPlanUsageCostReservationStart::Rejected(rejection)) + } + } +} + +fn attach_plan_usage_reservation_token( + report_context: Option, + reservation_token: &str, +) -> Option { + let mut object = match report_context { + Some(Value::Object(object)) => object, + Some(other) => Map::from_iter([("seed".to_string(), other)]), + None => Map::new(), + }; + // Always overwrite a client-derived/seed value with the server-issued + // token. Terminal reconciliation must never trust caller-controlled JSON. + object.insert( + "plan_usage_reservation_token".to_string(), + Value::String(reservation_token.to_string()), + ); + // The token and its reconciliation controls are server-owned. A seed or + // client event must never be able to defer reconciliation on its own. + object.insert( + PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY.to_string(), + Value::Bool(false), + ); + Some(Value::Object(object)) +} + +fn websocket_plan_usage_rejection_error( + rejection: &crate::plan_usage_policy::PlanUsagePolicyRejection, +) -> GatewayError { + GatewayError::PlanUsageLimited(rejection.clone()) +} + +async fn release_responses_plan_usage_cost( + state: &AppState, + control_decision: &GatewayControlDecision, + plan: &ExecutionPlan, + reservation_token: &str, + _reason: &'static str, +) -> Result<(), GatewayError> { + crate::plan_usage_policy::release_plan_usage_policy_cost( + state, + control_decision, + plan, + reservation_token, + current_unix_ms() / 1_000, + ) + .await +} + +async fn release_responses_plan_usage_cost_best_effort( + state: &AppState, + control_decision: &GatewayControlDecision, + plan: &ExecutionPlan, + reservation_token: &str, + reason: &'static str, +) { + if let Err(error) = + release_responses_plan_usage_cost(state, control_decision, plan, reservation_token, reason) + .await + { + warn!( + event_name = "responses_websocket_plan_usage_cost_release_failed", + log_type = "ops", + request_id = %plan.request_id, + candidate_id = ?plan.candidate_id, + reservation_token, + reason, + error = ?error, + "gateway failed to release a Responses WebSocket plan cost reservation" + ); + } } /// 组装一轮 turn 的 decision。 @@ -307,6 +523,7 @@ pub(super) async fn begin_unowned_responses_websocket_turn( control_decision: &GatewayControlDecision, decision: AiExecutionDecision, client_event: &Value, + plan_usage_policy_snapshot: Option, ) -> Result { let planned_report_context = decision.report_context.clone(); let attempt = match build_openai_responses_stream_plan_from_decision( @@ -415,6 +632,68 @@ pub(super) async fn begin_unowned_responses_websocket_turn( } }; + let reservation_token = uuid::Uuid::new_v4().to_string(); + let cost_reservation = match reserve_responses_plan_usage_cost_owned( + state, + control_decision, + &plan, + report_context.as_ref(), + plan_usage_policy_snapshot.as_ref(), + reservation_token, + ) + .await + { + Ok(ResponsesPlanUsageCostReservationStart::NotRequired) => None, + Ok(ResponsesPlanUsageCostReservationStart::Reserved(guard)) => { + report_context = + attach_plan_usage_reservation_token(report_context, guard.reservation_token()); + Some(guard) + } + Ok(ResponsesPlanUsageCostReservationStart::Rejected(rejection)) => { + let error = websocket_plan_usage_rejection_error(&rejection); + admission.release().await; + release_then_record_responses_websocket_admission_failure( + release_local_pool_key_lease( + state, + LocalExecutionEffectContext { + plan: &plan, + report_context: report_context.as_ref(), + }, + ), + record_responses_websocket_admission_failure( + state, + &plan, + report_context.as_ref(), + candidate_started_at_unix_ms, + &error, + ), + ) + .await; + return Err(error); + } + Err(error) => { + admission.release().await; + release_then_record_responses_websocket_admission_failure( + release_local_pool_key_lease( + state, + LocalExecutionEffectContext { + plan: &plan, + report_context: report_context.as_ref(), + }, + ), + record_responses_websocket_admission_failure( + state, + &plan, + report_context.as_ref(), + candidate_started_at_unix_ms, + &error, + ), + ) + .await; + return Err(error); + } + }; + let lifecycle = ExecutionAttemptLifecycle::begin( state, AttemptLifecycleSeed { @@ -427,6 +706,7 @@ pub(super) async fn begin_unowned_responses_websocket_turn( }, ) .await; + let plan_usage_cost_reservation = cost_reservation.map(|guard| guard.disarm()); Ok(ResponsesProviderAttempt { lifecycle, @@ -442,10 +722,11 @@ pub(super) async fn begin_unowned_responses_websocket_turn( first_event_timeout, terminal_timeout, admission: Some(admission), + plan_usage_cost_reservation, terminal_error_body: None, provider_outcome: None, client_delivery: AttemptClientDelivery::Complete, - upstream_request_sent: false, + upstream_request_state: UpstreamRequestState::NotStarted, }) } @@ -496,6 +777,11 @@ fn responses_websocket_admission_failure_update( "gateway_admission_timeout", format!("gateway admission gate {gate} timed out after {queue_budget_ms}ms"), ), + GatewayError::PlanUsageLimited(_) => ( + StatusCode::TOO_MANY_REQUESTS.as_u16(), + "plan_usage_limit_exceeded", + "subscription plan usage limit exceeded".to_string(), + ), other => ( StatusCode::INTERNAL_SERVER_ERROR.as_u16(), "gateway_admission_failed", @@ -577,6 +863,64 @@ impl ResponsesProviderAttempt { } } + /// Releases this attempt's reserved cost before a transparent retry. The + /// terminal usage write that follows may reconcile the same token again; + /// repository reconciliation is idempotent, so the first Released state + /// wins and the replacement attempt can reserve independently. + pub(super) async fn release_plan_usage_cost_for_retry( + &mut self, + state: &AppState, + ) -> Result<(), GatewayError> { + if self.upstream_request_state.may_have_reached_provider() + && self.provider_outcome.is_none() + { + return Err(GatewayError::Internal( + "cannot release an unresolved plan cost reservation after the upstream request may have been sent" + .to_string(), + )); + } + let Some(reservation) = self.plan_usage_cost_reservation.take() else { + return Ok(()); + }; + let release = release_responses_plan_usage_cost( + state, + &reservation.control_decision, + self.lifecycle.plan(), + reservation.reservation_token.as_str(), + "transparent_retry", + ) + .await; + if release.is_err() { + self.plan_usage_cost_reservation = Some(reservation); + } + release + } + + /// A transparent replacement failed before its request reached an + /// upstream. No terminal usage event can safely own the still-reserved + /// estimate, so release it explicitly before the orphan attempt is + /// finalized or dropped. + pub(super) async fn release_plan_usage_cost_before_upstream_send( + &mut self, + state: &AppState, + reason: &'static str, + ) { + if self.upstream_request_state.may_have_reached_provider() { + return; + } + let Some(reservation) = self.plan_usage_cost_reservation.take() else { + return; + }; + release_responses_plan_usage_cost_best_effort( + state, + &reservation.control_decision, + self.lifecycle.plan(), + reservation.reservation_token.as_str(), + reason, + ) + .await; + } + pub(super) fn set_provider_response_headers(&mut self, headers: BTreeMap) { let observed_at_unix_ms = current_unix_ms(); let report_context = attach_provider_response_headers_to_report_context( @@ -592,12 +936,33 @@ impl ResponsesProviderAttempt { /// Starts the per-turn response deadlines only after the corresponding /// `response.create` has been accepted by the upstream socket writer. - pub(super) fn mark_upstream_request_sent(&mut self) { - self.upstream_request_sent = true; - self.started_at = Instant::now(); - self.provider_request_started_at_unix_ms = current_unix_ms(); - self.provider_request_order_id = uuid::Uuid::now_v7().to_string(); - self.first_event_elapsed_ms = None; + pub(super) fn record_upstream_request_state(&mut self, state: UpstreamRequestState) { + match state { + UpstreamRequestState::NotStarted => { + debug_assert!( + false, + "upstream request state cannot move back to not started" + ); + } + UpstreamRequestState::PossiblySent => { + debug_assert_eq!( + self.upstream_request_state, + UpstreamRequestState::NotStarted + ); + self.upstream_request_state = state; + self.started_at = Instant::now(); + self.provider_request_started_at_unix_ms = current_unix_ms(); + self.provider_request_order_id = uuid::Uuid::now_v7().to_string(); + self.first_event_elapsed_ms = None; + } + UpstreamRequestState::Sent => { + debug_assert_eq!( + self.upstream_request_state, + UpstreamRequestState::PossiblySent + ); + self.upstream_request_state = state; + } + } } /// Selects a cancellation-safe fallback for an attempt whose owner task @@ -605,7 +970,9 @@ impl ResponsesProviderAttempt { /// after the write it remains a gateway relay failure because provider /// work may already have started. pub(super) const fn abandonment_outcome(&self) -> ResponsesWebSocketTurnOutcome { - ResponsesWebSocketTurnOutcome::relay_task_abandonment(self.upstream_request_sent) + ResponsesWebSocketTurnOutcome::relay_task_abandonment( + self.upstream_request_state.may_have_reached_provider(), + ) } pub(super) fn deadline(&self) -> ResponsesWebSocketTurnDeadline { @@ -734,6 +1101,14 @@ impl ResponsesProviderAttempt { /// 这里只提供 WS 观察到的终态事实。 async fn settle(mut self, state: &AppState, outcome: ResponsesWebSocketTurnOutcome) { let facts = attempt_facts_for_outcome(self.provider_outcome, self.client_delivery, outcome); + if self.upstream_request_state.may_have_reached_provider() + && !facts.provider.is_terminal() + && self.plan_usage_cost_reservation.is_some() + { + let report_context = + attach_plan_usage_reservation_deferred(self.lifecycle.take_report_context()); + self.lifecycle.set_report_context(report_context); + } if let Some(reason) = facts.delivery.aborted_reason() { let report_context = attach_client_delivery_to_report_context( self.lifecycle.take_report_context(), @@ -951,6 +1326,19 @@ fn attach_client_delivery_to_report_context( Some(Value::Object(object)) } +fn attach_plan_usage_reservation_deferred(report_context: Option) -> Option { + let mut object = match report_context { + Some(Value::Object(object)) => object, + Some(other) => Map::from_iter([("seed".to_string(), other)]), + None => Map::new(), + }; + object.insert( + PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY.to_string(), + Value::Bool(true), + ); + Some(Value::Object(object)) +} + fn provider_terminal_outcome( frame: &ParsedResponsesWebSocketFrame<'_>, ) -> Option { @@ -1016,7 +1404,8 @@ mod tests { attempt_facts_for_outcome, settle_signal_for_client_delivery_failure, }; use super::{ - attach_client_delivery_to_report_context, prepare_websocket_report_context, + attach_client_delivery_to_report_context, attach_plan_usage_reservation_deferred, + attach_plan_usage_reservation_token, prepare_websocket_report_context, provider_terminal_outcome, release_then_record_responses_websocket_admission_failure, resolve_responses_websocket_turn_timeouts, responses_websocket_admission_failure_update, websocket_event_as_sse_line, ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnOutcome, @@ -1028,6 +1417,36 @@ mod tests { }; use crate::GatewayError; + #[test] + fn server_plan_usage_reservation_token_overrides_seed_value() { + let context = attach_plan_usage_reservation_token( + Some(json!({ + "candidate_index": 2, + "plan_usage_reservation_token": "client-controlled-token" + })), + "server-reservation-token", + ) + .expect("reservation context"); + + assert_eq!(context["candidate_index"], 2); + assert_eq!( + context["plan_usage_reservation_token"], + "server-reservation-token" + ); + assert_eq!(context["plan_usage_reservation_deferred"], false); + } + + #[test] + fn gateway_can_defer_a_server_owned_reservation_after_transport_handoff() { + let context = attach_plan_usage_reservation_deferred(Some(json!({ + "plan_usage_reservation_token": "server-reservation-token", + "plan_usage_reservation_deferred": false + }))) + .expect("reservation context"); + + assert_eq!(context["plan_usage_reservation_deferred"], true); + } + #[tokio::test] async fn admission_failure_releases_pool_lease_before_recording_candidate_terminal() { let phase = Arc::new(AtomicU8::new(0)); diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs index ebcd5f2df..07f08f64f 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs @@ -11,6 +11,7 @@ use serde_json::Value; use super::control::ResponsesWebSocketTurnControl; use super::lifecycle::ActiveProviderAttempt; use super::request::response_create_has_previous_response_id; +use crate::plan_usage_policy::PlanUsagePolicySnapshot; /// 客户端一次 `response.create` 对应的 logical turn。 /// @@ -18,7 +19,7 @@ use super::request::response_create_has_previous_response_id; /// 换一条上游连接重放同一份客户端事件,但对客户端始终是同一轮请求。 /// `client_event` 保存的必须是**已脱敏**的事件(见 `super::redaction`), /// 因为透明重试直接重放它。 -#[derive(Debug, Clone)] +#[derive(Debug)] pub(super) struct LogicalTurn { pub(super) client_event: Value, /// Effective `store` after provider body rules and WebSocket framing. Only @@ -37,6 +38,14 @@ pub(super) struct LogicalTurn { /// this logical turn. Quota retries reuse it instead of falling back to the /// connection's Upgrade-time authorization snapshot. pub(super) turn_control: Option, + /// Subscription-plan concurrency is per logical `response.create`, not per + /// provider attempt. Keeping the permit here makes transparent retries + /// retain the same slot through `Replanning` until the logical terminal. + pub(super) plan_usage_permit: Option, + /// Cost limits are evaluated against the policy and entitlement window + /// that admitted this logical response.create. Every transparent provider + /// attempt reuses this snapshot while receiving its own reservation token. + pub(super) plan_usage_policy_snapshot: Option, } impl LogicalTurn { @@ -51,6 +60,8 @@ impl LogicalTurn { retry_attempted: false, retry_unsafe_reason: None, turn_control: None, + plan_usage_permit: None, + plan_usage_policy_snapshot: None, } } @@ -59,6 +70,22 @@ impl LogicalTurn { self } + pub(super) fn with_plan_usage_permit( + mut self, + permit: Option, + ) -> Self { + self.plan_usage_permit = permit; + self + } + + pub(super) fn with_plan_usage_policy_snapshot( + mut self, + snapshot: Option, + ) -> Self { + self.plan_usage_policy_snapshot = snapshot; + self + } + pub(super) fn with_codex_fingerprint_context( mut self, context: CodexFingerprintConvergenceContext, diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs index 8644588c6..460817ff5 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs @@ -7,18 +7,33 @@ use wreq::ws::message::Message as WreqWsMessage; use super::adapter::ResponsesWebSocketProtocolAdapter; use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError}; +use super::plan_admission::responses_websocket_plan_permit_is_healthy; use super::redaction::ResponsesWebSocketRedactionRestorer; use super::request::{planned_request_uses_codex_responses_lite, planned_response_create_event}; use super::state::{ BoundResponsesConnection, ContinuationResponseIds, ExhaustedResponsesWebSocketExclusions, }; +use super::turn::UpstreamRequestState; use super::turn_state::ResponsesTurnState; use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization}; use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS; use crate::handlers::proxy::websocket::transport::{ - close_upstream_socket, connect_upstream_websocket, send_upstream_message, + close_upstream_socket, connect_upstream_websocket, feed_upstream_message, + flush_upstream_messages, WebSocketWriteError, }; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum ResponsesWebSocketUpstreamSendError { + PlanUsagePermitLost, + Transport(WebSocketWriteError), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum ResponsesWebSocketUpstreamBindError { + PlanUsagePermitLost, + Transport(&'static str), +} + /// 上游 WebSocket 握手的默认绝对 deadline(30 秒)。 /// 覆盖 DNS → TCP connect → TLS → HTTP 101 Upgrade → 发送首条 event 的完整链路。 /// 如果 decision 配置了更短的 first_byte_ms 或 total_ms,取其与此值的较小者。 @@ -39,59 +54,100 @@ pub(super) fn resolve_upstream_handshake_deadline(decision: &AiExecutionDecision Duration::from_millis(deadline_ms) } -pub(super) async fn bind_responses_upstream( +pub(super) async fn bind_responses_upstream( decision: &AiExecutionDecision, normalization: ResponsesWebSocketBodyNormalization, initial_event: &Value, adapter: &'static dyn ResponsesWebSocketProtocolAdapter, -) -> Result { + plan_usage_permit: Option<&aether_runtime::AdmissionPermit>, + record_upstream_request_state: F, +) -> Result +where + F: FnMut(UpstreamRequestState), +{ // 绝对 deadline:从此刻起必须在限定时间内完成握手 + 首条事件发送, // 防止慢 TLS / 慢 HTTP Upgrade 无限占用 connection permit。 let handshake_deadline = resolve_upstream_handshake_deadline(decision); tokio::time::timeout( handshake_deadline, - bind_responses_upstream_inner(decision, normalization, initial_event, adapter), + bind_responses_upstream_inner( + decision, + normalization, + initial_event, + adapter, + plan_usage_permit, + record_upstream_request_state, + ), ) .await - .map_err(|_| "responses_websocket_upstream_handshake_timeout")? + .map_err(|_| { + ResponsesWebSocketUpstreamBindError::Transport( + "responses_websocket_upstream_handshake_timeout", + ) + })? } /// 实际执行握手 + 首条事件发送的内部函数,由外层 timeout 包裹。 -async fn bind_responses_upstream_inner( +async fn bind_responses_upstream_inner( decision: &AiExecutionDecision, normalization: ResponsesWebSocketBodyNormalization, initial_event: &Value, adapter: &'static dyn ResponsesWebSocketProtocolAdapter, -) -> Result { + plan_usage_permit: Option<&aether_runtime::AdmissionPermit>, + record_upstream_request_state: F, +) -> Result +where + F: FnMut(UpstreamRequestState), +{ let binding_identity = - UpstreamBindingIdentity::from_decision(adapter, decision).map_err(|error| match error { - UpstreamBindingIdentityError::MissingUpstreamUrl => { - adapter.upstream_errors().upstream_url_missing - } - UpstreamBindingIdentityError::InvalidUpstreamUrl => { - adapter.upstream_errors().upstream_url_invalid - } - UpstreamBindingIdentityError::InvalidHandshakeHeaders => { - adapter.upstream_errors().headers_invalid - } + UpstreamBindingIdentity::from_decision(adapter, decision).map_err(|error| { + ResponsesWebSocketUpstreamBindError::Transport(match error { + UpstreamBindingIdentityError::MissingUpstreamUrl => { + adapter.upstream_errors().upstream_url_missing + } + UpstreamBindingIdentityError::InvalidUpstreamUrl => { + adapter.upstream_errors().upstream_url_invalid + } + UpstreamBindingIdentityError::InvalidHandshakeHeaders => { + adapter.upstream_errors().headers_invalid + } + }) })?; let mut upstream = connect_upstream_websocket( decision, RESPONSES_WEBSOCKET_SESSION_LIMITS, adapter.upstream_errors(), ) - .await?; - let first_event = planned_response_create_event(decision, &normalization, initial_event)?; - send_upstream_message(&mut upstream.socket, WreqWsMessage::text(first_event)) - .await - .map_err(|_| "responses_websocket_initial_send_failed")?; + .await + .map_err(ResponsesWebSocketUpstreamBindError::Transport)?; + let first_event = planned_response_create_event(decision, &normalization, initial_event) + .map_err(ResponsesWebSocketUpstreamBindError::Transport)?; + send_responses_websocket_upstream_message( + &mut upstream.socket, + WreqWsMessage::text(first_event), + plan_usage_permit, + record_upstream_request_state, + ) + .await + .map_err(|error| match error { + ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost => { + ResponsesWebSocketUpstreamBindError::PlanUsagePermitLost + } + ResponsesWebSocketUpstreamSendError::Transport(_) => { + ResponsesWebSocketUpstreamBindError::Transport( + "responses_websocket_initial_send_failed", + ) + } + })?; let client_model = initial_event .get("model") .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .ok_or("responses_websocket_model_missing")? + .ok_or(ResponsesWebSocketUpstreamBindError::Transport( + "responses_websocket_model_missing", + ))? .to_string(); let provider_model = decision .provider_request_body @@ -107,7 +163,9 @@ async fn bind_responses_upstream_inner( .map(str::trim) .filter(|value| !value.is_empty()) }) - .ok_or("responses_websocket_mapped_model_missing")? + .ok_or(ResponsesWebSocketUpstreamBindError::Transport( + "responses_websocket_mapped_model_missing", + ))? .to_string(); let responses_lite_static_config = @@ -139,6 +197,64 @@ async fn bind_responses_upstream_inner( }) } +pub(super) async fn send_responses_websocket_upstream_message( + upstream: &mut wreq::ws::WebSocket, + message: WreqWsMessage, + plan_usage_permit: Option<&aether_runtime::AdmissionPermit>, + record_upstream_request_state: F, +) -> Result<(), ResponsesWebSocketUpstreamSendError> +where + F: FnMut(UpstreamRequestState), +{ + let mut record_upstream_request_state = record_upstream_request_state; + if !responses_websocket_plan_permit_is_healthy(plan_usage_permit) { + return Err(ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost); + } + feed_upstream_message(upstream, message) + .await + .map_err(ResponsesWebSocketUpstreamSendError::Transport)?; + // No await between successful start_send and the state transition: an + // outer supervisor cannot cancel this attempt in the ambiguity window. + record_upstream_request_state(UpstreamRequestState::PossiblySent); + flush_upstream_messages(upstream) + .await + .map_err(ResponsesWebSocketUpstreamSendError::Transport)?; + record_upstream_request_state(UpstreamRequestState::Sent); + Ok(()) +} + +#[cfg(test)] +async fn complete_responses_websocket_upstream_send( + queue: Q, + flush: F, + plan_usage_permit: Option<&aether_runtime::AdmissionPermit>, + mut record_upstream_request_state: R, +) -> Result<(), ResponsesWebSocketUpstreamSendError> +where + Q: std::future::Future>, + F: std::future::Future>, + R: FnMut(UpstreamRequestState), +{ + // Concurrency lease health is an admission condition, so check it before + // the socket writer is polled. Once polling begins, SinkExt::send may have + // completed start_send while still waiting for flush; cancelling it at + // that point cannot prove the provider did not receive response.create. + if !responses_websocket_plan_permit_is_healthy(plan_usage_permit) { + return Err(ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost); + } + queue + .await + .map_err(ResponsesWebSocketUpstreamSendError::Transport)?; + // No await between successful start_send and the state transition: an + // outer supervisor cannot cancel this attempt in the ambiguity window. + record_upstream_request_state(UpstreamRequestState::PossiblySent); + flush + .await + .map_err(ResponsesWebSocketUpstreamSendError::Transport)?; + record_upstream_request_state(UpstreamRequestState::Sent); + Ok(()) +} + pub(super) async fn receive_optional_upstream( upstream: &mut Option, ) -> Option> { @@ -181,13 +297,175 @@ pub(super) fn decision_bound_upstream_change_fields( #[cfg(test)] mod tests { + use std::future::Future; + use std::pin::Pin; + use std::sync::atomic::{AtomicBool, Ordering}; + use std::sync::Arc; + use std::task::{Context, Poll}; use std::time::Duration; use aether_contracts::ExecutionTimeouts; use crate::ai_serving::AiExecutionDecision; - use super::{resolve_upstream_handshake_deadline, DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS}; + use super::{ + complete_responses_websocket_upstream_send, resolve_upstream_handshake_deadline, + ResponsesWebSocketUpstreamBindError, DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS, + }; + use crate::handlers::proxy::websocket::responses::turn::UpstreamRequestState; + use crate::handlers::proxy::websocket::transport::WebSocketWriteError; + + struct ReadySendThatRequestsSupervisorCancellation { + cancellation_requested: Arc, + } + + struct MutablePermitHealth(Arc); + + impl aether_runtime::AdmissionPermitHealth for MutablePermitHealth { + fn is_healthy(&self) -> bool { + self.0.load(Ordering::Acquire) + } + } + + impl Future for ReadySendThatRequestsSupervisorCancellation { + type Output = Result<(), WebSocketWriteError>; + + fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll { + // Model the socket writer waking an already-running outer + // supervisor in the same poll in which it reports a successful + // provider transfer. The supervisor cannot drop this future until + // control returns from this poll. + self.cancellation_requested.store(true, Ordering::Release); + Poll::Ready(Ok(())) + } + } + + #[tokio::test] + async fn successful_send_marks_attempt_before_returning_to_cancelling_supervisor() { + let cancellation_requested = Arc::new(AtomicBool::new(false)); + let upstream_request_sent = Arc::new(AtomicBool::new(false)); + let state_by_handoff = Arc::clone(&upstream_request_sent); + + let result = complete_responses_websocket_upstream_send( + ReadySendThatRequestsSupervisorCancellation { + cancellation_requested: Arc::clone(&cancellation_requested), + }, + async { Ok(()) }, + None, + move |state| { + if state == UpstreamRequestState::Sent { + state_by_handoff.store(true, Ordering::Release); + } + }, + ) + .await; + + assert_eq!(result, Ok(())); + assert!(cancellation_requested.load(Ordering::Acquire)); + assert!( + upstream_request_sent.load(Ordering::Acquire), + "a successful provider write must transfer lifecycle ownership before an outer supervisor can cancel the send future" + ); + } + + #[tokio::test] + async fn unhealthy_plan_permit_rejects_before_polling_the_upstream_send() { + let healthy = Arc::new(AtomicBool::new(false)); + let permit = + aether_runtime::AdmissionPermit::from_parts(None, Some(MutablePermitHealth(healthy))) + .expect("test permit"); + let send_polled = Arc::new(AtomicBool::new(false)); + let send_polled_by_future = Arc::clone(&send_polled); + + let result = complete_responses_websocket_upstream_send( + async move { + send_polled_by_future.store(true, Ordering::Release); + Ok(()) + }, + async { Ok(()) }, + Some(&permit), + |_| {}, + ) + .await; + + assert_eq!( + result, + Err(super::ResponsesWebSocketUpstreamSendError::PlanUsagePermitLost) + ); + assert!(!send_polled.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn plan_permit_loss_after_send_start_does_not_cancel_the_send() { + let healthy = Arc::new(AtomicBool::new(true)); + let permit = aether_runtime::AdmissionPermit::from_parts( + None, + Some(MutablePermitHealth(Arc::clone(&healthy))), + ) + .expect("test permit"); + let flush_started = Arc::new(tokio::sync::Notify::new()); + let allow_flush_to_finish = Arc::new(tokio::sync::Notify::new()); + let flush_started_by_future = Arc::clone(&flush_started); + let allow_flush_to_finish_by_future = Arc::clone(&allow_flush_to_finish); + let marked_possibly_sent = Arc::new(AtomicBool::new(false)); + let marked_possibly_sent_by_callback = Arc::clone(&marked_possibly_sent); + let marked_sent = Arc::new(AtomicBool::new(false)); + let marked_sent_by_callback = Arc::clone(&marked_sent); + + let send = complete_responses_websocket_upstream_send( + async { Ok(()) }, + async move { + flush_started_by_future.notify_one(); + allow_flush_to_finish_by_future.notified().await; + Ok(()) + }, + Some(&permit), + move |state| match state { + UpstreamRequestState::PossiblySent => { + marked_possibly_sent_by_callback.store(true, Ordering::Release) + } + UpstreamRequestState::Sent => { + marked_sent_by_callback.store(true, Ordering::Release) + } + UpstreamRequestState::NotStarted => {} + }, + ); + let revoke = async move { + flush_started.notified().await; + healthy.store(false, Ordering::Release); + allow_flush_to_finish.notify_one(); + }; + let (result, ()) = tokio::join!(send, revoke); + + assert_eq!(result, Ok(())); + assert!(marked_possibly_sent.load(Ordering::Acquire)); + assert!(marked_sent.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn flush_failure_keeps_the_attempt_possibly_sent() { + let observed = Arc::new(std::sync::Mutex::new(Vec::new())); + let observed_by_callback = Arc::clone(&observed); + + let result = complete_responses_websocket_upstream_send( + async { Ok(()) }, + async { Err(WebSocketWriteError::TimedOut) }, + None, + move |state| observed_by_callback.lock().expect("states").push(state), + ) + .await; + + assert_eq!( + result, + Err(super::ResponsesWebSocketUpstreamSendError::Transport( + WebSocketWriteError::TimedOut + )) + ); + assert_eq!( + *observed.lock().expect("states"), + vec![UpstreamRequestState::PossiblySent] + ); + } fn sample_decision() -> AiExecutionDecision { AiExecutionDecision { @@ -336,12 +614,16 @@ mod tests { ResponsesWebSocketBodyNormalization::for_tests("test-model"), &json!({"type": "response.create", "model": "test-model"}), adapter, + None, + |_| {}, ) .await; assert_eq!( result.err().expect("bind should fail with timeout"), - "responses_websocket_upstream_handshake_timeout" + ResponsesWebSocketUpstreamBindError::Transport( + "responses_websocket_upstream_handshake_timeout" + ) ); } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs b/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs index 4fe1f30aa..4da39352d 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs @@ -7,6 +7,7 @@ use std::collections::BTreeMap; use std::time::Duration; +use aether_contracts::ProxySnapshot; use axum::extract::ws::{CloseFrame as AxumCloseFrame, Message as AxumWsMessage, WebSocket}; use axum::http::header::{ ACCEPT, ACCEPT_ENCODING, CONNECTION, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE, HOST, @@ -22,7 +23,8 @@ use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage}; use crate::ai_serving::AiExecutionDecision; use crate::execution_runtime::transport::{ - build_browser_wreq_client, build_request_headers, ExecutionTransportControls, + build_browser_wreq_client, build_request_headers, normalize_execution_proxy_url, + ExecutionTransportControls, }; use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error; use crate::handlers::proxy::websocket::session::{ @@ -64,7 +66,7 @@ pub(crate) async fn connect_upstream_websocket( )?; let headers = websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?; - let client = build_websocket_client(decision, errors)?; + let client = build_websocket_client(decision, &upstream_url, errors).await?; let response = client .websocket(upstream_url.as_str()) .headers(headers) @@ -100,14 +102,26 @@ fn guarded_websocket_upstream_url( } fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); headers .iter() .filter(|(name, _)| websocket_response_header_is_safe_to_retain(name)) .filter_map(|(name, value)| { + let normalized = name.as_str().to_ascii_lowercase(); + if crate::headers::should_skip_response_header(&normalized) + || connection_declared.contains(&normalized) + { + return None; + } value .to_str() .ok() - .map(|value| (name.as_str().to_string(), value.to_string())) + .map(|value| (normalized, value.to_string())) }) .collect() } @@ -135,16 +149,25 @@ pub(crate) fn websocket_upstream_url( invalid_code: &'static str, ) -> Result { let mut url = Url::parse(raw).map_err(|_| invalid_code)?; - if url.host_str().is_none() || !url.username().is_empty() || url.password().is_some() { + if url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + || url.fragment().is_some() + { return Err(invalid_code); } let websocket_scheme = match url.scheme() { "https" => "wss", "http" => "ws", - "wss" | "ws" => return Ok(url), + "wss" => return Ok(url), + "ws" if aether_http::url_has_literal_loopback_host(&url) => return Ok(url), + "ws" => return Err(invalid_code), _ => return Err(invalid_code), }; url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?; + if url.scheme() == "ws" && !aether_http::url_has_literal_loopback_host(&url) { + return Err(invalid_code); + } Ok(url) } @@ -200,11 +223,13 @@ pub(crate) fn websocket_handshake_headers( Ok(headers) } -fn build_websocket_client( +async fn build_websocket_client( decision: &AiExecutionDecision, + upstream_url: &Url, errors: UpstreamWebSocketErrorCodes, ) -> Result { let timeouts = websocket_timeouts(decision); + let proxy_url = resolve_websocket_proxy_url(decision.proxy.as_ref(), errors)?; if let Some(profile) = decision.transport_profile.as_ref() { return build_browser_wreq_client( timeouts.as_ref(), @@ -216,30 +241,99 @@ fn build_websocket_client( .map_err(|_| errors.client_build_failed); } - let mut builder = wreq::Client::builder(); + let mut builder = wreq::Client::builder().no_proxy(); if let Some(connect_ms) = timeouts.as_ref().and_then(|timeouts| timeouts.connect_ms) { builder = builder.connect_timeout(Duration::from_millis(connect_ms)); } - if let Some(proxy) = decision - .proxy - .as_ref() - .filter(|proxy| proxy.enabled != Some(false)) - { - if let Some(proxy_url) = proxy - .url - .as_deref() - .map(str::trim) - .filter(|url| !url.is_empty()) - { - let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?; - builder = builder.proxy(proxy); - } else if proxy.node_id.is_some() || proxy.mode.as_deref() == Some("tunnel") { - return Err(errors.tunnel_proxy_unsupported); + if let Some(proxy_url) = proxy_url { + let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?; + builder = builder.proxy(proxy); + } else { + // Pin every direct WebSocket connection to the DNS answers validated + // here. This also covers the explicitly permitted loopback `ws://` + // form; otherwise the client would perform a second lookup and a + // rebinding could escape the loopback-only policy. + let host = upstream_url.host_str().ok_or(errors.upstream_url_invalid)?; + let port = upstream_url + .port_or_known_default() + .ok_or(errors.upstream_url_invalid)?; + let addresses = if let Ok(ip) = host.parse::() { + vec![std::net::SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits( + host, + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| errors.upstream_url_invalid)? + }; + let allows_loopback = host.trim_end_matches('.').eq_ignore_ascii_case("localhost") + || host + .parse::() + .map(|ip| ip.is_loopback()) + .unwrap_or(false); + let unsafe_answer = if allows_loopback { + addresses.iter().any(|address| !address.ip().is_loopback()) + } else { + addresses + .iter() + .any(|address| aether_http::is_private_or_reserved_ip(address.ip())) + }; + if addresses.is_empty() || unsafe_answer { + return Err(errors.upstream_url_invalid); } + builder = builder.resolve_to_addrs(host.to_string(), addresses.iter().copied()); } builder.build().map_err(|_| errors.client_build_failed) } +fn resolve_websocket_proxy_url( + proxy: Option<&ProxySnapshot>, + errors: UpstreamWebSocketErrorCodes, +) -> Result, &'static str> { + let Some(proxy) = proxy else { + return Ok(None); + }; + if proxy.enabled == Some(false) { + return Ok(None); + } + if let Some(proxy_url) = proxy + .url + .as_deref() + .map(str::trim) + .filter(|url| !url.is_empty()) + { + let parsed = Url::parse(proxy_url).map_err(|_| errors.proxy_invalid)?; + if !matches!( + parsed.scheme().to_ascii_lowercase().as_str(), + "http" | "https" | "socks5" | "socks5h" + ) || parsed.host_str().is_none() + || !matches!(parsed.path(), "" | "/") + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return Err(errors.proxy_invalid); + } + // Manual proxy nodes bind credentials to the node identity before a + // snapshot reaches this path. Reject userinfo on an otherwise + // unbound snapshot so an arbitrary decision cannot smuggle proxy + // credentials through a URL; preserve the established node-auth URL + // form for authenticated manual proxy nodes. + if (!parsed.username().is_empty() || parsed.password().is_some()) && proxy.node_id.is_none() + { + return Err(errors.proxy_invalid); + } + let normalized = + normalize_execution_proxy_url(proxy_url).map_err(|_| errors.proxy_invalid)?; + return Ok(Some(normalized)); + } + if proxy.node_id.is_some() || proxy.mode.as_deref() == Some("tunnel") { + return Err(errors.tunnel_proxy_unsupported); + } + Err(errors.proxy_invalid) +} + pub(crate) fn websocket_timeouts( decision: &AiExecutionDecision, ) -> Option { @@ -355,6 +449,23 @@ pub(crate) async fn send_upstream_message( bounded_send(RELAY_WRITE_TIMEOUT, upstream.send(message).map_err(|_| ())).await } +/// Queues one frame in the upstream sink without flushing it. Completion means +/// `start_send` succeeded, so callers must conservatively treat the frame as +/// possibly delivered even when a later flush fails or is cancelled. +pub(crate) async fn feed_upstream_message( + upstream: &mut wreq::ws::WebSocket, + message: WreqWsMessage, +) -> Result<(), WebSocketWriteError> { + bounded_send(RELAY_WRITE_TIMEOUT, upstream.feed(message).map_err(|_| ())).await +} + +/// Flushes frames previously queued with [`feed_upstream_message`]. +pub(crate) async fn flush_upstream_messages( + upstream: &mut wreq::ws::WebSocket, +) -> Result<(), WebSocketWriteError> { + bounded_send(RELAY_WRITE_TIMEOUT, upstream.flush().map_err(|_| ())).await +} + /// Best-effort teardown write. The caller is already ending the session, so /// the outcome only matters for keeping the wait bounded. async fn send_teardown_message(write: F) @@ -567,13 +678,15 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16 #[cfg(test)] mod tests { use super::{ - bounded_send, guarded_websocket_upstream_url, responses_websocket_error_event, - responses_websocket_error_event_with_stream_id, websocket_handshake_headers, - websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url, - WebSocketRelayPumpControl, WebSocketRelayQueueError, WebSocketWriteError, - RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT, + bounded_send, guarded_websocket_upstream_url, resolve_websocket_proxy_url, + responses_websocket_error_event, responses_websocket_error_event_with_stream_id, + websocket_handshake_headers, websocket_relay_frame_queue, websocket_response_headers, + websocket_upstream_url, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl, + WebSocketRelayQueueError, WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY, + RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT, }; use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url; + use aether_contracts::ProxySnapshot; use axum::http::HeaderMap; use std::collections::BTreeMap; use std::time::Duration; @@ -747,6 +860,68 @@ mod tests { assert!(websocket_upstream_url("https://token@example.test/responses", "invalid").is_err()); } + #[test] + fn remote_websocket_requires_wss_but_loopback_ws_is_allowed() { + for allowed in [ + "wss://example.test/v1/responses", + "https://example.test/v1/responses", + "ws://localhost:8080/v1/responses", + "http://127.42.0.1:8080/v1/responses", + "ws://[::1]:8080/v1/responses", + ] { + assert!( + websocket_upstream_url(allowed, "invalid").is_ok(), + "{allowed}" + ); + } + for rejected in [ + "ws://example.test/v1/responses", + "http://10.0.0.1/v1/responses", + "ws://0.0.0.0:8080/v1/responses", + "ws://[::ffff:127.0.0.1]:8080/v1/responses", + "wss://example.test/v1/responses#secret", + ] { + assert!( + websocket_upstream_url(rejected, "invalid").is_err(), + "{rejected}" + ); + } + } + + #[test] + fn active_websocket_proxy_without_a_target_fails_closed() { + let errors = UpstreamWebSocketErrorCodes { + upstream_url_missing: "missing", + upstream_url_invalid: "upstream_invalid", + frontdoor_self_loop: "frontdoor_self_loop", + headers_invalid: "headers_invalid", + client_build_failed: "client_build_failed", + proxy_invalid: "proxy_invalid", + tunnel_proxy_unsupported: "tunnel_unsupported", + handshake_failed: "handshake_failed", + upgrade_rejected: "upgrade_rejected", + upgrade_failed: "upgrade_failed", + }; + let missing = ProxySnapshot { + enabled: Some(true), + ..ProxySnapshot::default() + }; + assert_eq!( + resolve_websocket_proxy_url(Some(&missing), errors), + Err("proxy_invalid") + ); + + let tunnel = ProxySnapshot { + enabled: Some(true), + mode: Some("tunnel".to_string()), + ..ProxySnapshot::default() + }; + assert_eq!( + resolve_websocket_proxy_url(Some(&tunnel), errors), + Err("tunnel_unsupported") + ); + } + #[test] fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() { let base_url = configured_gateway_frontdoor_base_url(); @@ -762,6 +937,64 @@ mod tests { ); } + #[test] + fn websocket_proxy_url_must_be_an_allowed_origin() { + let errors = UpstreamWebSocketErrorCodes { + upstream_url_missing: "missing", + upstream_url_invalid: "upstream_invalid", + frontdoor_self_loop: "frontdoor_self_loop", + headers_invalid: "headers_invalid", + client_build_failed: "client_build_failed", + proxy_invalid: "proxy_invalid", + tunnel_proxy_unsupported: "tunnel_unsupported", + handshake_failed: "handshake_failed", + upgrade_rejected: "upgrade_rejected", + upgrade_failed: "upgrade_failed", + }; + + for value in [ + "file:///tmp/proxy", + "http://proxy.example:8080/path", + "http://proxy.example:8080?token=secret", + "http://proxy.example:8080#fragment", + "http://alice:password@proxy.example:8080", + ] { + let proxy = ProxySnapshot { + enabled: Some(true), + url: Some(value.to_string()), + ..ProxySnapshot::default() + }; + assert_eq!( + resolve_websocket_proxy_url(Some(&proxy), errors), + Err("proxy_invalid"), + "proxy URL should be rejected: {value}" + ); + } + + let authenticated_node = ProxySnapshot { + enabled: Some(true), + node_id: Some("manual-node-1".to_string()), + url: Some("http://alice:password@proxy.example:8080".to_string()), + ..ProxySnapshot::default() + }; + assert_eq!( + resolve_websocket_proxy_url(Some(&authenticated_node), errors), + Ok(Some( + "http://alice:password@proxy.example:8080/".to_string() + )) + ); + + let socks = ProxySnapshot { + enabled: Some(true), + url: Some("socks5://proxy.example:1080".to_string()), + ..ProxySnapshot::default() + }; + assert_eq!( + resolve_websocket_proxy_url(Some(&socks), errors), + Ok(Some("socks5h://proxy.example:1080".to_string())) + ); + } + #[test] fn rejects_live_direct_and_sideband_frontdoor_self_loops_before_connecting() { let base_url = configured_gateway_frontdoor_base_url(); diff --git a/apps/aether-gateway/src/handlers/public/ai_public.rs b/apps/aether-gateway/src/handlers/public/ai_public.rs index b689cc30a..1a3f2e711 100644 --- a/apps/aether-gateway/src/handlers/public/ai_public.rs +++ b/apps/aether-gateway/src/handlers/public/ai_public.rs @@ -4,8 +4,14 @@ use crate::ai_serving::{ use crate::async_task::CancelVideoTaskError; use crate::control::GatewayControlDecision; use crate::control::GatewayPublicRequestContext; +use crate::handlers::shared::{ + find_multipart_boundary, find_multipart_boundary_after_crlf, parse_multipart_boundary, + query_param_value, unix_ms_to_rfc3339, unix_secs_to_rfc3339, MAX_MULTIPART_PARTS, + MAX_MULTIPART_PART_HEADER_BYTES, +}; use crate::image_capabilities::openai_image_gateway_max_generation_count; use crate::{AppState, GatewayError}; +use aether_data_contracts::repository::gemini_file_mappings::StoredGeminiFileMapping; use aether_data_contracts::repository::video_tasks::{ StoredVideoTask, VideoTaskQueryFilter, VideoTaskStatus, }; @@ -16,8 +22,12 @@ use axum::Json; use serde_json::{json, Value}; const GEMINI_VIDEO_TASK_NOT_FOUND_DETAIL: &str = "Video task not found"; +const GEMINI_FILE_NOT_FOUND_DETAIL: &str = "File not found"; +const GEMINI_FILES_DATA_UNAVAILABLE_DETAIL: &str = "Gemini Files data is unavailable"; const AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL: &str = "Method not allowed"; const AI_PUBLIC_UNAUTHORIZED_DETAIL: &str = "Unauthorized"; +const AI_PUBLIC_INTERNAL_ERROR_DETAIL: &str = "Service temporarily unavailable"; +const AI_PUBLIC_UPSTREAM_ERROR_DETAIL: &str = "Upstream request failed"; const OPENAI_IMAGE_PROMPT_DETAIL: &str = "图片生成/编辑请求缺少 prompt"; const OPENAI_IMAGE_EDIT_INPUT_DETAIL: &str = "图片编辑请求至少需要 1 张输入图片"; const OPENAI_IMAGE_PARTIAL_IMAGES_DETAIL: &str = @@ -150,9 +160,269 @@ pub(crate) async fn maybe_build_local_ai_public_response( return Some(response); } + if let Some(response) = + maybe_build_local_gemini_files_response(state, request_context, decision).await + { + return Some(response); + } + maybe_build_local_gemini_video_operations_response(state, request_context, decision).await } +async fn maybe_build_local_gemini_files_response( + state: &AppState, + request_context: &GatewayPublicRequestContext, + decision: &GatewayControlDecision, +) -> Option> { + if decision.route_family.as_deref() != Some("gemini") + || decision.route_kind.as_deref() != Some("files") + || !request_context.request_path.starts_with("/v1beta/files") + { + return None; + } + + let Some(user_id) = allowed_ai_public_user_id(decision) else { + return Some(build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + )); + }; + + if request_context.request_path == "/v1beta/files" { + return Some(match request_context.request_method { + http::Method::GET => { + build_local_gemini_files_list_response(state, request_context, user_id).await + } + _ => build_ai_public_error_response( + http::StatusCode::METHOD_NOT_ALLOWED, + AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL, + ), + }); + } + + if !matches!( + request_context.request_method, + http::Method::GET | http::Method::DELETE + ) { + return Some(build_ai_public_error_response( + http::StatusCode::METHOD_NOT_ALLOWED, + AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL, + )); + } + + let file_name = normalize_gemini_file_request_path(request_context.request_path.as_str()); + if let Some(short_id) = file_name + .as_deref() + .and_then(|value| value.strip_prefix("files/")) + .and_then(|file_id| file_id.strip_prefix("aev_")) + .filter(|value| !value.is_empty()) + { + return Some( + build_local_gemini_video_file_response( + state, + request_context.request_method.clone(), + user_id, + short_id, + ) + .await, + ); + } + + if !state.has_gemini_file_mapping_data_reader() { + return Some(build_ai_public_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + GEMINI_FILES_DATA_UNAVAILABLE_DETAIL, + )); + } + let Some(file_name) = file_name else { + return Some(build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + )); + }; + let mapping = match state + .find_active_gemini_file_mapping_for_user( + file_name.as_str(), + user_id, + crate::clock::current_unix_secs(), + ) + .await + { + Ok(Some(mapping)) => mapping, + Ok(None) => { + return Some(build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + )); + } + Err(_) => { + return Some(build_ai_public_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + GEMINI_FILES_DATA_UNAVAILABLE_DETAIL, + )); + } + }; + if mapping.user_id.as_deref().map(str::trim) != Some(user_id) { + return Some(build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + )); + } + + None +} + +fn allowed_ai_public_user_id(decision: &GatewayControlDecision) -> Option<&str> { + decision + .auth_context + .as_ref() + .filter(|auth_context| auth_context.access_allowed) + .map(|auth_context| auth_context.user_id.trim()) + .filter(|value| !value.is_empty()) +} + +async fn build_local_gemini_video_file_response( + state: &AppState, + method: http::Method, + user_id: &str, + short_id: &str, +) -> Response { + if method != http::Method::GET { + return build_ai_public_error_response( + http::StatusCode::METHOD_NOT_ALLOWED, + AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL, + ); + } + + let task = match state + .find_video_task_by_short_id_for_user(short_id, user_id) + .await + { + Ok(Some(task)) if is_gemini_video_task(&task) => task, + Ok(_) => { + return build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + ); + } + Err(_) => { + return build_ai_public_internal_error_response( + "gemini_video_file_lookup", + "data_store_unavailable", + ); + } + }; + + let source = match crate::async_task::video_task_video_source_from_task(state, &task).await { + Ok(Some(source)) => source, + Ok(None) => { + return build_ai_public_error_response( + http::StatusCode::NOT_FOUND, + GEMINI_FILE_NOT_FOUND_DETAIL, + ); + } + Err(_) => { + return build_ai_public_internal_error_response( + "gemini_video_file_source", + "video_source_unavailable", + ); + } + }; + + match crate::async_task::build_video_task_video_response(state, &task.id, source).await { + Ok(response) => response, + Err(_) => build_ai_public_internal_error_response( + "gemini_video_file_delivery", + "video_delivery_failed", + ), + } +} + +async fn build_local_gemini_files_list_response( + state: &AppState, + request_context: &GatewayPublicRequestContext, + user_id: &str, +) -> Response { + if !state.has_gemini_file_mapping_data_reader() { + return build_ai_public_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + GEMINI_FILES_DATA_UNAVAILABLE_DETAIL, + ); + } + let page_size = query_param_value(request_context.request_query_string.as_deref(), "pageSize") + .and_then(|value| value.parse::().ok()) + .filter(|value| *value > 0) + .unwrap_or(10) + .min(100); + let offset = query_param_value(request_context.request_query_string.as_deref(), "pageToken") + .and_then(|value| value.parse::().ok()) + .unwrap_or(0); + let mappings = match state + .list_gemini_file_mappings( + &aether_data::repository::gemini_file_mappings::GeminiFileMappingListQuery { + user_id: Some(user_id.to_string()), + include_expired: false, + search: None, + offset, + limit: page_size, + now_unix_secs: crate::clock::current_unix_secs(), + }, + ) + .await + { + Ok(mappings) => mappings, + Err(_) => { + return build_ai_public_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + GEMINI_FILES_DATA_UNAVAILABLE_DETAIL, + ); + } + }; + let files = mappings + .items + .iter() + .map(build_gemini_file_mapping_payload) + .collect::>(); + let next_offset = offset.saturating_add(files.len()); + let mut payload = json!({ "files": files }); + if next_offset < mappings.total { + payload["nextPageToken"] = Value::String(next_offset.to_string()); + } + Json(payload).into_response() +} + +fn build_gemini_file_mapping_payload(mapping: &StoredGeminiFileMapping) -> Value { + let mut payload = serde_json::Map::new(); + payload.insert("name".to_string(), Value::String(mapping.file_name.clone())); + if let Some(display_name) = mapping.display_name.as_ref() { + payload.insert( + "displayName".to_string(), + Value::String(display_name.clone()), + ); + } + if let Some(mime_type) = mapping.mime_type.as_ref() { + payload.insert("mimeType".to_string(), Value::String(mime_type.clone())); + } + if let Some(created_at) = unix_ms_to_rfc3339(mapping.created_at_unix_ms) { + payload.insert("createTime".to_string(), Value::String(created_at)); + } + if let Some(expires_at) = unix_secs_to_rfc3339(mapping.expires_at_unix_secs) { + payload.insert("expirationTime".to_string(), Value::String(expires_at)); + } + payload.insert("state".to_string(), Value::String("ACTIVE".to_string())); + Value::Object(payload) +} + +fn normalize_gemini_file_request_path(path: &str) -> Option { + let suffix = path.strip_prefix("/v1beta/files/")?.trim_matches('/'); + let suffix = suffix.strip_suffix(":download").unwrap_or(suffix).trim(); + let suffix = suffix.strip_prefix("files/").unwrap_or(suffix).trim(); + if suffix.is_empty() || suffix.contains('/') { + return None; + } + Some(format!("files/{suffix}")) +} + fn maybe_build_local_openai_request_validation_response( request_context: &GatewayPublicRequestContext, request_body: Option<&Bytes>, @@ -675,7 +945,8 @@ fn parse_openai_image_validation_input_from_multipart( request_body: &Bytes, content_type: &str, ) -> Result { - let boundary = multipart_boundary(content_type).ok_or(OPENAI_IMAGE_INVALID_MULTIPART_DETAIL)?; + let boundary = + parse_multipart_boundary(content_type).ok_or(OPENAI_IMAGE_INVALID_MULTIPART_DETAIL)?; let fields = parse_multipart_fields(request_body, &boundary); if fields.is_empty() { return Err(OPENAI_IMAGE_INVALID_MULTIPART_DETAIL); @@ -780,67 +1051,203 @@ fn parse_multipart_fields(body: &[u8], boundary: &str) -> Vec { let delimiter = format!("--{boundary}").into_bytes(); let mut parts = Vec::new(); let mut cursor = 0usize; + let mut part_count = 0usize; - while let Some(index) = find_subslice(&body[cursor..], &delimiter) { + while let Some(index) = find_multipart_boundary(&body[cursor..], &delimiter) { let start = cursor + index + delimiter.len(); if body.get(start..start + 2) == Some(b"--") { + let closing_suffix = body.get(start + 2..).unwrap_or_default(); + if !(closing_suffix.is_empty() || closing_suffix.starts_with(b"\r\n")) { + return Vec::new(); + } break; } + part_count = part_count.saturating_add(1); + if part_count > MAX_MULTIPART_PARTS { + return Vec::new(); + } let mut part = &body[start..]; if part.starts_with(b"\r\n") { part = &part[2..]; } - let Some(next) = find_subslice(part, &delimiter) else { - break; + // A multipart body is only valid once a real (CRLF-delimited) next + // boundary has been found. Returning the fields parsed so far here + // would let a truncated request pass validation. + let Some(next) = find_multipart_boundary_after_crlf(part, &delimiter) else { + return Vec::new(); }; let raw = &part[..next]; let raw = raw.strip_suffix(b"\r\n").unwrap_or(raw); - if let Some(field) = parse_multipart_field(raw) { - parts.push(field); + if find_subslice(raw, b"\r\n\r\n") + .is_some_and(|header_end| header_end > MAX_MULTIPART_PART_HEADER_BYTES) + { + return Vec::new(); } + let Some(field) = parse_multipart_field(raw) else { + return Vec::new(); + }; + parts.push(field); cursor = start + next; } parts } -fn multipart_boundary(content_type: &str) -> Option { - content_type.split(';').find_map(|segment| { - let (key, value) = segment.trim().split_once('=')?; - if !key.trim().eq_ignore_ascii_case("boundary") { - return None; - } - let boundary = value.trim().trim_matches('"').trim(); - (!boundary.is_empty()).then(|| boundary.to_string()) - }) -} - fn parse_multipart_field(raw: &[u8]) -> Option { let header_end = find_subslice(raw, b"\r\n\r\n")?; let headers = &raw[..header_end]; let data = raw.get(header_end + 4..)?.to_vec(); - let header_text = String::from_utf8_lossy(headers); + let header_text = std::str::from_utf8(headers).ok()?; let mut name = None; - for line in header_text.lines() { - let trimmed = line.trim(); - if trimmed - .to_ascii_lowercase() - .starts_with("content-disposition:") - { - name = extract_quoted_header_value(trimmed, "name"); + let mut disposition_seen = false; + for line in header_text.split("\r\n") { + let (header_name, header_value) = line.split_once(':')?; + let header_name = header_name.trim(); + if header_name.eq_ignore_ascii_case("content-disposition") { + if disposition_seen { + return None; + } + disposition_seen = true; + name = parse_multipart_content_disposition_name(header_value.trim()); } } Some(MultipartField { name: name?, data }) } -fn extract_quoted_header_value(header: &str, key: &str) -> Option { - let pattern = format!("{key}=\""); - let start = header.find(&pattern)? + pattern.len(); - let rest = &header[start..]; - let end = rest.find('"')?; - Some(rest[..end].to_string()) +fn parse_multipart_content_disposition_name(value: &str) -> Option { + let segments = split_multipart_header_parameters(value)?; + let disposition = segments.first()?.trim(); + if !disposition.eq_ignore_ascii_case("form-data") { + return None; + } + + let mut seen_keys = Vec::new(); + let mut name = None; + for segment in segments.into_iter().skip(1) { + let segment = segment.trim(); + if segment.is_empty() { + return None; + } + let (raw_key, raw_value) = segment.split_once('=')?; + let key = raw_key.trim(); + if key.is_empty() || !key.as_bytes().iter().copied().all(is_multipart_token_byte) { + return None; + } + if seen_keys + .iter() + .any(|seen: &String| seen.eq_ignore_ascii_case(key)) + { + return None; + } + seen_keys.push(key.to_ascii_lowercase()); + + let parsed_value = parse_multipart_header_parameter_value(raw_value.trim())?; + if key.eq_ignore_ascii_case("name") { + if parsed_value.is_empty() { + return None; + } + name = Some(parsed_value); + } + } + + name +} + +fn split_multipart_header_parameters(value: &str) -> Option> { + let mut segments = Vec::new(); + let mut start = 0usize; + let mut in_quotes = false; + let mut escaped = false; + + for (index, byte) in value.as_bytes().iter().copied().enumerate() { + if in_quotes { + if escaped { + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + in_quotes = false; + } + } else if byte == b'"' { + in_quotes = true; + } else if byte == b';' { + segments.push(&value[start..index]); + start = index + 1; + } + } + + if in_quotes || escaped { + return None; + } + segments.push(&value[start..]); + Some(segments) +} + +fn parse_multipart_header_parameter_value(value: &str) -> Option { + if value.is_empty() { + return None; + } + if value.starts_with('"') { + if value.len() < 2 || !value.ends_with('"') { + return None; + } + let inner = &value[1..value.len() - 1]; + let mut parsed = String::with_capacity(inner.len()); + let mut escaped = false; + for character in inner.chars() { + if escaped { + if character.is_control() { + return None; + } + parsed.push(character); + escaped = false; + } else if character == '\\' { + escaped = true; + } else { + if character == '"' || character.is_control() { + return None; + } + parsed.push(character); + } + } + if escaped { + return None; + } + return Some(parsed); + } + + value + .as_bytes() + .iter() + .copied() + .all(is_multipart_token_byte) + .then(|| value.to_string()) +} + +fn is_multipart_token_byte(byte: u8) -> bool { + matches!( + byte, + b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + | b'!' + | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) } fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option { @@ -1339,12 +1746,7 @@ async fn build_local_gemini_video_operations_list_response( state: &AppState, decision: &GatewayControlDecision, ) -> Response { - let Some(user_id) = decision - .auth_context - .as_ref() - .map(|auth_context| auth_context.user_id.trim()) - .filter(|value| !value.is_empty()) - else { + let Some(user_id) = allowed_ai_public_user_id(decision) else { return build_ai_public_error_response( http::StatusCode::UNAUTHORIZED, AI_PUBLIC_UNAUTHORIZED_DETAIL, @@ -1359,10 +1761,10 @@ async fn build_local_gemini_video_operations_list_response( }; let tasks = match state.list_video_task_page(&filter, 0, 100).await { Ok(tasks) => tasks, - Err(err) => { - return build_ai_public_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("{err:?}"), + Err(_) => { + return build_ai_public_internal_error_response( + "gemini_video_operations_list", + "data_store_unavailable", ); } }; @@ -1389,10 +1791,10 @@ async fn build_local_gemini_video_operation_detail_response( GEMINI_VIDEO_TASK_NOT_FOUND_DETAIL, ); } - Err(err) => { - return build_ai_public_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("{err:?}"), + Err(_) => { + return build_ai_public_internal_error_response( + "gemini_video_operation_detail", + "data_store_unavailable", ); } }; @@ -1405,6 +1807,12 @@ async fn build_local_gemini_video_operation_cancel_response( decision: &GatewayControlDecision, operation_path: &str, ) -> Response { + let Some(user_id) = allowed_ai_public_user_id(decision) else { + return build_ai_public_error_response( + http::StatusCode::UNAUTHORIZED, + AI_PUBLIC_UNAUTHORIZED_DETAIL, + ); + }; let task = match find_user_gemini_video_task_for_operation(state, decision, operation_path).await { Ok(Some(task)) => task, @@ -1414,15 +1822,15 @@ async fn build_local_gemini_video_operation_cancel_response( GEMINI_VIDEO_TASK_NOT_FOUND_DETAIL, ); } - Err(err) => { - return build_ai_public_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("{err:?}"), + Err(_) => { + return build_ai_public_internal_error_response( + "gemini_video_operation_cancel_lookup", + "data_store_unavailable", ); } }; - match crate::async_task::cancel_video_task_record(state, &task.id).await { + match crate::async_task::cancel_video_task_record_for_user(state, &task.id, user_id).await { Ok(_) => Json(json!({})).into_response(), Err(CancelVideoTaskError::NotFound) => build_ai_public_error_response( http::StatusCode::NOT_FOUND, @@ -1435,10 +1843,12 @@ async fn build_local_gemini_video_operation_cancel_response( video_task_status_name(status) ), ), - Err(CancelVideoTaskError::Response(response)) => response, - Err(CancelVideoTaskError::Gateway(err)) => build_ai_public_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("{err:?}"), + Err(CancelVideoTaskError::Response(response)) => { + build_ai_public_upstream_error_response(response, "gemini_video_operation_cancel") + } + Err(CancelVideoTaskError::Gateway(_)) => build_ai_public_internal_error_response( + "gemini_video_operation_cancel", + "cancel_execution_failed", ), } } @@ -1448,21 +1858,19 @@ async fn find_user_gemini_video_task_for_operation( decision: &GatewayControlDecision, operation_path: &str, ) -> Result, GatewayError> { - let Some(user_id) = decision - .auth_context - .as_ref() - .map(|auth_context| auth_context.user_id.trim()) - .filter(|value| !value.is_empty()) - else { + let Some(user_id) = allowed_ai_public_user_id(decision) else { return Ok(None); }; let Some(short_id) = extract_short_id_from_gemini_operation_path(operation_path) else { return Ok(None); }; - let Some(task) = state.find_video_task_by_short_id(short_id).await? else { + let Some(task) = state + .find_video_task_by_short_id_for_user(short_id, user_id) + .await? + else { return Ok(None); }; - if task.user_id.as_deref().map(str::trim) != Some(user_id) || !is_gemini_video_task(&task) { + if !is_gemini_video_task(&task) { return Ok(None); } Ok(Some(task)) @@ -1482,13 +1890,7 @@ fn extract_short_id_from_gemini_operation_path(operation_path: &str) -> Option<& } fn is_gemini_video_task(task: &StoredVideoTask) -> bool { - matches!( - task.provider_api_format - .as_deref() - .or(task.client_api_format.as_deref()) - .map(str::trim), - Some("gemini:video") - ) + task.effective_api_format() == Some("gemini:video") } fn build_gemini_video_operation_payload(task: &StoredVideoTask) -> serde_json::Value { @@ -1515,13 +1917,7 @@ fn build_gemini_video_operation_payload(task: &StoredVideoTask) -> serde_json::V VideoTaskStatus::Failed | VideoTaskStatus::Expired => json!({ "name": gemini_video_operation_name(task), "done": true, - "error": { - "code": task.error_code.clone().unwrap_or_else(|| "UNKNOWN".to_string()), - "message": task - .error_message - .clone() - .unwrap_or_else(|| "Video generation failed".to_string()), - } + "error": gemini_video_task_error_projection(task), }), _ => json!({ "name": gemini_video_operation_name(task), @@ -1531,6 +1927,28 @@ fn build_gemini_video_operation_payload(task: &StoredVideoTask) -> serde_json::V } } +fn gemini_video_task_error_projection(task: &StoredVideoTask) -> serde_json::Value { + let code = match task.error_code.as_deref().map(str::trim) { + Some("authentication_error") => "authentication_error", + Some("content_policy_violation") => "content_policy_violation", + Some("expired") => "expired", + Some("invalid_request") => "invalid_request", + Some("not_found") => "not_found", + Some("permission_denied") => "permission_denied", + Some("poll_permanent_error") => "poll_permanent_error", + Some("poll_timeout") => "poll_timeout", + Some("provider_error") => "provider_error", + Some("rate_limit_exceeded") => "rate_limit_exceeded", + Some("server_error") => "server_error", + Some("unknown") => "unknown", + _ => "provider_error", + }; + json!({ + "code": code, + "message": "Video generation failed", + }) +} + fn gemini_video_operation_name(task: &StoredVideoTask) -> String { format!( "models/{}/operations/{}", @@ -1594,6 +2012,67 @@ fn video_task_status_name(status: VideoTaskStatus) -> &'static str { fn build_ai_public_error_response( status: http::StatusCode, detail: impl Into, +) -> Response { + let detail = detail.into(); + let public_detail = if status.is_server_error() { + tracing::error!( + event_name = "ai_public_internal_error", + %status, + "internal AI public API error hidden from client" + ); + AI_PUBLIC_INTERNAL_ERROR_DETAIL.to_string() + } else { + detail + }; + build_ai_public_error_payload(status, public_detail) +} + +fn build_ai_public_internal_error_response( + operation: &'static str, + error_category: &'static str, +) -> Response { + tracing::error!( + event_name = "ai_public_internal_error", + operation, + error_category, + "internal AI public operation failed" + ); + build_ai_public_error_payload( + http::StatusCode::INTERNAL_SERVER_ERROR, + AI_PUBLIC_INTERNAL_ERROR_DETAIL, + ) +} + +fn build_ai_public_upstream_error_response( + response: Response, + operation: &'static str, +) -> Response { + let upstream_status = response.status(); + tracing::warn!( + event_name = "ai_public_upstream_error", + operation, + error_category = "upstream_response_projected", + "upstream error response body discarded" + ); + + // Preserve a provider 4xx status for caller retry/validation semantics, but never + // forward its body. Non-4xx responses are represented as a gateway failure. + let status = if upstream_status.is_client_error() { + upstream_status + } else { + http::StatusCode::BAD_GATEWAY + }; + let detail = if status.is_server_error() { + AI_PUBLIC_INTERNAL_ERROR_DETAIL + } else { + AI_PUBLIC_UPSTREAM_ERROR_DETAIL + }; + build_ai_public_error_payload(status, detail) +} + +fn build_ai_public_error_payload( + status: http::StatusCode, + detail: impl Into, ) -> Response { (status, Json(json!({ "detail": detail.into() }))).into_response() } @@ -1601,14 +2080,138 @@ fn build_ai_public_error_response( #[cfg(test)] mod tests { use super::{ - parse_openai_image_validation_input, validate_claude_count_tokens_request, - validate_openai_image_n, OpenAiImageOperation, CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL, - CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL, CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL, - CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL, + build_ai_public_error_response, build_ai_public_upstream_error_response, + build_gemini_file_mapping_payload, gemini_video_task_error_projection, + parse_multipart_fields, parse_openai_image_validation_input, + validate_claude_count_tokens_request, validate_openai_image_n, OpenAiImageOperation, + StoredGeminiFileMapping, AI_PUBLIC_UPSTREAM_ERROR_DETAIL, + CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL, CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL, + CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL, CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL, + MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES, + OPENAI_IMAGE_INVALID_MULTIPART_DETAIL, }; - use axum::body::Bytes; + use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus}; + use axum::body::{to_bytes, Body, Bytes}; + use axum::http::StatusCode; use serde_json::json; + #[tokio::test] + async fn ai_public_server_errors_do_not_expose_internal_details() { + let response = build_ai_public_error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "database connection failed: password=internal-secret", + ); + + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("error response body should read"); + let payload: serde_json::Value = + serde_json::from_slice(&body).expect("error response should be JSON"); + assert_eq!(payload["detail"], "Service temporarily unavailable"); + assert!(!String::from_utf8_lossy(&body).contains("internal-secret")); + } + + #[tokio::test] + async fn ai_public_upstream_errors_discard_the_upstream_response_body() { + let upstream_response = axum::http::Response::builder() + .status(StatusCode::UNAUTHORIZED) + .body(Body::from( + r#"{"error":{"message":"Bearer upstream-secret at https://internal.test"}}"#, + )) + .expect("upstream response should build"); + + let response = build_ai_public_upstream_error_response(upstream_response, "test_operation"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("projected response body should read"); + let payload: serde_json::Value = + serde_json::from_slice(&body).expect("projected response should be JSON"); + assert_eq!(payload["detail"], AI_PUBLIC_UPSTREAM_ERROR_DETAIL); + let body = String::from_utf8_lossy(&body); + for secret in ["upstream-secret", "Bearer", "internal.test"] { + assert!(!body.contains(secret)); + } + } + + #[test] + fn gemini_video_error_projection_discards_historical_sensitive_diagnostics() { + let mut task = StoredVideoTask::new( + "task-1".to_string(), + Some("short-1".to_string()), + "request-1".to_string(), + Some("user-1".to_string()), + None, + None, + None, + Some("operations/upstream-1".to_string()), + Some("provider-1".to_string()), + Some("endpoint-1".to_string()), + Some("key-1".to_string()), + Some("gemini:video".to_string()), + Some("gemini:video".to_string()), + false, + Some("veo-3".to_string()), + None, + None, + None, + None, + None, + None, + VideoTaskStatus::Failed, + 100, + None, + 0, + 10, + None, + 1, + 360, + 1, + Some(1), + Some(2), + 2, + Some("provider_error".to_string()), + None, + None, + None, + ) + .expect("task should be valid"); + task.error_code = Some( + "Authorization: Bearer code-secret at https://internal.test/?key=secret".to_string(), + ); + task.error_message = Some( + "Authorization: Bearer message-secret at https://internal.test/?token=secret" + .to_string(), + ); + + let payload = gemini_video_task_error_projection(&task); + + assert_eq!(payload["code"], "provider_error"); + assert_eq!(payload["message"], "Video generation failed"); + let encoded = payload.to_string(); + for sensitive in ["Bearer", "code-secret", "message-secret", "internal.test"] { + assert!(!encoded.contains(sensitive)); + } + } + + #[test] + fn gemini_file_payload_formats_millisecond_timestamps_as_milliseconds() { + let mapping = StoredGeminiFileMapping::new( + "mapping-1".to_string(), + "files/file-1".to_string(), + "key-1".to_string(), + 1_710_000_123_456, + 1_710_003_723, + ) + .expect("mapping should be valid"); + + let payload = build_gemini_file_mapping_payload(&mapping); + + assert_eq!(payload["createTime"], "2024-03-09T16:02:03.456Z"); + } + #[test] fn count_tokens_validation_rejects_only_structurally_invalid_requests() { assert_eq!( @@ -1684,6 +2287,143 @@ mod tests { assert_eq!(validation.image_count, 1); } + #[test] + fn image_validation_rejects_malformed_or_oversized_multipart_boundaries() { + let body = Bytes::from_static(b"not-a-multipart-body"); + for content_type in [ + "multipart/form-data; boundary=bad boundary", + "multipart/form-data; boundary=bad\"quote", + "multipart/form-data; boundary=\"unterminated", + "multipart/form-data; boundary=aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + ] { + assert!(matches!( + parse_openai_image_validation_input( + OpenAiImageOperation::Generate, + Some(content_type), + &body, + ), + Err(OPENAI_IMAGE_INVALID_MULTIPART_DETAIL) + )); + } + } + + #[test] + fn multipart_parser_caps_part_count_and_header_size() { + let boundary = "bounded-parts"; + let mut accepted_body = Vec::new(); + for index in 0..MAX_MULTIPART_PARTS { + accepted_body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"field-{index}\"\r\n\r\nvalue\r\n" + ) + .as_bytes(), + ); + } + accepted_body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes()); + assert_eq!( + parse_multipart_fields(&accepted_body, boundary).len(), + MAX_MULTIPART_PARTS + ); + + let mut body = Vec::new(); + for index in 0..(MAX_MULTIPART_PARTS + 1) { + body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"field-{index}\"\r\n\r\nvalue\r\n" + ) + .as_bytes(), + ); + } + body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes()); + assert!(parse_multipart_fields(&body, boundary).is_empty()); + + let mut oversized_header = + format!("--{boundary}\r\nContent-Disposition: form-data; name=\"field\"; x=\"") + .into_bytes(); + oversized_header.extend(std::iter::repeat_n(b'x', MAX_MULTIPART_PART_HEADER_BYTES)); + oversized_header + .extend_from_slice(format!("\"\r\n\r\nvalue\r\n--{boundary}--\r\n").as_bytes()); + assert!(parse_multipart_fields(&oversized_header, boundary).is_empty()); + } + + #[test] + fn multipart_parser_preserves_boundary_like_payload_and_fails_closed() { + let boundary = "payload-boundary"; + let body = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"\r\n\r\n", + "prefix\r\n--{boundary}X\r\nsuffix--{boundary}\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + let fields = parse_multipart_fields(body.as_bytes(), boundary); + assert_eq!(fields.len(), 1); + assert_eq!(fields[0].name, "prompt"); + assert_eq!( + fields[0].data, + format!("prefix\r\n--{boundary}X\r\nsuffix--{boundary}").into_bytes() + ); + + // A valid first part must not make a truncated second part appear + // valid. The only marker after the second part has an invalid + // suffix and there is no closing boundary. + let malformed = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"first\"\r\n\r\n", + "ok\r\n", + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"second\"\r\n\r\n", + "truncated\r\n--{boundary}X\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(malformed.as_bytes(), boundary).is_empty()); + } + + #[test] + fn multipart_parser_does_not_extract_name_from_filename_and_rejects_duplicates() { + let boundary = "header-parameters"; + let filename_only = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; filename=\"name=\\\"prompt\\\"\"\r\n\r\n", + "attacker-value\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(filename_only.as_bytes(), boundary).is_empty()); + + let duplicate_name = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"; name=\"image\"\r\n\r\n", + "ambiguous-value\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(duplicate_name.as_bytes(), boundary).is_empty()); + } + + #[test] + fn multipart_parser_rejects_garbage_after_closing_boundary() { + let boundary = "closing-suffix"; + let body = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"\r\n\r\n", + "value\r\n", + "--{boundary}--junk" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(body.as_bytes(), boundary).is_empty()); + } + #[test] fn image_validation_applies_the_global_count_limit_before_model_mapping() { let openai_body = Bytes::from_static(br#"{"model":"gpt-image-2","prompt":"draw","n":2}"#); diff --git a/apps/aether-gateway/src/handlers/public/catalog_helpers.rs b/apps/aether-gateway/src/handlers/public/catalog_helpers.rs index 9fc0a5e10..d14722f0b 100644 --- a/apps/aether-gateway/src/handlers/public/catalog_helpers.rs +++ b/apps/aether-gateway/src/handlers/public/catalog_helpers.rs @@ -1,13 +1,12 @@ use crate::api::ai::public_api_format_local_path; -use crate::handlers::shared::{ - query_param_optional_bool, query_param_value, unix_ms_to_rfc3339, unix_secs_to_rfc3339, -}; +use crate::handlers::shared::{query_param_value, unix_ms_to_rfc3339, unix_secs_to_rfc3339}; use crate::provider_key_auth::{ provider_key_configured_api_formats, provider_key_effective_api_formats, }; use crate::AppState; use aether_data_contracts::repository::candidates::{ - PublicHealthTimelineBucket, RequestCandidateStatus, StoredRequestCandidate, + sanitize_request_candidate_error_type, PublicHealthTimelineBucket, RequestCandidateStatus, + StoredRequestCandidate, }; use aether_data_contracts::repository::global_models::{ PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, StoredPublicCatalogModel, @@ -123,34 +122,378 @@ pub(crate) fn normalize_admin_base_url(base_url: &str) -> Result if trimmed.is_empty() { return Err("base_url 不能为空".to_string()); } - let normalized = trimmed.trim_end_matches('/'); - let lower = normalized.to_ascii_lowercase(); - if !lower.starts_with("http://") && !lower.starts_with("https://") { + let mut parsed = url::Url::parse(trimmed).map_err(|_| "base_url 不是有效 URL".to_string())?; + if !matches!(parsed.scheme(), "http" | "https") { return Err("URL 必须以 http:// 或 https:// 开头".to_string()); } - Ok(normalized.to_string()) + if parsed.host_str().is_none() { + return Err("base_url 必须包含有效主机".to_string()); + } + if !aether_http::is_https_or_loopback_http_url(&parsed) { + return Err("base_url 必须使用 HTTPS;HTTP 仅允许字面量 loopback 主机".to_string()); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err("base_url 不允许包含用户名或密码".to_string()); + } + if parsed.query().is_some() || parsed.fragment().is_some() { + return Err("base_url 不允许包含查询参数或片段".to_string()); + } + parsed.set_query(None); + parsed.set_fragment(None); + let normalized = parsed.to_string(); + Ok(normalized.trim_end_matches('/').to_string()) +} + +#[cfg(test)] +mod normalize_admin_base_url_tests { + use super::normalize_admin_base_url; + + #[test] + fn endpoint_base_url_rejects_embedded_credentials_and_hidden_suffixes() { + for value in [ + "https://user:password@api.example.test/v1", + "https://api.example.test/v1?key=secret", + "https://api.example.test/v1#secret", + "http://api.example.test/v1", + "http://10.0.0.1/v1", + "http://[::ffff:127.0.0.1]/v1", + "https://", + "https://api.example.test:invalid/v1", + ] { + assert!(normalize_admin_base_url(value).is_err(), "accepted {value}"); + } + } + + #[test] + fn endpoint_base_url_is_parsed_and_normalized() { + assert_eq!( + normalize_admin_base_url(" HTTPS://API.EXAMPLE.TEST/v1/ ").expect("valid base URL"), + "https://api.example.test/v1" + ); + assert_eq!( + normalize_admin_base_url("http://127.0.0.1:18181/v1/") + .expect("literal loopback HTTP should remain available"), + "http://127.0.0.1:18181/v1" + ); + assert_eq!( + normalize_admin_base_url("http://[::1]:18181/v1/") + .expect("IPv6 loopback HTTP should remain available"), + "http://[::1]:18181/v1" + ); + } } pub(crate) fn sanitize_public_model_config_for_user( config: Option, ) -> Option { - let Some(mut config) = config else { + let Some(config) = config else { return None; }; - if let Some(object) = config.as_object_mut() { - for key in [ - "model_mappings", - "model_mapping", - "global_model_mappings", - "provider_model_mappings", - "provider_model_aliases", - "mapping_preview", - "model_mapping_preview", - ] { - object.remove(key); + let object = config.as_object()?; + let mut public = serde_json::Map::new(); + + if let Some(description) = object + .get("description") + .and_then(serde_json::Value::as_str) + { + public.insert( + "description".to_string(), + serde_json::Value::String(description.to_string()), + ); + } + for key in [ + "streaming", + "image_generation", + "vision", + "function_calling", + "extended_thinking", + "embedding", + "rerank", + ] { + if let Some(value) = object.get(key).and_then(serde_json::Value::as_bool) { + public.insert(key.to_string(), serde_json::Value::Bool(value)); } } - Some(config) + for key in ["model_type", "type"] { + if let Some(value) = object.get(key).and_then(serde_json::Value::as_str) { + public.insert( + key.to_string(), + serde_json::Value::String(value.to_string()), + ); + } + } + for key in ["api_formats", "capabilities", "supported_capabilities"] { + if let Some(values) = public_string_array(object.get(key)) { + public.insert(key.to_string(), serde_json::Value::Array(values)); + } + } + if let Some(video_billing) = public_video_billing(object.get("billing")) { + public.insert("billing".to_string(), video_billing); + } + + (!public.is_empty()).then(|| serde_json::Value::Object(public)) +} + +pub(crate) fn sanitize_public_model_capabilities( + capabilities: Option, +) -> Option { + public_string_array(capabilities.as_ref()).map(serde_json::Value::Array) +} + +pub(crate) fn sanitize_public_tiered_pricing( + pricing: Option, +) -> Option { + let pricing = pricing?.as_object()?.clone(); + let mut public = serde_json::Map::new(); + + if let Some(tiers) = pricing.get("tiers").and_then(serde_json::Value::as_array) { + let tiers = tiers + .iter() + .filter_map(public_pricing_tier) + .collect::>(); + if !tiers.is_empty() { + public.insert("tiers".to_string(), serde_json::Value::Array(tiers)); + } + } + copy_public_price(&pricing, &mut public, "image_output_price_default", false); + copy_public_price_matrix(&pricing, &mut public, "image_output_prices"); + copy_public_price_ranges(&pricing, &mut public, "image_output_price_ranges"); + + if let Some(processing_tiers) = pricing + .get("processing_tiers") + .and_then(serde_json::Value::as_object) + { + let processing_tiers = processing_tiers + .iter() + .filter_map(|(name, pricing)| { + let mut pricing_without_nested_tiers = pricing.as_object()?.clone(); + pricing_without_nested_tiers.remove("processing_tiers"); + let mut sanitized = sanitize_public_tiered_pricing(Some( + serde_json::Value::Object(pricing_without_nested_tiers), + )) + .unwrap_or_else(|| serde_json::Value::Object(serde_json::Map::new())); + let sanitized = sanitized.as_object_mut()?; + if let Some(multiplier) = pricing + .get("price_multiplier") + .and_then(nonnegative_finite_number) + { + sanitized.insert("price_multiplier".to_string(), multiplier); + } + (!sanitized.is_empty()) + .then(|| (name.clone(), serde_json::Value::Object(sanitized.clone()))) + }) + .collect::>(); + if !processing_tiers.is_empty() { + public.insert( + "processing_tiers".to_string(), + serde_json::Value::Object(processing_tiers), + ); + } + } + + (!public.is_empty()).then(|| serde_json::Value::Object(public)) +} + +fn public_string_array(value: Option<&serde_json::Value>) -> Option> { + let values = value?.as_array()?; + let values = values + .iter() + .filter_map(serde_json::Value::as_str) + .map(|value| serde_json::Value::String(value.to_string())) + .collect::>(); + (!values.is_empty()).then_some(values) +} + +fn public_pricing_tier(value: &serde_json::Value) -> Option { + let value = value.as_object()?; + let mut public = serde_json::Map::new(); + match value.get("up_to") { + Some(serde_json::Value::Null) => { + public.insert("up_to".to_string(), serde_json::Value::Null); + } + Some(value) => { + public.insert("up_to".to_string(), nonnegative_integer_number(value)?); + } + None => return None, + } + for key in [ + "input_price_per_1m", + "output_price_per_1m", + "cache_creation_price_per_1m", + "cache_read_price_per_1m", + ] { + copy_public_price(value, &mut public, key, false); + } + if let Some(cache_ttl_pricing) = value + .get("cache_ttl_pricing") + .and_then(serde_json::Value::as_array) + { + let cache_ttl_pricing = cache_ttl_pricing + .iter() + .filter_map(public_cache_ttl_price) + .collect::>(); + if !cache_ttl_pricing.is_empty() { + public.insert( + "cache_ttl_pricing".to_string(), + serde_json::Value::Array(cache_ttl_pricing), + ); + } + } + Some(serde_json::Value::Object(public)) +} + +fn public_cache_ttl_price(value: &serde_json::Value) -> Option { + let value = value.as_object()?; + let mut public = serde_json::Map::new(); + public.insert( + "ttl_minutes".to_string(), + nonnegative_integer_number(value.get("ttl_minutes")?)?, + ); + for key in ["cache_creation_price_per_1m", "cache_read_price_per_1m"] { + copy_public_price(value, &mut public, key, false); + } + Some(serde_json::Value::Object(public)) +} + +fn copy_public_price( + source: &serde_json::Map, + target: &mut serde_json::Map, + key: &str, + allow_null: bool, +) { + match source.get(key) { + Some(serde_json::Value::Null) if allow_null => { + target.insert(key.to_string(), serde_json::Value::Null); + } + Some(value) => { + if let Some(value) = nonnegative_finite_number(value) { + target.insert(key.to_string(), value); + } + } + None => {} + } +} + +fn copy_public_price_matrix( + source: &serde_json::Map, + target: &mut serde_json::Map, + key: &str, +) { + let Some(matrix) = source.get(key).and_then(serde_json::Value::as_object) else { + return; + }; + let matrix = matrix + .iter() + .filter_map(|(size, prices)| { + let prices = prices.as_object()?; + let prices = prices + .iter() + .filter_map(|(quality, price)| { + nonnegative_finite_number(price).map(|price| (quality.clone(), price)) + }) + .collect::>(); + (!prices.is_empty()).then(|| (size.clone(), serde_json::Value::Object(prices))) + }) + .collect::>(); + if !matrix.is_empty() { + target.insert(key.to_string(), serde_json::Value::Object(matrix)); + } +} + +fn copy_public_price_ranges( + source: &serde_json::Map, + target: &mut serde_json::Map, + key: &str, +) { + let Some(ranges) = source.get(key).and_then(serde_json::Value::as_array) else { + return; + }; + let ranges = ranges + .iter() + .filter_map(|range| { + let range = range.as_object()?; + let mut public = serde_json::Map::new(); + match range.get("up_to_pixels") { + Some(serde_json::Value::Null) => { + public.insert("up_to_pixels".to_string(), serde_json::Value::Null); + } + Some(value) => { + public.insert( + "up_to_pixels".to_string(), + nonnegative_integer_number(value)?, + ); + } + None => return None, + } + if let Some(label) = range.get("label").and_then(serde_json::Value::as_str) { + public.insert( + "label".to_string(), + serde_json::Value::String(label.to_string()), + ); + } + let prices = range.get("prices")?.as_object()?; + let prices = prices + .iter() + .filter_map(|(quality, price)| { + nonnegative_finite_number(price).map(|price| (quality.clone(), price)) + }) + .collect::>(); + if prices.is_empty() { + return None; + } + public.insert("prices".to_string(), serde_json::Value::Object(prices)); + Some(serde_json::Value::Object(public)) + }) + .collect::>(); + if !ranges.is_empty() { + target.insert(key.to_string(), serde_json::Value::Array(ranges)); + } +} + +fn nonnegative_integer_number(value: &serde_json::Value) -> Option { + value + .as_u64() + .map(serde_json::Number::from) + .map(serde_json::Value::Number) +} + +fn nonnegative_finite_number(value: &serde_json::Value) -> Option { + value + .as_f64() + .filter(|value| value.is_finite() && *value >= 0.0) + .and_then(serde_json::Number::from_f64) + .map(serde_json::Value::Number) +} + +fn public_video_billing(value: Option<&serde_json::Value>) -> Option { + let prices = value? + .get("video")? + .get("price_per_second_by_resolution")? + .as_object()?; + let prices = prices + .iter() + .filter_map(|(resolution, price)| { + price + .as_f64() + .filter(|price| price.is_finite() && *price >= 0.0) + .map(|price| { + ( + resolution.clone(), + serde_json::Number::from_f64(price) + .map(serde_json::Value::Number) + .expect("finite price should serialize"), + ) + }) + }) + .collect::>(); + if prices.is_empty() { + return None; + } + Some(json!({ + "video": { + "price_per_second_by_resolution": prices, + } + })) } pub(crate) fn admin_requested_force_stream(value: &serde_json::Value) -> bool { @@ -192,7 +535,6 @@ pub(crate) async fn build_public_providers_payload( return None; } - let is_active = query_param_optional_bool(query, "is_active"); let skip = query_param_value(query, "skip") .and_then(|value| value.parse::().ok()) .unwrap_or(0); @@ -201,15 +543,11 @@ pub(crate) async fn build_public_providers_payload( .filter(|value| *value > 0 && *value <= 1000) .unwrap_or(100); - let active_only = is_active.unwrap_or(true); let mut providers = state - .list_provider_catalog_providers(active_only) + .list_provider_catalog_providers(true) .await .ok() .unwrap_or_default(); - if matches!(is_active, Some(false)) { - providers.retain(|provider| !provider.is_active); - } providers.sort_by(|left, right| { left.provider_priority .cmp(&right.provider_priority) @@ -241,16 +579,14 @@ pub(crate) async fn build_public_providers_payload( let mut endpoints_count_by_provider = BTreeMap::::new(); let mut active_endpoints_count_by_provider = BTreeMap::::new(); let mut api_formats = BTreeSet::::new(); - for endpoint in &endpoints { + for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) { *endpoints_count_by_provider .entry(endpoint.provider_id.clone()) .or_default() += 1; - if endpoint.is_active { - *active_endpoints_count_by_provider - .entry(endpoint.provider_id.clone()) - .or_default() += 1; - api_formats.insert(endpoint.api_format.clone()); - } + *active_endpoints_count_by_provider + .entry(endpoint.provider_id.clone()) + .or_default() += 1; + api_formats.insert(endpoint.api_format.clone()); } let mut models_by_provider = BTreeMap::>::new(); @@ -621,19 +957,7 @@ pub(crate) async fn build_api_format_health_monitor_payload( let avg_first_byte_ms = model_health_average_first_byte_ms(&usage_events); let events = attempts .into_iter() - .filter_map(|candidate| { - let timestamp = candidate - .finished_at_unix_ms - .or(candidate.started_at_unix_ms) - .unwrap_or(candidate.created_at_unix_ms); - Some(json!({ - "timestamp": unix_ms_to_rfc3339(timestamp)?, - "status": request_candidate_status_label(candidate.status), - "status_code": candidate.status_code, - "latency_ms": candidate.latency_ms, - "error_type": candidate.error_type, - })) - }) + .filter_map(public_request_candidate_health_event) .collect::>(); let empty_timeline = BTreeMap::new(); let timeline_source = timeline_by_format @@ -688,6 +1012,22 @@ pub(crate) async fn build_api_format_health_monitor_payload( })) } +fn public_request_candidate_health_event( + candidate: StoredRequestCandidate, +) -> Option { + let timestamp = candidate + .finished_at_unix_ms + .or(candidate.started_at_unix_ms) + .unwrap_or(candidate.created_at_unix_ms); + Some(json!({ + "timestamp": unix_ms_to_rfc3339(timestamp)?, + "status": request_candidate_status_label(candidate.status), + "status_code": candidate.status_code, + "latency_ms": candidate.latency_ms, + "error_type": sanitize_request_candidate_error_type(candidate.error_type), + })) +} + pub(crate) async fn build_model_health_monitor_payload( state: &AppState, lookback_hours: u64, @@ -1980,7 +2320,11 @@ pub(crate) fn api_format_display_name(api_format: &str) -> String { #[cfg(test)] mod tests { - use super::request_candidate_event_unix_ms; + use super::{ + public_request_candidate_health_event, request_candidate_event_unix_ms, + sanitize_public_model_capabilities, sanitize_public_model_config_for_user, + sanitize_public_tiered_pricing, + }; use crate::handlers::shared::unix_ms_to_rfc3339; use aether_data_contracts::repository::candidates::{ RequestCandidateStatus, StoredRequestCandidate, @@ -2023,4 +2367,114 @@ mod tests { Some("2023-11-14T22:13:20.123Z") ); } + + #[test] + fn public_candidate_health_event_classifies_untrusted_error_text() { + let mut candidate = StoredRequestCandidate::new( + "cand-unsafe".to_string(), + "req-1".to_string(), + None, + None, + None, + None, + 0, + 0, + Some("provider-1".to_string()), + Some("endpoint-1".to_string()), + Some("key-1".to_string()), + RequestCandidateStatus::Failed, + None, + false, + Some(500), + None, + None, + Some(42), + Some(1), + None, + None, + 1_700_000_000_000, + Some(1_700_000_000_111), + Some(1_700_000_000_123), + ) + .expect("candidate should build"); + candidate.error_type = Some("Bearer public-health-secret".to_string()); + + let event = public_request_candidate_health_event(candidate) + .expect("candidate health event should build"); + + assert_eq!(event["error_type"], "unclassified_error"); + assert!(!event.to_string().contains("public-health-secret")); + } + + #[test] + fn public_model_metadata_uses_typed_allowlists() { + let config = sanitize_public_model_config_for_user(Some(serde_json::json!({ + "description": "Public model", + "streaming": true, + "api_formats": ["openai:responses", {"secret": "nested"}], + "client_secret": "hidden", + "billing": { + "video": { + "price_per_second_by_resolution": { + "720p": 0.12, + "internal": "hidden" + }, + "private_key": "hidden" + } + } + }))) + .expect("public config should remain"); + assert_eq!(config["description"], "Public model"); + assert_eq!(config["streaming"], true); + assert_eq!( + config["api_formats"], + serde_json::json!(["openai:responses"]) + ); + assert_eq!( + config["billing"]["video"]["price_per_second_by_resolution"], + serde_json::json!({"720p": 0.12}) + ); + assert!(config.get("client_secret").is_none()); + assert!(config["billing"]["video"].get("private_key").is_none()); + + assert_eq!( + sanitize_public_model_capabilities(Some(serde_json::json!([ + "vision", + {"secret": "hidden"} + ]))), + Some(serde_json::json!(["vision"])) + ); + + let pricing = sanitize_public_tiered_pricing(Some(serde_json::json!({ + "tiers": [{ + "up_to": null, + "input_price_per_1m": 3.0, + "output_price_per_1m": 15.0, + "internal_note": "hidden", + "cache_ttl_pricing": [{ + "ttl_minutes": 60, + "cache_creation_price_per_1m": 4.0, + "secret": "hidden" + }] + }], + "processing_tiers": { + "priority": { + "price_multiplier": 1.5, + "private_note": "hidden" + } + }, + "internal_pricing": {"secret": "hidden"} + }))) + .expect("public pricing should remain"); + assert_eq!(pricing["tiers"][0]["input_price_per_1m"], 3.0); + assert!(pricing.get("internal_pricing").is_none()); + assert!(pricing["tiers"][0].get("internal_note").is_none()); + assert!(pricing["tiers"][0]["cache_ttl_pricing"][0] + .get("secret") + .is_none()); + assert_eq!( + pricing["processing_tiers"]["priority"], + serde_json::json!({"price_multiplier": 1.5}) + ); + } } diff --git a/apps/aether-gateway/src/handlers/public/mod.rs b/apps/aether-gateway/src/handlers/public/mod.rs index d666a3dd9..4c304a597 100644 --- a/apps/aether-gateway/src/handlers/public/mod.rs +++ b/apps/aether-gateway/src/handlers/public/mod.rs @@ -13,8 +13,9 @@ pub(crate) use self::catalog_helpers::{ build_public_health_timeline, build_public_health_timeline_details, build_public_providers_payload, build_related_health_monitor_payload, normalize_admin_base_url, provider_key_api_formats, request_candidate_event_unix_ms, request_candidate_status_label, - sanitize_public_model_config_for_user, ApiFormatHealthMonitorOptions, - HealthMonitorRelationDimension, ModelHealthMonitorOptions, + sanitize_public_model_capabilities, sanitize_public_model_config_for_user, + sanitize_public_tiered_pricing, ApiFormatHealthMonitorOptions, HealthMonitorRelationDimension, + ModelHealthMonitorOptions, }; pub(crate) use self::system_modules_helpers::{ build_admin_keys_grouped_by_format_payload, build_public_auth_modules_status_payload, @@ -28,5 +29,5 @@ pub(crate) use self::support::{ build_api_key_install_session_response, build_proxy_node_install_session_response, build_unhandled_public_support_response, matches_model_mapping_for_models, maybe_build_local_admin_announcements_response, maybe_build_local_public_support_response, - CreateApiKeyInstallSessionRequest, + vscodex_ws_proxy, CreateApiKeyInstallSessionRequest, }; diff --git a/apps/aether-gateway/src/handlers/public/support.rs b/apps/aether-gateway/src/handlers/public/support.rs index 4a14e2819..ac38a5040 100644 --- a/apps/aether-gateway/src/handlers/public/support.rs +++ b/apps/aether-gateway/src/handlers/public/support.rs @@ -3,17 +3,19 @@ use super::{ build_public_auth_modules_status_payload, build_public_catalog_models_payload, build_public_catalog_search_models_payload, build_public_providers_payload, build_related_health_monitor_payload, capability_detail_by_name, ldap_module_config_is_valid, - sanitize_public_model_config_for_user, serialize_public_capability, supported_capability_names, + sanitize_public_model_capabilities, sanitize_public_model_config_for_user, + sanitize_public_tiered_pricing, serialize_public_capability, supported_capability_names, ApiFormatHealthMonitorOptions, HealthMonitorRelationDimension, ModelHealthMonitorOptions, PUBLIC_CAPABILITY_DEFINITIONS, }; use crate::control::GatewayPublicRequestContext; use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, - escape_admin_email_template_html, module_available_from_env, query_param_bool, - query_param_optional_bool, query_param_value, read_admin_email_template_payload, - render_admin_email_template_html, system_config_bool, system_config_string, - unix_secs_to_rfc3339, + decrypt_catalog_secret_with_fallbacks, decrypt_or_migrate_auth_api_key_secret, + decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_system_config_secret, + encrypt_catalog_secret_with_fallbacks, escape_admin_email_template_html, + module_available_from_env, query_param_bool, query_param_optional_bool, query_param_value, + read_admin_email_template_payload, render_admin_email_template_html, system_config_bool, + system_config_string, unix_secs_to_rfc3339, }; use crate::{AppState, GatewayError}; use aether_data_contracts::repository::global_models::PublicGlobalModelQuery; @@ -48,6 +50,8 @@ mod support_payment; mod support_test_connection; #[path = "support/user_me.rs"] mod support_user_me; +#[path = "support/user_me_vscodex.rs"] +mod support_vscodex; #[path = "support/wallet.rs"] mod support_wallet; @@ -66,9 +70,11 @@ use self::support_auth::auth_session::{ build_auth_wallet_summary_payload, handle_auth_me, resolve_authenticated_local_user, AuthenticatedLocalUserContext, }; +use self::support_auth::{auth_email_is_verified, consume_auth_email_registration_proof}; use self::support_auth::{ - build_auth_error_response, build_auth_json_response, build_auth_registration_settings_payload, - build_auth_settings_payload, extract_client_device_id, maybe_build_local_auth_response, + build_auth_error_response, build_auth_json_response, build_auth_refresh_cookie_clear_header, + build_auth_registration_settings_payload, build_auth_settings_payload, + extract_client_device_id, mark_sensitive_response_no_store, maybe_build_local_auth_response, }; use self::support_billing::maybe_build_local_billing_response; use self::support_ccswitch::maybe_build_local_ccswitch_response; @@ -89,12 +95,16 @@ use self::support_oauth::maybe_build_local_oauth_response; use self::support_payment::maybe_build_local_payment_callback_response; use self::support_test_connection::maybe_build_local_test_connection_response; use self::support_user_me::maybe_build_local_users_me_response; +pub(crate) use self::support_vscodex::vscodex_ws_proxy; +use self::support_vscodex::{handle_users_me_vscodex_request, maybe_build_local_vscodex_response}; use self::support_wallet::{ build_wallet_balance_payload_for_auth_scope, build_wallet_balance_payload_for_user, build_wallet_live_today_usage_payload_for_api_key, build_wallet_live_today_usage_payload_for_user, direct_gateway_channels, - maybe_build_local_wallet_response, sanitize_wallet_gateway_response, - wallet_normalize_optional_string_field, + maybe_build_local_wallet_response, prepare_billing_gateway_response_for_storage, + prepare_wallet_gateway_response_for_storage, resolve_direct_gateway_channel, + sanitize_wallet_gateway_response, wallet_normalize_optional_string_field, + wallet_payment_instructions_from_checkout, wallet_payment_instructions_from_stored, }; pub(crate) fn build_unhandled_public_support_response( @@ -120,7 +130,8 @@ pub(crate) async fn maybe_build_local_public_support_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + remote_addr: &std::net::SocketAddr, + client_ip: std::net::IpAddr, request_body: Option<&Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -133,15 +144,21 @@ pub(crate) async fn maybe_build_local_public_support_response( state, request_context, headers, - cf_connecting_ip, + client_ip, request_body, ) .await; } if decision.route_family.as_deref() == Some("oauth") { - return maybe_build_local_oauth_response(state, request_context, headers, request_body) - .await; + return maybe_build_local_oauth_response( + state, + request_context, + headers, + client_ip, + request_body, + ) + .await; } if decision.route_family.as_deref() == Some("dashboard") { @@ -163,8 +180,14 @@ pub(crate) async fn maybe_build_local_public_support_response( } if decision.route_family.as_deref() == Some("wallet") { - if let Some(response) = - maybe_build_local_wallet_response(state, request_context, headers, request_body).await + if let Some(response) = maybe_build_local_wallet_response( + state, + request_context, + headers, + client_ip, + request_body, + ) + .await { return Some(response); } @@ -172,8 +195,14 @@ pub(crate) async fn maybe_build_local_public_support_response( } if decision.route_family.as_deref() == Some("billing") { - if let Some(response) = - maybe_build_local_billing_response(state, request_context, headers, request_body).await + if let Some(response) = maybe_build_local_billing_response( + state, + request_context, + headers, + client_ip, + request_body, + ) + .await { return Some(response); } @@ -188,7 +217,18 @@ pub(crate) async fn maybe_build_local_public_support_response( } if decision.route_family.as_deref() == Some("users_me") { - return maybe_build_local_users_me_response(state, request_context, headers, request_body) + return maybe_build_local_users_me_response( + state, + request_context, + headers, + remote_addr, + request_body, + ) + .await; + } + + if decision.route_family.as_deref() == Some("vscodex") { + return maybe_build_local_vscodex_response(state, request_context, client_ip, request_body) .await; } @@ -369,10 +409,6 @@ pub(crate) async fn maybe_build_local_public_support_response( .and_then(|value| value.parse::().ok()) .filter(|value| *value > 0 && *value <= 1000) .unwrap_or(100); - let is_active = query_param_optional_bool( - request_context.request_query_string.as_deref(), - "is_active", - ); let search = query_param_value(request_context.request_query_string.as_deref(), "search"); @@ -380,7 +416,7 @@ pub(crate) async fn maybe_build_local_public_support_response( .list_public_global_models(&PublicGlobalModelQuery { offset: skip, limit, - is_active, + is_active: Some(true), search, }) .await @@ -396,8 +432,8 @@ pub(crate) async fn maybe_build_local_public_support_response( "display_name": model.display_name, "is_active": model.is_active, "default_price_per_request": model.default_price_per_request, - "default_tiered_pricing": model.default_tiered_pricing, - "supported_capabilities": model.supported_capabilities, + "default_tiered_pricing": sanitize_public_tiered_pricing(model.default_tiered_pricing), + "supported_capabilities": sanitize_public_model_capabilities(model.supported_capabilities), "config": sanitize_public_model_config_for_user(model.config), "usage_count": model.usage_count, }) @@ -743,16 +779,11 @@ pub(crate) async fn maybe_build_local_public_support_response( "include_endpoints", false, ); - let active_only = query_param_bool( - request_context.request_query_string.as_deref(), - "active_only", - true, - ); if include_models { return None; } let providers = state - .list_provider_catalog_providers(active_only) + .list_provider_catalog_providers(true) .await .ok() .unwrap_or_default(); @@ -784,7 +815,10 @@ pub(crate) async fn maybe_build_local_public_support_response( payload["endpoints"] = serde_json::Value::Array( endpoints .iter() - .filter(|endpoint| endpoint.provider_id == provider_id) + .filter(|endpoint| { + endpoint.provider_id == provider_id + && endpoint.is_active + }) .map(|endpoint| json!({ "id": endpoint.id, "api_format": endpoint.api_format, @@ -833,7 +867,7 @@ pub(crate) async fn maybe_build_local_public_support_response( None }; let provider = match provider { - Some(provider) => provider, + Some(provider) if provider.is_active => provider, None => { return Some( ( @@ -843,6 +877,15 @@ pub(crate) async fn maybe_build_local_public_support_response( .into_response(), ); } + Some(_) => { + return Some( + ( + http::StatusCode::NOT_FOUND, + Json(json!({ "detail": "Provider not found" })), + ) + .into_response(), + ); + } }; let provider_id = provider.id.clone(); let mut payload = json!({ @@ -861,6 +904,7 @@ pub(crate) async fn maybe_build_local_public_support_response( payload["endpoints"] = serde_json::Value::Array( endpoints .into_iter() + .filter(|endpoint| endpoint.is_active) .map(|endpoint| { json!({ "id": endpoint.id, @@ -877,7 +921,8 @@ pub(crate) async fn maybe_build_local_public_support_response( if decision.route_kind.as_deref() == Some("test_connection") && request_context.request_path == "/v1/test-connection" { - return maybe_build_local_test_connection_response(state, request_context).await; + return maybe_build_local_test_connection_response(state, request_context, headers) + .await; } if decision.route_kind.as_deref() == Some("test_connection") diff --git a/apps/aether-gateway/src/handlers/public/support/announcements/public_routes.rs b/apps/aether-gateway/src/handlers/public/support/announcements/public_routes.rs index 0b762569d..388ed8238 100644 --- a/apps/aether-gateway/src/handlers/public/support/announcements/public_routes.rs +++ b/apps/aether-gateway/src/handlers/public/support/announcements/public_routes.rs @@ -10,7 +10,7 @@ use crate::AppState; use super::super::build_unhandled_public_support_response; use super::announcements_shared::{ - announcements_bad_request_response, announcements_internal_detail, + announcement_is_public_at, announcements_bad_request_response, announcements_internal_detail, announcements_internal_error_response, announcements_not_found_response, build_public_announcement_list_payload, build_public_announcement_payload, parse_public_announcements_query, public_announcement_id_from_path, @@ -37,7 +37,7 @@ pub(crate) async fn maybe_build_local_public_announcements_response( "/api/announcements" | "/api/announcements/" ) => { - let query = match parse_public_announcements_query( + let mut query = match parse_public_announcements_query( request_context.request_query_string.as_deref(), true, 50, @@ -45,6 +45,10 @@ pub(crate) async fn maybe_build_local_public_announcements_response( Ok(value) => value, Err(detail) => return Some(announcements_bad_request_response(detail)), }; + // This route is anonymous. Drafts and scheduled announcements are + // only visible through the authenticated admin surface. + query.active_only = true; + query.now_unix_secs = Some(chrono::Utc::now().timestamp().max(0) as u64); let page = match state.list_announcements(&query).await { Ok(value) => value, Err(err) => { @@ -93,6 +97,10 @@ pub(crate) async fn maybe_build_local_public_announcements_response( )) } }; + let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64; + if !announcement_is_public_at(&announcement, now_unix_secs) { + return Some(announcements_not_found_response()); + } Some(Json(build_public_announcement_payload(&announcement)).into_response()) } _ => Some(build_unhandled_public_support_response(request_context)), diff --git a/apps/aether-gateway/src/handlers/public/support/announcements/shared.rs b/apps/aether-gateway/src/handlers/public/support/announcements/shared.rs index 548a768b7..cfdd4454a 100644 --- a/apps/aether-gateway/src/handlers/public/support/announcements/shared.rs +++ b/apps/aether-gateway/src/handlers/public/support/announcements/shared.rs @@ -13,6 +13,19 @@ use aether_data::repository::announcements::{ use crate::handlers::shared::{query_param_optional_bool, query_param_value}; use crate::GatewayError; +pub(super) fn announcement_is_public_at( + announcement: &StoredAnnouncement, + now_unix_secs: u64, +) -> bool { + announcement.is_active + && announcement + .start_time_unix_secs + .is_none_or(|value| value <= now_unix_secs) + && announcement + .end_time_unix_secs + .is_none_or(|value| value >= now_unix_secs) +} + pub(super) fn parse_public_announcements_query( query: Option<&str>, default_active_only: bool, @@ -128,9 +141,10 @@ pub(super) fn announcements_not_found_response() -> Response { } pub(super) fn announcements_internal_error_response(detail: impl Into) -> Response { + let _ = detail.into(); ( http::StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({ "detail": detail.into() })), + Json(json!({ "detail": "Service temporarily unavailable" })), ) .into_response() } diff --git a/apps/aether-gateway/src/handlers/public/support/announcements/user_routes.rs b/apps/aether-gateway/src/handlers/public/support/announcements/user_routes.rs index 06e54dee6..07daf760d 100644 --- a/apps/aether-gateway/src/handlers/public/support/announcements/user_routes.rs +++ b/apps/aether-gateway/src/handlers/public/support/announcements/user_routes.rs @@ -11,7 +11,7 @@ use crate::AppState; use super::super::{build_unhandled_public_support_response, resolve_authenticated_local_user}; use super::announcements_shared::{ - announcements_bad_request_response, announcements_internal_detail, + announcement_is_public_at, announcements_bad_request_response, announcements_internal_detail, announcements_internal_error_response, announcements_not_found_response, build_public_announcement_payload, read_status_announcement_id_from_path, }; @@ -136,7 +136,9 @@ pub(crate) async fn maybe_build_local_announcement_user_response( None => return Some(build_unhandled_public_support_response(request_context)), }; match state.find_announcement_by_id(announcement_id).await { - Ok(Some(_)) => {} + Ok(Some(announcement)) + if announcement_is_public_at(&announcement, now_unix_secs) => {} + Ok(Some(_)) => return Some(announcements_not_found_response()), Ok(None) => return Some(announcements_not_found_response()), Err(err) => { return Some(announcements_internal_error_response( diff --git a/apps/aether-gateway/src/handlers/public/support/auth.rs b/apps/aether-gateway/src/handlers/public/support/auth.rs index cdcdcdf42..9a3677087 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth.rs @@ -1,5 +1,6 @@ pub(super) use super::{ build_unhandled_public_support_response, decrypt_catalog_secret_with_fallbacks, + decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_system_config_secret, escape_admin_email_template_html, ldap_module_config_is_valid, module_available_from_env, read_admin_email_template_payload, render_admin_email_template_html, system_config_bool, system_config_string, AppState, GatewayError, GatewayPublicRequestContext, @@ -14,10 +15,17 @@ pub(super) use regex::Regex; use serde::Deserialize; pub(super) use serde_json::json; +const AUTH_DUMMY_PASSWORD_HASH: &str = + "$2y$10$.OBQfixAECpsb8V/VS3csOMf00x2E/jD/gnud20t6RG0yiQosyOZ2"; + #[path = "auth_helpers.rs"] mod auth_helpers; pub(crate) use auth_helpers::*; +#[path = "auth_rate_limit.rs"] +mod auth_rate_limit; +use auth_rate_limit::*; + #[path = "auth_turnstile.rs"] mod auth_turnstile; use auth_turnstile::*; @@ -25,11 +33,21 @@ use auth_turnstile::*; #[path = "auth_email.rs"] mod auth_email; use auth_email::*; +pub(in crate::handlers::public::support) use auth_email::{ + auth_email_is_verified, consume_auth_email_registration_proof, +}; #[path = "auth_ldap.rs"] mod auth_ldap; use auth_ldap::*; +pub(super) async fn local_password_login_allowed_for_user( + state: &AppState, + user: Option<&aether_data::repository::users::StoredUserAuthRecord>, +) -> Result { + auth_ldap::auth_local_login_allowed_for_user(state, user).await +} + #[path = "auth_session.rs"] pub(super) mod auth_session; use auth_session::*; @@ -38,7 +56,7 @@ use auth_session::*; pub(super) mod auth_registration; use auth_registration::*; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct AuthLoginRequest { email: String, password: String, @@ -87,12 +105,41 @@ fn system_config_string_list(value: Option<&serde_json::Value>) -> Vec { _ => Vec::new(), } } + +pub(super) fn build_auth_internal_error_response( + event_name: &'static str, + _error: impl std::fmt::Debug, + clear_cookie: bool, +) -> Response { + tracing::error!(event_name = event_name, "authentication operation failed"); + build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "认证服务暂不可用,请稍后重试", + clear_cookie, + ) +} + +async fn verify_auth_password( + password: &str, + password_hash: &str, +) -> Result { + let password = password.to_string(); + let password_hash = password_hash.to_string(); + tokio::task::spawn_blocking(move || bcrypt::verify(password, &password_hash).unwrap_or(false)) + .await +} + async fn handle_auth_login( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { + if let Err(response) = enforce_auth_ip_rate_limit(state, AUTH_LOGIN_RATE_LIMIT, client_ip).await + { + return response; + } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少登录请求体", false); }; @@ -114,17 +161,87 @@ async fn handle_auth_login( false, ); } + if let Err(response) = + enforce_auth_identity_rate_limit(state, AUTH_LOGIN_RATE_LIMIT, &identifier).await + { + return response; + } if let Err(detail) = validate_auth_login_password(&payload.password) { return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } - let client_device_id = match extract_client_device_id(request_context, headers) { + // Login creates a cookie-backed session. Requiring a custom header prevents + // cross-site forms from logging a victim's browser into an attacker's account. + let client_device_id = match extract_client_device_id_header(headers) { Ok(value) => value, Err(response) => return response, }; let auth_type = payload.auth_type.trim().to_ascii_lowercase(); - let user = match auth_type.as_str() { + let (user, expected_password_hash) = match auth_type.as_str() { "local" => { let user = match state.find_user_auth_by_identifier(&identifier).await { + Ok(user) => user, + Err(err) => { + return build_auth_internal_error_response( + "auth_login_user_lookup_failed", + err, + false, + ) + } + }; + let basic_user_is_eligible = user.as_ref().is_some_and(|user| { + !user.is_deleted + && user.is_active + && user.auth_source.eq_ignore_ascii_case("local") + && user + .password_hash + .as_deref() + .is_some_and(|hash| !hash.is_empty()) + }); + let local_login_is_allowed = + match auth_local_login_allowed_for_user(state, user.as_ref()).await { + Ok(allowed) => allowed, + Err(err) => { + return build_auth_internal_error_response( + "auth_login_settings_lookup_failed", + err, + false, + ) + } + }; + let user_is_eligible = basic_user_is_eligible && local_login_is_allowed; + let password_hash = user + .as_ref() + .filter(|_| user_is_eligible) + .and_then(|user| user.password_hash.as_deref()) + .unwrap_or(AUTH_DUMMY_PASSWORD_HASH) + .to_string(); + let password_matches = + match verify_auth_password(&payload.password, &password_hash).await { + Ok(matches) => matches, + Err(err) => { + return build_auth_internal_error_response( + "auth_login_password_worker_failed", + err, + false, + ) + } + }; + if !user_is_eligible || !password_matches { + return build_auth_error_response( + http::StatusCode::UNAUTHORIZED, + "邮箱或密码错误", + false, + ); + } + ( + user.expect("eligible authenticated local user should exist"), + Some(password_hash), + ) + } + "ldap" => { + let ldap_user = match authenticate_auth_ldap_user(state, &identifier, &payload.password) + .await + { Ok(Some(user)) => user, Ok(None) => { return build_auth_error_response( @@ -134,75 +251,9 @@ async fn handle_auth_login( ) } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth user lookup failed: {err:?}"), - false, - ) + return build_auth_internal_error_response("auth_ldap_login_failed", err, false) } }; - if user.is_deleted - || !user.is_active - || !user.auth_source.eq_ignore_ascii_case("local") - || user.password_hash.as_deref().is_none_or(str::is_empty) - { - return build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "邮箱或密码错误", - false, - ); - } - match auth_local_login_allowed_for_user(state, &user).await { - Ok(true) => {} - Ok(false) => { - return build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "邮箱或密码错误", - false, - ) - } - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, - ) - } - } - let password_hash = user - .password_hash - .as_deref() - .expect("validated password hash should exist"); - let password_matches = - bcrypt::verify(&payload.password, password_hash).unwrap_or(false); - if !password_matches { - return build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "邮箱或密码错误", - false, - ); - } - user - } - "ldap" => { - let ldap_user = - match authenticate_auth_ldap_user(state, &identifier, &payload.password).await { - Ok(Some(user)) => user, - Ok(None) => { - return build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "邮箱或密码错误", - false, - ) - } - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth ldap login failed: {err:?}"), - false, - ) - } - }; let _ = &ldap_user.display_name; let initial_gift = match state .read_system_config_json_value("default_user_initial_gift_usd") @@ -210,14 +261,14 @@ async fn handle_auth_login( { Ok(value) => system_config_f64(value.as_ref(), 10.0), Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), + return build_auth_internal_error_response( + "auth_ldap_settings_lookup_failed", + err, false, ) } }; - match state + let user = match state .get_or_create_ldap_auth_user( ldap_user.email, ldap_user.username, @@ -238,13 +289,14 @@ async fn handle_auth_login( ) } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth ldap user sync failed: {err:?}"), + return build_auth_internal_error_response( + "auth_ldap_user_sync_failed", + err, false, ) } - } + }; + (user, None) } _ => { return build_auth_error_response( @@ -255,14 +307,22 @@ async fn handle_auth_login( } }; - build_auth_login_success_response(state, headers, client_device_id, user).await + build_auth_login_success_response( + state, + headers, + client_ip, + client_device_id, + user, + expected_password_hash.as_deref(), + ) + .await } pub(super) async fn maybe_build_local_auth_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + client_ip: std::net::IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -274,24 +334,21 @@ pub(super) async fn maybe_build_local_auth_response( Some("send_verification_code") if request_context.request_path == "/api/auth/send-verification-code" => { - Some( - handle_auth_send_verification_code(state, headers, cf_connecting_ip, request_body) - .await, - ) + Some(handle_auth_send_verification_code(state, headers, client_ip, request_body).await) } Some("login") if request_context.request_path == "/api/auth/login" => { - Some(handle_auth_login(state, request_context, headers, request_body).await) + Some(handle_auth_login(state, request_context, headers, client_ip, request_body).await) } Some("register") if request_context.request_path == "/api/auth/register" => { - Some(handle_auth_register(state, headers, cf_connecting_ip, request_body).await) + Some(handle_auth_register(state, headers, client_ip, request_body).await) } Some("verify_email") if request_context.request_path == "/api/auth/verify-email" => { - Some(handle_auth_verify_email(state, request_body).await) + Some(handle_auth_verify_email(state, client_ip, request_body).await) } Some("verification_status") if request_context.request_path == "/api/auth/verification-status" => { - Some(handle_auth_verification_status(state, request_body).await) + Some(handle_auth_verification_status(state, client_ip, request_body).await) } Some("me") if request_context.request_path == "/api/auth/me" => { Some(handle_auth_me(state, request_context, headers).await) @@ -308,10 +365,43 @@ pub(super) async fn maybe_build_local_auth_response( #[cfg(test)] mod tests { - use super::{maybe_build_local_auth_response, AppState, GatewayPublicRequestContext}; + use super::{ + build_auth_internal_error_response, maybe_build_local_auth_response, verify_auth_password, + AppState, GatewayPublicRequestContext, AUTH_DUMMY_PASSWORD_HASH, + }; use crate::control::GatewayControlDecision; use axum::body::to_bytes; use axum::http::{HeaderMap, Method, StatusCode, Uri}; + use std::net::{IpAddr, Ipv4Addr}; + + #[tokio::test] + async fn auth_dummy_password_hash_has_fixed_bcrypt_work() { + assert!(verify_auth_password("secret123", AUTH_DUMMY_PASSWORD_HASH) + .await + .expect("password worker should complete")); + assert!( + !verify_auth_password("incorrect-password", AUTH_DUMMY_PASSWORD_HASH) + .await + .expect("password worker should complete") + ); + } + + #[tokio::test] + async fn auth_internal_error_response_does_not_expose_error_detail() { + let response = build_auth_internal_error_response( + "auth_internal_error_response_test", + "database secret diagnostic", + false, + ); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + 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!(payload["detail"], "服务暂不可用,请稍后重试"); + assert!(!String::from_utf8_lossy(&body).contains("database secret diagnostic")); + } fn request_context(method: Method, uri: &str, route_kind: &str) -> GatewayPublicRequestContext { GatewayPublicRequestContext::from_request_parts( @@ -337,7 +427,7 @@ mod tests { &state, &request_context, &HeaderMap::new(), - None, + IpAddr::V4(Ipv4Addr::LOCALHOST), None, ) .await diff --git a/apps/aether-gateway/src/handlers/public/support/auth_email.rs b/apps/aether-gateway/src/handlers/public/support/auth_email.rs index f7cfdc288..594c981e2 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_email.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_email.rs @@ -6,22 +6,47 @@ use super::{ use crate::email_delivery::{ read_smtp_delivery_config, send_smtp_email, ComposedEmail, SmtpDeliveryConfig, }; +use aether_admin::system::{ + admin_email_template_subject_is_valid, ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES, + ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES, +}; +use hmac::Mac; +use sha2::{Digest, Sha256}; -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[derive(Clone, serde::Serialize, serde::Deserialize)] pub(super) struct StoredAuthEmailVerificationCode { - pub(super) code: String, + pub(super) code_hash: String, pub(super) created_at: String, + pub(super) verification_token_hash: String, } pub(super) type AuthSmtpConfig = SmtpDeliveryConfig; pub(super) type AuthComposedEmail = ComposedEmail; -pub(super) fn auth_email_verification_key(email: &str) -> String { - format!("{AUTH_EMAIL_VERIFICATION_PREFIX}{email}") +fn auth_email_storage_key_digest(domain: &str, parts: &[&str]) -> Result { + let secret = super::auth_jwt_secret().map_err(GatewayError::Internal)?; + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .map_err(|_| GatewayError::Internal("auth email HMAC key invalid".to_string()))?; + mac.update(b"aether-auth-email-storage-v1\0"); + mac.update(domain.as_bytes()); + for part in parts { + mac.update(b"\0"); + mac.update(part.as_bytes()); + } + Ok(mac + .finalize() + .into_bytes() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect()) } -pub(super) fn auth_email_verified_key(email: &str) -> String { - format!("{AUTH_EMAIL_VERIFIED_PREFIX}{email}") +pub(super) fn auth_email_verification_key(email: &str) -> Result { + let email = email.trim().to_ascii_lowercase(); + Ok(format!( + "{AUTH_EMAIL_VERIFICATION_PREFIX}{}", + auth_email_storage_key_digest("pending", &[email.as_str()])? + )) } pub(super) fn record_auth_email_delivery_for_tests( @@ -46,13 +71,98 @@ pub(super) fn generate_auth_verification_code() -> String { format!("{:06}", uuid::Uuid::new_v4().as_u128() % 1_000_000) } +pub(super) fn generate_auth_verification_token() -> String { + format!( + "{}{}", + uuid::Uuid::new_v4().simple(), + uuid::Uuid::new_v4().simple() + ) +} + +pub(super) fn auth_verification_token_hash(token: &str) -> String { + format!("{:x}", Sha256::digest(token.trim().as_bytes())) +} + +pub(super) fn auth_verification_code_hash(verification_token: &str, code: &str) -> String { + format!( + "{:x}", + Sha256::digest( + format!( + "aether-email-verification\0{}\0{}", + verification_token.trim(), + code.trim() + ) + .as_bytes() + ) + ) +} + +fn constant_time_eq(left: &str, right: &str) -> bool { + if left.len() != right.len() { + return false; + } + left.as_bytes() + .iter() + .zip(right.as_bytes()) + .fold(0u8, |difference, (left, right)| difference | (left ^ right)) + == 0 +} + +pub(super) fn auth_email_registration_proof_key( + email: &str, + verification_token: &str, +) -> Result { + let email = email.trim().to_ascii_lowercase(); + Ok(format!( + "{AUTH_EMAIL_VERIFIED_PREFIX}{}", + auth_email_storage_key_digest( + "registration-proof", + &[email.as_str(), verification_token.trim()] + )? + )) +} + +pub(super) fn auth_verification_token_matches( + stored: &StoredAuthEmailVerificationCode, + verification_token: &str, +) -> bool { + !verification_token.trim().is_empty() + && constant_time_eq( + &stored.verification_token_hash, + &auth_verification_token_hash(verification_token), + ) +} + +pub(super) fn auth_verification_code_matches( + stored: &StoredAuthEmailVerificationCode, + verification_token: &str, + code: &str, +) -> bool { + constant_time_eq( + &stored.code_hash, + &auth_verification_code_hash(verification_token, code), + ) +} + fn render_auth_template_string( template: &str, variables: &std::collections::BTreeMap, escape_html: bool, ) -> Result { + if template.len() > ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES + || template.bytes().any(|byte| byte < 0x20 || byte == 0x7f) + { + return Err(GatewayError::Internal( + "email template subject is invalid or oversized".to_string(), + )); + } let mut rendered = template.to_string(); for (key, value) in variables { + if value.len() > ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES { + return Err(GatewayError::Internal( + "email template variable is oversized".to_string(), + )); + } let pattern = regex::Regex::new(&format!(r"\{{\{{\s*{}\s*\}}\}}", regex::escape(key))) .map_err(|err| GatewayError::Internal(err.to_string()))?; let replacement = if escape_html { @@ -60,9 +170,37 @@ fn render_auth_template_string( } else { value.clone() }; + let (matched_bytes, occurrences) = pattern + .find_iter(&rendered) + .fold((0usize, 0usize), |(matched, count), found| { + (matched.saturating_add(found.as_str().len()), count + 1) + }); + let prospective_len = occurrences + .checked_mul(replacement.len()) + .and_then(|bytes| { + rendered + .len() + .checked_sub(matched_bytes)? + .checked_add(bytes) + }); + if prospective_len.is_none_or(|length| length > ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES) { + return Err(GatewayError::Internal( + "rendered email template subject is oversized".to_string(), + )); + } rendered = pattern - .replace_all(&rendered, replacement.as_str()) + .replace_all(&rendered, regex::NoExpand(replacement.as_str())) .into_owned(); + if rendered.len() > ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES { + return Err(GatewayError::Internal( + "rendered email template subject is oversized".to_string(), + )); + } + } + if !admin_email_template_subject_is_valid(&rendered) { + return Err(GatewayError::Internal( + "rendered email template subject is invalid or oversized".to_string(), + )); } Ok(rendered) } @@ -82,7 +220,7 @@ pub(super) async fn read_auth_email_verification_code( state: &AppState, email: &str, ) -> Result, GatewayError> { - let key = auth_email_verification_key(email); + let key = auth_email_verification_key(email)?; let raw = state.runtime_kv_get(&key).await?; raw.map(|value| { serde_json::from_str::(&value) @@ -91,55 +229,72 @@ pub(super) async fn read_auth_email_verification_code( .transpose() } -pub(super) async fn auth_email_is_verified( +pub(in crate::handlers::public::support) async fn auth_email_is_verified( state: &AppState, email: &str, + verification_token: &str, ) -> Result { - let key = auth_email_verified_key(email); + let key = auth_email_registration_proof_key(email, verification_token)?; state.runtime_kv_exists(&key).await } pub(super) async fn mark_auth_email_verified( state: &AppState, email: &str, + verification_token: &str, ) -> Result { - let key = auth_email_verified_key(email); + let key = auth_email_registration_proof_key(email, verification_token)?; state .runtime_kv_setex(&key, "verified", AUTH_EMAIL_VERIFIED_TTL_SECS) .await?; Ok(true) } +pub(super) async fn consume_auth_email_verification_code( + state: &AppState, + email: &str, +) -> Result, GatewayError> { + let key = auth_email_verification_key(email)?; + state + .runtime_kv_getdel(&key) + .await? + .map(|value| { + serde_json::from_str::(&value) + .map_err(|err| GatewayError::Internal(err.to_string())) + }) + .transpose() +} + +pub(in crate::handlers::public::support) async fn consume_auth_email_registration_proof( + state: &AppState, + email: &str, + verification_token: &str, +) -> Result { + let key = auth_email_registration_proof_key(email, verification_token)?; + Ok(state.runtime_kv_getdel(&key).await?.as_deref() == Some("verified")) +} + pub(super) async fn clear_auth_email_pending_code( state: &AppState, email: &str, ) -> Result { - let verification_key = auth_email_verification_key(email); + let verification_key = auth_email_verification_key(email)?; state.runtime_kv_del(&verification_key).await } -pub(super) async fn clear_auth_email_verification( - state: &AppState, - email: &str, -) -> Result { - let verification_key = auth_email_verification_key(email); - let verified_key = auth_email_verified_key(email); - let deleted_pending = state.runtime_kv_del(&verification_key).await?; - let deleted_verified = state.runtime_kv_del(&verified_key).await?; - Ok(deleted_pending || deleted_verified) -} - pub(super) async fn store_auth_email_verification_code( state: &AppState, email: &str, code: &str, + verification_token: &str, created_at: chrono::DateTime, ttl_seconds: u64, ) -> Result { - let key = auth_email_verification_key(email); + let key = auth_email_verification_key(email)?; let value = json!({ - "code": code, + "code_hash": auth_verification_code_hash(verification_token, code), "created_at": created_at.to_rfc3339(), + "verification_token_hash": auth_verification_token_hash(verification_token), }) .to_string(); state.runtime_kv_setex(&key, &value, ttl_seconds).await?; @@ -226,3 +381,63 @@ pub(super) async fn auth_registration_email_configured( ) -> Result { Ok(read_smtp_delivery_config(state).await?.is_some()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn auth_email_runtime_keys_are_keyed_and_contain_no_plaintext_identifiers() { + let email = "Alice+security@example.com"; + let token = "verification-token-security-test"; + let pending = auth_email_verification_key(email).expect("pending key should build"); + let proof = + auth_email_registration_proof_key(email, token).expect("proof key should build"); + + for key in [&pending, &proof] { + assert!(!key.to_ascii_lowercase().contains("alice")); + assert!(!key.contains("example.com")); + assert!(!key.contains(token)); + } + let plain_email_hash = format!("{:x}", Sha256::digest(email.to_ascii_lowercase())); + assert!(!pending.contains(&plain_email_hash)); + assert_ne!(pending, proof); + } + + #[tokio::test] + async fn auth_email_challenges_and_registration_proofs_are_consumed_once_concurrently() { + let state = AppState::new().expect("app state should build"); + let email = "atomic@example.com"; + let token = "atomic-verification-token"; + store_auth_email_verification_code(&state, email, "123456", token, chrono::Utc::now(), 300) + .await + .expect("challenge should store"); + + let (first_challenge, second_challenge) = tokio::join!( + consume_auth_email_verification_code(&state, email), + consume_auth_email_verification_code(&state, email) + ); + assert_eq!( + [first_challenge, second_challenge] + .into_iter() + .filter(|result| result.as_ref().is_ok_and(Option::is_some)) + .count(), + 1 + ); + + mark_auth_email_verified(&state, email, token) + .await + .expect("proof should store"); + let (first_proof, second_proof) = tokio::join!( + consume_auth_email_registration_proof(&state, email, token), + consume_auth_email_registration_proof(&state, email, token) + ); + assert_eq!( + [first_proof, second_proof] + .into_iter() + .filter(|result| result.as_ref().is_ok_and(|consumed| *consumed)) + .count(), + 1 + ); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/auth_helpers.rs b/apps/aether-gateway/src/handlers/public/support/auth_helpers.rs index 7950a5d2d..ce46b0f3d 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_helpers.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_helpers.rs @@ -52,8 +52,7 @@ pub(crate) async fn build_auth_registration_settings_payload( .filter(|value| !value.is_empty()) .is_some(); let enable_registration = system_config_bool(enable_registration.as_ref(), false); - let require_email_verification = - system_config_bool(require_email_verification.as_ref(), false) && email_configured; + let require_email_verification = system_config_bool(require_email_verification.as_ref(), false); let password_policy_level = match system_config_string(password_policy_level_config.as_ref()) { Some(value) if matches!(value.as_str(), "weak" | "medium" | "strong") => value, _ => "weak".to_string(), @@ -127,6 +126,7 @@ pub(crate) fn build_auth_json_response( payload: serde_json::Value, set_cookie: Option, ) -> Response { + let has_set_cookie = set_cookie.is_some(); let mut response = (status, Json(payload)).into_response(); if let Some(set_cookie) = set_cookie { if let Ok(value) = axum::http::HeaderValue::from_str(&set_cookie) { @@ -135,6 +135,22 @@ pub(crate) fn build_auth_json_response( .append(axum::http::header::SET_COOKIE, value); } } + if has_set_cookie { + mark_sensitive_response_no_store(response) + } else { + response + } +} + +pub(crate) fn mark_sensitive_response_no_store(mut response: Response) -> Response { + response.headers_mut().insert( + axum::http::header::CACHE_CONTROL, + axum::http::HeaderValue::from_static("no-store"), + ); + response.headers_mut().insert( + axum::http::header::PRAGMA, + axum::http::HeaderValue::from_static("no-cache"), + ); response } @@ -143,29 +159,69 @@ pub(crate) fn build_auth_error_response( detail: impl Into, clear_cookie: bool, ) -> Response { + let detail = detail.into(); + let public_detail = if status.is_server_error() { + tracing::error!( + event_name = "public_api_internal_error", + %status, + "internal public API error hidden from client" + ); + "服务暂不可用,请稍后重试".to_string() + } else { + detail + }; let cookie = clear_cookie.then(build_auth_refresh_cookie_clear_header); - build_auth_json_response(status, json!({ "detail": detail.into() }), cookie) + build_auth_json_response(status, json!({ "detail": public_detail }), cookie) } -fn auth_environment() -> String { +#[cfg(test)] +mod public_error_projection_tests { + use super::{build_auth_error_response, build_auth_json_response}; + use axum::body::to_bytes; + + #[tokio::test] + async fn public_server_errors_never_return_internal_details() { + let response = build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "database failed for https://user:password@db.example?token=secret", + false, + ); + let body = to_bytes(response.into_body(), 4096).await.expect("body"); + let encoded = String::from_utf8_lossy(&body); + assert!(encoded.contains("服务暂不可用")); + for secret in ["user", "password", "token", "db.example"] { + assert!(!encoded.contains(secret), "leaked {secret}"); + } + } + + #[test] + fn cookie_authenticated_responses_are_never_cacheable() { + let response = build_auth_json_response( + http::StatusCode::OK, + serde_json::json!({ "access_token": "secret" }), + Some("session=secret; HttpOnly".to_string()), + ); + + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&http::HeaderValue::from_static("no-store")) + ); + assert_eq!( + response.headers().get(http::header::PRAGMA), + Some(&http::HeaderValue::from_static("no-cache")) + ); + } +} + +fn auth_environment() -> Option { std::env::var("ENVIRONMENT") .ok() .map(|value| value.trim().to_string()) .filter(|value| !value.is_empty()) - .unwrap_or_else(|| "development".to_string()) } pub(super) fn auth_jwt_secret() -> Result { - if let Ok(value) = std::env::var("JWT_SECRET_KEY") { - let value = value.trim(); - if !value.is_empty() { - return Ok(value.to_string()); - } - } - if auth_environment().eq_ignore_ascii_case("production") { - return Err("JWT_SECRET_KEY 未配置".to_string()); - } - Ok("aether-rust-dev-jwt-secret".to_string()) + crate::local_auth_token::local_auth_jwt_secret() } pub(super) fn auth_access_token_expiry_hours() -> i64 { @@ -200,11 +256,28 @@ pub(super) fn auth_refresh_cookie_name() -> String { .unwrap_or_else(|| "aether_refresh_token".to_string()) } -fn auth_refresh_cookie_secure() -> bool { - std::env::var("AUTH_REFRESH_COOKIE_SECURE") - .ok() - .map(|value| value.trim().eq_ignore_ascii_case("true")) - .unwrap_or_else(|| auth_environment().eq_ignore_ascii_case("production")) +pub(crate) fn auth_refresh_cookie_secure() -> bool { + auth_refresh_cookie_secure_from_values( + std::env::var("AUTH_REFRESH_COOKIE_SECURE").ok().as_deref(), + auth_environment().as_deref(), + ) +} + +fn auth_refresh_cookie_secure_from_values( + explicit_secure: Option<&str>, + environment: Option<&str>, +) -> bool { + match explicit_secure.map(str::trim) { + Some(value) if value.eq_ignore_ascii_case("true") => true, + Some(value) if value.eq_ignore_ascii_case("false") => false, + Some(_) => true, + None => !environment.is_some_and(|value| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "development" | "test" | "local" + ) + }), + } } fn auth_refresh_cookie_samesite() -> &'static str { @@ -212,7 +285,12 @@ fn auth_refresh_cookie_samesite() -> &'static str { Ok(value) if value.trim().eq_ignore_ascii_case("strict") => "Strict", Ok(value) if value.trim().eq_ignore_ascii_case("none") => "None", Ok(value) if value.trim().eq_ignore_ascii_case("lax") => "Lax", - _ if auth_environment().eq_ignore_ascii_case("production") => "None", + _ if auth_environment() + .as_deref() + .is_some_and(|value| value.eq_ignore_ascii_case("production")) => + { + "None" + } _ => "Lax", } } @@ -231,7 +309,7 @@ pub(super) fn build_auth_refresh_cookie_header(refresh_token: &str) -> String { cookie } -pub(super) fn build_auth_refresh_cookie_clear_header() -> String { +pub(crate) fn build_auth_refresh_cookie_clear_header() -> String { let mut cookie = format!( "{}=; Path=/api/auth; HttpOnly; SameSite={}; Max-Age=0", auth_refresh_cookie_name(), @@ -254,23 +332,42 @@ fn auth_non_empty_string(value: Option) -> Option { } pub(super) fn extract_bearer_token(headers: &http::HeaderMap) -> Option { - let value = crate::headers::header_value_str(headers, http::header::AUTHORIZATION.as_str())?; + let mut values = headers.get_all(http::header::AUTHORIZATION).iter(); + let value = values.next()?.to_str().ok()?.trim(); + if values.next().is_some() { + return None; + } let (scheme, token) = value.split_once(' ')?; if !scheme.eq_ignore_ascii_case("bearer") { return None; } - auth_non_empty_string(Some(token.to_string())) + let token = token.trim(); + if token.is_empty() + || !token + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) + { + return None; + } + Some(token.to_string()) } -pub(super) fn extract_cookie_value(headers: &http::HeaderMap, cookie_name: &str) -> Option { - let header = crate::headers::header_value_str(headers, http::header::COOKIE.as_str())?; - for pair in header.split(';') { - let (name, value) = pair.trim().split_once('=')?; - if name.trim() == cookie_name { - return auth_non_empty_string(Some(value.to_string())); +pub(crate) fn extract_cookie_value(headers: &http::HeaderMap, cookie_name: &str) -> Option { + let mut found = None; + for header in headers.get_all(http::header::COOKIE).iter() { + let header = header.to_str().ok()?; + for pair in header.split(';') { + let (name, value) = pair.trim().split_once('=')?; + if name.trim() != cookie_name { + continue; + } + let value = auth_non_empty_string(Some(value.to_string()))?; + if found.replace(value).is_some() { + return None; + } } } - None + found } pub(crate) fn extract_client_device_id( @@ -303,41 +400,32 @@ pub(crate) fn extract_client_device_id( Ok(candidate.to_string()) } +pub(crate) fn extract_client_device_id_header( + headers: &http::HeaderMap, +) -> Result> { + let candidate = + crate::headers::header_value_str(headers, "x-client-device-id").unwrap_or_default(); + let candidate = candidate.trim(); + if candidate.is_empty() + || candidate.len() > 128 + || !candidate + .chars() + .all(|ch| ch.is_ascii_alphanumeric() || ch == '-' || ch == '_') + { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "缺少或无效的设备标识", + false, + )); + } + Ok(candidate.to_string()) +} + pub(super) fn auth_user_agent(headers: &http::HeaderMap) -> Option { crate::headers::header_value_str(headers, http::header::USER_AGENT.as_str()) .map(|value| value.chars().take(1000).collect()) } -pub(super) fn auth_client_ip(headers: &http::HeaderMap) -> Option { - crate::headers::header_value_str(headers, "x-forwarded-for") - .and_then(|value| { - value - .split(',') - .next() - .map(|segment| segment.trim().to_string()) - }) - .filter(|value| !value.is_empty()) - .map(|value| value.chars().take(45).collect()) - .or_else(|| { - crate::headers::header_value_str(headers, "x-real-ip") - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.chars().take(45).collect()) - }) -} - -pub(super) fn auth_client_ip_with_cf( - headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, -) -> Option { - cf_connecting_ip - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.chars().take(45).collect()) - .or_else(|| auth_client_ip(headers)) -} - pub(super) fn normalize_auth_login_identifier(value: &str) -> String { let normalized = value.trim(); if normalized.contains('@') { @@ -356,3 +444,90 @@ pub(super) fn validate_auth_login_password(password: &str) -> Result<(), String> } Ok(()) } + +#[cfg(test)] +mod tests { + use super::{ + auth_refresh_cookie_secure_from_values, extract_bearer_token, extract_cookie_value, + }; + use axum::http::{header, HeaderMap, HeaderValue}; + + #[test] + fn bearer_token_requires_one_unambiguous_authorization_value() { + let mut headers = HeaderMap::new(); + headers.insert( + header::AUTHORIZATION, + HeaderValue::from_static("bEaReR abc.def_ghi-jkl"), + ); + assert_eq!( + extract_bearer_token(&headers).as_deref(), + Some("abc.def_ghi-jkl") + ); + + headers.append( + header::AUTHORIZATION, + HeaderValue::from_static("Bearer second-token"), + ); + assert_eq!(extract_bearer_token(&headers), None); + + headers.clear(); + headers.insert( + header::AUTHORIZATION, + HeaderValue::from_static("Bearer first-token, Bearer second-token"), + ); + assert_eq!(extract_bearer_token(&headers), None); + } + + #[test] + fn cookie_value_rejects_duplicate_cookie_names() { + let mut headers = HeaderMap::new(); + headers.insert( + header::COOKIE, + HeaderValue::from_static("theme=dark; refresh=first-token"), + ); + assert_eq!( + extract_cookie_value(&headers, "refresh").as_deref(), + Some("first-token") + ); + + headers.append( + header::COOKIE, + HeaderValue::from_static("refresh=second-token"), + ); + assert_eq!(extract_cookie_value(&headers, "refresh"), None); + + headers.clear(); + headers.insert( + header::COOKIE, + HeaderValue::from_static("refresh=first-token; refresh=second-token"), + ); + assert_eq!(extract_cookie_value(&headers, "refresh"), None); + } + + #[test] + fn refresh_cookie_secure_fails_closed_when_environment_is_missing_or_unknown() { + assert!(auth_refresh_cookie_secure_from_values(None, None)); + assert!(auth_refresh_cookie_secure_from_values( + None, + Some("staging") + )); + assert!(auth_refresh_cookie_secure_from_values( + Some("invalid"), + Some("development") + )); + } + + #[test] + fn refresh_cookie_secure_allows_only_explicit_local_or_override_opt_out() { + assert!(!auth_refresh_cookie_secure_from_values( + None, + Some("development") + )); + assert!(!auth_refresh_cookie_secure_from_values(None, Some("test"))); + assert!(!auth_refresh_cookie_secure_from_values(Some("false"), None)); + assert!(auth_refresh_cookie_secure_from_values( + Some("true"), + Some("development") + )); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/auth_ldap.rs b/apps/aether-gateway/src/handlers/public/support/auth_ldap.rs index 1be574dba..ee2ecff3b 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_ldap.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_ldap.rs @@ -1,9 +1,9 @@ use super::{ - decrypt_catalog_secret_with_fallbacks, ldap_config_is_enabled, module_available_from_env, + decrypt_or_migrate_ldap_bind_password, ldap_config_is_enabled, module_available_from_env, normalize_auth_login_identifier, system_config_bool, AppState, GatewayError, }; -#[derive(Debug, Clone)] +#[derive(Clone)] pub(super) struct AuthLdapRuntimeConfig { server_url: String, bind_dn: String, @@ -28,7 +28,7 @@ pub(super) struct AuthLdapAuthenticatedUser { pub(super) async fn auth_local_login_allowed_for_user( state: &AppState, - user: &aether_data::repository::users::StoredUserAuthRecord, + user: Option<&aether_data::repository::users::StoredUserAuthRecord>, ) -> Result { let ldap_enabled_config = state .read_system_config_json_value("module.ldap.enabled") @@ -45,7 +45,9 @@ pub(super) async fn auth_local_login_allowed_for_user( if !ldap_exclusive { return Ok(true); } - Ok(user.role.eq_ignore_ascii_case("admin") && user.auth_source.eq_ignore_ascii_case("local")) + Ok(user.is_some_and(|user| { + user.role.eq_ignore_ascii_case("admin") && user.auth_source.eq_ignore_ascii_case("local") + })) } fn auth_ldap_default_search_filter(username_attr: &str) -> String { @@ -85,31 +87,8 @@ fn auth_ldap_escape_filter(value: &str) -> Result { Ok(escaped) } -fn auth_ldap_normalize_server_url(server_url: &str) -> Option { - let server_url = server_url.trim(); - if server_url.is_empty() { - return None; - } - if server_url.contains("://") { - return Some(server_url.to_string()); - } - Some(format!("ldap://{server_url}")) -} - -fn auth_ldap_decrypt_bind_password( - state: &AppState, - config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, -) -> Option { - config - .bind_password_encrypted - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), value) - .unwrap_or_else(|| value.to_string()) - }) - .filter(|value| !value.trim().is_empty()) +fn auth_ldap_normalize_server_url(server_url: &str, use_starttls: bool) -> Option { + crate::handlers::shared::normalize_ldap_transport_server_url(server_url, use_starttls) } async fn read_auth_ldap_runtime_config( @@ -126,10 +105,11 @@ async fn read_auth_ldap_runtime_config( }) else { return Ok(None); }; - let Some(server_url) = auth_ldap_normalize_server_url(&config.server_url) else { + let Some(server_url) = auth_ldap_normalize_server_url(&config.server_url, config.use_starttls) + else { return Ok(None); }; - let Some(bind_password) = auth_ldap_decrypt_bind_password(state, &config) else { + let Some(bind_password) = decrypt_or_migrate_ldap_bind_password(state, &config).await? else { return Ok(None); }; diff --git a/apps/aether-gateway/src/handlers/public/support/auth_rate_limit.rs b/apps/aether-gateway/src/handlers/public/support/auth_rate_limit.rs new file mode 100644 index 000000000..f5dca1285 --- /dev/null +++ b/apps/aether-gateway/src/handlers/public/support/auth_rate_limit.rs @@ -0,0 +1,318 @@ +use super::{build_auth_json_response, http, json, AppState, Body, Response}; +use aether_runtime_state::{UsageLimitCheck, UsageLimitInput, UsageLimitRule}; +use hmac::Mac; +use sha2::Sha256; +use std::net::IpAddr; +use std::time::{SystemTime, UNIX_EPOCH}; + +const AUTH_RATE_LIMIT_UNAVAILABLE_DETAIL: &str = "认证安全服务暂不可用,请稍后重试"; +const AUTH_RATE_LIMITED_DETAIL: &str = "请求过于频繁,请稍后重试"; +const AUTH_VERIFICATION_MAX_FAILURES: u32 = 5; + +#[derive(Debug, Clone, Copy)] +pub(super) struct AuthRateLimitPolicy { + action: &'static str, + window_seconds: u64, + ip_limit: u32, + identity_limit: u32, +} + +impl AuthRateLimitPolicy { + const fn new( + action: &'static str, + window_seconds: u64, + ip_limit: u32, + identity_limit: u32, + ) -> Self { + Self { + action, + window_seconds, + ip_limit, + identity_limit, + } + } +} + +pub(super) const AUTH_LOGIN_RATE_LIMIT: AuthRateLimitPolicy = + AuthRateLimitPolicy::new("login", 60, 60, 10); +pub(super) const AUTH_SEND_VERIFICATION_RATE_LIMIT: AuthRateLimitPolicy = + AuthRateLimitPolicy::new("send-verification-code", 3_600, 20, 3); +pub(super) const AUTH_REGISTER_RATE_LIMIT: AuthRateLimitPolicy = + AuthRateLimitPolicy::new("register", 3_600, 10, 5); +pub(super) const AUTH_VERIFY_EMAIL_RATE_LIMIT: AuthRateLimitPolicy = + AuthRateLimitPolicy::new("verify-email", 300, 30, 10); +pub(super) const AUTH_VERIFICATION_STATUS_RATE_LIMIT: AuthRateLimitPolicy = + AuthRateLimitPolicy::new("verification-status", 60, 60, 20); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AuthRateLimitCheck { + Allowed, + Rejected { retry_after: u64 }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum AuthVerificationFailureDecision { + Incorrect, + Exhausted { retry_after: u64 }, +} + +fn unix_timestamp_millis() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis() + .try_into() + .unwrap_or(u64::MAX) +} + +fn digest_subject_with_key( + secret: &[u8], + action: &str, + dimension: &str, + subject: &str, +) -> Result { + let mut mac = hmac::Hmac::::new_from_slice(secret) + .map_err(|_| "auth rate-limit HMAC key invalid".to_string())?; + mac.update(b"aether-auth-rate-limit-v1\0"); + mac.update(action.as_bytes()); + mac.update(b"\0"); + mac.update(dimension.as_bytes()); + mac.update(b"\0"); + mac.update(subject.as_bytes()); + Ok(mac + .finalize() + .into_bytes() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect()) +} + +fn rate_limit_key_with_key( + secret: &[u8], + action: &str, + dimension: &str, + subject: &str, + window_seconds: u64, +) -> Result { + let digest = digest_subject_with_key(secret, action, dimension, subject)?; + Ok(format!( + "auth:rate-limit:v2:{{{digest}}}:{action}:{dimension}:{}", + window_seconds.max(1) + )) +} + +fn rate_limit_key( + action: &str, + dimension: &str, + subject: &str, + window_seconds: u64, +) -> Result { + let secret = super::auth_jwt_secret()?; + rate_limit_key_with_key( + secret.as_bytes(), + action, + dimension, + subject, + window_seconds, + ) +} + +async fn consume_auth_rate_limit( + state: &AppState, + action: &str, + dimension: &str, + subject: &str, + limit: u32, + window_seconds: u64, +) -> Result { + let window_seconds = window_seconds.max(1); + let now_unix_ms = unix_timestamp_millis(); + let counter_key = rate_limit_key(action, dimension, subject, window_seconds)?; + let event_id = uuid::Uuid::new_v4().simple().to_string(); + let rule = UsageLimitRule { + key: &counter_key, + limit: u64::from(limit.max(1)), + window_seconds, + retention_seconds: window_seconds, + }; + let result = state + .runtime_state() + .check_and_consume_usage_limits(UsageLimitInput { + rules: std::slice::from_ref(&rule), + event_id: &event_id, + now_unix_ms, + }) + .await + .map_err(|err| err.to_string())?; + + Ok(match result { + UsageLimitCheck::Allowed => AuthRateLimitCheck::Allowed, + UsageLimitCheck::Rejected { retry_after, .. } => AuthRateLimitCheck::Rejected { + retry_after: retry_after.max(1), + }, + }) +} + +pub(super) fn build_auth_rate_limited_response( + detail: &'static str, + retry_after: u64, +) -> Response { + let retry_after = retry_after.max(1); + let mut response = build_auth_json_response( + http::StatusCode::TOO_MANY_REQUESTS, + json!({ + "detail": detail, + "retry_after": retry_after, + }), + None, + ); + if let Ok(value) = http::HeaderValue::from_str(&retry_after.to_string()) { + response + .headers_mut() + .insert(http::header::RETRY_AFTER, value); + } + response +} + +fn build_auth_rate_limit_unavailable_response() -> Response { + build_auth_json_response( + http::StatusCode::SERVICE_UNAVAILABLE, + json!({ "detail": AUTH_RATE_LIMIT_UNAVAILABLE_DETAIL }), + None, + ) +} + +async fn enforce_auth_rate_limit( + state: &AppState, + policy: AuthRateLimitPolicy, + dimension: &'static str, + subject: &str, + limit: u32, +) -> Result<(), Response> { + match consume_auth_rate_limit( + state, + policy.action, + dimension, + subject, + limit, + policy.window_seconds, + ) + .await + { + Ok(AuthRateLimitCheck::Allowed) => Ok(()), + Ok(AuthRateLimitCheck::Rejected { retry_after }) => Err(build_auth_rate_limited_response( + AUTH_RATE_LIMITED_DETAIL, + retry_after, + )), + Err(_) => { + tracing::warn!( + event_name = "auth_rate_limit_check_failed", + action = policy.action, + dimension, + "authentication request rejected because the rate-limit backend failed" + ); + Err(build_auth_rate_limit_unavailable_response()) + } + } +} + +pub(super) async fn enforce_auth_ip_rate_limit( + state: &AppState, + policy: AuthRateLimitPolicy, + client_ip: IpAddr, +) -> Result<(), Response> { + enforce_auth_rate_limit(state, policy, "ip", &client_ip.to_string(), policy.ip_limit).await +} + +pub(super) async fn enforce_auth_identity_rate_limit( + state: &AppState, + policy: AuthRateLimitPolicy, + identity: &str, +) -> Result<(), Response> { + enforce_auth_rate_limit(state, policy, "identity", identity, policy.identity_limit).await +} + +pub(super) async fn record_auth_verification_failure( + state: &AppState, + email: &str, + challenge_created_at: &str, + challenge_code_hash: &str, + challenge_ttl_seconds: u64, +) -> Result> { + let challenge_subject = format!("{email}\0{challenge_created_at}\0{challenge_code_hash}"); + match consume_auth_rate_limit( + state, + "verify-email-challenge", + "challenge", + &challenge_subject, + AUTH_VERIFICATION_MAX_FAILURES.saturating_sub(1), + challenge_ttl_seconds.max(1), + ) + .await + { + Ok(AuthRateLimitCheck::Rejected { .. }) => Ok(AuthVerificationFailureDecision::Exhausted { + retry_after: challenge_ttl_seconds.max(1), + }), + Ok(AuthRateLimitCheck::Allowed) => Ok(AuthVerificationFailureDecision::Incorrect), + Err(_) => { + tracing::warn!( + event_name = "auth_verification_failure_count_failed", + "email verification challenge rejected because the failure counter failed" + ); + Err(build_auth_rate_limit_unavailable_response()) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use sha2::Digest; + + #[test] + fn rate_limit_keys_use_keyed_subject_digests_without_plaintext() { + let secret = b"fixed-test-secret-with-at-least-32-bytes"; + let cases = [ + ("login", "ip", "203.0.113.42"), + ("login", "identity", "target@example.com"), + ( + "verify-email-challenge", + "challenge", + "target@example.com\0created-at\0code-hash", + ), + ]; + + for (action, dimension, subject) in cases { + let digest = digest_subject_with_key(secret, action, dimension, subject) + .expect("fixed HMAC key should be accepted"); + let enumerable_digest = format!( + "{:x}", + Sha256::digest(format!("{dimension}\0{subject}").as_bytes()) + ); + let key = rate_limit_key_with_key(secret, action, dimension, subject, 60) + .expect("fixed HMAC key should be accepted"); + + assert_ne!(digest, enumerable_digest); + assert!(key.contains(&format!("{{{digest}}}"))); + assert!(!key.contains(&enumerable_digest)); + for plaintext_part in subject.split('\0') { + assert!(!key.contains(plaintext_part)); + } + } + } + + #[test] + fn rate_limit_subject_digest_is_domain_separated() { + let secret = b"fixed-test-secret-with-at-least-32-bytes"; + let subject = "same-subject"; + let login_identity = digest_subject_with_key(secret, "login", "identity", subject) + .expect("fixed HMAC key should be accepted"); + let register_identity = digest_subject_with_key(secret, "register", "identity", subject) + .expect("fixed HMAC key should be accepted"); + let login_ip = digest_subject_with_key(secret, "login", "ip", subject) + .expect("fixed HMAC key should be accepted"); + + assert_ne!(login_identity, register_identity); + assert_ne!(login_identity, login_ip); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/auth_registration.rs b/apps/aether-gateway/src/handlers/public/support/auth_registration.rs index d7f35a442..48f8182f7 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_registration.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_registration.rs @@ -1,18 +1,36 @@ use super::{ auth_email_is_verified, auth_now, auth_registration_email_configured, auth_verification_code_expire_minutes, auth_verification_send_cooldown_seconds, - build_auth_error_response, build_auth_json_response, build_auth_verification_email, - clear_auth_email_pending_code, clear_auth_email_verification, generate_auth_verification_code, - http, json, mark_auth_email_verified, read_auth_email_verification_code, read_auth_smtp_config, - send_auth_email, store_auth_email_verification_code, system_config_bool, system_config_f64, - system_config_string, system_config_string_list, verify_auth_turnstile, AppState, - AuthTurnstileAction, Body, GatewayError, Regex, Response, + build_auth_error_response, build_auth_json_response, build_auth_rate_limited_response, + build_auth_verification_email, clear_auth_email_pending_code, + consume_auth_email_registration_proof, consume_auth_email_verification_code, + enforce_auth_identity_rate_limit, enforce_auth_ip_rate_limit, generate_auth_verification_code, + generate_auth_verification_token, http, json, mark_auth_email_verified, + mark_sensitive_response_no_store, read_auth_email_verification_code, read_auth_smtp_config, + record_auth_verification_failure, send_auth_email, store_auth_email_verification_code, + system_config_bool, system_config_f64, system_config_string, system_config_string_list, + verify_auth_turnstile, AppState, AuthTurnstileAction, AuthVerificationFailureDecision, Body, + GatewayError, Regex, Response, AUTH_REGISTER_RATE_LIMIT, AUTH_SEND_VERIFICATION_RATE_LIMIT, + AUTH_VERIFICATION_STATUS_RATE_LIMIT, AUTH_VERIFY_EMAIL_RATE_LIMIT, }; use serde::Deserialize; +use std::net::IpAddr; const AUTH_REGISTRATION_STORAGE_UNAVAILABLE_DETAIL: &str = "注册数据存储暂不可用"; +const AUTH_REGISTRATION_UNAVAILABLE_DETAIL: &str = "注册暂不可用,请稍后重试"; +const AUTH_REGISTRATION_IDENTITY_UNAVAILABLE_DETAIL: &str = "无法使用该注册信息"; +const AUTH_VERIFICATION_UNAVAILABLE_DETAIL: &str = "邮箱验证暂不可用,请稍后重试"; -#[derive(Debug, Deserialize)] +fn build_auth_internal_error_response( + event_name: &'static str, + _err: &GatewayError, + detail: &'static str, +) -> Response { + tracing::warn!(event_name, "public authentication request failed"); + build_auth_error_response(http::StatusCode::INTERNAL_SERVER_ERROR, detail, false) +} + +#[derive(Deserialize)] struct AuthRegisterRequest { email: Option, username: String, @@ -21,18 +39,26 @@ struct AuthRegisterRequest { invite_code: Option, privacy_policy_accepted: Option, privacy_policy_version: Option, + email_verification_token: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct AuthEmailRequest { email: String, turnstile_token: Option, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct AuthVerifyEmailRequest { email: String, code: String, + verification_token: String, +} + +#[derive(Deserialize)] +struct AuthVerificationStatusRequest { + email: String, + verification_token: String, } fn normalize_auth_email(value: &str) -> Option { @@ -203,9 +229,14 @@ async fn validate_auth_email_suffix( pub(super) async fn handle_auth_send_verification_code( state: &AppState, headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + client_ip: IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { + if let Err(response) = + enforce_auth_ip_rate_limit(state, AUTH_SEND_VERIFICATION_RATE_LIMIT, client_ip).await + { + return response; + } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "请求数据验证失败", false); }; @@ -222,11 +253,15 @@ pub(super) async fn handle_auth_send_verification_code( let Some(email) = normalize_auth_email(&payload.email) else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "邮箱格式无效", false); }; + if let Err(response) = + enforce_auth_identity_rate_limit(state, AUTH_SEND_VERIFICATION_RATE_LIMIT, &email).await + { + return response; + } if let Err(response) = verify_auth_turnstile( state, - headers, - cf_connecting_ip, + client_ip, payload.turnstile_token.as_deref(), AuthTurnstileAction::SendVerificationCode, ) @@ -241,28 +276,14 @@ pub(super) async fn handle_auth_send_verification_code( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_send_verification_email_policy_lookup_failed", + &err, + AUTH_VERIFICATION_UNAVAILABLE_DETAIL, ); } } - if state - .find_user_auth_by_identifier(&email) - .await - .ok() - .flatten() - .is_some() - { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "该邮箱已被注册,请直接登录或使用其他邮箱", - false, - ); - } - let smtp_config = match read_auth_smtp_config(state).await { Ok(Some(value)) => value, Ok(None) => { @@ -273,10 +294,10 @@ pub(super) async fn handle_auth_send_verification_code( ); } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth smtp settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_send_verification_smtp_lookup_failed", + &err, + AUTH_VERIFICATION_UNAVAILABLE_DETAIL, ); } }; @@ -306,14 +327,15 @@ pub(super) async fn handle_auth_send_verification_code( let expire_minutes = auth_verification_code_expire_minutes(); let code = generate_auth_verification_code(); + let verification_token = generate_auth_verification_token(); let email_message = match build_auth_verification_email(state, &email, &code, expire_minutes).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth verification email render failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_verification_email_render_failed", + &err, + AUTH_VERIFICATION_UNAVAILABLE_DETAIL, ); } }; @@ -322,19 +344,25 @@ pub(super) async fn handle_auth_send_verification_code( state, &email, &code, + &verification_token, now, u64::try_from(expire_minutes.saturating_mul(60)).unwrap_or(300), ) .await { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth verification code save failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_verification_challenge_store_failed", + &err, + AUTH_VERIFICATION_UNAVAILABLE_DETAIL, ); } - if let Err(_err) = send_auth_email(state, smtp_config, email_message).await { + if let Err(err) = send_auth_email(state, smtp_config, email_message).await { + tracing::warn!( + event_name = "auth_verification_email_send_failed", + error = ?err, + "failed to send authentication verification email" + ); let _ = clear_auth_email_pending_code(state, &email).await; return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, @@ -343,23 +371,29 @@ pub(super) async fn handle_auth_send_verification_code( ); } - build_auth_json_response( + mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "message": "验证码已发送,请查收邮件", "success": true, "expire_minutes": expire_minutes, + "verification_token": verification_token, }), None, - ) + )) } pub(super) async fn handle_auth_register( state: &AppState, headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + client_ip: IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { + if let Err(response) = + enforce_auth_ip_rate_limit(state, AUTH_REGISTER_RATE_LIMIT, client_ip).await + { + return response; + } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少请求体", false); }; @@ -381,13 +415,30 @@ pub(super) async fn handle_auth_register( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } }; + if let Some(email) = email.as_deref() { + if let Err(response) = + enforce_auth_identity_rate_limit(state, AUTH_REGISTER_RATE_LIMIT, email).await + { + return response; + } + } + let username_rate_limit_identity = format!("username:{}", username.to_ascii_lowercase()); + if let Err(response) = enforce_auth_identity_rate_limit( + state, + AUTH_REGISTER_RATE_LIMIT, + &username_rate_limit_identity, + ) + .await + { + return response; + } let password_policy = match auth_password_policy_level(state).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_password_policy_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; @@ -401,10 +452,10 @@ pub(super) async fn handle_auth_register( { Ok(value) => system_config_bool(value.as_ref(), false), Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_enabled_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; @@ -414,10 +465,10 @@ pub(super) async fn handle_auth_register( let privacy_policy = match read_registration_privacy_policy_settings(state).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_privacy_policy_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; @@ -439,8 +490,7 @@ pub(super) async fn handle_auth_register( if let Err(response) = verify_auth_turnstile( state, - headers, - cf_connecting_ip, + client_ip, payload.turnstile_token.as_deref(), AuthTurnstileAction::Register, ) @@ -452,10 +502,10 @@ pub(super) async fn handle_auth_register( let email_configured = match auth_registration_email_configured(state).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_email_configuration_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; @@ -463,15 +513,26 @@ pub(super) async fn handle_auth_register( .read_system_config_json_value("require_email_verification") .await { - Ok(value) => system_config_bool(value.as_ref(), false) && email_configured, + Ok(value) => system_config_bool(value.as_ref(), false), Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_verification_policy_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; + if require_verification && !email_configured { + tracing::error!( + event_name = "auth_registration_verification_channel_unavailable", + "email verification is required but SMTP delivery is not configured" + ); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, + false, + ); + } if require_verification && email.is_none() { return build_auth_error_response( @@ -480,26 +541,17 @@ pub(super) async fn handle_auth_register( false, ); } - if require_verification { - if let Some(email) = email.as_deref() { - let is_verified = match auth_email_is_verified(state, email).await { - Ok(value) => value, - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth verification lookup failed: {err:?}"), - false, - ); - } - }; - if !is_verified { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "请先完成邮箱验证。请发送验证码并验证后再注册。", - false, - ); - } - } + let email_verification_token = payload + .email_verification_token + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + if require_verification && email_verification_token.is_none() { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "请先完成邮箱验证。请发送验证码并验证后再注册。", + false, + ); } if let Some(email) = email.as_deref() { match validate_auth_email_suffix(state, email).await { @@ -508,39 +560,68 @@ pub(super) async fn handle_auth_register( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_email_policy_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } } - if state - .find_user_auth_by_identifier(email) - .await - .ok() - .flatten() - .is_some() - { + if require_verification { + let verification_token = email_verification_token + .expect("verified registration requires verification token"); + match auth_email_is_verified(state, email, verification_token).await { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "邮箱验证凭据无效或已过期,请重新验证", + false, + ); + } + Err(err) => { + return build_auth_internal_error_response( + "auth_registration_proof_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, + ); + } + } + } + match state.find_user_auth_by_identifier(email).await { + Ok(Some(_)) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + AUTH_REGISTRATION_IDENTITY_UNAVAILABLE_DETAIL, + false, + ); + } + Ok(None) => {} + Err(err) => { + return build_auth_internal_error_response( + "auth_registration_email_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, + ); + } + } + } + match state.find_user_auth_by_identifier(&username).await { + Ok(Some(_)) => { return build_auth_error_response( http::StatusCode::BAD_REQUEST, - format!("邮箱已存在: {email}"), + AUTH_REGISTRATION_IDENTITY_UNAVAILABLE_DETAIL, false, ); } - } - if state - .find_user_auth_by_identifier(&username) - .await - .ok() - .flatten() - .is_some() - { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - format!("用户名已存在: {username}"), - false, - ); + Ok(None) => {} + Err(err) => { + return build_auth_internal_error_response( + "auth_registration_username_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, + ); + } } let password_hash = match bcrypt::hash(&payload.password, bcrypt::DEFAULT_COST) { @@ -559,15 +640,44 @@ pub(super) async fn handle_auth_register( { Ok(value) => system_config_f64(value.as_ref(), 10.0), Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth settings lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_initial_gift_lookup_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }; - let Some((user, _wallet)) = (match state - .register_local_auth_user( + if require_verification { + let email = email + .as_deref() + .expect("verified registration requires email"); + let verification_token = + email_verification_token.expect("verified registration requires verification token"); + match consume_auth_email_registration_proof(state, email, verification_token).await { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "邮箱验证凭据无效或已过期,请重新验证", + false, + ); + } + Err(err) => { + tracing::warn!( + event_name = "auth_email_registration_proof_consume_failed", + error = ?err, + "failed to consume email registration proof" + ); + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "注册暂不可用,请稍后重试", + false, + ); + } + } + } + let Some((user, wallet, wallet_created)) = (match state + .register_local_auth_user_with_wallet_outcome( email.clone(), require_verification && email.is_some(), username.clone(), @@ -579,10 +689,10 @@ pub(super) async fn handle_auth_register( { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth register failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_registration_create_user_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } }) else { @@ -592,15 +702,18 @@ pub(super) async fn handle_auth_register( false, ); }; + let owned_wallet_id = wallet_created.then(|| wallet.id.clone()); if let Err(err) = state .assign_default_group_to_self_registered_user(&user.id) .await { - let _ = state.delete_local_auth_user(&user.id).await; - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth default user group assignment failed: {err:?}"), - false, + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await; + return build_auth_internal_error_response( + "auth_registration_default_group_assignment_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } if privacy_policy.enabled { @@ -610,7 +723,12 @@ pub(super) async fn handle_auth_register( { Ok(true) => {} Ok(false) => { - let _ = state.delete_local_auth_user(&user.id).await; + let _ = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await; return build_auth_error_response( http::StatusCode::SERVICE_UNAVAILABLE, AUTH_REGISTRATION_STORAGE_UNAVAILABLE_DETAIL, @@ -618,11 +736,16 @@ pub(super) async fn handle_auth_register( ); } Err(err) => { - let _ = state.delete_local_auth_user(&user.id).await; - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth privacy policy acceptance failed: {err:?}"), - false, + let _ = state + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await; + return build_auth_internal_error_response( + "auth_registration_privacy_policy_record_failed", + &err, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, ); } } @@ -635,7 +758,7 @@ pub(super) async fn handle_auth_register( if invite_code.is_some() { let source = json!({ "channel": "registration", - "ip": cf_connecting_ip, + "ip": client_ip.to_string(), "user_agent": headers .get(http::header::USER_AGENT) .and_then(|value| value.to_str().ok()), @@ -649,23 +772,23 @@ pub(super) async fn handle_auth_register( ) .await { - let _ = state.delete_local_auth_user(&user.id).await; + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await; let (status, detail) = match err { GatewayError::Client { status, message } => (status, message), - other => ( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth referral binding failed: {other:?}"), - ), + other => { + return build_auth_internal_error_response( + "auth_registration_referral_binding_failed", + &other, + AUTH_REGISTRATION_UNAVAILABLE_DETAIL, + ); + } }; return build_auth_error_response(status, detail, false); } } - if require_verification { - if let Some(email) = email.as_deref() { - let _ = clear_auth_email_verification(state, email).await; - } - } build_auth_json_response( http::StatusCode::OK, json!({ @@ -680,8 +803,14 @@ pub(super) async fn handle_auth_register( pub(super) async fn handle_auth_verify_email( state: &AppState, + client_ip: IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { + if let Err(response) = + enforce_auth_ip_rate_limit(state, AUTH_VERIFY_EMAIL_RATE_LIMIT, client_ip).await + { + return response; + } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少请求体", false); }; @@ -694,7 +823,20 @@ pub(super) async fn handle_auth_verify_email( let Some(email) = normalize_auth_email(&payload.email) else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "邮箱格式无效", false); }; + if let Err(response) = + enforce_auth_identity_rate_limit(state, AUTH_VERIFY_EMAIL_RATE_LIMIT, &email).await + { + return response; + } let code = payload.code.trim(); + let verification_token = payload.verification_token.trim(); + if verification_token.is_empty() { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证会话无效或已过期", + false, + ); + } if code.len() != 6 || !code.chars().all(|ch| ch.is_ascii_digit()) { return build_auth_error_response( http::StatusCode::BAD_REQUEST, @@ -705,20 +847,27 @@ pub(super) async fn handle_auth_verify_email( let pending = match read_auth_email_verification_code(state, &email).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("verification lookup failed: {err:?}"), - false, + return build_auth_internal_error_response( + "auth_verification_challenge_lookup_failed", + &err, + AUTH_VERIFICATION_UNAVAILABLE_DETAIL, ) } }; let Some(pending) = pending else { return build_auth_error_response( http::StatusCode::BAD_REQUEST, - "验证码不存在或已过期", + "验证会话无效或已过期", false, ); }; + if !super::auth_verification_token_matches(&pending, verification_token) { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证会话无效或已过期", + false, + ); + } let created_at = chrono::DateTime::parse_from_rfc3339(&pending.created_at) .ok() .map(|value| value.with_timezone(&chrono::Utc)); @@ -728,17 +877,86 @@ pub(super) async fn handle_auth_verify_email( let _ = clear_auth_email_pending_code(state, &email).await; return build_auth_error_response( http::StatusCode::BAD_REQUEST, - "验证码不存在或已过期", + "验证会话无效或已过期", false, ); } - if pending.code != code { - return build_auth_error_response(http::StatusCode::BAD_REQUEST, "验证码错误", false); + if !super::auth_verification_code_matches(&pending, verification_token, code) { + let challenge_ttl_seconds = expires_at + .map(|value| value.signed_duration_since(auth_now()).num_seconds().max(1) as u64) + .unwrap_or_else(|| { + u64::try_from(auth_verification_code_expire_minutes().saturating_mul(60)) + .unwrap_or(300) + .max(1) + }); + match record_auth_verification_failure( + state, + &email, + &pending.created_at, + &pending.code_hash, + challenge_ttl_seconds, + ) + .await + { + Ok(AuthVerificationFailureDecision::Incorrect) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证码错误", + false, + ); + } + Ok(AuthVerificationFailureDecision::Exhausted { retry_after }) => { + // The counter and pending challenge use different runtime structures. The + // counter transition is atomic; invalidating the challenge is best-effort. + let _ = clear_auth_email_pending_code(state, &email).await; + return build_auth_rate_limited_response( + "验证码尝试次数过多,请重新获取", + retry_after, + ); + } + Err(response) => { + let _ = clear_auth_email_pending_code(state, &email).await; + return response; + } + } } - if mark_auth_email_verified(state, &email).await.ok() != Some(true) { - return build_auth_error_response(http::StatusCode::BAD_REQUEST, "系统错误", false); + let consumed = match consume_auth_email_verification_code(state, &email).await { + Ok(Some(consumed)) + if consumed.code_hash == pending.code_hash + && super::auth_verification_token_matches(&consumed, verification_token) => + { + true + } + Ok(_) => false, + Err(err) => { + tracing::warn!( + event_name = "auth_email_verification_consume_failed", + error = ?err, + "failed to consume email verification challenge" + ); + false + } + }; + if !consumed + || match mark_auth_email_verified(state, &email, verification_token).await { + Ok(true) => false, + Ok(false) => true, + Err(err) => { + tracing::warn!( + event_name = "auth_registration_proof_store_failed", + error = ?err, + "failed to store email registration proof" + ); + true + } + } + { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证会话无效或已过期", + false, + ); } - let _ = clear_auth_email_pending_code(state, &email).await; build_auth_json_response( http::StatusCode::OK, json!({ "message": "邮箱验证成功", "success": true }), @@ -748,12 +966,18 @@ pub(super) async fn handle_auth_verify_email( pub(super) async fn handle_auth_verification_status( state: &AppState, + client_ip: IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { + if let Err(response) = + enforce_auth_ip_rate_limit(state, AUTH_VERIFICATION_STATUS_RATE_LIMIT, client_ip).await + { + return response; + } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少请求体", false); }; - let payload = match serde_json::from_slice::(request_body) { + let payload = match serde_json::from_slice::(request_body) { Ok(value) => value, Err(_) => { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "输入验证失败", false) @@ -762,11 +986,35 @@ pub(super) async fn handle_auth_verification_status( let Some(email) = normalize_auth_email(&payload.email) else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "邮箱格式无效", false); }; + let verification_token = payload.verification_token.trim(); + if verification_token.is_empty() { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证会话无效或已过期", + false, + ); + } + if let Err(response) = + enforce_auth_identity_rate_limit(state, AUTH_VERIFICATION_STATUS_RATE_LIMIT, &email).await + { + return response; + } let pending = read_auth_email_verification_code(state, &email) .await .ok() .flatten(); - let is_verified = auth_email_is_verified(state, &email).await.unwrap_or(false); + let pending = pending + .filter(|pending| super::auth_verification_token_matches(pending, verification_token)); + let is_verified = auth_email_is_verified(state, &email, verification_token) + .await + .unwrap_or(false); + if pending.is_none() && !is_verified { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "验证会话无效或已过期", + false, + ); + } let now = auth_now(); let (has_pending_code, cooldown_remaining, code_expires_in) = if let Some(pending) = pending { let created_at = chrono::DateTime::parse_from_rfc3339(&pending.created_at) diff --git a/apps/aether-gateway/src/handlers/public/support/auth_session.rs b/apps/aether-gateway/src/handlers/public/support/auth_session.rs index a49894c65..e57c56e6f 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_session.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_session.rs @@ -1,142 +1,39 @@ use super::{ - auth_access_token_expiry_hours, auth_client_ip, auth_jwt_secret, auth_now, - auth_refresh_cookie_name, auth_user_agent, build_auth_error_response, build_auth_json_response, + auth_access_token_expiry_hours, auth_now, auth_refresh_cookie_name, auth_user_agent, + build_auth_error_response, build_auth_internal_error_response, build_auth_json_response, build_auth_refresh_cookie_clear_header, build_auth_refresh_cookie_header, extract_bearer_token, - extract_client_device_id, extract_cookie_value, http, json, AppState, Body, - GatewayPublicRequestContext, Response, AUTH_REFRESH_TOKEN_EXPIRATION_DAYS, + extract_client_device_id, extract_client_device_id_header, extract_cookie_value, http, json, + AppState, Body, GatewayPublicRequestContext, Response, AUTH_REFRESH_TOKEN_EXPIRATION_DAYS, }; +use crate::handlers::public::support::mark_sensitive_response_no_store; +use crate::local_auth_token::LocalAuthTokenType; use crate::GatewayUserSessionView; use uuid::Uuid; -fn base64url_encode(bytes: &[u8]) -> String { - use base64::Engine; - - base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes) -} - -fn base64url_decode(value: &str) -> Result, String> { - use base64::Engine; - - base64::engine::general_purpose::URL_SAFE_NO_PAD - .decode(value) - .map_err(|_| "无效的Token".to_string()) +fn auth_token_error_is_internal(detail: &str) -> bool { + !matches!(detail, "无效的Token" | "Token已过期") && !detail.starts_with("Token类型错误:") } pub(crate) fn create_auth_token( - token_type: &str, - mut payload: serde_json::Map, + token_type: LocalAuthTokenType, + payload: serde_json::Map, expires_at: chrono::DateTime, ) -> Result { - use hmac::Mac; - - let secret = auth_jwt_secret()?; - let header = serde_json::json!({ "alg": "HS256", "typ": "JWT" }); - payload.insert("exp".to_string(), json!(expires_at.timestamp())); - payload.insert("type".to_string(), json!(token_type)); - let header_segment = base64url_encode( - serde_json::to_vec(&header) - .map_err(|_| "无法序列化JWT header".to_string())? - .as_slice(), - ); - let payload_segment = base64url_encode( - serde_json::to_vec(&payload) - .map_err(|_| "无法序列化JWT payload".to_string())? - .as_slice(), - ); - let signing_input = format!("{header_segment}.{payload_segment}"); - let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) - .map_err(|_| "JWT secret 无效".to_string())?; - mac.update(signing_input.as_bytes()); - let signature = mac.finalize().into_bytes(); - Ok(format!( - "{header_segment}.{payload_segment}.{}", - base64url_encode(signature.as_slice()) - )) + crate::local_auth_token::create_local_auth_token(token_type, payload, expires_at) } pub(crate) fn decode_auth_token( token: &str, - expected_type: &str, + expected_type: LocalAuthTokenType, ) -> Result, String> { - use hmac::Mac; - - let secret = auth_jwt_secret()?; - let mut parts = token.split('.'); - let Some(header_segment) = parts.next() else { - return Err("无效的Token".to_string()); - }; - let Some(payload_segment) = parts.next() else { - return Err("无效的Token".to_string()); - }; - let Some(signature_segment) = parts.next() else { - return Err("无效的Token".to_string()); - }; - if parts.next().is_some() { - return Err("无效的Token".to_string()); - } - - let signing_input = format!("{header_segment}.{payload_segment}"); - let signature = base64url_decode(signature_segment)?; - let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) - .map_err(|_| "JWT secret 无效".to_string())?; - mac.update(signing_input.as_bytes()); - mac.verify_slice(&signature) - .map_err(|_| "无效的Token".to_string())?; - - let payload_bytes = base64url_decode(payload_segment)?; - let payload = serde_json::from_slice::(&payload_bytes) - .map_err(|_| "无效的Token".to_string())?; - let payload = payload - .as_object() - .cloned() - .ok_or_else(|| "无效的Token".to_string())?; - let actual_type = payload - .get("type") - .and_then(serde_json::Value::as_str) - .unwrap_or_default(); - if actual_type != expected_type { - return Err(format!( - "Token类型错误: 期望 {expected_type}, 实际 {actual_type}" - )); - } - let exp = payload - .get("exp") - .and_then(serde_json::Value::as_i64) - .ok_or_else(|| "无效的Token".to_string())?; - if exp <= auth_now().timestamp() { - return Err("Token已过期".to_string()); - } - Ok(payload) + crate::local_auth_token::decode_local_auth_token(token, expected_type) } pub(super) fn auth_token_identity_matches_user( payload: &serde_json::Map, user: &aether_data::repository::users::StoredUserAuthRecord, ) -> bool { - if let Some(token_email) = payload.get("email").and_then(serde_json::Value::as_str) { - if user - .email - .as_deref() - .is_some_and(|email| email != token_email) - { - return false; - } - } - - let Some(token_created_at) = payload - .get("created_at") - .and_then(serde_json::Value::as_str) - else { - return true; - }; - let Some(user_created_at) = user.created_at else { - return true; - }; - let Ok(token_created_at) = chrono::DateTime::parse_from_rfc3339(token_created_at) else { - return false; - }; - let token_created_at = token_created_at.with_timezone(&chrono::Utc); - (user_created_at - token_created_at).num_seconds().abs() <= 1 + crate::local_auth_token::local_auth_token_identity_matches_user(payload, user) } pub(crate) fn build_auth_wallet_summary_payload( @@ -217,14 +114,21 @@ pub(crate) async fn resolve_authenticated_local_user( false, )); }; - let claims = match decode_auth_token(&token, "access") { + let claims = match decode_auth_token(&token, LocalAuthTokenType::Access) { Ok(value) => value, Err(detail) => { + if auth_token_error_is_internal(&detail) { + return Err(build_auth_internal_error_response( + "auth_access_token_decode_failed", + detail, + false, + )); + } return Err(build_auth_error_response( http::StatusCode::UNAUTHORIZED, - detail, + "无效的用户令牌", false, - )) + )); } }; let Some(user_id) = claims.get("user_id").and_then(serde_json::Value::as_str) else { @@ -251,9 +155,9 @@ pub(crate) async fn resolve_authenticated_local_user( )) } Err(err) => { - return Err(build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth user lookup failed: {err:?}"), + return Err(build_auth_internal_error_response( + "auth_session_user_lookup_failed", + err, false, )) } @@ -280,9 +184,9 @@ pub(crate) async fn resolve_authenticated_local_user( let Some(session) = (match state.find_user_session(user_id, session_id).await { Ok(value) => value, Err(err) => { - return Err(build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth session lookup failed: {err:?}"), + return Err(build_auth_internal_error_response( + "auth_session_lookup_failed", + err, false, )) } @@ -293,7 +197,10 @@ pub(crate) async fn resolve_authenticated_local_user( false, )); }; - if session.is_revoked() || session.is_expired(now) { + if session.is_revoked() + || session.is_expired(now) + || session.security_version != user.security_version + { return Err(build_auth_error_response( http::StatusCode::UNAUTHORIZED, "登录会话已失效,请重新登录", @@ -341,18 +248,18 @@ pub(crate) async fn handle_auth_me( let feature_settings = match state.read_user_feature_settings(&auth.user.id).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user feature settings lookup failed: {err:?}"), + return build_auth_internal_error_response( + "auth_user_feature_settings_lookup_failed", + err, false, ) } }; - build_auth_json_response( + mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, build_auth_me_payload(&auth.user, wallet.as_ref(), feature_settings), None, - ) + )) } pub(super) async fn handle_auth_refresh( @@ -374,15 +281,17 @@ pub(super) async fn handle_auth_refresh( let Some(refresh_token) = extract_cookie_value(headers, &cookie_name) else { return build_auth_error_response(http::StatusCode::UNAUTHORIZED, "缺少刷新令牌", true); }; - let claims = match decode_auth_token(&refresh_token, "refresh") { + let claims = match decode_auth_token(&refresh_token, LocalAuthTokenType::Refresh) { Ok(value) => value, Err(detail) => { - let detail = if detail == "Token已过期" || detail == "无效的Token" { - "刷新令牌失败".to_string() - } else { - detail - }; - return build_auth_error_response(http::StatusCode::UNAUTHORIZED, detail, true); + if auth_token_error_is_internal(&detail) { + return build_auth_internal_error_response( + "auth_refresh_token_decode_failed", + detail, + true, + ); + } + return build_auth_error_response(http::StatusCode::UNAUTHORIZED, "刷新令牌失败", true); } }; let Some(user_id) = claims.get("user_id").and_then(serde_json::Value::as_str) else { @@ -401,11 +310,7 @@ pub(super) async fn handle_auth_refresh( ) } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth user lookup failed: {err:?}"), - true, - ) + return build_auth_internal_error_response("auth_refresh_user_lookup_failed", err, true) } }; if !user.is_active { @@ -417,7 +322,9 @@ pub(super) async fn handle_auth_refresh( if !auth_token_identity_matches_user(&claims, &user) { return build_auth_error_response(http::StatusCode::UNAUTHORIZED, "无效的刷新令牌", true); } - let client_device_id = match extract_client_device_id(request_context, headers) { + // Refresh is cookie-authenticated. Requiring a non-simple custom header keeps + // cross-site forms from rotating a victim's session via a query parameter. + let client_device_id = match extract_client_device_id_header(headers) { Ok(value) => value, Err(response) => return response, }; @@ -425,9 +332,9 @@ pub(super) async fn handle_auth_refresh( let Some(session) = (match state.find_user_session(user_id, session_id).await { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth session lookup failed: {err:?}"), + return build_auth_internal_error_response( + "auth_refresh_session_lookup_failed", + err, true, ) } @@ -438,7 +345,10 @@ pub(super) async fn handle_auth_refresh( true, ); }; - if session.is_revoked() || session.is_expired(now) { + if session.is_revoked() + || session.is_expired(now) + || session.security_version != user.security_version + { return build_auth_error_response( http::StatusCode::UNAUTHORIZED, "登录会话已失效,请重新登录", @@ -454,19 +364,43 @@ pub(super) async fn handle_auth_refresh( } let (is_valid, is_prev) = session.verify_refresh_token(&refresh_token, now); if !is_valid { - let _ = state + match state .revoke_user_session(user_id, session_id, now, "refresh_token_reused") - .await; + .await + { + Ok(true) => {} + Ok(false) => { + return build_auth_internal_error_response( + "auth_refresh_replay_revoke_failed", + "refresh-token replay session was not revoked", + true, + ) + } + Err(err) => { + return build_auth_internal_error_response( + "auth_refresh_replay_revoke_failed", + err, + true, + ) + } + } return build_auth_error_response( http::StatusCode::UNAUTHORIZED, "登录会话已失效,请重新登录", true, ); } + if is_prev { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "刷新令牌已轮换,请重试请求", + false, + ); + } let access_expires_at = now + chrono::Duration::hours(auth_access_token_expiry_hours()); let access_token = match create_auth_token( - "access", + LocalAuthTokenType::Access, serde_json::Map::from_iter([ ("user_id".to_string(), json!(user.id)), ("role".to_string(), json!(user.role)), @@ -480,51 +414,66 @@ pub(super) async fn handle_auth_refresh( ) { Ok(value) => value, Err(detail) => { - return build_auth_error_response(http::StatusCode::INTERNAL_SERVER_ERROR, detail, true) + return build_auth_internal_error_response( + "auth_refresh_access_token_create_failed", + detail, + true, + ) } }; - let mut set_cookie = None; - if !is_prev { - let new_refresh_token = match create_auth_token( - "refresh", - serde_json::Map::from_iter([ - ("user_id".to_string(), json!(user.id)), - ( - "created_at".to_string(), - json!(user.created_at.map(|value| value.to_rfc3339())), - ), - ("session_id".to_string(), json!(session.id)), - ("jti".to_string(), json!(uuid::Uuid::new_v4().to_string())), - ]), + let new_refresh_token = match create_auth_token( + LocalAuthTokenType::Refresh, + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!(session.id)), + ("jti".to_string(), json!(uuid::Uuid::new_v4().to_string())), + ]), + now + chrono::Duration::days(AUTH_REFRESH_TOKEN_EXPIRATION_DAYS), + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_internal_error_response( + "auth_refresh_token_create_failed", + detail, + true, + ) + } + }; + let rotated = state + .rotate_user_session_refresh_token( + user_id, + session_id, + &session.refresh_token_hash, + &GatewayUserSessionView::hash_refresh_token(&new_refresh_token), + now, now + chrono::Duration::days(AUTH_REFRESH_TOKEN_EXPIRATION_DAYS), - ) { - Ok(value) => value, - Err(detail) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - detail, - true, - ) - } - }; - let rotated = state - .rotate_user_session_refresh_token( - user_id, - session_id, - &session.refresh_token_hash, - &GatewayUserSessionView::hash_refresh_token(&new_refresh_token), - now, - now + chrono::Duration::days(AUTH_REFRESH_TOKEN_EXPIRATION_DAYS), - None, - auth_user_agent(headers).as_deref(), + None, + auth_user_agent(headers).as_deref(), + ) + .await; + match rotated { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "刷新令牌已轮换,请重试请求", + false, + ) + } + Err(err) => { + return build_auth_internal_error_response( + "auth_refresh_token_rotation_failed", + err, + true, ) - .await; - if rotated.ok() != Some(true) { - return build_auth_error_response(http::StatusCode::UNAUTHORIZED, "刷新令牌失败", true); } - set_cookie = Some(build_auth_refresh_cookie_header(&new_refresh_token)); } + let set_cookie = Some(build_auth_refresh_cookie_header(&new_refresh_token)); build_auth_json_response( http::StatusCode::OK, @@ -540,23 +489,23 @@ pub(super) async fn handle_auth_refresh( pub(crate) async fn build_auth_login_success_response( state: &AppState, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, client_device_id: String, user: aether_data::repository::users::StoredUserAuthRecord, + expected_password_hash: Option<&str>, ) -> Response { let now = auth_now(); - if let Err(err) = state.touch_auth_user_last_login(&user.id, now).await { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth last login update failed: {err:?}"), - false, - ); + if expected_password_hash.is_none() { + if let Err(err) = state.touch_auth_user_last_login(&user.id, now).await { + return build_auth_internal_error_response("auth_last_login_update_failed", err, false); + } } let session_id = Uuid::new_v4().to_string(); let access_expires_at = now + chrono::Duration::hours(auth_access_token_expiry_hours()); let refresh_expires_at = now + chrono::Duration::days(AUTH_REFRESH_TOKEN_EXPIRATION_DAYS); let access_token = match create_auth_token( - "access", + LocalAuthTokenType::Access, serde_json::Map::from_iter([ ("user_id".to_string(), json!(user.id.clone())), ("role".to_string(), json!(user.role.clone())), @@ -570,15 +519,15 @@ pub(crate) async fn build_auth_login_success_response( ) { Ok(value) => value, Err(detail) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, + return build_auth_internal_error_response( + "auth_login_access_token_create_failed", detail, false, ) } }; let refresh_token = match create_auth_token( - "refresh", + LocalAuthTokenType::Refresh, serde_json::Map::from_iter([ ("user_id".to_string(), json!(user.id.clone())), ( @@ -592,8 +541,8 @@ pub(crate) async fn build_auth_login_success_response( ) { Ok(value) => value, Err(detail) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, + return build_auth_internal_error_response( + "auth_login_refresh_token_create_failed", detail, false, ) @@ -611,33 +560,43 @@ pub(crate) async fn build_auth_login_success_response( Some(refresh_expires_at), None, None, - auth_client_ip(headers), + Some(client_ip.to_string()), auth_user_agent(headers), Some(now), Some(now), - ) { + ) + .and_then(|session| session.with_security_version(user.security_version)) + { Ok(value) => value, Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth session build failed: {err:?}"), + return build_auth_internal_error_response( + "auth_login_session_build_failed", + err, false, ) } }; - let created = match state.create_user_session(session).await { + let created_result = match expected_password_hash { + Some(expected_password_hash) => { + state + .create_user_session_if_password_matches(session, expected_password_hash) + .await + } + None => state.create_user_session(session).await, + }; + let created = match created_result { Ok(Some(session)) => session, Ok(None) => { return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - "auth session backend unavailable", + http::StatusCode::UNAUTHORIZED, + "邮箱或密码错误", false, - ) + ); } Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth session create failed: {err:?}"), + return build_auth_internal_error_response( + "auth_login_session_create_failed", + err, false, ) } @@ -665,31 +624,65 @@ async fn try_auth_logout_with_access_token( headers: &http::HeaderMap, ) -> Option> { let token = extract_bearer_token(headers)?; - let claims = decode_auth_token(&token, "access").ok()?; + let claims = decode_auth_token(&token, LocalAuthTokenType::Access).ok()?; let user_id = claims.get("user_id").and_then(serde_json::Value::as_str)?; let session_id = claims .get("session_id") .and_then(serde_json::Value::as_str)?; - let user = state.find_user_auth_by_id(user_id).await.ok().flatten()?; + let user = match state.find_user_auth_by_id(user_id).await { + Ok(Some(user)) => user, + Ok(None) => return None, + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_user_lookup_failed", + err, + true, + )) + } + }; if !user.is_active || user.is_deleted || !auth_token_identity_matches_user(&claims, &user) { return None; } let client_device_id = extract_client_device_id(request_context, headers).ok()?; let now = auth_now(); - let session = state - .find_user_session(user_id, session_id) - .await - .ok() - .flatten()?; + let session = match state.find_user_session(user_id, session_id).await { + Ok(Some(session)) => session, + Ok(None) => return None, + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_session_lookup_failed", + err, + true, + )) + } + }; if session.is_revoked() || session.is_expired(now) + || session.security_version != user.security_version || session.client_device_id != client_device_id { return None; } - let _ = state + match state .revoke_user_session(user_id, session_id, now, "user_logout") - .await; + .await + { + Ok(true) => {} + Ok(false) => { + return Some(build_auth_internal_error_response( + "auth_logout_session_revoke_failed", + "auth session was not revoked", + true, + )) + } + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_session_revoke_failed", + err, + true, + )) + } + } Some(build_auth_json_response( http::StatusCode::OK, json!({ "message": "登出成功", "success": true }), @@ -699,30 +692,76 @@ async fn try_auth_logout_with_access_token( async fn try_auth_logout_with_refresh_cookie( state: &AppState, - request_context: &GatewayPublicRequestContext, + _request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, ) -> Option> { let refresh_token = extract_cookie_value(headers, &auth_refresh_cookie_name())?; - let claims = decode_auth_token(&refresh_token, "refresh").ok()?; + let claims = decode_auth_token(&refresh_token, LocalAuthTokenType::Refresh).ok()?; let user_id = claims.get("user_id").and_then(serde_json::Value::as_str)?; let session_id = claims .get("session_id") .and_then(serde_json::Value::as_str)?; - let client_device_id = extract_client_device_id(request_context, headers).ok()?; + let user = match state.find_user_auth_by_id(user_id).await { + Ok(Some(user)) => user, + Ok(None) => return None, + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_refresh_user_lookup_failed", + err, + true, + )) + } + }; + if !user.is_active || user.is_deleted || !auth_token_identity_matches_user(&claims, &user) { + return None; + } + // The cookie fallback must not accept the device binding from the URL: a + // cross-site form can submit query parameters but cannot set this header. + let client_device_id = match extract_client_device_id_header(headers) { + Ok(value) => value, + Err(response) => return Some(response), + }; let now = auth_now(); - if let Some(session) = state - .find_user_session(user_id, session_id) - .await - .ok() - .flatten() + let session = match state.find_user_session(user_id, session_id).await { + Ok(Some(session)) => session, + Ok(None) => return None, + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_refresh_session_lookup_failed", + err, + true, + )) + } + }; + if session.is_revoked() + || session.is_expired(now) + || session.security_version != user.security_version + || session.client_device_id != client_device_id { - if !session.is_revoked() - && !session.is_expired(now) - && session.client_device_id == client_device_id - { - let _ = state - .revoke_user_session(user_id, session_id, now, "user_logout") - .await; + return None; + } + let (refresh_token_is_valid, _) = session.verify_refresh_token(&refresh_token, now); + if !refresh_token_is_valid { + return None; + } + match state + .revoke_user_session(user_id, session_id, now, "user_logout") + .await + { + Ok(true) => {} + Ok(false) => { + return Some(build_auth_internal_error_response( + "auth_logout_refresh_session_revoke_failed", + "auth session was not revoked", + true, + )) + } + Err(err) => { + return Some(build_auth_internal_error_response( + "auth_logout_refresh_session_revoke_failed", + err, + true, + )) } } Some(build_auth_json_response( diff --git a/apps/aether-gateway/src/handlers/public/support/auth_turnstile.rs b/apps/aether-gateway/src/handlers/public/support/auth_turnstile.rs index 1e963fd0c..b78a8e818 100644 --- a/apps/aether-gateway/src/handlers/public/support/auth_turnstile.rs +++ b/apps/aether-gateway/src/handlers/public/support/auth_turnstile.rs @@ -1,12 +1,13 @@ use super::{ - auth_client_ip_with_cf, build_auth_error_response, decrypt_catalog_secret_with_fallbacks, http, - system_config_bool, system_config_string, system_config_string_list, AppState, Body, Response, + build_auth_error_response, decrypt_or_migrate_system_config_secret, http, system_config_bool, + system_config_string, system_config_string_list, AppState, Body, Response, }; use serde::{Deserialize, Serialize}; use std::time::Duration; use tracing::warn; const TURNSTILE_SITEVERIFY_URL: &str = "https://challenges.cloudflare.com/turnstile/v0/siteverify"; +const MAX_TURNSTILE_RESPONSE_BYTES: usize = 64 * 1024; const TURNSTILE_TOKEN_MAX_LEN: usize = 2048; const TURNSTILE_SITEVERIFY_TIMEOUT: Duration = Duration::from_secs(10); @@ -25,7 +26,6 @@ impl AuthTurnstileAction { } } -#[derive(Debug)] struct AuthTurnstileConfig { enabled: bool, site_key: Option, @@ -33,7 +33,7 @@ struct AuthTurnstileConfig { allowed_hostnames: Vec, } -#[derive(Debug, Serialize)] +#[derive(Serialize)] struct TurnstileSiteverifyRequest<'a> { secret: &'a str, response: &'a str, @@ -74,12 +74,11 @@ impl AuthTurnstileFailure { pub(super) async fn verify_auth_turnstile( state: &AppState, - headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + client_ip: std::net::IpAddr, token: Option<&str>, action: AuthTurnstileAction, ) -> Result<(), Response> { - match verify_auth_turnstile_inner(state, headers, cf_connecting_ip, token, action).await { + match verify_auth_turnstile_inner(state, client_ip, token, action).await { Ok(()) => Ok(()), Err(err) => Err(err.into_response()), } @@ -87,8 +86,7 @@ pub(super) async fn verify_auth_turnstile( async fn verify_auth_turnstile_inner( state: &AppState, - headers: &http::HeaderMap, - cf_connecting_ip: Option<&str>, + client_ip: std::net::IpAddr, token: Option<&str>, action: AuthTurnstileAction, ) -> Result<(), AuthTurnstileFailure> { @@ -118,11 +116,11 @@ async fn verify_auth_turnstile_inner( return Err(AuthTurnstileFailure::BadRequest("人机验证失败,请重试")); } - let remoteip = auth_client_ip_with_cf(headers, cf_connecting_ip); + let remoteip = client_ip.to_string(); let siteverify_request = TurnstileSiteverifyRequest { secret: secret_key, response: token, - remoteip: remoteip.as_deref(), + remoteip: Some(remoteip.as_str()), idempotency_key: uuid::Uuid::new_v4().to_string(), }; let siteverify_url = turnstile_siteverify_url(state); @@ -139,8 +137,8 @@ async fn verify_auth_turnstile_inner( warn!("turnstile siteverify request timed out"); AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") })? - .map_err(|err| { - warn!(error = %err, "turnstile siteverify request failed"); + .map_err(|_| { + warn!("turnstile siteverify request failed"); AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") })?; if !response.status().is_success() { @@ -153,13 +151,16 @@ async fn verify_auth_turnstile_inner( "人机验证服务暂不可用,请稍后重试", )); } - let payload = response - .json::() + let body = aether_http::read_response_bytes_with_limit(response, MAX_TURNSTILE_RESPONSE_BYTES) .await - .map_err(|err| { - warn!(error = %err, "turnstile siteverify response decode failed"); + .map_err(|_| { + warn!("turnstile siteverify response read failed"); AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") })?; + let payload = serde_json::from_slice::(&body).map_err(|_| { + warn!("turnstile siteverify response decode failed"); + AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") + })?; if !payload.success { warn!( @@ -246,9 +247,17 @@ async fn read_auth_turnstile_config( AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") })?; - let secret_key = system_config_string(secret_key.as_ref()).map(|value| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value) - }); + let secret_key = match system_config_string(secret_key.as_ref()) { + Some(value) => Some( + decrypt_or_migrate_system_config_secret(state, "turnstile_secret_key", value) + .await + .map_err(|error| { + warn!(error = ?error, "turnstile secret key migration failed"); + AuthTurnstileFailure::ServiceUnavailable("人机验证服务暂不可用,请稍后重试") + })?, + ), + None => None, + }; Ok(AuthTurnstileConfig { enabled: system_config_bool(enabled.as_ref(), false), diff --git a/apps/aether-gateway/src/handlers/public/support/billing.rs b/apps/aether-gateway/src/handlers/public/support/billing.rs index 531a5aa71..07ddbcf14 100644 --- a/apps/aether-gateway/src/handlers/public/support/billing.rs +++ b/apps/aether-gateway/src/handlers/public/support/billing.rs @@ -3,12 +3,16 @@ use super::support_payment::payment_epay::{ EpayCheckoutInput, }; use super::{ - build_auth_error_response, build_auth_json_response, resolve_authenticated_local_user, - sanitize_wallet_gateway_response, unix_secs_to_rfc3339, AppState, GatewayPublicRequestContext, + build_auth_error_response, build_auth_json_response, mark_sensitive_response_no_store, + prepare_billing_gateway_response_for_storage, resolve_authenticated_local_user, + resolve_direct_gateway_channel, sanitize_wallet_gateway_response, unix_secs_to_rfc3339, + wallet_payment_instructions_from_checkout, wallet_payment_instructions_from_stored, AppState, + GatewayPublicRequestContext, }; +use crate::handlers::shared::normalize_payment_currency; use crate::handlers::shared::{ - create_alipay_direct_checkout, create_stripe_direct_checkout, create_wxpay_direct_checkout, - direct_payment_client_ip, DirectPaymentCheckoutInput, + close_direct_gateway_checkout, create_alipay_direct_checkout, create_stripe_direct_checkout, + create_wxpay_direct_checkout, DirectPaymentCheckoutError, DirectPaymentCheckoutInput, }; use axum::{ body::{Body, Bytes}, @@ -19,9 +23,15 @@ use axum::{ use chrono::Utc; use serde::Deserialize; use serde_json::{json, Value}; +use sha2::{Digest, Sha256}; +use tracing::warn; use uuid::Uuid; +use aether_data::repository::wallet::stored_timestamp_unix_secs; + const BILLING_STORAGE_UNAVAILABLE_DETAIL: &str = "套餐后端暂不可用"; +const BILLING_ORDER_IDENTITY_BUCKET_SECS: u64 = 30 * 60; +const PLAN_PAYMENT_AMOUNT_EPSILON: f64 = 0.000_001; #[derive(Debug, Deserialize, Default)] struct BillingPlanCheckoutRequest { @@ -48,17 +58,31 @@ fn billing_storage_unavailable_response() -> Response { ) } -fn normalize_optional_checkout_string(value: Option, max_len: usize) -> Option { - value - .map(|value| value.trim().to_ascii_lowercase()) - .filter(|value| !value.is_empty() && value.chars().count() <= max_len) +fn normalize_optional_checkout_string( + value: Option, + max_len: usize, +) -> Result, &'static str> { + let Some(value) = value else { + return Ok(None); + }; + let value = value.trim().to_ascii_lowercase(); + if value.is_empty() { + return Ok(None); + } + if value.chars().count() > max_len { + return Err("输入验证失败"); + } + Ok(Some(value)) } fn normalize_checkout_request( payload: BillingPlanCheckoutRequest, ) -> Result { - let payment_provider = normalize_optional_checkout_string(payload.payment_provider, 30) - .or_else(|| normalize_optional_checkout_string(payload.payment_method.clone(), 30)) + let requested_provider = normalize_optional_checkout_string(payload.payment_provider, 30)?; + let requested_method = normalize_optional_checkout_string(payload.payment_method, 30)?; + let payment_provider = requested_provider + .clone() + .or(requested_method.clone()) .unwrap_or_else(|| "epay".to_string()); if !matches!( payment_provider.as_str(), @@ -66,10 +90,32 @@ fn normalize_checkout_request( ) { return Err("unsupported payment_provider"); } - let payment_method = normalize_optional_checkout_string(payload.payment_method, 30) - .unwrap_or_else(|| payment_provider.clone()); - let payment_channel = normalize_optional_checkout_string(payload.payment_channel, 30) - .or_else(|| (payment_method != "epay").then_some(payment_method.clone())); + let payment_method = requested_method.unwrap_or_else(|| payment_provider.clone()); + let payment_channel = normalize_optional_checkout_string(payload.payment_channel, 30)?; + if payment_provider == "epay" { + if !matches!(payment_method.as_str(), "epay" | "alipay" | "wxpay") { + return Err("payment_method 与 payment_provider 不匹配"); + } + if payment_method != "epay" + && payment_channel + .as_deref() + .is_some_and(|channel| channel != payment_method) + { + return Err("payment_method 与 payment_channel 不匹配"); + } + } else if payment_method != payment_provider { + return Err("payment_method 与 payment_provider 不匹配"); + } + // Only the legacy EPay shorthand uses `payment_method` as a channel. + // Direct providers have their own channel allowlists (for example + // `native`/`h5` for WxPay and `card`/`link` for Stripe); passing the + // provider name through as a channel would make an otherwise valid + // request fail resolution before the configured default can be selected. + let payment_channel = if payment_provider == "epay" && payment_method != "epay" { + payment_channel.or_else(|| Some(payment_method.clone())) + } else { + payment_channel + }; Ok(NormalizedBillingPlanCheckoutRequest { payment_method: payment_provider.clone(), payment_provider, @@ -88,14 +134,421 @@ fn plan_id_from_checkout_path(path: &str) -> Option { } } -fn billing_order_no(now: chrono::DateTime) -> String { - format!( - "pp_{}_{}", - now.format("%Y%m%d%H%M%S%6f"), - &Uuid::new_v4().simple().to_string()[..12] +/// Derive a stable merchant order number for one checkout identity. The +/// short time bucket lets a later attempt create a fresh order after the +/// original 30-minute pending window expires, while retries in the same +/// window reuse the provider-side idempotency key (Stripe) and merchant order +/// number (EPay/Alipay/WxPay). +fn billing_order_no( + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + now: chrono::DateTime, +) -> String { + billing_order_no_with_salt( + user_id, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + None, + now, ) } +fn billing_order_no_with_salt( + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + salt: Option<&str>, + now: chrono::DateTime, +) -> String { + let bucket = now.timestamp().max(0) as u64 / BILLING_ORDER_IDENTITY_BUCKET_SECS; + let identity = json!({ + "version": 2, + "bucket": bucket, + "user_id": user_id, + "plan_id": plan_id, + "payment_method": payment_method.trim().to_ascii_lowercase(), + "payment_provider": payment_provider.trim().to_ascii_lowercase(), + "payment_channel": payment_channel.trim().to_ascii_lowercase(), + "amount_usd": format!("{amount_usd:.8}"), + "pay_amount": format!("{pay_amount:.8}"), + "pay_currency": pay_currency.trim().to_ascii_lowercase(), + "exchange_rate": format!("{exchange_rate:.8}"), + "salt": salt, + }) + .to_string(); + let digest = format!("{:x}", Sha256::digest(identity.as_bytes())); + // 3-byte prefix + 56 hex characters stays below every adapter's 64-byte + // order number limit while retaining ample collision resistance. + format!("pp_{}", &digest[..56]) +} + +fn billing_order_no_with_nonce( + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + nonce: &str, + now: chrono::DateTime, +) -> String { + billing_order_no_with_salt( + user_id, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + Some(nonce), + now, + ) +} + +fn payment_amounts_match(expected: f64, stored: f64) -> bool { + expected.is_finite() + && stored.is_finite() + && (expected - stored).abs() <= PLAN_PAYMENT_AMOUNT_EPSILON +} + +fn plan_order_metadata_string<'a>( + order: &'a aether_data::repository::wallet::StoredAdminPaymentOrder, + key: &str, +) -> Option<&'a str> { + order + .gateway_response + .as_ref() + .and_then(Value::as_object) + .and_then(|object| object.get(key)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +/// Check every payment identity component before replaying a pending plan +/// order. The repository query deliberately stays broad for compatibility +/// with older rows; this boundary check prevents a payment-method, channel, +/// currency, rate, or amount change from receiving stale instructions. +fn plan_purchase_order_matches_checkout( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, +) -> bool { + if !order.status.eq_ignore_ascii_case("pending") + || !order + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at > Utc::now().timestamp().max(0) as u64) + || !payment_amounts_match(amount_usd, order.amount_usd) + || !payment_amounts_match(pay_amount, order.pay_amount.unwrap_or(f64::NAN)) + || !order + .pay_currency + .as_deref() + .is_some_and(|stored| stored.eq_ignore_ascii_case(pay_currency.trim())) + || !payment_amounts_match(exchange_rate, order.exchange_rate.unwrap_or(f64::NAN)) + { + return false; + } + + let stored_provider = plan_order_metadata_string(order, "gateway") + .or_else(|| plan_order_metadata_string(order, "payment_provider")); + let provider_matches = stored_provider + .map(|stored| stored.eq_ignore_ascii_case(payment_provider.trim())) + .unwrap_or_else(|| { + order + .payment_method + .eq_ignore_ascii_case(payment_provider.trim()) + }); + + let stored_channel = plan_order_metadata_string(order, "payment_channel"); + let channel_matches = stored_channel + .map(|stored| stored.eq_ignore_ascii_case(payment_channel.trim())) + // Pre-provider plan rows may have stored the EPay channel as the + // payment method. Accept that legacy shape only when the requested + // provider is EPay and the channel still agrees exactly. + .unwrap_or_else(|| { + payment_provider.eq_ignore_ascii_case("epay") + && order + .payment_method + .eq_ignore_ascii_case(payment_channel.trim()) + }); + + let method_matches = order + .payment_method + .eq_ignore_ascii_case(payment_method.trim()) + || (payment_provider.eq_ignore_ascii_case("epay") + && payment_method.eq_ignore_ascii_case("epay") + && order + .payment_method + .eq_ignore_ascii_case(payment_channel.trim())); + + let product_matches = plan_order_metadata_string(order, "product_id") + .map(|stored| stored == plan_id) + .unwrap_or(true); + + provider_matches && channel_matches && method_matches && product_matches +} + +async fn find_matching_pending_plan_purchase_order( + state: &AppState, + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, +) -> Result, String> { + let order = state + .find_pending_plan_purchase_order_by_user_id(user_id, plan_id) + .await + .map_err(|err| format!("pending billing checkout lookup failed: {err:?}"))?; + Ok(order.filter(|order| { + plan_purchase_order_matches_checkout( + order, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + ) + })) +} + +enum PlanOrderNoChoice { + Reuse(aether_data::repository::wallet::StoredAdminPaymentOrder), + Fresh(String), +} + +/// Reserve a merchant order number identity before contacting a provider. +/// Deterministic retries keep their original number, while a terminal (or +/// unrelated) row occupying that number receives a fresh high-entropy number. +async fn choose_plan_order_no( + state: &AppState, + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + now: chrono::DateTime, +) -> Result { + let base = billing_order_no( + user_id, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + now, + ); + let existing = state + .find_payment_order_by_order_no(&base) + .await + .map_err(|err| format!("billing order number lookup failed: {err:?}"))?; + if let Some(existing) = existing { + if existing.user_id.as_deref() == Some(user_id) + && plan_purchase_order_matches_checkout( + &existing, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + ) + { + return Ok(PlanOrderNoChoice::Reuse(existing)); + } + + // A deterministic identity may already belong to a completed order. + // Keep trying bounded random identities; the database uniqueness check + // remains the final authority for a concurrent writer. + for _ in 0..4 { + let nonce = Uuid::new_v4().simple().to_string(); + let candidate = billing_order_no_with_nonce( + user_id, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + &nonce, + now, + ); + if state + .find_payment_order_by_order_no(&candidate) + .await + .map_err(|err| format!("billing order number lookup failed: {err:?}"))? + .is_none() + { + return Ok(PlanOrderNoChoice::Fresh(candidate)); + } + } + return Err("无法生成唯一支付订单号,请稍后重试".to_string()); + } + Ok(PlanOrderNoChoice::Fresh(base)) +} + +enum PlanCreateFailureResolution { + Replay(Response), + Missing, + Occupied, +} + +async fn fresh_plan_order_no( + state: &AppState, + user_id: &str, + plan_id: &str, + payment_method: &str, + payment_provider: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + now: chrono::DateTime, +) -> Result { + for _ in 0..4 { + let nonce = Uuid::new_v4().simple().to_string(); + let candidate = billing_order_no_with_nonce( + user_id, + plan_id, + payment_method, + payment_provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + &nonce, + now, + ); + if state + .find_payment_order_by_order_no(&candidate) + .await + .map_err(|err| format!("billing order number lookup failed: {err:?}"))? + .is_none() + { + return Ok(candidate); + } + } + Err("无法生成唯一支付订单号,请稍后重试".to_string()) +} + +async fn resolve_plan_order_after_create_failure( + state: &AppState, + user_id: &str, + plan: &aether_data_contracts::repository::billing::BillingPlanRecord, + provider: &str, + payment_method: &str, + payment_channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, + order_no: &str, +) -> Result { + let Some(order) = state + .find_payment_order_by_order_no(order_no) + .await + .map_err(|err| format!("billing order reconciliation failed: {err:?}"))? + else { + return Ok(PlanCreateFailureResolution::Missing); + }; + if order.user_id.as_deref() != Some(user_id) + || !plan_purchase_order_matches_checkout( + &order, + &plan.id, + payment_method, + provider, + payment_channel, + amount_usd, + pay_amount, + pay_currency, + exchange_rate, + ) + { + return Ok(PlanCreateFailureResolution::Occupied); + } + Ok(PlanCreateFailureResolution::Replay( + plan_checkout_replay_response(state, &order, plan).await, + )) +} + +async fn close_abandoned_direct_plan_checkout( + state: &AppState, + provider: &str, + order_no: &str, + gateway_order_id: Option<&str>, + failure_stage: &'static str, +) { + match close_direct_gateway_checkout(state, provider, order_no, gateway_order_id).await { + Ok(Some(_)) => warn!( + payment_provider = provider, + order_no, + failure_stage, + "closed direct gateway checkout after local plan checkout failure" + ), + Ok(None) => warn!( + payment_provider = provider, + order_no, failure_stage, "direct gateway does not support checkout compensation" + ), + Err(error) => warn!( + payment_provider = provider, + order_no, + failure_stage, + error, + "failed to close direct gateway checkout after local plan checkout failure" + ), + } +} + fn billing_plan_payload( record: &aether_data_contracts::repository::billing::BillingPlanRecord, ) -> serde_json::Value { @@ -141,7 +594,7 @@ fn plan_has_package_rights( items.iter().any(|item| { matches!( item.get("type").and_then(|value| value.as_str()), - Some("daily_quota" | "membership_group") + Some("daily_quota" | "membership_group" | "usage_policy") ) }) }) @@ -151,6 +604,18 @@ fn payment_order_payload( record: &aether_data::repository::wallet::StoredAdminPaymentOrder, plan: &aether_data_contracts::repository::billing::BillingPlanRecord, ) -> serde_json::Value { + // A gateway response is a live checkout capability, not durable order + // history. Once the order is paid, terminal, or expired, suppress URLs, + // form parameters, and provider metadata from the public payload. + let gateway_response = if record.status.eq_ignore_ascii_case("pending") + && record + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at > Utc::now().timestamp().max(0) as u64) + { + record.gateway_response.clone() + } else { + None + }; json!({ "id": record.id, "order_no": record.order_no, @@ -162,18 +627,35 @@ fn payment_order_payload( "exchange_rate": record.exchange_rate, "payment_method": record.payment_method, "gateway_order_id": record.gateway_order_id, - "gateway_response": sanitize_wallet_gateway_response(record.gateway_response.clone()), + "gateway_response": sanitize_wallet_gateway_response(gateway_response), "status": record.status, "order_kind": "plan_purchase", "product_id": plan.id, "product": billing_plan_payload(plan), - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "paid_at": record.paid_at_unix_secs.and_then(unix_secs_to_rfc3339), "credited_at": record.credited_at_unix_secs.and_then(unix_secs_to_rfc3339), "expires_at": record.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), }) } +async fn plan_checkout_replay_response( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + plan: &aether_data_contracts::repository::billing::BillingPlanRecord, +) -> Response { + let payment_instructions = wallet_payment_instructions_from_stored(state, order).await; + mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": payment_order_payload(order, plan), + "payment_instructions": payment_instructions, + "reused_pending_order": true, + }), + None, + )) +} + fn entitlement_payload( record: &aether_data_contracts::repository::billing::UserPlanEntitlementRecord, ) -> serde_json::Value { @@ -196,18 +678,55 @@ fn compute_plan_payment_amounts( pay_currency: &str, usd_exchange_rate: f64, ) -> Result<(f64, f64), &'static str> { - if !plan.price_amount.is_finite() || plan.price_amount <= 0.0 || usd_exchange_rate <= 0.0 { + if !plan.price_amount.is_finite() + || plan.price_amount <= 0.0 + || !usd_exchange_rate.is_finite() + || usd_exchange_rate <= 0.0 + { return Err("套餐价格配置无效"); } - if plan.price_currency.eq_ignore_ascii_case(pay_currency) { + let plan_currency = normalize_payment_currency(&plan.price_currency, "price_currency") + .map_err(|_| "套餐币种配置无效")?; + let pay_currency = normalize_payment_currency(pay_currency, "pay_currency") + .map_err(|_| "支付网关币种配置无效")?; + // USD is the canonical settlement currency. A USD-priced plan paid in + // USD must not be divided by the configured non-USD conversion rate (the + // default is 7.2), even though the two normalized currencies are equal. + if plan_currency == "USD" && pay_currency == "USD" { + let amount_usd = (plan.price_amount * 100_000_000.0).round() / 100_000_000.0; + let pay_amount = (plan.price_amount * 100.0).round() / 100.0; + if !amount_usd.is_finite() + || amount_usd <= 0.0 + || !pay_amount.is_finite() + || pay_amount <= 0.0 + { + return Err("套餐价格配置无效"); + } + return Ok((amount_usd, pay_amount)); + } + if plan_currency == pay_currency { let amount_usd = (plan.price_amount / usd_exchange_rate * 100_000_000.0).round() / 100_000_000.0; let pay_amount = (plan.price_amount * 100.0).round() / 100.0; + if !amount_usd.is_finite() + || amount_usd <= 0.0 + || !pay_amount.is_finite() + || pay_amount <= 0.0 + { + return Err("套餐价格配置无效"); + } return Ok((amount_usd, pay_amount)); } - if plan.price_currency.eq_ignore_ascii_case("USD") { + if plan_currency == "USD" { let amount_usd = (plan.price_amount * 100_000_000.0).round() / 100_000_000.0; let pay_amount = (plan.price_amount * usd_exchange_rate * 100.0).round() / 100.0; + if !amount_usd.is_finite() + || amount_usd <= 0.0 + || !pay_amount.is_finite() + || pay_amount <= 0.0 + { + return Err("套餐价格配置无效"); + } return Ok((amount_usd, pay_amount)); } Err("套餐币种与支付网关币种不匹配") @@ -273,6 +792,7 @@ pub(super) async fn handle_billing_plan_checkout( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, request_body: Option<&Bytes>, ) -> Response { let auth = match resolve_authenticated_local_user(state, request_context, headers).await { @@ -327,47 +847,8 @@ pub(super) async fn handle_billing_plan_checkout( false, ); } - match state - .find_pending_plan_purchase_order_by_user_id(&auth.user.id, &plan.id) - .await - { - Ok(Some(order)) => { - return build_auth_json_response( - http::StatusCode::OK, - json!({ - "order": payment_order_payload(&order, &plan), - "payment_instructions": sanitize_wallet_gateway_response( - order.gateway_response.clone() - ), - "reused_pending_order": true, - }), - None, - ) - } - Ok(None) => {} - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("pending billing checkout lookup failed: {err:?}"), - false, - ) - } - } - let now = Utc::now(); - let order_no = billing_order_no(now); - let expires_at = now + chrono::Duration::minutes(30); let requested_provider = checkout_request.payment_provider.as_str(); let payment_method = checkout_request.payment_method.clone(); - let payment_channel = - checkout_request - .payment_channel - .clone() - .or_else(|| match requested_provider { - "alipay" => Some("alipay".to_string()), - "wxpay" => Some("native".to_string()), - "stripe" => Some("card".to_string()), - _ => None, - }); if requested_provider == "epay" { let config = match load_epay_config(state).await { Ok(value) => value, @@ -393,18 +874,70 @@ pub(super) async fn handle_billing_plan_checkout( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) } }; - let Some(callback_base_url) = epay_callback_base_url( - config.callback_base_url.as_deref(), - headers, - request_context, - ) else { + let pending_order = match find_matching_pending_plan_purchase_order( + state, + &auth.user.id, + &plan.id, + &payment_method, + requested_provider, + &payment_channel_id, + amount_usd, + pay_amount, + &config.pay_currency, + config.usd_exchange_rate, + ) + .await + { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + if let Some(order) = pending_order { + return plan_checkout_replay_response(state, &order, &plan).await; + } + let now = Utc::now(); + let expires_at = now + chrono::Duration::minutes(30); + let order_no = match choose_plan_order_no( + state, + &auth.user.id, + &plan.id, + &payment_method, + requested_provider, + &payment_channel_id, + amount_usd, + pay_amount, + &config.pay_currency, + config.usd_exchange_rate, + now, + ) + .await + { + Ok(PlanOrderNoChoice::Fresh(order_no)) => order_no, + Ok(PlanOrderNoChoice::Reuse(order)) => { + return plan_checkout_replay_response(state, &order, &plan).await + } + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + let Some(callback_base_url) = epay_callback_base_url(config.callback_base_url.as_deref()) + else { return build_auth_error_response( http::StatusCode::BAD_REQUEST, "epay callback_base_url is required", false, ); }; - let checkout = build_epay_checkout_url( + let checkout = match build_epay_checkout_url( &config, &EpayCheckoutInput { order_no: order_no.clone(), @@ -414,7 +947,28 @@ pub(super) async fn handle_billing_plan_checkout( notify_url: format!("{callback_base_url}/api/payment/epay/notify"), return_url: format!("{callback_base_url}/api/payment/epay/return"), }, - ); + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } + }; + let stored_gateway_response = match prepare_billing_gateway_response_for_storage( + state, + "epay", + &order_no, + &auth.user.id, + &checkout, + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; let outcome = match state .create_plan_purchase_order( aether_data::repository::wallet::CreatePlanPurchaseOrderInput { @@ -426,9 +980,9 @@ pub(super) async fn handle_billing_plan_checkout( exchange_rate: config.usd_exchange_rate, payment_method: payment_method.clone(), payment_provider: Some(checkout_request.payment_provider.clone()), - payment_channel: Some(payment_channel_id), + payment_channel: Some(payment_channel_id.clone()), gateway_order_id: order_no.clone(), - gateway_response: checkout.clone(), + gateway_response: stored_gateway_response, order_no: order_no.clone(), product_id: plan.id.clone(), product_snapshot: billing_plan_snapshot(&plan), @@ -440,11 +994,37 @@ pub(super) async fn handle_billing_plan_checkout( Ok(Some(value)) => value, Ok(None) => return billing_storage_unavailable_response(), Err(err) => { + match resolve_plan_order_after_create_failure( + state, + &auth.user.id, + &plan, + requested_provider, + &payment_method, + &payment_channel_id, + amount_usd, + pay_amount, + &config.pay_currency, + config.usd_exchange_rate, + &order_no, + ) + .await + { + Ok(PlanCreateFailureResolution::Replay(response)) => return response, + Ok( + PlanCreateFailureResolution::Missing + | PlanCreateFailureResolution::Occupied, + ) => {} + Err(reconcile_error) => warn!( + order_no, + error = reconcile_error, + "failed to reconcile EPay plan order after create error" + ), + } return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("billing checkout create failed: {err:?}"), false, - ) + ); } }; let order = match outcome { @@ -466,14 +1046,14 @@ pub(super) async fn handle_billing_plan_checkout( ) } }; - build_auth_json_response( + mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "order": order, "payment_instructions": sanitize_wallet_gateway_response(Some(checkout)), }), None, - ) + )) } else { let (payment_channel, display_name, pay_currency, usd_exchange_rate, callback_base_url) = { let record = match state.find_payment_gateway_config(requested_provider).await { @@ -493,34 +1073,33 @@ pub(super) async fn handle_billing_plan_checkout( ) } }; - let payment_channel = - payment_channel - .clone() - .unwrap_or_else(|| match requested_provider { - "alipay" => "alipay".to_string(), - "wxpay" => "native".to_string(), - "stripe" => "card".to_string(), - _ => "alipay".to_string(), - }); - let display_name = match requested_provider { - "alipay" => "支付宝官方".to_string(), - "wxpay" => match payment_channel.as_str() { - "h5" => "微信 H5".to_string(), - "jsapi" => "微信 JSAPI".to_string(), - _ => "微信 Native".to_string(), - }, - "stripe" => match payment_channel.as_str() { - "alipay" => "Stripe Alipay".to_string(), - "wechat_pay" => "Stripe WeChat Pay".to_string(), - "link" => "Stripe Link".to_string(), - _ => "Stripe Card".to_string(), - }, - _ => "支付".to_string(), + let resolved_channel = match resolve_direct_gateway_channel( + requested_provider, + &record, + checkout_request.payment_channel.as_deref(), + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } }; + let payment_channel = resolved_channel.channel; + let display_name = resolved_channel.display_name; + let pay_currency = + match normalize_payment_currency(&record.pay_currency, "pay_currency") { + Ok(value) => value, + Err(_) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "支付网关币种配置无效", + false, + ) + } + }; ( payment_channel, display_name, - record.pay_currency, + pay_currency, record.usd_exchange_rate, record.callback_base_url, ) @@ -532,52 +1111,184 @@ pub(super) async fn handle_billing_plan_checkout( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) } }; - let Some(callback_base_url) = - epay_callback_base_url(callback_base_url.as_deref(), headers, request_context) - else { + let pending_order = match find_matching_pending_plan_purchase_order( + state, + &auth.user.id, + &plan.id, + &payment_method, + requested_provider, + &payment_channel, + amount_usd, + pay_amount, + &pay_currency, + usd_exchange_rate, + ) + .await + { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + if let Some(order) = pending_order { + return plan_checkout_replay_response(state, &order, &plan).await; + } + let now = Utc::now(); + let expires_at = now + chrono::Duration::minutes(30); + let mut order_no = match choose_plan_order_no( + state, + &auth.user.id, + &plan.id, + &payment_method, + requested_provider, + &payment_channel, + amount_usd, + pay_amount, + &pay_currency, + usd_exchange_rate, + now, + ) + .await + { + Ok(PlanOrderNoChoice::Fresh(order_no)) => order_no, + Ok(PlanOrderNoChoice::Reuse(order)) => { + return plan_checkout_replay_response(state, &order, &plan).await + } + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + let Some(callback_base_url) = epay_callback_base_url(callback_base_url.as_deref()) else { return build_auth_error_response( http::StatusCode::BAD_REQUEST, "支付网关 callback_base_url is required", false, ); }; - let direct_input = DirectPaymentCheckoutInput { - payment_channel: payment_channel.clone(), - display_name, - order_no: order_no.clone(), - subject: plan.title.clone(), - pay_amount, - pay_currency: pay_currency.clone(), - notify_url: format!("{callback_base_url}/api/payment/{requested_provider}/notify"), - return_url: Some(format!("{callback_base_url}/dashboard/billing")), - client_ip: direct_payment_client_ip(headers), - expires_at, + let mut checkout_attempt = 0_u8; + let checkout = loop { + let direct_input = DirectPaymentCheckoutInput { + payment_channel: payment_channel.clone(), + display_name: display_name.clone(), + order_no: order_no.clone(), + subject: plan.title.clone(), + pay_amount, + pay_currency: pay_currency.clone(), + notify_url: format!("{callback_base_url}/api/payment/{requested_provider}/notify"), + return_url: Some(format!("{callback_base_url}/dashboard/billing")), + client_ip: Some(client_ip.to_string()), + expires_at, + }; + let checkout_result: Result = + match requested_provider { + "alipay" => create_alipay_direct_checkout(state, &direct_input).await, + "wxpay" => create_wxpay_direct_checkout(state, &direct_input).await, + "stripe" => create_stripe_direct_checkout(state, &direct_input).await, + _ => Err(DirectPaymentCheckoutError::Failed( + "unsupported payment provider".to_string(), + )), + }; + match checkout_result { + Ok(value) => break value, + Err(DirectPaymentCheckoutError::Canceled) + if requested_provider == "stripe" && checkout_attempt == 0 => + { + // A prior failed local write may have cancelled the + // PaymentIntent retained by Stripe for this idempotency + // key. Move to a fresh merchant identity before retrying; + // otherwise Stripe would return the same unusable intent. + checkout_attempt = 1; + order_no = match fresh_plan_order_no( + state, + &auth.user.id, + &plan.id, + &payment_method, + requested_provider, + &payment_channel, + amount_usd, + pay_amount, + &pay_currency, + usd_exchange_rate, + now, + ) + .await + { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + } + Err(DirectPaymentCheckoutError::Canceled) => { + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "Stripe PaymentIntent 已取消,请稍后重试", + false, + ) + } + Err(DirectPaymentCheckoutError::Uncertain(detail)) => { + // A provider may accept the merchant order before the + // client loses or rejects its response. Alipay and WxPay + // can still be closed by merchant order number; Stripe + // remains retry-recoverable through its deterministic + // idempotency key when no intent id was received. + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + None, + "provider_checkout_failed", + ) + .await; + return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false); + } + Err(DirectPaymentCheckoutError::Failed(detail)) => { + return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false); + } + } }; - let checkout = match requested_provider { - "alipay" => match create_alipay_direct_checkout(state, &direct_input).await { - Ok(value) => value, - Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false) - } - }, - "wxpay" => match create_wxpay_direct_checkout(state, &direct_input).await { - Ok(value) => value, - Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false) - } - }, - "stripe" => match create_stripe_direct_checkout(state, &direct_input).await { - Ok(value) => value, - Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false) - } - }, - _ => { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "unsupported payment provider", - false, + let checkout_gateway_order_id = checkout + .get("gateway_order_id") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let payment_instructions = + wallet_payment_instructions_from_checkout(requested_provider, &checkout); + let stored_gateway_response = match prepare_billing_gateway_response_for_storage( + state, + requested_provider, + &order_no, + &auth.user.id, + &checkout, + ) { + Ok(value) => value, + Err(detail) => { + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + checkout_gateway_order_id.as_deref(), + "gateway_response_projection_failed", ) + .await; + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ); } }; let outcome = match state @@ -589,15 +1300,13 @@ pub(super) async fn handle_billing_plan_checkout( pay_amount, pay_currency: pay_currency.clone(), exchange_rate: usd_exchange_rate, - payment_method, + payment_method: payment_method.clone(), payment_provider: Some(requested_provider.to_string()), payment_channel: Some(payment_channel.clone()), - gateway_order_id: checkout - .get("gateway_order_id") - .and_then(Value::as_str) - .unwrap_or(&order_no) - .to_string(), - gateway_response: checkout.clone(), + gateway_order_id: checkout_gateway_order_id + .clone() + .unwrap_or_else(|| order_no.clone()), + gateway_response: stored_gateway_response, order_no: order_no.clone(), product_id: plan.id.clone(), product_snapshot: billing_plan_snapshot(&plan), @@ -607,13 +1316,62 @@ pub(super) async fn handle_billing_plan_checkout( .await { Ok(Some(value)) => value, - Ok(None) => return billing_storage_unavailable_response(), + Ok(None) => { + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + checkout_gateway_order_id.as_deref(), + "billing_storage_unavailable", + ) + .await; + return billing_storage_unavailable_response(); + } Err(err) => { + let create_error = format!("billing checkout create failed: {err:?}"); + match resolve_plan_order_after_create_failure( + state, + &auth.user.id, + &plan, + requested_provider, + &payment_method, + &payment_channel, + amount_usd, + pay_amount, + &pay_currency, + usd_exchange_rate, + &order_no, + ) + .await + { + Ok(PlanCreateFailureResolution::Replay(response)) => return response, + Ok(PlanCreateFailureResolution::Missing) => { + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + checkout_gateway_order_id.as_deref(), + "billing_order_create_failed_no_local_order", + ) + .await; + } + Ok(PlanCreateFailureResolution::Occupied) => { + warn!( + order_no, + "direct plan checkout create conflicted with an existing order; external checkout was left untouched" + ); + } + Err(reconcile_error) => warn!( + order_no, + error = reconcile_error, + "could not determine whether direct plan checkout was persisted" + ), + } return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, - format!("billing checkout create failed: {err:?}"), + create_error, false, - ) + ); } }; let order = match outcome { @@ -621,6 +1379,14 @@ pub(super) async fn handle_billing_plan_checkout( payment_order_payload(&order, &plan) } aether_data::repository::wallet::CreatePlanPurchaseOrderOutcome::WalletInactive => { + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + checkout_gateway_order_id.as_deref(), + "wallet_inactive", + ) + .await; return build_auth_error_response( http::StatusCode::BAD_REQUEST, "wallet is not active", @@ -628,6 +1394,14 @@ pub(super) async fn handle_billing_plan_checkout( ) } aether_data::repository::wallet::CreatePlanPurchaseOrderOutcome::ActivePlanLimitReached => { + close_abandoned_direct_plan_checkout( + state, + requested_provider, + &order_no, + checkout_gateway_order_id.as_deref(), + "active_plan_limit_reached", + ) + .await; return build_auth_error_response( http::StatusCode::CONFLICT, "套餐购买限制已达到上限", @@ -635,14 +1409,14 @@ pub(super) async fn handle_billing_plan_checkout( ) } }; - build_auth_json_response( + mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "order": order, - "payment_instructions": sanitize_wallet_gateway_response(Some(checkout)), + "payment_instructions": payment_instructions, }), None, - ) + )) } } @@ -650,6 +1424,7 @@ pub(super) async fn maybe_build_local_billing_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, request_body: Option<&Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -660,12 +1435,436 @@ pub(super) async fn maybe_build_local_billing_response( Some("plans") if request_context.request_path == "/api/billing/plans" => { Some(handle_billing_plans_list(state).await) } - Some("plan_checkout") => { - Some(handle_billing_plan_checkout(state, request_context, headers, request_body).await) - } + Some("plan_checkout") => Some( + handle_billing_plan_checkout(state, request_context, headers, client_ip, request_body) + .await, + ), Some("entitlements") if request_context.request_path == "/api/billing/entitlements" => { Some(handle_billing_entitlements(state, request_context, headers).await) } _ => None, } } + +#[cfg(test)] +mod tests { + use super::super::prepare_wallet_gateway_response_for_storage; + use super::{ + billing_order_no, billing_order_no_with_nonce, compute_plan_payment_amounts, + normalize_checkout_request, plan_purchase_order_matches_checkout, + prepare_billing_gateway_response_for_storage, wallet_payment_instructions_from_stored, + AppState, BillingPlanCheckoutRequest, + }; + use aether_data_contracts::repository::{ + billing::BillingPlanRecord, wallet::StoredAdminPaymentOrder, + }; + use chrono::Utc; + use serde_json::json; + + fn request( + payment_method: Option<&str>, + payment_provider: Option<&str>, + payment_channel: Option<&str>, + ) -> BillingPlanCheckoutRequest { + BillingPlanCheckoutRequest { + payment_method: payment_method.map(str::to_string), + payment_provider: payment_provider.map(str::to_string), + payment_channel: payment_channel.map(str::to_string), + } + } + + #[test] + fn checkout_normalization_preserves_canonical_and_legacy_epay_requests() { + let canonical = + normalize_checkout_request(request(Some("EPAY"), Some("epay"), Some("Alipay"))) + .expect("canonical EPay request should normalize"); + assert_eq!(canonical.payment_method, "epay"); + assert_eq!(canonical.payment_provider, "epay"); + assert_eq!(canonical.payment_channel.as_deref(), Some("alipay")); + + let legacy = normalize_checkout_request(request(Some("alipay"), Some("epay"), None)) + .expect("legacy EPay shorthand should normalize"); + // EPay is the durable payment namespace; the legacy method selects + // only the aggregator channel. + assert_eq!(legacy.payment_method, "epay"); + assert_eq!(legacy.payment_provider, "epay"); + assert_eq!(legacy.payment_channel.as_deref(), Some("alipay")); + + let defaulted = normalize_checkout_request(request(None, None, None)) + .expect("empty checkout request should use the EPay default"); + assert_eq!(defaulted.payment_method, "epay"); + assert_eq!(defaulted.payment_provider, "epay"); + assert!(defaulted.payment_channel.is_none()); + } + + #[test] + fn checkout_normalization_rejects_provider_method_and_legacy_channel_conflicts() { + assert!(normalize_checkout_request(request(Some("alipay"), Some("stripe"), None)).is_err()); + assert!(normalize_checkout_request(request(Some("stripe"), Some("epay"), None)).is_err()); + assert!( + normalize_checkout_request(request(Some("alipay"), Some("epay"), Some("wxpay"))) + .is_err() + ); + assert!(normalize_checkout_request(request(Some("unknown"), Some("epay"), None)).is_err()); + } + + #[test] + fn checkout_normalization_leaves_direct_provider_channel_for_config_resolution() { + let wxpay = normalize_checkout_request(request(Some("wxpay"), Some("wxpay"), None)) + .expect("wxpay checkout should normalize without a channel"); + assert_eq!(wxpay.payment_provider, "wxpay"); + assert_eq!(wxpay.payment_method, "wxpay"); + assert!(wxpay.payment_channel.is_none()); + + let stripe = normalize_checkout_request(request(Some("stripe"), Some("stripe"), None)) + .expect("stripe checkout should normalize without a channel"); + assert_eq!(stripe.payment_provider, "stripe"); + assert_eq!(stripe.payment_method, "stripe"); + assert!(stripe.payment_channel.is_none()); + + let explicit = + normalize_checkout_request(request(Some("wxpay"), Some("wxpay"), Some("h5"))) + .expect("an explicit direct channel should be preserved"); + assert_eq!(explicit.payment_channel.as_deref(), Some("h5")); + } + + #[test] + fn checkout_normalization_rejects_overlong_fields_instead_of_falling_back() { + let too_long = "x".repeat(31); + assert!(normalize_checkout_request(request(None, Some(&too_long), None)).is_err()); + assert!(normalize_checkout_request(request(Some(&too_long), None, None)).is_err()); + assert!(normalize_checkout_request(request(None, None, Some(&too_long))).is_err()); + } + + fn test_plan(price_amount: f64, price_currency: &str) -> BillingPlanRecord { + BillingPlanRecord { + id: "plan-test".to_string(), + title: "Test".to_string(), + description: None, + price_amount, + price_currency: price_currency.to_string(), + duration_unit: "month".to_string(), + duration_value: 1, + enabled: true, + sort_order: 0, + max_active_per_user: 0, + purchase_limit_scope: "none".to_string(), + entitlements_json: json!({}), + created_at_unix_secs: 0, + updated_at_unix_secs: 0, + } + } + + #[test] + fn plan_payment_amounts_reject_non_finite_exchange_rate_and_results() { + assert!(compute_plan_payment_amounts(&test_plan(10.0, "USD"), "CNY", f64::NAN).is_err()); + assert!( + compute_plan_payment_amounts(&test_plan(f64::MAX, "USD"), "CNY", f64::MAX).is_err() + ); + assert!( + compute_plan_payment_amounts(&test_plan(10.0, "USD"), "CNY", 7.25) + .expect("valid plan amounts") + .0 + .is_finite() + ); + } + + #[test] + fn plan_payment_amounts_normalize_currency_at_the_public_boundary() { + let amounts = compute_plan_payment_amounts(&test_plan(10.0, " USD "), " cny ", 7.25) + .expect("trimmed ASCII currencies should be accepted"); + assert_eq!(amounts.0, 10.0); + assert_eq!(amounts.1, 72.5); + assert!(compute_plan_payment_amounts(&test_plan(10.0, "US$"), "CNY", 7.25).is_err()); + assert!(compute_plan_payment_amounts(&test_plan(10.0, "美元"), "CNY", 7.25).is_err()); + assert!(compute_plan_payment_amounts(&test_plan(10.0, "USD"), "CN", 7.25).is_err()); + } + + #[test] + fn usd_plan_paid_in_usd_ignores_non_usd_exchange_rate() { + let amounts = compute_plan_payment_amounts(&test_plan(10.0, "USD"), "USD", 7.2) + .expect("USD/USD plan checkout should be valid"); + assert_eq!(amounts, (10.0, 10.0)); + } + + fn pending_order( + payment_method: &str, + gateway_response: serde_json::Value, + ) -> StoredAdminPaymentOrder { + let now = Utc::now().timestamp().max(0) as u64; + StoredAdminPaymentOrder { + id: "order-test".to_string(), + order_no: "pp-test".to_string(), + wallet_id: "wallet-test".to_string(), + user_id: Some("user-test".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: payment_method.to_string(), + payment_provider: Some( + gateway_response + .get("gateway") + .and_then(serde_json::Value::as_str) + .unwrap_or(payment_method) + .to_string(), + ), + order_kind: "plan_purchase".to_string(), + gateway_order_id: Some("gateway-test".to_string()), + gateway_response: Some(gateway_response), + status: "pending".to_string(), + created_at_unix_ms: now.saturating_mul(1000), + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(now.saturating_add(600)), + } + } + + #[tokio::test] + async fn plan_stripe_client_secret_cannot_be_copied_to_another_order_or_wallet_kind() { + let state = AppState::new().expect("test app state"); + let client_secret = "pi_plan_source_secret_capability"; + let stored = prepare_billing_gateway_response_for_storage( + &state, + "stripe", + "pp-source", + "user-source", + &json!({ + "gateway": "stripe", + "client_secret": client_secret, + "publishable_key": "pk_test_public", + }), + ) + .expect("source plan checkout should encrypt"); + let mut source = pending_order("stripe", stored); + source.id = "plan-order-source".to_string(); + source.order_no = "pp-source".to_string(); + source.user_id = Some("user-source".to_string()); + assert_eq!( + wallet_payment_instructions_from_stored(&state, &source).await["client_secret"], + client_secret + ); + + let mut foreign = source.clone(); + foreign.id = "plan-order-foreign".to_string(); + foreign.order_no = "pp-foreign".to_string(); + assert!(wallet_payment_instructions_from_stored(&state, &foreign) + .await + .get("client_secret") + .is_none()); + + let wallet_ciphertext = prepare_wallet_gateway_response_for_storage( + &state, + "stripe", + "po-wallet-source", + "user-source", + &json!({ + "gateway": "stripe", + "client_secret": "pi_wallet_source_secret_capability", + "publishable_key": "pk_test_public", + }), + ) + .expect("wallet checkout should encrypt"); + let mut plan_with_wallet_ciphertext = pending_order("stripe", wallet_ciphertext); + plan_with_wallet_ciphertext.id = "plan-with-wallet-ciphertext".to_string(); + plan_with_wallet_ciphertext.order_no = "po-wallet-source".to_string(); + plan_with_wallet_ciphertext.user_id = Some("user-source".to_string()); + assert!( + wallet_payment_instructions_from_stored(&state, &plan_with_wallet_ciphertext) + .await + .get("client_secret") + .is_none() + ); + } + + #[test] + fn pending_plan_reuse_requires_the_complete_payment_identity() { + let order = pending_order( + "stripe", + json!({ + "gateway": "stripe", + "payment_channel": "card", + "product_id": "plan-test" + }), + ); + let matches = |order: &StoredAdminPaymentOrder| { + plan_purchase_order_matches_checkout( + order, + "plan-test", + "stripe", + "stripe", + "card", + 10.0, + 72.5, + "CNY", + 7.25, + ) + }; + assert!(matches(&order)); + + let mut changed = order.clone(); + changed.payment_method = "alipay".to_string(); + assert!(!matches(&changed)); + let mut changed = order.clone(); + changed.gateway_response = Some(json!({ + "gateway": "wxpay", + "payment_channel": "native", + "product_id": "plan-test" + })); + assert!(!matches(&changed)); + let mut changed = order.clone(); + changed.gateway_response = Some(json!({ + "gateway": "stripe", + "payment_channel": "link", + "product_id": "plan-test" + })); + assert!(!matches(&changed)); + let mut changed = order.clone(); + changed.pay_amount = Some(72.0); + assert!(!matches(&changed)); + let mut changed = order.clone(); + changed.pay_currency = Some("USD".to_string()); + assert!(!matches(&changed)); + let mut changed = order; + changed.exchange_rate = Some(7.0); + assert!(!matches(&changed)); + } + + #[test] + fn pending_plan_reuse_accepts_legacy_epay_method_only_for_same_channel() { + let order = pending_order( + "alipay", + json!({ + "gateway": "epay", + "payment_channel": "alipay", + "product_id": "plan-test" + }), + ); + assert!(plan_purchase_order_matches_checkout( + &order, + "plan-test", + "epay", + "epay", + "alipay", + 10.0, + 72.5, + "CNY", + 7.25, + )); + assert!(!plan_purchase_order_matches_checkout( + &order, + "plan-test", + "epay", + "epay", + "wxpay", + 10.0, + 72.5, + "CNY", + 7.25, + )); + } + + #[test] + fn billing_order_number_is_stable_for_retries_and_changes_with_payment_identity() { + let now = Utc::now(); + let first = billing_order_no( + "user-test", + "plan-test", + "stripe", + "stripe", + "card", + 10.0, + 72.5, + "CNY", + 7.25, + now, + ); + let retry = billing_order_no( + "user-test", + "plan-test", + "stripe", + "stripe", + "card", + 10.0, + 72.5, + "CNY", + 7.25, + now, + ); + assert_eq!(first, retry); + assert!(first.starts_with("pp_")); + assert!(first.len() <= 64); + let different_channel = billing_order_no( + "user-test", + "plan-test", + "stripe", + "stripe", + "link", + 10.0, + 72.5, + "CNY", + 7.25, + now, + ); + assert_ne!(first, different_channel); + } + + #[test] + fn cancelled_stripe_retry_uses_a_new_merchant_order_identity() { + let now = Utc::now(); + let original = billing_order_no( + "user-test", + "plan-test", + "stripe", + "stripe", + "card", + 10.0, + 72.5, + "CNY", + 7.25, + now, + ); + let recovered = billing_order_no_with_nonce( + "user-test", + "plan-test", + "stripe", + "stripe", + "card", + 10.0, + 72.5, + "CNY", + 7.25, + "stripe-cancelled-recovery", + now, + ); + assert_ne!(original, recovered); + assert!(recovered.starts_with("pp_")); + assert!(recovered.len() <= 64); + } + + #[test] + fn direct_channel_resolution_uses_the_configured_allowlist() { + let record = aether_data_contracts::repository::billing::PaymentGatewayConfigRecord { + provider: "wxpay".to_string(), + enabled: true, + endpoint_url: "https://pay.example.test".to_string(), + callback_base_url: Some("https://app.example.test".to_string()), + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some("encrypted".to_string()), + pay_currency: "CNY".to_string(), + usd_exchange_rate: 7.0, + min_recharge_usd: 1.0, + channels_json: json!([ + {"channel": "native", "display_name": "Native", "fee_rate": 0.0} + ]), + created_at_unix_secs: 0, + updated_at_unix_secs: 0, + }; + let selected = super::resolve_direct_gateway_channel("wxpay", &record, None) + .expect("configured native channel should be selected"); + assert_eq!(selected.channel, "native"); + assert!(super::resolve_direct_gateway_channel("wxpay", &record, Some("h5")).is_err()); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/install.rs b/apps/aether-gateway/src/handlers/public/support/install.rs index 67682d936..aa4d764e5 100644 --- a/apps/aether-gateway/src/handlers/public/support/install.rs +++ b/apps/aether-gateway/src/handlers/public/support/install.rs @@ -6,15 +6,19 @@ use axum::{ }; use serde::{Deserialize, Serialize}; use serde_json::json; +use sha2::{Digest, Sha256}; use super::{ build_auth_error_response, decrypt_catalog_secret_with_fallbacks, - resolve_authenticated_local_user, AppState, GatewayPublicRequestContext, + decrypt_or_migrate_auth_api_key_secret, encrypt_catalog_secret_with_fallbacks, + mark_sensitive_response_no_store, resolve_authenticated_local_user, AppState, + GatewayPublicRequestContext, }; const INSTALL_SESSION_TTL_SECS: u64 = 15 * 60; const INSTALL_SESSION_KEY_PREFIX: &str = "install:session:"; const TUNNEL_INSTALL_SESSION_KEY_PREFIX: &str = "tunnel-install:session:"; +const INSTALL_SESSION_ENVELOPE_PREFIX: &str = "aether-install-session-v1:"; const TUNNEL_INSTALL_UNIX_SCRIPT_URL: &str = "https://raw.githubusercontent.com/fawney19/Aether/refs/heads/main/apps/aether-tunnel/install.sh"; const TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL: &str = @@ -45,9 +49,11 @@ pub(crate) struct CreateApiKeyInstallSessionRequest { #[derive(Debug, Serialize, Deserialize)] struct StoredInstallSession { + install_code: String, api_key_id: String, - api_key_name: String, - api_key: String, + api_key_owner_user_id: String, + api_key_is_standalone: bool, + api_key_hash: String, base_url: String, target_cli: InstallTargetCli, target_system: InstallTargetSystem, @@ -56,7 +62,10 @@ struct StoredInstallSession { #[derive(Debug, Serialize, Deserialize)] struct StoredTunnelInstallSession { + install_code: String, aether_url: String, + management_token_snapshot: aether_data::repository::management_tokens::StoredManagementToken, + management_token_user_security_version: i64, management_token: String, node_name: String, tunnel_security: String, @@ -90,7 +99,7 @@ fn install_code_from_path(request_path: &str) -> Option<(String, bool)> { } let is_powershell = raw.ends_with(".ps1"); let code = raw.strip_suffix(".ps1").unwrap_or(raw).trim(); - (!code.is_empty()).then(|| (code.to_string(), is_powershell)) + is_valid_install_code(code).then(|| (code.to_string(), is_powershell)) } fn tunnel_install_code_from_path(request_path: &str) -> Option<(String, bool)> { @@ -103,15 +112,25 @@ fn tunnel_install_code_from_path(request_path: &str) -> Option<(String, bool)> { } let is_powershell = raw.ends_with(".ps1"); let code = raw.strip_suffix(".ps1").unwrap_or(raw).trim(); - (!code.is_empty()).then(|| (code.to_string(), is_powershell)) + is_valid_install_code(code).then(|| (code.to_string(), is_powershell)) +} + +fn is_valid_install_code(code: &str) -> bool { + code.len() == 24 + && code + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) } fn install_session_runtime_key(code: &str) -> String { - format!("{INSTALL_SESSION_KEY_PREFIX}{code}") + format!("{INSTALL_SESSION_KEY_PREFIX}sha256:{}", sha256_hex(code)) } fn tunnel_install_session_runtime_key(code: &str) -> String { - format!("{TUNNEL_INSTALL_SESSION_KEY_PREFIX}{code}") + format!( + "{TUNNEL_INSTALL_SESSION_KEY_PREFIX}sha256:{}", + sha256_hex(code) + ) } fn generate_install_code() -> String { @@ -134,39 +153,253 @@ fn generate_tunnel_encryption_key() -> String { base64::engine::general_purpose::STANDARD.encode(key) } +fn seal_install_session(state: &AppState, plaintext: &str) -> Option { + encrypt_catalog_secret_with_fallbacks(state, plaintext) + .map(|ciphertext| format!("{INSTALL_SESSION_ENVELOPE_PREFIX}{ciphertext}")) +} + +fn open_install_session(state: &AppState, stored: &str) -> Option { + let ciphertext = stored.strip_prefix(INSTALL_SESSION_ENVELOPE_PREFIX)?; + decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) +} + fn unix_secs_now() -> u64 { chrono::Utc::now().timestamp().max(0) as u64 } +fn api_key_record_is_current_for_install( + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + now_unix_secs: u64, +) -> bool { + record.is_active + && record + .expires_at_unix_secs + .is_none_or(|expires_at| expires_at > now_unix_secs) +} + +fn api_key_record_matches_current_snapshot( + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + snapshot: &aether_data::repository::auth::ResolvedAuthApiKeySnapshot, + now_unix_secs: u64, +) -> bool { + snapshot.api_key_id == record.api_key_id + && snapshot.user_id == record.user_id + && snapshot.api_key_is_standalone == record.is_standalone + && snapshot.api_key_is_active == record.is_active + && snapshot.api_key_expires_at_unix_secs == record.expires_at_unix_secs + && snapshot.user_is_active + && !snapshot.user_is_deleted + && snapshot.api_key_is_active + && !snapshot.api_key_is_locked + && snapshot + .api_key_expires_at_unix_secs + .is_none_or(|expires_at| expires_at > now_unix_secs) + && api_key_record_is_current_for_install(record, now_unix_secs) +} + +async fn api_key_record_has_current_snapshot( + state: &AppState, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + now_unix_secs: u64, +) -> Result { + let snapshot = state + .data + .read_auth_api_key_snapshot_strong(&record.user_id, &record.api_key_id, now_unix_secs) + .await + .map_err(|err| crate::GatewayError::Internal(err.to_string()))?; + Ok(snapshot.is_some_and(|snapshot| { + api_key_record_matches_current_snapshot(record, &snapshot, now_unix_secs) + })) +} + +fn install_session_matches_api_key_record( + session: &StoredInstallSession, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + now_unix_secs: u64, +) -> bool { + record.api_key_id == session.api_key_id + && record.user_id == session.api_key_owner_user_id + && record.is_standalone == session.api_key_is_standalone + && record.key_hash == session.api_key_hash + && api_key_record_is_current_for_install(record, now_unix_secs) +} + +async fn resolve_current_install_session_api_key( + state: &AppState, + session: &StoredInstallSession, + now_unix_secs: u64, +) -> Result, crate::GatewayError> { + let mut matching = state + .list_auth_api_key_export_records_by_ids(std::slice::from_ref(&session.api_key_id)) + .await? + .into_iter() + .filter(|record| record.api_key_id == session.api_key_id); + let Some(record) = matching.next() else { + return Ok(None); + }; + if matching.next().is_some() + || !install_session_matches_api_key_record(session, &record, now_unix_secs) + || !api_key_record_has_current_snapshot(state, &record, now_unix_secs).await? + { + return Ok(None); + } + let Some(ciphertext) = record + .key_encrypted + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Ok(None); + }; + let current_api_key = match decrypt_or_migrate_auth_api_key_secret(state, &record).await { + Ok(value) => value, + Err(_) => return Ok(None), + }; + if sha256_hex(¤t_api_key) != record.key_hash { + return Ok(None); + } + Ok(Some(current_api_key)) +} + +fn sha256_hex(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} + +async fn activate_tunnel_install_management_token( + state: &AppState, + session: &StoredTunnelInstallSession, + now_unix_secs: u64, +) -> Result { + if session.management_token_snapshot.permissions != Some(json!(["admin:proxy_nodes:write"])) + || session + .management_token_snapshot + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at <= now_unix_secs) + { + return Ok(false); + } + + state + .activate_management_token_if_matches(&tunnel_install_token_mutation( + session, + now_unix_secs, + )) + .await +} + +fn tunnel_install_token_mutation( + session: &StoredTunnelInstallSession, + now_unix_secs: u64, +) -> aether_data::repository::management_tokens::ActivateManagementTokenIfMatches { + aether_data::repository::management_tokens::ActivateManagementTokenIfMatches { + expected_token: session.management_token_snapshot.clone(), + token_hash: sha256_hex(&session.management_token), + expected_user_security_version: session.management_token_user_security_version, + now_unix_secs, + } +} + +async fn discard_tunnel_install_session_token( + state: &AppState, + session: &StoredTunnelInstallSession, +) { + let _ = state + .delete_inactive_management_token_if_matches(&tunnel_install_token_mutation( + session, + unix_secs_now(), + )) + .await; +} + +async fn discard_unused_tunnel_install_management_token( + state: &AppState, + record: &aether_data::repository::management_tokens::StoredManagementToken, + management_token: &str, + expected_user_security_version: i64, +) { + let _ = state + .delete_inactive_management_token_if_matches( + &aether_data::repository::management_tokens::ActivateManagementTokenIfMatches { + expected_token: record.clone(), + token_hash: sha256_hex(management_token), + expected_user_security_version, + now_unix_secs: unix_secs_now(), + }, + ) + .await; +} + pub(crate) fn base_url_from_request( headers: &http::HeaderMap, request_context: &GatewayPublicRequestContext, + remote_addr: &std::net::SocketAddr, ) -> String { if let Some(value) = std::env::var("AETHER_PUBLIC_BASE_URL") .ok() .or_else(|| std::env::var("PUBLIC_BASE_URL").ok()) - .map(|value| value.trim().trim_end_matches('/').to_string()) - .filter(|value| value.starts_with("https://") || value.starts_with("http://")) + .and_then(|value| normalize_public_base_url(&value)) { return value; } - let host = crate::headers::header_value_str(headers, "x-forwarded-host") - .or_else(|| request_context.host_header.clone()) - .map(|value| value.trim().trim_end_matches('/').to_string()) - .filter(|value| { - !value.is_empty() - && !value.contains('/') - && !value.contains('\\') - && !value.contains('@') - && !value.contains(char::is_whitespace) + base_url_from_request_metadata( + headers, + request_context.host_header.as_deref(), + crate::headers::trusted_proxy_ip(remote_addr.ip()), + ) +} + +fn base_url_from_request_metadata( + headers: &http::HeaderMap, + host_header: Option<&str>, + trusted_peer: bool, +) -> String { + if !trusted_peer { + return "http://localhost".to_string(); + } + + let host = forwarded_header_last(headers, "x-forwarded-host") + .or_else(|| { + host_header + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) }) .unwrap_or_else(|| "localhost".to_string()); - let proto = crate::headers::header_value_str(headers, "x-forwarded-proto") - .map(|value| value.trim().trim_end_matches(':').to_ascii_lowercase()) + let proto = forwarded_header_last(headers, "x-forwarded-proto") + .map(|value| value.trim_end_matches(':').to_ascii_lowercase()) .filter(|value| value == "http" || value == "https") .unwrap_or_else(|| "http".to_string()); - format!("{proto}://{host}") + normalize_public_base_url(&format!("{proto}://{host}")) + .unwrap_or_else(|| "http://localhost".to_string()) +} + +fn forwarded_header_last(headers: &http::HeaderMap, name: &str) -> Option { + headers + .get_all(name) + .iter() + .filter_map(|value| value.to_str().ok()) + .flat_map(|value| value.split(',')) + .map(str::trim) + .rfind(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn normalize_public_base_url(value: &str) -> Option { + let value = value.trim().trim_end_matches('/'); + let parsed = url::Url::parse(value).ok()?; + if !aether_http::is_https_or_loopback_http_url(&parsed) + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return None; + } + Some(parsed.as_str().trim_end_matches('/').to_string()) } fn shell_single_quote(value: &str) -> String { @@ -257,7 +490,7 @@ fn cli_binary(target_cli: InstallTargetCli) -> &'static str { } } -fn build_unix_script(session: &StoredInstallSession) -> String { +fn build_unix_script(session: &StoredInstallSession, api_key: &str) -> String { let target_cli = match session.target_cli { InstallTargetCli::ClaudeCode => "claude_code", InstallTargetCli::CodexCli => "codex_cli", @@ -409,14 +642,14 @@ say "$CLI_LABEL 已配置到 Aether。执行 $CLI_BIN --version 验证安装。" target_cli = target_cli, target_system = target_system, base_url = shell_single_quote(&session.base_url), - api_key = shell_single_quote(&session.api_key), + api_key = shell_single_quote(api_key), label = shell_single_quote(cli_label(session.target_cli)), binary = shell_single_quote(cli_binary(session.target_cli)), npm_package = shell_single_quote(npm_package(session.target_cli)), ) } -fn build_powershell_script(session: &StoredInstallSession) -> String { +fn build_powershell_script(session: &StoredInstallSession, api_key: &str) -> String { let target_cli = match session.target_cli { InstallTargetCli::ClaudeCode => "claude_code", InstallTargetCli::CodexCli => "codex_cli", @@ -523,7 +756,7 @@ Say "$CliLabel 已配置到 Aether。执行 $CliBin --version 验证安装。" InstallTargetSystem::Auto => "auto", }), base_url = powershell_single_quote(&session.base_url), - api_key = powershell_single_quote(&session.api_key), + api_key = powershell_single_quote(api_key), label = powershell_single_quote(cli_label(session.target_cli)), binary = powershell_single_quote(cli_binary(session.target_cli)), npm_package = powershell_single_quote(npm_package(session.target_cli)), @@ -534,6 +767,7 @@ pub(super) async fn handle_users_me_api_key_install_session_create( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { let auth = match resolve_authenticated_local_user(state, request_context, headers).await { @@ -572,12 +806,56 @@ pub(super) async fn handle_users_me_api_key_install_session_create( ) } }; - let Some(record) = records + let mut matching_records = records .into_iter() - .find(|record| !record.is_standalone && record.api_key_id == api_key_id) - else { + .filter(|record| !record.is_standalone && record.api_key_id == api_key_id); + let Some(record) = matching_records.next() else { return build_auth_error_response(http::StatusCode::NOT_FOUND, "API密钥不存在", false); }; + if matching_records.next().is_some() { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "API密钥记录不唯一,拒绝创建安装会话", + false, + ); + } + build_api_key_install_session_response( + state, + request_context, + headers, + remote_addr, + &record, + payload, + ) + .await +} + +pub(crate) async fn build_api_key_install_session_response( + state: &AppState, + request_context: &GatewayPublicRequestContext, + headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + payload: CreateApiKeyInstallSessionRequest, +) -> Response { + let now_unix_secs = unix_secs_now(); + match api_key_record_has_current_snapshot(state, record, now_unix_secs).await { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "该密钥已停用、锁定、过期或所有者不可用", + false, + ) + } + Err(err) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + format!("install session key snapshot validation failed: {err:?}"), + false, + ) + } + } let Some(ciphertext) = record .key_encrypted .as_deref() @@ -590,43 +868,33 @@ pub(super) async fn handle_users_me_api_key_install_session_create( false, ); }; - let Some(api_key) = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - else { + let current_api_key = match decrypt_or_migrate_auth_api_key_secret(state, record).await { + Ok(value) => value, + Err(_) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "解密或校验密钥失败", + false, + ) + } + }; + if sha256_hex(¤t_api_key) != record.key_hash { return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - "解密密钥失败", + http::StatusCode::BAD_REQUEST, + "该密钥完整性校验失败", false, ); - }; + } - build_api_key_install_session_response( - state, - request_context, - headers, - record.api_key_id.clone(), - record.name.unwrap_or_else(|| "API Key".to_string()), - api_key, - payload, - ) - .await -} - -pub(crate) async fn build_api_key_install_session_response( - state: &AppState, - request_context: &GatewayPublicRequestContext, - headers: &http::HeaderMap, - api_key_id: String, - api_key_name: String, - api_key: String, - payload: CreateApiKeyInstallSessionRequest, -) -> Response { let code = generate_install_code(); - let expires_at_unix_secs = unix_secs_now().saturating_add(INSTALL_SESSION_TTL_SECS); + let expires_at_unix_secs = now_unix_secs.saturating_add(INSTALL_SESSION_TTL_SECS); let session = StoredInstallSession { - api_key_id, - api_key_name, - api_key, - base_url: base_url_from_request(headers, request_context), + install_code: code.clone(), + api_key_id: record.api_key_id.clone(), + api_key_owner_user_id: record.user_id.clone(), + api_key_is_standalone: record.is_standalone, + api_key_hash: record.key_hash.clone(), + base_url: base_url_from_request(headers, request_context, remote_addr), target_cli: payload.target_cli, target_system: payload.target_system, expires_at_unix_secs, @@ -641,10 +909,17 @@ pub(crate) async fn build_api_key_install_session_response( ) } }; + let Some(sealed) = seal_install_session(state, &serialized) else { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "安装会话加密不可用", + false, + ); + }; if let Err(err) = state .runtime_kv_setex( &install_session_runtime_key(&code), - &serialized, + &sealed, INSTALL_SESSION_TTL_SECS, ) .await @@ -657,31 +932,97 @@ pub(crate) async fn build_api_key_install_session_response( } let base_url = session.base_url.trim_end_matches('/'); - Json(json!({ - "install_code": code, - "expires_at_unix_secs": expires_at_unix_secs, - "expires_in_seconds": INSTALL_SESSION_TTL_SECS, - "target_cli": session.target_cli, - "target_cli_label": cli_label(session.target_cli), - "target_system": session.target_system, - "target_system_label": system_label(session.target_system), - "unix_command": format!("curl -fsSL {base_url}/install/{code} | sh"), - "powershell_command": format!("irm {base_url}/install/{code}.ps1 | iex"), - })) - .into_response() + let unix_url = shell_single_quote(&format!("{base_url}/install/{code}")); + let powershell_url = powershell_single_quote(&format!("{base_url}/install/{code}.ps1")); + mark_sensitive_response_no_store( + Json(json!({ + "install_code": code, + "expires_at_unix_secs": expires_at_unix_secs, + "expires_in_seconds": INSTALL_SESSION_TTL_SECS, + "target_cli": session.target_cli, + "target_cli_label": cli_label(session.target_cli), + "target_system": session.target_system, + "target_system_label": system_label(session.target_system), + "unix_command": format!("curl -fsSL {unix_url} | sh"), + "powershell_command": format!("irm {powershell_url} | iex"), + })) + .into_response(), + ) } pub(crate) async fn build_proxy_node_install_session_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, node_name: String, + management_token_record: &aether_data::repository::management_tokens::StoredManagementToken, management_token: String, ) -> Response { + let current_user = match state + .find_user_auth_by_id(&management_token_record.user_id) + .await + { + Ok(Some(user)) => user, + Ok(None) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "隧道安装令牌所有者不可用", + false, + ) + } + Err(err) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + format!("tunnel install administrator snapshot lookup failed: {err:?}"), + false, + ) + } + }; + if current_user.id != management_token_record.user_id + || !current_user.is_active + || current_user.is_deleted + || !current_user.role.eq_ignore_ascii_case("admin") + || current_user.security_version < 0 + { + if current_user.security_version >= 0 { + discard_unused_tunnel_install_management_token( + state, + management_token_record, + &management_token, + current_user.security_version, + ) + .await; + } + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "隧道安装令牌所有者不再具有有效管理员身份", + false, + ); + } + if management_token_record.is_active + || management_token_record.permissions != Some(json!(["admin:proxy_nodes:write"])) + { + discard_unused_tunnel_install_management_token( + state, + management_token_record, + &management_token, + current_user.security_version, + ) + .await; + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "隧道安装令牌初始化状态无效", + false, + ); + } let code = generate_install_code(); let expires_at_unix_secs = unix_secs_now().saturating_add(INSTALL_SESSION_TTL_SECS); let session = StoredTunnelInstallSession { - aether_url: base_url_from_request(headers, request_context), + install_code: code.clone(), + aether_url: base_url_from_request(headers, request_context, remote_addr), + management_token_snapshot: management_token_record.clone(), + management_token_user_security_version: current_user.security_version, management_token, node_name, tunnel_security: "non_tls_required".to_string(), @@ -691,21 +1032,49 @@ pub(crate) async fn build_proxy_node_install_session_response( let serialized = match serde_json::to_string(&session) { Ok(value) => value, Err(err) => { + discard_unused_tunnel_install_management_token( + state, + management_token_record, + &session.management_token, + session.management_token_user_security_version, + ) + .await; return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("tunnel install session serialize failed: {err:?}"), false, - ) + ); } }; + let Some(sealed) = seal_install_session(state, &serialized) else { + discard_unused_tunnel_install_management_token( + state, + management_token_record, + &session.management_token, + session.management_token_user_security_version, + ) + .await; + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "隧道安装会话加密不可用", + false, + ); + }; if let Err(err) = state .runtime_kv_setex( &tunnel_install_session_runtime_key(&code), - &serialized, + &sealed, INSTALL_SESSION_TTL_SECS, ) .await { + discard_unused_tunnel_install_management_token( + state, + management_token_record, + &session.management_token, + session.management_token_user_security_version, + ) + .await; return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("tunnel install session create failed: {err:?}"), @@ -714,16 +1083,20 @@ pub(crate) async fn build_proxy_node_install_session_response( } let base_url = session.aether_url.trim_end_matches('/'); - Json(json!({ - "install_code": code, - "expires_at_unix_secs": expires_at_unix_secs, - "expires_in_seconds": INSTALL_SESSION_TTL_SECS, - "node_name": session.node_name, - "aether_url": session.aether_url, - "unix_command": format!("curl -fsSL {base_url}/install-tunnel/{code} | sh"), - "powershell_command": format!("irm {base_url}/install-tunnel/{code}.ps1 | iex"), - })) - .into_response() + let unix_url = shell_single_quote(&format!("{base_url}/install-tunnel/{code}")); + let powershell_url = powershell_single_quote(&format!("{base_url}/install-tunnel/{code}.ps1")); + mark_sensitive_response_no_store( + Json(json!({ + "install_code": code, + "expires_at_unix_secs": expires_at_unix_secs, + "expires_in_seconds": INSTALL_SESSION_TTL_SECS, + "node_name": session.node_name, + "aether_url": session.aether_url, + "unix_command": format!("curl -fsSL {unix_url} | sh"), + "powershell_command": format!("irm {powershell_url} | iex"), + })) + .into_response(), + ) } pub(super) async fn maybe_build_local_install_response( @@ -765,6 +1138,13 @@ pub(super) async fn maybe_build_local_install_response( )) } }; + let Some(raw) = open_install_session(state, &raw) else { + return Some(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "install code 数据无效", + false, + )); + }; let session = match serde_json::from_str::(&raw) { Ok(value) => value, Err(_) => { @@ -775,6 +1155,13 @@ pub(super) async fn maybe_build_local_install_response( )) } }; + if session.install_code != code { + return Some(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "install code 数据绑定无效", + false, + )); + } if session.expires_at_unix_secs <= unix_secs_now() { return Some(build_auth_error_response( http::StatusCode::NOT_FOUND, @@ -782,10 +1169,28 @@ pub(super) async fn maybe_build_local_install_response( false, )); } + let api_key = + match resolve_current_install_session_api_key(state, &session, unix_secs_now()).await { + Ok(Some(api_key)) => api_key, + Ok(None) => { + return Some(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "install code 已失效或关联密钥不可用", + false, + )) + } + Err(err) => { + return Some(build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + format!("install session key validation failed: {err:?}"), + false, + )) + } + }; let body = if wants_powershell { - build_powershell_script(&session) + build_powershell_script(&session, &api_key) } else { - build_unix_script(&session) + build_unix_script(&session, &api_key) }; let content_type = if wants_powershell { "text/plain; charset=utf-8" @@ -845,6 +1250,13 @@ async fn maybe_build_local_tunnel_install_response( ) } }; + let Some(raw) = open_install_session(state, &raw) else { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "tunnel install code 数据无效", + false, + ); + }; let session = match serde_json::from_str::(&raw) { Ok(value) => value, Err(_) => { @@ -855,13 +1267,41 @@ async fn maybe_build_local_tunnel_install_response( ) } }; + if session.install_code != code { + discard_tunnel_install_session_token(state, &session).await; + return build_auth_error_response( + http::StatusCode::NOT_FOUND, + "tunnel install code 数据绑定无效", + false, + ); + } if session.expires_at_unix_secs <= unix_secs_now() { + discard_tunnel_install_session_token(state, &session).await; return build_auth_error_response( http::StatusCode::NOT_FOUND, "tunnel install code 已过期", false, ); } + match activate_tunnel_install_management_token(state, &session, unix_secs_now()).await { + Ok(true) => {} + Ok(false) => { + discard_tunnel_install_session_token(state, &session).await; + return build_auth_error_response( + http::StatusCode::NOT_FOUND, + "tunnel install code 已失效或关联令牌不可用", + false, + ); + } + Err(err) => { + discard_tunnel_install_session_token(state, &session).await; + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + format!("tunnel install token activation failed: {err:?}"), + false, + ); + } + } let body = if wants_powershell { build_tunnel_powershell_script(&session) } else { @@ -895,12 +1335,21 @@ async fn maybe_build_local_tunnel_install_response( #[cfg(test)] mod tests { use super::*; + use std::sync::Arc; + + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; + use aether_data::repository::auth::{ + AuthApiKeyWriteRepository, InMemoryAuthApiKeySnapshotRepository, + StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, + }; fn test_session(target_cli: InstallTargetCli) -> StoredInstallSession { StoredInstallSession { + install_code: "0123456789abcdef01234567".to_string(), api_key_id: "key-1".to_string(), - api_key_name: "Key 1".to_string(), - api_key: "sk-test".to_string(), + api_key_owner_user_id: "user-1".to_string(), + api_key_is_standalone: false, + api_key_hash: "hash-key-1".to_string(), base_url: "http://localhost:8084".to_string(), target_cli, target_system: InstallTargetSystem::Linux, @@ -910,7 +1359,20 @@ mod tests { fn test_tunnel_session() -> StoredTunnelInstallSession { StoredTunnelInstallSession { + install_code: "0123456789abcdef01234567".to_string(), aether_url: "https://aether.example".to_string(), + management_token_snapshot: + aether_data::repository::management_tokens::StoredManagementToken::new( + "token-1".to_string(), + "admin-1".to_string(), + "tunnel install token".to_string(), + ) + .expect("management token should build") + .with_display_fields(None, Some("ae-test".to_string()), None) + .with_permissions(Some(json!(["admin:proxy_nodes:write"]))) + .with_runtime_fields(None, None, None, 0, false) + .with_timestamps(Some(1_700_000_000), Some(1_700_000_000)), + management_token_user_security_version: 3, management_token: "ae-test-token".to_string(), node_name: "jp-proxy-01".to_string(), tunnel_security: "non_tls_required".to_string(), @@ -919,17 +1381,257 @@ mod tests { } } + fn test_user_api_key_snapshot() -> StoredAuthApiKeySnapshot { + StoredAuthApiKeySnapshot::new( + "user-1".to_string(), + "user".to_string(), + Some("user@example.com".to_string()), + "user".to_string(), + "local".to_string(), + true, + false, + None, + None, + None, + "key-1".to_string(), + Some("Key 1".to_string()), + true, + false, + false, + None, + None, + Some(4_102_444_800), + None, + None, + None, + ) + .expect("snapshot should build") + } + + fn test_user_api_key_export(api_key: &str) -> StoredAuthApiKeyExportRecord { + StoredAuthApiKeyExportRecord::new( + "user-1".to_string(), + "key-1".to_string(), + sha256_hex(api_key), + Some( + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, api_key) + .expect("API key should encrypt"), + ), + Some("Key 1".to_string()), + None, + None, + None, + None, + None, + None, + true, + Some(4_102_444_800), + false, + 0, + 0, + 0.0, + false, + ) + .expect("export record should build") + } + + #[test] + fn install_session_runtime_envelope_encrypts_credentials_and_round_trips() { + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let plaintext = serde_json::to_string(&test_session(InstallTargetCli::CodexCli)) + .expect("session should serialize"); + assert!(!plaintext.contains("sk-test")); + assert!(!plaintext.contains("\"api_key\":")); + + let sealed = seal_install_session(&state, &plaintext) + .expect("configured state should seal install session"); + assert!(sealed.starts_with(INSTALL_SESSION_ENVELOPE_PREFIX)); + assert_eq!( + open_install_session(&state, &sealed).as_deref(), + Some(plaintext.as_str()) + ); + } + + #[test] + fn install_session_runtime_keys_do_not_retain_bearer_codes() { + let code = "0123456789abcdef01234567"; + for runtime_key in [ + install_session_runtime_key(code), + tunnel_install_session_runtime_key(code), + ] { + assert!(runtime_key.contains("sha256:")); + assert!(!runtime_key.contains(code)); + } + } + + #[test] + fn tunnel_install_mutation_binds_full_token_and_admin_security_snapshots() { + let session = test_tunnel_session(); + let mutation = tunnel_install_token_mutation(&session, 1_700_000_100); + + assert_eq!(mutation.expected_token, session.management_token_snapshot); + assert_eq!( + mutation.expected_user_security_version, + session.management_token_user_security_version + ); + assert_eq!(mutation.token_hash, sha256_hex(&session.management_token)); + assert_eq!(mutation.now_unix_secs, 1_700_000_100); + } + + #[tokio::test] + async fn install_session_strong_revalidation_rejects_new_lock_with_stale_cache() { + let api_key = "sk-test-lock-revalidation"; + let repository = Arc::new( + InMemoryAuthApiKeySnapshotRepository::seed([( + Some(sha256_hex(api_key)), + test_user_api_key_snapshot(), + )]) + .with_export_records([test_user_api_key_export(api_key)]), + ); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_cached_auth_api_key_repository_for_tests( + Arc::clone(&repository), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let mut session = test_session(InstallTargetCli::CodexCli); + session.api_key_hash = sha256_hex(api_key); + + let now_unix_secs = unix_secs_now(); + let cached = state + .data + .read_auth_api_key_snapshot("user-1", "key-1", now_unix_secs) + .await + .expect("initial snapshot should load") + .expect("initial snapshot should exist"); + assert!(!cached.api_key_is_locked); + assert_eq!( + resolve_current_install_session_api_key(&state, &session, now_unix_secs) + .await + .expect("initial strong validation should complete") + .as_deref(), + Some(api_key), + ); + + repository + .set_user_api_key_locked("user-1", "key-1", true) + .await + .expect("lock update should succeed"); + let stale = state + .data + .read_auth_api_key_snapshot("user-1", "key-1", now_unix_secs) + .await + .expect("cached snapshot should load") + .expect("cached snapshot should exist"); + assert!( + !stale.api_key_is_locked, + "test must retain a stale cache entry" + ); + + assert_eq!( + resolve_current_install_session_api_key(&state, &session, now_unix_secs,) + .await + .expect("strong validation should complete"), + None, + ); + } + + #[test] + fn install_session_reader_rejects_legacy_plaintext() { + let state = AppState::new().expect("state should build"); + let legacy = r#"{"api_key":"legacy-secret"}"#; + assert!(open_install_session(&state, legacy).is_none()); + } + #[test] fn tunnel_install_path_accepts_shell_and_powershell_codes() { + let code = "0123456789abcdef01234567"; assert_eq!( - tunnel_install_code_from_path("/install-tunnel/abc123"), - Some(("abc123".to_string(), false)) + tunnel_install_code_from_path(&format!("/install-tunnel/{code}")), + Some((code.to_string(), false)) ); assert_eq!( - tunnel_install_code_from_path("/install-tunnel/abc123.ps1"), - Some(("abc123".to_string(), true)) + tunnel_install_code_from_path(&format!("/install-tunnel/{code}.ps1")), + Some((code.to_string(), true)) ); assert_eq!(tunnel_install_code_from_path("/install-tunnel/a/b"), None); + assert_eq!( + tunnel_install_code_from_path("/install-tunnel/abc123"), + None + ); + assert_eq!( + tunnel_install_code_from_path("/install-tunnel/0123456789ABCDEF01234567"), + None + ); + } + + #[test] + fn public_base_url_rejects_credentials_queries_and_shell_metacharacters() { + assert_eq!( + normalize_public_base_url(" https://aether.example/base/ "), + Some("https://aether.example/base".to_string()) + ); + assert_eq!( + normalize_public_base_url("http://127.0.0.1:8084/"), + Some("http://127.0.0.1:8084".to_string()) + ); + for value in [ + "http://aether.example", + "http://10.0.0.8:8084", + "https://user:secret@aether.example", + "https://aether.example?token=secret", + "https://aether.example/#fragment", + "https://aether.example';touch /tmp/pwned;'", + "javascript:alert(1)", + ] { + assert_eq!( + normalize_public_base_url(value), + None, + "unsafe URL: {value}" + ); + } + } + + #[test] + fn untrusted_peer_cannot_poison_generated_install_base_url() { + let mut headers = http::HeaderMap::new(); + headers.insert("x-forwarded-host", "attacker.example".parse().unwrap()); + headers.insert("x-forwarded-proto", "https".parse().unwrap()); + + assert_eq!( + base_url_from_request_metadata(&headers, Some("also-attacker.example"), false), + "http://localhost" + ); + assert_eq!( + base_url_from_request_metadata(&headers, Some("gateway.example"), true), + "https://attacker.example" + ); + } + + #[test] + fn trusted_proxy_uses_rightmost_forwarded_host_and_proto() { + let mut headers = http::HeaderMap::new(); + headers.append( + "x-forwarded-host", + "client-injected.example, stale-proxy.example" + .parse() + .unwrap(), + ); + headers.append("x-forwarded-host", "gateway.example".parse().unwrap()); + headers.append("x-forwarded-proto", "http, http".parse().unwrap()); + headers.append("x-forwarded-proto", "https".parse().unwrap()); + + assert_eq!( + base_url_from_request_metadata(&headers, Some("ignored.example"), true), + "https://gateway.example" + ); } #[test] @@ -966,7 +1668,7 @@ mod tests { #[test] fn codex_unix_script_preserves_config_and_uses_responses_bearer_token() { - let script = build_unix_script(&test_session(InstallTargetCli::CodexCli)); + let script = build_unix_script(&test_session(InstallTargetCli::CodexCli), "sk-test"); assert!(script.contains("path.read_text() if path.exists() else ''")); assert!(script.contains("stripped == '[model_providers.aether]'")); @@ -981,7 +1683,7 @@ mod tests { #[test] fn codex_powershell_script_preserves_config_and_uses_responses_bearer_token() { - let script = build_powershell_script(&test_session(InstallTargetCli::CodexCli)); + let script = build_powershell_script(&test_session(InstallTargetCli::CodexCli), "sk-test"); assert!(script.contains("Get-Content $Path -Raw")); assert!(script.contains("$Stripped -eq '[model_providers.aether]'")); diff --git a/apps/aether-gateway/src/handlers/public/support/monitoring/audit_logs.rs b/apps/aether-gateway/src/handlers/public/support/monitoring/audit_logs.rs index 167aa5dc4..0eaf38754 100644 --- a/apps/aether-gateway/src/handlers/public/support/monitoring/audit_logs.rs +++ b/apps/aether-gateway/src/handlers/public/support/monitoring/audit_logs.rs @@ -123,14 +123,10 @@ pub(super) async fn handle_user_audit_logs( .await { Ok(value) => value, - Err(err) => { - let detail = match err { - crate::GatewayError::Internal(message) => message, - other => format!("{other:?}"), - }; + Err(_err) => { return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, - detail, + "audit logs unavailable", false, ); } diff --git a/apps/aether-gateway/src/handlers/public/support/oauth.rs b/apps/aether-gateway/src/handlers/public/support/oauth.rs index d38e356ef..d0200cfa7 100644 --- a/apps/aether-gateway/src/handlers/public/support/oauth.rs +++ b/apps/aether-gateway/src/handlers/public/support/oauth.rs @@ -1,24 +1,144 @@ -use super::support_auth::auth_session::{ - build_auth_login_success_response, create_auth_token, decode_auth_token, -}; +use super::support_auth::auth_session::build_auth_login_success_response; +use super::support_auth::local_password_login_allowed_for_user; +use super::support_auth::{auth_refresh_cookie_secure, extract_cookie_value}; use super::{ build_auth_error_response, build_auth_json_response, extract_client_device_id, http, json, - query_param_value, resolve_authenticated_local_user, AppState, Body, Bytes, + mark_sensitive_response_no_store, resolve_authenticated_local_user, AppState, Body, Bytes, GatewayPublicRequestContext, IntoResponse, Json, Response, }; -use aether_oauth::core::{generate_pkce_verifier, pkce_s256, OAuthError}; +use aether_oauth::core::{generate_oauth_nonce, generate_pkce_verifier, pkce_s256, OAuthError}; use aether_oauth::identity::{ IdentityClaims, IdentityOAuthExchangeContext, IdentityOAuthService, IdentityOAuthStartContext, }; -use axum::body::to_bytes; use axum::http::header::{LOCATION, SET_COOKIE}; use axum::http::HeaderValue; +use hmac::{Hmac, Mac}; +use sha2::{Digest, Sha256}; use url::form_urlencoded; +const OAUTH_LOGIN_COOKIE_NAME_PREFIX: &str = "aether_oauth_login_"; +const OAUTH_LOGIN_HOST_COOKIE_NAME_PREFIX: &str = "__Host-aether_oauth_login_"; +const OAUTH_LOGIN_COOKIE_MAX_AGE_SECS: u64 = 10 * 60; +const OAUTH_LOGIN_BINDING_COMPARE_KEY: &[u8] = b"aether-oauth-login-binding-v1"; + +type HmacSha256 = Hmac; + +fn browser_binding_hash(binding: &str) -> String { + format!("{:x}", Sha256::digest(binding.as_bytes())) +} + +fn oauth_login_cookie_name(state_nonce: &str) -> String { + oauth_login_cookie_name_for_security(state_nonce, auth_refresh_cookie_secure()) +} + +fn oauth_login_cookie_name_for_security(state_nonce: &str, secure: bool) -> String { + let prefix = if secure { + OAUTH_LOGIN_HOST_COOKIE_NAME_PREFIX + } else { + OAUTH_LOGIN_COOKIE_NAME_PREFIX + }; + format!("{prefix}{:x}", Sha256::digest(state_nonce.as_bytes()),) +} + +fn browser_binding_matches(expected_hash: Option<&str>, binding: Option<&str>) -> bool { + let Some(expected_hash) = expected_hash else { + return false; + }; + let candidate_hash = browser_binding_hash(binding.unwrap_or_default()); + + // HMAC verification performs the digest comparison in constant time while still + // handling malformed/legacy state values as a normal mismatch. + let mut expected_mac = HmacSha256::new_from_slice(OAUTH_LOGIN_BINDING_COMPARE_KEY) + .expect("static OAuth binding comparison key should be valid"); + expected_mac.update(expected_hash.as_bytes()); + let expected_tag = expected_mac.finalize().into_bytes(); + let mut candidate_mac = HmacSha256::new_from_slice(OAUTH_LOGIN_BINDING_COMPARE_KEY) + .expect("static OAuth binding comparison key should be valid"); + candidate_mac.update(candidate_hash.as_bytes()); + candidate_mac.verify_slice(&expected_tag).is_ok() +} + +fn build_oauth_login_cookie_header(state_nonce: &str, binding: &str) -> String { + build_oauth_login_cookie_header_for_security(state_nonce, binding, auth_refresh_cookie_secure()) +} + +fn build_oauth_login_cookie_header_for_security( + state_nonce: &str, + binding: &str, + secure: bool, +) -> String { + let path = if secure { "/" } else { "/api/oauth" }; + let mut cookie = format!( + "{}={}; Path={path}; HttpOnly; SameSite=Lax; Max-Age={}", + oauth_login_cookie_name_for_security(state_nonce, secure), + binding, + OAUTH_LOGIN_COOKIE_MAX_AGE_SECS, + ); + if secure { + cookie.push_str("; Secure"); + } + cookie +} + +fn build_oauth_login_cookie_clear_header(state_nonce: &str) -> String { + build_oauth_login_cookie_clear_header_for_security(state_nonce, auth_refresh_cookie_secure()) +} + +fn build_oauth_login_cookie_clear_header_for_security(state_nonce: &str, secure: bool) -> String { + let path = if secure { "/" } else { "/api/oauth" }; + let mut cookie = format!( + "{}=; Path={path}; HttpOnly; SameSite=Lax; Max-Age=0", + oauth_login_cookie_name_for_security(state_nonce, secure), + ); + if secure { + cookie.push_str("; Secure"); + } + cookie +} + +fn append_oauth_login_cookie( + mut response: Response, + state_nonce: &str, + binding: &str, +) -> Response { + response.headers_mut().append( + SET_COOKIE, + HeaderValue::from_str(&build_oauth_login_cookie_header(state_nonce, binding)) + .expect("OAuth login cookie header should be valid"), + ); + mark_sensitive_response_no_store(response) +} + +fn append_oauth_login_cookie_clear( + mut response: Response, + state_nonce: &str, +) -> Response { + response.headers_mut().append( + SET_COOKIE, + HeaderValue::from_str(&build_oauth_login_cookie_clear_header(state_nonce)) + .expect("OAuth login cookie clear header should be valid"), + ); + mark_sensitive_response_no_store(response) +} + +fn redirect_oauth_error_for_mode( + frontend_callback_url: Option<&str>, + code: &str, + login_cookie_state_nonce: Option<&str>, +) -> Response { + let response = redirect_oauth_error(frontend_callback_url, code); + if let Some(state_nonce) = login_cookie_state_nonce { + append_oauth_login_cookie_clear(response, state_nonce) + } else { + response + } +} + pub(super) async fn maybe_build_local_oauth_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, _request_body: Option<&Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -37,7 +157,7 @@ pub(super) async fn maybe_build_local_oauth_response( Some(handle_oauth_authorize(state, request_context, headers).await) } Some("callback") if request_context.request_method == http::Method::GET => { - Some(handle_oauth_callback(state, request_context, headers).await) + Some(handle_oauth_callback(state, request_context, headers, client_ip).await) } Some("bindable_providers") if request_context.request_method == http::Method::GET @@ -87,10 +207,12 @@ async fn handle_oauth_bindable_providers( Err(response) => return response, }; if auth.user.auth_source.eq_ignore_ascii_case("ldap") { - return Json(json!({ "providers": [] })).into_response(); + return mark_sensitive_response_no_store(Json(json!({ "providers": [] })).into_response()); } match crate::oauth::list_bindable_identity_oauth_providers(state, &auth.user.id).await { - Ok(providers) => Json(json!({ "providers": providers })).into_response(), + Ok(providers) => mark_sensitive_response_no_store( + Json(json!({ "providers": providers })).into_response(), + ), Err(err) => build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("oauth provider lookup failed: {err:?}"), @@ -109,7 +231,9 @@ async fn handle_oauth_links( Err(response) => return response, }; match crate::oauth::list_identity_oauth_links(state, &auth.user.id).await { - Ok(links) => Json(json!({ "links": links })).into_response(), + Ok(links) => { + mark_sensitive_response_no_store(Json(json!({ "links": links })).into_response()) + } Err(err) => build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("oauth link lookup failed: {err:?}"), @@ -183,29 +307,38 @@ async fn handle_oauth_bind_token( } Err(err) => return oauth_account_error_response(err), } - let token = match create_auth_token( - "oauth_bind", - serde_json::Map::from_iter([ - ("user_id".to_string(), json!(auth.user.id)), - ("session_id".to_string(), json!(auth.session_id)), - ("provider_type".to_string(), json!(provider_type)), - ]), - chrono::Utc::now() + chrono::Duration::minutes(10), - ) { + let client_device_id = match extract_client_device_id(request_context, headers) { Ok(value) => value, - Err(detail) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - detail, - false, - ) - } + Err(response) => return response, }; - build_auth_json_response(http::StatusCode::OK, json!({ "bind_token": token }), None) + let user_id = auth.user.id.clone(); + let session_id = auth.session_id.clone(); + let (authorize_url, browser_cookie) = match build_identity_oauth_authorize_url( + state, + &provider_type, + client_device_id, + crate::oauth::IdentityOAuthStateMode::Bind, + Some(user_id), + Some(session_id), + ) + .await + { + Ok(value) => value, + Err(response) => return response, + }; + // The response contains only the provider authorization URL. The binding + // capability remains server-side in the one-time OAuth state record. + let response = build_auth_json_response( + http::StatusCode::OK, + json!({ "authorize_url": authorize_url }), + None, + ); + let (state_nonce, binding) = browser_cookie; + append_oauth_login_cookie(response, &state_nonce, &binding) } async fn handle_oauth_bind_start( - state: &AppState, + _state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, ) -> Response { @@ -217,55 +350,37 @@ async fn handle_oauth_bind_start( false, ); }; - let client_device_id = match extract_client_device_id(request_context, headers) { - Ok(value) => value, - Err(response) => return response, - }; - let Some(bind_token) = query_param_value( - request_context.request_query_string.as_deref(), - "bind_token", - ) else { - return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少绑定令牌", false); - }; - let bind = - match validate_bind_token(state, &provider_type, &client_device_id, &bind_token).await { - Ok(value) => value, - Err(response) => return response, - }; - start_identity_oauth( - state, - &provider_type, - client_device_id, - crate::oauth::IdentityOAuthStateMode::Bind, - Some(bind.user_id), - Some(bind.session_id), + // Bind state is created by the authenticated POST /bind-token endpoint. + // A browser navigation must never carry a bearer binding token in its URL. + // The state record is selected by the provider authorization URL returned by + // that endpoint, so this legacy GET endpoint is intentionally unavailable. + let _ = headers; + build_auth_error_response( + http::StatusCode::GONE, + "OAuth 绑定流程已更新,请重新发起绑定", + false, ) - .await } async fn handle_oauth_callback( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, ) -> Response { let Some(provider_type) = public_oauth_provider_from_path(&request_context.request_path, "callback") else { return redirect_oauth_error(None, "provider_unavailable"); }; - let params = callback_params(request_context); - if params - .get("error") - .is_some_and(|value| value.eq_ignore_ascii_case("access_denied")) - { - return redirect_oauth_error(None, "authorization_denied"); - } - let Some(code) = params - .get("code") - .map(String::as_str) - .filter(|value| !value.is_empty()) - else { - return redirect_oauth_error(None, "invalid_callback"); + let params = match callback_params(request_context.request_query_string.as_deref()) { + Ok(params) => params, + Err(CallbackParamsError::DuplicateState) => { + return redirect_oauth_error(None, "invalid_state") + } + Err(CallbackParamsError::DuplicateCallbackParameter) => { + return redirect_oauth_error(None, "invalid_callback") + } }; let Some(nonce) = params .get("state") @@ -274,20 +389,63 @@ async fn handle_oauth_callback( else { return redirect_oauth_error(None, "invalid_state"); }; - let stored = match crate::oauth::consume_identity_oauth_state(state, nonce).await { + let preview = match crate::oauth::load_identity_oauth_state(state, nonce).await { Ok(Some(value)) => value, Ok(None) => return redirect_oauth_error(None, "invalid_state"), Err(_) => return redirect_oauth_error(None, "invalid_state"), }; - if stored.provider_type != provider_type { - return redirect_oauth_error(None, "invalid_state"); + let browser_cookie_state_nonce = Some(preview.nonce.as_str()); + let login_cookie_name = oauth_login_cookie_name(&preview.nonce); + let browser_binding_matches = browser_binding_matches( + preview.browser_binding_hash.as_deref(), + extract_cookie_value(headers, &login_cookie_name).as_deref(), + ); + // Both login and account binding states carry authorization authority. A + // callback must prove it originated in the browser that started the flow. + if !browser_binding_matches { + return redirect_oauth_error_for_mode(None, "invalid_state", browser_cookie_state_nonce); } + if preview.provider_type != provider_type { + return redirect_oauth_error_for_mode(None, "invalid_state", browser_cookie_state_nonce); + } + let stored = match crate::oauth::consume_identity_oauth_state(state, nonce).await { + Ok(Some(value)) if value == preview => value, + Ok(Some(_)) | Ok(None) | Err(_) => { + return redirect_oauth_error_for_mode(None, "invalid_state", browser_cookie_state_nonce) + } + }; + let browser_cookie_state_nonce = Some(stored.nonce.as_str()); + if params + .get("error") + .is_some_and(|value| value.eq_ignore_ascii_case("access_denied")) + { + return redirect_oauth_error_for_mode( + None, + "authorization_denied", + browser_cookie_state_nonce, + ); + } + let Some(code) = params + .get("code") + .map(String::as_str) + .filter(|value| !value.is_empty()) + else { + return redirect_oauth_error_for_mode(None, "invalid_callback", browser_cookie_state_nonce); + }; let config = match crate::oauth::get_enabled_identity_oauth_provider_config(state, &provider_type).await { Ok(Some(value)) => value, - Ok(None) => return redirect_oauth_error(None, "provider_disabled"), - Err(err) => return redirect_oauth_error(None, err.code()), + Ok(None) => { + return redirect_oauth_error_for_mode( + None, + "provider_disabled", + browser_cookie_state_nonce, + ) + } + Err(err) => { + return redirect_oauth_error_for_mode(None, err.code(), browser_cookie_state_nonce) + } }; let network = crate::oauth::resolve_identity_oauth_network_context(state).await; let exchange_ctx = IdentityOAuthExchangeContext { @@ -301,26 +459,32 @@ async fn handle_oauth_callback( let claims = match service.login(&executor, &config, &exchange_ctx).await { Ok(outcome) => outcome.claims, Err(err) => { - return redirect_oauth_error( + return redirect_oauth_error_for_mode( Some(&config.frontend_callback_url), oauth_error_code(&err), + browser_cookie_state_nonce, ) } }; match stored.mode { crate::oauth::IdentityOAuthStateMode::Login => { - complete_oauth_login( + let response = complete_oauth_login( state, headers, + client_ip, &config.frontend_callback_url, stored.client_device_id, claims, ) - .await + .await; + append_oauth_login_cookie_clear(response, &stored.nonce) } crate::oauth::IdentityOAuthStateMode::Bind => { - complete_oauth_bind(state, &config.frontend_callback_url, stored, claims).await + let state_nonce = stored.nonce.clone(); + let response = + complete_oauth_bind(state, &config.frontend_callback_url, stored, claims).await; + append_oauth_login_cookie_clear(response, &state_nonce) } } } @@ -343,7 +507,25 @@ async fn handle_oauth_unbind( Ok(value) => value, Err(response) => return response, }; - match crate::oauth::unbind_identity_oauth(state, &auth.user, &provider_type).await { + let local_password_login_allowed = + match local_password_login_allowed_for_user(state, Some(&auth.user)).await { + Ok(allowed) => allowed, + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("auth login policy lookup failed: {err:?}"), + false, + ) + } + }; + match crate::oauth::unbind_identity_oauth( + state, + &auth.user, + &provider_type, + local_password_login_allowed, + ) + .await + { Ok(true) => Json(json!({ "message": "解绑成功" })).into_response(), Ok(false) => { build_auth_error_response(http::StatusCode::NOT_FOUND, "OAuth 绑定不存在", false) @@ -352,6 +534,98 @@ async fn handle_oauth_unbind( } } +async fn build_identity_oauth_authorize_url( + state: &AppState, + provider_type: &str, + client_device_id: String, + mode: crate::oauth::IdentityOAuthStateMode, + bind_user_id: Option, + bind_session_id: Option, +) -> Result<(String, (String, String)), Response> { + let config = match crate::oauth::get_enabled_identity_oauth_provider_config( + state, + provider_type, + ) + .await + { + Ok(Some(value)) => value, + Ok(None) => { + return Err(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "OAuth Provider 不存在或已禁用", + false, + )) + } + Err(err) => return Err(oauth_account_error_response(err)), + }; + let pkce_verifier = generate_pkce_verifier(); + let code_challenge = pkce_s256(&pkce_verifier); + let browser_binding = generate_oauth_nonce(); + let stored = match mode { + crate::oauth::IdentityOAuthStateMode::Login => { + crate::oauth::StoredIdentityOAuthState::login( + provider_type, + client_device_id, + Some(pkce_verifier), + Some(browser_binding_hash(&browser_binding)), + ) + } + crate::oauth::IdentityOAuthStateMode::Bind => { + let Some(user_id) = bind_user_id else { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "缺少绑定用户", + false, + )); + }; + let Some(session_id) = bind_session_id else { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "缺少绑定会话", + false, + )); + }; + crate::oauth::StoredIdentityOAuthState::bind( + provider_type, + client_device_id, + Some(pkce_verifier), + browser_binding_hash(&browser_binding), + user_id, + session_id, + ) + } + }; + if crate::oauth::save_identity_oauth_state(state, &stored) + .await + .is_err() + { + return Err(build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "OAuth 状态存储不可用", + false, + )); + } + let network = crate::oauth::resolve_identity_oauth_network_context(state).await; + let state_nonce = stored.nonce.clone(); + let start_ctx = IdentityOAuthStartContext { + state: stored.nonce, + code_challenge: Some(code_challenge), + network, + }; + let authorize = match IdentityOAuthService::with_builtin_providers().start(&config, &start_ctx) + { + Ok(value) => value, + Err(_) => { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "OAuth Provider 不可用", + false, + )) + } + }; + Ok((authorize.authorize_url, (state_nonce, browser_binding))) +} + async fn start_identity_oauth( state: &AppState, provider_type: &str, @@ -360,89 +634,28 @@ async fn start_identity_oauth( bind_user_id: Option, bind_session_id: Option, ) -> Response { - let config = match crate::oauth::get_enabled_identity_oauth_provider_config( + let (authorize_url, browser_cookie) = match build_identity_oauth_authorize_url( state, provider_type, + client_device_id, + mode, + bind_user_id, + bind_session_id, ) .await - { - Ok(Some(value)) => value, - Ok(None) => { - return build_auth_error_response( - http::StatusCode::NOT_FOUND, - "OAuth Provider 不存在或已禁用", - false, - ) - } - Err(err) => return oauth_account_error_response(err), - }; - let pkce_verifier = generate_pkce_verifier(); - let code_challenge = pkce_s256(&pkce_verifier); - let stored = match mode { - crate::oauth::IdentityOAuthStateMode::Login => { - crate::oauth::StoredIdentityOAuthState::login( - provider_type, - client_device_id, - Some(pkce_verifier), - ) - } - crate::oauth::IdentityOAuthStateMode::Bind => { - let Some(user_id) = bind_user_id else { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "缺少绑定用户", - false, - ); - }; - let Some(session_id) = bind_session_id else { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "缺少绑定会话", - false, - ); - }; - crate::oauth::StoredIdentityOAuthState::bind( - provider_type, - client_device_id, - Some(pkce_verifier), - user_id, - session_id, - ) - } - }; - if crate::oauth::save_identity_oauth_state(state, &stored) - .await - .is_err() - { - return build_auth_error_response( - http::StatusCode::SERVICE_UNAVAILABLE, - "OAuth 状态存储不可用", - false, - ); - } - let network = crate::oauth::resolve_identity_oauth_network_context(state).await; - let start_ctx = IdentityOAuthStartContext { - state: stored.nonce, - code_challenge: Some(code_challenge), - network, - }; - let authorize = match IdentityOAuthService::with_builtin_providers().start(&config, &start_ctx) { Ok(value) => value, - Err(_) => { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "OAuth Provider 不可用", - false, - ) - } + Err(response) => return response, }; - redirect_to(&authorize.authorize_url, None) + let response = redirect_to(&authorize_url, None); + let (state_nonce, binding) = browser_cookie; + append_oauth_login_cookie(response, &state_nonce, &binding) } async fn complete_oauth_login( state: &AppState, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, frontend_callback_url: &str, client_device_id: String, claims: IdentityClaims, @@ -453,44 +666,28 @@ async fn complete_oauth_login( Err(err) => return redirect_oauth_error(Some(frontend_callback_url), err.code()), }; let login_response = - build_auth_login_success_response(state, headers, client_device_id, user).await; + build_auth_login_success_response(state, headers, client_ip, client_device_id, user, None) + .await; if login_response.status() != http::StatusCode::OK { return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"); } + redirect_oauth_login_success(frontend_callback_url, &login_response) +} + +fn redirect_oauth_login_success( + frontend_callback_url: &str, + login_response: &Response, +) -> Response { let set_cookies = login_response .headers() .get_all(SET_COOKIE) .iter() .cloned() .collect::>(); - let body = login_response.into_body(); - let body = match to_bytes(body, crate::headers::max_internal_buffered_body_bytes()).await { - Ok(value) => value, - Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"), - }; - let payload = match serde_json::from_slice::(&body) { - Ok(value) => value, - Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"), - }; - let Some(access_token) = payload - .get("access_token") - .and_then(serde_json::Value::as_str) - else { - return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"); - }; - let expires_in = payload - .get("expires_in") - .and_then(serde_json::Value::as_i64) - .unwrap_or(24 * 60 * 60) - .to_string(); - let mut response = redirect_to( - frontend_callback_url, - Some(RedirectParams::Fragment(vec![ - ("access_token", access_token.to_string()), - ("token_type", "bearer".to_string()), - ("expires_in", expires_in), - ])), - ); + // The callback only carries the HttpOnly refresh cookie. The frontend + // exchanges that cookie for an in-memory access token after navigation, so + // bearer credentials never enter the URL, browser history, or referrer. + let mut response = redirect_to(frontend_callback_url, None); for cookie in set_cookies { response.headers_mut().append(SET_COOKIE, cookie); } @@ -520,11 +717,24 @@ async fn complete_oauth_bind( let now = chrono::Utc::now(); if session.is_revoked() || session.is_expired(now) + || session.security_version != user.security_version || session.client_device_id != stored.client_device_id { return redirect_oauth_error(Some(frontend_callback_url), "invalid_state"); } - if let Err(err) = crate::oauth::bind_identity_oauth_to_user(state, &user, &claims).await { + let session_expectation = + match aether_data::repository::users::BindUserOAuthLinkSessionExpectation::new( + session.id, + stored.client_device_id, + user.security_version, + now, + ) { + Ok(expectation) => expectation, + Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "invalid_state"), + }; + if let Err(err) = + crate::oauth::bind_identity_oauth_to_user(state, &user, &claims, &session_expectation).await + { return redirect_oauth_error(Some(frontend_callback_url), err.code()); } redirect_to( @@ -536,116 +746,257 @@ async fn complete_oauth_bind( ) } -#[derive(Debug, Clone)] -struct ValidatedBindToken { - user_id: String, - session_id: String, +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn login_browser_binding_uses_a_hash_and_constant_time_match() { + let binding = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"; + let hash = browser_binding_hash(binding); + assert_ne!(hash, binding); + assert_eq!(hash.len(), 64); + assert!(browser_binding_matches(Some(&hash), Some(binding))); + assert!(!browser_binding_matches(Some(&hash), Some("wrong-binding"))); + assert!(!browser_binding_matches(None, Some(binding))); + assert!(!browser_binding_matches(Some("malformed"), Some(binding))); + } + + #[test] + fn secure_login_cookie_uses_host_prefix_and_required_scope() { + let state_nonce = "state-nonce"; + let binding = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"; + let cookie_name = oauth_login_cookie_name_for_security(state_nonce, true); + let set_cookie = build_oauth_login_cookie_header_for_security(state_nonce, binding, true); + assert!(set_cookie.starts_with(&format!("{cookie_name}="))); + assert!(cookie_name + .strip_prefix(OAUTH_LOGIN_HOST_COOKIE_NAME_PREFIX) + .is_some_and( + |suffix| suffix.len() == 64 && suffix.chars().all(|ch| ch.is_ascii_hexdigit()) + )); + assert!(set_cookie.contains("Path=/")); + assert!(set_cookie.contains("HttpOnly")); + assert!(set_cookie.contains("SameSite=Lax")); + assert!(set_cookie.contains("Max-Age=600")); + assert!(set_cookie.contains("Secure")); + assert!(!set_cookie.contains("Domain=")); + + let clear_cookie = build_oauth_login_cookie_clear_header_for_security(state_nonce, true); + assert!(clear_cookie.starts_with(&format!("{cookie_name}=;"))); + assert!(clear_cookie.contains("Path=/")); + assert!(clear_cookie.contains("Max-Age=0")); + assert!(clear_cookie.contains("Secure")); + assert!(!clear_cookie.contains("Domain=")); + assert!(!clear_cookie.contains("aether_refresh_token")); + } + + #[test] + fn insecure_local_login_cookie_keeps_legacy_name_and_narrow_path() { + let state_nonce = "state-nonce"; + let binding = "local-binding"; + let cookie_name = oauth_login_cookie_name_for_security(state_nonce, false); + let set_cookie = build_oauth_login_cookie_header_for_security(state_nonce, binding, false); + + assert!(cookie_name.starts_with(OAUTH_LOGIN_COOKIE_NAME_PREFIX)); + assert!(!cookie_name.starts_with(OAUTH_LOGIN_HOST_COOKIE_NAME_PREFIX)); + assert!(set_cookie.starts_with(&format!("{cookie_name}="))); + assert!(set_cookie.contains("Path=/api/oauth")); + assert!(!set_cookie.contains("Secure")); + assert!(!set_cookie.contains("Domain=")); + + let clear_cookie = build_oauth_login_cookie_clear_header_for_security(state_nonce, false); + assert!(clear_cookie.starts_with(&format!("{cookie_name}=;"))); + assert!(clear_cookie.contains("Path=/api/oauth")); + assert!(!clear_cookie.contains("Secure")); + assert!(!clear_cookie.contains("Domain=")); + } + + #[test] + fn login_cookie_names_are_distinct_per_state_nonce() { + let first = oauth_login_cookie_name_for_security("first-state", true); + let second = oauth_login_cookie_name_for_security("second-state", true); + + assert_ne!(first, second); + assert_ne!( + build_oauth_login_cookie_clear_header_for_security("first-state", true), + build_oauth_login_cookie_clear_header_for_security("second-state", true) + ); + } + + #[test] + fn oauth_provider_paths_require_exactly_one_provider_segment() { + assert_eq!( + public_oauth_provider_from_path("/api/oauth/LinuxDo/authorize", "authorize"), + Some("linuxdo".to_string()) + ); + assert_eq!( + user_oauth_provider_from_path("/api/user/oauth/LinuxDo/bind-token", "bind-token"), + Some("linuxdo".to_string()) + ); + assert_eq!( + user_oauth_provider_from_path_without_suffix("/api/user/oauth/LinuxDo"), + Some("linuxdo".to_string()) + ); + + for path in [ + "/api/oauth/linuxdo/extra/authorize", + "/api/oauth//authorize", + "/api/oauth/linuxdo//authorize", + ] { + assert_eq!(public_oauth_provider_from_path(path, "authorize"), None); + } + for path in [ + "/api/user/oauth/linuxdo/extra/bind-token", + "/api/user/oauth//bind-token", + "/api/user/oauth/linuxdo//bind-token", + ] { + assert_eq!(user_oauth_provider_from_path(path, "bind-token"), None); + } + for path in [ + "/api/user/oauth/", + "/api/user/oauth/ ", + "/api/user/oauth/linuxdo/extra", + ] { + assert_eq!(user_oauth_provider_from_path_without_suffix(path), None); + } + } + + #[test] + fn oauth_callback_rejects_duplicate_security_parameters() { + assert_eq!( + callback_params(Some("state=first&state=second")), + Err(CallbackParamsError::DuplicateState) + ); + assert_eq!( + callback_params(Some("state=first&st%61te=second")), + Err(CallbackParamsError::DuplicateState) + ); + for query in [ + "state=state&code=first&code=second", + "state=state&error=first&error=second", + ] { + assert_eq!( + callback_params(Some(query)), + Err(CallbackParamsError::DuplicateCallbackParameter) + ); + } + + let params = callback_params(Some("state=state&code=code")) + .expect("unique callback parameters should parse"); + assert_eq!(params.get("state").map(String::as_str), Some("state")); + assert_eq!(params.get("code").map(String::as_str), Some("code")); + } + + #[test] + fn redirect_location_appends_parameters_to_relative_targets() { + let location = build_redirect_location( + "/auth/callback", + Some(RedirectParams::Query(vec![( + "error_code", + "invalid_state".to_string(), + )])), + ); + assert_eq!(location, "/auth/callback?error_code=invalid_state"); + + let location = build_redirect_location( + "/auth/callback?existing=1#old", + Some(RedirectParams::Query(vec![( + "error_code", + "provider unavailable".to_string(), + )])), + ); + assert_eq!( + location, + "/auth/callback?existing=1&error_code=provider+unavailable#old" + ); + } + + #[test] + fn oauth_login_success_redirect_keeps_cookie_but_carries_no_access_token() { + let login_response = Response::builder() + .status(http::StatusCode::OK) + .header( + SET_COOKIE, + "aether_refresh_token=refresh-secret; Path=/api/auth; HttpOnly", + ) + .body(Body::from( + r#"{"access_token":"must-never-enter-callback-url"}"#, + )) + .expect("login response should build"); + let response = redirect_oauth_login_success( + "https://frontend.example/auth/callback?source=oauth", + &login_response, + ); + let location = response + .headers() + .get(LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("OAuth login redirect should have a location"); + + assert_eq!( + location, + "https://frontend.example/auth/callback?source=oauth" + ); + assert!(!location.contains("access_token")); + assert!(!location.contains("must-never-enter-callback-url")); + assert!(!location.contains('#')); + let set_cookie = response + .headers() + .get(SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .expect("refresh cookie should be preserved"); + assert!(set_cookie.starts_with("aether_refresh_token=refresh-secret")); + } } -async fn validate_bind_token( - state: &AppState, - provider_type: &str, - client_device_id: &str, - bind_token: &str, -) -> Result> { - let payload = decode_auth_token(bind_token, "oauth_bind").map_err(|detail| { - build_auth_error_response(http::StatusCode::UNAUTHORIZED, detail, false) - })?; - let token_provider = payload - .get("provider_type") - .and_then(serde_json::Value::as_str) - .unwrap_or_default(); - if token_provider != provider_type { - return Err(build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "绑定令牌不匹配", - false, - )); - } - let Some(user_id) = payload.get("user_id").and_then(serde_json::Value::as_str) else { - return Err(build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "绑定令牌无效", - false, - )); - }; - let Some(session_id) = payload - .get("session_id") - .and_then(serde_json::Value::as_str) - else { - return Err(build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "绑定令牌无效", - false, - )); - }; - let session = state - .find_user_session(user_id, session_id) - .await - .map_err(|err| { - build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("auth session lookup failed: {err:?}"), - false, - ) - })? - .ok_or_else(|| { - build_auth_error_response(http::StatusCode::UNAUTHORIZED, "绑定会话已失效", false) - })?; - if session.is_revoked() - || session.is_expired(chrono::Utc::now()) - || session.client_device_id != client_device_id - { - return Err(build_auth_error_response( - http::StatusCode::UNAUTHORIZED, - "绑定会话已失效", - false, - )); - } - Ok(ValidatedBindToken { - user_id: user_id.to_string(), - session_id: session_id.to_string(), - }) +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum CallbackParamsError { + DuplicateState, + DuplicateCallbackParameter, } fn callback_params( - request_context: &GatewayPublicRequestContext, -) -> std::collections::BTreeMap { - request_context - .request_query_string - .as_deref() - .map(|query| { - form_urlencoded::parse(query.as_bytes()) - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect() - }) - .unwrap_or_default() + query: Option<&str>, +) -> Result, CallbackParamsError> { + let mut params = std::collections::BTreeMap::new(); + for (key, value) in query + .into_iter() + .flat_map(|query| form_urlencoded::parse(query.as_bytes())) + { + let key = key.into_owned(); + if matches!(key.as_str(), "state" | "code" | "error") && params.contains_key(&key) { + return Err(if key == "state" { + CallbackParamsError::DuplicateState + } else { + CallbackParamsError::DuplicateCallbackParameter + }); + } + params.insert(key, value.into_owned()); + } + Ok(params) } fn public_oauth_provider_from_path(path: &str, suffix: &str) -> Option { - path.strip_prefix("/api/oauth/")? - .strip_suffix(&format!("/{suffix}"))? - .split('/') - .next() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_ascii_lowercase) + oauth_provider_from_path(path, "/api/oauth/", suffix) } fn user_oauth_provider_from_path(path: &str, suffix: &str) -> Option { - path.strip_prefix("/api/user/oauth/")? + oauth_provider_from_path(path, "/api/user/oauth/", suffix) +} + +fn oauth_provider_from_path(path: &str, prefix: &str, suffix: &str) -> Option { + let provider_type = path + .strip_prefix(prefix)? .strip_suffix(&format!("/{suffix}"))? - .split('/') - .next() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_ascii_lowercase) + .trim(); + (!provider_type.is_empty() && !provider_type.contains('/')) + .then(|| provider_type.to_ascii_lowercase()) } fn user_oauth_provider_from_path_without_suffix(path: &str) -> Option { - let provider_type = path.strip_prefix("/api/user/oauth/")?; + let provider_type = path.strip_prefix("/api/user/oauth/")?.trim(); (!provider_type.is_empty() && !provider_type.contains('/')) - .then(|| provider_type.trim().to_ascii_lowercase()) + .then(|| provider_type.to_ascii_lowercase()) } fn oauth_error_code(error: &OAuthError) -> &'static str { @@ -678,7 +1029,6 @@ fn oauth_account_error_response(error: crate::oauth::IdentityOAuthAccountError) enum RedirectParams { Query(Vec<(&'static str, String)>), - Fragment(Vec<(&'static str, String)>), } fn redirect_oauth_error(frontend_callback_url: Option<&str>, code: &str) -> Response { @@ -702,7 +1052,16 @@ fn redirect_to(target: &str, params: Option) -> Response { } fn build_redirect_location(target: &str, params: Option) -> String { - let Ok(mut url) = url::Url::parse(target) else { + let relative_target = + url::Url::parse(target).is_err() && target.starts_with('/') && !target.starts_with("//"); + let parsed_target = url::Url::parse(target).or_else(|_| { + if relative_target { + url::Url::parse("http://aether.invalid").and_then(|base| base.join(target)) + } else { + Err(url::ParseError::RelativeUrlWithoutBase) + } + }); + let Ok(mut url) = parsed_target else { return target.to_string(); }; match params { @@ -713,16 +1072,26 @@ fn build_redirect_location(target: &str, params: Option) -> Stri query.append_pair(key, &value); } } - url.to_string() - } - Some(RedirectParams::Fragment(items)) => { - let mut serializer = form_urlencoded::Serializer::new(String::new()); - for (key, value) in items { - serializer.append_pair(key, &value); + if relative_target { + relative_url_string(&url) + } else { + url.to_string() } - url.set_fragment(Some(&serializer.finish())); - url.to_string() } + None if relative_target => relative_url_string(&url), None => url.to_string(), } } + +fn relative_url_string(url: &url::Url) -> String { + let mut target = url.path().to_string(); + if let Some(query) = url.query() { + target.push('?'); + target.push_str(query); + } + if let Some(fragment) = url.fragment() { + target.push('#'); + target.push_str(fragment); + } + target +} diff --git a/apps/aether-gateway/src/handlers/public/support/payment.rs b/apps/aether-gateway/src/handlers/public/support/payment.rs index 798167a35..31585d3b6 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment.rs @@ -28,16 +28,34 @@ use self::payment_repository::{ handle_payment_callback_input_with_wallet_repository, handle_payment_callback_with_wallet_repository, process_payment_callback_input_with_wallet_repository, + reconcile_payment_callback_referral_rewards, }; use self::payment_shared::NormalizedPaymentCallbackRequest; const PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL: &str = "支付回调存储暂不可用"; +const PAYMENT_CALLBACK_PROCESSING_FAILED_DETAIL: &str = "支付回调处理失败"; fn build_payment_callback_storage_unavailable_response() -> Response { - build_auth_error_response( + // This is an intentional, stable public error. Do not route it through + // the generic 5xx sanitizer: payment providers use the detail to decide + // whether a callback should be retried, and the message contains no + // implementation detail. + build_auth_json_response( http::StatusCode::SERVICE_UNAVAILABLE, - PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL, - false, + serde_json::json!({ + "detail": PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL, + }), + None, + ) +} + +fn build_payment_callback_processing_failed_response() -> Response { + build_auth_json_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + serde_json::json!({ + "detail": PAYMENT_CALLBACK_PROCESSING_FAILED_DETAIL, + }), + None, ) } @@ -59,8 +77,9 @@ pub(super) async fn maybe_build_local_payment_callback_response( #[cfg(test)] mod tests { use super::{ + build_payment_callback_processing_failed_response, build_payment_callback_storage_unavailable_response, - PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL, + PAYMENT_CALLBACK_PROCESSING_FAILED_DETAIL, PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL, }; use axum::body::to_bytes; use axum::http; @@ -81,4 +100,20 @@ mod tests { json!({ "detail": PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL }) ); } + + #[tokio::test] + async fn payment_callback_processing_failure_is_a_stable_public_error() { + let response = build_payment_callback_processing_failed_response(); + + assert_eq!(response.status(), http::StatusCode::INTERNAL_SERVER_ERROR); + 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!( + payload, + json!({ "detail": PAYMENT_CALLBACK_PROCESSING_FAILED_DETAIL }) + ); + } } diff --git a/apps/aether-gateway/src/handlers/public/support/payment/alipay.rs b/apps/aether-gateway/src/handlers/public/support/payment/alipay.rs index 64bc2c513..9633c7d1b 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/alipay.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/alipay.rs @@ -1,7 +1,8 @@ use axum::{body::Body, http, response::Response}; use super::{ - process_payment_callback_input_with_wallet_repository, AppState, GatewayPublicRequestContext, + process_payment_callback_input_with_wallet_repository, + reconcile_payment_callback_referral_rewards, AppState, GatewayPublicRequestContext, }; use tracing::warn; @@ -24,24 +25,22 @@ pub(super) async fn handle_alipay_notify( let input = match crate::handlers::shared::verify_alipay_notify_callback(state, request_body).await { Ok(value) => value, - Err(detail) => { - warn!(error = %detail, "alipay notify verification failed"); + Err(_) => { + warn!( + error_category = "callback_verification_failed", + "alipay notify verification failed" + ); return alipay_plain(http::StatusCode::OK, "fail"); } }; - match process_payment_callback_input_with_wallet_repository(state, input).await { - Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { - order, - order_id, - .. - }) => { - if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await { - warn!( - error = ?err, - order_id = %order_id, - "failed to apply referral rewards for alipay callback" - ); - } + let outcome = process_payment_callback_input_with_wallet_repository(state, input).await; + if let Ok(outcome) = &outcome { + if !reconcile_payment_callback_referral_rewards(state, outcome, "alipay").await { + return alipay_plain(http::StatusCode::OK, "fail"); + } + } + match outcome { + Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { .. }) => { alipay_plain(http::StatusCode::OK, "success") } Ok( @@ -52,13 +51,21 @@ pub(super) async fn handle_alipay_notify( .. }, ) => alipay_plain(http::StatusCode::OK, "success"), - Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { - error, - .. - }) => { - warn!(error = %error, path = %request_context.request_path, "alipay notify processing failed"); + Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { .. }) => { + warn!( + error_category = "callback_rejected", + path = %request_context.request_path, + "alipay notify processing failed" + ); + alipay_plain(http::StatusCode::OK, "fail") + } + Err(response) => { + warn!( + status = %response.status(), + path = %request_context.request_path, + "alipay notify storage processing failed" + ); alipay_plain(http::StatusCode::OK, "fail") } - Err(response) => response, } } diff --git a/apps/aether-gateway/src/handlers/public/support/payment/epay.rs b/apps/aether-gateway/src/handlers/public/support/payment/epay.rs index fe0ea493a..4400502db 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/epay.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/epay.rs @@ -1,13 +1,19 @@ use std::collections::BTreeMap; use axum::{body::Body, http, response::Response}; +use hmac::{Hmac, Mac}; use md5::{Digest, Md5}; use serde_json::json; +use sha2::Sha256; use tracing::warn; use super::{payment_shared::payment_callback_payload_hash, AppState, GatewayPublicRequestContext}; -#[derive(Debug, Clone)] +const MAX_EPAY_ORDER_NO_BYTES: usize = 64; +const MAX_EPAY_GATEWAY_ORDER_ID_BYTES: usize = 123; +const MAX_EPAY_CHANNEL_BYTES: usize = 64; + +#[derive(Clone)] pub(crate) struct EpayMerchantConfig { pub(crate) endpoint_url: String, pub(crate) callback_base_url: Option, @@ -19,6 +25,25 @@ pub(crate) struct EpayMerchantConfig { pub(crate) channels: serde_json::Value, } +impl std::fmt::Debug for EpayMerchantConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("EpayMerchantConfig") + .field("endpoint_url", &"[REDACTED]") + .field( + "callback_base_url", + &self.callback_base_url.as_ref().map(|_| "[REDACTED]"), + ) + .field("merchant_id", &self.merchant_id) + .field("merchant_key", &"[REDACTED]") + .field("pay_currency", &self.pay_currency) + .field("usd_exchange_rate", &self.usd_exchange_rate) + .field("min_recharge_usd", &self.min_recharge_usd) + .field("channels", &"[REDACTED]") + .finish() + } +} + #[derive(Debug, Clone, PartialEq)] pub(crate) struct EpayChannelConfig { pub(crate) channel: String, @@ -121,82 +146,45 @@ pub(crate) fn epay_signature_valid(params: &BTreeMap, merchant_k let Some(sign) = params.get("sign") else { return false; }; - epay_sign(params, merchant_key).eq_ignore_ascii_case(sign.trim()) + let expected = epay_sign(params, merchant_key); + let provided = sign.trim().to_ascii_lowercase(); + let mut expected_mac = Hmac::::new_from_slice(b"aether-epay-signature-compare") + .expect("static epay comparison key should be valid"); + expected_mac.update(expected.as_bytes()); + let expected_tag = expected_mac.finalize().into_bytes(); + let mut provided_mac = Hmac::::new_from_slice(b"aether-epay-signature-compare") + .expect("static epay comparison key should be valid"); + provided_mac.update(provided.as_bytes()); + provided_mac.verify_slice(&expected_tag).is_ok() } -fn epay_submit_url(endpoint_url: &str) -> String { - let trimmed = endpoint_url.trim(); - if trimmed.is_empty() { - return trimmed.to_string(); - } - let Ok(mut url) = url::Url::parse(trimmed) else { - return trimmed.trim_end_matches('/').to_string(); - }; +fn epay_submit_url(endpoint_url: &str) -> Result { + let normalized = + crate::handlers::shared::normalize_payment_https_url(endpoint_url, "endpoint_url")?; + let mut url = url::Url::parse(&normalized) + .map_err(|_| "endpoint_url must be an absolute HTTPS URL".to_string())?; let path = url.path(); if path.is_empty() || path == "/" { url.set_path("submit.php"); } - url.to_string() + Ok(url.to_string()) } -fn normalize_epay_base_url(value: &str) -> Option { - let trimmed = value.trim().trim_end_matches('/'); - let parsed = url::Url::parse(trimmed).ok()?; - if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() { - return None; - } - Some(trimmed.to_string()) -} - -fn forwarded_header_first(value: String) -> Option { - value - .split(',') - .next() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) -} - -pub(crate) fn epay_callback_base_url( - configured: Option<&str>, - headers: &http::HeaderMap, - request_context: &GatewayPublicRequestContext, -) -> Option { - if let Some(value) = configured.and_then(normalize_epay_base_url) { - return Some(value); +pub(crate) fn epay_callback_base_url(configured: Option<&str>) -> Option { + if let Some(configured) = configured { + return crate::handlers::shared::normalize_payment_callback_base_url(configured).ok(); } - if let Some(value) = std::env::var("AETHER_PUBLIC_BASE_URL") + std::env::var("AETHER_PUBLIC_BASE_URL") .ok() .or_else(|| std::env::var("PUBLIC_BASE_URL").ok()) - .and_then(|value| normalize_epay_base_url(&value)) - { - return Some(value); - } - - let host = crate::headers::header_value_str(headers, crate::constants::FORWARDED_HOST_HEADER) - .and_then(forwarded_header_first) - .or_else(|| request_context.host_header.clone()) - .map(|value| value.trim().trim_end_matches('/').to_string()) - .filter(|value| { - !value.is_empty() - && !value.contains('/') - && !value.contains('\\') - && !value.contains('@') - && !value.contains(char::is_whitespace) - })?; - let proto = crate::headers::header_value_str(headers, crate::constants::FORWARDED_PROTO_HEADER) - .and_then(forwarded_header_first) - .map(|value| value.trim().trim_end_matches(':').to_ascii_lowercase()) - .filter(|value| value == "http" || value == "https") - .unwrap_or_else(|| "http".to_string()); - normalize_epay_base_url(&format!("{proto}://{host}")) + .and_then(|value| crate::handlers::shared::normalize_payment_callback_base_url(&value).ok()) } pub(crate) fn build_epay_checkout_url( config: &EpayMerchantConfig, input: &EpayCheckoutInput, -) -> serde_json::Value { +) -> Result { let money = format!("{:.2}", input.pay_amount); let mut params = BTreeMap::new(); params.insert("pid".to_string(), config.merchant_id.clone()); @@ -210,12 +198,12 @@ pub(crate) fn build_epay_checkout_url( let sign = epay_sign(¶ms, &config.merchant_key); params.insert("sign".to_string(), sign); - let payment_url = epay_submit_url(&config.endpoint_url); + let payment_url = epay_submit_url(&config.endpoint_url)?; let payment_params = params .iter() .map(|(key, value)| (key.clone(), serde_json::Value::String(value.clone()))) .collect::>(); - json!({ + Ok(json!({ "gateway": "epay", "display_name": "易支付", "gateway_order_id": input.order_no, @@ -226,7 +214,7 @@ pub(crate) fn build_epay_checkout_url( "pay_amount": input.pay_amount, "pay_currency": config.pay_currency, "payment_channel": input.channel, - }) + })) } pub(crate) fn parse_epay_params( @@ -243,34 +231,131 @@ pub(crate) fn parse_epay_params( .collect() } +fn epay_callback_projection( + _params: &BTreeMap, + order_no: &str, + gateway_order_id: Option<&str>, + pay_amount: f64, + pay_currency: &str, + payment_channel: Option<&str>, +) -> serde_json::Value { + json!({ + "gateway": "epay", + "event_id": gateway_order_id, + "gateway_order_id": gateway_order_id, + "order_no": order_no, + "amount": pay_amount, + "currency": pay_currency, + "payment_channel": payment_channel, + "status": "success", + "signature_valid": true, + }) +} + +fn bounded_epay_callback_value(value: Option<&String>, max_bytes: usize) -> Option { + let value = value?.trim(); + (!value.is_empty() + && value.len() <= max_bytes + && !value.bytes().any(|byte| byte.is_ascii_control())) + .then(|| value.to_string()) +} + pub(crate) async fn load_epay_config(state: &AppState) -> Result { let Some(record) = state .find_payment_gateway_config("epay") .await - .map_err(|err| format!("epay config lookup failed: {err:?}"))? + .map_err(|_| "epay config lookup failed".to_string())? else { return Err("epay is not configured".to_string()); }; if !record.enabled { return Err("epay is disabled".to_string()); } + let pay_currency = crate::handlers::shared::normalize_payment_currency( + &record.pay_currency, + "epay pay_currency", + ) + .map_err(|_| "epay pay_currency is invalid".to_string())?; + let usd_exchange_rate = crate::handlers::shared::effective_payment_exchange_rate( + &pay_currency, + record.usd_exchange_rate, + ) + .map_err(|_| "epay usd_exchange_rate is invalid".to_string())?; + if !record.min_recharge_usd.is_finite() || record.min_recharge_usd < 0.0 { + return Err("epay min_recharge_usd is invalid".to_string()); + } + let endpoint_url = + crate::handlers::shared::normalize_payment_https_url(&record.endpoint_url, "endpoint_url")?; + let callback_base_url = record + .callback_base_url + .as_deref() + .map(crate::handlers::shared::normalize_payment_callback_base_url) + .transpose()?; let Some(encrypted_key) = record.merchant_key_encrypted.as_deref() else { return Err("epay merchant key is missing".to_string()); }; - let Some(merchant_key) = crate::handlers::shared::decrypt_catalog_secret_with_fallbacks( - state.encryption_key(), - encrypted_key, - ) else { - return Err("epay merchant key decrypt failed".to_string()); + let binding = crate::handlers::shared::PaymentGatewaySecretBinding::from_record(&record) + .map_err(|_| "epay merchant key binding is invalid".to_string())?; + let merchant_key = + crate::handlers::shared::open_payment_gateway_secret(state, &binding, encrypted_key) + .map_err(|_| "epay merchant key decrypt failed".to_string())? + .plaintext; + Ok(EpayMerchantConfig { + endpoint_url, + callback_base_url, + merchant_id: record.merchant_id, + merchant_key, + pay_currency, + usd_exchange_rate, + min_recharge_usd: record.min_recharge_usd, + channels: record.channels_json, + }) +} + +/// Callback authentication must continue to work for an order created before +/// an administrator disabled EPay or changed checkout pricing. Only the +/// merchant identity and signing secret are needed to authenticate a notify; +/// settlement values are resolved from the stored payment order below. +pub(crate) async fn load_epay_callback_config( + state: &AppState, +) -> Result { + let Some(record) = state + .find_payment_gateway_config("epay") + .await + .map_err(|_| "epay config lookup failed".to_string())? + else { + return Err("epay is not configured".to_string()); }; + let Some(encrypted_key) = record.merchant_key_encrypted.as_deref() else { + return Err("epay merchant key is missing".to_string()); + }; + let binding = crate::handlers::shared::PaymentGatewaySecretBinding::from_record(&record) + .map_err(|_| "epay merchant key binding is invalid".to_string())?; + let merchant_key = + crate::handlers::shared::open_payment_gateway_secret(state, &binding, encrypted_key) + .map_err(|_| "epay merchant key decrypt failed".to_string())? + .plaintext; + // EPay notifications do not carry a currency or exchange-rate field. Use + // conservative finite fallbacks only when the order lookup below cannot + // provide the values; an unknown order is rejected by the repository. + let pay_currency = crate::handlers::shared::normalize_payment_currency( + &record.pay_currency, + "epay pay_currency", + ) + .unwrap_or_else(|_| "CNY".to_string()); + let usd_exchange_rate = crate::handlers::shared::effective_payment_exchange_rate( + &pay_currency, + record.usd_exchange_rate, + ) + .unwrap_or(1.0); Ok(EpayMerchantConfig { endpoint_url: record.endpoint_url, callback_base_url: record.callback_base_url, merchant_id: record.merchant_id, merchant_key, - pay_currency: record.pay_currency, - usd_exchange_rate: record.usd_exchange_rate, - min_recharge_usd: record.min_recharge_usd, + pay_currency, + usd_exchange_rate, + min_recharge_usd: 0.0, channels: record.channels_json, }) } @@ -326,7 +411,7 @@ pub(super) async fn handle_epay_notify( request_context: &GatewayPublicRequestContext, request_body: Option<&axum::body::Bytes>, ) -> Response { - let config = match load_epay_config(state).await { + let config = match load_epay_callback_config(state).await { Ok(value) => value, Err(_) => return epay_plain(http::StatusCode::OK, "fail"), }; @@ -337,50 +422,83 @@ pub(super) async fn handle_epay_notify( if !epay_signature_valid(¶ms, &config.merchant_key) { return epay_plain(http::StatusCode::OK, "fail"); } + if params.get("pid").map(String::as_str) != Some(config.merchant_id.as_str()) { + return epay_plain(http::StatusCode::OK, "fail"); + } if params.get("trade_status").map(String::as_str) != Some("TRADE_SUCCESS") { return epay_plain(http::StatusCode::OK, "fail"); } - let Some(order_no) = params.get("out_trade_no").cloned() else { + let Some(order_no) = + bounded_epay_callback_value(params.get("out_trade_no"), MAX_EPAY_ORDER_NO_BYTES) + else { return epay_plain(http::StatusCode::OK, "fail"); }; let Some(pay_amount) = params .get("money") .and_then(|value| value.parse::().ok()) + .filter(|value| value.is_finite() && *value > 0.0) else { return epay_plain(http::StatusCode::OK, "fail"); }; - let channel = params - .get("type") - .map(|value| value.trim().to_ascii_lowercase()) - .filter(|value| !value.is_empty()); - let payload = serde_json::to_value(¶ms).unwrap_or_else(|_| json!({})); - let payload_hash = match payment_callback_payload_hash(&payload) { + let Some(channel) = bounded_epay_callback_value(params.get("type"), MAX_EPAY_CHANNEL_BYTES) + .map(|value| value.to_ascii_lowercase()) + else { + return epay_plain(http::StatusCode::OK, "fail"); + }; + // Do not require the channel to remain enabled after checkout creation. + // The repository binds this signed value to the channel stored on the + // order, so removing a channel only blocks new checkouts and cannot strand + // an already-paid order. + let Some(gateway_order_id) = + bounded_epay_callback_value(params.get("trade_no"), MAX_EPAY_GATEWAY_ORDER_ID_BYTES) + else { + return epay_plain(http::StatusCode::OK, "fail"); + }; + let raw_payload = serde_json::to_value(¶ms).unwrap_or_else(|_| json!({})); + let payload_hash = match payment_callback_payload_hash(&raw_payload) { Ok(value) => value, Err(_) => return epay_plain(http::StatusCode::OK, "fail"), }; - let callback_key = params - .get("trade_no") - .cloned() - .unwrap_or_else(|| format!("epay:{order_no}:{payload_hash}")); - let amount_usd = if config.usd_exchange_rate > 0.0 { - pay_amount / config.usd_exchange_rate - } else { - pay_amount + let callback_key = format!("epay:{gateway_order_id}"); + let order = match crate::handlers::shared::find_payment_callback_order(state, &order_no).await { + Ok(value) => value, + Err(_) => return epay_plain(http::StatusCode::OK, "fail"), }; + let (amount_usd, exchange_rate) = + match crate::handlers::shared::payment_callback_settlement_values( + order.as_ref(), + pay_amount, + Some(config.usd_exchange_rate), + ) { + Ok(value) => value, + Err(_) => return epay_plain(http::StatusCode::OK, "fail"), + }; + let pay_currency = order + .as_ref() + .and_then(|order| order.pay_currency.clone()) + .unwrap_or_else(|| config.pay_currency.clone()); + let payload = epay_callback_projection( + ¶ms, + &order_no, + Some(&gateway_order_id), + pay_amount, + &pay_currency, + Some(&channel), + ); let outcome = state .process_payment_callback( aether_data::repository::wallet::ProcessPaymentCallbackInput { payment_method: "epay".to_string(), payment_provider: Some("epay".to_string()), - payment_channel: channel, + payment_channel: Some(channel), callback_key, order_no: Some(order_no), - gateway_order_id: params.get("trade_no").cloned(), + gateway_order_id: Some(gateway_order_id), amount_usd, pay_amount: Some(pay_amount), - pay_currency: Some(config.pay_currency), - exchange_rate: Some(config.usd_exchange_rate), + pay_currency: Some(pay_currency), + exchange_rate, payload_hash, payload, signature_valid: true, @@ -388,21 +506,18 @@ pub(super) async fn handle_epay_notify( ) .await; + if let Ok(Some(callback_outcome)) = &outcome { + if !super::reconcile_payment_callback_referral_rewards(state, callback_outcome, "epay") + .await + { + return epay_plain(http::StatusCode::OK, "fail"); + } + } + match outcome { Ok(Some(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { - order, - order_id, .. - })) => { - if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await { - warn!( - error = ?err, - order_id = %order_id, - "failed to apply referral rewards for epay callback" - ); - } - epay_plain(http::StatusCode::OK, "success") - } + })) => epay_plain(http::StatusCode::OK, "success"), Ok(Some( aether_data::repository::wallet::ProcessPaymentCallbackOutcome::AlreadyCredited { .. @@ -426,7 +541,7 @@ pub(super) async fn handle_epay_return( request_context.request_query_string.as_deref(), request_body, ); - let signature_valid = load_epay_config(state) + let signature_valid = load_epay_callback_config(state) .await .ok() .is_some_and(|config| epay_signature_valid(¶ms, &config.merchant_key)); @@ -436,8 +551,9 @@ pub(super) async fn handle_epay_return( #[cfg(test)] mod tests { use super::{ - build_epay_checkout_url, configured_epay_channels, epay_sign, epay_signature_valid, - resolve_epay_channel, EpayCheckoutInput, EpayMerchantConfig, + build_epay_checkout_url, configured_epay_channels, epay_callback_base_url, + epay_callback_projection, epay_sign, epay_signature_valid, resolve_epay_channel, + EpayCheckoutInput, EpayMerchantConfig, }; use chrono::Utc; use serde_json::json; @@ -457,6 +573,26 @@ mod tests { assert!(!epay_signature_valid(¶ms, "wrong")); } + #[test] + fn epay_config_debug_output_redacts_merchant_credentials() { + let mut config = test_epay_config(json!({"private": "epay-channel-canary"})); + config.endpoint_url = "https://pay.example/?token=epay-endpoint-canary".to_string(); + config.callback_base_url = + Some("https://callback.example/epay-callback-canary".to_string()); + config.merchant_key = "epay-merchant-key-canary".to_string(); + + let debug = format!("{config:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "epay-channel-canary", + "epay-endpoint-canary", + "epay-callback-canary", + "epay-merchant-key-canary", + ] { + assert!(!debug.contains(secret), "debug output leaked {secret}"); + } + } + #[test] fn configured_epay_channels_do_not_invent_defaults() { let mut config = test_epay_config(json!([ @@ -509,7 +645,8 @@ mod tests { notify_url: "https://aether.example.com/api/payment/epay/notify".to_string(), return_url: "https://aether.example.com/api/payment/epay/return".to_string(), }, - ); + ) + .expect("valid HTTPS endpoint should build checkout"); assert_eq!( checkout["payment_url"], @@ -536,13 +673,89 @@ mod tests { notify_url: "https://aether.example.com/api/payment/epay/notify".to_string(), return_url: "https://aether.example.com/api/payment/epay/return".to_string(), }, - ); + ) + .expect("valid HTTPS endpoint should build checkout"); assert_eq!( checkout["payment_url"], "https://pay.example.com/submit.php" ); } + #[test] + fn epay_checkout_rejects_executable_or_insecure_endpoint_urls() { + for endpoint_url in [ + "javascript:alert(document.domain)", + "data:text/html,attack", + "http://pay.example.com/submit.php", + "/submit.php", + ] { + let mut config = test_epay_config(json!([])); + config.endpoint_url = endpoint_url.to_string(); + let result = build_epay_checkout_url( + &config, + &EpayCheckoutInput { + order_no: "po_unsafe".to_string(), + channel: "alipay".to_string(), + subject: "wallet recharge".to_string(), + pay_amount: 1.0, + notify_url: "https://aether.example/api/payment/epay/notify".to_string(), + return_url: "https://aether.example/api/payment/epay/return".to_string(), + }, + ); + assert!( + result.is_err(), + "unsafe endpoint should fail: {endpoint_url}" + ); + } + } + + #[test] + fn epay_callback_base_requires_explicit_https_configuration() { + assert_eq!( + epay_callback_base_url(Some("https://aether.example/")), + Some("https://aether.example".to_string()) + ); + assert_eq!(epay_callback_base_url(Some("http://aether.example")), None); + assert_eq!( + epay_callback_base_url(Some("https://user:secret@aether.example")), + None + ); + } + + #[test] + fn epay_callback_projection_does_not_persist_signature_or_payer_fields() { + let params = BTreeMap::from([ + ("trade_no".to_string(), "gateway-1".to_string()), + ("sign".to_string(), "replayable-signature".to_string()), + ("buyer_email".to_string(), "payer@example.com".to_string()), + ("name".to_string(), "private subject".to_string()), + ]); + + let projection = epay_callback_projection( + ¶ms, + "po_1", + Some("gateway-1"), + 72.0, + "CNY", + Some("alipay"), + ); + assert_eq!(projection["event_id"], "gateway-1"); + assert_eq!(projection["gateway_order_id"], "gateway-1"); + assert_eq!(projection["order_no"], "po_1"); + assert_eq!(projection["payment_channel"], "alipay"); + assert!(projection.get("sign").is_none()); + + let encoded = projection.to_string(); + for forbidden in [ + "replayable-signature", + "buyer_email", + "payer@example.com", + "private subject", + ] { + assert!(!encoded.contains(forbidden), "persisted {forbidden}"); + } + } + fn test_epay_config(channels: serde_json::Value) -> EpayMerchantConfig { EpayMerchantConfig { endpoint_url: "https://pay.example.com/submit.php".to_string(), diff --git a/apps/aether-gateway/src/handlers/public/support/payment/repository.rs b/apps/aether-gateway/src/handlers/public/support/payment/repository.rs index 0ef871cee..cd0134431 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/repository.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/repository.rs @@ -1,15 +1,15 @@ use super::payment_shared::{ - payment_callback_mark_failed_response, payment_callback_payload_hash, + generic_payment_callback_method_allowed, payment_callback_mark_failed_response, + payment_callback_namespaced_key, payment_callback_payload_hash, + payment_callback_persistence_projection, payment_callback_success_response, NormalizedPaymentCallbackRequest, }; use aether_data::repository::wallet::{ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome}; use axum::{body::Body, http, response::Response}; -use serde_json::json; -use super::super::build_auth_json_response; use super::{ - build_auth_error_response, build_payment_callback_storage_unavailable_response, AppState, - GatewayPublicRequestContext, + build_auth_error_response, build_payment_callback_processing_failed_response, + build_payment_callback_storage_unavailable_response, AppState, GatewayPublicRequestContext, }; use tracing::warn; @@ -20,24 +20,34 @@ pub(super) async fn handle_payment_callback_with_wallet_repository( payload: &NormalizedPaymentCallbackRequest, signature_valid: bool, ) -> Response { + if !generic_payment_callback_method_allowed(payment_method) { + return build_auth_error_response( + http::StatusCode::NOT_FOUND, + "payment callback route not found", + false, + ); + } + let payment_method = payment_method.trim().to_ascii_lowercase(); if !state.has_database_wallet_data_writer() { return build_payment_callback_storage_unavailable_response(); } let callback_payload_hash = match payment_callback_payload_hash(&payload.payload) { Ok(value) => value, - Err(err) => { - return build_auth_error_response(http::StatusCode::INTERNAL_SERVER_ERROR, err, false) - } + Err(_) => return build_payment_callback_processing_failed_response(), }; + let persisted_payload = + payment_callback_persistence_projection(&payment_method, payload, signature_valid); handle_payment_callback_input_with_wallet_repository( state, request_context, ProcessPaymentCallbackInput { - payment_method: payment_method.to_string(), + payment_method: payment_method.clone(), payment_provider: None, payment_channel: None, - callback_key: payload.callback_key.clone(), + // Keep idempotency keys in a provider namespace. A merchant can + // legitimately reuse an event id across separate callback routes. + callback_key: payment_callback_namespaced_key(&payment_method, &payload.callback_key), order_no: payload.order_no.clone(), gateway_order_id: payload.gateway_order_id.clone(), amount_usd: payload.amount_usd, @@ -45,7 +55,7 @@ pub(super) async fn handle_payment_callback_with_wallet_repository( pay_currency: payload.pay_currency.clone(), exchange_rate: payload.exchange_rate, payload_hash: callback_payload_hash, - payload: payload.payload.clone(), + payload: persisted_payload, signature_valid, }, ) @@ -54,7 +64,7 @@ pub(super) async fn handle_payment_callback_with_wallet_repository( pub(super) async fn handle_payment_callback_input_with_wallet_repository( state: &AppState, - request_context: &GatewayPublicRequestContext, + _request_context: &GatewayPublicRequestContext, input: ProcessPaymentCallbackInput, ) -> Response { let payment_method = input.payment_method.clone(); @@ -62,84 +72,72 @@ pub(super) async fn handle_payment_callback_input_with_wallet_repository( Ok(value) => value, Err(response) => return response, }; + if !reconcile_payment_callback_referral_rewards(state, &outcome, &payment_method).await { + // The payment credit is durable, but referral creation/crediting uses + // its own transaction. Do not acknowledge the webhook until that + // obligation is either applied or durably represented: providers can + // then replay the same callback and the idempotent payment path will + // retry referral coordination without crediting the order twice. + return build_payment_callback_processing_failed_response(); + } match outcome { aether_data::repository::wallet::ProcessPaymentCallbackOutcome::DuplicateProcessed { - order_id, - } => build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": true, - "credited": false, - "order_id": order_id, - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ), - aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { - duplicate, - error, - } => payment_callback_mark_failed_response( - duplicate, - &error, - &payment_method, - &request_context.request_path, - ), + .. + } => payment_callback_success_response(), + aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { .. } => { + payment_callback_mark_failed_response() + } aether_data::repository::wallet::ProcessPaymentCallbackOutcome::AlreadyCredited { - duplicate, - order_id, - order_no, - wallet_id, - } => build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": duplicate, - "credited": false, - "order_id": order_id, - "order_no": order_no, - "status": "credited", - "wallet_id": wallet_id, - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ), - aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { - duplicate, - order_id, - order_no, - wallet_id, - order, - } => { - if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await { - warn!( - error = ?err, - order_id = %order_id, - "failed to apply referral rewards for credited payment order" - ); - } - build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": duplicate, - "credited": true, - "order_id": order_id, - "order_no": order_no, - "status": order.status, - "wallet_id": wallet_id, - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ) + .. + } => payment_callback_success_response(), + aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { .. } => { + payment_callback_success_response() } } } +pub(super) async fn reconcile_payment_callback_referral_rewards( + state: &AppState, + outcome: &ProcessPaymentCallbackOutcome, + payment_method: &str, +) -> bool { + let result = match outcome { + ProcessPaymentCallbackOutcome::Applied { order, .. } => { + state.apply_referral_rewards_for_paid_order(order).await + } + ProcessPaymentCallbackOutcome::AlreadyCredited { order_id, .. } + | ProcessPaymentCallbackOutcome::DuplicateProcessed { + order_id: Some(order_id), + } => { + state + .apply_referral_rewards_for_payment_order_id(order_id) + .await + } + ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id: None } => { + // A processed callback must be linked to the credited order. A + // missing link is corrupted durable state, so acknowledging it + // would permanently skip referral reconciliation. + warn!( + error_category = "payment_callback_missing_order_link", + payment_method, "processed payment callback has no payment order link" + ); + return false; + } + ProcessPaymentCallbackOutcome::Failed { .. } => return true, + }; + if let Err(error) = result { + warn!( + error = ?error, + error_category = "referral_reward_apply_failed", + payment_method, + "failed to reconcile referral rewards for credited payment order" + ); + return false; + } + true +} + pub(super) async fn process_payment_callback_input_with_wallet_repository( state: &AppState, input: ProcessPaymentCallbackInput, @@ -151,21 +149,25 @@ pub(super) async fn process_payment_callback_input_with_wallet_repository( match state.process_payment_callback(input).await { Ok(Some(value)) => Ok(value), Ok(None) => Err(build_payment_callback_storage_unavailable_response()), - Err(err) => Err(build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("payment callback failed: {err:?}"), - false, - )), + Err(_) => { + warn!( + error_category = "wallet_repository_failed", + "payment callback repository processing failed" + ); + Err(build_payment_callback_processing_failed_response()) + } } } #[cfg(test)] mod tests { use super::{ - handle_payment_callback_with_wallet_repository, AppState, NormalizedPaymentCallbackRequest, + handle_payment_callback_with_wallet_repository, + reconcile_payment_callback_referral_rewards, AppState, NormalizedPaymentCallbackRequest, }; use crate::control::GatewayPublicRequestContext; use crate::handlers::public::support::support_payment::PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL; + use aether_data::repository::wallet::ProcessPaymentCallbackOutcome; use axum::body::to_bytes; use axum::http::{HeaderMap, Method, Uri}; use serde_json::json; @@ -176,7 +178,7 @@ mod tests { let request_context = GatewayPublicRequestContext::from_request_parts( "trace-payment-callback-wallet-writer-missing", &Method::POST, - &"/api/payment/callback/alipay" + &"/api/payment/callback/manual" .parse::() .expect("uri should parse"), &HeaderMap::new(), @@ -195,7 +197,7 @@ mod tests { let response = handle_payment_callback_with_wallet_repository( &state, - "alipay", + "manual", &request_context, &payload, true, @@ -213,4 +215,54 @@ mod tests { json!({ "detail": PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL }) ); } + + #[tokio::test] + async fn generic_repository_entrypoint_rejects_official_payment_providers() { + let state = AppState::new().expect("state should build"); + let request_context = GatewayPublicRequestContext::from_request_parts( + "trace-official-payment-callback-bypass", + &Method::POST, + &"/api/payment/callback/alipay" + .parse::() + .expect("uri should parse"), + &HeaderMap::new(), + None, + ); + let payload = NormalizedPaymentCallbackRequest { + callback_key: "callback-key-1".to_string(), + order_no: Some("order-no-1".to_string()), + gateway_order_id: Some("gateway-order-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload: json!({ "status": "paid" }), + }; + + for provider in ["alipay", "wxpay", "stripe", "epay"] { + let response = handle_payment_callback_with_wallet_repository( + &state, + provider, + &request_context, + &payload, + true, + ) + .await; + assert_eq!(response.status(), http::StatusCode::NOT_FOUND); + } + } + + #[tokio::test] + async fn processed_callback_without_order_link_is_not_acknowledged() { + let state = AppState::new().expect("state should build"); + + assert!( + !reconcile_payment_callback_referral_rewards( + &state, + &ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id: None }, + "stripe", + ) + .await + ); + } } diff --git a/apps/aether-gateway/src/handlers/public/support/payment/route.rs b/apps/aether-gateway/src/handlers/public/support/payment/route.rs index e6182ae37..0ebd27e5e 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/route.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/route.rs @@ -2,7 +2,8 @@ use axum::{body::Body, http, response::Response}; use super::payment_gateway::{PaymentGatewayRegistry, VerifyCallbackInput}; use super::payment_shared::{ - payment_callback_payment_method_from_path, payment_callback_secret, PaymentCallbackRequest, + generic_payment_callback_method_allowed, payment_callback_payment_method_from_path, + payment_callback_secret, payment_callback_secret_matches, PaymentCallbackRequest, PAYMENT_CALLBACK_SIGNATURE_HEADER, PAYMENT_CALLBACK_TOKEN_HEADER, }; use super::{ @@ -53,6 +54,23 @@ pub(super) async fn maybe_build_local_payment_callback_route_response( return None; } + let Some(payment_method) = + payment_callback_payment_method_from_path(&request_context.request_path) + else { + return Some(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "payment_method is required", + false, + )); + }; + if !generic_payment_callback_method_allowed(&payment_method) { + return Some(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "payment callback route not found", + false, + )); + } + let Some(secret) = payment_callback_secret() else { return Some(build_auth_error_response( http::StatusCode::SERVICE_UNAVAILABLE, @@ -69,7 +87,7 @@ pub(super) async fn maybe_build_local_payment_callback_route_response( false, )); }; - if provided_token.trim() != secret { + if !payment_callback_secret_matches(&provided_token, &secret) { return Some(build_auth_error_response( http::StatusCode::UNAUTHORIZED, "invalid payment callback token", @@ -102,15 +120,6 @@ pub(super) async fn maybe_build_local_payment_callback_route_response( )); } }; - let Some(payment_method) = - payment_callback_payment_method_from_path(&request_context.request_path) - else { - return Some(build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "payment_method is required", - false, - )); - }; let Some(adapter) = PaymentGatewayRegistry::get(&payment_method) else { return Some(build_auth_error_response( http::StatusCode::BAD_REQUEST, @@ -165,3 +174,18 @@ pub(super) async fn maybe_build_local_payment_callback_route_response( Some(build_payment_callback_storage_unavailable_response()) } } + +#[cfg(test)] +mod tests { + use super::generic_payment_callback_method_allowed; + + #[test] + fn generic_hmac_callback_excludes_official_direct_providers() { + for provider in ["alipay", "ALIPAY", " wxpay ", "Stripe", "epay"] { + assert!(!generic_payment_callback_method_allowed(provider)); + } + for payment_method in ["manual", "wechat"] { + assert!(generic_payment_callback_method_allowed(payment_method)); + } + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/payment/shared.rs b/apps/aether-gateway/src/handlers/public/support/payment/shared.rs index b93330233..7d9266ef1 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/shared.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/shared.rs @@ -5,11 +5,15 @@ use serde_json::json; use sha2::{Digest, Sha256}; use super::super::{build_auth_json_response, wallet_normalize_optional_string_field}; +use crate::handlers::shared::normalize_payment_currency; pub(super) const PAYMENT_CALLBACK_TOKEN_HEADER: &str = "x-payment-callback-token"; pub(super) const PAYMENT_CALLBACK_SIGNATURE_HEADER: &str = "x-payment-callback-signature"; +const PAYMENT_CALLBACK_SECRET_MIN_BYTES: usize = 32; +const PAYMENT_CALLBACK_SECRET_MIN_UNIQUE_BYTES: usize = 8; +const PAYMENT_CALLBACK_KEY_MAX_CHARS: usize = 128; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] pub(crate) struct PaymentCallbackRequest { pub(crate) callback_key: String, #[serde(default)] @@ -27,7 +31,23 @@ pub(crate) struct PaymentCallbackRequest { pub(crate) payload: Option>, } -#[derive(Debug, Clone)] +impl std::fmt::Debug for PaymentCallbackRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PaymentCallbackRequest") + .field("callback_key", &"[REDACTED]") + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field("amount_usd", &self.amount_usd) + .field("pay_amount", &self.pay_amount) + .field("pay_currency", &self.pay_currency) + .field("exchange_rate", &self.exchange_rate) + .field("payload", &self.payload.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + +#[derive(Clone)] pub(crate) struct NormalizedPaymentCallbackRequest { pub(crate) callback_key: String, pub(crate) order_no: Option, @@ -39,11 +59,46 @@ pub(crate) struct NormalizedPaymentCallbackRequest { pub(crate) payload: serde_json::Value, } +impl std::fmt::Debug for NormalizedPaymentCallbackRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("NormalizedPaymentCallbackRequest") + .field("callback_key", &"[REDACTED]") + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field("amount_usd", &self.amount_usd) + .field("pay_amount", &self.pay_amount) + .field("pay_currency", &self.pay_currency) + .field("exchange_rate", &self.exchange_rate) + .field("payload", &"[REDACTED]") + .finish() + } +} + pub(super) fn payment_callback_secret() -> Option { std::env::var("PAYMENT_CALLBACK_SECRET") .ok() .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) + .filter(|value| payment_callback_secret_is_strong(value)) +} + +fn payment_callback_secret_is_strong(value: &str) -> bool { + let bytes = value.as_bytes(); + if bytes.len() < PAYMENT_CALLBACK_SECRET_MIN_BYTES + || bytes.iter().any(|byte| byte.is_ascii_control()) + { + return false; + } + let mut seen = [false; 256]; + let mut unique = 0usize; + for byte in bytes { + let index = usize::from(*byte); + if !seen[index] { + seen[index] = true; + unique += 1; + } + } + unique >= PAYMENT_CALLBACK_SECRET_MIN_UNIQUE_BYTES } pub(super) fn payment_callback_payment_method_from_path(path: &str) -> Option { @@ -56,6 +111,13 @@ pub(super) fn payment_callback_payment_method_from_path(path: &str) -> Option bool { + !matches!( + payment_method.trim().to_ascii_lowercase().as_str(), + "alipay" | "wxpay" | "stripe" | "epay" + ) +} + fn normalize_payment_callback_optional_string( value: Option, max_chars: usize, @@ -82,26 +144,23 @@ pub(super) fn normalize_payment_callback_request( let order_no = normalize_payment_callback_optional_string(payload.order_no, 64)?; let gateway_order_id = normalize_payment_callback_optional_string(payload.gateway_order_id, 128)?; - let pay_currency = normalize_payment_callback_optional_string(payload.pay_currency, 3)?; - if matches!(pay_currency.as_deref(), Some(value) if value.chars().count() != 3) { - return Err("输入验证失败"); - } + let pay_currency = normalize_payment_callback_optional_string(payload.pay_currency, 3)? + .map(|value| normalize_payment_currency(&value, "pay_currency")) + .transpose() + .map_err(|_| "输入验证失败")?; - let payload_value = payload - .payload - .map(serde_json::Value::Object) - .unwrap_or_else(|| { - json!({ - "callback_key": callback_key, - "order_no": order_no, - "gateway_order_id": gateway_order_id, - "amount_usd": payload.amount_usd, - "pay_amount": payload.pay_amount, - "pay_currency": pay_currency, - "exchange_rate": payload.exchange_rate, - "payload": serde_json::Value::Null, - }) - }); + // The signature covers the complete settlement envelope. Signing only the + // provider-specific payload would leave the order and amount fields mutable. + let payload_value = json!({ + "callback_key": callback_key, + "order_no": order_no, + "gateway_order_id": gateway_order_id, + "amount_usd": payload.amount_usd, + "pay_amount": payload.pay_amount, + "pay_currency": pay_currency, + "exchange_rate": payload.exchange_rate, + "payload": payload.payload.map(serde_json::Value::Object), + }); Ok(NormalizedPaymentCallbackRequest { callback_key: callback_key.to_string(), @@ -136,17 +195,32 @@ fn payment_callback_canonicalize_json(value: &serde_json::Value) -> serde_json:: } } -fn payment_callback_signature_hex( - payload: &serde_json::Value, - secret: &str, -) -> Result { - let canonical = serde_json::to_string(payload) - .map_err(|err| format!("payment callback canonicalization failed: {err}"))?; - let mut mac = Hmac::::new_from_slice(secret.as_bytes()) - .map_err(|err| format!("payment callback hmac init failed: {err}"))?; - mac.update(canonical.as_bytes()); - let bytes = mac.finalize().into_bytes(); - Ok(bytes.iter().map(|byte| format!("{byte:02x}")).collect()) +fn decode_payment_callback_signature(value: &str) -> Option> { + let value = value.trim().strip_prefix("sha256=").unwrap_or(value.trim()); + if value.len() != 64 || !value.is_ascii() { + return None; + } + value + .as_bytes() + .chunks_exact(2) + .map(|chunk| { + std::str::from_utf8(chunk) + .ok() + .and_then(|hex| u8::from_str_radix(hex, 16).ok()) + }) + .collect() +} + +pub(super) fn payment_callback_secret_matches(provided: &str, expected: &str) -> bool { + let mut expected_mac = Hmac::::new_from_slice(b"aether-payment-callback-token") + .expect("static payment callback comparison key should be valid"); + expected_mac.update(expected.as_bytes()); + let expected_tag = expected_mac.finalize().into_bytes(); + + let mut provided_mac = Hmac::::new_from_slice(b"aether-payment-callback-token") + .expect("static payment callback comparison key should be valid"); + provided_mac.update(provided.trim().as_bytes()); + provided_mac.verify_slice(&expected_tag).is_ok() } pub(super) fn payment_callback_signature_matches( @@ -154,37 +228,179 @@ pub(super) fn payment_callback_signature_matches( provided_signature: &str, secret: &str, ) -> Result { - let expected = payment_callback_signature_hex(payload, secret)?; - let provided = provided_signature - .trim() - .strip_prefix("sha256=") - .unwrap_or(provided_signature.trim()) - .to_ascii_lowercase(); - Ok(provided == expected) + let Some(provided) = decode_payment_callback_signature(provided_signature) else { + return Ok(false); + }; + let canonical = serde_json::to_string(payload) + .map_err(|_| "payment callback canonicalization failed".to_string())?; + let mut mac = Hmac::::new_from_slice(secret.as_bytes()) + .map_err(|_| "payment callback hmac init failed".to_string())?; + mac.update(canonical.as_bytes()); + Ok(mac.verify_slice(&provided).is_ok()) } pub(super) fn payment_callback_payload_hash(payload: &serde_json::Value) -> Result { let encoded = serde_json::to_vec(payload) - .map_err(|err| format!("payment callback payload encode failed: {err}"))?; + .map_err(|_| "payment callback payload encode failed".to_string())?; let digest = Sha256::digest(&encoded); Ok(digest.iter().map(|byte| format!("{byte:02x}")).collect()) } -pub(super) fn payment_callback_mark_failed_response( - duplicate: bool, - error: &str, +pub(super) fn payment_callback_namespaced_key(payment_method: &str, callback_key: &str) -> String { + let payment_method = payment_method.trim().to_ascii_lowercase(); + let prefix = format!("{payment_method}:"); + if prefix.chars().count() + callback_key.chars().count() <= PAYMENT_CALLBACK_KEY_MAX_CHARS { + return format!("{prefix}{callback_key}"); + } + + let digest = Sha256::digest(callback_key.as_bytes()); + let digest = digest + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::(); + format!("{prefix}{digest}") +} + +pub(super) fn payment_callback_persistence_projection( payment_method: &str, - request_path: &str, -) -> Response { + payload: &NormalizedPaymentCallbackRequest, + signature_valid: bool, +) -> serde_json::Value { + json!({ + "gateway": payment_method, + "order_no": payload.order_no, + "gateway_order_id": payload.gateway_order_id, + "amount_usd": payload.amount_usd, + "pay_amount": payload.pay_amount, + "pay_currency": payload.pay_currency, + "exchange_rate": payload.exchange_rate, + "signature_valid": signature_valid, + }) +} + +pub(super) fn payment_callback_success_response() -> Response { + build_auth_json_response(http::StatusCode::OK, json!({ "ok": true }), None) +} + +pub(super) fn payment_callback_mark_failed_response() -> Response { build_auth_json_response( http::StatusCode::OK, json!({ "ok": false, - "duplicate": duplicate, - "error": error, - "payment_method": payment_method, - "request_path": request_path, + "error": "payment callback rejected", }), None, ) } + +#[cfg(test)] +mod tests { + use super::{ + payment_callback_mark_failed_response, payment_callback_namespaced_key, + payment_callback_persistence_projection, payment_callback_secret_is_strong, + payment_callback_success_response, NormalizedPaymentCallbackRequest, + }; + use axum::body::to_bytes; + use serde_json::json; + + #[test] + fn payment_callback_secret_rejects_short_or_obviously_low_entropy_values() { + assert!(!payment_callback_secret_is_strong("callback-secret-test")); + assert!(!payment_callback_secret_is_strong(&"a".repeat(64))); + assert!(!payment_callback_secret_is_strong( + "0123456789abcdef\n0123456789abcdef" + )); + assert!(payment_callback_secret_is_strong( + "test-callback-secret-0123456789abcdef" + )); + } + + #[test] + fn payment_callback_key_namespace_is_canonical_and_bounded() { + assert_eq!( + payment_callback_namespaced_key(" MANUAL ", "event-1"), + "manual:event-1" + ); + + let first = payment_callback_namespaced_key("MANUAL", &"a".repeat(128)); + let second = payment_callback_namespaced_key("manual", &"b".repeat(128)); + assert!(first.starts_with("manual:")); + assert!(first.chars().count() <= 128); + assert_ne!(first, second); + } + + #[tokio::test] + async fn payment_callback_responses_are_fixed_and_do_not_expose_processing_state() { + let response = payment_callback_mark_failed_response(); + 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("body should be JSON"); + assert_eq!( + payload, + json!({ "ok": false, "error": "payment callback rejected" }) + ); + + let response = payment_callback_success_response(); + 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("body should be JSON"); + assert_eq!(payload, json!({ "ok": true })); + } + + #[test] + fn payment_callback_persistence_projection_excludes_arbitrary_signed_payload() { + let payload = NormalizedPaymentCallbackRequest { + callback_key: "event-1".to_string(), + order_no: Some("po_1".to_string()), + gateway_order_id: Some("gateway-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payload: json!({ + "payload": { + "client_secret": "pi_1_secret_replayable", + "customer_email": "payer@example.com", + "authorization": "Bearer upstream-secret" + } + }), + }; + + let debug = format!("{payload:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "event-1", + "pi_1_secret_replayable", + "payer@example.com", + "upstream-secret", + ] { + assert!(!debug.contains(secret), "debug output leaked {secret}"); + } + + let projected = payment_callback_persistence_projection("manual", &payload, true); + assert_eq!(projected["gateway"], "manual"); + assert_eq!(projected["order_no"], "po_1"); + assert_eq!(projected["gateway_order_id"], "gateway-1"); + assert_eq!(projected["pay_currency"], "CNY"); + assert!(projected.get("status").is_none()); + let encoded = projected.to_string(); + for forbidden in [ + "client_secret", + "replayable", + "customer_email", + "payer@example.com", + "authorization", + "upstream-secret", + ] { + assert!(!encoded.contains(forbidden), "persisted {forbidden}"); + } + + let rejected = payment_callback_persistence_projection("manual", &payload, false); + assert_eq!(rejected["signature_valid"], false); + assert!(rejected.get("status").is_none()); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/payment/stripe.rs b/apps/aether-gateway/src/handlers/public/support/payment/stripe.rs index da7b5ac45..fa66847b2 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/stripe.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/stripe.rs @@ -4,9 +4,12 @@ use serde_json::{json, Value}; use sha2::Sha256; use super::{ - build_auth_error_response, build_auth_json_response, - handle_payment_callback_input_with_wallet_repository, - payment_shared::payment_callback_payload_hash, AppState, GatewayPublicRequestContext, + build_auth_error_response, handle_payment_callback_input_with_wallet_repository, + payment_shared::{ + payment_callback_namespaced_key, payment_callback_payload_hash, + payment_callback_success_response, + }, + AppState, GatewayPublicRequestContext, }; const STRIPE_SIGNATURE_HEADER: &str = "stripe-signature"; @@ -21,12 +24,12 @@ fn decrypt_gateway_secrets( let Some(encrypted) = record.merchant_key_encrypted.as_deref() else { return Err("Stripe webhook_secret 未配置".to_string()); }; - let Some(plaintext) = crate::handlers::shared::decrypt_catalog_secret_with_fallbacks( - state.encryption_key(), - encrypted, - ) else { - return Err("Stripe 密钥解密失败".to_string()); - }; + let binding = crate::handlers::shared::PaymentGatewaySecretBinding::from_record(record) + .map_err(|_| "Stripe 密钥绑定无效".to_string())?; + let plaintext = + crate::handlers::shared::open_payment_gateway_secret(state, &binding, encrypted) + .map_err(|_| "Stripe 密钥解密失败".to_string())? + .plaintext; serde_json::from_str::(&plaintext) .ok() .and_then(|value| value.as_object().cloned()) @@ -46,13 +49,10 @@ async fn stripe_webhook_secret(state: &AppState) -> Result { let Some(record) = state .find_payment_gateway_config("stripe") .await - .map_err(|err| format!("Stripe 配置读取失败: {err:?}"))? + .map_err(|_| "Stripe 配置读取失败".to_string())? else { return Err("Stripe 未配置".to_string()); }; - if !record.enabled { - return Err("Stripe 未启用".to_string()); - } let secrets = decrypt_gateway_secrets(state, &record)?; gateway_secret_string(&secrets, "webhook_secret") .ok_or_else(|| "Stripe webhook_secret 未配置".to_string()) @@ -105,7 +105,7 @@ fn stripe_signature_matches_at( if signatures.is_empty() { return Ok(false); } - if (now_unix_secs - timestamp).abs() > STRIPE_SIGNATURE_TOLERANCE_SECONDS { + if now_unix_secs.abs_diff(timestamp) > STRIPE_SIGNATURE_TOLERANCE_SECONDS as u64 { return Ok(false); } @@ -114,7 +114,7 @@ fn stripe_signature_matches_at( continue; }; let mut mac = HmacSha256::new_from_slice(secret.as_bytes()) - .map_err(|err| format!("Stripe webhook HMAC 初始化失败: {err}"))?; + .map_err(|_| "Stripe webhook HMAC 初始化失败".to_string())?; mac.update(timestamp.to_string().as_bytes()); mac.update(b"."); mac.update(body); @@ -125,18 +125,6 @@ fn stripe_signature_matches_at( Ok(false) } -fn stripe_amount_multiplier(currency: &str) -> f64 { - match currency.trim().to_ascii_lowercase().as_str() { - "bif" | "clp" | "djf" | "gnf" | "jpy" | "kmf" | "krw" | "mga" | "pyg" | "rwf" | "ugx" - | "vnd" | "vuv" | "xaf" | "xof" | "xpf" => 1.0, - _ => 100.0, - } -} - -fn stripe_amount_to_major(amount_minor: i64, currency: &str) -> f64 { - amount_minor as f64 / stripe_amount_multiplier(currency) -} - fn stripe_string_field<'a>(value: &'a Value, key: &str) -> Option<&'a str> { value .get(key) @@ -157,6 +145,27 @@ fn stripe_payment_intent_channel(intent: &Value) -> Option { .map(ToOwned::to_owned) } +fn stripe_callback_projection( + event: &Value, + intent_id: &str, + order_no: &str, + pay_amount: f64, + currency: &str, + payment_channel: Option<&str>, +) -> Value { + json!({ + "gateway": "stripe", + "event_id": stripe_string_field(event, "id"), + "gateway_order_id": intent_id, + "order_no": order_no, + "amount": pay_amount, + "currency": currency, + "payment_channel": payment_channel, + "status": "success", + "signature_valid": true, + }) +} + async fn build_stripe_callback_input( state: &AppState, event: Value, @@ -173,6 +182,9 @@ async fn build_stripe_callback_input( .get("data") .and_then(|value| value.get("object")) .ok_or_else(|| "Stripe 事件缺少 PaymentIntent".to_string())?; + if stripe_string_field(intent, "status") != Some("succeeded") { + return Err("Stripe PaymentIntent 不是成功状态".to_string()); + } let intent_id = stripe_string_field(intent, "id") .ok_or_else(|| "Stripe PaymentIntent 缺少 id".to_string())? .to_string(); @@ -187,43 +199,51 @@ async fn build_stripe_callback_input( let Some(order_no) = order_no else { return Err("Stripe PaymentIntent 缺少 metadata.order_no".to_string()); }; - let currency = stripe_string_field(intent, "currency") - .unwrap_or("usd") - .to_ascii_uppercase(); + let currency = crate::handlers::shared::normalize_payment_currency( + stripe_string_field(intent, "currency") + .ok_or_else(|| "Stripe PaymentIntent 缺少 currency".to_string())?, + "Stripe PaymentIntent currency", + ) + .map_err(|_| "Stripe PaymentIntent 币种无效".to_string())?; let amount_minor = intent .get("amount_received") .or_else(|| intent.get("amount")) .and_then(Value::as_i64) .filter(|value| *value > 0) .ok_or_else(|| "Stripe PaymentIntent 金额无效".to_string())?; - let pay_amount = stripe_amount_to_major(amount_minor, ¤cy); + let pay_amount = crate::handlers::shared::stripe_amount_to_major(amount_minor, ¤cy); - let record = state - .find_payment_gateway_config("stripe") - .await - .map_err(|err| format!("Stripe 配置读取失败: {err:?}"))? - .ok_or_else(|| "Stripe 未配置".to_string())?; - let exchange_rate = record.usd_exchange_rate; - let amount_usd = if exchange_rate > 0.0 { - pay_amount / exchange_rate - } else { - pay_amount - }; + let order = crate::handlers::shared::find_payment_callback_order(state, &order_no).await?; + let (amount_usd, exchange_rate) = crate::handlers::shared::payment_callback_settlement_values( + order.as_ref(), + pay_amount, + None, + )?; + let payment_channel = stripe_payment_intent_channel(intent) + .ok_or_else(|| "Stripe PaymentIntent 缺少 payment_method_types".to_string())?; let payload_hash = payment_callback_payload_hash(&event)?; + let payload = stripe_callback_projection( + &event, + &intent_id, + &order_no, + pay_amount, + ¤cy, + Some(&payment_channel), + ); Ok(Some( aether_data::repository::wallet::ProcessPaymentCallbackInput { payment_method: "stripe".to_string(), payment_provider: Some("stripe".to_string()), - payment_channel: stripe_payment_intent_channel(intent), - callback_key: event_id, + payment_channel: Some(payment_channel), + callback_key: payment_callback_namespaced_key("stripe", &event_id), order_no: Some(order_no), gateway_order_id: Some(intent_id), amount_usd, pay_amount: Some(pay_amount), pay_currency: Some(currency), - exchange_rate: Some(exchange_rate), + exchange_rate, payload_hash, - payload: event, + payload, signature_valid: true, }, )) @@ -283,13 +303,7 @@ pub(super) async fn handle_stripe_webhook( }; let input = match build_stripe_callback_input(state, event).await { Ok(Some(value)) => value, - Ok(None) => { - return build_auth_json_response( - http::StatusCode::OK, - json!({ "ok": true, "ignored": true, "payment_method": "stripe" }), - None, - ) - } + Ok(None) => return payment_callback_success_response(), Err(detail) => { return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) } @@ -300,8 +314,10 @@ pub(super) async fn handle_stripe_webhook( #[cfg(test)] mod tests { - use super::{stripe_amount_to_major, stripe_signature_matches_at}; + use super::{stripe_callback_projection, stripe_signature_matches_at}; + use crate::handlers::shared::stripe_amount_to_major; use hmac::{Hmac, Mac}; + use serde_json::json; use sha2::Sha256; #[test] @@ -329,11 +345,53 @@ mod tests { !stripe_signature_matches_at("wrong", &header, body, timestamp) .expect("signature check should run") ); + assert!(!stripe_signature_matches_at( + secret, + "t=-9223372036854775808,v1=00", + body, + i64::MAX, + ) + .expect("extreme timestamp should be rejected without overflow")); } #[test] - fn stripe_amount_handles_zero_decimal_currencies() { + fn stripe_callback_amount_handles_zero_and_two_decimal_currencies() { assert_eq!(stripe_amount_to_major(1234, "usd"), 12.34); assert_eq!(stripe_amount_to_major(1234, "jpy"), 1234.0); } + + #[test] + fn stripe_callback_projection_does_not_persist_raw_customer_or_secret_fields() { + let event = json!({ + "id": "evt_1", + "account": "acct_connected", + "data": { + "object": { + "client_secret": "pi_1_secret_replayable", + "receipt_email": "payer@example.com", + "shipping": {"address": {"line1": "private address"}} + } + } + }); + + let projection = + stripe_callback_projection(&event, "pi_1", "po_1", 12.34, "USD", Some("card")); + assert_eq!(projection["event_id"], "evt_1"); + assert_eq!(projection["gateway_order_id"], "pi_1"); + assert_eq!(projection["order_no"], "po_1"); + assert_eq!(projection["payment_channel"], "card"); + + let encoded = projection.to_string(); + for forbidden in [ + "client_secret", + "replayable", + "receipt_email", + "payer@example.com", + "shipping", + "private address", + "acct_connected", + ] { + assert!(!encoded.contains(forbidden), "persisted {forbidden}"); + } + } } diff --git a/apps/aether-gateway/src/handlers/public/support/payment/test_support.rs b/apps/aether-gateway/src/handlers/public/support/payment/test_support.rs index 9923fe275..ba0bfd4f9 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/test_support.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/test_support.rs @@ -1,18 +1,18 @@ use super::super::support_wallet::wallet_test_recharge_store; -use axum::{body::Body, http, response::Response}; +use axum::{body::Body, response::Response}; use chrono::Utc; use serde_json::json; -use super::super::{build_auth_json_response, sanitize_wallet_gateway_response}; +use super::super::sanitize_wallet_gateway_response; use super::payment_shared::{ - payment_callback_mark_failed_response, NormalizedPaymentCallbackRequest, + payment_callback_mark_failed_response, payment_callback_success_response, + NormalizedPaymentCallbackRequest, }; use super::GatewayPublicRequestContext; #[derive(Debug, Clone)] struct PaymentTestCallbackRecord { callback_key: String, - payment_order_id: Option, status: String, } @@ -24,45 +24,25 @@ fn payment_test_callback_store() -> &'static std::sync::Mutex Response { let mut callback_store = payment_test_callback_store() .lock() .expect("payment test callback store should lock"); - if let Some(existing) = callback_store + if callback_store .iter() - .find(|entry| entry.callback_key == payload.callback_key && entry.status == "processed") + .any(|entry| entry.callback_key == payload.callback_key && entry.status == "processed") { - return build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": true, - "credited": false, - "order_id": existing.payment_order_id, - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ); + return payment_callback_success_response(); } - let duplicate = callback_store - .iter() - .any(|entry| entry.callback_key == payload.callback_key); if !signature_valid { callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: None, status: "failed".to_string(), }); - return payment_callback_mark_failed_response( - duplicate, - "invalid callback signature", - payment_method, - &request_context.request_path, - ); + return payment_callback_mark_failed_response(); } let mut recharge_store = wallet_test_recharge_store() @@ -75,15 +55,9 @@ pub(super) async fn handle_payment_callback_with_test_store( let Some(order) = order else { callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: None, status: "failed".to_string(), }); - return payment_callback_mark_failed_response( - duplicate, - "payment order not found", - payment_method, - &request_context.request_path, - ); + return payment_callback_mark_failed_response(); }; let order_payment_method = order.payload["payment_method"] .as_str() @@ -92,30 +66,18 @@ pub(super) async fn handle_payment_callback_with_test_store( if !order_payment_method.eq_ignore_ascii_case(payment_method) { callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: order.payload["id"].as_str().map(ToOwned::to_owned), status: "failed".to_string(), }); - return payment_callback_mark_failed_response( - duplicate, - "payment method mismatch", - payment_method, - &request_context.request_path, - ); + return payment_callback_mark_failed_response(); } let order_amount = order.payload["amount_usd"].as_f64().unwrap_or_default(); if (payload.amount_usd - order_amount).abs() > f64::EPSILON { callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: order.payload["id"].as_str().map(ToOwned::to_owned), status: "failed".to_string(), }); - return payment_callback_mark_failed_response( - duplicate, - "callback amount mismatch", - payment_method, - &request_context.request_path, - ); + return payment_callback_mark_failed_response(); } let current_status = order.payload["status"] @@ -123,27 +85,11 @@ pub(super) async fn handle_payment_callback_with_test_store( .unwrap_or_default() .to_string(); if current_status == "credited" { - let order_id = order.payload["id"].as_str().map(ToOwned::to_owned); callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: order_id.clone(), status: "processed".to_string(), }); - return build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": duplicate, - "credited": false, - "order_id": order_id, - "order_no": order.payload["order_no"], - "status": "credited", - "wallet_id": order.payload["wallet_id"], - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ); + return payment_callback_success_response(); } let now = Utc::now().to_rfc3339(); @@ -168,27 +114,9 @@ pub(super) async fn handle_payment_callback_with_test_store( order.payload["refundable_amount_usd"] = json!(order_amount); order.payload["paid_at"] = json!(now.clone()); order.payload["credited_at"] = json!(now); - let order_id = order.payload["id"].as_str().map(ToOwned::to_owned); - let order_no = order.payload["order_no"].as_str().map(ToOwned::to_owned); - let wallet_id = order.payload["wallet_id"].as_str().map(ToOwned::to_owned); callback_store.push(PaymentTestCallbackRecord { callback_key: payload.callback_key.clone(), - payment_order_id: order_id.clone(), status: "processed".to_string(), }); - build_auth_json_response( - http::StatusCode::OK, - json!({ - "ok": true, - "duplicate": duplicate, - "credited": true, - "order_id": order_id, - "order_no": order_no, - "status": "credited", - "wallet_id": wallet_id, - "payment_method": payment_method, - "request_path": request_context.request_path, - }), - None, - ) + payment_callback_success_response() } diff --git a/apps/aether-gateway/src/handlers/public/support/payment/wxpay.rs b/apps/aether-gateway/src/handlers/public/support/payment/wxpay.rs index b780e2433..a2e737823 100644 --- a/apps/aether-gateway/src/handlers/public/support/payment/wxpay.rs +++ b/apps/aether-gateway/src/handlers/public/support/payment/wxpay.rs @@ -1,7 +1,8 @@ use axum::{body::Body, http, response::Response}; use super::{ - process_payment_callback_input_with_wallet_repository, AppState, GatewayPublicRequestContext, + process_payment_callback_input_with_wallet_repository, + reconcile_payment_callback_referral_rewards, AppState, GatewayPublicRequestContext, }; use serde_json::json; use tracing::warn; @@ -37,24 +38,26 @@ pub(super) async fn handle_wxpay_notify( .await { Ok(value) => value, - Err(detail) => { - warn!(error = %detail, "wxpay notify verification failed"); - return wxpay_json(http::StatusCode::BAD_REQUEST, "FAIL", detail); + Err(_) => { + warn!( + error_category = "callback_verification_failed", + "wxpay notify verification failed" + ); + return wxpay_json(http::StatusCode::BAD_REQUEST, "FAIL", "支付通知验证失败"); } }; - match process_payment_callback_input_with_wallet_repository(state, input).await { - Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { - order, - order_id, - .. - }) => { - if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await { - warn!( - error = ?err, - order_id = %order_id, - "failed to apply referral rewards for wxpay callback" - ); - } + let outcome = process_payment_callback_input_with_wallet_repository(state, input).await; + if let Ok(outcome) = &outcome { + if !reconcile_payment_callback_referral_rewards(state, outcome, "wxpay").await { + return wxpay_json( + http::StatusCode::INTERNAL_SERVER_ERROR, + "FAIL", + "支付通知关联处理失败", + ); + } + } + match outcome { + Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied { .. }) => { wxpay_json(http::StatusCode::OK, "SUCCESS", "成功") } Ok( @@ -65,13 +68,26 @@ pub(super) async fn handle_wxpay_notify( .. }, ) => wxpay_json(http::StatusCode::OK, "SUCCESS", "成功"), - Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { - error, - .. - }) => { - warn!(error = %error, path = %request_context.request_path, "wxpay notify processing failed"); - wxpay_json(http::StatusCode::INTERNAL_SERVER_ERROR, "FAIL", error) + Ok(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Failed { .. }) => { + warn!( + error_category = "callback_rejected", + path = %request_context.request_path, + "wxpay notify processing failed" + ); + wxpay_json( + http::StatusCode::INTERNAL_SERVER_ERROR, + "FAIL", + "支付通知处理失败", + ) + } + Err(response) => { + let status = response.status(); + warn!( + status = %status, + path = %request_context.request_path, + "wxpay notify storage processing failed" + ); + wxpay_json(status, "FAIL", "支付通知处理失败") } - Err(response) => response, } } diff --git a/apps/aether-gateway/src/handlers/public/support/test_connection.rs b/apps/aether-gateway/src/handlers/public/support/test_connection.rs index b7abe5bfd..9103caa09 100644 --- a/apps/aether-gateway/src/handlers/public/support/test_connection.rs +++ b/apps/aether-gateway/src/handlers/public/support/test_connection.rs @@ -1,6 +1,9 @@ -use axum::{body::Body, response::Response}; +use axum::{body::Body, http, response::Response}; -pub(super) use super::{query_param_value, AppState, GatewayPublicRequestContext}; +pub(super) use super::{ + build_auth_error_response, query_param_value, resolve_authenticated_local_user, AppState, + GatewayPublicRequestContext, +}; use crate::handlers::shared::provider_catalog_key_supports_format; #[path = "test_connection/route.rs"] @@ -11,7 +14,22 @@ mod test_connection_shared; pub(super) async fn maybe_build_local_test_connection_response( state: &AppState, request_context: &GatewayPublicRequestContext, + headers: &http::HeaderMap, ) -> Option> { + if request_context.request_path != "/v1/test-connection" { + return None; + } + let auth = match resolve_authenticated_local_user(state, request_context, headers).await { + Ok(value) => value, + Err(response) => return Some(response), + }; + if !crate::roles::can_write_admin_console(&auth.user.role) { + return Some(build_auth_error_response( + http::StatusCode::FORBIDDEN, + "仅管理员可以测试供应商连接", + false, + )); + } test_connection_route::maybe_build_local_test_connection_route_response(state, request_context) .await } diff --git a/apps/aether-gateway/src/handlers/public/support/test_connection/route.rs b/apps/aether-gateway/src/handlers/public/support/test_connection/route.rs index b879f62b0..6d2772fe5 100644 --- a/apps/aether-gateway/src/handlers/public/support/test_connection/route.rs +++ b/apps/aether-gateway/src/handlers/public/support/test_connection/route.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::net::{IpAddr, SocketAddr}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use axum::{ @@ -15,6 +16,102 @@ use super::{ }; const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000; +const MAX_TEST_CONNECTION_RESPONSE_BYTES: usize = 256 * 1024; + +#[cfg(test)] +fn build_test_connection_client() -> Result { + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(Duration::from_secs(10)) + .http2_adaptive_window(true) + .build() +} + +#[derive(Debug)] +struct ResolvedTestConnectionTarget { + url: reqwest::Url, + host: String, + addresses: Vec, +} + +/// Resolve the provider endpoint once and pin reqwest to that answer. The +/// test-connection route is reachable through the public front door, so it +/// must not perform an unbounded DNS lookup on every connect (which would +/// permit DNS rebinding into private/reserved networks). +async fn resolve_test_connection_target( + raw_url: &str, + allow_private_targets: bool, +) -> Result { + let url = reqwest::Url::parse(raw_url).map_err(|_| "provider endpoint URL is invalid")?; + let literal_loopback = aether_http::url_has_literal_loopback_host(&url); + if !matches!(url.scheme(), "http" | "https") + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + || url.fragment().is_some() + { + return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment"); + } + if url.scheme() == "http" && !(allow_private_targets && literal_loopback) { + return Err("provider endpoint must use HTTPS"); + } + let host = url + .host_str() + .ok_or("provider endpoint is missing a host")? + .to_string(); + let literal_ip = host.parse::().ok(); + let port = url + .port_or_known_default() + .ok_or("provider endpoint is missing a port")?; + let addresses = if let Some(ip) = literal_ip { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits( + host.as_str(), + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| "provider endpoint DNS resolution failed")? + }; + if addresses.is_empty() { + return Err("provider endpoint DNS resolution returned no addresses"); + } + let has_private_answer = addresses + .iter() + .any(|address| aether_http::is_private_or_reserved_ip(address.ip())); + // `allow_private_targets` is only enabled for in-process test fixtures. + // Keep that escape hatch narrowly scoped to literal loopback URLs whose + // every DNS answer is loopback; otherwise a test-only build (or an + // accidentally reused helper) could turn this public route into a + // private-network HTTP client. + let test_loopback_target = allow_private_targets + && literal_loopback + && addresses.iter().all(|address| address.ip().is_loopback()); + if has_private_answer && !test_loopback_target { + return Err("provider endpoint resolves to a private or reserved address"); + } + Ok(ResolvedTestConnectionTarget { + url, + host, + addresses, + }) +} + +fn build_pinned_test_connection_client( + target: &ResolvedTestConnectionTarget, +) -> Result { + let mut builder = reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(Duration::from_secs(10)) + .http2_adaptive_window(true); + if target.host.parse::().is_err() { + builder = builder.resolve_to_addrs(&target.host, &target.addresses); + } + builder.build() +} pub(super) async fn maybe_build_local_test_connection_route_response( state: &AppState, @@ -290,7 +387,42 @@ pub(super) async fn maybe_build_local_test_connection_route_response( ); } - let mut upstream_request = state.client.post(&upstream_url); + // Resolve and pin the endpoint before constructing the request. This + // keeps the public health-check route subject to the same DNS/SSRF + // boundary as the main execution transport. Unit-test fixtures may use + // loopback listeners; production requests never opt into private targets. + let target = match resolve_test_connection_target(&upstream_url, cfg!(test)).await { + Ok(target) => target, + Err(reason) => { + tracing::warn!( + event_name = "provider_test_connection_target_rejected", + provider_id = %provider.id, + endpoint_id = %endpoint.id, + reason, + "provider connection test target was rejected" + ); + return Some( + ( + http::StatusCode::SERVICE_UNAVAILABLE, + Json(json!({"detail": "Provider connection test unavailable"})), + ) + .into_response(), + ); + } + }; + let test_client = match build_pinned_test_connection_client(&target) { + Ok(client) => client, + Err(_) => { + return Some( + ( + http::StatusCode::SERVICE_UNAVAILABLE, + Json(json!({"detail": "Provider connection test unavailable"})), + ) + .into_response(), + ); + } + }; + let mut upstream_request = test_client.post(target.url); for (name, value) in &provider_request_headers { upstream_request = upstream_request.header(name, value); } @@ -301,31 +433,41 @@ pub(super) async fn maybe_build_local_test_connection_route_response( let response = match upstream_request.json(&provider_request_body).send().await { Ok(response) => response, - Err(error) => { + Err(_) => { + tracing::warn!( + event_name = "provider_test_connection_request_failed", + provider_id = %provider.id, + endpoint_id = %endpoint.id, + "provider connection test request failed" + ); return Some( ( http::StatusCode::SERVICE_UNAVAILABLE, - Json(json!({"detail": error.to_string()})), + Json(json!({"detail": "Provider connection test failed"})), ) .into_response(), - ) + ); } }; let status = response.status(); - let response_json = response.json::().await.ok(); + let response_json = + aether_http::read_response_bytes_with_limit(response, MAX_TEST_CONNECTION_RESPONSE_BYTES) + .await + .ok() + .and_then(|body| serde_json::from_slice::(&body).ok()); if !status.is_success() { - let detail = response_json - .as_ref() - .and_then(|value| value.get("error")) - .and_then(|value| value.get("message").or(Some(value))) - .and_then(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .unwrap_or_else(|| format!("upstream returned HTTP {status}")); + tracing::warn!( + event_name = "provider_test_connection_upstream_rejected", + provider_id = %provider.id, + endpoint_id = %endpoint.id, + upstream_status = %status, + "provider connection test upstream returned an error" + ); return Some( ( http::StatusCode::SERVICE_UNAVAILABLE, - Json(json!({ "detail": detail })), + Json(json!({ "detail": "Provider connection test failed" })), ) .into_response(), ); @@ -353,3 +495,138 @@ pub(super) async fn maybe_build_local_test_connection_route_response( .into_response(), ) } + +#[cfg(test)] +mod tests { + use super::{build_test_connection_client, resolve_test_connection_target}; + use axum::{ + body::Body, + http::{header, Request, StatusCode}, + response::IntoResponse, + routing::post, + Router, + }; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + + #[tokio::test] + async fn test_connection_client_never_forwards_credentials_across_redirects() { + let redirected_hits = Arc::new(AtomicUsize::new(0)); + let redirected_hits_for_route = Arc::clone(&redirected_hits); + let redirected_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect target listener"); + let redirected_addr = redirected_listener + .local_addr() + .expect("redirect target addr"); + let redirected_app = Router::new().route( + "/capture", + post(move |request: Request| { + let hits = Arc::clone(&redirected_hits_for_route); + async move { + hits.fetch_add(1, Ordering::SeqCst); + assert!(request.headers().get(header::AUTHORIZATION).is_none()); + StatusCode::OK + } + }), + ); + let redirected_server = tokio::spawn(async move { + axum::serve(redirected_listener, redirected_app) + .await + .expect("redirect target server"); + }); + + let source_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect source listener"); + let source_addr = source_listener.local_addr().expect("redirect source addr"); + let location = format!("http://{redirected_addr}/capture"); + let source_app = Router::new().route( + "/redirect", + post(move || { + let location = location.clone(); + async move { (StatusCode::FOUND, [(header::LOCATION, location)]).into_response() } + }), + ); + let source_server = tokio::spawn(async move { + axum::serve(source_listener, source_app) + .await + .expect("redirect source server"); + }); + + let response = build_test_connection_client() + .expect("client") + .post(format!("http://{source_addr}/redirect")) + .header(header::AUTHORIZATION, "Bearer stored-provider-secret") + .send() + .await + .expect("redirect response"); + assert_eq!(response.status(), StatusCode::FOUND); + assert_eq!(redirected_hits.load(Ordering::SeqCst), 0); + + source_server.abort(); + redirected_server.abort(); + } + + #[tokio::test] + async fn test_connection_target_rejects_private_addresses_in_production_mode() { + for raw_url in [ + "http://127.0.0.1:8080/v1/chat/completions", + "https://10.0.0.1/v1/chat/completions", + "https://[::1]/v1/chat/completions", + "https://localhost/v1/chat/completions", + "http://8.8.8.8/v1/chat/completions", + ] { + assert!( + resolve_test_connection_target(raw_url, false) + .await + .is_err(), + "private provider target should be rejected: {raw_url}" + ); + } + } + + #[tokio::test] + async fn test_connection_target_allows_loopback_only_for_test_fixtures() { + let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true) + .await + .expect("test fixture target should resolve"); + assert_eq!(target.host, "127.0.0.1"); + assert_eq!(target.addresses.len(), 1); + assert!( + resolve_test_connection_target("http://8.8.8.8/v1/chat", true) + .await + .is_err(), + "test mode must not make cleartext public endpoints acceptable" + ); + assert!( + resolve_test_connection_target("https://10.0.0.1/v1/chat", true) + .await + .is_err(), + "test mode must not make private non-loopback endpoints acceptable" + ); + assert!( + resolve_test_connection_target("http://localhost:8080/v1/chat", true) + .await + .is_ok(), + "literal localhost should remain available for local fixtures" + ); + } + + #[tokio::test] + async fn test_connection_target_rejects_url_credentials_and_fragments() { + for raw_url in [ + "https://user:pass@example.com/v1/chat", + "https://example.com/v1/chat#fragment", + ] { + assert!( + resolve_test_connection_target(raw_url, false) + .await + .is_err(), + "unsafe provider target should be rejected: {raw_url}" + ); + } + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/user_me.rs b/apps/aether-gateway/src/handlers/public/support/user_me.rs index c9672bbdc..54d196599 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me.rs @@ -1,9 +1,12 @@ use super::{ - auth_password_policy_level, base_url_from_request, build_auth_error_response, - build_auth_wallet_summary_payload, decrypt_catalog_secret_with_fallbacks, - encrypt_catalog_secret_with_fallbacks, handle_auth_me, - handle_users_me_api_key_install_session_create, query_param_optional_bool, query_param_value, - resolve_authenticated_local_user, sanitize_public_model_config_for_user, unix_secs_to_rfc3339, + auth_email_is_verified, auth_password_policy_level, base_url_from_request, + build_auth_error_response, build_auth_json_response, build_auth_refresh_cookie_clear_header, + build_auth_wallet_summary_payload, consume_auth_email_registration_proof, + decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, handle_auth_me, + handle_users_me_api_key_install_session_create, handle_users_me_vscodex_request, + query_param_optional_bool, query_param_value, resolve_authenticated_local_user, + sanitize_public_model_capabilities, sanitize_public_model_config_for_user, + sanitize_public_tiered_pricing, unix_secs_to_rfc3339, users_me_api_key_install_sessions_path_matches, validate_auth_register_password, AppState, AuthenticatedLocalUserContext, GatewayPublicRequestContext, PUBLIC_CAPABILITY_DEFINITIONS, }; diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs b/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs index c296e5be9..baec48c2f 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs @@ -10,22 +10,30 @@ use serde::Deserialize; use serde_json::json; use crate::handlers::shared::{ - api_key_placeholder_display, deserialize_optional_json_patch, - deserialize_optional_string_list_patch, generate_gateway_api_key_plaintext, - masked_gateway_api_key_display, normalize_feature_settings, normalize_ip_rules, - normalize_optional_api_key_concurrent_limit, + api_key_placeholder_display, decrypt_or_migrate_auth_api_key_secret, + deserialize_optional_json_patch, deserialize_optional_string_list_patch, + generate_gateway_api_key_plaintext, masked_gateway_api_key_display, normalize_feature_settings, + normalize_ip_rules, normalize_optional_api_key_concurrent_limit, open_auth_api_key_secret, + seal_auth_api_key_secret, }; use super::{ - build_auth_error_response, decrypt_catalog_secret_with_fallbacks, - encrypt_catalog_secret_with_fallbacks, format_users_me_optional_unix_secs_iso8601, - known_capability_names, normalize_user_model_capability_settings_input, - query_param_optional_bool, resolve_authenticated_local_user, - user_configurable_capability_names, AppState, GatewayPublicRequestContext, + build_auth_error_response, format_users_me_optional_unix_secs_iso8601, known_capability_names, + normalize_user_model_capability_settings_input, query_param_optional_bool, + resolve_authenticated_local_user, user_configurable_capability_names, AppState, + GatewayPublicRequestContext, }; const USERS_ME_API_KEY_WRITE_UNAVAILABLE_DETAIL: &str = "用户 API 密钥写入暂不可用"; +fn users_me_api_key_secret_response(mut response: Response) -> Response { + response.headers_mut().insert( + http::header::CACHE_CONTROL, + http::HeaderValue::from_static("no-store"), + ); + response +} + #[derive(Debug, Deserialize)] struct UsersMeCreateApiKeyRequest { name: String, @@ -132,15 +140,14 @@ pub(super) fn users_me_api_key_capabilities_path_matches(request_path: &str) -> users_me_api_key_nested_id_from_path(request_path, "capabilities").is_some() } -fn users_me_masked_api_key_display(state: &AppState, ciphertext: Option<&str>) -> String { - let Some(ciphertext) = ciphertext.map(str::trim).filter(|value| !value.is_empty()) else { +fn users_me_masked_api_key_display( + state: &AppState, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, +) -> String { + let Ok(projection) = open_auth_api_key_secret(state, record) else { return api_key_placeholder_display(); }; - let Some(full_key) = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - else { - return api_key_placeholder_display(); - }; - masked_gateway_api_key_display(Some(full_key.as_str())) + masked_gateway_api_key_display(Some(projection.plaintext.as_str())) } fn build_users_me_api_key_writer_unavailable_response() -> Response { @@ -151,6 +158,14 @@ fn build_users_me_api_key_writer_unavailable_response() -> Response { ) } +fn build_users_me_api_key_mutation_conflict_response() -> Response { + build_auth_error_response( + http::StatusCode::CONFLICT, + "API密钥已锁定或状态已发生变化,请刷新后重试", + false, + ) +} + fn build_users_me_api_key_list_payload( state: &AppState, record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, @@ -159,7 +174,7 @@ fn build_users_me_api_key_list_payload( json!({ "id": record.api_key_id, "name": record.name, - "key_display": users_me_masked_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": users_me_masked_api_key_display(state, record), "is_active": record.is_active, "is_locked": is_locked, "last_used_at": format_users_me_optional_unix_secs_iso8601(record.last_used_at_unix_secs), @@ -183,7 +198,7 @@ fn build_users_me_api_key_detail_payload( json!({ "id": record.api_key_id, "name": record.name, - "key_display": users_me_masked_api_key_display(state, record.key_encrypted.as_deref()), + "key_display": users_me_masked_api_key_display(state, record), "is_active": record.is_active, "is_locked": is_locked, "allowed_providers": record.allowed_providers, @@ -413,16 +428,17 @@ pub(super) async fn handle_users_me_api_key_detail_get( false, ); } - let Some(full_key) = - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - else { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - "解密密钥失败", - false, - ); + let full_key = match decrypt_or_migrate_auth_api_key_secret(state, &record).await { + Ok(value) => value, + Err(_) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + "解密或校验密钥失败", + false, + ); + } }; - return Json(json!({ "key": full_key })).into_response(); + return users_me_api_key_secret_response(Json(json!({ "key": full_key })).into_response()); } let snapshot_ids = vec![api_key_id.clone()]; @@ -565,7 +581,16 @@ pub(super) async fn handle_users_me_api_key_create( }; let plaintext_key = generate_users_me_api_key_plaintext(); - let Some(key_encrypted) = encrypt_catalog_secret_with_fallbacks(state, &plaintext_key) else { + let api_key_id = uuid::Uuid::new_v4().to_string(); + let key_hash = hash_users_me_api_key(&plaintext_key); + let Ok(key_encrypted) = seal_auth_api_key_secret( + state, + &auth.user.id, + &api_key_id, + &key_hash, + false, + &plaintext_key, + ) else { return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, "API密钥加密失败", @@ -574,8 +599,8 @@ pub(super) async fn handle_users_me_api_key_create( }; let record = aether_data::repository::auth::CreateUserApiKeyRecord { user_id: auth.user.id.clone(), - api_key_id: uuid::Uuid::new_v4().to_string(), - key_hash: hash_users_me_api_key(&plaintext_key), + api_key_id, + key_hash, key_encrypted: Some(key_encrypted), name: Some(name.clone()), allowed_providers: None, @@ -585,6 +610,7 @@ pub(super) async fn handle_users_me_api_key_create( rate_limit, concurrent_limit, force_capabilities: None, + feature_settings, is_active: true, expires_at_unix_secs: None, auto_delete_on_expiry: false, @@ -604,47 +630,27 @@ pub(super) async fn handle_users_me_api_key_create( }) else { return build_users_me_api_key_writer_unavailable_response(); }; - let created = if feature_settings.is_some() { - match state - .set_user_api_key_feature_settings( - &auth.user.id, - &created.api_key_id, - feature_settings.clone(), - ) - .await - { - Ok(Some(record)) => record, - Ok(None) => return build_users_me_api_key_writer_unavailable_response(), - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user api key feature settings update failed: {err:?}"), - false, - ) - } - } - } else { - created - }; - Json(json!({ - "id": created.api_key_id, - "name": created.name, - "key": plaintext_key, - "key_display": users_me_masked_api_key_display(state, created.key_encrypted.as_deref()), - "is_active": created.is_active, - "is_locked": false, - "rate_limit": created.rate_limit, - "concurrent_limit": created.concurrent_limit, - "ip_rules": created.ip_rules, - "feature_settings": created.feature_settings, - "last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs), - "created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs), - "total_requests": created.total_requests, - "total_cost_usd": created.total_cost_usd, - "message": "API密钥创建成功", - })) - .into_response() + users_me_api_key_secret_response( + Json(json!({ + "id": created.api_key_id, + "name": created.name, + "key": plaintext_key, + "key_display": users_me_masked_api_key_display(state, &created), + "is_active": created.is_active, + "is_locked": false, + "rate_limit": created.rate_limit, + "concurrent_limit": created.concurrent_limit, + "ip_rules": created.ip_rules, + "feature_settings": created.feature_settings, + "last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs), + "created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs), + "total_requests": created.total_requests, + "total_cost_usd": created.total_cost_usd, + "message": "API密钥创建成功", + })) + .into_response(), + ) } pub(super) async fn handle_users_me_api_key_update( @@ -727,16 +733,27 @@ pub(super) async fn handle_users_me_api_key_update( }, None => None, }; + let name_present = name.is_some(); + let rate_limit_present = rate_limit.is_some(); + let concurrent_limit_present = concurrent_limit.is_some(); let Some(updated) = (match state - .update_user_api_key_basic(aether_data::repository::auth::UpdateUserApiKeyBasicRecord { - user_id: auth.user.id.clone(), - api_key_id: snapshot.api_key_id.clone(), - name, - rate_limit, - concurrent_limit, - ip_rules, - }) + .update_user_api_key_basic_if_unlocked( + aether_data::repository::auth::UpdateUserApiKeyBasicRecord { + user_id: auth.user.id.clone(), + api_key_id: snapshot.api_key_id.clone(), + key_encrypted: None, + key_encrypted_present: false, + name, + name_present, + rate_limit, + rate_limit_present, + concurrent_limit, + concurrent_limit_present, + ip_rules, + feature_settings, + }, + ) .await { Ok(value) => value, @@ -748,29 +765,7 @@ pub(super) async fn handle_users_me_api_key_update( ) } }) else { - return build_users_me_api_key_writer_unavailable_response(); - }; - let updated = if let Some(feature_settings) = feature_settings { - match state - .set_user_api_key_feature_settings( - &auth.user.id, - &snapshot.api_key_id, - feature_settings, - ) - .await - { - Ok(Some(record)) => record, - Ok(None) => return build_users_me_api_key_writer_unavailable_response(), - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user api key feature settings update failed: {err:?}"), - false, - ) - } - } - } else { - updated + return build_users_me_api_key_mutation_conflict_response(); }; let mut payload = @@ -820,7 +815,7 @@ pub(super) async fn handle_users_me_api_key_patch( }; let Some(updated) = (match state - .set_user_api_key_active(&auth.user.id, &snapshot.api_key_id, desired_is_active) + .set_user_api_key_active_if_unlocked(&auth.user.id, &snapshot.api_key_id, desired_is_active) .await { Ok(value) => value, @@ -832,7 +827,7 @@ pub(super) async fn handle_users_me_api_key_patch( ) } }) else { - return build_users_me_api_key_writer_unavailable_response(); + return build_users_me_api_key_mutation_conflict_response(); }; Json(json!({ @@ -869,11 +864,11 @@ pub(super) async fn handle_users_me_api_key_delete( } match state - .delete_user_api_key(&auth.user.id, &snapshot.api_key_id) + .delete_user_api_key_if_unlocked(&auth.user.id, &snapshot.api_key_id) .await { Ok(true) => Json(json!({ "message": "API密钥已删除" })).into_response(), - Ok(false) => build_auth_error_response(http::StatusCode::NOT_FOUND, "API密钥不存在", false), + Ok(false) => build_users_me_api_key_mutation_conflict_response(), Err(err) => build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, format!("user api key delete failed: {err:?}"), @@ -995,7 +990,11 @@ pub(super) async fn handle_users_me_api_key_providers_put( }; let Some(updated) = (match state - .set_user_api_key_allowed_providers(&auth.user.id, &snapshot.api_key_id, allowed_providers) + .set_user_api_key_allowed_providers_if_unlocked( + &auth.user.id, + &snapshot.api_key_id, + allowed_providers, + ) .await { Ok(value) => value, @@ -1007,7 +1006,7 @@ pub(super) async fn handle_users_me_api_key_providers_put( ) } }) else { - return build_users_me_api_key_writer_unavailable_response(); + return build_users_me_api_key_mutation_conflict_response(); }; Json(json!({ @@ -1065,7 +1064,7 @@ pub(super) async fn handle_users_me_api_key_capabilities_put( }; let Some(updated) = (match state - .set_user_api_key_force_capabilities( + .set_user_api_key_force_capabilities_if_unlocked( &auth.user.id, &snapshot.api_key_id, force_capabilities, @@ -1081,7 +1080,7 @@ pub(super) async fn handle_users_me_api_key_capabilities_put( ) } }) else { - return build_users_me_api_key_writer_unavailable_response(); + return build_users_me_api_key_mutation_conflict_response(); }; Json(json!({ @@ -1093,9 +1092,24 @@ pub(super) async fn handle_users_me_api_key_capabilities_put( #[cfg(test)] mod tests { - use super::{normalize_users_me_ip_rules, UsersMeUpdateApiKeyRequest}; + use axum::{response::IntoResponse, Json}; + + use super::{ + normalize_users_me_ip_rules, users_me_api_key_secret_response, UsersMeUpdateApiKeyRequest, + }; use serde_json::json; + #[test] + fn plaintext_api_key_responses_are_never_cacheable() { + let response = + users_me_api_key_secret_response(Json(json!({ "key": "sk-secret" })).into_response()); + + assert_eq!( + response.headers().get(axum::http::header::CACHE_CONTROL), + Some(&axum::http::HeaderValue::from_static("no-store")) + ); + } + #[test] fn normalize_ip_rules_trims_ip_and_cidr_values() { let values = normalize_users_me_ip_rules(Some(vec![ diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_catalog.rs b/apps/aether-gateway/src/handlers/public/support/user_me_catalog.rs index 27ded4022..1e287528a 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_catalog.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_catalog.rs @@ -14,7 +14,8 @@ use serde_json::json; use super::{ build_admin_endpoint_health_status_payload, build_auth_error_response, query_param_value, - resolve_authenticated_local_user, sanitize_public_model_config_for_user, AppState, + resolve_authenticated_local_user, sanitize_public_model_capabilities, + sanitize_public_model_config_for_user, sanitize_public_tiered_pricing, AppState, GatewayPublicRequestContext, USERS_ME_AVAILABLE_MODELS_FETCH_LIMIT, }; @@ -26,10 +27,18 @@ fn build_users_me_available_model_payload( model: StoredPublicGlobalModel, hide_mapping_config: bool, ) -> serde_json::Value { - let config = if hide_mapping_config { - sanitize_public_model_config_for_user(model.config) + let (default_tiered_pricing, supported_capabilities, config) = if hide_mapping_config { + ( + sanitize_public_tiered_pricing(model.default_tiered_pricing), + sanitize_public_model_capabilities(model.supported_capabilities), + sanitize_public_model_config_for_user(model.config), + ) } else { - model.config + ( + model.default_tiered_pricing, + model.supported_capabilities, + model.config, + ) }; json!({ "id": model.id, @@ -37,8 +46,8 @@ fn build_users_me_available_model_payload( "display_name": model.display_name, "is_active": model.is_active, "default_price_per_request": model.default_price_per_request, - "default_tiered_pricing": model.default_tiered_pricing, - "supported_capabilities": model.supported_capabilities, + "default_tiered_pricing": default_tiered_pricing, + "supported_capabilities": supported_capabilities, "config": config, "usage_count": model.usage_count, }) diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_management_tokens.rs b/apps/aether-gateway/src/handlers/public/support/user_me_management_tokens.rs index 78a9db1bf..3f35d0b97 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_management_tokens.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_management_tokens.rs @@ -20,6 +20,7 @@ use super::{ GatewayPublicRequestContext, }; use crate::control::normalize_assignable_management_token_permissions; +use crate::handlers::public::support::mark_sensitive_response_no_store; use crate::handlers::shared::{generate_gateway_secret_plaintext, parse_json_ip_rules}; use crate::LocalMutationOutcome; @@ -34,6 +35,10 @@ const USERS_ME_MANAGEMENT_TOKEN_READ_UNAVAILABLE_DETAIL: &str = const USERS_ME_MANAGEMENT_TOKEN_WRITE_UNAVAILABLE_DETAIL: &str = "用户 Management Token 写入暂不可用"; +fn users_me_management_token_secret_response(response: Response) -> Response { + mark_sensitive_response_no_store(response) +} + #[derive(Debug, Clone)] struct UsersMeManagementTokenCreateInput { name: String, @@ -141,10 +146,10 @@ fn hash_users_me_management_token(value: &str) -> String { fn users_me_management_token_prefix(value: &str) -> Option { (!value.is_empty()).then(|| { - value[..value - .len() - .min(USERS_ME_MANAGEMENT_TOKEN_DISPLAY_PREFIX_LEN)] - .to_string() + value + .chars() + .take(USERS_ME_MANAGEMENT_TOKEN_DISPLAY_PREFIX_LEN) + .collect() }) } @@ -164,6 +169,18 @@ fn build_users_me_management_token_writer_unavailable_response() -> Response bool { + crate::roles::can_write_admin_console(role) +} + +fn build_users_me_management_token_write_forbidden_response(action: &str) -> Response { + build_auth_error_response( + http::StatusCode::FORBIDDEN, + format!("仅管理员可以{action} Management Token"), + false, + ) +} + fn users_me_management_token_limit(query: Option<&str>) -> usize { query_param_value(query, "limit") .and_then(|value| value.parse::().ok()) @@ -446,12 +463,8 @@ pub(super) async fn handle_users_me_management_token_create( Ok(value) => value, Err(response) => return response, }; - if !auth.user.role.eq_ignore_ascii_case("admin") { - return build_auth_error_response( - http::StatusCode::FORBIDDEN, - "仅管理员可以创建 Management Token", - false, - ); + if !users_me_management_token_write_allowed(&auth.user.role) { + return build_users_me_management_token_write_forbidden_response("创建"); } let Some(request_body) = request_body else { return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少请求体", false); @@ -513,15 +526,17 @@ pub(super) async fn handle_users_me_management_token_create( }; match state.create_management_token(&record).await { - Ok(LocalMutationOutcome::Applied(token)) => ( - http::StatusCode::CREATED, - Json(json!({ - "message": "Management Token 创建成功", - "token": raw_token, - "data": build_management_token_payload(&token, None), - })), - ) - .into_response(), + Ok(LocalMutationOutcome::Applied(token)) => users_me_management_token_secret_response( + ( + http::StatusCode::CREATED, + Json(json!({ + "message": "Management Token 创建成功", + "token": raw_token, + "data": build_management_token_payload(&token, None), + })), + ) + .into_response(), + ), Ok(LocalMutationOutcome::Invalid(detail)) => { build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) } @@ -578,6 +593,9 @@ pub(super) async fn handle_users_me_management_token_update( Ok(value) => value, Err(response) => return response, }; + if !users_me_management_token_write_allowed(&auth.user.role) { + return build_users_me_management_token_write_forbidden_response("管理"); + } let Some(token_id) = users_me_management_token_id_from_path(&request_context.request_path) else { return build_auth_error_response( @@ -646,7 +664,10 @@ pub(super) async fn handle_users_me_management_token_update( is_active: None, }; - match state.update_management_token(&record).await { + match state + .update_management_token_for_user(&record, &auth.user.id) + .await + { Ok(LocalMutationOutcome::Applied(token)) => Json(json!({ "message": "更新成功", "data": build_management_token_payload(&token, None), @@ -684,6 +705,9 @@ pub(super) async fn handle_users_me_management_token_delete( Ok(value) => value, Err(response) => return response, }; + if !users_me_management_token_write_allowed(&auth.user.role) { + return build_users_me_management_token_write_forbidden_response("管理"); + } let Some(token_id) = users_me_management_token_id_from_path(&request_context.request_path) else { return build_auth_error_response( @@ -696,7 +720,10 @@ pub(super) async fn handle_users_me_management_token_delete( Ok(value) => value, Err(response) => return response, }; - match state.delete_management_token(&existing.token.id).await { + match state + .delete_management_token_for_user(&existing.token.id, &auth.user.id) + .await + { Ok(true) => Json(json!({ "message": "删除成功" })).into_response(), Ok(false) => build_auth_error_response( http::StatusCode::NOT_FOUND, @@ -724,6 +751,9 @@ pub(super) async fn handle_users_me_management_token_toggle( Ok(value) => value, Err(response) => return response, }; + if !users_me_management_token_write_allowed(&auth.user.role) { + return build_users_me_management_token_write_forbidden_response("管理"); + } let Some(token_id) = users_me_management_token_status_id_from_path(&request_context.request_path) else { @@ -738,7 +768,11 @@ pub(super) async fn handle_users_me_management_token_toggle( Err(response) => return response, }; match state - .set_management_token_active(&existing.token.id, !existing.token.is_active) + .set_management_token_active_for_user( + &existing.token.id, + &auth.user.id, + !existing.token.is_active, + ) .await { Ok(Some(token)) => Json(json!({ @@ -772,6 +806,9 @@ pub(super) async fn handle_users_me_management_token_regenerate( Ok(value) => value, Err(response) => return response, }; + if !users_me_management_token_write_allowed(&auth.user.role) { + return build_users_me_management_token_write_forbidden_response("管理"); + } let Some(token_id) = users_me_management_token_regenerate_id_from_path(&request_context.request_path) else { @@ -792,13 +829,18 @@ pub(super) async fn handle_users_me_management_token_regenerate( token_prefix: users_me_management_token_prefix(&raw_token), }; - match state.regenerate_management_token_secret(&mutation).await { - Ok(LocalMutationOutcome::Applied(token)) => Json(json!({ - "message": "Token 已重新生成", - "token": raw_token, - "data": build_management_token_payload(&token, None), - })) - .into_response(), + match state + .regenerate_management_token_secret_for_user(&mutation, &auth.user.id) + .await + { + Ok(LocalMutationOutcome::Applied(token)) => users_me_management_token_secret_response( + Json(json!({ + "message": "Token 已重新生成", + "token": raw_token, + "data": build_management_token_payload(&token, None), + })) + .into_response(), + ), Ok(LocalMutationOutcome::NotFound) => build_auth_error_response( http::StatusCode::NOT_FOUND, "Management Token 不存在", @@ -817,3 +859,31 @@ pub(super) async fn handle_users_me_management_token_regenerate( ), } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn management_token_writes_require_the_current_full_admin_role() { + assert!(users_me_management_token_write_allowed("admin")); + assert!(!users_me_management_token_write_allowed("audit_admin")); + assert!(!users_me_management_token_write_allowed("user")); + } + + #[test] + fn plaintext_management_token_responses_are_never_cacheable() { + let response = users_me_management_token_secret_response( + Json(json!({ "token": "ae-secret" })).into_response(), + ); + + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&http::HeaderValue::from_static("no-store")) + ); + assert_eq!( + response.headers().get(http::header::PRAGMA), + Some(&http::HeaderValue::from_static("no-cache")) + ); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_profile.rs b/apps/aether-gateway/src/handlers/public/support/user_me_profile.rs index 964298b2d..4a51d45b4 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_profile.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_profile.rs @@ -12,25 +12,28 @@ use crate::handlers::shared::{ }; use super::{ - auth_password_policy_level, base_url_from_request, build_auth_error_response, - resolve_authenticated_local_user, validate_auth_register_password, AppState, - GatewayPublicRequestContext, + auth_email_is_verified, auth_password_policy_level, base_url_from_request, + build_auth_error_response, build_auth_json_response, build_auth_refresh_cookie_clear_header, + consume_auth_email_registration_proof, resolve_authenticated_local_user, + validate_auth_register_password, AppState, GatewayPublicRequestContext, }; const USERS_ME_PROFILE_STORAGE_UNAVAILABLE_DETAIL: &str = "用户资料存储暂不可用"; const USERS_ME_CREDENTIAL_STORAGE_UNAVAILABLE_DETAIL: &str = "用户凭证存储暂不可用"; -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct UsersMeUpdateProfileRequest { #[serde(default)] email: Option, #[serde(default)] + email_verification_token: Option, + #[serde(default)] username: Option, #[serde(default, deserialize_with = "deserialize_optional_json_patch")] feature_settings: Option>, } -#[derive(Debug, Deserialize)] +#[derive(Deserialize)] struct UsersMeChangePasswordRequest { #[serde(default, alias = "current_password")] old_password: Option, @@ -45,6 +48,7 @@ pub(super) async fn handle_users_me_client_config_get( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, ) -> Response { if let Err(response) = resolve_authenticated_local_user(state, request_context, headers).await { return response; @@ -59,7 +63,7 @@ pub(super) async fn handle_users_me_client_config_get( .unwrap_or_else(|| "Aether".to_string()); Json(json!({ - "base_url": base_url_from_request(headers, request_context), + "base_url": base_url_from_request(headers, request_context, remote_addr), "site_name": site_name, })) .into_response() @@ -90,7 +94,15 @@ pub(super) async fn handle_users_me_detail_put( }; let email = normalize_users_me_optional_non_empty_string(payload.email); + let email_verification_token = payload + .email_verification_token + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); let username = normalize_users_me_optional_non_empty_string(payload.username); + let email_changed = email + .as_deref() + .is_some_and(|value| auth.user.email.as_deref() != Some(value)); let feature_settings = match payload.feature_settings { Some(value) => { let current = match state.read_user_feature_settings(&auth.user.id).await { @@ -136,6 +148,39 @@ pub(super) async fn handle_users_me_detail_put( } } + if email_changed { + let Some(verification_token) = email_verification_token else { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "修改邮箱前请先完成新邮箱验证", + false, + ); + }; + match auth_email_is_verified( + state, + email.as_deref().unwrap_or_default(), + verification_token, + ) + .await + { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "邮箱验证凭据无效或已过期,请重新验证", + false, + ); + } + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("email verification lookup failed: {err:?}"), + false, + ); + } + } + } + if let Some(username) = username.as_deref() { match state .is_other_user_auth_username_taken(username, &auth.user.id) @@ -159,39 +204,74 @@ pub(super) async fn handle_users_me_detail_put( } } + if email_changed { + let verification_token = + email_verification_token.expect("email change token checked above"); + match consume_auth_email_registration_proof( + state, + email.as_deref().unwrap_or_default(), + verification_token, + ) + .await + { + Ok(true) => {} + Ok(false) => { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "邮箱验证凭据已被使用,请重新验证", + false, + ); + } + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("email verification consume failed: {err:?}"), + false, + ); + } + } + } + match state - .update_local_auth_user_profile(&auth.user.id, email, username) + .update_local_auth_user_profile( + &auth.user.id, + email.is_some(), + email, + email_changed.then_some(true), + username, + ) .await { - Ok(Some(_)) => { - if let Some(feature_settings) = feature_settings { - match state - .update_user_feature_settings(&auth.user.id, feature_settings) - .await - { - Ok(_) => {} - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user feature settings update failed: {err:?}"), - false, - ) - } - } - } - Json(json!({ "message": "个人信息更新成功" })).into_response() + Ok(Some(_)) => {} + Ok(None) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + USERS_ME_PROFILE_STORAGE_UNAVAILABLE_DETAIL, + false, + ) + } + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("user profile update failed: {err:?}"), + false, + ) } - Ok(None) => build_auth_error_response( - http::StatusCode::SERVICE_UNAVAILABLE, - USERS_ME_PROFILE_STORAGE_UNAVAILABLE_DETAIL, - false, - ), - Err(err) => build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user profile update failed: {err:?}"), - false, - ), } + + if let Some(feature_settings) = feature_settings { + if let Err(err) = state + .update_user_feature_settings(&auth.user.id, feature_settings) + .await + { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("user feature settings update failed: {err:?}"), + false, + ); + } + } + Json(json!({ "message": "个人信息更新成功" })).into_response() } pub(super) async fn handle_users_me_password_patch( @@ -281,15 +361,21 @@ pub(super) async fn handle_users_me_password_patch( }; let updated_at = chrono::Utc::now(); match state - .update_local_auth_user_password_hash(&auth.user.id, password_hash, updated_at) + .change_local_auth_password_and_revoke_sessions( + &auth.user.id, + &auth.session_id, + current_password_hash, + password_hash, + updated_at, + ) .await { - Ok(Some(_)) => {} - Ok(None) => { + Ok(true) => {} + Ok(false) => { return build_auth_error_response( http::StatusCode::SERVICE_UNAVAILABLE, USERS_ME_CREDENTIAL_STORAGE_UNAVAILABLE_DETAIL, - false, + true, ) } Err(err) => { @@ -301,36 +387,14 @@ pub(super) async fn handle_users_me_password_patch( } } - let sessions = match state.list_user_sessions(&auth.user.id).await { - Ok(value) => value, - Err(err) => { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user session lookup failed: {err:?}"), - false, - ) - } - }; - for session in sessions { - if session.id == auth.session_id { - continue; - } - if let Err(err) = state - .revoke_user_session(&auth.user.id, &session.id, updated_at, "password_changed") - .await - { - return build_auth_error_response( - http::StatusCode::INTERNAL_SERVER_ERROR, - format!("user session revoke failed: {err:?}"), - false, - ); - } - } - let action = if current_password_hash.is_some() { "修改" } else { "设置" }; - Json(json!({ "message": format!("密码{action}成功") })).into_response() + build_auth_json_response( + http::StatusCode::OK, + json!({ "message": format!("密码{action}成功") }), + Some(build_auth_refresh_cookie_clear_header()), + ) } diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs b/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs index 377b12a1a..42b930491 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs @@ -1,4 +1,5 @@ use crate::handlers::public::support::build_unhandled_public_support_response; +use crate::handlers::public::support::mark_sensitive_response_no_store; use axum::{body::Body, http, response::Response}; use super::{ @@ -18,9 +19,10 @@ use super::{ handle_users_me_preferences_put, handle_users_me_providers_get, handle_users_me_referral_get, handle_users_me_sessions_get, handle_users_me_update_session, handle_users_me_usage_active_get, handle_users_me_usage_get, handle_users_me_usage_heatmap_get, - handle_users_me_usage_interval_timeline_get, users_me_api_key_capabilities_path_matches, - users_me_api_key_detail_path_matches, users_me_api_key_install_sessions_path_matches, - users_me_api_key_providers_path_matches, users_me_management_token_detail_path_matches, + handle_users_me_usage_interval_timeline_get, handle_users_me_vscodex_request, + users_me_api_key_capabilities_path_matches, users_me_api_key_detail_path_matches, + users_me_api_key_install_sessions_path_matches, users_me_api_key_providers_path_matches, + users_me_management_token_detail_path_matches, users_me_management_token_regenerate_path_matches, users_me_management_token_toggle_path_matches, users_me_management_tokens_root, users_me_session_detail_path_matches, AppState, GatewayPublicRequestContext, @@ -30,6 +32,7 @@ pub(crate) async fn maybe_build_local_users_me_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + remote_addr: &std::net::SocketAddr, request_body: Option<&axum::body::Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -55,6 +58,14 @@ pub(crate) async fn maybe_build_local_users_me_response( { Some(handle_users_me_delete_other_sessions(state, request_context, headers).await) } + Some( + "vscodex_devices_list" + | "vscodex_pairing_create" + | "vscodex_device_delete" + | "vscodex_ws_ticket_create", + ) => Some( + handle_users_me_vscodex_request(state, request_context, headers, request_body).await, + ), Some("session_delete") if users_me_session_detail_path_matches(&request_context.request_path) => { @@ -73,7 +84,9 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_tokens_list") if users_me_management_tokens_root(&request_context.request_path) => { - Some(handle_users_me_management_tokens_list(state, request_context, headers).await) + Some(mark_sensitive_response_no_store( + handle_users_me_management_tokens_list(state, request_context, headers).await, + )) } Some("api_keys_create") if request_context.request_path == "/api/users/me/api-keys" => { Some( @@ -83,7 +96,7 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_tokens_create") if users_me_management_tokens_root(&request_context.request_path) => { - Some( + Some(mark_sensitive_response_no_store( handle_users_me_management_token_create( state, request_context, @@ -91,7 +104,7 @@ pub(crate) async fn maybe_build_local_users_me_response( request_body, ) .await, - ) + )) } Some("api_key_detail") if users_me_api_key_detail_path_matches(&request_context.request_path) => @@ -101,7 +114,9 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_token_detail") if users_me_management_token_detail_path_matches(&request_context.request_path) => { - Some(handle_users_me_management_token_detail_get(state, request_context, headers).await) + Some(mark_sensitive_response_no_store( + handle_users_me_management_token_detail_get(state, request_context, headers).await, + )) } Some("api_key_update") if users_me_api_key_detail_path_matches(&request_context.request_path) => @@ -113,7 +128,7 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_token_update") if users_me_management_token_detail_path_matches(&request_context.request_path) => { - Some( + Some(mark_sensitive_response_no_store( handle_users_me_management_token_update( state, request_context, @@ -121,7 +136,7 @@ pub(crate) async fn maybe_build_local_users_me_response( request_body, ) .await, - ) + )) } Some("api_key_patch") if users_me_api_key_detail_path_matches(&request_context.request_path) => @@ -131,7 +146,9 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_token_toggle") if users_me_management_token_toggle_path_matches(&request_context.request_path) => { - Some(handle_users_me_management_token_toggle(state, request_context, headers).await) + Some(mark_sensitive_response_no_store( + handle_users_me_management_token_toggle(state, request_context, headers).await, + )) } Some("api_key_delete") if users_me_api_key_detail_path_matches(&request_context.request_path) => @@ -141,7 +158,9 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_token_delete") if users_me_management_token_detail_path_matches(&request_context.request_path) => { - Some(handle_users_me_management_token_delete(state, request_context, headers).await) + Some(mark_sensitive_response_no_store( + handle_users_me_management_token_delete(state, request_context, headers).await, + )) } Some("api_key_providers_update") if users_me_api_key_providers_path_matches(&request_context.request_path) => @@ -177,6 +196,7 @@ pub(crate) async fn maybe_build_local_users_me_response( state, request_context, headers, + remote_addr, request_body, ) .await, @@ -185,7 +205,9 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("management_token_regenerate") if users_me_management_token_regenerate_path_matches(&request_context.request_path) => { - Some(handle_users_me_management_token_regenerate(state, request_context, headers).await) + Some(mark_sensitive_response_no_store( + handle_users_me_management_token_regenerate(state, request_context, headers).await, + )) } Some("usage") if request_context.request_path == "/api/users/me/usage" => { Some(handle_users_me_usage_get(state, request_context, headers).await) @@ -221,7 +243,10 @@ pub(crate) async fn maybe_build_local_users_me_response( Some(handle_users_me_available_models(state, request_context, headers).await) } Some("client_config") if request_context.request_path == "/api/users/me/client-config" => { - Some(handle_users_me_client_config_get(state, request_context, headers).await) + Some( + handle_users_me_client_config_get(state, request_context, headers, remote_addr) + .await, + ) } Some("model_capabilities") if request_context.request_path == "/api/users/me/model-capabilities" => diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_sessions.rs b/apps/aether-gateway/src/handlers/public/support/user_me_sessions.rs index 6adc52fe6..0dbe5ecd2 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_sessions.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_sessions.rs @@ -12,6 +12,7 @@ use super::{ format_users_me_required_session_datetime_iso8601, resolve_authenticated_local_user, AppState, GatewayPublicRequestContext, }; +use crate::handlers::public::support::mark_sensitive_response_no_store; use crate::GatewayUserSessionView; #[derive(Debug, Deserialize)] @@ -79,13 +80,15 @@ pub(super) async fn handle_users_me_sessions_get( } }; - Json( - sessions - .into_iter() - .map(|session| build_users_me_session_payload(session, &auth.session_id)) - .collect::>(), + mark_sensitive_response_no_store( + Json( + sessions + .into_iter() + .map(|session| build_users_me_session_payload(session, &auth.session_id)) + .collect::>(), + ) + .into_response(), ) - .into_response() } pub(super) async fn handle_users_me_delete_other_sessions( 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 970761fad..3cca499df 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 @@ -31,6 +31,38 @@ use super::{ }; const USERS_ME_USAGE_DATA_UNAVAILABLE_DETAIL: &str = "用户用量数据暂不可用"; +// The active-usage endpoint accepts an explicit list of request IDs. Keep +// this list bounded before it reaches the repository layer: SQLite/MySQL +// expand every value into a bind parameter, while PostgreSQL still has to +// materialize the complete array. Request IDs are normally UUIDs, but a +// generous per-item bound preserves compatibility with provider-generated +// identifiers without allowing query amplification. +const MAX_USERS_ME_USAGE_IDS: usize = 256; +const MAX_USERS_ME_USAGE_ID_BYTES: usize = 256; +const MAX_USERS_ME_USAGE_IDS_QUERY_BYTES: usize = 64 * 1024; + +fn users_me_usage_public_error_message(item: &StoredRequestUsageAudit) -> Option { + if item + .error_message + .as_deref() + .is_none_or(|value| value.trim().is_empty()) + && item.status_code.is_none_or(|status| status < 400) + { + return None; + } + + item.error_category + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| { + item.status_code + .filter(|status| *status >= 400) + .map(|status| format!("http_{status}")) + }) + .or_else(|| Some("request_failed".to_string())) +} fn build_users_me_usage_reader_unavailable_response() -> Response { build_auth_error_response( @@ -133,15 +165,37 @@ fn parse_users_me_usage_timeline_limit(query: Option<&str>) -> Result) -> Option> { - let ids = query_param_value(query, "ids")?; - let values = ids +fn parse_users_me_usage_ids(query: Option<&str>) -> Result>, String> { + let Some(ids) = query_param_value(query, "ids") else { + return Ok(None); + }; + if ids.len() > MAX_USERS_ME_USAGE_IDS_QUERY_BYTES { + return Err(format!( + "ids query value must not exceed {MAX_USERS_ME_USAGE_IDS_QUERY_BYTES} bytes" + )); + } + + let mut values = BTreeSet::new(); + let mut item_count = 0usize; + for value in ids .split(',') .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .collect::>(); - (!values.is_empty()).then_some(values) + { + item_count = item_count.saturating_add(1); + if item_count > MAX_USERS_ME_USAGE_IDS { + return Err(format!( + "ids must contain at most {MAX_USERS_ME_USAGE_IDS} identifiers" + )); + } + if value.len() > MAX_USERS_ME_USAGE_ID_BYTES { + return Err(format!( + "each id must not exceed {MAX_USERS_ME_USAGE_ID_BYTES} bytes" + )); + } + values.insert(value.to_owned()); + } + Ok((!values.is_empty()).then_some(values)) } fn users_me_usage_cache_creation_tokens(item: &StoredRequestUsageAudit) -> u64 { @@ -546,7 +600,7 @@ fn build_users_me_usage_record_payload( "cache_creation_ephemeral_1h_input_tokens": item.cache_creation_ephemeral_1h_input_tokens, "cache_read_input_tokens": item.cache_read_input_tokens, "status_code": item.status_code, - "error_message": item.error_message, + "error_message": users_me_usage_public_error_message(item), "request_type": item.request_type, "input_price_per_1m": input_price_per_1m, "output_price_per_1m": output_price_per_1m, @@ -609,7 +663,7 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_ "updated_at": unix_secs_to_rfc3339(item.updated_at_unix_secs), "response_time_updated_at": users_me_usage_response_time_updated_at(item), "status_code": item.status_code, - "error_message": item.error_message, + "error_message": users_me_usage_public_error_message(item), "api_format": item.api_format, "endpoint_api_format": item.endpoint_api_format, "is_stream": item.is_stream, @@ -746,9 +800,6 @@ fn users_me_usage_terminal_candidate_state_override( if let Some(status_code) = candidate.status_code { payload["status_code"] = json!(status_code); } - if let Some(error_message) = candidate.error_message.as_ref() { - payload["error_message"] = json!(error_message); - } Some(payload) } @@ -1338,7 +1389,12 @@ pub(super) async fn handle_users_me_usage_active_get( Ok(value) => value, Err(response) => return response, }; - let ids = parse_users_me_usage_ids(request_context.request_query_string.as_deref()); + let ids = match parse_users_me_usage_ids(request_context.request_query_string.as_deref()) { + Ok(value) => value, + Err(message) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, message, false); + } + }; // When polling for active (pending/streaming) requests without specific ids, // limit to the last 1 hour to avoid scanning all historical records. let items = match ids.as_ref() { @@ -1628,11 +1684,37 @@ mod tests { use super::{ build_users_me_usage_active_payload, build_users_me_usage_record_payload, - parse_users_me_usage_record_filter, users_me_usage_client_is_stream, - users_me_usage_is_failed, users_me_usage_terminal_candidate_state_override, - users_me_usage_upstream_is_stream, + parse_users_me_usage_ids, parse_users_me_usage_record_filter, + users_me_usage_client_is_stream, users_me_usage_is_failed, + users_me_usage_terminal_candidate_state_override, users_me_usage_upstream_is_stream, + MAX_USERS_ME_USAGE_IDS, MAX_USERS_ME_USAGE_ID_BYTES, }; + #[test] + fn user_active_usage_ids_are_trimmed_and_deduplicated() { + let ids = parse_users_me_usage_ids(Some("ids=req-2,%20req-1,req-2")) + .expect("bounded ids should parse") + .expect("non-empty ids should be present"); + + assert_eq!(ids.into_iter().collect::>(), vec!["req-1", "req-2"]); + assert_eq!( + parse_users_me_usage_ids(Some("other=value")).expect("missing ids should parse"), + None + ); + } + + #[test] + fn user_active_usage_ids_reject_query_amplification() { + let too_many = (0..=MAX_USERS_ME_USAGE_IDS) + .map(|index| format!("req-{index}")) + .collect::>() + .join(","); + assert!(parse_users_me_usage_ids(Some(&format!("ids={too_many}"))).is_err()); + + let oversized = "x".repeat(MAX_USERS_ME_USAGE_ID_BYTES + 1); + assert!(parse_users_me_usage_ids(Some(&format!("ids={oversized}"))).is_err()); + } + #[test] fn users_me_usage_transport_statuses_are_disjoint_server_side_filters() { let live_websocket = parse_users_me_usage_record_filter(Some( @@ -1931,6 +2013,26 @@ mod tests { assert!(users_me_usage_is_failed(&item)); } + #[test] + fn user_usage_payload_does_not_return_historical_raw_error_text() { + let item = StoredRequestUsageAudit { + status_code: Some(401), + error_message: Some( + "upstream said Authorization: Bearer live-secret at https://api.example?key=secret" + .to_string(), + ), + error_category: Some("authentication_error".to_string()), + ..sample_usage("failed") + }; + + let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false); + let active = build_users_me_usage_active_payload(&item); + assert_eq!(record["error_message"], "authentication_error"); + assert_eq!(active["error_message"], "authentication_error"); + assert!(!record.to_string().contains("live-secret")); + assert!(!active.to_string().contains("live-secret")); + } + #[test] fn user_usage_payloads_include_symmetric_stream_fields() { let item = StoredRequestUsageAudit { diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_vscodex.rs b/apps/aether-gateway/src/handlers/public/support/user_me_vscodex.rs new file mode 100644 index 000000000..df2014e05 --- /dev/null +++ b/apps/aether-gateway/src/handlers/public/support/user_me_vscodex.rs @@ -0,0 +1,931 @@ +use std::collections::HashMap; +use std::net::IpAddr; +use std::sync::{Arc, LazyLock, Mutex}; +use std::time::Duration; + +use axum::body::{Body, Bytes}; +use axum::extract::{ + ws::{CloseFrame as AxumCloseFrame, Message as AxumMessage, WebSocket, WebSocketUpgrade}, + ConnectInfo, State, +}; +use axum::http::{self, header}; +use axum::response::{IntoResponse, Response}; +use futures_util::{SinkExt, StreamExt}; +use serde::Deserialize; +use serde_json::{json, Map, Value}; +use tokio::sync::Semaphore; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::protocol::{ + CloseFrame as TungsteniteCloseFrame, WebSocketConfig, +}; +use tokio_tungstenite::tungstenite::Message as TungsteniteMessage; +use tracing::warn; + +use super::{ + build_auth_error_response, build_auth_json_response, module_available_from_env, + resolve_authenticated_local_user, AppState, GatewayPublicRequestContext, +}; + +const VSCODEX_ENABLED_ENV: &str = "AETHER_VSCODEX_ENABLED"; +const VSCODEX_INTERNAL_URL_ENV: &str = "AETHER_VSCODEX_INTERNAL_URL"; +const VSCODEX_INTERNAL_TOKEN_ENV: &str = "AETHER_VSCODEX_INTERNAL_TOKEN"; +const VSCODEX_REQUEST_TIMEOUT: Duration = Duration::from_secs(15); +const VSCODEX_MAX_RESPONSE_BYTES: usize = 1024 * 1024; +const VSCODEX_WS_MAX_MESSAGE_BYTES: usize = 16 * 1024 * 1024; +const VSCODEX_WS_MAX_CONNECTIONS: usize = 256; +const VSCODEX_WS_MAX_CONNECTIONS_PER_IP: usize = 16; +const VSCODEX_DEVICE_PATH_PREFIX: &str = "/api/users/me/vscodex/devices/"; +const VSCODEX_CLIENT_IP_HEADER: &str = "x-aether-client-ip"; + +static VSCODEX_HTTP_CLIENT: LazyLock> = + LazyLock::new(|| { + reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + }); +static VSCODEX_WS_CONNECTIONS: LazyLock> = + LazyLock::new(|| Arc::new(Semaphore::new(VSCODEX_WS_MAX_CONNECTIONS))); +static VSCODEX_WS_CONNECTIONS_BY_IP: LazyLock> = + LazyLock::new(|| { + Arc::new(VscodexWsIpConnectionLimiter::new( + VSCODEX_WS_MAX_CONNECTIONS_PER_IP, + )) + }); + +#[derive(Debug)] +struct VscodexWsIpConnectionLimiter { + max_connections: usize, + active: Mutex>, +} + +impl VscodexWsIpConnectionLimiter { + fn new(max_connections: usize) -> Self { + Self { + max_connections: max_connections.max(1), + active: Mutex::new(HashMap::new()), + } + } + + fn try_acquire(self: &Arc, client_ip: IpAddr) -> Option { + let mut active = self + .active + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let current = active.get(&client_ip).copied().unwrap_or_default(); + if current >= self.max_connections { + return None; + } + active.insert(client_ip, current.saturating_add(1)); + Some(VscodexWsIpConnectionPermit { + limiter: Arc::clone(self), + client_ip, + }) + } + + fn release(&self, client_ip: IpAddr) { + let mut active = self + .active + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let Some(current) = active.get_mut(&client_ip) else { + return; + }; + if *current <= 1 { + active.remove(&client_ip); + } else { + *current -= 1; + } + } + + #[cfg(test)] + fn active_ip_count(&self) -> usize { + self.active + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .len() + } +} + +#[derive(Debug)] +struct VscodexWsIpConnectionPermit { + limiter: Arc, + client_ip: IpAddr, +} + +impl Drop for VscodexWsIpConnectionPermit { + fn drop(&mut self) { + self.limiter.release(self.client_ip); + } +} + +#[derive(Debug)] +struct VscodexSidecarConfig { + base_url: reqwest::Url, + authorization: reqwest::header::HeaderValue, + http_client: reqwest::Client, +} + +#[derive(Debug, Default, Deserialize)] +struct CreatePairingRequest { + name: Option, +} + +#[derive(Debug, Deserialize)] +struct CreateWsTicketRequest { + device_id: String, +} + +#[derive(Debug, Deserialize)] +struct ExchangePairingRequest { + code: String, + name: Option, +} + +pub(crate) async fn vscodex_ws_proxy( + State(state): State, + ConnectInfo(remote_addr): ConnectInfo, + ws: WebSocketUpgrade, + headers: http::HeaderMap, +) -> Response { + let request_permit = match state.try_acquire_request_permit().await { + Ok(value) => value, + Err(err) => { + warn!(error = ?err, "VS Codex WebSocket request admission rejected"); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "服务繁忙,请稍后重试", + false, + ); + } + }; + let client_ip = crate::headers::effective_client_ip(&headers, &remote_addr); + match state.admin_security_ip_blacklisted(client_ip).await { + Ok(true) => { + return build_auth_error_response( + http::StatusCode::FORBIDDEN, + "当前 IP 已被禁止访问", + false, + ) + } + Ok(false) => {} + Err(err) => warn!( + client_ip = %client_ip, + error = ?err, + "VS Codex WebSocket IP blacklist check failed open" + ), + } + let connection_permit = match Arc::clone(&VSCODEX_WS_CONNECTIONS).try_acquire_owned() { + Ok(value) => value, + Err(_) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 连接数已达上限", + false, + ) + } + }; + // Only active connections have entries, and each already owns one of the 256 global slots. + let ip_connection_permit = match VSCODEX_WS_CONNECTIONS_BY_IP.try_acquire(client_ip) { + Some(value) => value, + None => { + warn!( + client_ip = %client_ip, + limit = VSCODEX_WS_MAX_CONNECTIONS_PER_IP, + "VS Codex per-IP WebSocket connection limit reached" + ); + let mut response = build_auth_error_response( + http::StatusCode::TOO_MANY_REQUESTS, + "当前 IP 的 VS Codex 连接数已达上限", + false, + ); + response + .headers_mut() + .insert(header::RETRY_AFTER, http::HeaderValue::from_static("1")); + return response; + } + }; + let config = match load_vscodex_sidecar_config() { + Ok(Some(value)) => value, + Ok(None) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务未启用", + false, + ) + } + Err(detail) => { + warn!(error = %detail, "VS Codex WebSocket sidecar configuration is invalid"); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + ); + } + }; + let sidecar_url = match build_vscodex_websocket_url(&config.base_url) { + Ok(value) => value, + Err(detail) => { + warn!(error = %detail, "could not build VS Codex sidecar WebSocket URL"); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + ); + } + }; + let mut sidecar_request = match sidecar_url.as_str().into_client_request() { + Ok(value) => value, + Err(err) => { + warn!(error = %err, "could not build VS Codex sidecar WebSocket request"); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + ); + } + }; + for header_name in [header::ORIGIN, header::SEC_WEBSOCKET_PROTOCOL] { + if let Some(value) = headers.get(&header_name) { + sidecar_request + .headers_mut() + .insert(header_name, value.clone()); + } + } + + let mut sidecar_config = WebSocketConfig::default(); + sidecar_config.max_message_size = Some(VSCODEX_WS_MAX_MESSAGE_BYTES); + sidecar_config.max_frame_size = Some(VSCODEX_WS_MAX_MESSAGE_BYTES); + let (sidecar_socket, sidecar_response) = match tokio::time::timeout( + VSCODEX_REQUEST_TIMEOUT, + tokio_tungstenite::connect_async_with_config(sidecar_request, Some(sidecar_config), true), + ) + .await + { + Ok(Ok(value)) => value, + Ok(Err(err)) => { + warn!(error = %err, "VS Codex sidecar WebSocket handshake failed"); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务暂时不可用", + false, + ); + } + Err(_) => { + warn!("VS Codex sidecar WebSocket handshake timed out"); + return build_auth_error_response( + http::StatusCode::GATEWAY_TIMEOUT, + "VS Codex 服务请求超时", + false, + ); + } + }; + + let selected_protocol = sidecar_response + .headers() + .get(header::SEC_WEBSOCKET_PROTOCOL) + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + let ws = ws + .max_message_size(VSCODEX_WS_MAX_MESSAGE_BYTES) + .max_frame_size(VSCODEX_WS_MAX_MESSAGE_BYTES); + let ws = match selected_protocol { + Some(protocol) => ws.protocols([protocol]), + None => ws, + }; + drop(request_permit); + ws.on_upgrade(move |browser_socket| async move { + let _connection_permit = connection_permit; + bridge_vscodex_websockets(browser_socket, sidecar_socket, ip_connection_permit).await; + }) +} + +pub(super) async fn maybe_build_local_vscodex_response( + _state: &AppState, + request_context: &GatewayPublicRequestContext, + client_ip: std::net::IpAddr, + request_body: Option<&Bytes>, +) -> Option> { + let decision = request_context.control_decision.as_ref()?; + if decision.route_family.as_deref() != Some("vscodex") { + return None; + } + if decision.route_kind.as_deref() != Some("pairing_exchange") + || !matches!( + request_context.request_path.as_str(), + "/api/vscodex/pair" | "/api/vscodex/pair/" + ) + { + return Some(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "VS Codex 接口不存在", + false, + )); + } + + let config = match load_vscodex_sidecar_config() { + Ok(Some(value)) => value, + Ok(None) => { + return Some(build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务未启用", + false, + )) + } + Err(detail) => { + warn!( + error = %detail, + "VS Codex sidecar configuration is invalid" + ); + return Some(build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + )); + } + }; + let payload = match parse_pairing_exchange_request(request_body) { + Ok(value) => value, + Err(response) => return Some(response), + }; + let url = match append_vscodex_sidecar_path(&config.base_url, &["v1", "pairings", "exchange"]) { + Ok(value) => value, + Err(detail) => { + warn!(error = %detail, "could not build VS Codex pairing exchange URL"); + return Some(build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + )); + } + }; + let request = + build_authenticated_sidecar_request(&config, reqwest::Method::POST, url, Some(payload)) + .header(VSCODEX_CLIENT_IP_HEADER, client_ip.to_string()); + Some(send_vscodex_sidecar_request(request, "public", "pairing_exchange").await) +} + +pub(super) async fn handle_users_me_vscodex_request( + state: &AppState, + request_context: &GatewayPublicRequestContext, + headers: &http::HeaderMap, + request_body: Option<&Bytes>, +) -> Response { + let auth = match resolve_authenticated_local_user(state, request_context, headers).await { + Ok(value) => value, + Err(response) => return response, + }; + let config = match load_vscodex_sidecar_config() { + Ok(Some(value)) => value, + Ok(None) => { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务未启用", + false, + ) + } + Err(detail) => { + warn!( + user_id = %auth.user.id, + error = %detail, + "VS Codex sidecar configuration is invalid" + ); + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + ); + } + }; + + let Some(route_kind) = request_context + .control_decision + .as_ref() + .and_then(|decision| decision.route_kind.as_deref()) + else { + return build_auth_error_response( + http::StatusCode::NOT_FOUND, + "VS Codex 接口不存在", + false, + ); + }; + + let request = match build_vscodex_sidecar_request( + &config, + &auth.user.id, + route_kind, + &request_context.request_path, + request_body, + ) { + Ok(value) => value, + Err(response) => return response, + }; + + send_vscodex_sidecar_request(request, &auth.user.id, route_kind).await +} + +fn load_vscodex_sidecar_config() -> Result, String> { + if !module_available_from_env(VSCODEX_ENABLED_ENV, false) { + return Ok(None); + } + + let raw_url = required_env(VSCODEX_INTERNAL_URL_ENV)?; + let base_url = reqwest::Url::parse(&raw_url) + .map_err(|err| format!("{VSCODEX_INTERNAL_URL_ENV} is invalid: {err}"))?; + if !matches!(base_url.scheme(), "http" | "https") + || !base_url.has_host() + || !base_url.username().is_empty() + || base_url.password().is_some() + || base_url.query().is_some() + || base_url.fragment().is_some() + || base_url.cannot_be_a_base() + { + return Err(format!( + "{VSCODEX_INTERNAL_URL_ENV} must be an HTTP(S) base URL without credentials, query, or fragment" + )); + } + + let token = required_env(VSCODEX_INTERNAL_TOKEN_ENV)?; + let authorization = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}")) + .map_err(|_| format!("{VSCODEX_INTERNAL_TOKEN_ENV} is not a valid HTTP credential"))?; + let http_client = VSCODEX_HTTP_CLIENT + .as_ref() + .map_err(|err| format!("could not initialize VS Codex HTTP client: {err}"))? + .clone(); + + Ok(Some(VscodexSidecarConfig { + base_url, + authorization, + http_client, + })) +} + +fn required_env(key: &str) -> Result { + std::env::var(key) + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .ok_or_else(|| format!("{key} is required")) +} + +fn build_vscodex_sidecar_request( + config: &VscodexSidecarConfig, + user_id: &str, + route_kind: &str, + request_path: &str, + request_body: Option<&Bytes>, +) -> Result> { + let (method, suffix, payload) = match route_kind { + "vscodex_devices_list" => (reqwest::Method::GET, vec!["devices"], None), + "vscodex_pairing_create" => ( + reqwest::Method::POST, + vec!["pairings"], + Some(parse_pairing_request(request_body)?), + ), + "vscodex_device_delete" => { + let Some(device_id) = vscodex_device_id_from_path(request_path) else { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "设备标识无效", + false, + )); + }; + (reqwest::Method::DELETE, vec!["devices", device_id], None) + } + "vscodex_ws_ticket_create" => ( + reqwest::Method::POST, + vec!["ws-tickets"], + Some(parse_ws_ticket_request(request_body)?), + ), + _ => { + return Err(build_auth_error_response( + http::StatusCode::NOT_FOUND, + "VS Codex 接口不存在", + false, + )) + } + }; + let url = build_vscodex_sidecar_url(&config.base_url, user_id, &suffix).map_err(|detail| { + warn!(user_id = %user_id, error = %detail, "could not build VS Codex sidecar URL"); + build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "VS Codex 服务配置不完整", + false, + ) + })?; + + Ok(build_authenticated_sidecar_request( + config, method, url, payload, + )) +} + +fn build_authenticated_sidecar_request( + config: &VscodexSidecarConfig, + method: reqwest::Method, + url: reqwest::Url, + payload: Option, +) -> reqwest::RequestBuilder { + let mut request = config + .http_client + .request(method, url) + .header(header::AUTHORIZATION, config.authorization.clone()) + .header(header::ACCEPT, "application/json") + .timeout(VSCODEX_REQUEST_TIMEOUT); + if let Some(payload) = payload { + request = request.json(&payload); + } + request +} + +fn build_vscodex_sidecar_url( + base_url: &reqwest::Url, + user_id: &str, + suffix: &[&str], +) -> Result { + let mut segments = vec!["internal", "v1", "users", user_id]; + segments.extend(suffix.iter().copied()); + append_vscodex_sidecar_path(base_url, &segments) +} + +fn append_vscodex_sidecar_path( + base_url: &reqwest::Url, + suffix: &[&str], +) -> Result { + let mut url = base_url.clone(); + let mut path_segments = url + .path_segments_mut() + .map_err(|_| "VS Codex sidecar URL cannot contain path segments".to_string())?; + path_segments.pop_if_empty(); + path_segments.extend(suffix.iter().copied()); + drop(path_segments); + Ok(url) +} + +fn build_vscodex_websocket_url(base_url: &reqwest::Url) -> Result { + let mut url = append_vscodex_sidecar_path(base_url, &["api", "vscodex", "ws"])?; + let scheme = match url.scheme() { + "http" => "ws", + "https" => "wss", + _ => return Err("VS Codex sidecar URL must use HTTP(S)".to_string()), + }; + url.set_scheme(scheme) + .map_err(|_| "could not convert VS Codex sidecar URL to WebSocket".to_string())?; + Ok(url) +} + +fn parse_pairing_request(request_body: Option<&Bytes>) -> Result> { + let payload = parse_json_request::(request_body, true)?; + let mut object = Map::new(); + if let Some(name) = payload.name { + object.insert("name".to_string(), Value::String(name)); + } + Ok(Value::Object(object)) +} + +fn parse_ws_ticket_request(request_body: Option<&Bytes>) -> Result> { + let payload = parse_json_request::(request_body, false)?; + let device_id = payload.device_id.trim(); + if !valid_vscodex_device_id(device_id) { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "设备标识无效", + false, + )); + } + Ok(json!({ "device_id": device_id })) +} + +fn parse_pairing_exchange_request(request_body: Option<&Bytes>) -> Result> { + let payload = parse_json_request::(request_body, false)?; + let code = payload.code.trim(); + if code.is_empty() || code.len() > 256 || code.chars().any(char::is_control) { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "配对码无效", + false, + )); + } + let mut object = Map::from_iter([("code".to_string(), Value::String(code.to_string()))]); + if let Some(name) = payload.name { + object.insert("name".to_string(), Value::String(name)); + } + Ok(Value::Object(object)) +} + +fn parse_json_request( + request_body: Option<&Bytes>, + empty_object_allowed: bool, +) -> Result> +where + T: serde::de::DeserializeOwned, +{ + let body = request_body.filter(|body| !body.is_empty()); + let result = match body { + Some(body) => serde_json::from_slice(body), + None if empty_object_allowed => serde_json::from_slice(b"{}"), + None => { + return Err(build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "缺少请求体", + false, + )) + } + }; + result.map_err(|_| { + build_auth_error_response(http::StatusCode::BAD_REQUEST, "请求数据验证失败", false) + }) +} + +fn vscodex_device_id_from_path(path: &str) -> Option<&str> { + let trimmed = path.trim_end_matches('/'); + let device_id = trimmed.strip_prefix(VSCODEX_DEVICE_PATH_PREFIX)?; + if device_id.contains('/') || !valid_vscodex_device_id(device_id) { + return None; + } + Some(device_id) +} + +fn valid_vscodex_device_id(value: &str) -> bool { + !value.is_empty() && value.len() <= 128 && !value.chars().any(char::is_control) +} + +async fn send_vscodex_sidecar_request( + request: reqwest::RequestBuilder, + request_scope: &str, + operation: &str, +) -> Response { + let mut upstream = match request.send().await { + Ok(value) => value, + Err(err) => { + warn!( + request_scope = %request_scope, + operation = %operation, + error = %err, + "VS Codex sidecar request failed" + ); + let (status, detail) = if err.is_timeout() { + (http::StatusCode::GATEWAY_TIMEOUT, "VS Codex 服务请求超时") + } else { + (http::StatusCode::BAD_GATEWAY, "VS Codex 服务暂时不可用") + }; + return build_auth_error_response(status, detail, false); + } + }; + let status = http::StatusCode::from_u16(upstream.status().as_u16()) + .unwrap_or(http::StatusCode::BAD_GATEWAY); + + if matches!( + status, + http::StatusCode::UNAUTHORIZED | http::StatusCode::FORBIDDEN + ) { + warn!( + request_scope = %request_scope, + operation = %operation, + upstream_status = status.as_u16(), + "VS Codex sidecar rejected gateway credentials" + ); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务鉴权失败", + false, + ); + } + if status.is_redirection() { + warn!( + request_scope = %request_scope, + operation = %operation, + upstream_status = status.as_u16(), + "VS Codex sidecar returned an unexpected redirect" + ); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务返回无效响应", + false, + ); + } + if status == http::StatusCode::NO_CONTENT { + return vscodex_no_store_response(status.into_response(), None); + } + + let mut response_body = Vec::new(); + while let Some(chunk) = match upstream.chunk().await { + Ok(value) => value, + Err(err) => { + warn!( + request_scope = %request_scope, + operation = %operation, + error = %err, + "could not read VS Codex sidecar response" + ); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务返回无效响应", + false, + ); + } + } { + if response_body.len().saturating_add(chunk.len()) > VSCODEX_MAX_RESPONSE_BYTES { + warn!( + request_scope = %request_scope, + operation = %operation, + "VS Codex sidecar response exceeded the size limit" + ); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务返回无效响应", + false, + ); + } + response_body.extend_from_slice(&chunk); + } + + if response_body.is_empty() { + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务返回无效响应", + false, + ); + } + let payload = match serde_json::from_slice(&response_body) { + Ok(value) => value, + Err(err) => { + warn!( + request_scope = %request_scope, + operation = %operation, + upstream_status = status.as_u16(), + error = %err, + "VS Codex sidecar returned non-JSON data" + ); + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + "VS Codex 服务返回无效响应", + false, + ); + } + }; + let retry_after = upstream.headers().get(header::RETRY_AFTER).cloned(); + vscodex_no_store_response(build_auth_json_response(status, payload, None), retry_after) +} + +fn vscodex_no_store_response( + mut response: Response, + retry_after: Option, +) -> Response { + response.headers_mut().insert( + header::CACHE_CONTROL, + http::HeaderValue::from_static("no-store"), + ); + if let Some(retry_after) = retry_after { + response + .headers_mut() + .insert(header::RETRY_AFTER, retry_after); + } + response +} + +async fn bridge_vscodex_websockets( + browser_socket: WebSocket, + sidecar_socket: S, + ip_connection_permit: VscodexWsIpConnectionPermit, +) where + S: futures_util::Stream< + Item = Result, + > + futures_util::Sink + + Unpin + + Send + + 'static, +{ + let (mut browser_tx, mut browser_rx) = browser_socket.split(); + let (mut sidecar_tx, mut sidecar_rx) = sidecar_socket.split(); + let mut ip_connection_permit = Some(ip_connection_permit); + + loop { + tokio::select! { + browser_message = browser_rx.next() => { + match browser_message { + Some(Ok(message)) => { + let close = matches!(message, AxumMessage::Close(_)); + if let Err(err) = sidecar_tx.send(axum_to_tungstenite_message(message)).await { + warn!(error = %err, "could not forward VS Codex browser WebSocket frame"); + break; + } + if close { + break; + } + } + Some(Err(err)) => { + warn!(error = %err, "VS Codex browser WebSocket read failed"); + break; + } + None => break, + } + } + sidecar_message = sidecar_rx.next() => { + match sidecar_message { + Some(Ok(TungsteniteMessage::Frame(_))) => continue, + Some(Ok(message)) => { + if ip_connection_permit.is_some() && vscodex_ws_authentication_succeeded(&message) { + ip_connection_permit.take(); + } + let close = matches!(message, TungsteniteMessage::Close(_)); + if let Err(err) = browser_tx.send(tungstenite_to_axum_message(message)).await { + warn!(error = %err, "could not forward VS Codex sidecar WebSocket frame"); + break; + } + if close { + break; + } + } + Some(Err(err)) => { + warn!(error = %err, "VS Codex sidecar WebSocket read failed"); + break; + } + None => break, + } + } + } + } + + let _ = sidecar_tx.close().await; + let _ = browser_tx.close().await; +} + +fn vscodex_ws_authentication_succeeded(message: &TungsteniteMessage) -> bool { + let TungsteniteMessage::Text(text) = message else { + return false; + }; + serde_json::from_str::(text.as_ref()) + .ok() + .and_then(|payload| { + payload + .get("type") + .and_then(Value::as_str) + .map(str::to_string) + }) + .as_deref() + == Some("auth.ok") +} + +fn axum_to_tungstenite_message(message: AxumMessage) -> TungsteniteMessage { + match message { + AxumMessage::Text(text) => TungsteniteMessage::Text(text.to_string().into()), + AxumMessage::Binary(bytes) => TungsteniteMessage::Binary(bytes), + AxumMessage::Ping(bytes) => TungsteniteMessage::Ping(bytes), + AxumMessage::Pong(bytes) => TungsteniteMessage::Pong(bytes), + AxumMessage::Close(frame) => { + TungsteniteMessage::Close(frame.map(|frame| TungsteniteCloseFrame { + code: frame.code.into(), + reason: frame.reason.to_string().into(), + })) + } + } +} + +fn tungstenite_to_axum_message(message: TungsteniteMessage) -> AxumMessage { + match message { + TungsteniteMessage::Text(text) => AxumMessage::Text(text.to_string().into()), + TungsteniteMessage::Binary(bytes) => AxumMessage::Binary(bytes), + TungsteniteMessage::Ping(bytes) => AxumMessage::Ping(bytes), + TungsteniteMessage::Pong(bytes) => AxumMessage::Pong(bytes), + TungsteniteMessage::Close(frame) => AxumMessage::Close(frame.map(|frame| AxumCloseFrame { + code: frame.code.into(), + reason: frame.reason.to_string().into(), + })), + TungsteniteMessage::Frame(_) => AxumMessage::Close(None), + } +} + +#[cfg(test)] +mod tests { + use super::{vscodex_ws_authentication_succeeded, VscodexWsIpConnectionLimiter}; + use std::sync::Arc; + use tokio_tungstenite::tungstenite::Message; + + #[test] + fn vscodex_ws_ip_limiter_releases_and_removes_inactive_ips() { + let limiter = Arc::new(VscodexWsIpConnectionLimiter::new(1)); + let client_ip = "198.51.100.10".parse().expect("IP should parse"); + + let permit = limiter + .try_acquire(client_ip) + .expect("first connection should acquire"); + assert_eq!(limiter.active_ip_count(), 1); + assert!(limiter.try_acquire(client_ip).is_none()); + + drop(permit); + assert_eq!(limiter.active_ip_count(), 0); + assert!(limiter.try_acquire(client_ip).is_some()); + } + + #[test] + fn vscodex_ws_ip_limiter_releases_only_after_sidecar_auth_success() { + assert!(vscodex_ws_authentication_succeeded(&Message::Text( + r#"{"type":"auth.ok","role":"operator"}"#.into() + ))); + assert!(!vscodex_ws_authentication_succeeded(&Message::Text( + r#"{"type":"auth","token":"client-controlled"}"#.into() + ))); + assert!(!vscodex_ws_authentication_succeeded(&Message::Binary( + br#"{"type":"auth.ok"}"#.to_vec().into() + ))); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/wallet.rs b/apps/aether-gateway/src/handlers/public/support/wallet.rs index ba454a432..56fe24958 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet.rs @@ -13,8 +13,8 @@ pub(crate) use self::test_support::wallet_test_recharge_store; #[cfg(test)] use self::test_support::{ record_wallet_test_recharge, record_wallet_test_refund, wallet_test_recharge_order_by_id, - wallet_test_recharge_orders_for_user, wallet_test_refund_by_id, - wallet_test_refund_by_idempotency, wallet_test_refunds_for_wallet, + wallet_test_recharge_order_by_order_no, wallet_test_recharge_orders_for_user, + wallet_test_refund_by_id, wallet_test_refund_by_idempotency, wallet_test_refunds_for_wallet, wallet_test_reserved_refund_amount, }; #[path = "wallet/flow.rs"] @@ -39,7 +39,12 @@ use self::reads::{ parse_wallet_limit, parse_wallet_offset, wallet_fixed_offset, wallet_transaction_payload_from_record, }; -pub(crate) use self::recharge::{direct_gateway_channels, sanitize_wallet_gateway_response}; +pub(crate) use self::recharge::{ + direct_gateway_channels, prepare_billing_gateway_response_for_storage, + prepare_wallet_gateway_response_for_storage, resolve_direct_gateway_channel, + sanitize_wallet_gateway_response, wallet_payment_instructions_from_checkout, + wallet_payment_instructions_from_stored, +}; use self::recharge::{ handle_wallet_create_recharge, handle_wallet_recharge_detail, handle_wallet_recharge_list, handle_wallet_recharge_options, wallet_recharge_detail_path_matches, @@ -72,12 +77,12 @@ const WALLET_SAFE_GATEWAY_RESPONSE_KEYS: &[&str] = &[ "code_url", "h5_url", "jsapi", - "client_secret", "publishable_key", "intent_id", "payment_method_types", "provider_label", "subject", + "instructions", "callback_url", "return_url", "integration_status", @@ -121,6 +126,7 @@ pub(super) async fn maybe_build_local_wallet_response( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Option> { let decision = request_context.control_decision.as_ref()?; @@ -184,7 +190,8 @@ pub(super) async fn maybe_build_local_wallet_response( && request_context.request_path == "/api/wallet/recharge" { return Some( - handle_wallet_create_recharge(state, request_context, headers, request_body).await, + handle_wallet_create_recharge(state, request_context, headers, client_ip, request_body) + .await, ); } @@ -220,7 +227,6 @@ mod tests { use super::{ build_wallet_recharge_storage_unavailable_response, build_wallet_refund_storage_unavailable_response, - WALLET_RECHARGE_STORAGE_UNAVAILABLE_DETAIL, WALLET_REFUND_STORAGE_UNAVAILABLE_DETAIL, }; use axum::body::to_bytes; use axum::http; @@ -236,10 +242,7 @@ mod tests { .expect("body should read"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); - assert_eq!( - payload, - json!({ "detail": WALLET_RECHARGE_STORAGE_UNAVAILABLE_DETAIL }) - ); + assert_eq!(payload, json!({ "detail": "服务暂不可用,请稍后重试" })); } #[tokio::test] @@ -252,9 +255,6 @@ mod tests { .expect("body should read"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); - assert_eq!( - payload, - json!({ "detail": WALLET_REFUND_STORAGE_UNAVAILABLE_DETAIL }) - ); + assert_eq!(payload, json!({ "detail": "服务暂不可用,请稍后重试" })); } } diff --git a/apps/aether-gateway/src/handlers/public/support/wallet/reads.rs b/apps/aether-gateway/src/handlers/public/support/wallet/reads.rs index d6da0c942..df9b3b3ba 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet/reads.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet/reads.rs @@ -4,6 +4,7 @@ use super::{ GatewayPublicRequestContext, Response, WALLET_LEGACY_TIMEZONE, }; use crate::handlers::shared::round_to; +use aether_data::repository::wallet::stored_timestamp_unix_secs; use aether_data_contracts::repository::usage::UsageSettledCostSummaryQuery; use chrono::{TimeZone, Utc}; use serde_json::json; @@ -309,7 +310,10 @@ pub(super) fn wallet_transaction_payload_from_record( "link_id": record.link_id.clone(), "operator_id": record.operator_id.clone(), "description": record.description.clone(), - "created_at": record.created_at_unix_ms.and_then(unix_secs_to_rfc3339), + "created_at": record + .created_at_unix_ms + .map(stored_timestamp_unix_secs) + .and_then(unix_secs_to_rfc3339), }) } diff --git a/apps/aether-gateway/src/handlers/public/support/wallet/recharge.rs b/apps/aether-gateway/src/handlers/public/support/wallet/recharge.rs index 68227d8a3..be91be7df 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet/recharge.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet/recharge.rs @@ -1,3 +1,4 @@ +use super::super::mark_sensitive_response_no_store; use super::super::support_payment::payment_epay::{ build_epay_checkout_url, configured_epay_channels, epay_callback_base_url, load_epay_config, resolve_epay_channel, EpayCheckoutInput, @@ -12,20 +13,86 @@ use super::{ wallet_normalize_optional_string_field, AppState, Body, GatewayPublicRequestContext, Response, WALLET_SAFE_GATEWAY_RESPONSE_KEYS, }; + +const MAX_PAYMENT_GATEWAY_RESPONSE_BYTES: usize = 1024 * 1024; +const WALLET_RECHARGE_ORDER_KIND: &str = "wallet_recharge"; +const PAYMENT_ORDER_STRIPE_SECRET_MIGRATION_RETRIES: usize = 8; +const PAYMENT_AMOUNT_EPSILON: f64 = 0.00000001; +const STRIPE_WALLET_IDEMPOTENCY_PREFIX: &str = "aether-payment-intent-"; +const STRIPE_CHECKOUT_UNCERTAIN_DETAIL: &str = "Stripe 支付服务暂时不可用"; +const STRIPE_CHECKOUT_FAILED_DETAIL: &str = "Stripe 支付请求被拒绝"; +const WALLET_SAFE_EPAY_PAYMENT_PARAM_KEYS: &[&str] = &[ + "pid", + "type", + "out_trade_no", + "notify_url", + "return_url", + "name", + "money", + "sign_type", + "sign", +]; + +#[derive(Debug, Clone, PartialEq, Eq)] +enum StripeWalletCheckoutError { + Canceled, + /// The provider request may have been accepted, but the local process + /// could not obtain a trustworthy response. Retrying with another gateway + /// identity could strand the first payment. + Uncertain(String), + Failed(String), +} + +impl From for StripeWalletCheckoutError { + fn from(value: String) -> Self { + Self::Failed(value) + } +} + +fn stripe_wallet_checkout_response_is_canceled(value: &Value) -> bool { + value + .get("status") + .and_then(Value::as_str) + .is_some_and(|status| status.trim().eq_ignore_ascii_case("canceled")) +} + +/// Keep the local merchant order number stable while rotating the provider +/// idempotency key once when Stripe has retained a canceled intent. The +/// suffix is deterministic so a later request can recover a response that was +/// lost after the retry was accepted, without creating a third intent. +fn stripe_wallet_idempotency_key(order_no: &str, retry: bool) -> String { + if retry { + format!("{STRIPE_WALLET_IDEMPOTENCY_PREFIX}{order_no}-retry-1") + } else { + format!("{STRIPE_WALLET_IDEMPOTENCY_PREFIX}{order_no}") + } +} + #[cfg(test)] use super::{ record_wallet_test_recharge, wallet_test_recharge_order_by_id, - wallet_test_recharge_orders_for_user, + wallet_test_recharge_order_by_order_no, wallet_test_recharge_orders_for_user, }; use crate::handlers::shared::{ - create_alipay_direct_checkout, create_wxpay_direct_checkout, direct_payment_client_ip, - DirectPaymentCheckoutInput, + create_alipay_direct_checkout, create_wxpay_direct_checkout, normalize_payment_currency, + normalize_stripe_client_secret, open_payment_order_stripe_client_secret, + public_payment_http_client, seal_payment_order_stripe_client_secret, + DirectPaymentCheckoutError, DirectPaymentCheckoutInput, PaymentOrderStripeSecretBinding, + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY, }; use chrono::Utc; use serde::Deserialize; use serde_json::{json, Value}; +use sha2::{Digest, Sha256}; use uuid::Uuid; +use aether_data::repository::wallet::{ + wallet_recharge_checkout_claim_response, wallet_recharge_order_created_at_unix_secs, + wallet_recharge_order_is_checkout_placeholder, + wallet_recharge_order_is_reclaimable_placeholder, FailWalletRechargeCheckoutInput, + ReclaimWalletRechargeCheckoutInput, +}; + #[derive(Debug, Deserialize)] struct WalletCreateRechargeRequest { amount_usd: f64, @@ -40,6 +107,8 @@ struct WalletCreateRechargeRequest { pay_currency: Option, #[serde(default)] exchange_rate: Option, + #[serde(default)] + idempotency_key: Option, } #[derive(Debug, Clone)] @@ -51,6 +120,7 @@ struct NormalizedWalletCreateRechargeRequest { pay_amount: Option, pay_currency: Option, exchange_rate: Option, + idempotency_key: Option, } fn normalize_wallet_create_recharge_request( @@ -67,6 +137,29 @@ fn normalize_wallet_create_recharge_request( .map(|value| value.to_ascii_lowercase()); let payment_channel = wallet_normalize_optional_string_field(payload.payment_channel, 30)? .map(|value| value.to_ascii_lowercase()); + match payment_provider.as_deref() { + Some("epay") => { + // EPay is an aggregator. Older clients used alipay/wxpay as the + // method to select the channel, while newer clients send epay for + // both method and provider and carry the channel separately. + if !matches!(payment_method.as_str(), "epay" | "alipay" | "wxpay") { + return Err("输入验证失败"); + } + if payment_method != "epay" + && payment_channel + .as_deref() + .is_some_and(|channel| channel != payment_method) + { + return Err("输入验证失败"); + } + } + Some(provider) if provider != payment_method => { + // Do not let a request select one gateway while recording another + // payment namespace. + return Err("输入验证失败"); + } + _ => {} + } if matches!(payload.pay_amount, Some(value) if !value.is_finite() || value <= 0.0) { return Err("输入验证失败"); } @@ -86,6 +179,7 @@ fn normalize_wallet_create_recharge_request( pay_amount: payload.pay_amount, pay_currency, exchange_rate: payload.exchange_rate, + idempotency_key: wallet_normalize_optional_string_field(payload.idempotency_key, 128)?, }) } @@ -97,6 +191,32 @@ fn wallet_build_order_no(now: chrono::DateTime) -> String { ) } +fn wallet_build_idempotent_order_no(user_id: &str, idempotency_key: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(b"wallet-recharge:"); + hasher.update(user_id.trim().as_bytes()); + hasher.update([0]); + hasher.update(idempotency_key.trim().as_bytes()); + let digest = hasher.finalize(); + // payment_orders.order_no is limited to 64 bytes. Keep the stable prefix + // and 56 hex characters (28 digest bytes) within that database contract. + let digest_hex: String = digest[..28] + .iter() + .map(|byte| format!("{byte:02x}")) + .collect(); + format!("po_idem_{digest_hex}") +} + +fn wallet_recharge_order_no( + user_id: &str, + idempotency_key: Option<&str>, + now: chrono::DateTime, +) -> String { + idempotency_key + .map(|key| wallet_build_idempotent_order_no(user_id, key)) + .unwrap_or_else(|| wallet_build_order_no(now)) +} + fn wallet_payment_return_url(callback_base_url: &str, provider: &str, order_no: &str) -> String { let mut serializer = url::form_urlencoded::Serializer::new(String::new()); serializer.append_pair("payment_provider", provider); @@ -122,6 +242,63 @@ pub(super) fn wallet_recharge_detail_path_matches(request_path: &str) -> bool { wallet_order_id_from_path(request_path).is_some() } +fn remove_stripe_client_secret_fields(value: &mut Value) { + match value { + Value::Object(object) => { + object.remove("client_secret"); + object.remove(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY); + for item in object.values_mut() { + remove_stripe_client_secret_fields(item); + } + } + Value::Array(items) => { + for item in items { + remove_stripe_client_secret_fields(item); + } + } + _ => {} + } +} + +fn sanitize_wallet_payment_params(value: &Value) -> Option { + let object = value.as_object()?; + let mut projected = serde_json::Map::new(); + for key in WALLET_SAFE_EPAY_PAYMENT_PARAM_KEYS { + let Some(item) = object.get(*key) else { + continue; + }; + // EPay signs and submits string parameters. Reject objects/arrays here + // so a provider response cannot smuggle an arbitrary nested payload + // through an otherwise allow-listed field. + if item.is_string() { + projected.insert((*key).to_string(), item.clone()); + } + } + (!projected.is_empty()).then_some(Value::Object(projected)) +} + +fn sanitize_wallet_gateway_value(key: &str, value: &Value) -> Option { + match key { + "payment_params" => sanitize_wallet_payment_params(value), + "payment_method_types" => { + let values = value.as_array()?; + let projected = values + .iter() + .filter_map(Value::as_str) + .filter(|item| !item.trim().is_empty()) + .map(ToOwned::to_owned) + .map(Value::String) + .collect::>(); + (!projected.is_empty()).then_some(Value::Array(projected)) + } + // All other public checkout fields are scalar values in the adapter + // contract. Do not retain an object/array supplied in a future + // response under one of those names. + _ if value.is_object() || value.is_array() => None, + _ => Some(value.clone()), + } +} + pub(crate) fn sanitize_wallet_gateway_response( value: Option, ) -> serde_json::Value { @@ -134,10 +311,331 @@ pub(crate) fn sanitize_wallet_gateway_response( let mut sanitized = serde_json::Map::new(); for key in WALLET_SAFE_GATEWAY_RESPONSE_KEYS { if let Some(item) = object.get(*key) { - sanitized.insert((*key).to_string(), item.clone()); + if let Some(item) = sanitize_wallet_gateway_value(key, item) { + sanitized.insert((*key).to_string(), item); + } } } - serde_json::Value::Object(sanitized) + let mut sanitized = serde_json::Value::Object(sanitized); + remove_stripe_client_secret_fields(&mut sanitized); + sanitized +} + +fn valid_stripe_client_secret(value: &str) -> Option<&str> { + normalize_stripe_client_secret(value) +} + +fn insert_stripe_client_secret(instructions: &mut Value, client_secret: &str) { + if let Some(object) = instructions.as_object_mut() { + object.insert( + "client_secret".to_string(), + Value::String(client_secret.to_string()), + ); + } +} + +pub(crate) fn prepare_wallet_gateway_response_for_storage( + state: &AppState, + payment_provider: &str, + order_no: &str, + user_id: &str, + checkout: &Value, +) -> Result { + let binding = payment_provider + .trim() + .eq_ignore_ascii_case("stripe") + .then(|| { + PaymentOrderStripeSecretBinding::new( + order_no, + Some(user_id), + WALLET_RECHARGE_ORDER_KIND, + payment_provider, + ) + }) + .transpose() + .map_err(|error| error.to_string())?; + prepare_gateway_response_for_storage_with_encrypt( + payment_provider, + checkout, + true, + |client_secret| { + binding.as_ref().and_then(|binding| { + seal_payment_order_stripe_client_secret(state, binding, client_secret).ok() + }) + }, + ) +} + +/// Prepare a provider response for a non-wallet payment order (for example a +/// plan purchase). Keep the same credential stripping/encryption guarantees as +/// wallet checkout storage, but do not stamp the response with the wallet +/// recharge discriminator. +pub(crate) fn prepare_billing_gateway_response_for_storage( + state: &AppState, + payment_provider: &str, + order_no: &str, + user_id: &str, + checkout: &Value, +) -> Result { + let binding = payment_provider + .trim() + .eq_ignore_ascii_case("stripe") + .then(|| { + PaymentOrderStripeSecretBinding::new( + order_no, + Some(user_id), + "plan_purchase", + payment_provider, + ) + }) + .transpose() + .map_err(|error| error.to_string())?; + prepare_gateway_response_for_storage_with_encrypt( + payment_provider, + checkout, + false, + |client_secret| { + binding.as_ref().and_then(|binding| { + seal_payment_order_stripe_client_secret(state, binding, client_secret).ok() + }) + }, + ) +} + +fn prepare_wallet_gateway_response_for_storage_with_encrypt( + payment_provider: &str, + checkout: &Value, + encrypt_secret: impl FnOnce(&str) -> Option, +) -> Result { + prepare_gateway_response_for_storage_with_encrypt( + payment_provider, + checkout, + true, + encrypt_secret, + ) +} + +fn prepare_gateway_response_for_storage_with_encrypt( + payment_provider: &str, + checkout: &Value, + include_wallet_order_kind: bool, + encrypt_secret: impl FnOnce(&str) -> Option, +) -> Result { + if !checkout.is_object() { + return Err("支付网关响应格式无效".to_string()); + } + let client_secret = checkout + .get("client_secret") + .and_then(Value::as_str) + .and_then(valid_stripe_client_secret) + .map(ToOwned::to_owned); + // Persist only the same allow-listed checkout fields exposed by the + // public wallet API. Gateway responses are an adapter boundary, so a + // future provider integration must not be able to write arbitrary + // response fields (or credentials) into payment_orders. + let mut stored = sanitize_wallet_gateway_response(Some(checkout.clone())); + remove_stripe_client_secret_fields(&mut stored); + + if include_wallet_order_kind { + if let Some(object) = stored.as_object_mut() { + object.insert( + "order_kind".to_string(), + Value::String(WALLET_RECHARGE_ORDER_KIND.to_string()), + ); + } + } + + // Claim metadata is generated by this handler and is needed by the + // repository to reject stale updates. Keep it in the persisted projection, + // while leaving it out of the public response projection above. + if let Some(token) = checkout + .get("checkout_claim_token") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty() && value.len() <= 128) + { + if let Some(object) = stored.as_object_mut() { + object.insert( + "checkout_claim_token".to_string(), + Value::String(token.to_string()), + ); + if let Some(claimed_at) = checkout + .get("checkout_claimed_at_unix_secs") + .and_then(Value::as_u64) + { + object.insert( + "checkout_claimed_at_unix_secs".to_string(), + Value::Number(serde_json::Number::from(claimed_at)), + ); + } + } + } + + if !payment_provider.trim().eq_ignore_ascii_case("stripe") { + return Ok(stored); + } + let client_secret = client_secret.ok_or_else(|| "Stripe client_secret 无效".to_string())?; + let encrypted = encrypt_secret(&client_secret) + .ok_or_else(|| "Stripe client_secret 加密失败".to_string())?; + let Some(object) = stored.as_object_mut() else { + return Err("支付网关响应格式无效".to_string()); + }; + object.insert( + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY.to_string(), + Value::String(encrypted), + ); + Ok(stored) +} + +pub(crate) fn wallet_payment_instructions_from_checkout( + payment_provider: &str, + checkout: &Value, +) -> Value { + let mut instructions = sanitize_wallet_gateway_response(Some(checkout.clone())); + if payment_provider == "stripe" { + if let Some(client_secret) = checkout + .get("client_secret") + .and_then(Value::as_str) + .and_then(valid_stripe_client_secret) + { + insert_stripe_client_secret(&mut instructions, client_secret); + } + } + instructions +} + +fn payment_order_stripe_secret_identity_matches( + expected: &aether_data::repository::wallet::StoredAdminPaymentOrder, + actual: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> bool { + expected.id == actual.id + && expected.order_no == actual.order_no + && expected.wallet_id == actual.wallet_id + && expected.user_id == actual.user_id + && expected.payment_method == actual.payment_method + && expected.payment_provider == actual.payment_provider + && expected.order_kind == actual.order_kind + && expected.gateway_order_id == actual.gateway_order_id +} + +fn payment_order_is_live_pending( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> bool { + let now_unix_secs = Utc::now().timestamp().max(0) as u64; + order.status.trim().eq_ignore_ascii_case("pending") + && order + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at > now_unix_secs) +} + +async fn decrypt_or_migrate_payment_order_stripe_client_secret( + state: &AppState, + initial: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> Result { + let identity = initial.clone(); + let mut current = initial.clone(); + + for _ in 0..PAYMENT_ORDER_STRIPE_SECRET_MIGRATION_RETRIES { + if !payment_order_stripe_secret_identity_matches(&identity, ¤t) { + return Err( + "payment order Stripe secret identity changed during migration".to_string(), + ); + } + if !payment_order_is_live_pending(¤t) { + return Err("payment order is no longer live and pending".to_string()); + } + let gateway_response = current + .gateway_response + .clone() + .ok_or_else(|| "payment order gateway response is unavailable".to_string())?; + let observed = gateway_response + .as_object() + .and_then(|object| object.get(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY)) + .and_then(Value::as_str) + .ok_or_else(|| "payment order Stripe secret ciphertext is unavailable".to_string())? + .to_string(); + let binding = PaymentOrderStripeSecretBinding::from_order(¤t) + .map_err(|error| error.to_string())?; + let projection = open_payment_order_stripe_client_secret(state, &binding, &observed) + .map_err(|error| error.to_string())?; + if !projection.migration_required { + return Ok(projection.plaintext); + } + + let mutation = + aether_data::repository::wallet::CompareAndSwapPaymentOrderStripeClientSecretInput { + order_id: current.id.clone(), + order_no: current.order_no.clone(), + wallet_id: current.wallet_id.clone(), + user_id: current.user_id.clone(), + payment_method: current.payment_method.clone(), + payment_provider: current.payment_provider.clone(), + order_kind: current.order_kind.clone(), + gateway_order_id: current.gateway_order_id.clone(), + expected_status: current.status.clone(), + expected_expires_at_unix_secs: current.expires_at_unix_secs, + expected_gateway_response: gateway_response, + expected_client_secret_encrypted: observed, + replacement_client_secret_encrypted: projection.protected, + }; + match state + .compare_and_swap_payment_order_stripe_client_secret(mutation) + .await + .map_err(|error| format!("payment order Stripe secret migration failed: {error:?}"))? + { + Some(true) => return Ok(projection.plaintext), + Some(false) => {} + None => { + return Err( + "payment order Stripe secret migration storage is unavailable".to_string(), + ) + } + } + + current = state + .find_payment_order_by_id(&identity.id) + .await + .map_err(|error| format!("payment order Stripe secret reread failed: {error:?}"))? + .ok_or_else(|| { + "payment order disappeared during Stripe secret migration".to_string() + })?; + } + + Err("payment order Stripe secret migration did not stabilize".to_string()) +} + +pub(crate) async fn wallet_payment_instructions_from_stored( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> Value { + if !payment_order_is_live_pending(order) { + return json!({}); + } + + let mut instructions = sanitize_wallet_gateway_response(order.gateway_response.clone()); + let payment_provider = order + .payment_provider + .as_deref() + .unwrap_or(order.payment_method.as_str()); + if !payment_provider.trim().eq_ignore_ascii_case("stripe") { + return instructions; + } + let client_secret = + match decrypt_or_migrate_payment_order_stripe_client_secret(state, order).await { + Ok(value) => value, + Err(error) => { + tracing::warn!( + order_id = %order.id, + order_no = %order.order_no, + error_category = "stripe_client_secret_open_failed", + reason = %error, + "payment order Stripe client secret was withheld" + ); + return instructions; + } + }; + insert_stripe_client_secret(&mut instructions, &client_secret); + instructions } fn build_wallet_payment_order_payload( @@ -185,6 +683,12 @@ fn build_wallet_payment_order_payload( fn wallet_payment_order_payload_from_record( record: &aether_data::repository::wallet::StoredAdminPaymentOrder, ) -> serde_json::Value { + // The order history is durable, but checkout capabilities are not. An + // expired, paid, or credited order must not keep advertising a URL/form + // that the callback path will no longer accept. + let gateway_response = wallet_recharge_order_is_pending_and_live(record) + .then(|| record.gateway_response.clone()) + .flatten(); build_wallet_payment_order_payload( record.id.clone(), record.order_no.clone(), @@ -198,15 +702,451 @@ fn wallet_payment_order_payload_from_record( record.refundable_amount_usd, record.payment_method.clone(), record.gateway_order_id.clone(), - record.gateway_response.clone(), + gateway_response, record.status.clone(), - Some(unix_secs_to_rfc3339(record.created_at_unix_ms)).flatten(), + unix_secs_to_rfc3339(wallet_recharge_order_created_at_unix_secs(record)), record.paid_at_unix_secs.and_then(unix_secs_to_rfc3339), record.credited_at_unix_secs.and_then(unix_secs_to_rfc3339), record.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), ) } +fn wallet_recharge_order_has_checkout( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> bool { + order + .gateway_response + .as_ref() + .and_then(Value::as_object) + .is_some_and(|object| { + object.keys().any(|key| { + matches!( + key.as_str(), + "payment_url" + | "payment_params" + | "qr_code" + | "code_url" + | "h5_url" + | "jsapi" + | "client_secret" + | STRIPE_CLIENT_SECRET_ENCRYPTED_KEY + | "intent_id" + ) + }) + }) +} + +fn wallet_recharge_order_is_pending_and_live( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> bool { + order.status.eq_ignore_ascii_case("pending") + && order + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at > Utc::now().timestamp().max(0) as u64) +} + +fn wallet_recharge_exchange_rates_match( + pay_currency: Option<&str>, + stored_rate: Option, + requested_rate: f64, +) -> bool { + let Some(pay_currency) = pay_currency else { + return false; + }; + let Some(stored_rate) = stored_rate else { + return false; + }; + let Ok(stored_rate) = + crate::handlers::shared::effective_payment_exchange_rate(pay_currency, stored_rate) + else { + return false; + }; + let Ok(requested_rate) = + crate::handlers::shared::effective_payment_exchange_rate(pay_currency, requested_rate) + else { + return false; + }; + (stored_rate - requested_rate).abs() <= PAYMENT_AMOUNT_EPSILON +} + +fn wallet_recharge_order_matches_request( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + payload: &NormalizedWalletCreateRechargeRequest, +) -> bool { + let amount_matches = order.amount_usd.is_finite() + && (order.amount_usd - payload.amount_usd).abs() <= PAYMENT_AMOUNT_EPSILON; + let requested_provider = payload + .payment_provider + .as_deref() + .unwrap_or(payload.payment_method.as_str()); + let stored_provider = order + .gateway_response + .as_ref() + .and_then(Value::as_object) + .and_then(|object| object.get("gateway")) + .and_then(Value::as_str); + // EPay exposes multiple channels behind one provider. When callers use + // the legacy shorthand (`payment_method: alipay|wxpay`) without an + // explicit payment_channel, that method is still part of the idempotent + // request identity and must not replay another channel's order. + let requested_channel = payload.payment_channel.as_deref().or_else(|| { + (requested_provider.eq_ignore_ascii_case("epay") + && !payload.payment_method.eq_ignore_ascii_case("epay")) + .then_some(payload.payment_method.as_str()) + }); + let metadata = order.gateway_response.as_ref().and_then(Value::as_object); + let stored_channel = metadata + .and_then(|object| object.get("payment_channel")) + .and_then(Value::as_str); + let payment_identity_matches = wallet_recharge_payment_identity_matches( + &order.payment_method, + stored_provider, + stored_channel, + requested_provider, + requested_channel, + ); + let pay_amount_matches = payload.pay_amount.is_none_or(|value| { + order.pay_amount.is_some_and(|stored| { + stored.is_finite() && (stored - value).abs() <= PAYMENT_AMOUNT_EPSILON + }) + }); + let currency_matches = payload.pay_currency.as_deref().is_none_or(|currency| { + order + .pay_currency + .as_deref() + .is_some_and(|stored| stored.eq_ignore_ascii_case(currency)) + }); + let exchange_rate_matches = payload.exchange_rate.is_none_or(|value| { + wallet_recharge_exchange_rates_match( + order.pay_currency.as_deref(), + order.exchange_rate, + value, + ) + }); + amount_matches + && payment_identity_matches + && pay_amount_matches + && currency_matches + && exchange_rate_matches +} + +fn wallet_recharge_payment_identity_matches( + stored_method: &str, + stored_provider: Option<&str>, + stored_channel: Option<&str>, + requested_provider: &str, + requested_channel: Option<&str>, +) -> bool { + let legacy_epay_method = stored_provider + .is_some_and(|provider| provider.eq_ignore_ascii_case("epay")) + && ["alipay", "wxpay"] + .iter() + .any(|method| stored_method.eq_ignore_ascii_case(method)); + let method_matches = stored_method.eq_ignore_ascii_case(requested_provider) + || (requested_provider.eq_ignore_ascii_case("epay") && legacy_epay_method); + let provider_matches = stored_provider.map_or_else( + || stored_method.eq_ignore_ascii_case(requested_provider), + |provider| provider.eq_ignore_ascii_case(requested_provider), + ); + let effective_stored_channel = + stored_channel.or_else(|| legacy_epay_method.then_some(stored_method)); + let channel_matches = requested_channel.is_none_or(|channel| { + effective_stored_channel.is_some_and(|stored| stored.eq_ignore_ascii_case(channel)) + }); + method_matches && provider_matches && channel_matches +} + +fn wallet_recharge_order_matches_effective_request( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + provider: &str, + channel: &str, + amount_usd: f64, + pay_amount: f64, + pay_currency: &str, + exchange_rate: f64, +) -> bool { + if !order.amount_usd.is_finite() + || !amount_usd.is_finite() + || (order.amount_usd - amount_usd).abs() > PAYMENT_AMOUNT_EPSILON + || !order.pay_amount.is_some_and(|value| { + value.is_finite() && (value - pay_amount).abs() <= PAYMENT_AMOUNT_EPSILON + }) + || !order + .pay_currency + .as_deref() + .is_some_and(|value| value.eq_ignore_ascii_case(pay_currency)) + || !wallet_recharge_exchange_rates_match( + order.pay_currency.as_deref(), + order.exchange_rate, + exchange_rate, + ) + { + return false; + } + let Some(metadata) = order.gateway_response.as_ref().and_then(Value::as_object) else { + return false; + }; + let stored_provider = metadata.get("gateway").and_then(Value::as_str); + let stored_channel = metadata.get("payment_channel").and_then(Value::as_str); + wallet_recharge_payment_identity_matches( + &order.payment_method, + stored_provider, + stored_channel, + provider, + Some(channel), + ) +} + +fn wallet_test_recharge_payload_matches_request( + order: &Value, + payload: &NormalizedWalletCreateRechargeRequest, +) -> bool { + let amount_matches = order + .get("amount_usd") + .and_then(Value::as_f64) + .is_some_and(|value| { + value.is_finite() && (value - payload.amount_usd).abs() <= PAYMENT_AMOUNT_EPSILON + }); + let stored_method = order.get("payment_method").and_then(Value::as_str); + let stored_provider = order + .get("payment_provider") + .and_then(Value::as_str) + .or_else(|| { + order + .get("gateway_response") + .and_then(Value::as_object) + .and_then(|object| object.get("gateway")) + .and_then(Value::as_str) + }); + let requested_provider = payload + .payment_provider + .as_deref() + .unwrap_or(payload.payment_method.as_str()); + let requested_channel = payload.payment_channel.as_deref().or_else(|| { + (requested_provider.eq_ignore_ascii_case("epay") + && !payload.payment_method.eq_ignore_ascii_case("epay")) + .then_some(payload.payment_method.as_str()) + }); + let metadata = order.get("gateway_response").and_then(Value::as_object); + let stored_channel = order + .get("payment_channel") + .and_then(Value::as_str) + .or_else(|| { + metadata + .and_then(|object| object.get("payment_channel")) + .and_then(Value::as_str) + }); + let payment_identity_matches = stored_method.is_some_and(|stored_method| { + wallet_recharge_payment_identity_matches( + stored_method, + stored_provider, + stored_channel, + requested_provider, + requested_channel, + ) + }); + let pay_amount_matches = payload.pay_amount.is_none_or(|value| { + order + .get("pay_amount") + .and_then(Value::as_f64) + .is_some_and(|stored| { + stored.is_finite() && (stored - value).abs() <= PAYMENT_AMOUNT_EPSILON + }) + }); + let currency_matches = payload.pay_currency.as_deref().is_none_or(|currency| { + order + .get("pay_currency") + .and_then(Value::as_str) + .is_some_and(|stored| stored.eq_ignore_ascii_case(currency)) + }); + let exchange_rate_matches = payload.exchange_rate.is_none_or(|value| { + wallet_recharge_exchange_rates_match( + order.get("pay_currency").and_then(Value::as_str), + order.get("exchange_rate").and_then(Value::as_f64), + value, + ) + }); + amount_matches + && payment_identity_matches + && pay_amount_matches + && currency_matches + && exchange_rate_matches +} + +fn wallet_test_recharge_replay_payment_instructions(order: &Value) -> Value { + let is_pending_and_live = order + .get("status") + .and_then(Value::as_str) + .is_some_and(|status| status.eq_ignore_ascii_case("pending")) + && order + .get("expires_at") + .and_then(Value::as_str) + .and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok()) + .is_some_and(|expires_at| expires_at.timestamp() > Utc::now().timestamp()); + if !is_pending_and_live { + return json!({}); + } + sanitize_wallet_gateway_response(order.get("gateway_response").cloned()) +} + +#[cfg(test)] +fn wallet_test_recharge_public_payload(mut order: Value) -> Value { + let is_pending_and_live = order + .get("status") + .and_then(Value::as_str) + .is_some_and(|status| status.eq_ignore_ascii_case("pending")) + && order + .get("expires_at") + .and_then(Value::as_str) + .and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok()) + .is_some_and(|expires_at| expires_at.timestamp() > Utc::now().timestamp()); + if !is_pending_and_live { + if let Some(object) = order.as_object_mut() { + object.insert("gateway_response".to_string(), json!({})); + } + } else if let Some(gateway_response) = order.get("gateway_response").cloned() { + order["gateway_response"] = sanitize_wallet_gateway_response(Some(gateway_response)); + } + order +} + +async fn wallet_recharge_replay_payment_instructions( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> Value { + // A replayed idempotency response must not hand out a stale payment URL. + // Providers can reject an expired order, but returning the URL still + // invites a user to submit a payment that cannot be credited safely. + if !wallet_recharge_order_is_pending_and_live(order) { + return json!({}); + } + wallet_payment_instructions_from_stored(state, order).await +} + +async fn wallet_recharge_replay_response( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> Response { + let payment_instructions = wallet_recharge_replay_payment_instructions(state, order).await; + mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": wallet_payment_order_payload_from_record(order), + "payment_instructions": payment_instructions, + "reused_idempotent_order": true, + }), + None, + )) +} + +/// A provider checkout can finish after its callback has already credited the +/// order. In that case the conditional checkout update reports a conflict, +/// but the client must observe the durable settled order instead of receiving a +/// misleading checkout error (or a second payment capability). +async fn settled_wallet_recharge_after_checkout_conflict( + state: &AppState, + user_id: &str, + order_id: &str, +) -> Option { + match state + .find_wallet_payment_order_by_user_id(user_id, order_id) + .await + { + Ok(Some(order)) if wallet_recharge_order_is_settled(&order) => Some(order), + _ => None, + } +} + +fn wallet_recharge_order_is_settled( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, +) -> bool { + matches!(order.status.as_str(), "paid" | "credited") +} + +fn wallet_recharge_claim_token() -> String { + format!("wrc_{}", Uuid::new_v4().simple()) +} + +fn wallet_recharge_claimed_placeholder( + value: &Value, + claim_token: &str, + claimed_at_unix_secs: u64, +) -> Result { + wallet_recharge_checkout_claim_response(value, claim_token, claimed_at_unix_secs) +} + +async fn best_effort_fail_wallet_recharge_checkout( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + claim_token: &str, + reason: &str, +) { + let _ = state + .fail_wallet_recharge_checkout(FailWalletRechargeCheckoutInput { + order_id: order.id.clone(), + claim_token: claim_token.to_string(), + reason: reason.to_string(), + provider_request_may_have_succeeded: false, + }) + .await; +} + +async fn best_effort_mark_wallet_recharge_checkout_uncertain( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + claim_token: &str, + reason: &str, +) { + let _ = state + .fail_wallet_recharge_checkout(FailWalletRechargeCheckoutInput { + order_id: order.id.clone(), + claim_token: claim_token.to_string(), + reason: reason.to_string(), + provider_request_may_have_succeeded: true, + }) + .await; +} + +async fn reclaim_wallet_recharge_checkout( + state: &AppState, + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + placeholder: Value, + claim_token: &str, + expires_at_unix_secs: u64, +) -> Result { + match state + .reclaim_wallet_recharge_checkout(ReclaimWalletRechargeCheckoutInput { + order_id: order.id.clone(), + claim_token: claim_token.to_string(), + gateway_response: placeholder, + expires_at_unix_secs, + }) + .await + { + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(order))) => { + Ok(order) + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound)) => { + Err("充值订单已不存在".to_string()) + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid(detail))) => { + Err(detail) + } + Ok(None) => Err("钱包充值后端暂不可用".to_string()), + Err(err) => Err(format!("wallet recharge checkout reclaim failed: {err:?}")), + } +} + +fn attach_wallet_recharge_claim_token(mut checkout: Value, claim_token: &str) -> Value { + if let Some(object) = checkout.as_object_mut() { + object.insert( + "checkout_claim_token".to_string(), + Value::String(claim_token.to_string()), + ); + } + checkout +} + #[derive(Debug, Clone)] pub(crate) struct DirectGatewayChannelConfig { pub(crate) channel: String, @@ -227,24 +1167,49 @@ fn configured_channel_fee_rate(value: Option<&Value>) -> f64 { } } -fn round_payment_amount(value: f64) -> f64 { - (value * 100.0).round() / 100.0 +fn round_payment_amount(value: f64) -> Option { + if !value.is_finite() { + return None; + } + let rounded = (value * 100.0).round() / 100.0; + rounded.is_finite().then_some(rounded) +} + +#[derive(Debug, Clone, Copy, PartialEq)] +struct WalletRechargePaymentBreakdown { + base_pay_amount: f64, + fee_amount: f64, + pay_amount: f64, + exchange_rate: f64, } fn wallet_recharge_payment_breakdown( amount_usd: f64, + pay_currency: &str, usd_exchange_rate: f64, fee_rate: f64, -) -> (f64, f64, f64) { - let safe_fee_rate = if fee_rate.is_finite() && fee_rate > 0.0 { - fee_rate - } else { - 0.0 - }; - let base_pay_amount = round_payment_amount(amount_usd * usd_exchange_rate); - let fee_amount = round_payment_amount(base_pay_amount * safe_fee_rate / 100.0); - let pay_amount = round_payment_amount(base_pay_amount + fee_amount); - (base_pay_amount, fee_amount, pay_amount) +) -> Result { + if !amount_usd.is_finite() || amount_usd <= 0.0 || !fee_rate.is_finite() || fee_rate < 0.0 { + return Err("充值金额配置无效"); + } + let exchange_rate = + crate::handlers::shared::effective_payment_exchange_rate(pay_currency, usd_exchange_rate) + .map_err(|_| "充值金额配置无效")?; + let base_pay_amount = round_payment_amount(amount_usd * exchange_rate) + .filter(|value| *value > 0.0) + .ok_or("充值金额配置无效")?; + let fee_amount = round_payment_amount(base_pay_amount * fee_rate / 100.0) + .filter(|value| *value >= 0.0) + .ok_or("充值金额配置无效")?; + let pay_amount = round_payment_amount(base_pay_amount + fee_amount) + .filter(|value| *value > 0.0) + .ok_or("充值金额配置无效")?; + Ok(WalletRechargePaymentBreakdown { + base_pay_amount, + fee_amount, + pay_amount, + exchange_rate, + }) } fn add_wallet_recharge_fee_metadata( @@ -283,24 +1248,26 @@ pub(crate) fn direct_gateway_channels( .filter(|value| !value.is_empty()) .unwrap_or(channel_id); Some(DirectGatewayChannelConfig { - channel: channel_id.to_string(), + channel: channel_id.to_ascii_lowercase(), display_name: display_name.to_string(), fee_rate: configured_channel_fee_rate(channel.get("fee_rate")), }) }) - .filter(|channel| match provider { - "alipay" => channel.channel == "alipay", - "wxpay" => matches!(channel.channel.as_str(), "native" | "h5" | "jsapi"), - "stripe" => matches!( - channel.channel.as_str(), - "card" | "alipay" | "wechat_pay" | "link" - ), - _ => false, - }) + .filter( + |channel| match provider.trim().to_ascii_lowercase().as_str() { + "alipay" => channel.channel == "alipay", + "wxpay" => matches!(channel.channel.as_str(), "native" | "h5" | "jsapi"), + "stripe" => matches!( + channel.channel.as_str(), + "card" | "alipay" | "wechat_pay" | "link" + ), + _ => false, + }, + ) .collect() } -fn resolve_direct_gateway_channel( +pub(crate) fn resolve_direct_gateway_channel( provider: &str, record: &aether_data_contracts::repository::billing::PaymentGatewayConfigRecord, requested: Option<&str>, @@ -338,12 +1305,12 @@ fn decrypt_direct_gateway_secrets( let Some(encrypted) = record.merchant_key_encrypted.as_deref() else { return Err("支付网关密钥未配置".to_string()); }; - let Some(plaintext) = crate::handlers::shared::decrypt_catalog_secret_with_fallbacks( - state.encryption_key(), - encrypted, - ) else { - return Err("支付网关密钥解密失败".to_string()); - }; + let binding = crate::handlers::shared::PaymentGatewaySecretBinding::from_record(record) + .map_err(|_| "支付网关密钥绑定无效".to_string())?; + let plaintext = + crate::handlers::shared::open_payment_gateway_secret(state, &binding, encrypted) + .map_err(|_| "支付网关密钥解密失败".to_string())? + .plaintext; serde_json::from_str::(&plaintext) .ok() .and_then(|value| value.as_object().cloned()) @@ -362,20 +1329,6 @@ fn direct_gateway_secret_string( .map(ToOwned::to_owned) } -fn stripe_minor_unit_amount(pay_amount: f64, pay_currency: &str) -> Result { - let currency = pay_currency.trim().to_ascii_lowercase(); - let multiplier = match currency.as_str() { - "bif" | "clp" | "djf" | "gnf" | "jpy" | "kmf" | "krw" | "mga" | "pyg" | "rwf" | "ugx" - | "vnd" | "vuv" | "xaf" | "xof" | "xpf" => 1.0, - _ => 100.0, - }; - let amount = (pay_amount * multiplier).round(); - if !amount.is_finite() || amount <= 0.0 { - return Err("Stripe 支付金额无效".to_string()); - } - Ok(amount as i64) -} - async fn create_stripe_wallet_recharge_checkout( state: &AppState, record: &aether_data_contracts::repository::billing::PaymentGatewayConfigRecord, @@ -384,16 +1337,18 @@ async fn create_stripe_wallet_recharge_checkout( order_no: &str, pay_amount: f64, expires_at: chrono::DateTime, -) -> Result { + idempotency_key: &str, +) -> Result { let secrets = decrypt_direct_gateway_secrets(state, record)?; let Some(secret_key) = direct_gateway_secret_string(&secrets, "secret_key") else { - return Err("Stripe secret_key 未配置".to_string()); + return Err("Stripe secret_key 未配置".to_string().into()); }; let Some(publishable_key) = direct_gateway_public_config_string(record, "publishable_key") else { - return Err("Stripe publishable_key 未配置".to_string()); + return Err("Stripe publishable_key 未配置".to_string().into()); }; - let amount = stripe_minor_unit_amount(pay_amount, &record.pay_currency)?; + let amount = crate::handlers::shared::stripe_amount_to_minor(pay_amount, &record.pay_currency) + .ok_or_else(|| "Stripe 支付金额无效".to_string())?; let currency = record.pay_currency.trim().to_ascii_lowercase(); let mut form = vec![ ("amount".to_string(), amount.to_string()), @@ -419,34 +1374,81 @@ async fn create_stripe_wallet_recharge_checkout( "web".to_string(), )); } - let response = state - .client - .post("https://api.stripe.com/v1/payment_intents") + let stripe_endpoint = + url::Url::parse("https://api.stripe.com/v1/payment_intents").map_err(|_| { + StripeWalletCheckoutError::Uncertain(STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string()) + })?; + let stripe_client = public_payment_http_client(&stripe_endpoint) + .await + .map_err(|_| { + StripeWalletCheckoutError::Uncertain(STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string()) + })?; + let response = match stripe_client + .post(stripe_endpoint) + .header("Idempotency-Key", idempotency_key) .basic_auth(secret_key, Some("")) .form(&form) .send() .await - .map_err(|err| format!("Stripe PaymentIntent 创建失败: {err}"))?; + { + Ok(response) => response, + Err(_) => { + tracing::warn!( + event_name = "stripe_wallet_payment_intent_request_failed", + "Stripe wallet PaymentIntent request failed" + ); + return Err(StripeWalletCheckoutError::Uncertain( + STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string(), + )); + } + }; let status = response.status(); - let body = response - .text() - .await - .map_err(|err| format!("Stripe 响应读取失败: {err}"))?; - let value = - serde_json::from_str::(&body).map_err(|_| "Stripe 响应格式无效".to_string())?; + let body = + aether_http::read_response_bytes_with_limit(response, MAX_PAYMENT_GATEWAY_RESPONSE_BYTES) + .await + .map_err(|_| { + tracing::warn!( + event_name = "stripe_wallet_payment_intent_response_read_failed", + upstream_status = %status, + "Stripe wallet PaymentIntent response could not be read" + ); + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + StripeWalletCheckoutError::Uncertain( + STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string(), + ) + } else { + StripeWalletCheckoutError::Failed(STRIPE_CHECKOUT_FAILED_DETAIL.to_string()) + } + })?; if !status.is_success() { - let message = value - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Stripe PaymentIntent 创建失败"); - return Err(message.to_string()); + tracing::warn!( + event_name = "stripe_wallet_payment_intent_upstream_rejected", + upstream_status = %status, + "Stripe wallet PaymentIntent request was rejected" + ); + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + return Err(StripeWalletCheckoutError::Uncertain( + STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string(), + )); + } + return Err(StripeWalletCheckoutError::Failed( + STRIPE_CHECKOUT_FAILED_DETAIL.to_string(), + )); + } + let value = serde_json::from_slice::(&body) + .map_err(|_| StripeWalletCheckoutError::Uncertain("Stripe 响应格式无效".to_string()))?; + if stripe_wallet_checkout_response_is_canceled(&value) { + return Err(StripeWalletCheckoutError::Canceled); } let Some(intent_id) = value.get("id").and_then(Value::as_str) else { - return Err("Stripe 响应缺少 PaymentIntent ID".to_string()); + return Err(StripeWalletCheckoutError::Uncertain( + "Stripe 响应缺少 PaymentIntent ID".to_string(), + )); }; let Some(client_secret) = value.get("client_secret").and_then(Value::as_str) else { - return Err("Stripe 响应缺少 client_secret".to_string()); + return Err(StripeWalletCheckoutError::Uncertain( + "Stripe 响应缺少 client_secret".to_string(), + )); }; Ok(json!({ "gateway": "stripe", @@ -468,6 +1470,7 @@ pub(super) async fn handle_wallet_create_recharge( state: &AppState, request_context: &GatewayPublicRequestContext, headers: &http::HeaderMap, + client_ip: std::net::IpAddr, request_body: Option<&axum::body::Bytes>, ) -> Response { let auth = match resolve_authenticated_local_user(state, request_context, headers).await { @@ -513,6 +1516,54 @@ pub(super) async fn handle_wallet_create_recharge( } }; + let now = Utc::now(); + let now_unix_secs = now.timestamp().max(0) as u64; + let order_no = wallet_recharge_order_no(&auth.user.id, payload.idempotency_key.as_deref(), now); + if payload.idempotency_key.is_some() { + match state + .find_wallet_recharge_order_by_order_no(&auth.user.id, &order_no) + .await + { + Ok(Some(order)) => { + if !wallet_recharge_order_matches_request(&order, &payload) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "idempotency_key 已用于其他充值订单", + false, + ); + } + let reclaimable = + wallet_recharge_order_is_reclaimable_placeholder(&order, now_unix_secs); + if wallet_recharge_order_has_checkout(&order) + || (!reclaimable && !wallet_recharge_order_is_pending_and_live(&order)) + { + return wallet_recharge_replay_response(state, &order).await; + } + if !reclaimable { + // A pending order without checkout data is an in-flight claim held by + // another request. Never call an external gateway from a retry: doing so + // would create duplicate provider-side orders while the first request is + // still waiting for its response. + return build_auth_error_response( + http::StatusCode::CONFLICT, + "充值订单正在创建,请稍后重试", + false, + ); + } + // The provider-specific branch below will atomically replace this + // placeholder after resolving the effective channel and amount. + } + Ok(None) => {} + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("wallet recharge idempotency lookup failed: {err:?}"), + false, + ) + } + } + } + if !state.has_database_wallet_data_writer() { #[cfg(test)] { @@ -530,9 +1581,31 @@ pub(super) async fn handle_wallet_create_recharge( false, ); } - let now = Utc::now(); + if let Some(idempotency_key) = payload.idempotency_key.as_deref() { + if let Some(existing) = + wallet_test_recharge_order_by_order_no(&auth.user.id, &order_no) + { + if !wallet_test_recharge_payload_matches_request(&existing, &payload) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "idempotency_key 已用于其他充值订单", + false, + ); + } + let instructions = wallet_test_recharge_replay_payment_instructions(&existing); + return mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": existing, + "payment_instructions": instructions, + "reused_idempotent_order": true, + }), + None, + )); + } + let _ = idempotency_key; + } let order_id = Uuid::new_v4().to_string(); - let order_no = wallet_build_order_no(now); let expires_at = now + chrono::Duration::minutes(30); let Some(adapter) = PaymentGatewayRegistry::get(&payload.payment_method) else { return build_auth_error_response( @@ -551,9 +1624,9 @@ pub(super) async fn handle_wallet_create_recharge( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } }; - let order_payload = build_wallet_payment_order_payload( + let mut order_payload = build_wallet_payment_order_payload( order_id, - order_no, + order_no.clone(), wallet.id.clone(), Some(auth.user.id.clone()), payload.amount_usd, @@ -562,7 +1635,7 @@ pub(super) async fn handle_wallet_create_recharge( payload.exchange_rate, 0.0, 0.0, - payload.payment_method, + payload.payment_method.clone(), Some(checkout.gateway_order_id.clone()), Some(checkout.gateway_response.clone()), "pending".to_string(), @@ -571,22 +1644,38 @@ pub(super) async fn handle_wallet_create_recharge( None, Some(expires_at.to_rfc3339()), ); - record_wallet_test_recharge(auth.user.id, order_payload.clone()); - return build_auth_json_response( + if let Some(object) = order_payload.as_object_mut() { + object.insert( + "payment_provider".to_string(), + Value::String(payload.payment_method.clone()), + ); + object.insert( + "payment_channel".to_string(), + payload + .payment_channel + .clone() + .map(Value::String) + .unwrap_or(Value::Null), + ); + object.insert( + "order_kind".to_string(), + Value::String(WALLET_RECHARGE_ORDER_KIND.to_string()), + ); + } + record_wallet_test_recharge(auth.user.id, order_no.clone(), order_payload.clone()); + return mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "order": order_payload, "payment_instructions": sanitize_wallet_gateway_response(Some(checkout.gateway_response)), }), None, - ); + )); } #[cfg(not(test))] return build_wallet_recharge_storage_unavailable_response(); } - let now = Utc::now(); - let order_no = wallet_build_order_no(now); let expires_at = now + chrono::Duration::minutes(30); let uses_epay = payload.payment_provider.as_deref() == Some("epay") || payload.payment_method == "epay"; @@ -613,96 +1702,336 @@ pub(super) async fn handle_wallet_create_recharge( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } }; - let (base_pay_amount, fee_amount, pay_amount) = wallet_recharge_payment_breakdown( + let WalletRechargePaymentBreakdown { + base_pay_amount, + fee_amount, + pay_amount, + exchange_rate, + } = match wallet_recharge_payment_breakdown( payload.amount_usd, + &config.pay_currency, config.usd_exchange_rate, payment_channel.fee_rate, - ); - let Some(callback_base_url) = epay_callback_base_url( - config.callback_base_url.as_deref(), - headers, - request_context, - ) else { + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } + }; + let Some(callback_base_url) = epay_callback_base_url(config.callback_base_url.as_deref()) + else { return build_auth_error_response( http::StatusCode::BAD_REQUEST, "epay callback_base_url is required", false, ); }; - let checkout = build_epay_checkout_url( + let payment_channel_id = payment_channel.channel.clone(); + let claim_token = wallet_recharge_claim_token(); + let expires_at_unix_secs = expires_at.timestamp().max(0) as u64; + let placeholder = json!({ + "gateway": "epay", + "gateway_order_id": order_no.clone(), + "order_kind": WALLET_RECHARGE_ORDER_KIND, + "payment_channel": payment_channel_id.clone(), + "pay_amount": pay_amount, + "pay_currency": config.pay_currency.clone(), + "exchange_rate": exchange_rate, + "integration_status": "checkout_pending", + }); + let placeholder = + match wallet_recharge_claimed_placeholder(&placeholder, &claim_token, now_unix_secs) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + let order_record = { + let outcome = match state + .create_wallet_recharge_order( + aether_data::repository::wallet::CreateWalletRechargeOrderInput { + preferred_wallet_id: wallet.as_ref().map(|value| value.id.clone()), + user_id: auth.user.id.clone(), + amount_usd: payload.amount_usd, + pay_amount: Some(pay_amount), + pay_currency: Some(config.pay_currency.clone()), + exchange_rate: Some(exchange_rate), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some(payment_channel_id.clone()), + gateway_order_id: order_no.clone(), + gateway_response: placeholder.clone(), + order_no: order_no.clone(), + expires_at_unix_secs, + }, + ) + .await + { + Ok(Some(value)) => value, + Ok(None) => return build_wallet_recharge_storage_unavailable_response(), + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("wallet recharge create failed: {err:?}"), + false, + ) + } + }; + match outcome { + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Created( + order, + ) => { + if wallet_recharge_order_has_checkout(&order) { + return wallet_recharge_replay_response(state, &order).await; + } + order + } + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Existing( + order, + ) => { + if !wallet_recharge_order_matches_effective_request( + &order, + "epay", + &payment_channel_id, + payload.amount_usd, + pay_amount, + &config.pay_currency, + exchange_rate, + ) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "幂等订单的支付参数已发生变化,请重新发起充值", + false, + ); + } + if wallet_recharge_order_has_checkout(&order) { + return wallet_recharge_replay_response(state, &order).await; + } + if !wallet_recharge_order_is_reclaimable_placeholder( + &order, + now_unix_secs, + ) { + if wallet_recharge_order_is_pending_and_live(&order) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "充值订单正在创建,请稍后重试", + false, + ); + } + return wallet_recharge_replay_response(state, &order).await; + } + match reclaim_wallet_recharge_checkout( + state, + &order, + placeholder.clone(), + &claim_token, + expires_at_unix_secs, + ) + .await + { + Ok(order) => order, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::CONFLICT, + detail, + false, + ) + } + } + } + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::WalletInactive => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "wallet is not active", + false, + ) + } + } + }; + if !wallet_recharge_order_matches_effective_request( + &order_record, + "epay", + &payment_channel_id, + payload.amount_usd, + pay_amount, + &config.pay_currency, + exchange_rate, + ) { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + "充值订单支付参数校验失败", + ) + .await; + return build_auth_error_response( + http::StatusCode::CONFLICT, + "幂等订单的支付参数已发生变化,请重新发起充值", + false, + ); + } + let checkout = match build_epay_checkout_url( &config, &EpayCheckoutInput { order_no: order_no.clone(), - channel: payment_channel.channel.clone(), + channel: payment_channel_id, subject: "钱包充值".to_string(), pay_amount, notify_url: format!("{callback_base_url}/api/payment/epay/notify"), return_url: format!("{callback_base_url}/api/payment/epay/return"), }, - ); + ) { + Ok(value) => value, + Err(detail) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); + } + }; let checkout = add_wallet_recharge_fee_metadata( checkout, base_pay_amount, payment_channel.fee_rate, fee_amount, ); - let outcome = match state - .create_wallet_recharge_order( - aether_data::repository::wallet::CreateWalletRechargeOrderInput { - preferred_wallet_id: wallet.as_ref().map(|value| value.id.clone()), - user_id: auth.user.id.clone(), - amount_usd: payload.amount_usd, - pay_amount: Some(pay_amount), - pay_currency: Some(config.pay_currency.clone()), - exchange_rate: Some(config.usd_exchange_rate), - payment_method: "epay".to_string(), - payment_provider: Some("epay".to_string()), - payment_channel: Some(payment_channel.channel), + let checkout_for_storage = + attach_wallet_recharge_claim_token(checkout.clone(), &claim_token); + let stored_gateway_response = match prepare_wallet_gateway_response_for_storage( + state, + "epay", + &order_no, + &auth.user.id, + &checkout_for_storage, + ) { + Ok(value) => value, + Err(detail) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ); + } + }; + let stored_order = match state + .update_wallet_recharge_checkout( + aether_data::repository::wallet::UpdateWalletRechargeCheckoutInput { + order_id: order_record.id.clone(), gateway_order_id: order_no.clone(), - gateway_response: checkout.clone(), - order_no, - expires_at_unix_secs: expires_at.timestamp().max(0) as u64, + gateway_response: stored_gateway_response, }, ) .await { - Ok(Some(value)) => value, - Ok(None) => return build_wallet_recharge_storage_unavailable_response(), + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(order))) => { + order + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound)) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + "充值订单已不存在", + ) + .await; + return build_auth_error_response( + http::StatusCode::CONFLICT, + "充值订单已不存在", + false, + ); + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid(detail))) => { + if let Some(order) = settled_wallet_recharge_after_checkout_conflict( + state, + &auth.user.id, + &order_record.id, + ) + .await + { + let order_payload = wallet_payment_order_payload_from_record(&order); + return mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": order_payload, + "payment_instructions": wallet_recharge_replay_payment_instructions( + state, &order + ).await, + }), + None, + )); + } + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response(http::StatusCode::CONFLICT, detail, false); + } + Ok(None) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + "钱包充值后端暂不可用", + ) + .await; + return build_wallet_recharge_storage_unavailable_response(); + } Err(err) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &format!("wallet recharge checkout update failed: {err:?}"), + ) + .await; return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, - format!("wallet recharge create failed: {err:?}"), + format!("wallet recharge checkout update failed: {err:?}"), false, - ) + ); } }; - let order_payload = match outcome { - aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Created(order) => { - wallet_payment_order_payload_from_record(&order) - } - aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::WalletInactive => { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "wallet is not active", - false, - ) - } - }; - return build_auth_json_response( + let order_payload = wallet_payment_order_payload_from_record(&stored_order); + // The provider call can race a payment callback. The repository may + // therefore return an already-credited order even though this request + // still has a checkout response in hand. Rebuild instructions from + // the final persisted state so settled/expired orders never receive a + // payment URL or client secret. + let payment_instructions = + wallet_recharge_replay_payment_instructions(state, &stored_order).await; + return mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "order": order_payload, - "payment_instructions": sanitize_wallet_gateway_response(Some(checkout)), + "payment_instructions": payment_instructions, }), None, - ); + )); } let requested_provider = payload .payment_provider .as_deref() .unwrap_or(payload.payment_method.as_str()); if matches!(requested_provider, "alipay" | "wxpay" | "stripe") { - let record = match state.find_payment_gateway_config(requested_provider).await { + let mut record = match state.find_payment_gateway_config(requested_provider).await { Ok(Some(value)) if value.enabled && value.merchant_key_encrypted.is_some() => value, Ok(Some(_)) | Ok(None) => { return build_auth_error_response( @@ -719,6 +2048,17 @@ pub(super) async fn handle_wallet_create_recharge( ) } }; + record.pay_currency = match normalize_payment_currency(&record.pay_currency, "pay_currency") + { + Ok(value) => value, + Err(_) => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "支付网关币种配置无效", + false, + ) + } + }; if payload.amount_usd < record.min_recharge_usd { return build_auth_error_response( http::StatusCode::BAD_REQUEST, @@ -736,34 +2076,250 @@ pub(super) async fn handle_wallet_create_recharge( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) } }; - let (base_pay_amount, fee_amount, pay_amount) = wallet_recharge_payment_breakdown( + let WalletRechargePaymentBreakdown { + base_pay_amount, + fee_amount, + pay_amount, + exchange_rate, + } = match wallet_recharge_payment_breakdown( payload.amount_usd, + &record.pay_currency, record.usd_exchange_rate, payment_channel.fee_rate, - ); - let checkout = if requested_provider == "stripe" { - match create_stripe_wallet_recharge_checkout( - state, - &record, - &payment_channel.channel, - &payment_channel.display_name, - &order_no, - pay_amount, - expires_at, - ) - .await - { + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } + }; + let payment_channel_id = payment_channel.channel.clone(); + let claim_token = wallet_recharge_claim_token(); + let expires_at_unix_secs = expires_at.timestamp().max(0) as u64; + let placeholder = json!({ + "gateway": requested_provider, + "gateway_order_id": order_no.clone(), + "order_kind": WALLET_RECHARGE_ORDER_KIND, + "payment_channel": payment_channel_id.clone(), + "pay_amount": pay_amount, + "pay_currency": record.pay_currency.clone(), + "exchange_rate": exchange_rate, + "integration_status": "checkout_pending", + }); + let placeholder = + match wallet_recharge_claimed_placeholder(&placeholder, &claim_token, now_unix_secs) { Ok(value) => value, Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false) + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ) + } + }; + let order_record = { + let outcome = match state + .create_wallet_recharge_order( + aether_data::repository::wallet::CreateWalletRechargeOrderInput { + preferred_wallet_id: wallet.as_ref().map(|value| value.id.clone()), + user_id: auth.user.id.clone(), + amount_usd: payload.amount_usd, + pay_amount: Some(pay_amount), + pay_currency: Some(record.pay_currency.clone()), + exchange_rate: Some(exchange_rate), + payment_method: requested_provider.to_string(), + payment_provider: Some(requested_provider.to_string()), + payment_channel: Some(payment_channel_id.clone()), + gateway_order_id: order_no.clone(), + gateway_response: placeholder.clone(), + order_no: order_no.clone(), + expires_at_unix_secs, + }, + ) + .await + { + Ok(Some(value)) => value, + Ok(None) => return build_wallet_recharge_storage_unavailable_response(), + Err(err) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("wallet recharge create failed: {err:?}"), + false, + ) + } + }; + match outcome { + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Created( + order, + ) => { + if wallet_recharge_order_has_checkout(&order) { + return wallet_recharge_replay_response(state, &order).await; + } + order + } + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Existing( + order, + ) => { + if !wallet_recharge_order_matches_effective_request( + &order, + requested_provider, + &payment_channel_id, + payload.amount_usd, + pay_amount, + &record.pay_currency, + exchange_rate, + ) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "幂等订单的支付参数已发生变化,请重新发起充值", + false, + ); + } + if wallet_recharge_order_has_checkout(&order) { + return wallet_recharge_replay_response(state, &order).await; + } + if !wallet_recharge_order_is_reclaimable_placeholder( + &order, + now_unix_secs, + ) { + if wallet_recharge_order_is_pending_and_live(&order) { + return build_auth_error_response( + http::StatusCode::CONFLICT, + "充值订单正在创建,请稍后重试", + false, + ); + } + return wallet_recharge_replay_response(state, &order).await; + } + match reclaim_wallet_recharge_checkout( + state, + &order, + placeholder.clone(), + &claim_token, + expires_at_unix_secs, + ) + .await + { + Ok(order) => order, + Err(detail) => { + return build_auth_error_response( + http::StatusCode::CONFLICT, + detail, + false, + ) + } + } + } + aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::WalletInactive => { + return build_auth_error_response( + http::StatusCode::BAD_REQUEST, + "wallet is not active", + false, + ) + } + } + }; + if !wallet_recharge_order_matches_effective_request( + &order_record, + requested_provider, + &payment_channel_id, + payload.amount_usd, + pay_amount, + &record.pay_currency, + exchange_rate, + ) { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + "充值订单支付参数校验失败", + ) + .await; + return build_auth_error_response( + http::StatusCode::CONFLICT, + "幂等订单的支付参数已发生变化,请重新发起充值", + false, + ); + } + let checkout = if requested_provider == "stripe" { + let mut retry_with_new_provider_key = false; + loop { + let idempotency_key = + stripe_wallet_idempotency_key(&order_no, retry_with_new_provider_key); + match create_stripe_wallet_recharge_checkout( + state, + &record, + &payment_channel_id, + &payment_channel.display_name, + &order_no, + pay_amount, + expires_at, + &idempotency_key, + ) + .await + { + Ok(value) => break value, + Err(StripeWalletCheckoutError::Canceled) if !retry_with_new_provider_key => { + // Stripe can retain a canceled PaymentIntent under an + // idempotency key. Rotate the provider key once while + // keeping the local order and metadata identity stable. + retry_with_new_provider_key = true; + } + Err(StripeWalletCheckoutError::Canceled) => { + let detail = "Stripe PaymentIntent 已取消,请稍后重试".to_string(); + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + detail, + false, + ); + } + Err(StripeWalletCheckoutError::Uncertain(detail)) => { + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + detail, + false, + ); + } + Err(StripeWalletCheckoutError::Failed(detail)) => { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response( + http::StatusCode::BAD_GATEWAY, + detail, + false, + ); + } } } } else { - let Some(callback_base_url) = epay_callback_base_url( - record.callback_base_url.as_deref(), - headers, - request_context, - ) else { + let Some(callback_base_url) = + epay_callback_base_url(record.callback_base_url.as_deref()) + else { + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + "支付网关 callback_base_url is required", + ) + .await; return build_auth_error_response( http::StatusCode::BAD_REQUEST, "支付网关 callback_base_url is required", @@ -771,7 +2327,7 @@ pub(super) async fn handle_wallet_create_recharge( ); }; let direct_input = DirectPaymentCheckoutInput { - payment_channel: payment_channel.channel.clone(), + payment_channel: payment_channel_id.clone(), display_name: payment_channel.display_name.clone(), order_no: order_no.clone(), subject: "钱包充值".to_string(), @@ -783,18 +2339,44 @@ pub(super) async fn handle_wallet_create_recharge( requested_provider, &order_no, )), - client_ip: direct_payment_client_ip(headers), + client_ip: Some(client_ip.to_string()), expires_at, }; let result = match requested_provider { "alipay" => create_alipay_direct_checkout(state, &direct_input).await, "wxpay" => create_wxpay_direct_checkout(state, &direct_input).await, - _ => Err("支付网关不支持".to_string()), + _ => Err(DirectPaymentCheckoutError::Failed( + "支付网关不支持".to_string(), + )), }; match result { Ok(value) => value, - Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false) + Err(DirectPaymentCheckoutError::Uncertain(detail)) => { + // The direct provider may have accepted the request before + // its response was lost. Keep the order non-reclaimable so + // a retry cannot create a second payment capability. + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false); + } + Err( + error @ (DirectPaymentCheckoutError::Canceled + | DirectPaymentCheckoutError::Failed(_)), + ) => { + let detail = error.into_detail(); + best_effort_fail_wallet_recharge_checkout( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response(http::StatusCode::BAD_GATEWAY, detail, false); } } }; @@ -809,56 +2391,130 @@ pub(super) async fn handle_wallet_create_recharge( .and_then(Value::as_str) .unwrap_or(&order_no) .to_string(); - let outcome = match state - .create_wallet_recharge_order( - aether_data::repository::wallet::CreateWalletRechargeOrderInput { - preferred_wallet_id: wallet.as_ref().map(|value| value.id.clone()), - user_id: auth.user.id.clone(), - amount_usd: payload.amount_usd, - pay_amount: Some(pay_amount), - pay_currency: Some(record.pay_currency.clone()), - exchange_rate: Some(record.usd_exchange_rate), - payment_method: requested_provider.to_string(), - payment_provider: Some(requested_provider.to_string()), - payment_channel: Some(payment_channel.channel), + let checkout_for_storage = + attach_wallet_recharge_claim_token(checkout.clone(), &claim_token); + let stored_gateway_response = match prepare_wallet_gateway_response_for_storage( + state, + requested_provider, + &order_no, + &auth.user.id, + &checkout_for_storage, + ) { + Ok(value) => value, + Err(detail) => { + // The provider has already returned a checkout response. A + // local projection/encryption failure therefore has an + // unknown provider outcome and must not permit a replacement + // checkout on retry. + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + detail, + false, + ); + } + }; + let stored_order = match state + .update_wallet_recharge_checkout( + aether_data::repository::wallet::UpdateWalletRechargeCheckoutInput { + order_id: order_record.id.clone(), gateway_order_id, - gateway_response: checkout.clone(), - order_no, - expires_at_unix_secs: expires_at.timestamp().max(0) as u64, + gateway_response: stored_gateway_response, }, ) .await { - Ok(Some(value)) => value, - Ok(None) => return build_wallet_recharge_storage_unavailable_response(), + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(order))) => { + order + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound)) => { + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + "充值订单已不存在", + ) + .await; + return build_auth_error_response( + http::StatusCode::CONFLICT, + "充值订单已不存在", + false, + ); + } + Ok(Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid(detail))) => { + if let Some(order) = settled_wallet_recharge_after_checkout_conflict( + state, + &auth.user.id, + &order_record.id, + ) + .await + { + let order_payload = wallet_payment_order_payload_from_record(&order); + return mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": order_payload, + "payment_instructions": wallet_recharge_replay_payment_instructions( + state, &order + ).await, + }), + None, + )); + } + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + &detail, + ) + .await; + return build_auth_error_response(http::StatusCode::CONFLICT, detail, false); + } + Ok(None) => { + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + "钱包充值后端暂不可用", + ) + .await; + return build_wallet_recharge_storage_unavailable_response(); + } Err(err) => { + best_effort_mark_wallet_recharge_checkout_uncertain( + state, + &order_record, + &claim_token, + &format!("wallet recharge checkout update failed: {err:?}"), + ) + .await; return build_auth_error_response( http::StatusCode::INTERNAL_SERVER_ERROR, - format!("wallet recharge create failed: {err:?}"), + format!("wallet recharge checkout update failed: {err:?}"), false, - ) + ); } }; - let order_payload = match outcome { - aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::Created(order) => { - wallet_payment_order_payload_from_record(&order) - } - aether_data::repository::wallet::CreateWalletRechargeOrderOutcome::WalletInactive => { - return build_auth_error_response( - http::StatusCode::BAD_REQUEST, - "wallet is not active", - false, - ) - } - }; - return build_auth_json_response( + let order_payload = wallet_payment_order_payload_from_record(&stored_order); + // See the EPay path above: a callback may have credited the order while + // the external checkout request was in flight. Only replay evidence + // from a still-live pending order. + let payment_instructions = + wallet_recharge_replay_payment_instructions(state, &stored_order).await; + return mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, json!({ "order": order_payload, - "payment_instructions": sanitize_wallet_gateway_response(Some(checkout)), + "payment_instructions": payment_instructions, }), None, - ); + )); } build_auth_error_response( http::StatusCode::BAD_REQUEST, @@ -894,12 +2550,24 @@ pub(super) async fn handle_wallet_recharge_options( } } for provider in ["alipay", "wxpay", "stripe"] { - let Ok(Some(record)) = state.find_payment_gateway_config(provider).await else { + let Ok(Some(mut record)) = state.find_payment_gateway_config(provider).await else { continue; }; if !record.enabled || record.merchant_key_encrypted.is_none() { continue; } + let Ok(pay_currency) = normalize_payment_currency(&record.pay_currency, "pay_currency") + else { + continue; + }; + record.pay_currency = pay_currency; + let Ok(exchange_rate) = crate::handlers::shared::effective_payment_exchange_rate( + &record.pay_currency, + record.usd_exchange_rate, + ) else { + continue; + }; + record.usd_exchange_rate = exchange_rate; for DirectGatewayChannelConfig { channel: payment_channel, display_name, @@ -981,7 +2649,14 @@ pub(super) async fn handle_wallet_recharge_list( #[cfg(test)] let (items, total) = if !state.has_database_wallet_data_writer() && items.is_empty() && total == 0 { - wallet_test_recharge_orders_for_user(&auth.user.id, limit, offset) + let (items, total) = wallet_test_recharge_orders_for_user(&auth.user.id, limit, offset); + ( + items + .into_iter() + .map(wallet_test_recharge_public_payload) + .collect(), + total, + ) } else { (items, total) }; @@ -1020,19 +2695,28 @@ pub(super) async fn handle_wallet_recharge_detail( .find_wallet_payment_order_by_user_id(&auth.user.id, &order_id) .await { - Ok(Some(order)) => build_auth_json_response( - http::StatusCode::OK, - json!({ "order": wallet_payment_order_payload_from_record(&order) }), - None, - ), + Ok(Some(order)) => { + let payment_instructions = wallet_payment_instructions_from_stored(state, &order).await; + mark_sensitive_response_no_store(build_auth_json_response( + http::StatusCode::OK, + json!({ + "order": wallet_payment_order_payload_from_record(&order), + "payment_instructions": payment_instructions, + }), + None, + )) + } Ok(None) => { #[cfg(test)] if let Some(order) = wallet_test_recharge_order_by_id(&auth.user.id, &order_id) { - return build_auth_json_response( + return mark_sensitive_response_no_store(build_auth_json_response( http::StatusCode::OK, - json!({ "order": order }), + json!({ + "order": wallet_test_recharge_public_payload(order.clone()), + "payment_instructions": wallet_test_recharge_replay_payment_instructions(&order), + }), None, - ); + )); } build_auth_error_response( http::StatusCode::NOT_FOUND, @@ -1047,3 +2731,1003 @@ pub(super) async fn handle_wallet_recharge_detail( ), } } + +#[cfg(test)] +mod tests { + use super::{ + attach_wallet_recharge_claim_token, prepare_wallet_gateway_response_for_storage, + prepare_wallet_gateway_response_for_storage_with_encrypt, sanitize_wallet_gateway_response, + stripe_wallet_checkout_response_is_canceled, stripe_wallet_idempotency_key, + wallet_payment_instructions_from_checkout, wallet_payment_instructions_from_stored, + wallet_payment_order_payload_from_record, wallet_recharge_order_is_settled, + wallet_recharge_order_matches_effective_request, wallet_recharge_order_matches_request, + wallet_recharge_payment_breakdown, wallet_recharge_replay_payment_instructions, + wallet_test_recharge_payload_matches_request, + wallet_test_recharge_replay_payment_instructions, AppState, + NormalizedWalletCreateRechargeRequest, WalletRechargePaymentBreakdown, + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY, WALLET_RECHARGE_ORDER_KIND, + }; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::auth::InMemoryAuthApiKeySnapshotRepository; + use aether_data::repository::wallet::{ + CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, InMemoryWalletRepository, + StoredAdminPaymentOrder, WalletReadRepository, WalletWriteRepository, + }; + use chrono::Utc; + use serde_json::json; + use std::sync::Arc; + + use crate::data::GatewayDataState; + use crate::handlers::shared::encrypt_catalog_secret_with_fallbacks; + + fn stored_checkout_order( + provider: &str, + status: &str, + expires_at_unix_secs: Option, + gateway_response: serde_json::Value, + ) -> StoredAdminPaymentOrder { + StoredAdminPaymentOrder { + id: "order-test".to_string(), + order_no: "po-test".to_string(), + wallet_id: "wallet-test".to_string(), + user_id: Some("user-test".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: provider.to_string(), + payment_provider: Some(provider.to_string()), + order_kind: super::WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("gateway-test".to_string()), + gateway_response: Some(gateway_response), + status: status.to_string(), + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs, + } + } + + #[test] + fn stripe_canceled_response_detection_is_case_and_whitespace_insensitive() { + assert!(stripe_wallet_checkout_response_is_canceled(&json!({ + "status": " canceled " + }))); + assert!(stripe_wallet_checkout_response_is_canceled(&json!({ + "status": "CANCELED" + }))); + assert!(!stripe_wallet_checkout_response_is_canceled(&json!({ + "status": "requires_payment_method" + }))); + assert!(!stripe_wallet_checkout_response_is_canceled(&json!({ + "status": null + }))); + assert!(!stripe_wallet_checkout_response_is_canceled(&json!({ + "status": ["canceled"] + }))); + } + + #[test] + fn stripe_wallet_retry_key_is_stable_distinct_and_keeps_order_identity() { + let order_no = "po_202608291234567890123456789012345678"; + let initial = stripe_wallet_idempotency_key(order_no, false); + let retry = stripe_wallet_idempotency_key(order_no, true); + + assert_ne!(initial, retry); + assert_eq!(initial, stripe_wallet_idempotency_key(order_no, false)); + assert_eq!(retry, stripe_wallet_idempotency_key(order_no, true)); + assert!(initial.contains(order_no)); + assert!(retry.contains(order_no)); + assert!(initial.len() <= 255); + assert!(retry.len() <= 255); + assert!(retry.ends_with("-retry-1")); + } + + #[tokio::test] + async fn stripe_client_secret_is_encrypted_at_rest_and_only_rehydrated_explicitly() { + let state = AppState::new().expect("test app state"); + let client_secret = "pi_test_123_secret_payment_capability"; + let checkout = json!({ + "gateway": "stripe", + "gateway_order_id": "pi_test_123", + "intent_id": "pi_test_123", + "client_secret": client_secret, + "publishable_key": "pk_test_public", + "payment_channel": "card", + }); + + let stored = prepare_wallet_gateway_response_for_storage( + &state, + "stripe", + "po-test", + "user-test", + &checkout, + ) + .expect("Stripe checkout should encrypt"); + assert!(stored.get("client_secret").is_none()); + assert!(stored + .get(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY) + .and_then(serde_json::Value::as_str) + .is_some()); + assert!(!stored.to_string().contains(client_secret)); + + let projected = sanitize_wallet_gateway_response(Some(stored.clone())); + assert!(projected.get("client_secret").is_none()); + assert!(projected.get(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY).is_none()); + + let fresh = wallet_payment_instructions_from_checkout("stripe", &checkout); + assert_eq!(fresh["client_secret"], client_secret); + let order = stored_checkout_order("stripe", "pending", Some(u64::MAX), stored); + let resumed = wallet_payment_instructions_from_stored(&state, &order).await; + assert_eq!(resumed["client_secret"], client_secret); + } + + #[tokio::test] + async fn wallet_stripe_client_secret_cannot_be_copied_to_another_order_or_kind() { + let state = AppState::new().expect("test app state"); + let client_secret = "pi_wallet_source_secret_capability"; + let stored = prepare_wallet_gateway_response_for_storage( + &state, + "stripe", + "po-source", + "user-source", + &json!({ + "gateway": "stripe", + "client_secret": client_secret, + "publishable_key": "pk_test_public", + }), + ) + .expect("source checkout should encrypt"); + let mut source = stored_checkout_order("stripe", "pending", Some(u64::MAX), stored.clone()); + source.id = "order-source".to_string(); + source.order_no = "po-source".to_string(); + source.user_id = Some("user-source".to_string()); + assert_eq!( + wallet_payment_instructions_from_stored(&state, &source).await["client_secret"], + client_secret + ); + + let mut foreign_order = source.clone(); + foreign_order.id = "order-foreign".to_string(); + foreign_order.order_no = "po-foreign".to_string(); + assert!( + wallet_payment_instructions_from_stored(&state, &foreign_order) + .await + .get("client_secret") + .is_none() + ); + + let mut foreign_owner = source.clone(); + foreign_owner.id = "order-foreign-owner".to_string(); + foreign_owner.user_id = Some("user-foreign".to_string()); + assert!( + wallet_payment_instructions_from_stored(&state, &foreign_owner) + .await + .get("client_secret") + .is_none() + ); + + let mut foreign_kind = source; + foreign_kind.id = "order-plan".to_string(); + foreign_kind.order_kind = "plan_purchase".to_string(); + assert!( + wallet_payment_instructions_from_stored(&state, &foreign_kind) + .await + .get("client_secret") + .is_none() + ); + } + + #[tokio::test] + async fn legacy_wallet_stripe_secret_is_atomically_migrated_but_unknown_envelope_is_not() { + let wallet_repository = Arc::new(InMemoryWalletRepository::default()); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let state = AppState::new() + .expect("test app state") + .with_data_state_for_tests( + GatewayDataState::with_auth_and_wallet_for_tests( + auth_repository, + Arc::clone(&wallet_repository), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let plaintext = "pi_legacy_wallet_secret_capability"; + let legacy = encrypt_catalog_secret_with_fallbacks(&state, plaintext) + .expect("legacy Fernet should encrypt"); + let expires_at = Utc::now().timestamp().max(0) as u64 + 600; + let order = match wallet_repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-legacy".to_string()), + user_id: "user-legacy".to_string(), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "pi-legacy".to_string(), + gateway_response: json!({ + "gateway": "stripe", + "publishable_key": "pk_test_public", + (STRIPE_CLIENT_SECRET_ENCRYPTED_KEY): legacy, + }), + order_no: "po-legacy".to_string(), + expires_at_unix_secs: expires_at, + }) + .await + .expect("legacy order creation should run") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + other => panic!("unexpected legacy order outcome: {other:?}"), + }; + + let instructions = wallet_payment_instructions_from_stored(&state, &order).await; + assert_eq!(instructions["client_secret"], plaintext); + let migrated = wallet_repository + .find_admin_payment_order(&order.id) + .await + .expect("migrated order should be readable") + .expect("migrated order should remain"); + assert!(migrated + .gateway_response + .as_ref() + .and_then(|response| response[STRIPE_CLIENT_SECRET_ENCRYPTED_KEY].as_str()) + .is_some_and(|value| value.starts_with( + "aether-payment-order-stripe-client-secret-v2:aether-runtime-secret-v1:" + ))); + + let unknown = "aether-payment-order-stripe-client-secret-v3:unknown"; + let unknown_order = match wallet_repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-unknown".to_string()), + user_id: "user-unknown".to_string(), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "pi-unknown".to_string(), + gateway_response: json!({ + "gateway": "stripe", + "publishable_key": "pk_test_public", + (STRIPE_CLIENT_SECRET_ENCRYPTED_KEY): unknown, + }), + order_no: "po-unknown".to_string(), + expires_at_unix_secs: expires_at, + }) + .await + .expect("unknown-envelope order creation should run") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + other => panic!("unexpected unknown-envelope outcome: {other:?}"), + }; + assert!( + wallet_payment_instructions_from_stored(&state, &unknown_order) + .await + .get("client_secret") + .is_none() + ); + let unchanged = wallet_repository + .find_admin_payment_order(&unknown_order.id) + .await + .expect("unknown-envelope order should be readable") + .expect("unknown-envelope order should remain"); + assert_eq!( + unchanged + .gateway_response + .as_ref() + .and_then(|response| response[STRIPE_CLIENT_SECRET_ENCRYPTED_KEY].as_str()), + Some(unknown) + ); + } + + #[test] + fn stripe_checkout_is_not_persistable_when_client_secret_encryption_fails() { + let checkout = json!({ + "gateway": "stripe", + "client_secret": "pi_test_123_secret_payment_capability", + "publishable_key": "pk_test_public", + }); + + let error = + prepare_wallet_gateway_response_for_storage_with_encrypt("stripe", &checkout, |_| None) + .expect_err("encryption failure must stop persistence"); + + assert_eq!(error, "Stripe client_secret 加密失败"); + } + + #[tokio::test] + async fn stored_legacy_plaintext_client_secret_is_not_replayed() { + let state = AppState::new().expect("test app state"); + let legacy = json!({ + "gateway": "stripe", + "client_secret": "pi_legacy_secret_plaintext", + "publishable_key": "pk_test_public", + }); + + let order = stored_checkout_order("stripe", "pending", Some(u64::MAX), legacy); + let instructions = wallet_payment_instructions_from_stored(&state, &order).await; + assert!(instructions.get("client_secret").is_none()); + assert_eq!(instructions["publishable_key"], "pk_test_public"); + } + + #[tokio::test] + async fn stored_stripe_client_secret_is_only_rehydrated_for_live_pending_orders() { + let state = AppState::new().expect("test app state"); + let checkout = json!({ + "gateway": "stripe", + "client_secret": "pi_test_123_secret_payment_capability", + "publishable_key": "pk_test_public", + }); + let stored = prepare_wallet_gateway_response_for_storage( + &state, + "stripe", + "po-test", + "user-test", + &checkout, + ) + .expect("Stripe checkout should encrypt"); + + for (status, expires_at) in [ + ("paid", Some(u64::MAX)), + ("credited", Some(u64::MAX)), + ("expired", Some(u64::MAX)), + ("pending", Some(0)), + ("pending", None), + ] { + let order = stored_checkout_order("stripe", status, expires_at, stored.clone()); + let instructions = wallet_payment_instructions_from_stored(&state, &order).await; + assert!( + instructions.get("client_secret").is_none(), + "secret was replayed for status={status}, expires_at={expires_at:?}", + ); + assert_eq!( + instructions, + json!({}), + "checkout capabilities were replayed for status={status}, expires_at={expires_at:?}" + ); + } + + let live_order = stored_checkout_order( + "stripe", + "pending", + Some(chrono::Utc::now().timestamp().max(0) as u64 + 60), + stored, + ); + let live = wallet_payment_instructions_from_stored(&state, &live_order).await; + assert_eq!(live["publishable_key"], "pk_test_public"); + assert_eq!( + live["client_secret"], + "pi_test_123_secret_payment_capability" + ); + } + + #[tokio::test] + async fn stored_non_stripe_payment_instructions_are_only_returned_for_live_pending_orders() { + let state = AppState::new().expect("test app state"); + let checkout = json!({ + "gateway": "alipay", + "payment_url": "https://pay.example.test/order", + "payment_params": { + "out_trade_no": "order-1", + "sign": "signed-value", + }, + }); + let now = chrono::Utc::now().timestamp().max(0) as u64; + + let live_order = stored_checkout_order( + "alipay", + "pending", + Some(now.saturating_add(60)), + checkout.clone(), + ); + let live = wallet_payment_instructions_from_stored(&state, &live_order).await; + assert_eq!(live["payment_url"], checkout["payment_url"]); + + for (status, expires_at) in [ + ("paid", Some(now.saturating_add(60))), + ("credited", Some(now.saturating_add(60))), + ("pending", Some(now.saturating_sub(1))), + ("pending", None), + ] { + let order = stored_checkout_order("alipay", status, expires_at, checkout.clone()); + let instructions = wallet_payment_instructions_from_stored(&state, &order).await; + assert_eq!( + instructions, + json!({}), + "status={status} expires={expires_at:?}" + ); + } + } + + #[test] + fn historical_order_payload_does_not_expose_checkout_capabilities() { + let order = StoredAdminPaymentOrder { + id: "historical-order".to_string(), + order_no: "po-historical".to_string(), + wallet_id: "wallet-historical".to_string(), + user_id: Some("user-historical".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("gateway-historical".to_string()), + gateway_response: Some(json!({ + "gateway": "alipay", + "payment_url": "https://pay.example.test/historical", + "payment_params": {"out_trade_no": "po-historical", "sign": "signed"}, + })), + status: "credited".to_string(), + // SQL adapters expose this field in seconds despite its legacy + // name; the public projection must preserve that meaning. + created_at_unix_ms: 1_700_000_000, + paid_at_unix_secs: Some(1_700_000_001), + credited_at_unix_secs: Some(1_700_000_002), + expires_at_unix_secs: Some(1_700_000_003), + }; + + let payload = wallet_payment_order_payload_from_record(&order); + assert_eq!(payload["gateway_response"], json!({})); + assert_eq!(payload["created_at"], "2023-11-14T22:13:20Z"); + } + + #[test] + fn unexpected_client_secret_is_removed_for_non_stripe_storage() { + let state = AppState::new().expect("test app state"); + let checkout = json!({ + "gateway": "alipay", + "payment_url": "https://pay.example.test/order", + "client_secret": "pi_injected_secret_should_not_survive", + "provider_private_token": "should_not_survive", + "payment_params": { + "client_secret": "pi_nested_secret_should_not_survive", + "order_no": "order-1", + "out_trade_no": "order-1", + "sign": "signed-value", + "nested": {"authorization": "Bearer secret"}, + "array": ["secret"], + }, + }); + + let stored = prepare_wallet_gateway_response_for_storage( + &state, + "alipay", + "po-test", + "user-test", + &checkout, + ) + .expect("non-Stripe checkout should project"); + assert!(stored.get("client_secret").is_none()); + assert!(stored.get("provider_private_token").is_none()); + assert!(stored["payment_params"].get("client_secret").is_none()); + assert!(stored["payment_params"].get("nested").is_none()); + assert!(stored["payment_params"].get("array").is_none()); + assert_eq!(stored["payment_params"]["out_trade_no"], "order-1"); + assert_eq!(stored["payment_params"]["sign"], "signed-value"); + assert_eq!(stored["payment_url"], checkout["payment_url"]); + } + + #[test] + fn checkout_claim_metadata_is_persisted_but_never_exposed() { + let state = AppState::new().expect("test app state"); + let checkout = attach_wallet_recharge_claim_token( + json!({ + "gateway": "alipay", + "payment_url": "https://pay.example.test/order", + "payment_channel": "alipay", + }), + "wrc_test_claim", + ); + + let stored = prepare_wallet_gateway_response_for_storage( + &state, + "alipay", + "po-test", + "user-test", + &checkout, + ) + .expect("checkout claim should be persistable"); + assert_eq!( + stored + .get("checkout_claim_token") + .and_then(serde_json::Value::as_str), + Some("wrc_test_claim") + ); + let public = sanitize_wallet_gateway_response(Some(stored)); + assert!(public.get("checkout_claim_token").is_none()); + } + + #[tokio::test] + async fn idempotent_replay_hides_expired_non_stripe_checkout() { + let state = AppState::new().expect("test app state"); + let order = StoredAdminPaymentOrder { + id: "expired-replay-order".to_string(), + order_no: "po_idem_expired".to_string(), + wallet_id: "wallet-expired-replay".to_string(), + user_id: Some("user-expired-replay".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("po_idem_expired".to_string()), + gateway_response: Some(json!({ + "gateway": "alipay", + "payment_url": "https://pay.example.test/expired", + "payment_channel": "alipay", + })), + status: "pending".to_string(), + created_at_unix_ms: 0, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(0), + }; + + assert_eq!( + wallet_recharge_replay_payment_instructions(&state, &order).await, + json!({}), + ); + } + + #[test] + fn effective_matcher_rejects_a_competing_request_with_a_different_usd_amount() { + let order = StoredAdminPaymentOrder { + id: "placeholder-order".to_string(), + order_no: "po_idem_placeholder".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("po_idem_placeholder".to_string()), + gateway_response: Some(json!({ + "gateway": "alipay", + "payment_channel": "alipay", + "integration_status": "checkout_pending", + })), + status: "pending".to_string(), + created_at_unix_ms: 0, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(u64::MAX), + }; + + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 11.0, 79.75, "CNY", 7.25, + )); + assert!(wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.5, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "wxpay", 10.0, 72.5, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.51, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.5, "USD", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.5, "CNY", 7.26, + )); + } + + #[test] + fn effective_matcher_still_runs_when_existing_order_has_checkout_evidence() { + let mut order = StoredAdminPaymentOrder { + id: "checkout-order".to_string(), + order_no: "po_idem_checkout".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("gw-1".to_string()), + gateway_response: Some(json!({ + "gateway": "alipay", + "payment_channel": "alipay", + "payment_url": "https://pay.example.test/gw-1" + })), + status: "pending".to_string(), + created_at_unix_ms: 1_700_000_000_000, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(u64::MAX), + }; + + assert!(wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.5, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 11.0, 79.75, "CNY", 7.25, + )); + } + + #[test] + fn effective_matcher_accepts_legacy_epay_channel_alias() { + let order = StoredAdminPaymentOrder { + id: "legacy-epay-order".to_string(), + order_no: "po_legacy_epay".to_string(), + wallet_id: "wallet-legacy-epay".to_string(), + user_id: Some("user-legacy-epay".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + // Older EPay rows stored the selected channel in payment_method. + payment_method: "alipay".to_string(), + payment_provider: Some("epay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("po_legacy_epay".to_string()), + gateway_response: Some(json!({ + "gateway": "epay", + "integration_status": "checkout_pending", + })), + status: "pending".to_string(), + created_at_unix_ms: 0, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(u64::MAX), + }; + + assert!(wallet_recharge_order_matches_effective_request( + &order, "epay", "alipay", 10.0, 72.5, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "epay", "wxpay", 10.0, 72.5, "CNY", 7.25, + )); + assert!(!wallet_recharge_order_matches_effective_request( + &order, "alipay", "alipay", 10.0, 72.5, "CNY", 7.25, + )); + } + + #[tokio::test] + async fn replay_instructions_are_withheld_after_checkout_race_credits_order() { + let state = AppState::new().expect("test app state"); + let mut order = StoredAdminPaymentOrder { + id: "credited-checkout-order".to_string(), + order_no: "po_credited_checkout".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("gw-1".to_string()), + gateway_response: Some(json!({ + "gateway": "alipay", + "payment_channel": "alipay", + "payment_url": "https://pay.example.test/gw-1" + })), + status: "credited".to_string(), + created_at_unix_ms: 1_700_000_000_000, + paid_at_unix_secs: Some(1_700_000_010), + credited_at_unix_secs: Some(1_700_000_010), + expires_at_unix_secs: Some(u64::MAX), + }; + + assert!(wallet_recharge_order_is_settled(&order)); + assert_eq!( + super::wallet_recharge_replay_payment_instructions(&state, &order).await, + json!({}), + ); + + order.status = "paid".to_string(); + assert!(wallet_recharge_order_is_settled(&order)); + order.status = "pending".to_string(); + assert!(!wallet_recharge_order_is_settled(&order)); + order.status = "failed".to_string(); + assert!(!wallet_recharge_order_is_settled(&order)); + } + + #[test] + fn test_recharge_replay_hides_expired_or_non_pending_instructions() { + let expired = json!({ + "status": "pending", + "expires_at": "2000-01-01T00:00:00Z", + "gateway_response": { + "gateway": "alipay", + "payment_url": "https://pay.example.test/expired" + } + }); + assert_eq!( + wallet_test_recharge_replay_payment_instructions(&expired), + json!({}) + ); + + let paid = json!({ + "status": "credited", + "expires_at": "2999-01-01T00:00:00Z", + "gateway_response": { + "gateway": "alipay", + "payment_url": "https://pay.example.test/paid" + } + }); + assert_eq!( + wallet_test_recharge_replay_payment_instructions(&paid), + json!({}) + ); + } + + #[test] + fn test_recharge_idempotency_matches_all_client_amount_fields() { + let order = json!({ + "amount_usd": 10.0, + "payment_method": "alipay", + "payment_provider": "alipay", + "payment_channel": "alipay", + "pay_amount": 72.5, + "pay_currency": "CNY", + "exchange_rate": 7.25, + "gateway_response": { + "gateway": "alipay", + "payment_channel": "alipay" + } + }); + let payload = NormalizedWalletCreateRechargeRequest { + amount_usd: 10.0, + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + idempotency_key: Some("recharge-match".to_string()), + }; + assert!(wallet_test_recharge_payload_matches_request( + &order, &payload + )); + + let mut changed = payload.clone(); + changed.pay_amount = Some(72.51); + assert!(!wallet_test_recharge_payload_matches_request( + &order, &changed + )); + changed = payload.clone(); + changed.pay_currency = Some("USD".to_string()); + assert!(!wallet_test_recharge_payload_matches_request( + &order, &changed + )); + changed = payload; + changed.exchange_rate = Some(7.26); + assert!(!wallet_test_recharge_payload_matches_request( + &order, &changed + )); + } + + #[test] + fn epay_idempotency_includes_legacy_payment_method_channel_shorthand() { + let order = StoredAdminPaymentOrder { + id: "epay-idempotent-order".to_string(), + order_no: "po_idem_epay".to_string(), + wallet_id: "wallet-epay".to_string(), + user_id: Some("user-epay".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + order_kind: WALLET_RECHARGE_ORDER_KIND.to_string(), + gateway_order_id: Some("po_idem_epay".to_string()), + gateway_response: Some(json!({ + "gateway": "epay", + "payment_channel": "alipay", + "integration_status": "checkout_pending" + })), + status: "pending".to_string(), + created_at_unix_ms: 0, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(u64::MAX), + }; + let payload = |payment_method: &str, payment_channel: Option<&str>| { + NormalizedWalletCreateRechargeRequest { + amount_usd: 10.0, + payment_method: payment_method.to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: payment_channel.map(str::to_string), + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + idempotency_key: Some("epay-channel-regression".to_string()), + } + }; + + assert!(wallet_recharge_order_matches_request( + &order, + &payload("alipay", None), + )); + assert!(!wallet_recharge_order_matches_request( + &order, + &payload("wxpay", None), + )); + assert!(!wallet_recharge_order_matches_request( + &order, + &payload("alipay", Some("wxpay")), + )); + + let legacy_order = StoredAdminPaymentOrder { + payment_method: "alipay".to_string(), + gateway_response: Some(json!({ + "gateway": "epay", + "integration_status": "checkout_pending" + })), + ..order + }; + assert!(wallet_recharge_order_matches_request( + &legacy_order, + &payload("alipay", None), + )); + assert!(!wallet_recharge_order_matches_request( + &legacy_order, + &payload("wxpay", None), + )); + + let mut direct_provider_payload = payload("alipay", None); + direct_provider_payload.payment_provider = Some("alipay".to_string()); + assert!(!wallet_recharge_order_matches_request( + &legacy_order, + &direct_provider_payload, + )); + } + + #[test] + fn test_epay_idempotency_includes_legacy_payment_method_channel_shorthand() { + let order = json!({ + "amount_usd": 10.0, + "payment_method": "epay", + "payment_provider": "epay", + "pay_amount": 72.5, + "pay_currency": "CNY", + "exchange_rate": 7.25, + "gateway_response": { + "gateway": "epay", + "payment_channel": "alipay" + } + }); + let payload = |payment_method: &str| NormalizedWalletCreateRechargeRequest { + amount_usd: 10.0, + payment_method: payment_method.to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: None, + pay_amount: Some(72.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.25), + idempotency_key: Some("epay-channel-regression-json".to_string()), + }; + assert!(wallet_test_recharge_payload_matches_request( + &order, + &payload("alipay") + )); + assert!(!wallet_test_recharge_payload_matches_request( + &order, + &payload("wxpay") + )); + + let legacy_order = json!({ + "amount_usd": 10.0, + "payment_method": "alipay", + "payment_provider": "epay", + "pay_amount": 72.5, + "pay_currency": "CNY", + "exchange_rate": 7.25, + "gateway_response": { + "gateway": "epay" + } + }); + assert!(wallet_test_recharge_payload_matches_request( + &legacy_order, + &payload("alipay") + )); + assert!(!wallet_test_recharge_payload_matches_request( + &legacy_order, + &payload("wxpay") + )); + } + + #[test] + fn direct_gateway_channels_normalize_configured_channel_case() { + let record = aether_data_contracts::repository::billing::PaymentGatewayConfigRecord { + provider: "wxpay".to_string(), + enabled: true, + endpoint_url: "https://pay.example.test".to_string(), + callback_base_url: Some("https://app.example.test".to_string()), + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some("encrypted".to_string()), + pay_currency: "CNY".to_string(), + usd_exchange_rate: 7.0, + min_recharge_usd: 1.0, + channels_json: json!([ + {"channel": "NATIVE", "display_name": "Native", "fee_rate": 0.0}, + {"channel": "H5", "display_name": "H5", "fee_rate": 0.0}, + {"channel": "APP", "display_name": "Unsupported", "fee_rate": 0.0} + ]), + created_at_unix_secs: 0, + updated_at_unix_secs: 0, + }; + + let channels = super::direct_gateway_channels("WXPAY", &record); + assert_eq!( + channels + .iter() + .map(|channel| channel.channel.as_str()) + .collect::>(), + vec!["native", "h5"] + ); + let selected = super::resolve_direct_gateway_channel("wxpay", &record, Some("NATIVE")) + .expect("configured channel should resolve case-insensitively"); + assert_eq!(selected.channel, "native"); + } + + #[test] + fn recharge_payment_breakdown_rejects_non_finite_inputs_and_overflow() { + assert!(wallet_recharge_payment_breakdown(10.0, "CNY", f64::NAN, 0.0).is_err()); + assert!(wallet_recharge_payment_breakdown(10.0, "CNY", 7.0, f64::INFINITY).is_err()); + assert!(wallet_recharge_payment_breakdown(f64::MAX, "CNY", f64::MAX, 0.0).is_err()); + let breakdown = wallet_recharge_payment_breakdown(10.0, "CNY", 7.0, 2.5) + .expect("valid recharge amounts"); + assert!(breakdown.base_pay_amount.is_finite()); + assert!(breakdown.fee_amount.is_finite()); + assert!(breakdown.pay_amount.is_finite()); + assert_eq!( + breakdown, + WalletRechargePaymentBreakdown { + base_pay_amount: 70.0, + fee_amount: 1.75, + pay_amount: 71.75, + exchange_rate: 7.0, + } + ); + } + + #[test] + fn usd_recharge_uses_unit_exchange_rate_for_amount_order_and_fee() { + let breakdown = wallet_recharge_payment_breakdown(10.0, "USD", 7.2, 2.5) + .expect("USD recharge should use canonical rate"); + assert_eq!(breakdown.exchange_rate, 1.0); + assert_eq!( + ( + breakdown.base_pay_amount, + breakdown.fee_amount, + breakdown.pay_amount + ), + (10.0, 0.25, 10.25) + ); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/wallet/redeem.rs b/apps/aether-gateway/src/handlers/public/support/wallet/redeem.rs index ad7bc4d05..4a6358b8d 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet/redeem.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet/redeem.rs @@ -8,6 +8,8 @@ use serde::Deserialize; use serde_json::json; use uuid::Uuid; +use aether_data::repository::wallet::stored_timestamp_unix_secs; + #[derive(Debug, Deserialize)] struct WalletRedeemRequest { code: String, @@ -39,7 +41,7 @@ fn build_wallet_payment_order_payload( "gateway_order_id": record.gateway_order_id, "gateway_response": sanitize_wallet_gateway_response(record.gateway_response.clone()), "status": record.status, - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "paid_at": record.paid_at_unix_secs.and_then(unix_secs_to_rfc3339), "credited_at": record.credited_at_unix_secs.and_then(unix_secs_to_rfc3339), "expires_at": record.expires_at_unix_secs.and_then(unix_secs_to_rfc3339), diff --git a/apps/aether-gateway/src/handlers/public/support/wallet/refunds.rs b/apps/aether-gateway/src/handlers/public/support/wallet/refunds.rs index 92daa05d5..0bd7f7a28 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet/refunds.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet/refunds.rs @@ -17,6 +17,9 @@ use uuid::Uuid; use crate::handlers::shared::{ payment_gateway_allow_user_refund, payment_gateway_provider_for_payment_method, }; +use aether_data::repository::wallet::{ + canonicalize_wallet_refund_fields, stored_timestamp_unix_secs, +}; const WALLET_REFUND_CONFIGURED_PROVIDERS: &[&str] = &["epay", "alipay", "wxpay", "stripe"]; @@ -150,14 +153,67 @@ fn wallet_refund_payload_from_record( "gateway_refund_id": record.gateway_refund_id, "payout_method": record.payout_method, "payout_reference": record.payout_reference, - "payout_proof": record.payout_proof, - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(record.updated_at_unix_secs), "processed_at": record.processed_at_unix_secs.and_then(unix_secs_to_rfc3339), "completed_at": record.completed_at_unix_secs.and_then(unix_secs_to_rfc3339), }) } +#[cfg(test)] +fn wallet_public_refund_payload(mut payload: serde_json::Value) -> serde_json::Value { + if let Some(object) = payload.as_object_mut() { + object.remove("payout_proof"); + } + payload +} + +#[cfg(test)] +mod tests { + use super::wallet_refund_payload_from_record; + use aether_data::repository::wallet::StoredAdminWalletRefund; + use serde_json::json; + + #[test] + fn public_refund_projection_excludes_payout_proof_and_upstream_payload() { + let record = StoredAdminWalletRefund { + id: "refund-1".to_string(), + refund_no: "rf_1".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + payment_order_id: Some("order-1".to_string()), + source_type: "payment_order".to_string(), + source_id: Some("order-1".to_string()), + refund_mode: "original_channel".to_string(), + amount_usd: 10.0, + status: "processing".to_string(), + reason: Some("requested".to_string()), + failure_reason: None, + gateway_refund_id: Some("gateway-refund-1".to_string()), + payout_method: None, + payout_reference: None, + payout_proof: Some(json!({ + "gateway_refund": { + "id": "gateway-refund-1", + "payload": {"payer": "sensitive", "credential": "secret"} + } + })), + requested_by: Some("user-1".to_string()), + approved_by: Some("admin-1".to_string()), + processed_by: Some("admin-1".to_string()), + created_at_unix_ms: 1, + updated_at_unix_secs: 1, + processed_at_unix_secs: Some(1), + completed_at_unix_secs: None, + }; + + let payload = wallet_refund_payload_from_record(&record); + assert!(payload.get("payout_proof").is_none()); + assert_eq!(payload["status"], "processing"); + assert_eq!(payload["gateway_refund_id"], "gateway-refund-1"); + } +} + pub(super) async fn handle_wallet_refunds_list( state: &AppState, request_context: &GatewayPublicRequestContext, @@ -240,8 +296,7 @@ pub(super) async fn handle_wallet_refunds_list( "gateway_refund_id": record.gateway_refund_id, "payout_method": record.payout_method, "payout_reference": record.payout_reference, - "payout_proof": record.payout_proof, - "created_at": unix_secs_to_rfc3339(record.created_at_unix_ms), + "created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)), "updated_at": unix_secs_to_rfc3339(record.updated_at_unix_secs), "processed_at": record.processed_at_unix_secs.and_then(unix_secs_to_rfc3339), "completed_at": record.completed_at_unix_secs.and_then(unix_secs_to_rfc3339), @@ -257,6 +312,7 @@ pub(super) async fn handle_wallet_refunds_list( .into_iter() .skip(offset) .take(limit) + .map(wallet_public_refund_payload) .collect::>(); (items, total) } else { @@ -324,7 +380,11 @@ pub(super) async fn handle_wallet_refund_detail( Ok(None) => { #[cfg(test)] if let Some(payload) = wallet_test_refund_by_id(&wallet.id, &refund_id) { - return build_auth_json_response(http::StatusCode::OK, payload, None); + return build_auth_json_response( + http::StatusCode::OK, + wallet_public_refund_payload(payload), + None, + ); } build_auth_error_response( http::StatusCode::NOT_FOUND, @@ -359,7 +419,7 @@ pub(super) async fn handle_wallet_create_refund( return build_auth_error_response(http::StatusCode::BAD_REQUEST, "输入验证失败", false) } }; - let payload = match normalize_wallet_create_refund_request(payload) { + let mut payload = match normalize_wallet_create_refund_request(payload) { Ok(value) => value, Err(detail) => { return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); @@ -389,6 +449,7 @@ pub(super) async fn handle_wallet_create_refund( ); }; + let mut resolved_payment_method = None; if let Some(payment_order_id) = payload.payment_order_id.as_deref() { let order = match state .find_wallet_payment_order_by_user_id(&auth.user.id, payment_order_id) @@ -437,8 +498,25 @@ pub(super) async fn handle_wallet_create_refund( false, ); } + resolved_payment_method = Some(order.payment_method.clone()); } + let canonical = match canonicalize_wallet_refund_fields( + payload.payment_order_id.as_deref(), + payload.source_type.as_deref(), + payload.source_id.as_deref(), + payload.refund_mode.as_deref(), + resolved_payment_method.as_deref(), + ) { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); + } + }; + payload.source_type = Some(canonical.source_type); + payload.source_id = canonical.source_id; + payload.refund_mode = Some(canonical.refund_mode); + if !state.has_database_wallet_data_writer() { #[cfg(test)] { @@ -446,11 +524,19 @@ pub(super) async fn handle_wallet_create_refund( if let Some(existing) = wallet_test_refund_by_idempotency(&auth.user.id, idempotency_key) { - return build_auth_json_response(http::StatusCode::OK, existing, None); + return build_auth_json_response( + http::StatusCode::OK, + wallet_public_refund_payload(existing), + None, + ); } } let reserved_amount = wallet_test_reserved_refund_amount(&wallet.id); - if payload.amount_usd > (wallet.balance - reserved_amount) { + let available_balance = wallet.balance - reserved_amount; + if !wallet.balance.is_finite() + || !available_balance.is_finite() + || payload.amount_usd > available_balance + { return build_auth_error_response( http::StatusCode::BAD_REQUEST, "refund amount exceeds available refundable recharge balance", @@ -472,7 +558,6 @@ pub(super) async fn handle_wallet_create_refund( "gateway_refund_id": serde_json::Value::Null, "payout_method": serde_json::Value::Null, "payout_reference": serde_json::Value::Null, - "payout_proof": serde_json::Value::Null, "created_at": now.to_rfc3339(), "updated_at": now.to_rfc3339(), "processed_at": serde_json::Value::Null, @@ -527,6 +612,9 @@ pub(super) async fn handle_wallet_create_refund( None, ) } + aether_data::repository::wallet::CreateWalletRefundRequestOutcome::InvalidInput(detail) => { + build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } aether_data::repository::wallet::CreateWalletRefundRequestOutcome::WalletMissing => { build_auth_error_response( http::StatusCode::BAD_REQUEST, diff --git a/apps/aether-gateway/src/handlers/public/support/wallet/test_support.rs b/apps/aether-gateway/src/handlers/public/support/wallet/test_support.rs index e46fabfbd..3b0fa6b9e 100644 --- a/apps/aether-gateway/src/handlers/public/support/wallet/test_support.rs +++ b/apps/aether-gateway/src/handlers/public/support/wallet/test_support.rs @@ -9,6 +9,7 @@ pub(super) struct WalletTestRefundRecord { #[derive(Debug, Clone)] pub(crate) struct WalletTestRechargeRecord { pub(crate) user_id: String, + pub(crate) order_no: String, pub(crate) payload: serde_json::Value, } @@ -81,8 +82,13 @@ pub(super) fn wallet_test_reserved_refund_amount(wallet_id: &str) -> f64 { Some("pending_approval" | "approved") ) }) - .map(|entry| entry.payload["amount_usd"].as_f64().unwrap_or_default()) - .sum::() + .filter_map(|entry| entry.payload["amount_usd"].as_f64()) + .filter(|amount| amount.is_finite() && *amount > 0.0) + .try_fold(0.0_f64, |total, amount| { + let next = total + amount; + next.is_finite().then_some(next) + }) + .unwrap_or(f64::INFINITY) } pub(super) fn record_wallet_test_refund( @@ -140,9 +146,36 @@ pub(super) fn wallet_test_recharge_order_by_id( .map(|entry| entry.payload.clone()) } -pub(super) fn record_wallet_test_recharge(user_id: String, payload: serde_json::Value) { +pub(super) fn wallet_test_recharge_order_by_order_no( + user_id: &str, + order_no: &str, +) -> Option { wallet_test_recharge_store() .lock() .expect("wallet test recharge store should lock") - .push(WalletTestRechargeRecord { user_id, payload }); + .iter() + .find(|entry| entry.user_id == user_id && entry.order_no == order_no) + .map(|entry| entry.payload.clone()) +} + +pub(super) fn record_wallet_test_recharge( + user_id: String, + order_no: String, + payload: serde_json::Value, +) { + let mut store = wallet_test_recharge_store() + .lock() + .expect("wallet test recharge store should lock"); + if let Some(existing) = store + .iter_mut() + .find(|entry| entry.user_id == user_id && entry.order_no == order_no) + { + existing.payload = payload; + return; + } + store.push(WalletTestRechargeRecord { + user_id, + order_no, + payload, + }); } diff --git a/apps/aether-gateway/src/handlers/public/system_modules_helpers/modules.rs b/apps/aether-gateway/src/handlers/public/system_modules_helpers/modules.rs index 0a2e2cbfa..c4d5eae5f 100644 --- a/apps/aether-gateway/src/handlers/public/system_modules_helpers/modules.rs +++ b/apps/aether-gateway/src/handlers/public/system_modules_helpers/modules.rs @@ -44,18 +44,7 @@ pub(crate) fn oauth_module_config_is_valid( pub(crate) fn ldap_module_config_is_valid( config: Option<&aether_data::repository::auth_modules::StoredLdapModuleConfig>, ) -> bool { - let Some(config) = config else { - return false; - }; - !config.server_url.trim().is_empty() - && !config.bind_dn.trim().is_empty() - && !config.base_dn.trim().is_empty() - && config - .bind_password_encrypted - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - .is_some() + crate::handlers::shared::ldap_module_config_is_valid(config) } pub(crate) async fn build_public_auth_modules_status_payload( diff --git a/apps/aether-gateway/src/handlers/shared/admin_proxy.rs b/apps/aether-gateway/src/handlers/shared/admin_proxy.rs index 770991459..8ee5d935a 100644 --- a/apps/aether-gateway/src/handlers/shared/admin_proxy.rs +++ b/apps/aether-gateway/src/handlers/shared/admin_proxy.rs @@ -50,3 +50,39 @@ pub(crate) fn attach_admin_audit_response( attach_admin_audit_event(&mut response, event_name, action, target_type, target_id); response } + +pub(crate) fn mark_sensitive_admin_response_no_store( + mut response: Response, +) -> Response { + response.headers_mut().insert( + http::header::CACHE_CONTROL, + http::HeaderValue::from_static("no-store"), + ); + response.headers_mut().insert( + http::header::PRAGMA, + http::HeaderValue::from_static("no-cache"), + ); + response +} + +#[cfg(test)] +mod tests { + use super::mark_sensitive_admin_response_no_store; + use axum::{http, response::IntoResponse, Json}; + + #[test] + fn plaintext_admin_secret_responses_are_never_cacheable() { + let response = mark_sensitive_admin_response_no_store( + Json(serde_json::json!({ "key": "secret" })).into_response(), + ); + + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&http::HeaderValue::from_static("no-store")) + ); + assert_eq!( + response.headers().get(http::header::PRAGMA), + Some(&http::HeaderValue::from_static("no-cache")) + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/api_keys.rs b/apps/aether-gateway/src/handlers/shared/api_keys.rs index 79bfa6ca7..ba1630ab0 100644 --- a/apps/aether-gateway/src/handlers/shared/api_keys.rs +++ b/apps/aether-gateway/src/handlers/shared/api_keys.rs @@ -68,18 +68,44 @@ pub(crate) fn generate_gateway_api_key_plaintext() -> String { generate_gateway_api_key_plaintext_with_prefix(&configured_api_key_prefix()) } +pub(crate) fn masked_secret_display( + value: &str, + preferred_prefix_chars: usize, + preferred_suffix_chars: usize, + separator: &str, +) -> String { + let char_count = value.chars().count(); + if char_count <= 4 { + return "***".to_string(); + } + + // Never reveal more than half of a short secret. For normal generated keys, + // the preferred prefix/suffix remains stable while a meaningful middle + // section is always hidden. + let visible_budget = + (char_count / 2).min(preferred_prefix_chars.saturating_add(preferred_suffix_chars)); + if visible_budget == 0 { + return "***".to_string(); + } + let suffix_chars = preferred_suffix_chars.min(visible_budget / 2); + let prefix_chars = preferred_prefix_chars.min(visible_budget.saturating_sub(suffix_chars)); + if prefix_chars == 0 && suffix_chars == 0 { + return "***".to_string(); + } + + let prefix = value.chars().take(prefix_chars).collect::(); + let suffix = value + .chars() + .skip(char_count.saturating_sub(suffix_chars)) + .collect::(); + format!("{prefix}{separator}{suffix}") +} + pub(crate) fn masked_gateway_api_key_display(full_key: Option<&str>) -> String { let Some(full_key) = full_key.map(str::trim).filter(|value| !value.is_empty()) else { return api_key_placeholder_display(); }; - let prefix_len = full_key.len().min(10); - let prefix = &full_key[..prefix_len]; - let suffix = if full_key.len() >= 4 { - &full_key[full_key.len().saturating_sub(4)..] - } else { - "" - }; - format!("{prefix}...{suffix}") + masked_secret_display(full_key, 10, 4, "...") } pub(crate) fn normalize_optional_api_key_concurrent_limit( @@ -96,7 +122,7 @@ mod tests { use super::{ api_key_placeholder_display_with_prefix, configured_api_key_prefix_from_lookup, generate_gateway_api_key_plaintext_with_prefix, generate_gateway_secret_plaintext, - masked_gateway_api_key_display, + masked_gateway_api_key_display, masked_secret_display, }; #[test] @@ -149,7 +175,17 @@ mod tests { fn masks_plaintext_api_key_without_changing_prefix() { assert_eq!( masked_gateway_api_key_display(Some("ak-1234567890abcdef")), - "ak-1234567...cdef".to_string() + "ak-12...cdef".to_string() ); } + + #[test] + fn masking_never_discloses_an_entire_short_or_unicode_secret() { + for secret in ["abc", "sk-short", "测试密钥一二三"] { + let masked = masked_secret_display(secret, 10, 4, "..."); + assert_ne!(masked, secret); + assert!(!masked.contains(secret)); + assert!(masked.contains('*') || masked.contains("...")); + } + } } diff --git a/apps/aether-gateway/src/handlers/shared/auth_api_key_secret.rs b/apps/aether-gateway/src/handlers/shared/auth_api_key_secret.rs new file mode 100644 index 000000000..bf9381635 --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/auth_api_key_secret.rs @@ -0,0 +1,359 @@ +use aether_crypto::looks_like_python_fernet_ciphertext; +use sha2::{Digest, Sha256}; + +use crate::{AppState, GatewayError}; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +const AUTH_API_KEY_SECRET_MIGRATION_RETRIES: usize = 8; +const AUTH_API_KEY_SECRET_ENVELOPE_FAMILY: &str = "aether-auth-api-key-secret-"; +const AUTH_API_KEY_SECRET_ENVELOPE_V2: &str = "aether-auth-api-key-secret-v2:"; +const AUTH_API_KEY_SECRET_PURPOSE_V2: &str = "auth-api-key-secret-bound-v2"; +const AETHER_ENVELOPE_FAMILY: &str = "aether-"; + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct AuthApiKeySecretProjection { + pub(crate) plaintext: String, + pub(crate) protected: String, + pub(crate) migration_required: bool, +} + +fn auth_api_key_secret_purpose( + user_id: &str, + api_key_id: &str, + key_hash: &str, + is_standalone: bool, +) -> Result { + if user_id.is_empty() { + return Err("API-key secret owner is empty"); + } + if api_key_id.is_empty() { + return Err("API-key secret record ID is empty"); + } + if key_hash.is_empty() { + return Err("API-key secret hash is empty"); + } + let scope = if is_standalone { "standalone" } else { "user" }; + Ok(format!( + "{AUTH_API_KEY_SECRET_PURPOSE_V2}\0scope={scope}\0owner-bytes={}\0{user_id}\0api-key-id-bytes={}\0{api_key_id}\0hash-bytes={}\0{key_hash}\0field=key", + user_id.len(), + api_key_id.len(), + key_hash.len(), + )) +} + +pub(crate) fn seal_auth_api_key_secret( + state: &AppState, + user_id: &str, + api_key_id: &str, + key_hash: &str, + is_standalone: bool, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("API-key plaintext contains reserved secret framing"); + } + if sha256_hex(plaintext) != key_hash { + return Err("API-key plaintext does not match its hash"); + } + let purpose = auth_api_key_secret_purpose(user_id, api_key_id, key_hash, is_standalone)?; + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("API-key encryption key is not configured")?; + Ok(format!("{AUTH_API_KEY_SECRET_ENVELOPE_V2}{sealed}")) +} + +pub(crate) fn open_auth_api_key_secret( + state: &AppState, + record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, +) -> Result { + let observed_raw = record + .key_encrypted + .as_deref() + .ok_or("API-key ciphertext is not stored")?; + let stored = observed_raw.trim(); + if stored.is_empty() { + return Err("API-key ciphertext is empty"); + } + let purpose = auth_api_key_secret_purpose( + &record.user_id, + &record.api_key_id, + &record.key_hash, + record.is_standalone, + )?; + + let (plaintext, protected, migration_required) = + if let Some(sealed) = stored.strip_prefix(AUTH_API_KEY_SECRET_ENVELOPE_V2) { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("API-key secret authentication failed")?; + ( + plaintext, + stored.to_string(), + observed_raw.as_bytes() != stored.as_bytes(), + ) + } else { + if stored.starts_with(AUTH_API_KEY_SECRET_ENVELOPE_FAMILY) { + return Err("unsupported API-key secret envelope"); + } + // Every purpose-bound secret family in Aether uses an `aether-` envelope. Never feed + // a foreign or future envelope into the legacy Fernet path. + if stored.starts_with(AETHER_ENVELOPE_FAMILY) { + return Err("secret envelope has the wrong purpose"); + } + if !looks_like_python_fernet_ciphertext(stored) { + return Err("API-key secret is not an authenticated ciphertext"); + } + let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("legacy API-key secret authentication failed")?; + let protected = seal_auth_api_key_secret( + state, + &record.user_id, + &record.api_key_id, + &record.key_hash, + record.is_standalone, + &plaintext, + )?; + (plaintext, protected, true) + }; + + if plaintext.contains('\0') { + return Err("API-key plaintext contains reserved secret framing"); + } + if sha256_hex(&plaintext) != record.key_hash { + return Err("API-key plaintext integrity check failed"); + } + Ok(AuthApiKeySecretProjection { + plaintext, + protected, + migration_required, + }) +} + +pub(crate) async fn decrypt_or_migrate_auth_api_key_secret( + state: &AppState, + initial: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, +) -> Result { + let identity = ( + initial.user_id.clone(), + initial.api_key_id.clone(), + initial.key_hash.clone(), + initial.is_standalone, + ); + let mut current = initial.clone(); + + for _ in 0..AUTH_API_KEY_SECRET_MIGRATION_RETRIES { + if current.user_id != identity.0 + || current.api_key_id != identity.1 + || current.key_hash != identity.2 + || current.is_standalone != identity.3 + { + return Err(api_key_secret_error( + "stored API-key secret identity changed during migration", + )); + } + let projection = open_auth_api_key_secret(state, ¤t) + .map_err(|message| api_key_secret_error(message))?; + if !projection.migration_required { + return Ok(projection.plaintext); + } + let observed = current.key_encrypted.clone().ok_or_else(|| { + api_key_secret_error("stored API-key ciphertext disappeared during migration") + })?; + let mutation = aether_data::repository::auth::CompareAndSwapAuthApiKeyCiphertext { + user_id: identity.0.clone(), + api_key_id: identity.1.clone(), + key_hash: identity.2.clone(), + is_standalone: identity.3, + expected_key_encrypted: observed, + key_encrypted: projection.protected, + }; + if state.compare_and_swap_api_key_ciphertext(&mutation).await? { + return Ok(projection.plaintext); + } + + let mut matches = state + .list_auth_api_key_export_records_by_ids(std::slice::from_ref(&identity.1)) + .await? + .into_iter() + .filter(|record| record.api_key_id == identity.1); + let Some(next) = matches.next() else { + return Err(api_key_secret_error( + "stored API-key secret is unavailable during migration", + )); + }; + if matches.next().is_some() { + return Err(api_key_secret_error( + "stored API-key identity is not unique during migration", + )); + } + current = next; + } + + Err(api_key_secret_error( + "stored API-key secret migration did not stabilize", + )) +} + +fn sha256_hex(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} + +fn api_key_secret_error(message: &str) -> GatewayError { + GatewayError::Internal(message.to_string()) +} + +#[cfg(test)] +mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + use super::{open_auth_api_key_secret, seal_auth_api_key_secret, sha256_hex}; + use crate::handlers::shared::encrypt_catalog_secret_with_fallbacks; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + fn record( + plaintext: &str, + user_id: &str, + api_key_id: &str, + is_standalone: bool, + key_encrypted: Option, + ) -> aether_data::repository::auth::StoredAuthApiKeyExportRecord { + aether_data::repository::auth::StoredAuthApiKeyExportRecord { + user_id: user_id.to_string(), + api_key_id: api_key_id.to_string(), + key_hash: sha256_hex(plaintext), + key_encrypted, + name: None, + allowed_providers: None, + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + rate_limit: None, + concurrent_limit: None, + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + last_used_at_unix_secs: None, + created_at_unix_secs: None, + updated_at_unix_secs: None, + is_standalone, + } + } + + #[test] + fn v2_ciphertext_is_bound_to_owner_record_scope_and_hash() { + let state = state_with_encryption_key(); + let plaintext = "sk-record-bound-secret"; + let key_hash = sha256_hex(plaintext); + let ciphertext = + seal_auth_api_key_secret(&state, "owner-a", "key-a", &key_hash, false, plaintext) + .expect("API-key secret should seal"); + let source = record( + plaintext, + "owner-a", + "key-a", + false, + Some(ciphertext.clone()), + ); + assert_eq!( + open_auth_api_key_secret(&state, &source) + .expect("source record should open") + .plaintext, + plaintext + ); + + for mut copied in [ + record( + plaintext, + "owner-b", + "key-a", + false, + Some(ciphertext.clone()), + ), + record( + plaintext, + "owner-a", + "key-b", + false, + Some(ciphertext.clone()), + ), + record( + plaintext, + "owner-a", + "key-a", + true, + Some(ciphertext.clone()), + ), + ] { + assert!(open_auth_api_key_secret(&state, &copied).is_err()); + copied.key_hash = sha256_hex("different-secret"); + assert!(open_auth_api_key_secret(&state, &copied).is_err()); + } + } + + #[test] + fn legacy_ciphertext_requires_matching_record_hash_before_migration() { + let state = state_with_encryption_key(); + let plaintext = "sk-legacy-record-secret"; + let ciphertext = encrypt_catalog_secret_with_fallbacks(&state, plaintext) + .expect("legacy secret should encrypt"); + let source = record( + plaintext, + "owner-a", + "key-a", + false, + Some(ciphertext.clone()), + ); + let projection = + open_auth_api_key_secret(&state, &source).expect("legacy secret should open"); + assert_eq!(projection.plaintext, plaintext); + assert!(projection.migration_required); + assert!(projection + .protected + .starts_with("aether-auth-api-key-secret-v2:")); + + let copied = record( + "another-record-secret", + "owner-b", + "key-b", + false, + Some(ciphertext), + ); + assert!(open_auth_api_key_secret(&state, &copied).is_err()); + } + + #[test] + fn foreign_and_unknown_envelopes_never_fall_back_to_legacy_decryption() { + let state = state_with_encryption_key(); + for ciphertext in [ + "aether-auth-api-key-secret-v3:unknown", + "aether-system-config-secret-v2:unknown", + "plaintext-secret", + ] { + let record = record( + "plaintext-secret", + "owner-a", + "key-a", + false, + Some(ciphertext.to_string()), + ); + assert!(open_auth_api_key_secret(&state, &record).is_err()); + } + } +} diff --git a/apps/aether-gateway/src/handlers/shared/catalog.rs b/apps/aether-gateway/src/handlers/shared/catalog.rs index 04dbe1a9c..c0649d493 100644 --- a/apps/aether-gateway/src/handlers/shared/catalog.rs +++ b/apps/aether-gateway/src/handlers/shared/catalog.rs @@ -1,4 +1,5 @@ -use crate::handlers::shared::{json_string_list, unix_secs_to_rfc3339}; +use crate::handlers::shared::{json_string_list, masked_secret_display, unix_secs_to_rfc3339}; +use crate::model_fetch::safe_model_fetch_error; use crate::provider_key_auth::{ provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization, provider_key_auth_semantics, provider_key_can_export_oauth, provider_key_can_refresh_oauth, @@ -6,10 +7,17 @@ use crate::provider_key_auth::{ }; use crate::AppState; use aether_admin::provider::quota as admin_provider_quota_pure; +use aether_admin::provider::redaction::{ + admin_provider_oauth_invalid_reason_safe_text, admin_provider_status_snapshot_safe_json, + admin_provider_upstream_metadata_safe_json, admin_secret_safe_json, admin_secret_safe_proxy, +}; use aether_admin::provider::status as admin_provider_status_pure; #[cfg(test)] use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; -use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext}; +use aether_crypto::{ + decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, + looks_like_python_fernet_ciphertext, +}; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; use aether_provider_pool::{ grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier, @@ -24,6 +32,12 @@ const OAUTH_EXPIRED_PREFIX: &str = "[OAUTH_EXPIRED] "; const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] "; const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] "; +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum StoredCatalogSecret { + Encrypted(String), + LegacyPlaintext(String), +} + pub(crate) fn provider_catalog_key_supports_format( key: &StoredProviderCatalogKey, provider_type: &str, @@ -73,8 +87,29 @@ pub(crate) fn decrypt_catalog_secret_with_fallbacks( None } +pub(crate) fn decrypt_catalog_secret_or_legacy_plaintext( + encryption_key: Option<&str>, + stored_value: &str, +) -> Result { + if let Some(plaintext) = decrypt_catalog_secret_with_fallbacks(encryption_key, stored_value) { + return Ok(StoredCatalogSecret::Encrypted(plaintext)); + } + if looks_like_python_fernet_ciphertext(stored_value) { + return Err(()); + } + Ok(StoredCatalogSecret::LegacyPlaintext( + stored_value.to_string(), + )) +} + pub(crate) fn effective_catalog_encryption_key(state: &AppState) -> Option> { - let encryption_key = state.encryption_key().map(str::trim).unwrap_or(""); + effective_catalog_encryption_key_from_config(state.encryption_key()) +} + +pub(crate) fn effective_catalog_encryption_key_from_config( + encryption_key: Option<&str>, +) -> Option> { + let encryption_key = encryption_key.map(str::trim).unwrap_or(""); if !encryption_key.is_empty() { return Some(Cow::Borrowed(encryption_key)); } @@ -103,7 +138,14 @@ pub(crate) fn encrypt_catalog_secret_with_fallbacks( state: &AppState, plaintext: &str, ) -> Option { - let encryption_key = effective_catalog_encryption_key(state)?; + encrypt_catalog_secret_with_configured_key_fallbacks(state.encryption_key(), plaintext) +} + +pub(crate) fn encrypt_catalog_secret_with_configured_key_fallbacks( + encryption_key: Option<&str>, + plaintext: &str, +) -> Option { + let encryption_key = effective_catalog_encryption_key_from_config(encryption_key)?; encrypt_python_fernet_plaintext(encryption_key.as_ref(), plaintext).ok() } @@ -151,18 +193,11 @@ pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatal else { return "[未设置]".to_string(); }; - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - .map(|value| { - if value.chars().count() <= 12 { - format!("{value}***") - } else { - format!( - "{}***{}", - take_secret_prefix(&value, 8), - take_secret_suffix(&value, 4) - ) - } - }) + state + .decrypt_provider_catalog_key_api_key(key) + .ok() + .flatten() + .map(|value| masked_secret_display(&value, 8, 4, "***")) .unwrap_or_else(|| "***ERROR***".to_string()) } } @@ -189,7 +224,10 @@ pub(crate) fn parse_catalog_auth_config_json( if ciphertext.is_empty() { return None; } - let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)?; + let plaintext = state + .decrypt_provider_catalog_key_auth_config(key) + .ok() + .flatten()?; serde_json::from_str::(&plaintext) .ok()? .as_object() @@ -399,7 +437,9 @@ pub(crate) fn sync_provider_key_oauth_status_snapshot( "oauth".to_string(), build_provider_key_oauth_status_snapshot(key), ); - Some(Value::Object(snapshot)) + Some(admin_provider_status_snapshot_safe_json(Some( + &Value::Object(snapshot), + ))) } fn build_provider_key_account_status_snapshot( @@ -654,6 +694,112 @@ fn antigravity_model_quota_window_snapshot( Some(window) } +fn antigravity_grouped_quota_window_snapshots( + metadata: &Map, + observed_at_unix_secs: Option, +) -> Vec { + let mut windows = Vec::new(); + let Some(groups) = metadata.get("quota_groups").and_then(Value::as_array) else { + return windows; + }; + + for (group_index, group) in groups.iter().filter_map(Value::as_object).enumerate() { + let group_code = group + .get("group_id") + .or_else(|| group.get("groupId")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| format!("group:{group_index}")); + let group_label = group + .get("display_name") + .or_else(|| group.get("displayName")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| format!("Quota group {}", group_index + 1)); + let group_description = group + .get("description") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + + for (bucket_index, bucket) in group + .get("buckets") + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter_map(Value::as_object) + .enumerate() + { + let bucket_id = bucket + .get("bucket_id") + .or_else(|| bucket.get("bucketId")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| format!("bucket-{}", bucket_index + 1)); + let Some(mut window) = + model_quota_window_snapshot(&bucket_id, bucket, observed_at_unix_secs) + else { + continue; + }; + let Some(window) = window.as_object_mut() else { + continue; + }; + let period = bucket + .get("window") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let bucket_detail = bucket + .get("display_name") + .or_else(|| bucket.get("displayName")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| period.clone()); + let bucket_label = bucket_detail + .filter(|detail| !detail.eq_ignore_ascii_case(&group_label)) + .map(|detail| format!("{group_label} · {detail}")) + .unwrap_or_else(|| group_label.clone()); + let description = bucket + .get("description") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| group_description.clone()); + + window.insert( + "code".to_string(), + json!(format!("group:{group_index}:{bucket_id}")), + ); + window.insert("label".to_string(), json!(bucket_label)); + window.insert("scope".to_string(), json!("quota_group")); + window.remove("model"); + window.insert("quota_group".to_string(), json!(group_code)); + window.insert("quota_group_label".to_string(), json!(group_label)); + window.insert("bucket_id".to_string(), json!(bucket_id)); + if let Some(period) = period { + window.insert("window".to_string(), json!(period)); + } + if let Some(description) = description { + window.insert("description".to_string(), json!(description)); + } + windows.push(Value::Object(window.clone())); + } + } + + windows +} + fn provider_quota_metadata_string( metadata: &Map, fields: &[&str], @@ -989,7 +1135,7 @@ fn build_codex_quota_status_snapshot( .and_then(admin_provider_quota_pure::coerce_json_bool); let reset_credits = build_codex_reset_credits_status_snapshot(metadata, observed_at_unix_secs); - let windows = [ + let mut windows = [ codex_quota_window_snapshot(metadata, "primary", "weekly", "周", observed_at_unix_secs), codex_quota_window_snapshot(metadata, "secondary", "5h", "5H", observed_at_unix_secs), codex_quota_window_snapshot( @@ -1010,6 +1156,17 @@ fn build_codex_quota_status_snapshot( .into_iter() .flatten() .collect::>(); + if let Some(additional_windows) = metadata + .get("additional_quota_windows") + .and_then(Value::as_array) + { + windows.extend( + additional_windows + .iter() + .filter(|window| window.is_object()) + .cloned(), + ); + } if windows.is_empty() && plan_type.is_none() @@ -1611,7 +1768,7 @@ fn build_antigravity_quota_status_snapshot( .map(str::trim) .filter(|value| !value.is_empty()) .map(ToOwned::to_owned); - let windows = provider_quota_model_bucket(metadata) + let mut windows = provider_quota_model_bucket(metadata) .map(|models| { models .iter() @@ -1625,6 +1782,13 @@ fn build_antigravity_quota_status_snapshot( .collect::>() }) .unwrap_or_default(); + let grouped_observed_at_unix_secs = + provider_quota_timestamp_unix_secs(metadata.get("quota_groups_updated_at")) + .or(observed_at_unix_secs); + windows.extend(antigravity_grouped_quota_window_snapshots( + metadata, + grouped_observed_at_unix_secs, + )); if windows.is_empty() && observed_at_unix_secs.is_none() && !is_forbidden { return None; @@ -2107,7 +2271,9 @@ pub(crate) fn sync_provider_key_quota_status_snapshot( .or_else(|| default_snapshot.as_object().cloned()) .unwrap_or_default(); snapshot.insert("quota".to_string(), quota); - Some(Value::Object(snapshot)) + Some(admin_provider_status_snapshot_safe_json(Some( + &Value::Object(snapshot), + ))) } fn quota_snapshot_has_materialized_data( @@ -2290,7 +2456,7 @@ pub(crate) fn provider_key_status_snapshot_payload( "account".to_string(), build_provider_key_account_status_snapshot(key, provider_type), ); - Value::Object(snapshot) + admin_provider_status_snapshot_safe_json(Some(&Value::Object(snapshot))) } pub(crate) fn provider_key_health_summary( @@ -2606,7 +2772,10 @@ pub(crate) fn build_admin_provider_key_response( ); payload.insert("oauth_header_auth".to_string(), json!(oauth_header_auth)); payload.insert("name".to_string(), json!(key.name)); - payload.insert("rate_multipliers".to_string(), json!(key.rate_multipliers)); + payload.insert( + "rate_multipliers".to_string(), + admin_secret_safe_json(key.rate_multipliers.as_ref()), + ); payload.insert( "internal_priority".to_string(), json!(key.internal_priority), @@ -2626,7 +2795,10 @@ pub(crate) fn build_admin_provider_key_response( .collect(), ), ); - payload.insert("capabilities".to_string(), json!(key.capabilities)); + payload.insert( + "capabilities".to_string(), + admin_secret_safe_json(key.capabilities.as_ref()), + ); payload.insert( "oauth_expires_at".to_string(), json!(auth_semantics @@ -2685,7 +2857,7 @@ pub(crate) fn build_admin_provider_key_response( ); payload.insert( "oauth_organizations".to_string(), - serde_json::Value::Array(oauth_organizations), + admin_secret_safe_json(Some(&serde_json::Value::Array(oauth_organizations))), ); payload.insert("oauth_temporary".to_string(), json!(oauth_temporary)); payload.insert( @@ -2699,12 +2871,17 @@ pub(crate) fn build_admin_provider_key_response( "oauth_invalid_reason".to_string(), json!(auth_semantics .can_show_oauth_metadata() - .then_some(key.oauth_invalid_reason.clone()) + .then(|| { + admin_provider_oauth_invalid_reason_safe_text(key.oauth_invalid_reason.as_deref()) + }) .flatten()), ); payload.insert( "status_snapshot".to_string(), - provider_key_status_snapshot_payload(key, provider_type), + admin_provider_status_snapshot_safe_json(Some(&provider_key_status_snapshot_payload( + key, + provider_type, + ))), ); payload.insert( "cache_ttl_minutes".to_string(), @@ -2714,10 +2891,13 @@ pub(crate) fn build_admin_provider_key_response( "max_probe_interval_minutes".to_string(), json!(key.max_probe_interval_minutes), ); - payload.insert("health_by_format".to_string(), json!(key.health_by_format)); + payload.insert( + "health_by_format".to_string(), + admin_secret_safe_json(key.health_by_format.as_ref()), + ); payload.insert( "circuit_breaker_by_format".to_string(), - json!(key.circuit_breaker_by_format), + admin_secret_safe_json(key.circuit_breaker_by_format.as_ref()), ); payload.insert("health_score".to_string(), json!(health_score)); payload.insert( @@ -2766,10 +2946,9 @@ pub(crate) fn build_admin_provider_key_response( ); payload.insert( "request_results_window".to_string(), - circuit_sample - .and_then(|value| value.get("request_results_window")) - .cloned() - .unwrap_or(serde_json::Value::Null), + admin_secret_safe_json( + circuit_sample.and_then(|value| value.get("request_results_window")), + ), ); payload.insert("request_count".to_string(), json!(request_count)); payload.insert("success_count".to_string(), json!(success_count)); @@ -2788,7 +2967,7 @@ pub(crate) fn build_admin_provider_key_response( payload.insert("effective_limit".to_string(), json!(effective_limit)); payload.insert( "utilization_samples".to_string(), - json!(key.utilization_samples), + admin_secret_safe_json(key.utilization_samples.as_ref()), ); payload.insert( "last_probe_increase_at".to_string(), @@ -2819,7 +2998,10 @@ pub(crate) fn build_admin_provider_key_response( ); payload.insert( "last_models_fetch_error".to_string(), - json!(key.last_models_fetch_error), + json!(key + .last_models_fetch_error + .as_deref() + .map(safe_model_fetch_error)), ); payload.insert("locked_models".to_string(), json!(key.locked_models)); payload.insert( @@ -2832,10 +3014,16 @@ pub(crate) fn build_admin_provider_key_response( ); payload.insert( "upstream_metadata".to_string(), - json!(key.upstream_metadata), + admin_provider_upstream_metadata_safe_json(key.upstream_metadata.as_ref()), + ); + payload.insert( + "proxy".to_string(), + admin_secret_safe_proxy(key.proxy.as_ref()), + ); + payload.insert( + "fingerprint".to_string(), + admin_secret_safe_json(key.fingerprint.as_ref()), ); - payload.insert("proxy".to_string(), json!(key.proxy)); - payload.insert("fingerprint".to_string(), json!(key.fingerprint)); payload.insert( "last_used_at".to_string(), json!(key.last_used_at_unix_secs.and_then(unix_secs_to_rfc3339)), @@ -2938,6 +3126,40 @@ mod tests { assert_ne!(masked, "***ERROR***"); } + #[test] + fn masked_catalog_api_key_never_returns_a_complete_short_secret() { + let state = AppState::new().expect("gateway should build"); + let plaintext = "sk-test-a"; + let encrypted_api_key = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, plaintext) + .expect("api key ciphertext should build"); + let key = StoredProviderCatalogKey::new( + "key-short".to_string(), + "provider-test".to_string(), + "default".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + Some(json!(["openai:chat"])), + encrypted_api_key, + None, + None, + None, + None, + None, + None, + None, + ) + .expect("key transport should build"); + + let masked = masked_catalog_api_key(&state, &key); + assert_ne!(masked, plaintext); + assert!(!masked.contains(plaintext)); + } + #[test] fn provider_aware_mask_labels_agent_identity_without_exposing_placeholder() { let state = AppState::new().expect("gateway should build"); @@ -3121,6 +3343,52 @@ mod tests { assert_eq!(spark_weekly.get("remaining_ratio"), Some(&json!(0.95))); } + #[test] + fn codex_wham_snapshot_does_not_duplicate_normalized_spark_windows() { + let codex = admin_provider_quota_pure::parse_codex_wham_usage_response( + &json!({ + "plan_type": "plus", + "rate_limit": { + "primary_window": {"used_percent": 25.0}, + "secondary_window": {"used_percent": 10.0} + }, + "additional_rate_limits": [{ + "limit_name": "GPT-5.3-Codex-Spark", + "rate_limit": { + "primary_window": { + "used_percent": 40.0, + "limit_window_seconds": 18_000 + }, + "secondary_window": { + "used_percent": 5.0, + "limit_window_seconds": 604_800 + } + } + }] + }), + 1_777_000_000, + ) + .expect("Codex WHAM quota should parse"); + let upstream_metadata = json!({"codex": codex}); + let payload = sync_provider_key_quota_status_snapshot( + None, + "codex", + Some(&upstream_metadata), + "refresh_api", + ) + .expect("Codex quota snapshot should sync"); + let windows = payload["quota"]["windows"] + .as_array() + .expect("quota windows should exist"); + let spark_codes = windows + .iter() + .filter_map(|window| window.get("code").and_then(Value::as_str)) + .filter(|code| code.starts_with("spark_")) + .collect::>(); + + assert_eq!(spark_codes, vec!["spark_5h", "spark_weekly"]); + } + #[test] fn provider_key_status_snapshot_payload_keeps_codex_free_window_quota_available() { let mut key = sample_catalog_key(); @@ -4155,6 +4423,68 @@ mod tests { assert_eq!(claude_window.get("reset_seconds"), Some(&json!(12_011u64))); } + #[test] + fn sync_provider_key_quota_status_snapshot_materializes_antigravity_group_windows() { + let upstream_metadata = json!({ + "antigravity": { + "updated_at": 1_777_000_000u64, + "models": { + "gemini-3.7-flash-tiered": { + "remaining_fraction": 0.9 + } + }, + "quota_groups": [{ + "display_name": "Claude and GPT models", + "description": "Shared quota", + "buckets": [{ + "bucket_id": "3p-5h", + "window": "5h", + "remaining_fraction": 0.25, + "reset_time": "2026-05-05T05:00:00Z", + "display_name": "5 hour" + }, { + "bucket_id": "3p-weekly", + "window": "weekly", + "remaining_fraction": 0.8, + "reset_time": "2026-05-11T00:00:00Z" + }] + }] + } + }); + + let payload = sync_provider_key_quota_status_snapshot( + None, + "antigravity", + Some(&upstream_metadata), + "refresh_api", + ) + .expect("Antigravity quota snapshot should sync"); + let windows = payload["quota"]["windows"] + .as_array() + .expect("quota windows should exist"); + let five_hour = windows + .iter() + .find(|window| window["code"] == "group:0:3p-5h") + .expect("5h grouped quota window should exist"); + let weekly = windows + .iter() + .find(|window| window["code"] == "group:0:3p-weekly") + .expect("weekly grouped quota window should exist"); + + assert_eq!(windows.len(), 3); + assert_eq!(five_hour["scope"], json!("quota_group")); + assert_eq!(five_hour["bucket_id"], json!("3p-5h")); + assert_eq!(five_hour["label"], json!("Claude and GPT models · 5 hour")); + assert_eq!( + five_hour["quota_group_label"], + json!("Claude and GPT models") + ); + assert_eq!(five_hour["remaining_ratio"], json!(0.25)); + assert_eq!(five_hour["used_ratio"], json!(0.75)); + assert!(five_hour.get("model").is_none()); + assert_eq!(weekly["window"], json!("weekly")); + } + #[test] fn provider_key_status_snapshot_payload_backfills_account_block_from_oauth_invalid_reason() { let mut key = sample_catalog_key(); @@ -4168,10 +4498,7 @@ mod tests { assert_eq!(account.get("code"), Some(&json!("account_disabled"))); assert_eq!(account.get("label"), Some(&json!("账号停用"))); - assert_eq!( - account.get("reason"), - Some(&json!("account has been deactivated")) - ); + assert_eq!(account.get("reason"), Some(&json!("Account is disabled"))); assert_eq!(account.get("blocked"), Some(&json!(true))); assert_eq!(account.get("source"), Some(&json!("oauth_invalid"))); } @@ -4217,4 +4544,79 @@ mod tests { assert_eq!(account.get("blocked"), Some(&json!(true))); assert_eq!(account.get("source"), Some(&json!("metadata"))); } + + #[test] + fn admin_provider_key_response_projects_historical_sensitive_diagnostics() { + let state = AppState::new().expect("gateway should build"); + let mut key = sample_catalog_key(); + key.auth_type = "oauth".to_string(); + key.oauth_invalid_at_unix_secs = Some(1_777_000_000); + key.oauth_invalid_reason = Some( + "[ACCOUNT_BLOCK] account has been deactivated: Authorization: Bearer upstream-secret https://user:password@internal.test?q=secret" + .to_string(), + ); + key.last_models_fetch_error = Some( + "request failed for https://user:password@internal.test/models?q=secret; Authorization: Bearer upstream-secret" + .to_string(), + ); + key.status_snapshot = Some(json!({ + "oauth": { + "code": "invalid", + "reason": "Authorization: Bearer upstream-secret" + }, + "account": { + "code": "account_disabled", + "reason": "https://user:password@internal.test?q=secret", + "blocked": true + }, + "quota": { + "provider_type": "codex", + "code": "cooldown", + "reason": "Authorization: Bearer upstream-secret", + "exhausted": false, + "reset_credits": { + "detail_error": "https://user:password@internal.test?q=secret" + }, + "unknown": {"body": "upstream-secret"} + } + })); + key.upstream_metadata = Some(json!({ + "codex": { + "primary_used_percent": 25.0, + "message": "Authorization: Bearer upstream-secret", + "reset_credits": { + "detail_error": "https://user:password@internal.test?q=secret" + } + } + })); + + let payload = build_admin_provider_key_response( + &state, + &key, + "codex", + &["openai:responses".to_string()], + 1_777_000_001, + ); + + assert_eq!( + payload["oauth_invalid_reason"], + json!("[ACCOUNT_BLOCK] Account is disabled") + ); + assert_eq!( + payload["last_models_fetch_error"], + json!("Upstream models fetch failed") + ); + assert_eq!( + payload.pointer("/status_snapshot/account/reason"), + Some(&json!("Account is disabled")) + ); + assert_eq!( + payload.pointer("/upstream_metadata/codex/primary_used_percent"), + Some(&json!(25.0)) + ); + let serialized = payload.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); + assert!(!serialized.contains("q=secret")); + } } diff --git a/apps/aether-gateway/src/handlers/shared/email_templates.rs b/apps/aether-gateway/src/handlers/shared/email_templates.rs index 91bd63cc9..2458b37c8 100644 --- a/apps/aether-gateway/src/handlers/shared/email_templates.rs +++ b/apps/aether-gateway/src/handlers/shared/email_templates.rs @@ -1,7 +1,14 @@ use super::system_config_string; use crate::{AppState, GatewayError}; +use aether_admin::system::{ + admin_email_template_html_is_valid, admin_email_template_subject_is_valid, + ADMIN_EMAIL_TEMPLATE_MAX_HTML_BYTES, ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES, +}; +use regex::Regex; use serde_json::json; +const MAX_RENDERED_EMAIL_HTML_BYTES: usize = 512 * 1024; + pub(crate) struct AdminEmailTemplateDefinition { pub(crate) template_type: &'static str, pub(crate) name: &'static str, @@ -173,10 +180,24 @@ pub(crate) async fn read_admin_email_template_payload( let html = state .read_system_config_json_value(&admin_email_template_html_key(definition.template_type)) .await?; - let subject = system_config_string(subject.as_ref()) - .unwrap_or_else(|| definition.default_subject.to_string()); - let html = - system_config_string(html.as_ref()).unwrap_or_else(|| definition.default_html.to_string()); + let subject = match system_config_string(subject.as_ref()) { + Some(value) if admin_email_template_subject_is_valid(&value) => value, + Some(_) => { + return Err(GatewayError::Internal( + "stored email template subject is invalid or oversized".to_string(), + )); + } + None => definition.default_subject.to_string(), + }; + let html = match system_config_string(html.as_ref()) { + Some(value) if admin_email_template_html_is_valid(&value) => value, + Some(_) => { + return Err(GatewayError::Internal( + "stored email template html is invalid or oversized".to_string(), + )); + } + None => definition.default_html.to_string(), + }; let is_custom = subject != definition.default_subject || html != definition.default_html; Ok(Some(json!({ @@ -204,13 +225,76 @@ pub(crate) fn render_admin_email_template_html( template_html: &str, variables: &std::collections::BTreeMap, ) -> Result { + if template_html.len() > ADMIN_EMAIL_TEMPLATE_MAX_HTML_BYTES { + return Err(GatewayError::Internal( + "email template html exceeds the allowed size".to_string(), + )); + } + if !admin_email_template_html_is_valid(template_html) { + return Err(GatewayError::Internal( + "email template html contains invalid control bytes".to_string(), + )); + } let mut rendered = template_html.to_string(); for (key, value) in variables { - let pattern = regex::Regex::new(&format!(r"\{{\{{\s*{}\s*\}}\}}", regex::escape(key))) + if value.len() > ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES { + return Err(GatewayError::Internal( + "email template variable is oversized".to_string(), + )); + } + let pattern = Regex::new(&format!(r"\{{\{{\s*{}\s*\}}\}}", regex::escape(key))) .map_err(|err| GatewayError::Internal(err.to_string()))?; + let escaped = escape_admin_email_template_html(value); + let (matched_bytes, occurrences) = pattern + .find_iter(&rendered) + .fold((0usize, 0usize), |(matched, count), found| { + (matched.saturating_add(found.as_str().len()), count + 1) + }); + let replacement_bytes = occurrences.checked_mul(escaped.len()).and_then(|bytes| { + rendered + .len() + .checked_sub(matched_bytes)? + .checked_add(bytes) + }); + if replacement_bytes.is_none_or(|bytes| bytes > MAX_RENDERED_EMAIL_HTML_BYTES) { + return Err(GatewayError::Internal( + "rendered email template exceeds the allowed size".to_string(), + )); + } rendered = pattern - .replace_all(&rendered, escape_admin_email_template_html(value)) + .replace_all(&rendered, regex::NoExpand(escaped.as_str())) .into_owned(); } + if rendered.len() > MAX_RENDERED_EMAIL_HTML_BYTES { + return Err(GatewayError::Internal( + "rendered email template exceeds the allowed size".to_string(), + )); + } Ok(rendered) } + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::BTreeMap; + + #[test] + fn renderer_escapes_values_without_regex_replacement_expansion() { + let variables = BTreeMap::from([(String::from("app_name"), String::from("$1<&"))]); + let rendered = render_admin_email_template_html("

{{app_name}}

", &variables) + .expect("normal template should render"); + assert_eq!(rendered, "

$1<&

"); + } + + #[test] + fn renderer_rejects_control_bytes_and_expansion_bombs() { + let controls = render_admin_email_template_html("

bad\u{0001}

", &BTreeMap::new()); + assert!(controls.is_err()); + + let template = "{{value}}".repeat(100_000); + let variables = BTreeMap::from([(String::from("value"), String::from("x".repeat(64)))]); + let error = render_admin_email_template_html(&template, &variables) + .expect_err("rendered output must remain bounded"); + assert!(format!("{error:?}").contains("exceeds")); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/identity_oauth_provider_secret.rs b/apps/aether-gateway/src/handlers/shared/identity_oauth_provider_secret.rs new file mode 100644 index 000000000..94d62c00d --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/identity_oauth_provider_secret.rs @@ -0,0 +1,634 @@ +use std::future::Future; + +use aether_crypto::looks_like_python_fernet_ciphertext; +use aether_data::repository::oauth_providers::{ + validate_oauth_redirect_uri, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord, +}; +use url::{Host, Url}; + +use crate::{AppState, GatewayError}; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +const IDENTITY_OAUTH_CLIENT_SECRET_MIGRATION_RETRIES: usize = 8; +const IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_FAMILY: &str = "aether-identity-oauth-client-secret-"; +const IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V2: &str = "aether-identity-oauth-client-secret-v2:"; +const IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3: &str = "aether-identity-oauth-client-secret-v3:"; +const IDENTITY_OAUTH_CLIENT_SECRET_PURPOSE_V2: &str = "identity-oauth-client-secret-bound-v2"; +const IDENTITY_OAUTH_CLIENT_SECRET_PURPOSE_V3: &str = "identity-oauth-client-secret-bound-v3"; +const IDENTITY_OAUTH_CLIENT_SECRET_FIELD: &str = "client_secret_encrypted"; +const LINUXDO_AUTHORIZATION_URL: &str = "https://connect.linux.do/oauth2/authorize"; +const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token"; +const LINUXDO_USERINFO_URL: &str = "https://connect.linux.do/api/user"; + +#[derive(Clone, PartialEq, Eq)] +struct IdentityOAuthClientSecretProjection { + plaintext: String, + protected: String, + migration_required: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct IdentityOAuthClientSecretBinding { + provider_type: String, + client_id: String, + authorization_url: String, + token_url: String, + userinfo_url: String, + redirect_uri: String, +} + +fn normalized_identity_oauth_provider_type(provider_type: &str) -> Result { + let provider_type = provider_type.trim().to_ascii_lowercase(); + if provider_type.is_empty() { + return Err("identity OAuth provider type is empty"); + } + if provider_type.contains('\0') { + return Err("identity OAuth provider type contains reserved framing"); + } + Ok(provider_type) +} + +fn identity_oauth_client_secret_purpose_v2(provider_type: &str) -> Result { + let provider_type = normalized_identity_oauth_provider_type(provider_type)?; + Ok(format!( + "{IDENTITY_OAUTH_CLIENT_SECRET_PURPOSE_V2}\0provider-type-bytes={}\0{provider_type}\0field-bytes={}\0{IDENTITY_OAUTH_CLIENT_SECRET_FIELD}", + provider_type.len(), + IDENTITY_OAUTH_CLIENT_SECRET_FIELD.len(), + )) +} + +fn canonical_identity_oauth_endpoint(raw: &str) -> Result { + if raw.contains('\0') { + return Err("identity OAuth endpoint contains reserved framing"); + } + let mut parsed = Url::parse(raw.trim()).map_err(|_| "identity OAuth endpoint is invalid")?; + if parsed.scheme() != "https" + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.fragment().is_some() + || matches!(parsed.host(), Some(Host::Ipv4(_)) | Some(Host::Ipv6(_))) + { + return Err("identity OAuth endpoint is not a canonical HTTPS DNS URL"); + } + let host = parsed + .host_str() + .map(|host| host.trim_end_matches('.').to_ascii_lowercase()) + .filter(|host| !host.is_empty()) + .ok_or("identity OAuth endpoint is missing a host")?; + parsed + .set_host(Some(&host)) + .map_err(|_| "identity OAuth endpoint host is invalid")?; + if parsed.port() == Some(443) { + parsed + .set_port(None) + .map_err(|_| "identity OAuth endpoint port is invalid")?; + } + Ok(parsed.to_string()) +} + +fn canonical_identity_oauth_redirect_uri(raw: &str) -> Result { + if raw.contains('\0') { + return Err("identity OAuth redirect URI contains reserved framing"); + } + let raw = raw.trim(); + validate_oauth_redirect_uri(raw).map_err(|_| "identity OAuth redirect URI is invalid")?; + Url::parse(raw) + .map(|url| url.to_string()) + .map_err(|_| "identity OAuth redirect URI is invalid") +} + +fn effective_identity_oauth_endpoint<'a>( + provider_type: &str, + override_value: Option<&'a str>, + linuxdo_default: &'static str, +) -> Result<&'a str, &'static str> { + if let Some(value) = override_value + .map(str::trim) + .filter(|value| !value.is_empty()) + { + return Ok(value); + } + if provider_type == "linuxdo" { + // The static default can be shortened to the caller lifetime. + return Ok(linuxdo_default); + } + Err("identity OAuth endpoint is missing") +} + +fn identity_oauth_client_secret_binding( + provider_type: &str, + client_id: &str, + authorization_url_override: Option<&str>, + token_url_override: Option<&str>, + userinfo_url_override: Option<&str>, + redirect_uri: &str, +) -> Result { + let provider_type = normalized_identity_oauth_provider_type(provider_type)?; + let client_id = client_id.trim(); + if client_id.is_empty() { + return Err("identity OAuth client ID is empty"); + } + if client_id.contains('\0') { + return Err("identity OAuth client ID contains reserved framing"); + } + let authorization_url = effective_identity_oauth_endpoint( + &provider_type, + authorization_url_override, + LINUXDO_AUTHORIZATION_URL, + )?; + let token_url = + effective_identity_oauth_endpoint(&provider_type, token_url_override, LINUXDO_TOKEN_URL)?; + let userinfo_url = effective_identity_oauth_endpoint( + &provider_type, + userinfo_url_override, + LINUXDO_USERINFO_URL, + )?; + Ok(IdentityOAuthClientSecretBinding { + provider_type, + client_id: client_id.to_string(), + authorization_url: canonical_identity_oauth_endpoint(authorization_url)?, + token_url: canonical_identity_oauth_endpoint(token_url)?, + userinfo_url: canonical_identity_oauth_endpoint(userinfo_url)?, + redirect_uri: canonical_identity_oauth_redirect_uri(redirect_uri)?, + }) +} + +fn stored_identity_oauth_client_secret_binding( + provider: &StoredOAuthProviderConfig, +) -> Result { + identity_oauth_client_secret_binding( + &provider.provider_type, + &provider.client_id, + provider.authorization_url_override.as_deref(), + provider.token_url_override.as_deref(), + provider.userinfo_url_override.as_deref(), + &provider.redirect_uri, + ) +} + +fn upsert_identity_oauth_client_secret_binding( + provider: &UpsertOAuthProviderConfigRecord, +) -> Result { + identity_oauth_client_secret_binding( + &provider.provider_type, + &provider.client_id, + provider.authorization_url_override.as_deref(), + provider.token_url_override.as_deref(), + provider.userinfo_url_override.as_deref(), + &provider.redirect_uri, + ) +} + +fn identity_oauth_client_secret_purpose_v3(binding: &IdentityOAuthClientSecretBinding) -> String { + format!( + "{IDENTITY_OAUTH_CLIENT_SECRET_PURPOSE_V3}\0provider-type-bytes={}\0{}\0client-id-bytes={}\0{}\0authorization-url-bytes={}\0{}\0token-url-bytes={}\0{}\0userinfo-url-bytes={}\0{}\0redirect-uri-bytes={}\0{}\0field-bytes={}\0{IDENTITY_OAUTH_CLIENT_SECRET_FIELD}", + binding.provider_type.len(), + binding.provider_type, + binding.client_id.len(), + binding.client_id, + binding.authorization_url.len(), + binding.authorization_url, + binding.token_url.len(), + binding.token_url, + binding.userinfo_url.len(), + binding.userinfo_url, + binding.redirect_uri.len(), + binding.redirect_uri, + IDENTITY_OAUTH_CLIENT_SECRET_FIELD.len(), + ) +} + +pub(crate) fn identity_oauth_provider_secret_binding_matches( + stored: &StoredOAuthProviderConfig, + replacement: &UpsertOAuthProviderConfigRecord, +) -> Result { + Ok(stored_identity_oauth_client_secret_binding(stored)? + == upsert_identity_oauth_client_secret_binding(replacement)?) +} + +pub(crate) fn seal_identity_oauth_provider_client_secret( + state: &AppState, + provider: &UpsertOAuthProviderConfigRecord, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("identity OAuth client secret contains reserved framing"); + } + let binding = upsert_identity_oauth_client_secret_binding(provider)?; + let purpose = identity_oauth_client_secret_purpose_v3(&binding); + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("identity OAuth client secret encryption key is not configured")?; + Ok(format!( + "{IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3}{sealed}" + )) +} + +fn seal_identity_oauth_provider_client_secret_for_binding( + state: &AppState, + binding: &IdentityOAuthClientSecretBinding, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("identity OAuth client secret contains reserved framing"); + } + let purpose = identity_oauth_client_secret_purpose_v3(binding); + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("identity OAuth client secret encryption key is not configured")?; + Ok(format!( + "{IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3}{sealed}" + )) +} + +fn open_identity_oauth_provider_client_secret( + state: &AppState, + provider: &StoredOAuthProviderConfig, + stored: &str, +) -> Result { + let binding = stored_identity_oauth_client_secret_binding(provider)?; + let purpose = identity_oauth_client_secret_purpose_v3(&binding); + if let Some(sealed) = stored.strip_prefix(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3) { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("identity OAuth client secret authentication failed")?; + if plaintext.contains('\0') { + return Err("identity OAuth client secret contains reserved framing"); + } + return Ok(IdentityOAuthClientSecretProjection { + plaintext, + protected: stored.to_string(), + migration_required: false, + }); + } + if let Some(sealed) = stored.strip_prefix(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V2) { + let legacy_purpose = identity_oauth_client_secret_purpose_v2(&provider.provider_type)?; + let plaintext = open_runtime_secret_payload(state, &legacy_purpose, sealed) + .ok_or("identity OAuth client secret authentication failed")?; + if plaintext.contains('\0') { + return Err("identity OAuth client secret contains reserved framing"); + } + let protected = + seal_identity_oauth_provider_client_secret_for_binding(state, &binding, &plaintext)?; + return Ok(IdentityOAuthClientSecretProjection { + plaintext, + protected, + migration_required: true, + }); + } + if stored.starts_with(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_FAMILY) { + return Err("unsupported identity OAuth client secret envelope"); + } + if stored.starts_with("aether-") { + return Err("Aether secret envelope has the wrong record binding"); + } + if !looks_like_python_fernet_ciphertext(stored) { + return Err("identity OAuth client secret is not an authenticated ciphertext"); + } + + let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("legacy identity OAuth client secret authentication failed")?; + if plaintext.contains('\0') { + return Err("legacy identity OAuth client secret contains reserved framing"); + } + let protected = + seal_identity_oauth_provider_client_secret_for_binding(state, &binding, &plaintext)?; + Ok(IdentityOAuthClientSecretProjection { + plaintext, + protected, + migration_required: true, + }) +} + +pub(crate) async fn decrypt_or_migrate_identity_oauth_provider_client_secret( + state: &AppState, + provider: &StoredOAuthProviderConfig, +) -> Result, GatewayError> { + decrypt_or_migrate_identity_oauth_provider_client_secret_with_before_compare( + state, + provider, + || async {}, + ) + .await +} + +async fn decrypt_or_migrate_identity_oauth_provider_client_secret_with_before_compare< + BeforeCompare, + CompareFuture, +>( + state: &AppState, + provider: &StoredOAuthProviderConfig, + before_compare: BeforeCompare, +) -> Result, GatewayError> +where + BeforeCompare: Fn() -> CompareFuture, + CompareFuture: Future, +{ + let provider_storage_key = provider.provider_type.trim(); + let original_binding = + stored_identity_oauth_client_secret_binding(provider).map_err(secret_error)?; + + for _ in 0..IDENTITY_OAUTH_CLIENT_SECRET_MIGRATION_RETRIES { + let current = state + .get_oauth_provider_config(provider_storage_key) + .await? + .ok_or_else(|| secret_error("identity OAuth provider is unavailable"))?; + if stored_identity_oauth_client_secret_binding(¤t).map_err(secret_error)? + != original_binding + { + return Err(secret_error( + "identity OAuth provider record binding changed unexpectedly", + )); + } + let Some(observed) = current.client_secret_encrypted.as_deref() else { + return Ok(None); + }; + if observed.is_empty() { + return Err(secret_error("stored identity OAuth client secret is empty")); + } + let projection = open_identity_oauth_provider_client_secret(state, ¤t, observed) + .map_err(secret_error)?; + if !projection.migration_required { + return Ok(Some(projection.plaintext)); + } + + before_compare().await; + if state + .compare_and_swap_oauth_provider_client_secret( + provider_storage_key, + observed, + &projection.protected, + ) + .await? + { + return Ok(Some(projection.plaintext)); + } + } + + Err(secret_error( + "identity OAuth client secret migration did not stabilize", + )) +} + +fn secret_error(message: &'static str) -> GatewayError { + GatewayError::Internal(message.to_string()) +} + +#[cfg(test)] +mod tests { + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::oauth_providers::{ + EncryptedSecretUpdate, InMemoryOAuthProviderRepository, OAuthProviderReadRepository, + OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord, + }; + + use super::{ + decrypt_or_migrate_identity_oauth_provider_client_secret, + decrypt_or_migrate_identity_oauth_provider_client_secret_with_before_compare, + open_identity_oauth_provider_client_secret, seal_identity_oauth_provider_client_secret, + IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3, + }; + use crate::handlers::shared::{ + encrypt_catalog_secret_with_fallbacks, seal_runtime_secret_payload, + }; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + fn sample_provider(provider_type: &str, encrypted: &str) -> StoredOAuthProviderConfig { + let normalized_provider_type = provider_type.trim().to_ascii_lowercase(); + StoredOAuthProviderConfig::new( + provider_type.to_string(), + format!("{normalized_provider_type} display"), + format!("{normalized_provider_type}-client"), + format!("https://{normalized_provider_type}.example.com/redirect"), + "https://frontend.example.com/auth/callback".to_string(), + ) + .expect("provider should build") + .with_config_fields( + Some(encrypted.to_string()), + Some("https://connect.linux.do/oauth2/authorize".to_string()), + Some("https://connect.linux.do/oauth2/token".to_string()), + None, + Some(vec!["openid".to_string()]), + None, + None, + None, + true, + ) + .with_timestamps(Some(10), Some(20)) + } + + fn sample_upsert( + provider_type: &str, + encrypted: EncryptedSecretUpdate, + display_name: &str, + ) -> UpsertOAuthProviderConfigRecord { + UpsertOAuthProviderConfigRecord { + provider_type: provider_type.to_string(), + display_name: display_name.to_string(), + client_id: format!("{provider_type}-client"), + client_secret_encrypted: encrypted, + authorization_url_override: Some( + "https://connect.linux.do/oauth2/authorize".to_string(), + ), + token_url_override: Some("https://connect.linux.do/oauth2/token".to_string()), + userinfo_url_override: None, + scopes: Some(vec!["openid".to_string()]), + redirect_uri: format!("https://{provider_type}.example.com/redirect"), + frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(), + attribute_mapping: None, + extra_config: None, + icon_url: None, + is_enabled: true, + } + } + + fn sample_binding_upsert() -> UpsertOAuthProviderConfigRecord { + sample_upsert("linuxdo", EncryptedSecretUpdate::Preserve, "Linux.do") + } + + #[test] + fn v2_round_trip_binds_normalized_provider_and_rejects_tampering() { + let state = state_with_encryption_key(); + let record = sample_binding_upsert(); + let sealed = seal_identity_oauth_provider_client_secret(&state, &record, "client-secret") + .expect("client secret should seal"); + assert!(sealed.starts_with(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3)); + let provider = sample_provider(" LinuxDo ", &sealed); + assert_eq!( + open_identity_oauth_provider_client_secret(&state, &provider, &sealed) + .expect("matching provider should open") + .plaintext, + "client-secret" + ); + let wrong_provider = sample_provider("github", &sealed); + assert!( + open_identity_oauth_provider_client_secret(&state, &wrong_provider, &sealed).is_err() + ); + + let mut tampered = sealed.into_bytes(); + let last = tampered + .last_mut() + .expect("sealed value should not be empty"); + *last = if *last == b'A' { b'B' } else { b'A' }; + let tampered = String::from_utf8(tampered).expect("ciphertext should remain UTF-8"); + assert!(open_identity_oauth_provider_client_secret(&state, &provider, &tampered).is_err()); + } + + #[test] + fn reader_rejects_unknown_and_cross_family_aether_envelopes() { + let state = state_with_encryption_key(); + for stored in [ + "aether-identity-oauth-client-secret-v3:unknown", + "aether-system-config-secret-v2:foreign", + "aether-proxy-node-secret-v2:foreign", + "plaintext-secret", + ] { + assert!( + open_identity_oauth_provider_client_secret( + &state, + &sample_provider("linuxdo", stored), + stored + ) + .is_err(), + "unexpectedly accepted {stored}" + ); + } + let other_runtime = seal_runtime_secret_payload(&state, "another-purpose", "secret") + .expect("runtime secret should seal"); + assert!(open_identity_oauth_provider_client_secret( + &state, + &sample_provider("linuxdo", &other_runtime), + &other_runtime, + ) + .is_err()); + let stripped = other_runtime + .strip_prefix("aether-runtime-secret-v1:") + .expect("runtime envelope should contain its Fernet payload"); + assert!( + open_identity_oauth_provider_client_secret( + &state, + &sample_provider("linuxdo", stripped), + stripped + ) + .is_err(), + "stripping a foreign runtime envelope must not turn it into a legacy secret", + ); + } + + #[tokio::test] + async fn legacy_fernet_is_migrated_to_record_bound_v2() { + let bootstrap = state_with_encryption_key(); + let legacy = encrypt_catalog_secret_with_fallbacks(&bootstrap, "legacy-secret") + .expect("legacy secret should encrypt"); + let provider = sample_provider("linuxdo", &legacy); + let repository = Arc::new(InMemoryOAuthProviderRepository::seed([provider.clone()])); + let state = AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::with_oauth_provider_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + + assert_eq!( + decrypt_or_migrate_identity_oauth_provider_client_secret(&state, &provider) + .await + .expect("legacy secret should migrate") + .as_deref(), + Some("legacy-secret") + ); + let stored = repository + .get_oauth_provider_config("linuxdo") + .await + .expect("provider should read") + .expect("provider should exist") + .client_secret_encrypted + .expect("secret should remain configured"); + assert!(stored.starts_with(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3)); + } + + #[tokio::test] + async fn migration_cas_miss_rereads_and_preserves_concurrent_non_secret_update() { + let bootstrap = state_with_encryption_key(); + let legacy_before = encrypt_catalog_secret_with_fallbacks(&bootstrap, "before") + .expect("legacy secret should encrypt"); + let legacy_after = encrypt_catalog_secret_with_fallbacks(&bootstrap, "after") + .expect("rotated legacy secret should encrypt"); + let provider = sample_provider("linuxdo", &legacy_before); + let repository = Arc::new(InMemoryOAuthProviderRepository::seed([provider.clone()])); + let state = AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::with_oauth_provider_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let first_compare = Arc::new(AtomicBool::new(true)); + let repository_for_race = Arc::clone(&repository); + let legacy_after_for_race = legacy_after.clone(); + + let plaintext = + decrypt_or_migrate_identity_oauth_provider_client_secret_with_before_compare( + &state, + &provider, + move || { + let first_compare = Arc::clone(&first_compare); + let repository = Arc::clone(&repository_for_race); + let legacy_after = legacy_after_for_race.clone(); + async move { + if first_compare.swap(false, Ordering::SeqCst) { + repository + .upsert_oauth_provider_config(&sample_upsert( + "linuxdo", + EncryptedSecretUpdate::Set(legacy_after), + "concurrent display", + )) + .await + .expect("concurrent provider update should persist"); + } + } + }, + ) + .await + .expect("migration should retry after CAS miss") + .expect("secret should remain configured"); + assert_eq!(plaintext, "after"); + + let current = repository + .get_oauth_provider_config("linuxdo") + .await + .expect("provider should read") + .expect("provider should exist"); + assert_eq!(current.display_name, "concurrent display"); + let stored = current + .client_secret_encrypted + .expect("secret should remain configured"); + assert!(stored.starts_with(IDENTITY_OAUTH_CLIENT_SECRET_ENVELOPE_V3)); + assert_eq!( + open_identity_oauth_provider_client_secret( + &state, + &repository + .get_oauth_provider_config("linuxdo") + .await + .unwrap() + .unwrap(), + &stored + ) + .expect("migrated secret should open") + .plaintext, + "after" + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/mod.rs b/apps/aether-gateway/src/handlers/shared/mod.rs index 00aabf45d..97638b093 100644 --- a/apps/aether-gateway/src/handlers/shared/mod.rs +++ b/apps/aether-gateway/src/handlers/shared/mod.rs @@ -1,35 +1,47 @@ mod admin_proxy; mod api_keys; +mod auth_api_key_secret; mod catalog; mod email_templates; mod external_models; +mod identity_oauth_provider_secret; +mod multipart; mod normalize; mod payloads; +mod payment_currency; mod payment_direct; mod payment_gateway_config; +mod payment_gateway_secret; +mod payment_order_stripe_secret; +mod provider_catalog_credential; +mod provider_ops_credential; pub(crate) mod provider_pool; mod request_utils; +mod runtime_secret; mod system_config_values; mod usage_stats; pub(crate) use self::admin_proxy::{ attach_admin_audit_response, build_admin_proxy_auth_required_response, - build_unhandled_admin_proxy_response, + build_unhandled_admin_proxy_response, mark_sensitive_admin_response_no_store, }; pub(crate) use self::api_keys::{ api_key_placeholder_display, configured_api_key_prefix, generate_gateway_api_key_plaintext, - generate_gateway_secret_plaintext, masked_gateway_api_key_display, + generate_gateway_secret_plaintext, masked_gateway_api_key_display, masked_secret_display, normalize_optional_api_key_concurrent_limit, }; +pub(crate) use self::auth_api_key_secret::{ + decrypt_or_migrate_auth_api_key_secret, open_auth_api_key_secret, seal_auth_api_key_secret, +}; pub(crate) use self::catalog::{ - build_admin_provider_key_response, decrypt_catalog_secret_with_fallbacks, - default_provider_key_status_snapshot, effective_catalog_encryption_key, - encrypt_catalog_secret_with_fallbacks, masked_catalog_api_key, - masked_catalog_api_key_for_provider, parse_catalog_auth_config_json, + build_admin_provider_key_response, decrypt_catalog_secret_or_legacy_plaintext, + decrypt_catalog_secret_with_fallbacks, default_provider_key_status_snapshot, + effective_catalog_encryption_key, encrypt_catalog_secret_with_fallbacks, + masked_catalog_api_key, masked_catalog_api_key_for_provider, parse_catalog_auth_config_json, provider_catalog_key_supports_format, provider_key_health_summary, provider_key_health_summary_at, provider_key_status_snapshot_payload, sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot, - take_secret_prefix, take_secret_suffix, + take_secret_prefix, take_secret_suffix, StoredCatalogSecret, }; pub(crate) use self::email_templates::{ admin_email_template_definition, admin_email_template_html_key, @@ -37,6 +49,14 @@ pub(crate) use self::email_templates::{ read_admin_email_template_payload, render_admin_email_template_html, }; pub(crate) use self::external_models::OFFICIAL_EXTERNAL_MODEL_PROVIDERS; +pub(crate) use self::identity_oauth_provider_secret::{ + decrypt_or_migrate_identity_oauth_provider_client_secret, + identity_oauth_provider_secret_binding_matches, seal_identity_oauth_provider_client_secret, +}; +pub(crate) use self::multipart::{ + find_multipart_boundary, find_multipart_boundary_after_crlf, parse_multipart_boundary, + MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES, +}; pub(crate) use self::normalize::{ deserialize_optional_json_patch, deserialize_optional_string_list_patch, ip_rule_pattern_matches, ip_rules_allow, json_ip_rules_allow, normalize_feature_settings, @@ -47,30 +67,69 @@ pub(crate) use self::payloads::{ InternalGatewayAuthContextRequest, InternalGatewayExecuteRequest, InternalGatewayResolveRequest, InternalTunnelHeartbeatRequest, InternalTunnelNodeStatusRequest, }; +pub(crate) use self::payment_currency::{ + effective_payment_exchange_rate, normalize_payment_currency, stripe_amount_to_major, + stripe_amount_to_minor, +}; pub(crate) use self::payment_direct::{ - close_direct_gateway_order, create_alipay_direct_checkout, create_stripe_direct_checkout, - create_wxpay_direct_checkout, direct_payment_client_ip, refund_direct_gateway_order, + close_direct_gateway_checkout, close_direct_gateway_order, create_alipay_direct_checkout, + create_stripe_direct_checkout, create_wxpay_direct_checkout, find_payment_callback_order, + payment_callback_settlement_values, public_payment_http_client, refund_direct_gateway_order, verify_alipay_notify_callback, verify_wxpay_notify_callback, DirectGatewayRefundResult, - DirectPaymentCheckoutInput, + DirectPaymentCheckoutError, DirectPaymentCheckoutInput, }; pub(crate) use self::payment_gateway_config::{ + normalize_payment_callback_base_url, normalize_payment_https_url, payment_gateway_allow_user_refund, payment_gateway_channels_config_json, payment_gateway_channels_json, payment_gateway_config_json, payment_gateway_provider_for_payment_method, payment_gateway_refund_enabled, payment_gateway_secret_keys_json, }; +pub(crate) use self::payment_gateway_secret::{ + open_payment_gateway_secret, payment_gateway_secret_is_legacy_unbound, + seal_payment_gateway_secret, PaymentGatewaySecretBinding, PaymentGatewaySecretProjection, +}; +pub(crate) use self::payment_order_stripe_secret::{ + normalize_stripe_client_secret, open_payment_order_stripe_client_secret, + seal_payment_order_stripe_client_secret, PaymentOrderStripeSecretBinding, + PaymentOrderStripeSecretProjection, STRIPE_CLIENT_SECRET_ENCRYPTED_KEY, +}; +pub(crate) use self::provider_catalog_credential::{ + open_provider_catalog_credential, seal_provider_catalog_credential, + ProviderCatalogCredentialField, ProviderCatalogCredentialProjection, +}; +pub(crate) use self::provider_ops_credential::{ + canonicalize_provider_ops_base_url, open_provider_ops_credential, + provider_ops_credential_binding_from_config, provider_ops_credential_field_is_secret, + provider_ops_outbound_policy_digest, resolve_provider_ops_same_origin_url, + seal_provider_ops_credential, ProviderOpsCanonicalDestination, ProviderOpsCredentialBinding, + ProviderOpsCredentialProjection, PROVIDER_OPS_PERSISTENT_SECRET_FIELDS, + PROVIDER_OPS_TRANSIENT_METADATA_FIELDS, PROVIDER_OPS_TRANSIENT_SECRET_FIELDS, +}; pub(crate) use self::request_utils::{ admin_proxy_local_requires_buffered_body, internal_proxy_local_requires_buffered_body, json_string_list, local_proxy_route_requires_buffered_body, mark_external_models_official_providers, public_support_local_requires_buffered_body, query_param_bool, query_param_optional_bool, query_param_value, request_enables_control_execute, rust_auth_terminates_provider_credentials, - sanitize_upstream_path_and_query, should_strip_forwarded_provider_credential_header, - should_strip_forwarded_trusted_admin_header, strip_query_param, unix_ms_to_rfc3339, - unix_secs_to_rfc3339, + sanitize_upstream_path_and_query, security_log_url_origin, + should_strip_forwarded_provider_credential_header, should_strip_forwarded_trusted_admin_header, + strip_query_param, unix_ms_to_rfc3339, unix_secs_to_rfc3339, +}; +pub(crate) use self::runtime_secret::{ + open_runtime_secret_payload, open_runtime_secret_payload_with_encryption_key, + runtime_secret_payload_is_sealed, seal_runtime_secret_payload, + seal_runtime_secret_payload_with_encryption_key, }; pub(crate) use self::system_config_values::{ - module_available_from_env, system_config_bool, system_config_string, + bark_device_key_binding, canonical_bark_server_url, decrypt_or_migrate_bark_device_key, + decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_smtp_password, + decrypt_or_migrate_system_config_secret, decrypt_system_config_secret, encrypt_bark_device_key, + encrypt_ldap_bind_password, encrypt_smtp_password, encrypt_system_config_secret, + ldap_attribute_description_is_valid, ldap_bind_password_binding_matches, + ldap_distinguished_name_is_valid, ldap_module_config_is_valid, ldap_search_filter_is_valid, + module_available_from_env, normalize_ldap_transport_server_url, smtp_password_binding, + system_config_bool, system_config_string, BarkDeviceKeyBinding, SmtpPasswordBinding, }; pub(crate) use self::usage_stats::{ admin_stats_bad_request_response, parse_bounded_u32, round_to, AdminStatsTimeRange, diff --git a/apps/aether-gateway/src/handlers/shared/multipart.rs b/apps/aether-gateway/src/handlers/shared/multipart.rs new file mode 100644 index 000000000..f6d76533a --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/multipart.rs @@ -0,0 +1,350 @@ +/// RFC 2046 limits a multipart boundary to at most 70 characters. Using +/// HTTP token syntax here also keeps the value safe to embed in the byte +/// delimiter used by the lightweight parsers. +pub(crate) const MAX_MULTIPART_BOUNDARY_BYTES: usize = 70; +/// Keep multipart metadata bounded independently of the file payload size. +pub(crate) const MAX_MULTIPART_PARTS: usize = 128; +pub(crate) const MAX_MULTIPART_PART_HEADER_BYTES: usize = 64 * 1024; + +/// Find a multipart delimiter at the beginning of a buffer or after CRLF. +/// The returned index points at the delimiter itself (not the preceding CRLF). +pub(crate) fn find_multipart_boundary(haystack: &[u8], delimiter: &[u8]) -> Option { + find_multipart_boundary_inner(haystack, delimiter, true) +} + +/// Find a multipart delimiter that is preceded by CRLF. This variant is used +/// while scanning a part payload, where an apparent delimiter at byte zero is +/// payload data rather than a valid framing boundary. +pub(crate) fn find_multipart_boundary_after_crlf( + haystack: &[u8], + delimiter: &[u8], +) -> Option { + find_multipart_boundary_inner(haystack, delimiter, false) +} + +fn find_multipart_boundary_inner( + haystack: &[u8], + delimiter: &[u8], + allow_start: bool, +) -> Option { + if delimiter.is_empty() { + return None; + } + if allow_start && multipart_boundary_is_valid_at(haystack, 0, delimiter) { + return Some(0); + } + + let window_len = delimiter.len().checked_add(2)?; + haystack + .windows(window_len) + .enumerate() + .find_map(|(index, window)| { + if &window[..2] != b"\r\n" || &window[2..] != delimiter { + return None; + } + let delimiter_index = index + 2; + multipart_boundary_is_valid_at(haystack, delimiter_index, delimiter) + .then_some(delimiter_index) + }) +} + +fn multipart_boundary_is_valid_at(haystack: &[u8], index: usize, delimiter: &[u8]) -> bool { + let suffix_start = index.checked_add(delimiter.len()); + let Some(suffix_start) = suffix_start else { + return false; + }; + if !haystack + .get(index..) + .is_some_and(|remaining| remaining.starts_with(delimiter)) + { + return false; + } + let Some(suffix) = haystack.get(suffix_start..) else { + return false; + }; + if suffix.starts_with(b"\r\n") { + return true; + } + suffix + .strip_prefix(b"--") + .is_some_and(|remaining| remaining.is_empty() || remaining.starts_with(b"\r\n")) +} + +/// Extract and validate a multipart boundary parameter. +/// +/// The parameter may use the usual optional surrounding quotes, but the +/// boundary value itself must be an ASCII HTTP token. Rejecting malformed +/// values at the content-type boundary prevents parser ambiguity and bounds +/// the work performed by downstream delimiter scans. +pub(crate) fn parse_multipart_boundary(content_type: &str) -> Option { + let segments = split_multipart_header_parameters(content_type)?; + let media_type = segments.first()?.trim(); + if !media_type.eq_ignore_ascii_case("multipart/form-data") { + return None; + } + + let mut boundary = None; + let mut seen_keys = Vec::new(); + for segment in segments.into_iter().skip(1) { + let segment = segment.trim(); + if segment.is_empty() { + return None; + } + let (raw_key, raw_value) = segment.split_once('=')?; + let key = raw_key.trim(); + if key.is_empty() || !key.as_bytes().iter().copied().all(is_http_token_byte) { + return None; + } + if seen_keys + .iter() + .any(|seen: &String| seen.eq_ignore_ascii_case(key)) + { + return None; + } + seen_keys.push(key.to_ascii_lowercase()); + + let (value, had_escape) = parse_multipart_parameter_value(raw_value.trim())?; + if !key.eq_ignore_ascii_case("boundary") { + continue; + } + // Keep boundary parsing deliberately narrower than generic quoted + // parameter parsing: escaped boundary values are ambiguous across + // HTTP stacks and are rejected here. + if had_escape { + return None; + } + if !is_valid_multipart_boundary(&value) { + return None; + } + boundary = Some(value); + } + + boundary +} + +/// Split a semicolon-delimited HTTP header while honoring quoted strings. +/// Returning `None` for an unterminated quote or escape prevents a malformed +/// parameter from being reinterpreted by a downstream parser. +fn split_multipart_header_parameters(value: &str) -> Option> { + let mut segments = Vec::new(); + let mut start = 0usize; + let mut in_quotes = false; + let mut escaped = false; + + for (index, character) in value.char_indices() { + if character.is_ascii_control() { + return None; + } + if in_quotes { + if escaped { + escaped = false; + } else if character == '\\' { + escaped = true; + } else if character == '"' { + in_quotes = false; + } + } else if character == '"' { + in_quotes = true; + } else if character == ';' { + segments.push(&value[start..index]); + start = index + character.len_utf8(); + } + } + + if in_quotes || escaped { + return None; + } + segments.push(&value[start..]); + Some(segments) +} + +fn parse_multipart_parameter_value(value: &str) -> Option<(String, bool)> { + if value.is_empty() { + return None; + } + if value.starts_with('"') { + if value.len() < 2 || !value.ends_with('"') { + return None; + } + let inner = &value[1..value.len() - 1]; + let mut parsed = String::with_capacity(inner.len()); + let mut escaped = false; + let mut had_escape = false; + for character in inner.chars() { + if escaped { + if character.is_ascii_control() { + return None; + } + parsed.push(character); + escaped = false; + had_escape = true; + } else if character == '\\' { + escaped = true; + } else { + if character == '"' || character.is_ascii_control() { + return None; + } + parsed.push(character); + } + } + if escaped { + return None; + } + return Some((parsed, had_escape)); + } + + value + .as_bytes() + .iter() + .copied() + .all(is_http_token_byte) + .then(|| (value.to_string(), false)) +} + +fn is_valid_multipart_boundary(value: &str) -> bool { + !value.is_empty() + && value.len() <= MAX_MULTIPART_BOUNDARY_BYTES + && value.as_bytes().iter().copied().all(is_http_token_byte) +} + +fn is_http_token_byte(byte: u8) -> bool { + matches!( + byte, + b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + | b'!' + | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) +} + +#[cfg(test)] +mod tests { + use super::{ + find_multipart_boundary, find_multipart_boundary_after_crlf, parse_multipart_boundary, + MAX_MULTIPART_BOUNDARY_BYTES, MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES, + }; + + #[test] + fn accepts_token_boundary_and_optional_quotes() { + assert_eq!( + parse_multipart_boundary("Multipart/Form-Data; boundary=----WebKitFormBoundaryabc123") + .as_deref(), + Some("----WebKitFormBoundaryabc123") + ); + assert_eq!( + parse_multipart_boundary("multipart/form-data; boundary=\"quoted-boundary\"") + .as_deref(), + Some("quoted-boundary") + ); + } + + #[test] + fn rejects_non_token_control_quote_and_oversized_boundaries() { + for content_type in [ + "multipart/form-data; boundary=", + "multipart/form-data; boundary=bad boundary", + "multipart/form-data; boundary=\"bad;boundary\"", + "multipart/form-data; boundary=bad\r\nvalue", + "multipart/form-data; boundary=bad\"quote", + "multipart/form-data; boundary=\"unterminated", + "multipart/form-data; foo", + "multipart/form-data; foo=\"unterminated; boundary=valid", + "multipart/form-data; boundary=valid trailing", + ] { + assert!( + parse_multipart_boundary(content_type).is_none(), + "{content_type:?}" + ); + } + + let oversized = "a".repeat(MAX_MULTIPART_BOUNDARY_BYTES + 1); + assert!( + parse_multipart_boundary(&format!("multipart/form-data; boundary={oversized}")) + .is_none() + ); + } + + #[test] + fn rejects_duplicate_boundary_parameters() { + for content_type in [ + "multipart/form-data; boundary=first; boundary=second", + "multipart/form-data; boundary=first; BOUNDARY=second", + ] { + assert!( + parse_multipart_boundary(content_type).is_none(), + "duplicate boundary parameters must be rejected: {content_type}" + ); + } + } + + #[test] + fn accepts_quoted_unknown_parameters_with_semicolons() { + assert_eq!( + parse_multipart_boundary( + "multipart/form-data; note=\"semi;colon\"; boundary=quoted-token" + ) + .as_deref(), + Some("quoted-token") + ); + } + + #[test] + fn rejects_escaped_boundary_and_duplicate_unknown_parameters() { + for content_type in [ + "multipart/form-data; boundary=\"escaped\\\"token\"", + "multipart/form-data; note=one; NOTE=two; boundary=token", + "multipart/form-data; note=\"unterminated; boundary=token", + "multipart/form-data; note=\"closed\"trailing; boundary=token", + ] { + assert!( + parse_multipart_boundary(content_type).is_none(), + "malformed content type must be rejected: {content_type:?}" + ); + } + } + + #[test] + fn rejects_non_multipart_media_types() { + assert!(parse_multipart_boundary("application/json; boundary=abc").is_none()); + assert!(parse_multipart_boundary("x-multipart/form-data; boundary=abc").is_none()); + } + + #[test] + fn multipart_metadata_limits_remain_bounded() { + assert_eq!(MAX_MULTIPART_PARTS, 128); + assert_eq!(MAX_MULTIPART_PART_HEADER_BYTES, 64 * 1024); + } + + #[test] + fn boundary_scanner_ignores_embedded_markers_and_invalid_suffixes() { + let delimiter = b"--boundary"; + let payload = b"prefix\r\n--boundaryX\r\nmore\r\n--boundary\r\n"; + assert_eq!( + find_multipart_boundary_after_crlf(payload, delimiter), + Some(payload.len() - delimiter.len() - 2) + ); + assert_eq!( + find_multipart_boundary(b"payload--boundary\r\n", delimiter), + None + ); + assert_eq!(find_multipart_boundary(b"--boundaryX\r\n", delimiter), None); + assert_eq!( + find_multipart_boundary_after_crlf(b"payload\r\n--boundaryX\r\n", delimiter), + None + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/normalize.rs b/apps/aether-gateway/src/handlers/shared/normalize.rs index 597faa7ac..582aa8197 100644 --- a/apps/aether-gateway/src/handlers/shared/normalize.rs +++ b/apps/aether-gateway/src/handlers/shared/normalize.rs @@ -151,12 +151,15 @@ pub(crate) fn ip_rules_allow(rules: Option<&[String]>, remote_ip: IpAddr) -> boo for raw in rules { let rule = raw.trim(); if rule.is_empty() { - continue; + return false; } let (deny, pattern) = match rule.strip_prefix('!') { Some(pattern) => (true, pattern.trim()), None => (false, rule), }; + if pattern.is_empty() || !valid_ip_rule_pattern(pattern) { + return false; + } let matched = ip_rule_pattern_matches(pattern, remote_ip); if deny && matched { return false; @@ -181,11 +184,14 @@ pub(crate) fn json_ip_rules_allow(value: Option<&Value>, remote_ip: IpAddr) -> b return true; }; if value.is_null() { - return true; + return false; } let Some(items) = value.as_array() else { return false; }; + if items.is_empty() { + return false; + } let mut rules = Vec::with_capacity(items.len()); for item in items { let Some(rule) = item.as_str() else { @@ -505,10 +511,24 @@ mod tests { #[test] fn json_ip_rules_allow_rejects_invalid_stored_shape() { + assert!(!json_ip_rules_allow( + Some(&serde_json::Value::Null), + v4(10, 0, 0, 1) + )); + assert!(!json_ip_rules_allow(Some(&json!([])), v4(10, 0, 0, 1))); assert!(!json_ip_rules_allow( Some(&json!({"bad": true})), v4(10, 0, 0, 1) )); assert!(!json_ip_rules_allow(Some(&json!([123])), v4(10, 0, 0, 1))); + assert!(!json_ip_rules_allow( + Some(&json!(["!not-an-ip"])), + v4(10, 0, 0, 1) + )); + assert!(!json_ip_rules_allow( + Some(&json!(["203.0.113.0/999"])), + v4(203, 0, 113, 1) + )); + assert!(!json_ip_rules_allow(Some(&json!([""])), v4(10, 0, 0, 1))); } } diff --git a/apps/aether-gateway/src/handlers/shared/payloads.rs b/apps/aether-gateway/src/handlers/shared/payloads.rs index ff9723da3..4d707e51d 100644 --- a/apps/aether-gateway/src/handlers/shared/payloads.rs +++ b/apps/aether-gateway/src/handlers/shared/payloads.rs @@ -4,6 +4,7 @@ use std::collections::BTreeMap; #[derive(Debug, Deserialize)] pub(crate) struct InternalTunnelHeartbeatRequest { pub(crate) node_id: String, + pub(crate) heartbeat_session_id: String, pub(crate) heartbeat_id: u64, #[serde(default)] pub(crate) heartbeat_interval: Option, diff --git a/apps/aether-gateway/src/handlers/shared/payment_currency.rs b/apps/aether-gateway/src/handlers/shared/payment_currency.rs new file mode 100644 index 000000000..7ea28fcba --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/payment_currency.rs @@ -0,0 +1,101 @@ +/// Normalize a payment currency at an external payment boundary. +/// +/// Payment-order storage uses a three-character ISO-style code. Keep the +/// boundary deliberately narrow so values cannot be truncated differently by +/// the individual database adapters or payment providers. +pub(crate) fn normalize_payment_currency(value: &str, field: &str) -> Result { + let trimmed = value.trim(); + if trimmed.len() != 3 || !trimmed.bytes().all(|byte| byte.is_ascii_alphabetic()) { + return Err(format!("{field} must be a 3-letter currency code")); + } + Ok(trimmed.to_ascii_uppercase()) +} + +/// Return the exchange rate that should be persisted and used for settlement. +/// USD is already the canonical accounting currency, so a configured +/// conversion rate (often the CNY default) must never be applied to USD +/// checkouts. +pub(crate) fn effective_payment_exchange_rate( + pay_currency: &str, + configured_rate: f64, +) -> Result { + let currency = normalize_payment_currency(pay_currency, "pay_currency")?; + if !configured_rate.is_finite() || configured_rate <= 0.0 { + return Err("usd_exchange_rate must be finite and positive".to_string()); + } + Ok(if currency == "USD" { + 1.0 + } else { + configured_rate + }) +} + +pub(crate) fn stripe_minor_unit_multiplier(currency: &str) -> f64 { + match currency.trim().to_ascii_lowercase().as_str() { + "bif" | "clp" | "djf" | "gnf" | "jpy" | "kmf" | "krw" | "mga" | "pyg" | "rwf" | "ugx" + | "vnd" | "vuv" | "xaf" | "xof" | "xpf" => 1.0, + "bhd" | "jod" | "kwd" | "omr" | "tnd" => 1_000.0, + _ => 100.0, + } +} + +pub(crate) fn stripe_amount_to_minor(amount_major: f64, currency: &str) -> Option { + let amount_minor = (amount_major * stripe_minor_unit_multiplier(currency)).round(); + if !amount_minor.is_finite() || amount_minor <= 0.0 || amount_minor >= i64::MAX as f64 { + return None; + } + Some(amount_minor as i64) +} + +pub(crate) fn stripe_amount_to_major(amount_minor: i64, currency: &str) -> f64 { + amount_minor as f64 / stripe_minor_unit_multiplier(currency) +} + +#[cfg(test)] +mod tests { + use super::{ + effective_payment_exchange_rate, normalize_payment_currency, stripe_amount_to_major, + stripe_amount_to_minor, + }; + + #[test] + fn payment_currency_is_trimmed_uppercased_and_bounded_to_ascii_three_letters() { + assert_eq!( + normalize_payment_currency(" cny ", "pay_currency"), + Ok("CNY".to_string()) + ); + for value in ["CN", "CNYY", "C1Y", "人民币", ""] { + assert!( + normalize_payment_currency(value, "pay_currency").is_err(), + "invalid currency should be rejected: {value}" + ); + } + } + + #[test] + fn usd_uses_unit_effective_exchange_rate() { + assert_eq!(effective_payment_exchange_rate(" usd ", 7.2), Ok(1.0)); + assert_eq!(effective_payment_exchange_rate("CNY", 7.2), Ok(7.2)); + assert!(effective_payment_exchange_rate("USD", f64::NAN).is_err()); + assert!(effective_payment_exchange_rate("CN", 7.2).is_err()); + } + + #[test] + fn stripe_amounts_handle_zero_two_and_three_decimal_currencies() { + assert_eq!(stripe_amount_to_minor(1234.0, "JPY"), Some(1234)); + assert_eq!(stripe_amount_to_major(1234, "jpy"), 1234.0); + + assert_eq!(stripe_amount_to_minor(12.34, "USD"), Some(1234)); + assert_eq!(stripe_amount_to_major(1234, "usd"), 12.34); + + assert_eq!(stripe_amount_to_minor(1.234, "KWD"), Some(1234)); + assert_eq!(stripe_amount_to_major(1234, "kwd"), 1.234); + } + + #[test] + fn stripe_minor_amount_rejects_invalid_or_overflowing_values() { + assert_eq!(stripe_amount_to_minor(f64::NAN, "usd"), None); + assert_eq!(stripe_amount_to_minor(0.0, "usd"), None); + assert_eq!(stripe_amount_to_minor(f64::MAX, "kwd"), None); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/payment_direct.rs b/apps/aether-gateway/src/handlers/shared/payment_direct.rs index cc0769574..408fb2f2f 100644 --- a/apps/aether-gateway/src/handlers/shared/payment_direct.rs +++ b/apps/aether-gateway/src/handlers/shared/payment_direct.rs @@ -3,26 +3,84 @@ use aes_gcm::{ aead::{Aead, Payload}, Aes256Gcm, KeyInit, Nonce, }; +use aether_crypto::{rsa_pkcs1_sha256_sign, rsa_pkcs1_sha256_verify, RsaPkcs1Sha256Error}; use axum::http; use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine as _}; use chrono::Utc; -use rsa::pkcs1::{DecodeRsaPrivateKey, DecodeRsaPublicKey}; -use rsa::pkcs1v15::{ - Signature as RsaPkcs1v15Signature, SigningKey as RsaPkcs1v15SigningKey, - VerifyingKey as RsaPkcs1v15VerifyingKey, -}; -use rsa::pkcs8::{DecodePrivateKey, DecodePublicKey}; -use rsa::signature::{SignatureEncoding, Signer, Verifier}; -use rsa::{RsaPrivateKey, RsaPublicKey}; use serde_json::{json, Value}; use sha2::{Digest, Sha256}; use std::collections::BTreeMap; +use std::net::{IpAddr, SocketAddr}; const ALIPAY_DEFAULT_GATEWAY_URL: &str = "https://openapi.alipay.com/gateway.do"; const WXPAY_DEFAULT_BASE_URL: &str = "https://api.mch.weixin.qq.com"; +const MAX_PAYMENT_GATEWAY_RESPONSE_BYTES: usize = 1024 * 1024; +const MAX_PAYMENT_ORDER_NO_BYTES: usize = 64; +const MAX_PAYMENT_GATEWAY_ID_BYTES: usize = 128; +const MAX_PAYMENT_CALLBACK_KEY_BYTES: usize = 128; +const MAX_PAYMENT_SIGNATURE_BYTES: usize = 16 * 1024; +const MAX_PAYMENT_CALLBACK_CIPHERTEXT_BYTES: usize = 1024 * 1024; const WXPAY_NOTIFY_SUCCESS: &str = "TRANSACTION.SUCCESS"; const WXPAY_TRADE_SUCCESS: &str = "SUCCESS"; const WXPAY_CURRENCY: &str = "CNY"; +const WXPAY_SIGNATURE_TOLERANCE_SECONDS: i64 = 5 * 60; +const STRIPE_CHECKOUT_UNCERTAIN_DETAIL: &str = "Stripe 支付服务暂时不可用"; +const STRIPE_CHECKOUT_FAILED_DETAIL: &str = "Stripe 支付请求被拒绝"; + +/// Errors returned by direct checkout creators are typed so callers can +/// distinguish a provider request whose outcome was not observed from a +/// deterministic validation/business failure. A cancelled Stripe intent must +/// never be replayed as a live checkout under the same merchant order number. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum DirectPaymentCheckoutError { + Canceled, + Uncertain(String), + Failed(String), +} + +impl From for DirectPaymentCheckoutError { + fn from(value: String) -> Self { + Self::Failed(value) + } +} + +impl DirectPaymentCheckoutError { + pub(crate) fn into_detail(self) -> String { + match self { + Self::Canceled => "支付订单已取消".to_string(), + Self::Uncertain(detail) | Self::Failed(detail) => detail, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum DirectPaymentRequestError { + Uncertain(String), + Failed(String), +} + +impl From for DirectPaymentRequestError { + fn from(value: String) -> Self { + Self::Failed(value) + } +} + +impl DirectPaymentRequestError { + fn into_detail(self) -> String { + match self { + Self::Uncertain(detail) | Self::Failed(detail) => detail, + } + } +} + +impl From for DirectPaymentCheckoutError { + fn from(error: DirectPaymentRequestError) -> Self { + match error { + DirectPaymentRequestError::Uncertain(detail) => Self::Uncertain(detail), + DirectPaymentRequestError::Failed(detail) => Self::Failed(detail), + } + } +} #[derive(Debug, Clone)] pub(crate) struct DirectPaymentCheckoutInput { @@ -42,35 +100,190 @@ pub(crate) struct DirectPaymentCheckoutInput { pub(crate) struct DirectGatewayRefundResult { pub(crate) gateway_refund_id: String, pub(crate) status: String, - pub(crate) payload: Value, + pub(crate) proof: Value, } -#[derive(Debug, Clone)] +impl DirectGatewayRefundResult { + pub(crate) fn is_succeeded(&self) -> bool { + self.status.eq_ignore_ascii_case("success") + } + + pub(crate) fn is_pending(&self) -> bool { + matches!( + self.status.to_ascii_lowercase().as_str(), + "pending" | "processing" + ) + } +} + +#[derive(Clone)] struct DirectGatewayConfig { record: aether_data_contracts::repository::billing::PaymentGatewayConfigRecord, config: serde_json::Map, secrets: serde_json::Map, } +impl std::fmt::Debug for DirectGatewayConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("DirectGatewayConfig") + .field("record", &self.record) + .field("config", &"[REDACTED]") + .field("secrets", &"[REDACTED]") + .finish() + } +} + +pub(crate) async fn find_payment_callback_order( + state: &AppState, + order_no: &str, +) -> Result, String> { + state + .find_payment_order_by_order_no(order_no) + .await + .map_err(|_| "支付订单查询失败".to_string()) +} + +fn valid_exchange_rate(value: Option) -> Option { + value.filter(|value| value.is_finite() && *value > 0.0) +} + +/// Resolve callback accounting values without allowing mutable gateway +/// configuration to rewrite an order that was already created. The provider +/// amount is still passed separately and is checked against the stored order +/// by the repository; this helper only supplies a finite amount for the +/// callback contract and a non-sensitive rate fallback. +pub(crate) fn payment_callback_settlement_values( + order: Option<&aether_data::repository::wallet::StoredAdminPaymentOrder>, + pay_amount: f64, + fallback_exchange_rate: Option, +) -> Result<(f64, Option), String> { + let exchange_rate = order + .and_then(|order| valid_exchange_rate(order.exchange_rate)) + .or_else(|| valid_exchange_rate(fallback_exchange_rate)); + let amount_usd = order + .map(|order| order.amount_usd) + .filter(|value| value.is_finite() && *value > 0.0) + .or_else(|| { + exchange_rate + .map(|rate| pay_amount / rate) + .filter(|value| value.is_finite() && *value > 0.0) + }) + .unwrap_or(pay_amount); + if !amount_usd.is_finite() || amount_usd <= 0.0 { + return Err("支付回调金额换算无效".to_string()); + } + Ok((amount_usd, exchange_rate)) +} + fn payment_payload_hash(payload: &Value) -> Result { let encoded = serde_json::to_vec(payload) - .map_err(|err| format!("payment callback payload encode failed: {err}"))?; + .map_err(|_| "payment callback payload encode failed".to_string())?; let digest = Sha256::digest(&encoded); Ok(digest.iter().map(|byte| format!("{byte:02x}")).collect()) } -fn stripe_minor_unit_amount(pay_amount: f64, pay_currency: &str) -> Result { - let currency = pay_currency.trim().to_ascii_lowercase(); - let multiplier = match currency.as_str() { - "bif" | "clp" | "djf" | "gnf" | "jpy" | "kmf" | "krw" | "mga" | "pyg" | "rwf" | "ugx" - | "vnd" | "vuv" | "xaf" | "xof" | "xpf" => 1.0, - _ => 100.0, - }; - let amount = (pay_amount * multiplier).round(); - if !amount.is_finite() || amount <= 0.0 { - return Err("Stripe 支付金额无效".to_string()); +fn payment_bytes_hash(payload: &[u8]) -> String { + let digest = Sha256::digest(payload); + digest.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn validated_payment_identifier( + value: &str, + field: &str, + max_bytes: usize, +) -> Result { + let value = value.trim(); + if value.is_empty() + || value.len() > max_bytes + || !value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) + { + return Err(format!("{field} 格式无效")); + } + Ok(value.to_string()) +} + +fn validated_optional_payment_identifier( + value: Option<&str>, + field: &str, + max_bytes: usize, +) -> Result, String> { + value + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(|value| validated_payment_identifier(value, field, max_bytes)) + .transpose() +} + +fn payment_callback_key( + gateway: &str, + candidate: Option<&str>, + payload_hash: &str, +) -> Result { + let candidate = validated_optional_payment_identifier( + candidate, + "支付通知事件 ID", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )?; + let prefix = format!("{gateway}:"); + if let Some(candidate) = candidate { + if prefix.len() + candidate.len() <= MAX_PAYMENT_CALLBACK_KEY_BYTES { + return Ok(format!("{prefix}{candidate}")); + } + } + Ok(format!("{prefix}{payload_hash}")) +} + +fn payment_callback_projection( + gateway: &str, + order_no: &str, + gateway_order_id: Option<&str>, + amount: f64, + currency: &str, + status: &str, +) -> Value { + json!({ + "gateway": gateway, + "order_no": order_no, + "gateway_order_id": gateway_order_id, + "amount": amount, + "currency": currency, + "status": status, + "signature_valid": true, + "processed_at": Utc::now().to_rfc3339(), + }) +} + +fn gateway_refund_proof( + gateway: &str, + gateway_refund_id: &str, + status: &str, + order_no: &str, + refund_no: &str, + amount: f64, + currency: &str, +) -> Value { + json!({ + "gateway": gateway, + "id": gateway_refund_id, + "status": status, + "order_no": order_no, + "refund_no": refund_no, + "amount": amount, + "currency": currency, + "processed_at": Utc::now().to_rfc3339(), + }) +} + +fn wxpay_refund_status(value: Option<&str>) -> &'static str { + match value.map(str::trim).map(str::to_ascii_uppercase).as_deref() { + Some("SUCCESS") => "success", + Some("PROCESSING") => "processing", + Some("CLOSED" | "ABNORMAL") | None => "failed", + Some(_) => "failed", } - Ok(amount as i64) } fn config_string(config: &serde_json::Map, key: &str) -> Option { @@ -91,6 +304,13 @@ fn secret_string(secrets: &serde_json::Map, key: &str) -> Option< .map(ToOwned::to_owned) } +fn stripe_checkout_response_is_canceled(value: &Value) -> bool { + value + .get("status") + .and_then(Value::as_str) + .is_some_and(|status| status.trim().eq_ignore_ascii_case("canceled")) +} + fn record_config_map( record: &aether_data_contracts::repository::billing::PaymentGatewayConfigRecord, ) -> serde_json::Map { @@ -104,22 +324,35 @@ async fn load_direct_gateway_config( state: &AppState, provider: &str, ) -> Result { - let record = state + let mut record = state .find_payment_gateway_config(provider) .await - .map_err(|err| format!("{provider} 配置读取失败: {err:?}"))? + .map_err(|_| format!("{provider} 配置读取失败"))? .ok_or_else(|| format!("{provider} 未配置"))?; if !record.enabled { return Err(format!("{provider} 未启用")); } + record.pay_currency = super::normalize_payment_currency(&record.pay_currency, "pay_currency")?; + record.usd_exchange_rate = + super::effective_payment_exchange_rate(&record.pay_currency, record.usd_exchange_rate) + .map_err(|_| format!("{provider} 汇率配置无效"))?; + if !record.endpoint_url.trim().is_empty() { + record.endpoint_url = + super::normalize_payment_https_url(&record.endpoint_url, "endpoint_url")?; + } + record.callback_base_url = record + .callback_base_url + .as_deref() + .map(super::normalize_payment_callback_base_url) + .transpose()?; let Some(encrypted) = record.merchant_key_encrypted.as_deref() else { return Err(format!("{provider} 密钥未配置")); }; - let Some(plaintext) = - super::decrypt_catalog_secret_with_fallbacks(state.encryption_key(), encrypted) - else { - return Err(format!("{provider} 密钥解密失败")); - }; + let binding = super::PaymentGatewaySecretBinding::from_record(&record) + .map_err(|_| format!("{provider} 密钥绑定无效"))?; + let plaintext = super::open_payment_gateway_secret(state, &binding, encrypted) + .map_err(|_| format!("{provider} 密钥解密失败"))? + .plaintext; let secrets = serde_json::from_str::(&plaintext) .ok() .and_then(|value| value.as_object().cloned()) @@ -132,52 +365,48 @@ async fn load_direct_gateway_config( }) } -fn pem_candidates(raw: &str, labels: &[&str]) -> Vec { - let trimmed = raw.trim(); - if trimmed.starts_with("-----BEGIN") { - return vec![trimmed.to_string()]; - } - labels - .iter() - .map(|label| format!("-----BEGIN {label}-----\n{trimmed}\n-----END {label}-----")) - .collect() -} - -fn decode_rsa_private_key(raw: &str) -> Result { - let mut errors = Vec::new(); - for candidate in pem_candidates(raw, &["PRIVATE KEY", "RSA PRIVATE KEY"]) { - match RsaPrivateKey::from_pkcs8_pem(&candidate) { - Ok(key) => return Ok(key), - Err(err) => errors.push(format!("pkcs8: {err}")), - } - match RsaPrivateKey::from_pkcs1_pem(&candidate) { - Ok(key) => return Ok(key), - Err(err) => errors.push(format!("pkcs1: {err}")), - } - } - Err(format!("RSA 私钥解析失败: {}", errors.join("; "))) -} - -fn decode_rsa_public_key(raw: &str) -> Result { - let mut errors = Vec::new(); - for candidate in pem_candidates(raw, &["PUBLIC KEY", "RSA PUBLIC KEY"]) { - match RsaPublicKey::from_public_key_pem(&candidate) { - Ok(key) => return Ok(key), - Err(err) => errors.push(format!("spki: {err}")), - } - match RsaPublicKey::from_pkcs1_pem(&candidate) { - Ok(key) => return Ok(key), - Err(err) => errors.push(format!("pkcs1: {err}")), - } - } - Err(format!("RSA 公钥解析失败: {}", errors.join("; "))) +/// Load only the credentials and provider identifiers required to authenticate +/// a callback. Gateway enablement, checkout endpoint, and current pricing +/// settings are mutable operational controls; changing them must not strand a +/// payment that was already accepted by the provider. +async fn load_direct_gateway_callback_config( + state: &AppState, + provider: &str, +) -> Result { + let record = state + .find_payment_gateway_config(provider) + .await + .map_err(|_| format!("{provider} 配置读取失败"))? + .ok_or_else(|| format!("{provider} 未配置"))?; + let Some(encrypted) = record.merchant_key_encrypted.as_deref() else { + return Err(format!("{provider} 密钥未配置")); + }; + let binding = super::PaymentGatewaySecretBinding::from_record(&record) + .map_err(|_| format!("{provider} 密钥绑定无效"))?; + let plaintext = super::open_payment_gateway_secret(state, &binding, encrypted) + .map_err(|_| format!("{provider} 密钥解密失败"))? + .plaintext; + let secrets = serde_json::from_str::(&plaintext) + .ok() + .and_then(|value| value.as_object().cloned()) + .ok_or_else(|| format!("{provider} 密钥格式无效"))?; + let config = record_config_map(&record); + Ok(DirectGatewayConfig { + record, + config, + secrets, + }) } fn rsa_sha256_sign_base64(private_key: &str, message: &str) -> Result { - let private_key = decode_rsa_private_key(private_key)?; - let signing_key = RsaPkcs1v15SigningKey::::new(private_key); - let signature = signing_key.sign(message.as_bytes()); - Ok(BASE64_STANDARD.encode(signature.to_bytes())) + let signature = + rsa_pkcs1_sha256_sign(private_key.as_bytes(), message.as_bytes()).map_err(|error| { + match error { + RsaPkcs1Sha256Error::InvalidPrivateKey => "RSA 私钥解析失败".to_string(), + _ => "RSA 签名失败".to_string(), + } + })?; + Ok(BASE64_STANDARD.encode(signature)) } fn rsa_sha256_verify_base64( @@ -185,14 +414,26 @@ fn rsa_sha256_verify_base64( message: &str, signature_base64: &str, ) -> Result { - let public_key = decode_rsa_public_key(public_key)?; - let signature_bytes = BASE64_STANDARD - .decode(signature_base64.trim()) - .map_err(|err| format!("签名 base64 解码失败: {err}"))?; - let signature = RsaPkcs1v15Signature::try_from(signature_bytes.as_slice()) - .map_err(|err| format!("签名格式无效: {err}"))?; - let verifying_key = RsaPkcs1v15VerifyingKey::::new(public_key); - Ok(verifying_key.verify(message.as_bytes(), &signature).is_ok()) + let signature_bytes = decode_payment_base64_with_limit( + signature_base64.trim(), + MAX_PAYMENT_SIGNATURE_BYTES, + "签名 base64 解码失败", + )?; + rsa_pkcs1_sha256_verify(public_key.as_bytes(), message.as_bytes(), &signature_bytes).map_err( + |error| match error { + RsaPkcs1Sha256Error::InvalidPublicKey => "RSA 公钥解析失败".to_string(), + _ => "签名格式无效".to_string(), + }, + ) +} + +fn decode_payment_base64_with_limit( + value: &str, + limit_bytes: usize, + invalid_message: &'static str, +) -> Result, String> { + crate::execution_runtime::transport::decode_base64_body_with_limit(value, limit_bytes) + .map_err(|_| invalid_message.to_string()) } fn alipay_timestamp() -> String { @@ -262,7 +503,7 @@ fn alipay_signed_params( params.insert( "biz_content".to_string(), serde_json::to_string(&biz_content) - .map_err(|err| format!("支付宝 biz_content 编码失败: {err}"))?, + .map_err(|_| "支付宝 biz_content 编码失败".to_string())?, ); if let Some(notify_url) = notify_url.map(str::trim).filter(|value| !value.is_empty()) { params.insert("notify_url".to_string(), notify_url.to_string()); @@ -277,7 +518,7 @@ fn alipay_signed_params( } fn url_with_query(base: &str, params: &BTreeMap) -> Result { - let mut url = url::Url::parse(base).map_err(|err| format!("支付网关地址无效: {err}"))?; + let mut url = url::Url::parse(base).map_err(|_| "支付网关地址无效".to_string())?; { let mut query = url.query_pairs_mut(); for (key, value) in params { @@ -287,29 +528,78 @@ fn url_with_query(base: &str, params: &BTreeMap) -> Result, +) -> Result { + let gateway_url = url::Url::parse(&alipay_gateway_url(config)) + .map_err(|_| DirectPaymentRequestError::Failed("支付宝网关 URL 无效".to_string()))?; + let client = public_payment_http_client(&gateway_url) + .await + .map_err(DirectPaymentRequestError::Failed)?; + let response = client + .post(gateway_url) + .form(params) + .send() + .await + .map_err(|_| DirectPaymentRequestError::Uncertain("支付宝请求失败".to_string()))?; + let status = response.status(); + let body = + aether_http::read_response_bytes_with_limit(response, MAX_PAYMENT_GATEWAY_RESPONSE_BYTES) + .await + .map_err(|_| DirectPaymentRequestError::Uncertain("支付宝响应读取失败".to_string()))?; + let value = serde_json::from_slice::(&body) + .map_err(|_| DirectPaymentRequestError::Uncertain("支付宝响应格式无效".to_string()))?; + if !status.is_success() { + let detail = format!("支付宝 HTTP 状态异常: {status}"); + return Err( + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + DirectPaymentRequestError::Uncertain(detail) + } else { + DirectPaymentRequestError::Failed(detail) + }, + ); + } + Ok(value) +} + async fn alipay_post( state: &AppState, config: &DirectGatewayConfig, params: &BTreeMap, ) -> Result { - let response = state - .client - .post(alipay_gateway_url(config)) - .form(params) - .send() + alipay_post_typed(state, config, params) .await - .map_err(|err| format!("支付宝请求失败: {err}"))?; - let status = response.status(); - let body = response - .text() - .await - .map_err(|err| format!("支付宝响应读取失败: {err}"))?; - let value = - serde_json::from_str::(&body).map_err(|_| "支付宝响应格式无效".to_string())?; - if !status.is_success() { - return Err(format!("支付宝 HTTP 状态异常: {status}")); + .map_err(DirectPaymentRequestError::into_detail) +} + +/// A parsed 4xx Alipay response is an explicit request/business rejection and +/// can safely use the page-pay fallback. System/transient sub-codes are +/// excluded even when Alipay wraps them in the generic 40004 business code: +/// those responses do not establish that the precreate request was rejected +/// before reaching the provider. +fn alipay_precreate_business_refusal(response: &Value) -> bool { + let code = response.get("code").and_then(Value::as_str).unwrap_or(""); + if !code.starts_with('4') { + return false; } - Ok(value) + let sub_code = response + .get("sub_code") + .and_then(Value::as_str) + .unwrap_or("") + .trim() + .to_ascii_uppercase(); + ![ + "SYSTEM_ERROR", + "TIMEOUT", + "REQUEST_TIMEOUT", + "NETWORK_ERROR", + "PROCESSING", + "UNKNOWN_ERROR", + ] + .iter() + .any(|marker| sub_code.contains(marker)) } fn alipay_response_success<'a>(value: &'a Value, key: &str) -> Result<&'a Value, String> { @@ -319,21 +609,20 @@ fn alipay_response_success<'a>(value: &'a Value, key: &str) -> Result<&'a Value, if response.get("code").and_then(Value::as_str) == Some("10000") { return Ok(response); } - let message = response - .get("sub_msg") - .or_else(|| response.get("msg")) - .and_then(Value::as_str) - .unwrap_or("支付宝业务请求失败"); - Err(message.to_string()) + Err("支付宝业务请求失败".to_string()) } pub(crate) async fn create_alipay_direct_checkout( state: &AppState, input: &DirectPaymentCheckoutInput, -) -> Result { +) -> Result { + let order_no = + validated_payment_identifier(&input.order_no, "支付订单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "alipay").await?; if !input.pay_currency.eq_ignore_ascii_case("CNY") { - return Err("支付宝官方直连当前仅支持 CNY".to_string()); + return Err(DirectPaymentCheckoutError::Failed( + "支付宝官方直连当前仅支持 CNY".to_string(), + )); } let mode = config_string(&config.config, "payment_mode") .unwrap_or_else(|| "precreate".to_string()) @@ -346,7 +635,7 @@ pub(crate) async fn create_alipay_direct_checkout( &config, "alipay.trade.page.pay", json!({ - "out_trade_no": input.order_no, + "out_trade_no": order_no, "total_amount": total_amount, "subject": input.subject, "product_code": "FAST_INSTANT_TRADE_PAY", @@ -357,7 +646,7 @@ pub(crate) async fn create_alipay_direct_checkout( return Ok(json!({ "gateway": "alipay", "display_name": input.display_name, - "gateway_order_id": input.order_no, + "gateway_order_id": order_no, "payment_url": url_with_query(&gateway_url, ¶ms)?, "submit_method": "GET", "qr_code": Value::Null, @@ -374,7 +663,7 @@ pub(crate) async fn create_alipay_direct_checkout( &config, "alipay.trade.wap.pay", json!({ - "out_trade_no": input.order_no, + "out_trade_no": order_no, "total_amount": total_amount, "subject": input.subject, "product_code": "QUICK_WAP_WAY", @@ -385,7 +674,7 @@ pub(crate) async fn create_alipay_direct_checkout( return Ok(json!({ "gateway": "alipay", "display_name": input.display_name, - "gateway_order_id": input.order_no, + "gateway_order_id": order_no, "payment_url": url_with_query(&gateway_url, ¶ms)?, "submit_method": "GET", "qr_code": Value::Null, @@ -402,7 +691,7 @@ pub(crate) async fn create_alipay_direct_checkout( &config, "alipay.trade.precreate", json!({ - "out_trade_no": input.order_no, + "out_trade_no": order_no, "total_amount": total_amount, "subject": input.subject, "product_code": "FACE_TO_FACE_PAYMENT", @@ -410,40 +699,27 @@ pub(crate) async fn create_alipay_direct_checkout( Some(&input.notify_url), None, )?; - match alipay_post(state, &config, ¶ms) + let value = alipay_post_typed(state, &config, ¶ms) .await - .and_then(|value| { - let response = alipay_response_success(&value, "alipay_trade_precreate_response")?; - response - .get("qr_code") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|qr_code| { - json!({ - "gateway": "alipay", - "display_name": input.display_name, - "gateway_order_id": input.order_no, - "payment_url": Value::Null, - "submit_method": "qrcode", - "qr_code": qr_code, - "pay_amount": input.pay_amount, - "pay_currency": input.pay_currency, - "payment_channel": input.payment_channel, - "callback_url": input.notify_url, - "return_url": return_url, - "expires_at": input.expires_at.to_rfc3339(), - }) - }) - .ok_or_else(|| "支付宝预下单响应缺少 qr_code".to_string()) - }) { - Ok(value) => Ok(value), - Err(precreate_err) => { + .map_err(DirectPaymentCheckoutError::from)?; + let Some(response) = value.get("alipay_trade_precreate_response") else { + return Err(DirectPaymentCheckoutError::Uncertain( + "支付宝预下单响应缺少业务结果".to_string(), + )); + }; + let response_code = response.get("code").and_then(Value::as_str); + if response_code != Some("10000") { + if !alipay_precreate_business_refusal(response) { + return Err(DirectPaymentCheckoutError::Uncertain( + "支付宝预下单结果不确定".to_string(), + )); + } + { let params = alipay_signed_params( &config, "alipay.trade.page.pay", json!({ - "out_trade_no": input.order_no, + "out_trade_no": order_no, "total_amount": total_amount, "subject": input.subject, "product_code": "FAST_INSTANT_TRADE_PAY", @@ -454,7 +730,7 @@ pub(crate) async fn create_alipay_direct_checkout( Ok(json!({ "gateway": "alipay", "display_name": input.display_name, - "gateway_order_id": input.order_no, + "gateway_order_id": order_no, "payment_url": url_with_query(&gateway_url, ¶ms)?, "submit_method": "GET", "qr_code": Value::Null, @@ -464,9 +740,34 @@ pub(crate) async fn create_alipay_direct_checkout( "callback_url": input.notify_url, "return_url": return_url, "expires_at": input.expires_at.to_rfc3339(), - "integration_status": format!("precreate_fallback: {precreate_err}"), + "integration_status": "precreate_fallback", })) } + } else { + let Some(qr_code) = response + .get("qr_code") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Err(DirectPaymentCheckoutError::Uncertain( + "支付宝预下单响应缺少 qr_code".to_string(), + )); + }; + Ok(json!({ + "gateway": "alipay", + "display_name": input.display_name, + "gateway_order_id": order_no, + "payment_url": Value::Null, + "submit_method": "qrcode", + "qr_code": qr_code, + "pay_amount": input.pay_amount, + "pay_currency": input.pay_currency, + "payment_channel": input.payment_channel, + "callback_url": input.notify_url, + "return_url": return_url, + "expires_at": input.expires_at.to_rfc3339(), + })) } } @@ -474,7 +775,7 @@ pub(crate) async fn verify_alipay_notify_callback( state: &AppState, body: &[u8], ) -> Result { - let config = load_direct_gateway_config(state, "alipay").await?; + let config = load_direct_gateway_callback_config(state, "alipay").await?; let raw = std::str::from_utf8(body).map_err(|_| "支付宝通知请求体不是 UTF-8".to_string())?; let params = url::form_urlencoded::parse(raw.as_bytes()) .map(|(key, value)| (key.into_owned(), value.into_owned())) @@ -487,10 +788,13 @@ pub(crate) async fn verify_alipay_notify_callback( if !rsa_sha256_verify_base64(&alipay_public_key(&config)?, &sign_content, signature)? { return Err("支付宝通知签名无效".to_string()); } - if let Some(app_id) = params.get("app_id").map(String::as_str) { - if app_id != alipay_app_id(&config)? { - return Err("支付宝通知 app_id 不匹配".to_string()); - } + let app_id = params + .get("app_id") + .map(String::as_str) + .filter(|value| !value.trim().is_empty()) + .ok_or_else(|| "支付宝通知缺少 app_id".to_string())?; + if app_id != alipay_app_id(&config)? { + return Err("支付宝通知 app_id 不匹配".to_string()); } if !matches!( params.get("trade_status").map(String::as_str), @@ -500,23 +804,51 @@ pub(crate) async fn verify_alipay_notify_callback( } let order_no = params .get("out_trade_no") - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) + .map(|value| { + validated_payment_identifier( + value, + "支付宝通知 out_trade_no", + MAX_PAYMENT_ORDER_NO_BYTES, + ) + }) + .transpose()? .ok_or_else(|| "支付宝通知缺少 out_trade_no".to_string())?; + let gateway_order_id = validated_optional_payment_identifier( + params.get("trade_no").map(String::as_str), + "支付宝通知 trade_no", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )? + .ok_or_else(|| "支付宝通知缺少 trade_no".to_string())?; let pay_amount = params .get("total_amount") .or_else(|| params.get("receipt_amount")) .and_then(|value| value.parse::().ok()) - .filter(|value| *value > 0.0) + .filter(|value| value.is_finite() && *value > 0.0) .ok_or_else(|| "支付宝通知金额无效".to_string())?; - let payload = serde_json::to_value(¶ms).unwrap_or_else(|_| json!({})); - let payload_hash = payment_payload_hash(&payload)?; - let callback_key = params - .get("notify_id") - .or_else(|| params.get("trade_no")) - .cloned() - .unwrap_or_else(|| format!("alipay:{order_no}:{payload_hash}")); - let exchange_rate = config.record.usd_exchange_rate; + let raw_payload = serde_json::to_value(¶ms).unwrap_or_else(|_| json!({})); + let payload_hash = payment_payload_hash(&raw_payload)?; + let callback_key = payment_callback_key( + "alipay", + params + .get("notify_id") + .map(String::as_str) + .or(Some(gateway_order_id.as_str())), + &payload_hash, + )?; + let order = find_payment_callback_order(state, &order_no).await?; + let (amount_usd, exchange_rate) = payment_callback_settlement_values( + order.as_ref(), + pay_amount, + Some(config.record.usd_exchange_rate), + )?; + let payload = payment_callback_projection( + "alipay", + &order_no, + Some(&gateway_order_id), + pay_amount, + "CNY", + "success", + ); Ok( aether_data::repository::wallet::ProcessPaymentCallbackInput { payment_method: "alipay".to_string(), @@ -524,15 +856,11 @@ pub(crate) async fn verify_alipay_notify_callback( payment_channel: Some("alipay".to_string()), callback_key, order_no: Some(order_no), - gateway_order_id: params.get("trade_no").cloned(), - amount_usd: if exchange_rate > 0.0 { - pay_amount / exchange_rate - } else { - pay_amount - }, + gateway_order_id: Some(gateway_order_id), + amount_usd, pay_amount: Some(pay_amount), - pay_currency: Some(config.record.pay_currency), - exchange_rate: Some(exchange_rate), + pay_currency: Some("CNY".to_string()), + exchange_rate, payload_hash, payload, signature_valid: true, @@ -549,6 +877,70 @@ fn wxpay_base_url(config: &DirectGatewayConfig) -> String { } } +pub(crate) async fn public_payment_http_client(url: &url::Url) -> Result { + if url.scheme() != "https" + || !url.username().is_empty() + || url.password().is_some() + || url.host_str().is_none() + || url.fragment().is_some() + { + return Err("支付网关必须是无凭据的 HTTPS URL".to_string()); + } + let host = url + .host_str() + .ok_or_else(|| "支付网关 URL 缺少主机名".to_string())?; + let port = url + .port_or_known_default() + .ok_or_else(|| "支付网关 URL 缺少端口".to_string())?; + let addrs = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .map_err(|_| "支付网关 DNS 解析失败".to_string())? + }; + validate_public_payment_resolved_addrs(url, &addrs)?; + + let mut builder = reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()); + if host.parse::().is_err() { + builder = builder.resolve_to_addrs(host, &addrs); + } + builder + .build() + .map_err(|_| "支付网关 HTTP 客户端初始化失败".to_string()) +} + +fn validate_public_payment_resolved_addrs( + url: &url::Url, + addrs: &[SocketAddr], +) -> Result<(), String> { + if addrs.is_empty() + || addrs.iter().any(|addr| { + aether_http::is_private_or_reserved_ip(addr.ip()) + && !(is_fixed_stripe_api_origin(url) + && aether_http::is_ipv4_benchmarking_fake_ip(addr.ip())) + }) + { + return Err("支付网关解析到私有或保留地址".to_string()); + } + Ok(()) +} + +fn is_fixed_stripe_api_origin(url: &url::Url) -> bool { + url.scheme() == "https" + && url.host_str().is_some_and(|host| { + host.trim_end_matches('.') + .eq_ignore_ascii_case("api.stripe.com") + }) + && url.port_or_known_default() == Some(443) + && url.username().is_empty() + && url.password().is_none() + && url.query().is_none() + && url.fragment().is_none() +} + fn wxpay_config_string(config: &DirectGatewayConfig, key: &str) -> Result { config_string(&config.config, key).ok_or_else(|| format!("微信支付 {key} 未配置")) } @@ -588,13 +980,18 @@ async fn wxpay_post_json( config: &DirectGatewayConfig, canonical_url: &str, body: Value, -) -> Result { - let body = - serde_json::to_string(&body).map_err(|err| format!("微信支付请求体编码失败: {err}"))?; - let auth = wxpay_authorization(config, "POST", canonical_url, &body)?; +) -> Result { + let body = serde_json::to_string(&body) + .map_err(|_| DirectPaymentRequestError::Failed("微信支付请求体编码失败".to_string()))?; + let auth = wxpay_authorization(config, "POST", canonical_url, &body) + .map_err(DirectPaymentRequestError::Failed)?; let url = format!("{}{}", wxpay_base_url(config), canonical_url); - let response = state - .client + let url = url::Url::parse(&url) + .map_err(|_| DirectPaymentRequestError::Failed("微信支付网关 URL 无效".to_string()))?; + let client = public_payment_http_client(&url) + .await + .map_err(DirectPaymentRequestError::Failed)?; + let response = client .post(url) .header(http::header::AUTHORIZATION, auth) .header(http::header::ACCEPT, "application/json") @@ -602,34 +999,34 @@ async fn wxpay_post_json( .body(body) .send() .await - .map_err(|err| format!("微信支付请求失败: {err}"))?; + .map_err(|_| DirectPaymentRequestError::Uncertain("微信支付请求失败".to_string()))?; let status = response.status(); - let text = response - .text() - .await - .map_err(|err| format!("微信支付响应读取失败: {err}"))?; + let body = + aether_http::read_response_bytes_with_limit(response, MAX_PAYMENT_GATEWAY_RESPONSE_BYTES) + .await + .map_err(|_| { + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + DirectPaymentRequestError::Uncertain("微信支付响应读取失败".to_string()) + } else { + DirectPaymentRequestError::Failed("微信支付响应读取失败".to_string()) + } + })?; + let text = String::from_utf8_lossy(&body); if !status.is_success() { - let detail = serde_json::from_str::(&text) - .ok() - .and_then(|value| { - value - .get("message") - .and_then(Value::as_str) - .map(ToOwned::to_owned) - .or_else(|| { - value - .get("code") - .and_then(Value::as_str) - .map(ToOwned::to_owned) - }) - }) - .unwrap_or_else(|| format!("微信支付 HTTP 状态异常: {status}")); - return Err(detail); + let detail = "微信支付业务请求失败".to_string(); + return Err( + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + DirectPaymentRequestError::Uncertain(detail) + } else { + DirectPaymentRequestError::Failed(detail) + }, + ); } if text.trim().is_empty() { return Ok(json!({})); } - serde_json::from_str::(&text).map_err(|_| "微信支付响应格式无效".to_string()) + serde_json::from_str::(&text) + .map_err(|_| DirectPaymentRequestError::Uncertain("微信支付响应格式无效".to_string())) } async fn wxpay_post_empty_success( @@ -637,39 +1034,21 @@ async fn wxpay_post_empty_success( config: &DirectGatewayConfig, canonical_url: &str, body: Value, -) -> Result { +) -> Result { wxpay_post_json(state, config, canonical_url, body).await } -fn wxpay_client_ip(headers: &http::HeaderMap) -> Option { - crate::headers::header_value_str(headers, "x-forwarded-for") - .and_then(|value| { - value - .split(',') - .map(str::trim) - .find(|segment| !segment.is_empty() && !segment.eq_ignore_ascii_case("unknown")) - .map(ToOwned::to_owned) - }) - .or_else(|| { - crate::headers::header_value_str(headers, "x-real-ip").and_then(|value| { - let value = value.trim(); - (!value.is_empty() && !value.eq_ignore_ascii_case("unknown")) - .then(|| value.to_string()) - }) - }) -} - -pub(crate) fn direct_payment_client_ip(headers: &http::HeaderMap) -> Option { - wxpay_client_ip(headers) -} - pub(crate) async fn create_wxpay_direct_checkout( state: &AppState, input: &DirectPaymentCheckoutInput, -) -> Result { +) -> Result { + let order_no = + validated_payment_identifier(&input.order_no, "支付订单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "wxpay").await?; if !input.pay_currency.eq_ignore_ascii_case(WXPAY_CURRENCY) { - return Err("微信支付官方直连当前仅支持 CNY".to_string()); + return Err(DirectPaymentCheckoutError::Failed( + "微信支付官方直连当前仅支持 CNY".to_string(), + )); } let total_fen = wxpay_money_to_fen(input.pay_amount)?; let app_id = wxpay_config_string(&config, "app_id")?; @@ -678,7 +1057,7 @@ pub(crate) async fn create_wxpay_direct_checkout( "appid": app_id, "mchid": mch_id, "description": input.subject, - "out_trade_no": input.order_no, + "out_trade_no": order_no, "notify_url": input.notify_url, "amount": { "total": total_fen, @@ -694,11 +1073,15 @@ pub(crate) async fn create_wxpay_direct_checkout( .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .ok_or_else(|| "微信 Native 响应缺少 code_url".to_string())?; + .ok_or_else(|| { + DirectPaymentCheckoutError::Uncertain( + "微信 Native 响应缺少 code_url".to_string(), + ) + })?; Ok(json!({ "gateway": "wxpay", "display_name": input.display_name, - "gateway_order_id": input.order_no, + "gateway_order_id": order_no, "payment_url": Value::Null, "submit_method": "qrcode", "qr_code": code_url, @@ -729,7 +1112,9 @@ pub(crate) async fn create_wxpay_direct_checkout( .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .ok_or_else(|| "微信 H5 响应缺少 h5_url".to_string())? + .ok_or_else(|| { + DirectPaymentCheckoutError::Uncertain("微信 H5 响应缺少 h5_url".to_string()) + })? .to_string(); if let Some(return_url) = input.return_url.as_deref() { let sep = if h5_url.contains('?') { "&" } else { "?" }; @@ -741,7 +1126,7 @@ pub(crate) async fn create_wxpay_direct_checkout( Ok(json!({ "gateway": "wxpay", "display_name": input.display_name, - "gateway_order_id": input.order_no, + "gateway_order_id": order_no, "payment_url": h5_url, "h5_url": h5_url, "submit_method": "GET", @@ -754,29 +1139,36 @@ pub(crate) async fn create_wxpay_direct_checkout( "expires_at": input.expires_at.to_rfc3339(), })) } - "jsapi" => Err("微信 JSAPI 需要前端提供 OpenID,当前充值入口尚未接入".to_string()), - _ => Err("微信支付通道不可用".to_string()), + "jsapi" => Err(DirectPaymentCheckoutError::Failed( + "微信 JSAPI 需要前端提供 OpenID,当前充值入口尚未接入".to_string(), + )), + _ => Err(DirectPaymentCheckoutError::Failed( + "微信支付通道不可用".to_string(), + )), } } pub(crate) async fn create_stripe_direct_checkout( state: &AppState, input: &DirectPaymentCheckoutInput, -) -> Result { +) -> Result { + let order_no = + validated_payment_identifier(&input.order_no, "支付订单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "stripe").await?; let Some(secret_key) = secret_string(&config.secrets, "secret_key") else { - return Err("Stripe secret_key 未配置".to_string()); + return Err("Stripe secret_key 未配置".to_string().into()); }; let Some(publishable_key) = config_string(&config.config, "publishable_key") else { - return Err("Stripe publishable_key 未配置".to_string()); + return Err("Stripe publishable_key 未配置".to_string().into()); }; - let amount = stripe_minor_unit_amount(input.pay_amount, &config.record.pay_currency)?; + let amount = super::stripe_amount_to_minor(input.pay_amount, &config.record.pay_currency) + .ok_or_else(|| "Stripe 支付金额无效".to_string())?; let currency = config.record.pay_currency.trim().to_ascii_lowercase(); let mut form = vec![ ("amount".to_string(), amount.to_string()), ("currency".to_string(), currency.clone()), ("description".to_string(), input.subject.clone()), - ("metadata[order_no]".to_string(), input.order_no.clone()), + ("metadata[order_no]".to_string(), order_no), ( "metadata[payment_provider]".to_string(), "stripe".to_string(), @@ -796,34 +1188,96 @@ pub(crate) async fn create_stripe_direct_checkout( "web".to_string(), )); } - let response = state - .client - .post("https://api.stripe.com/v1/payment_intents") + let stripe_endpoint = + url::Url::parse("https://api.stripe.com/v1/payment_intents").map_err(|_| { + DirectPaymentCheckoutError::Uncertain(STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string()) + })?; + let stripe_client = public_payment_http_client(&stripe_endpoint) + .await + .map_err(|_| { + DirectPaymentCheckoutError::Uncertain(STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string()) + })?; + let response = match stripe_client + .post(stripe_endpoint) + .header( + "Idempotency-Key", + format!("aether-payment-intent-{}", input.order_no), + ) .basic_auth(secret_key, Some("")) .form(&form) .send() .await - .map_err(|err| format!("Stripe PaymentIntent 创建失败: {err}"))?; + { + Ok(response) => response, + Err(_) => { + // Do not surface reqwest's error string: it can contain provider + // response details, proxy information, or other deployment data. + tracing::warn!( + event_name = "stripe_payment_intent_request_failed", + "Stripe PaymentIntent request failed" + ); + return Err(DirectPaymentCheckoutError::Uncertain( + STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string(), + )); + } + }; let status = response.status(); - let body = response - .text() - .await - .map_err(|err| format!("Stripe 响应读取失败: {err}"))?; - let value = - serde_json::from_str::(&body).map_err(|_| "Stripe 响应格式无效".to_string())?; + let body = + aether_http::read_response_bytes_with_limit(response, MAX_PAYMENT_GATEWAY_RESPONSE_BYTES) + .await + .map_err(|_| { + tracing::warn!( + event_name = "stripe_payment_intent_response_read_failed", + upstream_status = %status, + "Stripe PaymentIntent response could not be read" + ); + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + DirectPaymentCheckoutError::Uncertain( + STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string(), + ) + } else { + DirectPaymentCheckoutError::Failed(STRIPE_CHECKOUT_FAILED_DETAIL.to_string()) + } + })?; if !status.is_success() { - let message = value - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Stripe PaymentIntent 创建失败"); - return Err(message.to_string()); + // Stripe error bodies may contain account identifiers or provider + // internals. Classify by status only and keep the body private. + tracing::warn!( + event_name = "stripe_payment_intent_upstream_rejected", + upstream_status = %status, + "Stripe PaymentIntent request was rejected" + ); + return Err( + if status.is_server_error() || status == http::StatusCode::TOO_MANY_REQUESTS { + DirectPaymentCheckoutError::Uncertain(STRIPE_CHECKOUT_UNCERTAIN_DETAIL.to_string()) + } else { + DirectPaymentCheckoutError::Failed(STRIPE_CHECKOUT_FAILED_DETAIL.to_string()) + }, + ); + } + let value = serde_json::from_slice::(&body) + .map_err(|_| DirectPaymentCheckoutError::Uncertain("Stripe 响应格式无效".to_string()))?; + // Stripe retains the response for an idempotency key even after the + // PaymentIntent is cancelled. Reusing that key would otherwise hand the + // caller a client secret that can never be confirmed. + if stripe_checkout_response_is_canceled(&value) { + return Err(DirectPaymentCheckoutError::Canceled); } let Some(intent_id) = value.get("id").and_then(Value::as_str) else { - return Err("Stripe 响应缺少 PaymentIntent ID".to_string()); + return Err(DirectPaymentCheckoutError::Uncertain( + "Stripe 响应缺少 PaymentIntent ID".to_string(), + )); }; + let intent_id = validated_payment_identifier( + intent_id, + "Stripe PaymentIntent ID", + MAX_PAYMENT_GATEWAY_ID_BYTES, + ) + .map_err(DirectPaymentCheckoutError::Uncertain)?; let Some(client_secret) = value.get("client_secret").and_then(Value::as_str) else { - return Err("Stripe 响应缺少 client_secret".to_string()); + return Err(DirectPaymentCheckoutError::Uncertain( + "Stripe 响应缺少 client_secret".to_string(), + )); }; Ok(json!({ "gateway": "stripe", @@ -855,6 +1309,12 @@ fn wxpay_verify_notify_headers( ) -> Result<(), String> { let signature = wxpay_header(headers, "wechatpay-signature")?; let timestamp = wxpay_header(headers, "wechatpay-timestamp")?; + let timestamp_unix = timestamp + .parse::() + .map_err(|_| "微信支付通知时间戳无效".to_string())?; + if Utc::now().timestamp().abs_diff(timestamp_unix) > WXPAY_SIGNATURE_TOLERANCE_SECONDS as u64 { + return Err("微信支付通知已过期".to_string()); + } let nonce = wxpay_header(headers, "wechatpay-nonce")?; let serial = wxpay_header(headers, "wechatpay-serial")?; if let Some(expected) = config_string(&config.config, "public_key_id") { @@ -888,6 +1348,9 @@ fn wxpay_decrypt_resource(config: &DirectGatewayConfig, resource: &Value) -> Res .get("nonce") .and_then(Value::as_str) .ok_or_else(|| "微信支付通知 resource.nonce 缺失".to_string())?; + if nonce.as_bytes().len() != 12 { + return Err("微信支付通知 resource.nonce 长度无效".to_string()); + } let associated_data = resource .get("associated_data") .and_then(Value::as_str) @@ -896,9 +1359,11 @@ fn wxpay_decrypt_resource(config: &DirectGatewayConfig, resource: &Value) -> Res .get("ciphertext") .and_then(Value::as_str) .ok_or_else(|| "微信支付通知 resource.ciphertext 缺失".to_string())?; - let ciphertext = BASE64_STANDARD - .decode(ciphertext) - .map_err(|err| format!("微信支付通知密文解码失败: {err}"))?; + let ciphertext = decode_payment_base64_with_limit( + ciphertext, + MAX_PAYMENT_CALLBACK_CIPHERTEXT_BYTES, + "微信支付通知密文解码失败", + )?; let cipher = Aes256Gcm::new_from_slice(api_v3_key.as_bytes()) .map_err(|_| "微信支付 api_v3_key 无效".to_string())?; let plaintext = cipher @@ -913,12 +1378,27 @@ fn wxpay_decrypt_resource(config: &DirectGatewayConfig, resource: &Value) -> Res serde_json::from_slice::(&plaintext).map_err(|_| "微信支付通知明文格式无效".to_string()) } +fn wxpay_notify_payment_channel(tx: &Value) -> Result { + let trade_type = tx + .get("trade_type") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| "微信支付通知缺少 trade_type".to_string())?; + match trade_type.to_ascii_uppercase().as_str() { + "NATIVE" => Ok("native".to_string()), + "MWEB" => Ok("h5".to_string()), + "JSAPI" => Ok("jsapi".to_string()), + _ => Err("微信支付通知 trade_type 不受支持".to_string()), + } +} + pub(crate) async fn verify_wxpay_notify_callback( state: &AppState, headers: &http::HeaderMap, body: &[u8], ) -> Result { - let config = load_direct_gateway_config(state, "wxpay").await?; + let config = load_direct_gateway_callback_config(state, "wxpay").await?; wxpay_verify_notify_headers(&config, headers, body)?; let payload = serde_json::from_slice::(body).map_err(|_| "微信支付通知请求体无效".to_string())?; @@ -934,56 +1414,85 @@ pub(crate) async fn verify_wxpay_notify_callback( if tx.get("trade_state").and_then(Value::as_str) != Some(WXPAY_TRADE_SUCCESS) { return Err("微信支付交易不是成功状态".to_string()); } + let expected_app_id = wxpay_config_string(&config, "app_id")?; + let expected_mch_id = wxpay_config_string(&config, "mch_id")?; + if tx.get("appid").and_then(Value::as_str) != Some(expected_app_id.as_str()) { + return Err("微信支付通知 appid 不匹配".to_string()); + } + if tx.get("mchid").and_then(Value::as_str) != Some(expected_mch_id.as_str()) { + return Err("微信支付通知 mchid 不匹配".to_string()); + } + let payment_channel = wxpay_notify_payment_channel(&tx)?; let order_no = tx .get("out_trade_no") .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .map(|value| { + validated_payment_identifier( + value, + "微信支付通知 out_trade_no", + MAX_PAYMENT_ORDER_NO_BYTES, + ) + }) + .transpose()? .ok_or_else(|| "微信支付通知缺少 out_trade_no".to_string())?; - let transaction_id = tx - .get("transaction_id") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned); + let transaction_id = validated_optional_payment_identifier( + tx.get("transaction_id").and_then(Value::as_str), + "微信支付通知 transaction_id", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )? + .ok_or_else(|| "微信支付通知缺少 transaction_id".to_string())?; let amount_fen = tx .get("amount") .and_then(|value| value.get("total")) .and_then(Value::as_i64) .filter(|value| *value > 0) .ok_or_else(|| "微信支付通知金额无效".to_string())?; - let pay_amount = amount_fen as f64 / 100.0; - let exchange_rate = config.record.usd_exchange_rate; - let callback_key = payload - .get("id") + let currency = tx + .get("amount") + .and_then(|value| value.get("currency")) .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .or_else(|| transaction_id.clone()) - .unwrap_or_else(|| format!("wxpay:{order_no}")); - let callback_payload = json!({ - "notification": payload, - "transaction": tx, - }); - let payload_hash = payment_payload_hash(&callback_payload)?; + .ok_or_else(|| "微信支付通知缺少币种".to_string())?; + if !currency.eq_ignore_ascii_case(WXPAY_CURRENCY) { + return Err("微信支付通知币种不匹配".to_string()); + } + let pay_amount = amount_fen as f64 / 100.0; + let order = find_payment_callback_order(state, &order_no).await?; + let (amount_usd, exchange_rate) = payment_callback_settlement_values( + order.as_ref(), + pay_amount, + Some(config.record.usd_exchange_rate), + )?; + let payload_hash = payment_bytes_hash(body); + let callback_key = payment_callback_key( + "wxpay", + payload + .get("id") + .and_then(Value::as_str) + .or(Some(transaction_id.as_str())), + &payload_hash, + )?; + let callback_payload = payment_callback_projection( + "wxpay", + &order_no, + Some(&transaction_id), + pay_amount, + WXPAY_CURRENCY, + "success", + ); Ok( aether_data::repository::wallet::ProcessPaymentCallbackInput { payment_method: "wxpay".to_string(), payment_provider: Some("wxpay".to_string()), - payment_channel: None, + payment_channel: Some(payment_channel), callback_key, order_no: Some(order_no), - gateway_order_id: transaction_id, - amount_usd: if exchange_rate > 0.0 { - pay_amount / exchange_rate - } else { - pay_amount - }, + gateway_order_id: Some(transaction_id), + amount_usd, pay_amount: Some(pay_amount), - pay_currency: Some(config.record.pay_currency), - exchange_rate: Some(exchange_rate), + pay_currency: Some(WXPAY_CURRENCY.to_string()), + exchange_rate, payload_hash, payload: callback_payload, signature_valid: true, @@ -991,51 +1500,122 @@ pub(crate) async fn verify_wxpay_notify_callback( ) } -pub(crate) async fn close_direct_gateway_order( +pub(crate) async fn close_direct_gateway_checkout( state: &AppState, - order: &crate::AdminWalletPaymentOrderRecord, + payment_provider: &str, + order_no: &str, + gateway_order_id: Option<&str>, ) -> Result, String> { - match order.payment_method.as_str() { + match payment_provider.trim().to_ascii_lowercase().as_str() { "alipay" => { + let order_no = + validated_payment_identifier(order_no, "支付订单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "alipay").await?; let params = alipay_signed_params( &config, "alipay.trade.close", - json!({ "out_trade_no": order.order_no }), + json!({ "out_trade_no": order_no }), None, None, )?; let value = alipay_post(state, &config, ¶ms).await?; let response = alipay_response_success(&value, "alipay_trade_close_response")?; + let gateway_order_id = validated_optional_payment_identifier( + response.get("trade_no").and_then(Value::as_str), + "支付宝 trade_no", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )?; Ok(Some(json!({ "gateway": "alipay", "closed": true, - "gateway_order_id": response.get("trade_no").and_then(Value::as_str), - "payload": value, + "gateway_order_id": gateway_order_id, }))) } "wxpay" => { + let order_no = + validated_payment_identifier(order_no, "支付订单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "wxpay").await?; - let canonical_url = - format!("/v3/pay/transactions/out-trade-no/{}/close", order.order_no); - let value = wxpay_post_empty_success( + let canonical_url = format!("/v3/pay/transactions/out-trade-no/{order_no}/close"); + wxpay_post_empty_success( state, &config, &canonical_url, json!({ "mchid": wxpay_config_string(&config, "mch_id")? }), ) - .await?; + .await + .map_err(DirectPaymentRequestError::into_detail)?; + let gateway_order_id = validated_optional_payment_identifier( + gateway_order_id, + "微信支付 transaction_id", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )?; Ok(Some(json!({ "gateway": "wxpay", "closed": true, - "gateway_order_id": order.gateway_order_id, - "payload": value, + "gateway_order_id": gateway_order_id, + }))) + } + "stripe" => { + let intent_id = validated_optional_payment_identifier( + gateway_order_id, + "Stripe PaymentIntent ID", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )? + .ok_or_else(|| "Stripe PaymentIntent ID 缺失,无法取消支付".to_string())?; + let config = load_direct_gateway_config(state, "stripe").await?; + let secret_key = secret_string(&config.secrets, "secret_key") + .ok_or_else(|| "Stripe secret_key 未配置".to_string())?; + let stripe_endpoint = url::Url::parse(&format!( + "https://api.stripe.com/v1/payment_intents/{intent_id}/cancel" + )) + .map_err(|_| "Stripe PaymentIntent 取消 URL 无效".to_string())?; + let stripe_client = public_payment_http_client(&stripe_endpoint).await?; + let response = stripe_client + .post(stripe_endpoint) + .basic_auth(secret_key, Some("")) + .send() + .await + .map_err(|_| "Stripe PaymentIntent 取消失败".to_string())?; + let status = response.status(); + let body = aether_http::read_response_bytes_with_limit( + response, + MAX_PAYMENT_GATEWAY_RESPONSE_BYTES, + ) + .await + .map_err(|_| "Stripe 取消响应读取失败".to_string())?; + let value = serde_json::from_slice::(&body) + .map_err(|_| "Stripe 取消响应格式无效".to_string())?; + if !status.is_success() { + return Err("Stripe PaymentIntent 取消失败".to_string()); + } + if value.get("id").and_then(Value::as_str) != Some(intent_id.as_str()) + || value.get("status").and_then(Value::as_str) != Some("canceled") + { + return Err("Stripe PaymentIntent 取消结果无效".to_string()); + } + Ok(Some(json!({ + "gateway": "stripe", + "closed": true, + "gateway_order_id": intent_id, }))) } _ => Ok(None), } } +pub(crate) async fn close_direct_gateway_order( + state: &AppState, + order: &crate::AdminWalletPaymentOrderRecord, +) -> Result, String> { + close_direct_gateway_checkout( + state, + &order.payment_method, + &order.order_no, + order.gateway_order_id.as_deref(), + ) + .await +} + fn refund_pay_amount( order: &crate::AdminWalletPaymentOrderRecord, amount_usd: f64, @@ -1047,7 +1627,11 @@ fn refund_pay_amount( let exchange_rate = order.exchange_rate.unwrap_or(1.0); order.amount_usd * exchange_rate }); - if order.amount_usd <= 0.0 || total_pay_amount <= 0.0 { + if !order.amount_usd.is_finite() + || order.amount_usd <= 0.0 + || !total_pay_amount.is_finite() + || total_pay_amount <= 0.0 + { return Err("原支付订单金额无效".to_string()); } Ok((amount_usd * total_pay_amount / order.amount_usd * 100.0).round() / 100.0) @@ -1062,13 +1646,23 @@ pub(crate) async fn refund_direct_gateway_order( ) -> Result, String> { match order.payment_method.as_str() { "alipay" => { + let order_no = validated_payment_identifier( + &order.order_no, + "支付订单号", + MAX_PAYMENT_ORDER_NO_BYTES, + )?; + let refund_no = + validated_payment_identifier(refund_no, "退款单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "alipay").await?; + if !config.record.pay_currency.eq_ignore_ascii_case("CNY") { + return Err("支付宝退款币种配置无效".to_string()); + } let refund_amount = refund_pay_amount(order, amount_usd)?; let params = alipay_signed_params( &config, "alipay.trade.refund", json!({ - "out_trade_no": order.order_no, + "out_trade_no": order_no, "refund_amount": format!("{refund_amount:.2}"), "refund_reason": reason.unwrap_or("wallet refund"), "out_request_no": refund_no, @@ -1078,19 +1672,46 @@ pub(crate) async fn refund_direct_gateway_order( )?; let value = alipay_post(state, &config, ¶ms).await?; let response = alipay_response_success(&value, "alipay_trade_refund_response")?; - let gateway_refund_id = response - .get("trade_no") - .and_then(Value::as_str) - .unwrap_or(&order.order_no) - .to_string(); + // Alipay's `trade_no` identifies the original payment transaction, + // not this refund. `out_request_no` is the merchant-scoped, + // idempotent refund identifier and is what we persist. + let gateway_refund_id = validated_payment_identifier( + &refund_no, + "支付宝退款 ID", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )?; + let status = "success".to_string(); + let proof = gateway_refund_proof( + "alipay", + &gateway_refund_id, + &status, + &order_no, + &refund_no, + refund_amount, + "CNY", + ); Ok(Some(DirectGatewayRefundResult { gateway_refund_id, - status: "success".to_string(), - payload: value, + status, + proof, })) } "wxpay" => { + let order_no = validated_payment_identifier( + &order.order_no, + "支付订单号", + MAX_PAYMENT_ORDER_NO_BYTES, + )?; + let refund_no = + validated_payment_identifier(refund_no, "退款单号", MAX_PAYMENT_ORDER_NO_BYTES)?; let config = load_direct_gateway_config(state, "wxpay").await?; + if !config + .record + .pay_currency + .eq_ignore_ascii_case(WXPAY_CURRENCY) + { + return Err("微信支付退款币种配置无效".to_string()); + } let total_pay_amount = order.pay_amount.ok_or_else(|| { "微信支付退款需要原订单 pay_amount,请确认订单已通过官方直连创建".to_string() })?; @@ -1100,7 +1721,7 @@ pub(crate) async fn refund_direct_gateway_order( &config, "/v3/refund/domestic/refunds", json!({ - "out_trade_no": order.order_no, + "out_trade_no": order_no, "out_refund_no": refund_no, "reason": reason.unwrap_or("wallet refund"), "amount": { @@ -1110,23 +1731,405 @@ pub(crate) async fn refund_direct_gateway_order( }, }), ) - .await?; + .await + .map_err(DirectPaymentRequestError::into_detail)?; let gateway_refund_id = value .get("refund_id") .and_then(Value::as_str) - .unwrap_or(refund_no) - .to_string(); - let status = value - .get("status") - .and_then(Value::as_str) - .unwrap_or("PROCESSING") - .to_ascii_lowercase(); + .unwrap_or(&refund_no); + let gateway_refund_id = validated_payment_identifier( + gateway_refund_id, + "微信支付退款 ID", + MAX_PAYMENT_GATEWAY_ID_BYTES, + )?; + let status = + wxpay_refund_status(value.get("status").and_then(Value::as_str)).to_string(); + let proof = gateway_refund_proof( + "wxpay", + &gateway_refund_id, + &status, + &order_no, + &refund_no, + refund_amount, + WXPAY_CURRENCY, + ); Ok(Some(DirectGatewayRefundResult { gateway_refund_id, status, - payload: value, + proof, })) } _ => Ok(None), } } + +#[cfg(test)] +mod tests { + use super::{ + alipay_precreate_business_refusal, decode_payment_base64_with_limit, gateway_refund_proof, + payment_callback_key, payment_callback_projection, payment_payload_hash, + public_payment_http_client, rsa_sha256_sign_base64, rsa_sha256_verify_base64, + validate_public_payment_resolved_addrs, validated_payment_identifier, + wxpay_notify_payment_channel, wxpay_refund_status, DirectGatewayConfig, + DirectGatewayRefundResult, MAX_PAYMENT_GATEWAY_ID_BYTES, + }; + use aws_lc_rs::encoding::{AsDer, Pkcs8V1Der, PublicKeyX509Der}; + use aws_lc_rs::rsa::{KeyPair as AwsRsaKeyPair, KeySize}; + use aws_lc_rs::signature::KeyPair as _; + use base64::engine::general_purpose::STANDARD; + use base64::Engine as _; + use serde_json::json; + use std::net::SocketAddr; + + #[test] + fn direct_gateway_config_debug_output_redacts_decrypted_secrets() { + let config = DirectGatewayConfig { + record: aether_data_contracts::repository::billing::PaymentGatewayConfigRecord { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: "https://example.test".to_string(), + callback_base_url: None, + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some("gateway-ciphertext-canary".to_string()), + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: json!({}), + created_at_unix_secs: 1, + updated_at_unix_secs: 1, + }, + config: serde_json::Map::from_iter([( + "private_config".to_string(), + json!("gateway-config-canary"), + )]), + secrets: serde_json::Map::from_iter([( + "private_key".to_string(), + json!("gateway-plaintext-secret-canary"), + )]), + }; + + let debug = format!("{config:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "gateway-ciphertext-canary", + "gateway-config-canary", + "gateway-plaintext-secret-canary", + ] { + assert!( + !debug.contains(secret), + "debug output leaked {secret}: {debug}" + ); + } + } + + #[test] + fn payment_rsa_sha256_keeps_pem_and_base64_compatibility() { + let key_pair = AwsRsaKeyPair::generate(KeySize::Rsa2048) + .expect("2048-bit test RSA private key should generate"); + let pkcs8 = AsDer::>::as_der(&key_pair) + .expect("test RSA private key should encode as PKCS#8"); + let spki = AsDer::>::as_der(key_pair.public_key()) + .expect("test RSA public key should encode as SPKI"); + let private_key = format!( + "-----BEGIN PRIVATE KEY-----\n{}\n-----END PRIVATE KEY-----", + STANDARD.encode(pkcs8.as_ref()) + ); + let public_key = STANDARD.encode(spki.as_ref()); + let message = "Aether payment signature compatibility"; + let signature = + rsa_sha256_sign_base64(&private_key, message).expect("PKCS#8 PEM key should sign"); + + assert!(rsa_sha256_verify_base64(&public_key, message, &signature) + .expect("bare-base64 SPKI key should verify")); + assert!( + !rsa_sha256_verify_base64(&public_key, "tampered", &signature) + .expect("valid but mismatched signature should return false") + ); + assert_eq!( + rsa_sha256_verify_base64(&public_key, message, "not-base64"), + Err("签名 base64 解码失败".to_string()) + ); + } + + #[test] + fn payment_base64_decode_enforces_its_allocation_limit() { + assert_eq!( + decode_payment_base64_with_limit("YWI=", 2, "invalid").expect("two decoded bytes"), + b"ab" + ); + assert_eq!( + decode_payment_base64_with_limit("AAAA", 2, "invalid"), + Err("invalid".to_string()) + ); + assert_eq!( + decode_payment_base64_with_limit("AAAAA", 2, "invalid"), + Err("invalid".to_string()) + ); + } + + #[test] + fn wxpay_notify_trade_type_maps_to_the_stored_checkout_channel() { + for (trade_type, expected) in [ + ("NATIVE", "native"), + ("MWEB", "h5"), + ("JSAPI", "jsapi"), + (" native ", "native"), + ] { + assert_eq!( + wxpay_notify_payment_channel(&json!({ "trade_type": trade_type })) + .expect("supported trade type should map"), + expected + ); + } + } + + #[test] + fn wxpay_notify_trade_type_is_required_and_rejects_uncreated_channels() { + assert!(wxpay_notify_payment_channel(&json!({})).is_err()); + assert!(wxpay_notify_payment_channel(&json!({ "trade_type": "APP" })).is_err()); + } + + #[test] + fn direct_refund_requires_an_explicit_success_terminal_state() { + for status in ["PROCESSING", "pending"] { + let result = DirectGatewayRefundResult { + gateway_refund_id: "refund-1".to_string(), + status: status.to_string(), + proof: json!({}), + }; + assert!(result.is_pending()); + assert!(!result.is_succeeded()); + } + + for status in ["SUCCESS", "success"] { + let result = DirectGatewayRefundResult { + gateway_refund_id: "refund-1".to_string(), + status: status.to_string(), + proof: json!({}), + }; + assert!(result.is_succeeded()); + assert!(!result.is_pending()); + } + + for status in ["CLOSED", "ABNORMAL", "unknown"] { + let result = DirectGatewayRefundResult { + gateway_refund_id: "refund-1".to_string(), + status: status.to_string(), + proof: json!({}), + }; + assert!(!result.is_succeeded()); + assert!(!result.is_pending()); + } + } + + #[test] + fn payment_identifiers_reject_secret_bearing_or_oversized_values() { + for value in [ + "Authorization: Bearer top-secret", + "https://internal.example/refund?token=top-secret", + "payer/openid", + "refund id with spaces", + ] { + assert!(validated_payment_identifier(value, "test", 128).is_err()); + } + assert!(validated_payment_identifier( + &"a".repeat(MAX_PAYMENT_GATEWAY_ID_BYTES + 1), + "test", + MAX_PAYMENT_GATEWAY_ID_BYTES, + ) + .is_err()); + assert_eq!( + validated_payment_identifier(" rf_123-ABC ", "test", 128) + .expect("safe identifier should pass"), + "rf_123-ABC" + ); + } + + #[test] + fn callback_projection_contains_only_the_persistence_allowlist() { + let raw = json!({ + "authorization": "Bearer payment-secret", + "url": "https://internal.example/callback?token=secret", + "payer": {"openid": "openid-secret"}, + "credential": "gateway-credential" + }); + let raw_hash = payment_payload_hash(&raw).expect("raw callback should be hashable"); + assert_eq!(raw_hash.len(), 64); + + let projection = payment_callback_projection( + "wxpay", + "order-1", + Some("transaction-1"), + 12.34, + "CNY", + "success", + ); + let object = projection + .as_object() + .expect("callback projection should be an object"); + assert_eq!(object.len(), 8); + for key in [ + "gateway", + "order_no", + "gateway_order_id", + "amount", + "currency", + "status", + "signature_valid", + "processed_at", + ] { + assert!(object.contains_key(key), "missing safe field: {key}"); + } + let encoded = projection.to_string(); + for sensitive in [ + "payment-secret", + "?token=secret", + "openid-secret", + "gateway-credential", + "payer", + "openid", + "credential", + "Bearer", + ] { + assert!(!encoded.contains(sensitive)); + } + } + + #[test] + fn gateway_refund_proof_contains_only_fixed_fields() { + let proof = gateway_refund_proof( + "alipay", + "gateway-refund-1", + "success", + "order-1", + "refund-1", + 8.5, + "CNY", + ); + let object = proof + .as_object() + .expect("gateway refund proof should be an object"); + assert_eq!(object.len(), 8); + assert!(object + .get("processed_at") + .and_then(|v| v.as_str()) + .is_some()); + for forbidden in ["payload", "payer", "openid", "credential", "message"] { + assert!(!proof.to_string().contains(forbidden)); + } + } + + #[test] + fn callback_keys_and_wxpay_refund_statuses_are_bounded_allowlists() { + let hash = "a".repeat(64); + let long_event_id = "b".repeat(MAX_PAYMENT_GATEWAY_ID_BYTES); + assert_eq!( + payment_callback_key("wxpay", Some(&long_event_id), &hash) + .expect("valid long event IDs should use a bounded hash key"), + format!("wxpay:{hash}") + ); + assert_eq!(wxpay_refund_status(Some("SUCCESS")), "success"); + assert_eq!(wxpay_refund_status(Some("PROCESSING")), "processing"); + for status in [ + Some("CLOSED"), + Some("ABNORMAL"), + Some("credential=secret"), + None, + ] { + assert_eq!(wxpay_refund_status(status), "failed"); + } + } + + #[test] + fn stripe_checkout_does_not_replay_cancelled_idempotent_intent() { + assert!(super::stripe_checkout_response_is_canceled(&json!({ + "id": "pi_cancelled", + "status": "canceled", + "client_secret": "pi_cancelled_secret" + }))); + assert!(super::stripe_checkout_response_is_canceled(&json!({ + "status": " CANCELED " + }))); + assert!(!super::stripe_checkout_response_is_canceled(&json!({ + "status": "requires_payment_method" + }))); + assert!(!super::stripe_checkout_response_is_canceled(&json!({ + "id": "pi_missing_status" + }))); + } + + #[test] + fn alipay_precreate_fallback_requires_a_parsed_business_refusal() { + assert!(alipay_precreate_business_refusal(&json!({ + "code": "40004", + "msg": "Business Failed", + "sub_code": "ACQ.INVALID_PARAMETER" + }))); + // Alipay sometimes omits sub_code for a regular 4xx rejection. The + // HTTP response is still an explicit business refusal in that case. + assert!(alipay_precreate_business_refusal(&json!({ + "code": "40001" + }))); + + for response in [ + json!({"code": "20000", "sub_code": "ACQ.SYSTEM_ERROR"}), + json!({"code": "40004", "sub_code": "ACQ.SYSTEM_ERROR"}), + json!({"code": "40004", "sub_code": "REQUEST_TIMEOUT"}), + json!({"code": "10000", "qr_code": "https://qr.example"}), + json!({"msg": "missing code"}), + ] { + assert!( + !alipay_precreate_business_refusal(&response), + "response must not trigger a page-pay fallback: {response}" + ); + } + } + + #[tokio::test] + async fn payment_http_client_rejects_private_targets_before_connecting() { + for target in [ + "https://127.0.0.1/payment", + "https://169.254.169.254/latest/meta-data", + "https://[::1]/payment", + ] { + let url = url::Url::parse(target).expect("test URL should parse"); + assert!( + public_payment_http_client(&url).await.is_err(), + "private payment target should fail: {target}" + ); + } + } + + #[test] + fn stripe_api_origin_allows_only_benchmarking_addresses() { + let fake = SocketAddr::from(([198, 18, 75, 234], 443)); + for raw_url in [ + "https://api.stripe.com/v1/payment_intents", + "https://API.STRIPE.COM:443/v1/payment_intents", + ] { + let url = url::Url::parse(raw_url).expect("Stripe URL should parse"); + assert!(validate_public_payment_resolved_addrs(&url, &[fake]).is_ok()); + } + assert!(validate_public_payment_resolved_addrs( + &url::Url::parse("https://api.stripe.com/v1/payment_intents").unwrap(), + &[fake, SocketAddr::from(([127, 0, 0, 1], 443))], + ) + .is_err()); + } + + #[test] + fn custom_or_non_default_payment_origins_reject_benchmarking_addresses() { + let fake = SocketAddr::from(([198, 18, 75, 234], 443)); + for raw_url in [ + "https://payments.example.test/v1/payment_intents", + "https://api.stripe.com:8443/v1/payment_intents", + "http://api.stripe.com/v1/payment_intents", + "https://api.stripe.com.evil.test/v1/payment_intents", + "https://api.stripe.com/v1/payment_intents?redirect=internal", + "https://api.stripe.com/v1/payment_intents#fragment", + ] { + let url = url::Url::parse(raw_url).expect("test URL should parse"); + assert!(validate_public_payment_resolved_addrs(&url, &[fake]).is_err()); + } + } +} diff --git a/apps/aether-gateway/src/handlers/shared/payment_gateway_config.rs b/apps/aether-gateway/src/handlers/shared/payment_gateway_config.rs index 84bda3569..fa0ac1e99 100644 --- a/apps/aether-gateway/src/handlers/shared/payment_gateway_config.rs +++ b/apps/aether-gateway/src/handlers/shared/payment_gateway_config.rs @@ -3,6 +3,43 @@ use serde_json::{json, Value}; const REFUND_ENABLED_KEY: &str = "refund_enabled"; const ALLOW_USER_REFUND_KEY: &str = "allow_user_refund"; +pub(crate) fn normalize_payment_https_url(value: &str, field: &str) -> Result { + let trimmed = value.trim(); + let parsed = + url::Url::parse(trimmed).map_err(|_| format!("{field} must be an absolute HTTPS URL"))?; + if parsed.scheme() != "https" + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.fragment().is_some() + { + return Err(format!( + "{field} must be an absolute HTTPS URL without credentials or a fragment" + )); + } + let literal_ip = match parsed.host() { + Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)), + Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)), + Some(url::Host::Domain(_)) | None => None, + }; + if literal_ip.is_some_and(aether_http::is_private_or_reserved_ip) { + return Err(format!( + "{field} must not target a private or reserved address" + )); + } + Ok(trimmed.to_string()) +} + +pub(crate) fn normalize_payment_callback_base_url(value: &str) -> Result { + let normalized = normalize_payment_https_url(value, "callback_base_url")?; + let parsed = url::Url::parse(&normalized) + .map_err(|_| "callback_base_url must be an absolute HTTPS URL".to_string())?; + if parsed.query().is_some() { + return Err("callback_base_url must not contain a query string".to_string()); + } + Ok(normalized.trim_end_matches('/').to_string()) +} + fn json_bool(value: Option<&Value>) -> bool { match value { Some(Value::Bool(value)) => *value, @@ -83,3 +120,43 @@ pub(crate) fn payment_gateway_provider_for_payment_method( _ => None, } } + +#[cfg(test)] +mod tests { + use super::{normalize_payment_callback_base_url, normalize_payment_https_url}; + + #[test] + fn payment_urls_require_absolute_https_without_embedded_credentials() { + assert_eq!( + normalize_payment_https_url(" https://pay.example/submit.php ", "endpoint_url"), + Ok("https://pay.example/submit.php".to_string()) + ); + + for value in [ + "javascript:alert(1)", + "data:text/html,attack", + "//pay.example/submit.php", + "/submit.php", + "http://pay.example/submit.php", + "https://user:secret@pay.example/submit.php", + "https://pay.example/submit.php#fragment", + "https://127.0.0.1/submit.php", + "https://169.254.169.254/latest/meta-data", + "https://[::1]/submit.php", + ] { + assert!( + normalize_payment_https_url(value, "endpoint_url").is_err(), + "unsafe URL should be rejected: {value}" + ); + } + } + + #[test] + fn callback_base_url_rejects_query_strings_and_trims_trailing_slashes() { + assert_eq!( + normalize_payment_callback_base_url("https://app.example/"), + Ok("https://app.example".to_string()) + ); + assert!(normalize_payment_callback_base_url("https://app.example/?tenant=one").is_err()); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/payment_gateway_secret.rs b/apps/aether-gateway/src/handlers/shared/payment_gateway_secret.rs new file mode 100644 index 000000000..003023689 --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/payment_gateway_secret.rs @@ -0,0 +1,428 @@ +use aether_crypto::looks_like_python_fernet_ciphertext; + +use crate::AppState; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +const PAYMENT_GATEWAY_SECRET_ENVELOPE_FAMILY: &str = "aether-payment-gateway-secret-"; +const PAYMENT_GATEWAY_SECRET_ENVELOPE_V2: &str = "aether-payment-gateway-secret-v2:"; +const PAYMENT_GATEWAY_SECRET_ENVELOPE_V3: &str = "aether-payment-gateway-secret-v3:"; +const PAYMENT_GATEWAY_SECRET_PURPOSE_V2: &str = "payment-gateway-secret-bound-v2"; +const PAYMENT_GATEWAY_SECRET_PURPOSE_V3: &str = "payment-gateway-secret-bound-v3"; +const RUNTIME_SECRET_ENVELOPE_FAMILY: &str = "aether-runtime-secret-"; + +const ALIPAY_DEFAULT_GATEWAY_URL: &str = "https://openapi.alipay.com/gateway.do"; +const WXPAY_DEFAULT_BASE_URL: &str = "https://api.mch.weixin.qq.com"; +const STRIPE_DEFAULT_API_URL: &str = "https://api.stripe.com"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PaymentGatewaySecretBinding { + pub(crate) provider: String, + pub(crate) endpoint_url: String, + pub(crate) merchant_id: String, +} + +impl PaymentGatewaySecretBinding { + pub(crate) fn new( + provider: &str, + endpoint_url: &str, + merchant_id: &str, + ) -> Result { + let provider = provider.trim().to_ascii_lowercase(); + if provider.is_empty() || provider.contains('\0') || provider.chars().any(char::is_control) + { + return Err("payment gateway secret provider is invalid"); + } + let endpoint_url = canonical_payment_gateway_endpoint(&provider, endpoint_url)?; + let merchant_id = merchant_id.trim().to_string(); + if merchant_id.chars().any(char::is_control) { + return Err("payment gateway secret merchant_id contains reserved framing"); + } + if merchant_id.len() > 256 { + return Err("payment gateway secret merchant_id is too long"); + } + Ok(Self { + provider, + endpoint_url, + merchant_id, + }) + } + + pub(crate) fn from_record( + record: &aether_data_contracts::repository::billing::PaymentGatewayConfigRecord, + ) -> Result { + Self::new(&record.provider, &record.endpoint_url, &record.merchant_id) + } +} + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct PaymentGatewaySecretProjection { + pub(crate) plaintext: String, + pub(crate) protected: String, + pub(crate) migration_required: bool, +} + +/// Returns whether a stored gateway secret predates destination binding. +/// +/// Legacy Fernet values carry no gateway identity at all, while the v2 +/// envelope authenticates only the provider. Neither format can prove that +/// a value belongs to a newly supplied endpoint/merchant pair, so callers +/// performing a destination-changing mutation must require an explicit +/// replacement secret instead of silently reusing it. +pub(crate) fn payment_gateway_secret_is_legacy_unbound(stored: &str) -> bool { + let stored = stored.trim(); + if stored.is_empty() { + return false; + } + if stored.starts_with(PAYMENT_GATEWAY_SECRET_ENVELOPE_V2) { + return true; + } + if stored.starts_with(PAYMENT_GATEWAY_SECRET_ENVELOPE_FAMILY) + || stored.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY) + || stored.starts_with("aether-") + { + return false; + } + looks_like_python_fernet_ciphertext(stored) +} + +fn payment_gateway_secret_purpose(provider: &str) -> Result { + let provider = provider.trim().to_ascii_lowercase(); + if provider.is_empty() { + return Err("payment gateway secret provider is empty"); + } + Ok(format!( + "{PAYMENT_GATEWAY_SECRET_PURPOSE_V2}\0provider-bytes={}\0{provider}\0field=merchant-key", + provider.len() + )) +} + +fn payment_gateway_secret_purpose_v3( + binding: &PaymentGatewaySecretBinding, +) -> Result { + for value in [ + binding.provider.as_str(), + binding.endpoint_url.as_str(), + binding.merchant_id.as_str(), + ] { + if value.contains('\0') { + return Err("payment gateway secret binding contains reserved framing"); + } + } + Ok(format!( + "{PAYMENT_GATEWAY_SECRET_PURPOSE_V3}\0provider-bytes={}\0{}\0endpoint-url-bytes={}\0{}\0merchant-id-bytes={}\0{}\0field=merchant-key", + binding.provider.len(), + binding.provider, + binding.endpoint_url.len(), + binding.endpoint_url, + binding.merchant_id.len(), + binding.merchant_id, + )) +} + +fn canonical_payment_gateway_endpoint( + provider: &str, + endpoint_url: &str, +) -> Result { + let endpoint_url = endpoint_url.trim(); + let endpoint_url = if endpoint_url.is_empty() { + match provider { + "alipay" => ALIPAY_DEFAULT_GATEWAY_URL, + "wxpay" => WXPAY_DEFAULT_BASE_URL, + "stripe" => STRIPE_DEFAULT_API_URL, + // EPay requires an explicit endpoint at checkout time. Keep an + // explicit marker for legacy records so their secret remains + // bound to the empty value instead of silently changing scope. + _ => return Ok("".to_string()), + } + } else { + endpoint_url + }; + let endpoint_url = super::normalize_payment_https_url(endpoint_url, "endpoint_url") + .map_err(|_| "payment gateway secret endpoint_url is invalid")?; + let mut parsed = url::Url::parse(&endpoint_url) + .map_err(|_| "payment gateway secret endpoint_url is invalid")?; + if parsed.scheme() != "https" + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return Err("payment gateway secret endpoint_url must be an HTTPS URL without credentials or a fragment"); + } + if let Some(host) = parsed.host_str() { + let host = host.trim_end_matches('.').to_ascii_lowercase(); + if host.is_empty() { + return Err("payment gateway secret endpoint_url host is empty"); + } + parsed + .set_host(Some(&host)) + .map_err(|_| "payment gateway secret endpoint_url host is invalid")?; + } + if parsed.port() == Some(443) { + parsed + .set_port(None) + .map_err(|_| "payment gateway secret endpoint_url port is invalid")?; + } + let canonical = parsed.to_string().trim_end_matches('/').to_string(); + // Stripe requests are intentionally sent to the official API origin in + // the checkout/refund implementations below. Accepting a configurable + // destination here would bind the credential to one host while sending + // it to another, which defeats the purpose of destination binding and + // could silently route a live secret through an unintended proxy. + if provider == "stripe" && canonical != STRIPE_DEFAULT_API_URL { + return Err("Stripe endpoint_url must use the official API endpoint"); + } + Ok(canonical) +} + +pub(crate) fn seal_payment_gateway_secret( + state: &AppState, + binding: &PaymentGatewaySecretBinding, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("payment gateway secret contains reserved framing"); + } + let purpose = payment_gateway_secret_purpose_v3(binding)?; + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("payment gateway secret encryption key is not configured")?; + Ok(format!("{PAYMENT_GATEWAY_SECRET_ENVELOPE_V3}{sealed}")) +} + +pub(crate) fn open_payment_gateway_secret( + state: &AppState, + binding: &PaymentGatewaySecretBinding, + stored: &str, +) -> Result { + let purpose = payment_gateway_secret_purpose_v3(binding)?; + if let Some(sealed) = stored.strip_prefix(PAYMENT_GATEWAY_SECRET_ENVELOPE_V3) { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("payment gateway secret authentication or binding failed")?; + if plaintext.contains('\0') { + return Err("payment gateway secret contains reserved framing"); + } + return Ok(PaymentGatewaySecretProjection { + plaintext, + protected: stored.to_string(), + migration_required: false, + }); + } + if let Some(sealed) = stored.strip_prefix(PAYMENT_GATEWAY_SECRET_ENVELOPE_V2) { + // v2 was bound only to provider. Authenticate it with the historical + // purpose, then immediately re-seal under the complete destination + // binding before returning the plaintext to a caller. + let plaintext = open_runtime_secret_payload( + state, + &payment_gateway_secret_purpose(&binding.provider)?, + sealed, + ) + .ok_or("legacy payment gateway secret authentication failed")?; + if plaintext.contains('\0') { + return Err("legacy payment gateway secret contains reserved framing"); + } + let protected = seal_payment_gateway_secret(state, binding, &plaintext)?; + return Ok(PaymentGatewaySecretProjection { + plaintext, + protected, + migration_required: true, + }); + } + if stored.starts_with(PAYMENT_GATEWAY_SECRET_ENVELOPE_FAMILY) { + return Err("unsupported payment gateway secret envelope"); + } + if stored.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY) { + return Err("runtime secret envelope has the wrong purpose"); + } + if !looks_like_python_fernet_ciphertext(stored) { + return Err("payment gateway secret is not an authenticated ciphertext"); + } + + let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("legacy payment gateway secret authentication failed")?; + if plaintext.contains('\0') { + return Err("legacy payment gateway secret contains reserved framing"); + } + let protected = seal_payment_gateway_secret(state, binding, &plaintext)?; + Ok(PaymentGatewaySecretProjection { + plaintext, + protected, + migration_required: true, + }) +} + +#[cfg(test)] +mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + use super::{ + open_payment_gateway_secret, seal_payment_gateway_secret, PaymentGatewaySecretBinding, + PAYMENT_GATEWAY_SECRET_ENVELOPE_V2, PAYMENT_GATEWAY_SECRET_ENVELOPE_V3, + STRIPE_DEFAULT_API_URL, + }; + use crate::handlers::shared::{ + encrypt_catalog_secret_with_fallbacks, seal_runtime_secret_payload, + }; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[test] + fn v3_round_trip_binds_destination_and_rejects_tampering() { + let state = state_with_encryption_key(); + let binding = PaymentGatewaySecretBinding::new( + " EPay ", + "https://payments.example.test:443/checkout/", + " merchant-1 ", + ) + .expect("payment gateway binding should build"); + let sealed = seal_payment_gateway_secret(&state, &binding, "secret-value") + .expect("payment gateway secret should seal"); + + let opened = open_payment_gateway_secret(&state, &binding, &sealed) + .expect("payment gateway secret should open"); + assert_eq!(opened.plaintext, "secret-value"); + assert!(!opened.migration_required); + assert!(open_payment_gateway_secret( + &state, + &PaymentGatewaySecretBinding::new( + "epay", + "https://payments.example.test/other", + "merchant-1", + ) + .unwrap(), + &sealed, + ) + .is_err()); + assert!(open_payment_gateway_secret( + &state, + &PaymentGatewaySecretBinding::new( + "epay", + "https://payments.example.test/checkout/", + "merchant-2", + ) + .unwrap(), + &sealed, + ) + .is_err()); + assert!(open_payment_gateway_secret( + &state, + &PaymentGatewaySecretBinding::new( + "alipay", + "https://payments.example.test/checkout/", + "merchant-1", + ) + .unwrap(), + &sealed, + ) + .is_err()); + + let stripped = sealed + .strip_prefix(PAYMENT_GATEWAY_SECRET_ENVELOPE_V3) + .and_then(|value| value.strip_prefix("aether-runtime-secret-v1:")) + .expect("test value should contain both envelope layers"); + assert!(open_payment_gateway_secret(&state, &binding, stripped).is_err()); + + let mut tampered = sealed.into_bytes(); + let last = tampered + .last_mut() + .expect("sealed value should not be empty"); + *last = if *last == b'A' { b'B' } else { b'A' }; + let tampered = String::from_utf8(tampered).expect("ciphertext should remain utf-8"); + assert!(open_payment_gateway_secret(&state, &binding, &tampered).is_err()); + } + + #[test] + fn authenticated_legacy_values_migrate_to_v3_destination_binding() { + let state = state_with_encryption_key(); + let binding = PaymentGatewaySecretBinding::new( + "epay", + "https://pay.example.test/submit.php", + "merchant-1", + ) + .unwrap(); + let legacy = encrypt_catalog_secret_with_fallbacks(&state, "legacy-secret") + .expect("legacy secret should encrypt"); + let opened = open_payment_gateway_secret(&state, &binding, &legacy) + .expect("legacy secret should migrate"); + assert_eq!(opened.plaintext, "legacy-secret"); + assert!(opened.migration_required); + assert!(opened + .protected + .starts_with("aether-payment-gateway-secret-v3:")); + + let old_v2 = seal_runtime_secret_payload( + &state, + "payment-gateway-secret-bound-v2\0provider-bytes=4\0epay\0field=merchant-key", + "v2-secret", + ) + .expect("legacy v2 secret should encrypt"); + let old_v2 = format!("{PAYMENT_GATEWAY_SECRET_ENVELOPE_V2}{old_v2}"); + let migrated = open_payment_gateway_secret(&state, &binding, &old_v2) + .expect("legacy v2 secret should migrate"); + assert_eq!(migrated.plaintext, "v2-secret"); + assert!(migrated.migration_required); + assert!(migrated + .protected + .starts_with(PAYMENT_GATEWAY_SECRET_ENVELOPE_V3)); + + assert!(open_payment_gateway_secret( + &state, + &binding, + "aether-payment-gateway-secret-v4:unknown", + ) + .is_err()); + assert!(open_payment_gateway_secret(&state, &binding, "plaintext-secret").is_err()); + + let other_runtime = seal_runtime_secret_payload(&state, "another-purpose", "secret") + .expect("runtime secret should seal"); + assert!(open_payment_gateway_secret(&state, &binding, &other_runtime).is_err()); + } + + #[test] + fn canonical_binding_uses_provider_defaults_and_rejects_unsafe_urls() { + let default_alipay = PaymentGatewaySecretBinding::new("ALIPAY", "", "merchant") + .expect("default Alipay endpoint should be accepted"); + let explicit_alipay = PaymentGatewaySecretBinding::new( + "alipay", + "https://OPENAPI.ALIPAY.COM:443/gateway.do", + "merchant", + ) + .expect("explicit Alipay endpoint should be accepted"); + assert_eq!(default_alipay.endpoint_url, explicit_alipay.endpoint_url); + let default_stripe = PaymentGatewaySecretBinding::new("stripe", "", "merchant") + .expect("default Stripe endpoint should be accepted"); + let explicit_stripe = + PaymentGatewaySecretBinding::new("stripe", "https://API.STRIPE.COM:443/", "merchant") + .expect("official Stripe endpoint should be accepted"); + assert_eq!(default_stripe.endpoint_url, STRIPE_DEFAULT_API_URL); + assert_eq!(default_stripe, explicit_stripe); + assert!(PaymentGatewaySecretBinding::new( + "stripe", + "https://stripe-proxy.example.test", + "merchant", + ) + .is_err()); + for endpoint in [ + "http://payments.example.test", + "https://user:password@payments.example.test", + "https://127.0.0.1/pay", + "https://payments.example.test/#fragment", + ] { + assert!( + PaymentGatewaySecretBinding::new("stripe", endpoint, "merchant").is_err(), + "unsafe endpoint should be rejected: {endpoint}" + ); + } + } +} diff --git a/apps/aether-gateway/src/handlers/shared/payment_order_stripe_secret.rs b/apps/aether-gateway/src/handlers/shared/payment_order_stripe_secret.rs new file mode 100644 index 000000000..a7748dd0e --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/payment_order_stripe_secret.rs @@ -0,0 +1,278 @@ +use aether_crypto::looks_like_python_fernet_ciphertext; + +use crate::AppState; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +pub(crate) const STRIPE_CLIENT_SECRET_ENCRYPTED_KEY: &str = "_stripe_client_secret_encrypted"; +const MAX_STRIPE_CLIENT_SECRET_BYTES: usize = 1024; +const PAYMENT_ORDER_STRIPE_SECRET_ENVELOPE_FAMILY: &str = + "aether-payment-order-stripe-client-secret-"; +const PAYMENT_ORDER_STRIPE_SECRET_ENVELOPE_V2: &str = + "aether-payment-order-stripe-client-secret-v2:"; +const PAYMENT_ORDER_STRIPE_SECRET_PURPOSE_V2: &str = "payment-order-stripe-client-secret-bound-v2"; +const AETHER_ENVELOPE_FAMILY: &str = "aether-"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PaymentOrderStripeSecretBinding { + pub(crate) order_no: String, + pub(crate) user_id: Option, + pub(crate) order_kind: String, + pub(crate) payment_provider: String, +} + +impl PaymentOrderStripeSecretBinding { + pub(crate) fn new( + order_no: &str, + user_id: Option<&str>, + order_kind: &str, + payment_provider: &str, + ) -> Result { + validate_identity_component(order_no, "payment order number", 128)?; + if let Some(user_id) = user_id { + validate_identity_component(user_id, "payment order user ID", 128)?; + } + let order_kind = order_kind.trim().to_ascii_lowercase(); + if !matches!(order_kind.as_str(), "wallet_recharge" | "plan_purchase") { + return Err("payment order kind is not eligible for a Stripe client secret"); + } + let payment_provider = payment_provider.trim().to_ascii_lowercase(); + if payment_provider != "stripe" { + return Err("payment order provider is not Stripe"); + } + Ok(Self { + order_no: order_no.to_string(), + user_id: user_id.map(ToOwned::to_owned), + order_kind, + payment_provider, + }) + } + + pub(crate) fn from_order( + order: &aether_data::repository::wallet::StoredAdminPaymentOrder, + ) -> Result { + Self::new( + &order.order_no, + order.user_id.as_deref(), + &order.order_kind, + order + .payment_provider + .as_deref() + .unwrap_or(order.payment_method.as_str()), + ) + } +} + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct PaymentOrderStripeSecretProjection { + pub(crate) plaintext: String, + pub(crate) protected: String, + pub(crate) migration_required: bool, +} + +fn validate_identity_component( + value: &str, + _label: &'static str, + max_bytes: usize, +) -> Result<(), &'static str> { + if value.is_empty() { + return Err("payment order secret binding contains an empty identity component"); + } + if value.as_bytes().len() > max_bytes { + return Err("payment order secret binding identity component is too long"); + } + if value.chars().any(char::is_control) { + return Err("payment order secret binding identity component contains control characters"); + } + Ok(()) +} + +fn payment_order_stripe_secret_purpose( + binding: &PaymentOrderStripeSecretBinding, +) -> Result { + let user_binding = match binding.user_id.as_deref() { + Some(user_id) => format!( + "user-id-present=1\0user-id-bytes={}\0{user_id}", + user_id.len() + ), + None => "user-id-present=0".to_string(), + }; + Ok(format!( + "{PAYMENT_ORDER_STRIPE_SECRET_PURPOSE_V2}\0provider-bytes={}\0{}\0order-no-bytes={}\0{}\0{}\0order-kind-bytes={}\0{}\0field-bytes={}\0{}", + binding.payment_provider.len(), + binding.payment_provider, + binding.order_no.len(), + binding.order_no, + user_binding, + binding.order_kind.len(), + binding.order_kind, + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY.len(), + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY, + )) +} + +pub(crate) fn normalize_stripe_client_secret(value: &str) -> Option<&str> { + let value = value.trim(); + (!value.is_empty() + && value.len() <= MAX_STRIPE_CLIENT_SECRET_BYTES + && value.starts_with("pi_") + && value.contains("_secret_") + && !value.chars().any(char::is_control)) + .then_some(value) +} + +pub(crate) fn seal_payment_order_stripe_client_secret( + state: &AppState, + binding: &PaymentOrderStripeSecretBinding, + plaintext: &str, +) -> Result { + let plaintext = normalize_stripe_client_secret(plaintext) + .ok_or("Stripe client secret format is invalid")?; + let purpose = payment_order_stripe_secret_purpose(binding)?; + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("payment order Stripe client secret encryption key is not configured")?; + Ok(format!("{PAYMENT_ORDER_STRIPE_SECRET_ENVELOPE_V2}{sealed}")) +} + +pub(crate) fn open_payment_order_stripe_client_secret( + state: &AppState, + binding: &PaymentOrderStripeSecretBinding, + stored: &str, +) -> Result { + let purpose = payment_order_stripe_secret_purpose(binding)?; + let observed = stored; + let stored = observed.trim(); + if stored.is_empty() { + return Err("payment order Stripe client secret ciphertext is empty"); + } + + let (plaintext, protected, migration_required) = if let Some(sealed) = + stored.strip_prefix(PAYMENT_ORDER_STRIPE_SECRET_ENVELOPE_V2) + { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("payment order Stripe client secret authentication failed")?; + ( + plaintext, + stored.to_string(), + observed.as_bytes() != stored.as_bytes(), + ) + } else { + if stored.starts_with(PAYMENT_ORDER_STRIPE_SECRET_ENVELOPE_FAMILY) { + return Err("unsupported payment order Stripe client secret envelope"); + } + if stored.starts_with(AETHER_ENVELOPE_FAMILY) { + return Err("secret envelope has the wrong payment order binding"); + } + if !looks_like_python_fernet_ciphertext(stored) { + return Err("payment order Stripe client secret is not an authenticated ciphertext"); + } + let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("legacy payment order Stripe client secret authentication failed")?; + if plaintext.contains('\0') { + return Err("legacy payment order Stripe client secret contains reserved framing"); + } + let protected = seal_payment_order_stripe_client_secret(state, binding, &plaintext)?; + (plaintext, protected, true) + }; + + let plaintext = normalize_stripe_client_secret(&plaintext) + .ok_or("Stripe client secret plaintext format is invalid")? + .to_string(); + Ok(PaymentOrderStripeSecretProjection { + plaintext, + protected, + migration_required, + }) +} + +#[cfg(test)] +mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + use super::{ + open_payment_order_stripe_client_secret, seal_payment_order_stripe_client_secret, + PaymentOrderStripeSecretBinding, + }; + use crate::handlers::shared::{ + encrypt_catalog_secret_with_fallbacks, seal_runtime_secret_payload, + }; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + fn binding(order_no: &str, user_id: &str, order_kind: &str) -> PaymentOrderStripeSecretBinding { + PaymentOrderStripeSecretBinding::new(order_no, Some(user_id), order_kind, "stripe") + .expect("test binding should be valid") + } + + #[test] + fn v2_ciphertext_is_bound_to_order_owner_kind_provider_and_field() { + let state = state_with_encryption_key(); + let source = binding("po-source", "user-a", "wallet_recharge"); + let sealed = + seal_payment_order_stripe_client_secret(&state, &source, "pi_source_secret_capability") + .expect("client secret should seal"); + + let opened = open_payment_order_stripe_client_secret(&state, &source, &sealed) + .expect("matching order should open"); + assert_eq!(opened.plaintext, "pi_source_secret_capability"); + assert!(!opened.migration_required); + for foreign in [ + binding("po-foreign", "user-a", "wallet_recharge"), + binding("po-source", "user-b", "wallet_recharge"), + binding("po-source", "user-a", "plan_purchase"), + ] { + assert!(open_payment_order_stripe_client_secret(&state, &foreign, &sealed).is_err()); + } + assert!(PaymentOrderStripeSecretBinding::new( + "po-source", + Some("user-a"), + "wallet_recharge", + "alipay", + ) + .is_err()); + } + + #[test] + fn reader_migrates_only_real_legacy_fernet_and_rejects_other_envelopes() { + let state = state_with_encryption_key(); + let binding = binding("po-source", "user-a", "wallet_recharge"); + let legacy = encrypt_catalog_secret_with_fallbacks(&state, "pi_legacy_secret_capability") + .expect("legacy secret should encrypt"); + let opened = open_payment_order_stripe_client_secret(&state, &binding, &legacy) + .expect("real legacy Fernet should open"); + assert!(opened.migration_required); + assert!(opened + .protected + .starts_with("aether-payment-order-stripe-client-secret-v2:")); + + for stored in [ + "plaintext-secret", + "aether-payment-order-stripe-client-secret-v3:unknown", + "aether-payment-gateway-secret-v2:foreign", + ] { + assert!(open_payment_order_stripe_client_secret(&state, &binding, stored).is_err()); + } + let foreign_runtime = + seal_runtime_secret_payload(&state, "another-purpose", "pi_x_secret_y") + .expect("runtime secret should seal"); + assert!( + open_payment_order_stripe_client_secret(&state, &binding, &foreign_runtime,).is_err() + ); + + let invalid_legacy = encrypt_catalog_secret_with_fallbacks(&state, "not-a-stripe-secret") + .expect("legacy value should encrypt"); + assert!( + open_payment_order_stripe_client_secret(&state, &binding, &invalid_legacy,).is_err() + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/provider_catalog_credential.rs b/apps/aether-gateway/src/handlers/shared/provider_catalog_credential.rs new file mode 100644 index 000000000..df36e427b --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/provider_catalog_credential.rs @@ -0,0 +1,281 @@ +use aether_crypto::looks_like_python_fernet_ciphertext; + +use crate::AppState; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY: &str = "aether-provider-catalog-credential-"; +const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2: &str = "aether-provider-catalog-credential-v2:"; +const PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2: &str = "provider-catalog-credential-bound-v2"; +const RUNTIME_SECRET_ENVELOPE_FAMILY: &str = "aether-runtime-secret-"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ProviderCatalogCredentialField { + ApiKey, + AuthConfig, +} + +impl ProviderCatalogCredentialField { + fn label(self) -> &'static str { + match self { + Self::ApiKey => "api-key", + Self::AuthConfig => "auth-config", + } + } +} + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct ProviderCatalogCredentialProjection { + pub(crate) plaintext: String, + pub(crate) protected: String, + pub(crate) migration_required: bool, +} + +fn provider_catalog_credential_purpose( + provider_id: &str, + key_id: &str, + field: ProviderCatalogCredentialField, +) -> Result { + if provider_id.is_empty() { + return Err("provider catalog credential provider_id is empty"); + } + if key_id.is_empty() { + return Err("provider catalog credential key_id is empty"); + } + Ok(format!( + "{PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2}\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={}", + provider_id.len(), + key_id.len(), + field.label(), + )) +} + +pub(crate) fn seal_provider_catalog_credential( + state: &AppState, + provider_id: &str, + key_id: &str, + field: ProviderCatalogCredentialField, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("provider catalog credential contains reserved framing"); + } + let purpose = provider_catalog_credential_purpose(provider_id, key_id, field)?; + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("provider catalog credential encryption key is not configured")?; + Ok(format!("{PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2}{sealed}")) +} + +pub(crate) fn open_provider_catalog_credential( + state: &AppState, + provider_id: &str, + key_id: &str, + field: ProviderCatalogCredentialField, + stored: &str, +) -> Result { + let purpose = provider_catalog_credential_purpose(provider_id, key_id, field)?; + if let Some(sealed) = stored.strip_prefix(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2) { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("provider catalog credential authentication failed")?; + if plaintext.contains('\0') { + return Err("provider catalog credential contains reserved framing"); + } + return Ok(ProviderCatalogCredentialProjection { + plaintext, + protected: stored.to_string(), + migration_required: false, + }); + } + if stored.starts_with(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY) { + return Err("unsupported provider catalog credential envelope"); + } + if stored.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY) || stored.starts_with("aether-") { + return Err("Aether secret envelope has the wrong record binding"); + } + if !looks_like_python_fernet_ciphertext(stored) { + return Err("provider catalog credential is not an authenticated ciphertext"); + } + + let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("legacy provider catalog credential authentication failed")?; + // A stripped runtime envelope decrypts to `purpose\0payload`. Rejecting + // reserved framing prevents it from being accepted as a legacy value. + if plaintext.contains('\0') { + return Err("legacy provider catalog credential contains reserved framing"); + } + let protected = + seal_provider_catalog_credential(state, provider_id, key_id, field, &plaintext)?; + Ok(ProviderCatalogCredentialProjection { + plaintext, + protected, + migration_required: true, + }) +} + +#[cfg(test)] +mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + use super::{ + open_provider_catalog_credential, seal_provider_catalog_credential, + ProviderCatalogCredentialField, PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2, + }; + use crate::handlers::shared::{ + encrypt_catalog_secret_with_fallbacks, seal_runtime_secret_payload, + }; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[test] + fn v2_round_trip_binds_provider_key_and_field() { + let state = state_with_encryption_key(); + for field in [ + ProviderCatalogCredentialField::ApiKey, + ProviderCatalogCredentialField::AuthConfig, + ] { + let sealed = seal_provider_catalog_credential( + &state, + "provider-1", + "key-1", + field, + "secret-value", + ) + .expect("credential should seal"); + assert!(sealed.starts_with(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2)); + assert_eq!( + open_provider_catalog_credential(&state, "provider-1", "key-1", field, &sealed,) + .expect("credential should open") + .plaintext, + "secret-value" + ); + assert!(open_provider_catalog_credential( + &state, + "provider-2", + "key-1", + field, + &sealed, + ) + .is_err()); + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-2", + field, + &sealed, + ) + .is_err()); + let other_field = match field { + ProviderCatalogCredentialField::ApiKey => { + ProviderCatalogCredentialField::AuthConfig + } + ProviderCatalogCredentialField::AuthConfig => { + ProviderCatalogCredentialField::ApiKey + } + }; + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + other_field, + &sealed, + ) + .is_err()); + } + } + + #[test] + fn v2_reader_rejects_tampering_unknown_envelopes_and_stripping() { + let state = state_with_encryption_key(); + let sealed = seal_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::ApiKey, + "secret-value", + ) + .expect("credential should seal"); + + let mut tampered = sealed.clone().into_bytes(); + let last = tampered + .last_mut() + .expect("sealed value should not be empty"); + *last = if *last == b'A' { b'B' } else { b'A' }; + let tampered = String::from_utf8(tampered).expect("ciphertext should remain UTF-8"); + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::ApiKey, + &tampered, + ) + .is_err()); + + for invalid in [ + "aether-provider-catalog-credential-v3:unknown", + "aether-payment-gateway-secret-v2:foreign", + "plaintext-secret", + ] { + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::ApiKey, + invalid, + ) + .is_err()); + } + let other_runtime = seal_runtime_secret_payload(&state, "another-purpose", "secret") + .expect("runtime secret should seal"); + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::ApiKey, + &other_runtime, + ) + .is_err()); + + let stripped = sealed + .strip_prefix(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2) + .and_then(|value| value.strip_prefix("aether-runtime-secret-v1:")) + .expect("test value should contain both envelope layers"); + assert!(open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::ApiKey, + stripped, + ) + .is_err()); + } + + #[test] + fn only_real_legacy_fernet_values_are_migrated() { + let state = state_with_encryption_key(); + let legacy = encrypt_catalog_secret_with_fallbacks(&state, "legacy-secret") + .expect("legacy credential should encrypt"); + let opened = open_provider_catalog_credential( + &state, + "provider-1", + "key-1", + ProviderCatalogCredentialField::AuthConfig, + &legacy, + ) + .expect("legacy credential should migrate"); + assert_eq!(opened.plaintext, "legacy-secret"); + assert!(opened.migration_required); + assert!(opened + .protected + .starts_with(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2)); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/provider_ops_credential.rs b/apps/aether-gateway/src/handlers/shared/provider_ops_credential.rs new file mode 100644 index 000000000..9d0251716 --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/provider_ops_credential.rs @@ -0,0 +1,381 @@ +use aether_crypto::looks_like_python_fernet_ciphertext; +use serde_json::Value; +use sha2::{Digest, Sha256}; + +use crate::AppState; + +use super::{ + decrypt_catalog_secret_with_fallbacks, open_runtime_secret_payload, seal_runtime_secret_payload, +}; + +const PROVIDER_OPS_CREDENTIAL_ENVELOPE_FAMILY: &str = "aether-provider-ops-credential-"; +const PROVIDER_OPS_CREDENTIAL_ENVELOPE_V2: &str = "aether-provider-ops-credential-v2:"; +const PROVIDER_OPS_CREDENTIAL_PURPOSE_V2: &str = "provider-ops-credential-bound-v2"; + +pub(crate) const PROVIDER_OPS_PERSISTENT_SECRET_FIELDS: &[&str] = &[ + "api_key", + "password", + "refresh_token", + "session_token", + "session_cookie", + "token_cookie", + "auth_cookie", + "cookie_string", + "cookie", +]; +pub(crate) const PROVIDER_OPS_TRANSIENT_SECRET_FIELDS: &[&str] = &["_cached_access_token"]; +pub(crate) const PROVIDER_OPS_TRANSIENT_METADATA_FIELDS: &[&str] = &["_cached_token_expires_at"]; + +pub(crate) fn provider_ops_credential_field_is_secret(field: &str) -> bool { + PROVIDER_OPS_PERSISTENT_SECRET_FIELDS.contains(&field) + || PROVIDER_OPS_TRANSIENT_SECRET_FIELDS.contains(&field) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ProviderOpsCanonicalDestination { + canonical_base_url: String, + canonical_origin: String, +} + +impl ProviderOpsCanonicalDestination { + pub(crate) fn base_url(&self) -> &str { + &self.canonical_base_url + } + + pub(crate) fn origin(&self) -> &str { + &self.canonical_origin + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ProviderOpsCredentialBinding { + pub(crate) provider_id: String, + pub(crate) architecture_id: String, + pub(crate) auth_type: String, + pub(crate) destination: ProviderOpsCanonicalDestination, + pub(crate) outbound_policy_digest: String, +} + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct ProviderOpsCredentialProjection { + pub(crate) plaintext: String, + pub(crate) protected: String, + pub(crate) migration_required: bool, +} + +pub(crate) fn canonicalize_provider_ops_base_url( + raw: &str, +) -> Result { + let raw = raw.trim(); + if raw.is_empty() { + return Err("Provider Ops base_url 不能为空"); + } + let mut parsed = url::Url::parse(raw).map_err(|_| "Provider Ops base_url 无效")?; + if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() { + return Err("Provider Ops base_url 必须是有效的 HTTP(S) URL"); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err("Provider Ops base_url 不允许包含认证信息"); + } + if parsed.query().is_some() || parsed.fragment().is_some() { + return Err("Provider Ops base_url 不允许包含 query 或 fragment"); + } + + let normalized_path = parsed.path().trim_end_matches('/').to_string(); + parsed.set_path(if normalized_path.is_empty() { + "/" + } else { + normalized_path.as_str() + }); + let canonical_origin = parsed.origin().ascii_serialization(); + let mut canonical_base_url = parsed.to_string(); + if parsed.path() == "/" { + canonical_base_url.truncate(canonical_base_url.len().saturating_sub(1)); + } + + Ok(ProviderOpsCanonicalDestination { + canonical_base_url, + canonical_origin, + }) +} + +pub(crate) fn resolve_provider_ops_same_origin_url( + destination: &ProviderOpsCanonicalDestination, + endpoint: &str, +) -> Result { + let endpoint = endpoint.trim(); + if endpoint.is_empty() { + return Ok(destination.canonical_base_url.clone()); + } + if endpoint.starts_with("//") { + return Err("Provider Ops endpoint 不允许使用 scheme-relative URL"); + } + + let candidate = if endpoint.starts_with("http://") || endpoint.starts_with("https://") { + url::Url::parse(endpoint).map_err(|_| "Provider Ops endpoint URL 无效")? + } else { + if !endpoint.starts_with('/') { + return Err("Provider Ops endpoint 必须是以 / 开头的路径或同源绝对 URL"); + } + let base = url::Url::parse(&format!("{}/", destination.canonical_origin)) + .map_err(|_| "Provider Ops canonical origin 无效")?; + base.join(endpoint) + .map_err(|_| "Provider Ops endpoint 路径无效")? + }; + if !matches!(candidate.scheme(), "http" | "https") + || candidate.host_str().is_none() + || !candidate.username().is_empty() + || candidate.password().is_some() + || candidate.fragment().is_some() + { + return Err("Provider Ops endpoint URL 无效"); + } + if candidate.origin().ascii_serialization() != destination.canonical_origin { + return Err("Provider Ops endpoint 必须与 base_url 同源"); + } + Ok(candidate.to_string()) +} + +pub(crate) fn provider_ops_outbound_policy_digest( + architecture_id: &str, + auth_type: &str, + destination: &ProviderOpsCanonicalDestination, + connector_config: Option<&Value>, + actions: Option<&Value>, +) -> String { + let policy = serde_json::json!({ + "architecture_id": architecture_id, + "auth_type": auth_type, + "canonical_base_url": destination.base_url(), + "canonical_origin": destination.origin(), + "connector_config": connector_config.cloned().unwrap_or_else(|| serde_json::json!({})), + "actions": actions.cloned().unwrap_or_else(|| serde_json::json!({})), + }); + let mut canonical = String::new(); + append_canonical_json(&policy, &mut canonical); + format!("{:x}", Sha256::digest(canonical.as_bytes())) +} + +pub(crate) fn provider_ops_credential_binding_from_config( + provider_id: &str, + provider_ops_config: &serde_json::Map, + effective_base_url: &str, +) -> Result { + if provider_id.trim().is_empty() { + return Err("Provider Ops provider_id 不能为空"); + } + let raw_architecture_id = provider_ops_config + .get("architecture_id") + .and_then(Value::as_str) + .unwrap_or("generic_api") + .trim(); + let architecture_id = + aether_admin::provider::ops::normalize_architecture_id(raw_architecture_id); + if !raw_architecture_id.is_empty() && raw_architecture_id != architecture_id { + return Err("Provider Ops architecture_id 无效"); + } + let connector = provider_ops_config + .get("connector") + .and_then(Value::as_object); + let auth_type = connector + .and_then(|connector| connector.get("auth_type")) + .and_then(Value::as_str) + .unwrap_or("api_key") + .trim(); + if !aether_admin::provider::ops::admin_provider_ops_is_supported_auth_type(auth_type) { + return Err("Provider Ops connector.auth_type 无效"); + } + let destination = canonicalize_provider_ops_base_url(effective_base_url)?; + let outbound_policy_digest = provider_ops_outbound_policy_digest( + architecture_id, + auth_type, + &destination, + connector.and_then(|connector| connector.get("config")), + provider_ops_config.get("actions"), + ); + Ok(ProviderOpsCredentialBinding { + provider_id: provider_id.to_string(), + architecture_id: architecture_id.to_string(), + auth_type: auth_type.to_string(), + destination, + outbound_policy_digest, + }) +} + +pub(crate) fn seal_provider_ops_credential( + state: &AppState, + binding: &ProviderOpsCredentialBinding, + field: &str, + plaintext: &str, +) -> Result { + if plaintext.contains('\0') { + return Err("Provider Ops credential 包含保留分隔符"); + } + let purpose = provider_ops_credential_purpose(binding, field)?; + let sealed = seal_runtime_secret_payload(state, &purpose, plaintext) + .ok_or("gateway 未配置 Provider Ops 加密密钥")?; + Ok(format!("{PROVIDER_OPS_CREDENTIAL_ENVELOPE_V2}{sealed}")) +} + +pub(crate) fn open_provider_ops_credential( + state: &AppState, + binding: &ProviderOpsCredentialBinding, + field: &str, + stored: &str, +) -> Result { + let purpose = provider_ops_credential_purpose(binding, field)?; + if let Some(sealed) = stored.strip_prefix(PROVIDER_OPS_CREDENTIAL_ENVELOPE_V2) { + let plaintext = open_runtime_secret_payload(state, &purpose, sealed) + .ok_or("Provider Ops credential 认证或绑定校验失败")?; + if plaintext.contains('\0') { + return Err("Provider Ops credential 包含保留分隔符"); + } + return Ok(ProviderOpsCredentialProjection { + plaintext, + protected: stored.to_string(), + migration_required: false, + }); + } + if stored.starts_with(PROVIDER_OPS_CREDENTIAL_ENVELOPE_FAMILY) { + return Err("不支持的 Provider Ops credential envelope 版本"); + } + if stored.starts_with("aether-") { + return Err("Aether secret envelope 的 Provider Ops 记录绑定错误"); + } + + let plaintext = if looks_like_python_fernet_ciphertext(stored) { + decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored) + .ok_or("历史 Provider Ops credential 密文无法解密")? + } else { + stored.to_string() + }; + if plaintext.contains('\0') { + return Err("历史 Provider Ops credential 包含保留分隔符"); + } + let protected = seal_provider_ops_credential(state, binding, field, &plaintext)?; + Ok(ProviderOpsCredentialProjection { + plaintext, + protected, + migration_required: !stored.is_empty(), + }) +} + +fn provider_ops_credential_purpose( + binding: &ProviderOpsCredentialBinding, + field: &str, +) -> Result { + for value in [ + binding.provider_id.as_str(), + binding.architecture_id.as_str(), + binding.auth_type.as_str(), + binding.destination.base_url(), + binding.destination.origin(), + binding.outbound_policy_digest.as_str(), + field, + ] { + if value.is_empty() || value.contains('\0') { + return Err("Provider Ops credential binding 无效"); + } + } + Ok(format!( + "{PROVIDER_OPS_CREDENTIAL_PURPOSE_V2}\0provider-id-bytes={}\0{}\0architecture-id-bytes={}\0{}\0auth-type-bytes={}\0{}\0base-url-bytes={}\0{}\0origin-bytes={}\0{}\0policy-sha256={}\0field-bytes={}\0{}", + binding.provider_id.len(), + binding.provider_id, + binding.architecture_id.len(), + binding.architecture_id, + binding.auth_type.len(), + binding.auth_type, + binding.destination.base_url().len(), + binding.destination.base_url(), + binding.destination.origin().len(), + binding.destination.origin(), + binding.outbound_policy_digest, + field.len(), + field, + )) +} + +fn append_canonical_json(value: &Value, output: &mut String) { + match value { + Value::Null => output.push_str("null"), + Value::Bool(value) => output.push_str(if *value { "true" } else { "false" }), + Value::Number(value) => output.push_str(&value.to_string()), + Value::String(value) => output.push_str( + &serde_json::to_string(value).expect("serializing a JSON string cannot fail"), + ), + Value::Array(items) => { + output.push('['); + for (index, item) in items.iter().enumerate() { + if index > 0 { + output.push(','); + } + append_canonical_json(item, output); + } + output.push(']'); + } + Value::Object(map) => { + output.push('{'); + let mut keys = map.keys().collect::>(); + keys.sort_unstable(); + for (index, key) in keys.into_iter().enumerate() { + if index > 0 { + output.push(','); + } + output.push_str( + &serde_json::to_string(key).expect("serializing a JSON key cannot fail"), + ); + output.push(':'); + append_canonical_json(&map[key], output); + } + output.push('}'); + } + } +} + +#[cfg(test)] +mod tests { + use super::{ + canonicalize_provider_ops_base_url, provider_ops_outbound_policy_digest, + resolve_provider_ops_same_origin_url, + }; + + #[test] + fn canonical_destination_normalizes_and_enforces_origin() { + let destination = canonicalize_provider_ops_base_url(" HTTPS://Example.COM:443/api/ ") + .expect("base URL should normalize"); + assert_eq!(destination.base_url(), "https://example.com/api"); + assert_eq!(destination.origin(), "https://example.com"); + assert!(resolve_provider_ops_same_origin_url(&destination, "/v1/me").is_ok()); + assert!( + resolve_provider_ops_same_origin_url(&destination, "https://example.com/v1/me").is_ok() + ); + assert!( + resolve_provider_ops_same_origin_url(&destination, "https://evil.test/v1/me").is_err() + ); + assert!(resolve_provider_ops_same_origin_url(&destination, "//evil.test/v1/me").is_err()); + } + + #[test] + fn policy_digest_is_independent_of_json_object_order() { + let destination = canonicalize_provider_ops_base_url("https://example.com") + .expect("base URL should normalize"); + let left = serde_json::json!({"b": 2, "a": 1}); + let right = serde_json::json!({"a": 1, "b": 2}); + assert_eq!( + provider_ops_outbound_policy_digest( + "generic_api", + "api_key", + &destination, + Some(&left), + None, + ), + provider_ops_outbound_policy_digest( + "generic_api", + "api_key", + &destination, + Some(&right), + None, + ) + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/request_utils.rs b/apps/aether-gateway/src/handlers/shared/request_utils.rs index 5c268e02f..62bec53aa 100644 --- a/apps/aether-gateway/src/handlers/shared/request_utils.rs +++ b/apps/aether-gateway/src/handlers/shared/request_utils.rs @@ -93,15 +93,29 @@ pub(crate) fn sanitize_upstream_path_and_query( let Some(decision) = decision else { return base; }; - if !rust_auth_terminates_provider_credentials(Some(decision)) - || decision.route_family.as_deref() != Some("gemini") - { + if !rust_auth_terminates_provider_credentials(Some(decision)) { return base; } strip_query_param(&base, "key") } +pub(crate) fn security_log_url_origin(value: &str) -> String { + let Ok(parsed) = url::Url::parse(value.trim()) else { + return "-".to_string(); + }; + if !matches!(parsed.scheme(), "http" | "https") { + return "-".to_string(); + } + let Some(host) = parsed.host_str() else { + return "-".to_string(); + }; + match parsed.port() { + Some(port) => format!("{}://{host}:{port}", parsed.scheme()), + None => format!("{}://{host}", parsed.scheme()), + } +} + pub(crate) fn strip_query_param(path_and_query: &str, key_to_strip: &str) -> String { let Some((path, query)) = path_and_query.split_once('?') else { return path_and_query.to_string(); @@ -495,8 +509,14 @@ pub(crate) fn public_support_local_requires_buffered_body( Some( "api_keys_create" | "api_key_install_session_create" - | "management_tokens_create", + | "management_tokens_create" + | "vscodex_pairing_create" + | "vscodex_ws_ticket_create", ), + ) | ( + Some("vscodex"), + http::Method::POST, + Some("pairing_exchange"), ) | ( Some("wallet"), http::Method::POST, @@ -525,3 +545,68 @@ pub(crate) fn local_proxy_route_requires_buffered_body( || internal_proxy_local_requires_buffered_body(request_context) || public_support_local_requires_buffered_body(request_context) } + +#[cfg(test)] +mod tests { + use super::sanitize_upstream_path_and_query; + use crate::control::{GatewayControlAuthContext, GatewayControlDecision}; + + fn authenticated_ai_decision( + route_family: &str, + path_and_query: &str, + ) -> GatewayControlDecision { + let (path, query) = path_and_query + .split_once('?') + .map_or((path_and_query, None), |(path, query)| (path, Some(query))); + let mut decision = GatewayControlDecision::synthetic( + path, + Some("ai_public".to_string()), + Some(route_family.to_string()), + Some("chat".to_string()), + Some(format!("{route_family}:chat")), + ); + decision.public_query_string = query.map(str::to_string); + decision.auth_context = Some(GatewayControlAuthContext { + user_id: "user-1".to_string(), + api_key_id: "key-1".to_string(), + username: None, + api_key_name: None, + balance_remaining: None, + access_allowed: true, + user_rate_limit: None, + api_key_rate_limit: None, + api_key_is_standalone: false, + admin_bypass_limits: false, + local_rejection: None, + allowed_models: None, + ip_rules: None, + verified_api_key_hash: None, + }); + decision + } + + #[test] + fn authenticated_ai_routes_strip_query_api_keys_across_formats() { + for route_family in ["openai", "claude", "gemini"] { + let decision = authenticated_ai_decision( + route_family, + "/v1/chat/completions?key=client-secret&stream=true", + ); + assert_eq!( + sanitize_upstream_path_and_query( + Some(&decision), + "/v1/chat/completions?key=client-secret&stream=true", + ), + "/v1/chat/completions?stream=true" + ); + } + } + + #[test] + fn unauthenticated_routes_preserve_query_parameters() { + assert_eq!( + sanitize_upstream_path_and_query(None, "/v1/chat/completions?key=passthrough"), + "/v1/chat/completions?key=passthrough" + ); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/runtime_secret.rs b/apps/aether-gateway/src/handlers/shared/runtime_secret.rs new file mode 100644 index 000000000..84977037d --- /dev/null +++ b/apps/aether-gateway/src/handlers/shared/runtime_secret.rs @@ -0,0 +1,130 @@ +use crate::AppState; + +use super::catalog::{ + decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_configured_key_fallbacks, + encrypt_catalog_secret_with_fallbacks, +}; + +const RUNTIME_SECRET_ENVELOPE_PREFIX: &str = "aether-runtime-secret-v1:"; + +pub(crate) fn seal_runtime_secret_payload( + state: &AppState, + purpose: &str, + plaintext: &str, +) -> Option { + if purpose.is_empty() || plaintext.contains('\0') { + return None; + } + let protected = format!("{purpose}\0{plaintext}"); + encrypt_catalog_secret_with_fallbacks(state, &protected) + .map(|ciphertext| format!("{RUNTIME_SECRET_ENVELOPE_PREFIX}{ciphertext}")) +} + +pub(crate) fn seal_runtime_secret_payload_with_encryption_key( + encryption_key: Option<&str>, + purpose: &str, + plaintext: &str, +) -> Option { + if purpose.is_empty() || plaintext.contains('\0') { + return None; + } + let protected = format!("{purpose}\0{plaintext}"); + encrypt_catalog_secret_with_configured_key_fallbacks(encryption_key, &protected) + .map(|ciphertext| format!("{RUNTIME_SECRET_ENVELOPE_PREFIX}{ciphertext}")) +} + +pub(crate) fn open_runtime_secret_payload( + state: &AppState, + purpose: &str, + stored: &str, +) -> Option { + open_runtime_secret_payload_with_encryption_key(state.encryption_key(), purpose, stored) +} + +pub(crate) fn open_runtime_secret_payload_with_encryption_key( + encryption_key: Option<&str>, + purpose: &str, + stored: &str, +) -> Option { + if purpose.is_empty() { + return None; + } + let ciphertext = stored.strip_prefix(RUNTIME_SECRET_ENVELOPE_PREFIX)?; + let protected = decrypt_catalog_secret_with_fallbacks(encryption_key, ciphertext)?; + let plaintext = protected.strip_prefix(purpose)?.strip_prefix('\0')?; + (!plaintext.contains('\0')).then(|| plaintext.to_owned()) +} + +pub(crate) fn runtime_secret_payload_is_sealed(value: &str) -> bool { + value.starts_with(RUNTIME_SECRET_ENVELOPE_PREFIX) +} + +#[cfg(test)] +mod tests { + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + use super::{open_runtime_secret_payload, seal_runtime_secret_payload}; + use crate::{data::GatewayDataState, AppState}; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[test] + fn runtime_secret_envelope_hides_payload_and_binds_purpose() { + let state = state_with_encryption_key(); + let payload = r#"{"pkce_verifier":"pkce-runtime-secret-marker"}"#; + + let sealed = seal_runtime_secret_payload(&state, "identity-oauth-state", payload) + .expect("runtime secret should encrypt"); + + assert!(sealed.starts_with("aether-runtime-secret-v1:")); + assert!(!sealed.contains("pkce-runtime-secret-marker")); + assert_eq!( + open_runtime_secret_payload(&state, "identity-oauth-state", &sealed).as_deref(), + Some(payload) + ); + assert!(open_runtime_secret_payload(&state, "provider-oauth-state", &sealed).is_none()); + } + + #[test] + fn runtime_secret_reader_rejects_legacy_plaintext() { + let state = state_with_encryption_key(); + let legacy = r#"{"pkce_verifier":"legacy-verifier"}"#; + + assert!(open_runtime_secret_payload(&state, "identity-oauth-state", legacy).is_none()); + } + + #[test] + fn runtime_secret_envelope_requires_exact_purpose_boundary() { + let state = state_with_encryption_key(); + + assert!(seal_runtime_secret_payload(&state, "", "secret").is_none()); + assert!(seal_runtime_secret_payload(&state, "purpose", "prefix\0secret").is_none()); + + // Structured, field-bound purposes intentionally contain NUL + // separators. A shorter prefix must not be accepted as that same + // purpose, while the complete structured purpose remains readable. + let structured_purpose = "purpose\0prefix"; + let structured = seal_runtime_secret_payload(&state, structured_purpose, "secret") + .expect("structured purpose should encrypt"); + assert_eq!( + open_runtime_secret_payload(&state, structured_purpose, &structured).as_deref(), + Some("secret") + ); + assert!(open_runtime_secret_payload(&state, "purpose", &structured).is_none()); + + // A legacy payload with an extra NUL in the plaintext is rejected; + // the final NUL is not allowed to silently redefine the boundary. + let protected = "purpose\0prefix\0secret"; + let ciphertext = super::encrypt_catalog_secret_with_fallbacks(&state, protected) + .expect("historical ambiguous payload should encrypt"); + let stored = format!("{}{}", super::RUNTIME_SECRET_ENVELOPE_PREFIX, ciphertext); + assert!(open_runtime_secret_payload(&state, "purpose", &stored).is_none()); + } +} diff --git a/apps/aether-gateway/src/handlers/shared/system_config_values.rs b/apps/aether-gateway/src/handlers/shared/system_config_values.rs index 854da27ba..293eac0e1 100644 --- a/apps/aether-gateway/src/handlers/shared/system_config_values.rs +++ b/apps/aether-gateway/src/handlers/shared/system_config_values.rs @@ -1,3 +1,317 @@ +use super::catalog::decrypt_catalog_secret_with_fallbacks; +use super::runtime_secret::{open_runtime_secret_payload, seal_runtime_secret_payload}; +use crate::{AppState, GatewayError}; +use aether_crypto::looks_like_python_fernet_ciphertext; +use std::future::Future; +use url::Url; + +const SYSTEM_CONFIG_SECRET_MIGRATION_RETRIES: usize = 8; +const LDAP_BIND_PASSWORD_MIGRATION_RETRIES: usize = 8; +const BARK_DEVICE_KEY_MIGRATION_RETRIES: usize = 8; +const SYSTEM_CONFIG_SECRET_ENVELOPE_FAMILY_PREFIX: &str = "aether-system-config-secret-"; +const SYSTEM_CONFIG_SECRET_V2_PREFIX: &str = "aether-system-config-secret-v2:"; +const SYSTEM_CONFIG_SECRET_BOUND_PURPOSE_VERSION: &str = "system-config-secret-bound-v2"; +const SMTP_PASSWORD_V3_PREFIX: &str = "aether-smtp-password-v3:"; +const SMTP_PASSWORD_BOUND_PURPOSE_V3: &str = "smtp-password-bound-v3"; +const LDAP_BIND_PASSWORD_ENVELOPE_FAMILY_PREFIX: &str = "aether-ldap-bind-password-"; +const LDAP_BIND_PASSWORD_V2_PREFIX: &str = "aether-ldap-bind-password-v2:"; +const LDAP_BIND_PASSWORD_V3_PREFIX: &str = "aether-ldap-bind-password-v3:"; +const LDAP_BIND_PASSWORD_BOUND_PURPOSE: &str = "ldap-bind-password-bound-v2"; +const LDAP_BIND_PASSWORD_BOUND_PURPOSE_V3: &str = "ldap-bind-password-bound-v3"; +const BARK_DEVICE_KEY_V2_PREFIX: &str = "aether-bark-device-key-v2:"; +const BARK_DEVICE_KEY_BOUND_PURPOSE_V2: &str = "bark-device-key-bound-v2"; +const BARK_DEVICE_KEY_CONFIG_KEY: &str = "module.bark_push.device_key"; +const RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX: &str = "aether-runtime-secret-"; + +fn system_config_secret_purpose(key: &str) -> String { + let key = aether_admin::system::normalize_admin_system_config_key(key); + format!( + "{SYSTEM_CONFIG_SECRET_BOUND_PURPOSE_VERSION}\0key-bytes={}\0{key}", + key.len() + ) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SmtpPasswordBinding { + pub(crate) host: String, + pub(crate) port: u16, + pub(crate) user: String, + pub(crate) use_tls: bool, + pub(crate) use_ssl: bool, +} + +pub(crate) fn smtp_password_binding( + host: &str, + port: u16, + user: Option<&str>, + use_tls: bool, + use_ssl: bool, +) -> Option { + let host = host.trim(); + let user = user.unwrap_or("").trim(); + if host.is_empty() + || user.is_empty() + || host.contains('\0') + || user.contains('\0') + || host.bytes().any(|byte| matches!(byte, b'\r' | b'\n')) + || user.bytes().any(|byte| matches!(byte, b'\r' | b'\n')) + { + return None; + } + Some(SmtpPasswordBinding { + host: host.to_ascii_lowercase(), + port, + user: user.to_string(), + use_tls, + use_ssl, + }) +} + +fn smtp_password_purpose(binding: &SmtpPasswordBinding) -> String { + format!( + "{SMTP_PASSWORD_BOUND_PURPOSE_V3}\0host-bytes={}\0{}\0port={}\0user-bytes={}\0{}\0tls={}\0ssl={}\0field-bytes={}\0smtp_password", + binding.host.len(), + binding.host, + binding.port, + binding.user.len(), + binding.user, + if binding.use_tls { 1 } else { 0 }, + if binding.use_ssl { 1 } else { 0 }, + "smtp_password".len(), + ) +} + +pub(crate) fn encrypt_smtp_password( + state: &AppState, + binding: &SmtpPasswordBinding, + plaintext: &str, +) -> Option { + if plaintext.contains('\0') { + return None; + } + seal_runtime_secret_payload(state, &smtp_password_purpose(binding), plaintext) + .map(|sealed| format!("{SMTP_PASSWORD_V3_PREFIX}{sealed}")) +} + +fn decrypt_smtp_password_v3( + state: &AppState, + binding: &SmtpPasswordBinding, + stored: &str, +) -> Option { + let sealed = stored.strip_prefix(SMTP_PASSWORD_V3_PREFIX)?; + open_runtime_secret_payload(state, &smtp_password_purpose(binding), sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +pub(crate) async fn decrypt_or_migrate_smtp_password( + state: &AppState, + binding: &SmtpPasswordBinding, + stored: String, +) -> Result { + if let Some(plaintext) = decrypt_smtp_password_v3(state, binding, stored.trim()) { + return Ok(plaintext); + } + if stored.trim().starts_with(SMTP_PASSWORD_V3_PREFIX) + || stored.trim().starts_with("aether-smtp-password-") + { + return Err(system_config_secret_error( + "stored SMTP password cannot be decrypted", + )); + } + let plaintext = decrypt_system_config_secret(state, "smtp_password", stored.trim()) + .or_else(|| { + (!stored.trim().is_empty() && !looks_like_python_fernet_ciphertext(stored.trim())) + .then(|| stored.trim().to_string()) + }) + .ok_or_else(|| system_config_secret_error("stored SMTP password cannot be decrypted"))?; + if plaintext.contains('\0') { + return Err(system_config_secret_error( + "stored SMTP password contains reserved secret framing", + )); + } + let replacement = encrypt_smtp_password(state, binding, &plaintext) + .ok_or_else(|| system_config_secret_error("SMTP password migration is unavailable"))?; + if state + .compare_and_set_system_config_string_value("smtp_password", stored.trim(), &replacement) + .await? + { + return Ok(plaintext); + } + let current = state + .read_system_config_json_value_strong("smtp_password") + .await? + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .ok_or_else(|| system_config_secret_error("stored SMTP password is unavailable"))?; + decrypt_smtp_password_v3(state, binding, current.trim()) + .ok_or_else(|| system_config_secret_error("stored SMTP password changed during migration")) +} + +pub(crate) fn encrypt_system_config_secret( + state: &AppState, + key: &str, + plaintext: &str, +) -> Option { + if plaintext.contains('\0') { + return None; + } + let purpose = system_config_secret_purpose(key); + seal_runtime_secret_payload(state, &purpose, plaintext) + .map(|sealed| format!("{SYSTEM_CONFIG_SECRET_V2_PREFIX}{sealed}")) +} + +pub(crate) fn decrypt_system_config_secret( + state: &AppState, + key: &str, + stored: &str, +) -> Option { + let sealed = stored.strip_prefix(SYSTEM_CONFIG_SECRET_V2_PREFIX)?; + let purpose = system_config_secret_purpose(key); + open_runtime_secret_payload(state, &purpose, sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct LdapBindPasswordBinding { + server_url: String, + bind_dn: String, + base_dn: String, + use_starttls: bool, +} + +/// Parse and canonicalize an LDAP transport endpoint. +/// +/// LDAP simple binds carry credentials, so plaintext `ldap://` is only valid +/// when StartTLS is explicitly requested. The parser deliberately does not +/// reject private or loopback hosts: LDAP deployments commonly run on an +/// internal network and this function is a transport-integrity check, not an +/// outbound SSRF policy. The `mockldap` scheme exists only in test builds. +pub(crate) fn normalize_ldap_transport_server_url(raw: &str, use_starttls: bool) -> Option { + #[cfg(test)] + { + // Gateway unit/integration fixtures use an in-process mock endpoint. Keep + // this exception behind the gateway test configuration; production code + // always uses the strict parser without custom schemes. + return aether_admin::system::normalize_ldap_transport_server_url_for_tests( + raw, + use_starttls, + ); + } + #[cfg(not(test))] + { + aether_admin::system::normalize_ldap_transport_server_url(raw, use_starttls) + } +} + +/// Return whether an LDAP user search filter is safe to use with the escaped +/// `{username}` substitution performed by the login path. LDAP filters are +/// always parenthesized; keeping the same bounded shape at every config +/// ingress prevents malformed/imported values from reaching the query layer. +pub(crate) fn ldap_search_filter_is_valid(value: &str) -> bool { + aether_admin::system::ldap_search_filter_is_valid(value) +} + +pub(crate) fn ldap_distinguished_name_is_valid(value: &str) -> bool { + aether_admin::system::ldap_distinguished_name_is_valid(value) +} + +pub(crate) fn ldap_attribute_description_is_valid(value: &str) -> bool { + aether_admin::system::ldap_attribute_description_is_valid(value) +} + +pub(crate) fn ldap_module_config_is_valid( + config: Option<&aether_data::repository::auth_modules::StoredLdapModuleConfig>, +) -> bool { + config.is_some_and(|config| { + normalize_ldap_transport_server_url(&config.server_url, config.use_starttls).is_some() + && aether_admin::system::ldap_module_config_fields_are_valid(config) + }) +} + +fn canonical_ldap_server_url(raw: &str, use_starttls: bool) -> Option { + normalize_ldap_transport_server_url(raw, use_starttls) +} + +fn ldap_bind_password_binding( + config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, +) -> Option { + let server_url = canonical_ldap_server_url(&config.server_url, config.use_starttls)?; + let bind_dn = config.bind_dn.trim(); + let base_dn = config.base_dn.trim(); + if !ldap_distinguished_name_is_valid(&config.bind_dn) + || !ldap_distinguished_name_is_valid(&config.base_dn) + { + return None; + } + Some(LdapBindPasswordBinding { + server_url, + bind_dn: bind_dn.to_string(), + base_dn: base_dn.to_string(), + use_starttls: config.use_starttls, + }) +} + +pub(crate) fn ldap_bind_password_binding_matches( + stored: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + replacement: &aether_data::repository::auth_modules::StoredLdapModuleConfig, +) -> Result { + let stored = + ldap_bind_password_binding(stored).ok_or("stored LDAP bind password binding is invalid")?; + let replacement = ldap_bind_password_binding(replacement) + .ok_or("replacement LDAP bind password binding is invalid")?; + Ok(stored == replacement) +} + +fn ldap_bind_password_purpose_v3(binding: &LdapBindPasswordBinding) -> String { + format!( + "{LDAP_BIND_PASSWORD_BOUND_PURPOSE_V3}\0server-url-bytes={}\0{}\0bind-dn-bytes={}\0{}\0base-dn-bytes={}\0{}\0starttls={}\0field-bytes={}\0bind_password_encrypted", + binding.server_url.len(), + binding.server_url, + binding.bind_dn.len(), + binding.bind_dn, + binding.base_dn.len(), + binding.base_dn, + if binding.use_starttls { 1 } else { 0 }, + "bind_password_encrypted".len(), + ) +} + +pub(crate) fn encrypt_ldap_bind_password( + state: &AppState, + config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + plaintext: &str, +) -> Option { + if plaintext.contains('\0') { + return None; + } + let binding = ldap_bind_password_binding(config)?; + seal_runtime_secret_payload(state, &ldap_bind_password_purpose_v3(&binding), plaintext) + .map(|sealed| format!("{LDAP_BIND_PASSWORD_V3_PREFIX}{sealed}")) +} + +fn decrypt_ldap_bind_password_v2(state: &AppState, stored: &str) -> Option { + let sealed = stored.strip_prefix(LDAP_BIND_PASSWORD_V2_PREFIX)?; + open_runtime_secret_payload(state, LDAP_BIND_PASSWORD_BOUND_PURPOSE, sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +fn decrypt_ldap_bind_password_v3( + state: &AppState, + config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + stored: &str, +) -> Option { + let sealed = stored.strip_prefix(LDAP_BIND_PASSWORD_V3_PREFIX)?; + let binding = ldap_bind_password_binding(config)?; + open_runtime_secret_payload(state, &ldap_bind_password_purpose_v3(&binding), sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +fn stored_secret_uses_known_envelope_family(value: &str) -> bool { + value.starts_with(SYSTEM_CONFIG_SECRET_ENVELOPE_FAMILY_PREFIX) + || value.starts_with(LDAP_BIND_PASSWORD_ENVELOPE_FAMILY_PREFIX) + || value.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX) + || value.starts_with("aether-") +} + pub(crate) fn module_available_from_env(env_key: &str, default_available: bool) -> bool { match std::env::var(env_key) { Ok(value) => matches!( @@ -38,3 +352,915 @@ pub(crate) fn system_config_string(value: Option<&serde_json::Value>) -> Option< _ => None, } } + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct BarkDeviceKeyBinding { + pub(crate) server_url: String, +} + +/// Canonicalize the Bark base URL before it participates in a secret binding. +/// Host spelling and default ports are normalized so equivalent destinations +/// use the same binding, while credentials and URL-controlled request data are +/// rejected entirely. +pub(crate) fn canonical_bark_server_url(raw: &str) -> Option { + if raw.contains('\0') { + return None; + } + let raw = raw.trim().trim_end_matches('/'); + if raw.is_empty() || raw.contains('@') { + return None; + } + let mut parsed = Url::parse(raw).ok()?; + if !matches!(parsed.scheme(), "https" | "http") + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return None; + } + let host = parsed + .host_str()? + .trim_end_matches('.') + .to_ascii_lowercase(); + if host.is_empty() { + return None; + } + parsed.set_host(Some(&host)).ok()?; + let default_port = if parsed.scheme() == "https" { 443 } else { 80 }; + if parsed.port() == Some(default_port) { + parsed.set_port(None).ok()?; + } + Some(parsed.as_str().trim_end_matches('/').to_string()) +} + +pub(crate) fn bark_device_key_binding(server_url: &str) -> Option { + Some(BarkDeviceKeyBinding { + server_url: canonical_bark_server_url(server_url)?, + }) +} + +fn bark_device_key_purpose(binding: &BarkDeviceKeyBinding) -> String { + format!( + "{BARK_DEVICE_KEY_BOUND_PURPOSE_V2}\0server-url-bytes={}\0{}\0field-bytes={}\0{}", + binding.server_url.len(), + binding.server_url, + BARK_DEVICE_KEY_CONFIG_KEY.len(), + BARK_DEVICE_KEY_CONFIG_KEY, + ) +} + +pub(crate) fn encrypt_bark_device_key( + state: &AppState, + binding: &BarkDeviceKeyBinding, + plaintext: &str, +) -> Option { + if plaintext.contains('\0') { + return None; + } + seal_runtime_secret_payload(state, &bark_device_key_purpose(binding), plaintext) + .map(|sealed| format!("{BARK_DEVICE_KEY_V2_PREFIX}{sealed}")) +} + +fn decrypt_bark_device_key_v2( + state: &AppState, + binding: &BarkDeviceKeyBinding, + stored: &str, +) -> Option { + let sealed = stored.strip_prefix(BARK_DEVICE_KEY_V2_PREFIX)?; + open_runtime_secret_payload(state, &bark_device_key_purpose(binding), sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +pub(crate) async fn decrypt_or_migrate_bark_device_key( + state: &AppState, + binding: &BarkDeviceKeyBinding, + stored_value: String, +) -> Result { + let mut observed_raw = stored_value; + for _ in 0..BARK_DEVICE_KEY_MIGRATION_RETRIES { + let current_raw = + read_strong_system_config_secret(state, BARK_DEVICE_KEY_CONFIG_KEY).await?; + if current_raw != observed_raw { + observed_raw = current_raw; + continue; + } + let observed = observed_raw.trim(); + if observed.starts_with(BARK_DEVICE_KEY_V2_PREFIX) { + return decrypt_bark_device_key_v2(state, binding, observed).ok_or_else(|| { + system_config_secret_error("stored Bark device key cannot be decrypted") + }); + } + + // Older Bark entries were written by the generic system-config secret + // path. They may migrate once, but all new writes are destination-bound. + let plaintext = if let Some(plaintext) = + decrypt_system_config_secret(state, BARK_DEVICE_KEY_CONFIG_KEY, observed) + { + plaintext + } else { + if stored_secret_uses_known_envelope_family(observed) { + return Err(system_config_secret_error( + "stored Bark device key cannot be decrypted", + )); + } + match decrypt_catalog_secret_with_fallbacks(state.encryption_key(), observed) { + Some(plaintext) => plaintext, + None if looks_like_python_fernet_ciphertext(observed) => { + return Err(system_config_secret_error( + "stored Bark device key cannot be decrypted", + )); + } + None => observed.to_string(), + } + }; + if plaintext.contains('\0') { + return Err(system_config_secret_error( + "stored Bark device key contains reserved secret framing", + )); + } + let encrypted = encrypt_bark_device_key(state, binding, &plaintext).ok_or_else(|| { + system_config_secret_error("Bark device key migration is unavailable") + })?; + if state + .compare_and_set_system_config_string_value( + BARK_DEVICE_KEY_CONFIG_KEY, + &observed_raw, + &encrypted, + ) + .await? + { + return Ok(plaintext); + } + observed_raw = read_strong_system_config_secret(state, BARK_DEVICE_KEY_CONFIG_KEY).await?; + } + + Err(system_config_secret_error( + "Bark device key migration did not stabilize", + )) +} + +pub(crate) async fn decrypt_or_migrate_system_config_secret( + state: &AppState, + key: &str, + stored_value: String, +) -> Result { + decrypt_or_migrate_system_config_secret_with_before_compare( + state, + key, + stored_value, + || async {}, + ) + .await +} + +pub(crate) async fn decrypt_or_migrate_ldap_bind_password( + state: &AppState, + config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, +) -> Result, GatewayError> { + let mut current = config.clone(); + for _ in 0..LDAP_BIND_PASSWORD_MIGRATION_RETRIES { + let Some(observed_raw) = current + .bind_password_encrypted + .as_deref() + .filter(|value| !value.trim().is_empty()) + else { + return Ok(None); + }; + let observed = observed_raw.trim(); + if observed.starts_with(LDAP_BIND_PASSWORD_V3_PREFIX) { + let plaintext = + decrypt_ldap_bind_password_v3(state, ¤t, observed).ok_or_else(|| { + ldap_bind_password_error("stored LDAP bind password cannot be decrypted") + })?; + return Ok((!plaintext.trim().is_empty()).then_some(plaintext)); + } + if observed.starts_with(LDAP_BIND_PASSWORD_V2_PREFIX) { + let plaintext = decrypt_ldap_bind_password_v2(state, observed).ok_or_else(|| { + ldap_bind_password_error("stored LDAP bind password cannot be decrypted") + })?; + if plaintext.trim().is_empty() { + return Ok(None); + } + // v2 was bound only to the LDAP field name. Re-seal it with the + // complete current LDAP destination before returning it so a + // legacy ciphertext cannot remain portable across configurations. + let encrypted = + encrypt_ldap_bind_password(state, ¤t, &plaintext).ok_or_else(|| { + ldap_bind_password_error("LDAP bind password migration is unavailable") + })?; + if state + .compare_and_swap_ldap_bind_password(observed_raw, &encrypted) + .await? + { + return Ok(Some(plaintext)); + } + current = state.get_ldap_module_config().await?.ok_or_else(|| { + ldap_bind_password_error("stored LDAP configuration is unavailable") + })?; + continue; + } + if stored_secret_uses_known_envelope_family(observed) { + return Err(ldap_bind_password_error( + "stored LDAP bind password cannot be decrypted", + )); + } + let plaintext = + match decrypt_catalog_secret_with_fallbacks(state.encryption_key(), observed) { + Some(plaintext) => plaintext, + None if looks_like_python_fernet_ciphertext(observed) => { + return Err(ldap_bind_password_error( + "stored LDAP bind password cannot be decrypted", + )); + } + None => observed.to_string(), + }; + if plaintext.contains('\0') { + return Err(ldap_bind_password_error( + "stored LDAP bind password contains reserved secret framing", + )); + } + let encrypted = + encrypt_ldap_bind_password(state, ¤t, &plaintext).ok_or_else(|| { + ldap_bind_password_error("LDAP bind password migration is unavailable") + })?; + if state + .compare_and_swap_ldap_bind_password(observed_raw, &encrypted) + .await? + { + return Ok((!plaintext.trim().is_empty()).then_some(plaintext)); + } + current = state + .get_ldap_module_config() + .await? + .ok_or_else(|| ldap_bind_password_error("stored LDAP configuration is unavailable"))?; + } + + Err(ldap_bind_password_error( + "LDAP bind password migration did not stabilize", + )) +} + +async fn decrypt_or_migrate_system_config_secret_with_before_compare( + state: &AppState, + key: &str, + stored_value: String, + before_compare: BeforeCompare, +) -> Result +where + BeforeCompare: Fn() -> CompareFuture, + CompareFuture: Future, +{ + let mut observed_raw = stored_value; + for _ in 0..SYSTEM_CONFIG_SECRET_MIGRATION_RETRIES { + let current_raw = read_strong_system_config_secret(state, key).await?; + if current_raw != observed_raw { + observed_raw = current_raw; + continue; + } + let observed = observed_raw.trim(); + + if observed.starts_with(SYSTEM_CONFIG_SECRET_V2_PREFIX) { + let plaintext = + decrypt_system_config_secret(state, key, observed).ok_or_else(|| { + system_config_secret_error( + "stored system configuration secret cannot be decrypted", + ) + })?; + return Ok(plaintext); + } + if stored_secret_uses_known_envelope_family(observed) { + return Err(system_config_secret_error( + "stored system configuration secret cannot be decrypted", + )); + } + let plaintext = + match decrypt_catalog_secret_with_fallbacks(state.encryption_key(), observed) { + Some(plaintext) => plaintext, + None if looks_like_python_fernet_ciphertext(observed) => { + return Err(system_config_secret_error( + "stored system configuration secret cannot be decrypted", + )); + } + None => observed.to_string(), + }; + if plaintext.contains('\0') { + return Err(system_config_secret_error( + "stored system configuration secret contains reserved secret framing", + )); + } + let encrypted = encrypt_system_config_secret(state, key, &plaintext).ok_or_else(|| { + system_config_secret_error("system configuration secret migration is unavailable") + })?; + + before_compare().await; + if state + .compare_and_set_system_config_string_value(key, &observed_raw, &encrypted) + .await? + { + return Ok(plaintext); + } + observed_raw = read_strong_system_config_secret(state, key).await?; + } + + Err(system_config_secret_error( + "system configuration secret migration did not stabilize", + )) +} + +async fn read_strong_system_config_secret( + state: &AppState, + key: &str, +) -> Result { + state + .read_system_config_json_value_strong(key) + .await? + .and_then(|value| { + value + .as_str() + .filter(|value| !value.trim().is_empty()) + .map(ToOwned::to_owned) + }) + .ok_or_else(|| { + system_config_secret_error("stored system configuration secret is unavailable") + }) +} + +fn system_config_secret_error(message: &str) -> GatewayError { + GatewayError::Internal(message.to_string()) +} + +fn ldap_bind_password_error(message: &str) -> GatewayError { + GatewayError::Internal(message.to_string()) +} + +#[cfg(test)] +mod tests { + use super::{ + bark_device_key_binding, decrypt_bark_device_key_v2, decrypt_ldap_bind_password_v2, + decrypt_ldap_bind_password_v3, decrypt_or_migrate_bark_device_key, + decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_system_config_secret, + decrypt_or_migrate_system_config_secret_with_before_compare, decrypt_system_config_secret, + encrypt_bark_device_key, encrypt_ldap_bind_password, encrypt_system_config_secret, + ldap_module_config_is_valid, normalize_ldap_transport_server_url, + LDAP_BIND_PASSWORD_V2_PREFIX, LDAP_BIND_PASSWORD_V3_PREFIX, SYSTEM_CONFIG_SECRET_V2_PREFIX, + }; + use crate::data::GatewayDataState; + use crate::AppState; + use aether_crypto::{ + encrypt_python_fernet_plaintext, looks_like_python_fernet_ciphertext, + DEVELOPMENT_ENCRYPTION_KEY, + }; + use aether_data::repository::auth_modules::{ + AuthModuleReadRepository, InMemoryAuthModuleReadRepository, StoredLdapModuleConfig, + }; + use futures_util::future::join_all; + use serde_json::json; + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; + use tokio::sync::Barrier; + + const TEST_KEY: &str = "smtp_password"; + const TEST_SECRET: &str = "legacy-plaintext-password"; + const BARK_DEVICE_KEY: &str = "module.bark_push.device_key"; + + fn state_with_stored_secret(value: &str) -> AppState { + state_with_named_stored_secret(TEST_KEY, value) + } + + fn state_with_named_stored_secret(key: &str, value: &str) -> AppState { + let data = GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests([(key.to_string(), json!(value))]); + let mut state = AppState::new().expect("gateway state should build"); + state.replace_data_state(Arc::new(data)); + state + } + + fn ldap_config(bind_password: &str) -> StoredLdapModuleConfig { + StoredLdapModuleConfig { + server_url: "ldaps://ldap.example.com".to_string(), + bind_dn: "cn=admin,dc=example,dc=com".to_string(), + bind_password_encrypted: Some(bind_password.to_string()), + base_dn: "dc=example,dc=com".to_string(), + user_search_filter: Some("(uid={username})".to_string()), + username_attr: Some("uid".to_string()), + email_attr: Some("mail".to_string()), + display_name_attr: Some("displayName".to_string()), + is_enabled: true, + is_exclusive: false, + use_starttls: false, + connect_timeout: Some(10), + } + } + + #[test] + fn ldap_transport_url_requires_tls_and_rejects_url_control_data() { + assert_eq!( + normalize_ldap_transport_server_url(" ldaps://LDAP.Example.COM:636/ ", false) + .as_deref(), + Some("ldaps://ldap.example.com") + ); + assert_eq!( + normalize_ldap_transport_server_url("ldap://10.20.30.40:389", true).as_deref(), + Some("ldap://10.20.30.40") + ); + assert!(normalize_ldap_transport_server_url("ldap://10.20.30.40", false).is_none()); + assert!( + normalize_ldap_transport_server_url("ldap://user:password@ldap.example.com", true) + .is_none() + ); + assert!(normalize_ldap_transport_server_url("ldaps://@ldap.example.com", false).is_none()); + assert!( + normalize_ldap_transport_server_url("ldaps://ldap.example.com?x=1", false).is_none() + ); + assert!( + normalize_ldap_transport_server_url("ldaps://ldap.example.com#fragment", false) + .is_none() + ); + assert!( + normalize_ldap_transport_server_url("ldaps://ldap.example.com/dc=example", false) + .is_none() + ); + assert!(normalize_ldap_transport_server_url("https://ldap.example.com", false).is_none()); + assert!(normalize_ldap_transport_server_url("ldaps://ldap.example.com\n", false).is_none()); + // The in-process mock is intentionally available only to test builds. + assert!( + normalize_ldap_transport_server_url("mockldap://ldap.example.com", false).is_some() + ); + } + + #[test] + fn gateway_ldap_validation_reuses_shared_field_rules_with_test_transport() { + let mut config = ldap_config("sealed-password"); + config.server_url = "mockldap://ldap.example.com".to_string(); + config.use_starttls = false; + assert!(ldap_module_config_is_valid(Some(&config))); + + config.user_search_filter = Some("(uid={username})(objectClass=*)".to_string()); + assert!(!ldap_module_config_is_valid(Some(&config))); + config.user_search_filter = Some("(uid={username})".to_string()); + config.username_attr = Some("uid)(|(objectClass=*)".to_string()); + assert!(!ldap_module_config_is_valid(Some(&config))); + config.username_attr = Some("uid".to_string()); + config.bind_dn = "cn=admin,dc=example,dc=com\n".to_string(); + assert!(!ldap_module_config_is_valid(Some(&config))); + } + + fn state_with_ldap_bind_password( + value: &str, + ) -> (AppState, Arc) { + let repository = Arc::new(InMemoryAuthModuleReadRepository::seed( + Vec::new(), + Some(ldap_config(value)), + )); + let data = GatewayDataState::with_auth_module_repository_for_tests(repository.clone()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let mut state = AppState::new().expect("gateway state should build"); + state.replace_data_state(Arc::new(data)); + (state, repository) + } + + #[tokio::test] + async fn legacy_ldap_bind_password_is_lazily_migrated() { + let (state, repository) = state_with_ldap_bind_password(TEST_SECRET); + let config = repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .expect("LDAP config should exist"); + + let plaintext = decrypt_or_migrate_ldap_bind_password(&state, &config) + .await + .expect("legacy LDAP bind password should migrate"); + assert_eq!(plaintext.as_deref(), Some(TEST_SECRET)); + + let stored = repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .and_then(|config| config.bind_password_encrypted) + .expect("LDAP bind password should exist"); + assert!(stored.starts_with(LDAP_BIND_PASSWORD_V3_PREFIX)); + assert_eq!( + decrypt_ldap_bind_password_v3(&state, &config, &stored) + .expect("migrated LDAP password should decrypt"), + TEST_SECRET + ); + } + + #[tokio::test] + async fn tampered_ldap_bind_password_ciphertext_fails_closed() { + let mut tampered = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_SECRET) + .expect("LDAP password should encrypt"); + tampered.replace_range(tampered.len() - 2.., "AA"); + assert!(looks_like_python_fernet_ciphertext(&tampered)); + let (state, repository) = state_with_ldap_bind_password(&tampered); + let config = repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .expect("LDAP config should exist"); + + let error = decrypt_or_migrate_ldap_bind_password(&state, &config) + .await + .expect_err("tampered LDAP ciphertext must not be used as plaintext"); + assert!(format!("{error:?}").contains("cannot be decrypted")); + assert_eq!( + repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .and_then(|config| config.bind_password_encrypted), + Some(tampered) + ); + } + + #[tokio::test] + async fn bound_secret_envelopes_cannot_move_between_ldap_and_system_config() { + let (ldap_state, _) = state_with_ldap_bind_password(TEST_SECRET); + let ldap_sealed = + encrypt_ldap_bind_password(&ldap_state, &ldap_config("legacy"), TEST_SECRET) + .expect("LDAP bind password should seal"); + let system_sealed = encrypt_system_config_secret(&ldap_state, TEST_KEY, TEST_SECRET) + .expect("system config secret should seal"); + + let (ldap_with_system_secret, repository) = state_with_ldap_bind_password(&system_sealed); + let config = repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .expect("LDAP config should exist"); + let ldap_error = decrypt_or_migrate_ldap_bind_password(&ldap_with_system_secret, &config) + .await + .expect_err("system config ciphertext must not become an LDAP password"); + assert!(ldap_error.into_message().contains("cannot be decrypted")); + + let system_with_ldap_secret = state_with_stored_secret(&ldap_sealed); + let system_error = decrypt_or_migrate_system_config_secret( + &system_with_ldap_secret, + TEST_KEY, + ldap_sealed.clone(), + ) + .await + .expect_err("LDAP ciphertext must not become a system config secret"); + assert!(system_error.into_message().contains("cannot be decrypted")); + + let stripped_system = system_sealed + .strip_prefix(SYSTEM_CONFIG_SECRET_V2_PREFIX) + .and_then(|value| value.strip_prefix("aether-runtime-secret-v1:")) + .expect("system secret should contain a nested runtime envelope"); + let (ldap_with_stripped_system, repository) = + state_with_ldap_bind_password(stripped_system); + let config = repository + .get_ldap_config() + .await + .expect("LDAP config should read") + .expect("LDAP config should exist"); + let error = decrypt_or_migrate_ldap_bind_password(&ldap_with_stripped_system, &config) + .await + .expect_err("stripped system framing must not become an LDAP password"); + assert!(error.into_message().contains("reserved secret framing")); + + let stripped_ldap = ldap_sealed + .strip_prefix(LDAP_BIND_PASSWORD_V3_PREFIX) + .and_then(|value| value.strip_prefix("aether-runtime-secret-v1:")) + .expect("LDAP secret should contain a nested runtime envelope"); + let system_with_stripped_ldap = state_with_stored_secret(stripped_ldap); + let error = decrypt_or_migrate_system_config_secret( + &system_with_stripped_ldap, + TEST_KEY, + stripped_ldap.to_string(), + ) + .await + .expect_err("stripped LDAP framing must not become a system secret"); + assert!(error.into_message().contains("reserved secret framing")); + } + + #[tokio::test] + async fn legacy_plaintext_secret_is_lazily_migrated() { + let state = state_with_stored_secret(TEST_SECRET); + + let plaintext = + decrypt_or_migrate_system_config_secret(&state, TEST_KEY, TEST_SECRET.to_string()) + .await + .expect("legacy secret should migrate"); + assert_eq!(plaintext, TEST_SECRET); + + let stored = state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored secret should be a string"); + assert_ne!(stored, TEST_SECRET); + assert!(stored.starts_with(SYSTEM_CONFIG_SECRET_V2_PREFIX)); + assert_eq!( + decrypt_system_config_secret(&state, TEST_KEY, &stored) + .expect("migrated secret should decrypt"), + TEST_SECRET + ); + } + + #[tokio::test] + async fn legacy_secret_migration_compares_the_untrimmed_stored_value() { + let stored_raw = format!(" {TEST_SECRET} "); + let state = state_with_stored_secret(&stored_raw); + + let plaintext = + decrypt_or_migrate_system_config_secret(&state, TEST_KEY, TEST_SECRET.to_string()) + .await + .expect("whitespace-wrapped legacy secret should migrate"); + assert_eq!(plaintext, TEST_SECRET); + + let stored = state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored secret should be a string"); + assert_eq!( + decrypt_system_config_secret(&state, TEST_KEY, &stored).as_deref(), + Some(TEST_SECRET) + ); + } + + #[tokio::test] + async fn concurrent_legacy_reads_do_not_double_encrypt() { + let state = state_with_stored_secret(TEST_SECRET); + let reads = (0..16).map(|_| { + decrypt_or_migrate_system_config_secret(&state, TEST_KEY, TEST_SECRET.to_string()) + }); + for result in join_all(reads).await { + assert_eq!( + result.expect("concurrent migration should succeed"), + TEST_SECRET + ); + } + + let stored = state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored secret should be a string"); + assert_eq!( + decrypt_system_config_secret(&state, TEST_KEY, &stored) + .expect("migrated secret should decrypt exactly once"), + TEST_SECRET + ); + } + + #[tokio::test] + async fn stale_cached_ciphertext_does_not_bypass_strong_read() { + let stale = encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + "credential-before-rotation", + ) + .expect("stale fixture should encrypt"); + let current = encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + "credential-after-rotation", + ) + .expect("current fixture should encrypt"); + let state = state_with_stored_secret(¤t); + + let plaintext = decrypt_or_migrate_system_config_secret(&state, TEST_KEY, stale) + .await + .expect("strong read should use the rotated credential"); + + assert_eq!(plaintext, "credential-after-rotation"); + let stored = state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored secret should be a string"); + assert_ne!(stored, current); + assert_eq!( + decrypt_system_config_secret(&state, TEST_KEY, &stored).as_deref(), + Some("credential-after-rotation") + ); + } + + #[tokio::test] + async fn administrator_rotation_wins_race_with_plaintext_migration() { + let state = state_with_stored_secret(TEST_SECRET); + let rotated_plaintext = "credential-after-administrator-rotation"; + let rotated_ciphertext = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, rotated_plaintext) + .expect("rotated fixture should encrypt"); + let reached_compare = Arc::new(Barrier::new(2)); + let resume_compare = Arc::new(Barrier::new(2)); + let first_compare = Arc::new(AtomicBool::new(true)); + let migration_reached_compare = Arc::clone(&reached_compare); + let migration_resume_compare = Arc::clone(&resume_compare); + let migration_first_compare = Arc::clone(&first_compare); + + let migration = decrypt_or_migrate_system_config_secret_with_before_compare( + &state, + TEST_KEY, + TEST_SECRET.to_string(), + move || { + let reached_compare = Arc::clone(&migration_reached_compare); + let resume_compare = Arc::clone(&migration_resume_compare); + let first_compare = Arc::clone(&migration_first_compare); + async move { + if first_compare.swap(false, Ordering::SeqCst) { + reached_compare.wait().await; + resume_compare.wait().await; + } + } + }, + ); + let rotation = async { + reached_compare.wait().await; + state + .upsert_system_config_json_value(TEST_KEY, &json!(rotated_ciphertext.clone()), None) + .await + .expect("administrator rotation should persist"); + resume_compare.wait().await; + }; + + let (migration_result, ()) = tokio::join!(migration, rotation); + assert_eq!( + migration_result.expect("migration should retry against the rotated value"), + rotated_plaintext + ); + let stored = state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("stored secret should be a string"); + assert_ne!(stored, rotated_ciphertext); + assert_eq!( + decrypt_system_config_secret(&state, TEST_KEY, &stored).as_deref(), + Some(rotated_plaintext) + ); + } + + #[tokio::test] + async fn failed_compare_and_set_invalidates_cached_system_config_value() { + let data = Arc::new( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests([( + TEST_KEY.to_string(), + json!("value-cached-on-this-node"), + )]), + ); + let mut state = AppState::new().expect("gateway state should build"); + state.replace_data_state(Arc::clone(&data)); + assert_eq!( + state + .read_system_config_json_value(TEST_KEY) + .await + .expect("initial config should cache"), + Some(json!("value-cached-on-this-node")) + ); + + data.upsert_system_config_value(TEST_KEY, &json!("value-rotated-by-another-node"), None) + .await + .expect("simulated remote rotation should persist"); + assert!(!state + .compare_and_set_system_config_string_value( + TEST_KEY, + "value-cached-on-this-node", + "stale-migration-replacement", + ) + .await + .expect("stale compare-and-set should complete")); + + assert_eq!( + state + .read_system_config_json_value(TEST_KEY) + .await + .expect("config should reload after failed compare-and-set"), + Some(json!("value-rotated-by-another-node")) + ); + } + + #[tokio::test] + async fn undecryptable_fernet_secret_fails_without_plaintext_fallback() { + let ciphertext = encrypt_python_fernet_plaintext("unavailable-historical-key", TEST_SECRET) + .expect("fixture should encrypt"); + let state = state_with_stored_secret(&ciphertext); + + let error = decrypt_or_migrate_system_config_secret(&state, TEST_KEY, ciphertext.clone()) + .await + .expect_err("unknown Fernet ciphertext must fail closed"); + let error_text = error.into_message(); + assert!(!error_text.contains(TEST_SECRET)); + assert!(!error_text.contains(&ciphertext)); + assert_eq!( + state + .read_system_config_json_value_strong(TEST_KEY) + .await + .expect("stored secret should read"), + Some(json!(ciphertext)) + ); + } + + #[tokio::test] + async fn system_config_secret_ciphertext_is_bound_to_its_config_key() { + let sealed = encrypt_system_config_secret( + &state_with_stored_secret(TEST_SECRET), + TEST_KEY, + TEST_SECRET, + ) + .expect("system config secret should seal"); + let state = state_with_stored_secret(&sealed); + + assert_eq!( + decrypt_or_migrate_system_config_secret(&state, TEST_KEY, sealed.clone(),) + .await + .expect("matching system config secret should open"), + TEST_SECRET + ); + let wrong_key = "backup_s3_secret_access_key"; + let wrong_state = state_with_named_stored_secret(wrong_key, &sealed); + let wrong_key_error = + decrypt_or_migrate_system_config_secret(&wrong_state, wrong_key, sealed.clone()) + .await + .expect_err("copied system config secret must fail closed"); + assert!(wrong_key_error + .into_message() + .contains("cannot be decrypted")); + assert_eq!( + decrypt_system_config_secret(&state, "SMTP_PASSWORD", &sealed).as_deref(), + Some(TEST_SECRET) + ); + } + + #[test] + fn bark_device_key_ciphertext_is_bound_to_canonical_server_url() { + let state = state_with_stored_secret(TEST_SECRET); + let binding = bark_device_key_binding(" HTTPS://Example.COM:443/api/// ") + .expect("Bark server URL should canonicalize"); + assert_eq!(binding.server_url, "https://example.com/api"); + let equivalent = bark_device_key_binding("https://example.com:443/api") + .expect("equivalent Bark server URL should canonicalize"); + let sealed = encrypt_bark_device_key(&state, &binding, TEST_SECRET) + .expect("Bark device key should seal"); + + assert_eq!( + decrypt_bark_device_key_v2(&state, &equivalent, &sealed).as_deref(), + Some(TEST_SECRET) + ); + let changed = bark_device_key_binding("https://example.com/other") + .expect("changed Bark server URL should parse"); + assert!(decrypt_bark_device_key_v2(&state, &changed, &sealed).is_none()); + } + + #[tokio::test] + async fn legacy_bark_device_key_migrates_to_destination_bound_envelope() { + let state = state_with_named_stored_secret(BARK_DEVICE_KEY, TEST_SECRET); + let binding = + bark_device_key_binding("https://api.day.app").expect("Bark URL should parse"); + let plaintext = + decrypt_or_migrate_bark_device_key(&state, &binding, TEST_SECRET.to_string()) + .await + .expect("legacy Bark device key should migrate"); + assert_eq!(plaintext, TEST_SECRET); + + let stored = state + .read_system_config_json_value_strong(BARK_DEVICE_KEY) + .await + .expect("Bark device key should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("Bark device key should remain a string"); + assert!(stored.starts_with("aether-bark-device-key-v2:")); + assert_eq!( + decrypt_bark_device_key_v2(&state, &binding, &stored).as_deref(), + Some(TEST_SECRET) + ); + } + + #[tokio::test] + async fn bark_device_key_rejects_a_ciphertext_bound_to_another_server_url() { + let state = state_with_stored_secret(TEST_SECRET); + let original = + bark_device_key_binding("https://api.day.app").expect("Bark URL should parse"); + let sealed = encrypt_bark_device_key(&state, &original, TEST_SECRET) + .expect("Bark device key should seal"); + let state = state_with_named_stored_secret(BARK_DEVICE_KEY, &sealed); + let changed = bark_device_key_binding("https://bark.example.test/api") + .expect("changed Bark URL should parse"); + + let error = decrypt_or_migrate_bark_device_key(&state, &changed, sealed.clone()) + .await + .expect_err("Bark device key must not move between destinations"); + assert!(error.into_message().contains("cannot be decrypted")); + assert_eq!( + state + .read_system_config_json_value_strong(BARK_DEVICE_KEY) + .await + .expect("Bark device key should remain readable"), + Some(json!(sealed)) + ); + } +} diff --git a/apps/aether-gateway/src/headers.rs b/apps/aether-gateway/src/headers.rs index 9c4c8042c..36e5b0757 100644 --- a/apps/aether-gateway/src/headers.rs +++ b/apps/aether-gateway/src/headers.rs @@ -17,30 +17,54 @@ const MAX_REQUEST_BODY_MB_ENV: &str = "AETHER_MAX_REQUEST_BODY_MB"; const MAX_REDACTED_SYNC_RESPONSE_BODY_MB_ENV: &str = "AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB"; const MAX_INTERNAL_BUFFERED_BODY_MB_ENV: &str = "AETHER_MAX_INTERNAL_BUFFERED_BODY_MB"; const TRUSTED_PROXY_CIDRS_ENV: &str = "AETHER_TRUSTED_PROXY_CIDRS"; +const DEFAULT_MAX_REQUEST_BODY_BYTES: u64 = 256 * 1024 * 1024; +const DEFAULT_MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES: u64 = 64 * 1024 * 1024; +const DEFAULT_MAX_INTERNAL_BUFFERED_BODY_BYTES: u64 = 64 * 1024 * 1024; +// A finite ceiling remains in force even when an operator uses the historical +// `0`/oversized value to request an effectively unlimited body. This protects +// direct execution-runtime and internal aggregation paths that do not hold a +// frontdoor body-budget permit. It does not apply to streaming bodies. +const MAX_CONFIGURED_BUFFERED_BODY_BYTES: u64 = 256 * 1024 * 1024; +const MAX_REQUEST_CONTENT_ENCODINGS: usize = 8; -/// Optional operator cap applied after Content-Encoding decoding, and to -/// uncompressed bodies as-is. Unset, zero, or invalid values disable the cap. -static MAX_REQUEST_BODY_BYTES: LazyLock = - LazyLock::new(|| body_limit_bytes_from_env(MAX_REQUEST_BODY_MB_ENV)); +/// Operator cap applied after Content-Encoding decoding, and to uncompressed +/// bodies as-is. A configured zero disables the optional lower cap, while the +/// finite safety ceiling still prevents an unbounded allocation. +static MAX_REQUEST_BODY_BYTES: LazyLock = LazyLock::new(|| { + body_limit_bytes_from_env(MAX_REQUEST_BODY_MB_ENV, DEFAULT_MAX_REQUEST_BODY_BYTES) +}); -static MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES: LazyLock = - LazyLock::new(|| body_limit_bytes_from_env(MAX_REDACTED_SYNC_RESPONSE_BODY_MB_ENV)); +static MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES: LazyLock = LazyLock::new(|| { + body_limit_bytes_from_env( + MAX_REDACTED_SYNC_RESPONSE_BODY_MB_ENV, + DEFAULT_MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES, + ) +}); -static MAX_INTERNAL_BUFFERED_BODY_BYTES: LazyLock = - LazyLock::new(|| body_limit_bytes_from_env(MAX_INTERNAL_BUFFERED_BODY_MB_ENV)); +static MAX_INTERNAL_BUFFERED_BODY_BYTES: LazyLock = LazyLock::new(|| { + body_limit_bytes_from_env( + MAX_INTERNAL_BUFFERED_BODY_MB_ENV, + DEFAULT_MAX_INTERNAL_BUFFERED_BODY_BYTES, + ) +}); -fn body_limit_bytes_from_env(name: &str) -> u64 { +fn body_limit_bytes_from_env(name: &str, default_bytes: u64) -> u64 { let value = std::env::var(name).ok(); - body_limit_bytes(value.as_deref()) + body_limit_bytes(value.as_deref(), default_bytes) } -fn body_limit_bytes(value: Option<&str>) -> u64 { - value - .map(str::trim) - .and_then(|value| value.parse::().ok()) - .filter(|value| *value > 0) - .map(|value| value.saturating_mul(1024 * 1024)) - .unwrap_or(u64::MAX) +fn body_limit_bytes(value: Option<&str>, default_bytes: u64) -> u64 { + let configured = match value.map(str::trim).filter(|value| !value.is_empty()) { + Some("0") => MAX_CONFIGURED_BUFFERED_BODY_BYTES, + Some(value) => match value.parse::() { + Ok(value) if value > 0 => value + .checked_mul(1024 * 1024) + .unwrap_or(MAX_CONFIGURED_BUFFERED_BODY_BYTES), + _ => default_bytes, + }, + None => default_bytes, + }; + configured.min(MAX_CONFIGURED_BUFFERED_BODY_BYTES) } static TRUSTED_PROXY_CIDRS: LazyLock> = LazyLock::new(|| { @@ -62,7 +86,9 @@ pub(crate) fn max_redacted_sync_response_body_bytes() -> u64 { } pub(crate) fn max_internal_buffered_body_bytes() -> usize { - usize::try_from(*MAX_INTERNAL_BUFFERED_BODY_BYTES).unwrap_or(usize::MAX) + usize::try_from(*MAX_INTERNAL_BUFFERED_BODY_BYTES) + .unwrap_or(usize::MAX) + .min(usize::try_from(MAX_CONFIGURED_BUFFERED_BODY_BYTES).unwrap_or(usize::MAX)) } pub(crate) fn extract_or_generate_trace_id(headers: &http::HeaderMap) -> String { @@ -86,13 +112,24 @@ pub(crate) fn header_value_u64(headers: &http::HeaderMap, key: &str) -> Option, pub(crate) user_agent: Option, + pub(crate) forwarded_headers_trusted: bool, } pub(crate) fn request_origin_from_headers(headers: &http::HeaderMap) -> RequestOrigin { + RequestOrigin { + client_ip: None, + user_agent: header_value_str(headers, http::header::USER_AGENT.as_str()) + .map(|value| truncate_chars(value.as_str(), 1_000)), + forwarded_headers_trusted: false, + } +} + +pub(crate) fn request_origin_from_trusted_headers(headers: &http::HeaderMap) -> RequestOrigin { RequestOrigin { client_ip: client_ip_from_headers(headers), user_agent: header_value_str(headers, http::header::USER_AGENT.as_str()) .map(|value| truncate_chars(value.as_str(), 1_000)), + forwarded_headers_trusted: true, } } @@ -104,6 +141,7 @@ pub(crate) fn request_origin_from_headers_and_remote_addr( client_ip: Some(effective_client_ip(headers, remote_addr).to_string()), user_agent: header_value_str(headers, http::header::USER_AGENT.as_str()) .map(|value| truncate_chars(value.as_str(), 1_000)), + forwarded_headers_trusted: trusted_proxy_ip(remote_addr.ip()), } } @@ -113,30 +151,35 @@ pub(crate) fn effective_client_ip(headers: &http::HeaderMap, remote_addr: &Socke return remote_ip; } - if let Some(real_ip) = - header_value_str(headers, "x-real-ip").and_then(|value| value.parse::().ok()) - { - return real_ip; - } - - let forwarded_ips = header_value_str(headers, "x-forwarded-for") - .map(|value| { - value - .split(',') - .filter_map(|segment| segment.trim().parse::().ok()) - .collect::>() - }) - .unwrap_or_default(); - forwarded_ips + let forwarded_ips = headers + .get_all("x-forwarded-for") + .iter() + .filter_map(|value| value.to_str().ok()) + .flat_map(|value| value.split(',')) + .filter_map(|segment| segment.trim().parse::().ok()) + .collect::>(); + if let Some(client_ip) = forwarded_ips .iter() .rev() .copied() .find(|ip| !trusted_proxy_ip(*ip)) - .or_else(|| forwarded_ips.first().copied()) - .unwrap_or(remote_ip) + { + return client_ip; + } + + let mut real_ip_values = headers.get_all("x-real-ip").iter(); + let real_ip = real_ip_values + .next() + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.trim().parse::().ok()); + if real_ip_values.next().is_none() { + return real_ip.unwrap_or(remote_ip); + } + + remote_ip } -fn trusted_proxy_ip(ip: IpAddr) -> bool { +pub(crate) fn trusted_proxy_ip(ip: IpAddr) -> bool { TRUSTED_PROXY_CIDRS .iter() .any(|pattern| ip_or_cidr_matches(pattern, ip)) @@ -269,30 +312,81 @@ pub(crate) fn should_skip_upstream_passthrough_header(name: &str) -> bool { } pub(crate) fn should_skip_response_header(name: &str) -> bool { - matches!( - name.to_ascii_lowercase().as_str(), - "connection" - | "keep-alive" - | "proxy-authenticate" - | "proxy-authorization" - | "proxy-connection" - | "te" - | "trailer" - | "transfer-encoding" - | "upgrade" - | "x-aether-control-executed" - | "x-aether-control-action" - ) + let name = name.to_ascii_lowercase(); + name == "set-cookie" + || name.starts_with("x-aether-") + // CORS is a gateway policy. If an upstream can supply these fields, + // it can opt an otherwise-disallowed browser origin into reading a + // credentialed gateway response after the CORS middleware declines + // to add its own headers. + || name.starts_with("access-control-") + // Browser security policy is owned by the gateway. Besides weakening + // active-content protections, an upstream-controlled report-only + // policy / Reporting API endpoint can make a browser disclose gateway + // URLs and diagnostics to an attacker-controlled collector. + || name.starts_with("content-security-policy") + // These response headers are interpreted by common reverse proxies + // and application servers as privileged internal redirects or local + // file-send instructions. Upstream providers are untrusted at this + // boundary and must not be able to make the gateway's front proxy + // fetch an internal URL or disclose a local file. + || name.starts_with("x-accel-") + || matches!( + name.as_str(), + "accept-ch" + | "alt-svc" + | "authentication-info" + | "connection" + | "clear-site-data" + | "content-length" + | "critical-ch" + | "keep-alive" + | "nel" + | "proxy-authenticate" + | "proxy-authentication-info" + | "proxy-authorization" + | "proxy-connection" + | "referrer-policy" + | "refresh" + | "report-to" + | "reporting-endpoints" + | "strict-transport-security" + | "te" + | "timing-allow-origin" + | "trailer" + | "transfer-encoding" + | "upgrade" + | "x-content-type-options" + | "x-httpd-send-file" + | "x-lighttpd-send-file" + | "x-litespeed-location" + | "x-reproxy-url" + | "x-send-file" + | "x-sendfile" + | "x-sendfile2" + ) } pub(crate) fn collect_control_headers(headers: &http::HeaderMap) -> BTreeMap { + let connection_declared = aether_http::connection_declared_header_names( + headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); headers .iter() .filter_map(|(name, value)| { + let normalized = name.as_str().to_ascii_lowercase(); + if normalized == http::header::CONNECTION.as_str() + || connection_declared.contains(&normalized) + { + return None; + } value .to_str() .ok() - .map(|value| (name.as_str().to_ascii_lowercase(), value.trim().to_string())) + .map(|value| (normalized, value.trim().to_string())) }) .collect() } @@ -305,6 +399,8 @@ pub(crate) fn is_json_request(headers: &http::HeaderMap) -> bool { #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) enum RequestBodyNormalizationError { + InvalidBodyFraming, + AmbiguousBodyFraming, UnsupportedContentEncoding(String), DecodeFailed { encoding: String, reason: String }, DecompressedBodyTooLarge { encoding: String, limit_bytes: u64 }, @@ -314,6 +410,9 @@ pub(crate) enum RequestBodyNormalizationError { impl RequestBodyNormalizationError { pub(crate) fn client_message(&self) -> String { match self { + Self::InvalidBodyFraming | Self::AmbiguousBodyFraming => { + "Invalid request body framing".to_string() + } Self::UnsupportedContentEncoding(encoding) => { format!("Unsupported request Content-Encoding: {encoding}") } @@ -334,6 +433,7 @@ impl RequestBodyNormalizationError { pub(crate) fn http_status(&self) -> http::StatusCode { match self { + Self::InvalidBodyFraming | Self::AmbiguousBodyFraming => http::StatusCode::BAD_REQUEST, Self::DecompressedBodyTooLarge { .. } | Self::RequestBodyTooLarge { .. } => { http::StatusCode::PAYLOAD_TOO_LARGE } @@ -347,6 +447,8 @@ impl RequestBodyNormalizationError { impl fmt::Display for RequestBodyNormalizationError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { + Self::InvalidBodyFraming => write!(f, "invalid request body framing"), + Self::AmbiguousBodyFraming => write!(f, "ambiguous request body framing"), Self::UnsupportedContentEncoding(encoding) => { write!(f, "unsupported request Content-Encoding: {encoding}") } @@ -412,8 +514,8 @@ pub(crate) fn check_request_content_length_with_limit( headers: &http::HeaderMap, limit: u64, ) -> Result<(), RequestBodyNormalizationError> { - let declared = header_value_str(headers, http::header::CONTENT_LENGTH.as_str()) - .and_then(|value| value.trim().parse::().ok()); + validate_request_body_framing(headers)?; + let declared = declared_request_content_length(headers)?; if declared.is_some_and(|value| value > limit) { return Err(RequestBodyNormalizationError::RequestBodyTooLarge { limit_bytes: limit }); } @@ -432,6 +534,7 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>( body_bytes: &'a [u8], limit: u64, ) -> Result, RequestBodyNormalizationError> { + validate_request_body_framing(headers)?; let encodings = request_content_encodings(headers); if encodings.is_empty() { if body_bytes.len() as u64 > limit { @@ -448,17 +551,73 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>( } fn request_content_encodings(headers: &http::HeaderMap) -> Vec { - header_value_str(headers, http::header::CONTENT_ENCODING.as_str()) - .map(|value| { - value - .split(',') - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_ascii_lowercase) - .filter(|value| value != "identity") - .collect() - }) - .unwrap_or_default() + headers + .get_all(http::header::CONTENT_ENCODING) + .iter() + .flat_map(|value| value.to_str().unwrap_or_default().split(',')) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_ascii_lowercase) + .filter(|value| value != "identity") + .collect() +} + +fn validate_request_body_framing( + headers: &http::HeaderMap, +) -> Result<(), RequestBodyNormalizationError> { + let _ = declared_request_content_length(headers)?; + if headers.contains_key(http::header::CONTENT_LENGTH) + && headers.contains_key(http::header::TRANSFER_ENCODING) + { + return Err(RequestBodyNormalizationError::AmbiguousBodyFraming); + } + let mut encoding_count = 0usize; + if headers + .get_all(http::header::CONTENT_ENCODING) + .iter() + .nth(1) + .is_some() + { + return Err(RequestBodyNormalizationError::AmbiguousBodyFraming); + } + for value in headers.get_all(http::header::CONTENT_ENCODING).iter() { + let value = value + .to_str() + .map_err(|_| RequestBodyNormalizationError::InvalidBodyFraming)?; + for encoding in value.split(',') { + if encoding.trim().is_empty() { + return Err(RequestBodyNormalizationError::InvalidBodyFraming); + } + encoding_count = encoding_count.saturating_add(1); + if encoding_count > MAX_REQUEST_CONTENT_ENCODINGS { + return Err(RequestBodyNormalizationError::InvalidBodyFraming); + } + } + } + Ok(()) +} + +fn declared_request_content_length( + headers: &http::HeaderMap, +) -> Result, RequestBodyNormalizationError> { + let mut values = headers.get_all(http::header::CONTENT_LENGTH).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err(RequestBodyNormalizationError::AmbiguousBodyFraming); + } + let value = value + .to_str() + .map_err(|_| RequestBodyNormalizationError::InvalidBodyFraming)? + .trim(); + if value.is_empty() || value.contains(',') { + return Err(RequestBodyNormalizationError::AmbiguousBodyFraming); + } + value + .parse::() + .map(Some) + .map_err(|_| RequestBodyNormalizationError::InvalidBodyFraming) } fn decode_single_request_body( @@ -594,8 +753,59 @@ mod tests { use super::{ decoded_request_body_bytes, effective_client_ip, normalize_request_body_headers_and_bytes, request_origin_from_headers, request_origin_from_headers_and_remote_addr, + request_origin_from_trusted_headers, should_skip_response_header, tls_fingerprint_from_headers, RequestBodyNormalizationError, RequestOrigin, }; + + #[test] + fn upstream_response_header_filter_blocks_privileged_server_control_headers() { + for name in [ + "set-cookie", + "Set-Cookie", + "x-aether-control-executed", + "X-Aether-Future-Control", + "Access-Control-Allow-Origin", + "Access-Control-Allow-Credentials", + "Access-Control-Expose-Headers", + "Accept-CH", + "Alt-Svc", + "Authentication-Info", + "Content-Security-Policy", + "Content-Security-Policy-Report-Only", + "Clear-Site-Data", + "Content-Length", + "Critical-CH", + "NEL", + "Proxy-Authentication-Info", + "Referrer-Policy", + "Refresh", + "Report-To", + "Reporting-Endpoints", + "Strict-Transport-Security", + "Timing-Allow-Origin", + "X-Content-Type-Options", + "X-Accel-Redirect", + "x-accel-expires", + "X-Sendfile", + "X-Sendfile2", + "X-Send-File", + "X-HTTPD-Send-File", + "X-LIGHTTPD-send-file", + "X-LiteSpeed-Location", + "X-Reproxy-URL", + ] { + assert!( + should_skip_response_header(name), + "upstream response header should be blocked: {name}" + ); + } + assert!(!should_skip_response_header("content-type")); + assert!(!should_skip_response_header("x-proxy-timing")); + // WWW-Authenticate is an end-to-end challenge used by legitimate + // provider APIs (for example Bearer realm/error challenges). It is not + // a proxy control header and must remain available to SDK clients. + assert!(!should_skip_response_header("WWW-Authenticate")); + } use flate2::{ write::{DeflateEncoder, GzEncoder, ZlibEncoder}, Compression, @@ -608,7 +818,7 @@ mod tests { }; #[test] - fn request_origin_prefers_first_forwarded_for_ip() { + fn trusted_request_origin_prefers_first_forwarded_for_ip() { let mut headers = HeaderMap::new(); headers.insert( "x-forwarded-for", @@ -621,12 +831,14 @@ mod tests { ); assert_eq!( - request_origin_from_headers(&headers), + request_origin_from_trusted_headers(&headers), RequestOrigin { client_ip: Some("203.0.113.8".to_string()), user_agent: Some("Claude-Code/1.0".to_string()), + forwarded_headers_trusted: true, } ); + assert_eq!(request_origin_from_headers(&headers).client_ip, None); } #[test] @@ -652,6 +864,21 @@ mod tests { effective_client_ip(&headers, &remote_addr), IpAddr::V4(Ipv4Addr::new(198, 51, 100, 4)) ); + assert!( + request_origin_from_headers_and_remote_addr(&headers, &remote_addr) + .forwarded_headers_trusted + ); + } + + #[test] + fn request_origin_does_not_trust_forwarded_metadata_from_public_peer() { + let headers = HeaderMap::new(); + let remote_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 10)), 443); + + assert!( + !request_origin_from_headers_and_remote_addr(&headers, &remote_addr) + .forwarded_headers_trusted + ); } #[test] @@ -669,6 +896,51 @@ mod tests { ); } + #[test] + fn effective_client_ip_does_not_trust_all_trusted_forwarded_chain() { + let mut headers = HeaderMap::new(); + headers.insert( + "x-forwarded-for", + HeaderValue::from_static("127.0.0.2, 127.0.0.3"), + ); + let remote_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + + assert_eq!( + effective_client_ip(&headers, &remote_addr), + remote_addr.ip(), + "a chain containing only trusted proxy addresses cannot establish the client IP" + ); + } + + #[test] + fn effective_client_ip_prefers_forwarded_chain_over_conflicting_real_ip() { + let mut headers = HeaderMap::new(); + headers.insert( + "x-forwarded-for", + HeaderValue::from_static("203.0.113.8, 127.0.0.2"), + ); + headers.insert("x-real-ip", HeaderValue::from_static("198.51.100.4")); + let remote_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + + assert_eq!( + effective_client_ip(&headers, &remote_addr), + IpAddr::V4(Ipv4Addr::new(203, 0, 113, 8)) + ); + } + + #[test] + fn effective_client_ip_rejects_ambiguous_real_ip_headers() { + let mut headers = HeaderMap::new(); + headers.append("x-real-ip", HeaderValue::from_static("198.51.100.4")); + headers.append("x-real-ip", HeaderValue::from_static("203.0.113.8")); + let remote_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + + assert_eq!( + effective_client_ip(&headers, &remote_addr), + remote_addr.ip() + ); + } + #[test] fn decoded_request_body_bytes_decodes_zstd() { let payload = br#"{"model":"gpt-5.4"}"#; @@ -862,31 +1134,90 @@ mod tests { } #[test] - fn body_limits_default_to_unlimited() { - assert_eq!(super::body_limit_bytes(None), u64::MAX); - assert_eq!(super::body_limit_bytes(Some("0")), u64::MAX); - assert_eq!(super::body_limit_bytes(Some("invalid")), u64::MAX); + fn request_body_framing_rejects_duplicate_content_length() { + let mut headers = HeaderMap::new(); + headers.append(http::header::CONTENT_LENGTH, HeaderValue::from_static("4")); + headers.append(http::header::CONTENT_LENGTH, HeaderValue::from_static("4")); + + let err = super::check_request_content_length_with_limit(&headers, 10) + .expect_err("duplicate Content-Length fields must be rejected"); + assert_eq!(err, RequestBodyNormalizationError::AmbiguousBodyFraming); + } + + #[test] + fn request_body_framing_rejects_content_length_and_transfer_encoding() { + let mut headers = HeaderMap::new(); + headers.insert(http::header::CONTENT_LENGTH, HeaderValue::from_static("4")); + headers.insert( + http::header::TRANSFER_ENCODING, + HeaderValue::from_static("chunked"), + ); + + let err = super::decoded_request_body_bytes_with_limit(&headers, b"body", 10) + .expect_err("Content-Length and Transfer-Encoding must not be combined"); + assert_eq!(err, RequestBodyNormalizationError::AmbiguousBodyFraming); + } + + #[test] + fn request_body_framing_rejects_duplicate_content_encoding_fields() { + let mut headers = HeaderMap::new(); + headers.append( + http::header::CONTENT_ENCODING, + HeaderValue::from_static("gzip"), + ); + headers.append( + http::header::CONTENT_ENCODING, + HeaderValue::from_static("identity"), + ); + + let err = super::decoded_request_body_bytes_with_limit(&headers, b"body", 10) + .expect_err("duplicate Content-Encoding fields must be rejected"); + assert_eq!(err, RequestBodyNormalizationError::AmbiguousBodyFraming); + } + + #[test] + fn body_limits_use_safe_default_and_finite_unlimited_override() { + let default = 64 * 1024 * 1024; + assert_eq!(super::body_limit_bytes(None, default), default); + assert_eq!(super::body_limit_bytes(Some("invalid"), default), default); + assert_eq!( + super::body_limit_bytes(Some("0"), default), + super::MAX_CONFIGURED_BUFFERED_BODY_BYTES + ); + assert_eq!( + super::body_limit_bytes(Some("999999999999"), default), + super::MAX_CONFIGURED_BUFFERED_BODY_BYTES + ); } #[test] fn positive_body_limit_is_converted_from_mibibytes() { - assert_eq!(super::body_limit_bytes(Some(" 8 ")), 8 * 1024 * 1024); + assert_eq!( + super::body_limit_bytes(Some(" 8 "), 64 * 1024 * 1024), + 8 * 1024 * 1024 + ); } #[test] - fn unlimited_limit_accepts_declared_and_buffered_body() { + fn finite_unlimited_limit_rejects_values_above_safety_ceiling() { let mut headers = HeaderMap::new(); headers.insert( http::header::CONTENT_LENGTH, - HeaderValue::from_static("18446744073709551615"), + HeaderValue::from_static("268435457"), ); - super::check_request_content_length_with_limit(&headers, u64::MAX) - .expect("unlimited mode should accept every representable content length"); + super::check_request_content_length_with_limit( + &headers, + super::MAX_CONFIGURED_BUFFERED_BODY_BYTES, + ) + .expect_err("body above the finite safety ceiling should be rejected"); - let body = b"body larger than the former default is admitted by the unlimited sentinel"; - let decoded = - super::decoded_request_body_bytes_with_limit(&HeaderMap::new(), body, u64::MAX) - .expect("unlimited mode should accept buffered bytes"); + let body = b"body remains bounded by the finite safety ceiling"; + let decoded = super::decoded_request_body_bytes_with_limit( + &HeaderMap::new(), + body, + super::MAX_CONFIGURED_BUFFERED_BODY_BYTES, + ) + .expect("body within the finite safety ceiling should pass"); assert_eq!(decoded.as_ref(), body); } @@ -933,13 +1264,15 @@ mod tests { } #[test] - fn request_origin_uses_real_ip_after_empty_forwarded_for_segments() { + fn trusted_request_origin_uses_real_ip_after_empty_forwarded_for_segments() { let mut headers = HeaderMap::new(); headers.insert("x-forwarded-for", HeaderValue::from_static(" , unknown ")); headers.insert("x-real-ip", HeaderValue::from_static("198.51.100.4")); assert_eq!( - request_origin_from_headers(&headers).client_ip.as_deref(), + request_origin_from_trusted_headers(&headers) + .client_ip + .as_deref(), Some("198.51.100.4") ); } diff --git a/apps/aether-gateway/src/important_notification.rs b/apps/aether-gateway/src/important_notification.rs index 1f5595409..70f2fcd79 100644 --- a/apps/aether-gateway/src/important_notification.rs +++ b/apps/aether-gateway/src/important_notification.rs @@ -24,6 +24,18 @@ pub(crate) const IMPORTANT_NOTIFICATION_DEFAULT_CHANNEL_KEY: &str = pub(crate) const IMPORTANT_NOTIFICATION_ITEMS_KEY: &str = "module.important_notification.items"; pub(crate) const PROVIDER_QUOTA_ALERT_ITEM_KEY: &str = "provider_quota_alert"; +// Notification configuration is administrator-controlled but is also read on +// request paths. Bound fan-out and template materialization so a malformed or +// compromised configuration cannot trigger unbounded work or SMTP payloads. +const MAX_NOTIFICATION_RECIPIENTS: usize = 128; +const MAX_NOTIFICATION_RECIPIENT_BYTES: usize = 320; +const MAX_NOTIFICATION_RECIPIENT_VALUE_BYTES: usize = 64 * 1024; +const MAX_NOTIFICATION_RECIPIENT_PARTS_PER_VALUE: usize = 256; +const MAX_NOTIFICATION_ITEMS: usize = 128; +const MAX_NOTIFICATION_ITEM_KEY_BYTES: usize = 128; +const MAX_NOTIFICATION_ITEM_NAME_BYTES: usize = 512; +const MAX_NOTIFICATION_TEMPLATE_BYTES: usize = 256 * 1024; + #[derive(Debug, Clone)] pub(crate) struct ImportantNotification { pub(crate) title: String, @@ -430,6 +442,7 @@ fn parse_notification_items(value: Option<&Value>) -> Vec>() } @@ -437,7 +450,7 @@ fn parse_notification_items(value: Option<&Value>) -> Vec Option { let item = value.as_object()?; let key = item.get("key")?.as_str()?.trim(); - if key.is_empty() { + if key.is_empty() || key.len() > MAX_NOTIFICATION_ITEM_KEY_BYTES { return None; } let name = item @@ -445,6 +458,7 @@ fn parse_notification_item(value: &Value) -> Option Option) -> Option { .map(ToOwned::to_owned) } +fn bounded_optional_string(value: Option<&Value>) -> Option { + optional_non_empty_string(value).filter(|value| value.len() <= MAX_NOTIFICATION_TEMPLATE_BYTES) +} + fn find_notification_item<'a>( config: &'a ImportantNotificationConfig, item_key: &str, @@ -765,7 +783,10 @@ fn parse_recipient_list(value: Option<&Value>) -> Vec { let mut recipients = Vec::new(); match value { Some(Value::Array(items)) => { - for item in items { + for item in items.iter().take(MAX_NOTIFICATION_RECIPIENTS * 4) { + if recipients.len() >= MAX_NOTIFICATION_RECIPIENTS { + break; + } if let Some(raw) = item.as_str() { push_recipient_parts(&mut recipients, raw); } @@ -780,11 +801,21 @@ fn parse_recipient_list(value: Option<&Value>) -> Vec { } fn push_recipient_parts(recipients: &mut Vec, raw: &str) { + if raw.len() > MAX_NOTIFICATION_RECIPIENT_VALUE_BYTES { + return; + } for item in raw .split([',', ';', '\n', '\r']) .map(str::trim) .filter(|value| !value.is_empty()) + .take(MAX_NOTIFICATION_RECIPIENT_PARTS_PER_VALUE) { + if recipients.len() >= MAX_NOTIFICATION_RECIPIENTS { + return; + } + if item.len() > MAX_NOTIFICATION_RECIPIENT_BYTES { + continue; + } recipients.push(item.to_string()); } } @@ -811,6 +842,8 @@ mod tests { use super::{ apply_notification_item_template, parse_channel_filter, parse_notification_items, parse_recipient_list, ImportantNotification, ImportantNotificationChannelFilter, + MAX_NOTIFICATION_ITEMS, MAX_NOTIFICATION_RECIPIENTS, MAX_NOTIFICATION_RECIPIENT_BYTES, + MAX_NOTIFICATION_TEMPLATE_BYTES, }; use serde_json::json; @@ -828,6 +861,18 @@ mod tests { ); } + #[test] + fn parse_recipient_list_bounds_fanout_and_drops_oversized_entries() { + let oversized = "x".repeat(MAX_NOTIFICATION_RECIPIENT_BYTES + 1); + let many = (0..(MAX_NOTIFICATION_RECIPIENTS + 20)) + .map(|index| format!("user{index}@example.com")) + .collect::>() + .join(","); + let recipients = parse_recipient_list(Some(&json!([oversized.clone(), many]))); + assert!(recipients.len() <= MAX_NOTIFICATION_RECIPIENTS); + assert!(!recipients.iter().any(|value| value == &oversized)); + } + #[test] fn parse_notification_items_reads_channel_and_user_email_flag() { let items = parse_notification_items(Some(&json!([ @@ -851,6 +896,22 @@ mod tests { assert!(items[0].user_email_enabled); } + #[test] + fn parse_notification_items_bounds_templates_and_count() { + let oversized = "x".repeat(MAX_NOTIFICATION_TEMPLATE_BYTES + 1); + let raw = (0..(MAX_NOTIFICATION_ITEMS + 10)) + .map(|index| { + json!({ + "key": format!("item_{index}"), + "title_template": oversized.clone(), + }) + }) + .collect::>(); + let items = parse_notification_items(Some(&json!(raw))); + assert!(items.len() <= MAX_NOTIFICATION_ITEMS); + assert!(items.iter().all(|item| item.title_template.is_none())); + } + #[test] fn parse_channel_filter_accepts_bark() { assert_eq!( diff --git a/apps/aether-gateway/src/internal_gateway_auth.rs b/apps/aether-gateway/src/internal_gateway_auth.rs new file mode 100644 index 000000000..ea0e5b10a --- /dev/null +++ b/apps/aether-gateway/src/internal_gateway_auth.rs @@ -0,0 +1,234 @@ +use std::fmt; +use std::sync::Arc; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use sha2::{Digest as _, Sha256}; + +use crate::AppState; + +use aether_contracts::internal_gateway::{ + verify_internal_gateway_request_signature, INTERNAL_GATEWAY_AUTH_NONCE_HEADER, + INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER, INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER, +}; + +pub(crate) const INTERNAL_GATEWAY_AUTH_SECRET_ENV: &str = "AETHER_INTERNAL_GATEWAY_AUTH_SECRET"; +const INTERNAL_GATEWAY_AUTH_SECRET_MIN_BYTES: usize = 32; +const INTERNAL_GATEWAY_AUTH_SECRET_MAX_BYTES: usize = 4096; +const INTERNAL_GATEWAY_AUTH_CLOCK_SKEW_SECS: u64 = 300; +const INTERNAL_GATEWAY_AUTH_NONCE_MIN_BYTES: usize = 16; +const INTERNAL_GATEWAY_AUTH_NONCE_MAX_BYTES: usize = 128; +const INTERNAL_GATEWAY_AUTH_NONCE_TTL: Duration = + Duration::from_secs(INTERNAL_GATEWAY_AUTH_CLOCK_SKEW_SECS * 2 + 30); +const INTERNAL_GATEWAY_AUTH_NONCE_KEY_PREFIX: &str = "internal:gateway:auth:nonce:"; + +#[derive(Clone)] +pub(crate) struct InternalGatewayAuthConfig { + mode: InternalGatewayAuthMode, +} + +#[derive(Clone)] +enum InternalGatewayAuthMode { + Disabled, + Misconfigured, + Hmac { + secret: Arc<[u8]>, + }, + #[cfg(test)] + LoopbackTestCompatibility, +} + +impl fmt::Debug for InternalGatewayAuthConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("InternalGatewayAuthConfig") + .field("status", &self.status()) + .finish() + } +} + +impl InternalGatewayAuthConfig { + pub(crate) fn for_process() -> Self { + #[cfg(test)] + { + Self { + mode: InternalGatewayAuthMode::LoopbackTestCompatibility, + } + } + #[cfg(not(test))] + { + Self::from_environment() + } + } + + fn from_environment() -> Self { + match std::env::var(INTERNAL_GATEWAY_AUTH_SECRET_ENV) { + Ok(value) => Self::from_secret_value(Some(value.as_str())), + Err(std::env::VarError::NotPresent) => Self::from_secret_value(None), + Err(std::env::VarError::NotUnicode(_)) => Self { + mode: InternalGatewayAuthMode::Misconfigured, + }, + } + } + + fn from_secret_value(value: Option<&str>) -> Self { + let Some(value) = value else { + return Self { + mode: InternalGatewayAuthMode::Disabled, + }; + }; + let secret = value.trim().as_bytes(); + if !(INTERNAL_GATEWAY_AUTH_SECRET_MIN_BYTES..=INTERNAL_GATEWAY_AUTH_SECRET_MAX_BYTES) + .contains(&secret.len()) + { + return Self { + mode: InternalGatewayAuthMode::Misconfigured, + }; + } + Self { + mode: InternalGatewayAuthMode::Hmac { + secret: Arc::from(secret), + }, + } + } + + pub(crate) fn status(&self) -> &'static str { + match &self.mode { + InternalGatewayAuthMode::Disabled => "disabled", + InternalGatewayAuthMode::Misconfigured => "misconfigured", + InternalGatewayAuthMode::Hmac { .. } => "hmac_authenticated", + #[cfg(test)] + InternalGatewayAuthMode::LoopbackTestCompatibility => "test_loopback_compatibility", + } + } + + #[cfg(test)] + pub(crate) fn disabled_for_tests() -> Self { + Self { + mode: InternalGatewayAuthMode::Disabled, + } + } + + #[cfg(test)] + pub(crate) fn with_secret_for_tests(secret: &str) -> Self { + let config = Self::from_secret_value(Some(secret)); + assert_eq!(config.status(), "hmac_authenticated"); + config + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum InternalGatewayAuthError { + Disabled, + Invalid, + Unavailable, +} + +pub(crate) async fn authenticate_internal_gateway_request( + state: &AppState, + remote_addr: &std::net::SocketAddr, + method: &http::Method, + path_and_query: &str, + headers: &http::HeaderMap, + body: &[u8], +) -> Result<(), InternalGatewayAuthError> { + let secret = match &state.internal_gateway_auth.mode { + InternalGatewayAuthMode::Disabled => return Err(InternalGatewayAuthError::Disabled), + InternalGatewayAuthMode::Misconfigured => { + return Err(InternalGatewayAuthError::Unavailable) + } + InternalGatewayAuthMode::Hmac { secret } => Arc::clone(secret), + #[cfg(test)] + InternalGatewayAuthMode::LoopbackTestCompatibility => { + return remote_addr + .ip() + .is_loopback() + .then_some(()) + .ok_or(InternalGatewayAuthError::Invalid) + } + }; + + let timestamp = unique_header(headers, INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER) + .and_then(|value| value.parse::().ok()) + .ok_or(InternalGatewayAuthError::Invalid)?; + let nonce = unique_header(headers, INTERNAL_GATEWAY_AUTH_NONCE_HEADER) + .filter(|value| valid_nonce(value)) + .ok_or(InternalGatewayAuthError::Invalid)?; + let signature = unique_header(headers, INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER) + .ok_or(InternalGatewayAuthError::Invalid)?; + let now = current_unix_secs().ok_or(InternalGatewayAuthError::Unavailable)?; + if now.abs_diff(timestamp) > INTERNAL_GATEWAY_AUTH_CLOCK_SKEW_SECS { + return Err(InternalGatewayAuthError::Invalid); + } + if !verify_internal_gateway_request_signature( + secret.as_ref(), + method.as_str(), + path_and_query, + timestamp, + nonce, + body, + signature, + ) { + return Err(InternalGatewayAuthError::Invalid); + } + + let nonce_digest = Sha256::digest(nonce.as_bytes()); + let nonce_key = format!("{INTERNAL_GATEWAY_AUTH_NONCE_KEY_PREFIX}{nonce_digest:x}"); + match state + .runtime_state + .kv_set_if_absent( + &nonce_key, + timestamp.to_string(), + INTERNAL_GATEWAY_AUTH_NONCE_TTL, + ) + .await + { + Ok(true) => Ok(()), + Ok(false) => Err(InternalGatewayAuthError::Invalid), + Err(_) => Err(InternalGatewayAuthError::Unavailable), + } +} + +fn unique_header<'a>(headers: &'a http::HeaderMap, name: &'static str) -> Option<&'a str> { + let mut values = headers.get_all(name).iter(); + let value = values.next()?.to_str().ok()?.trim(); + if value.is_empty() || values.next().is_some() { + return None; + } + Some(value) +} + +fn valid_nonce(nonce: &str) -> bool { + (INTERNAL_GATEWAY_AUTH_NONCE_MIN_BYTES..=INTERNAL_GATEWAY_AUTH_NONCE_MAX_BYTES) + .contains(&nonce.len()) + && nonce + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) +} + +fn current_unix_secs() -> Option { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|v| v.as_secs()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn secret_configuration_fails_closed_and_debug_never_contains_secret() { + assert_eq!( + InternalGatewayAuthConfig::from_secret_value(None).status(), + "disabled" + ); + assert_eq!( + InternalGatewayAuthConfig::from_secret_value(Some("short")).status(), + "misconfigured" + ); + let secret = "internal-gateway-test-secret-32-bytes-minimum"; + let config = InternalGatewayAuthConfig::from_secret_value(Some(secret)); + assert_eq!(config.status(), "hmac_authenticated"); + assert!(!format!("{config:?}").contains(secret)); + } +} diff --git a/apps/aether-gateway/src/lib.rs b/apps/aether-gateway/src/lib.rs index b22c9387f..630717169 100644 --- a/apps/aether-gateway/src/lib.rs +++ b/apps/aether-gateway/src/lib.rs @@ -52,12 +52,16 @@ mod headers; mod hooks; mod image_capabilities; mod important_notification; +mod internal_gateway_auth; +mod local_auth_token; mod log_ids; mod maintenance; +mod management_token_auth; pub(crate) mod middleware; mod model_fetch; mod oauth; mod orchestration; +mod plan_usage_policy; mod privacy; mod process_metrics; mod provider_key_auth; @@ -95,6 +99,11 @@ pub(crate) use self::ai_serving::{ AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt, }; pub use self::async_task::VideoTaskTruthSourceMode; +pub use self::backup::{ + apply_restored_backup, restore_backup_json, BackupApplyError, BackupDecryptionKey, + BackupRestoreError, BackupRestoreLimits, BackupRestoreScope, RestoredBackupJson, + DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES, DEFAULT_BACKUP_MAX_JSON_BYTES, +}; pub use self::data::GatewayDataConfig; pub(crate) use self::error::GatewayError; pub(crate) use self::execution_runtime::{ @@ -125,6 +134,75 @@ pub use self::tunnel::{ build_tunnel_runtime_router_with_state, tunnel_protocol, TunnelConnConfig, TunnelControlPlaneClient, TunnelRuntimeState, }; +#[cfg(feature = "testkit")] +pub fn configure_test_tunnel_security( + state: &mut AppState, + node_id: &str, + tunnel_generation: &str, + key: &str, +) { + use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; + use std::sync::Arc; + + let node = StoredProxyNode::new( + node_id.to_string(), + "test tunnel node".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "offline".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + true, + false, + 0, + ) + .expect("test tunnel node should be valid") + .with_runtime_fields( + None, + None, + None, + None, + Some(serde_json::json!({ + "tunnel_security": {"encryption_key": key} + })), + None, + None, + None, + None, + None, + None, + ) + .with_tunnel_generation(tunnel_generation.to_string()); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let data = self::data::GatewayDataState::with_proxy_node_repository_for_testkit( + repository, + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, + ); + state.replace_data_state(Arc::new(data)); +} +#[cfg(feature = "testkit")] +pub fn configure_test_tunnel_runtime_auth( + state: TunnelRuntimeState, + node_id: &str, + tunnel_generation: &str, + raw_management_token: &str, +) -> Result { + use std::sync::Arc; + + let data = self::data::GatewayDataState::with_tunnel_management_auth_for_testkit( + node_id, + tunnel_generation, + raw_management_token, + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, + ) + .map_err(|error| format!("failed to configure tunnel harness authentication: {error}"))?; + Ok(state.with_data(Arc::new(data))) +} pub use self::usage::UsageRuntimeConfig; use axum::http::header::{HeaderName, HeaderValue}; diff --git a/apps/aether-gateway/src/local_auth_token.rs b/apps/aether-gateway/src/local_auth_token.rs new file mode 100644 index 000000000..bf513ddc9 --- /dev/null +++ b/apps/aether-gateway/src/local_auth_token.rs @@ -0,0 +1,457 @@ +use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; +use hmac::Mac; +use serde_json::{Map, Value}; + +const TEST_JWT_SECRET: &str = "aether-rust-test-jwt-secret-32-bytes-minimum"; +const INSECURE_JWT_SECRETS: &[&str] = &[ + "change-this-to-a-secure-random-string", + "aether-rust-dev-jwt-secret", + TEST_JWT_SECRET, +]; +const INVALID_TOKEN: &str = "无效的Token"; +const MAX_LOCAL_AUTH_TOKEN_BYTES: usize = 128 * 1024; +const MAX_LOCAL_AUTH_JWT_PART_BYTES: usize = 64 * 1024; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum LocalAuthTokenType { + Access, + Refresh, +} + +impl LocalAuthTokenType { + fn as_str(self) -> &'static str { + match self { + Self::Access => "access", + Self::Refresh => "refresh", + } + } +} + +fn validate_jwt_secret_value(value: Option<&str>) -> Result { + let Some(value) = value else { + return Err("JWT_SECRET_KEY 未配置".to_string()); + }; + let value = value.trim(); + if value.as_bytes().len() < 32 || INSECURE_JWT_SECRETS.contains(&value) { + return Err("JWT_SECRET_KEY 必须是至少32字节的非默认随机密钥".to_string()); + } + Ok(value.to_string()) +} + +pub(crate) fn local_auth_jwt_secret() -> Result { + match std::env::var("JWT_SECRET_KEY") { + Ok(value) => validate_jwt_secret_value(Some(&value)), + Err(std::env::VarError::NotPresent) => { + #[cfg(test)] + { + return Ok(TEST_JWT_SECRET.to_string()); + } + + #[cfg(not(test))] + Err("JWT_SECRET_KEY 未配置".to_string()) + } + Err(std::env::VarError::NotUnicode(_)) => { + Err("JWT_SECRET_KEY 必须是有效的UTF-8字符串".to_string()) + } + } +} + +fn base64url_encode(bytes: &[u8]) -> String { + URL_SAFE_NO_PAD.encode(bytes) +} + +fn base64url_decode(value: &str) -> Result, String> { + let max_encoded_len = MAX_LOCAL_AUTH_JWT_PART_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4); + if value.len() > max_encoded_len { + return Err(INVALID_TOKEN.to_string()); + } + let decoded = URL_SAFE_NO_PAD + .decode(value) + .map_err(|_| INVALID_TOKEN.to_string())?; + (decoded.len() <= MAX_LOCAL_AUTH_JWT_PART_BYTES) + .then_some(decoded) + .ok_or_else(|| INVALID_TOKEN.to_string()) +} + +fn non_empty_string_claim<'a>(payload: &'a Map, name: &str) -> Option<&'a str> { + payload + .get(name) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn local_auth_claims_are_valid( + payload: &Map, + token_type: LocalAuthTokenType, +) -> bool { + if non_empty_string_claim(payload, "user_id").is_none() + || non_empty_string_claim(payload, "session_id").is_none() + { + return false; + } + if !matches!( + payload.get("created_at"), + Some(Value::String(value)) if chrono::DateTime::parse_from_rfc3339(value).is_ok() + ) { + return false; + } + match token_type { + LocalAuthTokenType::Access => non_empty_string_claim(payload, "role").is_some(), + LocalAuthTokenType::Refresh => non_empty_string_claim(payload, "jti").is_some(), + } +} + +pub(crate) fn create_local_auth_token( + token_type: LocalAuthTokenType, + mut payload: Map, + expires_at: chrono::DateTime, +) -> Result { + let secret = local_auth_jwt_secret()?; + let header = serde_json::json!({ "alg": "HS256", "typ": "JWT" }); + payload.insert("exp".to_string(), serde_json::json!(expires_at.timestamp())); + payload.insert("type".to_string(), serde_json::json!(token_type.as_str())); + if !local_auth_claims_are_valid(&payload, token_type) { + return Err("无法签发缺少必要身份声明的Token".to_string()); + } + let header_segment = base64url_encode( + &serde_json::to_vec(&header).map_err(|_| "无法序列化JWT header".to_string())?, + ); + let payload_segment = base64url_encode( + &serde_json::to_vec(&payload).map_err(|_| "无法序列化JWT payload".to_string())?, + ); + let signing_input = format!("{header_segment}.{payload_segment}"); + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .map_err(|_| "JWT secret 无效".to_string())?; + mac.update(signing_input.as_bytes()); + let signature = mac.finalize().into_bytes(); + Ok(format!( + "{signing_input}.{}", + base64url_encode(signature.as_slice()) + )) +} + +pub(crate) fn decode_local_auth_token( + token: &str, + expected_type: LocalAuthTokenType, +) -> Result, String> { + if token.len() > MAX_LOCAL_AUTH_TOKEN_BYTES { + return Err(INVALID_TOKEN.to_string()); + } + let mut parts = token.split('.'); + let (Some(header_segment), Some(payload_segment), Some(signature_segment)) = + (parts.next(), parts.next(), parts.next()) + else { + return Err(INVALID_TOKEN.to_string()); + }; + if header_segment.is_empty() + || payload_segment.is_empty() + || signature_segment.is_empty() + || parts.next().is_some() + { + return Err(INVALID_TOKEN.to_string()); + } + + let signature = base64url_decode(signature_segment)?; + if signature.len() != 32 { + return Err(INVALID_TOKEN.to_string()); + } + let secret = local_auth_jwt_secret()?; + let signing_input = format!("{header_segment}.{payload_segment}"); + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .map_err(|_| "JWT secret 无效".to_string())?; + mac.update(signing_input.as_bytes()); + mac.verify_slice(&signature) + .map_err(|_| INVALID_TOKEN.to_string())?; + + let header_bytes = base64url_decode(header_segment)?; + let header = + serde_json::from_slice::(&header_bytes).map_err(|_| INVALID_TOKEN.to_string())?; + if header.get("alg").and_then(Value::as_str) != Some("HS256") + || header.get("typ").and_then(Value::as_str) != Some("JWT") + { + return Err(INVALID_TOKEN.to_string()); + } + + let payload_bytes = base64url_decode(payload_segment)?; + let payload = + serde_json::from_slice::(&payload_bytes).map_err(|_| INVALID_TOKEN.to_string())?; + let payload = payload + .as_object() + .cloned() + .ok_or_else(|| INVALID_TOKEN.to_string())?; + let actual_type = payload + .get("type") + .and_then(Value::as_str) + .unwrap_or_default(); + if actual_type != expected_type.as_str() { + return Err(format!( + "Token类型错误: 期望 {}, 实际 {actual_type}", + expected_type.as_str() + )); + } + let exp = payload + .get("exp") + .and_then(Value::as_i64) + .ok_or_else(|| INVALID_TOKEN.to_string())?; + if exp <= chrono::Utc::now().timestamp() { + return Err("Token已过期".to_string()); + } + if !local_auth_claims_are_valid(&payload, expected_type) { + return Err(INVALID_TOKEN.to_string()); + } + Ok(payload) +} + +pub(crate) fn local_auth_token_identity_matches_user( + payload: &Map, + user: &aether_data::repository::users::StoredUserAuthRecord, +) -> bool { + if non_empty_string_claim(payload, "user_id") != Some(user.id.as_str()) { + return false; + } + + // Access tokens carry the role used by downstream authorization checks. + // Refresh tokens intentionally omit it, so only validate the claim when + // present. A present-but-malformed role must never be treated as absent. + if let Some(token_role) = payload.get("role") { + if token_role.as_str() != Some(user.role.as_str()) { + return false; + } + } + + if let Some(token_email) = payload.get("email") { + match (token_email.as_str(), user.email.as_deref()) { + (Some(token_email), Some(email)) if token_email == email => {} + (None, None) if token_email.is_null() => {} + _ => return false, + } + } + + match (payload.get("created_at"), user.created_at) { + (Some(Value::String(token_created_at)), Some(user_created_at)) => { + let Ok(token_created_at) = chrono::DateTime::parse_from_rfc3339(token_created_at) + else { + return false; + }; + let token_created_at = token_created_at.with_timezone(&chrono::Utc); + user_created_at == token_created_at + } + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::{ + create_local_auth_token, decode_local_auth_token, local_auth_jwt_secret, + local_auth_token_identity_matches_user, validate_jwt_secret_value, LocalAuthTokenType, + MAX_LOCAL_AUTH_TOKEN_BYTES, TEST_JWT_SECRET, + }; + use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; + use hmac::Mac; + use serde_json::{json, Map, Value}; + + fn signed_token(header: Value, payload: Value) -> String { + let header = + URL_SAFE_NO_PAD.encode(serde_json::to_vec(&header).expect("header should encode")); + let payload = + URL_SAFE_NO_PAD.encode(serde_json::to_vec(&payload).expect("payload should encode")); + let signing_input = format!("{header}.{payload}"); + let secret = local_auth_jwt_secret().expect("test JWT secret should resolve"); + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .expect("test JWT secret should sign"); + mac.update(signing_input.as_bytes()); + format!( + "{signing_input}.{}", + URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()) + ) + } + + fn access_claims() -> Map { + Map::from_iter([ + ("user_id".to_string(), json!("user-1")), + ("session_id".to_string(), json!("session-1")), + ("role".to_string(), json!("user")), + ( + "created_at".to_string(), + json!(chrono::Utc::now().to_rfc3339()), + ), + ]) + } + + #[test] + fn jwt_secret_validation_fails_closed() { + assert!(validate_jwt_secret_value(None).is_err()); + assert!(validate_jwt_secret_value(Some("short-secret")).is_err()); + assert!(validate_jwt_secret_value(Some("change-this-to-a-secure-random-string")).is_err()); + assert!(validate_jwt_secret_value(Some("aether-rust-dev-jwt-secret")).is_err()); + assert!(validate_jwt_secret_value(Some(TEST_JWT_SECRET)).is_err()); + assert!(validate_jwt_secret_value(Some( + "a-valid-local-auth-secret-with-at-least-32-bytes" + )) + .is_ok()); + } + + #[test] + fn access_and_refresh_tokens_cannot_be_interchanged() { + let expires_at = chrono::Utc::now() + chrono::Duration::minutes(5); + let access_claims = access_claims(); + let mut refresh_claims = access_claims.clone(); + refresh_claims.remove("role"); + refresh_claims.insert("jti".to_string(), json!("refresh-1")); + let access = create_local_auth_token(LocalAuthTokenType::Access, access_claims, expires_at) + .expect("access token should encode"); + let refresh = + create_local_auth_token(LocalAuthTokenType::Refresh, refresh_claims, expires_at) + .expect("refresh token should encode"); + + assert!(decode_local_auth_token(&access, LocalAuthTokenType::Access).is_ok()); + assert!(decode_local_auth_token(&refresh, LocalAuthTokenType::Refresh).is_ok()); + assert!(decode_local_auth_token(&access, LocalAuthTokenType::Refresh).is_err()); + assert!(decode_local_auth_token(&refresh, LocalAuthTokenType::Access).is_err()); + } + + #[test] + fn decoder_rejects_expired_and_malformed_tokens() { + let expired = create_local_auth_token( + LocalAuthTokenType::Access, + access_claims(), + chrono::Utc::now() - chrono::Duration::seconds(1), + ) + .expect("expired token should still encode"); + + assert_eq!( + decode_local_auth_token(&expired, LocalAuthTokenType::Access), + Err("Token已过期".to_string()) + ); + for malformed in ["", "one", "one.two", "one.two.three.four", ".."] { + assert!(decode_local_auth_token(malformed, LocalAuthTokenType::Access).is_err()); + } + + let oversized = "a".repeat(MAX_LOCAL_AUTH_TOKEN_BYTES + 1); + assert!(decode_local_auth_token(&oversized, LocalAuthTokenType::Access).is_err()); + } + + #[test] + fn decoder_rejects_wrong_header_and_required_claim_shapes() { + let exp = (chrono::Utc::now() + chrono::Duration::minutes(5)).timestamp(); + let valid_payload = json!({ + "user_id": "user-1", + "session_id": "session-1", + "role": "user", + "created_at": chrono::Utc::now().to_rfc3339(), + "type": "access", + "exp": exp, + }); + for header in [ + json!({"alg": "none", "typ": "JWT"}), + json!({"alg": "HS256", "typ": "JWS"}), + ] { + let token = signed_token(header, valid_payload.clone()); + assert!(decode_local_auth_token(&token, LocalAuthTokenType::Access).is_err()); + } + + for payload in [ + json!({"session_id": "session-1", "role": "user", "created_at": null, "type": "access", "exp": exp}), + json!({"user_id": "user-1", "session_id": "session-1", "created_at": null, "type": "access", "exp": exp}), + json!({"user_id": "user-1", "session_id": "session-1", "role": "user", "type": "access", "exp": exp}), + json!({"user_id": "user-1", "session_id": "session-1", "role": "user", "created_at": null, "type": "access", "exp": exp}), + json!({"user_id": "user-1", "session_id": "session-1", "role": "user", "created_at": null, "type": "access", "exp": "later"}), + ] { + let token = signed_token(json!({"alg": "HS256", "typ": "JWT"}), payload); + assert!(decode_local_auth_token(&token, LocalAuthTokenType::Access).is_err()); + } + } + + #[test] + fn issuer_rejects_null_created_at_identity_binding() { + let mut claims = access_claims(); + claims.insert("created_at".to_string(), Value::Null); + + assert!(create_local_auth_token( + LocalAuthTokenType::Access, + claims, + chrono::Utc::now() + chrono::Duration::minutes(5), + ) + .is_err()); + } + + #[test] + fn identity_binding_rejects_even_subsecond_user_replacement() { + let created_at = chrono::Utc::now(); + let user = aether_data::repository::users::StoredUserAuthRecord::new( + "user-1".to_string(), + Some("user-1@example.com".to_string()), + true, + "user-1".to_string(), + None, + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(created_at + chrono::Duration::milliseconds(1)), + None, + ) + .expect("test user should be valid"); + let claims = Map::from_iter([ + ("created_at".to_string(), json!(created_at.to_rfc3339())), + ("user_id".to_string(), json!("user-1")), + ]); + + assert!(!local_auth_token_identity_matches_user(&claims, &user)); + + let mut replaced_id_claims = claims; + replaced_id_claims.insert("created_at".to_string(), json!(user.created_at)); + replaced_id_claims.insert("user_id".to_string(), json!("user-2")); + assert!(!local_auth_token_identity_matches_user( + &replaced_id_claims, + &user + )); + } + + #[test] + fn identity_binding_validates_present_role_claim() { + let created_at = chrono::Utc::now(); + let user = aether_data::repository::users::StoredUserAuthRecord::new( + "user-1".to_string(), + Some("user-1@example.com".to_string()), + true, + "user-1".to_string(), + None, + "admin".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(created_at), + None, + ) + .expect("test user should be valid"); + let mut claims = Map::from_iter([ + ("created_at".to_string(), json!(created_at.to_rfc3339())), + ("user_id".to_string(), json!("user-1")), + ]); + + // Refresh tokens intentionally omit role, so this remains valid. + assert!(local_auth_token_identity_matches_user(&claims, &user)); + + claims.insert("role".to_string(), json!("admin")); + assert!(local_auth_token_identity_matches_user(&claims, &user)); + + claims.insert("role".to_string(), json!("user")); + assert!(!local_auth_token_identity_matches_user(&claims, &user)); + + claims.insert("role".to_string(), Value::Null); + assert!(!local_auth_token_identity_matches_user(&claims, &user)); + } +} diff --git a/apps/aether-gateway/src/main.rs b/apps/aether-gateway/src/main.rs index 01fb9743d..571dee507 100644 --- a/apps/aether-gateway/src/main.rs +++ b/apps/aether-gateway/src/main.rs @@ -2,24 +2,123 @@ #[global_allocator] static GLOBAL: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc; -use std::path::PathBuf; +use std::fs; +#[cfg(unix)] +use std::io::Write; +use std::io::{self, Read}; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use axum::{body::Body, extract::Request}; use clap::{Args as ClapArgs, Parser, Subcommand, ValueEnum}; use hyper::body::Incoming; use hyper_util::{ - rt::{TokioExecutor, TokioIo}, + rt::{TokioExecutor, TokioIo, TokioTimer}, server::conn::auto::Builder as HyperServerBuilder, service::TowerToHyperService, }; use tower::{Service as _, ServiceExt as _}; use tracing::{debug, info, warn}; +/// Coordinates the connection-level deadline that covers protocol detection and +/// the first request header block. Hyper's HTTP/1 timer starts only after the +/// auto protocol detector has finished, while HTTP/2 has no equivalent header +/// timer. Keeping this gate outside the parser closes that initial gap without +/// imposing a deadline on request or response bodies. +#[derive(Clone)] +struct GatewayFirstRequestGate { + seen: Arc, + notify: Arc, +} + +impl GatewayFirstRequestGate { + fn new() -> Self { + Self { + seen: Arc::new(AtomicBool::new(false)), + notify: Arc::new(tokio::sync::Notify::new()), + } + } + + fn mark_seen(&self) { + if !self.seen.swap(true, Ordering::Release) { + self.notify.notify_one(); + } + } + + fn is_seen(&self) -> bool { + self.seen.load(Ordering::Acquire) + } +} + +#[derive(Clone)] +struct GatewayFirstRequestService { + inner: S, + gate: GatewayFirstRequestGate, +} + +impl tower::Service for GatewayFirstRequestService +where + S: tower::Service, +{ + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, request: Request) -> Self::Future { + self.gate.mark_seen(); + self.inner.call(request) + } +} + +/// Drive one Hyper connection while enforcing the deadline for its first +/// request. This is kept generic so the timeout behavior can be regression +/// tested independently from the listener and application state. +async fn drive_gateway_connection( + connection: F, + first_request_gate: GatewayFirstRequestGate, + first_request_timeout: std::time::Duration, +) -> Result<(), E> +where + F: std::future::Future>, +{ + if first_request_gate.is_seen() { + return connection.await; + } + + let mut connection = Box::pin(connection); + let first_request_timeout = tokio::time::sleep(first_request_timeout); + tokio::pin!(first_request_timeout); + let first_request_notified = first_request_gate.notify.notified(); + tokio::pin!(first_request_notified); + + tokio::select! { + result = &mut connection => result, + _ = &mut first_request_timeout => { + if first_request_gate.is_seen() { + (&mut connection).await + } else { + tracing::debug!( + "gateway connection closed before the first request header completed" + ); + Ok(()) + } + } + _ = &mut first_request_notified => (&mut connection).await, + } +} + use aether_crypto::warm_python_fernet_secret; use aether_data::lifecycle::export::{ copy_database_records, export_database_jsonl, import_database_jsonl, DataCopyOptions, - ExportDomain, + ExportDomain, MAX_JSONL_INPUT_BYTES, }; use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL}; use aether_gateway::{ @@ -37,6 +136,28 @@ use aether_runtime_state::{ RuntimeStateConfig, }; +const MIN_GATEWAY_DATA_ENCRYPTION_KEY_BYTES: usize = 32; +const INSECURE_GATEWAY_DATA_ENCRYPTION_KEYS: &[&str] = &[ + "change-this-to-another-secure-random-string", + "change-this-to-a-secure-random-string", + "dev-encryption-key-do-not-use-in-production", +]; + +fn validate_gateway_data_encryption_key(value: Option<&str>) -> Result<(), &'static str> { + let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else { + return Ok(()); + }; + if value.len() < MIN_GATEWAY_DATA_ENCRYPTION_KEY_BYTES { + return Err("gateway data encryption key must contain at least 32 bytes"); + } + if INSECURE_GATEWAY_DATA_ENCRYPTION_KEYS.contains(&value) { + return Err( + "gateway data encryption key must not use a published example or development value", + ); + } + Ok(()) +} + #[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] enum VideoTaskTruthSourceArg { PythonSyncReport, @@ -76,6 +197,29 @@ enum DatabaseDriverArg { Postgres, } +#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] +enum DatabaseModeArg { + Auto, + VerifyOnly, +} + +fn resolve_database_mode( + configured: Option, + legacy_auto_prepare: Option, +) -> DatabaseModeArg { + if let Some(configured) = configured { + return configured; + } + if let Some(legacy_auto_prepare) = legacy_auto_prepare { + return if legacy_auto_prepare { + DatabaseModeArg::Auto + } else { + DatabaseModeArg::VerifyOnly + }; + } + DatabaseModeArg::Auto +} + #[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] enum ExportDomainArg { Users, @@ -279,6 +423,18 @@ const MAX_GATEWAY_LISTENER_SHARDS: usize = 64; const DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 16_384; const MIN_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 200; const MAX_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 1_000_000; +// These limits protect the connection parser from slow-header and header-bomb +// attacks. They apply to request metadata only and do not cap body size or +// HTTP/2 stream concurrency. +const DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS: u64 = 30_000; +const MIN_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS: u64 = 1_000; +const MAX_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS: u64 = 300_000; +const DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES: usize = 64 * 1024; +const MIN_GATEWAY_HTTP_HEADER_MAX_BYTES: usize = 8 * 1024; +const MAX_GATEWAY_HTTP_HEADER_MAX_BYTES: usize = 16 * 1024 * 1024; +const DEFAULT_GATEWAY_HTTP_MAX_HEADERS: usize = 256; +const MIN_GATEWAY_HTTP_MAX_HEADERS: usize = 16; +const MAX_GATEWAY_HTTP_MAX_HEADERS: usize = 4_096; const AUTO_GATEWAY_REQUESTS_PER_CPU: usize = 1_024; const MIN_AUTO_GATEWAY_REQUEST_CONCURRENCY: usize = 512; const MAX_AUTO_GATEWAY_REQUEST_CONCURRENCY: usize = 65_536; @@ -538,51 +694,81 @@ fn automatic_sql_pool_config_for_parallelism( #[derive(ClapArgs, Debug, Clone)] struct GatewayDataArgs { - #[arg(long, env = "AETHER_DATABASE_DRIVER")] + #[arg(long, env = "AETHER_DATABASE_DRIVER", global = true)] database_driver: Option, - #[arg(long, env = "AETHER_DATABASE_URL")] + #[arg(long, env = "AETHER_DATABASE_URL", global = true)] database_url: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_URL")] + #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_URL", global = true)] postgres_url: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_ENCRYPTION_KEY")] + #[arg(long, env = "AETHER_GATEWAY_DATA_ENCRYPTION_KEY", global = true)] encryption_key: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_REDIS_URL")] + #[arg(long, env = "AETHER_GATEWAY_DATA_REDIS_URL", global = true)] redis_url: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_REDIS_KEY_PREFIX")] + #[arg(long, env = "AETHER_GATEWAY_DATA_REDIS_KEY_PREFIX", global = true)] redis_key_prefix: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS", + global = true + )] postgres_min_connections: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS", + global = true + )] postgres_max_connections: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_ACQUIRE_TIMEOUT_MS")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_ACQUIRE_TIMEOUT_MS", + global = true + )] postgres_acquire_timeout_ms: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_IDLE_TIMEOUT_MS")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_IDLE_TIMEOUT_MS", + global = true + )] postgres_idle_timeout_ms: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_MAX_LIFETIME_MS")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_MAX_LIFETIME_MS", + global = true + )] postgres_max_lifetime_ms: Option, - #[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_CACHE_CAPACITY")] + #[arg( + long, + env = "AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_CACHE_CAPACITY", + global = true + )] postgres_statement_cache_capacity: Option, #[arg( long, env = "AETHER_GATEWAY_DATA_POSTGRES_REQUIRE_SSL", - default_value_t = false + default_value_t = false, + global = true )] postgres_require_ssl: bool, } impl GatewayDataArgs { + fn validate_encryption_key(&self) -> Result<(), std::io::Error> { + validate_gateway_data_encryption_key(self.effective_encryption_key().as_deref()) + .map_err(|message| std::io::Error::new(std::io::ErrorKind::InvalidInput, message)) + } + fn effective_database_driver(&self) -> Option { self.database_driver.map(Into::into).or_else(|| { self.database_url @@ -1100,7 +1286,9 @@ struct GatewayRateLimitArgs { #[arg(long, env = "RPM_KEY_TTL_SECONDS", default_value_t = 120)] key_ttl_seconds: u64, - #[arg(long, env = "RATE_LIMIT_FAIL_OPEN", default_value_t = true)] + /// Explicitly allow requests when the shared RPM backend is unavailable. + /// Keep the secure fail-closed behavior as the production default. + #[arg(long, env = "RATE_LIMIT_FAIL_OPEN", default_value_t = false)] fail_open: bool, } @@ -1144,25 +1332,39 @@ enum DataCommand { Import(DataImportArgs), /// Copy persistent SQL data directly between two databases without a JSONL file. Copy(DataCopyArgs), + /// Inspect or prepare the configured database. + Db(DatabaseCommandArgs), +} + +#[derive(ClapArgs, Debug, Clone)] +struct DatabaseCommandArgs { + #[command(subcommand)] + command: DatabaseCommand, +} + +#[derive(Subcommand, Debug, Clone)] +enum DatabaseCommand { + /// Show whether schema migrations and data backfills are current. + Status, + /// Apply pending schema migrations and data backfills. + Prepare, } #[derive(ClapArgs, Debug, Clone)] struct DataExportArgs { - #[command(flatten)] - data: GatewayDataArgs, - #[arg(long)] output: PathBuf, + /// Atomically replace an existing regular output owned by the current user. + #[arg(long)] + overwrite: bool, + #[arg(long, value_enum, value_delimiter = ',')] domains: Vec, } #[derive(ClapArgs, Debug, Clone)] struct DataImportArgs { - #[command(flatten)] - data: GatewayDataArgs, - #[arg(long)] input: PathBuf, } @@ -1175,12 +1377,22 @@ struct DataCopyArgs { #[arg(long)] source_url: String, + /// Permit a cleartext source connection for a non-loopback database. + /// Leave unset to require TLS for remote MySQL/Postgres URLs. + #[arg(long)] + source_allow_insecure: bool, + #[arg(long, value_enum)] target_driver: DatabaseDriverArg, #[arg(long)] target_url: String, + /// Permit a cleartext target connection for a non-loopback database. + /// Leave unset to require TLS for remote MySQL/Postgres URLs. + #[arg(long)] + target_allow_insecure: bool, + #[arg(long, value_enum, value_delimiter = ',')] domains: Vec, @@ -1256,6 +1468,30 @@ struct Args { )] http2_max_concurrent_streams: u32, + #[arg( + long, + env = "AETHER_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS", + default_value_t = DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS + )] + /// Maximum time allowed to receive one complete HTTP request header block. + http_header_read_timeout_ms: u64, + + #[arg( + long, + env = "AETHER_GATEWAY_HTTP_HEADER_MAX_BYTES", + default_value_t = DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES + )] + /// Maximum HTTP request header bytes (HTTP/2 uses decompressed list size). + http_header_max_bytes: usize, + + #[arg( + long, + env = "AETHER_GATEWAY_HTTP_MAX_HEADERS", + default_value_t = DEFAULT_GATEWAY_HTTP_MAX_HEADERS + )] + /// Maximum number of HTTP/1 request header fields. + http_max_headers: usize, + /// 容器内健康检查入口:根据当前 bind 端口探测本地 /health。 #[arg(long, hide = true, default_value_t = false)] healthcheck: bool, @@ -1284,18 +1520,25 @@ struct Args { )] node_role: NodeRoleArg, - #[arg(long, default_value_t = false)] + #[arg(long, hide = true, default_value_t = false)] migrate: bool, - #[arg(long, default_value_t = false)] + #[arg(long, hide = true, default_value_t = false)] apply_backfills: bool, + /// Database startup policy. Defaults to auto when neither this nor the legacy setting is set. + #[arg(long, env = "AETHER_GATEWAY_DATABASE_MODE", value_enum)] + database_mode: Option, + + /// Legacy compatibility switch. Prefer --database-mode. #[arg( long, env = "AETHER_GATEWAY_AUTO_PREPARE_DATABASE", - default_value_t = false + hide = true, + num_args = 0..=1, + default_missing_value = "true" )] - auto_prepare_database: bool, + auto_prepare_database: Option, /// Path to frontend static files directory (SPA). When set, the gateway /// serves the frontend directly without nginx. @@ -1405,6 +1648,10 @@ struct Args { } impl Args { + fn effective_database_mode(&self) -> DatabaseModeArg { + resolve_database_mode(self.database_mode, self.auto_prepare_database) + } + fn effective_runtime_backend( &self, database: Option<&SqlDatabaseConfig>, @@ -1475,15 +1722,7 @@ impl Args { } fn runtime_config(&self) -> Result { - let default_log_filter = if self.command.is_some() - || self.migrate - || self.apply_backfills - || self.auto_prepare_database - { - "aether_gateway=info,aether_data=info" - } else { - "aether_gateway=info" - }; + let default_log_filter = "aether_gateway=info,aether_data=info"; let config = self .logging .apply_to_runtime_config(ServiceRuntimeConfig::new( @@ -1553,6 +1792,24 @@ fn gateway_http2_max_concurrent_streams(streams: u32) -> u32 { ) } +fn gateway_http_header_read_timeout_ms(timeout_ms: u64) -> u64 { + timeout_ms.clamp( + MIN_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS, + MAX_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS, + ) +} + +fn gateway_http_header_max_bytes(bytes: usize) -> usize { + bytes.clamp( + MIN_GATEWAY_HTTP_HEADER_MAX_BYTES, + MAX_GATEWAY_HTTP_HEADER_MAX_BYTES, + ) +} + +fn gateway_http_max_headers(headers: usize) -> usize { + headers.clamp(MIN_GATEWAY_HTTP_MAX_HEADERS, MAX_GATEWAY_HTTP_MAX_HEADERS) +} + fn gateway_listener( bind_addr: std::net::SocketAddr, backlog: i32, @@ -1604,14 +1861,29 @@ async fn serve_gateway_router( listeners: Vec, router: axum::Router, http2_max_concurrent_streams: u32, + http_header_read_timeout_ms: u64, + http_header_max_bytes: usize, + http_max_headers: usize, ) -> Result<(), Box> { let http2_max_concurrent_streams = gateway_http2_max_concurrent_streams(http2_max_concurrent_streams); + let http_header_read_timeout_ms = + gateway_http_header_read_timeout_ms(http_header_read_timeout_ms); + let http_header_max_bytes = gateway_http_header_max_bytes(http_header_max_bytes); + let http_max_headers = gateway_http_max_headers(http_max_headers); let mut servers = tokio::task::JoinSet::new(); for listener in listeners { let router = router.clone(); servers.spawn(async move { - serve_gateway_listener(listener, router, http2_max_concurrent_streams).await + serve_gateway_listener( + listener, + router, + http2_max_concurrent_streams, + http_header_read_timeout_ms, + http_header_max_bytes, + http_max_headers, + ) + .await }); } if let Some(result) = servers.join_next().await { @@ -1627,6 +1899,9 @@ async fn serve_gateway_listener( listener: tokio::net::TcpListener, router: axum::Router, http2_max_concurrent_streams: u32, + http_header_read_timeout_ms: u64, + http_header_max_bytes: usize, + http_max_headers: usize, ) -> Result<(), std::io::Error> { let mut make_service = router.into_make_service_with_connect_info::(); loop { @@ -1636,19 +1911,49 @@ async fn serve_gateway_listener( .await .unwrap_or_else(|err| match err {}) .map_request(|req: Request| req.map(Body::new)); - let hyper_service = TowerToHyperService::new(tower_service); + let first_request_gate = GatewayFirstRequestGate::new(); + let hyper_service = TowerToHyperService::new(GatewayFirstRequestService { + inner: tower_service, + gate: first_request_gate.clone(), + }); let io = TokioIo::new(io); tokio::spawn(async move { let mut builder = HyperServerBuilder::new(TokioExecutor::new()); + // Hyper's HTTP/1 header timer is opt-in when using the custom + // connection builder. Configure both protocol parsers explicitly: + // HTTP/1 gets a slow-header deadline and bounded parser buffer; + // HTTP/2 gets a decompressed header-list limit. The timer is + // connection metadata protection and does not affect request body + // streaming or the configured stream concurrency. + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(std::time::Duration::from_millis( + http_header_read_timeout_ms, + )) + .max_buf_size(http_header_max_bytes) + .max_headers(http_max_headers); builder.http2().enable_connect_protocol(); builder .http2() - .max_concurrent_streams(http2_max_concurrent_streams); - if let Err(err) = builder - .serve_connection_with_upgrades(io, hyper_service) - .await - { + .timer(TokioTimer::new()) + .max_concurrent_streams(http2_max_concurrent_streams) + .max_header_list_size(u32::try_from(http_header_max_bytes).unwrap_or(u32::MAX)); + + // The auto builder reads the HTTP/2 preface before Hyper's H1 + // header timer starts, and H2 has no header-read timer of its own. + // Race the whole connection until the first valid request reaches + // the service so a peer cannot hold a socket open while dribbling + // protocol bytes or an initial header block. Once the gate opens, + // request and response bodies remain fully streaming. + let connection_result = drive_gateway_connection( + builder.serve_connection_with_upgrades(io, hyper_service), + first_request_gate, + std::time::Duration::from_millis(http_header_read_timeout_ms), + ) + .await; + if let Err(err) = connection_result { tracing::trace!(error = ?err, "gateway connection closed with error"); } }); @@ -1669,6 +1974,7 @@ async fn run_healthcheck( ) -> Result<(), Box> { let url = resolve_healthcheck_url(app_port)?; reqwest::Client::builder() + .no_proxy() .timeout(std::time::Duration::from_millis( healthcheck_timeout_ms.max(1), )) @@ -1786,7 +2092,13 @@ async fn run() -> Result<(), Box> { let args = Args::parse(); if let Some(command) = args.command.as_ref() { init_service_runtime(args.runtime_config()?)?; - return run_data_command(command).await; + // Data export/import can decrypt and persist sensitive credentials; + // apply the same encryption-key policy as the normal gateway path + // before touching the selected database. + if matches!(command, DataCommand::Export(_) | DataCommand::Import(_)) { + args.data.validate_encryption_key()?; + } + return run_data_command(command, &args.data).await; } if args.migrate { init_service_runtime(args.runtime_config()?)?; @@ -1814,6 +2126,7 @@ async fn run() -> Result<(), Box> { runtime_redis_url.as_deref(), runtime_backend, )?; + args.data.validate_encryption_key()?; let data_config = args.data.to_config(); let isolate_background_database = args.node_role.isolates_background_database(); let background_database_config = if isolate_background_database { @@ -2086,7 +2399,7 @@ async fn run() -> Result<(), Box> { execution_runtime_configured = state.execution_runtime_configured(), "aether-gateway data layer configured" ); - prepare_database_startup_requirements(&state, args.auto_prepare_database).await?; + prepare_database_startup_requirements(&state, args.effective_database_mode()).await?; state.warm_database_pools().await?; let reset_stale_proxy_nodes = state.reset_stale_proxy_node_tunnel_statuses().await?; if reset_stale_proxy_nodes > 0 { @@ -2096,6 +2409,17 @@ async fn run() -> Result<(), Box> { ); } state.bootstrap_admin_from_env().await?; + match state.ensure_system_default_routing_group().await { + Ok(Some(group)) => { + info!( + group_id = %group.id, + group_name = %group.name, + "created system default routing group from routing strategy defaults" + ); + } + Ok(None) => {} + Err(err) => return Err(err.into()), + } match state.prewarm_chat_pii_redaction_runtime_config().await { Ok(enabled) => { info!( @@ -2110,6 +2434,10 @@ async fn run() -> Result<(), Box> { ); } } + match state.prewarm_execution_extra_trusted_dns_hosts().await { + Ok(_) => info!("prewarmed execution Fake-IP DNS allowlist"), + Err(err) => warn!(error = %err, "failed to prewarm execution Fake-IP DNS allowlist"), + } match prewarm_direct_h2c_sender_cache_from_env_for_startup().await { Ok(Some(report)) => { if report.failed_targets > 0 { @@ -2185,21 +2513,92 @@ async fn run() -> Result<(), Box> { "aether-gateway ready" ); - serve_gateway_router(listeners, router, args.http2_max_concurrent_streams).await?; + serve_gateway_router( + listeners, + router, + args.http2_max_concurrent_streams, + args.http_header_read_timeout_ms, + args.http_header_max_bytes, + args.http_max_headers, + ) + .await?; if let Some(background_tasks) = background_tasks { background_tasks.shutdown().await; } Ok(()) } -async fn run_data_command(command: &DataCommand) -> Result<(), Box> { +async fn run_data_command( + command: &DataCommand, + data: &GatewayDataArgs, +) -> Result<(), Box> { match command { - DataCommand::Export(args) => run_data_export(args).await, - DataCommand::Import(args) => run_data_import(args).await, + DataCommand::Export(args) => run_data_export(args, data).await, + DataCommand::Import(args) => run_data_import(args, data).await, DataCommand::Copy(args) => run_data_copy(args).await, + DataCommand::Db(args) => run_database_command(args, data).await, } } +async fn run_database_command( + args: &DatabaseCommandArgs, + data: &GatewayDataArgs, +) -> Result<(), Box> { + match args.command { + DatabaseCommand::Status => run_database_status(data).await, + DatabaseCommand::Prepare => run_database_prepare(data).await, + } +} + +fn database_maintenance_state( + data: &GatewayDataArgs, +) -> Result<(DatabaseDriver, AppState), Box> { + let database = required_sql_database_config(data)?; + let driver = database.driver; + let state = AppState::new()?.with_data_config(data.to_config())?; + Ok((driver, state)) +} + +async fn run_database_status(data: &GatewayDataArgs) -> Result<(), Box> { + let (driver, state) = database_maintenance_state(data)?; + let pending_migrations = state + .pending_database_migrations() + .await? + .unwrap_or_default(); + + if let Some(next) = pending_migrations.first() { + println!("database {driver}: preparation required"); + println!("pending migrations: {}", pending_migrations.len()); + println!("next migration: {} ({})", next.version, next.description); + println!("pending backfills: not checked until migrations are current"); + println!("run `aether-gateway db prepare`"); + return Ok(()); + } + + let pending_backfills = state + .pending_database_backfills() + .await? + .unwrap_or_default(); + if let Some(next) = pending_backfills.first() { + println!("database {driver}: preparation required"); + println!("pending migrations: 0"); + println!("pending backfills: {}", pending_backfills.len()); + println!("next backfill: {} ({})", next.version, next.description); + println!("run `aether-gateway db prepare`"); + return Ok(()); + } + + println!("database {driver}: ready (schema and backfills are current)"); + Ok(()) +} + +async fn run_database_prepare(data: &GatewayDataArgs) -> Result<(), Box> { + let (driver, state) = database_maintenance_state(data)?; + prepare_database_startup_requirements(&state, DatabaseModeArg::Auto).await?; + println!("database {driver}: ready (schema and backfills are current)"); + Ok(()) +} + fn required_sql_database_config( data: &GatewayDataArgs, ) -> Result> { @@ -2226,14 +2625,17 @@ fn current_unix_secs() -> Result { .as_secs()) } -async fn run_data_export(args: &DataExportArgs) -> Result<(), Box> { - let database = required_sql_database_config(&args.data)?; +async fn run_data_export( + args: &DataExportArgs, + data: &GatewayDataArgs, +) -> Result<(), Box> { + let database = required_sql_database_config(data)?; let driver = database.driver; let domains = requested_export_domains(args); let created_at_unix_secs = current_unix_secs()?; let encoded = export_database_jsonl(database, domains, created_at_unix_secs).await?; - tokio::fs::write(&args.output, encoded.as_bytes()).await?; + write_atomic_private_export(&args.output, encoded.as_bytes(), args.overwrite)?; info!( driver = %driver, output = %args.output.display(), @@ -2249,10 +2651,290 @@ async fn run_data_export(args: &DataExportArgs) -> Result<(), Box Result<(), Box> { - let database = required_sql_database_config(&args.data)?; +fn write_atomic_private_export(path: &Path, bytes: &[u8], overwrite: bool) -> io::Result<()> { + #[cfg(not(unix))] + { + let _ = (path, bytes, overwrite); + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "private atomic database exports currently require Unix filesystem checks", + )); + } + + #[cfg(unix)] + { + use std::ffi::CString; + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::ffi::OsStrExt; + use std::os::unix::fs::PermissionsExt; + + let file_name = path.file_name().ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidInput, "export path must name a file") + })?; + let input_parent = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + let parent = open_private_export_directory(input_parent)?; + let output_name = CString::new(file_name.as_bytes()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "export path must not contain an embedded NUL byte", + ) + })?; + + // Check the target through the already-open parent directory. The + // later renameat/linkat calls use that same descriptor, so replacing a + // writable ancestor cannot redirect the export to another directory. + if let Some(stat) = private_export_stat_at(&parent, &output_name)? { + let effective_uid = unsafe { libc::geteuid() }; + if stat.st_mode & libc::S_IFMT != libc::S_IFREG + || stat.st_uid != effective_uid + || stat.st_nlink != 1 + { + return Err(io::Error::other( + "export output must be a regular, single-link file owned by the current user", + )); + } + if !overwrite { + return Err(io::Error::new( + io::ErrorKind::AlreadyExists, + "export output already exists; pass --overwrite to replace it", + )); + } + } + + let temporary_name = CString::new(format!( + ".aether-data-export-{}-{}.tmp", + std::process::id(), + uuid::Uuid::new_v4() + )) + .expect("generated temporary export name cannot contain NUL"); + // O_EXCL + O_NOFOLLOW makes creation of the temporary file independent + // of any attacker-controlled directory entry with the same name. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + temporary_name.as_ptr(), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_CLOEXEC | libc::O_NOFOLLOW, + 0o600, + ) + }; + if descriptor < 0 { + return Err(io::Error::last_os_error()); + } + let mut file = unsafe { fs::File::from_raw_fd(descriptor) }; + let result = (|| -> io::Result<()> { + file.set_permissions(fs::Permissions::from_mode(0o600))?; + file.write_all(bytes)?; + file.sync_all()?; + drop(file); + + if overwrite { + // renameat replaces the directory entry and never dereferences + // a destination symlink. No attacker-selected file is opened + // or truncated even if the target changed after the check. + private_export_rename_at(&parent, &temporary_name, &output_name)?; + } else { + private_export_link_at(&parent, &temporary_name, &output_name).map_err( + |error| { + if error.kind() == io::ErrorKind::AlreadyExists { + io::Error::new( + io::ErrorKind::AlreadyExists, + "export output already exists; pass --overwrite to replace it", + ) + } else { + error + } + }, + )?; + private_export_unlink_at(&parent, &temporary_name)?; + } + parent.sync_all() + })(); + if result.is_err() { + let _ = private_export_unlink_at(&parent, &temporary_name); + } + result + } +} + +#[cfg(unix)] +fn open_private_export_directory(path: &Path) -> io::Result { + use std::ffi::CString; + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::ffi::OsStrExt; + use std::path::Component; + + // Walk one component at a time and retain the final descriptor. We allow + // trusted system symlink components (for example macOS `/var`), but + // validate the directory reached by every open before continuing. The + // descriptor remains pinned even if the symlink is exchanged later. + let mut directory = fs::File::open(if path.is_absolute() { "/" } else { "." })?; + validate_private_export_directory_fd(&directory, path)?; + for component in path.components() { + let name = match component { + Component::RootDir | Component::CurDir => continue, + Component::Normal(name) => name, + Component::ParentDir => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "export path must not contain '..' components", + )) + } + Component::Prefix(_) => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "export path uses an unsupported prefix", + )) + } + }; + let name = CString::new(name.as_bytes()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "export directory path must not contain an embedded NUL byte", + ) + })?; + let descriptor = unsafe { + libc::openat( + directory.as_raw_fd(), + name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_DIRECTORY, + ) + }; + if descriptor < 0 { + return Err(io::Error::last_os_error()); + } + let next = unsafe { fs::File::from_raw_fd(descriptor) }; + validate_private_export_directory_fd(&next, path)?; + directory = next; + } + Ok(directory) +} + +#[cfg(unix)] +fn validate_private_export_directory_fd( + directory: &fs::File, + display_path: &Path, +) -> io::Result<()> { + use std::mem::MaybeUninit; + use std::os::fd::AsRawFd; + + let mut stat = MaybeUninit::::uninit(); + let result = unsafe { libc::fstat(directory.as_raw_fd(), stat.as_mut_ptr()) }; + if result != 0 { + return Err(io::Error::last_os_error()); + } + let stat = unsafe { stat.assume_init() }; + let effective_uid = unsafe { libc::geteuid() }; + let mode = stat.st_mode; + if mode & libc::S_IFMT != libc::S_IFDIR + || (stat.st_uid != effective_uid && stat.st_uid != 0) + || (mode & 0o022 != 0 && mode & 0o1000 == 0) + { + return Err(io::Error::other(format!( + "export output directory '{}' has unsafe ownership or permissions", + display_path.display() + ))); + } + Ok(()) +} + +#[cfg(unix)] +fn private_export_stat_at( + parent: &fs::File, + name: &std::ffi::CStr, +) -> io::Result> { + use std::mem::MaybeUninit; + use std::os::fd::AsRawFd; + + let mut stat = MaybeUninit::::uninit(); + let result = unsafe { + libc::fstatat( + parent.as_raw_fd(), + name.as_ptr(), + stat.as_mut_ptr(), + libc::AT_SYMLINK_NOFOLLOW, + ) + }; + if result == 0 { + return Ok(Some(unsafe { stat.assume_init() })); + } + let error = io::Error::last_os_error(); + if error.kind() == io::ErrorKind::NotFound { + Ok(None) + } else { + Err(error) + } +} + +#[cfg(unix)] +fn private_export_link_at( + parent: &fs::File, + source: &std::ffi::CStr, + destination: &std::ffi::CStr, +) -> io::Result<()> { + use std::os::fd::AsRawFd; + + let result = unsafe { + libc::linkat( + parent.as_raw_fd(), + source.as_ptr(), + parent.as_raw_fd(), + destination.as_ptr(), + 0, + ) + }; + if result == 0 { + Ok(()) + } else { + Err(io::Error::last_os_error()) + } +} + +#[cfg(unix)] +fn private_export_rename_at( + parent: &fs::File, + source: &std::ffi::CStr, + destination: &std::ffi::CStr, +) -> io::Result<()> { + use std::os::fd::AsRawFd; + + let result = unsafe { + libc::renameat( + parent.as_raw_fd(), + source.as_ptr(), + parent.as_raw_fd(), + destination.as_ptr(), + ) + }; + if result == 0 { + Ok(()) + } else { + Err(io::Error::last_os_error()) + } +} + +#[cfg(unix)] +fn private_export_unlink_at(parent: &fs::File, name: &std::ffi::CStr) -> io::Result<()> { + use std::os::fd::AsRawFd; + + let result = unsafe { libc::unlinkat(parent.as_raw_fd(), name.as_ptr(), 0) }; + if result == 0 { + Ok(()) + } else { + Err(io::Error::last_os_error()) + } +} + +async fn run_data_import( + args: &DataImportArgs, + data: &GatewayDataArgs, +) -> Result<(), Box> { + let database = required_sql_database_config(data)?; let driver = database.driver; - let input = tokio::fs::read_to_string(&args.input).await?; + let input_path = args.input.clone(); + let input = tokio::task::spawn_blocking(move || read_data_import_input(&input_path)).await??; let imported = import_database_jsonl(database, &input).await?; info!( @@ -2270,10 +2952,225 @@ async fn run_data_import(args: &DataImportArgs) -> Result<(), Box io::Result { + read_data_import_input_with_limit(path, MAX_JSONL_INPUT_BYTES) +} + +fn read_data_import_input_with_limit(path: &Path, limit: usize) -> io::Result { + let mut file = open_data_import_file(path)?; + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "data import input '{}' must be a regular file", + path.display() + ), + )); + } + + let limit_u64 = u64::try_from(limit).unwrap_or(u64::MAX); + if metadata.len() > limit_u64 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "data import input '{}' exceeds the {} byte limit", + path.display(), + limit + ), + )); + } + + // A file can grow after metadata() returns. Reading one extra byte catches + // that race without allowing the input buffer to exceed the configured + // parser budget. + let read_limit = limit.saturating_add(1); + // Do not reserve the whole metadata length: sparse or concurrently grown + // files can advertise a huge size while containing little data, and a + // single capacity reservation would otherwise become a local DoS vector. + const MAX_INITIAL_IMPORT_READ_CAPACITY: usize = 8 * 1024 * 1024; + let initial_capacity = usize::try_from(metadata.len()) + .unwrap_or(limit) + .min(limit) + .min(MAX_INITIAL_IMPORT_READ_CAPACITY); + let mut bytes = Vec::with_capacity(initial_capacity.min(read_limit)); + // Read in fixed-size chunks instead of `read_to_end`: the latter may use a + // file's attacker-controlled size hint to reserve a large buffer before + // the limit check runs. Reserve only the exact next chunk so capacity + // stays close to the configured `limit + 1` budget. + let mut chunk = [0_u8; 32 * 1024]; + while bytes.len() < read_limit { + let remaining = read_limit - bytes.len(); + let chunk_len = remaining.min(chunk.len()); + let read = file.read(&mut chunk[..chunk_len])?; + if read == 0 { + break; + } + bytes.try_reserve_exact(read).map_err(|error| { + io::Error::other(format!( + "data import input buffer allocation failed: {error}" + )) + })?; + bytes.extend_from_slice(&chunk[..read]); + } + if bytes.len() > limit { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "data import input '{}' exceeds the {} byte limit", + path.display(), + limit + ), + )); + } + + String::from_utf8(bytes).map_err(|error| { + io::Error::new( + io::ErrorKind::InvalidData, + format!( + "data import input '{}' is not valid UTF-8: {error}", + path.display() + ), + ) + }) +} + +#[cfg(unix)] +fn open_data_import_file(path: &Path) -> io::Result { + use std::ffi::CString; + use std::os::fd::{AsRawFd, FromRawFd}; + use std::os::unix::ffi::OsStrExt; + + let file_name = path.file_name().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "data import input path must name a file", + ) + })?; + let input_parent = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + // Keep the directory descriptor alive through openat. This prevents a + // concurrent rename of an ancestor from changing which directory is used. + let parent = open_private_export_directory(input_parent)?; + let file_name = CString::new(file_name.as_bytes()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "data import input path must not contain an embedded NUL byte", + ) + })?; + + // O_NONBLOCK is important even though imports require regular files: + // opening a FIFO without it can block before fstat has a chance to reject + // the special file. O_NOFOLLOW makes the final path component race-free. + let descriptor = unsafe { + libc::openat( + parent.as_raw_fd(), + file_name.as_ptr(), + libc::O_RDONLY | libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_NONBLOCK, + ) + }; + if descriptor < 0 { + let error = io::Error::last_os_error(); + if error.raw_os_error() == Some(libc::ELOOP) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "data import input '{}' must not be a symbolic link", + path.display() + ), + )); + } + return Err(error); + } + + // SAFETY: openat returned a new descriptor owned by this function; no + // other owner exists and File closes it on every return path. + Ok(unsafe { fs::File::from_raw_fd(descriptor) }) +} + +#[cfg(not(unix))] +fn open_data_import_file(path: &Path) -> io::Result { + let metadata = fs::symlink_metadata(path)?; + if metadata.file_type().is_symlink() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "data import input '{}' must not be a symbolic link", + path.display() + ), + )); + } + fs::OpenOptions::new().read(true).open(path) +} + +fn copy_database_host_is_literal_loopback(host: &str) -> bool { + let host = host.trim().trim_start_matches('[').trim_end_matches(']'); + if host.eq_ignore_ascii_case("localhost") { + return true; + } + let Ok(address) = host.parse::() else { + return false; + }; + match address { + std::net::IpAddr::V4(address) => address.is_loopback(), + std::net::IpAddr::V6(address) => { + address.is_loopback() + || address + .to_ipv4_mapped() + .is_some_and(|mapped| mapped.is_loopback()) + } + } +} + +fn copy_database_url_is_literal_loopback( + driver: DatabaseDriver, + url: &str, + label: &str, +) -> Result { + // Parse with the same SQLx driver that will open the pool. This preserves + // query-parameter overrides such as PostgreSQL `host`/`hostaddr` and + // MySQL/PostgreSQL Unix `socket` paths, which URL authority inspection + // alone would miss. + match driver { + DatabaseDriver::Postgres => { + let options = url + .parse::() + .map_err(|error| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("{label} database URL is invalid: {error}"), + ) + })?; + Ok(options.get_socket().is_some() + || copy_database_host_is_literal_loopback(options.get_host())) + } + DatabaseDriver::Mysql => { + let options = url + .parse::() + .map_err(|error| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("{label} database URL is invalid: {error}"), + ) + })?; + Ok(options.get_socket().is_some() + || copy_database_host_is_literal_loopback(options.get_host())) + } + DatabaseDriver::Sqlite => Ok(false), + } +} + fn copy_database_config( driver: DatabaseDriverArg, url: &str, label: &str, + allow_insecure: bool, ) -> Result> { let url = url.trim(); if url.is_empty() { @@ -2284,19 +3181,39 @@ fn copy_database_config( .into()); } let driver = DatabaseDriver::from(driver); + // A loopback exception preserves the existing local-development workflow, + // while every named/remote SQL host defaults to an encrypted connection. + // `allow_insecure` is deliberately endpoint-specific so a local source + // does not silently downgrade a remote target (or vice versa). + let literal_loopback = if driver == DatabaseDriver::Sqlite { + false + } else { + copy_database_url_is_literal_loopback(driver, url, label)? + }; + let require_ssl = driver != DatabaseDriver::Sqlite && !allow_insecure && !literal_loopback; Ok(SqlDatabaseConfig::new( driver, url, SqlPoolConfig { - require_ssl: false, + require_ssl, ..SqlPoolConfig::default() }, )?) } async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box> { - let source = copy_database_config(args.source_driver, &args.source_url, "source")?; - let target = copy_database_config(args.target_driver, &args.target_url, "target")?; + let source = copy_database_config( + args.source_driver, + &args.source_url, + "source", + args.source_allow_insecure, + )?; + let target = copy_database_config( + args.target_driver, + &args.target_url, + "target", + args.target_allow_insecure, + )?; let source_driver = source.driver; let target_driver = target.driver; let domains = requested_domains(&args.domains); @@ -2340,6 +3257,7 @@ async fn run_explicit_migrations(args: &Args) -> Result<(), Box Result<(), Box Result<(), Box Result<(), Box> { - if !auto_prepare_database { + if matches!(database_mode, DatabaseModeArg::VerifyOnly) { ensure_database_schema_is_current(state).await?; ensure_database_backfills_are_current(state).await?; return Ok(()); } - info!( - "auto database preparation enabled; applying pending migrations and backfills before serving traffic" - ); + info!("database preparation enabled; applying pending migrations and backfills"); let Some(pending_migrations) = state.prepare_database_for_startup().await? else { return Ok(()); @@ -2434,10 +3351,10 @@ async fn prepare_database_startup_requirements( next_version = next.version, next_description = %next.description, pending_versions = %format_pending_migrations(&pending_migrations), - "running database migrations during service startup..." + "running database migrations during database preparation..." ); if state.run_database_migrations().await? { - info!("database migrations complete during service startup"); + info!("database migrations complete"); } } @@ -2456,10 +3373,10 @@ async fn prepare_database_startup_requirements( next_version = next.version, next_description = %next.description, pending_versions = %format_pending_backfills(&pending_backfills), - "running database backfills during service startup..." + "running database backfills during database preparation..." ); if state.run_database_backfills().await? { - info!("database backfills complete during service startup"); + info!("database backfills complete"); } Ok(()) @@ -2504,7 +3421,7 @@ async fn ensure_database_backfills_are_current( async fn ensure_database_schema_is_current( state: &AppState, ) -> Result<(), Box> { - let Some(pending) = state.prepare_database_for_startup().await? else { + let Some(pending) = state.pending_database_migrations().await? else { return Ok(()); }; if pending.is_empty() { @@ -2523,7 +3440,7 @@ fn pending_schema_error( next_description: &str, ) -> std::io::Error { std::io::Error::other(format!( - "database schema is behind by {} migration(s); next pending migration is {} ({})\nrun `aether-gateway --migrate` before starting the service", + "database schema is behind by {} migration(s); next pending migration is {} ({})\nrun `aether-gateway db prepare` before starting the service", pending_count, next_version, next_description )) } @@ -2534,7 +3451,7 @@ fn pending_backfills_error( next_description: &str, ) -> std::io::Error { std::io::Error::other(format!( - "database backfills are behind by {} backfill(s); next pending backfill is {} ({})\nrun `aether-gateway --apply-backfills` before starting the service", + "database backfills are behind by {} backfill(s); next pending backfill is {} ({})\nrun `aether-gateway db prepare` before starting the service", pending_count, next_version, next_description )) } @@ -2545,19 +3462,32 @@ mod tests { automatic_gateway_request_concurrency_for_capacity, automatic_gateway_request_concurrency_for_parallelism, automatic_sql_pool_config, automatic_sql_pool_config_for_parallelism, automatic_usage_queue_workers_for_parallelism, - ensure_database_backfills_are_current, ensure_database_schema_is_current, - pending_backfills_error, pending_schema_error, resolve_healthcheck_url, - usage_database_config_for_role, Args, DatabaseDriverArg, DeploymentTopologyArg, - GatewayDataArgs, GatewayFrontdoorArgs, GatewayLogDestinationArg, GatewayLogFormatArg, - GatewayLogRotationArg, GatewayLoggingArgs, GatewayRateLimitArgs, GatewayUsageArgs, - NodeRoleArg, RuntimeBackendArg, VideoTaskTruthSourceArg, - DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, DEFAULT_GATEWAY_LISTENER_SHARDS, + copy_database_config, ensure_database_backfills_are_current, + ensure_database_schema_is_current, pending_backfills_error, pending_schema_error, + read_data_import_input_with_limit, resolve_database_mode, resolve_healthcheck_url, + usage_database_config_for_role, validate_gateway_data_encryption_key, + write_atomic_private_export, Args, DataCommand, DatabaseCommand, DatabaseDriverArg, + DatabaseModeArg, DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs, + GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs, + GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg, + VideoTaskTruthSourceArg, DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, + DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES, DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS, + DEFAULT_GATEWAY_HTTP_MAX_HEADERS, DEFAULT_GATEWAY_LISTENER_SHARDS, DEFAULT_GATEWAY_LISTEN_BACKLOG, MAX_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, MAX_GATEWAY_LISTENER_SHARDS, MAX_GATEWAY_LISTEN_BACKLOG, MIN_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, MIN_GATEWAY_LISTEN_BACKLOG, }; use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig}; use aether_gateway::AppState; + use bytes::Bytes; + use clap::Parser; + use http_body_util::{BodyExt, Full}; + use hyper::body::Incoming as HyperIncoming; + use hyper::{Request as HyperRequest, Response as HyperResponse}; + use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; + use hyper_util::server::conn::auto::Builder as HyperServerBuilder; + use std::convert::Infallible; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; fn test_args() -> Args { Args { @@ -2566,13 +3496,17 @@ mod tests { listen_backlog: DEFAULT_GATEWAY_LISTEN_BACKLOG, listener_shards: DEFAULT_GATEWAY_LISTENER_SHARDS, http2_max_concurrent_streams: DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, + http_header_read_timeout_ms: DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS, + http_header_max_bytes: DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES, + http_max_headers: DEFAULT_GATEWAY_HTTP_MAX_HEADERS, healthcheck: false, healthcheck_timeout_ms: 3_000, deployment_topology: DeploymentTopologyArg::SingleNode, node_role: NodeRoleArg::All, migrate: false, apply_backfills: false, - auto_prepare_database: false, + database_mode: None, + auto_prepare_database: None, static_dir: None, video_task_truth_source_mode: VideoTaskTruthSourceArg::PythonSyncReport, video_task_poller_interval_ms: 5_000, @@ -2642,7 +3576,7 @@ mod tests { rate_limit: GatewayRateLimitArgs { bucket_seconds: 60, key_ttl_seconds: 120, - fail_open: true, + fail_open: false, }, logging: GatewayLoggingArgs { log_format: GatewayLogFormatArg::Pretty, @@ -2674,6 +3608,19 @@ mod tests { .expect("test database config should build") } + fn temporary_sqlite_args(label: &str) -> (Args, std::path::PathBuf) { + let mut args = test_args(); + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("clock should be available") + .as_nanos(); + let database_path = + std::env::temp_dir().join(format!("aether-{label}-{}-{nonce}.db", std::process::id())); + args.data.database_driver = Some(DatabaseDriverArg::Sqlite); + args.data.database_url = Some(format!("sqlite://{}", database_path.display())); + (args, database_path) + } + #[test] fn resolves_healthcheck_url_from_app_port() { assert_eq!( @@ -2741,6 +3688,40 @@ mod tests { ); } + #[test] + fn clamps_http_header_security_settings_without_touching_stream_concurrency() { + assert_eq!( + super::gateway_http_header_read_timeout_ms(0), + super::MIN_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS + ); + assert_eq!( + super::gateway_http_header_read_timeout_ms(u64::MAX), + super::MAX_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS + ); + assert_eq!( + super::gateway_http_header_max_bytes(1), + super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES + ); + assert_eq!( + super::gateway_http_header_max_bytes(usize::MAX), + super::MAX_GATEWAY_HTTP_HEADER_MAX_BYTES + ); + assert_eq!( + super::gateway_http_max_headers(0), + super::MIN_GATEWAY_HTTP_MAX_HEADERS + ); + assert_eq!( + super::gateway_http_max_headers(usize::MAX), + super::MAX_GATEWAY_HTTP_MAX_HEADERS + ); + assert_eq!( + super::gateway_http2_max_concurrent_streams( + super::DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS + ), + 16_384 + ); + } + #[test] fn auto_gateway_request_concurrency_scales_and_clamps() { assert_eq!( @@ -2785,11 +3766,14 @@ mod tests { } #[test] - fn normal_runtime_config_keeps_gateway_only_logs() { + fn normal_runtime_config_includes_database_lifecycle_logs() { let config = test_args() .runtime_config() .expect("runtime config should build"); - assert_eq!(config.default_log_filter, "aether_gateway=info"); + assert_eq!( + config.default_log_filter, + "aether_gateway=info,aether_data=info" + ); } #[test] @@ -2806,7 +3790,7 @@ mod tests { #[test] fn auto_prepare_database_runtime_config_enables_data_logs() { let mut args = test_args(); - args.auto_prepare_database = true; + args.auto_prepare_database = Some(true); let config = args.runtime_config().expect("runtime config should build"); assert_eq!( config.default_log_filter, @@ -2814,6 +3798,87 @@ mod tests { ); } + #[test] + fn database_mode_defaults_to_auto_and_preserves_legacy_false() { + assert_eq!(resolve_database_mode(None, None), DatabaseModeArg::Auto); + assert_eq!( + resolve_database_mode(None, Some(false)), + DatabaseModeArg::VerifyOnly + ); + assert_eq!( + resolve_database_mode(Some(DatabaseModeArg::Auto), Some(false)), + DatabaseModeArg::Auto + ); + } + + #[test] + fn parses_database_commands_and_verify_only_mode() { + let status = Args::try_parse_from(["aether-gateway", "db", "status"]) + .expect("db status should parse"); + assert!(matches!( + status.command, + Some(DataCommand::Db(args)) + if matches!(args.command, DatabaseCommand::Status) + )); + + let verify_only = + Args::try_parse_from(["aether-gateway", "--database-mode", "verify-only"]) + .expect("verify-only mode should parse"); + assert_eq!( + verify_only.effective_database_mode(), + DatabaseModeArg::VerifyOnly + ); + + let legacy_false = + Args::try_parse_from(["aether-gateway", "--auto-prepare-database=false"]) + .expect("legacy false setting should parse"); + assert_eq!( + legacy_false.effective_database_mode(), + DatabaseModeArg::VerifyOnly + ); + + let prepare = Args::try_parse_from(["aether-gateway", "db", "prepare"]) + .expect("db prepare should parse"); + assert!(matches!( + prepare.command, + Some(DataCommand::Db(args)) + if matches!(args.command, DatabaseCommand::Prepare) + )); + } + + #[test] + fn database_arguments_are_global_for_database_commands() { + let before = Args::try_parse_from([ + "aether-gateway", + "--database-driver", + "sqlite", + "--database-url", + "sqlite:///tmp/before.db", + "db", + "status", + ]) + .expect("database arguments before db should parse"); + assert_eq!( + before.data.database_url.as_deref(), + Some("sqlite:///tmp/before.db") + ); + + let after = Args::try_parse_from([ + "aether-gateway", + "db", + "prepare", + "--database-driver", + "sqlite", + "--database-url", + "sqlite:///tmp/after.db", + ]) + .expect("database arguments after db prepare should parse"); + assert_eq!( + after.data.database_url.as_deref(), + Some("sqlite:///tmp/after.db") + ); + } + #[test] fn gateway_data_pool_auto_sizes_sqlite_to_single_connection() { let mut args = test_args(); @@ -3298,6 +4363,280 @@ mod tests { ); } + #[test] + fn gateway_data_encryption_key_rejects_weak_and_published_values() { + assert!(validate_gateway_data_encryption_key(None).is_ok()); + assert!( + validate_gateway_data_encryption_key(Some("0123456789abcdef0123456789abcdef")).is_ok() + ); + + for insecure in [ + "short-secret", + "change-this-to-another-secure-random-string", + "change-this-to-a-secure-random-string", + "dev-encryption-key-do-not-use-in-production", + ] { + assert!( + validate_gateway_data_encryption_key(Some(insecure)).is_err(), + "accepted insecure key: {insecure}" + ); + } + } + + #[test] + fn data_copy_requires_tls_for_remote_sql_but_preserves_loopback_compatibility() { + let remote_postgres = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres://user:pass@db.example/aether", + "source", + false, + ) + .expect("remote postgres config should build"); + assert!(remote_postgres.pool.require_ssl); + + let remote_mysql = copy_database_config( + DatabaseDriverArg::Mysql, + "mysql://user:pass@192.0.2.10:3306/aether", + "target", + false, + ) + .expect("remote mysql config should build"); + assert!(remote_mysql.pool.require_ssl); + + for url in [ + "postgres://user:pass@localhost/aether", + "postgres://user:pass@127.42.17.9/aether", + "postgres://user:pass@[::1]/aether", + "postgres://user:pass@[::ffff:127.0.0.1]/aether", + ] { + let config = copy_database_config(DatabaseDriverArg::Postgres, url, "source", false) + .expect("literal loopback config should build"); + assert!( + !config.pool.require_ssl, + "loopback URL unexpectedly requires TLS: {url}" + ); + } + + let explicitly_insecure = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres://user:pass@db.example/aether", + "source", + true, + ) + .expect("explicit insecure opt-out should build"); + assert!(!explicitly_insecure.pool.require_ssl); + + let sqlite = copy_database_config( + DatabaseDriverArg::Sqlite, + "sqlite://./data/aether.db", + "target", + false, + ) + .expect("sqlite config should build"); + assert!(!sqlite.pool.require_ssl); + } + + #[test] + fn data_copy_tls_policy_is_conservative_for_non_loopback_and_query_hosts() { + let query_remote = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres:///aether?host=db.example", + "source", + false, + ) + .expect("query-host postgres config should build"); + assert!(query_remote.pool.require_ssl); + + let authority_loopback_query_remote = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres://user:pass@localhost/aether?host=db.example", + "source", + false, + ) + .expect("query-host override config should build"); + assert!(authority_loopback_query_remote.pool.require_ssl); + + let authority_loopback_hostaddr_remote = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres://user:pass@localhost/aether?hostaddr=192.0.2.10", + "source", + false, + ) + .expect("hostaddr override config should build"); + assert!(authority_loopback_hostaddr_remote.pool.require_ssl); + + let query_loopback = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres://user:pass@db.example/aether?host=127.0.0.1", + "source", + false, + ) + .expect("query-loopback config should build"); + assert!(!query_loopback.pool.require_ssl); + + let hostless = copy_database_config( + DatabaseDriverArg::Postgres, + "postgres:///aether", + "source", + false, + ) + .expect("hostless postgres config should build"); + // SQLx resolves a hostless PostgreSQL URL to its local socket or + // localhost default; that path is safe to keep plaintext for local + // development, just like an explicit loopback URL. + assert!(!hostless.pool.require_ssl); + } + + #[test] + fn data_copy_cli_accepts_independent_insecure_opt_outs() { + let parsed = Args::try_parse_from([ + "aether-gateway", + "copy", + "--source-driver", + "postgres", + "--source-url", + "postgres://user:pass@db.example/aether", + "--source-allow-insecure", + "--target-driver", + "mysql", + "--target-url", + "mysql://user:pass@db.example/aether", + ]) + .expect("copy command should parse endpoint-specific TLS flags"); + + let Some(DataCommand::Copy(copy)) = parsed.command else { + panic!("expected copy command"); + }; + assert!(copy.source_allow_insecure); + assert!(!copy.target_allow_insecure); + } + + #[cfg(unix)] + #[test] + fn database_export_output_is_private_atomic_and_no_clobber_by_default() { + use std::os::unix::fs::{symlink, MetadataExt, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-data-export-output-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&root).unwrap(); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700)).unwrap(); + let output = root.join("export.jsonl"); + + write_atomic_private_export(&output, b"first\n", false).unwrap(); + let metadata = std::fs::symlink_metadata(&output).unwrap(); + assert_eq!(metadata.mode() & 0o777, 0o600); + assert_eq!(metadata.nlink(), 1); + assert!(write_atomic_private_export(&output, b"second\n", false).is_err()); + assert_eq!(std::fs::read(&output).unwrap(), b"first\n"); + + write_atomic_private_export(&output, b"second\n", true).unwrap(); + assert_eq!(std::fs::read(&output).unwrap(), b"second\n"); + + let victim = root.join("victim"); + std::fs::write(&victim, b"known-good").unwrap(); + std::fs::remove_file(&output).unwrap(); + symlink(&victim, &output).unwrap(); + assert!(write_atomic_private_export(&output, b"replace\n", true).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_file(&output).unwrap(); + std::fs::hard_link(&victim, &output).unwrap(); + assert!(write_atomic_private_export(&output, b"replace\n", true).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn database_export_rejects_an_unsafe_symbolic_link_parent_target() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-data-export-parent-link-test-{}", + uuid::Uuid::new_v4() + )); + let real_parent = root.join("real"); + let linked_parent = root.join("linked"); + std::fs::create_dir_all(&real_parent).unwrap(); + // A symlink into a directory writable by other users must not become + // an escape hatch for the private export. + std::fs::set_permissions(&real_parent, std::fs::Permissions::from_mode(0o777)).unwrap(); + symlink(&real_parent, &linked_parent).unwrap(); + + let output = linked_parent.join("export.jsonl"); + assert!(write_atomic_private_export(&output, b"must not be written", false).is_err()); + assert!(!real_parent.join("export.jsonl").exists()); + + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn database_import_input_is_bounded_and_rejects_final_symlinks() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-data-import-input-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&root).unwrap(); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700)).unwrap(); + + let input = root.join("input.jsonl"); + std::fs::write(&input, b"0123456789").unwrap(); + let error = read_data_import_input_with_limit(&input, 4) + .expect_err("an input larger than the configured limit must be rejected"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert!(error.to_string().contains("4 byte limit")); + + let target = root.join("target.jsonl"); + std::fs::write(&target, b"safe\n").unwrap(); + let link = root.join("input-link.jsonl"); + symlink(&target, &link).unwrap(); + let error = read_data_import_input_with_limit(&link, 1024) + .expect_err("a final symbolic link must not be followed"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert!(error.to_string().contains("must not be a symbolic link")); + + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn database_import_input_rejects_fifo_without_blocking() { + use std::ffi::CString; + use std::os::unix::ffi::OsStrExt; + use std::os::unix::fs::PermissionsExt; + + let root = std::env::temp_dir().join(format!( + "aether-data-import-fifo-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&root).unwrap(); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700)).unwrap(); + let fifo = root.join("input.fifo"); + let fifo_name = CString::new(fifo.as_os_str().as_bytes()).unwrap(); + let result = unsafe { libc::mkfifo(fifo_name.as_ptr(), 0o600) }; + assert_eq!( + result, + 0, + "mkfifo failed: {}", + std::io::Error::last_os_error() + ); + + // O_NONBLOCK in open_data_import_file means this call reaches the + // regular-file check immediately even when no FIFO writer exists. + let error = read_data_import_input_with_limit(&fifo, 1024) + .expect_err("FIFO input must be rejected before reading"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert!(error.to_string().contains("regular file")); + + std::fs::remove_dir_all(root).unwrap(); + } + #[test] fn redis_runtime_config_owns_redis_connection() { let mut args = test_args(); @@ -3484,17 +4823,17 @@ mod tests { } #[test] - fn pending_schema_error_mentions_explicit_migrate_command() { + fn pending_schema_error_mentions_database_prepare_command() { let error = pending_schema_error(2, 20260413020000, "squash usage schema split"); let message = error.to_string(); assert!(message.contains("database schema is behind by 2 migration(s)")); assert!(message.contains("20260413020000")); assert!(message.contains("squash usage schema split")); - assert!(message.contains("aether-gateway --migrate")); + assert!(message.contains("aether-gateway db prepare")); } #[test] - fn pending_backfills_error_mentions_explicit_apply_backfills_command() { + fn pending_backfills_error_mentions_database_prepare_command() { let message = pending_backfills_error( 1, 20260422110000, @@ -3504,7 +4843,7 @@ mod tests { assert!(message.contains("database backfills are behind by 1 backfill(s)")); assert!(message.contains("20260422110000")); assert!(message.contains("backfill stats aggregate read path support")); - assert!(message.contains("aether-gateway --apply-backfills")); + assert!(message.contains("aether-gateway db prepare")); assert!(message.contains("before starting the service")); } @@ -3527,11 +4866,79 @@ mod tests { #[tokio::test] async fn auto_prepare_database_is_noop_without_database_pool() { let state = AppState::new().expect("state should build"); - super::prepare_database_startup_requirements(&state, true) + super::prepare_database_startup_requirements(&state, DatabaseModeArg::Auto) .await .expect("disabled data backend should not block startup"); } + #[tokio::test] + async fn verify_only_does_not_prepare_fresh_sqlite_database() { + let (args, database_path) = temporary_sqlite_args("verify-only"); + let state = AppState::new() + .expect("state should build") + .with_data_config(args.data.to_config()) + .expect("sqlite state should build"); + let pending_before = state + .pending_database_migrations() + .await + .expect("pending migrations should load") + .expect("sqlite should expose migration state"); + assert!(!pending_before.is_empty()); + + let error = + super::prepare_database_startup_requirements(&state, DatabaseModeArg::VerifyOnly) + .await + .expect_err("verify-only should reject a fresh database"); + assert!(error.to_string().contains("aether-gateway db prepare")); + + let pending_after = state + .pending_database_migrations() + .await + .expect("pending migrations should reload") + .expect("sqlite should expose migration state"); + assert_eq!(pending_after, pending_before); + drop(state); + let _ = std::fs::remove_file(database_path); + } + + #[tokio::test] + async fn auto_mode_prepares_fresh_sqlite_database() { + let (args, database_path) = temporary_sqlite_args("auto-prepare"); + let state = AppState::new() + .expect("state should build") + .with_data_config(args.data.to_config()) + .expect("sqlite state should build"); + + super::prepare_database_startup_requirements(&state, DatabaseModeArg::Auto) + .await + .expect("auto mode should prepare a fresh database"); + assert!(state + .pending_database_migrations() + .await + .expect("pending migrations should load") + .expect("sqlite should expose migration state") + .is_empty()); + assert!(state + .pending_database_backfills() + .await + .expect("pending backfills should load") + .expect("sqlite should expose backfill state") + .is_empty()); + drop(state); + let _ = std::fs::remove_file(database_path); + } + + #[tokio::test] + async fn database_prepare_requires_database_url() { + let data = test_args().data; + let error = super::run_database_prepare(&data) + .await + .expect_err("missing database URL should fail"); + assert!(error + .to_string() + .contains("AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL")); + } + #[tokio::test] async fn explicit_migrate_requires_database_url() { let args = test_args(); @@ -3586,4 +4993,174 @@ mod tests { .expect("sqlite backfills should be an explicit no-op"); let _ = std::fs::remove_file(database_path); } + + #[tokio::test] + async fn first_request_gate_closes_a_connection_that_never_reaches_service() { + let gate = super::GatewayFirstRequestGate::new(); + let result = tokio::time::timeout( + std::time::Duration::from_secs(1), + super::drive_gateway_connection( + std::future::pending::>(), + gate, + std::time::Duration::from_millis(5), + ), + ) + .await + .expect("first-request deadline should fire promptly"); + assert!(result.is_ok()); + } + + #[tokio::test] + async fn first_request_gate_does_not_deadline_a_streaming_connection() { + let gate = super::GatewayFirstRequestGate::new(); + gate.mark_seen(); + // Model a body that remains active after request headers arrive. + let connection = async { + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + Ok::<(), ()>(()) + }; + let result = tokio::time::timeout(std::time::Duration::from_secs(1), async move { + super::drive_gateway_connection(connection, gate, std::time::Duration::from_millis(5)) + .await + }) + .await + .expect("streaming connection should finish"); + assert!(result.is_ok()); + } + + #[tokio::test] + async fn first_request_deadline_covers_partial_http1_and_h2_preface() { + let prefixes: &[&[u8]] = &[ + b"G", + b"GET / HTTP/1.1\r\nHost: localhost\r\n", + b"PRI * HTTP/2.0\r\n\r\nSM\r\n", + ]; + for prefix in prefixes { + let (mut client, server) = tokio::io::duplex(16 * 1024); + client + .write_all(prefix) + .await + .expect("fixture prefix should be writable"); + let gate = super::GatewayFirstRequestGate::new(); + let service = tower::service_fn(|_request: HyperRequest| async { + Ok::<_, Infallible>(HyperResponse::new(Full::new(Bytes::from_static(b"ok")))) + }); + let mut builder = HyperServerBuilder::new(TokioExecutor::new()); + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(std::time::Duration::from_millis(10)) + .max_buf_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES) + .max_headers(super::MIN_GATEWAY_HTTP_MAX_HEADERS); + builder + .http2() + .timer(TokioTimer::new()) + .max_concurrent_streams(super::DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS) + .max_header_list_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES as u32); + + let result = tokio::time::timeout( + std::time::Duration::from_secs(1), + super::drive_gateway_connection( + builder.serve_connection_with_upgrades( + TokioIo::new(server), + super::TowerToHyperService::new(super::GatewayFirstRequestService { + inner: service, + gate: gate.clone(), + }), + ), + gate, + std::time::Duration::from_millis(5), + ), + ) + .await + .expect("partial protocol input should hit the first-request deadline"); + assert!( + result.is_ok(), + "deadline should close without a parser error" + ); + + let mut byte = [0u8; 1]; + let read = + tokio::time::timeout(std::time::Duration::from_secs(1), client.read(&mut byte)) + .await + .expect("timed-out connection should close its peer"); + assert!(matches!(read, Ok(0) | Err(_))); + } + } + + #[tokio::test] + async fn first_request_deadline_does_not_cut_off_a_delayed_http1_body() { + let (mut client, server) = tokio::io::duplex(16 * 1024); + let gate = super::GatewayFirstRequestGate::new(); + let (headers_seen_tx, mut headers_seen_rx) = tokio::sync::mpsc::unbounded_channel(); + let service = tower::service_fn(move |request: HyperRequest| { + let _ = headers_seen_tx.send(()); + async move { + let body = request + .into_body() + .collect() + .await + .expect("test body should decode") + .to_bytes(); + Ok::<_, Infallible>(HyperResponse::new(Full::new(body))) + } + }); + let mut builder = HyperServerBuilder::new(TokioExecutor::new()); + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(std::time::Duration::from_millis(20)) + .max_buf_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES) + .max_headers(super::MIN_GATEWAY_HTTP_MAX_HEADERS); + builder + .http2() + .timer(TokioTimer::new()) + .max_concurrent_streams(super::DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS) + .max_header_list_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES as u32); + + let server_task = tokio::spawn(async move { + super::drive_gateway_connection( + builder.serve_connection_with_upgrades( + TokioIo::new(server), + super::TowerToHyperService::new(super::GatewayFirstRequestService { + inner: service, + gate: gate.clone(), + }), + ), + gate, + std::time::Duration::from_millis(20), + ) + .await + }); + client + .write_all( + b"POST / HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\nContent-Length: 5\r\n\r\n", + ) + .await + .expect("request headers should be writable"); + tokio::time::timeout(std::time::Duration::from_secs(1), headers_seen_rx.recv()) + .await + .expect("request headers should reach the service") + .expect("service notification should remain available"); + tokio::time::sleep(std::time::Duration::from_millis(40)).await; + client + .write_all(b"hello") + .await + .expect("body should remain writable after the header deadline"); + let mut response = Vec::new(); + tokio::time::timeout( + std::time::Duration::from_secs(1), + client.read_to_end(&mut response), + ) + .await + .expect("streaming response should complete") + .expect("response should be readable"); + let result = server_task + .await + .expect("server connection task should join"); + assert!(result.is_ok()); + assert!(response + .windows(b"hello".len()) + .any(|window| window == b"hello")); + } } diff --git a/apps/aether-gateway/src/maintenance/mod.rs b/apps/aether-gateway/src/maintenance/mod.rs index 5cb50439b..868fafe0a 100644 --- a/apps/aether-gateway/src/maintenance/mod.rs +++ b/apps/aether-gateway/src/maintenance/mod.rs @@ -9,6 +9,7 @@ pub(crate) use runtime::{ perform_oauth_token_refresh_once, perform_pool_quota_probe_once, perform_provider_checkin_once, perform_provider_quota_alert_once, pool_quota_probe_target_count, preview_manual_usage_cleanup, rebuild_admin_stats_once, record_completed_cleanup_run, record_proxy_upgrade_traffic_success, + record_proxy_upgrade_traffic_success_for_generation, restore_proxy_upgrade_rollout_skipped_nodes, retry_proxy_upgrade_rollout_node, run_admin_system_cleanup_once, run_manual_usage_cleanup_once, skip_proxy_upgrade_rollout_node, spawn_account_self_check_worker, spawn_audit_cleanup_worker, spawn_db_maintenance_worker, diff --git a/apps/aether-gateway/src/maintenance/runtime.rs b/apps/aether-gateway/src/maintenance/runtime.rs index b513c2a82..66d104b65 100644 --- a/apps/aether-gateway/src/maintenance/runtime.rs +++ b/apps/aether-gateway/src/maintenance/runtime.rs @@ -101,12 +101,13 @@ use proxy_upgrade_rollout::*; pub(crate) use proxy_upgrade_rollout::{ cancel_proxy_upgrade_rollout, clear_proxy_upgrade_rollout_conflicts, collect_proxy_upgrade_rollout_probes, inspect_proxy_upgrade_rollout, - record_proxy_upgrade_traffic_success, restore_proxy_upgrade_rollout_skipped_nodes, - retry_proxy_upgrade_rollout_node, skip_proxy_upgrade_rollout_node, start_proxy_upgrade_rollout, - ProxyUpgradeRolloutCancelSummary, ProxyUpgradeRolloutConflictClearSummary, - ProxyUpgradeRolloutNodeActionSummary, ProxyUpgradeRolloutPendingProbe, - ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary, - ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutSummary, ProxyUpgradeRolloutTrackedNodeState, + record_proxy_upgrade_traffic_success, record_proxy_upgrade_traffic_success_for_generation, + restore_proxy_upgrade_rollout_skipped_nodes, retry_proxy_upgrade_rollout_node, + skip_proxy_upgrade_rollout_node, start_proxy_upgrade_rollout, ProxyUpgradeRolloutCancelSummary, + ProxyUpgradeRolloutConflictClearSummary, ProxyUpgradeRolloutNodeActionSummary, + ProxyUpgradeRolloutPendingProbe, ProxyUpgradeRolloutProbeConfig, + ProxyUpgradeRolloutSkippedRestoreSummary, ProxyUpgradeRolloutStatus, + ProxyUpgradeRolloutSummary, ProxyUpgradeRolloutTrackedNodeState, }; use request_candidate_cleanup::*; use runners::*; diff --git a/apps/aether-gateway/src/maintenance/runtime/account_self_check.rs b/apps/aether-gateway/src/maintenance/runtime/account_self_check.rs index e16a63202..55feeba08 100644 --- a/apps/aether-gateway/src/maintenance/runtime/account_self_check.rs +++ b/apps/aether-gateway/src/maintenance/runtime/account_self_check.rs @@ -92,22 +92,19 @@ impl AccountSelfCheckWorkerConfig { enum AccountSelfCheckOutcome { Success { status_code: Option, - message: Option, }, Blocked { status_code: Option, - message: String, }, AutoRemoved { status_code: Option, - message: String, }, Failed { status_code: Option, - message: String, + category: &'static str, }, Skipped { - message: String, + category: &'static str, }, } @@ -132,13 +129,12 @@ impl AccountSelfCheckOutcome { } } - fn message(&self) -> Option<&str> { + fn category(&self) -> &'static str { match self { - Self::Success { message, .. } => message.as_deref(), - Self::Blocked { message, .. } - | Self::AutoRemoved { message, .. } - | Self::Failed { message, .. } - | Self::Skipped { message, .. } => Some(message.as_str()), + Self::Success { .. } => "quota_refresh_succeeded", + Self::Blocked { .. } => "account_blocked", + Self::AutoRemoved { .. } => "account_auto_removed", + Self::Failed { category, .. } | Self::Skipped { category } => category, } } } @@ -370,13 +366,13 @@ fn quota_payload_result_for_key(key_id: &str, payload: Option) -> Account let Some(payload) = payload else { return AccountSelfCheckOutcome::Failed { status_code: None, - message: "quota refresh returned no payload".to_string(), + category: "missing_payload", }; }; let Some(results) = payload.get("results").and_then(Value::as_array) else { return AccountSelfCheckOutcome::Failed { status_code: None, - message: "quota refresh returned no result list".to_string(), + category: "invalid_payload", }; }; let Some(item) = results.iter().find(|item| { @@ -386,7 +382,7 @@ fn quota_payload_result_for_key(key_id: &str, payload: Option) -> Account }) else { return AccountSelfCheckOutcome::Failed { status_code: None, - message: "quota refresh result missing key".to_string(), + category: "missing_key_result", }; }; @@ -404,34 +400,34 @@ fn quota_payload_result_for_key(key_id: &str, payload: Option) -> Account .get("message") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned); + .filter(|value| !value.is_empty()); let auto_removed = item .get("auto_removed") .and_then(Value::as_bool) .unwrap_or(false); if status == "success" { - return AccountSelfCheckOutcome::Success { - status_code, - message, - }; + return AccountSelfCheckOutcome::Success { status_code }; } if auto_removed { - return AccountSelfCheckOutcome::AutoRemoved { - status_code, - message: message.unwrap_or_else(|| "已自动删除".to_string()), - }; + return AccountSelfCheckOutcome::AutoRemoved { status_code }; } - if quota_result_status_is_blocked(&status, status_code, message.as_deref()) { - return AccountSelfCheckOutcome::Blocked { - status_code, - message: message.unwrap_or_else(|| status.clone()), - }; + if quota_result_status_is_blocked(&status, status_code, message) { + return AccountSelfCheckOutcome::Blocked { status_code }; } AccountSelfCheckOutcome::Failed { status_code, - message: message.unwrap_or_else(|| status.clone()), + category: quota_refresh_failure_category(&status, status_code), + } +} + +fn quota_refresh_failure_category(status: &str, status_code: Option) -> &'static str { + match status_code { + Some(408 | 504) => "upstream_timed_out", + Some(429) => "rate_limited", + Some(500..=599) => "upstream_error", + _ if matches!(status, "unsupported" | "not_supported") => "unsupported", + _ => "quota_refresh_failed", } } @@ -546,21 +542,7 @@ async fn record_score_probe_result_for_key( succeeded, hard_state, probe_status, - score_reason_patch: Some(json!({ - "last_probe": { - "source": "account_self_check", - "status": outcome.score_status(), - "status_code": outcome.status_code(), - "message": outcome.message() - }, - "last_self_check": { - "source": "account_self_check", - "status": outcome.score_status(), - "status_code": outcome.status_code(), - "message": outcome.message(), - "attempted_at": attempted_at - } - })), + score_reason_patch: Some(score_reason_patch_for_outcome(outcome, attempted_at)), }; if let Err(err) = state.data.record_pool_member_probe_result(result).await { debug!( @@ -572,6 +554,24 @@ async fn record_score_probe_result_for_key( } } +fn score_reason_patch_for_outcome(outcome: &AccountSelfCheckOutcome, attempted_at: u64) -> Value { + json!({ + "last_probe": { + "source": "account_self_check", + "status": outcome.score_status(), + "status_code": outcome.status_code(), + "category": outcome.category() + }, + "last_self_check": { + "source": "account_self_check", + "status": outcome.score_status(), + "status_code": outcome.status_code(), + "category": outcome.category(), + "attempted_at": attempted_at + } + }) +} + fn endpoint_for_self_check( provider_type: &str, endpoints: &[StoredProviderCatalogEndpoint], @@ -579,8 +579,21 @@ fn endpoint_for_self_check( provider_quota_refresh_endpoint_for_provider(provider_type, endpoints, true) } -fn gateway_error_message(err: GatewayError) -> String { - err.into_message() +fn gateway_error_category(err: &GatewayError) -> &'static str { + match err { + GatewayError::UpstreamUnavailable { .. } => "upstream_unavailable", + GatewayError::ControlUnavailable { .. } => "control_unavailable", + GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timed_out", + GatewayError::AdmissionTimeout { .. } => "admission_timed_out", + GatewayError::Client { status, .. } if status.as_u16() == 429 => "rate_limited", + GatewayError::Client { status, .. } if status.is_server_error() => "upstream_error", + GatewayError::Client { .. } => "request_rejected", + GatewayError::PlanUsageLimited(_) => "plan_usage_limited", + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + "operation_rejected" + } + GatewayError::Internal(_) => "internal_error", + } } fn update_summary_from_outcome( @@ -727,7 +740,7 @@ pub(crate) async fn perform_account_self_check_once_with_config( Ok(outcome) => outcome, Err(err) => AccountSelfCheckOutcome::Failed { status_code: None, - message: gateway_error_message(err), + category: gateway_error_category(&err), }, }; record_score_probe_result_for_key(state, &provider.id, &key.id, now_ts, &outcome).await; @@ -797,7 +810,12 @@ pub(crate) fn spawn_account_self_check_worker( #[cfg(test)] mod tests { - use super::select_account_self_check_key_ids; + use super::{ + gateway_error_category, quota_payload_result_for_key, score_reason_patch_for_outcome, + select_account_self_check_key_ids, + }; + use crate::GatewayError; + use serde_json::json; use std::collections::BTreeMap; #[test] @@ -813,4 +831,38 @@ mod tests { assert_eq!(selected, vec!["never".to_string(), "stale".to_string()]); } + + #[test] + fn account_self_check_patch_does_not_persist_upstream_message() { + let secret = "Authorization: Bearer quota-secret from /srv/aether/credentials.json"; + let outcome = quota_payload_result_for_key( + "key-secret-regression", + Some(json!({ + "results": [{ + "key_id": "key-secret-regression", + "status": "error", + "status_code": 502, + "message": secret + }] + })), + ); + + let patch = score_reason_patch_for_outcome(&outcome, 1_777_000_000); + let serialized = patch.to_string(); + + assert_eq!(patch["last_self_check"]["category"], "upstream_error"); + assert!(patch["last_self_check"].get("message").is_none()); + assert!(!serialized.contains("quota-secret")); + assert!(!serialized.contains("/srv/aether/credentials.json")); + } + + #[test] + fn account_self_check_gateway_error_category_drops_internal_details() { + let error = GatewayError::Internal( + "postgresql://admin:database-secret@db.internal/aether".to_string(), + ); + + assert_eq!(gateway_error_category(&error), "internal_error"); + assert!(!gateway_error_category(&error).contains("database-secret")); + } } diff --git a/apps/aether-gateway/src/maintenance/runtime/cleanup_runs.rs b/apps/aether-gateway/src/maintenance/runtime/cleanup_runs.rs index 437e0d5e1..7fecc768f 100644 --- a/apps/aether-gateway/src/maintenance/runtime/cleanup_runs.rs +++ b/apps/aether-gateway/src/maintenance/runtime/cleanup_runs.rs @@ -15,6 +15,15 @@ const CLEANUP_RUN_HISTORY_KEY: &str = "admin_cleanup_run_history"; const CLEANUP_RUN_HISTORY_LIMIT: usize = 50; const REQUEST_BODY_PROGRESS_UPDATE_BATCHES: usize = 10; +const CLEANUP_ERROR_INVALID_CONFIGURATION: &str = "invalid_configuration"; +const CLEANUP_ERROR_INVALID_INPUT: &str = "invalid_input"; +const CLEANUP_ERROR_POSTGRES: &str = "postgres"; +const CLEANUP_ERROR_REDIS: &str = "redis"; +const CLEANUP_ERROR_SQL: &str = "sql"; +const CLEANUP_ERROR_TIMED_OUT: &str = "timed_out"; +const CLEANUP_ERROR_UNEXPECTED_VALUE: &str = "unexpected_value"; +const CLEANUP_ERROR_INTERNAL: &str = "internal_error"; + pub(crate) const USAGE_CLEANUP_KIND: &str = "usage_cleanup"; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub(crate) struct AdminCleanupRunRecord { @@ -30,6 +39,73 @@ pub(crate) struct AdminCleanupRunRecord { pub(crate) error: Option, } +pub(super) fn cleanup_data_layer_error_category(error: &DataLayerError) -> &'static str { + match error { + DataLayerError::InvalidConfiguration(_) => CLEANUP_ERROR_INVALID_CONFIGURATION, + DataLayerError::InvalidInput(_) => CLEANUP_ERROR_INVALID_INPUT, + DataLayerError::Postgres(_) => CLEANUP_ERROR_POSTGRES, + DataLayerError::Redis(_) => CLEANUP_ERROR_REDIS, + DataLayerError::Sql(_) => CLEANUP_ERROR_SQL, + DataLayerError::TimedOut(_) => CLEANUP_ERROR_TIMED_OUT, + DataLayerError::UnexpectedValue(_) => CLEANUP_ERROR_UNEXPECTED_VALUE, + } +} + +fn cleanup_gateway_error_category(error: &GatewayError) -> &'static str { + match error { + GatewayError::UpstreamUnavailable { .. } => "upstream_unavailable", + GatewayError::ControlUnavailable { .. } => "control_unavailable", + GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timed_out", + GatewayError::AdmissionTimeout { .. } => "admission_timed_out", + GatewayError::Client { status, .. } if status.is_server_error() => "upstream_error", + GatewayError::Client { .. } => "request_rejected", + GatewayError::PlanUsageLimited(_) => "plan_usage_limited", + GatewayError::LastActiveAdminUpdateDenied | GatewayError::LastActiveAdminDeleteDenied => { + "operation_rejected" + } + GatewayError::Internal(_) => CLEANUP_ERROR_INTERNAL, + } +} + +fn normalize_stored_cleanup_error(error: &str) -> &'static str { + let error = error.trim(); + match error { + CLEANUP_ERROR_INVALID_CONFIGURATION => CLEANUP_ERROR_INVALID_CONFIGURATION, + CLEANUP_ERROR_INVALID_INPUT => CLEANUP_ERROR_INVALID_INPUT, + CLEANUP_ERROR_POSTGRES => CLEANUP_ERROR_POSTGRES, + CLEANUP_ERROR_REDIS => CLEANUP_ERROR_REDIS, + CLEANUP_ERROR_SQL => CLEANUP_ERROR_SQL, + CLEANUP_ERROR_TIMED_OUT => CLEANUP_ERROR_TIMED_OUT, + CLEANUP_ERROR_UNEXPECTED_VALUE => CLEANUP_ERROR_UNEXPECTED_VALUE, + CLEANUP_ERROR_INTERNAL => CLEANUP_ERROR_INTERNAL, + "upstream_unavailable" => "upstream_unavailable", + "control_unavailable" => "control_unavailable", + "planning_timed_out" => "planning_timed_out", + "admission_timed_out" => "admission_timed_out", + "upstream_error" => "upstream_error", + "request_rejected" => "request_rejected", + "plan_usage_limited" => "plan_usage_limited", + "operation_rejected" => "operation_rejected", + _ if error.starts_with("invalid configuration:") => CLEANUP_ERROR_INVALID_CONFIGURATION, + _ if error.starts_with("invalid input:") => CLEANUP_ERROR_INVALID_INPUT, + _ if error.starts_with("postgres error:") => CLEANUP_ERROR_POSTGRES, + _ if error.starts_with("redis error:") => CLEANUP_ERROR_REDIS, + _ if error.starts_with("sql error:") => CLEANUP_ERROR_SQL, + _ if error.starts_with("operation timed out:") => CLEANUP_ERROR_TIMED_OUT, + _ if error.starts_with("unexpected database value:") => CLEANUP_ERROR_UNEXPECTED_VALUE, + _ => CLEANUP_ERROR_INTERNAL, + } +} + +fn normalize_cleanup_run_record(mut record: AdminCleanupRunRecord) -> AdminCleanupRunRecord { + record.error = record + .error + .as_deref() + .map(normalize_stored_cleanup_error) + .map(str::to_string); + record +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum AdminCleanupTaskKind { Config, @@ -225,7 +301,7 @@ pub(crate) async fn record_failed_cleanup_run( .unwrap_or(u64::MAX), ), summary: json!({}), - error: Some(error.to_string()), + error: Some(cleanup_data_layer_error_category(error).to_string()), }; if let Err(err) = record_cleanup_run(data, record).await { warn!(error = %err, kind, "failed to record failed cleanup run"); @@ -277,7 +353,7 @@ async fn run_request_body_cleanup_task( batch_size, &total, Some(started_at), - Some(err.to_string()), + Some(cleanup_data_layer_error_category(&err).to_string()), ); if let Err(record_err) = record_cleanup_run(&data, failed).await { warn!(error = %record_err, "failed to record request body cleanup failure"); @@ -350,7 +426,7 @@ async fn run_admin_system_purge_task( kind.failure_message().to_string(), json!({}), Some(started_at), - Some(format!("{err:?}")), + Some(cleanup_gateway_error_category(&err).to_string()), ); if let Err(record_err) = record_cleanup_run(&data, failed).await { warn!(error = %record_err, "failed to record admin system purge task failure"); @@ -493,6 +569,7 @@ pub(crate) async fn record_admin_cleanup_run( data: &GatewayDataState, record: AdminCleanupRunRecord, ) -> Result<(), DataLayerError> { + let record = normalize_cleanup_run_record(record); let mut records = list_admin_cleanup_run_records(data).await?; records.retain(|existing| existing.id != record.id); records.insert(0, record); @@ -522,5 +599,50 @@ fn parse_cleanup_run_records(value: Value) -> Vec { .into_iter() .flat_map(|items| items.iter()) .filter_map(|item| serde_json::from_value::(item.clone()).ok()) + .map(normalize_cleanup_run_record) .collect() } + +#[cfg(test)] +mod tests { + use super::{ + cleanup_data_layer_error_category, parse_cleanup_run_records, AdminCleanupRunRecord, + }; + use aether_data_contracts::DataLayerError; + use serde_json::json; + + #[test] + fn cleanup_error_categories_do_not_include_data_layer_details() { + let error = DataLayerError::Postgres( + "connection failed for postgresql://admin:database-secret@db.internal/aether" + .to_string(), + ); + + assert_eq!(cleanup_data_layer_error_category(&error), "postgres"); + assert!(!cleanup_data_layer_error_category(&error).contains("database-secret")); + } + + #[test] + fn cleanup_history_read_projection_removes_legacy_error_details() { + let secret = "postgres error: password=database-secret path=/srv/aether/private.db"; + let stored = AdminCleanupRunRecord { + id: "cleanup-secret-regression".to_string(), + kind: "usage_cleanup".to_string(), + trigger: "manual".to_string(), + status: "failed".to_string(), + message: "请求记录手动清理失败".to_string(), + started_at_unix_secs: 1, + completed_at_unix_secs: Some(2), + duration_ms: Some(10), + summary: json!({}), + error: Some(secret.to_string()), + }; + + let records = parse_cleanup_run_records(json!([stored])); + let serialized = serde_json::to_string(&records).expect("cleanup records should serialize"); + + assert_eq!(records[0].error.as_deref(), Some("postgres")); + assert!(!serialized.contains("database-secret")); + assert!(!serialized.contains("/srv/aether/private.db")); + } +} diff --git a/apps/aether-gateway/src/maintenance/runtime/fixed_provider_reconciliation.rs b/apps/aether-gateway/src/maintenance/runtime/fixed_provider_reconciliation.rs index c1a38d466..106f7c967 100644 --- a/apps/aether-gateway/src/maintenance/runtime/fixed_provider_reconciliation.rs +++ b/apps/aether-gateway/src/maintenance/runtime/fixed_provider_reconciliation.rs @@ -289,7 +289,7 @@ mod tests { assert_eq!(responses.max_retries, Some(9)); assert_eq!( responses.proxy, - Some(json!({"url": "http://proxy.internal:8080"})) + Some(json!({"url": "http://proxy.internal:8080/"})) ); assert_eq!( responses @@ -312,7 +312,7 @@ mod tests { assert_eq!(live.max_retries, Some(6)); assert_eq!( live.proxy, - Some(json!({"url": "http://voice-proxy.internal:8080"})) + Some(json!({"url": "http://voice-proxy.internal:8080/"})) ); assert_eq!( live.config diff --git a/apps/aether-gateway/src/maintenance/runtime/oauth_token_refresh.rs b/apps/aether-gateway/src/maintenance/runtime/oauth_token_refresh.rs index 7ec147a7e..7eff1577e 100644 --- a/apps/aether-gateway/src/maintenance/runtime/oauth_token_refresh.rs +++ b/apps/aether-gateway/src/maintenance/runtime/oauth_token_refresh.rs @@ -39,7 +39,11 @@ pub(crate) async fn perform_oauth_token_refresh_once( return Ok(OAuthTokenRefreshRunSummary::default()); } - let providers = state.list_provider_catalog_providers(true).await?; + // Maintenance must not let one malformed historical proxy credential + // abort the scan for every provider. Read the rows first, then open each + // row in isolation so a bad record can be skipped while database errors + // and missing encryption configuration still fail closed. + let providers = read_oauth_maintenance_providers(state).await?; let provider_ids = providers .iter() .map(|provider| provider.id.clone()) @@ -48,12 +52,16 @@ pub(crate) async fn perform_oauth_token_refresh_once( return Ok(OAuthTokenRefreshRunSummary::default()); } - let endpoints = state - .list_provider_catalog_endpoints_by_provider_ids(&provider_ids) - .await?; + let endpoints = read_oauth_maintenance_endpoints(state, &provider_ids).await?; + // Read the catalog rows without opening/decrypting credentials in bulk. + // A single legacy/plaintext row must not abort refresh for every healthy + // key, and this maintenance scan must not trigger the normal lazy v2 + // credential rewrite path. Each candidate is opened in isolation below. let keys = state + .data .list_provider_catalog_keys_by_provider_ids(&provider_ids) - .await?; + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; let endpoints_by_provider = group_endpoints_by_provider(endpoints); let keys_by_provider = group_keys_by_provider(keys); let mut summary = OAuthTokenRefreshRunSummary::default(); @@ -85,12 +93,31 @@ pub(crate) async fn perform_oauth_token_refresh_once( continue; }; - let Some(transport) = state + let transport = match state .read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id) - .await? - else { - summary.skipped = summary.skipped.saturating_add(1); - continue; + .await + { + Ok(Some(transport)) => transport, + Ok(None) => { + summary.skipped = summary.skipped.saturating_add(1); + continue; + } + Err(err) if is_nonfatal_legacy_catalog_credential_error(&err) => { + // Keep malformed historical credentials untouched. They + // are intentionally skipped while other keys continue. + summary.skipped = summary.skipped.saturating_add(1); + warn!( + event_name = "oauth_token_refresh_skipped_invalid_credential", + log_type = "ops", + worker = "oauth_token_refresh", + provider_id = %provider.id, + key_id = %key.id, + reason = "invalid_stored_credential", + "gateway skipped oauth refresh for an invalid stored credential" + ); + continue; + } + Err(err) => return Err(err), }; let is_agent_identity = crate::provider_transport::is_codex_agent_identity_transport(&transport); @@ -128,7 +155,7 @@ pub(crate) async fn perform_oauth_token_refresh_once( Ok(None) => { summary.skipped = summary.skipped.saturating_add(1); } - Err(err) => { + Err(_) => { summary.failed = summary.failed.saturating_add(1); warn!( event_name = "oauth_token_refresh_failed", @@ -136,7 +163,6 @@ pub(crate) async fn perform_oauth_token_refresh_once( worker = "oauth_token_refresh", provider_id = %provider.id, key_id = %key.id, - error = ?err, "gateway oauth token auto refresh failed" ); } @@ -162,6 +188,81 @@ pub(crate) async fn perform_oauth_token_refresh_once( Ok(summary) } +async fn read_oauth_maintenance_providers( + state: &AppState, +) -> Result, GatewayError> { + let stored = state + .data + .list_provider_catalog_providers(true) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut opened = Vec::with_capacity(stored.len()); + for provider in stored { + let provider_id = provider.id.clone(); + match state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) + .await + { + Ok(mut rows) => { + if let Some(row) = rows.pop() { + opened.push(row); + } + } + Err(error) if is_nonfatal_stored_proxy_error(&error) => { + warn!( + event_name = "oauth_token_refresh_skipped_invalid_provider_proxy", + log_type = "ops", + worker = "oauth_token_refresh", + provider_id = %provider_id, + reason = "invalid_stored_proxy_credential", + "gateway skipped oauth refresh for a provider with an invalid stored proxy credential" + ); + } + Err(error) => return Err(error), + } + } + Ok(opened) +} + +async fn read_oauth_maintenance_endpoints( + state: &AppState, + provider_ids: &[String], +) -> Result, GatewayError> { + let stored = state + .data + .list_provider_catalog_endpoints_by_provider_ids(provider_ids) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut opened = Vec::with_capacity(stored.len()); + for endpoint in stored { + let provider_id = endpoint.provider_id.clone(); + let endpoint_id = endpoint.id.clone(); + match state + .read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id)) + .await + { + Ok(mut rows) => { + if let Some(row) = rows.pop() { + opened.push(row); + } + } + Err(error) if is_nonfatal_stored_proxy_error(&error) => { + warn!( + event_name = "oauth_token_refresh_skipped_invalid_endpoint_proxy", + log_type = "ops", + worker = "oauth_token_refresh", + provider_id = %provider_id, + endpoint_id = %endpoint_id, + reason = "invalid_stored_proxy_credential", + "gateway skipped oauth refresh for an endpoint with an invalid stored proxy credential" + ); + } + Err(error) => return Err(error), + } + } + Ok(opened) +} + fn group_endpoints_by_provider( endpoints: Vec, ) -> BTreeMap> { @@ -237,8 +338,10 @@ async fn provider_key_credentials_changed( before: &StoredProviderCatalogKey, ) -> Result { let Some(after) = state + .data .list_provider_catalog_keys_by_ids(std::slice::from_ref(&before.id)) - .await? + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? .into_iter() .next() else { @@ -275,6 +378,58 @@ fn now_unix_secs() -> u64 { .unwrap_or_default() } +/// Credential decoding errors are expected for rows written by older +/// versions of the service. They are non-fatal for a best-effort maintenance +/// scan, but normal request/admin paths still fail closed on the same error. +fn is_nonfatal_legacy_catalog_credential_error(error: &GatewayError) -> bool { + is_nonfatal_legacy_provider_key_credential_error(error) || is_nonfatal_stored_proxy_error(error) +} + +fn is_nonfatal_legacy_provider_key_credential_error(error: &GatewayError) -> bool { + let GatewayError::Internal(message) = error else { + return false; + }; + let message = message.to_ascii_lowercase(); + // Missing encryption configuration is an operational failure and must + // remain fail-closed. Only errors that identify a stored field or a + // malformed legacy ciphertext are safe to isolate to one key. + if message.contains("encryption key is not configured") { + return false; + } + message.contains("provider_api_keys.api_key") + || message.contains("provider_api_keys.auth_config") + || message.contains("provider_api_keys.api_formats") + || message.contains("provider_api_keys.allowed_models") + || message.contains("legacy provider catalog credential") + || message.contains("stored provider catalog credential is empty") + || message.contains("aether secret envelope has the wrong record binding") + || message.contains("provider catalog credential is not an authenticated ciphertext") + || message.contains("provider catalog credential contains reserved framing") + || message.contains("provider catalog credential authentication failed") + || message.contains("provider catalog credential envelope") + || message + .contains("provider catalog key provider binding changed during credential migration") +} + +/// Stored provider/endpoint/key proxy secrets are opened independently by the +/// maintenance scan. A malformed historical row is safe to isolate, while +/// encryption/configuration failures remain fatal so operators are alerted. +fn is_nonfatal_stored_proxy_error(error: &GatewayError) -> bool { + let GatewayError::Internal(message) = error else { + return false; + }; + let message = message.to_ascii_lowercase(); + message.contains("stored provider proxy credentials cannot be decrypted") + || message.contains("stored endpoint proxy credentials cannot be decrypted") + || message.contains("stored key proxy credentials cannot be decrypted") + || message.contains("stored provider proxy changed during credential migration") + || message.contains("stored endpoint proxy changed during credential migration") + || message.contains("stored key changed during credential migration") + || message.contains("stored provider proxy credential migration did not stabilize") + || message.contains("stored endpoint proxy credential migration did not stabilize") + || message.contains("stored key proxy credential migration did not stabilize") +} + #[cfg(test)] mod tests { use aether_data_contracts::repository::provider_catalog::{ @@ -282,8 +437,10 @@ mod tests { }; use super::{ - agent_identity_needs_task_recovery, auth_config_has_refresh_token, oauth_refresh_candidate, + agent_identity_needs_task_recovery, auth_config_has_refresh_token, + is_nonfatal_legacy_catalog_credential_error, oauth_refresh_candidate, }; + use crate::GatewayError; #[test] fn legacy_antigravity_refresh_token_is_refreshable() { @@ -336,4 +493,54 @@ mod tests { Some("[REFRESH_FAILED] temporary"), )); } + + #[test] + fn only_stored_catalog_credential_errors_are_non_fatal() { + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal( + "provider catalog credential is not an authenticated ciphertext".to_string(), + ) + )); + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal( + "provider_api_keys.auth_config has an invalid provider catalog credential envelope" + .to_string(), + ) + )); + assert!(!is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal("postgres error: connection refused".to_string(),) + )); + assert!(!is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal( + "provider catalog credential encryption key is not configured".to_string(), + ) + )); + for scope in ["provider", "endpoint", "key"] { + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal(format!( + "stored {scope} proxy credentials cannot be decrypted" + )) + )); + } + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal("stored provider catalog credential is empty".to_string()) + )); + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal( + "Aether secret envelope has the wrong record binding".to_string() + ) + )); + for field in ["api_formats", "allowed_models"] { + assert!(is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal(format!( + "provider_api_keys.{field} contains a malformed value" + )) + )); + } + assert!(!is_nonfatal_legacy_catalog_credential_error( + &GatewayError::Internal( + "endpoint proxy credential encryption is unavailable".to_string(), + ) + )); + } } diff --git a/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs b/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs index e0f9d918d..f71c9031d 100644 --- a/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs +++ b/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs @@ -1042,6 +1042,17 @@ async fn refresh_provider_probe_keys( .await } +fn pool_quota_probe_worker_error_score_reason(_error: &GatewayError) -> Value { + serde_json::json!({ + "last_probe": { + "source": "pool_quota_probe", + "status": "worker_error", + "category": "worker_error", + "message": "Provider quota probe worker failed" + } + }) +} + fn update_summary_from_payload( summary: &mut PoolQuotaProbeRunSummary, selected_count: usize, @@ -1213,6 +1224,18 @@ fn probe_result_hard_state(item: &Value) -> Option { .unwrap_or_default() .trim() .to_ascii_lowercase(); + // Providers frequently encode an exhausted account as HTTP 429 with a + // provider-specific status/message (for example RESOURCE_EXHAUSTED or + // `quota exceeded`) rather than the normalized `quota_exhausted` status. + // Preserve that distinction in the score so the member stays out of the + // scheduler until a successful quota probe observes a reset. + let serialized_item = item.to_string().to_ascii_lowercase(); + let status_code = item.get("status_code").and_then(Value::as_u64); + if status == "quota_exhausted" + || (status_code == Some(429) && contains_quota_exhaustion_marker(&serialized_item)) + { + return Some(PoolMemberHardState::QuotaExhausted); + } match status.as_str() { "auth_invalid" | "forbidden" => Some(PoolMemberHardState::AuthInvalid), "workspace_deactivated" => Some(PoolMemberHardState::Banned), @@ -1226,6 +1249,26 @@ fn probe_result_hard_state(item: &Value) -> Option { } } +fn contains_quota_exhaustion_marker(value: &str) -> bool { + [ + "quota exhausted", + "quota_exhausted", + "quota exceeded", + "quota_exceeded", + "insufficient_quota", + "resource exhausted", + "resource has been exhausted", + "resource_exhausted", + "usage_limit_reached", + "limit_reached", + "quota limit reached", + "credits exhausted", + "insufficient credits", + ] + .iter() + .any(|marker| value.contains(marker)) +} + async fn perform_pool_quota_probe_for_provider( state: &AppState, admin_state: &AdminAppState<'_>, @@ -1417,20 +1460,14 @@ async fn perform_pool_quota_probe_for_provider( now_ts, false, Some(PoolMemberHardState::Cooldown), - serde_json::json!({ - "last_probe": { - "source": "pool_quota_probe", - "status": "worker_error", - "message": format!("{err:?}") - } - }), + pool_quota_probe_worker_error_score_reason(&err), ) .await; warn!( provider_id = %provider_short_id, provider_type, key_id, - error = ?err, + error_category = "provider_quota_probe_failed", "gateway pool quota probe failed" ); } @@ -1710,9 +1747,12 @@ pub(crate) fn spawn_pool_quota_probe_worker( interval.tick().await; loop { interval.tick().await; - if let Err(err) = perform_pool_quota_probe_once_with_config(&state, config).await { + if perform_pool_quota_probe_once_with_config(&state, config) + .await + .is_err() + { warn!( - error = ?err, + error_category = "pool_quota_probe_worker_failed", "gateway pool quota probe worker tick failed" ); } @@ -1726,6 +1766,30 @@ mod tests { use super::*; use serde_json::json; + #[test] + fn worker_error_score_reason_drops_runtime_error_details() { + let secret = "postgresql://admin:db-secret@db.internal/aether; Authorization: Bearer quota-secret; https://user:password@upstream.test/quota?q=secret"; + let patch = + pool_quota_probe_worker_error_score_reason(&GatewayError::Internal(secret.to_string())); + let serialized = patch.to_string(); + + assert_eq!(patch["last_probe"]["status"], "worker_error"); + assert_eq!(patch["last_probe"]["category"], "worker_error"); + assert_eq!( + patch["last_probe"]["message"], + "Provider quota probe worker failed" + ); + for sensitive in [ + "db-secret", + "quota-secret", + "user:password", + "db.internal", + "upstream.test", + ] { + assert!(!serialized.contains(sensitive)); + } + } + fn key( id: &str, provider_id: &str, @@ -1771,6 +1835,26 @@ mod tests { assert_eq!(selected, vec!["never".to_string(), "old".to_string()]); } + #[test] + fn quota_markers_in_429_probe_results_are_hard_exhaustion() { + assert_eq!( + probe_result_hard_state(&json!({ + "status": "rate_limited", + "status_code": 429, + "message": "RESOURCE_EXHAUSTED: quota exceeded" + })), + Some(PoolMemberHardState::QuotaExhausted) + ); + assert_eq!( + probe_result_hard_state(&json!({ + "status": "rate_limited", + "status_code": 429, + "message": "temporary rate limit" + })), + Some(PoolMemberHardState::Cooldown) + ); + } + fn score( member_id: &str, hard_state: PoolMemberHardState, diff --git a/apps/aether-gateway/src/maintenance/runtime/provider_quota_alert.rs b/apps/aether-gateway/src/maintenance/runtime/provider_quota_alert.rs index 0025967d1..c869b0911 100644 --- a/apps/aether-gateway/src/maintenance/runtime/provider_quota_alert.rs +++ b/apps/aether-gateway/src/maintenance/runtime/provider_quota_alert.rs @@ -199,7 +199,14 @@ async fn run_provider_quota_alert_for_provider( let Some(total_available) = extract_total_available(&payload) else { warn!( provider_id = %provider_id, - payload = %payload, + payload_status = payload + .get("status") + .and_then(serde_json::Value::as_str) + .unwrap_or("unknown"), + action_type = payload + .get("action_type") + .and_then(serde_json::Value::as_str) + .unwrap_or("unknown"), "provider quota alert skipped because balance payload has no total_available" ); write_checked_runtime_state_without_balance(state, &provider_id, now_unix_secs).await; diff --git a/apps/aether-gateway/src/maintenance/runtime/proxy_node_staleness.rs b/apps/aether-gateway/src/maintenance/runtime/proxy_node_staleness.rs index 728b41988..85201b324 100644 --- a/apps/aether-gateway/src/maintenance/runtime/proxy_node_staleness.rs +++ b/apps/aether-gateway/src/maintenance/runtime/proxy_node_staleness.rs @@ -53,6 +53,7 @@ pub(super) async fn cleanup_stale_proxy_nodes_once( ); let mutation = ProxyNodeTunnelStatusMutation { node_id: node.id.clone(), + expected_tunnel_generation: Some(node.tunnel_generation.clone()), connected: false, conn_count: 0, detail: Some(detail), diff --git a/apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs b/apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs index 14805a2b3..b873782a5 100644 --- a/apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs +++ b/apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs @@ -27,6 +27,8 @@ struct ProxyUpgradeRolloutPlan { #[serde(default)] skipped_node_ids: Vec, #[serde(default)] + skipped_node_generations: std::collections::BTreeMap, + #[serde(default)] tracked_nodes: Vec, } @@ -39,6 +41,8 @@ pub(crate) struct ProxyUpgradeRolloutProbeConfig { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] struct ProxyUpgradeRolloutTrackedNode { node_id: String, + #[serde(default)] + tunnel_generation: String, dispatched_at_unix_secs: u64, version_confirmed_at_unix_secs: Option, #[serde(default)] @@ -51,6 +55,7 @@ struct ProxyUpgradeRolloutTrackedNode { #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct ProxyUpgradeRolloutPendingProbe { pub(crate) node_id: String, + pub(crate) tunnel_generation: String, pub(crate) url: String, pub(crate) timeout_secs: u64, } @@ -167,6 +172,7 @@ struct RolloutSnapshot { completed: Vec, pending: Vec, pending_conflicts: Vec, + pending_conflict_nodes: Vec, skipped: Vec, available: Vec, ready_to_finalize: Vec, @@ -245,6 +251,11 @@ pub(crate) async fn start_proxy_upgrade_rollout( .filter(|_| preserve_existing) .map(|plan| plan.skipped_node_ids.clone()) .unwrap_or_default(), + skipped_node_generations: existing + .as_ref() + .filter(|_| preserve_existing) + .map(|plan| plan.skipped_node_generations.clone()) + .unwrap_or_default(), tracked_nodes: existing .as_ref() .filter(|_| preserve_existing) @@ -307,6 +318,7 @@ pub(crate) async fn collect_proxy_upgrade_rollout_probes( } Some(ProxyUpgradeRolloutPendingProbe { node_id: tracked.node_id, + tunnel_generation: tracked.tunnel_generation, url: probe.url.clone(), timeout_secs: probe.timeout_secs, }) @@ -430,10 +442,11 @@ pub(crate) async fn clear_proxy_upgrade_rollout_conflicts( } let mut cleared_node_ids = Vec::new(); - for node_id in snapshot.pending_conflicts { + for node in snapshot.pending_conflict_nodes { let Some(updated) = data .update_proxy_node_remote_config(&ProxyNodeRemoteConfigMutation { - node_id: node_id.clone(), + node_id: node.id, + expected_tunnel_generation: Some(node.tunnel_generation), node_name: None, allowed_ports: None, log_level: None, @@ -484,6 +497,8 @@ pub(crate) async fn skip_proxy_upgrade_rollout_node( plan.skipped_node_ids.sort(); plan.skipped_node_ids.dedup(); } + plan.skipped_node_generations + .insert(node_id.to_string(), node.tunnel_generation.clone()); plan.tracked_nodes .retain(|tracked| tracked.node_id != node_id); plan.updated_at_unix_secs = now; @@ -495,6 +510,7 @@ pub(crate) async fn skip_proxy_upgrade_rollout_node( let _ = data .update_proxy_node_remote_config(&ProxyNodeRemoteConfigMutation { node_id: node_id.to_string(), + expected_tunnel_generation: Some(node.tunnel_generation), node_name: None, allowed_ports: None, log_level: None, @@ -554,6 +570,7 @@ pub(crate) async fn restore_proxy_upgrade_rollout_skipped_nodes( } plan.skipped_node_ids.clear(); + plan.skipped_node_generations.clear(); plan.updated_at_unix_secs = now; save_proxy_upgrade_rollout_plan(data, &plan).await?; @@ -584,14 +601,15 @@ pub(crate) async fn retry_proxy_upgrade_rollout_node( let Some(mut plan) = load_proxy_upgrade_rollout_plan(data).await? else { return Ok(None); }; - let Some(_node) = data.find_proxy_node(node_id).await? else { + let Some(node) = data.find_proxy_node(node_id).await? else { return Ok(None); }; let now = now_unix_secs(); - let _ = data + let Some(updated) = data .update_proxy_node_remote_config(&ProxyNodeRemoteConfigMutation { node_id: node_id.to_string(), + expected_tunnel_generation: Some(node.tunnel_generation), node_name: None, allowed_ports: None, log_level: None, @@ -599,13 +617,18 @@ pub(crate) async fn retry_proxy_upgrade_rollout_node( scheduling_state: None, upgrade_to: Some(Some(plan.version.clone())), }) - .await?; + .await? + else { + return Ok(None); + }; plan.skipped_node_ids.retain(|id| id != node_id); + plan.skipped_node_generations.remove(node_id); plan.tracked_nodes .retain(|tracked| tracked.node_id != node_id); plan.tracked_nodes.push(ProxyUpgradeRolloutTrackedNode { node_id: node_id.to_string(), + tunnel_generation: updated.tunnel_generation, dispatched_at_unix_secs: now, version_confirmed_at_unix_secs: None, traffic_confirmed_at_unix_secs: None, @@ -659,6 +682,7 @@ async fn advance_proxy_upgrade_rollout( let _ = data .update_proxy_node_remote_config(&ProxyNodeRemoteConfigMutation { node_id: node.id.clone(), + expected_tunnel_generation: Some(node.tunnel_generation.clone()), node_name: None, allowed_ports: None, log_level: None, @@ -711,11 +735,12 @@ async fn advance_proxy_upgrade_rollout( return Ok(summary); } - let mut updated_node_ids = Vec::with_capacity(selected.len()); + let mut updated_nodes = Vec::with_capacity(selected.len()); for node in selected { let Some(updated) = data .update_proxy_node_remote_config(&ProxyNodeRemoteConfigMutation { node_id: node.id.clone(), + expected_tunnel_generation: Some(node.tunnel_generation), node_name: None, allowed_ports: None, log_level: None, @@ -727,28 +752,30 @@ async fn advance_proxy_upgrade_rollout( else { continue; }; - updated_node_ids.push(updated.id); + updated_nodes.push((updated.id, updated.tunnel_generation)); } plan.last_dispatched_at_unix_secs = Some(now); plan.updated_at_unix_secs = now; plan.tracked_nodes - .extend( - updated_node_ids - .iter() - .cloned() - .map(|node_id| ProxyUpgradeRolloutTrackedNode { - node_id, - dispatched_at_unix_secs: now, - version_confirmed_at_unix_secs: None, - traffic_confirmed_at_unix_secs: None, - confirm_failed_requests: None, - confirm_dns_failures: None, - confirm_stream_errors: None, - }), - ); + .extend(updated_nodes.iter().map(|(node_id, tunnel_generation)| { + ProxyUpgradeRolloutTrackedNode { + node_id: node_id.clone(), + tunnel_generation: tunnel_generation.clone(), + dispatched_at_unix_secs: now, + version_confirmed_at_unix_secs: None, + traffic_confirmed_at_unix_secs: None, + confirm_failed_requests: None, + confirm_dns_failures: None, + confirm_stream_errors: None, + } + })); save_proxy_upgrade_rollout_plan(data, &plan).await?; + let updated_node_ids = updated_nodes + .into_iter() + .map(|(node_id, _)| node_id) + .collect::>(); summary.updated = updated_node_ids.len(); summary.skipped = summary.skipped.saturating_sub(summary.updated); summary.node_ids = updated_node_ids.clone(); @@ -785,7 +812,12 @@ fn build_rollout_snapshot( snapshot.online_eligible_total = snapshot.online_eligible_total.saturating_add(1); } - if skipped_node_ids.contains(node.id.as_str()) { + if skipped_node_ids.contains(node.id.as_str()) + && plan + .skipped_node_generations + .get(node.id.as_str()) + .is_some_and(|generation| generation == &node.tunnel_generation) + { snapshot.skipped.push(node.id.clone()); tracked_by_node_id.remove(node.id.as_str()); continue; @@ -794,7 +826,13 @@ fn build_rollout_snapshot( let reported_version = proxy_reported_version(node.proxy_metadata.as_ref()); let pending_target = remote_config_upgrade_target(node.remote_config.as_ref()); - if let Some(mut tracked) = tracked_by_node_id.remove(node.id.as_str()) { + if let Some(mut tracked) = tracked_by_node_id + .remove(node.id.as_str()) + .filter(|tracked| { + !tracked.tunnel_generation.is_empty() + && tracked.tunnel_generation == node.tunnel_generation + }) + { snapshot.remaining_total = snapshot.remaining_total.saturating_add(1); if reported_version.as_deref() == Some(target_version.as_str()) { @@ -845,6 +883,7 @@ fn build_rollout_snapshot( snapshot.pending.push(node.id.clone()); } else { snapshot.pending_conflicts.push(node.id.clone()); + snapshot.pending_conflict_nodes.push(node); } continue; } @@ -905,9 +944,24 @@ async fn save_proxy_upgrade_rollout_plan( Ok(()) } +/// Records rollout traffic only when the callback is bound to the exact tunnel +/// generation that was dispatched. Callers handling a live connection should +/// use [`record_proxy_upgrade_traffic_success_for_generation`]. pub(crate) async fn record_proxy_upgrade_traffic_success( data: &GatewayDataState, node_id: &str, +) -> Result { + let Some(node) = data.find_proxy_node(node_id).await? else { + return Ok(false); + }; + record_proxy_upgrade_traffic_success_for_generation(data, node_id, &node.tunnel_generation) + .await +} + +pub(crate) async fn record_proxy_upgrade_traffic_success_for_generation( + data: &GatewayDataState, + node_id: &str, + tunnel_generation: &str, ) -> Result { if !data.has_system_config_store() { return Ok(false); @@ -916,13 +970,17 @@ pub(crate) async fn record_proxy_upgrade_traffic_success( let Some(mut plan) = load_proxy_upgrade_rollout_plan(data).await? else { return Ok(false); }; + let Some(node) = data.find_proxy_node(node_id).await? else { + return Ok(false); + }; + if tunnel_generation.is_empty() || node.tunnel_generation != tunnel_generation { + return Ok(false); + } let now = now_unix_secs(); - let Some(tracked) = plan - .tracked_nodes - .iter_mut() - .find(|tracked| tracked.node_id == node_id) - else { + let Some(tracked) = plan.tracked_nodes.iter_mut().find(|tracked| { + tracked.node_id == node_id && tracked.tunnel_generation == tunnel_generation + }) else { return Ok(false); }; let Some(version_confirmed_at_unix_secs) = tracked.version_confirmed_at_unix_secs else { diff --git a/apps/aether-gateway/src/maintenance/runtime/runners.rs b/apps/aether-gateway/src/maintenance/runtime/runners.rs index 615d3588c..410234c47 100644 --- a/apps/aether-gateway/src/maintenance/runtime/runners.rs +++ b/apps/aether-gateway/src/maintenance/runtime/runners.rs @@ -7,17 +7,19 @@ use tracing::{info, warn}; use crate::data::GatewayDataState; use crate::{AppState, GatewayError}; +use super::cleanup_runs::cleanup_data_layer_error_category; use super::{ advance_proxy_upgrade_rollout_once, cleanup_audit_logs_once, cleanup_expired_gemini_file_mappings_once, cleanup_proxy_node_metrics_once, cleanup_request_candidates_once, cleanup_stale_pending_requests_once, cleanup_stale_proxy_nodes_once, collect_proxy_upgrade_rollout_probes, now_unix_secs, - perform_db_maintenance_once, perform_manual_usage_cleanup_once, perform_provider_checkin_once, + perform_automatic_usage_cleanup_once, perform_db_maintenance_once, + perform_manual_usage_cleanup_once, perform_provider_checkin_once, perform_stats_aggregation_once, perform_stats_hourly_aggregation_once, - perform_usage_cleanup_once, perform_wallet_daily_usage_aggregation_once, - record_admin_cleanup_run, record_completed_cleanup_run, record_failed_cleanup_run, - record_proxy_upgrade_traffic_success, summarize_database_pool, AdminCleanupRunRecord, - ManualUsageCleanupOptions, + perform_wallet_daily_usage_aggregation_once, record_admin_cleanup_run, + record_completed_cleanup_run, record_failed_cleanup_run, + record_proxy_upgrade_traffic_success_for_generation, summarize_database_pool, + AdminCleanupRunRecord, ManualUsageCleanupOptions, }; pub(super) async fn run_audit_cleanup_once(data: &GatewayDataState) -> Result<(), DataLayerError> { @@ -120,14 +122,19 @@ pub(super) async fn run_proxy_upgrade_rollout_once(state: &AppState) -> Result<( .await { Ok(status) if (200..300).contains(&status) => { - let _ = record_proxy_upgrade_traffic_success(&state.data, &probe.node_id).await?; + let _ = record_proxy_upgrade_traffic_success_for_generation( + &state.data, + &probe.node_id, + &probe.tunnel_generation, + ) + .await?; probe_recorded = true; info!( event_name = "proxy_upgrade_rollout_probe_succeeded", log_type = "ops", worker = "proxy_upgrade_rollout", node_id = %probe.node_id, - url = %probe.url, + probe_origin = %crate::handlers::shared::security_log_url_origin(&probe.url), status, "gateway confirmed proxy upgrade health probe" ); @@ -138,7 +145,7 @@ pub(super) async fn run_proxy_upgrade_rollout_once(state: &AppState) -> Result<( log_type = "ops", worker = "proxy_upgrade_rollout", node_id = %probe.node_id, - url = %probe.url, + probe_origin = %crate::handlers::shared::security_log_url_origin(&probe.url), status, "gateway proxy upgrade health probe returned non-success status" ); @@ -149,7 +156,7 @@ pub(super) async fn run_proxy_upgrade_rollout_once(state: &AppState) -> Result<( log_type = "ops", worker = "proxy_upgrade_rollout", node_id = %probe.node_id, - url = %probe.url, + probe_origin = %crate::handlers::shared::security_log_url_origin(&probe.url), error = %error, "gateway proxy upgrade health probe failed" ); @@ -237,7 +244,7 @@ pub(super) async fn run_stats_aggregation_once( pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<(), DataLayerError> { let started_at_unix_secs = now_unix_secs(); let started_at = Instant::now(); - let summary = match perform_usage_cleanup_once(data).await { + let summary = match perform_automatic_usage_cleanup_once(data).await { Ok(summary) => summary, Err(err) => { record_failed_cleanup_run( @@ -265,6 +272,8 @@ pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<() "header_cleaned": summary.header_cleaned, "keys_cleaned": summary.keys_cleaned, "records_deleted": summary.records_deleted, + "cost_reservations_deleted": summary.cost_reservations_deleted, + "request_admissions_deleted": summary.request_admissions_deleted, }), format!( "请求记录自动清理完成,影响 {} 项", @@ -275,6 +284,8 @@ pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<() .saturating_add(summary.header_cleaned) .saturating_add(summary.keys_cleaned) .saturating_add(summary.records_deleted) + .saturating_add(summary.cost_reservations_deleted) + .saturating_add(summary.request_admissions_deleted) ), ) .await; @@ -284,6 +295,8 @@ pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<() || summary.header_cleaned > 0 || summary.keys_cleaned > 0 || summary.records_deleted > 0 + || summary.cost_reservations_deleted > 0 + || summary.request_admissions_deleted > 0 { info!( event_name = "usage_cleanup_completed", @@ -295,6 +308,8 @@ pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<() header_cleaned = summary.header_cleaned, keys_cleaned = summary.keys_cleaned, records_deleted = summary.records_deleted, + cost_reservations_deleted = summary.cost_reservations_deleted, + request_admissions_deleted = summary.request_admissions_deleted, "gateway finished usage cleanup" ); } @@ -384,6 +399,8 @@ pub(crate) async fn run_manual_usage_cleanup_once( "header_cleaned": summary.header_cleaned, "keys_cleaned": summary.keys_cleaned, "records_deleted": summary.records_deleted, + "cost_reservations_deleted": summary.cost_reservations_deleted, + "request_admissions_deleted": summary.request_admissions_deleted, "mode": options.mode.as_str(), "requested_older_than_days": options.requested_older_than_days, "targets": options.targets, @@ -401,6 +418,8 @@ pub(crate) async fn run_manual_usage_cleanup_once( requested_older_than_days = options.requested_older_than_days, actor_user_id = actor_user_id.as_deref(), total_affected = total, + cost_reservations_deleted = summary.cost_reservations_deleted, + request_admissions_deleted = summary.request_admissions_deleted, "gateway finished manual usage cleanup" ); Ok(summary) @@ -476,7 +495,7 @@ async fn run_manual_usage_cleanup_task( None, actor_user_id.as_deref(), ), - error: Some(err.to_string()), + error: Some(cleanup_data_layer_error_category(&err).to_string()), }; if let Err(record_err) = record_admin_cleanup_run(&data, record).await { warn!(error = %record_err, "failed to record manual usage cleanup failure"); @@ -496,6 +515,8 @@ fn usage_cleanup_total( .saturating_add(summary.header_cleaned) .saturating_add(summary.keys_cleaned) .saturating_add(summary.records_deleted) + .saturating_add(summary.cost_reservations_deleted) + .saturating_add(summary.request_admissions_deleted) } fn manual_usage_cleanup_start_message(options: ManualUsageCleanupOptions) -> String { @@ -549,6 +570,8 @@ fn manual_usage_cleanup_progress_summary( "header_cleaned": summary.header_cleaned, "keys_cleaned": summary.keys_cleaned, "records_deleted": summary.records_deleted, + "cost_reservations_deleted": summary.cost_reservations_deleted, + "request_admissions_deleted": summary.request_admissions_deleted, "total": usage_cleanup_total(summary), "actor_user_id": actor_user_id, }) @@ -564,7 +587,7 @@ impl std::fmt::Display for ManualUsageCleanupError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Self::AlreadyRunning => f.write_str("a usage cleanup run is already in progress"), - Self::DataLayer(err) => write!(f, "{err}"), + Self::DataLayer(_) => f.write_str("usage cleanup data operation failed"), } } } @@ -597,20 +620,49 @@ pub(super) fn run_pool_monitor_once(data: &GatewayDataState) { ); } -pub(super) async fn run_pending_cleanup_once( - data: &GatewayDataState, -) -> Result<(), DataLayerError> { - let summary = cleanup_stale_pending_requests_once(data).await?; - if summary.failed > 0 || summary.recovered > 0 { +pub(super) async fn run_pending_cleanup_once(app: &AppState) -> Result<(), DataLayerError> { + let data = &app.data; + let cleanup_result = cleanup_stale_pending_requests_once(data).await; + let referral_result = if data.has_referral_data_backend() { + Some(app.reconcile_referral_rewards_once().await) + } else { + None + }; + let summary = cleanup_result.as_ref().copied().unwrap_or_default(); + let referral_summary = referral_result + .as_ref() + .and_then(|result| result.as_ref().ok().copied()) + .unwrap_or_default(); + if summary.failed > 0 + || summary.recovered > 0 + || cleanup_result.is_err() + || referral_result.as_ref().is_some_and(Result::is_err) + || referral_summary.order_attempted > 0 + || referral_summary.reward_attempted > 0 + || referral_summary.reversal_attempted > 0 + { info!( event_name = "pending_cleanup_completed", log_type = "ops", worker = "pending_cleanup", failed = summary.failed, recovered = summary.recovered, + referral_order_attempted = referral_summary.order_attempted, + referral_order_repaired = referral_summary.order_repaired, + referral_reward_attempted = referral_summary.reward_attempted, + referral_reward_applied = referral_summary.reward_applied, + referral_reversal_attempted = referral_summary.reversal_attempted, + referral_reversal_applied = referral_summary.reversal_applied, + referral_deferred = referral_summary.deferred, "gateway cleaned stale pending and streaming requests" ); } + if let Err(error) = cleanup_result { + return Err(error); + } + if let Some(Err(error)) = referral_result { + return Err(DataLayerError::UnexpectedValue(error.into_message())); + } Ok(()) } diff --git a/apps/aether-gateway/src/maintenance/runtime/tests.rs b/apps/aether-gateway/src/maintenance/runtime/tests.rs index fbe0ddcff..03106ced3 100644 --- a/apps/aether-gateway/src/maintenance/runtime/tests.rs +++ b/apps/aether-gateway/src/maintenance/runtime/tests.rs @@ -2,10 +2,16 @@ use std::collections::{HashSet, VecDeque}; use std::sync::{Arc, Mutex}; use std::time::Duration; +use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::proxy_nodes::{ bucket_start_unix_secs, InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeWriteRepository, StoredProxyNode, }; +use aether_data::repository::settlement::{ + InMemorySettlementRepository, ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput, + SettlementWriteRepository, UsagePolicyCostWindow, UsagePolicyRequestWindow, +}; use aether_runtime::bounded_queue; use axum::extract::ws::Message; use chrono::{DateTime, Utc}; @@ -18,7 +24,8 @@ use super::{ cleanup_proxy_node_metrics_once, cleanup_stale_proxy_nodes_once, inspect_proxy_upgrade_rollout, next_daily_run_after, next_db_maintenance_run_after, next_stats_aggregation_run_after, next_stats_hourly_aggregation_run_after, pending_cleanup_batch_size, - pending_cleanup_timeout_minutes, plan_pending_cleanup_batch, provider_checkin_schedule, + pending_cleanup_timeout_minutes, perform_automatic_usage_cleanup_once, + perform_manual_usage_cleanup_once, plan_pending_cleanup_batch, provider_checkin_schedule, proxy_node_metrics_cleanup_settings, record_proxy_upgrade_traffic_success, run_db_maintenance_with, run_proxy_upgrade_rollout_once, spawn_account_self_check_worker, spawn_audit_cleanup_worker, spawn_db_maintenance_worker, @@ -167,6 +174,29 @@ fn sample_connected_proxy_node( ) } +const ACTIVE_PROBE_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; +const ACTIVE_PROBE_TUNNEL_TEST_GENERATION: &str = "active-probe-test-generation-1"; + +fn active_probe_tunnel_metadata(version: &str) -> serde_json::Value { + json!({ + "version": version, + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": ACTIVE_PROBE_TUNNEL_TEST_PSK, + } + }) +} + +async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, +) -> Message { + tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) +} + #[tokio::test] async fn stale_proxy_node_cleanup_marks_timed_out_tunnel_offline() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ @@ -230,6 +260,7 @@ async fn proxy_upgrade_rollout_advances_next_wave_after_version_health_confirmat repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-alpha".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), @@ -267,6 +298,7 @@ async fn proxy_upgrade_rollout_advances_next_wave_after_version_health_confirmat repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-zeta".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), @@ -378,6 +410,7 @@ async fn proxy_upgrade_rollout_blocks_next_wave_after_post_upgrade_transport_err repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-alpha".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), @@ -406,6 +439,7 @@ async fn proxy_upgrade_rollout_blocks_next_wave_after_post_upgrade_transport_err repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-alpha".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), @@ -436,9 +470,10 @@ async fn proxy_upgrade_rollout_blocks_next_wave_after_post_upgrade_transport_err #[tokio::test] async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_confirmation() { - let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000); + let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000) + .with_tunnel_generation(ACTIVE_PROBE_TUNNEL_TEST_GENERATION.to_string()); alpha.name = "alpha".to_string(); - alpha.proxy_metadata = Some(json!({"version": "1.0.0"})); + alpha.proxy_metadata = Some(active_probe_tunnel_metadata("1.0.0")); alpha.remote_config = None; alpha.config_version = 0; @@ -450,7 +485,8 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha])); let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) - .with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()); + .with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); let state = AppState::new() .expect("gateway state should build") .with_data_state_for_tests(data.clone()); @@ -472,6 +508,7 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-alpha".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), @@ -479,7 +516,7 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con failed_requests_delta: Some(0), dns_failures_delta: Some(0), stream_errors_delta: Some(0), - proxy_metadata: Some(json!({"version": "2.0.0"})), + proxy_metadata: Some(active_probe_tunnel_metadata("2.0.0")), proxy_version: Some("2.0.0".to_string()), }) .await @@ -488,9 +525,8 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con let tunnel_state = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_state - .hub - .register_proxy(Arc::new(crate::tunnel::TunnelProxyConn::new( + tunnel_state.hub.register_proxy(Arc::new( + crate::tunnel::TunnelProxyConn::new( 700, "node-alpha".to_string(), "Node Alpha".to_string(), @@ -498,14 +534,18 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con proxy_close_tx, 16, 2, - ))); + ) + .with_tunnel_generation(ACTIVE_PROBE_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(ACTIVE_PROBE_TUNNEL_TEST_PSK.to_string()), + )); let responder_hub = tunnel_state.hub.clone(); let responder = tokio::spawn(async move { - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { - Message::Binary(data) => data, - other => panic!("unexpected message: {other:?}"), - }; + let request_headers = + match recv_tunnel_test_frame(&mut proxy_rx, "probe headers frame").await { + Message::Binary(data) => data, + other => panic!("unexpected message: {other:?}"), + }; let request_header = crate::tunnel::tunnel_protocol::FrameHeader::parse(&request_headers) .expect("probe request headers should parse"); assert_eq!( @@ -513,7 +553,7 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con crate::tunnel::tunnel_protocol::REQUEST_HEADERS ); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "probe body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -594,6 +634,106 @@ async fn spawn_usage_cleanup_worker_skips_when_usage_writer_unavailable() { assert!(spawn_usage_cleanup_worker(state).is_none()); } +#[tokio::test] +async fn spawn_usage_cleanup_worker_starts_for_settlement_writer_only() { + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_settlement_writer_for_tests(Arc::new( + InMemorySettlementRepository::default(), + )), + ); + let handle = spawn_usage_cleanup_worker(state) + .expect("settlement ledger cleanup should have a worker without a usage writer"); + handle.abort(); +} + +#[tokio::test] +async fn automatic_usage_cleanup_collects_ledgers_when_usage_cleanup_is_disabled() { + let repository = Arc::new(InMemorySettlementRepository::default()); + repository + .reserve_usage_policy_request(ReserveUsagePolicyRequestInput { + request_id: "expired-request".to_string(), + subject_id: "user-1".to_string(), + event_token: "expired-event".to_string(), + admitted_at_unix_secs: 1, + retain_until_unix_secs: 2, + windows: vec![UsagePolicyRequestWindow { + starts_at_unix_secs: 0, + ends_at_unix_secs: 2, + limit_requests: 10, + }], + }) + .await + .expect("request admission should be seeded"); + repository + .reserve_usage_policy_cost(ReserveUsagePolicyCostInput { + request_id: "expired-cost".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "expired-reservation".to_string(), + admitted_at_unix_secs: 1, + reserved_cost_units: 1, + reservation_expires_at_unix_secs: 2, + retain_until_unix_secs: 2, + windows: vec![UsagePolicyCostWindow { + window_id: "expired-window".to_string(), + starts_at_unix_secs: 0, + ends_at_unix_secs: 2, + limit_cost_units: 10, + }], + }) + .await + .expect("cost reservation should be seeded"); + let data = GatewayDataState::disabled() + .with_settlement_writer_for_tests(repository) + .with_system_config_values_for_tests([ + ("enable_auto_cleanup".to_string(), json!(false)), + ("cleanup_batch_size".to_string(), json!(1)), + ]); + + let summary = perform_automatic_usage_cleanup_once(&data) + .await + .expect("automatic ledger cleanup should succeed"); + + assert_eq!(summary.cost_reservations_deleted, 1); + assert_eq!(summary.request_admissions_deleted, 1); +} + +#[tokio::test] +async fn manual_usage_cleanup_does_not_collect_policy_ledgers() { + let repository = Arc::new(InMemorySettlementRepository::default()); + repository + .reserve_usage_policy_request(ReserveUsagePolicyRequestInput { + request_id: "manual-request".to_string(), + subject_id: "user-1".to_string(), + event_token: "manual-event".to_string(), + admitted_at_unix_secs: 1, + retain_until_unix_secs: 2, + windows: vec![UsagePolicyRequestWindow { + starts_at_unix_secs: 0, + ends_at_unix_secs: 2, + limit_requests: 10, + }], + }) + .await + .expect("request admission should be seeded"); + let data = GatewayDataState::disabled().with_settlement_writer_for_tests(repository.clone()); + + let summary = + perform_manual_usage_cleanup_once(&data, super::ManualUsageCleanupOptions::policy()) + .await + .expect("manual usage cleanup should succeed"); + + assert_eq!(summary.request_admissions_deleted, 0); + assert_eq!( + repository + .cleanup_usage_policy_request_admissions(u64::MAX, 10) + .await + .expect("direct cleanup should find the retained admission"), + 1 + ); +} + #[tokio::test] async fn spawn_wallet_daily_usage_aggregation_worker_skips_when_wallet_daily_usage_backend_unavailable( ) { @@ -851,6 +991,7 @@ async fn proxy_node_metrics_cleanup_deletes_expired_buckets_in_batches() { repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: node_id.to_string(), + expected_tunnel_generation: None, heartbeat_interval: Some(30), active_connections: Some(i32::try_from(idx + 1).unwrap()), total_requests_delta: None, @@ -915,6 +1056,7 @@ async fn proxy_node_metrics_cleanup_respects_auto_cleanup_toggle() { repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-metrics-disabled".to_string(), + expected_tunnel_generation: None, heartbeat_interval: Some(30), active_connections: Some(1), total_requests_delta: None, diff --git a/apps/aether-gateway/src/maintenance/runtime/usage_cleanup.rs b/apps/aether-gateway/src/maintenance/runtime/usage_cleanup.rs index b7f20e62e..38f6b40dc 100644 --- a/apps/aether-gateway/src/maintenance/runtime/usage_cleanup.rs +++ b/apps/aether-gateway/src/maintenance/runtime/usage_cleanup.rs @@ -67,6 +67,14 @@ pub(super) async fn perform_usage_cleanup_once( perform_usage_cleanup_once_with_override(data, None, true).await } +pub(super) async fn perform_automatic_usage_cleanup_once( + data: &GatewayDataState, +) -> Result { + let mut summary = perform_usage_cleanup_once(data).await?; + cleanup_usage_policy_ledgers(data, &mut summary).await?; + Ok(summary) +} + pub(super) async fn perform_usage_cleanup_once_with_override( data: &GatewayDataState, override_older_than: Option, @@ -128,6 +136,41 @@ async fn perform_usage_cleanup_once_with_options( .await } +async fn cleanup_usage_policy_ledgers( + data: &GatewayDataState, + summary: &mut UsageCleanupSummary, +) -> Result<(), DataLayerError> { + if !data.has_settlement_writer() { + return Ok(()); + } + + let settings = usage_cleanup_settings(data).await?; + let now_unix_secs = u64::try_from(Utc::now().timestamp()).unwrap_or(0); + let cleanup_batch_size = settings.batch_size.max(1); + // Bound each maintenance run while still draining faster than one day's normal growth. + for _ in 0..32 { + let deleted = data + .cleanup_usage_policy_cost_reservations(now_unix_secs, cleanup_batch_size) + .await?; + summary.cost_reservations_deleted = + summary.cost_reservations_deleted.saturating_add(deleted); + if deleted < cleanup_batch_size { + break; + } + } + for _ in 0..32 { + let deleted = data + .cleanup_usage_policy_request_admissions(now_unix_secs, cleanup_batch_size) + .await?; + summary.request_admissions_deleted = + summary.request_admissions_deleted.saturating_add(deleted); + if deleted < cleanup_batch_size { + break; + } + } + Ok(()) +} + pub(crate) async fn preview_manual_usage_cleanup( data: &GatewayDataState, options: ManualUsageCleanupOptions, diff --git a/apps/aether-gateway/src/maintenance/runtime/workers.rs b/apps/aether-gateway/src/maintenance/runtime/workers.rs index d3b0febfd..7275a9350 100644 --- a/apps/aether-gateway/src/maintenance/runtime/workers.rs +++ b/apps/aether-gateway/src/maintenance/runtime/workers.rs @@ -279,7 +279,7 @@ pub(crate) fn spawn_stats_aggregation_worker(app: AppState) -> Option Option> { - if !app.data.has_usage_writer() { + if !app.data.has_usage_writer() && !app.data.has_settlement_writer() { return None; } @@ -595,7 +595,7 @@ pub(crate) fn spawn_gemini_file_mapping_cleanup_worker( } pub(crate) fn spawn_pending_cleanup_worker(app: AppState) -> Option> { - if !app.data.has_usage_writer() { + if !app.data.has_usage_writer() && !app.data.has_referral_data_backend() { return None; } @@ -603,8 +603,8 @@ pub(crate) fn spawn_pending_cleanup_worker(app: AppState) -> Option Option, + pub(crate) verified_token_hash: VerifiedManagementTokenHash, +} + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct VerifiedManagementTokenHash(String); + +impl VerifiedManagementTokenHash { + pub(crate) fn new(value: String) -> Self { + Self(value) + } + + pub(crate) fn as_str(&self) -> &str { + self.0.as_str() + } +} + +impl std::fmt::Debug for VerifiedManagementTokenHash { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("VerifiedManagementTokenHash([REDACTED])") + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ManagementTokenAuthError { + Missing, + Invalid, + Unavailable, +} + +#[async_trait::async_trait] +pub(crate) trait ManagementTokenAuthSource { + async fn find_management_token( + &self, + token_hash: &str, + ) -> Result, ManagementTokenAuthError>; + + async fn find_current_user( + &self, + user_id: &str, + ) -> Result, ManagementTokenAuthError>; +} + +#[async_trait::async_trait] +impl ManagementTokenAuthSource for GatewayDataState { + async fn find_management_token( + &self, + token_hash: &str, + ) -> Result, ManagementTokenAuthError> { + self.get_management_token_with_user_by_hash(token_hash) + .await + .map_err(|_| ManagementTokenAuthError::Unavailable) + } + + async fn find_current_user( + &self, + user_id: &str, + ) -> Result, ManagementTokenAuthError> { + self.find_user_auth_by_id(user_id) + .await + .map_err(|_| ManagementTokenAuthError::Unavailable) + } +} + +#[async_trait::async_trait] +impl ManagementTokenAuthSource for crate::AppState { + async fn find_management_token( + &self, + token_hash: &str, + ) -> Result, ManagementTokenAuthError> { + self.get_management_token_with_user_by_hash(token_hash) + .await + .map_err(|_| ManagementTokenAuthError::Unavailable) + } + + async fn find_current_user( + &self, + user_id: &str, + ) -> Result, ManagementTokenAuthError> { + self.find_user_auth_by_id(user_id) + .await + .map_err(|_| ManagementTokenAuthError::Unavailable) + } +} + +pub(crate) async fn authenticate_management_token( + source: &S, + headers: &HeaderMap, + remote_ip: IpAddr, +) -> Result +where + S: ManagementTokenAuthSource + Sync + ?Sized, +{ + let token = extract_unique_management_token_bearer(headers)?; + let token_hash = VerifiedManagementTokenHash::new(hash_management_token(token)); + authenticate_management_token_hash(source, &token_hash, remote_ip).await +} + +pub(crate) async fn authenticate_management_token_hash( + source: &S, + token_hash: &VerifiedManagementTokenHash, + remote_ip: IpAddr, +) -> Result +where + S: ManagementTokenAuthSource + Sync + ?Sized, +{ + let token_with_user = source + .find_management_token(token_hash.as_str()) + .await? + .ok_or(ManagementTokenAuthError::Invalid)?; + if token_with_user.token.user_id != token_with_user.user.id { + return Err(ManagementTokenAuthError::Invalid); + } + + let now = chrono::Utc::now().timestamp().max(0) as u64; + if !token_with_user.token.is_active + || token_with_user + .token + .expires_at_unix_secs + .is_some_and(|expires_at| expires_at <= now) + || !crate::handlers::shared::json_ip_rules_allow( + token_with_user.token.allowed_ips.as_ref(), + remote_ip, + ) + { + return Err(ManagementTokenAuthError::Invalid); + } + + let user = source + .find_current_user(&token_with_user.token.user_id) + .await? + .ok_or(ManagementTokenAuthError::Invalid)?; + if !user.is_active || user.is_deleted || !crate::roles::can_access_admin_console(&user.role) { + return Err(ManagementTokenAuthError::Invalid); + } + + let permissions = crate::control::management_token_permission_keys_from_value( + token_with_user.token.permissions.as_ref(), + ) + .map_err(|_| ManagementTokenAuthError::Invalid)? + .unwrap_or_else(crate::control::legacy_full_management_token_permissions); + + Ok(AuthenticatedManagementToken { + token: token_with_user.token, + user, + permissions, + verified_token_hash: token_hash.clone(), + }) +} + +fn extract_unique_management_token_bearer( + headers: &HeaderMap, +) -> Result<&str, ManagementTokenAuthError> { + let mut values = headers.get_all(http::header::AUTHORIZATION).iter(); + let value = values.next().ok_or(ManagementTokenAuthError::Missing)?; + if values.next().is_some() { + return Err(ManagementTokenAuthError::Invalid); + } + let value = value + .to_str() + .map_err(|_| ManagementTokenAuthError::Invalid)?; + let (scheme, token) = value + .split_once(' ') + .ok_or(ManagementTokenAuthError::Invalid)?; + let token = token.trim(); + if !scheme.eq_ignore_ascii_case("bearer") + || token.is_empty() + || token.len() > MANAGEMENT_TOKEN_MAX_LEN + || (!token.starts_with("ae-") && !token.starts_with("ae_")) + || token.chars().any(char::is_whitespace) + { + return Err(ManagementTokenAuthError::Invalid); + } + Ok(token) +} + +fn hash_management_token(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} + +#[cfg(test)] +mod tests { + use super::{ + extract_unique_management_token_bearer, ManagementTokenAuthError, + VerifiedManagementTokenHash, + }; + use axum::http::{header, HeaderMap, HeaderValue}; + + #[test] + fn extracts_one_case_insensitive_bearer_value() { + let mut headers = HeaderMap::new(); + headers.insert( + header::AUTHORIZATION, + HeaderValue::from_static("bEaReR ae-valid-token"), + ); + assert_eq!( + extract_unique_management_token_bearer(&headers), + Ok("ae-valid-token") + ); + } + + #[test] + fn rejects_duplicate_or_ambiguous_authorization() { + let mut headers = HeaderMap::new(); + headers.append( + header::AUTHORIZATION, + HeaderValue::from_static("Bearer ae-first"), + ); + headers.append( + header::AUTHORIZATION, + HeaderValue::from_static("Bearer ae-second"), + ); + assert_eq!( + extract_unique_management_token_bearer(&headers), + Err(ManagementTokenAuthError::Invalid) + ); + + let mut headers = HeaderMap::new(); + headers.insert( + header::AUTHORIZATION, + HeaderValue::from_static("Bearer ae-valid extra"), + ); + assert_eq!( + extract_unique_management_token_bearer(&headers), + Err(ManagementTokenAuthError::Invalid) + ); + } + + #[test] + fn verified_management_token_hash_debug_is_redacted() { + let hash = VerifiedManagementTokenHash::new("credential-equivalent-hash".to_string()); + let rendered = format!("{hash:?}"); + assert!(rendered.contains("[REDACTED]")); + assert!(!rendered.contains("credential-equivalent-hash")); + } +} diff --git a/apps/aether-gateway/src/middleware/mod.rs b/apps/aether-gateway/src/middleware/mod.rs index 070e4b080..04c29fb55 100644 --- a/apps/aether-gateway/src/middleware/mod.rs +++ b/apps/aether-gateway/src/middleware/mod.rs @@ -5,6 +5,6 @@ pub(crate) use access_log::{ access_log_middleware, sanitize_access_log_path, should_downgrade_access_log, GatewayRequestAcceptedAt, RequestLogEmitted, }; +pub(crate) use aether_gateway_frontdoor::apply_cf_header_stripping; pub use aether_gateway_frontdoor::strip_cf_headers_middleware; -pub(crate) use aether_gateway_frontdoor::{apply_cf_header_stripping, CfConnectingIp}; pub(crate) use frontdoor_cors::frontdoor_cors_middleware; diff --git a/apps/aether-gateway/src/model_fetch/catalog.rs b/apps/aether-gateway/src/model_fetch/catalog.rs index c4ab9347b..922c8eee8 100644 --- a/apps/aether-gateway/src/model_fetch/catalog.rs +++ b/apps/aether-gateway/src/model_fetch/catalog.rs @@ -1494,6 +1494,133 @@ pub(crate) async fn read_recent_codex_catalog_client_version( (!normalized.used_fallback()).then_some(normalized.value) } +/// Management is not tied to a downstream client's compatibility version. Keep its +/// directory at least as new as the built-in fingerprint and successful catalogs. +pub(crate) struct CodexManagementCatalog { + pub(crate) client_version: String, + pub(crate) models: Option>, + target: CodexCatalogTarget, +} + +pub(crate) async fn read_codex_management_catalog( + runtime: &R, + provider_id: &str, + key_id: &str, +) -> Option +where + R: CodexCatalogRuntime + ?Sized, +{ + let target = bind_codex_catalog_target( + runtime, + &CodexCatalogTarget { + identity: CodexCatalogIdentity::new(provider_id, key_id), + endpoint_ids: Vec::new(), + credential_scope: None, + }, + ) + .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()?; + if let Some(recent) = + read_recent_codex_catalog_client_version(state, provider_id, key_id, scope).await + { + version = version.max(Version::parse(&recent).ok()?); + } + // The most recently seen client can be older than an already successful catalog. + // Only consider the current credential generation's bounded success index. + for member in state + .score_range_by_min(&catalog_versions_key(&target.identity), 0.0) + .await + .unwrap_or_default() + { + if let Some((stored_scope, stored_version)) = parse_catalog_version_member(&member) { + if stored_scope == scope { + if let Ok(candidate) = Version::parse(stored_version) { + version = version.max(candidate); + } + } + } + } + let client_version = version.to_string(); + let models = if let Some(snapshot) = read_lkg_snapshot(runtime, &target, &client_version).await + { + let fresh = state + .kv_get(&catalog_fresh_key(&target, &client_version)) + .await + .ok() + .flatten(); + (fresh.as_deref() == Some(snapshot.content_sha256.as_str())).then_some(snapshot.models) + } else { + None + }; + if !codex_catalog_credential_scope_is_current(runtime, &target).await { + return None; + } + Some(CodexManagementCatalog { + client_version, + models, + target, + }) +} + +/// Publish a successful management fetch into the same versioned directory used by +/// clients. A credential replacement while the request is in flight must not leak +/// the previous account's catalog into the new account. +pub(crate) async fn store_codex_management_catalog( + runtime: &R, + catalog: &CodexManagementCatalog, + transports: &[GatewayProviderTransportSnapshot], + models: Vec, + etag: Option<&str>, +) where + R: CodexCatalogRuntime + ?Sized, +{ + let target = &catalog.target; + if transports.is_empty() + || transports.iter().any(|transport| { + codex_catalog_credential_scope_from_transport(transport).as_deref() + != target.credential_scope() + }) + || !codex_catalog_credential_scope_is_current(runtime, target).await + { + return; + } + let Ok(serialized) = validate_catalog_models(&models) else { + return; + }; + let now = current_unix_secs(); + let snapshot = CodexCatalogSnapshot { + schema_version: CODEX_CATALOG_SCHEMA_VERSION, + provider_id: target.identity.provider_id.clone(), + key_id: target.identity.key_id.clone(), + credential_scope: target.credential_scope().unwrap_or_default().to_string(), + client_version: catalog.client_version.clone(), + models, + etag: etag.and_then(normalize_etag), + fetched_at_unix_secs: now, + last_checked_at_unix_secs: now, + content_sha256: sha256_hex(&serialized), + }; + persist_catalog_success( + runtime, + target, + &catalog.client_version, + &snapshot, + Some(200), + Duration::ZERO, + ) + .await; + if !codex_catalog_credential_scope_is_current(runtime, target).await { + discard_catalog_success_after_scope_change( + runtime.codex_catalog_runtime_state(), + target, + &catalog.client_version, + ) + .await; + } +} + async fn retain_catalog_version( state: &RuntimeState, target: &CodexCatalogTarget, @@ -2327,6 +2454,101 @@ mod tests { assert!(!prerelease.used_fallback()); } + #[tokio::test] + async fn management_catalog_uses_version_floor_and_newest_success_not_last_client() { + let runtime = TestRuntime::new(vec![successful_execution("gpt-new", "etag")]); + remember_seen_version(&runtime.state, &target(), "0.144.1").await; + 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!(initial.models.is_none()); + + seed_catalog(&runtime, &version("0.200.0")).await; + remember_seen_version(&runtime.state, &target(), "0.144.1").await; + let current = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) + .await + .expect("management context"); + assert_eq!(current.client_version, "0.200.0"); + assert_eq!(current.models.unwrap()[0]["slug"], "gpt-new"); + assert_eq!( + runtime.execution_count(), + 1, + "management read does not fetch upstream" + ); + + runtime + .state + .kv_delete(&catalog_fresh_key(&target(), "0.200.0")) + .await + .unwrap(); + let stale = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) + .await + .unwrap(); + assert_eq!(stale.client_version, "0.200.0"); + assert!( + stale.models.is_none(), + "expired LKG must not become a fresh admin cache" + ); + + runtime.set_credential_generation(TEST_CREDENTIAL_GENERATION_B); + 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!(rebound.models.is_none()); + } + + #[tokio::test] + async fn management_refresh_updates_shared_catalog_but_rejects_replaced_credentials() { + let runtime = TestRuntime::new(vec![]); + let context = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) + .await + .unwrap(); + let transports = vec![sample_codex_transport()]; + for slug in ["gpt-old", "gpt-6-astra"] { + store_codex_management_catalog( + &runtime, + &context, + &transports, + vec![codex_model(slug)], + Some("etag"), + ) + .await; + let shared = read_lkg_snapshot(&runtime, &target(), &context.client_version) + .await + .unwrap(); + assert_eq!(shared.models[0]["slug"], slug); + let admin = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) + .await + .unwrap(); + assert_eq!(admin.models.unwrap()[0]["slug"], slug); + } + runtime.set_credential_generation(TEST_CREDENTIAL_GENERATION_B); + store_codex_management_catalog( + &runtime, + &context, + &transports, + vec![codex_model("gpt-leaked")], + None, + ) + .await; + let rebound = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) + .await + .unwrap(); + assert!(rebound.models.is_none()); + let old = read_raw_snapshot(&runtime, &target(), &context.client_version) + .await + .unwrap(); + assert_eq!(old.models[0]["slug"], "gpt-6-astra"); + } + #[test] fn invalid_or_oversized_client_versions_use_bounded_fallback_identity() { for raw in [ diff --git a/apps/aether-gateway/src/model_fetch/mod.rs b/apps/aether-gateway/src/model_fetch/mod.rs index 545078399..e5438473a 100644 --- a/apps/aether-gateway/src/model_fetch/mod.rs +++ b/apps/aether-gateway/src/model_fetch/mod.rs @@ -6,12 +6,12 @@ mod tests; pub(crate) use aether_model_fetch::ModelFetchRunSummary; pub(crate) use catalog::{ codex_catalog_credential_scope_from_stored_key, codex_catalog_targets, load_codex_catalogs, - normalize_codex_client_version, read_recent_codex_catalog_client_version, - refresh_codex_catalog_target, CodexCatalogLoad, CodexCatalogRuntime, CodexCatalogTarget, + normalize_codex_client_version, read_codex_management_catalog, refresh_codex_catalog_target, + store_codex_management_catalog, CodexCatalogLoad, CodexCatalogRuntime, CodexCatalogTarget, NormalizedCodexClientVersion, }; pub(crate) use runtime::state::ModelFetchRuntimeState; pub(crate) use runtime::{ perform_model_fetch_for_key, perform_model_fetch_for_keys, perform_model_fetch_once, - spawn_model_fetch_worker, + safe_model_fetch_error, spawn_model_fetch_worker, }; diff --git a/apps/aether-gateway/src/model_fetch/runtime.rs b/apps/aether-gateway/src/model_fetch/runtime.rs index fc032d24f..cce13f195 100644 --- a/apps/aether-gateway/src/model_fetch/runtime.rs +++ b/apps/aether-gateway/src/model_fetch/runtime.rs @@ -7,7 +7,7 @@ use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use aether_model_fetch::{ - apply_model_filters, fetch_models_from_transports_for_client_version, json_string_list, + apply_model_filters, fetch_models_from_transports_for_management, json_string_list, model_catalog_upstream_metadata, model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_startup_enabled, preset_models_for_provider, selected_models_fetch_endpoints, sync_provider_model_whitelist_associations, @@ -44,7 +44,10 @@ pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option Option>(); let mut endpoints_by_provider = HashMap::>::new(); for endpoint in state - .list_provider_catalog_endpoints_by_provider_ids(&provider_ids) + .list_provider_catalog_endpoints_for_model_fetch(&provider_ids) .await? { endpoints_by_provider @@ -147,12 +153,10 @@ where .push(endpoint); } let mut keys_by_provider = HashMap::>::new(); - for key in ::list_provider_catalog_keys_by_provider_ids( - state, - &provider_ids, - ) - .await - .map_err(GatewayError::Internal)? + for key in state + .list_provider_catalog_keys_for_model_fetch(&provider_ids) + .await + .map_err(GatewayError::Internal)? { keys_by_provider .entry(key.provider_id.clone()) @@ -164,8 +168,12 @@ where for provider in providers { let endpoints = endpoints_by_provider .remove(&provider.id) - .unwrap_or_default(); + .unwrap_or_default() + .into_iter() + .map(sanitize_model_fetch_endpoint) + .collect::>(); let keys = keys_by_provider.remove(&provider.id).unwrap_or_default(); + let provider = sanitize_model_fetch_provider(provider); for key in keys { if key_id_filter.is_some_and(|key_ids| !key_ids.contains(&key.id)) { continue; @@ -174,6 +182,7 @@ where continue; } let selected_endpoints = selected_models_fetch_endpoints(&endpoints, &key); + let key = sanitize_model_fetch_key(key); targets.push(SelectedFetchTarget { provider: provider.clone(), key, @@ -184,6 +193,82 @@ where Ok(targets) } +/// Keep only the key metadata needed after target collection. Raw catalog +/// rows contain encrypted credentials and transport secrets; model discovery +/// reopens a single snapshot by id when it actually needs to make a request. +fn sanitize_model_fetch_key(mut key: StoredProviderCatalogKey) -> StoredProviderCatalogKey { + // `SelectedFetchTarget` lives across endpoint selection and the complete + // fetch/persist operation. Keep only fields consumed by that operation; + // in particular, do not retain historical diagnostics, scheduling state, + // usage counters, or transport configuration copied from a raw database + // row. The actual credential/proxy snapshot is reopened by id for one + // endpoint at a time. + key.capabilities = None; + key.auth_type_by_format = None; + key.allow_auth_channel_mismatch_formats = None; + key.encrypted_api_key = None; + key.encrypted_auth_config = None; + key.note = None; + key.internal_priority = 0; + key.rate_multipliers = None; + key.global_priority_by_format = None; + key.expires_at_unix_secs = None; + key.cache_ttl_minutes = 0; + key.max_probe_interval_minutes = 0; + key.proxy = None; + key.fingerprint = None; + key.rpm_limit = None; + key.concurrent_limit = None; + key.learned_rpm_limit = None; + key.concurrent_429_count = None; + key.rpm_429_count = None; + key.last_429_at_unix_secs = None; + key.last_429_type = None; + key.adjustment_history = None; + key.utilization_samples = None; + key.last_probe_increase_at_unix_secs = None; + key.last_rpm_peak = None; + key.request_count = None; + key.total_tokens = 0; + key.total_cost_usd = 0.0; + key.success_count = None; + key.error_count = None; + key.total_response_time_ms = None; + key.last_used_at_unix_secs = None; + key.last_models_fetch_at_unix_secs = None; + key.last_models_fetch_error = None; + key.oauth_invalid_at_unix_secs = None; + key.oauth_invalid_reason = None; + key.status_snapshot = None; + key +} + +/// Keep only non-secret provider metadata in a background fetch target. The +/// authoritative transport snapshot is reopened by ID immediately before a +/// request, so carrying stored proxy/config JSON here would needlessly retain +/// credentials and could expose malformed historical values to later stages. +fn sanitize_model_fetch_provider( + mut provider: StoredProviderCatalogProvider, +) -> StoredProviderCatalogProvider { + provider.proxy = None; + provider.config = None; + provider +} + +/// Endpoint selection needs only activity, format, and identity. Clear +/// transport rules/proxy data because those are reloaded from the snapshot +/// just before execution. +fn sanitize_model_fetch_endpoint( + mut endpoint: StoredProviderCatalogEndpoint, +) -> StoredProviderCatalogEndpoint { + endpoint.header_rules = None; + endpoint.body_rules = None; + endpoint.config = None; + endpoint.format_acceptance_config = None; + endpoint.proxy = None; + endpoint +} + async fn execute_fetch_targets( state: &S, targets: Vec, @@ -287,13 +372,14 @@ async fn fetch_and_persist_key_models( } let mut transports = Vec::new(); + let mut skipped_invalid_credential = false; for endpoint in &target.endpoints { match state .read_provider_transport_snapshot(&target.provider.id, &endpoint.id, &target.key.id) - .await? + .await { - Some(transport) => transports.push(transport), - None => { + Ok(Some(transport)) => transports.push(transport), + Ok(None) => { warn!( provider_id = %target.provider.id, endpoint_id = %endpoint.id, @@ -301,9 +387,29 @@ async fn fetch_and_persist_key_models( "gateway model fetch transport snapshot unavailable" ); } + Err(error) if is_nonfatal_legacy_credential_error(&error) => { + skipped_invalid_credential = true; + warn!( + event_name = "model_fetch_skipped_invalid_credential", + log_type = "ops", + provider_id = %target.provider.id, + endpoint_id = %endpoint.id, + key_id = %target.key.id, + reason = "invalid_stored_credential", + "gateway skipped model fetch for an invalid stored credential" + ); + } + Err(error) => return Err(error), } } + // A malformed legacy credential is isolated to its key. Do not turn it + // into a cycle-wide failure (or rewrite the row merely to record a fetch + // error), and let other eligible keys continue through the worker. + if transports.is_empty() && skipped_invalid_credential { + return Ok(KeyFetchDisposition::Skipped); + } + if transports.is_empty() { persist_key_fetch_failure( state, @@ -327,7 +433,7 @@ async fn fetch_and_persist_key_models( } else { None }; - let result = match fetch_models_from_transports_for_client_version( + let result = match fetch_models_from_transports_for_management( state, &transports, codex_client_version.as_deref(), @@ -336,11 +442,13 @@ async fn fetch_and_persist_key_models( { Ok(result) => result, Err(err) => { - persist_key_fetch_failure(state, &target.key, now_unix_secs, err.clone()).await?; + let safe_error = safe_model_fetch_error(&err); + persist_key_fetch_failure(state, &target.key, now_unix_secs, safe_error.clone()) + .await?; warn!( provider_id = %target.provider.id, key_id = %target.key.id, - message = %err, + error = %safe_error, "gateway model fetch failed" ); return Ok(KeyFetchDisposition::Failed); @@ -353,11 +461,12 @@ async fn fetch_and_persist_key_models( } else { result.errors.join("; ") }; - persist_key_fetch_failure(state, &target.key, now_unix_secs, error.clone()).await?; + let safe_error = safe_model_fetch_error(&error); + persist_key_fetch_failure(state, &target.key, now_unix_secs, safe_error.clone()).await?; warn!( provider_id = %target.provider.id, key_id = %target.key.id, - message = %error, + error = %safe_error, "gateway model fetch failed" ); return Ok(KeyFetchDisposition::Failed); @@ -392,18 +501,144 @@ async fn persist_key_fetch_failure( now_unix_secs: u64, error: String, ) -> Result<(), GatewayError> { + let safe_error = safe_model_fetch_error(&error); state .update_provider_catalog_key_model_fetch_state( &key.id, key.allowed_models.as_ref(), Some(now_unix_secs), - Some(&error), + Some(&safe_error), Some(now_unix_secs), ) .await?; Ok(()) } +pub(crate) fn safe_model_fetch_error(error: &str) -> String { + let trimmed = error.trim(); + match trimmed { + "No supported endpoint for Rust models fetch" + | "Provider transport snapshot unavailable" => return trimmed.to_string(), + _ => {} + } + + let lower = trimmed.to_ascii_lowercase(); + if let Some(status) = model_fetch_error_http_status(&lower) { + return match status { + 401 => "Upstream models fetch authentication failed (status 401)".to_string(), + 403 => "Upstream models fetch authorization failed (status 403)".to_string(), + 404 => "Upstream models fetch endpoint not found (status 404)".to_string(), + 408 => "Upstream models fetch timed out (status 408)".to_string(), + 429 => "Upstream models fetch rate limited (status 429)".to_string(), + _ => format!("Upstream models fetch failed (status {status})"), + }; + } + if lower.contains("unauthorized") + || lower.contains("authentication failed") + || lower.contains("invalid api key") + || lower.contains("invalid token") + { + return "Upstream models fetch authentication failed".to_string(); + } + if lower.contains("forbidden") || lower.contains("authorization failed") { + return "Upstream models fetch authorization failed".to_string(); + } + if lower.contains("rate limit") || lower.contains("too many requests") { + return "Upstream models fetch rate limited".to_string(); + } + if lower.contains("timeout") || lower.contains("timed out") { + return "Upstream models fetch timed out".to_string(); + } + if lower.contains("missing api key") + || lower.contains("missing access token") + || (lower.contains("requires") && lower.contains("auth")) + { + return "Provider credentials unavailable for models fetch".to_string(); + } + if lower.contains("private_key") + || lower.contains("auth_config") + || lower.contains("configuration") + { + return "Provider models fetch configuration is invalid".to_string(); + } + if lower.contains("response body") + || lower.contains("json") + || lower.contains("malformed") + || lower.contains("parse") + || lower.contains("no models") + || lower.contains("invalid response") + { + return "Upstream models fetch response was invalid".to_string(); + } + if lower.contains("connect") + || lower.contains("connection") + || lower.contains("network") + || lower.contains("dns") + || lower.contains("tls") + || lower.contains("certificate") + { + return "Upstream models fetch connection failed".to_string(); + } + "Upstream models fetch failed".to_string() +} + +/// Credential decoding failures from old catalog rows are isolated by the +/// background model-fetch worker. Normal request/admin paths remain fail +/// closed; this predicate only controls whether one maintenance item may be +/// skipped without aborting the whole cycle. +fn is_nonfatal_legacy_credential_error(error: &GatewayError) -> bool { + let GatewayError::Internal(message) = error else { + return false; + }; + let message = message.to_ascii_lowercase(); + // Missing encryption configuration is an operational failure and must + // remain fail-closed. Only errors that identify a stored field or a + // malformed legacy ciphertext are safe to isolate to one key. + if message.contains("encryption key is not configured") { + return false; + } + message.contains("provider_api_keys.api_key") + || message.contains("provider_api_keys.auth_config") + || message.contains("stored provider proxy credentials cannot be decrypted") + || message.contains("stored endpoint proxy credentials cannot be decrypted") + || message.contains("stored key proxy credentials cannot be decrypted") + || message.contains("stored provider proxy changed during credential migration") + || message.contains("stored endpoint proxy changed during credential migration") + || message.contains("stored key changed during credential migration") + || message.contains("stored provider proxy credential migration did not stabilize") + || message.contains("stored endpoint proxy credential migration did not stabilize") + || message.contains("stored key proxy credential migration did not stabilize") + || message.contains("legacy provider catalog credential") + || message.contains("provider catalog credential is not an authenticated ciphertext") + || message.contains("provider catalog credential contains reserved framing") + || message.contains("provider catalog credential authentication failed") + || message.contains("provider catalog credential envelope") +} + +fn model_fetch_error_http_status(error: &str) -> Option { + [ + "http ", + "status ", + "status=", + "status:", + "status_code=", + "status_code:", + ] + .into_iter() + .find_map(|marker| { + let suffix = error.split_once(marker)?.1.trim_start(); + let digits = suffix + .bytes() + .take_while(u8::is_ascii_digit) + .take(3) + .collect::>(); + (digits.len() == 3) + .then(|| std::str::from_utf8(&digits).ok()?.parse::().ok()) + .flatten() + .filter(|status| (400..600).contains(status)) + }) +} + async fn persist_key_fetch_success( state: &(impl ModelFetchRuntimeState + ?Sized), key: &StoredProviderCatalogKey, @@ -450,7 +685,10 @@ fn now_unix_secs() -> u64 { #[cfg(test)] mod tests { - use super::{perform_model_fetch_once_with_state, state::ModelFetchRuntimeState}; + use super::{ + perform_model_fetch_once_with_state, safe_model_fetch_error, sanitize_model_fetch_key, + state::ModelFetchRuntimeState, + }; use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot}; use aether_data_contracts::repository::global_models::{ AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModelPage, @@ -481,6 +719,7 @@ mod tests { endpoints: Arc>, keys: Arc>>, transports: Arc>, + transport_errors: Arc>, execution_results: Arc>>, executed_plans: Arc>>, cached_models: Arc>>>, @@ -500,6 +739,7 @@ mod tests { endpoints: Arc::new(endpoints), keys: Arc::new(Mutex::new(keys)), transports: Arc::new(transports), + transport_errors: Arc::new(HashMap::new()), execution_results: Arc::new(Mutex::new(VecDeque::from(execution_results))), executed_plans: Arc::new(Mutex::new(Vec::new())), cached_models: Arc::new(Mutex::new(HashMap::new())), @@ -507,6 +747,14 @@ mod tests { } } + fn with_transport_errors( + mut self, + transport_errors: HashMap<(String, String, String), String>, + ) -> Self { + self.transport_errors = Arc::new(transport_errors); + self + } + fn key(&self, key_id: &str) -> StoredProviderCatalogKey { self.keys .lock() @@ -654,6 +902,13 @@ mod tests { endpoint_id: &str, key_id: &str, ) -> Result, GatewayError> { + if let Some(error) = self.transport_errors.get(&( + provider_id.to_string(), + endpoint_id.to_string(), + key_id.to_string(), + )) { + return Err(GatewayError::Internal(error.clone())); + } Ok(self .transports .get(&( @@ -880,10 +1135,14 @@ mod tests { } fn execution_result(body: Value) -> ExecutionResult { + execution_result_with_status(200, body) + } + + fn execution_result_with_status(status_code: u16, body: Value) -> ExecutionResult { ExecutionResult { request_id: "req-1".to_string(), candidate_id: None, - status_code: 200, + status_code, headers: Default::default(), response_observation: None, body: Some(aether_contracts::ResponseBody { @@ -895,6 +1154,56 @@ mod tests { } } + #[test] + fn model_fetch_error_projection_discards_transport_credentials_and_urls() { + let error = "connection failed for https://user:password@example.test/v1/models?key=\ + query-secret: Authorization: Bearer transport-secret-token-value"; + + let safe_error = safe_model_fetch_error(error); + + assert_eq!(safe_error, "Upstream models fetch connection failed"); + for secret in [ + "user", + "password", + "query-secret", + "transport-secret-token-value", + "Bearer", + "example.test", + ] { + assert!(!safe_error.contains(secret)); + } + } + + #[test] + fn model_fetch_error_projection_discards_unclassified_details() { + let error = + "opaque failure at https://user:password@example.test/private?key=query-secret; \ + Authorization: Bearer transport-secret-token-value"; + + let safe_error = safe_model_fetch_error(error); + + assert_eq!(safe_error, "Upstream models fetch failed"); + for secret in [ + "user", + "password", + "query-secret", + "transport-secret-token-value", + "Bearer", + "example.test", + ] { + assert!(!safe_error.contains(secret)); + } + } + + #[test] + fn invalid_credential_classifier_does_not_swallow_missing_key_configuration() { + assert!(!super::is_nonfatal_legacy_credential_error( + &GatewayError::Internal( + "provider catalog credential encryption key is not configured".to_string(), + ) + )); + } + #[tokio::test] async fn gateway_runtime_state_supports_shared_models_fetch_plan_builder() { let state = TestState::default(); @@ -1233,4 +1542,211 @@ mod tests { Some("Provider transport snapshot unavailable") ); } + + #[tokio::test] + async fn model_fetch_isolates_malformed_legacy_key_from_healthy_key() { + let provider = sample_provider("provider-openai", "openai"); + let endpoint = sample_endpoint( + "endpoint-openai-responses", + "provider-openai", + "openai:responses", + ); + let mut malformed = sample_key( + "key-openai-malformed", + "provider-openai", + "api_key", + &["openai:responses"], + ); + malformed.encrypted_api_key = Some("legacy-plaintext-or-corrupt".to_string()); + malformed.allowed_models = Some(json!(["legacy-model"])); + malformed.last_models_fetch_error = Some("previous error".to_string()); + let malformed_ciphertext = malformed.encrypted_api_key.clone(); + let healthy = sample_key( + "key-openai-healthy", + "provider-openai", + "api_key", + &["openai:responses"], + ); + let healthy_transport = sample_transport( + "openai", + "provider-openai", + "endpoint-openai-responses", + "key-openai-healthy", + "openai:responses", + "api_key", + None, + ); + let state = TestState::new( + vec![provider], + vec![endpoint], + vec![malformed, healthy], + HashMap::from([( + ( + "provider-openai".to_string(), + "endpoint-openai-responses".to_string(), + "key-openai-healthy".to_string(), + ), + healthy_transport, + )]), + vec![execution_result(json!({ + "data": [{"id": "gpt-healthy"}] + }))], + ) + .with_transport_errors(HashMap::from([( + ( + "provider-openai".to_string(), + "endpoint-openai-responses".to_string(), + "key-openai-malformed".to_string(), + ), + "provider_api_keys.api_key is not an authenticated ciphertext".to_string(), + )])); + + let summary = perform_model_fetch_once_with_state(&state) + .await + .expect("one malformed key must not abort the cycle"); + + assert_eq!(summary.attempted, 2); + assert_eq!(summary.skipped, 1); + assert_eq!(summary.succeeded, 1); + let malformed_after = state.key("key-openai-malformed"); + assert_eq!(malformed_after.encrypted_api_key, malformed_ciphertext); + assert_eq!( + malformed_after.allowed_models, + Some(json!(["legacy-model"])) + ); + assert_eq!( + malformed_after.last_models_fetch_error.as_deref(), + Some("previous error") + ); + assert_eq!( + state.key("key-openai-healthy").allowed_models, + Some(json!(["gpt-healthy"])) + ); + } + + #[test] + fn sanitized_model_fetch_key_drops_raw_transport_and_diagnostic_state() { + let mut key = sample_key("key-sanitize", "provider", "api_key", &["openai:responses"]); + key.capabilities = Some(json!({"secret": "capability"})); + key.note = Some("operator note".to_string()); + key.proxy = Some(json!({"url": "http://user:pass@example.test"})); + key.last_models_fetch_error = Some("upstream detail".to_string()); + key.oauth_invalid_reason = Some("token detail".to_string()); + key.allowed_models = Some(json!(["keep-model-filter"])); + key.upstream_metadata = Some(json!({"provider": {"quota": 1}})); + + let sanitized = sanitize_model_fetch_key(key); + + assert_eq!(sanitized.encrypted_api_key, None); + assert_eq!(sanitized.encrypted_auth_config, None); + assert_eq!(sanitized.proxy, None); + assert_eq!(sanitized.fingerprint, None); + assert_eq!(sanitized.note, None); + assert_eq!(sanitized.last_models_fetch_error, None); + assert_eq!(sanitized.oauth_invalid_reason, None); + assert_eq!(sanitized.allowed_models, Some(json!(["keep-model-filter"]))); + assert_eq!( + sanitized.upstream_metadata, + Some(json!({"provider": {"quota": 1}})) + ); + } + + #[test] + fn sanitized_model_fetch_provider_and_endpoint_drop_transport_secrets() { + let mut provider = sample_provider("provider-sanitize", "openai"); + provider.proxy = Some(json!({"url": "http://user:pass@example.test"})); + provider.config = Some(json!({"api_key": "provider-secret"})); + let sanitized_provider = super::sanitize_model_fetch_provider(provider); + assert_eq!(sanitized_provider.proxy, None); + assert_eq!(sanitized_provider.config, None); + + let mut endpoint = + sample_endpoint("endpoint-sanitize", "provider-sanitize", "openai:responses"); + endpoint.header_rules = Some(json!({"authorization": "Bearer endpoint-secret"})); + endpoint.body_rules = Some(json!({"token": "endpoint-secret"})); + endpoint.config = Some(json!({"password": "endpoint-secret"})); + endpoint.format_acceptance_config = Some(json!({"secret": "endpoint-secret"})); + endpoint.proxy = Some(json!({"url": "http://user:pass@example.test"})); + let sanitized_endpoint = super::sanitize_model_fetch_endpoint(endpoint); + assert_eq!(sanitized_endpoint.header_rules, None); + assert_eq!(sanitized_endpoint.body_rules, None); + assert_eq!(sanitized_endpoint.config, None); + assert_eq!(sanitized_endpoint.format_acceptance_config, None); + assert_eq!(sanitized_endpoint.proxy, None); + assert_eq!(sanitized_endpoint.id, "endpoint-sanitize"); + assert_eq!(sanitized_endpoint.api_format, "openai:responses"); + } + + #[tokio::test] + async fn model_fetch_failure_does_not_persist_upstream_error_body_credentials() { + const UPSTREAM_SECRET: &str = "upstream-secret-token-value"; + let provider = sample_provider("provider-openai", "openai"); + let endpoint = sample_endpoint( + "endpoint-openai-responses", + "provider-openai", + "openai:responses", + ); + let key = sample_key( + "key-openai-responses", + "provider-openai", + "api_key", + &["openai:responses"], + ); + let transport = sample_transport( + "openai", + "provider-openai", + "endpoint-openai-responses", + "key-openai-responses", + "openai:responses", + "api_key", + None, + ); + let state = TestState::new( + vec![provider], + vec![endpoint], + vec![key], + HashMap::from([( + ( + "provider-openai".to_string(), + "endpoint-openai-responses".to_string(), + "key-openai-responses".to_string(), + ), + transport, + )]), + vec![execution_result_with_status( + 401, + json!({ + "error": { + "message": format!( + "Authorization: Bearer {UPSTREAM_SECRET}; api_key=query-secret; \ + config=/srv/provider/private.json" + ) + } + }), + )], + ); + + let summary = perform_model_fetch_once_with_state(&state) + .await + .expect("fetch should finish with a projected failure"); + + assert_eq!(summary.failed, 1); + let persisted = state + .key("key-openai-responses") + .last_models_fetch_error + .expect("safe failure should be persisted"); + assert_eq!( + persisted, + "Upstream models fetch authentication failed (status 401)" + ); + for secret in [ + UPSTREAM_SECRET, + "query-secret", + "/srv/provider/private.json", + "Bearer", + "api_key", + ] { + assert!(!persisted.contains(secret)); + } + } } diff --git a/apps/aether-gateway/src/model_fetch/runtime/state.rs b/apps/aether-gateway/src/model_fetch/runtime/state.rs index e90d720fa..a05820414 100644 --- a/apps/aether-gateway/src/model_fetch/runtime/state.rs +++ b/apps/aether-gateway/src/model_fetch/runtime/state.rs @@ -26,11 +26,46 @@ pub(crate) trait ModelFetchRuntimeState: active_only: bool, ) -> Result, GatewayError>; + /// Return provider rows for the background fetcher without opening or + /// migrating stored proxy credentials. Production implementations should + /// use the raw repository projection so one malformed historical row does + /// not abort the entire cycle and a read does not rewrite old data. + async fn list_provider_catalog_providers_for_model_fetch( + &self, + active_only: bool, + ) -> Result, GatewayError> { + self.list_provider_catalog_providers(active_only).await + } + async fn list_provider_catalog_endpoints_by_provider_ids( &self, provider_ids: &[String], ) -> Result, GatewayError>; + /// Raw endpoint counterpart to + /// [`Self::list_provider_catalog_providers_for_model_fetch`]. + async fn list_provider_catalog_endpoints_for_model_fetch( + &self, + provider_ids: &[String], + ) -> Result, GatewayError> { + self.list_provider_catalog_endpoints_by_provider_ids(provider_ids) + .await + } + + /// Return raw catalog key rows for the background fetcher. Production + /// implementations should avoid the normal bulk credential-opening + /// wrapper here: one malformed legacy row must not prevent healthy keys + /// from being considered, and a maintenance scan must not lazily rewrite + /// historical ciphertext. Test implementations can use the association + /// store's existing method via this default. + async fn list_provider_catalog_keys_for_model_fetch( + &self, + provider_ids: &[String], + ) -> Result, String> { + self.list_provider_catalog_keys_by_provider_ids(provider_ids) + .await + } + async fn read_provider_transport_snapshot( &self, provider_id: &str, diff --git a/apps/aether-gateway/src/oauth/http_executor.rs b/apps/aether-gateway/src/oauth/http_executor.rs index bfe43bf0d..a00843a98 100644 --- a/apps/aether-gateway/src/oauth/http_executor.rs +++ b/apps/aether-gateway/src/oauth/http_executor.rs @@ -1,16 +1,43 @@ use crate::admin_api::AdminAppState; use crate::{AppState, GatewayError}; use aether_contracts::{ - ExecutionPlan, ExecutionResult, ExecutionTimeouts, RequestBody, + ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, }; use aether_oauth::core::OAuthError; -use aether_oauth::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse}; +use aether_oauth::network::{ + OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkPolicy, OAuthTimeouts, +}; use async_trait::async_trait; use base64::{engine::general_purpose::STANDARD, Engine as _}; use flate2::read::{DeflateDecoder, GzDecoder}; +use futures_util::StreamExt; +use reqwest::header::{HeaderMap, HeaderName, HeaderValue, CONTENT_ENCODING, CONTENT_TYPE}; use std::collections::BTreeMap; use std::io::Read; +use std::net::{IpAddr, SocketAddr}; +use std::time::Duration; + +const OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; + +// Some local DNS interception deployments return RFC 2544's benchmarking +// range (198.18.0.0/15) for well-known public services. Identity OAuth is a +// sensitive direct-connection path, so the range is never accepted globally: +// only exact, built-in public origins may use it. URL validation below still +// requires HTTPS, strips credentials/fragments, pins the resolved addresses, +// and disables redirects. +const TRUSTED_IDENTITY_BENCHMARKING_DNS_HOSTS: &[&str] = &[ + "accounts.google.com", + "auth.openai.com", + "claude.ai", + "connect.linux.do", + "connect.linuxdo.org", + "oauth2.googleapis.com", + "platform.claude.com", + "register.windsurf.com", + "server.self-serve.windsurf.com", + "windsurf.com", +]; #[derive(Clone)] pub(crate) struct GatewayOAuthHttpExecutor<'a> { @@ -37,57 +64,438 @@ impl<'a> GatewayOAuthHttpExecutor<'a> { #[async_trait] impl<'a> OAuthHttpExecutor for GatewayOAuthHttpExecutor<'a> { async fn execute(&self, request: OAuthHttpRequest) -> Result { - let body = if let Some(json_body) = request.json_body { - RequestBody::from_json(json_body) - } else { - RequestBody { - json_body: None, - body_bytes_b64: request.body_bytes.map(|bytes| STANDARD.encode(bytes)), - body_ref: None, + match request.network.policy { + OAuthNetworkPolicy::DirectOnly | OAuthNetworkPolicy::DirectOrSystemProxy => { + match identity_oauth_route(request.network.policy, request.network.proxy.as_ref())? + { + IdentityOAuthRoute::Direct => { + execute_direct_identity_oauth(&self.app, request).await + } + } } - }; - let timeouts = request.network.timeouts; - let mut headers = request.headers; + OAuthNetworkPolicy::ProviderOperationProxy => { + #[cfg(test)] + if self.app.execution_runtime_override_base_url().is_none() + && identity_oauth_endpoint_policy(&self.app, &request.url) + == IdentityOAuthEndpointPolicy::ExplicitTestLoopback + { + return execute_direct_identity_oauth(&self.app, request).await; + } + execute_provider_oauth_via_runtime(&self.app, request).await + } + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IdentityOAuthRoute { + Direct, +} + +fn identity_oauth_route( + policy: OAuthNetworkPolicy, + proxy: Option<&ProxySnapshot>, +) -> Result { + let Some(proxy) = proxy.filter(|proxy| proxy.enabled != Some(false)) else { + return Ok(IdentityOAuthRoute::Direct); + }; + + if policy == OAuthNetworkPolicy::DirectOnly { + return Err(OAuthError::transport( + "identity OAuth direct-only transport cannot use a proxy", + )); + } + + let has_proxy_url = proxy + .url + .as_deref() + .map(str::trim) + .is_some_and(|value| !value.is_empty()); + if has_proxy_url { + return Err(OAuthError::transport( + "identity OAuth HTTP/SOCKS proxies are disabled; use a controlled tunnel or direct transport", + )); + } + + Err(OAuthError::transport( + "identity OAuth cannot use a configured proxy or tunnel; use direct transport", + )) +} + +async fn execute_provider_oauth_via_runtime( + app: &AppState, + request: OAuthHttpRequest, +) -> Result { + let plan = oauth_execution_plan(request, false); + let result = crate::execution_runtime::execute_execution_runtime_sync_plan(app, None, &plan) + .await + .map_err(gateway_error_to_oauth_error)?; + Ok(execution_result_to_oauth_response(&result)) +} + +fn oauth_execution_plan(request: OAuthHttpRequest, force_disable_redirects: bool) -> ExecutionPlan { + let OAuthHttpRequest { + request_id, + method, + url, + mut headers, + content_type, + json_body, + body_bytes, + network, + transport_profile, + } = request; + let body = if let Some(json_body) = json_body { + RequestBody::from_json(json_body) + } else { + RequestBody { + json_body: None, + body_bytes_b64: body_bytes.map(|bytes| STANDARD.encode(bytes)), + body_ref: None, + } + }; + let timeouts = network.timeouts; + if force_disable_redirects { + headers.insert( + EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), + "false".to_string(), + ); + } else { headers .entry(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string()) .or_insert_with(|| "true".to_string()); - let plan = ExecutionPlan { - request_id: request.request_id, - candidate_id: None, - provider_name: Some("oauth".to_string()), - provider_id: String::new(), - endpoint_id: String::new(), - key_id: String::new(), - method: request.method.as_str().to_string(), - url: request.url, - headers, - content_type: request.content_type, - content_encoding: None, - body, - stream: false, - client_api_format: "oauth:exchange".to_string(), - provider_api_format: "oauth:exchange".to_string(), - model_name: Some("oauth-exchange".to_string()), - proxy: request.network.proxy, - transport_profile: request.transport_profile, - timeouts: Some(ExecutionTimeouts { - connect_ms: Some(timeouts.connect_ms), - read_ms: Some(timeouts.read_ms), - write_ms: Some(timeouts.write_ms), - pool_ms: Some(timeouts.connect_ms), - total_ms: Some(timeouts.total_ms), - ..ExecutionTimeouts::default() - }), - }; - let result = - crate::execution_runtime::execute_execution_runtime_sync_plan(&self.app, None, &plan) - .await - .map_err(gateway_error_to_oauth_error)?; - Ok(OAuthHttpResponse { - status_code: result.status_code, - body_text: execution_body_text(&result), - json_body: execution_json_body(&result), + } + let plan = ExecutionPlan { + request_id, + candidate_id: None, + provider_name: Some("oauth".to_string()), + provider_id: String::new(), + endpoint_id: String::new(), + key_id: String::new(), + method: method.as_str().to_string(), + url, + headers, + content_type, + content_encoding: None, + body, + stream: false, + client_api_format: "oauth:exchange".to_string(), + provider_api_format: "oauth:exchange".to_string(), + model_name: Some("oauth-exchange".to_string()), + proxy: network.proxy, + transport_profile, + timeouts: Some(ExecutionTimeouts { + connect_ms: Some(timeouts.connect_ms), + read_ms: Some(timeouts.read_ms), + write_ms: Some(timeouts.write_ms), + pool_ms: Some(timeouts.connect_ms), + total_ms: Some(timeouts.total_ms), + ..ExecutionTimeouts::default() + }), + }; + crate::execution_runtime::transport::with_upstream_response_body_limit( + &plan, + OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) +} + +#[derive(Debug)] +struct ResolvedIdentityOAuthEndpoint { + url: reqwest::Url, + host: String, + addrs: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IdentityOAuthEndpointPolicy { + PublicHttps, + #[cfg(test)] + ExplicitTestLoopback, +} + +fn identity_oauth_endpoint_policy(app: &AppState, raw_url: &str) -> IdentityOAuthEndpointPolicy { + #[cfg(test)] + { + let parsed_url = reqwest::Url::parse(raw_url).ok(); + let is_explicit_test_url = app + .provider_oauth_token_url_overrides + .lock() + .expect("provider oauth token URL overrides should lock") + .values() + .any(|value| { + parsed_url + .as_ref() + .is_some_and(|url| test_loopback_override_allows_url(value, url)) + }); + if is_explicit_test_url { + return IdentityOAuthEndpointPolicy::ExplicitTestLoopback; + } + } + + let _ = (app, raw_url); + IdentityOAuthEndpointPolicy::PublicHttps +} + +#[cfg(test)] +fn test_loopback_override_allows_url(override_url: &str, target: &reqwest::Url) -> bool { + let Ok(registered) = reqwest::Url::parse(override_url) else { + return false; + }; + let Some(target_ip) = target + .host_str() + .and_then(|host| host.parse::().ok()) + .filter(|address| address.is_loopback()) + else { + return false; + }; + let same_origin = registered.scheme() == target.scheme() + && registered + .host_str() + .and_then(|host| host.parse::().ok()) + == Some(target_ip) + && registered.port_or_known_default() == target.port_or_known_default(); + if !same_origin || registered.query().is_some() || registered.fragment().is_some() { + return registered == *target; + } + + let registered_path = registered.path().trim_end_matches('/'); + registered == *target + || registered_path.is_empty() + || target + .path() + .strip_prefix(registered_path) + .is_some_and(|suffix| suffix.starts_with('/')) +} + +fn parse_identity_oauth_endpoint( + raw_url: &str, + policy: IdentityOAuthEndpointPolicy, +) -> Result<(reqwest::Url, String, u16), OAuthError> { + let url = reqwest::Url::parse(raw_url) + .map_err(|_| OAuthError::transport("identity OAuth endpoint URL is invalid"))?; + if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() { + return Err(OAuthError::transport( + "identity OAuth endpoint must not contain credentials or a fragment", + )); + } + let host = url + .host_str() + .map(ToOwned::to_owned) + .ok_or_else(|| OAuthError::transport("identity OAuth endpoint is missing a host"))?; + match policy { + IdentityOAuthEndpointPolicy::PublicHttps if url.scheme() != "https" => { + return Err(OAuthError::transport( + "identity OAuth endpoint must use HTTPS", + )); + } + #[cfg(test)] + IdentityOAuthEndpointPolicy::ExplicitTestLoopback => { + let is_loopback_literal = host + .parse::() + .is_ok_and(|address| address.is_loopback()); + if !matches!(url.scheme(), "http" | "https") || !is_loopback_literal { + return Err(OAuthError::transport( + "test identity OAuth endpoint must use a loopback IP literal", + )); + } + } + _ => {} + } + let port = url + .port_or_known_default() + .ok_or_else(|| OAuthError::transport("identity OAuth endpoint is missing a port"))?; + Ok((url, host, port)) +} + +fn validate_identity_oauth_resolved_addrs( + url: &reqwest::Url, + addrs: &[SocketAddr], + policy: IdentityOAuthEndpointPolicy, +) -> Result<(), OAuthError> { + if addrs.is_empty() { + return Err(OAuthError::transport( + "identity OAuth endpoint DNS resolution returned no addresses", + )); + } + #[cfg(test)] + if policy == IdentityOAuthEndpointPolicy::ExplicitTestLoopback { + if addrs.iter().all(|addr| addr.ip().is_loopback()) { + return Ok(()); + } + return Err(OAuthError::transport( + "test identity OAuth endpoint must resolve only to loopback addresses", + )); + } + let allows_benchmarking_dns = policy == IdentityOAuthEndpointPolicy::PublicHttps + && identity_oauth_origin_allows_benchmarking_dns(url); + if addrs.iter().any(|addr| { + aether_http::is_private_or_reserved_ip(addr.ip()) + && !(allows_benchmarking_dns && aether_http::is_ipv4_benchmarking_fake_ip(addr.ip())) + }) { + return Err(OAuthError::transport( + "identity OAuth endpoint resolves to a private or reserved address", + )); + } + Ok(()) +} + +fn identity_oauth_origin_allows_benchmarking_dns(url: &reqwest::Url) -> bool { + url.scheme() == "https" + && url.port_or_known_default() == Some(443) + && url.username().is_empty() + && url.password().is_none() + && url.fragment().is_none() + && url.host_str().is_some_and(|host| { + let host = host.trim_end_matches('.'); + TRUSTED_IDENTITY_BENCHMARKING_DNS_HOSTS + .iter() + .any(|trusted| trusted.eq_ignore_ascii_case(host)) }) +} + +async fn resolve_identity_oauth_endpoint( + raw_url: &str, + policy: IdentityOAuthEndpointPolicy, +) -> Result { + let (url, host, port) = parse_identity_oauth_endpoint(raw_url, policy)?; + let addrs = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + aether_http::lookup_host_with_limits( + host.as_str(), + port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|_| OAuthError::transport("identity OAuth endpoint DNS resolution failed"))? + }; + validate_identity_oauth_resolved_addrs(&url, &addrs, policy)?; + Ok(ResolvedIdentityOAuthEndpoint { url, host, addrs }) +} + +fn build_pinned_identity_oauth_client( + host: &str, + addrs: &[SocketAddr], + timeouts: OAuthTimeouts, +) -> Result { + reqwest::Client::builder() + .no_proxy() + .redirect(identity_oauth_redirect_policy()) + .connect_timeout(Duration::from_millis(timeouts.connect_ms)) + .read_timeout(Duration::from_millis(timeouts.read_ms)) + .timeout(Duration::from_millis(timeouts.total_ms)) + .resolve_to_addrs(host, addrs) + .build() + .map_err(|_| OAuthError::transport("identity OAuth HTTP client initialization failed")) +} + +fn identity_oauth_redirect_policy() -> reqwest::redirect::Policy { + reqwest::redirect::Policy::none() +} + +fn identity_oauth_headers( + headers: &BTreeMap, + content_type: Option<&str>, +) -> Result { + let mut result = HeaderMap::new(); + for (name, value) in headers { + if name.eq_ignore_ascii_case(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) { + continue; + } + let name = HeaderName::from_bytes(name.as_bytes()) + .map_err(|_| OAuthError::transport("identity OAuth request has an invalid header"))?; + let value = HeaderValue::from_str(value) + .map_err(|_| OAuthError::transport("identity OAuth request has an invalid header"))?; + result.insert(name, value); + } + if !result.contains_key(CONTENT_TYPE) { + if let Some(content_type) = content_type { + let value = HeaderValue::from_str(content_type).map_err(|_| { + OAuthError::transport("identity OAuth request has an invalid content type") + })?; + result.insert(CONTENT_TYPE, value); + } + } + Ok(result) +} + +async fn execute_direct_identity_oauth( + app: &AppState, + request: OAuthHttpRequest, +) -> Result { + let endpoint_policy = identity_oauth_endpoint_policy(app, &request.url); + let target = resolve_identity_oauth_endpoint(&request.url, endpoint_policy).await?; + let client = build_pinned_identity_oauth_client( + target.host.as_str(), + &target.addrs, + request.network.timeouts, + )?; + let headers = identity_oauth_headers(&request.headers, request.content_type.as_deref())?; + let mut builder = client.request(request.method, target.url).headers(headers); + if let Some(json_body) = request.json_body { + builder = builder.json(&json_body); + } else if let Some(body_bytes) = request.body_bytes { + builder = builder.body(body_bytes); + } + + let response = builder + .send() + .await + .map_err(|_| OAuthError::transport("identity OAuth request failed"))?; + let status_code = response.status().as_u16(); + let content_encoding = response + .headers() + .get(CONTENT_ENCODING) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + let body_bytes = collect_identity_oauth_response_body(response).await?; + let decoded = decode_response_bytes_with_limit( + &body_bytes, + content_encoding.as_deref(), + OAUTH_RESPONSE_BODY_LIMIT_BYTES, + )? + .unwrap_or(body_bytes); + Ok(OAuthHttpResponse { + status_code, + body_text: String::from_utf8_lossy(&decoded).to_string(), + json_body: serde_json::from_slice(&decoded).ok(), + }) +} + +async fn collect_identity_oauth_response_body( + response: reqwest::Response, +) -> Result, OAuthError> { + if response + .content_length() + .is_some_and(|length| length > OAUTH_RESPONSE_BODY_LIMIT_BYTES as u64) + { + return Err(oauth_response_too_large()); + } + + let mut body = Vec::new(); + let mut stream = response.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = + chunk.map_err(|_| OAuthError::transport("identity OAuth response body read failed"))?; + if chunk.len() > OAUTH_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) { + return Err(oauth_response_too_large()); + } + body.extend_from_slice(&chunk); + } + Ok(body) +} + +fn oauth_response_too_large() -> OAuthError { + OAuthError::transport(format!( + "OAuth response body exceeds {OAUTH_RESPONSE_BODY_LIMIT_BYTES} bytes" + )) +} + +fn execution_result_to_oauth_response(result: &ExecutionResult) -> OAuthHttpResponse { + OAuthHttpResponse { + status_code: result.status_code, + body_text: execution_body_text(result), + json_body: execution_json_body(result), } } @@ -125,15 +533,28 @@ fn execution_body_bytes( headers: &BTreeMap, body: &aether_contracts::ResponseBody, ) -> Option> { - let bytes = body - .body_bytes_b64 - .as_deref() - .and_then(|value| STANDARD.decode(value).ok())?; - decode_response_bytes(&bytes, headers.get("content-encoding").map(String::as_str)) - .or(Some(bytes)) + let bytes = body.body_bytes_b64.as_deref().and_then(|value| { + crate::execution_runtime::transport::decode_base64_body_with_limit( + value, + OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + })?; + decode_response_bytes_with_limit( + &bytes, + headers.get("content-encoding").map(String::as_str), + OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + .flatten() + .or(Some(bytes)) } -fn decode_response_bytes(bytes: &[u8], content_encoding: Option<&str>) -> Option> { +fn decode_response_bytes_with_limit( + bytes: &[u8], + content_encoding: Option<&str>, + limit_bytes: usize, +) -> Result>, OAuthError> { match content_encoding .map(str::trim) .filter(|value| !value.is_empty()) @@ -142,20 +563,320 @@ fn decode_response_bytes(bytes: &[u8], content_encoding: Option<&str>) -> Option { Some("gzip") => { let mut decoder = GzDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) + read_oauth_response_decoder(&mut decoder, limit_bytes).map(Some) } Some("deflate") => { let mut decoder = DeflateDecoder::new(bytes); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) + read_oauth_response_decoder(&mut decoder, limit_bytes).map(Some) } - _ => None, + _ => Ok(None), } } +fn read_oauth_response_decoder( + decoder: &mut impl Read, + limit_bytes: usize, +) -> Result, OAuthError> { + let read_limit = u64::try_from(limit_bytes) + .unwrap_or(u64::MAX) + .saturating_add(1); + let mut limited = decoder.take(read_limit); + let mut out = Vec::new(); + limited + .read_to_end(&mut out) + .map_err(|_| OAuthError::transport("OAuth response body decompression failed"))?; + if out.len() > limit_bytes { + return Err(oauth_response_too_large()); + } + Ok(out) +} + fn gateway_error_to_oauth_error(error: GatewayError) -> OAuthError { OAuthError::Transport(error.into_message()) } + +#[cfg(test)] +mod tests { + use super::{ + decode_response_bytes_with_limit, execution_body_bytes, identity_oauth_endpoint_policy, + identity_oauth_origin_allows_benchmarking_dns, identity_oauth_redirect_policy, + identity_oauth_route, oauth_execution_plan, parse_identity_oauth_endpoint, + resolve_identity_oauth_endpoint, validate_identity_oauth_resolved_addrs, + IdentityOAuthEndpointPolicy, IdentityOAuthRoute, + }; + use aether_contracts::{ + ProxySnapshot, ResponseBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, + }; + use aether_oauth::network::{ + OAuthHttpRequest, OAuthNetworkContext, OAuthNetworkPolicy, OAuthTimeouts, + }; + use std::collections::BTreeMap; + use std::io::Write; + use std::net::SocketAddr; + + fn identity_request(proxy: Option) -> OAuthHttpRequest { + OAuthHttpRequest { + request_id: "identity-oauth:test".to_string(), + method: reqwest::Method::GET, + url: "https://oauth.example.test/userinfo".to_string(), + headers: BTreeMap::new(), + content_type: None, + json_body: None, + body_bytes: None, + network: OAuthNetworkContext { + policy: OAuthNetworkPolicy::DirectOrSystemProxy, + requirement: aether_oauth::network::NetworkRequirement::Optional, + proxy, + timeouts: OAuthTimeouts::DIRECT_DEFAULT, + }, + transport_profile: None, + } + } + + #[tokio::test] + async fn private_ip_literal_is_rejected_before_connect() { + let error = resolve_identity_oauth_endpoint( + "https://127.0.0.1/token", + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .await + .expect_err("loopback identity endpoint must be rejected"); + + assert!(error.to_string().contains("private or reserved")); + } + + #[test] + fn public_https_endpoint_and_resolved_addresses_are_accepted_without_dns() { + let (url, host, port) = parse_identity_oauth_endpoint( + "https://oauth.example.test/token?flow=login", + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .expect("public HTTPS URL should parse"); + let addrs = ["8.8.8.8:443".parse::().unwrap()]; + + assert_eq!(url.scheme(), "https"); + assert_eq!(host, "oauth.example.test"); + assert_eq!(port, 443); + validate_identity_oauth_resolved_addrs( + &url, + &addrs, + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .expect("controlled public address should pass validation"); + } + + #[test] + fn identity_oauth_rejects_credentials_fragments_and_any_private_dns_answer() { + assert!(parse_identity_oauth_endpoint( + "https://user:secret@example.test/token", + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + assert!(parse_identity_oauth_endpoint( + "https://example.test/token#secret", + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + assert!(parse_identity_oauth_endpoint( + "http://example.test/token", + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + + let mixed = [ + "8.8.8.8:443".parse::().unwrap(), + "10.0.0.4:443".parse::().unwrap(), + ]; + assert!(validate_identity_oauth_resolved_addrs( + &reqwest::Url::parse("https://example.test/token").unwrap(), + &mixed, + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + } + + #[test] + fn explicit_test_endpoint_policy_only_allows_loopback_literals() { + let (_, host, port) = parse_identity_oauth_endpoint( + "http://127.0.0.1:32123/token", + IdentityOAuthEndpointPolicy::ExplicitTestLoopback, + ) + .expect("explicit test loopback URL should parse"); + assert_eq!(host, "127.0.0.1"); + assert_eq!(port, 32123); + validate_identity_oauth_resolved_addrs( + &reqwest::Url::parse("http://127.0.0.1:32123/token").unwrap(), + &["127.0.0.1:32123".parse().unwrap()], + IdentityOAuthEndpointPolicy::ExplicitTestLoopback, + ) + .expect("loopback resolution should be accepted for an explicit test endpoint"); + + assert!(parse_identity_oauth_endpoint( + "http://10.0.0.1/token", + IdentityOAuthEndpointPolicy::ExplicitTestLoopback, + ) + .is_err()); + assert!(parse_identity_oauth_endpoint( + "http://localhost/token", + IdentityOAuthEndpointPolicy::ExplicitTestLoopback, + ) + .is_err()); + assert!(validate_identity_oauth_resolved_addrs( + &reqwest::Url::parse("http://10.0.0.1/token").unwrap(), + &["10.0.0.1:80".parse().unwrap()], + IdentityOAuthEndpointPolicy::ExplicitTestLoopback, + ) + .is_err()); + } + + #[test] + fn identity_oauth_benchmarking_dns_is_exact_origin_only() { + let fake = "198.18.75.234:443".parse::().unwrap(); + let trusted = reqwest::Url::parse("https://CONNECT.LINUX.DO/oauth2/token").unwrap(); + assert!(identity_oauth_origin_allows_benchmarking_dns(&trusted)); + assert!(validate_identity_oauth_resolved_addrs( + &trusted, + &[fake], + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_ok()); + + for raw_url in [ + "https://connect.linux.do:8443/oauth2/token", + "https://connect.linux.do.evil.test/oauth2/token", + "http://connect.linux.do/oauth2/token", + "https://oauth.example.test/token", + ] { + let url = reqwest::Url::parse(raw_url).unwrap(); + assert!(!identity_oauth_origin_allows_benchmarking_dns(&url)); + assert!(validate_identity_oauth_resolved_addrs( + &url, + &[fake], + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + } + + let mixed = [fake, "127.0.0.1:443".parse::().unwrap()]; + assert!(validate_identity_oauth_resolved_addrs( + &trusted, + &mixed, + IdentityOAuthEndpointPolicy::PublicHttps, + ) + .is_err()); + } + + #[test] + fn test_loopback_policy_requires_a_registered_loopback_origin_and_path() { + let app = crate::AppState::new() + .expect("gateway state should build") + .with_provider_oauth_token_url_for_tests("codex", "http://127.0.0.1:32123/oauth") + .with_provider_oauth_token_url_for_tests("bad", "http://10.0.0.1/token"); + + assert_eq!( + identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32123/oauth/token"), + IdentityOAuthEndpointPolicy::ExplicitTestLoopback + ); + assert_eq!( + identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32123/oauth2/token"), + IdentityOAuthEndpointPolicy::PublicHttps + ); + assert_eq!( + identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32124/oauth/token"), + IdentityOAuthEndpointPolicy::PublicHttps + ); + assert_eq!( + identity_oauth_endpoint_policy(&app, "http://10.0.0.1/token"), + IdentityOAuthEndpointPolicy::PublicHttps + ); + } + + #[test] + fn identity_oauth_disables_redirects_for_direct_and_tunnel_requests() { + assert_eq!( + format!("{:?}", identity_oauth_redirect_policy()), + "Policy(None)" + ); + + let tunnel = ProxySnapshot { + enabled: Some(true), + mode: Some("tunnel".to_string()), + node_id: Some("node-1".to_string()), + ..ProxySnapshot::default() + }; + assert!( + identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&tunnel)).is_err() + ); + let plan = oauth_execution_plan(identity_request(None), true); + assert_eq!( + plan.headers + .get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) + .map(String::as_str), + Some("false") + ); + assert_eq!( + crate::execution_runtime::transport::execution_plan_response_body_limit_bytes(&plan), + super::OAUTH_RESPONSE_BODY_LIMIT_BYTES + ); + } + + #[test] + fn identity_oauth_rejects_forward_proxy_snapshots() { + let proxy = ProxySnapshot { + enabled: Some(true), + mode: Some("http".to_string()), + url: Some("http://proxy.example.test:8080".to_string()), + ..ProxySnapshot::default() + }; + + assert!( + identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&proxy)).is_err() + ); + assert_eq!( + identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, None).unwrap(), + IdentityOAuthRoute::Direct + ); + } + + #[test] + fn identity_oauth_rejects_enabled_tunnel_even_when_node_id_is_present() { + let proxy = ProxySnapshot { + enabled: Some(true), + mode: Some("tunnel".to_string()), + node_id: Some("node-1".to_string()), + ..ProxySnapshot::default() + }; + assert!( + identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&proxy)).is_err() + ); + } + + #[test] + fn oauth_response_decoder_rejects_decompression_bombs() { + let payload = b"123456789"; + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default()); + encoder + .write_all(payload) + .expect("gzip payload should write"); + let encoded = encoder.finish().expect("gzip payload should finish"); + + let error = decode_response_bytes_with_limit(&encoded, Some("gzip"), 8) + .expect_err("decoded OAuth body above the limit must fail closed"); + + assert!(error.to_string().contains("exceeds")); + } + + #[test] + fn oauth_execution_body_rejects_oversized_base64_before_decode() { + let encoded_limit = + crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit( + super::OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ); + let body = ResponseBody { + json_body: None, + body_bytes_b64: Some("A".repeat(encoded_limit + 1)), + }; + + assert!(execution_body_bytes(&BTreeMap::new(), &body).is_none()); + } +} diff --git a/apps/aether-gateway/src/oauth/identity_repo.rs b/apps/aether-gateway/src/oauth/identity_repo.rs index d3e965aab..c7950952c 100644 --- a/apps/aether-gateway/src/oauth/identity_repo.rs +++ b/apps/aether-gateway/src/oauth/identity_repo.rs @@ -1,10 +1,16 @@ use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, module_available_from_env, + decrypt_or_migrate_identity_oauth_provider_client_secret, module_available_from_env, system_config_bool as system_config_bool_with_default, }; use crate::{AppState, GatewayError}; -use aether_data::repository::oauth_providers::StoredOAuthProviderConfig; -use aether_data::repository::users::{StoredUserAuthRecord, StoredUserOAuthLinkSummary}; +use aether_data::repository::oauth_providers::{ + validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config, + validate_oauth_redirect_uri, StoredOAuthProviderConfig, +}; +use aether_data::repository::users::{ + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserOAuthLinkSummary, +}; use aether_oauth::identity::{IdentityClaims, IdentityOAuthProviderConfig}; use chrono::Utc; use serde::Serialize; @@ -15,6 +21,12 @@ const LINUXDO_AUTHORIZE_URL: &str = "https://connect.linux.do/oauth2/authorize"; const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token"; const LINUXDO_USERINFO_URL: &str = "https://connect.linux.do/api/user"; +static IDENTITY_OAUTH_MUTATION_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); + +pub(crate) async fn lock_identity_oauth_mutation() -> tokio::sync::MutexGuard<'static, ()> { + IDENTITY_OAUTH_MUTATION_LOCK.lock().await +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize)] pub(crate) struct IdentityOAuthProviderSummary { pub(crate) provider_type: String, @@ -45,6 +57,7 @@ pub(crate) enum IdentityOAuthAccountError { AlreadyBoundProvider, LastOAuthBinding, LastLoginMethod, + BindingSessionUnavailable, Storage(String), } @@ -60,6 +73,7 @@ impl IdentityOAuthAccountError { Self::AlreadyBoundProvider => "already_bound_provider", Self::LastOAuthBinding => "last_oauth_binding", Self::LastLoginMethod => "last_login_method", + Self::BindingSessionUnavailable => "invalid_state", } } @@ -112,7 +126,9 @@ pub(crate) async fn get_enabled_identity_oauth_provider_config( let Some(config) = config.filter(|config| config.is_enabled) else { return Ok(None); }; - stored_provider_config_to_identity_config(state, config).map(Some) + stored_provider_config_to_identity_config(state, config) + .await + .map(Some) } async fn identity_oauth_module_enabled(state: &AppState) -> Result { @@ -161,27 +177,39 @@ pub(crate) async fn resolve_identity_oauth_login_user( state: &AppState, claims: &IdentityClaims, ) -> Result { + let _mutation_guard = lock_identity_oauth_mutation().await; let now = Utc::now(); - if let Some(user) = state + let provider_enabled = state + .get_oauth_provider_config(&claims.provider_type) + .await + .map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))? + .is_some_and(|provider| provider.is_enabled); + let verified_email = claims + .email_verified + .then(|| normalize_identity_email(claims.email.as_deref())) + .flatten(); + match state .data - .find_oauth_linked_user(&claims.provider_type, &claims.subject) + .resolve_enabled_oauth_linked_user( + &claims.provider_type, + &claims.subject, + claims.username.as_deref(), + claims.email.as_deref(), + None, + verified_email.as_deref(), + now, + provider_enabled, + ) .await .map_err(repo_data_error)? { - state - .data - .touch_oauth_link( - &claims.provider_type, - &claims.subject, - claims.username.as_deref(), - claims.email.as_deref(), - Some(claims.raw.clone()), - now, - ) - .await - .map_err(repo_data_error)?; - return Ok(user); + ResolveOAuthLinkedUserOutcome::Linked(user) => return Ok(user), + ResolveOAuthLinkedUserOutcome::NotLinked => {} + ResolveOAuthLinkedUserOutcome::ProviderUnavailable => { + return Err(IdentityOAuthAccountError::ProviderUnavailable); + } } + drop(_mutation_guard); let email = normalize_identity_email(claims.email.as_deref()); if let Some(email) = email.as_deref() { @@ -219,35 +247,44 @@ pub(crate) async fn resolve_identity_oauth_login_user( .unwrap_or(10.0); let username = unique_oauth_username(state, claims).await?; + let email_verified = email.is_some() && claims.email_verified; let user = state .data - .create_oauth_auth_user(email, username, now) + .create_oauth_auth_user(email, email_verified, username, now) .await .map_err(repo_data_error)? .ok_or_else(|| IdentityOAuthAccountError::Storage("oauth user not created".to_string()))?; - match state - .initialize_auth_user_wallet(&user.id, initial_gift, false) + let owned_wallet_id = match state + .initialize_auth_user_wallet_with_outcome(&user.id, initial_gift, false) .await { - Ok(Some(_wallet)) => {} + Ok(Some(outcome)) => outcome.created.then_some(outcome.wallet.id), Ok(None) => { - let _ = state.delete_local_auth_user(&user.id).await; + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; return Err(IdentityOAuthAccountError::ProviderUnavailable); } Err(err) => { - let _ = state.delete_local_auth_user(&user.id).await; + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; return Err(IdentityOAuthAccountError::Storage(format!("{err:?}"))); } - } + }; if let Err(err) = state .assign_default_group_to_self_registered_user(&user.id) .await { - let _ = state.delete_local_auth_user(&user.id).await; + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await; return Err(IdentityOAuthAccountError::Storage(format!("{err:?}"))); } - if let Err(err) = upsert_oauth_link(state, &user.id, claims, now).await { - let _ = state.delete_local_auth_user(&user.id).await; + if let Err(err) = bind_oauth_link_for_new_user(state, &user.id, claims, now).await { + let _ = state + .rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref()) + .await; return Err(err); } Ok(user) @@ -257,29 +294,39 @@ pub(crate) async fn bind_identity_oauth_to_user( state: &AppState, user: &StoredUserAuthRecord, claims: &IdentityClaims, + session_expectation: &BindUserOAuthLinkSessionExpectation, ) -> Result<(), IdentityOAuthAccountError> { if user.auth_source.eq_ignore_ascii_case("ldap") { return Err(IdentityOAuthAccountError::EmailIsLdap); } - if let Some(owner) = state - .data - .find_oauth_link_owner(&claims.provider_type, &claims.subject) - .await - .map_err(repo_data_error)? - { - if owner != user.id { + let now = Utc::now(); + match bind_oauth_link(state, &user.id, claims, now, Some(session_expectation)).await? { + BindUserOAuthLinkOutcome::Bound => {} + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + | BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider => { + return Err(IdentityOAuthAccountError::AlreadyBoundProvider); + } + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser => { return Err(IdentityOAuthAccountError::OAuthAlreadyBound); } + BindUserOAuthLinkOutcome::SessionUnavailable => { + return Err(IdentityOAuthAccountError::BindingSessionUnavailable); + } + BindUserOAuthLinkOutcome::UserNotFound + | BindUserOAuthLinkOutcome::ProviderNotFound + | BindUserOAuthLinkOutcome::ProviderDisabled => { + return Err(IdentityOAuthAccountError::ProviderUnavailable); + } } - if state - .data - .has_user_oauth_provider_link(&user.id, &claims.provider_type) - .await - .map_err(repo_data_error)? - { - return Err(IdentityOAuthAccountError::AlreadyBoundProvider); + if claims.email_verified { + if let Some(verified_email) = normalize_identity_email(claims.email.as_deref()) { + state + .data + .upgrade_oauth_email_verification_if_matches(&user.id, &verified_email, now) + .await + .map_err(repo_data_error)?; + } } - upsert_oauth_link(state, &user.id, claims, Utc::now()).await?; Ok(()) } @@ -287,32 +334,58 @@ pub(crate) async fn unbind_identity_oauth( state: &AppState, user: &StoredUserAuthRecord, provider_type: &str, + local_password_login_allowed: bool, ) -> Result { if user.auth_source.eq_ignore_ascii_case("ldap") { return Err(IdentityOAuthAccountError::EmailIsLdap); } - let link_count = state + let _mutation_guard = lock_identity_oauth_mutation().await; + let enabled_provider_types_snapshot = state + .list_oauth_provider_configs() + .await + .map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))? + .into_iter() + .filter(|provider| provider.is_enabled) + .map(|provider| provider.provider_type) + .collect::>(); + let outcome = state .data - .count_user_oauth_links(&user.id) + .delete_user_oauth_link( + &user.id, + provider_type.trim(), + local_password_login_allowed, + &enabled_provider_types_snapshot, + ) .await .map_err(repo_data_error)?; - if user.auth_source.eq_ignore_ascii_case("oauth") && link_count <= 1 { - return Err(IdentityOAuthAccountError::LastOAuthBinding); + match outcome { + DeleteUserOAuthLinkOutcome::Deleted => Ok(true), + DeleteUserOAuthLinkOutcome::NotFound => Ok(false), + DeleteUserOAuthLinkOutcome::LastOAuthBinding => { + Err(IdentityOAuthAccountError::LastOAuthBinding) + } + DeleteUserOAuthLinkOutcome::LastLoginMethod => { + Err(IdentityOAuthAccountError::LastLoginMethod) + } } - if !user.auth_source.eq_ignore_ascii_case("local") && link_count <= 1 { - return Err(IdentityOAuthAccountError::LastLoginMethod); - } - state - .data - .delete_user_oauth_link(&user.id, provider_type.trim()) - .await - .map_err(repo_data_error) } -fn stored_provider_config_to_identity_config( +async fn stored_provider_config_to_identity_config( state: &AppState, config: StoredOAuthProviderConfig, ) -> Result { + validate_oauth_redirect_uri(config.redirect_uri.trim()) + .map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?; + validate_oauth_frontend_callback_url(config.frontend_callback_url.trim()) + .map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?; + validate_oauth_provider_endpoint_config( + &config.provider_type, + config.authorization_url_override.as_deref(), + config.token_url_override.as_deref(), + config.userinfo_url_override.as_deref(), + config.extra_config.as_ref(), + ) + .map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?; let defaults = identity_provider_defaults(&config.provider_type); let authorization_url = config .authorization_url_override @@ -330,13 +403,9 @@ fn stored_provider_config_to_identity_config( .userinfo_url_override .clone() .or_else(|| defaults.map(|defaults| defaults.2.to_string())); - let client_secret = match config.client_secret_encrypted.as_deref() { - Some(ciphertext) => Some( - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext) - .ok_or(IdentityOAuthAccountError::ProviderUnavailable)?, - ), - None => None, - }; + let client_secret = decrypt_or_migrate_identity_oauth_provider_client_secret(state, &config) + .await + .map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?; Ok(IdentityOAuthProviderConfig { provider_type: config.provider_type, @@ -421,27 +490,65 @@ async fn unique_oauth_username( Ok(format!("oauth_{}", short_uuid())) } -async fn upsert_oauth_link( +async fn bind_oauth_link( state: &AppState, user_id: &str, claims: &IdentityClaims, now: chrono::DateTime, -) -> Result<(), IdentityOAuthAccountError> { + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, +) -> Result { + let _mutation_guard = lock_identity_oauth_mutation().await; + let provider_enabled = state + .get_oauth_provider_config(&claims.provider_type) + .await + .map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))? + .is_some_and(|provider| provider.is_enabled); + if !provider_enabled { + return Ok(BindUserOAuthLinkOutcome::ProviderNotFound); + } state .data - .upsert_user_oauth_link( + .bind_user_oauth_link( user_id, &claims.provider_type, &claims.subject, claims.username.as_deref(), claims.email.as_deref(), - Some(claims.raw.clone()), + None, now, + provider_enabled, + session_expectation, ) .await .map_err(repo_data_error) } +async fn bind_oauth_link_for_new_user( + state: &AppState, + user_id: &str, + claims: &IdentityClaims, + now: chrono::DateTime, +) -> Result<(), IdentityOAuthAccountError> { + match bind_oauth_link(state, user_id, claims, now, None).await? { + BindUserOAuthLinkOutcome::Bound => Ok(()), + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + | BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser => { + Err(IdentityOAuthAccountError::OAuthAlreadyBound) + } + BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider => { + Err(IdentityOAuthAccountError::AlreadyBoundProvider) + } + BindUserOAuthLinkOutcome::SessionUnavailable => { + Err(IdentityOAuthAccountError::BindingSessionUnavailable) + } + BindUserOAuthLinkOutcome::UserNotFound + | BindUserOAuthLinkOutcome::ProviderNotFound + | BindUserOAuthLinkOutcome::ProviderDisabled => { + Err(IdentityOAuthAccountError::ProviderUnavailable) + } + } +} + fn normalize_identity_email(value: Option<&str>) -> Option { value .map(str::trim) diff --git a/apps/aether-gateway/src/oauth/mod.rs b/apps/aether-gateway/src/oauth/mod.rs index 30b814bc3..f1f362f64 100644 --- a/apps/aether-gateway/src/oauth/mod.rs +++ b/apps/aether-gateway/src/oauth/mod.rs @@ -8,14 +8,14 @@ pub(crate) use http_executor::GatewayOAuthHttpExecutor; pub(crate) use identity_repo::{ bind_identity_oauth_to_user, get_enabled_identity_oauth_provider_config, list_bindable_identity_oauth_providers, list_enabled_identity_oauth_providers, - list_identity_oauth_links, resolve_identity_oauth_login_user, unbind_identity_oauth, - IdentityOAuthAccountError, + list_identity_oauth_links, lock_identity_oauth_mutation, resolve_identity_oauth_login_user, + unbind_identity_oauth, IdentityOAuthAccountError, }; pub(crate) use provider_repo::ProviderOAuthRepository; pub(crate) use proxy::{ resolve_identity_oauth_network_context, resolve_provider_oauth_operation_proxy_snapshot, }; pub(crate) use state_store::{ - consume_identity_oauth_state, save_identity_oauth_state, IdentityOAuthStateMode, - StoredIdentityOAuthState, + consume_identity_oauth_state, identity_oauth_state_storage_key, load_identity_oauth_state, + save_identity_oauth_state, IdentityOAuthStateMode, StoredIdentityOAuthState, }; diff --git a/apps/aether-gateway/src/oauth/provider_repo.rs b/apps/aether-gateway/src/oauth/provider_repo.rs index 48fd62ef9..6defa500a 100644 --- a/apps/aether-gateway/src/oauth/provider_repo.rs +++ b/apps/aether-gateway/src/oauth/provider_repo.rs @@ -13,24 +13,6 @@ use aether_data_contracts::repository::provider_catalog::{ pub(crate) struct ProviderOAuthRepository; impl ProviderOAuthRepository { - pub(crate) async fn update_provider_catalog_key_oauth_credentials( - state: &AdminAppState<'_>, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - state - .app() - .update_provider_catalog_key_oauth_credentials( - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } - pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker( state: &AdminAppState<'_>, key_id: &str, diff --git a/apps/aether-gateway/src/oauth/proxy.rs b/apps/aether-gateway/src/oauth/proxy.rs index 486246f54..ca4b71f61 100644 --- a/apps/aether-gateway/src/oauth/proxy.rs +++ b/apps/aether-gateway/src/oauth/proxy.rs @@ -24,11 +24,18 @@ pub(crate) async fn resolve_provider_oauth_operation_proxy_snapshot( temporary_proxy_node_id: Option<&str>, configured_proxies: &[Option<&serde_json::Value>], ) -> Option { - if let Some(snapshot) = state - .resolve_admin_proxy_node_snapshot(temporary_proxy_node_id) - .await + if let Some(temporary_proxy_node_id) = temporary_proxy_node_id + .map(str::trim) + .filter(|value| !value.is_empty()) { - return Some(snapshot); + return state + .resolve_admin_proxy_node_snapshot(Some(temporary_proxy_node_id)) + .await + .or_else(|| { + Some(crate::state::unavailable_proxy_snapshot( + "temporary_proxy_node_unavailable", + )) + }); } for proxy in configured_proxies { diff --git a/apps/aether-gateway/src/oauth/state_store.rs b/apps/aether-gateway/src/oauth/state_store.rs index 766e6f3ea..b6e50a114 100644 --- a/apps/aether-gateway/src/oauth/state_store.rs +++ b/apps/aether-gateway/src/oauth/state_store.rs @@ -1,8 +1,11 @@ use crate::{AppState, GatewayError}; use aether_oauth::core::{current_unix_secs, generate_oauth_nonce}; use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; const IDENTITY_OAUTH_STATE_TTL_SECS: u64 = 10 * 60; +const IDENTITY_OAUTH_STATE_SECRET_PURPOSE: &str = "identity-oauth-state"; +const IDENTITY_OAUTH_STATE_MAX_CLOCK_SKEW_SECS: u64 = 60; #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] @@ -11,13 +14,15 @@ pub(crate) enum IdentityOAuthStateMode { Bind, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub(crate) struct StoredIdentityOAuthState { pub(crate) nonce: String, pub(crate) provider_type: String, pub(crate) mode: IdentityOAuthStateMode, pub(crate) client_device_id: String, #[serde(default, skip_serializing_if = "Option::is_none")] + pub(crate) browser_binding_hash: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] pub(crate) pkce_verifier: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub(crate) bind_user_id: Option, @@ -26,17 +31,36 @@ pub(crate) struct StoredIdentityOAuthState { pub(crate) created_at: u64, } +impl std::fmt::Debug for StoredIdentityOAuthState { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredIdentityOAuthState") + .field("nonce", &"[REDACTED]") + .field("provider_type", &self.provider_type) + .field("mode", &self.mode) + .field("client_device_id", &"[REDACTED]") + .field("browser_binding_hash", &"[REDACTED]") + .field("pkce_verifier", &"[REDACTED]") + .field("bind_user_id", &"[REDACTED]") + .field("bind_session_id", &"[REDACTED]") + .field("created_at", &self.created_at) + .finish() + } +} + impl StoredIdentityOAuthState { pub(crate) fn login( provider_type: impl Into, client_device_id: impl Into, pkce_verifier: Option, + browser_binding_hash: Option, ) -> Self { Self { nonce: generate_oauth_nonce(), provider_type: provider_type.into(), mode: IdentityOAuthStateMode::Login, client_device_id: client_device_id.into(), + browser_binding_hash, pkce_verifier, bind_user_id: None, bind_session_id: None, @@ -48,6 +72,7 @@ impl StoredIdentityOAuthState { provider_type: impl Into, client_device_id: impl Into, pkce_verifier: Option, + browser_binding_hash: String, user_id: impl Into, session_id: impl Into, ) -> Self { @@ -56,6 +81,7 @@ impl StoredIdentityOAuthState { provider_type: provider_type.into(), mode: IdentityOAuthStateMode::Bind, client_device_id: client_device_id.into(), + browser_binding_hash: Some(browser_binding_hash), pkce_verifier, bind_user_id: Some(user_id.into()), bind_session_id: Some(session_id.into()), @@ -65,16 +91,43 @@ impl StoredIdentityOAuthState { } pub(crate) fn identity_oauth_state_storage_key(nonce: &str) -> String { + format!( + "identity_oauth_state:sha256:{:x}", + Sha256::digest(nonce.trim().as_bytes()) + ) +} + +fn identity_oauth_state_secret_purpose(nonce: &str) -> String { + format!( + "{IDENTITY_OAUTH_STATE_SECRET_PURPOSE}:sha256:{:x}", + Sha256::digest(nonce.trim().as_bytes()) + ) +} + +fn legacy_identity_oauth_state_storage_key(nonce: &str) -> String { format!("identity_oauth_state:{}", nonce.trim()) } +fn is_generated_oauth_nonce(nonce: &str) -> bool { + let nonce = nonce.trim(); + nonce.len() == 64 + && nonce + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) +} + pub(crate) async fn save_identity_oauth_state( state: &AppState, record: &StoredIdentityOAuthState, ) -> Result<(), GatewayError> { let key = identity_oauth_state_storage_key(&record.nonce); - let value = + let plaintext = serde_json::to_string(record).map_err(|err| GatewayError::Internal(err.to_string()))?; + let purpose = identity_oauth_state_secret_purpose(&record.nonce); + let value = crate::handlers::shared::seal_runtime_secret_payload(state, &purpose, &plaintext) + .ok_or_else(|| { + GatewayError::Internal("identity OAuth state encryption unavailable".to_string()) + })?; state .runtime_kv_setex(&key, &value, IDENTITY_OAUTH_STATE_TTL_SECS) .await @@ -85,10 +138,226 @@ pub(crate) async fn consume_identity_oauth_state( nonce: &str, ) -> Result, GatewayError> { let key = identity_oauth_state_storage_key(nonce); - let raw = state.runtime_kv_getdel(&key).await?; - raw.map(|value| { - serde_json::from_str::(&value) - .map_err(|err| GatewayError::Internal(err.to_string())) - }) - .transpose() + let raw = match state.runtime_kv_getdel(&key).await? { + Some(value) => Some(value), + None if is_generated_oauth_nonce(nonce) => { + state + .runtime_kv_getdel(&legacy_identity_oauth_state_storage_key(nonce)) + .await? + } + None => None, + }; + raw.map(|value| decode_identity_oauth_state(state, nonce, &value)) + .transpose() +} + +pub(crate) async fn load_identity_oauth_state( + state: &AppState, + nonce: &str, +) -> Result, GatewayError> { + let key = identity_oauth_state_storage_key(nonce); + let raw = match state.runtime_kv_get(&key).await? { + Some(value) => Some(value), + None if is_generated_oauth_nonce(nonce) => { + state + .runtime_kv_get(&legacy_identity_oauth_state_storage_key(nonce)) + .await? + } + None => None, + }; + raw.map(|value| decode_identity_oauth_state(state, nonce, &value)) + .transpose() +} + +fn decode_identity_oauth_state( + state: &AppState, + expected_nonce: &str, + stored: &str, +) -> Result { + let expected_nonce = expected_nonce.trim(); + let purpose = identity_oauth_state_secret_purpose(expected_nonce); + let plaintext = crate::handlers::shared::open_runtime_secret_payload(state, &purpose, stored) + // States created immediately before a rolling upgrade still live under the + // legacy key and were sealed with the fixed purpose. They remain safe to + // accept for their short TTL because the decoded record is checked against + // the callback nonce and all authority-bearing fields below. + .or_else(|| { + crate::handlers::shared::open_runtime_secret_payload( + state, + IDENTITY_OAUTH_STATE_SECRET_PURPOSE, + stored, + ) + }) + .ok_or_else(|| GatewayError::Internal("identity OAuth state is invalid".to_string()))?; + let record = serde_json::from_str::(&plaintext) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + validate_identity_oauth_state(expected_nonce, &record)?; + Ok(record) +} + +fn validate_identity_oauth_state( + expected_nonce: &str, + record: &StoredIdentityOAuthState, +) -> Result<(), GatewayError> { + let invalid = || GatewayError::Internal("identity OAuth state is invalid".to_string()); + if !is_generated_oauth_nonce(expected_nonce) + || record.nonce != expected_nonce + || !is_generated_oauth_nonce(&record.nonce) + || record.provider_type.trim().is_empty() + || record.provider_type != record.provider_type.trim().to_ascii_lowercase() + || record.client_device_id.trim().is_empty() + || !record + .browser_binding_hash + .as_deref() + .is_some_and(is_lower_hex_sha256) + || !record + .pkce_verifier + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + return Err(invalid()); + } + + let mode_is_valid = match record.mode { + IdentityOAuthStateMode::Login => { + record.bind_user_id.is_none() && record.bind_session_id.is_none() + } + IdentityOAuthStateMode::Bind => { + record + .bind_user_id + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + && record + .bind_session_id + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + } + }; + let now = current_unix_secs(); + let time_is_valid = record.created_at + <= now.saturating_add(IDENTITY_OAUTH_STATE_MAX_CLOCK_SKEW_SECS) + && now.saturating_sub(record.created_at) <= IDENTITY_OAUTH_STATE_TTL_SECS; + if !mode_is_valid || !time_is_valid { + return Err(invalid()); + } + Ok(()) +} + +fn is_lower_hex_sha256(value: &str) -> bool { + value.len() == 64 + && value + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte)) +} + +#[cfg(test)] +mod tests { + use super::{ + decode_identity_oauth_state, identity_oauth_state_secret_purpose, StoredIdentityOAuthState, + }; + use crate::{data::GatewayDataState, AppState}; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + + fn state_with_encryption_key() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[test] + fn identity_oauth_state_ciphertext_is_bound_to_its_nonce() { + let state = state_with_encryption_key(); + let record = StoredIdentityOAuthState::login( + "linuxdo", + "device-1", + Some("pkce-verifier".to_string()), + Some("a".repeat(64)), + ); + let plaintext = serde_json::to_string(&record).expect("state should serialize"); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &state, + &identity_oauth_state_secret_purpose(&record.nonce), + &plaintext, + ) + .expect("state should seal"); + + assert_eq!( + decode_identity_oauth_state(&state, &record.nonce, &sealed) + .expect("matching state should open"), + record + ); + assert!(decode_identity_oauth_state(&state, &"b".repeat(64), &sealed).is_err()); + } + + #[test] + fn identity_oauth_state_debug_redacts_authorization_material() { + let record = StoredIdentityOAuthState::login( + "linuxdo", + "debug-secret-device", + Some("debug-secret-pkce".to_string()), + Some("debug-secret-binding-hash".to_string()), + ); + let nonce = record.nonce.clone(); + let rendered = format!("{record:?}"); + + for secret in [ + nonce.as_str(), + "debug-secret-device", + "debug-secret-pkce", + "debug-secret-binding-hash", + ] { + assert!(!rendered.contains(secret), "Debug output leaked {secret}"); + } + assert!(rendered.contains("[REDACTED]")); + assert!(rendered.contains("linuxdo")); + } + + #[test] + fn identity_oauth_state_accepts_valid_legacy_ciphertext_during_ttl_window() { + let state = state_with_encryption_key(); + let record = StoredIdentityOAuthState::login( + "linuxdo", + "device-1", + Some("pkce-verifier".to_string()), + Some("a".repeat(64)), + ); + let plaintext = serde_json::to_string(&record).expect("state should serialize"); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &state, + super::IDENTITY_OAUTH_STATE_SECRET_PURPOSE, + &plaintext, + ) + .expect("legacy state should seal"); + + assert_eq!( + decode_identity_oauth_state(&state, &record.nonce, &sealed) + .expect("valid legacy state should open"), + record + ); + assert!(decode_identity_oauth_state(&state, &"b".repeat(64), &sealed).is_err()); + } + + #[test] + fn identity_oauth_state_rejects_mode_field_confusion() { + let state = state_with_encryption_key(); + let mut record = StoredIdentityOAuthState::login( + "linuxdo", + "device-1", + Some("pkce-verifier".to_string()), + Some("a".repeat(64)), + ); + record.bind_user_id = Some("unexpected-user".to_string()); + let plaintext = serde_json::to_string(&record).expect("state should serialize"); + let sealed = crate::handlers::shared::seal_runtime_secret_payload( + &state, + &identity_oauth_state_secret_purpose(&record.nonce), + &plaintext, + ) + .expect("state should seal"); + + assert!(decode_identity_oauth_state(&state, &record.nonce, &sealed).is_err()); + } } diff --git a/apps/aether-gateway/src/orchestration/attempt.rs b/apps/aether-gateway/src/orchestration/attempt.rs index 1ae62662f..aa4c1d0be 100644 --- a/apps/aether-gateway/src/orchestration/attempt.rs +++ b/apps/aether-gateway/src/orchestration/attempt.rs @@ -1,8 +1,9 @@ +use aether_ai_serving::{AiExecutionAttempt, STICKY_KEY_ATTEMPTS_REPORT_FIELD}; +use aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS; use aether_runtime_state::RuntimeLockLease; use aether_scheduler_core::parse_request_candidate_report_context; use serde_json::Value; - -use crate::provider_transport::GatewayProviderTransportSnapshot; +use uuid::Uuid; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) struct ExecutionAttemptIdentity { @@ -32,6 +33,9 @@ pub(crate) struct LocalExecutionCandidateMetadata { pub(crate) pool_key_index: Option, pub(crate) pool_key_lease: Option, pub(crate) scheduler_affinity_epoch: Option, + /// Routing-policy `sticky_key_attempts` in effect for this request. `None` + /// means the policy default applies. + pub(crate) sticky_key_attempts: Option, } pub(crate) const SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD: &str = "scheduler_affinity_epoch"; @@ -42,6 +46,11 @@ pub(crate) const POOL_KEY_LEASE_TOKEN_REPORT_FIELD: &str = "pool_key_lease_token pub(crate) const POOL_KEY_LEASE_FENCING_REPORT_FIELD: &str = "pool_key_lease_fencing_token"; pub(crate) const POOL_KEY_LEASE_TTL_MS_REPORT_FIELD: &str = "pool_key_lease_ttl_ms"; +/// Pool-expanded keys encode `pool_key_index * STRIDE + retry_index` into the +/// persisted `retry_index` so a pool group's keys stay ordered in one candidate +/// slot. Same-key retries on a pool key are therefore bounded by the stride. +pub(crate) const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100; + pub(crate) fn attempt_identity_from_report_context( report_context: Option<&Value>, ) -> Option { @@ -72,6 +81,10 @@ pub(crate) fn local_execution_candidate_metadata_from_report_context( scheduler_affinity_epoch: report_context .and_then(|value| value.get(SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD)) .and_then(Value::as_u64), + sticky_key_attempts: report_context + .and_then(|value| value.get(STICKY_KEY_ATTEMPTS_REPORT_FIELD)) + .and_then(Value::as_u64) + .and_then(|value| u32::try_from(value).ok()), } } @@ -140,58 +153,55 @@ fn pool_key_lease_from_report_context(report_context: Option<&Value>) -> Option< }) } -pub(crate) fn build_local_attempt_identities( - candidate_index: u32, - transport: &GatewayProviderTransportSnapshot, -) -> Vec { - let attempt_slots = local_attempt_slot_count(transport); - (0..attempt_slots) - .map(|retry_index| ExecutionAttemptIdentity::new(candidate_index, retry_index)) - .collect() +/// Retry index of the next same-key attempt, or `None` when the sticky-key +/// budget for this candidate is used up. +/// +/// Only the first-ranked candidate (index `0`, the cache-affinity sticky key) +/// is retried on the same key; every later candidate gets exactly one attempt +/// so that once failover has started it keeps advancing. `sticky_key_attempts` +/// is the *total* attempt count on that key: `2` means one retry, `0` and `1` +/// mean none. There is no upper bound: attempts are derived one at a time +/// after each failure, never materialized ahead of time. +/// +/// Inside a pool group only the first key (`pool_key_index == 0`) is treated +/// as sticky, and its retries stay below `POOL_KEY_RETRY_INDEX_STRIDE` so the +/// encoded retry index never collides with the next pool key. +pub(crate) fn next_same_key_retry_index( + identity: ExecutionAttemptIdentity, + sticky_key_attempts: Option, +) -> Option { + if identity.candidate_index != 0 { + return None; + } + let pool_limit = match identity.pool_key_index { + None => u32::MAX, + Some(0) => POOL_KEY_RETRY_INDEX_STRIDE, + Some(_) => return None, + }; + let budget = sticky_key_attempts.unwrap_or(DEFAULT_STICKY_KEY_ATTEMPTS); + let attempts_so_far = identity.retry_index.checked_add(1)?; + if attempts_so_far >= budget || attempts_so_far >= pool_limit { + return None; + } + Some(attempts_so_far) } -pub(crate) fn local_attempt_slot_count(transport: &GatewayProviderTransportSnapshot) -> u32 { - local_attempt_slots_from_transport(transport).unwrap_or(1) -} - -/// For endpoint/provider table fields, `2` is the legacy admin default and is -/// treated as "not explicitly configured" so existing local-execution behaviour -/// (one attempt slot per candidate) stays unchanged. Values `0`, `1`, and `>2` -/// are treated as explicit. -const LEGACY_DEFAULT_MAX_RETRIES: u32 = 2; - -/// Upper bound on local attempt slots. This is intentionally stricter than -/// admin max_retries validation to prevent unbounded pre-materialization from -/// arbitrarily large JSON config values. -const MAX_LOCAL_ATTEMPT_SLOTS: u32 = 99; - -fn local_attempt_slots_from_transport(transport: &GatewayProviderTransportSnapshot) -> Option { - let rules = transport - .provider - .config - .as_ref() - .and_then(|config| config.get("failover_rules")) - .and_then(Value::as_object); - - rules - .and_then(|value| value.get("max_retries")) - .and_then(Value::as_u64) - .and_then(|value| u32::try_from(value).ok()) - .or_else(|| { - transport - .endpoint - .max_retries - .and_then(|value| u32::try_from(value).ok()) - .filter(|&value| value != LEGACY_DEFAULT_MAX_RETRIES) - }) - .or_else(|| { - transport - .provider - .max_retries - .and_then(|value| u32::try_from(value).ok()) - .filter(|&value| value != LEGACY_DEFAULT_MAX_RETRIES) - }) - .map(|value| value.clamp(1, MAX_LOCAL_ATTEMPT_SLOTS)) +/// Derive the next same-key attempt for `attempt` after a candidate-scoped +/// failure, reading the attempt identity and sticky budget from its report +/// context. Returns `None` when no further same-key retry is allowed. +pub(crate) fn next_same_key_retry_attempt(attempt: &A) -> Option { + let owned_report_context = attempt + .report_context_ref() + .is_none() + .then(|| attempt.report_context()) + .flatten(); + let report_context = attempt + .report_context_ref() + .or(owned_report_context.as_ref()); + let identity = attempt_identity_from_report_context(report_context)?; + let metadata = local_execution_candidate_metadata_from_report_context(report_context); + let retry_index = next_same_key_retry_index(identity, metadata.sticky_key_attempts)?; + attempt.with_same_key_retry(retry_index, Uuid::new_v4().to_string()) } #[cfg(test)] @@ -199,277 +209,143 @@ mod tests { use serde_json::json; use super::{ - attempt_identity_from_report_context, build_local_attempt_identities, - local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity, - LocalExecutionCandidateMetadata, - }; - use crate::provider_transport::snapshot::{ - GatewayProviderTransportEndpoint, GatewayProviderTransportKey, - GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, + attempt_identity_from_report_context, + local_execution_candidate_metadata_from_report_context, next_same_key_retry_attempt, + next_same_key_retry_index, ExecutionAttemptIdentity, LocalExecutionCandidateMetadata, + POOL_KEY_RETRY_INDEX_STRIDE, }; + use aether_ai_serving::{AiExecutionAttempt, AiSyncAttempt}; use aether_runtime_state::RuntimeLockLease; - fn sample_transport( - provider_max_retries: Option, - endpoint_max_retries: Option, - provider_config: Option, - ) -> GatewayProviderTransportSnapshot { - GatewayProviderTransportSnapshot { - provider: GatewayProviderTransportProvider { - id: "provider-1".to_string(), - name: "OpenAI".to_string(), - provider_type: "llm".to_string(), - website: None, - is_active: true, - keep_priority_on_conversion: false, - enable_format_conversion: true, - concurrent_limit: None, - max_retries: provider_max_retries, - proxy: None, - request_timeout_secs: None, - stream_first_byte_timeout_secs: None, - config: provider_config, - }, - endpoint: GatewayProviderTransportEndpoint { - id: "endpoint-1".to_string(), - provider_id: "provider-1".to_string(), - api_format: "openai:chat".to_string(), - api_family: Some("openai".to_string()), - endpoint_kind: Some("chat".to_string()), - is_active: true, - base_url: "https://example.com".to_string(), - header_rules: None, - body_rules: None, - max_retries: endpoint_max_retries, - custom_path: None, - config: None, - format_acceptance_config: None, - proxy: None, - }, - key: GatewayProviderTransportKey { - id: "key-1".to_string(), - provider_id: "provider-1".to_string(), - name: "primary".to_string(), - auth_type: "bearer".to_string(), - is_active: true, - api_formats: None, - auth_type_by_format: None, - allow_auth_channel_mismatch_formats: None, + #[test] + fn first_candidate_defaults_to_one_same_key_retry() { + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 0), None), + Some(1) + ); + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 1), None), + 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: "secret".to_string(), - decrypted_auth_config: None, - }, + #[test] + fn first_candidate_uses_policy_sticky_key_attempts_without_upper_bound() { + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 2), Some(3)), + None + ); + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 1), Some(3)), + Some(2) + ); + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 4_999), Some(10_000)), + Some(5_000) + ); + } + + #[test] + fn zero_and_one_mean_single_attempt() { + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 0), Some(0)), + None + ); + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(0, 0), Some(1)), + None + ); + } + + #[test] + fn failover_candidates_never_retry_on_the_same_key() { + for candidate_index in 1..5 { + assert_eq!( + next_same_key_retry_index(ExecutionAttemptIdentity::new(candidate_index, 0), None), + None + ); + assert_eq!( + next_same_key_retry_index( + ExecutionAttemptIdentity::new(candidate_index, 0), + Some(50) + ), + None + ); } } #[test] - fn build_local_attempt_identities_defaults_to_single_attempt() { - let identities = build_local_attempt_identities(3, &sample_transport(None, None, None)); - - assert_eq!(identities, vec![ExecutionAttemptIdentity::new(3, 0)]); - } - - #[test] - fn build_local_attempt_identities_prefer_failover_rules_over_endpoint_and_provider() { - let identities = build_local_attempt_identities( - 1, - &sample_transport( - Some(5), - Some(4), - Some(json!({ - "failover_rules": { - "max_retries": 2 - } - })), - ), - ); + fn pool_groups_only_retry_their_first_key_within_the_stride() { + let first_pool_key = ExecutionAttemptIdentity::new(0, 0).with_pool_key_index(Some(0)); + assert_eq!(next_same_key_retry_index(first_pool_key, Some(3)), Some(1)); + let at_stride_limit = ExecutionAttemptIdentity::new(0, POOL_KEY_RETRY_INDEX_STRIDE - 1) + .with_pool_key_index(Some(0)); assert_eq!( - identities, - vec![ - ExecutionAttemptIdentity::new(1, 0), - ExecutionAttemptIdentity::new(1, 1), - ] + next_same_key_retry_index(at_stride_limit, Some(10_000)), + None ); - } - - #[test] - fn build_local_attempt_identities_falls_back_to_endpoint_max_retries() { - let identities = - build_local_attempt_identities(2, &sample_transport(Some(5), Some(3), None)); + let second_pool_key = ExecutionAttemptIdentity::new(0, POOL_KEY_RETRY_INDEX_STRIDE) + .with_pool_key_index(Some(1)); assert_eq!( - identities, - vec![ - ExecutionAttemptIdentity::new(2, 0), - ExecutionAttemptIdentity::new(2, 1), - ExecutionAttemptIdentity::new(2, 2), - ] + next_same_key_retry_index(second_pool_key, Some(10_000)), + None ); } #[test] - fn build_local_attempt_identities_falls_back_to_provider_max_retries() { - let identities = build_local_attempt_identities(0, &sample_transport(Some(4), None, None)); + fn next_same_key_retry_attempt_rewrites_candidate_id_and_retry_index() { + let attempt = AiSyncAttempt { + plan: aether_contracts::ExecutionPlan { + request_id: "trace-1".to_string(), + candidate_id: Some("candidate-a".to_string()), + provider_name: None, + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-1".to_string(), + method: "POST".to_string(), + url: "https://example.com".to_string(), + headers: Default::default(), + content_type: None, + content_encoding: None, + body: aether_contracts::RequestBody { + json_body: None, + body_bytes_b64: None, + body_ref: None, + }, + stream: false, + client_api_format: "openai:chat".to_string(), + provider_api_format: "openai:chat".to_string(), + model_name: None, + proxy: None, + transport_profile: None, + timeouts: None, + }, + report_kind: None, + report_context: Some(json!({ + "candidate_id": "candidate-a", + "candidate_index": 0, + "retry_index": 0, + "sticky_key_attempts": 2, + })), + }; - assert_eq!( - identities, - vec![ - ExecutionAttemptIdentity::new(0, 0), - ExecutionAttemptIdentity::new(0, 1), - ExecutionAttemptIdentity::new(0, 2), - ExecutionAttemptIdentity::new(0, 3), - ] + let retry = next_same_key_retry_attempt(&attempt).expect("one same-key retry remains"); + let retry_candidate_id = retry.plan.candidate_id.clone().expect("fresh candidate id"); + assert_ne!(retry_candidate_id, "candidate-a"); + assert_eq!(retry.plan.key_id, "key-1"); + let context = retry.report_context_ref().expect("context retained"); + assert_eq!(context["candidate_id"], json!(retry_candidate_id)); + assert_eq!(context["retry_index"], json!(1)); + assert_eq!(context["candidate_index"], json!(0)); + + assert!( + next_same_key_retry_attempt(&retry).is_none(), + "budget of 2 attempts is exhausted after one retry" ); } - #[test] - fn build_local_attempt_identities_endpoint_overrides_provider() { - let identities = - build_local_attempt_identities(7, &sample_transport(Some(10), Some(3), None)); - - assert_eq!( - identities, - vec![ - ExecutionAttemptIdentity::new(7, 0), - ExecutionAttemptIdentity::new(7, 1), - ExecutionAttemptIdentity::new(7, 2), - ] - ); - } - - #[test] - fn build_local_attempt_identities_default_two_treated_as_unset() { - let identities = - build_local_attempt_identities(5, &sample_transport(Some(2), Some(2), None)); - - assert_eq!(identities, vec![ExecutionAttemptIdentity::new(5, 0)]); - } - - #[test] - fn build_local_attempt_identities_endpoint_two_falls_back_to_provider_ten() { - let identities = - build_local_attempt_identities(1, &sample_transport(Some(10), Some(2), None)); - - assert_eq!( - identities, - vec![ - ExecutionAttemptIdentity::new(1, 0), - ExecutionAttemptIdentity::new(1, 1), - ExecutionAttemptIdentity::new(1, 2), - ExecutionAttemptIdentity::new(1, 3), - ExecutionAttemptIdentity::new(1, 4), - ExecutionAttemptIdentity::new(1, 5), - ExecutionAttemptIdentity::new(1, 6), - ExecutionAttemptIdentity::new(1, 7), - ExecutionAttemptIdentity::new(1, 8), - ExecutionAttemptIdentity::new(1, 9), - ] - ); - } - - #[test] - fn build_local_attempt_identities_failover_rules_zero_produces_one_slot() { - let identities = build_local_attempt_identities( - 1, - &sample_transport( - Some(5), - Some(4), - Some(json!({ - "failover_rules": { - "max_retries": 0 - } - })), - ), - ); - - assert_eq!(identities, vec![ExecutionAttemptIdentity::new(1, 0)]); - } - - #[test] - fn build_local_attempt_identities_endpoint_zero_produces_one_slot() { - let identities = - build_local_attempt_identities(3, &sample_transport(Some(5), Some(0), None)); - - assert_eq!(identities, vec![ExecutionAttemptIdentity::new(3, 0)]); - } - - #[test] - fn build_local_attempt_identities_provider_zero_produces_one_slot() { - let identities = build_local_attempt_identities(3, &sample_transport(Some(0), None, None)); - - assert_eq!(identities, vec![ExecutionAttemptIdentity::new(3, 0)]); - } - - #[test] - fn build_local_attempt_identities_provider_ten_creates_ten_slots() { - let identities = build_local_attempt_identities(2, &sample_transport(Some(10), None, None)); - - assert_eq!(identities.len(), 10); - assert_eq!(identities[0], ExecutionAttemptIdentity::new(2, 0)); - assert_eq!(identities[9], ExecutionAttemptIdentity::new(2, 9)); - } - - #[test] - fn build_local_attempt_identities_failover_rules_over_limit_clamped_to_max() { - let identities = build_local_attempt_identities( - 0, - &sample_transport( - Some(3), - Some(5), - Some(json!({ - "failover_rules": { - "max_retries": 1000 - } - })), - ), - ); - - assert_eq!(identities.len(), 99); - } - - #[test] - fn build_local_attempt_identities_failover_rules_u32_max_clamped_to_max() { - let identities = build_local_attempt_identities( - 0, - &sample_transport( - None, - None, - Some(json!({ - "failover_rules": { - "max_retries": u32::MAX - } - })), - ), - ); - - assert_eq!(identities.len(), 99); - } - - #[test] - fn build_local_attempt_identities_endpoint_over_limit_clamped_to_max() { - let identities = - build_local_attempt_identities(0, &sample_transport(None, Some(2000), None)); - - assert_eq!(identities.len(), 99); - } - - #[test] - fn build_local_attempt_identities_provider_over_limit_clamped_to_max() { - let identities = - build_local_attempt_identities(0, &sample_transport(Some(5000), None, None)); - - assert_eq!(identities.len(), 99); - } - #[test] fn parse_attempt_identity_from_report_context_reads_candidate_and_retry_indices() { let identity = attempt_identity_from_report_context(Some(&json!({ @@ -499,6 +375,7 @@ mod tests { "pool_key_lease_token": "gateway-1:token-1", "pool_key_lease_fencing_token": 7, "pool_key_lease_ttl_ms": 900000, + "sticky_key_attempts": 3, }))); assert_eq!( @@ -514,6 +391,7 @@ mod tests { ttl_ms: 900000, }), scheduler_affinity_epoch: None, + sticky_key_attempts: Some(3), } ); } diff --git a/apps/aether-gateway/src/orchestration/effects.rs b/apps/aether-gateway/src/orchestration/effects.rs index d439a421b..d43f23058 100644 --- a/apps/aether-gateway/src/orchestration/effects.rs +++ b/apps/aether-gateway/src/orchestration/effects.rs @@ -53,7 +53,6 @@ use crate::scheduler::affinity::{ scheduler_affinity_policy_context_from_report_context, SCHEDULER_AFFINITY_POLICY_REPORT_FIELD, SCHEDULER_AFFINITY_TTL, }; -use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode}; use crate::AppState; const POOL_SCORE_FEEDBACK_GATE_MAX_ENTRIES: usize = 50_000; @@ -763,36 +762,19 @@ async fn local_scheduler_affinity_matches_failed_target( local_execution_plan_uses_pool(state, plan).await } -async fn scheduler_cache_affinity_enabled( - state: &AppState, - report_context: Option<&Value>, -) -> bool { - if report_context +fn scheduler_cache_affinity_enabled(report_context: Option<&Value>) -> bool { + report_context .and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD)) .is_some() - { - return scheduler_affinity_policy_context_from_report_context(report_context) - .is_some_and(|context| context.cache_affinity_enabled()); - } - match read_scheduler_ordering_config(state).await { - Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, - Err(error) => { - warn!( - event_name = "orchestration_scheduler_affinity_config_load_failed", - log_type = "event", - error = ?error, - "failed to load scheduler config while checking cache affinity mode" - ); - SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity - } - } + && scheduler_affinity_policy_context_from_report_context(report_context) + .is_some_and(|context| context.cache_affinity_enabled()) } async fn remember_successful_local_scheduler_affinity( state: &AppState, context: LocalExecutionEffectContext<'_>, ) { - if !scheduler_cache_affinity_enabled(state, context.report_context).await { + if !scheduler_cache_affinity_enabled(context.report_context) { return; } let Some(cache_key) = local_scheduler_affinity_cache_key(context.report_context) else { @@ -2060,6 +2042,15 @@ fn pool_score_hard_state_for_status( return Some(pool_score_hard_state_for_terminal_error_reason(&reason)); } + // A number of providers report account quota exhaustion as HTTP 429 rather + // than 402. Keep those members out of the score-based pool fallback until + // the provider's quota probe observes a reset; treating every 429 as a + // generic cooldown otherwise lets the member re-enter as soon as the short + // transient cooldown expires. + if status_code == 429 && error_body_indicates_quota_exhaustion(error_body) { + return Some(PoolMemberHardState::QuotaExhausted); + } + match status_code { 401 | 403 => Some(PoolMemberHardState::AuthInvalid), 402 => Some(PoolMemberHardState::QuotaExhausted), @@ -2082,6 +2073,27 @@ fn pool_score_hard_state_for_status( } } +fn error_body_indicates_quota_exhaustion(error_body: Option<&str>) -> bool { + let body = error_body.unwrap_or_default().to_ascii_lowercase(); + [ + "quota exhausted", + "quota_exhausted", + "quota exceeded", + "quota_exceeded", + "insufficient_quota", + "resource exhausted", + "resource has been exhausted", + "resource_exhausted", + "usage_limit_reached", + "limit_reached", + "quota limit reached", + "credits exhausted", + "insufficient credits", + ] + .iter() + .any(|marker| body.contains(marker)) +} + fn pool_score_hard_state_for_terminal_error_reason(reason: &str) -> PoolMemberHardState { if reason.starts_with("payment_required_") { PoolMemberHardState::QuotaExhausted @@ -2274,6 +2286,9 @@ mod tests { "api_key_id": "api-key-1", "client_api_format": "openai:chat", "model": "gpt-5", + "scheduler_affinity_policy": { + "scheduling_mode": "cache_affinity" + }, "client_session_affinity": { "client_family": "generic", "session_key": "session=session-1;agent=coder" @@ -2288,6 +2303,17 @@ mod tests { }) } + fn cache_affinity_report_context() -> Value { + json!({ + "api_key_id": "api-key-1", + "client_api_format": "openai:chat", + "model": "gpt-5", + "scheduler_affinity_policy": { + "scheduling_mode": "cache_affinity" + } + }) + } + fn session_scheduler_affinity_cache_key() -> String { build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( "api-key-1", @@ -2389,14 +2415,27 @@ mod tests { } fn sample_codex_key() -> StoredProviderCatalogKey { - let encrypted_auth_config = encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","refresh_token":"rt-codex-local-123"}"#, - ) - .expect("auth config should encrypt"); + let provider_id = "provider-codex-cli-local-1"; + let key_id = "key-codex-cli-local-1"; + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key(provider_id, key_id, "codex-access-token") + .expect("access token should encrypt"); + let encrypted_auth_config = credential_state + .seal_provider_catalog_key_auth_config( + provider_id, + key_id, + r#"{"provider_type":"codex","refresh_token":"rt-codex-local-123"}"#, + ) + .expect("auth config should encrypt"); StoredProviderCatalogKey::new( - "key-codex-cli-local-1".to_string(), - "provider-codex-cli-local-1".to_string(), + key_id.to_string(), + provider_id.to_string(), "oauth".to_string(), "oauth".to_string(), None, @@ -2405,8 +2444,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(serde_json::json!(["openai:responses"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "codex-access-token") - .expect("access token should encrypt"), + encrypted_api_key, Some(encrypted_auth_config), None, Some(serde_json::json!({"openai:responses": 1})), @@ -3007,11 +3045,7 @@ mod tests { async fn stream_success_effect_helper_projects_health_and_scheduler_affinity() { let state = AppState::new().expect("gateway state should build"); let plan = sample_plan(); - let report_context = json!({ - "api_key_id": "api-key-1", - "client_api_format": "openai:chat", - "model": "gpt-5", - }); + let report_context = cache_affinity_report_context(); let cache_key = build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") .expect("scheduler affinity cache key should build"); @@ -3122,11 +3156,7 @@ mod tests { async fn success_remembers_scheduler_affinity_cache_for_final_candidate() { let state = AppState::new().expect("gateway state should build"); let plan = sample_plan(); - let report_context = json!({ - "api_key_id": "api-key-1", - "client_api_format": "openai:chat", - "model": "gpt-5", - }); + let report_context = cache_affinity_report_context(); let cache_key = build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") .expect("scheduler affinity cache key should build"); @@ -3293,11 +3323,7 @@ mod tests { async fn health_success_keeps_scheduler_affinity_after_health_state_update() { let state = health_state(); let plan = sample_plan(); - let report_context = json!({ - "api_key_id": "api-key-1", - "client_api_format": "openai:chat", - "model": "gpt-5", - }); + let report_context = cache_affinity_report_context(); let cache_key = build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") .expect("scheduler affinity cache key should build"); @@ -3324,19 +3350,15 @@ mod tests { #[tokio::test] async fn load_balance_success_does_not_remember_scheduler_affinity_cache() { - let state = AppState::new() - .expect("gateway state should build") - .with_data_state_for_tests( - GatewayDataState::disabled().with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("load_balance"), - )]), - ); + let state = AppState::new().expect("gateway state should build"); let plan = sample_plan(); let report_context = json!({ "api_key_id": "api-key-1", "client_api_format": "openai:chat", "model": "gpt-5", + "scheduler_affinity_policy": { + "scheduling_mode": "load_balance" + } }); let cache_key = build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") @@ -3399,11 +3421,7 @@ mod tests { success_plan.provider_id = "prov-2".to_string(); success_plan.endpoint_id = "ep-2".to_string(); success_plan.key_id = "key-2".to_string(); - let report_context = json!({ - "api_key_id": "api-key-1", - "client_api_format": "openai:chat", - "model": "gpt-5", - }); + let report_context = cache_affinity_report_context(); let cache_key = build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") .expect("scheduler affinity cache key should build"); @@ -3659,6 +3677,13 @@ mod tests { ), Some(PoolMemberHardState::QuotaExhausted) ); + assert_eq!( + pool_score_hard_state_for_status( + 429, + Some(r#"{"error":{"status":"RESOURCE_EXHAUSTED","message":"quota exhausted"}}"#), + ), + Some(PoolMemberHardState::QuotaExhausted) + ); } #[tokio::test] @@ -3895,7 +3920,7 @@ mod tests { assert!(stored_key.oauth_invalid_at_unix_secs.is_some()); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] session expired") + Some("[OAUTH_EXPIRED] Codex Token 已过期") ); assert_eq!( stored_key @@ -4207,7 +4232,7 @@ mod tests { assert!(stored_key.oauth_invalid_at_unix_secs.is_some()); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] Codex Token 已失效 (403): forbidden") + Some("[OAUTH_EXPIRED] Codex Token 已失效 (403)") ); assert_eq!( stored_key @@ -4250,7 +4275,7 @@ mod tests { assert!(stored_key.oauth_invalid_at_unix_secs.is_some()); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] Personal access token owner is inactive.") + Some("[OAUTH_EXPIRED] Codex Token 已失效") ); assert_eq!( stored_key @@ -4292,7 +4317,7 @@ mod tests { .expect("recoverable token invalidation should retain the key"); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] Personal access token owner is inactive.") + Some("[OAUTH_EXPIRED] Codex Token 已失效") ); } @@ -4352,7 +4377,7 @@ mod tests { .expect("recoverable expired token should be retained"); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] session expired") + Some("[OAUTH_EXPIRED] Codex Token 已过期") ); } @@ -4555,7 +4580,7 @@ mod tests { .expect("stored key should exist"); assert_eq!( stored_key.oauth_invalid_reason.as_deref(), - Some("[OAUTH_EXPIRED] session expired") + Some("[OAUTH_EXPIRED] Codex Token 已过期") ); assert!(stored_key.oauth_invalid_at_unix_secs.is_some()); assert_eq!( @@ -5019,8 +5044,13 @@ mod tests { #[tokio::test] async fn health_success_projection_is_rate_limited_until_failure_resets_gate() { - let state = health_state(); - let plan = sample_plan(); + // Keep this test's process-wide persistence gate isolated from the other + // effect tests, which intentionally exercise the same health key in parallel. + let mut plan = sample_plan(); + plan.key_id = format!("health-success-rate-limit-{}", uuid::Uuid::new_v4()); + let mut key = sample_health_key(); + key.id = plan.key_id.clone(); + let state = health_state_with_key(key); apply_local_execution_effect( &state, diff --git a/apps/aether-gateway/src/orchestration/mod.rs b/apps/aether-gateway/src/orchestration/mod.rs index 8b86eec98..4c2decec4 100644 --- a/apps/aether-gateway/src/orchestration/mod.rs +++ b/apps/aether-gateway/src/orchestration/mod.rs @@ -20,11 +20,10 @@ pub(crate) use self::adaptive::{ LocalAdaptiveRateLimitProjection, LocalAdaptiveSuccessProjection, }; pub(crate) use self::attempt::{ - attempt_identity_from_report_context, build_local_attempt_identities, - insert_pool_key_lease_report_context_fields, local_attempt_slot_count, - local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity, - LocalExecutionCandidateMetadata, ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, - SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, + attempt_identity_from_report_context, insert_pool_key_lease_report_context_fields, + local_execution_candidate_metadata_from_report_context, next_same_key_retry_attempt, + ExecutionAttemptIdentity, LocalExecutionCandidateMetadata, POOL_KEY_RETRY_INDEX_STRIDE, + ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, }; pub(crate) use self::classifier::{ classify_anthropic_failure_disposition, classify_failure_disposition, classify_local_failover, @@ -57,10 +56,11 @@ pub(crate) use self::oauth_error::{ }; pub(crate) use self::policy::{ append_local_failover_policy_to_value, codex_cyber_flag_passthrough_enabled, - cyber_continue_failover_enabled, local_failover_policy_from_report_context, - local_failover_policy_from_transport, resolve_local_failover_policy, - responses_websocket_adapter, LocalFailoverPolicy, LocalFailoverRegexRule, - ResponsesWebSocketAdapter, CYBER_CONTINUE_FAILOVER_CONFIG_KEY, RESPONSES_WEBSOCKET_CONFIG_KEY, + local_failover_policy_from_report_context, local_failover_policy_from_transport, + resolve_local_failover_policy, responses_websocket_adapter, + routing_execution_policy_from_report_context, LocalFailoverPolicy, LocalFailoverRegexRule, + ResponsesWebSocketAdapter, RESPONSES_WEBSOCKET_CONFIG_KEY, + ROUTING_EXECUTION_POLICY_REPORT_FIELD, }; pub(crate) use self::recovery::{ analyze_local_failover, analyze_local_transport_error, apply_provider_failure_disposition, @@ -264,7 +264,16 @@ fn mask_trace_header_value(name: &str, value: &str) -> String { if value.len() <= 8 { return "****".to_string(); } - format!("{}****{}", &value[..4], &value[value.len() - 4..]) + let prefix = value.chars().take(4).collect::(); + let suffix = value + .chars() + .rev() + .take(4) + .collect::() + .chars() + .rev() + .collect::(); + format!("{prefix}****{suffix}") } fn trace_header_is_sensitive(name: &str) -> bool { @@ -280,3 +289,24 @@ fn trace_header_is_sensitive(name: &str) -> bool { .iter() .any(|candidate| name.trim().eq_ignore_ascii_case(candidate)) } + +#[cfg(test)] +mod tests { + use super::mask_trace_header_value; + + #[test] + fn sensitive_header_masking_is_safe_for_unicode_values() { + let masked = mask_trace_header_value("authorization", "令牌值-абвгдеж"); + assert!(masked.starts_with("令牌值-")); + assert!(masked.contains("****")); + assert!(masked.ends_with("гдеж")); + } + + #[test] + fn sensitive_header_masking_preserves_ascii_shape() { + assert_eq!( + mask_trace_header_value("x-api-key", "abcdefghijk"), + "abcd****hijk" + ); + } +} diff --git a/apps/aether-gateway/src/orchestration/policy.rs b/apps/aether-gateway/src/orchestration/policy.rs index f1d3731a3..0722a11cd 100644 --- a/apps/aether-gateway/src/orchestration/policy.rs +++ b/apps/aether-gateway/src/orchestration/policy.rs @@ -4,11 +4,13 @@ use aether_contracts::ExecutionPlan; use serde_json::{json, Value}; use tracing::debug; +use aether_routing_core::RoutingExecutionPolicy; + use crate::provider_transport::GatewayProviderTransportSnapshot; use crate::AppState; -pub(crate) const CYBER_CONTINUE_FAILOVER_CONFIG_KEY: &str = "cyber_continue_failover"; pub(crate) const RESPONSES_WEBSOCKET_CONFIG_KEY: &str = "responses_websocket"; +pub(crate) const ROUTING_EXECUTION_POLICY_REPORT_FIELD: &str = "routing_execution_policy"; #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct LocalFailoverPolicy { @@ -50,7 +52,7 @@ pub(crate) struct LocalFailoverRegexRule { pub(crate) async fn resolve_local_failover_policy( state: &AppState, plan: &ExecutionPlan, - _report_context: Option<&serde_json::Value>, + report_context: Option<&serde_json::Value>, ) -> LocalFailoverPolicy { let mut policy = match state .read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id) @@ -59,7 +61,8 @@ pub(crate) async fn resolve_local_failover_policy( Ok(Some(transport)) => local_failover_policy_from_transport(&transport), Ok(None) | Err(_) => LocalFailoverPolicy::default(), }; - let cyber_continue_failover = cyber_continue_failover_enabled(state).await; + let cyber_continue_failover = routing_execution_policy_from_report_context(report_context) + .is_some_and(|policy| policy.cyber_continue_failover); policy.stop_cyber_policy_errors = !cyber_continue_failover; debug!( event_name = "local_failover_policy_loaded", @@ -83,15 +86,13 @@ pub(crate) async fn resolve_local_failover_policy( policy } -pub(crate) async fn cyber_continue_failover_enabled(state: &AppState) -> bool { - state - .read_system_config_json_value(CYBER_CONTINUE_FAILOVER_CONFIG_KEY) - .await - .ok() - .flatten() - .as_ref() - .and_then(Value::as_bool) - .unwrap_or(false) +pub(crate) fn routing_execution_policy_from_report_context( + report_context: Option<&Value>, +) -> Option { + report_context + .and_then(Value::as_object) + .and_then(|object| object.get(ROUTING_EXECUTION_POLICY_REPORT_FIELD)) + .and_then(|value| serde_json::from_value(value.clone()).ok()) } pub(crate) fn local_failover_policy_from_transport( diff --git a/apps/aether-gateway/src/orchestration/report_effects.rs b/apps/aether-gateway/src/orchestration/report_effects.rs index 12cef61c7..7e28a9e78 100644 --- a/apps/aether-gateway/src/orchestration/report_effects.rs +++ b/apps/aether-gateway/src/orchestration/report_effects.rs @@ -6,11 +6,10 @@ use aether_admin::provider::quota as admin_provider_quota_pure; use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate; use aether_provider_pool::grok_quota_window_key_for_model; use aether_usage_runtime::{ - extract_gemini_file_mapping_entries, gemini_file_mapping_cache_key, normalize_gemini_file_name, - report_request_id, GatewayStreamReportRequest, GatewaySyncReportRequest, - GEMINI_FILE_MAPPING_TTL_SECONDS, + decode_internal_report_body_base64, extract_gemini_file_mapping_entries, + gemini_file_mapping_cache_key, normalize_gemini_file_name, report_request_id, + GatewayStreamReportRequest, GatewaySyncReportRequest, GEMINI_FILE_MAPPING_TTL_SECONDS, }; -use base64::Engine as _; use regex::Regex; use serde_json::{json, Value}; use tracing::warn; @@ -241,9 +240,7 @@ fn gemini_cli_credits_from_stream_payload( now_unix_secs: u64, ) -> Option { let body_base64 = payload.provider_body_base64.as_deref()?; - let body = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .ok()?; + let body = decode_internal_report_body_base64(body_base64).ok()?; let text = std::str::from_utf8(&body).ok()?; let mut latest = None::; for raw_line in text.lines() { @@ -272,9 +269,7 @@ fn codex_websocket_quota_from_stream_payload( now_unix_secs: u64, ) -> Option { let body_base64 = payload.provider_body_base64.as_deref()?; - let body = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .ok()?; + let body = decode_internal_report_body_base64(body_base64).ok()?; let text = std::str::from_utf8(&body).ok()?; let mut latest = None::; for raw_line in text.lines() { @@ -742,8 +737,30 @@ async fn apply_local_gemini_file_mapping_report_effect( let Some(file_name) = file_name else { return; }; + let user_id = payload + .report_context + .as_ref() + .and_then(|context| context.get("user_id")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + let key_id = payload + .report_context + .as_ref() + .and_then(|context| context.get("key_id")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + let Some(user_id) = user_id else { + return; + }; + let Some(key_id) = key_id else { + return; + }; - if let Err(err) = delete_local_gemini_file_mapping(state, file_name.as_str()).await { + if let Err(err) = + delete_local_gemini_file_mapping(state, file_name.as_str(), key_id, user_id).await + { warn!( event_name = "gemini_file_mapping_delete_failed", log_type = "ops", @@ -772,8 +789,8 @@ pub(crate) async fn store_local_gemini_file_mapping( }; let expires_at_unix_secs = current_unix_secs().saturating_add(GEMINI_FILE_MAPPING_TTL_SECONDS); - let _stored = state - .upsert_gemini_file_mapping( + let stored = state + .upsert_gemini_file_mapping_if_owner_matches( aether_data::repository::gemini_file_mappings::UpsertGeminiFileMappingRecord { id: Uuid::new_v4().to_string(), file_name: file_name.clone(), @@ -786,6 +803,17 @@ pub(crate) async fn store_local_gemini_file_mapping( }, ) .await?; + if stored.is_none() { + warn!( + event_name = "gemini_file_mapping_atomic_owner_mismatch", + log_type = "security", + file_name = %file_name, + requested_key_id = %key_id, + requested_user_id = user_id.unwrap_or_default(), + "gateway refused to reassign a Gemini file mapping during the atomic write" + ); + return Ok(()); + } state .cache_set_string_with_ttl( gemini_file_mapping_cache_key(file_name.as_str()).as_str(), @@ -799,14 +827,26 @@ pub(crate) async fn store_local_gemini_file_mapping( async fn delete_local_gemini_file_mapping( state: &AppState, file_name: &str, + key_id: &str, + user_id: &str, ) -> Result<(), GatewayError> { let Some(file_name) = normalize_gemini_file_name(file_name) else { return Ok(()); }; - - let _deleted = state - .delete_gemini_file_mapping_by_file_name(file_name.as_str()) + let deleted = state + .delete_gemini_file_mapping_by_file_name_for_owner(file_name.as_str(), key_id, user_id) .await?; + if !deleted { + warn!( + event_name = "gemini_file_mapping_atomic_delete_owner_mismatch", + log_type = "security", + file_name = %file_name, + requested_key_id = %key_id, + requested_user_id = %user_id, + "gateway refused to delete a Gemini file mapping after its owner changed" + ); + return Ok(()); + } state .cache_delete_key(gemini_file_mapping_cache_key(file_name.as_str()).as_str()) .await?; diff --git a/apps/aether-gateway/src/plan_usage_policy.rs b/apps/aether-gateway/src/plan_usage_policy.rs new file mode 100644 index 000000000..5d2ef8d16 --- /dev/null +++ b/apps/aether-gateway/src/plan_usage_policy.rs @@ -0,0 +1,1487 @@ +use std::collections::BTreeMap; +use std::sync::Arc; +use std::time::Duration; + +use aether_data_contracts::repository::billing::{ + nonnegative_usd_to_usage_policy_cost_units, parse_usage_policy_entitlements, UsagePolicyMetric, + UsagePolicyWindow, UserPlanEntitlementRecord, USAGE_POLICY_COST_UNITS_PER_USD, +}; +use aether_data_contracts::repository::settlement::{ + ReconcileUsagePolicyCostInput, ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, + ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome, + UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestWindow, +}; +use aether_runtime::AdmissionPermit; +use aether_runtime_state::{ + RuntimeSemaphoreConfig, RuntimeSemaphoreError, UsageLimitCheck, UsageLimitInput, + UsageLimitReleaseInput, UsageLimitRule, +}; +use chrono::{Datelike, TimeZone, Utc, Weekday}; +use tracing::warn; + +use crate::control::GatewayControlDecision; +use crate::{AppState, GatewayError}; + +const POLICY_CACHE_TTL: Duration = Duration::from_secs(5); +const CONCURRENCY_GATE: &str = "plan_usage_concurrency"; +const DEFAULT_CALENDAR_TIMEZONE: &str = "Asia/Shanghai"; +const COST_RESERVATION_TTL_SECS: u64 = 24 * 60 * 60; +const COST_RESERVATION_SAFE_HISTORY_SECS: u64 = 32 * 24 * 60 * 60; + +#[derive(Debug, Clone)] +pub(crate) struct PlanUsagePolicySnapshot { + pub(crate) admitted_at_unix_secs: u64, + subject_id: Arc, + policy: Arc, +} + +impl PlanUsagePolicySnapshot { + fn for_admission( + subject_id: &str, + policy: EffectivePlanUsagePolicy, + admitted_at_unix_secs: u64, + ) -> Option { + if policy.cost_rules.is_empty() { + return None; + } + Some(Self { + admitted_at_unix_secs, + subject_id: subject_id.to_string().into(), + policy: Arc::new(policy), + }) + } + + pub(crate) fn new_reservation_context(&self) -> PlanUsageReservationContext { + PlanUsageReservationContext { + policy_snapshot: self.clone(), + token: uuid::Uuid::new_v4().to_string().into(), + } + } + + pub(crate) fn subject_id(&self) -> &str { + self.subject_id.as_ref() + } + + pub(crate) fn policy(&self) -> &EffectivePlanUsagePolicy { + self.policy.as_ref() + } +} + +#[derive(Debug, Clone)] +pub(crate) struct PlanUsageReservationContext { + policy_snapshot: PlanUsagePolicySnapshot, + token: Arc, +} + +impl PlanUsageReservationContext { + #[cfg(test)] + pub(crate) fn for_test( + subject_id: impl Into>, + token: impl Into>, + admitted_at_unix_secs: u64, + policy: EffectivePlanUsagePolicy, + ) -> Self { + Self { + policy_snapshot: PlanUsagePolicySnapshot { + admitted_at_unix_secs, + subject_id: subject_id.into(), + policy: Arc::new(policy), + }, + token: token.into(), + } + } + + pub(crate) fn subject_id(&self) -> &str { + self.policy_snapshot.subject_id() + } + + pub(crate) fn token(&self) -> &str { + self.token.as_ref() + } + + pub(crate) fn policy(&self) -> &EffectivePlanUsagePolicy { + self.policy_snapshot.policy() + } + + pub(crate) const fn admitted_at_unix_secs(&self) -> u64 { + self.policy_snapshot.admitted_at_unix_secs + } +} + +#[derive(Debug)] +pub(crate) struct HttpPlanUsageAdmission { + pub(crate) permit: Option, + pub(crate) reservation_context: Option, +} + +#[derive(Debug)] +pub(crate) struct PlanUsageAdmission { + pub(crate) permit: Option, + pub(crate) policy_snapshot: Option, +} + +#[derive(Debug, Clone, PartialEq, Default)] +pub(crate) struct EffectivePlanUsagePolicy { + request_rules: Vec, + cost_rules: Vec, + concurrency_limit: Option, + valid_until_unix_secs: Option, +} + +#[derive(Debug, Clone, PartialEq)] +struct EffectiveRequestRule { + identity: String, + window: UsagePolicyWindow, + limit: u64, + entitlement_period: Option<(u64, u64)>, +} + +#[derive(Debug, Clone, PartialEq)] +struct EffectiveCostRule { + identity: String, + window: UsagePolicyWindow, + limit_cost_units: u64, + entitlement_period: Option<(u64, u64)>, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct RuntimeRequestRule { + key: String, + limit: u64, + window_seconds: u64, + retention_seconds: u64, + retry_after: u64, + window: &'static str, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct DurableRequestRule { + window: UsagePolicyRequestWindow, + retry_after: u64, + label: &'static str, + influence_ends_at_unix_secs: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct RuntimeCostRule { + window: UsagePolicyCostWindow, + retry_after: u64, + label: &'static str, + influence_ends_at_unix_secs: u64, +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct PlanUsagePolicyRejection { + pub(crate) metric: &'static str, + pub(crate) limit: f64, + pub(crate) retry_after: u64, + pub(crate) window: &'static str, +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) enum PlanUsageCostReservationOutcome { + NotRequired, + Reserved, + Rejected(PlanUsagePolicyRejection), +} + +pub(crate) async fn reserve_admitted_http_plan_usage_policy_cost( + state: &AppState, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + report_context: Option<&serde_json::Value>, + reservation: Option<&PlanUsageReservationContext>, +) -> Result { + let Some(auth) = plan_usage_auth_context(decision) else { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + }; + let Some(reservation) = reservation else { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + }; + if reservation.subject_id() != auth.user_id { + return Err(GatewayError::Internal( + "plan usage reservation subject does not match the admitted request".to_string(), + )); + } + reserve_plan_usage_policy_cost_with_policy( + state, + decision, + plan, + report_context, + reservation.policy(), + reservation.admitted_at_unix_secs(), + reservation.token(), + ) + .await +} + +pub(crate) async fn reserve_admitted_plan_usage_policy_cost( + state: &AppState, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + report_context: Option<&serde_json::Value>, + snapshot: Option<&PlanUsagePolicySnapshot>, + reservation_token: &str, +) -> Result { + let Some(auth) = plan_usage_auth_context(decision) else { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + }; + let Some(snapshot) = snapshot else { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + }; + if snapshot.subject_id() != auth.user_id { + return Err(GatewayError::Internal( + "plan usage reservation subject does not match the admitted request".to_string(), + )); + } + reserve_plan_usage_policy_cost_with_policy( + state, + decision, + plan, + report_context, + snapshot.policy(), + snapshot.admitted_at_unix_secs, + reservation_token, + ) + .await +} + +fn plan_usage_auth_context( + decision: &GatewayControlDecision, +) -> Option<&crate::control::GatewayControlAuthContext> { + decision.auth_context.as_ref().filter(|auth| { + decision.route_class.as_deref() == Some("ai_public") + && !auth.admin_bypass_limits + && !auth.api_key_is_standalone + }) +} + +#[allow(clippy::too_many_arguments)] +async fn reserve_plan_usage_policy_cost_with_policy( + state: &AppState, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + report_context: Option<&serde_json::Value>, + policy: &EffectivePlanUsagePolicy, + admitted_at_unix_secs: u64, + reservation_token: &str, +) -> Result { + let Some(auth) = plan_usage_auth_context(decision) else { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + }; + if policy.cost_rules.is_empty() { + return Ok(PlanUsageCostReservationOutcome::NotRequired); + } + ensure_cost_policy_usage_runtime_enabled(state.usage_runtime.is_enabled())?; + let estimated_cost_usd = + crate::control::estimate_execution_plan_cost_upper_bound_usd(state, plan, report_context) + .await? + .ok_or_else(|| { + GatewayError::Internal( + "a hard plan cost limit is active, but this request cost cannot be estimated safely" + .to_string(), + ) + })?; + let reserved_cost_units = nonnegative_usd_to_usage_policy_cost_units(estimated_cost_usd) + .ok_or_else(|| { + GatewayError::Internal( + "estimated request cost is outside the supported plan limit range".to_string(), + ) + })?; + let runtime_rules = policy + .cost_rules + .iter() + .map(|rule| runtime_cost_rule(rule, admitted_at_unix_secs)) + .collect::, _>>()?; + let reservation_expires_at_unix_secs = + admitted_at_unix_secs.saturating_add(COST_RESERVATION_TTL_SECS); + let retain_until_unix_secs = runtime_rules + .iter() + .map(|rule| rule.influence_ends_at_unix_secs) + .chain(std::iter::once(reservation_expires_at_unix_secs)) + .chain(std::iter::once( + admitted_at_unix_secs.saturating_add(COST_RESERVATION_SAFE_HISTORY_SECS), + )) + .max() + .unwrap_or(reservation_expires_at_unix_secs); + let outcome = state + .data + .reserve_usage_policy_cost(ReserveUsagePolicyCostInput { + request_id: plan.request_id.clone(), + reservation_token: reservation_token.to_string(), + subject_id: auth.user_id.clone(), + admitted_at_unix_secs, + reserved_cost_units, + reservation_expires_at_unix_secs, + retain_until_unix_secs, + windows: runtime_rules + .iter() + .map(|rule| rule.window.clone()) + .collect(), + }) + .await + .map_err(|error| GatewayError::Internal(error.to_string()))? + .ok_or_else(|| { + GatewayError::Internal( + "plan cost limits require a settlement write repository".to_string(), + ) + })?; + + match outcome { + ReserveUsagePolicyCostOutcome::Allowed { .. } => { + Ok(PlanUsageCostReservationOutcome::Reserved) + } + ReserveUsagePolicyCostOutcome::Rejected { + window_index, + limit_cost_units, + .. + } => { + let rule = runtime_rules.get(window_index).ok_or_else(|| { + GatewayError::Internal( + "usage policy cost reservation returned an invalid window index".to_string(), + ) + })?; + Ok(PlanUsageCostReservationOutcome::Rejected( + PlanUsagePolicyRejection { + metric: "actual_cost_usd", + limit: limit_cost_units as f64 / USAGE_POLICY_COST_UNITS_PER_USD as f64, + retry_after: rule.retry_after, + window: rule.label, + }, + )) + } + ReserveUsagePolicyCostOutcome::AlreadyTerminal { state } => { + Err(GatewayError::Internal(format!( + "plan cost reservation for request {} is already {}", + plan.request_id, + state.as_str() + ))) + } + ReserveUsagePolicyCostOutcome::Conflict => Err(GatewayError::Internal( + "plan cost reservation request identity conflict".to_string(), + )), + } +} + +pub(crate) async fn release_plan_usage_policy_cost( + state: &AppState, + decision: &GatewayControlDecision, + plan: &aether_contracts::ExecutionPlan, + reservation_token: &str, + finalized_at_unix_secs: u64, +) -> Result<(), GatewayError> { + let Some(auth) = decision + .auth_context + .as_ref() + .filter(|_| decision.route_class.as_deref() == Some("ai_public")) + else { + return Ok(()); + }; + if auth.admin_bypass_limits || auth.api_key_is_standalone { + return Ok(()); + } + + state + .data + .reconcile_usage_policy_cost(build_usage_policy_cost_release_input( + plan.request_id.as_str(), + auth.user_id.as_str(), + reservation_token, + finalized_at_unix_secs, + )) + .await + .map_err(|error| GatewayError::Internal(error.to_string()))?; + Ok(()) +} + +fn build_usage_policy_cost_release_input( + request_id: &str, + subject_id: &str, + reservation_token: &str, + finalized_at_unix_secs: u64, +) -> ReconcileUsagePolicyCostInput { + ReconcileUsagePolicyCostInput { + request_id: request_id.to_string(), + subject_id: subject_id.to_string(), + reservation_token: reservation_token.to_string(), + actual_cost_units: 0, + terminal_state: UsagePolicyCostReservationState::Released, + finalized_at_unix_secs, + } +} + +fn ensure_cost_policy_usage_runtime_enabled(enabled: bool) -> Result<(), GatewayError> { + if enabled { + Ok(()) + } else { + Err(GatewayError::Internal( + "plan cost limits require the usage runtime to be enabled for terminal reconciliation" + .to_string(), + )) + } +} + +#[derive(Debug)] +pub(crate) enum PlanUsageAdmissionError { + Rejected(PlanUsagePolicyRejection), + Runtime(RuntimeSemaphoreError), + Gateway(GatewayError), +} + +impl From for PlanUsageAdmissionError { + fn from(value: GatewayError) -> Self { + Self::Gateway(value) + } +} + +pub(crate) async fn check_and_acquire_plan_usage_policy_admission( + state: &AppState, + decision: Option<&GatewayControlDecision>, + event_id: &str, + now_unix_ms: u64, +) -> Result { + let Some(auth) = decision + .filter(|decision| decision.route_class.as_deref() == Some("ai_public")) + .and_then(|decision| decision.auth_context.as_ref()) + else { + return Ok(PlanUsageAdmission { + permit: None, + policy_snapshot: None, + }); + }; + if auth.admin_bypass_limits || auth.api_key_is_standalone { + return Ok(PlanUsageAdmission { + permit: None, + policy_snapshot: None, + }); + } + + let admitted_at_unix_secs = now_unix_ms / 1_000; + let policy = load_effective_policy(state, &auth.user_id, admitted_at_unix_secs).await?; + let permit = check_and_acquire_compiled_plan_usage_policy( + state, + &auth.user_id, + &policy, + event_id, + now_unix_ms, + ) + .await?; + let policy_snapshot = + PlanUsagePolicySnapshot::for_admission(&auth.user_id, policy, admitted_at_unix_secs); + Ok(PlanUsageAdmission { + permit, + policy_snapshot, + }) +} + +async fn check_and_acquire_compiled_plan_usage_policy( + state: &AppState, + subject_id: &str, + policy: &EffectivePlanUsagePolicy, + event_id: &str, + now_unix_ms: u64, +) -> Result, PlanUsageAdmissionError> { + if policy.request_rules.is_empty() && policy.concurrency_limit.is_none() { + return Ok(None); + } + + let now_unix_secs = now_unix_ms / 1_000; + + let plan_permit = if let Some(limit) = policy.concurrency_limit { + let limit = usize::try_from(limit).map_err(|_| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "plan usage concurrency limit exceeds platform capacity".to_string(), + )) + })?; + let gate = state + .runtime_state + .keyed_semaphore( + CONCURRENCY_GATE, + format!("admission:{CONCURRENCY_GATE}:user:{{{subject_id}}}"), + limit, + RuntimeSemaphoreConfig::default(), + ) + .map_err(PlanUsageAdmissionError::Runtime)?; + Some( + gate.try_acquire() + .await + .map_err(PlanUsageAdmissionError::Runtime)?, + ) + } else { + None + }; + + let (runtime_request_rules, durable_request_rules): (Vec<_>, Vec<_>) = policy + .request_rules + .iter() + .partition(|rule| request_rule_uses_runtime_state(rule)); + let runtime_rules = runtime_request_rules + .iter() + .map(|rule| runtime_request_rule(subject_id, rule, now_unix_secs)) + .collect::, _>>()?; + let runtime_inputs = runtime_rules + .iter() + .map(|rule| UsageLimitRule { + key: rule.key.as_str(), + limit: rule.limit, + window_seconds: rule.window_seconds, + retention_seconds: rule.retention_seconds, + }) + .collect::>(); + + // Consume the short-lived runtime windows first. If the durable admission below rejects or + // errors, remove this event from every short window. This ordering bounds a process-crash + // compensation gap to the QPS/RPM retention (at most 60 seconds) instead of leaking a + // calendar/subscription-period admission for weeks or months. + if !runtime_inputs.is_empty() { + match state + .runtime_state + .check_and_consume_usage_limits(UsageLimitInput { + rules: &runtime_inputs, + event_id, + now_unix_ms, + }) + .await + .map_err(|error| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal(error.to_string())) + })? { + UsageLimitCheck::Allowed => {} + UsageLimitCheck::Rejected { + rule_index, + limit, + retry_after, + } => { + let rule = runtime_rules.get(rule_index).ok_or_else(|| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "usage policy runtime returned an invalid rule index".to_string(), + )) + })?; + return Err(PlanUsageAdmissionError::Rejected( + PlanUsagePolicyRejection { + metric: "request_count", + limit: limit as f64, + retry_after: retry_after.min(rule.retry_after).max(1), + window: rule.window, + }, + )); + } + } + } + + let durable_rules = durable_request_rules + .iter() + .map(|rule| durable_request_rule(rule, now_unix_secs)) + .collect::, _>>()?; + if !durable_rules.is_empty() { + let retain_until_unix_secs = durable_rules + .iter() + .map(|rule| rule.influence_ends_at_unix_secs) + .chain(std::iter::once( + now_unix_secs.saturating_add(COST_RESERVATION_SAFE_HISTORY_SECS), + )) + .max() + .unwrap_or(now_unix_secs.saturating_add(COST_RESERVATION_SAFE_HISTORY_SECS)); + let outcome = state + .data + .reserve_usage_policy_request(ReserveUsagePolicyRequestInput { + request_id: event_id.to_string(), + subject_id: subject_id.to_string(), + event_token: event_id.to_string(), + admitted_at_unix_secs: now_unix_secs, + retain_until_unix_secs, + windows: durable_rules + .iter() + .map(|rule| rule.window.clone()) + .collect(), + }) + .await; + let outcome = match outcome { + Ok(Some(outcome)) => outcome, + Ok(None) => { + release_runtime_usage_limits_best_effort( + state, + &runtime_inputs, + event_id, + "settlement_writer_unavailable", + ) + .await; + return Err(PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "long-window plan request limits require a settlement write repository" + .to_string(), + ))); + } + Err(error) => { + release_runtime_usage_limits_best_effort( + state, + &runtime_inputs, + event_id, + "durable_admission_error", + ) + .await; + return Err(PlanUsageAdmissionError::Gateway(GatewayError::Internal( + error.to_string(), + ))); + } + }; + match outcome { + ReserveUsagePolicyRequestOutcome::Allowed => {} + ReserveUsagePolicyRequestOutcome::Rejected { + window_index, + limit_requests, + .. + } => { + release_runtime_usage_limits_best_effort( + state, + &runtime_inputs, + event_id, + "durable_admission_rejected", + ) + .await; + let rule = durable_rules.get(window_index).ok_or_else(|| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "usage policy request ledger returned an invalid rule index".to_string(), + )) + })?; + return Err(PlanUsageAdmissionError::Rejected( + PlanUsagePolicyRejection { + metric: "request_count", + limit: limit_requests as f64, + retry_after: rule.retry_after, + window: rule.label, + }, + )); + } + ReserveUsagePolicyRequestOutcome::AlreadyReleased => { + release_runtime_usage_limits_best_effort( + state, + &runtime_inputs, + event_id, + "durable_admission_already_released", + ) + .await; + return Err(PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "plan request admission token is already released".to_string(), + ))); + } + ReserveUsagePolicyRequestOutcome::Conflict => { + release_runtime_usage_limits_best_effort( + state, + &runtime_inputs, + event_id, + "durable_admission_conflict", + ) + .await; + return Err(PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "plan request admission identity conflict".to_string(), + ))); + } + } + } + + Ok(AdmissionPermit::from_parts(None, plan_permit)) +} + +pub(crate) async fn check_and_acquire_http_plan_usage_policy( + state: &AppState, + decision: Option<&GatewayControlDecision>, + event_id: &str, + now_unix_ms: u64, +) -> Result { + let Some(auth) = decision + .filter(|decision| decision.route_class.as_deref() == Some("ai_public")) + .and_then(|decision| decision.auth_context.as_ref()) + else { + return Ok(HttpPlanUsageAdmission { + permit: None, + reservation_context: None, + }); + }; + if auth.admin_bypass_limits || auth.api_key_is_standalone { + return Ok(HttpPlanUsageAdmission { + permit: None, + reservation_context: None, + }); + } + + let admitted_at_unix_secs = now_unix_ms / 1_000; + let policy = load_effective_policy(state, &auth.user_id, admitted_at_unix_secs).await?; + let permit = check_and_acquire_compiled_plan_usage_policy( + state, + &auth.user_id, + &policy, + event_id, + now_unix_ms, + ) + .await?; + let reservation_context = + PlanUsagePolicySnapshot::for_admission(&auth.user_id, policy, admitted_at_unix_secs) + .map(|snapshot| snapshot.new_reservation_context()); + Ok(HttpPlanUsageAdmission { + permit, + reservation_context, + }) +} + +fn request_rule_uses_runtime_state(rule: &&EffectiveRequestRule) -> bool { + matches!(rule.window, UsagePolicyWindow::Rolling { seconds } if seconds <= 60) +} + +async fn release_runtime_usage_limits_best_effort( + state: &AppState, + rules: &[UsageLimitRule<'_>], + event_id: &str, + reason: &'static str, +) { + if rules.is_empty() { + return; + } + if let Err(error) = state + .runtime_state + .release_usage_limits(UsageLimitReleaseInput { rules, event_id }) + .await + { + warn!( + event_name = "plan_usage_runtime_compensation_failed", + log_type = "ops", + event_id, + reason, + error = %error, + "gateway failed to compensate short-window plan usage after durable admission failed" + ); + } +} + +async fn load_effective_policy( + state: &AppState, + user_id: &str, + now_unix_secs: u64, +) -> Result { + let cache_key = user_id.to_string(); + let policy = state + .auth_plan_usage_policy_cache + .get_or_load(cache_key.clone(), POLICY_CACHE_TTL, || async move { + let _permit = state.acquire_auth_snapshot_load_gate().await?; + let entitlements = state + .list_user_plan_entitlements(user_id) + .await? + .unwrap_or_default(); + compile_effective_policy(&entitlements, now_unix_secs).map(Some) + }) + .await + .map(|policy| policy.unwrap_or_default())?; + if policy + .valid_until_unix_secs + .is_some_and(|valid_until| valid_until <= now_unix_secs) + { + state.auth_plan_usage_policy_cache.clear(); + let _permit = state.acquire_auth_snapshot_load_gate().await?; + let entitlements = state + .list_user_plan_entitlements(user_id) + .await? + .unwrap_or_default(); + let policy = compile_effective_policy(&entitlements, now_unix_secs)?; + state.auth_plan_usage_policy_cache.insert( + cache_key, + Some(policy.clone()), + POLICY_CACHE_TTL, + ); + return Ok(policy); + } + Ok(policy) +} + +fn compile_effective_policy( + entitlements: &[UserPlanEntitlementRecord], + now_unix_secs: u64, +) -> Result { + let mut request_rules = BTreeMap::::new(); + let mut cost_rules = BTreeMap::::new(); + let mut concurrency_limit = None::; + let mut valid_until_unix_secs = None::; + + for entitlement in entitlements.iter().filter(|entitlement| { + entitlement.status == "active" + && entitlement.starts_at_unix_secs <= now_unix_secs + && entitlement.expires_at_unix_secs > now_unix_secs + }) { + valid_until_unix_secs = Some( + valid_until_unix_secs.map_or(entitlement.expires_at_unix_secs, |current| { + current.min(entitlement.expires_at_unix_secs) + }), + ); + let policies = parse_usage_policy_entitlements(&entitlement.entitlements_snapshot) + .map_err(|error| GatewayError::Internal(error.to_string()))?; + for policy in policies { + for rule in policy.rules { + match rule.metric { + UsagePolicyMetric::Concurrency => { + let limit = rule.request_limit().ok_or_else(|| { + GatewayError::Internal( + "validated concurrency usage policy lost its integer limit" + .to_string(), + ) + })?; + concurrency_limit = + Some(concurrency_limit.map_or(limit, |current| current.min(limit))); + } + UsagePolicyMetric::RequestCount => { + let limit = rule.request_limit().ok_or_else(|| { + GatewayError::Internal( + "validated request-count usage policy lost its integer limit" + .to_string(), + ) + })?; + let (identity, entitlement_period) = match &rule.window { + UsagePolicyWindow::SubscriptionPeriod => ( + format!("subscription:{}", entitlement.id), + Some(( + entitlement.starts_at_unix_secs, + entitlement.expires_at_unix_secs, + )), + ), + window => (window_identity(window), None), + }; + let candidate = EffectiveRequestRule { + identity: identity.clone(), + window: rule.window, + limit, + entitlement_period, + }; + request_rules + .entry(identity) + .and_modify(|current| { + if candidate.limit < current.limit { + *current = candidate.clone(); + } + }) + .or_insert(candidate); + } + UsagePolicyMetric::ActualCostUsd => { + let limit_cost_units = rule.cost_limit_units().ok_or_else(|| { + GatewayError::Internal( + "validated cost usage policy lost its fixed-point limit" + .to_string(), + ) + })?; + let (identity, entitlement_period) = match &rule.window { + UsagePolicyWindow::SubscriptionPeriod => ( + format!("subscription:{}", entitlement.id), + Some(( + entitlement.starts_at_unix_secs, + entitlement.expires_at_unix_secs, + )), + ), + window => (window_identity(window), None), + }; + let candidate = EffectiveCostRule { + identity: identity.clone(), + window: rule.window, + limit_cost_units, + entitlement_period, + }; + cost_rules + .entry(identity) + .and_modify(|current| { + if candidate.limit_cost_units < current.limit_cost_units { + *current = candidate.clone(); + } + }) + .or_insert(candidate); + } + } + } + } + } + + Ok(EffectivePlanUsagePolicy { + request_rules: request_rules.into_values().collect(), + cost_rules: cost_rules.into_values().collect(), + concurrency_limit, + valid_until_unix_secs, + }) +} + +fn window_identity(window: &UsagePolicyWindow) -> String { + match window { + UsagePolicyWindow::Rolling { seconds } => format!("rolling:{seconds}"), + UsagePolicyWindow::CalendarDay { timezone } => { + format!( + "calendar_day:{}", + effective_timezone_name(timezone.as_deref()) + ) + } + UsagePolicyWindow::CalendarWeek { + timezone, + week_start, + } => format!( + "calendar_week:{}:{week_start}", + effective_timezone_name(timezone.as_deref()) + ), + UsagePolicyWindow::CalendarMonth { timezone } => { + format!( + "calendar_month:{}", + effective_timezone_name(timezone.as_deref()) + ) + } + UsagePolicyWindow::SubscriptionPeriod => "subscription_period".to_string(), + UsagePolicyWindow::Concurrent => "concurrent".to_string(), + } +} + +fn runtime_request_rule( + user_id: &str, + rule: &EffectiveRequestRule, + now_unix_secs: u64, +) -> Result { + let (bucket, window_seconds, retry_after, window) = match &rule.window { + UsagePolicyWindow::Rolling { seconds } => ( + "sliding".to_string(), + *seconds, + *seconds, + rolling_label(*seconds), + ), + UsagePolicyWindow::CalendarDay { timezone } => { + calendar_bucket(now_unix_secs, timezone.as_deref(), CalendarWindow::Day, 1)? + } + UsagePolicyWindow::CalendarWeek { + timezone, + week_start, + } => calendar_bucket( + now_unix_secs, + timezone.as_deref(), + CalendarWindow::Week, + *week_start, + )?, + UsagePolicyWindow::CalendarMonth { timezone } => { + calendar_bucket(now_unix_secs, timezone.as_deref(), CalendarWindow::Month, 1)? + } + UsagePolicyWindow::SubscriptionPeriod => { + let (starts_at, expires_at) = rule.entitlement_period.ok_or_else(|| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "subscription period usage policy is missing its entitlement period" + .to_string(), + )) + })?; + let window_seconds = expires_at.saturating_sub(starts_at).max(1); + let remaining = expires_at.saturating_sub(now_unix_secs).max(1); + ( + starts_at.to_string(), + window_seconds, + remaining, + "subscription_period", + ) + } + UsagePolicyWindow::Concurrent => { + return Err(PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "concurrent window reached request counter compiler".to_string(), + ))); + } + }; + Ok(RuntimeRequestRule { + key: format!("plan-usage:user:{{{user_id}}}:{}:{bucket}", rule.identity), + limit: rule.limit, + window_seconds: window_seconds.max(1), + retention_seconds: retry_after.max(1), + retry_after: retry_after.max(1), + window, + }) +} + +fn durable_request_rule( + rule: &EffectiveRequestRule, + admitted_at_unix_secs: u64, +) -> Result { + let (starts_at, ends_at, retry_after, label) = + usage_window_bounds(&rule.window, rule.entitlement_period, admitted_at_unix_secs) + .map_err(PlanUsageAdmissionError::Gateway)?; + Ok(DurableRequestRule { + window: UsagePolicyRequestWindow { + starts_at_unix_secs: starts_at, + ends_at_unix_secs: ends_at, + limit_requests: rule.limit, + }, + retry_after, + label, + influence_ends_at_unix_secs: match &rule.window { + UsagePolicyWindow::Rolling { seconds } => admitted_at_unix_secs + .saturating_add(*seconds) + .saturating_add(1), + _ => ends_at, + }, + }) +} + +fn runtime_cost_rule( + rule: &EffectiveCostRule, + admitted_at_unix_secs: u64, +) -> Result { + let (starts_at, ends_at, retry_after, label) = + usage_window_bounds(&rule.window, rule.entitlement_period, admitted_at_unix_secs)?; + Ok(RuntimeCostRule { + window: UsagePolicyCostWindow { + window_id: format!("{}:{starts_at}", rule.identity), + starts_at_unix_secs: starts_at, + ends_at_unix_secs: ends_at, + limit_cost_units: rule.limit_cost_units, + }, + retry_after, + label, + influence_ends_at_unix_secs: match &rule.window { + UsagePolicyWindow::Rolling { seconds } => admitted_at_unix_secs + .saturating_add(*seconds) + .saturating_add(1), + _ => ends_at, + }, + }) +} + +fn usage_window_bounds( + window: &UsagePolicyWindow, + entitlement_period: Option<(u64, u64)>, + now_unix_secs: u64, +) -> Result<(u64, u64, u64, &'static str), GatewayError> { + match window { + UsagePolicyWindow::Rolling { seconds } => Ok(( + now_unix_secs.saturating_sub(*seconds), + now_unix_secs.saturating_add(1), + *seconds, + rolling_label(*seconds), + )), + UsagePolicyWindow::CalendarDay { timezone } => { + calendar_window_bounds(now_unix_secs, timezone.as_deref(), CalendarWindow::Day, 1) + } + UsagePolicyWindow::CalendarWeek { + timezone, + week_start, + } => calendar_window_bounds( + now_unix_secs, + timezone.as_deref(), + CalendarWindow::Week, + *week_start, + ), + UsagePolicyWindow::CalendarMonth { timezone } => { + calendar_window_bounds(now_unix_secs, timezone.as_deref(), CalendarWindow::Month, 1) + } + UsagePolicyWindow::SubscriptionPeriod => { + let (starts_at, ends_at) = entitlement_period.ok_or_else(|| { + GatewayError::Internal( + "subscription period cost policy is missing its entitlement period".to_string(), + ) + })?; + Ok(( + starts_at, + ends_at, + ends_at.saturating_sub(now_unix_secs).max(1), + "subscription_period", + )) + } + UsagePolicyWindow::Concurrent => Err(GatewayError::Internal( + "concurrent window reached cost policy compiler".to_string(), + )), + } +} + +fn rolling_label(seconds: u64) -> &'static str { + match seconds { + 1 => "qps", + 60 => "rpm", + _ => "rolling", + } +} + +#[derive(Clone, Copy)] +enum CalendarWindow { + Day, + Week, + Month, +} + +fn calendar_bucket( + now_unix_secs: u64, + timezone: Option<&str>, + kind: CalendarWindow, + week_start: u8, +) -> Result<(String, u64, u64, &'static str), PlanUsageAdmissionError> { + let timezone = effective_timezone_name(timezone) + .parse::() + .map_err(|_| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "usage policy contains an invalid timezone".to_string(), + )) + })?; + let now = Utc + .timestamp_opt(i64::try_from(now_unix_secs).unwrap_or(i64::MAX), 0) + .single() + .ok_or_else(|| { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "usage policy timestamp is out of range".to_string(), + )) + })? + .with_timezone(&timezone); + let current_date = now.date_naive(); + let (start_date, end_date, label) = match kind { + CalendarWindow::Day => ( + current_date, + current_date.succ_opt().ok_or_else(date_overflow)?, + "calendar_day", + ), + CalendarWindow::Week => { + let current = weekday_number(current_date.weekday()); + let elapsed = (current + 7 - u32::from(week_start)) % 7; + let start = current_date - chrono::Duration::days(i64::from(elapsed)); + ( + start, + start + .checked_add_days(chrono::Days::new(7)) + .ok_or_else(date_overflow)?, + "calendar_week", + ) + } + CalendarWindow::Month => { + let start = current_date.with_day(1).ok_or_else(date_overflow)?; + let (year, month) = if start.month() == 12 { + (start.year() + 1, 1) + } else { + (start.year(), start.month() + 1) + }; + ( + start, + chrono::NaiveDate::from_ymd_opt(year, month, 1).ok_or_else(date_overflow)?, + "calendar_month", + ) + } + }; + let start = timezone + .from_local_datetime(&start_date.and_hms_opt(0, 0, 0).ok_or_else(date_overflow)?) + .earliest() + .ok_or_else(date_overflow)?; + let end = timezone + .from_local_datetime(&end_date.and_hms_opt(0, 0, 0).ok_or_else(date_overflow)?) + .latest() + .ok_or_else(date_overflow)?; + let end_unix = u64::try_from(end.timestamp()).map_err(|_| date_overflow())?; + let start_unix = u64::try_from(start.timestamp()).map_err(|_| date_overflow())?; + let window_seconds = end_unix.saturating_sub(start_unix).max(1); + let remaining = end_unix.saturating_sub(now_unix_secs).max(1); + Ok(( + start.timestamp().to_string(), + window_seconds, + remaining, + label, + )) +} + +fn calendar_window_bounds( + now_unix_secs: u64, + timezone: Option<&str>, + kind: CalendarWindow, + week_start: u8, +) -> Result<(u64, u64, u64, &'static str), GatewayError> { + let (start, window_seconds, retry_after, label) = + calendar_bucket(now_unix_secs, timezone, kind, week_start).map_err( + |error| match error { + PlanUsageAdmissionError::Gateway(error) => error, + PlanUsageAdmissionError::Runtime(error) => { + GatewayError::Internal(error.to_string()) + } + PlanUsageAdmissionError::Rejected(_) => GatewayError::Internal( + "calendar usage policy compilation unexpectedly rejected a request".to_string(), + ), + }, + )?; + let starts_at_unix_secs = start.parse::().map_err(|_| { + GatewayError::Internal("calendar usage policy produced an invalid epoch".to_string()) + })?; + Ok(( + starts_at_unix_secs, + starts_at_unix_secs.saturating_add(window_seconds), + retry_after, + label, + )) +} + +fn effective_timezone_name(explicit: Option<&str>) -> String { + explicit + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| { + std::env::var("APP_TIMEZONE") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| value.parse::().is_ok()) + }) + .unwrap_or_else(|| DEFAULT_CALENDAR_TIMEZONE.to_string()) +} + +fn weekday_number(weekday: Weekday) -> u32 { + weekday.num_days_from_monday() + 1 +} + +fn date_overflow() -> PlanUsageAdmissionError { + PlanUsageAdmissionError::Gateway(GatewayError::Internal( + "usage policy calendar window is out of range".to_string(), + )) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn entitlement(id: &str, snapshot: serde_json::Value) -> UserPlanEntitlementRecord { + UserPlanEntitlementRecord { + id: id.to_string(), + user_id: "user-1".to_string(), + plan_id: "plan-1".to_string(), + payment_order_id: "order-1".to_string(), + status: "active".to_string(), + starts_at_unix_secs: 1_000, + expires_at_unix_secs: 10_000, + entitlements_snapshot: snapshot, + created_at_unix_secs: 1_000, + updated_at_unix_secs: 1_000, + } + } + + #[test] + fn compiles_week_only_and_combined_rules() { + let policy = compile_effective_policy( + &[entitlement( + "ent-1", + json!([{ + "type": "usage_policy", + "rules": [ + {"metric":"request_count","window":{"kind":"rolling","seconds":18000},"limit":500}, + {"metric":"request_count","window":{"kind":"calendar_week","timezone":"Asia/Shanghai"},"limit":10000}, + {"metric":"concurrency","window":{"kind":"concurrent"},"limit":4} + ] + }]), + )], + 2_000, + ) + .expect("policy"); + assert_eq!(policy.request_rules.len(), 2); + assert_eq!(policy.concurrency_limit, Some(4)); + } + + #[test] + fn policy_snapshot_exists_only_for_cost_rules_and_derives_unique_tokens() { + assert!(PlanUsagePolicySnapshot::for_admission( + "user-1", + EffectivePlanUsagePolicy::default(), + 2_000, + ) + .is_none()); + + let policy = compile_effective_policy( + &[entitlement( + "ent-cost", + json!([{"type":"usage_policy","rules":[ + {"metric":"actual_cost_usd","window":{"kind":"subscription_period"},"limit":10.0} + ]}]), + )], + 2_000, + ) + .expect("cost policy"); + let snapshot = PlanUsagePolicySnapshot::for_admission("user-1", policy.clone(), 2_000) + .expect("cost rules require a policy snapshot"); + let reservation = snapshot.new_reservation_context(); + let retry_reservation = snapshot.new_reservation_context(); + + assert_eq!(reservation.subject_id(), "user-1"); + assert_eq!(reservation.admitted_at_unix_secs(), 2_000); + assert_eq!(reservation.policy(), &policy); + assert!(!reservation.token().trim().is_empty()); + assert_ne!(reservation.token(), retry_reservation.token()); + assert_eq!(retry_reservation.admitted_at_unix_secs(), 2_000); + assert!(std::ptr::eq( + reservation.policy(), + retry_reservation.policy() + )); + } + + #[test] + fn admitted_cost_snapshot_keeps_original_window_after_entitlement_expiry() { + let policy = compile_effective_policy( + &[entitlement( + "ent-cost", + json!([{"type":"usage_policy","rules":[ + {"metric":"actual_cost_usd","window":{"kind":"subscription_period"},"limit":10.0} + ]}]), + )], + 9_999, + ) + .expect("policy before expiry"); + let snapshot = PlanUsagePolicySnapshot::for_admission("user-1", policy, 9_999) + .expect("cost policy snapshot"); + + let rule = snapshot + .policy() + .cost_rules + .first() + .expect("snapshotted cost rule"); + let runtime = runtime_cost_rule(rule, snapshot.admitted_at_unix_secs) + .expect("snapshot remains compilable at its admitted timestamp"); + assert_eq!(runtime.window.starts_at_unix_secs, 1_000); + assert_eq!(runtime.window.ends_at_unix_secs, 10_000); + assert_eq!(runtime.retry_after, 1); + + assert!(compile_effective_policy( + &[entitlement( + "ent-cost", + json!([{"type":"usage_policy","rules":[ + {"metric":"actual_cost_usd","window":{"kind":"subscription_period"},"limit":10.0} + ]}]), + )], + 10_000, + ) + .expect("policy at expiry") + .cost_rules + .is_empty()); + } + + #[test] + fn same_window_uses_strictest_limit_but_subscription_periods_stay_independent() { + let first = entitlement( + "ent-a", + json!([{"type":"usage_policy","rules":[ + {"metric":"request_count","window":{"kind":"rolling","seconds":60},"limit":100}, + {"metric":"request_count","window":{"kind":"subscription_period"},"limit":1000} + ]}]), + ); + let second = entitlement( + "ent-b", + json!([{"type":"usage_policy","rules":[ + {"metric":"request_count","window":{"kind":"rolling","seconds":60},"limit":40}, + {"metric":"request_count","window":{"kind":"subscription_period"},"limit":2000} + ]}]), + ); + let policy = compile_effective_policy(&[first, second], 2_000).expect("policy"); + assert_eq!(policy.request_rules.len(), 3); + assert_eq!( + policy + .request_rules + .iter() + .find(|rule| rule.identity == "rolling:60") + .map(|rule| rule.limit), + Some(40) + ); + } + + #[test] + fn rolling_windows_use_stable_sliding_keys_and_calendar_windows_use_epoch_keys() { + let rolling = EffectiveRequestRule { + identity: "rolling:18000".to_string(), + window: UsagePolicyWindow::Rolling { seconds: 18_000 }, + limit: 10, + entitlement_period: None, + }; + let first = runtime_request_rule("user-1", &rolling, 20_000).expect("first rolling"); + let second = runtime_request_rule("user-1", &rolling, 40_000).expect("second rolling"); + assert_eq!(first.key, second.key); + assert_eq!(first.window_seconds, 18_000); + assert_eq!(first.retention_seconds, 18_000); + + let weekly = EffectiveRequestRule { + identity: "calendar_week:Asia/Shanghai:1".to_string(), + window: UsagePolicyWindow::CalendarWeek { + timezone: Some("Asia/Shanghai".to_string()), + week_start: 1, + }, + limit: 20, + entitlement_period: None, + }; + let sunday = chrono::DateTime::parse_from_rfc3339("2026-08-16T12:00:00+08:00") + .unwrap() + .timestamp() as u64; + let monday = chrono::DateTime::parse_from_rfc3339("2026-08-17T12:00:00+08:00") + .unwrap() + .timestamp() as u64; + let first = runtime_request_rule("user-1", &weekly, sunday).expect("first week"); + let second = runtime_request_rule("user-1", &weekly, monday).expect("second week"); + assert_ne!(first.key, second.key); + assert_eq!(first.window, "calendar_week"); + assert_eq!(first.retention_seconds, first.retry_after); + assert!(first.retention_seconds < first.window_seconds); + } + + #[test] + fn subscription_period_counter_keeps_the_entitlement_epoch_isolated() { + let rule = EffectiveRequestRule { + identity: "subscription:ent-1".to_string(), + window: UsagePolicyWindow::SubscriptionPeriod, + limit: 5, + entitlement_period: Some((1_000, 10_000)), + }; + let runtime = runtime_request_rule("user-1", &rule, 9_500).expect("subscription rule"); + assert!(runtime.key.ends_with(":subscription:ent-1:1000")); + assert_eq!(runtime.window_seconds, 9_000); + assert_eq!(runtime.retention_seconds, 500); + assert_eq!(runtime.retry_after, 500); + } + + #[test] + fn calendar_bucket_retention_uses_only_the_short_remaining_period() { + let daily = EffectiveRequestRule { + identity: "calendar_day:Asia/Shanghai".to_string(), + window: UsagePolicyWindow::CalendarDay { + timezone: Some("Asia/Shanghai".to_string()), + }, + limit: 5, + entitlement_period: None, + }; + let near_midnight = chrono::DateTime::parse_from_rfc3339("2026-08-18T23:59:59+08:00") + .unwrap() + .timestamp() as u64; + let runtime = runtime_request_rule("user-1", &daily, near_midnight).expect("daily rule"); + + assert_eq!(runtime.window_seconds, 24 * 60 * 60); + assert_eq!(runtime.retention_seconds, 1); + assert_eq!(runtime.retry_after, 1); + } + + #[test] + fn cost_policy_fails_closed_when_usage_runtime_is_disabled() { + assert!(ensure_cost_policy_usage_runtime_enabled(true).is_ok()); + assert!(matches!( + ensure_cost_policy_usage_runtime_enabled(false), + Err(GatewayError::Internal(message)) + if message.contains("usage runtime") + && message.contains("terminal reconciliation") + )); + } + + #[test] + fn cost_retention_uses_exclusive_rolling_influence_end() { + let rule = EffectiveCostRule { + identity: "rolling:18000".to_string(), + window: UsagePolicyWindow::Rolling { seconds: 18_000 }, + limit_cost_units: 100, + entitlement_period: None, + }; + let runtime = runtime_cost_rule(&rule, 20_000).expect("runtime cost rule"); + assert_eq!(runtime.influence_ends_at_unix_secs, 38_001); + } + + #[test] + fn explicit_cost_release_uses_server_token_and_zero_actual_cost() { + let input = build_usage_policy_cost_release_input( + "request-1", + "user-1", + "server-reservation-token", + 12_345, + ); + + assert_eq!(input.request_id, "request-1"); + assert_eq!(input.subject_id, "user-1"); + assert_eq!(input.reservation_token, "server-reservation-token"); + assert_eq!(input.actual_cost_units, 0); + assert_eq!( + input.terminal_state, + UsagePolicyCostReservationState::Released + ); + assert_eq!(input.finalized_at_unix_secs, 12_345); + input.validate().expect("release input should validate"); + } +} diff --git a/apps/aether-gateway/src/privacy/mod.rs b/apps/aether-gateway/src/privacy/mod.rs index 3d66c47ec..2809ceec4 100644 --- a/apps/aether-gateway/src/privacy/mod.rs +++ b/apps/aether-gateway/src/privacy/mod.rs @@ -1,6 +1,7 @@ use std::collections::{BTreeMap, HashMap, HashSet}; use std::convert::Infallible; use std::fmt; +use std::io::Write; use std::net::{Ipv4Addr, Ipv6Addr}; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, LazyLock, Mutex}; @@ -26,6 +27,10 @@ const MAX_SENTINEL_NAMESPACE_LEN: usize = 32; const DIRECT_RESTORE_SENTINEL_LIMIT: usize = 32; const MAX_CACHE_SENTINEL_BYTES: usize = 128; const MAX_CACHE_RECORD_BYTES: usize = 512; +// A configured response buffer may be explicitly disabled, so restoration +// needs an independent ceiling for bytes introduced by sentinel replacement. +// This is not a response-body limit: an ordinary large response remains valid. +pub(crate) const MAX_SYNC_RESTORE_EXPANSION_BYTES: usize = 64 * 1024 * 1024; const CHAT_PII_REDACTION_RUNTIME_CONFIG_CACHE_TTL: Duration = Duration::from_secs(5); static EMAIL_REGEX: LazyLock = LazyLock::new(|| { @@ -495,11 +500,16 @@ impl RedactionSession { self.mappings.values() } - fn restore_text(&self, input: &str) -> RestoredText { - if self.mapping_count() <= DIRECT_RESTORE_SENTINEL_LIMIT { - return restore_text_direct_longest_first(input, self); - } - restore_text_with_matcher(input, &SentinelMatcher::new(self)) + fn restore_text(&self, input: &str) -> Result { + let matcher = SentinelMatcher::new(self); + let mut budget = RestoreExpansionBudget::new(MAX_SYNC_RESTORE_EXPANSION_BYTES); + let (bytes, restored) = restore_bytes_with_budget(input.as_bytes(), &matcher, &mut budget)?; + let text = String::from_utf8(bytes).map_err(|_| { + GatewayError::Internal( + "redaction response restoration produced invalid UTF-8".to_string(), + ) + })?; + Ok(RestoredText { text, restored }) } fn sentinel_for_candidate(&mut self, source_text: &str, candidate: &Candidate) -> String { @@ -2589,10 +2599,358 @@ struct RestoredText { restored: bool, } +pub(crate) struct RestoreExpansionBudget { + limit: usize, + used: usize, +} + +impl RestoreExpansionBudget { + /// The caller creates a fresh budget for one sync response, stream + /// `push_chunk`/`finish`, SSE event, or WebSocket frame. It intentionally + /// never accumulates across a long-lived stream. + pub(crate) fn new(limit: usize) -> Self { + Self { limit, used: 0 } + } + + fn charge(&mut self, expansion: usize) -> Result<(), GatewayError> { + let next = self.used.checked_add(expansion).ok_or_else(|| { + GatewayError::Internal("redaction response restoration expansion overflow".to_string()) + })?; + if next > self.limit { + return Err(GatewayError::Internal(format!( + "redaction response restoration expansion exceeds {}/{} bytes", + next, self.limit + ))); + } + self.used = next; + Ok(()) + } + + fn checkpoint(&self) -> usize { + self.used + } + + fn available_from(&self, checkpoint: usize) -> usize { + self.limit.saturating_sub(checkpoint) + } + + fn replace_charge(&mut self, checkpoint: usize, expansion: usize) -> Result<(), GatewayError> { + self.used = checkpoint.min(self.used); + self.charge(expansion) + } +} + +struct RestoreScan { + consumed: usize, + output_len: usize, + expansion: usize, + restored: bool, +} + +fn scan_restore_bytes( + input: &[u8], + matcher: &SentinelMatcher<'_>, +) -> Result { + scan_restore_prefix(input, input.len(), matcher) +} + +fn scan_restore_prefix( + input: &[u8], + scan_limit: usize, + matcher: &SentinelMatcher<'_>, +) -> Result { + let scan_limit = scan_limit.min(input.len()); + let mut output_len = 0usize; + let mut index = 0usize; + let mut restored = false; + while index < scan_limit { + if let Some(mapping) = matcher.matching_mapping_at(input, index) { + let next_index = index + .checked_add(mapping.sentinel.len()) + .filter(|next| *next > index && *next <= input.len()) + .ok_or_else(|| { + GatewayError::Internal( + "redaction response restoration sentinel length overflow".to_string(), + ) + })?; + output_len = output_len + .checked_add(mapping.original.len()) + .ok_or_else(|| { + GatewayError::Internal( + "redaction response restoration output length overflow".to_string(), + ) + })?; + index = next_index; + restored = true; + } else { + output_len = output_len.checked_add(1).ok_or_else(|| { + GatewayError::Internal( + "redaction response restoration output length overflow".to_string(), + ) + })?; + index += 1; + } + } + Ok(RestoreScan { + consumed: index, + output_len, + expansion: output_len.saturating_sub(index), + restored, + }) +} + +fn build_restored_bytes( + input: &[u8], + matcher: &SentinelMatcher<'_>, + scan: &RestoreScan, +) -> Vec { + let mut output = Vec::with_capacity(scan.output_len); + let mut index = 0usize; + while index < scan.consumed { + if let Some(mapping) = matcher.matching_mapping_at(input, index) { + output.extend_from_slice(mapping.original.as_bytes()); + index += mapping.sentinel.len(); + } else { + output.push(input[index]); + index += 1; + } + } + output +} + +fn restore_bytes_with_budget( + input: &[u8], + matcher: &SentinelMatcher<'_>, + budget: &mut RestoreExpansionBudget, +) -> Result<(Vec, bool), GatewayError> { + let scan = scan_restore_bytes(input, matcher)?; + if !scan.restored { + return Ok((input.to_vec(), false)); + } + // Charge before allocating the replacement buffer. A provider can repeat + // a short sentinel many times, so checking after `Vec` construction is too + // late to protect the process from an amplification attack. + budget.charge(scan.expansion)?; + Ok((build_restored_bytes(input, matcher, &scan), true)) +} + +fn measure_json_restore( + value: &Value, + matcher: &SentinelMatcher<'_>, +) -> Result<(usize, bool), GatewayError> { + match value { + Value::String(text) => { + let scan = scan_restore_bytes(text.as_bytes(), matcher)?; + Ok((scan.expansion, scan.restored)) + } + Value::Array(values) => { + let mut expansion = 0usize; + let mut restored = false; + for value in values { + let (value_expansion, value_restored) = measure_json_restore(value, matcher)?; + expansion = expansion.checked_add(value_expansion).ok_or_else(|| { + GatewayError::Internal( + "redaction response restoration expansion overflow".to_string(), + ) + })?; + restored |= value_restored; + } + Ok((expansion, restored)) + } + Value::Object(values) => { + let mut expansion = 0usize; + let mut restored = false; + for value in values.values() { + let (value_expansion, value_restored) = measure_json_restore(value, matcher)?; + expansion = expansion.checked_add(value_expansion).ok_or_else(|| { + GatewayError::Internal( + "redaction response restoration expansion overflow".to_string(), + ) + })?; + restored |= value_restored; + } + Ok((expansion, restored)) + } + _ => Ok((0, false)), + } +} + +fn restore_json_strings_inner( + value: &mut Value, + matcher: &SentinelMatcher<'_>, + preserve_opaque_reasoning_state: bool, +) -> Result { + match value { + Value::String(text) => { + let scan = scan_restore_bytes(text.as_bytes(), matcher)?; + if !scan.restored { + return Ok(false); + } + let bytes = build_restored_bytes(text.as_bytes(), matcher, &scan); + *text = String::from_utf8(bytes).map_err(|_| { + GatewayError::Internal( + "redaction response restoration produced invalid UTF-8".to_string(), + ) + })?; + Ok(true) + } + Value::Array(values) => { + let mut restored = false; + for value in values { + restored |= + restore_json_strings_inner(value, matcher, preserve_opaque_reasoning_state)?; + } + Ok(restored) + } + Value::Object(values) => { + if preserve_opaque_reasoning_state + && openai_responses_object_is_opaque_reasoning_state(values) + { + return Ok(false); + } + let mut restored = false; + for value in values.values_mut() { + restored |= + restore_json_strings_inner(value, matcher, preserve_opaque_reasoning_state)?; + } + Ok(restored) + } + _ => Ok(false), + } +} + +pub(crate) fn restore_json_strings_with_budget( + value: &mut Value, + session: &RedactionSession, + budget: &mut RestoreExpansionBudget, +) -> Result { + let matcher = SentinelMatcher::new(session); + let (expansion, restored) = measure_json_restore(value, &matcher)?; + if !restored { + return Ok(false); + } + budget.charge(expansion)?; + restore_json_strings_inner( + value, + &matcher, + session.preserves_deepseek_opaque_reasoning_state(), + ) +} + +/// Compatibility wrapper for crate-local callers that only need the historical +/// boolean result. The bounded implementation remains the source of truth; +/// failures are treated as "not restored" so callers cannot accidentally emit +/// a partially restored value. +pub(crate) fn restore_json_strings(value: &mut Value, session: &RedactionSession) -> bool { + let mut budget = RestoreExpansionBudget::new(MAX_SYNC_RESTORE_EXPANSION_BYTES); + restore_json_strings_with_budget(value, session, &mut budget).unwrap_or(false) +} + +fn openai_responses_object_is_opaque_reasoning_state( + object: &serde_json::Map, +) -> bool { + let Some(item_type) = object.get("type").and_then(Value::as_str) else { + return false; + }; + // A provider-owned reasoning item is replayed byte-for-byte. The + // encrypted continuation token is only considered opaque when it is + // paired with the reasoning text item shape; ordinary reasoning text must + // still participate in PII masking/restoration. + let id_is_absent_or_empty = object + .get("id") + .is_none_or(|id| id.is_null() || id.as_str().is_some_and(|value| value.trim().is_empty())); + item_type == "reasoning" + && id_is_absent_or_empty + && object + .get("encrypted_content") + .and_then(Value::as_str) + .is_some_and(|state| !state.trim().is_empty()) + && object + .get("content") + .and_then(Value::as_array) + .is_some_and(|content| { + content.iter().any(|part| { + part.get("type").and_then(Value::as_str) == Some("reasoning_text") + && part.get("text").is_some_and(Value::is_string) + }) + }) +} + +fn effective_sync_restore_expansion_limit(configured_limit: u64) -> usize { + usize::try_from(configured_limit) + .unwrap_or(usize::MAX) + .min(MAX_SYNC_RESTORE_EXPANSION_BYTES) +} + +fn sync_restore_expansion_limit() -> usize { + effective_sync_restore_expansion_limit(crate::headers::max_redacted_sync_response_body_bytes()) +} + +struct LimitedJsonWriter { + bytes: Vec, + limit: usize, + exceeded: bool, +} + +impl LimitedJsonWriter { + fn new(limit: usize) -> Self { + Self { + bytes: Vec::with_capacity(limit.min(16 * 1024)), + limit, + exceeded: false, + } + } +} + +impl Write for LimitedJsonWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.limit.saturating_sub(self.bytes.len()) { + self.exceeded = true; + return Err(std::io::Error::new( + std::io::ErrorKind::WriteZero, + "redaction response exceeds output limit", + )); + } + self.bytes.extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +pub(crate) fn serialize_json_value_with_limit( + value: &Value, + limit: usize, +) -> Result { + let mut writer = LimitedJsonWriter::new(limit); + match serde_json::to_writer(&mut writer, value) { + Ok(()) => String::from_utf8(writer.bytes).map_err(|_| { + GatewayError::Internal( + "redaction response serialization produced invalid UTF-8".to_string(), + ) + }), + Err(_) if writer.exceeded => Err(GatewayError::Internal(format!( + "redaction response serialized body exceeds {limit} bytes" + ))), + Err(error) => Err(GatewayError::Internal(error.to_string())), + } +} + pub(crate) fn restore_sync_response_body( headers: &mut BTreeMap, body: &[u8], session: &RedactionSession, +) -> Result { + restore_sync_response_body_with_limit(headers, body, session, sync_restore_expansion_limit()) +} + +fn restore_sync_response_body_with_limit( + headers: &mut BTreeMap, + body: &[u8], + session: &RedactionSession, + expansion_limit: usize, ) -> Result { if session.mapping_count() == 0 { return Ok(RestoredSyncResponseBody { @@ -2604,9 +2962,9 @@ pub(crate) fn restore_sync_response_body( ensure_identity_response_encoding(headers)?; let restored = if response_body_is_json(headers, body) { - restore_json_response_body(body, session)? + restore_json_response_body(body, session, expansion_limit)? } else { - restore_text_response_body(body, session) + restore_text_response_body(body, session, expansion_limit)? }; if restored.restored { set_content_length(headers, restored.body.len()); @@ -2654,17 +3012,22 @@ fn content_type_is_json(content_type: &str) -> bool { fn restore_json_response_body( body: &[u8], session: &RedactionSession, + expansion_limit: usize, ) -> Result { let Ok(mut value) = serde_json::from_slice::(body) else { - return Ok(restore_text_response_body(body, session)); + return restore_text_response_body(body, session, expansion_limit); }; - if !restore_json_strings(&mut value, session) { + let mut budget = RestoreExpansionBudget::new(expansion_limit); + if !restore_json_strings_with_budget(&mut value, session, &mut budget)? { return Ok(RestoredSyncResponseBody { body: body.to_vec(), restored: false, }); } - let body = serde_json::to_vec(&value).map_err(|err| GatewayError::Internal(err.to_string()))?; + let output_limit = body.len().checked_add(expansion_limit).ok_or_else(|| { + GatewayError::Internal("redaction response output length overflow".to_string()) + })?; + let body = serialize_json_value_with_limit(&value, output_limit)?.into_bytes(); Ok(RestoredSyncResponseBody { body, restored: true, @@ -2678,89 +3041,21 @@ fn restore_json_response_body( /// 两边因此保持同一套还原语义:未映射的占位符原样保留,`type` / `model` / `id` /// 这类协议字段虽然也被遍历,但它们不可能包含本 session 派生出的 sentinel, /// 所以不会被改写。 -/// -/// Provider-owned Responses reasoning state is the exception. DeepSeek-style -/// items bind `reasoning_text` to `encrypted_content` and require both to be -/// replayed unchanged. Restoring a sentinel in only the text half would make -/// the client return a different continuation item, so those objects/events -/// stay opaque even when another response field is restored. -pub(crate) fn restore_json_strings(value: &mut Value, session: &RedactionSession) -> bool { - match value { - Value::String(text) => { - let restored = session.restore_text(text); - if !restored.restored { - return false; - } - *text = restored.text; - true - } - Value::Array(values) => { - let mut restored = false; - for value in values { - restored = restore_json_strings(value, session) || restored; - } - restored - } - Value::Object(values) => { - if session.preserves_deepseek_opaque_reasoning_state() - && openai_responses_object_is_opaque_reasoning_state(values) - { - return false; - } - let mut restored = false; - for value in values.values_mut() { - restored = restore_json_strings(value, session) || restored; - } - restored - } - _ => false, - } -} - -fn openai_responses_object_is_opaque_reasoning_state( - object: &serde_json::Map, -) -> bool { - let Some(item_type) = object.get("type").and_then(Value::as_str) else { - return false; - }; - // Never treat a `reasoning_text` content part or delta/done event name by - // itself as proof of opaque state: ordinary OpenAI reasoning text may - // contain mask placeholders that still need client-side restoration. The - // binding evidence is the parent reasoning item carrying provider-owned - // encrypted continuation state. Returning here keeps that whole object, - // including its nested reasoning_text parts, value-for-value unchanged. - let id_is_absent_or_empty = object - .get("id") - .is_none_or(|id| id.is_null() || id.as_str().is_some_and(|value| value.trim().is_empty())); - item_type == "reasoning" - && id_is_absent_or_empty - && object - .get("encrypted_content") - .and_then(Value::as_str) - .is_some_and(|state| !state.trim().is_empty()) - && object - .get("content") - .and_then(Value::as_array) - .is_some_and(|content| { - content.iter().any(|part| { - part.get("type").and_then(Value::as_str) == Some("reasoning_text") - && part.get("text").is_some_and(Value::is_string) - }) - }) -} - -fn restore_text_response_body(body: &[u8], session: &RedactionSession) -> RestoredSyncResponseBody { +fn restore_text_response_body( + body: &[u8], + session: &RedactionSession, + expansion_limit: usize, +) -> Result { let Ok(text) = std::str::from_utf8(body) else { - return RestoredSyncResponseBody { + return Ok(RestoredSyncResponseBody { body: body.to_vec(), restored: false, - }; + }); }; - let restored = session.restore_text(text); - RestoredSyncResponseBody { - body: restored.text.into_bytes(), - restored: restored.restored, - } + let matcher = SentinelMatcher::new(session); + let mut budget = RestoreExpansionBudget::new(expansion_limit); + let (body, restored) = restore_bytes_with_budget(text.as_bytes(), &matcher, &mut budget)?; + Ok(RestoredSyncResponseBody { body, restored }) } #[derive(Clone, Copy, Eq, PartialEq)] @@ -2777,8 +3072,22 @@ pub(crate) struct StreamingResponseRestorer<'a> { text_carry: Vec, sse_line_buffer: Vec, sse_json_event_lines: Vec>, + sse_json_event_bytes: usize, + max_expansion_bytes: usize, } +// A provider-controlled SSE line/event must not be allowed to grow forever +// while the stream restorer waits for a newline or blank-record separator. +// This bounds parser carry state only; complete response bodies and ordinary +// streaming chunks are still passed through without a body/concurrency cap. +const MAX_STREAM_RESTORE_BUFFER_BYTES: usize = 16 * 1024 * 1024; +// This is an expansion budget, not a body or duration limit. A normal +// pass-through chunk can be any size; only bytes introduced by replacing a +// sentinel with its original value consume this budget. +pub(crate) const MAX_STREAM_RESTORE_EXPANSION_BYTES: usize = 16 * 1024 * 1024; +pub(crate) const MAX_STREAM_RESTORE_OUTPUT_BYTES: usize = + (16usize * 1024 * 1024).saturating_add(MAX_STREAM_RESTORE_EXPANSION_BYTES); + impl<'a> StreamingResponseRestorer<'a> { pub(crate) fn new( headers: &BTreeMap, @@ -2805,6 +3114,26 @@ impl<'a> StreamingResponseRestorer<'a> { Self::with_mode(session, StreamRestoreMode::Text) } + #[cfg(test)] + fn for_text_with_expansion_limit( + session: &'a RedactionSession, + max_expansion_bytes: usize, + ) -> Self { + let mut restorer = Self::with_mode(session, StreamRestoreMode::Text); + restorer.max_expansion_bytes = max_expansion_bytes; + restorer + } + + #[cfg(test)] + fn for_sse_with_expansion_limit( + session: &'a RedactionSession, + max_expansion_bytes: usize, + ) -> Self { + let mut restorer = Self::with_mode(session, StreamRestoreMode::Sse); + restorer.max_expansion_bytes = max_expansion_bytes; + restorer + } + fn with_mode(session: &'a RedactionSession, mode: StreamRestoreMode) -> Self { let max_sentinel_len = session .mappings() @@ -2819,6 +3148,8 @@ impl<'a> StreamingResponseRestorer<'a> { text_carry: Vec::new(), sse_line_buffer: Vec::new(), sse_json_event_lines: Vec::new(), + sse_json_event_bytes: 0, + max_expansion_bytes: MAX_STREAM_RESTORE_EXPANSION_BYTES, } } @@ -2826,9 +3157,12 @@ impl<'a> StreamingResponseRestorer<'a> { if self.max_sentinel_len == 0 { return Ok(chunk.to_vec()); } + // Per-call expansion protection only: do not turn this into a total + // stream byte or duration limit. + let mut budget = RestoreExpansionBudget::new(self.max_expansion_bytes); match self.mode { - StreamRestoreMode::Sse => self.push_sse_chunk(chunk), - StreamRestoreMode::Text => Ok(self.push_text_chunk(chunk)), + StreamRestoreMode::Sse => self.push_sse_chunk(chunk, &mut budget), + StreamRestoreMode::Text => self.push_text_chunk(chunk, &mut budget), } } @@ -2836,9 +3170,11 @@ impl<'a> StreamingResponseRestorer<'a> { if self.max_sentinel_len == 0 { return Ok(Vec::new()); } + // `finish` is a separate bounded flush, not a cumulative stream cap. + let mut budget = RestoreExpansionBudget::new(self.max_expansion_bytes); match self.mode { - StreamRestoreMode::Sse => self.finish_sse(), - StreamRestoreMode::Text => Ok(self.flush_text_carry()), + StreamRestoreMode::Sse => self.finish_sse(&mut budget), + StreamRestoreMode::Text => self.flush_text_carry(&mut budget), } } @@ -2852,16 +3188,27 @@ impl<'a> StreamingResponseRestorer<'a> { self.text_carry.len() } - fn push_text_chunk(&mut self, chunk: &[u8]) -> Vec { + fn push_text_chunk( + &mut self, + chunk: &[u8], + budget: &mut RestoreExpansionBudget, + ) -> Result, GatewayError> { self.text_carry.extend_from_slice(chunk); - self.restore_available_text(false) + self.restore_available_text(false, budget) } - fn flush_text_carry(&mut self) -> Vec { - self.restore_available_text(true) + fn flush_text_carry( + &mut self, + budget: &mut RestoreExpansionBudget, + ) -> Result, GatewayError> { + self.restore_available_text(true, budget) } - fn restore_available_text(&mut self, flush: bool) -> Vec { + fn restore_available_text( + &mut self, + flush: bool, + budget: &mut RestoreExpansionBudget, + ) -> Result, GatewayError> { let scan_limit = if flush { self.text_carry.len() } else { @@ -2869,48 +3216,56 @@ impl<'a> StreamingResponseRestorer<'a> { .len() .saturating_sub(self.max_sentinel_len.saturating_sub(1)) }; - let mut output = Vec::with_capacity(scan_limit); - let mut index = 0; - while index < scan_limit { - if let Some(mapping) = self.matcher.matching_mapping_at(&self.text_carry, index) { - output.extend_from_slice(mapping.original.as_bytes()); - index += mapping.sentinel.len(); - } else { - output.push(self.text_carry[index]); - index += 1; - } - } - self.text_carry.drain(..index); - output + let scan = scan_restore_prefix(&self.text_carry, scan_limit, &self.matcher)?; + budget.charge(scan.expansion)?; + let output = build_restored_bytes(&self.text_carry, &self.matcher, &scan); + self.text_carry.drain(..scan.consumed); + Ok(output) } - fn push_sse_chunk(&mut self, chunk: &[u8]) -> Result, GatewayError> { - self.sse_line_buffer.extend_from_slice(chunk); + fn push_sse_chunk( + &mut self, + chunk: &[u8], + budget: &mut RestoreExpansionBudget, + ) -> Result, GatewayError> { let mut output = Vec::new(); - while let Some(line_end) = self.sse_line_buffer.iter().position(|byte| *byte == b'\n') { - let line = self.sse_line_buffer.drain(..=line_end).collect::>(); - self.push_sse_line(line, &mut output)?; + let mut remaining = chunk; + while !remaining.is_empty() { + let Some(line_end) = remaining.iter().position(|byte| *byte == b'\n') else { + append_bounded_stream_restore_bytes(&mut self.sse_line_buffer, remaining)?; + break; + }; + let line_len = line_end + 1; + append_bounded_stream_restore_bytes(&mut self.sse_line_buffer, &remaining[..line_len])?; + remaining = &remaining[line_len..]; + let line = std::mem::take(&mut self.sse_line_buffer); + self.push_sse_line(line, &mut output, budget)?; } Ok(output) } - fn finish_sse(&mut self) -> Result, GatewayError> { + fn finish_sse(&mut self, budget: &mut RestoreExpansionBudget) -> Result, GatewayError> { let mut output = Vec::new(); if !self.sse_line_buffer.is_empty() { let line = std::mem::take(&mut self.sse_line_buffer); - self.push_sse_line(line, &mut output)?; + self.push_sse_line(line, &mut output, budget)?; } - self.flush_sse_json_event(&mut output)?; + self.flush_sse_json_event(&mut output, budget)?; Ok(output) } - fn push_sse_line(&mut self, line: Vec, output: &mut Vec) -> Result<(), GatewayError> { + fn push_sse_line( + &mut self, + line: Vec, + output: &mut Vec, + budget: &mut RestoreExpansionBudget, + ) -> Result<(), GatewayError> { if !self.sse_json_event_lines.is_empty() { if sse_line_is_blank(&line) { - self.flush_sse_json_event(output)?; + self.flush_sse_json_event(output, budget)?; output.extend_from_slice(&line); } else { - self.sse_json_event_lines.push(line); + self.push_sse_json_event_line(line)?; } return Ok(()); } @@ -2929,19 +3284,25 @@ impl<'a> StreamingResponseRestorer<'a> { return Ok(()); } if sse_data_value_may_be_json(value) { - self.sse_json_event_lines.push(line); + self.push_sse_json_event_line(line)?; return Ok(()); } - output.extend(self.restore_sse_text_data_line(&line)); + output.extend(self.restore_sse_text_data_line(&line, budget)?); Ok(()) } - fn flush_sse_json_event(&mut self, output: &mut Vec) -> Result<(), GatewayError> { + fn flush_sse_json_event( + &mut self, + output: &mut Vec, + budget: &mut RestoreExpansionBudget, + ) -> Result<(), GatewayError> { if self.sse_json_event_lines.is_empty() { return Ok(()); } let event_lines = std::mem::take(&mut self.sse_json_event_lines); + self.sse_json_event_bytes = 0; + let event_budget_checkpoint = budget.checkpoint(); let Some(payload) = sse_event_data_payload(&event_lines) else { output.extend(event_lines.into_iter().flatten()); return Ok(()); @@ -2954,7 +3315,7 @@ impl<'a> StreamingResponseRestorer<'a> { let Ok(mut value) = serde_json::from_str::(&payload) else { for line in event_lines { if sse_data_line_value(&line).is_some() { - output.extend(self.restore_sse_text_data_line(&line)); + output.extend(self.restore_sse_text_data_line(&line, budget)?); } else { output.extend_from_slice(&line); } @@ -2962,44 +3323,127 @@ impl<'a> StreamingResponseRestorer<'a> { return Ok(()); }; - if !restore_json_strings(&mut value, self.session) { + if !restore_json_strings_with_budget(&mut value, self.session, budget)? { output.extend(event_lines.into_iter().flatten()); return Ok(()); } + let event_input_bytes = event_lines + .iter() + .map(Vec::len) + .try_fold(0usize, |sum, len| { + sum.checked_add(len).ok_or_else(|| { + GatewayError::Internal("stream redaction SSE event length overflow".to_string()) + }) + })?; + let event_output_limit = event_input_bytes + .checked_add(budget.available_from(event_budget_checkpoint)) + .ok_or_else(|| { + GatewayError::Internal( + "stream redaction SSE restored output length overflow".to_string(), + ) + })?; let restored_payload = - serde_json::to_string(&value).map_err(|err| GatewayError::Internal(err.to_string()))?; + serialize_json_value_with_limit(&value, event_output_limit)?.into_bytes(); let mut wrote_data_line = false; + let mut output_len = 0usize; for line in event_lines { if sse_data_line_value(&line).is_some() { if !wrote_data_line { - output.extend(replace_sse_data_line_value( - &line, - restored_payload.as_bytes(), - )); + let replaced = replace_sse_data_line_value(&line, &restored_payload); + output_len = output_len.checked_add(replaced.len()).ok_or_else(|| { + GatewayError::Internal( + "stream redaction SSE restored output length overflow".to_string(), + ) + })?; + if output_len > event_output_limit { + return Err(GatewayError::Internal(format!( + "stream redaction SSE restored event exceeds {event_output_limit} bytes" + ))); + } + output.extend(replaced); wrote_data_line = true; } } else { + output_len = output_len.checked_add(line.len()).ok_or_else(|| { + GatewayError::Internal( + "stream redaction SSE restored output length overflow".to_string(), + ) + })?; + if output_len > event_output_limit { + return Err(GatewayError::Internal(format!( + "stream redaction SSE restored event exceeds {event_output_limit} bytes" + ))); + } output.extend_from_slice(&line); } } + let actual_expansion = output_len.saturating_sub(event_input_bytes); + budget.replace_charge(event_budget_checkpoint, actual_expansion)?; Ok(()) } - fn restore_sse_text_data_line(&self, line: &[u8]) -> Vec { + fn push_sse_json_event_line(&mut self, line: Vec) -> Result<(), GatewayError> { + let next_len = self + .sse_json_event_bytes + .checked_add(line.len()) + .ok_or_else(|| { + GatewayError::Internal("stream redaction SSE event length overflow".to_string()) + })?; + if next_len > MAX_STREAM_RESTORE_BUFFER_BYTES { + return Err(GatewayError::Internal(format!( + "stream redaction SSE event exceeds {MAX_STREAM_RESTORE_BUFFER_BYTES} bytes" + ))); + } + self.sse_json_event_lines.push(line); + self.sse_json_event_bytes = next_len; + Ok(()) + } + + fn restore_sse_text_data_line( + &self, + line: &[u8], + budget: &mut RestoreExpansionBudget, + ) -> Result, GatewayError> { let Some(range) = sse_data_line_value_range(line) else { - return line.to_vec(); + return Ok(line.to_vec()); }; let (body, ending) = split_sse_line_ending(line); - let restored = restore_known_sentinels_in_bytes(&body[range.clone()], self.session); - let mut output = Vec::with_capacity(line.len()); + let (restored, _) = restore_bytes_with_budget(&body[range.clone()], &self.matcher, budget)?; + let output_len = body + .len() + .saturating_sub(range.len()) + .checked_add(restored.len()) + .and_then(|len| len.checked_add(ending.len())) + .ok_or_else(|| { + GatewayError::Internal( + "stream redaction SSE restored output length overflow".to_string(), + ) + })?; + let mut output = Vec::with_capacity(output_len); output.extend_from_slice(&body[..range.start]); output.extend(restored); output.extend_from_slice(ending); - output + Ok(output) } } +fn append_bounded_stream_restore_bytes( + buffered: &mut Vec, + chunk: &[u8], +) -> Result<(), GatewayError> { + let next_len = buffered.len().checked_add(chunk.len()).ok_or_else(|| { + GatewayError::Internal("stream redaction buffer length overflow".to_string()) + })?; + if next_len > MAX_STREAM_RESTORE_BUFFER_BYTES { + return Err(GatewayError::Internal(format!( + "stream redaction buffer exceeds {MAX_STREAM_RESTORE_BUFFER_BYTES} bytes" + ))); + } + buffered.extend_from_slice(chunk); + Ok(()) +} + enum SentinelMatcher<'a> { Direct(Vec<&'a RedactionMapping>), Trie(SentinelTrie<'a>), @@ -3086,55 +3530,6 @@ impl<'a> SentinelTrie<'a> { } } -fn restore_text_direct_longest_first(input: &str, session: &RedactionSession) -> RestoredText { - let mut mappings = session.mappings().collect::>(); - mappings.sort_by_key(|mapping| std::cmp::Reverse(mapping.sentinel.len())); - let mut text = input.to_string(); - let mut restored = false; - for mapping in mappings { - if text.contains(&mapping.sentinel) { - text = text.replace(&mapping.sentinel, &mapping.original); - restored = true; - } - } - RestoredText { text, restored } -} - -fn restore_text_with_matcher(input: &str, matcher: &SentinelMatcher<'_>) -> RestoredText { - let input_bytes = input.as_bytes(); - let mut output = Vec::with_capacity(input_bytes.len()); - let mut restored = false; - let mut index = 0; - while index < input_bytes.len() { - if let Some(mapping) = matcher.matching_mapping_at(input_bytes, index) { - output.extend_from_slice(mapping.original.as_bytes()); - index += mapping.sentinel.len(); - restored = true; - } else { - output.push(input_bytes[index]); - index += 1; - } - } - let text = String::from_utf8(output).unwrap_or_else(|_| input.to_string()); - RestoredText { text, restored } -} - -fn restore_known_sentinels_in_bytes(input: &[u8], session: &RedactionSession) -> Vec { - let matcher = SentinelMatcher::new(session); - let mut output = Vec::with_capacity(input.len()); - let mut index = 0; - while index < input.len() { - if let Some(mapping) = matcher.matching_mapping_at(input, index) { - output.extend_from_slice(mapping.original.as_bytes()); - index += mapping.sentinel.len(); - } else { - output.push(input[index]); - index += 1; - } - } - output -} - fn content_type_is_sse(content_type: &str) -> bool { content_type .split(';') @@ -4572,7 +4967,10 @@ mod tests { assert!(redacted.text.contains(sentinel)); assert!(!redacted.text.contains(", status_update: SchedulerRequestCandidateStatusUpdate, ) { - let Some(record) = + let Some(mut record) = build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput { plan, report_context, @@ -372,6 +372,8 @@ pub(crate) async fn record_local_request_candidate_status( else { return; }; + record.skip_reason = + local_request_candidate_skip_reason(record.status, record.error_type.as_deref()); persist_local_request_candidate_status_record(state, record).await; } @@ -429,6 +431,7 @@ fn build_local_request_candidate_status_snapshot_record( started_at_unix_ms, finished_at_unix_ms, } = status_update; + let skip_reason = local_request_candidate_skip_reason(status, error_type.as_deref()); UpsertRequestCandidateRecord { id: snapshot.candidate_id.clone(), request_id: snapshot.request_id.clone(), @@ -442,7 +445,7 @@ fn build_local_request_candidate_status_snapshot_record( endpoint_id: Some(snapshot.endpoint_id.clone()), key_id: Some(snapshot.key_id.clone()), status, - skip_reason: None, + skip_reason, is_cached: None, status_code, error_type, @@ -457,6 +460,17 @@ fn build_local_request_candidate_status_snapshot_record( } } +fn local_request_candidate_skip_reason( + status: RequestCandidateStatus, + error_type: Option<&str>, +) -> Option { + (status == RequestCandidateStatus::Skipped) + .then_some(error_type) + .flatten() + .filter(|reason| *reason == "provider_key_concurrency_limit_reached") + .map(ToOwned::to_owned) +} + pub(crate) fn try_enqueue_local_request_candidate_status_snapshot( state: &(impl RequestCandidateRuntimeWriter + ?Sized), snapshot: &LocalRequestCandidateStatusSnapshot, @@ -1078,6 +1092,40 @@ mod tests { assert_eq!(records[0].status_code, Some(200)); } + #[test] + fn saturated_provider_key_snapshot_persists_capacity_skip_reason() { + let mut plan = sample_plan(); + plan.candidate_id = Some("candidate-provider-key-saturated".to_string()); + let snapshot = snapshot_local_request_candidate_status(&plan, None) + .expect("candidate snapshot should build"); + let writer = SynchronousStatusWriter::default(); + + try_enqueue_local_request_candidate_status_snapshot( + &writer, + &snapshot, + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Skipped, + status_code: Some(429), + error_type: Some("provider_key_concurrency_limit_reached".to_string()), + error_message: Some("provider key concurrency limit reached: 1".to_string()), + latency_ms: Some(0), + started_at_unix_ms: Some(123), + finished_at_unix_ms: Some(123), + }, + ) + .expect("saturated status should use the synchronous enqueue path"); + + let records = writer + .records + .lock() + .expect("synchronous status records lock"); + assert_eq!(records.len(), 1); + assert_eq!( + records[0].skip_reason.as_deref(), + Some("provider_key_concurrency_limit_reached") + ); + } + fn sample_minimal_candidate() -> SchedulerMinimalCandidateSelectionCandidate { SchedulerMinimalCandidateSelectionCandidate { provider_id: "provider-1".to_string(), @@ -1280,7 +1328,10 @@ mod tests { assert_eq!(stored[0].status, RequestCandidateStatus::Success); assert_eq!(stored[0].status_code, Some(200)); assert_eq!(stored[0].latency_ms, Some(25)); - assert_eq!(stored[0].started_at_unix_ms, Some(101)); + // Report updates preserve the original attempt start timestamp from + // the persisted slot; the update's value is only used when no start + // timestamp has been recorded yet. + assert_eq!(stored[0].started_at_unix_ms, Some(100_000)); assert_eq!(stored[0].finished_at_unix_ms, Some(102)); } diff --git a/apps/aether-gateway/src/router.rs b/apps/aether-gateway/src/router.rs index b3486af81..3c78123db 100644 --- a/apps/aether-gateway/src/router.rs +++ b/apps/aether-gateway/src/router.rs @@ -1,12 +1,20 @@ use std::path::PathBuf; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; +use std::time::Duration; +use axum::body::Body; use axum::extract::Request; use axum::http::header::{CACHE_CONTROL, EXPIRES, PRAGMA}; use axum::http::{HeaderValue, Method}; use axum::response::{IntoResponse, Response}; use axum::routing::any; use axum::Router; -use tower::ServiceExt; +use hyper::body::Incoming; +use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; +use hyper_util::server::conn::auto::Builder as HyperServerBuilder; +use hyper_util::service::TowerToHyperService; +use tower::{Service as _, ServiceExt}; use tower_http::services::{ServeDir, ServeFile}; use tracing::warn; @@ -15,6 +23,98 @@ use aether_runtime_state::RuntimeSemaphoreError; use super::{api, handlers::proxy::proxy_request, middleware, state::AppState}; +// Keep the compatibility `serve_tcp` entry point subject to the same parser +// protections as the configured binary listener. These are metadata limits; +// request and response bodies remain streaming after the first request gate +// opens. The HTTP/2 stream default intentionally stays high for capable hosts. +const DEFAULT_TCP_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 16_384; +const DEFAULT_TCP_HTTP_HEADER_READ_TIMEOUT: Duration = Duration::from_secs(30); +const DEFAULT_TCP_HTTP_HEADER_MAX_BYTES: usize = 64 * 1024; +const DEFAULT_TCP_HTTP_MAX_HEADERS: usize = 256; + +#[derive(Clone)] +struct FirstRequestGate { + seen: Arc, + notify: Arc, +} + +impl FirstRequestGate { + fn new() -> Self { + Self { + seen: Arc::new(AtomicBool::new(false)), + notify: Arc::new(tokio::sync::Notify::new()), + } + } + + fn mark_seen(&self) { + if !self.seen.swap(true, Ordering::Release) { + self.notify.notify_one(); + } + } + + fn is_seen(&self) -> bool { + self.seen.load(Ordering::Acquire) + } +} + +#[derive(Clone)] +struct FirstRequestService { + inner: S, + gate: FirstRequestGate, +} + +impl tower::Service for FirstRequestService +where + S: tower::Service, +{ + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, request: Req) -> Self::Future { + self.gate.mark_seen(); + self.inner.call(request) + } +} + +async fn drive_first_request_gate( + connection: F, + gate: FirstRequestGate, + timeout: Duration, +) -> Result<(), E> +where + F: std::future::Future>, +{ + if gate.is_seen() { + return connection.await; + } + + let mut connection = Box::pin(connection); + let timeout = tokio::time::sleep(timeout); + tokio::pin!(timeout); + let notified = gate.notify.notified(); + tokio::pin!(notified); + + tokio::select! { + result = &mut connection => result, + _ = &mut timeout => { + if gate.is_seen() { + (&mut connection).await + } else { + Ok(()) + } + } + _ = &mut notified => (&mut connection).await, + } +} + pub fn build_router() -> Result { Ok(build_router_with_state(AppState::new()?)) } @@ -29,7 +129,7 @@ pub fn build_router_with_state(state: AppState) -> Router { let cors_state = state.clone(); let mut router = Router::::new(); router = api::mount_core_routes(router); - router = api::mount_operational_routes(router); + router = api::mount_operational_routes(router, state.clone()); router = api::mount_ai_routes(router); router = api::mount_public_support_routes(router); router = api::mount_oauth_routes(router); @@ -144,10 +244,45 @@ pub(crate) enum RequestAdmissionError { pub async fn serve_tcp(bind: &str) -> Result<(), Box> { let listener = tokio::net::TcpListener::bind(bind).await?; let router = build_router()?; - axum::serve( - listener, - router.into_make_service_with_connect_info::(), - ) - .await?; - Ok(()) + let mut make_service = router.into_make_service_with_connect_info::(); + loop { + let (io, remote_addr) = listener.accept().await?; + let tower_service = make_service + .call(remote_addr) + .await + .unwrap_or_else(|err| match err {}) + .map_request(|request: http::Request| request.map(Body::new)); + let first_request_gate = FirstRequestGate::new(); + let hyper_service = TowerToHyperService::new(FirstRequestService { + inner: tower_service, + gate: first_request_gate.clone(), + }); + let io = TokioIo::new(io); + + tokio::spawn(async move { + let mut builder = HyperServerBuilder::new(TokioExecutor::new()); + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(DEFAULT_TCP_HTTP_HEADER_READ_TIMEOUT) + .max_buf_size(DEFAULT_TCP_HTTP_HEADER_MAX_BYTES) + .max_headers(DEFAULT_TCP_HTTP_MAX_HEADERS); + builder + .http2() + .timer(TokioTimer::new()) + .enable_connect_protocol() + .max_concurrent_streams(DEFAULT_TCP_HTTP2_MAX_CONCURRENT_STREAMS) + .max_header_list_size(DEFAULT_TCP_HTTP_HEADER_MAX_BYTES as u32); + + let result = drive_first_request_gate( + builder.serve_connection_with_upgrades(io, hyper_service), + first_request_gate, + DEFAULT_TCP_HTTP_HEADER_READ_TIMEOUT, + ) + .await; + if let Err(error) = result { + tracing::trace!(error = ?error, "compatibility gateway connection closed with error"); + } + }); + } } diff --git a/apps/aether-gateway/src/routing/mutations.rs b/apps/aether-gateway/src/routing/mutations.rs index 65d2ad517..bb2b44599 100644 --- a/apps/aether-gateway/src/routing/mutations.rs +++ b/apps/aether-gateway/src/routing/mutations.rs @@ -8,22 +8,25 @@ use serde_json::Value; use crate::GatewayError; +const INVALID_ROUTING_MUTATION_MESSAGE: &str = "invalid routing mutation"; + pub(crate) fn apply_routing_mutation_plan( body: &mut Value, headers: &mut HeaderMap, plan: &MutationPlan, ) -> Result<(), GatewayError> { - apply_json_patch_operations(body, &plan.body_patch).map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: err.to_string(), - })?; - apply_header_patch(headers, &plan.header_patch).map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: err.to_string(), - })?; + apply_json_patch_operations(body, &plan.body_patch).map_err(|_| invalid_routing_mutation())?; + apply_header_patch(headers, &plan.header_patch).map_err(|_| invalid_routing_mutation())?; Ok(()) } +fn invalid_routing_mutation() -> GatewayError { + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: INVALID_ROUTING_MUTATION_MESSAGE.to_string(), + } +} + fn apply_header_patch( headers: &mut HeaderMap, patch: &[RoutingHeaderPatch], @@ -47,3 +50,64 @@ fn apply_header_patch( } Ok(()) } + +#[cfg(test)] +mod tests { + use aether_routing_core::RoutingJsonPatchOperation; + use serde_json::json; + + use super::*; + + #[test] + fn body_mutation_errors_do_not_echo_json_pointer() { + let secret = "https://internal.example/?token=Bearer-secret"; + let plan = MutationPlan { + body_patch: vec![RoutingJsonPatchOperation::Replace { + path: secret.to_string(), + value: json!("replacement"), + }], + ..MutationPlan::default() + }; + + let error = apply_routing_mutation_plan( + &mut json!({"model": "test"}), + &mut HeaderMap::new(), + &plan, + ) + .expect_err("invalid pointer should fail"); + + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + ref message, + } if message == INVALID_ROUTING_MUTATION_MESSAGE && !message.contains(secret) + )); + } + + #[test] + fn header_mutation_errors_do_not_echo_header_name() { + let secret = "Authorization: Bearer secret"; + let plan = MutationPlan { + header_patch: vec![RoutingHeaderPatch::Remove { + name: secret.to_string(), + }], + ..MutationPlan::default() + }; + + let error = apply_routing_mutation_plan( + &mut json!({"model": "test"}), + &mut HeaderMap::new(), + &plan, + ) + .expect_err("invalid header should fail"); + + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + ref message, + } if message == INVALID_ROUTING_MUTATION_MESSAGE && !message.contains(secret) + )); + } +} diff --git a/apps/aether-gateway/src/routing/resolver.rs b/apps/aether-gateway/src/routing/resolver.rs index 451bbedd7..b636abd02 100644 --- a/apps/aether-gateway/src/routing/resolver.rs +++ b/apps/aether-gateway/src/routing/resolver.rs @@ -1,7 +1,7 @@ use aether_routing_core::{ resolve_routing_policy, MutationPlan, RankingOverlay, ResolvedRoutingPolicy, - RoutingGroupConfig, RoutingPolicyInput, RoutingRulePhase, RoutingSchedulingMode, - RoutingSetPriorityMode, + RoutingDefaultPolicy, RoutingGroupConfig, RoutingPolicyError, RoutingPolicyInput, + RoutingRulePhase, RoutingSchedulingMode, RoutingSetPriorityMode, DEFAULT_STICKY_KEY_ATTEMPTS, }; use http::StatusCode; use serde_json::Value; @@ -9,6 +9,10 @@ use std::collections::BTreeMap; use crate::GatewayError; +const INVALID_ROUTING_GROUP_CONFIG_MESSAGE: &str = "invalid routing group config"; +const INVALID_ROUTING_MUTATION_MESSAGE: &str = "invalid routing mutation"; +const ROUTING_MODEL_NOT_ALLOWED_MESSAGE: &str = "requested model is not allowed by routing policy"; + #[derive(Debug, Clone)] pub(crate) struct GatewayRoutingPolicyInput<'a> { pub group_id: Option<&'a str>, @@ -52,10 +56,7 @@ pub(crate) fn resolve_gateway_routing_policy( } let config = serde_json::from_value::(input.group_config_json.clone()) - .map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: format!("invalid routing group config: {err}"), - })?; + .map_err(|_| invalid_routing_group_config())?; resolve_routing_policy( &config, RoutingPolicyInput { @@ -72,18 +73,13 @@ pub(crate) fn resolve_gateway_routing_policy( phase: input.phase, }, ) - .map_err(|err| GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: err.to_string(), - }) + .map_err(routing_policy_error) } pub(crate) fn resolve_gateway_static_default_routing_policy( input: GatewayStaticRoutingPolicyInput<'_>, ) -> Result, GatewayError> { - let Some((priority_mode, scheduling_mode, keep_priority_on_conversion)) = - static_default_policy_fields(input.group_config_json)? - else { + let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else { return Ok(None); }; @@ -93,9 +89,11 @@ pub(crate) fn resolve_gateway_static_default_routing_policy( selection_source: input.selection_source.to_string(), requested_model: input.requested_model.to_string(), resolved_model: input.resolved_model.to_string(), - priority_mode, - scheduling_mode, - keep_priority_on_conversion, + priority_mode: default_policy.priority_mode, + scheduling_mode: default_policy.scheduling_mode, + keep_priority_on_conversion: default_policy.keep_priority_on_conversion, + sticky_key_attempts: default_policy.sticky_key_attempts, + execution_policy: default_policy.execution_policy, ranking_overlay: RankingOverlay::default(), mutation_plan: MutationPlan::default(), pool_policy_overrides: BTreeMap::new(), @@ -105,23 +103,21 @@ pub(crate) fn resolve_gateway_static_default_routing_policy( fn static_default_policy_fields( config_json: &Value, -) -> Result, GatewayError> { +) -> Result, GatewayError> { let Some(object) = config_json.as_object() else { return Ok(None); }; - if !routing_array_field_is_missing_or_empty(object, "allowed_models") - || !routing_array_field_is_missing_or_empty(object, "model_policies") + // A strategy's default policy applies to every model. Only model policies + // and rules require the request-context-aware resolver; unknown legacy + // fields (including the removed group allowlist) are intentionally ignored. + if !routing_array_field_is_missing_or_empty(object, "model_policies") || !routing_array_field_is_missing_or_empty(object, "rules") { return Ok(None); } let Some(default_policy) = object.get("default_policy") else { - return Ok(Some(( - RoutingSetPriorityMode::default(), - RoutingSchedulingMode::default(), - false, - ))); + return Ok(Some(RoutingDefaultPolicy::default())); }; let Some(default_policy) = default_policy.as_object() else { return Ok(None); @@ -136,17 +132,53 @@ fn static_default_policy_fields( RoutingSchedulingMode::default, )?; let keep_priority_on_conversion = match default_policy.get("keep_priority_on_conversion") { - Some(value) => value.as_bool().ok_or_else(|| { - invalid_routing_group_config("keep_priority_on_conversion must be a boolean") - })?, + Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config)?, None => false, }; + let sticky_key_attempts = match default_policy.get("sticky_key_attempts") { + Some(value) => value + .as_u64() + .and_then(|value| u32::try_from(value).ok()) + .ok_or_else(invalid_routing_group_config)?, + None => DEFAULT_STICKY_KEY_ATTEMPTS, + }; + let enable_cf_heartbeat = routing_bool_field( + default_policy.get("enable_cf_heartbeat"), + "enable_cf_heartbeat", + )?; + // Older strategies stored separate image/text heartbeat flags. Treat + // either legacy flag as enabling the unified CF heartbeat setting while + // allowing newly saved strategies to use only the canonical key. + let legacy_image_heartbeat = routing_bool_field( + default_policy.get("enable_openai_image_sync_heartbeat"), + "enable_openai_image_sync_heartbeat", + )?; + let legacy_text_heartbeat = routing_bool_field( + default_policy.get("enable_standard_text_sync_heartbeat"), + "enable_standard_text_sync_heartbeat", + )?; + let execution_policy = aether_routing_core::RoutingExecutionPolicy { + enable_cf_heartbeat: enable_cf_heartbeat || legacy_image_heartbeat || legacy_text_heartbeat, + cyber_continue_failover: routing_bool_field( + default_policy.get("cyber_continue_failover"), + "cyber_continue_failover", + )?, + }; - Ok(Some(( + Ok(Some(RoutingDefaultPolicy { priority_mode, scheduling_mode, keep_priority_on_conversion, - ))) + sticky_key_attempts, + execution_policy, + })) +} + +fn routing_bool_field(value: Option<&Value>, _field: &str) -> Result { + match value { + Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config), + None => Ok(false), + } } fn routing_array_field_is_missing_or_empty( @@ -168,17 +200,29 @@ where T: serde::de::DeserializeOwned, { match value { - Some(value) => serde_json::from_value(value.clone()).map_err(|err| { - invalid_routing_group_config(format!("invalid default routing policy: {err}")) - }), + Some(value) => { + serde_json::from_value(value.clone()).map_err(|_| invalid_routing_group_config()) + } None => Ok(default()), } } -fn invalid_routing_group_config(message: impl Into) -> GatewayError { +fn routing_policy_error(error: RoutingPolicyError) -> GatewayError { + let message = match error { + RoutingPolicyError::InvalidConfig(_) => INVALID_ROUTING_GROUP_CONFIG_MESSAGE, + RoutingPolicyError::InvalidMutation(_) => INVALID_ROUTING_MUTATION_MESSAGE, + RoutingPolicyError::ModelNotAllowed(_) => ROUTING_MODEL_NOT_ALLOWED_MESSAGE, + }; GatewayError::Client { status: StatusCode::BAD_REQUEST, - message: format!("invalid routing group config: {}", message.into()), + message: message.to_string(), + } +} + +fn invalid_routing_group_config() -> GatewayError { + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: INVALID_ROUTING_GROUP_CONFIG_MESSAGE.to_string(), } } @@ -196,7 +240,7 @@ mod tests { "scheduling_mode": "load_balance", "keep_priority_on_conversion": true }, - "allowed_models": [], + "allowed_models": ["legacy-model"], "model_policies": [], "rules": [] }); @@ -269,4 +313,72 @@ mod tests { assert!(policy.is_none()); } + + #[test] + fn routing_config_errors_do_not_echo_config_values() { + let secret = "https://internal.example/?token=Bearer-secret"; + let config = json!({ + "default_policy": { + "priority_mode": secret + } + }); + + let error = + resolve_gateway_static_default_routing_policy(GatewayStaticRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: "mock-model", + resolved_model: "mock-model", + }) + .expect_err("invalid config should fail"); + + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + ref message, + } if message == INVALID_ROUTING_GROUP_CONFIG_MESSAGE && !message.contains(secret) + )); + } + + #[test] + fn routing_policy_errors_do_not_echo_requested_model() { + let secret = "model?token=Bearer-secret"; + let config = json!({ + "rules": [{ + "id": "restrict-model", + "conditions": {}, + "actions": [{ + "type": "restrict_models", + "models": ["allowed-model"] + }] + }] + }); + + let error = resolve_gateway_routing_policy(GatewayRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: secret, + resolved_model: secret, + api_format: "openai:chat", + user_id: Some("user-1"), + api_key_id: Some("key-1"), + headers: &json!({}), + body: &json!({"model": secret}), + phase: RoutingRulePhase::ClientRequest, + }) + .expect_err("disallowed model should fail"); + + assert!(matches!( + error, + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + ref message, + } if message == ROUTING_MODEL_NOT_ALLOWED_MESSAGE && !message.contains(secret) + )); + } } diff --git a/apps/aether-gateway/src/routing/selection.rs b/apps/aether-gateway/src/routing/selection.rs index bef1ccfcd..2d3d2e953 100644 --- a/apps/aether-gateway/src/routing/selection.rs +++ b/apps/aether-gateway/src/routing/selection.rs @@ -8,6 +8,8 @@ pub(crate) const ROUTING_GROUP_HEADER: &str = "x-aether-scheduler-group"; #[derive(Debug, Error, Clone, PartialEq, Eq)] pub(crate) enum GatewayRoutingSelectionError { + #[error("no enabled routing strategy is configured for this request")] + NoDefault, #[error("routing group was explicitly requested but was not found: {0}")] NotFound(String), #[error("routing group was explicitly requested but is not enabled: {0}")] @@ -239,6 +241,7 @@ mod tests { description: None, enabled: true, is_system_default: false, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, @@ -287,6 +290,7 @@ mod tests { description: None, enabled: true, is_system_default: true, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, @@ -322,6 +326,7 @@ mod tests { description: None, enabled: true, is_system_default: false, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, @@ -444,6 +449,7 @@ mod tests { description: None, enabled: false, is_system_default: false, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, @@ -481,6 +487,7 @@ mod tests { description: None, enabled: true, is_system_default: false, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, diff --git a/apps/aether-gateway/src/scheduler/candidate/mod.rs b/apps/aether-gateway/src/scheduler/candidate/mod.rs index 1da74c2ad..7a53b268f 100644 --- a/apps/aether-gateway/src/scheduler/candidate/mod.rs +++ b/apps/aether-gateway/src/scheduler/candidate/mod.rs @@ -1,7 +1,8 @@ use self::selection::{ - collect_selectable_candidates, collect_selectable_candidates_with_skip_reasons, + collect_selectable_candidates, collect_selectable_candidates_with_skip_reasons_and_ordering, collect_selectable_enumerated_candidates_with_skip_reasons, }; +use super::config::SchedulerOrderingConfig; use super::state::SchedulerRuntimeState; mod affinity; @@ -53,6 +54,9 @@ enum RequiredCapabilityMatchMode { Exclusive, } +/// `ordering_config` carries the request's routing-policy derived scheduler +/// config. Every production scheduling pass must provide this snapshot. +#[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -64,6 +68,7 @@ pub(crate) async fn list_selectable_candidates( client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { collect_selectable_candidates( selection_row_source, @@ -76,6 +81,7 @@ pub(crate) async fn list_selectable_candidates( client_session_affinity, now_unix_secs, enable_model_directives, + ordering_config, ) .await } @@ -87,6 +93,7 @@ pub(crate) fn is_exact_all_skipped_by_auth_limit( selection::is_exact_all_skipped_by_auth_limit(selected, skipped) } +#[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_with_skip_reasons( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -98,6 +105,7 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons( client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -105,7 +113,7 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons( ), GatewayError, > { - collect_selectable_candidates_with_skip_reasons( + collect_selectable_candidates_with_skip_reasons_and_ordering( selection_row_source, runtime_state, api_format, @@ -117,10 +125,12 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons( now_unix_secs, enable_model_directives, None, + ordering_config, ) .await } +#[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_operation( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -133,6 +143,7 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_ope now_unix_secs: u64, enable_model_directives: bool, request_operation: Option<&str>, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -140,7 +151,7 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_ope ), GatewayError, > { - collect_selectable_candidates_with_skip_reasons( + collect_selectable_candidates_with_skip_reasons_and_ordering( selection_row_source, runtime_state, api_format, @@ -152,6 +163,7 @@ pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_ope now_unix_secs, enable_model_directives, request_operation, + ordering_config, ) .await } @@ -166,6 +178,7 @@ pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, + ordering_config: SchedulerOrderingConfig, ) -> Result< ( Vec, @@ -173,7 +186,6 @@ pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons( ), GatewayError, > { - let ordering_config = runtime_state.read_scheduler_ordering_config().await?; let priority_affinity_key = selection::scheduling_priority_affinity_key( auth_snapshot, client_session_affinity, @@ -194,6 +206,7 @@ pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons( .await } +#[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_for_required_capability_without_requested_model( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -203,6 +216,7 @@ pub(crate) async fn list_selectable_candidates_for_required_capability_without_r auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { Ok( list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal( @@ -214,12 +228,14 @@ pub(crate) async fn list_selectable_candidates_for_required_capability_without_r auth_snapshot, client_session_affinity, now_unix_secs, + ordering_config, ) .await? .0, ) } +#[allow(clippy::too_many_arguments)] pub(crate) async fn list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -229,6 +245,7 @@ pub(crate) async fn list_selectable_candidates_for_required_capability_without_r auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, + ordering_config: SchedulerOrderingConfig, ) -> Result<(Vec, bool), GatewayError> { let normalized_api_format = normalize_api_format(candidate_api_format); if normalized_api_format.is_empty() { @@ -262,20 +279,22 @@ pub(crate) async fn list_selectable_candidates_for_required_capability_without_r let mut all_attempts_blocked_by_auth_limit = !model_names.is_empty(); for global_model_name in model_names { - let (candidates, skipped_candidates) = collect_selectable_candidates_with_skip_reasons( - selection_row_source, - runtime_state, - &normalized_api_format, - &global_model_name, - require_streaming, - required_capabilities.as_ref(), - auth_snapshot, - client_session_affinity, - now_unix_secs, - false, - None, - ) - .await?; + let (candidates, skipped_candidates) = + collect_selectable_candidates_with_skip_reasons_and_ordering( + selection_row_source, + runtime_state, + &normalized_api_format, + &global_model_name, + require_streaming, + required_capabilities.as_ref(), + auth_snapshot, + client_session_affinity, + now_unix_secs, + false, + None, + ordering_config, + ) + .await?; all_attempts_blocked_by_auth_limit &= is_exact_all_skipped_by_auth_limit(&candidates, &skipped_candidates); match capability_mode { diff --git a/apps/aether-gateway/src/scheduler/candidate/runtime.rs b/apps/aether-gateway/src/scheduler/candidate/runtime.rs index 0227abd5d..25bb44260 100644 --- a/apps/aether-gateway/src/scheduler/candidate/runtime.rs +++ b/apps/aether-gateway/src/scheduler/candidate/runtime.rs @@ -341,18 +341,33 @@ fn read_key_account_quota_exhaustion_map( candidates .iter() .map(|candidate| { - let exhausted = provider_skip_exhausted_accounts - .get(candidate.provider_id.as_str()) - .copied() - .unwrap_or(false) - && provider_key_rpm_states - .get(candidate.key_id.as_str()) - .is_some_and(|key| { - admin_provider_pool_pure::admin_pool_key_account_quota_exhausted( + let exhausted = provider_key_rpm_states + .get(candidate.key_id.as_str()) + .is_some_and(|key| { + let account_exhausted = + admin_provider_pool_pure::admin_pool_key_model_quota_exhausted( key, candidate.provider_type.as_str(), + candidate.selected_provider_model_name.as_str(), ) - }); + .unwrap_or_else(|| { + admin_provider_pool_pure::admin_pool_key_account_quota_exhausted( + key, + candidate.provider_type.as_str(), + ) + }); + let hard_blocked = + admin_provider_pool_pure::admin_pool_key_model_quota_hard_blocked( + key, + candidate.provider_type.as_str(), + candidate.selected_provider_model_name.as_str(), + ); + let skip_configured = provider_skip_exhausted_accounts + .get(candidate.provider_id.as_str()) + .copied() + .unwrap_or(false); + hard_blocked || (skip_configured && account_exhausted) + }); (candidate.key_id.clone(), exhausted) }) .collect() diff --git a/apps/aether-gateway/src/scheduler/candidate/selection.rs b/apps/aether-gateway/src/scheduler/candidate/selection.rs index 5a784dced..36c46f056 100644 --- a/apps/aether-gateway/src/scheduler/candidate/selection.rs +++ b/apps/aether-gateway/src/scheduler/candidate/selection.rs @@ -1,7 +1,7 @@ use crate::data::auth::GatewayAuthApiKeySnapshot; use crate::data::candidate_selection::MinimalCandidateSelectionRowSource; use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL; -use crate::scheduler::config::SchedulerSchedulingMode; +use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode}; use crate::GatewayError; use aether_scheduler_core::ClientSessionAffinity; @@ -47,7 +47,7 @@ pub(super) fn is_exact_all_skipped_by_auth_limit( .all(|candidate| is_auth_api_key_concurrency_limit_skip_reason(candidate.skip_reason)) } -#[cfg_attr(not(test), allow(dead_code))] +#[cfg(test)] pub(super) async fn select_minimal_candidate( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -59,20 +59,8 @@ pub(super) async fn select_minimal_candidate( client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { - let affinity_epoch = runtime_state.scheduler_affinity_epoch(); - let ordering_config = runtime_state.read_scheduler_ordering_config().await?; - let affinity_cache_key = build_scheduler_affinity_cache_key( - auth_snapshot, - api_format, - global_model_name, - client_session_affinity, - ); - let priority_affinity_key = scheduling_priority_affinity_key( - auth_snapshot, - client_session_affinity, - ordering_config.scheduling_mode, - ); let candidates = enumerate_scheduler_candidates( selection_row_source, api_format, @@ -84,7 +72,7 @@ pub(super) async fn select_minimal_candidate( None, ) .await?; - let selected = collect_selectable_enumerated_candidates_with_skip_reasons( + Ok(collect_selectable_enumerated_candidates_with_skip_reasons( runtime_state, api_format, global_model_name, @@ -94,27 +82,19 @@ pub(super) async fn select_minimal_candidate( client_session_affinity, now_unix_secs, ordering_config, - priority_affinity_key, + scheduling_priority_affinity_key( + auth_snapshot, + client_session_affinity, + ordering_config.scheduling_mode, + ), ) .await? .0 .into_iter() - .next(); - if ordering_config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity - && has_explicit_session_affinity(client_session_affinity) - { - if let Some(candidate) = selected.as_ref() { - remember_scheduler_affinity( - affinity_cache_key.as_deref(), - runtime_state, - candidate, - Some(affinity_epoch), - ); - } - } - Ok(selected) + .next()) } +#[allow(clippy::too_many_arguments)] pub(super) async fn collect_selectable_candidates( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -126,24 +106,30 @@ pub(super) async fn collect_selectable_candidates( client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, enable_model_directives: bool, + ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { - Ok(collect_selectable_candidates_with_skip_reasons( - selection_row_source, - runtime_state, - api_format, - global_model_name, - require_streaming, - required_capabilities, - auth_snapshot, - client_session_affinity, - now_unix_secs, - enable_model_directives, - None, + Ok( + collect_selectable_candidates_with_skip_reasons_and_ordering( + selection_row_source, + runtime_state, + api_format, + global_model_name, + require_streaming, + required_capabilities, + auth_snapshot, + client_session_affinity, + now_unix_secs, + enable_model_directives, + None, + ordering_config, + ) + .await? + .0, ) - .await? - .0) } +#[allow(clippy::too_many_arguments)] +#[cfg(test)] pub(super) async fn collect_selectable_candidates_with_skip_reasons( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &impl SchedulerRuntimeState, @@ -163,7 +149,44 @@ pub(super) async fn collect_selectable_candidates_with_skip_reasons( ), GatewayError, > { - let ordering_config = runtime_state.read_scheduler_ordering_config().await?; + collect_selectable_candidates_with_skip_reasons_and_ordering( + selection_row_source, + runtime_state, + api_format, + global_model_name, + require_streaming, + required_capabilities, + auth_snapshot, + client_session_affinity, + now_unix_secs, + enable_model_directives, + request_operation, + SchedulerOrderingConfig::default(), + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn collect_selectable_candidates_with_skip_reasons_and_ordering( + selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), + runtime_state: &impl SchedulerRuntimeState, + api_format: &str, + global_model_name: &str, + require_streaming: bool, + required_capabilities: Option<&serde_json::Value>, + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, + client_session_affinity: Option<&ClientSessionAffinity>, + now_unix_secs: u64, + enable_model_directives: bool, + request_operation: Option<&str>, + ordering_config: SchedulerOrderingConfig, +) -> Result< + ( + Vec, + Vec, + ), + GatewayError, +> { let priority_affinity_key = scheduling_priority_affinity_key( auth_snapshot, client_session_affinity, @@ -205,7 +228,7 @@ pub(super) async fn collect_selectable_enumerated_candidates_with_skip_reasons( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, client_session_affinity: Option<&ClientSessionAffinity>, now_unix_secs: u64, - ordering_config: crate::scheduler::config::SchedulerOrderingConfig, + ordering_config: SchedulerOrderingConfig, priority_affinity_key: Option<&str>, ) -> Result< ( diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/affinity.rs b/apps/aether-gateway/src/scheduler/candidate/tests/affinity.rs index 0fe2b2467..ec991c795 100644 --- a/apps/aether-gateway/src/scheduler/candidate/tests/affinity.rs +++ b/apps/aether-gateway/src/scheduler/candidate/tests/affinity.rs @@ -23,6 +23,7 @@ use crate::data::candidate_selection::{ read_requested_model_rows, MinimalCandidateSelectionRowSource, }; use crate::data::GatewayDataState; +use crate::scheduler::config::SchedulerOrderingConfig; use crate::{AppState, GatewayError}; use super::super::affinity::build_scheduler_affinity_cache_key; @@ -50,6 +51,7 @@ async fn select_candidate( client_session_affinity, now_unix_secs, false, + SchedulerOrderingConfig::default(), ) .await } diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/required_capability.rs b/apps/aether-gateway/src/scheduler/candidate/tests/required_capability.rs index 2fb9e5f2e..89e7db9e7 100644 --- a/apps/aether-gateway/src/scheduler/candidate/tests/required_capability.rs +++ b/apps/aether-gateway/src/scheduler/candidate/tests/required_capability.rs @@ -65,6 +65,7 @@ async fn compatible_required_capability_prefers_matching_keys_without_hard_filte None, None, 100, + crate::scheduler::config::SchedulerOrderingConfig::default(), ) .await .expect("selection should succeed"); @@ -120,6 +121,7 @@ async fn exclusive_required_capability_keeps_hard_filtering_only_matching_keys() None, None, 100, + crate::scheduler::config::SchedulerOrderingConfig::default(), ) .await .expect("selection should succeed"); @@ -196,6 +198,7 @@ async fn required_capability_without_model_uses_session_scoped_affinity() { Some(&auth_snapshot), Some(&client_session_affinity), 100, + crate::scheduler::config::SchedulerOrderingConfig::default(), ) .await .expect("selection should succeed"); @@ -273,6 +276,7 @@ async fn required_capability_reports_auth_limit_signal_when_every_model_is_block Some(&auth_snapshot), None, 100, + crate::scheduler::config::SchedulerOrderingConfig::default(), ) .await .expect("selection should succeed"); diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs b/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs index b1578c777..6efd0a329 100644 --- a/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs +++ b/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs @@ -1,10 +1,12 @@ use std::sync::Arc; use std::time::Duration; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::quota::InMemoryProviderQuotaRepository; +use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; use aether_data_contracts::repository::candidate_selection::{ StoredMinimalCandidateSelectionRow, StoredProviderModelMapping, }; @@ -13,6 +15,9 @@ use aether_data_contracts::repository::candidates::{ }; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot; +use aether_data_contracts::repository::routing_profiles::{ + CreateRoutingGroupRecord, RoutingGroupWriteRepository, +}; use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate}; use serde_json::json; @@ -20,6 +25,7 @@ use crate::cache::SchedulerAffinityTarget; use crate::data::auth::GatewayAuthApiKeySnapshot; use crate::data::candidate_selection::MinimalCandidateSelectionRowSource; use crate::data::GatewayDataState; +use crate::scheduler::config::SchedulerOrderingConfig; use crate::{AppState, GatewayError}; use super::super::affinity::build_scheduler_affinity_cache_key; @@ -31,6 +37,39 @@ use super::super::selection::{ }; use super::support::{sample_auth_snapshot, sample_key, sample_provider, sample_row}; +async fn state_with_routing_default_policy( + data_state: GatewayDataState, + default_policy: serde_json::Value, +) -> AppState { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + repository + .create_routing_group(CreateRoutingGroupRecord { + id: "selection-test-default".to_string(), + name: "selection-test-default".to_string(), + description: None, + enabled: true, + is_system_default: true, + sort_order: 0, + config_json: json!({"default_policy": default_policy}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .expect("routing strategy should be created"); + AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state.with_routing_group_repository_for_tests(repository)) +} + +async fn ordering_config(state: &AppState) -> SchedulerOrderingConfig { + crate::scheduler::config::read_system_default_routing_ordering_config(state) + .await + .expect("routing strategy should load") + .unwrap_or_default() +} + async fn select_candidate( selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync), runtime_state: &AppState, @@ -40,6 +79,7 @@ async fn select_candidate( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, now_unix_secs: u64, ) -> Result, GatewayError> { + let ordering_config = ordering_config(runtime_state).await; select_candidate_impl( selection_row_source, runtime_state, @@ -51,6 +91,7 @@ async fn select_candidate( None, now_unix_secs, false, + ordering_config, ) .await } @@ -64,6 +105,7 @@ async fn collect_selectable_candidates( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, now_unix_secs: u64, ) -> Result, GatewayError> { + let ordering_config = ordering_config(runtime_state).await; collect_selectable_candidates_impl( selection_row_source, runtime_state, @@ -75,6 +117,7 @@ async fn collect_selectable_candidates( None, now_unix_secs, false, + ordering_config, ) .await } @@ -140,6 +183,21 @@ fn provider_key_with_concurrent_limit( key } +fn sealed_kiro_auth_config() -> String { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + credential_state + .seal_provider_catalog_key_auth_config( + "provider-kiro", + "key-kiro", + r#"{"refresh_token":"refreshable-session"}"#, + ) + .expect("auth config should encrypt") +} + fn active_provider_key_candidate( candidate_id: &str, request_id: &str, @@ -288,15 +346,11 @@ async fn selects_by_provider_priority_when_priority_mode_is_provider() { global_key_first, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "provider_priority_mode".to_string(), - json!("provider"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"priority_mode": "provider"}), + ) + .await; let selected = select_candidate( state.data.as_ref(), @@ -342,15 +396,11 @@ async fn selects_by_global_key_priority_when_priority_mode_is_global_key() { global_key_first, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "provider_priority_mode".to_string(), - json!("global_key"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"priority_mode": "global_key"}), + ) + .await; let selected = select_candidate( state.data.as_ref(), @@ -414,6 +464,7 @@ async fn scheduler_selection_prefers_required_capability_matches_before_priority None, 100, false, + SchedulerOrderingConfig::default(), ) .await .expect("selection should succeed") @@ -449,15 +500,11 @@ async fn fixed_order_ignores_cached_scheduler_affinity_promotion() { first, second, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("fixed_order"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "fixed_order"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); state.remember_scheduler_affinity_target( @@ -514,15 +561,11 @@ async fn fixed_order_disables_same_priority_affinity_hash_tiebreaker() { first, second, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("fixed_order"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "fixed_order"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); let selection = collect_selectable_candidates( @@ -568,15 +611,11 @@ async fn cache_affinity_promotes_cached_scheduler_affinity_candidate_when_enable first, second, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("cache_affinity"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "cache_affinity"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); let client_session_affinity = ClientSessionAffinity::from_session_key("session-1"); @@ -609,6 +648,7 @@ async fn cache_affinity_promotes_cached_scheduler_affinity_candidate_when_enable Some(&client_session_affinity), 100, false, + ordering_config(&state).await, ) .await .expect("selection should succeed") @@ -644,15 +684,11 @@ async fn cache_affinity_ignores_cached_scheduler_affinity_without_client_session first, second, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("cache_affinity"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "cache_affinity"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); state.remember_scheduler_affinity_target( @@ -690,15 +726,11 @@ async fn load_balance_selection_does_not_remember_scheduler_affinity() { row, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("load_balance"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "load_balance"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); let client_session_affinity = ClientSessionAffinity::from_session_key("session-1"); let cache_key = build_scheduler_affinity_cache_key( @@ -720,6 +752,7 @@ async fn load_balance_selection_does_not_remember_scheduler_affinity() { Some(&client_session_affinity), 100, false, + ordering_config(&state).await, ) .await .expect("selection should succeed") @@ -757,15 +790,11 @@ async fn load_balance_ignores_provider_priority_and_cached_affinity() { first, second, ])); let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![])); - let state = AppState::new() - .expect("state should build") - .with_data_state_for_tests( - GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas) - .with_system_config_values_for_tests(vec![( - "scheduling_mode".to_string(), - json!("load_balance"), - )]), - ); + let state = state_with_routing_default_policy( + GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas), + json!({"scheduling_mode": "load_balance"}), + ) + .await; let auth_snapshot = sample_auth_snapshot("affinity-key-1"); state.remember_scheduler_affinity_target( @@ -2211,7 +2240,7 @@ async fn keeps_refreshable_kiro_candidate_selectable_with_runtime_oauth_invalid_ vec![{ let mut key = sample_key("key-kiro", "provider-kiro", Some(10)); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some("encrypted-refreshable-session".to_string()); + key.encrypted_auth_config = Some(sealed_kiro_auth_config()); key.oauth_invalid_at_unix_secs = Some(1_710_000_000); key.oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string()); key @@ -2227,7 +2256,8 @@ async fn keeps_refreshable_kiro_candidate_selectable_with_runtime_oauth_invalid_ provider_catalog, quotas, request_candidates, - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let (selected, skipped) = collect_selectable_candidates_with_skip_reasons( @@ -2271,7 +2301,7 @@ async fn keeps_refreshable_kiro_candidate_selectable_when_oauth_token_expired() vec![{ let mut key = sample_key("key-kiro", "provider-kiro", Some(10)); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some("encrypted-refreshable-session".to_string()); + key.encrypted_auth_config = Some(sealed_kiro_auth_config()); key.expires_at_unix_secs = Some(1_710_000_000); key }], @@ -2286,7 +2316,8 @@ async fn keeps_refreshable_kiro_candidate_selectable_when_oauth_token_expired() provider_catalog, quotas, request_candidates, - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let (selected, skipped) = collect_selectable_candidates_with_skip_reasons( @@ -2330,7 +2361,7 @@ async fn keeps_kiro_candidate_selectable_after_refresh_token_failure_until_acces vec![{ let mut key = sample_key("key-kiro", "provider-kiro", Some(10)); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some("encrypted-refreshable-session".to_string()); + key.encrypted_auth_config = Some(sealed_kiro_auth_config()); key.expires_at_unix_secs = Some(1_710_000_200); key.oauth_invalid_at_unix_secs = Some(1_710_000_000); key.oauth_invalid_reason = Some( @@ -2350,7 +2381,8 @@ async fn keeps_kiro_candidate_selectable_after_refresh_token_failure_until_acces provider_catalog, quotas, request_candidates, - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let (selected, skipped) = collect_selectable_candidates_with_skip_reasons( @@ -2394,7 +2426,7 @@ async fn skips_kiro_candidate_after_refresh_token_failure_and_access_token_expir vec![{ let mut key = sample_key("key-kiro", "provider-kiro", Some(10)); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some("encrypted-refreshable-session".to_string()); + key.encrypted_auth_config = Some(sealed_kiro_auth_config()); key.expires_at_unix_secs = Some(1_710_000_000); key.oauth_invalid_at_unix_secs = Some(1_710_000_000); key.oauth_invalid_reason = Some( @@ -2414,7 +2446,8 @@ async fn skips_kiro_candidate_after_refresh_token_failure_and_access_token_expir provider_catalog, quotas, request_candidates, - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let (selected, skipped) = collect_selectable_candidates_with_skip_reasons( @@ -2459,7 +2492,7 @@ async fn skips_refreshable_kiro_candidate_when_oauth_marker_is_account_block() { vec![{ let mut key = sample_key("key-kiro", "provider-kiro", Some(10)); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some("encrypted-refreshable-session".to_string()); + key.encrypted_auth_config = Some(sealed_kiro_auth_config()); key.oauth_invalid_at_unix_secs = Some(1_710_000_000); key.oauth_invalid_reason = Some("账户已封禁: account banned".to_string()); key @@ -2475,7 +2508,8 @@ async fn skips_refreshable_kiro_candidate_when_oauth_marker_is_account_block() { provider_catalog, quotas, request_candidates, - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let (selected, skipped) = collect_selectable_candidates_with_skip_reasons( diff --git a/apps/aether-gateway/src/scheduler/config.rs b/apps/aether-gateway/src/scheduler/config.rs index 1d566c478..0b43c38a7 100644 --- a/apps/aether-gateway/src/scheduler/config.rs +++ b/apps/aether-gateway/src/scheduler/config.rs @@ -1,4 +1,10 @@ +use aether_data_contracts::repository::routing_profiles::RoutingGroupLookupKey; +use aether_routing_core::{ + ResolvedRoutingPolicy, RoutingDefaultPolicy, RoutingSchedulingMode, RoutingSetPriorityMode, + DEFAULT_STICKY_KEY_ATTEMPTS, +}; use aether_scheduler_core::SchedulerPriorityMode; +use tracing::warn; use crate::{AppState, GatewayError}; @@ -10,11 +16,30 @@ pub(crate) enum SchedulerSchedulingMode { LoadBalance, } +impl SchedulerSchedulingMode { + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::FixedOrder => "fixed_order", + Self::CacheAffinity => "cache_affinity", + Self::LoadBalance => "load_balance", + } + } +} + +pub(crate) fn scheduler_priority_mode_as_str(mode: SchedulerPriorityMode) -> &'static str { + match mode { + SchedulerPriorityMode::Provider => "provider", + SchedulerPriorityMode::GlobalKey => "global_key", + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) struct SchedulerOrderingConfig { pub(crate) priority_mode: SchedulerPriorityMode, pub(crate) scheduling_mode: SchedulerSchedulingMode, pub(crate) keep_priority_on_conversion: bool, + /// Total attempts on the first-ranked (sticky) candidate before failover. + pub(crate) sticky_key_attempts: u32, } impl Default for SchedulerOrderingConfig { @@ -23,69 +48,283 @@ impl Default for SchedulerOrderingConfig { priority_mode: SchedulerPriorityMode::Provider, scheduling_mode: SchedulerSchedulingMode::CacheAffinity, keep_priority_on_conversion: false, + sticky_key_attempts: DEFAULT_STICKY_KEY_ATTEMPTS, } } } -pub(crate) fn parse_scheduler_priority_mode( - value: Option<&serde_json::Value>, -) -> SchedulerPriorityMode { - match value - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()) - .as_deref() - { - Some("global_key") => SchedulerPriorityMode::GlobalKey, - _ => SchedulerPriorityMode::Provider, +impl SchedulerOrderingConfig { + /// Ordering config derived from a resolved routing policy. The policy is + /// the single source of truth for request scheduling. + pub(crate) fn from_routing_policy(policy: &ResolvedRoutingPolicy) -> Self { + Self { + priority_mode: scheduler_priority_mode_from_routing(policy.priority_mode), + scheduling_mode: scheduler_scheduling_mode_from_routing(policy.scheduling_mode), + keep_priority_on_conversion: policy.keep_priority_on_conversion, + sticky_key_attempts: policy.sticky_key_attempts, + } + } + + pub(crate) fn from_routing_default_policy(policy: &RoutingDefaultPolicy) -> Self { + Self { + priority_mode: scheduler_priority_mode_from_routing(policy.priority_mode), + scheduling_mode: scheduler_scheduling_mode_from_routing(policy.scheduling_mode), + keep_priority_on_conversion: policy.keep_priority_on_conversion, + sticky_key_attempts: policy.sticky_key_attempts, + } + } + + pub(crate) fn to_routing_default_policy(self) -> RoutingDefaultPolicy { + RoutingDefaultPolicy { + priority_mode: match self.priority_mode { + SchedulerPriorityMode::Provider => RoutingSetPriorityMode::Provider, + SchedulerPriorityMode::GlobalKey => RoutingSetPriorityMode::GlobalKey, + }, + scheduling_mode: match self.scheduling_mode { + SchedulerSchedulingMode::FixedOrder => RoutingSchedulingMode::FixedOrder, + SchedulerSchedulingMode::CacheAffinity => RoutingSchedulingMode::CacheAffinity, + SchedulerSchedulingMode::LoadBalance => RoutingSchedulingMode::LoadBalance, + }, + keep_priority_on_conversion: self.keep_priority_on_conversion, + sticky_key_attempts: self.sticky_key_attempts, + execution_policy: aether_routing_core::RoutingExecutionPolicy::default(), + } + } + + pub(crate) fn priority_mode_str(self) -> &'static str { + scheduler_priority_mode_as_str(self.priority_mode) + } + + pub(crate) fn scheduling_mode_str(self) -> &'static str { + self.scheduling_mode.as_str() } } -pub(crate) fn parse_keep_priority_on_conversion(value: Option<&serde_json::Value>) -> bool { - value.and_then(serde_json::Value::as_bool).unwrap_or(false) -} - -pub(crate) fn parse_scheduler_scheduling_mode( - value: Option<&serde_json::Value>, -) -> SchedulerSchedulingMode { - match value - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()) - .as_deref() - { - Some("fixed_order") => SchedulerSchedulingMode::FixedOrder, - Some("load_balance") => SchedulerSchedulingMode::LoadBalance, - _ => SchedulerSchedulingMode::CacheAffinity, +fn scheduler_priority_mode_from_routing(mode: RoutingSetPriorityMode) -> SchedulerPriorityMode { + match mode { + RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider, + RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey, } } -pub(crate) async fn read_scheduler_ordering_config( +fn scheduler_scheduling_mode_from_routing(mode: RoutingSchedulingMode) -> SchedulerSchedulingMode { + match mode { + RoutingSchedulingMode::FixedOrder => SchedulerSchedulingMode::FixedOrder, + RoutingSchedulingMode::CacheAffinity => SchedulerSchedulingMode::CacheAffinity, + RoutingSchedulingMode::LoadBalance => SchedulerSchedulingMode::LoadBalance, + } +} + +/// Ordering config from the enabled system-default routing group, if any. +pub(crate) async fn read_system_default_routing_ordering_config( state: &AppState, -) -> Result { - let priority_mode = parse_scheduler_priority_mode( - state - .read_system_config_json_value("provider_priority_mode") - .await? - .as_ref(), - ); - let scheduling_mode = parse_scheduler_scheduling_mode( - state - .read_system_config_json_value("scheduling_mode") - .await? - .as_ref(), - ); - let keep_priority_on_conversion = parse_keep_priority_on_conversion( - state - .read_system_config_json_value("keep_priority_on_conversion") - .await? - .as_ref(), - ); - Ok(SchedulerOrderingConfig { - priority_mode, - scheduling_mode, - keep_priority_on_conversion, - }) +) -> Result, GatewayError> { + let Some(group) = state + .find_routing_group(RoutingGroupLookupKey::SystemDefault) + .await? + .filter(|group| group.enabled) + else { + return Ok(None); + }; + let default_policy = match group.config_json.get("default_policy") { + None | Some(serde_json::Value::Null) => RoutingDefaultPolicy::default(), + Some(value) => match serde_json::from_value::(value.clone()) { + Ok(policy) => policy, + Err(error) => { + warn!( + event_name = "scheduler_system_default_routing_policy_invalid", + log_type = "event", + group_id = %group.id, + error = %error, + "system default routing group has an invalid default_policy; ignoring it" + ); + return Ok(None); + } + }, + }; + Ok(Some(SchedulerOrderingConfig::from_routing_default_policy( + &default_policy, + ))) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; + use aether_data_contracts::repository::routing_profiles::{ + CreateRoutingGroupRecord, RoutingGroupLookupKey, RoutingGroupReadRepository, + RoutingGroupWriteRepository, + }; + use serde_json::json; + + use super::*; + use crate::data::GatewayDataState; + + async fn create_system_default( + repository: &InMemoryRoutingGroupRepository, + enabled: bool, + config_json: serde_json::Value, + ) { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: "system-default".to_string(), + name: "system-default".to_string(), + description: None, + enabled, + is_system_default: true, + sort_order: 0, + config_json, + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + + #[tokio::test] + async fn system_default_routing_group_exposes_strategy_ordering() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + create_system_default( + &repository, + true, + json!({ + "default_policy": { + "priority_mode": "provider", + "scheduling_mode": "fixed_order", + "keep_priority_on_conversion": false + } + }), + ) + .await; + let state = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::disabled().with_routing_group_repository_for_tests(repository), + ); + + let config = read_system_default_routing_ordering_config(&state) + .await + .unwrap() + .unwrap(); + + assert_eq!(config.priority_mode, SchedulerPriorityMode::Provider); + assert_eq!(config.scheduling_mode, SchedulerSchedulingMode::FixedOrder); + assert!(!config.keep_priority_on_conversion); + } + + #[tokio::test] + async fn missing_default_policy_in_system_default_group_uses_routing_defaults() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + create_system_default(&repository, true, json!({})).await; + let state = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::disabled().with_routing_group_repository_for_tests(repository), + ); + + let config = read_system_default_routing_ordering_config(&state) + .await + .unwrap() + .unwrap(); + + assert_eq!(config, SchedulerOrderingConfig::default()); + } + + #[tokio::test] + async fn disabled_or_missing_system_default_group_uses_routing_defaults() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + create_system_default( + &repository, + false, + json!({"default_policy": {"scheduling_mode": "fixed_order"}}), + ) + .await; + let with_disabled_group = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::disabled().with_routing_group_repository_for_tests(repository), + ); + let without_repository = AppState::new() + .unwrap() + .with_data_state_for_tests(GatewayDataState::disabled()); + + for state in [with_disabled_group, without_repository] { + let config = read_system_default_routing_ordering_config(&state) + .await + .unwrap(); + assert!(config.is_none()); + } + } + + #[tokio::test] + async fn bootstrap_creates_system_default_group_from_routing_defaults_once() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + let state = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::disabled() + .with_routing_group_repository_for_tests(repository.clone()), + ); + + let created = state + .ensure_system_default_routing_group_inner() + .await + .unwrap() + .expect("first bootstrap should create the system default group"); + assert!(created.enabled); + assert!(created.is_system_default); + assert_eq!( + created.config_json["default_policy"], + json!({ + "priority_mode": "provider", + "scheduling_mode": "cache_affinity", + "keep_priority_on_conversion": false, + "sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS + }) + ); + + let second = state + .ensure_system_default_routing_group_inner() + .await + .unwrap(); + assert!(second.is_none(), "bootstrap must be idempotent"); + assert_eq!( + repository + .find_routing_group(RoutingGroupLookupKey::SystemDefault) + .await + .unwrap() + .map(|group| group.id), + Some(created.id) + ); + + let config = read_system_default_routing_ordering_config(&state) + .await + .unwrap() + .unwrap(); + assert_eq!(config, SchedulerOrderingConfig::default()); + } + + #[tokio::test] + async fn bootstrap_does_not_migrate_legacy_scheduler_keys() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + let state = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::disabled() + .with_system_config_values_for_tests([ + ("provider_priority_mode".to_string(), json!("global_key")), + ("scheduling_mode".to_string(), json!("load_balance")), + ("keep_priority_on_conversion".to_string(), json!(true)), + ]) + .with_routing_group_repository_for_tests(repository), + ); + + let created = state + .ensure_system_default_routing_group_inner() + .await + .unwrap() + .expect("bootstrap should create the strategy"); + assert_eq!( + created.config_json["default_policy"], + json!({ + "priority_mode": "provider", + "scheduling_mode": "cache_affinity", + "keep_priority_on_conversion": false, + "sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS + }) + ); + } } diff --git a/apps/aether-gateway/src/scheduler/state.rs b/apps/aether-gateway/src/scheduler/state.rs index c08297b35..9a3c51c1c 100644 --- a/apps/aether-gateway/src/scheduler/state.rs +++ b/apps/aether-gateway/src/scheduler/state.rs @@ -10,8 +10,6 @@ use async_trait::async_trait; use crate::GatewayError; -use super::config::SchedulerOrderingConfig; - #[async_trait] pub(crate) trait SchedulerRuntimeState { async fn read_provider_quota_snapshot( @@ -60,7 +58,4 @@ pub(crate) trait SchedulerRuntimeState { max_entries: usize, expected_epoch: Option, ) -> bool; - - async fn read_scheduler_ordering_config(&self) - -> Result; } diff --git a/apps/aether-gateway/src/server_chan_push.rs b/apps/aether-gateway/src/server_chan_push.rs index e1acc82e0..ab6c0bf21 100644 --- a/apps/aether-gateway/src/server_chan_push.rs +++ b/apps/aether-gateway/src/server_chan_push.rs @@ -1,8 +1,10 @@ use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, system_config_bool, system_config_string, + decrypt_or_migrate_system_config_secret, system_config_bool, system_config_string, }; use crate::{AppState, GatewayError}; use serde_json::Value; +use std::net::SocketAddr; +use std::time::Duration; pub(crate) const SERVER_CHAN_PUSH_ENABLED_KEY: &str = "module.server_chan_push.enabled"; pub(crate) const SERVER_CHAN_PUSH_SEND_KEY_KEY: &str = "module.server_chan_push.send_key"; @@ -15,14 +17,87 @@ pub(crate) const LEGACY_SERVER_CHAN_TEMPLATE_KEY: &str = "module.important_notification.server_chan_template"; const SERVER_CHAN_API_BASE: &str = "https://sctapi.ftqq.com"; +const SERVER_CHAN_API_HOST: &str = "sctapi.ftqq.com"; +const SERVER_CHAN_API_PORT: u16 = 443; +const SERVER_CHAN_DNS_TIMEOUT: Duration = Duration::from_secs(10); +const MAX_PUSH_RESPONSE_BYTES: usize = 64 * 1024; +const MAX_SERVER_CHAN_SEND_KEY_BYTES: usize = 512; +const MAX_SERVER_CHAN_TEMPLATE_BYTES: usize = 256 * 1024; +const MAX_SERVER_CHAN_TITLE_BYTES: usize = 512; +const MAX_SERVER_CHAN_BODY_BYTES: usize = 2 * 1024 * 1024; +const MAX_SERVER_CHAN_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024; -#[derive(Debug, Clone)] +#[cfg(test)] +fn build_server_chan_client() -> Result { + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(std::time::Duration::from_secs(10)) + .timeout(std::time::Duration::from_secs(300)) + .build() +} + +fn validate_server_chan_resolved_addresses( + addresses: &[SocketAddr], + allow_benchmarking_ip: bool, +) -> Result<(), &'static str> { + if addresses.is_empty() { + return Err("Server Chan API DNS resolution returned no addresses"); + } + if addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(allow_benchmarking_ip && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) { + return Err("Server Chan API resolved to a private or reserved address"); + } + Ok(()) +} + +/// Build the production client with the API hostname pinned to one validated +/// DNS answer. The SendKey is carried in the request path; allowing reqwest +/// to resolve the host again at connect time would let a DNS rebinding or a +/// poisoned resolver redirect that credential to an unintended destination. +async fn build_pinned_server_chan_client() -> Result { + let addresses = aether_http::lookup_host_with_limits( + SERVER_CHAN_API_HOST, + SERVER_CHAN_API_PORT, + SERVER_CHAN_DNS_TIMEOUT, + ) + .await + .map_err(|error| match error.kind() { + std::io::ErrorKind::TimedOut => "Server Chan API DNS resolution timed out", + _ => "Server Chan API DNS resolution failed", + })?; + validate_server_chan_resolved_addresses(&addresses, true)?; + + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(SERVER_CHAN_DNS_TIMEOUT) + .timeout(Duration::from_secs(300)) + .resolve_to_addrs(SERVER_CHAN_API_HOST, &addresses) + .build() + .map_err(|_| "Server Chan HTTP client initialization failed") +} + +#[derive(Clone)] pub(crate) struct ServerChanPushConfig { pub(crate) enabled: bool, pub(crate) send_key: Option, pub(crate) template: Option, } +impl std::fmt::Debug for ServerChanPushConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ServerChanPushConfig") + .field("enabled", &self.enabled) + .field("send_key", &self.send_key.as_ref().map(|_| "[REDACTED]")) + .field("template", &self.template.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + pub(crate) async fn server_chan_push_module_enabled( state: &AppState, ) -> Result { @@ -49,16 +124,7 @@ pub(crate) async fn read_server_chan_push_config( state: &AppState, ) -> Result { let enabled = server_chan_push_module_enabled(state).await?; - let send_key = read_server_chan_value( - state, - SERVER_CHAN_PUSH_SEND_KEY_KEY, - LEGACY_SERVER_CHAN_SEND_KEY_KEY, - ) - .await? - .and_then(|value| system_config_string(Some(&value))) - .map(|value| { - decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value) - }); + let send_key = read_server_chan_secret(state).await?; let template = read_server_chan_value( state, SERVER_CHAN_PUSH_TEMPLATE_KEY, @@ -66,6 +132,9 @@ pub(crate) async fn read_server_chan_push_config( ) .await? .and_then(|value| system_config_string(Some(&value))); + if let Some(template) = template.as_deref() { + validate_server_chan_field("template", template, MAX_SERVER_CHAN_TEMPLATE_BYTES)?; + } Ok(ServerChanPushConfig { enabled, @@ -74,6 +143,24 @@ pub(crate) async fn read_server_chan_push_config( }) } +async fn read_server_chan_secret(state: &AppState) -> Result, GatewayError> { + for key in [ + SERVER_CHAN_PUSH_SEND_KEY_KEY, + LEGACY_SERVER_CHAN_SEND_KEY_KEY, + ] { + let Some(value) = state.read_system_config_json_value(key).await? else { + continue; + }; + let Some(value) = system_config_string(Some(&value)) else { + return Ok(None); + }; + let value = decrypt_or_migrate_system_config_secret(state, key, value).await?; + validate_server_chan_field("send_key", &value, MAX_SERVER_CHAN_SEND_KEY_BYTES)?; + return Ok(Some(value)); + } + Ok(None) +} + async fn read_server_chan_value( state: &AppState, canonical_key: &str, @@ -87,7 +174,7 @@ async fn read_server_chan_value( } pub(crate) async fn send_server_chan_push( - state: &AppState, + _state: &AppState, config: &ServerChanPushConfig, title: &str, markdown_body: &str, @@ -103,73 +190,391 @@ pub(crate) async fn send_server_chan_push( "Server 酱 SendKey 不能为空".to_string(), )); } - let desp = render_server_chan_desp(config.template.as_deref(), title, markdown_body); + if !is_safe_server_chan_send_key(send_key) { + return Err(GatewayError::Internal( + "Server 酱 SendKey 格式无效".to_string(), + )); + } + validate_server_chan_field("send_key", send_key, MAX_SERVER_CHAN_SEND_KEY_BYTES)?; + validate_server_chan_field("title", title, MAX_SERVER_CHAN_TITLE_BYTES)?; + validate_server_chan_field("body", markdown_body, MAX_SERVER_CHAN_BODY_BYTES)?; + let desp = render_server_chan_desp(config.template.as_deref(), title, markdown_body)?; let url = format!("{SERVER_CHAN_API_BASE}/{send_key}.send"); - let response = state - .client + let client = build_pinned_server_chan_client() + .await + .map_err(|message| GatewayError::Internal(message.to_string()))?; + let response = client .post(url) .form(&[("title", title), ("desp", desp.as_str())]) .send() .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + .map_err(|err| GatewayError::Internal(server_chan_request_error_message(&err)))?; let status = response.status(); - let text = response - .text() + let body = aether_http::read_response_bytes_with_limit(response, MAX_PUSH_RESPONSE_BYTES) .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; - if !status.is_success() { - return Err(GatewayError::Internal(format!( - "Server 酱返回 HTTP {status}: {text}" - ))); - } - if let Ok(payload) = serde_json::from_str::(&text) { - let code_is_ok = payload - .get("code") - .and_then(|value| { - value - .as_i64() - .map(|code| code == 0) - .or_else(|| value.as_str().map(|code| code.trim() == "0")) - }) - .unwrap_or(true); - if !code_is_ok { - return Err(GatewayError::Internal(format!( - "Server 酱返回失败: {payload}" - ))); - } + .map_err(|err| GatewayError::Internal(server_chan_response_body_error_message(&err)))?; + let text = String::from_utf8_lossy(&body); + if let Some(message) = server_chan_response_failure_message(status, &text) { + return Err(GatewayError::Internal(message)); } Ok(()) } -fn render_server_chan_desp(template: Option<&str>, title: &str, markdown_body: &str) -> String { - match template { - Some(template) if !template.trim().is_empty() => template - .replace("{title}", title) - .replace("{body}", markdown_body), - _ => markdown_body.to_string(), +fn is_safe_server_chan_send_key(send_key: &str) -> bool { + !send_key.is_empty() + && send_key.len() <= MAX_SERVER_CHAN_SEND_KEY_BYTES + && send_key + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) +} + +fn server_chan_response_failure_message(status: http::StatusCode, text: &str) -> Option { + if !status.is_success() { + return Some(format!("Server Chan returned HTTP {status}")); } + let payload = serde_json::from_str::(text).ok()?; + let code_is_ok = payload + .get("code") + .and_then(|value| { + value + .as_i64() + .map(|code| code == 0) + .or_else(|| value.as_str().map(|code| code.trim() == "0")) + }) + .unwrap_or(true); + (!code_is_ok).then(|| "Server Chan returned failure".to_string()) +} + +fn server_chan_request_error_message(error: &reqwest::Error) -> String { + // Reqwest errors may include the request URL. ServerChan authenticates with + // the SendKey in that URL's path, so forwarding the source error would leak + // the credential into logs and notification test responses. + format!("Server Chan request failed ({})", reqwest_error_kind(error)) +} + +fn server_chan_response_body_error_message(error: &aether_http::ResponseBodyReadError) -> String { + match error { + aether_http::ResponseBodyReadError::TooLarge { max_bytes } => { + format!("Server Chan response exceeds {max_bytes} bytes") + } + aether_http::ResponseBodyReadError::Read(error) => format!( + "Server Chan response read failed ({})", + reqwest_error_kind(error) + ), + } +} + +fn reqwest_error_kind(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else { + "transport" + } +} + +fn validate_server_chan_field( + field: &str, + value: &str, + max_bytes: usize, +) -> Result<(), GatewayError> { + if value.len() > max_bytes || value.bytes().any(|byte| byte == 0) { + return Err(GatewayError::Internal(format!( + "Server Chan {field} exceeds the allowed size or contains a NUL byte" + ))); + } + Ok(()) +} + +fn render_server_chan_desp( + template: Option<&str>, + title: &str, + markdown_body: &str, +) -> Result { + validate_server_chan_field("title", title, MAX_SERVER_CHAN_TITLE_BYTES)?; + validate_server_chan_field("body", markdown_body, MAX_SERVER_CHAN_BODY_BYTES)?; + let template = template + .filter(|value| !value.trim().is_empty()) + .unwrap_or("{body}"); + validate_server_chan_field("template", template, MAX_SERVER_CHAN_TEMPLATE_BYTES)?; + + let mut rendered = String::with_capacity( + template + .len() + .saturating_add(title.len()) + .saturating_add(markdown_body.len()) + .min(MAX_SERVER_CHAN_RENDERED_BODY_BYTES), + ); + let mut cursor = 0usize; + while cursor < template.len() { + let remaining = &template[cursor..]; + let title_match = remaining.find("{title}"); + let body_match = remaining.find("{body}"); + let next = match (title_match, body_match) { + (None, None) => { + append_server_chan_rendered_part(&mut rendered, remaining)?; + cursor = template.len(); + continue; + } + (Some(index), None) => (index, "{title}", title), + (None, Some(index)) => (index, "{body}", markdown_body), + (Some(title_index), Some(body_index)) if title_index <= body_index => { + (title_index, "{title}", title) + } + (Some(_), Some(body_index)) => (body_index, "{body}", markdown_body), + }; + append_server_chan_rendered_part(&mut rendered, &remaining[..next.0])?; + append_server_chan_rendered_part(&mut rendered, next.2)?; + cursor += next.0 + next.1.len(); + } + Ok(rendered) +} + +fn append_server_chan_rendered_part(output: &mut String, part: &str) -> Result<(), GatewayError> { + let next_len = output.len().checked_add(part.len()).ok_or_else(|| { + GatewayError::Internal("Server Chan rendered body is too large".to_string()) + })?; + if next_len > MAX_SERVER_CHAN_RENDERED_BODY_BYTES { + return Err(GatewayError::Internal( + "Server Chan rendered body exceeds the allowed size".to_string(), + )); + } + output.push_str(part); + Ok(()) } #[cfg(test)] mod tests { - use super::render_server_chan_desp; + use super::{ + build_server_chan_client, is_safe_server_chan_send_key, read_server_chan_push_config, + render_server_chan_desp, server_chan_request_error_message, + server_chan_response_body_error_message, server_chan_response_failure_message, + validate_server_chan_resolved_addresses, LEGACY_SERVER_CHAN_SEND_KEY_KEY, + }; + use crate::data::GatewayDataState; + use crate::handlers::shared::decrypt_system_config_secret; + use crate::AppState; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use axum::{ + http::{header, StatusCode}, + response::IntoResponse, + routing::post, + Router, + }; + use std::net::SocketAddr; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + + #[tokio::test] + async fn legacy_send_key_is_migrated_at_its_original_config_key() { + let plaintext = "SCT-legacy-plaintext-send-key"; + let data = GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests([( + LEGACY_SERVER_CHAN_SEND_KEY_KEY.to_string(), + serde_json::json!(plaintext), + )]); + let mut state = AppState::new().expect("gateway state should build"); + state.replace_data_state(Arc::new(data)); + + let config = read_server_chan_push_config(&state) + .await + .expect("legacy config should read"); + assert_eq!(config.send_key.as_deref(), Some(plaintext)); + + let stored = state + .read_system_config_json_value_strong(LEGACY_SERVER_CHAN_SEND_KEY_KEY) + .await + .expect("migrated legacy key should read") + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .expect("migrated legacy key should remain a string"); + assert_ne!(stored, plaintext); + assert_eq!( + decrypt_system_config_secret(&state, LEGACY_SERVER_CHAN_SEND_KEY_KEY, &stored) + .expect("migrated legacy key should decrypt"), + plaintext + ); + } + + #[tokio::test] + async fn server_chan_client_never_forwards_send_key_across_redirects() { + let redirected_hits = Arc::new(AtomicUsize::new(0)); + let redirected_hits_for_route = Arc::clone(&redirected_hits); + let redirected_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect target listener"); + let redirected_addr = redirected_listener + .local_addr() + .expect("redirect target addr"); + let redirected_app = Router::new().route( + "/capture", + post(move || { + let hits = Arc::clone(&redirected_hits_for_route); + async move { + hits.fetch_add(1, Ordering::SeqCst); + StatusCode::OK + } + }), + ); + let redirected_server = tokio::spawn(async move { + axum::serve(redirected_listener, redirected_app) + .await + .expect("redirect target server"); + }); + + let source_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect source listener"); + let source_addr = source_listener.local_addr().expect("redirect source addr"); + let location = format!("http://{redirected_addr}/capture"); + let source_app = Router::new().route( + "/SCT-secret-send-key.send", + post(move || { + let location = location.clone(); + async move { (StatusCode::FOUND, [(header::LOCATION, location)]).into_response() } + }), + ); + let source_server = tokio::spawn(async move { + axum::serve(source_listener, source_app) + .await + .expect("redirect source server"); + }); + + let response = build_server_chan_client() + .expect("client") + .post(format!("http://{source_addr}/SCT-secret-send-key.send")) + .send() + .await + .expect("redirect response"); + assert_eq!(response.status(), StatusCode::FOUND); + assert_eq!(redirected_hits.load(Ordering::SeqCst), 0); + + source_server.abort(); + redirected_server.abort(); + } + + #[tokio::test] + async fn server_chan_transport_error_does_not_expose_send_key_url() { + let send_key = "SCT-secret-send-key"; + let error = reqwest::Client::new() + .post(format!("ftp://sctapi.ftqq.com/{send_key}.send")) + .send() + .await + .expect_err("unsupported URL scheme should fail before network I/O"); + + let message = server_chan_request_error_message(&error); + assert!(message.starts_with("Server Chan request failed (")); + assert!(!message.contains(send_key)); + assert!(!message.contains("sctapi.ftqq.com")); + assert!(!message.contains(".send")); + + let body_error = aether_http::ResponseBodyReadError::Read(error); + let message = server_chan_response_body_error_message(&body_error); + assert!(message.starts_with("Server Chan response read failed (")); + assert!(!message.contains(send_key)); + assert!(!message.contains("sctapi.ftqq.com")); + assert!(!message.contains(".send")); + } + + #[test] + fn server_chan_dns_answers_must_be_public() { + assert!(validate_server_chan_resolved_addresses( + &[SocketAddr::from(([1, 1, 1, 1], 443))], + false, + ) + .is_ok()); + for address in [ + SocketAddr::from(([127, 0, 0, 1], 443)), + SocketAddr::from(([10, 0, 0, 1], 443)), + SocketAddr::from(([169, 254, 169, 254], 443)), + ] { + assert!( + validate_server_chan_resolved_addresses(&[address], false).is_err(), + "private Server Chan DNS answer should be rejected: {address}" + ); + } + assert!(validate_server_chan_resolved_addresses(&[], false).is_err()); + } + + #[test] + fn server_chan_dns_allows_benchmarking_ip_for_builtin_host() { + let fake = SocketAddr::from(([198, 18, 75, 234], 443)); + assert!(validate_server_chan_resolved_addresses(&[fake], true).is_ok()); + assert!(validate_server_chan_resolved_addresses( + &[fake, SocketAddr::from(([127, 0, 0, 1], 443))], + true, + ) + .is_err()); + assert!(validate_server_chan_resolved_addresses(&[fake], false).is_err()); + } + + #[test] + fn server_chan_failure_response_does_not_expose_arbitrary_body() { + let secret_body = "upstream echoed SCT-secret-send-key and internal details"; + let message = + server_chan_response_failure_message(http::StatusCode::BAD_GATEWAY, secret_body) + .expect("non-success response should fail"); + assert_eq!(message, "Server Chan returned HTTP 502 Bad Gateway"); + assert!(!message.contains(secret_body)); + assert!(!message.contains("SCT-secret-send-key")); + + let message = server_chan_response_failure_message( + http::StatusCode::OK, + r#"{"code":"SCT-secret-send-key","message":"internal details"}"#, + ) + .expect("non-zero business response should fail"); + assert_eq!(message, "Server Chan returned failure"); + } + + #[test] + fn server_chan_send_key_cannot_escape_the_request_path() { + assert!(is_safe_server_chan_send_key("SCT123abc-._")); + for unsafe_key in [ + "SCT/other", + "SCT?token=secret", + "SCT#fragment", + "SCT\\other", + "SCT key", + "SCT\r\nX-Injected: yes", + ] { + assert!( + !is_safe_server_chan_send_key(unsafe_key), + "unsafe key: {unsafe_key:?}" + ); + } + } #[test] fn server_chan_desp_uses_template_when_provided() { let rendered = - render_server_chan_desp(Some("**{title}**\n\n{body}\n\n--end--"), "告警", "原始正文"); + render_server_chan_desp(Some("**{title}**\n\n{body}\n\n--end--"), "告警", "原始正文") + .expect("template should render"); assert_eq!(rendered, "**告警**\n\n原始正文\n\n--end--"); } #[test] fn server_chan_desp_falls_back_to_markdown_body_for_empty_template() { assert_eq!( - render_server_chan_desp(None, "告警", "原始正文"), + render_server_chan_desp(None, "告警", "原始正文").expect("fallback should render"), "原始正文" ); assert_eq!( - render_server_chan_desp(Some(" "), "告警", "原始正文"), + render_server_chan_desp(Some(" "), "告警", "原始正文") + .expect("fallback should render"), "原始正文" ); } + + #[test] + fn server_chan_desp_rejects_expansion_bombs_and_oversized_content() { + let template = "{body}".repeat(super::MAX_SERVER_CHAN_TEMPLATE_BYTES / 6 + 1); + assert!(render_server_chan_desp(Some(&template), "告警", "正文").is_err()); + let body = "x".repeat(super::MAX_SERVER_CHAN_BODY_BYTES + 1); + assert!(render_server_chan_desp(None, "告警", &body).is_err()); + } } diff --git a/apps/aether-gateway/src/state/admin_types.rs b/apps/aether-gateway/src/state/admin_types.rs index be06e404e..d8eb3f4b6 100644 --- a/apps/aether-gateway/src/state/admin_types.rs +++ b/apps/aether-gateway/src/state/admin_types.rs @@ -6,6 +6,7 @@ pub(crate) use aether_data::repository::wallet::{ pub(crate) use aether_data_contracts::repository::billing::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, + BillingPlanRecord, BillingPlanWriteInput, PaymentGatewayConfigCasWriteInput, + PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, + UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; diff --git a/apps/aether-gateway/src/state/app.rs b/apps/aether-gateway/src/state/app.rs index 1b816563b..3a8a63ce8 100644 --- a/apps/aether-gateway/src/state/app.rs +++ b/apps/aether-gateway/src/state/app.rs @@ -34,7 +34,6 @@ use super::{ ProviderTransportSnapshotFlight, }; -const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000; const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000; const MAX_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 600_000; const REQUEST_BODY_READ_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS"; @@ -96,7 +95,7 @@ impl std::fmt::Debug for TestExecutionRuntimeSyncOverride { #[derive(Debug, Clone)] pub(crate) struct FrontdoorRuntimeGuardConfig { - pub(crate) request_body_read_timeout: Duration, + pub(crate) request_body_read_timeout: Option, pub(crate) request_body_buffer_budget_bytes: usize, pub(crate) request_body_buffer_budget_permits: usize, pub(crate) local_execution_planning_timeout: Duration, @@ -113,9 +112,8 @@ pub(crate) const METRIC_SNAPSHOT_TTL: Duration = Duration::from_secs(2); impl FrontdoorRuntimeGuardConfig { pub(crate) fn from_env() -> Self { Self { - request_body_read_timeout: env_duration_ms( + request_body_read_timeout: optional_env_duration_ms( REQUEST_BODY_READ_TIMEOUT_MS_ENV, - DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS, MIN_REQUEST_BODY_READ_TIMEOUT_MS, MAX_REQUEST_BODY_READ_TIMEOUT_MS, ), @@ -148,7 +146,7 @@ impl FrontdoorRuntimeGuardConfig { #[cfg(test)] pub(crate) fn for_tests( - request_body_read_timeout: Duration, + request_body_read_timeout: Option, local_execution_planning_timeout: Duration, ) -> Self { Self { @@ -189,6 +187,19 @@ fn request_body_buffer_budget_permits_from_env() -> usize { / REQUEST_BODY_BUFFER_PERMIT_BYTES } +fn optional_env_duration_ms(key: &str, min_ms: u64, max_ms: u64) -> Option { + let raw = std::env::var(key).ok(); + parse_optional_duration_ms(raw.as_deref(), min_ms, max_ms) +} + +fn parse_optional_duration_ms(raw: Option<&str>, min_ms: u64, max_ms: u64) -> Option { + let parsed = raw?.trim().parse::().ok()?; + if parsed == 0 { + return None; + } + Some(Duration::from_millis(parsed.clamp(min_ms, max_ms))) +} + fn env_duration_ms(key: &str, default_ms: u64, min_ms: u64, max_ms: u64) -> Duration { let ms = std::env::var(key) .ok() @@ -371,6 +382,7 @@ pub struct AppState { pub(crate) background_data: Arc, pub(crate) background_data_isolated: bool, pub(crate) runtime_state: Arc, + pub(crate) internal_gateway_auth: Arc, pub(crate) usage_runtime: Arc, pub(crate) video_tasks: Arc, pub(crate) video_task_poller: Option, @@ -397,6 +409,8 @@ pub struct AppState { pub(crate) auth_api_key_feature_settings_cache: Arc>, pub(crate) auth_daily_quota_availability_cache: Arc>, + pub(crate) auth_plan_usage_policy_cache: + Arc>, pub(crate) auth_wallet_snapshot_cache: Arc>, pub(crate) auth_request_cost_upper_bound_cache: Arc>, @@ -518,6 +532,61 @@ mod tests { fd_soft_limit: 1_048_576, }; + #[test] + fn request_body_read_timeout_parser_defaults_to_disabled() { + assert_eq!( + parse_optional_duration_ms( + None, + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + None + ); + } + + #[test] + fn request_body_read_timeout_parser_disables_zero_and_invalid_values() { + for value in ["", "invalid", "-1", "0", " 0 "] { + assert_eq!( + parse_optional_duration_ms( + Some(value), + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + None, + "{value:?} should disable the optional timeout" + ); + } + } + + #[test] + fn request_body_read_timeout_parser_clamps_nonzero_values() { + assert_eq!( + parse_optional_duration_ms( + Some("1"), + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + Some(Duration::from_millis(MIN_REQUEST_BODY_READ_TIMEOUT_MS)) + ); + assert_eq!( + parse_optional_duration_ms( + Some("120000"), + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + Some(Duration::from_millis(120_000)) + ); + assert_eq!( + parse_optional_duration_ms( + Some("900000"), + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + Some(Duration::from_millis(MAX_REQUEST_BODY_READ_TIMEOUT_MS)) + ); + } + #[test] fn gate_limit_parser_defaults_to_auto() { assert_eq!( diff --git a/apps/aether-gateway/src/state/bootstrap_admin.rs b/apps/aether-gateway/src/state/bootstrap_admin.rs index 03db6e899..e59185a2d 100644 --- a/apps/aether-gateway/src/state/bootstrap_admin.rs +++ b/apps/aether-gateway/src/state/bootstrap_admin.rs @@ -7,13 +7,24 @@ const BOOTSTRAP_ADMIN_EMAIL_ENVS: &[&str] = &["ADMIN_EMAIL"]; const BOOTSTRAP_ADMIN_USERNAME_ENVS: &[&str] = &["ADMIN_USERNAME"]; const BOOTSTRAP_ADMIN_PASSWORD_ENVS: &[&str] = &["ADMIN_PASSWORD"]; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] struct BootstrapAdminConfig { email: Option, username: String, password: String, } +impl std::fmt::Debug for BootstrapAdminConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("BootstrapAdminConfig") + .field("email", &self.email) + .field("username", &self.username) + .field("password", &"[REDACTED]") + .finish() + } +} + impl BootstrapAdminConfig { fn from_env() -> Result, GatewayError> { Self::from_lookup(|key| { @@ -495,6 +506,14 @@ mod tests { ); } + #[test] + fn bootstrap_admin_config_debug_output_redacts_password() { + let config = bootstrap_config(); + let debug = format!("{config:?}"); + assert!(debug.contains("[REDACTED]")); + assert!(!debug.contains("Secret123!")); + } + #[test] fn bootstrap_admin_config_rejects_partial_env() { let vars = diff --git a/apps/aether-gateway/src/state/catalog.rs b/apps/aether-gateway/src/state/catalog.rs index 5b3ee60ec..025329c9b 100644 --- a/apps/aether-gateway/src/state/catalog.rs +++ b/apps/aether-gateway/src/state/catalog.rs @@ -37,20 +37,24 @@ impl AppState { &self, active_only: bool, ) -> Result, GatewayError> { - self.data + let providers = self + .data .list_provider_catalog_providers(active_only) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_providers(providers).await } pub(crate) async fn list_provider_catalog_endpoints_by_provider_ids( &self, provider_ids: &[String], ) -> Result, GatewayError> { - self.data + let endpoints = self + .data .list_provider_catalog_endpoints_by_provider_ids(provider_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_endpoints(endpoints).await } pub(crate) async fn list_public_global_models( @@ -128,6 +132,20 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn update_management_token_for_user( + &self, + record: &aether_data::repository::management_tokens::UpdateManagementTokenRecord, + user_id: &str, + ) -> Result< + LocalMutationOutcome, + GatewayError, + > { + self.data + .update_management_token_for_user(record, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn delete_management_token( &self, token_id: &str, @@ -138,6 +156,17 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + self.data + .delete_management_token_for_user(token_id, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn record_management_token_usage( &self, token_id: &str, @@ -166,6 +195,41 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result< + Option, + GatewayError, + > { + self.data + .set_management_token_active_for_user(token_id, user_id, is_active) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn activate_management_token_if_matches( + &self, + mutation: &aether_data::repository::management_tokens::ActivateManagementTokenIfMatches, + ) -> Result { + self.data + .activate_management_token_if_matches(mutation) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn delete_inactive_management_token_if_matches( + &self, + mutation: &aether_data::repository::management_tokens::ActivateManagementTokenIfMatches, + ) -> Result { + self.data + .delete_inactive_management_token_if_matches(mutation) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn regenerate_management_token_secret( &self, mutation: &aether_data::repository::management_tokens::RegenerateManagementTokenSecret, @@ -179,6 +243,20 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn regenerate_management_token_secret_for_user( + &self, + mutation: &aether_data::repository::management_tokens::RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result< + LocalMutationOutcome, + GatewayError, + > { + self.data + .regenerate_management_token_secret_for_user(mutation, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn get_public_global_model_by_name( &self, model_name: &str, @@ -443,20 +521,24 @@ impl AppState { &self, provider_ids: &[String], ) -> Result, GatewayError> { - self.data + let keys = self + .data .list_provider_catalog_keys_by_provider_ids(provider_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_keys(keys).await } pub(crate) async fn list_provider_catalog_key_summaries_by_provider_ids( &self, provider_ids: &[String], ) -> Result, GatewayError> { - self.data + let keys = self + .data .list_provider_catalog_key_summaries_by_provider_ids(provider_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_keys(keys).await } pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids( @@ -474,30 +556,37 @@ impl AppState { &self, key_ids: &[String], ) -> Result, GatewayError> { - self.data + let keys = self + .data .list_provider_catalog_keys_by_ids(key_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_keys(keys).await } pub(crate) async fn list_provider_catalog_keys_by_ids_strong( &self, key_ids: &[String], ) -> Result, GatewayError> { - self.data + let keys = self + .data .list_provider_catalog_keys_by_ids_strong(key_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_keys(keys).await } pub(crate) async fn list_provider_catalog_key_page( &self, query: &provider_catalog::ProviderCatalogKeyListQuery, ) -> Result { - self.data + let mut page = self + .data .list_provider_catalog_key_page(query) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + page.items = self.open_provider_catalog_keys(page.items).await?; + Ok(page) } pub(crate) async fn list_provider_catalog_key_stats_by_provider_ids( @@ -514,15 +603,19 @@ impl AppState { &self, key: &provider_catalog::StoredProviderCatalogKey, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_key(key)?; let created = self .data - .create_provider_catalog_key(key) + .create_provider_catalog_key(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if created.is_some() { self.invalidate_provider_routing_caches(); } - Ok(created) + match created { + Some(key) => self.open_provider_catalog_key(key).await.map(Some), + None => Ok(None), + } } pub(crate) async fn create_provider_catalog_provider( @@ -530,29 +623,58 @@ impl AppState { provider: &provider_catalog::StoredProviderCatalogProvider, shift_existing_priorities_from: Option, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_provider(provider)?; let created = self .data - .create_provider_catalog_provider(provider, shift_existing_priorities_from) + .create_provider_catalog_provider(&protected, shift_existing_priorities_from) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if created.is_some() { self.invalidate_provider_routing_caches(); } - Ok(created) + match created { + Some(provider) => self + .open_provider_catalog_provider(provider) + .await + .map(Some), + None => Ok(None), + } } pub(crate) async fn update_provider_catalog_provider( &self, provider: &provider_catalog::StoredProviderCatalogProvider, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_provider(provider)?; let updated = self .data - .update_provider_catalog_provider(provider) + .update_provider_catalog_provider(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if updated.is_some() { self.invalidate_provider_routing_caches(); } + match updated { + Some(provider) => self + .open_provider_catalog_provider(provider) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn compare_and_swap_provider_catalog_provider_config( + &self, + update: &provider_catalog::ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + let updated = self + .data + .compare_and_swap_provider_catalog_provider_config(update) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if updated { + self.invalidate_provider_routing_caches(); + } Ok(updated) } @@ -616,30 +738,44 @@ impl AppState { &self, endpoint: &provider_catalog::StoredProviderCatalogEndpoint, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_endpoint(endpoint)?; let created = self .data - .create_provider_catalog_endpoint(endpoint) + .create_provider_catalog_endpoint(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if created.is_some() { self.invalidate_provider_routing_caches(); } - Ok(created) + match created { + Some(endpoint) => self + .open_provider_catalog_endpoint(endpoint) + .await + .map(Some), + None => Ok(None), + } } pub(crate) async fn update_provider_catalog_endpoint( &self, endpoint: &provider_catalog::StoredProviderCatalogEndpoint, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_endpoint(endpoint)?; let updated = self .data - .update_provider_catalog_endpoint(endpoint) + .update_provider_catalog_endpoint(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if updated.is_some() { self.invalidate_provider_routing_caches(); } - Ok(updated) + match updated { + Some(endpoint) => self + .open_provider_catalog_endpoint(endpoint) + .await + .map(Some), + None => Ok(None), + } } pub(crate) async fn delete_provider_catalog_endpoint( @@ -661,24 +797,30 @@ impl AppState { &self, key: &provider_catalog::StoredProviderCatalogKey, ) -> Result, GatewayError> { + let protected = self.protect_provider_catalog_key(key)?; let updated = self .data - .update_provider_catalog_key(key) + .update_provider_catalog_key(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if updated.is_some() { self.invalidate_provider_routing_caches(); } - Ok(updated) + match updated { + Some(key) => self.open_provider_catalog_key(key).await.map(Some), + None => Ok(None), + } } pub(crate) async fn compare_and_update_provider_catalog_key_admin_state( &self, update: &provider_catalog::ProviderCatalogKeyAdminCasUpdate, ) -> Result { + let mut protected = update.clone(); + protected.key = self.protect_provider_catalog_key(&update.key)?; let updated = self .data - .compare_and_update_provider_catalog_key_admin_state(update) + .compare_and_update_provider_catalog_key_admin_state(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; // A conflict means another instance changed credentials. Invalidate on @@ -691,15 +833,22 @@ impl AppState { &self, keys: &[provider_catalog::StoredProviderCatalogKey], ) -> Result>, GatewayError> { + let protected = keys + .iter() + .map(|key| self.protect_provider_catalog_key(key)) + .collect::, _>>()?; let updated = self .data - .update_provider_catalog_keys(keys) + .update_provider_catalog_keys(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if updated.as_ref().is_some_and(|keys| !keys.is_empty()) { self.invalidate_provider_routing_caches(); } - Ok(updated) + match updated { + Some(keys) => self.open_provider_catalog_keys(keys).await.map(Some), + None => Ok(None), + } } pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state( @@ -1041,30 +1190,36 @@ impl AppState { &self, provider_ids: &[String], ) -> Result, GatewayError> { - self.data + let providers = self + .data .list_provider_catalog_providers_by_ids(provider_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_providers(providers).await } pub(crate) async fn read_provider_catalog_endpoints_by_ids( &self, endpoint_ids: &[String], ) -> Result, GatewayError> { - self.data + let endpoints = self + .data .list_provider_catalog_endpoints_by_ids(endpoint_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_endpoints(endpoints).await } pub(crate) async fn read_provider_catalog_keys_by_ids( &self, key_ids: &[String], ) -> Result, GatewayError> { - self.data + let keys = self + .data .list_provider_catalog_keys_by_ids(key_ids) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.open_provider_catalog_keys(keys).await } pub(crate) async fn update_provider_catalog_key_format_health( diff --git a/apps/aether-gateway/src/state/catalog_credentials.rs b/apps/aether-gateway/src/state/catalog_credentials.rs new file mode 100644 index 000000000..462d901a5 --- /dev/null +++ b/apps/aether-gateway/src/state/catalog_credentials.rs @@ -0,0 +1,395 @@ +use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogKeyCredentialsCasUpdate, StoredProviderCatalogKey, +}; + +use super::AppState; +use crate::handlers::shared::{ + open_provider_catalog_credential, seal_provider_catalog_credential, + ProviderCatalogCredentialField, ProviderCatalogCredentialProjection, +}; +use crate::GatewayError; + +impl AppState { + pub(super) fn protect_provider_catalog_key_credentials( + &self, + key: &StoredProviderCatalogKey, + ) -> Result { + let mut protected = key.clone(); + protected.encrypted_api_key = self + .project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::ApiKey, + key.encrypted_api_key.as_deref(), + )? + .map(|projection| projection.protected); + protected.encrypted_auth_config = self + .project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::AuthConfig, + key.encrypted_auth_config.as_deref(), + )? + .map(|projection| projection.protected); + Ok(protected) + } + + pub(super) async fn open_provider_catalog_key_credentials_once( + &self, + key: &mut StoredProviderCatalogKey, + ) -> Result { + let observed_api_key = key.encrypted_api_key.clone(); + let observed_auth_config = key.encrypted_auth_config.clone(); + let api_key = self.project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::ApiKey, + observed_api_key.as_deref(), + )?; + let auth_config = self.project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::AuthConfig, + observed_auth_config.as_deref(), + )?; + let migration_required = api_key + .as_ref() + .is_some_and(|projection| projection.migration_required) + || auth_config + .as_ref() + .is_some_and(|projection| projection.migration_required); + if !migration_required { + return Ok(true); + } + if !self.has_provider_catalog_data_writer() { + return Err(provider_catalog_credential_error( + "stored provider catalog credentials require migration but the catalog writer is unavailable", + )); + } + + let protected_api_key = api_key.map(|projection| projection.protected); + let protected_auth_config = auth_config.map(|projection| projection.protected); + let updated = self + .data + .compare_and_swap_provider_catalog_key_credentials( + &ProviderCatalogKeyCredentialsCasUpdate { + key_id: key.id.clone(), + expected_provider_id: key.provider_id.clone(), + expected_encrypted_api_key: observed_api_key, + expected_encrypted_auth_config: observed_auth_config, + encrypted_api_key: protected_api_key.clone(), + encrypted_auth_config: protected_auth_config.clone(), + }, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if updated { + key.encrypted_api_key = protected_api_key; + key.encrypted_auth_config = protected_auth_config; + } + Ok(updated) + } + + pub(crate) fn decrypt_provider_catalog_key_api_key( + &self, + key: &StoredProviderCatalogKey, + ) -> Result, GatewayError> { + self.project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::ApiKey, + key.encrypted_api_key.as_deref(), + ) + .map(|projection| projection.map(|projection| projection.plaintext)) + } + + pub(crate) fn decrypt_provider_catalog_key_auth_config( + &self, + key: &StoredProviderCatalogKey, + ) -> Result, GatewayError> { + self.project_provider_catalog_key_credential( + key, + ProviderCatalogCredentialField::AuthConfig, + key.encrypted_auth_config.as_deref(), + ) + .map(|projection| projection.map(|projection| projection.plaintext)) + } + + pub(crate) fn seal_provider_catalog_key_api_key( + &self, + provider_id: &str, + key_id: &str, + plaintext: &str, + ) -> Result { + seal_provider_catalog_credential( + self, + provider_id, + key_id, + ProviderCatalogCredentialField::ApiKey, + plaintext, + ) + .map_err(provider_catalog_credential_error) + } + + pub(crate) fn seal_provider_catalog_key_auth_config( + &self, + provider_id: &str, + key_id: &str, + plaintext: &str, + ) -> Result { + seal_provider_catalog_credential( + self, + provider_id, + key_id, + ProviderCatalogCredentialField::AuthConfig, + plaintext, + ) + .map_err(provider_catalog_credential_error) + } + + pub(super) fn validate_protected_provider_catalog_key_api_key( + &self, + provider_id: &str, + key_id: &str, + stored: &str, + ) -> Result<(), GatewayError> { + self.validate_protected_provider_catalog_key_credential( + provider_id, + key_id, + ProviderCatalogCredentialField::ApiKey, + stored, + ) + } + + pub(super) fn validate_protected_provider_catalog_key_auth_config( + &self, + provider_id: &str, + key_id: &str, + stored: &str, + ) -> Result<(), GatewayError> { + self.validate_protected_provider_catalog_key_credential( + provider_id, + key_id, + ProviderCatalogCredentialField::AuthConfig, + stored, + ) + } + + fn validate_protected_provider_catalog_key_credential( + &self, + provider_id: &str, + key_id: &str, + field: ProviderCatalogCredentialField, + stored: &str, + ) -> Result<(), GatewayError> { + let projection = open_provider_catalog_credential(self, provider_id, key_id, field, stored) + .map_err(provider_catalog_credential_error)?; + if projection.migration_required || projection.protected != stored { + return Err(provider_catalog_credential_error( + "provider catalog credential write requires a bound v2 ciphertext", + )); + } + Ok(()) + } + + fn project_provider_catalog_key_credential( + &self, + key: &StoredProviderCatalogKey, + field: ProviderCatalogCredentialField, + stored: Option<&str>, + ) -> Result, GatewayError> { + let Some(stored) = stored else { + return Ok(None); + }; + if stored.is_empty() { + return Err(provider_catalog_credential_error( + "stored provider catalog credential is empty", + )); + } + open_provider_catalog_credential(self, &key.provider_id, &key.id, field, stored) + .map(Some) + .map_err(provider_catalog_credential_error) + } +} + +fn provider_catalog_credential_error(message: &'static str) -> GatewayError { + GatewayError::Internal(message.to_string()) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; + use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogReadRepository, + ProviderCatalogWriteRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider, + }; + + use crate::{data::GatewayDataState, AppState}; + + fn sample_provider(id: &str) -> StoredProviderCatalogProvider { + StoredProviderCatalogProvider::new( + id.to_string(), + format!("Provider {id}"), + Some("https://example.test".to_string()), + "openai".to_string(), + ) + .expect("provider should build") + } + + fn sample_key( + id: &str, + provider_id: &str, + encrypted_api_key: Option, + encrypted_auth_config: Option, + ) -> StoredProviderCatalogKey { + StoredProviderCatalogKey::new( + id.to_string(), + provider_id.to_string(), + format!("Key {id}"), + "oauth".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + None, + encrypted_api_key, + encrypted_auth_config, + None, + None, + None, + None, + None, + None, + ) + .expect("key transport should build") + } + + fn state_with_repository(repository: Arc) -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[tokio::test] + async fn app_state_migrates_both_legacy_fields_with_one_exact_cas() { + let legacy_api = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-api-key") + .expect("legacy API key should encrypt"); + let legacy_auth = encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + r#"{"refresh_token":"legacy-refresh"}"#, + ) + .expect("legacy auth config should encrypt"); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-1")], + Vec::new(), + vec![sample_key( + "key-1", + "provider-1", + Some(legacy_api), + Some(legacy_auth), + )], + )); + let state = state_with_repository(Arc::clone(&repository)); + + let opened = state + .list_provider_catalog_keys_by_ids(&["key-1".to_string()]) + .await + .expect("legacy key should migrate") + .into_iter() + .next() + .expect("key should exist"); + assert_eq!( + state + .decrypt_provider_catalog_key_api_key(&opened) + .expect("API key should open") + .as_deref(), + Some("legacy-api-key") + ); + assert_eq!( + state + .decrypt_provider_catalog_key_auth_config(&opened) + .expect("auth config should open") + .as_deref(), + Some(r#"{"refresh_token":"legacy-refresh"}"#) + ); + + let stored = repository + .list_keys_by_ids(&["key-1".to_string()]) + .await + .expect("stored key should read") + .into_iter() + .next() + .expect("stored key should exist"); + assert!(stored + .encrypted_api_key + .as_deref() + .is_some_and(|value| value.starts_with("aether-provider-catalog-credential-v2:"))); + assert!(stored + .encrypted_auth_config + .as_deref() + .is_some_and(|value| value.starts_with("aether-provider-catalog-credential-v2:"))); + } + + #[tokio::test] + async fn app_state_rejects_ciphertext_copied_to_another_key() { + let empty_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-1")], + Vec::new(), + Vec::new(), + )); + let bootstrap = state_with_repository(Arc::clone(&empty_repository)); + let copied = bootstrap + .seal_provider_catalog_key_api_key("provider-1", "key-1", "secret") + .expect("credential should seal"); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-1")], + Vec::new(), + vec![sample_key("key-2", "provider-1", Some(copied), None)], + )); + let state = state_with_repository(repository); + + assert!(state + .list_provider_catalog_keys_by_ids(&["key-2".to_string()]) + .await + .is_err()); + } + + #[tokio::test] + async fn credential_cas_fences_provider_and_both_ciphertexts() { + let repository = InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-1"), sample_provider("provider-2")], + Vec::new(), + vec![sample_key( + "key-1", + "provider-2", + Some("api-before".to_string()), + Some("auth-before".to_string()), + )], + ); + let update = ProviderCatalogKeyCredentialsCasUpdate { + key_id: "key-1".to_string(), + expected_provider_id: "provider-1".to_string(), + expected_encrypted_api_key: Some("api-before".to_string()), + expected_encrypted_auth_config: Some("auth-before".to_string()), + encrypted_api_key: Some("api-after".to_string()), + encrypted_auth_config: Some("auth-after".to_string()), + }; + + assert!(!repository + .compare_and_swap_key_credentials(&update) + .await + .expect("provider-fenced CAS should execute")); + let stored = repository + .list_keys_by_ids(&["key-1".to_string()]) + .await + .expect("key should read") + .into_iter() + .next() + .expect("key should exist"); + assert_eq!(stored.encrypted_api_key.as_deref(), Some("api-before")); + assert_eq!(stored.encrypted_auth_config.as_deref(), Some("auth-before")); + } +} diff --git a/apps/aether-gateway/src/state/catalog_proxy.rs b/apps/aether-gateway/src/state/catalog_proxy.rs new file mode 100644 index 000000000..8355acbcc --- /dev/null +++ b/apps/aether-gateway/src/state/catalog_proxy.rs @@ -0,0 +1,1860 @@ +use aether_admin::provider::redaction::{ + admin_json_field_has_contextual_secrets, admin_json_field_is_sensitive, + admin_proxy_credential_field, AdminProxyCredentialField, +}; +use aether_crypto::looks_like_python_fernet_ciphertext; +use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogProxyCasUpdate, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogProvider, +}; +use http::StatusCode; +use percent_encoding::percent_decode_str; +use serde_json::{Map, Value}; +use url::Url; + +use super::AppState; +use crate::handlers::shared::{ + open_runtime_secret_payload, runtime_secret_payload_is_sealed, seal_runtime_secret_payload, +}; +use crate::GatewayError; + +const PROVIDER_PROXY_USERNAME_PURPOSE: &str = "provider-catalog-provider-proxy-username"; +const PROVIDER_PROXY_PASSWORD_PURPOSE: &str = "provider-catalog-provider-proxy-password"; +const ENDPOINT_PROXY_USERNAME_PURPOSE: &str = "provider-catalog-endpoint-proxy-username"; +const ENDPOINT_PROXY_PASSWORD_PURPOSE: &str = "provider-catalog-endpoint-proxy-password"; +const KEY_PROXY_USERNAME_PURPOSE: &str = "provider-catalog-key-proxy-username"; +const KEY_PROXY_PASSWORD_PURPOSE: &str = "provider-catalog-key-proxy-password"; +const CATALOG_PROXY_SECRET_V2_PREFIX: &str = "aether-provider-catalog-proxy-secret-v2:"; +const CATALOG_PROXY_BOUND_PURPOSE_VERSION: &str = "provider-catalog-proxy-credential-v2"; +const CATALOG_PROXY_MIGRATION_RETRIES: usize = 8; + +#[derive(Debug, Clone, Copy)] +enum CatalogProxyScope { + Provider, + Endpoint, + Key, +} + +impl CatalogProxyScope { + fn label(self) -> &'static str { + match self { + Self::Provider => "provider", + Self::Endpoint => "endpoint", + Self::Key => "key", + } + } + + fn username_purpose(self) -> &'static str { + match self { + Self::Provider => PROVIDER_PROXY_USERNAME_PURPOSE, + Self::Endpoint => ENDPOINT_PROXY_USERNAME_PURPOSE, + Self::Key => KEY_PROXY_USERNAME_PURPOSE, + } + } + + fn password_purpose(self) -> &'static str { + match self { + Self::Provider => PROVIDER_PROXY_PASSWORD_PURPOSE, + Self::Endpoint => ENDPOINT_PROXY_PASSWORD_PURPOSE, + Self::Key => KEY_PROXY_PASSWORD_PURPOSE, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum CatalogProxySource { + Incoming, + Stored, +} + +struct CatalogProxyProjection { + runtime: Option, + protected: Option, + migration_required: bool, +} + +struct CatalogProxyCredential { + plaintext: String, + protected: String, + migration_required: bool, +} + +impl AppState { + pub(super) fn protect_provider_catalog_provider( + &self, + provider: &StoredProviderCatalogProvider, + ) -> Result { + let mut protected = provider.clone(); + protected.proxy = self.protect_catalog_proxy( + CatalogProxyScope::Provider, + &provider.id, + provider.proxy.as_ref(), + )?; + Ok(protected) + } + + pub(super) fn protect_provider_catalog_endpoint( + &self, + endpoint: &StoredProviderCatalogEndpoint, + ) -> Result { + let mut protected = endpoint.clone(); + protected.proxy = self.protect_catalog_proxy( + CatalogProxyScope::Endpoint, + &endpoint.id, + endpoint.proxy.as_ref(), + )?; + Ok(protected) + } + + pub(super) fn protect_provider_catalog_key( + &self, + key: &StoredProviderCatalogKey, + ) -> Result { + let mut protected = self.protect_provider_catalog_key_credentials(key)?; + protected.proxy = + self.protect_catalog_proxy(CatalogProxyScope::Key, &key.id, key.proxy.as_ref())?; + Ok(protected) + } + + pub(super) async fn open_provider_catalog_providers( + &self, + providers: Vec, + ) -> Result, GatewayError> { + let mut opened = Vec::with_capacity(providers.len()); + for provider in providers { + opened.push(self.open_provider_catalog_provider(provider).await?); + } + Ok(opened) + } + + pub(super) async fn open_provider_catalog_endpoints( + &self, + endpoints: Vec, + ) -> Result, GatewayError> { + let mut opened = Vec::with_capacity(endpoints.len()); + for endpoint in endpoints { + opened.push(self.open_provider_catalog_endpoint(endpoint).await?); + } + Ok(opened) + } + + pub(super) async fn open_provider_catalog_keys( + &self, + keys: Vec, + ) -> Result, GatewayError> { + let mut opened = Vec::with_capacity(keys.len()); + for key in keys { + opened.push(self.open_provider_catalog_key(key).await?); + } + Ok(opened) + } + + pub(super) async fn open_provider_catalog_provider( + &self, + mut provider: StoredProviderCatalogProvider, + ) -> Result { + for _ in 0..CATALOG_PROXY_MIGRATION_RETRIES { + let projection = self.project_catalog_proxy( + CatalogProxyScope::Provider, + &provider.id, + provider.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if !projection.migration_required { + provider.proxy = projection.runtime; + return Ok(provider); + } + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Provider)?; + let update = ProviderCatalogProxyCasUpdate { + record_id: provider.id.clone(), + expected_proxy: provider.proxy.clone(), + proxy: projection.protected, + }; + if self + .data + .compare_and_swap_provider_catalog_provider_proxy(&update) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + provider.proxy = projection.runtime; + return Ok(provider); + } + provider = self + .data + .list_provider_catalog_providers_by_ids(std::slice::from_ref(&provider.id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .next() + .ok_or_else(|| { + catalog_proxy_migration_changed_error(CatalogProxyScope::Provider) + })?; + } + Err(catalog_proxy_migration_unstable_error( + CatalogProxyScope::Provider, + )) + } + + pub(super) async fn open_provider_catalog_endpoint( + &self, + mut endpoint: StoredProviderCatalogEndpoint, + ) -> Result { + for _ in 0..CATALOG_PROXY_MIGRATION_RETRIES { + let projection = self.project_catalog_proxy( + CatalogProxyScope::Endpoint, + &endpoint.id, + endpoint.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if !projection.migration_required { + endpoint.proxy = projection.runtime; + return Ok(endpoint); + } + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Endpoint)?; + let update = ProviderCatalogProxyCasUpdate { + record_id: endpoint.id.clone(), + expected_proxy: endpoint.proxy.clone(), + proxy: projection.protected, + }; + if self + .data + .compare_and_swap_provider_catalog_endpoint_proxy(&update) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + endpoint.proxy = projection.runtime; + return Ok(endpoint); + } + endpoint = self + .data + .list_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint.id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .next() + .ok_or_else(|| { + catalog_proxy_migration_changed_error(CatalogProxyScope::Endpoint) + })?; + } + Err(catalog_proxy_migration_unstable_error( + CatalogProxyScope::Endpoint, + )) + } + + pub(super) async fn open_provider_catalog_key( + &self, + mut key: StoredProviderCatalogKey, + ) -> Result { + let initial_provider_id = key.provider_id.clone(); + for _ in 0..CATALOG_PROXY_MIGRATION_RETRIES { + if !self + .open_provider_catalog_key_credentials_once(&mut key) + .await? + { + key = self + .data + .list_provider_catalog_keys_by_ids_strong(std::slice::from_ref(&key.id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .next() + .ok_or_else(|| catalog_proxy_migration_changed_error(CatalogProxyScope::Key))?; + if key.provider_id != initial_provider_id { + return Err(GatewayError::Internal( + "provider catalog key provider binding changed during credential migration" + .to_string(), + )); + } + continue; + } + let projection = self.project_catalog_proxy( + CatalogProxyScope::Key, + &key.id, + key.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if !projection.migration_required { + key.proxy = projection.runtime; + return Ok(key); + } + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Key)?; + let update = ProviderCatalogProxyCasUpdate { + record_id: key.id.clone(), + expected_proxy: key.proxy.clone(), + proxy: projection.protected, + }; + if self + .data + .compare_and_swap_provider_catalog_key_proxy(&update) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + key.proxy = projection.runtime; + return Ok(key); + } + key = self + .data + .list_provider_catalog_keys_by_ids_strong(std::slice::from_ref(&key.id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .next() + .ok_or_else(|| catalog_proxy_migration_changed_error(CatalogProxyScope::Key))?; + if key.provider_id != initial_provider_id { + return Err(GatewayError::Internal( + "provider catalog key provider binding changed during proxy migration" + .to_string(), + )); + } + } + Err(catalog_proxy_migration_unstable_error( + CatalogProxyScope::Key, + )) + } + + pub(super) async fn open_provider_transport_snapshot_once( + &self, + snapshot: &mut crate::provider_transport::GatewayProviderTransportSnapshot, + ) -> Result { + let Some(mut stored_key) = self + .data + .list_provider_catalog_keys_by_ids_strong(std::slice::from_ref(&snapshot.key.id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .next() + else { + return Ok(false); + }; + if stored_key.provider_id != snapshot.provider.id + || stored_key.provider_id != snapshot.key.provider_id + { + return Ok(false); + } + if !self + .open_provider_catalog_key_credentials_once(&mut stored_key) + .await? + { + return Ok(false); + } + + let provider = self.project_catalog_proxy( + CatalogProxyScope::Provider, + &snapshot.provider.id, + snapshot.provider.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if provider.migration_required { + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Provider)?; + if !self + .data + .compare_and_swap_provider_catalog_provider_proxy(&ProviderCatalogProxyCasUpdate { + record_id: snapshot.provider.id.clone(), + expected_proxy: snapshot.provider.proxy.clone(), + proxy: provider.protected, + }) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + return Ok(false); + } + } + snapshot.provider.proxy = provider.runtime; + + let endpoint = self.project_catalog_proxy( + CatalogProxyScope::Endpoint, + &snapshot.endpoint.id, + snapshot.endpoint.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if endpoint.migration_required { + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Endpoint)?; + if !self + .data + .compare_and_swap_provider_catalog_endpoint_proxy(&ProviderCatalogProxyCasUpdate { + record_id: snapshot.endpoint.id.clone(), + expected_proxy: snapshot.endpoint.proxy.clone(), + proxy: endpoint.protected, + }) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + return Ok(false); + } + } + snapshot.endpoint.proxy = endpoint.runtime; + + let key = self.project_catalog_proxy( + CatalogProxyScope::Key, + &snapshot.key.id, + snapshot.key.proxy.as_ref(), + CatalogProxySource::Stored, + )?; + if key.migration_required { + self.require_catalog_proxy_migration_writer(CatalogProxyScope::Key)?; + if !self + .data + .compare_and_swap_provider_catalog_key_proxy(&ProviderCatalogProxyCasUpdate { + record_id: snapshot.key.id.clone(), + expected_proxy: snapshot.key.proxy.clone(), + proxy: key.protected, + }) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + return Ok(false); + } + } + snapshot.key.proxy = key.runtime; + Ok(true) + } + + fn require_catalog_proxy_migration_writer( + &self, + scope: CatalogProxyScope, + ) -> Result<(), GatewayError> { + if self.has_provider_catalog_data_writer() { + Ok(()) + } else { + Err(GatewayError::Internal(format!( + "stored {} proxy credentials require migration but the catalog writer is unavailable", + scope.label() + ))) + } + } + + fn protect_catalog_proxy( + &self, + scope: CatalogProxyScope, + record_id: &str, + proxy: Option<&Value>, + ) -> Result, GatewayError> { + self.project_catalog_proxy(scope, record_id, proxy, CatalogProxySource::Incoming) + .map(|projection| projection.protected) + } + + fn project_catalog_proxy( + &self, + scope: CatalogProxyScope, + record_id: &str, + proxy: Option<&Value>, + source: CatalogProxySource, + ) -> Result { + let Some(proxy) = proxy else { + return Ok(CatalogProxyProjection { + runtime: None, + protected: None, + migration_required: false, + }); + }; + if proxy.is_null() { + return Ok(CatalogProxyProjection { + runtime: None, + protected: None, + migration_required: source == CatalogProxySource::Stored, + }); + } + + let proxy_was_string = proxy.is_string(); + let mut object = match proxy { + Value::Object(object) => object.clone(), + Value::String(url) => { + Map::from_iter([("url".to_string(), Value::String(url.to_string()))]) + } + _ => return Err(catalog_proxy_representation_error(scope, source)), + }; + + let mut raw_usernames = Vec::new(); + let mut raw_passwords = Vec::new(); + let mut migration_required = scrub_catalog_proxy_sensitive_fields( + scope, + source, + &mut object, + true, + &mut raw_usernames, + &mut raw_passwords, + )?; + + let mut normalized_url = None; + let mut url_username = None; + let mut url_password = None; + let url_keys = object + .keys() + .filter(|key| matches!(compact_catalog_proxy_key(key).as_str(), "url" | "proxyurl")) + .cloned() + .collect::>(); + for key in url_keys { + let value = object + .remove(&key) + .expect("catalog proxy URL key should still exist"); + if value.is_null() { + migration_required |= source == CatalogProxySource::Stored; + continue; + } + let Some(raw_url) = value.as_str() else { + return Err(catalog_proxy_url_error(scope, source)); + }; + let (candidate_url, username, password) = + parse_catalog_proxy_url(scope, source, raw_url)?; + merge_catalog_proxy_url(scope, source, &mut normalized_url, candidate_url.clone())?; + merge_catalog_proxy_url_credential( + scope, + source, + "username", + &mut url_username, + username, + )?; + merge_catalog_proxy_url_credential( + scope, + source, + "password", + &mut url_password, + password, + )?; + if key != "url" + || (candidate_url != raw_url + && (!proxy_was_string || stored_catalog_proxy_url_requires_cleanup(raw_url))) + { + migration_required |= source == CatalogProxySource::Stored; + } + } + if let Some(url) = normalized_url { + object.insert("url".to_string(), Value::String(url)); + } + + let mut explicit_username = None; + for value in raw_usernames { + let candidate = self.catalog_proxy_credential( + scope, + record_id, + source, + "username", + Some(&value), + scope.username_purpose(), + )?; + merge_catalog_proxy_explicit_credential( + scope, + source, + "username", + &mut explicit_username, + candidate, + )?; + } + let mut explicit_password = None; + for value in raw_passwords { + let candidate = self.catalog_proxy_credential( + scope, + record_id, + source, + "password", + Some(&value), + scope.password_purpose(), + )?; + merge_catalog_proxy_explicit_credential( + scope, + source, + "password", + &mut explicit_password, + candidate, + )?; + } + + let username = merge_catalog_proxy_credential( + scope, + record_id, + source, + "username", + explicit_username, + url_username, + self, + )?; + let password = merge_catalog_proxy_credential( + scope, + record_id, + source, + "password", + explicit_password, + url_password, + self, + )?; + + let mut runtime = object.clone(); + let mut protected = object; + + migration_required |= apply_catalog_proxy_credential( + source, + "username", + username, + &mut runtime, + &mut protected, + ); + migration_required |= apply_catalog_proxy_credential( + source, + "password", + password, + &mut runtime, + &mut protected, + ); + + Ok(CatalogProxyProjection { + runtime: Some(Value::Object(runtime)), + protected: Some(Value::Object(protected)), + migration_required, + }) + } + + fn catalog_proxy_credential( + &self, + scope: CatalogProxyScope, + record_id: &str, + source: CatalogProxySource, + field: &'static str, + value: Option<&Value>, + legacy_purpose: &'static str, + ) -> Result, GatewayError> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_null() { + return Ok(Some(CatalogProxyCredential { + plaintext: String::new(), + protected: String::new(), + migration_required: source == CatalogProxySource::Stored, + })); + } + let Some(value) = value.as_str() else { + return Err(catalog_proxy_credential_error(scope, source, field)); + }; + if value.contains('\0') { + return Err(match source { + CatalogProxySource::Incoming => { + catalog_proxy_credential_error(scope, source, field) + } + CatalogProxySource::Stored => catalog_proxy_storage_error(scope), + }); + } + if value.is_empty() { + return Ok(Some(CatalogProxyCredential { + plaintext: String::new(), + protected: String::new(), + migration_required: source == CatalogProxySource::Stored, + })); + } + + if source == CatalogProxySource::Stored && catalog_proxy_secret_is_v2(value) { + let plaintext = open_catalog_proxy_secret_v2(self, scope, record_id, field, value) + .ok_or_else(|| catalog_proxy_storage_error(scope))?; + return Ok(Some(CatalogProxyCredential { + plaintext, + protected: value.to_string(), + migration_required: false, + })); + } + if source == CatalogProxySource::Stored && runtime_secret_payload_is_sealed(value) { + let plaintext = open_runtime_secret_payload(self, legacy_purpose, value) + .ok_or_else(|| catalog_proxy_storage_error(scope))?; + let protected = seal_catalog_proxy_secret_v2(self, scope, record_id, field, &plaintext) + .ok_or_else(|| catalog_proxy_encryption_error(scope))?; + return Ok(Some(CatalogProxyCredential { + plaintext, + protected, + migration_required: true, + })); + } + if catalog_proxy_secret_is_v2(value) + || runtime_secret_payload_is_sealed(value) + || value.starts_with("aether-") + || looks_like_python_fernet_ciphertext(value) + { + return Err(match source { + CatalogProxySource::Incoming => { + catalog_proxy_credential_error(scope, source, field) + } + CatalogProxySource::Stored => catalog_proxy_storage_error(scope), + }); + } + + let protected = seal_catalog_proxy_secret_v2(self, scope, record_id, field, value) + .ok_or_else(|| catalog_proxy_encryption_error(scope))?; + Ok(Some(CatalogProxyCredential { + plaintext: value.to_string(), + protected, + migration_required: source == CatalogProxySource::Stored, + })) + } +} + +fn catalog_proxy_secret_is_v2(value: &str) -> bool { + value.starts_with(CATALOG_PROXY_SECRET_V2_PREFIX) +} + +fn catalog_proxy_bound_purpose(scope: CatalogProxyScope, record_id: &str, field: &str) -> String { + format!( + "{CATALOG_PROXY_BOUND_PURPOSE_VERSION}\0scope={}\0field={field}\0record-id-bytes={}\0{record_id}", + scope.label(), + record_id.len() + ) +} + +fn seal_catalog_proxy_secret_v2( + state: &AppState, + scope: CatalogProxyScope, + record_id: &str, + field: &str, + plaintext: &str, +) -> Option { + if plaintext.contains('\0') { + return None; + } + let purpose = catalog_proxy_bound_purpose(scope, record_id, field); + seal_runtime_secret_payload(state, &purpose, plaintext) + .map(|sealed| format!("{CATALOG_PROXY_SECRET_V2_PREFIX}{sealed}")) +} + +fn open_catalog_proxy_secret_v2( + state: &AppState, + scope: CatalogProxyScope, + record_id: &str, + field: &str, + stored: &str, +) -> Option { + // The distinct outer envelope is security-significant: a v2 binding failure + // must never be retried with the unbound legacy purpose. + let sealed = stored.strip_prefix(CATALOG_PROXY_SECRET_V2_PREFIX)?; + let purpose = catalog_proxy_bound_purpose(scope, record_id, field); + open_runtime_secret_payload(state, &purpose, sealed) + .filter(|plaintext| !plaintext.contains('\0')) +} + +fn scrub_catalog_proxy_sensitive_fields( + scope: CatalogProxyScope, + source: CatalogProxySource, + object: &mut Map, + at_root: bool, + usernames: &mut Vec, + passwords: &mut Vec, +) -> Result { + let mut changed = false; + let keys = object.keys().cloned().collect::>(); + for key in keys { + let compact_key = compact_catalog_proxy_key(&key); + if compact_key == "hascredentials" { + object.remove(&key); + changed = true; + continue; + } + + if let Some(field) = admin_proxy_credential_field(&key) { + let value = object + .remove(&key) + .expect("catalog proxy credential key should still exist"); + let canonical = at_root + && key + == match field { + AdminProxyCredentialField::Username => "username", + AdminProxyCredentialField::Password => "password", + }; + changed |= !canonical || catalog_proxy_sensitive_value_is_unset(&value); + match field { + AdminProxyCredentialField::Username => usernames.push(value), + AdminProxyCredentialField::Password => passwords.push(value), + } + continue; + } + + let is_credential_container = matches!( + compact_key.as_str(), + "auth" | "proxyauth" | "credentials" | "proxycredentials" + ); + if is_credential_container + && object + .get(&key) + .is_some_and(|value| value.is_object() || value.is_array()) + { + let nested_changed = scrub_catalog_proxy_sensitive_value( + scope, + source, + object + .get_mut(&key) + .expect("catalog proxy credential container should still exist"), + usernames, + passwords, + )?; + changed |= nested_changed; + let is_empty = object.get(&key).is_some_and(|value| match value { + Value::Object(value) => value.is_empty(), + Value::Array(value) => value.is_empty(), + _ => false, + }); + if is_empty { + object.remove(&key); + changed = true; + continue; + } + } + + let value = object + .get(&key) + .expect("catalog proxy field should still exist"); + let is_root_proxy_url = at_root && matches!(compact_key.as_str(), "url" | "proxyurl"); + if !is_root_proxy_url && admin_json_field_has_contextual_secrets(&key, value) { + return Err(catalog_proxy_unsupported_sensitive_error( + scope, source, &key, + )); + } + if admin_json_field_is_sensitive(&key, value) { + if catalog_proxy_sensitive_value_is_unset(value) { + object.remove(&key); + changed = true; + continue; + } + return Err(catalog_proxy_unsupported_sensitive_error( + scope, source, &key, + )); + } + + changed |= scrub_catalog_proxy_sensitive_value( + scope, + source, + object + .get_mut(&key) + .expect("catalog proxy nested field should still exist"), + usernames, + passwords, + )?; + } + Ok(changed && source == CatalogProxySource::Stored) +} + +fn scrub_catalog_proxy_sensitive_value( + scope: CatalogProxyScope, + source: CatalogProxySource, + value: &mut Value, + usernames: &mut Vec, + passwords: &mut Vec, +) -> Result { + match value { + Value::Object(object) => { + scrub_catalog_proxy_sensitive_fields(scope, source, object, false, usernames, passwords) + } + Value::Array(values) => { + let mut changed = false; + for value in values { + changed |= scrub_catalog_proxy_sensitive_value( + scope, source, value, usernames, passwords, + )?; + } + Ok(changed) + } + _ => Ok(false), + } +} + +fn catalog_proxy_sensitive_value_is_unset(value: &Value) -> bool { + match value { + Value::Null => true, + Value::String(value) => value.is_empty(), + Value::Array(value) => value.is_empty(), + Value::Object(value) => value.is_empty(), + _ => false, + } +} + +fn compact_catalog_proxy_key(key: &str) -> String { + key.chars() + .filter(|character| character.is_ascii_alphanumeric()) + .map(|character| character.to_ascii_lowercase()) + .collect() +} + +fn parse_catalog_proxy_url( + scope: CatalogProxyScope, + source: CatalogProxySource, + raw_url: &str, +) -> Result<(String, Option, Option), GatewayError> { + let raw_url = raw_url.trim(); + let mut parsed = Url::parse(raw_url).map_err(|_| catalog_proxy_url_error(scope, source))?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + { + return Err(catalog_proxy_url_error(scope, source)); + } + + let has_non_root_path = !matches!(parsed.path(), "" | "/"); + let has_disallowed_components = + has_non_root_path || parsed.query().is_some() || parsed.fragment().is_some(); + if has_disallowed_components && source == CatalogProxySource::Incoming { + return Err(catalog_proxy_url_error(scope, source)); + } + + let username = (!parsed.username().is_empty() || parsed.password().is_some()) + .then(|| decode_catalog_proxy_userinfo(scope, source, parsed.username())) + .transpose()?; + let password = parsed + .password() + .map(|value| decode_catalog_proxy_userinfo(scope, source, value)) + .transpose()? + .filter(|value| !value.is_empty()); + if username.is_some() || parsed.password().is_some() { + parsed + .set_username("") + .map_err(|_| catalog_proxy_url_error(scope, source))?; + parsed + .set_password(None) + .map_err(|_| catalog_proxy_url_error(scope, source))?; + } + parsed.set_path(""); + parsed.set_query(None); + parsed.set_fragment(None); + Ok((parsed.to_string(), username, password)) +} + +fn stored_catalog_proxy_url_requires_cleanup(raw_url: &str) -> bool { + let Ok(parsed) = Url::parse(raw_url.trim()) else { + return true; + }; + !parsed.username().is_empty() + || parsed.password().is_some() + || !matches!(parsed.path(), "" | "/") + || parsed.query().is_some() + || parsed.fragment().is_some() +} + +fn merge_catalog_proxy_url( + scope: CatalogProxyScope, + source: CatalogProxySource, + current: &mut Option, + candidate: String, +) -> Result<(), GatewayError> { + if current + .as_ref() + .is_some_and(|current| current != &candidate) + { + return Err(catalog_proxy_ambiguous_url_error(scope, source)); + } + *current = Some(candidate); + Ok(()) +} + +fn decode_catalog_proxy_userinfo( + scope: CatalogProxyScope, + source: CatalogProxySource, + value: &str, +) -> Result { + percent_decode_str(value) + .decode_utf8() + .map(|value| value.into_owned()) + .map_err(|_| catalog_proxy_url_error(scope, source)) +} + +fn merge_catalog_proxy_url_credential( + scope: CatalogProxyScope, + source: CatalogProxySource, + field: &'static str, + current: &mut Option, + candidate: Option, +) -> Result<(), GatewayError> { + let Some(candidate) = candidate else { + return Ok(()); + }; + if current + .as_ref() + .is_some_and(|current| current != &candidate) + { + return Err(catalog_proxy_ambiguous_credential_error( + scope, source, field, + )); + } + *current = Some(candidate); + Ok(()) +} + +fn merge_catalog_proxy_credential( + scope: CatalogProxyScope, + record_id: &str, + source: CatalogProxySource, + field: &'static str, + explicit: Option, + from_url: Option, + state: &AppState, +) -> Result, GatewayError> { + let explicit = explicit.filter(|credential| !credential.plaintext.is_empty()); + if let (Some(explicit), Some(from_url)) = (explicit.as_ref(), from_url.as_ref()) { + if &explicit.plaintext != from_url { + return Err(catalog_proxy_ambiguous_credential_error( + scope, source, field, + )); + } + } + if explicit.is_some() { + return Ok(explicit); + } + let Some(from_url) = from_url.filter(|value| !value.is_empty()) else { + return Ok(None); + }; + if from_url.contains('\0') + || catalog_proxy_secret_is_v2(&from_url) + || runtime_secret_payload_is_sealed(&from_url) + || from_url.starts_with("aether-") + || looks_like_python_fernet_ciphertext(&from_url) + { + return Err(match source { + CatalogProxySource::Incoming => catalog_proxy_credential_error(scope, source, field), + CatalogProxySource::Stored => catalog_proxy_storage_error(scope), + }); + } + let protected = seal_catalog_proxy_secret_v2(state, scope, record_id, field, &from_url) + .ok_or_else(|| catalog_proxy_encryption_error(scope))?; + Ok(Some(CatalogProxyCredential { + plaintext: from_url, + protected, + migration_required: source == CatalogProxySource::Stored, + })) +} + +fn merge_catalog_proxy_explicit_credential( + scope: CatalogProxyScope, + source: CatalogProxySource, + field: &'static str, + current: &mut Option, + candidate: Option, +) -> Result<(), GatewayError> { + let Some(candidate) = candidate.filter(|credential| !credential.plaintext.is_empty()) else { + return Ok(()); + }; + let Some(existing) = current.as_mut() else { + *current = Some(candidate); + return Ok(()); + }; + if existing.plaintext != candidate.plaintext { + return Err(catalog_proxy_ambiguous_credential_error( + scope, source, field, + )); + } + existing.migration_required |= candidate.migration_required; + Ok(()) +} + +fn apply_catalog_proxy_credential( + source: CatalogProxySource, + field: &'static str, + credential: Option, + runtime: &mut Map, + protected: &mut Map, +) -> bool { + let existing = runtime.contains_key(field) || protected.contains_key(field); + let Some(credential) = credential else { + runtime.remove(field); + protected.remove(field); + return existing && source == CatalogProxySource::Stored; + }; + runtime.insert(field.to_string(), Value::String(credential.plaintext)); + protected.insert(field.to_string(), Value::String(credential.protected)); + credential.migration_required +} + +fn catalog_proxy_representation_error( + scope: CatalogProxyScope, + source: CatalogProxySource, +) -> GatewayError { + catalog_proxy_error( + scope, + source, + "proxy must be an object, URL string, or null", + ) +} + +fn catalog_proxy_url_error(scope: CatalogProxyScope, source: CatalogProxySource) -> GatewayError { + catalog_proxy_error( + scope, + source, + "proxy URL must be an http, https, socks5, or socks5h origin without path, query, or fragment", + ) +} + +fn catalog_proxy_ambiguous_url_error( + scope: CatalogProxyScope, + source: CatalogProxySource, +) -> GatewayError { + catalog_proxy_error(scope, source, "proxy contains conflicting URL fields") +} + +fn catalog_proxy_unsupported_sensitive_error( + scope: CatalogProxyScope, + source: CatalogProxySource, + field: &str, +) -> GatewayError { + catalog_proxy_error( + scope, + source, + &format!("proxy contains unsupported sensitive field {field}"), + ) +} + +fn catalog_proxy_credential_error( + scope: CatalogProxyScope, + source: CatalogProxySource, + field: &'static str, +) -> GatewayError { + catalog_proxy_error( + scope, + source, + &format!("proxy {field} must be gateway-managed plaintext"), + ) +} + +fn catalog_proxy_ambiguous_credential_error( + scope: CatalogProxyScope, + source: CatalogProxySource, + field: &'static str, +) -> GatewayError { + catalog_proxy_error( + scope, + source, + &format!("proxy URL and proxy {field} contain conflicting credentials"), + ) +} + +fn catalog_proxy_error( + scope: CatalogProxyScope, + source: CatalogProxySource, + detail: &str, +) -> GatewayError { + match source { + CatalogProxySource::Incoming => GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: format!("{} {detail}", scope.label()), + }, + CatalogProxySource::Stored => catalog_proxy_storage_error(scope), + } +} + +fn catalog_proxy_storage_error(scope: CatalogProxyScope) -> GatewayError { + GatewayError::Internal(format!( + "stored {} proxy credentials cannot be decrypted", + scope.label() + )) +} + +fn catalog_proxy_encryption_error(scope: CatalogProxyScope) -> GatewayError { + GatewayError::Internal(format!( + "{} proxy credential encryption is unavailable", + scope.label() + )) +} + +fn catalog_proxy_migration_changed_error(scope: CatalogProxyScope) -> GatewayError { + GatewayError::Internal(format!( + "stored {} proxy changed during credential migration", + scope.label() + )) +} + +fn catalog_proxy_migration_unstable_error(scope: CatalogProxyScope) -> GatewayError { + GatewayError::Internal(format!( + "stored {} proxy credential migration did not stabilize", + scope.label() + )) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; + use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogReadRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogProvider, + }; + use serde_json::{json, Value}; + + use super::{ + catalog_proxy_secret_is_v2, open_catalog_proxy_secret_v2, seal_catalog_proxy_secret_v2, + CatalogProxyScope, ENDPOINT_PROXY_PASSWORD_PURPOSE, PROVIDER_PROXY_PASSWORD_PURPOSE, + }; + use crate::data::GatewayDataState; + use crate::handlers::shared::seal_runtime_secret_payload; + use crate::{AppState, GatewayError}; + + fn sample_provider(id: &str, proxy: Option) -> StoredProviderCatalogProvider { + let mut provider = StoredProviderCatalogProvider::new( + id.to_string(), + format!("Provider {id}"), + Some("https://example.test".to_string()), + "openai".to_string(), + ) + .expect("provider should build"); + provider.proxy = proxy; + provider + } + + fn sample_endpoint(proxy: Option) -> StoredProviderCatalogEndpoint { + StoredProviderCatalogEndpoint::new( + "endpoint-1".to_string(), + "provider-1".to_string(), + "openai:chat".to_string(), + Some("openai".to_string()), + Some("chat".to_string()), + true, + ) + .expect("endpoint should build") + .with_transport_fields( + "https://api.example.test/v1".to_string(), + None, + None, + None, + None, + None, + None, + proxy, + ) + .expect("endpoint transport should build") + } + + fn sample_key(proxy: Option) -> StoredProviderCatalogKey { + StoredProviderCatalogKey::new( + "key-1".to_string(), + "provider-1".to_string(), + "Key 1".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + None, + None::, + None, + None, + None, + None, + None, + proxy, + None, + ) + .expect("key transport should build") + } + + fn state_with_repository(repository: Arc) -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + fn encryption_state() -> AppState { + AppState::new() + .expect("test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + } + + #[tokio::test] + async fn provider_proxy_credentials_are_sealed_before_repository_write() { + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + Vec::new(), + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(Arc::clone(&repository)); + let provider = sample_provider( + "provider-create", + Some(json!({ + "url": "http://alice%40example.test:p%3Ass@proxy.example.test:8080", + "enabled": true + })), + ); + + let created = state + .create_provider_catalog_provider(&provider, None) + .await + .expect("provider create should succeed") + .expect("provider writer should exist"); + assert_eq!( + created.proxy.as_ref().unwrap()["username"], + "alice@example.test" + ); + assert_eq!(created.proxy.as_ref().unwrap()["password"], "p:ss"); + assert_eq!( + created.proxy.as_ref().unwrap()["url"], + "http://proxy.example.test:8080/" + ); + + let stored = repository + .list_providers_by_ids(&["provider-create".to_string()]) + .await + .expect("stored provider should load") + .pop() + .expect("stored provider should exist"); + let serialized = serde_json::to_string(&stored.proxy).expect("proxy should serialize"); + assert!(!serialized.contains("alice@example.test")); + assert!(!serialized.contains("p:ss")); + let stored_proxy = stored.proxy.as_ref().expect("stored proxy should exist"); + assert!(stored_proxy["username"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + assert!(stored_proxy["password"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + assert_eq!(stored_proxy["url"], "http://proxy.example.test:8080/"); + } + + #[tokio::test] + async fn legacy_proxy_url_userinfo_is_migrated_with_field_level_cas() { + let provider = sample_provider( + "provider-legacy", + Some(Value::String( + "http://legacy-user:legacy-pass@proxy.example.test:8080".to_string(), + )), + ); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(Arc::clone(&repository)); + + let opened = state + .list_provider_catalog_providers(false) + .await + .expect("legacy provider should migrate") + .pop() + .expect("legacy provider should exist"); + assert_eq!(opened.proxy.as_ref().unwrap()["username"], "legacy-user"); + assert_eq!(opened.proxy.as_ref().unwrap()["password"], "legacy-pass"); + + let stored = repository + .list_providers_by_ids(&["provider-legacy".to_string()]) + .await + .expect("migrated provider should load") + .pop() + .expect("migrated provider should exist"); + let stored_proxy = stored.proxy.as_ref().expect("stored proxy should exist"); + assert_eq!(stored_proxy["url"], "http://proxy.example.test:8080/"); + assert!(stored_proxy["username"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + assert!(stored_proxy["password"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + } + + #[tokio::test] + async fn legacy_unbound_proxy_ciphertext_is_migrated_to_record_bound_v2_with_cas() { + let encryptor = encryption_state(); + let legacy_password = seal_runtime_secret_payload( + &encryptor, + PROVIDER_PROXY_PASSWORD_PURPOSE, + "legacy-unbound-password", + ) + .expect("legacy password should seal"); + let provider = sample_provider( + "provider-legacy-envelope", + Some(json!({ + "url": "http://proxy.example.test:8080", + "password": legacy_password + })), + ); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(Arc::clone(&repository)); + + let opened = state + .list_provider_catalog_providers(false) + .await + .expect("legacy ciphertext should migrate") + .pop() + .expect("legacy provider should exist"); + assert_eq!( + opened.proxy.as_ref().unwrap()["password"], + "legacy-unbound-password" + ); + + let stored = repository + .list_providers_by_ids(&["provider-legacy-envelope".to_string()]) + .await + .expect("migrated provider should load") + .pop() + .expect("migrated provider should exist"); + let migrated = stored.proxy.as_ref().unwrap()["password"] + .as_str() + .expect("migrated password should be a string"); + assert!(catalog_proxy_secret_is_v2(migrated)); + assert_eq!( + open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Provider, + "provider-legacy-envelope", + "password", + migrated, + ) + .as_deref(), + Some("legacy-unbound-password") + ); + } + + #[tokio::test] + async fn record_bound_proxy_ciphertext_copied_to_another_record_fails_closed() { + let state = encryption_state(); + let protected = state + .protect_provider_catalog_provider(&sample_provider( + "provider-source", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "bound-user", + "password": "bound-password" + })), + )) + .expect("source proxy should seal"); + let copied = sample_provider("provider-target", protected.proxy.clone()); + + let error = state + .open_provider_catalog_provider(copied) + .await + .expect_err("cross-record ciphertext copy must fail closed"); + assert!(matches!(error, GatewayError::Internal(_))); + } + + #[test] + fn provider_endpoint_and_key_proxy_credentials_use_distinct_purposes() { + let state = encryption_state(); + let proxy = Some(json!({ + "url": "socks5h://proxy.example.test:1080", + "username": "purpose-user", + "password": "purpose-password" + })); + let provider = state + .protect_provider_catalog_provider(&sample_provider("provider-1", proxy.clone())) + .expect("provider proxy should seal"); + let endpoint = state + .protect_provider_catalog_endpoint(&sample_endpoint(proxy.clone())) + .expect("endpoint proxy should seal"); + let key = state + .protect_provider_catalog_key(&sample_key(proxy)) + .expect("key proxy should seal"); + + let provider_password = provider.proxy.as_ref().unwrap()["password"] + .as_str() + .unwrap(); + let endpoint_password = endpoint.proxy.as_ref().unwrap()["password"] + .as_str() + .unwrap(); + let key_password = key.proxy.as_ref().unwrap()["password"].as_str().unwrap(); + assert_eq!( + open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Provider, + "provider-1", + "password", + provider_password, + ) + .as_deref(), + Some("purpose-password") + ); + assert!(open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Endpoint, + "endpoint-1", + "password", + provider_password, + ) + .is_none()); + assert_eq!( + open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Endpoint, + "endpoint-1", + "password", + endpoint_password, + ) + .as_deref(), + Some("purpose-password") + ); + assert_eq!( + open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Key, + "key-1", + "password", + key_password, + ) + .as_deref(), + Some("purpose-password") + ); + } + + #[tokio::test] + async fn proxy_url_string_uses_runtime_compatible_object_shape() { + let state = encryption_state(); + let protected = state + .protect_provider_catalog_provider(&sample_provider( + "provider-string", + Some(Value::String("http://proxy.example.test:8080".to_string())), + )) + .expect("proxy string should normalize"); + + assert_eq!( + protected.proxy, + Some(json!({"url": "http://proxy.example.test:8080/"})) + ); + + let opened = state + .open_provider_catalog_provider(sample_provider( + "provider-legacy-string", + Some(Value::String("http://proxy.example.test:8080".to_string())), + )) + .await + .expect("credential-free legacy proxy string should open without a writer"); + assert_eq!( + opened.proxy, + Some(json!({"url": "http://proxy.example.test:8080/"})) + ); + } + + #[tokio::test] + async fn proxy_credentials_preserve_significant_whitespace() { + let state = encryption_state(); + let protected = state + .protect_provider_catalog_provider(&sample_provider( + "provider-spaces", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": " alice ", + "password": " pass phrase " + })), + )) + .expect("proxy credentials should seal"); + let stored = serde_json::to_string(&protected.proxy).expect("proxy should serialize"); + assert!(!stored.contains(" alice ")); + assert!(!stored.contains(" pass phrase ")); + + let opened = state + .open_provider_catalog_provider(protected) + .await + .expect("sealed proxy credentials should open"); + assert_eq!(opened.proxy.as_ref().unwrap()["username"], " alice "); + assert_eq!(opened.proxy.as_ref().unwrap()["password"], " pass phrase "); + } + + #[tokio::test] + async fn proxy_credential_aliases_are_migrated_to_encrypted_canonical_fields() { + let state = encryption_state(); + let protected = state + .protect_provider_catalog_provider(&sample_provider( + "provider-aliases", + Some(json!({ + "proxy_url": "socks5h://proxy.example.test:1080", + "proxy_auth": { + "proxy_user": "alias-user", + "proxy_passphrase": "alias-password" + }, + "region": "test" + })), + )) + .expect("supported proxy credential aliases should normalize"); + let stored_proxy = protected.proxy.as_ref().expect("proxy should exist"); + assert_eq!(stored_proxy["url"], "socks5h://proxy.example.test:1080"); + assert_eq!(stored_proxy["region"], "test"); + assert!(stored_proxy.get("proxy_url").is_none()); + assert!(stored_proxy.get("proxy_auth").is_none()); + assert!(stored_proxy["username"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + assert!(stored_proxy["password"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + let serialized = stored_proxy.to_string(); + assert!(!serialized.contains("alias-user")); + assert!(!serialized.contains("alias-password")); + + let opened = state + .open_provider_catalog_provider(protected) + .await + .expect("canonical proxy credentials should open"); + assert_eq!(opened.proxy.as_ref().unwrap()["username"], "alias-user"); + assert_eq!(opened.proxy.as_ref().unwrap()["password"], "alias-password"); + } + + #[test] + fn unsupported_proxy_sensitive_fields_fail_closed() { + let state = encryption_state(); + for (id, proxy) in [ + ( + "provider-token", + json!({ + "url": "http://proxy.example.test:8080", + "token": "must-not-persist" + }), + ), + ( + "provider-nested-secret", + json!({ + "url": "http://proxy.example.test:8080", + "options": {"clientSecret": "must-not-enter-extra"} + }), + ), + ( + "provider-header-map", + json!({ + "url": "http://proxy.example.test:8080", + "headers": {"x-custom-auth": "must-not-enter-extra"} + }), + ), + ( + "provider-nested-url", + json!({ + "url": "http://proxy.example.test:8080", + "options": { + "health_url": "https://alice:secret@health.example.test/?token=secret" + } + }), + ), + ] { + let error = state + .protect_provider_catalog_provider(&sample_provider(id, Some(proxy))) + .expect_err("unsupported sensitive proxy fields must be rejected"); + assert!(matches!(error, GatewayError::Client { .. })); + } + } + + #[test] + fn incoming_proxy_url_is_origin_only_and_scheme_allowlisted() { + let state = encryption_state(); + for (id, url) in [ + ("provider-ftp", "ftp://proxy.example.test:21"), + ("provider-path", "http://proxy.example.test:8080/path"), + ( + "provider-query", + "http://proxy.example.test:8080/?token=secret", + ), + ( + "provider-fragment", + "socks5://proxy.example.test:1080/#secret", + ), + ] { + let error = state + .protect_provider_catalog_provider(&sample_provider(id, Some(json!({"url": url})))) + .expect_err("non-origin proxy URL must be rejected"); + assert!(matches!(error, GatewayError::Client { .. })); + } + } + + #[tokio::test] + async fn stored_proxy_url_components_are_removed_with_cas_migration() { + let provider = sample_provider( + "provider-url-cleanup", + Some(json!({ + "url": "http://legacy-user:legacy-pass@proxy.example.test:8080/path?token=secret#fragment" + })), + ); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(Arc::clone(&repository)); + + let opened = state + .list_provider_catalog_providers(false) + .await + .expect("legacy proxy URL should migrate") + .pop() + .expect("provider should exist"); + assert_eq!( + opened.proxy.as_ref().unwrap()["url"], + "http://proxy.example.test:8080/" + ); + assert_eq!(opened.proxy.as_ref().unwrap()["username"], "legacy-user"); + assert_eq!(opened.proxy.as_ref().unwrap()["password"], "legacy-pass"); + + let stored = repository + .list_providers_by_ids(&["provider-url-cleanup".to_string()]) + .await + .expect("migrated provider should load") + .pop() + .expect("migrated provider should exist"); + let serialized = stored.proxy.expect("stored proxy should exist").to_string(); + for secret in ["legacy-user", "legacy-pass", "token=secret", "fragment"] { + assert!(!serialized.contains(secret)); + } + } + + #[tokio::test] + async fn password_only_proxy_credentials_remain_compatible() { + let state = encryption_state(); + let protected = state + .protect_provider_catalog_provider(&sample_provider( + "provider-password-only", + Some(json!({ + "url": "http://proxy.example.test:8080", + "password": "legacy-password" + })), + )) + .expect("password-only proxy should seal"); + assert!(protected.proxy.as_ref().unwrap().get("username").is_none()); + assert!(!protected + .proxy + .as_ref() + .unwrap() + .to_string() + .contains("legacy-password")); + + let opened = state + .open_provider_catalog_provider(protected) + .await + .expect("password-only proxy should open"); + assert!(opened.proxy.as_ref().unwrap().get("username").is_none()); + assert_eq!( + opened.proxy.as_ref().unwrap()["password"], + "legacy-password" + ); + + let protected_from_url = state + .protect_provider_catalog_provider(&sample_provider( + "provider-password-only-url", + Some(Value::String( + "http://:url-password@proxy.example.test:8080".to_string(), + )), + )) + .expect("password-only URL userinfo should normalize"); + let stored_proxy = protected_from_url.proxy.as_ref().unwrap(); + assert_eq!(stored_proxy["url"], "http://proxy.example.test:8080/"); + assert!(stored_proxy.get("username").is_none()); + assert!(stored_proxy["password"] + .as_str() + .is_some_and(catalog_proxy_secret_is_v2)); + assert!(!stored_proxy.to_string().contains("url-password")); + } + + #[tokio::test] + async fn damaged_or_wrong_purpose_stored_proxy_ciphertext_fails_closed() { + let encryptor = encryption_state(); + let wrong_purpose = seal_runtime_secret_payload( + &encryptor, + ENDPOINT_PROXY_PASSWORD_PURPOSE, + "must-not-open-as-provider-password", + ) + .expect("test ciphertext should seal"); + let providers = vec![ + sample_provider( + "provider-wrong-purpose", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "alice", + "password": wrong_purpose + })), + ), + sample_provider( + "provider-damaged", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "alice", + "password": "aether-runtime-secret-v1:not-a-fernet-token" + })), + ), + sample_provider( + "provider-foreign-envelope", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "alice", + "password": "aether-payment-gateway-secret-v2:foreign" + })), + ), + ]; + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + providers, + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(repository); + + for provider_id in [ + "provider-wrong-purpose", + "provider-damaged", + "provider-foreign-envelope", + ] { + let error = state + .read_provider_catalog_providers_by_ids(&[provider_id.to_string()]) + .await + .expect_err("invalid stored ciphertext must fail closed"); + assert!(matches!(&error, GatewayError::Internal(_))); + assert!(!error.into_message().contains("must-not-open")); + } + } + + #[tokio::test] + async fn incoming_proxy_ciphertext_and_ambiguous_userinfo_are_rejected() { + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + Vec::new(), + Vec::new(), + Vec::new(), + )); + let state = state_with_repository(Arc::clone(&repository)); + let forged = sample_provider( + "provider-forged", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "alice", + "password": "aether-runtime-secret-v1:not-client-controlled" + })), + ); + let ambiguous = sample_provider( + "provider-ambiguous", + Some(json!({ + "url": "http://url-user:url-pass@proxy.example.test:8080", + "username": "different-user", + "password": "url-pass" + })), + ); + let foreign_envelope = sample_provider( + "provider-foreign-envelope", + Some(json!({ + "url": "http://proxy.example.test:8080", + "username": "alice", + "password": "aether-payment-gateway-secret-v2:foreign" + })), + ); + + for provider in [&forged, &foreign_envelope, &ambiguous] { + let error = state + .create_provider_catalog_provider(provider, None) + .await + .expect_err("untrusted proxy credential representation must be rejected"); + assert!(matches!(&error, GatewayError::Client { .. })); + } + assert!(repository + .list_providers(false) + .await + .expect("repository should remain readable") + .is_empty()); + } + + #[test] + fn provider_username_envelope_is_bound_separately_from_password() { + let state = encryption_state(); + let sealed = seal_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Provider, + "provider-field-binding", + "username", + "separate-user", + ) + .expect("username should seal"); + assert!(open_catalog_proxy_secret_v2( + &state, + CatalogProxyScope::Provider, + "provider-field-binding", + "password", + &sealed, + ) + .is_none()); + } +} diff --git a/apps/aether-gateway/src/state/core.rs b/apps/aether-gateway/src/state/core.rs index ef4b6a4c7..236073550 100644 --- a/apps/aether-gateway/src/state/core.rs +++ b/apps/aether-gateway/src/state/core.rs @@ -13,7 +13,7 @@ use aether_data::repository::proxy_nodes::{ use aether_data_contracts::repository::usage::{ UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot, }; -use aether_http::{build_http_client, HttpClientConfig}; +use aether_http::{apply_http_client_config, HttpClientConfig}; use aether_runtime::{ service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind, MetricLabel, MetricSample, @@ -79,12 +79,7 @@ const SYSTEM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(30); // five minutes of total age. Direct database edits that bypass AppState // invalidation can therefore take at most this bounded interval to appear. const SYSTEM_CONFIG_CACHE_MAX_STALENESS: Duration = Duration::from_secs(5 * 60); -const SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &[ - "enable_format_conversion", - "keep_priority_on_conversion", - "provider_priority_mode", - "scheduling_mode", -]; +const SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &["enable_format_conversion"]; const AUTH_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &[ crate::constants::DEFAULT_USER_GROUP_CONFIG_KEY, crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY, @@ -109,6 +104,24 @@ const USAGE_COUNTER_EXACT_HEALTH_METRICS_TTL: Duration = Duration::from_secs(5 * const USAGE_COUNTER_EXACT_HEALTH_METRICS_MAX_STALENESS: Duration = Duration::from_secs(10 * 60); const USAGE_COUNTER_EXACT_HEALTH_METRICS_RETRY_BACKOFF: Duration = Duration::from_secs(5); +const ADMIN_USAGE_AGGREGATE_INVALID_INPUT_DETAIL: &str = + "usage aggregate import payload is invalid"; + +fn admin_usage_aggregate_import_error(detail: String) -> GatewayError { + // Import validation errors can include source row IDs, table names, and + // adapter-specific details. Keep those details in process memory only; + // the public/admin response receives a stable client-safe message. + warn!( + event_name = "admin_usage_aggregate_import_rejected", + error_length = detail.len(), + "usage aggregate import input was rejected" + ); + GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + message: ADMIN_USAGE_AGGREGATE_INVALID_INPUT_DETAIL.to_string(), + } +} + fn system_config_key_affects_scheduler(key: &str) -> bool { let key = key.trim(); SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key) @@ -141,6 +154,15 @@ impl AppState { .map_err(|err| format!("{err:?}")) } + pub async fn prewarm_execution_extra_trusted_dns_hosts(&self) -> Result<(), String> { + self.read_system_config_json_value( + aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY, + ) + .await + .map(|_| ()) + .map_err(|err| format!("{err:?}")) + } + fn usage_worker_queue_for( runtime_state: &Arc, ) -> Option> { @@ -292,18 +314,38 @@ impl AppState { GatewayDataState::disabled() .with_usage_worker_queue(Self::usage_worker_queue_for(&runtime_state)), ); - let client = build_http_client(&HttpClientConfig { - connect_timeout_ms: Some(10_000), - request_timeout_ms: Some(300_000), - http2_adaptive_window: true, - ..HttpClientConfig::default() - })?; - let owner_forward_client = build_http_client(&HttpClientConfig { - connect_timeout_ms: Some(10_000), - http2_adaptive_window: true, - ..HttpClientConfig::default() - })?; + let client = apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &HttpClientConfig { + connect_timeout_ms: Some(10_000), + request_timeout_ms: Some(300_000), + http2_adaptive_window: true, + ..HttpClientConfig::default() + }, + ) + .build()?; + let owner_forward_client = apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &HttpClientConfig { + connect_timeout_ms: Some(10_000), + http2_adaptive_window: true, + ..HttpClientConfig::default() + }, + ) + .build()?; let frontdoor_runtime_guards = Arc::new(FrontdoorRuntimeGuardConfig::from_env()); + let internal_gateway_auth = + Arc::new(crate::internal_gateway_auth::InternalGatewayAuthConfig::for_process()); + if internal_gateway_auth.status() == "misconfigured" { + warn!( + environment_variable = crate::internal_gateway_auth::INTERNAL_GATEWAY_AUTH_SECRET_ENV, + "internal gateway control plane is fail-closed because its authentication secret is invalid" + ); + } Ok(Self { #[cfg(test)] execution_runtime_override_base_url: execution_runtime_override_base_url @@ -315,6 +357,7 @@ impl AppState { background_data: Arc::clone(&data), background_data_isolated: false, runtime_state: runtime_state.clone(), + internal_gateway_auth, usage_runtime: Arc::new(usage::UsageRuntime::disabled()), video_tasks: Arc::new(VideoTaskService::new( VideoTaskTruthSourceMode::PythonSyncReport, @@ -354,6 +397,7 @@ impl AppState { auth_api_key_force_capabilities_cache: Arc::new(JsonValueCache::default()), auth_api_key_feature_settings_cache: Arc::new(JsonValueCache::default()), auth_daily_quota_availability_cache: Arc::new(ValueCache::default()), + auth_plan_usage_policy_cache: Arc::new(ValueCache::default()), auth_wallet_snapshot_cache: Arc::new(ValueCache::default()), auth_request_cost_upper_bound_cache: Arc::new(ValueCache::default()), provider_quota_snapshot_cache: Arc::new(ValueCache::default()), @@ -456,6 +500,10 @@ impl AppState { true } + pub(crate) fn internal_gateway_auth_status(&self) -> &'static str { + self.internal_gateway_auth.status() + } + #[cfg(test)] pub(crate) fn execution_runtime_override_base_url(&self) -> Option<&str> { self.execution_runtime_override_base_url.as_deref() @@ -730,6 +778,18 @@ impl AppState { .expect("admin monitoring error stats reset cache should lock") } + fn refresh_execution_extra_trusted_dns_hosts( + &self, + key: &str, + value: Option<&serde_json::Value>, + ) { + if key.eq_ignore_ascii_case( + aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY, + ) { + crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(value); + } + } + pub(crate) fn mark_admin_monitoring_error_stats_reset(&self, now_unix_secs: u64) { let mut reset_at = self .admin_monitoring_error_stats_reset_at @@ -742,22 +802,42 @@ impl AppState { &self, key: &str, ) -> Result, GatewayError> { - self.read_system_config_json_value_with_cache_windows( - key, - SYSTEM_CONFIG_CACHE_TTL, - SYSTEM_CONFIG_CACHE_MAX_STALENESS, - ) - .await + let value = self + .read_system_config_json_value_with_cache_windows( + key, + SYSTEM_CONFIG_CACHE_TTL, + SYSTEM_CONFIG_CACHE_MAX_STALENESS, + ) + .await?; + self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref()); + Ok(value) } pub(crate) async fn read_system_config_json_value_strong( &self, key: &str, ) -> Result, GatewayError> { - self.data + let value = self + .data .find_system_config_value_strong(key) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref()); + Ok(value) + } + + pub(crate) async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + let result = self + .data + .compare_and_set_system_config_string_value(key, expected, replacement) + .await; + self.system_config_cache.invalidate(key); + result.map_err(|err| GatewayError::Internal(err.to_string())) } async fn read_system_config_json_value_with_cache_windows( @@ -881,6 +961,7 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string()))?; self.system_config_cache .insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_MAX_STALENESS); + self.refresh_execution_extra_trusted_dns_hosts(key, None); if deleted && system_config_key_affects_scheduler(key) { self.invalidate_scheduler_affinity_cache(); } @@ -950,6 +1031,7 @@ impl AppState { self.auth_api_key_force_capabilities_cache.clear(); self.auth_api_key_feature_settings_cache.clear(); self.auth_daily_quota_availability_cache.clear(); + self.auth_plan_usage_policy_cache.clear(); self.auth_wallet_snapshot_cache.clear(); self.auth_request_cost_upper_bound_cache.clear(); self.provider_quota_snapshot_cache.clear(); @@ -961,6 +1043,7 @@ impl AppState { } fn remember_system_config_write(&self, key: &str, value: Option) { + self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref()); self.system_config_cache .insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_MAX_STALENESS); if system_config_key_affects_scheduler(key) { @@ -1008,6 +1091,7 @@ impl AppState { | aether_data::repository::system::AdminSystemPurgeTarget::Stats ) { self.system_config_cache.clear(); + crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(None); self.invalidate_provider_routing_caches(); } Ok(summary) @@ -1035,10 +1119,9 @@ impl AppState { .import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode) .await .map_err(|err| match err { - aether_data::DataLayerError::InvalidInput(detail) => GatewayError::Client { - status: http::StatusCode::BAD_REQUEST, - message: detail, - }, + aether_data::DataLayerError::InvalidInput(detail) => { + admin_usage_aggregate_import_error(detail) + } other => GatewayError::Internal(other.to_string()), }) } @@ -1129,18 +1212,27 @@ impl AppState { &self, mutation: &aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation, ) -> Result, GatewayError> { - self.data - .register_proxy_node(mutation) - .await - .map_err(|err| GatewayError::Internal(err.to_string())) + self.register_proxy_node_with_bound_secrets(mutation).await } pub(crate) async fn create_manual_proxy_node( &self, mutation: &ProxyNodeManualCreateMutation, ) -> Result, GatewayError> { + let mut protected = mutation.clone(); + let node_id = mutation + .node_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + protected.node_id = Some(node_id.clone()); + if let Some(password) = mutation.proxy_password.as_deref() { + protected.proxy_password = Some(self.protect_proxy_node_password(&node_id, password)?); + } self.data - .create_manual_proxy_node(mutation) + .create_manual_proxy_node(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -1149,8 +1241,17 @@ impl AppState { &self, mutation: &ProxyNodeManualUpdateMutation, ) -> Result, GatewayError> { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + let mut protected = mutation.clone(); + protected.node_id = existing.id.clone(); + if let Some(password) = mutation.proxy_password.as_deref() { + protected.proxy_password = + Some(self.protect_proxy_node_password(&existing.id, password)?); + } self.data - .update_manual_proxy_node(mutation) + .update_manual_proxy_node(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -1183,6 +1284,9 @@ impl AppState { &self, mutation: &ProxyNodeHeartbeatMutation, ) -> Result, GatewayError> { + crate::state::decrypt_or_migrate_proxy_tunnel_psk(&self.data, &mutation.node_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; self.data .apply_proxy_node_heartbeat(mutation) .await @@ -2020,9 +2124,16 @@ impl AppState { mut self, path: impl Into, ) -> std::io::Result { + let encryption_key = self.data.encryption_key().ok_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "AETHER_GATEWAY_VIDEO_TASK_STORE_PATH requires a configured encryption key", + ) + })?; self.video_tasks = Arc::new(VideoTaskService::with_file_store( self.video_tasks.truth_source_mode(), path, + encryption_key, )?); Ok(self) } @@ -3848,7 +3959,8 @@ mod tests { use serde_json::json; use super::{ - database_bounded_auth_load_limit, merge_usage_counter_health_snapshots, + admin_usage_aggregate_import_error, database_bounded_auth_load_limit, + merge_usage_counter_health_snapshots, usage_counter_pending_health_metric_samples_with_timeout, usage_queue_health_metric_samples_with_timeout, usage_runtime_metric_samples, AppState, MetricKind, MetricSample, METRIC_SNAPSHOT_TTL, @@ -3857,6 +3969,23 @@ mod tests { use crate::cache::SchedulerAffinityTarget; use crate::data::{GatewayDataConfig, GatewayDataState}; + #[test] + fn usage_aggregate_import_invalid_input_is_projected_to_a_safe_client_message() { + let error = admin_usage_aggregate_import_error( + "postgres table stats_user_daily row secret-user contains column password".to_string(), + ); + + match error { + super::GatewayError::Client { status, message } => { + assert_eq!(status, http::StatusCode::BAD_REQUEST); + assert_eq!(message, "usage aggregate import payload is invalid"); + assert!(!message.contains("secret-user")); + assert!(!message.contains("password")); + } + other => panic!("expected client-safe import error, got {other:?}"), + } + } + #[test] fn auth_load_gate_reserves_half_of_foreground_database_pool() { assert_eq!( @@ -4350,12 +4479,13 @@ mod tests { } #[tokio::test] - async fn system_config_entry_write_refreshes_cache_and_scheduler_affinity_for_routing_keys() { + async fn system_config_entry_write_refreshes_cache_and_scheduler_affinity_for_format_conversion( + ) { let state = AppState::new() .expect("app state should build") .with_data_state_for_tests( GatewayDataState::disabled().with_system_config_values_for_tests([( - "keep_priority_on_conversion".to_string(), + "enable_format_conversion".to_string(), json!(false), )]), ); @@ -4364,7 +4494,7 @@ mod tests { assert_eq!( state - .read_system_config_json_value("keep_priority_on_conversion") + .read_system_config_json_value("enable_format_conversion") .await .expect("system config read should succeed"), Some(json!(false)) @@ -4385,13 +4515,13 @@ mod tests { let initial_epoch = state.scheduler_affinity_epoch(); state - .upsert_system_config_entry("keep_priority_on_conversion", &json!(true), None) + .upsert_system_config_entry("enable_format_conversion", &json!(true), None) .await .expect("admin config write should succeed"); assert_eq!( state - .read_system_config_json_value("keep_priority_on_conversion") + .read_system_config_json_value("enable_format_conversion") .await .expect("system config read should use refreshed cache"), Some(json!(true)) diff --git a/apps/aether-gateway/src/state/integrations.rs b/apps/aether-gateway/src/state/integrations.rs index 9c3397c6b..8d61f0a77 100644 --- a/apps/aether-gateway/src/state/integrations.rs +++ b/apps/aether-gateway/src/state/integrations.rs @@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot; +use aether_data_contracts::DataLayerError; use aether_model_fetch::{ aggregate_models_for_cache, build_antigravity_load_code_assist_plan, fetch_models_from_transports, merge_upstream_metadata, model_fetch_interval_minutes, @@ -25,7 +26,7 @@ use tracing::{debug, warn}; use super::{AppState, GatewayError}; use crate::clock::current_unix_secs; -use crate::model_fetch::{CodexCatalogRuntime, ModelFetchRuntimeState}; +use crate::model_fetch::{safe_model_fetch_error, CodexCatalogRuntime, ModelFetchRuntimeState}; use crate::provider_transport::{GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth}; use crate::request_candidate_runtime::{ RequestCandidateRuntimeCapabilityReader, RequestCandidateRuntimeReader, @@ -36,6 +37,44 @@ use crate::{execution_runtime, provider_transport}; const MODEL_FETCH_RESPONSE_BODY_LIMIT_BYTES: usize = 8 * 1024 * 1024; +#[async_trait] +impl provider_transport::ProviderTransportSnapshotSource for AppState { + fn encryption_key(&self) -> Option<&str> { + AppState::encryption_key(self) + } + + async fn list_provider_catalog_providers_by_ids( + &self, + ids: &[String], + ) -> Result, DataLayerError> { + self.read_provider_catalog_providers_by_ids(ids) + .await + .map_err(provider_transport_snapshot_data_error) + } + + async fn list_provider_catalog_endpoints_by_ids( + &self, + ids: &[String], + ) -> Result, DataLayerError> { + self.read_provider_catalog_endpoints_by_ids(ids) + .await + .map_err(provider_transport_snapshot_data_error) + } + + async fn list_provider_catalog_keys_by_ids( + &self, + ids: &[String], + ) -> Result, DataLayerError> { + self.read_provider_catalog_keys_by_ids(ids) + .await + .map_err(provider_transport_snapshot_data_error) + } +} + +fn provider_transport_snapshot_data_error(error: GatewayError) -> DataLayerError { + DataLayerError::UnexpectedValue(error.into_message()) +} + impl AppState { pub(crate) async fn hydrate_antigravity_project_metadata_for_transport( &self, @@ -159,7 +198,7 @@ impl AppState { provider_id = %transport.provider.id, endpoint_id = %transport.endpoint.id, key_id = %transport.key.id, - error = %err, + error = %safe_model_fetch_error(&err), "gemini_cli project metadata hydration failed" ); return None; @@ -324,31 +363,12 @@ impl CodexCatalogRuntime for AppState { return Ok(Some(scope)); } - let decrypted_auth_config = match key.encrypted_auth_config.as_deref() { - Some(ciphertext) => Some( - crate::handlers::shared::decrypt_catalog_secret_with_fallbacks( - self.encryption_key(), - ciphertext, - ) - .ok_or_else(|| { - "Codex catalog auth config could not be verified for credential fencing" - .to_string() - })?, - ), - None => None, - }; - let decrypted_api_key = match key.encrypted_api_key.as_deref() { - Some(ciphertext) => Some( - crate::handlers::shared::decrypt_catalog_secret_with_fallbacks( - self.encryption_key(), - ciphertext, - ) - .ok_or_else(|| { - "Codex catalog API key could not be verified for credential fencing".to_string() - })?, - ), - None => None, - }; + let decrypted_auth_config = self + .decrypt_provider_catalog_key_auth_config(&key) + .map_err(GatewayError::into_message)?; + let decrypted_api_key = self + .decrypt_provider_catalog_key_api_key(&key) + .map_err(GatewayError::into_message)?; Ok( crate::model_fetch::codex_catalog_credential_scope_from_stored_key( @@ -377,6 +397,16 @@ impl ModelFetchRuntimeState for AppState { AppState::list_provider_catalog_providers(self, active_only).await } + async fn list_provider_catalog_providers_for_model_fetch( + &self, + active_only: bool, + ) -> Result, GatewayError> { + self.data + .list_provider_catalog_providers(active_only) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + async fn list_provider_catalog_endpoints_by_provider_ids( &self, provider_ids: &[String], @@ -384,6 +414,26 @@ impl ModelFetchRuntimeState for AppState { AppState::list_provider_catalog_endpoints_by_provider_ids(self, provider_ids).await } + async fn list_provider_catalog_endpoints_for_model_fetch( + &self, + provider_ids: &[String], + ) -> Result, GatewayError> { + self.data + .list_provider_catalog_endpoints_by_provider_ids(provider_ids) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + async fn list_provider_catalog_keys_for_model_fetch( + &self, + provider_ids: &[String], + ) -> Result, String> { + self.data + .list_provider_catalog_keys_by_provider_ids(provider_ids) + .await + .map_err(|err| err.to_string()) + } + async fn read_provider_transport_snapshot( &self, provider_id: &str, @@ -405,22 +455,9 @@ impl ModelFetchRuntimeState for AppState { provider_id: &str, key_id: &str, ) -> Option { - let credential_scope = - ::read_codex_catalog_credential_scope_strong( - self, - provider_id, - key_id, - ) + crate::model_fetch::read_codex_management_catalog(self, provider_id, key_id) .await - .ok() - .flatten()?; - crate::model_fetch::read_recent_codex_catalog_client_version( - self.runtime_state.as_ref(), - provider_id, - key_id, - &credential_scope, - ) - .await + .map(|catalog| catalog.client_version) } async fn update_provider_catalog_key_model_fetch_state( @@ -680,10 +717,4 @@ impl SchedulerRuntimeState for AppState { expected_epoch, ) } - - async fn read_scheduler_ordering_config( - &self, - ) -> Result { - crate::scheduler::config::read_scheduler_ordering_config(self).await - } } diff --git a/apps/aether-gateway/src/state/mod.rs b/apps/aether-gateway/src/state/mod.rs index c342a7c52..4bd15edbb 100644 --- a/apps/aether-gateway/src/state/mod.rs +++ b/apps/aether-gateway/src/state/mod.rs @@ -6,6 +6,8 @@ mod app; mod bootstrap_admin; mod cache; mod catalog; +mod catalog_credentials; +mod catalog_proxy; mod core; mod cors; mod integrations; @@ -23,7 +25,8 @@ pub(crate) use self::admin_types::{ AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AdminPaymentCallbackRecord, AdminSecurityBlacklistEntry, AdminWalletPaymentOrderRecord, AdminWalletRefundRecord, AdminWalletTransactionRecord, BillingPlanRecord, - BillingPlanWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + BillingPlanWriteInput, PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, + PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; pub use self::app::AppState; @@ -42,10 +45,15 @@ pub(crate) use self::oauth::{ provider_transport_context_allows_credential_rotation, AgentIdentityAuthConfigFence, CodexRuntimeOAuthObservation, ProviderTransportCredentialFence, }; +pub(crate) use self::proxy::{ + decrypt_or_migrate_proxy_tunnel_psk, decrypt_or_migrate_proxy_tunnel_psk_binding, + unavailable_proxy_snapshot, +}; pub(crate) use self::types::{ AdminWalletMutationOutcome, GatewayAdminPaymentCallbackView, GatewayUserPreferenceView, GatewayUserSessionView, LocalExecutionRuntimeMissDiagnostic, LocalMutationOutcome, LocalProviderDeleteTaskState, }; +pub(crate) use self::video::VideoTaskRouteAccess; use super::provider_transport::provider_transport_snapshot_looks_refreshed; pub(crate) use super::provider_transport::ProviderTransportSnapshotCacheKey; diff --git a/apps/aether-gateway/src/state/oauth.rs b/apps/aether-gateway/src/state/oauth.rs index f541e59b9..0accef522 100644 --- a/apps/aether-gateway/src/state/oauth.rs +++ b/apps/aether-gateway/src/state/oauth.rs @@ -4,9 +4,7 @@ use super::{ ProviderTransportSnapshotFlightResult, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL, }; -use crate::handlers::shared::{ - decrypt_catalog_secret_with_fallbacks, default_provider_key_status_snapshot, -}; +use crate::handlers::shared::default_provider_key_status_snapshot; use crate::provider_transport::LocalOAuthHttpExecutor; use super::super::provider_transport; @@ -24,11 +22,9 @@ use aether_data_contracts::repository::provider_catalog::{ use aether_runtime_state::RuntimeLockLease; use base64::{engine::general_purpose::STANDARD, Engine as _}; use dashmap::{mapref::entry::Entry as DashMapEntry, DashMap}; -use flate2::read::{DeflateDecoder, GzDecoder}; use serde_json::{json, Map, Value}; use sha2::{Digest, Sha256}; use std::collections::BTreeMap; -use std::io::Read; use std::sync::atomic::Ordering; use std::sync::Arc; use std::time::{Duration, SystemTime, UNIX_EPOCH}; @@ -38,6 +34,7 @@ use aether_crypto::{ }; const LOCAL_OAUTH_HTTP_TIMEOUT_MS: u64 = 30_000; +const LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; const REMOTE_OAUTH_REFRESH_WAIT_TIMEOUT: Duration = Duration::from_secs(35); const REMOTE_OAUTH_REFRESH_POLL_INTERVAL: Duration = Duration::from_millis(100); const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] "; @@ -46,12 +43,22 @@ const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] "; const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] "; const CODEX_OAUTH_INVALIDATION_CAS_MAX_ATTEMPTS: usize = 16; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub(crate) struct ProviderTransportCredentialFence { pub(crate) encrypted_auth_config: String, pub(crate) credential: ProviderCatalogKeyOAuthCredentialFence, } +impl std::fmt::Debug for ProviderTransportCredentialFence { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderTransportCredentialFence") + .field("encrypted_auth_config", &"[REDACTED]") + .field("credential", &self.credential) + .finish() + } +} + #[derive(Debug, Clone, Copy)] pub(crate) struct CodexRuntimeOAuthObservation<'a> { pub(crate) request_started_at_unix_ms: u64, @@ -262,11 +269,11 @@ fn local_oauth_request_refresh_token_fingerprint( } fn local_oauth_log_excerpt(body: &str) -> String { - let body = body.trim(); - if body.is_empty() { - return "-".to_string(); + if body.trim().is_empty() { + "-".to_string() + } else { + "[redacted]".to_string() } - body.chars().take(300).collect() } fn local_oauth_proxy_is_tunnel(proxy: Option<&ProxySnapshot>) -> bool { @@ -331,7 +338,6 @@ fn normalize_local_oauth_refresh_error_message( ) -> String { let mut message = None::; let mut error_code = None::; - let mut error_type = None::; if let Some(body_excerpt) = body_excerpt { if let Ok(value) = serde_json::from_str::(body_excerpt) { @@ -352,12 +358,6 @@ fn normalize_local_oauth_refresh_error_message( .map(str::trim) .filter(|value| !value.is_empty()) .map(|value| value.to_ascii_lowercase()); - error_type = error_object - .get("type") - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); } if message.is_none() { message = object @@ -376,14 +376,6 @@ fn normalize_local_oauth_refresh_error_message( .filter(|value| !value.is_empty()) .map(|value| value.to_ascii_lowercase()); } - if error_type.is_none() { - error_type = object - .get("type") - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); - } } } } @@ -398,7 +390,6 @@ fn normalize_local_oauth_refresh_error_message( .unwrap_or_default(); let lowered = message.to_ascii_lowercase(); let error_code = error_code.unwrap_or_default(); - let error_type = error_type.unwrap_or_default(); if error_code == "refresh_token_reused" || lowered.contains("already been used to generate a new access token") @@ -416,15 +407,9 @@ fn normalize_local_oauth_refresh_error_message( { return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string(); } - if error_type == "invalid_request_error" && !message.is_empty() { - return message; - } - if !message.is_empty() { - return message; - } status_code .map(|status_code| format!("HTTP {status_code}")) - .unwrap_or_else(|| "未知错误".to_string()) + .unwrap_or_else(|| "Token 刷新失败".to_string()) } fn merge_local_oauth_refresh_failure_reason( @@ -915,10 +900,14 @@ impl AppState { key_id: &str, ) -> Result, GatewayError> { - self.data - .read_provider_transport_snapshot(provider_id, endpoint_id, key_id) - .await - .map_err(|err| GatewayError::Internal(err.to_string())) + crate::provider_transport::read_provider_transport_snapshot( + self, + provider_id, + endpoint_id, + key_id, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) } async fn apply_global_format_conversion_override( @@ -961,13 +950,38 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } - pub(crate) async fn upsert_ldap_module_config( + pub(crate) async fn compare_and_swap_ldap_module_config( &self, - config: &aether_data::repository::auth_modules::StoredLdapModuleConfig, - ) -> Result, GatewayError> - { + expected: Option<&aether_data::repository::auth_modules::StoredLdapModuleConfig>, + replacement: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + bind_password_update: &aether_data::repository::auth_modules::LdapBindPasswordUpdate, + ) -> Result< + Option, + GatewayError, + > { self.data - .upsert_ldap_module_config(config) + .compare_and_swap_ldap_module_config(expected, replacement, bind_password_update) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn delete_ldap_module_config_if_matches( + &self, + expected: &aether_data::repository::auth_modules::StoredLdapModuleConfig, + ) -> Result { + self.data + .delete_ldap_module_config_if_matches(expected) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + self.data + .compare_and_swap_ldap_bind_password(expected, replacement) .await .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -993,6 +1007,16 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + self.data + .has_oauth_links_for_provider(provider_type) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn get_oauth_provider_config( &self, provider_type: &str, @@ -1006,6 +1030,18 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn compare_and_swap_oauth_provider_client_secret( + &self, + provider_type: &str, + expected: &str, + replacement: &str, + ) -> Result { + self.data + .compare_and_swap_oauth_provider_client_secret(provider_type, expected, replacement) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn count_locked_users_if_oauth_provider_disabled( &self, provider_type: &str, @@ -1024,18 +1060,95 @@ impl AppState { Option, GatewayError, > { + let _mutation_guard = crate::oauth::lock_identity_oauth_mutation().await; + let ldap_exclusive = self.get_ldap_module_config().await?.is_some_and(|config| { + config.is_enabled + && config.is_exclusive + && config + .bind_password_encrypted + .as_deref() + .map(str::trim) + .is_some_and(|value| !value.is_empty()) + }); + let locked_users_snapshot = if record.is_enabled { + 0 + } else { + self.count_locked_users_if_oauth_provider_disabled( + &record.provider_type, + ldap_exclusive, + ) + .await? + }; + match self + .data + .upsert_oauth_provider_config(record, ldap_exclusive, false, locked_users_snapshot) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + Some( + aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::Upserted( + provider, + ), + ) => Ok(Some(provider)), + Some( + aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count, + }, + ) => Err(GatewayError::Client { + status: axum::http::StatusCode::CONFLICT, + message: format!( + "禁用 OAuth Provider '{}' 会导致 {affected_count} 个用户无法登录", + record.provider_type + ), + }), + None => Ok(None), + } + } + + pub(crate) async fn upsert_oauth_provider_config_with_force_disable( + &self, + record: &aether_data::repository::oauth_providers::UpsertOAuthProviderConfigRecord, + force_disable: bool, + ) -> Result< + Option, + GatewayError, + > { + let _mutation_guard = crate::oauth::lock_identity_oauth_mutation().await; + let ldap_exclusive = self.get_ldap_module_config().await?.is_some_and(|config| { + config.is_enabled + && config.is_exclusive + && config + .bind_password_encrypted + .as_deref() + .map(str::trim) + .is_some_and(|value| !value.is_empty()) + }); + let locked_users_snapshot = if record.is_enabled || force_disable { + 0 + } else { + self.count_locked_users_if_oauth_provider_disabled( + &record.provider_type, + ldap_exclusive, + ) + .await? + }; self.data - .upsert_oauth_provider_config(record) + .upsert_oauth_provider_config( + record, + ldap_exclusive, + force_disable, + locked_users_snapshot, + ) .await .map_err(|err| GatewayError::Internal(err.to_string())) } - pub(crate) async fn delete_oauth_provider_config( + pub(crate) async fn delete_oauth_provider_config_if_unlinked( &self, provider_type: &str, ) -> Result { self.data - .delete_oauth_provider_config(provider_type) + .delete_oauth_provider_config_if_unlinked(provider_type) .await .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -1307,35 +1420,11 @@ impl AppState { .map(|snapshot| (*snapshot).clone())) } - pub(crate) async fn update_provider_catalog_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - let updated = self - .data - .update_provider_catalog_key_oauth_credentials( - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; - if updated { - self.clear_provider_transport_snapshot_cache(); - } - Ok(updated) - } - pub(crate) async fn update_provider_catalog_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { let updated = self @@ -1344,7 +1433,6 @@ impl AppState { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -1359,9 +1447,29 @@ impl AppState { &self, update: &ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ) -> Result { + let mut protected = update.clone(); + let expected_credential = update.expected_credential.clone().ok_or_else(|| { + GatewayError::Internal( + "provider catalog OAuth credential update requires a complete credential fence" + .to_string(), + ) + })?; + self.validate_protected_provider_catalog_key_auth_config( + &expected_credential.provider_id, + &update.key_id, + &update.encrypted_auth_config, + )?; + if let Some(encrypted_api_key) = update.encrypted_api_key_update.as_deref() { + self.validate_protected_provider_catalog_key_api_key( + &expected_credential.provider_id, + &update.key_id, + encrypted_api_key, + )?; + } + protected.expected_credential = Some(expected_credential); let updated = self .data - .compare_and_update_provider_catalog_key_oauth_runtime_state(update) + .compare_and_update_provider_catalog_key_oauth_runtime_state(&protected) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; // A conflict means another instance/admin changed the credential. @@ -1430,7 +1538,7 @@ impl AppState { body_excerpt, .. }) if matches!(status_code, 400 | 401 | 403) => { - if let Err(err) = self + if self .persist_local_oauth_refresh_failure_state( ¤t_transport, status_code, @@ -1438,17 +1546,21 @@ impl AppState { false, ) .await + .is_err() { tracing::warn!( key_id = %current_transport.key.id, provider_type = %current_transport.provider.provider_type, - error = ?err, "gateway local oauth refresh failure persistence failed" ); } return Ok(None); } - Err(err) => return Err(GatewayError::Internal(err.to_string())), + Err(_) => { + return Err(GatewayError::Internal( + "local oauth refresh failed".to_string(), + )); + } }; if resolution @@ -1485,18 +1597,18 @@ impl AppState { .await; return Ok(None); } - if let Err(err) = self + if self .persist_local_oauth_refresh_entry( ¤t_transport, &refreshed_entry, expected_credential_fence.as_ref(), ) .await + .is_err() { tracing::warn!( key_id = %current_transport.key.id, provider_type = %current_transport.provider.provider_type, - error = ?err, "gateway local oauth refresh persistence failed" ); let _ = self @@ -1682,18 +1794,18 @@ impl AppState { .await; return Ok(None); } - if let Err(err) = self + if self .persist_local_oauth_refresh_entry( ¤t_transport, &refreshed_entry, expected_credential_fence.as_ref(), ) .await + .is_err() { tracing::warn!( key_id = %current_transport.key.id, provider_type = %current_transport.provider.provider_type, - error = ?err, "gateway manual oauth refresh persistence failed" ); let _ = self @@ -1708,7 +1820,7 @@ impl AppState { return Err( provider_transport::LocalOAuthRefreshError::InvalidResponse { provider_type: "gateway", - message: format!("local oauth refresh persistence failed: {err:?}"), + message: "local oauth refresh persistence failed".to_string(), }, ); } @@ -1745,10 +1857,9 @@ impl AppState { let Some(lease) = lease else { return; }; - if let Err(err) = self.runtime_state.lock_release(&lease).await { + if self.runtime_state.lock_release(&lease).await.is_err() { tracing::warn!( key_id = %lease.key, - error = ?err, "gateway local oauth refresh distributed lease release failed" ); } @@ -1814,18 +1925,13 @@ impl AppState { return Ok(None); } - let stored_api_key = match stored.encrypted_api_key.as_deref() { - Some(ciphertext) => Some( - decrypt_catalog_secret_with_fallbacks(self.data.encryption_key(), ciphertext) - .ok_or_else(|| { - GatewayError::Internal( - "provider api_key could not be verified for runtime fencing" - .to_string(), - ) - })?, - ), - None => None, - }; + let stored_api_key = self + .decrypt_provider_catalog_key_api_key(&stored) + .map_err(|_| { + GatewayError::Internal( + "provider api_key could not be verified for runtime fencing".to_string(), + ) + })?; let transport_api_key = (!transport.key.decrypted_api_key.is_empty()) .then_some(transport.key.decrypted_api_key.as_str()); if stored_api_key.as_deref() != transport_api_key { @@ -1835,14 +1941,18 @@ impl AppState { let Some(ciphertext) = stored.encrypted_auth_config.as_deref() else { return Ok(None); }; - let plaintext = - decrypt_catalog_secret_with_fallbacks(self.data.encryption_key(), ciphertext) - .ok_or_else(|| { - GatewayError::Internal( - "provider auth_config could not be verified for runtime fencing" - .to_string(), - ) - })?; + let plaintext = self + .decrypt_provider_catalog_key_auth_config(&stored) + .map_err(|_| { + GatewayError::Internal( + "provider auth_config could not be verified for runtime fencing".to_string(), + ) + })? + .ok_or_else(|| { + GatewayError::Internal( + "provider auth_config could not be verified for runtime fencing".to_string(), + ) + })?; let config = serde_json::from_str::(&plaintext) .map_err(|err| GatewayError::Internal(err.to_string()))?; let transport_config = transport @@ -1927,7 +2037,6 @@ impl AppState { key_id, latest_key.oauth_invalid_at_unix_secs, latest_key.oauth_invalid_reason.as_deref(), - None, latest_key.updated_at_unix_secs, ) .await?; @@ -2367,9 +2476,9 @@ impl AppState { return Ok(()); } - let Some(encryption_key) = self.data.encryption_key() else { + if self.data.encryption_key().is_none() { return Ok(()); - }; + } if provider_transport::is_codex_agent_identity_cached_entry(entry) { let metadata = entry.metadata.as_ref().ok_or_else(|| { @@ -2381,9 +2490,11 @@ impl AppState { .map_err(GatewayError::Internal)?; let auth_config = serde_json::to_string(metadata) .map_err(|err| GatewayError::Internal(err.to_string()))?; - let encrypted_auth_config = - encrypt_python_fernet_plaintext(encryption_key, &auth_config) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let encrypted_auth_config = self.seal_provider_catalog_key_auth_config( + &transport.provider.id, + key_id, + &auth_config, + )?; let source_fingerprint = entry.source_fingerprint.as_deref().ok_or_else(|| { GatewayError::Internal( @@ -2437,16 +2548,14 @@ impl AppState { .to_string(), )); } - let latest_auth_config = decrypt_catalog_secret_with_fallbacks( - Some(encryption_key), - expected_encrypted_auth_config.as_str(), - ) - .and_then(|value| serde_json::from_str::(&value).ok()) - .ok_or_else(|| { - GatewayError::Internal( - "Agent Identity current auth_config could not be verified".to_string(), - ) - })?; + let latest_auth_config = self + .decrypt_provider_catalog_key_auth_config(&latest_key)? + .and_then(|value| serde_json::from_str::(&value).ok()) + .ok_or_else(|| { + GatewayError::Internal( + "Agent Identity current auth_config could not be verified".to_string(), + ) + })?; let latest_fingerprint = provider_transport::codex_agent_identity_credential_fingerprint( &latest_auth_config, @@ -2527,26 +2636,32 @@ impl AppState { ) })?; - let encrypted_api_key = encrypt_python_fernet_plaintext(encryption_key, access_token) - .map_err(|err| GatewayError::Internal(err.to_string()))?; + let encrypted_api_key = + self.seal_provider_catalog_key_api_key(&transport.provider.id, key_id, access_token)?; let encrypted_auth_config = entry .metadata .as_ref() .map(|value| serde_json::to_string(value)) .transpose() .map_err(|err| GatewayError::Internal(err.to_string()))? - .map(|value| encrypt_python_fernet_plaintext(encryption_key, value.as_str())) - .transpose() - .map_err(|err| GatewayError::Internal(err.to_string()))?; - let requires_fenced_persistence = - provider_transport::supports_local_oauth_request_auth_resolution(transport); - if requires_fenced_persistence - && (expected_credential_fence.is_none() || encrypted_auth_config.is_none()) - { - return Err(GatewayError::Internal( + .map(|value| { + self.seal_provider_catalog_key_auth_config( + &transport.provider.id, + key_id, + value.as_str(), + ) + }) + .transpose()?; + let expected_credential_fence = expected_credential_fence.ok_or_else(|| { + GatewayError::Internal( "OAuth refresh persistence is missing its credential fence".to_string(), - )); - } + ) + })?; + let encrypted_auth_config = encrypted_auth_config.ok_or_else(|| { + GatewayError::Internal( + "OAuth refresh persistence is missing its auth_config".to_string(), + ) + })?; let Some(mut latest_key) = self .data @@ -2559,15 +2674,14 @@ impl AppState { return Ok(()); }; - let observed_credential_matches = expected_credential_fence.is_none_or(|expected| { - latest_key.encrypted_auth_config.as_deref() - == Some(expected.encrypted_auth_config.as_str()) - && latest_key.encrypted_api_key == expected.credential.encrypted_api_key - && latest_key.auth_type == expected.credential.auth_type - && latest_key.provider_id == expected.credential.provider_id - }); + let observed_credential_matches = latest_key.encrypted_auth_config.as_deref() + == Some(expected_credential_fence.encrypted_auth_config.as_str()) + && latest_key.encrypted_api_key + == expected_credential_fence.credential.encrypted_api_key + && latest_key.auth_type == expected_credential_fence.credential.auth_type + && latest_key.provider_id == expected_credential_fence.credential.provider_id; latest_key.encrypted_api_key = Some(encrypted_api_key.clone()); - latest_key.encrypted_auth_config = encrypted_auth_config.clone(); + latest_key.encrypted_auth_config = Some(encrypted_auth_config.clone()); latest_key.expires_at_unix_secs = entry.expires_at_unix_secs; let (oauth_invalid_at_unix_secs, oauth_invalid_reason) = local_oauth_refresh_success_invalid_state(&latest_key); @@ -2583,70 +2697,33 @@ impl AppState { let current_status_snapshot = latest_key.status_snapshot.take(); latest_key.status_snapshot = sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key); - let used_fenced_persistence = - expected_credential_fence.is_some() && encrypted_auth_config.is_some(); - let updated = if let (Some(expected_credential_fence), Some(encrypted_auth_config)) = - (expected_credential_fence, encrypted_auth_config.as_deref()) - { - if !observed_credential_matches { - false - } else { - self.compare_and_update_provider_catalog_key_oauth_runtime_state( - &ProviderCatalogKeyOAuthRuntimeStateCasUpdate { - key_id: key_id.to_string(), - expected_encrypted_auth_config: Some( - expected_credential_fence.encrypted_auth_config.clone(), - ), - expected_credential: Some(expected_credential_fence.credential.clone()), - expected_upstream_metadata_namespace: None, - encrypted_auth_config: encrypted_auth_config.to_string(), - encrypted_api_key_update: Some(encrypted_api_key.clone()), - expires_at_unix_secs_update: Some(entry.expires_at_unix_secs), - oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs, - oauth_invalid_reason: latest_key.oauth_invalid_reason.clone(), - upstream_metadata_patch: None, - upstream_metadata_namespace_to_remove: None, - status_snapshot_patch: provider_key_oauth_status_snapshot_update( - &latest_key, - ) - .status_snapshot_patch, - reset_error_count: false, - updated_at_unix_secs: latest_key.updated_at_unix_secs, - }, - ) - .await? - } + let updated = if !observed_credential_matches { + false } else { - let mut updated = self - .update_provider_catalog_key_oauth_credentials( - key_id, - &encrypted_api_key, - encrypted_auth_config.as_deref(), - entry.expires_at_unix_secs, - ) - .await?; - if updated { - updated = self - .update_provider_catalog_key_oauth_runtime_state( - key_id, - latest_key.oauth_invalid_at_unix_secs, - latest_key.oauth_invalid_reason.as_deref(), - None, - latest_key.updated_at_unix_secs, - ) - .await?; - } - if updated { - updated = self - .update_provider_catalog_key_status_snapshot( - &provider_key_oauth_status_snapshot_update(&latest_key), - ) - .await?; - self.clear_provider_transport_snapshot_cache(); - } - updated + self.compare_and_update_provider_catalog_key_oauth_runtime_state( + &ProviderCatalogKeyOAuthRuntimeStateCasUpdate { + key_id: key_id.to_string(), + expected_encrypted_auth_config: Some( + expected_credential_fence.encrypted_auth_config.clone(), + ), + expected_credential: Some(expected_credential_fence.credential.clone()), + expected_upstream_metadata_namespace: None, + encrypted_auth_config, + encrypted_api_key_update: Some(encrypted_api_key.clone()), + expires_at_unix_secs_update: Some(entry.expires_at_unix_secs), + oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs, + oauth_invalid_reason: latest_key.oauth_invalid_reason.clone(), + upstream_metadata_patch: None, + upstream_metadata_namespace_to_remove: None, + status_snapshot_patch: provider_key_oauth_status_snapshot_update(&latest_key) + .status_snapshot_patch, + reset_error_count: false, + updated_at_unix_secs: latest_key.updated_at_unix_secs, + }, + ) + .await? }; - if !updated && (requires_fenced_persistence || used_fenced_persistence) { + if !updated { return Err(GatewayError::Internal( "OAuth credential changed during refresh persistence".to_string(), )); @@ -2993,7 +3070,6 @@ impl AppState { key_id, latest_key.oauth_invalid_at_unix_secs, latest_key.oauth_invalid_reason.as_deref(), - None, latest_key.updated_at_unix_secs, ) .await?; @@ -3151,7 +3227,7 @@ impl AppState { provider_type, request_id = %request.request_id, method = %plan.method, - token_url = %plan.url, + token_origin = %crate::handlers::shared::security_log_url_origin(&plan.url), content_type = plan.content_type.as_deref().unwrap_or("-"), body_bytes_len = ?request.body_bytes.as_ref().map(Vec::len), json_body_present = request.json_body.is_some(), @@ -3191,15 +3267,22 @@ impl AppState { .unwrap_or("-"), "gateway local oauth execution request prepared" ); - let result = - crate::execution_runtime::execute_execution_runtime_sync_plan(self, None, &plan) - .await - .map_err( - |err| provider_transport::LocalOAuthRefreshError::InvalidResponse { - provider_type, - message: err.into_message(), - }, - )?; + let bounded_plan = crate::execution_runtime::transport::with_upstream_response_body_limit( + &plan, + LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ); + let result = crate::execution_runtime::execute_execution_runtime_sync_plan( + self, + None, + &bounded_plan, + ) + .await + .map_err( + |err| provider_transport::LocalOAuthRefreshError::InvalidResponse { + provider_type, + message: err.into_message(), + }, + )?; let response_body_text = local_oauth_execution_body_text(&result); if (200..300).contains(&result.status_code) { tracing::info!( @@ -3351,31 +3434,20 @@ fn local_oauth_execution_body_bytes( headers: &BTreeMap, body: &aether_contracts::ResponseBody, ) -> Option> { - let bytes = body - .body_bytes_b64 - .as_deref() - .and_then(|value| STANDARD.decode(value).ok())?; - let encoding = headers - .get("content-encoding") - .map(String::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.to_ascii_lowercase()); - match encoding.as_deref() { - Some("gzip") => { - let mut decoder = GzDecoder::new(bytes.as_slice()); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) - } - Some("deflate") => { - let mut decoder = DeflateDecoder::new(bytes.as_slice()); - let mut out = Vec::new(); - decoder.read_to_end(&mut out).ok()?; - Some(out) - } - _ => Some(bytes), - } + let bytes = body.body_bytes_b64.as_deref().and_then(|value| { + crate::execution_runtime::transport::decode_base64_body_with_limit( + value, + LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + })?; + crate::execution_runtime::transport::decode_response_body_bytes_with_limit( + headers, + &bytes, + LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES, + ) + .ok() + .map(std::borrow::Cow::into_owned) } fn local_oauth_request_uses_direct_client(url: &str) -> bool { @@ -3394,13 +3466,10 @@ fn local_oauth_request_uses_direct_client(url: &str) -> bool { #[cfg(test)] mod tests { use std::sync::atomic::{AtomicUsize, Ordering}; - use std::sync::Arc; + use std::sync::{Arc, OnceLock}; use std::time::{Duration, Instant}; - use aether_crypto::{ - decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, - DEVELOPMENT_ENCRYPTION_KEY, - }; + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyListQuery, @@ -3456,10 +3525,40 @@ mod tests { .expect("endpoint transport should build") } - fn sample_key() -> StoredProviderCatalogKey { + fn sample_key_for_provider(provider_id: &str) -> StoredProviderCatalogKey { + let encrypted_api_key = if provider_id == "provider-1" { + static ENCRYPTED_API_KEY: OnceLock = OnceLock::new(); + ENCRYPTED_API_KEY + .get_or_init(|| { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + credential_state + .seal_provider_catalog_key_api_key( + "provider-1", + "key-1", + "plain-upstream-key", + ) + .expect("api key should encrypt") + }) + .clone() + } else { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + credential_state + .seal_provider_catalog_key_api_key(provider_id, "key-1", "plain-upstream-key") + .expect("api key should encrypt") + }; StoredProviderCatalogKey::new( "key-1".to_string(), - "provider-1".to_string(), + provider_id.to_string(), "default".to_string(), "api_key".to_string(), None, @@ -3468,7 +3567,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:chat"])), - "plain-upstream-key".to_string(), + encrypted_api_key, None, None, Some(json!({"openai:chat": 1})), @@ -3480,6 +3579,10 @@ mod tests { .expect("key transport should build") } + fn sample_key() -> StoredProviderCatalogKey { + sample_key_for_provider("provider-1") + } + fn codex_oauth_state( auth_config: &serde_json::Value, access_token: &str, @@ -3548,12 +3651,18 @@ mod tests { "private_key": "TEST-PRIVATE-KEY", "project_id": "demo-project" }); - let encrypted_auth_config = - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &auth_config.to_string()) - .expect("Vertex auth config should encrypt"); - let encrypted_api_key = - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__") - .expect("Vertex placeholder should encrypt"); + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_auth_config = credential_state + .seal_provider_catalog_key_auth_config("provider-1", "key-1", &auth_config.to_string()) + .expect("Vertex auth config should encrypt"); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key("provider-1", "key-1", "__placeholder__") + .expect("Vertex placeholder should encrypt"); let key = StoredProviderCatalogKey::new( "key-1".to_string(), "provider-1".to_string(), @@ -3631,7 +3740,7 @@ mod tests { )); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( repository, - "test-encryption-key", + DEVELOPMENT_ENCRYPTION_KEY, ) .with_system_config_values_for_tests(vec![( "enable_format_conversion".to_string(), @@ -3842,7 +3951,7 @@ mod tests { ))); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( reader.clone(), - "test-encryption-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -3906,7 +4015,7 @@ mod tests { )); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( reader.clone(), - "test-encryption-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -4111,7 +4220,7 @@ mod tests { ))); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( reader.clone(), - "test-encryption-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -4149,8 +4258,7 @@ mod tests { #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn transport_snapshot_error_is_broadcast_to_all_followers() { const REQUESTS: usize = 64; - let mut mismatched_key = sample_key(); - mismatched_key.provider_id = "provider-other".to_string(); + let mismatched_key = sample_key_for_provider("provider-other"); let inner = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider()], vec![sample_endpoint()], @@ -4161,7 +4269,7 @@ mod tests { ))); let data_state = GatewayDataState::with_provider_transport_reader_for_tests( reader.clone(), - "test-encryption-key", + DEVELOPMENT_ENCRYPTION_KEY, ); let state = AppState::new() .expect("state should build") @@ -4400,6 +4508,18 @@ mod tests { ); } + #[test] + fn local_oauth_diagnostics_do_not_reflect_unknown_response_credentials() { + let body = r#"{"error":{"message":"authorization=Bearer upstream-secret https://user:pass@example.test?q=secret","type":"invalid_request_error","code":"unexpected"}}"#; + + let normalized = super::normalize_local_oauth_refresh_error_message(Some(502), Some(body)); + assert_eq!(normalized, "HTTP 502"); + assert_eq!(super::local_oauth_log_excerpt(body), "[redacted]"); + for secret in ["upstream-secret", "user:pass", "q=secret"] { + assert!(!normalized.contains(secret), "leaked {secret}"); + } + } + #[test] fn local_refresh_failure_is_appended_to_access_token_expired_marker() { assert_eq!( @@ -4625,22 +4745,14 @@ mod tests { .expect("key should reload") .pop() .expect("key should remain"); - let access_token = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored - .encrypted_api_key - .as_deref() - .expect("access token should persist"), - ) - .expect("access token should decrypt"); - let auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored - .encrypted_auth_config - .as_deref() - .expect("auth config should persist"), - ) - .expect("auth config should decrypt"); + let access_token = state + .decrypt_provider_catalog_key_api_key(&stored) + .expect("access token should decrypt") + .expect("access token should persist"); + let auth_config = state + .decrypt_provider_catalog_key_auth_config(&stored) + .expect("auth config should decrypt") + .expect("auth config should persist"); let auth_config: serde_json::Value = serde_json::from_str(&auth_config).expect("auth config should parse"); @@ -4957,7 +5069,6 @@ mod tests { "key-1", Some(1), Some("[OAUTH_EXPIRED] response raced with quota persistence"), - None, Some(1), ) .await @@ -5124,7 +5235,6 @@ mod tests { "key-1", Some(1), Some("[ACCOUNT_BLOCK] account deactivated"), - None, Some(1), ) .await @@ -5171,7 +5281,6 @@ mod tests { "key-1", Some(1), Some("[OAUTH_EXPIRED] current generation invalid"), - None, Some(1), ) .await @@ -5226,7 +5335,6 @@ mod tests { "key-1", Some(1), Some("[OAUTH_EXPIRED] replacement generation invalid"), - None, Some(1), ) .await @@ -5253,7 +5361,6 @@ mod tests { "key-1", Some(1), Some("[OAUTH_EXPIRED] replacement generation invalid"), - None, Some(1), ) .await @@ -5294,7 +5401,6 @@ mod tests { "key-1", Some(1), Some("[OAUTH_EXPIRED] unordered response"), - None, Some(1), ) .await diff --git a/apps/aether-gateway/src/state/proxy.rs b/apps/aether-gateway/src/state/proxy.rs index 84362f212..07f49b733 100644 --- a/apps/aether-gateway/src/state/proxy.rs +++ b/apps/aether-gateway/src/state/proxy.rs @@ -1,15 +1,484 @@ -use aether_contracts::ProxySnapshot; -use aether_data::repository::proxy_nodes::proxy_node_accepts_new_tunnels; +use aether_contracts::{ProxySnapshot, PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY}; +use aether_crypto::looks_like_python_fernet_ciphertext; +use aether_data::repository::proxy_nodes::{ + proxy_node_accepts_new_tunnels, ProxyNodeRegistrationMutation, +}; use serde_json::{json, Map, Value}; +use std::future::Future; use super::AppState; +use crate::data::GatewayDataState; +use crate::handlers::shared::{ + open_runtime_secret_payload_with_encryption_key, runtime_secret_payload_is_sealed, + seal_runtime_secret_payload_with_encryption_key, +}; use crate::provider_transport::{GatewayProviderTransportSnapshot, TransportTunnelAffinityLookup}; +use crate::GatewayError; const TUNNEL_BASE_URL_EXTRA_KEY: &str = "tunnel_base_url"; const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id"; const TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY: &str = "tunnel_owner_observed_at_unix_secs"; +const PROXY_NODE_PASSWORD_LEGACY_SECRET_PURPOSE: &str = "proxy-node-password"; +const PROXY_TUNNEL_PSK_LEGACY_SECRET_PURPOSE: &str = "proxy-node-tunnel-psk"; +const PROXY_NODE_SECRET_ENVELOPE_FAMILY_PREFIX: &str = "aether-proxy-node-secret-"; +const PROXY_NODE_SECRET_V2_PREFIX: &str = "aether-proxy-node-secret-v2:"; +const RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX: &str = "aether-runtime-secret-"; +const PROXY_NODE_BOUND_PURPOSE_VERSION: &str = "proxy-node-secret-bound-v2"; +const PROXY_NODE_PASSWORD_SCOPE: &str = "manual-proxy"; +const PROXY_NODE_PASSWORD_FIELD: &str = "password"; +const PROXY_TUNNEL_PSK_SCOPE: &str = "tunnel-security"; +const PROXY_TUNNEL_PSK_FIELD: &str = "pre-shared-key"; +const PROXY_NODE_SECRET_MIGRATION_RETRIES: usize = 8; +const PROXY_UNAVAILABLE_REASON_EXTRA_KEY: &str = "aether_proxy_unavailable_reason"; + +pub(crate) fn unavailable_proxy_snapshot(reason: &str) -> ProxySnapshot { + let extra = Map::from_iter([( + PROXY_UNAVAILABLE_REASON_EXTRA_KEY.to_string(), + Value::String(reason.to_string()), + )]); + ProxySnapshot { + enabled: Some(true), + mode: Some("unavailable".to_string()), + node_id: None, + label: None, + url: None, + extra: Some(Value::Object(extra)), + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ProxyTunnelPskBinding { + pub key: String, + pub tunnel_generation: String, +} + +pub(crate) async fn decrypt_or_migrate_proxy_tunnel_psk( + data: &GatewayDataState, + node_id: &str, +) -> Result, aether_data::DataLayerError> { + Ok(decrypt_or_migrate_proxy_tunnel_psk_binding(data, node_id) + .await? + .map(|binding| binding.key)) +} + +pub(crate) async fn decrypt_or_migrate_proxy_tunnel_psk_binding( + data: &GatewayDataState, + node_id: &str, +) -> Result, aether_data::DataLayerError> { + for _ in 0..PROXY_NODE_SECRET_MIGRATION_RETRIES { + let Some(node) = data.find_proxy_node(node_id).await? else { + return Ok(None); + }; + let persistent_node_id = node.id.clone(); + let tunnel_generation = node.tunnel_generation.clone(); + let Some(observed_metadata) = node.proxy_metadata else { + return Ok(None); + }; + let encrypted = observed_metadata + .pointer("/tunnel_security/encryption_key_encrypted") + .map(|value| value.as_str().ok_or_else(proxy_tunnel_psk_storage_error)) + .transpose()?; + let legacy = observed_metadata + .pointer("/tunnel_security/encryption_key") + .map(|value| value.as_str().ok_or_else(proxy_tunnel_psk_storage_error)) + .transpose()?; + + if let Some(stored) = encrypted { + let (plaintext, replacement_ciphertext) = if proxy_node_secret_is_v2(stored) { + ( + open_proxy_node_secret_v2_with_encryption_key( + data.encryption_key(), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + &persistent_node_id, + stored, + ) + .ok_or_else(proxy_tunnel_psk_storage_error)?, + None, + ) + } else if runtime_secret_payload_is_sealed(stored) { + let plaintext = open_runtime_secret_payload_with_encryption_key( + data.encryption_key(), + PROXY_TUNNEL_PSK_LEGACY_SECRET_PURPOSE, + stored, + ) + .ok_or_else(proxy_tunnel_psk_storage_error)?; + let replacement = seal_proxy_node_secret_v2_with_encryption_key( + data.encryption_key(), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + &persistent_node_id, + &plaintext, + ) + .ok_or_else(proxy_tunnel_psk_migration_error)?; + (plaintext, Some(replacement)) + } else { + return Err(proxy_tunnel_psk_storage_error()); + }; + aether_contracts::tunnel_security::decode_psk(&plaintext) + .map_err(|_| proxy_tunnel_psk_storage_error())?; + if replacement_ciphertext.is_none() && legacy.is_none() { + return Ok(Some(ProxyTunnelPskBinding { + key: plaintext, + tunnel_generation, + })); + } + + let replacement = proxy_metadata_with_encrypted_tunnel_psk( + observed_metadata.clone(), + replacement_ciphertext.unwrap_or_else(|| stored.to_string()), + )?; + if data + .compare_and_set_proxy_node_metadata( + &persistent_node_id, + &observed_metadata, + &replacement, + ) + .await? + { + return Ok(Some(ProxyTunnelPskBinding { + key: plaintext, + tunnel_generation, + })); + } + continue; + } + + let Some(legacy) = legacy else { + return Ok(None); + }; + if looks_like_python_fernet_ciphertext(legacy) + || legacy.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX) + || legacy.starts_with(PROXY_NODE_SECRET_ENVELOPE_FAMILY_PREFIX) + { + return Err(proxy_tunnel_psk_storage_error()); + } + aether_contracts::tunnel_security::decode_psk(legacy) + .map_err(|_| proxy_tunnel_psk_storage_error())?; + let encrypted = seal_proxy_node_secret_v2_with_encryption_key( + data.encryption_key(), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + &persistent_node_id, + legacy, + ) + .ok_or_else(proxy_tunnel_psk_migration_error)?; + let replacement = + proxy_metadata_with_encrypted_tunnel_psk(observed_metadata.clone(), encrypted)?; + if data + .compare_and_set_proxy_node_metadata( + &persistent_node_id, + &observed_metadata, + &replacement, + ) + .await? + { + return Ok(Some(ProxyTunnelPskBinding { + key: legacy.to_string(), + tunnel_generation, + })); + } + } + + Err(proxy_tunnel_psk_migration_error()) +} + +fn proxy_metadata_with_encrypted_tunnel_psk( + mut metadata: Value, + encrypted: String, +) -> Result { + let tunnel_security = metadata + .as_object_mut() + .and_then(|metadata| metadata.get_mut("tunnel_security")) + .and_then(Value::as_object_mut) + .ok_or_else(proxy_tunnel_psk_storage_error)?; + tunnel_security.remove("encryption_key"); + tunnel_security.insert( + "encryption_key_encrypted".to_string(), + Value::String(encrypted), + ); + Ok(metadata) +} + +fn proxy_tunnel_psk_storage_error() -> aether_data::DataLayerError { + aether_data::DataLayerError::UnexpectedValue( + "stored proxy tunnel security key cannot be decrypted".to_string(), + ) +} + +fn proxy_tunnel_psk_migration_error() -> aether_data::DataLayerError { + aether_data::DataLayerError::UnexpectedValue( + "proxy tunnel security key migration did not stabilize".to_string(), + ) +} + +fn proxy_node_password_error() -> GatewayError { + GatewayError::Internal("stored proxy node password cannot be decrypted".to_string()) +} + +fn proxy_node_secret_is_v2(value: &str) -> bool { + value.starts_with(PROXY_NODE_SECRET_V2_PREFIX) +} + +fn proxy_node_bound_secret_purpose(scope: &str, field: &str, node_id: &str) -> String { + format!( + "{PROXY_NODE_BOUND_PURPOSE_VERSION}\0scope-bytes={}\0{scope}\0field-bytes={}\0{field}\0node-id-bytes={}\0{node_id}", + scope.len(), + field.len(), + node_id.len(), + ) +} + +fn seal_proxy_node_secret_v2_with_encryption_key( + encryption_key: Option<&str>, + scope: &str, + field: &str, + node_id: &str, + plaintext: &str, +) -> Option { + let purpose = proxy_node_bound_secret_purpose(scope, field, node_id); + seal_runtime_secret_payload_with_encryption_key(encryption_key, &purpose, plaintext) + .map(|sealed| format!("{PROXY_NODE_SECRET_V2_PREFIX}{sealed}")) +} + +fn open_proxy_node_secret_v2_with_encryption_key( + encryption_key: Option<&str>, + scope: &str, + field: &str, + node_id: &str, + stored: &str, +) -> Option { + // The distinct outer envelope is security-significant. A v2 binding + // failure must never be retried with the unbound legacy purpose. + let sealed = stored.strip_prefix(PROXY_NODE_SECRET_V2_PREFIX)?; + let purpose = proxy_node_bound_secret_purpose(scope, field, node_id); + open_runtime_secret_payload_with_encryption_key(encryption_key, &purpose, sealed) +} + +fn incoming_proxy_node_secret_is_ciphertext(value: &str) -> bool { + let value = value.trim(); + value.starts_with(PROXY_NODE_SECRET_ENVELOPE_FAMILY_PREFIX) + || value.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX) + || looks_like_python_fernet_ciphertext(value) +} impl AppState { + pub(super) async fn register_proxy_node_with_bound_secrets( + &self, + mutation: &ProxyNodeRegistrationMutation, + ) -> Result, GatewayError> { + let mut new_node_id = mutation + .node_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + + for _ in 0..PROXY_NODE_SECRET_MIGRATION_RETRIES { + let nodes = self + .data + .list_proxy_nodes() + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let node_id = proxy_node_registration_identity(&nodes, mutation) + .unwrap_or_else(|| new_node_id.clone()); + let protected = self.protect_proxy_node_registration_mutation(&node_id, mutation)?; + match self.data.register_proxy_node(&protected).await { + Ok(Some(node)) if node.id == node_id => return Ok(Some(node)), + Ok(Some(_)) => { + return Err(GatewayError::Internal( + "proxy node repository changed the protected node identity".to_string(), + )); + } + Ok(None) => return Ok(None), + Err(error) => { + let latest = self + .data + .list_proxy_nodes() + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let Some(latest_node_id) = proxy_node_registration_identity(&latest, mutation) + else { + return Err(GatewayError::Internal(error.to_string())); + }; + if latest_node_id == node_id { + return Err(GatewayError::Internal(error.to_string())); + } + new_node_id = latest_node_id; + } + } + } + + Err(GatewayError::Internal( + "proxy node registration identity did not stabilize".to_string(), + )) + } + + pub(super) fn protect_proxy_node_registration_mutation( + &self, + node_id: &str, + mutation: &ProxyNodeRegistrationMutation, + ) -> Result { + let mut protected = mutation.clone(); + protected.node_id = Some(node_id.to_string()); + let Some(metadata) = protected.proxy_metadata.as_mut() else { + return Ok(protected); + }; + let Some(tunnel_security) = metadata + .as_object_mut() + .and_then(|metadata| metadata.get_mut("tunnel_security")) + .and_then(Value::as_object_mut) + else { + return Ok(protected); + }; + if tunnel_security.contains_key("encryption_key_encrypted") { + return Err(GatewayError::Internal( + "proxy tunnel security ciphertext must be created by the gateway".to_string(), + )); + } + let plaintext = tunnel_security + .remove("encryption_key") + .map(|value| { + value.as_str().map(ToOwned::to_owned).ok_or_else(|| { + GatewayError::Internal( + "proxy tunnel security key has an invalid stored representation" + .to_string(), + ) + }) + }) + .transpose()?; + let Some(plaintext) = plaintext else { + return Ok(protected); + }; + if incoming_proxy_node_secret_is_ciphertext(&plaintext) { + return Err(GatewayError::Internal( + "proxy tunnel security ciphertext must be created by the gateway".to_string(), + )); + } + aether_contracts::tunnel_security::decode_psk(&plaintext).map_err(|_| { + GatewayError::Internal("proxy tunnel security key is invalid".to_string()) + })?; + let encrypted = seal_proxy_node_secret_v2_with_encryption_key( + self.encryption_key(), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + node_id, + plaintext.as_str(), + ) + .ok_or_else(|| { + GatewayError::Internal("proxy tunnel security encryption is unavailable".to_string()) + })?; + tunnel_security.insert( + "encryption_key_encrypted".to_string(), + Value::String(encrypted), + ); + Ok(protected) + } + + pub(super) fn protect_proxy_node_password( + &self, + node_id: &str, + plaintext: &str, + ) -> Result { + if incoming_proxy_node_secret_is_ciphertext(plaintext) { + return Err(GatewayError::Internal( + "proxy node password ciphertext must be created by the gateway".to_string(), + )); + } + seal_proxy_node_secret_v2_with_encryption_key( + self.encryption_key(), + PROXY_NODE_PASSWORD_SCOPE, + PROXY_NODE_PASSWORD_FIELD, + node_id, + plaintext, + ) + .ok_or_else(|| { + GatewayError::Internal("proxy node password encryption is unavailable".to_string()) + }) + } + + pub(crate) async fn decrypt_proxy_node_password( + &self, + node_id: &str, + ) -> Result, GatewayError> { + self.decrypt_proxy_node_password_with_before_compare(node_id, || async {}) + .await + } + + async fn decrypt_proxy_node_password_with_before_compare( + &self, + node_id: &str, + before_compare: BeforeCompare, + ) -> Result, GatewayError> + where + BeforeCompare: Fn() -> CompareFuture, + CompareFuture: Future, + { + for _ in 0..PROXY_NODE_SECRET_MIGRATION_RETRIES { + let Some(node) = self.find_proxy_node(node_id).await? else { + return Ok(None); + }; + let persistent_node_id = node.id.clone(); + let Some(observed) = node.proxy_password else { + return Ok(None); + }; + if proxy_node_secret_is_v2(&observed) { + return open_proxy_node_secret_v2_with_encryption_key( + self.encryption_key(), + PROXY_NODE_PASSWORD_SCOPE, + PROXY_NODE_PASSWORD_FIELD, + &persistent_node_id, + &observed, + ) + .map(Some) + .ok_or_else(proxy_node_password_error); + } + let observed_shape = observed.trim(); + if observed_shape.starts_with(PROXY_NODE_SECRET_ENVELOPE_FAMILY_PREFIX) + || (observed_shape.starts_with(RUNTIME_SECRET_ENVELOPE_FAMILY_PREFIX) + && !runtime_secret_payload_is_sealed(&observed)) + || looks_like_python_fernet_ciphertext(&observed) + { + return Err(proxy_node_password_error()); + } + + let plaintext = if runtime_secret_payload_is_sealed(&observed) { + open_runtime_secret_payload_with_encryption_key( + self.encryption_key(), + PROXY_NODE_PASSWORD_LEGACY_SECRET_PURPOSE, + &observed, + ) + .ok_or_else(proxy_node_password_error)? + } else { + observed.clone() + }; + let encrypted = seal_proxy_node_secret_v2_with_encryption_key( + self.encryption_key(), + PROXY_NODE_PASSWORD_SCOPE, + PROXY_NODE_PASSWORD_FIELD, + &persistent_node_id, + &plaintext, + ) + .ok_or_else(|| { + GatewayError::Internal("proxy node password encryption is unavailable".to_string()) + })?; + before_compare().await; + if self + .data + .compare_and_set_proxy_node_password(&persistent_node_id, &observed, &encrypted) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + return Ok(Some(plaintext)); + } + } + + Err(GatewayError::Internal( + "proxy node password migration did not stabilize".to_string(), + )) + } + pub(crate) async fn read_system_proxy_node_id(&self) -> Option { self.read_system_config_json_value("system_proxy_node_id") .await @@ -54,6 +523,10 @@ impl AppState { } else if !self.tunnel.has_local_proxy(node_id) { return None; } + extra.insert( + PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY.to_string(), + Value::String(node.tunnel_generation.clone()), + ); return Some(ProxySnapshot { enabled: Some(true), mode: Some("tunnel".to_string()), @@ -70,29 +543,41 @@ impl AppState { if !node.is_manual { return None; } + let proxy_password = match self.decrypt_proxy_node_password(&node.id).await { + Ok(password) => password, + Err(_) => return None, + }; let proxy_url = node .proxy_url .as_deref() .map(str::trim) .filter(|value| !value.is_empty())?; + let proxy_url = proxy_url_with_node_auth( + proxy_url, + node.proxy_username.as_deref(), + proxy_password.as_deref(), + )?; + let mut extra = Map::new(); + extra.insert( + PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY.to_string(), + Value::String(node.tunnel_generation.clone()), + ); Some(ProxySnapshot { enabled: Some(true), - mode: proxy_mode_from_url(Some(proxy_url)), + mode: proxy_mode_from_url(Some(&proxy_url)), node_id: Some(node.id), label: Some(node.name), - url: proxy_url_with_node_auth( - proxy_url, - node.proxy_username.as_deref(), - node.proxy_password.as_deref(), - ) - .or_else(|| Some(proxy_url.to_string())), - extra: None, + url: Some(proxy_url), + extra: Some(Value::Object(extra)), }) } pub(crate) async fn resolve_system_proxy_snapshot(&self) -> Option { let node_id = self.read_system_proxy_node_id().await; - self.resolve_proxy_node_snapshot(node_id.as_deref()).await + let node_id = node_id.as_deref()?; + self.resolve_proxy_node_snapshot(Some(node_id)) + .await + .or_else(|| Some(unavailable_proxy_snapshot("system_proxy_node_unavailable"))) } pub(crate) async fn resolve_transport_proxy_snapshot_with_tunnel_affinity( @@ -122,19 +607,31 @@ impl AppState { return None; } - let node_id = json_string_field(object, "node_id"); - if let Some(snapshot) = self.resolve_proxy_node_snapshot(node_id.as_deref()).await { - return Some(snapshot); - } - if let Some(node_id) = node_id.as_deref() { - if !proxy_object_has_inline_url(object) - && self.find_proxy_node(node_id).await.ok().flatten().is_some() - { - return None; + if json_field_is_explicit(object, "node_id") { + let node_id = json_string_field(object, "node_id"); + if let Some(snapshot) = self.resolve_proxy_node_snapshot(node_id.as_deref()).await { + return Some(snapshot); } + return Some(unavailable_proxy_snapshot( + "configured_proxy_node_unavailable", + )); } - proxy_snapshot_from_object(object) + if json_field_is_explicit(object, "url") || json_field_is_explicit(object, "proxy_url") { + let Some(snapshot) = proxy_snapshot_from_object(object) else { + return Some(unavailable_proxy_snapshot( + "configured_proxy_url_unavailable", + )); + }; + if snapshot.url.is_none() { + return Some(unavailable_proxy_snapshot( + "configured_proxy_url_unavailable", + )); + } + return Some(snapshot); + } + + None } async fn resolve_transport_proxy_with_source_with_tunnel_affinity( @@ -169,6 +666,22 @@ impl AppState { } } +fn proxy_node_registration_identity( + nodes: &[aether_data::repository::proxy_nodes::StoredProxyNode], + mutation: &ProxyNodeRegistrationMutation, +) -> Option { + nodes + .iter() + .filter(|node| !node.is_manual && node.ip == mutation.ip && node.port == mutation.port) + .min_by(|left, right| { + left.created_at_unix_ms + .unwrap_or(u64::MAX) + .cmp(&right.created_at_unix_ms.unwrap_or(u64::MAX)) + .then(left.id.cmp(&right.id)) + }) + .map(|node| node.id.clone()) +} + fn proxy_enabled(object: &Map) -> bool { object .get("enabled") @@ -180,7 +693,15 @@ fn proxy_snapshot_from_object(object: &Map) -> Option) -> Option) -> Option) -> bool { - json_string_field(object, "url").is_some() || json_string_field(object, "proxy_url").is_some() +fn json_field_is_explicit(object: &Map, key: &str) -> bool { + object.get(key).is_some_and(|value| !value.is_null()) } fn json_string_field(object: &Map, key: &str) -> Option { @@ -224,6 +752,13 @@ fn json_string_field(object: &Map, key: &str) -> Option { .map(ToOwned::to_owned) } +fn json_proxy_credential_field<'a>(object: &'a Map, key: &str) -> Option<&'a str> { + object + .get(key) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) +} + fn proxy_mode_from_url(proxy_url: Option<&str>) -> Option { let proxy_url = proxy_url?.trim(); if proxy_url.is_empty() { @@ -245,12 +780,21 @@ fn proxy_url_with_node_auth( username: Option<&str>, password: Option<&str>, ) -> Option { - let username = username.map(str::trim).filter(|value| !value.is_empty())?; + let username = username.filter(|value| !value.is_empty()); + let password = password.filter(|value| !value.is_empty()); let mut parsed = url::Url::parse(proxy_url).ok()?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + { + return None; + } + if username.is_none() && password.is_none() { + return Some(parsed.to_string()); + } + let username = username.unwrap_or(""); if parsed.set_username(username).is_err() { return None; } - let password = password.map(str::trim).filter(|value| !value.is_empty()); if parsed.set_password(password).is_err() { return None; } @@ -259,20 +803,751 @@ fn proxy_url_with_node_auth( #[cfg(test)] mod tests { - use std::sync::Arc; + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; - use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; + use aether_data::repository::proxy_nodes::{ + InMemoryProxyNodeRepository, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, + ProxyNodeReadRepository, ProxyNodeRegistrationMutation, ProxyNodeWriteRepository, + StoredProxyNode, + }; + use aether_provider_transport::snapshot::{ + GatewayProviderTransportEndpoint, GatewayProviderTransportKey, + GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, + }; use serde_json::json; + use tokio::sync::Barrier; - use super::proxy_url_with_node_auth; + use super::{ + decrypt_or_migrate_proxy_tunnel_psk, open_proxy_node_secret_v2_with_encryption_key, + proxy_node_bound_secret_purpose, proxy_node_secret_is_v2, proxy_url_with_node_auth, + seal_proxy_node_secret_v2_with_encryption_key, PROXY_NODE_PASSWORD_FIELD, + PROXY_NODE_PASSWORD_LEGACY_SECRET_PURPOSE, PROXY_NODE_PASSWORD_SCOPE, + PROXY_TUNNEL_PSK_FIELD, PROXY_TUNNEL_PSK_LEGACY_SECRET_PURPOSE, PROXY_TUNNEL_PSK_SCOPE, + PROXY_UNAVAILABLE_REASON_EXTRA_KEY, + }; + use crate::handlers::shared::seal_runtime_secret_payload_with_encryption_key; use crate::{data::GatewayDataState, AppState}; + const VALID_TUNNEL_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const ROTATED_TUNNEL_PSK: &str = "CAgICAgICAgICAgICAgICAgICAgICAgICAgICAgICAg="; + + #[test] + fn proxy_node_bound_secret_purpose_has_unambiguous_component_lengths() { + assert_ne!( + proxy_node_bound_secret_purpose("a", "bc", "node"), + proxy_node_bound_secret_purpose("ab", "c", "node"), + ); + assert_ne!( + proxy_node_bound_secret_purpose("scope", "field", "a\0b"), + proxy_node_bound_secret_purpose("scope", "field\0a", "b"), + ); + } + #[test] fn proxy_url_with_node_auth_omits_empty_password_separator() { assert_eq!( proxy_url_with_node_auth("socks5://proxy.example:1080", Some("alice"), None).as_deref(), Some("socks5://alice@proxy.example:1080") ); + assert_eq!( + proxy_url_with_node_auth("http://proxy.example:8080", None, None).as_deref(), + Some("http://proxy.example:8080/") + ); + assert_eq!( + proxy_url_with_node_auth("http://proxy.example:8080", None, Some("secret")).as_deref(), + Some("http://:secret@proxy.example:8080/") + ); + } + + #[tokio::test] + async fn manual_proxy_credentials_fail_closed_when_url_cannot_accept_auth() { + let mut malformed = sample_manual_node("manual-malformed-auth-url"); + malformed.proxy_url = Some("not a proxy url".to_string()); + malformed.proxy_username = Some("alice".to_string()); + + let mut unsupported = sample_manual_node("manual-unsupported-auth-url"); + unsupported.proxy_url = Some("mailto:proxy@example.com".to_string()); + unsupported.proxy_username = Some("alice".to_string()); + + let mut password_only = sample_manual_node("manual-password-only"); + password_only.proxy_username = None; + password_only.proxy_password = Some("legacy-password".to_string()); + + let repository = Arc::new(InMemoryProxyNodeRepository::seed([ + malformed, + unsupported, + password_only, + ])); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + + for node_id in ["manual-malformed-auth-url", "manual-unsupported-auth-url"] { + assert!( + state + .resolve_proxy_node_snapshot(Some(node_id)) + .await + .is_none(), + "configured credentials must not fall back to an unauthenticated URL: {node_id}" + ); + } + let password_only = state + .resolve_proxy_node_snapshot(Some("manual-password-only")) + .await + .expect("legacy password-only proxy should remain usable"); + assert_eq!( + password_only.url.as_deref(), + Some("http://:legacy-password@proxy.example:8080/") + ); + } + + #[tokio::test] + async fn manual_proxy_password_is_encrypted_before_repository_write() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + + let created = state + .create_manual_proxy_node(&ProxyNodeManualCreateMutation { + node_id: None, + name: "manual-proxy".to_string(), + ip: "127.0.0.1".to_string(), + port: 8080, + region: None, + proxy_url: "http://proxy.example:8080".to_string(), + proxy_username: Some("alice".to_string()), + proxy_password: Some("manual-password-marker".to_string()), + registered_by: None, + }) + .await + .expect("manual proxy create should succeed") + .expect("manual proxy should be stored"); + let stored = repository + .find_proxy_node(&created.id) + .await + .expect("stored proxy should read") + .expect("stored proxy should exist") + .proxy_password + .expect("stored proxy password should exist"); + + assert!(proxy_node_secret_is_v2(&stored)); + assert!(!stored.contains("manual-password-marker")); + assert_eq!( + state + .decrypt_proxy_node_password(&created.id) + .await + .expect("stored password should decrypt") + .as_deref(), + Some("manual-password-marker") + ); + } + + #[tokio::test] + async fn manual_proxy_duplicate_node_id_fails_without_overwriting_bound_password() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let mut first = sample_manual_create_mutation( + "manual-duplicate-first", + "127.0.0.41", + 8141, + "first-password", + ); + first.node_id = Some("fixed-manual-node-id".to_string()); + let created = state + .create_manual_proxy_node(&first) + .await + .expect("first proxy should create") + .expect("first proxy should persist"); + assert_eq!(created.id, "fixed-manual-node-id"); + + let mut duplicate = sample_manual_create_mutation( + "manual-duplicate-second", + "127.0.0.42", + 8142, + "replacement-password", + ); + duplicate.node_id = Some("fixed-manual-node-id".to_string()); + assert!(state.create_manual_proxy_node(&duplicate).await.is_err()); + assert_eq!( + state + .decrypt_proxy_node_password("fixed-manual-node-id") + .await + .expect("original password should remain readable") + .as_deref(), + Some("first-password") + ); + let nodes = repository + .list_proxy_nodes() + .await + .expect("repository should list"); + assert_eq!(nodes.len(), 1); + assert_eq!(nodes[0].name, "manual-duplicate-first"); + } + + #[tokio::test] + async fn administrator_password_rotation_wins_legacy_migration_race() { + let mut node = sample_manual_node("manual-race"); + node.proxy_password = Some("legacy-password".to_string()); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let reached_compare = Arc::new(Barrier::new(2)); + let resume_compare = Arc::new(Barrier::new(2)); + let first_compare = Arc::new(AtomicBool::new(true)); + + let migration = state.decrypt_proxy_node_password_with_before_compare("manual-race", { + let reached_compare = Arc::clone(&reached_compare); + let resume_compare = Arc::clone(&resume_compare); + let first_compare = Arc::clone(&first_compare); + move || { + let reached_compare = Arc::clone(&reached_compare); + let resume_compare = Arc::clone(&resume_compare); + let first_compare = Arc::clone(&first_compare); + async move { + if first_compare.swap(false, Ordering::SeqCst) { + reached_compare.wait().await; + resume_compare.wait().await; + } + } + } + }); + let rotation = async { + reached_compare.wait().await; + state + .update_manual_proxy_node(&ProxyNodeManualUpdateMutation { + node_id: "manual-race".to_string(), + name: None, + ip: None, + port: None, + region: None, + proxy_url: None, + proxy_username: None, + proxy_password: Some("rotated-password".to_string()), + }) + .await + .expect("administrator rotation should persist") + .expect("manual proxy should exist"); + resume_compare.wait().await; + }; + + let (migrated, ()) = tokio::join!(migration, rotation); + assert_eq!( + migrated + .expect("migration should retry the rotated value") + .as_deref(), + Some("rotated-password") + ); + let stored = repository + .find_proxy_node("manual-race") + .await + .expect("stored proxy should read") + .expect("stored proxy should exist") + .proxy_password + .expect("stored password should exist"); + assert!(proxy_node_secret_is_v2(&stored)); + assert!(!stored.contains("legacy-password")); + } + + #[tokio::test] + async fn damaged_proxy_password_ciphertext_fails_closed() { + let mut node = sample_manual_node("manual-damaged"); + node.proxy_password = Some("aether-runtime-secret-v1:not-a-fernet-token".to_string()); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + + assert!(state + .decrypt_proxy_node_password("manual-damaged") + .await + .is_err()); + assert!(state + .resolve_proxy_node_snapshot(Some("manual-damaged")) + .await + .is_none()); + } + + #[tokio::test] + async fn proxy_node_password_v2_rejects_cross_node_and_cross_field_copy() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let first = state + .create_manual_proxy_node(&sample_manual_create_mutation( + "manual-bound-first", + "127.0.0.11", + 8111, + "first-password", + )) + .await + .expect("first proxy should create") + .expect("first proxy should exist"); + let second = state + .create_manual_proxy_node(&sample_manual_create_mutation( + "manual-bound-second", + "127.0.0.12", + 8112, + "second-password", + )) + .await + .expect("second proxy should create") + .expect("second proxy should exist"); + let first_ciphertext = repository + .find_proxy_node(&first.id) + .await + .expect("first proxy should read") + .and_then(|node| node.proxy_password) + .expect("first ciphertext should exist"); + let second_ciphertext = repository + .find_proxy_node(&second.id) + .await + .expect("second proxy should read") + .and_then(|node| node.proxy_password) + .expect("second ciphertext should exist"); + + assert!(proxy_node_secret_is_v2(&first_ciphertext)); + assert!(open_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + &first.id, + &first_ciphertext, + ) + .is_none()); + assert!(repository + .compare_and_set_proxy_password(&second.id, &second_ciphertext, &first_ciphertext) + .await + .expect("ciphertext copy should persist for the adversarial fixture")); + assert_eq!( + state + .decrypt_proxy_node_password(&first.id) + .await + .expect("original ciphertext should open") + .as_deref(), + Some("first-password") + ); + assert!(state.decrypt_proxy_node_password(&second.id).await.is_err()); + } + + #[tokio::test] + async fn proxy_node_password_legacy_runtime_envelope_migrates_to_bound_v2() { + let legacy = seal_runtime_secret_payload_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_NODE_PASSWORD_LEGACY_SECRET_PURPOSE, + "legacy-runtime-password", + ) + .expect("legacy password should seal"); + let mut node = sample_manual_node("manual-runtime-legacy"); + node.proxy_password = Some(legacy); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + + assert_eq!( + state + .decrypt_proxy_node_password("manual-runtime-legacy") + .await + .expect("legacy password should migrate") + .as_deref(), + Some("legacy-runtime-password") + ); + let migrated = repository + .find_proxy_node("manual-runtime-legacy") + .await + .expect("migrated proxy should read") + .and_then(|node| node.proxy_password) + .expect("migrated password should exist"); + assert!(proxy_node_secret_is_v2(&migrated)); + } + + #[tokio::test] + async fn tunnel_psk_registration_and_legacy_migration_store_only_ciphertext() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data.clone()); + let registered = state + .register_proxy_node(&ProxyNodeRegistrationMutation { + node_id: None, + name: "tunnel-new".to_string(), + ip: "127.0.0.1".to_string(), + port: 0, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key": VALID_TUNNEL_PSK, + } + })), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("tunnel registration should succeed") + .expect("tunnel should be stored"); + let stored = repository + .find_proxy_node(®istered.id) + .await + .expect("stored tunnel should read") + .expect("stored tunnel should exist") + .proxy_metadata + .expect("stored tunnel metadata should exist"); + assert!(stored.pointer("/tunnel_security/encryption_key").is_none()); + let ciphertext = stored + .pointer("/tunnel_security/encryption_key_encrypted") + .and_then(serde_json::Value::as_str) + .expect("encrypted tunnel key should exist"); + assert!(!ciphertext.contains(VALID_TUNNEL_PSK)); + assert_eq!( + open_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + ®istered.id, + ciphertext, + ) + .as_deref(), + Some(VALID_TUNNEL_PSK) + ); + + let mut legacy = sample_tunnel_node("tunnel-legacy"); + legacy.proxy_metadata = Some(json!({ + "version": "1.0.0", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key": VALID_TUNNEL_PSK, + } + })); + let legacy_repository = Arc::new(InMemoryProxyNodeRepository::seed([legacy])); + let legacy_data = + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&legacy_repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + assert_eq!( + decrypt_or_migrate_proxy_tunnel_psk(&legacy_data, "tunnel-legacy") + .await + .expect("legacy tunnel key should migrate") + .as_deref(), + Some(VALID_TUNNEL_PSK) + ); + let migrated = legacy_repository + .find_proxy_node("tunnel-legacy") + .await + .expect("migrated tunnel should read") + .expect("migrated tunnel should exist") + .proxy_metadata + .expect("migrated metadata should exist"); + assert!(migrated + .pointer("/tunnel_security/encryption_key") + .is_none()); + assert!(migrated + .pointer("/tunnel_security/encryption_key_encrypted") + .and_then(serde_json::Value::as_str) + .is_some_and(proxy_node_secret_is_v2)); + } + + #[tokio::test] + async fn damaged_tunnel_psk_ciphertext_fails_closed() { + let mut node = sample_tunnel_node("tunnel-damaged"); + node.proxy_metadata = Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-runtime-secret-v1:not-a-fernet-token", + } + })); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + + assert!(decrypt_or_migrate_proxy_tunnel_psk(&data, "tunnel-damaged") + .await + .is_err()); + } + + #[tokio::test] + async fn proxy_tunnel_psk_v2_rejects_cross_node_and_cross_field_copy() { + let first_ciphertext = seal_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + "tunnel-bound-first", + VALID_TUNNEL_PSK, + ) + .expect("first tunnel key should seal"); + let mut first = sample_tunnel_node("tunnel-bound-first"); + first.proxy_metadata = Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": first_ciphertext, + } + })); + let mut second = sample_tunnel_node("tunnel-bound-second"); + second.proxy_metadata = first.proxy_metadata.clone(); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([first, second])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + + assert_eq!( + decrypt_or_migrate_proxy_tunnel_psk(&data, "tunnel-bound-first") + .await + .expect("original tunnel key should open") + .as_deref(), + Some(VALID_TUNNEL_PSK) + ); + assert!( + decrypt_or_migrate_proxy_tunnel_psk(&data, "tunnel-bound-second") + .await + .is_err() + ); + let copied = data + .find_proxy_node("tunnel-bound-first") + .await + .expect("first tunnel should read") + .and_then(|node| node.proxy_metadata) + .and_then(|metadata| { + metadata + .pointer("/tunnel_security/encryption_key_encrypted") + .and_then(serde_json::Value::as_str) + .map(ToOwned::to_owned) + }) + .expect("tunnel ciphertext should exist"); + assert!(open_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_NODE_PASSWORD_SCOPE, + PROXY_NODE_PASSWORD_FIELD, + "tunnel-bound-first", + &copied, + ) + .is_none()); + } + + #[tokio::test] + async fn wrong_binding_v2_tunnel_psk_never_falls_back_to_legacy_plaintext() { + let wrong_binding = seal_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + "different-node", + VALID_TUNNEL_PSK, + ) + .expect("wrong-binding fixture should seal"); + let mut node = sample_tunnel_node("tunnel-no-v2-fallback"); + node.proxy_metadata = Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": wrong_binding, + "encryption_key": VALID_TUNNEL_PSK, + } + })); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + + assert!( + decrypt_or_migrate_proxy_tunnel_psk(&data, "tunnel-no-v2-fallback") + .await + .is_err() + ); + } + + #[tokio::test] + async fn proxy_tunnel_psk_legacy_runtime_envelope_migrates_to_bound_v2() { + let legacy = seal_runtime_secret_payload_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_LEGACY_SECRET_PURPOSE, + VALID_TUNNEL_PSK, + ) + .expect("legacy tunnel key should seal"); + let mut node = sample_tunnel_node("tunnel-runtime-legacy"); + node.proxy_metadata = Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": legacy, + } + })); + let repository = Arc::new(InMemoryProxyNodeRepository::seed([node])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + + assert_eq!( + decrypt_or_migrate_proxy_tunnel_psk(&data, "tunnel-runtime-legacy") + .await + .expect("legacy tunnel key should migrate") + .as_deref(), + Some(VALID_TUNNEL_PSK) + ); + let migrated = repository + .find_proxy_node("tunnel-runtime-legacy") + .await + .expect("migrated tunnel should read") + .and_then(|node| node.proxy_metadata) + .and_then(|metadata| { + metadata + .pointer("/tunnel_security/encryption_key_encrypted") + .and_then(serde_json::Value::as_str) + .map(ToOwned::to_owned) + }) + .expect("migrated tunnel ciphertext should exist"); + assert!(proxy_node_secret_is_v2(&migrated)); + } + + #[tokio::test] + async fn incoming_proxy_node_ciphertext_is_rejected_before_repository_write() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let password_ciphertext = seal_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_NODE_PASSWORD_SCOPE, + PROXY_NODE_PASSWORD_FIELD, + "attacker-selected-id", + "copied-password", + ) + .expect("password fixture should seal"); + assert!(state + .create_manual_proxy_node(&sample_manual_create_mutation( + "manual-ciphertext-input", + "127.0.0.21", + 8121, + &password_ciphertext, + )) + .await + .is_err()); + let legacy_runtime_ciphertext = seal_runtime_secret_payload_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_NODE_PASSWORD_LEGACY_SECRET_PURPOSE, + "legacy-copied-password", + ) + .expect("legacy password fixture should seal"); + let fernet_ciphertext = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "fernet-copied-password") + .expect("fernet password fixture should seal"); + for (name, ip, port, ciphertext) in [ + ( + "manual-v1-ciphertext-input", + "127.0.0.23", + 8123, + legacy_runtime_ciphertext.as_str(), + ), + ( + "manual-fernet-ciphertext-input", + "127.0.0.24", + 8124, + fernet_ciphertext.as_str(), + ), + ] { + assert!(state + .create_manual_proxy_node(&sample_manual_create_mutation( + name, ip, port, ciphertext, + )) + .await + .is_err()); + } + + let tunnel_ciphertext = seal_proxy_node_secret_v2_with_encryption_key( + Some(DEVELOPMENT_ENCRYPTION_KEY), + PROXY_TUNNEL_PSK_SCOPE, + PROXY_TUNNEL_PSK_FIELD, + "attacker-selected-id", + VALID_TUNNEL_PSK, + ) + .expect("tunnel fixture should seal"); + assert!(state + .register_proxy_node(&sample_tunnel_registration_mutation( + "tunnel-ciphertext-input", + "127.0.0.22", + 8122, + &tunnel_ciphertext, + )) + .await + .is_err()); + assert!(repository + .list_proxy_nodes() + .await + .expect("repository should list") + .is_empty()); + } + + #[tokio::test] + async fn tunnel_reregistration_keeps_node_identity_and_rebinds_rotated_psk() { + let repository = Arc::new(InMemoryProxyNodeRepository::default()); + let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data.clone()); + let first = state + .register_proxy_node(&sample_tunnel_registration_mutation( + "tunnel-stable-first", + "127.0.0.31", + 8131, + VALID_TUNNEL_PSK, + )) + .await + .expect("first registration should succeed") + .expect("first registration should persist"); + let rotated = state + .register_proxy_node(&sample_tunnel_registration_mutation( + "tunnel-stable-rotated", + "127.0.0.31", + 8131, + ROTATED_TUNNEL_PSK, + )) + .await + .expect("rotated registration should succeed") + .expect("rotated registration should persist"); + + assert_eq!(rotated.id, first.id); + assert_eq!( + decrypt_or_migrate_proxy_tunnel_psk(&data, &first.id) + .await + .expect("rotated key should open") + .as_deref(), + Some(ROTATED_TUNNEL_PSK) + ); } #[tokio::test] @@ -295,9 +1570,9 @@ mod tests { #[tokio::test] async fn resolve_proxy_node_snapshot_keeps_tunnel_node_with_owner_hint() { - let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_tunnel_node( - "proxy-node-owned", - )])); + let node = sample_tunnel_node("proxy-node-owned"); + let tunnel_generation = node.tunnel_generation.clone(); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node])); let state = AppState::new() .expect("state should build") .with_data_state_for_tests( @@ -307,6 +1582,7 @@ mod tests { json!({ "gateway_instance_id": "gateway-owner", "relay_base_url": "http://gateway-owner.internal", + "tunnel_generation": tunnel_generation, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -331,7 +1607,7 @@ mod tests { } #[tokio::test] - async fn resolve_configured_proxy_snapshot_rejects_unroutable_stored_tunnel_reference() { + async fn resolve_configured_proxy_snapshot_blocks_unroutable_stored_tunnel_reference() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_tunnel_node( "proxy-node-stale", )])); @@ -346,13 +1622,24 @@ mod tests { "node_id": "proxy-node-stale", "enabled": true, }))) - .await; + .await + .expect("explicit unavailable proxy must remain represented"); - assert_eq!(snapshot, None); + assert_eq!(snapshot.enabled, Some(true)); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); + assert_eq!( + snapshot + .extra + .as_ref() + .and_then(|extra| extra.get(PROXY_UNAVAILABLE_REASON_EXTRA_KEY)) + .and_then(serde_json::Value::as_str), + Some("configured_proxy_node_unavailable") + ); } #[tokio::test] - async fn resolve_configured_proxy_snapshot_keeps_inline_url_when_stored_node_is_unroutable() { + async fn explicit_stored_node_does_not_fall_back_to_inline_url_when_unroutable() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_tunnel_node( "proxy-node-stale", )])); @@ -369,10 +1656,156 @@ mod tests { "enabled": true, }))) .await - .expect("inline proxy URL should still resolve"); + .expect("explicit unavailable proxy must remain represented"); - assert_eq!(snapshot.node_id.as_deref(), Some("proxy-node-stale")); - assert_eq!(snapshot.url.as_deref(), Some("http://proxy.example:8080")); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.node_id.is_none()); + assert!(snapshot.url.is_none()); + } + + #[tokio::test] + async fn malformed_explicit_inline_proxy_remains_fail_closed() { + let state = AppState::new().expect("state should build"); + let snapshot = state + .resolve_configured_proxy_snapshot_with_tunnel_affinity(Some(&json!({ + "enabled": true, + "url": " ", + }))) + .await + .expect("malformed explicit proxy must remain represented"); + + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); + } + + #[tokio::test] + async fn inline_proxy_credentials_are_injected_without_copying_them_to_extra() { + let state = AppState::new().expect("state should build"); + let snapshot = state + .resolve_configured_proxy_snapshot_with_tunnel_affinity(Some(&json!({ + "enabled": true, + "url": "http://proxy.example:8080", + "username": " alice ", + "password": " p:ss ", + "region": "test", + }))) + .await + .expect("authenticated inline proxy should resolve"); + + assert_eq!( + snapshot.url.as_deref(), + Some("http://%20alice%20:%20p%3Ass%20@proxy.example:8080/") + ); + assert_eq!(snapshot.extra, Some(json!({"region": "test"}))); + } + + #[tokio::test] + async fn unavailable_key_proxy_blocks_endpoint_and_provider_fallbacks() { + let repository = Arc::new(InMemoryProxyNodeRepository::seed([sample_tunnel_node( + "proxy-node-stale", + )])); + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests( + repository, + )); + let transport = sample_transport( + Some(json!({"enabled": true, "node_id": "proxy-node-stale"})), + Some(json!({"enabled": true, "url": "http://endpoint-proxy:8080"})), + Some(json!({"enabled": true, "url": "http://provider-proxy:8080"})), + ); + + let snapshot = state + .resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport) + .await + .expect("unavailable key proxy must block fallback"); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); + assert_eq!( + state + .resolve_transport_proxy_source_with_tunnel_affinity(&transport) + .await, + Some("key") + ); + + let transport = sample_transport( + Some(json!({"enabled": false, "node_id": "proxy-node-stale"})), + Some(json!({"enabled": true, "url": "http://endpoint-proxy:8080"})), + None, + ); + let snapshot = state + .resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport) + .await + .expect("disabled key proxy should allow endpoint resolution"); + assert_eq!(snapshot.url.as_deref(), Some("http://endpoint-proxy:8080/")); + } + + #[tokio::test] + async fn unavailable_system_proxy_node_does_not_become_direct_transport() { + let state = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_system_config_values_for_tests([( + "system_proxy_node_id".to_string(), + json!("missing-system-proxy"), + )]), + ); + + let snapshot = state + .resolve_system_proxy_snapshot() + .await + .expect("configured unavailable system proxy must remain represented"); + assert_eq!(snapshot.mode.as_deref(), Some("unavailable")); + assert!(snapshot.url.is_none()); + } + + fn sample_manual_create_mutation( + name: &str, + ip: &str, + port: i32, + password: &str, + ) -> ProxyNodeManualCreateMutation { + ProxyNodeManualCreateMutation { + node_id: None, + name: name.to_string(), + ip: ip.to_string(), + port, + region: None, + proxy_url: "http://proxy.example:8080".to_string(), + proxy_username: Some("alice".to_string()), + proxy_password: Some(password.to_string()), + registered_by: None, + } + } + + fn sample_tunnel_registration_mutation( + name: &str, + ip: &str, + port: i32, + psk: &str, + ) -> ProxyNodeRegistrationMutation { + ProxyNodeRegistrationMutation { + node_id: None, + name: name.to_string(), + ip: ip.to_string(), + port, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key": psk, + } + })), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + } } fn sample_tunnel_node(id: &str) -> StoredProxyNode { @@ -395,4 +1828,90 @@ mod tests { ) .expect("sample tunnel node should build") } + + fn sample_manual_node(id: &str) -> StoredProxyNode { + StoredProxyNode::new( + id.to_string(), + id.to_string(), + "127.0.0.1".to_string(), + 8080, + true, + "online".to_string(), + 0, + 0, + 0, + 0, + 0, + 0, + false, + false, + 0, + ) + .expect("sample manual node should build") + .with_manual_proxy_fields( + Some("http://proxy.example:8080".to_string()), + Some("alice".to_string()), + None, + ) + } + + fn sample_transport( + key_proxy: Option, + endpoint_proxy: Option, + provider_proxy: Option, + ) -> GatewayProviderTransportSnapshot { + GatewayProviderTransportSnapshot { + provider: GatewayProviderTransportProvider { + id: "provider-1".to_string(), + name: "provider".to_string(), + provider_type: "openai".to_string(), + website: None, + is_active: true, + keep_priority_on_conversion: false, + enable_format_conversion: false, + concurrent_limit: None, + max_retries: None, + proxy: provider_proxy, + request_timeout_secs: None, + stream_first_byte_timeout_secs: None, + config: None, + }, + endpoint: GatewayProviderTransportEndpoint { + id: "endpoint-1".to_string(), + provider_id: "provider-1".to_string(), + api_format: "openai:chat_completions".to_string(), + api_family: None, + endpoint_kind: None, + is_active: true, + base_url: "https://api.example.test".to_string(), + header_rules: None, + body_rules: None, + max_retries: None, + custom_path: None, + config: None, + format_acceptance_config: None, + proxy: endpoint_proxy, + }, + key: GatewayProviderTransportKey { + id: "key-1".to_string(), + provider_id: "provider-1".to_string(), + name: "key".to_string(), + auth_type: "api_key".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: key_proxy, + fingerprint: None, + upstream_metadata: None, + decrypted_api_key: "test-key".to_string(), + decrypted_auth_config: None, + }, + } + } } diff --git a/apps/aether-gateway/src/state/routing_profiles.rs b/apps/aether-gateway/src/state/routing_profiles.rs index 65fd868aa..8bbb4ff9b 100644 --- a/apps/aether-gateway/src/state/routing_profiles.rs +++ b/apps/aether-gateway/src/state/routing_profiles.rs @@ -4,11 +4,85 @@ use aether_data_contracts::repository::routing_profiles::{ StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion, UpdateRoutingGroupBindingRecord, UpdateRoutingGroupRecord, }; +use aether_routing_core::RoutingGroupConfig; use std::sync::Arc; +use tracing::warn; use super::{AppState, GatewayError}; +const BOOTSTRAP_SYSTEM_DEFAULT_ROUTING_GROUP_NAME: &str = "system-default"; + impl AppState { + /// Make sure an enabled system-default routing group exists. + /// + /// When none exists, one is created from the routing defaults. + /// Returns the created group, or `None` when nothing had to be created (no + /// routing storage, no writer, or a system default already exists). + pub async fn ensure_system_default_routing_group( + &self, + ) -> Result, std::io::Error> { + self.ensure_system_default_routing_group_inner() + .await + .map_err(|err| std::io::Error::other(format!("{err:?}"))) + } + + pub(crate) async fn ensure_system_default_routing_group_inner( + &self, + ) -> Result, GatewayError> { + if !self.has_routing_group_data_reader() { + return Ok(None); + } + if self + .find_routing_group(RoutingGroupLookupKey::SystemDefault) + .await? + .is_some() + { + return Ok(None); + } + if !self.has_routing_group_data_writer() { + warn!( + event_name = "routing_system_default_bootstrap_skipped", + log_type = "event", + "no system default routing group exists and routing storage is read-only; scheduler uses routing defaults" + ); + return Ok(None); + } + + let config = RoutingGroupConfig::default(); + let config_json = serde_json::to_value(config) + .map_err(|err| GatewayError::Internal(format!("serialize routing config: {err}")))?; + + let name = if self + .find_routing_group(RoutingGroupLookupKey::Name( + BOOTSTRAP_SYSTEM_DEFAULT_ROUTING_GROUP_NAME, + )) + .await? + .is_some() + { + format!( + "{BOOTSTRAP_SYSTEM_DEFAULT_ROUTING_GROUP_NAME}-{}", + &uuid::Uuid::new_v4().simple().to_string()[..8] + ) + } else { + BOOTSTRAP_SYSTEM_DEFAULT_ROUTING_GROUP_NAME.to_string() + }; + let now = crate::clock::current_unix_secs() as i64; + self.create_routing_group(CreateRoutingGroupRecord { + id: uuid::Uuid::new_v4().to_string(), + name, + description: Some("系统默认调度策略".to_string()), + enabled: true, + is_system_default: true, + sort_order: 0, + config_json, + version: 1, + created_at: now, + updated_at: now, + published_at: Some(now), + }) + .await + } + pub(crate) fn has_routing_group_data_reader(&self) -> bool { self.data.has_routing_group_reader() } diff --git a/apps/aether-gateway/src/state/runtime/api_key_exports.rs b/apps/aether-gateway/src/state/runtime/api_key_exports.rs index c89f6ea83..5e7dca7d0 100644 --- a/apps/aether-gateway/src/state/runtime/api_key_exports.rs +++ b/apps/aether-gateway/src/state/runtime/api_key_exports.rs @@ -234,6 +234,23 @@ impl AppState { record: aether_data::repository::auth::CreateUserApiKeyRecord, ) -> Result, GatewayError> { + #[cfg(test)] + { + // Unit-test AppState instances keep users and API keys in separate in-memory + // repositories. Bridge them with the authoritative user record only; an unknown, + // inactive, or deleted user must never be synthesized from the key request. + let Some(user) = self.find_user_auth_by_id(&record.user_id).await? else { + return Ok(None); + }; + if !user.is_active || user.is_deleted { + return Ok(None); + } + self.data + .synchronize_user_api_key_owner_for_tests(&user) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + } + let api_key = self .data .create_user_api_key(record) @@ -277,6 +294,32 @@ impl AppState { Ok(api_key) } + pub(crate) async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &aether_data::repository::auth::CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + self.data + .compare_and_swap_api_key_ciphertext(mutation) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn update_user_api_key_basic_if_unlocked( + &self, + record: aether_data::repository::auth::UpdateUserApiKeyBasicRecord, + ) -> Result, GatewayError> + { + let api_key = self + .data + .update_user_api_key_basic_if_unlocked(record) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if api_key.is_some() { + self.invalidate_auth_context_cache(); + } + Ok(api_key) + } + pub(crate) async fn update_standalone_api_key_basic( &self, record: aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord, @@ -293,6 +336,22 @@ impl AppState { Ok(api_key) } + pub(crate) async fn restore_api_key_if_matches( + &self, + expected: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + restored: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, + ) -> Result { + let restored = self + .data + .restore_api_key_if_matches(expected, restored) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if restored { + self.invalidate_auth_context_cache(); + } + Ok(restored) + } + pub(crate) async fn set_user_api_key_active( &self, user_id: &str, @@ -311,6 +370,24 @@ impl AppState { Ok(api_key) } + pub(crate) async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, GatewayError> + { + let api_key = self + .data + .set_user_api_key_active_if_unlocked(user_id, api_key_id, is_active) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if api_key.is_some() { + self.invalidate_auth_context_cache(); + } + Ok(api_key) + } + pub(crate) async fn set_standalone_api_key_active( &self, api_key_id: &str, @@ -363,6 +440,24 @@ impl AppState { Ok(api_key) } + pub(crate) async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, GatewayError> + { + let api_key = self + .data + .set_user_api_key_allowed_providers_if_unlocked(user_id, api_key_id, allowed_providers) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if api_key.is_some() { + self.invalidate_auth_context_cache(); + } + Ok(api_key) + } + pub(crate) async fn set_user_api_key_force_capabilities( &self, user_id: &str, @@ -381,6 +476,28 @@ impl AppState { Ok(api_key) } + pub(crate) async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, GatewayError> + { + let api_key = self + .data + .set_user_api_key_force_capabilities_if_unlocked( + user_id, + api_key_id, + force_capabilities, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if api_key.is_some() { + self.invalidate_auth_context_cache(); + } + Ok(api_key) + } + pub(crate) async fn set_user_api_key_feature_settings( &self, user_id: &str, @@ -399,6 +516,24 @@ impl AppState { Ok(api_key) } + pub(crate) async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, GatewayError> + { + let api_key = self + .data + .set_user_api_key_feature_settings_if_unlocked(user_id, api_key_id, feature_settings) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if api_key.is_some() { + self.invalidate_auth_context_cache(); + } + Ok(api_key) + } + pub(crate) async fn set_api_key_usage_totals( &self, api_key_id: &str, @@ -451,6 +586,22 @@ impl AppState { Ok(deleted) } + pub(crate) async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + let deleted = self + .data + .delete_user_api_key_if_unlocked(user_id, api_key_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if deleted { + self.invalidate_auth_context_cache(); + } + Ok(deleted) + } + pub(crate) async fn delete_standalone_api_key( &self, api_key_id: &str, @@ -466,3 +617,190 @@ impl AppState { Ok(deleted) } } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_data::repository::auth::{ + AuthApiKeyLookupKey, AuthApiKeyReadRepository, CreateUserApiKeyRecord, + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, + }; + use aether_data::repository::users::StoredUserAuthRecord; + + use crate::data::GatewayDataState; + use crate::AppState; + + fn authoritative_user( + user_id: &str, + is_active: bool, + is_deleted: bool, + ) -> StoredUserAuthRecord { + StoredUserAuthRecord::new( + user_id.to_string(), + Some(format!("{user_id}@example.com")), + true, + format!("owner-{user_id}"), + Some("server-managed-password-hash".to_string()), + "admin".to_string(), + "oauth".to_string(), + Some(serde_json::json!(["openai"])), + Some(serde_json::json!(["openai:chat"])), + Some(serde_json::json!(["gpt-5"])), + is_active, + is_deleted, + None, + None, + ) + .expect("authoritative user should build") + .with_security_version(41) + .expect("security version should be valid") + } + + fn create_record(user_id: &str, api_key_id: &str) -> CreateUserApiKeyRecord { + CreateUserApiKeyRecord { + user_id: user_id.to_string(), + api_key_id: api_key_id.to_string(), + key_hash: format!("hash-{api_key_id}"), + key_encrypted: Some(format!("encrypted-{api_key_id}")), + name: Some("first key".to_string()), + allowed_providers: Some(vec!["anthropic".to_string()]), + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + rate_limit: 0, + concurrent_limit: None, + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + } + } + + fn state_with_users( + repository: Arc, + users: I, + ) -> AppState + where + I: IntoIterator, + { + AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests( + repository, + )) + .with_auth_users_for_tests(users) + } + + #[tokio::test] + async fn test_gateway_first_user_key_requires_authoritative_active_owner() { + let unknown_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let unknown_state = state_with_users( + Arc::clone(&unknown_repository), + Vec::::new(), + ); + assert!(unknown_state + .create_user_api_key(create_record("missing-user", "missing-key")) + .await + .expect("unknown owner creation should resolve") + .is_none()); + assert!(unknown_repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("missing-key")) + .await + .expect("unknown key lookup should resolve") + .is_none()); + + for (user_id, is_active, is_deleted) in [ + ("inactive-user", false, false), + ("deleted-user", true, true), + ] { + let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let state = state_with_users( + Arc::clone(&repository), + [authoritative_user(user_id, is_active, is_deleted)], + ); + let api_key_id = format!("{user_id}-key"); + assert!(state + .create_user_api_key(create_record(user_id, &api_key_id)) + .await + .expect("ineligible owner creation should resolve") + .is_none()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId(&api_key_id)) + .await + .expect("ineligible key lookup should resolve") + .is_none()); + } + } + + #[tokio::test] + async fn test_gateway_first_user_key_syncs_owner_without_mutating_authority() { + let stale_owner = StoredAuthApiKeySnapshot::new( + "active-user".to_string(), + "request-derived-owner".to_string(), + None, + "user".to_string(), + "local".to_string(), + false, + false, + None, + None, + None, + "ignored-owner-fixture".to_string(), + None, + true, + false, + false, + None, + None, + None, + None, + None, + None, + ) + .expect("stale owner fixture should build"); + let repository = Arc::new( + InMemoryAuthApiKeySnapshotRepository::default().with_owner_snapshots([stale_owner]), + ); + let authoritative = authoritative_user("active-user", true, false); + let state = state_with_users(Arc::clone(&repository), [authoritative.clone()]); + + state + .create_user_api_key(create_record("active-user", "active-key")) + .await + .expect("active owner creation should resolve") + .expect("active authoritative owner should allow its first key"); + + let snapshot = repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("active-key")) + .await + .expect("created key lookup should resolve") + .expect("created key should exist"); + assert_eq!(snapshot.user_role, "admin"); + assert!(snapshot.user_is_active); + assert!(!snapshot.user_is_deleted); + assert_eq!( + snapshot.user_allowed_providers, + Some(vec!["openai".to_string()]) + ); + assert_eq!( + snapshot.api_key_allowed_providers, + Some(vec!["anthropic".to_string()]) + ); + + let unchanged = state + .find_user_auth_by_id("active-user") + .await + .expect("authoritative owner lookup should resolve") + .expect("authoritative owner should remain"); + assert_eq!(unchanged, authoritative); + assert_eq!(unchanged.role, authoritative.role); + assert_eq!(unchanged.is_active, authoritative.is_active); + assert_eq!(unchanged.is_deleted, authoritative.is_deleted); + assert_eq!(unchanged.security_version, 41); + } +} diff --git a/apps/aether-gateway/src/state/runtime/auth/sessions.rs b/apps/aether-gateway/src/state/runtime/auth/sessions.rs index 25ef1e077..a4d47e9be 100644 --- a/apps/aether-gateway/src/state/runtime/auth/sessions.rs +++ b/apps/aether-gateway/src/state/runtime/auth/sessions.rs @@ -122,13 +122,44 @@ impl AppState { { let session = session.into(); #[cfg(test)] - if let Some(store) = self.auth_session_store.as_ref() { + if let (Some(user_store), Some(session_store)) = ( + self.auth_user_store.as_ref(), + self.auth_session_store.as_ref(), + ) { + let existing = { + user_store + .lock() + .expect("auth user store should lock") + .get(&session.user_id) + .cloned() + }; + let existing = match existing { + Some(user) => Some(user), + None => self + .data + .find_user_auth_by_id(&session.user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?, + }; + let Some(existing) = existing else { + return Ok(None); + }; + let mut users = user_store.lock().expect("auth user store should lock"); + let user = users.entry(session.user_id.clone()).or_insert(existing); + if !user.is_active + || user.is_deleted + || user.security_version != session.security_version + { + return Ok(None); + } let now = session .created_at .or(session.updated_at) .or(session.last_seen_at) .unwrap_or_else(chrono::Utc::now); - let mut guard = store.lock().expect("auth session store should lock"); + let mut guard = session_store + .lock() + .expect("auth session store should lock"); for existing in guard.values_mut() { if existing.user_id == session.user_id && existing.client_device_id == session.client_device_id @@ -155,12 +186,96 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn create_user_session_if_password_matches( + &self, + session: T, + expected_password_hash: &str, + ) -> Result, GatewayError> + where + T: Into, + { + let session = session.into(); + #[cfg(test)] + if self.auth_session_store.is_some() && self.auth_user_store.is_some() { + let existing = { + self.auth_user_store + .as_ref() + .expect("checked auth user store") + .lock() + .expect("auth user store should lock") + .get(&session.user_id) + .cloned() + }; + let existing = match existing { + Some(user) => Some(user), + None => self + .data + .find_user_auth_by_id(&session.user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?, + }; + let Some(existing) = existing else { + return Ok(None); + }; + let mut users = self + .auth_user_store + .as_ref() + .expect("checked auth user store") + .lock() + .expect("auth user store should lock"); + let user = users.entry(session.user_id.clone()).or_insert(existing); + if user.password_hash.as_deref() != Some(expected_password_hash) + || !user.auth_source.eq_ignore_ascii_case("local") + || !user.is_active + || user.is_deleted + || user.security_version != session.security_version + { + return Ok(None); + } + let now = session + .created_at + .or(session.updated_at) + .or(session.last_seen_at) + .unwrap_or_else(chrono::Utc::now); + user.last_login_at = Some(now); + let mut sessions = self + .auth_session_store + .as_ref() + .expect("checked auth session store") + .lock() + .expect("auth session store should lock"); + for existing in sessions.values_mut() { + if existing.user_id == session.user_id + && existing.client_device_id == session.client_device_id + && !existing.is_revoked() + && !existing.is_expired(now) + { + existing.revoked_at = Some(now); + existing.revoke_reason = Some("replaced_by_new_login".to_string()); + existing.updated_at = Some(now); + } + } + sessions.insert( + format!("{}:{}", session.user_id, session.id), + session.clone().into(), + ); + return Ok(Some(session)); + } + + let raw_session: crate::data::state::StoredUserSessionRecord = session.into(); + self.data + .create_user_session_if_password_matches(&raw_session, expected_password_hash) + .await + .map(|value| value.map(Into::into)) + .map_err(|err| GatewayError::Internal(err.to_string())) + } + #[allow(clippy::too_many_arguments)] pub(crate) async fn rotate_user_session_refresh_token( &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: chrono::DateTime, expires_at: chrono::DateTime, @@ -171,8 +286,12 @@ impl AppState { if let Some(store) = self.auth_session_store.as_ref() { let key = format!("{user_id}:{session_id}"); let mut guard = store.lock().expect("auth session store should lock"); - if let Some(session) = guard.get_mut(&key) { - session.prev_refresh_token_hash = Some(previous_refresh_token_hash.to_string()); + if let Some(session) = guard.get_mut(&key).filter(|session| { + session.refresh_token_hash == expected_refresh_token_hash + && !session.is_revoked() + && !session.is_expired(rotated_at) + }) { + session.prev_refresh_token_hash = Some(expected_refresh_token_hash.to_string()); session.refresh_token_hash = next_refresh_token_hash.to_string(); session.rotated_at = Some(rotated_at); session.expires_at = Some(expires_at); @@ -193,7 +312,7 @@ impl AppState { .rotate_user_session_refresh_token( user_id, session_id, - previous_refresh_token_hash, + expected_refresh_token_hash, next_refresh_token_hash, rotated_at, expires_at, diff --git a/apps/aether-gateway/src/state/runtime/auth/user_lifecycle.rs b/apps/aether-gateway/src/state/runtime/auth/user_lifecycle.rs index bc5cca72d..f2c397489 100644 --- a/apps/aether-gateway/src/state/runtime/auth/user_lifecycle.rs +++ b/apps/aether-gateway/src/state/runtime/auth/user_lifecycle.rs @@ -288,6 +288,22 @@ impl AppState { Ok(group) } + pub(crate) async fn restore_user_group_if_matches( + &self, + expected: &aether_data::repository::users::StoredUserGroup, + restored: &aether_data::repository::users::StoredUserGroup, + ) -> Result { + let restored = self + .data + .restore_user_group_if_matches(expected, restored) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if restored { + self.invalidate_auth_context_cache(); + } + Ok(restored) + } + pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result { let deleted = self .data @@ -369,6 +385,23 @@ impl AppState { Ok(groups) } + pub(crate) async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + let restored = self + .data + .restore_user_groups_if_matches(user_id, expected_group_ids, restored_group_ids) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if restored { + self.invalidate_auth_context_cache(); + } + Ok(restored) + } + pub(crate) async fn add_user_to_group( &self, group_id: &str, @@ -434,7 +467,9 @@ impl AppState { pub(crate) async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, GatewayError> { #[cfg(test)] @@ -457,8 +492,11 @@ impl AppState { let Some(mut user) = existing else { return Ok(None); }; - if let Some(email) = email { - user.email = Some(email); + if email_present { + user.email = email; + } + if let Some(email_verified) = email_verified { + user.email_verified = email_verified; } if let Some(username) = username { user.username = username; @@ -473,7 +511,7 @@ impl AppState { let user = self .data - .update_local_auth_user_profile(user_id, email, username) + .update_local_auth_user_profile(user_id, email_present, email, email_verified, username) .await .map_err(|err| GatewayError::Internal(err.to_string()))?; if user.is_some() { @@ -482,6 +520,159 @@ impl AppState { Ok(user) } + #[allow(clippy::too_many_arguments)] + pub(crate) async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &aether_data::repository::users::StoredUserAuthRecord, + restored_auth: &aether_data::repository::users::StoredUserAuthRecord, + expected_export: &aether_data::repository::users::StoredUserExportRow, + restored_export: &aether_data::repository::users::StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + #[cfg(test)] + if let Some(store) = self.auth_user_store.as_ref() { + // The gateway test harness uses an auth overlay for users while most auxiliary data + // remains in the repository. Keep the compare-and-write atomic for that overlay too; + // otherwise import rollback tests would silently exercise a different, unconditional + // path than production. + let current_feature = self + .data + .read_user_feature_settings(&expected_auth.id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut users = store.lock().expect("auth user store should lock"); + if users.contains_key(&expected_auth.id) { + if expected_auth.id != restored_auth.id + || expected_export.id != expected_auth.id + || restored_export.id != restored_auth.id + { + return Ok(false); + } + let Some(current) = users.get(&expected_auth.id).cloned() else { + return Ok(false); + }; + if !current.matches_restore_state(expected_auth) { + return Ok(false); + } + + let current_model = + self.auth_user_model_capability_store + .as_ref() + .and_then(|settings| { + settings + .lock() + .expect("auth user model capability store should lock") + .get(&expected_auth.id) + .cloned() + }); + if current_model.as_ref() != expected_model_capability_settings { + return Ok(false); + } + + // Feature settings have no separate test overlay. When the backing repository has + // a row, still honor the snapshot comparison; an absent row is represented by + // `None`, which is the normal overlay case. + if current_feature.as_ref() != expected_feature_settings { + return Ok(false); + } + + let security_state_changed = current.role != restored_auth.role + || current.is_active != restored_auth.is_active; + let removes_active_admin = current.role.eq_ignore_ascii_case("admin") + && current.is_active + && !current.is_deleted + && (!restored_auth.role.eq_ignore_ascii_case("admin") + || !restored_auth.is_active); + if removes_active_admin + && users + .values() + .filter(|user| { + user.role.eq_ignore_ascii_case("admin") + && user.is_active + && !user.is_deleted + }) + .count() + <= 1 + { + return Err(GatewayError::LastActiveAdminUpdateDenied); + } + + let mut updated = restored_auth.clone(); + // Server-managed credentials and timestamps are deliberately not part of this + // aggregate restore. The password has its own nullable CAS operation. + updated.password_hash = current.password_hash; + updated.security_version = current.security_version; + updated.created_at = current.created_at; + updated.last_login_at = current.last_login_at; + updated.auth_source = current.auth_source; + updated.is_deleted = current.is_deleted; + if security_state_changed { + updated.security_version = + updated.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; + } + users.insert(updated.id.clone(), updated.clone()); + drop(users); + + if let Some(settings) = self.auth_user_model_capability_store.as_ref() { + let mut guard = settings + .lock() + .expect("auth user model capability store should lock"); + match restored_model_capability_settings { + Some(value) => { + guard.insert(updated.id.clone(), value); + } + None => { + guard.remove(&updated.id); + } + } + } + if security_state_changed { + if let Some(sessions) = self.auth_session_store.as_ref() { + let now = chrono::Utc::now(); + for session in sessions + .lock() + .expect("auth session store should lock") + .values_mut() + .filter(|session| { + session.user_id == updated.id && session.revoked_at.is_none() + }) + { + session.revoked_at = Some(now); + session.revoke_reason = Some("user_security_state_changed".to_string()); + session.updated_at = Some(now); + } + } + } + self.invalidate_auth_context_cache(); + return Ok(true); + } + } + + let restored = self + .data + .restore_local_auth_user_state_if_matches( + expected_auth, + restored_auth, + expected_export, + restored_export, + expected_model_capability_settings, + restored_model_capability_settings, + expected_feature_settings, + restored_feature_settings, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if restored { + self.invalidate_auth_context_cache(); + } + Ok(restored) + } + pub(crate) async fn update_local_auth_user_password_hash( &self, user_id: &str, @@ -509,6 +700,9 @@ impl AppState { return Ok(None); }; user.password_hash = Some(password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; store .lock() .expect("auth user store should lock") @@ -522,6 +716,193 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: chrono::DateTime, + ) -> Result { + #[cfg(test)] + if let Some(store) = self.auth_user_store.as_ref() { + let existing = { + store + .lock() + .expect("auth user store should lock") + .get(user_id) + .cloned() + }; + let existing = match existing { + Some(user) => Some(user), + None => self + .data + .find_user_auth_by_id(user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?, + }; + let Some(mut user) = existing else { + return Ok(false); + }; + if user.password_hash.as_deref() != expected_password_hash { + return Ok(false); + } + user.password_hash = password_hash; + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; + store + .lock() + .expect("auth user store should lock") + .insert(user.id.clone(), user); + self.invalidate_auth_context_cache(); + return Ok(true); + } + + self.data + .restore_local_auth_user_password_hash_if_matches( + user_id, + expected_password_hash, + password_hash, + updated_at, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + #[cfg(test)] + if let (Some(user_store), Some(session_store)) = ( + self.auth_user_store.as_ref(), + self.auth_session_store.as_ref(), + ) { + let existing = { + user_store + .lock() + .expect("auth user store should lock") + .get(user_id) + .cloned() + }; + let existing = match existing { + Some(user) => Some(user), + None => self + .data + .find_user_auth_by_id(user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?, + }; + let Some(existing) = existing.filter(|user| !user.is_deleted) else { + return Ok(false); + }; + let mut users = user_store.lock().expect("auth user store should lock"); + let mut sessions = session_store + .lock() + .expect("auth session store should lock"); + let user = users.entry(user_id.to_string()).or_insert(existing); + user.password_hash = Some(password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && !session.is_revoked()) + { + session.revoked_at = Some(changed_at); + session.revoke_reason = Some("admin_password_reset".to_string()); + session.updated_at = Some(changed_at); + } + self.invalidate_auth_context_cache(); + return Ok(true); + } + + let reset = self + .data + .reset_local_auth_user_password_and_revoke_sessions(user_id, password_hash, changed_at) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if reset { + self.invalidate_auth_context_cache(); + } + Ok(reset) + } + + pub(crate) async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + #[cfg(test)] + if let (Some(user_store), Some(session_store)) = ( + self.auth_user_store.as_ref(), + self.auth_session_store.as_ref(), + ) { + let existing = { + user_store + .lock() + .expect("auth user store should lock") + .get(user_id) + .cloned() + }; + let existing = match existing { + Some(user) => Some(user), + None => self + .data + .find_user_auth_by_id(user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?, + }; + let Some(existing) = existing else { + return Ok(false); + }; + let mut users = user_store.lock().expect("auth user store should lock"); + let mut sessions = session_store + .lock() + .expect("auth session store should lock"); + let user = users.entry(user_id.to_string()).or_insert(existing); + if user.password_hash.as_deref() != expected_password_hash { + return Ok(false); + } + let current_key = format!("{user_id}:{current_session_id}"); + if !sessions + .get(¤t_key) + .is_some_and(|session| !session.is_revoked() && !session.is_expired(changed_at)) + { + return Ok(false); + } + user.password_hash = Some(next_password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && !session.is_revoked()) + { + session.revoked_at = Some(changed_at); + session.revoke_reason = Some("password_changed".to_string()); + session.updated_at = Some(changed_at); + } + return Ok(true); + } + + self.data + .change_local_auth_password_and_revoke_sessions( + user_id, + current_session_id, + expected_password_hash, + next_password_hash, + changed_at, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn create_local_auth_user( &self, email: Option, @@ -642,10 +1023,18 @@ impl AppState { ) -> Result, GatewayError> { #[cfg(test)] if let Some(store) = self.auth_user_store.as_ref() { - let mut guard = store.lock().expect("auth user store should lock"); - let Some(user) = guard.get_mut(user_id) else { + let mut users = store.lock().expect("auth user store should lock"); + let Some(user) = users.get_mut(user_id) else { return Ok(None); }; + let security_state_changed = role + .as_deref() + .is_some_and(|next_role| !user.role.eq_ignore_ascii_case(next_role)) + || is_active.is_some_and(|next_active| user.is_active != next_active); + let mut sessions = self + .auth_session_store + .as_ref() + .map(|sessions| sessions.lock().expect("auth session store should lock")); if let Some(role) = role { user.role = role; } @@ -661,9 +1050,26 @@ impl AppState { if let Some(is_active) = is_active { user.is_active = is_active; } + if security_state_changed { + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + GatewayError::Internal("users.security_version overflow".to_string()) + })?; + let revoked_at = chrono::Utc::now(); + if let Some(sessions) = sessions.as_mut() { + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && !session.is_revoked()) + { + session.revoked_at = Some(revoked_at); + session.revoke_reason = Some("user_security_state_changed".to_string()); + session.updated_at = Some(revoked_at); + } + } + } let _ = (rate_limit_present, rate_limit); let user = user.clone(); - drop(guard); + drop(sessions); + drop(users); self.invalidate_auth_context_cache(); return Ok(Some(user)); } @@ -684,7 +1090,13 @@ impl AppState { is_active, ) .await - .map_err(|err| GatewayError::Internal(err.to_string()))?; + .map_err(|err| { + if aether_data::repository::users::is_last_active_admin_update_denied(&err) { + GatewayError::LastActiveAdminUpdateDenied + } else { + GatewayError::Internal(err.to_string()) + } + })?; if user.is_some() { self.invalidate_auth_context_cache(); } @@ -788,10 +1200,16 @@ impl AppState { self.data .delete_local_auth_user(user_id) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| { + if aether_data::repository::users::is_last_active_admin_delete_denied(&err) { + GatewayError::LastActiveAdminDeleteDenied + } else { + GatewayError::Internal(err.to_string()) + } + }) } - pub(crate) async fn register_local_auth_user( + pub(crate) async fn register_local_auth_user_with_wallet_outcome( &self, email: Option, email_verified: bool, @@ -803,6 +1221,7 @@ impl AppState { Option<( aether_data::repository::users::StoredUserAuthRecord, aether_data::repository::wallet::StoredWalletSnapshot, + bool, )>, GatewayError, > { @@ -868,11 +1287,17 @@ impl AppState { .lock() .expect("auth wallet store should lock") .insert(wallet.id.clone(), wallet.clone()); - return Ok(Some((user, wallet))); + super::user_provisioning::record_test_initial_gift_transaction( + self, + &wallet, + &user.id, + "用户初始赠款", + ); + return Ok(Some((user, wallet, true))); } self.data - .register_local_auth_user( + .register_local_auth_user_with_wallet_outcome( email, email_verified, username, @@ -883,6 +1308,34 @@ impl AppState { .await .map_err(|err| GatewayError::Internal(err.to_string())) } + + pub(crate) async fn register_local_auth_user( + &self, + email: Option, + email_verified: bool, + username: String, + password_hash: String, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result< + Option<( + aether_data::repository::users::StoredUserAuthRecord, + aether_data::repository::wallet::StoredWalletSnapshot, + )>, + GatewayError, + > { + Ok(self + .register_local_auth_user_with_wallet_outcome( + email, + email_verified, + username, + password_hash, + initial_gift_usd, + unlimited, + ) + .await? + .map(|(user, wallet, _created)| (user, wallet))) + } } fn normalized_user_group_ids(group_ids: &[String]) -> BTreeSet { @@ -944,6 +1397,7 @@ mod tests { local_rejection: None, allowed_models: Some(vec!["gpt-4.1".to_string()]), ip_rules: None, + verified_api_key_hash: None, } } diff --git a/apps/aether-gateway/src/state/runtime/auth/user_provisioning.rs b/apps/aether-gateway/src/state/runtime/auth/user_provisioning.rs index 60f9e1730..cba1f5e7b 100644 --- a/apps/aether-gateway/src/state/runtime/auth/user_provisioning.rs +++ b/apps/aether-gateway/src/state/runtime/auth/user_provisioning.rs @@ -4,7 +4,540 @@ use crate::{AppState, GatewayError}; const USER_RUNTIME_JSON_CACHE_TTL: Duration = Duration::from_secs(30); +fn normalize_ldap_identity_email(value: &str) -> Option { + let normalized = value.trim().to_ascii_lowercase(); + (!normalized.is_empty()).then_some(normalized) +} + +/// LDAP synchronization may create a wallet as a side effect. Keep the +/// wallet id only when this invocation created that row so later compensation +/// can never delete a pre-existing wallet. +pub(crate) struct LdapAuthProvisioningResult { + pub(crate) user: aether_data::repository::users::StoredUserAuthRecord, + pub(crate) owned_wallet_id: Option, +} + +#[cfg(test)] +pub(super) fn record_test_initial_gift_transaction( + state: &AppState, + wallet: &aether_data::repository::wallet::StoredWalletSnapshot, + owner_id: &str, + description: &str, +) { + if wallet.gift_balance <= 0.0 { + return; + } + let Some(store) = state.admin_wallet_transaction_store.as_ref() else { + return; + }; + let created_at_unix_ms = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() + .min(u64::MAX as u128) as u64; + let transaction = crate::AdminWalletTransactionRecord { + id: uuid::Uuid::new_v4().to_string(), + wallet_id: wallet.id.clone(), + category: "gift".to_string(), + reason_code: "gift_initial".to_string(), + amount: wallet.gift_balance, + balance_before: 0.0, + balance_after: wallet.gift_balance, + recharge_balance_before: 0.0, + recharge_balance_after: 0.0, + gift_balance_before: 0.0, + gift_balance_after: wallet.gift_balance, + link_type: Some("system_task".to_string()), + link_id: Some(owner_id.to_string()), + operator_id: None, + description: Some(description.to_string()), + created_at_unix_ms, + }; + store + .lock() + .expect("admin wallet transaction store should lock") + .insert(transaction.id.clone(), transaction); +} + +#[cfg(test)] impl AppState { + /// The test-only auth stores keep users and API keys in separate + /// repositories, just like the production data layer. Check both wallet + /// owner forms before a compensating user delete so an API-key wallet can + /// never be orphaned by the in-memory path. + async fn test_api_key_wallet_exists_for_user( + &self, + user_id: &str, + ) -> Result { + let Some(store) = self.auth_wallet_store.as_ref() else { + return Ok(false); + }; + + if self.data.has_auth_api_key_reader() { + let user_ids = vec![user_id.to_string()]; + let api_key_ids = self + .data + .list_auth_api_key_export_records_by_user_ids(&user_ids) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + .into_iter() + .map(|record| record.api_key_id) + .collect::>(); + if api_key_ids.is_empty() { + return Ok(false); + } + let wallets = store.lock().expect("auth wallet store should lock"); + return Ok(wallets.values().any(|wallet| { + wallet + .api_key_id + .as_deref() + .is_some_and(|api_key_id| api_key_ids.contains(api_key_id)) + })); + } + + // No key repository means ownership cannot be resolved. Treat any + // API-key wallet as a reference and fail closed. + let wallets = store.lock().expect("auth wallet store should lock"); + Ok(wallets.values().any(|wallet| wallet.api_key_id.is_some())) + } +} + +impl AppState { + pub(crate) async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + if wallet_id.trim().is_empty() { + return Ok(false); + } + + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + let owner_matches = + |wallet: &aether_data::repository::wallet::StoredWalletSnapshot| match owner { + aether_data::repository::wallet::WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.id == wallet_id + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.id == wallet_id + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::WalletId(_) => false, + }; + if matches!( + owner, + aether_data::repository::wallet::WalletLookupKey::WalletId(_) + ) { + return Err(GatewayError::Internal( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )); + } + let wallet = store + .lock() + .expect("auth wallet store should lock") + .values() + .find(|wallet| owner_matches(wallet)) + .cloned(); + let Some(wallet) = wallet else { + return Ok(false); + }; + let untouched = wallet.balance == 0.0 + && wallet.gift_balance == 0.0 + && wallet.total_recharged == 0.0 + && wallet.total_consumed == 0.0 + && wallet.total_refunded == 0.0 + && wallet.total_adjusted == 0.0 + && matches!(wallet.limit_mode.as_str(), "finite" | "unlimited") + && wallet.currency == "USD" + && wallet.status == "active"; + if !untouched { + return Ok(false); + } + let has_order = self + .admin_wallet_payment_order_store + .as_ref() + .is_some_and(|orders| { + orders + .lock() + .expect("admin wallet payment order store should lock") + .values() + .any(|order| order.wallet_id == wallet.id) + }); + let has_transaction = + self.admin_wallet_transaction_store + .as_ref() + .is_some_and(|transactions| { + transactions + .lock() + .expect("admin wallet transaction store should lock") + .values() + .any(|transaction| transaction.wallet_id == wallet.id) + }); + let has_refund = self + .admin_wallet_refund_store + .as_ref() + .is_some_and(|refunds| { + refunds + .lock() + .expect("admin wallet refund store should lock") + .values() + .any(|refund| refund.wallet_id == wallet.id) + }); + if has_order || has_transaction || has_refund { + return Ok(false); + } + let mut wallets = store.lock().expect("auth wallet store should lock"); + if !wallets + .get(&wallet.id) + .is_some_and(|current| owner_matches(current)) + { + return Ok(false); + } + let removed = wallets.remove(&wallet.id).is_some(); + if removed { + self.invalidate_auth_context_cache(); + } + return Ok(removed); + } + + self.data + .delete_wallet_if_unreferenced(wallet_id, owner) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &aether_data::repository::wallet::StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + if expected.id.trim().is_empty() { + return Ok(false); + } + + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + let owner_matches = + |wallet: &aether_data::repository::wallet::StoredWalletSnapshot| match owner { + aether_data::repository::wallet::WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.id == expected.id + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.id == expected.id + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::WalletId(_) => false, + }; + if matches!( + owner, + aether_data::repository::wallet::WalletLookupKey::WalletId(_) + ) { + return Err(GatewayError::Internal( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )); + } + let current = store + .lock() + .expect("auth wallet store should lock") + .values() + .find(|wallet| owner_matches(wallet)) + .cloned(); + if current.as_ref() != Some(expected) { + return Ok(false); + } + let has_order = self + .admin_wallet_payment_order_store + .as_ref() + .is_some_and(|orders| { + orders + .lock() + .expect("admin wallet payment order store should lock") + .values() + .any(|order| order.wallet_id == expected.id) + }); + let has_transaction = + self.admin_wallet_transaction_store + .as_ref() + .is_some_and(|transactions| { + transactions + .lock() + .expect("admin wallet transaction store should lock") + .values() + .any(|transaction| transaction.wallet_id == expected.id) + }); + let has_refund = self + .admin_wallet_refund_store + .as_ref() + .is_some_and(|refunds| { + refunds + .lock() + .expect("admin wallet refund store should lock") + .values() + .any(|refund| refund.wallet_id == expected.id) + }); + if has_order || has_transaction || has_refund { + return Ok(false); + } + let mut wallets = store.lock().expect("auth wallet store should lock"); + if wallets + .get(&expected.id) + .is_some_and(|current| owner_matches(current) && current == expected) + { + let removed = wallets.remove(&expected.id).is_some(); + if removed { + self.invalidate_auth_context_cache(); + } + return Ok(removed); + } + return Ok(false); + } + + self.data + .delete_wallet_if_snapshot_matches_and_unreferenced(expected, owner) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn restore_wallet_if_snapshot_matches( + &self, + before: &aether_data::repository::wallet::StoredWalletSnapshot, + after: &aether_data::repository::wallet::StoredWalletSnapshot, + owner: aether_data::repository::wallet::WalletLookupKey<'_>, + ) -> Result { + if before.id.trim().is_empty() || after.id.trim().is_empty() { + return Ok(false); + } + if before.id != after.id { + return Err(GatewayError::Internal( + "wallet restore snapshots must reference the same wallet".to_string(), + )); + } + + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + let owner_matches = + |wallet: &aether_data::repository::wallet::StoredWalletSnapshot| match owner { + aether_data::repository::wallet::WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + aether_data::repository::wallet::WalletLookupKey::WalletId(_) => false, + }; + if matches!( + owner, + aether_data::repository::wallet::WalletLookupKey::WalletId(_) + ) { + return Err(GatewayError::Internal( + "wallet restore requires an explicit user or API-key owner".to_string(), + )); + } + if !owner_matches(before) || !owner_matches(after) { + return Ok(false); + } + let mut guard = store.lock().expect("auth wallet store should lock"); + let Some(current) = guard.get(&after.id) else { + return Ok(false); + }; + if current != after { + return Ok(false); + } + guard.insert(before.id.clone(), before.clone()); + drop(guard); + self.invalidate_auth_context_cache(); + return Ok(true); + } + + let restored = self + .data + .restore_wallet_if_snapshot_matches(before, after, owner) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + if restored { + self.invalidate_auth_context_cache(); + } + Ok(restored) + } + + pub(crate) async fn rollback_provisional_auth_user( + &self, + user_id: &str, + ) -> Result<(), GatewayError> { + self.rollback_provisional_auth_user_with_wallet(user_id, None) + .await + } + + pub(crate) async fn rollback_provisional_auth_user_with_wallet( + &self, + user_id: &str, + wallet_id: Option<&str>, + ) -> Result<(), GatewayError> { + if wallet_id.is_some_and(|wallet_id| wallet_id.trim().is_empty()) { + return Err(GatewayError::Internal( + "wallet compensation wallet id cannot be empty".to_string(), + )); + } + // Keep this ordering: wallets reference the user with SET NULL, so + // purge the guarded provisioning wallet before deleting its owner. + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + if self.test_api_key_wallet_exists_for_user(user_id).await? { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: wallet ownership is unknown" + ))); + } + let wallet = store + .lock() + .expect("auth wallet store should lock") + .values() + .find(|wallet| { + wallet.user_id.as_deref() == Some(user_id) + && wallet_id.is_none_or(|wallet_id| wallet.id == wallet_id) + }) + .cloned(); + if wallet_id.is_none() && wallet.is_some() { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: wallet ownership is unknown" + ))); + } + if wallet_id.is_some() && wallet.is_none() { + let supplied_wallet_exists = store + .lock() + .expect("auth wallet store should lock") + .values() + .any(|wallet| wallet_id.is_some_and(|expected| wallet.id == expected)); + if supplied_wallet_exists { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: supplied wallet still exists" + ))); + } + let any_wallet = store + .lock() + .expect("auth wallet store should lock") + .values() + .any(|wallet| wallet.user_id.as_deref() == Some(user_id)); + if any_wallet { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: wallet ownership is unknown" + ))); + } + } + if let Some(wallet) = wallet { + let structurally_removable = wallet.api_key_id.is_none() + && wallet.balance == 0.0 + && wallet.gift_balance >= 0.0 + && wallet.total_recharged == 0.0 + && wallet.total_consumed == 0.0 + && wallet.total_refunded == 0.0 + && wallet.total_adjusted == wallet.gift_balance + && wallet.status == "active" + && matches!(wallet.limit_mode.as_str(), "finite" | "unlimited") + && wallet.currency == "USD"; + if !structurally_removable { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: wallet is not eligible for rollback" + ))); + } + let wallet_id = wallet.id; + let has_order = + self.admin_wallet_payment_order_store + .as_ref() + .is_some_and(|orders| { + orders + .lock() + .expect("admin wallet payment order store should lock") + .values() + .any(|order| order.wallet_id == wallet_id) + }); + let has_transaction = + self.admin_wallet_transaction_store + .as_ref() + .is_some_and(|transactions| { + let wallet_transactions = transactions + .lock() + .expect("admin wallet transaction store should lock") + .values() + .filter(|transaction| transaction.wallet_id == wallet_id) + .cloned() + .collect::>(); + if wallet_transactions.is_empty() { + return false; + } + let gift_balance = store + .lock() + .expect("auth wallet store should lock") + .get(&wallet_id) + .map(|wallet| wallet.gift_balance) + .unwrap_or_default(); + !(wallet_transactions.len() == 1 + && gift_balance > 0.0 + && wallet_transactions[0].category == "gift" + && wallet_transactions[0].reason_code == "gift_initial" + && wallet_transactions[0].amount == gift_balance + && wallet_transactions[0].balance_before == 0.0 + && wallet_transactions[0].balance_after == gift_balance + && wallet_transactions[0].recharge_balance_before == 0.0 + && wallet_transactions[0].recharge_balance_after == 0.0 + && wallet_transactions[0].gift_balance_before == 0.0 + && wallet_transactions[0].gift_balance_after == gift_balance + && wallet_transactions[0].link_type.as_deref() + == Some("system_task") + && wallet_transactions[0].link_id.as_deref() == Some(user_id) + && wallet_transactions[0].operator_id.is_none()) + }); + let has_refund = self + .admin_wallet_refund_store + .as_ref() + .is_some_and(|refunds| { + refunds + .lock() + .expect("admin wallet refund store should lock") + .values() + .any(|refund| refund.wallet_id == wallet_id) + }); + if has_order || has_transaction || has_refund { + return Err(GatewayError::Internal(format!( + "refusing to delete provisional auth user {user_id}: wallet has financial activity" + ))); + } + if let Some(transactions) = self.admin_wallet_transaction_store.as_ref() { + transactions + .lock() + .expect("admin wallet transaction store should lock") + .retain(|_, transaction| transaction.wallet_id != wallet_id); + } + store + .lock() + .expect("auth wallet store should lock") + .remove(&wallet_id); + } + self.delete_local_auth_user(user_id).await?; + return Ok(()); + } + + self.data + .rollback_provisional_auth_user_with_wallet(user_id, wallet_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + Ok(()) + } + pub(crate) async fn read_user_model_capability_settings( &self, user_id: &str, @@ -122,114 +655,160 @@ impl AppState { initial_gift_usd: f64, unlimited: bool, ) -> Result, GatewayError> { + Ok(self + .get_or_create_ldap_auth_user_with_wallet_outcome( + email, + username, + ldap_dn, + ldap_username, + logged_in_at, + initial_gift_usd, + unlimited, + ) + .await? + .map(|result| result.user)) + } + + #[allow(clippy::too_many_arguments)] + pub(crate) async fn get_or_create_ldap_auth_user_with_wallet_outcome( + &self, + email: String, + username: String, + ldap_dn: Option, + ldap_username: Option, + logged_in_at: chrono::DateTime, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, GatewayError> { + let Some(email) = normalize_ldap_identity_email(&email) else { + return Ok(None); + }; #[cfg(test)] - if let (Some(user_store), Some(wallet_store)) = ( + if let (Some(user_store), Some(_wallet_store)) = ( self.auth_user_store.as_ref(), self.auth_wallet_store.as_ref(), ) { - let mut users = user_store.lock().expect("auth user store should lock"); - let existing_id = users - .values() - .find(|user| { - user.email.as_deref() == Some(email.as_str()) - || user.username == username - || ldap_username - .as_deref() - .is_some_and(|value| user.username == value) - }) - .map(|user| user.id.clone()); - - if let Some(existing_id) = existing_id { - let Some(user) = users.get_mut(&existing_id) else { - return Ok(None); - }; - if user.is_deleted || !user.is_active { - return Ok(None); - } - if !user.auth_source.eq_ignore_ascii_case("ldap") { - return Ok(None); - } - user.email = Some(email); - user.email_verified = true; - user.last_login_at = Some(logged_in_at); - return Ok(Some(user.clone())); + if !initial_gift_usd.is_finite() { + return Err(GatewayError::Internal( + "initial gift amount must be finite".to_string(), + )); } + let user = { + let mut users = user_store.lock().expect("auth user store should lock"); + let existing_id = users + .values() + .find(|user| { + user.email.as_deref() == Some(email.as_str()) + || user.username == username + || ldap_username + .as_deref() + .is_some_and(|value| user.username == value) + }) + .map(|user| user.id.clone()); - let base_username = ldap_username - .as_deref() - .filter(|value| !value.trim().is_empty()) - .unwrap_or(username.as_str()) - .trim() - .to_string(); - let mut candidate_username = base_username.clone(); - while users - .values() - .any(|user| user.username == candidate_username) - { - let suffix = uuid::Uuid::new_v4().simple().to_string(); - candidate_username = format!( - "{}_ldap_{}{}", - base_username, - logged_in_at.timestamp(), - &suffix[..4] - ); - } + if let Some(existing_id) = existing_id { + let Some(user) = users.get_mut(&existing_id) else { + return Ok(None); + }; + if user.is_deleted || !user.is_active { + return Ok(None); + } + if !user.auth_source.eq_ignore_ascii_case("ldap") { + return Ok(None); + } + user.email = Some(email); + user.email_verified = true; + user.last_login_at = Some(logged_in_at); + return Ok(Some(LdapAuthProvisioningResult { + user: user.clone(), + owned_wallet_id: None, + })); + } - let user = aether_data::repository::users::StoredUserAuthRecord::new( - uuid::Uuid::new_v4().to_string(), - Some(email), - true, - candidate_username, - None, - "user".to_string(), - "ldap".to_string(), - None, - None, - None, - true, - false, - Some(logged_in_at), - Some(logged_in_at), - ) - .map_err(|err| GatewayError::Internal(err.to_string()))?; - users.insert(user.id.clone(), user.clone()); - drop(users); + let base_username = ldap_username + .as_deref() + .filter(|value| !value.trim().is_empty()) + .unwrap_or(username.as_str()) + .trim() + .to_string(); + let mut candidate_username = base_username.clone(); + while users + .values() + .any(|user| user.username == candidate_username) + { + let suffix = uuid::Uuid::new_v4().simple().to_string(); + candidate_username = format!( + "{}_ldap_{}{}", + base_username, + logged_in_at.timestamp(), + &suffix[..4] + ); + } - let gift_balance = if unlimited { - 0.0 - } else { - initial_gift_usd.max(0.0) + let user = aether_data::repository::users::StoredUserAuthRecord::new( + uuid::Uuid::new_v4().to_string(), + Some(email), + true, + candidate_username, + None, + "user".to_string(), + "ldap".to_string(), + None, + None, + None, + true, + false, + Some(logged_in_at), + Some(logged_in_at), + ) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + users.insert(user.id.clone(), user.clone()); + user }; - let wallet = aether_data::repository::wallet::StoredWalletSnapshot::new( - uuid::Uuid::new_v4().to_string(), - Some(user.id.clone()), - None, - 0.0, - gift_balance, - if unlimited { - "unlimited".to_string() - } else { - "finite".to_string() - }, - "USD".to_string(), - "active".to_string(), - 0.0, - 0.0, - 0.0, - gift_balance, - logged_in_at.timestamp(), - ) - .map_err(|err| GatewayError::Internal(err.to_string()))?; - wallet_store - .lock() - .expect("auth wallet store should lock") - .insert(wallet.id.clone(), wallet); let _ = ldap_dn; - return Ok(Some(user)); + + let initialized = match self + .initialize_auth_user_wallet_with_outcome(&user.id, initial_gift_usd, unlimited) + .await + { + Ok(Some(initialized)) => initialized, + Ok(None) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; + return Ok(None); + } + Err(err) => { + let _ = self + .rollback_provisional_auth_user_with_wallet(&user.id, None) + .await; + return Err(err); + } + }; + let wallet_is_user_owned = initialized.wallet.user_id.as_deref() + == Some(user.id.as_str()) + && initialized.wallet.api_key_id.is_none(); + if !wallet_is_user_owned { + let owned_wallet_id = initialized.created.then(|| initialized.wallet.id.clone()); + let _ = self + .rollback_provisional_auth_user_with_wallet( + &user.id, + owned_wallet_id.as_deref(), + ) + .await; + return Err(GatewayError::Internal( + "LDAP user wallet owner does not match the provisioned user".to_string(), + )); + } + return Ok(Some(LdapAuthProvisioningResult { + user, + owned_wallet_id: initialized.created.then(|| initialized.wallet.id), + })); } - self.data - .get_or_create_ldap_auth_user( + let result = self + .data + .get_or_create_ldap_auth_user_with_wallet_outcome( email, username, ldap_dn, @@ -239,7 +818,11 @@ impl AppState { unlimited, ) .await - .map_err(|err| GatewayError::Internal(err.to_string())) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + Ok(result.map(|result| LdapAuthProvisioningResult { + user: result.user, + owned_wallet_id: result.owned_wallet_id, + })) } pub(crate) async fn initialize_auth_user_wallet( @@ -250,6 +833,16 @@ impl AppState { ) -> Result, GatewayError> { #[cfg(test)] if let Some(store) = self.auth_wallet_store.as_ref() { + if user_id.trim().is_empty() { + return Err(GatewayError::Internal( + "user id is required to initialize a wallet".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(GatewayError::Internal( + "initial gift amount must be finite".to_string(), + )); + } let gift_balance = if unlimited { 0.0 } else { @@ -279,10 +872,17 @@ impl AppState { now_unix_secs, ) .map_err(|err| GatewayError::Internal(err.to_string()))?; - store - .lock() - .expect("auth wallet store should lock") - .insert(wallet.id.clone(), wallet.clone()); + let mut wallets = store.lock().expect("auth wallet store should lock"); + if let Some(existing) = wallets.values().find(|existing| { + existing.user_id.as_deref() == Some(user_id) && existing.api_key_id.is_none() + }) { + let existing = existing.clone(); + drop(wallets); + return Ok(Some(existing)); + } + wallets.insert(wallet.id.clone(), wallet.clone()); + drop(wallets); + record_test_initial_gift_transaction(self, &wallet, user_id, "用户初始赠款"); self.invalidate_auth_context_cache(); return Ok(Some(wallet)); } @@ -298,6 +898,83 @@ impl AppState { Ok(wallet) } + pub(crate) async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, GatewayError> + { + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + if user_id.trim().is_empty() { + return Err(GatewayError::Internal( + "user id is required to initialize a wallet".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(GatewayError::Internal( + "initial gift amount must be finite".to_string(), + )); + } + let gift_balance = if unlimited { + 0.0 + } else { + initial_gift_usd.max(0.0) + }; + let now_unix_secs = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64; + let wallet = aether_data::repository::wallet::StoredWalletSnapshot::new( + uuid::Uuid::new_v4().to_string(), + Some(user_id.to_string()), + None, + 0.0, + gift_balance, + if unlimited { + "unlimited".to_string() + } else { + "finite".to_string() + }, + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + gift_balance, + now_unix_secs, + ) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut wallets = store.lock().expect("auth wallet store should lock"); + if let Some(existing) = wallets.values().find(|existing| { + existing.user_id.as_deref() == Some(user_id) && existing.api_key_id.is_none() + }) { + return Ok(Some( + aether_data::repository::wallet::InitializeAuthWalletOutcome { + wallet: existing.clone(), + created: false, + }, + )); + } + wallets.insert(wallet.id.clone(), wallet.clone()); + drop(wallets); + record_test_initial_gift_transaction(self, &wallet, user_id, "用户初始赠款"); + self.invalidate_auth_context_cache(); + return Ok(Some( + aether_data::repository::wallet::InitializeAuthWalletOutcome { + wallet, + created: true, + }, + )); + } + + self.data + .initialize_auth_user_wallet_with_outcome(user_id, initial_gift_usd, unlimited) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn initialize_auth_api_key_wallet( &self, api_key_id: &str, @@ -306,6 +983,16 @@ impl AppState { ) -> Result, GatewayError> { #[cfg(test)] if let Some(store) = self.auth_wallet_store.as_ref() { + if api_key_id.trim().is_empty() { + return Err(GatewayError::Internal( + "api key id is required to initialize a wallet".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(GatewayError::Internal( + "initial gift amount must be finite".to_string(), + )); + } let gift_balance = if unlimited { 0.0 } else { @@ -335,10 +1022,22 @@ impl AppState { now_unix_secs, ) .map_err(|err| GatewayError::Internal(err.to_string()))?; - store - .lock() - .expect("auth wallet store should lock") - .insert(wallet.id.clone(), wallet.clone()); + let mut wallets = store.lock().expect("auth wallet store should lock"); + if let Some(existing) = wallets.values().find(|existing| { + existing.api_key_id.as_deref() == Some(api_key_id) && existing.user_id.is_none() + }) { + let existing = existing.clone(); + drop(wallets); + return Ok(Some(existing)); + } + wallets.insert(wallet.id.clone(), wallet.clone()); + drop(wallets); + record_test_initial_gift_transaction( + self, + &wallet, + api_key_id, + "独立余额 Key 初始赠款", + ); self.invalidate_auth_context_cache(); return Ok(Some(wallet)); } @@ -354,6 +1053,88 @@ impl AppState { Ok(wallet) } + pub(crate) async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, GatewayError> + { + #[cfg(test)] + if let Some(store) = self.auth_wallet_store.as_ref() { + if api_key_id.trim().is_empty() { + return Err(GatewayError::Internal( + "api key id is required to initialize a wallet".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(GatewayError::Internal( + "initial gift amount must be finite".to_string(), + )); + } + let gift_balance = if unlimited { + 0.0 + } else { + initial_gift_usd.max(0.0) + }; + let now_unix_secs = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64; + let wallet = aether_data::repository::wallet::StoredWalletSnapshot::new( + uuid::Uuid::new_v4().to_string(), + None, + Some(api_key_id.to_string()), + 0.0, + gift_balance, + if unlimited { + "unlimited".to_string() + } else { + "finite".to_string() + }, + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + gift_balance, + now_unix_secs, + ) + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let mut wallets = store.lock().expect("auth wallet store should lock"); + if let Some(existing) = wallets.values().find(|existing| { + existing.api_key_id.as_deref() == Some(api_key_id) && existing.user_id.is_none() + }) { + return Ok(Some( + aether_data::repository::wallet::InitializeAuthWalletOutcome { + wallet: existing.clone(), + created: false, + }, + )); + } + wallets.insert(wallet.id.clone(), wallet.clone()); + drop(wallets); + record_test_initial_gift_transaction( + self, + &wallet, + api_key_id, + "独立余额 Key 初始赠款", + ); + self.invalidate_auth_context_cache(); + return Ok(Some( + aether_data::repository::wallet::InitializeAuthWalletOutcome { + wallet, + created: true, + }, + )); + } + + self.data + .initialize_auth_api_key_wallet_with_outcome(api_key_id, initial_gift_usd, unlimited) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn update_auth_user_wallet_limit_mode( &self, user_id: &str, @@ -552,3 +1333,314 @@ impl AppState { Ok(wallet) } } + +#[cfg(test)] +mod tests { + use super::*; + use aether_data::repository::users::StoredUserAuthRecord; + use aether_data::repository::wallet::{StoredWalletSnapshot, WalletLookupKey}; + + #[test] + fn ldap_identity_email_uses_the_canonical_account_namespace() { + assert_eq!( + normalize_ldap_identity_email(" Alice@Example.COM ").as_deref(), + Some("alice@example.com") + ); + assert_eq!(normalize_ldap_identity_email(" "), None); + } + + fn provisional_user(user_id: &str) -> StoredUserAuthRecord { + let now = chrono::Utc::now(); + StoredUserAuthRecord::new( + user_id.to_string(), + Some(format!("{user_id}@example.com")), + true, + user_id.to_string(), + Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build") + } + + #[tokio::test] + async fn ldap_provisioning_preserves_only_new_wallet_id() { + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests(Vec::::new()) + .with_auth_wallets_for_tests(Vec::::new()); + let logged_in_at = chrono::Utc::now(); + + let first = state + .get_or_create_ldap_auth_user_with_wallet_outcome( + " Alice@Example.COM ".to_string(), + "alice".to_string(), + Some("uid=alice,dc=example".to_string()), + Some("alice".to_string()), + logged_in_at, + 10.0, + false, + ) + .await + .expect("LDAP provisioning should succeed") + .expect("new LDAP user should be returned"); + let first_wallet_id = first + .owned_wallet_id + .clone() + .expect("new provisioning should report its wallet id"); + assert_eq!(first.user.email.as_deref(), Some("alice@example.com")); + assert!(state + .find_wallet(WalletLookupKey::WalletId(&first_wallet_id)) + .await + .expect("wallet lookup should succeed") + .is_some()); + + let replay = state + .get_or_create_ldap_auth_user_with_wallet_outcome( + "ALICE@example.com".to_string(), + "alice".to_string(), + None, + Some("alice".to_string()), + logged_in_at, + 99.0, + false, + ) + .await + .expect("LDAP replay should succeed") + .expect("existing LDAP user should be returned"); + assert_eq!(replay.user.id, first.user.id); + assert_eq!(replay.owned_wallet_id, None); + assert_eq!( + state + .find_wallet(WalletLookupKey::WalletId(&first_wallet_id)) + .await + .expect("wallet lookup should succeed") + .expect("existing wallet should remain") + .gift_balance, + 10.0 + ); + } + + #[tokio::test] + async fn test_store_rollback_preserves_user_with_active_wallet() { + let user_id = "active-provisional-user"; + let wallet = StoredWalletSnapshot::new( + "active-provisional-wallet".to_string(), + Some(user_id.to_string()), + None, + 1.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 1.0, + 0.0, + 0.0, + 0.0, + chrono::Utc::now().timestamp(), + ) + .expect("wallet should build"); + let wallet_id = wallet.id.clone(); + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests([provisional_user(user_id)]) + .with_auth_wallets_for_tests([wallet]); + + assert!(state + .rollback_provisional_auth_user_with_wallet(user_id, Some(wallet_id.as_str())) + .await + .is_err()); + assert!(state + .find_user_auth_by_id(user_id) + .await + .expect("user lookup should succeed") + .is_some()); + assert!(state + .find_wallet(WalletLookupKey::UserId(user_id)) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn test_store_rollback_rejects_wallet_id_owned_by_another_user() { + let target_user_id = "target-provisional-user"; + let other_user_id = "other-user"; + let other_wallet = StoredWalletSnapshot::new( + "other-user-wallet".to_string(), + Some(other_user_id.to_string()), + None, + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + chrono::Utc::now().timestamp(), + ) + .expect("wallet should build"); + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests([provisional_user(target_user_id)]) + .with_auth_wallets_for_tests([other_wallet]); + + assert!(state + .rollback_provisional_auth_user_with_wallet(target_user_id, Some("other-user-wallet"),) + .await + .is_err()); + assert!(state + .find_user_auth_by_id(target_user_id) + .await + .expect("user lookup should succeed") + .is_some()); + assert!(state + .find_wallet(WalletLookupKey::WalletId("other-user-wallet")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn test_store_rollback_removes_user_when_wallet_is_absent() { + let user_id = "no-wallet-provisional-user"; + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests([provisional_user(user_id)]); + + state + .rollback_provisional_auth_user_with_wallet(user_id, Some("missing-wallet")) + .await + .expect("confirmed wallet absence should allow rollback"); + assert!(state + .find_user_auth_by_id(user_id) + .await + .expect("user lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn test_store_rollback_preserves_user_when_api_key_wallet_exists() { + let user_id = "api-key-wallet-provisional-user"; + let api_key_wallet = StoredWalletSnapshot::new( + "api-key-wallet-provisional".to_string(), + None, + Some("api-key-owned-by-user".to_string()), + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + chrono::Utc::now().timestamp(), + ) + .expect("api-key wallet should build"); + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests([provisional_user(user_id)]) + .with_auth_wallets_for_tests([api_key_wallet.clone()]); + + assert!(state + .rollback_provisional_auth_user_with_wallet(user_id, None) + .await + .is_err()); + assert!(state + .find_user_auth_by_id(user_id) + .await + .expect("user lookup should succeed") + .is_some()); + assert!(state + .find_wallet(WalletLookupKey::ApiKeyId("api-key-owned-by-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn test_store_wallet_initialization_records_initial_gift_once_and_rolls_back() { + let user_id = "gift-provisional-user"; + let state = AppState::new() + .expect("state should build") + .with_auth_users_for_tests([provisional_user(user_id)]); + + let first = state + .initialize_auth_user_wallet_with_outcome(user_id, 7.5, false) + .await + .expect("wallet initialization should resolve") + .expect("wallet should be available"); + assert!(first.created); + + let transaction_store = state + .admin_wallet_transaction_store + .as_ref() + .expect("transaction store should be available") + .clone(); + { + let transactions = transaction_store + .lock() + .expect("transaction store should lock"); + assert_eq!(transactions.len(), 1); + let transaction = transactions + .values() + .next() + .expect("initial gift transaction should exist"); + assert_eq!(transaction.wallet_id, first.wallet.id); + assert_eq!(transaction.category, "gift"); + assert_eq!(transaction.reason_code, "gift_initial"); + assert_eq!(transaction.amount, 7.5); + assert_eq!(transaction.balance_before, 0.0); + assert_eq!(transaction.balance_after, 7.5); + assert_eq!(transaction.gift_balance_before, 0.0); + assert_eq!(transaction.gift_balance_after, 7.5); + assert_eq!(transaction.link_type.as_deref(), Some("system_task")); + assert_eq!(transaction.link_id.as_deref(), Some(user_id)); + assert!(transaction.operator_id.is_none()); + } + + let replay = state + .initialize_auth_user_wallet_with_outcome(user_id, 99.0, false) + .await + .expect("wallet replay should resolve") + .expect("existing wallet should be returned"); + assert!(!replay.created); + assert_eq!(replay.wallet.id, first.wallet.id); + assert_eq!( + transaction_store + .lock() + .expect("transaction store should lock") + .len(), + 1 + ); + + state + .rollback_provisional_auth_user_with_wallet(user_id, Some(first.wallet.id.as_str())) + .await + .expect("rollback should remove the untouched wallet"); + assert!(state + .find_wallet(WalletLookupKey::UserId(user_id)) + .await + .expect("wallet lookup should resolve") + .is_none()); + assert!(state + .find_user_auth_by_id(user_id) + .await + .expect("user lookup should resolve") + .is_none()); + assert!(transaction_store + .lock() + .expect("transaction store should lock") + .is_empty()); + } +} diff --git a/apps/aether-gateway/src/state/runtime/billing/admin.rs b/apps/aether-gateway/src/state/runtime/billing/admin.rs index e40b9f671..5e05c1e79 100644 --- a/apps/aether-gateway/src/state/runtime/billing/admin.rs +++ b/apps/aether-gateway/src/state/runtime/billing/admin.rs @@ -2,10 +2,12 @@ use super::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState, BillingPlanRecord, BillingPlanWriteInput, GatewayError, LocalMutationOutcome, - PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord, - UserPlanEntitlementRecord, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; +const PAYMENT_GATEWAY_SECRET_MIGRATION_MAX_ATTEMPTS: usize = 8; + fn data_error(err: impl ToString) -> GatewayError { GatewayError::Internal(err.to_string()) } @@ -446,9 +448,73 @@ impl AppState { &self, provider: &str, ) -> Result, GatewayError> { - self.data - .find_payment_gateway_config(provider) + let provider = provider.trim().to_ascii_lowercase(); + let mut record = self + .data + .find_payment_gateway_config(&provider) .await + .map_err(data_error)?; + for _ in 0..PAYMENT_GATEWAY_SECRET_MIGRATION_MAX_ATTEMPTS { + let Some(mut current) = record else { + return Ok(None); + }; + let Some(observed) = current.merchant_key_encrypted.as_deref() else { + return Ok(Some(current)); + }; + let projection = crate::handlers::shared::open_payment_gateway_secret( + self, + &crate::handlers::shared::PaymentGatewaySecretBinding::from_record(¤t) + .map_err(|detail| { + GatewayError::Internal(format!( + "payment gateway secret binding is invalid for {}: {detail}", + current.provider + )) + })?, + observed, + ) + .map_err(|detail| { + GatewayError::Internal(format!( + "payment gateway secret integrity check failed for {}: {detail}", + current.provider + )) + })?; + if !projection.migration_required { + return Ok(Some(current)); + } + + let update = PaymentGatewaySecretCasUpdate { + provider: current.provider.clone(), + expected_merchant_key_encrypted: observed.to_string(), + merchant_key_encrypted: projection.protected.clone(), + }; + if self + .data + .compare_and_swap_payment_gateway_secret(&update) + .await + .map_err(data_error)? + { + current.merchant_key_encrypted = Some(projection.protected); + return Ok(Some(current)); + } + record = self + .data + .find_payment_gateway_config_strong(&provider) + .await + .map_err(data_error)?; + } + Err(GatewayError::Internal(format!( + "payment gateway secret migration did not converge for {provider}" + ))) + } + + pub(crate) async fn compare_and_swap_payment_gateway_config( + &self, + input: &PaymentGatewayConfigCasWriteInput, + ) -> Result, GatewayError> { + self.data + .compare_and_swap_payment_gateway_config(input) + .await + .map(local_mutation_outcome) .map_err(data_error) } @@ -499,11 +565,16 @@ impl AppState { plan_id: &str, input: &BillingPlanWriteInput, ) -> Result, GatewayError> { - self.data + let outcome = self + .data .update_billing_plan(plan_id, input) .await .map(local_mutation_outcome) - .map_err(data_error) + .map_err(data_error)?; + if matches!(&outcome, LocalMutationOutcome::Applied(_)) { + self.invalidate_auth_context_cache(); + } + Ok(outcome) } pub(crate) async fn set_billing_plan_enabled( @@ -539,6 +610,23 @@ impl AppState { .map_err(data_error) } + pub(crate) async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, GatewayError> { + let outcome = self + .data + .revoke_user_plan_entitlement(user_id, entitlement_id) + .await + .map(local_mutation_outcome) + .map_err(data_error)?; + if matches!(&outcome, LocalMutationOutcome::Applied(_)) { + self.invalidate_auth_context_cache(); + } + Ok(outcome) + } + pub(crate) async fn find_user_daily_quota_availability( &self, user_id: &str, @@ -586,13 +674,20 @@ impl AppState { #[cfg(test)] mod tests { + use std::sync::Arc; use std::time::Duration; + use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; + use aether_data::repository::billing::InMemoryBillingReadRepository; + use aether_data_contracts::repository::billing::{ + BillingReadRepository, PaymentGatewayConfigWriteInput, + }; use serde_json::json; use super::{ AdminBillingCollectorWriteInput, AdminBillingRuleWriteInput, AppState, LocalMutationOutcome, }; + use crate::data::GatewayDataState; const CACHE_KEY: &str = "billing-mutation-test"; @@ -714,4 +809,59 @@ mod tests { .get(&CACHE_KEY.to_string(), Duration::from_secs(60),) .is_some()); } + + #[tokio::test] + async fn gateway_lookup_lazily_migrates_legacy_secret_without_touching_other_fields() { + let repository = Arc::new(InMemoryBillingReadRepository::default()); + let legacy = encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + r#"{"secret_key":"legacy-value"}"#, + ) + .expect("legacy secret should encrypt"); + repository + .upsert_payment_gateway_config(&PaymentGatewayConfigWriteInput { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: "https://api.stripe.com".to_string(), + callback_base_url: Some("https://example.com".to_string()), + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some(legacy), + preserve_existing_secret: false, + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: json!({"channels": []}), + }) + .await + .expect("gateway seed should succeed"); + let before = repository + .find_payment_gateway_config("stripe") + .await + .expect("gateway lookup should succeed") + .expect("gateway should exist"); + let data = GatewayDataState::with_billing_reader_for_tests(repository.clone()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new() + .expect("app state should build") + .with_data_state_for_tests(data); + + let migrated = state + .find_payment_gateway_config("stripe") + .await + .expect("gateway migration should succeed") + .expect("gateway should exist"); + assert!(migrated + .merchant_key_encrypted + .as_deref() + .is_some_and(|value| value.starts_with("aether-payment-gateway-secret-v3:"))); + let stored = repository + .find_payment_gateway_config("stripe") + .await + .expect("gateway lookup should succeed") + .expect("gateway should exist"); + assert_eq!(stored, migrated); + assert_eq!(stored.updated_at_unix_secs, before.updated_at_unix_secs); + assert_eq!(stored.endpoint_url, before.endpoint_url); + assert_eq!(stored.channels_json, before.channels_json); + } } diff --git a/apps/aether-gateway/src/state/runtime/billing/finance_queries.rs b/apps/aether-gateway/src/state/runtime/billing/finance_queries.rs index aea35cf79..5c0d5954d 100644 --- a/apps/aether-gateway/src/state/runtime/billing/finance_queries.rs +++ b/apps/aether-gateway/src/state/runtime/billing/finance_queries.rs @@ -74,7 +74,7 @@ impl AppState { let effective_status = if order.status == "pending" && order .expires_at_unix_secs - .is_some_and(|value| value < now_unix_secs) + .is_some_and(|value| value <= now_unix_secs) { "expired" } else { diff --git a/apps/aether-gateway/src/state/runtime/billing/mod.rs b/apps/aether-gateway/src/state/runtime/billing/mod.rs index c8a81e3c8..096d65735 100644 --- a/apps/aether-gateway/src/state/runtime/billing/mod.rs +++ b/apps/aether-gateway/src/state/runtime/billing/mod.rs @@ -2,8 +2,8 @@ use super::super::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState, BillingPlanRecord, BillingPlanWriteInput, GatewayError, LocalMutationOutcome, - PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord, - UserPlanEntitlementRecord, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; mod admin; diff --git a/apps/aether-gateway/src/state/runtime/candidate_queries.rs b/apps/aether-gateway/src/state/runtime/candidate_queries.rs index edc728a76..eef5ef3f4 100644 --- a/apps/aether-gateway/src/state/runtime/candidate_queries.rs +++ b/apps/aether-gateway/src/state/runtime/candidate_queries.rs @@ -111,8 +111,9 @@ impl AppState { pub(crate) async fn upsert_request_candidate( &self, - candidate: candidates::UpsertRequestCandidateRecord, + mut candidate: candidates::UpsertRequestCandidateRecord, ) -> Result, GatewayError> { + candidate.sanitize_for_persistence(); if let Some(queue) = self.request_candidate_queue.as_ref() { let stored = stored_request_candidate_from_upsert(&candidate)?; queue @@ -136,8 +137,9 @@ impl AppState { /// the async queue is enabled. pub(crate) async fn enqueue_request_candidate_status( &self, - candidate: candidates::UpsertRequestCandidateRecord, + mut candidate: candidates::UpsertRequestCandidateRecord, ) -> Result, GatewayError> { + candidate.sanitize_for_persistence(); if let Some(queue) = self.request_candidate_queue.as_ref() { queue .enqueue_or_fallback(candidate) @@ -158,8 +160,9 @@ impl AppState { /// when the queue is disabled or closed. pub(crate) fn try_enqueue_request_candidate_status( &self, - candidate: candidates::UpsertRequestCandidateRecord, + mut candidate: candidates::UpsertRequestCandidateRecord, ) -> Result<(), candidates::UpsertRequestCandidateRecord> { + candidate.sanitize_for_persistence(); let Some(queue) = self.request_candidate_queue.as_ref() else { return Err(candidate); }; diff --git a/apps/aether-gateway/src/state/runtime/gemini_files.rs b/apps/aether-gateway/src/state/runtime/gemini_files.rs index 2fcdd4875..c37c4de9c 100644 --- a/apps/aether-gateway/src/state/runtime/gemini_files.rs +++ b/apps/aether-gateway/src/state/runtime/gemini_files.rs @@ -14,6 +14,19 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn upsert_gemini_file_mapping_if_owner_matches( + &self, + record: aether_data::repository::gemini_file_mappings::UpsertGeminiFileMappingRecord, + ) -> Result< + Option, + GatewayError, + > { + self.data + .upsert_gemini_file_mapping_if_owner_matches(record) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn list_gemini_file_mappings( &self, query: &aether_data::repository::gemini_file_mappings::GeminiFileMappingListQuery, @@ -27,6 +40,50 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn find_gemini_file_mapping_by_file_name( + &self, + file_name: &str, + ) -> Result< + Option, + GatewayError, + > { + self.data + .find_gemini_file_mapping_by_file_name(file_name) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn find_active_gemini_file_mapping_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result< + Option, + GatewayError, + > { + self.data + .find_active_gemini_file_mapping_for_user(file_name, user_id, now_unix_secs) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn find_active_gemini_file_mapping_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result< + Option, + GatewayError, + > { + self.data + .find_active_gemini_file_mapping_for_owner(file_name, key_id, user_id, now_unix_secs) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn summarize_gemini_file_mappings( &self, now_unix_secs: u64, @@ -48,6 +105,29 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + self.data + .delete_gemini_file_mapping_by_file_name_for_user(file_name, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + self.data + .delete_gemini_file_mapping_by_file_name_for_owner(file_name, key_id, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn delete_gemini_file_mapping_by_id( &self, mapping_id: &str, diff --git a/apps/aether-gateway/src/state/runtime/monitoring.rs b/apps/aether-gateway/src/state/runtime/monitoring.rs index 9d5ab5058..f79e6925d 100644 --- a/apps/aether-gateway/src/state/runtime/monitoring.rs +++ b/apps/aether-gateway/src/state/runtime/monitoring.rs @@ -1,5 +1,6 @@ use std::collections::BTreeMap; +use aether_admin::observability::usage::admin_usage_safe_metadata_value; use aether_data::repository::audit::AuditLogListQuery; use chrono::{DateTime, Utc}; use serde_json::{json, Value}; @@ -43,8 +44,8 @@ impl AppState { "description": record.description, "ip_address": record.ip_address, "status_code": record.status_code, - "error_message": record.error_message, - "metadata": record.metadata, + "error_message": record.error_message.as_ref().map(|_| "audit_event_failed"), + "metadata": sanitize_admin_audit_metadata(record.metadata.as_ref()), "created_at": record.created_at_rfc3339(), }) }) @@ -72,7 +73,7 @@ impl AppState { "user_id": record.user_id, "description": record.description, "ip_address": record.ip_address, - "metadata": record.metadata, + "metadata": sanitize_admin_audit_metadata(record.metadata.as_ref()), "created_at": record.created_at_rfc3339(), }) }) @@ -134,3 +135,41 @@ impl AppState { fn cutoff_unix_secs(cutoff_time: DateTime) -> u64 { cutoff_time.timestamp().max(0) as u64 } + +fn sanitize_admin_audit_metadata(metadata: Option<&Value>) -> Value { + metadata + .map(admin_usage_safe_metadata_value) + .unwrap_or(Value::Null) +} + +#[cfg(test)] +mod tests { + use super::sanitize_admin_audit_metadata; + use serde_json::json; + + #[test] + fn admin_audit_metadata_drops_credentials_and_url_components() { + let metadata = sanitize_admin_audit_metadata(Some(&json!({ + "category": "security", + "authorization": "Bearer audit-secret", + "nested": { + "refresh_token": "refresh-secret", + "endpoint_url": "https://user:password@example.test/v1?token=query-secret#fragment", + "safe_count": 2 + } + }))); + + assert_eq!(metadata["category"], "security"); + assert!(metadata.get("authorization").is_none()); + assert!(metadata["nested"].get("refresh_token").is_none()); + assert_eq!( + metadata["nested"]["endpoint_url"], + "https://example.test/v1" + ); + assert_eq!(metadata["nested"]["safe_count"], 2); + let encoded = metadata.to_string(); + for secret in ["audit-secret", "refresh-secret", "password", "query-secret"] { + assert!(!encoded.contains(secret), "leaked {secret}"); + } + } +} diff --git a/apps/aether-gateway/src/state/runtime/payments.rs b/apps/aether-gateway/src/state/runtime/payments.rs index 0ac9e4f74..11e93e222 100644 --- a/apps/aether-gateway/src/state/runtime/payments.rs +++ b/apps/aether-gateway/src/state/runtime/payments.rs @@ -142,7 +142,7 @@ impl AppState { } if order .expires_at_unix_secs - .is_some_and(|value| value < chrono::Utc::now().timestamp().max(0) as u64) + .is_some_and(|value| value <= chrono::Utc::now().timestamp().max(0) as u64) { return Ok(AdminWalletMutationOutcome::Invalid( "payment order expired".to_string(), diff --git a/apps/aether-gateway/src/state/runtime/referrals.rs b/apps/aether-gateway/src/state/runtime/referrals.rs index e1e1b9c1f..fdcc70529 100644 --- a/apps/aether-gateway/src/state/runtime/referrals.rs +++ b/apps/aether-gateway/src/state/runtime/referrals.rs @@ -4,13 +4,39 @@ use crate::data::state::{ }; use crate::{AppState, GatewayError}; use axum::http::StatusCode; +use tracing::warn; + +const REFERRAL_INVALID_INPUT_FALLBACK: &str = "返利请求无效"; + +fn safe_referral_invalid_input_detail(detail: &str) -> &'static str { + // These messages are deliberate domain-level validation responses. Any + // future adapter/storage detail must stay server-side instead of becoming + // an oracle for database state or schema information. + match detail { + "邀请码无效" => "邀请码无效", + "不能使用自己的邀请码注册" => "不能使用自己的邀请码注册", + "仅失败返利可以补发" => "仅失败返利可以补发", + "返利金额无效,无法补发" => "返利金额无效,无法补发", + _ => REFERRAL_INVALID_INPUT_FALLBACK, + } +} fn referral_data_error(err: aether_data::DataLayerError) -> GatewayError { match err { - aether_data::DataLayerError::InvalidInput(detail) => GatewayError::Client { - status: StatusCode::BAD_REQUEST, - message: detail, - }, + aether_data::DataLayerError::InvalidInput(detail) => { + let safe_detail = safe_referral_invalid_input_detail(&detail); + if safe_detail == REFERRAL_INVALID_INPUT_FALLBACK { + warn!( + event_name = "referral_invalid_input_hidden", + error_length = detail.len(), + "referral data-layer validation detail hidden from client" + ); + } + GatewayError::Client { + status: StatusCode::BAD_REQUEST, + message: safe_detail.to_string(), + } + } other => GatewayError::Internal(other.to_string()), } } @@ -51,6 +77,15 @@ fn config_f64(value: Option<&serde_json::Value>, default: f64) -> f64 { } } +fn config_percent(value: Option<&serde_json::Value>) -> f64 { + let value = config_f64(value, 0.0); + if value.is_finite() && value > 0.0 && value <= 100.0 { + value + } else { + 0.0 + } +} + impl AppState { pub(crate) fn has_referral_data_backend(&self) -> bool { self.data.has_referral_data_backend() @@ -93,7 +128,7 @@ impl AppState { config_string(headcount_trigger.as_ref()).unwrap_or_else(|| "registration".to_string()); Ok(Some(ReferralRewardConfig { percent_enabled: matches!(mode.as_str(), "percent" | "both"), - percent_rate: config_f64(percent.as_ref(), 0.0), + percent_rate: config_percent(percent.as_ref()), headcount_enabled: matches!(mode.as_str(), "headcount" | "both"), headcount_amount_usd: config_f64(headcount_amount.as_ref(), 0.0), headcount_trigger, @@ -238,4 +273,43 @@ impl AppState { .await .map_err(|err| GatewayError::Internal(err.to_string())) } + + pub(crate) async fn reconcile_referral_rewards_once( + &self, + ) -> Result { + let config = self.referral_reward_config().await?; + self.data + .reconcile_referral_rewards_once(config) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } +} + +#[cfg(test)] +mod tests { + use super::{referral_data_error, REFERRAL_INVALID_INPUT_FALLBACK}; + + #[test] + fn referral_invalid_input_projection_allowlists_domain_messages() { + let known = super::referral_data_error(aether_data::DataLayerError::InvalidInput( + "邀请码无效".to_string(), + )); + match known { + crate::GatewayError::Client { message, .. } => assert_eq!(message, "邀请码无效"), + other => panic!("expected client error, got {other:?}"), + } + + let secret = "database table referral_rewards row reward-secret has invalid wallet"; + let unknown = referral_data_error(aether_data::DataLayerError::InvalidInput( + secret.to_string(), + )); + match unknown { + crate::GatewayError::Client { message, .. } => { + assert_eq!(message, REFERRAL_INVALID_INPUT_FALLBACK); + assert!(!message.contains("reward-secret")); + assert!(!message.contains("referral_rewards")); + } + other => panic!("expected client error, got {other:?}"), + } + } } diff --git a/apps/aether-gateway/src/state/runtime/wallet/mutations.rs b/apps/aether-gateway/src/state/runtime/wallet/mutations.rs index bfd2138e1..995ba626f 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/mutations.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/mutations.rs @@ -1,10 +1,12 @@ use aether_data::repository::wallet::{ - AdjustWalletBalanceInput, CompleteAdminWalletRefundInput, CreateManualWalletRechargeInput, - CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + AdjustWalletBalanceInput, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, + CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, - WalletMutationOutcome, + FailWalletRechargeCheckoutInput, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, + ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, + UpdateAdminWalletRefundGatewayInput, UpdateWalletRechargeCheckoutInput, WalletMutationOutcome, }; use crate::{AppState, GatewayError}; @@ -20,6 +22,55 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result< + Option>, + GatewayError, + > { + self.data + .update_wallet_recharge_checkout(input) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result, GatewayError> { + self.data + .compare_and_swap_payment_order_stripe_client_secret(input) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result< + Option>, + GatewayError, + > { + self.data + .fail_wallet_recharge_checkout(input) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result< + Option>, + GatewayError, + > { + self.data + .reclaim_wallet_recharge_checkout(input) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn create_plan_purchase_order( &self, input: CreatePlanPurchaseOrderInput, @@ -134,6 +185,19 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result< + Option>, + GatewayError, + > { + self.data + .update_admin_wallet_refund_gateway(input) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn fail_admin_wallet_refund( &self, input: FailAdminWalletRefundInput, diff --git a/apps/aether-gateway/src/state/runtime/wallet/reads.rs b/apps/aether-gateway/src/state/runtime/wallet/reads.rs index dd1d92ac3..90896a419 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/reads.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/reads.rs @@ -191,6 +191,18 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, GatewayError> + { + self.data + .find_wallet_recharge_order_by_order_no(user_id, order_no) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, @@ -203,6 +215,28 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, GatewayError> + { + self.data + .find_payment_order_by_order_no(order_no) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn find_payment_order_by_id( + &self, + order_id: &str, + ) -> Result, GatewayError> + { + self.data + .find_admin_payment_order(order_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn find_wallet_refund( &self, wallet_id: &str, diff --git a/apps/aether-gateway/src/state/runtime/wallet/refund_lifecycle.rs b/apps/aether-gateway/src/state/runtime/wallet/refund_lifecycle.rs index 011ca3f69..81858fe8c 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/refund_lifecycle.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/refund_lifecycle.rs @@ -2,6 +2,9 @@ use super::{ AdminWalletMutationOutcome, AdminWalletRefundRecord, AdminWalletTransactionRecord, AppState, GatewayError, }; +use aether_data::repository::wallet::{ + payment_order_refund_amounts_are_consistent, wallet_refund_proof_is_success, +}; impl AppState { pub(crate) async fn admin_process_wallet_refund( @@ -39,6 +42,11 @@ impl AppState { else { return Ok(AdminWalletMutationOutcome::NotFound); }; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(AdminWalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if !matches!(refund.status.as_str(), "approved" | "pending_approval") { return Ok(AdminWalletMutationOutcome::Invalid( "refund status is not approvable".to_string(), @@ -49,8 +57,26 @@ impl AppState { let mut updated_wallet = wallet.clone(); let before_recharge = updated_wallet.balance; let before_gift = updated_wallet.gift_balance; + let before_total_refunded = updated_wallet.total_refunded; let before_total = before_recharge + before_gift; let after_recharge = before_recharge - amount_usd; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded + amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + { + return Ok(AdminWalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } if after_recharge < 0.0 { return Ok(AdminWalletMutationOutcome::Invalid( "refund amount exceeds refundable recharge balance".to_string(), @@ -72,20 +98,41 @@ impl AppState { "payment order not found".to_string(), )); }; - if amount_usd > order.refundable_amount_usd { + if order.wallet_id != wallet_id || order.status != "credited" { return Ok(AdminWalletMutationOutcome::Invalid( - "refund amount exceeds refundable amount".to_string(), + "payment order is not refundable for this wallet".to_string(), + )); + } + let order_amount = order.amount_usd; + let refunded_before = order.refunded_amount_usd; + let refundable_before = order.refundable_amount_usd; + let refunded_after = refunded_before + amount_usd; + let refundable_after = refundable_before - amount_usd; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_before, + refundable_before, + ) || !refunded_after.is_finite() + || !refundable_after.is_finite() + || amount_usd > refundable_before + || refunded_after < 0.0 + || refunded_after > order_amount + || refundable_after < 0.0 + || refundable_after > order_amount + { + return Ok(AdminWalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), )); } let mut order = order; - order.refunded_amount_usd += amount_usd; - order.refundable_amount_usd -= amount_usd; + order.refunded_amount_usd = refunded_after; + order.refundable_amount_usd = refundable_after; updated_order = Some(order); } let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64; updated_wallet.balance = after_recharge; - updated_wallet.total_refunded = (updated_wallet.total_refunded + amount_usd).max(0.0); + updated_wallet.total_refunded = after_total_refunded; updated_wallet.updated_at_unix_secs = now_unix_secs; let transaction = AdminWalletTransactionRecord { @@ -95,7 +142,7 @@ impl AppState { reason_code: "refund_out".to_string(), amount: -amount_usd, balance_before: before_total, - balance_after: after_recharge + before_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -130,6 +177,12 @@ impl AppState { .expect("admin wallet payment order store should lock") .insert(updated_order.id.clone(), updated_order); } + if let Some(transaction_store) = self.admin_wallet_transaction_store.as_ref() { + transaction_store + .lock() + .expect("admin wallet transaction store should lock") + .insert(transaction.id.clone(), transaction.clone()); + } self.invalidate_auth_context_cache(); return Ok(AdminWalletMutationOutcome::Applied(( @@ -187,6 +240,23 @@ impl AppState { else { return Ok(AdminWalletMutationOutcome::NotFound); }; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(AdminWalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let (Some(existing_id), Some(incoming_id)) = + (refund.gateway_refund_id.as_deref(), gateway_refund_id) + { + if existing_id != incoming_id { + return Ok(AdminWalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if refund.status == "succeeded" { + return Ok(AdminWalletMutationOutcome::Applied(refund)); + } if refund.status != "processing" { return Ok(AdminWalletMutationOutcome::Invalid( "refund status must be processing before completion".to_string(), @@ -195,9 +265,22 @@ impl AppState { let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64; let mut updated_refund = refund; updated_refund.status = "succeeded".to_string(); - updated_refund.gateway_refund_id = gateway_refund_id.map(ToOwned::to_owned); - updated_refund.payout_reference = payout_reference.map(ToOwned::to_owned); - updated_refund.payout_proof = payout_proof; + updated_refund.gateway_refund_id = gateway_refund_id + .map(ToOwned::to_owned) + .or_else(|| updated_refund.gateway_refund_id.clone()); + updated_refund.payout_reference = payout_reference + .map(ToOwned::to_owned) + .or_else(|| updated_refund.payout_reference.clone()); + // Keep the durable provider response on ordinary retries. A + // terminal success proof is the only completion payload allowed + // to upgrade an earlier processing proof. + if updated_refund.payout_proof.is_none() + || payout_proof + .as_ref() + .is_some_and(wallet_refund_proof_is_success) + { + updated_refund.payout_proof = payout_proof; + } updated_refund.completed_at_unix_secs = Some(now_unix_secs); updated_refund.updated_at_unix_secs = now_unix_secs; refund_store @@ -269,6 +352,12 @@ impl AppState { return Ok(AdminWalletMutationOutcome::NotFound); }; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(AdminWalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64; if matches!(refund.status.as_str(), "pending_approval" | "approved") { let mut updated_refund = refund; @@ -292,15 +381,95 @@ impl AppState { ))); } + // The in-memory implementation mirrors the database contract: + // only an explicitly offline payout without external evidence may + // release its reservation. + if refund.gateway_refund_id.is_some() + || refund.payout_proof.is_some() + || !refund + .refund_mode + .trim() + .eq_ignore_ascii_case("offline_payout") + { + return Ok(AdminWalletMutationOutcome::Invalid( + "cannot fail refund while gateway settlement is processing".to_string(), + )); + } + let amount_usd = refund.amount_usd; let before_recharge = wallet.balance; let before_gift = wallet.gift_balance; + let before_total_refunded = wallet.total_refunded; let before_total = before_recharge + before_gift; let after_recharge = before_recharge + amount_usd; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded - amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || before_total_refunded < amount_usd + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + || after_total_refunded < 0.0 + { + return Ok(AdminWalletMutationOutcome::Invalid( + "wallet balance is invalid for refund recovery".to_string(), + )); + } + + let mut updated_order = None; + if let Some(payment_order_id) = refund.payment_order_id.clone() { + let Some(order_store) = self.admin_wallet_payment_order_store.as_ref() else { + return Ok(AdminWalletMutationOutcome::Unavailable); + }; + let Some(order) = order_store + .lock() + .expect("admin wallet payment order store should lock") + .get(&payment_order_id) + .cloned() + else { + return Ok(AdminWalletMutationOutcome::Invalid( + "payment order not found".to_string(), + )); + }; + if order.wallet_id != wallet_id || order.status != "credited" { + return Ok(AdminWalletMutationOutcome::Invalid( + "payment order is not refundable for this wallet".to_string(), + )); + } + let refunded_before = order.refunded_amount_usd; + let refundable_before = order.refundable_amount_usd; + let refunded_after = refunded_before - amount_usd; + let refundable_after = refundable_before + amount_usd; + if !payment_order_refund_amounts_are_consistent( + order.amount_usd, + refunded_before, + refundable_before, + ) || !refunded_before.is_finite() + || refunded_before < amount_usd + || !refunded_after.is_finite() + || refunded_after < 0.0 + || refundable_after < 0.0 + || refundable_after > order.amount_usd + { + return Ok(AdminWalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + let mut order = order; + order.refunded_amount_usd = refunded_after; + order.refundable_amount_usd = refundable_after; + updated_order = Some(order); + } let mut updated_wallet = wallet.clone(); updated_wallet.balance = after_recharge; - updated_wallet.total_refunded = (updated_wallet.total_refunded - amount_usd).max(0.0); + updated_wallet.total_refunded = after_total_refunded; updated_wallet.updated_at_unix_secs = now_unix_secs; let transaction = AdminWalletTransactionRecord { @@ -310,7 +479,7 @@ impl AppState { reason_code: "refund_revert".to_string(), amount: amount_usd, balance_before: before_total, - balance_after: after_recharge + before_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -322,23 +491,13 @@ impl AppState { created_at_unix_ms: now_unix_secs, }; - if let Some(payment_order_id) = refund.payment_order_id.clone() { - let Some(order_store) = self.admin_wallet_payment_order_store.as_ref() else { - return Ok(AdminWalletMutationOutcome::Unavailable); - }; - let maybe_order = order_store + if let Some(updated_order) = updated_order { + self.admin_wallet_payment_order_store + .as_ref() + .expect("admin wallet payment order store should exist") .lock() .expect("admin wallet payment order store should lock") - .get(&payment_order_id) - .cloned(); - if let Some(mut order) = maybe_order { - order.refunded_amount_usd -= amount_usd; - order.refundable_amount_usd += amount_usd; - order_store - .lock() - .expect("admin wallet payment order store should lock") - .insert(order.id.clone(), order); - } + .insert(updated_order.id.clone(), updated_order); } let mut updated_refund = refund; @@ -354,6 +513,12 @@ impl AppState { .lock() .expect("admin wallet refund store should lock") .insert(updated_refund.id.clone(), updated_refund.clone()); + if let Some(transaction_store) = self.admin_wallet_transaction_store.as_ref() { + transaction_store + .lock() + .expect("admin wallet transaction store should lock") + .insert(transaction.id.clone(), transaction.clone()); + } self.invalidate_auth_context_cache(); return Ok(AdminWalletMutationOutcome::Applied(( @@ -446,3 +611,103 @@ fn stored_admin_wallet_transaction_to_gateway( created_at_unix_ms: transaction.created_at_unix_ms.unwrap_or_default(), } } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn refund_with_proof(proof: serde_json::Value) -> AdminWalletRefundRecord { + AdminWalletRefundRecord { + id: "refund-proof-lifecycle".to_string(), + refund_no: "rf-proof-lifecycle".to_string(), + wallet_id: "wallet-proof-lifecycle".to_string(), + user_id: Some("user-proof-lifecycle".to_string()), + payment_order_id: None, + source_type: "manual".to_string(), + source_id: None, + refund_mode: "original_channel".to_string(), + amount_usd: 4.0, + status: "processing".to_string(), + reason: Some("proof lifecycle regression".to_string()), + failure_reason: None, + gateway_refund_id: Some("gateway-proof-lifecycle".to_string()), + payout_method: Some("wxpay".to_string()), + payout_reference: None, + payout_proof: Some(proof), + requested_by: Some("user-proof-lifecycle".to_string()), + approved_by: Some("admin-proof-lifecycle".to_string()), + processed_by: Some("admin-proof-lifecycle".to_string()), + created_at_unix_ms: 1_710_000_000, + updated_at_unix_secs: 1_710_000_000, + processed_at_unix_secs: Some(1_710_000_010), + completed_at_unix_secs: None, + } + } + + #[tokio::test] + async fn completion_preserves_processing_proof_on_non_terminal_retry() { + let processing_proof = json!({ + "gateway": "wxpay", + "id": "gateway-proof-lifecycle", + "status": "processing" + }); + let state = AppState::new() + .expect("gateway state should build") + .with_admin_wallet_refunds_for_tests([refund_with_proof(processing_proof.clone())]); + + let outcome = state + .admin_complete_wallet_refund( + "wallet-proof-lifecycle", + "refund-proof-lifecycle", + Some("gateway-proof-lifecycle"), + None, + Some(json!({ + "gateway": "wxpay", + "id": "gateway-proof-lifecycle", + "status": "pending", + "attempt": 2 + })), + ) + .await + .expect("completion should resolve"); + let AdminWalletMutationOutcome::Applied(refund) = outcome else { + panic!("completion should apply"); + }; + assert_eq!(refund.status, "succeeded"); + assert_eq!(refund.payout_proof, Some(processing_proof)); + } + + #[tokio::test] + async fn completion_allows_terminal_success_proof_to_upgrade_processing_evidence() { + let state = AppState::new() + .expect("gateway state should build") + .with_admin_wallet_refunds_for_tests([refund_with_proof(json!({ + "gateway": "wxpay", + "id": "gateway-proof-lifecycle", + "status": "processing" + }))]); + let success_proof = json!({ + "gateway": "wxpay", + "id": "gateway-proof-lifecycle", + "status": "succeeded", + "processed_at": "2026-08-29T00:00:00Z" + }); + + let outcome = state + .admin_complete_wallet_refund( + "wallet-proof-lifecycle", + "refund-proof-lifecycle", + Some("gateway-proof-lifecycle"), + None, + Some(success_proof.clone()), + ) + .await + .expect("completion should resolve"); + let AdminWalletMutationOutcome::Applied(refund) = outcome else { + panic!("completion should apply"); + }; + assert_eq!(refund.status, "succeeded"); + assert_eq!(refund.payout_proof, Some(success_proof)); + } +} diff --git a/apps/aether-gateway/src/state/testing.rs b/apps/aether-gateway/src/state/testing.rs index 504ee8c6f..bdbf662e9 100644 --- a/apps/aether-gateway/src/state/testing.rs +++ b/apps/aether-gateway/src/state/testing.rs @@ -9,14 +9,64 @@ use aether_data_contracts::repository::usage::{UsageReadRepository, UsageReposit use aether_data_contracts::repository::video_tasks::{ VideoTaskReadRepository, VideoTaskRepository, }; +use hmac::Mac; use serde_json::json; +use sha2::{Digest, Sha256}; use super::{AppState, FrontdoorRuntimeGuardConfig, GatewayDataState}; use crate::{provider_transport, usage}; +fn auth_email_storage_key_digest_for_tests(domain: &str, parts: &[&str]) -> String { + let secret = std::env::var("JWT_SECRET_KEY") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string()); + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .expect("HMAC should accept the test auth key"); + mac.update(b"aether-auth-email-storage-v1\0"); + mac.update(domain.as_bytes()); + for part in parts { + mac.update(b"\0"); + mac.update(part.as_bytes()); + } + mac.finalize() + .into_bytes() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect() +} + #[cfg(test)] impl AppState { + pub(crate) fn with_internal_gateway_auth_secret_for_tests(mut self, secret: &str) -> Self { + self.internal_gateway_auth = Arc::new( + crate::internal_gateway_auth::InternalGatewayAuthConfig::with_secret_for_tests(secret), + ); + self + } + + pub(crate) fn without_internal_gateway_for_tests(mut self) -> Self { + self.internal_gateway_auth = + Arc::new(crate::internal_gateway_auth::InternalGatewayAuthConfig::disabled_for_tests()); + self + } + pub(crate) fn with_data_state_for_tests(mut self, data_state: GatewayDataState) -> Self { + // Request-execution fixtures provide candidate and provider data but + // bypass the production startup bootstrap that creates the enabled + // system-default routing group. Keep those isolated states aligned + // with the real gateway contract while leaving intentionally disabled + // or routing-only fixtures untouched. + let data_state = if data_state.has_minimal_candidate_selection_reader() + && data_state.has_provider_catalog_reader() + && !data_state.has_routing_group_reader() + && !data_state.has_routing_group_writer() + { + data_state.with_system_default_routing_group_for_tests() + } else { + data_state + }; self.replace_data_state(Arc::new(data_state)); self.request_candidate_queue = None; self @@ -57,6 +107,38 @@ impl AppState { self } + pub(crate) fn with_tunnel_identity_and_relay_secret_for_tests( + mut self, + instance_id: &str, + relay_base_url: Option<&str>, + relay_auth_secret: &str, + ) -> Self { + self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_directory_for_tests( + Arc::clone(&self.data), + crate::tunnel::TunnelAttachmentDirectory::for_tests(instance_id, relay_base_url, 90), + relay_auth_secret, + ); + self + } + + pub(crate) fn with_tunnel_identity_runtime_state_and_relay_secret_for_tests( + mut self, + instance_id: &str, + relay_base_url: Option<&str>, + runtime_state: Arc, + relay_auth_secret: &str, + ) -> Self { + self.tunnel = + crate::tunnel::EmbeddedTunnelState::with_data_identity_runtime_state_and_relay_secret_for_tests( + Arc::clone(&self.data), + instance_id, + relay_base_url, + runtime_state, + relay_auth_secret, + ); + self + } + pub(crate) fn with_video_task_data_reader_for_tests( mut self, repository: Arc, @@ -225,16 +307,22 @@ impl AppState { nonce: &str, payload: serde_json::Value, ) -> Self { + let key = aether_data::repository::provider_oauth::provider_oauth_state_storage_key(nonce); + let plaintext = payload.to_string(); + let purpose = format!("provider-oauth-state:{key}"); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext) + .expect("test provider OAuth state should seal"); let store = self .provider_oauth_state_store .get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new()))); store .lock() .expect("provider oauth state store should lock") - .insert(format!("provider_oauth_state:{nonce}"), payload.to_string()); + .insert(key.clone(), plaintext); self.runtime_state.kv_set_local_nowait( - &format!("provider_oauth_state:{nonce}"), - payload.to_string(), + &key, + sealed, Some(Duration::from_secs( aether_data::repository::provider_oauth::PROVIDER_OAUTH_STATE_TTL_SECS, )), @@ -245,23 +333,43 @@ impl AppState { pub(crate) fn with_provider_oauth_device_session_entry_for_tests( mut self, session_id: &str, - payload: serde_json::Value, + mut payload: serde_json::Value, ) -> Self { + if let Some(payload) = payload.as_object_mut() { + payload + .entry("session_id".to_string()) + .or_insert_with(|| json!(session_id)); + payload + .entry("initiated_by_user_id".to_string()) + .or_insert_with(|| json!("admin-user-123")); + payload + .entry("initiated_by_session_id".to_string()) + .or_insert_with(|| json!("session-123")); + payload + .entry("initiated_by_management_token_id".to_string()) + .or_insert_with(|| json!("management-token-123")); + } + let key = + aether_data::repository::provider_oauth::provider_oauth_device_session_storage_key( + session_id, + ); + let plaintext = payload.to_string(); + let purpose = + aether_data::repository::provider_oauth::provider_oauth_device_session_secret_purpose( + session_id, + ); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext) + .expect("test provider OAuth device session should seal"); let store = self .provider_oauth_device_session_store .get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new()))); store .lock() .expect("provider oauth device session store should lock") - .insert( - format!("device_auth_session:{session_id}"), - payload.to_string(), - ); - self.runtime_state.kv_set_local_nowait( - &format!("device_auth_session:{session_id}"), - payload.to_string(), - Some(Duration::from_secs(3600)), - ); + .insert(key.clone(), plaintext); + self.runtime_state + .kv_set_local_nowait(&key, sealed, Some(Duration::from_secs(3600))); self } @@ -270,19 +378,26 @@ impl AppState { task_id: &str, payload: serde_json::Value, ) -> Self { + let key = + aether_data::repository::provider_oauth::provider_oauth_batch_task_storage_key(task_id); + let plaintext = payload.to_string(); + let purpose = + aether_data::repository::provider_oauth::provider_oauth_batch_task_secret_purpose( + task_id, + ); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext) + .expect("test provider OAuth batch task should seal"); let store = self .provider_oauth_batch_task_store .get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new()))); store .lock() .expect("provider oauth batch task store should lock") - .insert( - format!("provider_oauth_batch_task:{task_id}"), - payload.to_string(), - ); + .insert(key.clone(), plaintext); self.runtime_state.kv_set_local_nowait( - &format!("provider_oauth_batch_task:{task_id}"), - payload.to_string(), + &key, + sealed, Some(Duration::from_secs( aether_data::repository::provider_oauth::PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, )), @@ -332,6 +447,11 @@ impl AppState { self } + pub(crate) fn without_auth_session_store_for_tests(mut self) -> Self { + self.auth_session_store = None; + self + } + pub(crate) fn without_auth_user_model_capability_store_for_tests(mut self) -> Self { self.auth_user_model_capability_store = None; self @@ -587,47 +707,67 @@ impl AppState { mut self, email: &str, code: &str, + verification_token: &str, created_at: chrono::DateTime, ) -> Self { + let code_hash = format!( + "{:x}", + Sha256::digest( + format!( + "aether-email-verification\0{}\0{}", + verification_token.trim(), + code.trim() + ) + .as_bytes() + ) + ); + let verification_token_hash = + format!("{:x}", Sha256::digest(verification_token.trim().as_bytes())); + let normalized_email = email.trim().to_ascii_lowercase(); + let key = format!( + "email:verification:{}", + auth_email_storage_key_digest_for_tests("pending", &[normalized_email.as_str()]) + ); + let value = json!({ + "code_hash": code_hash, + "created_at": created_at.to_rfc3339(), + "verification_token_hash": verification_token_hash, + }) + .to_string(); let store = self .auth_email_verification_store .get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new()))); store .lock() .expect("auth email verification store should lock") - .insert( - format!("email:verification:{}", email.trim().to_ascii_lowercase()), - json!({ - "code": code, - "created_at": created_at.to_rfc3339(), - }) - .to_string(), - ); - self.runtime_state.kv_set_local_nowait( - &format!("email:verification:{}", email.trim().to_ascii_lowercase()), - json!({ - "code": code, - "created_at": created_at.to_rfc3339(), - }) - .to_string(), - Some(Duration::from_secs(600)), - ); + .insert(key.clone(), value.clone()); + self.runtime_state + .kv_set_local_nowait(&key, value, Some(Duration::from_secs(600))); self } - pub(crate) fn with_auth_email_verified_for_tests(mut self, email: &str) -> Self { + pub(crate) fn with_auth_email_verified_for_tests( + mut self, + email: &str, + verification_token: &str, + ) -> Self { + let normalized_email = email.trim().to_ascii_lowercase(); + let key = format!( + "email:verified:{}", + auth_email_storage_key_digest_for_tests( + "registration-proof", + &[normalized_email.as_str(), verification_token.trim()] + ) + ); let store = self .auth_email_verification_store .get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new()))); store .lock() .expect("auth email verification store should lock") - .insert( - format!("email:verified:{}", email.trim().to_ascii_lowercase()), - "verified".to_string(), - ); + .insert(key.clone(), "verified".to_string()); self.runtime_state.kv_set_local_nowait( - &format!("email:verified:{}", email.trim().to_ascii_lowercase()), + &key, "verified".to_string(), Some(Duration::from_secs(3600)), ); diff --git a/apps/aether-gateway/src/state/types.rs b/apps/aether-gateway/src/state/types.rs index 1106ef8cd..4c1f3a318 100644 --- a/apps/aether-gateway/src/state/types.rs +++ b/apps/aether-gateway/src/state/types.rs @@ -61,7 +61,7 @@ pub(crate) enum AdminWalletMutationOutcome { Unavailable, } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub(crate) struct GatewayUserSessionView { pub(crate) id: String, pub(crate) user_id: String, @@ -78,6 +78,7 @@ pub(crate) struct GatewayUserSessionView { pub(crate) user_agent: Option, pub(crate) created_at: Option>, pub(crate) updated_at: Option>, + pub(crate) security_version: i64, } impl GatewayUserSessionView { @@ -131,9 +132,18 @@ impl GatewayUserSessionView { user_agent, created_at, updated_at, + security_version: 0, }) } + pub(crate) fn with_security_version(mut self, security_version: i64) -> Result { + if security_version < 0 { + return Err("user_sessions.security_version is negative".to_string()); + } + self.security_version = security_version; + Ok(self) + } + pub(crate) fn hash_refresh_token(token: &str) -> String { use sha2::Digest; @@ -157,8 +167,10 @@ impl GatewayUserSessionView { let Some(rotated_at) = self.rotated_at else { return (false, false); }; + let age = now.signed_duration_since(rotated_at); if prev_hash == &token_hash - && now.signed_duration_since(rotated_at).num_seconds() <= Self::REFRESH_GRACE_SECONDS + && age >= chrono::Duration::zero() + && age <= chrono::Duration::seconds(Self::REFRESH_GRACE_SECONDS) { return (true, true); } @@ -183,6 +195,48 @@ impl GatewayUserSessionView { } } +#[cfg(test)] +mod gateway_user_session_view_tests { + use super::GatewayUserSessionView; + use chrono::{Duration, Utc}; + + fn session_with_rotation(rotated_at: chrono::DateTime) -> GatewayUserSessionView { + GatewayUserSessionView::new( + "session-1".to_string(), + "user-1".to_string(), + "device-1".to_string(), + None, + GatewayUserSessionView::hash_refresh_token("current-token"), + Some(GatewayUserSessionView::hash_refresh_token("previous-token")), + Some(rotated_at), + None, + None, + None, + None, + None, + None, + None, + None, + ) + .expect("session should build") + } + + #[test] + fn refresh_grace_rejects_future_rotation_timestamps() { + let now = Utc::now(); + assert_eq!( + session_with_rotation(now - Duration::seconds(1)) + .verify_refresh_token("previous-token", now), + (true, true) + ); + assert_eq!( + session_with_rotation(now + Duration::milliseconds(1)) + .verify_refresh_token("previous-token", now), + (false, false) + ); + } +} + impl From for GatewayUserSessionView { fn from(value: crate::data::state::StoredUserSessionRecord) -> Self { Self { @@ -201,6 +255,7 @@ impl From for GatewayUserSessionVie user_agent: value.user_agent, created_at: value.created_at, updated_at: value.updated_at, + security_version: value.security_version, } } } @@ -229,6 +284,7 @@ impl From for crate::data::state::StoredUserSessionRecor user_agent: value.user_agent, created_at: value.created_at, updated_at: value.updated_at, + security_version: value.security_version, } } } @@ -308,7 +364,7 @@ impl From for crate::data::state::StoredUserPreferenc } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub(crate) struct GatewayAdminPaymentCallbackView { pub(crate) id: String, pub(crate) payment_order_id: Option, @@ -325,6 +381,33 @@ pub(crate) struct GatewayAdminPaymentCallbackView { pub(crate) processed_at_unix_secs: Option, } +impl std::fmt::Debug for GatewayAdminPaymentCallbackView { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GatewayAdminPaymentCallbackView") + .field("id", &self.id) + .field("payment_order_id", &self.payment_order_id) + .field("payment_method", &self.payment_method) + .field("callback_key", &"[REDACTED]") + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field( + "payload_hash", + &self.payload_hash.as_ref().map(|_| "[REDACTED]"), + ) + .field("signature_valid", &self.signature_valid) + .field("status", &self.status) + .field("payload", &self.payload.as_ref().map(|_| "[REDACTED]")) + .field( + "error_message", + &self.error_message.as_ref().map(|_| "[REDACTED]"), + ) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .field("processed_at_unix_secs", &self.processed_at_unix_secs) + .finish() + } +} + impl From for GatewayAdminPaymentCallbackView { fn from(value: super::AdminPaymentCallbackRecord) -> Self { Self { diff --git a/apps/aether-gateway/src/state/video.rs b/apps/aether-gateway/src/state/video.rs index e5c83b2a0..c04271f3f 100644 --- a/apps/aether-gateway/src/state/video.rs +++ b/apps/aether-gateway/src/state/video.rs @@ -6,6 +6,13 @@ use aether_data_contracts::repository::video_tasks::{ VideoTaskQueryFilter, VideoTaskStatusCount, }; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum VideoTaskRouteAccess { + Allowed, + NotFound, + Denied, +} + impl AppState { pub(crate) async fn read_data_backed_video_task_response( &self, @@ -18,6 +25,18 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn read_data_backed_video_task_response_for_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + ) -> Result, GatewayError> { + self.data + .read_video_task_response_for_user(route_family, request_path, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn find_video_task_by_id( &self, task_id: &str, @@ -38,12 +57,72 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn find_video_task_by_id_for_user( + &self, + task_id: &str, + user_id: &str, + ) -> Result, GatewayError> { + self.data + .find_video_task_for_user(VideoTaskLookupKey::Id(task_id), user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + + pub(crate) async fn find_video_task_by_short_id_for_user( + &self, + short_id: &str, + user_id: &str, + ) -> Result, GatewayError> { + self.data + .find_video_task_for_user(VideoTaskLookupKey::ShortId(short_id), user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn upsert_video_task_snapshot( &self, snapshot: &video_tasks::LocalVideoTaskSnapshot, ) -> Result, GatewayError> { + let mut record = snapshot.to_upsert_record(); + // Reconstructed snapshots intentionally omit sensitive/request-only fields. Preserve the + // persisted row's immutable identity and request-shape scalars before writing lifecycle + // changes back, so the repository can continue enforcing immutable-field integrity. + let existing_by_id = self + .data + .find_video_task(VideoTaskLookupKey::Id(record.id.as_str())) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + let existing = if existing_by_id.is_some() { + existing_by_id + } else if let Some(short_id) = record.short_id.as_deref() { + self.data + .find_video_task(VideoTaskLookupKey::ShortId(short_id)) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + } else { + None + }; + if let Some(existing) = existing { + record.id = existing.id; + record.short_id = existing.short_id; + record.request_id = existing.request_id; + record.user_id = existing.user_id; + record.api_key_id = existing.api_key_id; + record.external_task_id = existing.external_task_id; + record.provider_id = existing.provider_id; + record.endpoint_id = existing.endpoint_id; + record.key_id = existing.key_id; + record.client_api_format = existing.client_api_format; + record.provider_api_format = existing.provider_api_format; + record.format_converted = existing.format_converted; + record.model = existing.model; + record.duration_seconds = existing.duration_seconds; + record.resolution = existing.resolution; + record.aspect_ratio = existing.aspect_ratio; + record.size = existing.size; + } self.data - .upsert_video_task(snapshot.to_upsert_record()) + .upsert_video_task(record) .await .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -77,6 +156,50 @@ impl AppState { Ok(true) } + pub(crate) async fn hydrate_video_task_for_route_for_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + ) -> Result { + let user_id = user_id.trim(); + if user_id.is_empty() { + return Ok(VideoTaskRouteAccess::Denied); + } + let Some(lookup) = + video_tasks::resolve_video_task_hydration_lookup_key(route_family, request_path) + else { + return Ok(VideoTaskRouteAccess::NotFound); + }; + + if let Some(task) = self + .data + .find_video_task_for_user(lookup, user_id) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))? + { + if !self.video_tasks.hydrate_from_stored_task(&task) { + if let Some(snapshot) = self.reconstruct_video_task_snapshot(&task).await? { + self.video_tasks.record_snapshot(snapshot); + } + } + return Ok(VideoTaskRouteAccess::Allowed); + } + + Ok( + match self + .video_tasks + .snapshot_for_route(route_family, request_path) + { + Some(snapshot) if snapshot.belongs_to_user(user_id) => { + VideoTaskRouteAccess::Allowed + } + Some(_) => VideoTaskRouteAccess::Denied, + None => VideoTaskRouteAccess::NotFound, + }, + ) + } + pub(crate) async fn reconstruct_video_task_snapshot( &self, task: &StoredVideoTask, diff --git a/apps/aether-gateway/src/task_runtime/mod.rs b/apps/aether-gateway/src/task_runtime/mod.rs index 33b0f7eca..e8d4c95e9 100644 --- a/apps/aether-gateway/src/task_runtime/mod.rs +++ b/apps/aether-gateway/src/task_runtime/mod.rs @@ -8,7 +8,7 @@ use aether_data_contracts::repository::background_tasks::{ use aether_runtime::task::spawn_named; use aether_task_runtime::{RetryPolicy, TaskDefinition, TaskKind}; pub(crate) use aether_task_runtime::{TaskSupervisor, TaskSupervisorMetrics}; -use serde_json::{json, Value}; +use serde_json::Value; use sha2::{Digest, Sha256}; use tokio::task::JoinHandle; use tracing::warn; @@ -56,7 +56,6 @@ const RETRY_ONCE: RetryPolicy = RetryPolicy { max_attempts: 1 }; const RETRY_THREE: RetryPolicy = RetryPolicy { max_attempts: 3 }; const BACKGROUND_TASK_RUN_ID_MAX_BYTES: usize = 64; const WORKER_BOOT_RUN_ID_HASH_HEX_BYTES: usize = 20; - fn build_worker_boot_run_id(task_key: &str) -> String { let full_run_id = format!("boot:{task_key}"); if full_run_id.len() <= BACKGROUND_TASK_RUN_ID_MAX_BYTES { @@ -104,10 +103,6 @@ fn build_worker_boot_run( } } -fn worker_boot_event_payload(gateway_instance_id: &str) -> Value { - json!({ "gateway_instance_id": gateway_instance_id }) -} - pub(crate) fn spawn_singleton_worker( app: AppState, task_key: &'static str, @@ -473,8 +468,11 @@ pub(crate) async fn upsert_run_with_logging( ) -> Option { match app.upsert_background_task_run(run).await { Ok(result) => result, - Err(error) => { - warn!(error = ?error, "failed to upsert background task run"); + Err(_) => { + warn!( + error_category = "task_run_persistence_failed", + "failed to upsert background task run" + ); None } } @@ -532,8 +530,12 @@ pub(crate) async fn append_event_with_logging( payload_json, created_at_unix_secs: now_unix_secs(), }; - if let Err(error) = app.upsert_background_task_event(event).await { - warn!(error = ?error, run_id = %run_id, "failed to upsert background task event"); + if app.upsert_background_task_event(event).await.is_err() { + warn!( + error_category = "task_event_persistence_failed", + run_id = %run_id, + "failed to upsert background task event" + ); } } @@ -545,7 +547,6 @@ pub(crate) fn spawn_record_worker_boot( ) -> JoinHandle<()> { spawn_named("task-runtime-record-worker-boot", async move { let now = now_unix_secs(); - let gateway_instance_id = app.tunnel.local_instance_id().to_string(); let run = build_worker_boot_run(task_key, kind, trigger, now); let run_id = run.id.clone(); if upsert_run_with_logging(&app, run).await.is_none() { @@ -556,7 +557,7 @@ pub(crate) fn spawn_record_worker_boot( &run_id, "worker_boot", "background worker supervisor started", - Some(worker_boot_event_payload(&gateway_instance_id)), + None, ) .await; }) @@ -753,10 +754,11 @@ pub(crate) async fn submit_provider_delete_task( ) .await; } - Err(err) => { + Err(_) => { warn!( - "gateway admin provider delete task failed for provider {}: {:?}", - provider_id, err + provider_id = %provider_id, + error_category = "provider_delete_failed", + "gateway admin provider delete task failed" ); app.put_provider_delete_task(crate::LocalProviderDeleteTaskState { task_id: run_id.clone(), @@ -767,7 +769,7 @@ pub(crate) async fn submit_provider_delete_task( deleted_keys: 0, total_endpoints: 0, deleted_endpoints: 0, - message: format!("provider delete failed: {err:?}"), + message: "provider delete failed".to_string(), }); let _ = update_run_status( &app, @@ -776,7 +778,7 @@ pub(crate) async fn submit_provider_delete_task( Some(100), Some("provider delete task failed".to_string()), None, - Some(format!("{err:?}")), + Some("provider_delete_failed".to_string()), None, Some(now_unix_secs()), ) @@ -786,7 +788,9 @@ pub(crate) async fn submit_provider_delete_task( &run_id, "failed", "provider delete task failed", - Some(serde_json::json!({ "error": format!("{err:?}") })), + Some(serde_json::json!({ + "error_code": "provider_delete_failed" + })), ) .await; } @@ -802,21 +806,18 @@ pub(crate) async fn submit_provider_delete_task( #[cfg(test)] mod worker_boot_run_id_tests { - use std::collections::BTreeSet; use std::sync::Arc; use super::{ build_worker_boot_run, build_worker_boot_run_id, spawn_record_worker_boot, - worker_boot_event_payload, BACKGROUND_TASK_RUN_ID_MAX_BYTES, + BACKGROUND_TASK_RUN_ID_MAX_BYTES, }; + use crate::{data::GatewayDataState, AppState}; use aether_data::repository::background_tasks::InMemoryBackgroundTaskRepository; use aether_data_contracts::repository::background_tasks::{ BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository, BackgroundTaskStatus, }; - use serde_json::json; - - use crate::{data::GatewayDataState, AppState}; #[test] fn worker_boot_run_id_is_keyed_only_by_task() { @@ -883,16 +884,8 @@ mod worker_boot_run_id_tests { assert_eq!(run.updated_at_unix_secs, 123); } - #[test] - fn worker_boot_event_keeps_the_observing_gateway_instance() { - assert_eq!( - worker_boot_event_payload("gateway-a"), - json!({ "gateway_instance_id": "gateway-a" }) - ); - } - #[tokio::test] - async fn worker_boot_registration_is_shared_but_events_keep_each_gateway() { + async fn worker_boot_registration_is_shared_without_instance_identifiers() { let repository = Arc::new(InMemoryBackgroundTaskRepository::default()); let state_for = |gateway_instance_id: &str| { AppState::new() @@ -943,20 +936,6 @@ mod worker_boot_run_id_tests { .expect("worker boot events should load"); assert_eq!(events.len(), 2); assert!(events.iter().all(|event| event.event_type == "worker_boot")); - let gateway_instances = events - .iter() - .filter_map(|event| { - event - .payload_json - .as_ref() - .and_then(|payload| payload.get("gateway_instance_id")) - .and_then(serde_json::Value::as_str) - .map(str::to_string) - }) - .collect::>(); - assert_eq!( - gateway_instances, - BTreeSet::from(["gateway-a".to_string(), "gateway-b".to_string()]) - ); + assert!(events.iter().all(|event| event.payload_json.is_none())); } } diff --git a/apps/aether-gateway/src/tests/ai_execute/control_execute.rs b/apps/aether-gateway/src/tests/ai_execute/control_execute.rs index 3f18c8363..5a32fd0bf 100644 --- a/apps/aether-gateway/src/tests/ai_execute/control_execute.rs +++ b/apps/aether-gateway/src/tests/ai_execute/control_execute.rs @@ -5,8 +5,36 @@ use super::{ EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, TRACE_ID_HEADER, }; -#[tokio::test] -async fn gateway_locally_denies_sync_ai_control_execute_when_opted_in_and_execution_runtime_missing( +fn run_async_test_on_large_stack(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("tokio runtime should build") + .block_on(future); + }) + .expect("large-stack control execute test thread should spawn"); + + if let Err(payload) = handle.join() { + std::panic::resume_unwind(payload); + } +} + +#[test] +fn gateway_locally_denies_sync_ai_control_execute_when_opted_in_and_execution_runtime_missing() { + run_async_test_on_large_stack( + "gateway_locally_denies_sync_ai_control_execute_when_opted_in_and_execution_runtime_missing", + gateway_locally_denies_sync_ai_control_execute_when_opted_in_and_execution_runtime_missing_impl(), + ); +} + +async fn gateway_locally_denies_sync_ai_control_execute_when_opted_in_and_execution_runtime_missing_impl( ) { let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits_clone = Arc::clone(&execute_hits); @@ -271,8 +299,16 @@ async fn gateway_locally_denies_stream_ai_control_execute_when_opted_in_and_exec upstream_handle.abort(); } -#[tokio::test] -async fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_sync_ai_routes( +#[test] +fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_sync_ai_routes( +) { + run_async_test_on_large_stack( + "gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_sync_ai_routes", + gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_sync_ai_routes_impl(), + ); +} + +async fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_sync_ai_routes_impl( ) { let plan_hits = Arc::new(Mutex::new(0usize)); let plan_hits_clone = Arc::clone(&plan_hits); @@ -404,8 +440,16 @@ async fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_exec upstream_handle.abort(); } -#[tokio::test] -async fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_stream_ai_routes( +#[test] +fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_stream_ai_routes( +) { + run_async_test_on_large_stack( + "gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_stream_ai_routes", + gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_stream_ai_routes_impl(), + ); +} + +async fn gateway_does_not_proxy_control_execute_over_http_when_opted_in_and_execution_runtime_misses_stream_ai_routes_impl( ) { let plan_hits = Arc::new(Mutex::new(0usize)); let plan_hits_clone = Arc::clone(&plan_hits); diff --git a/apps/aether-gateway/src/tests/ai_execute/fallback.rs b/apps/aether-gateway/src/tests/ai_execute/fallback.rs index d66ec63ae..5f7ec2f9f 100644 --- a/apps/aether-gateway/src/tests/ai_execute/fallback.rs +++ b/apps/aether-gateway/src/tests/ai_execute/fallback.rs @@ -1,8 +1,9 @@ use super::{ any, build_router, build_router_with_execution_runtime_override, json, start_server, Arc, Body, HeaderValue, Json, Mutex, Request, Response, Router, StatusCode, DEPENDENCY_REASON_HEADER, - EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, - LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER, + EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_AI_PUBLIC, + EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, + TRACE_ID_HEADER, }; #[tokio::test] @@ -486,6 +487,8 @@ async fn assert_ai_route_locally_denied_after_execution_runtime_miss_with_reques let gateway = build_router_with_execution_runtime_override(execution_runtime_url); let (gateway_url, gateway_handle) = start_server(gateway).await; + let is_gemini_files_local_read = + method == reqwest::Method::GET && route_family == "gemini" && route_kind == "files"; let mut request = reqwest::Client::new().request(method, format!("{gateway_url}{request_path}")); if let Some(request_body) = request_body { @@ -495,7 +498,46 @@ async fn assert_ai_route_locally_denied_after_execution_runtime_miss_with_reques } let response = request.send().await.expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!( + response.status(), + if is_gemini_files_local_read { + StatusCode::NOT_FOUND + } else { + StatusCode::SERVICE_UNAVAILABLE + } + ); + if is_gemini_files_local_read { + assert_eq!( + response + .headers() + .get(EXECUTION_PATH_HEADER) + .and_then(|value| value.to_str().ok()), + Some(EXECUTION_PATH_LOCAL_AI_PUBLIC) + ); + assert_eq!( + response + .headers() + .get(DEPENDENCY_REASON_HEADER) + .and_then(|value| value.to_str().ok()), + None + ); + let payload: serde_json::Value = response.json().await.expect("body should parse"); + assert_eq!(payload, json!({"detail": "File not found"})); + assert_eq!(*control_execute_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); + assert_eq!( + public_execution_path + .lock() + .expect("mutex should lock") + .as_deref(), + None + ); + + gateway_handle.abort(); + execution_runtime_handle.abort(); + upstream_handle.abort(); + return; + } assert_eq!( response .headers() 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 c2103c842..3bac90c5d 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs @@ -1877,7 +1877,8 @@ async fn gateway_executes_openai_chat_antigravity_cross_format_sync_via_local_fi Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .with_system_default_routing_group_for_tests(), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -1949,7 +1950,7 @@ async fn gateway_executes_openai_chat_antigravity_cross_format_sync_via_local_fi "Bearer imported-antigravity-chat-token" ); assert_eq!(seen_execution_runtime_request.x_client_name, "antigravity"); - assert_eq!(seen_execution_runtime_request.x_client_version, "1.2.3"); + assert_eq!(seen_execution_runtime_request.x_client_version, "4.3.0"); assert_eq!( seen_execution_runtime_request.x_vscode_sessionid, "sess-antigravity-chat-local-123" diff --git a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/compact.rs b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/compact.rs index 2232af5ca..b6d665143 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/compact.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/compact.rs @@ -6,6 +6,7 @@ use super::{ EXECUTION_PATH_HEADER, TRACE_ID_HEADER, }; use crate::data::GatewayDataState; +use aether_ai_formats::openai_responses_message_item_id; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, @@ -413,7 +414,7 @@ async fn gateway_executes_openai_responses_compact_openai_family_upstream_stream "output_text": "Hello Compact", "output": [{ "type": "message", - "id": "resp_compact_openai_family_123_msg", + "id": openai_responses_message_item_id("resp_compact_openai_family_123", 0), "role": "assistant", "status": "completed", "content": [{ 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 03471fe8b..0cedf3111 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 @@ -6,6 +6,7 @@ use super::{ EXECUTION_PATH_HEADER, TRACE_ID_HEADER, }; use crate::data::GatewayDataState; +use aether_ai_formats::openai_responses_message_item_id; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, @@ -412,7 +413,7 @@ async fn gateway_executes_openai_responses_cross_format_upstream_stream_via_loca "output_text": "Hello Gemini CLI", "output": [{ "type": "message", - "id": "upstream-cli-stream-123_msg", + "id": openai_responses_message_item_id("upstream-cli-stream-123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -881,7 +882,7 @@ async fn gateway_executes_openai_responses_cross_format_function_call_upstream_s "output": [ { "type": "message", - "id": "upstream-cli-tool-123_msg", + "id": openai_responses_message_item_id("upstream-cli-tool-123", 0), "role": "assistant", "status": "completed", "content": [{ @@ -1411,7 +1412,12 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "antigravity", + "test-antigravity-client-id", + "test-antigravity-client-secret", + ), ), ]); let gateway_state = @@ -1424,7 +1430,8 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .with_system_default_routing_group_for_tests(), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); let gateway = build_router_with_state(gateway_state); @@ -1465,7 +1472,7 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str "output_text": "Hello Antigravity", "output": [{ "type": "message", - "id": "resp-local-stream_msg", + "id": openai_responses_message_item_id("resp-local-stream", 0), "role": "assistant", "status": "completed", "content": [{ @@ -1500,12 +1507,12 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-antigravity-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-antigravity-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-antigravity-cli-local-123")); @@ -1537,7 +1544,7 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str ); assert_eq!( seen_remote_execution_runtime_request.x_client_version, - "1.2.3" + "4.3.0" ); assert_eq!( seen_remote_execution_runtime_request.x_vscode_sessionid, diff --git a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/direct.rs b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/direct.rs index 21db78caf..8091f57a3 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/direct.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/direct.rs @@ -6,6 +6,7 @@ use super::{ TRACE_ID_HEADER, }; use crate::data::GatewayDataState; +use aether_ai_formats::openai_responses_message_item_id; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, @@ -428,7 +429,7 @@ async fn gateway_executes_openai_responses_sync_upstream_stream_via_local_finali "output_text": "Hello", "output": [{ "type": "message", - "id": "resp_stream_001_msg", + "id": openai_responses_message_item_id("resp_stream_001", 0), "role": "assistant", "status": "completed", "content": [{ @@ -983,6 +984,11 @@ async fn gateway_executes_kiro_claude_cli_sync_upstream_stream_via_local_finaliz Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), DEVELOPMENT_ENCRYPTION_KEY, + ) + .attach_proxy_node_repository_for_tests( + crate::tests::ai_execute::ai_execute_proxy_node_repository([ + "proxy-node-kiro-cli-finalize-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 8fa7f5745..7c3575c2f 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 @@ -1991,7 +1991,12 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_ crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "antigravity", + "test-antigravity-client-id", + "test-antigravity-client-secret", + ), ), ]); let gateway_state = @@ -2004,7 +2009,8 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_ Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .with_system_default_routing_group_for_tests(), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); let gateway = build_router_with_state(gateway_state); @@ -2086,12 +2092,12 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_ assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-antigravity-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-antigravity-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-antigravity-cli-local-123")); @@ -2123,7 +2129,7 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_ ); assert_eq!( seen_remote_execution_runtime_request.x_client_version, - "1.2.3" + "4.3.0" ); assert_eq!( seen_remote_execution_runtime_request.x_vscode_sessionid, diff --git a/apps/aether-gateway/src/tests/ai_execute/lifecycle.rs b/apps/aether-gateway/src/tests/ai_execute/lifecycle.rs index c3cae80b1..6e26a8976 100644 --- a/apps/aether-gateway/src/tests/ai_execute/lifecycle.rs +++ b/apps/aether-gateway/src/tests/ai_execute/lifecycle.rs @@ -14,6 +14,9 @@ use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadReposi use aether_data_contracts::repository::candidate_selection::{ StoredMinimalCandidateSelectionRow, StoredProviderModelMapping, }; +use aether_data_contracts::repository::candidates::{ + RequestCandidateReadRepository, RequestCandidateStatus, +}; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; @@ -427,6 +430,112 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() { upstream_handle.abort(); } +#[test] +fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byte() { + run_lifecycle_test( + "gateway_settles_stream_attempt_when_client_disconnects_before_first_byte", + gateway_settles_stream_attempt_when_client_disconnects_before_first_byte_impl, + ); +} + +async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byte_impl() { + // The execution runtime accepts the plan and then goes quiet, so the attempt + // is parked between its `pending` rows and the first upstream byte. + let execution_runtime = Router::new().route( + "/v1/execute/stream", + any(|_request: Request| async move { + tokio::time::sleep(std::time::Duration::from_secs(30)).await; + StatusCode::OK + }), + ); + + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("sk-client-openai-stream-precommit-disconnect")), + sample_local_openai_auth_snapshot( + "api-key-openai-lifecycle-local-1", + "user-openai-lifecycle-local-1", + ), + )])); + let candidate_selection_repository = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + sample_local_openai_candidate_row(), + ])); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_local_openai_provider()], + vec![sample_local_openai_endpoint()], + vec![sample_local_openai_key()], + )); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let gateway = build_router_with_state( + build_state_with_execution_runtime_override(execution_runtime_url) + .with_data_state_for_tests( + 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, + ), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let request = reqwest::Client::new() + .post(format!("{gateway_url}/v1/chat/completions")) + .header(http::header::CONTENT_TYPE, "application/json") + .header( + http::header::AUTHORIZATION, + "Bearer sk-client-openai-stream-precommit-disconnect", + ) + .header( + TRACE_ID_HEADER, + "trace-openai-chat-stream-precommit-disconnect-123", + ) + .body("{\"model\":\"gpt-5\",\"messages\":[],\"stream\":true}") + .send(); + + // Drop the in-flight request the way a downstream client does when its own + // first-byte timeout fires, before any response header exists. + assert!( + tokio::time::timeout(std::time::Duration::from_millis(750), request) + .await + .is_err(), + "the execution runtime should not have answered before the client gave up" + ); + + let mut stored_candidates = Vec::new(); + for _ in 0..200 { + stored_candidates = request_candidate_repository + .list_by_request_id("trace-openai-chat-stream-precommit-disconnect-123") + .await + .expect("request candidate trace should read"); + if stored_candidates + .iter() + .any(|candidate| candidate.status == RequestCandidateStatus::Cancelled) + { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + } + + let cancelled = stored_candidates + .iter() + .find(|candidate| candidate.status == RequestCandidateStatus::Cancelled) + .unwrap_or_else(|| { + panic!("dropped stream attempt should settle as cancelled: {stored_candidates:?}") + }); + assert_eq!(cancelled.status_code, Some(499)); + assert_eq!( + cancelled.error_type.as_deref(), + Some("local_stream_attempt_cancelled") + ); + assert!(cancelled.finished_at_unix_ms.is_some()); + + gateway_handle.abort(); + execution_runtime_handle.abort(); +} + #[test] fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error() { run_lifecycle_test( diff --git a/apps/aether-gateway/src/tests/ai_execute/mod.rs b/apps/aether-gateway/src/tests/ai_execute/mod.rs index 05cba7124..f48939ca4 100644 --- a/apps/aether-gateway/src/tests/ai_execute/mod.rs +++ b/apps/aether-gateway/src/tests/ai_execute/mod.rs @@ -1,6 +1,8 @@ use std::convert::Infallible; use std::sync::{Arc, Mutex}; +use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; +use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider; use axum::body::{to_bytes, Body, Bytes}; use axum::response::Response; use axum::routing::any; @@ -12,8 +14,9 @@ use serde_json::json; use crate::constants::{ CONTROL_EXECUTED_HEADER, CONTROL_EXECUTE_FALLBACK_HEADER, DEPENDENCY_REASON_HEADER, EXECUTION_PATH_EXECUTION_RUNTIME_STREAM, EXECUTION_PATH_EXECUTION_RUNTIME_SYNC, - EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, - LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER, + EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_AI_PUBLIC, + EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, + TRACE_ID_HEADER, }; use super::{ @@ -24,6 +27,85 @@ use super::{ VideoTaskTruthSourceMode, }; +/// Build the real proxy-node repository used by execution-runtime fixtures. +/// +/// Production proxy resolution deliberately fails closed when a configured +/// node id is missing or stale. These tests exercise the resolved-node path, +/// so each fixture must seed an online manual node just as deployment state +/// would. Manual nodes do not need a password when the configured URL has no +/// credentials. +pub(super) fn ai_execute_proxy_node_repository( + node_ids: I, +) -> Arc +where + I: IntoIterator, + S: AsRef, +{ + let nodes = node_ids.into_iter().map(|node_id| { + let node_id = node_id.as_ref(); + StoredProxyNode::new( + node_id.to_string(), + format!("ai-execute-{node_id}"), + "127.0.0.1".to_string(), + 1, + true, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + false, + false, + 1, + ) + .expect("ai_execute proxy node should build") + .with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None) + .with_tunnel_generation(format!("ai-execute-generation-{node_id}")) + }); + Arc::new(InMemoryProxyNodeRepository::seed(nodes)) +} + +/// Add a test-only non-retryable status rule to a provider fixture. +/// +/// Execution-runtime error fixtures represent a single upstream response. A +/// `429` is normally retryable by the production policy, so without an +/// explicit stop rule the candidate loop would consume the fixture and return +/// a synthetic 503 instead of the upstream error the test is exercising. +pub(super) fn ai_execute_provider_stop_on_status_code( + mut provider: StoredProviderCatalogProvider, + status_code: u16, +) -> StoredProviderCatalogProvider { + let mut config = provider + .config + .take() + .unwrap_or_else(|| serde_json::json!({})); + let object = config + .as_object_mut() + .expect("provider test config should be a JSON object"); + let rules = object + .entry("failover_rules".to_string()) + .or_insert_with(|| serde_json::json!({})); + let rules_object = rules + .as_object_mut() + .expect("provider failover_rules test config should be a JSON object"); + let statuses = rules_object + .entry("stop_on_status_codes".to_string()) + .or_insert_with(|| serde_json::json!([])); + let status_values = statuses + .as_array_mut() + .expect("provider stop_on_status_codes test config should be an array"); + if !status_values + .iter() + .any(|value| value.as_u64() == Some(u64::from(status_code))) + { + status_values.push(serde_json::json!(status_code)); + } + provider.config = Some(config); + provider +} + mod control_execute; mod fallback; mod finalize_local; diff --git a/apps/aether-gateway/src/tests/ai_execute/stream/decision.rs b/apps/aether-gateway/src/tests/ai_execute/stream/decision.rs index 4759ab296..ceec7e1bd 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream/decision.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream/decision.rs @@ -1801,7 +1801,12 @@ async fn gateway_executes_openai_chat_stream_with_custom_path_via_local_decision .with_system_config_values_for_tests(vec![( "provider_priority_mode".to_string(), json!("global_key"), - )]), + )]) + .attach_proxy_node_repository_for_tests( + crate::tests::ai_execute::ai_execute_proxy_node_repository([ + "proxy-node-openai-custom-stream", + ]), + ), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -2220,7 +2225,9 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable .to_string(), }); - let frames = if attempt == 1 { + // The primary key gets two attempts under the default + // sticky_key_attempts; both must fail to reach the backup. + let frames = if attempt <= 2 { concat!( "{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":429,\"headers\":{\"content-type\":\"application/json\"}}}\n", "{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"{\\\"error\\\":{\\\"message\\\":\\\"rate limited\\\",\\\"type\\\":\\\"rate_limit_error\\\"}}\"}}\n", @@ -2362,14 +2369,25 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable .lock() .expect("mutex should lock") .len() - >= 2 + >= 3 }) .await; let seen_execution_runtime_requests = seen_execution_runtime .lock() .expect("mutex should lock") .clone(); - assert_eq!(seen_execution_runtime_requests.len(), 2); + // Default sticky_key_attempts is 2: the primary key is retried once on + // the same key, then failover moves to the backup with a single attempt. + assert_eq!(seen_execution_runtime_requests.len(), 3); + assert_eq!( + seen_execution_runtime_requests + .iter() + .filter(|request| { + request.url == "https://api.openai.primary.example/chat/completions" + }) + .count(), + 2 + ); let primary_request = seen_execution_runtime_requests .iter() .find(|request| request.url == "https://api.openai.primary.example/chat/completions") @@ -2404,7 +2422,15 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable .list_by_request_id("trace-openai-chat-local-stream-failover-123") .await .expect("request candidate trace should read"); - assert_eq!(stored_candidates.len(), 2); + assert_eq!(stored_candidates.len(), 3); + assert_eq!( + stored_candidates + .iter() + .filter(|candidate| candidate.status == RequestCandidateStatus::Failed) + .count(), + 2, + "both sticky-key attempts on the primary should be recorded as failed" + ); let failed_candidate = stored_candidates .iter() .find(|candidate| { @@ -2434,24 +2460,15 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable failed_candidate.error_type.as_deref(), Some("retryable_upstream_status") ); - assert_eq!( - failed_candidate.error_message.as_deref(), - Some("execution runtime stream returned retryable status 429") - ); + assert!(failed_candidate.error_message.is_none()); let failed_upstream_response = failed_candidate .extra_data .as_ref() .and_then(|value| value.get("upstream_response")) .expect("failed stream candidate should keep its upstream response"); assert_eq!(failed_upstream_response["status_code"], json!(429)); - assert_eq!( - failed_upstream_response["body"]["error"]["message"], - json!("rate limited") - ); - assert_eq!( - failed_upstream_response["body"]["error"]["type"], - json!("rate_limit_error") - ); + assert!(failed_upstream_response.get("headers").is_none()); + assert!(failed_upstream_response.get("body").is_none()); assert_eq!(success_candidate.status, RequestCandidateStatus::Success); assert_eq!(success_candidate.status_code, Some(200)); assert!(success_candidate.started_at_unix_ms.is_some()); @@ -2465,7 +2482,7 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable assert_eq!( execution_runtime_hits.load(std::sync::atomic::Ordering::SeqCst), - 2 + 3 ); assert_eq!(*decision_hits.lock().expect("mutex should lock"), 0); assert_eq!(*plan_hits.lock().expect("mutex should lock"), 0); 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 8ede08204..b090f56f0 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"], - "codex_cli_rs/0.144.1" + "codex_cli_rs/0.153.3" ); assert_eq!( seen_execution_runtime_request.headers["originator"], diff --git a/apps/aether-gateway/src/tests/ai_execute/stream_cli/compact.rs b/apps/aether-gateway/src/tests/ai_execute/stream_cli/compact.rs index 1e1580c11..fb0081cbe 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_cli/compact.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_cli/compact.rs @@ -499,6 +499,11 @@ async fn gateway_executes_openai_responses_compact_as_unary_request_impl() { 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-compact-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 8143f5af9..3f5963bc1 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs @@ -508,6 +508,11 @@ async fn gateway_executes_kiro_claude_cli_stream_via_local_provider_catalog_cand 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-kiro-cli-local-stream", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -976,6 +981,11 @@ async fn gateway_executes_claude_cli_stream_via_local_decision_gate_without_wait 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-claude-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -1488,6 +1498,11 @@ async fn gateway_executes_claude_code_cli_stream_via_local_decision_gate_with_lo 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-claude-code-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -1979,6 +1994,11 @@ async fn gateway_executes_claude_chat_stream_via_local_decision_gate_with_local_ 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-claude-chat-stream", + ]), ), ); let gateway = build_router_with_state(gateway_state); diff --git a/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_chat.rs b/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_chat.rs index 3b2a21f91..202ca08b9 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_chat.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_chat.rs @@ -393,6 +393,11 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_ 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-gemini-chat-stream", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 7b7d98ad1..d34092b35 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 @@ -368,6 +368,11 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s 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-gemini-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -884,7 +889,12 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_ crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("gemini_cli", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("gemini_cli", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "gemini_cli", + "test-gemini-client-id", + "test-gemini-client-secret", + ), ), ]); let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url.clone()) @@ -895,6 +905,11 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_ 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-gemini-cli-oauth-local", + ]), ), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); @@ -934,12 +949,12 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_ assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-gemini-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-gemini-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-gemini-cli-stream-local-123")); @@ -1904,7 +1919,12 @@ async fn gateway_executes_antigravity_gemini_cli_stream_via_local_decision_gate_ crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "antigravity", + "test-antigravity-client-id", + "test-antigravity-client-secret", + ), ), ]); let data_state = @@ -1915,6 +1935,7 @@ async fn gateway_executes_antigravity_gemini_cli_stream_via_local_decision_gate_ Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, ) + .with_system_default_routing_group_for_tests() .with_system_config_values_for_tests([( crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY.to_string(), json!({ @@ -1988,12 +2009,12 @@ async fn gateway_executes_antigravity_gemini_cli_stream_via_local_decision_gate_ assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-antigravity-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-antigravity-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-antigravity-cli-stream-local-123")); @@ -2017,7 +2038,7 @@ async fn gateway_executes_antigravity_gemini_cli_stream_via_local_decision_gate_ "Bearer refreshed-antigravity-cli-stream-access-token" ); assert_eq!(seen_execution_runtime_request.x_client_name, "antigravity"); - assert_eq!(seen_execution_runtime_request.x_client_version, "1.2.3"); + assert_eq!(seen_execution_runtime_request.x_client_version, "4.3.0"); assert_eq!( seen_execution_runtime_request.x_vscode_sessionid, "sess-antigravity-stream-local-123" diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/chat/failover.rs b/apps/aether-gateway/src/tests/ai_execute/sync/chat/failover.rs index cc8de1774..61b73baaa 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/chat/failover.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/chat/failover.rs @@ -112,7 +112,7 @@ async fn gateway_skips_unsupported_local_openai_chat_sync_candidate_before_tryin false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -134,7 +134,7 @@ async fn gateway_skips_unsupported_local_openai_chat_sync_candidate_before_tryin "https://api.openai.skip.example".to_string(), None, None, - Some(2), + Some(1), None, None, None, @@ -520,7 +520,7 @@ async fn gateway_surfaces_local_execution_runtime_miss_reason_when_all_openai_ch false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -542,7 +542,7 @@ async fn gateway_surfaces_local_execution_runtime_miss_reason_when_all_openai_ch "https://chatgpt.com/backend-api/codex".to_string(), None, None, - Some(2), + Some(1), None, None, None, @@ -802,7 +802,7 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -828,7 +828,7 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur base_url.to_string(), None, None, - Some(2), + Some(1), None, None, None, @@ -970,7 +970,9 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur .to_string(), }); - if attempt == 1 { + // The primary key gets two attempts under the default + // sticky_key_attempts; both must fail to reach the backup. + if attempt <= 2 { return Json(json!({ "request_id": "trace-openai-chat-local-failover-123", "status_code": 401, @@ -1125,56 +1127,58 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur .lock() .expect("mutex should lock") .clone(); - assert_eq!(seen_execution_runtime_requests.len(), 2); + // Default sticky_key_attempts is 2: the primary key is retried once on + // the same key, then failover moves to the backup with a single attempt. + assert_eq!(seen_execution_runtime_requests.len(), 3); + for primary_request in &seen_execution_runtime_requests[..2] { + assert_eq!( + primary_request.trace_id, + "trace-openai-chat-local-failover-123" + ); + assert_eq!( + primary_request.url, + "https://api.openai.primary.example/chat/completions" + ); + assert_eq!( + primary_request.authorization, + "Bearer sk-upstream-openai-primary" + ); + } assert_eq!( - seen_execution_runtime_requests[0].trace_id, - "trace-openai-chat-local-failover-123" - ); - assert_eq!( - seen_execution_runtime_requests[0].url, - "https://api.openai.primary.example/chat/completions" - ); - assert_eq!( - seen_execution_runtime_requests[0].authorization, - "Bearer sk-upstream-openai-primary" - ); - assert_eq!( - seen_execution_runtime_requests[1].url, + seen_execution_runtime_requests[2].url, "https://api.openai.backup.example/chat/completions" ); assert_eq!( - seen_execution_runtime_requests[1].model, + seen_execution_runtime_requests[2].model, "gpt-5-upstream-backup" ); assert_eq!( - seen_execution_runtime_requests[1].authorization, + seen_execution_runtime_requests[2].authorization, "Bearer sk-upstream-openai-backup" ); let stored_candidates = request_candidate_repository .list_by_request_id("trace-openai-chat-local-failover-123") .await .expect("request candidate trace should read"); - assert_eq!(stored_candidates.len(), 2); - assert_eq!(stored_candidates[0].candidate_index, 0); - assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Failed); - assert_eq!(stored_candidates[0].status_code, Some(401)); - assert_eq!( - stored_candidates[0].error_message.as_deref(), - Some("invalid auth token") - ); - let failed_upstream_response = stored_candidates[0] - .extra_data - .as_ref() - .and_then(|value| value.get("upstream_response")) - .expect("failed candidate should keep its upstream response"); - assert_eq!(failed_upstream_response["status_code"], json!(401)); - assert_eq!( - failed_upstream_response["body"]["error"]["message"], - json!("invalid auth token") - ); - assert_eq!(stored_candidates[1].candidate_index, 1); - assert_eq!(stored_candidates[1].status, RequestCandidateStatus::Success); - assert_eq!(stored_candidates[1].status_code, Some(200)); + assert_eq!(stored_candidates.len(), 3); + for (retry_index, failed_candidate) in stored_candidates[..2].iter().enumerate() { + assert_eq!(failed_candidate.candidate_index, 0); + assert_eq!(failed_candidate.retry_index, retry_index as u32); + assert_eq!(failed_candidate.status, RequestCandidateStatus::Failed); + assert_eq!(failed_candidate.status_code, Some(401)); + assert!(failed_candidate.error_message.is_none()); + let failed_upstream_response = failed_candidate + .extra_data + .as_ref() + .and_then(|value| value.get("upstream_response")) + .expect("failed candidate should keep its upstream response"); + assert_eq!(failed_upstream_response["status_code"], json!(401)); + assert!(failed_upstream_response.get("headers").is_none()); + assert!(failed_upstream_response.get("body").is_none()); + } + assert_eq!(stored_candidates[2].candidate_index, 1); + assert_eq!(stored_candidates[2].status, RequestCandidateStatus::Success); + assert_eq!(stored_candidates[2].status_code, Some(200)); tokio::time::sleep(std::time::Duration::from_millis(100)).await; assert!( @@ -1184,7 +1188,7 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur assert_eq!( *execution_runtime_hits.lock().expect("mutex should lock"), - 2 + 3 ); assert_eq!(*decision_hits.lock().expect("mutex should lock"), 0); assert_eq!(*plan_hits.lock().expect("mutex should lock"), 0); 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 2e04fae22..9d97ebd83 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 @@ -1734,6 +1734,9 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_cli_syn Some(serde_json::json!({ "claude_code_advanced": { "cli_only_enabled": false + }, + "failover_rules": { + "stop_on_status_codes": [429] } })), ) @@ -1907,7 +1910,7 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_cli_syn }); Json(json!({ "request_id": "trace-openai-chat-claude-cli-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -1940,7 +1943,12 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_cli_syn sample_candidate_row(), ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -2185,7 +2193,11 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_cli_syn None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -2362,7 +2374,7 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_cli_syn }); Json(json!({ "request_id": "trace-openai-chat-gemini-cli-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -2394,7 +2406,12 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_cli_syn sample_candidate_row(), ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -2615,7 +2632,11 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_sync_fa None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -2778,7 +2799,7 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_sync_fa }); Json(json!({ "request_id": "trace-openai-chat-claude-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -2811,7 +2832,12 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_sync_fa sample_candidate_row(), ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -3024,7 +3050,11 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_sync_fa None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -3220,7 +3250,7 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_sync_fa }); Json(json!({ "request_id": "trace-openai-chat-gemini-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -3252,7 +3282,12 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_gemini_sync_fa sample_candidate_row(), ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -3726,6 +3761,11 @@ async fn gateway_executes_openai_chat_sync_with_custom_path_via_local_decision_g provider_catalog_repository, Arc::new(InMemoryRequestCandidateRepository::default()), DEVELOPMENT_ENCRYPTION_KEY, + ) + .attach_proxy_node_repository_for_tests( + crate::tests::ai_execute::ai_execute_proxy_node_repository([ + "proxy-node-openai-custom-path", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 7c5687444..4665f0b51 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 @@ -473,6 +473,11 @@ async fn gateway_executes_claude_code_cli_sync_via_local_decision_gate_with_loca 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-claude-code-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/claude/kiro.rs b/apps/aether-gateway/src/tests/ai_execute/sync/claude/kiro.rs index 04a775038..f98c45bca 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/claude/kiro.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/claude/kiro.rs @@ -528,6 +528,11 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid Arc::clone(&request_candidate_repository), Arc::clone(&usage_repository), DEVELOPMENT_ENCRYPTION_KEY, + ) + .attach_proxy_node_repository_for_tests( + crate::tests::ai_execute::ai_execute_proxy_node_repository([ + "proxy-node-kiro-cli-local-sync", + ]), ), ) .with_usage_runtime_for_tests(UsageRuntimeConfig { @@ -1145,6 +1150,11 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid 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-kiro-cli-local-sync", + ]), ), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_chat.rs b/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_chat.rs index ad3952b00..fd3b4e18c 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_chat.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_chat.rs @@ -422,6 +422,11 @@ async fn gateway_executes_claude_chat_sync_via_local_decision_gate_with_local_sy 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-claude-chat-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -980,7 +985,7 @@ async fn gateway_returns_claude_chat_error_for_local_sync_failure_impl() { any(move |_request: Request| async move { Json(json!({ "request_id": "trace-claude-chat-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -1026,6 +1031,11 @@ async fn gateway_returns_claude_chat_error_for_local_sync_failure_impl() { 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-claude-chat-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_cli.rs b/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_cli.rs index fcc657b3e..f90fb5cc9 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_cli.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/claude/local_cli.rs @@ -396,6 +396,11 @@ async fn gateway_executes_claude_cli_sync_via_local_decision_gate_with_local_syn 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-claude-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -675,7 +680,7 @@ async fn gateway_returns_claude_cli_error_for_local_sync_failure_impl() { any(move |_request: Request| async move { Json(json!({ "request_id": "trace-claude-cli-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -723,6 +728,11 @@ async fn gateway_returns_claude_cli_error_for_local_sync_failure_impl() { 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-claude-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 b15c80c62..3fce6a6c7 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs @@ -455,6 +455,11 @@ async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_loc 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); @@ -543,6 +548,9 @@ async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_loc .expect("request candidate trace should read"); assert_eq!(stored_candidates.len(), 1); assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Success); + // Candidate persistence receives proxy metadata from the report context; + // this local fixture intentionally supplies only the plan-level node id. + // The runtime assertion above verifies that the plan proxy was resolved. assert_eq!( stored_candidates[0] .extra_data @@ -550,7 +558,7 @@ async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_loc .and_then(|value| value.get("proxy")) .and_then(|value| value.get("node_id")) .and_then(serde_json::Value::as_str), - Some("proxy-node-openai-cli-local") + None ); assert_eq!( stored_candidates[0] @@ -1540,7 +1548,7 @@ async fn gateway_returns_openai_responses_error_for_local_sync_failure_impl() { any(move |_request: Request| async move { Json(json!({ "request_id": "trace-openai-cli-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -1572,7 +1580,12 @@ async fn gateway_returns_openai_responses_error_for_local_sync_failure_impl() { sample_candidate_row(), ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -1587,6 +1600,11 @@ async fn gateway_returns_openai_responses_error_for_local_sync_failure_impl() { 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); @@ -1754,7 +1772,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cl None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -1914,7 +1936,7 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cl }); Json(json!({ "request_id": "trace-openai-cli-gemini-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -1962,6 +1984,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cl 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); @@ -2161,7 +2188,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sy None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -2307,7 +2338,7 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sy }); Json(json!({ "request_id": "trace-openai-cli-claude-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -2356,6 +2387,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sy 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); @@ -2547,7 +2583,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_ch None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -2693,7 +2733,7 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_ch }); Json(json!({ "request_id": "trace-openai-cli-claude-chat-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -2742,6 +2782,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_ch 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); @@ -2936,7 +2981,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_ch None, Some(20.0), None, - None, + Some(serde_json::json!({ + "failover_rules": { + "stop_on_status_codes": [429] + } + })), ) } @@ -3082,7 +3131,7 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_ch }); Json(json!({ "request_id": "trace-openai-cli-gemini-chat-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -3130,6 +3179,11 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_ch 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); 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 770bd5846..8abd3a094 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 @@ -397,6 +397,11 @@ async fn gateway_executes_gemini_cli_sync_via_local_decision_gate_with_local_syn 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-gemini-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -666,7 +671,7 @@ async fn gateway_returns_gemini_cli_error_for_local_sync_failure_impl() { any(move |_request: Request| async move { Json(json!({ "request_id": "trace-gemini-cli-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -698,7 +703,12 @@ async fn gateway_returns_gemini_cli_error_for_local_sync_failure_impl() { ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -713,6 +723,11 @@ async fn gateway_returns_gemini_cli_error_for_local_sync_failure_impl() { 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-gemini-cli-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -1208,7 +1223,12 @@ async fn gateway_executes_gemini_cli_sync_via_local_decision_gate_after_oauth_re crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("gemini_cli", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("gemini_cli", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "gemini_cli", + "test-gemini-client-id", + "test-gemini-client-secret", + ), ), ]); let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url.clone()) @@ -1219,6 +1239,11 @@ async fn gateway_executes_gemini_cli_sync_via_local_decision_gate_after_oauth_re 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-gemini-cli-oauth-local", + ]), ), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); @@ -1256,12 +1281,12 @@ async fn gateway_executes_gemini_cli_sync_via_local_decision_gate_after_oauth_re assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-gemini-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-gemini-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-gemini-cli-local-123")); @@ -2223,7 +2248,12 @@ async fn gateway_executes_antigravity_gemini_cli_sync_via_local_decision_gate_af crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ Arc::new( crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")), + .with_token_url_for_tests("antigravity", format!("{refresh_url}/oauth/token")) + .with_oauth_credentials_for_tests( + "antigravity", + "test-antigravity-client-id", + "test-antigravity-client-secret", + ), ), ]); let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url.clone()) @@ -2234,7 +2264,8 @@ async fn gateway_executes_antigravity_gemini_cli_sync_via_local_decision_gate_af provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .with_system_default_routing_group_for_tests(), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh); let gateway = build_router_with_state(gateway_state); @@ -2292,12 +2323,12 @@ async fn gateway_executes_antigravity_gemini_cli_sync_via_local_decision_gate_af assert!(seen_refresh_request .body .contains("grant_type=refresh_token")); - assert!(seen_refresh_request.body.contains( - "client_id=1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com" - )); assert!(seen_refresh_request .body - .contains("client_secret=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")); + .contains("client_id=test-antigravity-client-id")); + assert!(seen_refresh_request + .body + .contains("client_secret=test-antigravity-client-secret")); assert!(seen_refresh_request .body .contains("refresh_token=rt-antigravity-cli-local-123")); @@ -2321,7 +2352,7 @@ async fn gateway_executes_antigravity_gemini_cli_sync_via_local_decision_gate_af "Bearer refreshed-antigravity-cli-access-token" ); assert_eq!(seen_execution_runtime_request.x_client_name, "antigravity"); - assert_eq!(seen_execution_runtime_request.x_client_version, "1.2.3"); + assert_eq!(seen_execution_runtime_request.x_client_version, "4.3.0"); assert_eq!( seen_execution_runtime_request.x_vscode_sessionid, "sess-antigravity-local-123" diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/gemini/local_chat.rs b/apps/aether-gateway/src/tests/ai_execute/sync/gemini/local_chat.rs index 79d89436e..54530f7a8 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/gemini/local_chat.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/gemini/local_chat.rs @@ -398,6 +398,11 @@ async fn gateway_executes_gemini_chat_sync_via_local_decision_gate_with_local_sy 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-gemini-chat-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); @@ -693,7 +698,7 @@ async fn gateway_returns_gemini_chat_error_for_local_sync_failure_impl() { any(move |_request: Request| async move { Json(json!({ "request_id": "trace-gemini-chat-local-error-123", - "status_code": 200, + "status_code": 429, "headers": { "content-type": "application/json" }, @@ -722,7 +727,12 @@ async fn gateway_returns_gemini_chat_error_for_local_sync_failure_impl() { ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider_catalog_provider()], + vec![ + crate::tests::ai_execute::ai_execute_provider_stop_on_status_code( + sample_provider_catalog_provider(), + 429, + ), + ], vec![sample_provider_catalog_endpoint()], vec![sample_provider_catalog_key()], )); @@ -738,6 +748,11 @@ async fn gateway_returns_gemini_chat_error_for_local_sync_failure_impl() { 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-gemini-chat-local", + ]), ), ); let gateway = build_router_with_state(gateway_state); 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 9f7d917e1..43409318a 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"], - "codex_cli_rs/0.144.1" + "codex_cli_rs/0.153.3" ); assert_eq!( seen_execution_runtime_request.headers["originator"], diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/search.rs b/apps/aether-gateway/src/tests/ai_execute/sync/search.rs index 3e3cfd577..5b2ec2067 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/search.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/search.rs @@ -140,7 +140,7 @@ async fn gateway_executes_codex_search_with_responses_permission_and_search_cont false, false, None, - Some(2), + Some(1), None, Some(900.0), None, @@ -162,7 +162,7 @@ async fn gateway_executes_codex_search_with_responses_permission_and_search_cont "https://chatgpt.com/backend-api/codex".to_string(), None, None, - Some(2), + Some(1), None, None, None, @@ -570,9 +570,12 @@ async fn gateway_executes_codex_search_with_responses_permission_and_search_cont .filter(|plan| plan["request_id"] == "trace-search-failover-1") .map(|plan| plan["provider_id"].clone()) .collect::>(); + // Default sticky_key_attempts is 2: the first provider is retried once on + // the same key before failover advances to the second provider. assert_eq!( failover_plans, vec![ + json!("provider-codex-search-1"), json!("provider-codex-search-1"), json!("provider-codex-search-2") ] @@ -581,17 +584,16 @@ async fn gateway_executes_codex_search_with_responses_permission_and_search_cont .list_by_request_id("trace-search-failover-1") .await .expect("failover request candidates should read"); - assert_eq!(failover_candidates.len(), 2); + assert_eq!(failover_candidates.len(), 3); + for failed_candidate in &failover_candidates[..2] { + assert_eq!(failed_candidate.status, RequestCandidateStatus::Failed); + assert_eq!(failed_candidate.status_code, Some(500)); + } assert_eq!( - failover_candidates[0].status, - RequestCandidateStatus::Failed - ); - assert_eq!(failover_candidates[0].status_code, Some(500)); - assert_eq!( - failover_candidates[1].status, + failover_candidates[2].status, RequestCandidateStatus::Success ); - assert_eq!(failover_candidates[1].status_code, Some(200)); + assert_eq!(failover_candidates[2].status_code, Some(200)); gateway_handle.abort(); execution_runtime_handle.abort(); diff --git a/apps/aether-gateway/src/tests/architecture/admin_provider.rs b/apps/aether-gateway/src/tests/architecture/admin_provider.rs index b67ad922e..14039106c 100644 --- a/apps/aether-gateway/src/tests/architecture/admin_provider.rs +++ b/apps/aether-gateway/src/tests/architecture/admin_provider.rs @@ -1163,7 +1163,8 @@ fn admin_provider_write_uses_specific_local_owners() { "crate::handlers::admin::provider::write::normalize::{", "normalize_auth_type,", "validate_vertex_api_formats,", - "encrypt_catalog_secret_with_fallbacks, json_string_list,", + ".seal_provider_catalog_key_api_key(", + ".seal_provider_catalog_key_auth_config(", "normalize_json_object, normalize_string_list,", ] { assert!( @@ -1243,7 +1244,6 @@ fn admin_provider_ops_providers_mod_stays_thin() { "apps/aether-gateway/src/handlers/admin/provider/ops/providers/support.rs", ); for pattern in [ - "pub(super) const ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS:", "pub(super) const ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE:", "pub(super) const ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE:", "pub(super) const ADMIN_PROVIDER_OPS_VERIFY_RUST_ONLY_MESSAGE:", @@ -1257,6 +1257,30 @@ fn admin_provider_ops_providers_mod_stays_thin() { "handlers/admin/provider/ops/providers/support.rs should own {pattern}" ); } + assert!( + !providers_support.contains("ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS"), + "provider ops support should not duplicate the shared credential sensitivity policy" + ); + + let provider_ops_credentials = + read_workspace_file("apps/aether-gateway/src/handlers/shared/provider_ops_credential.rs"); + for pattern in [ + "pub(crate) const PROVIDER_OPS_PERSISTENT_SECRET_FIELDS:", + "pub(crate) const PROVIDER_OPS_TRANSIENT_SECRET_FIELDS:", + "pub(crate) fn provider_ops_credential_field_is_secret(", + ] { + assert!( + provider_ops_credentials.contains(pattern), + "handlers/shared/provider_ops_credential.rs should own provider credential sensitivity policy {pattern}" + ); + } + let providers_config = read_workspace_file( + "apps/aether-gateway/src/handlers/admin/provider/ops/providers/config.rs", + ); + assert!( + providers_config.contains("provider_ops_credential_field_is_secret"), + "provider ops config should consume the shared credential sensitivity policy" + ); for path in [ "apps/aether-gateway/src/maintenance/runtime.rs", @@ -1963,8 +1987,9 @@ fn admin_provider_oauth_quota_mod_stays_thin() { "handlers/admin/provider/oauth/quota/antigravity.rs should import common quota helpers from shared.rs" ); assert!( - quota_antigravity - .contains("use aether_provider_pool::build_antigravity_pool_quota_request;"), + quota_antigravity.contains("use aether_provider_pool::{") + && quota_antigravity.contains("build_antigravity_pool_quota_request") + && quota_antigravity.contains("build_antigravity_pool_quota_summary_request"), "handlers/admin/provider/oauth/quota/antigravity.rs should delegate antigravity quota request construction to aether-provider-pool" ); let quota_chatgpt_web = read_workspace_file( diff --git a/apps/aether-gateway/src/tests/architecture/admin_shared.rs b/apps/aether-gateway/src/tests/architecture/admin_shared.rs index 1794a6e0e..cf429d3b9 100644 --- a/apps/aether-gateway/src/tests/architecture/admin_shared.rs +++ b/apps/aether-gateway/src/tests/architecture/admin_shared.rs @@ -208,7 +208,6 @@ fn admin_wrapped_state_owns_provider_oauth_capabilities() { "pub(crate) async fn create_provider_oauth_catalog_key(", "pub(crate) async fn update_existing_provider_oauth_catalog_key(", "pub(crate) async fn refresh_provider_oauth_account_state_after_update(", - "pub(crate) async fn update_provider_catalog_key_oauth_credentials(", ] { assert!( admin_request.contains(pattern), @@ -841,12 +840,12 @@ fn admin_proxy_uses_single_admin_routes_entrypoint() { "pub(crate) fn has_usage_data_reader(&self) -> bool", "pub(crate) fn has_auth_module_writer(&self) -> bool", "pub(crate) async fn get_ldap_module_config(", - "pub(crate) async fn upsert_ldap_module_config(", + "pub(crate) async fn compare_and_swap_ldap_module_config(", "pub(crate) async fn count_active_local_admin_users_with_valid_password(", "pub(crate) async fn list_oauth_provider_configs(", "pub(crate) async fn get_oauth_provider_config(", "pub(crate) async fn upsert_oauth_provider_config(", - "pub(crate) async fn delete_oauth_provider_config(", + "pub(crate) async fn delete_oauth_provider_config_if_unlinked(", "pub(crate) async fn get_management_token_with_user(", "pub(crate) async fn delete_management_token(", "pub(crate) async fn remove_admin_security_blacklist(", diff --git a/apps/aether-gateway/src/tests/architecture/ai_serving.rs b/apps/aether-gateway/src/tests/architecture/ai_serving.rs index eed95eceb..31d2e5ffd 100644 --- a/apps/aether-gateway/src/tests/architecture/ai_serving.rs +++ b/apps/aether-gateway/src/tests/architecture/ai_serving.rs @@ -717,6 +717,9 @@ fn ai_serving_crate_api_is_confined_to_root_seams() { if relative == "apps/aether-gateway/src/ai_serving/pure/mod.rs" || relative == "apps/aether-gateway/src/ai_serving/transport.rs" || relative == "apps/aether-gateway/src/ai_serving/api.rs" + // This module is included by the finalize implementation solely as a + // test fixture. Keep the architecture rule focused on production seams. + || relative == "apps/aether-gateway/src/ai_serving/finalize/tests_sync.rs" || relative.ends_with("/tests.rs") || relative.contains("/tests/") || relative.starts_with("apps/aether-gateway/src/tests/") @@ -1714,7 +1717,6 @@ fn ai_serving_candidate_materialization_owns_affinity_and_candidate_runtime_pers "pub fn ai_should_persist_available_candidate_for_pool_key", "pub fn ai_should_persist_skipped_candidate_for_pool_membership", "pub fn ai_candidate_extra_data_with_ranking", - "attempt_slot_count", "should_persist_available_candidate", "persist_available_candidate", "build_attempt", @@ -2894,6 +2896,22 @@ fn ai_serving_standard_attempts_consume_eligible_local_candidates_without_transp ); } + let standard_family_request = read_workspace_file( + "apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs", + ); + for pattern in [ + "is_antigravity_provider_transport(", + "build_antigravity_v1internal_provider_request(", + "is_gemini_cli_provider_transport(", + "build_gemini_cli_v1internal_provider_request(", + ] { + assert!( + standard_family_request.contains(pattern), + "standard family request preparation should build v1internal envelopes through {pattern} \ + so a cross-format client never posts a bare Gemini body to a v1internal URL" + ); + } + let provider_transport_standard = read_workspace_file("crates/aether-provider/transport/src/standard/mod.rs"); for pattern in [ @@ -4227,7 +4245,7 @@ fn ai_serving_finalize_standard_sync_products_are_owned_by_format_crate() { "pub struct OpenAiImageSyncFinalizeProduct", "OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND", "CODEX_OPENAI_IMAGE_DEFAULT_OUTPUT_FORMAT", - "base64::engine::general_purpose::STANDARD.decode", + "decode_sync_report_body_base64", ] { assert!( surface_openai_image_stream.contains(expected), diff --git a/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs b/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs index 823e197c1..8772263c8 100644 --- a/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs +++ b/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs @@ -1537,6 +1537,68 @@ fn hotspot_modules_do_not_log_sensitive_payload_like_fields() { } } +#[test] +fn identity_oauth_does_not_persist_raw_userinfo_claims() { + let source = read_workspace_file("apps/aether-gateway/src/oauth/identity_repo.rs"); + assert!( + !source.contains("claims.raw"), + "identity OAuth must persist only dedicated claim fields, never raw userinfo JSON" + ); +} + +#[test] +fn gateway_does_not_log_raw_request_queries_or_credential_bearing_urls() { + let patterns = [ + "request_query_string = %", + "request_query_string = ?", + "request_path_and_query = %request_context.request_path_and_query()", + "request_path_and_query = ?request_context.request_path_and_query()", + "upstream_url = %", + "upstream_url = ?", + "plan_url = %", + "plan_url = ?", + "token_url = %", + "token_url = ?", + "endpoint_base_url = %", + "endpoint_base_url = ?", + "url = %plan.url", + "url = ?plan.url", + ]; + + for root in [ + "src/ai_serving", + "src/execution_runtime", + "src/handlers", + "src/maintenance", + "src/oauth", + "src/state", + ] { + assert_no_sensitive_log_patterns(root, &patterns); + } + + let provider_transport_root = + Path::new(env!("CARGO_MANIFEST_DIR")).join("../../crates/aether-provider/transport/src"); + let mut provider_transport_files = Vec::new(); + collect_rust_files(&provider_transport_root, &mut provider_transport_files); + let violations = provider_transport_files + .into_iter() + .filter_map(|path| { + let source = std::fs::read_to_string(&path).expect("source file should be readable"); + let hits = patterns + .iter() + .filter(|pattern| source.contains(**pattern)) + .copied() + .collect::>(); + (!hits.is_empty()).then(|| format!("{} -> {}", path.display(), hits.join(", "))) + }) + .collect::>(); + assert!( + violations.is_empty(), + "provider transport must log URL origins/hosts instead of raw URLs or queries:\n{}", + violations.join("\n") + ); +} + #[test] fn execution_runtime_video_finalize_paths_depend_on_shared_video_task_core() { let response = diff --git a/apps/aether-gateway/src/tests/async_task.rs b/apps/aether-gateway/src/tests/async_task.rs index 40f63d33d..2af4bfa87 100644 --- a/apps/aether-gateway/src/tests/async_task.rs +++ b/apps/aether-gateway/src/tests/async_task.rs @@ -5,13 +5,14 @@ use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use aether_data_contracts::repository::video_tasks::{ - UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, + UpsertVideoTask, VideoTaskLookupKey, VideoTaskReadRepository, VideoTaskStatus, + VideoTaskWriteRepository, }; use super::{ - any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server, - to_bytes, AppState, Arc, Body, Bytes, HeaderValue, Json, Mutex, Request, Response, Router, - StatusCode, + any, authenticated_operational_client, build_state_with_execution_runtime_override, json, + start_authenticated_operational_server, start_server, to_bytes, AppState, Arc, Body, Bytes, + HeaderValue, Json, Mutex, Request, Response, Router, StatusCode, }; fn sample_video_task( @@ -75,6 +76,93 @@ fn sample_video_task( } } +fn bound_provider_api_key(provider_id: &str, key_id: &str, plaintext: &str) -> String { + let purpose = format!( + "provider-catalog-credential-bound-v2\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field=api-key", + provider_id.len(), + key_id.len(), + ); + let protected = format!("{purpose}\0{plaintext}"); + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &protected) + .expect("bound provider api key should encrypt"); + format!("aether-provider-catalog-credential-v2:aether-runtime-secret-v1:{ciphertext}") +} + +fn video_provider_catalog_repository( + provider_id: &str, + provider_type: &str, + endpoint_id: &str, + api_format: &str, + endpoint_base_url: &str, + key_id: &str, + upstream_api_key: &str, +) -> Arc { + let provider = StoredProviderCatalogProvider::new( + provider_id.to_string(), + format!("video-{provider_type}"), + Some("https://example.com".to_string()), + provider_type.to_string(), + ) + .expect("provider should build") + .with_transport_fields( + true, + false, + false, + None, + Some(2), + None, + Some(20.0), + None, + None, + ); + let endpoint = StoredProviderCatalogEndpoint::new( + endpoint_id.to_string(), + provider_id.to_string(), + api_format.to_string(), + Some(provider_type.to_string()), + Some("video".to_string()), + true, + ) + .expect("endpoint should build") + .with_transport_fields( + endpoint_base_url.to_string(), + None, + None, + Some(2), + None, + None, + None, + None, + ) + .expect("endpoint transport should build"); + let key = StoredProviderCatalogKey::new( + key_id.to_string(), + provider_id.to_string(), + "primary".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + Some(json!([api_format])), + bound_provider_api_key(provider_id, key_id, upstream_api_key), + None, + None, + Some(json!({api_format: 1})), + None, + None, + None, + None, + ) + .expect("key transport should build"); + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![key], + )) +} + #[tokio::test] async fn gateway_lists_video_tasks_via_internal_async_task_endpoint() { let repository = Arc::new(InMemoryVideoTaskRepository::default()); @@ -112,14 +200,14 @@ async fn gateway_lists_video_tasks_via_internal_async_task_endpoint() { .await .expect("upsert should succeed"); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(repository), - ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(repository); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks?status=completed&page=1&page_size=1" )) @@ -169,14 +257,14 @@ async fn gateway_exposes_video_task_stats_via_internal_async_task_endpoint() { .await .expect("upsert should succeed"); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(repository), - ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(repository); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/stats" )) @@ -212,14 +300,14 @@ async fn gateway_reads_video_task_detail_via_internal_async_task_endpoint() { .await .expect("upsert should succeed"); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(repository), - ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(repository); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-1" )) @@ -238,7 +326,7 @@ async fn gateway_reads_video_task_detail_via_internal_async_task_endpoint() { } #[tokio::test] -async fn gateway_redirects_direct_video_task_video_from_internal_async_task_endpoint() { +async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endpoint() { let repository = Arc::new(InMemoryVideoTaskRepository::default()); let mut task = sample_video_task( "task-redirect", @@ -248,23 +336,20 @@ async fn gateway_redirects_direct_video_task_video_from_internal_async_task_endp "user-1", "openai:video", ); - task.video_url = Some("https://cdn.example.com/video-task-redirect.mp4".to_string()); - repository + task.video_url = Some("https://8.8.8.8/video-task-redirect.mp4".to_string()); + let stored = repository .upsert(task) .await .expect("upsert should succeed"); + assert_eq!(stored.video_url, None); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(repository), - ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(repository); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("client should build"); + let client = authenticated_operational_client(&access_token); let response = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-redirect/video" @@ -273,20 +358,13 @@ async fn gateway_redirects_direct_video_task_video_from_internal_async_task_endp .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT); - assert_eq!( - response - .headers() - .get(http::header::LOCATION) - .and_then(|value| value.to_str().ok()), - Some("https://cdn.example.com/video-task-redirect.mp4") - ); + assert_eq!(response.status(), StatusCode::NOT_FOUND); gateway_handle.abort(); } #[tokio::test] -async fn gateway_proxies_gemini_video_task_video_from_internal_async_task_endpoint() { +async fn gateway_rejects_private_gemini_video_target_without_sending_provider_key() { let seen_api_key = Arc::new(Mutex::new(None::)); let seen_api_key_clone = Arc::clone(&seen_api_key); @@ -384,8 +462,11 @@ async fn gateway_proxies_gemini_video_task_video_from_internal_async_task_endpoi .expect("key should build") .with_transport_fields( Some(json!(["gemini:video"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "gemini-upstream-secret") - .expect("api key should encrypt"), + bound_provider_api_key( + "provider-gemini-video-1", + "key-gemini-video-1", + "gemini-upstream-secret", + ), None, None, Some(json!({"gemini:video": 1})), @@ -397,20 +478,20 @@ async fn gateway_proxies_gemini_video_task_video_from_internal_async_task_endpoi .expect("key transport should build")], )); - let gateway = build_router_with_state( - AppState::new() + let state = AppState::new() .expect("gateway state should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( - repository, - provider_catalog_repository, - DEVELOPMENT_ENCRYPTION_KEY, - ), + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + repository, + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, ), ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-proxy/video" )) @@ -418,28 +499,10 @@ async fn gateway_proxies_gemini_video_task_video_from_internal_async_task_endpoi .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .headers() - .get(http::header::CONTENT_TYPE) - .and_then(|value| value.to_str().ok()), - Some("video/mp4") - ); - assert_eq!( - response - .headers() - .get(http::header::CONTENT_DISPOSITION) - .and_then(|value| value.to_str().ok()), - Some("inline; filename=\"video_task-proxy.mp4\"") - ); - assert_eq!( - response.bytes().await.expect("body should read"), - Bytes::from_static(b"proxied-video-bytes") - ); + assert_eq!(response.status(), StatusCode::BAD_GATEWAY); assert_eq!( seen_api_key.lock().expect("mutex should lock").as_deref(), - Some("gemini-upstream-secret") + None ); gateway_handle.abort(); @@ -563,14 +626,29 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint() { .await .expect("upsert should succeed"); - let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; - let gateway = build_router_with_state( - build_state_with_execution_runtime_override(execution_runtime_url) - .with_video_task_data_repository_for_tests(Arc::clone(&repository)), + let provider_catalog_repository = video_provider_catalog_repository( + "provider-1", + "openai", + "endpoint-1", + "openai:video", + "https://api.openai.example/v1", + "provider-key-1", + "sk-upstream-openai-video", ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let state = build_state_with_execution_runtime_override(execution_runtime_url) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + Arc::clone(&repository), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .post(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-openai-cancel/cancel" )) @@ -606,7 +684,7 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint() { "Bearer sk-upstream-openai-video" ); - let detail = reqwest::Client::new() + let detail = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-openai-cancel" )) @@ -617,10 +695,10 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint() { let detail_json: serde_json::Value = detail.json().await.expect("detail should parse"); assert_eq!(detail_json["status"], "Cancelled"); assert!(detail_json["next_poll_at_unix_secs"].is_null()); - assert_eq!( - detail_json["request_metadata"]["rust_local_snapshot"]["OpenAi"]["status"], - "Cancelled" - ); + assert!(detail_json["original_request_body"].is_null()); + assert!(detail_json["progress_message"].is_null()); + assert!(detail_json["error_message"].is_null()); + assert!(detail_json["request_metadata"].is_null()); gateway_handle.abort(); execution_runtime_handle.abort(); @@ -726,14 +804,29 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint_with .await .expect("upsert should succeed"); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(Arc::clone(&repository)), + let provider_catalog_repository = video_provider_catalog_repository( + "provider-1", + "openai", + "endpoint-1", + "openai:video", + &upstream_api_root, + "provider-key-1", + "sk-upstream-openai-video", ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + Arc::clone(&repository), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .post(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-openai-cancel-direct/cancel" )) @@ -763,7 +856,7 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint_with }) ); - let detail = reqwest::Client::new() + let detail = client .get(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-openai-cancel-direct" )) @@ -774,10 +867,10 @@ async fn gateway_cancels_openai_video_task_via_internal_async_task_endpoint_with let detail_json: serde_json::Value = detail.json().await.expect("detail should parse"); assert_eq!(detail_json["status"], "Cancelled"); assert!(detail_json["next_poll_at_unix_secs"].is_null()); - assert_eq!( - detail_json["request_metadata"]["rust_local_snapshot"]["OpenAi"]["status"], - "Cancelled" - ); + assert!(detail_json["original_request_body"].is_null()); + assert!(detail_json["progress_message"].is_null()); + assert!(detail_json["error_message"].is_null()); + assert!(detail_json["request_metadata"].is_null()); gateway_handle.abort(); upstream_handle.abort(); @@ -798,14 +891,14 @@ async fn gateway_rejects_terminal_video_task_cancel_via_internal_async_task_endp .await .expect("upsert should succeed"); - let gateway = build_router_with_state( - AppState::new() - .expect("gateway state should build") - .with_video_task_data_repository_for_tests(repository), - ); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(repository); + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = client .post(format!( "{gateway_url}/_gateway/async-tasks/video-tasks/task-cancelled-already/cancel" )) @@ -828,3 +921,41 @@ async fn gateway_rejects_terminal_video_task_cancel_via_internal_async_task_endp gateway_handle.abort(); } + +#[tokio::test] +async fn user_scoped_video_task_cancel_rechecks_owner_before_side_effects() { + let repository = Arc::new(InMemoryVideoTaskRepository::default()); + repository + .upsert(sample_video_task( + "task-owner-bound-cancel", + VideoTaskStatus::Processing, + 100, + "sora-2", + "user-owner", + "openai:video", + )) + .await + .expect("upsert should succeed"); + let state = AppState::new() + .expect("gateway state should build") + .with_video_task_data_repository_for_tests(Arc::clone(&repository)); + + let error = crate::async_task::cancel_video_task_record_for_user( + &state, + "task-owner-bound-cancel", + "user-foreign", + ) + .await + .expect_err("a non-owner cancellation must be hidden as not found"); + + assert!(matches!( + error, + crate::async_task::CancelVideoTaskError::NotFound + )); + let stored = repository + .find(VideoTaskLookupKey::Id("task-owner-bound-cancel")) + .await + .expect("task lookup should succeed") + .expect("task should remain present"); + assert_eq!(stored.status, VideoTaskStatus::Processing); +} diff --git a/apps/aether-gateway/src/tests/audit.rs b/apps/aether-gateway/src/tests/audit.rs index 776a0203e..02e316451 100644 --- a/apps/aether-gateway/src/tests/audit.rs +++ b/apps/aether-gateway/src/tests/audit.rs @@ -22,8 +22,9 @@ use serde_json::Value; use sha2::{Digest, Sha256}; use super::{ - any, build_router_with_state, json, start_server, AppState, Json, Request, Router, StatusCode, - UsageRuntimeConfig, CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER, TRACE_ID_HEADER, + any, authenticated_operational_client, json, start_authenticated_operational_server, + start_server, AppState, Json, Request, Router, StatusCode, UsageRuntimeConfig, + CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER, TRACE_ID_HEADER, }; fn hash_api_key(value: &str) -> String { @@ -239,8 +240,9 @@ async fn gateway_exposes_request_id_header_for_local_execution_response_impl() { enabled: true, ..UsageRuntimeConfig::default() }); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); let response = reqwest::Client::new() .post(format!("{gateway_url}/v1/chat/completions")) @@ -276,7 +278,7 @@ async fn gateway_exposes_request_id_header_for_local_execution_response_impl() { tokio::time::sleep(std::time::Duration::from_millis(10)).await; } - let audit_response = reqwest::Client::new() + let audit_response = operational_client .get(format!( "{gateway_url}/_gateway/audit/request-audit/{request_id}?attempted_only=true" )) @@ -444,10 +446,11 @@ async fn gateway_exposes_request_usage_via_internal_audit_endpoint() { let gateway_state = AppState::new() .expect("gateway state should build") .with_usage_data_reader_for_tests(repository); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!( "{gateway_url}/_gateway/audit/request-usage/req-usage-2" )) @@ -501,10 +504,11 @@ async fn gateway_exposes_request_audit_bundle_via_internal_audit_endpoint() { provider_catalog, usage_repository, ); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!( "{gateway_url}/_gateway/audit/request-audit/req-audit-1?attempted_only=true" )) @@ -554,10 +558,11 @@ async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() { let gateway_state = AppState::new() .expect("gateway state should build") .with_request_candidate_data_reader_for_tests(repository); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!( "{gateway_url}/_gateway/audit/request-candidates/req-trace-1?attempted_only=true" )) @@ -603,10 +608,11 @@ async fn gateway_exposes_decision_trace_via_internal_audit_endpoint() { let gateway_state = AppState::new() .expect("gateway state should build") .with_decision_trace_data_readers_for_tests(request_candidates, provider_catalog); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!( "{gateway_url}/_gateway/audit/decision-trace/req-trace-2?attempted_only=true" )) @@ -621,7 +627,7 @@ async fn gateway_exposes_decision_trace_via_internal_audit_endpoint() { assert_eq!(payload["candidates"][0]["provider_name"], "OpenAI"); assert_eq!( payload["candidates"][0]["provider_website"], - "https://openai.com" + "https://openai.com/" ); assert_eq!( payload["candidates"][0]["endpoint_api_format"], @@ -650,10 +656,11 @@ async fn gateway_exposes_auth_api_key_snapshot_via_internal_audit_endpoint() { let gateway_state = AppState::new() .expect("gateway state should build") .with_auth_api_key_data_reader_for_tests(repository); - let gateway = build_router_with_state(gateway_state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(gateway_state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!( "{gateway_url}/_gateway/audit/auth/users/user-1/api-keys/key-1" )) diff --git a/apps/aether-gateway/src/tests/concurrency.rs b/apps/aether-gateway/src/tests/concurrency.rs index 6cc596217..9df8302e2 100644 --- a/apps/aether-gateway/src/tests/concurrency.rs +++ b/apps/aether-gateway/src/tests/concurrency.rs @@ -5,9 +5,10 @@ use super::usage::{ sample_local_openai_endpoint, sample_local_openai_key, sample_local_openai_provider, }; use super::{ - any, build_router_with_state, build_state_with_execution_runtime_override, start_server, - wait_until, AppState, Arc, Body, Bytes, GatewayFallbackMetricKind, GatewayFallbackReason, - HeaderValue, Infallible, Request, Response, Router, StatusCode, + any, authenticated_operational_client, build_router_with_state, + build_state_with_execution_runtime_override, start_authenticated_operational_server, + start_server, wait_until, AppState, Arc, Body, Bytes, GatewayFallbackMetricKind, + GatewayFallbackReason, HeaderValue, Infallible, Request, Response, Router, StatusCode, EXECUTION_PATH_DISTRIBUTED_OVERLOADED, EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, EXECUTION_PATH_LOCAL_OVERLOADED, }; @@ -389,6 +390,43 @@ fn gateway_exposes_request_concurrency_metrics() { ); } +#[test] +fn gateway_rejects_anonymous_operational_metrics() { + run_concurrency_test( + "gateway_rejects_anonymous_operational_metrics", + gateway_rejects_anonymous_operational_metrics_impl, + ); +} + +async fn gateway_rejects_anonymous_operational_metrics_impl() { + let gateway = build_router_with_state(AppState::new().expect("gateway state should build")); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + assert_eq!( + response + .headers() + .get(http::header::WWW_AUTHENTICATE) + .and_then(|value| value.to_str().ok()), + Some("Bearer") + ); + + gateway_handle.abort(); +} + async fn gateway_exposes_request_concurrency_metrics_impl() { let state = AppState::new() .expect("gateway state should build") @@ -403,10 +441,11 @@ async fn gateway_exposes_request_concurrency_metrics_impl() { 9, )); assert!(state.prewarm_metric_snapshot().await); - let gateway = build_router_with_state(state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!("{gateway_url}/_gateway/metrics")) .send() .await @@ -532,10 +571,11 @@ async fn gateway_exposes_fallback_metrics_impl() { GatewayFallbackReason::LocalExecutionPathRequired, ); assert!(state.prewarm_metric_snapshot().await); - let gateway = build_router_with_state(state); - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (gateway_url, gateway_handle, access_token) = + start_authenticated_operational_server(state).await; + let operational_client = authenticated_operational_client(&access_token); - let response = reqwest::Client::new() + let response = operational_client .get(format!("{gateway_url}/_gateway/metrics")) .send() .await diff --git a/apps/aether-gateway/src/tests/control/admin/api_keys.rs b/apps/aether-gateway/src/tests/control/admin/api_keys.rs index 5c04ba601..fa84737e4 100644 --- a/apps/aether-gateway/src/tests/control/admin/api_keys.rs +++ b/apps/aether-gateway/src/tests/control/admin/api_keys.rs @@ -1,8 +1,9 @@ use std::sync::{Arc, Mutex}; -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::auth::{ - InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, + AuthApiKeyWriteRepository, InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, + StoredAuthApiKeySnapshot, }; use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot}; @@ -12,6 +13,7 @@ use axum::routing::any; use axum::{extract::Request, Router}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use super::super::{build_router_with_state, start_server, AppState}; use crate::constants::{ @@ -28,6 +30,12 @@ fn admin_request(builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder { .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") } +fn hash_api_key(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} + async fn start_api_keys_upstream( path: &'static str, ) -> (String, Arc>, tokio::task::JoinHandle<()>) { @@ -85,14 +93,26 @@ fn sample_standalone_export_record( plaintext_key: &str, is_active: bool, ) -> StoredAuthApiKeyExportRecord { + let key_hash = hash_api_key(plaintext_key); + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let key_encrypted = crate::handlers::shared::seal_auth_api_key_secret( + &bootstrap, + user_id, + api_key_id, + &key_hash, + true, + plaintext_key, + ) + .expect("key should encrypt"); let mut record = StoredAuthApiKeyExportRecord::new( user_id.to_string(), api_key_id.to_string(), - format!("hash-{api_key_id}"), - Some( - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, plaintext_key) - .expect("key should encrypt"), - ), + key_hash, + Some(key_encrypted), Some(format!("key-{api_key_id}")), Some(json!(["openai"])), Some(json!(["openai:chat"])), @@ -240,10 +260,7 @@ async fn gateway_handles_admin_api_keys_list_locally_with_trusted_admin_principa assert_eq!(payload["skip"], json!(0)); assert_eq!(payload["api_keys"][0]["id"], json!("key-1")); assert_eq!(payload["api_keys"][0]["is_standalone"], json!(true)); - assert_eq!( - payload["api_keys"][0]["key_display"], - json!("sk-key-1-p...text") - ); + assert_eq!(payload["api_keys"][0]["key_display"], json!("sk-ke...text")); assert_eq!(payload["api_keys"][0]["total_requests"], json!(7)); assert_eq!(payload["api_keys"][0]["total_tokens"], json!(0)); assert_eq!( @@ -317,7 +334,7 @@ async fn gateway_handles_admin_api_keys_detail_locally_with_trusted_admin_princi assert_eq!(payload["wallet"]["id"], json!("wallet-key-1")); assert_eq!(payload["wallet"]["unlimited"], json!(true)); assert_eq!(payload["wallet"]["balance"], json!(20.0)); - assert_eq!(payload["key_display"], json!("sk-key-1-p...text")); + assert_eq!(payload["key_display"], json!("sk-ke...text")); assert_eq!(payload["total_tokens"], json!(77)); assert_eq!(payload["created_at"], json!("2024-03-21T05:48:20+00:00")); assert_eq!(payload["last_used_at"], json!("2024-03-21T05:48:22+00:00")); @@ -442,8 +459,10 @@ async fn gateway_handles_admin_api_key_install_session_locally_with_trusted_admi AppState::new() .expect("gateway should build") .with_data_state_for_tests( - GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) - .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone( + &auth_repository, + )) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -479,15 +498,78 @@ async fn gateway_handles_admin_api_key_install_session_locally_with_trusted_admi assert_eq!( payload["unix_command"], json!(format!( - "curl -fsSL https://aether.example/install/{install_code} | sh" + "curl -fsSL 'https://aether.example/install/{install_code}' | sh" )) ); assert_eq!( payload["powershell_command"], json!(format!( - "irm https://aether.example/install/{install_code}.ps1 | iex" + "irm 'https://aether.example/install/{install_code}.ps1' | iex" )) ); + + let second_response = admin_request(reqwest::Client::new().post(format!( + "{gateway_url}/api/admin/api-keys/key-1/install-sessions" + ))) + .header("x-forwarded-host", "aether.example") + .header("x-forwarded-proto", "https") + .json(&json!({ + "target_cli": "codex_cli", + "target_system": "linux", + })) + .send() + .await + .expect("second install session should be created"); + assert_eq!(second_response.status(), StatusCode::OK); + let second_payload: serde_json::Value = second_response + .json() + .await + .expect("second install session response should parse"); + let second_install_code = second_payload["install_code"] + .as_str() + .expect("second install code should be returned"); + + let script_response = reqwest::Client::new() + .get(format!("{gateway_url}/install/{install_code}")) + .send() + .await + .expect("install script should resolve"); + assert_eq!(script_response.status(), StatusCode::OK); + assert_eq!( + script_response + .headers() + .get("cache-control") + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + assert!(script_response + .text() + .await + .expect("install script should be readable") + .contains("sk-key-1-plaintext")); + + let replay = reqwest::Client::new() + .get(format!("{gateway_url}/install/{install_code}")) + .send() + .await + .expect("install replay should receive a response"); + assert_eq!(replay.status(), StatusCode::NOT_FOUND); + + assert!(auth_repository + .delete_standalone_api_key("key-1") + .await + .expect("test API key deletion should succeed")); + let deleted_key_session = reqwest::Client::new() + .get(format!("{gateway_url}/install/{second_install_code}")) + .send() + .await + .expect("deleted-key install session should receive a response"); + assert_eq!(deleted_key_session.status(), StatusCode::NOT_FOUND); + assert!(!deleted_key_session + .text() + .await + .expect("deleted-key response should be readable") + .contains("sk-key-1-plaintext")); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); 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 44c706b45..1043eb702 100644 --- a/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs +++ b/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs @@ -1,9 +1,7 @@ use std::sync::{Arc, Mutex}; use aether_contracts::ExecutionPlan; -use aether_crypto::{ - decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, -}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, StoredProviderCatalogEndpoint, @@ -18,8 +16,8 @@ use http::StatusCode; use serde_json::json; use super::super::super::{ - build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint, - sample_key, sample_provider, start_server, AppState, + build_router_with_state, build_state_with_execution_runtime_override, sample_bound_auth_config, + sample_bound_key, sample_endpoint, sample_provider, start_server, AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -29,6 +27,18 @@ use crate::data::GatewayDataState; const PROVIDER_KEYS_TEST_STACK_BYTES: usize = 16 * 1024 * 1024; +fn open_provider_catalog_api_key_for_test(key: &StoredProviderCatalogKey) -> String { + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + state + .decrypt_provider_catalog_key_api_key(key) + .expect("provider catalog API key should decrypt") + .expect("provider catalog API key should be present") +} + fn run_provider_keys_test(test_name: &'static str, make_future: F) where F: FnOnce() -> Fut + Send + 'static, @@ -172,7 +182,7 @@ async fn gateway_handles_admin_provider_keys_locally_with_trusted_admin_principa }), ); - let mut key_a = sample_key( + let mut key_a = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -193,7 +203,7 @@ async fn gateway_handles_admin_provider_keys_locally_with_trusted_admin_principa "quota": {"code": "unknown", "exhausted": false} })); - let mut key_b = sample_key( + let mut key_b = sample_bound_key( "key-openai-b", "provider-openai", "openai:chat", @@ -245,7 +255,7 @@ async fn gateway_handles_admin_provider_keys_locally_with_trusted_admin_principa assert_eq!(items[0]["success_count"], 9); assert_eq!(items[0]["error_count"], 3); assert_eq!(items[0]["note"], "primary key"); - assert_eq!(items[0]["api_key_masked"], "sk-test-a***"); + assert_eq!(items[0]["api_key_masked"], "sk***-a"); assert_eq!(items[1]["id"], "key-openai-b"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -255,25 +265,26 @@ async fn gateway_handles_admin_provider_keys_locally_with_trusted_admin_principa #[tokio::test] async fn gateway_provider_keys_expose_circuit_breaker_and_recover_clears_it() { - let key = sample_key("key-1", "provider-1", "openai:chat", "sk-test-a").with_health_fields( - Some(json!({"openai:chat": { - "health_score": 0.2, - "consecutive_failures": 8, - "last_failure_at": "2026-03-26T12:00:00+00:00" - }})), - Some(json!({"openai:chat": { - "open": true, - "open_at": "2026-03-26T12:00:00+00:00", - "reason": "consecutive_failures_8", - "next_probe_at": "2099-03-26T12:01:00+00:00", - "next_probe_at_unix_secs": 4078209660u64, - "probe_interval_minutes": 1, - "max_probe_interval_minutes": 32, - "half_open_until": null, - "half_open_successes": 0, - "half_open_failures": 0 - }})), - ); + let key = sample_bound_key("key-1", "provider-1", "openai:chat", "sk-test-a") + .with_health_fields( + Some(json!({"openai:chat": { + "health_score": 0.2, + "consecutive_failures": 8, + "last_failure_at": "2026-03-26T12:00:00+00:00" + }})), + Some(json!({"openai:chat": { + "open": true, + "open_at": "2026-03-26T12:00:00+00:00", + "reason": "consecutive_failures_8", + "next_probe_at": "2099-03-26T12:01:00+00:00", + "next_probe_at_unix_secs": 4078209660u64, + "probe_interval_minutes": 1, + "max_probe_interval_minutes": 32, + "half_open_until": null, + "half_open_successes": 0, + "half_open_failures": 0 + }})), + ); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-1", "openai", 10)], vec![sample_endpoint( @@ -369,7 +380,7 @@ async fn gateway_handles_admin_provider_keys_page_locally_with_total() { }), ); - let mut key_a = sample_key( + let mut key_a = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -378,7 +389,7 @@ async fn gateway_handles_admin_provider_keys_page_locally_with_total() { key_a.internal_priority = 10; key_a.created_at_unix_ms = Some(1_711_000_000); - let mut key_b = sample_key( + let mut key_b = sample_bound_key( "key-openai-b", "provider-openai", "openai:chat", @@ -387,7 +398,7 @@ async fn gateway_handles_admin_provider_keys_page_locally_with_total() { key_b.internal_priority = 20; key_b.created_at_unix_ms = Some(1_711_100_000); - let mut key_c = sample_key( + let mut key_c = sample_bound_key( "key-openai-c", "provider-openai", "openai:chat", @@ -455,24 +466,22 @@ async fn gateway_admin_provider_keys_prefers_upstream_plan_type_over_auth_config let mut provider = sample_provider("provider-codex", "codex", 10); provider.provider_type = "codex".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-oauth", "provider-codex", "openai:responses", "oauth-placeholder", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - &json!({ - "plan_type": "free", - "account_id": "acct-codex-legacy" - }) - .to_string(), - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-oauth", + &json!({ + "plan_type": "free", + "account_id": "acct-codex-legacy" + }) + .to_string(), + )); key.upstream_metadata = Some(json!({ "codex": { "plan_type": "plus", @@ -539,20 +548,18 @@ async fn gateway_admin_provider_keys_marks_oauth_header_auth() { let mut provider = sample_provider("provider-codex", "codex", 10); provider.provider_type = "codex".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-oauth-header", "provider-codex", "openai:responses", "imported-session-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-oauth-header", + r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -657,7 +664,7 @@ async fn gateway_creates_admin_provider_key_locally_with_trusted_admin_principal assert_eq!(payload["name"], "created key"); assert_eq!(payload["internal_priority"], 15); assert_eq!(payload["api_formats"], json!(["openai:chat"])); - assert_eq!(payload["api_key_masked"], "sk-creat***enai"); + assert_eq!(payload["api_key_masked"], "sk-c***enai"); assert_eq!(payload["request_count"], 0); assert_eq!(payload["success_count"], 0); assert_eq!(payload["error_count"], 0); @@ -682,7 +689,7 @@ async fn generic_key_routes_reject_agent_identity_credential_writes() { let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-codex", "codex", 10)], vec![], - vec![sample_key( + vec![sample_bound_key( "key-codex-existing", "provider-codex", "openai:responses", @@ -776,20 +783,18 @@ async fn generic_key_routes_reject_agent_identity_credential_writes() { #[tokio::test] async fn generic_codex_key_credential_switch_rotates_generation_and_clears_quota() { - let mut existing_key = sample_key( + let mut existing_key = sample_bound_key( "key-codex-existing", "provider-codex", "openai:responses", "old-oauth-access-token", ); existing_key.auth_type = "oauth".to_string(); - existing_key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","refresh_token":"old-refresh-token"}"#, - ) - .expect("old auth config should encrypt"), - ); + existing_key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-existing", + r#"{"provider_type":"codex","refresh_token":"old-refresh-token"}"#, + )); existing_key.upstream_metadata = Some(json!({ "codex": { "credential_generation": "generation-before-switch", @@ -1038,7 +1043,7 @@ async fn provider_key_concurrent_limit_create_and_list_responses() { #[tokio::test] async fn provider_key_concurrent_limit_reads_existing_list_response() { - let mut key_a = sample_key( + let mut key_a = sample_bound_key( "provider-key-a", "test-provider-a", "openai:chat", @@ -1046,7 +1051,7 @@ async fn provider_key_concurrent_limit_reads_existing_list_response() { ); key_a.concurrent_limit = Some(1); - let mut key_b = sample_key( + let mut key_b = sample_bound_key( "provider-key-b", "test-provider-a", "openai:chat", @@ -1232,7 +1237,7 @@ async fn gateway_reveals_admin_provider_key_locally_with_trusted_admin_principal let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-openai", "openai", 10)], vec![], - vec![sample_key( + vec![sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -1290,20 +1295,18 @@ async fn gateway_exports_admin_provider_key_locally_with_trusted_admin_principal }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-kiro-a", "provider-kiro", "claude:messages", "oauth-access-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"kiro","auth_method":"idc","refresh_token":"rt-kiro-123"}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + "key-kiro-a", + r#"{"provider_type":"kiro","auth_method":"idc","refresh_token":"rt-kiro-123"}"#, + )); key.upstream_metadata = Some(json!({"kiro": {"email": "alice@example.com"}})); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( @@ -1366,20 +1369,18 @@ async fn gateway_exports_admin_provider_key_access_token_when_refresh_token_is_m }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-a", "provider-codex", "openai:responses", "codex-access-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","email":"codex@example.com","updated_at":1710000000}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-a", + r#"{"provider_type":"codex","email":"codex@example.com","updated_at":1710000000}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-codex", "codex", 10)], @@ -1440,20 +1441,18 @@ async fn gateway_export_does_not_emit_access_token_from_imported_authorization_h }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-a", "provider-codex", "openai:responses", "imported-session-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","email":"codex@example.com","headers":{"authorization":"Bearer imported-session-token"}}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-a", + r#"{"provider_type":"codex","email":"codex@example.com","headers":{"authorization":"Bearer imported-session-token"}}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-codex", "codex", 10)], @@ -1516,20 +1515,18 @@ async fn gateway_export_preserves_distinct_imported_access_token_with_authorizat }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-a", "provider-codex", "openai:responses", "jwt-access-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","email":"codex@example.com","access_token":"jwt-access-token","headers":{"authorization":"Bearer imported-session-token"}}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-a", + r#"{"provider_type":"codex","email":"codex@example.com","access_token":"jwt-access-token","headers":{"authorization":"Bearer imported-session-token"}}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-codex", "codex", 10)], @@ -1581,27 +1578,25 @@ async fn gateway_generic_export_rejects_agent_identity_without_exposing_private_ let private_key = "agent-private-key-must-not-leak"; let mut provider = sample_provider("provider-codex", "codex", 10); provider.provider_type = "codex".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-agent", "provider-codex", "openai:responses", "__placeholder__", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - &json!({ - "provider_type": "codex", - "auth_mode": "agentIdentity", - "agent_runtime_id": "runtime-must-not-leak", - "agent_private_key": private_key, - "task_id": "task-must-not-leak" - }) - .to_string(), - ) - .expect("Agent Identity auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-agent", + &json!({ + "provider_type": "codex", + "auth_mode": "agentIdentity", + "agent_runtime_id": "runtime-must-not-leak", + "agent_private_key": private_key, + "task_id": "task-must-not-leak" + }) + .to_string(), + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![], @@ -1706,7 +1701,7 @@ async fn gateway_clears_admin_provider_key_oauth_invalid_locally_with_trusted_ad }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -1829,7 +1824,7 @@ async fn gateway_noops_admin_provider_key_oauth_invalid_clear_when_marker_absent let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-openai", "openai", 10)], vec![], - vec![sample_key( + vec![sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -1886,7 +1881,7 @@ async fn gateway_updates_admin_provider_key_locally_with_trusted_admin_principal }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -1961,14 +1956,7 @@ async fn gateway_updates_admin_provider_key_locally_with_trusted_admin_principal assert_eq!(reloaded[0].allowed_models, None); assert_eq!(reloaded[0].note.as_deref(), Some("updated from rust")); assert!(!reloaded[0].is_active); - let decrypted = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - reloaded[0] - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("ciphertext should decrypt"); + let decrypted = open_provider_catalog_api_key_for_test(&reloaded[0]); assert_eq!(decrypted, "sk-updated-openai"); gateway_handle.abort(); @@ -1977,7 +1965,7 @@ async fn gateway_updates_admin_provider_key_locally_with_trusted_admin_principal #[tokio::test] async fn provider_key_concurrent_limit_update_presence_semantics() { - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2123,7 +2111,7 @@ async fn gateway_clears_allowed_models_when_disabling_auto_fetch_on_provider_key }), ); - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2229,7 +2217,7 @@ async fn gateway_overwrites_allowed_models_immediately_when_enabling_auto_fetch_ ); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2345,7 +2333,7 @@ async fn gateway_fetches_allowed_models_immediately_when_enabling_auto_fetch_fro ); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2460,7 +2448,7 @@ async fn gateway_refreshes_allowed_models_when_updating_include_patterns_with_au ); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2572,7 +2560,7 @@ async fn gateway_refreshes_allowed_models_when_updating_exclude_patterns_with_au ); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; - let mut key = sample_key( + let mut key = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2667,13 +2655,13 @@ async fn gateway_rejects_admin_provider_key_update_when_api_key_duplicates_exist vec![sample_provider("provider-openai", "openai", 10)], vec![], vec![ - sample_key( + sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", "sk-test-a", ), - sample_key( + sample_bound_key( "key-openai-b", "provider-openai", "openai:chat", @@ -2740,7 +2728,7 @@ async fn gateway_deletes_admin_provider_key_locally_with_trusted_admin_principal let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-openai", "openai", 10)], vec![], - vec![sample_key( + vec![sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2807,13 +2795,13 @@ async fn gateway_batch_deletes_admin_provider_keys_locally_with_trusted_admin_pr vec![sample_provider("provider-openai", "openai", 10)], vec![], vec![ - sample_key( + sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", "sk-test-a", ), - sample_key( + sample_bound_key( "key-openai-b", "provider-openai", "openai:chat", @@ -2884,7 +2872,7 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin }), ); - let mut key_a = sample_key( + let mut key_a = sample_bound_key( "key-openai-a", "provider-openai", "openai:chat", @@ -2900,7 +2888,7 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin key_a.health_by_format = Some(json!({"openai:chat": {"health_score": 0.8}})); key_a.circuit_breaker_by_format = Some(json!({"openai:chat": {"open": false}})); - let mut key_b = sample_key( + let mut key_b = sample_bound_key( "key-claude-a", "provider-claude", "claude:messages", @@ -2912,20 +2900,18 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin key_b.created_at_unix_ms = Some(1_711_100_000); key_b.updated_at_unix_secs = Some(1_711_100_100); - let mut key_agent = sample_key( + let mut key_agent = sample_bound_key( "key-codex-agent", "provider-codex", "openai:responses", "__placeholder__", ); key_agent.auth_type = "oauth".to_string(); - key_agent.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#, - ) - .expect("Agent Identity auth config should encrypt"), - ); + key_agent.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-agent", + r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#, + )); let mut codex_provider = sample_provider("provider-codex", "codex", 30); codex_provider.provider_type = "codex".to_string(); 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 0c6bc508b..c266a944a 100644 --- a/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs +++ b/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs @@ -4,8 +4,12 @@ use std::sync::{Arc, Mutex}; use aether_crypto::{ decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, }; +use aether_data::repository::global_models::InMemoryGlobalModelReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; +use aether_data_contracts::repository::global_models::{ + AdminProviderModelListQuery, GlobalModelReadRepository, +}; use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; @@ -16,8 +20,8 @@ use http::StatusCode; use serde_json::json; use super::super::super::{ - build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint, - sample_key, sample_proxy_node, start_server, + build_router_with_state, build_state_with_execution_runtime_override, sample_bound_auth_config, + sample_bound_key, sample_endpoint, sample_key, sample_proxy_node, start_server, AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -27,6 +31,28 @@ use crate::data::GatewayDataState; const PROVIDER_QUOTA_TEST_STACK_BYTES: usize = 16 * 1024 * 1024; +fn open_provider_catalog_credential_for_test( + key: &StoredProviderCatalogKey, + field: &str, +) -> String { + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + match field { + "api_key" => state + .decrypt_provider_catalog_key_api_key(key) + .expect("provider catalog API key should decrypt") + .expect("provider catalog API key should be present"), + "auth_config" => state + .decrypt_provider_catalog_key_auth_config(key) + .expect("provider catalog auth config should decrypt") + .expect("provider catalog auth config should be present"), + _ => panic!("unsupported provider catalog credential field: {field}"), + } +} + fn run_provider_quota_test(test_name: &'static str, make_future: F) where F: FnOnce() -> Fut + Send + 'static, @@ -486,23 +512,9 @@ async fn gateway_codex_quota_refresh_persists_after_automatic_oauth_token_refres .await .expect("key should reload"); let persisted = reloaded.first().expect("key should remain installed"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should persist"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = open_provider_catalog_credential_for_test(persisted, "api_key"); assert_eq!(decrypted_api_key, "refreshed-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should persist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = open_provider_catalog_credential_for_test(persisted, "auth_config"); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["refresh_token"], "rotated-codex-refresh-token"); @@ -1390,40 +1402,23 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_ad }), ); - let encrypted_auth_config = encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, + let mut key = sample_bound_key( + "key-kiro-a", + "provider-kiro", + "claude:messages", + "__placeholder__", + ); + key.auth_type = "bearer".to_string(); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + "key-kiro-a", r#"{ "access_token":"kiro-access-token", "api_region":"us-west-2", "machine_id":"123e4567-e89b-12d3-a456-426614174000", "kiro_version":"1.2.3" }"#, - ) - .expect("auth config ciphertext should build"); - let encrypted_api_key = - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__") - .expect("api key ciphertext should build"); - let key = StoredProviderCatalogKey::new( - "key-kiro-a".to_string(), - "provider-kiro".to_string(), - "default".to_string(), - "bearer".to_string(), - None, - true, - ) - .expect("key should build") - .with_transport_fields( - Some(json!(["claude:messages"])), - encrypted_api_key, - Some(encrypted_auth_config), - None, - None, - None, - None, - None, - None, - ) - .expect("key transport should build"); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![StoredProviderCatalogProvider::new( @@ -1479,6 +1474,11 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_ad ); assert_eq!( payload["results"][0]["quota_snapshot"]["plan_type"], + serde_json::Value::Null, + "quota snapshot token projection must reject unsafely formatted plan labels" + ); + assert_eq!( + payload["results"][0]["metadata"]["subscription_title"], "KIRO PRO+" ); assert_eq!( @@ -1556,7 +1556,8 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_ad .as_ref() .and_then(|value| value.get("quota")) .and_then(|value| value.get("plan_type")), - Some(&json!("KIRO PRO+")) + None, + "persisted admin-safe status must omit unsafely formatted plan labels" ); assert_eq!( reloaded[0] @@ -1771,35 +1772,18 @@ async fn gateway_refresh_kiro_quota_reconciles_missing_fixed_endpoint_before_ref }), ); - let encrypted_auth_config = encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, + let mut key = sample_bound_key( + "key-kiro-reconcile", + "provider-kiro-reconcile", + "claude:messages", + "__placeholder__", + ); + key.auth_type = "bearer".to_string(); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro-reconcile", + "key-kiro-reconcile", r#"{"access_token":"kiro-access-token","api_region":"us-west-2"}"#, - ) - .expect("auth config ciphertext should build"); - let encrypted_api_key = - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__") - .expect("api key ciphertext should build"); - let key = StoredProviderCatalogKey::new( - "key-kiro-reconcile".to_string(), - "provider-kiro-reconcile".to_string(), - "default".to_string(), - "bearer".to_string(), - None, - true, - ) - .expect("key should build") - .with_transport_fields( - Some(json!(["claude:messages"])), - encrypted_api_key, - Some(encrypted_auth_config), - None, - None, - None, - None, - None, - None, - ) - .expect("key transport should build"); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![StoredProviderCatalogProvider::new( @@ -2234,7 +2218,7 @@ async fn gateway_reports_codex_quota_runtime_failures_locally_without_falling_ba assert!(payload["results"][0]["message"] .as_str() .expect("message should be string") - .contains("wham/usage 请求执行失败: execution runtime returned HTTP 500")); + .contains("wham/usage 请求执行失败")); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); let reloaded = provider_catalog_repository @@ -2264,6 +2248,8 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru struct SeenExecutionRuntimeRequest { url: String, authorization: String, + user_agent: String, + x_client_version: String, provider_api_format: String, request_body: Option, } @@ -2281,7 +2267,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru }), ); - let seen_execution_runtime = Arc::new(Mutex::new(None::)); + let seen_execution_runtime = Arc::new(Mutex::new(Vec::::new())); let seen_execution_runtime_clone = Arc::clone(&seen_execution_runtime); let execution_runtime = Router::new().route( "/v1/execute/sync", @@ -2294,26 +2280,30 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru .expect("body should read"), ) .expect("plan should parse"); - *seen_execution_runtime_inner + let request_body = plan.body.json_body.clone(); + seen_execution_runtime_inner .lock() - .expect("mutex should lock") = Some(SeenExecutionRuntimeRequest { - url: plan.url.clone(), - authorization: plan - .headers - .get("authorization") - .cloned() - .unwrap_or_default(), - provider_api_format: plan.provider_api_format.clone(), - request_body: plan.body.json_body.clone(), - }); - let result = aether_contracts::ExecutionResult { - request_id: plan.request_id, - candidate_id: None, - status_code: 200, - headers: BTreeMap::new(), - response_observation: None, - body: Some(aether_contracts::ResponseBody { - json_body: Some(json!({ + .expect("mutex should lock") + .push(SeenExecutionRuntimeRequest { + url: plan.url.clone(), + authorization: plan + .headers + .get("authorization") + .cloned() + .unwrap_or_default(), + user_agent: plan.headers.get("user-agent").cloned().unwrap_or_default(), + x_client_version: plan + .headers + .get("x-client-version") + .cloned() + .unwrap_or_default(), + provider_api_format: plan.provider_api_format.clone(), + request_body: request_body.clone(), + }); + let (status_code, json_body) = match plan.provider_api_format.as_str() { + "antigravity:fetch_available_models" => ( + 200, + json!({ "models": { "claude-sonnet-4": { "displayName": "Claude Sonnet 4", @@ -2324,9 +2314,55 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru }, "gemini-2.5-pro": { "displayName": "Gemini 2.5 Pro" + }, + "gemini-3.7-flash-tiered": { + "displayName": "Gemini 3.7 Flash" + }, + "chat_23310": { + "displayName": "Internal Chat" } } - })), + }), + ), + "antigravity:retrieve_user_quota_summary" + if request_body + .as_ref() + .and_then(|body| body.get("project")) + .is_some() => + { + (403, json!({"error": {"message": "project not accepted"}})) + } + "antigravity:retrieve_user_quota_summary" => ( + 200, + json!({ + "groups": [{ + "displayName": "Claude and GPT models", + "description": "Shared quota", + "buckets": [{ + "bucketId": "3p-5h", + "window": "5h", + "remainingFraction": 0.25, + "resetTime": "2026-05-05T05:00:00Z", + "displayName": "5 hour" + }, { + "bucketId": "3p-weekly", + "window": "weekly", + "remainingFraction": 0.8, + "resetTime": "2026-05-11T00:00:00Z" + }] + }] + }), + ), + unexpected => panic!("unexpected quota request format: {unexpected}"), + }; + let result = aether_contracts::ExecutionResult { + request_id: plan.request_id, + candidate_id: None, + status_code, + headers: BTreeMap::new(), + response_observation: None, + body: Some(aether_contracts::ResponseBody { + json_body: Some(json_body), body_bytes_b64: None, }), telemetry: None, @@ -2385,6 +2421,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru )], vec![key], )); + let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::default()); let (upstream_url, upstream_handle) = start_server(upstream).await; let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; @@ -2394,6 +2431,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog_repository.clone(), ) + .with_global_model_repository_for_tests(global_model_repository.clone()) .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ), ); @@ -2429,31 +2467,55 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru payload["results"][0]["quota_snapshot"]["windows"] .as_array() .map(Vec::len), - Some(1usize) + Some(3usize) ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); - let seen_execution_runtime_request = seen_execution_runtime + let seen_execution_runtime_requests = seen_execution_runtime .lock() .expect("mutex should lock") - .clone() - .expect("execution runtime request should be captured"); + .clone(); + assert_eq!(seen_execution_runtime_requests.len(), 3); + let fetch_models_request = &seen_execution_runtime_requests[0]; assert_eq!( - seen_execution_runtime_request.url, + fetch_models_request.url, "https://daily-cloudcode-pa.googleapis.com/v1internal:fetchAvailableModels" ); + assert_eq!(fetch_models_request.authorization, "Bearer ya29.ant-token"); assert_eq!( - seen_execution_runtime_request.authorization, - "Bearer ya29.ant-token" + fetch_models_request.user_agent, + "vscode/1.X.X (Antigravity/4.3.0)" ); + assert_eq!(fetch_models_request.x_client_version, "4.3.0"); assert_eq!( - seen_execution_runtime_request.provider_api_format, + fetch_models_request.provider_api_format, "antigravity:fetch_available_models" ); assert_eq!( - seen_execution_runtime_request.request_body, + fetch_models_request.request_body, Some(json!({ "project": "project-ant-123" })) ); + let grouped_with_project = &seen_execution_runtime_requests[1]; + let grouped_without_project = &seen_execution_runtime_requests[2]; + assert_eq!( + grouped_with_project.url, + "https://daily-cloudcode-pa.googleapis.com/v1internal:retrieveUserQuotaSummary" + ); + assert_eq!( + grouped_with_project.provider_api_format, + "antigravity:retrieve_user_quota_summary" + ); + assert_eq!( + grouped_with_project.request_body, + Some(json!({"project": "project-ant-123"})) + ); + assert_eq!(grouped_without_project.url, grouped_with_project.url); + assert_eq!(grouped_without_project.request_body, Some(json!({}))); + assert!(seen_execution_runtime_requests.iter().all(|request| { + request.authorization == "Bearer ya29.ant-token" + && request.user_agent == "vscode/1.X.X (Antigravity/4.3.0)" + && request.x_client_version == "4.3.0" + })); let reloaded = provider_catalog_repository .list_keys_by_ids(&["key-antigravity-a".to_string()]) @@ -2466,21 +2528,58 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru .upstream_metadata .as_ref() .and_then(|value| value.get("antigravity")) - .and_then(|value| value.get("models")) + .and_then(|value| value.get("quota_by_model")) .and_then(|value| value.get("claude-sonnet-4")) .and_then(|value| value.get("remaining_fraction")), Some(&json!(0.25)) ); + let imported_provider_models = global_model_repository + .list_admin_provider_models(&AdminProviderModelListQuery { + provider_id: "provider-antigravity".to_string(), + is_active: None, + offset: 0, + limit: 100, + }) + .await + .expect("imported Antigravity provider models should read"); + let imported_model_names = imported_provider_models + .iter() + .map(|model| model.provider_model_name.as_str()) + .collect::>(); + assert!(imported_model_names.contains("claude-sonnet-4")); + assert!(imported_model_names.contains("gemini-2.5-pro")); + assert!(imported_model_names.contains("gemini-3.7-flash-tiered")); + assert!(!imported_model_names.contains("chat_23310")); assert_eq!( reloaded[0] .upstream_metadata .as_ref() .and_then(|value| value.get("antigravity")) - .and_then(|value| value.get("models")) + .and_then(|value| value.get("quota_by_model")) .and_then(|value| value.get("claude-sonnet-4")) .and_then(|value| value.get("used_percent")), Some(&json!(75.0)) ); + assert_eq!( + reloaded[0] + .upstream_metadata + .as_ref() + .and_then(|value| value.pointer("/antigravity/quota_groups/0/buckets/0/bucket_id")), + Some(&json!("3p-5h")) + ); + assert_eq!( + reloaded[0] + .upstream_metadata + .as_ref() + .and_then(|value| value.pointer("/antigravity/project_id")), + Some(&json!("project-ant-123")) + ); + assert!(reloaded[0] + .upstream_metadata + .as_ref() + .and_then(|value| value.pointer("/antigravity/quota_groups_updated_at")) + .and_then(serde_json::Value::as_u64) + .is_some()); assert_eq!( reloaded[0] .status_snapshot @@ -2505,7 +2604,19 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru .and_then(|value| value.get("windows")) .and_then(|value| value.as_array()) .map(Vec::len), - Some(1usize) + Some(3usize) + ); + assert_eq!( + reloaded[0] + .status_snapshot + .as_ref() + .and_then(|value| value.pointer("/quota/windows")) + .and_then(serde_json::Value::as_array) + .and_then(|windows| windows + .iter() + .find(|window| window["code"] == "group:0:3p-5h")) + .and_then(|window| window.get("remaining_ratio")), + Some(&json!(0.25)) ); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/control/admin/endpoints/routes.rs b/apps/aether-gateway/src/tests/control/admin/endpoints/routes.rs index 83e4c9517..a97d67492 100644 --- a/apps/aether-gateway/src/tests/control/admin/endpoints/routes.rs +++ b/apps/aether-gateway/src/tests/control/admin/endpoints/routes.rs @@ -9,7 +9,8 @@ use http::StatusCode; use serde_json::json; use super::super::super::{ - build_router_with_state, sample_endpoint, sample_key, sample_provider, start_server, AppState, + build_router_with_state, sample_bound_key as sample_key, sample_endpoint, sample_provider, + start_server, AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -541,7 +542,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin assert_eq!(payload["max_retries"], 5); assert_eq!(payload["total_keys"], 0); assert_eq!(payload["active_keys"], 0); - assert_eq!(payload["proxy"]["url"], "http://proxy.internal"); + assert_eq!(payload["proxy"]["url"], "http://proxy.internal/"); assert_eq!(payload["proxy"]["password"], "***"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -741,8 +742,8 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin assert_eq!(payload["is_active"], false); assert_eq!(payload["total_keys"], 1); assert_eq!(payload["active_keys"], 1); - assert_eq!(payload["proxy"]["url"], "http://proxy-2.internal"); - assert_eq!(payload["proxy"]["password"], "***"); + assert_eq!(payload["proxy"]["url"], "http://proxy-2.internal/"); + assert_eq!(payload["proxy"]["password"], serde_json::Value::Null); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); let endpoints = provider_catalog_repository @@ -757,7 +758,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin assert_eq!(endpoints[0].config, Some(json!({"foo":"new"}))); assert_eq!( endpoints[0].proxy, - Some(json!({"url":"http://proxy-2.internal","password":"secret"})) + Some(json!({"url":"http://proxy-2.internal/"})) ); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/control/admin/health_access.rs b/apps/aether-gateway/src/tests/control/admin/health_access.rs index 521782a9b..a42db2a2e 100644 --- a/apps/aether-gateway/src/tests/control/admin/health_access.rs +++ b/apps/aether-gateway/src/tests/control/admin/health_access.rs @@ -4,7 +4,9 @@ use std::time::{SystemTime, UNIX_EPOCH}; use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::auth_modules::InMemoryAuthModuleReadRepository; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; -use aether_data::repository::management_tokens::InMemoryManagementTokenRepository; +use aether_data::repository::management_tokens::{ + InMemoryManagementTokenRepository, ManagementTokenListQuery, ManagementTokenReadRepository, +}; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data_contracts::repository::candidates::RequestCandidateStatus; use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository; @@ -15,9 +17,10 @@ use http::StatusCode; use serde_json::json; use super::super::{ - build_router_with_state, hash_management_token, issue_test_admin_access_token, sample_endpoint, - sample_key, sample_ldap_module_config, sample_management_token, sample_oauth_module_provider, - sample_provider, sample_request_candidate, start_server, AppState, + build_router_with_state, hash_management_token, issue_test_admin_access_token, + sample_bound_key, sample_endpoint, sample_ldap_module_config, sample_management_token, + sample_oauth_module_provider, sample_provider, sample_request_candidate, start_server, + AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -96,7 +99,7 @@ async fn gateway_handles_admin_health_api_formats_locally_with_trusted_admin_pri "openai:chat", "https://api.openai.example", )], - vec![sample_key( + vec![sample_bound_key( "key-openai", "provider-openai", "openai:chat", @@ -248,7 +251,7 @@ async fn gateway_handles_admin_health_summary_locally_with_trusted_admin_princip .with_health_score(0.2), ], vec![ - sample_key( + sample_bound_key( "key-openai-active", "provider-openai", "openai:chat", @@ -258,7 +261,7 @@ async fn gateway_handles_admin_health_summary_locally_with_trusted_admin_princip Some(json!({"openai:chat": {"health_score": 0.9}})), Some(json!({"openai:chat": {"open": false}})), ), - sample_key( + sample_bound_key( "key-openai-circuit", "provider-openai", "openai:chat", @@ -334,7 +337,7 @@ async fn gateway_handles_admin_key_health_locally_with_trusted_admin_principal() "https://api.openai.example", )], vec![ - sample_key("key-openai", "provider-openai", "openai:chat", "sk-test") + sample_bound_key("key-openai", "provider-openai", "openai:chat", "sk-test") .with_rate_limit_fields(None, None, None, None, None, None, None, Some(10), Some(7)) .with_usage_fields(Some(3), Some(2100)) .with_health_fields( @@ -432,7 +435,7 @@ async fn gateway_admin_key_health_summary_treats_expired_unix_circuit_as_closed( "https://api.openai.example", )], vec![ - sample_key("key-openai", "provider-openai", "openai:chat", "sk-test") + sample_bound_key("key-openai", "provider-openai", "openai:chat", "sk-test") .with_health_fields( Some(json!({"openai:chat": { "health_score": 0.7, @@ -511,7 +514,7 @@ async fn gateway_recovers_admin_key_health_locally_with_trusted_admin_principal( "https://api.openai.example", )], vec![ - sample_key("key-openai", "provider-openai", "openai:chat", "sk-test") + sample_bound_key("key-openai", "provider-openai", "openai:chat", "sk-test") .with_health_fields( Some(json!({"openai:chat": { "health_score": 0.2, @@ -620,7 +623,7 @@ async fn gateway_recovers_all_admin_key_health_locally_with_trusted_admin_princi "https://api.openai.example", )], vec![ - sample_key( + sample_bound_key( "key-openai-circuit", "provider-openai", "openai:chat", @@ -630,7 +633,7 @@ async fn gateway_recovers_all_admin_key_health_locally_with_trusted_admin_princi Some(json!({"openai:chat": {"health_score": 0.3}})), Some(json!({"openai:chat": {"open": true}})), ), - sample_key( + sample_bound_key( "key-openai-healthy", "provider-openai", "openai:chat", @@ -729,7 +732,7 @@ async fn gateway_handles_admin_health_status_locally_with_trusted_admin_principa "openai:chat", "https://api.openai.example", )], - vec![sample_key( + vec![sample_bound_key( "key-openai", "provider-openai", "openai:chat", @@ -1308,11 +1311,11 @@ async fn gateway_handles_admin_management_tokens_locally_with_trusted_admin_prin } #[tokio::test] -async fn gateway_allows_full_management_token_to_fetch_permission_catalog() { +async fn gateway_rejects_constrained_full_management_token_creating_unconstrained_child() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( - "/api/admin/management-tokens/permissions/catalog", + "/api/admin/management-tokens", any(move |_request: Request| { let upstream_hits_inner = Arc::clone(&upstream_hits_clone); async move { @@ -1341,38 +1344,116 @@ async fn gateway_allows_full_management_token_to_fetch_permission_catalog() { let raw_token = "ae-management-full-access"; let mut management_token = sample_management_token("mt-admin-full", &admin_user.id, "management-full", true); - management_token.token.allowed_ips = None; + management_token.token.allowed_ips = Some(json!(["127.0.0.1"])); + management_token.token.expires_at_unix_secs = Some( + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("clock should be after epoch") + .as_secs() + + 3_600, + ); management_token.token.permissions = Some(json!(all_assignable_management_token_permissions())); + let legacy_raw_token = "ae-management-legacy-full-access"; + let mut legacy_management_token = sample_management_token( + "mt-admin-legacy-full", + &admin_user.id, + "management-legacy-full", + true, + ); + legacy_management_token.token.allowed_ips = Some(json!(["127.0.0.1"])); + legacy_management_token.token.expires_at_unix_secs = + management_token.token.expires_at_unix_secs; + legacy_management_token.token.permissions = None; let management_token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( - vec![management_token], - vec![( - hash_management_token(raw_token), - "mt-admin-full".to_string(), - )], + vec![management_token, legacy_management_token], + vec![ + ( + hash_management_token(raw_token), + "mt-admin-full".to_string(), + ), + ( + hash_management_token(legacy_raw_token), + "mt-admin-legacy-full".to_string(), + ), + ], )); let (upstream_url, upstream_handle) = start_server(upstream).await; let gateway = build_router_with_state(state.with_data_state_for_tests( - GatewayDataState::with_management_token_repository_for_tests(management_token_repository), + GatewayDataState::with_management_token_repository_for_tests(Arc::clone( + &management_token_repository, + )), )); let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() - .get(format!( - "{gateway_url}/api/admin/management-tokens/permissions/catalog" - )) + .post(format!("{gateway_url}/api/admin/management-tokens")) .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") .bearer_auth(raw_token) + .json(&json!({ + "name": "unconstrained-child", + "permissions": all_assignable_management_token_permissions(), + })) .send() .await .expect("request should succeed"); let status = response.status(); let body = response.text().await.expect("body should read"); - assert_eq!(status, StatusCode::OK, "body={body}"); + assert_eq!(status, StatusCode::FORBIDDEN, "body={body}"); let payload: serde_json::Value = serde_json::from_str(&body).expect("json body should parse"); - assert!(payload["items"].is_array()); + assert_eq!(payload["detail"], "management token permission denied"); + assert_eq!( + payload["required_permission"], + "admin:management_tokens:admin" + ); + + let legacy_response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/management-tokens")) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(legacy_raw_token) + .json(&json!({ + "name": "unconstrained-legacy-child", + "permissions": all_assignable_management_token_permissions(), + })) + .send() + .await + .expect("legacy request should succeed"); + let legacy_status = legacy_response.status(); + let legacy_body = legacy_response.text().await.expect("body should read"); + assert_eq!(legacy_status, StatusCode::FORBIDDEN, "body={legacy_body}"); + let legacy_payload: serde_json::Value = + serde_json::from_str(&legacy_body).expect("json body should parse"); + assert_eq!( + legacy_payload["detail"], + "management token permission denied" + ); + assert_eq!( + legacy_payload["required_permission"], + "admin:management_tokens:admin" + ); + + let tokens = management_token_repository + .list_management_tokens(&ManagementTokenListQuery { + user_id: None, + is_active: None, + offset: 0, + limit: 10, + }) + .await + .expect("management token list should succeed"); + assert_eq!( + tokens.total, 2, + "unconstrained children must not be created" + ); + let mut token_ids = tokens + .items + .iter() + .map(|item| item.token.id.as_str()) + .collect::>(); + token_ids.sort_unstable(); + assert_eq!(token_ids, vec!["mt-admin-full", "mt-admin-legacy-full"]); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/control/admin/ldap.rs b/apps/aether-gateway/src/tests/control/admin/ldap.rs index 2e6a2d977..4826a8796 100644 --- a/apps/aether-gateway/src/tests/control/admin/ldap.rs +++ b/apps/aether-gateway/src/tests/control/admin/ldap.rs @@ -223,3 +223,69 @@ async fn gateway_tests_admin_ldap_connection_locally_with_trusted_admin_principa gateway_handle.abort(); upstream_handle.abort(); } + +#[tokio::test] +async fn gateway_rejects_ldap_filter_and_attribute_injection_on_admin_update() { + let auth_module_repository = Arc::new(InMemoryAuthModuleReadRepository::seed( + Vec::::new(), + Some(sample_ldap_module_config()), + )); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_auth_module_repository_for_tests( + auth_module_repository, + )), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for (search_filter, username_attr, expected_detail) in [ + ( + "(uid={username})(objectClass=*)", + "uid", + "搜索过滤器格式无效", + ), + ( + "(uid={username})", + "uid)(|(objectClass=*)", + "用户名属性必须是有效的 LDAP 属性名称", + ), + ] { + let response = client + .put(format!("{gateway_url}/api/admin/ldap/config")) + .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!({ + "server_url": "mockldap://ldap.internal.example.com", + "bind_dn": "cn=svc,dc=example,dc=com", + "bind_password": "secret123", + "base_dn": "ou=people,dc=example,dc=com", + "user_search_filter": search_filter, + "username_attr": username_attr, + "email_attr": "mail", + "display_name_attr": "cn", + "is_enabled": true, + "is_exclusive": false, + "use_starttls": false, + "connect_timeout": 20 + })) + .send() + .await + .expect("invalid LDAP update should complete locally"); + + let status = response.status(); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::BAD_REQUEST, "payload={payload}"); + assert!( + payload["detail"] + .as_str() + .is_some_and(|detail| detail.contains(expected_detail)), + "payload={payload}" + ); + } + + gateway_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/control/admin/models/external.rs b/apps/aether-gateway/src/tests/control/admin/models/external.rs index 710729b19..e47acde49 100644 --- a/apps/aether-gateway/src/tests/control/admin/models/external.rs +++ b/apps/aether-gateway/src/tests/control/admin/models/external.rs @@ -1,6 +1,7 @@ use std::sync::{Arc, Mutex}; use std::time::Duration; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use axum::body::Body; use axum::extract::ws::Message; @@ -11,7 +12,10 @@ use http::StatusCode; use serde_json::json; use tokio::sync::watch; -use super::super::super::{build_router_with_state, sample_proxy_node, start_server, AppState}; +use super::super::super::{ + build_router_with_state, sample_proxy_node, start_server, with_tunnel_control_plane_key, + AppState, TUNNEL_CONTROL_PLANE_TEST_GENERATION, TUNNEL_CONTROL_PLANE_TEST_PSK, +}; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, TRUSTED_ADMIN_USER_ROLE_HEADER, @@ -314,7 +318,10 @@ async fn gateway_fetches_external_models_through_connected_tunnel_node() { let source_url = "https://models.dev.test/api.json"; let _guard = set_admin_external_models_source_url_for_tests(source_url); - let mut tunnel_node = sample_proxy_node("tunnel-node"); + let mut tunnel_node = with_tunnel_control_plane_key( + sample_proxy_node("tunnel-node"), + TUNNEL_CONTROL_PLANE_TEST_PSK, + ); tunnel_node.name = "Tunnel Node".to_string(); tunnel_node.status = "online".to_string(); tunnel_node.tunnel_mode = true; @@ -322,6 +329,7 @@ async fn gateway_fetches_external_models_through_connected_tunnel_node() { tunnel_node.remote_config = None; let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![tunnel_node])); let data_state = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) .with_system_config_values_for_tests([( "external_models_proxy_node_id".to_string(), json!("tunnel-node"), @@ -332,9 +340,8 @@ async fn gateway_fetches_external_models_through_connected_tunnel_node() { let tunnel_state = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_state - .hub - .register_proxy(Arc::new(TunnelProxyConn::new( + tunnel_state.hub.register_proxy(Arc::new( + TunnelProxyConn::new( 700, "tunnel-node".to_string(), "Tunnel Node".to_string(), @@ -342,7 +349,10 @@ async fn gateway_fetches_external_models_through_connected_tunnel_node() { proxy_close_tx, 16, 2, - ))); + ) + .with_tunnel_generation(TUNNEL_CONTROL_PLANE_TEST_GENERATION.to_string()) + .with_authenticated_key(TUNNEL_CONTROL_PLANE_TEST_PSK.to_string()), + )); let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -371,7 +381,7 @@ async fn gateway_fetches_external_models_through_connected_tunnel_node() { serde_json::from_slice(&meta_payload).expect("request meta should parse"); assert_eq!(meta.method, "GET"); assert_eq!(meta.url, source_url); - assert_eq!(meta.follow_redirects, Some(true)); + assert_eq!(meta.follow_redirects, Some(false)); assert_eq!( meta.headers.get("accept").map(String::as_str), Some("application/json") diff --git a/apps/aether-gateway/src/tests/control/admin/models/global.rs b/apps/aether-gateway/src/tests/control/admin/models/global.rs index 38027db47..fb03012e0 100644 --- a/apps/aether-gateway/src/tests/control/admin/models/global.rs +++ b/apps/aether-gateway/src/tests/control/admin/models/global.rs @@ -13,7 +13,7 @@ use serde_json::json; use super::super::super::{ build_router_with_state, issue_test_admin_access_token, sample_admin_global_model, - sample_admin_provider_model, sample_endpoint, sample_key, sample_provider, start_server, + sample_admin_provider_model, sample_bound_key, sample_endpoint, sample_provider, start_server, AppState, }; use crate::constants::{ @@ -772,7 +772,7 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p ), ], { - let mut primary_key = sample_key( + let mut primary_key = sample_bound_key( "key-openai-routing", "provider-openai", "openai:chat", @@ -792,7 +792,7 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p "openai:chat": {"open": true, "next_probe_at": "2099-03-27T15:00:00Z"} })); - let mut mapped_key = sample_key( + let mut mapped_key = sample_bound_key( "key-alt-routing", "provider-alt", "openai:chat", @@ -804,7 +804,7 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p mapped_key.allowed_models = Some(json!(["gpt-5-upstream"])); mapped_key.rpm_limit = Some(120); - let mut unlinked_key = sample_key( + let mut unlinked_key = sample_bound_key( "key-unlinked-routing", "provider-unlinked", "openai:chat", @@ -850,10 +850,7 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p provider_catalog_repository, ) .with_global_model_repository_for_tests(global_model_repository) - .with_system_config_values_for_tests(vec![ - ("scheduling_mode".to_string(), json!("fixed_order")), - ("provider_priority_mode".to_string(), json!("global_key")), - ]), + .with_system_default_routing_group_for_tests(), ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -876,8 +873,8 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p assert_eq!(payload["global_model_name"], "gpt-5"); assert_eq!(payload["display_name"], "GPT 5"); assert_eq!(payload["global_model_mappings"], json!(["gpt-5-upstream"])); - assert_eq!(payload["scheduling_mode"], "fixed_order"); - assert_eq!(payload["priority_mode"], "global_key"); + assert_eq!(payload["scheduling_mode"], "cache_affinity"); + assert_eq!(payload["priority_mode"], "provider"); assert_eq!(payload["total_providers"], 2); assert_eq!(payload["active_providers"], 2); @@ -898,7 +895,7 @@ async fn gateway_handles_admin_global_model_routing_locally_with_trusted_admin_p let openai_keys = openai_endpoints[0]["keys"].as_array().expect("keys array"); assert_eq!(openai_keys.len(), 1); assert_eq!(openai_keys[0]["name"], "primary"); - assert_eq!(openai_keys[0]["masked_key"], "sk-opena***1234"); + assert_eq!(openai_keys[0]["masked_key"], "sk-open***1234"); assert_eq!(openai_keys[0]["is_adaptive"], true); assert_eq!(openai_keys[0]["effective_rpm"], 77); assert_eq!(openai_keys[0]["allowed_models"], json!(["gpt-5"])); @@ -955,7 +952,7 @@ async fn gateway_global_model_routing_counts_image_provider_keys_by_provider_mod image_provider.provider_type = "chatgpt_web".to_string(); let grok_provider = sample_provider("provider-grok", "grok2api", 20); - let mut image_key = sample_key( + let mut image_key = sample_bound_key( "key-image-routing", "provider-image", "legacy:mismatch", @@ -965,7 +962,7 @@ async fn gateway_global_model_routing_counts_image_provider_keys_by_provider_mod image_key.auth_type = "oauth".to_string(); image_key.allowed_models = Some(json!(["gpt-image-2"])); - let mut grok_key = sample_key( + let mut grok_key = sample_bound_key( "key-grok-routing", "provider-grok", "openai:chat", diff --git a/apps/aether-gateway/src/tests/control/admin/monitoring.rs b/apps/aether-gateway/src/tests/control/admin/monitoring.rs index 732bb9441..d09847f23 100644 --- a/apps/aether-gateway/src/tests/control/admin/monitoring.rs +++ b/apps/aether-gateway/src/tests/control/admin/monitoring.rs @@ -528,7 +528,7 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm assert_eq!(payload["candidates"][0]["provider_name"], json!("OpenAI")); assert_eq!( payload["candidates"][0]["provider_website"], - json!("https://openai.com") + json!("https://openai.com/") ); assert_eq!( payload["candidates"][0]["endpoint_name"], @@ -761,11 +761,12 @@ async fn gateway_handles_admin_monitoring_cache_affinities_locally_with_trusted_ AppState::new() .expect("gateway should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_provider_catalog_reader_for_tests( + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog, ) .with_user_reader(user_repository) - .with_auth_api_key_reader(auth_repository), + .with_auth_api_key_reader(auth_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_admin_monitoring_cache_affinity_entry_for_tests( "cache_affinity:user-key-1:openai:model-alpha", diff --git a/apps/aether-gateway/src/tests/control/admin/oauth.rs b/apps/aether-gateway/src/tests/control/admin/oauth.rs index 8f330ca32..b15aef987 100644 --- a/apps/aether-gateway/src/tests/control/admin/oauth.rs +++ b/apps/aether-gateway/src/tests/control/admin/oauth.rs @@ -3,9 +3,7 @@ use std::sync::{Arc, Mutex}; use aether_contracts::{ ExecutionPlan, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, }; -use aether_crypto::{ - decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, -}; +use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::background_tasks::InMemoryBackgroundTaskRepository; use aether_data::repository::management_tokens::{ InMemoryManagementTokenRepository, ManagementTokenReadRepository, @@ -16,12 +14,14 @@ use aether_data::repository::oauth_providers::{ use aether_data::repository::pool_scores::InMemoryPoolMemberScoreRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; +use aether_data::repository::users::{InMemoryUserReadRepository, UserReadRepository}; use aether_data_contracts::repository::background_tasks::BackgroundTaskReadRepository; use aether_data_contracts::repository::pool_scores::{ GetPoolMemberScoresByIdsQuery, PoolMemberHardState, PoolMemberIdentity, PoolScoreReadRepository, }; use aether_data_contracts::repository::provider_catalog::{ - ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogReadRepository, + ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, }; use axum::body::{to_bytes, Body, Bytes}; use axum::response::{IntoResponse, Response}; @@ -32,8 +32,8 @@ use serde_json::{json, Value}; use super::super::{ build_router_with_state, build_state_with_execution_runtime_override, hash_management_token, - sample_endpoint, sample_key, sample_management_token, sample_oauth_provider_config, - sample_provider, sample_proxy_node, start_server, AppState, + sample_bound_key, sample_endpoint, sample_key, sample_management_token, + sample_oauth_provider_config, sample_provider, sample_proxy_node, start_server, AppState, }; use crate::admin_api::{ maybe_build_local_admin_provider_oauth_response, AdminAppState, AdminRequestContext, @@ -73,6 +73,28 @@ where } } +fn decrypt_persisted_provider_api_key(key: &StoredProviderCatalogKey) -> String { + AppState::new() + .expect("credential test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .decrypt_provider_catalog_key_api_key(key) + .expect("persisted api key should decrypt") + .expect("persisted api key should exist") +} + +fn decrypt_persisted_provider_auth_config(key: &StoredProviderCatalogKey) -> String { + AppState::new() + .expect("credential test state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .decrypt_provider_catalog_key_auth_config(key) + .expect("persisted auth config should decrypt") + .expect("persisted auth config should exist") +} + fn trusted_admin_headers() -> HeaderMap { let mut headers = HeaderMap::new(); headers.insert(GATEWAY_HEADER, HeaderValue::from_static("rust-phase3b")); @@ -429,6 +451,11 @@ async fn gateway_authorizes_claude_cookie_without_persisting_cookie_impl() { vec![endpoint], vec![], )); + let mut tunnel_node = sample_proxy_node("proxy-node-claude"); + tunnel_node.status = "online".to_string(); + tunnel_node.tunnel_mode = true; + tunnel_node.tunnel_connected = true; + let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![tunnel_node])); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let state = build_state_with_execution_runtime_override(execution_runtime_url) @@ -436,6 +463,17 @@ async fn gateway_authorizes_claude_cookie_without_persisting_cookie_impl() { GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog_repository.clone(), ) + .attach_proxy_node_repository_for_tests(proxy_node_repository) + .with_system_config_values_for_tests(vec![( + "tunnel.attachments.proxy-node-claude".to_string(), + json!({ + "gateway_instance_id": "gateway-owner", + "relay_base_url": "http://gateway-owner.internal", + "tunnel_generation": "test-generation-1", + "conn_count": 1, + "observed_at_unix_secs": 4_102_444_800u64, + }), + )]) .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_provider_oauth_token_url_for_tests( @@ -475,23 +513,9 @@ async fn gateway_authorizes_claude_cookie_without_persisting_cookie_impl() { .await .expect("keys should load"); assert_eq!(keys.len(), 1); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_api_key - .as_deref() - .expect("api key should be encrypted"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&keys[0]); assert_eq!(decrypted_api_key, "sk-ant-oat01-created"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_auth_config - .as_deref() - .expect("auth config should be encrypted"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&keys[0]); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["refresh_token"], "sk-ant-ort01-created"); @@ -638,12 +662,28 @@ async fn gateway_batch_authorizes_claude_cookies_as_redacted_task_impl() { vec![], )); let background_task_repository = Arc::new(InMemoryBackgroundTaskRepository::default()); + let mut tunnel_node = sample_proxy_node("proxy-node-claude"); + tunnel_node.status = "online".to_string(); + tunnel_node.tunnel_mode = true; + tunnel_node.tunnel_connected = true; + let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![tunnel_node])); let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let state = build_state_with_execution_runtime_override(execution_runtime_url) .with_data_state_for_tests( GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog_repository.clone(), ) + .attach_proxy_node_repository_for_tests(proxy_node_repository) + .with_system_config_values_for_tests(vec![( + "tunnel.attachments.proxy-node-claude".to_string(), + json!({ + "gateway_instance_id": "gateway-owner", + "relay_base_url": "http://gateway-owner.internal", + "tunnel_generation": "test-generation-1", + "conn_count": 1, + "observed_at_unix_secs": 4_102_444_800u64, + }), + )]) .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) .with_background_task_repository_for_tests(background_task_repository.clone()), ) @@ -741,6 +781,20 @@ async fn gateway_batch_authorizes_claude_cookies_as_redacted_task_impl() { "raw task state contains {forbidden}" ); } + let persisted_task_state = state + .runtime_kv_get(format!("provider_oauth_batch_task:{task_id}").as_str()) + .await + .expect("runtime task state lookup should succeed") + .expect("runtime task state should exist"); + assert!(crate::handlers::shared::runtime_secret_payload_is_sealed( + &persisted_task_state + )); + for forbidden in ["sessionKey", "batch-sid-one", "batch-sid-two"] { + assert!( + !persisted_task_state.contains(forbidden), + "persisted task state contains {forbidden}" + ); + } let background_run = background_task_repository .find_run(&task_id) .await @@ -761,13 +815,7 @@ async fn gateway_batch_authorizes_claude_cookies_as_redacted_task_impl() { .expect("keys should load"); assert_eq!(keys.len(), 2); for key in &keys { - let auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - key.encrypted_auth_config - .as_deref() - .expect("auth config should be encrypted"), - ) - .expect("auth config should decrypt"); + let auth_config = decrypt_persisted_provider_auth_config(key); assert!(!auth_config.contains("batch-sid")); assert!(!auth_config.contains("sessionKey")); } @@ -1070,23 +1118,9 @@ async fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_ persisted.proxy, Some(json!({"node_id": "proxy-node-windsurf", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted); assert_eq!(decrypted_api_key, "devin-session-token$registered"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + 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"], "windsurf"); @@ -1551,6 +1585,10 @@ async fn gateway_handles_admin_provider_oauth_device_poll_locally_with_trusted_a .header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123") .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") + .header( + TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER, + "management-token-123", + ) .json(&json!({ "session_id": "session-123" })) @@ -1593,23 +1631,9 @@ async fn gateway_handles_admin_provider_oauth_device_poll_locally_with_trusted_a persisted.proxy, Some(json!({"node_id": "proxy-node-kiro", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted); assert_eq!(decrypted_api_key, expected_access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + 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"], "kiro"); @@ -1729,6 +1753,10 @@ async fn gateway_handles_admin_provider_oauth_device_poll_for_kiro_social_callba .header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123") .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") + .header( + TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER, + "management-token-123", + ) .json(&json!({ "session_id": "session-social", "callback_url": "http://localhost:49153/signin/callback?login_option=github&code=social-code-123&state=session-social" @@ -1780,23 +1808,9 @@ async fn gateway_handles_admin_provider_oauth_device_poll_for_kiro_social_callba .expect("persisted key should exist"); assert_eq!(persisted.name, "social@example.com (Github)"); assert_eq!(persisted.auth_type, "oauth"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted); assert_eq!(decrypted_api_key, expected_access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + 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"], "kiro"); @@ -2053,6 +2067,10 @@ async fn gateway_revalidates_kiro_device_poll_via_idc_refresh_and_backfills_emai .header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123") .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") + .header( + TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER, + "management-token-123", + ) .json(&json!({ "session_id": "session-refresh-email" })) @@ -2111,23 +2129,9 @@ async fn gateway_revalidates_kiro_device_poll_via_idc_refresh_and_backfills_emai persisted.proxy, Some(json!({"node_id": "proxy-node-kiro", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted); assert_eq!(decrypted_api_key, expected_refreshed_access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + 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"], "kiro"); @@ -2459,7 +2463,7 @@ async fn gateway_handles_admin_provider_oauth_start_key_locally_with_trusted_adm let mut provider = sample_provider("provider-codex", "codex", 10); provider.provider_type = "codex".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-oauth", "provider-codex", "openai:chat", @@ -2504,6 +2508,14 @@ async fn gateway_handles_admin_provider_oauth_start_key_locally_with_trusted_adm .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("no-store")) + ); + assert_eq!( + response.headers().get(http::header::PRAGMA), + Some(&HeaderValue::from_static("no-cache")) + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["provider_type"], "codex"); assert_eq!( @@ -2574,6 +2586,14 @@ async fn gateway_handles_admin_provider_oauth_start_provider_locally_with_truste .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response.headers().get(http::header::CACHE_CONTROL), + Some(&HeaderValue::from_static("no-store")) + ); + assert_eq!( + response.headers().get(http::header::PRAGMA), + Some(&HeaderValue::from_static("no-cache")) + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["provider_type"], "codex"); assert_eq!( @@ -2668,7 +2688,7 @@ async fn gateway_handles_admin_provider_oauth_batch_import_task_status_locally_w assert_eq!(payload["status"], "completed"); assert_eq!(payload["total"], 2); assert_eq!(payload["processed"], 2); - assert_eq!(payload["success"], 1); + assert_eq!(payload["success"], 1, "payload={payload}"); assert_eq!(payload["failed"], 1); assert_eq!(payload["created_count"], 0); assert_eq!(payload["replaced_count"], 1); @@ -2906,7 +2926,7 @@ async fn gateway_batch_imports_admin_provider_oauth_locally_with_trusted_admin_p 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_eq!(payload["success"], 1); + assert_eq!(payload["success"], 1, "payload={payload}"); assert_eq!(payload["failed"], 1); let results = payload["results"] .as_array() @@ -2936,14 +2956,7 @@ async fn gateway_batch_imports_admin_provider_oauth_locally_with_trusted_admin_p persisted.proxy, Some(json!({"node_id": "proxy-node-batch-import", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, "batch-imported-codex-access-token"); gateway_handle.abort(); @@ -3041,7 +3054,7 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints_impl() let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(status, StatusCode::OK, "payload={payload}"); assert_eq!(payload["total"], 1); - assert_eq!(payload["success"], 1); + assert_eq!(payload["success"], 1, "payload={payload}"); assert_eq!(payload["failed"], 0); assert_eq!(*token_hits.lock().expect("mutex should lock"), 0); @@ -3051,23 +3064,9 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints_impl() .expect("keys should load"); let persisted = reloaded.first().expect("persisted key should exist"); assert_eq!(persisted.expires_at_unix_secs, Some(2_100_000_000)); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, "chatgpt-web-batch-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "chatgpt_web"); @@ -3577,19 +3576,39 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_provider_oauth_state_entry_for_tests( - "nonce-codex-123", + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", json!({ - "nonce": "nonce-codex-123", + "nonce": "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "key_id": "key-codex-oauth", "provider_id": "provider-codex", "provider_type": "codex", "pkce_verifier": "verifier-codex-123", + "initiated_by_user_id": "admin-user-123", + "initiated_by_session_id": "session-123", + "created_at": aether_admin::provider::state::current_unix_secs(), }), ) .with_provider_oauth_token_url_for_tests("codex", format!("{token_url}/oauth/token")), ); let (gateway_url, gateway_handle) = start_server(gateway).await; + let wrong_session_response = reqwest::Client::new() + .post(format!( + "{gateway_url}/api/admin/provider-oauth/keys/key-codex-oauth/complete" + )) + .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-attacker") + .json(&json!({ + "callback_url": "http://localhost:1455/auth/callback?code=code-codex-123&state=aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" + })) + .send() + .await + .expect("mismatched-session request should succeed"); + assert_eq!(wrong_session_response.status(), StatusCode::BAD_REQUEST); + assert_eq!(*token_hits.lock().expect("mutex should lock"), 0); + let response = reqwest::Client::new() .post(format!( "{gateway_url}/api/admin/provider-oauth/keys/key-codex-oauth/complete" @@ -3599,7 +3618,7 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") .json(&json!({ - "callback_url": "http://localhost:1455/auth/callback?code=code-codex-123&state=nonce-codex-123" + "callback_url": "http://localhost:1455/auth/callback?code=code-codex-123&state=aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" })) .send() .await @@ -3686,23 +3705,9 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p scores[0].hard_state.schedulable(), "OAuth completion should replace AuthInvalid with a schedulable score" ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted); assert_eq!(decrypted_api_key, "new-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -3822,13 +3827,16 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_provider_oauth_state_entry_for_tests( - "nonce-provider-codex-123", + "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", json!({ - "nonce": "nonce-provider-codex-123", + "nonce": "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", "key_id": "", "provider_id": "provider-codex", "provider_type": "codex", "pkce_verifier": "verifier-provider-codex-123", + "initiated_by_user_id": "admin-user-123", + "initiated_by_session_id": "session-123", + "created_at": aether_admin::provider::state::current_unix_secs(), }), ) .with_provider_oauth_token_url_for_tests("codex", format!("{token_url}/oauth/token")), @@ -3844,7 +3852,7 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") .json(&json!({ - "callback_url": "http://localhost:1455/auth/callback?code=provider-code-123&state=nonce-provider-codex-123", + "callback_url": "http://localhost:1455/auth/callback?code=provider-code-123&state=bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", "proxy_node_id": "proxy-node-codex-oauth", "name": "should-not-override-inactive-name" })) @@ -3887,23 +3895,9 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad persisted.proxy, Some(json!({"node_id": "proxy-node-codex-oauth", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, "provider-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -3918,6 +3912,171 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad upstream_handle.abort(); } +#[test] +fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email() { + run_admin_oauth_test( + "gateway_names_new_antigravity_oauth_account_from_google_userinfo_email", + gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_impl, + ); +} + +async fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_impl() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().fallback(any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + })); + + let token_hits = Arc::new(Mutex::new(0usize)); + let token_hits_clone = Arc::clone(&token_hits); + let user_info_hits = Arc::new(Mutex::new(0usize)); + let user_info_hits_clone = Arc::clone(&user_info_hits); + let seen_user_info_authorization = Arc::new(Mutex::new(None::)); + let seen_user_info_authorization_clone = Arc::clone(&seen_user_info_authorization); + let google_server = Router::new() + .route( + "/oauth/token", + post(move || { + let token_hits_inner = Arc::clone(&token_hits_clone); + async move { + *token_hits_inner.lock().expect("mutex should lock") += 1; + Json(json!({ + "access_token": "antigravity-access-token", + "refresh_token": "antigravity-refresh-token", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "https://www.googleapis.com/auth/userinfo.email" + })) + } + }), + ) + .route( + "/oauth/userinfo", + get(move |headers: HeaderMap| { + let user_info_hits_inner = Arc::clone(&user_info_hits_clone); + let seen_authorization_inner = Arc::clone(&seen_user_info_authorization_clone); + async move { + *user_info_hits_inner.lock().expect("mutex should lock") += 1; + *seen_authorization_inner.lock().expect("mutex should lock") = headers + .get(http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + Json(json!({ + "email": "new-antigravity@example.com", + "verified_email": true, + "name": "Antigravity User" + })) + } + }), + ); + + let mut provider = sample_provider("provider-antigravity", "antigravity", 10); + provider.provider_type = "antigravity".to_string(); + let endpoint = sample_endpoint( + "endpoint-antigravity", + "provider-antigravity", + "gemini:generate_content", + "https://daily-cloudcode-pa.googleapis.com", + ); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![], + )); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let (google_url, google_handle) = start_server(google_server).await; + let gateway = build_router_with_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_state_entry_for_tests( + "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc", + json!({ + "nonce": "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc", + "key_id": "", + "provider_id": "provider-antigravity", + "provider_type": "antigravity", + "pkce_verifier": "verifier-antigravity-123", + "initiated_by_user_id": "admin-user-123", + "initiated_by_session_id": "session-123", + "created_at": aether_admin::provider::state::current_unix_secs(), + }), + ) + .with_provider_oauth_token_url_for_tests( + "antigravity", + format!("{google_url}/oauth/token"), + ) + .with_provider_oauth_token_url_for_tests( + "antigravity_user_info", + format!("{google_url}/oauth/userinfo"), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!( + "{gateway_url}/api/admin/provider-oauth/providers/provider-antigravity/complete" + )) + .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") + .json(&json!({ + "callback_url": "http://localhost:51121/oauth2callback?code=antigravity-code-123&state=cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc" + })) + .send() + .await + .expect("request should succeed"); + + let status = response.status(); + let payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + assert_eq!(payload["provider_type"], "antigravity"); + assert_eq!(payload["email"], "new-antigravity@example.com"); + assert_eq!(payload["replaced"], false); + assert_eq!(*token_hits.lock().expect("mutex should lock"), 1); + assert_eq!(*user_info_hits.lock().expect("mutex should lock"), 1); + assert_eq!( + seen_user_info_authorization + .lock() + .expect("mutex should lock") + .as_deref(), + Some("Bearer antigravity-access-token") + ); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + let key_id = payload["key_id"] + .as_str() + .expect("created key id should be returned") + .to_string(); + let persisted_keys = provider_catalog_repository + .list_keys_by_ids(std::slice::from_ref(&key_id)) + .await + .expect("created key should load"); + let persisted = persisted_keys.first().expect("created key should exist"); + assert_eq!(persisted.name, "new-antigravity@example.com"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(persisted); + let auth_config: Value = + serde_json::from_str(&decrypted_auth_config).expect("auth config json should parse"); + assert_eq!(auth_config["email"], "new-antigravity@example.com"); + assert_eq!(auth_config["refresh_token"], "antigravity-refresh-token"); + + gateway_handle.abort(); + google_handle.abort(); + upstream_handle.abort(); + drop(upstream_url); +} + #[test] fn gateway_imports_admin_provider_oauth_refresh_token_locally_with_trusted_admin_principal() { run_admin_oauth_test( @@ -4072,23 +4231,9 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_locally_with_trusted persisted.proxy, Some(json!({"node_id": "proxy-node-codex-import", "enabled": true})) ); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, "imported-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -4188,23 +4333,9 @@ async fn gateway_imports_codex_access_token_without_refresh_token_as_temporary_a .expect("keys should load"); let persisted = reloaded.first().expect("persisted key should exist"); assert_eq!(persisted.expires_at_unix_secs, Some(2_000_000_000)); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -4284,23 +4415,9 @@ async fn gateway_imports_codex_header_authorization_without_overwriting_payload_ .await .expect("keys should load"); let persisted = reloaded.first().expect("persisted key should exist"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["email"], "profile@example.com"); @@ -4404,23 +4521,9 @@ async fn gateway_imports_chatgpt_web_access_token_without_refresh_token_as_tempo .expect("keys should load"); let persisted = reloaded.first().expect("persisted key should exist"); assert_eq!(persisted.expires_at_unix_secs, Some(2_000_000_000)); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "chatgpt_web"); @@ -4501,14 +4604,7 @@ async fn gateway_imports_codex_access_token_with_payload_expires_at_when_token_h .expect("keys should load"); let persisted = reloaded.first().expect("persisted key should exist"); assert_eq!(persisted.expires_at_unix_secs, Some(2_100_000_000)); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["expires_at"], 2_100_000_000u64); @@ -4676,23 +4772,9 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_over_active_expired_ ); assert_eq!(persisted.oauth_invalid_at_unix_secs, None); assert_eq!(persisted.oauth_invalid_reason, None); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(persisted); assert_eq!(decrypted_api_key, "imported-expired-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - persisted - .encrypted_auth_config - .as_deref() - .expect("auth config should be stored"), - ) - .expect("auth config should decrypt"); + 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 json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -4827,7 +4909,6 @@ async fn gateway_import_invalidate_cached_oauth_entry_before_followup_resolution "key-codex-import-cache-duplicate", Some(1_700_000_000), Some("[OAUTH_EXPIRED] token invalidated"), - None, Some(1_700_000_000), ) .await @@ -5376,11 +5457,8 @@ async fn gateway_import_refresh_token_surfaces_execution_runtime_error_detail_im let status = response.status(); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(status, StatusCode::BAD_REQUEST, "payload={payload}"); - assert!( - payload["detail"] - .as_str() - .expect("detail should be string") - .contains("execution runtime returned HTTP 500"), + assert_eq!( + payload["detail"], "Refresh Token 验证失败: token exchange 失败", "payload={payload}" ); @@ -5493,7 +5571,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_locally_with_trusted_ad assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("payload should parse"); assert_eq!(payload["total"], 1); - assert_eq!(payload["success"], 1); + assert_eq!(payload["success"], 1, "payload={payload}"); assert_eq!(payload["failed"], 0); let results = payload["results"] .as_array() @@ -5518,14 +5596,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_locally_with_trusted_ad stored_key.proxy, Some(json!({"node_id": "proxy-node-kiro-batch", "enabled": true})) ); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["provider_type"], "kiro"); @@ -5644,7 +5715,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_over_active_expired_dup assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("payload should parse"); assert_eq!(payload["total"], 1); - assert_eq!(payload["success"], 1); + assert_eq!(payload["success"], 1, "payload={payload}"); assert_eq!(payload["failed"], 0); let results = payload["results"] .as_array() @@ -5671,14 +5742,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_over_active_expired_dup ); assert_eq!(stored_key.oauth_invalid_at_unix_secs, None); assert_eq!(stored_key.oauth_invalid_reason, None); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["provider_type"], "kiro"); @@ -5875,14 +5939,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_via_execution_runtime_p stored_key.proxy, Some(json!({"node_id": "proxy-node-kiro-batch-runtime", "enabled": true})) ); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["email"], "kiro-runtime@example.com"); @@ -6399,14 +6456,7 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p .into_iter() .next() .expect("refreshed key should exist"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("refreshed api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&stored_key); assert_eq!(decrypted_api_key, "refreshed-codex-access-token"); if account_state_recheck_attempted && payload["account_state_recheck_error"] == "wham/usage API 返回状态码 401" @@ -6429,14 +6479,7 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p assert_eq!(stored_key.oauth_invalid_reason, None); } - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("refreshed auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!( @@ -6479,7 +6522,7 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p ); assert_eq!( oauth_snapshot.get("reason"), - Some(&json!("Codex Token 已过期 (401)")) + Some(&json!("OAuth token has expired")) ); assert_eq!( oauth_snapshot @@ -6725,23 +6768,9 @@ async fn gateway_manual_codex_oauth_refresh_reconciles_missing_fixed_endpoint_im .into_iter() .next() .expect("refreshed key should exist"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_api_key - .as_deref() - .expect("api key should exist"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&stored_key); assert_eq!(decrypted_api_key, "refreshed-codex-access-token"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -6968,23 +6997,9 @@ async fn run_gateway_manual_kiro_oauth_refresh_maintenance_endpoint_test( .into_iter() .next() .expect("refreshed key should exist"); - let decrypted_api_key = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_api_key - .as_deref() - .expect("api key should exist"), - ) - .expect("api key should decrypt"); + let decrypted_api_key = decrypt_persisted_provider_api_key(&stored_key); assert_eq!(decrypted_api_key, expected_access_token); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["provider_type"], "kiro"); @@ -7170,9 +7185,7 @@ async fn gateway_marks_manual_oauth_refresh_failures_as_invalid_in_pool_payload_ assert_eq!(keys.len(), 1); assert_eq!( keys[0]["oauth_invalid_reason"], - json!( - "[REFRESH_FAILED] Token 续期失败 (401): refresh_token 已被使用并轮换,请重新登录授权" - ) + json!("[REFRESH_FAILED] OAuth token refresh failed") ); assert_eq!( keys[0]["status_snapshot"]["oauth"]["code"], @@ -7180,7 +7193,7 @@ async fn gateway_marks_manual_oauth_refresh_failures_as_invalid_in_pool_payload_ ); assert_eq!( keys[0]["status_snapshot"]["oauth"]["reason"], - json!("Token 续期失败 (401): refresh_token 已被使用并轮换,请重新登录授权") + json!("OAuth token refresh failed") ); assert_eq!( keys[0]["status_snapshot"]["oauth"]["source"], @@ -7588,6 +7601,11 @@ async fn gateway_refreshes_admin_provider_oauth_key_tunnel_proxy_with_direct_ref vec![endpoint], vec![key], )); + let mut tunnel_node = sample_proxy_node("proxy-node-tunnel"); + tunnel_node.status = "online".to_string(); + tunnel_node.tunnel_mode = true; + tunnel_node.tunnel_connected = true; + let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![tunnel_node])); let oauth_refresh = crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![ @@ -7604,6 +7622,17 @@ async fn gateway_refreshes_admin_provider_oauth_key_tunnel_proxy_with_direct_ref GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog_repository, ) + .attach_proxy_node_repository_for_tests(proxy_node_repository) + .with_system_config_values_for_tests(vec![( + "tunnel.attachments.proxy-node-tunnel".to_string(), + json!({ + "gateway_instance_id": "gateway-owner", + "relay_base_url": "http://gateway-owner.internal", + "tunnel_generation": "test-generation-1", + "conn_count": 1, + "observed_at_unix_secs": 4_102_444_800u64, + }), + )]) .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) .with_oauth_refresh_coordinator_for_tests(oauth_refresh), @@ -7841,14 +7870,7 @@ async fn gateway_consecutive_manual_oauth_refresh_uses_rotated_refresh_token_imp .into_iter() .next() .expect("refreshed key should exist"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!( @@ -8075,14 +8097,7 @@ async fn gateway_concurrent_manual_oauth_refresh_reuses_winner_after_lock_wait_i .into_iter() .next() .expect("refreshed key should exist"); - let decrypted_auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - stored_key - .encrypted_auth_config - .as_deref() - .expect("auth config should exist"), - ) - .expect("auth config should decrypt"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&stored_key); let auth_config: serde_json::Value = serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); assert_eq!(auth_config["refresh_token"], "rotated-codex-refresh-token"); @@ -8266,21 +8281,32 @@ async fn gateway_manual_oauth_refresh_prefers_fresher_transport_auth_config_over "Bearer cached-codex-access-token" ); - let fresh_auth_config = encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","refresh_token":"fresh-codex-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1,"updated_at":4102444810}"#, - ) - .expect("updated auth config ciphertext should build"); - assert!(provider_catalog_repository - .update_key_oauth_runtime_state( + let fresh_auth_config = app_state + .seal_provider_catalog_key_auth_config( + "provider-codex", "key-codex-oauth-stale-cache", - None, - None, - Some(&fresh_auth_config), - Some(4_102_444_810), + r#"{"provider_type":"codex","refresh_token":"fresh-codex-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1,"updated_at":4102444810}"#, ) + .expect("updated auth config ciphertext should build"); + let remotely_updated_key = provider_catalog_repository + .list_keys_by_ids(&["key-codex-oauth-stale-cache".to_string()]) .await - .expect("OAuth runtime state should update")); + .expect("provider key should reload") + .pop() + .expect("provider key should exist"); + let expected_encrypted_api_key = remotely_updated_key.encrypted_api_key.clone(); + let expected_encrypted_auth_config = remotely_updated_key.encrypted_auth_config.clone(); + assert!(provider_catalog_repository + .compare_and_swap_key_credentials(&ProviderCatalogKeyCredentialsCasUpdate { + key_id: remotely_updated_key.id.clone(), + expected_provider_id: remotely_updated_key.provider_id.clone(), + expected_encrypted_api_key, + expected_encrypted_auth_config, + encrypted_api_key: remotely_updated_key.encrypted_api_key.clone(), + encrypted_auth_config: Some(fresh_auth_config), + }) + .await + .expect("OAuth credential should update remotely")); let gateway = build_router_with_state(app_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -8785,6 +8811,187 @@ async fn gateway_upserts_admin_oauth_provider_locally_with_trusted_admin_princip upstream_handle.abort(); } +#[test] +fn gateway_requires_oauth_admin_management_token_to_configure_frontend_callback() { + run_admin_oauth_test( + "gateway_requires_oauth_admin_management_token_to_configure_frontend_callback", + gateway_requires_oauth_admin_management_token_to_configure_frontend_callback_impl, + ); +} + +async fn gateway_requires_oauth_admin_management_token_to_configure_frontend_callback_impl() { + let state = AppState::new().expect("gateway should build"); + let admin_user = state + .create_local_auth_user_with_settings( + Some("oauth-callback-admin@example.com".to_string()), + true, + "oauth-callback-admin".to_string(), + "hash".to_string(), + "admin".to_string(), + None, + None, + None, + None, + ) + .await + .expect("admin user should be created") + .expect("admin user should exist"); + + let write_raw_token = "ae-oauth-callback-write"; + let mut write_token = sample_management_token( + "token-oauth-callback-write", + &admin_user.id, + "oauth-callback-write", + true, + ); + write_token.token.allowed_ips = None; + write_token.token.permissions = Some(json!(["admin:oauth:write"])); + + let admin_raw_token = "ae-oauth-callback-admin"; + let mut oauth_admin_token = sample_management_token( + "token-oauth-callback-admin", + &admin_user.id, + "oauth-callback-admin", + true, + ); + oauth_admin_token.token.allowed_ips = None; + oauth_admin_token.token.permissions = Some(json!(["admin:oauth:admin"])); + + let management_tokens = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![write_token, oauth_admin_token], + vec![ + ( + hash_management_token(write_raw_token), + "token-oauth-callback-write".to_string(), + ), + ( + hash_management_token(admin_raw_token), + "token-oauth-callback-admin".to_string(), + ), + ], + )); + let oauth_providers = Arc::new(InMemoryOAuthProviderRepository::default()); + let data = GatewayDataState::with_management_token_repository_for_tests(management_tokens) + .attach_oauth_provider_repository_for_tests(Arc::clone(&oauth_providers)); + let gateway = build_router_with_state(state.with_data_state_for_tests(data)); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let provider_payload = |frontend_callback_url: &str| { + json!({ + "display_name": "Linux Do", + "client_id": "client-id", + "authorization_url_override": "https://connect.linux.do/oauth2/authorize", + "token_url_override": "https://connect.linux.do/oauth2/token", + "userinfo_url_override": "https://connect.linux.do/api/user", + "scopes": ["openid", "profile"], + "redirect_uri": "https://backend.example.com/oauth/callback", + "frontend_callback_url": frontend_callback_url, + "attribute_mapping": {"email": "email"}, + "extra_config": {"team": true}, + "is_enabled": true, + "force": false + }) + }; + let provider_url = format!("{gateway_url}/api/admin/oauth/providers/linuxdo"); + + let denied_create = client + .put(&provider_url) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(write_raw_token) + .json(&provider_payload("https://attacker.example/auth/callback")) + .send() + .await + .expect("write-token request should complete"); + let denied_create_status = denied_create.status(); + let denied_create_payload: Value = denied_create.json().await.expect("denial should be JSON"); + assert_eq!(denied_create_status, StatusCode::FORBIDDEN); + assert_eq!( + denied_create_payload["required_permission"], + "admin:oauth:admin" + ); + assert!(oauth_providers + .get_oauth_provider_config("linuxdo") + .await + .expect("provider lookup should succeed") + .is_none()); + + let allowed_create = client + .put(&provider_url) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(admin_raw_token) + .json(&provider_payload( + "https://configured.example/auth/callback", + )) + .send() + .await + .expect("oauth-admin request should complete"); + let allowed_create_status = allowed_create.status(); + let allowed_create_body = allowed_create.text().await.expect("body should read"); + assert_eq!( + allowed_create_status, + StatusCode::OK, + "body={allowed_create_body}" + ); + + let denied_change = client + .put(&provider_url) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(write_raw_token) + .json(&provider_payload( + "https://second-attacker.example/auth/callback", + )) + .send() + .await + .expect("write-token change request should complete"); + assert_eq!(denied_change.status(), StatusCode::FORBIDDEN); + assert_eq!( + oauth_providers + .get_oauth_provider_config("linuxdo") + .await + .expect("provider lookup should succeed") + .expect("provider should exist") + .frontend_callback_url, + "https://configured.example/auth/callback" + ); + + let allowed_session_change = client + .put(&provider_url) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .header(TRUSTED_ADMIN_USER_ID_HEADER, admin_user.id.as_str()) + .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") + .header( + TRUSTED_ADMIN_SESSION_ID_HEADER, + "session-oauth-callback-admin", + ) + .json(&provider_payload( + "https://frontend.example.com/auth/callback", + )) + .send() + .await + .expect("admin session request should complete"); + let allowed_session_status = allowed_session_change.status(); + let allowed_session_body = allowed_session_change + .text() + .await + .expect("body should read"); + assert_eq!( + allowed_session_status, + StatusCode::OK, + "body={allowed_session_body}" + ); + assert_eq!( + oauth_providers + .get_oauth_provider_config("linuxdo") + .await + .expect("provider lookup should succeed") + .expect("provider should exist") + .frontend_callback_url, + "https://frontend.example.com/auth/callback" + ); + + gateway_handle.abort(); +} + #[test] fn gateway_rejects_custom_oidc_without_allowed_domains() { run_admin_oauth_test( @@ -9032,14 +9239,14 @@ async fn gateway_upserts_multiple_custom_oidc_configs_impl() { } #[test] -fn gateway_tests_admin_oauth_linuxdo_endpoints_locally_with_configured_secret() { +fn gateway_blocks_admin_oauth_test_ssrf_to_private_endpoints() { run_admin_oauth_test( - "gateway_tests_admin_oauth_linuxdo_endpoints_locally_with_configured_secret", - gateway_tests_admin_oauth_linuxdo_endpoints_locally_with_configured_secret_impl, + "gateway_blocks_admin_oauth_test_ssrf_to_private_endpoints", + gateway_blocks_admin_oauth_test_ssrf_to_private_endpoints_impl, ); } -async fn gateway_tests_admin_oauth_linuxdo_endpoints_locally_with_configured_secret_impl() { +async fn gateway_blocks_admin_oauth_test_ssrf_to_private_endpoints_impl() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( @@ -9116,11 +9323,11 @@ async fn gateway_tests_admin_oauth_linuxdo_endpoints_locally_with_configured_sec assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["authorization_url_reachable"], true); - assert_eq!(payload["token_url_reachable"], true); + assert_eq!(payload["authorization_url_reachable"], false); + assert_eq!(payload["token_url_reachable"], false); assert_eq!(payload["secret_status"], "configured"); - assert_eq!(*authorization_hits.lock().expect("mutex should lock"), 1); - assert_eq!(*token_hits.lock().expect("mutex should lock"), 1); + assert_eq!(*authorization_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*token_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -9251,6 +9458,85 @@ async fn gateway_deletes_admin_oauth_provider_locally_with_trusted_admin_princip upstream_handle.abort(); } +#[test] +fn gateway_rejects_deleting_disabled_oauth_provider_with_user_bindings() { + run_admin_oauth_test( + "gateway_rejects_deleting_disabled_oauth_provider_with_user_bindings", + gateway_rejects_deleting_disabled_oauth_provider_with_user_bindings_impl, + ); +} + +async fn gateway_rejects_deleting_disabled_oauth_provider_with_user_bindings_impl() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().route( + "/api/admin/oauth/providers/linuxdo", + any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + }), + ); + + let mut provider = sample_oauth_provider_config("linuxdo"); + provider.is_enabled = false; + let provider_repository = Arc::new(InMemoryOAuthProviderRepository::seed([provider])); + let user_repository = Arc::new(InMemoryUserReadRepository::default()); + let now = chrono::Utc::now(); + let user = user_repository + .create_oauth_auth_user( + Some("bound@example.com".to_string()), + false, + "bound-user".to_string(), + now, + ) + .await + .expect("OAuth user should create") + .expect("OAuth user should exist"); + user_repository + .bind_user_oauth_link(&user.id, "linuxdo", "bound-subject", None, None, None, now) + .await + .expect("OAuth link should bind"); + + let data = GatewayDataState::with_oauth_provider_repository_for_tests(Arc::clone( + &provider_repository, + )) + .with_user_reader(user_repository); + let (upstream_url, upstream_handle) = start_server(upstream).await; + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .delete(format!("{gateway_url}/api/admin/oauth/providers/linuxdo")) + .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") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::CONFLICT); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["error"]["type"], "provider_has_bindings"); + assert!(provider_repository + .get_oauth_provider_config("linuxdo") + .await + .expect("provider lookup should succeed") + .is_some()); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); + let _ = upstream_url; +} + #[test] fn gateway_handles_admin_management_token_detail_locally_with_trusted_admin_principal() { run_admin_oauth_test( @@ -9497,7 +9783,11 @@ async fn gateway_allows_management_token_with_pool_write_for_provider_oauth_batc true, ); management_token.token.allowed_ips = None; - management_token.token.permissions = Some(json!(["admin:pool:read", "admin:pool:write"])); + management_token.token.permissions = Some(json!([ + "admin:pool:read", + "admin:pool:write", + "admin:provider_oauth:admin" + ])); let management_token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( vec![management_token], @@ -9604,6 +9894,7 @@ async fn gateway_prevents_pool_write_token_from_importing_agent_identity_via_bat for path in [ "/api/admin/provider-oauth/providers/provider-codex/batch-import", "/api/admin/provider-oauth/providers/provider-codex/batch-import/tasks", + "/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks", ] { let response = client .post(format!("{gateway_url}{path}")) @@ -9618,40 +9909,13 @@ async fn gateway_prevents_pool_write_token_from_importing_agent_identity_via_bat let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!( status, - StatusCode::BAD_REQUEST, + StatusCode::FORBIDDEN, "path={path} payload={payload}" ); - assert_eq!( - payload["detail"], - "Agent Identity JSON 必须使用专属导入接口" - ); + assert_eq!(payload["detail"], "management token permission denied"); + assert_eq!(payload["required_permission"], "admin:provider_oauth:admin"); } - let dedicated_response = client - .post(format!( - "{gateway_url}/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks" - )) - .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") - .bearer_auth(raw_token) - .json(&json!({ "credentials": credentials })) - .send() - .await - .expect("request should succeed"); - let dedicated_status = dedicated_response.status(); - let dedicated_payload: serde_json::Value = dedicated_response - .json() - .await - .expect("json body should parse"); - assert_eq!( - dedicated_status, - StatusCode::FORBIDDEN, - "payload={dedicated_payload}" - ); - assert_eq!( - dedicated_payload["required_permission"], - "admin:provider_oauth:write" - ); - let keys = provider_catalog_repository .list_keys_by_provider_ids(&["provider-codex".to_string()]) .await @@ -9723,7 +9987,7 @@ async fn gateway_rejects_management_token_without_pool_write_for_provider_oauth_ assert_eq!(response.status(), StatusCode::FORBIDDEN); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["detail"], "management token permission denied"); - assert_eq!(payload["required_permission"], "admin:pool:write"); + assert_eq!(payload["required_permission"], "admin:provider_oauth:admin"); assert_eq!(payload["route_family"], "provider_oauth_manage"); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/control/admin/payments.rs b/apps/aether-gateway/src/tests/control/admin/payments.rs index 67ce86e38..5f3539ef7 100644 --- a/apps/aether-gateway/src/tests/control/admin/payments.rs +++ b/apps/aether-gateway/src/tests/control/admin/payments.rs @@ -249,10 +249,9 @@ async fn gateway_handles_admin_payments_expire_order_locally_with_trusted_admin_ let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["expired"], true); assert_eq!(payload["order"]["status"], "expired"); - assert_eq!( - payload["order"]["gateway_response"]["expire_reason"], - "admin_mark_expired" - ); + assert!(payload["order"]["gateway_response"] + .get("expire_reason") + .is_none()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -305,15 +304,13 @@ async fn gateway_handles_admin_payments_credit_order_locally_with_trusted_admin_ assert_eq!(payload["order"]["pay_amount"], 91.25); assert_eq!(payload["order"]["pay_currency"], "CNY"); assert_eq!(payload["order"]["exchange_rate"], 7.3); - assert_eq!( - payload["order"]["gateway_response"]["channel"], - "manual-review" - ); + assert!(payload["order"]["gateway_response"] + .get("channel") + .is_none()); assert_eq!(payload["order"]["gateway_response"]["manual_credit"], true); - assert_eq!( - payload["order"]["gateway_response"]["credited_by"], - "admin-user-123" - ); + assert!(payload["order"]["gateway_response"] + .get("credited_by") + .is_none()); assert!(payload["order"]["paid_at"].as_str().is_some()); assert!(payload["order"]["credited_at"].as_str().is_some()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -354,10 +351,9 @@ async fn gateway_handles_admin_payments_fail_order_locally_with_trusted_admin_pr assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["order"]["status"], "failed"); - assert_eq!( - payload["order"]["gateway_response"]["failure_reason"], - "admin_mark_failed" - ); + assert!(payload["order"]["gateway_response"] + .get("failure_reason") + .is_none()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -410,7 +406,8 @@ async fn gateway_handles_admin_payments_callbacks_locally_with_trusted_admin_pri assert_eq!(items[0]["callback_key"], "callback-key-1"); assert_eq!(items[0]["signature_valid"], true); assert_eq!(items[0]["status"], "processed"); - assert_eq!(items[0]["payload"]["source"], "callback-1"); + assert!(items[0]["payload"].is_null()); + assert_eq!(items[0]["payload_summary"]["objects"], 1); assert!(items[0]["processed_at"].as_str().is_some()); assert_eq!(payload["total"], 1); assert_eq!(payload["limit"], 10); diff --git a/apps/aether-gateway/src/tests/control/admin/pool.rs b/apps/aether-gateway/src/tests/control/admin/pool.rs index d35f5fb0b..aec3221f9 100644 --- a/apps/aether-gateway/src/tests/control/admin/pool.rs +++ b/apps/aether-gateway/src/tests/control/admin/pool.rs @@ -1,6 +1,6 @@ use std::sync::{Arc, Mutex}; -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::pool_scores::InMemoryPoolMemberScoreRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::usage::InMemoryUsageReadRepository; @@ -16,7 +16,8 @@ use http::{HeaderMap, HeaderValue, StatusCode}; use serde_json::json; use super::super::{ - build_router_with_state, sample_endpoint, sample_key, sample_provider, start_server, AppState, + build_router_with_state, sample_bound_auth_config, sample_bound_key as sample_key, + sample_endpoint, sample_provider, start_server, AppState, }; use crate::admin_api::{maybe_build_local_admin_pool_response, AdminAppState, AdminRequestContext}; use crate::ai_serving::{provider_key_pool_score_id, provider_key_pool_score_scope}; @@ -1683,13 +1684,11 @@ async fn gateway_handles_admin_pool_list_keys_with_quota_compatibility_fields() key.name = "quota-key".to_string(); key.auth_type = "oauth".to_string(); key.expires_at_unix_secs = Some(1_775_556_730); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"plan_type":"pro","account_id":"acct-antigravity-1","account_name":"quota-user","account_user_id":"quota-user-1","organizations":[{"id":"org-1","name":"Org One"}]}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-antigravity", + "key-antigravity-a", + r#"{"plan_type":"pro","account_id":"acct-antigravity-1","account_name":"quota-user","account_user_id":"quota-user-1","organizations":[{"id":"org-1","name":"Org One"}]}"#, + )); key.status_snapshot = Some(json!({ "oauth": { "code": "expired", @@ -1806,20 +1805,18 @@ async fn gateway_includes_pool_quota_and_compat_fields_in_list_keys_response() { key.name = "quota key".to_string(); key.auth_type = "oauth".to_string(); key.expires_at_unix_secs = Some(1_775_556_730); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - &json!({ - "plan_type": "pro", - "account_id": "acct-demo-001", - "account_name": "Demo Account", - "account_user_id": "user-demo-001", - "organizations": [], - }) - .to_string(), - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-antigravity", + "key-antigravity-oauth", + &json!({ + "plan_type": "pro", + "account_id": "acct-demo-001", + "account_name": "Demo Account", + "account_user_id": "user-demo-001", + "organizations": [], + }) + .to_string(), + )); key.upstream_metadata = Some(json!({ "antigravity": { "updated_at": 1_775_553_285u64, @@ -2942,17 +2939,15 @@ async fn gateway_pool_prefers_upstream_plan_type_over_auth_config() { "oauth-placeholder", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - &json!({ - "plan_type": "free", - "account_id": "acct-codex-legacy" - }) - .to_string(), - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-precedence", + &json!({ + "plan_type": "free", + "account_id": "acct-codex-legacy" + }) + .to_string(), + )); key.upstream_metadata = Some(json!({ "codex": { "plan_type": "plus", @@ -3017,17 +3012,15 @@ async fn gateway_pool_plan_free_selector_prefers_upstream_plan_type() { "oauth-placeholder", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - &json!({ - "plan_type": "free", - "account_id": "acct-codex-legacy" - }) - .to_string(), - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-selector", + &json!({ + "plan_type": "free", + "account_id": "acct-codex-legacy" + }) + .to_string(), + )); key.upstream_metadata = Some(json!({ "codex": { "plan_type": "plus", @@ -3099,13 +3092,11 @@ async fn gateway_pool_keys_classify_oauth_credentials() { "imported-session-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-oauth-header", + r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, + )); let mut agent_key = sample_key( "key-codex-agent-identity", "provider-codex", @@ -3113,13 +3104,11 @@ async fn gateway_pool_keys_classify_oauth_credentials() { "", ); agent_key.auth_type = "oauth".to_string(); - agent_key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#, - ) - .expect("Agent Identity auth config should encrypt"), - ); + agent_key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-agent-identity", + r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -3191,13 +3180,11 @@ async fn gateway_pool_resolve_selection_marks_oauth_header_auth() { "imported-session-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex", + "key-codex-oauth-header", + r#"{"provider_type":"codex","headers":{"authorization":"Bearer imported-session-token"}}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -3265,7 +3252,7 @@ async fn gateway_handles_admin_pool_resolve_selection_locally_with_trusted_admin proxy_key.name = "alpha proxy".to_string(); proxy_key.proxy = Some(json!({ "mode": "direct", - "url": "https://proxy.example.com" + "url": "https://proxy.example.com/" })); let mut plain_key = sample_key("key-openai-b", "provider-openai", "openai:chat", "sk-b"); plain_key.name = "beta".to_string(); @@ -3275,7 +3262,7 @@ async fn gateway_handles_admin_pool_resolve_selection_locally_with_trusted_admin disabled_proxy_key.is_active = false; disabled_proxy_key.proxy = Some(json!({ "mode": "direct", - "url": "https://proxy-disabled.example.com" + "url": "https://proxy-disabled.example.com/" })); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -3455,13 +3442,11 @@ async fn gateway_resolve_selection_marks_legacy_kiro_bearer_keys_as_oauth_manage "kiro-access-token", ); key.auth_type = "bearer".to_string(); - key.encrypted_auth_config = Some( - encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"kiro","email":"legacy-kiro@example.com","refresh_token":"legacy-kiro-refresh-token"}"#, - ) - .expect("auth config ciphertext should build"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + "key-kiro-legacy", + r#"{"provider_type":"kiro","email":"legacy-kiro@example.com","refresh_token":"legacy-kiro-refresh-token"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], 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 decfc4297..ccc2fee1f 100644 --- a/apps/aether-gateway/src/tests/control/admin/provider_ops.rs +++ b/apps/aether-gateway/src/tests/control/admin/provider_ops.rs @@ -3,9 +3,7 @@ use std::sync::{Arc, Mutex}; use aether_contracts::{ ExecutionPlan, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, }; -use aether_crypto::{ - decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, -}; +use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository; @@ -31,9 +29,40 @@ use crate::constants::{ TRUSTED_ADMIN_USER_ROLE_HEADER, }; use crate::data::{GatewayDataConfig, GatewayDataState}; +use crate::handlers::shared::{ + open_provider_ops_credential, provider_ops_credential_binding_from_config, +}; const PROVIDER_OPS_TEST_STACK_BYTES: usize = 16 * 1024 * 1024; +fn open_stored_provider_ops_credential( + provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider, + field: &str, + stored: &str, +) -> String { + let provider_ops = provider + .config + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|config| config.get("provider_ops")) + .and_then(serde_json::Value::as_object) + .expect("provider ops config should be present"); + let base_url = provider_ops + .get("base_url") + .and_then(serde_json::Value::as_str) + .expect("provider ops base_url should be present"); + let binding = provider_ops_credential_binding_from_config(&provider.id, provider_ops, base_url) + .expect("provider ops credential binding should be valid"); + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + open_provider_ops_credential(&state, &binding, field, stored) + .expect("provider ops credential should decrypt") + .plaintext +} + fn run_provider_ops_test(test_name: &'static str, make_future: F) where F: FnOnce() -> Fut + Send + 'static, @@ -548,7 +577,10 @@ async fn gateway_saves_admin_provider_ops_config_locally_with_trusted_admin_prin "tenant": "acme" }, "credentials": { - "refresh_token": "************", + // A binding change (architecture/auth type/base URL) must be + // accompanied by a newly supplied secret; masked values are + // intentionally rejected by the handler. + "refresh_token": "new-refresh-secret", "api_key": "live-secret-api-key", } }, @@ -615,9 +647,8 @@ async fn gateway_saves_admin_provider_ops_config_locally_with_trusted_admin_prin .expect("refresh token should be string"); assert_ne!(stored_refresh, "refresh-secret-1234"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_refresh) - .expect("refresh token should decrypt"), - "refresh-secret-1234" + open_stored_provider_ops_credential(&stored_provider, "refresh_token", stored_refresh), + "new-refresh-secret" ); let stored_api_key = credentials .get("api_key") @@ -625,8 +656,7 @@ async fn gateway_saves_admin_provider_ops_config_locally_with_trusted_admin_prin .expect("api key should be string"); assert_ne!(stored_api_key, "live-secret-api-key"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_api_key) - .expect("api key should decrypt"), + open_stored_provider_ops_credential(&stored_provider, "api_key", stored_api_key), "live-secret-api-key" ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -1301,7 +1331,7 @@ async fn gateway_verifies_admin_provider_ops_locally_for_anyrouter_proxy_mode_im plan.headers .get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) .map(String::as_str), - None + Some("false") ); Json(json!({ "request_id": plan.request_id, @@ -1717,7 +1747,9 @@ async fn gateway_verifies_admin_provider_ops_locally_for_new_api_with_trusted_ad assert_eq!(payload["data"]["used_quota"], 12.5); assert_eq!(payload["data"]["request_count"], 9); assert_eq!(payload["data"]["email"], ""); - assert_eq!(payload["data"]["extra"]["group"], "default"); + // The safe verification projection intentionally drops unknown upstream + // fields such as `group`; only the documented fields are returned. + assert_eq!(payload["data"]["extra"], json!({})); assert_eq!(payload["updated_credentials"], serde_json::Value::Null); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -2536,8 +2568,11 @@ async fn gateway_verifies_admin_provider_ops_sub2api_persists_rotated_runtime_cr .and_then(serde_json::Value::as_str) .expect("refresh token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_refresh_token) - .expect("refresh token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "refresh_token", + stored_refresh_token + ), "refresh-token-new" ); let stored_cached_access_token = credentials @@ -2545,8 +2580,11 @@ async fn gateway_verifies_admin_provider_ops_sub2api_persists_rotated_runtime_cr .and_then(serde_json::Value::as_str) .expect("cached access token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_cached_access_token,) - .expect("cached access token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "_cached_access_token", + stored_cached_access_token, + ), "access-token-new" ); assert!( @@ -2933,7 +2971,7 @@ async fn gateway_handles_admin_provider_ops_balance_locally_for_generic_api_prox assert_eq!(payload["data"]["total_available"], 5.0); assert_eq!(payload["data"]["total_used"], 1.0); assert_eq!(payload["data"]["extra"]["checkin_success"], true); - assert_eq!(payload["data"]["extra"]["checkin_message"], "代理签到成功"); + assert_eq!(payload["data"]["extra"]["checkin_message"], "签到成功"); let plans = execution_plans.lock().expect("mutex should lock"); assert_eq!(plans.len(), 2); @@ -3078,10 +3116,7 @@ async fn gateway_handles_admin_provider_ops_balance_locally_without_proxy_via_ex assert_eq!(payload["data"]["total_available"], 5.0); assert_eq!(payload["data"]["total_used"], 1.0); assert_eq!(payload["data"]["extra"]["checkin_success"], true); - assert_eq!( - payload["data"]["extra"]["checkin_message"], - "执行层签到成功" - ); + assert_eq!(payload["data"]["extra"]["checkin_message"], "签到成功"); let plans = execution_plans.lock().expect("mutex should lock"); assert_eq!(plans.len(), 2); @@ -3208,7 +3243,7 @@ async fn gateway_handles_admin_provider_ops_checkin_locally_with_trusted_admin_p assert_eq!(payload["action_type"], "checkin"); assert_eq!(payload["data"]["reward"], 1.5); assert_eq!(payload["data"]["streak_days"], 3); - assert_eq!(payload["data"]["message"], "今日签到完成"); + assert_eq!(payload["data"]["message"], "签到成功"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -3341,7 +3376,7 @@ async fn gateway_handles_admin_provider_ops_checkin_locally_for_generic_api_prox assert_eq!(payload["action_type"], "checkin"); assert_eq!(payload["data"]["reward"], 1.5); assert_eq!(payload["data"]["streak_days"], 3); - assert_eq!(payload["data"]["message"], "今日签到完成"); + assert_eq!(payload["data"]["message"], "签到成功"); assert_eq!(execution_plans.lock().expect("mutex should lock").len(), 1); gateway_handle.abort(); @@ -3582,7 +3617,7 @@ async fn gateway_handles_admin_provider_ops_batch_balance_locally_with_trusted_a ); assert_eq!( payload["provider-anyrouter"]["data"]["extra"]["checkin_message"], - "Anyrouter 签到成功" + "签到成功" ); assert_eq!(payload["provider-missing"]["status"], "not_configured"); assert_eq!(payload["provider-missing"]["message"], "未配置操作设置"); @@ -4785,8 +4820,11 @@ async fn gateway_handles_admin_provider_ops_sub2api_balance_with_refresh_token_r .and_then(serde_json::Value::as_str) .expect("refresh token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_refresh_token) - .expect("refresh token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "refresh_token", + stored_refresh_token + ), "refresh-token-new" ); let stored_cached_access_token = credentials @@ -4794,8 +4832,11 @@ async fn gateway_handles_admin_provider_ops_sub2api_balance_with_refresh_token_r .and_then(serde_json::Value::as_str) .expect("cached access token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_cached_access_token,) - .expect("cached access token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "_cached_access_token", + stored_cached_access_token, + ), "access-token-new" ); assert!( @@ -5187,8 +5228,11 @@ async fn gateway_handles_admin_provider_ops_sub2api_balance_with_session_login_i .and_then(serde_json::Value::as_str) .expect("refresh token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_refresh_token) - .expect("refresh token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "refresh_token", + stored_refresh_token + ), "login-refresh-token" ); let stored_cached_access_token = credentials @@ -5196,8 +5240,11 @@ async fn gateway_handles_admin_provider_ops_sub2api_balance_with_session_login_i .and_then(serde_json::Value::as_str) .expect("cached access token should be string"); assert_eq!( - decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, stored_cached_access_token,) - .expect("cached access token should decrypt"), + open_stored_provider_ops_credential( + &stored_provider, + "_cached_access_token", + stored_cached_access_token, + ), "login-access-token" ); assert!( 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 e18848467..2535c9cdc 100644 --- a/apps/aether-gateway/src/tests/control/admin/provider_query.rs +++ b/apps/aether-gateway/src/tests/control/admin/provider_query.rs @@ -20,8 +20,8 @@ use serde_json::json; use super::super::{ build_router_with_state, build_state_with_execution_runtime_override, - sample_admin_provider_model, sample_endpoint, sample_key, sample_provider, start_server, - AppState, + sample_admin_provider_model, sample_bound_auth_config, sample_bound_key, sample_endpoint, + sample_provider, start_server, AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -216,7 +216,7 @@ async fn gateway_handles_admin_provider_query_models_fetches_upstream_for_select None, ) .expect("endpoint transport should build")], - vec![sample_key( + vec![sample_bound_key( "key-openai-selected", "provider-openai", "openai:chat", @@ -342,7 +342,7 @@ async fn gateway_handles_admin_provider_query_models_fetches_windsurf_model_conf let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-windsurf", "Windsurf", 10); provider.provider_type = "windsurf".to_string(); - let mut windsurf_key = sample_key( + let mut windsurf_key = sample_bound_key( "key-windsurf-selected", "provider-windsurf", "openai:chat", @@ -487,7 +487,7 @@ async fn gateway_handles_admin_provider_query_models_with_openai_responses_endpo None, ) .expect("endpoint transport should build")], - vec![sample_key( + vec![sample_bound_key( "key-openai-responses", "provider-openai", "openai:responses", @@ -538,14 +538,14 @@ async fn gateway_handles_admin_provider_query_models_with_openai_responses_endpo } #[test] -fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache() { +fn gateway_recovers_codex_slug_only_models_from_a_stale_legacy_cache() { run_provider_query_test( - "gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache", - gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl, + "gateway_recovers_codex_slug_only_models_from_a_stale_legacy_cache", + gateway_recovers_codex_slug_only_models_from_a_stale_legacy_cache_impl, ); } -async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl() { +async fn gateway_recovers_codex_slug_only_models_from_a_stale_legacy_cache_impl() { let execution_runtime_hits = Arc::new(Mutex::new(0usize)); let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits); let execution_runtime = Router::new().route( @@ -558,7 +558,11 @@ async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl .expect("mutex should lock") += 1; assert_eq!( plan.url, - "https://chatgpt.com/backend-api/codex/models?client_version=0.144.1" + "https://chatgpt.com/backend-api/codex/models?client_version=0.153.3" + ); + assert_eq!( + plan.headers.get("user-agent").map(String::as_str), + Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) ); assert_eq!(plan.provider_api_format, "openai:responses"); Json(json!({ @@ -600,7 +604,7 @@ async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl "openai:responses", "https://chatgpt.com/backend-api/codex", )], - vec![sample_key( + vec![sample_bound_key( "key-codex-dynamic", "provider-codex-dynamic", "openai:responses", @@ -617,16 +621,21 @@ async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl .runtime_state() .kv_set( "upstream_models:provider-codex-dynamic:key-codex-dynamic", - "[]".to_string(), + json!([{"id": "gpt-stale-legacy"}]).to_string(), None, ) .await - .expect("empty legacy cache should seed"); + .expect("stale legacy cache should seed"); let cache_state = state.clone(); let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; - for (request_index, expected_from_cache) in [(0usize, false), (1usize, true)] { + for (request_index, expected_from_cache) in [ + (0usize, false), + (1usize, true), + (2usize, false), + (3usize, true), + ] { let response = reqwest::Client::new() .post(format!("{gateway_url}/api/admin/provider-query/models")) .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") @@ -635,7 +644,8 @@ async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") .json(&json!({ "provider_id": "provider-codex-dynamic", - "api_key_id": "key-codex-dynamic" + "api_key_id": "key-codex-dynamic", + "force_refresh": request_index == 2, })) .send() .await @@ -672,10 +682,47 @@ async fn gateway_recovers_codex_slug_only_models_from_an_empty_legacy_cache_impl ); assert_eq!( *execution_runtime_hits.lock().expect("mutex should lock"), - 1, + if request_index < 2 { 1 } else { 2 }, "request {request_index} must not cause another upstream fetch" ); + assert_eq!( + payload["data"]["models"][0]["display_name"], + if request_index == 1 { + "Updated shared catalog" + } else { + "Future Dynamic" + } + ); if request_index == 0 { + let context = crate::model_fetch::read_codex_management_catalog( + &cache_state, + "provider-codex-dynamic", + "key-codex-dynamic", + ) + .await + .unwrap(); + let mut models = context + .models + .clone() + .expect("admin fetch publishes shared catalog"); + models[0]["display_name"] = json!("Updated shared catalog"); + let transport = cache_state + .read_provider_transport_snapshot( + "provider-codex-dynamic", + "endpoint-codex-dynamic", + "key-codex-dynamic", + ) + .await + .unwrap() + .unwrap(); + crate::model_fetch::store_codex_management_catalog( + &cache_state, + &context, + &[transport], + models, + None, + ) + .await; ::write_upstream_models_cache( &cache_state, "provider-codex-dynamic", @@ -712,7 +759,7 @@ async fn gateway_handles_admin_provider_query_models_falls_back_to_codex_preset_ .expect("mutex should lock") += 1; assert_eq!( plan.url, - "https://chatgpt.com/backend-api/codex/models?client_version=0.144.1" + "https://chatgpt.com/backend-api/codex/models?client_version=0.153.3" ); Json(json!({ "request_id": "req-provider-query-codex-invalidated", @@ -743,7 +790,7 @@ async fn gateway_handles_admin_provider_query_models_falls_back_to_codex_preset_ "openai:responses", "https://chatgpt.com/backend-api/codex", )], - vec![sample_key( + vec![sample_bound_key( "key-codex-invalidated", "provider-codex", "openai:responses", @@ -782,7 +829,13 @@ async fn gateway_handles_admin_provider_query_models_falls_back_to_codex_preset_ .as_str() .expect("Codex fallback warning should be present"); assert!(warning.contains("Codex 动态模型目录不可用")); - assert!(warning.contains("invalidated")); + // Model-fetch diagnostics are intentionally projected to a credential-safe + // category before being returned from the admin endpoint. The raw + // upstream invalidation text must not cross the response boundary. + assert!( + warning.contains("authorization failed (status 403)"), + "warning={warning}" + ); let model_ids = payload["data"]["models"] .as_array() .expect("models should be an array") @@ -900,7 +953,7 @@ async fn gateway_handles_admin_provider_query_models_respecting_key_api_formats_ ) .expect("endpoint transport should build"), ], - vec![sample_key( + vec![sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", @@ -1036,13 +1089,13 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys_imp ) .expect("endpoint transport should build")], vec![ - sample_key( + sample_bound_key( "key-openai-1", "provider-openai", "openai:chat", "sk-test-1", ), - sample_key( + sample_bound_key( "key-openai-2", "provider-openai", "openai:chat", @@ -1131,7 +1184,7 @@ async fn gateway_handles_admin_provider_query_models_for_fixed_provider_without_ let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![], - vec![sample_key( + vec![sample_bound_key( "key-codex-oauth", "provider-codex", "openai:responses", @@ -1251,7 +1304,7 @@ async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_ad "openai:chat", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-primary", "provider-openai", "openai:chat", @@ -1365,7 +1418,7 @@ async fn gateway_handles_admin_provider_query_embedding_model_test_impl() { "openai:embedding", "https://api.siliconflow.example", )], - vec![sample_key( + vec![sample_bound_key( "key-siliconflow-embedding", "provider-siliconflow", "openai:embedding", @@ -1486,7 +1539,7 @@ async fn gateway_handles_admin_provider_query_doubao_text_embedding_model_test_i "doubao:embedding", "https://ark.volces.example/api/v3", )], - vec![sample_key( + vec![sample_bound_key( "key-doubao-embedding", "provider-doubao", "doubao:embedding", @@ -1608,7 +1661,7 @@ async fn gateway_handles_admin_provider_query_gemini_embedding_model_test_impl() "gemini:embedding", "https://generativelanguage.googleapis.com/v1beta", )], - vec![sample_key( + vec![sample_bound_key( "key-gemini-embedding", "provider-gemini", "gemini:embedding", @@ -1726,7 +1779,7 @@ async fn gateway_handles_admin_provider_query_vertex_gemini_embedding_model_test let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-vertex-ai", "Vertex AI", 10); provider.provider_type = "vertex_ai".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-vertex-gemini-embedding", "provider-vertex-ai", "gemini:embedding", @@ -1873,7 +1926,7 @@ async fn gateway_handles_admin_provider_query_jina_embedding_model_test_impl() { "jina:embedding", "https://api.jina.example", )], - vec![sample_key( + vec![sample_bound_key( "key-jina-embedding", "provider-jina-embedding", "jina:embedding", @@ -1995,7 +2048,7 @@ async fn gateway_handles_admin_provider_query_openai_rerank_model_test_impl() { "openai:rerank", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-rerank", "provider-openai-rerank", "openai:rerank", @@ -2123,7 +2176,7 @@ async fn gateway_handles_admin_provider_query_rerank_model_test_impl() { "jina:rerank", "https://api.jina.example", )], - vec![sample_key( + vec![sample_bound_key( "key-jina-rerank", "provider-jina", "jina:rerank", @@ -2254,7 +2307,7 @@ async fn gateway_maps_admin_provider_model_before_model_list_test_request_impl() "openai:chat", "https://api.minimax.example", )], - vec![sample_key( + vec![sample_bound_key( "key-minimax-primary", "provider-minimax", "openai:chat", @@ -2498,7 +2551,7 @@ async fn gateway_streams_codex_openai_responses_upstream_for_admin_pool_model_te let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![endpoint], - vec![sample_key( + vec![sample_bound_key( "key-codex-primary", "provider-codex", "openai:responses", @@ -2656,20 +2709,18 @@ async fn gateway_executes_codex_search_admin_pool_model_test_with_search_contrac "https://chatgpt.com/backend-api/codex", ); endpoint.config = Some(json!({"upstream_stream_policy": "force_stream"})); - let mut key = sample_key( + let mut key = sample_bound_key( "key-codex-search", "provider-codex-search", "openai:search", "codex-search-access-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"codex","account_id":"account-search-admin","is_fedramp":true}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-codex-search", + "key-codex-search", + r#"{"provider_type":"codex","account_id":"account-search-admin","is_fedramp":true}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![endpoint], @@ -2796,24 +2847,22 @@ async fn gateway_routes_grok_responses_admin_pool_model_test_through_grok_runtim let mut provider = sample_provider("provider-grok", "Grok", 10); provider.provider_type = "grok".to_string(); provider.config = Some(json!({"pool_advanced": {}})); - let mut key = sample_key( + let mut key = sample_bound_key( "key-grok-oauth", "provider-grok", "openai:responses", "__placeholder__", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-grok", + "key-grok-oauth", + r#"{ "provider_type":"grok", "sso_token":"grok-sso", "sso_rw_token":"grok-rw" }"#, - ) - .expect("auth config should encrypt"), - ); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -2933,7 +2982,7 @@ async fn gateway_streams_windsurf_connect_upstream_for_admin_model_test_impl() { let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-windsurf", "Windsurf", 10); provider.provider_type = "windsurf".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-windsurf-primary", "provider-windsurf", "openai:chat", @@ -3045,7 +3094,7 @@ async fn gateway_uses_pool_scheduler_order_for_admin_pool_model_test_impl() { ] } })); - let mut free_key = sample_key( + let mut free_key = sample_bound_key( "key-codex-free", "provider-codex", "openai:responses", @@ -3060,7 +3109,7 @@ async fn gateway_uses_pool_scheduler_order_for_admin_pool_model_test_impl() { "usage_ratio": 0.1 } })); - let mut plus_key = sample_key( + let mut plus_key = sample_bound_key( "key-codex-plus", "provider-codex", "openai:responses", @@ -3253,13 +3302,13 @@ async fn gateway_handles_admin_provider_query_test_model_failover_locally_with_t "https://api.openai.example/v1", )], vec![ - sample_key( + sample_bound_key( "key-openai-first", "provider-openai", "openai:chat", "sk-test-first", ), - sample_key( + sample_bound_key( "key-openai-second", "provider-openai", "openai:chat", @@ -3382,17 +3431,17 @@ async fn gateway_handles_admin_provider_query_test_model_for_kiro_locally_impl() let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-kiro", "Kiro", 10); provider.provider_type = "kiro".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-kiro-primary", "provider-kiro", "claude:messages", "__placeholder__", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + "key-kiro-primary", + r#"{ "provider_type":"kiro", "auth_method":"idc", "access_token":"cached-kiro-token", @@ -3403,9 +3452,7 @@ async fn gateway_handles_admin_provider_query_test_model_for_kiro_locally_impl() "client_id":"client-id", "client_secret":"client-secret" }"#, - ) - .expect("auth config should encrypt"), - ); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -3518,17 +3565,17 @@ async fn gateway_uses_kiro_mapped_model_name_for_explicit_model_mapping_test_imp let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-kiro", "Kiro", 10); provider.provider_type = "kiro".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-kiro-primary", "provider-kiro", "claude:messages", "__placeholder__", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + "key-kiro-primary", + r#"{ "provider_type":"kiro", "auth_method":"idc", "access_token":"cached-kiro-token", @@ -3539,9 +3586,7 @@ async fn gateway_uses_kiro_mapped_model_name_for_explicit_model_mapping_test_imp "client_id":"client-id", "client_secret":"client-secret" }"#, - ) - .expect("auth config should encrypt"), - ); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], @@ -3728,12 +3773,12 @@ async fn gateway_handles_admin_provider_query_test_model_failover_for_kiro_local let mut provider = sample_provider("provider-kiro", "Kiro", 10); provider.provider_type = "kiro".to_string(); let build_key = |id: &str| { - let mut key = sample_key(id, "provider-kiro", "claude:messages", "__placeholder__"); + let mut key = sample_bound_key(id, "provider-kiro", "claude:messages", "__placeholder__"); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + id, + r#"{ "provider_type":"kiro", "auth_method":"idc", "access_token":"cached-kiro-token", @@ -3744,9 +3789,7 @@ async fn gateway_handles_admin_provider_query_test_model_failover_for_kiro_local "client_id":"client-id", "client_secret":"client-secret" }"#, - ) - .expect("auth config should encrypt"), - ); + )); key }; @@ -3890,12 +3933,12 @@ async fn gateway_retries_kiro_failover_after_http_error_without_message_impl() { let mut provider = sample_provider("provider-kiro", "Kiro", 10); provider.provider_type = "kiro".to_string(); let build_key = |id: &str| { - let mut key = sample_key(id, "provider-kiro", "claude:messages", "__placeholder__"); + let mut key = sample_bound_key(id, "provider-kiro", "claude:messages", "__placeholder__"); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-kiro", + id, + r#"{ "provider_type":"kiro", "auth_method":"idc", "access_token":"cached-kiro-token", @@ -3906,9 +3949,7 @@ async fn gateway_retries_kiro_failover_after_http_error_without_message_impl() { "client_id":"client-id", "client_secret":"client-secret" }"#, - ) - .expect("auth config should encrypt"), - ); + )); key }; @@ -4055,7 +4096,7 @@ async fn gateway_handles_non_kiro_multi_model_failover_locally_impl() { "openai:chat", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-primary", "provider-openai", "openai:chat", @@ -4235,7 +4276,7 @@ async fn gateway_handles_openai_responses_test_model_locally_impl() { "openai:responses", "https://tiger.bookapi.cc/codex", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", @@ -4367,7 +4408,7 @@ async fn gateway_handles_openai_image_test_model_locally_impl() { "openai:image", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-image", "provider-openai", "openai:image", @@ -4431,7 +4472,7 @@ async fn gateway_reports_transport_unsupported_reason_for_non_kiro_provider_impl "gemini:generate_content", "https://cloudcode-pa.googleapis.com", )], - vec![sample_key( + vec![sample_bound_key( "key-antigravity-gemini", "provider-antigravity", "gemini:generate_content", @@ -4582,26 +4623,24 @@ async fn gateway_handles_antigravity_endpoint_test_model_locally_impl() { let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-antigravity", "Antigravity", 10); provider.provider_type = "antigravity".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-antigravity-gemini", "provider-antigravity", "gemini:generate_content", "cached-antigravity-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{ + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-antigravity", + "key-antigravity-gemini", + r#"{ "provider_type":"antigravity", "project_id":"project-ant-123", "client_version":"1.2.3", "session_id":"sess-ant-123", "refresh_token":"rt-ant-123" }"#, - ) - .expect("auth config should encrypt"), - ); + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -4764,20 +4803,18 @@ async fn gateway_hydrates_antigravity_project_id_from_load_code_assist_for_test_ let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-antigravity", "Antigravity", 10); provider.provider_type = "antigravity".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-antigravity-gemini", "provider-antigravity", "gemini:generate_content", "cached-antigravity-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"antigravity","refresh_token":"rt-antigravity-123"}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-antigravity", + "key-antigravity-gemini", + r#"{"provider_type":"antigravity","refresh_token":"rt-antigravity-123"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -4906,13 +4943,13 @@ async fn gateway_prefers_supported_non_kiro_endpoint_when_api_format_is_omitted_ ), ], vec![ - sample_key( + sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", "sk-test-cli", ), - sample_key( + sample_bound_key( "key-openai-chat", "provider-openai", "openai:chat", @@ -5020,7 +5057,7 @@ async fn gateway_prefers_transport_supported_non_kiro_endpoint_when_api_format_i "https://api.openai.example/v1", ), ], - vec![sample_key( + vec![sample_bound_key( "key-openai-chat", "provider-openai", "openai:chat", @@ -5123,7 +5160,7 @@ async fn gateway_prefers_supported_non_kiro_endpoint_with_compatible_key_when_ap "https://api.openai.example/v1", ), ], - vec![sample_key( + vec![sample_bound_key( "key-openai-chat", "provider-openai", "openai:chat", @@ -5225,7 +5262,7 @@ async fn gateway_uses_compatible_cli_endpoint_when_api_format_is_omitted_impl() "https://api.openai.example/v1", ), ], - vec![sample_key( + vec![sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", @@ -5325,7 +5362,7 @@ async fn gateway_uses_runnable_cli_endpoint_after_chat_preference_when_api_forma "openai:responses", "https://api.openai.example/v1", ); - let mut shared_key = sample_key( + let mut shared_key = sample_bound_key( "key-openai-shared", "provider-openai", "openai:chat", @@ -5426,7 +5463,7 @@ async fn gateway_handles_openai_responses_test_model_failover_locally_impl() { "openai:responses", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", @@ -5527,7 +5564,7 @@ async fn gateway_handles_claude_cli_test_model_locally_impl() { "claude:messages", "https://api.anthropic.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-claude-cli", "provider-claude", "claude:messages", @@ -5622,7 +5659,7 @@ async fn gateway_uses_compatible_claude_cli_endpoint_when_api_format_is_omitted_ "claude:messages", "https://api.anthropic.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-claude-cli", "provider-claude", "claude:messages", @@ -5718,7 +5755,7 @@ async fn gateway_handles_claude_cli_test_model_failover_locally_impl() { "claude:messages", "https://api.anthropic.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-claude-cli", "provider-claude", "claude:messages", @@ -5822,7 +5859,7 @@ async fn gateway_handles_gemini_cli_test_model_locally_impl() { "gemini:generate_content", "https://generativelanguage.googleapis.com", )], - vec![sample_key( + vec![sample_bound_key( "key-gemini-cli", "provider-gemini", "gemini:generate_content", @@ -5931,20 +5968,18 @@ async fn gateway_handles_gemini_cli_test_model_with_oauth_header_fallback_impl() let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-gemini", "Gemini", 10); provider.provider_type = "gemini_cli".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-gemini-cli", "provider-gemini", "gemini:generate_content", "cached-gemini-cli-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"gemini_cli","project_id":"project-1"}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-gemini", + "key-gemini-cli", + r#"{"provider_type":"gemini_cli","project_id":"project-1"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -6086,20 +6121,18 @@ async fn gateway_hydrates_gemini_cli_project_id_from_load_code_assist_for_test_m let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-gemini", "Gemini", 10); provider.provider_type = "gemini_cli".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-gemini-cli", "provider-gemini", "gemini:generate_content", "cached-gemini-cli-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"gemini_cli","refresh_token":"rt-gemini-cli-123"}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-gemini", + "key-gemini-cli", + r#"{"provider_type":"gemini_cli","refresh_token":"rt-gemini-cli-123"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -6219,7 +6252,7 @@ async fn gateway_uses_compatible_gemini_cli_endpoint_when_api_format_is_omitted_ "gemini:generate_content", "https://generativelanguage.googleapis.com", )], - vec![sample_key( + vec![sample_bound_key( "key-gemini-cli", "provider-gemini", "gemini:generate_content", @@ -6315,7 +6348,7 @@ async fn gateway_handles_gemini_cli_test_model_failover_locally_impl() { "gemini:generate_content", "https://generativelanguage.googleapis.com", )], - vec![sample_key( + vec![sample_bound_key( "key-gemini-cli", "provider-gemini", "gemini:generate_content", @@ -6431,20 +6464,18 @@ async fn gateway_unwraps_gemini_cli_v1internal_response_for_failover_model_test_ let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let mut provider = sample_provider("provider-gemini-cli", "Gemini CLI", 10); provider.provider_type = "gemini_cli".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-gemini-cli", "provider-gemini-cli", "gemini:generate_content", "cached-gemini-cli-token", ); key.auth_type = "oauth".to_string(); - key.encrypted_auth_config = Some( - aether_crypto::encrypt_python_fernet_plaintext( - DEVELOPMENT_ENCRYPTION_KEY, - r#"{"provider_type":"gemini_cli","project_id":"project-1"}"#, - ) - .expect("auth config should encrypt"), - ); + key.encrypted_auth_config = Some(sample_bound_auth_config( + "provider-gemini-cli", + "key-gemini-cli", + r#"{"provider_type":"gemini_cli","project_id":"project-1"}"#, + )); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![provider], vec![sample_endpoint( @@ -6546,7 +6577,7 @@ async fn gateway_handles_admin_provider_query_test_model_failover_with_single_mo "openai:chat", "https://api.openai.example/v1", )], - vec![sample_key( + vec![sample_bound_key( "key-openai-alias", "provider-openai", "openai:chat", @@ -6667,13 +6698,13 @@ async fn gateway_retries_non_kiro_failover_after_http_error_without_message_impl "https://api.openai.example/v1", )], vec![ - sample_key( + sample_bound_key( "key-openai-first", "provider-openai", "openai:chat", "sk-test-first", ), - sample_key( + sample_bound_key( "key-openai-second", "provider-openai", "openai:chat", @@ -6794,13 +6825,13 @@ async fn gateway_retries_non_kiro_failover_after_success_status_without_body_imp "https://api.openai.example/v1", )], vec![ - sample_key( + sample_bound_key( "key-openai-first", "provider-openai", "openai:chat", "sk-test-first", ), - sample_key( + sample_bound_key( "key-openai-second", "provider-openai", "openai:chat", @@ -6850,7 +6881,9 @@ async fn gateway_retries_non_kiro_failover_after_success_status_without_body_imp assert_eq!(attempts[0]["status_code"], json!(200)); assert_eq!( attempts[0]["error_message"], - json!("Provider returned HTTP 200 without a model-test response body") + // Attempt diagnostics intentionally expose only the status class; + // detailed provider response text is not returned to the admin UI. + json!("HTTP 200") ); assert_eq!(attempts[1]["status"], json!("success")); diff --git a/apps/aether-gateway/src/tests/control/admin/providers.rs b/apps/aether-gateway/src/tests/control/admin/providers.rs index e65167bd8..2ad7f5fd7 100644 --- a/apps/aether-gateway/src/tests/control/admin/providers.rs +++ b/apps/aether-gateway/src/tests/control/admin/providers.rs @@ -1,7 +1,6 @@ use std::sync::{Arc, Mutex}; use std::time::{SystemTime, UNIX_EPOCH}; -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; use aether_data::repository::global_models::InMemoryGlobalModelReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; @@ -23,9 +22,10 @@ use serde_json::json; use super::super::{ build_router_with_state, issue_test_admin_access_token, sample_admin_provider_model, - sample_endpoint, sample_key, sample_provider, sample_provider_active_global_model, - sample_provider_model_stats, sample_provider_quota, sample_public_global_model_with_mappings, - sample_request_candidate, start_server, AppState, + sample_bound_key, sample_bound_key as sample_key, sample_bound_provider_proxy, sample_endpoint, + sample_provider, sample_provider_active_global_model, sample_provider_model_stats, + sample_provider_quota, sample_public_global_model_with_mappings, sample_request_candidate, + start_server, AppState, }; use crate::admin_api::{ maybe_build_local_admin_providers_response, AdminAppState, AdminRequestContext, @@ -233,7 +233,11 @@ async fn gateway_handles_admin_provider_summary_locally_with_trusted_admin_princ true, None, Some(4), - Some(json!({"host": "proxy.example", "password": "secret"})), + Some(sample_bound_provider_proxy( + "provider-openai", + "proxy.example", + "secret", + )), Some(45.0), Some(12.0), Some(json!({ @@ -269,25 +273,12 @@ async fn gateway_handles_admin_provider_summary_locally_with_trusted_admin_princ "sk-test-chat", ) .with_health_fields(Some(json!({"openai:chat": {"health_score": 0.25}})), None), - sample_key( + sample_bound_key( "key-openai-cli", "provider-openai", "openai:responses", - "sk-test-cli", + "sk-test-cli-2", ) - .with_transport_fields( - Some(json!(["openai:responses"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-test-cli-2") - .expect("api key ciphertext should build"), - None, - None, - None, - None, - None, - None, - None, - ) - .expect("key transport should build") .with_health_fields( Some(json!({"openai:responses": {"health_score": 0.75}})), None, @@ -780,6 +771,14 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() { let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![ sample_provider("provider-openai", "openai", 10) + .with_billing_fields( + Some("free_tier".to_string()), + None, + None, + Some(30), + None, + None, + ) .with_transport_fields( true, false, @@ -871,7 +870,7 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() { Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64) ); assert_eq!(payload["stream_first_byte_timeout"], 11.0); - assert_eq!(payload["proxy"], json!({"url": "https://proxy.example"})); + assert_eq!(payload["proxy"], json!({"url": "https://proxy.example/"})); assert_eq!(payload["claude_code_advanced"], json!({"pool_size": 2})); assert_eq!(payload["pool_advanced"], json!({})); assert_eq!(payload["failover_rules"], json!({"strategy": "ordered"})); @@ -974,6 +973,7 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() { .iter() .find(|provider| provider.id == "provider-openai") .expect("provider should exist"); + assert_eq!(updated_provider.billing_type.as_deref(), Some("free_tier")); assert_eq!( updated_provider.request_timeout_secs, Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64) @@ -1106,6 +1106,7 @@ async fn gateway_creates_admin_provider_locally_with_trusted_admin_principal() { .find(|provider| provider.id == "provider-existing") .expect("existing provider should remain"); assert_eq!(created.provider_type, "codex"); + assert_eq!(created.billing_type.as_deref(), Some("pay_as_you_go")); assert_eq!(created.provider_priority, 0); assert_eq!(existing.provider_priority, 1); assert_eq!(created.website.as_deref(), Some("https://codex.example")); @@ -1741,7 +1742,7 @@ async fn gateway_handles_admin_provider_mapping_preview_locally_with_trusted_adm let keys = payload["keys"].as_array().expect("keys should be an array"); assert_eq!(keys.len(), 1); assert_eq!(keys[0]["key_id"], "key-openai-preview"); - assert_eq!(keys[0]["masked_key"], "sk-p***1234"); + assert_eq!(keys[0]["masked_key"], "sk-p***234"); assert_eq!(keys[0]["allowed_models"], json!(["gpt-5", "gpt-4.1-mini"])); let matches = keys[0]["matching_global_models"] diff --git a/apps/aether-gateway/src/tests/control/admin/proxy_nodes.rs b/apps/aether-gateway/src/tests/control/admin/proxy_nodes.rs index 04c069278..d4078d131 100644 --- a/apps/aether-gateway/src/tests/control/admin/proxy_nodes.rs +++ b/apps/aether-gateway/src/tests/control/admin/proxy_nodes.rs @@ -1,11 +1,15 @@ use std::sync::{Arc, Mutex}; -use std::time::{SystemTime, UNIX_EPOCH}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; -use aether_data::repository::management_tokens::InMemoryManagementTokenRepository; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; +use aether_data::repository::management_tokens::{ + InMemoryManagementTokenRepository, ManagementTokenListQuery, ManagementTokenReadRepository, +}; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::proxy_nodes::{ InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, StoredProxyNodeEvent, }; +use aether_data::repository::users::InMemoryUserReadRepository; use axum::body::Body; use axum::extract::ws::Message; use axum::routing::any; @@ -16,8 +20,10 @@ use serde_json::json; use tokio::sync::watch; use super::super::{ - build_router_with_state, hash_management_token, sample_endpoint, sample_key, - sample_management_token, sample_provider, sample_proxy_node, start_server, AppState, + authenticated_tunnel_control_plane_request, build_router_with_state, hash_management_token, + sample_endpoint, sample_key, sample_management_token, sample_provider, sample_proxy_node, + start_server, with_tunnel_control_plane_key, AppState, TUNNEL_CONTROL_PLANE_TEST_GENERATION, + TUNNEL_CONTROL_PLANE_TEST_PSK, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -30,6 +36,16 @@ use crate::maintenance::{ }; use crate::tunnel::{tunnel_protocol, TunnelProxyConn}; +async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, +) -> Message { + tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) +} + #[tokio::test] async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -103,7 +119,8 @@ async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal( assert_eq!(items[0]["is_manual"], true); assert_eq!(items[0]["proxy_url"], "http://proxy.example:8080"); assert_eq!(items[0]["proxy_username"], "alice"); - assert_eq!(items[0]["proxy_password"], "su****et"); + assert_eq!(items[0]["has_proxy_password"], true); + assert!(items[0].get("proxy_password").is_none()); assert!(items[0]["created_at"].is_string()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -112,7 +129,7 @@ async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal( } #[tokio::test] -async fn gateway_returns_full_manual_proxy_node_detail_locally_with_trusted_admin_principal() { +async fn gateway_does_not_return_manual_proxy_password_in_node_detail() { let mut manual_node = sample_proxy_node("proxy-node-manual"); manual_node.name = "alpha-manual".to_string(); manual_node.status = "online".to_string(); @@ -149,7 +166,8 @@ async fn gateway_returns_full_manual_proxy_node_detail_locally_with_trusted_admi let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["node"]["id"], "proxy-node-manual"); assert_eq!(payload["node"]["proxy_username"], "alice"); - assert_eq!(payload["node"]["proxy_password"], "supersecret"); + assert_eq!(payload["node"]["has_proxy_password"], true); + assert!(payload["node"].get("proxy_password").is_none()); gateway_handle.abort(); } @@ -191,6 +209,7 @@ async fn gateway_reports_active_proxy_upgrade_rollout_in_proxy_node_list() { data_state .apply_proxy_node_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-alpha".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: None, total_requests_delta: None, @@ -354,6 +373,7 @@ async fn gateway_clears_proxy_upgrade_rollout_conflicts_locally() { .update_proxy_node_remote_config( &aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation { node_id: "node-beta".to_string(), + expected_tunnel_generation: None, node_name: None, allowed_ports: None, log_level: None, @@ -932,7 +952,177 @@ async fn gateway_rejects_management_token_without_required_admin_route_permissio } #[tokio::test] -async fn gateway_registers_proxy_node_with_management_token_when_allowed_ips_is_json_null() { +async fn proxy_node_install_session_requires_admin_and_preserves_parent_token_constraints() { + let write_raw_token = "ae-proxy-install-write-only"; + let admin_raw_token = "ae-proxy-install-admin"; + let state = AppState::new().expect("gateway should build"); + let admin_user = state + .create_local_auth_user_with_settings( + Some("proxy-install-admin@example.com".to_string()), + true, + "admin".to_string(), + "hash".to_string(), + "admin".to_string(), + None, + None, + None, + None, + ) + .await + .expect("admin user should be created") + .expect("admin user should exist"); + + let parent_allowed_ips = json!(["127.0.0.1"]); + let parent_expires_at = 4_102_444_800; + let mut write_parent = sample_management_token( + "token-proxy-install-write", + &admin_user.id, + "proxy-install-write", + true, + ); + write_parent.token.allowed_ips = Some(parent_allowed_ips.clone()); + write_parent.token.permissions = Some(json!(["admin:proxy_nodes:write"])); + write_parent.token.expires_at_unix_secs = Some(parent_expires_at); + let mut admin_parent = sample_management_token( + "token-proxy-install-admin", + &admin_user.id, + "proxy-install-admin", + true, + ); + admin_parent.token.allowed_ips = Some(parent_allowed_ips.clone()); + admin_parent.token.permissions = Some(json!(["admin:proxy_nodes:admin"])); + admin_parent.token.expires_at_unix_secs = Some(parent_expires_at); + + let management_token_repository = + Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![write_parent, admin_parent], + vec![ + ( + hash_management_token(write_raw_token), + "token-proxy-install-write".to_string(), + ), + ( + hash_management_token(admin_raw_token), + "token-proxy-install-admin".to_string(), + ), + ], + )); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([ + admin_user.clone() + ])); + let state = state.with_data_state_for_tests( + GatewayDataState::with_management_token_repository_for_tests(Arc::clone( + &management_token_repository, + )) + .with_user_reader(user_repository), + ); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let denied = client + .post(format!( + "{gateway_url}/api/admin/proxy-nodes/install-sessions" + )) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(write_raw_token) + .json(&json!({ "node_name": "write-only-node" })) + .send() + .await + .expect("write-only install request should complete"); + assert_eq!(denied.status(), StatusCode::FORBIDDEN); + let denied_payload: serde_json::Value = denied.json().await.expect("denial should be json"); + assert_eq!( + denied_payload["required_permission"], + json!("admin:proxy_nodes:admin") + ); + + let accepted = client + .post(format!( + "{gateway_url}/api/admin/proxy-nodes/install-sessions" + )) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(admin_raw_token) + .json(&json!({ "node_name": "constrained-node" })) + .send() + .await + .expect("admin install request should complete"); + let accepted_status = accepted.status(); + let accepted_body = accepted.text().await.expect("response body should read"); + assert_eq!( + accepted_status, + StatusCode::OK, + "unexpected response body: {accepted_body}" + ); + let accepted_payload: serde_json::Value = + serde_json::from_str(&accepted_body).expect("install response should be JSON"); + let install_code = accepted_payload["install_code"] + .as_str() + .expect("install code should be returned") + .to_string(); + + let tokens = management_token_repository + .list_management_tokens(&ManagementTokenListQuery { + user_id: Some(admin_user.id.clone()), + is_active: None, + offset: 0, + limit: 10, + }) + .await + .expect("management tokens should list"); + assert_eq!(tokens.total, 3); + let child = tokens + .items + .iter() + .find(|item| { + !matches!( + item.token.id.as_str(), + "token-proxy-install-write" | "token-proxy-install-admin" + ) + }) + .expect("install session should create one child management token"); + assert_eq!(child.token.allowed_ips, Some(parent_allowed_ips)); + assert_eq!(child.token.expires_at_unix_secs, Some(parent_expires_at)); + assert_eq!( + child.token.permissions, + Some(json!(["admin:proxy_nodes:write"])) + ); + assert!( + !child.token.is_active, + "unused install-session token must remain disabled" + ); + let child_id = child.token.id.clone(); + + // The in-memory repository cannot atomically verify the administrator row and therefore + // rejects one-time activation. SQL-backed repositories cover the successful atomic path. + let install_script = client + .get(format!("{gateway_url}/install-tunnel/{install_code}")) + .send() + .await + .expect("tunnel install script should receive a response"); + assert_eq!(install_script.status(), StatusCode::NOT_FOUND); + + let consumed = management_token_repository + .get_management_token_with_user(&child_id) + .await + .expect("consumed token lookup should succeed"); + assert!( + consumed.is_none(), + "failed one-time activation must discard the pending bearer" + ); + + let replay = client + .get(format!("{gateway_url}/install-tunnel/{install_code}")) + .send() + .await + .expect("tunnel install replay should receive a response"); + assert_eq!(replay.status(), StatusCode::NOT_FOUND); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_proxy_node_registration_when_allowed_ips_is_json_null() { let raw_token = "ae_proxy_register_json_null"; let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::default()); let state = AppState::new().expect("gateway should build"); @@ -989,16 +1179,11 @@ async fn gateway_registers_proxy_node_with_management_token_when_allowed_ips_is_ .await .expect("request should succeed"); - assert_eq!(register_response.status(), StatusCode::OK); - let register_payload: serde_json::Value = register_response - .json() + assert_eq!(register_response.status(), StatusCode::UNAUTHORIZED); + let _ = register_response + .bytes() .await - .expect("json body should parse"); - assert_eq!(register_payload["node"]["name"], "proxy-json-null"); - assert_eq!( - register_payload["node"]["registered_by"], - json!(admin_user.id) - ); + .expect("error body should be readable"); gateway_handle.abort(); } @@ -1071,7 +1256,8 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() { assert_eq!(create_payload["node"]["status"], "online"); assert_eq!(create_payload["node"]["proxy_url"], proxy_url); assert_eq!(create_payload["node"]["proxy_username"], "alice"); - assert_eq!(create_payload["node"]["proxy_password"], "su****et"); + assert_eq!(create_payload["node"]["has_proxy_password"], true); + assert!(create_payload["node"].get("proxy_password").is_none()); let test_url_response = client .post(format!("{gateway_url}/api/admin/proxy-nodes/test-url")) @@ -1202,22 +1388,25 @@ async fn gateway_tests_connected_tunnel_proxy_nodes_with_active_probe() { "https://probe.example/cdn-cgi/trace", ); - let mut node = sample_proxy_node("node-online"); + let mut node = with_tunnel_control_plane_key( + sample_proxy_node("node-online"), + TUNNEL_CONTROL_PLANE_TEST_PSK, + ); node.status = "online".to_string(); node.tunnel_connected = true; let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node])); let state = AppState::new() .expect("gateway should build") - .with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests( - proxy_node_repository, - )); + .with_data_state_for_tests( + GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); let tunnel_state = state.tunnel.app_state(); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - tunnel_state - .hub - .register_proxy(Arc::new(TunnelProxyConn::new( + tunnel_state.hub.register_proxy(Arc::new( + TunnelProxyConn::new( 500, "node-online".to_string(), "Node Online".to_string(), @@ -1225,7 +1414,10 @@ async fn gateway_tests_connected_tunnel_proxy_nodes_with_active_probe() { proxy_close_tx, 16, 2, - ))); + ) + .with_tunnel_generation(TUNNEL_CONTROL_PLANE_TEST_GENERATION.to_string()) + .with_authenticated_key(TUNNEL_CONTROL_PLANE_TEST_PSK.to_string()), + )); let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -1246,7 +1438,7 @@ async fn gateway_tests_connected_tunnel_proxy_nodes_with_active_probe() { } }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "probe headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -1261,7 +1453,7 @@ async fn gateway_tests_connected_tunnel_proxy_nodes_with_active_probe() { assert_eq!(meta.url, "https://probe.example/cdn-cgi/trace"); assert_eq!(meta.follow_redirects, Some(false)); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "probe body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -1602,10 +1794,9 @@ async fn gateway_handles_admin_proxy_node_events_locally_with_trusted_admin_prin #[tokio::test] async fn gateway_reports_proxy_node_metrics_and_filters_events_locally() { - let proxy_node_repository = - Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( - "node-1", - )])); + let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + with_tunnel_control_plane_key(sample_proxy_node("node-1"), TUNNEL_CONTROL_PLANE_TEST_PSK), + ])); let gateway = build_router_with_state( AppState::new() .expect("gateway should build") @@ -1620,61 +1811,73 @@ async fn gateway_reports_proxy_node_metrics_and_filters_events_locally() { .expect("system time should be after epoch") .as_secs(); - let baseline_heartbeat_response = client - .post(format!("{gateway_url}/api/internal/tunnel/heartbeat")) - .json(&json!({ - "node_id": "node-1", - "heartbeat_id": 90, - "heartbeat_interval": 30, - "active_connections": 0, - "proxy_metadata": { - "tunnel_metrics": { - "connect_errors": 0, - "disconnects": 0, - "error_events_total": 0, - "ws_in_bytes": 0, - "ws_out_bytes": 0, - "ws_in_frames": 0, - "ws_out_frames": 0, - "heartbeat_rtt_last_ms": 0 - } - }, - "proxy_version": "2.0.0" - })) - .send() - .await - .expect("baseline heartbeat request should succeed"); + let baseline_heartbeat = json!({ + "node_id": "node-1", + "heartbeat_session_id": "node-1-session", + "heartbeat_id": 90, + "heartbeat_interval": 30, + "active_connections": 0, + "proxy_metadata": { + "tunnel_metrics": { + "connect_errors": 0, + "disconnects": 0, + "error_events_total": 0, + "ws_in_bytes": 0, + "ws_out_bytes": 0, + "ws_in_frames": 0, + "ws_out_frames": 0, + "heartbeat_rtt_last_ms": 0 + } + }, + "proxy_version": "2.0.0" + }); + let baseline_heartbeat_response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}/api/internal/tunnel/heartbeat"), + "/api/internal/tunnel/heartbeat", + "node-1", + &baseline_heartbeat, + ) + .send() + .await + .expect("baseline heartbeat request should succeed"); assert_eq!(baseline_heartbeat_response.status(), StatusCode::OK); - let heartbeat_response = client - .post(format!("{gateway_url}/api/internal/tunnel/heartbeat")) - .json(&json!({ - "node_id": "node-1", - "heartbeat_id": 91, - "heartbeat_interval": 30, - "active_connections": 7, - "proxy_metadata": { - "tunnel_metrics": { - "connect_errors": 3, - "disconnects": 1, - "error_events_total": 1, - "ws_in_bytes": 1000, - "ws_out_bytes": 2000, - "ws_in_frames": 10, - "ws_out_frames": 20, - "heartbeat_rtt_last_ms": 42 - }, - "recent_tunnel_errors": [{ - "timestamp_unix_secs": now_unix_secs, - "category": "tcp_connect_timeout", - "message": "tunnel TCP connect timeout" - }] + let heartbeat = json!({ + "node_id": "node-1", + "heartbeat_session_id": "node-1-session", + "heartbeat_id": 91, + "heartbeat_interval": 30, + "active_connections": 7, + "proxy_metadata": { + "tunnel_metrics": { + "connect_errors": 3, + "disconnects": 1, + "error_events_total": 1, + "ws_in_bytes": 1000, + "ws_out_bytes": 2000, + "ws_in_frames": 10, + "ws_out_frames": 20, + "heartbeat_rtt_last_ms": 42 }, - "proxy_version": "2.0.0" - })) - .send() - .await - .expect("heartbeat request should succeed"); + "recent_tunnel_errors": [{ + "timestamp_unix_secs": now_unix_secs, + "category": "tcp_connect_timeout", + "message": "tunnel TCP connect timeout" + }] + }, + "proxy_version": "2.0.0" + }); + let heartbeat_response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}/api/internal/tunnel/heartbeat"), + "/api/internal/tunnel/heartbeat", + "node-1", + &heartbeat, + ) + .send() + .await + .expect("heartbeat request should succeed"); assert_eq!(heartbeat_response.status(), StatusCode::OK); let from = now_unix_secs.saturating_sub(120); @@ -1794,6 +1997,7 @@ async fn gateway_updates_proxy_node_config_and_dispatches_upgrade_targets_locall ); let mut online_node = sample_proxy_node("node-online"); + online_node = with_tunnel_control_plane_key(online_node, TUNNEL_CONTROL_PLANE_TEST_PSK); online_node.status = "online".to_string(); online_node.tunnel_connected = true; let mut online_node_2 = sample_proxy_node("node-zeta"); @@ -1895,21 +2099,27 @@ async fn gateway_updates_proxy_node_config_and_dispatches_upgrade_targets_locall assert_eq!(blocked_upgrade_payload["updated"], 0); assert_eq!(blocked_upgrade_payload["skipped"], 3); - let heartbeat_response = client - .post(format!("{gateway_url}/api/internal/tunnel/heartbeat")) - .json(&json!({ - "node_id": "node-online", - "heartbeat_id": 77, - "heartbeat_interval": 45, - "active_connections": 3, - "total_requests": 5, - "avg_latency_ms": 10.0, - "proxy_metadata": { "arch": "arm64" }, - "proxy_version": "2.0.0" - })) - .send() - .await - .expect("request should succeed"); + let heartbeat = json!({ + "node_id": "node-online", + "heartbeat_session_id": "node-online-session", + "heartbeat_id": 77, + "heartbeat_interval": 45, + "active_connections": 3, + "total_requests": 5, + "avg_latency_ms": 10.0, + "proxy_metadata": { "arch": "arm64" }, + "proxy_version": "2.0.0" + }); + let heartbeat_response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}/api/internal/tunnel/heartbeat"), + "/api/internal/tunnel/heartbeat", + "node-online", + &heartbeat, + ) + .send() + .await + .expect("request should succeed"); assert_eq!(heartbeat_response.status(), StatusCode::OK); let heartbeat_payload: serde_json::Value = heartbeat_response .json() diff --git a/apps/aether-gateway/src/tests/control/admin/security.rs b/apps/aether-gateway/src/tests/control/admin/security.rs index 2e8d1cded..2df981435 100644 --- a/apps/aether-gateway/src/tests/control/admin/security.rs +++ b/apps/aether-gateway/src/tests/control/admin/security.rs @@ -7,6 +7,9 @@ use http::{HeaderMap, HeaderValue, StatusCode}; use http_body_util::BodyExt; use serde_json::json; +use aether_runtime_state::{RedisClientConfig, RuntimeState}; +use aether_test_support::ManagedRedisServer; + use super::super::super::send_request; use super::super::{build_router_with_state, start_server, AppState}; use crate::admin_api::{ @@ -108,6 +111,59 @@ async fn gateway_blocks_forwarded_ip_from_trusted_proxy() { assert_eq!(response.status(), StatusCode::FORBIDDEN); } +#[tokio::test] +async fn gateway_fails_closed_when_ip_blacklist_state_is_unavailable() { + let mut redis = match ManagedRedisServer::start().await { + Ok(redis) => redis, + Err(error) if error.to_string().contains("No such file or directory") => { + eprintln!("skipping IP blacklist Redis outage test: {error}"); + return; + } + Err(error) => panic!("Redis test server should start: {error}"), + }; + let runtime_state = Arc::new( + RuntimeState::redis( + RedisClientConfig { + url: redis.redis_url().to_string(), + key_prefix: Some(format!("blacklist-outage-{}", std::process::id())), + }, + Some(250), + ) + .await + .expect("Redis runtime state should build"), + ); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_runtime_state(runtime_state), + ); + redis.stop().expect("Redis test server should stop"); + + let request = Request::builder() + .uri("/api/public/system") + .body(Body::empty()) + .expect("request should build"); + let response = send_request(gateway, request).await; + + assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!( + response + .headers() + .get(crate::constants::EXECUTION_PATH_HEADER) + .and_then(|value| value.to_str().ok()), + Some(crate::constants::EXECUTION_PATH_LOCAL_AUTH_DENIED) + ); + let payload = response + .into_body() + .collect() + .await + .expect("body should collect") + .to_bytes(); + let payload: serde_json::Value = + serde_json::from_slice(&payload).expect("response should be json"); + assert_eq!(payload["error"]["message"], "IP 访问控制暂时不可用"); +} + #[tokio::test] async fn admin_security_whitelist_matches_cidr() { let state = AppState::new() diff --git a/apps/aether-gateway/src/tests/control/admin/system.rs b/apps/aether-gateway/src/tests/control/admin/system.rs index d2763f71a..a60973a96 100644 --- a/apps/aether-gateway/src/tests/control/admin/system.rs +++ b/apps/aether-gateway/src/tests/control/admin/system.rs @@ -23,21 +23,28 @@ use axum::routing::{any, delete, get, post, put}; use axum::{extract::Request, Router}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use super::super::{ build_router_with_state, issue_test_admin_access_token, sample_admin_global_model, - sample_admin_provider_model, sample_endpoint, sample_key, sample_ldap_module_config, - sample_oauth_provider_config, sample_provider, sample_proxy_node, + sample_admin_provider_model, sample_bound_key, sample_endpoint, sample_key, + sample_ldap_module_config, sample_oauth_provider_config, sample_provider, sample_proxy_node, sample_recent_key_rpm_candidate, start_server, AppState, }; +use crate::admin_api::AdminAppState; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, TRUSTED_ADMIN_USER_ROLE_HEADER, }; use crate::data::GatewayDataState; +use crate::handlers::admin::SystemExportMode; static SYSTEM_UPDATE_TEST_MUTEX: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); +fn sha256_hex(value: &str) -> String { + format!("{:x}", Sha256::digest(value.as_bytes())) +} + #[tokio::test] async fn gateway_handles_admin_system_version_locally_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -231,7 +238,7 @@ async fn gateway_prepares_admin_system_update_locally() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::PRECONDITION_REQUIRED); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert!(payload["detail"].is_string()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -270,7 +277,7 @@ async fn gateway_rejects_admin_system_apply_update_without_prepared_version() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(response.status(), StatusCode::PRECONDITION_REQUIRED); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert!(payload["detail"].is_string()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -391,7 +398,7 @@ async fn gateway_rejects_admin_system_apply_update_with_nonexistent_version() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(response.status(), StatusCode::PRECONDITION_REQUIRED); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert!(payload["detail"].is_string()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -747,9 +754,14 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p "gpt-5", )]), ); + let mut ldap_config = sample_ldap_module_config(); + ldap_config.bind_password_encrypted = Some( + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "ldap-bind-secret") + .expect("LDAP password should encrypt"), + ); let auth_module_repository = Arc::new(InMemoryAuthModuleReadRepository::seed( Vec::::new(), - Some(sample_ldap_module_config()), + Some(ldap_config), )); let oauth_provider_repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![ sample_oauth_provider_config("linuxdo"), @@ -767,27 +779,109 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p let data_state = GatewayDataState::disabled() .attach_provider_catalog_repository_for_tests(provider_catalog_repository) .with_global_model_repository_for_tests(global_model_repository) - .attach_auth_module_reader_for_tests(auth_module_repository) + .attach_auth_module_repository_for_tests(auth_module_repository) .attach_oauth_provider_repository_for_tests(oauth_provider_repository) .attach_proxy_node_repository_for_tests(proxy_node_repository) .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) .with_system_config_values_for_tests(vec![ - ( - "smtp_password".to_string(), - json!( - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "smtp-secret",) - .expect("smtp secret should encrypt") - ), - ), + ("smtp_password".to_string(), json!("smtp-secret")), + ("smtp_host".to_string(), json!("smtp.example.test")), + ("turnstile_secret_key".to_string(), serde_json::Value::Null), + ("smtp_user".to_string(), json!("smtp-user")), ("site_name".to_string(), json!("Aether Test")), ]); - let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state( - AppState::new() - .expect("gateway should build") - .with_data_state_for_tests(data_state), + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state); + let recovery_payload = AdminAppState::new(&state) + .build_admin_system_config_export_payload(SystemExportMode::RecoveryBackup) + .await + .expect("recovery config export should build"); + assert!(recovery_payload.get("credential_state").is_none()); + assert_eq!( + recovery_payload["providers"][0]["config"]["provider_ops"]["connector"]["credentials"] + ["refresh_token"], + "provider-refresh-token" ); + assert_eq!( + recovery_payload["providers"][0]["api_keys"][0]["api_key"], + "live-api-key" + ); + assert_eq!( + recovery_payload["providers"][0]["api_keys"][0]["auth_config"], + r#"{"refresh_token":"oauth-refresh"}"# + ); + assert_eq!( + recovery_payload["providers"][0]["api_keys"][0]["is_active"], + true + ); + assert_eq!( + recovery_payload["ldap_config"]["bind_password"], + "ldap-bind-secret" + ); + assert_eq!(recovery_payload["ldap_config"]["is_enabled"], true); + assert_eq!( + recovery_payload["oauth_providers"][0]["client_secret"], + "secret-value" + ); + assert_eq!(recovery_payload["oauth_providers"][0]["is_enabled"], true); + assert_eq!( + recovery_payload["proxy_nodes"][0]["proxy_username"], + "proxy-user" + ); + assert_eq!( + recovery_payload["proxy_nodes"][0]["proxy_password"], + "proxy-pass" + ); + let recovery_system_configs = recovery_payload["system_configs"] + .as_array() + .expect("recovery system configs should be an array"); + assert_eq!( + recovery_system_configs + .iter() + .find(|entry| entry["key"] == "smtp_user") + .expect("SMTP user should be recoverable")["value"], + "smtp-user" + ); + assert_eq!( + recovery_system_configs + .iter() + .find(|entry| entry["key"] == "smtp_password") + .expect("SMTP password should be recoverable")["value"], + "smtp-secret" + ); + assert_eq!( + recovery_system_configs + .iter() + .find(|entry| entry["key"] == "turnstile_secret_key") + .expect("unset Turnstile secret should remain recoverable")["value"], + serde_json::Value::Null + ); + + let rollback_checkpoint = AdminAppState::new(&state) + .build_admin_system_config_export_payload(SystemExportMode::RollbackCheckpoint) + .await + .expect("rollback config checkpoint should build"); + assert_eq!( + rollback_checkpoint["providers"][0]["api_keys"][0]["is_active"], + true + ); + assert_eq!( + rollback_checkpoint["providers"][0]["api_keys"][0]["credential_state"], + "not_exported" + ); + assert!(rollback_checkpoint["providers"][0]["api_keys"][0] + .get("api_key") + .is_none()); + assert_eq!(rollback_checkpoint["ldap_config"]["is_enabled"], true); + assert_eq!( + rollback_checkpoint["oauth_providers"][0]["is_enabled"], + true + ); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() @@ -801,25 +895,35 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["version"], "2.3"); assert!(payload["exported_at"].as_str().is_some()); assert_eq!(payload["global_models"][0]["name"], "gpt-5"); assert_eq!(payload["global_models"][0]["usage_count"], json!(7)); assert_eq!(payload["providers"][0]["name"], "openai"); + assert_eq!(payload["credential_state"], "not_exported"); assert_eq!( - payload["providers"][0]["config"]["provider_ops"]["connector"]["credentials"] - ["refresh_token"], - "provider-refresh-token" + payload["providers"][0]["config"]["provider_ops"]["connector"]["credentials"], + "***" ); + assert!(payload["providers"][0]["api_keys"][0] + .get("api_key") + .is_none()); + assert!(payload["providers"][0]["api_keys"][0] + .get("auth_config") + .is_none()); assert_eq!( - payload["providers"][0]["api_keys"][0]["api_key"], - "live-api-key" - ); - assert_eq!( - payload["providers"][0]["api_keys"][0]["auth_config"], - r#"{"refresh_token":"oauth-refresh"}"# + payload["providers"][0]["api_keys"][0]["credential_state"], + "not_exported" ); + assert_eq!(payload["providers"][0]["api_keys"][0]["is_active"], false); assert_eq!( payload["providers"][0]["api_keys"][0]["supported_endpoints"], json!(["openai:chat"]) @@ -828,29 +932,77 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p payload["providers"][0]["models"][0]["global_model_name"], "gpt-5" ); - assert_eq!(payload["ldap_config"]["bind_password"], ""); - assert_eq!( - payload["oauth_providers"][0]["client_secret"], - "secret-value" - ); + assert!(payload["ldap_config"].get("bind_password").is_none()); + assert_eq!(payload["ldap_config"]["is_enabled"], false); + assert!(payload["oauth_providers"][0].get("client_secret").is_none()); + assert_eq!(payload["oauth_providers"][0]["is_enabled"], false); assert_eq!( payload["proxy_nodes"][0]["proxy_url"], "http://proxy.local:8080" ); + assert!(payload["proxy_nodes"][0].get("proxy_username").is_none()); + assert!(payload["proxy_nodes"][0].get("proxy_password").is_none()); let smtp_password = payload["system_configs"] .as_array() .expect("system configs should be array") .iter() .find(|entry| entry["key"] == "smtp_password") - .cloned() - .expect("smtp_password should exist"); - assert_eq!(smtp_password["value"], "smtp-secret"); + .cloned(); + assert!(smtp_password.is_none()); + let serialized = payload.to_string(); + for secret in [ + "provider-refresh-token", + "live-api-key", + "oauth-refresh", + "secret-value", + "proxy-user", + "proxy-pass", + "smtp-user", + "smtp-secret", + ] { + assert!(!serialized.contains(secret), "leaked secret: {secret}"); + } assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); upstream_handle.abort(); } +#[tokio::test] +async fn gateway_rejects_sensitive_system_exports_for_audit_admin() { + let gateway = build_router_with_state(AppState::new().expect("gateway should build")); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + for path in [ + "/api/admin/system/config/export", + "/api/admin/system/users/export", + "/api/admin/system/data/export", + ] { + let response = reqwest::Client::new() + .get(format!("{gateway_url}{path}")) + .header(GATEWAY_HEADER, "rust-phase3b") + .header(TRUSTED_ADMIN_USER_ID_HEADER, "audit-admin-123") + .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "audit_admin") + .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-audit-123") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!( + payload["detail"], "management token permission denied", + "path: {path}" + ); + assert_eq!( + payload["required_permission"], "admin:system:admin", + "path: {path}" + ); + } + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -917,7 +1069,7 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr StoredAuthApiKeyExportRecord::new( "user-1".to_string(), "key-user-1".to_string(), - "hash-user-1".to_string(), + sha256_hex("ak-user-live-1"), Some( encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "ak-user-live-1") .expect("user api key should encrypt"), @@ -941,7 +1093,7 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr StoredAuthApiKeyExportRecord::new( "admin-owner".to_string(), "key-standalone-1".to_string(), - "hash-standalone-1".to_string(), + sha256_hex("ak-standalone-live-1"), Some( encrypt_python_fernet_plaintext( DEVELOPMENT_ENCRYPTION_KEY, @@ -1001,17 +1153,59 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr ) .expect("standalone wallet should build"), ])); - let data_state = - GatewayDataState::with_auth_and_wallet_for_tests(auth_repository, wallet_repository) - .with_user_reader(user_repository) - .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let data_state = GatewayDataState::with_auth_and_wallet_for_tests( + auth_repository.clone(), + wallet_repository, + ) + .attach_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state); + let recovery_payload = AdminAppState::new(&state) + .build_admin_system_users_export_payload(SystemExportMode::RecoveryBackup) + .await + .expect("recovery users export should build"); + let migrated_keys = state + .list_auth_api_key_export_records_by_ids(&[ + "key-user-1".to_string(), + "key-standalone-1".to_string(), + ]) + .await + .expect("migrated API keys should reload"); + assert_eq!(migrated_keys.len(), 2); + assert!(migrated_keys.iter().all(|key| key + .key_encrypted + .as_deref() + .is_some_and(|value| value.starts_with("aether-auth-api-key-secret-v2:")))); + assert_eq!(recovery_payload["version"], "1.5"); + assert_eq!(recovery_payload["users"][0]["password_hash"], "argon2-hash"); + assert_eq!( + recovery_payload["users"][0]["api_keys"][0]["key_hash"], + sha256_hex("ak-user-live-1") + ); + assert_eq!( + recovery_payload["users"][0]["api_keys"][0]["key"], + "ak-user-live-1" + ); + assert_eq!( + recovery_payload["users"][0]["api_keys"][0]["is_active"], + true + ); + assert_eq!( + recovery_payload["standalone_keys"][0]["key_hash"], + sha256_hex("ak-standalone-live-1") + ); + assert_eq!( + recovery_payload["standalone_keys"][0]["key"], + "ak-standalone-live-1" + ); + assert_eq!(recovery_payload["standalone_keys"][0]["is_active"], true); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state( - AppState::new() - .expect("gateway should build") - .with_data_state_for_tests(data_state), - ); + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() @@ -1025,8 +1219,15 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["version"], "1.5"); + assert_eq!(payload["version"], "1.6"); assert!(payload["exported_at"].as_str().is_some()); assert_eq!(payload["user_groups"][0]["name"], "Restricted GPT"); assert!(payload["user_groups"][0].get("priority").is_none()); @@ -1035,6 +1236,8 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr json!(["gpt-5"]) ); assert_eq!(payload["users"][0]["email"], "alice@example.com"); + assert!(payload["users"][0].get("password_hash").is_none()); + assert!(!payload.to_string().contains("argon2-hash")); assert_eq!( payload["users"][0]["allowed_models_mode"], json!("specific") @@ -1058,11 +1261,29 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr json!(10.0) ); assert_eq!(payload["users"][0]["unlimited"], json!(false)); - assert_eq!(payload["users"][0]["api_keys"][0]["key"], "ak-user-live-1"); assert_eq!( - payload["users"][0]["api_keys"][0]["key_hash"], - "hash-user-1" + payload["users"][0]["api_keys"][0]["credential_state"], + "not_exported" ); + assert_eq!(payload["users"][0]["api_keys"][0]["is_active"], false); + for credential_field in ["key", "key_hash", "key_encrypted"] { + assert!(payload["users"][0]["api_keys"][0] + .get(credential_field) + .is_none()); + assert!(payload["standalone_keys"][0] + .get(credential_field) + .is_none()); + } + let serialized = payload.to_string(); + for secret in ["ak-user-live-1", "ak-standalone-live-1"] { + assert!(!serialized.contains(secret)); + } + for secret_hash in [ + sha256_hex("ak-user-live-1"), + sha256_hex("ak-standalone-live-1"), + ] { + assert!(!serialized.contains(&secret_hash)); + } assert_eq!( payload["users"][0]["api_keys"][0]["is_standalone"], json!(false) @@ -1076,9 +1297,10 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr json!(420) ); assert_eq!( - payload["standalone_keys"][0]["key"], - json!("ak-standalone-live-1") + payload["standalone_keys"][0]["credential_state"], + "not_exported" ); + assert_eq!(payload["standalone_keys"][0]["is_active"], false); assert_eq!( payload["standalone_keys"][0]["api_key_id"], json!("key-standalone-1") @@ -1948,7 +2170,46 @@ async fn gateway_validates_chat_pii_redaction_system_config_locally_with_trusted } #[tokio::test] -async fn gateway_handles_admin_system_provider_priority_mode_locally_with_bearer_admin_session() { +async fn gateway_handles_admin_system_config_locally_with_bearer_admin_session() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().route( + "/api/admin/system/configs/site_name", + any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + }), + ); + + let (_upstream_url, upstream_handle) = start_server(upstream).await; + let state = AppState::new().expect("gateway should build"); + let access_token = issue_test_admin_access_token(&state, "device-admin-config").await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/api/admin/system/configs/site_name")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-admin-config") + .send() + .await + .expect("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["key"], "site_name"); + assert_eq!(payload["value"], "Aether"); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_removed_admin_system_provider_priority_mode_with_bearer_admin_session() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( @@ -1962,7 +2223,7 @@ async fn gateway_handles_admin_system_provider_priority_mode_locally_with_bearer }), ); - let (upstream_url, upstream_handle) = start_server(upstream).await; + let (_upstream_url, upstream_handle) = start_server(upstream).await; let state = AppState::new().expect("gateway should build"); let access_token = issue_test_admin_access_token(&state, "device-admin-config").await; let gateway = build_router_with_state(state); @@ -1978,10 +2239,7 @@ async fn gateway_handles_admin_system_provider_priority_mode_locally_with_bearer .await .expect("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["key"], "provider_priority_mode"); - assert_eq!(payload["value"], "provider"); + assert_eq!(response.status(), StatusCode::NOT_FOUND); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -2003,8 +2261,12 @@ async fn gateway_sets_admin_system_config_locally_with_trusted_admin_principal() }), ); - let data_state = - GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let data_state = GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests(vec![ + ("smtp_host".to_string(), json!("smtp.example.com")), + ("smtp_user".to_string(), json!("smtp-user")), + ]); let (upstream_url, upstream_handle) = start_server(upstream).await; let gateway = build_router_with_state( AppState::new() @@ -2146,7 +2408,7 @@ async fn gateway_handles_admin_key_rpm_locally_with_trusted_admin_principal() { "https://api.openai.example", )], vec![ - sample_key("key-openai", "provider-openai", "openai:chat", "sk-test") + sample_bound_key("key-openai", "provider-openai", "openai:chat", "sk-test") .with_rate_limit_fields(Some(60), None, None, None, None, None, None, None, None), ], )); @@ -2233,7 +2495,7 @@ async fn gateway_resets_admin_key_rpm_locally_with_trusted_admin_principal() { "https://api.openai.example", )], vec![ - sample_key("key-openai", "provider-openai", "openai:chat", "sk-test") + sample_bound_key("key-openai", "provider-openai", "openai:chat", "sk-test") .with_rate_limit_fields(Some(60), None, None, None, None, None, None, None, None), ], )); 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 1a436ce44..3919c3f58 100644 --- a/apps/aether-gateway/src/tests/control/admin/system_import.rs +++ b/apps/aether-gateway/src/tests/control/admin/system_import.rs @@ -1,9 +1,8 @@ use std::sync::{Arc, Mutex}; +use std::time::{SystemTime, UNIX_EPOCH}; use aether_contracts::ExecutionPlan; -use aether_crypto::{ - decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, -}; +use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, }; @@ -41,11 +40,38 @@ use super::super::{ use crate::ai_serving::{ build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope, }; +use crate::backup::executor::{encrypt_backup_bytes, BackupDecryptionKey, BackupRestoreLimits}; +use crate::backup::{apply_restored_backup, BackupRestoreScope}; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, TRUSTED_ADMIN_USER_ROLE_HEADER, }; use crate::data::GatewayDataState; +use crate::handlers::shared::{ + decrypt_or_migrate_identity_oauth_provider_client_secret, + decrypt_or_migrate_ldap_bind_password, open_auth_api_key_secret, + open_provider_catalog_credential, ProviderCatalogCredentialField, +}; +use crate::restore_backup_json; + +fn decrypt_test_provider_catalog_credential( + key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey, + field: ProviderCatalogCredentialField, +) -> String { + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let ciphertext = match field { + ProviderCatalogCredentialField::ApiKey => key.encrypted_api_key.as_deref(), + ProviderCatalogCredentialField::AuthConfig => key.encrypted_auth_config.as_deref(), + } + .expect("provider credential should be present"); + open_provider_catalog_credential(&state, &key.provider_id, &key.id, field, ciphertext) + .expect("provider credential should decrypt with its record binding") + .plaintext +} fn build_admin_system_data_state_with_repositories( provider_catalog_repository: Arc, @@ -202,6 +228,16 @@ fn sample_system_import_payload() -> Value { "value": "Imported Aether", "description": "Site name" }, + { + "key": "smtp_host", + "value": "smtp.example.com", + "description": "SMTP host" + }, + { + "key": "smtp_user", + "value": "smtp-user", + "description": "SMTP user" + }, { "key": "smtp_password", "value": "smtp-secret", @@ -273,6 +309,31 @@ fn sample_import_admin_user(user_id: &str) -> StoredUserAuthRecord { .expect("admin user should build") } +fn sample_import_user_session( + user_id: &str, + session_id: &str, +) -> crate::data::state::StoredUserSessionRecord { + let now = chrono::Utc::now(); + crate::data::state::StoredUserSessionRecord::new( + session_id.to_string(), + user_id.to_string(), + format!("device-{session_id}"), + None, + format!("refresh-{session_id}"), + None, + None, + Some(now), + Some(now + chrono::Duration::days(7)), + None, + None, + Some("127.0.0.1".to_string()), + Some("system-import-test".to_string()), + Some(now), + Some(now), + ) + .expect("session should build") +} + const ADMIN_SYSTEM_IMPORT_TEST_STACK_BYTES: usize = 16 * 1024 * 1024; fn run_admin_system_import_test(test_name: &'static str, make_future: F) @@ -305,6 +366,62 @@ fn gateway_imports_admin_system_config_locally_and_persists_data() { ); } +#[test] +fn gateway_rejects_ldap_filter_and_attribute_injection_before_system_import() { + run_admin_system_import_test( + "gateway_rejects_ldap_filter_and_attribute_injection_before_system_import", + gateway_rejects_ldap_filter_and_attribute_injection_before_system_import_impl, + ); +} + +async fn gateway_rejects_ldap_filter_and_attribute_injection_before_system_import_impl() { + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(build_empty_admin_system_data_state()), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for (field, value, expected_detail) in [ + ( + "user_search_filter", + "(uid={username})(objectClass=*)", + "LDAP 搜索过滤器格式无效", + ), + ( + "username_attr", + "uid)(|(objectClass=*)", + "LDAP 用户名、邮箱或显示名称属性格式无效", + ), + ] { + let mut import_payload = sample_system_import_payload(); + import_payload["ldap_config"][field] = json!(value); + let response = client + .post(format!("{gateway_url}/api/admin/system/config/import")) + .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(&import_payload) + .send() + .await + .expect("invalid LDAP import should complete locally"); + + let status = response.status(); + let payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::BAD_REQUEST, "payload={payload}"); + assert!( + payload["detail"] + .as_str() + .is_some_and(|detail| detail.contains(expected_detail)), + "payload={payload}" + ); + } + + gateway_handle.abort(); +} + async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); @@ -342,11 +459,10 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state( - AppState::new() - .expect("gateway should build") - .with_data_state_for_tests(data_state), - ); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state); + let gateway = build_router_with_state(state.clone()); let (gateway_url, gateway_handle) = start_server(gateway).await; let client = reqwest::Client::new(); @@ -381,7 +497,7 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { assert_eq!(payload["stats"]["models"]["created"], json!(1)); assert_eq!(payload["stats"]["ldap"]["created"], json!(1)); assert_eq!(payload["stats"]["oauth"]["created"], json!(1)); - assert_eq!(payload["stats"]["system_configs"]["created"], json!(3)); + assert_eq!(payload["stats"]["system_configs"]["created"], json!(5)); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); let global_models = global_model_repository @@ -422,16 +538,20 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .expect("keys should load"); assert_eq!(keys.len(), 1); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"), + decrypt_test_provider_catalog_credential(&keys[0], ProviderCatalogCredentialField::ApiKey,), "sk-import-123" ); + assert!(open_provider_catalog_credential( + &state, + &keys[0].provider_id, + "different-destination-key-id", + ProviderCatalogCredentialField::ApiKey, + keys[0] + .encrypted_api_key + .as_deref() + .expect("api key should be present"), + ) + .is_err()); assert_eq!(keys[0].api_formats, Some(json!(["openai:chat"]))); assert_eq!( keys[0].auth_type_by_format, @@ -465,14 +585,10 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .expect("ldap config should exist"); assert_eq!(ldap_config.server_url, "ldaps://ldap.example.com"); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - ldap_config - .bind_password_encrypted - .as_deref() - .expect("bind password should exist"), - ) - .expect("ldap password should decrypt"), + decrypt_or_migrate_ldap_bind_password(&state, &ldap_config) + .await + .expect("ldap password should decrypt") + .expect("ldap password should exist"), "bind-secret" ); @@ -483,14 +599,10 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .expect("oauth config should exist"); assert_eq!(oauth_provider.client_id, "linuxdo-client"); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - oauth_provider - .client_secret_encrypted - .as_deref() - .expect("oauth secret should exist"), - ) - .expect("oauth secret should decrypt"), + decrypt_or_migrate_identity_oauth_provider_client_secret(&state, &oauth_provider) + .await + .expect("oauth secret should decrypt") + .expect("oauth secret should exist"), "linuxdo-secret" ); @@ -514,20 +626,22 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .and_then(|items| items.first()) .expect("provider export should exist"); assert_eq!( - exported_provider["config"]["provider_ops"]["connector"]["credentials"]["api_key"], - "ops-secret" + exported_provider["config"]["provider_ops"]["connector"]["credentials"], + "***" ); let exported_ldap = export_payload["ldap_config"] .as_object() .expect("ldap export should exist"); - assert_eq!(exported_ldap["bind_password"], "bind-secret"); + assert!(exported_ldap.get("bind_password").is_none()); + assert_eq!(exported_ldap["is_enabled"], false); let exported_oauth = export_payload["oauth_providers"] .as_array() .and_then(|items| items.first()) .expect("oauth export should exist"); - assert_eq!(exported_oauth["client_secret"], "linuxdo-secret"); + assert!(exported_oauth.get("client_secret").is_none()); + assert_eq!(exported_oauth["is_enabled"], false); let exported_system_configs = export_payload["system_configs"] .as_array() @@ -538,18 +652,21 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data_impl() { .expect("site_name should exist"); let exported_smtp_password = exported_system_configs .iter() - .find(|entry| entry["key"] == "smtp_password") - .expect("smtp_password should exist"); + .find(|entry| entry["key"] == "smtp_password"); let exported_external_models_proxy = exported_system_configs .iter() .find(|entry| entry["key"] == "external_models_proxy_node_id") .expect("external models proxy should exist"); assert_eq!(exported_site_name["value"], "Imported Aether"); - assert_eq!(exported_smtp_password["value"], "smtp-secret"); + assert!(exported_smtp_password.is_none()); assert_eq!( exported_external_models_proxy["value"], serde_json::Value::Null ); + let serialized = export_payload.to_string(); + for secret in ["ops-secret", "bind-secret", "linuxdo-secret", "smtp-secret"] { + assert!(!serialized.contains(secret), "leaked secret: {secret}"); + } gateway_handle.abort(); upstream_handle.abort(); @@ -849,14 +966,7 @@ async fn assert_legacy_admin_system_config_import_model_test_succeeds( .expect("keys should load"); assert_eq!(keys.len(), 1); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("api key should decrypt"), + decrypt_test_provider_catalog_credential(&keys[0], ProviderCatalogCredentialField::ApiKey,), expected_api_key ); @@ -952,6 +1062,1243 @@ async fn gateway_rejects_unknown_admin_system_config_import_versions_impl() { gateway_handle.abort(); } +#[test] +fn gateway_rejects_unsafe_oauth_targets_before_system_import_mutates_data() { + run_admin_system_import_test( + "gateway_rejects_unsafe_oauth_targets_before_system_import_mutates_data", + gateway_rejects_unsafe_oauth_targets_before_system_import_mutates_data_impl, + ); +} + +async fn gateway_rejects_unsafe_oauth_targets_before_system_import_mutates_data_impl() { + let oauth_repository = Arc::new(InMemoryOAuthProviderRepository::default()); + let data = build_empty_admin_system_data_state() + .attach_oauth_provider_repository_for_tests(Arc::clone(&oauth_repository)); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + for (field, value) in [ + ( + "frontend_callback_url", + "http://attacker.example/auth/callback", + ), + ("token_url_override", "https://attacker.example/oauth/token"), + ] { + let mut payload = sample_system_import_payload(); + payload["oauth_providers"][0][field] = json!(value); + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/config/import")) + .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(&payload) + .send() + .await + .expect("unsafe OAuth import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST, "field={field}"); + assert!(oauth_repository + .list_oauth_provider_configs() + .await + .expect("OAuth provider list should load") + .is_empty()); + } + + gateway_handle.abort(); +} + +#[test] +fn gateway_prevalidates_aggregate_user_data_before_config_mutation() { + run_admin_system_import_test( + "gateway_prevalidates_aggregate_user_data_before_config_mutation", + gateway_prevalidates_aggregate_user_data_before_config_mutation_impl, + ); +} + +async fn gateway_prevalidates_aggregate_user_data_before_config_mutation_impl() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + Vec::new(), + Vec::new(), + Vec::new(), + )); + let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::< + StoredPublicGlobalModel, + >::new())); + let data = + GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository)) + .with_user_reader(user_repository) + .attach_provider_catalog_repository_for_tests(provider_catalog_repository) + .with_global_model_repository_for_tests(global_model_repository) + .with_system_config_values_for_tests(Vec::<(String, Value)>::new()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/data/import")) + .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!({ + "version": "1.0", + "merge_mode": "overwrite", + "config_data": { + "version": "2.3", + "global_models": [], + "providers": [], + "system_configs": [{ + "key": "site_name", + "value": "must-not-be-written" + }] + }, + "user_data": { + "version": "9.9", + "users": [], + "standalone_keys": [] + } + })) + .send() + .await + .expect("aggregate import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!( + state + .read_system_config_json_value_strong("site_name") + .await + .expect("system config lookup should succeed"), + None + ); + + gateway_handle.abort(); +} + +#[test] +fn gateway_prevalidates_provider_key_duplicates_before_config_mutation() { + run_admin_system_import_test( + "gateway_prevalidates_provider_key_duplicates_before_config_mutation", + gateway_prevalidates_provider_key_duplicates_before_config_mutation_impl, + ); +} + +async fn gateway_prevalidates_provider_key_duplicates_before_config_mutation_impl() { + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + Vec::new(), + Vec::new(), + Vec::new(), + )); + let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::< + StoredPublicGlobalModel, + >::new())); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(build_admin_system_data_state_with_repositories( + Arc::clone(&provider_catalog_repository), + Arc::clone(&global_model_repository), + )), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/config/import")) + .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!({ + "version": "2.3", + "merge_mode": "overwrite", + "global_models": [{ + "name": "must-not-exist", + "display_name": "Must Not Exist", + "is_active": true + }], + "providers": [{ + "name": "must-not-exist-provider", + "provider_type": "custom", + "is_active": true, + "endpoints": [{ + "api_format": "openai:chat", + "base_url": "https://api.example.com", + "is_active": true + }], + "api_keys": [{ + "name": "first", + "auth_type": "api_key", + "api_key": "duplicate-secret", + "api_formats": ["openai:chat"], + "is_active": true + }, { + "name": "second", + "auth_type": "bearer", + "api_key": "duplicate-secret", + "api_formats": ["openai:chat"], + "is_active": true + }], + "models": [] + }] + })) + .send() + .await + .expect("config import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert!(provider_catalog_repository + .list_providers(false) + .await + .expect("providers should load") + .is_empty()); + assert!(global_model_repository + .list_admin_global_models(&AdminGlobalModelListQuery { + offset: 0, + limit: 10_000, + is_active: None, + search: None, + }) + .await + .expect("global models should load") + .items + .is_empty()); + + gateway_handle.abort(); +} + +#[test] +fn gateway_prevalidates_nested_user_key_before_user_mutation() { + run_admin_system_import_test( + "gateway_prevalidates_nested_user_key_before_user_mutation", + gateway_prevalidates_nested_user_key_before_user_mutation_impl, + ); +} + +async fn gateway_prevalidates_nested_user_key_before_user_mutation_impl() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "overwrite", + "users": [{ + "email": "must-not-exist@example.com", + "username": "must-not-exist", + "role": "user", + "is_active": true, + "api_keys": [{ + "key": "sk-invalid-concurrency", + "concurrent_limit": -1 + }] + }], + "standalone_keys": [] + })) + .send() + .await + .expect("user import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert!(state + .find_user_auth_by_identifier("must-not-exist@example.com") + .await + .expect("user lookup should succeed") + .is_none()); + + gateway_handle.abort(); +} + +#[test] +fn gateway_rejects_credentials_in_v16_user_import_before_mutation() { + run_admin_system_import_test( + "gateway_rejects_credentials_in_v16_user_import_before_mutation", + gateway_rejects_credentials_in_v16_user_import_before_mutation_impl, + ); +} + +async fn gateway_rejects_credentials_in_v16_user_import_before_mutation_impl() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let cases = [ + ("password-null", json!({ "password_hash": null })), + ("plaintext-key", json!({ "key": "sk-must-not-import" })), + ("key-hash", json!({ "key_hash": "legacy-hash" })), + ("encrypted-null", json!({ "key_encrypted": null })), + ]; + for (case_name, forbidden_fields) in cases { + let email = format!("{case_name}@example.com"); + let mut api_key = serde_json::Map::from_iter([ + ( + "api_key_id".to_string(), + json!(format!("source-{case_name}")), + ), + ("credential_state".to_string(), json!("not_exported")), + ]); + if case_name != "password-null" { + api_key.extend( + forbidden_fields + .as_object() + .expect("forbidden API key fields should be an object") + .clone(), + ); + } + let mut user = serde_json::Map::from_iter([ + ("id".to_string(), json!(format!("source-user-{case_name}"))), + ("email".to_string(), json!(email.clone())), + ("username".to_string(), json!(case_name)), + ("role".to_string(), json!("user")), + ("is_active".to_string(), json!(true)), + ( + "api_keys".to_string(), + Value::Array(vec![Value::Object(api_key)]), + ), + ]); + if case_name == "password-null" { + user.extend( + forbidden_fields + .as_object() + .expect("forbidden user fields should be an object") + .clone(), + ); + } + + let response = client + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.6", + "merge_mode": "overwrite", + "users": [Value::Object(user)], + "standalone_keys": [] + })) + .send() + .await + .expect("user import should complete locally"); + + assert_eq!( + response.status(), + StatusCode::BAD_REQUEST, + "case={case_name}" + ); + assert!(state + .find_user_auth_by_identifier(&email) + .await + .expect("user lookup should succeed") + .is_none()); + } + + gateway_handle.abort(); +} + +#[test] +fn gateway_prevalidates_usage_integer_storage_before_user_mutation() { + run_admin_system_import_test( + "gateway_prevalidates_usage_integer_storage_before_user_mutation", + gateway_prevalidates_usage_integer_storage_before_user_mutation_impl, + ); +} + +async fn gateway_prevalidates_usage_integer_storage_before_user_mutation_impl() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "overwrite", + "users": [{ + "id": "source-overflow-user", + "email": "must-not-exist@example.com", + "username": "must-not-exist", + "role": "user", + "is_active": true, + "api_keys": [] + }], + "standalone_keys": [], + "usage_aggregates": { + "stats_daily": [{ + "date_unix_secs": 86400, + "total_requests": 1, + "success_requests": 1, + "error_requests": 0, + "input_tokens": 9223372036854775808_u64, + "output_tokens": 0, + "cache_creation_tokens": 0, + "cache_read_tokens": 0, + "total_cost": 0.0, + "actual_total_cost": 0.0, + "is_complete": true, + "aggregated_at_unix_secs": 86400 + }], + "stats_user_daily": [], + "stats_daily_api_key": [] + } + })) + .send() + .await + .expect("user import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert!(state + .find_user_auth_by_identifier("must-not-exist@example.com") + .await + .expect("user lookup should succeed") + .is_none()); + + gateway_handle.abort(); +} + +#[test] +fn gateway_prevalidates_duplicate_usage_dimensions_before_user_mutation() { + run_admin_system_import_test( + "gateway_prevalidates_duplicate_usage_dimensions_before_user_mutation", + gateway_prevalidates_duplicate_usage_dimensions_before_user_mutation_impl, + ); +} + +async fn gateway_prevalidates_duplicate_usage_dimensions_before_user_mutation_impl() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let daily = json!({ + "date_unix_secs": 86400, + "total_requests": 1, + "success_requests": 1, + "error_requests": 0, + "input_tokens": 1, + "output_tokens": 0, + "cache_creation_tokens": 0, + "cache_read_tokens": 0, + "total_cost": 0.0, + "actual_total_cost": 0.0, + "is_complete": true, + "aggregated_at_unix_secs": 86400 + }); + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "error", + "users": [{ + "id": "source-duplicate-user", + "email": "must-not-exist@example.com", + "username": "must-not-exist", + "role": "user", + "is_active": true, + "api_keys": [] + }], + "standalone_keys": [], + "usage_aggregates": { + "stats_daily": [daily.clone(), daily], + "stats_user_daily": [], + "stats_daily_api_key": [] + } + })) + .send() + .await + .expect("user import should complete locally"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert!(state + .find_user_auth_by_identifier("must-not-exist@example.com") + .await + .expect("user lookup should succeed") + .is_none()); + + gateway_handle.abort(); +} + +#[test] +fn gateway_user_import_allows_chained_username_release() { + run_admin_system_import_test( + "gateway_user_import_allows_chained_username_release", + gateway_user_import_allows_chained_username_release_impl, + ); +} + +async fn gateway_user_import_allows_chained_username_release_impl() { + let first = StoredUserAuthRecord::new( + "user-a".to_string(), + Some("a@example.com".to_string()), + true, + "a".to_string(), + Some("hash-a".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(chrono::Utc::now()), + None, + ) + .expect("first user should build"); + let second = StoredUserAuthRecord::new( + "user-b".to_string(), + Some("b@example.com".to_string()), + true, + "b".to_string(), + Some("hash-b".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(chrono::Utc::now()), + None, + ) + .expect("second user should build"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123"), first, second]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "overwrite", + "users": [{ + "email": "b@example.com", + "username": "c", + "role": "user", + "is_active": true, + "api_keys": [] + }, { + "email": "a@example.com", + "username": "b", + "role": "user", + "is_active": true, + "api_keys": [] + }], + "standalone_keys": [] + })) + .send() + .await + .expect("user import should complete locally"); + + let status = response.status(); + let payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + assert_eq!( + state + .find_user_auth_by_identifier("b") + .await + .expect("user lookup should succeed") + .expect("released username should be reassigned") + .id, + "user-a" + ); + assert_eq!( + state + .find_user_auth_by_identifier("c") + .await + .expect("user lookup should succeed") + .expect("second user should be renamed") + .id, + "user-b" + ); + + gateway_handle.abort(); +} + +#[test] +fn gateway_user_import_preserves_omitted_email_in_simulated_state() { + run_admin_system_import_test( + "gateway_user_import_preserves_omitted_email_in_simulated_state", + gateway_user_import_preserves_omitted_email_in_simulated_state_impl, + ); +} + +async fn gateway_user_import_preserves_omitted_email_in_simulated_state_impl() { + let existing = StoredUserAuthRecord::new( + "user-existing".to_string(), + Some("existing@example.com".to_string()), + true, + "existing".to_string(), + Some("existing-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(chrono::Utc::now()), + None, + ) + .expect("existing user should build"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123"), existing]) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "overwrite", + "users": [{ + "username": "existing", + "role": "user", + "is_active": true, + "api_keys": [] + }, { + "email": "existing@example.com", + "username": "renamed", + "role": "user", + "is_active": true, + "api_keys": [] + }], + "standalone_keys": [] + })) + .send() + .await + .expect("user import should complete locally"); + + let status = response.status(); + let payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + let updated = state + .find_user_auth_by_identifier("existing@example.com") + .await + .expect("user lookup should succeed") + .expect("existing email should remain attached"); + assert_eq!(updated.id, "user-existing"); + assert_eq!(updated.username, "renamed"); + + gateway_handle.abort(); +} + +#[test] +fn gateway_user_import_password_hash_overwrite_revokes_existing_sessions() { + run_admin_system_import_test( + "gateway_user_import_password_hash_overwrite_revokes_existing_sessions", + gateway_user_import_password_hash_overwrite_revokes_existing_sessions_impl, + ); +} + +#[test] +fn authenticated_recovery_restores_password_and_api_key_login_material() { + run_admin_system_import_test( + "authenticated_recovery_restores_password_and_api_key_login_material", + authenticated_recovery_restores_password_and_api_key_login_material_impl, + ); +} + +async fn authenticated_recovery_restores_password_and_api_key_login_material_impl() { + let password = "recovered-password-123"; + let password_hash = bcrypt::hash(password, 4).expect("password should hash"); + let plaintext_key = "sk-recovered-user-key"; + let key_hash = hash_api_key(plaintext_key); + let payload = json!({ + "version": "1.5", + "exported_at": "2026-08-22T12:00:00Z", + "merge_mode": "overwrite", + "user_groups": [], + "users": [{ + "id": "source-recovered-user", + "email": "recovered@example.com", + "email_verified": true, + "username": "recovered", + "password_hash": password_hash, + "role": "user", + "is_active": true, + "api_keys": [{ + "api_key_id": "source-recovered-key", + "key": plaintext_key, + "key_hash": key_hash, + "name": "Recovered Key", + "is_active": true + }] + }], + "standalone_keys": [], + "usage_aggregates": {} + }); + let object_key = "prod/aether-users-backup-20260822-120000.json.zst.aes256gcm"; + let compressed = zstd::stream::encode_all( + serde_json::to_vec(&payload) + .expect("payload should serialize") + .as_slice(), + 0, + ) + .expect("payload should compress"); + let (envelope, _) = encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed) + .expect("backup should encrypt"); + let restored = restore_backup_json( + object_key, + &envelope, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY) + .expect("restore key should build")], + BackupRestoreLimits::default(), + ) + .expect("backup should authenticate"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository)) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123")]) + .with_auth_wallets_for_tests(Vec::::new()); + let result = apply_restored_backup( + &state, + restored, + BackupRestoreScope::Users, + Some("admin-user-123"), + ) + .await + .expect("restore should execute") + .expect("restore payload should be valid"); + assert_eq!(result["stats"]["users"]["created"], json!(1)); + assert_eq!(result["stats"]["api_keys"]["created"], json!(1)); + + let user = state + .find_user_auth_by_identifier("recovered@example.com") + .await + .expect("user lookup should succeed") + .expect("recovered user should exist"); + assert!(bcrypt::verify( + password, + user.password_hash + .as_deref() + .expect("password hash should be restored") + ) + .expect("bcrypt hash should verify")); + let keys = state + .list_auth_api_key_export_records_by_user_ids(std::slice::from_ref(&user.id)) + .await + .expect("keys should load"); + assert_eq!(keys.len(), 1); + assert_eq!(keys[0].key_hash, key_hash); + assert!(keys[0].is_active); + assert!(keys[0] + .key_encrypted + .as_deref() + .is_some_and(|value| value.starts_with("aether-auth-api-key-secret-v2:"))); + let decrypted = open_auth_api_key_secret(&state, &keys[0]) + .ok() + .map(|projection| projection.plaintext); + assert_eq!(decrypted.as_deref(), Some(plaintext_key)); + let mut copied_record = keys[0].clone(); + copied_record.api_key_id = "different-destination-record".to_string(); + assert!(open_auth_api_key_secret(&state, &copied_record).is_err()); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time should be valid") + .as_secs(); + let authenticated = state + .data + .read_auth_api_key_snapshot_by_key_hash_strong(&hash_api_key(plaintext_key), now) + .await + .expect("API Key lookup should succeed") + .expect("restored API Key should authenticate"); + assert!(authenticated.api_key_is_active); + assert_eq!(authenticated.user_id, user.id); +} + +#[test] +fn authenticated_recovery_reencrypts_existing_api_key() { + run_admin_system_import_test( + "authenticated_recovery_reencrypts_existing_api_key", + authenticated_recovery_reencrypts_existing_api_key_impl, + ); +} + +async fn authenticated_recovery_reencrypts_existing_api_key_impl() { + let plaintext_key = "sk-recovered-existing-key"; + let key_hash = hash_api_key(plaintext_key); + let password_hash = + bcrypt::hash("recovered-existing-password", 4).expect("password should hash"); + let existing_user = StoredUserAuthRecord::new( + "user-recovered-existing".to_string(), + Some("recovered-existing@example.com".to_string()), + true, + "recovered-existing".to_string(), + Some(bcrypt::hash("old-password", 4).expect("old password should hash")), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(chrono::Utc::now()), + None, + ) + .expect("existing user should build"); + let existing_snapshot = StoredAuthApiKeySnapshot::new( + existing_user.id.clone(), + existing_user.username.clone(), + existing_user.email.clone(), + existing_user.role.clone(), + existing_user.auth_source.clone(), + true, + false, + None, + None, + None, + "key-recovered-existing".to_string(), + Some("Old Key".to_string()), + false, + false, + false, + None, + None, + None, + None, + None, + None, + ) + .expect("existing key snapshot should build"); + let auth_repository = Arc::new( + InMemoryAuthApiKeySnapshotRepository::seed([(Some(key_hash.clone()), existing_snapshot)]) + .with_export_records([StoredAuthApiKeyExportRecord::new( + existing_user.id.clone(), + "key-recovered-existing".to_string(), + key_hash.clone(), + Some("old-unusable-ciphertext".to_string()), + Some("Old Key".to_string()), + None, + None, + None, + None, + None, + None, + false, + None, + false, + 0, + 0, + 0.0, + false, + ) + .expect("existing key export should build")]), + ); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository)) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([ + sample_import_admin_user("admin-user-123"), + existing_user.clone(), + ]) + .with_auth_wallets_for_tests(Vec::::new()); + let payload = json!({ + "version": "1.5", + "exported_at": "2026-08-22T12:00:00Z", + "merge_mode": "overwrite", + "user_groups": [], + "users": [{ + "id": "source-recovered-existing", + "email": existing_user.email, + "email_verified": true, + "username": existing_user.username, + "password_hash": password_hash, + "role": "user", + "is_active": true, + "api_keys": [{ + "api_key_id": "source-recovered-existing-key", + "key": plaintext_key, + "key_hash": key_hash, + "name": "Recovered Existing Key", + "is_active": true + }] + }], + "standalone_keys": [], + "usage_aggregates": {} + }); + let object_key = "prod/aether-users-backup-20260822-120001.json.zst.aes256gcm"; + let compressed = zstd::stream::encode_all( + serde_json::to_vec(&payload) + .expect("payload should serialize") + .as_slice(), + 0, + ) + .expect("payload should compress"); + let (envelope, _) = encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed) + .expect("backup should encrypt"); + let restored = restore_backup_json( + object_key, + &envelope, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY) + .expect("restore key should build")], + BackupRestoreLimits::default(), + ) + .expect("backup should authenticate"); + + let result = apply_restored_backup( + &state, + restored, + BackupRestoreScope::Users, + Some("admin-user-123"), + ) + .await + .expect("restore should execute") + .expect("restore payload should be valid"); + assert_eq!(result["stats"]["users"]["updated"], json!(1)); + assert_eq!(result["stats"]["api_keys"]["updated"], json!(1)); + + let restored_key = state + .list_auth_api_key_export_records_by_ids(&["key-recovered-existing".to_string()]) + .await + .expect("existing key should reload") + .into_iter() + .next() + .expect("existing key should remain"); + assert!(restored_key.is_active); + assert_ne!( + restored_key.key_encrypted.as_deref(), + Some("old-unusable-ciphertext") + ); + let decrypted = open_auth_api_key_secret(&state, &restored_key) + .ok() + .map(|projection| projection.plaintext); + assert_eq!(decrypted.as_deref(), Some(plaintext_key)); +} + +#[test] +fn authenticated_recovery_rejects_proxy_nodes_before_config_mutation() { + run_admin_system_import_test( + "authenticated_recovery_rejects_proxy_nodes_before_config_mutation", + authenticated_recovery_rejects_proxy_nodes_before_config_mutation_impl, + ); +} + +async fn authenticated_recovery_rejects_proxy_nodes_before_config_mutation_impl() { + let payload = json!({ + "version": "2.3", + "exported_at": "2026-08-22T12:00:00Z", + "merge_mode": "overwrite", + "global_models": [], + "providers": [], + "proxy_nodes": [{ + "id": "deployment-local-node", + "name": "Local Node", + "ip": "127.0.0.1", + "port": 8080, + "proxy_username": "proxy-user", + "proxy_password": "proxy-password" + }], + "oauth_providers": [], + "system_configs": [{ + "key": "site_name", + "value": "must-not-be-written" + }] + }); + let object_key = "prod/aether-config-backup-20260822-120000.json.zst.aes256gcm"; + let compressed = zstd::stream::encode_all( + serde_json::to_vec(&payload) + .expect("payload should serialize") + .as_slice(), + 0, + ) + .expect("payload should compress"); + let (envelope, _) = encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed) + .expect("backup should encrypt"); + let restored = restore_backup_json( + object_key, + &envelope, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY) + .expect("restore key should build")], + BackupRestoreLimits::default(), + ) + .expect("backup should authenticate"); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(build_empty_admin_system_data_state()); + let error = apply_restored_backup(&state, restored, BackupRestoreScope::Config, None) + .await + .expect("restore should execute") + .expect_err("proxy nodes must fail closed"); + assert_eq!(error.0, StatusCode::BAD_REQUEST); + assert!(error.1["detail"] + .as_str() + .is_some_and(|detail| detail.contains("不支持安全恢复代理节点"))); + assert!(state + .read_system_config_json_value_strong("site_name") + .await + .expect("config lookup should succeed") + .is_none()); +} + +#[test] +fn authenticated_recovery_rejects_ldap_filter_injection_before_config_mutation() { + run_admin_system_import_test( + "authenticated_recovery_rejects_ldap_filter_injection_before_config_mutation", + authenticated_recovery_rejects_ldap_filter_injection_before_config_mutation_impl, + ); +} + +async fn authenticated_recovery_rejects_ldap_filter_injection_before_config_mutation_impl() { + let payload = json!({ + "version": "2.3", + "exported_at": "2026-08-31T12:00:00Z", + "merge_mode": "overwrite", + "global_models": [], + "providers": [], + "proxy_nodes": [], + "ldap_config": { + "server_url": "ldaps://ldap.example.com", + "bind_dn": "cn=admin,dc=example,dc=com", + "bind_password": "bind-secret", + "base_dn": "dc=example,dc=com", + "user_search_filter": "(uid={username})(objectClass=*)", + "username_attr": "uid", + "email_attr": "mail", + "display_name_attr": "displayName", + "is_enabled": true, + "is_exclusive": false, + "use_starttls": false, + "connect_timeout": 10 + }, + "oauth_providers": [], + "system_configs": [{ + "key": "site_name", + "value": "must-not-be-written" + }] + }); + let object_key = "prod/aether-config-backup-20260831-120000.json.zst.aes256gcm"; + let compressed = zstd::stream::encode_all( + serde_json::to_vec(&payload) + .expect("payload should serialize") + .as_slice(), + 0, + ) + .expect("payload should compress"); + let (envelope, _) = encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed) + .expect("backup should encrypt"); + let restored = restore_backup_json( + object_key, + &envelope, + &[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY) + .expect("restore key should build")], + BackupRestoreLimits::default(), + ) + .expect("backup should authenticate"); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(build_empty_admin_system_data_state()); + + let error = apply_restored_backup(&state, restored, BackupRestoreScope::Config, None) + .await + .expect("restore should execute") + .expect_err("LDAP filter injection must fail closed"); + assert_eq!(error.0, StatusCode::BAD_REQUEST); + assert!(error.1["detail"] + .as_str() + .is_some_and(|detail| detail.contains("LDAP 搜索过滤器格式无效"))); + assert!(state + .read_system_config_json_value_strong("site_name") + .await + .expect("config lookup should succeed") + .is_none()); +} + +async fn gateway_user_import_password_hash_overwrite_revokes_existing_sessions_impl() { + let existing_user = StoredUserAuthRecord::new( + "user-existing".to_string(), + Some("existing@example.com".to_string()), + true, + "existing".to_string(), + Some("old-password-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(chrono::Utc::now()), + None, + ) + .expect("existing user should build"); + let old_session = sample_import_user_session("user-existing", "old-session"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let user_repository = + Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_user_reader(user_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_auth_users_for_tests([sample_import_admin_user("admin-user-123"), existing_user]) + .with_auth_session_for_tests(old_session) + .with_auth_wallets_for_tests(Vec::::new()); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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!({ + "version": "1.5", + "merge_mode": "overwrite", + "users": [{ + "email": "existing@example.com", + "username": "existing", + "password_hash": "new-password-hash", + "role": "user", + "is_active": true, + "api_keys": [] + }], + "standalone_keys": [] + })) + .send() + .await + .expect("user import should complete locally"); + + let status = response.status(); + let payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + let updated = state + .find_user_auth_by_identifier("existing@example.com") + .await + .expect("user lookup should succeed") + .expect("user should exist"); + assert!(updated + .password_hash + .as_deref() + .is_some_and(|hash| hash.starts_with("$aether-import-revoked$"))); + assert_ne!(updated.password_hash.as_deref(), Some("new-password-hash")); + let session = state + .find_user_session("user-existing", "old-session") + .await + .expect("session lookup should succeed") + .expect("session should remain auditable"); + assert!(session.is_revoked()); + assert_eq!( + session.revoke_reason.as_deref(), + Some("admin_password_reset") + ); + + gateway_handle.abort(); +} + #[test] fn gateway_imports_admin_system_users_locally_and_persists_data() { run_admin_system_import_test( @@ -1105,7 +2452,11 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data_impl() { .expect("user lookup should succeed") .expect("imported user should exist"); assert_eq!(imported_user.username, "alice"); - assert_eq!( + assert!(imported_user + .password_hash + .as_deref() + .is_some_and(|hash| hash.starts_with("$aether-import-revoked$"))); + assert_ne!( imported_user.password_hash.as_deref(), Some("argon2:imported-user-hash") ); @@ -1172,17 +2523,12 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data_impl() { user_api_keys[0].allowed_api_formats, Some(vec!["openai:chat".to_string()]) ); - assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - user_api_keys[0] - .key_encrypted - .as_deref() - .expect("encrypted user api key should exist"), - ) - .expect("user api key should decrypt"), - "sk-user-import-1" - ); + assert!(!user_api_keys[0].is_active); + assert!(user_api_keys[0].key_encrypted.is_none()); + assert!(user_api_keys[0] + .key_hash + .starts_with("$aether-import-revoked$")); + assert_ne!(user_api_keys[0].key_hash, hash_api_key("sk-user-import-1")); let standalone_keys = state .list_auth_api_key_export_standalone_records() @@ -1196,16 +2542,14 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data_impl() { assert_eq!(standalone_keys[0].total_requests, 3); assert_eq!(standalone_keys[0].total_tokens, 789); assert_eq!(standalone_keys[0].total_cost_usd, 0.75); - assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - standalone_keys[0] - .key_encrypted - .as_deref() - .expect("encrypted standalone api key should exist"), - ) - .expect("standalone api key should decrypt"), - "sk-standalone-import-1" + assert!(!standalone_keys[0].is_active); + assert!(standalone_keys[0].key_encrypted.is_none()); + assert!(standalone_keys[0] + .key_hash + .starts_with("$aether-import-revoked$")); + assert_ne!( + standalone_keys[0].key_hash, + hash_api_key("sk-standalone-import-1") ); let standalone_wallet = state @@ -1235,14 +2579,14 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data_impl() { } #[test] -fn gateway_overwrites_existing_admin_system_user_key_usage_totals() { +fn gateway_legacy_user_import_does_not_mutate_live_keys_by_exported_hash() { run_admin_system_import_test( - "gateway_overwrites_existing_admin_system_user_key_usage_totals", - gateway_overwrites_existing_admin_system_user_key_usage_totals_impl, + "gateway_legacy_user_import_does_not_mutate_live_keys_by_exported_hash", + gateway_legacy_user_import_does_not_mutate_live_keys_by_exported_hash_impl, ); } -async fn gateway_overwrites_existing_admin_system_user_key_usage_totals_impl() { +async fn gateway_legacy_user_import_does_not_mutate_live_keys_by_exported_hash_impl() { let user_key_hash = hash_api_key("sk-existing-user-key"); let standalone_key_hash = hash_api_key("sk-existing-standalone-key"); let existing_user = StoredUserAuthRecord::new( @@ -1374,42 +2718,44 @@ async fn gateway_overwrites_existing_admin_system_user_key_usage_totals_impl() { let gateway = build_router_with_state(state.clone()); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() + let import_payload = json!({ + "version": "1.4", + "merge_mode": "overwrite", + "users": [{ + "id": "source-user-existing", + "email": "existing@example.com", + "username": "existing", + "password_hash": "existing-hash", + "role": "user", + "is_active": true, + "api_keys": [{ + "api_key_id": "source-user-key", + "key_hash": user_key_hash, + "name": "Imported User Key", + "is_active": true, + "total_requests": 222, + "total_tokens": 3333, + "total_cost_usd": 4.56 + }] + }], + "standalone_keys": [{ + "api_key_id": "source-standalone-key", + "key_hash": standalone_key_hash, + "name": "Imported Standalone Key", + "is_active": true, + "total_requests": 444, + "total_tokens": 5555, + "total_cost_usd": 6.78 + }] + }); + let client = reqwest::Client::new(); + let response = client .post(format!("{gateway_url}/api/admin/system/users/import")) .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!({ - "version": "1.4", - "merge_mode": "overwrite", - "users": [{ - "id": "source-user-existing", - "email": "existing@example.com", - "username": "existing", - "password_hash": "existing-hash", - "role": "user", - "is_active": true, - "api_keys": [{ - "api_key_id": "source-user-key", - "key_hash": user_key_hash, - "name": "Imported User Key", - "is_active": true, - "total_requests": 222, - "total_tokens": 3333, - "total_cost_usd": 4.56 - }] - }], - "standalone_keys": [{ - "api_key_id": "source-standalone-key", - "key_hash": standalone_key_hash, - "name": "Imported Standalone Key", - "is_active": true, - "total_requests": 444, - "total_tokens": 5555, - "total_cost_usd": 6.78 - }] - })) + .json(&import_payload) .send() .await .expect("request should succeed"); @@ -1418,8 +2764,27 @@ async fn gateway_overwrites_existing_admin_system_user_key_usage_totals_impl() { let payload: Value = response.json().await.expect("json body should parse"); assert_eq!(status, StatusCode::OK, "payload={payload}"); assert_eq!(payload["stats"]["users"]["updated"], json!(1)); - assert_eq!(payload["stats"]["api_keys"]["updated"], json!(1)); - assert_eq!(payload["stats"]["standalone_keys"]["updated"], json!(1)); + assert_eq!(payload["stats"]["api_keys"]["created"], json!(1)); + assert_eq!(payload["stats"]["standalone_keys"]["created"], json!(1)); + + let response = client + .post(format!("{gateway_url}/api/admin/system/users/import")) + .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(&import_payload) + .send() + .await + .expect("repeated request should succeed"); + let status = response.status(); + let repeated_payload: Value = response.json().await.expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={repeated_payload}"); + assert_eq!(repeated_payload["stats"]["api_keys"]["updated"], json!(1)); + assert_eq!( + repeated_payload["stats"]["standalone_keys"]["updated"], + json!(1) + ); let updated_records = state .list_auth_api_key_export_records_by_ids(&[ @@ -1432,16 +2797,54 @@ async fn gateway_overwrites_existing_admin_system_user_key_usage_totals_impl() { .iter() .find(|record| record.api_key_id == "key-user-existing") .expect("updated user key should exist"); - assert_eq!(user_key.total_requests, 222); - assert_eq!(user_key.total_tokens, 3333); - assert_eq!(user_key.total_cost_usd, 4.56); + assert!(user_key.is_active); + assert_eq!(user_key.total_requests, 1); + assert_eq!(user_key.total_tokens, 2); + assert_eq!(user_key.total_cost_usd, 0.03); let standalone_key = updated_records .iter() .find(|record| record.api_key_id == "key-standalone-existing") .expect("updated standalone key should exist"); - assert_eq!(standalone_key.total_requests, 444); - assert_eq!(standalone_key.total_tokens, 5555); - assert_eq!(standalone_key.total_cost_usd, 6.78); + assert!(standalone_key.is_active); + assert_eq!(standalone_key.total_requests, 4); + assert_eq!(standalone_key.total_tokens, 5); + assert_eq!(standalone_key.total_cost_usd, 0.06); + + let user_keys = state + .list_auth_api_key_export_records_by_user_ids(&["user-existing".to_string()]) + .await + .expect("user key export records should load"); + assert_eq!(user_keys.len(), 2); + let imported_user_key = user_keys + .iter() + .find(|record| record.api_key_id != "key-user-existing") + .expect("disabled imported user key should exist"); + assert!(!imported_user_key.is_active); + assert!(imported_user_key.key_encrypted.is_none()); + assert!(imported_user_key + .key_hash + .starts_with("$aether-import-revoked$")); + assert_eq!(imported_user_key.total_requests, 222); + assert_eq!(imported_user_key.total_tokens, 3333); + assert_eq!(imported_user_key.total_cost_usd, 4.56); + + let standalone_keys = state + .list_auth_api_key_export_standalone_records() + .await + .expect("standalone key export records should load"); + assert_eq!(standalone_keys.len(), 2); + let imported_standalone_key = standalone_keys + .iter() + .find(|record| record.api_key_id != "key-standalone-existing") + .expect("disabled imported standalone key should exist"); + assert!(!imported_standalone_key.is_active); + assert!(imported_standalone_key.key_encrypted.is_none()); + assert!(imported_standalone_key + .key_hash + .starts_with("$aether-import-revoked$")); + assert_eq!(imported_standalone_key.total_requests, 444); + assert_eq!(imported_standalone_key.total_tokens, 5555); + assert_eq!(imported_standalone_key.total_cost_usd, 6.78); gateway_handle.abort(); } @@ -2063,24 +3466,13 @@ async fn gateway_imports_oauth_provider_key_credentials_from_admin_system_config assert_eq!(keys.len(), 1); assert_eq!(keys[0].auth_type, "oauth"); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("oauth access token should decrypt"), + decrypt_test_provider_catalog_credential(&keys[0], ProviderCatalogCredentialField::ApiKey,), "oauth-access-token-1" ); - let auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_auth_config - .as_deref() - .expect("oauth auth config should exist"), - ) - .expect("oauth auth config should decrypt"); + let auth_config = decrypt_test_provider_catalog_credential( + &keys[0], + ProviderCatalogCredentialField::AuthConfig, + ); let auth_config: Value = serde_json::from_str(&auth_config).expect("oauth auth config json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -2165,24 +3557,13 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp assert_eq!(keys.len(), 1); assert_eq!(keys[0].name, "oauth-primary"); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("oauth access token should decrypt"), + decrypt_test_provider_catalog_credential(&keys[0], ProviderCatalogCredentialField::ApiKey,), "oauth-access-token-new" ); - let auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - keys[0] - .encrypted_auth_config - .as_deref() - .expect("oauth auth config should exist"), - ) - .expect("oauth auth config should decrypt"); + let auth_config = decrypt_test_provider_catalog_credential( + &keys[0], + ProviderCatalogCredentialField::AuthConfig, + ); let auth_config: Value = serde_json::from_str(&auth_config).expect("oauth auth config json should parse"); assert_eq!(auth_config["refresh_token"], "oauth-refresh-token-new"); @@ -2192,7 +3573,6 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp &keys[0].id, Some(1_700_000_001), Some("[REFRESH_FAILED] imported token remains invalid"), - None, Some(1_700_000_001), ) .await @@ -2454,23 +3834,12 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp Some(&Value::Null) ); assert_eq!( - decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - key.encrypted_api_key - .as_deref() - .expect("api key should be present"), - ) - .expect("oauth access token should decrypt"), + decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::ApiKey,), "oauth-access-token-new" ); - let auth_config = decrypt_python_fernet_ciphertext( - DEVELOPMENT_ENCRYPTION_KEY, - key.encrypted_auth_config - .as_deref() - .expect("oauth auth config should exist"), - ) - .expect("oauth auth config should decrypt"); + let auth_config = + decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::AuthConfig); let auth_config: Value = serde_json::from_str(&auth_config).expect("oauth auth config json should parse"); assert_eq!(auth_config["provider_type"], "codex"); @@ -2673,7 +4042,7 @@ async fn gateway_preserves_manual_proxy_configs_while_skipping_proxy_nodes_durin providers[0].proxy, Some(json!({ "enabled": true, - "url": "https://proxy.example" + "url": "https://proxy.example/" })) ); diff --git a/apps/aether-gateway/src/tests/control/admin/usage.rs b/apps/aether-gateway/src/tests/control/admin/usage.rs index 7429df7f4..4d2a430ea 100644 --- a/apps/aether-gateway/src/tests/control/admin/usage.rs +++ b/apps/aether-gateway/src/tests/control/admin/usage.rs @@ -20,8 +20,8 @@ use http::{HeaderMap, HeaderValue, StatusCode}; use serde_json::json; use super::super::{ - build_router_with_state, issue_test_admin_access_token, sample_endpoint, sample_key, - sample_provider, start_server, AppState, + build_router_with_state, issue_test_admin_access_token, sample_bound_key, sample_endpoint, + sample_key, sample_provider, start_server, AppState, }; use crate::admin_api::{ maybe_build_local_admin_usage_response, AdminAppState, AdminRequestContext, @@ -986,7 +986,8 @@ async fn gateway_handles_admin_usage_active_locally_with_trusted_admin_principal DAY_1_UNIX_SECS, ), ])); - let mut provider_key = sample_key("provider-key-1", "provider-1", "openai:chat", "sk-upstream"); + let mut provider_key = + sample_bound_key("provider-key-1", "provider-1", "openai:chat", "sk-upstream"); provider_key.name = "upstream-primary".to_string(); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-1", "OpenAI", 10)], @@ -1283,7 +1284,8 @@ async fn gateway_handles_admin_usage_records_locally_with_trusted_admin_principa DAY_2_UNIX_SECS, ), ])); - let mut provider_key = sample_key("provider-key-1", "provider-1", "openai:chat", "sk-upstream"); + let mut provider_key = + sample_bound_key("provider-key-1", "provider-1", "openai:chat", "sk-upstream"); provider_key.name = "upstream-primary".to_string(); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-1", "OpenAI", 10)], @@ -3460,6 +3462,7 @@ async fn gateway_handles_admin_usage_cache_affinity_interval_timeline_with_legac allowed_models_mode: "unrestricted".to_string(), is_active: true, is_deleted: false, + security_version: 0, created_at: None, last_login_at: None, }]), diff --git a/apps/aether-gateway/src/tests/control/admin/users.rs b/apps/aether-gateway/src/tests/control/admin/users.rs index 86558d82e..8134f2106 100644 --- a/apps/aether-gateway/src/tests/control/admin/users.rs +++ b/apps/aether-gateway/src/tests/control/admin/users.rs @@ -4,6 +4,7 @@ use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY} use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, }; +use aether_data::repository::management_tokens::InMemoryManagementTokenRepository; use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::users::{ InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserExportRow, UpsertUserGroupRecord, @@ -20,7 +21,8 @@ use http::StatusCode; use serde_json::json; use super::super::{ - build_router_with_state, hash_api_key, issue_test_admin_access_token, start_server, AppState, + build_router_with_state, hash_api_key, hash_management_token, issue_test_admin_access_token, + sample_management_token, start_server, AppState, }; use crate::constants::{ GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER, @@ -216,6 +218,27 @@ fn sample_admin_api_key_snapshot(user_id: &str, api_key_id: &str) -> StoredAuthA .expect("api key snapshot should build") } +/// Build an auth API-key fixture with the provider/user-bound envelope used by +/// production writes. Read-only test repositories cannot perform the legacy +/// Fernet migration, so ordinary list/reveal fixtures must use this format. +fn sample_bound_admin_api_key_secret( + user_id: &str, + api_key_id: &str, + plaintext: &str, +) -> (String, String) { + let key_hash = hash_api_key(plaintext); + let state = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted = crate::handlers::shared::seal_auth_api_key_secret( + &state, user_id, api_key_id, &key_hash, false, plaintext, + ) + .expect("bound API-key ciphertext should build"); + (key_hash, encrypted) +} + #[tokio::test] async fn gateway_sorts_admin_users_by_created_at() { let oldest = Utc @@ -1231,6 +1254,280 @@ async fn gateway_handles_admin_users_root_locally_with_bearer_admin_session() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_rejects_users_write_management_token_for_administrator_account_mutations() { + let raw_token = "ae-users-write-cannot-escalate"; + let token_owner = sample_admin_user_with_role( + "token-owner", + "admin", + "token-owner@example.com", + "token_owner", + ); + let target_user = + sample_admin_user_with_role("target-user", "user", "target@example.com", "target_user"); + let target_admin = sample_admin_user_with_role( + "target-admin", + "admin", + "target-admin@example.com", + "target_admin", + ); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ + token_owner.clone(), + target_user.clone(), + target_admin.clone(), + ])); + let mut token = sample_management_token( + "token-users-write-cannot-escalate", + &token_owner.id, + &token_owner.username, + true, + ); + token.token.allowed_ips = None; + token.token.permissions = Some(json!(["admin:users:write"])); + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![token], + vec![( + hash_management_token(raw_token), + "token-users-write-cannot-escalate".to_string(), + )], + )); + let data = GatewayDataState::with_management_token_repository_for_tests(token_repository) + .with_user_reader(user_repository.clone()); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data) + .with_auth_users_for_tests([ + token_owner.clone(), + target_user.clone(), + target_admin.clone(), + ]); + let inspection_state = state.clone(); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let cases = [ + ( + reqwest::Method::POST, + "/api/admin/users", + json!({ + "email": "created-admin@example.com", + "username": "created_admin", + "password": "CreatedAdmin123!", + "role": "admin" + }), + "create_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-user", + json!({ "role": "admin" }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "password": "ResetAdmin123!" }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "role": "user" }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "is_active": false }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "email": "hijacked-admin@example.com" }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "group_ids": [] }), + "update_user", + ), + ( + reqwest::Method::PUT, + "/api/admin/users/target-admin", + json!({ "unlimited": true }), + "update_user", + ), + ( + reqwest::Method::DELETE, + "/api/admin/users/target-admin", + json!({}), + "delete_user", + ), + ( + reqwest::Method::POST, + "/api/admin/users/batch-action", + json!({ + "selection": { "user_ids": ["target-user"] }, + "action": "update_role", + "payload": { "role": "admin" } + }), + "batch_action_users", + ), + ( + reqwest::Method::POST, + "/api/admin/users/batch-action", + json!({ + "selection": { "user_ids": ["target-admin"] }, + "action": "update_role", + "payload": { "role": "user" } + }), + "batch_action_users", + ), + ( + reqwest::Method::POST, + "/api/admin/users/batch-action", + json!({ + "selection": { "user_ids": ["target-user", "target-admin"] }, + "action": "disable" + }), + "batch_action_users", + ), + ]; + + for (method, path, body, route_kind) in cases { + let response = client + .request(method, format!("{gateway_url}{path}")) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(raw_token) + .json(&body) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "management token permission denied"); + assert_eq!(payload["required_permission"], "admin:users:admin"); + assert_eq!(payload["route_kind"], route_kind); + } + + let ordinary_create_response = client + .post(format!("{gateway_url}/api/admin/users")) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(raw_token) + .json(&json!({ + "email": "delegated-user@example.com", + "username": "delegated_user", + "password": "DelegatedUser123!", + "role": "user" + })) + .send() + .await + .expect("ordinary user creation request should succeed"); + assert_ne!(ordinary_create_response.status(), StatusCode::FORBIDDEN); + + let ordinary_update_response = client + .put(format!("{gateway_url}/api/admin/users/target-user")) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(raw_token) + .json(&json!({ "is_active": false })) + .send() + .await + .expect("ordinary user update request should succeed"); + assert_eq!(ordinary_update_response.status(), StatusCode::OK); + + assert!(inspection_state + .find_user_auth_by_identifier("created_admin") + .await + .expect("created user lookup should succeed") + .is_none()); + assert_eq!( + inspection_state + .find_user_auth_by_id("target-user") + .await + .expect("target user lookup should succeed") + .expect("target user should still exist") + .is_active, + false + ); + let stored_target_admin = inspection_state + .find_user_auth_by_id("target-admin") + .await + .expect("target admin lookup should succeed") + .expect("target admin should still exist"); + assert_eq!(stored_target_admin.role, "admin"); + assert!(stored_target_admin.is_active); + assert_eq!( + stored_target_admin.password_hash, + target_admin.password_hash + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_requires_users_admin_to_revoke_privileged_user_sessions() { + for (role, suffix) in [("admin", "admin"), ("audit_admin", "audit-admin")] { + let raw_token = format!("ae-users-write-session-{suffix}"); + let token_owner = sample_admin_user_with_role( + &format!("token-owner-{suffix}"), + "admin", + &format!("token-owner-{suffix}@example.com"), + &format!("token_owner_{suffix}"), + ); + let target_id = format!("target-{suffix}"); + let target = sample_admin_user_with_role( + &target_id, + role, + &format!("target-{suffix}@example.com"), + &format!("target_{suffix}"), + ); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ + token_owner.clone(), + target.clone(), + ])); + let token_id = format!("token-users-write-session-{suffix}"); + let mut token = + sample_management_token(&token_id, &token_owner.id, &token_owner.username, true); + token.token.allowed_ips = None; + token.token.permissions = Some(json!(["admin:users:write"])); + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![token], + vec![(hash_management_token(&raw_token), token_id)], + )); + 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]) + .with_auth_session_for_tests(sample_admin_session(&target_id, "session-1")); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for path in [ + format!("/api/admin/users/{target_id}/sessions/session-1"), + format!("/api/admin/users/{target_id}/sessions"), + ] { + let response = client + .delete(format!("{gateway_url}{path}")) + .header(GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(&raw_token) + .send() + .await + .expect("session revocation request should complete"); + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["required_permission"], "admin:users:admin"); + } + + gateway_handle.abort(); + } +} + #[tokio::test] async fn gateway_handles_admin_user_detail_routes_locally_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -1244,16 +1541,16 @@ async fn gateway_handles_admin_user_detail_routes_locally_with_trusted_admin_pri })); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state( - AppState::new() - .expect("gateway should build") - .with_auth_users_for_tests([ - sample_admin_user("user-1"), - sample_admin_user_with_role("admin-1", "admin", "admin1@example.com", "admin_one"), - sample_admin_user_with_role("admin-2", "admin", "admin2@example.com", "admin_two"), - ]) - .with_auth_wallets_for_tests([sample_admin_wallet("user-1", "finite")]), - ); + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([ + sample_admin_user("user-1"), + sample_admin_user_with_role("admin-1", "admin", "admin1@example.com", "admin_one"), + sample_admin_user_with_role("admin-2", "admin", "admin2@example.com", "admin_two"), + ]) + .with_auth_wallets_for_tests([sample_admin_wallet("user-1", "finite")]); + let inspection_state = state.clone(); + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; let client = reqwest::Client::new(); @@ -1286,6 +1583,17 @@ async fn gateway_handles_admin_user_detail_routes_locally_with_trusted_admin_pri assert_eq!(update_payload["unlimited"], true); assert_eq!(update_payload["is_active"], false); + let updated_user = inspection_state + .find_user_auth_by_id("user-1") + .await + .expect("updated user lookup should succeed") + .expect("updated user should exist"); + assert_eq!( + updated_user.email.as_deref(), + Some("alice-updated@example.com") + ); + assert!(!updated_user.email_verified); + let hidden_update_response = client .put(format!("{gateway_url}/api/admin/users/user-1")) .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") @@ -1357,7 +1665,7 @@ async fn gateway_rejects_demoting_the_last_active_admin() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::BAD_REQUEST); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "不能降级最后一个管理员账户"); + assert_eq!(payload["detail"], "不能降级或停用最后一个管理员账户"); } gateway_handle.abort(); @@ -2392,17 +2700,16 @@ async fn gateway_lists_admin_user_api_keys_locally_with_trusted_admin_principal( }), ); - let encrypted = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-user-1") - .expect("ciphertext should build"); + let (key_hash, encrypted) = sample_bound_admin_api_key_secret("user-1", "key-1", "sk-user-1"); let auth_repository = Arc::new( InMemoryAuthApiKeySnapshotRepository::seed(vec![( - Some("hash-key-1".to_string()), + Some(key_hash.clone()), sample_admin_api_key_snapshot("user-1", "key-1"), )]) .with_export_records(vec![StoredAuthApiKeyExportRecord::new( "user-1".to_string(), "key-1".to_string(), - "hash-key-1".to_string(), + key_hash, Some(encrypted), Some("default".to_string()), Some(json!(["openai"])), @@ -2459,7 +2766,7 @@ async fn gateway_lists_admin_user_api_keys_locally_with_trusted_admin_principal( assert_eq!(payload["username"], "alice"); assert_eq!(payload["api_keys"][0]["id"], "key-1"); assert_eq!(payload["api_keys"][0]["name"], "default"); - assert_eq!(payload["api_keys"][0]["key_display"], "sk-user-1...er-1"); + 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]["total_requests"], 9); @@ -2597,17 +2904,16 @@ async fn gateway_reveals_admin_user_full_key_locally_with_trusted_admin_principa }), ); - let encrypted = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-user-1") - .expect("ciphertext should build"); + let (key_hash, encrypted) = sample_bound_admin_api_key_secret("user-1", "key-1", "sk-user-1"); let auth_repository = Arc::new( InMemoryAuthApiKeySnapshotRepository::seed(vec![( - Some("hash-key-1".to_string()), + Some(key_hash.clone()), sample_admin_api_key_snapshot("user-1", "key-1"), )]) .with_export_records(vec![StoredAuthApiKeyExportRecord::new( "user-1".to_string(), "key-1".to_string(), - "hash-key-1".to_string(), + key_hash, Some(encrypted), Some("default".to_string()), Some(json!(["openai"])), diff --git a/apps/aether-gateway/src/tests/control/admin/video_tasks.rs b/apps/aether-gateway/src/tests/control/admin/video_tasks.rs index 02e9bf6aa..5a91004b8 100644 --- a/apps/aether-gateway/src/tests/control/admin/video_tasks.rs +++ b/apps/aether-gateway/src/tests/control/admin/video_tasks.rs @@ -15,8 +15,8 @@ use http::{HeaderMap, HeaderValue, StatusCode}; use serde_json::json; use super::super::{ - build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint, - sample_key, sample_provider, start_server, AppState, + build_router_with_state, build_state_with_execution_runtime_override, sample_bound_key, + sample_endpoint, sample_provider, start_server, AppState, }; use crate::admin_api::{ maybe_build_local_admin_video_tasks_response, AdminAppState, AdminRequestContext, @@ -130,7 +130,7 @@ fn sample_admin_video_task( updated_at_unix_secs: created_at_unix_ms + 5, error_code: None, error_message: None, - video_url: Some(format!("https://example.com/{id}.mp4")), + video_url: Some(format!("https://8.8.8.8/{id}.mp4")), request_metadata: None, } } @@ -221,12 +221,13 @@ async fn gateway_handles_admin_video_tasks_list_locally_with_trusted_admin_princ assert_eq!(payload["pages"], json!(1)); assert_eq!(payload["items"].as_array().map(Vec::len), Some(1)); assert_eq!(payload["items"][0]["id"], "task-completed"); - assert_eq!(payload["items"][0]["username"], "alice"); + // Video-task persistence intentionally drops user-facing PII. The admin + // projection must therefore use the privacy-safe fallback when no separate + // user snapshot is joined. + assert_eq!(payload["items"][0]["username"], "Unknown"); assert_eq!(payload["items"][0]["provider_name"], "OpenAI"); assert_eq!(payload["items"][0]["status"], "completed"); - assert!(payload["items"][0]["prompt"] - .as_str() - .is_some_and(|value| value.ends_with("..."))); + assert!(payload["items"][0]["prompt"].is_null()); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -392,7 +393,7 @@ async fn gateway_handles_admin_video_task_detail_locally_with_trusted_admin_prin assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["id"], "task-detail"); - assert_eq!(payload["username"], "charlie"); + assert_eq!(payload["username"], "Unknown"); assert_eq!(payload["provider_name"], "OpenAI"); assert_eq!(payload["endpoint"]["id"], "endpoint-1"); assert_eq!(payload["endpoint"]["api_format"], "openai:video"); @@ -592,11 +593,36 @@ async fn gateway_cancels_admin_video_task_locally_with_trusted_admin_principal() .await .expect("upsert should succeed"); + // The persisted task intentionally omits its request/transport snapshot. + // Reconstructing an admin cancellation must therefore resolve the current + // provider catalog, including a record-bound credential that a read-only + // fixture can decrypt without a migration writer. + let provider_catalog_repository: Arc = + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-openai", "OpenAI", 10)], + vec![sample_endpoint( + "endpoint-1", + "provider-openai", + "openai:video", + "https://api.openai.example/v1", + )], + vec![sample_bound_key( + "provider-key-1", + "provider-openai", + "openai:video", + "sk-upstream-openai-video", + )], + )); + let (upstream_url, upstream_handle) = start_server(upstream).await; let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let gateway = build_router_with_state( build_state_with_execution_runtime_override(execution_runtime_url) - .with_video_task_data_repository_for_tests(Arc::clone(&repository)), + .with_video_task_repository_and_provider_transport_for_tests( + Arc::clone(&repository), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -639,10 +665,10 @@ async fn gateway_cancels_admin_video_task_locally_with_trusted_admin_principal() assert_eq!(detail.status(), StatusCode::OK); let detail_json: serde_json::Value = detail.json().await.expect("detail should parse"); assert_eq!(detail_json["status"], "cancelled"); - assert_eq!( - detail_json["request_metadata"]["rust_local_snapshot"]["OpenAi"]["status"], - "Cancelled" - ); + assert!(detail_json["original_request_body"].is_null()); + assert!(detail_json["progress_message"].is_null()); + assert!(detail_json["error_message"].is_null()); + assert!(detail_json["request_metadata"].is_null()); let seen_execution_runtime_request = seen_execution_runtime .lock() @@ -708,7 +734,7 @@ async fn local_admin_video_task_cancel_attaches_explicit_audit() { } #[tokio::test] -async fn gateway_redirects_admin_video_task_video_locally_with_trusted_admin_principal() { +async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( @@ -723,7 +749,7 @@ async fn gateway_redirects_admin_video_task_video_locally_with_trusted_admin_pri ); let repository = Arc::new(InMemoryVideoTaskRepository::default()); - repository + let stored = repository .upsert(sample_admin_video_task( "task-redirect", VideoTaskStatus::Completed, @@ -736,8 +762,9 @@ async fn gateway_redirects_admin_video_task_video_locally_with_trusted_admin_pri )) .await .expect("task should upsert"); + assert_eq!(stored.video_url, None); - let (upstream_url, upstream_handle) = start_server(upstream).await; + let (_upstream_url, upstream_handle) = start_server(upstream).await; let gateway = build_router_with_state( AppState::new() .expect("gateway state should build") @@ -761,14 +788,7 @@ async fn gateway_redirects_admin_video_task_video_locally_with_trusted_admin_pri .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT); - assert_eq!( - response - .headers() - .get(http::header::LOCATION) - .and_then(|value| value.to_str().ok()), - Some("https://example.com/task-redirect.mp4") - ); + assert_eq!(response.status(), StatusCode::NOT_FOUND); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -776,9 +796,9 @@ async fn gateway_redirects_admin_video_task_video_locally_with_trusted_admin_pri } #[tokio::test] -async fn local_admin_video_task_video_redirect_attaches_explicit_audit() { +async fn local_admin_video_task_video_is_unavailable_after_openai_url_sanitization() { let repository = Arc::new(InMemoryVideoTaskRepository::default()); - repository + let stored = repository .upsert(sample_admin_video_task( "task-video-audit", VideoTaskStatus::Completed, @@ -791,6 +811,7 @@ async fn local_admin_video_task_video_redirect_attaches_explicit_audit() { )) .await .expect("task should upsert"); + assert_eq!(stored.video_url, None); let state = AppState::new() .expect("gateway state should build") @@ -804,20 +825,12 @@ async fn local_admin_video_task_video_redirect_attaches_explicit_audit() { ) .await; - assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT); - let audit = response - .extensions() - .get::() - .cloned() - .expect("video task video should attach audit"); - assert_eq!(audit.event_name, "admin_video_task_video_viewed"); - assert_eq!(audit.action, "view_video_task_video"); - assert_eq!(audit.target_type, "video_task_video"); - assert_eq!(audit.target_id, "task-video-audit"); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + assert!(response.extensions().get::().is_none()); } #[tokio::test] -async fn gateway_proxies_admin_video_task_video_locally_with_trusted_admin_principal() { +async fn gateway_rejects_cross_origin_admin_gemini_video_without_sending_provider_key() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let seen_api_key = Arc::new(Mutex::new(None::)); @@ -887,7 +900,7 @@ async fn gateway_proxies_admin_video_task_video_locally_with_trusted_admin_princ "gemini:video", "https://generativelanguage.googleapis.com", )], - vec![sample_key( + vec![sample_bound_key( "key-gemini", "provider-gemini", "gemini:video", @@ -919,28 +932,10 @@ async fn gateway_proxies_admin_video_task_video_locally_with_trusted_admin_princ .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .headers() - .get(http::header::CONTENT_TYPE) - .and_then(|value| value.to_str().ok()), - Some("video/mp4") - ); - assert_eq!( - response - .headers() - .get(http::header::CONTENT_DISPOSITION) - .and_then(|value| value.to_str().ok()), - Some("inline; filename=\"video_task-proxy.mp4\"") - ); - assert_eq!( - response.bytes().await.expect("body should read"), - Bytes::from_static(b"proxied-video-bytes") - ); + assert_eq!(response.status(), StatusCode::BAD_GATEWAY); assert_eq!( seen_api_key.lock().expect("mutex should lock").as_deref(), - Some("gemini-upstream-secret") + None ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); diff --git a/apps/aether-gateway/src/tests/control/admin/wallets.rs b/apps/aether-gateway/src/tests/control/admin/wallets.rs index d68d9248f..e80eac5b7 100644 --- a/apps/aether-gateway/src/tests/control/admin/wallets.rs +++ b/apps/aether-gateway/src/tests/control/admin/wallets.rs @@ -198,7 +198,7 @@ fn sample_refund_record( payment_order_id: payment_order_id.map(ToOwned::to_owned), source_type: "payment_order".to_string(), source_id: payment_order_id.map(ToOwned::to_owned), - refund_mode: "original".to_string(), + refund_mode: "original_channel".to_string(), amount_usd, status: status.to_string(), reason: Some("用户申请退款".to_string()), @@ -1300,6 +1300,29 @@ async fn gateway_handles_admin_wallets_process_refund_locally_with_trusted_admin assert_eq!(payload["transaction"]["reason_code"], json!("refund_out")); assert_eq!(payload["transaction"]["amount"], json!(-4.0)); assert_eq!(payload["transaction"]["description"], json!("退款占款")); + + let ledger_response = reqwest::Client::new() + .get(format!( + "{gateway_url}/api/admin/wallets/wallet-123/transactions?limit=20&offset=0" + )) + .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") + .send() + .await + .expect("transaction list request should succeed"); + assert_eq!(ledger_response.status(), StatusCode::OK); + let ledger_payload: serde_json::Value = ledger_response + .json() + .await + .expect("transaction list response should be json"); + assert_eq!(ledger_payload["total"], json!(1)); + assert_eq!( + ledger_payload["items"][0]["reason_code"], + json!("refund_out") + ); + assert_eq!(ledger_payload["items"][0]["amount"], json!(-4.0)); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1425,7 +1448,7 @@ async fn gateway_handles_admin_wallets_complete_refund_locally_with_trusted_admi } #[tokio::test] -async fn gateway_handles_admin_wallets_fail_refund_locally_with_trusted_admin_principal() { +async fn gateway_rejects_admin_wallets_fail_processing_refund_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( @@ -1485,20 +1508,115 @@ async fn gateway_handles_admin_wallets_fail_refund_locally_with_trusted_admin_pr .await .expect("request should succeed"); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!( + payload["detail"], + json!("cannot fail refund while gateway settlement is processing") + ); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_releases_offline_processing_refund_without_gateway_evidence() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().route( + "/api/admin/wallets/wallet-123/refunds/refund-1/fail", + any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + }), + ); + + let mut wallet = sample_wallet_snapshot("wallet-123", Some("user-1"), None, "finite"); + wallet.balance = 8.5; + wallet.total_refunded = 7.0; + let mut refund = sample_refund_record( + "refund-1", + "wallet-123", + Some("user-1"), + Some("po-1"), + 4.0, + "processing", + Some("admin-user-123"), + Some("admin-user-123"), + Some(1_710_000_500), + ); + refund.refund_mode = "offline_payout".to_string(); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_auth_wallets_for_tests([wallet]) + .with_admin_wallet_payment_orders_for_tests([sample_payment_order_record( + "po-1", + "wallet-123", + Some("user-1"), + 10.0, + 5.0, + 5.0, + )]) + .with_admin_wallet_refunds_for_tests([refund]), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!( + "{gateway_url}/api/admin/wallets/wallet-123/refunds/refund-1/fail" + )) + .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") + .json(&json!({ + "reason": "线下退款未打款" + })) + .send() + .await + .expect("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["wallet"]["id"], json!("wallet-123")); assert_eq!(payload["wallet"]["balance"], json!(15.0)); assert_eq!(payload["wallet"]["recharge_balance"], json!(12.5)); assert_eq!(payload["wallet"]["total_refunded"], json!(3.0)); assert_eq!(payload["refund"]["status"], json!("failed")); - assert_eq!(payload["refund"]["failure_reason"], json!("原路退款失败")); assert_eq!( payload["transaction"]["reason_code"], json!("refund_revert") ); assert_eq!(payload["transaction"]["amount"], json!(4.0)); - assert_eq!(payload["transaction"]["description"], json!("退款失败回补")); + + let ledger_response = reqwest::Client::new() + .get(format!( + "{gateway_url}/api/admin/wallets/wallet-123/transactions?limit=20&offset=0" + )) + .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") + .send() + .await + .expect("transaction list request should succeed"); + assert_eq!(ledger_response.status(), StatusCode::OK); + let ledger_payload: serde_json::Value = ledger_response + .json() + .await + .expect("transaction list response should be json"); + assert_eq!(ledger_payload["total"], json!(1)); + assert_eq!( + ledger_payload["items"][0]["reason_code"], + json!("refund_revert") + ); + assert_eq!(ledger_payload["items"][0]["amount"], json!(4.0)); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/control/helpers.rs b/apps/aether-gateway/src/tests/control/helpers.rs index 6c91c8bcd..6c31ed981 100644 --- a/apps/aether-gateway/src/tests/control/helpers.rs +++ b/apps/aether-gateway/src/tests/control/helpers.rs @@ -13,7 +13,96 @@ use super::{ StoredProviderModelStats, StoredProviderQuotaSnapshot, StoredProxyNode, StoredPublicGlobalModel, StoredRequestCandidate, }; -use crate::AppState; +use crate::{data::GatewayDataState, AppState}; + +pub(super) const TUNNEL_CONTROL_PLANE_TEST_PSK: &str = + "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; +pub(super) const TUNNEL_CONTROL_PLANE_TEST_GENERATION: &str = "test-generation-1"; + +pub(super) fn with_tunnel_control_plane_key( + mut node: StoredProxyNode, + key: &str, +) -> StoredProxyNode { + let mut metadata = node.proxy_metadata.take().unwrap_or_else(|| json!({})); + let metadata = metadata + .as_object_mut() + .expect("sample proxy metadata should be an object"); + metadata.insert( + "tunnel_security".to_string(), + json!({ "encryption_key": key }), + ); + node.proxy_metadata = Some(serde_json::Value::Object(metadata.clone())); + node +} + +pub(super) fn authenticated_tunnel_control_plane_request( + client: &reqwest::Client, + url: String, + path: &str, + node_id: &str, + body: &serde_json::Value, +) -> reqwest::RequestBuilder { + authenticated_tunnel_control_plane_request_for_generation( + client, + url, + path, + node_id, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, + body, + ) +} + +pub(super) fn authenticated_tunnel_control_plane_request_for_generation( + client: &reqwest::Client, + url: String, + path: &str, + node_id: &str, + tunnel_generation: &str, + body: &serde_json::Value, +) -> reqwest::RequestBuilder { + let body = serde_json::to_vec(body).expect("control-plane payload should encode"); + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let signature = + aether_contracts::tunnel_security::sign_tunnel_control_plane_request_for_generation( + TUNNEL_CONTROL_PLANE_TEST_PSK, + "POST", + path, + node_id, + tunnel_generation, + timestamp, + &nonce, + &body, + ) + .expect("control-plane request should sign"); + client + .post(url) + .header("content-type", "application/json") + .header( + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, + node_id, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_GENERATION_HEADER, + tunnel_generation, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, + timestamp, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_NONCE_HEADER, + nonce, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, + signature, + ) + .body(body) +} pub(super) fn sample_currently_usable_auth_snapshot( api_key_id: &str, @@ -78,7 +167,7 @@ pub(super) fn test_auth_secret() -> String { .ok() .map(|value| value.trim().to_string()) .filter(|value| !value.is_empty()) - .unwrap_or_else(|| "aether-rust-dev-jwt-secret".to_string()) + .unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string()) } pub(super) fn build_test_auth_token( @@ -112,7 +201,7 @@ pub(super) fn build_test_auth_token( ) } -pub(super) async fn issue_test_admin_access_token( +pub(in crate::tests) async fn issue_test_admin_access_token( state: &AppState, client_device_id: &str, ) -> String { @@ -222,6 +311,7 @@ pub(super) fn sample_proxy_node(node_id: &str) -> StoredProxyNode { Some(1_709_000_000), Some(1_710_000_100), ) + .with_tunnel_generation(TUNNEL_CONTROL_PLANE_TEST_GENERATION.to_string()) } pub(super) fn sample_provider_quota(provider_id: &str) -> StoredProviderQuotaSnapshot { @@ -501,6 +591,67 @@ pub(super) fn sample_key( .expect("key transport should build") } +/// Build a provider catalog key using the same provider/key-bound v2 envelope +/// that production writes use. Most control-plane tests intentionally mount +/// a read-only catalog repository; using a legacy Fernet fixture there would +/// force the reader to migrate the credential and fail closed when no writer +/// is available. Keep `sample_key` for tests that explicitly exercise the +/// legacy migration path, and use this helper for ordinary catalog fixtures. +pub(super) fn sample_bound_key( + id: &str, + provider_id: &str, + api_format: &str, + secret: &str, +) -> StoredProviderCatalogKey { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let mut key = sample_key(id, provider_id, api_format, secret); + key.encrypted_api_key = Some( + bootstrap + .seal_provider_catalog_key_api_key(provider_id, id, secret) + .expect("bound provider api key ciphertext should build"), + ); + key +} + +/// Build the provider-scoped proxy representation used by the catalog reader. +/// Stored proxy credentials are record-bound before they are accepted by +/// read-only repositories, so fixtures must use the same envelope. +pub(super) fn sample_bound_provider_proxy(provider_id: &str, host: &str, password: &str) -> Value { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let purpose = format!( + "provider-catalog-proxy-credential-v2\0scope=provider\0field=password\0record-id-bytes={}\0{provider_id}", + provider_id.len(), + ); + let sealed = + crate::handlers::shared::seal_runtime_secret_payload(&bootstrap, &purpose, password) + .expect("provider proxy password should seal"); + json!({ + "host": host, + "password": format!("aether-provider-catalog-proxy-secret-v2:{sealed}"), + }) +} + +/// Seal an auth-config fixture with the provider/key-bound v2 envelope. +/// Ordinary read-only catalog tests must not rely on the migration writer. +pub(super) fn sample_bound_auth_config(provider_id: &str, key_id: &str, plaintext: &str) -> String { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + bootstrap + .seal_provider_catalog_key_auth_config(provider_id, key_id, plaintext) + .expect("bound provider auth config ciphertext should build") +} + pub(super) fn sample_request_candidate( id: &str, request_id: &str, diff --git a/apps/aether-gateway/src/tests/control/internal.rs b/apps/aether-gateway/src/tests/control/internal.rs index 093de6390..00fb52864 100644 --- a/apps/aether-gateway/src/tests/control/internal.rs +++ b/apps/aether-gateway/src/tests/control/internal.rs @@ -1,7 +1,12 @@ use std::io; use std::sync::{Arc, Mutex}; -use aether_contracts::tunnel::RequestMeta; +use aether_contracts::tunnel::{ + sign_tunnel_relay_request, tunnel_relay_payload_digest, RequestMeta, + 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_data::repository::proxy_nodes::ProxyNodeReadRepository; use axum::body::Body; use axum::routing::{any, post}; @@ -12,13 +17,234 @@ use http::header::HeaderValue; use http::StatusCode; use serde_json::json; use std::collections::HashMap; -use std::time::Duration; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes"; +const RELAY_TEST_SENDER: &str = "execution-runtime-test"; +const RELAY_TEST_NODE_ID: &str = "node-123"; +const RELAY_TEST_TUNNEL_GENERATION: &str = "relay-test-generation-node-123"; +const INTERNAL_GATEWAY_TEST_SECRET: &str = "internal-gateway-test-secret-at-least-32-bytes"; use super::{ - build_router_with_state, sample_proxy_node, start_server, AppState, GatewayDataState, - InMemoryProxyNodeRepository, TRACE_ID_HEADER, + authenticated_tunnel_control_plane_request, + authenticated_tunnel_control_plane_request_for_generation, build_router_with_state, + hash_management_token, sample_management_token, sample_proxy_node, start_server, + with_tunnel_control_plane_key, AppState, GatewayDataState, InMemoryManagementTokenRepository, + InMemoryProxyNodeRepository, InMemoryUserReadRepository, TRACE_ID_HEADER, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, TUNNEL_CONTROL_PLANE_TEST_PSK, }; +const TUNNEL_HEARTBEAT_PATH: &str = "/api/internal/tunnel/heartbeat"; +const TUNNEL_NODE_STATUS_PATH: &str = "/api/internal/tunnel/node-status"; + +fn authenticated_internal_gateway_request( + client: &reqwest::Client, + url: String, + path_and_query: &str, + body: &[u8], + timestamp: u64, + nonce: &str, +) -> reqwest::RequestBuilder { + use aether_contracts::internal_gateway::{ + sign_internal_gateway_request, INTERNAL_GATEWAY_AUTH_NONCE_HEADER, + INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER, INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER, + }; + + let signature = sign_internal_gateway_request( + INTERNAL_GATEWAY_TEST_SECRET.as_bytes(), + "POST", + path_and_query, + timestamp, + nonce, + body, + ); + client + .post(url) + .header("content-type", "application/json") + .header(INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER, timestamp) + .header(INTERNAL_GATEWAY_AUTH_NONCE_HEADER, nonce) + .header(INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER, signature) + .body(body.to_vec()) +} + +#[tokio::test] +async fn internal_gateway_requires_hmac_even_when_peer_is_loopback_and_rejects_replay() { + const PATH: &str = "/api/internal/gateway/resolve"; + const NONCE: &str = "internal-gateway-nonce-00000001"; + + let state = AppState::new() + .expect("gateway should build") + .with_internal_gateway_auth_secret_for_tests(INTERNAL_GATEWAY_TEST_SECRET); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = reqwest::Client::new(); + let body = serde_json::to_vec(&json!({ + "method": "GET", + "path": "/not-a-public-route", + "headers": {} + })) + .expect("request body should encode"); + + let unsigned = client + .post(format!("{gateway_url}{PATH}")) + .header("content-type", "application/json") + .body(body.clone()) + .send() + .await + .expect("unsigned loopback request should complete"); + assert_eq!(unsigned.status(), StatusCode::FORBIDDEN); + + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let accepted = authenticated_internal_gateway_request( + &client, + format!("{gateway_url}{PATH}"), + PATH, + &body, + timestamp, + NONCE, + ) + .send() + .await + .expect("signed request should complete"); + assert_eq!(accepted.status(), StatusCode::OK); + + let replay = authenticated_internal_gateway_request( + &client, + format!("{gateway_url}{PATH}"), + PATH, + &body, + timestamp, + NONCE, + ) + .send() + .await + .expect("replayed request should complete"); + assert_eq!(replay.status(), StatusCode::FORBIDDEN); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn internal_gateway_signature_binds_body_and_disabled_mode_is_not_discoverable() { + const PATH: &str = "/api/internal/gateway/resolve"; + const NONCE: &str = "internal-gateway-nonce-00000002"; + + let state = AppState::new() + .expect("gateway should build") + .with_internal_gateway_auth_secret_for_tests(INTERNAL_GATEWAY_TEST_SECRET); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = reqwest::Client::new(); + let signed_body = serde_json::to_vec(&json!({ + "method": "GET", + "path": "/v1/models", + "headers": {} + })) + .expect("signed body should encode"); + let tampered_body = serde_json::to_vec(&json!({ + "method": "GET", + "path": "/api/admin/providers", + "headers": {} + })) + .expect("tampered body should encode"); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let signature = aether_contracts::internal_gateway::sign_internal_gateway_request( + INTERNAL_GATEWAY_TEST_SECRET.as_bytes(), + "POST", + PATH, + timestamp, + NONCE, + &signed_body, + ); + let tampered = client + .post(format!("{gateway_url}{PATH}")) + .header("content-type", "application/json") + .header( + aether_contracts::internal_gateway::INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER, + timestamp, + ) + .header( + aether_contracts::internal_gateway::INTERNAL_GATEWAY_AUTH_NONCE_HEADER, + NONCE, + ) + .header( + aether_contracts::internal_gateway::INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER, + signature, + ) + .body(tampered_body) + .send() + .await + .expect("tampered request should complete"); + assert_eq!(tampered.status(), StatusCode::FORBIDDEN); + gateway_handle.abort(); + + let disabled = AppState::new() + .expect("gateway should build") + .without_internal_gateway_for_tests(); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(disabled)).await; + let hidden = client + .post(format!("{gateway_url}{PATH}")) + .header("content-type", "application/json") + .body(signed_body) + .send() + .await + .expect("disabled request should complete"); + assert_eq!(hidden.status(), StatusCode::NOT_FOUND); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn internal_gateway_rejects_caller_supplied_user_identity() { + const PATH: &str = "/api/internal/gateway/decision-sync"; + const NONCE: &str = "internal-gateway-nonce-identity01"; + + let state = AppState::new() + .expect("gateway should build") + .with_internal_gateway_auth_secret_for_tests(INTERNAL_GATEWAY_TEST_SECRET); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let body = serde_json::to_vec(&json!({ + "method": "POST", + "path": "/v1/chat/completions", + "headers": { "content-type": "application/json" }, + "body_json": { "model": "gpt-5", "messages": [] }, + "auth_context": { + "user_id": "enumerated-user-id", + "api_key_id": "enumerated-api-key-id", + "access_allowed": true + } + })) + .expect("request body should encode"); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let response = authenticated_internal_gateway_request( + &reqwest::Client::new(), + format!("{gateway_url}{PATH}"), + PATH, + &body, + timestamp, + NONCE, + ) + .send() + .await + .expect("signed identity-injection request should complete"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let payload: serde_json::Value = response.json().await.expect("error should be JSON"); + assert_eq!( + payload["detail"], + "supplied auth_context is not accepted; authenticate through request headers" + ); + + gateway_handle.abort(); +} + fn relay_request_meta( stream: bool, request_timeout_ms: Option, @@ -50,6 +276,105 @@ fn relay_envelope(meta: &RequestMeta, body: &[u8]) -> Vec { envelope } +fn relay_metadata_envelope(envelope: &[u8]) -> &[u8] { + let meta_len = u32::from_be_bytes(envelope[..4].try_into().expect("metadata prefix")) as usize; + &envelope[..4 + meta_len] +} + +fn authenticated_relay_request( + client: &reqwest::Client, + url: String, + owner: &str, + node_id: &str, + envelope: &[u8], +) -> reqwest::RequestBuilder { + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let metadata = relay_metadata_envelope(envelope); + let payload_digest = tunnel_relay_payload_digest(metadata, &envelope[metadata.len()..]); + let signature = sign_tunnel_relay_request( + RELAY_TEST_SECRET.as_bytes(), + RELAY_TEST_SENDER, + owner, + node_id, + "", + false, + timestamp, + &nonce, + &payload_digest, + ); + client + .post(url) + .header(TUNNEL_RELAY_AUTH_SENDER_HEADER, RELAY_TEST_SENDER) + .header(TUNNEL_RELAY_OWNER_INSTANCE_HEADER, owner) + .header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp) + .header(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce) + .header( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + payload_digest.encode_header_value(), + ) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature) +} + +fn authenticated_forwarded_relay_request( + client: &reqwest::Client, + url: String, + owner: &str, + node_id: &str, + envelope: &[u8], +) -> reqwest::RequestBuilder { + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let metadata = relay_metadata_envelope(envelope); + let payload_digest = tunnel_relay_payload_digest(metadata, &envelope[metadata.len()..]); + let signature = sign_tunnel_relay_request( + RELAY_TEST_SECRET.as_bytes(), + RELAY_TEST_SENDER, + owner, + node_id, + RELAY_TEST_SENDER, + false, + timestamp, + &nonce, + &payload_digest, + ); + client + .post(url) + .header(TUNNEL_RELAY_AUTH_SENDER_HEADER, RELAY_TEST_SENDER) + .header(TUNNEL_RELAY_OWNER_INSTANCE_HEADER, owner) + .header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp) + .header(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce) + .header( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + payload_digest.encode_header_value(), + ) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature) + .header( + aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER, + RELAY_TEST_SENDER, + ) +} + +fn relay_test_proxy_node_repository() -> Arc { + Arc::new(InMemoryProxyNodeRepository::seed([sample_proxy_node( + RELAY_TEST_NODE_ID, + ) + .with_tunnel_generation(RELAY_TEST_TUNNEL_GENERATION.to_string())])) +} + +fn relay_test_data_state( + config_values: impl IntoIterator, +) -> GatewayDataState { + GatewayDataState::with_proxy_node_repository_for_tests(relay_test_proxy_node_repository()) + .with_system_config_values_for_tests(config_values) +} + #[tokio::test] async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -65,11 +390,11 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() { }), ); - let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( - "node-123", - )])); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + with_tunnel_control_plane_key(sample_proxy_node("node-123"), TUNNEL_CONTROL_PLANE_TEST_PSK), + ])); - let (upstream_url, upstream_handle) = start_server(upstream).await; + let (_upstream_url, upstream_handle) = start_server(upstream).await; let gateway = build_router_with_state( AppState::new() .expect("gateway should build") @@ -79,28 +404,75 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() { ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/heartbeat")) - .json(&json!({ - "node_id": "node-123", - "heartbeat_id": 77, - "heartbeat_interval": 45, - "active_connections": 5, - "total_requests": 100, - "avg_latency_ms": 12.5, - "failed_requests": 20, - "dns_failures": 30, - "stream_errors": 40, - "window_total_requests": 9, - "window_failed_requests": 1, - "window_dns_failures": 2, - "window_stream_errors": 3, - "proxy_metadata": {"arch": "arm64"}, - "proxy_version": "2.0.0", - })) + let client = reqwest::Client::new(); + let heartbeat = json!({ + "node_id": "node-123", + "heartbeat_session_id": "session-77", + "heartbeat_id": 77, + "heartbeat_interval": 45, + "active_connections": 5, + "total_requests": 100, + "avg_latency_ms": 12.5, + "failed_requests": 20, + "dns_failures": 30, + "stream_errors": 40, + "window_total_requests": 9, + "window_failed_requests": 1, + "window_dns_failures": 2, + "window_stream_errors": 3, + "proxy_metadata": {"arch": "arm64"}, + "proxy_version": "2.0.0" + }); + let anonymous = client + .post(format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}")) + .json(&heartbeat) .send() .await - .expect("request should succeed"); + .expect("anonymous loopback request should complete"); + assert_eq!(anonymous.status(), StatusCode::FORBIDDEN); + + let forged = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}"), + TUNNEL_HEARTBEAT_PATH, + "different-node", + &heartbeat, + ) + .send() + .await + .expect("forged identity request should complete"); + assert_eq!(forged.status(), StatusCode::FORBIDDEN); + + let stale_generation = authenticated_tunnel_control_plane_request_for_generation( + &client, + format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}"), + TUNNEL_HEARTBEAT_PATH, + "node-123", + "deleted-node-generation", + &heartbeat, + ) + .send() + .await + .expect("stale generation request should complete"); + assert_eq!(stale_generation.status(), StatusCode::FORBIDDEN); + let unchanged = repository + .find_proxy_node("node-123") + .await + .expect("node lookup should succeed") + .expect("node should exist"); + assert_eq!(unchanged.total_requests, 0); + assert_eq!(unchanged.active_connections, 0); + + let response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}"), + TUNNEL_HEARTBEAT_PATH, + "node-123", + &heartbeat, + ) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); @@ -119,15 +491,55 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() { assert_eq!(node.dns_failures, 2); assert_eq!(node.stream_errors, 3); + let replay = json!({ + "node_id": "node-123", + "heartbeat_session_id": "session-77", + "heartbeat_id": 77, + "heartbeat_interval": 45, + "active_connections": 5, + "window_total_requests": 9, + "window_failed_requests": 1, + "window_dns_failures": 2, + "window_stream_errors": 3, + "proxy_metadata": {"arch": "arm64"}, + "proxy_version": "2.0.0" + }); + let replay_response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}"), + TUNNEL_HEARTBEAT_PATH, + "node-123", + &replay, + ) + .send() + .await + .expect("replayed request should receive the original ACK"); + assert_eq!(replay_response.status(), StatusCode::OK); + let replay_payload: serde_json::Value = replay_response + .json() + .await + .expect("replayed ACK should be JSON"); + assert_eq!(replay_payload["heartbeat_id"], 77); + + let node_after_replay = repository + .find_proxy_node("node-123") + .await + .expect("node lookup should succeed") + .expect("node should exist"); + assert_eq!(node_after_replay.total_requests, 9); + assert_eq!(node_after_replay.failed_requests, 1); + assert_eq!(node_after_replay.dns_failures, 2); + assert_eq!(node_after_replay.stream_errors, 3); + gateway_handle.abort(); upstream_handle.abort(); } #[tokio::test] async fn gateway_rejects_internal_tunnel_heartbeat_without_heartbeat_id() { - let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( - "node-123", - )])); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + with_tunnel_control_plane_key(sample_proxy_node("node-123"), TUNNEL_CONTROL_PLANE_TEST_PSK), + ])); let gateway = build_router_with_state( AppState::new() @@ -138,16 +550,21 @@ async fn gateway_rejects_internal_tunnel_heartbeat_without_heartbeat_id() { ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/heartbeat")) - .json(&json!({ - "node_id": "node-123", - "heartbeat_interval": 45, - "active_connections": 5 - })) - .send() - .await - .expect("request should succeed"); + let body = json!({ + "node_id": "node-123", + "heartbeat_interval": 45, + "active_connections": 5 + }); + let response = authenticated_tunnel_control_plane_request( + &reqwest::Client::new(), + format!("{gateway_url}{TUNNEL_HEARTBEAT_PATH}"), + TUNNEL_HEARTBEAT_PATH, + "node-123", + &body, + ) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::BAD_REQUEST); @@ -169,11 +586,11 @@ async fn gateway_handles_internal_tunnel_node_status_locally_with_loopback() { }), ); - let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( - "node-123", - )])); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + with_tunnel_control_plane_key(sample_proxy_node("node-123"), TUNNEL_CONTROL_PLANE_TEST_PSK), + ])); - let (upstream_url, upstream_handle) = start_server(upstream).await; + let (_upstream_url, upstream_handle) = start_server(upstream).await; let gateway = build_router_with_state( AppState::new() .expect("gateway should build") @@ -183,17 +600,43 @@ async fn gateway_handles_internal_tunnel_node_status_locally_with_loopback() { ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/node-status")) - .json(&json!({ - "node_id": "node-123", - "connected": true, - "conn_count": 4, - "observed_at_unix_secs": 1_800_000_321u64, - })) + let status_body = json!({ + "node_id": "node-123", + "connected": true, + "conn_count": 4, + "observed_at_unix_secs": 1_800_000_321u64 + }); + let client = reqwest::Client::new(); + let anonymous = client + .post(format!("{gateway_url}{TUNNEL_NODE_STATUS_PATH}")) + .json(&status_body) .send() .await - .expect("request should succeed"); + .expect("anonymous loopback request should complete"); + assert_eq!(anonymous.status(), StatusCode::FORBIDDEN); + + let forged = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}{TUNNEL_NODE_STATUS_PATH}"), + TUNNEL_NODE_STATUS_PATH, + "different-node", + &status_body, + ) + .send() + .await + .expect("forged identity request should complete"); + assert_eq!(forged.status(), StatusCode::FORBIDDEN); + + let response = authenticated_tunnel_control_plane_request( + &client, + format!("{gateway_url}{TUNNEL_NODE_STATUS_PATH}"), + TUNNEL_NODE_STATUS_PATH, + "node-123", + &status_body, + ) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); @@ -225,24 +668,230 @@ async fn gateway_owns_proxy_tunnel_path_without_proxying_upstream() { }), ); - let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state(AppState::new().expect("gateway should build")); + let (_upstream_url, upstream_handle) = start_server(upstream).await; + const NODE_ID: &str = "node-123"; + const SESSION: &str = "0123456789abcdef0123456789abcdef"; + const NONCE: &str = "abcdef0123456789abcdef0123456789"; + const PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + let repository = Arc::new(InMemoryProxyNodeRepository::seed([ + with_tunnel_control_plane_key(sample_proxy_node(NODE_ID), PSK), + ])); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests( + repository, + )); + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() + let unsigned = reqwest::Client::new() .get(format!("{gateway_url}/api/internal/proxy-tunnel")) - .header("x-node-id", "node-123") + .header(http::header::CONNECTION, "upgrade") + .header(http::header::UPGRADE, "websocket") + .header("sec-websocket-version", "13") + .header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==") + .header("x-node-id", NODE_ID) + .header( + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, + ) + .header( + aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, + aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION_STR, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER, + aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_SESSION_HEADER, + SESSION, + ) .send() .await .expect("request should succeed"); + assert_eq!(unsigned.status(), StatusCode::UNAUTHORIZED); + + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs(); + let signature = + aether_contracts::tunnel_security::sign_tunnel_security_handshake_for_generation( + PSK, + NODE_ID, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, + aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED, + SESSION, + aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION, + timestamp, + NONCE, + ) + .expect("handshake proof should sign"); + let send_signed_upgrade = || { + reqwest::Client::new() + .get(format!("{gateway_url}/api/internal/proxy-tunnel")) + .header(http::header::CONNECTION, "upgrade") + .header(http::header::UPGRADE, "websocket") + .header("sec-websocket-version", "13") + .header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==") + .header("x-node-id", NODE_ID) + .header( + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, + ) + .header( + aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, + aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION_STR, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER, + aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_SESSION_HEADER, + SESSION, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, + timestamp.to_string(), + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_NONCE_HEADER, + NONCE, + ) + .header( + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + signature.clone(), + ) + }; + + let accepted = send_signed_upgrade() + .send() + .await + .expect("signed WebSocket upgrade should complete"); + assert_eq!(accepted.status(), StatusCode::SWITCHING_PROTOCOLS); + drop(accepted); + + let replay = send_signed_upgrade() + .send() + .await + .expect("replayed WebSocket upgrade should complete"); + assert_eq!(replay.status(), StatusCode::UNAUTHORIZED); - 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_requires_current_authorized_management_token_for_default_proxy_tunnel() { + const NODE_ID: &str = "node-default-auth"; + const PROTOCOL_VERSION: &str = aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION_STR; + let raw_token = "ae-default-tunnel-auth-token"; + + let state = AppState::new().expect("gateway should build"); + let admin_user = state + .create_local_auth_user_with_settings( + Some("tunnel-admin@example.com".to_string()), + true, + "tunnel-admin".to_string(), + "hash".to_string(), + "admin".to_string(), + None, + None, + None, + None, + ) + .await + .expect("admin user should be created") + .expect("admin user should exist"); + let mut token = sample_management_token( + "token-default-tunnel-auth", + &admin_user.id, + "tunnel-admin", + true, + ); + token.token.allowed_ips = None; + token.token.permissions = Some(json!(["admin:proxy_nodes:admin"])); + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + [token], + [( + hash_management_token(raw_token), + "token-default-tunnel-auth".to_string(), + )], + )); + let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed([sample_proxy_node( + NODE_ID, + ) + .with_runtime_fields( + Some("test".to_string()), + Some("different-registering-admin".to_string()), + Some(1_710_000_000), + None, + None, + None, + None, + Some(1_710_000_010), + None, + Some(1_709_000_000), + Some(1_710_000_100), + )])); + let state = state.with_data_state_for_tests( + GatewayDataState::with_management_token_repository_for_tests(token_repository) + .attach_proxy_node_repository_for_tests(proxy_node_repository) + .with_user_reader(Arc::new(InMemoryUserReadRepository::seed_auth_users([ + admin_user, + ]))), + ); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let upgrade = |token: Option<&str>| { + let request = reqwest::Client::new() + .get(format!("{gateway_url}/api/internal/proxy-tunnel")) + .header(http::header::CONNECTION, "upgrade") + .header(http::header::UPGRADE, "websocket") + .header("sec-websocket-version", "13") + .header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==") + .header("x-node-id", NODE_ID) + .header( + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + TUNNEL_CONTROL_PLANE_TEST_GENERATION, + ) + .header( + aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, + PROTOCOL_VERSION, + ); + match token { + Some(token) => request.bearer_auth(token), + None => request, + } + }; + + let missing = upgrade(None) + .send() + .await + .expect("missing-token upgrade should complete"); + assert_eq!(missing.status(), StatusCode::UNAUTHORIZED); + + let invalid = upgrade(Some("ae-invalid-default-tunnel-token")) + .send() + .await + .expect("invalid-token upgrade should complete"); + assert_eq!(invalid.status(), StatusCode::UNAUTHORIZED); + + let accepted = upgrade(Some(raw_token)) + .send() + .await + .expect("valid-token upgrade should complete"); + assert_eq!(accepted.status(), StatusCode::SWITCHING_PROTOCOLS); + drop(accepted); + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_handles_internal_tunnel_relay_locally_without_proxying_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -259,16 +908,35 @@ async fn gateway_handles_internal_tunnel_relay_locally_without_proxying_upstream ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router_with_state(AppState::new().expect("gateway should build")); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ), + ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .body(Vec::::new()) - .send() - .await - .expect("request should succeed"); + // Authenticate the request so the relay handler reaches its local body + // parser. An empty metadata envelope is malformed and must be rejected + // locally without ever attempting to proxy to an upstream gateway. + let envelope = 0u32.to_be_bytes().to_vec(); + let response = authenticated_relay_request( + &reqwest::Client::new(), + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .body(envelope) + .send() + .await + .expect("request should succeed"); + // Once authenticated, malformed relay metadata is rejected locally; the + // request must never be dispatched to an upstream gateway. assert_eq!(response.status(), StatusCode::BAD_REQUEST); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -309,11 +977,12 @@ async fn gateway_forwards_tunnel_relay_to_attachment_owner() { ); let (owner_url, owner_handle) = start_server(owner).await; - let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![( + let data_state = relay_test_data_state([( "tunnel.attachments.node-123".to_string(), json!({ "gateway_instance_id": "gateway-b", "relay_base_url": owner_url, + "tunnel_generation": RELAY_TEST_TUNNEL_GENERATION, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -322,20 +991,30 @@ async fn gateway_forwards_tunnel_relay_to_attachment_owner() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")), + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; let envelope = relay_envelope(&relay_request_meta(false, None, None), b"relay-envelope"); - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .header(TRACE_ID_HEADER, "trace-owner-forward") - .header(http::header::CONTENT_TYPE, "application/octet-stream") - .body(envelope.clone()) - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = authenticated_relay_request( + &client, + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .header(TRACE_ID_HEADER, "trace-owner-forward") + .header(http::header::CONTENT_TYPE, "application/octet-stream") + .body(envelope.clone()) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); assert_eq!( @@ -369,11 +1048,12 @@ async fn gateway_owner_relay_uses_non_stream_timeout_from_envelope() { ); let (owner_url, owner_handle) = start_server(owner).await; - let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![( + let data_state = relay_test_data_state([( "tunnel.attachments.node-123".to_string(), json!({ "gateway_instance_id": "gateway-b", "relay_base_url": owner_url, + "tunnel_generation": RELAY_TEST_TUNNEL_GENERATION, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -381,7 +1061,11 @@ async fn gateway_owner_relay_uses_non_stream_timeout_from_envelope() { let mut state = AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ); let short_timeout_client = reqwest::Client::builder() .timeout(Duration::from_millis(10)) .build() @@ -400,15 +1084,21 @@ async fn gateway_owner_relay_uses_non_stream_timeout_from_envelope() { Ok::(Bytes::copy_from_slice(&envelope[split_at..])), ])); - let response = reqwest::Client::builder() + let client = reqwest::Client::builder() .timeout(Duration::from_secs(1)) .build() - .expect("request client should build") - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .body(request_body) - .send() - .await - .expect("request should succeed"); + .expect("request client should build"); + let response = authenticated_relay_request( + &client, + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .body(request_body) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); assert_eq!( @@ -444,12 +1134,13 @@ async fn gateway_streams_tunnel_relay_body_to_attachment_owner() { ); let (owner_url, owner_handle) = start_server(owner).await; - let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![ + let data_state = relay_test_data_state([ ( "tunnel.attachments.node-123".to_string(), json!({ "gateway_instance_id": "gateway-b", "relay_base_url": owner_url, + "tunnel_generation": RELAY_TEST_TUNNEL_GENERATION, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -459,7 +1150,11 @@ async fn gateway_streams_tunnel_relay_body_to_attachment_owner() { let mut state = AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ); state.client = reqwest::Client::builder() .timeout(Duration::from_millis(10)) .build() @@ -473,12 +1168,18 @@ async fn gateway_streams_tunnel_relay_body_to_attachment_owner() { let request_body = reqwest::Body::wrap_stream(stream::iter(vec![Ok::( expected_envelope.clone(), )])); - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .body(request_body) - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = authenticated_relay_request( + &client, + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .body(request_body) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); assert_eq!( @@ -507,11 +1208,12 @@ async fn gateway_does_not_forward_tunnel_relay_twice() { ); let (owner_url, owner_handle) = start_server(owner).await; - let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![( + let data_state = relay_test_data_state([( "tunnel.attachments.node-123".to_string(), json!({ "gateway_instance_id": "gateway-b", "relay_base_url": owner_url, + "tunnel_generation": RELAY_TEST_TUNNEL_GENERATION, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -520,19 +1222,27 @@ async fn gateway_does_not_forward_tunnel_relay_twice() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")), + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .header( - aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER, - "gateway-z", - ) - .send() - .await - .expect("request should succeed"); + let envelope = relay_envelope(&relay_request_meta(false, None, None), &[]); + let client = reqwest::Client::new(); + let response = authenticated_forwarded_relay_request( + &client, + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .body(envelope) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(*owner_hits.lock().expect("mutex should lock"), 0); @@ -562,12 +1272,13 @@ async fn gateway_forwards_owner_relay_body_above_recording_limit() { ); let (owner_url, owner_handle) = start_server(owner).await; - let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![ + let data_state = relay_test_data_state([ ( "tunnel.attachments.node-123".to_string(), json!({ "gateway_instance_id": "gateway-b", "relay_base_url": owner_url, + "tunnel_generation": RELAY_TEST_TUNNEL_GENERATION, "conn_count": 1, "observed_at_unix_secs": 4_102_444_800u64, }), @@ -581,7 +1292,11 @@ async fn gateway_forwards_owner_relay_body_above_recording_limit() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")), + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -591,12 +1306,18 @@ async fn gateway_forwards_owner_relay_body_above_recording_limit() { &request_payload, ); assert!(envelope.len() > RECORDING_LIMIT_BYTES); - let response = reqwest::Client::new() - .post(format!("{gateway_url}/api/internal/tunnel/relay/node-123")) - .body(envelope.clone()) - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = authenticated_relay_request( + &client, + format!("{gateway_url}/api/internal/tunnel/relay/node-123"), + "gateway-a", + "node-123", + &envelope, + ) + .body(envelope.clone()) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); assert_eq!( diff --git a/apps/aether-gateway/src/tests/control/mod.rs b/apps/aether-gateway/src/tests/control/mod.rs index 54267bbdd..f8da45035 100644 --- a/apps/aether-gateway/src/tests/control/mod.rs +++ b/apps/aether-gateway/src/tests/control/mod.rs @@ -23,6 +23,7 @@ use aether_data::repository::proxy_nodes::{ InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode, StoredProxyNodeEvent, }; use aether_data::repository::quota::InMemoryProviderQuotaRepository; +use aether_data::repository::users::InMemoryUserReadRepository; use aether_data::repository::wallet::InMemoryWalletRepository; use aether_data_contracts::repository::{ candidates::{RequestCandidateStatus, StoredRequestCandidate}, @@ -49,6 +50,8 @@ mod helpers; mod internal; mod proxy; +pub(super) use helpers::issue_test_admin_access_token as issue_shared_test_admin_access_token; + use super::{ build_router, build_router_with_execution_runtime_override, build_router_with_state, build_state_with_execution_runtime_override, start_server, wait_until, AppState, diff --git a/apps/aether-gateway/src/tests/control/proxy/embeddings.rs b/apps/aether-gateway/src/tests/control/proxy/embeddings.rs index dbe2e8edd..473997192 100644 --- a/apps/aether-gateway/src/tests/control/proxy/embeddings.rs +++ b/apps/aether-gateway/src/tests/control/proxy/embeddings.rs @@ -10,7 +10,7 @@ use serde_json::json; use super::super::{ any, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key, - sample_currently_usable_auth_snapshot, sample_endpoint, sample_key, sample_provider, + sample_bound_key, sample_currently_usable_auth_snapshot, sample_endpoint, sample_provider, start_server, AppState, GatewayDataState, InMemoryAuthApiKeySnapshotRepository, InMemoryProviderCatalogReadRepository, Json, Router, }; @@ -74,7 +74,7 @@ fn embedding_success_state(execution_runtime_url: String) -> AppState { "openai:embedding", "https://api.openai.example", )], - vec![sample_key( + vec![sample_bound_key( "key-upstream-embedding", "provider-embedding", "openai:embedding", @@ -135,7 +135,7 @@ fn gemini_embedding_success_state( "gemini:embedding", "https://generativelanguage.googleapis.com/v1beta", )], - vec![sample_key( + vec![sample_bound_key( "key-upstream-gemini-embedding", "provider-gemini-embedding", "gemini:embedding", @@ -175,7 +175,7 @@ fn vertex_gemini_embedding_success_state(execution_runtime_url: String) -> AppSt ])); let mut provider = sample_provider("provider-vertex-gemini-embedding", "Vertex AI", 1); provider.provider_type = "vertex_ai".to_string(); - let mut key = sample_key( + let mut key = sample_bound_key( "key-upstream-vertex-gemini-embedding", "provider-vertex-gemini-embedding", "gemini:embedding", @@ -233,7 +233,7 @@ fn aliyun_embedding_success_state(execution_runtime_url: String) -> AppState { "aliyun:multimodal_embedding", "https://dashscope.aliyuncs.com", )], - vec![sample_key( + vec![sample_bound_key( "key-upstream-aliyun-embedding", "provider-aliyun-embedding", "aliyun:multimodal_embedding", @@ -301,13 +301,13 @@ fn mixed_embedding_success_state(execution_runtime_url: String) -> AppState { ), ], vec![ - sample_key( + sample_bound_key( "key-upstream-embedding", "provider-embedding", "openai:embedding", "sk-upstream-embedding", ), - sample_key( + sample_bound_key( "key-upstream-aliyun-embedding", "provider-aliyun-embedding", "aliyun:multimodal_embedding", diff --git a/apps/aether-gateway/src/tests/control/proxy/local_denials.rs b/apps/aether-gateway/src/tests/control/proxy/local_denials.rs index fe65d1d76..10eb224dd 100644 --- a/apps/aether-gateway/src/tests/control/proxy/local_denials.rs +++ b/apps/aether-gateway/src/tests/control/proxy/local_denials.rs @@ -1,9 +1,16 @@ use std::sync::{Arc, Mutex}; +use std::time::{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_FORWARDED_BY_HEADER, TUNNEL_RELAY_OWNER_INSTANCE_HEADER, +}; use axum::body::Body; use axum::routing::any; use axum::{extract::Request, Json, Router}; -use http::StatusCode; +use http::{HeaderMap, HeaderValue, Method, StatusCode, Uri}; use serde_json::json; use super::super::{ @@ -14,10 +21,144 @@ use super::super::{ }; use crate::constants::{ CONTROL_ROUTE_CLASS_HEADER, EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_AUTH_DENIED, - GATEWAY_HEADER, TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + FORWARDED_FOR_HEADER, GATEWAY_HEADER, TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER, TRUSTED_AUTH_BALANCE_HEADER, TRUSTED_AUTH_USER_ID_HEADER, + TUNNEL_AFFINITY_FORWARDED_BY_HEADER, TUNNEL_AFFINITY_NODE_ID_HEADER, + TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, }; +const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes"; +const RELAY_TEST_SENDER: &str = "gateway-a"; +const RELAY_TEST_OWNER: &str = "gateway-b"; +const RELAY_TEST_NODE_ID: &str = "node-1"; +const AFFINITY_TEST_BODY: &str = "{\"model\":\"gpt-5\",\"messages\":[]}"; + +fn signed_affinity_headers( + method: &Method, + uri: &Uri, + user_id: &str, + api_key_id: &str, + access_allowed: bool, + balance: Option<&str>, + body: &[u8], +) -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert( + GATEWAY_HEADER, + HeaderValue::from_static("rust-phase3b-affinity"), + ); + headers.insert( + TUNNEL_RELAY_FORWARDED_BY_HEADER, + HeaderValue::from_static(RELAY_TEST_SENDER), + ); + headers.insert( + TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + HeaderValue::from_static(RELAY_TEST_SENDER), + ); + headers.insert( + TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + HeaderValue::from_static(RELAY_TEST_OWNER), + ); + headers.insert( + TUNNEL_AFFINITY_NODE_ID_HEADER, + HeaderValue::from_static(RELAY_TEST_NODE_ID), + ); + headers.insert( + TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_str(user_id).expect("trusted user ID should be a valid header value"), + ); + headers.insert( + TRUSTED_AUTH_API_KEY_ID_HEADER, + HeaderValue::from_str(api_key_id) + .expect("trusted API key ID should be a valid header value"), + ); + headers.insert( + TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + HeaderValue::from_static(if access_allowed { "true" } else { "false" }), + ); + headers.insert( + FORWARDED_FOR_HEADER, + HeaderValue::from_static("203.0.113.10"), + ); + if let Some(balance) = balance { + headers.insert( + TRUSTED_AUTH_BALANCE_HEADER, + HeaderValue::from_str(balance).expect("trusted balance should be a valid header value"), + ); + } + + let metadata = crate::tunnel::build_tunnel_affinity_auth_metadata(method, uri, &headers) + .expect("affinity authentication metadata should build"); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock should be after epoch") + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let payload_digest = tunnel_relay_payload_digest(&metadata, body); + let signature = sign_tunnel_relay_request( + RELAY_TEST_SECRET.as_bytes(), + RELAY_TEST_SENDER, + RELAY_TEST_OWNER, + RELAY_TEST_NODE_ID, + RELAY_TEST_SENDER, + false, + timestamp, + &nonce, + &payload_digest, + ); + headers.insert( + TUNNEL_RELAY_AUTH_SENDER_HEADER, + HeaderValue::from_static(RELAY_TEST_SENDER), + ); + headers.insert( + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + HeaderValue::from_static(RELAY_TEST_OWNER), + ); + headers.insert( + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + HeaderValue::from_str(×tamp.to_string()).expect("timestamp should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_NONCE_HEADER, + HeaderValue::from_str(&nonce).expect("nonce should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + HeaderValue::from_str(&payload_digest.encode_header_value()) + .expect("payload digest should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + HeaderValue::from_str(&signature).expect("signature should be a valid header"), + ); + headers +} + +fn signed_affinity_request( + client: &reqwest::Client, + url: String, + path: &str, + user_id: &str, + api_key_id: &str, + access_allowed: bool, + balance: Option<&str>, + body: &[u8], +) -> reqwest::RequestBuilder { + let method = Method::POST; + let uri = path.parse::().expect("request path should be valid"); + client + .request(method.clone(), url) + .headers(signed_affinity_headers( + &method, + &uri, + user_id, + api_key_id, + access_allowed, + balance, + body, + )) +} + #[tokio::test] async fn gateway_locally_denies_explicit_trusted_balance_failure_without_hitting_control_or_upstream( ) { @@ -63,23 +204,32 @@ async fn gateway_locally_denies_explicit_trusted_balance_failure_without_hitting let gateway = build_router_with_state( AppState::new() .expect("gateway state should build") - .with_auth_api_key_data_reader_for_tests(repository), + .with_auth_api_key_data_reader_for_tests(repository) + .with_tunnel_identity_and_relay_secret_for_tests( + RELAY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/v1/chat/completions")) - .header(http::header::CONTENT_TYPE, "application/json") - .header(TRACE_ID_HEADER, "trace-control-balance-denied-1") - .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") - .header(TRUSTED_AUTH_USER_ID_HEADER, "user-123") - .header(TRUSTED_AUTH_API_KEY_ID_HEADER, "key-123") - .header(TRUSTED_AUTH_BALANCE_HEADER, "0") - .header(TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, "false") - .body("{\"model\":\"gpt-5\",\"messages\":[]}") - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = signed_affinity_request( + &client, + format!("{gateway_url}/v1/chat/completions"), + "/v1/chat/completions", + "user-123", + "key-123", + false, + Some("0"), + AFFINITY_TEST_BODY.as_bytes(), + ) + .header(http::header::CONTENT_TYPE, "application/json") + .header(TRACE_ID_HEADER, "trace-control-balance-denied-1") + .body(AFFINITY_TEST_BODY) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); assert_eq!( @@ -152,21 +302,32 @@ async fn gateway_locally_denies_invalid_trusted_snapshot_without_hitting_control let gateway = build_router_with_state( AppState::new() .expect("gateway state should build") - .with_auth_api_key_data_reader_for_tests(repository), + .with_auth_api_key_data_reader_for_tests(repository) + .with_tunnel_identity_and_relay_secret_for_tests( + RELAY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/v1/chat/completions")) - .header(http::header::CONTENT_TYPE, "application/json") - .header(TRACE_ID_HEADER, "trace-control-invalid-trusted-1") - .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") - .header(TRUSTED_AUTH_USER_ID_HEADER, "user-123") - .header(TRUSTED_AUTH_API_KEY_ID_HEADER, "key-123") - .body("{\"model\":\"gpt-5\",\"messages\":[]}") - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = signed_affinity_request( + &client, + format!("{gateway_url}/v1/chat/completions"), + "/v1/chat/completions", + "user-123", + "key-123", + true, + None, + AFFINITY_TEST_BODY.as_bytes(), + ) + .header(http::header::CONTENT_TYPE, "application/json") + .header(TRACE_ID_HEADER, "trace-control-invalid-trusted-1") + .body(AFFINITY_TEST_BODY) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::UNAUTHORIZED); assert_eq!( @@ -227,21 +388,32 @@ async fn gateway_locally_denies_missing_wallet_without_hitting_control_or_upstre let gateway = build_router_with_state( AppState::new() .expect("gateway state should build") - .with_data_state_for_tests(data_state), + .with_data_state_for_tests(data_state) + .with_tunnel_identity_and_relay_secret_for_tests( + RELAY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/v1/chat/completions")) - .header(http::header::CONTENT_TYPE, "application/json") - .header(TRACE_ID_HEADER, "trace-control-wallet-missing-1") - .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") - .header(TRUSTED_AUTH_USER_ID_HEADER, "user-123") - .header(TRUSTED_AUTH_API_KEY_ID_HEADER, "key-123") - .body("{\"model\":\"gpt-5\",\"messages\":[]}") - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = signed_affinity_request( + &client, + format!("{gateway_url}/v1/chat/completions"), + "/v1/chat/completions", + "user-123", + "key-123", + true, + None, + AFFINITY_TEST_BODY.as_bytes(), + ) + .header(http::header::CONTENT_TYPE, "application/json") + .header(TRACE_ID_HEADER, "trace-control-wallet-missing-1") + .body(AFFINITY_TEST_BODY) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::FORBIDDEN); assert_eq!( @@ -757,21 +929,32 @@ async fn gateway_locally_denies_locked_trusted_snapshot_without_hitting_control_ let gateway = build_router_with_state( AppState::new() .expect("gateway state should build") - .with_auth_api_key_data_reader_for_tests(repository), + .with_auth_api_key_data_reader_for_tests(repository) + .with_tunnel_identity_and_relay_secret_for_tests( + RELAY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() - .post(format!("{gateway_url}/v1/chat/completions")) - .header(http::header::CONTENT_TYPE, "application/json") - .header(TRACE_ID_HEADER, "trace-control-locked-trusted-1") - .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") - .header(TRUSTED_AUTH_USER_ID_HEADER, "user-locked-123") - .header(TRUSTED_AUTH_API_KEY_ID_HEADER, "key-locked-123") - .body("{\"model\":\"gpt-5\",\"messages\":[]}") - .send() - .await - .expect("request should succeed"); + let client = reqwest::Client::new(); + let response = signed_affinity_request( + &client, + format!("{gateway_url}/v1/chat/completions"), + "/v1/chat/completions", + "user-locked-123", + "key-locked-123", + true, + None, + AFFINITY_TEST_BODY.as_bytes(), + ) + .header(http::header::CONTENT_TYPE, "application/json") + .header(TRACE_ID_HEADER, "trace-control-locked-trusted-1") + .body(AFFINITY_TEST_BODY) + .send() + .await + .expect("request should succeed"); assert_eq!(response.status(), StatusCode::FORBIDDEN); assert_eq!( diff --git a/apps/aether-gateway/src/tests/control/proxy/rerank.rs b/apps/aether-gateway/src/tests/control/proxy/rerank.rs index 81c0fadf9..5d7cb95f3 100644 --- a/apps/aether-gateway/src/tests/control/proxy/rerank.rs +++ b/apps/aether-gateway/src/tests/control/proxy/rerank.rs @@ -10,7 +10,7 @@ use serde_json::json; use super::super::{ any, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key, - sample_currently_usable_auth_snapshot, sample_endpoint, sample_key, sample_provider, + sample_bound_key, sample_currently_usable_auth_snapshot, sample_endpoint, sample_provider, start_server, AppState, GatewayDataState, InMemoryAuthApiKeySnapshotRepository, InMemoryProviderCatalogReadRepository, Json, Router, }; @@ -69,7 +69,7 @@ fn rerank_success_state(execution_runtime_url: String) -> AppState { "openai:rerank", "https://api.openai.example", )], - vec![sample_key( + vec![sample_bound_key( "key-upstream-rerank", "provider-rerank", "openai:rerank", diff --git a/apps/aether-gateway/src/tests/files/mod.rs b/apps/aether-gateway/src/tests/files/mod.rs index bc374ba78..a9c9b1e7c 100644 --- a/apps/aether-gateway/src/tests/files/mod.rs +++ b/apps/aether-gateway/src/tests/files/mod.rs @@ -4,13 +4,17 @@ use super::{ Mutex, Request, Response, Router, StatusCode, CONTROL_EXECUTED_HEADER, CONTROL_EXECUTE_FALLBACK_HEADER, EXECUTION_PATH_HEADER, TRACE_ID_HEADER, }; -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, }; use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; +use aether_data::repository::gemini_file_mappings::{ + InMemoryGeminiFileMappingRepository, StoredGeminiFileMapping, +}; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; +use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use aether_data_contracts::repository::candidate_selection::{ StoredMinimalCandidateSelectionRow, StoredProviderModelMapping, }; @@ -168,6 +172,12 @@ fn sample_files_provider_catalog_endpoint() -> StoredProviderCatalogEndpoint { } fn sample_files_provider_catalog_key() -> StoredProviderCatalogKey { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); StoredProviderCatalogKey::new( "key-gemini-files-local-1".to_string(), "provider-gemini-files-local-1".to_string(), @@ -179,7 +189,12 @@ fn sample_files_provider_catalog_key() -> StoredProviderCatalogKey { .expect("key should build") .with_transport_fields( Some(serde_json::json!(["gemini:files"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-upstream-gemini-files") + bootstrap + .seal_provider_catalog_key_api_key( + "provider-gemini-files-local-1", + "key-gemini-files-local-1", + "sk-upstream-gemini-files", + ) .expect("api key should encrypt"), None, Some(serde_json::json!({"gemini_files": true})), @@ -192,6 +207,52 @@ fn sample_files_provider_catalog_key() -> StoredProviderCatalogKey { .expect("key transport should build") } +pub(super) fn sample_files_proxy_node_repository( + node_ids: I, +) -> Arc +where + I: IntoIterator, + S: AsRef, +{ + let nodes = node_ids.into_iter().map(|node_id| { + let node_id = node_id.as_ref(); + StoredProxyNode::new( + node_id.to_string(), + format!("files-test-{node_id}"), + "127.0.0.1".to_string(), + 1, + true, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + false, + false, + 1, + ) + .expect("files test proxy node should build") + .with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None) + .with_tunnel_generation(format!("files-test-generation-{node_id}")) + }); + Arc::new(InMemoryProxyNodeRepository::seed(nodes)) +} + +fn sample_owned_file_mapping(file_name: &str, user_id: &str) -> StoredGeminiFileMapping { + let mut mapping = StoredGeminiFileMapping::new( + format!("mapping-{user_id}-{file_name}"), + file_name.to_string(), + "key-gemini-files-local-1".to_string(), + 1_700_000_000_000, + 4_102_444_800, + ) + .expect("file mapping should build"); + mapping.user_id = Some(user_id.to_string()); + mapping +} + #[test] fn gateway_locally_denies_gemini_files_download_control_sync_even_with_opt_in_headers_when_execution_runtime_missing( ) { @@ -253,13 +314,9 @@ async fn gateway_locally_denies_gemini_files_download_control_sync_even_with_opt .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"], - "当前 Gemini Files 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload["detail"], "File not found"); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); @@ -337,13 +394,9 @@ async fn gateway_locally_denies_gemini_files_download_control_sync_without_opt_i .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"], - "当前 Gemini Files 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload["detail"], "File not found"); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!( @@ -416,13 +469,9 @@ async fn gateway_skips_gemini_files_download_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"], - "当前 Gemini Files 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload["detail"], "File not found"); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); @@ -431,14 +480,15 @@ async fn gateway_skips_gemini_files_download_control_sync_without_opt_in_header_ } #[test] -fn gateway_executes_gemini_files_get_via_local_decision_gate_with_local_planning_only() { +fn gateway_rejects_gemini_files_get_key_mismatch_without_fallback_and_allows_creation_key() { run_files_test( - "gateway_executes_gemini_files_get_via_local_decision_gate_with_local_planning_only", - gateway_executes_gemini_files_get_via_local_decision_gate_with_local_planning_only_impl, + "gateway_rejects_gemini_files_get_key_mismatch_without_fallback_and_allows_creation_key", + gateway_rejects_gemini_files_get_key_mismatch_without_fallback_and_allows_creation_key_impl, ); } -async fn gateway_executes_gemini_files_get_via_local_decision_gate_with_local_planning_only_impl() { +async fn gateway_rejects_gemini_files_get_key_mismatch_without_fallback_and_allows_creation_key_impl( +) { #[derive(Debug, Clone)] struct SeenExecutionRuntimeSyncRequest { method: String, @@ -566,15 +616,29 @@ async fn gateway_executes_gemini_files_get_via_local_decision_gate_with_local_pl }), ); - let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( - Some(hash_api_key("client-files-local-key")), - sample_auth_snapshot("key-files-local-123", "user-files-local-123"), - )])); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![ + ( + Some(hash_api_key("client-files-local-key")), + sample_auth_snapshot("key-files-local-123", "user-files-local-123"), + ), + ( + Some(hash_api_key("client-files-local-rotated-key")), + sample_auth_snapshot("key-files-local-rotated-123", "user-files-local-123"), + ), + ])); let candidate_selection_repository = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ sample_files_candidate_row(), ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let mut mismatched_key_mapping = + sample_owned_file_mapping("files/key-mismatch", "user-files-local-123"); + mismatched_key_mapping.key_id = "key-gemini-files-other".to_string(); + let gemini_file_mapping_repository = Arc::new(InMemoryGeminiFileMappingRepository::seed([ + sample_owned_file_mapping("files/abc-123", "user-files-local-123"), + sample_owned_file_mapping("files/foreign", "user-files-foreign-123"), + mismatched_key_mapping, + ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_files_provider_catalog_provider()], vec![sample_files_provider_catalog_endpoint()], @@ -586,20 +650,76 @@ async fn gateway_executes_gemini_files_get_via_local_decision_gate_with_local_pl 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( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests( auth_repository, candidate_selection_repository, provider_catalog_repository, Arc::clone(&request_candidate_repository), + gemini_file_mapping_repository, DEVELOPMENT_ENCRYPTION_KEY, ), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() + let mismatch_response = reqwest::Client::new() .get(format!( - "{gateway_url}/v1beta/files/files/abc-123?view=FULL&key=client-files-local-key" + "{gateway_url}/v1beta/files/key-mismatch?key=client-files-local-key" + )) + .header("x-goog-api-key", "client-header-key") + .header(TRACE_ID_HEADER, "trace-gemini-files-key-mismatch-local-123") + .send() + .await + .expect("key mismatch request should complete"); + + assert_eq!(mismatch_response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert!( + seen_execution_runtime + .lock() + .expect("mutex should lock") + .is_none(), + "a candidate using a different creation key must not execute" + ); + assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); + + let client = reqwest::Client::new(); + for (method, path) in [ + ("GET", "/v1beta/files/foreign"), + ("DELETE", "/v1beta/files/foreign"), + ("GET", "/v1beta/files/foreign:download?alt=media"), + ] { + let request = match method { + "DELETE" => client.delete(format!("{gateway_url}{path}?key=client-files-local-key")), + _ if path.contains('?') => { + client.get(format!("{gateway_url}{path}&key=client-files-local-key")) + } + _ => client.get(format!("{gateway_url}{path}?key=client-files-local-key")), + }; + let foreign_response = request + .header("x-goog-api-key", "client-header-key") + .send() + .await + .expect("foreign file request should complete"); + assert_eq!(foreign_response.status(), StatusCode::NOT_FOUND); + assert_eq!( + foreign_response + .json::() + .await + .expect("foreign response should be JSON"), + json!({"detail": "File not found"}) + ); + } + assert!( + seen_execution_runtime + .lock() + .expect("mutex should lock") + .is_none(), + "cross-user file requests must not reach execution runtime" + ); + + let response = client + .get(format!( + "{gateway_url}/v1beta/files/files/abc-123?view=FULL&key=client-files-local-rotated-key" )) .header("x-goog-api-key", "client-header-key") .header(TRACE_ID_HEADER, "trace-gemini-files-local-123") diff --git a/apps/aether-gateway/src/tests/files/stream.rs b/apps/aether-gateway/src/tests/files/stream.rs index 895f304dd..884927092 100644 --- a/apps/aether-gateway/src/tests/files/stream.rs +++ b/apps/aether-gateway/src/tests/files/stream.rs @@ -5,8 +5,8 @@ use super::{ any, build_router, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key, json, sample_auth_snapshot, sample_files_candidate_row, sample_files_provider_catalog_endpoint, sample_files_provider_catalog_key, - sample_files_provider_catalog_provider, start_server, to_bytes, Arc, Body, Bytes, HeaderName, - HeaderValue, InMemoryAuthApiKeySnapshotRepository, + sample_files_provider_catalog_provider, sample_files_proxy_node_repository, start_server, + to_bytes, Arc, Body, Bytes, HeaderName, HeaderValue, InMemoryAuthApiKeySnapshotRepository, InMemoryMinimalCandidateSelectionReadRepository, InMemoryProviderCatalogReadRepository, InMemoryRequestCandidateRepository, Infallible, Json, Mutex, Request, RequestCandidateReadRepository, RequestCandidateStatus, Response, Router, StatusCode, @@ -193,6 +193,10 @@ async fn gateway_executes_gemini_files_download_via_local_decision_gate_with_loc sample_files_candidate_row(), ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let gemini_file_mapping_repository = + Arc::new(super::InMemoryGeminiFileMappingRepository::seed([ + super::sample_owned_file_mapping("files/file-123", "user-files-download-local-123"), + ])); let mut provider = sample_files_provider_catalog_provider(); provider.proxy = Some(serde_json::json!({"url":"http://provider-proxy.internal:8080"})); let mut endpoint = sample_files_provider_catalog_endpoint(); @@ -216,13 +220,17 @@ async fn gateway_executes_gemini_files_download_via_local_decision_gate_with_loc 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( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests( auth_repository, candidate_selection_repository, provider_catalog_repository, Arc::clone(&request_candidate_repository), + gemini_file_mapping_repository, DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .attach_proxy_node_repository_for_tests(sample_files_proxy_node_repository([ + "proxy-node-gemini-files-download-local", + ])), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; diff --git a/apps/aether-gateway/src/tests/files/sync.rs b/apps/aether-gateway/src/tests/files/sync.rs index a9f3f28f0..f4001a858 100644 --- a/apps/aether-gateway/src/tests/files/sync.rs +++ b/apps/aether-gateway/src/tests/files/sync.rs @@ -3,11 +3,12 @@ use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine as _}; use super::{ any, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key, json, sample_auth_snapshot, sample_files_candidate_row, sample_files_provider_catalog_endpoint, - sample_files_provider_catalog_key, sample_files_provider_catalog_provider, start_server, - to_bytes, Arc, Body, InMemoryAuthApiKeySnapshotRepository, - InMemoryMinimalCandidateSelectionReadRepository, InMemoryProviderCatalogReadRepository, - InMemoryRequestCandidateRepository, Json, Mutex, Request, RequestCandidateReadRepository, - RequestCandidateStatus, Router, StatusCode, DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER, + sample_files_provider_catalog_key, sample_files_provider_catalog_provider, + sample_files_proxy_node_repository, start_server, to_bytes, Arc, Body, + InMemoryAuthApiKeySnapshotRepository, InMemoryMinimalCandidateSelectionReadRepository, + InMemoryProviderCatalogReadRepository, InMemoryRequestCandidateRepository, Json, Mutex, + Request, RequestCandidateReadRepository, RequestCandidateStatus, Router, StatusCode, + DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER, }; #[test] @@ -228,7 +229,10 @@ async fn gateway_executes_gemini_files_upload_via_local_decision_gate_with_local provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .attach_proxy_node_repository_for_tests(sample_files_proxy_node_repository([ + "proxy-node-gemini-files-upload-local", + ])), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -318,19 +322,12 @@ fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_plannin async fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_planning_only_impl() { - #[derive(Debug, Clone)] - struct SeenExecutionRuntimeSyncRequest { - method: String, - url: String, - auth_header_value: String, - } - let decision_hits = Arc::new(Mutex::new(0usize)); let decision_hits_clone = Arc::clone(&decision_hits); let plan_hits = Arc::new(Mutex::new(0usize)); let plan_hits_clone = Arc::clone(&plan_hits); - let seen_execution_runtime = Arc::new(Mutex::new(None::)); - let seen_execution_runtime_clone = Arc::clone(&seen_execution_runtime); + let execution_runtime_hits = Arc::new(Mutex::new(0usize)); + let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits); let report_hits = Arc::new(Mutex::new(0usize)); let report_hits_clone = Arc::clone(&report_hits); let public_hits = Arc::new(Mutex::new(0usize)); @@ -399,48 +396,13 @@ async fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_p let execution_runtime = Router::new().route( "/v1/execute/sync", - any(move |request: Request| { - let seen_execution_runtime_inner = Arc::clone(&seen_execution_runtime_clone); + any(move |_request: Request| { + let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_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 + *execution_runtime_hits_inner .lock() - .expect("mutex should lock") = Some(SeenExecutionRuntimeSyncRequest { - method: payload - .get("method") - .and_then(|value| value.as_str()) - .unwrap_or_default() - .to_string(), - url: payload - .get("url") - .and_then(|value| value.as_str()) - .unwrap_or_default() - .to_string(), - auth_header_value: payload - .get("headers") - .and_then(|value| value.get("x-goog-api-key")) - .and_then(|value| value.as_str()) - .unwrap_or_default() - .to_string(), - }); - Json(json!({ - "request_id": "trace-gemini-files-list-local-123", - "status_code": 200, - "headers": { - "content-type": "application/json" - }, - "body": { - "json_body": { - "files": [] - } - }, - "telemetry": { - "elapsed_ms": 14 - } - })) + .expect("mutex should lock") += 1; + StatusCode::IM_A_TEAPOT } }), ); @@ -454,6 +416,11 @@ async fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_p sample_files_candidate_row(), ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let gemini_file_mapping_repository = + Arc::new(super::InMemoryGeminiFileMappingRepository::seed([ + super::sample_owned_file_mapping("files/owned-list-file", "user-files-list-local-123"), + super::sample_owned_file_mapping("files/other-list-file", "user-other"), + ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_files_provider_catalog_provider()], vec![sample_files_provider_catalog_endpoint()], @@ -465,11 +432,12 @@ async fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_p 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( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests( auth_repository, candidate_selection_repository, provider_catalog_repository, Arc::clone(&request_candidate_repository), + gemini_file_mapping_repository, DEVELOPMENT_ENCRYPTION_KEY, ), ); @@ -487,32 +455,19 @@ async fn gateway_executes_gemini_files_list_via_local_decision_gate_with_local_p .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = response.json().await.expect("body should parse"); + assert_eq!(payload["files"].as_array().map(Vec::len), Some(1)); + assert_eq!(payload["files"][0]["name"], "files/owned-list-file"); assert_eq!( - response.text().await.expect("body should read"), - "{\"files\":[]}" - ); - - 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.method, "GET"); - assert_eq!( - seen_execution_runtime_request.url, - "https://generativelanguage.googleapis.com/v1beta/files?pageSize=20" - ); - assert_eq!( - seen_execution_runtime_request.auth_header_value, - "sk-upstream-gemini-files" + *execution_runtime_hits.lock().expect("mutex should lock"), + 0 ); let stored_candidates = request_candidate_repository .list_by_request_id("trace-gemini-files-list-local-123") .await .expect("request candidate trace should read"); - assert_eq!(stored_candidates.len(), 1); - assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Success); + assert!(stored_candidates.is_empty()); assert_eq!(*report_hits.lock().expect("mutex should lock"), 0); assert_eq!(*decision_hits.lock().expect("mutex should lock"), 0); @@ -668,6 +623,10 @@ async fn gateway_executes_gemini_files_delete_via_local_decision_gate_with_local sample_files_candidate_row(), ])); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let gemini_file_mapping_repository = + Arc::new(super::InMemoryGeminiFileMappingRepository::seed([ + super::sample_owned_file_mapping("files/abc-123", "user-files-delete-local-123"), + ])); let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_files_provider_catalog_provider()], vec![sample_files_provider_catalog_endpoint()], @@ -679,11 +638,12 @@ async fn gateway_executes_gemini_files_delete_via_local_decision_gate_with_local 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( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests( auth_repository, candidate_selection_repository, provider_catalog_repository, Arc::clone(&request_candidate_repository), + gemini_file_mapping_repository, DEVELOPMENT_ENCRYPTION_KEY, ), ); diff --git a/apps/aether-gateway/src/tests/fixtures/admin_system/config_export_v22.json b/apps/aether-gateway/src/tests/fixtures/admin_system/config_export_v22.json index bd1a2f146..d7857f47b 100644 --- a/apps/aether-gateway/src/tests/fixtures/admin_system/config_export_v22.json +++ b/apps/aether-gateway/src/tests/fixtures/admin_system/config_export_v22.json @@ -88,6 +88,16 @@ "value": "Legacy Fixture v22", "description": "Site name from fixture" }, + { + "key": "smtp_host", + "value": "smtp.example.com", + "description": "SMTP host from fixture" + }, + { + "key": "smtp_user", + "value": "smtp-user", + "description": "SMTP user from fixture" + }, { "key": "smtp_password", "value": "smtp-secret-v22", diff --git a/apps/aether-gateway/src/tests/frontdoor/ai.rs b/apps/aether-gateway/src/tests/frontdoor/ai.rs index 13bcd07e4..1ba25afc6 100644 --- a/apps/aether-gateway/src/tests/frontdoor/ai.rs +++ b/apps/aether-gateway/src/tests/frontdoor/ai.rs @@ -5,6 +5,7 @@ use super::{ InMemoryVideoTaskRepository, StoredAuthApiKeySnapshot, UpsertVideoTask, VideoTaskLookupKey, VideoTaskReadRepository, VideoTaskStatus, VideoTaskWriteRepository, DEVELOPMENT_ENCRYPTION_KEY, }; +use crate::data::GatewayDataState; use crate::image_capabilities::openai_image_gateway_max_generation_count; use crate::tests::{ any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server, @@ -234,6 +235,12 @@ fn codex_catalog_key( key_id: &str, allowed_models: &[&str], ) -> StoredProviderCatalogKey { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); let mut key = StoredProviderCatalogKey::new( key_id.to_string(), provider_id.to_string(), @@ -245,7 +252,8 @@ fn codex_catalog_key( .expect("Codex key should build") .with_transport_fields( Some(json!(["openai:responses"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "oauth-upstream-secret") + bootstrap + .seal_provider_catalog_key_api_key(provider_id, key_id, "oauth-upstream-secret") .expect("Codex test token should encrypt"), None, None, @@ -411,6 +419,84 @@ fn sample_gemini_video_task( } } +fn gemini_video_catalog_repository() -> Arc { + const PROVIDER_ID: &str = "provider-gemini-video-local-1"; + const ENDPOINT_ID: &str = "endpoint-gemini-video-local-1"; + const KEY_ID: &str = "key-gemini-video-local-1"; + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let provider = StoredProviderCatalogProvider::new( + PROVIDER_ID.to_string(), + "gemini-video".to_string(), + Some("https://generativelanguage.googleapis.com".to_string()), + "gemini".to_string(), + ) + .expect("Gemini video provider should build") + .with_transport_fields( + true, + false, + false, + None, + Some(2), + None, + Some(20.0), + None, + None, + ); + let endpoint = StoredProviderCatalogEndpoint::new( + ENDPOINT_ID.to_string(), + PROVIDER_ID.to_string(), + "gemini:video".to_string(), + Some("gemini".to_string()), + Some("video".to_string()), + true, + ) + .expect("Gemini video endpoint should build") + .with_transport_fields( + "https://generativelanguage.googleapis.com".to_string(), + None, + None, + Some(2), + None, + None, + None, + None, + ) + .expect("Gemini video endpoint transport should build"); + let key = StoredProviderCatalogKey::new( + KEY_ID.to_string(), + PROVIDER_ID.to_string(), + "prod".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("Gemini video key should build") + .with_transport_fields( + Some(json!(["gemini:video"])), + bootstrap + .seal_provider_catalog_key_api_key(PROVIDER_ID, KEY_ID, "sk-upstream-gemini-video") + .expect("Gemini video api key should encrypt"), + None, + None, + None, + None, + None, + None, + None, + ) + .expect("Gemini video key transport should build"); + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![key], + )) +} + struct PendingMinimalCandidateSelectionReadRepository; impl PendingMinimalCandidateSelectionReadRepository { @@ -657,13 +743,19 @@ fn gateway_versioned_models_fail_closed_when_cached_auth_becomes_unusable_or_mis } async fn run_versioned_models_auth_race_scenario() { + let mut auth_race_snapshot = codex_models_snapshot( + "key-codex-models-auth-race", + "user-codex-models-auth-race", + &["future-alias"], + ); + // This scenario exercises API-key cache invalidation, not provider + // allowlist resolution. Keep the catalog dependency absent so the warm + // request remains on the local models route while the key is usable. + auth_race_snapshot.user_allowed_providers = None; + auth_race_snapshot.api_key_allowed_providers = None; let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( Some(hash_api_key("sk-codex-models-auth-race")), - codex_models_snapshot( - "key-codex-models-auth-race", - "user-codex-models-auth-race", - &["future-alias"], - ), + auth_race_snapshot, )])); let candidate_repository = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed( Vec::new(), @@ -3477,7 +3569,10 @@ async fn gateway_does_not_locally_reject_image_model_name_on_chat_completions() let gateway = build_router_with_state( AppState::new() .expect("gateway should build") - .with_auth_api_key_data_reader_for_tests(auth_repository), + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository) + .with_system_default_routing_group_for_tests(), + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -3704,6 +3799,231 @@ async fn gateway_handles_gemini_operation_detail_without_hitting_fallback_probe( fallback_probe_handle.abort(); } +#[tokio::test] +async fn gateway_hides_gemini_operation_detail_and_cancel_from_non_owner() { + let fallback_probe_hits = Arc::new(Mutex::new(0usize)); + let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits); + let fallback_probe = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let fallback_probe_hits_inner = Arc::clone(&fallback_probe_hits_clone); + async move { + *fallback_probe_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Json(json!({"proxied": true}))).into_response() + } + }), + ); + + let execution_runtime_hits = Arc::new(Mutex::new(0usize)); + let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits); + let execution_runtime = Router::new().route( + "/v1/execute/sync", + any(move |_request: Request| { + let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_clone); + async move { + *execution_runtime_hits_inner + .lock() + .expect("mutex should lock") += 1; + Json(json!({ + "request_id": "unexpected-cross-user-cancel", + "status_code": 200, + "headers": {}, + "body": { "json_body": {} } + })) + } + }), + ); + + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("sk-gemini-operation-non-owner")), + unrestricted_models_snapshot( + "key-gemini-operation-non-owner", + "user-gemini-operation-non-owner", + ), + )])); + let repository = Arc::new(InMemoryVideoTaskRepository::default()); + repository + .upsert(sample_gemini_video_task( + "task-gemini-operation-owner", + "opshort-owner-only", + "user-gemini-operation-owner", + "key-gemini-operation-owner", + "operations/ext-owner-only", + VideoTaskStatus::Submitted, + )) + .await + .expect("upsert should succeed"); + + let (_fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await; + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let gateway = build_router_with_state( + build_state_with_execution_runtime_override(execution_runtime_url) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( + auth_repository, + Arc::clone(&repository), + ), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let detail_response = client + .get(format!( + "{gateway_url}/v1beta/operations/opshort-owner-only?key=sk-gemini-operation-non-owner" + )) + .send() + .await + .expect("cross-user detail request should complete"); + assert_eq!(detail_response.status(), StatusCode::NOT_FOUND); + assert_eq!( + detail_response + .json::() + .await + .expect("json body should parse"), + json!({ "detail": "Video task not found" }) + ); + + let cancel_response = client + .post(format!( + "{gateway_url}/v1beta/operations/opshort-owner-only:cancel" + )) + .header("x-goog-api-key", "sk-gemini-operation-non-owner") + .header(http::header::CONTENT_TYPE, "application/json") + .body("{}") + .send() + .await + .expect("cross-user cancel request should complete"); + assert_eq!(cancel_response.status(), StatusCode::NOT_FOUND); + assert_eq!( + cancel_response + .json::() + .await + .expect("json body should parse"), + json!({ "detail": "Video task not found" }) + ); + + let stored = repository + .find(VideoTaskLookupKey::Id("task-gemini-operation-owner")) + .await + .expect("task lookup should succeed") + .expect("task should exist"); + assert_eq!(stored.status, VideoTaskStatus::Submitted); + assert_eq!( + *execution_runtime_hits.lock().expect("mutex should lock"), + 0 + ); + assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + execution_runtime_handle.abort(); + fallback_probe_handle.abort(); +} + +#[tokio::test] +async fn gateway_hides_gemini_video_file_from_non_owner_and_allows_owner_rotated_key() { + let fallback_probe_hits = Arc::new(Mutex::new(0usize)); + let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits); + let fallback_probe = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let fallback_probe_hits_inner = Arc::clone(&fallback_probe_hits_clone); + async move { + *fallback_probe_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Json(json!({"proxied": true}))).into_response() + } + }), + ); + + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![ + ( + Some(hash_api_key("sk-gemini-video-file-foreign")), + unrestricted_models_snapshot( + "key-gemini-video-file-foreign", + "user-gemini-video-file-foreign", + ), + ), + ( + Some(hash_api_key("sk-gemini-video-file-owner")), + unrestricted_models_snapshot( + "key-gemini-video-file-owner-rotated", + "user-gemini-video-file-owner", + ), + ), + ])); + let repository = Arc::new(InMemoryVideoTaskRepository::default()); + let mut task = sample_gemini_video_task( + "task-gemini-video-file-owner", + "opshort-file-owner", + "user-gemini-video-file-owner", + "key-gemini-video-file-original", + "operations/ext-file-owner", + VideoTaskStatus::Completed, + ); + task.provider_api_format = None; + task.video_url = Some("https://8.8.8.8/video-owner.mp4".to_string()); + repository + .upsert(task) + .await + .expect("upsert should succeed"); + + let (_unused_fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await; + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( + auth_repository, + repository, + ), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("client should build"); + + let foreign_response = client + .get(format!( + "{gateway_url}/v1beta/files/aev_opshort-file-owner:download?alt=media&key=sk-gemini-video-file-foreign" + )) + .send() + .await + .expect("foreign request should complete"); + assert_eq!(foreign_response.status(), StatusCode::NOT_FOUND); + assert_eq!( + foreign_response + .json::() + .await + .expect("json body should parse"), + json!({"detail": "File not found"}) + ); + assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0); + + let owner_response = client + .get(format!( + "{gateway_url}/v1beta/files/aev_opshort-file-owner:download?alt=media&key=sk-gemini-video-file-owner" + )) + .send() + .await + .expect("owner request should complete"); + assert_ne!(owner_response.status(), StatusCode::NOT_FOUND); + assert_ne!(owner_response.status(), StatusCode::TEMPORARY_REDIRECT); + assert_eq!(owner_response.status(), StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!( + owner_response + .json::() + .await + .expect("json body should parse"), + json!({"detail": "Service temporarily unavailable"}) + ); + assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + fallback_probe_handle.abort(); +} + #[tokio::test] async fn gateway_lists_gemini_operations_without_hitting_fallback_probe() { let fallback_probe_hits = Arc::new(Mutex::new(0usize)); @@ -3906,13 +4226,16 @@ async fn gateway_cancels_gemini_operation_without_hitting_fallback_probe() { let (fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await; let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let provider_catalog_repository = gemini_video_catalog_repository(); let gateway = build_router_with_state( build_state_with_execution_runtime_override(execution_runtime_url) .with_data_state_for_tests( - crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( - auth_repository, + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( Arc::clone(&repository), - ), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_auth_api_key_reader(auth_repository), ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; diff --git a/apps/aether-gateway/src/tests/frontdoor/internal.rs b/apps/aether-gateway/src/tests/frontdoor/internal.rs index c2fba96f9..bce477a65 100644 --- a/apps/aether-gateway/src/tests/frontdoor/internal.rs +++ b/apps/aether-gateway/src/tests/frontdoor/internal.rs @@ -3,8 +3,10 @@ use super::{ sample_models_candidate_row, sample_provider, unrestricted_models_snapshot, InMemoryAuthApiKeySnapshotRepository, InMemoryMinimalCandidateSelectionReadRepository, InMemoryProviderCatalogReadRepository, InMemoryRequestCandidateRepository, + InMemoryVideoTaskRepository, UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, DEVELOPMENT_ENCRYPTION_KEY, }; +use crate::data::GatewayDataState; use crate::tests::{ any, build_router, build_router_with_state, build_state_with_execution_runtime_override, json, start_server, strip_sse_keepalive_comments, AppState, Arc, Body, HeaderValue, Json, Mutex, @@ -17,6 +19,330 @@ use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot}; use base64::Engine as _; +const INTERNAL_REPORT_CAPABILITY_FIELD: &str = "_aether_internal_report_capability"; +const INTERNAL_REPORT_CLIENT_KEY: &str = "sk-internal-report-capability"; +const INTERNAL_REPORT_USER_ID: &str = "user-internal-report-capability"; +const INTERNAL_REPORT_API_KEY_ID: &str = "api-key-internal-report-capability"; + +fn internal_report_planner_state() -> AppState { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)), + unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID), + )])); + let candidate_repository = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + sample_models_candidate_row( + "provider-internal-report", + "openai", + "openai:chat", + "gpt-5", + 10, + ), + ])); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-internal-report", "openai", 10)], + vec![sample_endpoint( + "endpoint-provider-internal-report", + "provider-internal-report", + "openai:chat", + "https://api.openai.example", + )], + vec![sample_key( + "key-provider-internal-report", + "provider-internal-report", + "openai:chat", + "sk-upstream-openai", + )], + )); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + + AppState::new() + .expect("state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests( + auth_repository, + candidate_repository, + provider_catalog_repository, + request_candidate_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ) +} + +fn internal_video_create_planner_state(api_format: &str, model: &str) -> AppState { + let family = api_format + .split_once(':') + .map(|(family, _)| family) + .expect("video api format should contain a family"); + let provider_id = format!("provider-internal-{family}-video"); + let mut candidate = sample_models_candidate_row(&provider_id, family, api_format, model, 10); + candidate.endpoint_api_family = Some(family.to_string()); + candidate.endpoint_kind = Some("video".to_string()); + candidate.global_model_supports_streaming = Some(false); + candidate.model_supports_streaming = Some(false); + + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)), + unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID), + )])); + let candidate_repository = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + candidate, + ])); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider(&provider_id, family, 10)], + vec![sample_endpoint( + &format!("endpoint-{provider_id}"), + &provider_id, + api_format, + if family == "gemini" { + "https://generativelanguage.googleapis.com" + } else { + "https://api.openai.example" + }, + )], + vec![sample_key( + &format!("key-{provider_id}"), + &provider_id, + api_format, + "sk-upstream-video", + )], + )); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + + AppState::new() + .expect("state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests( + auth_repository, + candidate_repository, + provider_catalog_repository, + request_candidate_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ) +} + +async fn internal_video_followup_planner_state( + api_format: &str, + task_id: &str, + short_id: Option<&str>, + external_task_id: &str, + model: &str, +) -> AppState { + let family = api_format + .split_once(':') + .map(|(family, _)| family) + .expect("video api format should contain a family"); + let provider_id = format!("provider-internal-{family}-video-followup"); + let endpoint_id = format!("endpoint-{provider_id}"); + let key_id = format!("key-{provider_id}"); + let repository = Arc::new(InMemoryVideoTaskRepository::default()); + repository + .upsert(UpsertVideoTask { + id: task_id.to_string(), + short_id: short_id.map(ToOwned::to_owned), + request_id: format!("request-{task_id}"), + user_id: Some(INTERNAL_REPORT_USER_ID.to_string()), + api_key_id: Some(INTERNAL_REPORT_API_KEY_ID.to_string()), + username: Some("alice".to_string()), + api_key_name: Some("default".to_string()), + external_task_id: Some(external_task_id.to_string()), + provider_id: Some(provider_id.clone()), + endpoint_id: Some(endpoint_id.clone()), + key_id: Some(key_id.clone()), + client_api_format: Some(api_format.to_string()), + provider_api_format: Some(api_format.to_string()), + format_converted: false, + model: Some(model.to_string()), + prompt: Some("internal capability video".to_string()), + original_request_body: Some(json!({ + "model": model, + "prompt": "internal capability video", + })), + duration_seconds: Some(4), + resolution: Some("720p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("1280x720".to_string()), + status: if family == "openai" { + VideoTaskStatus::Completed + } else { + VideoTaskStatus::Submitted + }, + progress_percent: if family == "openai" { 100 } else { 0 }, + progress_message: None, + retry_count: 0, + poll_interval_seconds: 10, + next_poll_at_unix_secs: (family != "openai").then_some(1_700_000_010), + poll_count: 0, + max_poll_count: 360, + created_at_unix_ms: 1_700_000_000_000, + submitted_at_unix_secs: Some(1_700_000_000), + completed_at_unix_secs: (family == "openai").then_some(1_700_000_100), + updated_at_unix_secs: 1_700_000_000, + error_code: None, + error_message: None, + video_url: None, + request_metadata: None, + }) + .await + .expect("video task should seed"); + + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)), + unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID), + )])); + let provider_catalog_repository = crate::tests::video::video_provider_catalog_repository( + &provider_id, + family, + &endpoint_id, + api_format, + if family == "gemini" { + "https://generativelanguage.googleapis.com" + } else { + "https://api.openai.example/v1" + }, + &key_id, + "sk-upstream-video", + ); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let data_state = crate::data::GatewayDataState::with_video_task_provider_transport_and_request_candidate_repository_for_tests( + repository, + provider_catalog_repository, + request_candidate_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_auth_api_key_reader(auth_repository); + + AppState::new() + .expect("state should build") + .with_video_task_truth_source_mode(crate::VideoTaskTruthSourceMode::RustAuthoritative) + .with_data_state_for_tests(data_state) +} + +async fn issue_internal_gateway_report_capability( + client: &reqwest::Client, + gateway_url: &str, + endpoint: &str, + trace_id: &str, + method: &str, + path: &str, + request_headers: serde_json::Value, + body_json: serde_json::Value, +) -> (String, serde_json::Value) { + let response = client + .post(format!("{gateway_url}/api/internal/gateway/{endpoint}")) + .json(&json!({ + "trace_id": trace_id, + "method": method, + "path": path, + "headers": request_headers, + "body_json": body_json, + })) + .send() + .await + .expect("planner request should succeed"); + let status = response.status(); + let payload: serde_json::Value = response.json().await.expect("planner body should parse"); + assert_eq!( + status, + StatusCode::OK, + "planner should issue a report capability: {payload}" + ); + let report_kind = payload["report_kind"] + .as_str() + .expect("planner should return a report kind") + .to_string(); + let report_context = payload["report_context"].clone(); + assert!( + report_context[INTERNAL_REPORT_CAPABILITY_FIELD] + .as_str() + .is_some_and(|value| !value.is_empty()), + "planner should return an opaque report capability: {payload}" + ); + (report_kind, report_context) +} + +async fn issue_openai_chat_report_capability( + client: &reqwest::Client, + gateway_url: &str, + endpoint: &str, + trace_id: &str, + stream: bool, +) -> (String, serde_json::Value) { + issue_internal_gateway_report_capability( + client, + gateway_url, + endpoint, + trace_id, + "POST", + "/v1/chat/completions", + json!({ + "content-type": "application/json", + "x-api-key": INTERNAL_REPORT_CLIENT_KEY, + }), + json!({ + "model": "gpt-5", + "messages": [], + "stream": stream, + }), + ) + .await +} + +async fn post_internal_sync_report( + client: &reqwest::Client, + gateway_url: &str, + trace_id: &str, + report_kind: &str, + report_context: serde_json::Value, +) -> reqwest::Response { + client + .post(format!("{gateway_url}/api/internal/gateway/report-sync")) + .json(&json!({ + "trace_id": trace_id, + "report_kind": report_kind, + "report_context": report_context, + "status_code": 200, + "headers": { + "content-type": "application/json", + }, + "body_json": { + "id": "chatcmpl-internal-capability", + "usage": { + "input_tokens": 1, + "output_tokens": 2, + "total_tokens": 3, + } + } + })) + .send() + .await + .expect("report request should succeed") +} + +async fn assert_internal_report_capability_rejected(response: reqwest::Response) { + assert_eq!(response.status(), StatusCode::CONFLICT); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!( + payload, + json!({ + "detail": "internal gateway report context does not carry a valid planner capability", + }) + ); +} + +async fn assert_supplied_auth_context_rejected(response: reqwest::Response) { + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!( + payload, + json!({ + "detail": "supplied auth_context is not accepted; authenticate through request headers", + }) + ); +} + #[tokio::test] async fn gateway_handles_internal_gateway_resolve_without_proxying_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -281,21 +607,11 @@ async fn gateway_handles_internal_gateway_execute_sync_locally_impl() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .headers() - .get(EXECUTION_PATH_HEADER) - .and_then(|value| value.to_str().ok()), - Some(EXECUTION_PATH_EXECUTION_RUNTIME_SYNC) - ); - let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["id"], "chatcmpl-local-execute-sync"); - assert_eq!(payload["object"], "chat.completion"); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!( *execution_runtime_hits.lock().expect("mutex should lock"), - 1 + 0 ); gateway_handle.abort(); @@ -460,22 +776,11 @@ async fn gateway_handles_internal_gateway_execute_stream_locally() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .headers() - .get(EXECUTION_PATH_HEADER) - .and_then(|value| value.to_str().ok()), - Some(EXECUTION_PATH_EXECUTION_RUNTIME_STREAM) - ); - assert_eq!( - strip_sse_keepalive_comments(&response.text().await.expect("body should read")), - "data: one\n\ndata: [DONE]\n\n" - ); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!( *execution_runtime_hits.lock().expect("mutex should lock"), - 1 + 0 ); gateway_handle.abort(); @@ -613,18 +918,25 @@ async fn gateway_handles_internal_gateway_report_sync_locally() { ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let gateway = build_router_with_state(internal_report_planner_state()); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-report-sync", + false, + ) + .await; + assert_eq!(report_kind, "openai_chat_sync_success"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/report-sync")) .json(&json!({ "trace_id": "trace-internal-report-sync", - "report_kind": "openai_chat_sync_success", - "report_context": { - "user_id": "user-report-sync", - "api_key_id": "api-key-report-sync", - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -667,18 +979,25 @@ async fn gateway_handles_internal_gateway_report_stream_locally() { ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let gateway = build_router_with_state(internal_report_planner_state()); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-stream", + "trace-internal-report-stream", + true, + ) + .await; + assert_eq!(report_kind, "openai_chat_stream_success"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/report-stream")) .json(&json!({ "trace_id": "trace-internal-report-stream", - "report_kind": "openai_chat_stream_success", - "report_context": { - "user_id": "user-report-stream", - "api_key_id": "api-key-report-stream", - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "text/event-stream", @@ -698,6 +1017,181 @@ async fn gateway_handles_internal_gateway_report_stream_locally() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_rejects_internal_gateway_report_with_tampered_protected_context() { + let gateway = build_router_with_state(internal_report_planner_state()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-report-tampered-context", + false, + ) + .await; + + for (field, forged_value) in [ + ("user_id", json!("user-unrelated-victim")), + ("api_key_id", json!("api-key-unrelated-victim")), + ("provider_id", json!("provider-unrelated-victim")), + ("endpoint_id", json!("endpoint-unrelated-victim")), + ("key_id", json!("key-unrelated-victim")), + ("client_api_format", json!("gemini:video")), + ("task_id", json!("task-unrelated-victim")), + ("local_task_id", json!("local-task-unrelated-victim")), + ("local_short_id", json!("short-unrelated-victim")), + ("file_name", json!("files/unrelated-victim")), + ("file_key_id", json!("file-key-unrelated-victim")), + ] { + let mut tampered_context = report_context.clone(); + tampered_context + .as_object_mut() + .expect("planner report context should be an object") + .insert(field.to_string(), forged_value); + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-tampered-context", + &report_kind, + tampered_context, + ) + .await; + assert_internal_report_capability_rejected(response).await; + } + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_internal_gateway_report_without_a_known_capability() { + let gateway = build_router_with_state(internal_report_planner_state()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-report-missing-capability", + false, + ) + .await; + + let mut missing_capability = report_context.clone(); + missing_capability + .as_object_mut() + .expect("planner report context should be an object") + .remove(INTERNAL_REPORT_CAPABILITY_FIELD); + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-missing-capability", + &report_kind, + missing_capability, + ) + .await; + assert_internal_report_capability_rejected(response).await; + + let mut unknown_capability = report_context; + unknown_capability[INTERNAL_REPORT_CAPABILITY_FIELD] = + json!("00000000-0000-4000-8000-000000000000"); + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-missing-capability", + &report_kind, + unknown_capability, + ) + .await; + assert_internal_report_capability_rejected(response).await; + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_internal_gateway_report_for_wrong_trace_or_scope() { + let gateway = build_router_with_state(internal_report_planner_state()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-report-boundary", + false, + ) + .await; + + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-wrong-trace", + &report_kind, + report_context.clone(), + ) + .await; + assert_internal_report_capability_rejected(response).await; + + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-boundary", + "openai_image_sync_success", + report_context, + ) + .await; + assert_internal_report_capability_rejected(response).await; + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_allows_internal_gateway_report_observation_fields() { + let gateway = build_router_with_state(internal_report_planner_state()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, mut report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-report-observations", + false, + ) + .await; + let context = report_context + .as_object_mut() + .expect("planner report context should be an object"); + context.insert( + "provider_response_headers".to_string(), + json!({"x-request-id": "upstream-request-123"}), + ); + context.insert( + "client_response_headers".to_string(), + json!({"content-type": "application/json"}), + ); + context.insert("upstream_response".to_string(), json!({"status_code": 200})); + context.insert("error_flow".to_string(), json!({"attempted": false})); + + let response = post_internal_sync_report( + &client, + &gateway_url, + "trace-internal-report-observations", + &report_kind, + report_context, + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .json::() + .await + .expect("json body should parse"), + json!({"ok": true}) + ); + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_handles_internal_gateway_finalize_sync_locally() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -714,20 +1208,24 @@ async fn gateway_handles_internal_gateway_finalize_sync_locally() { ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let gateway = build_router_with_state(internal_report_planner_state()); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (_report_kind, report_context) = issue_openai_chat_report_capability( + &client, + &gateway_url, + "plan-sync", + "trace-internal-finalize-sync", + false, + ) + .await; - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/finalize-sync")) .json(&json!({ "trace_id": "trace-internal-finalize-sync", "report_kind": "openai_chat_sync_finalize", - "report_context": { - "user_id": "user-finalize-sync", - "api_key_id": "api-key-finalize-sync", - "client_api_format": "openai:chat", - "provider_api_format": "openai:chat", - }, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -781,24 +1279,37 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_locally() { ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let gateway = build_router_with_state(internal_video_create_planner_state( + "openai:video", + "sora-2", + )); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_internal_gateway_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-finalize-video", + "POST", + "/v1/videos", + json!({ + "authorization": format!("Bearer {INTERNAL_REPORT_CLIENT_KEY}"), + "content-type": "application/json", + }), + json!({ + "model": "sora-2", + "prompt": "make a trailer", + }), + ) + .await; + assert_eq!(report_kind, "openai_video_create_sync_finalize"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/finalize-sync")) .json(&json!({ "trace_id": "trace-internal-finalize-video", - "report_kind": "openai_video_create_sync_finalize", - "report_context": { - "user_id": "user-finalize-video", - "api_key_id": "api-key-finalize-video", - "model": "sora-2", - "local_task_id": "local-video-task-123", - "local_created_at": 1712345678u64, - "original_request_body": { - "prompt": "make a trailer" - } - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -821,18 +1332,15 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_locally() { "true" ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!( - payload, - json!({ - "id": "local-video-task-123", - "object": "video", - "status": "queued", - "progress": 0, - "created_at": 1712345678u64, - "model": "sora-2", - "prompt": "make a trailer", - }) - ); + assert!(payload["id"] + .as_str() + .is_some_and(|value| !value.is_empty() && value != "vid-ext-123")); + assert_eq!(payload["object"], "video"); + assert_eq!(payload["status"], "queued"); + assert_eq!(payload["progress"], 0); + assert!(payload["created_at"].as_u64().is_some()); + assert_eq!(payload["model"], "sora-2"); + assert_eq!(payload["prompt"], "make a trailer"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -855,20 +1363,34 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_locally() { ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let gateway = + build_router_with_state(internal_video_create_planner_state("gemini:video", "veo-3")); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_internal_gateway_report_capability( + &client, + &gateway_url, + "plan-sync", + "trace-internal-finalize-gemini-video", + "POST", + "/v1beta/models/veo-3:predictLongRunning", + json!({ + "content-type": "application/json", + "x-goog-api-key": INTERNAL_REPORT_CLIENT_KEY, + }), + json!({ + "prompt": "make a gemini trailer", + }), + ) + .await; + assert_eq!(report_kind, "gemini_video_create_sync_finalize"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/finalize-sync")) .json(&json!({ "trace_id": "trace-internal-finalize-gemini-video", - "report_kind": "gemini_video_create_sync_finalize", - "report_context": { - "user_id": "user-finalize-gemini-video", - "api_key_id": "api-key-finalize-gemini-video", - "model": "veo-3", - "local_short_id": "gemini-short-123" - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -892,14 +1414,11 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_locally() { "true" ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!( - payload, - json!({ - "name": "models/veo-3/operations/gemini-short-123", - "done": false, - "metadata": {}, - }) - ); + assert!(payload["name"] + .as_str() + .is_some_and(|value| value.starts_with("models/veo-3/operations/"))); + assert_eq!(payload["done"], false); + assert_eq!(payload["metadata"], json!({})); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -922,17 +1441,39 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_delete_loca ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let state = internal_video_followup_planner_state( + "openai:video", + "video-delete-123", + None, + "ext-video-delete-123", + "sora-2", + ) + .await; + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_internal_gateway_report_capability( + &client, + &gateway_url, + "decision-sync", + "trace-internal-finalize-video-delete", + "DELETE", + "/v1/videos/video-delete-123", + json!({ + "authorization": format!("Bearer {INTERNAL_REPORT_CLIENT_KEY}"), + "content-type": "application/json", + }), + json!({}), + ) + .await; + assert_eq!(report_kind, "openai_video_delete_sync_finalize"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/finalize-sync")) .json(&json!({ "trace_id": "trace-internal-finalize-video-delete", - "report_kind": "openai_video_delete_sync_finalize", - "report_context": { - "task_id": "video-delete-123" - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -982,18 +1523,39 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_cancel_loca ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let gateway = build_router().expect("gateway should build"); + let state = internal_video_followup_planner_state( + "gemini:video", + "gemini-cancel-task-record", + Some("gemini-cancel-123"), + "operations/ext-gemini-cancel-123", + "veo-3", + ) + .await; + let gateway = build_router_with_state(state); let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let (report_kind, report_context) = issue_internal_gateway_report_capability( + &client, + &gateway_url, + "plan-sync", + "trace-internal-finalize-gemini-video-cancel", + "POST", + "/v1beta/models/veo-3/operations/gemini-cancel-123:cancel", + json!({ + "content-type": "application/json", + "x-goog-api-key": INTERNAL_REPORT_CLIENT_KEY, + }), + json!({}), + ) + .await; + assert_eq!(report_kind, "gemini_video_cancel_sync_finalize"); - let response = reqwest::Client::new() + let response = client .post(format!("{gateway_url}/api/internal/gateway/finalize-sync")) .json(&json!({ "trace_id": "trace-internal-finalize-gemini-video-cancel", - "report_kind": "gemini_video_cancel_sync_finalize", - "report_context": { - "task_id": "gemini-cancel-123", - "model": "veo-3" - }, + "report_kind": report_kind, + "report_context": report_context, "status_code": 200, "headers": { "content-type": "application/json", @@ -1148,17 +1710,7 @@ async fn gateway_handles_internal_gateway_decision_sync_locally_with_supplied_au .await .expect("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["action"], "execution_runtime_sync_decision"); - assert_eq!(payload["decision_kind"], "openai_chat_sync"); - assert_eq!(payload["provider_id"], "provider-1"); - assert_eq!(payload["endpoint_id"], "endpoint-provider-1"); - assert_eq!(payload["key_id"], "key-provider-1"); - assert_eq!(payload["provider_api_format"], "openai:chat"); - assert_eq!(payload["client_api_format"], "openai:chat"); - assert_eq!(payload["model_name"], "gpt-5"); - assert_eq!(payload["auth_context"], serde_json::Value::Null); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1267,10 +1819,7 @@ async fn gateway_internal_decision_sync_revalidates_supplied_auth_context_wallet .await .expect("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["action"], "fallback_plan"); - assert_eq!(payload["auth_context"], serde_json::Value::Null); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1301,7 +1850,10 @@ async fn gateway_returns_internal_gateway_decision_sync_fallback_with_resolved_a let gateway = build_router_with_state( AppState::new() .expect("gateway should build") - .with_auth_api_key_data_reader_for_tests(auth_repository), + .with_data_state_for_tests( + GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository) + .with_system_default_routing_group_for_tests(), + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -1418,14 +1970,7 @@ async fn gateway_handles_internal_gateway_decision_stream_locally_with_supplied_ .await .expect("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["action"], "execution_runtime_stream_decision"); - assert_eq!(payload["decision_kind"], "openai_chat_stream"); - assert_eq!(payload["provider_id"], "provider-stream-1"); - assert_eq!(payload["endpoint_id"], "endpoint-provider-stream-1"); - assert_eq!(payload["key_id"], "key-provider-stream-1"); - assert_eq!(payload["auth_context"], serde_json::Value::Null); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1517,17 +2062,7 @@ async fn gateway_handles_internal_gateway_plan_sync_locally_with_supplied_auth_c .await .expect("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["action"], "execution_runtime_sync"); - assert_eq!(payload["plan_kind"], "openai_chat_sync"); - assert_eq!(payload["plan"]["provider_id"], "provider-plan-sync-1"); - assert_eq!( - payload["plan"]["endpoint_id"], - "endpoint-provider-plan-sync-1" - ); - assert_eq!(payload["plan"]["key_id"], "key-provider-plan-sync-1"); - assert_eq!(payload["auth_context"], serde_json::Value::Null); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1620,18 +2155,7 @@ async fn gateway_handles_internal_gateway_plan_stream_locally_with_supplied_auth .await .expect("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["action"], "execution_runtime_stream"); - assert_eq!(payload["plan_kind"], "openai_chat_stream"); - assert_eq!(payload["plan"]["provider_id"], "provider-plan-stream-1"); - assert_eq!( - payload["plan"]["endpoint_id"], - "endpoint-provider-plan-stream-1" - ); - assert_eq!(payload["plan"]["key_id"], "key-provider-plan-stream-1"); - assert_eq!(payload["plan"]["stream"], true); - assert_eq!(payload["auth_context"], serde_json::Value::Null); + assert_supplied_auth_context_rejected(response).await; assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); diff --git a/apps/aether-gateway/src/tests/frontdoor/oauth.rs b/apps/aether-gateway/src/tests/frontdoor/oauth.rs index 25a6a9846..87016273a 100644 --- a/apps/aether-gateway/src/tests/frontdoor/oauth.rs +++ b/apps/aether-gateway/src/tests/frontdoor/oauth.rs @@ -5,6 +5,13 @@ use crate::tests::{ use aether_data::repository::oauth_providers::{ InMemoryOAuthProviderRepository, StoredOAuthProviderConfig, }; +use aether_data::repository::users::{ + InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserSessionRecord, +}; +use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; +use base64::Engine as _; +use hmac::Mac as _; +use sha2::{Digest, Sha256}; fn sample_identity_oauth_provider(provider_type: &str) -> StoredOAuthProviderConfig { StoredOAuthProviderConfig::new( @@ -28,6 +35,63 @@ fn sample_identity_oauth_provider(provider_type: &str) -> StoredOAuthProviderCon ) } +fn oauth_login_cookie_name(state_nonce: &str) -> String { + format!( + "__Host-aether_oauth_login_{:x}", + Sha256::digest(state_nonce.as_bytes()) + ) +} + +fn oauth_state_from_authorize_location(location: &str) -> String { + url::Url::parse(location) + .expect("authorize location should be a URL") + .query_pairs() + .find_map(|(key, value)| (key == "state").then(|| value.into_owned())) + .expect("authorize location should include state") +} + +fn cookie_pair_from_set_cookie(set_cookie: &str) -> String { + set_cookie + .split(';') + .next() + .expect("Set-Cookie should include a cookie pair") + .to_string() +} + +fn build_oauth_test_access_token( + user: &StoredUserAuthRecord, + session_id: &str, + expires_at: chrono::DateTime, +) -> String { + let header = serde_json::json!({ "alg": "HS256", "typ": "JWT" }); + let payload = serde_json::json!({ + "exp": expires_at.timestamp(), + "type": "access", + "user_id": user.id, + "role": user.role, + "created_at": user.created_at.map(|value| value.to_rfc3339()), + "session_id": session_id, + }); + let header_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(serde_json::to_vec(&header).expect("JWT header should serialize")); + let payload_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(serde_json::to_vec(&payload).expect("JWT payload should serialize")); + let signing_input = format!("{header_segment}.{payload_segment}"); + let secret = std::env::var("JWT_SECRET_KEY") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string()); + let mut mac = hmac::Hmac::::new_from_slice(secret.as_bytes()) + .expect("test JWT secret should be valid"); + mac.update(signing_input.as_bytes()); + let signature = mac.finalize().into_bytes(); + format!( + "{signing_input}.{}", + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(signature) + ) +} + #[tokio::test] async fn gateway_serves_oauth_public_providers_locally_without_hitting_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -243,6 +307,294 @@ async fn gateway_serves_configured_oauth_provider_when_oauth_module_enabled() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_rejects_oauth_callback_without_browser_binding_without_consuming_state() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + }), + ); + + let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![ + sample_identity_oauth_provider("linuxdo"), + ])); + let data_state = + crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository) + .with_system_config_values_for_tests(vec![( + "module.oauth.enabled".to_string(), + serde_json::json!(true), + )]); + let runtime_state = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state) + .with_runtime_state(runtime_state); + + let binding_hash = format!("{:x}", Sha256::digest(b"correct-cookie")); + let missing_cookie_state = crate::oauth::StoredIdentityOAuthState::login( + "linuxdo", + "device-csrf-test", + Some("pkce-verifier".to_string()), + Some(binding_hash.clone()), + ); + let wrong_cookie_state = crate::oauth::StoredIdentityOAuthState::login( + "linuxdo", + "device-csrf-test", + Some("pkce-verifier".to_string()), + Some(binding_hash), + ); + crate::oauth::save_identity_oauth_state(&state, &missing_cookie_state) + .await + .expect("missing-cookie state should save"); + crate::oauth::save_identity_oauth_state(&state, &wrong_cookie_state) + .await + .expect("wrong-cookie state should save"); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("client should build"); + + for (state_nonce, cookie_header) in [ + (missing_cookie_state.nonce.as_str(), None), + ( + wrong_cookie_state.nonce.as_str(), + Some(format!( + "{}=wrong-cookie", + oauth_login_cookie_name(&wrong_cookie_state.nonce) + )), + ), + ] { + let mut request = client.get(format!( + "{gateway_url}/api/oauth/linuxdo/callback?code=provider-code&state={state_nonce}" + )); + if let Some(cookie_header) = cookie_header.as_deref() { + request = request.header(http::header::COOKIE, cookie_header); + } + let response = request + .send() + .await + .expect("callback request should succeed"); + assert_eq!(response.status(), StatusCode::FOUND); + let location = response + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("invalid callback should redirect"); + assert!( + location.contains("invalid_state"), + "unexpected location: {location}" + ); + let expected_cookie_name = oauth_login_cookie_name(state_nonce); + let clear_cookies = response + .headers() + .get_all(http::header::SET_COOKIE) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect::>(); + assert_eq!(clear_cookies.len(), 1); + assert!(clear_cookies[0].starts_with(&format!("{expected_cookie_name}=;"))); + assert!(clear_cookies[0].contains("Max-Age=0")); + } + + for state_nonce in [&missing_cookie_state.nonce, &wrong_cookie_state.nonce] { + assert!(state + .runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key(state_nonce)) + .await + .expect("OAuth state lookup should succeed") + .is_some()); + + let cookie_header = format!("{}=correct-cookie", oauth_login_cookie_name(state_nonce)); + let response = client + .get(format!( + "{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={state_nonce}" + )) + .header(http::header::COOKIE, &cookie_header) + .send() + .await + .expect("bound callback request should succeed"); + assert_eq!(response.status(), StatusCode::FOUND); + let location = response + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("denied callback should redirect"); + assert!(location.contains("authorization_denied")); + assert!(state + .runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key(state_nonce)) + .await + .expect("OAuth state lookup should succeed") + .is_none()); + + let replay = client + .get(format!( + "{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={state_nonce}" + )) + .header(http::header::COOKIE, cookie_header) + .send() + .await + .expect("replayed callback request should succeed"); + let replay_location = replay + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("replayed callback should redirect"); + assert!(replay_location.contains("invalid_state")); + } + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_keeps_parallel_oauth_login_cookies_independent_and_clears_only_consumed_state() { + let upstream_hits = Arc::new(Mutex::new(0usize)); + let upstream_hits_clone = Arc::clone(&upstream_hits); + let upstream = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let upstream_hits_inner = Arc::clone(&upstream_hits_clone); + async move { + *upstream_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::OK, Body::from("unexpected upstream hit")) + } + }), + ); + let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![ + sample_identity_oauth_provider("linuxdo"), + ])); + let data_state = + crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository) + .with_system_config_values_for_tests(vec![( + "module.oauth.enabled".to_string(), + serde_json::json!(true), + )]); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state) + .with_runtime_state(Arc::new(RuntimeState::memory( + MemoryRuntimeStateConfig::default(), + ))); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("client should build"); + let first_authorize = client + .get(format!("{gateway_url}/api/oauth/linuxdo/authorize")) + .header("x-client-device-id", "parallel-device"); + let second_authorize = client + .get(format!("{gateway_url}/api/oauth/linuxdo/authorize")) + .header("x-client-device-id", "parallel-device"); + let (first_response, second_response) = + tokio::join!(first_authorize.send(), second_authorize.send()); + let first_response = first_response.expect("first authorize request should succeed"); + let second_response = second_response.expect("second authorize request should succeed"); + assert_eq!(first_response.status(), StatusCode::FOUND); + assert_eq!(second_response.status(), StatusCode::FOUND); + + let first_location = first_response + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("first authorize location should exist"); + let second_location = second_response + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("second authorize location should exist"); + let first_state = oauth_state_from_authorize_location(first_location); + let second_state = oauth_state_from_authorize_location(second_location); + assert_ne!(first_state, second_state); + + let first_cookie_name = oauth_login_cookie_name(&first_state); + let second_cookie_name = oauth_login_cookie_name(&second_state); + assert_ne!(first_cookie_name, second_cookie_name); + let first_cookie = first_response + .headers() + .get(http::header::SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .map(cookie_pair_from_set_cookie) + .expect("first authorize response should set a login cookie"); + let second_cookie = second_response + .headers() + .get(http::header::SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .map(cookie_pair_from_set_cookie) + .expect("second authorize response should set a login cookie"); + assert!(first_cookie.starts_with(&format!("{first_cookie_name}="))); + assert!(second_cookie.starts_with(&format!("{second_cookie_name}="))); + + let combined_cookie_header = format!("{first_cookie}; {second_cookie}"); + let first_callback = client + .get(format!( + "{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={first_state}" + )) + .header(http::header::COOKIE, combined_cookie_header) + .send() + .await + .expect("first callback request should succeed"); + assert_eq!(first_callback.status(), StatusCode::FOUND); + let first_callback_location = first_callback + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("first callback location should exist"); + assert!(first_callback_location.contains("authorization_denied")); + let first_clears = first_callback + .headers() + .get_all(http::header::SET_COOKIE) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect::>(); + assert_eq!(first_clears.len(), 1); + assert!(first_clears[0].starts_with(&format!("{first_cookie_name}=;"))); + assert!(!first_clears[0].starts_with(&format!("{second_cookie_name}=;"))); + + let second_callback = client + .get(format!( + "{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={second_state}" + )) + .header(http::header::COOKIE, second_cookie) + .send() + .await + .expect("second callback request should succeed"); + assert_eq!(second_callback.status(), StatusCode::FOUND); + let second_callback_location = second_callback + .headers() + .get(http::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .expect("second callback location should exist"); + assert!(second_callback_location.contains("authorization_denied")); + let second_clears = second_callback + .headers() + .get_all(http::header::SET_COOKIE) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect::>(); + assert_eq!(second_clears.len(), 1); + assert!(second_clears[0].starts_with(&format!("{second_cookie_name}=;"))); + assert!(!second_clears[0].starts_with(&format!("{first_cookie_name}=;"))); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_requires_auth_for_oauth_user_bindable_providers_without_hitting_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -277,6 +629,102 @@ async fn gateway_requires_auth_for_oauth_user_bindable_providers_without_hitting upstream_handle.abort(); } +#[tokio::test] +async fn gateway_oauth_account_lists_are_never_cacheable() { + let provider_repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![ + sample_identity_oauth_provider("linuxdo"), + ])); + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "oauth-list-user".to_string(), + Some("oauth-list@example.com".to_string()), + true, + "oauth-list-user".to_string(), + Some("unused-password-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + Some(now), + ) + .expect("test user should build"); + let session_id = "oauth-list-session"; + let device_id = "oauth-list-device"; + let session = StoredUserSessionRecord::new( + session_id.to_string(), + user.id.clone(), + device_id.to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("unused-refresh-token"), + None, + None, + Some(now), + Some(now + chrono::Duration::days(1)), + None, + None, + Some("127.0.0.1".to_string()), + Some("oauth-list-test".to_string()), + Some(now), + Some(now), + ) + .expect("test session should build"); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user.clone()])); + let data_state = crate::data::GatewayDataState::with_oauth_provider_repository_for_tests( + provider_repository, + ) + .with_system_config_values_for_tests(vec![( + "module.oauth.enabled".to_string(), + serde_json::json!(true), + )]) + .with_user_reader(user_repository); + let access_token = + build_oauth_test_access_token(&user, session_id, now + chrono::Duration::hours(1)); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state) + .with_auth_users_for_tests([user]) + .with_auth_session_for_tests(session); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for path in [ + "/api/user/oauth/bindable-providers", + "/api/user/oauth/links", + ] { + let response = client + .get(format!("{gateway_url}{path}")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", device_id) + .send() + .await + .expect("OAuth account list request should succeed"); + assert_eq!(response.status(), StatusCode::OK, "{path}"); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store"), + "{path}" + ); + assert_eq!( + response + .headers() + .get(http::header::PRAGMA) + .and_then(|value| value.to_str().ok()), + Some("no-cache"), + "{path}" + ); + } + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_requires_auth_for_oauth_user_bind_token_without_hitting_upstream() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -310,3 +758,123 @@ async fn gateway_requires_auth_for_oauth_user_bind_token_without_hitting_upstrea gateway_handle.abort(); upstream_handle.abort(); } + +#[tokio::test] +async fn gateway_oauth_bind_token_returns_authorize_url_and_stores_bound_browser_state() { + let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![ + sample_identity_oauth_provider("linuxdo"), + ])); + let data_state = + crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository) + .with_system_config_values_for_tests(vec![( + "module.oauth.enabled".to_string(), + serde_json::json!(true), + )]); + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "oauth-bind-user".to_string(), + Some("oauth-bind@example.com".to_string()), + true, + "oauth-bind-user".to_string(), + Some("unused-password-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + Some(now), + ) + .expect("test user should build"); + let session_id = "oauth-bind-session"; + let device_id = "oauth-bind-device"; + let session = StoredUserSessionRecord::new( + session_id.to_string(), + user.id.clone(), + device_id.to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("unused-refresh-token"), + None, + None, + Some(now), + Some(now + chrono::Duration::days(1)), + None, + None, + Some("127.0.0.1".to_string()), + Some("oauth-bind-test".to_string()), + Some(now), + Some(now), + ) + .expect("test session should build"); + let access_token = + build_oauth_test_access_token(&user, session_id, now + chrono::Duration::hours(1)); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state) + .with_auth_users_for_tests([user.clone()]) + .with_auth_session_for_tests(session); + let gateway = build_router_with_state(state.clone()); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/user/oauth/linuxdo/bind-token")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", device_id) + .header("user-agent", "AetherOAuthBindTest/1.0") + .send() + .await + .expect("bind-token request should succeed"); + + assert_eq!(response.status(), StatusCode::OK); + let set_cookie = response + .headers() + .get(http::header::SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .expect("bind-token response should set the browser binding cookie") + .to_string(); + assert!(set_cookie.contains("Path=/")); + assert!(set_cookie.contains("HttpOnly")); + assert!(set_cookie.contains("SameSite=Lax")); + let cookie_pair = cookie_pair_from_set_cookie(&set_cookie); + let (cookie_name, browser_binding) = cookie_pair + .split_once('=') + .expect("browser binding cookie should contain a value"); + + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert!(payload.get("bind_token").is_none()); + let authorize_url = payload["authorize_url"] + .as_str() + .expect("bind-token response should include authorize_url"); + let state_nonce = oauth_state_from_authorize_location(authorize_url); + assert_eq!(cookie_name, oauth_login_cookie_name(&state_nonce)); + + let raw_state = state + .runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key( + &state_nonce, + )) + .await + .expect("OAuth state lookup should succeed") + .expect("OAuth bind state should be stored"); + assert!(crate::handlers::shared::runtime_secret_payload_is_sealed( + &raw_state + )); + assert!(!raw_state.contains("pkce_verifier")); + let stored = crate::oauth::load_identity_oauth_state(&state, &state_nonce) + .await + .expect("OAuth state lookup should succeed") + .expect("OAuth state should decrypt"); + assert_eq!(stored.mode, crate::oauth::IdentityOAuthStateMode::Bind); + assert_eq!(stored.provider_type, "linuxdo"); + assert_eq!(stored.client_device_id, device_id); + assert_eq!(stored.bind_user_id.as_deref(), Some(user.id.as_str())); + assert_eq!(stored.bind_session_id.as_deref(), Some(session_id)); + let expected_binding_hash = format!("{:x}", Sha256::digest(browser_binding.as_bytes())); + assert_eq!( + stored.browser_binding_hash.as_deref(), + Some(expected_binding_hash.as_str()) + ); + + gateway_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/frontdoor/ops.rs b/apps/aether-gateway/src/tests/frontdoor/ops.rs index 81fd14e2b..72613f263 100644 --- a/apps/aether-gateway/src/tests/frontdoor/ops.rs +++ b/apps/aether-gateway/src/tests/frontdoor/ops.rs @@ -184,7 +184,7 @@ async fn gateway_exposes_frontdoor_manifest_without_proxying_upstream() { .any(|value| value == "/v1internal:streamGenerateContent")); assert_eq!( payload["rust_frontdoor"]["internal_gateway"]["status"], - "rust_native_control_plane" + "test_loopback_compatibility" ); assert_eq!( payload["rust_frontdoor"]["internal_gateway"]["path_prefixes"][0], diff --git a/apps/aether-gateway/src/tests/frontdoor/public_support.rs b/apps/aether-gateway/src/tests/frontdoor/public_support.rs index 6eb1fd7a1..465ec9451 100644 --- a/apps/aether-gateway/src/tests/frontdoor/public_support.rs +++ b/apps/aether-gateway/src/tests/frontdoor/public_support.rs @@ -26,8 +26,8 @@ use aether_data::repository::auth_modules::{ }; use aether_data::repository::billing::InMemoryBillingReadRepository; use aether_data::repository::management_tokens::{ - InMemoryManagementTokenRepository, StoredManagementToken, StoredManagementTokenUserSummary, - StoredManagementTokenWithUser, + InMemoryManagementTokenRepository, ManagementTokenReadRepository, StoredManagementToken, + StoredManagementTokenUserSummary, StoredManagementTokenWithUser, }; use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::users::{ @@ -47,8 +47,13 @@ use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageRep use axum::response::IntoResponse; use chrono::{TimeZone, Utc}; +const TEST_EMAIL_VERIFICATION_TOKEN: &str = + "test-email-verification-token-00000000000000000000000000000000"; + #[path = "public_support/dashboard.rs"] mod dashboard; +#[path = "public_support/vscodex.rs"] +mod vscodex; #[tokio::test] async fn gateway_handles_public_announcements_list_without_proxying_upstream() { @@ -133,7 +138,7 @@ async fn gateway_handles_public_announcements_list_without_proxying_upstream() { let response = reqwest::Client::new() .get(format!( - "{gateway_url}/api/announcements?active_only=true&limit=50&offset=0" + "{gateway_url}/api/announcements?active_only=false&limit=50&offset=0" )) .send() .await @@ -252,6 +257,7 @@ async fn gateway_handles_public_announcement_detail_without_proxying_upstream() }), ); + let now = Utc::now().timestamp().max(0) as u64; let announcement_repository = Arc::new(InMemoryAnnouncementReadRepository::seed(vec![ StoredAnnouncement::new( "announcement-1".to_string(), @@ -264,8 +270,8 @@ async fn gateway_handles_public_announcement_detail_without_proxying_upstream() false, Some("admin-1".to_string()), Some("admin".to_string()), - Some(1_711_000_000), - Some(1_711_003_600), + Some(now.saturating_sub(60) as i64), + Some(now.saturating_add(3600) as i64), 1_711_000_000, 1_711_000_100, ) @@ -304,6 +310,62 @@ async fn gateway_handles_public_announcement_detail_without_proxying_upstream() upstream_handle.abort(); } +#[tokio::test] +async fn gateway_hides_non_public_announcement_details() { + let now = Utc::now().timestamp().max(0) as u64; + let announcements = [ + ("inactive", false, None, None), + ("future", true, Some(now.saturating_add(3600) as i64), None), + ("expired", true, None, Some(now.saturating_sub(1) as i64)), + ] + .into_iter() + .map(|(id, is_active, starts_at, ends_at)| { + StoredAnnouncement::new( + format!("announcement-{id}"), + format!("{id} title"), + format!("{id} private content"), + "info".to_string(), + 1, + is_active, + false, + false, + Some("admin-private".to_string()), + Some("admin".to_string()), + starts_at, + ends_at, + now as i64, + now as i64, + ) + .expect("announcement should build") + }) + .collect::>(); + let announcement_repository = Arc::new(InMemoryAnnouncementReadRepository::seed(announcements)); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::with_announcement_reader_for_tests( + announcement_repository, + ), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let client = reqwest::Client::new(); + for id in ["inactive", "future", "expired"] { + let response = client + .get(format!("{gateway_url}/api/announcements/announcement-{id}")) + .send() + .await + .expect("request should succeed"); + assert_eq!(response.status(), StatusCode::NOT_FOUND, "{id}"); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "Announcement not found", "{id}"); + } + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_creates_announcement_locally_with_trusted_admin_principal() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -721,10 +783,20 @@ async fn gateway_handles_public_catalog_providers_without_proxying_upstream() { }), ); + let mut inactive_provider = sample_provider("provider-draft", "draft", 30); + inactive_provider.is_active = false; + let mut inactive_endpoint = sample_endpoint( + "endpoint-openai-disabled", + "provider-openai", + "openai:responses", + "https://internal-openai.example", + ); + inactive_endpoint.is_active = false; let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![ sample_provider("provider-openai", "openai", 10), sample_provider("provider-claude", "claude", 20), + inactive_provider, ], vec![ sample_endpoint( @@ -745,6 +817,7 @@ async fn gateway_handles_public_catalog_providers_without_proxying_upstream() { "claude:messages", "https://api.anthropic.example", ), + inactive_endpoint, ], vec![], )); @@ -762,7 +835,9 @@ async fn gateway_handles_public_catalog_providers_without_proxying_upstream() { let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() - .get(format!("{gateway_url}/api/public/providers?limit=10")) + .get(format!( + "{gateway_url}/api/public/providers?is_active=false&limit=10" + )) .send() .await .expect("request should succeed"); @@ -783,6 +858,9 @@ async fn gateway_handles_public_catalog_providers_without_proxying_upstream() { assert_eq!(providers[1]["id"], "provider-claude"); assert_eq!(providers[1]["endpoints_count"], 2); assert_eq!(providers[1]["active_endpoints_count"], 2); + assert!(providers + .iter() + .all(|provider| provider["id"] != "provider-draft")); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1123,9 +1201,36 @@ async fn gateway_handles_public_global_models_without_proxying_upstream() { let mut gpt_model = sample_public_global_model("gm-3", "gpt-5", "GPT 5", true); gpt_model.config = Some(json!({ "description": "Public description", + "streaming": true, + "api_formats": ["openai:responses", {"secret": "capability-leak"}], + "billing": { + "video": { + "price_per_second_by_resolution": { + "720p": 0.12, + "internal": "video-billing-secret" + }, + "private_key": "video-secret" + } + }, + "client_secret": "config-secret", "model_mappings": ["gpt-5-upstream"], "provider_model_mappings": [{"name": "provider-gpt-5"}], })); + gpt_model.supported_capabilities = Some(json!(["vision", {"secret": "capability-secret"}])); + gpt_model.default_tiered_pricing = Some(json!({ + "tiers": [{ + "up_to": null, + "input_price_per_1m": 3.0, + "output_price_per_1m": 15.0, + "internal_note": "tier-secret", + "cache_ttl_pricing": [{ + "ttl_minutes": 60, + "cache_creation_price_per_1m": 4.0, + "secret": "ttl-secret" + }] + }], + "internal_pricing": {"secret": "pricing-secret"} + })); let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(vec![ sample_public_global_model("gm-1", "claude-sonnet-4-5", "Claude Sonnet 4.5", true), sample_public_global_model("gm-2", "disabled-model", "Disabled Model", false), @@ -1145,7 +1250,9 @@ async fn gateway_handles_public_global_models_without_proxying_upstream() { let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() - .get(format!("{gateway_url}/api/public/global-models?search=gpt")) + .get(format!( + "{gateway_url}/api/public/global-models?is_active=false&search=gpt" + )) .send() .await .expect("request should succeed"); @@ -1166,6 +1273,35 @@ async fn gateway_handles_public_global_models_without_proxying_upstream() { assert!(payload["models"][0]["config"] .get("provider_model_mappings") .is_none()); + assert!(payload["models"][0]["config"] + .get("client_secret") + .is_none()); + assert_eq!( + payload["models"][0]["config"]["api_formats"], + json!(["openai:responses"]) + ); + assert_eq!( + payload["models"][0]["config"]["billing"]["video"]["price_per_second_by_resolution"], + json!({"720p": 0.12}) + ); + assert!(payload["models"][0]["config"]["billing"]["video"] + .get("private_key") + .is_none()); + assert_eq!( + payload["models"][0]["supported_capabilities"], + json!(["vision"]) + ); + assert!(payload["models"][0]["default_tiered_pricing"] + .get("internal_pricing") + .is_none()); + assert!(payload["models"][0]["default_tiered_pricing"]["tiers"][0] + .get("internal_note") + .is_none()); + assert!( + payload["models"][0]["default_tiered_pricing"]["tiers"][0]["cache_ttl_pricing"][0] + .get("secret") + .is_none() + ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1779,8 +1915,13 @@ async fn gateway_handles_public_provider_detail_without_proxying_upstream() { }), ); + let mut inactive_provider = sample_provider("provider-disabled", "disabled", 20); + inactive_provider.is_active = false; let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider("provider-1", "openai", 10)], + vec![ + sample_provider("provider-1", "openai", 10), + inactive_provider, + ], vec![], vec![], )); @@ -1814,6 +1955,18 @@ async fn gateway_handles_public_provider_detail_without_proxying_upstream() { .await .expect("request should succeed"); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "Provider not found"); + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/v1/providers/provider-disabled?include_endpoints=true" + )) + .send() + .await + .expect("request should succeed"); + assert_eq!(response.status(), StatusCode::NOT_FOUND); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["detail"], "Provider not found"); @@ -1838,14 +1991,24 @@ async fn gateway_handles_public_providers_with_endpoints_without_proxying_upstre }), ); + let mut inactive_endpoint = sample_endpoint( + "endpoint-disabled", + "provider-1", + "openai:responses", + "https://internal-openai.example", + ); + inactive_endpoint.is_active = false; let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-1", "openai", 10)], - vec![sample_endpoint( - "endpoint-1", - "provider-1", - "openai:chat", - "https://api.openai.example", - )], + vec![ + sample_endpoint( + "endpoint-1", + "provider-1", + "openai:chat", + "https://api.openai.example", + ), + inactive_endpoint, + ], vec![], )); let (upstream_url, upstream_handle) = start_server(upstream).await; @@ -1876,6 +2039,12 @@ async fn gateway_handles_public_providers_with_endpoints_without_proxying_upstre payload["providers"][0]["endpoints"][0]["api_format"], "openai:chat" ); + assert_eq!( + payload["providers"][0]["endpoints"] + .as_array() + .map(Vec::len), + Some(1) + ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -1920,7 +2089,157 @@ async fn gateway_handles_test_connection_alias_without_proxying_upstream() { } #[tokio::test] -async fn gateway_handles_public_test_connection_without_hitting_fallback_probe() { +async fn gateway_rejects_unauthenticated_test_connection_before_using_provider_credentials() { + let provider_hits = Arc::new(Mutex::new(0usize)); + let provider_hits_clone = Arc::clone(&provider_hits); + let provider = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let provider_hits_inner = Arc::clone(&provider_hits_clone); + async move { + *provider_hits_inner.lock().expect("mutex should lock") += 1; + Json(json!({"id": "must-not-run"})).into_response() + } + }), + ); + let (provider_url, provider_handle) = start_server(provider).await; + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-auth-required", "openai", 10)], + vec![sample_endpoint( + "endpoint-auth-required", + "provider-auth-required", + "openai:chat", + &provider_url, + )], + vec![sample_key( + "key-auth-required", + "provider-auth-required", + "openai:chat", + "stored-provider-secret", + )], + )); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_provider_transport_reader_for_tests( + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + )), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/v1/test-connection?provider=provider-auth-required&model=gpt-5&api_format=openai:chat" + )) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "缺少用户凭证"); + assert_eq!(*provider_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + provider_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_non_admin_test_connection_before_using_provider_credentials() { + let provider_hits = Arc::new(Mutex::new(0usize)); + let provider_hits_clone = Arc::clone(&provider_hits); + let provider = Router::new().route( + "/{*path}", + any(move |_request: Request| { + let provider_hits_inner = Arc::clone(&provider_hits_clone); + async move { + *provider_hits_inner.lock().expect("mutex should lock") += 1; + Json(json!({"id": "must-not-run"})).into_response() + } + }), + ); + let (provider_url, provider_handle) = start_server(provider).await; + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider("provider-admin-only", "openai", 10)], + vec![sample_endpoint( + "endpoint-admin-only", + "provider-admin-only", + "openai:chat", + &provider_url, + )], + vec![sample_key( + "key-admin-only", + "provider-admin-only", + "openai:chat", + "stored-provider-secret", + )], + )); + + let now = Utc::now(); + let user = sample_auth_user(now); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ( + "session_id".to_string(), + json!("session-test-connection-non-admin"), + ), + ]), + now + chrono::Duration::hours(1), + ); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user])); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_transport_reader_for_tests( + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_user_reader(user_repository), + ) + .with_auth_session_for_tests(sample_auth_session( + "user-auth-1", + "session-test-connection-non-admin", + "device-test-connection-non-admin", + "refresh-token-placeholder", + now, + )), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/v1/test-connection?provider=provider-admin-only&model=gpt-5&api_format=openai:chat" + )) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-test-connection-non-admin", + ) + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "仅管理员可以测试供应商连接"); + assert_eq!(*provider_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + provider_handle.abort(); +} + +#[tokio::test] +async fn gateway_handles_authenticated_test_connection_without_hitting_fallback_probe() { let fallback_probe_hits = Arc::new(Mutex::new(0usize)); let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits); let fallback_probe = Router::new().route( @@ -1983,15 +2302,43 @@ async fn gateway_handles_public_test_connection_without_hitting_fallback_probe() )); let (_unused_fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await; + let now = Utc::now(); + let mut user = sample_auth_user(now); + user.role = "admin".to_string(); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ( + "session_id".to_string(), + json!("session-test-connection-openai"), + ), + ]), + now + chrono::Duration::hours(1), + ); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user])); let gateway = build_router_with_state( AppState::new() .expect("gateway should build") .with_data_state_for_tests( - crate::data::GatewayDataState::with_provider_transport_reader_for_tests( + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests( provider_catalog_repository, - DEVELOPMENT_ENCRYPTION_KEY, - ), - ), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_user_reader(user_repository), + ) + .with_auth_session_for_tests(sample_auth_session( + "user-auth-1", + "session-test-connection-openai", + "device-test-connection-openai", + "refresh-token-placeholder", + now, + )), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -1999,6 +2346,9 @@ async fn gateway_handles_public_test_connection_without_hitting_fallback_probe() .get(format!( "{gateway_url}/v1/test-connection?provider=provider-1&model=gpt-5&api_format=openai:chat" )) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-test-connection-openai") + .header("user-agent", "AetherTest/1.0") .send() .await .expect("request should succeed"); @@ -2073,12 +2423,42 @@ async fn gateway_gemini_test_connection_does_not_force_low_max_output_tokens() { )], )); + let now = Utc::now(); + let mut user = sample_auth_user(now); + user.role = "admin".to_string(); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ( + "session_id".to_string(), + json!("session-test-connection-gemini"), + ), + ]), + now + chrono::Duration::hours(1), + ); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user])); let gateway = build_router_with_state( AppState::new() .expect("gateway should build") - .with_data_state_for_tests(GatewayDataState::with_provider_transport_reader_for_tests( - provider_catalog_repository, - DEVELOPMENT_ENCRYPTION_KEY, + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog_repository, + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_user_reader(user_repository), + ) + .with_auth_session_for_tests(sample_auth_session( + "user-auth-1", + "session-test-connection-gemini", + "device-test-connection-gemini", + "refresh-token-placeholder", + now, )), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -2087,6 +2467,9 @@ async fn gateway_gemini_test_connection_does_not_force_low_max_output_tokens() { .get(format!( "{gateway_url}/v1/test-connection?provider=provider-gemini&model=gemini-3-flash-preview&api_format=gemini:generate_content" )) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-test-connection-gemini") + .header("user-agent", "AetherTest/1.0") .send() .await .expect("request should succeed"); @@ -2212,7 +2595,7 @@ fn test_auth_secret() -> String { .ok() .map(|value| value.trim().to_string()) .filter(|value| !value.is_empty()) - .unwrap_or_else(|| "aether-rust-dev-jwt-secret".to_string()) + .unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string()) } fn test_base64url_encode(bytes: &[u8]) -> String { @@ -2273,6 +2656,14 @@ fn set_test_env_var(key: &'static str, value: &str) -> TestEnvVarGuard { TestEnvVarGuard { key, previous } } +#[cfg(test)] +fn payment_callback_env_lock() -> &'static std::sync::Mutex<()> { + static LOCK: std::sync::OnceLock> = std::sync::OnceLock::new(); + LOCK.get_or_init(|| std::sync::Mutex::new(())) +} + +const TEST_PAYMENT_CALLBACK_SECRET: &str = "test-callback-secret-0123456789abcdef"; + fn canonicalize_test_json(value: &serde_json::Value) -> serde_json::Value { match value { serde_json::Value::Object(map) => { @@ -3578,6 +3969,89 @@ async fn gateway_returns_not_found_for_missing_announcement_read_status_locally( upstream_handle.abort(); } +#[tokio::test] +async fn gateway_cannot_mark_future_required_announcement_as_read() { + let now = Utc::now(); + let user = sample_auth_user(now); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ( + "session_id".to_string(), + json!("session-announcement-future"), + ), + ]), + now + chrono::Duration::hours(1), + ); + let starts_at = now.timestamp().max(0) as u64 + 3600; + let announcement_repository = Arc::new(InMemoryAnnouncementReadRepository::seed(vec![ + StoredAnnouncement::new( + "announcement-future-required".to_string(), + "Future required notice".to_string(), + "Not public yet".to_string(), + "warning".to_string(), + 10, + true, + false, + true, + Some("admin-1".to_string()), + Some("admin".to_string()), + Some(starts_at as i64), + None, + now.timestamp(), + now.timestamp(), + ) + .expect("announcement should build"), + ])); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_announcement_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-announcement-future", + "device-announcement-future", + "refresh-token-placeholder", + now, + )], + Arc::clone(&announcement_repository), + ) + .await; + + let response = reqwest::Client::new() + .patch(format!( + "{gateway_url}/api/announcements/announcement-future-required/read-status" + )) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-announcement-future") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ "is_read": true })) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::NOT_FOUND); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "Announcement not found"); + assert_eq!( + announcement_repository + .count_unread_active_announcements("user-auth-1", starts_at) + .await + .expect("unread count should load"), + 1 + ); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_rejects_invalid_nested_announcement_paths_as_local_not_found_without_hitting_upstream( ) { @@ -3822,7 +4296,7 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea .header("user-agent", "AetherTest/1.0") .json(&json!({ "amount_usd": 10.0, - "payment_method": "alipay", + "payment_method": "manual", "pay_amount": 72.5, "pay_currency": "CNY", "exchange_rate": 7.25, @@ -3831,6 +4305,13 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea .await .expect("request should succeed"); assert_eq!(create_response.status(), StatusCode::OK); + assert_eq!( + create_response + .headers() + .get("cache-control") + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let create_payload: serde_json::Value = create_response .json() .await @@ -3847,17 +4328,15 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea assert_eq!(create_payload["order"]["pay_amount"], 72.5); assert_eq!(create_payload["order"]["pay_currency"], "CNY"); assert_eq!(create_payload["order"]["exchange_rate"], 7.25); - assert_eq!(create_payload["order"]["payment_method"], "alipay"); + assert_eq!(create_payload["order"]["payment_method"], "manual"); assert_eq!(create_payload["order"]["status"], "pending"); - assert_eq!(create_payload["payment_instructions"]["gateway"], "alipay"); - assert!(create_payload["payment_instructions"]["payment_url"] - .as_str() - .unwrap_or_default() - .contains("/pay/mock/alipay/")); - assert!(create_payload["payment_instructions"]["qr_code"] - .as_str() - .unwrap_or_default() - .contains("mock://alipay/")); + assert_eq!(create_payload["payment_instructions"]["gateway"], "manual"); + assert!(create_payload["payment_instructions"]["payment_url"].is_null()); + assert!(create_payload["payment_instructions"]["qr_code"].is_null()); + assert_eq!( + create_payload["payment_instructions"]["instructions"], + "请线下确认到账后由管理员处理" + ); let list_response = client .get(format!( @@ -3891,6 +4370,13 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea .await .expect("request should succeed"); assert_eq!(detail_response.status(), StatusCode::OK); + assert_eq!( + detail_response + .headers() + .get("cache-control") + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let detail_payload: serde_json::Value = detail_response .json() .await @@ -3898,7 +4384,7 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea assert_eq!(detail_payload["order"]["id"], order_id); assert_eq!( detail_payload["order"]["gateway_response"]["gateway"], - "alipay" + "manual" ); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -4065,6 +4551,13 @@ async fn gateway_reuses_pending_billing_plan_checkout_order_without_proxying_ups .await .expect("first checkout request should succeed"); assert_eq!(first_response.status(), StatusCode::OK); + assert_eq!( + first_response + .headers() + .get("cache-control") + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let first_payload: serde_json::Value = first_response .json() .await @@ -4093,6 +4586,13 @@ async fn gateway_reuses_pending_billing_plan_checkout_order_without_proxying_ups .await .expect("second checkout request should succeed"); assert_eq!(second_response.status(), StatusCode::OK); + assert_eq!( + second_response + .headers() + .get("cache-control") + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let second_payload: serde_json::Value = second_response .json() .await @@ -4947,7 +5447,7 @@ async fn gateway_returns_service_unavailable_for_wallet_today_cost_without_usage assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "钱包今日费用数据暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -4989,13 +5489,33 @@ async fn gateway_updates_users_me_detail_locally_without_proxying_upstream() { .await; let client = reqwest::Client::new(); + let unverified_response = client + .put(format!("{gateway_url}/api/users/me")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-users-me-update-detail") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "alice+unverified@example.com", + "username": "alice-updated" + })) + .send() + .await + .expect("request should succeed"); + + assert_eq!(unverified_response.status(), StatusCode::BAD_REQUEST); + let unverified_payload: serde_json::Value = unverified_response + .json() + .await + .expect("json body should parse"); + assert_eq!(unverified_payload["detail"], "修改邮箱前请先完成新邮箱验证"); + let update_response = client .put(format!("{gateway_url}/api/users/me")) .header("authorization", format!("Bearer {access_token}")) .header("x-client-device-id", "device-users-me-update-detail") .header("user-agent", "AetherTest/1.0") .json(&json!({ - "email": "alice+updated@example.com", + "email": "alice@example.com", "username": "alice-updated", "feature_settings": { "chat_pii_redaction": { @@ -5025,7 +5545,7 @@ async fn gateway_updates_users_me_detail_locally_without_proxying_upstream() { assert_eq!(get_response.status(), StatusCode::OK); let get_payload: serde_json::Value = get_response.json().await.expect("json body should parse"); - assert_eq!(get_payload["email"], "alice+updated@example.com"); + assert_eq!(get_payload["email"], "alice@example.com"); assert_eq!(get_payload["username"], "alice-updated"); assert_eq!(get_payload["auth_source"], "local"); assert_eq!(get_payload["has_password"], true); @@ -5092,9 +5612,9 @@ async fn gateway_returns_service_unavailable_for_users_me_detail_update_without_ .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户资料存储暂不可用"); + assert_eq!(payload["detail"], "修改邮箱前请先完成新邮箱验证"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -5158,11 +5678,19 @@ async fn gateway_changes_users_me_password_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(update_response.status(), StatusCode::OK); + let clear_cookie = update_response + .headers() + .get(http::header::SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .expect("password change should clear the refresh cookie") + .to_string(); let update_payload: serde_json::Value = update_response .json() .await .expect("json body should parse"); assert_eq!(update_payload["message"], "密码修改成功"); + assert!(clear_cookie.contains("aether_refresh_token=")); + assert!(clear_cookie.contains("Max-Age=0")); let sessions_response = client .get(format!("{gateway_url}/api/users/me/sessions")) @@ -5173,16 +5701,12 @@ async fn gateway_changes_users_me_password_locally_without_proxying_upstream() { .await .expect("request should succeed"); - assert_eq!(sessions_response.status(), StatusCode::OK); + assert_eq!(sessions_response.status(), StatusCode::UNAUTHORIZED); let sessions_payload: serde_json::Value = sessions_response .json() .await .expect("json body should parse"); - let sessions = sessions_payload - .as_array() - .expect("sessions should be array"); - assert_eq!(sessions.len(), 1); - assert_eq!(sessions[0]["id"], "session-users-me-password-current"); + assert_eq!(sessions_payload["detail"], "登录会话已失效,请重新登录"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -5245,7 +5769,7 @@ async fn gateway_returns_service_unavailable_for_users_me_password_change_withou assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户凭证存储暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -5454,7 +5978,7 @@ async fn gateway_returns_service_unavailable_for_users_me_endpoint_status_withou assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户端点健康数据暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -5619,7 +6143,7 @@ async fn gateway_returns_service_unavailable_for_users_me_preferences_update_wit assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json should parse"); - assert_eq!(payload["detail"], "用户偏好设置存储暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -6215,7 +6739,7 @@ async fn gateway_returns_service_unavailable_for_users_me_usage_routes_without_r .expect("request should succeed"); assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE, "{path}"); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户用量数据暂不可用", "{path}"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试", "{path}"); } assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); @@ -6290,6 +6814,13 @@ async fn gateway_handles_users_me_sessions_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); let sessions = payload.as_array().expect("sessions should be array"); assert_eq!(sessions.len(), 2); @@ -6532,6 +7063,125 @@ async fn gateway_updates_users_me_session_label_locally_without_proxying_upstrea upstream_handle.abort(); } +#[tokio::test] +async fn gateway_cannot_update_or_revoke_another_users_session() { + let now = Utc::now(); + let user = sample_auth_user(now); + let mut foreign_user = user.clone(); + foreign_user.id = "user-auth-2".to_string(); + foreign_user.email = Some("bob@example.com".to_string()); + foreign_user.username = "bob".to_string(); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ( + "session_id".to_string(), + json!("session-users-me-owner-current"), + ), + ]), + now + chrono::Duration::hours(1), + ); + let foreign_access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(foreign_user.id)), + ("role".to_string(), json!(foreign_user.role)), + ( + "created_at".to_string(), + json!(foreign_user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-users-me-foreign")), + ]), + now + chrono::Duration::hours(1), + ); + let owner_session = sample_auth_session( + "user-auth-1", + "session-users-me-owner-current", + "device-users-me-owner-current", + "refresh-token-owner-current", + now, + ); + let mut foreign_session = sample_auth_session( + "user-auth-2", + "session-users-me-foreign", + "device-users-me-foreign", + "refresh-token-foreign", + now, + ); + foreign_session.device_label = Some("Bob phone".to_string()); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ + user, + foreign_user, + ])); + + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_builder(|| { + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_user_reader_for_tests( + user_repository, + )) + .with_auth_sessions_for_tests([owner_session, foreign_session]) + }) + .await; + let client = reqwest::Client::new(); + + let update_response = client + .patch(format!( + "{gateway_url}/api/users/me/sessions/session-users-me-foreign" + )) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-users-me-owner-current") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ "device_label": "Compromised" })) + .send() + .await + .expect("cross-user update request should complete"); + assert_eq!(update_response.status(), StatusCode::NOT_FOUND); + + let revoke_response = client + .delete(format!( + "{gateway_url}/api/users/me/sessions/session-users-me-foreign" + )) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-users-me-owner-current") + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("cross-user revoke request should complete"); + assert_eq!(revoke_response.status(), StatusCode::NOT_FOUND); + + let foreign_sessions_response = client + .get(format!("{gateway_url}/api/users/me/sessions")) + .header("authorization", format!("Bearer {foreign_access_token}")) + .header("x-client-device-id", "device-users-me-foreign") + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("foreign owner session request should complete"); + assert_eq!(foreign_sessions_response.status(), StatusCode::OK); + let foreign_sessions: serde_json::Value = foreign_sessions_response + .json() + .await + .expect("json body should parse"); + let foreign_sessions = foreign_sessions + .as_array() + .expect("sessions should be an array"); + assert_eq!(foreign_sessions.len(), 1); + assert_eq!(foreign_sessions[0]["id"], "session-users-me-foreign"); + assert_eq!(foreign_sessions[0]["device_label"], "Bob phone"); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_rejects_invalid_users_me_session_delete_path_as_local_not_found_without_hitting_upstream( ) { @@ -6646,7 +7296,7 @@ async fn gateway_handles_users_me_api_keys_locally_without_proxying_upstream() { .expect("ciphertext should build"); let auth_repository = Arc::new( InMemoryAuthApiKeySnapshotRepository::seed(vec![( - Some("hash-user-key-1".to_string()), + Some("87e376524de9f1d78f2b7d95d20bcc3faed619d572ccc9a1d9a343ab83e33834".to_string()), StoredAuthApiKeySnapshot::new( "user-auth-1".to_string(), "alice".to_string(), @@ -6675,7 +7325,7 @@ async fn gateway_handles_users_me_api_keys_locally_without_proxying_upstream() { .with_export_records(vec![StoredAuthApiKeyExportRecord::new( "user-auth-1".to_string(), "user-key-1".to_string(), - "hash-user-key-1".to_string(), + "87e376524de9f1d78f2b7d95d20bcc3faed619d572ccc9a1d9a343ab83e33834".to_string(), Some(encrypted), Some("primary".to_string()), Some(json!(["openai"])), @@ -6706,7 +7356,8 @@ async fn gateway_handles_users_me_api_keys_locally_without_proxying_upstream() { start_auth_gateway_with_builder(|| { let data_state = crate::data::GatewayDataState::with_user_reader_for_tests(user_repository) - .with_auth_api_key_reader(auth_repository); + .attach_auth_api_key_repository_for_tests(auth_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) @@ -6737,7 +7388,7 @@ async fn gateway_handles_users_me_api_keys_locally_without_proxying_upstream() { assert_eq!(api_keys.len(), 1); assert_eq!(api_keys[0]["id"], "user-key-1"); assert_eq!(api_keys[0]["name"], "primary"); - assert_eq!(api_keys[0]["key_display"], "sk-user-li...ve-1"); + assert_eq!(api_keys[0]["key_display"], "sk-u...e-1"); assert_eq!(api_keys[0]["total_requests"], 9); assert_eq!(api_keys[0]["total_cost_usd"], 1.5); assert_eq!(api_keys[0]["created_at"], "2024-03-21T05:48:20+00:00"); @@ -6755,6 +7406,13 @@ async fn gateway_handles_users_me_api_keys_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(detail_response.status(), StatusCode::OK); + assert_eq!( + detail_response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let detail_payload: serde_json::Value = detail_response .json() .await @@ -7118,13 +7776,25 @@ async fn gateway_handles_users_me_api_key_writes_locally_without_proxying_upstre .header("user-agent", "AetherTest/1.0") .json(&json!({ "name": "writer-key", - "rate_limit": 120 + "rate_limit": 120, + "feature_settings": { + "chat_pii_redaction": { + "enabled": true + } + } })) .send() .await .expect("request should succeed"); assert_eq!(create_response.status(), StatusCode::OK); + assert_eq!( + create_response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let create_payload: serde_json::Value = create_response .json() .await @@ -7136,7 +7806,10 @@ async fn gateway_handles_users_me_api_key_writes_locally_without_proxying_upstre assert_eq!(create_payload["name"], "writer-key"); assert_eq!(create_payload["rate_limit"], 120); assert_eq!(create_payload["concurrent_limit"], serde_json::Value::Null); - assert_eq!(create_payload["feature_settings"], serde_json::Value::Null); + assert_eq!( + create_payload["feature_settings"]["chat_pii_redaction"]["enabled"], + true + ); assert_eq!(create_payload["message"], "API密钥创建成功"); let created_at = create_payload["created_at"] .as_str() @@ -7159,7 +7832,7 @@ async fn gateway_handles_users_me_api_key_writes_locally_without_proxying_upstre "concurrent_limit": 4, "feature_settings": { "chat_pii_redaction": { - "enabled": true + "enabled": false } } })) @@ -7176,7 +7849,7 @@ async fn gateway_handles_users_me_api_key_writes_locally_without_proxying_upstre assert_eq!(update_payload["concurrent_limit"], 4); assert_eq!( update_payload["feature_settings"]["chat_pii_redaction"]["enabled"], - true + false ); assert_eq!(update_payload["message"], "API密钥已更新"); @@ -7260,7 +7933,7 @@ async fn gateway_handles_users_me_api_key_writes_locally_without_proxying_upstre assert_eq!(detail_payload["created_at"], created_at); assert_eq!( detail_payload["feature_settings"]["chat_pii_redaction"]["enabled"], - true + false ); let delete_response = client @@ -7415,7 +8088,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(create_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(create_payload["detail"], "服务暂不可用,请稍后重试"); let update_response = client .put(format!( @@ -7436,7 +8109,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(update_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(update_payload["detail"], "服务暂不可用,请稍后重试"); let toggle_response = client .patch(format!( @@ -7454,7 +8127,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(toggle_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(toggle_payload["detail"], "服务暂不可用,请稍后重试"); let providers_response = client .put(format!( @@ -7472,7 +8145,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(providers_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(providers_payload["detail"], "服务暂不可用,请稍后重试"); let capabilities_response = client .put(format!( @@ -7493,7 +8166,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(capabilities_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(capabilities_payload["detail"], "服务暂不可用,请稍后重试"); let delete_response = client .delete(format!( @@ -7510,7 +8183,7 @@ async fn gateway_returns_service_unavailable_for_users_me_api_key_writes_without .json() .await .expect("json body should parse"); - assert_eq!(delete_payload["detail"], "用户 API 密钥写入暂不可用"); + assert_eq!(delete_payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -7547,7 +8220,7 @@ async fn gateway_handles_users_me_management_token_reads_locally_without_proxyin start_auth_gateway_with_builder(|| { let data_state = crate::data::GatewayDataState::with_management_token_repository_for_tests( - repository, + Arc::clone(&repository), ) .with_user_reader(user_repository); AppState::new() @@ -7622,6 +8295,75 @@ async fn gateway_handles_users_me_management_token_reads_locally_without_proxyin .await .expect("request should succeed"); assert_eq!(foreign_response.status(), StatusCode::NOT_FOUND); + + let foreign_before = repository + .get_management_token_with_user("mt-user-2") + .await + .expect("foreign token lookup should succeed") + .expect("foreign token should exist"); + let update_response = client + .put(format!("{gateway_url}/api/me/management-tokens/mt-user-2")) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-users-me-management-token-reads", + ) + .header("user-agent", "AetherTest/1.0") + .json(&json!({ "name": "cross-user-update" })) + .send() + .await + .expect("cross-user token update should complete"); + assert_eq!(update_response.status(), StatusCode::FORBIDDEN); + + let toggle_response = client + .patch(format!( + "{gateway_url}/api/me/management-tokens/mt-user-2/status" + )) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-users-me-management-token-reads", + ) + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("cross-user token toggle should complete"); + assert_eq!(toggle_response.status(), StatusCode::FORBIDDEN); + + let regenerate_response = client + .post(format!( + "{gateway_url}/api/me/management-tokens/mt-user-2/regenerate" + )) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-users-me-management-token-reads", + ) + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("cross-user token regenerate should complete"); + assert_eq!(regenerate_response.status(), StatusCode::FORBIDDEN); + + let delete_response = client + .delete(format!("{gateway_url}/api/me/management-tokens/mt-user-2")) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-users-me-management-token-reads", + ) + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("cross-user token delete should complete"); + assert_eq!(delete_response.status(), StatusCode::FORBIDDEN); + + let foreign_after = repository + .get_management_token_with_user("mt-user-2") + .await + .expect("foreign token lookup should succeed") + .expect("foreign token should still exist"); + assert_eq!(foreign_after, foreign_before); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -7907,7 +8649,7 @@ async fn gateway_handles_users_me_management_token_writes_locally_without_proxyi } #[tokio::test] -async fn gateway_rejects_users_me_management_token_create_for_non_admin_user_without_proxying_upstream( +async fn gateway_rejects_users_me_management_token_writes_for_non_admin_user_without_proxying_upstream( ) { let now = Utc::now(); let user = sample_auth_user(now); @@ -7975,6 +8717,41 @@ async fn gateway_rejects_users_me_management_token_create_for_non_admin_user_wit create_payload["detail"], json!("仅管理员可以创建 Management Token") ); + + let client = reqwest::Client::new(); + for (method, path) in [ + (reqwest::Method::PUT, "/api/me/management-tokens/mt-owned"), + ( + reqwest::Method::PATCH, + "/api/me/management-tokens/mt-owned/status", + ), + ( + reqwest::Method::POST, + "/api/me/management-tokens/mt-owned/regenerate", + ), + ( + reqwest::Method::DELETE, + "/api/me/management-tokens/mt-owned", + ), + ] { + let response = client + .request(method, format!("{gateway_url}{path}")) + .header("authorization", format!("Bearer {access_token}")) + .header( + "x-client-device-id", + "device-users-me-management-token-create-denied", + ) + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("request should succeed"); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!( + payload["detail"], + json!("仅管理员可以管理 Management Token") + ); + } assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -8035,7 +8812,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_reads assert_eq!(list_response.status(), StatusCode::SERVICE_UNAVAILABLE); let list_payload: serde_json::Value = list_response.json().await.expect("json body should parse"); - assert_eq!(list_payload["detail"], "用户 Management Token 数据暂不可用"); + assert_eq!(list_payload["detail"], "服务暂不可用,请稍后重试"); let detail_response = client .get(format!("{gateway_url}/api/me/management-tokens/mt-missing")) @@ -8053,10 +8830,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_reads .json() .await .expect("json body should parse"); - assert_eq!( - detail_payload["detail"], - "用户 Management Token 数据暂不可用" - ); + assert_eq!(detail_payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -8129,10 +8903,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_write .json() .await .expect("json body should parse"); - assert_eq!( - create_payload["detail"], - "用户 Management Token 写入暂不可用" - ); + assert_eq!(create_payload["detail"], "服务暂不可用,请稍后重试"); let update_response = client .put(format!( @@ -8155,10 +8926,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_write .json() .await .expect("json body should parse"); - assert_eq!( - update_payload["detail"], - "用户 Management Token 写入暂不可用" - ); + assert_eq!(update_payload["detail"], "服务暂不可用,请稍后重试"); let toggle_response = client .patch(format!( @@ -8178,10 +8946,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_write .json() .await .expect("json body should parse"); - assert_eq!( - toggle_payload["detail"], - "用户 Management Token 写入暂不可用" - ); + assert_eq!(toggle_payload["detail"], "服务暂不可用,请稍后重试"); let regenerate_response = client .post(format!( @@ -8204,10 +8969,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_write .json() .await .expect("json body should parse"); - assert_eq!( - regenerate_payload["detail"], - "用户 Management Token 写入暂不可用" - ); + assert_eq!(regenerate_payload["detail"], "服务暂不可用,请稍后重试"); let delete_response = client .delete(format!( @@ -8227,10 +8989,7 @@ async fn gateway_returns_service_unavailable_for_users_me_management_token_write .json() .await .expect("json body should parse"); - assert_eq!( - delete_payload["detail"], - "用户 Management Token 写入暂不可用" - ); + assert_eq!(delete_payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -8410,6 +9169,13 @@ async fn gateway_handles_users_me_detail_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["id"], "user-auth-1"); assert_eq!(payload["email"], "alice@example.com"); @@ -8446,6 +9212,13 @@ async fn gateway_handles_auth_login_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let set_cookie = response .headers() .get(reqwest::header::SET_COOKIE) @@ -8479,6 +9252,155 @@ async fn gateway_handles_auth_login_locally_without_proxying_upstream() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_auth_login_rejects_query_only_device_binding_without_creating_session() { + let now = Utc::now(); + let user = sample_auth_user(now); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state(user, sample_auth_wallet("user-auth-1", now), []).await; + + let client = reqwest::Client::new(); + let rejected = client + .post(format!( + "{gateway_url}/api/auth/login?client_device_id=device-auth-login-csrf" + )) + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "alice@example.com", + "password": "secret123", + "auth_type": "local", + })) + .send() + .await + .expect("query-only login request should complete"); + + assert_eq!(rejected.status(), StatusCode::BAD_REQUEST); + assert!(!rejected.headers().contains_key(http::header::SET_COOKIE)); + + let legitimate = client + .post(format!("{gateway_url}/api/auth/login")) + .header("x-client-device-id", "device-auth-login-csrf") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "alice@example.com", + "password": "secret123", + "auth_type": "local", + })) + .send() + .await + .expect("legitimate login should complete"); + assert_eq!(legitimate.status(), StatusCode::OK); + assert!(legitimate.headers().contains_key(http::header::SET_COOKIE)); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_rate_limits_repeated_auth_login_attempts_by_identifier() { + let now = Utc::now(); + let user = sample_auth_user(now); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state(user, sample_auth_wallet("user-auth-1", now), []).await; + + let client = reqwest::Client::new(); + for attempt in 1..=10 { + let response = client + .post(format!("{gateway_url}/api/auth/login")) + .header("x-client-device-id", "device-auth-login-rate-limit") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "missing@example.com", + "password": "incorrect-password", + "auth_type": "local", + })) + .send() + .await + .expect("login request should succeed"); + assert_eq!( + response.status(), + StatusCode::UNAUTHORIZED, + "attempt {attempt} should reach authentication" + ); + } + + let rejected = client + .post(format!("{gateway_url}/api/auth/login")) + .header("x-client-device-id", "device-auth-login-rate-limit") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "missing@example.com", + "password": "incorrect-password", + "auth_type": "local", + })) + .send() + .await + .expect("rate-limited login request should succeed"); + + assert_eq!(rejected.status(), StatusCode::TOO_MANY_REQUESTS); + let retry_after = rejected + .headers() + .get(reqwest::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + .expect("Retry-After should be a positive integer"); + assert!(retry_after > 0 && retry_after <= 60); + let payload: serde_json::Value = rejected.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "请求过于频繁,请稍后重试"); + assert_eq!(payload["retry_after"], retry_after); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_auth_login_identifier_rate_limit_is_atomic_under_concurrency() { + let now = Utc::now(); + let user = sample_auth_user(now); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state(user, sample_auth_wallet("user-auth-1", now), []).await; + + let client = reqwest::Client::new(); + let mut tasks = Vec::new(); + for _ in 0..12 { + let client = client.clone(); + let gateway_url = gateway_url.clone(); + tasks.push(tokio::spawn(async move { + client + .post(format!("{gateway_url}/api/auth/login")) + .header("x-client-device-id", "device-auth-login-concurrent") + .header("user-agent", "AetherTest/1.0") + .json(&json!({ + "email": "concurrent-missing@example.com", + "password": "incorrect-password", + "auth_type": "local", + })) + .send() + .await + .expect("concurrent login request should succeed") + .status() + })); + } + + let mut unauthorized = 0; + let mut rate_limited = 0; + for task in tasks { + match task.await.expect("login task should complete") { + StatusCode::UNAUTHORIZED => unauthorized += 1, + StatusCode::TOO_MANY_REQUESTS => rate_limited += 1, + status => panic!("unexpected concurrent login status: {status}"), + } + } + assert_eq!(unauthorized, 10); + assert_eq!(rate_limited, 2); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_rejects_unsupported_auth_login_type_locally_without_proxying_upstream() { let now = Utc::now(); @@ -8534,7 +9456,7 @@ async fn gateway_handles_auth_ldap_login_locally_without_proxying_upstream() { )); let data_state = crate::data::GatewayDataState::disabled() .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) - .attach_auth_module_reader_for_tests(auth_module_repository) + .attach_auth_module_repository_for_tests(auth_module_repository) .with_system_config_values_for_tests(vec![ ("module.ldap.enabled".to_string(), json!(true)), ("default_user_initial_gift_usd".to_string(), json!(9.5)), @@ -8616,7 +9538,10 @@ async fn gateway_handles_auth_register_locally_without_proxying_upstream() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) }) .await; @@ -8627,6 +9552,7 @@ async fn gateway_handles_auth_register_locally_without_proxying_upstream() { "email": "alice@example.com", "username": "alice", "password": "secret123", + "email_verification_token": TEST_EMAIL_VERIFICATION_TOKEN, })) .send() .await @@ -8705,7 +9631,10 @@ async fn gateway_rejects_auth_register_without_current_privacy_policy_acceptance AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) }) .await; @@ -8805,7 +9734,10 @@ async fn gateway_rejects_auth_register_without_turnstile_token_when_enabled() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) }) .await; @@ -8815,6 +9747,7 @@ async fn gateway_rejects_auth_register_without_turnstile_token_when_enabled() { "email": "alice@example.com", "username": "alice", "password": "secret123", + "email_verification_token": TEST_EMAIL_VERIFICATION_TOKEN, })) .send() .await @@ -8862,7 +9795,10 @@ async fn gateway_rejects_auth_register_with_oversized_turnstile_token() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) }) .await; @@ -8904,7 +9840,10 @@ async fn gateway_returns_service_unavailable_when_turnstile_keys_are_incomplete( AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) }) .await; @@ -8922,7 +9861,7 @@ async fn gateway_returns_service_unavailable_when_turnstile_keys_are_incomplete( assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "人机验证服务暂不可用,请稍后重试"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -8930,7 +9869,7 @@ async fn gateway_returns_service_unavailable_when_turnstile_keys_are_incomplete( } #[tokio::test] -async fn gateway_allows_auth_register_after_successful_turnstile_verification() { +async fn gateway_ignores_spoofed_cf_ip_during_successful_turnstile_verification() { let (siteverify_url, turnstile_requests, turnstile_handle) = start_turnstile_siteverify_server( json!({ "success": true, @@ -8947,7 +9886,10 @@ async fn gateway_allows_auth_register_after_successful_turnstile_verification() AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -8961,6 +9903,7 @@ async fn gateway_allows_auth_register_after_successful_turnstile_verification() "username": "alice", "password": "secret123", "turnstile_token": "valid-token", + "email_verification_token": TEST_EMAIL_VERIFICATION_TOKEN, })) .send() .await @@ -8987,7 +9930,7 @@ async fn gateway_allows_auth_register_after_successful_turnstile_verification() ); assert_eq!( requests[0].get("remoteip").map(String::as_str), - Some("203.0.113.10") + Some("127.0.0.1") ); assert!(requests[0].contains_key("idempotency_key")); drop(requests); @@ -9017,7 +9960,10 @@ async fn gateway_rejects_auth_register_when_turnstile_action_mismatches() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -9063,7 +10009,10 @@ async fn gateway_rejects_auth_register_when_turnstile_siteverify_rejects_token() AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -9109,7 +10058,10 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_reports_s AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -9129,7 +10081,7 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_reports_s assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "人机验证服务暂不可用,请稍后重试"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -9156,7 +10108,10 @@ async fn gateway_rejects_auth_register_when_turnstile_hostname_mismatches() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -9199,7 +10154,10 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_fails() { AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) } }) @@ -9219,7 +10177,7 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_fails() { assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "人机验证服务暂不可用,请稍后重试"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -9247,7 +10205,10 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_times_out AppState::new() .expect("gateway should build") .with_data_state_for_tests(turnstile_enabled_data_state()) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .with_turnstile_siteverify_url_for_tests(&siteverify_url) .with_turnstile_siteverify_timeout_for_tests(Duration::from_millis(20)) } @@ -9268,7 +10229,7 @@ async fn gateway_returns_service_unavailable_when_turnstile_siteverify_times_out assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "人机验证服务暂不可用,请稍后重试"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -9291,7 +10252,10 @@ async fn gateway_returns_service_unavailable_for_auth_register_without_storage() AppState::new() .expect("gateway should build") .with_data_state_for_tests(data_state) - .with_auth_email_verified_for_tests("alice@example.com") + .with_auth_email_verified_for_tests( + "alice@example.com", + TEST_EMAIL_VERIFICATION_TOKEN, + ) .without_auth_user_store_for_tests() }) .await; @@ -9302,6 +10266,7 @@ async fn gateway_returns_service_unavailable_for_auth_register_without_storage() "email": "alice@example.com", "username": "alice", "password": "secret123", + "email_verification_token": TEST_EMAIL_VERIFICATION_TOKEN, })) .send() .await @@ -9309,7 +10274,7 @@ async fn gateway_returns_service_unavailable_for_auth_register_without_storage() assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "注册数据存储暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -9343,14 +10308,28 @@ async fn gateway_handles_auth_send_verification_code_locally_without_proxying_up .expect("request should succeed"); assert_eq!(send_response.status(), StatusCode::OK); + assert_eq!( + send_response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let send_payload: serde_json::Value = send_response.json().await.expect("json body should parse"); assert_eq!(send_payload["success"], true); assert_eq!(send_payload["expire_minutes"], 5); + let verification_token = send_payload["verification_token"] + .as_str() + .expect("verification token should exist") + .to_string(); let status_response = client .post(format!("{gateway_url}/api/auth/verification-status")) - .json(&json!({ "email": "alice@example.com" })) + .json(&json!({ + "email": "alice@example.com", + "verification_token": verification_token, + })) .send() .await .expect("status request should succeed"); @@ -9507,13 +10486,21 @@ async fn gateway_handles_auth_verification_status_locally_without_proxying_upstr start_auth_gateway_with_builder(|| { AppState::new() .expect("gateway should build") - .with_auth_email_verification_pending_for_tests("alice@example.com", "123456", now) + .with_auth_email_verification_pending_for_tests( + "alice@example.com", + "123456", + TEST_EMAIL_VERIFICATION_TOKEN, + now, + ) }) .await; let response = reqwest::Client::new() .post(format!("{gateway_url}/api/auth/verification-status")) - .json(&json!({ "email": "alice@example.com" })) + .json(&json!({ + "email": "alice@example.com", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) .send() .await .expect("request should succeed"); @@ -9538,14 +10525,23 @@ async fn gateway_handles_auth_verify_email_locally_without_proxying_upstream() { start_auth_gateway_with_builder(|| { AppState::new() .expect("gateway should build") - .with_auth_email_verification_pending_for_tests("alice@example.com", "123456", now) + .with_auth_email_verification_pending_for_tests( + "alice@example.com", + "123456", + TEST_EMAIL_VERIFICATION_TOKEN, + now, + ) }) .await; let client = reqwest::Client::new(); let verify_response = client .post(format!("{gateway_url}/api/auth/verify-email")) - .json(&json!({ "email": "alice@example.com", "code": "123456" })) + .json(&json!({ + "email": "alice@example.com", + "code": "123456", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) .send() .await .expect("request should succeed"); @@ -9559,7 +10555,10 @@ async fn gateway_handles_auth_verify_email_locally_without_proxying_upstream() { let status_response = client .post(format!("{gateway_url}/api/auth/verification-status")) - .json(&json!({ "email": "alice@example.com" })) + .json(&json!({ + "email": "alice@example.com", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) .send() .await .expect("status request should succeed"); @@ -9577,6 +10576,86 @@ async fn gateway_handles_auth_verify_email_locally_without_proxying_upstream() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_invalidates_email_verification_after_five_incorrect_codes() { + let now = Utc::now(); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_builder(|| { + AppState::new() + .expect("gateway should build") + .with_auth_email_verification_pending_for_tests( + "locked@example.com", + "123456", + TEST_EMAIL_VERIFICATION_TOKEN, + now, + ) + }) + .await; + + let client = reqwest::Client::new(); + for attempt in 1..=4 { + let response = client + .post(format!("{gateway_url}/api/auth/verify-email")) + .json(&json!({ + "email": "locked@example.com", + "code": "000000", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) + .send() + .await + .expect("incorrect verification request should succeed"); + assert_eq!( + response.status(), + StatusCode::BAD_REQUEST, + "attempt {attempt} should report an incorrect code" + ); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "验证码错误"); + } + + let exhausted = client + .post(format!("{gateway_url}/api/auth/verify-email")) + .json(&json!({ + "email": "locked@example.com", + "code": "000000", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) + .send() + .await + .expect("exhausted verification request should succeed"); + assert_eq!(exhausted.status(), StatusCode::TOO_MANY_REQUESTS); + assert!(exhausted + .headers() + .contains_key(reqwest::header::RETRY_AFTER)); + let exhausted_payload: serde_json::Value = + exhausted.json().await.expect("json body should parse"); + assert_eq!( + exhausted_payload["detail"], + "验证码尝试次数过多,请重新获取" + ); + + let correct_after_exhaustion = client + .post(format!("{gateway_url}/api/auth/verify-email")) + .json(&json!({ + "email": "locked@example.com", + "code": "123456", + "verification_token": TEST_EMAIL_VERIFICATION_TOKEN, + })) + .send() + .await + .expect("verification request after exhaustion should succeed"); + assert_eq!(correct_after_exhaustion.status(), StatusCode::BAD_REQUEST); + let payload: serde_json::Value = correct_after_exhaustion + .json() + .await + .expect("json body should parse"); + assert_eq!(payload["detail"], "验证会话无效或已过期"); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_handles_auth_me_locally_without_proxying_upstream() { let now = Utc::now(); @@ -9618,6 +10697,13 @@ async fn gateway_handles_auth_me_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let payload: serde_json::Value = response.json().await.expect("json body should parse"); assert_eq!(payload["id"], "user-auth-1"); assert_eq!(payload["email"], "alice@example.com"); @@ -9630,6 +10716,87 @@ async fn gateway_handles_auth_me_locally_without_proxying_upstream() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_rejects_access_and_refresh_tokens_from_a_stale_security_version() { + let now = Utc::now(); + let user = sample_auth_user(now) + .with_security_version(1) + .expect("auth user security version should update"); + let refresh_token = build_test_auth_token( + "refresh", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-stale")), + ("jti".to_string(), json!("jti-auth-stale")), + ]), + now + chrono::Duration::days(7), + ); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-stale")), + ]), + now + chrono::Duration::hours(1), + ); + let stale_session = sample_auth_session( + "user-auth-1", + "session-auth-stale", + "device-auth-stale", + &refresh_token, + now, + ); + assert_eq!(stale_session.security_version, 0); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [stale_session], + ) + .await; + let client = reqwest::Client::new(); + + let me_response = client + .get(format!("{gateway_url}/api/auth/me")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-auth-stale") + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("stale access request should complete"); + assert_eq!(me_response.status(), StatusCode::UNAUTHORIZED); + let me_payload: serde_json::Value = me_response.json().await.expect("json body should parse"); + assert_eq!(me_payload["detail"], "登录会话已失效,请重新登录"); + + let refresh_response = client + .post(format!("{gateway_url}/api/auth/refresh")) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .header("x-client-device-id", "device-auth-stale") + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("stale refresh request should complete"); + assert_eq!(refresh_response.status(), StatusCode::UNAUTHORIZED); + let refresh_payload: serde_json::Value = refresh_response + .json() + .await + .expect("json body should parse"); + assert_eq!(refresh_payload["detail"], "登录会话已失效,请重新登录"); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_handles_users_me_available_models_locally_without_proxying_upstream() { let now = Utc::now(); @@ -10015,7 +11182,7 @@ async fn gateway_returns_service_unavailable_for_users_me_available_models_witho assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户提供商目录暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -10221,7 +11388,7 @@ async fn gateway_returns_service_unavailable_for_users_me_model_capabilities_upd assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); let payload: serde_json::Value = response.json().await.expect("json body should parse"); - assert_eq!(payload["detail"], "用户模型能力配置存储暂不可用"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); @@ -10269,6 +11436,13 @@ async fn gateway_handles_auth_refresh_locally_without_proxying_upstream() { .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); let set_cookie = response .headers() .get(http::header::SET_COOKIE) @@ -10291,6 +11465,139 @@ async fn gateway_handles_auth_refresh_locally_without_proxying_upstream() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_auth_refresh_rejects_query_only_device_binding_without_rotating_session() { + let now = Utc::now(); + let user = sample_auth_user(now); + let refresh_token = build_test_auth_token( + "refresh", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-refresh-csrf")), + ("jti".to_string(), json!("jti-auth-refresh-csrf")), + ]), + now + chrono::Duration::days(7), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-auth-refresh-csrf", + "device-auth-refresh-csrf", + &refresh_token, + now, + )], + ) + .await; + + let client = reqwest::Client::new(); + let rejected = client + .post(format!( + "{gateway_url}/api/auth/refresh?client_device_id=device-auth-refresh-csrf" + )) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .send() + .await + .expect("query-only refresh request should complete"); + + assert_eq!(rejected.status(), StatusCode::BAD_REQUEST); + assert!(!rejected.headers().contains_key(http::header::SET_COOKIE)); + + let legitimate = client + .post(format!("{gateway_url}/api/auth/refresh")) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .header("x-client-device-id", "device-auth-refresh-csrf") + .send() + .await + .expect("legitimate refresh should complete"); + assert_eq!(legitimate.status(), StatusCode::OK); + assert!(legitimate.headers().contains_key(http::header::SET_COOKIE)); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_allows_only_one_concurrent_refresh_rotation() { + let now = Utc::now(); + let user = sample_auth_user(now); + let refresh_token = build_test_auth_token( + "refresh", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-refresh-race")), + ("jti".to_string(), json!("jti-auth-refresh-race")), + ]), + now + chrono::Duration::days(7), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-auth-refresh-race", + "device-auth-refresh-race", + &refresh_token, + now, + )], + ) + .await; + + let client = reqwest::Client::new(); + let build_request = || { + client + .post(format!("{gateway_url}/api/auth/refresh")) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .header("x-client-device-id", "device-auth-refresh-race") + .header("user-agent", "AetherTest/1.0") + .send() + }; + let (first, second) = tokio::join!(build_request(), build_request()); + let responses = [ + first.expect("first refresh request should complete"), + second.expect("second refresh request should complete"), + ]; + + let mut success_count = 0; + let mut conflict_count = 0; + for response in responses { + let status = response.status(); + let has_set_cookie = response.headers().contains_key(http::header::SET_COOKIE); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + match status { + StatusCode::OK => { + success_count += 1; + assert!(has_set_cookie); + assert!(payload["access_token"].as_str().is_some()); + } + StatusCode::CONFLICT => { + conflict_count += 1; + assert!(!has_set_cookie); + assert!(payload.get("access_token").is_none()); + } + status => panic!("unexpected concurrent refresh status: {status}"), + } + } + assert_eq!(success_count, 1); + assert_eq!(conflict_count, 1); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_handles_auth_logout_locally_without_proxying_upstream() { let now = Utc::now(); @@ -10359,9 +11666,217 @@ async fn gateway_handles_auth_logout_locally_without_proxying_upstream() { upstream_handle.abort(); } +#[tokio::test] +async fn gateway_logout_rejects_tokens_bound_to_a_replaced_user_identity() { + let now = Utc::now(); + let mut user = sample_auth_user(now); + let original_created_at = user + .created_at + .expect("test user should have creation time"); + let replacement_created_at = original_created_at + chrono::Duration::milliseconds(1); + user.created_at = Some(replacement_created_at); + + let build_access_token = |created_at: chrono::DateTime| { + build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ("created_at".to_string(), json!(created_at.to_rfc3339())), + ( + "session_id".to_string(), + json!("session-auth-replaced-user"), + ), + ]), + now + chrono::Duration::hours(1), + ) + }; + let stale_access_token = build_access_token(original_created_at); + let current_access_token = build_access_token(replacement_created_at); + let stale_refresh_token = build_test_auth_token( + "refresh", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(original_created_at.to_rfc3339()), + ), + ( + "session_id".to_string(), + json!("session-auth-replaced-user"), + ), + ("jti".to_string(), json!("jti-auth-replaced-user")), + ]), + now + chrono::Duration::days(7), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-auth-replaced-user", + "device-auth-replaced-user", + &stale_refresh_token, + now, + )], + ) + .await; + + let client = reqwest::Client::new(); + let stale_response = client + .post(format!("{gateway_url}/api/auth/logout")) + .header("authorization", format!("Bearer {stale_access_token}")) + .header( + "cookie", + format!("aether_refresh_token={stale_refresh_token}"), + ) + .header("x-client-device-id", "device-auth-replaced-user") + .send() + .await + .expect("stale logout request should complete"); + assert_eq!(stale_response.status(), StatusCode::UNAUTHORIZED); + + let current_response = client + .post(format!("{gateway_url}/api/auth/logout")) + .header("authorization", format!("Bearer {current_access_token}")) + .header("x-client-device-id", "device-auth-replaced-user") + .send() + .await + .expect("current logout request should complete"); + assert_eq!(current_response.status(), StatusCode::OK); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_cookie_logout_rejects_query_only_device_binding_without_revoking_session() { + let now = Utc::now(); + let user = sample_auth_user(now); + let refresh_token = build_test_auth_token( + "refresh", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-logout-csrf")), + ("jti".to_string(), json!("jti-auth-logout-csrf")), + ]), + now + chrono::Duration::days(7), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-auth-logout-csrf", + "device-auth-logout-csrf", + &refresh_token, + now, + )], + ) + .await; + + let client = reqwest::Client::new(); + let rejected = client + .post(format!( + "{gateway_url}/api/auth/logout?client_device_id=device-auth-logout-csrf" + )) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .send() + .await + .expect("query-only logout request should complete"); + + assert_eq!(rejected.status(), StatusCode::BAD_REQUEST); + assert!(!rejected.headers().contains_key(http::header::SET_COOKIE)); + + let legitimate = client + .post(format!("{gateway_url}/api/auth/logout")) + .header("cookie", format!("aether_refresh_token={refresh_token}")) + .header("x-client-device-id", "device-auth-logout-csrf") + .send() + .await + .expect("legitimate logout should complete"); + assert_eq!(legitimate.status(), StatusCode::OK); + let clear_cookie = legitimate + .headers() + .get(http::header::SET_COOKIE) + .and_then(|value| value.to_str().ok()) + .expect("legitimate logout should clear refresh cookie"); + assert!(clear_cookie.contains("Max-Age=0")); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_does_not_report_logout_success_when_session_revoke_is_rejected() { + let now = Utc::now(); + let user = sample_auth_user(now); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-auth-read-only")), + ]), + now + chrono::Duration::hours(1), + ); + let session = sample_auth_session( + "user-auth-1", + "session-auth-read-only", + "device-auth-read-only", + "refresh-token-read-only", + now, + ); + let repository = Arc::new( + InMemoryUserReadRepository::seed_auth_users(vec![user]) + .with_user_sessions([session]) + .read_only(), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_builder(|| { + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_user_reader_for_tests(repository)) + .without_auth_session_store_for_tests() + }) + .await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/auth/logout")) + .header("authorization", format!("Bearer {access_token}")) + .header("x-client-device-id", "device-auth-read-only") + .header("user-agent", "AetherTest/1.0") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["detail"], "服务暂不可用,请稍后重试"); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + #[tokio::test] async fn gateway_handles_payment_callback_route_locally_without_proxying_upstream() { - let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", "callback-secret-test"); + let _env_lock = payment_callback_env_lock() + .lock() + .expect("payment callback test env lock should not be poisoned"); + let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET); let now = Utc::now(); let user = StoredUserAuthRecord::new( "user-wallet-callback".to_string(), @@ -10431,7 +11946,7 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea .header("user-agent", "AetherTest/1.0") .json(&json!({ "amount_usd": 10.0, - "payment_method": "alipay", + "payment_method": "manual", "pay_amount": 72.5, "pay_currency": "CNY", "exchange_rate": 7.25, @@ -10464,10 +11979,10 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea "payload": serde_json::Value::Null, }); let callback_signature = - build_test_payment_callback_signature(&callback_body, "callback-secret-test"); + build_test_payment_callback_signature(&callback_body, TEST_PAYMENT_CALLBACK_SECRET); let callback_response = client - .post(format!("{gateway_url}/api/payment/callback/alipay")) - .header("x-payment-callback-token", "callback-secret-test") + .post(format!("{gateway_url}/api/payment/callback/manual")) + .header("x-payment-callback-token", TEST_PAYMENT_CALLBACK_SECRET) .header("x-payment-callback-signature", callback_signature) .json(&callback_body) .send() @@ -10478,13 +11993,7 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea .json() .await .expect("json body should parse"); - assert_eq!(callback_payload["ok"], true); - assert_eq!(callback_payload["credited"], true); - assert_eq!(callback_payload["payment_method"], "alipay"); - assert_eq!( - callback_payload["request_path"], - "/api/payment/callback/alipay" - ); + assert_eq!(callback_payload, json!({ "ok": true })); let detail_response = client .get(format!("{gateway_url}/api/wallet/recharge/{order_id}")) @@ -10509,7 +12018,10 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea #[tokio::test] async fn gateway_rejects_payment_callback_with_mismatched_payment_method_locally() { - let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", "callback-secret-test"); + let _env_lock = payment_callback_env_lock() + .lock() + .expect("payment callback test env lock should not be poisoned"); + let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET); let now = Utc::now(); let user = StoredUserAuthRecord::new( "user-wallet-callback-mismatch".to_string(), @@ -10615,10 +12127,10 @@ async fn gateway_rejects_payment_callback_with_mismatched_payment_method_locally "payload": serde_json::Value::Null, }); let callback_signature = - build_test_payment_callback_signature(&callback_body, "callback-secret-test"); + build_test_payment_callback_signature(&callback_body, TEST_PAYMENT_CALLBACK_SECRET); let callback_response = client .post(format!("{gateway_url}/api/payment/callback/wechat")) - .header("x-payment-callback-token", "callback-secret-test") + .header("x-payment-callback-token", TEST_PAYMENT_CALLBACK_SECRET) .header("x-payment-callback-signature", callback_signature) .json(&callback_body) .send() @@ -10630,8 +12142,11 @@ async fn gateway_rejects_payment_callback_with_mismatched_payment_method_locally .await .expect("json body should parse"); assert_eq!(callback_payload["ok"], false); - assert_eq!(callback_payload["error"], "payment method mismatch"); - assert_eq!(callback_payload["payment_method"], "wechat"); + assert_eq!(callback_payload["error"], "payment callback rejected"); + assert_eq!( + callback_payload.as_object().map(|value| value.len()), + Some(2) + ); let detail_response = client .get(format!("{gateway_url}/api/wallet/recharge/{order_id}")) diff --git a/apps/aether-gateway/src/tests/frontdoor/public_support/vscodex.rs b/apps/aether-gateway/src/tests/frontdoor/public_support/vscodex.rs new file mode 100644 index 000000000..b6e1d0aa2 --- /dev/null +++ b/apps/aether-gateway/src/tests/frontdoor/public_support/vscodex.rs @@ -0,0 +1,581 @@ +use super::{ + any, build_router_with_state, build_test_auth_token, json, sample_auth_session, + sample_auth_user, sample_auth_wallet, set_test_env_var, start_auth_gateway_with_state, + start_server, AppState, Arc, Json, Mutex, Request, Router, StatusCode, Utc, +}; +use axum::extract::ws::{Message as AxumWsMessage, WebSocketUpgrade}; +use axum::response::IntoResponse; +use futures_util::{SinkExt, StreamExt}; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::Message as TungsteniteMessage; + +#[derive(Debug, Clone, PartialEq)] +struct CapturedSidecarRequest { + method: http::Method, + path: String, + authorization: Option, + client_ip: Option, + body: Option, +} + +#[test] +fn gateway_authenticates_and_proxies_vscodex_bff_routes() { + std::thread::Builder::new() + .name("vscodex-gateway-test".to_string()) + .stack_size(32 * 1024 * 1024) + .spawn(|| { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .thread_stack_size(32 * 1024 * 1024) + .build() + .expect("test runtime should build") + .block_on(run_vscodex_gateway_integration()); + }) + .expect("test thread should spawn") + .join() + .expect("test thread should complete"); +} + +async fn run_vscodex_gateway_integration() { + let captured_requests = Arc::new(Mutex::new(Vec::::new())); + let captured_requests_for_handler = Arc::clone(&captured_requests); + let captured_ws_handshake = Arc::new(Mutex::new(None::<(Option, Option)>)); + let captured_ws_handshake_for_handler = Arc::clone(&captured_ws_handshake); + let sidecar = Router::new() + .route( + "/api/vscodex/ws", + any(move |ws: WebSocketUpgrade, headers: http::HeaderMap| { + let captured_ws_handshake = Arc::clone(&captured_ws_handshake_for_handler); + async move { + let origin = headers + .get(http::header::ORIGIN) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + let authorization = headers + .get(http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + *captured_ws_handshake + .lock() + .expect("WebSocket handshake store should lock") = + Some((origin, authorization)); + ws.protocols(["vscodex.v1"]) + .on_upgrade(|mut socket| async move { + while let Some(Ok(message)) = socket.next().await { + match message { + AxumWsMessage::Text(text) => { + let response = if text + == r#"{"type":"auth","token":"test-auth-ok"}"# { + r#"{"type":"auth.ok","role":"operator"}"#.to_string() + } else { + format!("echo:{text}") + }; + if socket + .send(AxumWsMessage::Text(response.into())) + .await + .is_err() + { + break; + } + } + AxumWsMessage::Binary(bytes) => { + if socket.send(AxumWsMessage::Binary(bytes)).await.is_err() + { + break; + } + } + AxumWsMessage::Close(frame) => { + let _ = socket.send(AxumWsMessage::Close(frame)).await; + break; + } + AxumWsMessage::Ping(_) | AxumWsMessage::Pong(_) => {} + } + } + }) + } + }), + ) + .route( + "/{*path}", + any(move |request: Request| { + let captured_requests = Arc::clone(&captured_requests_for_handler); + async move { + let method = request.method().clone(); + let path = request.uri().path().to_string(); + let authorization = request + .headers() + .get(http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + let client_ip = request + .headers() + .get("x-aether-client-ip") + .and_then(|value| value.to_str().ok()) + .map(ToOwned::to_owned); + let body = axum::body::to_bytes(request.into_body(), 1024 * 1024) + .await + .expect("sidecar request body should be readable"); + let body = (!body.is_empty()).then(|| { + serde_json::from_slice(&body).expect("sidecar request body should be JSON") + }); + captured_requests + .lock() + .expect("captured request store should lock") + .push(CapturedSidecarRequest { + method: method.clone(), + path: path.clone(), + authorization, + client_ip, + body, + }); + + let (status, payload) = match (method, path.as_str()) { + (http::Method::GET, "/internal/v1/users/user-auth-1/devices") => ( + StatusCode::OK, + json!({ "devices": [{ "id": "host-1", "name": "My Mac" }] }), + ), + (http::Method::POST, "/internal/v1/users/user-auth-1/pairings") => { + (StatusCode::CREATED, json!({ "code": "PAIR-123" })) + } + (http::Method::POST, "/internal/v1/users/user-auth-1/ws-tickets") => ( + StatusCode::CREATED, + json!({ + "ticket": "ticket-123", + "ws_url": "wss://aether.example/vscodex/ws" + }), + ), + (http::Method::DELETE, "/internal/v1/users/user-auth-1/devices/host-1") => { + (StatusCode::NO_CONTENT, json!({})) + } + ( + http::Method::DELETE, + "/internal/v1/users/user-auth-1/devices/missing", + ) => (StatusCode::NOT_FOUND, json!({ "detail": "设备不存在" })), + ( + http::Method::DELETE, + "/internal/v1/users/user-auth-1/devices/internal-denied", + ) => ( + StatusCode::UNAUTHORIZED, + json!({ "detail": "internal token invalid" }), + ), + ( + http::Method::DELETE, + "/internal/v1/users/user-auth-1/devices/redirect", + ) => (StatusCode::TEMPORARY_REDIRECT, json!({ "redirect": true })), + ( + http::Method::DELETE, + "/internal/v1/users/user-auth-1/devices/empty-ok", + ) => return StatusCode::OK.into_response(), + (http::Method::POST, "/v1/pairings/exchange") => ( + StatusCode::CREATED, + json!({ "device_id": "host-2", "device_token": "host-secret" }), + ), + _ => ( + StatusCode::NOT_FOUND, + json!({ "detail": "unexpected path" }), + ), + }; + let mut response = (status, Json(payload)).into_response(); + if status == StatusCode::TEMPORARY_REDIRECT { + response.headers_mut().insert( + http::header::LOCATION, + "/redirect-must-not-be-followed".parse().unwrap(), + ); + } + response + } + }), + ); + let (sidecar_url, sidecar_handle) = start_server(sidecar).await; + let _enabled = set_test_env_var("AETHER_VSCODEX_ENABLED", "true"); + let _internal_url = set_test_env_var("AETHER_VSCODEX_INTERNAL_URL", &sidecar_url); + let _internal_token = set_test_env_var("AETHER_VSCODEX_INTERNAL_TOKEN", "sidecar-secret"); + + let now = Utc::now(); + let user = sample_auth_user(now); + let access_token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".to_string(), json!(user.id)), + ("role".to_string(), json!(user.role)), + ( + "created_at".to_string(), + json!(user.created_at.map(|value| value.to_rfc3339())), + ), + ("session_id".to_string(), json!("session-vscodex")), + ]), + now + chrono::Duration::hours(1), + ); + let (gateway_url, upstream_hits, gateway_handle, upstream_handle) = + start_auth_gateway_with_state( + user, + sample_auth_wallet("user-auth-1", now), + [sample_auth_session( + "user-auth-1", + "session-vscodex", + "browser-device-vscodex", + "refresh-vscodex", + now, + )], + ) + .await; + let client = reqwest::Client::new(); + + let unauthenticated = client + .get(format!("{gateway_url}/api/users/me/vscodex/devices")) + .send() + .await + .expect("unauthenticated request should complete"); + assert_eq!(unauthenticated.status(), StatusCode::UNAUTHORIZED); + assert!(captured_requests + .lock() + .expect("captured request store should lock") + .is_empty()); + + let devices = client + .get(format!("{gateway_url}/api/users/me/vscodex/devices")) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("devices request should complete"); + assert_eq!(devices.status(), StatusCode::OK); + let devices_payload: serde_json::Value = + devices.json().await.expect("devices body should be JSON"); + assert_eq!(devices_payload["devices"][0]["id"], "host-1"); + + let pairing = client + .post(format!("{gateway_url}/api/users/me/vscodex/pairings")) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .json(&json!({ "name": "My Mac", "user_id": "attacker" })) + .send() + .await + .expect("pairing request should complete"); + assert_eq!(pairing.status(), StatusCode::CREATED); + assert_eq!( + pairing + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + let pairing_payload: serde_json::Value = + pairing.json().await.expect("pairing body should be JSON"); + assert_eq!(pairing_payload["code"], "PAIR-123"); + + let ticket = client + .post(format!("{gateway_url}/api/users/me/vscodex/ws-tickets")) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .json(&json!({ "device_id": "host-1", "user_id": "attacker" })) + .send() + .await + .expect("ticket request should complete"); + assert_eq!(ticket.status(), StatusCode::CREATED); + assert_eq!( + ticket + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + let ticket_payload: serde_json::Value = + ticket.json().await.expect("ticket body should be JSON"); + assert_eq!(ticket_payload["ticket"], "ticket-123"); + assert_eq!(ticket_payload["ws_url"], "wss://aether.example/vscodex/ws"); + + let deleted = client + .delete(format!("{gateway_url}/api/users/me/vscodex/devices/host-1")) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("delete request should complete"); + assert_eq!(deleted.status(), StatusCode::NO_CONTENT); + assert_eq!( + deleted + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + + let missing = client + .delete(format!( + "{gateway_url}/api/users/me/vscodex/devices/missing" + )) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("missing device request should complete"); + assert_eq!(missing.status(), StatusCode::NOT_FOUND); + let missing_payload: serde_json::Value = + missing.json().await.expect("missing body should be JSON"); + assert_eq!(missing_payload["detail"], "设备不存在"); + + let internal_denied = client + .delete(format!( + "{gateway_url}/api/users/me/vscodex/devices/internal-denied" + )) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("internal auth failure request should complete"); + assert_eq!(internal_denied.status(), StatusCode::BAD_GATEWAY); + let internal_denied_payload: serde_json::Value = internal_denied + .json() + .await + .expect("internal auth failure body should be JSON"); + assert_eq!( + internal_denied_payload["detail"], + "服务暂不可用,请稍后重试" + ); + + let redirected = client + .delete(format!( + "{gateway_url}/api/users/me/vscodex/devices/redirect" + )) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("redirecting sidecar request should complete"); + assert_eq!(redirected.status(), StatusCode::BAD_GATEWAY); + + let empty_success = client + .delete(format!( + "{gateway_url}/api/users/me/vscodex/devices/empty-ok" + )) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("empty sidecar success should complete"); + assert_eq!(empty_success.status(), StatusCode::BAD_GATEWAY); + + let pairing_exchange = client + .post(format!("{gateway_url}/api/vscodex/pair")) + .header("x-aether-client-ip", "203.0.113.99") + .json(&json!({ + "code": "PAIR-123", + "name": "Office Mac", + "user_id": "attacker", + "device_token": "stolen" + })) + .send() + .await + .expect("public pairing exchange should complete"); + assert_eq!(pairing_exchange.status(), StatusCode::CREATED); + assert_eq!( + pairing_exchange + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + let pairing_exchange_payload: serde_json::Value = pairing_exchange + .json() + .await + .expect("pairing exchange body should be JSON"); + assert_eq!(pairing_exchange_payload["device_id"], "host-2"); + assert_eq!(pairing_exchange_payload["device_token"], "host-secret"); + + let mut websocket_request = format!("{gateway_url}/api/vscodex/ws") + .replace("http://", "ws://") + .into_client_request() + .expect("WebSocket request should build"); + websocket_request.headers_mut().insert( + http::header::ORIGIN, + "https://aether.example".parse().unwrap(), + ); + websocket_request.headers_mut().insert( + http::header::AUTHORIZATION, + "Bearer browser-aether-jwt".parse().unwrap(), + ); + websocket_request.headers_mut().insert( + http::header::SEC_WEBSOCKET_PROTOCOL, + "vscodex.v1".parse().unwrap(), + ); + let (mut websocket, websocket_response) = tokio_tungstenite::connect_async(websocket_request) + .await + .expect("gateway WebSocket should connect"); + assert_eq!( + websocket_response + .headers() + .get(http::header::SEC_WEBSOCKET_PROTOCOL) + .and_then(|value| value.to_str().ok()), + Some("vscodex.v1") + ); + websocket + .send(TungsteniteMessage::Text( + "{\"type\":\"auth\",\"ticket\":\"one-time-ticket\"}".into(), + )) + .await + .expect("ticket frame should send"); + let echoed = websocket + .next() + .await + .expect("echoed frame should arrive") + .expect("echoed frame should be valid"); + assert_eq!( + echoed, + TungsteniteMessage::Text("echo:{\"type\":\"auth\",\"ticket\":\"one-time-ticket\"}".into()) + ); + websocket.close(None).await.expect("WebSocket should close"); + assert_eq!( + captured_ws_handshake + .lock() + .expect("WebSocket handshake store should lock") + .clone(), + Some((Some("https://aether.example".to_string()), None)) + ); + + let limited_gateway = build_router_with_state( + AppState::new() + .expect("limited gateway state should build") + .with_request_concurrency_limit(1), + ); + let (limited_gateway_url, limited_gateway_handle) = start_server(limited_gateway).await; + let limited_ws_url = + format!("{limited_gateway_url}/api/vscodex/ws").replace("http://", "ws://"); + let limited_ws_request = || { + let mut request = limited_ws_url + .as_str() + .into_client_request() + .expect("limited WebSocket request should build"); + request + .headers_mut() + .insert("x-real-ip", "198.51.100.50".parse().unwrap()); + request + }; + let mut held_websockets = Vec::new(); + for index in 0..16 { + let (websocket, _) = tokio_tungstenite::connect_async(limited_ws_request()) + .await + .unwrap_or_else(|err| panic!("limited WebSocket {index} should connect: {err}")); + held_websockets.push(websocket); + } + let per_ip_limit_error = tokio_tungstenite::connect_async(limited_ws_request()) + .await + .expect_err("seventeenth WebSocket from one IP should be rejected"); + match per_ip_limit_error { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!( + response + .headers() + .get(http::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()), + Some("1") + ); + } + other => panic!("expected HTTP per-IP limit rejection, got {other:?}"), + } + + held_websockets[0] + .send(TungsteniteMessage::Text( + r#"{"type":"auth","token":"test-auth-ok"}"#.into(), + )) + .await + .expect("test authentication frame should send"); + let auth_ok = + tokio::time::timeout(std::time::Duration::from_secs(1), held_websockets[0].next()) + .await + .expect("test authentication response should arrive in time") + .expect("test authentication response should contain a frame") + .expect("test authentication response should be valid"); + assert_eq!( + auth_ok, + TungsteniteMessage::Text(r#"{"type":"auth.ok","role":"operator"}"#.into()) + ); + tokio::time::sleep(std::time::Duration::from_millis(25)).await; + + let (mut replacement_websocket, _) = tokio_tungstenite::connect_async(limited_ws_request()) + .await + .expect("sidecar auth success should release one pending per-IP slot"); + replacement_websocket + .close(None) + .await + .expect("replacement WebSocket should close"); + for mut websocket in held_websockets { + websocket + .close(None) + .await + .expect("held WebSocket should close"); + } + limited_gateway_handle.abort(); + + let blacklisted_gateway = build_router_with_state( + AppState::new() + .expect("blacklisted gateway state should build") + .with_admin_security_blacklist_for_tests([( + "127.0.0.1".to_string(), + "blocked".to_string(), + )]), + ); + let (blacklisted_gateway_url, blacklisted_gateway_handle) = + start_server(blacklisted_gateway).await; + let blacklisted_ws_url = + format!("{blacklisted_gateway_url}/api/vscodex/ws").replace("http://", "ws://"); + let blacklisted_error = tokio_tungstenite::connect_async(&blacklisted_ws_url) + .await + .expect_err("blacklisted WebSocket should be rejected"); + match blacklisted_error { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert_eq!(response.status(), StatusCode::FORBIDDEN) + } + other => panic!("expected HTTP blacklist rejection, got {other:?}"), + } + blacklisted_gateway_handle.abort(); + + let requests = captured_requests + .lock() + .expect("captured request store should lock") + .clone(); + assert_eq!(requests.len(), 9); + assert!(requests + .iter() + .all(|request| request.authorization.as_deref() == Some("Bearer sidecar-secret"))); + assert!(requests[..8] + .iter() + .all(|request| request.path.starts_with("/internal/v1/users/user-auth-1/"))); + assert_eq!(requests[8].path, "/v1/pairings/exchange"); + assert_eq!(requests[1].body, Some(json!({ "name": "My Mac" }))); + assert_eq!(requests[2].body, Some(json!({ "device_id": "host-1" }))); + assert_eq!( + requests[8].body, + Some(json!({ "code": "PAIR-123", "name": "Office Mac" })) + ); + assert_eq!(requests[8].client_ip.as_deref(), Some("127.0.0.1")); + assert!(requests[..8] + .iter() + .all(|request| request.client_ip.is_none())); + + let _disabled = set_test_env_var("AETHER_VSCODEX_ENABLED", "false"); + let disabled = client + .get(format!("{gateway_url}/api/users/me/vscodex/devices")) + .bearer_auth(&access_token) + .header("x-client-device-id", "browser-device-vscodex") + .send() + .await + .expect("disabled feature request should complete"); + assert_eq!(disabled.status(), StatusCode::SERVICE_UNAVAILABLE); + let disabled_payload: serde_json::Value = + disabled.json().await.expect("disabled body should be JSON"); + assert_eq!(disabled_payload["detail"], "服务暂不可用,请稍后重试"); + assert_eq!( + captured_requests + .lock() + .expect("captured request store should lock") + .len(), + 9 + ); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); + sidecar_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/mod.rs b/apps/aether-gateway/src/tests/mod.rs index 9d13640fa..15d90bcf9 100644 --- a/apps/aether-gateway/src/tests/mod.rs +++ b/apps/aether-gateway/src/tests/mod.rs @@ -17,6 +17,7 @@ mod concurrency; mod control; mod files; mod frontdoor; +mod operational_auth; mod proxy; mod usage; mod video; @@ -45,6 +46,44 @@ pub(super) async fn start_server(app: Router) -> (String, tokio::task::JoinHandl (format!("http://{addr}"), handle) } +pub(super) const OPERATIONAL_ADMIN_DEVICE_ID: &str = "device-operational-admin"; + +pub(super) async fn start_authenticated_operational_server( + state: AppState, +) -> (String, tokio::task::JoinHandle<()>, String) { + let access_token = + control::issue_shared_test_admin_access_token(&state, OPERATIONAL_ADMIN_DEVICE_ID).await; + let (url, handle) = start_server(build_router_with_state(state)).await; + (url, handle, access_token) +} + +pub(super) fn authenticated_operational_client(access_token: &str) -> reqwest::Client { + authenticated_operational_client_with_builder(reqwest::Client::builder(), access_token) +} + +pub(super) fn authenticated_operational_client_with_builder( + builder: reqwest::ClientBuilder, + access_token: &str, +) -> reqwest::Client { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert( + reqwest::header::AUTHORIZATION, + format!("Bearer {access_token}") + .parse() + .expect("operational authorization header should build"), + ); + headers.insert( + "x-client-device-id", + OPERATIONAL_ADMIN_DEVICE_ID + .parse() + .expect("operational device header should build"), + ); + builder + .default_headers(headers) + .build() + .expect("operational client should build") +} + pub(super) async fn send_request(app: Router, mut request: Request) -> Response { use tower::ServiceExt; diff --git a/apps/aether-gateway/src/tests/operational_auth.rs b/apps/aether-gateway/src/tests/operational_auth.rs new file mode 100644 index 000000000..d2ba6270a --- /dev/null +++ b/apps/aether-gateway/src/tests/operational_auth.rs @@ -0,0 +1,717 @@ +use std::sync::Arc; + +use aether_data::repository::management_tokens::{ + InMemoryManagementTokenRepository, StoredManagementToken, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, +}; +use base64::Engine as _; +use hmac::Mac; +use sha2::{Digest, Sha256}; + +use super::{ + build_router_with_state, send_request, start_server, AppState, Body, Request, StatusCode, + OPERATIONAL_ADMIN_DEVICE_ID, +}; +use crate::data::GatewayDataState; + +fn hash_token(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} + +async fn issue_operational_session_access_token( + state: &AppState, + client_device_id: &str, + role: &str, +) -> String { + issue_operational_session_access_token_and_user(state, client_device_id, role) + .await + .0 +} + +async fn issue_operational_session_access_token_and_user( + state: &AppState, + client_device_id: &str, + role: &str, +) -> (String, aether_data::repository::users::StoredUserAuthRecord) { + let user = state + .create_local_auth_user_with_settings( + Some(format!("operational-{role}@example.com")), + true, + format!("operational_{role}"), + "hash".to_string(), + role.to_string(), + None, + None, + None, + None, + ) + .await + .expect("operational user should be created") + .expect("operational user should exist"); + let now = chrono::Utc::now(); + let session_id = format!("session-operational-{role}"); + let refresh_token = format!("refresh-{session_id}"); + let session = crate::data::state::StoredUserSessionRecord::new( + session_id.clone(), + user.id.clone(), + client_device_id.to_string(), + None, + crate::data::state::StoredUserSessionRecord::hash_refresh_token(&refresh_token), + None, + None, + Some(now), + Some(now + chrono::Duration::days(7)), + None, + None, + Some("127.0.0.1".to_string()), + Some("operational-auth-test".to_string()), + Some(now), + Some(now), + ) + .expect("session should build"); + state + .create_user_session(session) + .await + .expect("session should persist") + .expect("session should exist"); + + let header = serde_json::json!({ "alg": "HS256", "typ": "JWT" }); + let payload = serde_json::json!({ + "user_id": user.id, + "role": role, + "created_at": user.created_at.map(|value| value.to_rfc3339()), + "session_id": session_id, + "exp": (now + chrono::Duration::hours(12)).timestamp(), + "type": "access", + }); + let header_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode( + serde_json::to_vec(&header) + .expect("jwt header should serialize") + .as_slice(), + ); + let payload_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode( + serde_json::to_vec(&payload) + .expect("jwt payload should serialize") + .as_slice(), + ); + let signing_input = format!("{header_segment}.{payload_segment}"); + let secret = std::env::var("JWT_SECRET_KEY") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string()); + let mut mac = + hmac::Hmac::::new_from_slice(secret.as_bytes()).expect("jwt secret should build"); + mac.update(signing_input.as_bytes()); + let signature = mac.finalize().into_bytes(); + ( + format!( + "{signing_input}.{}", + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(signature.as_slice()) + ), + user, + ) +} + +fn management_token( + token_id: &str, + user_id: &str, + permissions: &[&str], +) -> StoredManagementTokenWithUser { + let token = StoredManagementToken::new( + token_id.to_string(), + user_id.to_string(), + token_id.to_string(), + ) + .expect("management token should build") + .with_permissions(Some(serde_json::json!(permissions))) + .with_runtime_fields(Some(4_102_444_800), None, None, 0, true); + let user = StoredManagementTokenUserSummary::new( + user_id.to_string(), + Some("operational-admin@example.com".to_string()), + "operational_admin".to_string(), + "admin".to_string(), + ) + .expect("management token user should build"); + StoredManagementTokenWithUser::new(token, user) +} + +async fn state_with_management_tokens( + tokens: Vec<(&'static str, StoredManagementTokenWithUser)>, +) -> AppState { + state_with_management_tokens_for_role(tokens, "admin").await +} + +async fn state_with_management_tokens_for_role( + tokens: Vec<(&'static str, StoredManagementTokenWithUser)>, + role: &str, +) -> AppState { + let state = AppState::new().expect("gateway state should build"); + let user = state + .create_local_auth_user_with_settings( + Some("operational-admin@example.com".to_string()), + true, + "operational_admin".to_string(), + "hash".to_string(), + role.to_string(), + None, + None, + None, + None, + ) + .await + .expect("admin user should be created") + .expect("admin user should exist"); + let items = tokens + .iter() + .map(|(_, token)| { + let mut token = token.clone(); + token.token.user_id = user.id.clone(); + token.user.id = user.id.clone(); + token + }) + .collect::>(); + let hashes = tokens + .into_iter() + .zip(items.iter()) + .map(|((raw, _), token)| (hash_token(raw), token.token.id.clone())) + .collect::>(); + let repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + items, hashes, + )); + state.with_data_state_for_tests( + GatewayDataState::with_management_token_repository_for_tests(repository), + ) +} + +#[tokio::test] +async fn operational_session_requires_matching_device_id() { + let state = AppState::new().expect("gateway state should build"); + let access_token = + super::control::issue_shared_test_admin_access_token(&state, OPERATIONAL_ADMIN_DEVICE_ID) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(access_token) + .header("x-client-device-id", "different-device") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + gateway_handle.abort(); +} + +#[tokio::test] +async fn operational_session_rejects_a_stale_user_security_version() { + let state = AppState::new().expect("gateway state should build"); + let (access_token, user) = issue_operational_session_access_token_and_user( + &state, + OPERATIONAL_ADMIN_DEVICE_ID, + "admin", + ) + .await; + let user = user + .with_security_version(1) + .expect("admin security version should update"); + let state = state.with_auth_users_for_tests([user]); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(access_token) + .header("x-client-device-id", OPERATIONAL_ADMIN_DEVICE_ID) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + gateway_handle.abort(); +} + +#[tokio::test] +async fn operational_session_rejects_a_replaced_user_identity() { + let state = AppState::new().expect("gateway state should build"); + let (access_token, mut user) = issue_operational_session_access_token_and_user( + &state, + OPERATIONAL_ADMIN_DEVICE_ID, + "admin", + ) + .await; + user.created_at = user + .created_at + .map(|created_at| created_at + chrono::Duration::days(1)); + let state = state.with_auth_users_for_tests([user]); + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(access_token) + .header("x-client-device-id", OPERATIONAL_ADMIN_DEVICE_ID) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + gateway_handle.abort(); +} + +#[tokio::test] +async fn operational_routes_reject_duplicate_authorization_headers() { + let raw_token = "ae-operational-duplicate-authorization"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-duplicate-authorization", + "placeholder-user", + &["admin:monitoring:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let mut request = Request::builder() + .uri("/_gateway/metrics") + .body(Body::empty()) + .expect("request should build"); + request.headers_mut().append( + http::header::AUTHORIZATION, + format!("Bearer {raw_token}") + .parse() + .expect("authorization should parse"), + ); + request.headers_mut().append( + http::header::AUTHORIZATION, + "Bearer ae-attacker-selected-token" + .parse() + .expect("authorization should parse"), + ); + + let response = send_request(gateway, request).await; + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!( + response + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); +} + +#[tokio::test] +async fn monitoring_management_token_cannot_read_video_tasks() { + let raw_token = "ae-operational-monitoring-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-monitoring-read", + "placeholder-user", + &["admin:monitoring:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let metrics = client + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(raw_token) + .send() + .await + .expect("metrics request should succeed"); + assert_eq!(metrics.status(), StatusCode::OK); + assert_eq!( + metrics + .headers() + .get(http::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()), + Some("no-store") + ); + + let video_tasks = client + .get(format!("{gateway_url}/_gateway/async-tasks/video-tasks")) + .bearer_auth(raw_token) + .send() + .await + .expect("video task request should succeed"); + assert_eq!(video_tasks.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = video_tasks.json().await.expect("response should parse"); + assert_eq!(payload["required_permission"], "admin:video_tasks:read"); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn monitoring_read_management_token_cannot_read_request_candidate_traces() { + let raw_token = "ae-operational-candidate-monitoring-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-candidate-monitoring-read", + "placeholder-user", + &["admin:monitoring:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for path in [ + "/_gateway/audit/request-candidates/not-found", + "/_gateway/audit/decision-trace/not-found", + ] { + let response = client + .get(format!("{gateway_url}{path}")) + .bearer_auth(raw_token) + .send() + .await + .expect("candidate trace request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!( + payload["required_permission"], "admin:monitoring:admin", + "path: {path}" + ); + } + + gateway_handle.abort(); +} + +#[tokio::test] +async fn monitoring_admin_management_token_may_reach_request_candidate_trace_handlers() { + let raw_token = "ae-operational-candidate-monitoring-admin"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-candidate-monitoring-admin", + "placeholder-user", + &["admin:monitoring:admin"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for path in [ + "/_gateway/audit/request-candidates/not-found", + "/_gateway/audit/decision-trace/not-found", + ] { + let response = client + .get(format!("{gateway_url}{path}")) + .bearer_auth(raw_token) + .send() + .await + .expect("candidate trace request should succeed"); + + assert_eq!(response.status(), StatusCode::NOT_FOUND, "path: {path}"); + } + + gateway_handle.abort(); +} + +#[tokio::test] +async fn usage_management_token_cannot_read_audit_bundle_candidate_trace() { + let raw_token = "ae-operational-usage-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-usage-read", + "placeholder-user", + &["admin:usage:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/_gateway/audit/request-audit/not-found" + )) + .bearer_auth(raw_token) + .send() + .await + .expect("audit bundle request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!(payload["required_permission"], "admin:monitoring:admin"); + gateway_handle.abort(); +} + +#[tokio::test] +async fn usage_and_api_key_management_token_cannot_read_audit_bundle_candidate_trace() { + let raw_token = "ae-operational-audit-bundle-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-audit-bundle-read", + "placeholder-user", + &["admin:usage:read", "admin:api_keys:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/_gateway/audit/request-audit/not-found" + )) + .bearer_auth(raw_token) + .send() + .await + .expect("audit bundle request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!(payload["required_permission"], "admin:monitoring:admin"); + gateway_handle.abort(); +} + +#[tokio::test] +async fn three_permission_management_token_may_reach_audit_bundle_handler() { + let raw_token = "ae-operational-audit-bundle-full-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-audit-bundle-full-read", + "placeholder-user", + &[ + "admin:monitoring:admin", + "admin:usage:read", + "admin:api_keys:read", + ], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!( + "{gateway_url}/_gateway/audit/request-audit/not-found" + )) + .bearer_auth(raw_token) + .send() + .await + .expect("audit bundle request should succeed"); + + assert_eq!(response.status(), StatusCode::NOT_FOUND); + gateway_handle.abort(); +} + +#[tokio::test] +async fn downgraded_audit_admin_management_token_cannot_read_full_admin_audit_routes() { + let raw_token = "ae-operational-downgraded-audit-admin"; + let state = state_with_management_tokens_for_role( + vec![( + raw_token, + management_token( + "operational-downgraded-audit-admin", + "placeholder-user", + &[ + "admin:monitoring:admin", + "admin:usage:read", + "admin:api_keys:read", + ], + ), + )], + "audit_admin", + ) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + for path in [ + "/_gateway/audit/request-candidates/not-found", + "/_gateway/audit/decision-trace/not-found", + "/_gateway/audit/request-audit/not-found", + ] { + let response = client + .get(format!("{gateway_url}{path}")) + .bearer_auth(raw_token) + .send() + .await + .expect("full-admin audit request should complete"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!( + payload["required_permission"], "admin:monitoring:admin", + "path: {path}" + ); + } + + gateway_handle.abort(); +} + +#[tokio::test] +async fn audit_admin_session_cannot_read_candidate_audit_routes() { + let state = AppState::new().expect("gateway state should build"); + let access_token = + issue_operational_session_access_token(&state, OPERATIONAL_ADMIN_DEVICE_ID, "audit_admin") + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let usage = client + .get(format!( + "{gateway_url}/_gateway/audit/request-usage/not-found" + )) + .bearer_auth(&access_token) + .header("x-client-device-id", OPERATIONAL_ADMIN_DEVICE_ID) + .send() + .await + .expect("usage audit request should succeed"); + assert_eq!(usage.status(), StatusCode::NOT_FOUND); + + for path in [ + "/_gateway/audit/request-candidates/not-found", + "/_gateway/audit/decision-trace/not-found", + "/_gateway/audit/request-audit/not-found", + ] { + let response = client + .get(format!("{gateway_url}{path}")) + .bearer_auth(&access_token) + .header("x-client-device-id", OPERATIONAL_ADMIN_DEVICE_ID) + .send() + .await + .expect("candidate audit request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN, "path: {path}"); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!( + payload["required_permission"], "admin:monitoring:admin", + "path: {path}" + ); + } + + gateway_handle.abort(); +} + +#[tokio::test] +async fn malformed_management_token_ip_rules_fail_closed() { + let raw_token = "ae-operational-malformed-ip-rules"; + let mut stored = management_token( + "operational-malformed-ip-rules", + "placeholder-user", + &["admin:monitoring:read"], + ); + stored.token.allowed_ips = Some(serde_json::json!(["!not-an-ip"])); + let state = state_with_management_tokens(vec![(raw_token, stored)]).await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(raw_token) + .send() + .await + .expect("metrics request should succeed"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + gateway_handle.abort(); +} + +#[tokio::test] +async fn json_null_management_token_permissions_fail_closed() { + let raw_token = "ae-operational-json-null-permissions"; + let mut stored = management_token( + "operational-json-null-permissions", + "placeholder-user", + &["admin:monitoring:read"], + ); + stored.token.permissions = Some(serde_json::Value::Null); + let state = state_with_management_tokens(vec![(raw_token, stored)]).await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/_gateway/metrics")) + .bearer_auth(raw_token) + .send() + .await + .expect("metrics request should complete"); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + gateway_handle.abort(); +} + +#[tokio::test] +async fn video_read_management_token_cannot_cancel_tasks() { + let raw_token = "ae-operational-video-read"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-video-read", + "placeholder-user", + &["admin:video_tasks:read"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!( + "{gateway_url}/_gateway/async-tasks/video-tasks/not-found/cancel" + )) + .bearer_auth(raw_token) + .send() + .await + .expect("cancel request should succeed"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + let payload: serde_json::Value = response.json().await.expect("response should parse"); + assert_eq!(payload["required_permission"], "admin:video_tasks:write"); + gateway_handle.abort(); +} + +#[tokio::test] +async fn video_admin_management_token_may_reach_cancel_handler() { + let raw_token = "ae-operational-video-admin"; + let state = state_with_management_tokens(vec![( + raw_token, + management_token( + "operational-video-admin", + "placeholder-user", + &["admin:video_tasks:admin"], + ), + )]) + .await; + let gateway = build_router_with_state(state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!( + "{gateway_url}/_gateway/async-tasks/video-tasks/not-found/cancel" + )) + .bearer_auth(raw_token) + .send() + .await + .expect("cancel request should succeed"); + + assert_eq!(response.status(), StatusCode::NOT_FOUND); + gateway_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/proxy.rs b/apps/aether-gateway/src/tests/proxy.rs index 81c9e6dda..281d69a60 100644 --- a/apps/aether-gateway/src/tests/proxy.rs +++ b/apps/aether-gateway/src/tests/proxy.rs @@ -1,31 +1,62 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH}; use super::{ - any, build_router, build_router_with_state, json, start_server, to_bytes, AppState, Arc, Body, - HeaderValue, Json, Mutex, Request, Response, Router, StatusCode, DEPENDENCY_REASON_HEADER, - EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_LOOP_DETECTED, - EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND, EXECUTION_RUNTIME_LOOP_GUARD_HEADER, - EXECUTION_RUNTIME_LOOP_GUARD_VALUE, FORWARDED_FOR_HEADER, GATEWAY_HEADER, TRACE_ID_HEADER, - TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER, - TRUSTED_AUTH_USER_ID_HEADER, TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + any, build_router, build_router_with_state, json, send_request, start_server, to_bytes, + AppState, Arc, Body, HeaderValue, Json, Mutex, Request, Response, Router, StatusCode, + DEPENDENCY_REASON_HEADER, EXECUTION_PATH_HEADER, EXECUTION_PATH_LOCAL_AUTH_DENIED, + EXECUTION_PATH_LOCAL_EXECUTION_LOOP_DETECTED, EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND, + EXECUTION_RUNTIME_LOOP_GUARD_HEADER, EXECUTION_RUNTIME_LOOP_GUARD_VALUE, FORWARDED_FOR_HEADER, + GATEWAY_HEADER, TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + TRUSTED_AUTH_API_KEY_ID_HEADER, TRUSTED_AUTH_USER_ID_HEADER, + TUNNEL_AFFINITY_FORWARDED_BY_HEADER, TUNNEL_AFFINITY_NODE_ID_HEADER, TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, }; +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_FORWARDED_BY_HEADER, TUNNEL_RELAY_OWNER_INSTANCE_HEADER, +}; use aether_data::repository::auth::{ InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, }; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; +use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; +use aether_runtime_state::{RedisClientConfig, RuntimeState}; +use aether_scheduler_core::{ + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, + SchedulerAffinityScope, +}; +use aether_test_support::ManagedRedisServer; use sha2::{Digest, Sha256}; +const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes"; +const AFFINITY_TEST_SENDER: &str = "gateway-a"; +const AFFINITY_TEST_OWNER: &str = "gateway-b"; +const AFFINITY_TEST_NODE: &str = "node-owner"; + fn hash_api_key(value: &str) -> String { let mut hasher = Sha256::new(); hasher.update(value.as_bytes()); format!("{:x}", hasher.finalize()) } +fn system_default_affinity_cache_key(api_key_id: &str, api_format: &str, model: &str) -> String { + let scope = SchedulerAffinityScope::new("system-default", Some(1)); + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id, + api_format, + model, + None, + Some(&scope), + ) + .expect("system-default affinity cache key should build") +} + fn sample_auth_snapshot( api_key_id: &str, user_id: &str, @@ -190,6 +221,22 @@ fn sample_key(key_id: &str, provider_id: &str, node_id: &str) -> StoredProviderC .expect("key transport should build") } +fn sample_bound_key(key_id: &str, provider_id: &str, node_id: &str) -> StoredProviderCatalogKey { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY), + ); + let mut key = sample_key(key_id, provider_id, node_id); + key.encrypted_api_key = Some( + bootstrap + .seal_provider_catalog_key_api_key(provider_id, key_id, "plain-upstream-key") + .expect("bound provider api key ciphertext should build"), + ); + key +} + fn sample_codex_key(key_id: &str, provider_id: &str, node_id: &str) -> StoredProviderCatalogKey { StoredProviderCatalogKey::new( key_id.to_string(), @@ -218,10 +265,52 @@ fn sample_codex_key(key_id: &str, provider_id: &str, node_id: &str) -> StoredPro .expect("key transport should build") } +fn sample_bound_codex_key( + key_id: &str, + provider_id: &str, + node_id: &str, +) -> StoredProviderCatalogKey { + let bootstrap = AppState::new() + .expect("bootstrap state should build") + .with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY), + ); + let mut key = sample_codex_key(key_id, provider_id, node_id); + key.encrypted_api_key = Some( + bootstrap + .seal_provider_catalog_key_api_key(provider_id, key_id, "plain-upstream-key") + .expect("bound provider api key ciphertext should build"), + ); + key +} + fn tunnel_attachment_key(node_id: &str) -> String { format!("tunnel.attachments.{node_id}") } +fn sample_tunnel_proxy_node(node_id: &str, tunnel_generation: &str) -> StoredProxyNode { + StoredProxyNode::new( + node_id.to_string(), + format!("proxy-{node_id}"), + "127.0.0.1".to_string(), + 1, + false, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + true, + true, + 1, + ) + .expect("tunnel proxy node should build") + .with_tunnel_generation(tunnel_generation.to_string()) +} + fn current_unix_secs() -> u64 { SystemTime::now() .duration_since(UNIX_EPOCH) @@ -229,6 +318,333 @@ fn current_unix_secs() -> u64 { .as_secs() } +fn signed_affinity_headers( + method: &http::Method, + uri: &http::Uri, + nonce: &str, + user_id: &str, + api_key_id: &str, + body: &[u8], +) -> http::HeaderMap { + let mut headers = http::HeaderMap::new(); + headers.insert( + GATEWAY_HEADER, + HeaderValue::from_static("rust-phase3b-affinity"), + ); + headers.insert( + TUNNEL_RELAY_FORWARDED_BY_HEADER, + HeaderValue::from_static(AFFINITY_TEST_SENDER), + ); + headers.insert( + TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + HeaderValue::from_static(AFFINITY_TEST_SENDER), + ); + headers.insert( + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + HeaderValue::from_static(AFFINITY_TEST_OWNER), + ); + headers.insert( + TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + HeaderValue::from_static(AFFINITY_TEST_OWNER), + ); + headers.insert( + TUNNEL_AFFINITY_NODE_ID_HEADER, + HeaderValue::from_static(AFFINITY_TEST_NODE), + ); + headers.insert( + TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_str(user_id).expect("trusted user id should be a valid header"), + ); + headers.insert( + TRUSTED_AUTH_API_KEY_ID_HEADER, + HeaderValue::from_str(api_key_id).expect("trusted API key id should be a valid header"), + ); + headers.insert( + TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + HeaderValue::from_static("true"), + ); + headers.insert( + FORWARDED_FOR_HEADER, + HeaderValue::from_static("203.0.113.10"), + ); + + let timestamp = current_unix_secs(); + let metadata = crate::tunnel::build_tunnel_affinity_auth_metadata(method, uri, &headers) + .expect("affinity authentication metadata should build"); + let payload_digest = tunnel_relay_payload_digest(&metadata, body); + let signature = sign_tunnel_relay_request( + RELAY_TEST_SECRET.as_bytes(), + AFFINITY_TEST_SENDER, + AFFINITY_TEST_OWNER, + AFFINITY_TEST_NODE, + AFFINITY_TEST_SENDER, + false, + timestamp, + nonce, + &payload_digest, + ); + headers.insert( + TUNNEL_RELAY_AUTH_SENDER_HEADER, + HeaderValue::from_static(AFFINITY_TEST_SENDER), + ); + headers.insert( + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + HeaderValue::from_str(×tamp.to_string()).expect("timestamp should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_NONCE_HEADER, + HeaderValue::from_str(nonce).expect("nonce should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + HeaderValue::from_str(&payload_digest.encode_header_value()) + .expect("payload digest should be a valid header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + HeaderValue::from_str(&signature).expect("signature should be a valid header"), + ); + headers +} + +fn affinity_request(uri: http::Uri, headers: http::HeaderMap, body: &'static str) -> Request { + let mut request = Request::builder() + .method(http::Method::POST) + .uri(uri) + .body(Body::from(body)) + .expect("affinity request should build"); + *request.headers_mut() = headers; + request.headers_mut().insert( + http::header::CONTENT_TYPE, + HeaderValue::from_static("application/json"), + ); + request +} + +#[tokio::test] +async fn gateway_rejects_incomplete_tunnel_affinity_trusted_auth_headers() { + let gateway = build_router_with_state( + AppState::new() + .expect("gateway state should build") + .with_tunnel_identity_and_relay_secret_for_tests( + AFFINITY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), + ); + let uri: http::Uri = "/v1/chat/completions" + .parse() + .expect("affinity URI should parse"); + let mut headers = http::HeaderMap::new(); + headers.insert( + TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_static("forged-user"), + ); + + let response = send_request( + gateway, + affinity_request(uri, headers, r#"{"model":"gpt-5","messages":[]}"#), + ) + .await; + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response + .headers() + .get(EXECUTION_PATH_HEADER) + .and_then(|value| value.to_str().ok()), + Some(EXECUTION_PATH_LOCAL_AUTH_DENIED) + ); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body should read"), + ) + .expect("response body should parse"); + assert_eq!( + payload["error"]["message"], + "invalid tunnel affinity authentication" + ); +} + +#[tokio::test] +async fn gateway_uses_signed_tunnel_affinity_identity_once_and_rejects_replay() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + None, + sample_auth_snapshot("affinity-key", "affinity-user", "gpt-4.1"), + )])); + let data_state = + crate::data::GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository.clone()); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests(data_state) + .with_tunnel_identity_and_relay_secret_for_tests( + AFFINITY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), + ); + let uri: http::Uri = "/v1/chat/completions?stream=false" + .parse() + .expect("affinity URI should parse"); + let body = r#"{"model":"gpt-5","messages":[]}"#; + let headers = signed_affinity_headers( + &http::Method::POST, + &uri, + "affinity-valid-once", + "affinity-user", + "affinity-key", + body.as_bytes(), + ); + + let first = send_request( + gateway.clone(), + affinity_request(uri.clone(), headers.clone(), body), + ) + .await; + assert_eq!(first.status(), StatusCode::FORBIDDEN); + let first_payload: serde_json::Value = serde_json::from_slice( + &to_bytes(first.into_body(), usize::MAX) + .await + .expect("first response body should read"), + ) + .expect("first response body should parse"); + assert_eq!(first_payload["error"]["type"], "http_error"); + assert!(first_payload["error"]["message"] + .as_str() + .is_some_and(|message| message.contains("gpt-5"))); + assert_eq!(auth_repository.snapshot_lookup_count("affinity-key"), 1); + + let replay = send_request(gateway, affinity_request(uri, headers, body)).await; + assert_eq!(replay.status(), StatusCode::FORBIDDEN); + let replay_payload: serde_json::Value = serde_json::from_slice( + &to_bytes(replay.into_body(), usize::MAX) + .await + .expect("replay response body should read"), + ) + .expect("replay response body should parse"); + assert_eq!( + replay_payload["error"]["message"], + "invalid tunnel affinity authentication" + ); + assert_eq!(auth_repository.snapshot_lookup_count("affinity-key"), 1); +} + +#[tokio::test] +async fn gateway_rejects_tunnel_affinity_path_and_trusted_identity_tampering() { + let gateway = build_router_with_state( + AppState::new() + .expect("gateway state should build") + .with_tunnel_identity_and_relay_secret_for_tests( + AFFINITY_TEST_OWNER, + None, + RELAY_TEST_SECRET, + ), + ); + let signed_uri: http::Uri = "/v1/chat/completions?stream=false" + .parse() + .expect("signed affinity URI should parse"); + let tampered_uri: http::Uri = "/v1/chat/completions?stream=true" + .parse() + .expect("tampered affinity URI should parse"); + let body = r#"{"model":"gpt-5","messages":[]}"#; + + let path_headers = signed_affinity_headers( + &http::Method::POST, + &signed_uri, + "affinity-path-tamper", + "affinity-user", + "affinity-key", + body.as_bytes(), + ); + let path_response = send_request( + gateway.clone(), + affinity_request(tampered_uri, path_headers, body), + ) + .await; + assert_eq!(path_response.status(), StatusCode::FORBIDDEN); + + let mut identity_headers = signed_affinity_headers( + &http::Method::POST, + &signed_uri, + "affinity-identity-tamper", + "affinity-user", + "affinity-key", + body.as_bytes(), + ); + identity_headers.insert( + TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_static("forged-user"), + ); + let identity_response = send_request( + gateway, + affinity_request(signed_uri, identity_headers, body), + ) + .await; + assert_eq!(identity_response.status(), StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn gateway_fails_closed_when_shared_tunnel_affinity_nonce_state_is_unavailable() { + let mut redis = match ManagedRedisServer::start().await { + Ok(redis) => redis, + Err(error) if error.to_string().contains("No such file or directory") => { + eprintln!("skipping affinity Redis outage test: {error}"); + return; + } + Err(error) => panic!("Redis test server should start: {error}"), + }; + let runtime_state = Arc::new( + RuntimeState::redis( + RedisClientConfig { + url: redis.redis_url().to_string(), + key_prefix: Some(format!("affinity-outage-{}", std::process::id())), + }, + Some(250), + ) + .await + .expect("Redis runtime state should build"), + ); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway state should build") + .with_tunnel_identity_runtime_state_and_relay_secret_for_tests( + AFFINITY_TEST_OWNER, + None, + runtime_state, + RELAY_TEST_SECRET, + ), + ); + let uri: http::Uri = "/v1/chat/completions" + .parse() + .expect("affinity URI should parse"); + let body = r#"{"model":"gpt-5","messages":[]}"#; + let headers = signed_affinity_headers( + &http::Method::POST, + &uri, + "affinity-runtime-unavailable", + "affinity-user", + "affinity-key", + body.as_bytes(), + ); + redis.stop().expect("Redis test server should stop"); + + let response = send_request(gateway, affinity_request(uri, headers, body)).await; + + assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body should read"), + ) + .expect("response body should parse"); + assert_eq!( + payload["error"]["message"], + "tunnel affinity authentication is unavailable" + ); +} + #[tokio::test] async fn gateway_rejects_unknown_path_locally_and_generates_trace_id() { let upstream_hits = Arc::new(Mutex::new(0usize)); @@ -505,6 +921,13 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ forwarded_for: String, forwarded_by: String, owner_instance_id: String, + relay_sender: String, + relay_owner: String, + relay_timestamp: String, + relay_nonce: String, + relay_signature: String, + cookie: String, + cookie2: String, } let fallback_probe_hits = Arc::new(Mutex::new(0usize)); @@ -595,6 +1018,48 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ .and_then(|value| value.to_str().ok()) .unwrap_or_default() .to_string(), + relay_sender: parts + .headers + .get(TUNNEL_RELAY_AUTH_SENDER_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + relay_owner: parts + .headers + .get(TUNNEL_RELAY_OWNER_INSTANCE_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + relay_timestamp: parts + .headers + .get(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + relay_nonce: parts + .headers + .get(TUNNEL_RELAY_AUTH_NONCE_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + relay_signature: parts + .headers + .get(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + cookie: parts + .headers + .get(http::header::COOKIE) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + cookie2: parts + .headers + .get("cookie2") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), }); ( StatusCode::OK, @@ -614,7 +1079,11 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ Some(0.1), )], vec![sample_endpoint("endpoint-owner", "provider-owner")], - vec![sample_key("key-owner", "provider-owner", "node-owner")], + vec![sample_bound_key( + "key-owner", + "provider-owner", + "node-owner", + )], )); let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( Some(hash_api_key("sk-client-openai-affinity")), @@ -623,32 +1092,42 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ let observed_at_unix_secs = current_unix_secs(); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( provider_catalog_repository, - "development-key", + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, ) .with_auth_api_key_reader(auth_repository) + .attach_proxy_node_repository_for_tests(Arc::new(InMemoryProxyNodeRepository::seed([ + sample_tunnel_proxy_node("node-owner", "test-generation-owner"), + ]))) .with_system_config_values_for_tests(vec![( tunnel_attachment_key("node-owner"), serde_json::to_value(crate::tunnel::TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: owner_url.clone(), + tunnel_generation: "test-generation-owner".to_string(), conn_count: 1, observed_at_unix_secs, }) .expect("attachment should serialize"), - )]); + )]) + .with_system_default_routing_group_for_tests(); let mut state = AppState::new().expect("gateway state should build"); state = state .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a:8080")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a:8080"), + RELAY_TEST_SECRET, + ); let short_timeout_client = reqwest::Client::builder() .timeout(Duration::from_millis(10)) .build() .expect("test client should build"); state.client = short_timeout_client.clone(); - state.owner_forward_client = short_timeout_client; + let affinity_cache_key = + system_default_affinity_cache_key("api-key-affinity-1", "openai:chat", "gpt-4.1"); state.remember_scheduler_affinity_target( - "scheduler_affinity:api-key-affinity-1:openai:chat:gpt-4.1", + &affinity_cache_key, crate::cache::SchedulerAffinityTarget { provider_id: "provider-owner".to_string(), endpoint_id: "endpoint-owner".to_string(), @@ -668,6 +1147,8 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ "Bearer sk-client-openai-affinity", ) .header(TRACE_ID_HEADER, "trace-tunnel-affinity-forward-1") + .header(http::header::COOKIE, "session=must-not-forward") + .header("cookie2", "legacy-session=must-not-forward") .body("{\"model\":\"gpt-4.1\",\"messages\":[]}") .send() .await @@ -679,7 +1160,7 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ .headers() .get(GATEWAY_HEADER) .and_then(|value| value.to_str().ok()), - Some("gateway-b-owner") + Some("rust-phase3b") ); assert_eq!( response @@ -720,12 +1201,68 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_ assert_eq!(owner_request.forwarded_for, "127.0.0.1"); assert_eq!(owner_request.forwarded_by, "gateway-a"); assert_eq!(owner_request.owner_instance_id, "gateway-b"); + assert_eq!(owner_request.relay_sender, "gateway-a"); + assert_eq!(owner_request.relay_owner, "gateway-b"); + assert!(!owner_request.relay_timestamp.is_empty()); + assert!(!owner_request.relay_nonce.is_empty()); + assert!(!owner_request.relay_signature.is_empty()); + assert_eq!(owner_request.cookie, ""); + assert_eq!(owner_request.cookie2, ""); gateway_handle.abort(); owner_handle.abort(); fallback_probe_handle.abort(); } +#[tokio::test] +async fn owner_forward_client_does_not_follow_redirects_with_relay_credentials() { + let redirected_hits = Arc::new(Mutex::new(0usize)); + let redirected_hits_clone = Arc::clone(&redirected_hits); + let redirected_target = Router::new().route( + "/captured", + any(move |_request: Request| { + let redirected_hits_inner = Arc::clone(&redirected_hits_clone); + async move { + *redirected_hits_inner.lock().expect("mutex should lock") += 1; + StatusCode::OK + } + }), + ); + let (redirected_url, redirected_handle) = start_server(redirected_target).await; + + let redirect_location = format!("{redirected_url}/captured"); + let redirect_source = Router::new().route( + "/relay", + any(move || { + let location = redirect_location.clone(); + async move { + Response::builder() + .status(StatusCode::TEMPORARY_REDIRECT) + .header(http::header::LOCATION, location) + .body(Body::empty()) + .expect("redirect response should build") + } + }), + ); + let (redirect_url, redirect_handle) = start_server(redirect_source).await; + + let state = AppState::new().expect("gateway state should build"); + let response = state + .owner_forward_client + .post(format!("{redirect_url}/relay")) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, "sensitive-signature") + .body("sensitive-relay-envelope") + .send() + .await + .expect("owner forward request should return the redirect response"); + + assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT); + assert_eq!(*redirected_hits.lock().expect("mutex should lock"), 0); + + redirect_handle.abort(); + redirected_handle.abort(); +} + #[tokio::test] async fn gateway_aggregates_sync_sse_from_remote_tunnel_owner_before_returning_to_client() { #[derive(Debug, Clone)] @@ -841,7 +1378,7 @@ async fn gateway_aggregates_sync_sse_from_remote_tunnel_owner_before_returning_t "endpoint-cli-owner", "provider-cli-owner", )], - vec![sample_codex_key( + vec![sample_bound_codex_key( "key-cli-owner", "provider-cli-owner", "node-cli-owner", @@ -850,26 +1387,37 @@ async fn gateway_aggregates_sync_sse_from_remote_tunnel_owner_before_returning_t let observed_at_unix_secs = current_unix_secs(); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( provider_catalog_repository, - "development-key", + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, ) .with_auth_api_key_reader(auth_repository) + .attach_proxy_node_repository_for_tests(Arc::new(InMemoryProxyNodeRepository::seed([ + sample_tunnel_proxy_node("node-cli-owner", "test-generation-cli-owner"), + ]))) .with_system_config_values_for_tests(vec![( tunnel_attachment_key("node-cli-owner"), serde_json::to_value(crate::tunnel::TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: owner_url.clone(), + tunnel_generation: "test-generation-cli-owner".to_string(), conn_count: 1, observed_at_unix_secs, }) .expect("attachment should serialize"), - )]); + )]) + .with_system_default_routing_group_for_tests(); let mut state = AppState::new().expect("gateway state should build"); state = state .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a:8080")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a:8080"), + RELAY_TEST_SECRET, + ); + let affinity_cache_key = + system_default_affinity_cache_key("api-key-affinity-cli-1", "openai:responses", "gpt-5.4"); state.remember_scheduler_affinity_target( - "scheduler_affinity:api-key-affinity-cli-1:openai:responses:gpt-5.4", + &affinity_cache_key, crate::cache::SchedulerAffinityTarget { provider_id: "provider-cli-owner".to_string(), endpoint_id: "endpoint-cli-owner".to_string(), @@ -904,7 +1452,7 @@ async fn gateway_aggregates_sync_sse_from_remote_tunnel_owner_before_returning_t .headers() .get(GATEWAY_HEADER) .and_then(|value| value.to_str().ok()), - Some("gateway-b-owner") + Some("rust-phase3b") ); assert_eq!( response @@ -1088,7 +1636,7 @@ async fn gateway_streamifies_sync_json_from_remote_tunnel_owner_before_returning "endpoint-cli-owner", "provider-cli-owner", )], - vec![sample_codex_key( + vec![sample_bound_codex_key( "key-cli-owner", "provider-cli-owner", "node-cli-owner", @@ -1097,30 +1645,41 @@ async fn gateway_streamifies_sync_json_from_remote_tunnel_owner_before_returning let observed_at_unix_secs = current_unix_secs(); let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( provider_catalog_repository, - "development-key", + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, ) .with_auth_api_key_reader(auth_repository) + .attach_proxy_node_repository_for_tests(Arc::new(InMemoryProxyNodeRepository::seed([ + sample_tunnel_proxy_node("node-cli-owner", "test-generation-cli-owner"), + ]))) .with_system_config_values_for_tests(vec![( tunnel_attachment_key("node-cli-owner"), serde_json::to_value(crate::tunnel::TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: owner_url.clone(), + tunnel_generation: "test-generation-cli-owner".to_string(), conn_count: 1, observed_at_unix_secs, }) .expect("attachment should serialize"), - )]); + )]) + .with_system_default_routing_group_for_tests(); let mut state = AppState::new().expect("gateway state should build"); state = state .with_data_state_for_tests(data_state) - .with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a:8080")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("http://gateway-a:8080"), + RELAY_TEST_SECRET, + ); state.client = reqwest::Client::builder() .timeout(Duration::from_millis(10)) .build() .expect("short shared client should build"); + let affinity_cache_key = + system_default_affinity_cache_key("api-key-affinity-cli-1", "openai:responses", "gpt-5.4"); state.remember_scheduler_affinity_target( - "scheduler_affinity:api-key-affinity-cli-1:openai:responses:gpt-5.4", + &affinity_cache_key, crate::cache::SchedulerAffinityTarget { provider_id: "provider-cli-owner".to_string(), endpoint_id: "endpoint-cli-owner".to_string(), @@ -1155,7 +1714,7 @@ async fn gateway_streamifies_sync_json_from_remote_tunnel_owner_before_returning .headers() .get(GATEWAY_HEADER) .and_then(|value| value.to_str().ok()), - Some("gateway-b-owner") + Some("rust-phase3b") ); assert_eq!( response diff --git a/apps/aether-gateway/src/tests/usage.rs b/apps/aether-gateway/src/tests/usage.rs index 77e865e61..43cd56bc0 100644 --- a/apps/aether-gateway/src/tests/usage.rs +++ b/apps/aether-gateway/src/tests/usage.rs @@ -122,7 +122,7 @@ pub(super) fn sample_local_openai_provider() -> StoredProviderCatalogProvider { false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -144,7 +144,7 @@ pub(super) fn sample_local_openai_endpoint() -> StoredProviderCatalogEndpoint { "https://api.openai.example/v1".to_string(), None, None, - Some(2), + Some(1), None, None, None, diff --git a/apps/aether-gateway/src/tests/usage/local.rs b/apps/aether-gateway/src/tests/usage/local.rs index d59d43a95..401252eed 100644 --- a/apps/aether-gateway/src/tests/usage/local.rs +++ b/apps/aether-gateway/src/tests/usage/local.rs @@ -12,7 +12,6 @@ use super::{ UsageReadRepository, UsageRuntimeConfig, DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER, }; use crate::constants::LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER; -use aether_data_contracts::repository::usage::UsageBodyCaptureState; fn deep_nested_metadata(levels: usize) -> serde_json::Value { let mut current = json!({"leaf": "value"}); @@ -403,32 +402,10 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im let stored_usage = stored_usage.expect("usage should be recorded"); assert_eq!(stored_usage.status, "completed"); assert_eq!(stored_usage.total_tokens, 5); - assert_eq!( - stored_usage - .request_body - .as_ref() - .and_then(|value| value.get("messages")) - .and_then(|value| value.as_array()) - .and_then(|messages| messages.first()) - .and_then(|value| value.get("content")) - .and_then(|value| value.as_str()) - .map(str::len), - Some(128 * 1024) - ); - assert_eq!( - stored_usage - .request_body - .as_ref() - .and_then(|value| value.get("metadata")) - .and_then(|value| value.get("child")) - .and_then(|value| value.get("child")) - .and_then(|value| value.get("child")) - .and_then(|value| value.get("child")) - .and_then(|value| value.get("child")) - .and_then(|value| value.as_object()) - .map(|value| value.contains_key("depth")), - Some(true) - ); + assert!(stored_usage.request_body.is_none()); + assert!(stored_usage.request_body_ref.is_none()); + assert!(stored_usage.request_body_state.is_none()); + assert!(stored_usage.request_headers.is_none()); gateway_handle.abort(); execution_runtime_handle.abort(); @@ -558,22 +535,12 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync ) .await; assert_eq!(stored_usage.total_tokens, 5); - assert_eq!( - stored_usage.request_body_state, - Some(UsageBodyCaptureState::Inline) - ); - assert_eq!( - stored_usage.provider_request_body_state, - Some(UsageBodyCaptureState::Inline) - ); - assert_ne!( - stored_usage - .request_body - .as_ref() - .and_then(|value| value.get("truncated")) - .and_then(|value| value.as_bool()), - Some(true) - ); + assert!(stored_usage.request_body.is_none()); + assert!(stored_usage.request_body_ref.is_none()); + assert!(stored_usage.request_body_state.is_none()); + assert!(stored_usage.provider_request_body.is_none()); + assert!(stored_usage.provider_request_body_ref.is_none()); + assert!(stored_usage.provider_request_body_state.is_none()); gateway_handle.abort(); execution_runtime_handle.abort(); @@ -856,32 +823,24 @@ async fn gateway_records_failed_usage_when_all_local_openai_chat_candidates_exha .and_then(|value| value.as_str()), Some("trace-openai-chat-local-report-sync-failure-123") ); - assert_eq!( - stored_usage - .response_body - .as_ref() - .and_then(|value| value.get("error")) - .and_then(|value| value.get("type")) - .and_then(|value| value.as_str()), - Some("upstream_error") - ); - assert_eq!( - stored_usage - .client_response_body - .as_ref() - .and_then(|value| value.get("error")) - .and_then(|value| value.get("type")) - .and_then(|value| value.as_str()), - Some("http_error") - ); + assert!(stored_usage.response_body.is_none()); + assert!(stored_usage.response_body_ref.is_none()); + assert!(stored_usage.response_body_state.is_none()); + assert!(stored_usage.client_response_body.is_none()); + assert!(stored_usage.client_response_body_ref.is_none()); + assert!(stored_usage.client_response_body_state.is_none()); let stored_candidates = request_candidate_repository .list_by_request_id("trace-openai-chat-local-report-sync-failure-123") .await .expect("request candidate trace should read"); - assert_eq!(stored_candidates.len(), 1); - assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Failed); - assert_eq!(stored_candidates[0].status_code, Some(503)); + // The only candidate is the sticky first key: the default policy retries + // it once on the same key before the request is exhausted. + assert_eq!(stored_candidates.len(), 2); + for candidate in &stored_candidates { + assert_eq!(candidate.status, RequestCandidateStatus::Failed); + assert_eq!(candidate.status_code, Some(503)); + } } #[test] @@ -957,7 +916,8 @@ async fn gateway_records_failed_usage_when_sync_runtime_transport_is_unavailable let response = send_request(gateway, request).await; assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); - assert_eq!(*execution_hits.lock().expect("mutex should lock"), 1); + // The sticky first key is retried once on the same key before exhaustion. + assert_eq!(*execution_hits.lock().expect("mutex should lock"), 2); let stored_usage = wait_for_usage_status( usage_repository.as_ref(), @@ -968,29 +928,23 @@ async fn gateway_records_failed_usage_when_sync_runtime_transport_is_unavailable assert_eq!(stored_usage.status, "failed"); assert_eq!(stored_usage.billing_status, "void"); assert_eq!(stored_usage.status_code, Some(503)); - assert_eq!( - stored_usage - .response_body - .as_ref() - .and_then(|value| value.get("error")) - .and_then(|value| value.get("type")) - .and_then(|value| value.as_str()), - Some("execution_runtime_unavailable") - ); + assert!(stored_usage.response_body.is_none()); + assert!(stored_usage.response_body_ref.is_none()); + assert!(stored_usage.response_body_state.is_none()); let stored_candidates = request_candidate_repository .list_by_request_id("trace-openai-chat-local-transport-unavailable-123") .await .expect("request candidate trace should read"); - assert_eq!(stored_candidates.len(), 1); - assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Failed); - assert!(stored_candidates[0] - .latency_ms - .is_some_and(|value| value >= 5)); - assert_eq!( - stored_candidates[0].error_type.as_deref(), - Some("execution_runtime_unavailable") - ); + assert_eq!(stored_candidates.len(), 2); + for candidate in &stored_candidates { + assert_eq!(candidate.status, RequestCandidateStatus::Failed); + assert!(candidate.latency_ms.is_some_and(|value| value >= 5)); + assert_eq!( + candidate.error_type.as_deref(), + Some("execution_runtime_unavailable") + ); + } } #[test] @@ -1316,15 +1270,10 @@ async fn gateway_records_failed_usage_for_claude_runtime_miss_without_execution_ .and_then(|value| value.as_str()), Some("trace-claude-runtime-miss-usage-123") ); - assert_eq!( - stored_usage - .client_response_body - .as_ref() - .and_then(|value| value.get("error")) - .and_then(|value| value.get("type")) - .and_then(|value| value.as_str()), - Some("overloaded_error") - ); + assert!(stored_usage.client_response_body.is_none()); + assert!(stored_usage.client_response_body_ref.is_none()); + assert!(stored_usage.client_response_body_state.is_none()); + assert!(stored_usage.error_message.is_none()); let stored_candidates = request_candidate_repository .list_by_request_id("trace-claude-runtime-miss-usage-123") @@ -1675,22 +1624,12 @@ async fn gateway_ignores_legacy_max_response_body_size_for_stream_usage_impl() { ) .await; assert_eq!(stored_usage.total_tokens, 6); - assert_eq!( - stored_usage.response_body_state, - Some(UsageBodyCaptureState::Inline) - ); - assert_eq!( - stored_usage.client_response_body_state, - Some(UsageBodyCaptureState::Inline) - ); - assert_ne!( - stored_usage - .response_body - .as_ref() - .and_then(|value| value.get("truncated")) - .and_then(|value| value.as_bool()), - Some(true) - ); + assert!(stored_usage.response_body.is_none()); + assert!(stored_usage.response_body_ref.is_none()); + assert!(stored_usage.response_body_state.is_none()); + assert!(stored_usage.client_response_body.is_none()); + assert!(stored_usage.client_response_body_ref.is_none()); + assert!(stored_usage.client_response_body_state.is_none()); gateway_handle.abort(); execution_runtime_handle.abort(); @@ -1786,7 +1725,7 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -1808,7 +1747,7 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s "https://right.codes/codex".to_string(), None, None, - Some(2), + Some(1), Some("/v1/messages".to_string()), None, None, @@ -1963,36 +1902,11 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s stored_usage.routing_local_execution_runtime_miss_reason(), Some("all_candidates_skipped") ); - assert_eq!( - stored_usage.error_message.as_deref(), - Some( - "找到 1 个支持模型 gpt-5.4 的候选提供商,但本次同步请求全部不可用:格式转换未启用 2 次(原因代码: all_candidates_skipped)" - ) - ); - assert_eq!( - stored_usage - .request_headers - .as_ref() - .and_then(|value| value.get("authorization")) - .and_then(|value| value.as_str()), - Some("Bear****miss") - ); - assert_eq!( - stored_usage - .request_headers - .as_ref() - .and_then(|value| value.get("content-type")) - .and_then(|value| value.as_str()), - Some("application/json") - ); - assert_eq!( - stored_usage - .request_body - .as_ref() - .and_then(|value| value.get("model")) - .and_then(|value| value.as_str()), - Some("gpt-5.4") - ); + assert!(stored_usage.error_message.is_none()); + assert!(stored_usage.request_headers.is_none()); + assert!(stored_usage.request_body.is_none()); + assert!(stored_usage.request_body_ref.is_none()); + assert!(stored_usage.request_body_state.is_none()); assert!(stored_usage.provider_request_body.is_none()); assert_eq!( stored_usage @@ -2111,7 +2025,7 @@ fn gateway_keeps_failed_usage_request_capture_lightweight_for_large_local_claude false, false, None, - Some(2), + Some(1), None, Some(20.0), None, @@ -2133,7 +2047,7 @@ fn gateway_keeps_failed_usage_request_capture_lightweight_for_large_local_claude "https://right.codes/codex".to_string(), None, None, - Some(2), + Some(1), Some("/v1/messages".to_string()), None, None, @@ -2255,18 +2169,9 @@ fn gateway_keeps_failed_usage_request_capture_lightweight_for_large_local_claude ) .await; assert_eq!(stored_usage.status, "failed"); - assert_eq!( - stored_usage.request_body_state, - Some(UsageBodyCaptureState::Inline) - ); - assert_eq!( - stored_usage - .request_body - .as_ref() - .and_then(|value| value.get("model")) - .and_then(|value| value.as_str()), - Some("gpt-5.4") - ); + assert!(stored_usage.request_body_state.is_none()); + assert!(stored_usage.request_body.is_none()); + assert!(stored_usage.request_body_ref.is_none()); assert!(stored_usage.provider_request_body.is_none()); assert_eq!( stored_usage diff --git a/apps/aether-gateway/src/tests/video/data_read.rs b/apps/aether-gateway/src/tests/video/data_read.rs index 6657e0fc2..04e4efee4 100644 --- a/apps/aether-gateway/src/tests/video/data_read.rs +++ b/apps/aether-gateway/src/tests/video/data_read.rs @@ -1,5 +1,8 @@ use std::sync::{Arc, Mutex}; +use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, +}; use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, @@ -9,9 +12,43 @@ use axum::routing::any; use axum::{extract::Request, Json, Router}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use super::{build_router_with_state, build_state_with_execution_runtime_override, start_server}; +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(), + "video-user".to_string(), + Some("video@example.com".to_string()), + "user".to_string(), + "local".to_string(), + true, + false, + None, + None, + None, + api_key_id.to_string(), + Some("default".to_string()), + true, + false, + false, + Some(60), + Some(5), + Some(4_102_444_800), + None, + None, + None, + ) + .expect("auth snapshot should build") +} + #[tokio::test] async fn gateway_reads_openai_video_task_via_data_read_side_without_hitting_public_route() { let public_hits = Arc::new(Mutex::new(0usize)); @@ -92,15 +129,25 @@ async fn gateway_reads_openai_video_task_via_data_read_side_without_hitting_publ }) .await .expect("upsert should succeed"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-video-read-owner-key")), + sample_auth_snapshot("api-key-video-db-rotated-123", "user-video-db-123"), + )])); let gateway = build_router_with_state( build_state_with_execution_runtime_override(upstream_url.clone()) - .with_video_task_data_reader_for_tests(repository), + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( + auth_repository, + repository, + ), + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; let response = reqwest::Client::new() .get(format!("{gateway_url}/v1/videos/task-db-123")) + .bearer_auth("client-video-read-owner-key") .send() .await .expect("request should succeed"); @@ -211,10 +258,19 @@ async fn gateway_reads_gemini_video_task_via_data_read_side_without_hitting_publ }) .await .expect("upsert should succeed"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-gemini-video-read-owner-key")), + sample_auth_snapshot("api-key-video-db-rotated-123", "user-video-db-123"), + )])); let gateway = build_router_with_state( build_state_with_execution_runtime_override(upstream_url.clone()) - .with_video_task_data_reader_for_tests(repository), + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( + auth_repository, + repository, + ), + ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -222,6 +278,7 @@ async fn gateway_reads_gemini_video_task_via_data_read_side_without_hitting_publ .get(format!( "{gateway_url}/v1beta/models/veo-3/operations/localshort123" )) + .header("x-goog-api-key", "client-gemini-video-read-owner-key") .send() .await .expect("request should succeed"); @@ -239,3 +296,118 @@ async fn gateway_reads_gemini_video_task_via_data_read_side_without_hitting_publ gateway_handle.abort(); upstream_handle.abort(); } + +#[tokio::test] +async fn gateway_hides_data_backed_video_task_from_non_owner() { + let public_hits = Arc::new(Mutex::new(0usize)); + let public_hits_clone = Arc::clone(&public_hits); + 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": "video", + "auth_endpoint_signature": "openai:video", + "execution_runtime_candidate": true, + "auth_context": { + "user_id": "user-video-foreign", + "api_key_id": "key-video-foreign", + "access_allowed": true + }, + "public_path": "/v1/videos/task-owned-123" + })) + }), + ) + .route( + "/v1/videos/task-owned-123", + any(move |_request: Request| { + let public_hits_inner = Arc::clone(&public_hits_clone); + async move { + *public_hits_inner.lock().expect("mutex should lock") += 1; + (StatusCode::IM_A_TEAPOT, Body::from("public-route-hit")) + } + }), + ); + + let (upstream_url, upstream_handle) = start_server(upstream).await; + let repository = Arc::new(InMemoryVideoTaskRepository::default()); + repository + .upsert(UpsertVideoTask { + id: "task-owned-123".to_string(), + short_id: None, + request_id: "request-owned-123".to_string(), + user_id: Some("user-video-owner".to_string()), + api_key_id: Some("key-video-owner".to_string()), + username: None, + api_key_name: None, + external_task_id: Some("ext-owned-123".to_string()), + provider_id: Some("provider-owned-123".to_string()), + endpoint_id: Some("endpoint-owned-123".to_string()), + key_id: Some("provider-key-owned-123".to_string()), + client_api_format: Some("openai:video".to_string()), + provider_api_format: Some("openai:video".to_string()), + format_converted: false, + model: Some("sora-2".to_string()), + prompt: Some("private video".to_string()), + original_request_body: Some(json!({"prompt": "private video"})), + duration_seconds: Some(4), + resolution: Some("720p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("1280x720".to_string()), + status: VideoTaskStatus::Processing, + progress_percent: 50, + progress_message: None, + retry_count: 0, + poll_interval_seconds: 10, + next_poll_at_unix_secs: Some(124), + poll_count: 1, + max_poll_count: 360, + created_at_unix_ms: 123, + submitted_at_unix_secs: Some(123), + completed_at_unix_secs: None, + updated_at_unix_secs: 124, + error_code: None, + error_message: None, + video_url: None, + request_metadata: None, + }) + .await + .expect("upsert should succeed"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-video-read-foreign-key")), + sample_auth_snapshot("key-video-foreign", "user-video-foreign"), + )])); + let gateway = build_router_with_state( + build_state_with_execution_runtime_override(upstream_url.clone()) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests( + auth_repository, + repository, + ), + ), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .get(format!("{gateway_url}/v1/videos/task-owned-123")) + .bearer_auth("client-video-read-foreign-key") + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::NOT_FOUND); + assert_eq!( + response + .json::() + .await + .expect("json body"), + json!({"detail": "Video task not found"}) + ); + assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/video/gemini_sync_create.rs b/apps/aether-gateway/src/tests/video/gemini_sync_create.rs index 3a6db5bd8..227264f1c 100644 --- a/apps/aether-gateway/src/tests/video/gemini_sync_create.rs +++ b/apps/aether-gateway/src/tests/video/gemini_sync_create.rs @@ -24,7 +24,10 @@ use std::sync::{Arc, Mutex}; use crate::constants::TRACE_ID_HEADER; -use super::{build_router_with_state, build_state_with_execution_runtime_override, start_server}; +use super::{ + build_router_with_state, build_state_with_execution_runtime_override, start_server, + video_proxy_node_repository, +}; #[tokio::test] async fn gateway_executes_gemini_video_create_via_local_decision_gate_with_local_planning_only() { @@ -404,7 +407,10 @@ async fn gateway_executes_gemini_video_create_via_local_decision_gate_with_local provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .attach_proxy_node_repository_for_tests(video_proxy_node_repository([ + "proxy-node-gemini-video-local", + ])), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; diff --git a/apps/aether-gateway/src/tests/video/gemini_sync_task.rs b/apps/aether-gateway/src/tests/video/gemini_sync_task.rs index 2c47ec702..b801dc257 100644 --- a/apps/aether-gateway/src/tests/video/gemini_sync_task.rs +++ b/apps/aether-gateway/src/tests/video/gemini_sync_task.rs @@ -1,13 +1,12 @@ -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; +use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, +}; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; -use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; use aether_data_contracts::repository::candidates::{ RequestCandidateReadRepository, RequestCandidateStatus, }; -use aether_data_contracts::repository::provider_catalog::{ - StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, -}; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, }; @@ -18,15 +17,49 @@ use axum::{extract::Request, Json, Router}; use http::header::{HeaderName, HeaderValue}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use std::sync::{Arc, Mutex}; use crate::constants::{CONTROL_EXECUTED_HEADER, CONTROL_EXECUTE_FALLBACK_HEADER, TRACE_ID_HEADER}; use super::{ build_router_with_state, build_state_with_execution_runtime_override, start_server, - VideoTaskTruthSourceMode, + video_provider_catalog_repository, VideoTaskTruthSourceMode, }; +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(), + "video-user".to_string(), + Some("video@example.com".to_string()), + "user".to_string(), + "local".to_string(), + true, + false, + Some(json!(["gemini"])), + Some(json!(["gemini:video"])), + Some(json!(["veo-3"])), + api_key_id.to_string(), + Some("default".to_string()), + true, + false, + false, + Some(60), + Some(5), + Some(4_102_444_800), + Some(json!(["gemini"])), + Some(json!(["gemini:video"])), + Some(json!(["veo-3"])), + ) + .expect("auth snapshot should build") +} + #[tokio::test] async fn gateway_executes_gemini_video_cancel_via_data_backed_local_follow_up_with_local_planning_only( ) { @@ -241,16 +274,35 @@ async fn gateway_executes_gemini_video_cancel_via_data_backed_local_follow_up_wi }) .await .expect("upsert should succeed"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-gemini-video-cancel-local-key")), + sample_auth_snapshot( + "key-gemini-video-cancel-rotated-local-123", + "user-gemini-video-cancel-local-123", + ), + )])); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-gemini-video-local-1", + "gemini", + "endpoint-gemini-video-local-1", + "gemini:video", + "https://generativelanguage.googleapis.com", + "key-gemini-video-local-1", + "sk-upstream-gemini-video", + ); 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) .with_data_state_for_tests( - crate::data::GatewayDataState::with_video_task_and_request_candidate_repository_for_tests( - repository, - Arc::clone(&request_candidate_repository), - ), - ); + crate::data::GatewayDataState::with_video_task_provider_transport_and_request_candidate_repository_for_tests( + repository, + provider_catalog_repository, + Arc::clone(&request_candidate_repository), + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_auth_api_key_reader(auth_repository), + ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -258,7 +310,7 @@ async fn gateway_executes_gemini_video_cancel_via_data_backed_local_follow_up_wi .post(format!( "{gateway_url}/v1beta/models/veo-3/operations/localshort123:cancel" )) - .header("x-goog-api-key", "client-key") + .header("x-goog-api-key", "client-gemini-video-cancel-local-key") .header(http::header::CONTENT_TYPE, "application/json") .header(TRACE_ID_HEADER, "trace-gemini-video-cancel-local-123") .body("{}") @@ -267,13 +319,8 @@ async fn gateway_executes_gemini_video_cancel_via_data_backed_local_follow_up_wi .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .json::() - .await - .expect("body should parse"), - json!({}) - ); + let response: serde_json::Value = response.json().await.expect("body should parse"); + assert_eq!(response, json!({})); let seen_execution_runtime_request = seen_execution_runtime .lock() @@ -318,75 +365,6 @@ async fn gateway_executes_gemini_video_cancel_via_reconstructed_data_backed_loca api_key: String, } - fn sample_provider() -> StoredProviderCatalogProvider { - StoredProviderCatalogProvider::new( - "provider-gemini-video-followup-1".to_string(), - "gemini".to_string(), - Some("https://example.com".to_string()), - "custom".to_string(), - ) - .expect("provider should build") - .with_transport_fields( - true, - false, - false, - None, - Some(2), - None, - Some(20.0), - None, - None, - ) - } - - fn sample_endpoint() -> StoredProviderCatalogEndpoint { - StoredProviderCatalogEndpoint::new( - "endpoint-gemini-video-followup-1".to_string(), - "provider-gemini-video-followup-1".to_string(), - "gemini:video".to_string(), - Some("gemini".to_string()), - Some("video".to_string()), - true, - ) - .expect("endpoint should build") - .with_transport_fields( - "https://generativelanguage.googleapis.com".to_string(), - None, - None, - Some(2), - None, - None, - None, - None, - ) - .expect("endpoint transport should build") - } - - fn sample_key() -> StoredProviderCatalogKey { - StoredProviderCatalogKey::new( - "key-gemini-video-followup-1".to_string(), - "provider-gemini-video-followup-1".to_string(), - "prod".to_string(), - "api_key".to_string(), - None, - true, - ) - .expect("key should build") - .with_transport_fields( - Some(json!(["gemini:video"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-upstream-gemini-video") - .expect("api key should encrypt"), - None, - None, - Some(json!({"gemini:video": 1})), - None, - None, - None, - None, - ) - .expect("key transport should build") - } - let decision_hits = Arc::new(Mutex::new(0usize)); let decision_hits_clone = Arc::clone(&decision_hits); let execute_hits = Arc::new(Mutex::new(0usize)); @@ -562,31 +540,43 @@ async fn gateway_executes_gemini_video_cancel_via_reconstructed_data_backed_loca }) .await .expect("upsert should succeed"); - let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider()], - vec![sample_endpoint()], - vec![sample_key()], - )); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-gemini-video-followup-1", + "gemini", + "endpoint-gemini-video-followup-1", + "gemini:video", + "https://generativelanguage.googleapis.com", + "key-gemini-video-followup-1", + "sk-upstream-gemini-video", + ); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-gemini-video-cancel-op-key")), + sample_auth_snapshot( + "key-gemini-video-cancel-rotated-op-123", + "user-gemini-video-cancel-op-123", + ), + )])); - let gateway = build_router_with_state( - build_state_with_execution_runtime_override(execution_runtime_url) + let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) .with_data_state_for_tests( crate::data::GatewayDataState::with_video_task_provider_transport_and_request_candidate_repository_for_tests( - repository, + repository.clone(), provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), - ), - ); + ) + .with_auth_api_key_reader(auth_repository), + ); + 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}/v1beta/models/veo-3/operations/opshort123:cancel" )) + .header("x-goog-api-key", "client-gemini-video-cancel-op-key") .header(CONTROL_EXECUTE_FALLBACK_HEADER, "true") .header(TRACE_ID_HEADER, "trace-gemini-video-cancel-op-123") .send() @@ -594,13 +584,8 @@ async fn gateway_executes_gemini_video_cancel_via_reconstructed_data_backed_loca .expect("request should succeed"); assert_eq!(response.status(), StatusCode::OK); - assert_eq!( - response - .json::() - .await - .expect("body should parse"), - json!({}) - ); + let response: serde_json::Value = response.json().await.expect("body should parse"); + assert_eq!(response, json!({})); let seen_execution_runtime_request = seen_execution_runtime .lock() diff --git a/apps/aether-gateway/src/tests/video/mod.rs b/apps/aether-gateway/src/tests/video/mod.rs index 4d5771a21..a93ac3b2d 100644 --- a/apps/aether-gateway/src/tests/video/mod.rs +++ b/apps/aether-gateway/src/tests/video/mod.rs @@ -1,6 +1,12 @@ use std::sync::{Arc, Mutex}; +use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; +use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, +}; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskLookupKey, VideoTaskReadRepository, VideoTaskWriteRepository, }; @@ -30,3 +36,137 @@ mod openai_sync_task; mod registry_poller; mod routing; mod stream; + +/// Seed online manual proxy nodes for video execution fixtures. +/// +/// Production resolution intentionally fails closed when a provider refers to +/// an unregistered node. Tests that exercise a configured node therefore need +/// 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, +{ + let nodes = node_ids.into_iter().map(|node_id| { + let node_id = node_id.as_ref(); + StoredProxyNode::new( + node_id.to_string(), + format!("video-test-{node_id}"), + "127.0.0.1".to_string(), + 1, + true, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + false, + false, + 1, + ) + .expect("video test proxy node should build") + .with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None) + .with_tunnel_generation(format!("video-test-generation-{node_id}")) + }); + Arc::new(InMemoryProxyNodeRepository::seed(nodes)) +} + +/// Build a provider catalog row set for tasks whose sensitive snapshot fields +/// have been removed by the persistence boundary. The credential is sealed +/// with the record-bound v2 envelope used by production, so a read-only test +/// state can reconstruct transport without relying on a migration writer. +pub(super) fn video_provider_catalog_repository( + provider_id: &str, + provider_type: &str, + endpoint_id: &str, + api_format: &str, + endpoint_base_url: &str, + key_id: &str, + upstream_api_key: &str, +) -> Arc { + fn seal_bound_credential( + provider_id: &str, + key_id: &str, + field: &str, + plaintext: &str, + ) -> String { + let purpose = format!( + "provider-catalog-credential-bound-v2\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={field}", + provider_id.len(), + key_id.len(), + ); + let protected = format!("{purpose}\0{plaintext}"); + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &protected) + .expect("provider test credential should encrypt"); + format!("aether-provider-catalog-credential-v2:aether-runtime-secret-v1:{ciphertext}") + } + + let provider = StoredProviderCatalogProvider::new( + provider_id.to_string(), + format!("video-{provider_type}"), + Some("https://example.com".to_string()), + provider_type.to_string(), + ) + .expect("provider should build") + .with_transport_fields( + true, + false, + false, + None, + Some(2), + None, + Some(20.0), + None, + None, + ); + let endpoint = StoredProviderCatalogEndpoint::new( + endpoint_id.to_string(), + provider_id.to_string(), + api_format.to_string(), + Some(provider_type.to_string()), + Some("video".to_string()), + true, + ) + .expect("endpoint should build") + .with_transport_fields( + endpoint_base_url.to_string(), + None, + None, + Some(2), + None, + None, + None, + None, + ) + .expect("endpoint transport should build"); + let key = StoredProviderCatalogKey::new( + key_id.to_string(), + provider_id.to_string(), + "prod".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + Some(serde_json::json!([api_format])), + seal_bound_credential(provider_id, key_id, "api-key", upstream_api_key), + None, + None, + Some(serde_json::json!({api_format: 1})), + None, + None, + None, + None, + ) + .expect("key transport should build"); + + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![key], + )) +} diff --git a/apps/aether-gateway/src/tests/video/openai_sync_create.rs b/apps/aether-gateway/src/tests/video/openai_sync_create.rs index 536c1939a..3d47e3da6 100644 --- a/apps/aether-gateway/src/tests/video/openai_sync_create.rs +++ b/apps/aether-gateway/src/tests/video/openai_sync_create.rs @@ -28,7 +28,10 @@ use std::sync::{Arc, Mutex}; use crate::constants::TRACE_ID_HEADER; -use super::{build_router_with_state, build_state_with_execution_runtime_override, start_server}; +use super::{ + build_router_with_state, build_state_with_execution_runtime_override, start_server, + video_provider_catalog_repository, video_proxy_node_repository, +}; #[tokio::test] async fn gateway_executes_openai_video_create_via_local_decision_gate_with_local_planning_only() { @@ -417,7 +420,10 @@ async fn gateway_executes_openai_video_create_via_local_decision_gate_with_local provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .attach_proxy_node_repository_for_tests(video_proxy_node_repository([ + "proxy-node-openai-video-local", + ])), ); let gateway = build_router_with_state(gateway_state); let (gateway_url, gateway_handle) = start_server(gateway).await; @@ -516,6 +522,39 @@ async fn gateway_executes_openai_video_remix_via_data_backed_local_follow_up_wit prompt: String, } + 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(), + "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!(["sora-2"])), + 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!(["sora-2"])), + ) + .expect("auth snapshot should build") + } + let decision_hits = Arc::new(Mutex::new(0usize)); let decision_hits_clone = Arc::clone(&decision_hits); let plan_hits = Arc::new(Mutex::new(0usize)); @@ -738,21 +777,42 @@ async fn gateway_executes_openai_video_remix_via_data_backed_local_follow_up_wit .await .expect("upsert should succeed"); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("client-openai-video-remix-local-key")), + sample_auth_snapshot( + "key-openai-video-remix-rotated-local-123", + "user-openai-video-remix-local-123", + ), + )])); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-openai-video-local-1", + "openai", + "endpoint-openai-video-local-1", + "openai:video", + "https://api.openai.example/v1", + "key-openai-video-local-1", + "sk-upstream-openai-video", + ); + 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) .with_data_state_for_tests( - crate::data::GatewayDataState::with_video_task_and_request_candidate_repository_for_tests( - repository, - Arc::clone(&request_candidate_repository), - ), - ); + crate::data::GatewayDataState::with_video_task_provider_transport_and_request_candidate_repository_for_tests( + repository, + provider_catalog_repository, + Arc::clone(&request_candidate_repository), + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_auth_api_key_reader(auth_repository), + ); 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/videos/task-local-123/remix")) .header(http::header::CONTENT_TYPE, "application/json") + .bearer_auth("client-openai-video-remix-local-key") .header(TRACE_ID_HEADER, "trace-openai-video-remix-local-123") .body("{\"prompt\":\"remix this\",\"model\":\"sora-2\"}") .send() diff --git a/apps/aether-gateway/src/tests/video/openai_sync_task.rs b/apps/aether-gateway/src/tests/video/openai_sync_task.rs index f6357fb47..f164d1a73 100644 --- a/apps/aether-gateway/src/tests/video/openai_sync_task.rs +++ b/apps/aether-gateway/src/tests/video/openai_sync_task.rs @@ -1,13 +1,12 @@ -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; +use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, +}; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; -use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; use aether_data_contracts::repository::candidates::{ RequestCandidateReadRepository, RequestCandidateStatus, }; -use aether_data_contracts::repository::provider_catalog::{ - StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, -}; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, }; @@ -18,13 +17,14 @@ use axum::{extract::Request, Json, Router}; use http::header::{HeaderName, HeaderValue}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use std::sync::{Arc, Mutex}; use crate::constants::{CONTROL_EXECUTED_HEADER, CONTROL_EXECUTE_FALLBACK_HEADER, TRACE_ID_HEADER}; use super::{ build_router_with_state, build_state_with_execution_runtime_override, start_server, - VideoTaskTruthSourceMode, + video_provider_catalog_repository, VideoTaskTruthSourceMode, }; #[tokio::test] @@ -37,73 +37,37 @@ async fn gateway_executes_openai_video_delete_via_reconstructed_data_backed_loca authorization: String, } - fn sample_provider() -> StoredProviderCatalogProvider { - StoredProviderCatalogProvider::new( - "provider-openai-video-followup-1".to_string(), - "openai".to_string(), - Some("https://example.com".to_string()), - "custom".to_string(), - ) - .expect("provider should build") - .with_transport_fields( - true, - false, - false, - None, - Some(2), - None, - Some(20.0), - None, - None, - ) + fn hash_api_key(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) } - fn sample_endpoint() -> StoredProviderCatalogEndpoint { - StoredProviderCatalogEndpoint::new( - "endpoint-openai-video-followup-1".to_string(), - "provider-openai-video-followup-1".to_string(), - "openai:video".to_string(), - Some("openai".to_string()), - Some("video".to_string()), + 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, - ) - .expect("endpoint should build") - .with_transport_fields( - "https://api.openai.example/v1".to_string(), - None, - None, - Some(2), - None, - None, - None, - None, - ) - .expect("endpoint transport should build") - } - - fn sample_key() -> StoredProviderCatalogKey { - StoredProviderCatalogKey::new( - "key-openai-video-followup-1".to_string(), - "provider-openai-video-followup-1".to_string(), - "prod".to_string(), - "api_key".to_string(), - None, - true, - ) - .expect("key should build") - .with_transport_fields( + false, + Some(json!(["openai"])), Some(json!(["openai:video"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-upstream-openai-video") - .expect("api key should encrypt"), - None, - None, - Some(json!({"openai:video": 1})), - None, - None, - None, - None, + Some(json!(["sora-2"])), + 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!(["sora-2"])), ) - .expect("key transport should build") + .expect("auth snapshot should build") } let decision_hits = Arc::new(Mutex::new(0usize)); @@ -116,7 +80,6 @@ async fn gateway_executes_openai_video_delete_via_reconstructed_data_backed_loca let report_hits_clone = Arc::clone(&report_hits); let seen_execution_runtime = Arc::new(Mutex::new(None::)); let seen_execution_runtime_clone = Arc::clone(&seen_execution_runtime); - let upstream = Router::new() .route( "/api/internal/gateway/resolve", @@ -128,11 +91,6 @@ async fn gateway_executes_openai_video_delete_via_reconstructed_data_backed_loca "route_kind": "video", "auth_endpoint_signature": "openai:video", "execution_runtime_candidate": true, - "auth_context": { - "user_id": "user-openai-video-delete-local-123", - "api_key_id": "key-openai-video-delete-local-123", - "access_allowed": true - }, "public_path": "/v1/videos/task-local-followup-123" })) }), @@ -279,12 +237,32 @@ async fn gateway_executes_openai_video_delete_via_reconstructed_data_backed_loca }) .await .expect("upsert should succeed"); - let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider()], - vec![sample_endpoint()], - vec![sample_key()], - )); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-openai-video-followup-1", + "openai", + "endpoint-openai-video-followup-1", + "openai:video", + "https://api.openai.example/v1", + "key-openai-video-followup-1", + "sk-upstream-openai-video", + ); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![ + ( + Some(hash_api_key("client-video-delete-foreign-key")), + sample_auth_snapshot( + "key-openai-video-delete-foreign-123", + "user-openai-video-delete-foreign-123", + ), + ), + ( + Some(hash_api_key("client-video-delete-owner-key")), + sample_auth_snapshot( + "key-openai-video-delete-rotated-local-123", + "user-openai-video-delete-local-123", + ), + ), + ])); let gateway = build_router_with_state( build_state_with_execution_runtime_override(execution_runtime_url) @@ -295,14 +273,44 @@ async fn gateway_executes_openai_video_delete_via_reconstructed_data_backed_loca provider_catalog_repository, Arc::clone(&request_candidate_repository), DEVELOPMENT_ENCRYPTION_KEY, - ), + ) + .with_auth_api_key_reader(auth_repository), ), ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() + let client = reqwest::Client::new(); + let foreign_response = client .delete(format!("{gateway_url}/v1/videos/task-local-followup-123")) .header(CONTROL_EXECUTE_FALLBACK_HEADER, "true") + .bearer_auth("client-video-delete-foreign-key") + .header( + TRACE_ID_HEADER, + "trace-openai-video-delete-foreign-local-123", + ) + .send() + .await + .expect("foreign delete request should complete"); + let foreign_status = foreign_response.status(); + let foreign_body = foreign_response.text().await.expect("body should read"); + assert_eq!( + foreign_status, + StatusCode::NOT_FOUND, + "unexpected foreign response body: {foreign_body}" + ); + assert!(seen_execution_runtime + .lock() + .expect("mutex should lock") + .is_none()); + assert_eq!(*decision_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*report_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); + + let response = client + .delete(format!("{gateway_url}/v1/videos/task-local-followup-123")) + .header(CONTROL_EXECUTE_FALLBACK_HEADER, "true") + .bearer_auth("client-video-delete-owner-key") .header(TRACE_ID_HEADER, "trace-openai-video-delete-local-123") .send() .await diff --git a/apps/aether-gateway/src/tests/video/registry_poller.rs b/apps/aether-gateway/src/tests/video/registry_poller.rs index e1dfd576c..1f176d20d 100644 --- a/apps/aether-gateway/src/tests/video/registry_poller.rs +++ b/apps/aether-gateway/src/tests/video/registry_poller.rs @@ -1,5 +1,6 @@ use std::sync::{Arc, Mutex}; +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskLookupKey, VideoTaskReadRepository, VideoTaskStatus, @@ -11,7 +12,8 @@ use axum::{extract::Request, Json, Router}; use serde_json::json; use super::{ - build_state_with_execution_runtime_override, start_server, AppState, VideoTaskTruthSourceMode, + build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository, + AppState, VideoTaskTruthSourceMode, }; fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask { @@ -168,9 +170,24 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep .upsert(sample_due_openai_task("https://api.openai.example/v1")) .await .expect("task upsert should succeed"); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-openai-video-local-1", + "openai", + "endpoint-openai-video-local-1", + "openai:video", + "https://api.openai.example/v1", + "key-openai-video-local-1", + "sk-upstream-openai-video", + ); let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url) - .with_video_task_data_repository_for_tests(Arc::clone(&repository)) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + Arc::clone(&repository), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) .with_video_task_poller_config(std::time::Duration::from_millis(25), 8); let background_tasks = gateway_state.spawn_background_tasks(); @@ -202,23 +219,10 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep stored.next_poll_at_unix_secs.is_some_and(|value| value > 0), "poller should push next poll into the future" ); - assert_eq!( - stored - .request_metadata - .as_ref() - .and_then(|value| value.get("rust_owner")) - .and_then(serde_json::Value::as_str), - Some("async_task") - ); - assert_eq!( - stored - .request_metadata - .as_ref() - .and_then(|value| value.get("poll_raw_response")) - .and_then(|value| value.get("status")) - .and_then(serde_json::Value::as_str), - Some("processing") - ); + assert!(stored.original_request_body.is_none()); + assert!(stored.progress_message.is_none()); + assert!(stored.error_message.is_none()); + assert!(stored.request_metadata.is_none()); assert_eq!( seen_execution_runtime_requests @@ -274,10 +278,25 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep .upsert(sample_due_openai_task(&upstream_api_root)) .await .expect("task upsert should succeed"); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-openai-video-local-1", + "openai", + "endpoint-openai-video-local-1", + "openai:video", + &upstream_api_root, + "key-openai-video-local-1", + "sk-upstream-openai-video", + ); let gateway_state = AppState::new() .expect("gateway state should build") - .with_video_task_data_repository_for_tests(Arc::clone(&repository)) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + Arc::clone(&repository), + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ), + ) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) .with_video_task_poller_config(std::time::Duration::from_millis(25), 8); let background_tasks = gateway_state.spawn_background_tasks(); diff --git a/apps/aether-gateway/src/tests/video/stream.rs b/apps/aether-gateway/src/tests/video/stream.rs index f41b009fd..1e50daecd 100644 --- a/apps/aether-gateway/src/tests/video/stream.rs +++ b/apps/aether-gateway/src/tests/video/stream.rs @@ -1,10 +1,9 @@ use aether_contracts::{StreamFrame, StreamFramePayload, StreamFrameType}; -use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; -use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; -use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; -use aether_data_contracts::repository::provider_catalog::{ - StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, +use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; +use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, }; +use aether_data::repository::video_tasks::InMemoryVideoTaskRepository; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, }; @@ -16,13 +15,14 @@ use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine as _}; use http::header::{HeaderName, HeaderValue}; use http::StatusCode; use serde_json::json; +use sha2::{Digest, Sha256}; use std::sync::{Arc, Mutex}; use crate::constants::{CONTROL_EXECUTED_HEADER, CONTROL_EXECUTE_FALLBACK_HEADER, TRACE_ID_HEADER}; use super::{ build_router_with_state, build_state_with_execution_runtime_override, start_server, - VideoTaskTruthSourceMode, + video_provider_catalog_repository, VideoTaskTruthSourceMode, }; #[tokio::test] @@ -34,73 +34,37 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with url: String, } - fn sample_provider() -> StoredProviderCatalogProvider { - StoredProviderCatalogProvider::new( - "provider-openai-video-content-followup-1".to_string(), - "openai".to_string(), - Some("https://example.com".to_string()), - "custom".to_string(), - ) - .expect("provider should build") - .with_transport_fields( - true, - false, - false, - None, - Some(2), - None, - Some(20.0), - None, - None, - ) + fn hash_api_key(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) } - fn sample_endpoint() -> StoredProviderCatalogEndpoint { - StoredProviderCatalogEndpoint::new( - "endpoint-openai-video-content-followup-1".to_string(), - "provider-openai-video-content-followup-1".to_string(), - "openai:video".to_string(), - Some("openai".to_string()), - Some("video".to_string()), + 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, - ) - .expect("endpoint should build") - .with_transport_fields( - "https://api.openai.example/v1".to_string(), - None, - None, - Some(2), - None, - None, - None, - None, - ) - .expect("endpoint transport should build") - } - - fn sample_key() -> StoredProviderCatalogKey { - StoredProviderCatalogKey::new( - "key-openai-video-content-followup-1".to_string(), - "provider-openai-video-content-followup-1".to_string(), - "prod".to_string(), - "api_key".to_string(), - None, - true, - ) - .expect("key should build") - .with_transport_fields( + false, + Some(json!(["openai"])), Some(json!(["openai:video"])), - encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-upstream-openai-video") - .expect("api key should encrypt"), - None, - None, - Some(json!({"openai:video": 1})), - None, - None, - None, - None, + Some(json!(["sora-2"])), + 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!(["sora-2"])), ) - .expect("key transport should build") + .expect("auth snapshot should build") } let decision_stream_hits = Arc::new(Mutex::new(0usize)); @@ -112,7 +76,6 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with let seen_execution_runtime_stream = Arc::new(Mutex::new(None::)); let seen_execution_runtime_stream_clone = Arc::clone(&seen_execution_runtime_stream); - let upstream = Router::new() .route( "/api/internal/gateway/resolve", @@ -124,11 +87,6 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with "route_kind": "video", "auth_endpoint_signature": "openai:video", "execution_runtime_candidate": true, - "auth_context": { - "user_id": "user-video-content-local-123", - "api_key_id": "key-video-content-local-123", - "access_allowed": true - }, "public_path": request.uri().path() })) }), @@ -299,28 +257,81 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with }) .await .expect("upsert should succeed"); - let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( - vec![sample_provider()], - vec![sample_endpoint()], - vec![sample_key()], - )); + let provider_catalog_repository = video_provider_catalog_repository( + "provider-openai-video-content-followup-1", + "openai", + "endpoint-openai-video-content-followup-1", + "openai:video", + "https://api.openai.example/v1", + "key-openai-video-content-followup-1", + "sk-upstream-openai-video", + ); + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![ + ( + Some(hash_api_key("client-video-content-foreign-key")), + sample_auth_snapshot( + "key-video-content-foreign-123", + "user-video-content-foreign-123", + ), + ), + ( + Some(hash_api_key("client-video-content-owner-key")), + sample_auth_snapshot( + "key-video-content-local-rotated-123", + "user-video-content-local-123", + ), + ), + ])); let gateway = build_router_with_state( build_state_with_execution_runtime_override(execution_runtime_url) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) - .with_video_task_repository_and_provider_transport_for_tests( - repository, - provider_catalog_repository, - DEVELOPMENT_ENCRYPTION_KEY, - ), + .with_data_state_for_tests( + crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests( + repository, + provider_catalog_repository, + DEVELOPMENT_ENCRYPTION_KEY, + ) + .with_auth_api_key_reader(auth_repository), + ) ); let (gateway_url, gateway_handle) = start_server(gateway).await; - let response = reqwest::Client::new() + let client = reqwest::Client::new(); + let foreign_response = client .get(format!( "{gateway_url}/v1/videos/task-content-local-123/content?variant=video" )) .header(CONTROL_EXECUTE_FALLBACK_HEADER, "true") + .bearer_auth("client-video-content-foreign-key") + .header( + TRACE_ID_HEADER, + "trace-openai-video-content-foreign-local-123", + ) + .send() + .await + .expect("foreign content request should complete"); + let foreign_status = foreign_response.status(); + let foreign_body = foreign_response.text().await.expect("body should read"); + assert_eq!( + foreign_status, + StatusCode::NOT_FOUND, + "unexpected foreign response body: {foreign_body}" + ); + assert!(seen_execution_runtime_stream + .lock() + .expect("mutex should lock") + .is_none()); + assert_eq!(*decision_stream_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*execute_stream_hits.lock().expect("mutex should lock"), 0); + assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); + + let response = client + .get(format!( + "{gateway_url}/v1/videos/task-content-local-123/content?variant=video" + )) + .header(CONTROL_EXECUTE_FALLBACK_HEADER, "true") + .bearer_auth("client-video-content-owner-key") .header(TRACE_ID_HEADER, "trace-openai-video-content-local-123") .send() .await @@ -347,7 +358,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with assert_eq!(seen_stream_request.method, "GET"); assert_eq!( seen_stream_request.url, - "https://cdn.example.com/video-content.mp4" + "https://api.openai.example/v1/videos/ext-video-content-followup-123/content" ); assert_eq!(*decision_stream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*execute_stream_hits.lock().expect("mutex should lock"), 0); diff --git a/apps/aether-gateway/src/tunnel/embedded/control_plane.rs b/apps/aether-gateway/src/tunnel/embedded/control_plane.rs index 580a6abb8..aaf086665 100644 --- a/apps/aether-gateway/src/tunnel/embedded/control_plane.rs +++ b/apps/aether-gateway/src/tunnel/embedded/control_plane.rs @@ -1,12 +1,29 @@ -use aether_http::{build_http_client, HttpClientConfig}; +use aether_http::{apply_http_client_config, HttpClientConfig}; use futures_util::future::BoxFuture; use reqwest::Client; use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; + +use aether_contracts::tunnel_security::{ + sign_tunnel_control_plane_request_for_generation, TUNNEL_CONTROL_PLANE_GENERATION_HEADER, + TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, TUNNEL_CONTROL_PLANE_NONCE_HEADER, + TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, +}; +use aether_gateway_tunnel::{TUNNEL_HEARTBEAT_PATH, TUNNEL_NODE_STATUS_PATH}; + +use super::hub::ProxyConn; + +const MAX_CONTROL_PLANE_RESPONSE_BYTES: usize = 256 * 1024; +const MAX_CONTROL_PLANE_BASE_URL_BYTES: usize = 2 * 1024; +pub(crate) const CONTROL_PLANE_CREDENTIAL_REVOKED: &str = "proxy tunnel credential revoked"; +pub(crate) const CONTROL_PLANE_CREDENTIAL_UNAVAILABLE: &str = + "proxy tunnel credential validation unavailable"; type HeartbeatAckCallback = - dyn Fn(Vec) -> BoxFuture<'static, Result, String>> + Send + Sync; -type NodeStatusCallback = - dyn Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + Send + Sync; + dyn Fn(Arc, Vec) -> BoxFuture<'static, Result, String>> + Send + Sync; +type NodeStatusCallback = dyn Fn(Arc, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + + Send + + Sync; enum ControlPlaneMode { Disabled, @@ -27,11 +44,32 @@ pub struct ControlPlaneClient { impl ControlPlaneClient { pub fn new(base_url: String) -> Self { - let client = build_http_client(&HttpClientConfig { - request_timeout_ms: Some(10_000), - user_agent: Some("aether-tunnel-standalone/control-plane".to_string()), - ..HttpClientConfig::default() - }) + // The control-plane base URL is operator supplied, but it is used for + // requests carrying a tunnel credential. Reject URL-controlled + // request components (userinfo/query/fragment) before concatenating + // endpoint paths; otherwise a typo such as `?token=...` can leak + // credentials or change the signed request target. Keep the + // established support for HTTP and private deployment hosts—the + // standalone tunnel commonly talks to an in-cluster gateway. + let Some(base_url) = normalize_control_plane_base_url(&base_url) else { + return Self { + inner: Arc::new(ControlPlaneMode::Http { + client: None, + base_url: String::new(), + }), + }; + }; + let client = apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &HttpClientConfig { + request_timeout_ms: Some(10_000), + user_agent: Some("aether-tunnel-standalone/control-plane".to_string()), + ..HttpClientConfig::default() + }, + ) + .build() .ok(); Self { inner: Arc::new(ControlPlaneMode::Http { client, base_url }), @@ -49,9 +87,11 @@ impl ControlPlaneClient { push_node_status: PushNodeStatus, ) -> Self where - HeartbeatAck: - Fn(Vec) -> BoxFuture<'static, Result, String>> + Send + Sync + 'static, - PushNodeStatus: Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + HeartbeatAck: Fn(Arc, Vec) -> BoxFuture<'static, Result, String>> + + Send + + Sync + + 'static, + PushNodeStatus: Fn(Arc, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + Send + Sync + 'static, @@ -64,88 +104,432 @@ impl ControlPlaneClient { } } - pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result, String> { + pub async fn heartbeat_ack( + &self, + authenticated_node_id: &str, + authenticated_key: Option<&str>, + authenticated_generation: &str, + payload: &[u8], + ) -> Result, String> { match self.inner.as_ref() { ControlPlaneMode::Disabled => Ok(b"{}".to_vec()), - ControlPlaneMode::Http { client, base_url } => { - let Some(client) = client else { - return Ok(b"{}".to_vec()); - }; - let url = format!( - "{}/api/internal/tunnel/heartbeat", - base_url.trim_end_matches('/') - ); - let response = client - .post(&url) - .header("content-type", "application/json") - .body(payload.to_vec()) - .send() - .await - .map_err(|e| format!("heartbeat callback request failed: {e}"))?; - if !response.status().is_success() { - return Err(format!( - "heartbeat callback failed with status {}", - response.status() - )); - } - response - .bytes() - .await - .map(|bytes| bytes.to_vec()) - .map_err(|e| format!("heartbeat callback body read failed: {e}")) + ControlPlaneMode::Http { .. } => { + self.heartbeat_ack_http( + authenticated_node_id, + authenticated_key, + authenticated_generation, + payload, + ) + .await + } + ControlPlaneMode::Local { .. } => { + Err("local heartbeat callback requires connection credential binding".to_string()) } - ControlPlaneMode::Local { heartbeat_ack, .. } => heartbeat_ack(payload.to_vec()).await, } } + pub async fn heartbeat_ack_for_connection( + &self, + connection: Arc, + payload: &[u8], + ) -> Result, String> { + match self.inner.as_ref() { + ControlPlaneMode::Disabled => Ok(b"{}".to_vec()), + ControlPlaneMode::Http { .. } => { + self.heartbeat_ack_http( + &connection.node_id, + connection.authenticated_key.as_deref(), + &connection.node_generation, + payload, + ) + .await + } + ControlPlaneMode::Local { heartbeat_ack, .. } => { + heartbeat_ack(connection, payload.to_vec()).await + } + } + } + + async fn heartbeat_ack_http( + &self, + authenticated_node_id: &str, + authenticated_key: Option<&str>, + authenticated_generation: &str, + payload: &[u8], + ) -> Result, String> { + let ControlPlaneMode::Http { client, base_url } = self.inner.as_ref() else { + return Err("HTTP heartbeat callback is unavailable".to_string()); + }; + let Some(client) = client else { + return Err("heartbeat callback HTTP client is unavailable".to_string()); + }; + let authenticated_key = authenticated_key + .ok_or_else(|| "heartbeat callback is missing authenticated tunnel key".to_string())?; + let url = format!("{}{TUNNEL_HEARTBEAT_PATH}", base_url.trim_end_matches('/')); + let request = client + .post(&url) + .header("content-type", "application/json") + .body(payload.to_vec()); + let response = sign_control_plane_request( + request, + authenticated_key, + TUNNEL_HEARTBEAT_PATH, + authenticated_node_id, + authenticated_generation, + payload, + )? + .send() + .await + .map_err(|error| { + format!( + "heartbeat callback request failed ({})", + control_plane_reqwest_error_kind(&error) + ) + })?; + if response.status() == reqwest::StatusCode::FORBIDDEN { + return Err(CONTROL_PLANE_CREDENTIAL_REVOKED.to_string()); + } + if !response.status().is_success() { + return Err(format!( + "heartbeat callback failed with status {}", + response.status() + )); + } + aether_http::read_response_bytes_with_limit(response, MAX_CONTROL_PLANE_RESPONSE_BYTES) + .await + .map_err(|e| format!("heartbeat callback body read failed: {e}")) + } + pub async fn push_node_status( &self, node_id: &str, + authenticated_key: Option<&str>, + authenticated_generation: &str, connected: bool, conn_count: usize, observed_at_unix_secs: u64, ) -> Result<(), String> { match self.inner.as_ref() { ControlPlaneMode::Disabled => Ok(()), - ControlPlaneMode::Http { client, base_url } => { - let Some(client) = client else { - return Ok(()); - }; - let url = format!( - "{}/api/internal/tunnel/node-status", - base_url.trim_end_matches('/') - ); - let response = client - .post(&url) - .json(&serde_json::json!({ - "node_id": node_id, - "connected": connected, - "conn_count": conn_count, - "observed_at_unix_secs": observed_at_unix_secs, - })) - .send() - .await - .map_err(|e| format!("node-status callback request failed: {e}"))?; - if response.status().is_success() { - Ok(()) - } else { - Err(format!( - "node-status callback failed with status {}", - response.status() - )) - } - } - ControlPlaneMode::Local { - push_node_status, .. - } => { - push_node_status( - node_id.to_string(), + ControlPlaneMode::Http { .. } => { + self.push_node_status_http( + node_id, + authenticated_key, + authenticated_generation, connected, conn_count, observed_at_unix_secs, ) .await } + ControlPlaneMode::Local { .. } => { + Err("local node-status callback requires connection credential binding".to_string()) + } + } + } + + pub async fn push_node_status_for_connection( + &self, + connection: Arc, + connected: bool, + conn_count: usize, + observed_at_unix_secs: u64, + ) -> Result<(), String> { + match self.inner.as_ref() { + ControlPlaneMode::Disabled => Ok(()), + ControlPlaneMode::Http { .. } => { + self.push_node_status_http( + &connection.node_id, + connection.authenticated_key.as_deref(), + &connection.node_generation, + connected, + conn_count, + observed_at_unix_secs, + ) + .await + } + ControlPlaneMode::Local { + push_node_status, .. + } => push_node_status(connection, connected, conn_count, observed_at_unix_secs).await, + } + } + + async fn push_node_status_http( + &self, + node_id: &str, + authenticated_key: Option<&str>, + authenticated_generation: &str, + connected: bool, + conn_count: usize, + observed_at_unix_secs: u64, + ) -> Result<(), String> { + let ControlPlaneMode::Http { client, base_url } = self.inner.as_ref() else { + return Err("HTTP node-status callback is unavailable".to_string()); + }; + let Some(client) = client else { + return Err("node-status callback HTTP client is unavailable".to_string()); + }; + let authenticated_key = authenticated_key.ok_or_else(|| { + "node-status callback is missing authenticated tunnel key".to_string() + })?; + let url = format!( + "{}{TUNNEL_NODE_STATUS_PATH}", + base_url.trim_end_matches('/') + ); + let payload = serde_json::to_vec(&serde_json::json!({ + "node_id": node_id, + "connected": connected, + "conn_count": conn_count, + "observed_at_unix_secs": observed_at_unix_secs, + })) + .map_err(|e| format!("node-status callback serialization failed: {e}"))?; + let request = client + .post(&url) + .header("content-type", "application/json") + .body(payload.clone()); + let response = sign_control_plane_request( + request, + authenticated_key, + TUNNEL_NODE_STATUS_PATH, + node_id, + authenticated_generation, + &payload, + )? + .send() + .await + .map_err(|error| { + format!( + "node-status callback request failed ({})", + control_plane_reqwest_error_kind(&error) + ) + })?; + if response.status() == reqwest::StatusCode::FORBIDDEN { + return Err(CONTROL_PLANE_CREDENTIAL_REVOKED.to_string()); + } + if response.status().is_success() { + Ok(()) + } else { + Err(format!( + "node-status callback failed with status {}", + response.status() + )) } } } + +fn normalize_control_plane_base_url(raw: &str) -> Option { + let raw = raw.trim(); + if raw.is_empty() + || raw.len() > MAX_CONTROL_PLANE_BASE_URL_BYTES + || raw.bytes().any(|byte| byte == 0 || byte.is_ascii_control()) + { + return None; + } + let parsed = url::Url::parse(raw).ok()?; + if !matches!(parsed.scheme(), "http" | "https") + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return None; + } + Some(raw.trim_end_matches('/').to_string()) +} + +pub(crate) fn is_credential_revoked_error(error: &str) -> bool { + error == CONTROL_PLANE_CREDENTIAL_REVOKED +} + +/// Return a stable transport category without rendering reqwest's error. +/// +/// `reqwest::Error`'s `Display` implementation may include the complete URL +/// (including path/query components). Control-plane errors are logged by the +/// tunnel hub, so forwarding that value could disclose operator deployment +/// details or credentials embedded in a path. Callers should use this helper +/// whenever a control-plane request fails. +fn control_plane_reqwest_error_kind(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else if error.is_body() { + "body" + } else if error.is_decode() { + "decode" + } else { + "transport" + } +} + +fn sign_control_plane_request( + request: reqwest::RequestBuilder, + authenticated_key: &str, + path: &str, + node_id: &str, + tunnel_generation: &str, + body: &[u8], +) -> Result { + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map_err(|_| "system clock is before the Unix epoch".to_string())? + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let signature = sign_tunnel_control_plane_request_for_generation( + authenticated_key, + "POST", + path, + node_id, + tunnel_generation, + timestamp, + &nonce, + body, + ) + .map_err(|error| format!("invalid authenticated tunnel key: {error}"))?; + Ok(request + .header(TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, node_id) + .header(TUNNEL_CONTROL_PLANE_GENERATION_HEADER, tunnel_generation) + .header(TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, timestamp) + .header(TUNNEL_CONTROL_PLANE_NONCE_HEADER, nonce) + .header(TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, signature)) +} + +#[cfg(test)] +mod tests { + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + + use axum::{ + body::Body, + http::{header, Response, StatusCode}, + routing::post, + Router, + }; + use base64::Engine as _; + + use super::{ + control_plane_reqwest_error_kind, normalize_control_plane_base_url, ControlPlaneClient, + }; + use aether_gateway_tunnel::{TUNNEL_HEARTBEAT_PATH, TUNNEL_NODE_STATUS_PATH}; + + #[test] + fn control_plane_base_url_rejects_credential_and_request_components() { + for value in [ + "https://user:secret@gateway.example", + "https://gateway.example?token=secret", + "https://gateway.example/control#fragment", + "file:///tmp/gateway", + "", + ] { + assert!( + normalize_control_plane_base_url(value).is_none(), + "unsafe control-plane URL should be rejected: {value:?}" + ); + } + } + + #[test] + fn control_plane_base_url_preserves_deployment_path_and_trims_slashes() { + assert_eq!( + normalize_control_plane_base_url(" https://gateway.example/control/// "), + Some("https://gateway.example/control".to_string()) + ); + assert_eq!( + normalize_control_plane_base_url("http://127.0.0.1:8084/"), + Some("http://127.0.0.1:8084".to_string()) + ); + } + + #[tokio::test] + async fn control_plane_transport_errors_do_not_render_request_url() { + let error = reqwest::Client::new() + .get("ftp://user:secret@example.invalid/control-plane") + .send() + .await + .expect_err("unsupported control-plane scheme should fail before a request"); + let rendered = format!( + "control-plane request failed ({})", + control_plane_reqwest_error_kind(&error) + ); + assert!(!rendered.contains("secret")); + assert!(!rendered.contains("example.invalid")); + assert!(rendered.starts_with("control-plane request failed (")); + } + + #[tokio::test] + async fn signed_control_plane_requests_never_follow_redirects() { + let redirected_hits = Arc::new(AtomicUsize::new(0)); + let redirected_hits_for_route = Arc::clone(&redirected_hits); + let redirected_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect target listener should bind"); + let redirected_addr = redirected_listener + .local_addr() + .expect("redirect target address should resolve"); + let redirected_app = Router::new().fallback(move || { + let hits = Arc::clone(&redirected_hits_for_route); + async move { + hits.fetch_add(1, Ordering::SeqCst); + StatusCode::OK + } + }); + let redirected_server = tokio::spawn(async move { + axum::serve(redirected_listener, redirected_app) + .await + .expect("redirect target server should run"); + }); + + let source_listener = crate::test_support::bind_loopback_listener() + .await + .expect("redirect source listener should bind"); + let source_addr = source_listener + .local_addr() + .expect("redirect source address should resolve"); + let location = format!("http://{redirected_addr}/captured"); + let redirect = move || { + let location = location.clone(); + async move { + Response::builder() + .status(StatusCode::TEMPORARY_REDIRECT) + .header(header::LOCATION, location) + .body(Body::empty()) + .expect("redirect response should build") + } + }; + let source_app = Router::new() + .route(TUNNEL_HEARTBEAT_PATH, post(redirect.clone())) + .route(TUNNEL_NODE_STATUS_PATH, post(redirect)); + let source_server = tokio::spawn(async move { + axum::serve(source_listener, source_app) + .await + .expect("redirect source server should run"); + }); + + let client = ControlPlaneClient::new(format!("http://{source_addr}")); + let key = base64::engine::general_purpose::STANDARD.encode([7_u8; 32]); + let heartbeat = client + .heartbeat_ack( + "node-1", + Some(&key), + "generation-1", + br#"{"node_id":"node-1"}"#, + ) + .await + .expect_err("heartbeat redirect should be returned as an error"); + assert!(heartbeat.contains("307 Temporary Redirect")); + let status = client + .push_node_status("node-1", Some(&key), "generation-1", true, 1, 1) + .await + .expect_err("node-status redirect should be returned as an error"); + assert!(status.contains("307 Temporary Redirect")); + assert_eq!(redirected_hits.load(Ordering::SeqCst), 0); + + source_server.abort(); + redirected_server.abort(); + } +} diff --git a/apps/aether-gateway/src/tunnel/embedded/hub.rs b/apps/aether-gateway/src/tunnel/embedded/hub.rs index f5e8e344d..69101027c 100644 --- a/apps/aether-gateway/src/tunnel/embedded/hub.rs +++ b/apps/aether-gateway/src/tunnel/embedded/hub.rs @@ -1,4 +1,5 @@ -use std::collections::{BTreeMap, HashMap}; +use std::collections::{BTreeMap, HashMap, HashSet}; +use std::net::IpAddr; use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicU8, AtomicUsize, Ordering}; use std::sync::{Arc, LazyLock}; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; @@ -19,6 +20,7 @@ use super::control_plane::ControlPlaneClient; use super::protocol; const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024; +const MAX_TUNNEL_CONTROL_PAYLOAD_SIZE: usize = 256 * 1024; const SOFT_AVOID_QUEUE_PRESSURE_PERCENT: u64 = 50; const SOFT_AVOID_STREAM_PRESSURE_PERCENT: u64 = 85; const OUTBOUND_BACKPRESSURE_TIMEOUT: Duration = Duration::from_secs(5); @@ -213,6 +215,9 @@ pub struct ProxyConn { pub id: u64, pub node_id: String, pub node_name: String, + pub node_generation: String, + pub authenticated_key: Option, + pub(crate) management_token_credential: Option, pub outbound: BoundedOutbound, next_stream_id: AtomicU32, pub stream_count: AtomicUsize, @@ -242,6 +247,9 @@ impl ProxyConn { id, node_id, node_name, + node_generation: String::new(), + authenticated_key: None, + management_token_credential: None, outbound: BoundedOutbound::new(tx, close_tx), next_stream_id: AtomicU32::new(2), stream_count: AtomicUsize::new(0), @@ -258,6 +266,37 @@ impl ProxyConn { } } + pub fn with_authenticated_key(mut self, authenticated_key: String) -> Self { + self.authenticated_key = Some(authenticated_key); + self + } + + pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self { + self.node_generation = tunnel_generation; + self + } + + pub(crate) fn with_management_token_credential( + mut self, + credential: ProxyManagementTokenCredential, + ) -> Self { + self.management_token_credential = Some(credential); + self + } + + pub(crate) fn credential_binding(&self) -> Option { + match ( + self.authenticated_key.as_deref(), + self.management_token_credential.as_ref(), + ) { + (Some(key), None) => Some(ProxyCredentialBinding::Psk(key.to_string())), + (None, Some(credential)) => { + Some(ProxyCredentialBinding::ManagementToken(credential.clone())) + } + _ => None, + } + } + pub fn record_write_latency(&self, elapsed: std::time::Duration) { let micros = u64::try_from(elapsed.as_micros()).unwrap_or(u64::MAX); self.write_latency_last_us.store(micros, Ordering::Relaxed); @@ -440,6 +479,44 @@ impl ProxyConn { } } +#[derive(Clone)] +pub(crate) struct ProxyManagementTokenCredential { + pub(crate) verified_token_hash: crate::management_token_auth::VerifiedManagementTokenHash, + pub(crate) token_id: String, + pub(crate) user_id: String, + pub(crate) remote_ip: IpAddr, +} + +#[derive(Clone)] +pub(crate) enum ProxyCredentialBinding { + Psk(String), + ManagementToken(ProxyManagementTokenCredential), +} + +impl std::fmt::Debug for ProxyCredentialBinding { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Psk(_) => formatter.write_str("ProxyCredentialBinding::Psk([REDACTED])"), + Self::ManagementToken(credential) => formatter + .debug_tuple("ProxyCredentialBinding::ManagementToken") + .field(credential) + .finish(), + } + } +} + +impl std::fmt::Debug for ProxyManagementTokenCredential { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProxyManagementTokenCredential") + .field("verified_token_hash", &"[REDACTED]") + .field("token_id", &self.token_id) + .field("user_id", &self.user_id) + .field("remote_ip", &self.remote_ip) + .finish() + } +} + #[derive(Debug, Clone, Copy)] struct ProxyConnSnapshot { conn_id: u64, @@ -503,6 +580,7 @@ struct LocalWaitState { pub struct LocalStream { pub id: u64, + tunnel_generation: String, proxy_conn_id: u64, proxy_stream_id: u32, request_window: StreamFlowWindow, @@ -515,10 +593,17 @@ pub struct LocalStream { } impl LocalStream { - fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32, initial_window_bytes: u32) -> Self { + fn new( + id: u64, + tunnel_generation: String, + proxy_conn_id: u64, + proxy_stream_id: u32, + initial_window_bytes: u32, + ) -> Self { let (body_tx, body_rx) = mpsc::channel(128); Self { id, + tunnel_generation, proxy_conn_id, proxy_stream_id, request_window: StreamFlowWindow::new(initial_window_bytes), @@ -531,6 +616,10 @@ impl LocalStream { } } + pub(crate) fn tunnel_generation(&self) -> &str { + &self.tunnel_generation + } + async fn acquire_request_window( &self, bytes: usize, @@ -676,9 +765,11 @@ pub struct HubRouter { drain_reasons: Mutex>, } -#[derive(Debug)] struct NodeStatusEvent { node_id: String, + authenticated_key: Option, + tunnel_generation: String, + connection: Option>, connected: bool, conn_count: usize, observed_at_unix_secs: u64, @@ -692,15 +783,37 @@ impl HubRouter { if let Ok(handle) = tokio::runtime::Handle::try_current() { handle.spawn(async move { while let Some(event) = node_status_rx.recv().await { - if let Err(error) = worker_control_plane - .push_node_status( - &event.node_id, - event.connected, - event.conn_count, - event.observed_at_unix_secs, - ) - .await - { + let connection = event.connection.clone(); + let result = match connection { + Some(connection) => { + worker_control_plane + .push_node_status_for_connection( + connection, + event.connected, + event.conn_count, + event.observed_at_unix_secs, + ) + .await + } + None => { + worker_control_plane + .push_node_status( + &event.node_id, + event.authenticated_key.as_deref(), + &event.tunnel_generation, + event.connected, + event.conn_count, + event.observed_at_unix_secs, + ) + .await + } + }; + if let Err(error) = result { + if super::control_plane::is_credential_revoked_error(&error) { + if let Some(connection) = event.connection.as_ref() { + connection.request_close(); + } + } warn!( node_id = %event.node_id, connected = event.connected, @@ -746,7 +859,7 @@ impl HubRouter { let healthy_count = { let mut map = self.proxy_conns.write(); - map.entry(node_id.clone()).or_default().push(conn); + map.entry(node_id.clone()).or_default().push(conn.clone()); available_conn_count(map.get(&node_id).map(Vec::as_slice).unwrap_or(&[])) }; @@ -758,11 +871,23 @@ impl HubRouter { "proxy connected" ); - self.notify_node_status(node_id, healthy_count > 0, healthy_count); + self.notify_node_status( + node_id, + conn.authenticated_key.clone(), + Some(conn), + healthy_count > 0, + healthy_count, + ); } pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) { - self.proxy_conns_by_id.remove(&conn_id); + let disconnected_connection = self + .proxy_conns_by_id + .remove(&conn_id) + .map(|(_, connection)| connection); + let disconnected_authenticated_key = disconnected_connection + .as_ref() + .and_then(|connection| connection.authenticated_key.clone()); let healthy_count = { let mut map = self.proxy_conns.write(); @@ -785,7 +910,20 @@ impl HubRouter { ); self.cancel_streams_for_proxy(conn_id); - self.notify_node_status(node_id.to_string(), healthy_count > 0, healthy_count); + let authenticated_key = self + .proxy_conns + .read() + .get(node_id) + .and_then(|connections| connections.first()) + .and_then(|connection| connection.authenticated_key.clone()) + .or(disconnected_authenticated_key); + self.notify_node_status( + node_id.to_string(), + authenticated_key, + disconnected_connection, + healthy_count > 0, + healthy_count, + ); } pub fn request_close_all_proxies(&self) -> usize { @@ -801,9 +939,52 @@ impl HubRouter { total } - fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) { + pub(crate) fn request_close_proxy(&self, conn_id: u64) -> bool { + let Some(conn) = self + .proxy_conns_by_id + .get(&conn_id) + .map(|entry| Arc::clone(entry.value())) + else { + return false; + }; + conn.request_close(); + true + } + + pub(crate) fn request_close_proxies_for_node(&self, node_id: &str) -> usize { + let conns = self.proxy_connections_for_node(node_id); + let total = conns.len(); + for conn in conns { + conn.request_close(); + } + total + } + + pub(crate) fn proxy_connections_for_node(&self, node_id: &str) -> Vec> { + self.proxy_conns + .read() + .get(node_id) + .map(|connections| connections.to_vec()) + .unwrap_or_default() + } + + fn notify_node_status( + &self, + node_id: String, + authenticated_key: Option, + connection: Option>, + connected: bool, + conn_count: usize, + ) { + let tunnel_generation = connection + .as_ref() + .map(|connection| connection.node_generation.clone()) + .unwrap_or_default(); let event = NodeStatusEvent { node_id, + authenticated_key, + tunnel_generation, + connection, connected, conn_count, observed_at_unix_secs: current_unix_secs(), @@ -837,13 +1018,24 @@ impl HubRouter { } fn notify_current_node_status(&self, node_id: &str) { - let healthy_count = { + let (healthy_count, connection, authenticated_key) = { let map = self.proxy_conns.read(); - map.get(node_id) - .map(|v| available_conn_count(v.as_slice())) - .unwrap_or(0) + let connections = map.get(node_id).map(Vec::as_slice).unwrap_or(&[]); + ( + available_conn_count(connections), + connections.first().cloned(), + connections + .first() + .and_then(|connection| connection.authenticated_key.clone()), + ) }; - self.notify_node_status(node_id.to_string(), healthy_count > 0, healthy_count); + self.notify_node_status( + node_id.to_string(), + authenticated_key, + connection, + healthy_count > 0, + healthy_count, + ); } fn record_stream_reset(&self, reason: &str) { @@ -856,7 +1048,11 @@ impl HubRouter { increment_reason(&self.drain_reasons, reason); } - fn ranked_proxy_conn_candidates(&self, node_id: &str) -> Vec { + fn ranked_proxy_conn_candidates( + &self, + node_id: &str, + authorized_conn_ids: Option<&HashSet>, + ) -> Vec { let conns = { let map = self.proxy_conns.read(); map.get(node_id) @@ -866,6 +1062,9 @@ impl HubRouter { let mut candidates = conns .into_iter() .filter_map(|conn| { + if authorized_conn_ids.is_some_and(|allowed| !allowed.contains(&conn.id)) { + return None; + } let snapshot = conn.snapshot(); snapshot .available @@ -877,7 +1076,7 @@ impl HubRouter { } pub fn has_local_proxy(&self, node_id: &str) -> bool { - !self.ranked_proxy_conn_candidates(node_id).is_empty() + !self.ranked_proxy_conn_candidates(node_id, None).is_empty() } pub async fn open_local_stream( @@ -885,7 +1084,17 @@ impl HubRouter { node_id: &str, meta: &protocol::RequestMeta, ) -> Result, String> { - let candidates = self.ranked_proxy_conn_candidates(node_id); + self.open_local_stream_with_authorized_connections(node_id, meta, None) + .await + } + + pub(crate) async fn open_local_stream_with_authorized_connections( + &self, + node_id: &str, + meta: &protocol::RequestMeta, + authorized_conn_ids: Option<&HashSet>, + ) -> Result, String> { + let candidates = self.ranked_proxy_conn_candidates(node_id, authorized_conn_ids); if candidates.is_empty() { self.selection_unavailable_total .fetch_add(1, Ordering::Relaxed); @@ -959,6 +1168,7 @@ impl HubRouter { let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed); let local_stream = Arc::new(LocalStream::new( local_stream_id, + proxy_conn.node_generation.clone(), proxy_conn.id, proxy_stream_id, *STREAM_INITIAL_WINDOW_BYTES, @@ -1160,8 +1370,14 @@ impl HubRouter { Some(h) => h, None => return, }; - let expected_len = protocol::HEADER_SIZE + header.payload_len as usize; - if data.len() < expected_len { + let Some(expected_len) = protocol::HEADER_SIZE.checked_add(header.payload_len as usize) + else { + return; + }; + if data.len() != expected_len { + if header.stream_id != 0 { + self.fail_proxy_stream(proxy_conn_id, header.stream_id, "invalid frame length"); + } return; } @@ -1176,22 +1392,32 @@ impl HubRouter { self.finish_proxy_stream(proxy_conn_id, header.stream_id); } protocol::STREAM_ERROR => { - let message = protocol::decode_payload(data, &header) - .ok() - .and_then(|payload| String::from_utf8(payload).ok()) - .unwrap_or_else(|| "stream error".to_string()); + let raw_message = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| String::from_utf8(payload).ok()) + .unwrap_or_else(|| "stream error".to_string()); + let message = safe_peer_stream_error(&raw_message); self.record_stream_reset(&message); self.fail_proxy_stream(proxy_conn_id, header.stream_id, message); } protocol::RESET_STREAM => { - let message = protocol::decode_payload(data, &header) - .ok() - .and_then(|payload| { - serde_json::from_slice::(&payload) - .ok() - .map(|payload| payload.reason) - }) - .unwrap_or_else(|| "stream reset".to_string()); + let raw_message = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| { + serde_json::from_slice::(&payload) + .ok() + .map(|payload| payload.reason) + }) + .unwrap_or_else(|| "stream reset".to_string()); + let message = safe_peer_stream_error(&raw_message); self.record_stream_reset(&message); self.fail_proxy_stream(proxy_conn_id, header.stream_id, message); } @@ -1217,17 +1443,19 @@ impl HubRouter { if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) { let first = pc.mark_draining(); if first { - let drain = - protocol::decode_payload(data, &header) - .ok() - .and_then(|payload| { - if payload.is_empty() { - None - } else { - serde_json::from_slice::(&payload) - .ok() - } - }); + let drain = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| { + if payload.is_empty() { + None + } else { + serde_json::from_slice::(&payload).ok() + } + }); let reason = drain .as_ref() .map(|payload| payload.reason.as_str()) @@ -1262,12 +1490,13 @@ impl HubRouter { } } protocol::HELLO => { - if let Some(payload) = - protocol::decode_payload(data, &header) - .ok() - .and_then(|payload| { - serde_json::from_slice::(&payload).ok() - }) + if let Some(payload) = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| serde_json::from_slice::(&payload).ok()) { if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) { pc.update_protocol_version(payload.protocol_version); @@ -1290,13 +1519,15 @@ impl HubRouter { self.handle_window_update(proxy_conn_id, header.stream_id, data, &header); } protocol::LOAD_REPORT => { - if let Some(payload) = - protocol::decode_payload(data, &header) - .ok() - .and_then(|payload| { - serde_json::from_slice::(&payload).ok() - }) - { + if let Some(payload) = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| { + serde_json::from_slice::(&payload).ok() + }) { if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) { pc.update_remote_health_score(payload.health_score); } @@ -1332,12 +1563,13 @@ impl HubRouter { data: &[u8], header: &protocol::FrameHeader, ) { - let Some(delta) = protocol::decode_payload(data, header) - .ok() - .and_then(|payload| { - serde_json::from_slice::(&payload).ok() - }) - .map(|payload| payload.delta_bytes) + let Some(delta) = + protocol::decode_payload_with_limit(data, header, MAX_TUNNEL_CONTROL_PAYLOAD_SIZE) + .ok() + .and_then(|payload| { + serde_json::from_slice::(&payload).ok() + }) + .map(|payload| payload.delta_bytes) else { return; }; @@ -1457,7 +1689,9 @@ impl HubRouter { let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else { return; }; - let Ok(payload) = protocol::decode_payload(data, &header) else { + let Ok(payload) = + protocol::decode_payload_with_limit(data, &header, MAX_TUNNEL_CONTROL_PAYLOAD_SIZE) + else { self.fail_proxy_stream( proxy_conn_id, header.stream_id, @@ -1487,7 +1721,11 @@ impl HubRouter { let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else { return; }; - let Ok(payload) = protocol::decode_payload(data, &header) else { + let Ok(payload) = protocol::decode_payload_with_limit( + data, + &header, + aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES, + ) else { self.fail_proxy_stream( proxy_conn_id, header.stream_id, @@ -1567,16 +1805,70 @@ impl HubRouter { data: &[u8], header: &protocol::FrameHeader, ) { - let payload = match protocol::decode_payload(data, header) { + let payload = match protocol::decode_payload_with_limit( + data, + header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) { Ok(payload) => payload, Err(error) => { warn!(proxy_conn_id = proxy_conn_id, error = %error, "failed to decode heartbeat payload"); return; } }; - let ack_payload = match self.control_plane.heartbeat_ack(&payload).await { + let Some(authenticated_node_id) = self + .proxy_conns_by_id + .get(&proxy_conn_id) + .map(|entry| entry.node_id.clone()) + else { + warn!( + proxy_conn_id, + "heartbeat rejected for unregistered proxy connection" + ); + return; + }; + let payload_node_id = match heartbeat_payload_node_id(&payload) { + Ok(node_id) => node_id, + Err(error) => { + warn!( + proxy_conn_id, + authenticated_node_id = %authenticated_node_id, + error = %error, + "heartbeat rejected before control-plane dispatch" + ); + return; + } + }; + if payload_node_id != authenticated_node_id { + warn!( + proxy_conn_id, + authenticated_node_id = %authenticated_node_id, + payload_node_id = %payload_node_id, + "heartbeat rejected because node identity does not match tunnel authentication" + ); + return; + } + let connection = self + .proxy_conns_by_id + .get(&proxy_conn_id) + .map(|entry| Arc::clone(entry.value())); + let Some(connection) = connection else { + warn!( + proxy_conn_id, + "heartbeat rejected for unregistered proxy connection" + ); + return; + }; + let ack_payload = match self + .control_plane + .heartbeat_ack_for_connection(connection.clone(), &payload) + .await + { Ok(payload) => payload, Err(error) => { + if super::control_plane::is_credential_revoked_error(&error) { + connection.request_close(); + } warn!( proxy_conn_id = proxy_conn_id, error = %error, @@ -1755,6 +2047,18 @@ impl HubRouter { } } +fn heartbeat_payload_node_id(payload: &[u8]) -> Result { + let value: serde_json::Value = serde_json::from_slice(payload) + .map_err(|_| "heartbeat payload is not valid JSON".to_string())?; + let node_id = value + .get("node_id") + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| "heartbeat payload is missing node_id".to_string())?; + Ok(node_id.to_string()) +} + fn current_unix_secs() -> u64 { SystemTime::now() .duration_since(UNIX_EPOCH) @@ -1830,6 +2134,56 @@ fn metric_reason(raw: &str) -> String { } } +fn safe_peer_stream_error(raw: &str) -> String { + const CLASSIFICATION_PREFIX_BYTES: usize = 4 * 1024; + + let prefix = if raw.len() <= CLASSIFICATION_PREFIX_BYTES { + raw + } else { + let mut end = CLASSIFICATION_PREFIX_BYTES; + while !raw.is_char_boundary(end) { + end = end.saturating_sub(1); + } + &raw[..end] + }; + let category = if contains_ascii_case_insensitive(prefix, "timed out") + || contains_ascii_case_insensitive(prefix, "timeout") + { + "timeout" + } else if contains_ascii_case_insensitive(prefix, "overloaded") + || contains_ascii_case_insensitive(prefix, "backpressure") + || contains_ascii_case_insensitive(prefix, "congested") + || contains_ascii_case_insensitive(prefix, "window") + { + "overloaded" + } else if contains_ascii_case_insensitive(prefix, "forbidden") + || contains_ascii_case_insensitive(prefix, "unauthorized") + || contains_ascii_case_insensitive(prefix, "authentication") + { + "forbidden" + } else if contains_ascii_case_insensitive(prefix, "dns") { + "dns" + } else if contains_ascii_case_insensitive(prefix, "connect") + || contains_ascii_case_insensitive(prefix, "socket") + { + "connect" + } else if contains_ascii_case_insensitive(prefix, "cancel") + || contains_ascii_case_insensitive(prefix, "reset") + { + "reset" + } else { + "relay" + }; + format!("tunnel stream {category} error") +} + +fn contains_ascii_case_insensitive(haystack: &str, needle: &str) -> bool { + haystack + .as_bytes() + .windows(needle.len()) + .any(|candidate| candidate.eq_ignore_ascii_case(needle.as_bytes())) +} + #[derive(serde::Serialize)] pub struct HubStats { pub proxy_connections: usize, @@ -2098,13 +2452,36 @@ impl HubStats { mod tests { use aether_runtime::bounded_queue; - use super::{protocol, ControlPlaneClient, HubRouter, ProxyConn, MAX_REQUEST_BODY_FRAME_SIZE}; + use super::{ + protocol, safe_peer_stream_error, ControlPlaneClient, HubRouter, ProxyConn, + MAX_REQUEST_BODY_FRAME_SIZE, + }; use axum::extract::ws::Message; use bytes::Bytes; use std::collections::HashMap; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; use tokio::sync::watch; + #[test] + fn peer_stream_errors_are_projected_to_finite_categories() { + let sensitive = safe_peer_stream_error( + "request failed for https://alice:secret@10.0.0.8/private?token=query-secret\r\nx: y", + ); + assert_eq!(sensitive, "tunnel stream relay error"); + for secret in ["alice", "secret", "10.0.0.8", "private", "query", "x: y"] { + assert!(!sensitive.contains(secret), "leaked {secret}: {sensitive}"); + } + assert_eq!( + safe_peer_stream_error("upstream connect timeout: Bearer secret"), + "tunnel stream timeout error" + ); + assert_eq!( + safe_peer_stream_error("outbound backpressure timeout"), + "tunnel stream timeout error" + ); + } + fn build_meta() -> protocol::RequestMeta { protocol::RequestMeta { provider_id: None, @@ -2358,8 +2735,10 @@ mod tests { #[tokio::test] async fn heartbeat_callback_failure_does_not_send_fake_ack() { let hub = HubRouter::new(ControlPlaneClient::local( - |_payload| Box::pin(async { Err("db unavailable".to_string()) }), - |_node_id, _connected, _conn_count, _observed_at_unix_secs| Box::pin(async { Ok(()) }), + |_connection, _payload| Box::pin(async { Err("db unavailable".to_string()) }), + |_connection, _connected, _conn_count, _observed_at_unix_secs| { + Box::pin(async { Ok(()) }) + }, )); let (proxy_tx, mut proxy_rx) = bounded_queue(8); @@ -2386,6 +2765,45 @@ mod tests { assert!(proxy_rx.try_recv().is_err()); } + #[tokio::test] + async fn heartbeat_for_another_node_is_rejected_before_callback_and_ack() { + let callback_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls_for_heartbeat = Arc::clone(&callback_calls); + let hub = HubRouter::new(ControlPlaneClient::local( + move |_connection, _payload| { + callback_calls_for_heartbeat.fetch_add(1, Ordering::Relaxed); + Box::pin(async { Ok(br#"{"heartbeat_id":99}"#.to_vec()) }) + }, + |_connection, _connected, _conn_count, _observed_at_unix_secs| { + Box::pin(async { Ok(()) }) + }, + )); + + let (proxy_tx, mut proxy_rx) = bounded_queue(8); + let (proxy_close_tx, _) = watch::channel(false); + let proxy = Arc::new(ProxyConn::new( + 301, + "authenticated-node".to_string(), + "Authenticated Node".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + )); + hub.register_proxy(proxy); + + let payload = serde_json::to_vec(&serde_json::json!({ + "node_id": "victim-node", + "heartbeat_id": 99u64, + })) + .expect("payload should serialize"); + let mut frame = protocol::encode_frame(1, protocol::HEARTBEAT_DATA, 0, &payload); + hub.handle_proxy_frame(301, &mut frame).await; + + assert_eq!(callback_calls.load(Ordering::Relaxed), 0); + assert!(proxy_rx.try_recv().is_err()); + } + #[tokio::test] async fn second_stream_works_after_first_completes_via_stream_end() { let hub = HubRouter::new(ControlPlaneClient::disabled()); diff --git a/apps/aether-gateway/src/tunnel/embedded/local_relay.rs b/apps/aether-gateway/src/tunnel/embedded/local_relay.rs index b1e9418f2..da6f23332 100644 --- a/apps/aether-gateway/src/tunnel/embedded/local_relay.rs +++ b/apps/aether-gateway/src/tunnel/embedded/local_relay.rs @@ -2,28 +2,23 @@ use std::io; use std::net::SocketAddr; use std::time::Duration; -use aether_contracts::tunnel::{ - resolve_tunnel_request_timeouts, try_decode_tunnel_relay_request_meta, - TUNNEL_RELAY_FORWARDED_BY_HEADER, -}; +use aether_contracts::tunnel::{resolve_tunnel_request_timeouts, TUNNEL_RELAY_FORWARDED_BY_HEADER}; use aether_runtime::{maybe_hold_axum_response_permit, AdmissionPermit}; use async_stream::stream; use axum::body::{Body, Bytes}; use axum::extract::{ConnectInfo, Path, Request, State}; use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode}; use axum::response::IntoResponse; -use bytes::BytesMut; -use futures_util::StreamExt; use tokio::sync::mpsc; use tracing::warn; use crate::api::response::apply_streaming_response_headers; use crate::headers::should_skip_response_header; -use crate::maintenance::record_proxy_upgrade_traffic_success; +use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation; use super::hub::{LocalBodyEvent, LocalStream}; use super::protocol; -use super::AppState; +use super::{AppState, RelayRequestAuthenticated}; pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error"; @@ -86,8 +81,7 @@ pub(crate) async fn open_direct_relay_stream( .await .map_err(map_request_admission_error)?; let stream = state - .hub - .open_local_stream(node_id, &meta) + .open_authorized_local_stream(node_id, &meta) .await .map_err(|error| format!("connect: {error}"))?; if let Err(error) = state @@ -107,7 +101,13 @@ pub(crate) async fn open_direct_relay_stream( return Err(format!("timeout: {error}")); } }; - if let Err(error) = record_proxy_upgrade_traffic_success(state.data.as_ref(), node_id).await { + if let Err(error) = record_proxy_upgrade_traffic_success_for_generation( + state.data.as_ref(), + node_id, + stream.tunnel_generation(), + ) + .await + { warn!( node_id = %node_id, error = %error, @@ -171,7 +171,7 @@ fn is_rollout_probe_request(headers: &HeaderMap, forwarded_by_gateway: bool) -> pub async fn relay_request( Path(node_id): Path, State(state): State, - ConnectInfo(addr): ConnectInfo, + ConnectInfo(_addr): ConnectInfo, request: Request, ) -> impl IntoResponse { let forwarded_by_gateway = request @@ -181,13 +181,10 @@ pub async fn relay_request( .map(str::trim) .is_some_and(|value| !value.is_empty()); let rollout_probe = is_rollout_probe_request(request.headers(), forwarded_by_gateway); - if !addr.ip().is_loopback() && !forwarded_by_gateway { - return tunnel_error_response( - StatusCode::FORBIDDEN, - "forbidden", - "local relay only accepts loopback requests", - ); - } + let already_authenticated = request + .extensions() + .get::() + .is_some(); let request_permit = match state.try_acquire_request_permit().await { Ok(permit) => permit, @@ -226,109 +223,77 @@ pub async fn relay_request( } }; - let mut body_stream = request.into_body().into_data_stream(); - let mut envelope_buf = BytesMut::new(); - let mut meta: Option = None; - let mut stream: Option> = None; + if !already_authenticated { + return release_permit_response( + tunnel_error_response( + StatusCode::FORBIDDEN, + "forbidden", + "relay request integrity must be verified before local dispatch", + ), + request_permit, + ); + } - while let Some(chunk_result) = body_stream.next().await { - let chunk = match chunk_result { - Ok(chunk) => chunk, + let Some(spool) = request + .extensions() + .get::() + .cloned() + else { + return release_permit_response( + tunnel_error_response( + StatusCode::FORBIDDEN, + "forbidden", + "verified relay payload is missing", + ), + request_permit, + ); + }; + let meta = spool.meta().clone(); + + let stream = match state.open_authorized_local_stream(&node_id, &meta).await { + Ok(stream) => stream, + Err(error) => { + return release_permit_response( + tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error), + request_permit, + ); + } + }; + let body_stream = match spool.body_stream().await { + Ok(stream) => stream, + Err(error) => { + state.hub.cancel_local_stream(stream.id, &error); + return release_permit_response( + tunnel_error_response(StatusCode::BAD_GATEWAY, "relay", &error), + request_permit, + ); + } + }; + futures_util::pin_mut!(body_stream); + while let Some(chunk) = futures_util::StreamExt::next(&mut body_stream).await { + let (chunk, end) = match chunk { + Ok(chunk) => (chunk, false), Err(error) => { - if let Some(active_stream) = &stream { - state - .hub - .cancel_local_stream(active_stream.id, "failed to read relay request body"); - } - warn!(error = %error, "failed to read local relay request body"); + let error = error.to_string(); + state.hub.cancel_local_stream(stream.id, &error); return release_permit_response( - tunnel_error_response( - StatusCode::BAD_GATEWAY, - "relay", - "failed to read relay request body", - ), + tunnel_error_response(StatusCode::BAD_GATEWAY, "relay", &error), request_permit, ); } }; - - if stream.is_none() { - envelope_buf.extend_from_slice(&chunk); - let Some((parsed_meta, body_offset)) = - (match try_decode_tunnel_relay_request_meta(&envelope_buf) { - Ok(result) => result, - Err(error) => { - return release_permit_response( - tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error), - request_permit, - ); - } - }) - else { - continue; - }; - - let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta).await { - Ok(stream) => stream, - Err(error) => { - return release_permit_response( - tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error), - request_permit, - ); - } - }; - - if envelope_buf.len() > body_offset { - let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]); - if let Err(error) = state - .hub - .push_local_request_body(opened_stream.id, first_body_chunk, false) - .await - { - state.hub.cancel_local_stream(opened_stream.id, &error); - return release_permit_response( - tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error), - request_permit, - ); - } - } - - envelope_buf.clear(); - meta = Some(parsed_meta); - stream = Some(opened_stream); - continue; - } - - let Some(active_stream) = &stream else { - continue; - }; if let Err(error) = state .hub - .push_local_request_body(active_stream.id, chunk, false) + .push_local_request_body(stream.id, chunk, end) .await { - state.hub.cancel_local_stream(active_stream.id, &error); + state.hub.cancel_local_stream(stream.id, &error); return release_permit_response( tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error), request_permit, ); } } - - let (meta, stream) = match (meta, stream) { - (Some(meta), Some(stream)) => (meta, stream), - _ => { - return release_permit_response( - tunnel_error_response( - StatusCode::BAD_REQUEST, - "bad_request", - "relay envelope metadata truncated", - ), - request_permit, - ); - } - }; - if let Err(error) = state .hub .push_local_request_body(stream.id, Bytes::new(), true) @@ -359,8 +324,12 @@ pub async fn relay_request( } }; if !rollout_probe { - if let Err(error) = - record_proxy_upgrade_traffic_success(state.data.as_ref(), &node_id).await + if let Err(error) = record_proxy_upgrade_traffic_success_for_generation( + state.data.as_ref(), + &node_id, + stream.tunnel_generation(), + ) + .await { warn!( node_id = %node_id, @@ -436,8 +405,16 @@ fn release_permit_response( } fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) { + let connection_declared = aether_http::connection_declared_header_names( + headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); for (name, value) in headers { - if should_skip_local_relay_response_header(name) { + if should_skip_local_relay_response_header(name) + || connection_declared.contains(&name.to_ascii_lowercase()) + { continue; } let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else { @@ -454,12 +431,21 @@ fn should_skip_local_relay_response_header(name: &str) -> bool { should_skip_response_header(name) || name.eq_ignore_ascii_case("content-length") } -fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response { +fn tunnel_error_response(status: StatusCode, kind: &str, _message: &str) -> Response { + let kind = safe_tunnel_error_kind(kind); + let message = match kind { + "overloaded" => "hub relay overloaded", + "forbidden" => "relay request forbidden", + "connect" => "tunnel connection failed", + "timeout" => "tunnel request timed out", + "unavailable" => "tunnel unavailable", + _ => "tunnel relay failed", + }; let mut builder = Response::builder().status(status); if let Some(headers) = builder.headers_mut() { headers.insert( HeaderName::from_static(TUNNEL_ERROR_HEADER), - HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")), + HeaderValue::from_static(kind), ); headers.insert( axum::http::header::CONTENT_TYPE, @@ -471,17 +457,32 @@ fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Respo .unwrap_or_else(|_| Response::new(Body::from("relay error"))) } +fn safe_tunnel_error_kind(kind: &str) -> &'static str { + match kind.trim().to_ascii_lowercase().as_str() { + "overloaded" => "overloaded", + "forbidden" => "forbidden", + "connect" => "connect", + "timeout" => "timeout", + "unavailable" => "unavailable", + "relay" => "relay", + _ => "relay", + } +} + #[cfg(test)] mod tests { use super::super::hub::ProxyConn; - use super::super::{protocol, AppState, ConnConfig, ControlPlaneClient}; + use super::super::{ + protocol, AppState, ConnConfig, ControlPlaneClient, RelayRequestAuthenticated, + }; use super::{ - is_rollout_probe_request, relay_header_timeout, relay_request, Body, HeaderMap, Request, - SocketAddr, StatusCode, TUNNEL_ERROR_HEADER, + is_rollout_probe_request, relay_header_timeout, relay_request, tunnel_error_response, Body, + HeaderMap, Request, SocketAddr, StatusCode, TUNNEL_ERROR_HEADER, }; use crate::data::GatewayDataState; use crate::maintenance::start_proxy_upgrade_rollout; - use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER; + use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::proxy_nodes::{ InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeWriteRepository, StoredProxyNode, @@ -496,6 +497,34 @@ mod tests { use std::time::Duration; use tokio::sync::watch; + const LOCAL_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const LOCAL_TUNNEL_TEST_GENERATION: &str = "local-relay-test-generation-1"; + + #[tokio::test] + async fn relay_error_response_drops_internal_and_peer_details() { + let response = tunnel_error_response( + StatusCode::BAD_GATEWAY, + "https://attacker.invalid/?token=header-secret", + "Bearer body-secret at http://10.0.0.8/private\r\nx-injected: true", + ); + + assert_eq!( + response + .headers() + .get(TUNNEL_ERROR_HEADER) + .and_then(|value| value.to_str().ok()), + Some("relay") + ); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("relay error response body should read"); + assert_eq!(body.as_ref(), b"tunnel relay failed"); + let body = String::from_utf8_lossy(&body); + assert!(!body.contains("body-secret")); + assert!(!body.contains("10.0.0.8")); + assert!(!body.contains("x-injected")); + } + #[test] fn rollout_probe_marker_is_only_trusted_from_a_forwarding_gateway() { let mut headers = HeaderMap::new(); @@ -522,6 +551,18 @@ mod tests { ) } + async fn authenticated_request(envelope: Vec) -> Request { + let spool = crate::tunnel::prepare_owner_relay_request_body(Body::from(envelope)) + .await + .expect("relay envelope should prepare"); + let mut request = Request::builder() + .body(Body::empty()) + .expect("request should build"); + request.extensions_mut().insert(RelayRequestAuthenticated); + request.extensions_mut().insert(spool); + request + } + #[test] fn relay_header_timeout_ignores_request_timeout_for_stream_requests() { let meta = protocol::RequestMeta { @@ -591,7 +632,12 @@ mod tests { None, Some(1_800_000_000), None, - None, + Some(json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": LOCAL_TUNNEL_TEST_PSK, + } + })), None, None, Some(1_800_000_000), @@ -599,6 +645,17 @@ mod tests { Some(1_800_000_000), Some(1_800_000_000), ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + } + + async fn recv_tunnel_test_frame( + proxy_rx: &mut aether_runtime::BoundedQueueReceiver, + description: &str, + ) -> Message { + tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv()) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + .unwrap_or_else(|| panic!("proxy channel closed before {description}")) } fn encode_relay_envelope(meta: &protocol::RequestMeta, body: &[u8]) -> Vec { @@ -611,9 +668,26 @@ mod tests { } #[tokio::test] - async fn relay_rejects_non_loopback_without_forwarded_header() { + async fn relay_rejects_unsigned_request_even_from_loopback() { let request = Request::builder() - .body(Body::empty()) + .body(Body::from(encode_relay_envelope( + &protocol::RequestMeta { + provider_id: None, + endpoint_id: None, + key_id: None, + method: "GET".to_string(), + url: "https://example.com/".to_string(), + headers: HashMap::new(), + stream: false, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 30, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }, + &[], + ))) .expect("request should build"); let response = relay_request( Path("node-123".to_string()), @@ -625,13 +699,40 @@ mod tests { .into_response(); assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response + .headers() + .get(TUNNEL_ERROR_HEADER) + .and_then(|value| value.to_str().ok()), + Some("forbidden") + ); } #[tokio::test] - async fn relay_accepts_forwarded_gateway_request_from_non_loopback() { + async fn relay_rejects_forged_forwarded_gateway_header() { let request = Request::builder() - .header(TUNNEL_RELAY_FORWARDED_BY_HEADER, "gateway-a") - .body(Body::empty()) + .header( + aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER, + "gateway-a", + ) + .body(Body::from(encode_relay_envelope( + &protocol::RequestMeta { + provider_id: None, + endpoint_id: None, + key_id: None, + method: "GET".to_string(), + url: "https://example.com/".to_string(), + headers: HashMap::new(), + stream: false, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 30, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }, + &[], + ))) .expect("request should build"); let response = relay_request( Path("node-123".to_string()), @@ -642,24 +743,29 @@ mod tests { .await .into_response(); - assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(response.status(), StatusCode::FORBIDDEN); assert_eq!( response .headers() .get(TUNNEL_ERROR_HEADER) .and_then(|value| value.to_str().ok()), - Some("bad_request") + Some("forbidden") ); } #[tokio::test] async fn relay_records_real_traffic_confirmation_for_upgrade_rollout() { let mut node = sample_connected_proxy_node("node-123"); - node.proxy_metadata = Some(json!({"version": "1.0.0"})); + node.proxy_metadata + .as_mut() + .and_then(serde_json::Value::as_object_mut) + .expect("proxy metadata should be an object") + .insert("version".to_string(), json!("1.0.0")); let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node])); let data = Arc::new( GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)) - .with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()), + .with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ); let started = start_proxy_upgrade_rollout(data.as_ref(), "2.0.0".to_string(), 1, 0, None) @@ -670,6 +776,7 @@ mod tests { repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-123".to_string(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(1), total_requests_delta: Some(1), @@ -677,7 +784,13 @@ mod tests { failed_requests_delta: Some(0), dns_failures_delta: Some(0), stream_errors_delta: Some(0), - proxy_metadata: Some(json!({"version": "2.0.0"})), + proxy_metadata: Some(json!({ + "version": "2.0.0", + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": LOCAL_TUNNEL_TEST_PSK, + } + })), proxy_version: Some("2.0.0".to_string()), }) .await @@ -692,15 +805,19 @@ mod tests { let state = test_app_state().with_data(Arc::clone(&data)); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - state.hub.register_proxy(Arc::new(ProxyConn::new( - 500, - "node-123".to_string(), - "Node 123".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + state.hub.register_proxy(Arc::new( + ProxyConn::new( + 500, + "node-123".to_string(), + "Node 123".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let meta = protocol::RequestMeta { provider_id: None, @@ -717,9 +834,7 @@ mod tests { http1_only: false, transport_profile: None, }; - let request = Request::builder() - .body(Body::from(encode_relay_envelope(&meta, &[]))) - .expect("request should build"); + let request = authenticated_request(encode_relay_envelope(&meta, &[])).await; let relay_state = state.clone(); let relay_task = tokio::spawn(async move { @@ -733,7 +848,7 @@ mod tests { .into_response() }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -741,7 +856,7 @@ mod tests { .expect("request header frame should parse"); assert_eq!(request_header.msg_type, protocol::REQUEST_HEADERS); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -807,18 +922,29 @@ mod tests { #[tokio::test] async fn relay_strips_hop_by_hop_and_stale_length_headers_from_proxy_response() { - let state = test_app_state(); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + sample_connected_proxy_node("node-123"), + ])); + let data = Arc::new( + GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let state = test_app_state().with_data(data); let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); let (proxy_close_tx, _) = watch::channel(false); - state.hub.register_proxy(Arc::new(ProxyConn::new( - 501, - "node-123".to_string(), - "Node 123".to_string(), - proxy_tx, - proxy_close_tx, - 16, - 2, - ))); + state.hub.register_proxy(Arc::new( + ProxyConn::new( + 501, + "node-123".to_string(), + "Node 123".to_string(), + proxy_tx, + proxy_close_tx, + 16, + 2, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()), + )); let meta = protocol::RequestMeta { provider_id: None, @@ -835,9 +961,7 @@ mod tests { http1_only: false, transport_profile: None, }; - let request = Request::builder() - .body(Body::from(encode_relay_envelope(&meta, &[]))) - .expect("request should build"); + let request = authenticated_request(encode_relay_envelope(&meta, &[])).await; let relay_state = state.clone(); let relay_task = tokio::spawn(async move { @@ -851,7 +975,7 @@ mod tests { .into_response() }); - let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") { + let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -859,7 +983,7 @@ mod tests { .expect("request header frame should parse"); assert_eq!(request_header.msg_type, protocol::REQUEST_HEADERS); - let request_body = match proxy_rx.recv().await.expect("body frame should arrive") { + let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await { Message::Binary(data) => data, other => panic!("unexpected message: {other:?}"), }; @@ -873,6 +997,17 @@ mod tests { ("content-length".to_string(), "999".to_string()), ("transfer-encoding".to_string(), "chunked".to_string()), ("connection".to_string(), "keep-alive".to_string()), + ( + "connection".to_string(), + "x-hop-private, x-accel-redirect".to_string(), + ), + ("x-hop-private".to_string(), "secret".to_string()), + ("x-accel-redirect".to_string(), "/internal".to_string()), + ("set-cookie".to_string(), "session=attacker".to_string()), + ( + "x-aether-future-control".to_string(), + "attacker".to_string(), + ), ("content-type".to_string(), "text/plain".to_string()), ( "x-proxy-timing".to_string(), @@ -905,6 +1040,10 @@ mod tests { assert!(response.headers().get("content-length").is_none()); assert!(response.headers().get("transfer-encoding").is_none()); assert!(response.headers().get("connection").is_none()); + assert!(response.headers().get("x-hop-private").is_none()); + assert!(response.headers().get("x-accel-redirect").is_none()); + assert!(response.headers().get("set-cookie").is_none()); + assert!(response.headers().get("x-aether-future-control").is_none()); assert_eq!( response .headers() diff --git a/apps/aether-gateway/src/tunnel/embedded/mod.rs b/apps/aether-gateway/src/tunnel/embedded/mod.rs index 768c8439c..ca570a1f0 100644 --- a/apps/aether-gateway/src/tunnel/embedded/mod.rs +++ b/apps/aether-gateway/src/tunnel/embedded/mod.rs @@ -4,6 +4,8 @@ mod local_relay; pub mod protocol; mod proxy_conn; +use std::collections::HashSet; +use std::net::{IpAddr, SocketAddr}; use std::sync::Arc; use aether_gateway_tunnel::{ @@ -13,14 +15,18 @@ use aether_runtime::{ hold_admission_permit_until, prometheus_response, service_up_sample, AdmissionPermit, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, MetricKind, MetricLabel, MetricSample, }; -use aether_runtime_state::{RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot}; +use aether_runtime_state::{ + MemoryRuntimeStateConfig, RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot, + RuntimeState, +}; use axum::extract::ws::WebSocketUpgrade; -use axum::extract::State; -use axum::http::HeaderMap; +use axum::extract::{ConnectInfo, State}; +use axum::http::{header, HeaderMap}; use axum::response::{IntoResponse, Json}; use axum::routing::{get, post}; use axum::Router; use dashmap::DashMap; +use sha2::{Digest as _, Sha256}; use tracing::warn; use crate::{data::GatewayDataState, middleware}; @@ -30,6 +36,41 @@ pub use hub::{ConnConfig, HubRouter, LocalBodyEvent, ProxyConn}; pub use local_relay::relay_request; pub(crate) use local_relay::{open_direct_relay_stream, DirectRelayResponse}; +const RELAY_AUTH_CLOCK_SKEW_SECS: u64 = 60; +const MAX_RELAY_AUTH_ID_LEN: usize = 200; +const MAX_RELAY_AUTH_NONCE_LEN: usize = 128; +const TUNNEL_SECURITY_PROOF_CLOCK_SKEW_SECS: u64 = 60; +const MAX_TUNNEL_SECURITY_NODE_ID_LEN: usize = 200; +const MAX_TUNNEL_SECURITY_SESSION_LEN: usize = 128; +const MAX_TUNNEL_SECURITY_PROOF_NONCE_LEN: usize = 128; +const MAX_TUNNEL_SECURITY_PROOF_SIGNATURE_LEN: usize = 128; +const CONTROL_PLANE_AUTH_CLOCK_SKEW_SECS: u64 = 60; +const MAX_CONTROL_PLANE_AUTH_NODE_ID_LEN: usize = 200; +const MAX_CONTROL_PLANE_AUTH_NONCE_LEN: usize = 128; +const MAX_CONTROL_PLANE_AUTH_SIGNATURE_LEN: usize = 128; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum RelayAuthError { + Unavailable, + Invalid, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ControlPlaneAuthError { + Unavailable, + Invalid, +} + +#[derive(Debug, Clone, Copy)] +pub(crate) struct RelayRequestAuthenticated; + +pub(crate) struct PendingRelayAuth { + pub(crate) payload_digest: aether_contracts::tunnel::TunnelRelayPayloadDigest, + nonce: String, + sender: String, + expires_at_unix_secs: u64, +} + #[derive(Clone)] pub struct AppState { pub hub: Arc, @@ -39,6 +80,9 @@ pub struct AppState { request_gate: Option>, distributed_request_gate: Option>, secure_tunnel_keys: Arc>, + relay_instance_id: Arc, + relay_auth_secret: Option>, + relay_auth_runtime_state: Arc, } #[derive(Debug)] @@ -47,6 +91,42 @@ enum RequestAdmissionError { Distributed(RuntimeSemaphoreError), } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProxyTunnelSecurityError { + MissingKey, + InvalidKey, + MissingMode, + UnsupportedMode, + MissingSession, + MissingProof, + MalformedHeader, + InvalidProof, + Replay, + MissingAuthorization, + InvalidAuthorization, + ManagementTokenUnavailable, + Unavailable, +} + +impl ProxyTunnelSecurityError { + fn status_code(self) -> axum::http::StatusCode { + match self { + Self::MissingKey + | Self::MissingMode + | Self::MissingSession + | Self::MissingProof + | Self::InvalidProof + | Self::Replay + | Self::MissingAuthorization + | Self::InvalidAuthorization => axum::http::StatusCode::UNAUTHORIZED, + Self::UnsupportedMode | Self::MalformedHeader => axum::http::StatusCode::BAD_REQUEST, + Self::InvalidKey | Self::ManagementTokenUnavailable | Self::Unavailable => { + axum::http::StatusCode::SERVICE_UNAVAILABLE + } + } + } +} + impl AppState { pub fn new( control_plane: ControlPlaneClient, @@ -61,9 +141,198 @@ impl AppState { request_gate: None, distributed_request_gate: None, secure_tunnel_keys: Arc::new(DashMap::new()), + relay_instance_id: Arc::from("standalone"), + relay_auth_secret: None, + relay_auth_runtime_state: Arc::new(RuntimeState::memory( + MemoryRuntimeStateConfig::default(), + )), } } + pub fn with_relay_auth( + mut self, + instance_id: impl Into, + secret: Option>>, + runtime_state: Arc, + ) -> Self { + self.relay_instance_id = Arc::from(instance_id.into()); + self.relay_auth_secret = secret + .map(Into::into) + .filter(|value: &Vec| value.len() >= 32) + .map(Arc::from); + self.relay_auth_runtime_state = runtime_state; + self + } + + async fn claim_relay_auth_nonce( + &self, + nonce: &str, + sender: &str, + now_unix_secs: u64, + expires_at_unix_secs: u64, + ) -> Result { + let ttl = std::time::Duration::from_secs( + expires_at_unix_secs + .saturating_sub(now_unix_secs) + .saturating_add(1) + .max(1), + ); + self.relay_auth_runtime_state + .kv_set_if_absent( + &format!("tunnel:relay:auth:nonce:{nonce}"), + sender.to_string(), + ttl, + ) + .await + .map_err(|_| RelayAuthError::Unavailable) + } + + pub(crate) async fn authenticate_relay_request_headers( + &self, + headers: &HeaderMap, + node_id: &str, + require_local_owner: bool, + ) -> Result { + let secret = self + .relay_auth_secret + .as_deref() + .ok_or(RelayAuthError::Unavailable)?; + let sender = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_AUTH_SENDER_HEADER, + MAX_RELAY_AUTH_ID_LEN, + )?; + let owner = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + MAX_RELAY_AUTH_ID_LEN, + )?; + if require_local_owner && owner != self.relay_instance_id.as_ref() { + return Err(RelayAuthError::Invalid); + } + let nonce = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_AUTH_NONCE_HEADER, + MAX_RELAY_AUTH_NONCE_LEN, + )?; + if !nonce + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) + { + return Err(RelayAuthError::Invalid); + } + let signature = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + 128, + )?; + let payload_digest = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + 96, + )?; + let payload_digest = + aether_contracts::tunnel::TunnelRelayPayloadDigest::decode_header_value(payload_digest) + .ok_or(RelayAuthError::Invalid)?; + let timestamp = relay_auth_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + 20, + )? + .parse::() + .map_err(|_| RelayAuthError::Invalid)?; + let now = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .map_err(|_| RelayAuthError::Unavailable)? + .as_secs(); + if now.abs_diff(timestamp) > RELAY_AUTH_CLOCK_SKEW_SECS { + return Err(RelayAuthError::Invalid); + } + + let forwarded_by = relay_auth_optional_header( + headers, + aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER, + MAX_RELAY_AUTH_ID_LEN, + )? + .unwrap_or_default(); + if !forwarded_by.is_empty() && forwarded_by != sender { + return Err(RelayAuthError::Invalid); + } + let rollout_probe = match relay_auth_optional_header( + headers, + crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_HEADER, + 1, + )? { + Some(crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_VALUE) => true, + Some(_) => return Err(RelayAuthError::Invalid), + None => false, + }; + if rollout_probe && forwarded_by.is_empty() { + return Err(RelayAuthError::Invalid); + } + if !aether_contracts::tunnel::verify_tunnel_relay_request_signature( + secret, + sender, + owner, + node_id, + forwarded_by, + rollout_probe, + timestamp, + nonce, + &payload_digest, + signature, + ) { + return Err(RelayAuthError::Invalid); + } + Ok(PendingRelayAuth { + payload_digest, + nonce: nonce.to_string(), + sender: sender.to_string(), + expires_at_unix_secs: timestamp.saturating_add(RELAY_AUTH_CLOCK_SKEW_SECS), + }) + } + + pub(crate) async fn commit_relay_auth( + &self, + pending: &PendingRelayAuth, + ) -> Result<(), RelayAuthError> { + let now = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .map_err(|_| RelayAuthError::Unavailable)? + .as_secs(); + if now > pending.expires_at_unix_secs { + return Err(RelayAuthError::Invalid); + } + if !self + .claim_relay_auth_nonce( + &pending.nonce, + &pending.sender, + now, + pending.expires_at_unix_secs, + ) + .await? + { + return Err(RelayAuthError::Invalid); + } + Ok(()) + } + + pub(crate) async fn authenticate_relay_request( + &self, + headers: &HeaderMap, + node_id: &str, + payload_digest: &aether_contracts::tunnel::TunnelRelayPayloadDigest, + require_local_owner: bool, + ) -> Result<(), RelayAuthError> { + let pending = self + .authenticate_relay_request_headers(headers, node_id, require_local_owner) + .await?; + if pending.payload_digest != *payload_digest { + return Err(RelayAuthError::Invalid); + } + self.commit_relay_auth(&pending).await + } + pub(crate) fn register_secure_tunnel_key( &self, node_id: impl Into, @@ -78,28 +347,394 @@ impl AppState { .map(|entry| entry.value().clone()) } - async fn secure_tunnel_key_for_node(&self, node_id: &str) -> Option { - if let Some(key) = self.secure_tunnel_key(node_id) { - return Some(key); + async fn secure_tunnel_binding_for_node( + &self, + node_id: &str, + ) -> Result, aether_data::DataLayerError> { + if !self.data.has_proxy_node_reader() { + return Ok(None); } - let key = self + let binding = + crate::state::decrypt_or_migrate_proxy_tunnel_psk_binding(self.data.as_ref(), node_id) + .await?; + if let Some(binding) = binding.as_ref() { + self.register_secure_tunnel_key(node_id.to_string(), binding.key.clone()); + } else { + self.secure_tunnel_keys.remove(node_id); + } + Ok(binding.map(|binding| (binding.key, binding.tunnel_generation))) + } + + async fn secure_tunnel_binding_for_handshake( + &self, + node_id: &str, + _requested_generation: &str, + ) -> Result, aether_data::DataLayerError> { + if !self.data.has_proxy_node_reader() { + return Err(aether_data::DataLayerError::InvalidConfiguration( + "proxy node reader is unavailable for tunnel authentication".to_string(), + )); + } + self.secure_tunnel_binding_for_node(node_id).await + } + + async fn authenticate_proxy_tunnel_management_token( + &self, + headers: &HeaderMap, + node_id: &str, + remote_ip: IpAddr, + ) -> Result<(hub::ProxyManagementTokenCredential, String), ProxyTunnelSecurityError> { + let authenticated = crate::management_token_auth::authenticate_management_token( + self.data.as_ref(), + headers, + remote_ip, + ) + .await + .map_err(|error| match error { + crate::management_token_auth::ManagementTokenAuthError::Missing => { + ProxyTunnelSecurityError::MissingAuthorization + } + crate::management_token_auth::ManagementTokenAuthError::Invalid => { + ProxyTunnelSecurityError::InvalidAuthorization + } + crate::management_token_auth::ManagementTokenAuthError::Unavailable => { + ProxyTunnelSecurityError::ManagementTokenUnavailable + } + })?; + if !crate::roles::can_write_admin_console(&authenticated.user.role) { + return Err(ProxyTunnelSecurityError::InvalidAuthorization); + } + + let Some(node) = self .data .find_proxy_node(node_id) .await - .ok() - .flatten() - .and_then(|node| { - node.proxy_metadata.and_then(|metadata| { - metadata - .pointer("/tunnel_security/encryption_key") - .and_then(|value| value.as_str()) - .map(str::to_string) - }) - }); - if let Some(key) = key.as_ref() { - self.register_secure_tunnel_key(node_id.to_string(), key.clone()); + .map_err(|_| ProxyTunnelSecurityError::ManagementTokenUnavailable)? + else { + return Err(ProxyTunnelSecurityError::InvalidAuthorization); + }; + if !node.tunnel_mode { + return Err(ProxyTunnelSecurityError::InvalidAuthorization); } - key + + if !management_token_may_connect_proxy_tunnel(&authenticated.permissions) { + return Err(ProxyTunnelSecurityError::InvalidAuthorization); + } + + if let Err(error) = self + .data + .record_management_token_usage(&authenticated.token.id, Some(&remote_ip.to_string())) + .await + { + warn!( + token_id = %authenticated.token.id, + error = %error, + "failed to record proxy tunnel management token usage" + ); + } + Ok(( + hub::ProxyManagementTokenCredential { + verified_token_hash: authenticated.verified_token_hash, + token_id: authenticated.token.id, + user_id: authenticated.user.id, + remote_ip, + }, + node.tunnel_generation, + )) + } + + async fn authorized_proxy_connections_for_new_stream( + &self, + node_id: &str, + ) -> Result>, String> { + if !self.data.has_proxy_node_reader() { + return Err(control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE.to_string()); + } + + let connections = self + .hub + .proxy_connections_for_node(node_id) + .into_iter() + .filter(|connection| connection.is_available()) + .collect::>(); + if connections.is_empty() { + return Ok(Some(HashSet::new())); + } + + let mut authorized = HashSet::new(); + let mut validation_unavailable = false; + for connection in connections { + match validate_proxy_connection_credential(self.data.as_ref(), &connection).await { + Ok(()) => { + authorized.insert(connection.id); + } + Err(error) if error == control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE => { + validation_unavailable = true; + } + Err(_) => { + warn!( + node_id = %node_id, + conn_id = connection.id, + "closing proxy connection after credential revocation" + ); + self.hub.request_close_proxy(connection.id); + } + } + } + + // Keep the node existence/mode read after credential validation. If a node is + // deleted while a token lookup is in flight, that lookup must not authorize a + // stream against the deleted node. + let node = self + .data + .find_proxy_node(node_id) + .await + .map_err(|_| "proxy tunnel credential validation unavailable".to_string())?; + if !node.is_some_and(|node| node.id == node_id && node.tunnel_mode) { + self.hub.request_close_proxies_for_node(node_id); + return Err("proxy tunnel credential was revoked".to_string()); + } + + if !authorized.is_empty() { + return Ok(Some(authorized)); + } + if validation_unavailable { + return Err("proxy tunnel credential validation unavailable".to_string()); + } + Err("proxy tunnel credential was revoked".to_string()) + } + + pub(crate) async fn open_authorized_local_stream( + &self, + node_id: &str, + meta: &protocol::RequestMeta, + ) -> Result, String> { + let authorized = self + .authorized_proxy_connections_for_new_stream(node_id) + .await?; + self.hub + .open_local_stream_with_authorized_connections(node_id, meta, authorized.as_ref()) + .await + } + + async fn authenticate_proxy_tunnel_security( + &self, + headers: &HeaderMap, + node_id: &str, + tunnel_generation: &str, + protocol_version: u8, + stored_security_key: Option, + ) -> Result<(String, String), ProxyTunnelSecurityError> { + let tunnel_security = required_proxy_tunnel_header( + headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER, + 32, + ProxyTunnelSecurityError::MissingMode, + )?; + let security_session = required_proxy_tunnel_header( + headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_SESSION_HEADER, + MAX_TUNNEL_SECURITY_SESSION_LEN, + ProxyTunnelSecurityError::MissingSession, + )?; + if !valid_proxy_tunnel_token(security_session, 16) { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + let (security_key, security_session) = resolve_proxy_tunnel_security( + stored_security_key, + Some(tunnel_security), + Some(security_session.to_string()), + )?; + + let timestamp = required_proxy_tunnel_header( + headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, + 20, + ProxyTunnelSecurityError::MissingProof, + )?; + if !timestamp.bytes().all(|byte| byte.is_ascii_digit()) { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + let timestamp = timestamp + .parse::() + .map_err(|_| ProxyTunnelSecurityError::MalformedHeader)?; + let nonce = required_proxy_tunnel_header( + headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_NONCE_HEADER, + MAX_TUNNEL_SECURITY_PROOF_NONCE_LEN, + ProxyTunnelSecurityError::MissingProof, + )?; + if !valid_proxy_tunnel_token(nonce, 16) { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + let signature = required_proxy_tunnel_header( + headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + MAX_TUNNEL_SECURITY_PROOF_SIGNATURE_LEN, + ProxyTunnelSecurityError::MissingProof, + )?; + if signature.len() != 43 + || !signature + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) + { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + + let now = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .map_err(|_| ProxyTunnelSecurityError::Unavailable)? + .as_secs(); + if now.abs_diff(timestamp) > TUNNEL_SECURITY_PROOF_CLOCK_SKEW_SECS { + return Err(ProxyTunnelSecurityError::InvalidProof); + } + if !aether_contracts::tunnel_security::verify_tunnel_security_handshake_for_generation( + &security_key, + node_id, + tunnel_generation, + tunnel_security, + &security_session, + protocol_version, + timestamp, + nonce, + signature, + ) { + return Err(ProxyTunnelSecurityError::InvalidProof); + } + + let ttl = std::time::Duration::from_secs( + timestamp + .saturating_add(TUNNEL_SECURITY_PROOF_CLOCK_SKEW_SECS) + .saturating_sub(now) + .saturating_add(1) + .max(1), + ); + let claimed = self + .relay_auth_runtime_state + .kv_set_if_absent( + &format!("tunnel:security:proof:nonce:{node_id}:{tunnel_generation}:{nonce}"), + security_session.clone(), + ttl, + ) + .await + .map_err(|_| ProxyTunnelSecurityError::Unavailable)?; + if !claimed { + return Err(ProxyTunnelSecurityError::Replay); + } + Ok((security_key, security_session)) + } + + pub(crate) async fn authenticate_control_plane_request( + &self, + headers: &HeaderMap, + method: &str, + path: &str, + payload_node_id: &str, + body: &[u8], + ) -> Result { + if !self.data.has_proxy_node_reader() { + return Err(ControlPlaneAuthError::Unavailable); + } + let authenticated_node_id = control_plane_auth_header( + headers, + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, + MAX_CONTROL_PLANE_AUTH_NODE_ID_LEN, + )?; + if authenticated_node_id != payload_node_id { + return Err(ControlPlaneAuthError::Invalid); + } + let authenticated_generation = control_plane_auth_header( + headers, + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_GENERATION_HEADER, + 128, + )?; + if !valid_proxy_tunnel_token(authenticated_generation, 1) { + return Err(ControlPlaneAuthError::Invalid); + } + let timestamp = control_plane_auth_header( + headers, + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, + 20, + )? + .parse::() + .map_err(|_| ControlPlaneAuthError::Invalid)?; + let nonce = control_plane_auth_header( + headers, + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_NONCE_HEADER, + MAX_CONTROL_PLANE_AUTH_NONCE_LEN, + )?; + if nonce.len() < 16 + || !nonce + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) + { + return Err(ControlPlaneAuthError::Invalid); + } + let signature = control_plane_auth_header( + headers, + aether_contracts::tunnel_security::TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, + MAX_CONTROL_PLANE_AUTH_SIGNATURE_LEN, + )?; + if signature.len() != 43 + || !signature + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) + { + return Err(ControlPlaneAuthError::Invalid); + } + + let now = std::time::SystemTime::now() + .duration_since(std::time::SystemTime::UNIX_EPOCH) + .map_err(|_| ControlPlaneAuthError::Unavailable)? + .as_secs(); + if now.abs_diff(timestamp) > CONTROL_PLANE_AUTH_CLOCK_SKEW_SECS { + return Err(ControlPlaneAuthError::Invalid); + } + let (security_key, current_generation) = self + .secure_tunnel_binding_for_node(payload_node_id) + .await + .map_err(|_| ControlPlaneAuthError::Unavailable)? + .ok_or(ControlPlaneAuthError::Invalid)?; + if current_generation != authenticated_generation { + return Err(ControlPlaneAuthError::Invalid); + } + if !aether_contracts::tunnel_security::verify_tunnel_control_plane_request_for_generation( + &security_key, + method, + path, + authenticated_node_id, + authenticated_generation, + timestamp, + nonce, + body, + signature, + ) { + return Err(ControlPlaneAuthError::Invalid); + } + + let ttl = std::time::Duration::from_secs( + timestamp + .saturating_add(CONTROL_PLANE_AUTH_CLOCK_SKEW_SECS) + .saturating_sub(now) + .saturating_add(1) + .max(1), + ); + let node_digest = Sha256::digest( + format!("{authenticated_node_id}\0{authenticated_generation}").as_bytes(), + ); + let claimed = self + .relay_auth_runtime_state + .kv_set_if_absent( + &format!("tunnel:control-plane:auth:nonce:{node_digest:x}:{nonce}"), + timestamp.to_string(), + ttl, + ) + .await + .map_err(|_| ControlPlaneAuthError::Unavailable)?; + if !claimed { + return Err(ControlPlaneAuthError::Invalid); + } + Ok(current_generation) } pub(crate) fn with_data(mut self, data: Arc) -> Self { @@ -181,6 +816,1279 @@ impl AppState { } } +fn management_token_may_connect_proxy_tunnel(permissions: &[String]) -> bool { + permissions + .iter() + .any(|permission| permission == "admin:proxy_nodes:admin") +} + +fn relay_auth_header<'a>( + headers: &'a HeaderMap, + name: &str, + max_len: usize, +) -> Result<&'a str, RelayAuthError> { + let mut values = headers.get_all(name).iter(); + let value = values.next().ok_or(RelayAuthError::Invalid)?; + if values.next().is_some() { + return Err(RelayAuthError::Invalid); + } + let value = value.to_str().map_err(|_| RelayAuthError::Invalid)?; + let trimmed = value.trim(); + if trimmed.is_empty() || trimmed.len() > max_len || trimmed != value { + return Err(RelayAuthError::Invalid); + } + Ok(trimmed) +} + +fn control_plane_auth_header<'a>( + headers: &'a HeaderMap, + name: &str, + max_len: usize, +) -> Result<&'a str, ControlPlaneAuthError> { + let mut values = headers.get_all(name).iter(); + let value = values.next().ok_or(ControlPlaneAuthError::Invalid)?; + if values.next().is_some() { + return Err(ControlPlaneAuthError::Invalid); + } + let value = value.to_str().map_err(|_| ControlPlaneAuthError::Invalid)?; + let trimmed = value.trim(); + if trimmed.is_empty() || trimmed.len() > max_len || trimmed != value { + return Err(ControlPlaneAuthError::Invalid); + } + Ok(trimmed) +} + +fn relay_auth_optional_header<'a>( + headers: &'a HeaderMap, + name: &str, + max_len: usize, +) -> Result, RelayAuthError> { + let mut values = headers.get_all(name).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err(RelayAuthError::Invalid); + } + let value = value.to_str().map_err(|_| RelayAuthError::Invalid)?; + let trimmed = value.trim(); + if trimmed.is_empty() || trimmed.len() > max_len || trimmed != value { + return Err(RelayAuthError::Invalid); + } + Ok(Some(trimmed)) +} + +fn required_proxy_tunnel_header<'a>( + headers: &'a HeaderMap, + name: &str, + max_len: usize, + missing_error: ProxyTunnelSecurityError, +) -> Result<&'a str, ProxyTunnelSecurityError> { + let mut values = headers.get_all(name).iter(); + let value = values.next().ok_or(missing_error)?; + if values.next().is_some() { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + let value = value + .to_str() + .map_err(|_| ProxyTunnelSecurityError::MalformedHeader)?; + let trimmed = value.trim(); + if trimmed.is_empty() || trimmed.len() > max_len || trimmed != value { + return Err(ProxyTunnelSecurityError::MalformedHeader); + } + Ok(trimmed) +} + +fn valid_proxy_tunnel_token(value: &str, min_len: usize) -> bool { + value.len() >= min_len + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) +} + +fn constant_time_secret_eq(left: &str, right: &str) -> bool { + if left.len() != right.len() { + return false; + } + left.as_bytes() + .iter() + .zip(right.as_bytes()) + .fold(0u8, |difference, (left, right)| difference | (left ^ right)) + == 0 +} + +pub(crate) async fn validate_proxy_connection_credential( + data: &GatewayDataState, + connection: &hub::ProxyConn, +) -> Result<(), &'static str> { + if !data.has_proxy_node_reader() { + return Err(control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE); + } + let node = data + .find_proxy_node(&connection.node_id) + .await + .map_err(|_| control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE)?; + let Some(node) = node.filter(|node| node.id == connection.node_id && node.tunnel_mode) else { + return Err(control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED); + }; + if connection.node_generation.is_empty() || connection.node_generation != node.tunnel_generation + { + return Err(control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED); + } + + match connection.credential_binding() { + Some(hub::ProxyCredentialBinding::Psk(authenticated_key)) => { + let current_key = crate::state::decrypt_or_migrate_proxy_tunnel_psk(data, &node.id) + .await + .map_err(|_| control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE)?; + if current_key + .as_deref() + .is_some_and(|current| constant_time_secret_eq(current, &authenticated_key)) + { + Ok(()) + } else { + Err(control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED) + } + } + Some(hub::ProxyCredentialBinding::ManagementToken(credential)) => { + let authenticated = crate::management_token_auth::authenticate_management_token_hash( + data, + &credential.verified_token_hash, + credential.remote_ip, + ) + .await + .map_err(|error| match error { + crate::management_token_auth::ManagementTokenAuthError::Unavailable => { + control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE + } + crate::management_token_auth::ManagementTokenAuthError::Invalid + | crate::management_token_auth::ManagementTokenAuthError::Missing => { + control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED + } + })?; + if authenticated.token.id == credential.token_id + && authenticated.token.user_id == credential.user_id + && authenticated.user.id == credential.user_id + && crate::roles::can_write_admin_console(&authenticated.user.role) + && management_token_may_connect_proxy_tunnel(&authenticated.permissions) + { + Ok(()) + } else { + Err(control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED) + } + } + None => Err(control_plane::CONTROL_PLANE_CREDENTIAL_REVOKED), + } +} + +fn resolve_proxy_tunnel_security( + stored_security_key: Option, + tunnel_security: Option<&str>, + security_session: Option, +) -> Result<(String, String), ProxyTunnelSecurityError> { + let key = stored_security_key.ok_or(ProxyTunnelSecurityError::MissingKey)?; + aether_contracts::tunnel_security::decode_psk(&key) + .map_err(|_| ProxyTunnelSecurityError::InvalidKey)?; + + match tunnel_security { + Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED) => {} + Some(_) => return Err(ProxyTunnelSecurityError::UnsupportedMode), + None => return Err(ProxyTunnelSecurityError::MissingMode), + } + + let session = security_session.ok_or(ProxyTunnelSecurityError::MissingSession)?; + Ok((key, session)) +} + +#[cfg(test)] +mod relay_auth_tests { + use super::{AppState, ConnConfig, ControlPlaneClient, RelayAuthError}; + 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 axum::http::{HeaderMap, HeaderValue}; + use std::time::{Duration, SystemTime, UNIX_EPOCH}; + + const SECRET: &[u8] = b"relay-auth-test-secret-at-least-32-bytes"; + + fn state() -> AppState { + AppState::new( + ControlPlaneClient::disabled(), + ConnConfig { + ping_interval: Duration::from_secs(15), + idle_timeout: Duration::ZERO, + outbound_queue_capacity: 8, + }, + 8, + ) + .with_relay_auth( + "gateway-b", + Some(SECRET.to_vec()), + std::sync::Arc::new(aether_runtime_state::RuntimeState::memory( + aether_runtime_state::MemoryRuntimeStateConfig::default(), + )), + ) + } + + fn signed_headers( + owner: &str, + node_id: &str, + metadata: &[u8], + body: &[u8], + timestamp: u64, + nonce: &str, + ) -> HeaderMap { + let sender = "gateway-a"; + let payload_digest = tunnel_relay_payload_digest(metadata, body); + let signature = sign_tunnel_relay_request( + SECRET, + sender, + owner, + node_id, + "", + false, + timestamp, + nonce, + &payload_digest, + ); + let mut headers = HeaderMap::new(); + headers.insert( + TUNNEL_RELAY_AUTH_SENDER_HEADER, + HeaderValue::from_static(sender), + ); + headers.insert( + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + HeaderValue::from_str(owner).expect("owner header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + HeaderValue::from_str(×tamp.to_string()).expect("timestamp header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_NONCE_HEADER, + HeaderValue::from_str(nonce).expect("nonce header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + HeaderValue::from_str(&payload_digest.encode_header_value()).expect("payload header"), + ); + headers.insert( + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + HeaderValue::from_str(&signature).expect("signature header"), + ); + headers + } + + fn now() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs() + } + + #[tokio::test] + async fn accepts_valid_relay_authentication_once() { + let state = state(); + let metadata = b"metadata-envelope"; + let body = b"request-body"; + let headers = signed_headers("gateway-b", "node-1", metadata, body, now(), "nonce-valid"); + + assert_eq!( + state + .authenticate_relay_request( + &headers, + "node-1", + &tunnel_relay_payload_digest(metadata, body), + true, + ) + .await, + Ok(()) + ); + } + + #[tokio::test] + async fn rejects_metadata_tampering_and_owner_mismatch() { + let state = state(); + let headers = signed_headers( + "gateway-b", + "node-1", + b"metadata", + b"body", + now(), + "nonce-tamper", + ); + assert_eq!( + state + .authenticate_relay_request( + &headers, + "node-1", + &tunnel_relay_payload_digest(b"tampered", b"body"), + true, + ) + .await, + Err(RelayAuthError::Invalid) + ); + + let headers = signed_headers( + "gateway-other", + "node-1", + b"metadata", + b"body", + now(), + "nonce-owner", + ); + assert_eq!( + state + .authenticate_relay_request( + &headers, + "node-1", + &tunnel_relay_payload_digest(b"metadata", b"body"), + true, + ) + .await, + Err(RelayAuthError::Invalid) + ); + } + + #[tokio::test] + async fn payload_tampering_does_not_consume_a_valid_nonce() { + let state = state(); + let metadata = b"metadata"; + let body = b"original-body"; + let headers = signed_headers( + "gateway-b", + "node-1", + metadata, + body, + now(), + "nonce-body-preserve", + ); + + assert_eq!( + state + .authenticate_relay_request( + &headers, + "node-1", + &tunnel_relay_payload_digest(metadata, b"tampered-body"), + true, + ) + .await, + Err(RelayAuthError::Invalid) + ); + assert_eq!( + state + .authenticate_relay_request( + &headers, + "node-1", + &tunnel_relay_payload_digest(metadata, body), + true, + ) + .await, + Ok(()) + ); + } + + #[tokio::test] + async fn rejects_replay_and_expired_timestamp() { + let state = state(); + let payload_digest = tunnel_relay_payload_digest(b"metadata", b"body"); + let headers = signed_headers( + "gateway-b", + "node-1", + b"metadata", + b"body", + now(), + "nonce-replay", + ); + assert_eq!( + state + .authenticate_relay_request(&headers, "node-1", &payload_digest, true) + .await, + Ok(()) + ); + assert_eq!( + state + .authenticate_relay_request(&headers, "node-1", &payload_digest, true) + .await, + Err(RelayAuthError::Invalid) + ); + + let expired = now().saturating_sub(super::RELAY_AUTH_CLOCK_SKEW_SECS + 1); + let headers = signed_headers( + "gateway-b", + "node-1", + b"metadata", + b"body", + expired, + "nonce-expired", + ); + assert_eq!( + state + .authenticate_relay_request(&headers, "node-1", &payload_digest, true) + .await, + Err(RelayAuthError::Invalid) + ); + } + + #[tokio::test] + async fn rejects_pending_relay_auth_that_expires_before_commit() { + let state = state(); + let metadata = b"metadata"; + let body = b"body"; + let nonce = "nonce-expired-before-commit"; + let headers = signed_headers("gateway-b", "node-1", metadata, body, now(), nonce); + let mut pending = state + .authenticate_relay_request_headers(&headers, "node-1", true) + .await + .expect("fresh signature headers should authenticate"); + pending.expires_at_unix_secs = now().saturating_sub(1); + + assert_eq!( + state.commit_relay_auth(&pending).await, + Err(RelayAuthError::Invalid) + ); + + let fresh_headers = signed_headers("gateway-b", "node-1", metadata, body, now(), nonce); + assert_eq!( + state + .authenticate_relay_request( + &fresh_headers, + "node-1", + &tunnel_relay_payload_digest(metadata, body), + true, + ) + .await, + Ok(()), + "an expired pending request must not consume the nonce" + ); + } +} + +#[cfg(test)] +mod proxy_tunnel_security_tests { + use super::{ + hub::ProxyManagementTokenCredential, management_token_may_connect_proxy_tunnel, + resolve_proxy_tunnel_security, AppState, ConnConfig, ControlPlaneAuthError, + ControlPlaneClient, ProxyConn, ProxyTunnelSecurityError, + TUNNEL_SECURITY_PROOF_CLOCK_SKEW_SECS, + }; + use aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION; + use aether_contracts::tunnel_security::{ + sign_tunnel_control_plane_request_for_generation, + sign_tunnel_security_handshake_for_generation, TUNNEL_CONTROL_PLANE_GENERATION_HEADER, + TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, TUNNEL_CONTROL_PLANE_NONCE_HEADER, + TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, + TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED, + TUNNEL_SECURITY_PROOF_NONCE_HEADER, TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, TUNNEL_SECURITY_SESSION_HEADER, + }; + use aether_data::repository::management_tokens::{ + CreateManagementTokenRecord, InMemoryManagementTokenRepository, + ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, + StoredManagementTokenUserSummary, StoredManagementTokenWithUser, + }; + use aether_data::repository::proxy_nodes::{ + InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeReadRepository, + ProxyNodeRegistrationMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, + StoredProxyNode, + }; + use aether_data::repository::users::{InMemoryUserReadRepository, StoredUserAuthRecord}; + use aether_runtime::{bounded_queue, BoundedQueueReceiver}; + use axum::extract::ws::Message; + use axum::http::{HeaderMap, HeaderValue}; + use serde_json::json; + use sha2::{Digest, Sha256}; + use std::net::{IpAddr, Ipv4Addr}; + use std::sync::Arc; + use std::time::{Duration, SystemTime, UNIX_EPOCH}; + use tokio::sync::watch; + + const VALID_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const ROTATED_PSK: &str = "CAgICAgICAgICAgICAgICAgICAgICAgICAgICAgICAg="; + const NODE_ID: &str = "node-proof-1"; + const SESSION: &str = "0123456789abcdef0123456789abcdef"; + const REVOCATION_NODE_ID: &str = "node-revocation-1"; + const REVOCATION_TOKEN_ID: &str = "token-revocation-1"; + const REVOCATION_USER_ID: &str = "user-revocation-1"; + const RAW_MANAGEMENT_TOKEN: &str = "ae-tunnel-revocation-original"; + const ROTATED_MANAGEMENT_TOKEN: &str = "ae-tunnel-revocation-rotated"; + const TEST_TUNNEL_GENERATION: &str = "test-generation-1"; + + fn state() -> AppState { + AppState::new( + ControlPlaneClient::disabled(), + ConnConfig { + ping_interval: Duration::from_secs(15), + idle_timeout: Duration::ZERO, + outbound_queue_capacity: 8, + }, + 8, + ) + } + + struct RevocationFixture { + state: AppState, + data: Arc, + token_repository: Arc, + node_repository: Arc, + } + + struct RegisteredTestProxy { + connection: Arc, + _outbound_rx: BoundedQueueReceiver, + close_rx: watch::Receiver, + } + + fn management_token_hash(raw_token: &str) -> String { + format!("{:x}", Sha256::digest(raw_token.as_bytes())) + } + + fn management_token_user_summary() -> StoredManagementTokenUserSummary { + StoredManagementTokenUserSummary::new( + REVOCATION_USER_ID.to_string(), + Some("tunnel-revocation@example.com".to_string()), + "tunnel_revocation_admin".to_string(), + "admin".to_string(), + ) + .expect("management token user summary should build") + } + + fn management_token_with_user() -> StoredManagementTokenWithUser { + let token = StoredManagementToken::new( + REVOCATION_TOKEN_ID.to_string(), + REVOCATION_USER_ID.to_string(), + "tunnel revocation token".to_string(), + ) + .expect("management token should build") + .with_permissions(Some(json!(["admin:proxy_nodes:admin"]))) + .with_runtime_fields(None, None, None, 0, true); + StoredManagementTokenWithUser::new(token, management_token_user_summary()) + } + + fn management_token_create_record(raw_token: &str) -> CreateManagementTokenRecord { + CreateManagementTokenRecord { + id: REVOCATION_TOKEN_ID.to_string(), + user_id: REVOCATION_USER_ID.to_string(), + user: management_token_user_summary(), + token_hash: management_token_hash(raw_token), + token_prefix: Some("ae-tunnel".to_string()), + name: "tunnel revocation token".to_string(), + description: None, + allowed_ips: None, + permissions: Some(json!(["admin:proxy_nodes:admin"])), + expires_at_unix_secs: None, + is_active: true, + } + } + + fn current_management_token_user() -> StoredUserAuthRecord { + StoredUserAuthRecord::new( + REVOCATION_USER_ID.to_string(), + Some("tunnel-revocation@example.com".to_string()), + true, + "tunnel_revocation_admin".to_string(), + None, + "admin".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + None, + None, + ) + .expect("current management token user should build") + } + + fn tunnel_node(psk: Option<&str>) -> StoredProxyNode { + StoredProxyNode::new( + REVOCATION_NODE_ID.to_string(), + "revocation test node".to_string(), + "127.0.0.1".to_string(), + 0, + false, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + true, + true, + 1, + ) + .expect("tunnel node should build") + .with_runtime_fields( + None, + None, + None, + None, + psk.map(|psk| { + json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": psk, + } + }) + }), + None, + None, + None, + None, + None, + None, + ) + .with_tunnel_generation(TEST_TUNNEL_GENERATION.to_string()) + } + + fn revocation_fixture(psk: Option<&str>) -> RevocationFixture { + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![management_token_with_user()], + vec![( + management_token_hash(RAW_MANAGEMENT_TOKEN), + REVOCATION_TOKEN_ID.to_string(), + )], + )); + let node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![tunnel_node(psk)])); + let data = Arc::new( + crate::data::GatewayDataState::with_management_token_repository_for_tests(Arc::clone( + &token_repository, + )) + .attach_proxy_node_repository_for_tests(Arc::clone(&node_repository)) + .with_user_reader(Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ + current_management_token_user(), + ]))) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY), + ); + let state = state().with_data(Arc::clone(&data)); + RevocationFixture { + state, + data, + token_repository, + node_repository, + } + } + + fn register_psk_connection(state: &AppState, psk: &str) -> RegisteredTestProxy { + let (outbound_tx, outbound_rx) = bounded_queue(8); + let (close_tx, close_rx) = watch::channel(false); + let connection = Arc::new( + ProxyConn::new( + state.hub.alloc_conn_id(), + REVOCATION_NODE_ID.to_string(), + "revocation test node".to_string(), + outbound_tx, + close_tx, + 8, + CURRENT_TUNNEL_PROTOCOL_VERSION, + ) + .with_tunnel_generation(TEST_TUNNEL_GENERATION.to_string()) + .with_authenticated_key(psk.to_string()), + ); + state.hub.register_proxy(Arc::clone(&connection)); + RegisteredTestProxy { + connection, + _outbound_rx: outbound_rx, + close_rx, + } + } + + fn register_management_token_connection( + state: &AppState, + raw_token: &str, + ) -> RegisteredTestProxy { + let (outbound_tx, outbound_rx) = bounded_queue(8); + let (close_tx, close_rx) = watch::channel(false); + let connection = Arc::new( + ProxyConn::new( + state.hub.alloc_conn_id(), + REVOCATION_NODE_ID.to_string(), + "revocation test node".to_string(), + outbound_tx, + close_tx, + 8, + CURRENT_TUNNEL_PROTOCOL_VERSION, + ) + .with_tunnel_generation(TEST_TUNNEL_GENERATION.to_string()) + .with_management_token_credential(ProxyManagementTokenCredential { + verified_token_hash: crate::management_token_auth::VerifiedManagementTokenHash::new( + management_token_hash(raw_token), + ), + token_id: REVOCATION_TOKEN_ID.to_string(), + user_id: REVOCATION_USER_ID.to_string(), + remote_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), + }), + ); + state.hub.register_proxy(Arc::clone(&connection)); + RegisteredTestProxy { + connection, + _outbound_rx: outbound_rx, + close_rx, + } + } + + async fn assert_connection_authorized(state: &AppState, proxy: &RegisteredTestProxy) { + let authorized = state + .authorized_proxy_connections_for_new_stream(REVOCATION_NODE_ID) + .await + .expect("current connection credential should validate") + .expect("repository-backed validation should return an authorization set"); + assert_eq!(authorized.len(), 1); + assert!(authorized.contains(&proxy.connection.id)); + assert!(proxy.connection.is_available()); + assert!(!*proxy.close_rx.borrow()); + } + + async fn assert_connection_revoked(state: &AppState, proxy: &RegisteredTestProxy) { + let error = state + .authorized_proxy_connections_for_new_stream(REVOCATION_NODE_ID) + .await + .expect_err("revoked connection must not authorize a new logical stream"); + assert_eq!(error, "proxy tunnel credential was revoked"); + assert!(!proxy.connection.is_available()); + assert!(*proxy.close_rx.borrow()); + } + + fn now() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock") + .as_secs() + } + + fn proof_headers(timestamp: u64, nonce: &str) -> HeaderMap { + let signature = sign_tunnel_security_handshake_for_generation( + VALID_PSK, + NODE_ID, + TEST_TUNNEL_GENERATION, + TUNNEL_SECURITY_NON_TLS_REQUIRED, + SESSION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + timestamp, + nonce, + ) + .expect("proof should sign"); + let mut headers = HeaderMap::new(); + headers.insert( + TUNNEL_SECURITY_HEADER, + HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED), + ); + headers.insert( + TUNNEL_SECURITY_SESSION_HEADER, + HeaderValue::from_static(SESSION), + ); + headers.insert( + TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, + HeaderValue::from_str(×tamp.to_string()).expect("timestamp header"), + ); + headers.insert( + TUNNEL_SECURITY_PROOF_NONCE_HEADER, + HeaderValue::from_str(nonce).expect("nonce header"), + ); + headers.insert( + TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + HeaderValue::from_str(&signature).expect("signature header"), + ); + headers + } + + #[test] + fn proxy_tunnel_security_rejects_missing_registered_psk() { + assert_eq!( + resolve_proxy_tunnel_security( + None, + Some(TUNNEL_SECURITY_NON_TLS_REQUIRED), + Some("session-1".to_string()), + ), + Err(ProxyTunnelSecurityError::MissingKey) + ); + assert_eq!( + resolve_proxy_tunnel_security(None, None, None), + Err(ProxyTunnelSecurityError::MissingKey) + ); + } + + #[test] + fn proxy_tunnel_security_rejects_invalid_registered_psk() { + assert_eq!( + resolve_proxy_tunnel_security( + Some("not-a-valid-32-byte-base64-key".to_string()), + Some(TUNNEL_SECURITY_NON_TLS_REQUIRED), + Some("session-1".to_string()), + ), + Err(ProxyTunnelSecurityError::InvalidKey) + ); + } + + #[test] + fn proxy_tunnel_security_requires_declared_supported_mode() { + assert_eq!( + resolve_proxy_tunnel_security( + Some(VALID_PSK.to_string()), + None, + Some("session-1".to_string()), + ), + Err(ProxyTunnelSecurityError::MissingMode) + ); + assert_eq!( + resolve_proxy_tunnel_security( + Some(VALID_PSK.to_string()), + Some("unsupported"), + Some("session-1".to_string()), + ), + Err(ProxyTunnelSecurityError::UnsupportedMode) + ); + } + + #[test] + fn proxy_tunnel_security_requires_session() { + assert_eq!( + resolve_proxy_tunnel_security( + Some(VALID_PSK.to_string()), + Some(TUNNEL_SECURITY_NON_TLS_REQUIRED), + None, + ), + Err(ProxyTunnelSecurityError::MissingSession) + ); + } + + #[test] + fn proxy_tunnel_security_accepts_valid_psk_mode_and_session() { + assert_eq!( + resolve_proxy_tunnel_security( + Some(VALID_PSK.to_string()), + Some(TUNNEL_SECURITY_NON_TLS_REQUIRED), + Some("session-1".to_string()), + ), + Ok((VALID_PSK.to_string(), "session-1".to_string())) + ); + } + + #[test] + fn proxy_tunnel_management_fallback_requires_admin_permission() { + assert!(!management_token_may_connect_proxy_tunnel(&[ + "admin:proxy_nodes:write".to_string(), + ])); + assert!(management_token_may_connect_proxy_tunnel(&[ + "admin:proxy_nodes:admin".to_string(), + ])); + } + + #[tokio::test] + async fn proxy_tunnel_security_accepts_valid_proof_once_and_rejects_replay() { + let state = state(); + let headers = proof_headers(now(), "nonce-valid-proof-0001"); + + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &headers, + NODE_ID, + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Ok((VALID_PSK.to_string(), SESSION.to_string())) + ); + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &headers, + NODE_ID, + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Err(ProxyTunnelSecurityError::Replay) + ); + } + + #[tokio::test] + async fn proxy_tunnel_handshake_fails_closed_without_proxy_node_reader() { + let state = state(); + state.register_secure_tunnel_key(NODE_ID, VALID_PSK); + assert_eq!(state.secure_tunnel_key(NODE_ID).as_deref(), Some(VALID_PSK)); + + let error = state + .secure_tunnel_binding_for_handshake(NODE_ID, "attacker-selected-generation") + .await + .expect_err("cached PSK must not replace the authoritative node generation"); + assert!(matches!( + error, + aether_data::DataLayerError::InvalidConfiguration(_) + )); + } + + #[tokio::test] + async fn existing_proxy_connection_cannot_open_stream_without_proxy_node_reader() { + let state = state(); + let proxy = register_psk_connection(&state, VALID_PSK); + + let error = state + .authorized_proxy_connections_for_new_stream(REVOCATION_NODE_ID) + .await + .expect_err("new streams require authoritative credential revalidation"); + assert_eq!( + error, + super::control_plane::CONTROL_PLANE_CREDENTIAL_UNAVAILABLE + ); + assert!(proxy.connection.is_available()); + } + + #[tokio::test] + async fn proxy_tunnel_security_rejects_tampering_expiry_and_duplicate_headers() { + let state = state(); + let timestamp = now(); + let headers = proof_headers(timestamp, "nonce-tamper-proof-0001"); + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &headers, + "node-proof-forged", + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Err(ProxyTunnelSecurityError::InvalidProof) + ); + + let expired = proof_headers( + timestamp.saturating_sub(TUNNEL_SECURITY_PROOF_CLOCK_SKEW_SECS + 1), + "nonce-expired-proof-01", + ); + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &expired, + NODE_ID, + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Err(ProxyTunnelSecurityError::InvalidProof) + ); + + let mut duplicate = proof_headers(timestamp, "nonce-duplicate-proof-1"); + duplicate.append( + TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + HeaderValue::from_static("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), + ); + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &duplicate, + NODE_ID, + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Err(ProxyTunnelSecurityError::MalformedHeader) + ); + } + + #[tokio::test] + async fn proxy_tunnel_security_requires_complete_proof_before_nonce_claim() { + let state = state(); + let mut headers = proof_headers(now(), "nonce-missing-proof-001"); + headers.remove(TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER); + + assert_eq!( + state + .authenticate_proxy_tunnel_security( + &headers, + NODE_ID, + TEST_TUNNEL_GENERATION, + CURRENT_TUNNEL_PROTOCOL_VERSION, + Some(VALID_PSK.to_string()), + ) + .await, + Err(ProxyTunnelSecurityError::MissingProof) + ); + } + + #[tokio::test] + async fn control_plane_authentication_fails_closed_without_proxy_node_reader() { + let state = state(); + state.register_secure_tunnel_key(NODE_ID, VALID_PSK); + let body = br#"{"node_id":"node-proof-1"}"#; + let timestamp = now(); + let nonce = "control-plane-no-reader-0001"; + let signature = sign_tunnel_control_plane_request_for_generation( + VALID_PSK, + "POST", + aether_gateway_tunnel::TUNNEL_HEARTBEAT_PATH, + NODE_ID, + TEST_TUNNEL_GENERATION, + timestamp, + nonce, + body, + ) + .expect("control-plane request should sign"); + let mut headers = HeaderMap::new(); + headers.insert( + TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, + HeaderValue::from_static(NODE_ID), + ); + headers.insert( + TUNNEL_CONTROL_PLANE_GENERATION_HEADER, + HeaderValue::from_static(TEST_TUNNEL_GENERATION), + ); + headers.insert( + TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, + HeaderValue::from_str(×tamp.to_string()).expect("timestamp header"), + ); + headers.insert( + TUNNEL_CONTROL_PLANE_NONCE_HEADER, + HeaderValue::from_static(nonce), + ); + headers.insert( + TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, + HeaderValue::from_str(&signature).expect("signature header"), + ); + + let result = state + .authenticate_control_plane_request( + &headers, + "POST", + aether_gateway_tunnel::TUNNEL_HEARTBEAT_PATH, + NODE_ID, + body, + ) + .await; + + assert_eq!(result, Err(ControlPlaneAuthError::Unavailable)); + } + + #[tokio::test] + async fn new_stream_revalidation_rejects_rotated_psk_without_local_close() { + let fixture = revocation_fixture(Some(VALID_PSK)); + let proxy = register_psk_connection(&fixture.state, VALID_PSK); + assert_connection_authorized(&fixture.state, &proxy).await; + + let node = fixture + .node_repository + .find_proxy_node(REVOCATION_NODE_ID) + .await + .expect("node lookup should succeed") + .expect("node should exist"); + let previous = node + .proxy_metadata + .expect("validated PSK should remain in protected metadata"); + let replacement = json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": ROTATED_PSK, + } + }); + assert!(fixture + .node_repository + .compare_and_set_proxy_metadata(REVOCATION_NODE_ID, &previous, &replacement) + .await + .expect("PSK rotation should succeed")); + + assert_connection_revoked(&fixture.state, &proxy).await; + } + + #[tokio::test] + async fn new_stream_revalidation_rejects_deleted_node_without_local_close() { + let fixture = revocation_fixture(Some(VALID_PSK)); + let proxy = register_psk_connection(&fixture.state, VALID_PSK); + assert_connection_authorized(&fixture.state, &proxy).await; + + fixture + .node_repository + .delete_node(REVOCATION_NODE_ID) + .await + .expect("node deletion should succeed") + .expect("node should exist before deletion"); + + assert_connection_revoked(&fixture.state, &proxy).await; + } + + #[tokio::test] + async fn same_node_id_recreated_with_same_psk_does_not_rebind_old_connection() { + let fixture = revocation_fixture(Some(VALID_PSK)); + let proxy = register_psk_connection(&fixture.state, VALID_PSK); + assert_connection_authorized(&fixture.state, &proxy).await; + let deleted_generation = proxy.connection.node_generation.clone(); + + fixture + .node_repository + .delete_node(REVOCATION_NODE_ID) + .await + .expect("node deletion should succeed") + .expect("node should exist before deletion"); + let replacement = fixture + .node_repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some(REVOCATION_NODE_ID.to_string()), + name: "replacement node".to_string(), + ip: "127.0.0.1".to_string(), + port: 0, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": TUNNEL_SECURITY_NON_TLS_REQUIRED, + "encryption_key": VALID_PSK, + } + })), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("same-id replacement node should register"); + + assert_eq!(replacement.id, REVOCATION_NODE_ID); + assert_ne!(replacement.tunnel_generation, deleted_generation); + assert_connection_revoked(&fixture.state, &proxy).await; + + let stale_heartbeat = fixture + .node_repository + .apply_heartbeat(&ProxyNodeHeartbeatMutation { + node_id: REVOCATION_NODE_ID.to_string(), + expected_tunnel_generation: Some(deleted_generation.clone()), + heartbeat_interval: Some(1), + active_connections: Some(99), + total_requests_delta: Some(500), + avg_latency_ms: Some(1.0), + failed_requests_delta: Some(5), + dns_failures_delta: Some(4), + stream_errors_delta: Some(3), + proxy_metadata: Some(json!({"forged": true})), + proxy_version: Some("stale".to_string()), + }) + .await + .expect("stale heartbeat should be handled without repository failure"); + assert!(stale_heartbeat.is_none()); + let stale_status = fixture + .node_repository + .update_tunnel_status(&ProxyNodeTunnelStatusMutation { + node_id: REVOCATION_NODE_ID.to_string(), + expected_tunnel_generation: Some(deleted_generation), + connected: true, + conn_count: 99, + detail: Some("stale".to_string()), + observed_at_unix_secs: Some(now()), + }) + .await + .expect("stale status should be handled without repository failure"); + assert!(stale_status.is_none()); + let unchanged = fixture + .node_repository + .find_proxy_node(REVOCATION_NODE_ID) + .await + .expect("replacement lookup should succeed") + .expect("replacement should remain"); + assert_eq!(unchanged.tunnel_generation, replacement.tunnel_generation); + assert_eq!(unchanged.active_connections, 0); + assert_eq!(unchanged.total_requests, 0); + assert!(unchanged + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.get("forged")) + .is_none()); + } + + #[tokio::test] + async fn management_token_disable_on_another_gateway_revokes_existing_connection() { + let fixture = revocation_fixture(None); + let shared_data = Arc::clone(&fixture.data); + let gateway_a = fixture.state; + let gateway_b = state().with_data(shared_data); + let proxy = register_management_token_connection(&gateway_b, RAW_MANAGEMENT_TOKEN); + assert_connection_authorized(&gateway_b, &proxy).await; + + let outcome = gateway_a + .data + .set_management_token_active(REVOCATION_TOKEN_ID, false) + .await + .expect("management token disable should succeed"); + assert!(outcome.is_some()); + + // gateway_b never receives an in-process close request. Its next-stream + // strong read of the shared repository must still observe the revocation. + assert_connection_revoked(&gateway_b, &proxy).await; + } + + #[tokio::test] + async fn new_stream_revalidation_rejects_deleted_management_token() { + let fixture = revocation_fixture(None); + let proxy = register_management_token_connection(&fixture.state, RAW_MANAGEMENT_TOKEN); + assert_connection_authorized(&fixture.state, &proxy).await; + + assert!(fixture + .token_repository + .delete_management_token(REVOCATION_TOKEN_ID) + .await + .expect("management token deletion should succeed")); + + assert_connection_revoked(&fixture.state, &proxy).await; + } + + #[tokio::test] + async fn new_stream_revalidation_rejects_regenerated_management_token_secret() { + let fixture = revocation_fixture(None); + let proxy = register_management_token_connection(&fixture.state, RAW_MANAGEMENT_TOKEN); + assert_connection_authorized(&fixture.state, &proxy).await; + + fixture + .token_repository + .regenerate_management_token_secret(&RegenerateManagementTokenSecret { + token_id: REVOCATION_TOKEN_ID.to_string(), + token_hash: management_token_hash(ROTATED_MANAGEMENT_TOKEN), + token_prefix: Some("ae-tunnel-rotated".to_string()), + }) + .await + .expect("management token regeneration should succeed") + .expect("management token should exist"); + + assert_connection_revoked(&fixture.state, &proxy).await; + let rotated_proxy = + register_management_token_connection(&fixture.state, ROTATED_MANAGEMENT_TOKEN); + assert_connection_authorized(&fixture.state, &rotated_proxy).await; + } + + #[tokio::test] + async fn same_token_id_recreated_with_different_hash_does_not_rebind_old_connection() { + let fixture = revocation_fixture(None); + let proxy = register_management_token_connection(&fixture.state, RAW_MANAGEMENT_TOKEN); + assert_connection_authorized(&fixture.state, &proxy).await; + + assert!(fixture + .token_repository + .delete_management_token(REVOCATION_TOKEN_ID) + .await + .expect("original management token deletion should succeed")); + fixture + .token_repository + .create_management_token(&management_token_create_record(ROTATED_MANAGEMENT_TOKEN)) + .await + .expect("same-id replacement management token should be created"); + + assert_connection_revoked(&fixture.state, &proxy).await; + let replacement_proxy = + register_management_token_connection(&fixture.state, ROTATED_MANAGEMENT_TOKEN); + assert_connection_authorized(&fixture.state, &replacement_proxy).await; + } +} + pub fn build_router_with_state(state: AppState) -> Router { middleware::apply_cf_header_stripping( Router::new() @@ -238,60 +2146,154 @@ async fn metrics(State(state): State) -> impl IntoResponse { pub async fn ws_proxy( ws: WebSocketUpgrade, State(state): State, + ConnectInfo(remote_addr): ConnectInfo, headers: HeaderMap, ) -> impl IntoResponse { - let node_id = headers - .get("x-node-id") - .and_then(|v| v.to_str().ok()) - .unwrap_or("") - .trim() - .to_string(); + let node_id = match required_proxy_tunnel_header( + &headers, + "x-node-id", + MAX_TUNNEL_SECURITY_NODE_ID_LEN, + ProxyTunnelSecurityError::MalformedHeader, + ) { + Ok(node_id) if valid_proxy_tunnel_token(node_id, 1) => node_id.to_string(), + _ => { + warn!("proxy connection rejected: invalid X-Node-ID header"); + return axum::http::StatusCode::BAD_REQUEST.into_response(); + } + }; let node_name = resolve_proxy_node_name(&headers, &node_id); - let max_streams = resolve_proxy_max_streams(&headers, state.max_streams); - let protocol_version = resolve_proxy_protocol_version(&headers); - let tunnel_security = headers - .get(aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER) - .and_then(|value| value.to_str().ok()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string); - let security_session = headers - .get(aether_contracts::tunnel_security::TUNNEL_SECURITY_SESSION_HEADER) - .and_then(|value| value.to_str().ok()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string); + let requested_generation = match required_proxy_tunnel_header( + &headers, + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + 128, + ProxyTunnelSecurityError::MalformedHeader, + ) { + Ok(generation) if valid_proxy_tunnel_token(generation, 1) => generation.to_string(), + _ => { + warn!(node_id = %node_id, "proxy connection rejected: invalid tunnel generation header"); + return axum::http::StatusCode::BAD_REQUEST.into_response(); + } + }; - if node_id.is_empty() { - warn!("proxy connection rejected: missing X-Node-ID header"); + let max_streams = resolve_proxy_max_streams(&headers, state.max_streams); + let raw_protocol_version = match required_proxy_tunnel_header( + &headers, + aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, + 3, + ProxyTunnelSecurityError::MalformedHeader, + ) + .and_then(|value| { + value + .parse::() + .ok() + .filter(|value| { + (1..=aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION).contains(value) + }) + .ok_or(ProxyTunnelSecurityError::MalformedHeader) + }) { + Ok(version) => version, + Err(error) => return error.status_code().into_response(), + }; + let protocol_version = resolve_proxy_protocol_version(&headers); + if protocol_version != raw_protocol_version { return axum::http::StatusCode::BAD_REQUEST.into_response(); } - let stored_security_key = state.secure_tunnel_key_for_node(&node_id).await; - let (security_key, security_session) = match tunnel_security.as_deref() { - Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED) => { - match stored_security_key { - Some(key) => { - let Some(session) = security_session else { - warn!(node_id = %node_id, "secure tunnel requested without a security session"); - return axum::http::StatusCode::BAD_REQUEST.into_response(); - }; - (Some(key), session) - } - None => { - warn!(node_id = %node_id, "secure tunnel requested but no PSK is registered"); - return axum::http::StatusCode::UNAUTHORIZED.into_response(); + let stored_security_binding = match state + .secure_tunnel_binding_for_handshake(&node_id, &requested_generation) + .await + { + Ok(binding) => binding, + Err(error) => { + warn!(node_id = %node_id, error = %error, "proxy connection rejected: tunnel security key lookup unavailable"); + return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response(); + } + }; + let tunnel_security = + if headers.contains_key(aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER) { + match required_proxy_tunnel_header( + &headers, + aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER, + 32, + ProxyTunnelSecurityError::MissingMode, + ) { + Ok(mode) => Some(mode), + Err(error) => return error.status_code().into_response(), + } + } else { + None + }; + let stored_security_key = stored_security_binding.as_ref().map(|(key, _)| key.clone()); + if stored_security_binding + .as_ref() + .is_some_and(|(_, generation)| generation != &requested_generation) + { + warn!(node_id = %node_id, "proxy connection rejected: stale tunnel generation"); + return ProxyTunnelSecurityError::InvalidProof + .status_code() + .into_response(); + } + let (security_key, security_session, management_token_credential, node_generation) = + match tunnel_security { + Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED) => { + match state + .authenticate_proxy_tunnel_security( + &headers, + &node_id, + &requested_generation, + protocol_version, + stored_security_key, + ) + .await + { + Ok(security) => ( + Some(security.0), + security.1, + None, + requested_generation.clone(), + ), + Err(error) => { + warn!(node_id = %node_id, ?error, "proxy connection rejected: invalid secure tunnel handshake"); + return error.status_code().into_response(); + } } } - } - Some(_) => return axum::http::StatusCode::BAD_REQUEST.into_response(), - None if stored_security_key.is_some() => { - warn!(node_id = %node_id, "proxy connection rejected: stored secure tunnel key requires encrypted frames"); - return axum::http::StatusCode::UNAUTHORIZED.into_response(); - } - None => (None, String::new()), - }; + Some(_) => { + return ProxyTunnelSecurityError::UnsupportedMode + .status_code() + .into_response() + } + None if stored_security_key.is_some() => { + warn!( + node_id = %node_id, + "proxy connection rejected: registered secure tunnel requires encrypted frames" + ); + return ProxyTunnelSecurityError::MissingMode + .status_code() + .into_response(); + } + None => { + let client_ip = crate::headers::effective_client_ip(&headers, &remote_addr); + let (credential, node_generation) = match state + .authenticate_proxy_tunnel_management_token(&headers, &node_id, client_ip) + .await + { + Ok(credential) => credential, + Err(error) => { + warn!(node_id = %node_id, ?error, "proxy connection rejected: invalid management token"); + return error.status_code().into_response(); + } + }; + if node_generation != requested_generation { + warn!(node_id = %node_id, "proxy connection rejected: stale tunnel generation"); + return ProxyTunnelSecurityError::InvalidAuthorization + .status_code() + .into_response(); + } + (None, String::new(), Some(credential), node_generation) + } + }; let request_permit = match state.try_acquire_request_permit().await { Ok(permit) => permit, @@ -326,10 +2328,12 @@ pub async fn ws_proxy( state.hub, node_id, node_name, + node_generation, max_streams, protocol_version, security_key, security_session, + management_token_credential, state.proxy_conn_cfg, ) .await diff --git a/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs b/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs index 8af2d0316..9834e510f 100644 --- a/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs +++ b/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs @@ -11,44 +11,80 @@ use futures_util::{SinkExt, StreamExt}; use tokio::sync::watch; use tracing::{debug, info, warn}; -use super::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus}; +use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus}; use super::protocol; -use aether_contracts::tunnel::Frame; +use aether_contracts::tunnel::{Frame, HelloPayload, MsgType}; use aether_contracts::tunnel_security::{SecureFrameCodec, TunnelSecurityRole}; /// Maximum single frame size: 64 MB const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024; +/// A connection that has passed the HTTP proof must still complete the +/// encrypted protocol handshake promptly. Keeping this deadline independent +/// from the normal idle timeout prevents half-open authenticated sockets from +/// holding an admission permit indefinitely. +const PROXY_HELLO_TIMEOUT: Duration = Duration::from_secs(10); +const MAX_PREAUTH_PINGS: usize = 8; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProxyHelloValidationError { + MalformedFrame, + DecryptionFailed, + UnexpectedFrame, + InvalidPayload, + ProtocolVersionMismatch, + SecuritySessionMismatch, +} pub async fn handle_proxy_connection( ws: WebSocket, hub: Arc, node_id: String, node_name: String, + node_generation: String, max_streams: usize, protocol_version: u8, security_key: Option, security_session: String, + management_token_credential: Option, cfg: ConnConfig, ) { let conn_id = hub.alloc_conn_id(); - let (mut ws_tx, ws_rx) = ws.split(); + let (mut ws_tx, mut ws_rx) = ws.split(); - let (tx, mut rx) = bounded_queue::(cfg.outbound_queue_capacity); - let (close_tx, mut close_rx) = watch::channel(false); - let security = match security_key.as_deref() { - Some(key) => { - match SecureFrameCodec::new(key, &security_session, TunnelSecurityRole::Server) { - Ok(codec) => Some(Arc::new(codec)), + let (security, initial_hello) = match security_key.as_deref() { + Some(security_key) => { + let security = match SecureFrameCodec::new( + security_key, + &security_session, + TunnelSecurityRole::Server, + ) { + Ok(codec) => Arc::new(codec), Err(error) => { warn!(conn_id, node_id = %node_id, error = %error, "secure tunnel codec initialization failed"); return; } - } + }; + let Some(hello) = read_authenticated_proxy_hello( + &mut ws_tx, + &mut ws_rx, + security.as_ref(), + protocol_version, + &security_session, + conn_id, + &node_id, + ) + .await + else { + return; + }; + (Some(security), Some(hello)) } - None => None, + None => (None, None), }; - let conn = Arc::new(ProxyConn::new( + let (tx, mut rx) = bounded_queue::(cfg.outbound_queue_capacity); + let (close_tx, mut close_rx) = watch::channel(false); + let conn = ProxyConn::new( conn_id, node_id.clone(), node_name.clone(), @@ -56,9 +92,21 @@ pub async fn handle_proxy_connection( close_tx, max_streams, protocol_version, - )); + ) + .with_tunnel_generation(node_generation); + let conn = match (security_key.clone(), management_token_credential) { + (Some(key), None) => Arc::new(conn.with_authenticated_key(key)), + (None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)), + (Some(_), Some(_)) | (None, None) => { + warn!(conn_id, node_id = %node_id, "proxy connection missing an unambiguous credential binding"); + return; + } + }; hub.register_proxy(conn.clone()); + if let Some(mut hello) = initial_hello { + hub.handle_proxy_frame(conn.id, &mut hello).await; + } let writer_conn_id = conn_id; let writer_conn = conn.clone(); @@ -75,7 +123,7 @@ pub async fn handle_proxy_connection( _ => 0, }; let send_started_at = std::time::Instant::now(); - let msg = match encrypt_message(msg, writer_security.as_deref()) { + let msg = match encrypt_message(msg, writer_security.as_deref()) { Ok(msg) => msg, Err(error) => { warn!(conn_id = writer_conn_id, error = %error, "failed to encrypt outbound proxy frame"); @@ -244,6 +292,120 @@ pub async fn handle_proxy_connection( let _ = writer.await; } +async fn read_authenticated_proxy_hello( + ws_tx: &mut futures_util::stream::SplitSink, + ws_rx: &mut futures_util::stream::SplitStream, + security: &SecureFrameCodec, + protocol_version: u8, + security_session: &str, + conn_id: u64, + node_id: &str, +) -> Option> { + let result = tokio::time::timeout(PROXY_HELLO_TIMEOUT, async { + let mut preauth_pings = 0usize; + loop { + match ws_rx.next().await { + Some(Ok(Message::Binary(data))) => { + return match validate_authenticated_proxy_hello( + data, + security, + protocol_version, + security_session, + ) { + Ok(hello) => Some(hello), + Err(error) => { + warn!( + conn_id, + node_id = %node_id, + ?error, + "proxy connection rejected: invalid encrypted HELLO" + ); + None + } + }; + } + Some(Ok(Message::Ping(payload))) => { + preauth_pings = preauth_pings.saturating_add(1); + if preauth_pings > MAX_PREAUTH_PINGS { + warn!( + conn_id, + node_id = %node_id, + "proxy connection rejected: too many WebSocket pings before encrypted HELLO" + ); + return None; + } + if let Err(error) = ws_tx.send(Message::Pong(payload)).await { + warn!(conn_id, node_id = %node_id, error = %error, "failed to answer WebSocket ping before proxy authentication"); + return None; + } + } + Some(Ok(Message::Pong(_))) => {} + Some(Ok(Message::Close(_))) | None => { + info!(conn_id, node_id = %node_id, "proxy disconnected before encrypted HELLO authentication"); + return None; + } + Some(Ok(Message::Text(_))) => { + warn!(conn_id, node_id = %node_id, "proxy connection rejected: text message received before encrypted HELLO"); + return None; + } + Some(Err(error)) => { + warn!(conn_id, node_id = %node_id, error = %error, "proxy WebSocket failed before encrypted HELLO authentication"); + return None; + } + } + } + }) + .await; + + match result { + Ok(hello) => hello, + Err(_) => { + warn!( + conn_id, + node_id = %node_id, + timeout_ms = PROXY_HELLO_TIMEOUT.as_millis(), + "proxy connection rejected: encrypted HELLO timed out" + ); + None + } + } +} + +fn validate_authenticated_proxy_hello( + data: bytes::Bytes, + security: &SecureFrameCodec, + protocol_version: u8, + security_session: &str, +) -> Result, ProxyHelloValidationError> { + let header = + protocol::FrameHeader::parse(&data).ok_or(ProxyHelloValidationError::MalformedFrame)?; + let expected_len = protocol::HEADER_SIZE + .checked_add(header.payload_len as usize) + .ok_or(ProxyHelloValidationError::MalformedFrame)?; + if expected_len != data.len() { + return Err(ProxyHelloValidationError::MalformedFrame); + } + + let frame = Frame::decode(data).map_err(|_| ProxyHelloValidationError::MalformedFrame)?; + let frame = security + .decrypt_frame(frame) + .map_err(|_| ProxyHelloValidationError::DecryptionFailed)?; + if frame.stream_id != 0 || frame.msg_type != MsgType::Hello || frame.flags != 0 { + return Err(ProxyHelloValidationError::UnexpectedFrame); + } + + let hello = serde_json::from_slice::(&frame.payload) + .map_err(|_| ProxyHelloValidationError::InvalidPayload)?; + if hello.protocol_version != protocol_version { + return Err(ProxyHelloValidationError::ProtocolVersionMismatch); + } + if hello.session_id.as_deref() != Some(security_session) { + return Err(ProxyHelloValidationError::SecuritySessionMismatch); + } + + Ok(frame.encode().to_vec()) +} + async fn run_proxy_reader( mut ws_rx: futures_util::stream::SplitStream, hub: Arc, @@ -358,3 +520,152 @@ fn decrypt_message( let frame = codec.decrypt_frame(frame)?; Ok(frame.encode().to_vec()) } + +#[cfg(test)] +mod tests { + use super::*; + + const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; + const SESSION: &str = "0123456789abcdef0123456789abcdef"; + const PROTOCOL_VERSION: u8 = 3; + + fn codecs() -> (SecureFrameCodec, SecureFrameCodec) { + ( + SecureFrameCodec::new(KEY, SESSION, TunnelSecurityRole::Client).expect("client codec"), + SecureFrameCodec::new(KEY, SESSION, TunnelSecurityRole::Server).expect("server codec"), + ) + } + + fn hello_frame(protocol_version: u8, session_id: &str) -> Frame { + Frame::control( + MsgType::Hello, + serde_json::to_vec(&HelloPayload { + protocol_version, + capabilities: vec!["flow-control".to_string()], + session_id: Some(session_id.to_string()), + replica_id: None, + }) + .expect("hello payload"), + ) + } + + #[test] + fn authenticated_proxy_hello_accepts_bound_encrypted_control_frame() { + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION)) + .expect("encrypted HELLO"); + + let clear = + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION) + .expect("authenticated HELLO"); + let frame = Frame::decode(bytes::Bytes::from(clear)).expect("clear HELLO"); + + assert_eq!(frame.stream_id, 0); + assert_eq!(frame.msg_type, MsgType::Hello); + } + + #[test] + fn authenticated_proxy_hello_advances_shared_receive_sequence() { + let (client, server) = codecs(); + let encrypted_hello = client + .encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION)) + .expect("encrypted HELLO"); + validate_authenticated_proxy_hello(encrypted_hello, &server, PROTOCOL_VERSION, SESSION) + .expect("authenticated HELLO"); + + let encrypted_settings = client + .encrypt_frame(Frame::control(MsgType::Settings, bytes::Bytes::new())) + .expect("encrypted SETTINGS"); + let settings = server + .decrypt_frame(Frame::decode(encrypted_settings).expect("wire SETTINGS")) + .expect("next sequence should decrypt"); + + assert_eq!(settings.msg_type, MsgType::Settings); + } + + #[test] + fn authenticated_proxy_hello_rejects_non_encrypted_or_wrong_frame() { + let (_, server) = codecs(); + let clear = hello_frame(PROTOCOL_VERSION, SESSION).encode(); + assert_eq!( + validate_authenticated_proxy_hello(clear, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::DecryptionFailed) + ); + + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(Frame::control(MsgType::Settings, bytes::Bytes::new())) + .expect("encrypted SETTINGS"); + assert_eq!( + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::UnexpectedFrame) + ); + + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(Frame::new( + 1, + MsgType::Hello, + 0, + hello_frame(PROTOCOL_VERSION, SESSION).payload, + )) + .expect("encrypted stream HELLO"); + assert_eq!( + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::UnexpectedFrame) + ); + } + + #[test] + fn authenticated_proxy_hello_rejects_protocol_or_session_mismatch() { + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(hello_frame(PROTOCOL_VERSION - 1, SESSION)) + .expect("encrypted HELLO"); + assert_eq!( + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::ProtocolVersionMismatch) + ); + + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(hello_frame(PROTOCOL_VERSION, "different-session")) + .expect("encrypted HELLO"); + assert_eq!( + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::SecuritySessionMismatch) + ); + } + + #[test] + fn authenticated_proxy_hello_rejects_malformed_or_ambiguous_frame() { + let (client, server) = codecs(); + let mut encrypted = client + .encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION)) + .expect("encrypted HELLO") + .to_vec(); + encrypted.push(0); + assert_eq!( + validate_authenticated_proxy_hello( + bytes::Bytes::from(encrypted), + &server, + PROTOCOL_VERSION, + SESSION, + ), + Err(ProxyHelloValidationError::MalformedFrame) + ); + + let (client, server) = codecs(); + let encrypted = client + .encrypt_frame(Frame::control( + MsgType::Hello, + bytes::Bytes::from_static(b"not-json"), + )) + .expect("encrypted malformed HELLO"); + assert_eq!( + validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION), + Err(ProxyHelloValidationError::InvalidPayload) + ); + } +} diff --git a/apps/aether-gateway/src/tunnel/mod.rs b/apps/aether-gateway/src/tunnel/mod.rs index b67b39418..bde158308 100644 --- a/apps/aether-gateway/src/tunnel/mod.rs +++ b/apps/aether-gateway/src/tunnel/mod.rs @@ -3,12 +3,19 @@ mod embedded; use std::collections::HashMap; use std::fmt; use std::io; -use std::sync::Arc; +use std::net::SocketAddr; +use std::path::PathBuf; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex, OnceLock}; use std::time::{Duration, SystemTime}; use aether_contracts::tunnel::{ - resolve_tunnel_request_timeouts, try_decode_tunnel_relay_request_meta, RequestMeta, - MAX_TUNNEL_RELAY_META_LEN, TUNNEL_RELAY_FORWARDED_BY_HEADER, + resolve_tunnel_request_timeouts, sign_tunnel_relay_request, + try_decode_tunnel_relay_request_meta, tunnel_relay_payload_digest, + tunnel_relay_payload_digest_from_hashes, RequestMeta, TunnelRelayPayloadDigest, + MAX_TUNNEL_RELAY_META_LEN, 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_FORWARDED_BY_HEADER, TUNNEL_RELAY_OWNER_INSTANCE_HEADER, }; use aether_data::repository::proxy_nodes::{ @@ -16,20 +23,22 @@ use aether_data::repository::proxy_nodes::{ }; use aether_gateway_tunnel::EmbeddedTunnelDefaults; use aether_runtime::MetricSample; -use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; +use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeLockLease, RuntimeState}; use async_stream::stream; use axum::body::{Body, Bytes}; use axum::extract::ws::WebSocketUpgrade; use axum::extract::{ConnectInfo, Path, Request, State}; -use axum::http::{HeaderMap, StatusCode}; +use axum::http::{HeaderMap, HeaderValue, Method, StatusCode, Uri}; use axum::response::IntoResponse; -use bytes::BytesMut; -use futures_util::StreamExt; use serde::{Deserialize, Serialize}; use serde_json::json; +use sha2::{Digest, Sha256}; +use tokio::io::{AsyncSeekExt, AsyncWriteExt}; +use tokio_util::io::ReaderStream; use tracing::warn; use self::embedded::{AppState as TunnelAppState, ConnConfig, ControlPlaneClient}; +pub(crate) use self::embedded::{ControlPlaneAuthError, RelayAuthError}; use super::api::response::{build_client_response, build_local_http_error_response}; use super::constants::TRACE_ID_HEADER; use super::data::GatewayDataState; @@ -55,9 +64,85 @@ const TUNNEL_ATTACHMENT_KEY_PREFIX: &str = "tunnel.attachments."; const TUNNEL_ATTACHMENT_REDIS_KEY_PREFIX: &str = "tunnel:attachments:"; const TUNNEL_INSTANCE_ID_ENV: &str = "AETHER_GATEWAY_INSTANCE_ID"; const TUNNEL_RELAY_BASE_URL_ENV: &str = "AETHER_TUNNEL_RELAY_BASE_URL"; +// Owner-forward relay URLs are deployment metadata, but a corrupted or +// stale attachment record must not turn the gateway into a generic internal +// network client. Private HTTPS relay targets therefore require an explicit +// operator opt-in, just like the tunnel's private upstream target escape +// hatch. Loopback HTTP remains supported for the local single-process path. +const TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS_ENV: &str = "AETHER_TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS"; +// Prefer this narrow host allowlist for private owner relay deployments. The +// legacy *_ALLOW_PRIVATE_TARGETS switch remains available for operators that +// intentionally trust every private address in their deployment, but an exact +// hostname list limits the blast radius of a stale or corrupted attachment. +const TUNNEL_RELAY_PRIVATE_HOST_ALLOWLIST_ENV: &str = "AETHER_TUNNEL_RELAY_PRIVATE_HOST_ALLOWLIST"; const TUNNEL_ATTACHMENT_TTL_ENV: &str = "AETHER_TUNNEL_ATTACHMENT_TTL_SECS"; +const TUNNEL_RELAY_AUTH_SECRET_ENV: &str = "AETHER_TUNNEL_RELAY_AUTH_SECRET"; +const TUNNEL_HEARTBEAT_STATE_KEY_PREFIX: &str = "tunnel:heartbeat:session:"; +const TUNNEL_HEARTBEAT_STATE_TTL: Duration = Duration::from_secs(30 * 24 * 60 * 60); +const TUNNEL_HEARTBEAT_LOCK_TTL: Duration = Duration::from_secs(60); +// Attachment updates touch both the shared runtime key and the durable +// system-config shadow. Keep the lease comfortably longer than the bounded +// operation so a slow database call cannot outlive the lock and race a new +// owner. RuntimeState also supports lock renewal, but these short operations +// are deliberately cancelled before renewal is needed. +const TUNNEL_ATTACHMENT_LOCK_TTL: Duration = Duration::from_secs(60); +const TUNNEL_ATTACHMENT_LOCK_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(10); +const TUNNEL_ATTACHMENT_OPERATION_TIMEOUT: Duration = Duration::from_secs(45); +const TUNNEL_ATTACHMENT_LOCK_RELEASE_TIMEOUT: Duration = Duration::from_secs(5); +const TUNNEL_RELAY_BODY_READ_TIMEOUT: Duration = Duration::from_secs(120); +const DEFAULT_TUNNEL_RELAY_MAX_BODY_BYTES: u64 = 256 * 1024 * 1024; +const MAX_TUNNEL_RELAY_BODY_BYTES: u64 = 1024 * 1024 * 1024; +const TUNNEL_RELAY_MAX_BODY_MB_ENV: &str = "AETHER_TUNNEL_RELAY_MAX_BODY_MB"; +const DEFAULT_TUNNEL_RELAY_SPOOL_BUDGET_BYTES: u64 = 1024 * 1024 * 1024; +const MAX_TUNNEL_RELAY_SPOOL_BUDGET_BYTES: u64 = 4 * 1024 * 1024 * 1024; +const TUNNEL_RELAY_SPOOL_BUDGET_MB_ENV: &str = "AETHER_TUNNEL_RELAY_SPOOL_BUDGET_MB"; +static TUNNEL_RELAY_SPOOL_BYTES_IN_USE: AtomicU64 = AtomicU64::new(0); +pub(crate) const TUNNEL_RELAY_AUTH_SECRET_MIN_BYTES: usize = 32; pub(crate) const TUNNEL_RELAY_ROLLOUT_PROBE_HEADER: &str = "x-aether-tunnel-rollout-probe"; pub(crate) const TUNNEL_RELAY_ROLLOUT_PROBE_VALUE: &str = "1"; +const TUNNEL_AFFINITY_AUTH_CONTEXT: &str = "aether.tunnel.affinity-auth.v1"; +const TUNNEL_AFFINITY_GATEWAY_MARKER: &str = "rust-phase3b-affinity"; +const MAX_TUNNEL_AFFINITY_ID_LEN: usize = 200; +const OWNER_FORWARD_DNS_LOOKUP_TIMEOUT: Duration = Duration::from_secs(10); +const MAX_OWNER_FORWARD_DNS_ADDRESSES: usize = 32; +// This bounds cached DNS answer/client combinations, not concurrent requests or HTTP/2 streams. +const MAX_OWNER_FORWARD_PINNED_CLIENT_CACHE_ENTRIES: usize = 256; + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct OwnerForwardDnsPinKey { + scheme: String, + host: String, + port: u16, + addresses: Vec, +} + +struct OwnerForwardPinnedClientCacheEntry { + client: reqwest::Client, + last_used: u64, +} + +fn owner_forward_pinned_client_cache( +) -> &'static Mutex> { + static CACHE: OnceLock< + Mutex>, + > = OnceLock::new(); + CACHE.get_or_init(|| Mutex::new(HashMap::new())) +} + +fn owner_forward_pinned_client_cache_clock() -> &'static AtomicU64 { + static CLOCK: OnceLock = OnceLock::new(); + CLOCK.get_or_init(|| AtomicU64::new(0)) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct TunnelAffinityAuthContext { + pub(crate) client_ip: std::net::IpAddr, +} + +pub(crate) struct PendingTunnelAffinityAuth { + context: TunnelAffinityAuthContext, + relay_auth: embedded::PendingRelayAuth, +} pub(crate) async fn send_owner_forward_request( request: reqwest::RequestBuilder, @@ -65,19 +150,287 @@ pub(crate) async fn send_owner_forward_request( ) -> Result { match first_byte_timeout { Some(timeout) => match tokio::time::timeout(timeout, request.send()).await { - Ok(result) => result.map_err(|error| error.to_string()), + Ok(result) => result.map_err(|error| owner_forward_request_error(&error)), Err(_) => Err(format!( "owner gateway first byte timeout after {} ms", timeout.as_millis() )), }, - None => request.send().await.map_err(|error| error.to_string()), + None => request + .send() + .await + .map_err(|error| owner_forward_request_error(&error)), } } +/// Project reqwest failures at the relay boundary. `reqwest::Error::to_string` +/// can include the complete request URL (including credentials, query tokens, +/// or an internal host), so it must never be copied into a client-facing +/// `GatewayError` or admin probe response. +pub(crate) fn owner_forward_request_error(error: &reqwest::Error) -> String { + if error.is_timeout() { + return "owner gateway request timed out".to_string(); + } + if error.is_connect() { + return "owner gateway connection failed".to_string(); + } + if error.is_redirect() { + return "owner gateway redirect was rejected".to_string(); + } + if error.is_body() { + return "owner gateway request body failed".to_string(); + } + if error.is_decode() { + return "owner gateway response decode failed".to_string(); + } + "owner gateway request failed".to_string() +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ResolvedOwnerForwardTarget { + scheme: String, + host: String, + port: u16, + addresses: Vec, + literal_host: bool, +} + +/// Resolve an owner URL once and return the exact addresses that must be used +/// for the ensuing request. A normal reqwest client performs another DNS +/// lookup when it opens a connection; that lookup can observe a different +/// answer after a validation pass (DNS rebinding). Callers use this result to +/// install a `resolve_to_addrs` override on the client that sends the request. +async fn resolve_owner_forward_target( + owner_url: &str, +) -> Result { + let url = url::Url::parse(owner_url.trim()) + .map_err(|error| format!("invalid owner gateway URL: {error}"))?; + validate_tunnel_relay_transport_url(&url)?; + if url.fragment().is_some() { + return Err("owner gateway URL must not include a fragment".to_string()); + } + let host = match url + .host() + .ok_or_else(|| "owner gateway URL must include a host".to_string())? + { + url::Host::Domain(host) => host.trim().to_string(), + url::Host::Ipv4(address) => address.to_string(), + url::Host::Ipv6(address) => address.to_string(), + }; + if host.is_empty() { + return Err("owner gateway URL must include a host".to_string()); + } + let port = url + .port_or_known_default() + .ok_or_else(|| "owner gateway URL must include a port".to_string())?; + let addresses = if let Ok(ip) = host.parse::() { + vec![SocketAddr::new(ip, port)] + } else { + let resolved = tokio::time::timeout( + OWNER_FORWARD_DNS_LOOKUP_TIMEOUT, + tokio::net::lookup_host((host.as_str(), port)), + ) + .await + .map_err(|_| "owner gateway DNS resolution timed out".to_string())? + .map_err(|_| "owner gateway DNS resolution failed".to_string())?; + // Keep the resolver iterator bounded before collecting it. A hostile + // or misconfigured resolver must not be able to force an unbounded + // address vector before the target policy gets a chance to inspect it. + resolved + .take(MAX_OWNER_FORWARD_DNS_ADDRESSES.saturating_add(1)) + .collect() + }; + build_owner_forward_target(url.scheme(), &host, port, addresses) +} + +fn build_owner_forward_target( + scheme: &str, + host: &str, + port: u16, + addresses: Vec, +) -> Result { + build_owner_forward_target_with_policy( + scheme, + host, + port, + addresses, + tunnel_relay_allows_private_targets() || tunnel_relay_private_host_is_allowlisted(host), + ) +} + +fn build_owner_forward_target_with_policy( + scheme: &str, + host: &str, + port: u16, + mut addresses: Vec, + allow_private_targets: bool, +) -> Result { + let host = host.trim(); + if host.is_empty() { + return Err("owner gateway URL must include a host".to_string()); + } + if !matches!(scheme.to_ascii_lowercase().as_str(), "http" | "https") { + return Err("owner gateway URL must use HTTP or HTTPS".to_string()); + } + if addresses.len() > MAX_OWNER_FORWARD_DNS_ADDRESSES { + return Err(format!( + "owner gateway DNS resolution returned too many addresses (maximum {})", + MAX_OWNER_FORWARD_DNS_ADDRESSES + )); + } + addresses.sort_unstable(); + addresses.dedup(); + if addresses.is_empty() { + return Err("owner gateway DNS resolution returned no addresses".to_string()); + } + if addresses.iter().any(|address| address.port() != port) { + return Err( + "owner gateway DNS resolution returned an address with the wrong port".to_string(), + ); + } + if scheme.eq_ignore_ascii_case("http") + && addresses.iter().any(|address| !address.ip().is_loopback()) + { + return Err("loopback HTTP owner gateway resolved to a non-loopback address".to_string()); + } + if !allow_private_targets + && addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + // The existing local-development exception is deliberately + // narrow: HTTPS never gets an implicit loopback/private + // exemption, while literal/localhost HTTP is checked above. + && !(scheme.eq_ignore_ascii_case("http") && address.ip().is_loopback()) + }) + { + return Err(format!( + "owner gateway DNS resolution returned a private or reserved address; set {TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS_ENV}=true only for a trusted internal relay" + )); + } + + Ok(ResolvedOwnerForwardTarget { + scheme: scheme.to_ascii_lowercase(), + host: host.to_string(), + port, + addresses, + literal_host: host.parse::().is_ok(), + }) +} + +fn tunnel_relay_allows_private_targets() -> bool { + std::env::var(TUNNEL_RELAY_ALLOW_PRIVATE_TARGETS_ENV) + .ok() + .is_some_and(|value| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "1" | "true" | "yes" | "on" + ) + }) +} + +/// Return whether a private owner relay hostname was explicitly allowlisted. +/// +/// Entries are comma-separated exact DNS names (case-insensitive); a trailing +/// dot is ignored for DNS canonicalisation. We intentionally do not support +/// arbitrary suffix or wildcard matching: allowing `.internal` would let a +/// compromised attachment redirect traffic to any sibling service in that +/// namespace. The broad `*_ALLOW_PRIVATE_TARGETS=true` switch above remains +/// the explicit escape hatch for deployments that need that behaviour. +fn tunnel_relay_private_host_is_allowlisted(host: &str) -> bool { + let host = host.trim().trim_end_matches('.').to_ascii_lowercase(); + if host.is_empty() { + return false; + } + std::env::var(TUNNEL_RELAY_PRIVATE_HOST_ALLOWLIST_ENV) + .ok() + .is_some_and(|value| tunnel_relay_private_host_matches_allowlist(&host, &value)) +} + +fn tunnel_relay_private_host_matches_allowlist(host: &str, allowlist: &str) -> bool { + let host = host.trim().trim_end_matches('.'); + !host.is_empty() + && allowlist.split(',').any(|entry| { + let entry = entry.trim().trim_end_matches('.'); + !entry.is_empty() && entry.eq_ignore_ascii_case(host) + }) +} + +fn build_owner_forward_pinned_client( + target: &ResolvedOwnerForwardTarget, +) -> Result { + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .connect_timeout(OWNER_FORWARD_DNS_LOOKUP_TIMEOUT) + .http2_adaptive_window(true) + .resolve_to_addrs(&target.host, &target.addresses) + .build() + .map_err(|_| "failed to build pinned owner gateway client".to_string()) +} + +/// Return a client whose DNS answer is fixed to the addresses resolved for +/// `owner_url`. Literal IP hosts are already immune to DNS rebinding and keep +/// using the configured shared client so test/deployment-specific client +/// settings (for example a deliberately short timeout) remain intact. +pub(crate) async fn owner_forward_client_for_url( + base_client: &reqwest::Client, + owner_url: &str, +) -> Result { + let target = resolve_owner_forward_target(owner_url).await?; + if target.literal_host { + return Ok(base_client.clone()); + } + + let key = OwnerForwardDnsPinKey { + scheme: target.scheme, + host: target.host, + port: target.port, + addresses: target.addresses, + }; + let clock = owner_forward_pinned_client_cache_clock().fetch_add(1, Ordering::Relaxed); + if let Ok(mut cache) = owner_forward_pinned_client_cache().lock() { + if let Some(entry) = cache.get_mut(&key) { + entry.last_used = clock; + return Ok(entry.client.clone()); + } + } + + let client = build_owner_forward_pinned_client(&ResolvedOwnerForwardTarget { + scheme: key.scheme.clone(), + host: key.host.clone(), + port: key.port, + addresses: key.addresses.clone(), + literal_host: false, + })?; + let mut cache = owner_forward_pinned_client_cache() + .lock() + .map_err(|_| "owner gateway DNS pin cache is unavailable".to_string())?; + if let Some(entry) = cache.get_mut(&key) { + entry.last_used = clock; + return Ok(entry.client.clone()); + } + if cache.len() >= MAX_OWNER_FORWARD_PINNED_CLIENT_CACHE_ENTRIES { + if let Some(oldest_key) = cache + .iter() + .min_by_key(|(_, entry)| entry.last_used) + .map(|(key, _)| key.clone()) + { + cache.remove(&oldest_key); + } + } + cache.insert( + key, + OwnerForwardPinnedClientCacheEntry { + client: client.clone(), + last_used: clock, + }, + ); + Ok(client) +} + #[derive(Debug, Deserialize)] struct InternalTunnelHeartbeatRequest { node_id: String, + heartbeat_session_id: String, heartbeat_id: u64, #[serde(default)] heartbeat_interval: Option, @@ -107,6 +460,10 @@ struct InternalTunnelHeartbeatRequest { proxy_version: Option, } +pub(crate) struct TunnelHeartbeatClaim { + lease: RuntimeLockLease, +} + #[derive(Debug, Clone)] pub(crate) struct TunnelInstanceIdentity { instance_id: String, @@ -171,13 +528,18 @@ impl TunnelAttachmentDirectory { &self.identity.instance_id } - async fn refresh_from_heartbeat( + async fn refresh_from_authenticated_heartbeat( &self, data: &GatewayDataState, + authenticated_node_id: &str, + authenticated_generation: &str, request_body: &[u8], ) -> Result<(), String> { let payload = parse_embedded_tunnel_heartbeat_request(request_body)?; let node_id = payload.node_id.trim(); + if node_id != authenticated_node_id { + return Err("heartbeat node_id does not match authenticated tunnel node".to_string()); + } let Some(node) = data .find_proxy_node(node_id) .await @@ -185,7 +547,7 @@ impl TunnelAttachmentDirectory { else { return Ok(()); }; - if !node.tunnel_connected { + if node.tunnel_generation != authenticated_generation || !node.tunnel_connected { return Ok(()); } @@ -203,6 +565,7 @@ impl TunnelAttachmentDirectory { &TunnelAttachmentRecord { gateway_instance_id: self.identity.instance_id.clone(), relay_base_url: relay_base_url.clone(), + tunnel_generation: authenticated_generation.to_string(), conn_count, observed_at_unix_secs: current_unix_secs(), }, @@ -214,6 +577,7 @@ impl TunnelAttachmentDirectory { &self, data: &GatewayDataState, node_id: &str, + node_generation: &str, connected: bool, conn_count: usize, observed_at_unix_secs: u64, @@ -223,7 +587,8 @@ impl TunnelAttachmentDirectory { return Ok(()); } if !connected || conn_count == 0 { - self.delete_attachment_record(data, node_id).await?; + self.delete_attachment_record_if_owned(data, node_id, Some(node_generation)) + .await?; return Ok(()); } @@ -236,6 +601,7 @@ impl TunnelAttachmentDirectory { &TunnelAttachmentRecord { gateway_instance_id: self.identity.instance_id.clone(), relay_base_url: relay_base_url.clone(), + tunnel_generation: node_generation.to_string(), conn_count, observed_at_unix_secs, }, @@ -251,6 +617,14 @@ impl TunnelAttachmentDirectory { let Some(record) = self.read_attachment_record(data, node_id).await? else { return Ok(None); }; + let current_generation = data + .find_proxy_node(node_id) + .await + .map_err(|err| format!("attachment node lookup failed: {err}"))? + .map(|node| node.tunnel_generation); + if current_generation.as_deref() != Some(record.tunnel_generation.as_str()) { + return Ok(None); + } if !record.is_routable(current_unix_secs(), self.identity.attachment_ttl_secs) { return Ok(None); } @@ -266,7 +640,8 @@ impl TunnelAttachmentDirectory { return Ok(()); }; if record.is_owned_by(&self.identity.instance_id) { - self.delete_attachment_record(data, node_id).await?; + self.delete_attachment_record_if_owned(data, node_id, None) + .await?; } Ok(()) } @@ -329,6 +704,29 @@ impl TunnelAttachmentDirectory { data: &GatewayDataState, node_id: &str, record: &TunnelAttachmentRecord, + ) -> Result<(), String> { + let lease = self.acquire_attachment_lock(node_id).await?; + let result = match tokio::time::timeout( + TUNNEL_ATTACHMENT_OPERATION_TIMEOUT, + self.write_attachment_record_unlocked(data, node_id, record), + ) + .await + { + Ok(result) => result, + Err(_) => Err(format!( + "attachment write timed out after {} seconds", + TUNNEL_ATTACHMENT_OPERATION_TIMEOUT.as_secs() + )), + }; + self.release_attachment_lock(lease).await; + result + } + + async fn write_attachment_record_unlocked( + &self, + data: &GatewayDataState, + node_id: &str, + record: &TunnelAttachmentRecord, ) -> Result<(), String> { let serialized = serde_json::to_string(record) .map_err(|err| format!("attachment serialization failed: {err}"))?; @@ -355,26 +753,105 @@ impl TunnelAttachmentDirectory { .map_err(|err| format!("attachment write failed: {err}")) } - async fn delete_attachment_record( + async fn acquire_attachment_lock(&self, node_id: &str) -> Result { + let key = format!("tunnel:attachments:lock:{}", node_id.trim()); + let owner = format!("tunnel-attachment:{}", self.identity.instance_id); + let acquired = tokio::time::timeout( + TUNNEL_ATTACHMENT_LOCK_ACQUIRE_TIMEOUT, + self.runtime_state + .lock_try_acquire(&key, &owner, TUNNEL_ATTACHMENT_LOCK_TTL), + ) + .await + .map_err(|_| { + format!( + "attachment lock acquisition timed out after {} seconds", + TUNNEL_ATTACHMENT_LOCK_ACQUIRE_TIMEOUT.as_secs() + ) + })? + .map_err(|err| format!("attachment lock failed: {err}"))?; + acquired.ok_or_else(|| "attachment update is busy; retry later".to_string()) + } + + async fn release_attachment_lock(&self, lease: RuntimeLockLease) { + match tokio::time::timeout( + TUNNEL_ATTACHMENT_LOCK_RELEASE_TIMEOUT, + self.runtime_state.lock_release(&lease), + ) + .await + { + Ok(Ok(_)) => {} + Ok(Err(error)) => { + warn!(error = %error, key = %lease.key, "failed to release tunnel attachment lock"); + } + Err(_) => { + warn!( + key = %lease.key, + timeout_secs = TUNNEL_ATTACHMENT_LOCK_RELEASE_TIMEOUT.as_secs(), + "timed out releasing tunnel attachment lock" + ); + } + } + } + + async fn delete_attachment_record_if_owned( &self, data: &GatewayDataState, node_id: &str, + expected_tunnel_generation: Option<&str>, ) -> Result<(), String> { - if let Err(error) = self - .runtime_state - .kv_delete(&tunnel_attachment_redis_key(node_id)) - .await + let lease = self.acquire_attachment_lock(node_id).await?; + + let result = match tokio::time::timeout(TUNNEL_ATTACHMENT_OPERATION_TIMEOUT, async { + // Read both shadows while holding the same distributed lock used + // by writers. If either store already belongs to another gateway, + // leave both records untouched; this prevents an old disconnect + // event from deleting a replacement owner. + let runtime_record = self.read_attachment_record_from_runtime(node_id).await?; + let config_record = self + .read_attachment_record_from_system_config(data, node_id) + .await?; + let expected_owner = self.identity.instance_id.as_str(); + if runtime_record.as_ref().is_some_and(|record| { + !record.is_owned_by(expected_owner) + || expected_tunnel_generation + .is_some_and(|expected| record.tunnel_generation != expected) + }) || config_record.as_ref().is_some_and(|record| { + !record.is_owned_by(expected_owner) + || expected_tunnel_generation + .is_some_and(|expected| record.tunnel_generation != expected) + }) { + return Ok(()); + } + + if runtime_record.is_some() + && !self + .runtime_state + .kv_delete(&tunnel_attachment_redis_key(node_id)) + .await + .map_err(|error| format!("attachment runtime delete failed: {error}"))? + { + return Err("attachment runtime record disappeared before delete".to_string()); + } + if config_record.is_some() + && !data + .delete_system_config_value(&tunnel_attachment_key(node_id)) + .await + .map_err(|err| format!("attachment delete failed: {err}"))? + { + return Err("attachment system record disappeared before delete".to_string()); + } + Ok(()) + }) + .await { - warn!( - error = %error, - node_id = %node_id, - "failed to delete tunnel attachment from runtime state; clearing system_config shadow anyway" - ); - } - data.delete_system_config_value(&tunnel_attachment_key(node_id)) - .await - .map(|_| ()) - .map_err(|err| format!("attachment delete failed: {err}")) + Ok(result) => result, + Err(_) => Err(format!( + "attachment delete timed out after {} seconds", + TUNNEL_ATTACHMENT_OPERATION_TIMEOUT.as_secs() + )), + }; + self.release_attachment_lock(lease).await; + result } } @@ -382,6 +859,7 @@ impl TunnelAttachmentDirectory { pub(crate) struct EmbeddedTunnelState { inner: TunnelAppState, attachment_directory: TunnelAttachmentDirectory, + relay_auth_secret: Result, Arc>, } #[derive(Debug, Clone, Copy, Serialize)] @@ -397,6 +875,300 @@ pub(crate) struct TunnelProbeResponse { pub(crate) body: String, } +pub(crate) struct RelayAuthHeaders { + sender_instance_id: String, + owner_instance_id: String, + timestamp_unix_secs: u64, + nonce: String, + payload_digest: String, + signature: String, +} + +#[derive(Clone)] +pub(crate) struct VerifiedRelaySpool { + inner: Arc, +} + +struct RelaySpoolInner { + path: PathBuf, + meta: RequestMeta, + metadata_envelope: Bytes, + body_offset: u64, + body_len: u64, + body_sha256: [u8; 32], + reserved_bytes: u64, +} + +impl Drop for RelaySpoolInner { + fn drop(&mut self) { + TUNNEL_RELAY_SPOOL_BYTES_IN_USE.fetch_sub(self.reserved_bytes, Ordering::AcqRel); + let _ = std::fs::remove_file(&self.path); + } +} + +impl VerifiedRelaySpool { + pub(crate) fn meta(&self) -> &RequestMeta { + &self.inner.meta + } + + async fn open_from_start(&self) -> Result { + tokio::fs::File::open(&self.inner.path) + .await + .map_err(|error| format!("failed to reopen tunnel relay spool: {error}")) + } + + async fn open_body(&self) -> Result { + let mut file = self.open_from_start().await?; + file.seek(std::io::SeekFrom::Start(self.inner.body_offset)) + .await + .map_err(|error| format!("failed to seek tunnel relay spool: {error}"))?; + Ok(file) + } + + fn reqwest_body(&self, file: tokio::fs::File) -> reqwest::Body { + let spool = self.clone(); + reqwest::Body::wrap_stream(async_stream::stream! { + let _spool = spool; + let mut stream = ReaderStream::new(file); + while let Some(chunk) = futures_util::StreamExt::next(&mut stream).await { + yield chunk; + } + }) + } + + pub(crate) async fn body_stream( + &self, + ) -> Result> + Send + 'static, String> + { + let file = self.open_body().await?; + let spool = self.clone(); + Ok(async_stream::stream! { + let _spool = spool; + let mut stream = ReaderStream::new(file); + while let Some(chunk) = futures_util::StreamExt::next(&mut stream).await { + yield chunk; + } + }) + } + + fn payload_digest(&self) -> TunnelRelayPayloadDigest { + tunnel_relay_payload_digest_from_hashes( + &self.inner.metadata_envelope, + self.inner.body_len, + self.inner.body_sha256, + ) + } +} + +impl RelayAuthHeaders { + pub(crate) fn apply(self, request: reqwest::RequestBuilder) -> reqwest::RequestBuilder { + request + .header(TUNNEL_RELAY_AUTH_SENDER_HEADER, self.sender_instance_id) + .header(TUNNEL_RELAY_OWNER_INSTANCE_HEADER, self.owner_instance_id) + .header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, self.timestamp_unix_secs) + .header(TUNNEL_RELAY_AUTH_NONCE_HEADER, self.nonce) + .header(TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, self.payload_digest) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, self.signature) + } + + pub(crate) fn apply_to_headers(self, headers: &mut HeaderMap) -> Result<(), String> { + insert_relay_auth_header( + headers, + TUNNEL_RELAY_AUTH_SENDER_HEADER, + &self.sender_instance_id, + )?; + insert_relay_auth_header( + headers, + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + &self.owner_instance_id, + )?; + insert_relay_auth_header( + headers, + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + &self.timestamp_unix_secs.to_string(), + )?; + insert_relay_auth_header(headers, TUNNEL_RELAY_AUTH_NONCE_HEADER, &self.nonce)?; + insert_relay_auth_header( + headers, + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + &self.payload_digest, + )?; + insert_relay_auth_header(headers, TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, &self.signature) + } +} + +fn insert_relay_auth_header( + headers: &mut HeaderMap, + name: &'static str, + value: &str, +) -> Result<(), String> { + let value = HeaderValue::from_str(value) + .map_err(|error| format!("invalid tunnel relay authentication header: {error}"))?; + headers.insert(http::header::HeaderName::from_static(name), value); + Ok(()) +} + +pub(crate) fn build_relay_auth_headers_from_environment( + owner_instance_id: &str, + node_id: &str, + metadata_envelope: &[u8], + body: &[u8], +) -> Result { + let secret = resolve_tunnel_relay_auth_secret_from_environment()?; + build_relay_auth_headers( + secret.as_bytes(), + &resolve_tunnel_instance_id(), + owner_instance_id, + node_id, + false, + false, + metadata_envelope, + body, + ) +} + +fn build_relay_auth_headers( + secret: &[u8], + sender_instance_id: &str, + owner_instance_id: &str, + node_id: &str, + forwarded_by: bool, + rollout_probe: bool, + metadata_envelope: &[u8], + body: &[u8], +) -> Result { + build_relay_auth_headers_for_digest( + secret, + sender_instance_id, + owner_instance_id, + node_id, + forwarded_by, + rollout_probe, + tunnel_relay_payload_digest(metadata_envelope, body), + ) +} + +fn build_relay_auth_headers_for_digest( + secret: &[u8], + sender_instance_id: &str, + owner_instance_id: &str, + node_id: &str, + forwarded_by: bool, + rollout_probe: bool, + payload_digest: TunnelRelayPayloadDigest, +) -> Result { + validate_tunnel_relay_auth_secret(secret)?; + let timestamp_unix_secs = current_unix_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let forwarded_by_value = if forwarded_by { sender_instance_id } else { "" }; + let signature = sign_tunnel_relay_request( + secret, + sender_instance_id, + owner_instance_id, + node_id, + forwarded_by_value, + rollout_probe, + timestamp_unix_secs, + &nonce, + &payload_digest, + ); + Ok(RelayAuthHeaders { + sender_instance_id: sender_instance_id.to_string(), + owner_instance_id: owner_instance_id.to_string(), + timestamp_unix_secs, + nonce, + payload_digest: payload_digest.encode_header_value(), + signature, + }) +} + +#[derive(Serialize)] +struct TunnelAffinityAuthMetadata { + context: &'static str, + method: String, + path_and_query: String, + gateway_marker: Option, + affinity_forwarded_by: Option, + affinity_owner_instance_id: Option, + affinity_node_id: Option, + forwarded_host: Option, + forwarded_for: Option, + forwarded_proto: Option, + trusted_user_id: Option, + trusted_api_key_id: Option, + trusted_access_allowed: Option, + trusted_balance_remaining: Option, +} + +pub(crate) fn build_tunnel_affinity_auth_metadata( + method: &Method, + uri: &Uri, + headers: &HeaderMap, +) -> Result, String> { + let metadata = TunnelAffinityAuthMetadata { + context: TUNNEL_AFFINITY_AUTH_CONTEXT, + method: method.as_str().to_string(), + path_and_query: uri + .path_and_query() + .map(|value| value.as_str()) + .unwrap_or("/") + .to_string(), + gateway_marker: canonical_header_value(headers, crate::constants::GATEWAY_HEADER)?, + affinity_forwarded_by: canonical_header_value( + headers, + crate::constants::TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + )?, + affinity_owner_instance_id: canonical_header_value( + headers, + crate::constants::TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + )?, + affinity_node_id: canonical_header_value( + headers, + crate::constants::TUNNEL_AFFINITY_NODE_ID_HEADER, + )?, + forwarded_host: canonical_header_value(headers, crate::constants::FORWARDED_HOST_HEADER)?, + forwarded_for: canonical_header_value(headers, crate::constants::FORWARDED_FOR_HEADER)?, + forwarded_proto: canonical_header_value(headers, crate::constants::FORWARDED_PROTO_HEADER)?, + trusted_user_id: canonical_header_value( + headers, + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + )?, + trusted_api_key_id: canonical_header_value( + headers, + crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER, + )?, + trusted_access_allowed: canonical_header_value( + headers, + crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + )?, + trusted_balance_remaining: canonical_header_value( + headers, + crate::constants::TRUSTED_AUTH_BALANCE_HEADER, + )?, + }; + serde_json::to_vec(&metadata) + .map_err(|error| format!("encode tunnel affinity authentication metadata failed: {error}")) +} + +fn canonical_header_value( + headers: &HeaderMap, + name: &'static str, +) -> Result, String> { + let mut values = headers.get_all(name).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err(format!( + "duplicate tunnel affinity authentication header: {name}" + )); + } + value + .to_str() + .map(|value| Some(value.to_string())) + .map_err(|_| format!("invalid tunnel affinity authentication header: {name}")) +} + impl EmbeddedTunnelState { pub(crate) fn new() -> Self { Self::with_data(Arc::new(GatewayDataState::disabled())) @@ -447,6 +1219,9 @@ impl EmbeddedTunnelState { attachment_directory: TunnelAttachmentDirectory, ) -> Self { let defaults = EmbeddedTunnelDefaults::default(); + let relay_auth_secret = resolve_tunnel_relay_auth_secret() + .map(Arc::<[u8]>::from) + .map_err(Arc::from); Self { inner: TunnelAppState::new( build_embedded_control_plane(Arc::clone(&data), attachment_directory.clone()), @@ -457,15 +1232,265 @@ impl EmbeddedTunnelState { }, defaults.max_streams, ) + .with_relay_auth( + attachment_directory.local_instance_id().to_string(), + relay_auth_secret + .as_ref() + .ok() + .map(|value| value.as_ref().to_vec()), + Arc::clone(&attachment_directory.runtime_state), + ) .with_data(data), attachment_directory, + relay_auth_secret, } } + #[cfg(test)] + pub(crate) fn with_data_and_directory_for_tests( + data: Arc, + attachment_directory: TunnelAttachmentDirectory, + relay_auth_secret: &str, + ) -> Self { + let defaults = EmbeddedTunnelDefaults::default(); + let relay_auth_secret = validate_tunnel_relay_auth_secret(relay_auth_secret.as_bytes()) + .map(|()| Arc::<[u8]>::from(relay_auth_secret.as_bytes())) + .map_err(Arc::from); + Self { + inner: TunnelAppState::new( + build_embedded_control_plane(Arc::clone(&data), attachment_directory.clone()), + ConnConfig { + ping_interval: defaults.ping_interval, + idle_timeout: defaults.proxy_idle_timeout, + outbound_queue_capacity: defaults.outbound_queue_capacity, + }, + defaults.max_streams, + ) + .with_relay_auth( + attachment_directory.local_instance_id().to_string(), + relay_auth_secret + .as_ref() + .ok() + .map(|value| value.as_ref().to_vec()), + Arc::clone(&attachment_directory.runtime_state), + ) + .with_data(data), + attachment_directory, + relay_auth_secret, + } + } + + #[cfg(test)] + pub(crate) fn with_data_identity_runtime_state_and_relay_secret_for_tests( + data: Arc, + instance_id: &str, + relay_base_url: Option<&str>, + runtime_state: Arc, + relay_auth_secret: &str, + ) -> Self { + Self::with_data_and_directory_for_tests( + data, + TunnelAttachmentDirectory::for_tests(instance_id, relay_base_url, 90) + .with_runtime_state(runtime_state), + relay_auth_secret, + ) + } + pub(crate) fn app_state(&self) -> TunnelAppState { self.inner.clone() } + pub(crate) async fn authenticate_relay_request( + &self, + headers: &HeaderMap, + node_id: &str, + metadata_envelope: &[u8], + body: &[u8], + require_local_owner: bool, + ) -> Result<(), embedded::RelayAuthError> { + self.inner + .authenticate_relay_request( + headers, + node_id, + &tunnel_relay_payload_digest(metadata_envelope, body), + require_local_owner, + ) + .await + } + + pub(crate) async fn prepare_tunnel_affinity_auth_request( + &self, + method: &Method, + uri: &Uri, + headers: &HeaderMap, + ) -> Result, RelayAuthError> { + if !has_tunnel_affinity_auth_headers(headers) { + return Ok(None); + } + if headers.contains_key(TUNNEL_RELAY_ROLLOUT_PROBE_HEADER) { + return Err(RelayAuthError::Invalid); + } + if headers.contains_key("x-real-ip") { + return Err(RelayAuthError::Invalid); + } + + let gateway_marker = required_affinity_header( + headers, + crate::constants::GATEWAY_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + if gateway_marker != TUNNEL_AFFINITY_GATEWAY_MARKER { + return Err(RelayAuthError::Invalid); + } + + let relay_sender = required_affinity_header( + headers, + TUNNEL_RELAY_AUTH_SENDER_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let relay_owner = required_affinity_header( + headers, + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let relay_forwarded_by = required_affinity_header( + headers, + TUNNEL_RELAY_FORWARDED_BY_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let affinity_forwarded_by = required_affinity_header( + headers, + crate::constants::TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let affinity_owner = required_affinity_header( + headers, + crate::constants::TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let affinity_node_id = required_affinity_header( + headers, + crate::constants::TUNNEL_AFFINITY_NODE_ID_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + if relay_sender != relay_forwarded_by + || relay_sender != affinity_forwarded_by + || relay_owner != affinity_owner + || affinity_owner != self.local_instance_id() + { + return Err(RelayAuthError::Invalid); + } + + required_affinity_header( + headers, + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + required_affinity_header( + headers, + crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER, + MAX_TUNNEL_AFFINITY_ID_LEN, + )?; + let access_allowed = required_affinity_header( + headers, + crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + 5, + )?; + if !matches!(access_allowed, "true" | "false") { + return Err(RelayAuthError::Invalid); + } + if let Some(balance) = + optional_affinity_header(headers, crate::constants::TRUSTED_AUTH_BALANCE_HEADER, 64)? + { + if !balance + .parse::() + .ok() + .is_some_and(|value| value.is_finite()) + { + return Err(RelayAuthError::Invalid); + } + } + let client_ip = + required_affinity_header(headers, crate::constants::FORWARDED_FOR_HEADER, 45)? + .parse::() + .map_err(|_| RelayAuthError::Invalid)?; + + let metadata = build_tunnel_affinity_auth_metadata(method, uri, headers) + .map_err(|_| RelayAuthError::Invalid)?; + let relay_auth = self + .inner + .authenticate_relay_request_headers(headers, affinity_node_id, true) + .await?; + if !relay_auth.payload_digest.matches_metadata(&metadata) { + return Err(RelayAuthError::Invalid); + } + Ok(Some(PendingTunnelAffinityAuth { + context: TunnelAffinityAuthContext { client_ip }, + relay_auth, + })) + } + + pub(crate) async fn commit_tunnel_affinity_auth_request( + &self, + pending: PendingTunnelAffinityAuth, + body: &[u8], + ) -> Result { + if !pending.relay_auth.payload_digest.matches_body(body) { + return Err(RelayAuthError::Invalid); + } + self.inner.commit_relay_auth(&pending.relay_auth).await?; + Ok(pending.context) + } + + pub(crate) fn build_relay_auth_headers( + &self, + owner_instance_id: &str, + node_id: &str, + forwarded_by: bool, + rollout_probe: bool, + metadata_envelope: &[u8], + body: &[u8], + ) -> Result { + let secret = self + .relay_auth_secret + .as_deref() + .map_err(|error| error.to_string())?; + let sender_instance_id = self.local_instance_id(); + build_relay_auth_headers( + secret, + sender_instance_id, + owner_instance_id, + node_id, + forwarded_by, + rollout_probe, + metadata_envelope, + body, + ) + } + + pub(crate) fn build_relay_auth_headers_for_digest( + &self, + owner_instance_id: &str, + node_id: &str, + forwarded_by: bool, + rollout_probe: bool, + payload_digest: TunnelRelayPayloadDigest, + ) -> Result { + let secret = self + .relay_auth_secret + .as_deref() + .map_err(|error| error.to_string())?; + build_relay_auth_headers_for_digest( + secret, + self.local_instance_id(), + owner_instance_id, + node_id, + forwarded_by, + rollout_probe, + payload_digest, + ) + } + pub(crate) fn register_secure_tunnel_key( &self, node_id: impl Into, @@ -474,6 +1499,19 @@ impl EmbeddedTunnelState { self.inner.register_secure_tunnel_key(node_id, key); } + pub(crate) async fn authenticate_control_plane_request( + &self, + headers: &HeaderMap, + method: &str, + path: &str, + payload_node_id: &str, + body: &[u8], + ) -> Result { + self.inner + .authenticate_control_plane_request(headers, method, path, payload_node_id, body) + .await + } + pub(crate) fn has_local_proxy(&self, node_id: &str) -> bool { self.inner.hub.has_local_proxy(node_id) } @@ -491,6 +1529,10 @@ impl EmbeddedTunnelState { self.inner.hub.request_close_all_proxies() } + pub(crate) fn request_close_proxies_for_node(&self, node_id: &str) -> usize { + self.inner.hub.request_close_proxies_for_node(node_id) + } + pub(crate) fn stats(&self) -> TunnelStatsSnapshot { let stats = self.inner.hub.stats(); TunnelStatsSnapshot { @@ -540,26 +1582,33 @@ impl EmbeddedTunnelState { } let timeout_secs = timeout_secs.clamp(5, 60); - let owner_url = build_owner_relay_url(&owner.relay_base_url, node_id) - .map_err(|error| format!("invalid owner tunnel probe URL: {error:?}"))?; + let owner_url = build_tunnel_owner_relay_url(&owner.relay_base_url, node_id) + .map_err(|error| format!("invalid owner tunnel probe URL: {error}"))?; let payload = encode_tunnel_relay_envelope(&build_tunnel_probe_meta(url, timeout_secs))?; - let response = state - .owner_forward_client + let relay_auth = self.build_relay_auth_headers( + &owner.gateway_instance_id, + node_id, + true, + true, + &payload, + &[], + )?; + let owner_client = + owner_forward_client_for_url(&state.owner_forward_client, &owner_url).await?; + let request = owner_client .post(owner_url) .header(TUNNEL_RELAY_FORWARDED_BY_HEADER, self.local_instance_id()) - .header( - TUNNEL_RELAY_OWNER_INSTANCE_HEADER, - owner.gateway_instance_id.as_str(), - ) .header( TUNNEL_RELAY_ROLLOUT_PROBE_HEADER, TUNNEL_RELAY_ROLLOUT_PROBE_VALUE, ) .timeout(Duration::from_secs(timeout_secs)) - .body(payload) + .body(payload); + let response = relay_auth + .apply(request) .send() .await - .map_err(|error| format!("owner tunnel probe failed: {error}"))?; + .map_err(|error| owner_forward_request_error(&error))?; Ok(response.status().as_u16()) } @@ -571,7 +1620,10 @@ impl EmbeddedTunnelState { ) -> Result { let timeout_secs = timeout_secs.clamp(5, 60); let meta = build_tunnel_probe_meta(url, timeout_secs); - let stream = self.inner.hub.open_local_stream(node_id, &meta).await?; + let stream = self + .inner + .open_authorized_local_stream(node_id, &meta) + .await?; let stream_id = stream.id; let result = async { self.inner @@ -641,6 +1693,56 @@ impl EmbeddedTunnelState { } } +fn has_tunnel_affinity_auth_headers(headers: &HeaderMap) -> bool { + [ + TUNNEL_RELAY_AUTH_SENDER_HEADER, + TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + TUNNEL_RELAY_AUTH_NONCE_HEADER, + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + TUNNEL_RELAY_FORWARDED_BY_HEADER, + TUNNEL_RELAY_ROLLOUT_PROBE_HEADER, + crate::constants::TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + crate::constants::TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + crate::constants::TUNNEL_AFFINITY_NODE_ID_HEADER, + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER, + crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + crate::constants::TRUSTED_AUTH_BALANCE_HEADER, + ] + .into_iter() + .any(|name| headers.contains_key(name)) +} + +fn required_affinity_header<'a>( + headers: &'a HeaderMap, + name: &'static str, + max_len: usize, +) -> Result<&'a str, RelayAuthError> { + optional_affinity_header(headers, name, max_len)?.ok_or(RelayAuthError::Invalid) +} + +fn optional_affinity_header<'a>( + headers: &'a HeaderMap, + name: &'static str, + max_len: usize, +) -> Result, RelayAuthError> { + let mut values = headers.get_all(name).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err(RelayAuthError::Invalid); + } + let value = value.to_str().map_err(|_| RelayAuthError::Invalid)?; + let trimmed = value.trim(); + if trimmed.is_empty() || trimmed.len() > max_len || trimmed != value { + return Err(RelayAuthError::Invalid); + } + Ok(Some(trimmed)) +} + fn build_tunnel_probe_meta(url: &str, timeout_secs: u64) -> tunnel_protocol::RequestMeta { tunnel_protocol::RequestMeta { provider_id: None, @@ -698,9 +1800,10 @@ impl fmt::Debug for EmbeddedTunnelState { pub(crate) async fn proxy_tunnel( ws: WebSocketUpgrade, State(state): State, + connect_info: ConnectInfo, headers: HeaderMap, ) -> impl IntoResponse { - embedded::ws_proxy(ws, State(state.tunnel.app_state()), headers).await + embedded::ws_proxy(ws, State(state.tunnel.app_state()), connect_info, headers).await } pub(crate) async fn relay_request( @@ -710,20 +1813,120 @@ pub(crate) async fn relay_request( request: Request, ) -> Result, GatewayError> { let node_id = path.0; - if state.tunnel.has_local_proxy(&node_id) { - return Ok(embedded::relay_request( - Path(node_id), - State(state.tunnel.app_state()), - connect_info, - request, - ) + let trace_id = extract_or_generate_trace_id(request.headers()); + let (mut parts, body) = request.into_parts(); + let relay_body_limit = tunnel_relay_body_limit_bytes(); + match declared_relay_body_len(&parts.headers) { + Ok(Some(length)) if length > relay_body_limit => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::PAYLOAD_TOO_LARGE, + &format!("tunnel relay body exceeds {relay_body_limit} bytes"), + ); + } + Err(error) => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::BAD_REQUEST, + &error, + ); + } + _ => {} + } + let pending_auth = match state + .tunnel + .app_state() + .authenticate_relay_request_headers(&parts.headers, &node_id, true) .await - .into_response()); + { + Ok(pending) => pending, + Err(embedded::RelayAuthError::Unavailable) => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::SERVICE_UNAVAILABLE, + "tunnel relay authentication is not configured", + ); + } + Err(embedded::RelayAuthError::Invalid) => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::FORBIDDEN, + "invalid tunnel relay authentication", + ); + } + }; + let authenticated_body = match prepare_owner_relay_request_body_with_limits( + body, + relay_body_limit, + TUNNEL_RELAY_BODY_READ_TIMEOUT, + ) + .await + { + Ok(prepared) => prepared, + Err(error) => { + return build_local_http_error_response( + &trace_id, + None, + if error.starts_with("tunnel relay body exceeds") { + StatusCode::PAYLOAD_TOO_LARGE + } else { + StatusCode::BAD_REQUEST + }, + &error, + ); + } + }; + if pending_auth.payload_digest != authenticated_body.payload_digest() { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::FORBIDDEN, + "invalid tunnel relay payload integrity", + ); + } + match state + .tunnel + .app_state() + .commit_relay_auth(&pending_auth) + .await + { + Ok(()) => {} + Err(embedded::RelayAuthError::Unavailable) => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::SERVICE_UNAVAILABLE, + "tunnel relay authentication is not configured", + ); + } + Err(embedded::RelayAuthError::Invalid) => { + return build_local_http_error_response( + &trace_id, + None, + StatusCode::FORBIDDEN, + "invalid tunnel relay authentication", + ); + } + } + parts.extensions.insert(embedded::RelayRequestAuthenticated); + + if state.tunnel.has_local_proxy(&node_id) { + return relay_spool_to_local_proxy( + state.tunnel.app_state(), + connect_info, + node_id, + parts, + authenticated_body, + ) + .await; } - let trace_id = extract_or_generate_trace_id(request.headers()); - let already_forwarded = request - .headers() + let already_forwarded = parts + .headers .get(TUNNEL_RELAY_FORWARDED_BY_HEADER) .and_then(|value| value.to_str().ok()) .map(str::trim) @@ -745,8 +1948,15 @@ pub(crate) async fn relay_request( .map_err(GatewayError::Internal)? { if owner.gateway_instance_id != state.tunnel.local_instance_id() { - return forward_relay_request_to_owner(&state, &node_id, request, &trace_id, &owner) - .await; + return forward_relay_request_to_owner( + &state, + &node_id, + parts, + authenticated_body, + &trace_id, + &owner, + ) + .await; } state .tunnel @@ -755,14 +1965,30 @@ pub(crate) async fn relay_request( .map_err(GatewayError::Internal)?; } - Ok(embedded::relay_request( - Path(node_id), - State(state.tunnel.app_state()), + relay_spool_to_local_proxy( + state.tunnel.app_state(), connect_info, - request, + node_id, + parts, + authenticated_body, ) .await - .into_response()) +} + +async fn relay_spool_to_local_proxy( + state: TunnelAppState, + connect_info: ConnectInfo, + node_id: String, + parts: http::request::Parts, + spool: VerifiedRelaySpool, +) -> Result, GatewayError> { + let mut request = Request::from_parts(parts, Body::empty()); + request.extensions_mut().insert(spool); + Ok( + embedded::relay_request(Path(node_id), State(state), connect_info, request) + .await + .into_response(), + ) } fn build_embedded_control_plane( @@ -774,13 +2000,33 @@ fn build_embedded_control_plane( let node_status_data = Arc::clone(&data); let node_status_directory = attachment_directory; ControlPlaneClient::local( - move |payload| { + move |connection, payload| { let data = Arc::clone(&heartbeat_data); let directory = heartbeat_directory.clone(); Box::pin(async move { - let ack = apply_embedded_tunnel_heartbeat(data.as_ref(), &payload).await?; + crate::tunnel::embedded::validate_proxy_connection_credential( + data.as_ref(), + &connection, + ) + .await + .map_err(str::to_string)?; + let authenticated_node_id = connection.node_id.clone(); + let authenticated_generation = connection.node_generation.clone(); + let ack = apply_embedded_tunnel_heartbeat( + data.as_ref(), + directory.runtime_state.as_ref(), + &authenticated_node_id, + &authenticated_generation, + &payload, + ) + .await?; if let Err(error) = directory - .refresh_from_heartbeat(data.as_ref(), &payload) + .refresh_from_authenticated_heartbeat( + data.as_ref(), + &authenticated_node_id, + &authenticated_generation, + &payload, + ) .await { warn!(error = %error, "failed to refresh tunnel attachment from heartbeat"); @@ -788,13 +2034,22 @@ fn build_embedded_control_plane( Ok(ack) }) }, - move |node_id, connected, conn_count, observed_at_unix_secs| { + move |connection, connected, conn_count, observed_at_unix_secs| { let data = Arc::clone(&node_status_data); let directory = node_status_directory.clone(); Box::pin(async move { + crate::tunnel::embedded::validate_proxy_connection_credential( + data.as_ref(), + &connection, + ) + .await + .map_err(str::to_string)?; + let node_id = connection.node_id.clone(); + let node_generation = connection.node_generation.clone(); apply_embedded_tunnel_node_status( data.as_ref(), &node_id, + &node_generation, connected, conn_count, Some(observed_at_unix_secs), @@ -804,6 +2059,7 @@ fn build_embedded_control_plane( .sync_node_status( data.as_ref(), &node_id, + &node_generation, connected, conn_count, observed_at_unix_secs, @@ -821,41 +2077,57 @@ fn build_embedded_control_plane( async fn forward_relay_request_to_owner( state: &AppState, node_id: &str, - request: Request, + parts: http::request::Parts, + prepared_body: VerifiedRelaySpool, trace_id: &str, owner: &TunnelAttachmentRecord, ) -> Result, GatewayError> { - let owner_url = build_owner_relay_url(&owner.relay_base_url, node_id)?; - let (parts, body) = request.into_parts(); - let prepared_body = match prepare_owner_relay_request_body(body).await { - Ok(prepared_body) => prepared_body, - Err(error) => { - return build_local_http_error_response( - trace_id, - None, - StatusCode::BAD_REQUEST, - &error, - ); - } - }; + let owner_url = build_tunnel_owner_relay_url(&owner.relay_base_url, node_id) + .map_err(GatewayError::Internal)?; + let payload_digest = prepared_body.payload_digest(); + let meta = prepared_body.meta().clone(); + let file = prepared_body + .open_from_start() + .await + .map_err(GatewayError::Internal)?; + let request_body = prepared_body.reqwest_body(file); + let relay_auth = state + .tunnel + .build_relay_auth_headers_for_digest( + &owner.gateway_instance_id, + node_id, + true, + false, + payload_digest, + ) + .map_err(GatewayError::Internal)?; - let mut upstream_request = state.owner_forward_client.post(owner_url); + let connection_declared = aether_http::connection_declared_header_names( + parts + .headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); + let owner_client = owner_forward_client_for_url(&state.owner_forward_client, &owner_url) + .await + .map_err(GatewayError::Internal)?; + let mut upstream_request = owner_client.post(owner_url); for (name, value) in &parts.headers { - if should_skip_request_header(name.as_str()) || name == http::header::HOST { + if should_skip_request_header(name.as_str()) + || name == http::header::HOST + || is_tunnel_relay_auth_header(name.as_str()) + || connection_declared.contains(&name.as_str().to_ascii_lowercase()) + { continue; } upstream_request = upstream_request.header(name, value); } - upstream_request = upstream_request - .header( - TUNNEL_RELAY_FORWARDED_BY_HEADER, - state.tunnel.local_instance_id(), - ) - .header( - TUNNEL_RELAY_OWNER_INSTANCE_HEADER, - owner.gateway_instance_id.as_str(), - ); - let resolved_timeouts = resolve_tunnel_request_timeouts(&prepared_body.meta); + upstream_request = relay_auth.apply(upstream_request).header( + TUNNEL_RELAY_FORWARDED_BY_HEADER, + state.tunnel.local_instance_id(), + ); + let resolved_timeouts = resolve_tunnel_request_timeouts(&meta); if let Some(timeout_ms) = resolved_timeouts.response_body_ms { upstream_request = upstream_request.timeout(Duration::from_millis(timeout_ms)); } @@ -863,104 +2135,315 @@ async fn forward_relay_request_to_owner( upstream_request = upstream_request.header(TRACE_ID_HEADER, trace_id); } - let first_byte_timeout = prepared_body - .meta + let first_byte_timeout = meta .stream .then_some(Duration::from_millis(resolved_timeouts.first_byte_ms)); - let upstream_response = match send_owner_forward_request( - upstream_request.body(prepared_body.body), - first_byte_timeout, - ) - .await - { - Ok(response) => response, - Err(err) => { - return Err(GatewayError::Internal(format!( - "owner tunnel relay failed: {err}" - ))); - } - }; + let upstream_response = + match send_owner_forward_request(upstream_request.body(request_body), first_byte_timeout) + .await + { + Ok(response) => response, + Err(err) => { + return Err(GatewayError::Internal(format!( + "owner tunnel relay failed: {err}" + ))); + } + }; build_client_response(upstream_response, trace_id, None) } -fn build_owner_relay_url(relay_base_url: &str, node_id: &str) -> Result { - let mut url = url::Url::parse(relay_base_url) - .map_err(|err| GatewayError::Internal(format!("invalid owner relay base url: {err}")))?; +pub(crate) fn is_tunnel_relay_auth_header(name: &str) -> bool { + name.eq_ignore_ascii_case(TUNNEL_RELAY_AUTH_SENDER_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_AUTH_NONCE_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_AUTH_PAYLOAD_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_OWNER_INSTANCE_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_FORWARDED_BY_HEADER) + || name.eq_ignore_ascii_case(TUNNEL_RELAY_ROLLOUT_PROBE_HEADER) +} + +pub(crate) fn build_tunnel_owner_relay_url( + relay_base_url: &str, + node_id: &str, +) -> Result { + let node_id = node_id.trim(); + let node_id_lower = node_id.to_ascii_lowercase(); + if node_id.is_empty() + || matches!( + node_id_lower.as_str(), + "." | ".." | "%2e" | "%2e%2e" | ".%2e" | "%2e." + ) { - let mut segments = url.path_segments_mut().map_err(|_| { - GatewayError::Internal("owner relay base url cannot be a base-less URL".to_string()) - })?; + return Err("tunnel relay node ID is invalid".to_string()); + } + let mut url = parse_tunnel_relay_base_url(relay_base_url)?; + { + let mut segments = url + .path_segments_mut() + .map_err(|_| "tunnel relay base URL cannot be a base-less URL".to_string())?; segments.pop_if_empty(); segments.push("api"); segments.push("internal"); segments.push("tunnel"); segments.push("relay"); - segments.push(node_id.trim()); + segments.push(node_id); } Ok(url.to_string()) } -struct PreparedOwnerRelayRequestBody { - body: reqwest::Body, - meta: RequestMeta, +pub(crate) fn build_tunnel_affinity_forward_url( + relay_base_url: &str, + request_uri: &Uri, +) -> Result { + let mut url = parse_tunnel_relay_base_url(relay_base_url)?; + url.set_path(request_uri.path()); + url.set_query(request_uri.query()); + url.set_fragment(None); + validate_tunnel_relay_transport_url(&url)?; + Ok(url.to_string()) } -async fn prepare_owner_relay_request_body( - body: Body, -) -> Result { - let mut body_stream = body.into_data_stream(); - let mut buffered_chunks = Vec::new(); - let mut meta_buffer = BytesMut::new(); - let mut meta = None; +fn parse_tunnel_relay_base_url(value: &str) -> Result { + let url = url::Url::parse(value.trim()) + .map_err(|error| format!("invalid tunnel relay base URL: {error}"))?; + validate_tunnel_relay_transport_url(&url)?; + if url.query().is_some() || url.fragment().is_some() { + return Err("tunnel relay base URL must not include a query or fragment".to_string()); + } + Ok(url) +} - while meta.is_none() { - let Some(next_chunk) = body_stream.next().await else { - return Err("incomplete tunnel relay metadata".to_string()); - }; - match next_chunk { - Ok(chunk) => { - let meta_buffer_limit = 4usize.saturating_add(MAX_TUNNEL_RELAY_META_LEN); - let remaining = meta_buffer_limit.saturating_sub(meta_buffer.len()); - if remaining > 0 { - meta_buffer.extend_from_slice(&chunk[..chunk.len().min(remaining)]); - } - buffered_chunks.push(chunk); - match try_decode_tunnel_relay_request_meta(&meta_buffer) { - Ok(Some((parsed, _))) => meta = Some(parsed), - Ok(None) if meta_buffer.len() < meta_buffer_limit => {} - Ok(None) => return Err("incomplete tunnel relay metadata".to_string()), - Err(error) => return Err(error), +pub(crate) fn validate_tunnel_relay_transport_url(url: &url::Url) -> Result<(), String> { + if !url.username().is_empty() || url.password().is_some() { + return Err("tunnel relay URL must not include credentials".to_string()); + } + let host = url + .host() + .ok_or_else(|| "tunnel relay URL must include a host".to_string())?; + match url.scheme() { + "https" => Ok(()), + "http" => { + let loopback = match host { + url::Host::Domain(host) => { + host.trim_end_matches('.').eq_ignore_ascii_case("localhost") } + url::Host::Ipv4(address) => address.is_loopback(), + url::Host::Ipv6(address) => address.is_loopback(), + }; + if loopback { + Ok(()) + } else { + Err("tunnel relay URL must use HTTPS unless the host is loopback".to_string()) } - Err(error) => { - return Err(format!("tunnel relay body read failed: {error}")); - } + } + _ => Err("tunnel relay URL must use HTTPS or loopback HTTP".to_string()), + } +} + +fn tunnel_relay_body_limit_bytes() -> u64 { + std::env::var(TUNNEL_RELAY_MAX_BODY_MB_ENV) + .ok() + .and_then(|value| value.trim().parse::().ok()) + .filter(|value| *value > 0) + .map(|value| { + value + .saturating_mul(1024 * 1024) + .min(MAX_TUNNEL_RELAY_BODY_BYTES) + }) + .unwrap_or(DEFAULT_TUNNEL_RELAY_MAX_BODY_BYTES) +} + +fn tunnel_relay_spool_budget_bytes() -> u64 { + std::env::var(TUNNEL_RELAY_SPOOL_BUDGET_MB_ENV) + .ok() + .and_then(|value| value.trim().parse::().ok()) + .filter(|value| *value > 0) + .map(|value| { + value + .saturating_mul(1024 * 1024) + .min(MAX_TUNNEL_RELAY_SPOOL_BUDGET_BYTES) + }) + .unwrap_or(DEFAULT_TUNNEL_RELAY_SPOOL_BUDGET_BYTES) +} + +fn reserve_tunnel_relay_spool_bytes(bytes: u64) -> Result<(), String> { + let budget = tunnel_relay_spool_budget_bytes(); + loop { + let current = TUNNEL_RELAY_SPOOL_BYTES_IN_USE.load(Ordering::Acquire); + let next = current + .checked_add(bytes) + .ok_or_else(|| "tunnel relay spool budget overflow".to_string())?; + if next > budget { + return Err(format!( + "tunnel relay spool budget exhausted at {budget} bytes" + )); + } + if TUNNEL_RELAY_SPOOL_BYTES_IN_USE + .compare_exchange_weak(current, next, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + return Ok(()); } } - let meta = meta.ok_or_else(|| "incomplete tunnel relay metadata".to_string())?; +} - let forwarded_body = reqwest::Body::wrap_stream(stream! { - for chunk in buffered_chunks { - yield Ok::(chunk); - } - while let Some(next_chunk) = body_stream.next().await { - match next_chunk { - Ok(chunk) => { - yield Ok::(chunk); - } - Err(error) => { - yield Err::(io::Error::other(error)); - break; +fn declared_relay_body_len(headers: &HeaderMap) -> Result, String> { + let mut values = headers.get_all(http::header::CONTENT_LENGTH).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err("duplicate content-length header on tunnel relay request".to_string()); + } + let value = value + .to_str() + .map_err(|_| "invalid content-length header on tunnel relay request".to_string())?; + value + .trim() + .parse::() + .map(Some) + .map_err(|_| "invalid content-length header on tunnel relay request".to_string()) +} + +async fn prepare_owner_relay_request_body(body: Body) -> Result { + prepare_owner_relay_request_body_with_limits( + body, + tunnel_relay_body_limit_bytes(), + TUNNEL_RELAY_BODY_READ_TIMEOUT, + ) + .await +} + +async fn prepare_owner_relay_request_body_with_limits( + body: Body, + max_envelope_bytes: u64, + read_timeout: Duration, +) -> Result { + let mut reserved_bytes = 0_u64; + let path = std::env::temp_dir().join(format!( + "aether-tunnel-relay-{}-{}", + std::process::id(), + uuid::Uuid::new_v4().simple() + )); + let mut options = tokio::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + options.mode(0o600); + } + let mut file = options + .open(&path) + .await + .map_err(|error| format!("failed to create tunnel relay spool: {error}"))?; + let mut cleanup = Some(path.clone()); + let result = async { + let mut body_stream = body.into_data_stream(); + let mut metadata_prefix = bytes::BytesMut::new(); + let mut decoded = None; + let mut body_hasher = Sha256::new(); + let mut body_len = 0_u64; + let mut envelope_len = 0_u64; + + loop { + let next_chunk = tokio::time::timeout( + read_timeout, + futures_util::StreamExt::next(&mut body_stream), + ) + .await + .map_err(|_| { + format!( + "tunnel relay body read timed out after {} seconds", + read_timeout.as_secs() + ) + })?; + let Some(chunk) = next_chunk else { + break; + }; + let chunk = chunk.map_err(|error| format!("tunnel relay body read failed: {error}"))?; + let chunk_start = envelope_len; + envelope_len = envelope_len + .checked_add(chunk.len() as u64) + .ok_or_else(|| "tunnel relay body length overflow".to_string())?; + if envelope_len > max_envelope_bytes { + return Err(format!( + "tunnel relay body exceeds {max_envelope_bytes} bytes" + )); + } + reserve_tunnel_relay_spool_bytes(chunk.len() as u64)?; + reserved_bytes = reserved_bytes.saturating_add(chunk.len() as u64); + tokio::time::timeout(read_timeout, file.write_all(&chunk)) + .await + .map_err(|_| { + format!( + "tunnel relay spool write timed out after {} seconds", + read_timeout.as_secs() + ) + })? + .map_err(|error| format!("failed to write tunnel relay spool: {error}"))?; + + if decoded.is_some() { + body_hasher.update(&chunk); + body_len = body_len.saturating_add(chunk.len() as u64); + continue; + } + + let prefix_limit = 4usize.saturating_add(MAX_TUNNEL_RELAY_META_LEN); + let remaining = prefix_limit.saturating_sub(metadata_prefix.len()); + metadata_prefix.extend_from_slice(&chunk[..chunk.len().min(remaining)]); + match try_decode_tunnel_relay_request_meta(&metadata_prefix)? { + Some((meta, body_offset)) => { + let body_offset_u64 = body_offset as u64; + if envelope_len > body_offset_u64 { + let body_in_chunk = body_offset_u64 + .saturating_sub(chunk_start) + .min(chunk.len() as u64) + as usize; + let first_body = &chunk[body_in_chunk..]; + body_hasher.update(first_body); + body_len = first_body.len() as u64; + } + decoded = Some((meta, body_offset)); } + None if metadata_prefix.len() < prefix_limit => {} + None => return Err("incomplete tunnel relay metadata".to_string()), } } - }); - - Ok(PreparedOwnerRelayRequestBody { - body: forwarded_body, - meta, - }) + tokio::time::timeout(read_timeout, file.flush()) + .await + .map_err(|_| { + format!( + "tunnel relay spool flush timed out after {} seconds", + read_timeout.as_secs() + ) + })? + .map_err(|error| format!("failed to flush tunnel relay spool: {error}"))?; + let (meta, body_offset) = + decoded.ok_or_else(|| "incomplete tunnel relay metadata".to_string())?; + let metadata_envelope = metadata_prefix.freeze().slice(..body_offset); + Ok(VerifiedRelaySpool { + inner: Arc::new(RelaySpoolInner { + path, + meta, + metadata_envelope, + body_offset: body_offset as u64, + body_len, + body_sha256: body_hasher.finalize().into(), + reserved_bytes, + }), + }) + } + .await; + if result.is_ok() { + cleanup = None; + } else if reserved_bytes > 0 { + TUNNEL_RELAY_SPOOL_BYTES_IN_USE.fetch_sub(reserved_bytes, Ordering::AcqRel); + } + if let Some(path) = cleanup { + let _ = tokio::fs::remove_file(path).await; + } + result } fn tunnel_attachment_key(node_id: &str) -> String { @@ -971,7 +2454,7 @@ fn tunnel_attachment_redis_key(node_id: &str) -> String { format!("{TUNNEL_ATTACHMENT_REDIS_KEY_PREFIX}{}", node_id.trim()) } -fn resolve_tunnel_instance_id() -> String { +pub(crate) fn resolve_tunnel_instance_id() -> String { std::env::var(TUNNEL_INSTANCE_ID_ENV) .ok() .map(|value| value.trim().to_string()) @@ -994,6 +2477,40 @@ fn normalize_relay_base_url(value: &str) -> Option { } } +fn resolve_tunnel_relay_auth_secret() -> Result, String> { + resolve_tunnel_relay_auth_secret_from_environment().map(String::into_bytes) +} + +fn resolve_tunnel_relay_auth_secret_from_environment() -> Result { + let value = std::env::var(TUNNEL_RELAY_AUTH_SECRET_ENV).map_err(|error| match error { + std::env::VarError::NotPresent => missing_tunnel_relay_auth_secret_error(), + std::env::VarError::NotUnicode(_) => { + format!("{TUNNEL_RELAY_AUTH_SECRET_ENV} must be valid UTF-8") + } + })?; + resolve_tunnel_relay_auth_secret_value(Some(&value)) +} + +fn resolve_tunnel_relay_auth_secret_value(value: Option<&str>) -> Result { + let value = value.ok_or_else(missing_tunnel_relay_auth_secret_error)?; + let value = value.trim(); + validate_tunnel_relay_auth_secret(value.as_bytes())?; + Ok(value.to_string()) +} + +fn missing_tunnel_relay_auth_secret_error() -> String { + format!("{TUNNEL_RELAY_AUTH_SECRET_ENV} is required for tunnel relay authentication") +} + +fn validate_tunnel_relay_auth_secret(secret: &[u8]) -> Result<(), String> { + if secret.len() < TUNNEL_RELAY_AUTH_SECRET_MIN_BYTES { + return Err(format!( + "{TUNNEL_RELAY_AUTH_SECRET_ENV} must contain at least {TUNNEL_RELAY_AUTH_SECRET_MIN_BYTES} bytes" + )); + } + Ok(()) +} + fn current_unix_secs() -> u64 { SystemTime::now() .duration_since(SystemTime::UNIX_EPOCH) @@ -1003,12 +2520,39 @@ fn current_unix_secs() -> u64 { async fn apply_embedded_tunnel_heartbeat( data: &GatewayDataState, + runtime_state: &RuntimeState, + authenticated_node_id: &str, + authenticated_generation: &str, request_body: &[u8], ) -> Result, String> { let payload = parse_embedded_tunnel_heartbeat_request(request_body)?; let node_id = payload.node_id.trim().to_string(); + if node_id != authenticated_node_id { + return Err("heartbeat node_id does not match authenticated tunnel node".to_string()); + } + let claim = claim_tunnel_heartbeat( + runtime_state, + &node_id, + &payload.heartbeat_session_id, + payload.heartbeat_id, + ) + .await?; + if claim.is_none() { + let node = data + .find_proxy_node(&node_id) + .await + .map_err(|err| format!("heartbeat duplicate lookup failed: {err}"))? + .filter(|node| node.tunnel_generation == authenticated_generation) + .ok_or_else(|| "proxy tunnel credential was revoked".to_string())?; + return Ok(build_embedded_tunnel_heartbeat_ack( + &node, + payload.heartbeat_id, + )); + } + let claim = claim.expect("fresh heartbeat claim should be present"); let mutation = ProxyNodeHeartbeatMutation { node_id: node_id.clone(), + expected_tunnel_generation: Some(authenticated_generation.to_string()), heartbeat_interval: payload.heartbeat_interval, active_connections: payload.active_connections, total_requests_delta: payload.window_total_requests.or(payload.total_requests), @@ -1020,11 +2564,26 @@ async fn apply_embedded_tunnel_heartbeat( proxy_version: payload.proxy_version, }; - let node = data + crate::state::decrypt_or_migrate_proxy_tunnel_psk(data, &node_id) + .await + .map_err(|err| format!("heartbeat security migration failed: {err}"))?; + let node_result = data .apply_proxy_node_heartbeat(&mutation) .await - .map_err(|err| format!("heartbeat sync failed: {err}"))? - .ok_or_else(|| format!("heartbeat sync failed: ProxyNode {node_id} 不存在"))?; + .map_err(|err| format!("heartbeat sync failed: {err}")) + .and_then(|node| { + node.ok_or_else(|| format!("heartbeat sync failed: ProxyNode {node_id} 不存在")) + }); + let node = match node_result { + Ok(node) => { + finish_tunnel_heartbeat_claim(runtime_state, claim).await; + node + } + Err(error) => { + finish_tunnel_heartbeat_claim(runtime_state, claim).await; + return Err(error); + } + }; Ok(build_embedded_tunnel_heartbeat_ack( &node, @@ -1032,15 +2591,107 @@ async fn apply_embedded_tunnel_heartbeat( )) } +pub(crate) async fn claim_tunnel_heartbeat( + runtime_state: &RuntimeState, + node_id: &str, + heartbeat_session_id: &str, + heartbeat_id: u64, +) -> Result, String> { + let state_key = tunnel_heartbeat_state_key(node_id, heartbeat_session_id)?; + let lock_key = format!("{state_key}:lock"); + let owner = format!("tunnel-heartbeat:{node_id}"); + let Some(lease) = runtime_state + .lock_try_acquire(&lock_key, &owner, TUNNEL_HEARTBEAT_LOCK_TTL) + .await + .map_err(|error| format!("heartbeat replay lock failed: {error}"))? + else { + return Err("heartbeat with this session is already being processed".to_string()); + }; + + let previous_heartbeat_id = match runtime_state.kv_get(&state_key).await { + Ok(Some(value)) => match value.parse::() { + Ok(value) => Some(value), + Err(_) => { + let _ = runtime_state.lock_release(&lease).await; + return Err("invalid heartbeat replay state".to_string()); + } + }, + Ok(None) => None, + Err(error) => { + let _ = runtime_state.lock_release(&lease).await; + return Err(format!("heartbeat replay state read failed: {error}")); + } + }; + + if previous_heartbeat_id.is_some_and(|previous| heartbeat_id <= previous) { + let _ = runtime_state.lock_release(&lease).await; + return Ok(None); + } + + if let Err(error) = runtime_state + .kv_set( + &state_key, + heartbeat_id.to_string(), + Some(TUNNEL_HEARTBEAT_STATE_TTL), + ) + .await + { + let _ = runtime_state.lock_release(&lease).await; + return Err(format!("heartbeat replay state write failed: {error}")); + } + + Ok(Some(TunnelHeartbeatClaim { lease })) +} + +pub(crate) async fn finish_tunnel_heartbeat_claim( + runtime_state: &RuntimeState, + claim: TunnelHeartbeatClaim, +) { + if let Err(error) = runtime_state.lock_release(&claim.lease).await { + warn!(error = %error, "failed to release heartbeat replay lock after commit"); + } +} + +fn tunnel_heartbeat_state_key(node_id: &str, heartbeat_session_id: &str) -> Result { + validate_tunnel_heartbeat_session_id(heartbeat_session_id)?; + let heartbeat_session_id = heartbeat_session_id.trim(); + + let mut digest = Sha256::new(); + digest.update(node_id.as_bytes()); + digest.update([0]); + digest.update(heartbeat_session_id.as_bytes()); + Ok(format!( + "{TUNNEL_HEARTBEAT_STATE_KEY_PREFIX}{:x}", + digest.finalize() + )) +} + +pub(crate) fn validate_tunnel_heartbeat_session_id( + heartbeat_session_id: &str, +) -> Result<(), String> { + let heartbeat_session_id = heartbeat_session_id.trim(); + if heartbeat_session_id.is_empty() + || heartbeat_session_id.len() > 128 + || !heartbeat_session_id + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b':')) + { + return Err("invalid heartbeat payload".to_string()); + } + Ok(()) +} + async fn apply_embedded_tunnel_node_status( data: &GatewayDataState, node_id: &str, + node_generation: &str, connected: bool, conn_count: usize, observed_at_unix_secs: Option, ) -> Result<(), String> { let mutation = ProxyNodeTunnelStatusMutation { node_id: node_id.trim().to_string(), + expected_tunnel_generation: Some(node_generation.to_string()), connected, conn_count: conn_count.min(i32::MAX as usize) as i32, detail: None, @@ -1049,8 +2700,9 @@ async fn apply_embedded_tunnel_node_status( data.update_proxy_node_tunnel_status(&mutation) .await - .map(|_| ()) - .map_err(|err| format!("node status sync failed: {err}")) + .map_err(|err| format!("node status sync failed: {err}"))? + .ok_or_else(|| "proxy tunnel credential was revoked".to_string())?; + Ok(()) } fn build_embedded_tunnel_heartbeat_ack(node: &StoredProxyNode, heartbeat_id: u64) -> Vec { @@ -1080,7 +2732,11 @@ fn parse_embedded_tunnel_heartbeat_request( .map_err(|_| "invalid heartbeat payload".to_string())?; let node_id = payload.node_id.trim(); - if node_id.is_empty() || node_id.len() > 36 || payload.heartbeat_id == 0 { + if node_id.is_empty() + || node_id.len() > 36 + || payload.heartbeat_id == 0 + || tunnel_heartbeat_state_key(node_id, &payload.heartbeat_session_id).is_err() + { return Err("invalid heartbeat payload".to_string()); } if payload @@ -1117,9 +2773,14 @@ fn parse_embedded_tunnel_heartbeat_request( mod tests { use super::{ apply_embedded_tunnel_heartbeat, apply_embedded_tunnel_node_status, - build_tunnel_probe_meta, current_unix_secs, encode_tunnel_relay_envelope, - prepare_owner_relay_request_body, tunnel_attachment_key, AppState, GatewayDataState, - TunnelAttachmentDirectory, TunnelAttachmentRecord, + build_owner_forward_target, build_owner_forward_target_with_policy, + build_tunnel_affinity_auth_metadata, build_tunnel_owner_relay_url, build_tunnel_probe_meta, + current_unix_secs, encode_tunnel_relay_envelope, owner_forward_client_for_url, + prepare_owner_relay_request_body, prepare_owner_relay_request_body_with_limits, + resolve_owner_forward_target, resolve_tunnel_relay_auth_secret_value, + tunnel_attachment_key, tunnel_relay_private_host_matches_allowlist, + validate_tunnel_relay_transport_url, AppState, GatewayDataState, TunnelAttachmentDirectory, + TunnelAttachmentRecord, MAX_OWNER_FORWARD_DNS_ADDRESSES, }; use aether_contracts::tunnel::{ try_decode_tunnel_relay_request_meta, TUNNEL_RELAY_FORWARDED_BY_HEADER, @@ -1128,12 +2789,325 @@ mod tests { use aether_data::repository::proxy_nodes::{ InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode, }; + use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; use axum::body::{Body, Bytes}; - use axum::http::{HeaderMap, StatusCode}; - use axum::routing::post; + use axum::http::{HeaderMap, HeaderValue, Method, StatusCode, Uri}; + use axum::routing::{any, post}; use axum::Router; use serde_json::json; + use std::io; + use std::net::SocketAddr; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; + use std::time::Duration; + + const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes"; + const TEST_TUNNEL_GENERATION: &str = "test-generation-1"; + + #[test] + fn tunnel_relay_secret_requires_an_independent_32_byte_value() { + let missing = resolve_tunnel_relay_auth_secret_value(None) + .expect_err("missing relay secret should fail"); + assert!(missing.contains("AETHER_TUNNEL_RELAY_AUTH_SECRET")); + + let short = resolve_tunnel_relay_auth_secret_value(Some(&"x".repeat(31))) + .expect_err("31-byte relay secret should fail"); + assert!(short.contains("at least 32 bytes")); + + assert_eq!( + resolve_tunnel_relay_auth_secret_value(Some(&"x".repeat(32))) + .expect("32-byte relay secret should pass"), + "x".repeat(32) + ); + } + + #[test] + fn tunnel_relay_url_policy_requires_https_except_for_loopback_http() { + for accepted in [ + "https://gateway.example.com/base", + "http://localhost:8084", + "http://localhost.:8084", + "http://127.0.0.1:8084", + "http://[::1]:8084", + ] { + let url = url::Url::parse(accepted).expect("accepted URL should parse"); + validate_tunnel_relay_transport_url(&url) + .unwrap_or_else(|error| panic!("{accepted} should be accepted: {error}")); + } + + for rejected in [ + "http://gateway.example.com", + "http://10.0.0.2:8084", + "ftp://gateway.example.com", + "https://user:password@gateway.example.com", + "file:///tmp/relay", + ] { + let url = url::Url::parse(rejected).expect("rejected URL should still parse"); + assert!( + validate_tunnel_relay_transport_url(&url).is_err(), + "{rejected} should be rejected" + ); + } + } + + #[test] + fn tunnel_owner_relay_url_uses_the_validated_base_and_escapes_node_id() { + assert_eq!( + build_tunnel_owner_relay_url("https://gateway.example.com/base", "node/a") + .expect("HTTPS relay URL should build"), + "https://gateway.example.com/base/api/internal/tunnel/relay/node%2Fa" + ); + assert!(build_tunnel_owner_relay_url("http://gateway.example.com", "node-1").is_err()); + for invalid_node_id in [ + "", " ", ".", "..", "%2e", "%2E", "%2e%2e", "%2E%2e", ".%2E", "%2e.", + ] { + assert!( + build_tunnel_owner_relay_url("https://gateway.example.com/base", invalid_node_id) + .is_err(), + "reserved or empty node ID {invalid_node_id:?} must not collapse the relay path" + ); + } + } + + #[test] + fn owner_forward_target_deduplicates_addresses_and_preserves_ports() { + let target = build_owner_forward_target( + "https", + "gateway.example.com", + 8443, + vec![ + "[2001:4860:4860::8888]:8443".parse().unwrap(), + "8.8.8.8:8443".parse().unwrap(), + "8.8.8.8:8443".parse().unwrap(), + ], + ) + .expect("valid DNS answers should produce a target"); + assert_eq!(target.host, "gateway.example.com"); + assert_eq!(target.port, 8443); + assert_eq!(target.addresses.len(), 2); + assert!(!target.literal_host); + } + + #[test] + fn owner_forward_target_rejects_empty_or_mismatched_dns_answers() { + assert!(build_owner_forward_target("https", "gateway.example.com", 443, vec![]).is_err()); + assert!(build_owner_forward_target( + "https", + "gateway.example.com", + 443, + vec!["192.0.2.10:8443".parse().unwrap()], + ) + .is_err()); + } + + #[test] + fn owner_forward_target_rejects_an_oversized_dns_answer_set() { + let addresses = (1..=MAX_OWNER_FORWARD_DNS_ADDRESSES + 1) + .map(|octet| SocketAddr::from(([192, 0, 2, octet as u8], 443))) + .collect(); + assert!( + build_owner_forward_target("https", "gateway.example.com", 443, addresses).is_err() + ); + } + + #[test] + fn owner_forward_target_rejects_non_loopback_http_dns_answers() { + assert!(build_owner_forward_target( + "http", + "localhost", + 8084, + vec!["192.0.2.10:8084".parse().unwrap()], + ) + .is_err()); + assert!(build_owner_forward_target( + "http", + "localhost", + 8084, + vec![ + "127.0.0.1:8084".parse().unwrap(), + "[::1]:8084".parse().unwrap(), + ], + ) + .is_ok()); + + // Private HTTPS owner deployments require an explicit deployment + // opt-in; the default path must not be an SSRF primitive. + assert!(build_owner_forward_target( + "https", + "gateway.internal", + 8443, + vec!["10.0.0.8:8443".parse().unwrap()], + ) + .is_err()); + assert!(build_owner_forward_target_with_policy( + "https", + "gateway.internal", + 8443, + vec!["10.0.0.8:8443".parse().unwrap()], + true, + ) + .is_ok()); + } + + #[test] + fn private_relay_host_allowlist_matches_exact_dns_names_only() { + assert!(tunnel_relay_private_host_matches_allowlist( + "gateway-a.internal.", + "gateway-a.internal, gateway-b.internal" + )); + assert!(tunnel_relay_private_host_matches_allowlist( + "GATEWAY-B.INTERNAL", + "gateway-a.internal, gateway-b.internal." + )); + assert!(!tunnel_relay_private_host_matches_allowlist( + "api.gateway-a.internal", + "gateway-a.internal" + )); + assert!(!tunnel_relay_private_host_matches_allowlist( + "gateway-a.internal.evil.example", + ".internal" + )); + } + + #[test] + fn owner_forward_target_marks_literal_ipv4_and_ipv6_hosts() { + let ipv4 = build_owner_forward_target( + "http", + "127.0.0.1", + 8084, + vec!["127.0.0.1:8084".parse().unwrap()], + ) + .expect("IPv4 literal should be valid"); + let ipv6 = + build_owner_forward_target("http", "::1", 8084, vec!["[::1]:8084".parse().unwrap()]) + .expect("IPv6 literal should be valid"); + assert!(ipv4.literal_host); + assert!(ipv6.literal_host); + } + + #[tokio::test] + async fn owner_forward_resolution_recognizes_bracketed_ipv6_literal_url() { + let target = resolve_owner_forward_target("http://[::1]:8084/api/internal/tunnel/relay/n") + .await + .expect("loopback IPv6 owner URL should resolve without DNS"); + assert_eq!(target.host, "::1"); + assert_eq!(target.addresses, vec!["[::1]:8084".parse().unwrap()]); + assert!(target.literal_host); + } + + #[tokio::test] + async fn owner_forward_domain_client_uses_pinned_transport_instead_of_shared_proxy() { + let owner_app = Router::new().route("/relay", post(|| async { StatusCode::NO_CONTENT })); + let owner_listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("owner listener should bind"); + let owner_port = owner_listener + .local_addr() + .expect("owner address should be available") + .port(); + let owner_server = tokio::spawn(async move { + axum::serve(owner_listener, owner_app) + .await + .expect("owner server should run"); + }); + + let proxy_hits = Arc::new(AtomicUsize::new(0)); + let proxy_hits_for_route = Arc::clone(&proxy_hits); + let proxy_app = Router::new().fallback(any(move || { + let proxy_hits = Arc::clone(&proxy_hits_for_route); + async move { + proxy_hits.fetch_add(1, Ordering::Relaxed); + StatusCode::BAD_GATEWAY + } + })); + let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("proxy listener should bind"); + let proxy_url = format!( + "http://{}", + proxy_listener + .local_addr() + .expect("proxy address should be available") + ); + let proxy_server = tokio::spawn(async move { + axum::serve(proxy_listener, proxy_app) + .await + .expect("proxy server should run"); + }); + + let shared_client = reqwest::Client::builder() + .proxy(reqwest::Proxy::all(&proxy_url).expect("proxy URL should build")) + .build() + .expect("shared client should build"); + let owner_url = format!("http://localhost:{owner_port}/relay"); + let pinned_client = owner_forward_client_for_url(&shared_client, &owner_url) + .await + .expect("localhost owner URL should resolve and build a pinned client"); + let response = pinned_client + .post(&owner_url) + .send() + .await + .expect("pinned owner request should bypass the shared proxy"); + + assert_eq!(response.status(), StatusCode::NO_CONTENT); + assert_eq!(proxy_hits.load(Ordering::Relaxed), 0); + + owner_server.abort(); + proxy_server.abort(); + let _ = owner_server.await; + let _ = proxy_server.await; + } + + #[test] + fn tunnel_affinity_auth_metadata_is_deterministic_and_binds_identity_and_path() { + let method = Method::POST; + let uri: Uri = "/v1/chat/completions?stream=false" + .parse() + .expect("URI should parse"); + let mut headers = HeaderMap::new(); + headers.insert( + crate::constants::GATEWAY_HEADER, + HeaderValue::from_static("rust-phase3b-affinity"), + ); + headers.insert( + crate::constants::TUNNEL_AFFINITY_FORWARDED_BY_HEADER, + HeaderValue::from_static("gateway-a"), + ); + headers.insert( + crate::constants::TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER, + HeaderValue::from_static("gateway-b"), + ); + headers.insert( + crate::constants::TUNNEL_AFFINITY_NODE_ID_HEADER, + HeaderValue::from_static("node-1"), + ); + headers.insert( + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_static("user-1"), + ); + + let first = build_tunnel_affinity_auth_metadata(&method, &uri, &headers) + .expect("metadata should build"); + let second = build_tunnel_affinity_auth_metadata(&method, &uri, &headers) + .expect("metadata should be deterministic"); + assert_eq!(first, second); + + let other_uri: Uri = "/v1/responses".parse().expect("URI should parse"); + assert_ne!( + first, + build_tunnel_affinity_auth_metadata(&method, &other_uri, &headers) + .expect("metadata should build") + ); + headers.insert( + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + HeaderValue::from_static("user-2"), + ); + assert_ne!( + first, + build_tunnel_affinity_auth_metadata(&method, &uri, &headers) + .expect("metadata should build") + ); + } fn sample_proxy_node(node_id: &str) -> StoredProxyNode { StoredProxyNode::new( @@ -1170,6 +3144,7 @@ mod tests { Some(1_700_000_000), Some(1_700_000_001), ) + .with_tunnel_generation(TEST_TUNNEL_GENERATION.to_string()) } #[test] @@ -1213,17 +3188,26 @@ mod tests { let owner = TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: owner_base_url, + tunnel_generation: TEST_TUNNEL_GENERATION.to_string(), conn_count: 1, observed_at_unix_secs: current_unix_secs(), }; - let data = GatewayDataState::disabled().with_system_config_values_for_tests(vec![( - tunnel_attachment_key("node-remote"), - serde_json::to_value(owner).expect("owner record should serialize"), - )]); + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( + "node-remote", + )])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_system_config_values_for_tests(vec![( + tunnel_attachment_key("node-remote"), + serde_json::to_value(owner).expect("owner record should serialize"), + )]); let state = AppState::new() .expect("app state should build") .with_data_state_for_tests(data) - .with_tunnel_identity("gateway-a", Some("http://gateway-a.internal")); + .with_tunnel_identity_and_relay_secret_for_tests( + "gateway-a", + Some("https://gateway-a.internal"), + RELAY_TEST_SECRET, + ); let status = state .tunnel @@ -1279,6 +3263,65 @@ mod tests { assert!(error.contains("invalid relay metadata")); } + #[tokio::test] + async fn owner_relay_body_preparation_hashes_a_large_body_in_one_chunk() { + let meta = build_tunnel_probe_meta("https://probe.example/health", 7); + let body = vec![b'x'; aether_contracts::tunnel::MAX_TUNNEL_RELAY_META_LEN + 1024]; + let envelope = { + let metadata = serde_json::to_vec(&meta).expect("metadata should encode"); + let mut envelope = Vec::with_capacity(4 + metadata.len() + body.len()); + envelope.extend_from_slice(&(metadata.len() as u32).to_be_bytes()); + envelope.extend_from_slice(&metadata); + envelope.extend_from_slice(&body); + envelope + }; + + let spool = prepare_owner_relay_request_body(Body::from(envelope)) + .await + .expect("single-chunk relay body should prepare"); + assert_eq!(spool.payload_digest().body_len(), body.len() as u64); + assert!(spool.payload_digest().matches_body(&body)); + } + + #[tokio::test] + async fn owner_relay_body_preparation_rejects_envelope_above_hard_limit() { + let meta = build_tunnel_probe_meta("https://probe.example/health", 7); + let envelope = encode_tunnel_relay_envelope(&meta).expect("metadata should encode"); + let limit = envelope.len() as u64; + let oversized = envelope + .into_iter() + .chain(std::iter::once(b'x')) + .collect::>(); + let error = prepare_owner_relay_request_body_with_limits( + Body::from(oversized), + limit, + Duration::from_secs(1), + ) + .await + .err() + .expect("oversized relay envelope should fail"); + assert!(error.contains("tunnel relay body exceeds")); + } + + #[tokio::test] + async fn owner_relay_body_preparation_times_out_slow_body_streams() { + let meta = build_tunnel_probe_meta("https://probe.example/health", 7); + let envelope = encode_tunnel_relay_envelope(&meta).expect("metadata should encode"); + let body = Body::from_stream(async_stream::stream! { + tokio::time::sleep(Duration::from_millis(50)).await; + yield Ok::(Bytes::from(envelope)); + }); + let error = prepare_owner_relay_request_body_with_limits( + body, + 1024 * 1024, + Duration::from_millis(10), + ) + .await + .err() + .expect("slow relay body should time out"); + assert!(error.contains("tunnel relay body read timed out")); + } + #[tokio::test] async fn embedded_tunnel_heartbeat_updates_proxy_node_repository() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( @@ -1288,8 +3331,12 @@ mod tests { let ack = apply_embedded_tunnel_heartbeat( &data, + &RuntimeState::memory(MemoryRuntimeStateConfig::default()), + "node-123", + TEST_TUNNEL_GENERATION, br#"{ "node_id": "node-123", + "heartbeat_session_id": "session-1", "heartbeat_id": 42, "heartbeat_interval": 45, "active_connections": 5, @@ -1334,6 +3381,49 @@ mod tests { ); } + #[tokio::test] + async fn embedded_tunnel_heartbeat_replay_does_not_repeat_counter_deltas() { + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( + "node-replay", + )])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)); + let runtime_state = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let heartbeat = br#"{ + "node_id": "node-replay", + "heartbeat_session_id": "session-replay", + "heartbeat_id": 7, + "total_requests": 9, + "failed_requests": 1, + "dns_failures": 2, + "stream_errors": 3 + }"#; + + for _ in 0..2 { + let ack = apply_embedded_tunnel_heartbeat( + &data, + &runtime_state, + "node-replay", + TEST_TUNNEL_GENERATION, + heartbeat, + ) + .await + .expect("original and replayed heartbeat should both receive an ACK"); + let ack: serde_json::Value = + serde_json::from_slice(&ack).expect("ack payload should parse"); + assert_eq!(ack["heartbeat_id"], 7); + } + + let node = repository + .find_proxy_node("node-replay") + .await + .expect("lookup should succeed") + .expect("node should exist"); + assert_eq!(node.total_requests, 9); + assert_eq!(node.failed_requests, 1); + assert_eq!(node.dns_failures, 2); + assert_eq!(node.stream_errors, 3); + } + #[tokio::test] async fn embedded_tunnel_heartbeat_rejects_missing_heartbeat_id() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( @@ -1343,8 +3433,12 @@ mod tests { let error = apply_embedded_tunnel_heartbeat( &data, + &RuntimeState::memory(MemoryRuntimeStateConfig::default()), + "node-123", + TEST_TUNNEL_GENERATION, br#"{ "node_id": "node-123", + "heartbeat_session_id": "session-1", "heartbeat_interval": 45, "active_connections": 5 }"#, @@ -1355,6 +3449,48 @@ mod tests { assert_eq!(error, "invalid heartbeat payload"); } + #[tokio::test] + async fn embedded_tunnel_heartbeat_rejects_a_different_authenticated_node() { + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( + "victim-node", + )])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)); + + let error = apply_embedded_tunnel_heartbeat( + &data, + &RuntimeState::memory(MemoryRuntimeStateConfig::default()), + "attacker-node", + TEST_TUNNEL_GENERATION, + br#"{ + "node_id": "victim-node", + "heartbeat_session_id": "session-1", + "heartbeat_id": 43, + "active_connections": 99, + "total_requests": 500, + "proxy_metadata": { + "tunnel_security": { + "mode": "disabled", + "encryption_key": "attacker-controlled" + } + } + }"#, + ) + .await + .expect_err("cross-node heartbeat must fail"); + + assert_eq!( + error, + "heartbeat node_id does not match authenticated tunnel node" + ); + let node = repository + .find_proxy_node("victim-node") + .await + .expect("lookup should succeed") + .expect("victim should still exist"); + assert_eq!(node.active_connections, 0); + assert_eq!(node.total_requests, 0); + } + #[tokio::test] async fn embedded_tunnel_node_status_updates_proxy_node_repository() { let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( @@ -1362,9 +3498,16 @@ mod tests { )])); let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository)); - apply_embedded_tunnel_node_status(&data, "node-123", true, 4, Some(1_800_000_123)) - .await - .expect("node status should succeed"); + apply_embedded_tunnel_node_status( + &data, + "node-123", + TEST_TUNNEL_GENERATION, + true, + 4, + Some(1_800_000_123), + ) + .await + .expect("node status should succeed"); let node = repository .find_proxy_node("node-123") @@ -1378,36 +3521,129 @@ mod tests { #[tokio::test] async fn tunnel_attachment_directory_syncs_and_clears_attachment_records() { - let data = GatewayDataState::disabled().with_system_config_values_for_tests(vec![]); - let directory = TunnelAttachmentDirectory::for_tests( - "gateway-a", - Some("http://gateway-a.internal"), - 90, - ); + tokio::time::timeout(Duration::from_secs(5), async { + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( + "node-123", + )])); + let data = GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_system_config_values_for_tests(vec![]); + let directory = TunnelAttachmentDirectory::for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + 90, + ); + let observed_at_unix_secs = current_unix_secs(); - directory - .sync_node_status(&data, "node-123", true, 2, 1_800_000_010) - .await - .expect("attachment should sync"); - let record = directory - .lookup_owner(&data, "node-123") - .await - .expect("lookup should succeed") - .expect("attachment should exist"); - assert_eq!(record.gateway_instance_id, "gateway-a"); - assert_eq!(record.relay_base_url, "http://gateway-a.internal"); - assert_eq!(record.conn_count, 2); - assert_eq!(record.observed_at_unix_secs, 1_800_000_010); + directory + .sync_node_status( + &data, + "node-123", + TEST_TUNNEL_GENERATION, + true, + 2, + observed_at_unix_secs, + ) + .await + .expect("attachment should sync"); + let record = directory + .lookup_owner(&data, "node-123") + .await + .expect("lookup should succeed") + .expect("attachment should exist"); + assert_eq!(record.gateway_instance_id, "gateway-a"); + assert_eq!(record.relay_base_url, "http://gateway-a.internal"); + assert_eq!(record.conn_count, 2); + assert_eq!(record.observed_at_unix_secs, observed_at_unix_secs); - directory - .sync_node_status(&data, "node-123", false, 0, 1_800_000_011) - .await - .expect("attachment should clear"); - assert!(directory - .lookup_owner(&data, "node-123") - .await - .expect("lookup should succeed") - .is_none()); + directory + .sync_node_status( + &data, + "node-123", + TEST_TUNNEL_GENERATION, + false, + 0, + observed_at_unix_secs.saturating_add(1), + ) + .await + .expect("attachment should clear"); + assert!(directory + .lookup_owner(&data, "node-123") + .await + .expect("lookup should succeed") + .is_none()); + }) + .await + .expect("attachment directory scenario should complete before timeout"); + } + + #[tokio::test] + async fn stale_gateway_disconnect_cannot_delete_new_attachment_owner() { + tokio::time::timeout(Duration::from_secs(5), async { + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node( + "node-123", + )])); + let data = Arc::new( + GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()), + ); + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let gateway_a = TunnelAttachmentDirectory::for_tests( + "gateway-a", + Some("http://gateway-a.internal"), + 90, + ) + .with_runtime_state(Arc::clone(&runtime)); + let gateway_b = TunnelAttachmentDirectory::for_tests( + "gateway-b", + Some("http://gateway-b.internal"), + 90, + ) + .with_runtime_state(runtime); + let observed_at_unix_secs = current_unix_secs(); + + gateway_a + .sync_node_status( + data.as_ref(), + "node-123", + TEST_TUNNEL_GENERATION, + true, + 1, + observed_at_unix_secs, + ) + .await + .expect("gateway A should publish attachment"); + gateway_b + .sync_node_status( + data.as_ref(), + "node-123", + TEST_TUNNEL_GENERATION, + true, + 1, + observed_at_unix_secs.saturating_add(1), + ) + .await + .expect("gateway B should publish replacement attachment"); + gateway_a + .sync_node_status( + data.as_ref(), + "node-123", + TEST_TUNNEL_GENERATION, + false, + 0, + observed_at_unix_secs.saturating_add(2), + ) + .await + .expect("stale gateway disconnect should be harmless"); + + let record = gateway_b + .lookup_owner(data.as_ref(), "node-123") + .await + .expect("attachment lookup should succeed") + .expect("new owner attachment should remain"); + assert_eq!(record.gateway_instance_id, "gateway-b"); + }) + .await + .expect("stale disconnect scenario should complete before timeout"); } #[tokio::test] @@ -1415,6 +3651,7 @@ mod tests { let stale = TunnelAttachmentRecord { gateway_instance_id: "gateway-b".to_string(), relay_base_url: "http://gateway-b.internal".to_string(), + tunnel_generation: "test-generation-stale".to_string(), conn_count: 1, observed_at_unix_secs: current_unix_secs().saturating_sub(120), }; diff --git a/apps/aether-gateway/src/upstream_admission.rs b/apps/aether-gateway/src/upstream_admission.rs index c73c16d2d..95217aab9 100644 --- a/apps/aether-gateway/src/upstream_admission.rs +++ b/apps/aether-gateway/src/upstream_admission.rs @@ -7,6 +7,7 @@ use aether_runtime::{ ConcurrencyError, ConcurrencyGate, ConcurrencyPermit, MetricKind, MetricLabel, MetricSample, }; use dashmap::DashMap; +use sha2::{Digest as _, Sha256}; use tokio::time::timeout; use url::Url; @@ -382,17 +383,61 @@ pub(crate) fn upstream_target_key_from_url( let proxy = proxy .map(str::trim) .filter(|value| !value.is_empty()) - .unwrap_or("-"); - Some(format!("{scheme}://{host}:{port}|proxy={proxy}")) + .map(safe_proxy_origin) + .unwrap_or_else(|| "-".to_string()); + Some(format!( + "{scheme}://{}:{port}|proxy={proxy}", + format_target_host(&host) + )) } fn fallback_target_key(plan: &ExecutionPlan) -> String { format!( - "unparsed|provider={}|endpoint={}|url={}", - plan.provider_id, plan.endpoint_id, plan.url + "unparsed|provider_sha256={}|endpoint_sha256={}|url_sha256={}", + short_target_hash(&plan.provider_id), + short_target_hash(&plan.endpoint_id), + short_target_hash(&plan.url) ) } +/// Return a proxy identity suitable for an in-memory key and a Prometheus +/// label. Proxy URLs frequently contain credentials, query-string tokens, or +/// fragments; none of those are relevant to target admission and must never be +/// copied into a metric label or log message. +fn safe_proxy_origin(raw_proxy: &str) -> String { + let Ok(proxy) = Url::parse(raw_proxy) else { + return "invalid".to_string(); + }; + let Some(host) = proxy.host_str() else { + return "invalid".to_string(); + }; + let scheme = proxy.scheme().trim().to_ascii_lowercase(); + if scheme.is_empty() { + return "invalid".to_string(); + } + let port = proxy + .port_or_known_default() + .map(|port| port.to_string()) + .unwrap_or_else(|| "-".to_string()); + format!("{scheme}://{}:{port}", format_target_host(host)) +} + +fn format_target_host(host: &str) -> String { + if host.contains(':') && !host.starts_with('[') { + format!("[{host}]") + } else { + host.to_string() + } +} + +fn short_target_hash(value: &str) -> String { + let digest = Sha256::digest(value.as_bytes()); + digest[..12] + .iter() + .map(|byte| format!("{byte:02x}")) + .collect() +} + fn upstream_target_metric_limit() -> usize { std::env::var(METRIC_TARGET_LIMIT_ENV) .ok() @@ -476,6 +521,48 @@ mod tests { ); } + #[test] + fn upstream_target_key_strips_proxy_credentials_and_url_components() { + let key = upstream_target_key_from_url( + "https://api.example.com/v1/chat/completions?api_key=upstream-secret#fragment", + Some("http://proxy-user:proxy-password@proxy.internal:8080/connect?token=proxy-secret#secret"), + ) + .expect("url should parse"); + + assert_eq!( + key, + "https://api.example.com:443|proxy=http://proxy.internal:8080" + ); + assert!(!key.contains("proxy-user")); + assert!(!key.contains("proxy-password")); + assert!(!key.contains("proxy-secret")); + assert!(!key.contains("upstream-secret")); + } + + #[test] + fn fallback_target_key_hashes_unparsed_url_and_identifiers() { + let mut plan = test_plan("not a valid https://user:password@example.test/?key=secret"); + plan.provider_id = "provider-secret".to_string(); + plan.endpoint_id = "endpoint-secret".to_string(); + + let key = upstream_target_key(&plan); + + assert!(key.starts_with("unparsed|provider_sha256=")); + assert!(!key.contains("provider-secret")); + assert!(!key.contains("endpoint-secret")); + assert!(!key.contains("password")); + assert!(!key.contains("key=secret")); + assert!(key.len() <= 160); + } + + #[test] + fn proxy_identity_handles_ipv6_without_ambiguity() { + assert_eq!( + safe_proxy_origin("http://user:pass@[2001:db8::1]:8080/path"), + "http://[2001:db8::1]:8080" + ); + } + #[test] fn target_queue_budget_defaults_to_short_budget() { assert_eq!( diff --git a/apps/aether-gateway/src/usage/http.rs b/apps/aether-gateway/src/usage/http.rs index af68f434c..8193918f2 100644 --- a/apps/aether-gateway/src/usage/http.rs +++ b/apps/aether-gateway/src/usage/http.rs @@ -52,7 +52,12 @@ pub(crate) async fn get_request_audit_bundle( .map_err(|err| GatewayError::Internal(err.to_string()).into_response())?; match bundle { - Some(bundle) => Ok(Json(bundle)), + Some(mut bundle) => { + if let Some(trace) = bundle.decision_trace.as_mut() { + trace.sanitize_sensitive_diagnostics(); + } + Ok(Json(bundle)) + } None => Err(( axum::http::StatusCode::NOT_FOUND, Json(json!({ diff --git a/apps/aether-gateway/src/usage/mod.rs b/apps/aether-gateway/src/usage/mod.rs index c0ad85977..cef83b8f5 100644 --- a/apps/aether-gateway/src/usage/mod.rs +++ b/apps/aether-gateway/src/usage/mod.rs @@ -11,6 +11,7 @@ pub(crate) use aether_usage_runtime::{ }; pub(crate) use aether_usage_runtime::{UsageQueueHealthSnapshot, UsageRuntimeMetricsSnapshot}; pub(crate) use reporting::{ + attach_internal_gateway_report_capability, resolve_bound_internal_gateway_report_context, spawn_sync_report, submit_stream_report, submit_sync_report, GatewayStreamReportRequest, GatewaySyncReportRequest, }; diff --git a/apps/aether-gateway/src/usage/reporting/context.rs b/apps/aether-gateway/src/usage/reporting/context.rs index bfb9375b1..8c81d8ca1 100644 --- a/apps/aether-gateway/src/usage/reporting/context.rs +++ b/apps/aether-gateway/src/usage/reporting/context.rs @@ -1,7 +1,13 @@ +use std::collections::BTreeMap; +use std::time::Duration; + use aether_data_contracts::repository::video_tasks::VideoTaskLookupKey; use aether_usage_runtime::build_locally_actionable_report_context_from_video_task; -use serde_json::Value; -use tokio::time::{sleep, Duration}; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; +use sha2::{Digest as _, Sha256}; +use tokio::time::sleep; +use uuid::Uuid; use crate::request_candidate_runtime::resolve_locally_actionable_request_candidate_report_context; use crate::video_tasks::{resolve_video_task_report_lookup, VideoTaskReportLookup}; @@ -11,6 +17,312 @@ pub(crate) use aether_usage_runtime::report_context_is_locally_actionable; const REQUEST_CANDIDATE_REPORT_CONTEXT_RETRY_ATTEMPTS: usize = 5; const REQUEST_CANDIDATE_REPORT_CONTEXT_RETRY_DELAY_MS: u64 = 50; +const INTERNAL_REPORT_CAPABILITY_FIELD: &str = "_aether_internal_report_capability"; +const INTERNAL_REPORT_CAPABILITY_KEY_PREFIX: &str = "internal:gateway:report-capability:"; +const INTERNAL_REPORT_CAPABILITY_VERSION: u8 = 1; +const INTERNAL_REPORT_CAPABILITY_TTL: Duration = Duration::from_secs(24 * 60 * 60); +const INTERNAL_REPORT_CAPABILITY_MINT_ATTEMPTS: usize = 4; +const PLAN_USAGE_RESERVATION_TOKEN_FIELD: &str = "plan_usage_reservation_token"; +const PLAN_USAGE_RESERVATION_DEFERRED_FIELD: &str = "plan_usage_reservation_deferred"; + +/// Fields produced while observing an upstream response. Everything else in the +/// planner-issued context is immutable and covered by the capability digest. +/// +/// This allowlist is intentionally top-level and fail-closed: adding a future +/// report-context side effect requires explicitly classifying it as an observation. +const INTERNAL_REPORT_OBSERVATION_FIELDS: &[&str] = &[ + "provider_response_headers", + "provider_request_started_at_unix_ms", + "provider_response_headers_observed_at_unix_ms", + "provider_request_order_id", + "client_response_status_code", + "client_response_headers", + "upstream_response", + "error_flow", + "transport_error", + "input_tokens", + "cache_creation_input_tokens", + "cache_read_input_tokens", + "kiro_simulated_cache_enabled", + "stage_timings_ms", + "db_timings_ms", + "end_to_end_time_ms", + "end_to_end_first_byte_time_ms", + "windsurf_native_runtime", + "windsurf_language_server_port", +]; + +#[derive(Debug, Serialize, Deserialize)] +struct InternalReportCapabilityRecord { + version: u8, + trace_id: String, + report_scope: String, + protected_context_sha256: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + kiro_web_search_context_sha256: Option, +} + +pub(crate) async fn attach_internal_gateway_report_capability( + state: &AppState, + trace_id: &str, + report_kind: Option<&str>, + provider_request_headers: &BTreeMap, + report_context: &mut Option, +) -> Result<(), crate::GatewayError> { + let Some(report_kind) = report_kind else { + return Ok(()); + }; + let Some(report_scope) = internal_report_capability_scope(report_kind) else { + return Ok(()); + }; + let Some(context) = report_context.as_mut().and_then(Value::as_object_mut) else { + return Err(crate::GatewayError::Internal( + "internal gateway report capability requires an object context".to_string(), + )); + }; + if context.contains_key(INTERNAL_REPORT_CAPABILITY_FIELD) { + return Err(crate::GatewayError::Internal( + "internal gateway planner produced a reserved report capability field".to_string(), + )); + } + context.insert( + "provider_request_headers".to_string(), + serde_json::to_value(provider_request_headers) + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?, + ); + + let protected_context_sha256 = protected_internal_report_context_sha256(context)?; + let kiro_web_search_context_sha256 = kiro_web_search_internal_report_context_sha256(context)?; + let record = InternalReportCapabilityRecord { + version: INTERNAL_REPORT_CAPABILITY_VERSION, + trace_id: trace_id.to_string(), + report_scope, + protected_context_sha256, + kiro_web_search_context_sha256, + }; + let serialized = serde_json::to_string(&record) + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?; + + for _ in 0..INTERNAL_REPORT_CAPABILITY_MINT_ATTEMPTS { + let capability = Uuid::new_v4().simple().to_string(); + let storage_key = internal_report_capability_storage_key(&capability); + let inserted = state + .runtime_state + .kv_set_if_absent( + &storage_key, + serialized.clone(), + INTERNAL_REPORT_CAPABILITY_TTL, + ) + .await + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?; + if inserted { + context.insert( + INTERNAL_REPORT_CAPABILITY_FIELD.to_string(), + Value::String(capability), + ); + return Ok(()); + } + } + + Err(crate::GatewayError::Internal( + "failed to allocate a unique internal gateway report capability".to_string(), + )) +} + +/// Validate and atomically consume a planner-issued report capability, then +/// return a context whose planner fields are equivalent after canonical JSON +/// normalization. +/// +/// The internal request HMAC authenticates a peer, but a peer must not choose a +/// candidate, user, provider key, video task, or file mapping target. The opaque +/// capability is independent from diagnostic candidate persistence, so `terminal` +/// and `none` persistence modes remain functional. +pub(crate) async fn resolve_bound_internal_gateway_report_context( + state: &AppState, + trace_id: &str, + report_kind: &str, + report_context: Option<&Value>, +) -> Result, crate::GatewayError> { + let Some(context) = report_context.and_then(Value::as_object) else { + return Ok(None); + }; + let Some(capability) = context + .get(INTERNAL_REPORT_CAPABILITY_FIELD) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Ok(None); + }; + if Uuid::parse_str(capability).is_err() { + return Ok(None); + } + if !internal_report_late_bound_settlement_fields_are_valid(context) + || !internal_report_windsurf_observation_is_valid(context) + { + return Ok(None); + } + + let storage_key = internal_report_capability_storage_key(capability); + let Some(serialized) = state + .runtime_state + .kv_get(&storage_key) + .await + .map_err(|error| crate::GatewayError::Internal(error.to_string()))? + else { + return Ok(None); + }; + let record: InternalReportCapabilityRecord = serde_json::from_str(&serialized) + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?; + let protected_context_sha256 = protected_internal_report_context_sha256(context)?; + let context_matches = protected_context_sha256 == record.protected_context_sha256 + || record.kiro_web_search_context_sha256.as_deref() + == Some(protected_context_sha256.as_str()); + if record.version != INTERNAL_REPORT_CAPABILITY_VERSION + || record.trace_id != trace_id + || internal_report_capability_scope(report_kind).as_deref() + != Some(record.report_scope.as_str()) + || !context_matches + { + return Ok(None); + } + + // A report can mutate billing, provider health, file mappings, and video + // tasks. Consume its capability so a signed peer cannot replay those side + // effects. The second read is atomic: concurrent valid submissions have a + // single winner, while invalid submissions above cannot burn the token. + let claimed = state + .runtime_state + .kv_take(&storage_key) + .await + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?; + if claimed.as_deref() != Some(serialized.as_str()) { + return Ok(None); + } + + let mut resolved = context.clone(); + resolved.remove(INTERNAL_REPORT_CAPABILITY_FIELD); + Ok(Some(Value::Object(resolved))) +} + +fn internal_report_capability_scope(report_kind: &str) -> Option { + let mut scope = report_kind.trim().to_ascii_lowercase(); + if scope.is_empty() || scope.len() > 160 { + return None; + } + for suffix in ["_success", "_error", "_failed", "_cancelled", "_finalize"] { + if let Some(value) = scope.strip_suffix(suffix) { + scope = value.to_string(); + break; + } + } + for suffix in ["_sync", "_stream"] { + if let Some(value) = scope.strip_suffix(suffix) { + scope = value.to_string(); + break; + } + } + (!scope.is_empty()).then_some(scope) +} + +fn internal_report_capability_storage_key(capability: &str) -> String { + let digest = Sha256::digest(capability.as_bytes()); + format!("{INTERNAL_REPORT_CAPABILITY_KEY_PREFIX}{digest:x}") +} + +fn protected_internal_report_context_sha256( + context: &Map, +) -> Result { + let mut protected = context.clone(); + protected.remove(INTERNAL_REPORT_CAPABILITY_FIELD); + protected.remove(PLAN_USAGE_RESERVATION_TOKEN_FIELD); + protected.remove(PLAN_USAGE_RESERVATION_DEFERRED_FIELD); + for field in INTERNAL_REPORT_OBSERVATION_FIELDS { + protected.remove(*field); + } + let canonical = canonicalize_internal_report_json(&Value::Object(protected)); + let encoded = serde_json::to_vec(&canonical) + .map_err(|error| crate::GatewayError::Internal(error.to_string()))?; + Ok(format!("{:x}", Sha256::digest(encoded))) +} + +/// Kiro MCP web-search execution emits a synthetic response that is already in +/// the client contract. The executor must therefore disable the planner's Kiro +/// envelope conversion before observing that response. Bind that exact, fixed +/// transformation when the capability is minted instead of making the affected +/// planner fields globally mutable. +fn kiro_web_search_internal_report_context_sha256( + context: &Map, +) -> Result, crate::GatewayError> { + let is_kiro_envelope = context + .get("envelope_name") + .and_then(Value::as_str) + .is_some_and(|value| { + value.eq_ignore_ascii_case(aether_provider_transport::kiro::KIRO_ENVELOPE_NAME) + }); + if !is_kiro_envelope { + return Ok(None); + } + + let mut synthetic = context.clone(); + synthetic.insert("has_envelope".to_string(), Value::Bool(false)); + synthetic.insert("needs_conversion".to_string(), Value::Bool(false)); + synthetic.remove("envelope_name"); + synthetic.insert("kiro_web_search_mcp".to_string(), Value::Bool(true)); + protected_internal_report_context_sha256(&synthetic).map(Some) +} + +/// HTTP candidate execution may create a plan-cost reservation only after the +/// planner has issued the report capability. Its opaque token is therefore +/// late-bound, but the peer cannot use it to select another request or user: +/// repository reconciliation also requires the capability-bound request and +/// subject identities. Deferred reconciliation is not a normal HTTP report +/// outcome and remains forbidden here; allowing it would let a peer strand a +/// reservation without submitting a terminal reconciliation. +fn internal_report_late_bound_settlement_fields_are_valid(context: &Map) -> bool { + match ( + context.get(PLAN_USAGE_RESERVATION_TOKEN_FIELD), + context.get(PLAN_USAGE_RESERVATION_DEFERRED_FIELD), + ) { + (None, None) => true, + (Some(Value::String(token)), Some(Value::Bool(false))) => { + Uuid::parse_str(token.trim()).is_ok() + } + _ => false, + } +} + +fn internal_report_windsurf_observation_is_valid(context: &Map) -> bool { + match ( + context.get("windsurf_native_runtime"), + context.get("windsurf_language_server_port"), + ) { + (None, None) => true, + (Some(Value::Bool(true)), Some(Value::Number(port))) => port + .as_u64() + .is_some_and(|port| u16::try_from(port).is_ok() && port != 0), + _ => false, + } +} + +fn canonicalize_internal_report_json(value: &Value) -> Value { + match value { + Value::Array(values) => Value::Array( + values + .iter() + .map(canonicalize_internal_report_json) + .collect(), + ), + Value::Object(object) => { + let mut entries = object.iter().collect::>(); + entries.sort_unstable_by_key(|(left, _)| *left); + Value::Object(Map::from_iter(entries.into_iter().map(|(key, value)| { + (key.clone(), canonicalize_internal_report_json(value)) + }))) + } + other => other.clone(), + } +} pub(crate) async fn resolve_locally_actionable_report_context( state: &AppState, @@ -78,9 +390,14 @@ async fn resolve_locally_actionable_report_context_from_video_task( state: &AppState, context: &Value, ) -> Option { + let requested_user_id = requested_report_user_id(context)?; let task = match resolve_video_task_report_lookup(context)? { VideoTaskReportLookup::Lookup(lookup) => { - state.data.find_video_task(lookup).await.ok()?? + let task = state.data.find_video_task(lookup).await.ok()??; + if !video_task_matches_requested_user(&task, requested_user_id) { + return None; + } + task } VideoTaskReportLookup::TaskIdOrExternal { task_id, user_id } => { if let Some(task) = state @@ -88,21 +405,47 @@ async fn resolve_locally_actionable_report_context_from_video_task( .find_video_task(VideoTaskLookupKey::Id(task_id)) .await .ok()? + .filter(|task| video_task_matches_requested_user(task, requested_user_id)) { task } else { let user_id = user_id?; - state + let task = state .data .find_video_task(VideoTaskLookupKey::UserExternal { user_id, external_task_id: task_id, }) .await - .ok()?? + .ok()??; + if !video_task_matches_requested_user(&task, requested_user_id) { + return None; + } + task } } }; build_locally_actionable_report_context_from_video_task(context, &task) } + +fn requested_report_user_id(context: &Value) -> Option> { + match context.get("user_id") { + None => Some(None), + Some(value) => value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(Some), + } +} + +fn video_task_matches_requested_user( + task: &aether_data_contracts::repository::video_tasks::StoredVideoTask, + requested_user_id: Option<&str>, +) -> bool { + let Some(requested_user_id) = requested_user_id else { + return true; + }; + task.user_id.as_deref().map(str::trim) == Some(requested_user_id) +} diff --git a/apps/aether-gateway/src/usage/reporting/mod.rs b/apps/aether-gateway/src/usage/reporting/mod.rs index 3e768d156..e2ea21e2f 100644 --- a/apps/aether-gateway/src/usage/reporting/mod.rs +++ b/apps/aether-gateway/src/usage/reporting/mod.rs @@ -13,6 +13,9 @@ use crate::task_runtime::{spawn_fire_and_forget, TASK_KEY_USAGE_SYNC_REPORT}; use crate::{AppState, GatewayError}; mod context; +pub(crate) use context::{ + attach_internal_gateway_report_capability, resolve_bound_internal_gateway_report_context, +}; use context::{report_context_is_locally_actionable, resolve_locally_actionable_report_context}; use aether_usage_runtime::{ @@ -307,6 +310,7 @@ mod tests { use std::collections::BTreeMap; use std::sync::Arc; + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; use aether_data::repository::candidates::InMemoryRequestCandidateRepository; use aether_data::repository::gemini_file_mappings::{ GeminiFileMappingReadRepository, InMemoryGeminiFileMappingRepository, @@ -328,6 +332,7 @@ mod tests { use serde_json::json; use super::{ + attach_internal_gateway_report_capability, resolve_bound_internal_gateway_report_context, resolve_locally_actionable_report_context, submit_stream_report, submit_sync_report, GatewayStreamReportRequest, GatewaySyncReportRequest, }; @@ -413,6 +418,42 @@ mod tests { ) } + async fn mint_internal_report_capability( + state: &AppState, + trace_id: &str, + report_kind: &str, + context: serde_json::Value, + ) -> serde_json::Value { + mint_internal_report_capability_with_headers( + state, + trace_id, + report_kind, + &BTreeMap::new(), + context, + ) + .await + } + + async fn mint_internal_report_capability_with_headers( + state: &AppState, + trace_id: &str, + report_kind: &str, + provider_request_headers: &BTreeMap, + context: serde_json::Value, + ) -> serde_json::Value { + let mut report_context = Some(context); + attach_internal_gateway_report_capability( + state, + trace_id, + Some(report_kind), + provider_request_headers, + &mut report_context, + ) + .await + .expect("report capability should mint"); + report_context.expect("report context should remain present") + } + fn build_video_test_state( video_repository: Arc, request_candidate_repository: Arc, @@ -447,7 +488,8 @@ mod tests { AppState::new() .expect("gateway state should build") .with_data_state_for_tests( - GatewayDataState::with_provider_catalog_repository_for_tests(repository), + GatewayDataState::with_provider_catalog_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), ) } @@ -465,6 +507,15 @@ mod tests { } fn sample_provider_catalog_key(key_id: &str, provider_id: &str) -> StoredProviderCatalogKey { + let credential_state = AppState::new() + .expect("credential state should build") + .with_data_state_for_tests( + GatewayDataState::disabled() + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let encrypted_api_key = credential_state + .seal_provider_catalog_key_api_key(provider_id, key_id, "sk-codex-test") + .expect("api key should encrypt"); StoredProviderCatalogKey::new( key_id.to_string(), provider_id.to_string(), @@ -476,7 +527,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:responses"])), - "sk-codex-test".to_string(), + encrypted_api_key, None, None, None, @@ -530,6 +581,7 @@ mod tests { repository: &InMemoryVideoTaskRepository, id: &str, short_id: Option<&str>, + external_task_id: &str, request_id: &str, user_id: &str, api_key_id: &str, @@ -548,7 +600,7 @@ mod tests { api_key_id: Some(api_key_id.to_string()), username: Some("video-user".to_string()), api_key_name: Some("video-key".to_string()), - external_task_id: Some("ext-video-task-reporting-123".to_string()), + external_task_id: Some(external_task_id.to_string()), provider_id: Some(provider_id.to_string()), endpoint_id: Some(endpoint_id.to_string()), key_id: Some(key_id.to_string()), @@ -598,6 +650,441 @@ mod tests { assert!(resolved.is_none()); } + #[tokio::test] + async fn internal_report_capability_validates_once_without_candidate_persistence() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let mut context = mint_internal_report_capability( + &state, + "trace-capability-valid-123", + "openai_chat_sync_success", + json!({ + "request_id": "req-capability-valid-123", + "candidate_id": "cand-capability-valid-123", + "user_id": "user-capability-valid-123", + "provider_id": "provider-reporting-tests-123", + "key_id": "key-reporting-tests-123", + }), + ) + .await; + context + .as_object_mut() + .expect("context should be an object") + .extend([ + ( + "provider_response_headers".to_string(), + json!({"x-ratelimit-limit": "100"}), + ), + ("upstream_response".to_string(), json!({"id": "resp-123"})), + ("error_flow".to_string(), json!({"stage": "upstream"})), + ( + "client_response_headers".to_string(), + json!({"content-type": "application/json"}), + ), + ]); + + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-valid-123", + "openai_chat_sync_error", + Some(&context), + ) + .await + .expect("capability lookup should succeed") + .expect("capability should validate"); + assert!(resolved.get("_aether_internal_report_capability").is_none()); + assert_eq!(resolved["upstream_response"]["id"], "resp-123"); + + let replay = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-valid-123", + "openai_chat_sync_success", + Some(&context), + ) + .await + .expect("capability lookup should succeed"); + assert!(replay.is_none(), "a consumed capability must not replay"); + } + + #[tokio::test] + async fn internal_report_capability_rejects_missing_unknown_trace_and_scope_mismatches() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let protected = json!({ + "request_id": "req-capability-reject-123", + "candidate_id": "cand-capability-reject-123", + "user_id": "user-capability-reject-123", + "provider_id": "provider-capability-reject-123", + "key_id": "key-capability-reject-123", + }); + let minted = mint_internal_report_capability( + &state, + "trace-capability-reject-123", + "openai_chat_stream_success", + protected.clone(), + ) + .await; + + let unknown = { + let mut value = minted.clone(); + value["_aether_internal_report_capability"] = + json!(uuid::Uuid::new_v4().simple().to_string()); + value + }; + for (trace_id, report_kind, context) in [ + ( + "trace-capability-reject-123", + "openai_chat_stream_success", + protected.clone(), + ), + ( + "trace-capability-reject-123", + "openai_chat_stream_success", + unknown, + ), + ( + "trace-capability-attacker-123", + "openai_chat_stream_success", + minted.clone(), + ), + ( + "trace-capability-reject-123", + "gemini_files_store_mapping", + minted.clone(), + ), + ( + "trace-capability-reject-123", + "openai_video_create_sync_finalize", + minted.clone(), + ), + ] { + let rejected = resolve_bound_internal_gateway_report_context( + &state, + trace_id, + report_kind, + Some(&context), + ) + .await + .expect("capability lookup should succeed"); + assert!( + rejected.is_none(), + "invalid capability use must be rejected" + ); + } + + let valid = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-reject-123", + "openai_chat_sync_finalize", + Some(&minted), + ) + .await + .expect("capability lookup should succeed"); + assert!( + valid.is_some(), + "invalid attempts must not consume the capability" + ); + } + + #[tokio::test] + async fn internal_report_capability_rejects_every_protected_identity_mutation() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let minted = mint_internal_report_capability( + &state, + "trace-capability-fields-123", + "openai_video_create_sync_finalize", + json!({ + "request_id": "req-capability-fields-123", + "candidate_id": "cand-capability-fields-123", + "user_id": "user-capability-fields-123", + "api_key_id": "api-key-capability-fields-123", + "provider_id": "provider-capability-fields-123", + "endpoint_id": "endpoint-capability-fields-123", + "key_id": "key-capability-fields-123", + "file_key_id": "file-key-capability-fields-123", + "task_id": "task-capability-fields-123", + "local_task_id": "local-task-capability-fields-123", + "local_short_id": "short-capability-fields-123", + "file_name": "files/capability-fields-123", + "client_api_format": "openai:video", + "upstream_url": "https://provider.example/v1/videos", + "has_envelope": false, + "needs_conversion": false, + }), + ) + .await; + + for field in [ + "user_id", + "api_key_id", + "provider_id", + "endpoint_id", + "key_id", + "file_key_id", + "task_id", + "local_task_id", + "local_short_id", + "file_name", + "client_api_format", + "upstream_url", + "has_envelope", + "needs_conversion", + ] { + let mut forged = minted.clone(); + forged[field] = json!(format!("attacker-{field}")); + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-fields-123", + "openai_video_create_sync_finalize", + Some(&forged), + ) + .await + .expect("capability lookup should succeed"); + assert!(resolved.is_none(), "mutating {field} must be rejected"); + } + + for report_kind in [ + "openai_video_delete_sync_finalize", + "openai_video_cancel_sync_finalize", + "gemini_video_create_sync_finalize", + ] { + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-fields-123", + report_kind, + Some(&minted), + ) + .await + .expect("capability lookup should succeed"); + assert!(resolved.is_none(), "cross-operation use must be rejected"); + } + + let valid = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-fields-123", + "openai_video_create_sync_error", + Some(&minted), + ) + .await + .expect("capability lookup should succeed"); + assert!(valid.is_some()); + } + + #[tokio::test] + async fn internal_report_capability_allows_only_the_bound_kiro_web_search_transform() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let minted = mint_internal_report_capability( + &state, + "trace-capability-kiro-search-123", + "openai_chat_stream_success", + json!({ + "request_id": "req-capability-kiro-search-123", + "candidate_id": "cand-capability-kiro-search-123", + "user_id": "user-capability-kiro-search-123", + "provider_id": "provider-capability-kiro-search-123", + "key_id": "key-capability-kiro-search-123", + "upstream_url": "https://kiro.example/generateAssistantResponse", + "has_envelope": true, + "needs_conversion": true, + "envelope_name": aether_provider_transport::kiro::KIRO_ENVELOPE_NAME, + }), + ) + .await; + + let mut forged_target = minted.clone(); + forged_target["upstream_url"] = json!("https://attacker.example/forged"); + forged_target["has_envelope"] = json!(false); + forged_target["needs_conversion"] = json!(false); + forged_target["kiro_web_search_mcp"] = json!(true); + forged_target + .as_object_mut() + .expect("context should be an object") + .remove("envelope_name"); + let rejected = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-kiro-search-123", + "openai_chat_stream_success", + Some(&forged_target), + ) + .await + .expect("capability lookup should succeed"); + assert!(rejected.is_none(), "the upstream target must remain bound"); + + let mut synthetic = minted; + synthetic["has_envelope"] = json!(false); + synthetic["needs_conversion"] = json!(false); + synthetic["kiro_web_search_mcp"] = json!(true); + synthetic + .as_object_mut() + .expect("context should be an object") + .remove("envelope_name"); + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-kiro-search-123", + "openai_chat_stream_success", + Some(&synthetic), + ) + .await + .expect("capability lookup should succeed") + .expect("the pre-bound Kiro web-search transform should validate"); + assert_eq!(resolved["kiro_web_search_mcp"], json!(true)); + assert_eq!( + resolved["upstream_url"], + json!("https://kiro.example/generateAssistantResponse") + ); + } + + #[tokio::test] + async fn internal_report_capability_binds_final_provider_request_headers() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let headers = BTreeMap::from([ + ( + "authorization".to_string(), + "Bearer final-token".to_string(), + ), + ("content-type".to_string(), "application/json".to_string()), + ]); + let minted = mint_internal_report_capability_with_headers( + &state, + "trace-capability-headers-123", + "openai_video_create_sync_success", + &headers, + json!({ + "request_id": "req-capability-headers-123", + "provider_request_headers": {"authorization": "Bearer stale-token"}, + }), + ) + .await; + assert_eq!(minted["provider_request_headers"], json!(headers)); + + let mut forged = minted.clone(); + forged["provider_request_headers"]["authorization"] = json!("Bearer attacker-token"); + let rejected = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-headers-123", + "openai_video_create_sync_success", + Some(&forged), + ) + .await + .expect("capability lookup should succeed"); + assert!( + rejected.is_none(), + "final request headers must remain bound" + ); + + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-headers-123", + "openai_video_create_sync_success", + Some(&minted), + ) + .await + .expect("capability lookup should succeed") + .expect("the authoritative final headers should validate"); + assert_eq!(resolved["provider_request_headers"], json!(headers)); + } + + #[tokio::test] + async fn internal_report_capability_validates_windsurf_native_observation_shape() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let minted = mint_internal_report_capability( + &state, + "trace-capability-windsurf-123", + "openai_responses_stream_success", + json!({"request_id": "req-capability-windsurf-123"}), + ) + .await; + + for (native_runtime, port) in [ + (json!(false), json!(42_137)), + (json!(true), json!(0)), + (json!(true), json!(65_536)), + (json!(true), json!("42137")), + ] { + let mut invalid = minted.clone(); + invalid["windsurf_native_runtime"] = native_runtime; + invalid["windsurf_language_server_port"] = port; + let rejected = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-windsurf-123", + "openai_responses_stream_success", + Some(&invalid), + ) + .await + .expect("capability lookup should succeed"); + assert!(rejected.is_none(), "invalid Windsurf metadata must fail"); + } + + let mut valid = minted; + valid["windsurf_native_runtime"] = json!(true); + valid["windsurf_language_server_port"] = json!(42_137); + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-windsurf-123", + "openai_responses_stream_success", + Some(&valid), + ) + .await + .expect("capability lookup should succeed") + .expect("valid Windsurf native metadata should pass"); + assert_eq!(resolved["windsurf_native_runtime"], json!(true)); + assert_eq!(resolved["windsurf_language_server_port"], json!(42_137)); + } + + #[tokio::test] + async fn internal_report_capability_accepts_only_non_deferred_late_bound_reservations() { + let repository = Arc::new(InMemoryRequestCandidateRepository::default()); + let state = build_test_state(repository); + let minted = mint_internal_report_capability( + &state, + "trace-capability-reservation-123", + "openai_chat_sync_success", + json!({ + "request_id": "req-capability-reservation-123", + "candidate_id": "cand-capability-reservation-123", + "user_id": "user-capability-reservation-123", + "provider_id": "provider-capability-reservation-123", + "key_id": "key-capability-reservation-123", + }), + ) + .await; + + let mut deferred = minted.clone(); + deferred["plan_usage_reservation_token"] = json!(uuid::Uuid::new_v4().to_string()); + deferred["plan_usage_reservation_deferred"] = json!(true); + let rejected = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-reservation-123", + "openai_chat_sync_success", + Some(&deferred), + ) + .await + .expect("capability lookup should succeed"); + assert!( + rejected.is_none(), + "a peer must not defer terminal reservation reconciliation" + ); + + let reservation_token = uuid::Uuid::new_v4().to_string(); + let mut terminal = minted; + terminal["plan_usage_reservation_token"] = json!(reservation_token); + terminal["plan_usage_reservation_deferred"] = json!(false); + let resolved = resolve_bound_internal_gateway_report_context( + &state, + "trace-capability-reservation-123", + "openai_chat_sync_success", + Some(&terminal), + ) + .await + .expect("capability lookup should succeed") + .expect("a server-issued terminal reservation token should validate"); + assert_eq!(resolved["plan_usage_reservation_deferred"], json!(false)); + } + #[tokio::test] async fn submit_sync_report_handles_request_id_only_context_locally_when_unique_candidate_exists( ) { @@ -821,10 +1308,7 @@ mod tests { stored[0].error_type.as_deref(), Some("stream_missing_terminal_event") ); - assert_eq!( - stored[0].error_message.as_deref(), - Some("execution runtime stream ended before provider terminal event") - ); + assert!(stored[0].error_message.is_none()); } #[tokio::test] @@ -1365,7 +1849,7 @@ mod tests { "req-gemini-files-delete-123", ), ])); - let gemini_file_mapping_repository = Arc::new(InMemoryGeminiFileMappingRepository::seed([ + let mut mapping = aether_data::repository::gemini_file_mappings::StoredGeminiFileMapping::new( "mapping-gemini-files-delete-123".to_string(), "files/delete-me".to_string(), @@ -1373,8 +1857,10 @@ mod tests { 1_700_000_000, 1_700_172_800, ) - .expect("gemini file mapping should build"), - ])); + .expect("gemini file mapping should build"); + mapping.user_id = Some("user-reporting-tests-123".to_string()); + let gemini_file_mapping_repository = + Arc::new(InMemoryGeminiFileMappingRepository::seed([mapping])); let state = build_gemini_file_mapping_test_state( Arc::clone(&request_candidate_repository), Arc::clone(&gemini_file_mapping_repository), @@ -1392,6 +1878,7 @@ mod tests { "provider_id": "provider-reporting-tests-123", "endpoint_id": "endpoint-reporting-tests-123", "key_id": "key-reporting-tests-123", + "user_id": "user-reporting-tests-123", "file_name": "delete-me", })), status_code: 204, @@ -1461,6 +1948,7 @@ mod tests { &video_repository, "task-openai-video-reporting-123", None, + "ext-video-task-reporting-123", "req-openai-video-reporting-123", "user-openai-video-reporting-123", "api-key-openai-video-reporting-123", @@ -1522,6 +2010,7 @@ mod tests { &video_repository, "task-gemini-video-reporting-123", Some("short-gemini-video-reporting-123"), + "ext-video-task-reporting-123", "req-gemini-video-reporting-123", "user-gemini-video-reporting-123", "api-key-gemini-video-reporting-123", @@ -1582,6 +2071,7 @@ mod tests { &video_repository, "task-openai-video-task-id-123", None, + "ext-video-task-reporting-123", "req-openai-video-task-id-123", "user-openai-video-task-id-123", "api-key-openai-video-task-id-123", @@ -1642,6 +2132,7 @@ mod tests { &video_repository, "task-gemini-video-external-id-123", Some("short-gemini-video-external-id-123"), + "models/veo-3/operations/ext-gemini-video-123", "req-gemini-video-external-id-123", "user-gemini-video-external-id-123", "api-key-gemini-video-external-id-123", @@ -1652,48 +2143,6 @@ mod tests { "gemini:video", ) .await; - video_repository - .upsert(UpsertVideoTask { - id: "task-gemini-video-external-id-123".to_string(), - short_id: Some("short-gemini-video-external-id-123".to_string()), - request_id: "req-gemini-video-external-id-123".to_string(), - user_id: Some("user-gemini-video-external-id-123".to_string()), - api_key_id: Some("api-key-gemini-video-external-id-123".to_string()), - username: Some("video-user".to_string()), - api_key_name: Some("video-key".to_string()), - external_task_id: Some("models/veo-3/operations/ext-gemini-video-123".to_string()), - provider_id: Some("provider-gemini-video-external-id-123".to_string()), - endpoint_id: Some("endpoint-gemini-video-external-id-123".to_string()), - key_id: Some("key-gemini-video-external-id-123".to_string()), - client_api_format: Some("gemini:video".to_string()), - provider_api_format: Some("gemini:video".to_string()), - format_converted: false, - model: Some("video-model".to_string()), - prompt: Some("video prompt".to_string()), - original_request_body: Some(json!({"prompt": "video prompt"})), - duration_seconds: Some(4), - resolution: Some("720p".to_string()), - aspect_ratio: Some("16:9".to_string()), - size: Some("1280x720".to_string()), - status: VideoTaskStatus::Submitted, - progress_percent: 0, - progress_message: None, - retry_count: 0, - poll_interval_seconds: 10, - next_poll_at_unix_secs: Some(1_700_000_010), - poll_count: 0, - max_poll_count: 360, - created_at_unix_ms: 1_700_000_000, - submitted_at_unix_secs: Some(1_700_000_000), - completed_at_unix_secs: None, - updated_at_unix_secs: 1_700_000_000, - error_code: None, - error_message: None, - video_url: None, - request_metadata: None, - }) - .await - .expect("video task should update external id"); let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); let state = build_video_test_state(video_repository, Arc::clone(&request_candidate_repository)); @@ -1737,4 +2186,67 @@ mod tests { Some("key-gemini-video-external-id-123") ); } + + #[tokio::test] + async fn report_context_task_id_collision_stays_bound_to_requested_user() { + let video_repository = Arc::new(InMemoryVideoTaskRepository::default()); + seed_video_task( + &video_repository, + "shared-task-id", + Some("victim-short-id"), + "victim-external-id", + "req-video-victim", + "user-video-victim", + "api-key-video-victim", + "provider-video-victim", + "endpoint-video-victim", + "key-video-victim", + "gemini:video", + "gemini:video", + ) + .await; + seed_video_task( + &video_repository, + "owner-local-id", + Some("owner-short-id"), + "shared-task-id", + "req-video-owner", + "user-video-owner", + "api-key-video-owner", + "provider-video-owner", + "endpoint-video-owner", + "key-video-owner", + "gemini:video", + "gemini:video", + ) + .await; + let state = build_video_test_state( + video_repository, + Arc::new(InMemoryRequestCandidateRepository::default()), + ); + + let resolved = resolve_locally_actionable_report_context( + &state, + Some(&json!({ + "task_id": "shared-task-id", + "user_id": "user-video-owner", + })), + ) + .await + .expect("the owner's external task id should resolve"); + + assert_eq!(resolved["request_id"], "req-video-owner"); + assert_eq!(resolved["provider_id"], "provider-video-owner"); + assert_eq!(resolved["user_id"], "user-video-owner"); + + let foreign_local_id = resolve_locally_actionable_report_context( + &state, + Some(&json!({ + "local_task_id": "shared-task-id", + "user_id": "user-video-owner", + })), + ) + .await; + assert!(foreign_local_id.is_none()); + } } diff --git a/apps/aether-gateway/src/video_tasks/mod.rs b/apps/aether-gateway/src/video_tasks/mod.rs index 677de6b3a..eb7c1012b 100644 --- a/apps/aether-gateway/src/video_tasks/mod.rs +++ b/apps/aether-gateway/src/video_tasks/mod.rs @@ -1,4 +1,5 @@ use crate::control::GatewayControlAuthContext; +use crate::GatewayError; use aether_contracts::{ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody}; pub(crate) use self::service::VideoTaskService; @@ -33,5 +34,16 @@ mod service; mod store; mod types; +pub(crate) fn not_found_body() -> serde_json::Value { + serde_json::json!({"detail": "Video task not found"}) +} + +pub(crate) fn not_found_error() -> GatewayError { + GatewayError::Client { + status: http::StatusCode::NOT_FOUND, + message: "Video task not found".to_string(), + } +} + #[cfg(test)] mod tests; diff --git a/apps/aether-gateway/src/video_tasks/service.rs b/apps/aether-gateway/src/video_tasks/service.rs index 024d071ef..c3eadcb7d 100644 --- a/apps/aether-gateway/src/video_tasks/service.rs +++ b/apps/aether-gateway/src/video_tasks/service.rs @@ -16,9 +16,10 @@ impl VideoTaskService { pub(crate) fn with_file_store( mode: VideoTaskTruthSourceMode, path: impl Into, + encryption_key: impl Into, ) -> std::io::Result { Ok(Self( - aether_video_tasks_core::VideoTaskService::with_file_store(mode, path)?, + aether_video_tasks_core::VideoTaskService::with_file_store(mode, path, encryption_key)?, )) } @@ -43,6 +44,43 @@ impl VideoTaskService { trace_id, ) } + + pub(crate) fn prepare_follow_up_sync_plan_for_user( + &self, + plan_kind: &str, + request_path: &str, + body_json: Option<&serde_json::Value>, + auth_context: Option<&GatewayControlAuthContext>, + trace_id: &str, + ) -> Option { + self.0.prepare_follow_up_sync_plan_for_user( + plan_kind, + request_path, + body_json, + auth_context.map(|value| value.user_id.as_str()), + auth_context.map(|value| value.api_key_id.as_str()), + trace_id, + ) + } + + pub(crate) fn prepare_follow_up_sync_plan_for_user_id( + &self, + plan_kind: &str, + request_path: &str, + body_json: Option<&serde_json::Value>, + user_id: &str, + api_key_id: Option<&str>, + trace_id: &str, + ) -> Option { + self.0.prepare_follow_up_sync_plan_for_user( + plan_kind, + request_path, + body_json, + Some(user_id), + api_key_id, + trace_id, + ) + } } impl Deref for VideoTaskService { diff --git a/apps/aether-gateway/src/video_tasks/tests/fixtures.rs b/apps/aether-gateway/src/video_tasks/tests/fixtures.rs index 9b588b813..349bc9d0d 100644 --- a/apps/aether-gateway/src/video_tasks/tests/fixtures.rs +++ b/apps/aether-gateway/src/video_tasks/tests/fixtures.rs @@ -55,6 +55,7 @@ pub(super) fn sample_auth_context() -> GatewayControlAuthContext { local_rejection: None, allowed_models: None, ip_rules: None, + verified_api_key_hash: None, } } diff --git a/apps/aether-gateway/src/video_tasks/tests/plans.rs b/apps/aether-gateway/src/video_tasks/tests/plans.rs index bb0042ffd..e40ec578a 100644 --- a/apps/aether-gateway/src/video_tasks/tests/plans.rs +++ b/apps/aether-gateway/src/video_tasks/tests/plans.rs @@ -62,6 +62,30 @@ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() { .and_then(Value::as_str), Some("task-local-123") ); + + let mut same_user_new_key = sample_auth_context(); + same_user_new_key.api_key_id = "key-rotated".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_cancel_sync", + "/v1/videos/task-local-123/cancel", + None, + Some(&same_user_new_key), + "trace-openai-cancel-owner-rotated-key", + ) + .is_some()); + let mut foreign_auth = sample_auth_context(); + foreign_auth.user_id = "user-foreign".to_string(); + foreign_auth.api_key_id = "key-foreign".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_cancel_sync", + "/v1/videos/task-local-123/cancel", + None, + Some(&foreign_auth), + "trace-openai-cancel-foreign", + ) + .is_none()); } #[test] @@ -123,6 +147,30 @@ fn rust_authoritative_service_builds_openai_remix_follow_up_plan() { .and_then(Value::as_str), Some("task-local-123") ); + + let mut same_user_new_key = sample_auth_context(); + same_user_new_key.api_key_id = "key-rotated".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_remix_sync", + "/v1/videos/task-local-123/remix", + Some(&remix_body), + Some(&same_user_new_key), + "trace-openai-remix-owner-rotated-key", + ) + .is_some()); + let mut foreign_auth = sample_auth_context(); + foreign_auth.user_id = "user-foreign".to_string(); + foreign_auth.api_key_id = "key-foreign".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_remix_sync", + "/v1/videos/task-local-123/remix", + Some(&remix_body), + Some(&foreign_auth), + "trace-openai-remix-foreign", + ) + .is_none()); } #[test] @@ -180,6 +228,30 @@ fn rust_authoritative_service_builds_openai_delete_follow_up_plan() { .and_then(Value::as_str), Some("task-local-123") ); + + let mut same_user_new_key = sample_auth_context(); + same_user_new_key.api_key_id = "key-rotated".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_delete_sync", + "/v1/videos/task-local-123", + None, + Some(&same_user_new_key), + "trace-openai-delete-owner-rotated-key", + ) + .is_some()); + let mut foreign_auth = sample_auth_context(); + foreign_auth.user_id = "user-foreign".to_string(); + foreign_auth.api_key_id = "key-foreign".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "openai_video_delete_sync", + "/v1/videos/task-local-123", + None, + Some(&foreign_auth), + "trace-openai-delete-foreign", + ) + .is_none()); } #[test] @@ -230,6 +302,30 @@ fn rust_authoritative_service_builds_gemini_cancel_follow_up_plan() { .and_then(Value::as_str), Some("localshort123") ); + + let mut same_user_new_key = sample_auth_context(); + same_user_new_key.api_key_id = "key-rotated".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "gemini_video_cancel_sync", + "/v1beta/models/veo-3/operations/localshort123:cancel", + None, + Some(&same_user_new_key), + "trace-gemini-cancel-owner-rotated-key", + ) + .is_some()); + let mut foreign_auth = sample_auth_context(); + foreign_auth.user_id = "user-foreign".to_string(); + foreign_auth.api_key_id = "key-foreign".to_string(); + assert!(service + .prepare_follow_up_sync_plan_for_user( + "gemini_video_cancel_sync", + "/v1beta/models/veo-3/operations/localshort123:cancel", + None, + Some(&foreign_auth), + "trace-gemini-cancel-foreign", + ) + .is_none()); } #[test] @@ -367,9 +463,13 @@ fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() fn file_video_task_store_persists_snapshots_across_service_rebuilds() { let store_path = std::env::temp_dir().join(format!("aether-video-task-store-{}.json", Uuid::new_v4())); - let service = - VideoTaskService::with_file_store(VideoTaskTruthSourceMode::RustAuthoritative, &store_path) - .expect("file-backed service should build"); + let encryption_key = "video-task-test-encryption-key"; + let service = VideoTaskService::with_file_store( + VideoTaskTruthSourceMode::RustAuthoritative, + &store_path, + encryption_key, + ) + .expect("file-backed service should build"); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { local_task_id: "task-file-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), @@ -393,9 +493,16 @@ fn file_video_task_store_persists_snapshots_across_service_rebuilds() { })); drop(service); - let reopened = - VideoTaskService::with_file_store(VideoTaskTruthSourceMode::RustAuthoritative, &store_path) - .expect("reopened file-backed service should build"); + let stored_bytes = std::fs::read(&store_path).expect("encrypted store should exist"); + assert!(stored_bytes.starts_with(b"aether-video-tasks-v2\n")); + assert!(!String::from_utf8_lossy(&stored_bytes).contains("api.openai.example")); + + let reopened = VideoTaskService::with_file_store( + VideoTaskTruthSourceMode::RustAuthoritative, + &store_path, + encryption_key, + ) + .expect("reopened file-backed service should build"); let response = reopened .read_response(Some("openai"), "/v1/videos/task-file-123") .expect("persisted read response should exist"); diff --git a/apps/aether-gateway/src/video_tasks/tests/projection.rs b/apps/aether-gateway/src/video_tasks/tests/projection.rs index 5738bab0e..fa7978bfd 100644 --- a/apps/aether-gateway/src/video_tasks/tests/projection.rs +++ b/apps/aether-gateway/src/video_tasks/tests/projection.rs @@ -57,8 +57,8 @@ fn rust_authoritative_service_projects_openai_status_into_local_read_response() "progress": 100, "completed_at": 1712345688u64, "error": { - "code": "upstream_failed", - "message": "provider failed" + "code": "Bearer code-secret", + "message": "request failed at https://internal.test/result?token=message-secret" } }) .as_object() @@ -79,8 +79,14 @@ fn rust_authoritative_service_projects_openai_status_into_local_read_response() .get("error") .and_then(Value::as_object) .and_then(|value| value.get("code")), - Some(&json!("upstream_failed")) + Some(&json!("provider_error")) ); + assert_eq!( + failed.body_json["error"]["message"], + "Video generation failed" + ); + assert!(!failed.body_json.to_string().contains("code-secret")); + assert!(!failed.body_json.to_string().contains("message-secret")); } #[test] @@ -125,6 +131,28 @@ fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_vide "https://cdn.example.com/ext-video-task-123.mp4" ); assert!(plan.headers.is_empty()); + + assert!(service + .prepare_openai_content_stream_action_for_user( + "/v1/videos/task-local-123/content", + Some("variant=video"), + "trace-foreign-content-123", + "user-foreign", + ) + .is_none()); + + let owner_action = service + .prepare_openai_content_stream_action_for_user( + "/v1/videos/task-local-123/content", + Some("variant=video"), + "trace-owner-content-123", + "user-123", + ) + .expect("owner content action should exist"); + assert!(matches!( + owner_action, + LocalVideoTaskContentAction::StreamPlan(_) + )); } #[test] @@ -197,7 +225,9 @@ fn rust_authoritative_service_projects_gemini_status_into_local_read_response() json!({ "done": false, "metadata": { - "state": "PROCESSING" + "state": "PROCESSING", + "authorization": "Bearer metadata-secret", + "debug_url": "https://internal.test/status?token=query-secret" } }) .as_object() @@ -211,10 +241,9 @@ fn rust_authoritative_service_projects_gemini_status_into_local_read_response() .expect("processing read response should exist"); assert_eq!(processing.status_code, 200); assert_eq!(processing.body_json.get("done"), Some(&json!(false))); - assert_eq!( - processing.body_json.get("metadata"), - Some(&json!({"state": "PROCESSING"})) - ); + assert_eq!(processing.body_json.get("metadata"), Some(&json!({}))); + assert!(!processing.body_json.to_string().contains("metadata-secret")); + assert!(!processing.body_json.to_string().contains("query-secret")); assert!(service.project_gemini_task_response( "localshort123", diff --git a/apps/aether-gateway/src/video_tasks/tests/sync.rs b/apps/aether-gateway/src/video_tasks/tests/sync.rs index de5d11834..9634d322a 100644 --- a/apps/aether-gateway/src/video_tasks/tests/sync.rs +++ b/apps/aether-gateway/src/video_tasks/tests/sync.rs @@ -253,6 +253,13 @@ fn rust_authoritative_service_reads_openai_task_from_local_registry() { response.body_json.get("status").and_then(Value::as_str), Some("queued") ); + + assert!(service + .read_response_for_user(Some("openai"), "/v1/videos/task-local-123", "user-foreign",) + .is_none()); + assert!(service + .read_response_for_user(Some("openai"), "/v1/videos/task-local-123", "user-123",) + .is_some()); } #[test] diff --git a/apps/aether-tunnel/Cargo.toml b/apps/aether-tunnel/Cargo.toml index 5d3629588..387f7565b 100644 --- a/apps/aether-tunnel/Cargo.toml +++ b/apps/aether-tunnel/Cargo.toml @@ -13,6 +13,7 @@ aether-runtime-state.workspace = true axum.workspace = true tokio = { version = "1", features = ["full"] } reqwest.workspace = true +semver.workspace = true hyper = { version = "1", features = ["client", "http1", "http2"] } hyper-util = { version = "0.1", features = ["client", "client-legacy", "http1", "http2", "tokio"] } http-body-util = "0.1" @@ -32,7 +33,7 @@ anyhow = "1" arc-swap = "1" toml = "0.8" rustls = { version = "0.23", features = ["ring"] } -ratatui = "0.30" +ratatui = "0.30.2" crossterm = "0.28" url = "2" sysinfo = "0.32" @@ -45,4 +46,4 @@ webpki-roots = "0.26" uuid.workspace = true [dev-dependencies] -aether-gateway.workspace = true +aether-gateway = { workspace = true, features = ["testkit"] } diff --git a/apps/aether-tunnel/README.md b/apps/aether-tunnel/README.md index 6299bb611..2d2895759 100644 --- a/apps/aether-tunnel/README.md +++ b/apps/aether-tunnel/README.md @@ -68,10 +68,14 @@ irm https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/in 可选变量:`AETHER_TUNNEL_RELEASE_TAG` 固定安装某个 `tunnel-v*` tag,`AETHER_TUNNEL_CONFIG` 指定配置文件路径,`AETHER_TUNNEL_INSTALL_DIR` 指定二进制安装目录。 ```bash -# 1. 首次安装配置(TUI 向导,勾选 Install Service 随系统启动服务) -sudo ./aether-tunnel setup +# 1. 注册 root 系统服务前,先把二进制和凭据配置放入 root 管理的路径 +sudo install -o root -g root -m 0755 ./aether-tunnel /usr/local/bin/aether-tunnel +sudo install -o root -g root -m 0700 -d /etc/aether-tunnel -# 2. 日常管理 (勾选 Install Service 作为系统服务的情况下) +# 2. 首次安装配置(TUI 向导,勾选 Install Service 随系统启动服务) +sudo /usr/local/bin/aether-tunnel setup /etc/aether-tunnel/aether-tunnel.toml + +# 3. 日常管理 (勾选 Install Service 作为系统服务的情况下) aether-tunnel status # 看状态 sudo aether-tunnel logs # 看日志 @@ -79,14 +83,14 @@ sudo aether-tunnel start # 启动服务 sudo aether-tunnel stop # 停止服务 sudo aether-tunnel restart # 重启服务 -# 3. 重新配置(改完自动重启服务) -sudo aether-tunnel setup +# 4. 重新配置(改完自动重启服务) +sudo aether-tunnel setup /etc/aether-tunnel/aether-tunnel.toml -# 4. 彻底卸载 +# 5. 彻底卸载 sudo aether-tunnel uninstall ``` -完成向导后, 配置自动保存到 `aether-tunnel.toml`,如果启用了 Install Service,将自动注册并启动当前系统支持的服务(`systemd` 或 `OpenRC`)。 +完成向导后,如果启用了 Install Service,将自动注册并启动当前系统支持的服务(`systemd` 或 `OpenRC`)。为避免 root 服务被普通用户替换二进制或凭据配置,服务模式只接受 root 所有、父目录不可由组或其他用户写入的二进制,以及权限为 `0600` 的单硬链接配置文件;直接运行模式不受此限制。 ### 直接运行 @@ -96,6 +100,10 @@ sudo aether-tunnel uninstall ./aether-tunnel ``` +### 安全更新 + +Linux/macOS 可运行 `sudo aether-tunnel upgrade [version]`。自更新只接受本仓库的非草稿 `tunnel-v*` / `proxy-v*` SemVer Release,下载当前平台的固定资产和同一 tag 下的 `SHA256SUMS.txt`,校验后在受保护的二进制目录内原子替换,并保留上一版本用于失败回滚。Windows 不执行进程内自更新;请重新运行上面的 PowerShell 安装脚本完成手工替换,避免二段重命名产生二进制缺失窗口。 + ## 配置 配置按以下优先级加载(高优先级覆盖低优先级): @@ -119,7 +127,7 @@ sudo aether-tunnel uninstall | `--node-region` | `AETHER_TUNNEL_NODE_REGION` | 自动检测 | 地区标识 | | `--heartbeat-interval` | `AETHER_TUNNEL_HEARTBEAT_INTERVAL` | `5` | 心跳间隔(秒) | | `--allowed-ports` | `AETHER_TUNNEL_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 | -| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `true` | 允许 private/reserved 目标地址,通过后仍受 `allowed_ports` 限制;设为 `false` 可恢复严格拦截 | +| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `false` | 默认拦截 private/reserved 目标地址;仅在明确需要访问内网服务时设为 `true`,通过后仍受 `allowed_ports` 限制 | #### Tunnel 连接 @@ -160,7 +168,7 @@ sudo aether-tunnel uninstall | `--upstream-tcp-nodelay` | `AETHER_TUNNEL_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY | | `--upstream-proxy-url` | `AETHER_TUNNEL_UPSTREAM_PROXY_URL` | 空 | 仅 provider 上游请求使用的出口代理 | -启用 `follow_redirects` 后,307/308 请求体重放不设置累计大小上限。 +启用 `follow_redirects` 后,同源 307/308 会在请求体不超过 5 MiB 时重放。首个上游请求始终流式传输;超过重放预算时不会拒绝或截断原请求,而是将 307/308 响应原样返回给调用方。 出口代理支持 `http://`、`socks5://`、`socks5h://`。配合 WARP sidecar 时可填写: @@ -183,7 +191,7 @@ upstream_proxy_url = "socks5h://microwarp:1080" | 参数 | 环境变量 | 默认值 | 说明 | |------|----------|--------|------| -| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `true` | 默认允许 private/reserved 目标地址;设为 `false` 可恢复拦截,且仅影响重启后的进程 | +| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `false` | 默认拦截 private/reserved 目标地址;仅在明确需要访问内网服务时设为 `true`,且仅影响重启后的进程 | | `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) | | `--dns-cache-capacity` | `AETHER_TUNNEL_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) | @@ -227,15 +235,15 @@ node_name = "jp-proxy-01" tunnel_security = "off" [[servers]] -aether_url = "http://aether-2.example.com" +aether_url = "http://127.0.0.1:8084" management_token = "ae_yyy" -node_name = "jp-proxy-02" +node_name = "local-dev-proxy" tunnel_encryption_key = "base64-32-bytes" ``` -`tunnel_security = "non_tls_required"` 是非 TLS secure tunnel 的 MVP 配置面:它要求同时提供当前 `[[servers]]` 条目的 `tunnel_encryption_key`,后续握手使用 `node_name` / `X-Node-Id` 查找对应 PSK,不引入 `tunnel_encryption_key_id`。`wss://` 仍是推荐方案;`ws:// + secure tunnel` 只加密注册完成后的 WebSocket tunnel frame,不保护安装脚本、注册请求、`management_token` 或 PSK 的首次分发;这些 bootstrap 凭据仍必须通过 HTTPS 或其他可信通道交付。它不等价于 HTTPS 伪装,也不覆盖 tunnel ↔ origin/provider 这段链路。 +`tunnel_security = "non_tls_required"` 是本机非 TLS secure tunnel 的兼容配置面:它要求同时提供当前 `[[servers]]` 条目的 `tunnel_encryption_key`,后续握手使用 `node_name` / `X-Node-Id` 查找对应 PSK,不引入 `tunnel_encryption_key_id`。公网或局域网 `aether_url` 必须使用 HTTPS;`http://` 只允许字面量 `localhost`、`127.0.0.0/8` 或 `::1`,避免明文泄漏 `management_token`。secure tunnel 只加密注册完成后的 WebSocket tunnel frame,不保护注册请求或 bootstrap 凭据,因此不能替代 HTTPS。 -如果 `aether_url` 使用 `http://` 且当前 `[[servers]]` 条目提供了 `tunnel_encryption_key`,省略 `tunnel_security` 时运行时会自动按 `non_tls_required` 生效;显式配置 `tunnel_security = "off"` 会关闭该自动推断。secure tunnel 会在 WebSocket tunnel 上加密所有二进制 tunnel frame;未配置 key 或显式关闭的旧节点仍按原明文协议工作。 +如果 loopback `aether_url` 使用 `http://` 且当前 `[[servers]]` 条目提供了 `tunnel_encryption_key`,省略 `tunnel_security` 时运行时会自动按 `non_tls_required` 生效;显式配置 `tunnel_security = "off"` 会关闭该自动推断。secure tunnel 会在 WebSocket tunnel 上加密所有二进制 tunnel frame;未配置 key 或显式关闭的旧节点仍按原明文协议工作。 ## 发布新版本 diff --git a/apps/aether-tunnel/install.ps1 b/apps/aether-tunnel/install.ps1 index d79ab799c..eaf2c0aa1 100644 --- a/apps/aether-tunnel/install.ps1 +++ b/apps/aether-tunnel/install.ps1 @@ -1,9 +1,11 @@ $ErrorActionPreference = 'Stop' +Add-Type -AssemblyName System.Net.Http $Repo = if ($env:AETHER_TUNNEL_RELEASE_REPO) { $env:AETHER_TUNNEL_RELEASE_REPO } else { 'fawney19/Aether' } $ReleaseTag = $env:AETHER_TUNNEL_RELEASE_TAG $InstallDir = $env:AETHER_TUNNEL_INSTALL_DIR $ConfigPath = $env:AETHER_TUNNEL_CONFIG +$TunnelReleaseTagPattern = '^tunnel-v(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)(?:-(?:0|[1-9][0-9]*|[0-9]*[A-Za-z-][0-9A-Za-z-]*)(?:\.(?:0|[1-9][0-9]*|[0-9]*[A-Za-z-][0-9A-Za-z-]*))*)?(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?$' function Say([string]$Message) { Write-Host "[Aether Tunnel] $Message" } function Fail([string]$Message) { throw "[Aether Tunnel] $Message" } @@ -15,15 +17,126 @@ function Prompt-IfEmpty([string]$Name, [string]$Value, [string]$Prompt) { return $Read } +function Assert-SafeReleaseRepo([string]$Value) { + if ([string]::IsNullOrWhiteSpace($Value) -or + $Value.Length -gt 200 -or + $Value -cnotmatch '^[A-Za-z0-9][A-Za-z0-9._-]*/[A-Za-z0-9][A-Za-z0-9._-]*$') { + Fail 'Release repository must be a safe GitHub OWNER/REPO identifier' + } +} + +function Assert-SafeTunnelReleaseTag([string]$Value) { + if ([string]::IsNullOrWhiteSpace($Value) -or + $Value.Length -gt 128 -or + $Value -cnotmatch $TunnelReleaseTagPattern) { + Fail 'Release tag must use tunnel-v followed by a valid semantic version' + } +} + +function Assert-TrustedGithubUri([Uri]$Uri) { + if (-not $Uri.IsAbsoluteUri -or + $Uri.Scheme -cne 'https' -or + -not [string]::IsNullOrEmpty($Uri.UserInfo) -or + -not [string]::IsNullOrEmpty($Uri.Fragment) -or + $Uri.Port -ne 443) { + Fail "GitHub downloads must use credential-free HTTPS on port 443: $Uri" + } + $HostName = $Uri.IdnHost.ToLowerInvariant() + $TrustedHost = $HostName -in @('api.github.com', 'github.com', 'objects.githubusercontent.com', 'release-assets.githubusercontent.com') -or + $HostName.EndsWith('.objects.githubusercontent.com', [StringComparison]::Ordinal) -or + $HostName.EndsWith('.release-assets.githubusercontent.com', [StringComparison]::Ordinal) + if (-not $TrustedHost) { Fail "GitHub download redirected to an untrusted host: $HostName" } +} + +function Get-TrustedGithubBytes([string]$UriText) { + $CurrentUri = [Uri]::new($UriText, [UriKind]::Absolute) + Assert-TrustedGithubUri $CurrentUri + $Handler = [System.Net.Http.HttpClientHandler]::new() + $Handler.AllowAutoRedirect = $false + $Client = [System.Net.Http.HttpClient]::new($Handler, $true) + $Client.DefaultRequestHeaders.UserAgent.ParseAdd('aether-tunnel-installer') + try { + for ($RedirectCount = 0; $RedirectCount -le 10; $RedirectCount++) { + $Response = $null + try { + $Response = $Client.GetAsync( + $CurrentUri, + [System.Net.Http.HttpCompletionOption]::ResponseHeadersRead + ).GetAwaiter().GetResult() + $StatusCode = [int]$Response.StatusCode + if ($StatusCode -in @(301, 302, 303, 307, 308)) { + if ($RedirectCount -eq 10) { Fail 'GitHub download redirected too many times' } + $Location = $Response.Headers.Location + if (-not $Location) { Fail 'GitHub redirect is missing the Location header' } + $CurrentUri = if ($Location.IsAbsoluteUri) { + $Location + } else { + [Uri]::new($CurrentUri, $Location) + } + Assert-TrustedGithubUri $CurrentUri + continue + } + if (-not $Response.IsSuccessStatusCode) { + Fail "GitHub download returned HTTP $StatusCode for $CurrentUri" + } + $Bytes = $Response.Content.ReadAsByteArrayAsync().GetAwaiter().GetResult() + return ,$Bytes + } finally { + if ($Response) { $Response.Dispose() } + } + } + Fail 'GitHub download redirected too many times' + } finally { + $Client.Dispose() + } +} + +function Read-TrustedGithubJson([string]$Uri) { + $Bytes = Get-TrustedGithubBytes $Uri + $Text = [Text.UTF8Encoding]::new($false, $true).GetString($Bytes) + return ($Text | ConvertFrom-Json) +} + +function Save-TrustedGithubFile([string]$Uri, [string]$Path) { + $Bytes = Get-TrustedGithubBytes $Uri + $Stream = [IO.File]::Open($Path, [IO.FileMode]::CreateNew, [IO.FileAccess]::Write, [IO.FileShare]::None) + try { + $Stream.Write($Bytes, 0, $Bytes.Length) + $Stream.Flush($true) + } finally { + $Stream.Dispose() + } +} + +function Assert-SafeNodeName([string]$Value) { + if ([string]::IsNullOrWhiteSpace($Value) -or + $Value.Length -gt 255 -or + $Value -ne $Value.Trim() -or + $Value -match '[\x00-\x1F\x7F]') { + Fail 'Node name must be 1 to 255 characters without surrounding whitespace or control characters' + } +} + function ConvertTo-TomlQuotedString([string]$Value) { return ($Value | ConvertTo-Json -Compress) } function Resolve-LatestTunnelTag { - if (-not [string]::IsNullOrWhiteSpace($ReleaseTag)) { return $ReleaseTag } + Assert-SafeReleaseRepo $Repo + if (-not [string]::IsNullOrWhiteSpace($ReleaseTag)) { + Assert-SafeTunnelReleaseTag $ReleaseTag + $RequestedUri = "https://api.github.com/repos/$Repo/releases/tags/$ReleaseTag" + $RequestedRelease = Read-TrustedGithubJson $RequestedUri + if ($RequestedRelease.draft -or ([string]$RequestedRelease.tag_name -cne $ReleaseTag)) { + Fail 'GitHub returned a draft or mismatched tunnel release' + } + return $ReleaseTag + } $Uri = "https://api.github.com/repos/$Repo/releases?per_page=100" - $Releases = Invoke-RestMethod -Uri $Uri -Headers @{ 'User-Agent' = 'aether-tunnel-installer' } - $TunnelReleases = @($Releases | Where-Object { -not $_.draft -and $_.tag_name -like 'tunnel-v*' } | Sort-Object published_at -Descending) + $Releases = Read-TrustedGithubJson $Uri + $TunnelReleases = @($Releases | Where-Object { + -not $_.draft -and -not $_.prerelease -and ([string]$_.tag_name -cmatch $TunnelReleaseTagPattern) + } | Sort-Object published_at -Descending) if ($TunnelReleases.Count -eq 0) { Fail "No tunnel-v* release found in $Repo" } return $TunnelReleases[0].tag_name } @@ -34,6 +147,159 @@ function Test-IsAdministrator { return $Principal.IsInRole([Security.Principal.WindowsBuiltInRole]::Administrator) } +function Protect-SensitiveConfigFile([string]$Path) { + if (-not (Test-Path -LiteralPath $Path -PathType Leaf)) { + Fail "Sensitive config file does not exist: $Path" + } + Assert-NotReparsePoint $Path 'Sensitive config file' + + $Identity = [Security.Principal.WindowsIdentity]::GetCurrent() + $CurrentSid = $Identity.User + if (-not $CurrentSid) { Fail 'Unable to resolve the current Windows account SID' } + + $Acl = [System.Security.AccessControl.FileSecurity]::new() + $Acl.SetAccessRuleProtection($true, $false) + $Acl.SetOwner($CurrentSid) + + $AllowedSids = @( + $CurrentSid.Value, + 'S-1-5-18', + 'S-1-5-32-544' + ) | Select-Object -Unique + foreach ($SidValue in $AllowedSids) { + $Sid = [Security.Principal.SecurityIdentifier]::new($SidValue) + $Rule = [System.Security.AccessControl.FileSystemAccessRule]::new( + $Sid, + [System.Security.AccessControl.FileSystemRights]::FullControl, + [System.Security.AccessControl.AccessControlType]::Allow + ) + $Acl.AddAccessRule($Rule) | Out-Null + } + + Set-Acl -LiteralPath $Path -AclObject $Acl +} + +function Protect-SensitiveConfigDirectory([string]$Path) { + if (-not (Test-Path -LiteralPath $Path -PathType Container)) { + Fail "Sensitive config directory does not exist: $Path" + } + Assert-NotReparsePoint $Path 'Sensitive config directory' + + $Identity = [Security.Principal.WindowsIdentity]::GetCurrent() + $CurrentSid = $Identity.User + if (-not $CurrentSid) { Fail 'Unable to resolve the current Windows account SID' } + + $Acl = [System.Security.AccessControl.DirectorySecurity]::new() + $Acl.SetAccessRuleProtection($true, $false) + $Acl.SetOwner($CurrentSid) + $AllowedSids = @($CurrentSid.Value, 'S-1-5-18', 'S-1-5-32-544') | Select-Object -Unique + foreach ($SidValue in $AllowedSids) { + $Sid = [Security.Principal.SecurityIdentifier]::new($SidValue) + $Rule = [System.Security.AccessControl.FileSystemAccessRule]::new( + $Sid, + [System.Security.AccessControl.FileSystemRights]::FullControl, + [System.Security.AccessControl.InheritanceFlags]'ContainerInherit, ObjectInherit', + [System.Security.AccessControl.PropagationFlags]::None, + [System.Security.AccessControl.AccessControlType]::Allow + ) + $Acl.AddAccessRule($Rule) | Out-Null + } + Set-Acl -LiteralPath $Path -AclObject $Acl +} + +function Protect-SensitiveConfigArtifacts([string]$Path) { + if (Test-Path -LiteralPath $Path -PathType Leaf) { + Protect-SensitiveConfigFile $Path + } + + $Directory = Split-Path -Parent $Path + $Leaf = Split-Path -Leaf $Path + if (-not (Test-Path -LiteralPath $Directory -PathType Container)) { return } + foreach ($Backup in Get-ChildItem -LiteralPath $Directory -File) { + if ($Backup.Name.StartsWith("$Leaf.bak.", [StringComparison]::OrdinalIgnoreCase)) { + Protect-SensitiveConfigFile $Backup.FullName + } + } +} + +function Assert-NotReparsePoint([string]$Path, [string]$Description) { + if (-not (Test-Path -LiteralPath $Path)) { return } + $Item = Get-Item -LiteralPath $Path -Force + if (($Item.Attributes -band [IO.FileAttributes]::ReparsePoint) -ne 0) { + Fail "$Description must not be a reparse point or symbolic link: $Path" + } +} + +function Assert-NoReparsePointAncestors([string]$Path, [string]$Description) { + $Current = [IO.Path]::GetFullPath($Path) + while (-not [string]::IsNullOrEmpty($Current)) { + if (Test-Path -LiteralPath $Current) { + Assert-NotReparsePoint $Current "$Description ancestor" + } + $Parent = Split-Path -Parent $Current + if ([string]::IsNullOrEmpty($Parent) -or $Parent -eq $Current) { break } + $Current = $Parent + } +} + +function Assert-NotHardLink([string]$Path, [string]$Description) { + if (-not (Test-Path -LiteralPath $Path -PathType Leaf)) { return } + $Item = Get-Item -LiteralPath $Path -Force + $LinkTypeProperty = $Item.PSObject.Properties['LinkType'] + if (-not $LinkTypeProperty) { + Fail "Unable to verify hard-link safety for $Description: $Path" + } + if ([string]$Item.LinkType -eq 'HardLink') { + Fail "$Description must not be a hard link: $Path" + } +} + +function Initialize-SecureConfigPath([string]$Path) { + $Directory = Split-Path -Parent $Path + Assert-NoReparsePointAncestors $Directory 'Config directory' + [IO.Directory]::CreateDirectory($Directory) | Out-Null + Assert-NotReparsePoint $Directory 'Config directory' + Protect-SensitiveConfigDirectory $Directory + Assert-NotReparsePoint $Path 'Config file' + if (Test-Path -LiteralPath $Path) { + if (-not (Test-Path -LiteralPath $Path -PathType Leaf)) { + Fail "Config path is not a regular file: $Path" + } + Assert-NotHardLink $Path 'Config file' + Protect-SensitiveConfigFile $Path + } +} + +function Write-SensitiveUtf8File([string]$Path, [string]$Content) { + $CreateStream = [IO.File]::Open( + $Path, + [IO.FileMode]::CreateNew, + [IO.FileAccess]::Write, + [IO.FileShare]::None + ) + $CreateStream.Dispose() + Protect-SensitiveConfigFile $Path + + $Stream = [IO.File]::Open( + $Path, + [IO.FileMode]::Open, + [IO.FileAccess]::Write, + [IO.FileShare]::None + ) + try { + $Encoding = [Text.UTF8Encoding]::new($false) + $Writer = [IO.StreamWriter]::new($Stream, $Encoding) + try { + $Writer.Write($Content) + $Writer.Flush() + } finally { + $Writer.Dispose() + } + } finally { + $Stream.Dispose() + } +} + function Initialize-Paths { if ([string]::IsNullOrWhiteSpace($script:InstallDir)) { if (Test-IsAdministrator) { @@ -49,9 +315,83 @@ function Initialize-Paths { $script:ConfigPath = Join-Path $env:APPDATA 'AetherTunnel\aether-tunnel.toml' } } + $script:InstallDir = [IO.Path]::GetFullPath($script:InstallDir) + $script:ConfigPath = [IO.Path]::GetFullPath($script:ConfigPath) +} + +function Install-VerifiedTunnelBinary([string]$SourceBinary) { + Assert-NoReparsePointAncestors $script:InstallDir 'Install directory' + [IO.Directory]::CreateDirectory($script:InstallDir) | Out-Null + Assert-NotReparsePoint $script:InstallDir 'Install directory' + + $TargetBinary = Join-Path $script:InstallDir 'aether-tunnel.exe' + Assert-NotReparsePoint $TargetBinary 'Install target' + if ((Test-Path -LiteralPath $TargetBinary) -and + -not (Test-Path -LiteralPath $TargetBinary -PathType Leaf)) { + Fail "Install target is not a regular file: $TargetBinary" + } + Assert-NotHardLink $TargetBinary 'Install target' + + $TempBinary = Join-Path $script:InstallDir ('.aether-tunnel.tmp.' + [Guid]::NewGuid().ToString('N')) + try { + [IO.File]::Copy($SourceBinary, $TempBinary, $false) + Assert-NotReparsePoint $TargetBinary 'Install target' + Assert-NotHardLink $TargetBinary 'Install target' + if (Test-Path -LiteralPath $TargetBinary -PathType Leaf) { + [IO.File]::Replace($TempBinary, $TargetBinary, $null, $true) + } elseif (Test-Path -LiteralPath $TargetBinary) { + Fail "Install target changed to a non-file during installation: $TargetBinary" + } else { + [IO.File]::Move($TempBinary, $TargetBinary) + } + } finally { + if (Test-Path -LiteralPath $TempBinary) { + Remove-Item -LiteralPath $TempBinary -Force + } + } +} + +function Expand-VerifiedTunnelArchive([string]$Archive, [string]$Destination) { + Add-Type -AssemblyName System.IO.Compression.FileSystem + $Zip = [System.IO.Compression.ZipFile]::OpenRead($Archive) + try { + $Entries = @($Zip.Entries) + if ($Entries.Count -ne 1) { Fail 'Release archive must contain exactly one file' } + + $Entry = $Entries[0] + if (($Entry.FullName -ne 'aether-tunnel.exe') -or ($Entry.Name -ne 'aether-tunnel.exe')) { + Fail 'Release archive must contain only aether-tunnel.exe at its root' + } + $UnixType = (($Entry.ExternalAttributes -shr 16) -band 0xF000) + if (($UnixType -ne 0) -and ($UnixType -ne 0x8000)) { + Fail 'aether-tunnel.exe in release archive is not a regular file' + } + if ($Entry.Length -le 0) { Fail 'aether-tunnel.exe in release archive is empty' } + + $InputStream = $Entry.Open() + try { + $OutputStream = [System.IO.File]::Open( + $Destination, + [System.IO.FileMode]::CreateNew, + [System.IO.FileAccess]::Write, + [System.IO.FileShare]::None + ) + try { + $InputStream.CopyTo($OutputStream) + } finally { + $OutputStream.Dispose() + } + } finally { + $InputStream.Dispose() + } + } finally { + $Zip.Dispose() + } } function Install-AetherTunnelBinary([string]$Tag, [string]$TempDir) { + Assert-SafeReleaseRepo $Repo + Assert-SafeTunnelReleaseTag $Tag if (-not [Environment]::Is64BitOperatingSystem) { Fail 'Windows release currently supports amd64 only' } $Asset = 'aether-tunnel-windows-amd64.zip' $Base = "https://github.com/$Repo/releases/download/$Tag" @@ -59,30 +399,36 @@ function Install-AetherTunnelBinary([string]$Tag, [string]$TempDir) { $Sums = Join-Path $TempDir 'SHA256SUMS.txt' Say "Downloading $Tag / $Asset" - Invoke-WebRequest -Uri "$Base/$Asset" -OutFile $Archive - try { Invoke-WebRequest -Uri "$Base/SHA256SUMS.txt" -OutFile $Sums } catch { $Sums = $null } + Save-TrustedGithubFile "$Base/$Asset" $Archive + Save-TrustedGithubFile "$Base/SHA256SUMS.txt" $Sums + if (-not (Test-Path -LiteralPath $Sums -PathType Leaf)) { Fail 'SHA256SUMS.txt download is missing' } - if ($Sums -and (Test-Path $Sums)) { - $ExpectedLine = Get-Content $Sums | Where-Object { $_ -match "\s$([regex]::Escape($Asset))$" } | Select-Object -First 1 - if ($ExpectedLine) { - $Expected = ($ExpectedLine -split '\s+')[0] - $Actual = (Get-FileHash -Algorithm SHA256 $Archive).Hash.ToLowerInvariant() - if ($Actual -ne $Expected.ToLowerInvariant()) { Fail "SHA256 verification failed for $Asset" } - } + $EscapedAsset = [regex]::Escape($Asset) + $ChecksumTargetPattern = '(?:^|\s)\*?' + $EscapedAsset + '(?:\s|$)' + $ExpectedLines = @(Get-Content -LiteralPath $Sums | Where-Object { $_ -match $ChecksumTargetPattern }) + if ($ExpectedLines.Count -ne 1) { + Fail "SHA256SUMS.txt must contain exactly one entry for $Asset" } + $ChecksumPattern = '^\s*(?\S+)\s+\*?' + $EscapedAsset + '\s*$' + $ExpectedMatch = [regex]::Match([string]$ExpectedLines[0], $ChecksumPattern) + if (-not $ExpectedMatch.Success) { Fail "SHA256SUMS.txt has an invalid entry for $Asset" } + $Expected = $ExpectedMatch.Groups['hash'].Value + if ($Expected -notmatch '^[0-9A-Fa-f]{64}$') { Fail "SHA256SUMS.txt has an invalid hash for $Asset" } + + $Actual = (Get-FileHash -Algorithm SHA256 -LiteralPath $Archive).Hash.ToLowerInvariant() + if ($Actual -ne $Expected.ToLowerInvariant()) { Fail "SHA256 verification failed for $Asset" } $ExtractDir = Join-Path $TempDir 'extract' - Expand-Archive -Path $Archive -DestinationPath $ExtractDir -Force + New-Item -ItemType Directory -Force -Path $ExtractDir | Out-Null $Binary = Join-Path $ExtractDir 'aether-tunnel.exe' - if (-not (Test-Path $Binary)) { Fail 'aether-tunnel.exe not found in release asset' } - New-Item -ItemType Directory -Force -Path $script:InstallDir | Out-Null - Copy-Item $Binary (Join-Path $script:InstallDir 'aether-tunnel.exe') -Force + Expand-VerifiedTunnelArchive $Archive $Binary + Install-VerifiedTunnelBinary $Binary Say "Installed binary: $(Join-Path $script:InstallDir 'aether-tunnel.exe')" } function Test-LegacySingleServerConfig([string]$Path) { - if (-not (Test-Path $Path)) { return $false } - foreach ($Line in Get-Content $Path) { + if (-not (Test-Path -LiteralPath $Path)) { return $false } + foreach ($Line in Get-Content -LiteralPath $Path) { if ($Line -match '^\s*\[') { return $false } if ($Line -match '^\s*(aether_url|management_token)\s*=') { return $true } } @@ -90,10 +436,10 @@ function Test-LegacySingleServerConfig([string]$Path) { } function Test-ServerExists([string]$Path, [string]$QuotedUrl, [string]$QuotedName) { - if (-not (Test-Path $Path)) { return $false } + if (-not (Test-Path -LiteralPath $Path)) { return $false } $FoundUrl = $false $FoundName = $false - foreach ($Line in Get-Content $Path) { + foreach ($Line in Get-Content -LiteralPath $Path) { if ($Line -match '^\s*\[\[servers\]\]\s*$') { if ($FoundUrl -and $FoundName) { return $true } $FoundUrl = $false @@ -106,8 +452,8 @@ function Test-ServerExists([string]$Path, [string]$QuotedUrl, [string]$QuotedNam } function Add-ServerConfig([string]$AetherUrl, [string]$ManagementToken, [string]$NodeName, [string]$TunnelSecurity, [string]$TunnelEncryptionKey) { - $ConfigDir = Split-Path -Parent $script:ConfigPath - New-Item -ItemType Directory -Force -Path $ConfigDir | Out-Null + Assert-SafeNodeName $NodeName + Initialize-SecureConfigPath $script:ConfigPath if (Test-LegacySingleServerConfig $script:ConfigPath) { Fail "Existing config uses removed top-level aether_url/management_token. Run aether-tunnel setup to migrate to [[servers]] first: $script:ConfigPath" @@ -123,11 +469,18 @@ function Add-ServerConfig([string]$AetherUrl, [string]$ManagementToken, [string] return } - if (Test-Path $script:ConfigPath) { - Copy-Item $script:ConfigPath "$script:ConfigPath.bak.$(Get-Date -Format yyyyMMddHHmmss)" -Force + $ConfigExists = Test-Path -LiteralPath $script:ConfigPath -PathType Leaf + $ExistingContent = if ($ConfigExists) { + [IO.File]::ReadAllText($script:ConfigPath) + } else { + '' + } + if ($ConfigExists) { + $BackupPath = "$script:ConfigPath.bak.$(Get-Date -Format yyyyMMddHHmmss).$([Guid]::NewGuid().ToString('N'))" + Write-SensitiveUtf8File $BackupPath $ExistingContent } - $Prefix = if ((Test-Path $script:ConfigPath) -and ((Get-Item $script:ConfigPath).Length -gt 0)) { "`n" } else { '' } + $Prefix = if ($ExistingContent.Length -gt 0) { "`n" } else { '' } $Block = @( "$Prefix# Added by Aether Tunnel one-click installer. Existing config is preserved.", '[[servers]]', @@ -142,15 +495,35 @@ function Add-ServerConfig([string]$AetherUrl, [string]$ManagementToken, [string] if ($TunnelEncryptionKey) { $Block += "`ntunnel_encryption_key = $QuotedTunnelEncryptionKey" } - Add-Content -Path $script:ConfigPath -Value ($Block + "`n") -Encoding UTF8 + $ConfigDir = Split-Path -Parent $script:ConfigPath + $TempPath = Join-Path $ConfigDir ("." + (Split-Path -Leaf $script:ConfigPath) + ".tmp." + [Guid]::NewGuid().ToString('N')) + try { + Write-SensitiveUtf8File $TempPath ($ExistingContent + $Block + "`n") + Assert-NotReparsePoint $script:ConfigPath 'Config file' + if ($ConfigExists) { + if (-not (Test-Path -LiteralPath $script:ConfigPath -PathType Leaf)) { + Fail "Config file changed while it was being updated: $script:ConfigPath" + } + [IO.File]::Replace($TempPath, $script:ConfigPath, $null, $true) + } else { + [IO.File]::Move($TempPath, $script:ConfigPath) + } + } finally { + if (Test-Path -LiteralPath $TempPath) { + Remove-Item -LiteralPath $TempPath -Force + } + } + Protect-SensitiveConfigArtifacts $script:ConfigPath Say "Appended [[servers]] to: $script:ConfigPath" } function Main { + Assert-SafeReleaseRepo $Repo Initialize-Paths $AetherUrl = Prompt-IfEmpty 'AETHER_TUNNEL_AETHER_URL' $env:AETHER_TUNNEL_AETHER_URL 'Aether URL' $ManagementToken = Prompt-IfEmpty 'AETHER_TUNNEL_MANAGEMENT_TOKEN' $env:AETHER_TUNNEL_MANAGEMENT_TOKEN 'Management token (ae_xxx)' $NodeName = Prompt-IfEmpty 'AETHER_TUNNEL_NODE_NAME' $env:AETHER_TUNNEL_NODE_NAME 'Node name' + Assert-SafeNodeName $NodeName $TunnelSecurity = if ($env:AETHER_TUNNEL_SECURITY) { $env:AETHER_TUNNEL_SECURITY } else { '' } $TunnelEncryptionKey = if ($env:AETHER_TUNNEL_ENCRYPTION_KEY) { $env:AETHER_TUNNEL_ENCRYPTION_KEY } else { '' } if ($TunnelSecurity -and ($TunnelSecurity -notin @('off', 'non_tls_required'))) { @@ -161,13 +534,21 @@ function Main { } $TempDir = Join-Path ([IO.Path]::GetTempPath()) ("aether-tunnel-" + [Guid]::NewGuid().ToString('N')) - New-Item -ItemType Directory -Force -Path $TempDir | Out-Null + Assert-NoReparsePointAncestors (Split-Path -Parent $TempDir) 'Temporary directory' + if (Test-Path -LiteralPath $TempDir) { Fail "Secure temporary path already exists: $TempDir" } + [IO.Directory]::CreateDirectory($TempDir) | Out-Null + Assert-NotReparsePoint $TempDir 'Temporary directory' + Protect-SensitiveConfigDirectory $TempDir try { $Tag = Resolve-LatestTunnelTag + Assert-SafeTunnelReleaseTag $Tag Install-AetherTunnelBinary $Tag $TempDir Add-ServerConfig $AetherUrl $ManagementToken $NodeName $TunnelSecurity $TunnelEncryptionKey } finally { - Remove-Item -Recurse -Force $TempDir -ErrorAction SilentlyContinue + if (Test-Path -LiteralPath $TempDir) { + Assert-NotReparsePoint $TempDir 'Temporary directory' + Remove-Item -Recurse -Force -LiteralPath $TempDir -ErrorAction SilentlyContinue + } } Say 'Complete. Start or configure the node with:' diff --git a/apps/aether-tunnel/install.sh b/apps/aether-tunnel/install.sh index e0cc62d76..9419406f8 100755 --- a/apps/aether-tunnel/install.sh +++ b/apps/aether-tunnel/install.sh @@ -1,16 +1,21 @@ #!/bin/sh set -eu +umask 077 REPO="${AETHER_TUNNEL_RELEASE_REPO:-fawney19/Aether}" TAG="${AETHER_TUNNEL_RELEASE_TAG:-}" INSTALL_DIR="${AETHER_TUNNEL_INSTALL_DIR:-}" CONFIG_PATH="${AETHER_TUNNEL_CONFIG:-}" TMP_DIR="" +CONFIG_TMP_PATH="" say() { printf '%s\n' "[Aether Tunnel] $1"; } fail() { printf '%s\n' "[Aether Tunnel] $1" >&2; exit 1; } cleanup() { + if [ -n "$CONFIG_TMP_PATH" ] && [ -f "$CONFIG_TMP_PATH" ] && [ ! -L "$CONFIG_TMP_PATH" ]; then + rm -f "$CONFIG_TMP_PATH" + fi if [ -n "$TMP_DIR" ] && [ -d "$TMP_DIR" ]; then rm -rf "$TMP_DIR" fi @@ -21,15 +26,166 @@ need_cmd() { command -v "$1" >/dev/null 2>&1 || fail "缺少命令:$1" } +validate_https_download_url() { + url="$1" + case "$url" in + https://*) ;; + *) fail "远程下载必须使用绝对 HTTPS URL" ;; + esac + case "$url" in + *'#'*) fail "远程下载 URL 不得包含 fragment" ;; + esac + authority=${url#https://} + authority=${authority%%[/?]*} + [ -n "$authority" ] || fail "远程下载 URL 的 host 不能为空" + case "$authority" in + *'@'*) fail "远程下载 URL 不得包含凭据" ;; + esac +} + +validate_trusted_github_download_url() { + url="$1" + validate_https_download_url "$url" + authority=${url#https://} + authority=${authority%%[/?]*} + case "$authority" in + *:*) fail "GitHub 下载 URL 不得使用非标准端口:$authority" ;; + esac + host=$(printf '%s' "$authority" | tr 'A-Z' 'a-z') + case "$host" in + api.github.com|github.com|objects.githubusercontent.com|*.objects.githubusercontent.com|release-assets.githubusercontent.com|*.release-assets.githubusercontent.com) ;; + *) fail "GitHub 下载重定向到了不受信任的主机:$host" ;; + esac +} + download() { url="$1" out="$2" - if command -v curl >/dev/null 2>&1; then - curl -fL --retry 3 --connect-timeout 10 -o "$out" "$url" - elif command -v wget >/dev/null 2>&1; then - wget -O "$out" "$url" - else - fail "需要 curl 或 wget 下载 release 制品" + validate_trusted_github_download_url "$url" + command -v curl >/dev/null 2>&1 || fail "需要 curl 安全下载 GitHub release 制品" + + download_dir=$(dirname "$out") + [ -d "$download_dir" ] && [ ! -L "$download_dir" ] \ + || fail "下载目标目录不是安全目录:$download_dir" + current_url="$url" + redirect_count=0 + while :; do + response_tmp=$(mktemp "$download_dir/.aether-download.body.XXXXXXXX") \ + || fail "无法创建安全下载临时文件" + headers_tmp=$(mktemp "$download_dir/.aether-download.headers.XXXXXXXX") || { + rm -f "$response_tmp" + fail "无法创建安全响应头临时文件" + } + if ! status=$(curl -sS --retry 3 --connect-timeout 10 \ + --proto '=https' --max-redirs 0 --dump-header "$headers_tmp" \ + --output "$response_tmp" --write-out '%{http_code}' "$current_url"); then + rm -f "$response_tmp" "$headers_tmp" + fail "GitHub 下载失败:$current_url" + fi + case "$status" in + 2??) + rm -f "$headers_tmp" + [ -f "$response_tmp" ] && [ ! -L "$response_tmp" ] \ + || fail "下载结果不是普通文件" + mv -f "$response_tmp" "$out" || { + rm -f "$response_tmp" + fail "无法原子保存下载结果:$out" + } + return + ;; + 301|302|303|307|308) + location=$(awk ' + tolower(substr($0, 1, 9)) == "location:" { + value = substr($0, 10) + sub(/^[[:space:]]*/, "", value) + sub(/\r$/, "", value) + } + END { print value } + ' "$headers_tmp") + rm -f "$response_tmp" "$headers_tmp" + [ -n "$location" ] || fail "GitHub 重定向缺少 Location 响应头" + validate_trusted_github_download_url "$location" + redirect_count=$((redirect_count + 1)) + [ "$redirect_count" -le 10 ] || fail "GitHub 下载重定向次数过多" + current_url="$location" + ;; + *) + rm -f "$response_tmp" "$headers_tmp" + fail "GitHub 下载返回 HTTP $status:$current_url" + ;; + esac + done +} + +validate_release_repo() { + value="$1" + [ -n "$value" ] && [ "${#value}" -le 200 ] \ + || fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" + case "$value" in + */*) ;; + *) fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" ;; + esac + owner=${value%%/*} + name=${value#*/} + case "$name" in + */*) fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" ;; + esac + case "$owner" in + [A-Za-z0-9]*) ;; + *) fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" ;; + esac + case "$name" in + [A-Za-z0-9]*) ;; + *) fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" ;; + esac + case "$owner$name" in + *[!A-Za-z0-9._-]*) fail "release 仓库必须是安全的 GitHub OWNER/REPO 标识符" ;; + esac +} + +validate_tunnel_release_tag() { + value="$1" + [ -n "$value" ] && [ "${#value}" -le 128 ] \ + || fail "release tag 包含不安全的 URL 或路径字符" + case "$value" in + *[!A-Za-z0-9._+-]*) fail "release tag 包含不安全的 URL 或路径字符" ;; + esac + version=${value#tunnel-v} + [ "$version" != "$value" ] || fail "release tag 必须使用 tunnel-v 格式" + semver_identifier='(0|[1-9][0-9]*|[0-9]*[A-Za-z-][0-9A-Za-z-]*)' + if ! printf '%s' "$version" | LC_ALL=C grep -Eq "^(0|[1-9][0-9]*)\\.(0|[1-9][0-9]*)\\.(0|[1-9][0-9]*)(-(${semver_identifier})(\\.${semver_identifier})*)?(\\+[0-9A-Za-z-]+(\\.[0-9A-Za-z-]+)*)?$"; then + fail "release tag 必须包含有效的 SemVer 版本" + fi +} + +validate_release_asset_name() { + value="$1" + [ -n "$value" ] && [ "${#value}" -le 200 ] \ + || fail "release 制品名称无效" + case "$value" in + aether-tunnel-*.tar.gz) ;; + *) fail "release 制品名称无效" ;; + esac + case "$value" in + *[!A-Za-z0-9._-]*) fail "release 制品名称包含不安全的路径字符" ;; + esac +} + +validate_node_name() { + value="$1" + [ -n "$value" ] && [ "${#value}" -le 255 ] \ + || fail "node name 必须是 1 到 255 个字符" + newline=' +' + carriage_return=$(printf '\r') + case "$value" in + *"$newline"*|*"$carriage_return"*) fail "node name 不得包含控制字符" ;; + esac + case "$value" in + [[:space:]]*|*[[:space:]]) fail "node name 不得包含首尾空白" ;; + esac + if printf '%s' "$value" | LC_ALL=C grep -q '[[:cntrl:]]'; then + fail "node name 不得包含控制字符" fi } @@ -54,29 +210,52 @@ prompt_if_empty() { toml_quote() { value="$1" if command -v python3 >/dev/null 2>&1; then - python3 -c 'import json,sys; print(json.dumps(sys.argv[1], ensure_ascii=False))' "$value" - else - escaped=$(printf '%s' "$value" | sed 's/\\/\\\\/g; s/"/\\"/g') - printf '"%s"\n' "$escaped" + quoted=$(python3 -c 'import json,sys; print(json.dumps(sys.argv[1], ensure_ascii=False))' "$value" 2>/dev/null || true) + if [ -n "$quoted" ]; then + printf '%s\n' "$quoted" + return + fi fi + printf '%s' "$value" | awk ' + BEGIN { printf "\"" } + { + if (NR > 1) printf "\\n" + gsub(/\\/, "\\\\") + gsub(/\"/, "\\\"") + gsub(/\t/, "\\t") + gsub(/\r/, "\\r") + printf "%s", $0 + } + END { print "\"" } + ' } resolve_latest_tunnel_tag() { - [ -n "$TAG" ] && { printf '%s\n' "$TAG"; return; } + validate_release_repo "$REPO" + if [ -n "$TAG" ]; then + validate_tunnel_release_tag "$TAG" + printf '%s\n' "$TAG" + return + fi api_url="https://api.github.com/repos/${REPO}/releases?per_page=100" releases="$TMP_DIR/releases.json" download "$api_url" "$releases" >/dev/null 2>&1 || fail "无法读取 GitHub Releases:$api_url" if command -v python3 >/dev/null 2>&1; then python3 - "$releases" <<'PY' -import json, sys +import json, re, sys releases = json.load(open(sys.argv[1], encoding='utf-8')) -tunnel = [r for r in releases if not r.get('draft') and str(r.get('tag_name', '')).startswith('tunnel-v')] +tag_pattern = re.compile(r'^tunnel-v(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)(?:-(?:0|[1-9][0-9]*|[0-9]*[A-Za-z-][0-9A-Za-z-]*)(?:\.(?:0|[1-9][0-9]*|[0-9]*[A-Za-z-][0-9A-Za-z-]*))*)?(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?$') +tunnel = [r for r in releases if not r.get('draft') and not r.get('prerelease') and tag_pattern.fullmatch(str(r.get('tag_name', '')))] tunnel.sort(key=lambda r: r.get('published_at') or r.get('created_at') or '', reverse=True) if tunnel: print(tunnel[0]['tag_name']) PY else - grep -o '"tag_name"[[:space:]]*:[[:space:]]*"tunnel-v[^"]*"' "$releases" | head -n 1 | sed 's/.*"\(tunnel-v[^"]*\)".*/\1/' + # Without a JSON parser, only accept an exact stable SemVer tag. This + # keeps the fallback from accidentally selecting a prerelease entry. + grep -Eo '"tag_name"[[:space:]]*:[[:space:]]*"tunnel-v(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)"' "$releases" | + head -n 1 | + sed 's/.*"\(tunnel-v[^"]*\)".*/\1/' fi } @@ -121,39 +300,237 @@ choose_paths() { fi } +stat_owner_id() { + stat -c '%u' "$1" 2>/dev/null || stat -f '%u' "$1" 2>/dev/null +} + +stat_mode() { + stat -c '%a' "$1" 2>/dev/null || stat -f '%Lp' "$1" 2>/dev/null +} + +stat_link_count() { + stat -c '%h' "$1" 2>/dev/null || stat -f '%l' "$1" 2>/dev/null +} + +require_single_link_regular_file() { + path="$1" + description="$2" + [ -f "$path" ] && [ ! -L "$path" ] || fail "$description 不是普通文件:$path" + link_count=$(stat_link_count "$path") || fail "无法读取 $description 的硬链接计数:$path" + [ "$link_count" = "1" ] || fail "$description 不得是多重硬链接文件:$path" +} + +validate_trusted_directory_ancestors() { + path="$1" + description="$2" + current="$path" + while :; do + [ ! -L "$current" ] || fail "$description 的祖先目录不得是符号链接:$current" + if [ -e "$current" ]; then + [ -d "$current" ] || fail "$description 的祖先路径不是目录:$current" + ancestor_owner=$(stat_owner_id "$current") \ + || fail "无法读取 $description 的祖先目录所有者:$current" + current_uid=$(id -u) + [ "$ancestor_owner" = "0" ] || [ "$ancestor_owner" = "$current_uid" ] \ + || fail "$description 的祖先目录必须属于 root 或当前用户:$current" + ancestor_mode=$(stat_mode "$current") \ + || fail "无法读取 $description 的祖先目录权限:$current" + case "$ancestor_mode" in + *[!0-7]*) fail "$description 的祖先目录权限格式无效:$ancestor_mode" ;; + esac + ancestor_permissions=$((0$ancestor_mode)) + [ $((ancestor_permissions & 0022)) -eq 0 ] \ + || fail "$description 的祖先目录不得允许组或其他用户写入:$current" + fi + parent=$(dirname "$current") + [ "$parent" != "$current" ] || break + current="$parent" + done +} + +prepare_secure_config_path() { + config_dir=$(dirname "$CONFIG_PATH") + config_name=$(basename "$CONFIG_PATH") + validate_trusted_directory_ancestors "$config_dir" "配置目录" + [ ! -L "$config_dir" ] || fail "配置目录不得是符号链接:$config_dir" + if [ -e "$config_dir" ] && [ ! -d "$config_dir" ]; then + fail "配置目录路径不是目录:$config_dir" + fi + if [ ! -d "$config_dir" ]; then + mkdir -p "$config_dir" + chmod 700 "$config_dir" + fi + [ ! -L "$config_dir" ] || fail "配置目录不得是符号链接:$config_dir" + config_dir=$(cd "$config_dir" && pwd -P) || fail "无法解析配置目录:$config_dir" + CONFIG_PATH="$config_dir/$config_name" + + config_dir_owner=$(stat_owner_id "$config_dir") || fail "无法读取配置目录所有者:$config_dir" + [ "$config_dir_owner" = "$(id -u)" ] || fail "配置目录必须属于当前用户:$config_dir" + config_dir_mode=$(stat_mode "$config_dir") || fail "无法读取配置目录权限:$config_dir" + case "$config_dir_mode" in + *[!0-7]*) fail "配置目录权限格式无效:$config_dir_mode" ;; + esac + config_dir_permissions=$((0$config_dir_mode)) + [ $((config_dir_permissions & 0022)) -eq 0 ] \ + || fail "配置目录不得允许组或其他用户写入:$config_dir" + + [ ! -L "$CONFIG_PATH" ] || fail "配置文件不得是符号链接:$CONFIG_PATH" + if [ -e "$CONFIG_PATH" ]; then + require_single_link_regular_file "$CONFIG_PATH" "配置文件" + config_owner=$(stat_owner_id "$CONFIG_PATH") || fail "无法读取配置文件所有者:$CONFIG_PATH" + [ "$config_owner" = "$(id -u)" ] || fail "配置文件必须属于当前用户:$CONFIG_PATH" + chmod 600 "$CONFIG_PATH" || fail "无法保护配置文件权限:$CONFIG_PATH" + fi +} + verify_checksum() { archive="$1" sums="$2" asset="$3" - [ -f "$sums" ] || return 0 - expected=$(awk -v asset="$asset" '$2 == asset { print $1 }' "$sums" | head -n 1) - [ -n "$expected" ] || return 0 + require_single_link_regular_file "$archive" "release 制品" + require_single_link_regular_file "$sums" "SHA256 校验文件" + matches=$(awk -v asset="$asset" ' + { + targets_asset = 0 + for (field = 2; field <= NF; field += 1) { + if ($field == asset || $field == "*" asset) targets_asset = 1 + } + if (targets_asset) { + if (NF == 2 && ($2 == asset || $2 == "*" asset)) print $1 + else print "INVALID" + } + } + ' "$sums") + match_count=$(printf '%s\n' "$matches" | awk 'NF { count += 1 } END { print count + 0 }') + [ "$match_count" -eq 1 ] \ + || fail "SHA256SUMS.txt 必须且只能包含一个制品条目:$asset" + expected=$(printf '%s\n' "$matches" | awk 'NF { print; exit }') + [ "${#expected}" -eq 64 ] || fail "SHA256SUMS.txt 中的哈希格式无效:$asset" + case "$expected" in + *[!0-9A-Fa-f]*) + fail "SHA256SUMS.txt 中的哈希格式无效:$asset" + ;; + esac if command -v sha256sum >/dev/null 2>&1; then actual=$(sha256sum "$archive" | awk '{print $1}') elif command -v shasum >/dev/null 2>&1; then actual=$(shasum -a 256 "$archive" | awk '{print $1}') else - say "未找到 sha256sum/shasum,跳过校验" - return 0 + fail "缺少 sha256sum 或 shasum,无法验证 release 制品" fi + actual=$(printf '%s' "$actual" | tr 'A-F' 'a-f') + expected=$(printf '%s' "$expected" | tr 'A-F' 'a-f') [ "$actual" = "$expected" ] || fail "SHA256 校验失败:$asset" } +extract_tunnel_binary() { + archive="$1" + destination="$2" + require_single_link_regular_file "$archive" "release 制品" + members=$(tar -tzf "$archive" 2>/dev/null) || fail "无法读取 release 制品" + member_count=$(printf '%s\n' "$members" | awk 'NF { count += 1 } END { print count + 0 }') + [ "$member_count" -eq 1 ] || fail "制品必须且只能包含一个 aether-tunnel 文件" + [ "$members" = "aether-tunnel" ] || fail "制品必须在根目录包含 aether-tunnel" + + listing=$(tar -tvzf "$archive" 2>/dev/null) || fail "无法读取 release 制品元数据" + member_type=$(printf '%s\n' "$listing" | awk 'NF { print substr($1, 1, 1); exit }') + [ "$member_type" = "-" ] || fail "制品中的 aether-tunnel 不是普通文件" + + destination_dir=$(dirname "$destination") + extract_tmp=$(mktemp "$destination_dir/.aether-tunnel.extract.XXXXXXXX") \ + || fail "无法创建安全解压临时文件" + if ! tar -xOzf "$archive" aether-tunnel > "$extract_tmp"; then + rm -f "$extract_tmp" + fail "无法安全提取 aether-tunnel" + fi + [ -s "$extract_tmp" ] || { + rm -f "$extract_tmp" + fail "制品中的 aether-tunnel 为空" + } + mv -f "$extract_tmp" "$destination" || { + rm -f "$extract_tmp" + fail "无法原子保存解压后的 aether-tunnel" + } + require_single_link_regular_file "$destination" "解压后的 aether-tunnel" +} + +prepare_secure_install_dir() { + validate_trusted_directory_ancestors "$INSTALL_DIR" "安装目录" + [ ! -L "$INSTALL_DIR" ] || fail "安装目录不得是符号链接:$INSTALL_DIR" + if [ -e "$INSTALL_DIR" ] && [ ! -d "$INSTALL_DIR" ]; then + fail "安装目录路径不是目录:$INSTALL_DIR" + fi + if [ ! -d "$INSTALL_DIR" ]; then + mkdir -p "$INSTALL_DIR" + fi + [ ! -L "$INSTALL_DIR" ] || fail "安装目录不得是符号链接:$INSTALL_DIR" + INSTALL_DIR=$(cd "$INSTALL_DIR" && pwd -P) || fail "无法解析安装目录:$INSTALL_DIR" + + install_dir_owner=$(stat_owner_id "$INSTALL_DIR") || fail "无法读取安装目录所有者:$INSTALL_DIR" + [ "$install_dir_owner" = "$(id -u)" ] || fail "安装目录必须属于当前用户:$INSTALL_DIR" + install_dir_mode=$(stat_mode "$INSTALL_DIR") || fail "无法读取安装目录权限:$INSTALL_DIR" + case "$install_dir_mode" in + *[!0-7]*) fail "安装目录权限格式无效:$install_dir_mode" ;; + esac + install_dir_permissions=$((0$install_dir_mode)) + [ $((install_dir_permissions & 0022)) -eq 0 ] \ + || fail "安装目录不得允许组或其他用户写入:$INSTALL_DIR" +} + +install_tunnel_binary_file() { + source_binary="$1" + require_single_link_regular_file "$source_binary" "待安装二进制" + prepare_secure_install_dir + target_binary="$INSTALL_DIR/aether-tunnel" + [ ! -L "$target_binary" ] || fail "安装目标不得是符号链接:$target_binary" + if [ -e "$target_binary" ]; then + require_single_link_regular_file "$target_binary" "安装目标" + target_owner=$(stat_owner_id "$target_binary") || fail "无法读取安装目标所有者:$target_binary" + [ "$target_owner" = "$(id -u)" ] || fail "安装目标必须属于当前用户:$target_binary" + fi + + install_tmp=$(mktemp "$INSTALL_DIR/.aether-tunnel.tmp.XXXXXXXX") \ + || fail "无法在安装目录中创建安全临时文件" + if ! cat "$source_binary" > "$install_tmp"; then + rm -f "$install_tmp" + fail "无法写入临时安装文件" + fi + chmod 0755 "$install_tmp" || { + rm -f "$install_tmp" + fail "无法设置临时安装文件权限" + } + + [ ! -L "$target_binary" ] || { + rm -f "$install_tmp" + fail "安装目标在写入期间变成了符号链接:$target_binary" + } + if [ -e "$target_binary" ]; then + if ! require_single_link_regular_file "$target_binary" "安装目标"; then + rm -f "$install_tmp" + fail "安装目标在写入期间变得不安全:$target_binary" + fi + fi + mv -f "$install_tmp" "$target_binary" || { + rm -f "$install_tmp" + fail "无法原子替换安装目标:$target_binary" + } +} + install_binary() { tag="$1" asset="$2" + validate_release_repo "$REPO" + validate_tunnel_release_tag "$tag" + validate_release_asset_name "$asset" base="https://github.com/${REPO}/releases/download/${tag}" archive="$TMP_DIR/$asset" say "下载 $tag / $asset" download "$base/$asset" "$archive" - download "$base/SHA256SUMS.txt" "$TMP_DIR/SHA256SUMS.txt" >/dev/null 2>&1 || true + download "$base/SHA256SUMS.txt" "$TMP_DIR/SHA256SUMS.txt" >/dev/null 2>&1 || fail "无法下载 SHA256SUMS.txt" verify_checksum "$archive" "$TMP_DIR/SHA256SUMS.txt" "$asset" - tar -xzf "$archive" -C "$TMP_DIR" - [ -f "$TMP_DIR/aether-tunnel" ] || fail "制品中未找到 aether-tunnel" - mkdir -p "$INSTALL_DIR" - cp "$TMP_DIR/aether-tunnel" "$INSTALL_DIR/aether-tunnel" - chmod +x "$INSTALL_DIR/aether-tunnel" + extract_tunnel_binary "$archive" "$TMP_DIR/aether-tunnel" + install_tunnel_binary_file "$TMP_DIR/aether-tunnel" say "已安装二进制:$INSTALL_DIR/aether-tunnel" } @@ -189,7 +566,8 @@ append_server_config() { tunnel_security="$4" tunnel_encryption_key="$5" - mkdir -p "$(dirname "$CONFIG_PATH")" + validate_node_name "$node_name" + prepare_secure_config_path quoted_url=$(toml_quote "$aether_url") quoted_token=$(toml_quote "$management_token") quoted_name=$(toml_quote "$node_name") @@ -207,12 +585,22 @@ append_server_config() { return fi + config_existed=false if [ -f "$CONFIG_PATH" ]; then - cp "$CONFIG_PATH" "$CONFIG_PATH.bak.$(date +%Y%m%d%H%M%S)" + config_existed=true + backup_path=$(mktemp "$CONFIG_PATH.bak.$(date +%Y%m%d%H%M%S).XXXXXXXX") \ + || fail "无法创建安全配置备份" + chmod 600 "$backup_path" || fail "无法保护配置备份权限:$backup_path" + cat "$CONFIG_PATH" > "$backup_path" || fail "无法备份配置文件:$CONFIG_PATH" fi + CONFIG_TMP_PATH=$(mktemp "$CONFIG_PATH.tmp.XXXXXXXX") || fail "无法创建安全配置临时文件" + chmod 600 "$CONFIG_TMP_PATH" || fail "无法保护配置临时文件权限" + if [ "$config_existed" = true ]; then + cat "$CONFIG_PATH" > "$CONFIG_TMP_PATH" || fail "无法读取现有配置:$CONFIG_PATH" + fi { - if [ -f "$CONFIG_PATH" ] && [ -s "$CONFIG_PATH" ]; then + if [ -s "$CONFIG_TMP_PATH" ]; then printf '\n' fi printf '# Added by Aether Tunnel one-click installer. Existing config is preserved.\n' @@ -226,19 +614,28 @@ append_server_config() { if [ -n "$tunnel_encryption_key" ]; then printf 'tunnel_encryption_key = %s\n' "$quoted_encryption_key" fi - } >> "$CONFIG_PATH" - chmod 600 "$CONFIG_PATH" 2>/dev/null || true + } >> "$CONFIG_TMP_PATH" + chmod 600 "$CONFIG_TMP_PATH" || fail "无法保护配置临时文件权限" + [ ! -L "$CONFIG_PATH" ] || fail "配置文件在写入期间变成了符号链接:$CONFIG_PATH" + if [ -e "$CONFIG_PATH" ]; then + require_single_link_regular_file "$CONFIG_PATH" "配置文件" + fi + mv -f "$CONFIG_TMP_PATH" "$CONFIG_PATH" || fail "无法原子替换配置文件:$CONFIG_PATH" + CONFIG_TMP_PATH="" + chmod 600 "$CONFIG_PATH" || fail "无法保护配置文件权限:$CONFIG_PATH" say "已追加 [[servers]] 到:$CONFIG_PATH" } main() { TMP_DIR=$(mktemp -d 2>/dev/null || mktemp -d -t aether-tunnel) need_cmd tar + validate_release_repo "$REPO" choose_paths aether_url=$(prompt_if_empty AETHER_TUNNEL_AETHER_URL "${AETHER_TUNNEL_AETHER_URL:-}" "Aether URL: ") management_token=$(prompt_if_empty AETHER_TUNNEL_MANAGEMENT_TOKEN "${AETHER_TUNNEL_MANAGEMENT_TOKEN:-}" "Management token (ae_xxx): ") node_name=$(prompt_if_empty AETHER_TUNNEL_NODE_NAME "${AETHER_TUNNEL_NODE_NAME:-}" "Node name: ") + validate_node_name "$node_name" tunnel_security="${AETHER_TUNNEL_SECURITY:-}" tunnel_encryption_key="${AETHER_TUNNEL_ENCRYPTION_KEY:-}" case "$tunnel_security" in @@ -251,6 +648,7 @@ main() { tag=$(resolve_latest_tunnel_tag) [ -n "$tag" ] || fail "没有找到可用的 tunnel-v* release" + validate_tunnel_release_tag "$tag" asset=$(detect_asset) install_binary "$tag" "$asset" append_server_config "$aether_url" "$management_token" "$node_name" "$tunnel_security" "$tunnel_encryption_key" diff --git a/apps/aether-tunnel/src/app.rs b/apps/aether-tunnel/src/app.rs index 77e0a0517..e4a422b56 100644 --- a/apps/aether-tunnel/src/app.rs +++ b/apps/aether-tunnel/src/app.rs @@ -19,8 +19,8 @@ use tokio::task::JoinHandle; use tracing::{error, info, warn}; use crate::config::{ - effective_tunnel_security, validate_tunnel_encryption_key, Config, ServerEntry, - TunnelPoolSizing, + aether_url_for_log, effective_tunnel_security, validate_tunnel_encryption_key, Config, + ServerEntry, TunnelPoolSizing, }; use crate::net; use crate::registration::client::AetherClient; @@ -96,6 +96,11 @@ struct DiagnosticsState { /// Run the full application lifecycle after config has been parsed. pub async fn run(mut config: Config, servers: Vec) -> anyhow::Result<()> { config.validate()?; + for (index, server) in servers.iter().enumerate() { + server + .validate() + .map_err(|error| anyhow::anyhow!("servers[{index}] invalid: {error}"))?; + } init_tracing(&config); info!( @@ -227,7 +232,7 @@ pub async fn run(mut config: Config, servers: Vec) -> anyhow::Resul if entry.aether_url.trim_start().starts_with("http://") { warn!( server = %label, - url = %entry.aether_url, + url = %aether_url_for_log(&entry.aether_url), "secure tunnel frame encryption starts after registration; deliver install and registration credentials over HTTPS or another trusted bootstrap channel" ); } @@ -245,16 +250,29 @@ pub async fn run(mut config: Config, servers: Vec) -> anyhow::Resul .register(&config, entry, &node_name, &public_ip, Some(&hw_info)) .await { - Ok(node_id) => { - info!(server = %label, node_id = %node_id, url = %entry.aether_url, node_name = %node_name, "registered"); + Ok(registration) => { + info!( + server = %label, + node_id = %registration.node_id, + tunnel_generation = %registration.tunnel_generation, + url = %aether_url_for_log(&entry.aether_url), + node_name = %node_name, + "registered" + ); server_contexts.lock().await.push(build_server_context( - &config, &label, entry, client, &node_name, node_id, + &config, + &label, + entry, + client, + &node_name, + registration.node_id, + registration.tunnel_generation, )); } Err(e) => { warn!( server = %label, - url = %entry.aether_url, + url = %aether_url_for_log(&entry.aether_url), error = %e, "registration failed, will retry in background" ); @@ -644,15 +662,16 @@ async fn retry_failed_registration( ) .await { - Ok(node_id) => { - info!(server = %label, node_id = %node_id, attempt, "registration retry succeeded"); + Ok(registration) => { + info!(server = %label, node_id = %registration.node_id, tunnel_generation = %registration.tunnel_generation, attempt, "registration retry succeeded"); let server = build_server_context( &state.config, &label, &entry, client, &node_name, - node_id, + registration.node_id, + registration.tunnel_generation, ); server_contexts.lock().await.push(Arc::clone(&server)); spawn_tunnel_pool_manager( @@ -712,6 +731,7 @@ fn build_server_context( client: Arc, node_name: &str, node_id: String, + tunnel_generation: String, ) -> Arc { let mut dynamic = DynamicConfig::from_config(config); dynamic.node_name = node_name.to_string(); @@ -727,6 +747,7 @@ fn build_server_context( tunnel_encryption_key: entry.tunnel_encryption_key.clone(), node_name: node_name.to_string(), node_id: Arc::new(RwLock::new(node_id)), + tunnel_generation, aether_client: client, dynamic: Arc::new(ArcSwap::from_pointee(dynamic)), active_connections: Arc::new(AtomicU64::new(0)), @@ -1290,7 +1311,10 @@ mod tests { register_hits.fetch_add(1, Ordering::SeqCst); ( AxumStatusCode::OK, - axum::Json(json!({ "node_id": "node-recovery" })), + axum::Json(json!({ + "node_id": "node-recovery", + "tunnel_generation": "test-generation-recovery" + })), ) } @@ -1351,6 +1375,7 @@ mod tests { client, &state.config.node_name, node_id.to_string(), + "test-generation-1".to_string(), ) } diff --git a/apps/aether-tunnel/src/config.rs b/apps/aether-tunnel/src/config.rs index f1acbde7e..31772fb4c 100644 --- a/apps/aether-tunnel/src/config.rs +++ b/apps/aether-tunnel/src/config.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::io::{self, Read, Write}; use std::net::SocketAddr; use std::path::Path; use std::str::FromStr; @@ -79,6 +80,12 @@ const TUNNEL_PROFILE_ENV: &str = "AETHER_TUNNEL_PROFILE"; const TUNNEL_STREAM_INITIAL_WINDOW_BYTES_ENV: &str = "AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES"; const TUNNEL_DRAIN_DEADLINE_MS_ENV: &str = "AETHER_TUNNEL_DRAIN_DEADLINE_MS"; +// The configuration contains only scalar settings and a bounded list of +// server entries. Refuse an unexpectedly large local file before TOML parsing +// so a replaced or corrupted config cannot force an unbounded allocation at +// service startup. +const MAX_CONFIG_FILE_BYTES: u64 = 1024 * 1024; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct TunnelPoolSizing { pub initial_connections: u32, @@ -193,12 +200,43 @@ pub fn effective_tunnel_security( TunnelSecurity::Off } +pub(crate) fn validate_aether_url(value: &str) -> anyhow::Result<()> { + let value = value.trim(); + let parsed = url::Url::parse(value) + .map_err(|_| anyhow::anyhow!("aether_url must be an absolute HTTP(S) URL"))?; + if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() { + anyhow::bail!("aether_url must be an absolute HTTP(S) URL"); + } + if !aether_http::is_https_or_loopback_http_url(&parsed) { + anyhow::bail!( + "aether_url must use HTTPS; HTTP is allowed only for a literal loopback host" + ); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + anyhow::bail!("aether_url must not contain embedded credentials"); + } + if parsed.query().is_some() || parsed.fragment().is_some() { + anyhow::bail!("aether_url must not contain a query string or fragment"); + } + Ok(()) +} + +pub(crate) fn aether_url_for_log(value: &str) -> String { + let Ok(parsed) = url::Url::parse(value.trim()) else { + return "".to_string(); + }; + if !matches!(parsed.scheme(), "http" | "https" | "ws" | "wss") || parsed.host_str().is_none() { + return "".to_string(); + } + parsed.origin().ascii_serialization() +} + /// Aether tunnel agent. /// /// Deployed on overseas VPS to relay API traffic for Aether instances /// behind the GFW. Connects to Aether via WebSocket tunnel, registers /// with Aether, and relays upstream requests. -#[derive(Parser, Debug, Clone)] +#[derive(Parser, Clone)] #[command(version, about)] pub struct Config { /// Aether server URL (e.g. https://aether.example.com) @@ -250,11 +288,12 @@ pub struct Config { )] pub allowed_ports: Vec, - /// Allow private/reserved upstream IP targets. Enabled by default. + /// Allow private/reserved upstream IP targets. Disabled by default; enable + /// explicitly only for deployments that require access to private services. #[arg( long, env = "AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS", - default_value_t = true + default_value_t = false )] pub allow_private_targets: bool, @@ -644,10 +683,31 @@ pub struct Config { pub tunnel_scale_down_grace_secs: u64, } +impl std::fmt::Debug for Config { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("Config") + .field("aether_url", &aether_url_for_log(&self.aether_url)) + .field("management_token", &"") + .field("node_name", &self.node_name) + .field("node_region", &self.node_region) + .field("tunnel_security", &self.tunnel_security) + .field( + "tunnel_encryption_key", + &self.tunnel_encryption_key.as_ref().map(|_| ""), + ) + .finish_non_exhaustive() + } +} + impl Config { /// Validate configuration values are within sane ranges. /// Called after parsing to catch misconfigurations early. pub fn validate(&self) -> anyhow::Result<()> { + validate_aether_url(&self.aether_url)?; + if self.management_token.trim().is_empty() { + anyhow::bail!("management_token must not be empty"); + } if self.heartbeat_interval == 0 { anyhow::bail!("heartbeat_interval must be > 0"); } @@ -698,6 +758,9 @@ impl Config { if matches!(self.tunnel_connections_max, Some(0)) { anyhow::bail!("tunnel_connections_max must be > 0"); } + if matches!(self.tunnel_max_streams, Some(0)) { + anyhow::bail!("tunnel_max_streams must be > 0"); + } if self.tunnel_stream_initial_window_bytes == 0 { anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0"); } @@ -906,7 +969,7 @@ impl Config { } /// Per-server connection config (used in multi-server TOML `[[servers]]`). -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Clone, Serialize, Deserialize)] #[serde(deny_unknown_fields)] pub struct ServerEntry { pub aether_url: String, @@ -921,13 +984,39 @@ pub struct ServerEntry { pub tunnel_encryption_key: Option, } +impl std::fmt::Debug for ServerEntry { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ServerEntry") + .field("aether_url", &aether_url_for_log(&self.aether_url)) + .field("management_token", &"") + .field("node_name", &self.node_name) + .field("tunnel_security", &self.tunnel_security) + .field( + "tunnel_encryption_key", + &self.tunnel_encryption_key.as_ref().map(|_| ""), + ) + .finish() + } +} + +impl ServerEntry { + pub(crate) fn validate(&self) -> anyhow::Result<()> { + validate_aether_url(&self.aether_url)?; + if self.management_token.trim().is_empty() { + anyhow::bail!("management_token must not be empty"); + } + Ok(()) + } +} + // --------------------------------------------------------------------------- // TOML config file support // --------------------------------------------------------------------------- /// Serializable config for TOML file persistence. /// All fields are optional -- only populated values are written. -#[derive(Debug, Default, Serialize, Deserialize)] +#[derive(Default, Serialize, Deserialize)] #[serde(deny_unknown_fields)] pub struct ConfigFile { #[serde(skip_serializing_if = "Option::is_none")] @@ -1048,17 +1137,46 @@ pub struct ConfigFile { pub servers: Vec, } +impl std::fmt::Debug for ConfigFile { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ConfigFile") + .field("node_name", &self.node_name) + .field("node_region", &self.node_region) + .field("server_count", &self.servers.len()) + .finish_non_exhaustive() + } +} + impl ConfigFile { /// Load from a TOML file. pub fn load(path: &Path) -> anyhow::Result { - let content = std::fs::read_to_string(path)?; + let mut file = std::fs::File::open(path)?; + let advertised_len = file.metadata()?.len(); + if advertised_len > MAX_CONFIG_FILE_BYTES { + anyhow::bail!( + "tunnel config exceeds the {} byte limit", + MAX_CONFIG_FILE_BYTES + ); + } + + let mut content = String::with_capacity(advertised_len as usize); + Read::by_ref(&mut file) + .take(MAX_CONFIG_FILE_BYTES.saturating_add(1)) + .read_to_string(&mut content)?; + if content.len() as u64 > MAX_CONFIG_FILE_BYTES { + anyhow::bail!( + "tunnel config exceeds the {} byte limit", + MAX_CONFIG_FILE_BYTES + ); + } parse_config_file_content(&content) } /// Save to a TOML file. pub fn save(&self, path: &Path) -> anyhow::Result<()> { let content = toml::to_string_pretty(self)?; - std::fs::write(path, content)?; + write_private_config_atomically(path, content.as_bytes())?; Ok(()) } @@ -1268,6 +1386,155 @@ impl ConfigFile { } } +fn write_private_config_atomically(path: &Path, content: &[u8]) -> io::Result<()> { + reject_config_symlink(path)?; + + let parent = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + let file_name = path.file_name().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "configuration path must name a file", + ) + })?; + if !parent.is_dir() { + return Err(io::Error::new( + io::ErrorKind::NotFound, + "configuration parent directory does not exist", + )); + } + + let mut temp_path = None; + let mut temp_file = None; + for _ in 0..16 { + let candidate = parent.join(format!( + ".{}.tmp-{}", + file_name.to_string_lossy(), + uuid::Uuid::new_v4() + )); + let mut options = std::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt as _; + options.mode(0o600); + } + match options.open(&candidate) { + Ok(file) => { + temp_path = Some(candidate); + temp_file = Some(file); + break; + } + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(error), + } + } + let temp_path = temp_path.ok_or_else(|| { + io::Error::new( + io::ErrorKind::AlreadyExists, + "could not allocate a unique configuration temporary file", + ) + })?; + let mut temp_file = temp_file.expect("temporary path and file are created together"); + + let replace_result = (|| -> io::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt as _; + temp_file.set_permissions(std::fs::Permissions::from_mode(0o600))?; + } + temp_file.write_all(content)?; + temp_file.sync_all()?; + drop(temp_file); + + // Do not silently replace a credential-bearing symlink. The second + // check also covers a target created while the temporary file was written. + reject_config_symlink(path)?; + replace_config_file(&temp_path, path)?; + + #[cfg(unix)] + std::fs::File::open(parent)?.sync_all()?; + Ok(()) + })(); + + if replace_result.is_err() { + let _ = std::fs::remove_file(&temp_path); + } + replace_result +} + +fn reject_config_symlink(path: &Path) -> io::Result<()> { + match std::fs::symlink_metadata(path) { + Ok(metadata) if config_metadata_is_link_like(&metadata) => Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "refusing to save configuration through a symbolic link or reparse point", + )), + Ok(_) => Ok(()), + Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()), + Err(error) => Err(error), + } +} + +fn config_metadata_is_link_like(metadata: &std::fs::Metadata) -> bool { + #[cfg(windows)] + { + use std::os::windows::fs::MetadataExt as _; + const FILE_ATTRIBUTE_REPARSE_POINT: u32 = 0x0400; + metadata.file_attributes() & FILE_ATTRIBUTE_REPARSE_POINT != 0 + } + #[cfg(not(windows))] + { + metadata.file_type().is_symlink() + } +} + +#[cfg(not(windows))] +fn replace_config_file(temp_path: &Path, path: &Path) -> io::Result<()> { + std::fs::rename(temp_path, path) +} + +#[cfg(windows)] +fn replace_config_file(temp_path: &Path, path: &Path) -> io::Result<()> { + use std::os::windows::ffi::OsStrExt as _; + + const MOVEFILE_REPLACE_EXISTING: u32 = 0x0000_0001; + const MOVEFILE_WRITE_THROUGH: u32 = 0x0000_0008; + + #[link(name = "kernel32")] + unsafe extern "system" { + fn MoveFileExW( + existing_file_name: *const u16, + new_file_name: *const u16, + flags: u32, + ) -> i32; + } + + let existing_file_name = temp_path + .as_os_str() + .encode_wide() + .chain(std::iter::once(0)) + .collect::>(); + let new_file_name = path + .as_os_str() + .encode_wide() + .chain(std::iter::once(0)) + .collect::>(); + let replaced = unsafe { + MoveFileExW( + existing_file_name.as_ptr(), + new_file_name.as_ptr(), + MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH, + ) + }; + if replaced == 0 { + Err(io::Error::last_os_error()) + } else { + Ok(()) + } +} + fn parse_config_file_content(content: &str) -> anyhow::Result { reject_removed_config_keys(content)?; let mut value: toml::Value = toml::from_str(content)?; @@ -1396,6 +1663,151 @@ mod tests { use super::*; use crate::hardware::HardwareInfo; + fn config_save_test_dir(label: &str) -> std::path::PathBuf { + let path = std::env::temp_dir().join(format!( + "aether-tunnel-config-{label}-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&path).expect("config save test directory should be created"); + path + } + + fn secret_bearing_config_file(secret: &str) -> ConfigFile { + ConfigFile { + node_name: Some("secure-save-test".to_string()), + servers: vec![ServerEntry { + aether_url: "https://example.com".to_string(), + management_token: secret.to_string(), + node_name: None, + tunnel_security: None, + tunnel_encryption_key: None, + }], + ..ConfigFile::default() + } + } + + #[test] + fn config_file_save_replaces_an_existing_config() { + let directory = config_save_test_dir("replace-existing"); + let path = directory.join("tunnel.toml"); + + secret_bearing_config_file("first-management-secret") + .save(&path) + .expect("initial config save should succeed"); + secret_bearing_config_file("second-management-secret") + .save(&path) + .expect("replacement config save should succeed"); + + let saved = std::fs::read_to_string(&path).expect("replacement config should be readable"); + assert!(saved.contains("second-management-secret")); + assert!(!saved.contains("first-management-secret")); + assert_eq!( + std::fs::read_dir(&directory) + .expect("test directory should be readable") + .count(), + 1, + "replacement save must not leave a temporary file" + ); + + std::fs::remove_dir_all(directory).expect("config save test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn config_file_save_atomically_replaces_with_owner_only_permissions() { + use std::io::Read as _; + use std::os::unix::fs::PermissionsExt as _; + + let directory = config_save_test_dir("atomic-private"); + let path = directory.join("tunnel.toml"); + std::fs::write(&path, "old configuration") + .expect("existing config fixture should be written"); + std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o644)) + .expect("existing config fixture should be world-readable"); + let mut old_handle = + std::fs::File::open(&path).expect("existing config handle should open"); + + secret_bearing_config_file("management-secret") + .save(&path) + .expect("config save should succeed"); + + let mode = std::fs::metadata(&path) + .expect("saved config metadata should exist") + .permissions() + .mode() + & 0o777; + assert_eq!(mode, 0o600); + let saved = std::fs::read_to_string(&path).expect("saved config should be readable"); + assert!(saved.contains("management-secret")); + + let mut old = String::new(); + old_handle + .read_to_string(&mut old) + .expect("old inode should remain readable through its open handle"); + assert_eq!(old, "old configuration"); + assert_eq!( + std::fs::read_dir(&directory) + .expect("test directory should be readable") + .count(), + 1, + "successful save must not leave a temporary file" + ); + + std::fs::remove_dir_all(directory).expect("config save test directory should be removed"); + } + + #[cfg(unix)] + #[test] + fn config_file_save_rejects_symbolic_link_targets() { + use std::os::unix::fs::symlink; + + let directory = config_save_test_dir("symlink"); + let target = directory.join("target.toml"); + let path = directory.join("tunnel.toml"); + std::fs::write(&target, "target sentinel") + .expect("symlink target fixture should be written"); + symlink(&target, &path).expect("config symlink fixture should be created"); + + let error = secret_bearing_config_file("management-secret") + .save(&path) + .expect_err("saving through a symlink must fail"); + + assert!(error.to_string().contains("symbolic link")); + assert_eq!( + std::fs::read_to_string(&target).expect("symlink target should remain readable"), + "target sentinel" + ); + assert!(std::fs::symlink_metadata(&path) + .expect("config symlink should still exist") + .file_type() + .is_symlink()); + + std::fs::remove_dir_all(directory).expect("config save test directory should be removed"); + } + + #[test] + fn config_file_save_cleans_temporary_file_when_replace_fails() { + let directory = config_save_test_dir("cleanup"); + let path = directory.join("destination-is-a-directory"); + std::fs::create_dir(&path).expect("destination directory fixture should be created"); + + secret_bearing_config_file("management-secret") + .save(&path) + .expect_err("replacing a directory with a config file must fail"); + + let entries = std::fs::read_dir(&directory) + .expect("test directory should be readable") + .map(|entry| { + entry + .expect("test directory entry should be readable") + .path() + }) + .collect::>(); + assert_eq!(entries, vec![path]); + + std::fs::remove_dir_all(directory).expect("config save test directory should be removed"); + } + #[test] fn config_file_load_ignores_removed_redirect_replay_budget() { let config = parse_config_file_content("redirect_replay_budget_bytes = \"1K\"") @@ -1404,6 +1816,19 @@ mod tests { assert!(!serialized.contains("redirect_replay_budget_bytes")); } + #[test] + fn config_file_load_rejects_oversized_files_before_parsing() { + let directory = config_save_test_dir("oversized-load"); + let path = directory.join("tunnel.toml"); + std::fs::write(&path, vec![b'a'; MAX_CONFIG_FILE_BYTES as usize + 1]) + .expect("oversized config fixture should be written"); + + let error = ConfigFile::load(&path).expect_err("oversized config must be rejected"); + assert!(error.to_string().contains("exceeds")); + + std::fs::remove_dir_all(directory).expect("config test directory should be removed"); + } + #[test] fn cli_accepts_but_hides_legacy_redirect_replay_budget() { let config = Config::parse_from([ @@ -1464,7 +1889,7 @@ tunnel_ipv6_only = false let cfg: ConfigFile = toml::from_str( r#" [[servers]] -aether_url = "http://aether.example.com" +aether_url = "http://127.0.0.1:8084" management_token = "ae_test" node_name = "jp-proxy-01" tunnel_security = "non_tls_required" @@ -1664,7 +2089,7 @@ node_name = "tunnel-test" } #[test] - fn cli_defaults_private_targets_to_enabled() { + fn cli_defaults_private_targets_to_disabled() { let config = Config::parse_from([ "aether-tunnel", "--aether-url", @@ -1674,6 +2099,21 @@ node_name = "tunnel-test" "--node-name", "tunnel-test", ]); + assert!(!config.allow_private_targets); + } + + #[test] + fn cli_allows_explicit_private_targets() { + let config = Config::parse_from([ + "aether-tunnel", + "--aether-url", + "https://example.com", + "--management-token", + "ae_test", + "--node-name", + "tunnel-test", + "--allow-private-targets", + ]); assert!(config.allow_private_targets); } @@ -1713,12 +2153,134 @@ node_name = "tunnel-test" assert!(config.tunnel_encryption_key.is_none()); } + #[test] + fn aether_url_validation_rejects_embedded_secrets_and_non_http_schemes() { + for value in [ + "https://alice:password@example.com", + "https://example.com?token=secret", + "https://example.com#secret-fragment", + "http://example.com", + "http://10.0.0.1:8084", + "http://[::ffff:127.0.0.1]:8084", + "file:///etc/passwd", + "not-a-url", + ] { + assert!( + validate_aether_url(value).is_err(), + "URL should be rejected: {value}" + ); + } + validate_aether_url("https://example.com/base/path") + .expect("ordinary HTTPS URL should validate"); + validate_aether_url("http://127.0.0.1:8084/base/path") + .expect("literal loopback HTTP should validate"); + validate_aether_url("http://[::1]:8084/base/path") + .expect("literal IPv6 loopback HTTP should validate"); + } + + #[test] + fn aether_url_log_projection_removes_credentials_query_and_fragment() { + let projected = aether_url_for_log( + "https://alice:password@example.com/base?token=query-secret#secret-fragment", + ); + + assert_eq!(projected, "https://example.com"); + for secret in ["alice", "password", "query-secret", "secret-fragment"] { + assert!(!projected.contains(secret)); + } + assert_eq!(aether_url_for_log("not-a-url"), ""); + } + + #[test] + fn config_and_server_entries_require_management_tokens() { + let config = Config::parse_from([ + "aether-tunnel", + "--aether-url", + "https://example.com", + "--management-token", + " ", + "--node-name", + "tunnel-test", + ]); + assert!(config.validate().is_err()); + + let entry = ServerEntry { + aether_url: "https://example.com".to_string(), + management_token: " ".to_string(), + node_name: None, + tunnel_security: None, + tunnel_encryption_key: None, + }; + assert!(entry.validate().is_err()); + } + + #[test] + fn secret_bearing_config_debug_output_is_redacted() { + let config = Config::parse_from([ + "aether-tunnel", + "--aether-url", + "https://alice:password@example.com/base?token=query-secret", + "--management-token", + "management-secret", + "--node-name", + "tunnel-test", + "--tunnel-encryption-key", + "psk-secret", + ]); + let config_debug = format!("{config:?}"); + for secret in [ + "alice", + "password", + "query-secret", + "management-secret", + "psk-secret", + ] { + assert!(!config_debug.contains(secret)); + } + + let entry = ServerEntry { + aether_url: + "https://alice:password@example.com/base?token=query-secret#fragment-secret" + .to_string(), + management_token: "management-secret".to_string(), + node_name: Some("edge-1".to_string()), + tunnel_security: Some(TunnelSecurity::NonTlsRequired), + tunnel_encryption_key: Some("psk-secret".to_string()), + }; + let entry_debug = format!("{entry:?}"); + for secret in [ + "alice", + "password", + "query-secret", + "fragment-secret", + "management-secret", + "psk-secret", + ] { + assert!(!entry_debug.contains(secret)); + } + + let file = ConfigFile { + upstream_proxy_url: Some("http://proxy-user:proxy-secret@proxy.example".to_string()), + servers: vec![entry], + ..ConfigFile::default() + }; + let file_debug = format!("{file:?}"); + for secret in [ + "proxy-user", + "proxy-secret", + "management-secret", + "psk-secret", + ] { + assert!(!file_debug.contains(secret)); + } + } + #[test] fn validate_requires_encryption_key_for_non_tls_security() { let config = Config::parse_from([ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", @@ -1735,7 +2297,7 @@ node_name = "tunnel-test" let with_key = Config::parse_from([ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", @@ -1755,7 +2317,7 @@ node_name = "tunnel-test" let config = Config::parse_from([ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", @@ -1791,7 +2353,7 @@ node_name = "tunnel-test" let config = Config::parse_from([ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", @@ -1913,6 +2475,26 @@ node_name = "tunnel-test" assert!(error.to_string().contains("tunnel_ipv4_only")); } + #[test] + fn validate_rejects_zero_tunnel_max_streams() { + let config = Config::parse_from([ + "aether-tunnel", + "--aether-url", + "https://example.com", + "--management-token", + "ae_test", + "--node-name", + "tunnel-test", + "--tunnel-max-streams", + "0", + ]); + + let error = config + .validate() + .expect_err("zero tunnel stream capacity must be rejected"); + assert!(error.to_string().contains("tunnel_max_streams")); + } + #[test] fn tunnel_fast_recovery_defaults_use_millisecond_values() { let config = Config::parse_from([ diff --git a/apps/aether-tunnel/src/egress_proxy.rs b/apps/aether-tunnel/src/egress_proxy.rs index 97498c5ec..f121be87c 100644 --- a/apps/aether-tunnel/src/egress_proxy.rs +++ b/apps/aether-tunnel/src/egress_proxy.rs @@ -40,7 +40,7 @@ pub(crate) enum UpstreamProxyScheme { Socks5h, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub(crate) struct UpstreamProxyConfig { raw: String, scheme: UpstreamProxyScheme, @@ -50,6 +50,20 @@ pub(crate) struct UpstreamProxyConfig { password: Option, } +impl std::fmt::Debug for UpstreamProxyConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("UpstreamProxyConfig") + .field("url", &self.redacted_url()) + .field("scheme", &self.scheme) + .field("host", &self.host) + .field("port", &self.port) + .field("username", &self.username.as_ref().map(|_| "[REDACTED]")) + .field("password", &self.password.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + impl UpstreamProxyConfig { pub(crate) fn parse(raw: &str) -> Result { let trimmed = raw.trim(); @@ -75,6 +89,18 @@ impl UpstreamProxyConfig { .filter(|value| !value.is_empty()) .ok_or_else(|| "upstream proxy URL must include a host".to_string())? .to_string(); + // A proxy URL identifies the proxy origin. Path/query/fragment + // components are not part of HTTP CONNECT or SOCKS negotiation and + // silently ignoring them can make the configured endpoint differ + // from what operators see in configuration and logs. + if !matches!(parsed.path(), "" | "/") + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return Err( + "upstream proxy URL must be an origin without path, query, or fragment".to_string(), + ); + } let port = parsed.port().unwrap_or(match scheme { UpstreamProxyScheme::Http => 80, UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => 1080, @@ -181,6 +207,38 @@ pub(crate) async fn connect_target_via_proxy( Ok(tcp) } +pub(crate) async fn connect_validated_target_via_proxy( + proxy: &UpstreamProxyConfig, + target_addr: SocketAddr, + options: ProxyConnectOptions, +) -> io::Result { + let mut tcp = connect_proxy_tcp( + proxy, + options.connect_timeout, + options.tcp_nodelay, + options.tcp_keepalive, + options.ip_family, + ) + .await?; + + match proxy.scheme() { + UpstreamProxyScheme::Http => { + http_connect(&mut tcp, &target_addr.to_string(), proxy).await?; + } + UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => { + socks5_connect( + &mut tcp, + proxy, + &target_addr.ip().to_string(), + target_addr.port(), + ) + .await?; + } + } + + Ok(tcp) +} + pub(crate) async fn connect_proxy_tcp( proxy: &UpstreamProxyConfig, connect_timeout: Duration, @@ -188,16 +246,19 @@ pub(crate) async fn connect_proxy_tcp( tcp_keepalive: Option, ip_family: IpFamily, ) -> io::Result { - let resolved = tokio::time::timeout( - connect_timeout, - tokio::net::lookup_host((proxy.host(), proxy.port())), - ) - .await - .map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "proxy DNS timeout"))? - .map_err(|err| io::Error::other(format!("proxy DNS failed: {err}")))?; + let resolved = + aether_http::lookup_host_with_limits(proxy.host(), proxy.port(), connect_timeout) + .await + .map_err(|err| { + if err.kind() == io::ErrorKind::TimedOut { + io::Error::new(io::ErrorKind::TimedOut, "proxy DNS timeout") + } else { + io::Error::other(format!("proxy DNS failed: {err}")) + } + })?; let mut last_error = None; - for addr in resolved.filter(|addr| ip_family.allows(*addr)) { + for addr in resolved.into_iter().filter(|addr| ip_family.allows(*addr)) { match tokio::time::timeout(connect_timeout, TcpStream::connect(addr)).await { Ok(Ok(stream)) => { configure_tcp_stream(&stream, tcp_nodelay, tcp_keepalive)?; @@ -395,12 +456,16 @@ pub(crate) async fn socks5_target_address( request.push(host.len() as u8); request.extend_from_slice(host); } else { - let mut resolved = tokio::net::lookup_host((target_host, target_port)) - .await - .map_err(|err| io::Error::other(format!("SOCKS5 target DNS failed: {err}")))?; - let addr = resolved - .next() - .ok_or_else(|| io::Error::other("SOCKS5 target DNS returned no addresses"))?; + let addr = aether_http::lookup_host_with_limits( + target_host, + target_port, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|err| io::Error::other(format!("SOCKS5 target DNS failed: {err}")))? + .into_iter() + .next() + .ok_or_else(|| io::Error::other("SOCKS5 target DNS returned no addresses"))?; push_socks5_socket_address(&mut request, addr); } request.extend_from_slice(&target_port.to_be_bytes()); @@ -481,6 +546,13 @@ mod tests { proxy.basic_auth_header().as_deref(), Some("Basic dXNlcjpwYXNz") ); + + let debug = format!("{proxy:?}"); + assert!(!debug.contains("user:pass")); + assert!(!debug.contains("Some(\"user\")")); + assert!(!debug.contains("Some(\"pass\")")); + assert!(!debug.contains("dXNlcjpwYXNz")); + assert!(debug.contains("[REDACTED]")); } #[test] @@ -490,4 +562,24 @@ mod tests { assert!(error.contains("unsupported upstream proxy scheme")); } + + #[test] + fn rejects_proxy_urls_with_non_origin_components() { + for value in [ + "http://proxy.example/path", + "http://proxy.example?token=secret", + "socks5://proxy.example#fragment", + ] { + let error = UpstreamProxyConfig::parse(value) + .expect_err("proxy URL with non-origin components should be rejected"); + assert!( + error.contains("without path, query, or fragment"), + "unexpected error for {value}: {error}" + ); + } + + for value in ["http://proxy.example", "http://proxy.example/"] { + UpstreamProxyConfig::parse(value).expect("root proxy origin should be accepted"); + } + } } diff --git a/apps/aether-tunnel/src/main.rs b/apps/aether-tunnel/src/main.rs index 68221558a..ee621be81 100644 --- a/apps/aether-tunnel/src/main.rs +++ b/apps/aether-tunnel/src/main.rs @@ -226,7 +226,7 @@ mod tests { let (config, tunnel_security) = parse_config_and_security(&[ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", @@ -252,7 +252,7 @@ mod tests { let (config, tunnel_security) = parse_config_and_security(&[ "aether-tunnel", "--aether-url", - "http://example.com", + "http://127.0.0.1:8084", "--management-token", "ae_test", "--node-name", diff --git a/apps/aether-tunnel/src/net.rs b/apps/aether-tunnel/src/net.rs index f367950e0..a3bcb1198 100644 --- a/apps/aether-tunnel/src/net.rs +++ b/apps/aether-tunnel/src/net.rs @@ -2,8 +2,15 @@ //! //! These are standalone helpers not tied to any specific client or service. -use aether_http::{build_http_client, HttpClientConfig}; +use std::net::IpAddr; + +use aether_http::{ + build_http_client, is_private_or_reserved_ip, read_response_bytes_with_limit, HttpClientConfig, +}; use tracing::{debug, info}; +use url::Url; + +const MAX_NETWORK_DISCOVERY_RESPONSE_BYTES: usize = 16 * 1024; /// Auto-detect public IP by querying external services. pub async fn detect_public_ip() -> anyhow::Result { @@ -22,10 +29,21 @@ pub async fn detect_public_ip() -> anyhow::Result { for endpoint in &endpoints { match client.get(*endpoint).send().await { Ok(resp) if resp.status().is_success() => { - let ip = resp.text().await?.trim().to_string(); - if !ip.is_empty() { - info!(ip = %ip, source = %endpoint, "detected public IP"); - return Ok(ip); + match read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES) + .await + { + Ok(body) => { + let text = String::from_utf8_lossy(&body); + if let Some(ip) = parse_ip(text.as_ref()) { + let ip = ip.to_string(); + info!(ip = %ip, source = %endpoint, "detected public IP"); + return Ok(ip); + } + debug!(endpoint = %endpoint, "IP detection response was not a valid IP"); + } + Err(error) => { + debug!(endpoint = %endpoint, error = %error, "IP detection response rejected"); + } } } Ok(resp) => { @@ -47,8 +65,20 @@ pub async fn detect_public_ip() -> anyhow::Result { /// This is best-effort and non-sensitive -- region detection should never /// block startup. pub async fn detect_region(ip: &str) -> Option { + // The value may come directly from the command line/environment. Parse + // it before putting it into a URL and avoid disclosing private or reserved + // addresses to third-party geolocation services. We intentionally do not + // reject such values from registration: controlled internal deployments + // can still advertise their configured address and region explicitly. + let ip = parse_ip(ip)?; + if is_private_or_reserved_ip(ip) { + debug!(ip = %ip, "skipping region detection for private or reserved IP"); + return None; + } + let ip = ip.to_string(); + // Try HTTPS provider first - let https_url = format!("https://ipinfo.io/{}/country", ip); + let https_url = ipinfo_url(&ip)?; let client = build_http_client(&HttpClientConfig { request_timeout_ms: Some(5_000), @@ -60,27 +90,29 @@ pub async fn detect_region(ip: &str) -> Option { // Try ipinfo.io (HTTPS, returns plain text country code) if let Ok(resp) = client.get(&https_url).send().await { if resp.status().is_success() { - if let Ok(text) = resp.text().await { - let code = text.trim(); - if !code.is_empty() && code.len() <= 3 { + if let Ok(body) = + read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES).await + { + let text = String::from_utf8_lossy(&body); + if let Some(code) = normalize_country_code(text.as_ref()) { info!(region = %code, ip = %ip, source = "ipinfo.io", "detected region"); - return Some(code.to_string()); + return Some(code); } } } } // Fallback: ip-api.com (HTTP only on free tier, non-sensitive data) - let http_url = format!("http://ip-api.com/json/{}?fields=countryCode", ip); + let http_url = ip_api_url(&ip)?; match client.get(&http_url).send().await { Ok(resp) if resp.status().is_success() => { - let body: serde_json::Value = resp.json().await.ok()?; - let code = body.get("countryCode")?.as_str()?; - if code.is_empty() { - return None; - } + let body = read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES) + .await + .ok()?; + let body: serde_json::Value = serde_json::from_slice(&body).ok()?; + let code = normalize_country_code(body.get("countryCode")?.as_str()?)?; info!(region = %code, ip = %ip, source = "ip-api.com", "detected region"); - Some(code.to_string()) + Some(code) } _ => { debug!(ip = %ip, "region detection failed"); @@ -88,3 +120,85 @@ pub async fn detect_region(ip: &str) -> Option { } } } + +fn parse_ip(value: &str) -> Option { + let value = value.trim(); + if value.is_empty() || value.len() > 45 { + return None; + } + value.parse().ok() +} + +fn normalize_country_code(value: &str) -> Option { + let value = value.trim(); + if !(2..=3).contains(&value.len()) || !value.bytes().all(|byte| byte.is_ascii_alphabetic()) { + return None; + } + Some(value.to_ascii_uppercase()) +} + +fn ipinfo_url(ip: &str) -> Option { + let mut url = Url::parse("https://ipinfo.io").ok()?; + url.path_segments_mut().ok()?.push(ip).push("country"); + Some(url.into()) +} + +fn ip_api_url(ip: &str) -> Option { + let mut url = Url::parse("http://ip-api.com").ok()?; + url.path_segments_mut().ok()?.push("json").push(ip); + url.query_pairs_mut().append_pair("fields", "countryCode"); + Some(url.into()) +} + +#[cfg(test)] +mod tests { + use super::{ip_api_url, ipinfo_url, normalize_country_code, parse_ip}; + + #[test] + fn parses_only_bounded_ip_values() { + assert_eq!(parse_ip(" 8.8.8.8\n"), Some("8.8.8.8".parse().unwrap())); + assert_eq!( + parse_ip("2001:4860:4860::8888"), + Some("2001:4860:4860::8888".parse().unwrap()) + ); + for value in ["", "8.8.8.8?x=1", "8.8.8.8\nX-Injected: yes", "not-an-ip"] { + assert!(parse_ip(value).is_none(), "accepted invalid IP: {value:?}"); + } + } + + #[test] + fn builds_urls_without_allowing_path_or_query_injection() { + let ip = "2001:4860:4860::8888"; + let info = ipinfo_url(ip).expect("fixed URL should parse"); + let info = url::Url::parse(&info).unwrap(); + assert_eq!(info.host_str(), Some("ipinfo.io")); + assert_eq!( + info.path_segments().unwrap().collect::>(), + ["2001:4860:4860::8888", "country"] + ); + assert_eq!(info.query(), None); + assert_eq!(info.fragment(), None); + + let api = ip_api_url(ip).expect("fixed URL should parse"); + let api = url::Url::parse(&api).unwrap(); + assert_eq!(api.host_str(), Some("ip-api.com")); + assert_eq!(api.query(), Some("fields=countryCode")); + assert_eq!( + api.path_segments().unwrap().collect::>(), + ["json", "2001:4860:4860::8888"] + ); + assert_eq!(api.fragment(), None); + } + + #[test] + fn country_codes_are_ascii_and_normalized() { + assert_eq!(normalize_country_code(" us\n"), Some("US".to_string())); + assert_eq!(normalize_country_code("GBR"), Some("GBR".to_string())); + for value in ["", "U", "US!", "US\nX", "中国"] { + assert!( + normalize_country_code(value).is_none(), + "accepted {value:?}" + ); + } + } +} diff --git a/apps/aether-tunnel/src/registration/client.rs b/apps/aether-tunnel/src/registration/client.rs index f28c0621e..64e0505a1 100644 --- a/apps/aether-tunnel/src/registration/client.rs +++ b/apps/aether-tunnel/src/registration/client.rs @@ -1,13 +1,21 @@ -use aether_http::{build_http_client, jittered_delay_for_retry, HttpClientConfig, HttpRetryConfig}; +use aether_http::{ + apply_http_client_config, jittered_delay_for_retry, read_response_bytes_with_limit, + HttpClientConfig, HttpRetryConfig, ResponseBodyReadError, +}; use aether_runtime::summarize_text_payload; use reqwest::{Client, StatusCode}; use serde::{Deserialize, Serialize}; use tokio::time::sleep; use tracing::{debug, error, info}; -use crate::config::{effective_tunnel_security, Config, ServerEntry, TunnelSecurity}; +use crate::config::{ + aether_url_for_log, effective_tunnel_security, validate_aether_url, Config, ServerEntry, + TunnelSecurity, +}; use crate::hardware::HardwareInfo; +const MAX_AETHER_CONTROL_RESPONSE_BYTES: usize = 256 * 1024; + #[derive(Debug, Serialize)] struct RegisterRequest { name: String, @@ -32,6 +40,7 @@ struct RegisterRequest { #[derive(Debug, Deserialize)] pub struct RegisterResponse { pub node_id: String, + pub tunnel_generation: String, } /// Remote configuration pushed by the Aether management backend. @@ -58,7 +67,7 @@ pub struct AetherClient { impl AetherClient { pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self { - let http = build_http_client(&HttpClientConfig { + let client_config = HttpClientConfig { connect_timeout_ms: Some(config.aether_connect_timeout_secs.saturating_mul(1_000)), request_timeout_ms: Some(config.aether_request_timeout_secs.saturating_mul(1_000)), pool_idle_timeout_ms: Some(config.aether_pool_idle_timeout_secs.saturating_mul(1_000)), @@ -75,8 +84,18 @@ impl AetherClient { .effective_aether_outbound_proxy_url() .map(str::to_string), ..HttpClientConfig::default() - }) - .expect("failed to create HTTP client"); + }; + let mut builder = apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()), + &client_config, + ); + if let Some(proxy_url) = client_config.proxy_url.as_deref() { + builder = builder + .proxy(reqwest::Proxy::all(proxy_url).expect("invalid Aether outbound proxy URL")); + } + let http = builder.build().expect("failed to create HTTP client"); let retry = HttpRetryConfig { max_attempts: config.aether_retry_max_attempts, @@ -87,7 +106,7 @@ impl AetherClient { Self { http, - base_url: aether_url.trim_end_matches('/').to_string(), + base_url: aether_url.trim().trim_end_matches('/').to_string(), token: management_token.to_string(), retry, } @@ -95,7 +114,7 @@ impl AetherClient { /// Register this node with Aether (idempotent upsert by ip:port). /// - /// Returns the stable node_id assigned by Aether. + /// Returns the stable node identity and the current server-issued tunnel generation. pub async fn register( &self, config: &Config, @@ -103,7 +122,8 @@ impl AetherClient { node_name: &str, public_ip: &str, hw: Option<&HardwareInfo>, - ) -> anyhow::Result { + ) -> anyhow::Result { + validate_aether_url(&self.base_url)?; let url = format!("{}/api/admin/proxy-nodes/register", self.base_url); let effective_security = effective_tunnel_security( &server.aether_url, @@ -130,7 +150,7 @@ impl AetherClient { }; info!( - url = %url, + url = %aether_url_for_log(&url), name = %body.name, ip = %body.ip, "registering with Aether" @@ -146,11 +166,22 @@ impl AetherClient { }, "register", ) - .await?; + .await + .map_err(|error| { + anyhow::anyhow!("register request failed ({})", reqwest_error_kind(&error)) + })?; let status = resp.status(); if !status.is_success() { - let text = resp.text().await.unwrap_or_default(); + let body = read_response_bytes_with_limit(resp, MAX_AETHER_CONTROL_RESPONSE_BYTES) + .await + .map_err(|error| { + anyhow::anyhow!( + "register failed (HTTP {status}): {}", + response_body_error_kind(&error) + ) + })?; + let text = String::from_utf8_lossy(&body); let summary = summarize_text_payload(&text); anyhow::bail!( "register failed (HTTP {}): response body redacted (bytes={}, sha256={})", @@ -160,13 +191,23 @@ impl AetherClient { ); } - let data: RegisterResponse = resp.json().await?; - info!(node_id = %data.node_id, "registered successfully"); - Ok(data.node_id) + let body = read_response_bytes_with_limit(resp, MAX_AETHER_CONTROL_RESPONSE_BYTES) + .await + .map_err(|error| { + anyhow::anyhow!( + "failed to read Aether registration response ({})", + response_body_error_kind(&error) + ) + })?; + let data: RegisterResponse = serde_json::from_slice(&body)?; + validate_registration_binding(&data)?; + info!(node_id = %data.node_id, tunnel_generation = %data.tunnel_generation, "registered successfully"); + Ok(data) } /// Unregister this node from Aether (graceful shutdown). pub async fn unregister(&self, node_id: &str) -> anyhow::Result<()> { + validate_aether_url(&self.base_url)?; let url = format!("{}/api/admin/proxy-nodes/unregister", self.base_url); let body = UnregisterRequest { node_id: node_id.to_string(), @@ -193,7 +234,15 @@ impl AetherClient { } Ok(r) => { let status = r.status(); - let text = r.text().await.unwrap_or_default(); + let body = read_response_bytes_with_limit(r, MAX_AETHER_CONTROL_RESPONSE_BYTES) + .await + .map_err(|error| { + anyhow::anyhow!( + "unregister failed (HTTP {status}): {}", + response_body_error_kind(&error) + ) + })?; + let text = String::from_utf8_lossy(&body); let summary = summarize_text_payload(&text); error!( status = %status, @@ -210,8 +259,9 @@ impl AetherClient { } Err(e) => { // Best-effort during shutdown - error!(error = %e, "unregister request failed"); - anyhow::bail!("unregister request failed: {}", e); + let kind = reqwest_error_kind(&e); + error!(error_kind = kind, "unregister request failed"); + anyhow::bail!("unregister request failed ({kind})"); } } } @@ -250,7 +300,7 @@ impl AetherClient { let sleep_for = jittered_delay_for_retry(self.retry, attempt - 1); debug!( attempt, - error = %e, + error_kind = reqwest_error_kind(&e), sleep_ms = sleep_for.as_millis(), label, "Aether request retrying" @@ -265,6 +315,55 @@ impl AetherClient { } } +fn validate_registration_binding(response: &RegisterResponse) -> anyhow::Result<()> { + for (field, value) in [ + ("node_id", response.node_id.as_str()), + ("tunnel_generation", response.tunnel_generation.as_str()), + ] { + if value.is_empty() + || value.len() > 128 + || !value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_')) + { + anyhow::bail!("registration response contains an invalid {field}"); + } + } + Ok(()) +} + +/// Reqwest error messages can include the complete request URL. Registration +/// URLs are derived from configuration, whose path may contain an operator +/// secret, so expose only a stable transport category at this boundary. +fn reqwest_error_kind(error: &reqwest::Error) -> &'static str { + if error.is_timeout() { + "timeout" + } else if error.is_connect() { + "connect" + } else if error.is_request() { + "request" + } else if error.is_redirect() { + "redirect" + } else if error.is_body() { + "body" + } else if error.is_decode() { + "decode" + } else { + "transport" + } +} + +fn response_body_error_kind(error: &ResponseBodyReadError) -> String { + match error { + ResponseBodyReadError::TooLarge { max_bytes } => { + format!("response too large (max {max_bytes} bytes)") + } + ResponseBodyReadError::Read(error) => { + format!("response read failed ({})", reqwest_error_kind(error)) + } + } +} + fn should_retry_status(status: StatusCode) -> bool { status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS diff --git a/apps/aether-tunnel/src/safe_dns.rs b/apps/aether-tunnel/src/safe_dns.rs deleted file mode 100644 index c46413a85..000000000 --- a/apps/aether-tunnel/src/safe_dns.rs +++ /dev/null @@ -1,66 +0,0 @@ -//! Safe DNS resolver for reqwest that reuses validated addresses from DnsCache. -//! -//! This resolver ensures reqwest connects only to addresses that have been -//! previously validated by `target_filter::validate_target()`, eliminating -//! the TOCTTOU gap where DNS rebinding could redirect traffic to private IPs. - -use std::net::SocketAddr; -use std::sync::Arc; - -use reqwest::dns::{Addrs, Name, Resolve, Resolving}; - -use crate::target_filter::{self, DnsCache}; - -/// A DNS resolver that serves validated public addresses from the shared DnsCache. -/// -/// When reqwest needs to resolve a hostname, this resolver returns addresses -/// from the cache (populated by `validate_target()` during request validation). -/// If the hostname is not in cache (shouldn't happen in normal flow), it -/// performs a fresh resolution with private-IP filtering. -pub struct SafeDnsResolver { - dns_cache: Arc, -} - -impl SafeDnsResolver { - pub fn new(dns_cache: Arc) -> Self { - Self { dns_cache } - } -} - -impl Resolve for SafeDnsResolver { - fn resolve(&self, name: Name) -> Resolving { - let dns_cache = Arc::clone(&self.dns_cache); - Box::pin(async move { - let host = name.as_str(); - - // Try cache first (should be populated by validate_target). - // reqwest resolves by hostname only (no port), so use host-only lookup. - if let Some(addrs) = dns_cache.get_by_host(host).await { - let socket_addrs: Vec = (*addrs).clone(); - return Ok(Box::new(socket_addrs.into_iter()) as Addrs); - } - - // Fallback: resolve with private-IP filtering (defensive). - // This path should rarely be hit since validate_target() runs first. - // We don't know the real port here (reqwest Resolve only gives hostname), - // so resolve directly without caching to avoid polluting the cache with - // an incorrect port-based key. - let addr_str = format!("{}:0", host); - let resolved: Vec = tokio::net::lookup_host(&addr_str) - .await - .map_err(|e| -> Box { Box::new(e) })? - .filter(|addr| !target_filter::is_private_ip(&addr.ip())) - .collect(); - - if resolved.is_empty() { - return Err(Box::new(std::io::Error::other(format!( - "all resolved addresses for {} are private/reserved", - host - ))) - as Box); - } - - Ok(Box::new(resolved.into_iter()) as Addrs) - }) - } -} diff --git a/apps/aether-tunnel/src/setup/service.rs b/apps/aether-tunnel/src/setup/service.rs index 7d3bbfbd9..e4e5f1f8c 100644 --- a/apps/aether-tunnel/src/setup/service.rs +++ b/apps/aether-tunnel/src/setup/service.rs @@ -3,8 +3,11 @@ //! Supports the host-native service manager we currently target: //! `systemd` on most Linux distributions and `OpenRC` on Alpine. +#[cfg(unix)] use std::fs::OpenOptions; use std::io::ErrorKind; +#[cfg(unix)] +use std::io::Write; use std::path::Path; use std::process::{Command, ExitStatus, Stdio}; @@ -18,6 +21,16 @@ const OPENRC_LOG_DIR: &str = "/var/log/aether-tunnel"; const OPENRC_STDOUT_LOG: &str = "/var/log/aether-tunnel/current.log"; const OPENRC_STDERR_LOG: &str = "/var/log/aether-tunnel/error.log"; +const SYSTEMCTL_BINS: &[&str] = &[ + "/usr/bin/systemctl", + "/bin/systemctl", + "/run/current-system/sw/bin/systemctl", +]; +const JOURNALCTL_BINS: &[&str] = &[ + "/usr/bin/journalctl", + "/bin/journalctl", + "/run/current-system/sw/bin/journalctl", +]; const OPENRC_RUN_BINS: &[&str] = &["/sbin/openrc-run", "/usr/sbin/openrc-run", "openrc-run"]; const OPENRC_SERVICE_BINS: &[&str] = &["/sbin/rc-service", "/usr/sbin/rc-service", "rc-service"]; const OPENRC_UPDATE_BINS: &[&str] = &["/sbin/rc-update", "/usr/sbin/rc-update", "rc-update"]; @@ -143,7 +156,7 @@ pub fn cmd_logs() -> anyhow::Result<()> { ensure_openrc_logs_readable()?; } let status = match manager { - ServiceManager::Systemd => Command::new("journalctl") + ServiceManager::Systemd => Command::new(journalctl_bin()) .args(["-u", SERVICE_NAME, "-f", "--no-pager", "-n", "100"]) .status()?, ServiceManager::OpenRc => Command::new(tail_bin()) @@ -281,9 +294,15 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> { .to_str() .unwrap_or("/"); + validate_service_unit_path(exe_str, "binary")?; + validate_service_unit_path(config_str, "config")?; + validate_service_unit_path(working_dir, "working directory")?; + validate_root_managed_service_file(&exe_path, "binary", false)?; + validate_root_managed_service_file(&config_abs, "config", true)?; + if Path::new(SYSTEMD_UNIT_PATH).exists() { eprintln!(" Stopping existing service..."); - let _ = Command::new("systemctl") + let _ = Command::new(systemctl_bin()) .args(["stop", SERVICE_NAME]) .status(); } @@ -293,34 +312,12 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> { eprintln!(" Config: {}", config_str); eprintln!(" WorkDir: {}", working_dir); - let unit_content = format!( - "[Unit]\n\ - Description=Aether Tunnel\n\ - After=network.target\n\ - \n\ - [Service]\n\ - Type=simple\n\ - WorkingDirectory={working_dir}\n\ - Environment=AETHER_TUNNEL_CONFIG={config_str}\n\ - Environment=AETHER_TUNNEL_SERVICE_MANAGER=systemd\n\ - Environment=AETHER_TUNNEL_LOG_DESTINATION=both\n\ - Environment=AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel\n\ - ExecStart={exe_str}\n\ - Restart=on-failure\n\ - RestartSec=5\n\ - LimitNOFILE=65535\n\ - UMask=0077\n\ - LogsDirectory=aether-tunnel\n\ - LogsDirectoryMode=0750\n\ - \n\ - [Install]\n\ - WantedBy=multi-user.target\n", - ); - std::fs::write(SYSTEMD_UNIT_PATH, &unit_content)?; + let unit_content = render_systemd_unit(exe_str, config_str, working_dir)?; + write_service_definition(SYSTEMD_UNIT_PATH, &unit_content, 0o644)?; eprintln!(" Enabling and starting service..."); - run_cmd("systemctl", &["daemon-reload"])?; - run_cmd("systemctl", &["enable", "--now", SERVICE_NAME])?; + run_cmd(systemctl_bin(), &["daemon-reload"])?; + run_cmd(systemctl_bin(), &["enable", "--now", SERVICE_NAME])?; eprintln!(); if manager_is_active(ServiceManager::Systemd) { @@ -333,6 +330,43 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> { Ok(()) } +fn render_systemd_unit( + exe_path: &str, + config_path: &str, + working_dir: &str, +) -> anyhow::Result { + validate_service_unit_path(exe_path, "binary")?; + validate_service_unit_path(config_path, "config")?; + validate_service_unit_path(working_dir, "working directory")?; + + let exe_path = systemd_quote(exe_path); + let working_dir = systemd_quote(working_dir); + let config_env = systemd_quote(&format!("AETHER_TUNNEL_CONFIG={config_path}")); + Ok(format!( + "[Unit]\n\ + Description=Aether Tunnel\n\ + After=network.target\n\ + \n\ + [Service]\n\ + Type=simple\n\ + WorkingDirectory={working_dir}\n\ + Environment={config_env}\n\ + Environment=AETHER_TUNNEL_SERVICE_MANAGER=systemd\n\ + Environment=AETHER_TUNNEL_LOG_DESTINATION=both\n\ + Environment=AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel\n\ + ExecStart={exe_path}\n\ + Restart=on-failure\n\ + RestartSec=5\n\ + LimitNOFILE=65535\n\ + UMask=0077\n\ + LogsDirectory=aether-tunnel\n\ + LogsDirectoryMode=0750\n\ + \n\ + [Install]\n\ + WantedBy=multi-user.target\n", + )) +} + fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> { let exe_path = std::env::current_exe()?.canonicalize()?; let exe_str = exe_path @@ -350,6 +384,12 @@ fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> { .to_str() .unwrap_or("/"); + validate_service_unit_path(exe_str, "binary")?; + validate_service_unit_path(config_str, "config")?; + validate_service_unit_path(working_dir, "working directory")?; + validate_root_managed_service_file(&exe_path, "binary", false)?; + validate_root_managed_service_file(&config_abs, "config", true)?; + if Path::new(OPENRC_INIT_PATH).exists() { eprintln!(" Stopping existing service..."); let _ = Command::new(openrc_service_bin()) @@ -357,12 +397,9 @@ fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> { .status(); } - std::fs::create_dir_all(OPENRC_LOG_DIR)?; - touch_log(OPENRC_STDOUT_LOG)?; - touch_log(OPENRC_STDERR_LOG)?; - set_mode(OPENRC_LOG_DIR, 0o750)?; - set_mode(OPENRC_STDOUT_LOG, 0o640)?; - set_mode(OPENRC_STDERR_LOG, 0o640)?; + ensure_private_service_directory(Path::new(OPENRC_LOG_DIR), 0o750)?; + open_private_service_log(Path::new(OPENRC_STDOUT_LOG), 0o640)?; + open_private_service_log(Path::new(OPENRC_STDERR_LOG), 0o640)?; eprintln!(" Generating OpenRC init script..."); eprintln!(" Binary: {}", exe_str); @@ -439,8 +476,7 @@ stop() {{ shell_quote("AETHER_TUNNEL_LOG_DESTINATION=both"), shell_quote(&format!("AETHER_TUNNEL_LOG_DIR={OPENRC_LOG_DIR}")), ); - std::fs::write(OPENRC_INIT_PATH, &init_content)?; - set_mode(OPENRC_INIT_PATH, 0o755)?; + write_service_definition(OPENRC_INIT_PATH, &init_content, 0o755)?; eprintln!(" Enabling and starting service..."); run_cmd(openrc_update_bin(), &["add", SERVICE_NAME, "default"])?; @@ -459,7 +495,7 @@ stop() {{ fn uninstall_systemd_service() -> anyhow::Result<()> { eprintln!(" Stopping and removing existing service..."); - let _ = Command::new("systemctl") + let _ = Command::new(systemctl_bin()) .args(["disable", "--now", SERVICE_NAME]) .status(); @@ -468,7 +504,7 @@ fn uninstall_systemd_service() -> anyhow::Result<()> { eprintln!(" Removed {}", SYSTEMD_UNIT_PATH); } - run_cmd("systemctl", &["daemon-reload"])?; + run_cmd(systemctl_bin(), &["daemon-reload"])?; eprintln!(" Service uninstalled."); Ok(()) } @@ -493,28 +529,28 @@ fn uninstall_openrc_service() -> anyhow::Result<()> { fn start_manager(manager: ServiceManager) -> anyhow::Result<()> { match manager { - ServiceManager::Systemd => run_cmd("systemctl", &["start", SERVICE_NAME]), + ServiceManager::Systemd => run_cmd(systemctl_bin(), &["start", SERVICE_NAME]), ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "start"]), } } fn stop_manager(manager: ServiceManager) -> anyhow::Result<()> { match manager { - ServiceManager::Systemd => run_cmd("systemctl", &["stop", SERVICE_NAME]), + ServiceManager::Systemd => run_cmd(systemctl_bin(), &["stop", SERVICE_NAME]), ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "stop"]), } } fn restart_manager(manager: ServiceManager) -> anyhow::Result<()> { match manager { - ServiceManager::Systemd => run_cmd("systemctl", &["restart", SERVICE_NAME]), + ServiceManager::Systemd => run_cmd(systemctl_bin(), &["restart", SERVICE_NAME]), ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "restart"]), } } fn manager_status(manager: ServiceManager) -> anyhow::Result { let status = match manager { - ServiceManager::Systemd => Command::new("systemctl") + ServiceManager::Systemd => Command::new(systemctl_bin()) .args(["status", SERVICE_NAME]) .status()?, ServiceManager::OpenRc => Command::new(openrc_service_bin()) @@ -528,7 +564,7 @@ fn manager_is_active(manager: ServiceManager) -> bool { match manager { ServiceManager::Systemd => { Path::new(SYSTEMD_UNIT_PATH).exists() - && Command::new("systemctl") + && Command::new(systemctl_bin()) .args(["is-active", "--quiet", SERVICE_NAME]) .stdout(Stdio::null()) .stderr(Stdio::null()) @@ -562,7 +598,8 @@ fn print_post_install_commands() { fn is_systemd_available() -> bool { Path::new("/run/systemd/system").exists() - && Command::new("systemctl") + && has_absolute_candidate(SYSTEMCTL_BINS) + && Command::new(systemctl_bin()) .arg("--version") .stdout(Stdio::null()) .stderr(Stdio::null()) @@ -610,28 +647,446 @@ fn pick_bin(candidates: &[&'static str]) -> &'static str { .iter() .copied() .find(|candidate| candidate.starts_with('/') && Path::new(candidate).exists()) - .unwrap_or_else(|| candidates[candidates.len() - 1]) + .or_else(|| { + candidates + .iter() + .copied() + .find(|candidate| candidate.starts_with('/')) + }) + .expect("trusted binary candidate list must include an absolute path") +} + +fn systemctl_bin() -> &'static str { + pick_bin(SYSTEMCTL_BINS) +} + +fn journalctl_bin() -> &'static str { + pick_bin(JOURNALCTL_BINS) +} + +fn validate_service_unit_path(value: &str, label: &str) -> anyhow::Result<()> { + if !Path::new(value).is_absolute() { + anyhow::bail!("{} path must be absolute", label); + } + if value.chars().any(char::is_control) || value.contains(['%', '$']) { + anyhow::bail!( + "{} path contains control or service-manager expansion characters", + label + ); + } + Ok(()) +} + +fn validate_root_managed_service_file( + path: &Path, + label: &str, + require_private_file: bool, +) -> anyhow::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + let file_metadata = std::fs::symlink_metadata(path)?; + if !file_metadata.is_file() || file_metadata.file_type().is_symlink() { + anyhow::bail!("service {} must be a regular non-symlink file", label); + } + if file_metadata.uid() != 0 || file_metadata.mode() & 0o022 != 0 { + anyhow::bail!( + "service {} must be owned by root and not writable by group or other users; use the official installer or move it to a root-managed path", + label + ); + } + if require_private_file && (file_metadata.mode() & 0o077 != 0 || file_metadata.nlink() != 1) + { + anyhow::bail!( + "service {} contains credentials and must be owner-only with exactly one hard link", + label + ); + } + + let mut ancestor = path.parent(); + while let Some(directory) = ancestor { + let metadata = std::fs::symlink_metadata(directory)?; + if !metadata.is_dir() + || metadata.file_type().is_symlink() + || metadata.uid() != 0 + || metadata.mode() & 0o022 != 0 + { + anyhow::bail!( + "service {} parent '{}' must be a root-owned directory that is not writable by group or other users", + label, + directory.display() + ); + } + ancestor = directory.parent(); + } + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = (path, label, require_private_file); + anyhow::bail!("managed tunnel services require Unix ownership checks") + } +} + +fn systemd_quote(value: &str) -> String { + format!("\"{}\"", value.replace('\\', "\\\\").replace('\"', "\\\"")) +} + +fn write_service_definition(path: &str, content: &str, mode: u32) -> anyhow::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; + + let requested_path = Path::new(path); + let requested_parent = requested_path + .parent() + .ok_or_else(|| anyhow::anyhow!("service definition path has no parent"))?; + let file_name = requested_path + .file_name() + .ok_or_else(|| anyhow::anyhow!("service definition path has no file name"))?; + let parent = std::fs::canonicalize(requested_parent)?; + let path = parent.join(file_name); + validate_private_service_directory(&parent)?; + validate_replaceable_service_file(&path)?; + + let temporary = parent.join(format!( + ".aether-tunnel-service-{}-{}.tmp", + std::process::id(), + uuid::Uuid::new_v4() + )); + let mut options = OpenOptions::new(); + options + .write(true) + .create_new(true) + .mode(mode) + .custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); + let mut file = options.open(&temporary)?; + let result = (|| -> anyhow::Result<()> { + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + let metadata = file.metadata()?; + if !metadata.is_file() || metadata.uid() != effective_uid || metadata.nlink() != 1 { + anyhow::bail!("temporary service definition has unsafe ownership or links"); + } + file.set_permissions(std::fs::Permissions::from_mode(mode))?; + file.write_all(content.as_bytes())?; + file.sync_all()?; + drop(file); + std::fs::rename(&temporary, &path)?; + std::fs::File::open(&parent)?.sync_all()?; + Ok(()) + })(); + if result.is_err() { + let _ = std::fs::remove_file(&temporary); + } + result + } + + #[cfg(not(unix))] + { + let _ = (path, content, mode); + anyhow::bail!("managed service definitions require Unix filesystem checks") + } } fn shell_quote(value: &str) -> String { format!("'{}'", value.replace('\'', "'\"'\"'")) } -fn touch_log(path: &str) -> anyhow::Result<()> { - OpenOptions::new().create(true).append(true).open(path)?; - Ok(()) -} - -fn set_mode(path: &str, mode: u32) -> anyhow::Result<()> { +fn validate_private_service_directory(path: &Path) -> anyhow::Result<()> { #[cfg(unix)] { - use std::os::unix::fs::PermissionsExt; - let perms = std::fs::Permissions::from_mode(mode); - std::fs::set_permissions(path, perms)?; + use std::os::unix::fs::MetadataExt; + + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + let metadata = std::fs::symlink_metadata(path)?; + if metadata.file_type().is_symlink() + || !metadata.is_dir() + || (metadata.uid() != effective_uid && metadata.uid() != 0) + || metadata.mode() & 0o022 != 0 + { + anyhow::bail!( + "service directory '{}' has unsafe ownership or permissions", + path.display() + ); + } + + let canonical = std::fs::canonicalize(path)?; + let mut ancestor = canonical.parent(); + while let Some(directory) = ancestor { + let metadata = std::fs::symlink_metadata(directory)?; + let mode = metadata.mode(); + if metadata.file_type().is_symlink() + || !metadata.is_dir() + || (metadata.uid() != effective_uid && metadata.uid() != 0) + || (mode & 0o022 != 0 && mode & 0o1000 == 0) + { + anyhow::bail!( + "service directory ancestor '{}' has unsafe ownership or permissions", + directory.display() + ); + } + ancestor = directory.parent(); + } + Ok(()) } #[cfg(not(unix))] - let _ = (path, mode); + { + let _ = path; + anyhow::bail!("managed service directories require Unix filesystem checks") + } +} - Ok(()) +fn validate_replaceable_service_file(path: &Path) -> anyhow::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + match std::fs::symlink_metadata(path) { + Ok(metadata) => { + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + if metadata.file_type().is_symlink() + || !metadata.is_file() + || metadata.uid() != effective_uid + || metadata.nlink() != 1 + { + anyhow::bail!( + "service file '{}' must be a regular single-link file owned by the current user", + path.display() + ); + } + } + Err(error) if error.kind() == ErrorKind::NotFound => {} + Err(error) => return Err(error.into()), + } + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = path; + anyhow::bail!("managed service files require Unix filesystem checks") + } +} + +fn ensure_private_service_directory(path: &Path, mode: u32) -> anyhow::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::{DirBuilderExt, MetadataExt, OpenOptionsExt, PermissionsExt}; + + let requested_parent = path + .parent() + .ok_or_else(|| anyhow::anyhow!("service directory has no parent"))?; + let file_name = path + .file_name() + .ok_or_else(|| anyhow::anyhow!("service directory has no file name"))?; + let parent = std::fs::canonicalize(requested_parent)?; + let path = parent.join(file_name); + validate_private_service_directory(&parent)?; + match std::fs::symlink_metadata(&path) { + Ok(_) => {} + Err(error) if error.kind() == ErrorKind::NotFound => { + let mut builder = std::fs::DirBuilder::new(); + builder.mode(mode).create(&path)?; + } + Err(error) => return Err(error.into()), + } + + let mut options = OpenOptions::new(); + options.read(true).custom_flags( + libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_DIRECTORY | libc::O_NONBLOCK, + ); + let directory = options.open(&path)?; + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + let metadata = directory.metadata()?; + if !metadata.is_dir() || metadata.uid() != effective_uid || metadata.mode() & 0o022 != 0 { + anyhow::bail!("service log directory has unsafe ownership or permissions"); + } + directory.set_permissions(std::fs::Permissions::from_mode(mode))?; + directory.sync_all()?; + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = (path, mode); + anyhow::bail!("managed service directories require Unix filesystem checks") + } +} + +fn open_private_service_log(path: &Path, mode: u32) -> anyhow::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; + + let requested_parent = path + .parent() + .ok_or_else(|| anyhow::anyhow!("service log path has no parent"))?; + let file_name = path + .file_name() + .ok_or_else(|| anyhow::anyhow!("service log path has no file name"))?; + let parent = std::fs::canonicalize(requested_parent)?; + let path = parent.join(file_name); + validate_private_service_directory(&parent)?; + let mut options = OpenOptions::new(); + options + .create(true) + .append(true) + .mode(mode) + .custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_NONBLOCK); + let file = options.open(&path)?; + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + let metadata = file.metadata()?; + if !metadata.is_file() || metadata.uid() != effective_uid || metadata.nlink() != 1 { + anyhow::bail!("service log must be a regular, single-link file owned by root"); + } + file.set_permissions(std::fs::Permissions::from_mode(mode))?; + file.sync_all()?; + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = (path, mode); + anyhow::bail!("managed service logs require Unix filesystem checks") + } +} + +#[cfg(test)] +mod tests { + use super::{ + ensure_private_service_directory, open_private_service_log, pick_bin, render_systemd_unit, + systemd_quote, validate_root_managed_service_file, validate_service_unit_path, + write_service_definition, + }; + + #[test] + fn systemd_unit_quotes_paths_without_changing_arguments() { + let unit = render_systemd_unit( + r#"/opt/Aether Tunnel/aether\"tunnel"#, + r#"/var/lib/aether tunnel/config\\node.toml"#, + "/var/lib/aether tunnel", + ) + .expect("safe absolute paths should render"); + + assert!(unit.contains(r#"ExecStart="/opt/Aether Tunnel/aether\\\"tunnel""#)); + assert!(unit.contains( + r#"Environment="AETHER_TUNNEL_CONFIG=/var/lib/aether tunnel/config\\\\node.toml""# + )); + assert!(unit.contains(r#"WorkingDirectory="/var/lib/aether tunnel""#)); + assert_eq!(systemd_quote("a\\b\"c"), r#""a\\b\"c""#); + } + + #[test] + fn service_unit_paths_reject_directive_and_expansion_injection() { + for value in [ + "relative/path", + "/tmp/config\nExecStart=/tmp/evil", + "/tmp/config\rEnvironment=EVIL=1", + "/tmp/%n/config", + "/tmp/$PATH/config", + ] { + assert!( + validate_service_unit_path(value, "test").is_err(), + "accepted unsafe service path: {value:?}" + ); + } + } + + #[test] + fn service_commands_never_fall_back_to_path_lookup() { + assert_eq!( + pick_bin(&["relative-tool", "/definitely/missing/trusted-tool"]), + "/definitely/missing/trusted-tool" + ); + } + + #[cfg(unix)] + #[test] + fn root_service_rejects_files_beneath_shared_writable_directories() { + let directory = + std::env::temp_dir().join(format!("aether-service-path-test-{}", std::process::id())); + let _ = std::fs::remove_dir_all(&directory); + std::fs::create_dir(&directory).expect("test directory should be created"); + { + use std::os::unix::fs::PermissionsExt; + std::fs::set_permissions(&directory, std::fs::Permissions::from_mode(0o777)) + .expect("test directory permissions should be set"); + } + let path = directory.join("aether-tunnel"); + std::fs::write(&path, b"binary").expect("test binary should be written"); + + let result = validate_root_managed_service_file(&path, "binary", false); + let _ = std::fs::remove_dir_all(&directory); + assert!(result.is_err()); + } + + #[cfg(unix)] + #[test] + fn service_definitions_and_logs_refuse_links_and_use_private_atomic_files() { + use std::io::Read; + use std::os::unix::fs::{symlink, MetadataExt, PermissionsExt}; + + let directory = std::env::temp_dir().join(format!( + "aether-service-write-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&directory).unwrap(); + std::fs::set_permissions(&directory, std::fs::Permissions::from_mode(0o700)).unwrap(); + + let definition = directory.join("aether-tunnel.service"); + write_service_definition(definition.to_str().unwrap(), "first", 0o644).unwrap(); + let metadata = std::fs::symlink_metadata(&definition).unwrap(); + assert_eq!(metadata.mode() & 0o777, 0o644); + assert_eq!(metadata.nlink(), 1); + let mut old_definition = std::fs::File::open(&definition).unwrap(); + write_service_definition(definition.to_str().unwrap(), "second", 0o644).unwrap(); + let mut old_contents = String::new(); + old_definition.read_to_string(&mut old_contents).unwrap(); + assert_eq!(old_contents, "first"); + assert_eq!(std::fs::read_to_string(&definition).unwrap(), "second"); + + let victim = directory.join("victim"); + std::fs::write(&victim, b"known-good").unwrap(); + std::fs::remove_file(&definition).unwrap(); + symlink(&victim, &definition).unwrap(); + assert!(write_service_definition(definition.to_str().unwrap(), "replace", 0o644).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_file(&definition).unwrap(); + std::fs::hard_link(&victim, &definition).unwrap(); + assert!(write_service_definition(definition.to_str().unwrap(), "replace", 0o644).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + std::fs::remove_file(&definition).unwrap(); + + let log_directory = directory.join("logs"); + ensure_private_service_directory(&log_directory, 0o750).unwrap(); + assert_eq!( + std::fs::symlink_metadata(&log_directory).unwrap().mode() & 0o777, + 0o750 + ); + let log = log_directory.join("current.log"); + open_private_service_log(&log, 0o640).unwrap(); + let metadata = std::fs::symlink_metadata(&log).unwrap(); + assert_eq!(metadata.mode() & 0o777, 0o640); + assert_eq!(metadata.nlink(), 1); + std::fs::remove_file(&log).unwrap(); + symlink(&victim, &log).unwrap(); + assert!(open_private_service_log(&log, 0o640).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_file(&log).unwrap(); + std::fs::hard_link(&victim, &log).unwrap(); + assert!(open_private_service_log(&log, 0o640).is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_dir_all(directory).unwrap(); + } } diff --git a/apps/aether-tunnel/src/setup/tui.rs b/apps/aether-tunnel/src/setup/tui.rs index f56066285..7f8d9da3a 100644 --- a/apps/aether-tunnel/src/setup/tui.rs +++ b/apps/aether-tunnel/src/setup/tui.rs @@ -200,11 +200,11 @@ impl App { Field { label: "Allow Private Targets", key: "allow_private_targets", - value: "true".into(), + value: "false".into(), kind: FieldKind::Bool, required: false, help: - "Allow proxying private/reserved upstream IPs by default; takes effect after restart", + "Allow proxying private/reserved upstream IPs; disabled by default and takes effect after restart", }, Field { label: "Heartbeat Interval", @@ -450,13 +450,6 @@ impl App { fn save(&mut self) -> anyhow::Result<()> { let cfg = self.to_config()?; cfg.save(&self.config_path)?; - // Restrict config file permissions to owner-only (contains management token). - #[cfg(unix)] - { - use std::os::unix::fs::PermissionsExt; - let _ = - std::fs::set_permissions(&self.config_path, std::fs::Permissions::from_mode(0o600)); - } self.modified = false; self.saved_once = true; self.message = Some(( diff --git a/apps/aether-tunnel/src/setup/upgrade.rs b/apps/aether-tunnel/src/setup/upgrade.rs index d7c689488..0ae6f690d 100644 --- a/apps/aether-tunnel/src/setup/upgrade.rs +++ b/apps/aether-tunnel/src/setup/upgrade.rs @@ -4,21 +4,73 @@ //! running binary atomically, and restarts the active managed service when //! applicable. +use std::io::{Read, Write}; +use std::net::SocketAddr; use std::path::{Path, PathBuf}; +use std::sync::Arc; -use aether_http::{apply_http_client_config, HttpClientConfig}; +use aether_http::{apply_http_client_config, read_response_bytes_with_limit, HttpClientConfig}; +use reqwest::dns::{Addrs, Name, Resolve, Resolving}; use sha2::{Digest, Sha256}; const GITHUB_API_BASE: &str = "https://api.github.com"; const GITHUB_REPO: &str = "fawney19/Aether"; const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION"); +const MAX_GITHUB_RELEASE_METADATA_BYTES: usize = 8 * 1024 * 1024; +const MAX_GITHUB_ERROR_RESPONSE_BYTES: usize = 256 * 1024; +const MAX_RELEASE_ARCHIVE_DOWNLOAD_BYTES: usize = 128 * 1024 * 1024; +const MAX_CHECKSUM_DOWNLOAD_BYTES: usize = 1024 * 1024; + +fn summarize_remote_error_body(body: &[u8]) -> String { + let digest = Sha256::digest(body); + format!( + "response body redacted (bytes={}, sha256={})", + body.len(), + hex::encode(digest) + ) +} // ── GitHub API types ───────────────────────────────────────────────────────── #[derive(serde::Deserialize)] struct GithubRelease { tag_name: String, - name: String, + #[serde(default)] + name: Option, + #[serde(default)] + draft: bool, + #[serde(default)] + prerelease: bool, +} + +fn tunnel_release_semver(tag: &str) -> anyhow::Result { + if tag.len() > 160 + || !tag + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-' | b'+')) + { + anyhow::bail!("invalid tunnel release tag") + } + let version = tag + .strip_prefix("tunnel-v") + .or_else(|| tag.strip_prefix("proxy-v")) + .ok_or_else(|| anyhow::anyhow!("tunnel release tag has an unsupported prefix"))?; + semver::Version::parse(version) + .map_err(|_| anyhow::anyhow!("tunnel release tag is not valid semantic versioning")) +} + +fn normalize_requested_release_tag(version: &str) -> anyhow::Result { + let version = version.trim(); + if version.is_empty() { + anyhow::bail!("upgrade version must not be empty"); + } + let tag = if version.starts_with("tunnel-v") || version.starts_with("proxy-v") { + version.to_string() + } else { + format!("tunnel-v{version}") + }; + tunnel_release_semver(&tag)?; + Ok(tag) } // ── Platform detection ─────────────────────────────────────────────────────── @@ -50,7 +102,7 @@ fn detect_platform() -> &'static str { // ── GitHub HTTP client ─────────────────────────────────────────────────────── -fn build_github_client() -> anyhow::Result { +fn build_github_api_client() -> anyhow::Result { let mut headers = reqwest::header::HeaderMap::new(); if let Ok(token) = std::env::var("GITHUB_TOKEN") { @@ -66,7 +118,11 @@ fn build_github_client() -> anyhow::Result { ); Ok(apply_http_client_config( - reqwest::Client::builder().default_headers(headers), + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::none()) + .dns_resolver(Arc::new(SafeGithubDnsResolver)) + .default_headers(headers), &HttpClientConfig { request_timeout_ms: Some(300_000), user_agent: Some(format!("aether-tunnel/{}", CURRENT_VERSION)), @@ -76,6 +132,112 @@ fn build_github_client() -> anyhow::Result { .build()?) } +fn build_github_download_client() -> anyhow::Result { + Ok(apply_http_client_config( + reqwest::Client::builder() + .no_proxy() + .redirect(reqwest::redirect::Policy::custom(|attempt| { + if attempt.previous().len() >= 10 { + return attempt.error("too many GitHub release redirects"); + } + if is_trusted_github_download_url(attempt.url()) { + attempt.follow() + } else { + attempt.error("GitHub release redirected to an untrusted URL") + } + })) + .dns_resolver(Arc::new(SafeGithubDnsResolver)), + &HttpClientConfig { + request_timeout_ms: Some(300_000), + user_agent: Some(format!("aether-tunnel/{}", CURRENT_VERSION)), + ..HttpClientConfig::default() + }, + ) + .build()?) +} + +#[derive(Debug)] +struct SafeGithubDnsResolver; + +impl Resolve for SafeGithubDnsResolver { + fn resolve(&self, name: Name) -> Resolving { + let host = name.as_str().trim_end_matches('.').to_ascii_lowercase(); + // Transparent DNS interception may map public domains into RFC 2544's + // 198.18.0.0/15 benchmark range. This resolver is used exclusively for + // built-in GitHub update hosts, so permit that synthetic range only for + // the same trusted host set used by redirect validation. + let allow_benchmarking_ip = is_trusted_github_host(&host); + Box::pin(async move { + let addresses = aether_http::lookup_host_with_limits( + host.as_str(), + 0, + aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT, + ) + .await + .map_err(|error| -> Box { Box::new(error) })?; + validate_github_resolved_addrs_with_fake_ip(&addresses, allow_benchmarking_ip) + .map_err(|message| { + Box::new(std::io::Error::other(message)) + as Box + })?; + Ok(Box::new(addresses.into_iter()) as Addrs) + }) + } +} + +#[cfg(test)] +fn validate_github_resolved_addrs(addresses: &[SocketAddr]) -> Result<(), &'static str> { + validate_github_resolved_addrs_with_fake_ip(addresses, false) +} + +fn validate_github_resolved_addrs_with_fake_ip( + addresses: &[SocketAddr], + allow_benchmarking_ip: bool, +) -> Result<(), &'static str> { + if addresses.is_empty() { + return Err("GitHub DNS resolution returned no addresses"); + } + if addresses.iter().any(|address| { + aether_http::is_private_or_reserved_ip(address.ip()) + && !(allow_benchmarking_ip && aether_http::is_ipv4_benchmarking_fake_ip(address.ip())) + }) { + return Err("GitHub DNS resolution returned a private or reserved address"); + } + Ok(()) +} + +fn is_trusted_github_host(host: &str) -> bool { + host.eq_ignore_ascii_case("github.com") + || host.eq_ignore_ascii_case("api.github.com") + || host.eq_ignore_ascii_case("objects.githubusercontent.com") + || host.ends_with(".objects.githubusercontent.com") + || host.eq_ignore_ascii_case("release-assets.githubusercontent.com") + || host.ends_with(".release-assets.githubusercontent.com") +} + +fn is_trusted_github_download_host(host: &str) -> bool { + host.eq_ignore_ascii_case("github.com") + || host.eq_ignore_ascii_case("objects.githubusercontent.com") + || host.ends_with(".objects.githubusercontent.com") + || host.eq_ignore_ascii_case("release-assets.githubusercontent.com") + || host.ends_with(".release-assets.githubusercontent.com") +} + +fn is_trusted_github_download_url(url: &url::Url) -> bool { + if url.scheme() != "https" + || !url.username().is_empty() + || url.password().is_some() + || url.fragment().is_some() + || url.port_or_known_default() != Some(443) + { + return false; + } + let Some(host) = url.host_str() else { + return false; + }; + is_trusted_github_download_host(host) +} + // ── Release fetching ───────────────────────────────────────────────────────── async fn fetch_release( @@ -85,22 +247,37 @@ async fn fetch_release( match version { Some(ver) => { // Accept both "tunnel-v0.2.0" and the legacy "proxy-v0.2.0". - let tag = if ver.starts_with("tunnel-v") || ver.starts_with("proxy-v") { - ver.to_string() - } else { - format!("tunnel-v{}", ver) - }; + let tag = normalize_requested_release_tag(ver)?; let url = format!( "{}/repos/{}/releases/tags/{}", GITHUB_API_BASE, GITHUB_REPO, tag ); let resp = client.get(&url).send().await?; - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - anyhow::bail!("release '{}' not found (HTTP {}): {}", tag, status, body); + let status = resp.status(); + let max_bytes = if status.is_success() { + MAX_GITHUB_RELEASE_METADATA_BYTES + } else { + MAX_GITHUB_ERROR_RESPONSE_BYTES + }; + let body = read_response_bytes_with_limit(resp, max_bytes) + .await + .map_err(|error| { + anyhow::anyhow!("failed to read GitHub release response: {error}") + })?; + if !status.is_success() { + anyhow::bail!( + "release '{}' not found (HTTP {}): {}", + tag, + status, + summarize_remote_error_body(&body) + ); } - Ok(resp.json().await?) + let release: GithubRelease = serde_json::from_slice(&body)?; + if release.draft || release.tag_name != tag { + anyhow::bail!("GitHub returned a draft or mismatched tunnel release"); + } + tunnel_release_semver(&release.tag_name)?; + Ok(release) } None => { // List releases and find the latest tunnel-v* tag @@ -109,15 +286,32 @@ async fn fetch_release( GITHUB_API_BASE, GITHUB_REPO ); let resp = client.get(&url).send().await?; - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - anyhow::bail!("failed to list releases (HTTP {}): {}", status, body); + let status = resp.status(); + let max_bytes = if status.is_success() { + MAX_GITHUB_RELEASE_METADATA_BYTES + } else { + MAX_GITHUB_ERROR_RESPONSE_BYTES + }; + let body = read_response_bytes_with_limit(resp, max_bytes) + .await + .map_err(|error| { + anyhow::anyhow!("failed to read GitHub releases response: {error}") + })?; + if !status.is_success() { + anyhow::bail!( + "failed to list releases (HTTP {}): {}", + status, + summarize_remote_error_body(&body) + ); } - let releases: Vec = resp.json().await?; + let releases: Vec = serde_json::from_slice(&body)?; releases .into_iter() - .find(|r| r.tag_name.starts_with("tunnel-v") || r.tag_name.starts_with("proxy-v")) + .find(|release| { + !release.draft + && !release.prerelease + && tunnel_release_semver(&release.tag_name).is_ok() + }) .ok_or_else(|| anyhow::anyhow!("no tunnel-v* release found")) } } @@ -131,12 +325,22 @@ async fn download_release_file( client: &reqwest::Client, tag: &str, filename: &str, + max_bytes: usize, ) -> anyhow::Result> { + tunnel_release_semver(tag)?; + if filename.is_empty() + || filename.len() > 160 + || !filename + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-' | b'_')) + { + anyhow::bail!("invalid GitHub release asset name"); + } let url = format!( "https://github.com/{}/releases/download/{}/{}", GITHUB_REPO, tag, filename ); - let resp = client + let mut resp = client .get(&url) .header(reqwest::header::ACCEPT, "application/octet-stream") .send() @@ -148,21 +352,58 @@ async fn download_release_file( resp.status(), ); } - Ok(resp.bytes().await?.to_vec()) + if resp + .content_length() + .is_some_and(|length| length > max_bytes as u64) + { + anyhow::bail!("download for '{}' exceeds {} bytes", filename, max_bytes); + } + + let mut bytes = Vec::new(); + while let Some(chunk) = resp.chunk().await? { + append_bounded_download_chunk(&mut bytes, &chunk, max_bytes, filename)?; + } + Ok(bytes) +} + +fn append_bounded_download_chunk( + bytes: &mut Vec, + chunk: &[u8], + max_bytes: usize, + filename: &str, +) -> anyhow::Result<()> { + if chunk.len() > max_bytes.saturating_sub(bytes.len()) { + anyhow::bail!("download for '{}' exceeds {} bytes", filename, max_bytes); + } + bytes.extend_from_slice(chunk); + Ok(()) } fn parse_checksum(sums_text: &str, filename: &str) -> anyhow::Result { + let mut matches = Vec::new(); for line in sums_text.lines() { // Format: " " (GNU coreutils convention) let mut parts = line.split_ascii_whitespace(); let (Some(hash), Some(name)) = (parts.next(), parts.next()) else { continue; }; - if name == filename || name.ends_with(filename) { - return Ok(hash.to_lowercase()); + if parts.next().is_some() { + continue; + } + let name = name.strip_prefix('*').unwrap_or(name); + if name == filename && hash.len() == 64 && hash.bytes().all(|byte| byte.is_ascii_hexdigit()) + { + matches.push(hash.to_ascii_lowercase()); } } - anyhow::bail!("checksum for '{}' not found in SHA256SUMS.txt", filename); + match matches.as_slice() { + [hash] => Ok(hash.clone()), + [] => anyhow::bail!("checksum for '{}' not found in SHA256SUMS.txt", filename), + _ => anyhow::bail!( + "SHA256SUMS.txt contains multiple valid entries for '{}'", + filename + ), + } } async fn download_and_verify( @@ -175,8 +416,13 @@ async fn download_and_verify( eprintln!(" Downloading {}...", archive_name); let (archive_bytes, checksum_bytes) = tokio::try_join!( - download_release_file(client, tag, &archive_name), - download_release_file(client, tag, "SHA256SUMS.txt"), + download_release_file( + client, + tag, + &archive_name, + MAX_RELEASE_ARCHIVE_DOWNLOAD_BYTES, + ), + download_release_file(client, tag, "SHA256SUMS.txt", MAX_CHECKSUM_DOWNLOAD_BYTES,), )?; let checksum_text = String::from_utf8(checksum_bytes)?; @@ -225,71 +471,335 @@ fn extract_binary(archive_bytes: &[u8], dest: &Path) -> anyhow::Result<()> { "aether-tunnel" }; - for entry in archive.entries()? { - let mut entry = entry?; - // Only accept regular files -- reject symlinks to prevent write-through attacks - if entry.header().entry_type() != tar::EntryType::Regular { - continue; + let mut entries = archive.entries()?; + let mut entry = entries + .next() + .transpose()? + .ok_or_else(|| anyhow::anyhow!("release archive is empty"))?; + if entry.header().entry_type() != tar::EntryType::Regular { + anyhow::bail!("release archive entry is not a regular file"); + } + let path = entry.path()?; + if path.as_ref() != Path::new(binary_name) { + anyhow::bail!( + "release archive must contain only '{}' at its root", + binary_name + ); + } + let size = entry.header().size()?; + if size == 0 || size > MAX_BINARY_SIZE { + anyhow::bail!( + "invalid binary size ({} bytes, expected 1..={} bytes)", + size, + MAX_BINARY_SIZE + ); + } + + let mut binary = Vec::with_capacity(size as usize); + entry + .by_ref() + .take(MAX_BINARY_SIZE + 1) + .read_to_end(&mut binary)?; + if binary.len() as u64 != size { + anyhow::bail!("release archive binary size does not match its header"); + } + drop(entry); + if entries.next().transpose()?.is_some() { + anyhow::bail!("release archive must contain exactly one entry"); + } + + let mut options = std::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(0o700); + } + let mut file = options.open(dest).map_err(|error| { + anyhow::anyhow!( + "refusing to overwrite upgrade staging path '{}': {}", + dest.display(), + error + ) + })?; + let write_result = (|| -> std::io::Result<()> { + file.write_all(&binary)?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + file.set_permissions(std::fs::Permissions::from_mode(0o755))?; } - let path = entry.path()?; - if path.file_name().and_then(|n| n.to_str()) == Some(binary_name) { - let size = entry.header().size()?; - if size > MAX_BINARY_SIZE { + file.sync_all() + })(); + drop(file); + if let Err(error) = write_result { + remove_upgrade_file_if_regular(dest); + return Err(error.into()); + } + + Ok(()) +} + +fn remove_upgrade_file_if_regular(path: &Path) { + if std::fs::symlink_metadata(path) + .is_ok_and(|metadata| !metadata.file_type().is_symlink() && metadata.is_file()) + { + let _ = std::fs::remove_file(path); + } +} + +fn validate_upgrade_storage(current_exe: &Path) -> anyhow::Result<()> { + let current_metadata = std::fs::symlink_metadata(current_exe)?; + if current_metadata.file_type().is_symlink() || !current_metadata.is_file() { + anyhow::bail!("current tunnel executable must be a regular file"); + } + let exe_dir = current_exe + .parent() + .ok_or_else(|| anyhow::anyhow!("cannot determine binary directory"))?; + let directory_metadata = std::fs::symlink_metadata(exe_dir)?; + if directory_metadata.file_type().is_symlink() || !directory_metadata.is_dir() { + anyhow::bail!("tunnel executable directory must be a real directory"); + } + + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + if current_metadata.uid() != effective_uid + || current_metadata.mode() & 0o7022 != 0 + || current_metadata.mode() & 0o100 == 0 + || current_metadata.nlink() != 1 + { + anyhow::bail!("current tunnel executable ownership or permissions are unsafe"); + } + let directory_mode = directory_metadata.mode(); + if directory_metadata.uid() != effective_uid + || directory_mode & 0o022 != 0 + || ((directory_mode >> 6) & 0o3) != 0o3 + { + anyhow::bail!("tunnel executable directory ownership or permissions are unsafe"); + } + + // The immediate directory is protected above. Every ancestor must also be + // controlled by this user or root; shared writable ancestors are accepted + // only when the sticky bit prevents unrelated users from replacing entries. + let canonical_exe_dir = std::fs::canonicalize(exe_dir)?; + let mut ancestor = canonical_exe_dir.parent(); + while let Some(directory) = ancestor { + let metadata = std::fs::symlink_metadata(directory)?; + let mode = metadata.mode(); + if metadata.file_type().is_symlink() + || !metadata.is_dir() + || (metadata.uid() != effective_uid && metadata.uid() != 0) + || (mode & 0o022 != 0 && mode & 0o1000 == 0) + { anyhow::bail!( - "binary too large ({} bytes, max {} bytes)", - size, - MAX_BINARY_SIZE + "tunnel executable ancestor '{}' has unsafe ownership or permissions", + directory.display() ); } - let mut file = std::fs::File::create(dest)?; - std::io::copy(&mut entry, &mut file)?; - - #[cfg(unix)] - { - use std::os::unix::fs::PermissionsExt; - std::fs::set_permissions(dest, std::fs::Permissions::from_mode(0o755))?; - } - - return Ok(()); + ancestor = directory.parent(); } } - anyhow::bail!("'{}' not found in archive", binary_name); + #[cfg(not(unix))] + anyhow::bail!( + "safe atomic tunnel self-upgrade is not supported on this platform; reinstall the release manually" + ); + + #[cfg(unix)] + Ok(()) +} + +fn probe_upgrade_directory_write(exe_dir: &Path) -> anyhow::Result<()> { + let probe_path = exe_dir.join(format!( + ".aether-tunnel.write-test-{}-{}", + std::process::id(), + uuid::Uuid::new_v4() + )); + let mut options = std::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(0o600); + } + let probe = options.open(&probe_path)?; + drop(probe); + std::fs::remove_file(&probe_path)?; + Ok(()) } // ── Atomic binary replacement ──────────────────────────────────────────────── fn atomic_replace(new_binary: &Path) -> anyhow::Result { let current_exe = std::env::current_exe()?.canonicalize()?; - let backup_path = current_exe.with_extension("bak"); + atomic_replace_paths(¤t_exe, new_binary) +} - // Remove stale backup - let _ = std::fs::remove_file(&backup_path); - - // current -> .bak - std::fs::rename(¤t_exe, &backup_path).map_err(|e| { - anyhow::anyhow!( - "failed to backup current binary '{}' -> '{}': {}", - current_exe.display(), - backup_path.display(), - e - ) - })?; - - // new -> current - if let Err(e) = std::fs::rename(new_binary, ¤t_exe) { - eprintln!(" ERROR: failed to place new binary, rolling back..."); - let _ = std::fs::rename(&backup_path, ¤t_exe); - anyhow::bail!( - "failed to install new binary '{}' -> '{}': {}", - new_binary.display(), - current_exe.display(), - e - ); +fn atomic_replace_paths(current_exe: &Path, new_binary: &Path) -> anyhow::Result { + let current_exe = std::fs::canonicalize(current_exe)?; + validate_upgrade_storage(¤t_exe)?; + let current_parent = current_exe + .parent() + .ok_or_else(|| anyhow::anyhow!("cannot determine current binary directory"))?; + let new_parent = new_binary + .parent() + .ok_or_else(|| anyhow::anyhow!("cannot determine staged binary directory"))?; + if std::fs::canonicalize(current_parent)? != std::fs::canonicalize(new_parent)? { + anyhow::bail!("staged and current tunnel binaries must share one directory"); } - eprintln!(" Binary replaced: {}", current_exe.display()); - Ok(backup_path) + let new_metadata = std::fs::symlink_metadata(new_binary)?; + if new_metadata.file_type().is_symlink() || !new_metadata.is_file() { + anyhow::bail!("staged tunnel upgrade must be a regular file"); + } + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + if new_metadata.uid() != effective_uid + || new_metadata.mode() & 0o7022 != 0 + || new_metadata.nlink() != 1 + { + anyhow::bail!("staged tunnel upgrade ownership or permissions are unsafe"); + } + } + + let backup_path = current_exe.with_extension("bak"); + + #[cfg(unix)] + { + let backup_staging = current_parent.join(format!( + ".aether-tunnel.backup-{}-{}", + std::process::id(), + uuid::Uuid::new_v4() + )); + std::fs::hard_link(¤t_exe, &backup_staging).map_err(|error| { + anyhow::anyhow!( + "failed to create a no-clobber backup of '{}': {}", + current_exe.display(), + error + ) + })?; + let prepare_result = std::fs::File::open(&backup_staging) + .and_then(|file| file.sync_all()) + .and_then(|_| sync_upgrade_directory(current_parent)); + if let Err(error) = prepare_result { + remove_upgrade_file_if_regular(&backup_staging); + return Err(anyhow::anyhow!( + "failed to durably stage the current tunnel binary backup: {error}" + )); + } + + if let Err(error) = std::fs::rename(new_binary, ¤t_exe) { + remove_upgrade_file_if_regular(&backup_staging); + anyhow::bail!( + "failed to atomically install new binary '{}' -> '{}': {}", + new_binary.display(), + current_exe.display(), + error + ); + } + if let Err(sync_error) = sync_upgrade_directory(current_parent) { + let rollback_result = std::fs::rename(&backup_staging, ¤t_exe) + .and_then(|_| sync_upgrade_directory(current_parent)); + if let Err(rollback_error) = rollback_result { + anyhow::bail!( + "installed tunnel binary but directory sync failed ({sync_error}); rollback also failed ({rollback_error})" + ); + } + anyhow::bail!( + "tunnel binary replacement was rolled back after directory sync failed: {sync_error}" + ); + } + + match std::fs::rename(&backup_staging, &backup_path) { + Ok(()) => { + if let Err(error) = sync_upgrade_directory(current_parent) { + eprintln!( + " WARNING: new binary is durable, but final backup-name sync failed: {}", + error + ); + } + } + Err(error) => { + eprintln!( + " WARNING: could not rotate prior backup ({}); keeping old binary at {}", + error, + backup_staging.display() + ); + eprintln!(" Binary replaced: {}", current_exe.display()); + return Ok(backup_staging); + } + } + + eprintln!(" Binary replaced: {}", current_exe.display()); + Ok(backup_path) + } + + #[cfg(not(unix))] + { + let _ = (new_binary, backup_path); + anyhow::bail!( + "safe atomic tunnel self-upgrade is not supported on this platform; reinstall the release manually" + ) + } +} + +fn sync_upgrade_directory(path: &Path) -> std::io::Result<()> { + std::fs::File::open(path)?.sync_all() +} + +fn restore_tunnel_backup(backup_path: &Path) -> anyhow::Result<()> { + let current_exe = std::env::current_exe()?.canonicalize()?; + restore_tunnel_backup_paths(¤t_exe, backup_path) +} + +fn restore_tunnel_backup_paths(current_exe: &Path, backup_path: &Path) -> anyhow::Result<()> { + let current_exe = std::fs::canonicalize(current_exe)?; + validate_upgrade_storage(¤t_exe)?; + + #[cfg(unix)] + { + use std::os::unix::fs::MetadataExt; + + let current_parent = current_exe + .parent() + .ok_or_else(|| anyhow::anyhow!("cannot determine current binary directory"))?; + if backup_path.parent() != Some(current_parent) { + anyhow::bail!("tunnel backup is outside the executable directory"); + } + let metadata = std::fs::symlink_metadata(backup_path)?; + // SAFETY: geteuid has no preconditions and does not retain pointers. + let effective_uid = unsafe { libc::geteuid() }; + if metadata.file_type().is_symlink() + || !metadata.is_file() + || metadata.uid() != effective_uid + || metadata.mode() & 0o7022 != 0 + || metadata.nlink() != 1 + { + anyhow::bail!("tunnel backup ownership or permissions are unsafe"); + } + std::fs::rename(backup_path, ¤t_exe)?; + if let Err(error) = sync_upgrade_directory(current_parent) { + eprintln!( + " WARNING: old binary was restored, but directory sync failed: {}", + error + ); + } + Ok(()) + } + + #[cfg(not(unix))] + { + let _ = (current_exe, backup_path); + anyhow::bail!("safe tunnel upgrade rollback is not supported on this platform") + } } // ── Public entry point ─────────────────────────────────────────────────────── @@ -310,25 +820,32 @@ async fn execute_upgrade( let exe_dir = current_exe .parent() .ok_or_else(|| anyhow::anyhow!("cannot determine binary directory"))?; - let temp_path = exe_dir.join(".aether-tunnel.upgrade.tmp"); - - if require_root { + if require_root && !super::service::is_root() { + anyhow::bail!("automatic upgrade requires root privileges"); + } + if let Err(error) = validate_upgrade_storage(¤t_exe) { if !super::service::is_root() { - anyhow::bail!("automatic upgrade requires root privileges"); + anyhow::bail!( + "no safe write access to {}: {}. Use: sudo aether-tunnel upgrade", + exe_dir.display(), + error + ); } - } else if !super::service::is_root() { + return Err(error); + } + let temp_path = exe_dir.join(format!( + ".aether-tunnel.upgrade-{}-{}.tmp", + std::process::id(), + uuid::Uuid::new_v4() + )); + + if !require_root && !super::service::is_root() { // Check write permission to binary directory for manual upgrade mode. - let test_path = exe_dir.join(".aether-tunnel.write-test"); - match std::fs::File::create(&test_path) { - Ok(_) => { - let _ = std::fs::remove_file(&test_path); - } - Err(_) => { - anyhow::bail!( - "no write access to {}. Use: sudo aether-tunnel upgrade", - exe_dir.display() - ); - } + if probe_upgrade_directory_write(exe_dir).is_err() { + anyhow::bail!( + "no safe write access to {}. Use: sudo aether-tunnel upgrade", + exe_dir.display() + ); } } @@ -336,36 +853,49 @@ async fn execute_upgrade( eprintln!(" Platform: {}", platform); eprintln!(" Current version: {}", CURRENT_VERSION); - let client = build_github_client()?; - let release = fetch_release(&client, version).await?; + let api_client = build_github_api_client()?; + let release = fetch_release(&api_client, version).await?; let target_tag = &release.tag_name; - let target_semver = target_tag - .strip_prefix("tunnel-v") - .or_else(|| target_tag.strip_prefix("proxy-v")) - .unwrap_or(target_tag); + let target_semver = tunnel_release_semver(target_tag)?; + let current_semver = semver::Version::parse(CURRENT_VERSION) + .map_err(|_| anyhow::anyhow!("current tunnel version is not valid semantic versioning"))?; - eprintln!(" Target version: {} ({})", target_tag, release.name); + eprintln!( + " Target version: {} ({})", + target_tag, + release.name.as_deref().unwrap_or("unnamed release") + ); - if target_semver == CURRENT_VERSION { + if target_semver == current_semver { eprintln!( " Already running version {}, nothing to do.", CURRENT_VERSION ); return Ok(()); } + if target_semver < current_semver { + anyhow::bail!( + "refusing to downgrade aether-tunnel from {} to {}", + current_semver, + target_semver + ); + } eprintln!(); eprintln!(" Upgrading: {} -> {}", CURRENT_VERSION, target_semver); eprintln!(); - if let Err(e) = download_and_verify(&client, target_tag, platform, &temp_path).await { - let _ = std::fs::remove_file(&temp_path); - return Err(e); + let download_client = build_github_download_client()?; + if let Err(error) = + download_and_verify(&download_client, target_tag, platform, &temp_path).await + { + remove_upgrade_file_if_regular(&temp_path); + return Err(error); } let backup_path = match atomic_replace(&temp_path) { Ok(backup) => backup, Err(e) => { - let _ = std::fs::remove_file(&temp_path); + remove_upgrade_file_if_regular(&temp_path); return Err(e); } }; @@ -394,12 +924,32 @@ async fn execute_upgrade( } } RestartMode::Required => { - if !super::service::is_root() { - anyhow::bail!("automatic upgrade requires root privileges"); - } eprintln!(" Restarting managed service..."); - super::service::restart_active_service()?; - eprintln!(" Service restarted."); + match super::service::restart_active_service() { + Ok(()) => eprintln!(" Service restarted."), + Err(restart_error) => { + eprintln!( + " ERROR: upgraded service did not restart; restoring the previous binary..." + ); + if let Err(rollback_error) = restore_tunnel_backup(&backup_path) { + anyhow::bail!( + "upgraded service restart failed ({restart_error}); binary rollback also failed ({rollback_error})" + ); + } + match super::service::restart_active_service() { + Ok(()) => { + anyhow::bail!( + "upgraded service restart failed and the previous binary was restored successfully: {restart_error}" + ); + } + Err(recovery_restart_error) => { + anyhow::bail!( + "upgraded service restart failed ({restart_error}); previous binary was restored but its restart also failed ({recovery_restart_error})" + ); + } + } + } + } } } @@ -424,3 +974,359 @@ pub async fn cmd_upgrade(version: Option) -> anyhow::Result<()> { pub async fn perform_upgrade(version: &str) -> anyhow::Result<()> { execute_upgrade(Some(version), true, RestartMode::Required).await } + +#[cfg(test)] +mod tests { + use super::{ + append_bounded_download_chunk, atomic_replace_paths, extract_binary, + is_trusted_github_download_url, is_trusted_github_host, normalize_requested_release_tag, + parse_checksum, probe_upgrade_directory_write, restore_tunnel_backup_paths, + summarize_remote_error_body, tunnel_release_semver, validate_github_resolved_addrs, + validate_github_resolved_addrs_with_fake_ip, validate_upgrade_storage, + }; + use flate2::{write::GzEncoder, Compression}; + use std::net::SocketAddr; + use std::path::Path; + + const HASH_A: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"; + const HASH_B: &str = "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"; + + #[test] + fn remote_error_body_summary_redacts_payload() { + let body = b"token=do-not-log&message=upstream-secret"; + let summary = summarize_remote_error_body(body); + + assert!(summary.starts_with("response body redacted (bytes=")); + assert!(summary.contains("sha256=")); + assert!(!summary.contains("do-not-log")); + assert!(!summary.contains("upstream-secret")); + } + + #[test] + fn release_redirects_require_trusted_https_hosts() { + for trusted in [ + "https://github.com/fawney19/Aether/releases/download/tag/archive.tar.gz", + "https://objects.githubusercontent.com/github-production-release-asset/archive", + "https://release-assets.githubusercontent.com/github-production-release-asset/archive", + ] { + assert!(is_trusted_github_download_url( + &url::Url::parse(trusted).unwrap() + )); + } + for untrusted in [ + "http://github.com/fawney19/Aether/archive.tar.gz", + "https://github.com:8443/fawney19/Aether/archive.tar.gz", + "https://github.com.evil.example/archive.tar.gz", + "https://user@github.com/archive.tar.gz", + "https://api.github.com/repos/fawney19/Aether/releases", + "https://example.com/archive.tar.gz", + ] { + assert!(!is_trusted_github_download_url( + &url::Url::parse(untrusted).unwrap() + )); + } + } + + #[test] + fn github_dns_rejects_private_or_mixed_answers() { + let public = "8.8.8.8:443".parse::().unwrap(); + let private = "127.0.0.1:443".parse::().unwrap(); + + assert!(validate_github_resolved_addrs(&[public]).is_ok()); + assert!(validate_github_resolved_addrs(&[private]).is_err()); + assert!(validate_github_resolved_addrs(&[public, private]).is_err()); + assert!(validate_github_resolved_addrs(&[]).is_err()); + } + + #[test] + fn github_dns_allows_benchmarking_ip_only_for_trusted_hosts() { + let fake = "198.18.75.234:443".parse::().unwrap(); + assert!(validate_github_resolved_addrs_with_fake_ip(&[fake], true).is_ok()); + assert!(validate_github_resolved_addrs_with_fake_ip( + &[fake, "127.0.0.1:443".parse().unwrap()], + true, + ) + .is_err()); + assert!(validate_github_resolved_addrs_with_fake_ip(&[fake], false).is_err()); + assert!(is_trusted_github_host("api.github.com")); + assert!(is_trusted_github_host("foo.objects.githubusercontent.com")); + assert!(!is_trusted_github_host("github.com.evil.example")); + } + + #[test] + fn release_download_bytes_are_bounded_without_trusting_content_length() { + let mut bytes = Vec::new(); + append_bounded_download_chunk(&mut bytes, b"1234", 8, "archive") + .expect("first chunk should fit"); + append_bounded_download_chunk(&mut bytes, b"5678", 8, "archive") + .expect("exact limit should fit"); + assert_eq!(bytes, b"12345678"); + assert!(append_bounded_download_chunk(&mut bytes, b"9", 8, "archive").is_err()); + assert_eq!(bytes, b"12345678"); + } + + fn archive(entries: &[(&str, tar::EntryType, &[u8])]) -> Vec { + let encoder = GzEncoder::new(Vec::new(), Compression::default()); + let mut builder = tar::Builder::new(encoder); + for (path, entry_type, body) in entries { + let mut header = tar::Header::new_gnu(); + header.set_entry_type(*entry_type); + header.set_mode(0o755); + header.set_size(body.len() as u64); + header.set_cksum(); + builder + .append_data(&mut header, path, *body) + .expect("test archive entry should append"); + } + builder + .into_inner() + .expect("test archive should finish") + .finish() + .expect("test gzip should finish") + } + + #[test] + fn checksum_requires_one_exact_valid_filename() { + let filename = "aether-tunnel-linux-amd64.tar.gz"; + assert_eq!( + parse_checksum(&format!("{HASH_A} {filename}\n"), filename) + .expect("exact checksum should parse"), + HASH_A + ); + assert!(parse_checksum(&format!("{HASH_A} nested/{filename}\n"), filename).is_err()); + assert!(parse_checksum( + &format!("{HASH_A} {filename}\n{HASH_B} *{filename}\n"), + filename + ) + .is_err()); + assert!(parse_checksum(&format!("not-a-hash {filename}\n"), filename).is_err()); + } + + #[test] + fn release_tags_require_bounded_semver_without_url_metacharacters() { + assert_eq!( + normalize_requested_release_tag("0.3.17").unwrap(), + "tunnel-v0.3.17" + ); + assert_eq!( + normalize_requested_release_tag("proxy-v0.3.17-rc.1").unwrap(), + "proxy-v0.3.17-rc.1" + ); + assert!(tunnel_release_semver("tunnel-v1.2.3+build.7").is_ok()); + for invalid in [ + "", + "latest", + "tunnel-v../other", + "tunnel-v1.2.3?x=1", + "tunnel-v1.2", + "other-v1.2.3", + ] { + assert!(normalize_requested_release_tag(invalid).is_err()); + } + } + + #[test] + fn extraction_rejects_ambiguous_or_non_root_archives() { + let root = std::env::temp_dir().join(format!( + "aether-tunnel-upgrade-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir_all(&root).expect("test directory should be created"); + let destination = root.join("candidate"); + let binary_name = if cfg!(target_os = "windows") { + "aether-tunnel.exe" + } else { + "aether-tunnel" + }; + + let nested_name = format!("nested/{binary_name}"); + assert!(extract_binary( + &archive(&[(&nested_name, tar::EntryType::Regular, b"binary")]), + &destination + ) + .is_err()); + assert!(extract_binary( + &archive(&[ + (binary_name, tar::EntryType::Regular, b"binary"), + ("extra", tar::EntryType::Regular, b"extra"), + ]), + &destination + ) + .is_err()); + assert!(!destination.exists()); + + std::fs::remove_dir_all(root).expect("test directory should be removed"); + } + + #[test] + fn extraction_refuses_to_follow_existing_staging_symlink() { + #[cfg(unix)] + { + use std::os::unix::fs::symlink; + + let root = std::env::temp_dir().join(format!( + "aether-tunnel-upgrade-symlink-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir_all(&root).expect("test directory should be created"); + let victim = root.join("victim"); + let destination = root.join("candidate"); + std::fs::write(&victim, b"keep").expect("victim should be written"); + symlink(&victim, &destination).expect("test symlink should be created"); + let binary_name = if cfg!(target_os = "windows") { + "aether-tunnel.exe" + } else { + "aether-tunnel" + }; + + assert!(extract_binary( + &archive(&[(binary_name, tar::EntryType::Regular, b"replace")]), + Path::new(&destination) + ) + .is_err()); + assert_eq!( + std::fs::read(&victim).expect("victim should remain readable"), + b"keep" + ); + std::fs::remove_dir_all(root).expect("test directory should be removed"); + } + } + + #[cfg(unix)] + #[test] + fn extraction_creates_a_private_synced_executable() { + use std::os::unix::fs::PermissionsExt; + + let root = std::env::temp_dir().join(format!( + "aether-tunnel-upgrade-mode-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir_all(&root).unwrap(); + let destination = root.join("candidate"); + extract_binary( + &archive(&[("aether-tunnel", tar::EntryType::Regular, b"new-binary")]), + &destination, + ) + .expect("valid archive should extract"); + + assert_eq!(std::fs::read(&destination).unwrap(), b"new-binary"); + assert_eq!( + std::fs::metadata(&destination) + .unwrap() + .permissions() + .mode() + & 0o777, + 0o755 + ); + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn upgrade_storage_and_write_probe_reject_unsafe_directories_without_fixed_files() { + use std::os::unix::fs::PermissionsExt; + + let root = std::env::temp_dir().join(format!( + "aether-tunnel-upgrade-storage-test-{}", + uuid::Uuid::new_v4() + )); + let executable_directory = root.join("bin"); + std::fs::create_dir_all(&executable_directory).unwrap(); + let current = executable_directory.join("aether-tunnel"); + std::fs::write(¤t, b"old").unwrap(); + std::fs::set_permissions(¤t, std::fs::Permissions::from_mode(0o755)).unwrap(); + + validate_upgrade_storage(¤t).expect("private owned storage should pass"); + probe_upgrade_directory_write(&executable_directory) + .expect("unique write probe should pass"); + assert!(std::fs::read_dir(&executable_directory) + .unwrap() + .all(|entry| { + !entry + .unwrap() + .file_name() + .to_string_lossy() + .starts_with(".aether-tunnel.write-test-") + })); + + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o777)).unwrap(); + assert!(validate_upgrade_storage(¤t).is_err()); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o755)).unwrap(); + std::fs::set_permissions( + &executable_directory, + std::fs::Permissions::from_mode(0o777), + ) + .unwrap(); + assert!(validate_upgrade_storage(¤t).is_err()); + std::fs::set_permissions( + &executable_directory, + std::fs::Permissions::from_mode(0o755), + ) + .unwrap(); + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn atomic_upgrade_replaces_in_one_step_and_preserves_safe_rollback() { + use std::os::unix::fs::{symlink, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-tunnel-atomic-upgrade-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir_all(&root).unwrap(); + let current = root.join("aether-tunnel"); + let staged = root.join("candidate"); + std::fs::write(¤t, b"old-binary").unwrap(); + std::fs::write(&staged, b"new-binary").unwrap(); + std::fs::set_permissions(¤t, std::fs::Permissions::from_mode(0o755)).unwrap(); + std::fs::set_permissions(&staged, std::fs::Permissions::from_mode(0o755)).unwrap(); + + let victim = root.join("victim"); + std::fs::write(&victim, b"known-good").unwrap(); + let fixed_backup = current.with_extension("bak"); + symlink(&victim, &fixed_backup).unwrap(); + + let backup = atomic_replace_paths(¤t, &staged).expect("upgrade should succeed"); + + assert_eq!(std::fs::read(¤t).unwrap(), b"new-binary"); + assert_eq!(std::fs::read(&backup).unwrap(), b"old-binary"); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + assert!(!std::fs::symlink_metadata(&backup) + .unwrap() + .file_type() + .is_symlink()); + assert!(!staged.exists()); + + restore_tunnel_backup_paths(¤t, &backup).expect("rollback should succeed"); + assert_eq!(std::fs::read(¤t).unwrap(), b"old-binary"); + assert!(!backup.exists()); + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn atomic_upgrade_rejects_hardlinked_staging_without_touching_current() { + use std::os::unix::fs::PermissionsExt; + + let root = std::env::temp_dir().join(format!( + "aether-tunnel-hardlink-upgrade-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir_all(&root).unwrap(); + let current = root.join("aether-tunnel"); + let staged = root.join("candidate"); + let outside = root.join("outside"); + std::fs::write(¤t, b"old-binary").unwrap(); + std::fs::write(&outside, b"new-binary").unwrap(); + std::fs::hard_link(&outside, &staged).unwrap(); + std::fs::set_permissions(¤t, std::fs::Permissions::from_mode(0o755)).unwrap(); + std::fs::set_permissions(&staged, std::fs::Permissions::from_mode(0o755)).unwrap(); + + assert!(atomic_replace_paths(¤t, &staged).is_err()); + assert_eq!(std::fs::read(¤t).unwrap(), b"old-binary"); + assert_eq!(std::fs::read(&outside).unwrap(), b"new-binary"); + std::fs::remove_dir_all(root).unwrap(); + } +} diff --git a/apps/aether-tunnel/src/state.rs b/apps/aether-tunnel/src/state.rs index 0f28d62fe..26b54a42d 100644 --- a/apps/aether-tunnel/src/state.rs +++ b/apps/aether-tunnel/src/state.rs @@ -55,6 +55,8 @@ pub struct ServerContext { pub node_name: String, /// Node ID assigned by this Aether server. pub node_id: Arc>, + /// Server-issued node incarnation bound into tunnel authentication. + pub tunnel_generation: String, /// API client for this server. pub aether_client: Arc, /// Dynamic config from this server's heartbeat ACKs. diff --git a/apps/aether-tunnel/src/target_filter.rs b/apps/aether-tunnel/src/target_filter.rs index 715c50dc1..7a09cc361 100644 --- a/apps/aether-tunnel/src/target_filter.rs +++ b/apps/aether-tunnel/src/target_filter.rs @@ -1,5 +1,5 @@ use std::collections::{HashMap, HashSet}; -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; +use std::net::{IpAddr, SocketAddr}; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -7,80 +7,11 @@ use tokio::sync::RwLock; /// Check if an IP address belongs to a private/reserved network. pub fn is_private_ip(ip: &IpAddr) -> bool { - match ip { - IpAddr::V4(v4) => is_private_ipv4(v4), - IpAddr::V6(v6) => is_private_ipv6(v6), - } -} - -fn is_private_ipv4(ip: &Ipv4Addr) -> bool { - let octets = ip.octets(); - // 10.0.0.0/8 - if octets[0] == 10 { - return true; - } - // 172.16.0.0/12 - if octets[0] == 172 && (16..=31).contains(&octets[1]) { - return true; - } - // 192.168.0.0/16 - if octets[0] == 192 && octets[1] == 168 { - return true; - } - // 127.0.0.0/8 - if octets[0] == 127 { - return true; - } - // 169.254.0.0/16 (link-local) - if octets[0] == 169 && octets[1] == 254 { - return true; - } - // 0.0.0.0/8 - if octets[0] == 0 { - return true; - } - // 100.64.0.0/10 (CGNAT / shared address space) - if octets[0] == 100 && (64..=127).contains(&octets[1]) { - return true; - } - // 192.0.0.0/24 (IETF protocol assignments) - if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 { - return true; - } - // 198.18.0.0/15 (benchmark testing) - if octets[0] == 198 && (18..=19).contains(&octets[1]) { - return true; - } - // 240.0.0.0/4 (reserved for future use) - if octets[0] >= 240 { - return true; - } - false -} - -fn is_private_ipv6(ip: &Ipv6Addr) -> bool { - // ::1 loopback - if ip.is_loopback() { - return true; - } - // :: unspecified - if ip.is_unspecified() { - return true; - } - let segments = ip.segments(); - // fc00::/7 (ULA) - first byte is 0xfc or 0xfd - if segments[0] & 0xfe00 == 0xfc00 { - return true; - } - // fe80::/10 (link-local) - if segments[0] & 0xffc0 == 0xfe80 { - return true; - } - // IPv4-mapped IPv6 (::ffff:x.x.x.x) - check the embedded IPv4 - if let Some(v4) = ip.to_ipv4_mapped() { - return is_private_ipv4(&v4); - } - false + // Keep every egress path on the same conservative classification as the + // gateway. This includes documentation, benchmarking, transition, and + // other reserved ranges that the standard `IpAddr::is_private` helpers do + // not cover. + aether_http::is_private_or_reserved_ip(*ip) } #[derive(Debug)] @@ -108,6 +39,13 @@ impl std::fmt::Display for FilterError { } } +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct DnsCacheKey { + host: String, + port: u16, + allow_private: bool, +} + struct DnsCacheEntry { addrs: Arc>, expires_at: Instant, @@ -115,12 +53,12 @@ struct DnsCacheEntry { } /// Lightweight DNS cache with TTL + capacity bounds. -/// Stores all public resolved addresses per host (used by SafeDnsResolver -/// to ensure reqwest connects to the same validated addresses). +/// Stores validated addresses per host and port so ACL checks and connection +/// setup can reuse the same resolution result. pub struct DnsCache { ttl: Duration, capacity: usize, - entries: RwLock>, + entries: RwLock>, } impl DnsCache { @@ -132,31 +70,30 @@ impl DnsCache { } } - /// Look up cached public addresses for a host (any port). + /// Look up cached public addresses for a host + port. /// - /// Used by `SafeDnsResolver` which only knows the hostname — returns the - /// first unexpired entry whose key starts with `host:`. - pub async fn get_by_host(&self, host: &str) -> Option>> { - if self.capacity == 0 || self.ttl.is_zero() { - return None; - } - let prefix = format!("{}:", host.to_ascii_lowercase()); - let now = Instant::now(); - let entries = self.entries.read().await; - for (key, entry) in entries.iter() { - if key.starts_with(&prefix) && entry.expires_at > now { - return Some(Arc::clone(&entry.addrs)); - } - } - None + /// This compatibility wrapper uses the restrictive (public-only) policy. + /// Callers that explicitly allow private targets must use + /// [`Self::get_for_policy`] so entries cannot cross the policy boundary. + #[allow(dead_code)] + pub async fn get(&self, host: &str, port: u16) -> Option>> { + self.get_for_policy(host, port, false).await } - /// Look up cached public addresses for a host + port. - pub async fn get(&self, host: &str, port: u16) -> Option>> { + /// Look up a cached resolution under the exact target policy used to + /// validate it. Private-target and public-only resolutions are kept in + /// separate entries; otherwise a cache populated while private targets + /// are enabled could bypass filtering after a policy change. + pub async fn get_for_policy( + &self, + host: &str, + port: u16, + allow_private: bool, + ) -> Option>> { if self.capacity == 0 || self.ttl.is_zero() { return None; } - let key = Self::key(host, port); + let key = Self::key(host, port, allow_private); let now = Instant::now(); // Fast path: read lock for cache hit @@ -175,12 +112,26 @@ impl DnsCache { None } - /// Insert resolved public addresses into cache. + /// Insert resolved public addresses into the restrictive (public-only) + /// cache. This compatibility wrapper preserves the original API; + /// policy-aware callers should use [`Self::insert_for_policy`]. + #[allow(dead_code)] pub async fn insert(&self, host: &str, port: u16, addrs: Arc>) { + self.insert_for_policy(host, port, false, addrs).await; + } + + /// Insert addresses under the exact target policy that produced them. + pub async fn insert_for_policy( + &self, + host: &str, + port: u16, + allow_private: bool, + addrs: Arc>, + ) { if self.capacity == 0 || self.ttl.is_zero() || addrs.is_empty() { return; } - let key = Self::key(host, port); + let key = Self::key(host, port, allow_private); let now = Instant::now(); let mut entries = self.entries.write().await; entries.retain(|_, entry| entry.expires_at > now); @@ -205,8 +156,12 @@ impl DnsCache { ); } - fn key(host: &str, port: u16) -> String { - format!("{}:{}", host.to_ascii_lowercase(), port) + fn key(host: &str, port: u16, allow_private: bool) -> DnsCacheKey { + DnsCacheKey { + host: host.to_ascii_lowercase(), + port, + allow_private, + } } } @@ -222,16 +177,16 @@ pub async fn resolve_public_addrs( dns_cache: &DnsCache, ) -> Result, FilterError> { // Cache hit - if let Some(addrs) = dns_cache.get(host, port).await { + if let Some(addrs) = dns_cache.get_for_policy(host, port, allow_private).await { return Ok((*addrs).clone()); } - // Async DNS resolution - let addr_str = format!("{}:{}", host, port); - let resolved: Vec = tokio::net::lookup_host(&addr_str) - .await - .map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))? - .collect(); + // Async DNS resolution. Keep resolver wait time and answer count + // bounded before applying the private-address policy below. + let resolved: Vec = + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?; if resolved.is_empty() { return Err(FilterError::DnsResolutionFailed(host.to_string())); @@ -253,15 +208,17 @@ pub async fn resolve_public_addrs( // Cache the validated public addresses let arc_addrs = Arc::new(public); - dns_cache.insert(host, port, Arc::clone(&arc_addrs)).await; + dns_cache + .insert_for_policy(host, port, allow_private, Arc::clone(&arc_addrs)) + .await; Ok((*arc_addrs).clone()) } /// Validate that the target host:port is allowed. /// /// Performs port whitelist check, private IP filtering, and DNS resolution -/// with caching. The resolved addresses are stored in the shared DnsCache -/// so that the SafeDnsResolver can reuse them, eliminating the TOCTTOU gap. +/// with caching. The caller must use the returned addresses for the actual +/// connection rather than resolving the hostname again. pub async fn validate_target( host: &str, port: u16, @@ -282,12 +239,14 @@ pub async fn validate_target( return Ok(vec![SocketAddr::new(ip, port)]); } - // Resolve and validate DNS (populates cache for SafeDnsResolver) + // Resolve and return the exact addresses authorized for this request. resolve_public_addrs(host, port, allow_private, dns_cache).await } #[cfg(test)] mod tests { + use std::net::{Ipv4Addr, Ipv6Addr}; + use super::*; fn ports() -> HashSet { @@ -318,9 +277,11 @@ mod tests { assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(198, 18, 0, 1)))); // Reserved assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(240, 0, 0, 1)))); + // Multicast + assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(224, 0, 0, 1)))); // Public assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)))); - assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)))); + assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)))); } #[test] @@ -335,6 +296,47 @@ mod tests { assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new( 0xfe80, 0, 0, 0, 0, 0, 0, 1 )))); + // fec0::/10 (deprecated site-local) + assert!(is_private_ip(&"fec0::1".parse().unwrap())); + assert!(is_private_ip( + &"feff:ffff:ffff:ffff:ffff:ffff:ffff:ffff".parse().unwrap() + )); + // ff00::/8 (multicast) + assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new( + 0xff02, 0, 0, 0, 0, 0, 0, 1 + )))); + // NAT64 well-known and local-use prefixes. + assert!(is_private_ip(&"64:ff9b::10.0.0.1".parse().unwrap())); + assert!(is_private_ip(&"64:ff9b::ffff:ffff".parse().unwrap())); + assert!(is_private_ip(&"64:ff9b:1::10.0.0.1".parse().unwrap())); + assert!(is_private_ip( + &"64:ff9b:1:ffff:ffff:ffff:ffff:ffff".parse().unwrap() + )); + assert!(!is_private_ip(&"64:ff9a:ffff::1".parse().unwrap())); + assert!(!is_private_ip(&"64:ff9b:0:1::1".parse().unwrap())); + assert!(!is_private_ip(&"64:ff9b:2::1".parse().unwrap())); + + // IPv6 transition formats with embedded IPv4 addresses. + assert!(is_private_ip(&"2002:0a00:0001::1".parse().unwrap())); + assert!(is_private_ip( + &"2002:ffff:ffff:ffff:ffff:ffff:ffff:ffff".parse().unwrap() + )); + assert!(!is_private_ip(&"2003::1".parse().unwrap())); + assert!(is_private_ip( + &"2001:0000:4136:e378:8000:63bf:3fff:fdd2".parse().unwrap() + )); + assert!(!is_private_ip(&"2001:1::1".parse().unwrap())); + assert!(is_private_ip(&"::192.0.2.1".parse().unwrap())); + assert!(is_private_ip(&"::ffff:0:192.0.2.1".parse().unwrap())); + assert!(is_private_ip(&"2001:db8::5efe:10.0.0.1".parse().unwrap())); + assert!(is_private_ip( + &"2001:db8::200:5efe:192.0.2.1".parse().unwrap() + )); + + // IPv4-mapped public addresses remain allowed, while private mapped + // addresses continue through the IPv4 classification. + assert!(!is_private_ip(&"::ffff:8.8.8.8".parse().unwrap())); + assert!(is_private_ip(&"::ffff:127.0.0.1".parse().unwrap())); } #[tokio::test] @@ -351,6 +353,24 @@ mod tests { assert!(matches!(result, Err(FilterError::PrivateIp(_)))); } + #[tokio::test] + async fn test_ipv6_site_local_blocked_unless_private_targets_allowed() { + let cache = cache(); + let result = validate_target("fec0::1", 443, &ports(), false, &cache).await; + assert!(matches!( + result, + Err(FilterError::PrivateIp(IpAddr::V6(ip))) if ip == "fec0::1".parse::().unwrap() + )); + + let result = validate_target("fec0::1", 443, &ports(), true, &cache) + .await + .unwrap(); + assert_eq!( + result, + vec![SocketAddr::new(IpAddr::V6("fec0::1".parse().unwrap()), 443)] + ); + } + #[tokio::test] async fn test_public_ip_allowed() { let cache = cache(); @@ -412,4 +432,40 @@ mod tests { let cached = cache.get("example.com", 443).await.unwrap(); assert_eq!(*cached, addrs); } + + #[tokio::test] + async fn test_cache_does_not_cross_private_target_policy() { + let cache = cache(); + let private = Arc::new(vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 80)]); + let public = Arc::new(vec![SocketAddr::new( + IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), + 80, + )]); + + cache + .insert_for_policy("example.com", 80, true, Arc::clone(&private)) + .await; + assert!(cache + .get_for_policy("example.com", 80, false) + .await + .is_none()); + assert_eq!( + *cache + .get_for_policy("example.com", 80, true) + .await + .expect("private-policy entry"), + *private + ); + + cache + .insert_for_policy("example.com", 80, false, Arc::clone(&public)) + .await; + assert_eq!( + *cache + .get_for_policy("EXAMPLE.COM", 80, false) + .await + .expect("public-policy entry"), + *public + ); + } } diff --git a/apps/aether-tunnel/src/tunnel/client.rs b/apps/aether-tunnel/src/tunnel/client.rs index 23765875e..5ccacc76e 100644 --- a/apps/aether-tunnel/src/tunnel/client.rs +++ b/apps/aether-tunnel/src/tunnel/client.rs @@ -3,7 +3,7 @@ use std::io; use std::net::SocketAddr; use std::sync::Arc; -use std::time::{Duration, Instant}; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use base64::Engine as _; use tokio::net::TcpStream; @@ -13,6 +13,7 @@ use tokio_tungstenite::tungstenite::http; use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; use tracing::{debug, info, warn}; +use crate::config::aether_url_for_log; use crate::egress_proxy::{ connect_target_via_proxy, IpFamily, ProxyConnectOptions, UpstreamProxyConfig, }; @@ -22,8 +23,10 @@ use aether_contracts::tunnel::{ TUNNEL_PROTOCOL_VERSION_HEADER, }; use aether_contracts::tunnel_security::{ - SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED, - TUNNEL_SECURITY_SESSION_HEADER, + sign_tunnel_security_handshake_for_generation, SecureFrameCodec, TunnelSecurityRole, + TUNNEL_GENERATION_HEADER, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED, + TUNNEL_SECURITY_PROOF_NONCE_HEADER, TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, TUNNEL_SECURITY_SESSION_HEADER, }; use super::{dispatcher, heartbeat, writer}; @@ -48,7 +51,7 @@ pub async fn connect_and_run( drain: watch::Receiver, ) -> Result { let ws_url = build_tunnel_url(server); - debug!(url = %ws_url, conn = conn_idx, "connecting tunnel"); + debug!(url = %aether_url_for_log(&ws_url), conn = conn_idx, "connecting tunnel"); // Build WebSocket request with auth headers let mut request = ws_url.clone().into_client_request()?; @@ -67,19 +70,38 @@ pub async fn connect_and_run( ); let node_id = server.node_id.read().unwrap().clone(); insert_ascii_header(headers, "X-Node-Id", &node_id, "node_id")?; - let security_session = uuid::Uuid::new_v4().simple().to_string(); - if server.tunnel_security == crate::config::TunnelSecurity::NonTlsRequired { - headers.insert( - TUNNEL_SECURITY_HEADER, - http::HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED), - ); - insert_ascii_header( - headers, - TUNNEL_SECURITY_SESSION_HEADER, - &security_session, - "tunnel security session", - )?; - } + insert_ascii_header( + headers, + TUNNEL_GENERATION_HEADER, + &server.tunnel_generation, + "tunnel_generation", + )?; + let security_session = + if server.tunnel_security == crate::config::TunnelSecurity::NonTlsRequired { + let key = server + .tunnel_encryption_key + .as_deref() + .ok_or_else(|| anyhow::anyhow!("secure tunnel requires tunnel_encryption_key"))?; + let session = uuid::Uuid::new_v4().simple().to_string(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map_err(|_| anyhow::anyhow!("system clock is before the Unix epoch"))? + .as_secs(); + insert_tunnel_security_handshake_headers( + headers, + key, + &node_id, + &server.tunnel_generation, + &session, + CURRENT_TUNNEL_PROTOCOL_VERSION, + timestamp, + &nonce, + )?; + session + } else { + String::new() + }; // Use dynamic node_name (may be updated by remote config) instead of // the static server.node_name, so that remote name changes take effect // on the next reconnect. @@ -449,9 +471,10 @@ async fn connect_direct_tunnel_tcp( port: u16, ip_family: IpFamily, ) -> io::Result { - let resolved = tokio::net::lookup_host((host, port)) - .await - .map_err(|err| io::Error::other(format!("tunnel DNS failed: {err}")))?; + let resolved = + aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .map_err(|err| io::Error::other(format!("tunnel DNS failed: {err}")))?; let addrs = filter_socket_addrs(resolved, ip_family); if addrs.is_empty() { @@ -540,6 +563,61 @@ fn insert_ascii_header( Ok(()) } +// The handshake transcript has a fixed set of wire fields. Keep the explicit +// arguments and ordering so the client remains interoperable with existing +// tunnel servers. +#[allow(clippy::too_many_arguments)] +fn insert_tunnel_security_handshake_headers( + headers: &mut http::HeaderMap, + key: &str, + node_id: &str, + tunnel_generation: &str, + session: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, +) -> anyhow::Result<()> { + let signature = sign_tunnel_security_handshake_for_generation( + key, + node_id, + tunnel_generation, + TUNNEL_SECURITY_NON_TLS_REQUIRED, + session, + protocol_version, + timestamp_unix_secs, + nonce, + )?; + headers.insert( + TUNNEL_SECURITY_HEADER, + http::HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED), + ); + insert_ascii_header( + headers, + TUNNEL_SECURITY_SESSION_HEADER, + session, + "tunnel security session", + )?; + insert_ascii_header( + headers, + TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, + ×tamp_unix_secs.to_string(), + "tunnel security proof timestamp", + )?; + insert_ascii_header( + headers, + TUNNEL_SECURITY_PROOF_NONCE_HEADER, + nonce, + "tunnel security proof nonce", + )?; + insert_ascii_header( + headers, + TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER, + &signature, + "tunnel security proof signature", + )?; + Ok(()) +} + fn insert_node_name_headers(headers: &mut http::HeaderMap, node_name: &str) -> anyhow::Result<()> { if node_name.is_ascii() { return insert_ascii_header(headers, "X-Node-Name", node_name, "node_name"); @@ -620,4 +698,46 @@ mod tests { .expect("encoded value should decode"); assert_eq!(decoded, "日本节点".as_bytes()); } + + #[test] + fn tunnel_security_headers_include_verifiable_psk_proof() { + let key = base64::engine::general_purpose::STANDARD.encode([7_u8; 32]); + let mut headers = http::HeaderMap::new(); + insert_tunnel_security_handshake_headers( + &mut headers, + &key, + "node-1", + "generation-1", + "0123456789abcdef0123456789abcdef", + CURRENT_TUNNEL_PROTOCOL_VERSION, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + ) + .expect("security proof headers"); + + let signature = headers[TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER] + .to_str() + .expect("signature header"); + assert_eq!( + headers[TUNNEL_SECURITY_HEADER], + TUNNEL_SECURITY_NON_TLS_REQUIRED + ); + assert_eq!( + headers[TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER], + "1700000000" + ); + assert!( + aether_contracts::tunnel_security::verify_tunnel_security_handshake_for_generation( + &key, + "node-1", + "generation-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + CURRENT_TUNNEL_PROTOCOL_VERSION, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + signature, + ) + ); + } } diff --git a/apps/aether-tunnel/src/tunnel/dispatcher.rs b/apps/aether-tunnel/src/tunnel/dispatcher.rs index 150713a17..73fb9d156 100644 --- a/apps/aether-tunnel/src/tunnel/dispatcher.rs +++ b/apps/aether-tunnel/src/tunnel/dispatcher.rs @@ -1,13 +1,14 @@ //! Frame dispatcher: reads incoming WebSocket frames and routes them. -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; +use std::mem::size_of; use std::sync::Arc; +use std::sync::LazyLock; use std::time::Duration; use bytes::Bytes; use futures_util::StreamExt; -use tokio::sync::mpsc; -use tokio::sync::watch; +use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore}; use tokio::task::JoinHandle; use tokio_tungstenite::tungstenite::Message; use tracing::{debug, error, info, warn}; @@ -15,12 +16,27 @@ use tracing::{debug, error, info, warn}; use crate::state::{AppState, ServerContext}; use super::heartbeat::HeartbeatHandle; -use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta}; +use super::protocol::{decompress_if_gzip_with_limit, Frame, MsgType, RequestMeta}; use super::stream_handler; use super::stream_handler::StreamSendWindow; use super::writer::FrameSender; use aether_contracts::tunnel_security::SecureFrameCodec; +const REQUEST_BODY_QUEUE_BUDGET_BYTES: usize = 256 * 1024 * 1024; +static REQUEST_BODY_QUEUE_BUDGET: LazyLock> = + LazyLock::new(|| Arc::new(Semaphore::new(REQUEST_BODY_QUEUE_BUDGET_BYTES))); + +struct BudgetedFramePayload { + bytes: Bytes, + _permit: OwnedSemaphorePermit, +} + +impl AsRef<[u8]> for BudgetedFramePayload { + fn as_ref(&self) -> &[u8] { + self.bytes.as_ref() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum StreamDispatchStatus { Delivered, @@ -34,6 +50,24 @@ struct StreamDispatchTarget { response_window: Arc, } +/// A request stream is identified by a non-zero id and may only be opened +/// once while its handler is active. Replacing an entry in `streams` would +/// orphan the old body channel while still spawning another handler, making +/// the active-stream limit ineffective and allowing unbounded task growth. +fn validate_request_stream_id( + streams: &HashMap, + active_handler_ids: &HashSet, + stream_id: u32, +) -> Result<(), &'static str> { + if stream_id == 0 { + return Err("invalid stream id"); + } + if streams.contains_key(&stream_id) || active_handler_ids.contains(&stream_id) { + return Err("duplicate stream id"); + } + Ok(()) +} + /// Run the dispatcher loop, reading from the WebSocket stream. #[allow(dead_code)] pub async fn run( @@ -70,6 +104,10 @@ where { // Active streams: stream_id -> body sender + response flow-control window. let mut streams: HashMap = HashMap::new(); + // A handler can outlive its routing entry when body dispatch fails. Keep + // its id reserved until the handler reports completion so a peer cannot + // reopen the same id and bypass the stream admission limit. + let mut active_handler_ids: HashSet = HashSet::new(); // Track spawned stream handlers so we can wait for them on shutdown let mut handler_handles: Vec> = Vec::new(); let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::(); @@ -85,7 +123,7 @@ where let mut draining = *drain.borrow(); let read_err = loop { - if draining && streams.is_empty() { + if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after in-flight streams completed"); break None; } @@ -109,8 +147,9 @@ where } finished = handler_finished_rx.recv() => { if let Some(stream_id) = finished { + active_handler_ids.remove(&stream_id); streams.remove(&stream_id); - if draining && streams.is_empty() { + if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after stream handler completion"); break None; } @@ -184,6 +223,20 @@ where match frame.msg_type { MsgType::RequestHeaders => { + if let Err(reason) = + validate_request_stream_id(&streams, &active_handler_ids, frame.stream_id) + { + warn!( + stream_id = frame.stream_id, + reason, "rejecting request headers with invalid stream id" + ); + // Zero is reserved for connection-level control frames, + // so do not emit a stream-scoped error using that id. + if frame.stream_id != 0 { + try_send_stream_error(&frame_tx, frame.stream_id, reason); + } + continue; + } if draining { if frame_tx .try_send(Frame::new( @@ -203,7 +256,10 @@ where } // Decompress if the frame is gzip-compressed, then parse metadata - let payload = match decompress_if_gzip(&frame) { + let payload = match decompress_if_gzip_with_limit( + &frame, + aether_contracts::tunnel::MAX_TUNNEL_RELAY_META_LEN, + ) { Ok(p) => p, Err(e) => { warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed"); @@ -233,7 +289,7 @@ where } }; - if streams.len() >= max_streams { + if active_handler_ids.len() >= max_streams { warn!( stream_id = frame.stream_id, "max concurrent streams reached" @@ -267,6 +323,7 @@ where response_window: Arc::clone(&response_window), }, ); + active_handler_ids.insert(frame.stream_id); let request_headers_end_stream = frame.is_end_stream(); let state_clone = Arc::clone(&state); @@ -321,7 +378,8 @@ where "tunnel request body dispatch stalled", ); } - if is_end && draining && streams.is_empty() { + if is_end && draining && streams.is_empty() && active_handler_ids.is_empty() + { info!("tunnel drained after request body completion"); break None; } @@ -333,7 +391,7 @@ where // Client-side cancellation or end if let Some(target) = streams.remove(&frame.stream_id) { let _ = dispatch_stream_frame(&target.body_tx, frame).await; - if draining && streams.is_empty() { + if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after stream termination"); break None; } @@ -399,7 +457,7 @@ where if frames_since_cleanup >= 64 || handler_handles.len() > max_streams { handler_handles.retain(|h| !h.is_finished()); frames_since_cleanup = 0; - if draining && streams.is_empty() { + if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after cleanup"); break None; } @@ -421,12 +479,18 @@ where async fn dispatch_stream_frame(tx: &mpsc::Sender, frame: Frame) -> StreamDispatchStatus { let stream_id = frame.stream_id; - match tokio::time::timeout(stream_frame_dispatch_timeout(), tx.send(frame)).await { - Ok(Ok(())) => StreamDispatchStatus::Delivered, - Ok(Err(_)) => { + let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async { + let frame = attach_request_body_queue_budget(frame).await?; + tx.send(frame).await.ok()?; + Some(()) + }) + .await; + match dispatched { + Ok(Some(())) => StreamDispatchStatus::Delivered, + Ok(None) => { warn!( stream_id, - "stream handler channel closed while dispatching tunnel frame" + "stream handler channel or request body budget closed while dispatching tunnel frame" ); StreamDispatchStatus::Closed } @@ -441,6 +505,50 @@ async fn dispatch_stream_frame(tx: &mpsc::Sender, frame: Frame) -> Stream } } +async fn attach_request_body_queue_budget(frame: Frame) -> Option { + attach_request_body_queue_budget_with( + frame, + Arc::clone(&REQUEST_BODY_QUEUE_BUDGET), + REQUEST_BODY_QUEUE_BUDGET_BYTES, + ) + .await +} + +async fn attach_request_body_queue_budget_with( + mut frame: Frame, + budget: Arc, + budget_bytes: usize, +) -> Option { + if frame.msg_type != MsgType::RequestBody { + return Some(frame); + } + let permits = request_body_queue_permits(&frame, budget_bytes)?; + let permit = budget.acquire_many_owned(permits).await.ok()?; + frame.payload = Bytes::from_owner(BudgetedFramePayload { + bytes: frame.payload, + _permit: permit, + }); + Some(frame) +} + +fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option { + let decoded_budget = if frame.is_gzip() { + aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES + } else { + 0 + }; + let retained_bytes = frame + .payload + .len() + .checked_add(decoded_budget)? + .checked_add(size_of::())? + .max(1); + if retained_bytes > budget_bytes { + return None; + } + u32::try_from(retained_bytes).ok() +} + /// Bound how long a single stream handler is allowed to block the shared /// WebSocket read loop while receiving request-body frames. fn stream_frame_dispatch_timeout() -> Duration { @@ -497,6 +605,7 @@ async fn drain_handlers(handles: Vec>) { #[cfg(test)] mod tests { use super::*; + use aether_contracts::tunnel::{compress_payload, flags}; use aether_runtime::bounded_queue; #[tokio::test] @@ -534,6 +643,70 @@ mod tests { assert_eq!(retained.payload, Bytes::from_static(b"first")); } + #[tokio::test] + async fn request_body_queue_budget_releases_when_frame_is_dropped() { + const BUDGET_BYTES: usize = 4096; + let budget = Arc::new(Semaphore::new(BUDGET_BYTES)); + let frame = Frame::new( + 7, + MsgType::RequestBody, + 0, + Bytes::from_static(b"request body"), + ); + let permits = request_body_queue_permits(&frame, BUDGET_BYTES).expect("permit count"); + + let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), BUDGET_BYTES) + .await + .expect("frame should fit the queue budget"); + assert_eq!(budget.available_permits(), BUDGET_BYTES - permits as usize); + + drop(frame); + assert_eq!(budget.available_permits(), BUDGET_BYTES); + } + + #[tokio::test] + async fn gzip_request_body_budget_follows_decoded_payload_lifetime() { + let (payload, frame_flags) = compress_payload(Bytes::from(vec![b'x'; 1024])); + assert_eq!(frame_flags, flags::GZIP_COMPRESSED); + let frame = Frame::new(7, MsgType::RequestBody, frame_flags, payload); + let required = request_body_queue_permits(&frame, REQUEST_BODY_QUEUE_BUDGET_BYTES) + .expect("gzip frame should fit the queue budget") as usize; + let budget = Arc::new(Semaphore::new(required)); + let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), required) + .await + .expect("frame should acquire the entire local budget"); + assert_eq!(budget.available_permits(), 0); + + let decoded = stream_handler::decode_request_body_frame(frame) + .expect("gzip request body should decode"); + assert_eq!(decoded, Bytes::from(vec![b'x'; 1024])); + assert_eq!(budget.available_permits(), 0); + + drop(decoded); + assert_eq!(budget.available_permits(), required); + } + + #[tokio::test] + async fn gzip_request_body_budget_releases_after_decode_error() { + let frame = Frame::new( + 7, + MsgType::RequestBody, + flags::GZIP_COMPRESSED, + Bytes::from_static(b"not gzip"), + ); + let required = request_body_queue_permits(&frame, REQUEST_BODY_QUEUE_BUDGET_BYTES) + .expect("gzip frame should fit the queue budget") as usize; + let budget = Arc::new(Semaphore::new(required)); + let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), required) + .await + .expect("frame should acquire the entire local budget"); + assert_eq!(budget.available_permits(), 0); + + stream_handler::decode_request_body_frame(frame) + .expect_err("invalid gzip request body should fail"); + assert_eq!(budget.available_permits(), required); + } + #[tokio::test] async fn try_send_stream_error_emits_stream_error_frame() { let (high_tx, mut high_rx) = bounded_queue::(4); @@ -581,4 +754,39 @@ mod tests { assert!(!streams.contains_key(&7)); assert!(streams.contains_key(&9)); } + + #[test] + fn request_stream_id_rejects_zero_and_active_duplicates() { + let (tx, _rx) = mpsc::channel::(1); + let streams = HashMap::from([( + 7, + StreamDispatchTarget { + body_tx: tx, + response_window: Arc::new(StreamSendWindow::new(1024)), + }, + )]); + let mut active_handler_ids = HashSet::from([7]); + + assert_eq!( + validate_request_stream_id(&streams, &active_handler_ids, 0), + Err("invalid stream id") + ); + assert_eq!( + validate_request_stream_id(&streams, &active_handler_ids, 7), + Err("duplicate stream id") + ); + assert_eq!( + validate_request_stream_id(&streams, &active_handler_ids, 9), + Ok(()) + ); + + // The routing entry may be removed after a dispatch failure while the + // handler is still running; its reservation must continue to reject + // a new request with the same id. + active_handler_ids.insert(11); + assert_eq!( + validate_request_stream_id(&HashMap::new(), &active_handler_ids, 11), + Err("duplicate stream id") + ); + } } diff --git a/apps/aether-tunnel/src/tunnel/heartbeat.rs b/apps/aether-tunnel/src/tunnel/heartbeat.rs index 42405d639..0a35cb09c 100644 --- a/apps/aether-tunnel/src/tunnel/heartbeat.rs +++ b/apps/aether-tunnel/src/tunnel/heartbeat.rs @@ -346,10 +346,12 @@ fn normalize_upgrade_target(raw: String) -> Option { .strip_prefix("tunnel-v") .or_else(|| trimmed.strip_prefix("proxy-v")) .unwrap_or(trimmed); - if normalized == CURRENT_VERSION { + let target = semver::Version::parse(normalized).ok()?; + let current = semver::Version::parse(CURRENT_VERSION).ok()?; + if target <= current { return None; } - Some(normalized.to_string()) + Some(target.to_string()) } fn maybe_trigger_upgrade(version: Option) { @@ -402,7 +404,10 @@ mod tests { use arc_swap::ArcSwap; use clap::Parser; - use super::{build_heartbeat_payload, handle_ack, AckDecision, HeartbeatSnapshot}; + use super::{ + build_heartbeat_payload, handle_ack, normalize_upgrade_target, AckDecision, + HeartbeatSnapshot, CURRENT_VERSION, + }; use crate::registration::client::AetherClient; use crate::runtime::DynamicConfig; use crate::state::{AppState, ServerContext, TunnelMetrics, TunnelRequestMetrics}; @@ -429,6 +434,7 @@ mod tests { tunnel_encryption_key: config.tunnel_encryption_key.clone(), node_name: config.node_name.clone(), node_id: Arc::new(RwLock::new("node-123".to_string())), + tunnel_generation: "test-generation-1".to_string(), aether_client: Arc::new(AetherClient::new( &config, &config.aether_url, @@ -489,6 +495,24 @@ mod tests { assert_eq!(server.dynamic.load().heartbeat_interval, 9); } + #[test] + fn remote_upgrade_accepts_only_strict_semver_upgrades() { + let current = semver::Version::parse(CURRENT_VERSION).expect("package version is semver"); + let target = semver::Version::new(current.major + 1, 0, 0); + + assert_eq!( + normalize_upgrade_target(format!("tunnel-v{target}")), + Some(target.to_string()) + ); + assert_eq!(normalize_upgrade_target(CURRENT_VERSION.to_string()), None); + assert_eq!(normalize_upgrade_target("0.0.1".to_string()), None); + assert_eq!( + normalize_upgrade_target("1.2.3/../../payload".to_string()), + None + ); + assert_eq!(normalize_upgrade_target("latest".to_string()), None); + } + #[tokio::test] async fn heartbeat_payload_reports_resource_usage_and_tunnel_error_diagnostics() { let config = sample_config(); diff --git a/apps/aether-tunnel/src/tunnel/mod.rs b/apps/aether-tunnel/src/tunnel/mod.rs index da2cf1ce2..052b873b6 100644 --- a/apps/aether-tunnel/src/tunnel/mod.rs +++ b/apps/aether-tunnel/src/tunnel/mod.rs @@ -231,8 +231,14 @@ fn mix_u64(mut x: u64) -> u64 { mod tests { use std::sync::atomic::AtomicU64; use std::sync::{Arc, Once}; - use std::time::Duration; + 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; @@ -303,7 +309,11 @@ mod tests { .await .expect("gateway should start"); - let state = sample_state(sample_config(&gateway_base_url)); + 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({ @@ -370,12 +380,45 @@ mod tests { gateway_base_url: &str, node_id: &str, ) -> Option<(StatusCode, String)> { + let payload = relay_probe_envelope(); + 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, + ); let response = reqwest::Client::new() .post(format!( "{gateway_base_url}/api/internal/tunnel/relay/{node_id}" )) .header("content-type", "application/octet-stream") - .body(relay_probe_envelope()) + .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()?; @@ -411,12 +454,40 @@ mod tests { async fn start_gateway_on_port( port: u16, ) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { - let state = GatewayAppState::new().expect("gateway test state should build"); + // The embedded gateway now fails closed when relay authentication is + // not configured. Keep this integration fixture explicitly authenticated. + 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); let router = build_router_with_state(state.clone()); let handle = spawn_router_on_port(port, router).await?; Ok((state, handle)) } + 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> { @@ -483,6 +554,7 @@ mod tests { 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, diff --git a/apps/aether-tunnel/src/tunnel/stream_handler.rs b/apps/aether-tunnel/src/tunnel/stream_handler.rs index 85f52c229..d77819713 100644 --- a/apps/aether-tunnel/src/tunnel/stream_handler.rs +++ b/apps/aether-tunnel/src/tunnel/stream_handler.rs @@ -4,14 +4,16 @@ //! and sends response frames back through the writer channel. use std::io; +use std::pin::Pin; use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::sync::Arc; use std::sync::Mutex; +use std::task::{Context, Poll}; use std::time::{Duration, Instant}; use aether_runtime::{AdmissionPermit, QueueSendError}; -use bytes::{Bytes, BytesMut}; +use bytes::Bytes; use futures_util::stream; use futures_util::StreamExt; use http_body_util::BodyExt; @@ -24,8 +26,8 @@ use crate::target_filter; use crate::upstream_client; use super::protocol::{ - compress_payload, decompress_if_gzip, flags, raw_payload, Frame as TunnelFrame, MsgType, - RequestMeta, ResetStreamPayload, ResponseMeta, + compress_payload, decompress_if_gzip_with_limit, flags, raw_payload, Frame as TunnelFrame, + MsgType, RequestMeta, ResetStreamPayload, ResponseMeta, }; use super::writer::FrameSender; @@ -39,6 +41,14 @@ const FLOW_CONTROL_WAIT_TIMEOUT: Duration = Duration::from_secs(5); const SLOW_STREAM_LOG_THRESHOLD: Duration = Duration::from_secs(2); const SUCCESS_LOG_SAMPLE_MODULO: u32 = 256; const REQUEST_BODY_SPOOL_QUEUE_CAPACITY: usize = 64; +/// Request bytes retained only to support same-origin 307/308 replay. This does +/// not limit or buffer the first upstream request, which remains streaming. +const REDIRECT_REPLAY_PER_REQUEST_BUDGET_BYTES: usize = 5 * 1024 * 1024; +const REDIRECT_REPLAY_MAX_CHUNKS: usize = 1024; +/// Bound replay retention across all active streams without reducing stream +/// admission or rejecting the original request when the cache is exhausted. +const REDIRECT_REPLAY_GLOBAL_BUDGET_BYTES: usize = 256 * 1024 * 1024; +static REDIRECT_REPLAY_BUFFERED_BYTES: AtomicUsize = AtomicUsize::new(0); #[derive(Debug)] pub(crate) struct StreamSendWindow { @@ -94,13 +104,73 @@ impl StreamSendWindow { } fn stream_reset_message(frame: &TunnelFrame) -> String { - if frame.msg_type == MsgType::ResetStream { - if let Ok(payload) = serde_json::from_slice::(&frame.payload) { - return payload.reason; - } + // Reset/error payloads are supplied by the peer and can contain arbitrary + // user data (or a provider error copied by the gateway). They are used as + // an `io::Error` below and may otherwise be echoed back in a later + // StreamError frame, so keep only a stable protocol-level category. + match frame.msg_type { + MsgType::ResetStream => "request reset by peer".to_string(), + MsgType::StreamError => "client cancelled request body".to_string(), + _ => "request body terminated".to_string(), } - String::from_utf8(frame.payload.to_vec()) - .unwrap_or_else(|_| "client cancelled request body".to_string()) +} + +/// Project an internal stream failure to a bounded, protocol-safe message. +/// +/// Hyper, URL, DNS, TLS, and proxy errors can include complete request URLs, +/// query credentials, private addresses, or implementation details. Tunnel +/// errors cross the authenticated tunnel and are eventually exposed by the +/// gateway, so never put those error strings on the wire (or in logs). +fn safe_stream_error_message(message: &str) -> &'static str { + let lower = message.trim().to_ascii_lowercase(); + if lower == "tunnel overloaded" { + return "tunnel overloaded"; + } + if lower == "tunnel admission unavailable" { + return "tunnel admission unavailable"; + } + if lower.contains("client cancelled") { + return "client cancelled request body"; + } + if lower.contains("response body timeout") { + return "upstream response body timeout"; + } + if lower.contains("flow_control_timeout") { + return "response flow-control timeout"; + } + if lower == "upstream timeout" || lower.contains("timed out") { + return "upstream timeout"; + } + if lower.contains("invalid") && lower.contains("url") { + return "invalid upstream URL"; + } + if lower.contains("unsupported") && lower.contains("scheme") { + return "unsupported upstream URL scheme"; + } + if lower.contains("target blocked") + || lower.contains("private/reserved") + || lower.contains("port not allowed") + || lower.contains("dns resolution") + || lower.contains("no public") + { + return "upstream target blocked"; + } + if lower.contains("gzip") || lower.contains("decompress") || lower.contains("request body") { + return "invalid request body"; + } + if lower.contains("redirect") { + return "upstream redirect failed"; + } + if lower.contains("connect") || lower.contains("tls") || lower.contains("proxy") { + return "upstream connect failed"; + } + if lower.contains("body") && (lower.contains("read") || lower.contains("response")) { + return "upstream response body failed"; + } + if lower.contains("request") { + return "upstream request failed"; + } + "upstream request failed" } fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) { @@ -162,14 +232,6 @@ const REDIRECT_DROP_BODY_HEADERS: &[&str] = &[ "content-type", "transfer-encoding", ]; -const REDIRECT_SENSITIVE_HEADERS: &[&str] = &[ - "authorization", - "cookie", - "cookie2", - "proxy-authorization", - "www-authenticate", -]; - #[derive(Debug, Clone)] enum ReplayableRequestBody { None, @@ -190,6 +252,8 @@ struct RequestTimeouts { #[derive(Debug)] struct RequestBodyReplayState { + budget_bytes: usize, + reserved_bytes: AtomicUsize, state: Mutex, ready: Notify, } @@ -200,15 +264,69 @@ enum RequestBodyReplayStatus { chunks: Vec, buffered_len: usize, }, - Ready(Bytes), + Ready { + chunks: Vec, + buffered_len: usize, + }, Empty, + NonReplayable, Error(String), } #[derive(Debug, Clone, PartialEq, Eq)] enum ReplayBodyResolution { Empty, - Replayable(Bytes), + Replayable { + chunks: Vec, + buffered_len: usize, + }, + NonReplayable, +} + +struct ReplayRequestBody { + chunks: std::vec::IntoIter, + remaining: u64, +} + +struct DecodedRequestBodyPayload { + decoded: Bytes, + _compressed_and_budget: Bytes, +} + +impl AsRef<[u8]> for DecodedRequestBodyPayload { + fn as_ref(&self) -> &[u8] { + self.decoded.as_ref() + } +} + +impl hyper::body::Body for ReplayRequestBody { + type Data = Bytes; + type Error = io::Error; + + fn poll_frame( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll, Self::Error>>> { + let body = self.get_mut(); + loop { + let Some(chunk) = body.chunks.next() else { + return Poll::Ready(None); + }; + if chunk.is_empty() { + continue; + } + body.remaining = body.remaining.saturating_sub(chunk.len() as u64); + return Poll::Ready(Some(Ok(BodyFrame::data(chunk)))); + } + } + + fn is_end_stream(&self) -> bool { + self.remaining == 0 + } + + fn size_hint(&self) -> hyper::body::SizeHint { + hyper::body::SizeHint::with_exact(self.remaining) + } } #[derive(Debug)] @@ -343,6 +461,7 @@ fn log_stream_success(ctx: StreamLogContext<'_>, status: u16, duration: Duration } fn log_stream_failure(ctx: StreamLogContext<'_>, error: &str, duration: Duration) { + let error = safe_stream_error_message(error); match ctx.url { Some(url) => { warn!( @@ -382,6 +501,29 @@ impl PreparedRequestBody { .take() .unwrap_or_else(empty_request_body) } + + /// Prefer the bounded replay snapshot for the first request when redirect + /// replay is enabled. A replay body advertises an exact size to Hyper, so + /// HTTP/1 requests get a correct `Content-Length` instead of implicit + /// chunked framing. Requests that exceed the replay budget remain streamed. + async fn resolve_initial_replay_body(&mut self, deadline: Instant) -> Result<(), String> { + let ReplayableRequestBody::Pending(state) = &self.replay_body else { + return Ok(()); + }; + match state.wait_for_resolution(deadline).await? { + ReplayBodyResolution::Empty => { + self.first_request_body = Some(empty_request_body()); + } + ReplayBodyResolution::Replayable { + chunks, + buffered_len, + } => { + self.first_request_body = Some(replay_request_body(chunks, buffered_len)); + } + ReplayBodyResolution::NonReplayable => {} + } + Ok(()) + } } async fn prepare_redirect_request_body( @@ -396,7 +538,11 @@ async fn prepare_redirect_request_body( ReplayableRequestBody::Pending(state) => { match state.wait_for_resolution(deadline).await? { ReplayBodyResolution::Empty => Ok(Some(empty_request_body())), - ReplayBodyResolution::Replayable(body) => Ok(Some(buffered_request_body(body))), + ReplayBodyResolution::Replayable { + chunks, + buffered_len, + } => Ok(Some(replay_request_body(chunks, buffered_len))), + ReplayBodyResolution::NonReplayable => Ok(None), } } ReplayableRequestBody::NonReplayable => Ok(None), @@ -405,8 +551,10 @@ async fn prepare_redirect_request_body( } impl RequestBodyReplayState { - fn new() -> Self { + fn new(budget_bytes: usize) -> Self { Self { + budget_bytes, + reserved_bytes: AtomicUsize::new(0), state: Mutex::new(RequestBodyReplayStatus::Collecting { chunks: Vec::new(), buffered_len: 0, @@ -416,14 +564,113 @@ impl RequestBodyReplayState { } fn push_chunk(&self, payload: Bytes) { + let mut disable_replay = false; let mut state = self.state.lock().expect("request body replay state lock"); if let RequestBodyReplayStatus::Collecting { chunks, buffered_len, } = &mut *state { - *buffered_len = buffered_len.saturating_add(payload.len()); - chunks.push(payload); + let Some(next_len) = buffered_len.checked_add(payload.len()) else { + chunks.clear(); + *state = RequestBodyReplayStatus::NonReplayable; + drop(state); + self.release_reserved_bytes(); + self.ready.notify_waiters(); + return; + }; + let accounted_bytes = payload.len().checked_add(std::mem::size_of::()); + if next_len > self.budget_bytes + || chunks.len() >= REDIRECT_REPLAY_MAX_CHUNKS + || accounted_bytes.is_none_or(|bytes| !self.try_reserve_bytes(bytes)) + { + disable_replay = true; + chunks.clear(); + *state = RequestBodyReplayStatus::NonReplayable; + } else { + *buffered_len = next_len; + chunks.push(payload); + } + } + drop(state); + if disable_replay { + self.release_reserved_bytes(); + self.ready.notify_waiters(); + } + } + + fn try_reserve_bytes(&self, bytes: usize) -> bool { + if bytes == 0 { + return true; + } + let mut current = REDIRECT_REPLAY_BUFFERED_BYTES.load(Ordering::Acquire); + loop { + let Some(next) = current.checked_add(bytes) else { + return false; + }; + if next > REDIRECT_REPLAY_GLOBAL_BUDGET_BYTES { + return false; + } + match REDIRECT_REPLAY_BUFFERED_BYTES.compare_exchange_weak( + current, + next, + Ordering::AcqRel, + Ordering::Acquire, + ) { + Ok(_) => { + self.reserved_bytes.fetch_add(bytes, Ordering::Release); + return true; + } + Err(observed) => current = observed, + } + } + } + + fn release_reserved_bytes(&self) { + let reserved = self.reserved_bytes.swap(0, Ordering::AcqRel); + if reserved > 0 { + REDIRECT_REPLAY_BUFFERED_BYTES.fetch_sub(reserved, Ordering::AcqRel); + } + } + + fn discard(&self) { + { + let mut state = self.state.lock().expect("request body replay state lock"); + match &*state { + RequestBodyReplayStatus::Collecting { .. } + | RequestBodyReplayStatus::Ready { .. } + | RequestBodyReplayStatus::Empty => { + *state = RequestBodyReplayStatus::NonReplayable; + } + RequestBodyReplayStatus::NonReplayable | RequestBodyReplayStatus::Error(_) => { + return; + } + } + } + self.release_reserved_bytes(); + self.ready.notify_waiters(); + } + + /// Disable replay while retaining the queued body for the first streaming + /// request. This prevents the preflight waiter from deadlocking when the + /// bounded spool queue fills before the request starts. + fn disable_replay(&self) { + let changed = { + let mut state = self.state.lock().expect("request body replay state lock"); + match &*state { + RequestBodyReplayStatus::Collecting { .. } + | RequestBodyReplayStatus::Ready { .. } => { + *state = RequestBodyReplayStatus::NonReplayable; + true + } + RequestBodyReplayStatus::Empty + | RequestBodyReplayStatus::NonReplayable + | RequestBodyReplayStatus::Error(_) => false, + } + }; + if changed { + self.release_reserved_bytes(); + self.ready.notify_waiters(); } } @@ -439,11 +686,10 @@ impl RequestBodyReplayState { if buffered_len == 0 { RequestBodyReplayStatus::Empty } else { - let mut buffered = BytesMut::with_capacity(buffered_len); - for chunk in chunks { - buffered.extend_from_slice(&chunk); + RequestBodyReplayStatus::Ready { + chunks, + buffered_len, } - RequestBodyReplayStatus::Ready(buffered.freeze()) } } terminal => terminal, @@ -461,19 +707,30 @@ impl RequestBodyReplayState { let mut state = self.state.lock().expect("request body replay state lock"); *state = RequestBodyReplayStatus::Error(message); } + self.release_reserved_bytes(); self.ready.notify_waiters(); } async fn wait_for_resolution(&self, deadline: Instant) -> Result { loop { + let notified = self.ready.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); let resolution = { let state = self.state.lock().expect("request body replay state lock"); match &*state { RequestBodyReplayStatus::Collecting { .. } => None, - RequestBodyReplayStatus::Ready(body) => { - Some(Ok(ReplayBodyResolution::Replayable(body.clone()))) - } + RequestBodyReplayStatus::Ready { + chunks, + buffered_len, + } => Some(Ok(ReplayBodyResolution::Replayable { + chunks: chunks.clone(), + buffered_len: *buffered_len, + })), RequestBodyReplayStatus::Empty => Some(Ok(ReplayBodyResolution::Empty)), + RequestBodyReplayStatus::NonReplayable => { + Some(Ok(ReplayBodyResolution::NonReplayable)) + } RequestBodyReplayStatus::Error(message) => Some(Err(message.clone())), } }; @@ -484,17 +741,82 @@ impl RequestBodyReplayState { let Some(remaining) = remaining_timeout(deadline) else { return Err("upstream timeout".to_string()); }; - tokio::time::timeout(remaining, self.ready.notified()) + tokio::time::timeout(remaining, &mut notified) .await .map_err(|_| "upstream timeout".to_string())?; } } } +impl Drop for RequestBodyReplayState { + fn drop(&mut self) { + self.release_reserved_bytes(); + } +} + fn follow_redirects_enabled(meta: &RequestMeta) -> bool { meta.follow_redirects == Some(true) } +/// Validate URL syntax at the tunnel trust boundary before any request body +/// is consumed or a connection is attempted. +/// +/// The gateway performs the same checks when it builds `RequestMeta`, but the +/// tunnel must not rely on a remote peer having constructed metadata through a +/// particular code path. In particular, URL userinfo and fragments are not +/// valid upstream request components: userinfo can alter authority parsing and +/// fragments must never cross an HTTP request boundary. +fn validate_tunnel_upstream_url( + url: &url::Url, + allow_private_targets: bool, +) -> Result<(), &'static str> { + if url.host_str().is_none() { + return Err("invalid upstream URL"); + } + if !matches!(url.scheme(), "http" | "https") { + return Err("unsupported upstream URL scheme"); + } + if !url.username().is_empty() || url.password().is_some() { + return Err("invalid upstream URL"); + } + if url.fragment().is_some() { + return Err("invalid upstream URL"); + } + + // Domain names are checked again against the resolved DNS answers by + // `target_filter::validate_target`. Reject literal private/reserved + // addresses here as well when the tunnel's private-target policy is + // disabled, so a metadata URL cannot bypass that policy through a + // different parser or connector. Explicitly enabled private-target + // deployments keep their existing behavior (including loopback HTTP + // endpoints). + if !allow_private_targets { + let literal_ip = match url.host() { + Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)), + Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)), + _ => None, + }; + if literal_ip.is_some_and(aether_http::is_private_or_reserved_ip) { + return Err("upstream target blocked"); + } + } + + Ok(()) +} + +fn validate_tunnel_redirect_url(url: &url::Url) -> Result<(), &'static str> { + if url.host_str().is_none() { + return Err("invalid upstream URL"); + } + if !matches!(url.scheme(), "http" | "https") { + return Err("unsupported upstream URL scheme"); + } + if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() { + return Err("invalid upstream URL"); + } + Ok(()) +} + fn request_likely_has_body( method: &hyper::Method, headers: &std::collections::HashMap, @@ -521,11 +843,19 @@ fn request_likely_has_body( fn sanitize_upstream_headers( headers: &std::collections::HashMap, ) -> Vec<(String, String)> { + let connection_declared = aether_http::connection_declared_header_names( + headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(hyper::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); headers .iter() .filter_map(|(key, value)| { let normalized = key.to_ascii_lowercase(); - if BLOCKED_HEADERS.contains(&normalized.as_str()) { + if BLOCKED_HEADERS.contains(&normalized.as_str()) + || connection_declared.contains(&normalized) + { None } else { Some((key.clone(), value.clone())) @@ -534,6 +864,44 @@ fn sanitize_upstream_headers( .collect() } +/// Return a single, syntactically valid request length that can safely be +/// applied to the streamed body. The original header is otherwise removed +/// because tunnel compression/decoding may change the bytes seen upstream. +/// Conflicting case variants and `Transfer-Encoding` are deliberately treated +/// as unknown framing rather than forwarded as an ambiguous pair. +fn validated_request_content_length( + headers: &std::collections::HashMap, +) -> Option { + if headers + .keys() + .any(|name| name.eq_ignore_ascii_case("transfer-encoding")) + { + return None; + } + + let values = headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("content-length")) + .map(|(_, value)| value.as_str()) + .collect::>(); + if values.is_empty() { + return None; + } + + let parsed = values + .iter() + .map(|value| { + let value = value.trim_matches(|character| matches!(character, ' ' | '\t')); + if value.is_empty() || !value.bytes().all(|byte| byte.is_ascii_digit()) { + return None; + } + value.parse::().ok() + }) + .collect::>>()?; + let first = *parsed.first()?; + parsed.iter().all(|value| *value == first).then_some(first) +} + fn apply_upstream_headers(headers: &mut hyper::HeaderMap, values: &[(String, String)]) { for (key, value) in values { if let (Ok(name), Ok(value)) = ( @@ -549,13 +917,43 @@ fn empty_request_body() -> upstream_client::UpstreamRequestBody { upstream_client::stream_request_body(stream::empty::, io::Error>>()) } -fn buffered_request_body(body: Bytes) -> upstream_client::UpstreamRequestBody { - upstream_client::full_request_body(body) +fn replay_request_body( + chunks: Vec, + buffered_len: usize, +) -> upstream_client::UpstreamRequestBody { + ReplayRequestBody { + chunks: chunks.into_iter(), + remaining: buffered_len as u64, + } + .boxed_unsync() +} + +pub(super) fn decode_request_body_frame(frame: TunnelFrame) -> Result { + if frame.is_gzip() { + let decoded = decompress_if_gzip_with_limit( + &frame, + aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES, + )?; + return Ok(Bytes::from_owner(DecodedRequestBodyPayload { + decoded, + _compressed_and_budget: frame.payload, + })); + } + if frame.payload.len() > aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!( + "decoded tunnel payload exceeds {} bytes", + aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES + ), + )); + } + Ok(frame.payload) } // Drain tunnel body frames on a detached task so the shared dispatcher is no -// longer coupled to upstream body polling. When requested, redirect replay -// retains a complete in-memory copy without a cumulative payload limit. +// longer coupled to upstream body polling. Redirect replay retains a bounded +// copy; crossing either replay budget only disables replay for this request. fn prepare_request_body( stream_id: u32, body_rx: mpsc::Receiver, @@ -565,7 +963,11 @@ fn prepare_request_body( frame_tx: FrameSender, ) -> PreparedRequestBody { let (spool_tx, spool_rx) = mpsc::channel(REQUEST_BODY_SPOOL_QUEUE_CAPACITY); - let replay_state = capture_for_redirects.then(|| Arc::new(RequestBodyReplayState::new())); + let replay_state = capture_for_redirects.then(|| { + Arc::new(RequestBodyReplayState::new( + REDIRECT_REPLAY_PER_REQUEST_BUDGET_BYTES, + )) + }); let replay_body = match replay_state.as_ref() { Some(state) => ReplayableRequestBody::Pending(Arc::clone(state)), None => ReplayableRequestBody::NonReplayable, @@ -602,55 +1004,6 @@ fn prepare_bodyless_request_body( } } -async fn collect_request_body_for_replay( - stream_id: u32, - mut body_rx: mpsc::Receiver, - body_size: Arc, - deadline: Instant, - frame_tx: &FrameSender, -) -> Result { - let mut body = BytesMut::new(); - - loop { - let frame = recv_body_frame_with_deadline(&mut body_rx, deadline).await?; - let Some(frame) = frame else { - return Ok(body.freeze()); - }; - - match frame.msg_type { - MsgType::RequestBody => { - let end_stream = frame.is_end_stream(); - let payload = decompress_if_gzip(&frame) - .map_err(|error| format!("gzip decompress failed: {error}"))?; - - if !payload.is_empty() { - body_size.fetch_add(payload.len(), Ordering::Relaxed); - try_send_window_update(frame_tx, stream_id, payload.len()); - body.extend_from_slice(&payload); - } - - if end_stream { - return Ok(body.freeze()); - } - } - MsgType::StreamError | MsgType::ResetStream => { - return Err(stream_reset_message(&frame)); - } - MsgType::StreamEnd => return Ok(body.freeze()), - _ => continue, - } - } -} - -fn replay_body_from_buffered(body: Bytes) -> ReplayableRequestBody { - let state = Arc::new(RequestBodyReplayState::new()); - if !body.is_empty() { - state.push_chunk(body); - } - state.finish(); - ReplayableRequestBody::Pending(state) -} - async fn recv_body_frame_with_deadline( body_rx: &mut mpsc::Receiver, deadline: Instant, @@ -692,7 +1045,12 @@ async fn spool_request_body( if let Some(state) = &replay_state { state.fail(message.clone()); } - let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::Error(message)).await; + let _ = send_spool_event( + &mut spool_tx, + SpoolBodyEvent::Error(message), + replay_state.as_ref(), + ) + .await; return; } }; @@ -701,22 +1059,27 @@ async fn spool_request_body( if let Some(state) = &replay_state { state.finish(); } - let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::End).await; + let _ = + send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()).await; return; }; match frame.msg_type { MsgType::RequestBody => { let end_stream = frame.is_end_stream(); - let payload = match decompress_if_gzip(&frame) { + let payload = match decode_request_body_frame(frame) { Ok(payload) => payload, Err(error) => { let message = format!("gzip decompress failed: {error}"); if let Some(state) = &replay_state { state.fail(message.clone()); } - let _ = - send_spool_event(&mut spool_tx, SpoolBodyEvent::Error(message)).await; + let _ = send_spool_event( + &mut spool_tx, + SpoolBodyEvent::Error(message), + replay_state.as_ref(), + ) + .await; return; } }; @@ -727,9 +1090,13 @@ async fn spool_request_body( if let Some(state) = &replay_state { state.push_chunk(payload.clone()); } - if send_spool_event(&mut spool_tx, SpoolBodyEvent::Data(payload)) - .await - .is_err() + if send_spool_event( + &mut spool_tx, + SpoolBodyEvent::Data(payload), + replay_state.as_ref(), + ) + .await + .is_err() { if let Some(state) = &replay_state { state.fail("request body replay channel closed".to_string()); @@ -742,7 +1109,9 @@ async fn spool_request_body( if let Some(state) = &replay_state { state.finish(); } - let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::End).await; + let _ = + send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()) + .await; return; } } @@ -751,14 +1120,20 @@ async fn spool_request_body( if let Some(state) = &replay_state { state.fail(message.clone()); } - let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::Error(message)).await; + let _ = send_spool_event( + &mut spool_tx, + SpoolBodyEvent::Error(message), + replay_state.as_ref(), + ) + .await; return; } MsgType::StreamEnd => { if let Some(state) = &replay_state { state.finish(); } - let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::End).await; + let _ = send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()) + .await; return; } _ => continue, @@ -769,8 +1144,18 @@ async fn spool_request_body( async fn send_spool_event( spool_tx: &mut mpsc::Sender, event: SpoolBodyEvent, + replay_state: Option<&Arc>, ) -> Result<(), ()> { - spool_tx.send(event).await.map_err(|_| ()) + match spool_tx.try_send(event) { + Ok(()) => Ok(()), + Err(mpsc::error::TrySendError::Closed(_)) => Err(()), + Err(mpsc::error::TrySendError::Full(event)) => { + if let Some(state) = replay_state { + state.disable_replay(); + } + spool_tx.send(event).await.map_err(|_| ()) + } + } } fn remove_headers_case_insensitive(headers: &mut Vec<(String, String)>, blocked: &[&str]) { @@ -780,16 +1165,13 @@ fn remove_headers_case_insensitive(headers: &mut Vec<(String, String)>, blocked: }); } -fn strip_sensitive_headers_for_redirect( - headers: &mut Vec<(String, String)>, - next: &url::Url, - previous: &url::Url, -) { - let cross_host = next.host_str() != previous.host_str() - || next.port_or_known_default() != previous.port_or_known_default(); - if cross_host { - remove_headers_case_insensitive(headers, REDIRECT_SENSITIVE_HEADERS); - } +fn redirect_urls_have_same_origin(left: &url::Url, right: &url::Url) -> bool { + left.scheme().eq_ignore_ascii_case(right.scheme()) + && left + .host_str() + .zip(right.host_str()) + .is_some_and(|(left, right)| left.eq_ignore_ascii_case(right)) + && left.port_or_known_default() == right.port_or_known_default() } fn resolve_redirect( @@ -830,16 +1212,21 @@ fn resolve_redirect( let Ok(next_url) = current_url.join(location) else { return RedirectDecision::Stop; }; - match next_url.scheme() { - "http" | "https" => {} - _ => return RedirectDecision::Stop, + // A same-origin redirect can still smuggle credentials or a fragment into + // the next request. Validate the resolved URL before considering it for + // replay; the target filter performs the address policy check when it is + // actually connected. + if validate_tunnel_redirect_url(&next_url).is_err() { + return RedirectDecision::Stop; } if redirects_followed >= MAX_REDIRECTS { return RedirectDecision::Error("too many redirects"); } + if !redirect_urls_have_same_origin(current_url, &next_url) { + return RedirectDecision::Stop; + } - strip_sensitive_headers_for_redirect(&mut next_headers, &next_url, current_url); RedirectDecision::Follow { method: next_method, url: next_url, @@ -866,9 +1253,9 @@ async fn execute_upstream_request( let port = current_url.port_or_known_default().unwrap_or(443); let dns_start = Instant::now(); - { + let validated_addrs = { let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports); - if let Err(error) = target_filter::validate_target( + match target_filter::validate_target( host, port, &allowed_ports, @@ -877,18 +1264,28 @@ async fn execute_upstream_request( ) .await { - server.metrics.dns_failures.fetch_add(1, Ordering::Release); - return Err(format!("target blocked: {error}")); + Ok(addrs) => addrs, + Err(_error) => { + server.metrics.dns_failures.fetch_add(1, Ordering::Release); + // Keep the detailed filter error out of the tunnel response; + // the request URL origin is already present in the structured + // failure log context. + return Err("upstream target blocked".to_string()); + } } - } + }; let dns_ms = dns_start.elapsed().as_millis() as u64; + let validated_target = + upstream_client::ValidatedUpstreamTarget::new(current_url, validated_addrs)?; + let client_key = upstream_client::upstream_client_pool_key( meta.provider_id.as_deref(), meta.endpoint_id.as_deref(), meta.key_id.as_deref(), meta.transport_profile.as_ref(), http1_only, + validated_target, ); let client = state.upstream_client_pool.get_or_build(client_key)?; @@ -896,19 +1293,8 @@ async fn execute_upstream_request( .method(method) .uri(current_url.as_str()) .body(request_body) - .map_err(|error| format!("invalid upstream request: {error}"))?; + .map_err(|_| "invalid upstream request".to_string())?; apply_upstream_headers(request.headers_mut(), headers); - if current_url.scheme() == "http" { - if let Some(value) = upstream_client::http_proxy_authorization_header( - state.config.upstream_proxy_url.as_deref(), - ) { - let value = hyper::header::HeaderValue::from_str(&value) - .map_err(|error| format!("invalid upstream proxy auth header: {error}"))?; - request - .headers_mut() - .insert(hyper::header::PROXY_AUTHORIZATION, value); - } - } let connection_start = Instant::now(); let mut captured_connection = upstream_client::capture_connection(&mut request); @@ -928,9 +1314,9 @@ async fn execute_upstream_request( .failed_requests .fetch_add(1, Ordering::Release); let message = if error.is_connect() { - format!("upstream connect error: {error}") + "upstream connect failed".to_string() } else { - format!("upstream error: {error}") + "upstream request failed".to_string() }; return Err(message); } @@ -1021,8 +1407,21 @@ where { let status = response.status().as_u16(); let ttfb_ms = total_elapsed.as_millis() as u64; + let connection_declared = aether_http::connection_declared_header_names( + response + .headers() + .get_all(hyper::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); let mut resp_headers: Vec<(String, String)> = Vec::with_capacity(response.headers().len() + 1); for (key, value) in response.headers() { + let normalized = key.as_str().to_ascii_lowercase(); + if BLOCKED_HEADERS.contains(&normalized.as_str()) + || connection_declared.contains(&normalized) + { + continue; + } if let Ok(value) = value.to_str() { resp_headers.push((key.as_str().to_string(), value.to_string())); } @@ -1206,8 +1605,9 @@ where } Err(error) => { server.metrics.stream_errors.fetch_add(1, Ordering::Release); - warn!(stream_id, error = %error, "upstream body read error"); - let error_message = format!("upstream body read error: {error}"); + let error_kind = safe_stream_error_message(&error.to_string()); + warn!(stream_id, error_kind, "upstream body read error"); + let error_message = error_kind; log_stream_failure( stream_log_context( server, @@ -1217,10 +1617,10 @@ where redirect_count, request_body_size.load(Ordering::Relaxed), ), - &error_message, + error_message, total_elapsed, ); - send_error(frame_tx, stream_id, &format!("body read error: {error}")).await; + send_error(frame_tx, stream_id, error_message).await; return Some(total_elapsed); } } @@ -1277,12 +1677,22 @@ where fn upstream_client_pool_key_for_request( meta: &RequestMeta, ) -> upstream_client::UpstreamClientPoolKey { + let target_url = url::Url::parse(&meta.url).expect("test request URL should parse"); + let port = target_url + .port_or_known_default() + .expect("test request URL should have a port"); + let validated_target = upstream_client::ValidatedUpstreamTarget::new( + &target_url, + vec![std::net::SocketAddr::from(([203, 0, 113, 1], port))], + ) + .expect("test target should validate"); upstream_client::upstream_client_pool_key( meta.provider_id.as_deref(), meta.endpoint_id.as_deref(), meta.key_id.as_deref(), meta.transport_profile.as_ref(), meta.http1_only, + validated_target, ) } @@ -1412,30 +1822,27 @@ async fn handle_stream_inner( let mut current_method: hyper::Method = parse_request_method(&meta.method); let mut current_url = match url::Url::parse(&meta.url) { Ok(u) => u, - Err(e) => { + Err(_) => { log_stream_failure( stream_log_context(server, stream_id, ¤t_method, None, 0, 0), - &format!("invalid URL: {e}"), + "invalid upstream URL", Duration::ZERO, ); - send_error(frame_tx, stream_id, &format!("invalid URL: {e}")).await; + send_error(frame_tx, stream_id, "invalid upstream URL").await; return None; } }; - // Only allow http/https schemes (block file://, data://, etc.) - match current_url.scheme() { - "http" | "https" => {} - other => { - let error_message = format!("unsupported URL scheme: {other}"); - log_stream_failure( - stream_log_context(server, stream_id, ¤t_method, Some(¤t_url), 0, 0), - &error_message, - Duration::ZERO, - ); - send_error(frame_tx, stream_id, &error_message).await; - return None; - } + if let Err(error_message) = + validate_tunnel_upstream_url(¤t_url, state.config.allow_private_targets) + { + log_stream_failure( + stream_log_context(server, stream_id, ¤t_method, Some(¤t_url), 0, 0), + error_message, + Duration::ZERO, + ); + send_error(frame_tx, stream_id, error_message).await; + return None; } let overall_start = Instant::now(); @@ -1445,57 +1852,30 @@ async fn handle_stream_inner( .response_body_timeout .map(|timeout| overall_start + timeout); let follow_redirects = follow_redirects_enabled(&meta); + let request_has_body = request_likely_has_body(¤t_method, &meta.headers); let mut current_headers = sanitize_upstream_headers(&meta.headers); + if request_has_body { + if let Some(content_length) = validated_request_content_length(&meta.headers) { + current_headers.push(( + hyper::header::CONTENT_LENGTH.as_str().to_string(), + content_length.to_string(), + )); + } + } let first_byte_timeout = request_timeouts.first_byte_timeout; let request_body_size = Arc::new(AtomicUsize::new(0)); - let request_has_body = request_likely_has_body(¤t_method, &meta.headers); - let can_buffer_redirect_body = request_has_body && follow_redirects; - let request_body_mode = if can_buffer_redirect_body { - "buffered_fixed" - } else if request_has_body { + let request_body_mode = if request_has_body { "streaming" } else { "empty" }; - let mut prepared_body = if can_buffer_redirect_body { - let buffered_body = match collect_request_body_for_replay( - stream_id, - body_rx, - Arc::clone(&request_body_size), - first_byte_deadline, - frame_tx, - ) - .await - { - Ok(body) => body, - Err(message) => { - log_stream_failure( - stream_log_context( - server, - stream_id, - ¤t_method, - Some(¤t_url), - 0, - request_body_size.load(Ordering::Relaxed), - ), - &message, - overall_start.elapsed(), - ); - send_error(frame_tx, stream_id, &message).await; - return None; - } - }; - PreparedRequestBody { - first_request_body: Some(buffered_request_body(buffered_body.clone())), - replay_body: replay_body_from_buffered(buffered_body), - } - } else if request_has_body { + let mut prepared_body = if request_has_body { prepare_request_body( stream_id, body_rx, Arc::clone(&request_body_size), first_byte_deadline, - false, + follow_redirects, frame_tx.clone(), ) } else { @@ -1506,8 +1886,36 @@ async fn handle_stream_inner( let mut redirects_followed = 0usize; let mut next_request_body = None::; + if follow_redirects { + if let Err(message) = prepared_body + .resolve_initial_replay_body(first_byte_deadline) + .await + { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } + log_stream_failure( + stream_log_context( + server, + stream_id, + ¤t_method, + Some(¤t_url), + 0, + request_body_size.load(Ordering::Relaxed), + ), + &message, + overall_start.elapsed(), + ); + send_error(frame_tx, stream_id, &message).await; + return None; + } + } + loop { let Some(remaining) = remaining_timeout(first_byte_deadline) else { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } log_stream_failure( stream_log_context( server, @@ -1542,6 +1950,9 @@ async fn handle_stream_inner( { Ok(context) => context, Err(message) => { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } log_stream_failure( stream_log_context( server, @@ -1570,6 +1981,9 @@ async fn handle_stream_inner( redirects_followed, ) { RedirectDecision::Stop => { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } drop(admission_permit.take()); return relay_upstream_response( server, @@ -1603,6 +2017,14 @@ async fn handle_stream_inner( .await { Ok(Some(body)) => { + if body_mode == RedirectBodyMode::Empty { + if let ReplayableRequestBody::Pending(state) = + &prepared_body.replay_body + { + state.discard(); + } + prepared_body.replay_body = ReplayableRequestBody::None; + } redirects_followed += 1; current_method = method; current_url = url; @@ -1611,6 +2033,9 @@ async fn handle_stream_inner( continue; } Ok(None) => { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } drop(admission_permit.take()); return relay_upstream_response( server, @@ -1632,6 +2057,9 @@ async fn handle_stream_inner( .await; } Err(message) => { + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } log_stream_failure( stream_log_context( server, @@ -1649,7 +2077,11 @@ async fn handle_stream_inner( } }, RedirectDecision::Error(message) => { - let error_message = format!("upstream redirect error: {message}"); + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } + let error_message = + safe_stream_error_message(&format!("upstream redirect error: {message}")); log_stream_failure( stream_log_context( server, @@ -1659,15 +2091,18 @@ async fn handle_stream_inner( redirects_followed, request_body_size.load(Ordering::Relaxed), ), - &error_message, + error_message, overall_start.elapsed(), ); - send_error(frame_tx, stream_id, &error_message).await; + send_error(frame_tx, stream_id, error_message).await; return None; } } } + if let ReplayableRequestBody::Pending(state) = &prepared_body.replay_body { + state.discard(); + } drop(admission_permit.take()); return relay_upstream_response( server, @@ -1692,21 +2127,23 @@ async fn handle_stream_inner( async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) { // Error frames use best-effort delivery — don't block if writer is congested + let safe_message = safe_stream_error_message(msg); let _ = send_frame( tx, TunnelFrame::new( stream_id, MsgType::StreamError, 0, - Bytes::from(msg.to_string()), + Bytes::from_static(safe_message.as_bytes()), ), ) .await; } async fn send_reset_stream(tx: &FrameSender, stream_id: u32, reason: &str) { + let safe_reason = safe_stream_error_message(reason); let payload = serde_json::to_vec(&ResetStreamPayload { - reason: reason.to_string(), + reason: safe_reason.to_string(), }) .expect("reset stream payload should serialize"); let _ = send_frame( @@ -1774,7 +2211,7 @@ fn build_prefixed_request_body( match frame.msg_type { MsgType::RequestBody => { let end_stream = frame.is_end_stream(); - let payload = match decompress_if_gzip(&frame) { + let payload = match decode_request_body_frame(frame) { Ok(payload) => payload, Err(error) => { let err = @@ -1828,6 +2265,7 @@ mod tests { use axum::http::{header, HeaderMap, Response, StatusCode}; use axum::routing::{get, post}; use axum::Router; + use bytes::BytesMut; use futures_util::Sink; use tokio::task::JoinHandle; use tokio_tungstenite::tungstenite::{Error as WebSocketError, Message}; @@ -1841,7 +2279,7 @@ mod tests { use crate::tunnel::client::build_tls_config; fn completed_replay_body(body: Bytes) -> ReplayableRequestBody { - let state = Arc::new(RequestBodyReplayState::new()); + let state = Arc::new(RequestBodyReplayState::new(body.len().max(1))); if !body.is_empty() { state.push_chunk(body); } @@ -1991,7 +2429,7 @@ mod tests { assert_eq!(second, Bytes::from_static(b"world")); assert!(body.frame().await.is_none()); - let mut replay = prepare_redirect_request_body( + let replay = prepare_redirect_request_body( prepared.replay_body.clone(), RedirectBodyMode::Replay, Instant::now() + Duration::from_secs(1), @@ -1999,16 +2437,12 @@ mod tests { .await .expect("redirect replay should resolve") .expect("body should be replayable"); - let frame = replay - .frame() + let replay = replay + .collect() .await - .expect("replayed frame should exist") - .expect("replayed frame should be ok"); - assert_eq!( - frame.into_data().expect("data frame"), - Bytes::from_static(b"hello world") - ); - assert!(replay.frame().await.is_none()); + .expect("replayed body should be readable") + .to_bytes(); + assert_eq!(replay, Bytes::from_static(b"hello world")); assert_eq!(body_size.load(Ordering::Relaxed), 11); let window_update_bytes = collect_emitted_frames(frame_tx, sent, writer_handle) .await @@ -2025,6 +2459,24 @@ mod tests { assert_eq!(window_update_bytes, 11); } + #[tokio::test] + async fn replay_state_disables_and_releases_cache_after_per_request_budget() { + let state = RequestBodyReplayState::new(5); + state.push_chunk(Bytes::from_static(b"123")); + assert!(state.reserved_bytes.load(Ordering::Acquire) > 0); + + state.push_chunk(Bytes::from_static(b"456")); + + assert_eq!(state.reserved_bytes.load(Ordering::Acquire), 0); + assert_eq!( + state + .wait_for_resolution(Instant::now() + Duration::from_secs(1)) + .await + .expect("over-budget replay should resolve without failing the request"), + ReplayBodyResolution::NonReplayable + ); + } + #[test] fn selects_http1_only_client_when_request_metadata_requires_it() { let default_meta = sample_request_meta(); @@ -2041,6 +2493,87 @@ mod tests { ); } + #[test] + fn stream_error_projection_never_returns_upstream_details() { + let secret_error = concat!( + "upstream connect error: error sending request for url (", + "https://user:password@example.test/v1/models?api_key=query-secret", + ")" + ); + assert_eq!( + safe_stream_error_message(secret_error), + "upstream connect failed" + ); + assert_eq!( + safe_stream_error_message( + "upstream body read error: authorization Bearer secret-token at 10.0.0.4" + ), + "upstream response body failed" + ); + assert_eq!( + safe_stream_error_message("invalid URL: https://user:pass@example.test/?token=secret"), + "invalid upstream URL" + ); + } + + #[test] + fn tunnel_upstream_url_validation_rejects_ambiguous_url_components() { + for raw in [ + "https://user:password@example.test/v1", + "https://user@example.test/v1", + "https://example.test/v1#fragment", + "file:///etc/passwd", + ] { + let url = url::Url::parse(raw).expect("fixture URL should parse"); + assert!( + validate_tunnel_upstream_url(&url, true).is_err(), + "URL should be rejected at the tunnel boundary: {raw}" + ); + } + + assert!(validate_tunnel_upstream_url( + &url::Url::parse("https://example.test/v1?api_key=query-secret") + .expect("query URL should parse"), + true, + ) + .is_ok()); + } + + #[test] + fn tunnel_upstream_url_validation_applies_literal_target_policy() { + let private = url::Url::parse("https://10.0.0.8/private").expect("private URL"); + assert!(validate_tunnel_upstream_url(&private, false).is_err()); + assert!(validate_tunnel_upstream_url(&private, true).is_ok()); + + // Loopback remains available to explicitly enabled local deployments; + // disabling private targets rejects it before connection setup. + let loopback = url::Url::parse("http://127.0.0.1:8080/local").expect("loopback URL"); + assert!(validate_tunnel_upstream_url(&loopback, false).is_err()); + assert!(validate_tunnel_upstream_url(&loopback, true).is_ok()); + } + + #[test] + fn peer_reset_payload_is_not_echoed_into_request_errors() { + let frame = TunnelFrame::new( + 1, + MsgType::StreamError, + 0, + Bytes::from_static(b"Authorization: Bearer secret-token"), + ); + assert_eq!( + stream_reset_message(&frame), + "client cancelled request body" + ); + + let reset = TunnelFrame::new( + 1, + MsgType::ResetStream, + 0, + Bytes::from_static(b"https://user:pass@example.test/?token=secret"), + ); + assert_eq!(stream_reset_message(&reset), "request reset by peer"); + } + #[test] fn upstream_client_pool_key_isolates_accounts() { let mut first = sample_request_meta(); @@ -2155,7 +2688,7 @@ mod tests { } #[test] - fn resolve_redirect_strips_sensitive_headers_for_cross_host_redirect() { + fn resolve_redirect_stops_cross_origin_redirects() { let current_url = url::Url::parse("https://redirect-a.test/start").expect("url"); let response = Response::builder() .status(StatusCode::FOUND) @@ -2169,26 +2702,186 @@ mod tests { &hyper::Method::GET, &[ ("authorization".into(), "Bearer secret".into()), + ("api-key".into(), "api-key-secret".into()), ("cookie".into(), "sid=123".into()), + ("x-api-key".into(), "x-api-key-secret".into()), + ("x-goog-api-key".into(), "google-secret".into()), ("x-custom".into(), "keep".into()), ], &ReplayableRequestBody::None, 0, ); - match decision { - RedirectDecision::Follow { headers, .. } => { - assert!(!headers - .iter() - .any(|(name, _)| name.eq_ignore_ascii_case("authorization"))); - assert!(!headers - .iter() - .any(|(name, _)| name.eq_ignore_ascii_case("cookie"))); - assert!(headers - .iter() - .any(|(name, value)| name.eq_ignore_ascii_case("x-custom") && value == "keep")); - } - other => panic!("unexpected redirect decision: {other:?}"), + assert_eq!(decision, RedirectDecision::Stop); + } + + #[test] + fn resolve_redirect_stops_https_to_http_downgrade() { + let current_url = url::Url::parse("https://redirect.test/start").expect("url"); + let response = Response::builder() + .status(StatusCode::FOUND) + .header(header::LOCATION, "http://redirect.test/final") + .body(()) + .expect("response"); + + let decision = resolve_redirect( + &response, + ¤t_url, + &hyper::Method::GET, + &[("authorization".into(), "Bearer secret".into())], + &ReplayableRequestBody::None, + 0, + ); + + assert_eq!(decision, RedirectDecision::Stop); + } + + #[test] + fn resolve_redirect_does_not_follow_userinfo_or_fragment_urls() { + let current_url = url::Url::parse("https://redirect.test/start").expect("url"); + for location in ["https://user:password@redirect.test/final", "/final#secret"] { + let response = Response::builder() + .status(StatusCode::FOUND) + .header(header::LOCATION, location) + .body(()) + .expect("response"); + + let decision = resolve_redirect( + &response, + ¤t_url, + &hyper::Method::GET, + &[], + &ReplayableRequestBody::None, + 0, + ); + assert_eq!(decision, RedirectDecision::Stop, "location: {location}"); + } + } + + #[test] + fn resolve_redirect_never_replays_post_body_cross_origin() { + let current_url = url::Url::parse("https://oauth.example/token").expect("url"); + let response = Response::builder() + .status(StatusCode::TEMPORARY_REDIRECT) + .header(header::LOCATION, "https://attacker.example/capture") + .body(()) + .expect("response"); + + let decision = resolve_redirect( + &response, + ¤t_url, + &hyper::Method::POST, + &[( + "content-type".into(), + "application/x-www-form-urlencoded".into(), + )], + &completed_replay_body(Bytes::from_static( + b"refresh_token=secret&client_secret=secret", + )), + 0, + ); + + assert_eq!(decision, RedirectDecision::Stop); + } + + #[test] + fn connection_declared_response_headers_are_not_relayed() { + let response = Response::builder() + .status(StatusCode::OK) + .header(header::CONNECTION, "x-hop-private, x-accel-redirect") + .header("x-hop-private", "secret") + .header("x-accel-redirect", "/internal") + .header("x-visible", "ok") + .body(()) + .expect("response"); + let declared = aether_http::connection_declared_header_names( + response + .headers() + .get_all(header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()), + ); + + assert!(declared.contains("x-hop-private")); + assert!(declared.contains("x-accel-redirect")); + assert!(!declared.contains("x-visible")); + } + + #[test] + fn connection_declared_request_headers_are_not_sent_upstream() { + let headers = std::collections::HashMap::from([ + ("Connection".to_string(), "x-hop-private".to_string()), + ("X-Hop-Private".to_string(), "secret".to_string()), + ("X-Visible".to_string(), "ok".to_string()), + ]); + + let sanitized = sanitize_upstream_headers(&headers); + + assert!(!sanitized + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("connection"))); + assert!(!sanitized + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("x-hop-private"))); + assert!(sanitized + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("x-visible") && value == "ok")); + } + + #[test] + fn validates_single_content_length_value() { + let headers = HashMap::from([("content-length".to_string(), "42".to_string())]); + + assert_eq!(validated_request_content_length(&headers), Some(42)); + } + + #[test] + fn accepts_identical_case_variant_content_lengths() { + let headers = HashMap::from([ + ("Content-Length".to_string(), " 42 ".to_string()), + ("content-length".to_string(), "42".to_string()), + ]); + + assert_eq!(validated_request_content_length(&headers), Some(42)); + } + + #[test] + fn rejects_conflicting_case_variant_content_lengths() { + let headers = HashMap::from([ + ("Content-Length".to_string(), "42".to_string()), + ("CONTENT-LENGTH".to_string(), "43".to_string()), + ]); + + assert_eq!(validated_request_content_length(&headers), None); + } + + #[test] + fn rejects_content_length_when_transfer_encoding_is_present() { + let headers = HashMap::from([ + ("Content-Length".to_string(), "42".to_string()), + ("Transfer-Encoding".to_string(), "chunked".to_string()), + ]); + + assert_eq!(validated_request_content_length(&headers), None); + } + + #[test] + fn rejects_empty_or_invalid_content_length_values() { + for value in [ + "", + " ", + "+42", + "-1", + "42, 42", + "not-a-length", + "18446744073709551616", + ] { + let headers = HashMap::from([("content-length".to_string(), value.to_string())]); + assert_eq!( + validated_request_content_length(&headers), + None, + "unexpectedly accepted Content-Length value {value:?}" + ); } } @@ -2507,6 +3200,67 @@ mod tests { assert_eq!(timing["redirect_count"], serde_json::json!(1)); } + #[tokio::test] + async fn cross_origin_redirect_is_preserved_without_a_second_connection() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener"); + let addr = listener.local_addr().expect("addr"); + let app = Router::new().route( + "/start", + get(move || async move { + Response::builder() + .status(StatusCode::FOUND) + .header( + header::LOCATION, + format!("http://127.0.0.1:{}/private", addr.port()), + ) + .body(Body::empty()) + .expect("redirect response") + }), + ); + let server = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("test server should run"); + }); + + let host = "redirect-to-private.test"; + let state = sample_state_for_port(addr.port()); + cache_test_host(&state, host, addr).await; + let server_ctx = sample_server(&state); + let (frame_tx, sent, writer_handle) = spawn_test_writer(); + let (_body_tx, body_rx) = mpsc::channel(1); + let mut meta = sample_request_meta(); + meta.url = format!("http://{host}:{}/start", addr.port()); + meta.follow_redirects = Some(true); + + handle_stream( + Arc::clone(&state), + server_ctx, + 19, + meta, + body_rx, + frame_tx.clone(), + test_response_window(), + ) + .await; + let result = collect_stream_result(frame_tx, sent, writer_handle).await; + server.abort(); + + assert!(result.error.is_none()); + let response = result.response.expect("redirect response metadata"); + assert_eq!(response.status, 302); + assert_eq!( + response + .headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("location")) + .map(|(_, value)| value.as_str()), + Some(format!("http://127.0.0.1:{}/private", addr.port()).as_str()) + ); + } + #[tokio::test] async fn preserves_redirect_response_when_follow_redirects_disabled() { let listener = tokio::net::TcpListener::bind("127.0.0.1:0") @@ -2571,7 +3325,7 @@ mod tests { } #[tokio::test] - async fn follows_307_redirect_for_fragmented_body_larger_than_legacy_replay_budget() { + async fn preserves_307_after_replay_budget_without_truncating_first_request() { const BODY_LEN: usize = 5 * 1024 * 1024 + 1; const REQUEST_FRAME_BYTES: usize = 32 * 1024; @@ -2610,10 +3364,9 @@ mod tests { .expect("test server should run"); }); - let host = "redirect-unlimited.test"; + let host = "redirect-over-budget.test"; let mut config = sample_config(); config.allowed_ports.push(addr.port()); - config.legacy_redirect_replay_budget_bytes_ignored = Some("1".to_string()); let state = sample_state_with_config(config); cache_test_host(&state, host, addr).await; let server_ctx = sample_server(&state); @@ -2665,8 +3418,16 @@ mod tests { result.error ); let response = result.response.expect("response metadata"); - assert_eq!(response.status, 200); - assert_eq!(result.body, Bytes::from_static(b"redirected")); + assert_eq!(response.status, 307); + assert_eq!( + response + .headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("location")) + .map(|(_, value)| value.as_str()), + Some("/final") + ); + assert!(result.body.is_empty()); } #[tokio::test] @@ -2828,6 +3589,7 @@ mod tests { tunnel_encryption_key: config.tunnel_encryption_key.clone(), node_name: config.node_name.clone(), node_id: Arc::new(std::sync::RwLock::new("node-1".to_string())), + tunnel_generation: "test-generation-1".to_string(), aether_client: Arc::new(AetherClient::new( &config, &config.aether_url, @@ -3004,13 +3766,21 @@ mod tests { for frame in collect_emitted_frames(frame_tx, sent, writer_handle).await { match frame.msg_type { MsgType::ResponseHeaders => { - let payload = decompress_if_gzip(&frame).expect("headers payload"); + let payload = decompress_if_gzip_with_limit( + &frame, + aether_contracts::tunnel::MAX_TUNNEL_RELAY_META_LEN, + ) + .expect("headers payload"); response = Some( serde_json::from_slice(&payload).expect("response metadata should decode"), ); } MsgType::ResponseBody => { - let payload = decompress_if_gzip(&frame).expect("body payload"); + let payload = decompress_if_gzip_with_limit( + &frame, + aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES, + ) + .expect("body payload"); body.extend_from_slice(&payload); } MsgType::StreamError => { diff --git a/apps/aether-tunnel/src/upstream_client.rs b/apps/aether-tunnel/src/upstream_client.rs index 1f0d1d044..a0f0c2967 100644 --- a/apps/aether-tunnel/src/upstream_client.rs +++ b/apps/aether-tunnel/src/upstream_client.rs @@ -1,8 +1,7 @@ use std::collections::HashMap; -use std::convert::Infallible; use std::future::Future; use std::io; -use std::net::IpAddr; +use std::net::{IpAddr, SocketAddr}; use std::pin::Pin; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; @@ -17,7 +16,7 @@ use aether_contracts::{ use bytes::Bytes; use futures_util::Stream; use http_body_util::combinators::UnsyncBoxBody; -use http_body_util::{BodyExt, Full, StreamBody}; +use http_body_util::{BodyExt, StreamBody}; use hyper::body::Frame; use hyper::rt; use hyper::Response; @@ -35,10 +34,9 @@ use tower_service::Service; use crate::config::Config; use crate::egress_proxy::{ - connect_proxy_tcp, http_connect, socks5_connect, ProxyConnectOptions, UpstreamProxyConfig, - UpstreamProxyScheme, + connect_validated_target_via_proxy, ProxyConnectOptions, UpstreamProxyConfig, }; -use crate::target_filter::{self, DnsCache}; +use crate::target_filter::DnsCache; type BoxError = Box; @@ -60,12 +58,75 @@ pub struct UpstreamClientPoolKey { pub profile_id: String, pub backend: String, pub http_mode: String, + pub validated_target: ValidatedUpstreamTarget, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct ValidatedUpstreamTarget { + scheme: String, + host: String, + port: u16, + addrs: Vec, +} + +impl ValidatedUpstreamTarget { + pub fn new(target_url: &url::Url, mut addrs: Vec) -> Result { + let scheme = target_url.scheme().to_ascii_lowercase(); + if !matches!(scheme.as_str(), "http" | "https") { + return Err(format!("unsupported upstream scheme {scheme}")); + } + let host = target_url + .host_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| "missing host in upstream URL".to_string())? + .trim_start_matches('[') + .trim_end_matches(']') + .to_ascii_lowercase(); + let port = target_url + .port_or_known_default() + .ok_or_else(|| "missing port in upstream URL".to_string())?; + if addrs.is_empty() { + return Err("validated upstream target has no addresses".to_string()); + } + if addrs.iter().any(|addr| addr.port() != port) { + return Err("validated upstream target address has the wrong port".to_string()); + } + addrs.sort_unstable(); + addrs.dedup(); + Ok(Self { + scheme, + host, + port, + addrs, + }) + } + + fn ensure_matches_uri(&self, uri: &Uri) -> Result<(), io::Error> { + let scheme = uri + .scheme_str() + .ok_or_else(|| io::Error::other("missing scheme"))?; + let host = uri_host(uri)?; + let port = uri_port_or_default(uri, scheme)?; + if !scheme.eq_ignore_ascii_case(&self.scheme) + || !host.eq_ignore_ascii_case(&self.host) + || port != self.port + { + return Err(io::Error::other( + "upstream connector target does not match its validated origin", + )); + } + Ok(()) + } + + fn addrs(&self) -> &[SocketAddr] { + &self.addrs + } } #[derive(Clone)] pub struct UpstreamClientPool { config: Arc, - dns_cache: Arc, clients: Arc>>, access_counter: Arc, } @@ -77,10 +138,9 @@ struct UpstreamClientPoolEntry { } impl UpstreamClientPool { - pub fn new(config: Arc, dns_cache: Arc) -> Self { + pub fn new(config: Arc, _dns_cache: Arc) -> Self { Self { config, - dns_cache, clients: Arc::new(Mutex::new(HashMap::new())), access_counter: Arc::new(AtomicU64::new(0)), } @@ -105,7 +165,7 @@ impl UpstreamClientPool { .eq_ignore_ascii_case(TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE); let client = build_upstream_client_with_protocol( &self.config, - Arc::clone(&self.dns_cache), + key.validated_target.clone(), http1_only, h2c_prior_knowledge, )?; @@ -156,6 +216,7 @@ pub fn upstream_client_pool_key( key_id: Option<&str>, profile: Option<&ResolvedTransportProfile>, http1_only: bool, + validated_target: ValidatedUpstreamTarget, ) -> UpstreamClientPoolKey { let profile_http_mode = profile .map(|profile| profile.http_mode.trim()) @@ -181,6 +242,7 @@ pub fn upstream_client_pool_key( .unwrap_or(DEFAULT_BACKEND) .to_string(), http_mode: http_mode.to_string(), + validated_target, } } @@ -201,18 +263,6 @@ fn validate_proxy_transport_backend(backend: &str) -> Result<(), String> { Err(format!("unsupported transport profile backend: {backend}")) } -pub fn http_proxy_authorization_header(proxy_url: Option<&str>) -> Option { - let proxy = proxy_url - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(|value| UpstreamProxyConfig::parse(value).ok())?; - if proxy.scheme() == UpstreamProxyScheme::Http { - proxy.basic_auth_header() - } else { - None - } -} - pub fn stream_request_body(stream: S) -> UpstreamRequestBody where S: Stream, io::Error>> + Send + 'static, @@ -220,9 +270,10 @@ where StreamBody::new(stream).boxed_unsync() } +#[cfg(test)] pub fn full_request_body(body: Bytes) -> UpstreamRequestBody { - Full::new(body) - .map_err(|err: Infallible| match err {}) + http_body_util::Full::new(body) + .map_err(|err: std::convert::Infallible| match err {}) .boxed_unsync() } @@ -241,18 +292,14 @@ pub struct RequestTiming { pub connection_reused: bool, } -#[derive(Clone)] -pub struct ValidatedResolver { - dns_cache: Arc, - allow_private: bool, +#[derive(Clone, Debug)] +struct PinnedResolver { + target: ValidatedUpstreamTarget, } -impl ValidatedResolver { - pub fn new(dns_cache: Arc, allow_private: bool) -> Self { - Self { - dns_cache, - allow_private, - } +impl PinnedResolver { + fn new(target: ValidatedUpstreamTarget) -> Self { + Self { target } } } @@ -268,7 +315,7 @@ impl Iterator for ValidatedAddrs { } } -impl Service for ValidatedResolver { +impl Service for PinnedResolver { type Response = ValidatedAddrs; type Error = io::Error; type Future = Pin> + Send>>; @@ -278,22 +325,16 @@ impl Service for ValidatedResolver { } fn call(&mut self, name: Name) -> Self::Future { - let dns_cache = Arc::clone(&self.dns_cache); - let allow_private = self.allow_private; - let host = name.as_str().to_string(); + let requested_host = name.as_str().to_string(); + let target = self.target.clone(); Box::pin(async move { - if let Some(addrs) = dns_cache.get_by_host(&host).await { - return Ok(ValidatedAddrs { - inner: (*addrs).clone().into_iter(), - }); + if !requested_host.eq_ignore_ascii_case(&target.host) { + return Err(io::Error::other( + "DNS request does not match the validated upstream host", + )); } - - let resolved = - target_filter::resolve_public_addrs(&host, 0, allow_private, dns_cache.as_ref()) - .await - .map_err(|err| io::Error::other(err.to_string()))?; Ok(ValidatedAddrs { - inner: resolved.into_iter(), + inner: target.addrs.into_iter(), }) }) } @@ -301,9 +342,10 @@ impl Service for ValidatedResolver { #[derive(Clone)] pub struct InstrumentedConnector { - http: HttpConnector, + http: HttpConnector, tls_config: Arc, proxy: Option, + validated_target: ValidatedUpstreamTarget, connect_timeout: Duration, tcp_nodelay: bool, tcp_keepalive: Option, @@ -319,9 +361,13 @@ impl Service for InstrumentedConnector { } fn call(&mut self, dst: Uri) -> Self::Future { + if let Err(error) = self.validated_target.ensure_matches_uri(&dst) { + return Box::pin(async move { Err(error.into()) }); + } let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase()); let tls_config = Arc::clone(&self.tls_config); if let Some(proxy) = self.proxy.clone() { + let validated_target = self.validated_target.clone(); let options = ProxyConnectOptions { connect_timeout: self.connect_timeout, tcp_nodelay: self.tcp_nodelay, @@ -330,7 +376,16 @@ impl Service for InstrumentedConnector { }; let connect_start = std::time::Instant::now(); return Box::pin(async move { - connect_via_proxy(dst, scheme, tls_config, proxy, options, connect_start).await + connect_via_proxy( + dst, + scheme, + tls_config, + proxy, + validated_target, + options, + connect_start, + ) + .await }); } let connecting = self.http.call(dst.clone()); @@ -381,39 +436,25 @@ async fn connect_via_proxy( scheme: Option, tls_config: Arc, proxy: UpstreamProxyConfig, + validated_target: ValidatedUpstreamTarget, options: ProxyConnectOptions, connect_start: std::time::Instant, ) -> Result { let scheme = scheme.ok_or_else(|| io::Error::other("missing scheme"))?; - let target_host = uri_host(&dst)?; - let target_port = uri_port_or_default(&dst, &scheme)?; - - let mut tcp = connect_proxy_tcp( - &proxy, - options.connect_timeout, - options.tcp_nodelay, - options.tcp_keepalive, - options.ip_family, - ) - .await?; - - match proxy.scheme() { - UpstreamProxyScheme::Http => { - if scheme == "https" { - http_connect( - &mut tcp, - &target_authority(&target_host, target_port), - &proxy, - ) - .await?; - } else if scheme != "http" { - return Err(io::Error::other(format!("unsupported scheme {scheme}")).into()); + let mut last_error = None; + let mut connected = None; + for target_addr in validated_target.addrs().iter().copied() { + match connect_validated_target_via_proxy(&proxy, target_addr, options).await { + Ok(tcp) => { + connected = Some(tcp); + break; } - } - UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => { - socks5_connect(&mut tcp, &proxy, &target_host, target_port).await?; + Err(error) => last_error = Some(error), } } + let tcp = connected.ok_or_else(|| { + last_error.unwrap_or_else(|| io::Error::other("validated upstream target has no addresses")) + })?; let connect_ms = connect_start.elapsed().as_millis() as u64; @@ -421,7 +462,7 @@ async fn connect_via_proxy( "http" => Ok(TimedConn::new( MaybeHttpsStream::Http { stream: TokioIo::new(tcp), - is_proxy: proxy.scheme() == UpstreamProxyScheme::Http, + is_proxy: false, }, ConnectTiming { connect_ms, @@ -465,24 +506,13 @@ fn uri_port_or_default(uri: &Uri, scheme: &str) -> Result { .ok_or_else(|| io::Error::other(format!("missing port for scheme {scheme}"))) } -fn target_authority(host: &str, port: u16) -> String { - if host.contains(':') && !host.starts_with('[') { - format!("[{host}]:{port}") - } else { - format!("{host}:{port}") - } -} - fn build_upstream_client_with_protocol( config: &Config, - dns_cache: Arc, + validated_target: ValidatedUpstreamTarget, http1_only: bool, h2c_prior_knowledge: bool, ) -> Result { - let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new( - dns_cache, - config.allow_private_targets, - )); + let mut http = HttpConnector::new_with_resolver(PinnedResolver::new(validated_target.clone())); http.enforce_http(false); http.set_connect_timeout(Some(Duration::from_secs( config.upstream_connect_timeout_secs, @@ -499,6 +529,7 @@ fn build_upstream_client_with_protocol( let connector = InstrumentedConnector { http, tls_config: build_tls_config(http1_only), + validated_target, proxy: config .upstream_proxy_url .as_deref() @@ -791,12 +822,19 @@ mod tests { header_fingerprint: None, extra: None, }; + let target_url = url::Url::parse("https://example.com/").expect("target URL"); + let validated_target = ValidatedUpstreamTarget::new( + &target_url, + vec![SocketAddr::from(([203, 0, 113, 10], 443))], + ) + .expect("validated target"); let pool_key = upstream_client_pool_key( Some("provider-1"), Some("endpoint-1"), Some("key-1"), Some(&profile), false, + validated_target, ); assert_eq!(pool_key.provider_id, "provider-1"); @@ -852,25 +890,20 @@ mod tests { assert!(!clients.contains_key(&key_b)); } - #[test] - fn http_proxy_authorization_header_uses_basic_auth_for_http_proxy() { - assert_eq!( - http_proxy_authorization_header(Some("http://user:pass@proxy.example:8080")).as_deref(), - Some("Basic dXNlcjpwYXNz") - ); - assert_eq!( - http_proxy_authorization_header(Some("socks5h://user:pass@127.0.0.1:1080")), - None - ); - } - fn test_pool_key(key_id: &str) -> UpstreamClientPoolKey { + let target_url = url::Url::parse("https://example.com/").expect("target URL"); + let validated_target = ValidatedUpstreamTarget::new( + &target_url, + vec![SocketAddr::from(([203, 0, 113, 10], 443))], + ) + .expect("validated target"); upstream_client_pool_key( Some("provider-1"), Some("endpoint-1"), Some(key_id), None, false, + validated_target, ) } @@ -892,10 +925,10 @@ mod tests { } #[tokio::test] - #[ignore = "requires loopback listener support"] - async fn upstream_client_sends_http_requests_through_http_proxy() { - let (proxy_url, request_rx) = spawn_http_proxy().await; - let client = proxied_client(&proxy_url); + async fn http_proxy_connects_to_pinned_ip_and_preserves_origin_host() { + let pinned_addr = SocketAddr::from(([203, 0, 113, 77], 80)); + let (proxy_url, connect_rx, request_rx) = spawn_http_proxy().await; + let client = proxied_client(&proxy_url, "http://example.com/", pinned_addr); let request = hyper::Request::builder() .method(hyper::Method::GET) .uri("http://example.com/tunnel-test") @@ -910,21 +943,66 @@ mod tests { .await .expect("body should collect") .to_bytes(); + let connect = connect_rx.await.expect("proxy should receive CONNECT"); let raw_request = request_rx.await.expect("proxy should receive request"); assert_eq!(status, hyper::StatusCode::OK); assert_eq!(&body[..], b"ok"); assert!( - raw_request.starts_with("GET http://example.com/tunnel-test HTTP/1.1\r\n"), + connect.starts_with("CONNECT 203.0.113.77:80 HTTP/1.1\r\n"), + "unexpected proxy CONNECT: {connect:?}" + ); + assert!( + raw_request.starts_with("GET /tunnel-test HTTP/1.1\r\n"), "unexpected proxy request: {raw_request:?}" ); + assert!( + raw_request + .to_ascii_lowercase() + .contains("\r\nhost: example.com\r\n"), + "original Host header should be preserved: {raw_request:?}" + ); } #[tokio::test] - #[ignore = "requires loopback listener support"] - async fn upstream_client_sends_http_requests_through_socks5h_proxy() { - let (proxy_url, target_rx, request_rx) = spawn_socks5h_proxy().await; - let client = proxied_client(&proxy_url); + async fn https_proxy_connects_to_pinned_ip_while_sni_uses_hostname() { + let pinned_addr = SocketAddr::from(([203, 0, 113, 78], 443)); + let (proxy_url, connect_rx) = spawn_connect_only_http_proxy().await; + let client = proxied_client(&proxy_url, "https://sni.example/", pinned_addr); + let request = hyper::Request::builder() + .method(hyper::Method::GET) + .uri("https://sni.example/secure") + .body(full_request_body(Bytes::new())) + .expect("request should build"); + + let _ = client.request(request).await; + let connect = connect_rx.await.expect("proxy should receive CONNECT"); + assert!( + connect.starts_with("CONNECT 203.0.113.78:443 HTTP/1.1\r\n"), + "unexpected proxy CONNECT: {connect:?}" + ); + let uri: Uri = "https://sni.example/secure".parse().expect("URI"); + match resolve_server_name(&uri).expect("server name") { + ServerName::DnsName(name) => assert_eq!(name.as_ref(), "sni.example"), + other => panic!("expected DNS SNI, got {other:?}"), + } + } + + #[tokio::test] + async fn socks5_proxy_connects_to_pinned_ip_and_preserves_origin_host() { + assert_socks_proxy_uses_pinned_ip("socks5").await; + } + + #[tokio::test] + async fn socks5h_proxy_connects_to_pinned_ip_and_preserves_origin_host() { + assert_socks_proxy_uses_pinned_ip("socks5h").await; + } + + async fn assert_socks_proxy_uses_pinned_ip(scheme: &str) { + let pinned_addr = SocketAddr::from(([203, 0, 113, 79], 80)); + let (proxy_addr, target_rx, request_rx) = spawn_socks5_proxy().await; + let proxy_url = format!("{scheme}://{proxy_addr}"); + let client = proxied_client(&proxy_url, "http://example.com/", pinned_addr); let request = hyper::Request::builder() .method(hyper::Method::GET) .uri("http://example.com/socks-test") @@ -944,14 +1022,24 @@ mod tests { .expect("SOCKS proxy should receive HTTP request"); assert_eq!(&body[..], b"ok"); - assert_eq!(target, ("example.com".to_string(), 80)); + assert_eq!(target, pinned_addr); assert!( raw_request.starts_with("GET /socks-test HTTP/1.1\r\n"), "unexpected SOCKS tunneled request: {raw_request:?}" ); + assert!( + raw_request + .to_ascii_lowercase() + .contains("\r\nhost: example.com\r\n"), + "original Host header should be preserved: {raw_request:?}" + ); } - fn proxied_client(proxy_url: &str) -> UpstreamClient { + fn proxied_client( + proxy_url: &str, + target_url: &str, + pinned_addr: SocketAddr, + ) -> UpstreamClient { let _ = rustls::crypto::ring::default_provider().install_default(); let config = Config::try_parse_from([ "aether-tunnel", @@ -967,23 +1055,32 @@ mod tests { "2", ]) .expect("config should parse"); - build_upstream_client_with_protocol( - &config, - Arc::new(DnsCache::new(Duration::from_secs(60), 16)), - true, - false, - ) - .expect("client should build") + let target_url = url::Url::parse(target_url).expect("target URL should parse"); + let validated_target = ValidatedUpstreamTarget::new(&target_url, vec![pinned_addr]) + .expect("target should validate"); + build_upstream_client_with_protocol(&config, validated_target, true, false) + .expect("client should build") } - async fn spawn_http_proxy() -> (String, tokio::sync::oneshot::Receiver) { + async fn spawn_http_proxy() -> ( + String, + tokio::sync::oneshot::Receiver, + tokio::sync::oneshot::Receiver, + ) { let listener = TcpListener::bind("127.0.0.1:0") .await .expect("listener should bind"); let addr = listener.local_addr().expect("local addr should exist"); + let (connect_tx, connect_rx) = tokio::sync::oneshot::channel(); let (request_tx, request_rx) = tokio::sync::oneshot::channel(); tokio::spawn(async move { let (mut stream, _) = listener.accept().await.expect("proxy should accept"); + let connect = read_http_headers(&mut stream).await; + let _ = connect_tx.send(connect); + stream + .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n") + .await + .expect("CONNECT response should write"); let request = read_http_headers(&mut stream).await; let _ = request_tx.send(request); stream @@ -991,12 +1088,30 @@ mod tests { .await .expect("proxy response should write"); }); - (format!("http://{addr}"), request_rx) + (format!("http://{addr}"), connect_rx, request_rx) } - async fn spawn_socks5h_proxy() -> ( - String, - tokio::sync::oneshot::Receiver<(String, u16)>, + async fn spawn_connect_only_http_proxy() -> (String, tokio::sync::oneshot::Receiver) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let addr = listener.local_addr().expect("local addr should exist"); + let (connect_tx, connect_rx) = tokio::sync::oneshot::channel(); + tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("proxy should accept"); + let connect = read_http_headers(&mut stream).await; + let _ = connect_tx.send(connect); + stream + .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n") + .await + .expect("CONNECT response should write"); + }); + (format!("http://{addr}"), connect_rx) + } + + async fn spawn_socks5_proxy() -> ( + SocketAddr, + tokio::sync::oneshot::Receiver, tokio::sync::oneshot::Receiver, ) { let listener = TcpListener::bind("127.0.0.1:0") @@ -1018,26 +1133,25 @@ mod tests { .await .expect("SOCKS method should write"); - let mut request_head = [0u8; 5]; + let mut request_head = [0u8; 4]; stream .read_exact(&mut request_head) .await .expect("SOCKS request head should read"); - assert_eq!(&request_head[..4], &[0x05, 0x01, 0x00, 0x03]); - let len = request_head[4] as usize; - let mut host = vec![0u8; len]; + assert_eq!(request_head, [0x05, 0x01, 0x00, 0x01]); + let mut ip = [0u8; 4]; stream - .read_exact(&mut host) + .read_exact(&mut ip) .await - .expect("SOCKS host should read"); + .expect("SOCKS IPv4 target should read"); let mut port = [0u8; 2]; stream .read_exact(&mut port) .await .expect("SOCKS port should read"); - let host = String::from_utf8(host).expect("SOCKS host should be UTF-8"); let port = u16::from_be_bytes(port); - let _ = target_tx.send((host, port)); + let target = SocketAddr::from((ip, port)); + let _ = target_tx.send(target); stream .write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0]) .await @@ -1050,7 +1164,7 @@ mod tests { .await .expect("SOCKS tunneled response should write"); }); - (format!("socks5h://{addr}"), target_rx, request_rx) + (addr, target_rx, request_rx) } async fn read_http_headers(stream: &mut TcpStream) -> String { diff --git a/crates/aether-admin/src/observability/monitoring.rs b/crates/aether-admin/src/observability/monitoring.rs index af0092894..4bd54fce4 100644 --- a/crates/aether-admin/src/observability/monitoring.rs +++ b/crates/aether-admin/src/observability/monitoring.rs @@ -1,5 +1,10 @@ use aether_data_contracts::repository::{ - candidates::{DecisionTrace, DecisionTraceCandidate, RequestCandidateStatus}, + candidates::{ + sanitize_request_candidate_api_formats, sanitize_request_candidate_error_type, + sanitize_request_candidate_extra_data, sanitize_request_candidate_required_capabilities, + sanitize_request_candidate_skip_reason, DecisionTrace, DecisionTraceCandidate, + RequestCandidateStatus, + }, provider_catalog::StoredProviderCatalogKey, usage::StoredRequestUsageAudit, }; @@ -9,7 +14,6 @@ use axum::{ response::{IntoResponse, Response}, Json, }; -use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine as _}; use serde_json::{json, Value}; use std::collections::BTreeMap; @@ -329,7 +333,19 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts( usage: Option<&StoredRequestUsageAudit>, key_accounts: &BTreeMap, ) -> Value { + let mut item = item.clone(); + item.sanitize_sensitive_diagnostics(); let candidate = &item.candidate; + let sanitized_extra_data = + build_admin_monitoring_trace_candidate_extra_data(candidate.extra_data.as_ref(), usage); + let sanitized_extra_data_ref = + (!sanitized_extra_data.is_null()).then_some(&sanitized_extra_data); + let sanitized_key_api_formats = + sanitize_request_candidate_api_formats(item.provider_key_api_formats.clone()); + let sanitized_key_capabilities = + sanitize_request_candidate_required_capabilities(item.provider_key_capabilities.clone()); + let sanitized_required_capabilities = + sanitize_request_candidate_required_capabilities(candidate.required_capabilities.clone()); let key_account = candidate .key_id .as_deref() @@ -349,32 +365,32 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts( "endpoint_name": item.endpoint_api_format, "endpoint_api_family": item.endpoint_api_family, "endpoint_kind": item.endpoint_kind, - "endpoint_format_acceptance_config": item.endpoint_format_acceptance_config, + "endpoint_format_acceptance_config": serde_json::Value::Null, "key_id": candidate.key_id, "key_name": item.provider_key_name, "key_account_label": key_account.and_then(|item| item.label.clone()), "key_preview": serde_json::Value::Null, "key_auth_type": item.provider_key_auth_type, - "key_api_formats": item.provider_key_api_formats, + "key_api_formats": sanitized_key_api_formats, "key_internal_priority": item.provider_key_internal_priority, - "key_global_priority_by_format": item.provider_key_global_priority_by_format, + "key_global_priority_by_format": serde_json::Value::Null, "key_oauth_plan_type": key_account.and_then(|item| item.oauth_plan_type.clone()), - "key_capabilities": item.provider_key_capabilities, - "required_capabilities": candidate.required_capabilities, + "key_capabilities": sanitized_key_capabilities, + "required_capabilities": sanitized_required_capabilities, "status": candidate.status, - "skip_reason": candidate.skip_reason, + "skip_reason": sanitize_request_candidate_skip_reason(candidate.skip_reason.clone()), "is_cached": candidate.is_cached, "status_code": candidate.status_code, - "error_type": candidate.error_type, - "error_message": candidate.error_message, + "error_type": sanitize_request_candidate_error_type(candidate.error_type.clone()), + "error_message": serde_json::Value::Null, "latency_ms": candidate.latency_ms, "concurrent_requests": candidate.concurrent_requests, - "ranking": build_admin_monitoring_trace_candidate_ranking(candidate.extra_data.as_ref()), - "image_progress": candidate.extra_data.as_ref() + "ranking": build_admin_monitoring_trace_candidate_ranking(sanitized_extra_data_ref), + "image_progress": sanitized_extra_data_ref .and_then(|value| value.get("image_progress")) .cloned() .unwrap_or(Value::Null), - "extra_data": build_admin_monitoring_trace_candidate_extra_data(candidate.extra_data.as_ref(), usage), + "extra_data": sanitized_extra_data, "created_at": unix_ms_to_rfc3339(candidate.created_at_unix_ms), "started_at": candidate.started_at_unix_ms.and_then(unix_ms_to_rfc3339), "finished_at": candidate.finished_at_unix_ms.and_then(unix_ms_to_rfc3339), @@ -502,11 +518,8 @@ fn build_admin_monitoring_trace_candidate_extra_data( existing: Option<&Value>, usage: Option<&StoredRequestUsageAudit>, ) -> Value { - let mut extra_data = match existing { - Some(Value::Object(object)) => Some(object.clone()), - Some(other) => return other.clone(), - None => None, - }; + let mut extra_data = sanitize_request_candidate_extra_data(existing.cloned()) + .and_then(|value| value.as_object().cloned()); if let Some(usage) = usage { let extra_object = extra_data.get_or_insert_with(serde_json::Map::new); @@ -534,9 +547,6 @@ fn build_admin_monitoring_trace_candidate_extra_data( if let Some(response) = admin_monitoring_trace_response_data( "upstream_response", usage.status_code, - usage.response_headers.as_ref(), - usage.response_body.as_ref(), - usage.response_body_ref.as_deref(), usage.response_body_state, ) { merge_admin_monitoring_trace_response(extra_object, "upstream_response", response); @@ -563,92 +573,25 @@ fn build_admin_monitoring_trace_candidate_extra_data( } } - match extra_data { - Some(object) => Value::Object(object), - None => Value::Null, - } + sanitize_request_candidate_extra_data(extra_data.map(Value::Object)).unwrap_or(Value::Null) } fn admin_monitoring_trace_response_data( source: &str, status_code: Option, - headers: Option<&Value>, - body: Option<&Value>, - body_ref: Option<&str>, body_state: Option, ) -> Option { - if status_code.is_none() - && headers.is_none() - && body.is_none() - && body_ref.is_none() - && body_state.is_none() - { + if status_code.is_none() && body_state.is_none() { return None; } - let body = admin_monitoring_trace_response_body(headers, body); Some(json!({ "source": source, "status_code": status_code, - "headers": headers.cloned().unwrap_or(Value::Null), - "body": body.unwrap_or(Value::Null), - "body_ref": body_ref, "body_state": body_state.map(|state| state.as_str()), })) } -fn admin_monitoring_trace_response_body( - headers: Option<&Value>, - body: Option<&Value>, -) -> Option { - let body = body?; - admin_monitoring_decode_connect_json_error_body(headers, body).or_else(|| Some(body.clone())) -} - -fn admin_monitoring_decode_connect_json_error_body( - headers: Option<&Value>, - body: &Value, -) -> Option { - if !admin_monitoring_headers_indicate_connect_json(headers) { - return None; - } - - let body_base64 = match body { - Value::String(value) => Some(value.as_str()), - Value::Object(object) => object - .get("encoding") - .and_then(Value::as_str) - .is_some_and(|value| value.eq_ignore_ascii_case("base64")) - .then(|| object.get("data").and_then(Value::as_str)) - .flatten(), - _ => None, - }? - .trim(); - if body_base64.is_empty() { - return None; - } - - let body_bytes = BASE64_STANDARD.decode(body_base64).ok()?; - aether_ai_formats::api::extract_provider_private_stream_error_body(None, &body_bytes) -} - -fn admin_monitoring_headers_indicate_connect_json(headers: Option<&Value>) -> bool { - headers - .and_then(Value::as_object) - .and_then(|object| { - object.iter().find_map(|(key, value)| { - key.eq_ignore_ascii_case("content-type") - .then(|| value.as_str()) - .flatten() - }) - }) - .map(str::trim) - .is_some_and(|value| { - let value = value.to_ascii_lowercase(); - value.contains("application/connect+json") || value.contains("+connect+json") - }) -} - fn merge_admin_monitoring_trace_response( extra_object: &mut serde_json::Map, key: &str, @@ -707,27 +650,25 @@ fn admin_monitoring_trace_request_path_and_query( fn admin_monitoring_usage_request_path(usage: &StoredRequestUsageAudit) -> Option { admin_monitoring_usage_metadata_string(usage, "request_path") + .and_then(|value| aether_ai_formats::api::sanitize_request_path(&value)) } fn admin_monitoring_usage_request_query_string(usage: &StoredRequestUsageAudit) -> Option { admin_monitoring_usage_metadata_string(usage, "request_query_string") - .map(|value| value.trim_start_matches('?').to_string()) - .filter(|value| !value.is_empty()) + .and_then(|value| aether_ai_formats::api::sanitize_request_query_string(&value)) } fn admin_monitoring_usage_request_path_and_query( usage: &StoredRequestUsageAudit, ) -> Option { - admin_monitoring_usage_metadata_string(usage, "request_path_and_query").or_else(|| { - let path = admin_monitoring_usage_metadata_string(usage, "request_path")?; - let query = admin_monitoring_usage_metadata_string(usage, "request_query_string") - .map(|value| value.trim_start_matches('?').to_string()) - .filter(|value| !value.is_empty()); - Some(match query { - Some(query) if !path.contains('?') => format!("{path}?{query}"), - _ => path, + admin_monitoring_usage_metadata_string(usage, "request_path_and_query") + .and_then(|value| aether_ai_formats::api::sanitize_request_path_and_query(&value, None)) + .or_else(|| { + let path = admin_monitoring_usage_metadata_string(usage, "request_path")?; + let query = admin_monitoring_usage_metadata_string(usage, "request_query_string") + .and_then(|value| aether_ai_formats::api::sanitize_request_query_string(&value)); + aether_ai_formats::api::sanitize_request_path_and_query(&path, query.as_deref()) }) - }) } fn admin_monitoring_usage_metadata_string( diff --git a/crates/aether-admin/src/observability/usage.rs b/crates/aether-admin/src/observability/usage.rs index 7215a08b4..85fbf540f 100644 --- a/crates/aether-admin/src/observability/usage.rs +++ b/crates/aether-admin/src/observability/usage.rs @@ -24,6 +24,29 @@ use url::form_urlencoded; pub const ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL: &str = "Admin usage data unavailable"; +fn admin_usage_safe_error_message(item: &StoredRequestUsageAudit) -> Option { + if item + .error_message + .as_deref() + .is_none_or(|value| value.trim().is_empty()) + && item.status_code.is_none_or(|status| status < 400) + { + return None; + } + + item.error_category + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| { + item.status_code + .filter(|status| *status >= 400) + .map(|status| format!("http_{status}")) + }) + .or_else(|| Some("request_failed".to_string())) +} + pub fn admin_usage_data_unavailable_response(detail: &'static str) -> Response { ( http::StatusCode::SERVICE_UNAVAILABLE, @@ -339,6 +362,77 @@ fn admin_usage_strip_trace_metadata(metadata: &mut serde_json::Map Value { + admin_usage_safe_metadata_value_inner(value, false) +} + +fn admin_usage_safe_metadata_value_inner(value: &Value, in_realtime_session: bool) -> Value { + match value { + Value::Object(object) => Value::Object( + object + .iter() + .filter_map(|(key, value)| { + let compact = key + .chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .map(|ch| ch.to_ascii_lowercase()) + .collect::(); + // The generic `token` key filter protects historical rows that may + // predate the persistence projection. Realtime audio usage counters + // are explicitly safe numeric metrics, but only when they occur in the + // structured realtime-session object. Do not broaden this exception to + // arbitrary token-shaped fields or string values. + let safe_realtime_counter = in_realtime_session + && matches!(compact.as_str(), "inputaudiotokens" | "outputaudiotokens") + && value.as_u64().is_some(); + if (compact.contains("authorization") + || compact.contains("credential") + || compact.contains("secret") + || compact.contains("token") + || compact.contains("apikey") + || compact.contains("password") + || matches!(compact.as_str(), "cookie" | "setcookie" | "headers")) + && !safe_realtime_counter + { + return None; + } + let child_in_realtime_session = compact == "realtimesession"; + Some(( + key.clone(), + admin_usage_safe_metadata_value_inner(value, child_in_realtime_session), + )) + }) + .collect(), + ), + Value::Array(items) => Value::Array( + items + .iter() + .map(|item| admin_usage_safe_metadata_value_inner(item, false)) + .collect(), + ), + Value::String(value) => { + if let Ok(mut url) = url::Url::parse(value.trim()) { + if !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + { + let _ = url.set_username(""); + let _ = url.set_password(None); + url.set_query(None); + url.set_fragment(None); + return Value::String(url.to_string()); + } + } + Value::String(value.clone()) + } + other => other.clone(), + } +} + fn admin_usage_string_field<'a>(value: &'a Value, field: &str) -> Option<&'a str> { value .get(field) @@ -1219,7 +1313,7 @@ fn admin_usage_active_request_json( "updated_at": unix_secs_to_rfc3339(item.updated_at_unix_secs), "response_time_updated_at": admin_usage_response_time_updated_at(item), "status_code": item.status_code, - "error_message": item.error_message, + "error_message": admin_usage_safe_error_message(item), "provider": item.provider_name, "api_key_name": api_key_name, "provider_key_name": provider_key_name, @@ -1334,7 +1428,7 @@ pub fn admin_usage_record_json( "cache_creation_price_per_1m": cache_creation_price_per_1m, "cache_read_price_per_1m": cache_read_price_per_1m, "status_code": item.status_code, - "error_message": item.error_message, + "error_message": admin_usage_safe_error_message(item), "status": item.status, "request_type": item.request_type, "has_fallback": admin_usage_has_fallback(item), @@ -2464,7 +2558,11 @@ pub fn build_admin_usage_detail_payload( auth_api_key_reader_available, provider_key_name, ); - let mut metadata = match item.request_metadata.clone() { + let mut metadata = match item + .request_metadata + .as_ref() + .map(admin_usage_safe_metadata_value) + { Some(Value::Object(object)) => Value::Object(object), Some(value) => json!({ "request_metadata": value }), None => json!({}), @@ -2626,7 +2724,8 @@ mod tests { admin_usage_has_fallback, admin_usage_is_failed, admin_usage_is_success, admin_usage_matches_search, admin_usage_matches_status, admin_usage_matches_username, admin_usage_record_json, admin_usage_resolve_request_capture_body, - admin_usage_total_tokens, admin_usage_upstream_is_stream, build_admin_usage_detail_payload, + admin_usage_safe_metadata_value, admin_usage_total_tokens, admin_usage_upstream_is_stream, + build_admin_usage_detail_payload, }; use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageBodyField}; @@ -2688,6 +2787,35 @@ mod tests { assert!(admin_usage_matches_status(&item, Some("completed"))); } + #[test] + fn admin_usage_summaries_do_not_return_historical_raw_error_text() { + let item = StoredRequestUsageAudit { + error_category: Some("authentication_error".to_string()), + ..sample_usage( + "failed", + Some(401), + Some( + "upstream said Authorization: Bearer live-secret at https://api.example?key=secret", + ), + ) + }; + + 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); + + assert_eq!(record["error_message"], "authentication_error"); + assert_eq!(active["error_message"], "authentication_error"); + assert!(!record.to_string().contains("live-secret")); + assert!(!active.to_string().contains("live-secret")); + } + #[test] fn client_requested_stream_prefers_request_metadata_flag() { let item = StoredRequestUsageAudit { @@ -2842,6 +2970,35 @@ mod tests { assert_eq!(payload["metadata"]["realtime_session"], realtime_session); } + #[test] + fn admin_usage_safe_metadata_keeps_only_numeric_realtime_audio_counters() { + let projected = admin_usage_safe_metadata_value(&json!({ + "input_audio_tokens": 99, + "refresh_token": "should-be-removed", + "realtime_session": { + "input_audio_tokens": 7, + "output_audio_tokens": 3, + "input_audio_tokens_text": "should-be-removed", + "refresh_token": 11, + "nested": { + "output_audio_tokens": 5 + } + } + })); + + assert!(projected.get("input_audio_tokens").is_none()); + assert!(projected.get("refresh_token").is_none()); + assert_eq!(projected["realtime_session"]["input_audio_tokens"], 7); + assert_eq!(projected["realtime_session"]["output_audio_tokens"], 3); + assert!(projected["realtime_session"] + .get("input_audio_tokens_text") + .is_none()); + assert!(projected["realtime_session"].get("refresh_token").is_none()); + assert!(projected["realtime_session"]["nested"] + .get("output_audio_tokens") + .is_none()); + } + #[test] fn admin_usage_record_infers_client_family_from_user_agent() { let item = StoredRequestUsageAudit { @@ -3313,6 +3470,46 @@ mod tests { assert_eq!(payload["has_client_response_body"], true); } + #[test] + fn detail_payload_sanitizes_historical_request_metadata() { + let item = StoredRequestUsageAudit { + request_metadata: Some(json!({ + "trace_id": "internal-trace", + "authorization": "Bearer live-secret", + "nested": { + "refresh_token": "refresh-secret", + "endpoint_url": "https://user:password@api.example.test/v1?key=secret#fragment", + "safe_count": 2 + } + })), + ..sample_usage("failed", Some(500), Some("failed")) + }; + + let payload = build_admin_usage_detail_payload( + &item, + &BTreeMap::new(), + &BTreeMap::new(), + false, + false, + None, + false, + None, + &BTreeMap::new(), + ); + let metadata = &payload["metadata"]; + assert!(metadata.get("authorization").is_none()); + assert!(metadata["nested"].get("refresh_token").is_none()); + assert_eq!( + metadata["nested"]["endpoint_url"], + "https://api.example.test/v1" + ); + assert_eq!(metadata["nested"]["safe_count"], 2); + let encoded = metadata.to_string(); + for secret in ["live-secret", "refresh-secret", "password", "key=secret"] { + assert!(!encoded.contains(secret), "leaked {secret}"); + } + } + #[test] fn detail_payload_separates_upstream_client_and_summary_errors() { let item = StoredRequestUsageAudit { @@ -3505,14 +3702,8 @@ mod tests { &BTreeMap::new(), ); - assert_eq!( - payload["client_error"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); - assert_eq!( - payload["failure_summary"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); + assert!(payload["client_error"]["message"].is_null()); + assert!(payload["failure_summary"].is_null()); assert_eq!( payload["scheduling_failure"]["title"], "本地调度失败:没有可调度候选" @@ -3522,10 +3713,7 @@ mod tests { "candidate_list_empty" ); assert!(payload["scheduling_failure"]["reason_summary"].is_null()); - assert_eq!( - payload["scheduling_failure"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); + assert!(payload["scheduling_failure"]["message"].is_null()); assert_eq!(payload["scheduling_failure"]["no_upstream_attempt"], true); } @@ -3562,14 +3750,8 @@ mod tests { &BTreeMap::new(), ); - assert_eq!( - payload["client_error"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); - assert_eq!( - payload["failure_summary"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); + assert!(payload["client_error"]["message"].is_null()); + assert!(payload["failure_summary"].is_null()); assert_eq!( payload["scheduling_failure"]["title"], "本地调度失败:所有候选均被跳过" @@ -3578,14 +3760,8 @@ mod tests { payload["scheduling_failure"]["reason"], "all_candidates_skipped" ); - assert_eq!( - payload["scheduling_failure"]["reason_summary"], - "provider_quota_blocked 2 次" - ); - assert_eq!( - payload["scheduling_failure"]["message"], - "没有可用提供商支持模型 gpt-5.4 的流式请求" - ); + assert!(payload["scheduling_failure"]["reason_summary"].is_null()); + assert!(payload["scheduling_failure"]["message"].is_null()); assert_eq!(payload["scheduling_failure"]["no_upstream_attempt"], true); } diff --git a/crates/aether-admin/src/provider/endpoints.rs b/crates/aether-admin/src/provider/endpoints.rs index 539600704..26d7643c0 100644 --- a/crates/aether-admin/src/provider/endpoints.rs +++ b/crates/aether-admin/src/provider/endpoints.rs @@ -6,6 +6,13 @@ use chrono::{TimeZone, Utc}; use serde_json::{json, Value}; use std::collections::BTreeMap; +use super::redaction::{ + admin_restore_secret_safe_body_rules, admin_restore_secret_safe_header_rules, + admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, admin_restore_secret_safe_url, + admin_secret_safe_body_rules, admin_secret_safe_header_rules, admin_secret_safe_json, + admin_secret_safe_proxy, admin_secret_safe_url, +}; + pub fn normalize_endpoint_api_format(api_format: &str) -> String { aether_ai_formats::normalize_api_format_alias(api_format) } @@ -213,21 +220,6 @@ mod endpoint_key_count_tests { } } -fn masked_proxy_value(proxy: Option<&serde_json::Value>) -> serde_json::Value { - let Some(proxy) = proxy.and_then(serde_json::Value::as_object) else { - return serde_json::Value::Null; - }; - let mut masked = proxy.clone(); - if masked - .get("password") - .and_then(serde_json::Value::as_str) - .is_some_and(|value| !value.trim().is_empty()) - { - masked.insert("password".to_string(), json!("***")); - } - serde_json::Value::Object(masked) -} - fn endpoint_timestamp_or_now(value: Option, now_unix_secs: u64) -> serde_json::Value { unix_secs_to_rfc3339(value.unwrap_or(now_unix_secs)) .map(serde_json::Value::String) @@ -246,15 +238,15 @@ pub fn build_admin_provider_endpoint_response( "provider_id": endpoint.provider_id, "provider_name": provider_name, "api_format": endpoint.api_format, - "base_url": endpoint.base_url, + "base_url": admin_secret_safe_url(Some(&endpoint.base_url)), "custom_path": endpoint.custom_path, - "header_rules": endpoint.header_rules, - "body_rules": endpoint.body_rules, + "header_rules": admin_secret_safe_header_rules(endpoint.header_rules.as_ref()), + "body_rules": admin_secret_safe_body_rules(endpoint.body_rules.as_ref()), "max_retries": endpoint.max_retries.unwrap_or(2), "is_active": endpoint.is_active, - "config": endpoint.config, - "proxy": masked_proxy_value(endpoint.proxy.as_ref()), - "format_acceptance_config": endpoint.format_acceptance_config, + "config": admin_secret_safe_json(endpoint.config.as_ref()), + "proxy": admin_secret_safe_proxy(endpoint.proxy.as_ref()), + "format_acceptance_config": admin_secret_safe_json(endpoint.format_acceptance_config.as_ref()), "total_keys": total_keys, "active_keys": active_keys, "created_at": endpoint_timestamp_or_now(endpoint.created_at_unix_ms, now_unix_secs), @@ -342,7 +334,8 @@ where "base_url 必须是字符串".to_string() }); }; - updated.base_url = base_url.to_string(); + updated.base_url = + admin_restore_secret_safe_url(Some(&existing_endpoint.base_url), base_url); } if contains_field("custom_path") { @@ -359,7 +352,10 @@ where if !header_rules.is_array() { return Err("header_rules 必须是数组或 null".to_string()); } - Some(header_rules.clone()) + Some(admin_restore_secret_safe_header_rules( + existing_endpoint.header_rules.as_ref(), + header_rules, + )) }; } @@ -373,7 +369,10 @@ where if !body_rules.is_array() { return Err("body_rules 必须是数组或 null".to_string()); } - Some(body_rules.clone()) + Some(admin_restore_secret_safe_body_rules( + existing_endpoint.body_rules.as_ref(), + body_rules, + )) }; } @@ -408,7 +407,10 @@ where if !config.is_object() { return Err("config 必须是对象或 null".to_string()); } - Some(config.clone()) + Some(admin_restore_secret_safe_json( + existing_endpoint.config.as_ref(), + config, + )) }; } @@ -416,26 +418,14 @@ where if is_null_field("proxy") { updated.proxy = None; } else { - let Some(mut proxy) = payload - .proxy - .clone() - .and_then(|value| value.as_object().cloned()) - else { + let Some(proxy) = payload.proxy.as_ref().and_then(Value::as_object) else { return Err("proxy 必须是对象或 null".to_string()); }; - if !proxy.contains_key("password") { - if let Some(old_password) = existing_endpoint - .proxy - .as_ref() - .and_then(Value::as_object) - .and_then(|proxy| proxy.get("password")) - .and_then(Value::as_str) - .filter(|value| !value.is_empty()) - { - proxy.insert("password".to_string(), json!(old_password)); - } - } - updated.proxy = Some(Value::Object(proxy)); + let restored = admin_restore_secret_safe_proxy( + existing_endpoint.proxy.as_ref(), + &Value::Object(proxy.clone()), + ); + updated.proxy = Some(restored); } } @@ -449,7 +439,10 @@ where if !config.is_object() { return Err("format_acceptance_config 必须是对象或 null".to_string()); } - Some(config.clone()) + Some(admin_restore_secret_safe_json( + existing_endpoint.format_acceptance_config.as_ref(), + config, + )) }; } diff --git a/crates/aether-admin/src/provider/mod.rs b/crates/aether-admin/src/provider/mod.rs index b41442da5..2a4333585 100644 --- a/crates/aether-admin/src/provider/mod.rs +++ b/crates/aether-admin/src/provider/mod.rs @@ -5,5 +5,6 @@ pub mod oauth; pub mod ops; pub mod pool; pub mod quota; +pub mod redaction; pub mod state; pub mod status; diff --git a/crates/aether-admin/src/provider/ops/actions.rs b/crates/aether-admin/src/provider/ops/actions.rs index 228bf7a3a..f6000b491 100644 --- a/crates/aether-admin/src/provider/ops/actions.rs +++ b/crates/aether-admin/src/provider/ops/actions.rs @@ -53,11 +53,7 @@ pub fn parse_sub2api_balance_payload( return Err("响应格式无效".to_string()); }; if me_payload.get("code").and_then(Value::as_i64).unwrap_or(-1) != 0 { - return Err(me_payload - .get("message") - .and_then(Value::as_str) - .unwrap_or("查询用户信息失败") - .to_string()); + return Err("查询用户信息失败".to_string()); } let Some(me_data) = me_payload.get("data").and_then(Value::as_object) else { return Err("响应格式无效".to_string()); @@ -126,11 +122,12 @@ pub fn attach_balance_checkin_outcome( .entry("extra".to_string()) .or_insert_with(|| Value::Object(Map::new())); if let Some(extra) = extra.as_object_mut() { + let message = stable_checkin_outcome_message(outcome); if outcome.cookie_expired { extra.insert("cookie_expired".to_string(), Value::Bool(true)); extra.insert( "cookie_expired_message".to_string(), - Value::String(outcome.message.clone()), + Value::String(message.to_string()), ); } else { extra.insert( @@ -139,7 +136,7 @@ pub fn attach_balance_checkin_outcome( ); extra.insert( "checkin_message".to_string(), - Value::String(outcome.message.clone()), + Value::String(message.to_string()), ); } } @@ -151,6 +148,17 @@ pub fn attach_balance_checkin_outcome( } } +fn stable_checkin_outcome_message(outcome: &ProviderOpsCheckinOutcome) -> &'static str { + if outcome.cookie_expired { + return "Cookie 已失效"; + } + match outcome.success { + Some(true) => "签到成功", + Some(false) => "签到失败", + None => "今日已签到", + } +} + pub fn build_balance_data( total_granted: Option, total_used: Option, @@ -177,11 +185,7 @@ fn parse_new_api_balance_payload( { response_json.get("data") } else if response_json.get("success").and_then(Value::as_bool) == Some(false) { - return Err(response_json - .get("message") - .and_then(Value::as_str) - .unwrap_or("业务状态码表示失败") - .to_string()); + return Err("业务状态码表示失败".to_string()); } else { Some(response_json) }; @@ -265,11 +269,7 @@ fn parse_cubence_balance_payload( { response_json.get("data") } else if response_json.get("success").and_then(Value::as_bool) == Some(false) { - return Err(response_json - .get("message") - .and_then(Value::as_str) - .unwrap_or("查询余额失败") - .to_string()); + return Err("查询余额失败".to_string()); } else { Some(response_json) }; @@ -694,4 +694,47 @@ mod tests { assert_eq!(payload["status"], json!("auth_expired")); assert_eq!(payload["data"]["extra"]["cookie_expired"], json!(true)); } + + #[test] + fn attach_balance_checkin_outcome_does_not_copy_upstream_message() { + let mut payload = json!({ + "status": "success", + "data": { "extra": {} } + }); + attach_balance_checkin_outcome( + &mut payload, + &ProviderOpsCheckinOutcome { + success: Some(true), + message: "authorization=Bearer upstream-secret".to_string(), + cookie_expired: false, + }, + ); + + assert_eq!( + payload["data"]["extra"]["checkin_message"], + json!("签到成功") + ); + assert!(!payload.to_string().contains("upstream-secret")); + } + + #[test] + fn balance_parsers_do_not_return_upstream_error_messages() { + let config = json!({}).as_object().cloned().expect("config"); + let secret = "authorization=Bearer upstream-secret"; + + let generic_error = parse_query_balance_payload( + "generic_api", + &config, + &json!({"success": false, "message": secret}), + ) + .expect_err("generic API failure should be rejected"); + let sub2api_error = + parse_sub2api_balance_payload(&config, &json!({"code": 401, "message": secret}), None) + .expect_err("Sub2API failure should be rejected"); + + assert_eq!(generic_error, "业务状态码表示失败"); + assert_eq!(sub2api_error, "查询用户信息失败"); + assert!(!generic_error.contains("upstream-secret")); + assert!(!sub2api_error.contains("upstream-secret")); + } } diff --git a/crates/aether-admin/src/provider/ops/verify.rs b/crates/aether-admin/src/provider/ops/verify.rs index bc4d531ca..616656513 100644 --- a/crates/aether-admin/src/provider/ops/verify.rs +++ b/crates/aether-admin/src/provider/ops/verify.rs @@ -5,6 +5,7 @@ use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; use serde_json::{json, Map, Value}; const ADMIN_PROVIDER_OPS_ANYROUTER_XOR_KEY: &str = "3000176000856006061501533003690027800375"; +const ADMIN_PROVIDER_OPS_ANYROUTER_SESSION_PART_MAX_BYTES: usize = 256 * 1024; const ADMIN_PROVIDER_OPS_ANYROUTER_UNSBOX_TABLE: [usize; 40] = [ 0xF, 0x23, 0x1D, 0x18, 0x21, 0x10, 0x1, 0x26, 0xA, 0x9, 0x13, 0x1F, 0x28, 0x1B, 0x16, 0x17, 0x19, 0xD, 0x6, 0xB, 0x27, 0x12, 0x14, 0x8, 0xE, 0x15, 0x20, 0x1A, 0x2, 0x1E, 0x7, 0x4, 0x11, @@ -186,7 +187,16 @@ pub fn admin_provider_ops_anyrouter_parse_session_user_id(cookie_input: &str) -> } fn decode_python_urlsafe_b64(input: &str) -> Option> { - let normalized = input.trim().replace('-', "+").replace('_', "/"); + let input = input.trim(); + let max_encoded_len = ADMIN_PROVIDER_OPS_ANYROUTER_SESSION_PART_MAX_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4); + if input.is_empty() || input.len() > max_encoded_len { + return None; + } + let normalized = input.replace('-', "+").replace('_', "/"); if normalized.is_empty() { return None; } @@ -195,7 +205,8 @@ fn decode_python_urlsafe_b64(input: &str) -> Option> { if remainder != 0 { padded.push_str(&"=".repeat(4 - remainder)); } - STANDARD.decode(padded.as_bytes()).ok() + let decoded = STANDARD.decode(padded.as_bytes()).ok()?; + (decoded.len() <= ADMIN_PROVIDER_OPS_ANYROUTER_SESSION_PART_MAX_BYTES).then_some(decoded) } pub fn admin_provider_ops_verify_failure(message: impl Into) -> Value { @@ -618,12 +629,7 @@ fn verify_payload_with_auth_messages( { response_json.get("data") } else if response_json.get("success").and_then(Value::as_bool) == Some(false) { - return admin_provider_ops_verify_failure( - response_json - .get("message") - .and_then(Value::as_str) - .unwrap_or("验证失败"), - ); + return admin_provider_ops_verify_failure("验证失败"); } else { Some(response_json) }; @@ -632,17 +638,6 @@ fn verify_payload_with_auth_messages( return admin_provider_ops_verify_failure("响应格式无效"); }; - let mut extra = Map::new(); - for (key, value) in user_data { - if matches!( - key.as_str(), - "username" | "display_name" | "email" | "quota" | "used_quota" | "request_count" - ) { - continue; - } - extra.insert(key.clone(), value.clone()); - } - admin_provider_ops_verify_success( admin_provider_ops_verify_user_payload_with_usage( user_data @@ -660,7 +655,7 @@ fn verify_payload_with_auth_messages( admin_provider_ops_value_as_f64(user_data.get("quota")), admin_provider_ops_value_as_f64(user_data.get("used_quota")), admin_provider_ops_value_as_u64(user_data.get("request_count")), - Some(extra), + None, ), None, ) @@ -685,12 +680,7 @@ pub fn admin_provider_ops_cubence_verify_payload( { response_json.get("data") } else if response_json.get("success").and_then(Value::as_bool) == Some(false) { - return admin_provider_ops_verify_failure( - response_json - .get("message") - .and_then(Value::as_str) - .unwrap_or("验证失败"), - ); + return admin_provider_ops_verify_failure("验证失败"); } else { Some(response_json) }; @@ -864,12 +854,7 @@ pub fn admin_provider_ops_sub2api_verify_payload( return admin_provider_ops_verify_failure("响应格式无效"); }; if payload.get("code").and_then(Value::as_i64).unwrap_or(-1) != 0 { - return admin_provider_ops_verify_failure( - payload - .get("message") - .and_then(Value::as_str) - .unwrap_or("验证失败"), - ); + return admin_provider_ops_verify_failure("验证失败"); } let Some(user_data) = payload.get("data").and_then(Value::as_object) else { @@ -915,9 +900,10 @@ mod tests { admin_provider_ops_anyrouter_compute_acw_sc_v2, admin_provider_ops_anyrouter_parse_session_user_id, admin_provider_ops_anyrouter_verify_payload, admin_provider_ops_cubence_verify_payload, - admin_provider_ops_frontend_updated_credentials, admin_provider_ops_sub2api_verify_payload, - admin_provider_ops_usage_api_verify_payload, admin_provider_ops_verify_headers, - parse_verify_payload, ADMIN_PROVIDER_OPS_USER_AGENT, + admin_provider_ops_frontend_updated_credentials, admin_provider_ops_generic_verify_payload, + admin_provider_ops_sub2api_verify_payload, admin_provider_ops_usage_api_verify_payload, + admin_provider_ops_verify_headers, parse_verify_payload, + ADMIN_PROVIDER_OPS_ANYROUTER_SESSION_PART_MAX_BYTES, ADMIN_PROVIDER_OPS_USER_AGENT, }; use http::StatusCode; use reqwest::header::COOKIE; @@ -951,6 +937,17 @@ mod tests { assert_eq!(actual.as_deref(), Some("42")); } + #[test] + fn anyrouter_parse_session_user_id_rejects_oversized_encoded_cookie() { + let encoded_limit = ADMIN_PROVIDER_OPS_ANYROUTER_SESSION_PART_MAX_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap() + .saturating_mul(4); + let cookie = format!("session={}", "A".repeat(encoded_limit + 1)); + assert!(admin_provider_ops_anyrouter_parse_session_user_id(&cookie).is_none()); + } + #[test] fn frontend_updated_credentials_omits_internal_runtime_fields() { let filtered = admin_provider_ops_frontend_updated_credentials(Map::from_iter([ @@ -1102,6 +1099,58 @@ mod tests { assert_eq!(auth_failed["message"], json!("Cookie 已失效,请重新配置")); } + #[test] + fn generic_verify_payload_does_not_reflect_upstream_errors_or_unknown_fields() { + let failed = admin_provider_ops_generic_verify_payload( + StatusCode::OK, + &json!({ + "success": false, + "message": "token=upstream-secret https://user:pass@example.com/path?q=secret" + }), + ); + assert_eq!(failed["success"], json!(false)); + assert_eq!(failed["message"], json!("验证失败")); + + let succeeded = admin_provider_ops_generic_verify_payload( + StatusCode::OK, + &json!({ + "username": "alice", + "quota": 12.5, + "access_token": "upstream-secret", + "profile": {"private_note": "sensitive"} + }), + ); + assert_eq!(succeeded["success"], json!(true)); + assert_eq!(succeeded["data"]["username"], json!("alice")); + assert_eq!(succeeded["data"]["quota"], json!(12.5)); + assert_eq!(succeeded["data"]["extra"], json!({})); + let serialized = succeeded.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("private_note")); + } + + #[test] + fn architecture_verify_failures_do_not_reflect_upstream_messages() { + let cubence = admin_provider_ops_cubence_verify_payload( + StatusCode::OK, + &json!({ + "success": false, + "message": "authorization=Bearer upstream-secret" + }), + ); + assert_eq!(cubence["message"], json!("验证失败")); + + let sub2api = admin_provider_ops_sub2api_verify_payload( + StatusCode::OK, + &json!({ + "code": 500, + "message": "cookie=session-secret" + }), + None, + ); + assert_eq!(sub2api["message"], json!("验证失败")); + } + #[test] fn done_hub_verify_payload_reads_wrapped_profile() { let payload = parse_verify_payload( diff --git a/crates/aether-admin/src/provider/pool.rs b/crates/aether-admin/src/provider/pool.rs index 97021e99a..18f18d2c7 100644 --- a/crates/aether-admin/src/provider/pool.rs +++ b/crates/aether-admin/src/provider/pool.rs @@ -6,6 +6,9 @@ use chrono::{TimeZone, Utc}; use serde_json::{json, Map, Value}; use std::collections::{BTreeMap, BTreeSet}; +use super::redaction::{ + admin_provider_status_snapshot_safe_json, admin_secret_safe_json, admin_secret_safe_proxy, +}; use super::status as provider_status; #[derive(Debug, Default, Clone, serde::Deserialize)] @@ -18,7 +21,7 @@ pub struct AdminPoolResolveSelectionRequest { pub quick_selectors: Vec, } -#[derive(Debug, Default, Clone, serde::Deserialize)] +#[derive(Default, Clone, serde::Deserialize)] pub struct AdminPoolBatchActionRequest { #[serde(default)] pub key_ids: Vec, @@ -28,6 +31,17 @@ pub struct AdminPoolBatchActionRequest { pub payload: Option, } +impl std::fmt::Debug for AdminPoolBatchActionRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPoolBatchActionRequest") + .field("key_ids", &self.key_ids) + .field("action", &self.action) + .field("payload", &self.payload.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AdminPoolBatchActionKind { Enable, @@ -39,7 +53,7 @@ pub enum AdminPoolBatchActionKind { Delete, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct AdminPoolBatchActionPlan { pub key_ids: Vec, pub action: AdminPoolBatchActionKind, @@ -48,6 +62,25 @@ pub struct AdminPoolBatchActionPlan { pub settings_payload: Option, } +impl std::fmt::Debug for AdminPoolBatchActionPlan { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPoolBatchActionPlan") + .field("key_ids", &self.key_ids) + .field("action", &self.action) + .field("action_label", &self.action_label) + .field( + "proxy_payload", + &self.proxy_payload.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "settings_payload", + &self.settings_payload.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + #[derive(Debug, Clone, Default)] pub struct AdminPoolKeyPayloadContext { pub cooldown_reason: Option, @@ -58,7 +91,7 @@ pub struct AdminPoolKeyPayloadContext { pub cost_limit: Option, } -#[derive(Debug, Default, Clone, serde::Deserialize)] +#[derive(Default, Clone, serde::Deserialize)] pub struct AdminPoolBatchImportRequest { #[serde(default)] pub keys: Vec, @@ -70,7 +103,19 @@ pub struct AdminPoolBatchImportRequest { pub settings: Option, } -#[derive(Debug, Default, Clone, serde::Deserialize)] +impl std::fmt::Debug for AdminPoolBatchImportRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPoolBatchImportRequest") + .field("keys", &self.keys) + .field("proxy_node_id", &self.proxy_node_id) + .field("api_formats", &self.api_formats) + .field("settings", &self.settings.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + +#[derive(Default, Clone, serde::Deserialize)] pub struct AdminPoolBatchImportItem { #[serde(default)] pub name: String, @@ -84,6 +129,19 @@ pub struct AdminPoolBatchImportItem { pub settings: Option, } +impl std::fmt::Debug for AdminPoolBatchImportItem { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPoolBatchImportItem") + .field("name", &self.name) + .field("api_key", &"[REDACTED]") + .field("auth_type", &self.auth_type) + .field("api_formats", &self.api_formats) + .field("settings", &self.settings.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + fn admin_pool_reason_indicates_ban(reason: &str) -> bool { let normalized = reason.trim().to_ascii_lowercase(); !normalized.is_empty() @@ -118,6 +176,30 @@ pub fn admin_pool_key_quota_hard_blocked( aether_provider_pool::provider_pool_key_quota_hard_blocked(key, provider_type) } +pub fn admin_pool_key_model_quota_exhausted( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: &str, +) -> Option { + aether_provider_pool::provider_pool_key_model_quota_exhausted( + key, + provider_type, + provider_model_name, + ) +} + +pub fn admin_pool_key_model_quota_hard_blocked( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: &str, +) -> bool { + aether_provider_pool::provider_pool_key_model_quota_hard_blocked( + key, + provider_type, + provider_model_name, + ) +} + fn admin_pool_has_proxy(key: &StoredProviderCatalogKey) -> bool { match key.proxy.as_ref() { Some(Value::Object(values)) => !values.is_empty(), @@ -154,6 +236,14 @@ fn admin_pool_json_object(value: Option<&Value>) -> Option) -> Value { + admin_pool_json_object(value) + .map(Value::Object) + .as_ref() + .map(|value| admin_secret_safe_json(Some(value))) + .unwrap_or(Value::Null) +} + fn admin_pool_health_score(key: &StoredProviderCatalogKey) -> f64 { let scores = key .health_by_format @@ -656,7 +746,8 @@ mod tests { apply_admin_pool_key_settings, build_admin_pool_batch_action_plan, build_admin_pool_batch_import_key_record, build_admin_pool_key_payload, resolve_admin_pool_key_settings, validate_admin_pool_key_settings_payload, - AdminPoolBatchActionKind, AdminPoolBatchActionRequest, AdminPoolKeyPayloadContext, + AdminPoolBatchActionKind, AdminPoolBatchActionRequest, AdminPoolBatchImportItem, + AdminPoolBatchImportRequest, AdminPoolKeyPayloadContext, }; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; use serde_json::json; @@ -675,6 +766,31 @@ mod tests { key } + #[test] + fn admin_pool_import_debug_output_redacts_api_keys_and_settings() { + let request = AdminPoolBatchImportRequest { + keys: vec![AdminPoolBatchImportItem { + name: "key".to_string(), + api_key: "pool-api-key-canary".to_string(), + auth_type: "api_key".to_string(), + api_formats: vec!["openai:chat".to_string()], + settings: Some(json!({"credential": "item-settings-canary"})), + }], + proxy_node_id: None, + api_formats: Vec::new(), + settings: Some(json!({"password": "request-settings-canary"})), + }; + let debug = format!("{request:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "pool-api-key-canary", + "item-settings-canary", + "request-settings-canary", + ] { + assert!(!debug.contains(secret), "debug output leaked {secret}"); + } + } + #[test] fn detects_codex_exhaustion_from_metadata() { assert!(admin_pool_key_account_quota_exhausted( @@ -863,6 +979,52 @@ mod tests { assert_eq!(payload["scheduling_label"], json!("可用")); } + #[test] + fn admin_pool_payload_projects_historical_status_and_cooldown_diagnostics() { + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "oauth": { + "code": "invalid", + "reason": "Authorization: Bearer upstream-secret" + }, + "account": { + "code": "account_disabled", + "blocked": true, + "reason": "https://user:password@internal.test?q=secret" + }, + "quota": { + "code": "cooldown", + "exhausted": false, + "reason": "Authorization: Bearer upstream-secret", + "reset_credits": { + "detail_error": "https://user:password@internal.test?q=secret" + } + } + })); + let context = AdminPoolKeyPayloadContext { + cooldown_reason: Some( + "Authorization: Bearer upstream-secret https://user:password@internal.test?q=secret" + .to_string(), + ), + ..AdminPoolKeyPayloadContext::default() + }; + + let payload = build_admin_pool_key_payload(&key, &context); + + assert_eq!( + payload["cooldown_reason"], + json!("Provider key is cooling down") + ); + assert_eq!( + payload.pointer("/status_snapshot/oauth/reason"), + Some(&json!("OAuth token is invalid")) + ); + let serialized = payload.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); + assert!(!serialized.contains("q=secret")); + } + #[test] fn validates_and_applies_shared_key_settings() { let settings = json!({ @@ -972,10 +1134,14 @@ pub fn build_admin_pool_key_payload( ) -> Value { let health_score = admin_pool_health_score(key); let circuit_breaker_open = false; + let cooldown_reason = context + .cooldown_reason + .as_ref() + .map(|_| "Provider key is cooling down".to_string()); let (scheduling_status, scheduling_reason, scheduling_label, scheduling_reasons) = admin_pool_scheduling_payload( key, - context.cooldown_reason.as_deref(), + cooldown_reason.as_deref(), context.cooldown_ttl_seconds, ); @@ -984,25 +1150,29 @@ pub fn build_admin_pool_key_payload( "key_name": key.name, "is_active": key.is_active, "auth_type": key.auth_type, - "status_snapshot": key.status_snapshot.clone().unwrap_or_else(|| json!({})), + "status_snapshot": key + .status_snapshot + .as_ref() + .map(|value| admin_provider_status_snapshot_safe_json(Some(value))) + .unwrap_or_else(|| json!({})), "health_score": health_score, "circuit_breaker_open": circuit_breaker_open, "api_formats": admin_pool_api_formats(key), - "rate_multipliers": admin_pool_json_object(key.rate_multipliers.as_ref()), + "rate_multipliers": admin_pool_secret_safe_json_object(key.rate_multipliers.as_ref()), "internal_priority": key.internal_priority, "rpm_limit": key.rpm_limit, "cache_ttl_minutes": key.cache_ttl_minutes, "max_probe_interval_minutes": key.max_probe_interval_minutes, "note": key.note, "allowed_models": admin_pool_string_list(key.allowed_models.as_ref()), - "capabilities": admin_pool_json_object(key.capabilities.as_ref()), + "capabilities": admin_pool_secret_safe_json_object(key.capabilities.as_ref()), "auto_fetch_models": key.auto_fetch_models, "locked_models": admin_pool_string_list(key.locked_models.as_ref()), "model_include_patterns": admin_pool_string_list(key.model_include_patterns.as_ref()), "model_exclude_patterns": admin_pool_string_list(key.model_exclude_patterns.as_ref()), - "proxy": key.proxy.clone(), - "fingerprint": key.fingerprint.clone(), - "cooldown_reason": context.cooldown_reason, + "proxy": admin_secret_safe_proxy(key.proxy.as_ref()), + "fingerprint": admin_secret_safe_json(key.fingerprint.as_ref()), + "cooldown_reason": cooldown_reason, "cooldown_ttl_seconds": context.cooldown_ttl_seconds, "cost_window_usage": context.cost_window_usage, "cost_limit": context.cost_limit, diff --git a/crates/aether-admin/src/provider/quota.rs b/crates/aether-admin/src/provider/quota.rs index 48e75ebef..2a4647159 100644 --- a/crates/aether-admin/src/provider/quota.rs +++ b/crates/aether-admin/src/provider/quota.rs @@ -10,6 +10,7 @@ const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] "; const OAUTH_EXPIRED_PREFIX: &str = "[OAUTH_EXPIRED] "; const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] "; const CODEX_SPARK_LIMIT_NAME: &str = "GPT-5.3-Codex-Spark"; +const CODEX_ADDITIONAL_QUOTA_WINDOWS_KEY: &str = "additional_quota_windows"; const CODEX_ACTIVE_LIMIT_HEADER: &str = "x-codex-active-limit"; const CODEX_HEADER_PREFIX: &str = "x-codex-"; const CODEX_LIMIT_NAME_HEADER_SUFFIX: &str = "-limit-name"; @@ -281,10 +282,128 @@ pub fn parse_antigravity_usage_response( "is_forbidden": false, "forbidden_reason": serde_json::Value::Null, "forbidden_at": serde_json::Value::Null, - "models": quota_by_model, + "quota_by_model": quota_by_model, })) } +pub fn parse_antigravity_quota_summary_response( + value: &serde_json::Value, +) -> Option { + let groups = value.get("groups")?.as_array()?; + let mut parsed_groups = Vec::new(); + + for (group_index, group) in groups.iter().enumerate() { + let Some(group) = group.as_object() else { + continue; + }; + let group_id = coerce_json_string( + group + .get("groupId") + .or_else(|| group.get("group_id")) + .or_else(|| group.get("id")), + ); + let display_name = coerce_json_string( + group + .get("displayName") + .or_else(|| group.get("display_name")), + ) + .unwrap_or_else(|| format!("Quota group {}", group_index + 1)); + let description = coerce_json_string(group.get("description")); + let mut parsed_buckets = Vec::new(); + + for (bucket_index, bucket) in group + .get("buckets") + .and_then(serde_json::Value::as_array) + .into_iter() + .flatten() + .enumerate() + { + let Some(bucket) = bucket.as_object() else { + continue; + }; + let bucket_id = coerce_json_string( + bucket + .get("bucketId") + .or_else(|| bucket.get("bucket_id")) + .or_else(|| bucket.get("id")), + ) + .unwrap_or_else(|| format!("bucket-{}", bucket_index + 1)); + let window = coerce_json_string(bucket.get("window")); + let remaining_fraction = bucket + .get("remainingFraction") + .or_else(|| bucket.get("remaining_fraction")) + .and_then(coerce_json_f64) + .map(|value| value.clamp(0.0, 1.0)); + let reset_time = bucket + .get("resetTime") + .or_else(|| bucket.get("reset_time")) + .cloned() + .filter(|value| !value.is_null()); + let bucket_display_name = coerce_json_string( + bucket + .get("displayName") + .or_else(|| bucket.get("display_name")), + ); + let bucket_description = coerce_json_string(bucket.get("description")); + + if window.is_none() + && remaining_fraction.is_none() + && reset_time.is_none() + && bucket_display_name.is_none() + && bucket_description.is_none() + { + continue; + } + + let mut parsed_bucket = serde_json::Map::new(); + parsed_bucket.insert("bucket_id".to_string(), json!(bucket_id)); + if let Some(window) = window { + parsed_bucket.insert("window".to_string(), json!(window)); + } + if let Some(remaining_fraction) = remaining_fraction { + parsed_bucket.insert("remaining_fraction".to_string(), json!(remaining_fraction)); + parsed_bucket.insert( + "used_percent".to_string(), + json!((1.0 - remaining_fraction) * 100.0), + ); + parsed_bucket.insert( + "is_exhausted".to_string(), + json!(remaining_fraction <= 1e-6), + ); + } + if let Some(reset_time) = reset_time { + parsed_bucket.insert("reset_time".to_string(), reset_time); + } + if let Some(display_name) = bucket_display_name { + parsed_bucket.insert("display_name".to_string(), json!(display_name)); + } + if let Some(description) = bucket_description { + parsed_bucket.insert("description".to_string(), json!(description)); + } + parsed_buckets.push(serde_json::Value::Object(parsed_bucket)); + } + + if parsed_buckets.is_empty() { + continue; + } + let mut parsed_group = serde_json::Map::new(); + if let Some(group_id) = group_id { + parsed_group.insert("group_id".to_string(), json!(group_id)); + } + parsed_group.insert("display_name".to_string(), json!(display_name)); + if let Some(description) = description { + parsed_group.insert("description".to_string(), json!(description)); + } + parsed_group.insert( + "buckets".to_string(), + serde_json::Value::Array(parsed_buckets), + ); + parsed_groups.push(serde_json::Value::Object(parsed_group)); + } + + (!parsed_groups.is_empty()).then_some(serde_json::Value::Array(parsed_groups)) +} + pub fn parse_gemini_cli_retrieve_user_quota_response( value: &serde_json::Value, updated_at_unix_secs: u64, @@ -1768,6 +1887,46 @@ fn codex_quota_is_account_status_key(key: &str) -> bool { matches!(key, "allowed" | "limit_reached") } +fn codex_quota_merge_reset_credits( + current_object: &serde_json::Map, + incoming: &serde_json::Value, +) -> serde_json::Value { + let Some(incoming_object) = incoming.as_object() else { + return incoming.clone(); + }; + let mut merged = current_object + .get("reset_credits") + .and_then(serde_json::Value::as_object) + .cloned() + .unwrap_or_default(); + let failed_detail = incoming_object + .get("detail_status") + .and_then(serde_json::Value::as_str) + .is_some_and(|status| status.trim().eq_ignore_ascii_case("failed")); + if !failed_detail { + return incoming.clone(); + } + + for (key, value) in incoming_object { + // A failed readonly-detail request contributes diagnostics, not an + // authoritative empty list. Keep the last known items and count so a + // transient 429 cannot make reset credits disappear from the UI. + if failed_detail + && key == "credits" + && value.as_array().is_some_and(|credits| credits.is_empty()) + && merged + .get("credits") + .and_then(serde_json::Value::as_array) + .is_some_and(|credits| !credits.is_empty()) + { + continue; + } + merged.insert(key.clone(), value.clone()); + } + + serde_json::Value::Object(merged) +} + /// Merge a parsed Codex quota observation into the stored flat metadata. /// /// Positive `window_minutes` values identify windows independently of the @@ -1842,7 +2001,12 @@ pub fn merge_codex_quota_metadata_snapshot( { continue; } - merged.insert(key.clone(), value.clone()); + let value = if key == "reset_credits" { + codex_quota_merge_reset_credits(¤t_object, value) + } else { + value.clone() + }; + merged.insert(key.clone(), value); } if let Some(incoming_order) = context.request_order().filter(|incoming| { codex_quota_request_order_is_newer(*incoming, stored_metadata_watermark) @@ -2022,6 +2186,119 @@ fn codex_find_spark_rate_limit( .and_then(serde_json::Value::as_object) } +/// Preserve every additional Codex rate-limit bucket in a model-addressable +/// representation. The upstream can add independent limits over time; do +/// not discard them merely because they are not one of the legacy flat +/// primary/secondary slots. +fn codex_additional_quota_windows( + root: &serde_json::Map, +) -> Vec { + let mut windows = Vec::new(); + let Some(items) = root + .get("additional_rate_limits") + .and_then(serde_json::Value::as_array) + else { + return windows; + }; + + for (index, item) in items.iter().enumerate() { + let Some(item_object) = item.as_object() else { + continue; + }; + if item_object + .get("limit_name") + .and_then(serde_json::Value::as_str) + .is_some_and(|name| name.trim() == CODEX_SPARK_LIMIT_NAME) + { + continue; + } + let Some(limit_name) = ["limit_name", "metered_feature", "name", "id"] + .iter() + .find_map(|key| item_object.get(*key).and_then(serde_json::Value::as_str)) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + continue; + }; + let model_identity = item_object + .get("model") + .or_else(|| item_object.get("model_name")) + .or_else(|| item_object.get("model_id")) + .cloned() + .unwrap_or_else(|| json!(limit_name)); + let model_identities = item_object + .get("models") + .or_else(|| item_object.get("model_ids")) + .cloned(); + let Some(rate_limit) = item_object + .get("rate_limit") + .and_then(serde_json::Value::as_object) + else { + continue; + }; + + for (slot, slot_name) in [ + ("primary_window", "primary"), + ("secondary_window", "secondary"), + ] { + let Some(source) = rate_limit.get(slot).and_then(serde_json::Value::as_object) else { + continue; + }; + let used_percent = source.get("used_percent").and_then(coerce_json_f64); + let reset_after_seconds = source + .get("reset_after_seconds") + .or_else(|| source.get("reset_seconds")) + .and_then(coerce_json_u64); + let reset_at = source.get("reset_at").and_then(coerce_json_u64); + let window_minutes = source + .get("limit_window_seconds") + .or_else(|| source.get("window_seconds")) + .and_then(coerce_json_u64) + .map(|seconds| seconds.saturating_add(59) / 60) + .filter(|minutes| *minutes > 0); + if used_percent.is_none() + && reset_after_seconds.is_none() + && reset_at.is_none() + && window_minutes.is_none() + { + continue; + } + + let used_ratio = used_percent.map(|value| (value / 100.0).clamp(0.0, 1.0)); + let mut window = serde_json::Map::new(); + window.insert( + "code".to_string(), + json!(format!("additional_{index}_{slot_name}")), + ); + window.insert("label".to_string(), json!(limit_name)); + window.insert("scope".to_string(), json!("model")); + // The upstream limit name is the only stable identity available + // for an additional bucket. Keeping it in both fields allows + // exact model matching and generic family matching downstream. + window.insert("model".to_string(), model_identity.clone()); + if let Some(model_identities) = model_identities.clone() { + window.insert("models".to_string(), model_identities); + } + window.insert("quota_group".to_string(), json!(limit_name)); + window.insert("limit_name".to_string(), json!(limit_name)); + window.insert("used_ratio".to_string(), json!(used_ratio)); + window.insert( + "remaining_ratio".to_string(), + json!(used_ratio.map(|value| (1.0 - value).max(0.0))), + ); + window.insert( + "is_exhausted".to_string(), + json!(used_ratio.is_some_and(|value| value >= 1.0 - 1e-6)), + ); + window.insert("reset_at".to_string(), json!(reset_at)); + window.insert("reset_seconds".to_string(), json!(reset_after_seconds)); + window.insert("window_minutes".to_string(), json!(window_minutes)); + windows.push(serde_json::Value::Object(window)); + } + } + windows +} + fn codex_reset_credits_container( root: &serde_json::Map, ) -> Option<&serde_json::Map> { @@ -2128,6 +2405,14 @@ pub fn parse_codex_wham_usage_response( } } + // Keep all additional limits, including future ones with names unknown to + // this version of the gateway, so scheduling can select their bucket by + // metadata rather than by a hard-coded model name. + result.insert( + CODEX_ADDITIONAL_QUOTA_WINDOWS_KEY.to_string(), + serde_json::Value::Array(codex_additional_quota_windows(root)), + ); + if let Some(credits) = root.get("credits").and_then(serde_json::Value::as_object) { if let Some(value) = credits.get("has_credits").and_then(coerce_json_bool) { result.insert("has_credits".to_string(), json!(value)); @@ -3056,54 +3341,24 @@ pub fn codex_structured_invalid_reason(status_code: u16, upstream_message: Optio return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}工作区已停用 (deactivated_workspace)"); } if codex_looks_like_account_deactivated(Some(message)) { - let detail = if message.is_empty() { - "OpenAI 账号已停用" - } else { - message - }; - return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}"); + return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}OpenAI 账号已停用"); } if codex_looks_like_token_invalidated(Some(message)) { - let detail = if message.is_empty() { - "Codex Token 已失效" - } else { - message - }; - return format!("{OAUTH_EXPIRED_PREFIX}{detail}"); + return format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已失效"); } if codex_looks_like_token_expired(Some(message)) { - let detail = if message.is_empty() { - "Codex Token 已过期" - } else { - message - }; - return format!("{OAUTH_EXPIRED_PREFIX}{detail}"); + return format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已过期"); } if status_code == 401 { - let detail = if message.is_empty() { - "Codex Token 已过期 (401)" - } else { - message - }; - return format!("{OAUTH_EXPIRED_PREFIX}{detail}"); + return format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已过期 (401)"); } if status_code == 403 { - let detail = if message.is_empty() { - "Codex 账户访问受限 (403)" - } else { - message - }; - return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}"); + return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}Codex 账户访问受限 (403)"); } if status_code == 402 { - let detail = if message.is_empty() { - "Codex 账户需要付款 (402)" - } else { - message - }; - return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}"); + return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}Codex 账户需要付款 (402)"); } - message.to_string() + format!("Codex 请求失败 ({status_code})") } pub fn codex_runtime_invalid_reason( @@ -3127,11 +3382,8 @@ pub fn codex_runtime_invalid_reason( } fn codex_generic_forbidden_runtime_invalid_reason(upstream_message: Option<&str>) -> String { - let detail = upstream_message - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|message| format!("Codex Token 已失效 (403): {message}")) - .unwrap_or_else(|| "Codex Token 已失效 (403)".to_string()); + let _ = upstream_message; + let detail = "Codex Token 已失效 (403)"; format!("{OAUTH_EXPIRED_PREFIX}{detail}") } @@ -3139,11 +3391,8 @@ pub fn codex_soft_request_failure_reason( status_code: u16, upstream_message: Option<&str>, ) -> String { - let detail = upstream_message - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .unwrap_or_else(|| format!("Codex 请求失败 ({status_code})")); + let _ = upstream_message; + let detail = format!("Codex 请求失败 ({status_code})"); format!("{OAUTH_REQUEST_FAILED_PREFIX}{detail}") } @@ -3456,33 +3705,21 @@ pub fn parse_windsurf_user_status_response( result.insert(target.to_string(), json!(found)); } } - for (target, aliases) in [ - ( - "ban_reason", - &[ - "banReason", - "ban_reason", - "blockedReason", - "blocked_reason", - "reason", - "message", - ][..], - ), + for (status_field, reason_field, fixed_reason) in [ + ("banned", "ban_reason", "Windsurf account is suspended"), ( + "quarantined", "quarantine_reason", - &["quarantineReason", "quarantine_reason", "reason", "message"][..], + "Windsurf account is quarantined", ), ( + "is_forbidden", "forbidden_reason", - &["forbiddenReason", "forbidden_reason", "reason", "message"][..], + "Windsurf account access is restricted", ), ] { - if let Some(found) = status_sources.iter().find_map(|source| { - aliases - .iter() - .find_map(|alias| coerce_json_string(source.get(*alias))) - }) { - result.insert(target.to_string(), json!(found)); + if result.get(status_field).and_then(coerce_json_bool) == Some(true) { + result.insert(reason_field.to_string(), json!(fixed_reason)); } } @@ -3668,20 +3905,19 @@ fn normalize_chatgpt_web_numeric_reset(value: f64, observed_at: u64) -> Option Vec { - value + let image_blocked = value .get("blocked_features") .or_else(|| value.get("blockedFeatures")) .and_then(serde_json::Value::as_array) - .map(|items| { - items - .iter() - .filter_map(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .collect::>() - }) - .unwrap_or_default() + .into_iter() + .flatten() + .filter_map(serde_json::Value::as_str) + .any(chatgpt_web_is_image_quota_feature); + if image_blocked { + vec!["image_generation".to_string()] + } else { + Vec::new() + } } pub fn parse_chatgpt_web_conversation_init_response( @@ -3731,11 +3967,9 @@ pub fn parse_chatgpt_web_conversation_init_response( json!(plan_type.to_ascii_lowercase()), ); } - result.insert("blocked_features".to_string(), json!(blocked_features)); - result.insert( - "limits_progress".to_string(), - serde_json::Value::Array(limits_progress), - ); + if !blocked_features.is_empty() { + result.insert("blocked_features".to_string(), json!(blocked_features)); + } if image_blocked { result.insert("image_quota_blocked".to_string(), json!(true)); @@ -3808,13 +4042,6 @@ pub fn parse_chatgpt_web_conversation_init_response( if let Some(reset_at) = reset_at { result.insert("image_quota_reset_at".to_string(), json!(reset_at)); } - if let Some(reset_after) = coerce_json_string( - image_limit - .get("reset_after") - .or_else(|| image_limit.get("resetAfter")), - ) { - result.insert("image_quota_reset_after".to_string(), json!(reset_after)); - } } else if image_blocked { result.insert("image_quota_remaining".to_string(), json!(0.0)); } @@ -3829,11 +4056,11 @@ mod tests { codex_rate_limit_metadata_exhausted, codex_runtime_invalid_reason, codex_websocket_response_has_usage_limit_error, codex_websocket_usage_limit_reset_at, extract_execution_error_detail, merge_codex_quota_metadata_snapshot, - normalize_codex_reset_credit_consume_outcome, parse_antigravity_usage_response, - parse_chatgpt_web_conversation_init_response, parse_codex_backend_me_response, - parse_codex_usage_headers, parse_codex_websocket_rate_limits_response, - parse_codex_wham_reset_credits_detail_response, parse_codex_wham_usage_response, - parse_gemini_cli_retrieve_user_quota_response, + normalize_codex_reset_credit_consume_outcome, parse_antigravity_quota_summary_response, + parse_antigravity_usage_response, parse_chatgpt_web_conversation_init_response, + parse_codex_backend_me_response, parse_codex_usage_headers, + parse_codex_websocket_rate_limits_response, parse_codex_wham_reset_credits_detail_response, + parse_codex_wham_usage_response, parse_gemini_cli_retrieve_user_quota_response, parse_gemini_cli_v1internal_credits_response, parse_windsurf_model_configs_response, parse_windsurf_rate_limit_response, parse_windsurf_user_status_response, provider_auto_remove_quota_exhausted_keys, quota_refresh_success_invalid_state, @@ -4537,6 +4764,60 @@ mod tests { assert_eq!(outcome.metadata["primary_reset_at"], json!(20_000u64)); } + #[test] + fn codex_quota_failed_reset_credit_detail_preserves_last_known_count_and_items() { + let current = json!({ + "reset_credits": { + "available_count": 2, + "updated_at": 100u64, + "detail_source": "wham_readonly", + "detail_status": "available", + "credits": [{ + "id": "credit-1", + "display_key": "credit", + "status": "available", + "expires_at": 20_000u64 + }] + }, + "updated_at": 100u64 + }); + let incoming = json!({ + "reset_credits": { + "updated_at": 110u64, + "detail_source": "wham_readonly", + "detail_status": "failed", + "detail_error": "HTTP 429", + "credits": [] + } + }); + + let outcome = merge_codex_quota( + Some(¤t), + &incoming, + 110, + 110_000, + CodexQuotaWindowCoverage::Patch, + ); + + assert!(outcome.changed); + assert_eq!( + outcome.metadata["reset_credits"]["available_count"], + json!(2u64) + ); + assert_eq!( + outcome.metadata["reset_credits"]["credits"][0]["id"], + json!("credit-1") + ); + assert_eq!( + outcome.metadata["reset_credits"]["detail_status"], + json!("failed") + ); + assert_eq!( + outcome.metadata["reset_credits"]["detail_error"], + json!("HTTP 429") + ); + } + #[test] fn codex_quota_explicit_reset_allows_usage_drop_with_same_deadline() { let current = json!({ @@ -5556,7 +5837,7 @@ mod tests { fn codex_runtime_invalid_reason_marks_401_as_expired() { assert_eq!( codex_runtime_invalid_reason(401, Some("session expired")), - Some(format!("{OAUTH_EXPIRED_PREFIX}session expired")) + Some(format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已过期")) ); } @@ -5564,9 +5845,7 @@ mod tests { fn codex_runtime_invalid_reason_marks_account_deactivated_403() { assert_eq!( codex_runtime_invalid_reason(403, Some("account has been deactivated")), - Some(format!( - "{OAUTH_ACCOUNT_BLOCK_PREFIX}account has been deactivated" - )) + Some(format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}OpenAI 账号已停用")) ); } @@ -5574,18 +5853,14 @@ mod tests { fn codex_runtime_invalid_reason_marks_inactive_pat_owner_403_as_token_invalid() { assert_eq!( codex_runtime_invalid_reason(403, Some("Personal access token owner is inactive.")), - Some(format!( - "{OAUTH_EXPIRED_PREFIX}Personal access token owner is inactive." - )) + Some(format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已失效")) ); assert_eq!( codex_runtime_invalid_reason( 403, Some("biscuit_baker_service_auth_credential_error_status") ), - Some(format!( - "{OAUTH_EXPIRED_PREFIX}biscuit_baker_service_auth_credential_error_status" - )) + Some(format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已失效")) ); } @@ -5593,9 +5868,7 @@ mod tests { fn codex_runtime_invalid_reason_marks_deleted_agent_runtime_as_invalid() { assert_eq!( codex_runtime_invalid_reason(403, Some("Agent runtime has been deleted.")), - Some(format!( - "{OAUTH_EXPIRED_PREFIX}Agent runtime has been deleted." - )) + Some(format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已失效")) ); } @@ -5603,7 +5876,9 @@ mod tests { fn codex_runtime_invalid_reason_marks_402_as_account_blocked() { assert_eq!( codex_runtime_invalid_reason(402, Some("payment required")), - Some(format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}payment required")) + Some(format!( + "{OAUTH_ACCOUNT_BLOCK_PREFIX}Codex 账户需要付款 (402)" + )) ); } @@ -5611,12 +5886,26 @@ mod tests { fn codex_runtime_invalid_reason_marks_generic_403_as_token_invalid() { assert_eq!( codex_runtime_invalid_reason(403, Some("forbidden")), - Some(format!( - "{OAUTH_EXPIRED_PREFIX}Codex Token 已失效 (403): forbidden" - )) + Some(format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已失效 (403)")) ); } + #[test] + fn codex_invalid_reason_does_not_persist_upstream_credentials() { + let reason = codex_runtime_invalid_reason( + 401, + Some("authorization=Bearer upstream-secret https://user:pass@example.test?q=secret"), + ) + .expect("401 should produce a reason"); + + assert_eq!( + reason, + format!("{OAUTH_EXPIRED_PREFIX}Codex Token 已过期 (401)") + ); + assert!(!reason.contains("upstream-secret")); + assert!(!reason.contains("user:pass")); + } + #[test] fn codex_invalid_state_appends_refresh_failure_to_oauth_expired() { let mut key = StoredProviderCatalogKey::new( @@ -6249,6 +6538,43 @@ mod tests { parsed.get("spark_secondary_window_minutes"), Some(&json!(10_080u64)) ); + let additional = parsed + .get("additional_quota_windows") + .and_then(serde_json::Value::as_array) + .expect("all additional quota windows should be preserved"); + assert!( + additional.is_empty(), + "the Spark windows are already normalized into the dedicated spark fields" + ); + } + + #[test] + fn parses_unknown_additional_quota_buckets_without_model_specific_code() { + let parsed = parse_codex_wham_usage_response( + &json!({ + "rate_limit": { + "primary_window": {"used_percent": 100.0}, + "secondary_window": {"used_percent": 0.0} + }, + "additional_rate_limits": [{ + "limit_name": "Future-Model-Limit", + "rate_limit": { + "primary_window": { + "used_percent": 100.0, + "limit_window_seconds": 3600 + } + } + }] + }), + 1_700_000_000, + ) + .expect("quota response should parse"); + let additional = parsed["additional_quota_windows"] + .as_array() + .expect("additional windows should be an array"); + assert_eq!(additional.len(), 1); + assert_eq!(additional[0]["model"], json!("Future-Model-Limit")); + assert_eq!(additional[0]["is_exhausted"], json!(true)); } #[test] @@ -6749,20 +7075,52 @@ mod tests { .expect("antigravity quota should parse"); assert_eq!( - parsed["models"]["RateLimitResetCredit_05cbb6eeeb9c81918e011d8300f9ebfb"] + parsed["quota_by_model"]["RateLimitResetCredit_05cbb6eeeb9c81918e011d8300f9ebfb"] ["display_name"], json!("Key-1") ); assert_eq!( - parsed["models"]["RateLimitResetCredit_05cbb6eeeb9c81918e011d8300f9ebfb"]["reset_time"], + parsed["quota_by_model"]["RateLimitResetCredit_05cbb6eeeb9c81918e011d8300f9ebfb"] + ["reset_time"], json!("2030-01-01T00:00:00Z") ); assert_eq!( - parsed["models"]["gemini-3-pro-preview"]["display_name"], + parsed["quota_by_model"]["gemini-3-pro-preview"]["display_name"], json!("Gemini 3 Pro Preview") ); } + #[test] + fn parses_antigravity_grouped_weekly_and_five_hour_quota() { + let groups = parse_antigravity_quota_summary_response(&json!({ + "groups": [{ + "displayName": "Claude and GPT models", + "description": "Shared quota", + "buckets": [{ + "bucketId": "3p-5h", + "window": "5h", + "remainingFraction": 0.25, + "resetTime": "2026-05-05T05:00:00Z", + "displayName": "5 hour" + }, { + "bucketId": "3p-weekly", + "window": "weekly", + "remainingFraction": 0.8, + "resetTime": "2026-05-11T00:00:00Z" + }] + }] + })) + .expect("grouped Antigravity quota should parse"); + + assert_eq!(groups[0]["display_name"], json!("Claude and GPT models")); + assert_eq!(groups[0]["description"], json!("Shared quota")); + assert_eq!(groups[0]["buckets"][0]["bucket_id"], json!("3p-5h")); + assert_eq!(groups[0]["buckets"][0]["window"], json!("5h")); + assert_eq!(groups[0]["buckets"][0]["remaining_fraction"], json!(0.25)); + assert_eq!(groups[0]["buckets"][0]["used_percent"], json!(75.0)); + assert_eq!(groups[0]["buckets"][1]["bucket_id"], json!("3p-weekly")); + } + #[test] fn parses_gemini_cli_retrieve_user_quota_buckets() { let parsed = parse_gemini_cli_retrieve_user_quota_response( @@ -6923,11 +7281,36 @@ mod tests { assert_eq!(parsed.get("quarantined"), Some(&json!(true))); assert_eq!( parsed.get("quarantine_reason"), - Some(&json!("quota review")) + Some(&json!("Windsurf account is quarantined")) ); assert_eq!(parsed.get("updated_at"), Some(&json!(1_770_000_000u64))); } + #[test] + fn windsurf_status_parser_never_persists_upstream_reason_or_message() { + let parsed = parse_windsurf_user_status_response( + &json!({ + "userStatus": { + "isBanned": true, + "reason": "Authorization: Bearer upstream-secret", + "message": "https://user:password@internal.test?q=secret", + "planStatus": {"dailyQuotaRemainingPercent": 10} + } + }), + 1_770_000_000, + ) + .expect("windsurf status should parse"); + + assert_eq!( + parsed.get("ban_reason"), + Some(&json!("Windsurf account is suspended")) + ); + let serialized = parsed.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); + assert!(!serialized.contains("q=secret")); + } + #[test] fn parses_windsurf_model_configs_response() { let parsed = parse_windsurf_model_configs_response( @@ -7016,7 +7399,11 @@ mod tests { fn parses_chatgpt_web_blocked_image_feature_as_zero_remaining() { let parsed = parse_chatgpt_web_conversation_init_response( &json!({ - "blocked_features": ["image_generation"], + "blocked_features": [ + "image_generation", + "Authorization: Bearer upstream-secret", + "https://user:password@internal.test?q=secret" + ], "limits_progress": [] }), 1_778_067_246, @@ -7025,5 +7412,43 @@ mod tests { assert_eq!(parsed.get("image_quota_blocked"), Some(&json!(true))); assert_eq!(parsed.get("image_quota_remaining"), Some(&json!(0.0))); + assert_eq!( + parsed.get("blocked_features"), + Some(&json!(["image_generation"])) + ); + assert!(parsed.get("limits_progress").is_none()); + assert!(!parsed.to_string().contains("upstream-secret")); + } + + #[test] + fn chatgpt_web_parser_projects_image_limit_scalars_only() { + let parsed = parse_chatgpt_web_conversation_init_response( + &json!({ + "limits_progress": [{ + "feature_name": "image_gen", + "remaining": 8, + "total": 12, + "reset_after": "60", + "message": "Authorization: Bearer upstream-secret", + "nested": { + "url": "https://user:password@internal.test?q=secret" + } + }] + }), + 1_778_067_246, + ) + .expect("image quota should parse"); + + assert_eq!(parsed.get("image_quota_remaining"), Some(&json!(8.0))); + assert_eq!(parsed.get("image_quota_total"), Some(&json!(12.0))); + assert_eq!( + parsed.get("image_quota_reset_at"), + Some(&json!(1_778_067_306u64)) + ); + assert!(parsed.get("limits_progress").is_none()); + assert!(parsed.get("image_quota_reset_after").is_none()); + let serialized = parsed.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); } } diff --git a/crates/aether-admin/src/provider/redaction.rs b/crates/aether-admin/src/provider/redaction.rs new file mode 100644 index 000000000..39890da95 --- /dev/null +++ b/crates/aether-admin/src/provider/redaction.rs @@ -0,0 +1,2545 @@ +use serde_json::{json, Map, Value}; +use std::collections::BTreeMap; + +const REDACTED_VALUE: &str = "***"; +const REDACTED_UPSTREAM_DIAGNOSTIC: &str = "[REDACTED upstream diagnostic]"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum AdminProxyCredentialField { + Username, + Password, +} + +impl AdminProxyCredentialField { + fn canonical_key(self) -> &'static str { + match self { + Self::Username => "username", + Self::Password => "password", + } + } +} + +/// Classifies the proxy credential spellings accepted by the catalog normalizer. +/// Keep this shared with the admin redactor so an accepted alias can never bypass +/// masking before it is migrated to the canonical encrypted field. +pub fn admin_proxy_credential_field(key: &str) -> Option { + proxy_credential_field_from_compact(&compact_json_key(key)) +} + +/// Returns whether the generic admin redactor treats this JSON field as secret. +/// Catalog persistence uses the same predicate to reject unsupported secret shapes. +pub fn admin_json_field_is_sensitive(key: &str, value: &Value) -> bool { + json_secret_key(&compact_json_key(key), value) +} + +/// Detects secrets that the admin projection hides because of surrounding JSON +/// semantics rather than the field name alone (for example URL userinfo/query, +/// header maps, mutation rules, or a nested proxy object). +pub fn admin_json_field_has_contextual_secrets(key: &str, value: &Value) -> bool { + let compact_key = compact_json_key(key); + if json_url_key(&compact_key) { + return value + .as_str() + .is_some_and(network_url_has_hidden_components); + } + if compact_key == "proxy" { + return admin_secret_safe_proxy(Some(value)) != *value; + } + if header_rules_key(&compact_key) { + return admin_secret_safe_header_rules(Some(value)) != *value; + } + if body_rules_key(&compact_key) { + return admin_secret_safe_body_rules(Some(value)) != *value; + } + if header_values_key(&compact_key) { + return value + .as_object() + .is_some_and(|headers| headers.values().any(secret_value_is_set)); + } + false +} + +#[derive(Clone, Copy)] +enum RuleKind { + Header, + Body, +} + +pub fn admin_secret_safe_json(value: Option<&Value>) -> Value { + value.map(redact_json_value).unwrap_or(Value::Null) +} + +pub fn admin_provider_oauth_invalid_reason_safe_text(reason: Option<&str>) -> Option { + let reason = reason.map(str::trim).filter(|value| !value.is_empty())?; + let lowered = reason.to_ascii_lowercase(); + + if tagged_diagnostic_reason(reason, "[ACCOUNT_BLOCK]") { + return Some(format!( + "[ACCOUNT_BLOCK] {}", + canonical_account_block_reason(&lowered) + )); + } + if tagged_diagnostic_reason(reason, "[OAUTH_EXPIRED]") { + let detail = if diagnostic_reason_is_hard_token_invalid(&lowered) { + "OAuth token is invalid" + } else { + "OAuth token has expired" + }; + return Some(format!("[OAUTH_EXPIRED] {detail}")); + } + if tagged_diagnostic_reason(reason, "[REFRESH_FAILED]") { + return Some("[REFRESH_FAILED] OAuth token refresh failed".to_string()); + } + if tagged_diagnostic_reason(reason, "[REQUEST_FAILED]") { + let detail = if lowered.contains("agent runtime has been deleted") { + "Agent runtime has been deleted" + } else { + "OAuth account status request failed" + }; + return Some(format!("[REQUEST_FAILED] {detail}")); + } + + Some( + if diagnostic_reason_is_hard_token_invalid(&lowered) { + "OAuth token is invalid" + } else if diagnostic_reason_is_token_expired(&lowered) { + "OAuth token has expired" + } else if diagnostic_reason_is_account_blocked(&lowered) { + canonical_account_block_reason(&lowered) + } else { + "OAuth credential is unavailable" + } + .to_string(), + ) +} + +pub fn admin_provider_upstream_metadata_safe_json(value: Option<&Value>) -> Value { + let mut projected = admin_secret_safe_json(value); + let Some(root) = projected.as_object_mut() else { + return projected; + }; + + if let Some(chatgpt_web) = root.get_mut("chatgpt_web") { + *chatgpt_web = project_chatgpt_web_metadata(chatgpt_web); + } + redact_upstream_diagnostic_fields(&mut projected); + projected +} + +pub fn admin_provider_metadata_bucket_safe_json( + provider_type: &str, + value: Option<&Value>, +) -> Value { + let provider_type = provider_type.trim().to_ascii_lowercase(); + let Some(value) = value else { + return Value::Null; + }; + if provider_type.is_empty() { + let mut projected = admin_secret_safe_json(Some(value)); + redact_upstream_diagnostic_fields(&mut projected); + return projected; + } + + let mut root = Map::new(); + root.insert(provider_type.clone(), value.clone()); + let mut projected = admin_provider_upstream_metadata_safe_json(Some(&Value::Object(root))); + projected + .as_object_mut() + .and_then(|object| object.remove(&provider_type)) + .unwrap_or(Value::Null) +} + +pub fn admin_provider_status_snapshot_safe_json(value: Option<&Value>) -> Value { + let Some(snapshot) = value.and_then(Value::as_object) else { + return Value::Null; + }; + + let mut projected = Map::new(); + projected.insert( + "oauth".to_string(), + project_oauth_status_snapshot(snapshot.get("oauth")), + ); + projected.insert( + "account".to_string(), + project_account_status_snapshot(snapshot.get("account")), + ); + projected.insert( + "quota".to_string(), + project_quota_status_snapshot(snapshot.get("quota")), + ); + Value::Object(projected) +} + +pub fn admin_restore_secret_safe_json(existing: Option<&Value>, incoming: &Value) -> Value { + restore_json_value(existing, incoming, None) +} + +pub fn admin_secret_safe_header_rules(rules: Option<&Value>) -> Value { + redact_rule_array(rules, redact_header_rule) +} + +pub fn admin_restore_secret_safe_header_rules(existing: Option<&Value>, incoming: &Value) -> Value { + restore_rule_array(existing, incoming, RuleKind::Header) +} + +pub fn admin_secret_safe_body_rules(rules: Option<&Value>) -> Value { + redact_rule_array(rules, redact_body_rule) +} + +pub fn admin_restore_secret_safe_body_rules(existing: Option<&Value>, incoming: &Value) -> Value { + restore_rule_array(existing, incoming, RuleKind::Body) +} + +pub fn admin_secret_safe_url(value: Option<&str>) -> Value { + value + .and_then(sanitize_network_url) + .map(Value::String) + .unwrap_or(Value::Null) +} + +pub fn admin_restore_secret_safe_url(existing: Option<&str>, incoming: &str) -> String { + if existing.and_then(sanitize_network_url).as_deref() == Some(incoming.trim()) { + return existing.unwrap_or(incoming).to_string(); + } + incoming.to_string() +} + +pub fn admin_secret_safe_proxy(proxy: Option<&Value>) -> Value { + let Some(proxy) = proxy else { + return Value::Null; + }; + if let Some(proxy_url) = proxy.as_str() { + return admin_secret_safe_url(Some(proxy_url)); + } + let Some(proxy) = proxy.as_object() else { + return Value::Null; + }; + + let mut projected = Map::new(); + let mut has_credentials = false; + for (key, value) in proxy { + let compact_key = compact_json_key(key); + if compact_key == "hascredentials" { + continue; + } + if matches!(compact_key.as_str(), "url" | "proxyurl") { + let sanitized = value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(|value| { + has_credentials |= network_url_has_hidden_components(value); + sanitize_network_url(value).map(Value::String) + }) + .unwrap_or(Value::Null); + projected.insert(key.clone(), sanitized); + continue; + } + if proxy_credential_key(&compact_key) { + has_credentials |= proxy_secret_value_is_set(value); + projected.insert(key.clone(), redact_proxy_secret_value(value)); + continue; + } + projected.insert(key.clone(), redact_proxy_json_value_for_key(key, value)); + } + if has_credentials { + projected.insert("has_credentials".to_string(), Value::Bool(true)); + } + + Value::Object(projected) +} + +pub fn admin_restore_secret_safe_proxy(existing: Option<&Value>, incoming: &Value) -> Value { + if let Some(incoming_url) = incoming.as_str() { + if admin_proxy_identities_match(existing, incoming) { + return existing + .and_then(proxy_url_value) + .map(|existing_url| { + Value::String(admin_restore_secret_safe_url( + Some(existing_url), + incoming_url, + )) + }) + .unwrap_or_else(|| Value::String(incoming_url.to_string())); + } + return Value::String(incoming_url.to_string()); + } + + let Some(incoming_object) = incoming.as_object() else { + return incoming.clone(); + }; + let existing_object = existing.and_then(Value::as_object); + let identities_match = admin_proxy_identities_match(existing, incoming); + let mut restored = Map::new(); + for (key, incoming_value) in incoming_object { + let compact_key = compact_json_key(key); + if compact_key == "hascredentials" { + continue; + } + if let Some(field) = admin_proxy_credential_field(key) { + if incoming_value.as_str() == Some(REDACTED_VALUE) { + if identities_match { + if let Some(value) = unambiguous_proxy_credential(existing, field) { + if proxy_secret_value_is_set(&value) { + restored.insert(key.clone(), value); + } + } + } + continue; + } + } + let existing_value = existing_object.and_then(|object| object.get(key)); + if let Some(value) = + restore_proxy_json_value(existing_value, incoming_value, Some(key), identities_match) + { + restored.insert(key.clone(), value); + } + } + + // Secret-safe admin projections normally contain masks, but partial clients may + // omit those fields. Preserve omitted credentials only while the normalized proxy + // authority and node identity are unchanged. Any identity change therefore needs + // explicit new credentials. + if identities_match { + for field in [ + AdminProxyCredentialField::Username, + AdminProxyCredentialField::Password, + ] { + if proxy_contains_credential_field(incoming, field) { + continue; + } + if let Some(existing_value) = unambiguous_proxy_credential(existing, field) { + if proxy_secret_value_is_set(&existing_value) { + restored.insert(field.canonical_key().to_string(), existing_value); + } + } + } + } + Value::Object(restored) +} + +fn restore_proxy_json_value( + existing: Option<&Value>, + incoming: &Value, + key: Option<&str>, + identities_match: bool, +) -> Option { + if let Some(key) = key { + let compact_key = compact_json_key(key); + if matches!(compact_key.as_str(), "url" | "proxyurl") { + return Some( + incoming + .as_str() + .map(|incoming_url| { + if identities_match { + restore_url_value(existing, incoming_url) + } else { + Value::String(incoming_url.to_string()) + } + }) + .unwrap_or_else(|| incoming.clone()), + ); + } + if proxy_credential_key(&compact_key) { + if incoming.as_str() == Some(REDACTED_VALUE) { + return identities_match + .then(|| { + existing + .filter(|value| proxy_secret_value_is_set(value)) + .cloned() + }) + .flatten(); + } + return Some(incoming.clone()); + } + } + + match incoming { + Value::Array(incoming_values) => { + let existing_values = existing.and_then(Value::as_array); + Some(Value::Array( + incoming_values + .iter() + .enumerate() + .filter_map(|(index, incoming_value)| { + restore_proxy_json_value( + existing_values.and_then(|values| values.get(index)), + incoming_value, + None, + identities_match, + ) + }) + .collect(), + )) + } + Value::Object(incoming_object) => { + let existing_object = existing.and_then(Value::as_object); + Some(Value::Object( + incoming_object + .iter() + .filter(|(key, _)| compact_json_key(key) != "hascredentials") + .filter_map(|(key, incoming_value)| { + restore_proxy_json_value( + existing_object.and_then(|object| object.get(key)), + incoming_value, + Some(key), + identities_match, + ) + .map(|value| (key.clone(), value)) + }) + .collect(), + )) + } + _ => Some(incoming.clone()), + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct AdminProxyIdentity { + url: Option<(String, String, u16)>, + node_id: Option, +} + +fn admin_proxy_identities_match(existing: Option<&Value>, incoming: &Value) -> bool { + existing + .and_then(admin_proxy_identity) + .zip(admin_proxy_identity(incoming)) + .is_some_and(|(existing, incoming)| { + existing == incoming && (existing.url.is_some() || existing.node_id.is_some()) + }) +} + +fn admin_proxy_identity(value: &Value) -> Option { + if let Some(url) = value.as_str() { + return normalized_proxy_url_identity(url).map(|url| AdminProxyIdentity { + url: Some(url), + node_id: None, + }); + } + let object = value.as_object()?; + + let mut url_field_seen = false; + let mut url = None; + let mut node_field_seen = false; + let mut node_id = None; + for (key, value) in object { + match compact_json_key(key).as_str() { + "url" | "proxyurl" => { + let candidate = if value.is_null() { + None + } else { + Some(normalized_proxy_url_identity(value.as_str()?)?) + }; + if url_field_seen && url != candidate { + return None; + } + url_field_seen = true; + url = candidate; + } + "nodeid" => { + let candidate = if value.is_null() { + None + } else { + value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }; + if !value.is_null() && candidate.is_none() { + return None; + } + if node_field_seen && node_id != candidate { + return None; + } + node_field_seen = true; + node_id = candidate; + } + _ => {} + } + } + Some(AdminProxyIdentity { url, node_id }) +} + +fn normalized_proxy_url_identity(value: &str) -> Option<(String, String, u16)> { + let parsed = url::Url::parse(value.trim()).ok()?; + let scheme = parsed.scheme().to_ascii_lowercase(); + if !matches!(scheme.as_str(), "http" | "https" | "socks5" | "socks5h") { + return None; + } + let host = parsed.host_str()?.to_ascii_lowercase(); + let port = parsed.port().or(match scheme.as_str() { + "http" => Some(80), + "https" => Some(443), + "socks5" | "socks5h" => Some(1080), + _ => None, + })?; + Some((scheme, host, port)) +} + +fn proxy_url_value(value: &Value) -> Option<&str> { + if let Some(url) = value.as_str() { + return Some(url); + } + value.as_object()?.iter().find_map(|(key, value)| { + matches!(compact_json_key(key).as_str(), "url" | "proxyurl") + .then(|| value.as_str()) + .flatten() + }) +} + +fn proxy_contains_credential_field(value: &Value, target: AdminProxyCredentialField) -> bool { + match value { + Value::Array(values) => values + .iter() + .any(|value| proxy_contains_credential_field(value, target)), + Value::Object(object) => object.iter().any(|(key, value)| { + admin_proxy_credential_field(key) == Some(target) + || proxy_contains_credential_field(value, target) + }), + _ => false, + } +} + +fn unambiguous_proxy_credential( + value: Option<&Value>, + target: AdminProxyCredentialField, +) -> Option { + fn collect(value: &Value, target: AdminProxyCredentialField, values: &mut Vec) { + match value { + Value::Array(items) => { + for item in items { + collect(item, target, values); + } + } + Value::Object(object) => { + for (key, value) in object { + if admin_proxy_credential_field(key) == Some(target) { + values.push(value.clone()); + } else { + collect(value, target, values); + } + } + } + _ => {} + } + } + + let mut values = Vec::new(); + collect(value?, target, &mut values); + let first = values.first()?.clone(); + values.iter().all(|value| value == &first).then_some(first) +} + +fn tagged_diagnostic_reason(reason: &str, tag: &str) -> bool { + reason + .lines() + .map(str::trim) + .any(|line| line.starts_with(tag)) +} + +fn diagnostic_reason_is_hard_token_invalid(lowered: &str) -> bool { + [ + "token invalid", + "token_invalid", + "token has been invalidated", + "token invalidated", + "token revoked", + "personal access token owner is inactive", + "auth_credential", + "invalid token", + "token 无效", + "token 失效", + "令牌无效", + "令牌失效", + ] + .iter() + .any(|keyword| lowered.contains(keyword)) +} + +fn diagnostic_reason_is_token_expired(lowered: &str) -> bool { + [ + "token expired", + "token has expired", + "session expired", + "access token expired", + "oauth_token_expired", + "token 过期", + "令牌过期", + "已过期", + ] + .iter() + .any(|keyword| lowered.contains(keyword)) +} + +fn diagnostic_reason_is_account_blocked(lowered: &str) -> bool { + [ + "banned", + "blocked", + "suspended", + "forbidden", + "disabled", + "deactivated", + "validation_required", + "verify your account", + "封禁", + "停用", + "受限", + "验证", + ] + .iter() + .any(|keyword| lowered.contains(keyword)) +} + +fn canonical_account_block_reason(lowered: &str) -> &'static str { + if lowered.contains("deactivated_workspace") || lowered.contains("workspace deactivated") { + "Workspace is deactivated" + } else if diagnostic_reason_is_hard_token_invalid(lowered) { + "OAuth token is invalid" + } else if diagnostic_reason_is_token_expired(lowered) { + "OAuth token has expired" + } else if lowered.contains("validation_required") || lowered.contains("verify your account") { + "Account verification is required" + } else if lowered.contains("disabled") || lowered.contains("account has been deactivated") { + "Account is disabled" + } else if lowered.contains("banned") || lowered.contains("suspended") { + "Account is suspended" + } else { + "Account access is restricted" + } +} + +fn project_chatgpt_web_metadata(value: &Value) -> Value { + let Some(source) = value.as_object() else { + return Value::Null; + }; + let mut projected = Map::new(); + + for field in [ + "updated_at", + "image_quota_remaining", + "image_quota_total", + "image_quota_used", + "image_quota_reset_at", + "image_quota_last_local_request_at", + "image_quota_local_request_count", + ] { + copy_json_number_or_null(source, &mut projected, field); + } + // This is an opaque idempotency key, not a diagnostic or credential. Keep it + // bounded and token-safe so quota request de-duplication survives persistence. + copy_safe_token_string_or_null(source, &mut projected, "image_quota_last_local_request_key"); + copy_json_bool_or_null(source, &mut projected, "image_quota_blocked"); + for field in [ + "default_model_slug", + "plan_type", + "email", + "account_id", + "account_user_id", + "user_id", + "image_quota_limit_source", + ] { + copy_safe_display_string_or_null(source, &mut projected, field); + } + + let image_blocked = source + .get("image_quota_blocked") + .and_then(Value::as_bool) + .unwrap_or(false) + || source + .get("blocked_features") + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter_map(Value::as_str) + .any(chatgpt_web_is_image_feature); + if image_blocked { + projected.insert("image_quota_blocked".to_string(), Value::Bool(true)); + projected.insert("blocked_features".to_string(), json!(["image_generation"])); + } + if source + .get("image_quota_feature_name") + .and_then(Value::as_str) + .is_some_and(chatgpt_web_is_image_feature) + { + projected.insert( + "image_quota_feature_name".to_string(), + Value::String("image_generation".to_string()), + ); + } + + Value::Object(projected) +} + +fn chatgpt_web_is_image_feature(value: &str) -> bool { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "image_gen" | "image_generation" | "image_edit" | "img_gen" + ) +} + +fn redact_upstream_diagnostic_fields(value: &mut Value) { + match value { + Value::Array(values) => { + for value in values { + redact_upstream_diagnostic_fields(value); + } + } + Value::Object(object) => { + for (key, value) in object { + if upstream_diagnostic_key(key) && diagnostic_value_has_content(value) { + *value = Value::String(REDACTED_UPSTREAM_DIAGNOSTIC.to_string()); + } else { + redact_upstream_diagnostic_fields(value); + } + } + } + _ => {} + } +} + +fn upstream_diagnostic_key(key: &str) -> bool { + let compact = compact_json_key(key); + matches!( + compact.as_str(), + "bodytext" + | "detail" + | "details" + | "error" + | "errors" + | "rawbody" + | "requestbody" + | "responsebody" + ) || compact.ends_with("error") + || compact.ends_with("message") + || compact.ends_with("reason") +} + +fn diagnostic_value_has_content(value: &Value) -> bool { + match value { + Value::String(value) => !value.trim().is_empty(), + Value::Array(value) => !value.is_empty(), + Value::Object(value) => !value.is_empty(), + _ => false, + } +} + +fn project_oauth_status_snapshot(value: Option<&Value>) -> Value { + let source = value.and_then(Value::as_object); + let code = source + .and_then(|source| source.get("code")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|code| { + matches!( + *code, + "none" + | "valid" + | "expiring" + | "expired" + | "invalid" + | "reauth_required" + | "check_failed" + ) + }) + .unwrap_or("none"); + let mut projected = Map::new(); + projected.insert("code".to_string(), json!(code)); + projected.insert( + "label".to_string(), + optional_static_text(oauth_status_label(code)), + ); + projected.insert( + "reason".to_string(), + optional_static_text(oauth_status_reason(code)), + ); + if let Some(source) = source { + for field in ["expires_at", "invalid_at"] { + copy_json_number_or_null(source, &mut projected, field); + } + for field in ["requires_reauth", "usable_until_expiry", "expiring_soon"] { + copy_json_bool_or_null(source, &mut projected, field); + } + copy_safe_token_string_or_null(source, &mut projected, "source"); + } + Value::Object(projected) +} + +fn oauth_status_label(code: &str) -> Option<&'static str> { + match code { + "valid" => Some("有效"), + "expiring" => Some("即将过期"), + "expired" => Some("已过期"), + "invalid" => Some("已失效"), + "reauth_required" => Some("续期失败"), + "check_failed" => Some("检查失败"), + _ => None, + } +} + +fn oauth_status_reason(code: &str) -> Option<&'static str> { + match code { + "expired" => Some("OAuth token has expired"), + "invalid" => Some("OAuth token is invalid"), + "reauth_required" => Some("OAuth token refresh failed"), + "check_failed" => Some("OAuth account status request failed"), + _ => None, + } +} + +fn project_account_status_snapshot(value: Option<&Value>) -> Value { + let source = value.and_then(Value::as_object); + let blocked = source + .and_then(|source| source.get("blocked")) + .and_then(Value::as_bool) + .unwrap_or(false); + let mut code = source + .and_then(|source| source.get("code")) + .and_then(Value::as_str) + .and_then(safe_status_token) + .unwrap_or(if blocked { "account_blocked" } else { "ok" }); + if code == "ok" && blocked { + code = "account_blocked"; + } + + let mut projected = Map::new(); + projected.insert("code".to_string(), json!(code)); + projected.insert( + "label".to_string(), + optional_static_text(account_status_label(code)), + ); + projected.insert( + "reason".to_string(), + optional_static_text(account_status_reason(code, blocked)), + ); + projected.insert("blocked".to_string(), Value::Bool(blocked)); + if let Some(source) = source { + copy_safe_token_string_or_null(source, &mut projected, "source"); + copy_json_bool_or_null(source, &mut projected, "recoverable"); + } + Value::Object(projected) +} + +fn account_status_label(code: &str) -> Option<&'static str> { + match code { + "account_banned" | "account_suspended" => Some("账号封禁"), + "account_quarantined" => Some("账号隔离"), + "workspace_deactivated" => Some("工作区停用"), + "account_disabled" => Some("账号停用"), + "account_forbidden" => Some("访问受限"), + "account_blocked" => Some("账号异常"), + "account_verification" => Some("需要验证"), + "oauth_token_invalid" => Some("Token 失效"), + "oauth_token_expired" => Some("Token 过期"), + "oauth_request_failed" => Some("请求失败"), + _ => None, + } +} + +fn account_status_reason(code: &str, blocked: bool) -> Option<&'static str> { + match code { + "account_banned" | "account_suspended" => Some("Account is suspended"), + "account_quarantined" => Some("Account is quarantined"), + "workspace_deactivated" => Some("Workspace is deactivated"), + "account_disabled" => Some("Account is disabled"), + "account_forbidden" => Some("Account access is restricted"), + "account_verification" => Some("Account verification is required"), + "oauth_token_invalid" => Some("OAuth token is invalid"), + "oauth_token_expired" => Some("OAuth token has expired"), + "oauth_request_failed" => Some("OAuth account status request failed"), + "account_blocked" => Some("Account is unavailable"), + _ if blocked => Some("Account is unavailable"), + _ => None, + } +} + +fn project_quota_status_snapshot(value: Option<&Value>) -> Value { + let source = value.and_then(Value::as_object); + let exhausted = source + .and_then(|source| source.get("exhausted")) + .and_then(Value::as_bool) + .unwrap_or(false); + let code = source + .and_then(|source| source.get("code")) + .and_then(Value::as_str) + .and_then(safe_status_token) + .unwrap_or(if exhausted { "exhausted" } else { "unknown" }); + + let mut projected = Map::new(); + projected.insert("code".to_string(), json!(code)); + projected.insert( + "label".to_string(), + optional_static_text(quota_status_label(code)), + ); + projected.insert( + "reason".to_string(), + optional_static_text(quota_status_reason(code, exhausted)), + ); + projected.insert("exhausted".to_string(), Value::Bool(exhausted)); + + let Some(source) = source else { + return Value::Object(projected); + }; + for field in [ + "version", + "observed_at", + "usage_ratio", + "updated_at", + "reset_at", + "reset_seconds", + "allowed_models_count", + ] { + copy_json_number_or_null(source, &mut projected, field); + } + for field in ["allowed", "limit_reached"] { + copy_json_bool_or_null(source, &mut projected, field); + } + for field in [ + "provider_type", + "freshness", + "source", + "plan_type", + "pool_tier", + ] { + copy_safe_token_string_or_null(source, &mut projected, field); + } + if let Some(credits) = source.get("credits") { + projected.insert("credits".to_string(), project_quota_credits(credits)); + } + if let Some(reset_credits) = source.get("reset_credits") { + projected.insert( + "reset_credits".to_string(), + project_quota_reset_credits(reset_credits), + ); + } + if let Some(rate_limit) = source.get("rate_limit") { + projected.insert( + "rate_limit".to_string(), + project_quota_rate_limit(rate_limit), + ); + } + if let Some(windows) = source.get("windows").and_then(Value::as_array) { + projected.insert( + "windows".to_string(), + Value::Array(windows.iter().filter_map(project_quota_window).collect()), + ); + } + Value::Object(projected) +} + +fn quota_status_label(code: &str) -> Option<&'static str> { + match code { + "banned" => Some("账号已封禁"), + "quarantined" => Some("账号隔离中"), + "forbidden" => Some("访问受限"), + "exhausted" => Some("额度耗尽"), + "cooldown" => Some("冷却中"), + "error" => Some("刷新失败"), + _ => None, + } +} + +fn quota_status_reason(code: &str, exhausted: bool) -> Option<&'static str> { + match code { + "banned" => Some("Account is suspended"), + "quarantined" => Some("Account is quarantined"), + "forbidden" => Some("Account access is restricted"), + "exhausted" => Some("Quota is exhausted"), + "cooldown" => Some("Quota is temporarily cooling down"), + "error" => Some("Quota refresh failed"), + _ if exhausted => Some("Quota is exhausted"), + _ => None, + } +} + +fn project_quota_credits(value: &Value) -> Value { + let Some(source) = value.as_object() else { + return Value::Null; + }; + let mut projected = Map::new(); + for field in ["has_credits", "unlimited"] { + copy_json_bool_or_null(source, &mut projected, field); + } + for field in ["balance", "remaining", "consumed", "total", "updated_at"] { + copy_json_number_or_null(source, &mut projected, field); + } + Value::Object(projected) +} + +fn project_quota_reset_credits(value: &Value) -> Value { + let Some(source) = value.as_object() else { + return Value::Null; + }; + let mut projected = Map::new(); + for field in ["available_count", "updated_at"] { + copy_json_number_or_null(source, &mut projected, field); + } + for field in ["detail_source", "detail_status"] { + copy_safe_token_string_or_null(source, &mut projected, field); + } + if let Some(detail_error) = source + .get("detail_error") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + { + projected.insert( + "detail_error".to_string(), + Value::String(canonical_quota_detail_error(detail_error).to_string()), + ); + } + if let Some(credits) = source.get("credits").and_then(Value::as_array) { + projected.insert( + "credits".to_string(), + Value::Array( + credits + .iter() + .filter_map(project_quota_reset_credit) + .collect(), + ), + ); + } + Value::Object(projected) +} + +fn canonical_quota_detail_error(value: &str) -> &'static str { + let lowered = value.to_ascii_lowercase(); + if lowered.contains("timeout") || lowered.contains("timed out") { + "Quota refresh timed out" + } else if ["connect", "connection", "dns", "tls", "certificate"] + .iter() + .any(|keyword| lowered.contains(keyword)) + { + "Quota refresh connection failed" + } else { + "Quota refresh failed" + } +} + +fn project_quota_reset_credit(value: &Value) -> Option { + let source = value.as_object()?; + let mut projected = Map::new(); + for field in ["id", "display_key", "status"] { + copy_safe_display_string_or_null(source, &mut projected, field); + } + for field in ["granted_at", "expires_at", "remaining_seconds"] { + copy_json_number_or_null(source, &mut projected, field); + } + (!projected.is_empty()).then_some(Value::Object(projected)) +} + +fn project_quota_rate_limit(value: &Value) -> Value { + let Some(source) = value.as_object() else { + return Value::Null; + }; + let mut projected = Map::new(); + for field in ["limited", "has_capacity"] { + copy_json_bool_or_null(source, &mut projected, field); + } + for field in [ + "messages_remaining", + "max_messages", + "retry_after_ms", + "reset_at", + "reset_seconds", + "updated_at", + ] { + copy_json_number_or_null(source, &mut projected, field); + } + Value::Object(projected) +} + +fn project_quota_window(value: &Value) -> Option { + let source = value.as_object()?; + let code = source + .get("code") + .and_then(Value::as_str) + .and_then(safe_status_token)?; + let mut projected = Map::new(); + projected.insert("code".to_string(), json!(code)); + for field in ["label", "model"] { + copy_safe_display_string_or_null(source, &mut projected, field); + } + for field in ["scope", "unit", "window", "quota_group", "bucket_id"] { + copy_safe_token_string_or_null(source, &mut projected, field); + } + for field in ["quota_group_label", "description"] { + copy_safe_display_string_or_null(source, &mut projected, field); + } + for field in [ + "used_ratio", + "remaining_ratio", + "used_value", + "remaining_value", + "limit_value", + "reset_at", + "reset_seconds", + "window_minutes", + "usage_reset_at", + ] { + copy_json_number_or_null(source, &mut projected, field); + } + copy_json_bool_or_null(source, &mut projected, "is_exhausted"); + if let Some(usage) = source.get("usage") { + projected.insert("usage".to_string(), project_quota_window_usage(usage)); + } + Some(Value::Object(projected)) +} + +fn project_quota_window_usage(value: &Value) -> Value { + let Some(source) = value.as_object() else { + return Value::Null; + }; + let mut projected = Map::new(); + for field in ["request_count", "total_tokens"] { + copy_json_number_or_null(source, &mut projected, field); + } + if let Some(value) = source.get("total_cost_usd") { + let safe = value.is_number() + || value.is_null() + || value + .as_str() + .is_some_and(|value| value.trim().parse::().is_ok()); + if safe { + projected.insert("total_cost_usd".to_string(), value.clone()); + } + } + Value::Object(projected) +} + +fn optional_static_text(value: Option<&'static str>) -> Value { + value + .map(|value| Value::String(value.to_string())) + .unwrap_or(Value::Null) +} + +fn copy_json_number_or_null( + source: &Map, + target: &mut Map, + key: &str, +) { + if let Some(value) = source + .get(key) + .filter(|value| value.is_number() || value.is_null()) + { + target.insert(key.to_string(), value.clone()); + } +} + +fn copy_json_bool_or_null(source: &Map, target: &mut Map, key: &str) { + if let Some(value) = source + .get(key) + .filter(|value| value.is_boolean() || value.is_null()) + { + target.insert(key.to_string(), value.clone()); + } +} + +fn copy_safe_token_string_or_null( + source: &Map, + target: &mut Map, + key: &str, +) { + let Some(value) = source.get(key) else { + return; + }; + if value.is_null() { + target.insert(key.to_string(), Value::Null); + } else if let Some(value) = value.as_str().and_then(safe_status_token) { + target.insert(key.to_string(), Value::String(value.to_string())); + } +} + +fn copy_safe_display_string_or_null( + source: &Map, + target: &mut Map, + key: &str, +) { + let Some(value) = source.get(key) else { + return; + }; + if value.is_null() { + target.insert(key.to_string(), Value::Null); + } else if let Some(value) = value.as_str().and_then(safe_display_text) { + target.insert(key.to_string(), Value::String(value.to_string())); + } +} + +fn safe_status_token(value: &str) -> Option<&str> { + let value = value.trim(); + (!value.is_empty() + && value.len() <= 160 + && value.chars().all(|character| { + character.is_ascii_alphanumeric() + || matches!(character, '_' | '-' | ':' | '.' | '/' | '+') + })) + .then_some(value) +} + +fn safe_display_text(value: &str) -> Option<&str> { + let value = value.trim(); + if value.is_empty() || value.len() > 256 || value.chars().any(char::is_control) { + return None; + } + let lowered = value.to_ascii_lowercase(); + (![ + "authorization", + "bearer ", + "password", + "secret", + "token=", + "://", + ] + .iter() + .any(|marker| lowered.contains(marker))) + .then_some(value) +} + +fn redact_rule_array(rules: Option<&Value>, redact: fn(&Value) -> Value) -> Value { + let Some(rules) = rules.and_then(Value::as_array) else { + return Value::Null; + }; + Value::Array(rules.iter().map(redact).collect()) +} + +fn redact_json_value(value: &Value) -> Value { + match value { + Value::Array(values) => Value::Array(values.iter().map(redact_json_value).collect()), + Value::Object(object) => Value::Object( + object + .iter() + .map(|(key, value)| (key.clone(), redact_json_value_for_key(key, value))) + .collect(), + ), + _ => value.clone(), + } +} + +fn redact_proxy_json_value(value: &Value) -> Value { + match value { + Value::Array(values) => Value::Array(values.iter().map(redact_proxy_json_value).collect()), + Value::Object(object) => Value::Object( + object + .iter() + .map(|(key, value)| (key.clone(), redact_proxy_json_value_for_key(key, value))) + .collect(), + ), + _ => value.clone(), + } +} + +fn redact_proxy_json_value_for_key(key: &str, value: &Value) -> Value { + let compact_key = compact_json_key(key); + if proxy_credential_key(&compact_key) { + return redact_proxy_secret_value(value); + } + if json_url_key(&compact_key) { + return value + .as_str() + .and_then(sanitize_network_url) + .map(Value::String) + .unwrap_or(Value::Null); + } + if compact_key == "proxy" { + return admin_secret_safe_proxy(Some(value)); + } + if header_rules_key(&compact_key) { + return admin_secret_safe_header_rules(Some(value)); + } + if body_rules_key(&compact_key) { + return admin_secret_safe_body_rules(Some(value)); + } + if header_values_key(&compact_key) && value.is_object() { + return redact_header_values(value); + } + redact_proxy_json_value(value) +} + +fn redact_json_value_for_key(key: &str, value: &Value) -> Value { + let compact_key = compact_json_key(key); + if json_secret_key(&compact_key, value) { + return redact_secret_value(value); + } + if json_url_key(&compact_key) { + return value + .as_str() + .and_then(sanitize_network_url) + .map(Value::String) + .unwrap_or(Value::Null); + } + if compact_key == "proxy" { + return admin_secret_safe_proxy(Some(value)); + } + if header_rules_key(&compact_key) { + return admin_secret_safe_header_rules(Some(value)); + } + if body_rules_key(&compact_key) { + return admin_secret_safe_body_rules(Some(value)); + } + if header_values_key(&compact_key) && value.is_object() { + return redact_header_values(value); + } + redact_json_value(value) +} + +fn redact_body_rule(rule: &Value) -> Value { + let Some(rule) = rule.as_object() else { + return redact_json_value(rule); + }; + let mut projected = rule + .iter() + .filter(|(key, _)| !is_rule_marker(key)) + .map(|(key, value)| (key.clone(), redact_json_value_for_key(key, value))) + .collect::>(); + if let Some(condition) = rule.get("condition") { + projected.insert("condition".to_string(), redact_condition(condition)); + } + + let target_is_sensitive = rule + .get("path") + .and_then(Value::as_str) + .is_some_and(json_path_targets_secret); + let action = normalized_string_field(rule, "action"); + if target_is_sensitive && matches!(action.as_deref(), Some("set" | "append" | "insert")) { + redact_rule_secret_field(rule, &mut projected, "value", "has_value"); + } + if target_is_sensitive && action.as_deref() == Some("regex_replace") { + redact_rule_secret_field(rule, &mut projected, "pattern", "has_pattern"); + redact_rule_secret_field(rule, &mut projected, "replacement", "has_replacement"); + } + Value::Object(projected) +} + +fn redact_header_rule(rule: &Value) -> Value { + let Some(rule) = rule.as_object() else { + return redact_json_value(rule); + }; + let mut projected = rule + .iter() + .filter(|(key, _)| !is_rule_marker(key)) + .map(|(key, value)| (key.clone(), redact_json_value_for_key(key, value))) + .collect::>(); + let is_set = normalized_string_field(rule, "action").as_deref() == Some("set"); + if is_set { + redact_rule_secret_field(rule, &mut projected, "value", "has_value"); + } + if let Some(condition) = rule.get("condition") { + projected.insert("condition".to_string(), redact_condition(condition)); + } + Value::Object(projected) +} + +fn redact_rule_secret_field( + source: &Map, + projected: &mut Map, + field: &str, + marker: &str, +) { + let Some(value) = source.get(field).filter(|value| secret_value_is_set(value)) else { + return; + }; + projected.insert(field.to_string(), redact_secret_value(value)); + projected.insert(marker.to_string(), Value::Bool(true)); +} + +fn redact_condition(condition: &Value) -> Value { + let Some(condition) = condition.as_object() else { + return redact_json_value(condition); + }; + let mut projected = condition + .iter() + .filter(|(key, _)| !is_rule_marker(key)) + .map(|(key, value)| { + let value = if matches!(key.as_str(), "all" | "any") { + value + .as_array() + .map(|items| Value::Array(items.iter().map(redact_condition).collect())) + .unwrap_or_else(|| redact_json_value(value)) + } else { + redact_json_value_for_key(key, value) + }; + (key.clone(), value) + }) + .collect::>(); + if condition_value_is_secret(condition) { + redact_rule_secret_field(condition, &mut projected, "value", "has_value"); + } + Value::Object(projected) +} + +fn redact_header_values(value: &Value) -> Value { + let Some(headers) = value.as_object() else { + return Value::Null; + }; + Value::Object( + headers + .iter() + .map(|(key, value)| (key.clone(), redact_secret_value(value))) + .collect(), + ) +} + +fn restore_json_value(existing: Option<&Value>, incoming: &Value, key: Option<&str>) -> Value { + if let Some(key) = key { + let compact_key = compact_json_key(key); + if compact_key == "proxy" { + return admin_restore_secret_safe_proxy(existing, incoming); + } + if header_rules_key(&compact_key) { + return admin_restore_secret_safe_header_rules(existing, incoming); + } + if body_rules_key(&compact_key) { + return admin_restore_secret_safe_body_rules(existing, incoming); + } + if header_values_key(&compact_key) && incoming.is_object() { + return restore_header_values(existing, incoming); + } + if json_url_key(&compact_key) { + return incoming + .as_str() + .map(|incoming_url| restore_url_value(existing, incoming_url)) + .unwrap_or_else(|| incoming.clone()); + } + if json_secret_key(&compact_key, incoming) { + return restore_masked_secret(existing, incoming, true); + } + } + + match incoming { + Value::Array(incoming_values) => { + let existing_values = existing.and_then(Value::as_array); + Value::Array( + incoming_values + .iter() + .enumerate() + .map(|(index, incoming_value)| { + restore_json_value( + existing_values.and_then(|values| values.get(index)), + incoming_value, + None, + ) + }) + .collect(), + ) + } + Value::Object(incoming_object) => { + let existing_object = existing.and_then(Value::as_object); + Value::Object( + incoming_object + .iter() + .map(|(key, incoming_value)| { + ( + key.clone(), + restore_json_value( + existing_object.and_then(|object| object.get(key)), + incoming_value, + Some(key), + ), + ) + }) + .collect(), + ) + } + _ => incoming.clone(), + } +} + +fn restore_rule_array(existing: Option<&Value>, incoming: &Value, kind: RuleKind) -> Value { + let Some(incoming_values) = incoming.as_array() else { + return incoming.clone(); + }; + let existing_values = existing.and_then(Value::as_array); + let incoming_identities = identity_counts(incoming_values, kind); + let existing_identities = existing_values + .map(|values| identity_counts(values, kind)) + .unwrap_or_default(); + + Value::Array( + incoming_values + .iter() + .map(|incoming_rule| { + let identity = rule_identity(incoming_rule, kind); + let existing_rule = identity.as_ref().and_then(|identity| { + (incoming_identities.get(identity) == Some(&1) + && existing_identities.get(identity) == Some(&1)) + .then(|| { + existing_values.and_then(|values| { + values.iter().find(|candidate| { + rule_identity(candidate, kind).as_ref() == Some(identity) + }) + }) + }) + .flatten() + }); + restore_rule(existing_rule, incoming_rule, kind) + }) + .collect(), + ) +} + +fn restore_rule(existing: Option<&Value>, incoming: &Value, kind: RuleKind) -> Value { + let Some(incoming_object) = incoming.as_object() else { + return restore_json_value(existing, incoming, None); + }; + let existing_object = existing.and_then(Value::as_object); + let mut restored = Map::new(); + for (key, incoming_value) in incoming_object { + if is_rule_marker(key) { + continue; + } + let existing_value = existing_object.and_then(|object| object.get(key)); + let value = if key == "condition" { + restore_condition(existing_value, incoming_value) + } else if rule_field_is_secret(incoming_object, kind, key) { + let marker = marker_for_secret_field(key); + let marker_set = marker + .and_then(|marker| incoming_object.get(marker)) + .and_then(Value::as_bool) + == Some(true); + restore_masked_secret(existing_value, incoming_value, marker_set) + } else { + restore_json_value(existing_value, incoming_value, Some(key)) + }; + restored.insert(key.clone(), value); + } + Value::Object(restored) +} + +fn restore_condition(existing: Option<&Value>, incoming: &Value) -> Value { + let Some(incoming_object) = incoming.as_object() else { + return restore_json_value(existing, incoming, None); + }; + let existing_object = existing.and_then(Value::as_object); + + for group_key in ["all", "any"] { + if let Some(incoming_children) = incoming_object.get(group_key) { + let matching_existing = existing_object + .filter(|object| object.contains_key(group_key)) + .and_then(|object| object.get(group_key)); + let mut restored = Map::new(); + for (key, incoming_value) in incoming_object { + if is_rule_marker(key) { + continue; + } + let value = if key == group_key { + restore_condition_array(matching_existing, incoming_children) + } else { + restore_json_value( + existing_object.and_then(|object| object.get(key)), + incoming_value, + Some(key), + ) + }; + restored.insert(key.clone(), value); + } + return Value::Object(restored); + } + } + + let identities_match = condition_identity(incoming) + .zip(existing.and_then(condition_identity)) + .is_some_and(|(incoming, existing)| incoming == existing); + let existing_object = identities_match.then_some(existing_object).flatten(); + let mut restored = Map::new(); + for (key, incoming_value) in incoming_object { + if is_rule_marker(key) { + continue; + } + let existing_value = existing_object.and_then(|object| object.get(key)); + let value = if key == "value" && condition_value_is_secret(incoming_object) { + let marker_set = + incoming_object.get("has_value").and_then(Value::as_bool) == Some(true); + restore_masked_secret(existing_value, incoming_value, marker_set) + } else { + restore_json_value(existing_value, incoming_value, Some(key)) + }; + restored.insert(key.clone(), value); + } + Value::Object(restored) +} + +fn restore_condition_array(existing: Option<&Value>, incoming: &Value) -> Value { + let Some(incoming_values) = incoming.as_array() else { + return incoming.clone(); + }; + let existing_values = existing.and_then(Value::as_array); + let incoming_counts = condition_identity_counts(incoming_values); + let existing_counts = existing_values + .map(|values| condition_identity_counts(values)) + .unwrap_or_default(); + + Value::Array( + incoming_values + .iter() + .map(|incoming_condition| { + let identity = condition_identity(incoming_condition); + let existing_condition = identity.as_ref().and_then(|identity| { + (incoming_counts.get(identity) == Some(&1) + && existing_counts.get(identity) == Some(&1)) + .then(|| { + existing_values.and_then(|values| { + values.iter().find(|candidate| { + condition_identity(candidate).as_ref() == Some(identity) + }) + }) + }) + .flatten() + }); + restore_condition(existing_condition, incoming_condition) + }) + .collect(), + ) +} + +fn restore_header_values(existing: Option<&Value>, incoming: &Value) -> Value { + let Some(incoming_headers) = incoming.as_object() else { + return incoming.clone(); + }; + let existing_headers = existing.and_then(Value::as_object); + Value::Object( + incoming_headers + .iter() + .map(|(key, value)| { + ( + key.clone(), + restore_masked_secret( + existing_headers.and_then(|headers| headers.get(key)), + value, + true, + ), + ) + }) + .collect(), + ) +} + +fn restore_masked_secret(existing: Option<&Value>, incoming: &Value, allow_restore: bool) -> Value { + if allow_restore && incoming.as_str() == Some(REDACTED_VALUE) { + return existing + .filter(|value| secret_value_is_set(value)) + .cloned() + .unwrap_or_else(|| incoming.clone()); + } + incoming.clone() +} + +fn restore_url_value(existing: Option<&Value>, incoming_url: &str) -> Value { + if let Some(existing_url) = existing.and_then(Value::as_str) { + return Value::String(admin_restore_secret_safe_url( + Some(existing_url), + incoming_url, + )); + } + Value::String(incoming_url.to_string()) +} + +fn identity_counts(values: &[Value], kind: RuleKind) -> BTreeMap { + let mut counts = BTreeMap::new(); + for identity in values.iter().filter_map(|value| rule_identity(value, kind)) { + *counts.entry(identity).or_insert(0) += 1; + } + counts +} + +fn condition_identity_counts(values: &[Value]) -> BTreeMap { + let mut counts = BTreeMap::new(); + for identity in values.iter().filter_map(condition_identity) { + *counts.entry(identity).or_insert(0) += 1; + } + counts +} + +fn rule_identity(value: &Value, kind: RuleKind) -> Option { + let value = value.as_object()?; + let action = normalized_string_field(value, "action")?; + match (kind, action.as_str()) { + (RuleKind::Header, "set" | "drop") => normalized_string_field(value, "key") + .map(|key| format!("header:{action}:{}", key.to_ascii_lowercase())), + (RuleKind::Header, "rename") => normalized_string_field(value, "from") + .zip(normalized_string_field(value, "to")) + .map(|(from, to)| { + format!( + "header:rename:{}:{}", + from.to_ascii_lowercase(), + to.to_ascii_lowercase() + ) + }), + (RuleKind::Body, "set" | "drop" | "append" | "insert" | "regex_replace") => { + trimmed_string_field(value, "path").map(|path| format!("body:{action}:{path}")) + } + (RuleKind::Body, "rename") => trimmed_string_field(value, "from") + .zip(trimmed_string_field(value, "to")) + .map(|(from, to)| format!("body:rename:{from}:{to}")), + _ => None, + } +} + +fn condition_identity(value: &Value) -> Option { + let value = value.as_object()?; + if value.contains_key("all") || value.contains_key("any") { + return None; + } + let path = trimmed_string_field(value, "path")?; + let op = normalized_string_field(value, "op")?; + let source = normalized_condition_source(value.get("source").and_then(Value::as_str)); + Some(format!("condition:{source}:{path}:{op}")) +} + +fn rule_field_is_secret(rule: &Map, kind: RuleKind, field: &str) -> bool { + match kind { + RuleKind::Header => { + normalized_string_field(rule, "action").as_deref() == Some("set") && field == "value" + } + RuleKind::Body => { + let target_is_sensitive = rule + .get("path") + .and_then(Value::as_str) + .is_some_and(json_path_targets_secret); + let action = normalized_string_field(rule, "action"); + target_is_sensitive + && (matches!(action.as_deref(), Some("set" | "append" | "insert")) + && field == "value" + || action.as_deref() == Some("regex_replace") + && matches!(field, "pattern" | "replacement")) + } + } +} + +fn marker_for_secret_field(field: &str) -> Option<&'static str> { + match field { + "value" => Some("has_value"), + "pattern" => Some("has_pattern"), + "replacement" => Some("has_replacement"), + _ => None, + } +} + +fn is_rule_marker(key: &str) -> bool { + matches!(key, "has_value" | "has_pattern" | "has_replacement") +} + +fn condition_value_is_secret(condition: &Map) -> bool { + let source = normalized_condition_source(condition.get("source").and_then(Value::as_str)); + source == "request_headers" + || condition + .get("path") + .and_then(Value::as_str) + .is_some_and(json_path_targets_secret) +} + +fn normalized_condition_source(source: Option<&str>) -> String { + match source + .map(str::trim) + .map(str::to_ascii_lowercase) + .as_deref() + { + Some("headers" | "request_headers") => "request_headers".to_string(), + Some(source) if !source.is_empty() => source.to_string(), + _ => "body".to_string(), + } +} + +fn normalized_string_field(object: &Map, key: &str) -> Option { + trimmed_string_field(object, key).map(|value| value.to_ascii_lowercase()) +} + +fn trimmed_string_field(object: &Map, key: &str) -> Option { + object + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn sanitize_network_url(value: &str) -> Option { + let value = value.trim(); + let mut parsed = url::Url::parse(value).ok()?; + parsed.host_str()?; + parsed.set_username("").ok()?; + parsed.set_password(None).ok()?; + parsed.set_query(None); + parsed.set_fragment(None); + + let root_path_only = parsed.path() == "/"; + let mut sanitized = parsed.to_string(); + if root_path_only + && !value + .split(['?', '#']) + .next() + .unwrap_or(value) + .ends_with('/') + { + sanitized.pop(); + } + Some(sanitized) +} + +fn network_url_has_hidden_components(value: &str) -> bool { + let Ok(parsed) = url::Url::parse(value.trim()) else { + return false; + }; + !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() +} + +fn compact_json_key(key: &str) -> String { + key.chars() + .filter(|character| character.is_ascii_alphanumeric()) + .map(|character| character.to_ascii_lowercase()) + .collect() +} + +fn proxy_credential_key(compact_key: &str) -> bool { + proxy_credential_field_from_compact(compact_key).is_some() + || json_secret_key(compact_key, &Value::String(String::new())) +} + +fn proxy_credential_field_from_compact(compact_key: &str) -> Option { + match compact_key { + "user" | "username" | "proxyuser" | "proxyusername" => { + Some(AdminProxyCredentialField::Username) + } + "password" | "passwd" | "passphrase" | "proxypassword" | "proxypasswd" + | "proxypassphrase" => Some(AdminProxyCredentialField::Password), + _ => None, + } +} + +fn json_secret_key(compact_key: &str, value: &Value) -> bool { + if value.is_boolean() + && (compact_key.starts_with("has") + || compact_key.starts_with("is") + || compact_key.starts_with("can") + || compact_key.starts_with("supports")) + { + return false; + } + + matches!( + compact_key, + "authorization" + | "proxyauthorization" + | "bearer" + | "cookie" + | "cookies" + | "credential" + | "credentials" + | "hmac" + | "passphrase" + | "passwd" + | "password" + | "psk" + | "secret" + | "sessionkey" + | "token" + | "username" + ) || [ + "apikey", + "accesstoken", + "refreshtoken", + "idtoken", + "sessiontoken", + "bearertoken", + "secretkey", + "accesskey", + "privatekey", + "signingkey", + "clientsecret", + "secret", + "credential", + "credentials", + "authorization", + "password", + "passphrase", + "cookie", + ] + .iter() + .any(|suffix| compact_key.ends_with(suffix)) +} + +fn header_rules_key(compact_key: &str) -> bool { + matches!( + compact_key, + "headerrules" | "requestheaderrules" | "responseheaderrules" + ) +} + +fn body_rules_key(compact_key: &str) -> bool { + matches!(compact_key, "bodyrules" | "requestbodyrules") +} + +fn json_url_key(compact_key: &str) -> bool { + compact_key == "url" || compact_key.ends_with("url") +} + +fn json_path_targets_secret(path: &str) -> bool { + body_path_key_segments(path) + .into_iter() + .map(|segment| compact_json_key(&segment)) + .any(|segment| json_secret_key(&segment, &Value::String(String::new()))) +} + +fn body_path_key_segments(path: &str) -> Vec { + let chars = path.trim().chars().collect::>(); + let mut segments = Vec::new(); + let mut current = String::new(); + let mut index = 0; + while index < chars.len() { + match chars[index] { + '\\' if index + 1 < chars.len() => { + current.push(chars[index + 1]); + index += 2; + } + '.' => { + push_path_segment(&mut segments, &mut current); + index += 1; + } + '[' => { + push_path_segment(&mut segments, &mut current); + let mut close = index + 1; + while close < chars.len() && chars[close] != ']' { + close += 1; + } + if close >= chars.len() { + break; + } + let inner = chars[index + 1..close] + .iter() + .collect::() + .trim() + .trim_matches(['\'', '"']) + .to_string(); + if (!inner.is_empty() + && inner != "*" + && inner.parse::().is_err() + && !inner.contains('-')) + || inner.contains(|character: char| character.is_ascii_alphabetic()) + { + segments.push(inner); + } + index = close + 1; + } + character => { + current.push(character); + index += 1; + } + } + } + push_path_segment(&mut segments, &mut current); + segments +} + +fn push_path_segment(segments: &mut Vec, current: &mut String) { + let segment = std::mem::take(current).trim().to_string(); + if !segment.is_empty() { + segments.push(segment); + } +} + +fn header_values_key(compact_key: &str) -> bool { + matches!( + compact_key, + "headers" + | "extraheaders" + | "requestheaders" + | "responseheaders" + | "staticheaders" + | "defaultheaders" + ) +} + +fn redact_secret_value(value: &Value) -> Value { + if secret_value_is_set(value) { + Value::String(REDACTED_VALUE.to_string()) + } else { + value.clone() + } +} + +fn redact_proxy_secret_value(value: &Value) -> Value { + if proxy_secret_value_is_set(value) { + Value::String(REDACTED_VALUE.to_string()) + } else { + value.clone() + } +} + +fn proxy_secret_value_is_set(value: &Value) -> bool { + match value { + Value::Null => false, + Value::String(value) => !value.is_empty(), + Value::Array(value) => !value.is_empty(), + Value::Object(value) => !value.is_empty(), + _ => true, + } +} + +fn secret_value_is_set(value: &Value) -> bool { + match value { + Value::Null => false, + Value::String(value) => !value.trim().is_empty(), + Value::Array(value) => !value.is_empty(), + Value::Object(value) => !value.is_empty(), + _ => true, + } +} + +#[cfg(test)] +mod tests { + use super::{ + admin_provider_metadata_bucket_safe_json, admin_provider_oauth_invalid_reason_safe_text, + admin_provider_status_snapshot_safe_json, admin_provider_upstream_metadata_safe_json, + admin_restore_secret_safe_body_rules, admin_restore_secret_safe_header_rules, + admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, + admin_secret_safe_body_rules, admin_secret_safe_header_rules, admin_secret_safe_json, + admin_secret_safe_proxy, admin_secret_safe_url, + }; + use serde_json::json; + + #[test] + fn proxy_projection_removes_all_url_credentials_and_nested_secrets() { + let existing = json!({ + "enabled": true, + "mode": "direct", + "node_id": "proxy-node-1", + "url": "http://alice:proxy-password@proxy.example:8080/path?token=query-secret#fragment", + "username": "alice", + "password": "proxy-password", + "options": { + "clientSecret": "nested-secret", + "region": "us-east-1" + } + }); + let projected = admin_secret_safe_proxy(Some(&existing)); + + assert_eq!(projected["url"], "http://proxy.example:8080/path"); + assert_eq!(projected["username"], "***"); + assert_eq!(projected["password"], "***"); + assert_eq!(projected["options"]["clientSecret"], "***"); + assert_eq!(projected["options"]["region"], "us-east-1"); + assert_eq!(projected["has_credentials"], true); + let serialized = projected.to_string(); + for secret in [ + "alice", + "proxy-password", + "query-secret", + "nested-secret", + "fragment", + ] { + assert!(!serialized.contains(secret)); + } + + let restored = admin_restore_secret_safe_proxy(Some(&existing), &projected); + assert_eq!(restored, existing); + assert!(restored.get("has_credentials").is_none()); + } + + #[test] + fn proxy_restore_drops_old_credentials_when_authority_changes() { + let existing = json!({ + "url": "http://proxy-old.example:8080", + "username": "alice", + "password": "old-password", + "options": { + "clientSecret": "old-nested-secret", + "region": "us-east-1" + } + }); + let mut incoming = admin_secret_safe_proxy(Some(&existing)); + incoming["url"] = json!("http://proxy-new.example:8080"); + + let restored = admin_restore_secret_safe_proxy(Some(&existing), &incoming); + + assert_eq!(restored["url"], "http://proxy-new.example:8080"); + assert!(restored.get("username").is_none()); + assert!(restored.get("password").is_none()); + assert!(restored["options"].get("clientSecret").is_none()); + assert_eq!(restored["options"]["region"], "us-east-1"); + assert!(restored.get("has_credentials").is_none()); + } + + #[test] + fn proxy_restore_preserves_omitted_credentials_only_for_same_identity() { + let existing = json!({ + "url": "HTTP://Proxy.Example:80", + "node_id": " node-a ", + "username": "alice", + "password": "old-password" + }); + let same_identity = json!({ + "url": "http://proxy.example", + "node_id": "node-a" + }); + let restored = admin_restore_secret_safe_proxy(Some(&existing), &same_identity); + assert_eq!(restored["username"], "alice"); + assert_eq!(restored["password"], "old-password"); + + let changed_node = json!({ + "url": "http://proxy.example", + "node_id": "node-b", + "username": "***", + "password": "***" + }); + let restored = admin_restore_secret_safe_proxy(Some(&existing), &changed_node); + assert!(restored.get("username").is_none()); + assert!(restored.get("password").is_none()); + } + + #[test] + fn proxy_projection_redacts_supported_nested_aliases() { + let projected = admin_secret_safe_proxy(Some(&json!({ + "url": "socks5h://proxy.example:1080", + "proxy_auth": { + "proxy_user": "alice", + "proxy_passphrase": "proxy-password" + } + }))); + + assert_eq!(projected["proxy_auth"]["proxy_user"], "***"); + assert_eq!(projected["proxy_auth"]["proxy_passphrase"], "***"); + assert!(!projected.to_string().contains("proxy-password")); + } + + #[test] + fn proxy_projection_treats_whitespace_only_credentials_as_significant() { + let existing = json!({ + "url": "http://proxy.example:8080", + "username": " ", + "password": " " + }); + let projected = admin_secret_safe_proxy(Some(&existing)); + assert_eq!(projected["username"], "***"); + assert_eq!(projected["password"], "***"); + assert_eq!(projected["has_credentials"], true); + assert!(!projected.to_string().contains("\" \"")); + + let restored = admin_restore_secret_safe_proxy(Some(&existing), &projected); + assert_eq!(restored["username"], " "); + assert_eq!(restored["password"], " "); + } + + #[test] + fn header_rule_projection_hides_set_and_header_condition_values() { + let projected = admin_secret_safe_header_rules(Some(&json!([ + { + "action": "set", + "key": "x-custom-auth", + "value": "static-secret", + "condition": { + "source": "request_headers", + "path": "x-tenant-marker", + "op": "eq", + "value": "tenant-secret" + } + }, + {"action": "drop", "key": "x-debug"} + ]))); + + assert_eq!(projected[0]["value"], "***"); + assert_eq!(projected[0]["has_value"], true); + assert_eq!(projected[0]["condition"]["value"], "***"); + assert_eq!(projected[0]["condition"]["has_value"], true); + assert!(projected[1].get("value").is_none()); + let serialized = projected.to_string(); + assert!(!serialized.contains("static-secret")); + assert!(!serialized.contains("tenant-secret")); + } + + #[test] + fn generic_projection_recurses_through_config_headers_and_proxy() { + let projected = admin_secret_safe_json(Some(&json!({ + "pool_size": 3, + "auth": { + "access_token": "access-secret", + "token": "exact-token-secret", + "username": "nested-user", + "webhook_secret": "webhook-secret", + "has_refresh_token": true + }, + "headers": { + "x-region": "us-east-1", + "authorization": "Bearer nested-secret" + }, + "proxy": "socks5://user:password@proxy.example:1080?key=secret", + "response_header_rules": [ + {"action": "set", "key": "x-response-marker", "value": "response-secret"} + ] + }))); + + assert_eq!(projected["pool_size"], 3); + assert_eq!(projected["auth"]["access_token"], "***"); + assert_eq!(projected["auth"]["token"], "***"); + assert_eq!(projected["auth"]["username"], "***"); + assert_eq!(projected["auth"]["webhook_secret"], "***"); + assert_eq!(projected["auth"]["has_refresh_token"], true); + assert_eq!(projected["headers"]["x-region"], "***"); + assert_eq!(projected["proxy"], "socks5://proxy.example:1080"); + assert_eq!(projected["response_header_rules"][0]["value"], "***"); + let serialized = projected.to_string(); + for secret in [ + "access-secret", + "exact-token-secret", + "nested-user", + "webhook-secret", + "nested-secret", + "password", + "response-secret", + ] { + assert!(!serialized.contains(secret)); + } + } + + #[test] + fn provider_metadata_projection_removes_historical_upstream_diagnostics() { + let projected = admin_provider_upstream_metadata_safe_json(Some(&json!({ + "codex": { + "primary_used_percent": 25.0, + "message": "Authorization: Bearer upstream-secret", + "reset_credits": { + "available_count": 2, + "detail_error": "https://user:password@internal.test?q=secret" + } + }, + "kiro": { + "is_banned": true, + "ban_reason": "Authorization: Bearer upstream-secret" + }, + "windsurf": { + "daily_remaining_percent": 70.0, + "last_error": "https://user:password@internal.test?q=secret", + "rate_limit": { + "limited": true, + "message": "Authorization: Bearer upstream-secret" + }, + "probe_warnings": [{ + "probe": "models", + "message": "Authorization: Bearer upstream-secret" + }] + }, + "chatgpt_web": { + "updated_at": 1_777_000_000, + "image_quota_remaining": 8, + "blocked_features": [ + "image_generation", + "Authorization: Bearer upstream-secret" + ], + "limits_progress": [{ + "message": "Authorization: Bearer upstream-secret", + "url": "https://user:password@internal.test?q=secret" + }] + } + }))); + + assert_eq!( + projected.pointer("/codex/primary_used_percent"), + Some(&json!(25.0)) + ); + assert_eq!( + projected.pointer("/windsurf/daily_remaining_percent"), + Some(&json!(70.0)) + ); + assert_eq!( + projected.pointer("/chatgpt_web/image_quota_remaining"), + Some(&json!(8)) + ); + assert_eq!( + projected.pointer("/chatgpt_web/blocked_features"), + Some(&json!(["image_generation"])) + ); + assert!(projected.pointer("/chatgpt_web/limits_progress").is_none()); + let serialized = projected.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); + assert!(!serialized.contains("q=secret")); + } + + #[test] + fn chatgpt_metadata_projection_keeps_bounded_quota_dedup_key() { + let projected = admin_provider_metadata_bucket_safe_json( + "chatgpt_web", + Some(&json!({ + "image_quota_last_local_request_key": "request-1:candidate-1" + })), + ); + assert_eq!( + projected["image_quota_last_local_request_key"], + "request-1:candidate-1" + ); + + let oversized = "x".repeat(161); + for rejected in [ + "Authorization: Bearer secret", + "request\nforged", + oversized.as_str(), + ] { + let projected = admin_provider_metadata_bucket_safe_json( + "chatgpt_web", + Some(&json!({"image_quota_last_local_request_key": rejected})), + ); + assert!( + projected + .get("image_quota_last_local_request_key") + .is_none(), + "unsafe quota dedup key must be dropped: {rejected:?}" + ); + } + } + + #[test] + fn provider_status_projection_rebuilds_diagnostic_text_from_codes() { + let projected = admin_provider_status_snapshot_safe_json(Some(&json!({ + "oauth": { + "code": "invalid", + "label": "Authorization: Bearer upstream-secret", + "reason": "https://user:password@internal.test?q=secret", + "invalid_at": 1_777_000_000, + "requires_reauth": true, + "unknown": {"body": "upstream-secret"} + }, + "account": { + "code": "account_disabled", + "label": "Authorization: Bearer upstream-secret", + "reason": "https://user:password@internal.test?q=secret", + "blocked": true, + "source": "metadata" + }, + "quota": { + "provider_type": "codex", + "code": "cooldown", + "reason": "Authorization: Bearer upstream-secret", + "exhausted": false, + "rate_limit": { + "limited": true, + "retry_after_ms": 5000, + "message": "https://user:password@internal.test?q=secret" + }, + "reset_credits": { + "available_count": 1, + "detail_error": "connect https://user:password@internal.test?q=secret" + }, + "unknown": {"body": "upstream-secret"} + } + }))); + + assert_eq!( + projected.pointer("/oauth/reason"), + Some(&json!("OAuth token is invalid")) + ); + assert_eq!( + projected.pointer("/account/reason"), + Some(&json!("Account is disabled")) + ); + assert_eq!( + projected.pointer("/quota/reason"), + Some(&json!("Quota is temporarily cooling down")) + ); + assert_eq!( + projected.pointer("/quota/reset_credits/detail_error"), + Some(&json!("Quota refresh connection failed")) + ); + assert_eq!( + projected.pointer("/quota/rate_limit/retry_after_ms"), + Some(&json!(5000)) + ); + assert!(projected.pointer("/quota/rate_limit/message").is_none()); + assert!(projected.pointer("/quota/unknown").is_none()); + let serialized = projected.to_string(); + assert!(!serialized.contains("upstream-secret")); + assert!(!serialized.contains("user:password")); + } + + #[test] + fn oauth_invalid_reason_projection_preserves_only_fixed_semantics() { + let projected = admin_provider_oauth_invalid_reason_safe_text(Some( + "[ACCOUNT_BLOCK] account has been deactivated: Authorization: Bearer upstream-secret https://user:password@internal.test?q=secret", + )) + .expect("reason should be projected"); + + assert_eq!(projected, "[ACCOUNT_BLOCK] Account is disabled"); + assert!(!projected.contains("upstream-secret")); + assert!(!projected.contains("user:password")); + + let bucket = admin_provider_metadata_bucket_safe_json( + "windsurf", + Some(&json!({ + "daily_remaining_percent": 50, + "last_error": "Authorization: Bearer upstream-secret" + })), + ); + assert_eq!(bucket["daily_remaining_percent"], json!(50)); + assert!(!bucket.to_string().contains("upstream-secret")); + } + + #[test] + fn body_rule_projection_detects_compound_and_bracketed_secret_paths() { + let projected = admin_secret_safe_body_rules(Some(&json!([ + { + "action": "set", + "path": "auth.api_key", + "value": "body-token-secret", + "condition": { + "source": "original", + "path": "auth['private-key']", + "op": "eq", + "value": "condition-secret" + } + }, + {"action": "set", "path": "metadata.region", "value": "us-east-1"} + ]))); + + assert_eq!(projected[0]["value"], "***"); + assert_eq!(projected[0]["has_value"], true); + assert_eq!(projected[0]["condition"]["value"], "***"); + assert_eq!(projected[0]["condition"]["has_value"], true); + assert_eq!(projected[1]["value"], "us-east-1"); + let serialized = projected.to_string(); + assert!(!serialized.contains("body-token-secret")); + assert!(!serialized.contains("condition-secret")); + } + + #[test] + fn restore_projection_matches_unique_reordered_rules_and_removes_markers() { + let existing = json!([ + {"action": "set", "key": "x-first", "value": "first-secret"}, + {"action": "set", "key": "x-second", "value": "second-secret"} + ]); + let incoming = json!([ + {"action": "set", "key": "x-second", "value": "***", "has_value": true}, + {"action": "set", "key": "x-first", "value": "replacement", "has_value": false} + ]); + + let restored = admin_restore_secret_safe_header_rules(Some(&existing), &incoming); + + assert_eq!(restored[0]["value"], "second-secret"); + assert!(restored[0].get("has_value").is_none()); + assert_eq!(restored[1]["value"], "replacement"); + assert!(restored[1].get("has_value").is_none()); + } + + #[test] + fn restore_does_not_move_a_secret_when_rule_identity_changes() { + let existing = json!([ + {"action": "set", "path": "auth.token", "value": "secret"} + ]); + let incoming = json!([ + {"action": "set", "path": "metadata.note", "value": "***", "has_value": true} + ]); + + let restored = admin_restore_secret_safe_body_rules(Some(&existing), &incoming); + + assert_eq!(restored[0]["value"], "***"); + assert!(restored[0].get("has_value").is_none()); + assert!(!restored.to_string().contains("secret")); + } + + #[test] + fn restore_does_not_move_a_secret_when_condition_semantics_change() { + let existing = json!([{ + "action": "set", + "key": "x-output", + "value": "header-secret", + "condition": { + "source": "request_headers", + "path": "x-tenant", + "op": "eq", + "value": "condition-secret" + } + }]); + let incoming = json!([{ + "action": "set", + "key": "x-output", + "value": "***", + "has_value": true, + "condition": { + "source": "body", + "path": "metadata.note", + "op": "eq", + "value": "***", + "has_value": true + } + }]); + + let restored = admin_restore_secret_safe_header_rules(Some(&existing), &incoming); + + assert_eq!(restored[0]["value"], "header-secret"); + assert_eq!(restored[0]["condition"]["value"], "***"); + assert!(!restored.to_string().contains("condition-secret")); + } + + #[test] + fn restore_requires_rule_marker_but_preserves_generic_sensitive_fields() { + let header_existing = json!([ + {"action": "set", "key": "x-secret", "value": "old"} + ]); + let header_incoming = json!([ + {"action": "set", "key": "x-secret", "value": "***"} + ]); + assert_eq!( + admin_restore_secret_safe_header_rules(Some(&header_existing), &header_incoming)[0] + ["value"], + "***" + ); + + let existing = json!({"note": "old note", "access_token": "old token"}); + let incoming = json!({"note": "***", "access_token": "***"}); + let restored = admin_restore_secret_safe_json(Some(&existing), &incoming); + assert_eq!(restored["note"], "***"); + assert_eq!(restored["access_token"], "old token"); + } + + #[test] + fn regex_markers_restore_only_the_same_sensitive_rule_and_are_not_persisted() { + let existing = json!([{ + "action": "regex_replace", + "path": "auth.token", + "pattern": "old-pattern", + "replacement": "old-replacement" + }]); + let projected = admin_secret_safe_body_rules(Some(&existing)); + let restored = admin_restore_secret_safe_body_rules(Some(&existing), &projected); + + assert_eq!(restored, existing); + assert!(restored[0].get("has_pattern").is_none()); + assert!(restored[0].get("has_replacement").is_none()); + } + + #[test] + fn empty_or_null_rule_values_do_not_gain_secret_markers() { + let projected = admin_secret_safe_header_rules(Some(&json!([ + {"action": "set", "key": "x-empty", "value": ""}, + {"action": "set", "key": "x-null", "value": null} + ]))); + + assert_eq!(projected[0]["value"], ""); + assert_eq!(projected[1]["value"], json!(null)); + assert!(projected[0].get("has_value").is_none()); + assert!(projected[1].get("has_value").is_none()); + } + + #[test] + fn generic_restore_preserves_hidden_url_parts_and_response_rule_secrets() { + let existing = json!({ + "callback_url": "https://alice:password@example.test/callback?token=secret#fragment", + "response_header_rules": [ + {"action": "set", "key": "authorization", "value": "Bearer secret"} + ] + }); + let projected = admin_secret_safe_json(Some(&existing)); + let restored = admin_restore_secret_safe_json(Some(&existing), &projected); + + assert_eq!(restored, existing); + } + + #[test] + fn url_projection_removes_userinfo_query_and_fragment() { + let projected = admin_secret_safe_url(Some( + "https://alice:password@api.example/v1?token=query-secret#fragment", + )); + + assert_eq!(projected, "https://api.example/v1"); + assert_eq!(admin_secret_safe_url(Some("not a url")), json!(null)); + } +} diff --git a/crates/aether-admin/src/provider/state.rs b/crates/aether-admin/src/provider/state.rs index e3450287e..ade0fd45f 100644 --- a/crates/aether-admin/src/provider/state.rs +++ b/crates/aether-admin/src/provider/state.rs @@ -8,6 +8,7 @@ use uuid::Uuid; const KIRO_DEVICE_DEFAULT_START_URL: &str = "https://view.awsapps.com/start"; const KIRO_DEVICE_DEFAULT_REGION: &str = "us-east-1"; +const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024; pub fn current_unix_secs() -> u64 { SystemTime::now() @@ -112,7 +113,18 @@ pub fn json_u64_value(value: Option<&Value>) -> Option { pub fn decode_jwt_claims(token: &str) -> Option> { 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() @@ -398,7 +410,10 @@ pub fn build_kiro_device_key_name(email: Option<&str>, refresh_token: Option<&st #[cfg(test)] mod tests { - use super::{enrich_admin_provider_oauth_auth_config, parse_provider_oauth_callback_params}; + use super::{ + decode_jwt_claims, enrich_admin_provider_oauth_auth_config, + parse_provider_oauth_callback_params, MAX_UNVERIFIED_JWT_CLAIMS_BYTES, + }; use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; use serde_json::json; @@ -516,4 +531,16 @@ mod tests { assert_eq!(auth_config.get("user_id"), Some(&json!("user-image"))); assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true))); } + + #[test] + fn decode_jwt_claims_rejects_oversized_payload_before_decode() { + let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap() + .saturating_mul(4); + let token = format!("header.{}.signature", "A".repeat(max_encoded_len + 1)); + + assert_eq!(decode_jwt_claims(&token), None); + } } diff --git a/crates/aether-admin/src/provider/status.rs b/crates/aether-admin/src/provider/status.rs index 598298888..a2981a4a0 100644 --- a/crates/aether-admin/src/provider/status.rs +++ b/crates/aether-admin/src/provider/status.rs @@ -24,6 +24,8 @@ const ACCOUNT_BLOCK_REASON_KEYWORDS: &[&str] = &[ "deactivated", "访问被禁止", "账户访问被禁止", + "账号已停用", + "账户已停用", "访问受限", "账户访问受限", "oauth_token_invalid", @@ -229,6 +231,8 @@ fn classify_block_reason(reason: &str) -> (&'static str, &'static str) { "organization_disabled", "访问被禁止", "账户访问被禁止", + "账号已停用", + "账户已停用", ] .iter() .any(|keyword| lowered.contains(keyword)) @@ -868,4 +872,26 @@ mod tests { assert!(!should_auto_remove_account_state(&state)); assert!(account_state_indicates_known_ban(&state)); } + + #[test] + fn canonical_chinese_account_disabled_state_is_auto_removed() { + for reason in [ + "[ACCOUNT_BLOCK] OpenAI 账号已停用", + "[ACCOUNT_BLOCK] OpenAI 账户已停用", + ] { + let state = resolve_pool_account_state(Some("codex"), None, Some(reason)); + + assert!(state.blocked); + assert_eq!(state.code.as_deref(), Some("account_disabled")); + assert!(should_auto_remove_account_state(&state)); + } + + let verification = resolve_pool_account_state( + Some("codex"), + None, + Some("[ACCOUNT_BLOCK] verify your account before continuing"), + ); + assert_eq!(verification.code.as_deref(), Some("account_verification")); + assert!(!should_auto_remove_account_state(&verification)); + } } diff --git a/crates/aether-admin/src/system.rs b/crates/aether-admin/src/system.rs index 83a46de99..9cafb2975 100644 --- a/crates/aether-admin/src/system.rs +++ b/crates/aether-admin/src/system.rs @@ -21,6 +21,7 @@ use semver::Version; use serde::{de, de::DeserializeOwned, Deserialize, Serialize}; use serde_json::{json, Map, Value}; use std::collections::BTreeSet; +use url::Url; #[derive(Debug, Clone)] pub struct AdminSystemSettingsUpdate { @@ -43,12 +44,43 @@ pub struct AdminEmailTemplateUpdate { pub html: Option, } +/// Email templates are administrator-configured, but they are rendered on +/// authentication and notification paths. Keep their resource and protocol +/// surface bounded before values reach storage or an SMTP header/body. +pub const ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES: usize = 512; +pub const ADMIN_EMAIL_TEMPLATE_MAX_HTML_BYTES: usize = 256 * 1024; +pub const ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_BYTES: usize = 512 * 1024; +pub const ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES: usize = 64 * 1024; + +pub fn admin_email_template_subject_is_valid(value: &str) -> bool { + value.len() <= ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES + && !value.bytes().any(|byte| byte < 0x20 || byte == 0x7f) +} + +pub fn admin_email_template_html_is_valid(value: &str) -> bool { + value.len() <= ADMIN_EMAIL_TEMPLATE_MAX_HTML_BYTES + && !value.bytes().any(|byte| { + byte == 0 || byte == 0x7f || (byte < 0x20 && !matches!(byte, b'\r' | b'\n' | b'\t')) + }) +} + pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.3"; pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] = &["2.0", "2.1", "2.2", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION]; -pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.5"; +pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str = "execution_extra_trusted_dns_hosts"; +pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES: usize = 128; +pub const EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES: usize = 253; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ExecutionExtraTrustedDnsHostsConfigError { + InvalidValue, + TooManyEntries, + InvalidHost, +} + +pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.6"; pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] = - &["1.3", "1.4", ADMIN_SYSTEM_USERS_EXPORT_VERSION]; + &["1.3", "1.4", "1.5", ADMIN_SYSTEM_USERS_EXPORT_VERSION]; pub const ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS: &[&str] = &[ "api_key", "password", @@ -383,7 +415,7 @@ pub struct AdminSystemConfigGlobalModel { pub is_active: bool, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigEndpoint { pub api_format: String, pub base_url: String, @@ -405,13 +437,40 @@ pub struct AdminSystemConfigEndpoint { pub proxy: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigEndpoint { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigEndpoint") + .field("api_format", &self.api_format) + .field("base_url", &self.base_url) + .field( + "header_rules", + &self.header_rules.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "body_rules", + &self.body_rules.as_ref().map(|_| "[REDACTED]"), + ) + .field("max_retries", &self.max_retries) + .field("is_active", &self.is_active) + .field("custom_path", &self.custom_path) + .field("config", &self.config.as_ref().map(|_| "[REDACTED]")) + .field( + "format_acceptance_config", + &self.format_acceptance_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("proxy", &self.proxy.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigProviderKey { - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub api_key: Option, #[serde(default)] pub auth_type: Option, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub auth_config: Option, #[serde(default)] pub name: Option, @@ -455,6 +514,27 @@ pub struct AdminSystemConfigProviderKey { pub proxy: Option, #[serde(default)] pub fingerprint: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub credential_state: Option, +} + +impl std::fmt::Debug for AdminSystemConfigProviderKey { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigProviderKey") + .field("api_key", &self.api_key.as_ref().map(|_| "[REDACTED]")) + .field("auth_type", &self.auth_type) + .field( + "auth_config", + &self.auth_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("name", &self.name) + .field("note", &self.note) + .field("api_formats", &self.api_formats) + .field("is_active", &self.is_active) + .field("credential_state", &self.credential_state) + .finish_non_exhaustive() + } } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] @@ -484,7 +564,7 @@ pub struct AdminSystemConfigProviderModel { pub config: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigProvider { pub name: String, #[serde(default)] @@ -527,7 +607,24 @@ pub struct AdminSystemConfigProvider { pub models: Vec, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigProvider { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigProvider") + .field("name", &self.name) + .field("provider_type", &self.provider_type) + .field("billing_type", &self.billing_type) + .field("is_active", &self.is_active) + .field("proxy", &self.proxy.as_ref().map(|_| "[REDACTED]")) + .field("config", &self.config.as_ref().map(|_| "[REDACTED]")) + .field("endpoints", &self.endpoints) + .field("api_keys", &self.api_keys) + .field("models", &self.models.len()) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigProxyNode { #[serde(default)] pub id: Option, @@ -543,9 +640,9 @@ pub struct AdminSystemConfigProxyNode { pub is_manual: Option, #[serde(default)] pub proxy_url: Option, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub proxy_username: Option, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub proxy_password: Option, #[serde(default)] pub tunnel_mode: Option, @@ -557,11 +654,38 @@ pub struct AdminSystemConfigProxyNode { pub config_version: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigProxyNode { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigProxyNode") + .field("id", &self.id) + .field("name", &self.name) + .field("ip", &self.ip) + .field("port", &self.port) + .field("region", &self.region) + .field("is_manual", &self.is_manual) + .field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]")) + .field("proxy_username", &self.proxy_username) + .field( + "proxy_password", + &self.proxy_password.as_ref().map(|_| "[REDACTED]"), + ) + .field("tunnel_mode", &self.tunnel_mode) + .field("heartbeat_interval", &self.heartbeat_interval) + .field( + "remote_config", + &self.remote_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("config_version", &self.config_version) + .finish() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigLdap { pub server_url: String, pub bind_dn: String, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub bind_password: Option, pub base_dn: String, #[serde(default)] @@ -582,12 +706,31 @@ pub struct AdminSystemConfigLdap { pub connect_timeout: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigLdap { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigLdap") + .field("server_url", &self.server_url) + .field("bind_dn", &self.bind_dn) + .field( + "bind_password", + &self.bind_password.as_ref().map(|_| "[REDACTED]"), + ) + .field("base_dn", &self.base_dn) + .field("is_enabled", &self.is_enabled) + .field("is_exclusive", &self.is_exclusive) + .field("use_starttls", &self.use_starttls) + .field("connect_timeout", &self.connect_timeout) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigOAuthProvider { pub provider_type: String, pub display_name: String, pub client_id: String, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub client_secret: Option, #[serde(default)] pub authorization_url_override: Option, @@ -607,7 +750,25 @@ pub struct AdminSystemConfigOAuthProvider { pub is_enabled: bool, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigOAuthProvider { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigOAuthProvider") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("client_id", &self.client_id) + .field( + "client_secret", + &self.client_secret.as_ref().map(|_| "[REDACTED]"), + ) + .field("redirect_uri", &self.redirect_uri) + .field("frontend_callback_url", &self.frontend_callback_url) + .field("is_enabled", &self.is_enabled) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigEntry { pub key: String, #[serde(default)] @@ -616,11 +777,26 @@ pub struct AdminSystemConfigEntry { pub description: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +impl std::fmt::Debug for AdminSystemConfigEntry { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let mut debug = formatter.debug_struct("AdminSystemConfigEntry"); + debug.field("key", &self.key); + if is_sensitive_admin_system_config_key(&self.key) { + debug.field("value", &"[REDACTED]"); + } else { + debug.field("value", &self.value); + } + debug.field("description", &self.description).finish() + } +} + +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigDocument { pub version: String, #[serde(default)] pub exported_at: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub credential_state: Option, #[serde(default)] pub global_models: Vec, #[serde(default)] @@ -635,6 +811,23 @@ pub struct AdminSystemConfigDocument { pub system_configs: Vec, } +impl std::fmt::Debug for AdminSystemConfigDocument { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminSystemConfigDocument") + .field("version", &self.version) + .field("exported_at", &self.exported_at) + .field("credential_state", &self.credential_state) + .field("global_models", &self.global_models.len()) + .field("providers", &self.providers) + .field("proxy_nodes", &self.proxy_nodes) + .field("ldap_config", &self.ldap_config) + .field("oauth_providers", &self.oauth_providers) + .field("system_configs", &self.system_configs) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct AdminSystemConfigImportRequest { #[serde(flatten)] @@ -643,18 +836,38 @@ pub struct AdminSystemConfigImportRequest { pub merge_mode: AdminImportMergeMode, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct ParsedAdminSystemConfigImportRequest { pub request: AdminSystemConfigImportRequest, pub root: Map, } -#[derive(Debug, Clone)] +impl std::fmt::Debug for ParsedAdminSystemConfigImportRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ParsedAdminSystemConfigImportRequest") + .field("request", &self.request) + .field("root", &"[REDACTED]") + .finish() + } +} + +#[derive(Clone)] pub struct ParsedAdminSystemConfigObject { pub raw: Map, pub value: T, } +impl std::fmt::Debug for ParsedAdminSystemConfigObject { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ParsedAdminSystemConfigObject") + .field("raw", &"[REDACTED]") + .field("value", &self.value) + .finish() + } +} + impl ParsedAdminSystemConfigObject { pub fn into_parts(self) -> (Map, T) { (self.raw, self.value) @@ -1172,6 +1385,12 @@ pub fn admin_email_template_not_found_error( pub fn parse_admin_email_template_update( request_body: &[u8], ) -> Result { + if request_body.len() > ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_BYTES { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板请求体超过允许大小" }), + )); + } let payload = match serde_json::from_slice::(request_body) { Ok(serde_json::Value::Object(payload)) => payload, _ => { @@ -1183,7 +1402,15 @@ pub fn parse_admin_email_template_update( }; let subject = match payload.get("subject") { - Some(serde_json::Value::String(value)) => Some(value.clone()), + Some(serde_json::Value::String(value)) if admin_email_template_subject_is_valid(value) => { + Some(value.clone()) + } + Some(serde_json::Value::String(_)) => { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板 subject 超过大小限制或包含非法控制字符" }), + )); + } Some(serde_json::Value::Null) | None => None, Some(_) => { return Err(( @@ -1193,7 +1420,15 @@ pub fn parse_admin_email_template_update( } }; let html = match payload.get("html") { - Some(serde_json::Value::String(value)) => Some(value.clone()), + Some(serde_json::Value::String(value)) if admin_email_template_html_is_valid(value) => { + Some(value.clone()) + } + Some(serde_json::Value::String(_)) => { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板 html 超过大小限制或包含非法控制字符" }), + )); + } Some(serde_json::Value::Null) | None => None, Some(_) => { return Err(( @@ -1217,8 +1452,40 @@ pub fn parse_admin_email_template_preview_payload( request_body: Option<&[u8]>, ) -> Result, (http::StatusCode, serde_json::Value)> { match request_body { + Some(bytes) if bytes.len() > ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_BYTES => Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板预览请求体超过允许大小" }), + )), Some(bytes) => match serde_json::from_slice::(bytes) { - Ok(serde_json::Value::Object(payload)) => Ok(payload), + Ok(serde_json::Value::Object(payload)) => { + if let Some(serde_json::Value::String(value)) = payload.get("html") { + if !admin_email_template_html_is_valid(value) { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板 html 超过大小限制或包含非法控制字符" }), + )); + } + } + if let Some(serde_json::Value::String(value)) = payload.get("subject") { + if !admin_email_template_subject_is_valid(value) { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板 subject 超过大小限制或包含非法控制字符" }), + )); + } + } + if payload.values().any(|value| { + value.as_str().is_some_and(|value| { + value.len() > ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES + }) + }) { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "模板预览变量超过允许大小" }), + )); + } + Ok(payload) + } Ok(serde_json::Value::Null) => Ok(serde_json::Map::new()), _ => Err(( http::StatusCode::BAD_REQUEST, @@ -1309,9 +1576,46 @@ pub fn ldap_module_config_is_valid(config: Option<&StoredLdapModuleConfig>) -> b let Some(config) = config else { return false; }; - !config.server_url.trim().is_empty() - && !config.bind_dn.trim().is_empty() - && !config.base_dn.trim().is_empty() + normalize_ldap_transport_server_url(&config.server_url, config.use_starttls).is_some() + && ldap_module_config_fields_are_valid(config) +} + +/// Validate LDAP configuration fields other than the transport endpoint. +/// +/// Keeping this separate lets gateway test builds use their in-process +/// `mockldap` transport while sharing every production data-shape check. +pub fn ldap_module_config_fields_are_valid(config: &StoredLdapModuleConfig) -> bool { + let search_filter = config + .user_search_filter + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("(uid={username})"); + let username_attr = config + .username_attr + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("uid"); + let email_attr = config + .email_attr + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("mail"); + let display_name_attr = config + .display_name_attr + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("cn"); + + ldap_distinguished_name_is_valid(&config.bind_dn) + && ldap_distinguished_name_is_valid(&config.base_dn) + && ldap_search_filter_is_valid(search_filter) + && ldap_attribute_description_is_valid(username_attr) + && ldap_attribute_description_is_valid(email_attr) + && ldap_attribute_description_is_valid(display_name_attr) && config .bind_password_encrypted .as_deref() @@ -1320,6 +1624,198 @@ pub fn ldap_module_config_is_valid(config: Option<&StoredLdapModuleConfig>) -> b .is_some() } +/// Validate the bounded LDAP user filter shape used by the gateway login +/// implementation. This pure helper is shared by admin status/validation +/// code so an imported malformed filter cannot be reported as active. +pub fn ldap_search_filter_is_valid(value: &str) -> bool { + if value.chars().any(char::is_control) { + return false; + } + let value = value.trim(); + if value.is_empty() + || value.len() > 200 + || !value.contains("{username}") + || !value.starts_with('(') + || !value.ends_with(')') + { + return false; + } + + let mut depth = 0i32; + let mut max_depth = 0i32; + let mut chars = value.chars().peekable(); + while let Some(ch) = chars.next() { + match ch { + '(' => { + depth += 1; + max_depth = max_depth.max(depth); + } + ')' => { + depth -= 1; + if depth < 0 { + return false; + } + // A valid LDAP filter has one outer pair. Reaching depth zero + // before EOF would allow a second top-level expression. + if depth == 0 && chars.peek().is_some() { + return false; + } + } + _ => {} + } + } + depth == 0 && max_depth <= 5 +} + +/// Validate a DN as bounded, protocol-safe configuration data. +/// +/// DN strings are encoded as their own LDAP protocol fields, so reparsing the +/// full RFC 4514 grammar here would add compatibility risk without preventing +/// filter injection. Raw control characters and unbounded values are the +/// relevant configuration hazards at this boundary. +pub fn ldap_distinguished_name_is_valid(value: &str) -> bool { + if value.chars().any(char::is_control) { + return false; + } + let value = value.trim(); + !value.is_empty() && value.len() <= 4096 +} + +/// Validate an RFC 4512 AttributeDescription used both for requested LDAP +/// attributes and, for the default username lookup, as filter syntax. +pub fn ldap_attribute_description_is_valid(value: &str) -> bool { + if value.chars().any(char::is_control) { + return false; + } + let value = value.trim(); + if value.is_empty() || value.len() > 128 || !value.is_ascii() { + return false; + } + + let mut parts = value.split(';'); + let Some(attribute_type) = parts.next() else { + return false; + }; + if !ldap_attribute_type_is_valid(attribute_type) { + return false; + } + parts.all(|option| { + !option.is_empty() + && option + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') + }) +} + +fn ldap_attribute_type_is_valid(value: &str) -> bool { + let mut bytes = value.bytes(); + if value + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphabetic) + { + return bytes.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-'); + } + + let mut components = value.split('.'); + let Some(first) = components.next() else { + return false; + }; + let Some(second) = components.next() else { + return false; + }; + ldap_oid_component_is_valid(first) + && ldap_oid_component_is_valid(second) + && components.all(ldap_oid_component_is_valid) +} + +fn ldap_oid_component_is_valid(value: &str) -> bool { + !value.is_empty() + && value.bytes().all(|byte| byte.is_ascii_digit()) + && (value.len() == 1 || !value.starts_with('0')) +} + +/// Parse and canonicalize an LDAP transport endpoint. +/// +/// LDAP simple binds carry credentials, therefore plaintext `ldap://` is only +/// accepted when StartTLS is explicitly enabled. This helper validates the +/// transport URL shape and removes URL-controlled request data (credentials, +/// query strings, fragments, and paths). It intentionally does not apply an +/// IP allow/deny policy: LDAP servers are commonly deployed on private +/// networks, and network reachability policy belongs to the outbound +/// transport layer rather than this pure configuration validator. +pub fn normalize_ldap_transport_server_url(raw: &str, use_starttls: bool) -> Option { + normalize_ldap_transport_server_url_inner(raw, use_starttls, false) +} + +/// Test-only variant used by gateway fixtures that emulate an LDAP endpoint +/// with the `mockldap://` scheme. Production callers must use +/// [`normalize_ldap_transport_server_url`]. +#[doc(hidden)] +pub fn normalize_ldap_transport_server_url_for_tests( + raw: &str, + use_starttls: bool, +) -> Option { + normalize_ldap_transport_server_url_inner(raw, use_starttls, true) +} + +fn normalize_ldap_transport_server_url_inner( + raw: &str, + use_starttls: bool, + allow_mockldap: bool, +) -> Option { + if raw.chars().any(char::is_control) { + return None; + } + let raw = raw.trim(); + if raw.is_empty() || raw.contains('@') { + return None; + } + let candidate = if raw.contains("://") { + raw.to_string() + } else { + format!("ldap://{raw}") + }; + let Ok(mut parsed) = Url::parse(&candidate) else { + return None; + }; + let scheme = parsed.scheme().to_ascii_lowercase(); + let secure_transport = match scheme.as_str() { + "ldaps" => true, + "ldap" => use_starttls, + "mockldap" => allow_mockldap, + _ => false, + }; + if !secure_transport + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + || (parsed.path() != "" && parsed.path() != "/") + { + return None; + } + let host = parsed + .host_str()? + .trim_end_matches('.') + .to_ascii_lowercase(); + if host.is_empty() { + return None; + } + parsed.set_host(Some(&host)).ok()?; + let default_port = match scheme.as_str() { + "ldaps" => Some(636), + "ldap" => Some(389), + "mockldap" => None, + _ => None, + }; + if default_port.is_some_and(|port| parsed.port() == Some(port)) { + parsed.set_port(None).ok()?; + } + Some(parsed.to_string().trim_end_matches('/').to_string()) +} + pub struct AdminModuleValidationInput<'a> { pub module_name: &'a str, pub oauth_providers: &'a [StoredOAuthProviderModuleConfig], @@ -1394,14 +1890,50 @@ pub fn build_admin_module_validation_result( let Some(config) = ldap_config else { return (false, Some("请先配置 LDAP 连接信息".to_string())); }; - if config.server_url.trim().is_empty() { + if normalize_ldap_transport_server_url(&config.server_url, config.use_starttls) + .is_none() + { return (false, Some("请配置 LDAP 服务器地址".to_string())); } - if config.bind_dn.trim().is_empty() { - return (false, Some("请配置绑定 DN".to_string())); + if !ldap_distinguished_name_is_valid(&config.bind_dn) { + return (false, Some("请配置有效的绑定 DN".to_string())); } - if config.base_dn.trim().is_empty() { - return (false, Some("请配置搜索基准 DN".to_string())); + if !ldap_distinguished_name_is_valid(&config.base_dn) { + return (false, Some("请配置有效的搜索基准 DN".to_string())); + } + let search_filter = config + .user_search_filter + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("(uid={username})"); + if !ldap_search_filter_is_valid(search_filter) { + return (false, Some("请配置有效的 LDAP 搜索过滤器".to_string())); + } + for (attribute, label) in [ + ( + config.username_attr.as_deref().unwrap_or("uid"), + "用户名属性", + ), + (config.email_attr.as_deref().unwrap_or("mail"), "邮箱属性"), + ( + config.display_name_attr.as_deref().unwrap_or("cn"), + "显示名称属性", + ), + ] { + let attribute = attribute.trim(); + let attribute = if attribute.is_empty() { + match label { + "用户名属性" => "uid", + "邮箱属性" => "mail", + _ => "cn", + } + } else { + attribute + }; + if !ldap_attribute_description_is_valid(attribute) { + return (false, Some(format!("请配置有效的 LDAP {label}"))); + } } if config .bind_password_encrypted @@ -1653,6 +2185,11 @@ pub fn normalize_admin_system_config_key(requested_key: &str) -> String { "module.server_chan_push.send_key".to_string() } else if trimmed.eq_ignore_ascii_case("module.important_notification.server_chan_template") { "module.server_chan_push.template".to_string() + } else if let Some(canonical) = SENSITIVE_SYSTEM_CONFIG_KEYS + .iter() + .find(|candidate| candidate.eq_ignore_ascii_case(trimmed)) + { + (*canonical).to_string() } else { trimmed.to_string() } @@ -1702,7 +2239,7 @@ pub fn admin_system_config_default_value(key: &str) -> Option "site_subtitle" => Some(json!("AI Gateway")), "default_user_initial_gift_usd" => Some(json!(10.0)), "password_policy_level" => Some(json!("weak")), - REQUEST_RECORD_LEVEL_KEY => Some(json!("full")), + REQUEST_RECORD_LEVEL_KEY => Some(json!("basic")), "max_request_body_size" => Some(json!(0)), "max_response_body_size" => Some(json!(0)), "sensitive_headers" => Some(json!([ @@ -1725,8 +2262,6 @@ pub fn admin_system_config_default_value(key: &str) -> Option "proxy_node_metrics_cleanup_batch_size" => Some(json!(5000)), "enable_provider_checkin" => Some(json!(true)), "provider_checkin_time" => Some(json!("01:05")), - "provider_priority_mode" => Some(json!("provider")), - "scheduling_mode" => Some(json!("cache_affinity")), "auto_delete_expired_keys" => Some(json!(false)), "turnstile_enabled" => Some(json!(false)), "turnstile_site_key" => Some(serde_json::Value::Null), @@ -1754,10 +2289,13 @@ pub fn admin_system_config_default_value(key: &str) -> Option "email_suffix_mode" => Some(json!("none")), "email_suffix_list" => Some(json!([])), "enable_format_conversion" => Some(json!(false)), - "cyber_continue_failover" => Some(json!(false)), + EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => Some(json!([])), "enable_model_directives" => Some(json!(false)), + // Failover after a provider-side Cyber policy refusal is an explicit + // opt-in. Keep the system-config fallback aligned with the routing + // policy default so an unset value cannot accidentally enable it. + "cyber_continue_failover" => Some(json!(false)), "model_directives" => Some(aether_ai_formats::default_model_directives_config()), - "keep_priority_on_conversion" => Some(json!(false)), "audit_log_retention_days" => Some(json!(30)), "enable_db_maintenance" => Some(json!(true)), "system_proxy_node_id" => Some(serde_json::Value::Null), @@ -1791,6 +2329,58 @@ pub fn admin_system_config_default_value(key: &str) -> Option } } +pub fn normalize_execution_extra_trusted_dns_hosts_config_value( + value: serde_json::Value, +) -> Result { + let values = match value { + Value::Null => Vec::new(), + Value::Array(values) => values, + _ => return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidValue), + }; + if values.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES { + return Err(ExecutionExtraTrustedDnsHostsConfigError::TooManyEntries); + } + + let mut hosts = BTreeSet::new(); + for value in values { + let host = value + .as_str() + .map(str::trim) + .ok_or(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost)?; + let host = host.trim_end_matches('.').to_ascii_lowercase(); + if !execution_extra_trusted_dns_host_is_valid(&host) { + return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost); + } + hosts.insert(host); + } + + Ok(Value::Array(hosts.into_iter().map(Value::String).collect())) +} + +fn execution_extra_trusted_dns_host_is_valid(host: &str) -> bool { + if host.is_empty() + || host.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES + || !host.is_ascii() + || host.parse::().is_ok() + { + return false; + } + + let labels = host.split('.').collect::>(); + if labels.len() < 2 { + return false; + } + labels.iter().all(|label| { + !label.is_empty() + && label.len() <= 63 + && !label.starts_with('-') + && !label.ends_with('-') + && label + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') + }) +} + pub fn build_admin_system_configs_payload( entries: &[StoredSystemConfigEntry], ) -> serde_json::Value { @@ -1975,6 +2565,17 @@ fn normalize_nullable_string_config_value( } } +fn normalize_smtp_control_config_value(value: serde_json::Value) -> Result { + let normalized = normalize_nullable_string_config_value(value)?; + if normalized + .as_str() + .is_some_and(|raw| raw.bytes().any(|byte| matches!(byte, b'\r' | b'\n' | 0))) + { + return Err(()); + } + Ok(normalized) +} + fn normalize_bark_server_url_config_value( value: serde_json::Value, ) -> Result { @@ -1985,10 +2586,20 @@ fn normalize_bark_server_url_config_value( if raw.is_empty() { return Ok(json!(DEFAULT_BARK_API_BASE)); } - if !raw.starts_with("https://") && !raw.starts_with("http://") { + if raw.len() > 2_048 { return Err(()); } - Ok(json!(raw)) + let parsed = url::Url::parse(raw).map_err(|_| ())?; + if !matches!(parsed.scheme(), "https" | "http") + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return Err(()); + } + Ok(json!(parsed.as_str().trim_end_matches('/'))) } _ => Err(()), } @@ -2236,8 +2847,8 @@ pub fn parse_admin_system_config_update( } match normalized_key.as_str() { - "cyber_continue_failover" - | "enable_model_directives" + "enable_model_directives" + | "cyber_continue_failover" | "module.important_notification.enabled" | "module.important_notification.email_enabled" | "module.server_chan_push.enabled" @@ -2261,6 +2872,15 @@ pub fn parse_admin_system_config_update( ) })?; } + EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => { + value = + normalize_execution_extra_trusted_dns_hosts_config_value(value).map_err(|_| { + ( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "额外可信 Fake-IP 域名配置格式无效" }), + ) + })?; + } "module.important_notification.default_channel" => { value = normalize_notification_channel_value(value).map_err(|_| { ( @@ -2329,6 +2949,14 @@ pub fn parse_admin_system_config_update( } }; } + "smtp_host" | "smtp_from_email" | "smtp_from_name" => { + value = normalize_smtp_control_config_value(value).map_err(|_| { + ( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "请求数据验证失败" }), + ) + })?; + } "model_directives" => { if value.is_null() { value = aether_ai_formats::default_model_directives_config(); @@ -2402,6 +3030,15 @@ pub fn parse_admin_system_config_update( _ => {} } + if is_sensitive_admin_system_config_key(&normalized_key) + && !matches!(&value, Value::Null | Value::String(_)) + { + return Err(( + http::StatusCode::BAD_REQUEST, + json!({ "detail": "请求数据验证失败" }), + )); + } + Ok(AdminSystemConfigUpdate { normalized_key, value, @@ -2719,6 +3356,9 @@ pub fn build_admin_proxy_nodes_not_found_response() -> Response { } pub fn build_admin_proxy_node_payload(node: &StoredProxyNode) -> serde_json::Value { + // Keep operational tunnel metadata, but never expose either legacy plaintext + // or encrypted credential material through an admin response. + let proxy_metadata = redact_admin_proxy_node_metadata(node.proxy_metadata.as_ref()); let mut payload = serde_json::Map::from_iter([ ("id".to_string(), json!(node.id)), ("name".to_string(), json!(node.name)), @@ -2755,7 +3395,7 @@ pub fn build_admin_proxy_node_payload(node: &StoredProxyNode) -> serde_json::Val ("failed_requests".to_string(), json!(node.failed_requests)), ("dns_failures".to_string(), json!(node.dns_failures)), ("stream_errors".to_string(), json!(node.stream_errors)), - ("proxy_metadata".to_string(), json!(node.proxy_metadata)), + ("proxy_metadata".to_string(), json!(proxy_metadata)), ("hardware_info".to_string(), json!(node.hardware_info)), ( "estimated_max_concurrency".to_string(), @@ -2776,17 +3416,34 @@ pub fn build_admin_proxy_node_payload(node: &StoredProxyNode) -> serde_json::Val if node.is_manual { payload.insert("proxy_url".to_string(), json!(node.proxy_url)); payload.insert("proxy_username".to_string(), json!(node.proxy_username)); - payload.insert( - "proxy_password".to_string(), - json!(mask_admin_proxy_node_password( - node.proxy_password.as_deref() - )), - ); } + payload.insert( + "has_proxy_password".to_string(), + json!(node + .proxy_password + .as_deref() + .is_some_and(|password| !password.is_empty())), + ); serde_json::Value::Object(payload) } +fn redact_admin_proxy_node_metadata( + metadata: Option<&serde_json::Value>, +) -> Option { + let metadata = metadata?; + let mut metadata = metadata.clone(); + if let Some(tunnel_security) = metadata + .as_object_mut() + .and_then(|object| object.get_mut("tunnel_security")) + .and_then(serde_json::Value::as_object_mut) + { + tunnel_security.remove("encryption_key"); + tunnel_security.remove("encryption_key_encrypted"); + } + Some(metadata) +} + pub fn build_admin_proxy_node_event_payload(event: &StoredProxyNodeEvent) -> serde_json::Value { json!({ "id": event.id, @@ -3114,25 +3771,248 @@ fn suffixed_path_identifier_from_path( .filter(|value| !value.is_empty() && !value.contains('/')) } -fn mask_admin_proxy_node_password(password: Option<&str>) -> Option { - let password = password?; - if password.is_empty() { - return None; - } - if password.len() < 8 { - return Some("****".to_string()); - } - Some(format!( - "{}****{}", - &password[..2], - &password[password.len() - 2..] - )) -} - #[cfg(test)] mod tests { use super::*; + #[test] + fn email_template_update_rejects_oversized_and_control_fields() { + let oversized_subject = serde_json::json!({ + "subject": "x".repeat(ADMIN_EMAIL_TEMPLATE_MAX_SUBJECT_BYTES + 1) + }); + assert!( + parse_admin_email_template_update(oversized_subject.to_string().as_bytes()).is_err() + ); + + let control_html = serde_json::json!({ "html": "

ok

\u{0001}" }); + assert!(parse_admin_email_template_update(control_html.to_string().as_bytes()).is_err()); + } + + #[test] + fn email_template_preview_bounds_html_and_variable_values() { + let oversized_html = serde_json::json!({ + "html": "x".repeat(ADMIN_EMAIL_TEMPLATE_MAX_HTML_BYTES + 1) + }); + assert!(parse_admin_email_template_preview_payload(Some( + oversized_html.to_string().as_bytes() + )) + .is_err()); + + let oversized_value = serde_json::json!({ + "app_name": "x".repeat(ADMIN_EMAIL_TEMPLATE_MAX_PREVIEW_VALUE_BYTES + 1) + }); + assert!(parse_admin_email_template_preview_payload(Some( + oversized_value.to_string().as_bytes() + )) + .is_err()); + } + + #[test] + fn email_template_field_validators_allow_normal_markup_and_subjects() { + assert!(admin_email_template_subject_is_valid("验证码")); + assert!(admin_email_template_html_is_valid("

{{app_name}}

\n")); + assert!(!admin_email_template_html_is_valid("

bad\u{007f}

")); + } + + #[test] + fn ldap_module_validation_requires_an_encrypted_transport() { + let config = |server_url: &str, use_starttls: bool| StoredLdapModuleConfig { + server_url: server_url.to_string(), + bind_dn: "cn=bind,dc=example,dc=com".to_string(), + bind_password_encrypted: Some("sealed-password".to_string()), + base_dn: "dc=example,dc=com".to_string(), + user_search_filter: Some("(uid={username})".to_string()), + username_attr: Some("uid".to_string()), + email_attr: Some("mail".to_string()), + display_name_attr: Some("cn".to_string()), + is_enabled: true, + is_exclusive: false, + use_starttls, + connect_timeout: Some(10), + }; + + assert!(ldap_module_config_is_valid(Some(&config( + "ldaps://ldap.internal.example:636", + false, + )))); + assert!(ldap_module_config_is_valid(Some(&config( + "ldap://10.0.0.7:389", + true, + )))); + assert!(!ldap_module_config_is_valid(Some(&config( + "ldap://10.0.0.7:389", + false, + )))); + assert!(!ldap_module_config_is_valid(Some(&config( + "ldaps://user:password@ldap.internal.example", + false, + )))); + assert!(!ldap_module_config_is_valid(Some(&config( + "ldaps://ldap.internal.example?secret=1", + false, + )))); + assert!(!ldap_module_config_is_valid(Some(&config( + "ldaps://ldap.internal.example#fragment", + false, + )))); + assert!( + normalize_ldap_transport_server_url("mockldap://ldap.internal.example", false,) + .is_none() + ); + assert!(normalize_ldap_transport_server_url_for_tests( + "mockldap://ldap.internal.example", + false, + ) + .is_some()); + } + + #[test] + fn ldap_search_filter_requires_one_outer_expression() { + assert!(ldap_search_filter_is_valid("(uid={username})")); + assert!(ldap_search_filter_is_valid( + "(&(objectClass=person)(uid={username}))" + )); + + assert!(!ldap_search_filter_is_valid( + "(uid={username})(objectClass=person)" + )); + assert!(!ldap_search_filter_is_valid( + "(uid={username}) (objectClass=person)" + )); + assert!(!ldap_search_filter_is_valid("(uid={username}))")); + assert!(!ldap_search_filter_is_valid("((uid={username})")); + assert!(!ldap_search_filter_is_valid("(uid=alice)")); + } + + #[test] + fn ldap_attribute_descriptions_follow_bounded_rfc4512_shape() { + for valid in [ + "uid", + "sAMAccountName", + "display-name", + "mail;lang-en", + "1.2.840.113556.1.4.221", + ] { + assert!(ldap_attribute_description_is_valid(valid), "{valid}"); + } + for invalid in [ + "", + "uid)(|(objectClass=*)", + "uid\nmail", + "-uid", + "01.2.3", + "1", + "uid;", + "uid;lang=en", + "用户名", + ] { + assert!(!ldap_attribute_description_is_valid(invalid), "{invalid}"); + } + assert!(!ldap_attribute_description_is_valid(&"a".repeat(129))); + } + + #[test] + fn ldap_distinguished_names_reject_control_and_unbounded_data() { + assert!(ldap_distinguished_name_is_valid( + "cn=service\\, account,ou=users,dc=example,dc=com" + )); + assert!(!ldap_distinguished_name_is_valid("")); + assert!(!ldap_distinguished_name_is_valid( + "cn=service\n,dc=example,dc=com" + )); + assert!(!ldap_distinguished_name_is_valid(&"a".repeat(4097))); + } + + #[test] + fn ldap_module_field_validation_rejects_filter_and_attribute_injection() { + let mut config = StoredLdapModuleConfig { + server_url: "ldaps://ldap.internal.example".to_string(), + bind_dn: "cn=bind,dc=example,dc=com".to_string(), + bind_password_encrypted: Some("sealed-password".to_string()), + base_dn: "dc=example,dc=com".to_string(), + user_search_filter: Some("(uid={username})".to_string()), + username_attr: Some("uid".to_string()), + email_attr: Some("mail".to_string()), + display_name_attr: Some("cn".to_string()), + is_enabled: true, + is_exclusive: false, + use_starttls: false, + connect_timeout: Some(10), + }; + assert!(ldap_module_config_fields_are_valid(&config)); + + config.user_search_filter = Some("(uid={username})(objectClass=*)".to_string()); + assert!(!ldap_module_config_fields_are_valid(&config)); + config.user_search_filter = Some("(uid={username})".to_string()); + config.username_attr = Some("uid)(|(objectClass=*)".to_string()); + assert!(!ldap_module_config_fields_are_valid(&config)); + config.username_attr = Some("uid".to_string()); + config.base_dn = "dc=example,dc=com\n".to_string(); + assert!(!ldap_module_config_fields_are_valid(&config)); + } + + #[test] + fn admin_proxy_node_payload_redacts_tunnel_psk_and_password() { + let mut node = StoredProxyNode::new( + "node-1".to_string(), + "edge-1".to_string(), + "127.0.0.1".to_string(), + 0, + true, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + false, + false, + 0, + ) + .expect("proxy node should build") + .with_manual_proxy_fields( + Some("http://proxy.example:8080".to_string()), + Some("alice".to_string()), + Some("supersecret".to_string()), + ); + node.proxy_metadata = Some(json!({ + "version": "1.2.3", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key": "secret-psk", + "encryption_key_encrypted": "sealed-secret-psk", + "rotation_id": "rotation-1" + } + })); + + let payload = build_admin_proxy_node_payload(&node); + + assert_eq!(payload["has_proxy_password"], json!(true)); + assert!(payload.get("proxy_password").is_none()); + assert_eq!( + payload["proxy_metadata"]["tunnel_security"]["mode"], + json!("non_tls_required") + ); + assert_eq!( + payload["proxy_metadata"]["tunnel_security"]["rotation_id"], + json!("rotation-1") + ); + assert!(payload["proxy_metadata"]["tunnel_security"] + .get("encryption_key") + .is_none()); + assert!(payload["proxy_metadata"]["tunnel_security"] + .get("encryption_key_encrypted") + .is_none()); + assert_eq!( + node.proxy_metadata + .as_ref() + .and_then(|value| value.pointer("/tunnel_security/encryption_key")) + .and_then(serde_json::Value::as_str), + Some("secret-psk") + ); + } + #[test] fn api_formats_payload_exposes_realtime_and_codex_live_separately() { let payload = build_admin_api_formats_payload(); @@ -3308,6 +4188,57 @@ mod tests { } } + #[test] + fn system_config_debug_output_redacts_exported_credentials() { + let secrets = [ + "debug-provider-api-key", + "debug-provider-auth-config", + "debug-proxy-password", + "debug-ldap-password", + "debug-oauth-client-secret", + "debug-smtp-password", + ]; + let document = serde_json::from_value::(json!({ + "version": "2.2", + "providers": [{ + "name": "provider", + "api_keys": [{ + "api_key": secrets[0], + "auth_config": {"refresh_token": secrets[1]} + }] + }], + "proxy_nodes": [{ + "name": "proxy", + "proxy_password": secrets[2] + }], + "ldap_config": { + "server_url": "ldaps://ldap.example.com", + "bind_dn": "cn=admin,dc=example,dc=com", + "bind_password": secrets[3], + "base_dn": "dc=example,dc=com" + }, + "oauth_providers": [{ + "provider_type": "linuxdo", + "display_name": "Linux.do", + "client_id": "client-id", + "client_secret": secrets[4], + "redirect_uri": "https://gateway.example/api/oauth/linuxdo/callback", + "frontend_callback_url": "https://frontend.example/auth/callback" + }], + "system_configs": [{ + "key": "smtp_password", + "value": secrets[5] + }] + })) + .expect("system config document should deserialize"); + let rendered = format!("{document:?}"); + + for secret in secrets { + assert!(!rendered.contains(secret), "Debug output leaked {secret}"); + } + assert!(rendered.contains("[REDACTED]")); + } + #[test] fn parse_admin_system_config_import_request_rejects_unknown_versions() { for version in ["1.9", "2.4"] { @@ -3469,6 +4400,42 @@ mod tests { assert!(!is_sensitive_admin_system_config_key("site_name")); } + #[test] + fn sensitive_admin_system_config_keys_use_canonical_storage_spelling() { + assert_eq!( + normalize_admin_system_config_key(" SMTP_PASSWORD "), + "smtp_password" + ); + assert_eq!( + normalize_admin_system_config_key("BACKUP_S3_SECRET_ACCESS_KEY"), + "backup_s3_secret_access_key" + ); + assert_eq!( + normalize_admin_system_config_key("MODULE.BARK_PUSH.DEVICE_KEY"), + "module.bark_push.device_key" + ); + } + + #[test] + fn sensitive_admin_system_config_values_reject_non_strings() { + for key in SENSITIVE_SYSTEM_CONFIG_KEYS { + for value in [ + json!(123), + json!(true), + json!(["secret"]), + json!({"secret": true}), + ] { + let body = serde_json::to_vec(&json!({ "value": value })) + .expect("sensitive config body should serialize"); + let error = parse_admin_system_config_update(key, &body) + .expect_err("non-string sensitive value must fail closed"); + assert_eq!(error.0, http::StatusCode::BAD_REQUEST, "key={key}"); + } + assert!(parse_admin_system_config_update(key, br#"{"value":null}"#).is_ok()); + assert!(parse_admin_system_config_update(key, br#"{"value":"secret"}"#).is_ok()); + } + } + #[test] fn s3_backup_secret_access_key_is_sensitive() { assert!(is_sensitive_admin_system_config_key( @@ -3530,6 +4497,28 @@ mod tests { .is_err()); } + #[test] + fn smtp_control_fields_reject_header_and_command_injection() { + for key in ["smtp_host", "smtp_from_email", "smtp_from_name"] { + for value in [ + "safe\r\nX-Injected: yes", + "safe\nMAIL FROM:", + "safe\0tail", + ] { + let body = serde_json::to_vec(&json!({ "value": value })) + .expect("SMTP config body should serialize"); + let error = parse_admin_system_config_update(key, &body) + .expect_err("SMTP control characters must be rejected"); + assert_eq!(error.0, http::StatusCode::BAD_REQUEST); + } + } + + let update = + parse_admin_system_config_update("smtp_from_name", br#"{"value":" Aether Mail "}"#) + .expect("normal SMTP display name should parse"); + assert_eq!(update.value, json!("Aether Mail")); + } + #[test] fn model_directives_update_accepts_legacy_and_current_config_shapes() { for body in [ @@ -3632,6 +4621,36 @@ mod tests { .is_err()); } + #[test] + fn extra_trusted_dns_hosts_update_normalizes_exact_hostnames() { + let update = parse_admin_system_config_update( + "execution_extra_trusted_dns_hosts", + br#"{"value":[" API.Example.COM. ","api.example.com"]}"#, + ) + .expect("valid extra trusted DNS hosts should parse"); + assert_eq!(update.value, json!(["api.example.com"])); + } + + #[test] + fn extra_trusted_dns_hosts_update_rejects_non_exact_hostnames() { + for value in [ + r#"["*.example.com"]"#, + r#"["example.com:443"]"#, + r#"["https://example.com/path"]"#, + r#"["10.0.0.1"]"#, + r#"["example..com"]"#, + r#"["localhost"]"#, + ] { + let body = format!(r#"{{"value":{value}}}"#); + let error = parse_admin_system_config_update( + "execution_extra_trusted_dns_hosts", + body.as_bytes(), + ) + .expect_err("non-exact hostname should be rejected"); + assert_eq!(error.0, http::StatusCode::BAD_REQUEST); + } + } + #[test] fn legacy_notification_email_config_key_normalizes_to_important_notification() { assert_eq!( diff --git a/crates/aether-ai/formats/src/api.rs b/crates/aether-ai/formats/src/api.rs index e75172785..de48589f7 100644 --- a/crates/aether-ai/formats/src/api.rs +++ b/crates/aether-ai/formats/src/api.rs @@ -78,6 +78,7 @@ pub use crate::formats::openai::{ OpenAiProviderRequestFinalization, }, responses::{ + normalize_openai_responses_message_item_ids, openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, strip_incompatible_openai_responses_reasoning_items, strip_incompatible_openai_responses_reasoning_items_with_policy, diff --git a/crates/aether-ai/formats/src/formats/conversion/response.rs b/crates/aether-ai/formats/src/formats/conversion/response.rs index 126fcf25f..646b99a9f 100644 --- a/crates/aether-ai/formats/src/formats/conversion/response.rs +++ b/crates/aether-ai/formats/src/formats/conversion/response.rs @@ -9,7 +9,7 @@ use serde_json::{json, Value}; use crate::formats::{ context::FormatContext, openai::responses::{ - openai_responses_synthetic_reasoning_item_id, + openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, response::ensure_modern_openai_responses_response_fields, }, registry, @@ -218,7 +218,7 @@ pub fn build_openai_responses_response_with_content( if !content.is_empty() { output.push(json!({ "type": "message", - "id": format!("{response_id}_msg"), + "id": openai_responses_message_item_id(response_id, 0), "role": "assistant", "status": "completed", "content": content 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 149a89a99..5b4208272 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 @@ -26,6 +26,7 @@ pub fn from_raw(body_json: &Value) -> Option { } let candidates = body.get("candidates")?.as_array()?; + let usage = gemini_usage_to_canonical(body.get("usageMetadata")); let mut outputs = Vec::new(); for (fallback_index, candidate) in candidates.iter().enumerate() { let candidate_object = candidate.as_object()?; @@ -79,7 +80,10 @@ pub fn from_raw(body_json: &Value) -> Option { extensions, }); } - outputs.retain(gemini_response_output_has_visible_content); + outputs.retain(|output| { + gemini_response_output_has_visible_content(output) + || gemini_response_output_is_reasoning_exhausted_terminal(output, usage.as_ref()) + }); if outputs.is_empty() { return None; } @@ -106,7 +110,7 @@ pub fn from_raw(body_json: &Value) -> Option { outputs, content, stop_reason, - usage: gemini_usage_to_canonical(body.get("usageMetadata")), + usage, extensions: gemini_extensions( body, &[ @@ -127,16 +131,36 @@ pub fn from_raw(body_json: &Value) -> Option { fn gemini_response_output_has_visible_content(output: &CanonicalResponseOutput) -> bool { output.content.iter().any(|block| match block { - CanonicalContentBlock::Text { text, .. } => !text.trim().is_empty(), + CanonicalContentBlock::Text { text, .. } | CanonicalContentBlock::Thinking { text, .. } => { + !text.trim().is_empty() + } CanonicalContentBlock::ToolUse { .. } | CanonicalContentBlock::ToolResult { .. } | CanonicalContentBlock::Image { .. } | CanonicalContentBlock::File { .. } | CanonicalContentBlock::Audio { .. } => true, - CanonicalContentBlock::Thinking { .. } | CanonicalContentBlock::Unknown { .. } => false, + CanonicalContentBlock::Unknown { .. } => false, }) } +fn gemini_response_output_is_reasoning_exhausted_terminal( + output: &CanonicalResponseOutput, + usage: Option<&CanonicalUsage>, +) -> bool { + matches!(output.stop_reason, Some(CanonicalStopReason::MaxTokens)) + && usage.is_some_and(|usage| usage.reasoning_tokens > 0) + && output.content.iter().any(|block| { + matches!( + block, + CanonicalContentBlock::Thinking { + text, + signature: Some(signature), + .. + } if text.trim().is_empty() && !signature.trim().is_empty() + ) + }) +} + pub fn to_raw(canonical: &CanonicalResponse, report_context: &Value) -> Option { let mut response = canonical_to_gemini_response(canonical, report_context)?; if let Some(object) = response.as_object_mut() { @@ -430,7 +454,7 @@ mod tests { } #[test] - fn gemini_response_with_only_thought_parts_is_not_success() { + fn gemini_response_with_only_thought_parts_is_success() { let body = json!({ "candidates": [{ "content": { @@ -443,6 +467,92 @@ mod tests { "responseId": "resp-thought-only" }); + let canonical = from_raw(&body).expect("thought text is representable output"); + assert!(matches!( + canonical.content.first(), + Some(CanonicalContentBlock::Thinking { text, .. }) if text == "hidden plan" + )); + assert!(matches!( + canonical.stop_reason, + Some(CanonicalStopReason::MaxTokens) + )); + + let openai = crate::canonical_to_openai_chat_response(&canonical); + assert_eq!( + openai["choices"][0]["message"]["reasoning_content"], + "hidden plan" + ); + assert_eq!(openai["choices"][0]["finish_reason"], "length"); + } + + #[test] + fn gemini_response_with_signature_only_reasoning_exhaustion_is_success() { + let body = json!({ + "candidates": [{ + "content": { + "role": "model", + "parts": [{ + "text": "", + "thoughtSignature": "opaque-thought-signature" + }] + }, + "finishReason": "MAX_TOKENS" + }], + "usageMetadata": { + "promptTokenCount": 22, + "thoughtsTokenCount": 29, + "totalTokenCount": 51 + }, + "modelVersion": "gemini-3.7-flash-tiered", + "responseId": "resp-signature-only" + }); + + let canonical = from_raw(&body).expect("reasoning exhaustion is a valid terminal"); + assert!(matches!( + canonical.content.first(), + Some(CanonicalContentBlock::Thinking { + text, + signature: Some(signature), + .. + }) if text.is_empty() && signature == "opaque-thought-signature" + )); + assert!(matches!( + canonical.stop_reason, + Some(CanonicalStopReason::MaxTokens) + )); + assert_eq!( + canonical.usage.as_ref().map(|usage| usage.reasoning_tokens), + Some(29) + ); + + let openai = crate::canonical_to_openai_chat_response(&canonical); + assert_eq!(openai["choices"][0]["finish_reason"], "length"); + assert_eq!( + openai["usage"]["completion_tokens_details"]["reasoning_tokens"], + 29 + ); + } + + #[test] + fn gemini_signature_only_terminal_without_reasoning_usage_is_not_success() { + let body = json!({ + "candidates": [{ + "content": { + "role": "model", + "parts": [{ + "text": "", + "thoughtSignature": "opaque-thought-signature" + }] + }, + "finishReason": "MAX_TOKENS" + }], + "usageMetadata": { + "promptTokenCount": 22, + "thoughtsTokenCount": 0, + "totalTokenCount": 22 + } + }); + assert!(from_raw(&body).is_none()); } 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 8b890cacb..5f5875ab2 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 @@ -104,10 +104,25 @@ impl GeminiProviderState { let Some(candidate_object) = candidate.as_object() else { continue; }; + let (response_id, response_model) = self.identity(report_context); + let terminal_error = gemini_stream_terminal_error_payload( + candidate_object, + response_id.as_str(), + response_model.as_str(), + event_object.get("usageMetadata"), + ); 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)); + self.finished = true; + } continue; }; let Some(parts) = content.get("parts").and_then(Value::as_array) else { + if let Some(payload) = terminal_error { + out.push(self.unknown_frame(report_context, payload)); + self.finished = true; + } continue; }; if !parts.is_empty() { @@ -129,7 +144,8 @@ impl GeminiProviderState { let is_reasoning = part_object .get("thought") .and_then(Value::as_bool) - .unwrap_or(false); + .unwrap_or(false) + || (text.trim().is_empty() && reasoning_signature.is_some()); let previous = if is_reasoning { self.reasoning_parts.entry(index).or_default() } else { @@ -305,6 +321,11 @@ impl GeminiProviderState { }); } } + if let Some(payload) = terminal_error { + out.push(self.unknown_frame(report_context, payload)); + self.finished = true; + continue; + } if let Some(finish_reason) = candidate_object.get("finishReason").and_then(Value::as_str) { @@ -349,6 +370,60 @@ impl GeminiProviderState { } } +fn gemini_stream_terminal_error_payload( + candidate: &Map, + response_id: &str, + model: &str, + usage_metadata: Option<&Value>, +) -> Option { + let finish_reason = candidate + .get("finishReason") + .or_else(|| candidate.get("finish_reason")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + matches!( + *value, + "MALFORMED_FUNCTION_CALL" + | "UNEXPECTED_TOOL_CALL" + | "TOO_MANY_TOOL_CALLS" + | "MISSING_THOUGHT_SIGNATURE" + | "MALFORMED_RESPONSE" + ) + })?; + let message = candidate + .get("finishMessage") + .or_else(|| candidate.get("finish_message")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| format!("Gemini stream ended with {finish_reason}")); + + let mut response = json!({ + "id": response_id, + "object": "response", + "model": model, + "status": "failed", + "error": { + "type": "upstream_gemini_finish_error", + "code": finish_reason, + "message": message, + "upstream_status": 200 + } + }); + if let Some(usage) = canonical_usage_from_gemini_usage(usage_metadata) + .map(|usage| openai_responses_usage_from_usage(&usage)) + { + response["usage"] = usage; + } + + Some(json!({ + "type": "response.failed", + "response": response + })) +} + fn map_gemini_stream_finish_reason(value: &str) -> Option<&str> { match value { "STOP" => Some("stop"), @@ -960,6 +1035,113 @@ mod tests { ))); } + #[test] + fn gemini_provider_state_preserves_signature_only_reasoning_terminal() { + let mut state = GeminiProviderState::default(); + let report_context = json!({}); + let frames = state + .push_line( + &report_context, + data_line(json!({ + "response": { + "responseId": "resp_signature_only_123", + "modelVersion": "gemini-3.7-flash-tiered", + "candidates": [{ + "index": 0, + "finishReason": "MAX_TOKENS", + "content": { + "role": "model", + "parts": [{ + "text": "", + "thoughtSignature": "opaque-thought-signature" + }] + } + }], + "usageMetadata": { + "promptTokenCount": 22, + "thoughtsTokenCount": 29, + "totalTokenCount": 51 + } + }, + "traceId": "trace-signature-only" + })), + ) + .expect("signature-only reasoning terminal should parse"); + + assert!(frames.iter().any(|frame| matches!( + frame.event, + CanonicalStreamEvent::ReasoningSignature(ref signature) + if signature == "opaque-thought-signature" + ))); + assert!(frames.iter().any(|frame| matches!( + frame.event, + CanonicalStreamEvent::Finish { + ref finish_reason, + usage: Some(CanonicalUsage { + reasoning_tokens: 29, + .. + }), + } if finish_reason.as_deref() == Some("length") + ))); + } + + #[test] + fn gemini_provider_state_emits_terminal_error_for_malformed_function_call() { + let mut state = GeminiProviderState::default(); + let report_context = json!({}); + let frames = state + .push_line( + &report_context, + data_line(json!({ + "response": { + "responseId": "resp_malformed_tool_call", + "modelVersion": "gemini-3.7-flash-tiered", + "candidates": [{ + "index": 0, + "content": { + "role": "model", + "parts": [{ + "text": "", + "thoughtSignature": "opaque-thought-signature" + }] + }, + "finishReason": "MALFORMED_FUNCTION_CALL", + "finishMessage": "Malformed function call: Function call is empty - no input to parse." + }], + "usageMetadata": { + "promptTokenCount": 206744, + "cachedContentTokenCount": 203947, + "thoughtsTokenCount": 1130, + "totalTokenCount": 207874 + } + } + })), + ) + .expect("malformed function call terminal should parse"); + + assert!(frames.iter().any(|frame| matches!( + &frame.event, + CanonicalStreamEvent::UnknownEvent(payload) + if payload["type"] == "response.failed" + && payload["response"]["status"] == "failed" + && payload["response"]["id"] == "resp_malformed_tool_call" + && payload["response"]["model"] == "gemini-3.7-flash-tiered" + && payload["response"]["error"]["code"] == "MALFORMED_FUNCTION_CALL" + && payload["response"]["error"]["message"] + == "Malformed function call: Function call is empty - no input to parse." + && payload["response"]["usage"]["input_tokens"] == 206744 + && payload["response"]["usage"]["output_tokens"] == 1130 + && payload["response"]["usage"]["total_tokens"] == 207874 + ))); + assert!(!frames + .iter() + .any(|frame| matches!(frame.event, CanonicalStreamEvent::Finish { .. }))); + assert!(state + .finish(&report_context) + .expect("finished error stream should not synthesize success") + .is_empty()); + } + #[test] fn gemini_provider_state_parses_function_response_as_tool_result() { let mut state = GeminiProviderState::default(); 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 632b83ab1..22dc3e502 100644 --- a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs +++ b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs @@ -4,7 +4,7 @@ use serde_json::{json, Map, Value}; use crate::formats::openai::namespace::NamespaceToolAliases; use crate::formats::openai::responses::{ - encode_gemini_tool_signature_carrier_with_direction, + encode_gemini_tool_signature_carrier_with_direction, openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, response::{ ensure_modern_openai_responses_response_fields, openai_responses_current_timestamp, @@ -1840,7 +1840,7 @@ impl OpenAIResponsesProviderState { }); self.finished = true; } - "keepalive" => {} + "keepalive" | "ping" => {} event_type if openai_responses_stream_event_is_known_noop(event_type) => { self.ensure_started(report_context, &mut out); } @@ -2328,7 +2328,7 @@ impl OpenAIResponsesClientEmitter { fn message_item_id(&self) -> String { self.message_item_id .clone() - .unwrap_or_else(|| format!("{}_msg", self.response_id())) + .unwrap_or_else(|| openai_responses_message_item_id(self.response_id(), 0)) } fn reasoning_item_id(&self) -> String { @@ -2347,7 +2347,7 @@ impl OpenAIResponsesClientEmitter { fn ensure_message_item_id(&mut self) -> String { if self.message_item_id.is_none() { - self.message_item_id = Some(format!("{}_msg", self.response_id())); + self.message_item_id = Some(openai_responses_message_item_id(self.response_id(), 0)); } self.message_item_id() } @@ -4357,7 +4357,7 @@ mod tests { assert!(sse.contains("event: response.output_item.done\n")); assert!(sse.contains("event: response.completed\n")); assert!(sse.contains("\"response_id\":\"resp_stream_123\"")); - assert!(sse.contains("\"item_id\":\"resp_stream_123_msg\"")); + assert!(sse.contains("\"item_id\":\"msg_aether_")); assert!(sse.contains("\"text\":\"Hello\"")); assert!(sse.contains("\"output_text\":\"Hello\"")); assert!(sse.contains("\"created_at\":")); @@ -4421,8 +4421,8 @@ mod tests { ); let sse = String::from_utf8(bytes).expect("sse should be utf8"); - assert!(sse.contains("\"item_id\":\"msg_first_msg\"")); - assert!(!sse.contains("\"item_id\":\"msg_second_msg\"")); + assert!(sse.contains("\"item_id\":\"msg_aether_")); + assert!(!sse.contains("msg_second_msg")); } #[test] diff --git a/crates/aether-ai/formats/src/formats/openai/image/mod.rs b/crates/aether-ai/formats/src/formats/openai/image/mod.rs index d87defe51..6d0a87b91 100644 --- a/crates/aether-ai/formats/src/formats/openai/image/mod.rs +++ b/crates/aether-ai/formats/src/formats/openai/image/mod.rs @@ -1,3 +1,99 @@ +use url::Url; + pub mod request; pub mod spec; pub mod stream; + +pub(crate) const MAX_OPENAI_IMAGE_DATA_BYTES: usize = 64 * 1024 * 1024; +pub(crate) const MAX_OPENAI_IMAGE_EXTERNAL_URL_BYTES: usize = 64 * 1024; +pub(crate) const MAX_OPENAI_IMAGE_REVISED_PROMPT_BYTES: usize = 256 * 1024; + +pub(crate) fn is_safe_openai_image_base64_payload(value: &str) -> bool { + if value.is_empty() + || value.len() > MAX_OPENAI_IMAGE_DATA_BYTES + || value.len() % 4 == 1 + || value + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte.is_ascii_control()) + { + return false; + } + let bytes = value.as_bytes(); + let first_padding = bytes.iter().position(|byte| *byte == b'='); + if let Some(index) = first_padding { + let padding = bytes.len() - index; + if padding > 2 + || bytes[index..].iter().any(|byte| *byte != b'=') + || !bytes.len().is_multiple_of(4) + { + return false; + } + } + bytes[..first_padding.unwrap_or(bytes.len())] + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(*byte, b'+' | b'/')) +} + +pub(crate) fn normalize_openai_image_output_format(value: &str) -> Option<&'static str> { + let value = value.trim(); + if value.eq_ignore_ascii_case("png") { + Some("png") + } else if value.eq_ignore_ascii_case("jpeg") || value.eq_ignore_ascii_case("jpg") { + Some("jpeg") + } else if value.eq_ignore_ascii_case("webp") { + Some("webp") + } else { + None + } +} + +pub(crate) fn bounded_openai_image_revised_prompt(value: &str) -> Option<&str> { + let value = value.trim(); + (!value.is_empty() && value.len() <= MAX_OPENAI_IMAGE_REVISED_PROMPT_BYTES).then_some(value) +} + +pub(crate) fn parse_safe_openai_image_data_url(value: &str) -> Option<(&'static str, &str)> { + let (metadata, payload) = value.trim().split_once(',')?; + let mime_type = metadata.strip_prefix("data:")?.strip_suffix(";base64")?; + let mime_type = safe_openai_image_mime_type(mime_type.trim())?; + (!payload.is_empty() && is_safe_openai_image_base64_payload(payload)) + .then_some((mime_type, payload)) +} + +pub(crate) fn safe_openai_image_mime_type(value: &str) -> Option<&'static str> { + if value.eq_ignore_ascii_case("image/png") { + Some("image/png") + } else if value.eq_ignore_ascii_case("image/jpeg") || value.eq_ignore_ascii_case("image/jpg") { + Some("image/jpeg") + } else if value.eq_ignore_ascii_case("image/webp") { + Some("image/webp") + } else { + None + } +} + +pub(crate) fn sanitize_openai_image_source_url(value: &str) -> Option { + let value = value.trim(); + if value.is_empty() + || value + .chars() + .any(|character| character.is_ascii_control() || character.is_whitespace()) + { + return None; + } + if let Some((mime_type, payload)) = parse_safe_openai_image_data_url(value) { + return Some(format!("data:{mime_type};base64,{payload}")); + } + if value.len() > MAX_OPENAI_IMAGE_EXTERNAL_URL_BYTES { + return None; + } + let parsed = Url::parse(value).ok()?; + if !matches!(parsed.scheme(), "http" | "https") + || parsed.host_str().is_none() + || !parsed.username().is_empty() + || parsed.password().is_some() + { + return None; + } + Some(value.to_string()) +} diff --git a/crates/aether-ai/formats/src/formats/openai/image/request.rs b/crates/aether-ai/formats/src/formats/openai/image/request.rs index a28050af5..384f526a3 100644 --- a/crates/aether-ai/formats/src/formats/openai/image/request.rs +++ b/crates/aether-ai/formats/src/formats/openai/image/request.rs @@ -71,6 +71,9 @@ impl OpenAiImageNormalizeOptions { pub const CHATGPT_WEB_IMAGE_MAX_AREA: u64 = 1_500_000; pub const OPENAI_IMAGE_MAX_GENERATION_COUNT: u64 = 10; +const MAX_MULTIPART_PARTS: usize = 128; +const MAX_MULTIPART_PART_HEADER_BYTES: usize = 64 * 1024; +const MAX_MULTIPART_BODY_BYTES: usize = 256 * 1024 * 1024; #[derive(Debug, Clone, PartialEq, Eq)] pub struct ChatGptWebImageRequestError { @@ -1653,34 +1656,69 @@ fn parse_multipart_fields_from_base64( .get(http::header::CONTENT_TYPE) .and_then(|value| value.to_str().ok())?; let boundary = multipart_boundary(content_type)?; - let body_bytes = base64::engine::general_purpose::STANDARD + let body_bytes = + decode_multipart_body_base64_with_limit(body_base64, MAX_MULTIPART_BODY_BYTES)?; + Some(parse_multipart_fields(&body_bytes, boundary.as_str())) +} + +fn decode_multipart_body_base64_with_limit( + body_base64: &str, + decoded_limit: usize, +) -> Option> { + let max_encoded_len = decoded_limit + .checked_add(2) + .and_then(|value| value.checked_div(3)) + .and_then(|value| value.checked_mul(4)) + .unwrap_or(usize::MAX); + if body_base64.len() > max_encoded_len { + return None; + } + let body = base64::engine::general_purpose::STANDARD .decode(body_base64) .ok()?; - Some(parse_multipart_fields(&body_bytes, boundary.as_str())) + (body.len() <= decoded_limit).then_some(body) } fn parse_multipart_fields(body: &[u8], boundary: &str) -> Vec { let delimiter = format!("--{boundary}").into_bytes(); let mut parts = Vec::new(); let mut cursor = 0usize; + let mut part_count = 0usize; - while let Some(index) = find_subslice(&body[cursor..], &delimiter) { + while let Some(index) = find_multipart_boundary(&body[cursor..], &delimiter, true) { let start = cursor + index + delimiter.len(); if body.get(start..start + 2) == Some(b"--") { + let closing_suffix = body.get(start + 2..).unwrap_or_default(); + if !(closing_suffix.is_empty() || closing_suffix.starts_with(b"\r\n")) { + return Vec::new(); + } break; } + part_count = part_count.saturating_add(1); + if part_count > MAX_MULTIPART_PARTS { + return Vec::new(); + } let mut part = &body[start..]; if part.starts_with(b"\r\n") { part = &part[2..]; } - let Some(next) = find_subslice(part, &delimiter) else { - break; + // Do not return fields parsed before a truncated part. Callers use + // an empty result as the invalid-multipart signal, so retaining a + // prefix would turn malformed input into an accepted request. + let Some(next) = find_multipart_boundary(part, &delimiter, false) else { + return Vec::new(); }; let raw = &part[..next]; let raw = raw.strip_suffix(b"\r\n").unwrap_or(raw); - if let Some(field) = parse_multipart_field(raw) { - parts.push(field); + if find_subslice(raw, b"\r\n\r\n") + .is_some_and(|header_end| header_end > MAX_MULTIPART_PART_HEADER_BYTES) + { + return Vec::new(); } + let Some(field) = parse_multipart_field(raw) else { + return Vec::new(); + }; + parts.push(field); cursor = start + next; } @@ -1688,34 +1726,219 @@ fn parse_multipart_fields(body: &[u8], boundary: &str) -> Vec { } fn multipart_boundary(content_type: &str) -> Option { - content_type.split(';').find_map(|segment| { - let (key, value) = segment.trim().split_once('=')?; - if !key.trim().eq_ignore_ascii_case("boundary") { + let segments = split_multipart_content_type_parameters(content_type)?; + let media_type = segments.first()?.trim(); + if !media_type.eq_ignore_ascii_case("multipart/form-data") { + return None; + } + + let mut boundary = None; + let mut seen_keys = Vec::new(); + for segment in segments.into_iter().skip(1) { + let segment = segment.trim(); + if segment.is_empty() { return None; } - let boundary = value.trim().trim_matches('"').trim(); - (!boundary.is_empty()).then(|| boundary.to_string()) - }) + let (raw_key, raw_value) = segment.split_once('=')?; + let key = raw_key.trim(); + if key.is_empty() || !key.as_bytes().iter().copied().all(is_http_token_byte) { + return None; + } + if seen_keys + .iter() + .any(|seen: &String| seen.eq_ignore_ascii_case(key)) + { + return None; + } + seen_keys.push(key.to_ascii_lowercase()); + + let (value, had_escape) = parse_multipart_content_type_parameter_value(raw_value.trim())?; + if !key.eq_ignore_ascii_case("boundary") { + continue; + } + if had_escape || !is_valid_multipart_boundary(&value) { + return None; + } + boundary = Some(value); + } + + boundary +} + +fn split_multipart_content_type_parameters(value: &str) -> Option> { + let mut segments = Vec::new(); + let mut start = 0usize; + let mut in_quotes = false; + let mut escaped = false; + + for (index, character) in value.char_indices() { + if character.is_ascii_control() { + return None; + } + if in_quotes { + if escaped { + escaped = false; + } else if character == '\\' { + escaped = true; + } else if character == '"' { + in_quotes = false; + } + } else if character == '"' { + in_quotes = true; + } else if character == ';' { + segments.push(&value[start..index]); + start = index + character.len_utf8(); + } + } + + if in_quotes || escaped { + return None; + } + segments.push(&value[start..]); + Some(segments) +} + +fn parse_multipart_content_type_parameter_value(value: &str) -> Option<(String, bool)> { + if value.is_empty() { + return None; + } + if value.starts_with('"') { + if value.len() < 2 || !value.ends_with('"') { + return None; + } + let inner = &value[1..value.len() - 1]; + let mut parsed = String::with_capacity(inner.len()); + let mut escaped = false; + let mut had_escape = false; + for character in inner.chars() { + if escaped { + if character.is_ascii_control() { + return None; + } + parsed.push(character); + escaped = false; + had_escape = true; + } else if character == '\\' { + escaped = true; + } else { + if character == '"' || character.is_ascii_control() { + return None; + } + parsed.push(character); + } + } + if escaped { + return None; + } + return Some((parsed, had_escape)); + } + + value + .as_bytes() + .iter() + .copied() + .all(is_http_token_byte) + .then(|| (value.to_string(), false)) +} + +fn find_multipart_boundary(haystack: &[u8], delimiter: &[u8], allow_start: bool) -> Option { + if delimiter.is_empty() { + return None; + } + let mut search_start = 0usize; + while search_start <= haystack.len() { + let relative = find_subslice(&haystack[search_start..], delimiter)?; + let index = search_start + relative; + let at_line_start = index == 0 + || (index >= 2 + && haystack + .get(index - 2..index) + .is_some_and(|prefix| prefix == b"\r\n")); + let allowed_position = if allow_start { + at_line_start + } else { + index >= 2 + && haystack + .get(index - 2..index) + .is_some_and(|prefix| prefix == b"\r\n") + }; + if allowed_position && multipart_boundary_suffix_is_valid(haystack, index, delimiter) { + return Some(index); + } + search_start = index.saturating_add(1); + } + None +} + +fn multipart_boundary_suffix_is_valid(haystack: &[u8], index: usize, delimiter: &[u8]) -> bool { + let suffix_start = index.saturating_add(delimiter.len()); + let Some(suffix) = haystack.get(suffix_start..) else { + return false; + }; + suffix.starts_with(b"\r\n") + || suffix + .strip_prefix(b"--") + .is_some_and(|remaining| remaining.is_empty() || remaining.starts_with(b"\r\n")) +} + +const MAX_MULTIPART_BOUNDARY_BYTES: usize = 70; + +fn is_valid_multipart_boundary(value: &str) -> bool { + !value.is_empty() + && value.len() <= MAX_MULTIPART_BOUNDARY_BYTES + && value.as_bytes().iter().copied().all(is_http_token_byte) +} + +fn is_http_token_byte(byte: u8) -> bool { + matches!( + byte, + b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + | b'!' + | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) } fn parse_multipart_field(raw: &[u8]) -> Option { let header_end = find_subslice(raw, b"\r\n\r\n")?; let headers = &raw[..header_end]; let data = raw.get(header_end + 4..)?.to_vec(); - let header_text = String::from_utf8_lossy(headers); + let header_text = std::str::from_utf8(headers).ok()?; let mut name = None; let mut content_type = None; - for line in header_text.lines() { - let trimmed = line.trim(); - let lower = trimmed.to_ascii_lowercase(); - if lower.starts_with("content-disposition:") { - name = extract_quoted_header_value(trimmed, "name"); - } else if lower.starts_with("content-type:") { - content_type = trimmed - .split_once(':') - .map(|(_, value)| value.trim().to_string()) - .filter(|value| !value.is_empty()); + let mut disposition_seen = false; + let mut content_type_seen = false; + for line in header_text.split("\r\n") { + let (header_name, header_value) = line.split_once(':')?; + let header_name = header_name.trim(); + let header_value = header_value.trim(); + if header_name.eq_ignore_ascii_case("content-disposition") { + if disposition_seen { + return None; + } + disposition_seen = true; + name = parse_multipart_content_disposition_name(header_value); + } else if header_name.eq_ignore_ascii_case("content-type") { + if content_type_seen || header_value.is_empty() { + return None; + } + content_type_seen = true; + content_type = Some(header_value.to_string()); } } @@ -1726,12 +1949,114 @@ fn parse_multipart_field(raw: &[u8]) -> Option { }) } -fn extract_quoted_header_value(header: &str, key: &str) -> Option { - let pattern = format!("{key}=\""); - let start = header.find(&pattern)? + pattern.len(); - let rest = &header[start..]; - let end = rest.find('"')?; - Some(rest[..end].to_string()) +fn parse_multipart_content_disposition_name(value: &str) -> Option { + let segments = split_multipart_header_parameters(value)?; + let disposition = segments.first()?.trim(); + if !disposition.eq_ignore_ascii_case("form-data") { + return None; + } + + let mut seen_keys = Vec::new(); + let mut name = None; + for segment in segments.into_iter().skip(1) { + let segment = segment.trim(); + if segment.is_empty() { + return None; + } + let (raw_key, raw_value) = segment.split_once('=')?; + let key = raw_key.trim(); + if key.is_empty() || !key.as_bytes().iter().copied().all(is_http_token_byte) { + return None; + } + if seen_keys + .iter() + .any(|seen: &String| seen.eq_ignore_ascii_case(key)) + { + return None; + } + seen_keys.push(key.to_ascii_lowercase()); + + let parsed_value = parse_multipart_header_parameter_value(raw_value.trim())?; + if key.eq_ignore_ascii_case("name") { + if parsed_value.is_empty() { + return None; + } + name = Some(parsed_value); + } + } + + name +} + +fn split_multipart_header_parameters(value: &str) -> Option> { + let mut segments = Vec::new(); + let mut start = 0usize; + let mut in_quotes = false; + let mut escaped = false; + + for (index, byte) in value.as_bytes().iter().copied().enumerate() { + if in_quotes { + if escaped { + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + in_quotes = false; + } + } else if byte == b'"' { + in_quotes = true; + } else if byte == b';' { + segments.push(&value[start..index]); + start = index + 1; + } + } + + if in_quotes || escaped { + return None; + } + segments.push(&value[start..]); + Some(segments) +} + +fn parse_multipart_header_parameter_value(value: &str) -> Option { + if value.is_empty() { + return None; + } + if value.starts_with('"') { + if value.len() < 2 || !value.ends_with('"') { + return None; + } + let inner = &value[1..value.len() - 1]; + let mut parsed = String::with_capacity(inner.len()); + let mut escaped = false; + for character in inner.chars() { + if escaped { + if character.is_control() { + return None; + } + parsed.push(character); + escaped = false; + } else if character == '\\' { + escaped = true; + } else { + if character == '"' || character.is_control() { + return None; + } + parsed.push(character); + } + } + if escaped { + return None; + } + return Some(parsed); + } + + value + .as_bytes() + .iter() + .copied() + .all(is_http_token_byte) + .then(|| value.to_string()) } fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option { @@ -1752,10 +2077,13 @@ mod tests { use super::{ build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body, build_openai_image_api_provider_request_body, build_openai_image_provider_request_body, - is_openai_image_stream_request, normalize_openai_image_quality, + decode_multipart_body_base64_with_limit, find_multipart_boundary, + is_openai_image_stream_request, multipart_boundary, normalize_openai_image_quality, normalize_openai_image_request, normalize_openai_image_request_with_options, - openai_image_operation_from_path, project_codex_openai_image_api_request_body, - project_openai_image_api_request_body, OpenAiImageNormalizeOptions, OpenAiImageOperation, + openai_image_operation_from_path, parse_multipart_fields, + project_codex_openai_image_api_request_body, project_openai_image_api_request_body, + OpenAiImageNormalizeOptions, OpenAiImageOperation, MAX_MULTIPART_BOUNDARY_BYTES, + MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES, }; use crate::formats::openai::image::spec::{resolve_stream_spec, resolve_sync_spec}; @@ -1833,6 +2161,210 @@ mod tests { )); } + #[test] + fn multipart_boundary_requires_rfc_token_and_length_limits() { + assert_eq!( + multipart_boundary("Multipart/Form-Data; boundary=quoted-token-123").as_deref(), + Some("quoted-token-123") + ); + assert_eq!( + multipart_boundary("multipart/form-data; boundary=\"quoted-token-123\"").as_deref(), + Some("quoted-token-123") + ); + + for content_type in [ + "multipart/form-data; boundary=bad boundary", + "multipart/form-data; boundary=bad\"quote", + "multipart/form-data; boundary=\"unterminated", + "multipart/form-data; boundary=first; boundary=second", + "multipart/form-data; foo", + "multipart/form-data; foo=\"unterminated; boundary=valid", + "multipart/form-data; boundary=valid trailing", + "application/json; boundary=valid-token", + ] { + assert!( + multipart_boundary(content_type).is_none(), + "{content_type:?}" + ); + } + + let oversized = "a".repeat(MAX_MULTIPART_BOUNDARY_BYTES + 1); + assert!( + multipart_boundary(&format!("multipart/form-data; boundary={oversized}")).is_none() + ); + } + + #[test] + fn multipart_boundary_accepts_quoted_unknown_parameter_with_semicolon() { + assert_eq!( + multipart_boundary("multipart/form-data; note=\"semi;colon\"; boundary=quoted-token") + .as_deref(), + Some("quoted-token") + ); + } + + #[test] + fn multipart_boundary_rejects_escaped_or_duplicate_content_type_parameters() { + for content_type in [ + "multipart/form-data; boundary=\"escaped\\\"token\"", + "multipart/form-data; note=one; NOTE=two; boundary=token", + "multipart/form-data; note=\"unterminated; boundary=token", + "multipart/form-data; note=\"closed\"trailing; boundary=token", + ] { + assert!( + multipart_boundary(content_type).is_none(), + "malformed content type must be rejected: {content_type:?}" + ); + } + } + + #[test] + fn multipart_parser_caps_part_count_and_header_size() { + let boundary = "bounded-parts"; + let mut accepted_body = Vec::new(); + for index in 0..MAX_MULTIPART_PARTS { + accepted_body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"field-{index}\"\r\n\r\nvalue\r\n" + ) + .as_bytes(), + ); + } + accepted_body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes()); + assert_eq!( + parse_multipart_fields(&accepted_body, boundary).len(), + MAX_MULTIPART_PARTS + ); + + let mut body = Vec::new(); + for index in 0..(MAX_MULTIPART_PARTS + 1) { + body.extend_from_slice( + format!( + "--{boundary}\r\nContent-Disposition: form-data; name=\"field-{index}\"\r\n\r\nvalue\r\n" + ) + .as_bytes(), + ); + } + body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes()); + assert!(parse_multipart_fields(&body, boundary).is_empty()); + + let mut oversized_header = + format!("--{boundary}\r\nContent-Disposition: form-data; name=\"field\"; x=\"") + .into_bytes(); + oversized_header.extend(std::iter::repeat_n(b'x', MAX_MULTIPART_PART_HEADER_BYTES)); + oversized_header + .extend_from_slice(format!("\"\r\n\r\nvalue\r\n--{boundary}--\r\n").as_bytes()); + assert!(parse_multipart_fields(&oversized_header, boundary).is_empty()); + } + + #[test] + fn multipart_body_base64_decode_enforces_allocation_limit() { + let exact = vec![b'x'; 6]; + let exact_encoded = base64::engine::general_purpose::STANDARD.encode(&exact); + assert_eq!( + decode_multipart_body_base64_with_limit(&exact_encoded, exact.len()), + Some(exact) + ); + + let oversized = vec![b'x'; 7]; + let oversized_encoded = base64::engine::general_purpose::STANDARD.encode(oversized); + assert_eq!( + decode_multipart_body_base64_with_limit(&oversized_encoded, 6), + None + ); + + let same_encoded_bucket = base64::engine::general_purpose::STANDARD.encode([b'x'; 6]); + assert_eq!( + decode_multipart_body_base64_with_limit(&same_encoded_bucket, 5), + None + ); + } + + #[test] + fn multipart_parser_preserves_boundary_like_payload_and_fails_closed() { + let boundary = "payload-boundary"; + let body = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"\r\n\r\n", + "prefix\r\n--{boundary}X\r\nsuffix--{boundary}\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + let fields = parse_multipart_fields(body.as_bytes(), boundary); + assert_eq!(fields.len(), 1); + assert_eq!(fields[0].name, "prompt"); + assert_eq!( + fields[0].data, + format!("prefix\r\n--{boundary}X\r\nsuffix--{boundary}").into_bytes() + ); + + let malformed = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"first\"\r\n\r\n", + "ok\r\n", + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"second\"\r\n\r\n", + "truncated\r\n--{boundary}X\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(malformed.as_bytes(), boundary).is_empty()); + } + + #[test] + fn multipart_parser_does_not_extract_name_from_filename_and_rejects_duplicates() { + let boundary = "header-parameters"; + let filename_only = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; filename=\"name=\\\"prompt\\\"\"\r\n\r\n", + "attacker-value\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(filename_only.as_bytes(), boundary).is_empty()); + + let duplicate_name = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"; name=\"image\"\r\n\r\n", + "ambiguous-value\r\n", + "--{boundary}--\r\n" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(duplicate_name.as_bytes(), boundary).is_empty()); + } + + #[test] + fn multipart_parser_rejects_garbage_after_closing_boundary() { + let boundary = "closing-suffix"; + let body = format!( + concat!( + "--{boundary}\r\n", + "Content-Disposition: form-data; name=\"prompt\"\r\n\r\n", + "value\r\n", + "--{boundary}--junk" + ), + boundary = boundary, + ); + assert!(parse_multipart_fields(body.as_bytes(), boundary).is_empty()); + } + + #[test] + fn multipart_scanner_skips_short_prefix_before_later_valid_boundary() { + let delimiter = b"--scanner-boundary"; + let haystack = b"x--scanner-boundary\r\ncontent\r\n--scanner-boundary\r\n"; + assert_eq!( + find_multipart_boundary(haystack, delimiter, false), + Some(haystack.len() - delimiter.len() - 2) + ); + } + #[test] fn openai_image_variation_path_is_not_supported() { let boundary = "boundary-variation-123"; diff --git a/crates/aether-ai/formats/src/formats/openai/image/stream.rs b/crates/aether-ai/formats/src/formats/openai/image/stream.rs index 93eca6b93..5a2868420 100644 --- a/crates/aether-ai/formats/src/formats/openai/image/stream.rs +++ b/crates/aether-ai/formats/src/formats/openai/image/stream.rs @@ -1,17 +1,33 @@ use std::collections::BTreeSet; use aether_contracts::{ExecutionStreamTerminalSummary, StandardizedUsage}; -use base64::Engine as _; use serde_json::{Map, Value}; +use sha2::{Digest as _, Sha256}; use crate::contracts::OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND; +use crate::formats::openai::image::{ + bounded_openai_image_revised_prompt, is_safe_openai_image_base64_payload, + normalize_openai_image_output_format, +}; use crate::formats::openai::responses::codex::CODEX_OPENAI_IMAGE_DEFAULT_OUTPUT_FORMAT; use crate::formats::shared::sse::{encode_done_sse, encode_json_sse}; use crate::formats::shared::stream_core::common::{ build_openai_chat_chunk, build_openai_chat_finish_chunk, build_openai_chat_usage_chunk_with_cache, }; -use crate::formats::shared::AiSurfaceFinalizeError; +use crate::formats::shared::{decode_sync_report_body_base64, AiSurfaceFinalizeError}; + +// Bound parser carry state while still allowing the largest supported image +// records. A 3840x2160 RGBA image is about 33 MiB before base64 encoding, so +// 64 MiB leaves room for encoding and event metadata. This is a per-stream +// parser limit, not a response-body or concurrency limit. +const MAX_STREAM_REWRITE_BUFFER_BYTES: usize = 64 * 1024 * 1024; + +// OpenAI image responses currently allow at most a small number of output +// images per request. Keep enough room for valid multi-image responses while +// preventing an untrusted provider from growing the de-duplication sets +// without bound. +const MAX_IMAGE_OUTPUT_KEYS: usize = 64; #[derive(Default)] pub struct OpenAiImageStreamState { @@ -49,7 +65,8 @@ struct OpenAiImageChatFrame { #[derive(Default)] pub struct OpenAiImageStreamTerminalState { event_name: Option, - data_lines: Vec, + data: Option, + buffered_data_bytes: usize, response_id: Option, model: Option, image_count: u64, @@ -65,14 +82,15 @@ impl OpenAiImageStreamState { report_context: &Value, chunk: &[u8], ) -> Result, AiSurfaceFinalizeError> { - self.buffered.extend_from_slice(chunk); - let mut output = Vec::new(); - while let Some(block_end) = find_sse_block_end(&self.buffered) { - let block = self.buffered.drain(..block_end).collect::>(); - output.extend(self.transform_block(report_context, &block)?); - drain_sse_separator(&mut self.buffered); - } - Ok(output) + let mut buffered = std::mem::take(&mut self.buffered); + let result = process_bounded_sse_chunk( + &mut buffered, + chunk, + MAX_STREAM_REWRITE_BUFFER_BYTES, + |block| self.transform_block(report_context, block), + ); + self.buffered = buffered; + result } pub fn finish(&mut self, report_context: &Value) -> Result, AiSurfaceFinalizeError> { @@ -91,16 +109,21 @@ impl OpenAiImageStreamState { let text = std::str::from_utf8(block) .map_err(|err| AiSurfaceFinalizeError::new(err.to_string()))?; let mut event_name = None::; - let mut data_lines = Vec::new(); + let mut data = None::; for raw_line in text.lines() { let line = raw_line.trim_end_matches('\r'); if let Some(value) = line.strip_prefix("event:") { event_name = Some(value.trim().to_string()); } else if let Some(value) = line.strip_prefix("data:") { - data_lines.push(value.trim().to_string()); + let had_data = data.is_some(); + let data_value = data.get_or_insert_with(String::new); + if had_data { + data_value.push('\n'); + } + data_value.push_str(value.trim()); } } - let data = data_lines.join("\n"); + let data = data.unwrap_or_default(); if data.is_empty() || data == "[DONE]" { return Ok(Vec::new()); } @@ -137,7 +160,7 @@ impl OpenAiImageStreamState { .or_else(|| event.get("b64_json")) .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) else { return Ok(Vec::new()); }; @@ -181,7 +204,7 @@ impl OpenAiImageStreamState { let Some(result) = item.get("result").and_then(Value::as_str).map(str::trim) else { return Ok(Vec::new()); }; - if result.is_empty() { + if !is_safe_openai_image_base64_payload(result) { return Ok(Vec::new()); } self.latest_image = Some(OpenAiImageFrame { @@ -274,14 +297,15 @@ impl OpenAiImageChatStreamState { report_context: &Value, chunk: &[u8], ) -> Result, AiSurfaceFinalizeError> { - self.buffered.extend_from_slice(chunk); - let mut output = Vec::new(); - while let Some(block_end) = find_sse_block_end(&self.buffered) { - let block = self.buffered.drain(..block_end).collect::>(); - output.extend(self.transform_block(report_context, &block)?); - drain_sse_separator(&mut self.buffered); - } - Ok(output) + let mut buffered = std::mem::take(&mut self.buffered); + let result = process_bounded_sse_chunk( + &mut buffered, + chunk, + MAX_STREAM_REWRITE_BUFFER_BYTES, + |block| self.transform_block(report_context, block), + ); + self.buffered = buffered; + result } pub fn finish(&mut self, report_context: &Value) -> Result, AiSurfaceFinalizeError> { @@ -305,16 +329,21 @@ impl OpenAiImageChatStreamState { let text = std::str::from_utf8(block) .map_err(|err| AiSurfaceFinalizeError::new(err.to_string()))?; let mut event_name = None::; - let mut data_lines = Vec::new(); + let mut data = None::; for raw_line in text.lines() { let line = raw_line.trim_end_matches('\r'); if let Some(value) = line.strip_prefix("event:") { event_name = Some(value.trim().to_string()); } else if let Some(value) = line.strip_prefix("data:") { - data_lines.push(value.trim().to_string()); + let had_data = data.is_some(); + let data_value = data.get_or_insert_with(String::new); + if had_data { + data_value.push('\n'); + } + data_value.push_str(value.trim()); } } - let data = data_lines.join("\n"); + let data = data.unwrap_or_default(); if data.is_empty() || data == "[DONE]" { return Ok(Vec::new()); } @@ -355,16 +384,15 @@ impl OpenAiImageChatStreamState { return Ok(Vec::new()); } if let Some(result) = item.get("result").and_then(Value::as_str).map(str::trim) { - if !result.is_empty() { + if is_safe_openai_image_base64_payload(result) { let key = image_chat_output_key(item, result); - if self.emitted_image_keys.insert(key) { + if insert_bounded_image_key(&mut self.emitted_image_keys, key) { self.latest_image = Some(OpenAiImageChatFrame { b64_json: result.to_string(), output_format: item .get("output_format") .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(normalize_openai_image_output_format) .map(ToOwned::to_owned), }); self.emitted_image_count = self.emitted_image_count.saturating_add(1); @@ -417,15 +445,14 @@ impl OpenAiImageChatStreamState { .or_else(|| event.get("result")) .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) { self.latest_image = Some(OpenAiImageChatFrame { b64_json: result.to_string(), output_format: event .get("output_format") .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(normalize_openai_image_output_format) .map(ToOwned::to_owned), }); self.emitted_image_count = self.emitted_image_count.max(1); @@ -592,7 +619,11 @@ impl OpenAiImageStreamTerminalState { if let Some(value) = trimmed.strip_prefix("event:") { self.event_name = Some(value.trim().to_string()); } else if let Some(value) = trimmed.strip_prefix("data:") { - self.data_lines.push(value.trim().to_string()); + append_bounded_sse_data_line( + &mut self.data, + &mut self.buffered_data_bytes, + value.trim(), + )?; } Ok(self.latest_summary(report_context)) } @@ -609,11 +640,12 @@ impl OpenAiImageStreamTerminalState { } fn flush_event(&mut self, report_context: &Value) -> Result<(), AiSurfaceFinalizeError> { - if self.data_lines.is_empty() { + let Some(data) = self.data.take() else { + self.buffered_data_bytes = 0; self.event_name = None; return Ok(()); - } - let data = std::mem::take(&mut self.data_lines).join("\n"); + }; + self.buffered_data_bytes = 0; let event_name = self.event_name.take(); if data.is_empty() || data == "[DONE]" { return Ok(()); @@ -660,12 +692,12 @@ impl OpenAiImageStreamTerminalState { .get("result") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) else { return; }; let key = image_chat_output_key(item, result); - if self.image_keys.insert(key) { + if insert_bounded_image_key(&mut self.image_keys, key) { self.image_count = self.image_count.saturating_add(1); } } @@ -695,7 +727,7 @@ impl OpenAiImageStreamTerminalState { .or_else(|| event.get("result")) .and_then(Value::as_str) .map(str::trim) - .is_some_and(|value| !value.is_empty()) + .is_some_and(is_safe_openai_image_base64_payload) { self.image_count = 1; } @@ -759,7 +791,7 @@ fn completed_response_image_chat_frame(response: &Value) -> Option Option u64 { item.get("result") .and_then(Value::as_str) .map(str::trim) - .is_some_and(|value| !value.is_empty()) + .is_some_and(is_safe_openai_image_base64_payload) }) + .take(MAX_IMAGE_OUTPUT_KEYS) .count() as u64 } fn image_chat_output_key(item: &Map, result: &str) -> String { - item.get("id") + let (source, value) = item + .get("id") .or_else(|| item.get("call_id")) .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .unwrap_or_else(|| result.to_string()) + .map(|value| ("id", value)) + .unwrap_or(("result", result.trim())); + let mut digest = Sha256::new(); + digest.update(source.as_bytes()); + digest.update([0]); + digest.update(value.as_bytes()); + format!("{source}:{:x}", digest.finalize()) +} + +fn insert_bounded_image_key(keys: &mut BTreeSet, key: String) -> bool { + if keys.contains(&key) || keys.len() >= MAX_IMAGE_OUTPUT_KEYS { + return false; + } + keys.insert(key) } fn openai_image_stream_standardized_usage( @@ -895,19 +940,13 @@ fn openai_image_usage_to_standardized_usage(value: &Value) -> Option String { - let mime_type = match frame - .output_format - .as_deref() - .unwrap_or("png") - .trim() - .to_ascii_lowercase() - .as_str() - { - "jpg" | "jpeg" => "image/jpeg".to_string(), - "webp" => "image/webp".to_string(), - "png" => "image/png".to_string(), - value if !value.is_empty() => format!("image/{value}"), - _ => "image/png".to_string(), + let mime_type = match frame.output_format.as_deref().map(str::trim) { + Some(value) if value.eq_ignore_ascii_case("jpg") || value.eq_ignore_ascii_case("jpeg") => { + "image/jpeg" + } + Some(value) if value.eq_ignore_ascii_case("webp") => "image/webp", + Some(value) if value.eq_ignore_ascii_case("png") => "image/png", + _ => "image/png", }; format!( "![generated image](data:{mime_type};base64,{})", @@ -1040,7 +1079,7 @@ fn completed_response_image_result(event: &Value) -> Option<&str> { .filter(|item| item.get("type").and_then(Value::as_str) == Some("image_generation_call")) .filter_map(|item| item.get("result").and_then(Value::as_str)) .map(str::trim) - .find(|value| !value.is_empty()) + .find(|value| is_safe_openai_image_base64_payload(value)) } fn requested_partial_images(report_context: &Value) -> u64 { @@ -1126,6 +1165,82 @@ fn image_bridge_model(report_context: Option<&Value>) -> Option { }) } +fn process_bounded_sse_chunk( + buffered: &mut Vec, + chunk: &[u8], + max_bytes: usize, + mut transform: F, +) -> Result, AiSurfaceFinalizeError> +where + F: FnMut(&[u8]) -> Result, AiSurfaceFinalizeError>, +{ + let mut remaining = chunk; + let mut output = Vec::new(); + loop { + if let Some(block_end) = find_sse_block_end(buffered) { + let block = buffered.drain(..block_end).collect::>(); + output.extend(transform(&block)?); + drain_sse_separator(buffered); + continue; + } + if remaining.is_empty() { + break; + } + let available = max_bytes.saturating_sub(buffered.len()); + if available == 0 { + return Err(AiSurfaceFinalizeError::new(format!( + "image stream buffer exceeds {max_bytes} bytes" + ))); + } + let take = available.min(remaining.len()); + append_bounded_stream_rewrite_chunk(buffered, &remaining[..take], max_bytes)?; + remaining = &remaining[take..]; + } + Ok(output) +} + +fn append_bounded_stream_rewrite_chunk( + buffered: &mut Vec, + chunk: &[u8], + max_bytes: usize, +) -> Result<(), AiSurfaceFinalizeError> { + let next_len = buffered + .len() + .checked_add(chunk.len()) + .ok_or_else(|| AiSurfaceFinalizeError::new("image stream buffer length overflow"))?; + if next_len > max_bytes { + return Err(AiSurfaceFinalizeError::new(format!( + "image stream buffer exceeds {max_bytes} bytes" + ))); + } + buffered.extend_from_slice(chunk); + Ok(()) +} + +fn append_bounded_sse_data_line( + data: &mut Option, + buffered_data_bytes: &mut usize, + value: &str, +) -> Result<(), AiSurfaceFinalizeError> { + let separator_bytes = usize::from(data.is_some()); + let next_len = buffered_data_bytes + .checked_add(separator_bytes) + .and_then(|length| length.checked_add(value.len())) + .ok_or_else(|| AiSurfaceFinalizeError::new("image stream data buffer length overflow"))?; + if next_len > MAX_STREAM_REWRITE_BUFFER_BYTES { + return Err(AiSurfaceFinalizeError::new(format!( + "image stream data buffer exceeds {MAX_STREAM_REWRITE_BUFFER_BYTES} bytes" + ))); + } + let data_value = data.get_or_insert_with(String::new); + if separator_bytes != 0 { + data_value.push('\n'); + } + data_value.push_str(value); + *buffered_data_bytes = next_len; + Ok(()) +} + fn find_sse_block_end(buffer: &[u8]) -> Option { buffer .windows(2) @@ -1198,8 +1313,16 @@ pub fn maybe_build_openai_image_sync_finalize_product( } if let Some(provider_body_json) = body_json { if openai_image_response_has_standard_data(provider_body_json) { + let Some(client_body_json) = crate::formats::shared::image_bridge:: + build_openai_image_response_from_standard_image_response( + provider_body_json, + Some(report_context), + ) + else { + return Ok(None); + }; return Ok(Some(OpenAiImageSyncFinalizeProduct { - client_body_json: provider_body_json.clone(), + client_body_json, provider_body_json: provider_body_json.clone(), })); } @@ -1223,16 +1346,16 @@ pub fn maybe_build_openai_image_sync_finalize_product( .get("image_request") .and_then(|value| value.get("output_format")) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(normalize_openai_image_output_format) .unwrap_or(CODEX_OPENAI_IMAGE_DEFAULT_OUTPUT_FORMAT); - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; let text = std::str::from_utf8(&body_bytes) .map_err(|err| AiSurfaceFinalizeError::new(err.to_string()))?; let mut created = None; let mut completed_response = None; let mut images = Vec::new(); + let mut image_keys = BTreeSet::new(); for raw_block in text.split("\n\n") { let block = raw_block.trim(); @@ -1271,10 +1394,30 @@ pub fn maybe_build_openai_image_sync_finalize_product( let Some(result) = item.get("result").and_then(Value::as_str) else { continue; }; + let result = result.trim(); + if !is_safe_openai_image_base64_payload(result) + || !insert_bounded_image_key( + &mut image_keys, + image_chat_output_key(item, result), + ) + { + continue; + } + let output_format = item + .get("output_format") + .and_then(Value::as_str) + .and_then(normalize_openai_image_output_format) + .unwrap_or(default_output_format); + let revised_prompt = item + .get("revised_prompt") + .and_then(Value::as_str) + .and_then(bounded_openai_image_revised_prompt) + .map(|value| Value::String(value.to_string())) + .unwrap_or(Value::Null); images.push(serde_json::json!({ "b64_json": result, - "output_format": item.get("output_format").cloned().unwrap_or(Value::String(default_output_format.to_string())), - "revised_prompt": item.get("revised_prompt").cloned().unwrap_or(Value::Null), + "output_format": output_format, + "revised_prompt": revised_prompt, })); } "response.completed" => { @@ -1357,10 +1500,16 @@ fn openai_image_response_has_standard_data(body_json: &Value) -> bool { #[cfg(test)] mod tests { - use base64::Engine as _; - use serde_json::json; + use std::collections::BTreeSet; - use super::{maybe_build_openai_image_sync_finalize_product, OpenAiImageStreamState}; + use base64::Engine as _; + use serde_json::{json, Value}; + + use super::{ + completed_response_image_count, image_chat_output_key, insert_bounded_image_key, + maybe_build_openai_image_sync_finalize_product, process_bounded_sse_chunk, + OpenAiImageChatStreamState, OpenAiImageStreamState, OpenAiImageStreamTerminalState, + }; fn utf8(bytes: Vec) -> String { String::from_utf8(bytes).expect("utf8 should decode") @@ -1514,6 +1663,308 @@ mod tests { .is_empty()); } + #[test] + fn image_stream_rejects_unbounded_incomplete_sse_buffer() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:image", + "needs_conversion": false, + "image_request": {"operation": "generate"} + }); + let mut rewriter = OpenAiImageStreamState::default(); + let oversized = vec![b'x'; super::MAX_STREAM_REWRITE_BUFFER_BYTES + 1]; + + let error = rewriter + .push_chunk(&report_context, &oversized) + .expect_err("incomplete image SSE block must be bounded"); + assert!(error.0.contains("image stream buffer exceeds")); + } + + #[test] + fn image_chat_stream_rejects_unbounded_incomplete_sse_buffer() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:chat", + "needs_conversion": true, + "image_request": {"operation": "generate"} + }); + let mut rewriter = OpenAiImageChatStreamState::default(); + let oversized = vec![b'x'; super::MAX_STREAM_REWRITE_BUFFER_BYTES + 1]; + + let error = rewriter + .push_chunk(&report_context, &oversized) + .expect_err("incomplete image SSE block must be bounded"); + assert!(error.0.contains("image stream buffer exceeds")); + } + + #[test] + fn image_terminal_observer_rejects_unbounded_data_lines() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:image", + "needs_conversion": false, + "image_request": {"operation": "generate"} + }); + let mut observer = OpenAiImageStreamTerminalState::default(); + let first_len = super::MAX_STREAM_REWRITE_BUFFER_BYTES / 2; + let second_len = super::MAX_STREAM_REWRITE_BUFFER_BYTES - first_len; + let data_line = |length: usize| { + let mut line = Vec::with_capacity(6 + length); + line.extend_from_slice(b"data: "); + line.extend(std::iter::repeat_n(b'x', length)); + line.push(b'\n'); + line + }; + + observer + .push_line(&report_context, data_line(first_len)) + .expect("first data line should fit"); + let error = observer + .push_line(&report_context, data_line(second_len)) + .expect_err("data lines without a blank separator must be bounded"); + assert!(error.0.contains("image stream data buffer exceeds")); + } + + #[test] + fn image_terminal_observer_preserves_multiline_data_semantics() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:image", + "needs_conversion": false, + "image_request": {"operation": "generate"} + }); + let mut observer = OpenAiImageStreamTerminalState::default(); + + // An empty first data line must still contribute the SSE newline when + // the following line contains the JSON event. + observer + .push_line(&report_context, b"data:\n".to_vec()) + .expect("empty data line should be accepted"); + observer + .push_line( + &report_context, + b"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_multiline\",\"model\":\"gpt-image-2\"}}\n".to_vec(), + ) + .expect("second data line should be accepted"); + observer + .push_line(&report_context, b"\n".to_vec()) + .expect("event separator should flush"); + + let summary = observer + .finish(&report_context) + .expect("observer finish should succeed") + .expect("completed event should produce a summary"); + assert_eq!(summary.response_id.as_deref(), Some("resp_multiline")); + assert_eq!(summary.model.as_deref(), Some("gpt-image-2")); + assert!(summary.observed_finish); + } + + #[test] + fn image_terminal_observer_compacts_empty_data_lines() { + let mut observer = OpenAiImageStreamTerminalState::default(); + let mut data = None; + let mut buffered_bytes = 0; + for _ in 0..10_000 { + super::append_bounded_sse_data_line(&mut data, &mut buffered_bytes, "") + .expect("empty data line should fit"); + } + + assert_eq!(buffered_bytes, 9_999); + let expected = "\n".repeat(9_999); + assert_eq!(data.as_deref(), Some(expected.as_str())); + // Keep the state type exercised as well; this guards against changing + // the compact representation back to per-line allocations. + observer.data = data; + observer.buffered_data_bytes = buffered_bytes; + assert_eq!(observer.data.as_ref().map(String::len), Some(9_999)); + } + + #[test] + fn image_output_keys_are_hashed_to_a_fixed_length() { + let long_id = "provider-id-".to_string() + &"x".repeat(32 * 1024); + let item_with_id = json!({"id": long_id}); + let id_key = image_chat_output_key( + item_with_id.as_object().expect("object item"), + "result-that-is-ignored-when-id-is-present", + ); + assert_eq!(id_key.len(), "id:".len() + 64); + assert!(id_key.starts_with("id:")); + assert!(!id_key.contains("provider-id-")); + + let long_result = "r".repeat(32 * 1024); + let item_without_id = json!({}); + let result_key = image_chat_output_key( + item_without_id.as_object().expect("object item"), + &long_result, + ); + assert_eq!(result_key.len(), "result:".len() + 64); + assert!(result_key.starts_with("result:")); + assert!(!result_key.contains(&long_result)); + } + + #[test] + fn image_output_key_set_is_bounded_and_keeps_dedupe_semantics() { + let mut keys = BTreeSet::new(); + let duplicate = "id:duplicate".to_string(); + assert!(insert_bounded_image_key(&mut keys, duplicate.clone())); + assert!(!insert_bounded_image_key(&mut keys, duplicate.clone())); + + for index in 1..super::MAX_IMAGE_OUTPUT_KEYS { + assert!(insert_bounded_image_key( + &mut keys, + format!("id:{index:064x}"), + )); + } + assert_eq!(keys.len(), super::MAX_IMAGE_OUTPUT_KEYS); + assert!(!insert_bounded_image_key( + &mut keys, + "id:overflow".to_string() + )); + assert_eq!(keys.len(), super::MAX_IMAGE_OUTPUT_KEYS); + } + + #[test] + fn image_chat_markdown_does_not_embed_untrusted_output_format() { + let frame = super::OpenAiImageChatFrame { + b64_json: "aGVsbG8=".to_string(), + output_format: Some("png);https://attacker.invalid/?x=(x".to_string()), + }; + + assert_eq!( + super::image_chat_markdown(&frame), + "![generated image](data:image/png;base64,aGVsbG8=)" + ); + } + + #[test] + fn completed_response_image_count_is_bounded() { + let output = (0..(super::MAX_IMAGE_OUTPUT_KEYS + 8)) + .map(|index| { + json!({ + "type": "image_generation_call", + "result": format!("aGVsbG{index:02x}"), + }) + }) + .collect::>(); + let response = json!({"output": output}); + + assert_eq!( + completed_response_image_count(&response), + super::MAX_IMAGE_OUTPUT_KEYS as u64 + ); + } + + #[test] + fn image_sync_finalize_dedupes_and_bounds_output_images() { + let report_context = json!({ + "client_api_format": "openai:image", + "provider_api_format": "openai:image", + "image_request": { + "operation": "generate", + "output_format": "png" + } + }); + let mut stream = String::new(); + let append_output_item = |stream: &mut String, id: &str, result: &str| { + stream.push_str("data: "); + stream.push_str( + &serde_json::to_string(&json!({ + "type": "response.output_item.done", + "item": { + "id": id, + "type": "image_generation_call", + "result": result, + } + })) + .expect("event should serialize"), + ); + stream.push_str("\n\n"); + }; + + append_output_item(&mut stream, "duplicate", "Zmlyc3QtaW1hZ2U="); + append_output_item(&mut stream, "duplicate", "c2Vjb25kLWltYWdl"); + for index in 0..super::MAX_IMAGE_OUTPUT_KEYS { + let id = format!("image-{index}"); + append_output_item(&mut stream, &id, "aGVsbG8="); + } + stream.push_str( + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-images\"}}\n\n", + ); + let body_base64 = base64::engine::general_purpose::STANDARD.encode(stream.as_bytes()); + + let product = maybe_build_openai_image_sync_finalize_product( + "openai_image_sync_finalize", + 200, + Some(&report_context), + None, + Some(&body_base64), + ) + .expect("finalize should succeed") + .expect("image stream should finalize"); + + assert_eq!( + product.client_body_json["data"] + .as_array() + .expect("client data array") + .len(), + super::MAX_IMAGE_OUTPUT_KEYS + ); + assert_eq!( + product.provider_body_json["output"] + .as_array() + .expect("provider output array") + .len(), + super::MAX_IMAGE_OUTPUT_KEYS + ); + assert_eq!( + product.client_body_json["data"][0]["b64_json"], + "Zmlyc3QtaW1hZ2U=" + ); + assert!(product.client_body_json["data"] + .as_array() + .expect("client data array") + .iter() + .all(|image| image["b64_json"] != "c2Vjb25kLWltYWdl")); + } + + #[test] + fn image_stream_consumes_complete_frames_before_chunk_limit() { + let frame = b"data: {\"type\":\"noop\"}\n\n"; + let mut chunk = Vec::with_capacity(frame.len() * 2); + chunk.extend_from_slice(frame); + chunk.extend_from_slice(frame); + let mut buffered = Vec::new(); + + let output = process_bounded_sse_chunk(&mut buffered, &chunk, frame.len(), |block| { + Ok(block.to_vec()) + }) + .expect("complete frames should be consumed even when the chunk is larger than the cap"); + + assert_eq!(output, chunk); + assert!(buffered.is_empty()); + } + + #[test] + fn image_stream_bounds_incomplete_frame_with_small_test_limit() { + let mut buffered = Vec::new(); + let error = process_bounded_sse_chunk(&mut buffered, b"123456789", 8, |_| Ok(Vec::new())) + .expect_err("an incomplete frame above the cap must be rejected"); + assert!(error.0.contains("image stream buffer exceeds 8 bytes")); + } + + #[test] + fn image_stream_buffer_budget_covers_gpt_image_2_max_resolution() { + // gpt-image-2 accepts up to 3840x2160. A worst-case raw RGBA payload + // still fits after base64 encoding with room for SSE/JSON metadata. + let raw_rgba_bytes = 3840usize * 2160 * 4; + let base64_bytes = raw_rgba_bytes.div_ceil(3) * 4; + let metadata_margin = 8 * 1024 * 1024; + assert!( + super::MAX_STREAM_REWRITE_BUFFER_BYTES >= base64_bytes + metadata_margin, + "image parser cap must cover the largest supported image payload" + ); + } + #[test] fn sync_finalize_product_maps_stream_response_to_client_and_provider_bodies() { let report_context = json!({ @@ -1594,4 +2045,87 @@ mod tests { assert_eq!(product.client_body_json["data"][0]["b64_json"], "aGVsbG8="); assert_eq!(product.provider_body_json, provider_body); } + + #[test] + fn sync_finalize_filters_untrusted_standard_openai_image_fields() { + let oversized_prompt = "p".repeat(256 * 1024 + 1); + let provider_body = json!({ + "created": 1779273523, + "model": "gpt-image-2", + "data": [ + {"url": "javascript:alert(1)"}, + {"url": "data:text/html;base64,PGh0bWw+"}, + { + "b64_json": "aGVsbG8=", + "output_format": "text/html;javascript:alert(1)", + "revised_prompt": oversized_prompt.clone() + }, + {"b64_json": "d29ybGQ=", "output_format": "JPG"} + ], + "usage": {"input_tokens": 1, "output_tokens": 2} + }); + let report_context = json!({ + "client_api_format": "openai:image", + "provider_api_format": "openai:image", + "image_request": {"operation": "generate", "output_format": "png"} + }); + + let product = maybe_build_openai_image_sync_finalize_product( + "openai_image_sync_finalize", + 200, + Some(&report_context), + Some(&provider_body), + None, + ) + .expect("standard image response should finalize") + .expect("at least one safe image should remain"); + + let client_data = product.client_body_json["data"] + .as_array() + .expect("client data array"); + assert_eq!(client_data.len(), 2); + assert_eq!(client_data[0]["b64_json"], "aGVsbG8="); + assert_eq!(client_data[0]["revised_prompt"], Value::Null); + assert!(client_data[0].get("output_format").is_none()); + assert_eq!(client_data[1]["b64_json"], "d29ybGQ="); + assert_eq!(client_data[1]["output_format"], "jpeg"); + let serialized = serde_json::to_string(&product.client_body_json).expect("json"); + assert!(!serialized.contains("javascript:")); + assert!(!serialized.contains("text/html")); + assert!(!serialized.contains(&oversized_prompt)); + assert_eq!(product.provider_body_json, provider_body); + } + + #[test] + fn sync_finalize_accepts_maximum_safe_base64_image_payload() { + let payload = "A".repeat(crate::formats::openai::image::MAX_OPENAI_IMAGE_DATA_BYTES); + let provider_body = json!({ + "created": 1779273523, + "data": [{"b64_json": payload}] + }); + let report_context = json!({ + "client_api_format": "openai:image", + "provider_api_format": "openai:image", + "image_request": {"operation": "generate"} + }); + + let product = maybe_build_openai_image_sync_finalize_product( + "openai_image_sync_finalize", + 200, + Some(&report_context), + Some(&provider_body), + None, + ) + .expect("maximum safe image response should finalize") + .expect("maximum safe image should be retained"); + + let returned = product.client_body_json["data"][0]["b64_json"] + .as_str() + .expect("base64 payload"); + assert_eq!( + returned.len(), + crate::formats::openai::image::MAX_OPENAI_IMAGE_DATA_BYTES + ); + assert!(returned.as_bytes().iter().all(|byte| *byte == b'A')); + } } 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 ebac54bd6..b145d7c69 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,13 @@ fn finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_ finalization.provider_api_format, reasoning_replay_policy, ); + if finalization + .provider_api_format + .trim() + .eq_ignore_ascii_case("openai:responses") + { + super::responses::normalize_openai_responses_message_item_ids(body); + } crate::enforce_request_body_stream_field( body, finalization.provider_api_format, @@ -374,10 +381,21 @@ mod tests { #[test] fn finalization_strips_non_replayable_responses_reasoning_history() { + let gemini_carrier = + crate::formats::openai::responses::encode_gemini_tool_signature_carrier( + "opaque-gemini-thought-signature", + ) + .expect("Gemini signature carrier"); let mut body = json!({ "model": "gpt-5.4", "input": [ {"type": "reasoning", "id": "rs_provider_123", "summary": []}, + { + "type": "reasoning", + "id": "rs_aether_55070860f6d45c6b8f6fa11efd9dff8a", + "summary": [], + "encrypted_content": gemini_carrier + }, { "type": "reasoning", "id": "item_72d3bd8d367d01977ace23f1", @@ -408,6 +426,39 @@ mod tests { assert_eq!(input[1]["type"], "message"); } + #[test] + fn finalization_repairs_legacy_responses_message_ids_for_same_format_upstream() { + let mut body = json!({ + "model": "gpt-5.4", + "input": [{ + "type": "message", + "id": "1c938e58-32a8-4d28-9c34-538d78076895_msg", + "role": "assistant", + "content": [{"type": "input_text", "text": "previous answer"}] + }] + }); + + finalize_openai_provider_request( + &mut body, + OpenAiProviderRequestFinalization { + source_api_format: "openai:responses", + provider_api_format: "openai:responses", + provider_type: "codex", + provider_model: "gpt-5.4", + source_model: "gpt-5.4", + body_rules: None, + upstream_is_stream: false, + require_body_stream_field: false, + }, + ) + .expect("legacy message IDs should be repaired before provider validation"); + + let repaired_id = body["input"][0]["id"] + .as_str() + .expect("message ID should be a string"); + assert!(repaired_id.starts_with("msg_")); + } + #[test] fn non_responses_sources_receive_codex_responses_reasoning_defaults() { for source_api_format in ["openai:chat", "claude:messages", "gemini:generate_content"] { 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 31260643f..5aa99bd1a 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/codex.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/codex.rs @@ -36,8 +36,8 @@ const CODEX_OPENAI_RESPONSES_COMPACT_BODY_FIELDS: &[&str] = &[ "prompt_cache_key", "text", ]; -pub const CODEX_CLIENT_VERSION: &str = "0.144.1"; -pub const CODEX_CLIENT_USER_AGENT: &str = "codex_cli_rs/0.144.1"; +pub const CODEX_CLIENT_VERSION: &str = "0.153.3"; +pub const CODEX_CLIENT_USER_AGENT: &str = "codex_cli_rs/0.153.3"; 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"; 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 dd42a6484..36053f67e 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/mod.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/mod.rs @@ -10,6 +10,7 @@ pub mod stream; const TOOL_ERROR_PREFIX: &str = "[tool error]"; const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_"; +const AETHER_MESSAGE_ITEM_ID_PREFIX: &str = "msg_aether_"; const GEMINI_TOOL_SIGNATURE_CARRIER_PREFIX: &str = "cpa-gemini-responses-carrier-v1:"; const MAX_GEMINI_THOUGHT_SIGNATURE_LEN: usize = 32 * 1024 * 1024; const MAX_GEMINI_THOUGHT_SIGNATURE_ENCODED_LEN: usize = @@ -102,12 +103,74 @@ pub fn openai_responses_synthetic_reasoning_item_id( ) } +/// Builds a stable, wire-compatible ID for a message item synthesized by Aether. +/// +/// Responses clients replay assistant message items verbatim on the next turn and +/// OpenAI requires those IDs to begin with `msg`. Upstream response IDs are not +/// guaranteed to have that prefix (Chat completion IDs and UUIDs are common), so +/// appending a suffix to the response ID is not sufficient. A deterministic UUID +/// keeps the ID stable across sync/stream projections while avoiding assumptions +/// about the upstream ID's shape or length. +pub fn openai_responses_message_item_id(response_id: &str, output_index: usize) -> String { + let seed = format!("{response_id}:{output_index}"); + format!( + "{AETHER_MESSAGE_ITEM_ID_PREFIX}{}", + uuid::Uuid::new_v5(&uuid::Uuid::NAMESPACE_OID, seed.as_bytes()).simple() + ) +} + +/// Repairs legacy/non-OpenAI message IDs in a Responses request in place. +/// +/// Aether versions before the `msg_` contract emitted IDs such as +/// `_msg`. Clients legitimately replay those assistant items on +/// the next turn, so merely fixing newly generated responses leaves existing +/// conversations broken. Preserve already-valid provider IDs and deterministically +/// remap only message items that do not begin with `msg`. +pub fn normalize_openai_responses_message_item_ids(body: &mut Value) -> usize { + let Some(items) = body.get_mut("input").and_then(Value::as_array_mut) else { + return 0; + }; + let mut repaired = 0usize; + for (index, item) in items.iter_mut().enumerate() { + let Some(object) = item.as_object_mut() else { + continue; + }; + if object.get("type").and_then(Value::as_str) != Some("message") { + continue; + } + let Some(raw_id) = object.get("id") else { + // IDs are optional for newly-authored input messages. Only repair + // an ID that a previous response actually supplied. + continue; + }; + let valid = object + .get("id") + .and_then(Value::as_str) + .is_some_and(|id| id.starts_with("msg")); + if valid { + continue; + } + let source_id = raw_id + .as_str() + .filter(|id| !id.trim().is_empty()) + .unwrap_or("missing") + .to_string(); + object.insert( + "id".to_string(), + Value::String(openai_responses_message_item_id(source_id.as_str(), index)), + ); + repaired += 1; + } + repaired +} + /// 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 -/// prefix. Foreign IDs (for example `item_...`) are therefore removed. Aether-synthesized -/// reasoning summaries are also removed unless they carry encrypted reasoning state that can be -/// replayed statelessly. +/// prefix. Foreign IDs (for example `item_...`) are therefore removed. Aether's Gemini signature +/// carriers are also removed: they are intentionally transported through the Responses +/// `encrypted_content` field so they can be restored on a later Gemini tool turn, but they are not +/// OpenAI ciphertext and must never be replayed to an OpenAI/Codex backend. pub fn strip_incompatible_openai_responses_reasoning_items( body: &mut Value, provider_api_format: &str, @@ -159,6 +222,13 @@ fn openai_responses_reasoning_item_is_replayable( if object.get("type").and_then(Value::as_str) != Some("reasoning") { return true; } + if object + .get("encrypted_content") + .and_then(Value::as_str) + .is_some_and(|value| value.starts_with(GEMINI_TOOL_SIGNATURE_CARRIER_PREFIX)) + { + return false; + } if policy == OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque && deepseek_opaque_reasoning_item_is_replayable(object) { @@ -255,6 +325,7 @@ 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, strip_incompatible_openai_responses_reasoning_items, strip_incompatible_openai_responses_reasoning_items_with_policy, @@ -348,6 +419,38 @@ mod tests { assert_ne!(first, other); } + #[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); + let second = openai_responses_message_item_id("1c938e58-32a8-4d28-9c34-538d78076895", 0); + let other = openai_responses_message_item_id("chatcmpl-123", 1); + + assert!(first.starts_with("msg_")); + assert_eq!(first, second); + assert_ne!(first, other); + } + + #[test] + fn normalizes_legacy_message_ids_but_preserves_valid_ids() { + let mut body = json!({ + "input": [ + {"type": "message", "id": "1c938e58-32a8-4d28-9c34-538d78076895_msg", "role": "assistant"}, + {"type": "message", "id": "msg_provider_123", "role": "assistant"}, + {"type": "function_call", "id": "legacy_call"}, + {"type": "message", "role": "user"} + ] + }); + + assert_eq!(normalize_openai_responses_message_item_ids(&mut body), 1); + let input = body["input"].as_array().expect("input array"); + assert!(input[0]["id"] + .as_str() + .is_some_and(|id| id.starts_with("msg_"))); + assert_eq!(input[1]["id"], "msg_provider_123"); + assert_eq!(input[2].get("id"), Some(&json!("legacy_call"))); + assert!(input[3].get("id").is_none()); + } + #[test] fn strips_foreign_and_non_replayable_synthetic_reasoning_items() { let portable_synthetic = openai_responses_synthetic_reasoning_item_id("resp_123", 1); @@ -380,6 +483,43 @@ mod tests { assert_eq!(input[2]["id"], "item_message_123"); } + #[test] + fn strips_gemini_signature_carriers_before_openai_replay() { + let gemini_item_id = openai_responses_synthetic_reasoning_item_id("resp_gemini", 0); + let openai_item_id = openai_responses_synthetic_reasoning_item_id("resp_openai", 0); + let carrier = encode_gemini_tool_signature_carrier_with_direction( + "opaque-gemini-thought-signature", + GeminiToolSignatureCarrierDirection::Next, + ) + .expect("Gemini signature carrier"); + let mut body = json!({ + "input": [ + { + "type": "reasoning", + "id": gemini_item_id, + "summary": [], + "encrypted_content": carrier + }, + { + "type": "reasoning", + "id": openai_item_id, + "summary": [], + "encrypted_content": "provider-encrypted-state" + }, + {"type": "reasoning", "id": "rs_provider_123", "summary": []} + ] + }); + + assert_eq!( + strip_incompatible_openai_responses_reasoning_items(&mut body, "openai:responses"), + 1 + ); + let input = body["input"].as_array().expect("input array"); + assert_eq!(input.len(), 2); + assert_eq!(input[0]["encrypted_content"], "provider-encrypted-state"); + assert_eq!(input[1]["id"], "rs_provider_123"); + } + #[test] fn reasoning_item_sanitizer_is_scoped_to_responses_targets() { let mut body = json!({ diff --git a/crates/aether-ai/formats/src/formats/shared/image_bridge.rs b/crates/aether-ai/formats/src/formats/shared/image_bridge.rs index 89e4b11e1..2735ecb73 100644 --- a/crates/aether-ai/formats/src/formats/shared/image_bridge.rs +++ b/crates/aether-ai/formats/src/formats/shared/image_bridge.rs @@ -1,7 +1,21 @@ use serde_json::{json, Map, Number, Value}; +use crate::formats::openai::image::{ + bounded_openai_image_revised_prompt, is_safe_openai_image_base64_payload, + normalize_openai_image_output_format, parse_safe_openai_image_data_url, + safe_openai_image_mime_type, sanitize_openai_image_source_url, +}; use crate::formats::shared::model_directives::extract_gemini_model_from_path; +const MAX_IMAGE_BRIDGE_OUTPUTS: usize = 64; +// Responses output items are provider-controlled and may contain deeply +// nested message/content arrays. Keep the projection bounded independently +// of the image count so a pathological text envelope cannot exhaust stack or +// heap while an otherwise valid image is being bridged. +const MAX_IMAGE_BRIDGE_PARTS: usize = 512; +const MAX_IMAGE_BRIDGE_TEXT_BYTES: usize = 256 * 1024; +const MAX_IMAGE_BRIDGE_RECURSION_DEPTH: usize = 32; + #[derive(Clone, Debug, PartialEq)] pub struct OpenAiImageRequestForGemini { pub requested_model: String, @@ -208,7 +222,7 @@ pub fn build_openai_image_response_from_gemini_response( .get("text") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(bounded_openai_image_revised_prompt) { revised_prompt = Some(Value::String(text.to_string())); } @@ -220,6 +234,12 @@ pub fn build_openai_image_response_from_gemini_response( "output_format": output_format_from_mime_type(&mime_type), "revised_prompt": revised_prompt.clone().unwrap_or(Value::Null), })); + if images.len() >= MAX_IMAGE_BRIDGE_OUTPUTS { + break; + } + } + if images.len() >= MAX_IMAGE_BRIDGE_OUTPUTS { + break; } } if images.is_empty() { @@ -249,17 +269,64 @@ pub fn build_openai_image_response_from_gemini_response( Some(Value::Object(response)) } +/// Projects a native OpenAI Images response before it is returned to a client. +/// +/// Native image responses normally do not need a format conversion, but they +/// still cross the provider trust boundary. Keep this projection separate +/// from the raw provider body retained for the conversion/audit report so +/// untrusted URLs, payloads, and metadata cannot be passed through unchanged. +pub(crate) fn build_openai_image_response_from_standard_image_response( + provider_body_json: &Value, + report_context: Option<&Value>, +) -> Option { + let data = provider_body_json.get("data")?.as_array()?; + let images = data + .iter() + .take(MAX_IMAGE_BRIDGE_OUTPUTS) + .filter_map(standard_openai_image_item_to_image_data) + .collect::>(); + if images.is_empty() { + return None; + } + + let created = provider_body_json + .get("created") + .and_then(Value::as_i64) + .unwrap_or_default(); + let mut response = Map::new(); + response.insert("created".to_string(), Value::Number(Number::from(created))); + response.insert("data".to_string(), Value::Array(images)); + if let Some(model) = provider_body_json + .get("model") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .or_else(|| report_context.and_then(context_model)) + { + response.insert("model".to_string(), Value::String(model.to_string())); + } + if let Some(usage) = provider_body_json.get("usage") { + response.insert("usage".to_string(), usage.clone()); + } + Some(Value::Object(response)) +} + pub fn build_gemini_image_response_from_openai_image_response( provider_body_json: &Value, report_context: Option<&Value>, ) -> Option { let mut parts = Vec::new(); - for item in provider_body_json.get("data")?.as_array()? { + for item in provider_body_json + .get("data")? + .as_array()? + .iter() + .take(MAX_IMAGE_BRIDGE_OUTPUTS) + { if let Some(prompt) = item .get("revised_prompt") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(bounded_openai_image_revised_prompt) { parts.push(json!({ "text": prompt })); } @@ -311,22 +378,23 @@ pub fn build_gemini_image_response_from_openai_responses_image_response( ) -> Option { let output = provider_body_json.get("output").and_then(Value::as_array)?; let mut parts = Vec::new(); - for item in output { + let mut budget = GeminiImagePartBudget::default(); + for item in output.iter().take(MAX_IMAGE_BRIDGE_OUTPUTS) { let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); if item_type == "image_generation_call" { if let Some(prompt) = item .get("revised_prompt") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(bounded_openai_image_revised_prompt) { - parts.push(json!({ "text": prompt })); + budget.push_text(&mut parts, prompt); } let Some(b64_json) = item .get("result") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) else { continue; }; @@ -335,19 +403,22 @@ pub fn build_gemini_image_response_from_openai_responses_image_response( .and_then(Value::as_str) .map(mime_type_from_output_format) .unwrap_or_else(|| "image/png".to_string()); - parts.push(json!({ - "inlineData": { - "mimeType": mime_type, - "data": b64_json, - } - })); + budget.push_part( + &mut parts, + json!({ + "inlineData": { + "mimeType": mime_type, + "data": b64_json, + } + }), + ); continue; } if matches!( item_type, "message" | "output_text" | "text" | "output_image" | "image_url" ) { - collect_openai_response_output_item_for_gemini(item, &mut parts); + collect_openai_response_output_item_for_gemini(item, &mut parts, &mut budget, 0); } } if !parts.iter().any(is_gemini_inline_image_part) { @@ -389,6 +460,7 @@ pub fn build_openai_image_response_from_response_stream_sync_body( let output = provider_body_json.get("output").and_then(Value::as_array)?; let images = output .iter() + .take(MAX_IMAGE_BRIDGE_OUTPUTS) .filter_map(openai_response_image_generation_item_to_image_data) .collect::>(); if images.is_empty() { @@ -438,28 +510,90 @@ fn openai_response_image_generation_item_to_image_data(item: &Value) -> Option { + Some(value) if value.trim_start().starts_with("data:") => { let (_, b64_json) = parse_data_url(value)?; image.insert("b64_json".to_string(), Value::String(b64_json)); } - Some(value) if value.starts_with("http://") || value.starts_with("https://") => { - image.insert("url".to_string(), Value::String(value.to_string())); - } Some(value) => { - image.insert("b64_json".to_string(), Value::String(value.to_string())); + if let Some(url) = sanitize_openai_image_source_url(value) { + if url.starts_with("data:") { + let (_, b64_json) = parse_data_url(&url)?; + image.insert("b64_json".to_string(), Value::String(b64_json)); + } else { + image.insert("url".to_string(), Value::String(url)); + } + } else if is_safe_openai_image_base64_payload(value) { + image.insert("b64_json".to_string(), Value::String(value.to_string())); + } else { + return None; + } } None => { let url = url?; - if let Some((_, b64_json)) = parse_data_url(url) { + let url = sanitize_openai_image_source_url(url)?; + if let Some((_, b64_json)) = parse_data_url(&url) { image.insert("b64_json".to_string(), Value::String(b64_json)); } else { - image.insert("url".to_string(), Value::String(url.to_string())); + image.insert("url".to_string(), Value::String(url)); } } } image.insert( "revised_prompt".to_string(), - item.get("revised_prompt").cloned().unwrap_or(Value::Null), + item.get("revised_prompt") + .and_then(Value::as_str) + .and_then(bounded_openai_image_revised_prompt) + .map(|value| Value::String(value.to_string())) + .unwrap_or(Value::Null), + ); + Some(Value::Object(image)) +} + +fn standard_openai_image_item_to_image_data(item: &Value) -> Option { + let object = item.as_object()?; + let b64_json = object + .get("b64_json") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| is_safe_openai_image_base64_payload(value)); + let url = object + .get("url") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(sanitize_openai_image_source_url); + + let mut image = Map::new(); + if let Some(b64_json) = b64_json { + image.insert("b64_json".to_string(), Value::String(b64_json.to_string())); + } else if let Some(url) = url { + if let Some((_, b64_json)) = parse_data_url(&url) { + image.insert("b64_json".to_string(), Value::String(b64_json)); + } else { + image.insert("url".to_string(), Value::String(url)); + } + } else { + return None; + } + + if let Some(output_format) = object + .get("output_format") + .and_then(Value::as_str) + .and_then(normalize_openai_image_output_format) + { + image.insert( + "output_format".to_string(), + Value::String(output_format.to_string()), + ); + } + image.insert( + "revised_prompt".to_string(), + object + .get("revised_prompt") + .and_then(Value::as_str) + .and_then(bounded_openai_image_revised_prompt) + .map(|value| Value::String(value.to_string())) + .unwrap_or(Value::Null), ); Some(Value::Object(image)) } @@ -474,12 +608,17 @@ pub fn build_openai_image_provider_body_from_response_stream_sync_body( } let output = data .iter() + .take(MAX_IMAGE_BRIDGE_OUTPUTS) .filter_map(|item| { extract_openai_image_response_item(item).map(|(mime_type, _)| { json!({ "type": "image_generation_call", "output_format": output_format_from_mime_type(&mime_type), - "revised_prompt": item.get("revised_prompt").cloned().unwrap_or(Value::Null), + "revised_prompt": item.get("revised_prompt") + .and_then(Value::as_str) + .and_then(bounded_openai_image_revised_prompt) + .map(|value| Value::String(value.to_string())) + .unwrap_or(Value::Null), }) }) }) @@ -578,7 +717,8 @@ fn openai_input_image_to_gemini_part(image: Value) -> Option { .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty())?; - if let Some((mime_type, data)) = parse_data_url(image_url) { + let image_url = sanitize_openai_image_source_url(image_url)?; + if let Some((mime_type, data)) = parse_data_url(&image_url) { return Some(json!({ "inlineData": { "mimeType": mime_type, @@ -588,7 +728,7 @@ fn openai_input_image_to_gemini_part(image: Value) -> Option { } Some(json!({ "fileData": { - "mimeType": mime_type_from_url(image_url), + "mimeType": mime_type_from_url(&image_url), "fileUri": image_url, } })) @@ -691,13 +831,53 @@ fn collect_gemini_part(part: &Value, text: &mut Vec, content: &mut Vec) { +#[derive(Default)] +struct GeminiImagePartBudget { + text_bytes: usize, +} + +impl GeminiImagePartBudget { + fn push_part(&mut self, parts: &mut Vec, part: Value) { + if parts.len() < MAX_IMAGE_BRIDGE_PARTS { + parts.push(part); + } + } + + fn push_text(&mut self, parts: &mut Vec, text: &str) { + let text_bytes = text.len(); + let Some(next_text_bytes) = self.text_bytes.checked_add(text_bytes) else { + return; + }; + if next_text_bytes > MAX_IMAGE_BRIDGE_TEXT_BYTES || parts.len() >= MAX_IMAGE_BRIDGE_PARTS { + return; + } + parts.push(json!({ "text": text })); + self.text_bytes = next_text_bytes; + } +} + +fn collect_openai_response_output_item_for_gemini( + item: &Value, + parts: &mut Vec, + budget: &mut GeminiImagePartBudget, + depth: usize, +) { + if depth >= MAX_IMAGE_BRIDGE_RECURSION_DEPTH { + return; + } + if let Value::Array(items) = item { + for child in items { + collect_openai_response_output_item_for_gemini(child, parts, budget, depth + 1); + if parts.len() >= MAX_IMAGE_BRIDGE_PARTS { + break; + } + } + return; + } let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); if item_type == "message" { - if let Some(content) = item.get("content").and_then(Value::as_array) { - for part in content { - collect_openai_response_output_item_for_gemini(part, parts); - } + if let Some(content) = item.get("content") { + collect_openai_response_output_item_for_gemini(content, parts, budget, depth + 1); } return; } @@ -708,7 +888,7 @@ fn collect_openai_response_output_item_for_gemini(item: &Value, parts: &mut Vec< .map(str::trim) .filter(|value| !value.is_empty()) { - parts.push(json!({ "text": text })); + budget.push_text(parts, text); } return; } @@ -727,12 +907,15 @@ fn collect_openai_response_output_item_for_gemini(item: &Value, parts: &mut Vec< .filter(|value| !value.is_empty()); if let Some(image_url) = image_url { if let Some((mime_type, data)) = parse_data_url(image_url) { - parts.push(json!({ - "inlineData": { - "mimeType": mime_type, - "data": data, - } - })); + budget.push_part( + parts, + json!({ + "inlineData": { + "mimeType": mime_type, + "data": data, + } + }), + ); } } } @@ -745,15 +928,14 @@ fn extract_gemini_inline_image(part: &Value) -> Option<(String, String)> { .get("mimeType") .or_else(|| object.get("mime_type")) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| value.starts_with("image/")) + .and_then(safe_openai_image_mime_type) .unwrap_or("image/png") .to_string(); let data = object .get("data") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty())? + .filter(|value| is_safe_openai_image_base64_payload(value))? .to_string(); Some((mime_type, data)) } @@ -764,13 +946,12 @@ fn extract_openai_image_response_item(item: &Value) -> Option<(String, String)> .get("b64_json") .and_then(Value::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) { let output_format = object .get("output_format") .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(normalize_openai_image_output_format) .unwrap_or("png"); return Some(( mime_type_from_output_format(output_format), @@ -786,13 +967,7 @@ fn extract_openai_image_response_item(item: &Value) -> Option<(String, String)> } fn parse_data_url(value: &str) -> Option<(String, String)> { - let (metadata, payload) = value.trim().split_once(',')?; - let metadata = metadata.strip_prefix("data:")?; - let mime_type = metadata.strip_suffix(";base64")?; - let payload = payload.trim(); - if payload.is_empty() { - return None; - } + let (mime_type, payload) = parse_safe_openai_image_data_url(value)?; Some((mime_type.to_string(), payload.to_string())) } @@ -822,12 +997,11 @@ fn output_format_from_mime_type(mime_type: &str) -> &'static str { } fn mime_type_from_output_format(output_format: &str) -> String { - match output_format.trim().to_ascii_lowercase().as_str() { - "jpeg" | "jpg" => "image/jpeg".to_string(), - "webp" => "image/webp".to_string(), - "png" => "image/png".to_string(), - other if other.starts_with("image/") => other.to_string(), - _ => "image/png".to_string(), + match normalize_openai_image_output_format(output_format) { + Some("jpeg") => "image/jpeg".to_string(), + Some("webp") => "image/webp".to_string(), + Some("png") | None => "image/png".to_string(), + Some(_) => "image/png".to_string(), } } @@ -901,7 +1075,7 @@ fn context_model(context: &Value) -> Option<&str> { #[cfg(test)] mod tests { use http::{Method, Request}; - use serde_json::json; + use serde_json::{json, Value}; use super::{ build_gemini_image_request_body_from_openai_image_request, @@ -909,6 +1083,7 @@ mod tests { build_openai_image_request_body_from_gemini_image_request, build_openai_image_response_from_gemini_response, build_openai_image_response_from_response_stream_sync_body, + build_openai_image_response_from_standard_image_response, gemini_request_is_image_generation, }; use crate::formats::openai::image::request::normalize_openai_image_request; @@ -1129,4 +1304,88 @@ mod tests { ); assert_eq!(converted["usageMetadata"]["totalTokenCount"], 3); } + + #[test] + fn standard_openai_image_bridge_filters_fields_and_bounds_outputs() { + let provider_body = json!({ + "created": 1779273523, + "model": "gpt-image-2", + "data": [ + {"url": "javascript:alert(1)"}, + {"url": "data:text/html;base64,PGh0bWw+"}, + { + "b64_json": "aGVsbG8=", + "output_format": "text/html", + "revised_prompt": "p".repeat(256 * 1024 + 1) + } + ] + }); + let converted = + build_openai_image_response_from_standard_image_response(&provider_body, None) + .expect("valid standard image should remain"); + assert_eq!(converted["data"].as_array().map(Vec::len), Some(1)); + assert_eq!(converted["data"][0]["b64_json"], "aGVsbG8="); + assert_eq!(converted["data"][0]["revised_prompt"], Value::Null); + assert!(converted["data"][0].get("output_format").is_none()); + let serialized = serde_json::to_string(&converted).expect("json"); + assert!(!serialized.contains("javascript:")); + assert!(!serialized.contains("text/html")); + + let outputs = (0..80) + .map(|index| json!({"b64_json": format!("image{index:02}=")})) + .collect::>(); + let bounded = build_openai_image_response_from_standard_image_response( + &json!({"data": outputs}), + None, + ) + .expect("bounded standard image response should convert"); + assert_eq!(bounded["data"].as_array().map(Vec::len), Some(64)); + } + + #[test] + fn responses_image_bridge_bounds_nested_text_without_losing_image() { + let mut nested = json!({"type": "output_text", "text": "nested"}); + for _ in 0..128 { + nested = json!({"type": "message", "content": [nested]}); + } + let converted = super::build_gemini_image_response_from_openai_responses_image_response( + &json!({ + "output": [ + nested, + {"type": "image_generation_call", "result": "aGVsbG8="} + ] + }), + None, + ) + .expect("nested response should still retain the image"); + let parts = converted["candidates"][0]["content"]["parts"] + .as_array() + .expect("gemini parts"); + assert!(parts.iter().any(super::is_gemini_inline_image_part)); + assert!(parts.len() <= super::MAX_IMAGE_BRIDGE_PARTS); + + let large_text = "t".repeat(super::MAX_IMAGE_BRIDGE_TEXT_BYTES / 2); + let output = (0..4) + .map(|_| json!({"type": "output_text", "text": large_text.clone()})) + .chain(std::iter::once(json!({ + "type": "image_generation_call", + "result": "aGVsbG8=" + }))) + .collect::>(); + let bounded = super::build_gemini_image_response_from_openai_responses_image_response( + &json!({"output": output}), + None, + ) + .expect("text budget should not suppress a valid image"); + let parts = bounded["candidates"][0]["content"]["parts"] + .as_array() + .expect("bounded parts"); + assert!(parts.iter().any(super::is_gemini_inline_image_part)); + let text_bytes = parts + .iter() + .filter_map(|part| part.get("text").and_then(Value::as_str)) + .map(str::len) + .sum::(); + assert!(text_bytes <= super::MAX_IMAGE_BRIDGE_TEXT_BYTES); + } } diff --git a/crates/aether-ai/formats/src/formats/shared/mod.rs b/crates/aether-ai/formats/src/formats/shared/mod.rs index ff4a69ac7..00fc97d12 100644 --- a/crates/aether-ai/formats/src/formats/shared/mod.rs +++ b/crates/aether-ai/formats/src/formats/shared/mod.rs @@ -1,5 +1,11 @@ +use base64::Engine as _; use std::fmt; +/// Report bodies are produced from execution responses, whose normal decoded +/// transport limit is 64 MiB. Keep format finalizers on the same boundary so +/// a base64 field cannot trigger an unchecked allocation before parsing. +pub(crate) const MAX_SYNC_REPORT_BODY_BYTES: usize = 64 * 1024 * 1024; + pub mod error_body; pub mod family; pub mod image_bridge; @@ -18,6 +24,35 @@ pub mod sync_products; pub mod sync_to_stream; pub mod video; +pub(crate) fn decode_sync_report_body_base64( + body_base64: &str, +) -> Result, AiSurfaceFinalizeError> { + if body_base64.is_empty() { + return Ok(Vec::new()); + } + + let max_encoded_len = MAX_SYNC_REPORT_BODY_BYTES + .checked_add(2) + .and_then(|value| value.checked_div(3)) + .and_then(|value| value.checked_mul(4)) + .unwrap_or(usize::MAX); + if body_base64.len() > max_encoded_len { + return Err(AiSurfaceFinalizeError::new(format!( + "sync report body exceeds {} decoded bytes", + MAX_SYNC_REPORT_BODY_BYTES + ))); + } + + let bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + if bytes.len() > MAX_SYNC_REPORT_BODY_BYTES { + return Err(AiSurfaceFinalizeError::new(format!( + "sync report body exceeds {} decoded bytes", + MAX_SYNC_REPORT_BODY_BYTES + ))); + } + Ok(bytes) +} + pub use self::sse::{encode_done_sse, encode_json_sse, map_claude_stop_reason}; pub use self::stream_core::{CanonicalStreamEvent, CanonicalStreamFrame}; pub use self::stream_rewrite::{ diff --git a/crates/aether-ai/formats/src/formats/shared/routing.rs b/crates/aether-ai/formats/src/formats/shared/routing.rs index cdb3047f1..3dec3d5b7 100644 --- a/crates/aether-ai/formats/src/formats/shared/routing.rs +++ b/crates/aether-ai/formats/src/formats/shared/routing.rs @@ -441,7 +441,19 @@ pub fn sanitize_request_path(path: &str) -> Option { // including for malformed routes that will later be rejected. return Some("/v1/live/{call_id}".to_string()); } - Some(path.to_string()) + Some(sanitize_sensitive_request_path(path)) +} + +fn sanitize_sensitive_request_path(path: &str) -> String { + for prefix in ["/install-tunnel/", "/install/", "/i/"] { + if path + .strip_prefix(prefix) + .is_some_and(|secret| !secret.is_empty()) + { + return format!("{prefix}[redacted]"); + } + } + path.to_string() } pub fn sanitize_request_query_string(query: &str) -> Option { @@ -1022,6 +1034,29 @@ mod tests { ); } + #[test] + fn request_path_metadata_sanitizer_redacts_install_session_codes() { + for (raw, expected) in [ + ("/install/secret-code", "/install/[redacted]"), + ("/install/secret-code.ps1", "/install/[redacted]"), + ("/i/secret-code", "/i/[redacted]"), + ( + "/install-tunnel/secret-code.ps1?token=also-secret", + "/install-tunnel/[redacted]", + ), + ] { + assert_eq!(sanitize_request_path(raw).as_deref(), Some(expected)); + assert_eq!( + sanitize_request_path_and_query(raw, None).as_deref(), + Some(expected) + ); + } + assert_eq!( + sanitize_request_path("/install/").as_deref(), + Some("/install/") + ); + } + #[test] fn stream_matching_requires_openai_stream_flag() { assert!(!is_matching_stream_request( 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 2f009554c..6983f7e94 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 @@ -17,8 +17,9 @@ use crate::formats::shared::error_body::{ }; use crate::formats::shared::sse::encode_json_sse; use crate::formats::shared::stream_core::common::{ - decode_json_data_line, openai_stream_terminal_error_body, openai_stream_terminal_error_message, - unsupported_stream_event_message, CanonicalStreamEvent, CanonicalStreamFrame, CanonicalUsage, + canonical_usage_from_openai_usage, decode_json_data_line, openai_stream_terminal_error_body, + openai_stream_terminal_error_message, unsupported_stream_event_message, CanonicalStreamEvent, + CanonicalStreamFrame, CanonicalUsage, }; use crate::formats::shared::AiSurfaceFinalizeError; @@ -134,7 +135,11 @@ impl StreamingStandardFormatMatrix { } if let CanonicalStreamEvent::UnknownEvent(payload) = &frame.event { self.terminated = true; - out.extend(client.emit_unknown_event(payload)?); + if openai_stream_terminal_error_body(payload).is_some() { + out.extend(client.emit_terminal_error_frame(frame)?); + } else { + out.extend(client.emit_unknown_event(payload)?); + } break; } if let CanonicalStreamEvent::OpenAiResponsesOutputItem { raw_event, .. } = &frame.event @@ -368,6 +373,10 @@ impl StreamingStandardTerminalObserver { summary.observed_finish = true; summary.finish_reason = Some("error".to_string()); summary.parser_error = openai_stream_terminal_error_message(&payload); + summary.standardized_usage = payload + .pointer("/response/usage") + .and_then(|usage| canonical_usage_from_openai_usage(Some(usage))) + .map(standardized_usage_from_canonical); } CanonicalStreamEvent::UnknownEvent(_) => { summary.unknown_event_count = summary.unknown_event_count.saturating_add(1); @@ -376,6 +385,13 @@ impl StreamingStandardTerminalObserver { finish_reason, usage, } => { + if let Some(parser_error) = finish_reason + .as_deref() + .filter(|reason| !canonical_stream_finish_reason_is_supported(reason)) + .map(|reason| format!("unsupported provider stream finish reason: {reason}")) + { + summary.parser_error.get_or_insert(parser_error); + } summary.finish_reason = finish_reason; summary.standardized_usage = usage.map(standardized_usage_from_canonical); summary.observed_finish = true; @@ -606,6 +622,44 @@ impl ClientStreamEmitter { self.emit_error(error_body) } + fn emit_terminal_error_frame( + &mut self, + frame: CanonicalStreamFrame, + ) -> Result, AiSurfaceFinalizeError> { + if matches!( + self, + ClientStreamEmitter::OpenAIChat(_) | ClientStreamEmitter::OpenAIResponses(_) + ) { + return self.emit(frame); + } + let CanonicalStreamEvent::UnknownEvent(payload) = frame.event else { + return self.emit(frame); + }; + let Some(source_error_body) = openai_stream_terminal_error_body(&payload) else { + return self.emit_unknown_event(&payload); + }; + let Some(error) = source_error_body.get("error") else { + return self.emit_unknown_event(&payload); + }; + let message = error + .get("message") + .and_then(Value::as_str) + .unwrap_or("Upstream stream ended with an error"); + let code = error.get("code").and_then(|value| match value { + Value::String(value) => Some(value.as_str()), + _ => None, + }); + let Some(error_body) = build_core_error_body_for_client_format( + self.api_format(), + message, + code, + LocalCoreSyncErrorKind::ServerError, + ) else { + return Ok(Vec::new()); + }; + self.emit_error(error_body) + } + fn emit_unsupported_finish_reason( &mut self, finish_reason: &str, @@ -785,6 +839,145 @@ mod tests { format!("event: {event}\n").into_bytes() } + #[test] + fn terminal_observer_marks_malformed_gemini_function_call_as_failure() { + let context = report_context("gemini:generate_content", "openai:responses"); + let mut observer = StreamingStandardTerminalObserver::default(); + observer + .push_line( + &context, + data_line(json!({ + "response": { + "responseId": "resp_malformed_tool_call", + "modelVersion": "gemini-3.7-flash-tiered", + "candidates": [{ + "index": 0, + "content": { + "role": "model", + "parts": [{"thoughtSignature": "signature", "text": ""}] + }, + "finishReason": "MALFORMED_FUNCTION_CALL", + "finishMessage": "Malformed function call: Function call is empty - no input to parse." + }], + "usageMetadata": { + "promptTokenCount": 206744, + "cachedContentTokenCount": 203947, + "thoughtsTokenCount": 1130, + "totalTokenCount": 207874 + } + }, + "responseId": "resp_malformed_tool_call" + })), + ) + .expect("Gemini terminal frame should parse"); + + let summary = observer + .finish(&context) + .expect("terminal observation should finish") + .expect("Gemini terminal frame should produce a summary"); + + assert!(summary.observed_finish); + assert_eq!(summary.finish_reason.as_deref(), Some("error")); + assert_eq!( + summary.parser_error.as_deref(), + Some("Malformed function call: Function call is empty - no input to parse.") + ); + let usage = summary + .standardized_usage + .expect("failed Gemini terminal should preserve usage"); + assert_eq!(usage.input_tokens, 206744); + assert_eq!(usage.output_tokens, 1130); + assert_eq!(usage.cache_read_tokens, 203947); + } + + #[test] + fn streams_gemini_thought_text_to_openai_responses_immediately() { + let context = report_context("gemini:generate_content", "openai:responses"); + let mut matrix = StreamingStandardFormatMatrix::default(); + let output = matrix + .transform_line( + &context, + data_line(json!({ + "response": { + "responseId": "resp_reasoning_123", + "modelVersion": "gemini-3.7-flash-tiered", + "candidates": [{ + "index": 0, + "content": { + "role": "model", + "parts": [{"thought": true, "text": "checking"}] + } + }] + } + })), + ) + .expect("first Gemini thought chunk should transform"); + let sse = String::from_utf8(output).expect("reasoning SSE should be utf8"); + + assert!( + sse.contains("event: response.reasoning_summary_text.delta\n"), + "{sse}" + ); + assert!(sse.contains("\"delta\":\"checking\""), "{sse}"); + } + + #[test] + fn transforms_malformed_gemini_function_call_to_responses_failed() { + let context = report_context("gemini:generate_content", "openai:responses"); + let mut matrix = StreamingStandardFormatMatrix::default(); + let output = matrix + .transform_line( + &context, + data_line(json!({ + "response": { + "responseId": "resp_malformed_tool_call", + "modelVersion": "gemini-3.7-flash-tiered", + "candidates": [{ + "index": 0, + "content": { + "role": "model", + "parts": [{ + "text": "", + "thoughtSignature": "opaque-thought-signature" + }] + }, + "finishReason": "MALFORMED_FUNCTION_CALL", + "finishMessage": "Malformed function call: Function call is empty - no input to parse." + }], + "usageMetadata": { + "promptTokenCount": 206744, + "cachedContentTokenCount": 203947, + "thoughtsTokenCount": 1130, + "totalTokenCount": 207874 + } + } + })), + ) + .expect("malformed Gemini terminal should transform to a stream error"); + let sse = String::from_utf8(output).expect("failed response SSE should be utf8"); + + assert!(sse.contains("event: response.failed\n"), "{sse}"); + assert!(sse.contains("\"type\":\"response.failed\""), "{sse}"); + assert!( + sse.contains("\"code\":\"MALFORMED_FUNCTION_CALL\""), + "{sse}" + ); + assert!( + sse.contains( + "\"message\":\"Malformed function call: Function call is empty - no input to parse.\"" + ), + "{sse}" + ); + assert!(sse.contains("\"input_tokens\":206744"), "{sse}"); + assert!(sse.contains("\"output_tokens\":1130"), "{sse}"); + assert!(sse.contains("\"cached_tokens\":203947"), "{sse}"); + assert!(!sse.contains("unsupported_finish_reason"), "{sse}"); + assert!(matrix + .finish(&context) + .expect("failed matrix should stay terminated") + .is_empty()); + } + #[test] fn event_only_stream_types_convert_across_standard_formats() { let responses_payload = json!({ @@ -1318,6 +1511,17 @@ mod tests { .expect("keepalive should be ignored"); assert!(keepalive.is_empty()); + let ping = matrix + .transform_line( + &report_context, + data_line(json!({ + "type": "ping", + "cost": "0", + })), + ) + .expect("provider ping should be ignored"); + assert!(ping.is_empty()); + for line in [ data_line(json!({ "type": "response.output_text.delta", 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 b9bd41f56..d766c5aa2 100644 --- a/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs +++ b/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs @@ -19,6 +19,12 @@ use crate::provider_compat::surfaces::{ provider_adaptation_should_unwrap_stream_envelope, KIRO_ENVELOPE_NAME, }; +// A stream record can legitimately contain a large tool payload, but a peer +// must not be able to keep the rewriter allocating forever by withholding the +// record separator. This is a parser carry-buffer bound, not a response-body +// or stream-concurrency limit. +const MAX_STREAM_REWRITE_BUFFER_BYTES: usize = 16 * 1024 * 1024; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum FinalizeStreamRewriteMode { EnvelopeUnwrap, @@ -294,7 +300,7 @@ impl AiSurfaceStreamRewriter<'_> { | AiSurfaceStreamRewriteState::ModelDirectiveDisplay | AiSurfaceStreamRewriteState::OpenAiResponsesCompat | AiSurfaceStreamRewriteState::Standard(_) => { - self.buffered.extend_from_slice(chunk); + append_bounded_stream_rewrite_chunk(&mut self.buffered, chunk)?; let mut output = Vec::new(); while let Some(line_end) = self.buffered.iter().position(|byte| *byte == b'\n') { let line = self.buffered.drain(..=line_end).collect::>(); @@ -399,7 +405,7 @@ impl ClaudeReadToolStreamSanitizer { report_context: &Value, chunk: &[u8], ) -> Result, AiSurfaceFinalizeError> { - self.buffered.extend_from_slice(chunk); + append_bounded_stream_rewrite_chunk(&mut self.buffered, chunk)?; let mut output = Vec::new(); while let Some(record) = drain_next_sse_record(&mut self.buffered) { output.extend(self.transform_record(report_context, record)?); @@ -523,6 +529,18 @@ impl ClaudeReadToolStreamSanitizer { return Ok(original_record); } if let Some(partial_json) = partial_json { + let next_len = state + .buffered_input_json + .len() + .checked_add(partial_json.len()) + .ok_or_else(|| { + AiSurfaceFinalizeError::new("stream rewrite tool input buffer length overflow") + })?; + if next_len > MAX_STREAM_REWRITE_BUFFER_BYTES { + return Err(AiSurfaceFinalizeError::new(format!( + "stream rewrite tool input buffer exceeds {MAX_STREAM_REWRITE_BUFFER_BYTES} bytes" + ))); + } state.buffered_input_json.push_str(partial_json); } Ok(Vec::new()) @@ -568,6 +586,23 @@ impl ClaudeReadToolStreamSanitizer { } } +fn append_bounded_stream_rewrite_chunk( + buffered: &mut Vec, + chunk: &[u8], +) -> Result<(), AiSurfaceFinalizeError> { + let next_len = buffered + .len() + .checked_add(chunk.len()) + .ok_or_else(|| AiSurfaceFinalizeError::new("stream rewrite buffer length overflow"))?; + if next_len > MAX_STREAM_REWRITE_BUFFER_BYTES { + return Err(AiSurfaceFinalizeError::new(format!( + "stream rewrite buffer exceeds {MAX_STREAM_REWRITE_BUFFER_BYTES} bytes" + ))); + } + buffered.extend_from_slice(chunk); + Ok(()) +} + fn sanitize_claude_tool_input_object(block: &mut Map, name: &str) -> bool { let Some(input) = block.get("input") else { return false; @@ -1344,6 +1379,48 @@ data: {\"type\":\"content_block_stop\",\"index\":0}\n\n", assert!(output.contains("event: content_block_stop")); } + #[test] + fn same_format_claude_stream_bounds_read_input_json_buffer() { + let report_context = json!({ + "provider_api_format": "claude:messages", + "client_api_format": "claude:messages", + "needs_conversion": false, + "anthropic_compatibility_profile": "claude_code_legacy", + }); + let mut rewriter = maybe_build_ai_surface_stream_rewriter(Some(&report_context)) + .expect("same-format claude sanitizer should exist"); + rewriter + .push_chunk( + b"event: content_block_start\n\ +data: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"tool_use\",\"id\":\"call_read_1\",\"name\":\"Read\",\"input\":{}}}\n\n", + ) + .expect("start should rewrite"); + + let build_delta = |partial_json: &str| { + let payload = json!({ + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "input_json_delta", + "partial_json": partial_json, + } + }); + let mut record = b"event: content_block_delta\ndata: ".to_vec(); + record.extend(serde_json::to_vec(&payload).expect("delta should serialize")); + record.extend_from_slice(b"\n\n"); + record + }; + let first = "x".repeat(super::MAX_STREAM_REWRITE_BUFFER_BYTES / 2); + let second = "x".repeat(super::MAX_STREAM_REWRITE_BUFFER_BYTES / 2 + 1); + rewriter + .push_chunk(&build_delta(&first)) + .expect("first partial JSON should fit"); + let error = rewriter + .push_chunk(&build_delta(&second)) + .expect_err("Read input JSON must be bounded across SSE records"); + assert!(error.0.contains("tool input buffer exceeds")); + } + #[test] fn same_format_claude_stream_preserves_other_tool_empty_pages() { 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 4f8a954f6..4a20b159f 100644 --- a/crates/aether-ai/formats/src/formats/shared/sync_products.rs +++ b/crates/aether-ai/formats/src/formats/shared/sync_products.rs @@ -7,8 +7,10 @@ use aether_ai_formats::formats::conversion::response::{ convert_openai_chat_response_to_openai_responses, convert_openai_responses_response_to_openai_chat, }; -use aether_ai_formats::formats::openai::responses::openai_responses_synthetic_reasoning_item_id; 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, +}; use aether_ai_formats::formats::registry::{convert_response, FormatContext, FormatError}; use aether_ai_formats::{ canonical_response_unknown_block_count, canonical_to_claude_response, @@ -21,7 +23,7 @@ use aether_ai_formats::{ }; use serde_json::{json, Map, Value}; -use super::AiSurfaceFinalizeError; +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::stream::{OpenAIChatProviderState, OpenAIResponsesProviderState}; @@ -77,7 +79,7 @@ pub fn maybe_build_standard_cross_format_sync_product_from_normalized_payload( let (aggregated_stream_body, aggregated_stream_api_format) = match body_base64 { Some(body_base64) => { - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; let provider_stream_event_api_format = provider_stream_event_api_format_for_report_context( report_context, @@ -469,6 +471,20 @@ pub fn maybe_build_standard_sync_finalize_product_from_normalized_payload( }; let body_base64 = body_base64.or(capture_stream_body_base64.as_deref()); + // Cross-format sync attempts can contain raw bytes because the plan requested a stream even + // though the provider returned one complete JSON response. Do not feed that response into an + // SSE aggregator. Capture envelopes and same-format responses retain their existing precedence. + let non_stream_capture_body_json = + if capture_envelope_used || !sync_finalize_needs_conversion(report_context) { + None + } else { + body_base64.and_then(decode_non_stream_sync_capture_body) + }; + let (body_json, body_base64) = match non_stream_capture_body_json.as_ref() { + Some(capture_body_json) => (body_json.or(Some(capture_body_json)), None), + None => (body_json, body_base64), + }; + if let Some(body_json) = maybe_build_standard_same_format_sync_body_from_normalized_payload( report_kind, status_code, @@ -596,7 +612,7 @@ pub fn maybe_build_embedding_cross_format_sync_product_from_normalized_payload( let provider_body_json = match body_base64 { Some(body_base64) => { - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; serde_json::from_slice::(&body_bytes).ok() } None => body_json.cloned(), @@ -741,7 +757,7 @@ fn maybe_build_standard_same_format_stream_sync_body( let Some(body_base64) = body_base64 else { return Ok(None); }; - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; let provider_stream_event_api_format = provider_stream_event_api_format_for_report_context(report_context, &provider_api_format); let Some(mut body) = try_aggregate_standard_chat_stream_sync_response( @@ -890,7 +906,7 @@ fn maybe_build_openai_responses_same_family_stream_sync_body( let Some(body_base64) = body_base64 else { return Ok(None); }; - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; // Same-family clients retain the authoritative terminal body verbatim, including future // output item fields, but unknown intermediate event types still fail closed. ensure_no_unknown_openai_responses_stream_events(&body_bytes, true)?; @@ -982,7 +998,7 @@ fn maybe_build_openai_cross_format_provider_body_from_normalized_payload( ) -> Result, AiSurfaceFinalizeError> { let aggregated_stream_body = match body_base64 { Some(body_base64) => { - let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?; + let body_bytes = decode_sync_report_body_base64(body_base64)?; let normalized_provider_api_format = normalize_openai_responses_family_api_format(provider_api_format); match normalized_provider_api_format.as_str() { @@ -1009,6 +1025,48 @@ fn maybe_build_openai_cross_format_provider_body_from_normalized_payload( })) } +fn sync_finalize_needs_conversion(report_context: Option<&Value>) -> bool { + report_context + .and_then(|report_context| report_context.get("needs_conversion")) + .and_then(Value::as_bool) + .unwrap_or(false) +} + +fn decode_non_stream_sync_capture_body(body_base64: &str) -> Option { + let body_bytes = base64::engine::general_purpose::STANDARD + .decode(body_base64) + .ok()?; + serde_json::from_slice::(&body_bytes) + .ok() + .filter(Value::is_object) + .filter(|body_json| !is_stream_event_object(body_json)) +} + +/// Unframed JSON events are accepted by the stream parsers and must not be mistaken for complete +/// provider response bodies merely because the entire capture parses as one JSON object. +fn is_stream_event_object(value: &Value) -> bool { + let Some(object) = value.as_object() else { + return false; + }; + if object + .get("object") + .and_then(Value::as_str) + .is_some_and(|object| object.ends_with(".chunk")) + { + return true; + } + + object + .get("type") + .and_then(Value::as_str) + .is_some_and(|event_type| { + event_type.contains('.') + || ["response", "message", "item", "delta", "content_block"] + .iter() + .any(|nested| object.contains_key(*nested)) + }) +} + fn is_error_like_sync_body(value: &Value) -> bool { let Some(object) = value.as_object() else { return false; @@ -2696,6 +2754,7 @@ fn aggregate_openai_responses_stream_sync_response_from_validated_terminal( if let Some(state) = message_states.remove(&output_index) { output.push(materialize_openai_responses_message_item( &response_id, + output_index, state, )); } @@ -3152,13 +3211,31 @@ fn resolve_openai_responses_tool_output_index( fn materialize_openai_responses_message_item( response_id: &str, + output_index: usize, state: OpenAIResponsesSyncMessageState, ) -> Value { let mut item = state.item; item.entry("type".to_string()) .or_insert_with(|| Value::String("message".to_string())); - item.entry("id".to_string()) - .or_insert_with(|| Value::String(format!("{response_id}_msg"))); + let message_id_is_valid = item + .get("id") + .and_then(Value::as_str) + .is_some_and(|id| id.starts_with("msg")); + if !message_id_is_valid { + let source_id = item + .get("id") + .and_then(Value::as_str) + .filter(|id| !id.trim().is_empty()) + .unwrap_or(response_id) + .to_string(); + item.insert( + "id".to_string(), + Value::String(openai_responses_message_item_id( + source_id.as_str(), + output_index, + )), + ); + } item.entry("role".to_string()) .or_insert_with(|| Value::String("assistant".to_string())); item.entry("status".to_string()) @@ -3958,7 +4035,8 @@ mod tests { aggregate_claude_stream_sync_response, aggregate_gemini_stream_sync_response, aggregate_openai_chat_stream_sync_response, aggregate_openai_responses_stream_sync_response, convert_standard_chat_response, - convert_standard_cli_response, materialize_openai_responses_reasoning_item, + convert_standard_cli_response, decode_non_stream_sync_capture_body, + materialize_openai_responses_reasoning_item, maybe_build_openai_chat_cross_format_sync_product_from_normalized_payload, maybe_build_openai_responses_cross_format_sync_product_from_normalized_payload, maybe_build_openai_responses_same_family_sync_body_from_normalized_payload, @@ -4236,6 +4314,57 @@ mod tests { assert_eq!(aggregated["usageMetadata"]["totalTokenCount"], 5); } + #[test] + fn aggregates_antigravity_signature_only_reasoning_exhaustion() { + let body = concat!( + "data: {\"response\":{\"responseId\":\"resp_signature_only_123\",\"modelVersion\":\"gemini-3.7-flash-tiered\",", + "\"candidates\":[{\"index\":0,\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"\",\"thoughtSignature\":\"opaque-thought-signature\"}]},\"finishReason\":\"MAX_TOKENS\"}],", + "\"usageMetadata\":{\"promptTokenCount\":22,\"thoughtsTokenCount\":29,\"totalTokenCount\":51}},", + "\"traceId\":\"trace-signature-only\"}\n\n", + ); + + let aggregated = aggregate_gemini_stream_sync_response(body.as_bytes()) + .expect("signature-only reasoning terminal should aggregate"); + + assert_eq!( + aggregated["candidates"][0]["content"]["parts"][0]["thought"], + true + ); + assert_eq!( + aggregated["candidates"][0]["content"]["parts"][0]["thoughtSignature"], + "opaque-thought-signature" + ); + assert_eq!(aggregated["candidates"][0]["finishReason"], "MAX_TOKENS"); + assert_eq!(aggregated["usageMetadata"]["thoughtsTokenCount"], 29); + assert!( + crate::formats::gemini::generate_content::response::from_raw(&aggregated).is_some() + ); + + let report_context = json!({ + "provider_api_format": "gemini:generate_content", + "client_api_format": "openai:chat", + "mapped_model": "gemini-3.7-flash-tiered", + }); + let product = maybe_build_standard_cross_format_sync_product_from_normalized_payload( + "openai_chat_sync_finalize", + 200, + Some(&report_context), + None, + Some(&base64::engine::general_purpose::STANDARD.encode(body)), + ) + .expect("signature-only reasoning terminal should convert") + .expect("cross-format product should exist"); + + assert_eq!( + product.client_body_json["choices"][0]["finish_reason"], + "length" + ); + assert_eq!( + product.client_body_json["usage"]["completion_tokens_details"]["reasoning_tokens"], + 29 + ); + } + #[test] fn gemini_stream_aggregation_rejects_unknown_parts() { let body = "data: {\"responseId\":\"resp_gem_unknown_123\",\"modelVersion\":\"gemini-2.5-pro\",\"candidates\":[{\"index\":0,\"content\":{\"role\":\"model\",\"parts\":[{\"futurePart\":{\"kept\":true}}]}}]}\n\n"; @@ -4455,6 +4584,132 @@ mod tests { ); } + #[test] + fn unframed_stream_events_are_not_mistaken_for_provider_bodies() { + for event in [ + json!({"type": "response.completed", "response": {"status": "completed"}}), + json!({"type": "response.output_text.delta", "delta": "hi"}), + json!({"type": "message_start", "message": {"id": "msg_1"}}), + json!({"type": "content_block_delta", "index": 0, "delta": {"text": "hi"}}), + json!({"object": "chat.completion.chunk", "choices": []}), + ] { + let body_base64 = base64::engine::general_purpose::STANDARD + .encode(serde_json::to_vec(&event).expect("serialize event")); + assert!( + decode_non_stream_sync_capture_body(&body_base64).is_none(), + "stream events belong to the aggregators: {event}" + ); + } + } + + #[test] + fn complete_provider_bodies_are_recovered_from_cross_format_captures() { + for body in [ + json!({"id": "resp_1", "object": "response", "status": "completed", "output": []}), + json!({"id": "chatcmpl_1", "object": "chat.completion", "choices": []}), + json!({"id": "msg_1", "type": "message", "role": "assistant", "content": []}), + json!({"candidates": [], "modelVersion": "probe-model"}), + ] { + let body_base64 = base64::engine::general_purpose::STANDARD + .encode(serde_json::to_vec(&body).expect("serialize provider body")); + assert_eq!( + decode_non_stream_sync_capture_body(&body_base64), + Some(body.clone()), + "a complete provider body is not a stream: {body}" + ); + } + } + + #[test] + fn recovers_cross_format_capture_that_is_a_complete_json_body() { + let report_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "claude:messages", + "needs_conversion": true, + "upstream_is_stream": true, + }); + let provider_body_json = json!({ + "id": "resp_1", + "object": "response", + "status": "completed", + "error": null, + "model": "probe-model", + "output": [{ + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "hello"}] + }], + "usage": {"input_tokens": 5, "output_tokens": 7, "total_tokens": 12} + }); + let body_base64 = base64::engine::general_purpose::STANDARD + .encode(serde_json::to_vec(&provider_body_json).expect("serialize provider body")); + + let product = maybe_build_standard_sync_finalize_product_from_normalized_payload( + "claude_chat_sync_finalize", + 200, + Some(&report_context), + None, + Some(&body_base64), + ) + .expect("a complete provider body must not fail the stream aggregator") + .expect("product should exist"); + + let StandardSyncFinalizeNormalizedProduct::CrossFormat(product) = product else { + panic!("cross-format attempt should produce a cross-format product"); + }; + assert_eq!(product.provider_body_json, provider_body_json); + assert_eq!(product.client_body_json["type"], "message"); + assert_eq!(product.client_body_json["content"][0]["text"], "hello"); + } + + #[test] + fn keeps_unframed_stream_event_on_the_aggregation_path() { + let report_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "claude:messages", + "needs_conversion": true, + "upstream_is_stream": true, + }); + let provider_body_json = json!({ + "id": "resp_1", + "object": "response", + "status": "completed", + "model": "probe-model", + "output": [{ + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "hello"}] + }], + "usage": {"input_tokens": 5, "output_tokens": 7, "total_tokens": 12} + }); + let event = json!({ + "type": "response.completed", + "response": provider_body_json.clone(), + }); + let body_base64 = base64::engine::general_purpose::STANDARD + .encode(serde_json::to_vec(&event).expect("serialize stream event")); + + let product = maybe_build_standard_sync_finalize_product_from_normalized_payload( + "claude_chat_sync_finalize", + 200, + Some(&report_context), + None, + Some(&body_base64), + ) + .expect("unframed stream event should aggregate") + .expect("product should exist"); + + let StandardSyncFinalizeNormalizedProduct::CrossFormat(product) = product else { + panic!("cross-format attempt should produce a cross-format product"); + }; + assert_eq!(product.provider_body_json["id"], provider_body_json["id"]); + assert_eq!(product.provider_body_json["object"], "response"); + assert!(product.provider_body_json.get("response").is_none()); + assert_eq!(product.client_body_json["type"], "message"); + } + #[test] fn builds_standard_same_format_body_from_stream_payload() { let body = concat!( diff --git a/crates/aether-ai/formats/src/formats/shared/sync_to_stream.rs b/crates/aether-ai/formats/src/formats/shared/sync_to_stream.rs index 9d2a38224..ff88144d1 100644 --- a/crates/aether-ai/formats/src/formats/shared/sync_to_stream.rs +++ b/crates/aether-ai/formats/src/formats/shared/sync_to_stream.rs @@ -12,6 +12,11 @@ use crate::formats::gemini::generate_content::stream::GeminiClientEmitter; use crate::formats::openai::chat::stream::{ OpenAIChatClientEmitter, OpenAIResponsesClientEmitter, OpenAIResponsesProviderState, }; +use crate::formats::openai::image::{ + bounded_openai_image_revised_prompt, is_safe_openai_image_base64_payload, + normalize_openai_image_output_format, parse_safe_openai_image_data_url, + sanitize_openai_image_source_url, +}; use crate::formats::openai::responses::history::{ record_converted_response_history, ResponseHistoryRecord, }; @@ -27,6 +32,8 @@ use crate::formats::shared::stream_core::{ use crate::formats::shared::stream_rewrite::maybe_build_ai_surface_stream_rewriter; use crate::formats::shared::AiSurfaceFinalizeError; +const MAX_OPENAI_IMAGE_OUTPUTS: usize = 64; + pub struct SyncToStreamBridgeOutcome { pub sse_body: Vec, pub terminal_summary: Option, @@ -198,7 +205,7 @@ fn maybe_bridge_openai_image_sync_json_to_stream( let Some(image) = outputs.iter().find_map(OpenAiImageOutput::b64_json) else { return Ok(None); }; - let image_count = openai_image_response_image_count(response).max(outputs.len() as u64); + let image_count = outputs.len() as u64; let usage = response.get("usage").cloned().unwrap_or(Value::Null); let event_name = openai_image_completed_event_name(report_context); let sse_body = encode_json_sse( @@ -238,7 +245,7 @@ fn maybe_bridge_openai_image_sync_json_to_chat_stream( return Ok(None); } - let image_count = openai_image_response_image_count(response).max(outputs.len() as u64); + let image_count = outputs.len() as u64; let summary = openai_image_terminal_summary(response, report_context, image_count); let response_id = openai_image_bridge_response_id(response, report_context, "chatcmpl-image"); let model = openai_image_bridge_response_model(response, report_context); @@ -349,7 +356,7 @@ fn maybe_bridge_openai_image_sync_json_to_responses_stream( }), )?); - let image_count = openai_image_response_image_count(response).max(outputs.len() as u64); + let image_count = outputs.len() as u64; Ok(Some(SyncToStreamBridgeOutcome { sse_body, terminal_summary: Some(openai_image_terminal_summary( @@ -372,17 +379,20 @@ struct OpenAiImageOutput { impl OpenAiImageOutput { fn b64_json(&self) -> Option { - self.b64_json - .clone() - .or_else(|| self.url.as_deref().and_then(extract_base64_from_data_url)) + let value = self + .b64_json + .as_deref() + .or_else(|| self.url.as_deref().and_then(extract_base64_from_data_url))?; + is_safe_openai_image_base64_payload(value).then(|| value.to_string()) } fn source_url(&self) -> Option { - self.url.clone().or_else(|| { + let source = self.url.clone().or_else(|| { self.b64_json .as_ref() .map(|value| format!("data:{};base64,{value}", self.mime_type)) - }) + })?; + sanitize_openai_image_source_url(&source) } fn markdown(&self, index: usize) -> String { @@ -391,7 +401,10 @@ impl OpenAiImageOutput { } else { format!("generated image {}", index + 1) }; - match self.source_url() { + match self + .source_url() + .and_then(|url| escape_markdown_image_destination(&url)) + { Some(url) => format!("![{alt}]({url})"), None => String::new(), } @@ -408,7 +421,11 @@ impl OpenAiImageOutput { Value::String("image_generation_call".to_string()), ); item.insert("status".to_string(), Value::String("completed".to_string())); - if let Some(result) = self.b64_json().or_else(|| self.url.clone()) { + if let Some(result) = self.b64_json().or_else(|| { + self.url + .as_deref() + .and_then(sanitize_openai_image_source_url) + }) { item.insert("result".to_string(), Value::String(result)); } if let Some(output_format) = self.output_format.as_ref() { @@ -470,6 +487,7 @@ fn collect_openai_image_outputs( .flatten() .filter_map(Value::as_object) .filter_map(|item| openai_image_output_from_item(item, report_context)) + .take(MAX_OPENAI_IMAGE_OUTPUTS) .collect() } @@ -483,7 +501,7 @@ fn openai_image_output_from_item( .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned); + .and_then(sanitize_openai_image_source_url); if b64_json.is_none() && url.is_none() { return None; } @@ -491,10 +509,13 @@ fn openai_image_output_from_item( .get("output_format") .or_else(|| item.get("format")) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(normalize_openai_image_output_format) .map(ToOwned::to_owned) - .or_else(|| image_request_output_format(report_context)); + .or_else(|| { + image_request_output_format(report_context).and_then(|value| { + normalize_openai_image_output_format(&value).map(ToOwned::to_owned) + }) + }); let mime_type = url .as_deref() .and_then(extract_mime_type_from_data_url) @@ -507,8 +528,7 @@ fn openai_image_output_from_item( let revised_prompt = item .get("revised_prompt") .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) + .and_then(bounded_openai_image_revised_prompt) .map(ToOwned::to_owned); Some(OpenAiImageOutput { @@ -520,14 +540,6 @@ fn openai_image_output_from_item( }) } -fn openai_image_response_image_count(response: &Map) -> u64 { - response - .get("data") - .and_then(Value::as_array) - .map(|items| items.len() as u64) - .unwrap_or(0) -} - fn openai_image_terminal_summary( response: &Map, report_context: Option<&Value>, @@ -680,11 +692,12 @@ fn image_request_quality(report_context: Option<&Value>) -> Option { } fn mime_type_from_image_output_format(output_format: &str) -> String { - match output_format.trim().to_ascii_lowercase().as_str() { - "jpg" | "jpeg" => "image/jpeg".to_string(), - "webp" => "image/webp".to_string(), - "png" => "image/png".to_string(), - value if !value.is_empty() => format!("image/{value}"), + match output_format.trim() { + value if value.eq_ignore_ascii_case("jpg") || value.eq_ignore_ascii_case("jpeg") => { + "image/jpeg".to_string() + } + value if value.eq_ignore_ascii_case("webp") => "image/webp".to_string(), + value if value.eq_ignore_ascii_case("png") => "image/png".to_string(), _ => "image/png".to_string(), } } @@ -913,29 +926,36 @@ fn extract_openai_image_sync_b64_json(item: &serde_json::Map) -> .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) + .filter(|value| is_safe_openai_image_base64_payload(value)) .map(ToOwned::to_owned) .or_else(|| { item.get("url") .and_then(Value::as_str) .and_then(extract_base64_from_data_url) + .map(ToOwned::to_owned) }) } -fn extract_base64_from_data_url(value: &str) -> Option { - let trimmed = value.trim(); - let (metadata, payload) = trimmed.split_once(',')?; - if !metadata.starts_with("data:") || !metadata.ends_with(";base64") { - return None; - } - (!payload.trim().is_empty()).then(|| payload.trim().to_string()) +fn extract_base64_from_data_url(value: &str) -> Option<&str> { + parse_safe_openai_image_data_url(value).map(|(_, payload)| payload) } fn extract_mime_type_from_data_url(value: &str) -> Option { - let trimmed = value.trim(); - let (metadata, _) = trimmed.split_once(',')?; - let mime_type = metadata.strip_prefix("data:")?.strip_suffix(";base64")?; - let mime_type = mime_type.trim(); - (!mime_type.is_empty()).then(|| mime_type.to_string()) + parse_safe_openai_image_data_url(value).map(|(mime_type, _)| mime_type.to_string()) +} + +fn escape_markdown_image_destination(value: &str) -> Option { + let mut escaped = String::with_capacity(value.len()); + for character in value.chars() { + if character.is_ascii_control() || character.is_whitespace() { + return None; + } + if matches!(character, '\\' | '(' | ')') { + escaped.push('\\'); + } + escaped.push(character); + } + Some(escaped) } fn openai_image_completed_event_name(report_context: Option<&Value>) -> &'static str { @@ -1320,7 +1340,10 @@ fn standardized_usage_from_openai_usage(value: &Value) -> Option) -> String { @@ -1890,6 +1913,98 @@ mod tests { assert!(output.contains("\"total_tokens\":9")); } + #[test] + fn bridges_openai_image_sync_http_url_with_markdown_escaping() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:chat", + "image_request": {"operation": "generate"} + }); + let outcome = maybe_bridge_standard_sync_json_to_stream( + &json!({ + "id": "img_url_123", + "data": [{ + "url": "https://cdn.example.test/generated/(image).png" + }] + }), + "openai:image", + "openai:chat", + Some(&report_context), + ) + .expect("bridge should succeed") + .expect("valid HTTP image URL should bridge"); + + let output = utf8(outcome.sse_body); + assert!(output.contains("generated/\\\\(image\\\\).png")); + } + + #[test] + fn rejects_non_http_image_urls_and_untrusted_data_mime_types() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:chat", + "image_request": {"operation": "generate"} + }); + for item in [ + json!({"url": "javascript:alert(1)"}), + json!({"url": "data:text/html;base64,PGh0bWw+"}), + ] { + let outcome = maybe_bridge_standard_sync_json_to_stream( + &json!({"data": [item]}), + "openai:image", + "openai:chat", + Some(&report_context), + ) + .expect("bridge should not error on an unsupported image source"); + assert!(outcome.is_none(), "unsupported source must not be emitted"); + } + + let output = OpenAiImageOutput { + b64_json: Some("aGVsbG8=".to_string()), + url: None, + mime_type: "image/png".to_string(), + output_format: Some("text/html);javascript:alert(1)".to_string()), + revised_prompt: None, + } + .markdown(0); + assert_eq!(output, "![generated image](data:image/png;base64,aGVsbG8=)"); + } + + #[test] + fn image_summary_counts_only_safe_emitted_outputs() { + let report_context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:chat", + "image_request": {"operation": "generate"} + }); + let outcome = maybe_bridge_standard_sync_json_to_stream( + &json!({ + "data": [ + {"b64_json": "not-base64!"}, + {"url": "javascript:alert(1)"}, + {"b64_json": "aGVsbG8="} + ], + "usage": {"input_tokens": 2, "output_tokens": 3, "total_tokens": 5} + }), + "openai:image", + "openai:chat", + Some(&report_context), + ) + .expect("bridge should ignore unsupported outputs") + .expect("valid output should still bridge"); + + let output = utf8(outcome.sse_body); + assert!(output.contains("data:image/png;base64,aGVsbG8=")); + assert!(!output.contains("not-base64")); + assert!(!output.contains("javascript:")); + let usage = outcome + .terminal_summary + .and_then(|summary| summary.standardized_usage) + .expect("safe output should produce usage"); + assert_eq!(usage.request_count, 1); + assert_eq!(usage.dimensions.get("image_count"), Some(&json!(1))); + } + #[test] fn bridges_aether_sse_response_capture_to_same_client_stream() { let captured_body = concat!( diff --git a/crates/aether-ai/formats/src/lib.rs b/crates/aether-ai/formats/src/lib.rs index 5dcce8c0b..be40c6d6f 100644 --- a/crates/aether-ai/formats/src/lib.rs +++ b/crates/aether-ai/formats/src/lib.rs @@ -57,6 +57,7 @@ pub use formats::openai::responses::request::{ validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation, }; 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, strip_incompatible_openai_responses_reasoning_items, strip_incompatible_openai_responses_reasoning_items_with_policy, diff --git a/crates/aether-ai/formats/src/protocol/canonical.rs b/crates/aether-ai/formats/src/protocol/canonical.rs index e16d1c7e4..55b0d7661 100644 --- a/crates/aether-ai/formats/src/protocol/canonical.rs +++ b/crates/aether-ai/formats/src/protocol/canonical.rs @@ -1,8 +1,10 @@ use std::collections::{BTreeMap, BTreeSet, VecDeque}; +use std::fmt; use serde::{Deserialize, Serialize}; use serde_json::{json, Map, Value}; +use crate::formats::openai::responses::openai_responses_message_item_id; use crate::formats::openai::responses::{ decode_gemini_tool_signature_carrier, GeminiToolSignatureCarrierDirection, }; @@ -56,7 +58,7 @@ pub enum CanonicalStopReason { Unknown, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum CanonicalToolChoice { Auto, @@ -65,7 +67,7 @@ pub enum CanonicalToolChoice { Tool { name: String }, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum CanonicalContentBlock { Text { @@ -147,7 +149,7 @@ pub enum CanonicalContentBlock { }, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalInstruction { pub role: CanonicalRole, #[serde(default)] @@ -156,7 +158,7 @@ pub struct CanonicalInstruction { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalMessage { pub role: CanonicalRole, #[serde(default)] @@ -165,7 +167,7 @@ pub struct CanonicalMessage { pub extensions: BTreeMap, } -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[derive(Clone, Default, PartialEq, Serialize, Deserialize)] pub struct CanonicalGenerationConfig { #[serde(default, skip_serializing_if = "Option::is_none")] pub max_tokens: Option, @@ -191,7 +193,7 @@ pub struct CanonicalGenerationConfig { pub top_logprobs: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalToolDefinition { pub name: String, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -204,7 +206,7 @@ pub struct CanonicalToolDefinition { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalThinkingConfig { #[serde(default)] pub enabled: bool, @@ -214,7 +216,7 @@ pub struct CanonicalThinkingConfig { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalResponseFormat { pub format_type: String, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -223,7 +225,7 @@ pub struct CanonicalResponseFormat { pub extensions: BTreeMap, } -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[derive(Clone, Default, PartialEq, Serialize, Deserialize)] pub struct CanonicalUsage { #[serde(default)] pub input_tokens: u64, @@ -253,7 +255,7 @@ fn is_false(value: &bool) -> bool { !*value } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] #[serde(untagged)] pub enum CanonicalEmbeddingInput { String(String), @@ -263,7 +265,7 @@ pub enum CanonicalEmbeddingInput { Multimodal(Vec), } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalEmbeddingContent { #[serde(default, skip_serializing_if = "Option::is_none")] pub text: Option, @@ -335,7 +337,7 @@ impl CanonicalEmbeddingContent { } } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalEmbeddingRequest { pub input: CanonicalEmbeddingInput, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -352,7 +354,7 @@ pub struct CanonicalEmbeddingRequest { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalRerankRequest { pub query: String, #[serde(default)] @@ -373,7 +375,7 @@ impl CanonicalRerankRequest { } } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalEmbedding { #[serde(default)] pub index: usize, @@ -383,7 +385,7 @@ pub struct CanonicalEmbedding { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalEmbeddingResponse { pub id: String, pub model: String, @@ -395,7 +397,7 @@ pub struct CanonicalEmbeddingResponse { pub extensions: BTreeMap, } -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[derive(Clone, Default, PartialEq, Serialize, Deserialize)] pub struct CanonicalRequest { #[serde(default)] pub model: String, @@ -427,7 +429,7 @@ pub struct CanonicalRequest { pub extensions: BTreeMap, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalResponseOutput { #[serde(default)] pub index: usize, @@ -453,7 +455,7 @@ impl Default for CanonicalResponseOutput { } } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct CanonicalResponse { pub id: String, pub model: String, @@ -469,6 +471,425 @@ pub struct CanonicalResponse { pub extensions: BTreeMap, } +fn debug_json_bytes(value: &Value) -> Option { + serde_json::to_vec(value).ok().map(|bytes| bytes.len()) +} + +fn debug_json_map_bytes(value: &Map) -> Option { + serde_json::to_vec(value).ok().map(|bytes| bytes.len()) +} + +fn debug_json_option_bytes(value: Option<&Value>) -> Option { + value.and_then(debug_json_bytes) +} + +fn debug_string_len(value: Option<&str>) -> Option { + value.map(str::len) +} + +fn debug_string_list_summary(value: Option<&Vec>) -> Option<(usize, usize)> { + value.map(|values| (values.len(), values.iter().map(String::len).sum::())) +} + +fn debug_json_list_summary(value: &[Value]) -> (usize, usize) { + ( + value.len(), + value.iter().filter_map(debug_json_bytes).sum::(), + ) +} + +impl fmt::Debug for CanonicalToolChoice { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("CanonicalToolChoice"); + match self { + Self::Auto => debug.field("kind", &"auto"), + Self::None => debug.field("kind", &"none"), + Self::Required => debug.field("kind", &"required"), + Self::Tool { name } => debug.field("kind", &"tool").field("name_len", &name.len()), + } + .finish() + } +} + +impl fmt::Debug for CanonicalContentBlock { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("CanonicalContentBlock"); + match self { + Self::Text { text, extensions } => debug + .field("kind", &"text") + .field("text_len", &text.len()) + .field("extension_count", &extensions.len()), + Self::Thinking { + text, + signature, + encrypted_content, + extensions, + } => debug + .field("kind", &"thinking") + .field("text_len", &text.len()) + .field("signature_len", &debug_string_len(signature.as_deref())) + .field( + "encrypted_content_len", + &debug_string_len(encrypted_content.as_deref()), + ) + .field("extension_count", &extensions.len()), + Self::Image { + data, + url, + media_type, + detail, + extensions, + } => debug + .field("kind", &"image") + .field("data_len", &debug_string_len(data.as_deref())) + .field("url_len", &debug_string_len(url.as_deref())) + .field("media_type", media_type) + .field("detail", detail) + .field("extension_count", &extensions.len()), + Self::File { + data, + file_id, + file_url, + media_type, + filename, + extensions, + } => debug + .field("kind", &"file") + .field("data_len", &debug_string_len(data.as_deref())) + .field("file_id_len", &debug_string_len(file_id.as_deref())) + .field("file_url_len", &debug_string_len(file_url.as_deref())) + .field("media_type", media_type) + .field("filename_len", &debug_string_len(filename.as_deref())) + .field("extension_count", &extensions.len()), + Self::Audio { + data, + media_type, + format, + extensions, + } => debug + .field("kind", &"audio") + .field("data_len", &debug_string_len(data.as_deref())) + .field("media_type", media_type) + .field("format", format) + .field("extension_count", &extensions.len()), + Self::ToolUse { + id, + name, + input, + extensions, + } => debug + .field("kind", &"tool_use") + .field("id_len", &id.len()) + .field("name_len", &name.len()) + .field("input_bytes", &debug_json_bytes(input)) + .field("extension_count", &extensions.len()), + Self::ToolResult { + tool_use_id, + name, + output, + content_text, + is_error, + extensions, + } => debug + .field("kind", &"tool_result") + .field("tool_use_id_len", &tool_use_id.len()) + .field("name_len", &debug_string_len(name.as_deref())) + .field("output_bytes", &debug_json_option_bytes(output.as_ref())) + .field( + "content_text_len", + &debug_string_len(content_text.as_deref()), + ) + .field("is_error", is_error) + .field("extension_count", &extensions.len()), + Self::Unknown { + raw_type, + payload, + extensions, + } => debug + .field("kind", &"unknown") + .field("raw_type_len", &raw_type.len()) + .field("payload_bytes", &debug_json_bytes(payload)) + .field("extension_count", &extensions.len()), + } + .finish() + } +} + +impl fmt::Debug for CanonicalInstruction { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalInstruction") + .field("role", &self.role) + .field("text_len", &self.text.len()) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalMessage { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalMessage") + .field("role", &self.role) + .field("content_count", &self.content.len()) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalGenerationConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalGenerationConfig") + .field("max_tokens", &self.max_tokens) + .field("temperature", &self.temperature) + .field("top_p", &self.top_p) + .field("top_k", &self.top_k) + .field( + "stop_sequences", + &debug_string_list_summary(self.stop_sequences.as_ref()), + ) + .field("n", &self.n) + .field("presence_penalty", &self.presence_penalty) + .field("frequency_penalty", &self.frequency_penalty) + .field("seed", &self.seed) + .field("logprobs", &self.logprobs) + .field("top_logprobs", &self.top_logprobs) + .finish() + } +} + +impl fmt::Debug for CanonicalToolDefinition { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalToolDefinition") + .field("name_len", &self.name.len()) + .field( + "description_len", + &debug_string_len(self.description.as_deref()), + ) + .field( + "parameters_bytes", + &debug_json_option_bytes(self.parameters.as_ref()), + ) + .field("strict", &self.strict) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalThinkingConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalThinkingConfig") + .field("enabled", &self.enabled) + .field("budget_tokens", &self.budget_tokens) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalResponseFormat { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalResponseFormat") + .field("format_type_len", &self.format_type.len()) + .field( + "json_schema_bytes", + &debug_json_option_bytes(self.json_schema.as_ref()), + ) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalUsage { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalUsage") + .field("input_tokens", &self.input_tokens) + .field( + "input_tokens_include_cache", + &self.input_tokens_include_cache, + ) + .field("output_tokens", &self.output_tokens) + .field("total_tokens", &self.total_tokens) + .field("cache_read_tokens", &self.cache_read_tokens) + .field("cache_write_tokens", &self.cache_write_tokens) + .field( + "cache_creation_ephemeral_5m_tokens", + &self.cache_creation_ephemeral_5m_tokens, + ) + .field( + "cache_creation_ephemeral_1h_tokens", + &self.cache_creation_ephemeral_1h_tokens, + ) + .field("reasoning_tokens", &self.reasoning_tokens) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalEmbeddingInput { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("CanonicalEmbeddingInput"); + match self { + Self::String(value) => debug + .field("kind", &"string") + .field("item_count", &1) + .field("total_text_bytes", &value.len()), + Self::StringArray(values) => debug + .field("kind", &"string_array") + .field("item_count", &values.len()) + .field( + "total_text_bytes", + &values.iter().map(String::len).sum::(), + ), + Self::TokenArray(values) => debug + .field("kind", &"token_array") + .field("item_count", &values.len()), + Self::TokenArrayArray(values) => debug + .field("kind", &"token_array_array") + .field("item_count", &values.len()) + .field( + "total_token_count", + &values.iter().map(Vec::len).sum::(), + ), + Self::Multimodal(values) => debug + .field("kind", &"multimodal") + .field("item_count", &values.len()), + } + .finish() + } +} + +impl fmt::Debug for CanonicalEmbeddingContent { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalEmbeddingContent") + .field("text_len", &debug_string_len(self.text.as_deref())) + .field("image_len", &debug_string_len(self.image.as_deref())) + .field("video_len", &debug_string_len(self.video.as_deref())) + .field( + "multi_images_summary", + &self + .multi_images + .as_ref() + .map(|images| (images.len(), images.iter().map(String::len).sum::())), + ) + .finish() + } +} + +impl fmt::Debug for CanonicalEmbeddingRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalEmbeddingRequest") + .field("input", &self.input) + .field("encoding_format", &self.encoding_format) + .field("dimensions", &self.dimensions) + .field("task_len", &debug_string_len(self.task.as_deref())) + .field("user_len", &debug_string_len(self.user.as_deref())) + .field( + "parameters_bytes", + &self.parameters.as_ref().and_then(debug_json_map_bytes), + ) + .field("parameter_count", &self.parameters.as_ref().map(Map::len)) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalRerankRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalRerankRequest") + .field("query_len", &self.query.len()) + .field("documents", &debug_json_list_summary(&self.documents)) + .field("top_n", &self.top_n) + .field("return_documents", &self.return_documents) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalEmbedding { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalEmbedding") + .field("index", &self.index) + .field("embedding_len", &self.embedding.len()) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalEmbeddingResponse { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalEmbeddingResponse") + .field("id_len", &self.id.len()) + .field("model_len", &self.model.len()) + .field("embedding_count", &self.embeddings.len()) + .field("usage", &self.usage) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalRequest") + .field("model_len", &self.model.len()) + .field("instruction_count", &self.instructions.len()) + .field("system_len", &debug_string_len(self.system.as_deref())) + .field("message_count", &self.messages.len()) + .field("embedding", &self.embedding) + .field("rerank", &self.rerank) + .field("generation", &self.generation) + .field("tool_count", &self.tools.len()) + .field("tool_choice", &self.tool_choice) + .field("thinking", &self.thinking) + .field("response_format", &self.response_format) + .field("parallel_tool_calls", &self.parallel_tool_calls) + .field( + "metadata_bytes", + &debug_json_option_bytes(self.metadata.as_ref()), + ) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalResponseOutput { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalResponseOutput") + .field("index", &self.index) + .field("role", &self.role) + .field("content_count", &self.content.len()) + .field("stop_reason", &self.stop_reason) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + +impl fmt::Debug for CanonicalResponse { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CanonicalResponse") + .field("id_len", &self.id.len()) + .field("model_len", &self.model.len()) + .field("output_count", &self.outputs.len()) + .field("content_count", &self.content.len()) + .field("stop_reason", &self.stop_reason) + .field("usage", &self.usage) + .field("extension_count", &self.extensions.len()) + .finish() + } +} + pub fn from_openai_chat_to_canonical_request(body_json: &Value) -> Option { crate::formats::openai::chat::request::from_raw(body_json) } @@ -1248,19 +1669,21 @@ pub(crate) fn gemini_part_to_canonical_block( ) -> Option { let part_object = part.as_object()?; if let Some(text) = part_object.get("text").and_then(Value::as_str) { - if part_object + let thought_signature = part_object + .get("thoughtSignature") + .or_else(|| part_object.get("thought_signature")) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let is_thinking = part_object .get("thought") .and_then(Value::as_bool) .unwrap_or(false) - { + || (text.trim().is_empty() && thought_signature.is_some()); + if is_thinking { return Some(CanonicalContentBlock::Thinking { text: text.to_string(), - signature: part_object - .get("thoughtSignature") - .or_else(|| part_object.get("thought_signature")) - .and_then(Value::as_str) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned), + signature: thought_signature, encrypted_content: None, extensions: gemini_extensions( part_object, @@ -4102,11 +4525,7 @@ pub(crate) fn flush_openai_responses_message_item( if message_content.is_empty() { return; } - let id = if *message_index == 0 { - format!("{response_id}_msg") - } else { - format!("{response_id}_msg_{message_index}") - }; + let id = openai_responses_message_item_id(response_id, *message_index); output.push(json!({ "type": "message", "id": id, @@ -4532,7 +4951,7 @@ pub(crate) type GeminiCanonicalTools = ( Option, ); -#[derive(Debug, Clone)] +#[derive(Clone)] pub(crate) struct GeminiGoogleSearchGrounding { pub source_field: &'static str, pub source_dialect: &'static str, @@ -4542,6 +4961,23 @@ pub(crate) struct GeminiGoogleSearchGrounding { pub output_payload: Value, } +impl fmt::Debug for GeminiGoogleSearchGrounding { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GeminiGoogleSearchGrounding") + .field("source_field", &self.source_field) + .field("source_dialect", &self.source_dialect) + .field("legacy", &self.legacy) + .field("payload_bytes", &debug_json_bytes(&self.payload)) + .field("raw_payload_bytes", &debug_json_bytes(&self.raw_payload)) + .field( + "output_payload_bytes", + &debug_json_bytes(&self.output_payload), + ) + .finish() + } +} + pub(crate) fn gemini_google_search_grounding( tool_object: &Map, ) -> 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 4eb01ad0f..d63f0e858 100644 --- a/crates/aether-ai/formats/src/provider_compat/private_envelope.rs +++ b/crates/aether-ai/formats/src/provider_compat/private_envelope.rs @@ -256,6 +256,10 @@ fn transform_provider_private_stream_line_with_event_state( const CONNECT_FRAME_HEADER_BYTES: usize = 5; const MAX_CONNECT_JSON_FRAME_BYTES: usize = 16 * 1024 * 1024; +// Keep malformed or incomplete provider streams from growing this parser's +// carry buffer without bound when no complete SSE/Connect record arrives. +const MAX_PRIVATE_STREAM_BUFFER_BYTES: usize = + MAX_CONNECT_JSON_FRAME_BYTES + CONNECT_FRAME_HEADER_BYTES; fn report_context_is_windsurf_envelope(report_context: &Value) -> bool { report_context @@ -424,6 +428,18 @@ impl ProviderPrivateStreamNormalizer<'_> { state.push_chunk(self.report_context, chunk) } ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => { + let next_len = self + .buffered + .len() + .checked_add(chunk.len()) + .ok_or_else(|| { + AiSurfaceFinalizeError::new("provider stream normalization buffer overflow") + })?; + if next_len > MAX_PRIVATE_STREAM_BUFFER_BYTES { + return Err(AiSurfaceFinalizeError::new(format!( + "provider stream normalization buffer exceeds {MAX_PRIVATE_STREAM_BUFFER_BYTES} bytes" + ))); + } self.buffered.extend_from_slice(chunk); if report_context_is_windsurf_envelope(self.report_context) && buffer_looks_like_connect_frame(&self.buffered) @@ -1282,4 +1298,21 @@ data: {"type":"response.failed","response":{"status":"failed","error":{"message" assert!(output_text.contains("\"message\":\"rate limited\"")); assert!(!output_text.contains("chat.completion.chunk")); } + + #[test] + fn private_stream_normalizer_rejects_unbounded_incomplete_buffer() { + let report_context = json!({ + "has_envelope": true, + "envelope_name": "antigravity:v1internal", + "provider_api_format": "gemini:generate_content", + }); + let mut normalizer = maybe_build_provider_private_stream_normalizer(Some(&report_context)) + .expect("normalizer should exist"); + let oversized = vec![b'x'; super::MAX_PRIVATE_STREAM_BUFFER_BYTES + 1]; + + let error = normalizer + .push_chunk(&oversized) + .expect_err("incomplete provider stream must be bounded"); + assert!(error.0.contains("buffer exceeds")); + } } diff --git a/crates/aether-ai/serving/src/attempt_loop.rs b/crates/aether-ai/serving/src/attempt_loop.rs index c2df49cfc..ed3c21a45 100644 --- a/crates/aether-ai/serving/src/attempt_loop.rs +++ b/crates/aether-ai/serving/src/attempt_loop.rs @@ -13,8 +13,23 @@ pub trait AiExecutionAttempt { fn report_context_ref(&self) -> Option<&serde_json::Value> { None } + + /// Re-issue this attempt against the same key as a fresh attempt with the + /// given retry index and candidate id. Attempt types that cannot be + /// re-issued return `None`, which disables same-key retries for them. + fn with_same_key_retry(&self, _retry_index: u32, _candidate_id: String) -> Option + where + Self: Sized, + { + None + } } +/// Report-context field carrying the routing policy's sticky-key attempt +/// budget for the request, so the attempt loop can derive same-key retries +/// lazily instead of pre-materializing them. +pub const STICKY_KEY_ATTEMPTS_REPORT_FIELD: &str = "sticky_key_attempts"; + #[derive(Debug)] pub enum AiAttemptLoopOutcome { Responded(Response), @@ -83,6 +98,17 @@ where Ok(()) } + /// After `attempt` failed with candidate scope, return the next attempt on + /// the same key, or `None` once the sticky-key budget is used up. Retries + /// are derived here on demand so no attempt is materialized before it is + /// actually needed. + async fn next_same_key_retry( + &self, + _attempt: &Attempt, + ) -> Result, Self::Error> { + Ok(None) + } + async fn mark_unused_attempts(&self, attempts: Vec) -> Result<(), Self::Error>; async fn build_exhaustion( @@ -101,11 +127,15 @@ where Attempt: AiExecutionAttempt + Send + Sync + 'static, { let mut remaining = attempts.into_iter(); + let mut pending_same_key_retry: Option = None; let mut last_attempted = None; let mut retry_filters: Vec = Vec::new(); let mut fallback_response = None; - while let Some(attempt) = remaining.next() { + loop { + let Some(attempt) = pending_same_key_retry.take().or_else(|| remaining.next()) else { + break; + }; if retry_filters.iter().any(|filter| filter.matches(&attempt)) || port.should_skip_attempt(&attempt).await? { @@ -133,7 +163,9 @@ where if attempt_fallback_response.is_some() { fallback_response = attempt_fallback_response; } - if scope != AiAttemptRetryScope::Candidate { + if scope == AiAttemptRetryScope::Candidate { + pending_same_key_retry = port.next_same_key_retry(&attempt).await?; + } else { retry_filters.push(AiAttemptRetryFilter::new(&attempt, scope)); } } @@ -188,6 +220,32 @@ impl AiAttemptRetryFilter { } } +/// Clone `plan`/`report_context` for a same-key retry: only the candidate id +/// and retry index change, everything else (url, headers, body) is reused. +fn same_key_retry_parts( + plan: &aether_contracts::ExecutionPlan, + report_context: Option<&serde_json::Value>, + retry_index: u32, + candidate_id: String, +) -> (aether_contracts::ExecutionPlan, Option) { + let mut plan = plan.clone(); + plan.candidate_id = Some(candidate_id.clone()); + let report_context = report_context.cloned().map(|mut value| { + if let Some(object) = value.as_object_mut() { + object.insert( + "candidate_id".to_string(), + serde_json::Value::String(candidate_id), + ); + object.insert( + "retry_index".to_string(), + serde_json::Value::Number(retry_index.into()), + ); + } + value + }); + (plan, report_context) +} + impl AiExecutionAttempt for crate::dto::AiSyncAttempt { fn execution_plan(&self) -> &aether_contracts::ExecutionPlan { &self.plan @@ -204,6 +262,20 @@ impl AiExecutionAttempt for crate::dto::AiSyncAttempt { fn report_context_ref(&self) -> Option<&serde_json::Value> { self.report_context.as_ref() } + + fn with_same_key_retry(&self, retry_index: u32, candidate_id: String) -> Option { + let (plan, report_context) = same_key_retry_parts( + &self.plan, + self.report_context.as_ref(), + retry_index, + candidate_id, + ); + Some(Self { + plan, + report_kind: self.report_kind.clone(), + report_context, + }) + } } impl AiExecutionAttempt for crate::dto::AiStreamAttempt { @@ -222,6 +294,20 @@ impl AiExecutionAttempt for crate::dto::AiStreamAttempt { fn report_context_ref(&self) -> Option<&serde_json::Value> { self.report_context.as_ref() } + + fn with_same_key_retry(&self, retry_index: u32, candidate_id: String) -> Option { + let (plan, report_context) = same_key_retry_parts( + &self.plan, + self.report_context.as_ref(), + retry_index, + candidate_id, + ); + Some(Self { + plan, + report_kind: self.report_kind.clone(), + report_context, + }) + } } #[cfg(test)] diff --git a/crates/aether-ai/serving/src/attempt_plan.rs b/crates/aether-ai/serving/src/attempt_plan.rs index 63bccbd3e..441b10af7 100644 --- a/crates/aether-ai/serving/src/attempt_plan.rs +++ b/crates/aether-ai/serving/src/attempt_plan.rs @@ -1,7 +1,8 @@ use std::collections::BTreeMap; +use std::fmt; use aether_ai_formats::api::ExecutionRuntimeAuthContext; -use aether_contracts::{ExecutionPlan, RequestBody}; +use aether_contracts::{redact_url_for_debug, ExecutionPlan, RequestBody}; use url::Url; use crate::dto::{AiExecutionDecision, AiRequestGzipPolicy}; @@ -18,13 +19,23 @@ pub struct AiDecisionPlanCore { pub client_api_format: String, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct AiUpstreamAuthPair { pub header: String, pub value: String, } -#[derive(Debug)] +impl fmt::Debug for AiUpstreamAuthPair { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiUpstreamAuthPair") + .field("header", &self.header) + .field("has_value", &(!self.value.is_empty())) + .field("value_len", &self.value.len()) + .finish() + } +} + pub struct AiExecutionPlanFromDecisionParts { pub core: AiDecisionPlanCore, pub method: String, @@ -35,7 +46,21 @@ pub struct AiExecutionPlanFromDecisionParts { pub stream: bool, } -#[derive(Debug)] +impl fmt::Debug for AiExecutionPlanFromDecisionParts { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiExecutionPlanFromDecisionParts") + .field("core", &self.core) + .field("method", &self.method) + .field("url", &redact_url_for_debug(&self.url)) + .field("header_names", &self.headers.keys().collect::>()) + .field("content_type", &self.content_type) + .field("body", &self.body) + .field("stream", &self.stream) + .finish() + } +} + pub struct AiExecutionDecisionFromPlanParts { pub action: String, pub decision_kind: Option, @@ -48,6 +73,33 @@ pub struct AiExecutionDecisionFromPlanParts { pub auth_context: Option, } +impl fmt::Debug for AiExecutionDecisionFromPlanParts { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiExecutionDecisionFromPlanParts") + .field("action", &self.action) + .field("decision_kind", &self.decision_kind) + .field("request_id", &self.request_id) + .field( + "upstream_base_url", + &self.upstream_base_url.as_deref().map(redact_url_for_debug), + ) + .field("include_auth_pair", &self.include_auth_pair) + .field("plan", &self.plan) + .field("report_kind", &self.report_kind) + .field("has_report_context", &self.report_context.is_some()) + .field( + "report_context_bytes", + &self + .report_context + .as_ref() + .and_then(|value| serde_json::to_vec(value).ok().map(|bytes| bytes.len())), + ) + .field("has_auth_context", &self.auth_context.is_some()) + .finish() + } +} + pub fn take_ai_non_empty_string(value: &mut Option) -> Option { value.take().filter(|value| !value.trim().is_empty()) } diff --git a/crates/aether-ai/serving/src/candidate_persistence.rs b/crates/aether-ai/serving/src/candidate_persistence.rs index 287b67d2f..7fe326f6c 100644 --- a/crates/aether-ai/serving/src/candidate_persistence.rs +++ b/crates/aether-ai/serving/src/candidate_persistence.rs @@ -9,8 +9,6 @@ pub trait AiAvailableCandidatePersistencePort: Send + Sync { type ExtraData: Clone + Send + Sync; type Error: Send; - fn attempt_slot_count(&self, candidate: &Self::Candidate) -> u32; - fn build_extra_data(&self, candidate: &Self::Candidate) -> Option; fn generate_candidate_id(&self) -> String; @@ -42,50 +40,28 @@ pub async fn run_ai_available_candidate_persistence( where Port: AiAvailableCandidatePersistencePort, { - let total_attempts = candidates - .iter() - .map(|candidate| port.attempt_slot_count(candidate) as usize) - .sum(); - let mut materialized = Vec::with_capacity(total_attempts); + // One attempt per candidate. Same-key retries are derived lazily by the + // attempt loop (`AiAttemptLoopPort::next_same_key_retry`) after a failure, + // so the sticky-key budget never inflates up-front materialization. + let mut materialized = Vec::with_capacity(candidates.len()); for (candidate_index, candidate) in candidates.into_iter().enumerate() { let candidate_index = candidate_index as u32; - let attempt_slots = port.attempt_slot_count(&candidate).max(1); let extra_data = port.build_extra_data(&candidate); - let mut owned_candidate = Some(candidate); - - for retry_index in 0..attempt_slots { - let candidate = owned_candidate - .as_ref() - .expect("candidate should remain available until final retry"); - let generated_candidate_id = port.generate_candidate_id(); - let candidate_id = if port.should_persist_available_candidate(candidate) { - port.persist_available_candidate( - candidate, - candidate_index, - retry_index, - generated_candidate_id.as_str(), - extra_data.clone(), - ) - .await? - } else { - generated_candidate_id - }; - - let candidate = if retry_index + 1 == attempt_slots { - owned_candidate - .take() - .expect("final retry should consume owned candidate") - } else { - candidate.clone() - }; - materialized.push(port.build_attempt( - candidate, + let generated_candidate_id = port.generate_candidate_id(); + let candidate_id = if port.should_persist_available_candidate(&candidate) { + port.persist_available_candidate( + &candidate, candidate_index, - retry_index, - candidate_id, - )); - } + 0, + generated_candidate_id.as_str(), + extra_data, + ) + .await? + } else { + generated_candidate_id + }; + materialized.push(port.build_attempt(candidate, candidate_index, 0, candidate_id)); } Ok(materialized) @@ -177,7 +153,6 @@ mod tests { #[derive(Debug, Clone, PartialEq, Eq)] struct TestCandidate { id: &'static str, - attempt_slots: u32, persist: bool, } @@ -216,10 +191,6 @@ mod tests { type ExtraData = String; type Error = std::convert::Infallible; - fn attempt_slot_count(&self, candidate: &Self::Candidate) -> u32 { - candidate.attempt_slots - } - fn build_extra_data(&self, candidate: &Self::Candidate) -> Option { Some(format!("extra:{}", candidate.id)) } @@ -299,7 +270,7 @@ mod tests { } #[tokio::test] - async fn available_persistence_expands_candidates_into_retry_attempts() { + async fn available_persistence_materializes_one_attempt_per_candidate() { let port = TestPort::default(); let attempts = run_ai_available_candidate_persistence( @@ -307,12 +278,10 @@ mod tests { vec![ TestCandidate { id: "a", - attempt_slots: 2, persist: true, }, TestCandidate { id: "b", - attempt_slots: 1, persist: false, }, ], @@ -320,6 +289,8 @@ mod tests { .await .unwrap(); + // Same-key retries are never pre-materialized; the attempt loop + // derives them on demand after a failure. assert_eq!( attempts, [ @@ -329,26 +300,17 @@ mod tests { retry_index: 0, candidate_id: "stored-candidate-1".to_string(), }, - TestAttempt { - id: "a", - candidate_index: 0, - retry_index: 1, - candidate_id: "stored-candidate-2".to_string(), - }, TestAttempt { id: "b", candidate_index: 1, retry_index: 0, - candidate_id: "candidate-3".to_string(), + candidate_id: "candidate-2".to_string(), }, ] ); assert_eq!( port.calls.lock().unwrap().as_slice(), - [ - "available:a:0:0:candidate-1:extra:a", - "available:a:0:1:candidate-2:extra:a", - ] + ["available:a:0:0:candidate-1:extra:a"] ); } diff --git a/crates/aether-ai/serving/src/candidate_preparation.rs b/crates/aether-ai/serving/src/candidate_preparation.rs index 936d8cee3..bbb46bc79 100644 --- a/crates/aether-ai/serving/src/candidate_preparation.rs +++ b/crates/aether-ai/serving/src/candidate_preparation.rs @@ -1,10 +1,24 @@ -#[derive(Debug, Clone, PartialEq, Eq)] +use std::fmt; + +#[derive(Clone, PartialEq, Eq)] pub struct AiPreparedHeaderAuthenticatedCandidate { pub auth_header: String, pub auth_value: String, pub mapped_model: String, } +impl fmt::Debug for AiPreparedHeaderAuthenticatedCandidate { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiPreparedHeaderAuthenticatedCandidate") + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field("auth_value_len", &self.auth_value.len()) + .field("mapped_model", &self.mapped_model) + .finish() + } +} + pub fn prepare_ai_header_authenticated_candidate( direct_auth: Option<(String, String)>, oauth_header_auth: Option<(String, String)>, diff --git a/crates/aether-ai/serving/src/dto.rs b/crates/aether-ai/serving/src/dto.rs index c8412e476..702e0efba 100644 --- a/crates/aether-ai/serving/src/dto.rs +++ b/crates/aether-ai/serving/src/dto.rs @@ -1,7 +1,10 @@ use std::collections::BTreeMap; +use std::fmt; use aether_ai_formats::api::ExecutionRuntimeAuthContext; -use aether_contracts::{ExecutionPlan, ExecutionTimeouts, ProxySnapshot, ResolvedTransportProfile}; +use aether_contracts::{ + redact_url_for_debug, ExecutionPlan, ExecutionTimeouts, ProxySnapshot, ResolvedTransportProfile, +}; use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] @@ -74,7 +77,7 @@ pub struct AiRequestGzipPolicy { pub min_bytes: Option, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Deserialize, Serialize)] pub struct AiExecutionPlanPayload { pub action: String, #[serde(default)] @@ -89,7 +92,25 @@ pub struct AiExecutionPlanPayload { pub auth_context: Option, } -#[derive(Debug, Clone, Deserialize, Serialize)] +impl fmt::Debug for AiExecutionPlanPayload { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiExecutionPlanPayload") + .field("action", &self.action) + .field("plan_kind", &self.plan_kind) + .field("plan", &self.plan) + .field("report_kind", &self.report_kind) + .field("has_report_context", &self.report_context.is_some()) + .field( + "report_context_bytes", + &json_value_len(&self.report_context), + ) + .field("has_auth_context", &self.auth_context.is_some()) + .finish() + } +} + +#[derive(Clone, Deserialize, Serialize)] pub struct AiExecutionDecision { pub action: String, #[serde(default)] @@ -166,20 +187,132 @@ pub struct AiExecutionDecision { pub auth_context: Option, } -#[derive(Debug)] +impl fmt::Debug for AiExecutionDecision { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("AiExecutionDecision"); + debug + .field("action", &self.action) + .field("decision_kind", &self.decision_kind) + .field("execution_strategy", &self.execution_strategy) + .field("conversion_mode", &self.conversion_mode) + .field("request_id", &self.request_id) + .field("candidate_id", &self.candidate_id) + .field("provider_name", &self.provider_name) + .field("provider_type", &self.provider_type) + .field("provider_id", &self.provider_id) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field( + "upstream_base_url", + &self.upstream_base_url.as_deref().map(redact_url_for_debug), + ) + .field( + "upstream_url", + &self.upstream_url.as_deref().map(redact_url_for_debug), + ) + .field("provider_request_method", &self.provider_request_method) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &self.auth_value.is_some()) + .field("auth_value_len", &self.auth_value.as_ref().map(String::len)) + .field("provider_api_format", &self.provider_api_format) + .field("client_api_format", &self.client_api_format) + .field("provider_contract", &self.provider_contract) + .field("client_contract", &self.client_contract) + .field("model_name", &self.model_name) + .field("mapped_model", &self.mapped_model) + .field("has_prompt_cache_key", &self.prompt_cache_key.is_some()) + .field( + "prompt_cache_key_len", + &self.prompt_cache_key.as_ref().map(String::len), + ) + .field( + "extra_header_names", + &self.extra_headers.keys().collect::>(), + ) + .field( + "provider_request_header_names", + &self.provider_request_headers.keys().collect::>(), + ) + .field( + "has_provider_request_body", + &self.provider_request_body.is_some(), + ) + .field( + "provider_request_body_bytes", + &json_value_len(&self.provider_request_body), + ) + .field( + "provider_request_body_base64_len", + &self.provider_request_body_base64.as_ref().map(String::len), + ) + .field("content_type", &self.content_type) + .field("content_encoding", &self.content_encoding) + .field("request_gzip", &self.request_gzip) + .field("proxy", &self.proxy) + .field("transport_profile", &self.transport_profile) + .field("timeouts", &self.timeouts) + .field("upstream_is_stream", &self.upstream_is_stream) + .field("report_kind", &self.report_kind) + .field("has_report_context", &self.report_context.is_some()) + .field( + "report_context_bytes", + &json_value_len(&self.report_context), + ) + .field("has_auth_context", &self.auth_context.is_some()) + .finish() + } +} + +#[derive(Clone)] pub struct AiSyncAttempt { pub plan: ExecutionPlan, pub report_kind: Option, pub report_context: Option, } -#[derive(Debug)] +impl fmt::Debug for AiSyncAttempt { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiSyncAttempt") + .field("plan", &self.plan) + .field("report_kind", &self.report_kind) + .field("has_report_context", &self.report_context.is_some()) + .field( + "report_context_bytes", + &json_value_len(&self.report_context), + ) + .finish() + } +} + +#[derive(Clone)] pub struct AiStreamAttempt { pub plan: ExecutionPlan, pub report_kind: Option, pub report_context: Option, } +impl fmt::Debug for AiStreamAttempt { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AiStreamAttempt") + .field("plan", &self.plan) + .field("report_kind", &self.report_kind) + .field("has_report_context", &self.report_context.is_some()) + .field( + "report_context_bytes", + &json_value_len(&self.report_context), + ) + .finish() + } +} + +fn json_value_len(value: &Option) -> Option { + value + .as_ref() + .and_then(|value| serde_json::to_vec(value).ok().map(|bytes| bytes.len())) +} + pub fn augment_sync_report_context( report_context: Option, provider_request_headers: &BTreeMap, diff --git a/crates/aether-ai/serving/src/lib.rs b/crates/aether-ai/serving/src/lib.rs index 502bfef9a..99dbac00d 100644 --- a/crates/aether-ai/serving/src/lib.rs +++ b/crates/aether-ai/serving/src/lib.rs @@ -54,7 +54,7 @@ pub use aether_pool_core::{ }; pub use attempt_loop::{ run_ai_attempt_loop, AiAttemptExecutionOutcome, AiAttemptLoopOutcome, AiAttemptLoopPort, - AiAttemptRetryScope, AiExecutionAttempt, + AiAttemptRetryScope, AiExecutionAttempt, STICKY_KEY_ATTEMPTS_REPORT_FIELD, }; pub use attempt_plan::{ build_ai_execution_decision_from_plan, build_ai_execution_plan_from_decision, diff --git a/crates/aether-contracts/Cargo.toml b/crates/aether-contracts/Cargo.toml index f5dead351..e5bab75c9 100644 --- a/crates/aether-contracts/Cargo.toml +++ b/crates/aether-contracts/Cargo.toml @@ -16,3 +16,4 @@ serde.workspace = true serde_json.workspace = true sha2.workspace = true thiserror.workspace = true +url.workspace = true diff --git a/crates/aether-contracts/src/internal_gateway.rs b/crates/aether-contracts/src/internal_gateway.rs new file mode 100644 index 000000000..10d6f1ef3 --- /dev/null +++ b/crates/aether-contracts/src/internal_gateway.rs @@ -0,0 +1,141 @@ +use base64::Engine as _; +use hmac::{Hmac, Mac}; +use sha2::{Digest as _, Sha256}; + +pub const INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER: &str = "x-aether-internal-gateway-timestamp"; +pub const INTERNAL_GATEWAY_AUTH_NONCE_HEADER: &str = "x-aether-internal-gateway-nonce"; +pub const INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER: &str = "x-aether-internal-gateway-signature"; + +const INTERNAL_GATEWAY_AUTH_CONTEXT: &[u8] = b"aether-internal-gateway-auth-v1"; + +type HmacSha256 = Hmac; + +pub fn sign_internal_gateway_request( + secret: &[u8], + method: &str, + path_and_query: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], +) -> String { + let mut mac = HmacSha256::new_from_slice(secret).expect("HMAC accepts keys of any size"); + update_auth_mac( + &mut mac, + method, + path_and_query, + timestamp_unix_secs, + nonce, + body, + ); + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()) +} + +pub fn verify_internal_gateway_request_signature( + secret: &[u8], + method: &str, + path_and_query: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], + signature: &str, +) -> bool { + let signature = signature.trim(); + if signature.len() > 43 { + return false; + } + let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else { + return false; + }; + let Ok(mut mac) = HmacSha256::new_from_slice(secret) else { + return false; + }; + update_auth_mac( + &mut mac, + method, + path_and_query, + timestamp_unix_secs, + nonce, + body, + ); + mac.verify_slice(&signature).is_ok() +} + +fn update_auth_mac( + mac: &mut HmacSha256, + method: &str, + path_and_query: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], +) { + mac.update(INTERNAL_GATEWAY_AUTH_CONTEXT); + update_auth_field(mac, method.as_bytes()); + update_auth_field(mac, path_and_query.as_bytes()); + mac.update(×tamp_unix_secs.to_be_bytes()); + update_auth_field(mac, nonce.as_bytes()); + mac.update(&(body.len() as u64).to_be_bytes()); + mac.update(&Sha256::digest(body)); +} + +fn update_auth_field(mac: &mut HmacSha256, value: &[u8]) { + mac.update(&(value.len() as u64).to_be_bytes()); + mac.update(value); +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn signature_binds_every_security_relevant_field() { + let secret = b"internal-gateway-test-secret-32-bytes-minimum"; + let body = br#"{"path":"/v1/models"}"#; + let timestamp = 1_800_000_000; + let nonce = "nonce-value-0000000000000001"; + let path = "/api/internal/gateway/resolve?mode=full"; + let signature = sign_internal_gateway_request(secret, "POST", path, timestamp, nonce, body); + + assert!(verify_internal_gateway_request_signature( + secret, "POST", path, timestamp, nonce, body, &signature, + )); + assert!(!verify_internal_gateway_request_signature( + secret, "GET", path, timestamp, nonce, body, &signature, + )); + assert!(!verify_internal_gateway_request_signature( + secret, + "POST", + "/api/internal/gateway/resolve?mode=brief", + timestamp, + nonce, + body, + &signature, + )); + assert!(!verify_internal_gateway_request_signature( + secret, + "POST", + path, + timestamp + 1, + nonce, + body, + &signature, + )); + assert!(!verify_internal_gateway_request_signature( + secret, + "POST", + path, + timestamp, + "different-nonce-0000000000001", + body, + &signature, + )); + assert!(!verify_internal_gateway_request_signature( + secret, + "POST", + path, + timestamp, + nonce, + br#"{"path":"/v1/providers"}"#, + &signature, + )); + } +} diff --git a/crates/aether-contracts/src/lib.rs b/crates/aether-contracts/src/lib.rs index f9a9233b1..a64f0a49f 100644 --- a/crates/aether-contracts/src/lib.rs +++ b/crates/aether-contracts/src/lib.rs @@ -1,5 +1,6 @@ mod error; mod frame; +pub mod internal_gateway; mod plan; mod result; pub mod tunnel; @@ -9,13 +10,14 @@ mod usage; pub use error::{ExecutionError, ExecutionErrorKind, ExecutionPhase}; pub use frame::{StreamFrame, StreamFramePayload, StreamFrameType}; pub use plan::{ - ExecutionPlan, ExecutionResponseBodyMode, ExecutionTimeouts, ProxySnapshot, RequestBody, - ResolvedTransportProfile, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, + redact_url_for_debug, ExecutionPlan, ExecutionResponseBodyMode, ExecutionTimeouts, + ProxySnapshot, RequestBody, ResolvedTransportProfile, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER, MAX_EXECUTION_REQUEST_TIMEOUT_MS, MAX_EXECUTION_REQUEST_TIMEOUT_SECS, MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_MS, - MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_SECS, TRANSPORT_BACKEND_BROWSER_WREQ, - TRANSPORT_BACKEND_HYPER_RUSTLS, TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, + MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_SECS, PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY, + TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_HYPER_RUSTLS, + TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY, TRANSPORT_POOL_SCOPE_KEY, }; diff --git a/crates/aether-contracts/src/plan.rs b/crates/aether-contracts/src/plan.rs index 443256521..be2bb745e 100644 --- a/crates/aether-contracts/src/plan.rs +++ b/crates/aether-contracts/src/plan.rs @@ -1,12 +1,11 @@ use std::collections::BTreeMap; +use std::fmt; use serde::{Deserialize, Serialize}; use serde_json::Value; pub const EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER: &str = "x-aether-execution-follow-redirects"; pub const EXECUTION_REQUEST_HTTP1_ONLY_HEADER: &str = "x-aether-execution-http1-only"; -pub const EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER: &str = - "x-aether-execution-accept-invalid-certs"; pub const EXECUTION_RESPONSE_BODY_MODE_HEADER: &str = "x-aether-execution-response-body-mode"; pub const MAX_EXECUTION_REQUEST_TIMEOUT_SECS: u64 = 1_200; pub const MAX_EXECUTION_REQUEST_TIMEOUT_MS: u64 = MAX_EXECUTION_REQUEST_TIMEOUT_SECS * 1_000; @@ -57,7 +56,7 @@ pub struct ExecutionTimeouts { pub total_ms: Option, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[derive(Clone, Serialize, Deserialize, PartialEq)] pub struct RequestBody { #[serde(default, skip_serializing_if = "Option::is_none")] pub json_body: Option, @@ -67,6 +66,27 @@ pub struct RequestBody { pub body_ref: Option, } +impl fmt::Debug for RequestBody { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("RequestBody") + .field("has_json_body", &self.json_body.is_some()) + .field( + "json_body_bytes", + &self + .json_body + .as_ref() + .and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())), + ) + .field( + "body_bytes_b64_len", + &self.body_bytes_b64.as_ref().map(String::len), + ) + .field("body_ref_len", &self.body_ref.as_ref().map(String::len)) + .finish() + } +} + impl RequestBody { pub fn from_json(json_body: Value) -> Self { Self { @@ -77,7 +97,7 @@ impl RequestBody { } } -#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)] +#[derive(Clone, Default, Serialize, Deserialize, PartialEq)] pub struct ProxySnapshot { #[serde(default, skip_serializing_if = "Option::is_none")] pub enabled: Option, @@ -93,6 +113,24 @@ pub struct ProxySnapshot { pub extra: Option, } +impl fmt::Debug for ProxySnapshot { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ProxySnapshot") + .field("enabled", &self.enabled) + .field("mode", &self.mode) + .field("node_id", &self.node_id) + .field("label", &self.label) + .field("url", &self.url.as_deref().map(redact_url_for_debug)) + .field("has_extra", &self.extra.is_some()) + .finish() + } +} + +/// Reserved internal metadata key used to fence proxy-node mutations against +/// a node incarnation recreated under the same stable id. +pub const PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY: &str = "proxy_node_tunnel_generation"; + pub const TRANSPORT_BACKEND_REQWEST_RUSTLS: &str = "reqwest_rustls"; pub const TRANSPORT_BACKEND_HYPER_RUSTLS: &str = "hyper_rustls"; pub const TRANSPORT_BACKEND_BROWSER_WREQ: &str = "browser_wreq"; @@ -101,7 +139,7 @@ pub const TRANSPORT_HTTP_MODE_HTTP1_ONLY: &str = "http1_only"; pub const TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE: &str = "h2c_prior_knowledge"; pub const TRANSPORT_POOL_SCOPE_KEY: &str = "key"; -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[derive(Clone, Serialize, Deserialize, PartialEq)] #[serde(default)] pub struct ResolvedTransportProfile { pub profile_id: String, @@ -127,7 +165,21 @@ impl Default for ResolvedTransportProfile { } } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +impl fmt::Debug for ResolvedTransportProfile { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ResolvedTransportProfile") + .field("profile_id", &self.profile_id) + .field("backend", &self.backend) + .field("http_mode", &self.http_mode) + .field("pool_scope", &self.pool_scope) + .field("has_header_fingerprint", &self.header_fingerprint.is_some()) + .field("has_extra", &self.extra.is_some()) + .finish() + } +} + +#[derive(Clone, Serialize, Deserialize, PartialEq)] pub struct ExecutionPlan { pub request_id: String, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -162,6 +214,63 @@ pub struct ExecutionPlan { pub timeouts: Option, } +impl fmt::Debug for ExecutionPlan { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ExecutionPlan") + .field("request_id", &self.request_id) + .field("candidate_id", &self.candidate_id) + .field("provider_name", &self.provider_name) + .field("provider_id", &self.provider_id) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field("method", &self.method) + .field("url", &redact_url_for_debug(&self.url)) + .field("header_names", &self.headers.keys().collect::>()) + .field("content_type", &self.content_type) + .field("content_encoding", &self.content_encoding) + .field("body", &self.body) + .field("stream", &self.stream) + .field("client_api_format", &self.client_api_format) + .field("provider_api_format", &self.provider_api_format) + .field("model_name", &self.model_name) + .field("proxy", &self.proxy) + .field("transport_profile", &self.transport_profile) + .field("timeouts", &self.timeouts) + .finish() + } +} + +/// Return a bounded URL representation suitable for diagnostics. +/// +/// URL userinfo, query parameters, and fragments are never emitted because +/// providers commonly put API keys or OAuth tokens in those locations. An +/// unparsable URL is represented only by its length rather than echoing input. +pub fn redact_url_for_debug(raw: &str) -> String { + const MAX_DEBUG_URL_CHARS: usize = 512; + let raw = raw.trim(); + if raw.is_empty() { + return String::new(); + } + let Ok(mut url) = url::Url::parse(raw) else { + return format!("[invalid-url len={}]", raw.len()); + }; + let _ = url.set_username(""); + let _ = url.set_password(None); + url.set_query(None); + url.set_fragment(None); + let rendered = url.to_string(); + if rendered.chars().count() <= MAX_DEBUG_URL_CHARS { + rendered + } else { + let prefix = rendered + .chars() + .take(MAX_DEBUG_URL_CHARS - 3) + .collect::(); + format!("{prefix}...") + } +} + #[cfg(test)] mod tests { use super::*; @@ -268,4 +377,75 @@ mod tests { Some(300_000) ); } + + #[test] + fn debug_redacts_plan_credentials_and_payloads() { + let plan = ExecutionPlan { + request_id: "req-debug".into(), + candidate_id: Some("candidate-debug".into()), + provider_name: Some("provider".into()), + provider_id: "provider-id".into(), + endpoint_id: "endpoint-id".into(), + key_id: "key-id".into(), + method: "POST".into(), + url: "https://proxy-user:proxy-password@example.test/v1?api_key=url-secret#fragment-secret".into(), + headers: BTreeMap::from([( + "authorization".into(), + "Bearer header-secret".into(), + )]), + content_type: Some("application/json".into()), + content_encoding: None, + body: RequestBody::from_json(serde_json::json!({ + "access_token": "body-secret", + "prompt": "hello" + })), + stream: false, + client_api_format: "openai:chat".into(), + provider_api_format: "openai:chat".into(), + model_name: Some("model".into()), + proxy: Some(ProxySnapshot { + enabled: Some(true), + mode: Some("http".into()), + node_id: Some("node".into()), + label: None, + url: Some("http://proxy-user:proxy-password@proxy.test?token=proxy-secret".into()), + extra: Some(serde_json::json!({"credential": "extra-secret"})), + }), + transport_profile: Some(ResolvedTransportProfile { + profile_id: "profile".into(), + backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(), + http_mode: TRANSPORT_HTTP_MODE_AUTO.into(), + pool_scope: TRANSPORT_POOL_SCOPE_KEY.into(), + header_fingerprint: Some(serde_json::json!({"authorization": "fingerprint-secret"})), + extra: Some(serde_json::json!({"secret": "profile-secret"})), + }), + timeouts: None, + }; + + let debug = format!("{plan:?}"); + for secret in [ + "proxy-user", + "proxy-password", + "url-secret", + "fragment-secret", + "header-secret", + "body-secret", + "proxy-secret", + "extra-secret", + "fingerprint-secret", + "profile-secret", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}: {debug}"); + } + assert!(debug.contains("header_names")); + assert!(debug.contains("has_json_body")); + assert!(debug.contains("profile")); + } + + #[test] + fn debug_url_redaction_fails_closed_for_invalid_urls() { + let redacted = redact_url_for_debug("not a url?token=secret"); + assert!(!redacted.contains("secret")); + assert!(redacted.starts_with("[invalid-url")); + } } diff --git a/crates/aether-contracts/src/result.rs b/crates/aether-contracts/src/result.rs index 3d4802dd9..b2312ceb8 100644 --- a/crates/aether-contracts/src/result.rs +++ b/crates/aether-contracts/src/result.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -22,7 +23,7 @@ pub struct ExecutionTelemetry { pub upstream_bytes: Option, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[derive(Clone, Serialize, Deserialize, PartialEq)] pub struct ResponseBody { #[serde(default, skip_serializing_if = "Option::is_none")] pub json_body: Option, @@ -30,7 +31,27 @@ pub struct ResponseBody { pub body_bytes_b64: Option, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +impl fmt::Debug for ResponseBody { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ResponseBody") + .field("has_json_body", &self.json_body.is_some()) + .field( + "json_body_bytes", + &self + .json_body + .as_ref() + .and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())), + ) + .field( + "body_bytes_b64_len", + &self.body_bytes_b64.as_ref().map(String::len), + ) + .finish() + } +} + +#[derive(Clone, Serialize, Deserialize, PartialEq)] pub struct ExecutionResult { pub request_id: String, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -47,3 +68,73 @@ pub struct ExecutionResult { #[serde(default, skip_serializing_if = "Option::is_none")] pub error: Option, } + +impl fmt::Debug for ExecutionResult { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("ExecutionResult"); + debug + .field("request_id", &self.request_id) + .field("candidate_id", &self.candidate_id) + .field("status_code", &self.status_code) + .field("header_names", &self.headers.keys().collect::>()) + .field("response_observation", &self.response_observation) + .field("body", &self.body) + .field("telemetry", &self.telemetry) + .field("has_error", &self.error.is_some()); + if let Some(error) = self.error.as_ref() { + // ExecutionError::message can contain an upstream response or URL. + debug + .field("error_kind", &error.kind) + .field("error_phase", &error.phase) + .field("error_upstream_status", &error.upstream_status) + .field("error_retryable", &error.retryable) + .field("error_failover_recommended", &error.failover_recommended); + } + debug.finish() + } +} + +#[cfg(test)] +mod tests { + use super::{ExecutionResult, ResponseBody}; + use crate::{ExecutionError, ExecutionErrorKind, ExecutionPhase}; + use std::collections::BTreeMap; + + #[test] + fn debug_does_not_render_response_headers_body_or_error_message() { + let result = ExecutionResult { + request_id: "request-1".into(), + candidate_id: None, + status_code: 401, + headers: BTreeMap::from([( + "set-cookie".into(), + "session=response-header-secret".into(), + )]), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(serde_json::json!({"access_token": "response-body-secret"})), + body_bytes_b64: None, + }), + telemetry: None, + error: Some(ExecutionError { + kind: ExecutionErrorKind::Upstream4xx, + phase: ExecutionPhase::Finalize, + message: "upstream detail error-secret".into(), + upstream_status: Some(401), + retryable: false, + failover_recommended: false, + }), + }; + + let debug = format!("{result:?}"); + for secret in [ + "response-header-secret", + "response-body-secret", + "error-secret", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}: {debug}"); + } + assert!(debug.contains("header_names")); + assert!(debug.contains("has_error")); + } +} diff --git a/crates/aether-contracts/src/tunnel.rs b/crates/aether-contracts/src/tunnel.rs index 97528adbc..115454b11 100644 --- a/crates/aether-contracts/src/tunnel.rs +++ b/crates/aether-contracts/src/tunnel.rs @@ -1,18 +1,205 @@ +use std::fmt; use std::io::Read; +use base64::Engine as _; use bytes::{Buf, BufMut, Bytes, BytesMut}; use flate2::read::GzDecoder; use flate2::write::GzEncoder; use flate2::Compression; +use hmac::{Hmac, Mac}; +use sha2::{Digest as _, Sha256}; pub const HEADER_SIZE: usize = 10; pub const TUNNEL_RELAY_FORWARDED_BY_HEADER: &str = "x-aether-tunnel-forwarded-by"; pub const TUNNEL_RELAY_OWNER_INSTANCE_HEADER: &str = "x-aether-tunnel-owner-instance-id"; +pub const TUNNEL_RELAY_AUTH_SENDER_HEADER: &str = "x-aether-tunnel-relay-sender"; +pub const TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER: &str = "x-aether-tunnel-relay-timestamp"; +pub const TUNNEL_RELAY_AUTH_NONCE_HEADER: &str = "x-aether-tunnel-relay-nonce"; +pub const TUNNEL_RELAY_AUTH_PAYLOAD_HEADER: &str = "x-aether-tunnel-relay-payload"; +pub const TUNNEL_RELAY_AUTH_SIGNATURE_HEADER: &str = "x-aether-tunnel-relay-signature"; pub const TUNNEL_PROTOCOL_VERSION_HEADER: &str = "x-aether-tunnel-protocol-version"; pub const TUNNEL_NODE_NAME_B64_HEADER: &str = "x-aether-tunnel-node-name-b64"; pub const CURRENT_TUNNEL_PROTOCOL_VERSION: u8 = 3; pub const CURRENT_TUNNEL_PROTOCOL_VERSION_STR: &str = "3"; pub const MAX_TUNNEL_RELAY_META_LEN: usize = 256 * 1024; +/// Keep decoded tunnel frames within the same size envelope enforced by the +/// WebSocket transports. This also bounds gzip expansion for untrusted peers. +pub const MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES: usize = 64 * 1024 * 1024; + +const TUNNEL_RELAY_AUTH_CONTEXT: &[u8] = b"aether-tunnel-relay-auth-v2"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct TunnelRelayPayloadDigest { + metadata_sha256: [u8; 32], + body_len: u64, + body_sha256: [u8; 32], +} + +impl TunnelRelayPayloadDigest { + pub fn body_len(self) -> u64 { + self.body_len + } + + pub fn encode_header_value(self) -> String { + let mut encoded = [0_u8; 72]; + encoded[..32].copy_from_slice(&self.metadata_sha256); + encoded[32..40].copy_from_slice(&self.body_len.to_be_bytes()); + encoded[40..].copy_from_slice(&self.body_sha256); + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(encoded) + } + + pub fn decode_header_value(value: &str) -> Option { + let value = value.trim(); + if value.len() > 96 { + return None; + } + let decoded = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(value) + .ok()?; + let encoded: [u8; 72] = decoded.try_into().ok()?; + Some(Self { + metadata_sha256: encoded[..32].try_into().ok()?, + body_len: u64::from_be_bytes(encoded[32..40].try_into().ok()?), + body_sha256: encoded[40..].try_into().ok()?, + }) + } + + pub fn matches_metadata(self, metadata_envelope: &[u8]) -> bool { + self.metadata_sha256 == <[u8; 32]>::from(Sha256::digest(metadata_envelope)) + } + + pub fn matches_body(self, body: &[u8]) -> bool { + self.matches_body_hash(body.len() as u64, Sha256::digest(body).into()) + } + + pub fn matches_body_hash(self, body_len: u64, body_sha256: [u8; 32]) -> bool { + self.body_len == body_len && self.body_sha256 == body_sha256 + } +} + +pub fn tunnel_relay_payload_digest( + metadata_envelope: &[u8], + body: &[u8], +) -> TunnelRelayPayloadDigest { + TunnelRelayPayloadDigest { + metadata_sha256: Sha256::digest(metadata_envelope).into(), + body_len: body.len() as u64, + body_sha256: Sha256::digest(body).into(), + } +} + +pub fn tunnel_relay_payload_digest_from_hashes( + metadata_envelope: &[u8], + body_len: u64, + body_sha256: [u8; 32], +) -> TunnelRelayPayloadDigest { + TunnelRelayPayloadDigest { + metadata_sha256: Sha256::digest(metadata_envelope).into(), + body_len, + body_sha256, + } +} + +// Keep the explicit protocol arguments in this public API: their order is +// reflected in the relay authentication MAC and changing it would break +// interoperability with deployed tunnel peers. +#[allow(clippy::too_many_arguments)] +pub fn sign_tunnel_relay_request( + secret: &[u8], + sender_instance_id: &str, + owner_instance_id: &str, + node_id: &str, + forwarded_by: &str, + rollout_probe: bool, + timestamp_unix_secs: u64, + nonce: &str, + payload_digest: &TunnelRelayPayloadDigest, +) -> String { + let mut mac = Hmac::::new_from_slice(secret).expect("HMAC accepts keys of any size"); + update_tunnel_relay_auth_mac( + &mut mac, + sender_instance_id, + owner_instance_id, + node_id, + forwarded_by, + rollout_probe, + timestamp_unix_secs, + nonce, + payload_digest, + ); + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()) +} + +// The verifier mirrors `sign_tunnel_relay_request` field-for-field so the +// authenticated transcript remains stable across crate versions. +#[allow(clippy::too_many_arguments)] +pub fn verify_tunnel_relay_request_signature( + secret: &[u8], + sender_instance_id: &str, + owner_instance_id: &str, + node_id: &str, + forwarded_by: &str, + rollout_probe: bool, + timestamp_unix_secs: u64, + nonce: &str, + payload_digest: &TunnelRelayPayloadDigest, + signature: &str, +) -> bool { + let signature = signature.trim(); + if signature.len() > 43 { + return false; + } + let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else { + return false; + }; + let Ok(mut mac) = Hmac::::new_from_slice(secret) else { + return false; + }; + update_tunnel_relay_auth_mac( + &mut mac, + sender_instance_id, + owner_instance_id, + node_id, + forwarded_by, + rollout_probe, + timestamp_unix_secs, + nonce, + payload_digest, + ); + mac.verify_slice(&signature).is_ok() +} + +// This helper deliberately accepts the wire fields separately to make the +// authenticated-field order visible next to the MAC construction. +#[allow(clippy::too_many_arguments)] +fn update_tunnel_relay_auth_mac( + mac: &mut Hmac, + sender_instance_id: &str, + owner_instance_id: &str, + node_id: &str, + forwarded_by: &str, + rollout_probe: bool, + timestamp_unix_secs: u64, + nonce: &str, + payload_digest: &TunnelRelayPayloadDigest, +) { + mac.update(TUNNEL_RELAY_AUTH_CONTEXT); + update_tunnel_relay_auth_field(mac, sender_instance_id.as_bytes()); + update_tunnel_relay_auth_field(mac, owner_instance_id.as_bytes()); + update_tunnel_relay_auth_field(mac, node_id.as_bytes()); + update_tunnel_relay_auth_field(mac, forwarded_by.as_bytes()); + mac.update(&[u8::from(rollout_probe)]); + mac.update(×tamp_unix_secs.to_be_bytes()); + update_tunnel_relay_auth_field(mac, nonce.as_bytes()); + mac.update(&payload_digest.metadata_sha256); + mac.update(&payload_digest.body_len.to_be_bytes()); + mac.update(&payload_digest.body_sha256); +} + +fn update_tunnel_relay_auth_field(mac: &mut Hmac, value: &[u8]) { + mac.update(&(value.len() as u64).to_be_bytes()); + mac.update(value); +} pub mod flags { pub const END_STREAM: u8 = 0x01; @@ -111,7 +298,7 @@ impl FrameHeader { } } -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct Frame { pub stream_id: u32, pub msg_type: MsgType, @@ -119,6 +306,18 @@ pub struct Frame { pub payload: Bytes, } +impl fmt::Debug for Frame { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("Frame") + .field("stream_id", &self.stream_id) + .field("msg_type", &self.msg_type) + .field("flags", &self.flags) + .field("payload_len", &self.payload.len()) + .finish() + } +} + impl Frame { pub fn new(stream_id: u32, msg_type: MsgType, flags: u8, payload: impl Into) -> Self { Self { @@ -170,6 +369,12 @@ impl Frame { actual: HEADER_SIZE + data.remaining(), }); } + if data.remaining() > payload_len { + return Err(ProtocolError::Trailing { + expected: HEADER_SIZE + payload_len, + actual: HEADER_SIZE + data.remaining(), + }); + } let msg_type = MsgType::from_u8(msg_type_raw).ok_or(ProtocolError::UnknownMsgType(msg_type_raw))?; @@ -190,11 +395,13 @@ pub enum ProtocolError { TooShort { expected: usize, actual: usize }, #[error("frame incomplete: expected {expected} bytes, got {actual}")] Incomplete { expected: usize, actual: usize }, + #[error("frame has trailing bytes: expected {expected} bytes, got {actual}")] + Trailing { expected: usize, actual: usize }, #[error("unknown message type: 0x{0:02x}")] UnknownMsgType(u8), } -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[derive(Clone, serde::Serialize, serde::Deserialize)] pub struct RequestMeta { #[serde(default, skip_serializing_if = "Option::is_none")] pub provider_id: Option, @@ -221,6 +428,30 @@ pub struct RequestMeta { pub transport_profile: Option, } +impl fmt::Debug for RequestMeta { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("RequestMeta") + .field("provider_id", &self.provider_id) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field("method", &self.method) + .field("url", &crate::redact_url_for_debug(&self.url)) + .field("header_names", &self.headers.keys().collect::>()) + .field("stream", &self.stream) + .field("request_timeout_ms", &self.request_timeout_ms) + .field( + "stream_first_byte_timeout_ms", + &self.stream_first_byte_timeout_ms, + ) + .field("timeout", &self.timeout) + .field("follow_redirects", &self.follow_redirects) + .field("http1_only", &self.http1_only) + .field("transport_profile", &self.transport_profile) + .finish() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct ResolvedTunnelRequestTimeouts { pub first_byte_ms: u64, @@ -316,12 +547,29 @@ where } } -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[derive(Clone, serde::Serialize, serde::Deserialize)] pub struct ResponseMeta { pub status: u16, pub headers: Vec<(String, String)>, } +impl fmt::Debug for ResponseMeta { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ResponseMeta") + .field("status", &self.status) + .field( + "header_names", + &self + .headers + .iter() + .map(|(name, _)| name) + .collect::>(), + ) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct HelloPayload { pub protocol_version: u8, @@ -453,30 +701,47 @@ fn encode_json_control(msg_type: u8, payload: &T) -> Vec(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> { let payload_len = header.payload_len as usize; let end = HEADER_SIZE.checked_add(payload_len)?; - if data.len() < end { + if data.len() != end { return None; } Some(&data[HEADER_SIZE..end]) } pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result, String> { + decode_payload_with_limit(data, header, MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES) +} + +pub fn decode_payload_with_limit( + data: &[u8], + header: &FrameHeader, + max_decoded_bytes: usize, +) -> Result, String> { let payload = frame_payload_by_header(data, header) .ok_or_else(|| "incomplete frame payload".to_string())?; if header.flags & FLAG_GZIP_COMPRESSED != 0 { - let mut decoder = GzDecoder::new(payload); - let mut decoded = Vec::new(); - decoder - .read_to_end(&mut decoded) - .map_err(|err| format!("failed to decompress payload: {err}"))?; - Ok(decoded) + decompress_gzip_with_limit(payload, max_decoded_bytes) + .map_err(|err| format!("failed to decompress payload: {err}")) + } else if payload.len() > max_decoded_bytes { + Err(format!( + "decoded tunnel payload exceeds {max_decoded_bytes} bytes" + )) } else { Ok(payload.to_vec()) } } pub fn decompress_if_gzip(frame: &Frame) -> Result { + decompress_if_gzip_with_limit(frame, MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES) +} + +pub fn decompress_if_gzip_with_limit( + frame: &Frame, + max_decoded_bytes: usize, +) -> Result { if frame.is_gzip() { - decompress_gzip(&frame.payload) + decompress_gzip_with_limit(&frame.payload, max_decoded_bytes).map(Bytes::from) + } else if frame.payload.len() > max_decoded_bytes { + Err(decoded_payload_too_large(max_decoded_bytes)) } else { Ok(frame.payload.clone()) } @@ -499,11 +764,43 @@ pub fn raw_payload(data: Bytes) -> (Bytes, u8) { const COMPRESS_MIN_SIZE: usize = 512; -fn decompress_gzip(data: &[u8]) -> Result { +fn decompress_gzip_with_limit( + data: &[u8], + max_decoded_bytes: usize, +) -> Result, std::io::Error> { let mut decoder = GzDecoder::new(data); - let mut buf = Vec::new(); - decoder.read_to_end(&mut buf)?; - Ok(Bytes::from(buf)) + let mut decoded = Vec::with_capacity(max_decoded_bytes.min(8 * 1024)); + let mut chunk = [0_u8; 8 * 1024]; + loop { + let remaining = max_decoded_bytes.saturating_sub(decoded.len()); + let read_len = remaining.saturating_add(1).min(chunk.len()); + let read = match decoder.read(&mut chunk[..read_len]) { + Ok(read) => read, + Err(error) if error.kind() == std::io::ErrorKind::Interrupted => continue, + Err(error) => return Err(error), + }; + if read == 0 { + return Ok(decoded); + } + if read > remaining { + return Err(decoded_payload_too_large(max_decoded_bytes)); + } + if decoded.capacity().saturating_sub(decoded.len()) < read { + decoded.try_reserve_exact(read).map_err(|error| { + std::io::Error::other(format!( + "failed to allocate decoded tunnel payload: {error}" + )) + })?; + } + decoded.extend_from_slice(&chunk[..read]); + } +} + +fn decoded_payload_too_large(max_decoded_bytes: usize) -> std::io::Error { + std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("decoded tunnel payload exceeds {max_decoded_bytes} bytes"), + ) } fn compress_gzip(data: &[u8]) -> Result { @@ -518,12 +815,15 @@ fn compress_gzip(data: &[u8]) -> Result { #[cfg(test)] mod tests { use super::{ - compress_payload, decode_payload, encode_frame, encode_goaway_v3, encode_ping, - encode_reset_stream, encode_window_update, raw_payload, resolve_tunnel_request_timeouts, - try_decode_tunnel_relay_request_meta, Frame, FrameHeader, GoAwayPayload, MsgType, - RequestMeta, ResetStreamPayload, WindowUpdatePayload, CURRENT_TUNNEL_PROTOCOL_VERSION, - CURRENT_TUNNEL_PROTOCOL_VERSION_STR, FLAG_GZIP_COMPRESSED, MAX_TUNNEL_RELAY_META_LEN, - REQUEST_HEADERS, TUNNEL_PROTOCOL_VERSION_HEADER, + compress_payload, decode_payload, decode_payload_with_limit, decompress_if_gzip_with_limit, + encode_frame, encode_goaway_v3, encode_ping, encode_reset_stream, encode_window_update, + frame_payload_by_header, raw_payload, resolve_tunnel_request_timeouts, + sign_tunnel_relay_request, try_decode_tunnel_relay_request_meta, + tunnel_relay_payload_digest, verify_tunnel_relay_request_signature, Frame, FrameHeader, + GoAwayPayload, MsgType, ProtocolError, RequestMeta, ResetStreamPayload, ResponseMeta, + WindowUpdatePayload, CURRENT_TUNNEL_PROTOCOL_VERSION, CURRENT_TUNNEL_PROTOCOL_VERSION_STR, + FLAG_GZIP_COMPRESSED, HEADER_SIZE, MAX_TUNNEL_RELAY_META_LEN, REQUEST_HEADERS, + RESPONSE_BODY, TUNNEL_PROTOCOL_VERSION_HEADER, }; use bytes::Bytes; @@ -628,6 +928,71 @@ mod tests { assert!(try_decode_tunnel_relay_request_meta(&oversized).is_err()); } + #[test] + fn tunnel_relay_signature_binds_routing_metadata_and_body() { + let digest = tunnel_relay_payload_digest(b"metadata", b"request-body"); + let signature = sign_tunnel_relay_request( + b"shared-secret", + "gateway-a", + "gateway-b", + "node-1", + "gateway-a", + false, + 123, + "nonce-1", + &digest, + ); + + assert!(verify_tunnel_relay_request_signature( + b"shared-secret", + "gateway-a", + "gateway-b", + "node-1", + "gateway-a", + false, + 123, + "nonce-1", + &digest, + &signature, + )); + assert!(!verify_tunnel_relay_request_signature( + b"shared-secret", + "gateway-a", + "gateway-b", + "node-2", + "gateway-a", + false, + 123, + "nonce-1", + &digest, + &signature, + )); + assert!(!verify_tunnel_relay_request_signature( + b"shared-secret", + "gateway-a", + "gateway-b", + "node-1", + "gateway-a", + false, + 123, + "nonce-1", + &tunnel_relay_payload_digest(b"tampered", b"request-body"), + &signature, + )); + assert!(!verify_tunnel_relay_request_signature( + b"shared-secret", + "gateway-a", + "gateway-b", + "node-1", + "gateway-a", + false, + 123, + "nonce-1", + &tunnel_relay_payload_digest(b"metadata", b"tampered-body"), + &signature, + )); + } + #[test] fn request_meta_accepts_integer_timeout() { let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15}"#; @@ -652,6 +1017,34 @@ mod tests { assert_eq!(decoded.payload, Bytes::from_static(b"hello")); } + #[test] + fn frame_decode_rejects_trailing_bytes() { + let frame = Frame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"hello")); + let mut encoded = frame.encode().to_vec(); + encoded.extend_from_slice(b"hidden"); + + let error = Frame::decode(Bytes::from(encoded)).expect_err("trailing bytes rejected"); + assert!(matches!( + error, + ProtocolError::Trailing { expected, actual } + if expected == HEADER_SIZE + 5 && actual == HEADER_SIZE + 11 + )); + } + + #[test] + fn frame_payload_lookup_rejects_trailing_bytes() { + let encoded = encode_frame(7, RESPONSE_BODY, 0, b"hello"); + let header = FrameHeader::parse(&encoded).expect("frame header should parse"); + let mut with_trailing = encoded.clone(); + with_trailing.push(0); + + assert!(frame_payload_by_header(&with_trailing, &header).is_none()); + assert_eq!( + frame_payload_by_header(&encoded, &header), + Some(&encoded[HEADER_SIZE..]) + ); + } + #[test] fn frame_header_parses_raw_ping_frame() { let encoded = encode_ping(); @@ -681,6 +1074,56 @@ mod tests { assert_eq!(decoded, control_payload.to_vec()); } + #[test] + fn compressed_tunnel_payload_is_rejected_before_exceeding_decode_limit() { + const LIMIT: usize = 1024; + + let at_limit = Bytes::from(vec![b'a'; LIMIT]); + let (at_limit_compressed, at_limit_flags) = compress_payload(at_limit.clone()); + assert_ne!(at_limit_flags & FLAG_GZIP_COMPRESSED, 0); + let at_limit_frame = Frame::new( + 1, + MsgType::RequestHeaders, + at_limit_flags, + at_limit_compressed, + ); + assert_eq!( + decompress_if_gzip_with_limit(&at_limit_frame, LIMIT) + .expect("payload at the limit should decode"), + at_limit + ); + + let over_limit = Bytes::from(vec![b'a'; LIMIT + 1]); + let (compressed, flags) = compress_payload(over_limit); + assert_ne!(flags & FLAG_GZIP_COMPRESSED, 0); + let frame = Frame::new(1, MsgType::RequestHeaders, flags, compressed.clone()); + let error = decompress_if_gzip_with_limit(&frame, LIMIT) + .expect_err("gzip expansion must be bounded"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(error.to_string().contains("exceeds 1024 bytes")); + + let encoded = encode_frame(1, REQUEST_HEADERS, flags, &compressed); + let header = FrameHeader::parse(&encoded).expect("frame should parse"); + let error = decode_payload_with_limit(&encoded, &header, LIMIT) + .expect_err("compatibility decoder must apply the same bound"); + assert!(error.contains("exceeds 1024 bytes")); + + let raw_at_limit = Frame::new(1, MsgType::RequestBody, 0, Bytes::from(vec![b'x'; LIMIT])); + assert!(decompress_if_gzip_with_limit(&raw_at_limit, LIMIT).is_ok()); + let raw_over_limit = Frame::new( + 1, + MsgType::RequestBody, + 0, + Bytes::from(vec![b'x'; LIMIT + 1]), + ); + assert_eq!( + decompress_if_gzip_with_limit(&raw_over_limit, LIMIT) + .expect_err("raw payloads must use the same bound") + .kind(), + std::io::ErrorKind::InvalidData + ); + } + #[test] fn tunnel_protocol_version_header_defaults_to_v2() { assert_eq!( @@ -719,4 +1162,50 @@ mod tests { assert_eq!(goaway_payload.drain_deadline_ms, 30_000); assert_eq!(goaway_payload.reason, "rolling restart"); } + + #[test] + fn debug_does_not_render_tunnel_credentials_or_payload_bytes() { + let meta = RequestMeta { + provider_id: Some("provider".into()), + endpoint_id: Some("endpoint".into()), + key_id: Some("key".into()), + method: "POST".into(), + url: "https://user:password@example.test/path?token=url-secret".into(), + headers: std::collections::HashMap::from([( + "authorization".into(), + "Bearer header-secret".into(), + )]), + stream: false, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 60, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }; + let response = ResponseMeta { + status: 200, + headers: vec![("set-cookie".into(), "session=response-secret".into())], + }; + let frame = Frame::new( + 1, + MsgType::RequestBody, + 0, + Bytes::from_static(b"request-body-secret"), + ); + + let debug = format!("{meta:?} {response:?} {frame:?}"); + for secret in [ + "user", + "password", + "url-secret", + "header-secret", + "response-secret", + "request-body-secret", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}: {debug}"); + } + assert!(debug.contains("header_names")); + assert!(debug.contains("payload_len")); + } } diff --git a/crates/aether-contracts/src/tunnel_security.rs b/crates/aether-contracts/src/tunnel_security.rs index 45e3b6614..2d03e9693 100644 --- a/crates/aether-contracts/src/tunnel_security.rs +++ b/crates/aether-contracts/src/tunnel_security.rs @@ -3,7 +3,7 @@ use aes_gcm::{Aes256Gcm, KeyInit, Nonce}; use base64::Engine; use bytes::{Buf, BufMut, Bytes, BytesMut}; use hmac::{Hmac, Mac}; -use sha2::Sha256; +use sha2::{Digest as _, Sha256}; use std::sync::atomic::{AtomicU64, Ordering}; use crate::tunnel::{Frame, MsgType, HEADER_SIZE}; @@ -12,10 +12,21 @@ type HmacSha256 = Hmac; pub const TUNNEL_SECURITY_HEADER: &str = "x-aether-tunnel-security"; pub const TUNNEL_SECURITY_SESSION_HEADER: &str = "x-aether-tunnel-security-session"; +pub const TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER: &str = "x-aether-tunnel-security-proof-timestamp"; +pub const TUNNEL_SECURITY_PROOF_NONCE_HEADER: &str = "x-aether-tunnel-security-proof-nonce"; +pub const TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER: &str = "x-aether-tunnel-security-proof-signature"; +pub const TUNNEL_GENERATION_HEADER: &str = "x-aether-tunnel-generation"; +pub const TUNNEL_CONTROL_PLANE_NODE_ID_HEADER: &str = "x-aether-tunnel-control-plane-node-id"; +pub const TUNNEL_CONTROL_PLANE_GENERATION_HEADER: &str = "x-aether-tunnel-control-plane-generation"; +pub const TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER: &str = "x-aether-tunnel-control-plane-timestamp"; +pub const TUNNEL_CONTROL_PLANE_NONCE_HEADER: &str = "x-aether-tunnel-control-plane-nonce"; +pub const TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER: &str = "x-aether-tunnel-control-plane-signature"; pub const TUNNEL_SECURITY_NON_TLS_REQUIRED: &str = "non_tls_required"; pub const FLAG_ENCRYPTED: u8 = 0x04; const CONTEXT: &[u8] = b"aether-tunnel-secure-v1"; +const HANDSHAKE_PROOF_CONTEXT: &[u8] = b"aether-tunnel-handshake-proof-v1"; +const CONTROL_PLANE_AUTH_CONTEXT: &[u8] = b"aether-tunnel-control-plane-auth-v1"; const CLIENT_TO_SERVER_LABEL: &[u8] = b"client-to-server"; const SERVER_TO_CLIENT_LABEL: &[u8] = b"server-to-client"; const CLIENT_TO_SERVER_NONCE_PREFIX: [u8; 4] = *b"c2s1"; @@ -41,6 +52,8 @@ pub enum TunnelSecurityError { PayloadTooShort, #[error("secure tunnel frame sequence is not the expected next value")] UnexpectedSequence, + #[error("secure tunnel frame sequence space is exhausted")] + SequenceExhausted, #[error("secure tunnel frame encryption failed")] Encrypt, #[error("secure tunnel frame decryption failed")] @@ -98,7 +111,12 @@ impl SecureFrameCodec { } pub fn encrypt_frame(&self, frame: Frame) -> Result { - let sequence = self.next_sequence.fetch_add(1, Ordering::Relaxed); + let sequence = self + .next_sequence + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| { + current.checked_add(1) + }) + .map_err(|_| TunnelSecurityError::SequenceExhausted)?; let nonce_bytes = nonce_bytes(self.seal_prefix, sequence); let nonce = Nonce::from_slice(&nonce_bytes); let clear_flags = frame.flags & !FLAG_ENCRYPTED; @@ -154,8 +172,17 @@ impl SecureFrameCodec { }, ) .map_err(|_| TunnelSecurityError::Decrypt)?; + let next_sequence = expected_sequence + .checked_add(1) + .ok_or(TunnelSecurityError::SequenceExhausted)?; self.next_open_sequence - .store(expected_sequence.wrapping_add(1), Ordering::Relaxed); + .compare_exchange( + expected_sequence, + next_sequence, + Ordering::AcqRel, + Ordering::Acquire, + ) + .map_err(|_| TunnelSecurityError::UnexpectedSequence)?; Ok(Frame::new( frame.stream_id, @@ -167,14 +194,299 @@ impl SecureFrameCodec { } pub fn decode_psk(key: &str) -> Result<[u8; 32], TunnelSecurityError> { + let key = key.trim(); + if key.len() > 44 { + return Err(TunnelSecurityError::InvalidKey); + } let decoded = base64::engine::general_purpose::STANDARD - .decode(key.trim()) + .decode(key) .map_err(|_| TunnelSecurityError::InvalidKey)?; decoded .try_into() .map_err(|_| TunnelSecurityError::InvalidKey) } +pub fn sign_tunnel_security_handshake( + key: &str, + node_id: &str, + security_mode: &str, + session_id: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, +) -> Result { + sign_tunnel_security_handshake_for_generation( + key, + node_id, + "", + security_mode, + session_id, + protocol_version, + timestamp_unix_secs, + nonce, + ) +} + +// These arguments are the versioned handshake transcript. Keep them explicit +// and ordered so existing clients and servers compute the same MAC. +#[allow(clippy::too_many_arguments)] +pub fn sign_tunnel_security_handshake_for_generation( + key: &str, + node_id: &str, + tunnel_generation: &str, + security_mode: &str, + session_id: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, +) -> Result { + let psk = decode_psk(key)?; + let mut mac = + ::new_from_slice(&psk).expect("HMAC accepts a 32-byte tunnel PSK"); + update_handshake_proof_mac( + &mut mac, + node_id, + tunnel_generation, + security_mode, + session_id, + protocol_version, + timestamp_unix_secs, + nonce, + ); + Ok(base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())) +} + +// The verifier preserves the legacy public signature while delegating to the +// generation-aware transcript implementation. +#[allow(clippy::too_many_arguments)] +pub fn verify_tunnel_security_handshake( + key: &str, + node_id: &str, + security_mode: &str, + session_id: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, + signature: &str, +) -> bool { + verify_tunnel_security_handshake_for_generation( + key, + node_id, + "", + security_mode, + session_id, + protocol_version, + timestamp_unix_secs, + nonce, + signature, + ) +} + +// This mirrors the signing API exactly; the argument order is part of the +// authenticated handshake format. +#[allow(clippy::too_many_arguments)] +pub fn verify_tunnel_security_handshake_for_generation( + key: &str, + node_id: &str, + tunnel_generation: &str, + security_mode: &str, + session_id: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, + signature: &str, +) -> bool { + let Ok(psk) = decode_psk(key) else { + return false; + }; + if signature.len() > 43 { + return false; + } + let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else { + return false; + }; + let Ok(mut mac) = ::new_from_slice(&psk) else { + return false; + }; + update_handshake_proof_mac( + &mut mac, + node_id, + tunnel_generation, + security_mode, + session_id, + protocol_version, + timestamp_unix_secs, + nonce, + ); + mac.verify_slice(&signature).is_ok() +} + +pub fn sign_tunnel_control_plane_request( + key: &str, + method: &str, + path: &str, + node_id: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], +) -> Result { + sign_tunnel_control_plane_request_for_generation( + key, + method, + path, + node_id, + "", + timestamp_unix_secs, + nonce, + body, + ) +} + +// Control-plane authentication signs these fields in this fixed order. Keep +// the public API stable instead of introducing a reordered parameter object. +#[allow(clippy::too_many_arguments)] +pub fn sign_tunnel_control_plane_request_for_generation( + key: &str, + method: &str, + path: &str, + node_id: &str, + tunnel_generation: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], +) -> Result { + let psk = decode_psk(key)?; + let mut mac = + ::new_from_slice(&psk).expect("HMAC accepts a 32-byte tunnel PSK"); + update_control_plane_auth_mac( + &mut mac, + method, + path, + node_id, + tunnel_generation, + timestamp_unix_secs, + nonce, + body, + ); + Ok(base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())) +} + +// Preserve the legacy verifier signature; it must feed the same transcript as +// the corresponding signing function. +#[allow(clippy::too_many_arguments)] +pub fn verify_tunnel_control_plane_request( + key: &str, + method: &str, + path: &str, + node_id: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], + signature: &str, +) -> bool { + verify_tunnel_control_plane_request_for_generation( + key, + method, + path, + node_id, + "", + timestamp_unix_secs, + nonce, + body, + signature, + ) +} + +// The generation-aware verifier intentionally mirrors the signer field order, +// which is part of the control-plane wire contract. +#[allow(clippy::too_many_arguments)] +pub fn verify_tunnel_control_plane_request_for_generation( + key: &str, + method: &str, + path: &str, + node_id: &str, + tunnel_generation: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], + signature: &str, +) -> bool { + let Ok(psk) = decode_psk(key) else { + return false; + }; + if signature.len() > 43 { + return false; + } + let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else { + return false; + }; + let Ok(mut mac) = ::new_from_slice(&psk) else { + return false; + }; + update_control_plane_auth_mac( + &mut mac, + method, + path, + node_id, + tunnel_generation, + timestamp_unix_secs, + nonce, + body, + ); + mac.verify_slice(&signature).is_ok() +} + +// Keep the MAC input fields separate and visibly ordered to avoid accidental +// changes to the authenticated control-plane transcript. +#[allow(clippy::too_many_arguments)] +fn update_control_plane_auth_mac( + mac: &mut HmacSha256, + method: &str, + path: &str, + node_id: &str, + tunnel_generation: &str, + timestamp_unix_secs: u64, + nonce: &str, + body: &[u8], +) { + mac.update(CONTROL_PLANE_AUTH_CONTEXT); + update_handshake_proof_field(mac, method.as_bytes()); + update_handshake_proof_field(mac, path.as_bytes()); + update_handshake_proof_field(mac, node_id.as_bytes()); + update_handshake_proof_field(mac, tunnel_generation.as_bytes()); + mac.update(×tamp_unix_secs.to_be_bytes()); + update_handshake_proof_field(mac, nonce.as_bytes()); + mac.update(&Sha256::digest(body)); +} + +// This helper encodes the handshake transcript in a fixed cryptographic order; +// grouping arguments into a struct would obscure that wire-level contract. +#[allow(clippy::too_many_arguments)] +fn update_handshake_proof_mac( + mac: &mut HmacSha256, + node_id: &str, + tunnel_generation: &str, + security_mode: &str, + session_id: &str, + protocol_version: u8, + timestamp_unix_secs: u64, + nonce: &str, +) { + mac.update(HANDSHAKE_PROOF_CONTEXT); + update_handshake_proof_field(mac, node_id.as_bytes()); + update_handshake_proof_field(mac, tunnel_generation.as_bytes()); + update_handshake_proof_field(mac, security_mode.as_bytes()); + update_handshake_proof_field(mac, session_id.as_bytes()); + mac.update(&[protocol_version]); + mac.update(×tamp_unix_secs.to_be_bytes()); + update_handshake_proof_field(mac, nonce.as_bytes()); +} + +fn update_handshake_proof_field(mac: &mut HmacSha256, value: &[u8]) { + mac.update(&(value.len() as u64).to_be_bytes()); + mac.update(value); +} + fn derive_key(psk: &[u8; 32], session_id: &[u8], label: &[u8]) -> [u8; 32] { let mut mac = ::new_from_slice(psk).expect("HMAC accepts 32-byte PSK"); mac.update(CONTEXT); @@ -276,6 +588,70 @@ mod tests { )); } + #[test] + fn secure_frame_accepts_a_concurrent_sequence_only_once() { + use std::sync::{Arc, Barrier}; + + let client = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client) + .expect("client codec"); + let server = Arc::new( + SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Server) + .expect("server codec"), + ); + let encrypted = client + .encrypt_frame(Frame::new( + 1, + MsgType::RequestBody, + 0, + Bytes::from_static(b"secret"), + )) + .expect("encrypt"); + let wire = Frame::decode(encrypted).expect("wire frame"); + let barrier = Arc::new(Barrier::new(3)); + let workers = (0..2) + .map(|_| { + let server = Arc::clone(&server); + let barrier = Arc::clone(&barrier); + let wire = wire.clone(); + std::thread::spawn(move || { + barrier.wait(); + server.decrypt_frame(wire) + }) + }) + .collect::>(); + barrier.wait(); + let results = workers + .into_iter() + .map(|worker| worker.join().expect("worker should not panic")) + .collect::>(); + + assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1); + assert_eq!( + results + .iter() + .filter(|result| matches!(result, Err(TunnelSecurityError::UnexpectedSequence))) + .count(), + 1 + ); + } + + #[test] + fn secure_frame_fails_closed_before_sequence_wrap() { + let codec = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client) + .expect("codec"); + codec.next_sequence.store(u64::MAX, Ordering::Relaxed); + + assert!(matches!( + codec.encrypt_frame(Frame::new( + 1, + MsgType::RequestBody, + 0, + Bytes::from_static(b"secret"), + )), + Err(TunnelSecurityError::SequenceExhausted) + )); + } + #[test] fn secure_frame_rejects_out_of_order_sequence_without_advancing() { let client = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client) @@ -335,4 +711,195 @@ mod tests { assert_ne!(encrypted_a, encrypted_b); } + + #[test] + fn handshake_proof_round_trips_and_binds_every_field() { + let key = test_key(); + let signature = sign_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + ) + .expect("sign handshake"); + assert_eq!(signature.len(), 43); + assert!(verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + &signature, + )); + + for valid in [ + verify_tunnel_security_handshake( + &key, + "node-2", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + &signature, + ), + verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "1123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + &signature, + ), + verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 2, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + &signature, + ), + verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_001, + "abcdef0123456789abcdef0123456789", + &signature, + ), + verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "bbcdef0123456789abcdef0123456789", + &signature, + ), + ] { + assert!(!valid, "tampered handshake field must fail verification"); + } + } + + #[test] + fn handshake_proof_rejects_wrong_key_and_malformed_signature() { + let key = test_key(); + let signature = sign_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + ) + .expect("sign handshake"); + let wrong_key = base64::engine::general_purpose::STANDARD.encode([8_u8; 32]); + assert!(!verify_tunnel_security_handshake( + &wrong_key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + &signature, + )); + assert!(!verify_tunnel_security_handshake( + &key, + "node-1", + TUNNEL_SECURITY_NON_TLS_REQUIRED, + "0123456789abcdef0123456789abcdef", + 3, + 1_700_000_000, + "abcdef0123456789abcdef0123456789", + "not-base64", + )); + } + + #[test] + fn control_plane_proof_round_trips_and_binds_identity_route_and_body() { + let key = test_key(); + let body = br#"{"node_id":"node-1","heartbeat_id":7}"#; + let signature = sign_tunnel_control_plane_request( + &key, + "POST", + "/api/internal/tunnel/heartbeat", + "node-1", + 1_700_000_000, + "nonce-0123456789abcdef", + body, + ) + .expect("sign control-plane request"); + + assert!(verify_tunnel_control_plane_request( + &key, + "POST", + "/api/internal/tunnel/heartbeat", + "node-1", + 1_700_000_000, + "nonce-0123456789abcdef", + body, + &signature, + )); + for valid in [ + verify_tunnel_control_plane_request( + &key, + "GET", + "/api/internal/tunnel/heartbeat", + "node-1", + 1_700_000_000, + "nonce-0123456789abcdef", + body, + &signature, + ), + verify_tunnel_control_plane_request( + &key, + "POST", + "/api/internal/tunnel/node-status", + "node-1", + 1_700_000_000, + "nonce-0123456789abcdef", + body, + &signature, + ), + verify_tunnel_control_plane_request( + &key, + "POST", + "/api/internal/tunnel/heartbeat", + "node-2", + 1_700_000_000, + "nonce-0123456789abcdef", + body, + &signature, + ), + verify_tunnel_control_plane_request( + &key, + "POST", + "/api/internal/tunnel/heartbeat", + "node-1", + 1_700_000_000, + "nonce-0123456789abcdef", + br#"{"node_id":"node-1","heartbeat_id":8}"#, + &signature, + ), + ] { + assert!( + !valid, + "tampered control-plane field must fail verification" + ); + } + } } diff --git a/crates/aether-crypto/Cargo.toml b/crates/aether-crypto/Cargo.toml index 369387d99..acf801039 100644 --- a/crates/aether-crypto/Cargo.toml +++ b/crates/aether-crypto/Cargo.toml @@ -8,6 +8,7 @@ description = "Shared crypto compatibility helpers for Rust migration" [dependencies] aes.workspace = true +aws-lc-rs.workspace = true base64.workspace = true cbc.workspace = true hmac.workspace = true diff --git a/crates/aether-crypto/src/lib.rs b/crates/aether-crypto/src/lib.rs index c36befdc6..9003cb89d 100644 --- a/crates/aether-crypto/src/lib.rs +++ b/crates/aether-crypto/src/lib.rs @@ -1,7 +1,9 @@ mod python_fernet; +mod rsa_pkcs1_sha256; pub use python_fernet::{ decrypt_python_fernet_ciphertext, derive_python_fernet_key, encrypt_python_fernet_plaintext, looks_like_python_fernet_ciphertext, warm_python_fernet_secret, PythonFernetCompat, PythonFernetError, APP_SALT_HEX, APP_SALT_SEED, DEVELOPMENT_ENCRYPTION_KEY, }; +pub use rsa_pkcs1_sha256::{rsa_pkcs1_sha256_sign, rsa_pkcs1_sha256_verify, RsaPkcs1Sha256Error}; diff --git a/crates/aether-crypto/src/python_fernet.rs b/crates/aether-crypto/src/python_fernet.rs index af113763e..899437b5d 100644 --- a/crates/aether-crypto/src/python_fernet.rs +++ b/crates/aether-crypto/src/python_fernet.rs @@ -19,8 +19,17 @@ const SIGNING_KEY_SIZE: usize = 16; const ENCRYPTION_KEY_SIZE: usize = 16; const MIN_CIPHERTEXT_SIZE: usize = 16; const MIN_TOKEN_SIZE: usize = 1 + 8 + IV_SIZE + MIN_CIPHERTEXT_SIZE + HMAC_SIZE; +const STANDARD_FERNET_TOKEN_PREFIX: &str = "gAAAA"; +const WRAPPED_FERNET_TOKEN_PREFIX: &str = "Z0FBQUFB"; const PBKDF2_ITERATIONS: u32 = 100_000; const MAX_CACHED_DERIVED_KEYS: usize = 16; +const MAX_FERNET_PLAINTEXT_BYTES: usize = 16 * 1024 * 1024; + +// A Fernet token contains fixed metadata, PKCS#7 padded ciphertext and an HMAC. +// Aether's compatibility format then base64-encodes that token twice. +const MAX_FERNET_TOKEN_BYTES: usize = + 1 + 8 + IV_SIZE + MAX_FERNET_PLAINTEXT_BYTES + AES_BLOCK_SIZE + HMAC_SIZE; +const AES_BLOCK_SIZE: usize = 16; pub const APP_SALT_SEED: &[u8] = b"aether-v1"; pub const APP_SALT_HEX: &str = "8797080a7a4b45b4810e934d1af36261"; @@ -55,7 +64,7 @@ fn minimum_wrapped_token_len() -> usize { base64_unpadded_len(base64_unpadded_len(MIN_TOKEN_SIZE)) } -#[derive(Debug, Default)] +#[derive(Default)] struct RawFernetKeyCache { entries: HashMap, [u8; 32]>, insertion_order: VecDeque>, @@ -99,14 +108,28 @@ pub enum PythonFernetError { InvalidPadding, #[error("invalid Python Fernet plaintext utf-8")] InvalidUtf8(#[from] std::string::FromUtf8Error), + #[error("Python Fernet plaintext exceeds {limit_bytes} bytes")] + PlaintextTooLarge { limit_bytes: usize }, + #[error("Python Fernet ciphertext exceeds the supported size")] + CiphertextTooLarge, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct PythonFernetCompat { signing_key: [u8; SIGNING_KEY_SIZE], encryption_key: [u8; ENCRYPTION_KEY_SIZE], } +impl std::fmt::Debug for PythonFernetCompat { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PythonFernetCompat") + .field("signing_key", &"[REDACTED]") + .field("encryption_key", &"[REDACTED]") + .finish() + } +} + impl PythonFernetCompat { pub fn from_secret(secret: &str) -> Self { let raw_key = raw_fernet_key(secret); @@ -115,13 +138,10 @@ impl PythonFernetCompat { pub fn decrypt_ciphertext(&self, ciphertext: &str) -> Result { if ciphertext.is_empty() { - return Ok(String::new()); + return Err(PythonFernetError::InvalidTokenStructure); } - let outer = - decode_urlsafe(ciphertext).map_err(|_| PythonFernetError::InvalidOuterBase64)?; - let token = - decode_urlsafe_bytes(&outer).map_err(|_| PythonFernetError::InvalidInnerBase64)?; + let token = decode_wrapped_fernet_token_with_limit(ciphertext, MAX_FERNET_TOKEN_BYTES)?; let plaintext = self.decrypt_token_bytes(token)?; String::from_utf8(plaintext).map_err(PythonFernetError::InvalidUtf8) } @@ -146,6 +166,9 @@ impl PythonFernetCompat { } fn decrypt_token_bytes(&self, mut token: Vec) -> Result, PythonFernetError> { + if token.len() > MAX_FERNET_TOKEN_BYTES { + return Err(PythonFernetError::CiphertextTooLarge); + } if token.len() < MIN_TOKEN_SIZE { return Err(PythonFernetError::InvalidTokenStructure); } @@ -174,6 +197,11 @@ impl PythonFernetCompat { .map_err(|_| PythonFernetError::InvalidPadding)? .len() }; + if plaintext_len > MAX_FERNET_PLAINTEXT_BYTES { + return Err(PythonFernetError::PlaintextTooLarge { + limit_bytes: MAX_FERNET_PLAINTEXT_BYTES, + }); + } token.copy_within(ciphertext_offset..ciphertext_offset + plaintext_len, 0); token.truncate(plaintext_len); @@ -187,6 +215,11 @@ impl PythonFernetCompat { iv: [u8; IV_SIZE], ) -> Result { let plaintext = plaintext.as_bytes(); + if plaintext.len() > MAX_FERNET_PLAINTEXT_BYTES { + return Err(PythonFernetError::PlaintextTooLarge { + limit_bytes: MAX_FERNET_PLAINTEXT_BYTES, + }); + } let mut padded = vec![0u8; plaintext.len() + IV_SIZE]; padded[..plaintext.len()].copy_from_slice(plaintext); let ciphertext = Aes128CbcEnc::new((&self.encryption_key).into(), (&iv).into()) @@ -227,14 +260,28 @@ pub fn decrypt_python_fernet_ciphertext( pub fn looks_like_python_fernet_ciphertext(ciphertext: &str) -> bool { let ciphertext = ciphertext.trim(); - if ciphertext.is_empty() || ciphertext.len() < minimum_wrapped_token_len() { + if ciphertext.is_empty() { + return false; + } + let max_inner_encoded_bytes = maximum_base64_len_for_decoded_limit(MAX_FERNET_TOKEN_BYTES); + if ciphertext.len() > maximum_base64_len_for_decoded_limit(max_inner_encoded_bytes) { + return false; + } + if (ciphertext.len() >= minimum_wrapped_token_len() + && ciphertext.starts_with(WRAPPED_FERNET_TOKEN_PREFIX)) + || (ciphertext.len() >= base64_unpadded_len(MIN_TOKEN_SIZE) + && ciphertext.starts_with(STANDARD_FERNET_TOKEN_PREFIX)) + { + return true; + } + if ciphertext.len() < minimum_wrapped_token_len() { return false; } - let Ok(outer) = decode_urlsafe(ciphertext) else { + let Ok(outer) = decode_urlsafe_with_limit(ciphertext, max_inner_encoded_bytes) else { return false; }; - let Ok(inner) = decode_urlsafe_bytes(&outer) else { + let Ok(inner) = decode_urlsafe_bytes_with_limit(&outer, MAX_FERNET_TOKEN_BYTES) else { return false; }; @@ -287,6 +334,9 @@ fn raw_fernet_key(secret: &str) -> [u8; 32] { } fn decode_direct_fernet_key(secret: &str) -> Result<[u8; 32], PythonFernetError> { + if secret.len() > base64_encoded_len(32) { + return Err(PythonFernetError::InvalidTokenStructure); + } let decoded = URL_SAFE .decode(secret) .or_else(|_| STANDARD.decode(secret)) @@ -298,32 +348,75 @@ fn decode_direct_fernet_key(secret: &str) -> Result<[u8; 32], PythonFernetError> Ok(raw_key) } -fn decode_urlsafe(value: &str) -> Result, base64::DecodeError> { - decode_with_engine_fallback(value.as_bytes()) +#[derive(Debug)] +enum BoundedBase64Error { + TooLarge, + Invalid, } -fn decode_urlsafe_bytes(value: &[u8]) -> Result, base64::DecodeError> { - decode_with_engine_fallback(value) +fn maximum_base64_len_for_decoded_limit(decoded_limit: usize) -> usize { + decoded_limit + .checked_add(2) + .and_then(|value| value.checked_div(3)) + .and_then(|value| value.checked_mul(4)) + .unwrap_or(usize::MAX) } -fn decode_with_engine_fallback(value: &[u8]) -> Result, base64::DecodeError> { +fn decode_urlsafe_with_limit( + value: &str, + decoded_limit: usize, +) -> Result, BoundedBase64Error> { + decode_urlsafe_bytes_with_limit(value.as_bytes(), decoded_limit) +} + +fn decode_urlsafe_bytes_with_limit( + value: &[u8], + decoded_limit: usize, +) -> Result, BoundedBase64Error> { + if value.len() > maximum_base64_len_for_decoded_limit(decoded_limit) { + return Err(BoundedBase64Error::TooLarge); + } let mut decoded = Vec::with_capacity(decoded_len_estimate(value.len())); match URL_SAFE.decode_vec(value, &mut decoded) { - Ok(()) => Ok(decoded), + Ok(()) if decoded.len() <= decoded_limit => Ok(decoded), + Ok(()) => Err(BoundedBase64Error::TooLarge), Err(_) => { decoded.clear(); - URL_SAFE_NO_PAD.decode_vec(value, &mut decoded)?; + URL_SAFE_NO_PAD + .decode_vec(value, &mut decoded) + .map_err(|_| BoundedBase64Error::Invalid)?; + if decoded.len() > decoded_limit { + return Err(BoundedBase64Error::TooLarge); + } Ok(decoded) } } } +fn decode_wrapped_fernet_token_with_limit( + ciphertext: &str, + max_token_bytes: usize, +) -> Result, PythonFernetError> { + let max_inner_encoded_bytes = maximum_base64_len_for_decoded_limit(max_token_bytes); + let outer = decode_urlsafe_with_limit(ciphertext, max_inner_encoded_bytes).map_err( + |error| match error { + BoundedBase64Error::TooLarge => PythonFernetError::CiphertextTooLarge, + BoundedBase64Error::Invalid => PythonFernetError::InvalidOuterBase64, + }, + )?; + decode_urlsafe_bytes_with_limit(&outer, max_token_bytes).map_err(|error| match error { + BoundedBase64Error::TooLarge => PythonFernetError::CiphertextTooLarge, + BoundedBase64Error::Invalid => PythonFernetError::InvalidInnerBase64, + }) +} + #[cfg(test)] mod tests { use super::{ - decrypt_python_fernet_ciphertext, derive_python_fernet_key, - encrypt_python_fernet_plaintext, looks_like_python_fernet_ciphertext, PythonFernetCompat, - PythonFernetError, APP_SALT_HEX, DEVELOPMENT_ENCRYPTION_KEY, + decode_wrapped_fernet_token_with_limit, decrypt_python_fernet_ciphertext, + derive_python_fernet_key, encrypt_python_fernet_plaintext, + looks_like_python_fernet_ciphertext, maximum_base64_len_for_decoded_limit, + PythonFernetCompat, PythonFernetError, APP_SALT_HEX, DEVELOPMENT_ENCRYPTION_KEY, }; #[test] @@ -341,6 +434,16 @@ mod tests { assert_eq!(derive_python_fernet_key(direct_key), direct_key); } + #[test] + fn debug_output_never_exposes_fernet_key_material() { + let crypto = PythonFernetCompat::from_secret(DEVELOPMENT_ENCRYPTION_KEY); + + assert_eq!( + format!("{crypto:?}"), + "PythonFernetCompat { signing_key: \"[REDACTED]\", encryption_key: \"[REDACTED]\" }" + ); + } + #[test] fn decrypts_legacy_python_ciphertext_with_direct_fernet_key() { let direct_key = "h0Wzfv1ieDOgsmGELpKEV7qH8QpFP+lRnPW4pI0g4/M="; @@ -400,6 +503,8 @@ mod tests { .expect("ciphertext should build"); ciphertext.replace_range(ciphertext.len() - 2.., "AA"); + assert!(looks_like_python_fernet_ciphertext(&ciphertext)); + let err = decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &ciphertext) .expect_err("tampered ciphertext should fail"); assert!(matches!( @@ -410,6 +515,35 @@ mod tests { )); } + #[test] + fn rejects_unauthenticated_empty_ciphertext() { + assert!(matches!( + decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, ""), + Err(PythonFernetError::InvalidTokenStructure) + )); + + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "") + .expect("empty plaintext should still have an authenticated token"); + assert_eq!( + decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &ciphertext) + .expect("authenticated empty plaintext should decrypt"), + "" + ); + } + + #[test] + fn wrapped_fernet_decode_rejects_oversized_outer_base64_before_allocation() { + let token_limit = 3; + let inner_encoded_limit = maximum_base64_len_for_decoded_limit(token_limit); + let outer_encoded_limit = maximum_base64_len_for_decoded_limit(inner_encoded_limit); + let oversized = "A".repeat(outer_encoded_limit + 1); + + assert!(matches!( + decode_wrapped_fernet_token_with_limit(&oversized, token_limit), + Err(PythonFernetError::CiphertextTooLarge) + )); + } + #[test] fn encrypt_and_decrypt_round_trip() { let ciphertext = diff --git a/crates/aether-crypto/src/rsa_pkcs1_sha256.rs b/crates/aether-crypto/src/rsa_pkcs1_sha256.rs new file mode 100644 index 000000000..777b41f6f --- /dev/null +++ b/crates/aether-crypto/src/rsa_pkcs1_sha256.rs @@ -0,0 +1,257 @@ +use aws_lc_rs::rand::SystemRandom; +use aws_lc_rs::rsa::{KeyPair, PublicKey, RsaParameters}; +use aws_lc_rs::signature::{UnparsedPublicKey, RSA_PKCS1_2048_8192_SHA256, RSA_PKCS1_SHA256}; +use base64::engine::general_purpose::STANDARD; +use base64::Engine as _; +use thiserror::Error; + +const MAX_RSA_KEY_INPUT_BYTES: usize = 64 * 1024; +const PRIVATE_KEY_PEM_LABELS: &[&str] = &["PRIVATE KEY", "RSA PRIVATE KEY"]; +const PUBLIC_KEY_PEM_LABELS: &[&str] = &["PUBLIC KEY", "RSA PUBLIC KEY"]; + +#[derive(Debug, Error, Clone, Copy, PartialEq, Eq)] +pub enum RsaPkcs1Sha256Error { + #[error("invalid RSA private key")] + InvalidPrivateKey, + #[error("invalid RSA public key")] + InvalidPublicKey, + #[error("RSA signing failed")] + SigningFailed, + #[error("invalid RSA signature encoding")] + InvalidSignature, +} + +fn decode_text_key_material(input: &[u8], labels: &[&str]) -> Option> { + let text = std::str::from_utf8(input).ok()?.trim(); + if text.is_empty() { + return None; + } + + let encoded = if text.starts_with("-----BEGIN ") { + labels.iter().find_map(|label| { + let header = format!("-----BEGIN {label}-----"); + let footer = format!("-----END {label}-----"); + text.strip_prefix(&header)?.strip_suffix(&footer) + })? + } else { + text + }; + + let mut compact = encoded + .bytes() + .filter(|byte| !byte.is_ascii_whitespace()) + .collect::>(); + let decoded = (!compact.is_empty()) + .then(|| STANDARD.decode(&compact).ok()) + .flatten(); + compact.fill(0); + decoded +} + +fn parse_private_key_der(input: &[u8]) -> Result { + KeyPair::from_pkcs8(input) + .or_else(|_| KeyPair::from_der(input)) + .map_err(|_| RsaPkcs1Sha256Error::InvalidPrivateKey) +} + +fn parse_private_key(input: &[u8]) -> Result { + if input.is_empty() || input.len() > MAX_RSA_KEY_INPUT_BYTES { + return Err(RsaPkcs1Sha256Error::InvalidPrivateKey); + } + if let Ok(key_pair) = parse_private_key_der(input) { + return Ok(key_pair); + } + + let mut der = decode_text_key_material(input, PRIVATE_KEY_PEM_LABELS) + .ok_or(RsaPkcs1Sha256Error::InvalidPrivateKey)?; + let result = parse_private_key_der(&der); + der.fill(0); + result +} + +fn parse_public_key_der(input: &[u8]) -> Result { + let public_key = + PublicKey::from_der(input).map_err(|_| RsaPkcs1Sha256Error::InvalidPublicKey)?; + let bits = RsaParameters::public_modulus_len(public_key.as_ref()) + .map_err(|_| RsaPkcs1Sha256Error::InvalidPublicKey)?; + if !(2048..=8192).contains(&bits) { + return Err(RsaPkcs1Sha256Error::InvalidPublicKey); + } + Ok(public_key) +} + +fn parse_public_key(input: &[u8]) -> Result { + if input.is_empty() || input.len() > MAX_RSA_KEY_INPUT_BYTES { + return Err(RsaPkcs1Sha256Error::InvalidPublicKey); + } + if let Ok(public_key) = parse_public_key_der(input) { + return Ok(public_key); + } + + let der = decode_text_key_material(input, PUBLIC_KEY_PEM_LABELS) + .ok_or(RsaPkcs1Sha256Error::InvalidPublicKey)?; + parse_public_key_der(&der) +} + +/// Signs `message` with RSASSA-PKCS1-v1_5 and SHA-256 using AWS-LC. +/// +/// The private key may be PKCS#8 or PKCS#1 DER, either PEM encoded or supplied +/// as bare standard-base64 DER. Raw DER bytes are accepted as well. +pub fn rsa_pkcs1_sha256_sign( + private_key: &[u8], + message: &[u8], +) -> Result, RsaPkcs1Sha256Error> { + let key_pair = parse_private_key(private_key)?; + let mut signature = vec![0; key_pair.public_modulus_len()]; + key_pair + .sign( + &RSA_PKCS1_SHA256, + &SystemRandom::new(), + message, + &mut signature, + ) + .map_err(|_| RsaPkcs1Sha256Error::SigningFailed)?; + Ok(signature) +} + +/// Verifies an RSASSA-PKCS1-v1_5 SHA-256 signature using AWS-LC. +/// +/// The public key may be PKCS#1 or X.509 SubjectPublicKeyInfo DER, either PEM +/// encoded or supplied as bare standard-base64 DER. Raw DER bytes are accepted +/// as well. +pub fn rsa_pkcs1_sha256_verify( + public_key: &[u8], + message: &[u8], + signature: &[u8], +) -> Result { + let public_key = parse_public_key(public_key)?; + let modulus_bits = RsaParameters::public_modulus_len(public_key.as_ref()) + .map_err(|_| RsaPkcs1Sha256Error::InvalidPublicKey)?; + let signature_len = (modulus_bits as usize).div_ceil(8); + if signature.len() != signature_len { + return Err(RsaPkcs1Sha256Error::InvalidSignature); + } + Ok( + UnparsedPublicKey::new(&RSA_PKCS1_2048_8192_SHA256, public_key.as_ref()) + .verify(message, signature) + .is_ok(), + ) +} + +#[cfg(test)] +mod tests { + use aws_lc_rs::encoding::{AsDer, Pkcs8V1Der, PublicKeyX509Der}; + use aws_lc_rs::rsa::{KeyPair, KeySize}; + use aws_lc_rs::signature::KeyPair as _; + use base64::engine::general_purpose::STANDARD; + use base64::Engine as _; + + use super::{rsa_pkcs1_sha256_sign, rsa_pkcs1_sha256_verify, RsaPkcs1Sha256Error}; + + fn read_der_tlv<'a>(input: &mut &'a [u8], expected_tag: u8) -> &'a [u8] { + assert_eq!(input.first().copied(), Some(expected_tag)); + let length_byte = input[1]; + let (header_len, value_len) = if length_byte & 0x80 == 0 { + (2, usize::from(length_byte)) + } else { + let length_bytes = usize::from(length_byte & 0x7f); + assert!((1..=4).contains(&length_bytes)); + let value_len = input[2..2 + length_bytes] + .iter() + .fold(0usize, |value, byte| (value << 8) | usize::from(*byte)); + (2 + length_bytes, value_len) + }; + let end = header_len + value_len; + assert!(end <= input.len()); + let value = &input[header_len..end]; + *input = &input[end..]; + value + } + + fn pkcs1_private_key_from_pkcs8(pkcs8: &[u8]) -> Vec { + let mut input = pkcs8; + let mut sequence = read_der_tlv(&mut input, 0x30); + assert!(input.is_empty()); + let _version = read_der_tlv(&mut sequence, 0x02); + let _algorithm = read_der_tlv(&mut sequence, 0x30); + read_der_tlv(&mut sequence, 0x04).to_vec() + } + + fn pem(label: &str, der: &[u8]) -> String { + format!( + "-----BEGIN {label}-----\n{}\n-----END {label}-----", + STANDARD.encode(der) + ) + } + + #[test] + fn signs_and_verifies_all_supported_rsa_key_encodings() { + let key_pair = KeyPair::generate(KeySize::Rsa2048).expect("RSA key should generate"); + let pkcs8 = AsDer::>::as_der(&key_pair) + .expect("PKCS#8 should encode") + .as_ref() + .to_vec(); + let pkcs1_private = pkcs1_private_key_from_pkcs8(&pkcs8); + let pkcs1_public = key_pair.public_key().as_ref().to_vec(); + let spki_public = AsDer::>::as_der(key_pair.public_key()) + .expect("SPKI should encode") + .as_ref() + .to_vec(); + let private_inputs = [ + pkcs8.clone(), + pkcs1_private.clone(), + pem("PRIVATE KEY", &pkcs8).into_bytes(), + pem("RSA PRIVATE KEY", &pkcs1_private).into_bytes(), + STANDARD.encode(&pkcs8).into_bytes(), + STANDARD.encode(&pkcs1_private).into_bytes(), + ]; + let public_inputs = [ + pkcs1_public.clone(), + spki_public.clone(), + pem("RSA PUBLIC KEY", &pkcs1_public).into_bytes(), + pem("PUBLIC KEY", &spki_public).into_bytes(), + STANDARD.encode(&pkcs1_public).into_bytes(), + STANDARD.encode(&spki_public).into_bytes(), + ]; + let message = b"Aether RSA-SHA256 compatibility vector"; + + let expected = + rsa_pkcs1_sha256_sign(&private_inputs[0], message).expect("PKCS#8 DER should sign"); + assert_eq!(expected.len(), 256); + for private_key in private_inputs { + assert_eq!( + rsa_pkcs1_sha256_sign(&private_key, message).expect("key format should sign"), + expected, + "PKCS#1 v1.5 output must remain deterministic across encodings" + ); + } + for public_key in public_inputs { + assert!(rsa_pkcs1_sha256_verify(&public_key, message, &expected) + .expect("key format should verify")); + } + assert!( + !rsa_pkcs1_sha256_verify(&pkcs1_public, b"tampered", &expected) + .expect("valid key with invalid signature should return false") + ); + assert_eq!( + rsa_pkcs1_sha256_verify(&pkcs1_public, message, &expected[..255]), + Err(RsaPkcs1Sha256Error::InvalidSignature) + ); + } + + #[test] + fn rejects_malformed_or_unsupported_rsa_keys() { + assert_eq!( + rsa_pkcs1_sha256_sign(b"not-a-key", b"message"), + Err(RsaPkcs1Sha256Error::InvalidPrivateKey) + ); + assert_eq!( + rsa_pkcs1_sha256_verify(b"not-a-key", b"message", b"signature"), + Err(RsaPkcs1Sha256Error::InvalidPublicKey) + ); + assert_eq!( + rsa_pkcs1_sha256_sign(&vec![b'A'; 64 * 1024 + 1], b"message"), + Err(RsaPkcs1Sha256Error::InvalidPrivateKey) + ); + } +} diff --git a/crates/aether-data/adapters/mysql/migrations/20260403000000_baseline.sql b/crates/aether-data/adapters/mysql/migrations/20260403000000_baseline.sql index 9c539a0ea..0f91573ad 100644 --- a/crates/aether-data/adapters/mysql/migrations/20260403000000_baseline.sql +++ b/crates/aether-data/adapters/mysql/migrations/20260403000000_baseline.sql @@ -578,7 +578,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( is_manual TINYINT(1) NOT NULL DEFAULT 0, proxy_url VARCHAR(500), proxy_username VARCHAR(255), - proxy_password VARCHAR(500), + proxy_password TEXT, created_at BIGINT NOT NULL, updated_at BIGINT NOT NULL, remote_config TEXT, diff --git a/crates/aether-data/adapters/mysql/migrations/20260814000000_add_usage_cost_reservations.sql b/crates/aether-data/adapters/mysql/migrations/20260814000000_add_usage_cost_reservations.sql new file mode 100644 index 000000000..aa0ac7c88 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260814000000_add_usage_cost_reservations.sql @@ -0,0 +1,35 @@ +CREATE TABLE IF NOT EXISTS usage_cost_reservations ( + `request_id` VARCHAR(128) NOT NULL, + `subject_id` VARCHAR(128) NOT NULL, + `reservation_token` VARCHAR(128) NOT NULL, + `admitted_at` BIGINT NOT NULL, + `reserved_cost_units` BIGINT NOT NULL, + `actual_cost_units` BIGINT, + `state` VARCHAR(20) NOT NULL, + `reservation_expires_at` BIGINT NOT NULL, + `retain_until` BIGINT NOT NULL, + `finalized_at` BIGINT, + `created_at` BIGINT NOT NULL, + `updated_at` BIGINT NOT NULL, + PRIMARY KEY (`reservation_token`), + CONSTRAINT usage_cost_reservations_state_check + CHECK (`state` IN ('reserved', 'finalized', 'released')), + CONSTRAINT usage_cost_reservations_reserved_cost_units_check + CHECK (`reserved_cost_units` >= 0), + CONSTRAINT usage_cost_reservations_actual_cost_units_check + CHECK (`actual_cost_units` IS NULL OR `actual_cost_units` >= 0), + CONSTRAINT usage_cost_reservations_expiry_check + CHECK (`reservation_expires_at` > `admitted_at`), + CONSTRAINT usage_cost_reservations_retention_check + CHECK (`retain_until` >= `reservation_expires_at`), + CONSTRAINT usage_cost_reservations_lifecycle_check CHECK ( + (`state` = 'reserved' AND `actual_cost_units` IS NULL AND `finalized_at` IS NULL) + OR (`state` = 'finalized' AND `actual_cost_units` IS NOT NULL AND `finalized_at` IS NOT NULL) + OR (`state` = 'released' AND `actual_cost_units` IS NOT NULL + AND `actual_cost_units` = 0 AND `finalized_at` IS NOT NULL) + ), + KEY usage_cost_reservations_request_id_idx (`request_id`), + KEY usage_cost_reservations_subject_admitted_at_idx (`subject_id`, `admitted_at`), + KEY usage_cost_reservations_reservation_expires_at_idx (`reservation_expires_at`), + KEY usage_cost_reservations_retain_until_token_idx (`retain_until`, `reservation_token`) +); diff --git a/crates/aether-data/adapters/mysql/migrations/20260815000000_add_usage_request_admissions.sql b/crates/aether-data/adapters/mysql/migrations/20260815000000_add_usage_request_admissions.sql new file mode 100644 index 000000000..7047f7711 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260815000000_add_usage_request_admissions.sql @@ -0,0 +1,22 @@ +CREATE TABLE IF NOT EXISTS usage_request_admissions ( + `request_id` VARCHAR(128) NOT NULL, + `subject_id` VARCHAR(128) NOT NULL, + `event_token` VARCHAR(128) NOT NULL, + `admitted_at` BIGINT NOT NULL, + `retain_until` BIGINT NOT NULL, + `state` VARCHAR(20) NOT NULL, + `released_at` BIGINT, + `created_at` BIGINT NOT NULL, + PRIMARY KEY (`event_token`), + CONSTRAINT usage_request_admissions_retention_check + CHECK (`retain_until` > `admitted_at`), + CONSTRAINT usage_request_admissions_state_check + CHECK (`state` IN ('active', 'released')), + CONSTRAINT usage_request_admissions_lifecycle_check CHECK ( + (`state` = 'active' AND `released_at` IS NULL) + OR (`state` = 'released' AND `released_at` IS NOT NULL + AND `released_at` >= `admitted_at`) + ), + KEY usage_request_admissions_subject_admitted_at_idx (`subject_id`, `admitted_at`), + KEY usage_request_admissions_retain_until_token_idx (`retain_until`, `event_token`) +); diff --git a/crates/aether-data/adapters/mysql/migrations/20260816000000_add_usage_cost_reservation_user_foreign_key.sql b/crates/aether-data/adapters/mysql/migrations/20260816000000_add_usage_cost_reservation_user_foreign_key.sql new file mode 100644 index 000000000..de9f0e7e3 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260816000000_add_usage_cost_reservation_user_foreign_key.sql @@ -0,0 +1,12 @@ +-- Preserve the already-published cost-ledger migration checksum. This +-- follow-up also upgrades databases which ran the feature branch before user +-- ownership was enforced. Keep one ALTER TABLE per migration because MySQL +-- DDL implicitly commits. +DELETE reservation +FROM usage_cost_reservations AS reservation +LEFT JOIN users AS app_user ON app_user.id = reservation.subject_id +WHERE app_user.id IS NULL; + +ALTER TABLE usage_cost_reservations + ADD CONSTRAINT usage_cost_reservations_subject_id_fkey + FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE; diff --git a/crates/aether-data/adapters/mysql/migrations/20260817000000_add_usage_request_admission_user_foreign_key.sql b/crates/aether-data/adapters/mysql/migrations/20260817000000_add_usage_request_admission_user_foreign_key.sql new file mode 100644 index 000000000..c294f4b6a --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260817000000_add_usage_request_admission_user_foreign_key.sql @@ -0,0 +1,10 @@ +-- Split from the cost-ledger foreign key so a MySQL implicit DDL commit cannot +-- leave two table changes behind one dirty migration record. +DELETE admission +FROM usage_request_admissions AS admission +LEFT JOIN users AS app_user ON app_user.id = admission.subject_id +WHERE app_user.id IS NULL; + +ALTER TABLE usage_request_admissions + ADD CONSTRAINT usage_request_admissions_subject_id_fkey + FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE; diff --git a/crates/aether-data/adapters/mysql/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql b/crates/aether-data/adapters/mysql/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql new file mode 100644 index 000000000..434e4c581 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql @@ -0,0 +1,55 @@ +-- A gateway transaction identifier may repeat across payment methods, but +-- must never identify two orders in the same normalized method. MySQL commits +-- persistent DDL implicitly, so reject historical conflicts before any +-- persistent UPDATE or ALTER TABLE. CREATE/DROP TEMPORARY TABLE do not cause an +-- implicit commit, and the leading DROP also makes a same-session retry safe. +-- Diagnose with: +-- SELECT LOWER(TRIM(payment_method)), +-- CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin, +-- COUNT(*) +-- FROM payment_orders +-- WHERE gateway_order_id IS NOT NULL +-- GROUP BY LOWER(TRIM(payment_method)), +-- CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin +-- HAVING COUNT(*) > 1; +DROP TEMPORARY TABLE IF EXISTS aether_payment_gateway_order_uniqueness_preflight; + +CREATE TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight ( + conflict_marker TINYINT NOT NULL PRIMARY KEY +); + +INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) +VALUES (1); + +-- Inserting the same marker fails on the first conflicting group. The grouping +-- mirrors the values and collations used by the normalization and final index: +-- payment methods use their existing column collation after LOWER/TRIM, while +-- opaque gateway identifiers use MySQL 8's case-sensitive binary collation. +INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) +SELECT 1 +FROM payment_orders +WHERE gateway_order_id IS NOT NULL +GROUP BY + LOWER(TRIM(payment_method)), + CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin +HAVING COUNT(*) > 1 +LIMIT 1; + +DROP TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight; + +UPDATE payment_orders +SET payment_method = LOWER(TRIM(payment_method)) +WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method)); + +UPDATE payment_callbacks +SET payment_method = LOWER(TRIM(payment_method)) +WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method)); + +-- Gateway identifiers are opaque and case-sensitive. Changing the column +-- collation and adding the unique index in one ALTER avoids a persistent +-- intermediate schema if either operation fails. +ALTER TABLE payment_orders + MODIFY COLUMN gateway_order_id VARCHAR(128) + CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NULL, + ADD UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id + (payment_method, gateway_order_id); diff --git a/crates/aether-data/adapters/mysql/migrations/20260821130000_add_user_security_version.sql b/crates/aether-data/adapters/mysql/migrations/20260821130000_add_user_security_version.sql new file mode 100644 index 000000000..9e68683bb --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260821130000_add_user_security_version.sql @@ -0,0 +1,5 @@ +ALTER TABLE users + ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0; + +ALTER TABLE user_sessions + ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0; diff --git a/crates/aether-data/adapters/mysql/migrations/20260827040000_expand_proxy_password_ciphertext.sql b/crates/aether-data/adapters/mysql/migrations/20260827040000_expand_proxy_password_ciphertext.sql new file mode 100644 index 000000000..c84e60c0b --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260827040000_expand_proxy_password_ciphertext.sql @@ -0,0 +1,3 @@ +-- Purpose-bound Fernet envelopes can exceed the former 500-character plaintext limit. +ALTER TABLE proxy_nodes + MODIFY COLUMN proxy_password TEXT NULL; diff --git a/crates/aether-data/adapters/mysql/migrations/20260827050000_anonymize_deleted_user_history.sql b/crates/aether-data/adapters/mysql/migrations/20260827050000_anonymize_deleted_user_history.sql new file mode 100644 index 000000000..f8a2b9cb2 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260827050000_anonymize_deleted_user_history.sql @@ -0,0 +1,93 @@ +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'user_plan_entitlements' + AND CONSTRAINT_NAME = 'user_plan_entitlements_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE user_plan_entitlements DROP FOREIGN KEY user_plan_entitlements_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'entitlement_usage_ledgers' + AND CONSTRAINT_NAME = 'entitlement_usage_ledgers_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE entitlement_usage_ledgers DROP FOREIGN KEY entitlement_usage_ledgers_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'user_referrals' + AND CONSTRAINT_NAME = 'user_referrals_inviter_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE user_referrals DROP FOREIGN KEY user_referrals_inviter_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'user_referrals' + AND CONSTRAINT_NAME = 'user_referrals_invitee_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE user_referrals DROP FOREIGN KEY user_referrals_invitee_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'referral_rewards' + AND CONSTRAINT_NAME = 'referral_rewards_inviter_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE referral_rewards DROP FOREIGN KEY referral_rewards_inviter_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +SET @aether_drop_fact_user_fk_sql := IF( + EXISTS ( + SELECT 1 FROM information_schema.TABLE_CONSTRAINTS + WHERE CONSTRAINT_SCHEMA = DATABASE() + AND TABLE_NAME = 'referral_rewards' + AND CONSTRAINT_NAME = 'referral_rewards_invitee_user_id_fkey' + AND CONSTRAINT_TYPE = 'FOREIGN KEY' + ), + 'ALTER TABLE referral_rewards DROP FOREIGN KEY referral_rewards_invitee_user_id_fkey', + 'DO 0' +); +PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql; +EXECUTE aether_drop_fact_user_fk_stmt; +DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt; + +-- Keep existing historical row values unchanged. Runtime deletion paths enforce +-- the current anonymization policy for newly deleted users. +SELECT 1; diff --git a/crates/aether-data/adapters/mysql/migrations/20260831000000_enforce_ldap_config_singleton.sql b/crates/aether-data/adapters/mysql/migrations/20260831000000_enforce_ldap_config_singleton.sql new file mode 100644 index 000000000..c8dae1f65 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260831000000_enforce_ldap_config_singleton.sql @@ -0,0 +1,13 @@ +-- LDAP configuration is a database-wide singleton. Preserve the row selected by the legacy +-- reader (the smallest id), remove historical duplicates, and let the database arbitrate +-- concurrent first creation. +DELETE FROM ldap_configs +WHERE id <> ( + SELECT keep_id + FROM (SELECT MIN(id) AS keep_id FROM ldap_configs) AS ldap_singleton_keeper +); + +ALTER TABLE ldap_configs + ADD COLUMN singleton_key INT NOT NULL DEFAULT 1, + ADD CONSTRAINT ldap_configs_singleton_key_check CHECK (singleton_key = 1), + ADD UNIQUE KEY ldap_configs_singleton_key_key (singleton_key); diff --git a/crates/aether-data/adapters/mysql/migrations/20260831010000_add_proxy_node_tunnel_generation.sql b/crates/aether-data/adapters/mysql/migrations/20260831010000_add_proxy_node_tunnel_generation.sql new file mode 100644 index 000000000..5b5d64b6e --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260831010000_add_proxy_node_tunnel_generation.sql @@ -0,0 +1,9 @@ +ALTER TABLE proxy_nodes + ADD COLUMN tunnel_generation VARCHAR(64) NULL AFTER id; + +UPDATE proxy_nodes +SET tunnel_generation = UUID() +WHERE tunnel_generation IS NULL OR TRIM(tunnel_generation) = ''; + +ALTER TABLE proxy_nodes + MODIFY COLUMN tunnel_generation VARCHAR(64) NOT NULL; diff --git a/crates/aether-data/adapters/mysql/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql b/crates/aether-data/adapters/mysql/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql new file mode 100644 index 000000000..ce2a46cfc --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql @@ -0,0 +1,10 @@ +-- A proxy endpoint has one stable node identity across manual and tunnel +-- registrations. MySQL 8 performs ALTER TABLE atomically; if historical +-- duplicates exist this migration fails without choosing or deleting a row. +-- Diagnose with: +-- SELECT ip, port, COUNT(*) +-- FROM proxy_nodes +-- GROUP BY ip, port +-- HAVING COUNT(*) > 1; +ALTER TABLE proxy_nodes + ADD UNIQUE INDEX uq_proxy_node_ip_port (ip, port); diff --git a/crates/aether-data/adapters/mysql/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql b/crates/aether-data/adapters/mysql/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql new file mode 100644 index 000000000..1d0fdf821 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql @@ -0,0 +1,2 @@ +ALTER TABLE usage_counter_deltas + ADD COLUMN target_tunnel_generation VARCHAR(64) NULL; diff --git a/crates/aether-data/adapters/mysql/migrations/20260903000000_add_routing_group_sort_order.sql b/crates/aether-data/adapters/mysql/migrations/20260903000000_add_routing_group_sort_order.sql new file mode 100644 index 000000000..f33fd1096 --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260903000000_add_routing_group_sort_order.sql @@ -0,0 +1,3 @@ +ALTER TABLE routing_groups + ADD COLUMN sort_order BIGINT NOT NULL DEFAULT 0, + ADD KEY routing_groups_enabled_sort_idx (enabled, sort_order, name, id); diff --git a/crates/aether-data/adapters/mysql/src/auth.rs b/crates/aether-data/adapters/mysql/src/auth.rs index 74527704c..94615d813 100644 --- a/crates/aether-data/adapters/mysql/src/auth.rs +++ b/crates/aether-data/adapters/mysql/src/auth.rs @@ -3,9 +3,9 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::auth::{ AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository, - AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, - StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, - UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, + AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, + CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, + StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; use aether_data_contracts::DataLayerError; @@ -69,6 +69,21 @@ SELECT FROM api_keys "#; +const MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE video_tasks SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE `usage` SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id = ?", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", +]; + +const MYSQL_DELETE_API_KEY_DEPENDENTS_SQL: &[&str] = + &["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"]; + #[derive(Debug, Clone)] pub struct MysqlAuthApiKeyReadRepository { pool: MysqlPool, @@ -111,6 +126,17 @@ impl MysqlAuthApiKeyReadRepository { record: CreateApiKeyInsertRecord, ) -> Result, DataLayerError> { let now = current_unix_secs(); + let mut tx = self.pool.begin().await.map_sql_err()?; + let owner_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? AND is_deleted = 0 FOR UPDATE") + .bind(&record.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if owner_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } sqlx::query( r#" INSERT INTO api_keys ( @@ -120,7 +146,7 @@ INSERT INTO api_keys ( total_requests, total_tokens, total_cost_usd, is_standalone, created_at, updated_at ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&record.api_key_id) @@ -150,6 +176,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) &record.force_capabilities, "api_keys.force_capabilities", )?) + .bind(optional_json_to_string( + &record.feature_settings, + "api_keys.feature_settings", + )?) .bind(record.is_active) .bind(optional_i64_from_u64( record.expires_at_unix_secs, @@ -165,11 +195,25 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(record.is_standalone) .bind(now as i64) .bind(now as i64) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - - self.reload_export_by_id(&record.api_key_id).await + let reload_sql = format!("{EXPORT_COLUMNS}\nWHERE api_keys.id = ?\nLIMIT 1"); + let row = sqlx::query(&reload_sql) + .bind(&record.api_key_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::UnexpectedValue(format!( + "created api_keys row is missing: {}", + record.api_key_id + ))); + }; + let created = map_auth_api_key_export_row(&row)?; + tx.commit().await.map_sql_err()?; + Ok(Some(created)) } } @@ -186,6 +230,7 @@ struct CreateApiKeyInsertRecord { rate_limit: Option, concurrent_limit: Option, force_capabilities: Option, + feature_settings: Option, is_active: bool, expires_at_unix_secs: Option, auto_delete_on_expiry: bool, @@ -337,6 +382,7 @@ impl AuthApiKeyReadRepository for MysqlAuthApiKeyReadRepository { if user_ids.is_empty() { return Ok(AuthApiKeyExportSummary::default()); } + let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let mut builder = QueryBuilder::::new( r#" @@ -345,7 +391,7 @@ SELECT SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >= "#, ); - builder.push_bind(now_unix_secs as i64); + builder.push_bind(now_unix_secs); builder.push( r#") THEN 1 ELSE 0 END) AS active FROM api_keys @@ -429,6 +475,7 @@ WHERE id = ? rate_limit: Some(record.rate_limit), concurrent_limit: record.concurrent_limit, force_capabilities: record.force_capabilities, + feature_settings: record.feature_settings, is_active: record.is_active, expires_at_unix_secs: record.expires_at_unix_secs, auto_delete_on_expiry: record.auto_delete_on_expiry, @@ -457,6 +504,7 @@ WHERE id = ? rate_limit: record.rate_limit, concurrent_limit: record.concurrent_limit, force_capabilities: record.force_capabilities, + feature_settings: None, is_active: record.is_active, expires_at_unix_secs: record.expires_at_unix_secs, auto_delete_on_expiry: record.auto_delete_on_expiry, @@ -472,35 +520,41 @@ WHERE id = ? &self, record: UpdateUserApiKeyBasicRecord, ) -> Result, DataLayerError> { - let now = current_unix_secs() as i64; - sqlx::query( + self.update_user_api_key_basic_scoped(record, false).await + } + + async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + let result = sqlx::query( r#" UPDATE api_keys -SET name = COALESCE(?, name), - rate_limit = COALESCE(?, rate_limit), - concurrent_limit = COALESCE(?, concurrent_limit), - ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END, - updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 +SET key_encrypted = ? +WHERE BINARY id = BINARY ? + AND BINARY user_id = BINARY ? + AND BINARY key_hash = BINARY ? + AND is_standalone = ? + AND BINARY key_encrypted = BINARY ? "#, ) - .bind(record.name.as_deref()) - .bind(record.rate_limit) - .bind(record.concurrent_limit) - .bind(record.ip_rules.is_some()) - .bind(json_string_from_nested_string_list( - &record.ip_rules, - "api_keys.ip_rules", - )?) - .bind(now) - .bind(&record.api_key_id) - .bind(&record.user_id) + .bind(&mutation.key_encrypted) + .bind(&mutation.api_key_id) + .bind(&mutation.user_id) + .bind(&mutation.key_hash) + .bind(mutation.is_standalone) + .bind(&mutation.expected_key_encrypted) .execute(&self.pool) .await .map_sql_err()?; - self.reload_export_by_id(&record.api_key_id).await + Ok(result.rows_affected() == 1) + } + + async fn update_user_api_key_basic_if_unlocked( + &self, + record: UpdateUserApiKeyBasicRecord, + ) -> Result, DataLayerError> { + self.update_user_api_key_basic_scoped(record, true).await } async fn update_standalone_api_key_basic( @@ -511,7 +565,9 @@ WHERE id = ? sqlx::query( r#" UPDATE api_keys -SET name = COALESCE(?, name), +SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END, + name = CASE WHEN ? THEN ? ELSE name END, + force_capabilities = CASE WHEN ? THEN ? ELSE force_capabilities END, rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END, allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END, @@ -525,7 +581,15 @@ WHERE id = ? AND is_standalone = 1 "#, ) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted.as_deref()) + .bind(record.name_present) .bind(record.name.as_deref()) + .bind(record.force_capabilities.is_some()) + .bind(optional_json_to_string( + &record.force_capabilities.clone().flatten(), + "api_keys.force_capabilities", + )?) .bind(record.rate_limit_present) .bind(record.rate_limit) .bind(record.concurrent_limit_present) @@ -565,13 +629,143 @@ WHERE id = ? self.reload_export_by_id(&record.api_key_id).await } + async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + if restored.api_key_id != expected.api_key_id + || restored.user_id != expected.user_id + || restored.key_hash != expected.key_hash + || restored.is_standalone != expected.is_standalone + { + return Ok(false); + } + + let mut tx = self.pool.begin().await.map_sql_err()?; + let select_sql = format!("{EXPORT_COLUMNS} WHERE api_keys.id = ? LIMIT 1 FOR UPDATE"); + let row = sqlx::query(&select_sql) + .bind(&expected.api_key_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_auth_api_key_export_row(&row)?; + if current != *expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = ?, + name = ?, + allowed_providers = ?, + allowed_api_formats = ?, + allowed_models = ?, + ip_rules = ?, + rate_limit = ?, + concurrent_limit = ?, + force_capabilities = ?, + feature_settings = ?, + is_active = ?, + expires_at = ?, + auto_delete_on_expiry = ?, + total_requests = ?, + total_tokens = ?, + total_cost_usd = ?, + last_used_at = ?, + updated_at = ? +WHERE id = ? + AND user_id = ? + AND key_hash = ? + AND is_standalone = ? +"#, + ) + .bind(restored.key_encrypted.as_deref()) + .bind(restored.name.as_deref()) + .bind(json_string_from_string_list( + restored.allowed_providers.as_ref(), + "api_keys.allowed_providers", + )?) + .bind(json_string_from_string_list( + restored.allowed_api_formats.as_ref(), + "api_keys.allowed_api_formats", + )?) + .bind(json_string_from_string_list( + restored.allowed_models.as_ref(), + "api_keys.allowed_models", + )?) + .bind(json_string_from_string_list( + restored.ip_rules.as_ref(), + "api_keys.ip_rules", + )?) + .bind(restored.rate_limit) + .bind(restored.concurrent_limit) + .bind(optional_json_to_string( + &restored.force_capabilities, + "api_keys.force_capabilities", + )?) + .bind(optional_json_to_string( + &restored.feature_settings, + "api_keys.feature_settings", + )?) + .bind(restored.is_active) + .bind(optional_i64_from_u64( + restored.expires_at_unix_secs, + "api_keys.expires_at", + )?) + .bind(restored.auto_delete_on_expiry) + .bind(i64_from_u64( + restored.total_requests, + "api_keys.total_requests", + )?) + .bind(i64_from_u64( + restored.total_tokens, + "api_keys.total_tokens", + )?) + .bind(restored.total_cost_usd) + .bind(optional_i64_from_u64( + restored.last_used_at_unix_secs, + "api_keys.last_used_at", + )?) + .bind(current_unix_secs() as i64) + .bind(&restored.api_key_id) + .bind(&restored.user_id) + .bind(&restored.key_hash) + .bind(restored.is_standalone) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn set_user_api_key_active( &self, user_id: &str, api_key_id: &str, is_active: bool, ) -> Result, DataLayerError> { - self.set_active(api_key_id, Some(user_id), is_active, false) + self.set_active(api_key_id, Some(user_id), is_active, false, false) + .await + } + + async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + self.set_active(api_key_id, Some(user_id), is_active, false, true) .await } @@ -580,7 +774,8 @@ WHERE id = ? api_key_id: &str, is_active: bool, ) -> Result, DataLayerError> { - self.set_active(api_key_id, None, is_active, true).await + self.set_active(api_key_id, None, is_active, true, false) + .await } async fn set_user_api_key_locked( @@ -615,26 +810,23 @@ WHERE id = ? api_key_id: &str, allowed_providers: Option>, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET allowed_providers = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, + self.set_user_api_key_allowed_providers_scoped( + user_id, + api_key_id, + allowed_providers, + false, ) - .bind(json_string_from_string_list( - allowed_providers.as_ref(), - "api_keys.allowed_providers", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, DataLayerError> { + self.set_user_api_key_allowed_providers_scoped(user_id, api_key_id, allowed_providers, true) + .await } async fn set_user_api_key_force_capabilities( @@ -643,26 +835,28 @@ WHERE id = ? api_key_id: &str, force_capabilities: Option, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET force_capabilities = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, + self.set_user_api_key_force_capabilities_scoped( + user_id, + api_key_id, + force_capabilities, + false, + ) + .await + } + + async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, DataLayerError> { + self.set_user_api_key_force_capabilities_scoped( + user_id, + api_key_id, + force_capabilities, + true, ) - .bind(optional_json_to_string( - &force_capabilities, - "api_keys.force_capabilities", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await } async fn set_user_api_key_feature_settings( @@ -671,26 +865,18 @@ WHERE id = ? api_key_id: &str, feature_settings: Option, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET feature_settings = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, - ) - .bind(optional_json_to_string( - &feature_settings, - "api_keys.feature_settings", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) - .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await + self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, false) + .await + } + + async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, DataLayerError> { + self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, true) + .await } async fn set_api_key_usage_totals( @@ -700,6 +886,11 @@ WHERE id = ? total_tokens: u64, total_cost_usd: f64, ) -> Result, DataLayerError> { + if !total_cost_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "api_keys.total_cost_usd is not finite".to_string(), + )); + } sqlx::query( r#" UPDATE api_keys @@ -710,8 +901,8 @@ SET total_requests = ?, WHERE id = ? "#, ) - .bind(total_requests as i64) - .bind(total_tokens as i64) + .bind(i64_from_u64(total_requests, "api_keys.total_requests")?) + .bind(i64_from_u64(total_tokens, "api_keys.total_tokens")?) .bind(total_cost_usd) .bind(current_unix_secs() as i64) .bind(api_key_id) @@ -726,11 +917,21 @@ WHERE id = ? user_id: &str, api_key_id: &str, ) -> Result { - self.delete_api_key(api_key_id, Some(user_id), false).await + self.delete_api_key(api_key_id, Some(user_id), false, false) + .await + } + + async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + self.delete_api_key(api_key_id, Some(user_id), false, true) + .await } async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result { - self.delete_api_key(api_key_id, None, true).await + self.delete_api_key(api_key_id, None, true, false).await } async fn set_standalone_api_key_feature_settings( @@ -760,12 +961,65 @@ WHERE id = ? } impl MysqlAuthApiKeyReadRepository { + async fn update_user_api_key_basic_scoped( + &self, + record: UpdateUserApiKeyBasicRecord, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END, + name = CASE WHEN ? THEN ? ELSE name END, + rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, + concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END, + ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END, + feature_settings = CASE WHEN ? THEN ? ELSE feature_settings END, + updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted.as_deref()) + .bind(record.name_present) + .bind(record.name.as_deref()) + .bind(record.rate_limit_present) + .bind(record.rate_limit) + .bind(record.concurrent_limit_present) + .bind(record.concurrent_limit) + .bind(record.ip_rules.is_some()) + .bind(json_string_from_nested_string_list( + &record.ip_rules, + "api_keys.ip_rules", + )?) + .bind(record.feature_settings.is_some()) + .bind(optional_json_to_string( + &record.feature_settings.clone().flatten(), + "api_keys.feature_settings", + )?) + .bind(current_unix_secs() as i64) + .bind(&record.api_key_id) + .bind(&record.user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(&record.api_key_id).await + } + async fn set_active( &self, api_key_id: &str, user_id: Option<&str>, is_active: bool, is_standalone: bool, + require_unlocked: bool, ) -> Result, DataLayerError> { let mut builder = QueryBuilder::::new("UPDATE api_keys SET is_active = "); builder @@ -779,7 +1033,115 @@ impl MysqlAuthApiKeyReadRepository { if let Some(user_id) = user_id { builder.push(" AND user_id = ").push_bind(user_id); } - builder.build().execute(&self.pool).await.map_sql_err()?; + if require_unlocked { + builder.push(" AND is_locked = ").push_bind(false); + } + let result = builder.build().execute(&self.pool).await.map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_allowed_providers_scoped( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET allowed_providers = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(json_string_from_string_list( + allowed_providers.as_ref(), + "api_keys.allowed_providers", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_force_capabilities_scoped( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET force_capabilities = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(optional_json_to_string( + &force_capabilities, + "api_keys.force_capabilities", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_feature_settings_scoped( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET feature_settings = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(optional_json_to_string( + &feature_settings, + "api_keys.feature_settings", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } self.reload_export_by_id(api_key_id).await } @@ -788,22 +1150,76 @@ impl MysqlAuthApiKeyReadRepository { api_key_id: &str, user_id: Option<&str>, is_standalone: bool, + require_unlocked: bool, ) -> Result { - let mut builder = QueryBuilder::::new("DELETE FROM api_keys WHERE id = "); - builder - .push_bind(api_key_id) - .push(" AND is_standalone = ") - .push_bind(is_standalone); - if let Some(user_id) = user_id { - builder.push(" AND user_id = ").push_bind(user_id); - } - let rows_affected = builder - .build() - .execute(&self.pool) + let mut tx = self.pool.begin().await.map_sql_err()?; + let matching_api_key = if let Some(user_id) = user_id { + if require_unlocked { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 AND is_locked = 0 FOR UPDATE", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 FOR UPDATE", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + } + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND is_standalone = 1 FOR UPDATE", + ) + .bind(api_key_id) + .fetch_optional(&mut *tx) .await .map_sql_err()? - .rows_affected(); - Ok(rows_affected > 0) + }; + if matching_api_key.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + sqlx::query( + "UPDATE wallets SET status = 'disabled', updated_at = UNIX_TIMESTAMP() WHERE api_key_id = ? AND status <> 'disabled'", + ) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + for sql in MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + for sql in MYSQL_DELETE_API_KEY_DEPENDENTS_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + let result = sqlx::query("DELETE FROM api_keys WHERE id = ? AND is_standalone = ?") + .bind(api_key_id) + .bind(is_standalone) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) } } @@ -827,6 +1243,7 @@ async fn summarize_api_keys( is_standalone: bool, now_unix_secs: u64, ) -> Result { + let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let row = sqlx::query( r#" SELECT @@ -836,7 +1253,7 @@ FROM api_keys WHERE is_standalone = ? "#, ) - .bind(now_unix_secs as i64) + .bind(now_unix_secs) .bind(is_standalone) .fetch_one(pool) .await @@ -1037,7 +1454,36 @@ fn map_auth_api_key_export_row( #[cfg(test)] mod tests { - use super::MysqlAuthApiKeyReadRepository; + use super::{ + MysqlAuthApiKeyReadRepository, MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL, + MYSQL_DELETE_API_KEY_DEPENDENTS_SQL, + }; + + #[test] + fn api_key_delete_sql_preserves_ids_and_removes_private_snapshots() { + for table in [ + "request_candidates", + "video_tasks", + "`usage`", + "stats_daily_api_key", + ] { + assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| { + sql.starts_with(&format!("UPDATE {table} ")) + && sql.contains("SET api_key_name = NULL") + && sql.ends_with("WHERE api_key_id = ?") + })); + } + assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL + .iter() + .any(|sql| sql + .starts_with("UPDATE audit_logs SET description = 'deleted API key event'"))); + assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| sql + .starts_with("UPDATE payment_callbacks SET payload = NULL, error_message = NULL"))); + assert_eq!( + MYSQL_DELETE_API_KEY_DEPENDENTS_SQL, + &["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"] + ); + } #[tokio::test] async fn repository_builds_from_lazy_pool() { diff --git a/crates/aether-data/adapters/mysql/src/auth_modules.rs b/crates/aether-data/adapters/mysql/src/auth_modules.rs index c59e78c9b..44ccc226c 100644 --- a/crates/aether-data/adapters/mysql/src/auth_modules.rs +++ b/crates/aether-data/adapters/mysql/src/auth_modules.rs @@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::auth_modules::*; use aether_data_contracts::DataLayerError; -use aether_data_query::{push_eq, push_limit, WhereClause}; +use aether_data_query::{push_eq, WhereClause}; use crate::error::SqlResultExt; use crate::MysqlPool; @@ -35,6 +35,87 @@ SELECT FROM ldap_configs "#; +const UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL: &str = r#" +UPDATE ldap_configs +SET + server_url = ?, + bind_dn = ?, + base_dn = ?, + user_search_filter = ?, + username_attr = ?, + email_attr = ?, + display_name_attr = ?, + is_enabled = ?, + is_exclusive = ?, + use_starttls = ?, + connect_timeout = ?, + updated_at = GREATEST(updated_at + 1, ?) +WHERE singleton_key = 1 + AND server_url <=> ? + AND bind_dn <=> ? + AND BINARY bind_password_encrypted <=> BINARY ? + AND base_dn <=> ? + AND user_search_filter <=> ? + AND username_attr <=> ? + AND email_attr <=> ? + AND display_name_attr <=> ? + AND is_enabled <=> ? + AND is_exclusive <=> ? + AND use_starttls <=> ? + AND connect_timeout <=> ? +"#; + +const UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL: &str = r#" +UPDATE ldap_configs +SET + server_url = ?, + bind_dn = ?, + bind_password_encrypted = ?, + base_dn = ?, + user_search_filter = ?, + username_attr = ?, + email_attr = ?, + display_name_attr = ?, + is_enabled = ?, + is_exclusive = ?, + use_starttls = ?, + connect_timeout = ?, + updated_at = GREATEST(updated_at + 1, ?) +WHERE singleton_key = 1 + AND server_url <=> ? + AND bind_dn <=> ? + AND BINARY bind_password_encrypted <=> BINARY ? + AND base_dn <=> ? + AND user_search_filter <=> ? + AND username_attr <=> ? + AND email_attr <=> ? + AND display_name_attr <=> ? + AND is_enabled <=> ? + AND is_exclusive <=> ? + AND use_starttls <=> ? + AND connect_timeout <=> ? +"#; + +const INSERT_LDAP_CONFIG_SQL: &str = r#" +INSERT INTO ldap_configs ( + singleton_key, + server_url, + bind_dn, + bind_password_encrypted, + base_dn, + user_search_filter, + username_attr, + email_attr, + display_name_attr, + is_enabled, + is_exclusive, + use_starttls, + connect_timeout, + created_at, + updated_at +) VALUES (1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +"#; + #[derive(Debug, Clone)] pub struct MysqlAuthModuleReadRepository { pool: MysqlPool, @@ -72,8 +153,7 @@ async fn get_ldap_config( pool: &MysqlPool, ) -> Result, DataLayerError> { let mut builder = QueryBuilder::::new(LDAP_CONFIG_COLUMNS); - builder.push(" ORDER BY id ASC"); - push_limit(&mut builder, 1); + builder.push(" WHERE singleton_key = 1"); let row = builder.build().fetch_optional(pool).await.map_sql_err()?; row.as_ref().map(map_ldap_row).transpose() } @@ -106,97 +186,208 @@ impl AuthModuleReadRepository for MysqlAuthModuleRepository { #[async_trait] impl AuthModuleWriteRepository for MysqlAuthModuleRepository { - async fn upsert_ldap_config( + async fn compare_and_swap_ldap_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, DataLayerError> { + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result { + let persisted = + ldap_config_after_password_update(expected, replacement, bind_password_update)?; let now = now_unix_secs(); - let updated = sqlx::query( + let Some(expected) = expected else { + let insert = sqlx::query(INSERT_LDAP_CONFIG_SQL) + .bind(&persisted.server_url) + .bind(&persisted.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&persisted.base_dn) + .bind(persisted.user_search_filter.as_deref()) + .bind(persisted.username_attr.as_deref()) + .bind(persisted.email_attr.as_deref()) + .bind(persisted.display_name_attr.as_deref()) + .bind(persisted.is_enabled) + .bind(persisted.is_exclusive) + .bind(persisted.use_starttls) + .bind(persisted.connect_timeout) + .bind(now as i64) + .bind(now as i64) + .execute(&self.pool) + .await; + return match insert { + Ok(result) if result.rows_affected() == 1 => { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } + Ok(_) => Ok(CompareAndSwapLdapConfigResult::Conflict), + Err(error) + if error + .as_database_error() + .is_some_and(|error| error.is_unique_violation()) => + { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } + Err(error) => Err(DataLayerError::sql(error)), + }; + }; + + let rows_affected = match bind_password_update { + LdapBindPasswordUpdate::Preserve => { + sqlx::query(UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected() + } + LdapBindPasswordUpdate::Set(_) | LdapBindPasswordUpdate::Clear => { + sqlx::query(UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected() + } + }; + if rows_affected == 1 { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } else { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } + } + + async fn delete_ldap_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result { + let rows_affected = sqlx::query( r#" -UPDATE ldap_configs -SET - server_url = ?, - bind_dn = ?, - bind_password_encrypted = ?, - base_dn = ?, - user_search_filter = ?, - username_attr = ?, - email_attr = ?, - display_name_attr = ?, - is_enabled = ?, - is_exclusive = ?, - use_starttls = ?, - connect_timeout = ?, - updated_at = ? -WHERE id = ( - SELECT id FROM ( - SELECT id - FROM ldap_configs - ORDER BY id ASC - LIMIT 1 - ) selected_ldap_config -) +DELETE FROM ldap_configs +WHERE singleton_key = 1 + AND server_url <=> ? + AND bind_dn <=> ? + AND BINARY bind_password_encrypted <=> BINARY ? + AND base_dn <=> ? + AND user_search_filter <=> ? + AND username_attr <=> ? + AND email_attr <=> ? + AND display_name_attr <=> ? + AND is_enabled <=> ? + AND is_exclusive <=> ? + AND use_starttls <=> ? + AND connect_timeout <=> ? "#, ) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) .execute(&self.pool) .await - .map_sql_err()?; - - if updated.rows_affected() == 0 { - sqlx::query( - r#" -INSERT INTO ldap_configs ( - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter, - username_attr, - email_attr, - display_name_attr, - is_enabled, - is_exclusive, - use_starttls, - connect_timeout, - created_at, - updated_at -) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) -"#, - ) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .bind(now as i64) - .bind(now as i64) - .execute(&self.pool) - .await - .map_sql_err()?; - } - - self.get_ldap_config().await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) } + + async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + let rows_affected = sqlx::query( + r#" +UPDATE ldap_configs +SET bind_password_encrypted = ?, updated_at = GREATEST(updated_at + 1, ?) +WHERE singleton_key = 1 + AND BINARY bind_password_encrypted = BINARY ? +"#, + ) + .bind(replacement) + .bind(now_unix_secs() as i64) + .bind(expected) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } +} + +fn ldap_config_after_password_update( + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, +) -> Result { + let bind_password_encrypted = match bind_password_update { + LdapBindPasswordUpdate::Preserve => expected + .ok_or_else(|| { + DataLayerError::InvalidConfiguration( + "LDAP bind password cannot be preserved while creating the singleton" + .to_string(), + ) + })? + .bind_password_encrypted + .clone(), + LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()), + LdapBindPasswordUpdate::Clear => None, + }; + Ok(StoredLdapModuleConfig { + bind_password_encrypted, + ..replacement.clone() + }) } fn now_unix_secs() -> u64 { diff --git a/crates/aether-data/adapters/mysql/src/background_tasks.rs b/crates/aether-data/adapters/mysql/src/background_tasks.rs index 88e68f85d..e4b8479a4 100644 --- a/crates/aether-data/adapters/mysql/src/background_tasks.rs +++ b/crates/aether-data/adapters/mysql/src/background_tasks.rs @@ -209,8 +209,9 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository { impl BackgroundTaskWriteRepository for MysqlBackgroundTaskRepository { async fn upsert_run( &self, - run: UpsertBackgroundTaskRun, + mut run: UpsertBackgroundTaskRun, ) -> Result { + run.sanitize_for_persistence(); run.validate()?; sqlx::query( r#" @@ -309,8 +310,9 @@ ON DUPLICATE KEY UPDATE async fn upsert_event( &self, - event: UpsertBackgroundTaskEvent, + mut event: UpsertBackgroundTaskEvent, ) -> Result { + event.sanitize_for_persistence(); event.validate()?; sqlx::query( r#" @@ -358,7 +360,7 @@ fn map_run_row(row: &MySqlRow) -> Result = row.try_get("finished_at_unix_secs").map_sql_err()?; let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?; - Ok(StoredBackgroundTaskRun { + let mut run = StoredBackgroundTaskRun { id: row.try_get("id").map_sql_err()?, task_key: row.try_get("task_key").map_sql_err()?, kind: BackgroundTaskKind::from_database(&kind)?, @@ -381,12 +383,14 @@ fn map_run_row(row: &MySqlRow) -> Result Result { let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?; - Ok(StoredBackgroundTaskEvent { + let mut event = StoredBackgroundTaskEvent { id: row.try_get("id").map_sql_err()?, run_id: row.try_get("run_id").map_sql_err()?, event_type: row.try_get("event_type").map_sql_err()?, @@ -396,7 +400,9 @@ fn map_event_row(row: &MySqlRow) -> Result Result { diff --git a/crates/aether-data/adapters/mysql/src/billing.rs b/crates/aether-data/adapters/mysql/src/billing.rs index 0b38c435d..fa3705087 100644 --- a/crates/aether-data/adapters/mysql/src/billing.rs +++ b/crates/aether-data/adapters/mysql/src/billing.rs @@ -4,8 +4,9 @@ use sqlx::{mysql::MySqlRow, Row}; use aether_data_contracts::repository::billing::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; use aether_data_contracts::DataLayerError; @@ -637,6 +638,157 @@ LIMIT 1 .transpose() } + async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE payment_gateway_configs +SET merchant_key_encrypted = ? +WHERE provider = ? + AND BINARY merchant_key_encrypted = BINARY ? + "#, + ) + .bind(&update.merchant_key_encrypted) + .bind(update.provider.trim().to_ascii_lowercase()) + .bind(&update.expected_merchant_key_encrypted) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_swap_payment_gateway_config( + &self, + mutation: &PaymentGatewayConfigCasWriteInput, + ) -> Result, DataLayerError> { + let input = &mutation.input; + let provider = input.provider.trim().to_ascii_lowercase(); + let now = current_unix_secs_i64(); + let mut tx = self.pool.begin().await.map_sql_err()?; + + if mutation.expected_existing { + let current = sqlx::query( + r#" +SELECT merchant_key_encrypted +FROM payment_gateway_configs +WHERE provider = ? +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(&provider) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let current_secret = match current.as_ref() { + Some(row) => row + .try_get::, _>("merchant_key_encrypted") + .map_sql_err()?, + None => { + tx.rollback().await.map_sql_err()?; + return Ok(AdminBillingMutationOutcome::NotFound); + } + }; + if current_secret != mutation.expected_merchant_key_encrypted { + tx.rollback().await.map_sql_err()?; + return Ok(AdminBillingMutationOutcome::NotFound); + } + + sqlx::query( + r#" +UPDATE payment_gateway_configs +SET + enabled = ?, + endpoint_url = ?, + callback_base_url = ?, + merchant_id = ?, + merchant_key_encrypted = CASE + WHEN ? THEN merchant_key_encrypted + ELSE ? + END, + pay_currency = ?, + usd_exchange_rate = ?, + min_recharge_usd = ?, + channels_json = ?, + updated_at = ? +WHERE provider = ? + "#, + ) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.preserve_existing_secret) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(json_to_string(&input.channels_json)?) + .bind(now) + .bind(&provider) + .execute(&mut *tx) + .await + .map_sql_err()?; + } else { + let inserted = sqlx::query( + r#" +INSERT INTO payment_gateway_configs ( + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd, + channels_json, created_at, updated_at +) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + "#, + ) + .bind(&provider) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(json_to_string(&input.channels_json)?) + .bind(now) + .bind(now) + .execute(&mut *tx) + .await; + if let Err(err) = inserted { + let unique = matches!( + &err, + sqlx::Error::Database(database_error) if database_error.is_unique_violation() + ); + tx.rollback().await.map_sql_err()?; + if unique { + return Ok(AdminBillingMutationOutcome::NotFound); + } + return Err(DataLayerError::sql(err)); + } + } + + let row = sqlx::query( + r#" +SELECT + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd, + channels_json, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs +FROM payment_gateway_configs +WHERE provider = ? +LIMIT 1 + "#, + ) + .bind(&provider) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let record = map_payment_gateway_config_mysql(&row)?; + tx.commit().await.map_sql_err()?; + Ok(AdminBillingMutationOutcome::Applied(record)) + } + async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -920,6 +1072,39 @@ ORDER BY expires_at ASC, created_at ASC )) } + async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, DataLayerError> { + let now = current_unix_secs_i64(); + let result = sqlx::query( + r#" +UPDATE user_plan_entitlements +SET status = 'revoked', + expires_at = LEAST(expires_at, ?), + updated_at = ? +WHERE id = ? + AND user_id = ? + AND status = 'active' + AND expires_at > ? + "#, + ) + .bind(now) + .bind(now) + .bind(entitlement_id) + .bind(user_id) + .bind(now) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + Ok(AdminBillingMutationOutcome::NotFound) + } else { + Ok(AdminBillingMutationOutcome::Applied(())) + } + } + async fn find_user_daily_quota_availability( &self, user_id: &str, @@ -927,13 +1112,19 @@ ORDER BY expires_at ASC, created_at ASC let now_unix_secs = current_unix_secs_i64(); let rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = ? - AND status = 'active' - AND starts_at <= ? - AND expires_at > ? -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = ? + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= ? + AND user_plan_entitlements.expires_at > ? +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC "#, ) .bind(user_id) @@ -948,9 +1139,13 @@ ORDER BY expires_at ASC, created_at ASC, id ASC let entitlement_id: String = row.try_get("id").map_sql_err()?; let entitlements = parse_json(row.try_get("entitlements_snapshot").ok().flatten())? .unwrap_or_else(|| serde_json::json!([])); + let plan_entitlements = + parse_json(row.try_get("plan_entitlements_json").ok().flatten())? + .unwrap_or_else(|| serde_json::json!([])); grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -1180,6 +1375,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -1205,15 +1401,27 @@ fn daily_quota_grants_from_entitlement( .and_then(serde_json::Value::as_str), now, )?, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + fn read_count_mysql(row: &MySqlRow) -> Result { Ok(row.try_get::("total").map_sql_err()?.max(0) as u64) } diff --git a/crates/aether-data/adapters/mysql/src/candidate_selection.rs b/crates/aether-data/adapters/mysql/src/candidate_selection.rs index 9b45fb690..5bdfe04a0 100644 --- a/crates/aether-data/adapters/mysql/src/candidate_selection.rs +++ b/crates/aether-data/adapters/mysql/src/candidate_selection.rs @@ -832,12 +832,12 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result, + field_name: &str, +) -> Result>, DataLayerError> { + let Some(raw) = raw else { + return Ok(None); + }; + let value = serde_json::from_str::(&raw).map_err(|err| { + DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}")) + })?; + parse_key_policy_string_list_value(&value, field_name) +} + +fn parse_key_policy_string_list_value( + value: &serde_json::Value, + field_name: &str, +) -> Result>, DataLayerError> { + match value { + serde_json::Value::Null => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains JSON null; use SQL NULL for an unset policy" + ))), + serde_json::Value::Array(array) => { + parse_key_policy_string_list_array(array, field_name).map(Some) + } + serde_json::Value::String(raw) => parse_embedded_key_policy_string_list(raw, field_name), + _ => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} is not a JSON array" + ))), + } +} + +fn parse_embedded_key_policy_string_list( + raw: &str, + field_name: &str, +) -> Result>, DataLayerError> { + let raw = raw.trim(); + if raw.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty string" + ))); + } + if raw.eq_ignore_ascii_case("null") { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains stringified JSON null; use SQL NULL for an unset policy" + ))); + } + + if let Ok(decoded) = serde_json::from_str::(raw) { + return parse_key_policy_string_list_value(&decoded, field_name); + } + + Ok(Some(vec![raw.to_string()])) +} + +fn parse_key_policy_string_list_array( + array: &[serde_json::Value], + field_name: &str, +) -> Result, DataLayerError> { + let mut items = Vec::with_capacity(array.len()); + for item in array { + let Some(item) = item.as_str() else { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains a non-string item" + ))); + }; + let item = item.trim(); + if item.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty item" + ))); + } + items.push(item.to_string()); + } + Ok(items) +} + fn parse_string_list_value( value: &serde_json::Value, field_name: &str, @@ -1108,7 +1184,8 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec { #[cfg(test)] mod tests { use super::{ - api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query, + api_format_page_query, parse_stored_key_policy_string_list, + pool_key_group_by_key_ids_query, pool_key_group_query, provider_model_mapping_api_format_covers, push_key_auth_channel_filter, requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator, MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT, @@ -1132,6 +1209,25 @@ mod tests { assert!(sql.contains("LIMIT ? OFFSET ?")); } + #[test] + fn malformed_key_policy_never_degrades_to_unrestricted() { + for raw in ["null", "\"null\"", "\"\"", "[\"openai:chat\",null]"] { + assert!(parse_stored_key_policy_string_list( + Some(raw.to_string()), + "provider_api_keys.api_formats", + ) + .is_err()); + } + assert_eq!( + parse_stored_key_policy_string_list( + Some("[\"openai:chat\"]".to_string()), + "provider_api_keys.api_formats", + ) + .expect("valid key policy should parse"), + Some(vec!["openai:chat".to_string()]) + ); + } + #[test] fn requested_model_page_query_filters_and_pages_before_fetch() { let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256); diff --git a/crates/aether-data/adapters/mysql/src/candidates.rs b/crates/aether-data/adapters/mysql/src/candidates.rs index b90cdcabc..37944e3c6 100644 --- a/crates/aether-data/adapters/mysql/src/candidates.rs +++ b/crates/aether-data/adapters/mysql/src/candidates.rs @@ -216,8 +216,9 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository { impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository { async fn upsert( &self, - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, ) -> Result { + candidate.sanitize_for_persistence(); candidate.validate()?; let mut tx = self.pool.begin().await.map_sql_err()?; match upsert_candidate_in_transaction(&mut tx, candidate).await { @@ -234,12 +235,13 @@ impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository { async fn upsert_many( &self, - candidates: Vec, + mut candidates: Vec, ) -> Result { if candidates.is_empty() { return Ok(0); } - for candidate in &candidates { + for candidate in &mut candidates { + candidate.sanitize_for_persistence(); candidate.validate()?; } @@ -423,26 +425,8 @@ ON DUPLICATE KEY UPDATE THEN status_code ELSE COALESCE(VALUES(status_code), status_code) END, - error_type = CASE - WHEN status IN ('success', 'failed', 'cancelled', 'skipped') - AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming') - THEN error_type - WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused') - THEN error_type - WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending') - THEN error_type - ELSE COALESCE(VALUES(error_type), error_type) - END, - error_message = CASE - WHEN status IN ('success', 'failed', 'cancelled', 'skipped') - AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming') - THEN error_message - WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused') - THEN error_message - WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending') - THEN error_message - ELSE COALESCE(VALUES(error_message), error_message) - END, + error_type = VALUES(error_type), + error_message = NULL, latency_ms = CASE WHEN status IN ('success', 'failed', 'cancelled', 'skipped') AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming') @@ -524,9 +508,10 @@ fn push_endpoint_in_clause<'args>( } fn merge_candidate( - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, existing: Option, ) -> Result { + candidate.sanitize_for_persistence(); let preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| { request_candidate_lifecycle_would_regress(value.status, candidate.status) }); @@ -538,15 +523,11 @@ fn merge_candidate( } else { candidate.status }; - let created_at_unix_ms = candidate - .created_at_unix_ms + let created_at_unix_ms = existing + .as_ref() + .map(|value| value.created_at_unix_ms) .filter(|value| *value > 1000) - .or_else(|| { - existing - .as_ref() - .map(|value| value.created_at_unix_ms) - .filter(|value| *value > 1000) - }) + .or_else(|| candidate.created_at_unix_ms.filter(|value| *value > 1000)) .or(candidate.started_at_unix_ms) .or(candidate.finished_at_unix_ms) .unwrap_or_else(current_unix_ms); @@ -561,35 +542,36 @@ fn merge_candidate( StoredRequestCandidate::new( id, candidate.request_id, - candidate - .user_id - .or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())), - candidate - .api_key_id - .or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())), - candidate - .username - .or_else(|| existing.as_ref().and_then(|value| value.username.clone())), - candidate.api_key_name.or_else(|| { - existing - .as_ref() - .and_then(|value| value.api_key_name.clone()) - }), + existing + .as_ref() + .and_then(|value| value.user_id.clone()) + .or(candidate.user_id), + existing + .as_ref() + .and_then(|value| value.api_key_id.clone()) + .or(candidate.api_key_id), + existing + .as_ref() + .and_then(|value| value.username.clone()) + .or(candidate.username), + existing + .as_ref() + .and_then(|value| value.api_key_name.clone()) + .or(candidate.api_key_name), to_i32(candidate.candidate_index)?, to_i32(candidate.retry_index)?, - candidate.provider_id.or_else(|| { - existing - .as_ref() - .and_then(|value| value.provider_id.clone()) - }), - candidate.endpoint_id.or_else(|| { - existing - .as_ref() - .and_then(|value| value.endpoint_id.clone()) - }), - candidate - .key_id - .or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())), + existing + .as_ref() + .and_then(|value| value.provider_id.clone()) + .or(candidate.provider_id), + existing + .as_ref() + .and_then(|value| value.endpoint_id.clone()) + .or(candidate.endpoint_id), + existing + .as_ref() + .and_then(|value| value.key_id.clone()) + .or(candidate.key_id), merged_status, candidate.skip_reason.or_else(|| { existing @@ -617,17 +599,7 @@ fn merge_candidate( .error_type .or_else(|| existing.as_ref().and_then(|value| value.error_type.clone())) }, - if preserve_existing_lifecycle { - existing - .as_ref() - .and_then(|value| value.error_message.clone()) - } else { - candidate.error_message.or_else(|| { - existing - .as_ref() - .and_then(|value| value.error_message.clone()) - }) - }, + None, if preserve_existing_lifecycle { match existing.as_ref().and_then(|value| value.latency_ms) { Some(value) => Some(to_i32_u64(value)?), @@ -657,9 +629,10 @@ fn merge_candidate( .and_then(|value| value.required_capabilities.clone()) }), u64_to_i64(created_at_unix_ms, "request candidate created_at")?, - candidate - .started_at_unix_ms - .or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms)) + existing + .as_ref() + .and_then(|value| value.started_at_unix_ms) + .or(candidate.started_at_unix_ms) .map(|value| u64_to_i64(value, "request candidate started_at")) .transpose()?, if preserve_existing_lifecycle { @@ -915,7 +888,7 @@ mod tests { &request_id, "initial", RequestCandidateStatus::Pending, - Some(json!({"initial": true})), + Some(json!({"gateway_execution_runtime": true})), 3_000_000, ); initial.is_cached = Some(false); @@ -923,6 +896,18 @@ mod tests { .upsert(initial) .await .expect("initial candidate should insert"); + sqlx::query( + "UPDATE request_candidates SET skip_reason = ?, error_type = ?, error_message = ?, extra_data = ?, required_capabilities = ? WHERE request_id = ?", + ) + .bind("legacy skip reason with tenant-secret") + .bind("legacy_error_type_with_token") + .bind("Bearer legacy-secret") + .bind(r#"{"gateway_execution_runtime":true,"request_body":{"password":"secret"}}"#) + .bind(r#"{"streaming":true,"internal_capability":"secret"}"#) + .bind(&request_id) + .execute(&pool) + .await + .expect("legacy diagnostics should be injected for the conflict test"); const WRITERS: usize = 8; let barrier = std::sync::Arc::new(tokio::sync::Barrier::new(WRITERS)); @@ -937,13 +922,22 @@ mod tests { } else { RequestCandidateStatus::Streaming }; - let mut extra_data = serde_json::Map::new(); - extra_data.insert(format!("writer_{writer}"), json!(writer)); + let extra_data = match writer { + 0 => json!({"stream_completed": true}), + 1 => json!({"cache_1h": true}), + 2 => json!({"first_byte_time_ms": 2}), + 3 => json!({"pool_key_index": 3}), + 4 => json!({"priority_slot": 4}), + 5 => json!({"ranking_index": 5}), + 6 => json!({"phase": "provider_request"}), + 7 => json!({"provider_api_format": "openai:responses"}), + _ => unreachable!("writer index is bounded by WRITERS"), + }; let mut candidate = sample_upsert( &request_id, format!("writer-{writer}").as_str(), status, - Some(serde_json::Value::Object(extra_data)), + Some(extra_data), 3_100_000 + u64::try_from(writer).expect("writer index should fit") * 10, ); if writer != 0 { @@ -970,25 +964,60 @@ mod tests { assert_eq!(candidate.status, RequestCandidateStatus::Success); assert_eq!(candidate.latency_ms, Some(123)); assert_eq!(candidate.finished_at_unix_ms, Some(3_100_002)); - let extra_data = candidate - .extra_data - .as_ref() - .and_then(serde_json::Value::as_object) - .expect("merged extra data should be an object"); - assert_eq!(extra_data.get("initial"), Some(&json!(true))); - for writer in 0..WRITERS { - assert_eq!( - extra_data.get(format!("writer_{writer}").as_str()), - Some(&json!(writer)) - ); - } + assert_eq!( + candidate.extra_data, + Some(json!({ + "cache_1h": true, + "first_byte_time_ms": 2, + "gateway_execution_runtime": true, + "phase": "provider_request", + "pool_key_index": 3, + "priority_slot": 4, + "provider_api_format": "openai:responses", + "ranking_index": 5, + "stream_completed": true + })) + ); + let raw = sqlx::query( + "SELECT skip_reason, error_type, error_message, extra_data, required_capabilities FROM request_candidates WHERE request_id = ?", + ) + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("raw candidate diagnostics should load"); + assert!( + sqlx::Row::try_get::, _>(&raw, "error_message") + .expect("error_message should decode") + .is_none() + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "skip_reason") + .expect("skip_reason should decode") + .as_deref(), + Some("unclassified_skip") + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "error_type") + .expect("error_type should decode") + .as_deref(), + Some("unclassified_error") + ); + let raw_extra = sqlx::Row::try_get::, _>(&raw, "extra_data") + .expect("extra_data should decode") + .and_then(|value| serde_json::from_str::(&value).ok()); + assert_eq!(raw_extra, candidate.extra_data); + let raw_capabilities = + sqlx::Row::try_get::, _>(&raw, "required_capabilities") + .expect("required_capabilities should decode") + .and_then(|value| serde_json::from_str::(&value).ok()); + assert_eq!(raw_capabilities, Some(json!({"streaming": true}))); let batch_request_id = format!("candidate-batch-{}", uuid::Uuid::new_v4()); let mut pending = sample_upsert( &batch_request_id, "batch-first", RequestCandidateStatus::Pending, - Some(json!({"pending": true})), + Some(json!({"gateway_execution_runtime": true})), 4_000_000, ); pending.is_cached = Some(false); @@ -996,7 +1025,7 @@ mod tests { &batch_request_id, "batch-second", RequestCandidateStatus::Streaming, - Some(json!({"streaming": true})), + Some(json!({"stream_completed": true})), 4_000_100, ); streaming.is_cached = None; @@ -1004,7 +1033,7 @@ mod tests { &batch_request_id, "batch-third", RequestCandidateStatus::Success, - Some(json!({"success": true})), + Some(json!({"cache_1h": true})), 4_000_200, ); success.is_cached = Some(true); @@ -1012,7 +1041,7 @@ mod tests { &batch_request_id, "batch-fourth", RequestCandidateStatus::Pending, - Some(json!({"late": true})), + Some(json!({"first_byte_time_ms": 42})), 4_000_300, ); late_pending.is_cached = None; @@ -1038,10 +1067,10 @@ mod tests { assert_eq!( batch_candidates[0].extra_data, Some(json!({ - "pending": true, - "streaming": true, - "success": true, - "late": true + "cache_1h": true, + "first_byte_time_ms": 42, + "gateway_execution_runtime": true, + "stream_completed": true })) ); @@ -1082,7 +1111,7 @@ mod tests { } #[test] - fn merge_candidate_keeps_terminal_status_when_streaming_arrives_late() { + fn merge_candidate_preserves_first_identity_and_terminal_fact() { let existing = StoredRequestCandidate::new( "candidate-1".to_string(), "request-1".to_string(), @@ -1103,7 +1132,7 @@ mod tests { None, Some(123), None, - Some(serde_json::json!({"terminal": true})), + Some(serde_json::json!({"stream_completed": true})), None, 1_000, Some(1_001), @@ -1115,24 +1144,24 @@ mod tests { UpsertRequestCandidateRecord { id: "candidate-late".to_string(), request_id: "request-1".to_string(), - user_id: Some("user-1".to_string()), - api_key_id: Some("key-1".to_string()), - username: None, - api_key_name: None, + user_id: Some("attacker-user".to_string()), + api_key_id: Some("attacker-api-key".to_string()), + username: Some("mallory".to_string()), + api_key_name: Some("attacker-key".to_string()), candidate_index: 0, retry_index: 0, - provider_id: Some("provider-1".to_string()), - endpoint_id: Some("endpoint-1".to_string()), - key_id: Some("provider-key-1".to_string()), - status: RequestCandidateStatus::Streaming, + provider_id: Some("attacker-provider".to_string()), + endpoint_id: Some("attacker-endpoint".to_string()), + key_id: Some("attacker-provider-key".to_string()), + status: RequestCandidateStatus::Failed, skip_reason: None, is_cached: Some(false), status_code: Some(200), error_type: None, - error_message: None, + error_message: Some("Bearer secret-token".to_string()), latency_ms: Some(9_999), concurrent_requests: None, - extra_data: Some(serde_json::json!({"late": true})), + extra_data: Some(serde_json::json!({"gateway_execution_runtime": true})), required_capabilities: None, created_at_unix_ms: Some(1_050), started_at_unix_ms: Some(1_051), @@ -1144,11 +1173,20 @@ mod tests { assert_eq!(merged.id, "candidate-1"); assert_eq!(merged.status, RequestCandidateStatus::Success); + assert_eq!(merged.user_id.as_deref(), Some("user-1")); + assert_eq!(merged.api_key_id.as_deref(), Some("key-1")); + assert_eq!(merged.provider_id.as_deref(), Some("provider-1")); + assert_eq!(merged.endpoint_id.as_deref(), Some("endpoint-1")); + assert_eq!(merged.key_id.as_deref(), Some("provider-key-1")); + assert!(merged.error_message.is_none()); assert_eq!(merged.latency_ms, Some(123)); assert_eq!(merged.finished_at_unix_ms, Some(1_123)); assert_eq!( merged.extra_data, - Some(serde_json::json!({"terminal": true, "late": true})) + Some(serde_json::json!({ + "gateway_execution_runtime": true, + "stream_completed": true + })) ); } diff --git a/crates/aether-data/adapters/mysql/src/gemini_file_mappings.rs b/crates/aether-data/adapters/mysql/src/gemini_file_mappings.rs index 2d2cbb098..13baa4492 100644 --- a/crates/aether-data/adapters/mysql/src/gemini_file_mappings.rs +++ b/crates/aether-data/adapters/mysql/src/gemini_file_mappings.rs @@ -12,6 +12,19 @@ use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, Whe use crate::error::SqlResultExt; use crate::MysqlPool; +const OWNER_GUARDED_UPSERT_SQL: &str = r#" +INSERT INTO gemini_file_mappings ( + id, file_name, key_id, user_id, display_name, mime_type, source_hash, + created_at, expires_at +) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) +ON DUPLICATE KEY UPDATE + display_name = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(display_name), display_name), + mime_type = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(mime_type), mime_type), + source_hash = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(source_hash), source_hash), + expires_at = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(expires_at), expires_at) +"#; + #[derive(Debug, Clone)] pub struct MysqlGeminiFileMappingRepository { pool: MysqlPool, @@ -63,6 +76,85 @@ LIMIT 1 row.as_ref().map(map_row).transpose() } + async fn find_active_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at AS created_at_unix_ms, + expires_at AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE BINARY file_name = BINARY ? + AND BINARY user_id = BINARY ? + AND expires_at > ? +LIMIT 1 +"#, + ) + .bind(file_name) + .bind(user_id) + .bind(i64_from_u64( + now_unix_secs, + "gemini_file_mappings.owner_read_now", + )?) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + + row.as_ref().map(map_row).transpose() + } + + async fn find_active_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at AS created_at_unix_ms, + expires_at AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE BINARY file_name = BINARY ? + AND BINARY key_id = BINARY ? + AND BINARY user_id = BINARY ? + AND expires_at > ? +LIMIT 1 +"#, + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .bind(i64_from_u64( + now_unix_secs, + "gemini_file_mappings.owner_read_now", + )?) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + + row.as_ref().map(map_row).transpose() + } + async fn list_mappings( &self, query: &GeminiFileMappingListQuery, @@ -185,6 +277,60 @@ ON DUPLICATE KEY UPDATE self.reload_by_file_name(&record.file_name).await } + async fn upsert_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, DataLayerError> { + record.validate()?; + let mut transaction = self.pool.begin().await.map_sql_err()?; + sqlx::query(OWNER_GUARDED_UPSERT_SQL) + .bind(&record.id) + .bind(&record.file_name) + .bind(&record.key_id) + .bind(&record.user_id) + .bind(&record.display_name) + .bind(&record.mime_type) + .bind(&record.source_hash) + .bind(current_unix_secs() as i64) + .bind(i64_from_u64( + record.expires_at_unix_secs, + "gemini_file_mappings.expires_at", + )?) + .execute(&mut *transaction) + .await + .map_sql_err()?; + + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at AS created_at_unix_ms, + expires_at AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE file_name = ? +LIMIT 1 +FOR UPDATE +"#, + ) + .bind(&record.file_name) + .fetch_one(&mut *transaction) + .await + .map_sql_err()?; + let stored = map_row(&row)?; + let owner_matches = stored.file_name == record.file_name + && stored.key_id == record.key_id + && stored.user_id == record.user_id; + transaction.commit().await.map_sql_err()?; + + Ok(owner_matches.then_some(stored)) + } + async fn delete_by_file_name(&self, file_name: &str) -> Result { let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?") .bind(file_name) @@ -195,6 +341,43 @@ ON DUPLICATE KEY UPDATE Ok(rows_affected > 0) } + async fn delete_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + let rows_affected = + sqlx::query( + "DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY user_id = BINARY ?", + ) + .bind(file_name) + .bind(user_id) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected > 0) + } + + async fn delete_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + let rows_affected = sqlx::query( + "DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY key_id = BINARY ? AND BINARY user_id = BINARY ?", + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected > 0) + } + async fn delete_by_id( &self, mapping_id: &str, @@ -282,6 +465,11 @@ fn apply_list_filters( where_clause: &mut WhereClause, query: &GeminiFileMappingListQuery, ) { + if let Some(user_id) = query.user_id.as_deref() { + where_clause.push_next(builder); + builder.push("BINARY user_id = BINARY "); + builder.push_bind(user_id.to_string()); + } if !query.include_expired { where_clause.push_next(builder); builder.push("expires_at > "); @@ -343,13 +531,17 @@ fn map_row(row: &MySqlRow) -> Result { #[cfg(test)] mod tests { - use super::{build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository}; + use super::{ + build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository, + OWNER_GUARDED_UPSERT_SQL, + }; use aether_data_contracts::repository::gemini_file_mappings::GeminiFileMappingListQuery; use sqlx::Execute; #[test] fn list_query_uses_shared_mysql_filter_and_pagination_rendering() { let query = GeminiFileMappingListQuery { + user_id: Some("user-1".to_string()), include_expired: false, search: Some(" Report ".to_string()), offset: 5, @@ -359,7 +551,9 @@ mod tests { let mut count = build_list_count_query(&query); let count_sql = count.build().sql().to_string(); - assert!(count_sql.contains(" WHERE expires_at > ? AND (LOWER(file_name) LIKE ?")); + assert!(count_sql.contains( + " WHERE BINARY user_id = BINARY ? AND expires_at > ? AND (LOWER(file_name) LIKE ?" + )); assert!(!count_sql.contains("WHERE 1=1")); let mut rows = build_list_rows_query(&query); @@ -368,6 +562,12 @@ mod tests { assert!(rows_sql.contains(" ORDER BY created_at DESC, file_name ASC LIMIT ? OFFSET ?")); } + #[test] + fn owner_guarded_upsert_uses_exact_file_key_and_user_identity() { + let identity = "BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id)"; + assert_eq!(OWNER_GUARDED_UPSERT_SQL.matches(identity).count(), 4); + } + #[tokio::test] async fn repository_builds_from_lazy_pool() { let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with( diff --git a/crates/aether-data/adapters/mysql/src/management_tokens.rs b/crates/aether-data/adapters/mysql/src/management_tokens.rs index 22e7be584..aff5831e2 100644 --- a/crates/aether-data/adapters/mysql/src/management_tokens.rs +++ b/crates/aether-data/adapters/mysql/src/management_tokens.rs @@ -2,10 +2,10 @@ use async_trait::async_trait; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser, - UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, UpdateManagementTokenRecord, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause}; @@ -38,8 +38,147 @@ impl MysqlManagementTokenRepository { .map_sql_err()?; row.as_ref().map(map_token_row).transpose() } + + async fn get_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + let mut builder = QueryBuilder::::new(TOKEN_COLUMNS); + let mut where_clause = WhereClause::new(); + push_eq(&mut builder, &mut where_clause, "id", token_id.to_string()); + push_optional_eq( + &mut builder, + &mut where_clause, + "user_id", + expected_user_id.map(ToOwned::to_owned), + ); + push_limit(&mut builder, 1); + let row = builder + .build() + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_token_row).transpose() + } + + async fn update_management_token_scoped( + &self, + record: &UpdateManagementTokenRecord, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + record.validate()?; + let allowed_ips = json_to_string(record.allowed_ips.as_ref())?; + let permissions = json_to_string(record.permissions.as_ref())?; + let now = now_unix_secs(); + + sqlx::query(UPDATE_MANAGEMENT_TOKEN_SQL) + .bind(record.name.as_deref()) + .bind(record.clear_description) + .bind(record.description.as_deref()) + .bind(record.clear_allowed_ips) + .bind(allowed_ips) + .bind(permissions) + .bind(record.clear_expires_at) + .bind( + record + .expires_at_unix_secs + .and_then(|value| i64::try_from(value).ok()), + ) + .bind(record.is_active) + .bind(now as i64) + .bind(&record.token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?; + self.get_token_scoped(&record.token_id, expected_user_id) + .await + } + + async fn delete_management_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> Result { + let result = sqlx::query( + "DELETE FROM management_tokens WHERE id = ? AND (? IS NULL OR user_id = ?)", + ) + .bind(token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() > 0) + } + + async fn set_management_token_active_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + is_active: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + "UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ? AND (? IS NULL OR user_id = ?)", + ) + .bind(is_active) + .bind(now_unix_secs() as i64) + .bind(token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.get_token_scoped(token_id, expected_user_id).await + } + + async fn regenerate_management_token_secret_scoped( + &self, + mutation: &RegenerateManagementTokenSecret, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + mutation.validate()?; + let result = sqlx::query( + r#" +UPDATE management_tokens +SET token_hash = ?, token_prefix = ?, updated_at = ? +WHERE id = ? AND (? IS NULL OR user_id = ?) +"#, + ) + .bind(&mutation.token_hash) + .bind(mutation.token_prefix.as_deref()) + .bind(now_unix_secs() as i64) + .bind(&mutation.token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.get_token_scoped(&mutation.token_id, expected_user_id) + .await + } } +const UPDATE_MANAGEMENT_TOKEN_SQL: &str = r#" +UPDATE management_tokens +SET name = COALESCE(?, name), + description = CASE WHEN ? THEN NULL ELSE COALESCE(?, description) END, + allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END, + permissions = COALESCE(?, permissions), + expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END, + is_active = COALESCE(?, is_active), + updated_at = ? +WHERE id = ? AND (? IS NULL OR user_id = ?) +"#; + const TOKEN_COLUMNS: &str = r#" SELECT id, @@ -59,6 +198,39 @@ SELECT FROM management_tokens "#; +const LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL: &str = r#" +SELECT id +FROM users +WHERE id = ? + AND is_active = TRUE + AND is_deleted = FALSE + AND LOWER(role) = 'admin' + AND security_version = ? +FOR UPDATE +"#; + +const LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL: &str = r#" +SELECT + id, + user_id, + token_hash, + name, + description, + token_prefix, + allowed_ips, + permissions, + expires_at AS expires_at_unix_secs, + last_used_at AS last_used_at_unix_secs, + last_used_ip, + COALESCE(usage_count, 0) AS usage_count, + is_active, + created_at AS created_at_unix_ms, + updated_at AS updated_at_unix_secs +FROM management_tokens +WHERE id = ? +FOR UPDATE +"#; + const TOKEN_WITH_USER_COLUMNS: &str = r#" SELECT mt.id, @@ -220,71 +392,29 @@ INSERT INTO management_tokens ( &self, record: &UpdateManagementTokenRecord, ) -> Result, DataLayerError> { - record.validate()?; - let current = self.get_token(&record.token_id).await?; - let Some(current) = current else { - return Ok(None); - }; - let name = record.name.as_deref().unwrap_or(¤t.name); - let description = if record.clear_description { - None - } else { - record - .description - .as_deref() - .or(current.description.as_deref()) - }; - let allowed_ips = if record.clear_allowed_ips { - None - } else { - record.allowed_ips.as_ref().or(current.allowed_ips.as_ref()) - }; - let permissions = record.permissions.as_ref().or(current.permissions.as_ref()); - let expires_at = if record.clear_expires_at { - None - } else { - record.expires_at_unix_secs.or(current.expires_at_unix_secs) - }; - let is_active = record.is_active.unwrap_or(current.is_active); - let now = now_unix_secs(); + self.update_management_token_scoped(record, None).await + } - let result = sqlx::query( - r#" -UPDATE management_tokens -SET name = ?, - description = ?, - allowed_ips = ?, - permissions = ?, - expires_at = ?, - is_active = ?, - updated_at = ? -WHERE id = ? -"#, - ) - .bind(name) - .bind(description) - .bind(json_to_string(allowed_ips)?) - .bind(json_to_string(permissions)?) - .bind(expires_at.and_then(|value| i64::try_from(value).ok())) - .bind(is_active) - .bind(now as i64) - .bind(&record.token_id) - .execute(&self.pool) - .await - .map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?; - if result.rows_affected() == 0 { - return Ok(None); - } - self.get_token(&record.token_id).await + async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, DataLayerError> { + self.update_management_token_scoped(record, Some(user_id)) + .await } async fn delete_management_token(&self, token_id: &str) -> Result { - let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?") - .bind(token_id) - .execute(&self.pool) + self.delete_management_token_scoped(token_id, None).await + } + + async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + self.delete_management_token_scoped(token_id, Some(user_id)) .await - .map_sql_err()?; - Ok(result.rows_affected() > 0) } async fn set_management_token_active( @@ -292,43 +422,135 @@ WHERE id = ? token_id: &str, is_active: bool, ) -> Result, DataLayerError> { - let result = - sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?") - .bind(is_active) - .bind(now_unix_secs() as i64) - .bind(token_id) - .execute(&self.pool) + self.set_management_token_active_scoped(token_id, None, is_active) + .await + } + + async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + self.set_management_token_active_scoped(token_id, Some(user_id), is_active) + .await + } + + async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + let eligible_user = + sqlx::query_scalar::<_, String>(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL) + .bind(&mutation.expected_token.user_id) + .bind(mutation.expected_user_security_version) + .fetch_optional(&mut *tx) .await .map_sql_err()?; - if result.rows_affected() == 0 { - return Ok(None); + if eligible_user.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); } - self.get_token(token_id).await + + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_sql_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE management_tokens +SET is_active = TRUE, updated_at = ? +WHERE id = ? + AND BINARY token_hash = BINARY ? + AND is_active = FALSE + AND (expires_at IS NULL OR expires_at > ?) +"#, + ) + .bind(now_unix_secs() as i64) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .bind(i64::try_from(mutation.now_unix_secs).unwrap_or(i64::MAX)) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_sql_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let result = sqlx::query( + "DELETE FROM management_tokens WHERE id = ? AND BINARY token_hash = BINARY ? AND is_active = FALSE", + ) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) } async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, ) -> Result, DataLayerError> { - mutation.validate()?; - let result = sqlx::query( - r#" -UPDATE management_tokens -SET token_hash = ?, token_prefix = ?, updated_at = ? -WHERE id = ? -"#, - ) - .bind(&mutation.token_hash) - .bind(mutation.token_prefix.as_deref()) - .bind(now_unix_secs() as i64) - .bind(&mutation.token_id) - .execute(&self.pool) - .await - .map_sql_err()?; - if result.rows_affected() == 0 { - return Ok(None); - } - self.get_token(&mutation.token_id).await + self.regenerate_management_token_secret_scoped(mutation, None) + .await + } + + async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, DataLayerError> { + self.regenerate_management_token_secret_scoped(mutation, Some(user_id)) + .await } async fn record_management_token_usage( @@ -365,8 +587,18 @@ fn now_unix_secs() -> u64 { chrono::Utc::now().timestamp().max(0) as u64 } -fn optional_unix_secs(value: Option) -> Option { - value.and_then(|value| u64::try_from(value).ok()) +fn non_negative_u64(value: i64, field_name: &str) -> Result { + u64::try_from(value).map_err(|_| { + DataLayerError::UnexpectedValue(format!( + "management_tokens.{field_name} must not be negative" + )) + }) +} + +fn optional_unix_secs(value: Option, field_name: &str) -> Result, DataLayerError> { + value + .map(|value| non_negative_u64(value, field_name)) + .transpose() } fn json_to_string(value: Option<&serde_json::Value>) -> Result, DataLayerError> { @@ -420,15 +652,30 @@ fn map_token_row(row: &MySqlRow) -> Result("usage_count").map_sql_err()?).unwrap_or(0), + non_negative_u64( + row.try_get::("usage_count").map_sql_err()?, + "usage_count", + )?, row.try_get("is_active").map_sql_err()?, ) .with_timestamps( - optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?), - optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?), + optional_unix_secs( + row.try_get("created_at_unix_ms").map_sql_err()?, + "created_at", + )?, + optional_unix_secs( + row.try_get("updated_at_unix_secs").map_sql_err()?, + "updated_at", + )?, )) } @@ -454,7 +701,11 @@ fn map_token_with_user_row( #[cfg(test)] mod tests { - use super::MysqlManagementTokenRepository; + use super::{ + non_negative_u64, optional_unix_secs, MysqlManagementTokenRepository, + LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL, LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL, + UPDATE_MANAGEMENT_TOKEN_SQL, + }; use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig}; #[tokio::test] @@ -468,6 +719,60 @@ mod tests { let _repository = MysqlManagementTokenRepository::new(pool); } + #[test] + fn mysql_install_activation_locks_admin_identity_and_token_snapshot() { + for predicate in [ + "is_active = TRUE", + "is_deleted = FALSE", + "LOWER(role) = 'admin'", + "security_version = ?", + "FOR UPDATE", + ] { + assert!(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL.contains(predicate)); + } + for column in [ + "token_hash", + "name", + "description", + "token_prefix", + "allowed_ips", + "permissions", + "expires_at_unix_secs", + "last_used_at_unix_secs", + "last_used_ip", + "usage_count", + "is_active", + "created_at_unix_ms", + "updated_at_unix_secs", + "FOR UPDATE", + ] { + assert!(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL.contains(column)); + } + } + + #[test] + fn mysql_management_token_mapping_rejects_negative_integer_state() { + assert!(optional_unix_secs(Some(-1), "expires_at").is_err()); + assert_eq!( + optional_unix_secs(None, "expires_at").expect("SQL NULL should remain optional"), + None + ); + assert!(non_negative_u64(-1, "usage_count").is_err()); + } + + #[test] + fn mysql_management_token_updates_patch_only_explicit_fields() { + for clause in [ + "name = COALESCE(?, name)", + "allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END", + "permissions = COALESCE(?, permissions)", + "expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END", + "is_active = COALESCE(?, is_active)", + ] { + assert!(UPDATE_MANAGEMENT_TOKEN_SQL.contains(clause)); + } + } + #[test] fn mysql_management_token_pool_config_remains_driver_specific() { let config = SqlDatabaseConfig { diff --git a/crates/aether-data/adapters/mysql/src/oauth_providers.rs b/crates/aether-data/adapters/mysql/src/oauth_providers.rs index 1090bb4d6..44ba5949d 100644 --- a/crates/aether-data/adapters/mysql/src/oauth_providers.rs +++ b/crates/aether-data/adapters/mysql/src/oauth_providers.rs @@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, Row}; use aether_data_contracts::repository::oauth_providers::{ OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig, - UpsertOAuthProviderConfigRecord, + UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord, }; use aether_data_contracts::DataLayerError; @@ -115,6 +115,13 @@ WHERE users.is_active = 1 ) "#; +const COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL: &str = r#" +UPDATE oauth_providers +SET client_secret_encrypted = ? +WHERE BINARY provider_type = BINARY ? + AND BINARY client_secret_encrypted = BINARY ? +"#; + #[async_trait] impl OAuthProviderReadRepository for MysqlOAuthProviderRepository { async fn list_oauth_provider_configs( @@ -158,12 +165,56 @@ impl OAuthProviderReadRepository for MysqlOAuthProviderRepository { #[async_trait] impl OAuthProviderWriteRepository for MysqlOAuthProviderRepository { - async fn upsert_oauth_provider_config( + async fn upsert_oauth_provider_config_guarded( &self, record: &UpsertOAuthProviderConfigRecord, - ) -> Result { + ldap_exclusive: bool, + force_disable: bool, + _locked_users_snapshot: usize, + ) -> Result { record.validate()?; let now = now_unix_secs(); + let mut tx = self.pool.begin().await.map_sql_err()?; + let existing_enabled: Option = if record.is_enabled || force_disable { + None + } else { + sqlx::query_scalar::<_, String>( + "SELECT provider_type FROM oauth_providers ORDER BY provider_type FOR UPDATE", + ) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?") + .bind(&record.provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + }; + if existing_enabled == Some(true) { + let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL) + .bind(&record.provider_type) + .bind(&record.provider_type) + .bind(ldap_exclusive) + .bind(&record.provider_type) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let affected_count = + usize::try_from(row.try_get::("locked_count").map_sql_err()?.max(0)) + .map_err(|_| { + DataLayerError::UnexpectedValue( + "oauth_providers.locked_user_count overflowed".to_string(), + ) + })?; + if affected_count > 0 { + tx.rollback().await.map_sql_err()?; + return Ok( + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count, + }, + ); + } + } sqlx::query( r#" INSERT INTO oauth_providers ( @@ -228,27 +279,65 @@ ON DUPLICATE KEY UPDATE .bind(now as i64) .bind(record.client_secret_encrypted.mode_name()) .bind(record.client_secret_encrypted.value()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - self.get_provider(&record.provider_type) - .await? - .ok_or_else(|| { - DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string()) - }) + let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL) + .bind(&record.provider_type) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let provider = map_oauth_provider_row(&row)?; + tx.commit().await.map_sql_err()?; + Ok(UpsertOAuthProviderConfigOutcome::Upserted(provider)) } - async fn delete_oauth_provider_config( + async fn compare_and_swap_oauth_provider_client_secret( &self, provider_type: &str, + expected: &str, + replacement: &str, ) -> Result { - let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?") + let result = sqlx::query(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL) + .bind(replacement) .bind(provider_type) + .bind(expected) .execute(&self.pool) .await .map_sql_err()?; - Ok(result.rows_affected() > 0) + Ok(result.rows_affected() == 1) + } + + async fn delete_oauth_provider_config_if_unlinked( + &self, + provider_type: &str, + has_links_snapshot: bool, + ) -> Result { + if has_links_snapshot { + return Ok(false); + } + let mut tx = self.pool.begin().await.map_sql_err()?; + let provider_exists: Option = sqlx::query_scalar( + "SELECT provider_type FROM oauth_providers WHERE provider_type = ? FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let result = sqlx::query( + "DELETE FROM oauth_providers WHERE provider_type = ? AND NOT EXISTS (SELECT 1 FROM user_oauth_links WHERE user_oauth_links.provider_type = oauth_providers.provider_type)", + ) + .bind(provider_type) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(result.rows_affected() == 1) } } @@ -378,7 +467,16 @@ fn map_oauth_provider_row(row: &MySqlRow) -> Result Result { - let ssl_mode = if self.config.pool.require_ssl { - MySqlSslMode::Required - } else { - MySqlSslMode::Preferred - }; MySqlConnectOptions::from_str(self.config.url.trim()) .map(|options| { + // Preserve explicit VERIFY_CA/VERIFY_IDENTITY from the URL. + // `require_ssl` is a minimum transport guarantee: upgrade + // weaker modes to Required, never downgrade verification. + let ssl_mode = if self.config.pool.require_ssl + && !matches!( + options.get_ssl_mode(), + MySqlSslMode::VerifyCa | MySqlSslMode::VerifyIdentity + ) { + MySqlSslMode::Required + } else { + options.get_ssl_mode() + }; options .ssl_mode(ssl_mode) .statement_cache_capacity(self.config.pool.statement_cache_capacity) @@ -78,6 +85,52 @@ impl MysqlPoolFactory { mod tests { use super::MysqlPoolFactory; use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig}; + use sqlx::mysql::MySqlSslMode; + + fn ssl_mode(url: &str, require_ssl: bool) -> MySqlSslMode { + MysqlPoolFactory::new(SqlDatabaseConfig { + driver: DatabaseDriver::Mysql, + url: url.to_string(), + pool: SqlPoolConfig { + require_ssl, + ..SqlPoolConfig::default() + }, + }) + .expect("mysql config should build") + .connect_options() + .expect("mysql options should parse") + .get_ssl_mode() + } + + #[test] + fn preserves_explicit_mysql_verification_modes() { + assert!(matches!( + ssl_mode( + "mysql://user:pass@localhost/aether?ssl-mode=VERIFY_IDENTITY", + false + ), + MySqlSslMode::VerifyIdentity + )); + assert!(matches!( + ssl_mode( + "mysql://user:pass@localhost/aether?ssl-mode=VERIFY_CA", + true + ), + MySqlSslMode::VerifyCa + )); + } + + #[test] + fn require_ssl_only_upgrades_weak_mysql_modes() { + for mode in ["DISABLED", "PREFERRED", "REQUIRED"] { + let url = format!("mysql://user:pass@localhost/aether?ssl-mode={mode}"); + assert!(matches!(ssl_mode(&url, true), MySqlSslMode::Required)); + } + assert!(matches!( + ssl_mode("mysql://user:pass@localhost/aether", false), + MySqlSslMode::Preferred + )); + } #[tokio::test] async fn factory_builds_lazy_pool_from_valid_config() { diff --git a/crates/aether-data/adapters/mysql/src/provider_catalog.rs b/crates/aether-data/adapters/mysql/src/provider_catalog.rs index 208d57156..b6b42f1e1 100644 --- a/crates/aether-data/adapters/mysql/src/provider_catalog.rs +++ b/crates/aether-data/adapters/mysql/src/provider_catalog.rs @@ -9,9 +9,11 @@ use sqlx::{ use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate, - ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate, + ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, + ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, @@ -550,6 +552,48 @@ WHERE id = ? self.reload_provider(&provider.id, "updated").await } + pub async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + validate_non_empty(&update.provider_id, "provider catalog provider_id")?; + let expected_config = + optional_json_to_string(&update.expected_config, "providers.expected_config")?; + let config = optional_json_to_string(&update.config, "providers.config")?; + let rows_affected = sqlx::query( + r#" +UPDATE providers +SET config = ?, updated_at = ? +WHERE id = ? + AND config <=> ? +"#, + ) + .bind(config) + .bind(current_unix_secs() as i64) + .bind(&update.provider_id) + .bind(expected_config) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + + pub async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog provider_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM providers WHERE id = ?", + "UPDATE providers SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?", + update, + "providers.proxy", + ) + .await + } + pub async fn delete_provider(&self, provider_id: &str) -> Result { validate_non_empty(provider_id, "provider catalog provider_id")?; let rows_affected = sqlx::query("DELETE FROM providers WHERE id = ?") @@ -754,6 +798,21 @@ WHERE id = ? self.reload_endpoint(&endpoint.id, "updated").await } + pub async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog endpoint_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM provider_endpoints WHERE id = ?", + "UPDATE provider_endpoints SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?", + update, + "provider_endpoints.proxy", + ) + .await + } + pub async fn delete_endpoint(&self, endpoint_id: &str) -> Result { validate_non_empty(endpoint_id, "provider catalog endpoint_id")?; let rows_affected = sqlx::query("DELETE FROM provider_endpoints WHERE id = ?") @@ -936,6 +995,46 @@ WHERE id = ? self.reload_key(&key.id, "updated").await } + pub async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog key_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM provider_api_keys WHERE id = ?", + "UPDATE provider_api_keys SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?", + update, + "provider_api_keys.proxy", + ) + .await + } + + pub async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + validate_non_empty(&update.key_id, "provider catalog key_id")?; + validate_non_empty( + &update.expected_provider_id, + "provider catalog expected provider_id", + )?; + let rows_affected = sqlx::query( + "UPDATE provider_api_keys SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?", + ) + .bind(update.encrypted_api_key.as_deref()) + .bind(update.encrypted_auth_config.as_deref()) + .bind(&update.key_id) + .bind(&update.expected_provider_id) + .bind(update.expected_encrypted_api_key.as_deref()) + .bind(update.expected_encrypted_auth_config.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + pub async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -1354,51 +1453,18 @@ WHERE id = ? Ok(rows_affected > 0) } - pub async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - validate_non_empty(key_id, "provider catalog key_id")?; - validate_non_empty(encrypted_api_key, "provider catalog oauth api_key")?; - let rows_affected = sqlx::query( - r#" -UPDATE provider_api_keys -SET api_key = ?, auth_config = ?, expires_at = ?, updated_at = ? -WHERE id = ? -"#, - ) - .bind(encrypted_api_key) - .bind(encrypted_auth_config) - .bind(optional_i64_from_u64( - expires_at_unix_secs, - "provider_api_keys.expires_at", - )?) - .bind(current_unix_secs() as i64) - .bind(key_id) - .execute(&self.pool) - .await - .map_sql_err()? - .rows_affected(); - Ok(rows_affected > 0) - } - pub async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { validate_non_empty(key_id, "provider catalog key_id")?; let rows_affected = sqlx::query( r#" UPDATE provider_api_keys -SET oauth_invalid_at = ?, oauth_invalid_reason = ?, - auth_config = COALESCE(?, auth_config), updated_at = ? +SET oauth_invalid_at = ?, oauth_invalid_reason = ?, updated_at = ? WHERE id = ? "#, ) @@ -1407,7 +1473,6 @@ WHERE id = ? "provider_api_keys.oauth_invalid_at", )?) .bind(oauth_invalid_reason) - .bind(encrypted_auth_config_update) .bind(updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64) .bind(key_id) .execute(&self.pool) @@ -1755,7 +1820,7 @@ WHERE id = ? update.expected_encrypted_auth_config.as_deref() { builder - .push(" AND auth_config <=> ") + .push(" AND BINARY auth_config <=> BINARY ") .push_bind(expected_encrypted_auth_config); } let rows_affected = builder @@ -1873,7 +1938,7 @@ SET health_by_format = ?, circuit_breaker_by_format = ?, updated_at = ? WHERE id = ? AND JSON_EXTRACT(health_by_format, '$') <=> CAST(? AS JSON) AND JSON_EXTRACT(circuit_breaker_by_format, '$') <=> CAST(? AS JSON) - AND (? IS NULL OR auth_config <=> ?) + AND (? IS NULL OR BINARY auth_config <=> BINARY ?) "#, ) .bind(optional_json_to_string( @@ -2042,6 +2107,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository { Self::update_provider(self, provider).await } + async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_config(self, update).await + } + + async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_proxy(self, update).await + } + async fn delete_provider(&self, provider_id: &str) -> Result { Self::delete_provider(self, provider_id).await } @@ -2077,6 +2156,13 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository { Self::update_endpoint(self, endpoint).await } + async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_endpoint_proxy(self, update).await + } + async fn delete_endpoint(&self, endpoint_id: &str) -> Result { Self::delete_endpoint(self, endpoint_id).await } @@ -2095,6 +2181,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository { Self::update_key(self, key).await } + async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_key_proxy(self, update).await + } + + async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + Self::compare_and_swap_key_credentials(self, update).await + } + async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -2189,29 +2289,11 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository { Self::clear_key_oauth_invalid_marker(self, key_id).await } - async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - Self::update_key_oauth_credentials( - self, - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } - async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { Self::update_key_oauth_runtime_state( @@ -2219,7 +2301,6 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -2490,6 +2571,44 @@ fn optional_json_to_string( optional_json_ref_to_string(value.as_ref(), field_name) } +async fn compare_and_swap_proxy_json( + pool: &MysqlPool, + select_sql: &'static str, + update_sql: &'static str, + update: &ProviderCatalogProxyCasUpdate, + field_name: &'static str, +) -> Result { + // Legacy catalog rows may contain semantically identical JSON with Python-style + // whitespace. Comparing a re-serialized serde_json::Value directly to a TEXT column + // would make lazy credential migration conflict forever. Compare the parsed value first, + // then fence the write against the exact raw bytes that were observed. + // Outer None means the row does not exist; inner None is an existing SQL NULL proxy. + let observed_raw: Option> = sqlx::query_scalar::<_, Option>(select_sql) + .bind(&update.record_id) + .fetch_optional(pool) + .await + .map_sql_err()?; + let Some(observed_raw) = observed_raw else { + return Ok(false); + }; + let observed = optional_json_from_string(observed_raw.clone(), field_name)?; + if observed != update.expected_proxy { + return Ok(false); + } + + let replacement = optional_json_to_string(&update.proxy, field_name)?; + let rows_affected = sqlx::query(update_sql) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(&update.record_id) + .bind(observed_raw) + .execute(pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) +} + fn build_in_query<'a>( select_sql: &'static str, column: &'static str, @@ -3134,11 +3253,11 @@ fn map_key_row(row: &MySqlRow) -> Result Result(key) + key.with_auth_channel_policy_fields( + auth_type_by_format, + allow_auth_channel_mismatch_formats, + ) })? } @@ -3238,6 +3360,14 @@ mod tests { assert!(sql.contains("binary auth_config <=> binary ?")); } + #[test] + fn credential_cas_migrates_legacy_encrypted_key_with_binary_fence() { + let source = include_str!("provider_catalog.rs"); + assert!(source.contains( + "SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?" + )); + } + #[test] fn admin_credential_cas_has_atomic_rotation_guards() { let source = include_str!("provider_catalog.rs"); diff --git a/crates/aether-data/adapters/mysql/src/proxy_nodes.rs b/crates/aether-data/adapters/mysql/src/proxy_nodes.rs index bbe461eae..53cd53fd3 100644 --- a/crates/aether-data/adapters/mysql/src/proxy_nodes.rs +++ b/crates/aether-data/adapters/mysql/src/proxy_nodes.rs @@ -3,7 +3,8 @@ use sqlx::{mysql::MySqlRow, Row}; use aether_data_contracts::repository::proxy_nodes::{ bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, - normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security, + merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata, + normalize_proxy_metadata, proxy_metadata_has_explicit_tunnel_security, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation, @@ -17,6 +18,8 @@ use aether_data_contracts::DataLayerError; use crate::error::SqlResultExt; use crate::MysqlPool; +const PROXY_NODE_REGISTRATION_CAS_RETRIES: usize = 8; + fn log_reported_tunnel_error_event( node_id: &str, event: &TunnelErrorEventRecord, @@ -49,19 +52,22 @@ impl MysqlProxyNodeReadRepository { Self { pool } } - async fn upsert_node(&self, node: &StoredProxyNode) -> Result<(), DataLayerError> { + async fn write_node( + &self, + node: &StoredProxyNode, + update_existing: bool, + ) -> Result<(), DataLayerError> { let now = current_unix_secs(); - sqlx::query( - r#" + let upsert_sql = r#" INSERT INTO proxy_nodes ( - id, name, ip, port, region, status, registered_by, last_heartbeat_at, + id, tunnel_generation, name, ip, port, region, status, registered_by, last_heartbeat_at, heartbeat_interval, active_connections, total_requests, avg_latency_ms, is_manual, proxy_url, proxy_username, proxy_password, created_at, updated_at, remote_config, config_version, hardware_info, estimated_max_concurrency, tunnel_mode, tunnel_connected, tunnel_connected_at, failed_requests, dns_failures, stream_errors, proxy_metadata ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON DUPLICATE KEY UPDATE name = VALUES(name), ip = VALUES(ip), @@ -90,58 +96,128 @@ ON DUPLICATE KEY UPDATE dns_failures = VALUES(dns_failures), stream_errors = VALUES(stream_errors), proxy_metadata = VALUES(proxy_metadata) -"#, - ) - .bind(&node.id) - .bind(&node.name) - .bind(&node.ip) - .bind(node.port) - .bind(&node.region) - .bind(&node.status) - .bind(&node.registered_by) - .bind(optional_i64_from_u64( - node.last_heartbeat_at_unix_secs, - "proxy_nodes.last_heartbeat_at", - )?) - .bind(node.heartbeat_interval) - .bind(node.active_connections) - .bind(node.total_requests) - .bind(node.avg_latency_ms) - .bind(node.is_manual) - .bind(&node.proxy_url) - .bind(&node.proxy_username) - .bind(&node.proxy_password) - .bind(node.created_at_unix_ms.unwrap_or(now) as i64) - .bind(node.updated_at_unix_secs.unwrap_or(now) as i64) - .bind(optional_json_to_string( - &node.remote_config, - "proxy_nodes.remote_config", - )?) - .bind(node.config_version) - .bind(optional_json_to_string( - &node.hardware_info, - "proxy_nodes.hardware_info", - )?) - .bind(node.estimated_max_concurrency) - .bind(node.tunnel_mode) - .bind(node.tunnel_connected) - .bind(optional_i64_from_u64( - node.tunnel_connected_at_unix_secs, - "proxy_nodes.tunnel_connected_at", - )?) - .bind(node.failed_requests) - .bind(node.dns_failures) - .bind(node.stream_errors) - .bind(optional_json_to_string( - &node.proxy_metadata, - "proxy_nodes.proxy_metadata", - )?) - .execute(&self.pool) - .await - .map_sql_err()?; +"#; + let sql = if update_existing { + upsert_sql + } else { + upsert_sql + .split_once("\nON DUPLICATE KEY UPDATE") + .map(|(insert_sql, _)| insert_sql) + .expect("proxy node upsert SQL should contain its conflict clause") + }; + sqlx::query(sql) + .bind(&node.id) + .bind(&node.tunnel_generation) + .bind(&node.name) + .bind(&node.ip) + .bind(node.port) + .bind(&node.region) + .bind(&node.status) + .bind(&node.registered_by) + .bind(optional_i64_from_u64( + node.last_heartbeat_at_unix_secs, + "proxy_nodes.last_heartbeat_at", + )?) + .bind(node.heartbeat_interval) + .bind(node.active_connections) + .bind(node.total_requests) + .bind(node.avg_latency_ms) + .bind(node.is_manual) + .bind(&node.proxy_url) + .bind(&node.proxy_username) + .bind(&node.proxy_password) + .bind(node.created_at_unix_ms.unwrap_or(now) as i64) + .bind(node.updated_at_unix_secs.unwrap_or(now) as i64) + .bind(optional_json_to_string( + &node.remote_config, + "proxy_nodes.remote_config", + )?) + .bind(node.config_version) + .bind(optional_json_to_string( + &node.hardware_info, + "proxy_nodes.hardware_info", + )?) + .bind(node.estimated_max_concurrency) + .bind(node.tunnel_mode) + .bind(node.tunnel_connected) + .bind(optional_i64_from_u64( + node.tunnel_connected_at_unix_secs, + "proxy_nodes.tunnel_connected_at", + )?) + .bind(node.failed_requests) + .bind(node.dns_failures) + .bind(node.stream_errors) + .bind(optional_json_to_string( + &node.proxy_metadata, + "proxy_nodes.proxy_metadata", + )?) + .execute(&self.pool) + .await + .map_sql_err()?; Ok(()) } + async fn insert_node(&self, node: &StoredProxyNode) -> Result<(), DataLayerError> { + self.write_node(node, false).await + } + + async fn update_existing_registration_if_unchanged( + &self, + mutation: &ProxyNodeRegistrationMutation, + existing: &StoredProxyNode, + replacement_proxy_metadata: Option<&serde_json::Value>, + now: u64, + ) -> Result { + let hardware_info = + optional_json_to_string(&mutation.hardware_info, "proxy_nodes.hardware_info")?; + let replacement_proxy_metadata_json = optional_json_to_string( + &replacement_proxy_metadata.cloned(), + "proxy_nodes.proxy_metadata", + )?; + let expected_proxy_metadata = + optional_json_to_string(&existing.proxy_metadata, "proxy_nodes.proxy_metadata")?; + let result = sqlx::query(UPDATE_PROXY_NODE_REGISTRATION_SQL) + .bind(&mutation.name) + .bind(&mutation.ip) + .bind(mutation.port) + .bind(mutation.region.as_deref()) + .bind(mutation.registered_by.as_deref()) + .bind(now as i64) + .bind(mutation.heartbeat_interval) + .bind(mutation.active_connections) + .bind(mutation.total_requests) + .bind(mutation.avg_latency_ms) + .bind(hardware_info) + .bind(mutation.estimated_max_concurrency) + .bind(mutation.tunnel_mode) + .bind(replacement_proxy_metadata_json) + .bind(now as i64) + .bind(&existing.id) + .bind(&existing.tunnel_generation) + .bind(&existing.ip) + .bind(existing.port) + .bind(expected_proxy_metadata.as_deref()) + .bind(expected_proxy_metadata.as_deref()) + .bind(expected_proxy_metadata.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() != 0 { + return Ok(true); + } + + let Some(current) = self.find_proxy_node(&existing.id).await? else { + return Ok(false); + }; + Ok(proxy_node_registration_matches( + ¤t, + mutation, + existing, + replacement_proxy_metadata, + now, + )) + } + async fn find_duplicate_proxy_node( &self, ip: &str, @@ -171,9 +247,26 @@ ON DUPLICATE KEY UPDATE row.as_ref().map(map_proxy_node_row).transpose() } + async fn find_registered_proxy_node_by_endpoint( + &self, + ip: &str, + port: i32, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE BINARY ip = BINARY ? AND port = ? AND is_manual = 0 ORDER BY created_at ASC, id ASC LIMIT 1" + )) + .bind(ip) + .bind(port) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_proxy_node_row).transpose() + } + async fn insert_event( &self, node_id: &str, + expected_tunnel_generation: Option<&str>, event_type: &str, detail: Option<&str>, event_metadata: Option<&serde_json::Value>, @@ -182,10 +275,11 @@ ON DUPLICATE KEY UPDATE sqlx::query( r#" INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at) -VALUES (?, ?, ?, ?, ?) +SELECT id, ?, ?, ?, ? +FROM proxy_nodes +WHERE id = ? AND (? IS NULL OR BINARY tunnel_generation = BINARY ?) "#, ) - .bind(node_id) .bind(event_type) .bind(detail) .bind(optional_json_to_string( @@ -193,6 +287,9 @@ VALUES (?, ?, ?, ?, ?) "proxy_node_events.event_metadata", )?) .bind(created_at_unix_secs.unwrap_or_else(current_unix_secs) as i64) + .bind(node_id) + .bind(expected_tunnel_generation) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_sql_err()?; @@ -203,6 +300,7 @@ VALUES (?, ?, ?, ?, ?) &self, table: &str, node_id: &str, + expected_tunnel_generation: Option<&str>, bucket_start: u64, sample: &TunnelMetricsSample, ) -> Result<(), DataLayerError> { @@ -225,7 +323,9 @@ INSERT INTO {table} ( ws_in_frames_delta, ws_out_frames_delta ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +SELECT ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? +FROM proxy_nodes +WHERE id = ? AND (? IS NULL OR BINARY tunnel_generation = BINARY ?) ON DUPLICATE KEY UPDATE samples = samples + VALUES(samples), uptime_samples = uptime_samples + VALUES(uptime_samples), @@ -257,6 +357,9 @@ ON DUPLICATE KEY UPDATE .bind(sample.ws_out_bytes_delta) .bind(sample.ws_in_frames_delta) .bind(sample.ws_out_frames_delta) + .bind(node_id) + .bind(expected_tunnel_generation) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_sql_err()?; @@ -330,6 +433,7 @@ ON DUPLICATE KEY UPDATE const PROXY_NODE_COLUMNS: &str = r#" SELECT id, + tunnel_generation, name, ip, port, @@ -361,6 +465,140 @@ SELECT FROM proxy_nodes "#; +const APPLY_HEARTBEAT_SQL: &str = r#" +UPDATE proxy_nodes +SET last_heartbeat_at = ?, + tunnel_connected_at = CASE + WHEN status <> 'online' OR tunnel_connected = 0 THEN ? + ELSE tunnel_connected_at + END, + updated_at = CASE + WHEN status <> 'online' OR tunnel_connected = 0 THEN ? + ELSE updated_at + END, + status = 'online', + tunnel_connected = 1, + heartbeat_interval = COALESCE(?, heartbeat_interval), + active_connections = COALESCE(?, active_connections), + avg_latency_ms = COALESCE(?, avg_latency_ms), + total_requests = total_requests + GREATEST(COALESCE(?, 0), 0), + failed_requests = failed_requests + GREATEST(COALESCE(?, 0), 0), + dns_failures = dns_failures + GREATEST(COALESCE(?, 0), 0), + stream_errors = stream_errors + GREATEST(COALESCE(?, 0), 0) +WHERE id = ? + AND tunnel_mode = 1 + AND BINARY tunnel_generation = BINARY ? +"#; + +const CAS_HEARTBEAT_PROXY_METADATA_SQL: &str = r#" +UPDATE proxy_nodes +SET proxy_metadata = ?, updated_at = ? +WHERE id = ? AND BINARY tunnel_generation = BINARY ? + AND ( + (proxy_metadata IS NULL AND ? IS NULL) + OR ( + proxy_metadata IS NOT NULL AND ? IS NOT NULL + AND JSON_VALID(proxy_metadata) = 1 + AND CAST(proxy_metadata AS JSON) = CAST(? AS JSON) + ) + ) +"#; + +const UPDATE_TUNNEL_STATUS_SQL: &str = r#" +UPDATE proxy_nodes +SET tunnel_connected = ?, + active_connections = CASE WHEN ? THEN active_connections ELSE 0 END, + tunnel_connected_at = ?, + status = CASE WHEN ? THEN 'online' ELSE 'offline' END, + updated_at = ? +WHERE id = ? + AND BINARY tunnel_generation = BINARY ? + AND (tunnel_connected_at IS NULL OR tunnel_connected_at <= ?) +"#; + +const UPDATE_MANUAL_PROXY_NODE_SQL: &str = r#" +UPDATE proxy_nodes +SET name = COALESCE(?, name), + ip = COALESCE(?, ip), + port = COALESCE(?, port), + region = COALESCE(?, region), + proxy_url = COALESCE(?, proxy_url), + proxy_username = COALESCE(?, proxy_username), + proxy_password = COALESCE(?, proxy_password), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND BINARY tunnel_generation = BINARY ? +"#; + +const UPDATE_PROXY_NODE_REGISTRATION_SQL: &str = r#" +UPDATE proxy_nodes +SET name = ?, ip = ?, port = ?, region = ?, registered_by = ?, + last_heartbeat_at = ?, heartbeat_interval = ?, + active_connections = COALESCE(?, active_connections), + total_requests = COALESCE(?, total_requests), + avg_latency_ms = COALESCE(?, avg_latency_ms), + hardware_info = COALESCE(?, hardware_info), + estimated_max_concurrency = COALESCE(?, estimated_max_concurrency), + tunnel_mode = ?, proxy_metadata = COALESCE(?, proxy_metadata), updated_at = ? +WHERE BINARY id = BINARY ? AND BINARY tunnel_generation = BINARY ? + AND is_manual = 0 AND BINARY ip = BINARY ? AND port = ? + AND ( + (proxy_metadata IS NULL AND ? IS NULL) + OR ( + proxy_metadata IS NOT NULL AND ? IS NOT NULL + AND JSON_VALID(proxy_metadata) = 1 + AND CAST(proxy_metadata AS JSON) = CAST(? AS JSON) + ) + ) +"#; + +const UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL: &str = r#" +UPDATE proxy_nodes +SET name = COALESCE(?, name), remote_config = ?, + config_version = config_version + 1, updated_at = ? +WHERE id = ? AND BINARY tunnel_generation = BINARY ? AND config_version = ? + AND is_manual = 0 +"#; + +const RECORD_PROXY_NODE_TRAFFIC_SQL: &str = r#" +UPDATE proxy_nodes +SET total_requests = total_requests + GREATEST(?, 0), + failed_requests = failed_requests + GREATEST(?, 0), + dns_failures = dns_failures + GREATEST(?, 0), + stream_errors = stream_errors + GREATEST(?, 0), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND BINARY tunnel_generation = BINARY ? +"#; + +const INCREMENT_MANUAL_PROXY_NODE_REQUESTS_SQL: &str = r#" +UPDATE proxy_nodes +SET total_requests = total_requests + GREATEST(?, 0), + failed_requests = failed_requests + GREATEST(?, 0), + avg_latency_ms = COALESCE(?, avg_latency_ms), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND BINARY tunnel_generation = BINARY ? +"#; + +const UNREGISTER_PROXY_NODE_SQL: &str = r#" +UPDATE proxy_nodes +SET status = 'offline', tunnel_connected = 0, active_connections = 0, + tunnel_connected_at = ?, updated_at = ? +WHERE id = ? + AND BINARY tunnel_generation = BINARY ? +"#; + +// Run after the parent delete commits so delete never waits on an outbox row +// already claimed by the flusher (which acquires locks in the opposite order). +const RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL: &str = r#" +DELETE FROM usage_counter_deltas +WHERE kind = 'proxy_node' + AND target_id = ? + AND BINARY target_tunnel_generation = BINARY ? + AND processed_at IS NULL +"#; + #[async_trait] impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository { async fn list_proxy_nodes(&self) -> Result, DataLayerError> { @@ -582,6 +820,61 @@ WHERE is_manual = 0 Ok(result.rows_affected() as usize) } + async fn compare_and_set_proxy_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_password = ?, updated_at = ? +WHERE id = ? AND BINARY proxy_password = BINARY ? +"#, + ) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(expected) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_set_proxy_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result { + let expected = serde_json::to_string(expected).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let replacement = serde_json::to_string(replacement).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_metadata = ?, updated_at = ? +WHERE id = ? + AND proxy_metadata IS NOT NULL + AND JSON_VALID(proxy_metadata) = 1 + AND CAST(proxy_metadata AS JSON) = CAST(? AS JSON) +"#, + ) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(&expected) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + async fn create_manual_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -593,9 +886,14 @@ WHERE is_manual = 0 return Err(duplicate_proxy_node_error(&existing)); } + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + if let Some(existing) = self.find_proxy_node(&node_id).await? { + return Err(proxy_node_id_in_use_error(&existing)); + } let now = Some(current_unix_secs()); let node = StoredProxyNode::new( - uuid::Uuid::new_v4().to_string(), + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -630,7 +928,18 @@ WHERE is_manual = 0 now, ); - self.upsert_node(&node).await?; + if let Err(error) = self.insert_node(&node).await { + if let Some(duplicate) = self + .find_duplicate_proxy_node(&mutation.ip, mutation.port, None) + .await? + { + return Err(duplicate_proxy_node_error(&duplicate)); + } + if let Some(owner) = self.find_proxy_node(&node.id).await? { + return Err(proxy_node_id_in_use_error(&owner)); + } + return Err(error); + } Ok(node) } @@ -638,17 +947,17 @@ WHERE is_manual = 0 &self, mutation: &ProxyNodeManualUpdateMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; - if !node.is_manual { + if !existing.is_manual { return Err(DataLayerError::InvalidInput( "只能编辑手动添加的代理节点".to_string(), )); } - let next_ip = mutation.ip.as_deref().unwrap_or(node.ip.as_str()); - let next_port = mutation.port.unwrap_or(node.port); + let next_ip = mutation.ip.as_deref().unwrap_or(existing.ip.as_str()); + let next_port = mutation.port.unwrap_or(existing.port); if let Some(existing) = self .find_duplicate_proxy_node(next_ip, next_port, Some(&mutation.node_id)) .await? @@ -656,56 +965,62 @@ WHERE is_manual = 0 return Err(duplicate_proxy_node_error(&existing)); } - if let Some(name) = mutation.name.as_ref() { - node.name = name.clone(); + let result = sqlx::query(UPDATE_MANUAL_PROXY_NODE_SQL) + .bind(mutation.name.as_deref()) + .bind(mutation.ip.as_deref()) + .bind(mutation.port) + .bind(mutation.region.as_deref()) + .bind(mutation.proxy_url.as_deref()) + .bind(mutation.proxy_username.as_deref()) + .bind(mutation.proxy_password.as_deref()) + .bind(current_unix_secs() as i64) + .bind(&mutation.node_id) + .bind(&existing.tunnel_generation) + .execute(&self.pool) + .await; + let result = match result { + Ok(result) => result, + Err(error) => { + if let Some(duplicate) = self + .find_duplicate_proxy_node(next_ip, next_port, Some(&mutation.node_id)) + .await? + { + return Err(duplicate_proxy_node_error(&duplicate)); + } + return Err(DataLayerError::sql(error)); + } + }; + if result.rows_affected() == 0 { + return Ok(None); } - if let Some(ip) = mutation.ip.as_ref() { - node.ip = ip.clone(); - } - if let Some(port) = mutation.port { - node.port = port; - } - if let Some(region) = mutation.region.as_ref() { - node.region = Some(region.clone()); - } - if let Some(proxy_url) = mutation.proxy_url.as_ref() { - node.proxy_url = Some(proxy_url.clone()); - } - if let Some(proxy_username) = mutation.proxy_username.as_ref() { - node.proxy_username = Some(proxy_username.clone()); - } - if let Some(proxy_password) = mutation.proxy_password.as_ref() { - node.proxy_password = Some(proxy_password.clone()); - } - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(Some(node)) + self.find_proxy_node(&mutation.node_id).await } async fn register_node( &self, mutation: &ProxyNodeRegistrationMutation, ) -> Result { - let now = Some(current_unix_secs()); + let requested_id = requested_proxy_node_id(mutation.node_id.as_deref())?; let normalized_proxy_metadata = normalize_proxy_metadata( mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), ); + let rotates_tunnel_security = + proxy_metadata_has_explicit_tunnel_security(normalized_proxy_metadata.as_ref()); - let existing = sqlx::query(&format!( - "{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? AND is_manual = 0 ORDER BY created_at ASC, id ASC LIMIT 1" - )) - .bind(&mutation.ip) - .bind(mutation.port) - .fetch_optional(&self.pool) - .await - .map_sql_err()?; - - let mut node = if let Some(row) = existing.as_ref() { - map_proxy_node_row(row)? - } else { - StoredProxyNode::new( - uuid::Uuid::new_v4().to_string(), + let Some(initial_existing) = self + .find_registered_proxy_node_by_endpoint(&mutation.ip, mutation.port) + .await? + else { + let node_id = requested_id + .clone() + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + if let Some(existing) = self.find_proxy_node(&node_id).await? { + return Err(proxy_node_id_in_use_error(&existing)); + } + let now = Some(current_unix_secs()); + let node = StoredProxyNode::new( + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -726,141 +1041,243 @@ WHERE is_manual = 0 mutation.registered_by.clone(), now, mutation.avg_latency_ms, - normalized_proxy_metadata.clone(), + merge_proxy_metadata_for_registration(None, normalized_proxy_metadata.clone()), mutation.hardware_info.clone(), mutation.estimated_max_concurrency, None, None, now, now, - ) + ); + if let Err(error) = self.insert_node(&node).await { + if let Some(winner) = self + .find_duplicate_proxy_node(&mutation.ip, mutation.port, None) + .await? + { + if winner.is_manual { + return Err(duplicate_proxy_node_error(&winner)); + } + if let Some(requested_id) = requested_id.as_deref() { + if requested_id != winner.id { + return Err(proxy_node_registration_identity_error( + requested_id, + &winner.id, + )); + } + } + return Ok(winner); + } + if let Some(owner) = self.find_proxy_node(&node.id).await? { + return Err(proxy_node_id_in_use_error(&owner)); + } + return Err(error); + } + return Ok(node); }; - node.name = mutation.name.clone(); - node.ip = mutation.ip.clone(); - node.port = mutation.port; - node.region = mutation.region.clone(); - node.registered_by = mutation.registered_by.clone(); - node.last_heartbeat_at_unix_secs = now; - node.heartbeat_interval = mutation.heartbeat_interval; - node.tunnel_mode = mutation.tunnel_mode; - if let Some(active_connections) = mutation.active_connections { - node.active_connections = active_connections; + if let Some(requested_id) = requested_id.as_deref() { + if requested_id != initial_existing.id { + return Err(proxy_node_registration_identity_error( + requested_id, + &initial_existing.id, + )); + } } - if let Some(total_requests) = mutation.total_requests { - node.total_requests = total_requests; + + let pinned_id = initial_existing.id.clone(); + let pinned_generation = initial_existing.tunnel_generation.clone(); + let mut existing = initial_existing; + for attempt in 0..PROXY_NODE_REGISTRATION_CAS_RETRIES { + if attempt != 0 { + existing = self + .find_registered_proxy_node_by_endpoint(&mutation.ip, mutation.port) + .await? + .ok_or_else(proxy_node_registration_changed_error)?; + } + if existing.id != pinned_id || existing.tunnel_generation != pinned_generation { + return Err(proxy_node_registration_changed_error()); + } + + let replacement_proxy_metadata = merge_proxy_metadata_for_registration( + existing.proxy_metadata.as_ref(), + normalized_proxy_metadata.clone(), + ); + let now = current_unix_secs(); + if self + .update_existing_registration_if_unchanged( + mutation, + &existing, + replacement_proxy_metadata.as_ref(), + now, + ) + .await? + { + return self + .find_proxy_node(&pinned_id) + .await? + .filter(|current| current.tunnel_generation == pinned_generation) + .ok_or_else(proxy_node_registration_changed_error); + } + if rotates_tunnel_security { + return Err(DataLayerError::UnexpectedValue( + "proxy node changed during explicit tunnel security rotation".to_string(), + )); + } } - if let Some(avg_latency_ms) = mutation.avg_latency_ms { - node.avg_latency_ms = Some(avg_latency_ms); - } - if let Some(hardware_info) = mutation.hardware_info.as_ref() { - node.hardware_info = Some(hardware_info.clone()); - } - if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency { - node.estimated_max_concurrency = Some(estimated_max_concurrency); - } - if let Some(proxy_metadata) = normalized_proxy_metadata { - node.proxy_metadata = Some(proxy_metadata); - } - if node.created_at_unix_ms.is_none() { - node.created_at_unix_ms = now; - } - node.updated_at_unix_secs = now; - self.upsert_node(&node).await?; - Ok(node) + + Err(DataLayerError::UnexpectedValue( + "proxy node registration changed during every CAS retry".to_string(), + )) } async fn apply_heartbeat( &self, mutation: &ProxyNodeHeartbeatMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; - if !node.tunnel_mode { + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != existing.tunnel_generation) + { + return Ok(None); + } + if !existing.tunnel_mode { return Err(DataLayerError::InvalidInput( "non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode" .to_string(), )); } - let previous_proxy_metadata = node.proxy_metadata.clone(); + let tunnel_generation = existing.tunnel_generation.clone(); let now_unix_secs = current_unix_secs(); - let now = Some(now_unix_secs); - node.last_heartbeat_at_unix_secs = now; - if node.status != "online" || !node.tunnel_connected { - node.status = "online".to_string(); - node.tunnel_connected = true; - node.tunnel_connected_at_unix_secs = now; - node.updated_at_unix_secs = now; - } - if let Some(value) = mutation.heartbeat_interval { - node.heartbeat_interval = value; - } - if let Some(value) = mutation.active_connections { - node.active_connections = value; - } - if let Some(value) = mutation.avg_latency_ms { - node.avg_latency_ms = Some(value); - } - let normalized_proxy_metadata = normalize_proxy_metadata( + let now = i64::try_from(now_unix_secs).unwrap_or(i64::MAX); + let has_proxy_metadata_update = normalize_heartbeat_proxy_metadata( + None, mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), - ); - let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security( - previous_proxy_metadata.as_ref(), - normalized_proxy_metadata, - ); - if let Some(value) = normalized_proxy_metadata { - node.proxy_metadata = Some(value); + ) + .is_some(); + + let result = sqlx::query(APPLY_HEARTBEAT_SQL) + .bind(now) + .bind(now) + .bind(now) + .bind(mutation.heartbeat_interval) + .bind(mutation.active_connections) + .bind(mutation.avg_latency_ms) + .bind(mutation.total_requests_delta) + .bind(mutation.failed_requests_delta) + .bind(mutation.dns_failures_delta) + .bind(mutation.stream_errors_delta) + .bind(&mutation.node_id) + .bind(&tunnel_generation) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); } - if let Some(value) = mutation.total_requests_delta.filter(|value| *value > 0) { - node.total_requests += value; + + let mut updated = None; + let mut tunnel_metrics_sample = None; + if has_proxy_metadata_update { + for _ in 0..8 { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + let Some(replacement) = normalize_heartbeat_proxy_metadata( + current.proxy_metadata.as_ref(), + mutation.proxy_metadata.as_ref(), + mutation.proxy_version.as_deref(), + ) else { + break; + }; + if current.proxy_metadata.as_ref() == Some(&replacement) { + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + Some(&replacement), + current.active_connections, + current.tunnel_connected, + ); + updated = Some(current); + break; + } + + let expected = + optional_json_to_string(¤t.proxy_metadata, "proxy_nodes.proxy_metadata")?; + let replacement_json = serde_json::to_string(&replacement).map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "proxy_nodes.proxy_metadata contains unserializable JSON: {error}" + )) + })?; + let result = sqlx::query(CAS_HEARTBEAT_PROXY_METADATA_SQL) + .bind(replacement_json) + .bind(now) + .bind(&mutation.node_id) + .bind(&tunnel_generation) + .bind(expected.as_deref()) + .bind(expected.as_deref()) + .bind(expected.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + continue; + } + let Some(after_cas) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if after_cas.tunnel_generation != tunnel_generation { + return Ok(None); + } + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + after_cas.proxy_metadata.as_ref(), + after_cas.active_connections, + after_cas.tunnel_connected, + ); + updated = Some(after_cas); + break; + } } - if let Some(value) = mutation.failed_requests_delta.filter(|value| *value > 0) { - node.failed_requests += value; - } - if let Some(value) = mutation.dns_failures_delta.filter(|value| *value > 0) { - node.dns_failures += value; - } - if let Some(value) = mutation.stream_errors_delta.filter(|value| *value > 0) { - node.stream_errors += value; - } - let reconciled_remote_config = reconcile_remote_config_after_heartbeat( - node.remote_config.as_ref(), - mutation.proxy_version.as_deref(), - ); - if reconciled_remote_config != node.remote_config { - node.remote_config = reconciled_remote_config; - node.config_version = node.config_version.saturating_add(1); - node.updated_at_unix_secs = now; - } - let tunnel_metrics_sample = build_tunnel_metrics_sample( - previous_proxy_metadata.as_ref(), - node.proxy_metadata.as_ref(), - node.active_connections, - node.tunnel_connected, - ); - self.upsert_node(&node).await?; + let updated = if let Some(updated) = updated { + updated + } else { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + current + }; if let Some(sample) = tunnel_metrics_sample.as_ref() { self.upsert_metrics_bucket( "proxy_node_metrics_1m", - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute), sample, ) .await?; self.upsert_metrics_bucket( "proxy_node_metrics_1h", - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour), sample, ) .await?; for error in &sample.recent_error_events { - log_reported_tunnel_error_event(&node.id, error, now_unix_secs); + log_reported_tunnel_error_event(&updated.id, error, now_unix_secs); let detail = build_tunnel_error_event_detail(error); let event_metadata = serde_json::json!({ "source": "heartbeat", @@ -874,7 +1291,8 @@ WHERE is_manual = 0 "timestamp_unix_ms": error.timestamp_unix_ms, }); self.insert_event( - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, Some(detail.as_str()), Some(&event_metadata), @@ -887,35 +1305,83 @@ WHERE is_manual = 0 .await?; } } - Ok(Some(node)) + if reconcile_remote_config_after_heartbeat( + updated.remote_config.as_ref(), + mutation.proxy_version.as_deref(), + ) != updated.remote_config + { + return self + .update_remote_config(&ProxyNodeRemoteConfigMutation { + node_id: mutation.node_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), + node_name: None, + allowed_ports: None, + log_level: None, + heartbeat_interval: None, + scheduling_state: None, + upgrade_to: Some(None), + }) + .await; + } + + Ok(Some(updated)) } async fn record_traffic( &self, mutation: &ProxyNodeTrafficMutation, ) -> Result { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(&mutation.node_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; return Ok(false); }; - if !node.is_manual { + let generation = map_proxy_node_row(&row)?.tunnel_generation; + let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + if expected_generation != generation { + tx.rollback().await.map_sql_err()?; return Ok(false); } - node.total_requests += mutation.total_requests_delta.max(0); - node.failed_requests += mutation.failed_requests_delta.max(0); - node.dns_failures += mutation.dns_failures_delta.max(0); - node.stream_errors += mutation.stream_errors_delta.max(0); - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(true) + let result = sqlx::query(RECORD_PROXY_NODE_TRAFFIC_SQL) + .bind(mutation.total_requests_delta) + .bind(mutation.failed_requests_delta) + .bind(mutation.dns_failures_delta) + .bind(mutation.stream_errors_delta) + .bind(current_unix_secs() as i64) + .bind(&mutation.node_id) + .bind(expected_generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + let applied = result.rows_affected() > 0; + tx.commit().await.map_sql_err()?; + Ok(applied) } async fn update_tunnel_status( &self, mutation: &ProxyNodeTunnelStatusMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(node) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } let event_time = mutation .observed_at_unix_secs @@ -932,108 +1398,211 @@ WHERE is_manual = 0 ) }); - if node - .tunnel_connected_at_unix_secs - .is_some_and(|last_transition| event_time < last_transition) - { - self.insert_event( - &mutation.node_id, - event_type, - Some(&format!("[stale_ignored] {event_detail}")), - None, - Some(current_unix_secs()), - ) - .await?; - return Ok(Some(node)); - } - - node.tunnel_connected = mutation.connected; - node.tunnel_connected_at_unix_secs = Some(event_time); - node.status = if mutation.connected { - "online".to_string() - } else { - "offline".to_string() + let event_time_i64 = i64::try_from(event_time).unwrap_or(i64::MAX); + let result = sqlx::query(UPDATE_TUNNEL_STATUS_SQL) + .bind(mutation.connected) + .bind(mutation.connected) + .bind(event_time_i64) + .bind(mutation.connected) + .bind(event_time_i64) + .bind(&mutation.node_id) + .bind(&node.tunnel_generation) + .bind(event_time_i64) + .execute(&self.pool) + .await + .map_sql_err()?; + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); }; - if !mutation.connected { - node.active_connections = 0; + if current.tunnel_generation != node.tunnel_generation { + return Ok(None); } - node.updated_at_unix_secs = Some(event_time); - self.upsert_node(&node).await?; + let stale = result.rows_affected() == 0 + && current + .tunnel_connected_at_unix_secs + .is_some_and(|last_transition| event_time < last_transition); + let persisted_detail = if stale { + format!("[stale_ignored] {event_detail}") + } else { + event_detail + }; self.insert_event( &mutation.node_id, + Some(node.tunnel_generation.as_str()), event_type, - Some(&event_detail), + Some(&persisted_detail), None, - Some(event_time), + Some(if stale { + current_unix_secs() + } else { + event_time + }), ) .await?; - Ok(Some(node)) + Ok(Some(current)) } async fn unregister_node( &self, node_id: &str, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(node_id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(node_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; return Ok(None); }; - let now = Some(current_unix_secs()); - node.status = "offline".to_string(); - node.tunnel_connected = false; - node.active_connections = 0; - node.tunnel_connected_at_unix_secs = now; - node.updated_at_unix_secs = now; - self.upsert_node(&node).await?; - Ok(Some(node)) + let generation = map_proxy_node_row(&row)?.tunnel_generation; + let now = current_unix_secs() as i64; + sqlx::query(UNREGISTER_PROXY_NODE_SQL) + .bind(now) + .bind(now) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + let updated = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err() + .and_then(|row| map_proxy_node_row(&row))?; + tx.commit().await.map_sql_err()?; + Ok(Some(updated)) } async fn delete_node(&self, node_id: &str) -> Result, DataLayerError> { - let existing = self.find_proxy_node(node_id).await?; - if existing.is_some() { - sqlx::query("DELETE FROM proxy_node_events WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_node_metrics_1m WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_node_metrics_1h WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_nodes WHERE id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(node_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(None); + }; + let existing = map_proxy_node_row(&row)?; + let generation = existing.tunnel_generation.as_str(); + + // Child tables do not carry generation; retain the parent identity check + // so cleanup cannot target a replacement row if schema constraints differ. + sqlx::query( + "DELETE FROM proxy_node_events WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes p WHERE p.id = ? AND BINARY p.tunnel_generation = BINARY ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "DELETE FROM proxy_node_metrics_1m WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes p WHERE p.id = ? AND BINARY p.tunnel_generation = BINARY ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "DELETE FROM proxy_node_metrics_1h WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes p WHERE p.id = ? AND BINARY p.tunnel_generation = BINARY ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + + let deleted = sqlx::query( + "DELETE FROM proxy_nodes WHERE id = ? AND BINARY tunnel_generation = BINARY ?", + ) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + if deleted.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(None); } - Ok(existing) + tx.commit().await.map_sql_err()?; + if let Err(error) = sqlx::query(RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL) + .bind(node_id) + .bind(generation) + .execute(&self.pool) + .await + .map_sql_err() + { + tracing::warn!( + node_id = %node_id, + tunnel_generation = %generation, + error = ?error, + "failed to retire deleted proxy node counter rows" + ); + } + Ok(Some(existing)) } async fn update_remote_config( &self, mutation: &ProxyNodeRemoteConfigMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { - return Ok(None); - }; - if node.is_manual { - return Err(DataLayerError::InvalidInput( - "手动节点不支持远程配置下发".to_string(), - )); + for _ in 0..8 { + let Some(node) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } + if node.is_manual { + return Err(DataLayerError::InvalidInput( + "手动节点不支持远程配置下发".to_string(), + )); + } + + let remote_config = + Self::normalize_remote_config(mutation, node.remote_config.as_ref()); + let remote_config = + optional_json_to_string(&remote_config, "proxy_nodes.remote_config")?; + let now = current_unix_secs() as i64; + let result = sqlx::query(UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL) + .bind(mutation.node_name.as_deref()) + .bind(remote_config) + .bind(now) + .bind(&mutation.node_id) + .bind(&node.tunnel_generation) + .bind(node.config_version) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + continue; + } + + let current = self.find_proxy_node(&mutation.node_id).await?; + return Ok( + current.filter(|current| current.tunnel_generation == node.tunnel_generation) + ); } - if let Some(node_name) = mutation.node_name.as_ref() { - node.name = node_name.clone(); - } - node.remote_config = Self::normalize_remote_config(mutation, node.remote_config.as_ref()); - node.config_version = node.config_version.saturating_add(1); - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(Some(node)) + + Err(DataLayerError::UnexpectedValue( + "proxy node remote config changed during every CAS retry".to_string(), + )) } async fn increment_manual_node_requests( @@ -1043,23 +1612,31 @@ WHERE is_manual = 0 failed_delta: i64, latency_ms: Option, ) -> Result<(), DataLayerError> { - let Some(mut node) = self.find_proxy_node(node_id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(node_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; return Ok(()); }; - if !node.is_manual { - return Ok(()); - } - if total_delta > 0 { - node.total_requests += total_delta; - } - if failed_delta > 0 { - node.failed_requests += failed_delta; - } - if let Some(ms) = latency_ms { - node.avg_latency_ms = Some(ms as f64); - } - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await + let generation = map_proxy_node_row(&row)?.tunnel_generation; + sqlx::query(INCREMENT_MANUAL_PROXY_NODE_REQUESTS_SQL) + .bind(total_delta) + .bind(failed_delta) + .bind(latency_ms.map(|value| value as f64)) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(()) } async fn cleanup_proxy_node_metrics( @@ -1150,6 +1727,76 @@ fn duplicate_proxy_node_error(node: &StoredProxyNode) -> DataLayerError { )) } +fn proxy_node_registration_matches( + current: &StoredProxyNode, + mutation: &ProxyNodeRegistrationMutation, + expected: &StoredProxyNode, + replacement_proxy_metadata: Option<&serde_json::Value>, + now: u64, +) -> bool { + current.id == expected.id + && current.tunnel_generation == expected.tunnel_generation + && !current.is_manual + && current.name == mutation.name + && current.ip == mutation.ip + && current.port == mutation.port + && current.region == mutation.region + && current.registered_by == mutation.registered_by + && current.last_heartbeat_at_unix_secs == Some(now) + && current.heartbeat_interval == mutation.heartbeat_interval + && mutation + .active_connections + .is_none_or(|value| current.active_connections == value) + && mutation + .total_requests + .is_none_or(|value| current.total_requests == value) + && mutation + .avg_latency_ms + .is_none_or(|value| current.avg_latency_ms == Some(value)) + && mutation + .hardware_info + .as_ref() + .is_none_or(|value| current.hardware_info.as_ref() == Some(value)) + && mutation + .estimated_max_concurrency + .is_none_or(|value| current.estimated_max_concurrency == Some(value)) + && current.tunnel_mode == mutation.tunnel_mode + && replacement_proxy_metadata + .is_none_or(|value| current.proxy_metadata.as_ref() == Some(value)) + && current.updated_at_unix_secs == Some(now) +} + +fn requested_proxy_node_id(value: Option<&str>) -> Result, DataLayerError> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_empty() || value.trim() != value { + return Err(DataLayerError::InvalidInput( + "proxy node id must be non-empty and unpadded".to_string(), + )); + } + Ok(Some(value.to_string())) +} + +fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node registration identity changed: requested {requested_id}, existing {existing_id}" + )) +} + +fn proxy_node_registration_changed_error() -> DataLayerError { + DataLayerError::UnexpectedValue( + "registered proxy node identity changed during registration".to_string(), + ) +} + +fn proxy_node_id_in_use_error(node: &StoredProxyNode) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node id is already in use: {} ({}:{})", + node.id, node.ip, node.port + )) +} + fn optional_json_from_string( value: Option, field_name: &str, @@ -1166,6 +1813,12 @@ fn optional_json_from_string( } fn map_proxy_node_row(row: &MySqlRow) -> Result { + let tunnel_generation: String = row.try_get("tunnel_generation").map_sql_err()?; + if tunnel_generation.trim().is_empty() { + return Err(DataLayerError::UnexpectedValue( + "proxy_nodes.tunnel_generation must not be empty".to_string(), + )); + } Ok(StoredProxyNode::new( row.try_get("id").map_sql_err()?, row.try_get("name").map_sql_err()?, @@ -1183,6 +1836,7 @@ fn map_proxy_node_row(row: &MySqlRow) -> Result row.try_get("tunnel_connected").map_sql_err()?, row.try_get("config_version").map_sql_err()?, )? + .with_tunnel_generation(tunnel_generation) .with_manual_proxy_fields( row.try_get("proxy_url").map_sql_err()?, row.try_get("proxy_username").map_sql_err()?, @@ -1278,6 +1932,13 @@ fn map_proxy_fleet_metric_row( #[cfg(test)] mod tests { use super::MysqlProxyNodeReadRepository; + use crate::run_migrations; + use aether_data_contracts::repository::proxy_nodes::{ + merge_proxy_metadata_for_registration, normalize_proxy_metadata, + ProxyNodeManualCreateMutation, ProxyNodeReadRepository, ProxyNodeRegistrationMutation, + ProxyNodeWriteRepository, + }; + use serde_json::json; #[tokio::test] async fn repository_builds_from_lazy_pool() { @@ -1289,4 +1950,248 @@ mod tests { let _repository = MysqlProxyNodeReadRepository::new(pool); } + + #[test] + fn proxy_node_mutation_sql_is_atomic_and_field_scoped() { + assert!(super::APPLY_HEARTBEAT_SQL + .contains("total_requests = total_requests + GREATEST(COALESCE(?, 0), 0)")); + assert!(super::APPLY_HEARTBEAT_SQL + .contains("failed_requests = failed_requests + GREATEST(COALESCE(?, 0), 0)")); + assert!(!super::APPLY_HEARTBEAT_SQL.contains("remote_config =")); + assert!(!super::APPLY_HEARTBEAT_SQL.contains("config_version =")); + assert!(super::APPLY_HEARTBEAT_SQL.contains("BINARY tunnel_generation = BINARY ?")); + + assert!(super::UPDATE_TUNNEL_STATUS_SQL + .contains("tunnel_connected_at IS NULL OR tunnel_connected_at <= ?")); + assert!(super::UPDATE_TUNNEL_STATUS_SQL.contains("BINARY tunnel_generation = BINARY ?")); + assert!(super::RECORD_PROXY_NODE_TRAFFIC_SQL + .contains("total_requests = total_requests + GREATEST(?, 0)")); + assert!(super::RECORD_PROXY_NODE_TRAFFIC_SQL.contains("is_manual = 1")); + assert!(super::UPDATE_MANUAL_PROXY_NODE_SQL.contains("name = COALESCE(?, name)")); + assert!(!super::UPDATE_MANUAL_PROXY_NODE_SQL.contains("remote_config")); + assert!(super::UPDATE_PROXY_NODE_REGISTRATION_SQL + .contains("BINARY id = BINARY ? AND BINARY tunnel_generation = BINARY ?")); + assert!(super::UPDATE_PROXY_NODE_REGISTRATION_SQL + .contains("is_manual = 0 AND BINARY ip = BINARY ? AND port = ?")); + assert!(super::UPDATE_PROXY_NODE_REGISTRATION_SQL + .contains("CAST(proxy_metadata AS JSON) = CAST(? AS JSON)")); + } + + #[tokio::test] + async fn proxy_metadata_cas_distinguishes_duplicate_array_elements_when_url_is_set() { + let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") + .ok() + .filter(|value| !value.trim().is_empty()) + else { + eprintln!( + "skipping mysql proxy metadata CAS test because AETHER_TEST_MYSQL_URL is unset" + ); + return; + }; + let pool = sqlx::mysql::MySqlPoolOptions::new() + .max_connections(2) + .connect(&database_url) + .await + .expect("mysql test pool should connect"); + run_migrations(&pool) + .await + .expect("mysql migrations should run"); + + let suffix = uuid::Uuid::new_v4().simple().to_string(); + let repository = MysqlProxyNodeReadRepository::new(pool.clone()); + let node = repository + .create_manual_node(&ProxyNodeManualCreateMutation { + node_id: None, + name: format!("metadata-cas-{suffix}"), + ip: format!("metadata-cas-{suffix}"), + port: 1, + region: None, + proxy_url: "http://127.0.0.1:1".to_string(), + proxy_username: None, + proxy_password: None, + registered_by: None, + }) + .await + .expect("mysql proxy fixture should insert"); + let stored = json!({"nested": {"values": [1, 1]}}); + sqlx::query("UPDATE proxy_nodes SET proxy_metadata = ? WHERE id = ?") + .bind(serde_json::to_string(&stored).expect("stored metadata should serialize")) + .bind(&node.id) + .execute(&pool) + .await + .expect("mysql proxy metadata fixture should update"); + + let updated = repository + .compare_and_set_proxy_metadata( + &node.id, + &json!({"nested": {"values": [1]}}), + &json!({"replacement": true}), + ) + .await + .expect("mysql proxy metadata CAS should execute"); + let persisted = repository + .find_proxy_node(&node.id) + .await + .expect("mysql proxy fixture should read") + .expect("mysql proxy fixture should exist") + .proxy_metadata; + let cleanup = sqlx::query("DELETE FROM proxy_nodes WHERE id = ?") + .bind(&node.id) + .execute(&pool) + .await; + + assert!(!updated, "different JSON arrays must not compare equal"); + assert_eq!(persisted, Some(stored)); + cleanup.expect("mysql proxy fixture should clean up"); + } + + #[tokio::test] + async fn registration_cas_rejects_stale_security_snapshot() { + let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") + .ok() + .filter(|value| !value.trim().is_empty()) + else { + eprintln!( + "skipping mysql registration CAS test because AETHER_TEST_MYSQL_URL is unset" + ); + return; + }; + let pool = sqlx::mysql::MySqlPoolOptions::new() + .max_connections(2) + .connect(&database_url) + .await + .expect("mysql test pool should connect"); + run_migrations(&pool) + .await + .expect("mysql migrations should run"); + + let suffix = uuid::Uuid::new_v4().simple().to_string(); + let node_id = format!("registration-security-{suffix}"); + let endpoint = format!("registration-security-{suffix}"); + let repository = MysqlProxyNodeReadRepository::new(pool.clone()); + let first = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some(node_id.clone()), + name: node_id.clone(), + ip: endpoint.clone(), + port: 7070, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + })), + proxy_version: Some("1.0.0".to_string()), + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("initial mysql registration should succeed"); + let stale = first.clone(); + + let rotated = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some(node_id.clone()), + name: node_id.clone(), + ip: endpoint.clone(), + port: 7070, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new" + } + })), + proxy_version: Some("2.0.0".to_string()), + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("mysql security rotation should succeed"); + + let stale_refresh = ProxyNodeRegistrationMutation { + node_id: Some(node_id.clone()), + name: node_id.clone(), + ip: endpoint, + port: 7070, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({"runtime": "stale-writer"})), + proxy_version: Some("2.1.0".to_string()), + registered_by: None, + tunnel_mode: true, + }; + let stale_replacement = merge_proxy_metadata_for_registration( + stale.proxy_metadata.as_ref(), + normalize_proxy_metadata( + stale_refresh.proxy_metadata.as_ref(), + stale_refresh.proxy_version.as_deref(), + ), + ); + assert!(!repository + .update_existing_registration_if_unchanged( + &stale_refresh, + &stale, + stale_replacement.as_ref(), + super::current_unix_secs(), + ) + .await + .expect("stale mysql registration CAS should execute")); + + let committed = repository + .register_node(&ProxyNodeRegistrationMutation { + name: format!("{node_id}-committed"), + proxy_metadata: Some(json!({"runtime": "committed-after-rotation"})), + ..stale_refresh + }) + .await + .expect("mysql metadata refresh should merge current security state"); + let cleanup = sqlx::query("DELETE FROM proxy_nodes WHERE id = ?") + .bind(&node_id) + .execute(&pool) + .await; + + assert_eq!( + rotated + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new") + ); + assert_eq!( + committed + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new") + ); + assert_eq!( + committed + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.get("runtime")), + Some(&json!("committed-after-rotation")) + ); + cleanup.expect("mysql proxy fixture should clean up"); + } } diff --git a/crates/aether-data/adapters/mysql/src/routing_profiles.rs b/crates/aether-data/adapters/mysql/src/routing_profiles.rs index 0d5b76e17..6b602d91e 100644 --- a/crates/aether-data/adapters/mysql/src/routing_profiles.rs +++ b/crates/aether-data/adapters/mysql/src/routing_profiles.rs @@ -15,6 +15,7 @@ SELECT description, enabled, is_system_default, + sort_order, config_json, version, created_at, @@ -61,10 +62,12 @@ impl MysqlRoutingGroupRepository { #[async_trait] impl RoutingGroupReadRepository for MysqlRoutingGroupRepository { async fn list_routing_groups(&self) -> Result, DataLayerError> { - let rows = sqlx::query(&format!("{ROUTING_GROUP_SELECT} ORDER BY name ASC, id ASC")) - .fetch_all(&self.pool) - .await - .map_sql_err()?; + let rows = sqlx::query(&format!( + "{ROUTING_GROUP_SELECT} ORDER BY enabled DESC, sort_order ASC, name ASC, id ASC" + )) + .fetch_all(&self.pool) + .await + .map_sql_err()?; rows.iter().map(map_group_row).collect() } @@ -173,10 +176,10 @@ impl RoutingGroupWriteRepository for MysqlRoutingGroupRepository { sqlx::query( r#" INSERT INTO routing_groups ( - id, name, description, enabled, is_system_default, config_json, + id, name, description, enabled, is_system_default, sort_order, config_json, version, created_at, updated_at, published_at ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&group.id) @@ -184,6 +187,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(json_to_string( &group.config_json, "routing_groups.config_json", @@ -241,6 +245,7 @@ SET name = ?, description = ?, enabled = ?, is_system_default = ?, + sort_order = ?, config_json = ?, version = ?, updated_at = ?, @@ -252,6 +257,7 @@ WHERE id = ? .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(json_to_string( &group.config_json, "routing_groups.config_json", @@ -460,6 +466,7 @@ fn map_group_row(row: &MySqlRow) -> Result { description: row.try_get("description").map_sql_err()?, enabled: row.try_get("enabled").map_sql_err()?, is_system_default: row.try_get("is_system_default").map_sql_err()?, + sort_order: row.try_get("sort_order").map_sql_err()?, config_json: json_from_string( row.try_get("config_json").map_sql_err()?, "routing_groups.config_json", diff --git a/crates/aether-data/adapters/mysql/src/settlement.rs b/crates/aether-data/adapters/mysql/src/settlement.rs index 02e0d8e93..adc141cd9 100644 --- a/crates/aether-data/adapters/mysql/src/settlement.rs +++ b/crates/aether-data/adapters/mysql/src/settlement.rs @@ -1,10 +1,14 @@ use async_trait::async_trait; -use sqlx::{mysql::MySqlRow, Row}; +use sqlx::{mysql::MySqlRow, Acquire, Row}; use aether_data_contracts::repository::settlement::{ finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd, - settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement, - UsageSettlementInput, SETTLEMENT_EPSILON_USD, + settlement_billing_status_for_usage_status, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation, + StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState, + UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD, }; use aether_data_contracts::DataLayerError; @@ -112,6 +116,142 @@ impl MysqlSettlementRepository { } } +fn usage_policy_cost_i64(value: u64, field: &str) -> Result { + i64::try_from(value) + .map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds the integer range"))) +} + +fn usage_policy_cost_u64(value: i64, field: &str) -> Result { + u64::try_from(value) + .map_err(|_| DataLayerError::UnexpectedValue(format!("{field} must not be negative"))) +} + +fn usage_policy_request_admission_from_mysql_row( + row: &MySqlRow, +) -> Result { + let state: String = row.try_get("state").map_sql_err()?; + Ok(StoredUsagePolicyRequestAdmission { + request_id: row.try_get("request_id").map_sql_err()?, + subject_id: row.try_get("subject_id").map_sql_err()?, + event_token: row.try_get("event_token").map_sql_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at_unix_secs").map_sql_err()?, + "usage policy request admitted_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until_unix_secs").map_sql_err()?, + "usage policy request retain_until", + )?, + state: UsagePolicyRequestAdmissionState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy request admission state {state}" + )) + })?, + released_at_unix_secs: row + .try_get::, _>("released_at_unix_secs") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy request released_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL: &str = r#" +SELECT request_id, subject_id, event_token, + admitted_at AS admitted_at_unix_secs, + retain_until AS retain_until_unix_secs, + state, released_at AS released_at_unix_secs +FROM usage_request_admissions +WHERE event_token = ? +FOR UPDATE +"#; + +const USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL: &str = + "SET TRANSACTION ISOLATION LEVEL READ COMMITTED"; + +fn usage_policy_cost_reservation_from_mysql_row( + row: &MySqlRow, +) -> Result { + let state: String = row.try_get("state").map_sql_err()?; + Ok(StoredUsagePolicyCostReservation { + request_id: row.try_get("request_id").map_sql_err()?, + subject_id: row.try_get("subject_id").map_sql_err()?, + reservation_token: row.try_get("reservation_token").map_sql_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at_unix_secs").map_sql_err()?, + "usage policy admitted_at", + )?, + reserved_cost_units: usage_policy_cost_u64( + row.try_get("reserved_cost_units").map_sql_err()?, + "usage policy reserved_cost_units", + )?, + actual_cost_units: row + .try_get::, _>("actual_cost_units") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy actual_cost_units")) + .transpose()?, + state: UsagePolicyCostReservationState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy reservation state {state}" + )) + })?, + reservation_expires_at_unix_secs: usage_policy_cost_u64( + row.try_get("reservation_expires_at_unix_secs") + .map_sql_err()?, + "usage policy reservation_expires_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until_unix_secs").map_sql_err()?, + "usage policy retain_until", + )?, + finalized_at_unix_secs: row + .try_get::, _>("finalized_at_unix_secs") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy finalized_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL: &str = r#" +SELECT + request_id, + subject_id, + reservation_token, + admitted_at AS admitted_at_unix_secs, + reserved_cost_units, + actual_cost_units, + state, + reservation_expires_at AS reservation_expires_at_unix_secs, + retain_until AS retain_until_unix_secs, + finalized_at AS finalized_at_unix_secs +FROM usage_cost_reservations +WHERE reservation_token = ? +FOR UPDATE +"#; + +async fn lock_usage_policy_subject_mysql( + tx: &mut sqlx::Transaction<'_, sqlx::MySql>, + subject_id: &str, +) -> Result { + let exists = sqlx::query_scalar::<_, String>( + r#" +SELECT id +FROM users +WHERE id = ? +FOR UPDATE + "#, + ) + .bind(subject_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .is_some(); + Ok(exists) +} + +fn usage_policy_subject_missing() -> DataLayerError { + DataLayerError::InvalidInput("usage policy subject does not exist".to_string()) +} + fn settlement_from_row(row: &MySqlRow) -> Result { Ok(StoredUsageSettlement { request_id: row.try_get("request_id").map_sql_err()?, @@ -205,6 +345,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -230,15 +371,27 @@ fn daily_quota_grants_from_entitlement( .and_then(serde_json::Value::as_str), now, )?, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + async fn consume_daily_quota_mysql( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, user_id: &str, @@ -248,18 +401,29 @@ async fn consume_daily_quota_mysql( wallet_can_overdraft: bool, now_unix_secs: i64, ) -> Result { - if total_cost_usd <= 0.0 { + if !total_cost_usd.is_finite() || total_cost_usd < 0.0 { + return Err(DataLayerError::InvalidInput( + "daily quota settlement cost must be finite and non-negative".to_string(), + )); + } + if total_cost_usd == 0.0 { return Ok(DailyQuotaDebitResult::default()); } let rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = ? - AND status = 'active' - AND starts_at <= ? - AND expires_at > ? -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = ? + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= ? + AND user_plan_entitlements.expires_at > ? +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC FOR UPDATE "#, ) @@ -280,9 +444,17 @@ FOR UPDATE "user_plan_entitlements.entitlements_snapshot invalid json: {err}" )) })?; + let plan_entitlements_raw: String = row.try_get("plan_entitlements_json").map_sql_err()?; + let plan_entitlements = serde_json::from_str::(&plan_entitlements_raw) + .map_err(|err| { + DataLayerError::UnexpectedValue(format!( + "billing_plans.entitlements_json invalid json: {err}" + )) + })?; grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -308,8 +480,18 @@ WHERE user_entitlement_id = ? .fetch_one(&mut **tx) .await .map_sql_err()?; + if !used.is_finite() || used < 0.0 { + return Err(DataLayerError::UnexpectedValue( + "daily quota usage ledger total is invalid".to_string(), + )); + } let remaining = (grant.daily_quota_usd - used).max(0.0); total_remaining += remaining; + if !total_remaining.is_finite() { + return Err(DataLayerError::UnexpectedValue( + "daily quota remaining total overflowed".to_string(), + )); + } grants_with_remaining.push((grant, remaining)); } let insufficient = (!allow_wallet_overage && total_remaining + 0.000_000_01 < total_cost_usd) @@ -359,6 +541,476 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) #[async_trait] impl SettlementWriteRepository for MysqlSettlementRepository { + async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result { + input.validate()?; + let admitted_at = usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy request admitted_at", + )?; + let retain_until = usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?; + let created_at = now_unix_secs()?; + // Different subjects lock different `users` rows. Under InnoDB's default REPEATABLE READ, + // two missing-token locking reads can retain compatible gap locks and then deadlock when + // both transactions try to insert the same unique event token. READ COMMITTED removes + // that gap-lock cycle while the subject row still serializes each subject's window count. + let mut connection = self.pool.acquire().await.map_sql_err()?; + sqlx::query(USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL) + .execute(&mut *connection) + .await + .map_sql_err()?; + let mut tx = connection.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + + let existing_row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL) + .bind(&input.event_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if let Some(row) = existing_row.as_ref() { + let existing = usage_policy_request_admission_from_mysql_row(row)?; + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at".to_string(), + )); + } + sqlx::query( + "UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?", + ) + .bind(retain_until) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + let outcome = match existing.state { + UsagePolicyRequestAdmissionState::Active => { + ReserveUsagePolicyRequestOutcome::Allowed + } + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }; + tx.commit().await.map_sql_err()?; + return Ok(outcome); + } + + for (window_index, window) in input.windows.iter().enumerate() { + let used_requests = sqlx::query_scalar::<_, i64>( + r#" +SELECT CAST(COUNT(*) AS SIGNED) +FROM usage_request_admissions +WHERE subject_id = ? + AND state = 'active' + AND admitted_at >= ? + AND admitted_at < ? + "#, + ) + .bind(&input.subject_id) + .bind(usage_policy_cost_i64( + window.starts_at_unix_secs, + "usage policy request window start", + )?) + .bind(usage_policy_cost_i64( + window.ends_at_unix_secs, + "usage policy request window end", + )?) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let used_requests = + usage_policy_cost_u64(used_requests, "usage policy request used_requests")?; + if used_requests >= window.limit_requests { + let outcome = ReserveUsagePolicyRequestOutcome::Rejected { + window_index, + limit_requests: window.limit_requests, + used_requests, + }; + tx.commit().await.map_sql_err()?; + return Ok(outcome); + } + } + + sqlx::query( + r#" +INSERT INTO usage_request_admissions ( + request_id, subject_id, event_token, admitted_at, retain_until, + state, released_at, created_at +) VALUES (?, ?, ?, ?, ?, 'active', NULL, ?) +ON DUPLICATE KEY UPDATE event_token = VALUES(event_token) + "#, + ) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.event_token) + .bind(admitted_at) + .bind(retain_until) + .bind(created_at) + .execute(&mut *tx) + .await + .map_sql_err()?; + + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL) + .bind(&input.event_token) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let existing = usage_policy_request_admission_from_mysql_row(&row)?; + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at".to_string(), + )); + } + sqlx::query( + "UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?", + ) + .bind(retain_until) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + let outcome = match existing.state { + UsagePolicyRequestAdmissionState::Active => ReserveUsagePolicyRequestOutcome::Allowed, + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }; + tx.commit().await.map_sql_err()?; + Ok(outcome) + } + + async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, DataLayerError> { + input.validate()?; + let released_at = usage_policy_cost_i64( + input.released_at_unix_secs, + "usage policy request released_at", + )?; + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL) + .bind(&input.event_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.commit().await.map_sql_err()?; + return Ok(None); + }; + let mut admission = usage_policy_request_admission_from_mysql_row(&row)?; + if admission.request_id != input.request_id || admission.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + if input.released_at_unix_secs < admission.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy released_at must not precede admitted_at".to_string(), + )); + } + if admission.state == UsagePolicyRequestAdmissionState::Active { + sqlx::query( + "UPDATE usage_request_admissions SET state = 'released', released_at = ? WHERE event_token = ? AND state = 'active'", + ) + .bind(released_at) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + admission.state = UsagePolicyRequestAdmissionState::Released; + admission.released_at_unix_secs = Some(input.released_at_unix_secs); + } + tx.commit().await.map_sql_err()?; + Ok(Some(admission)) + } + + async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy request cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_request_admissions +WHERE retain_until <= ? +ORDER BY retain_until, event_token +LIMIT ? + "#, + ) + .bind(now) + .bind(limit) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() as usize) + } + + async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result { + input.validate()?; + let admitted_at = + usage_policy_cost_i64(input.admitted_at_unix_secs, "usage policy admitted_at")?; + let reservation_expires_at = usage_policy_cost_i64( + input.reservation_expires_at_unix_secs, + "usage policy reservation_expires_at", + )?; + let updated_at = now_unix_secs()?; + + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + let existing_row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let existing = existing_row + .as_ref() + .map(usage_policy_cost_reservation_from_mysql_row) + .transpose()?; + if let Some(existing) = existing.as_ref() { + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyCostOutcome::Conflict); + } + if existing.state != UsagePolicyCostReservationState::Reserved { + let outcome = ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: existing.state, + }; + tx.commit().await.map_sql_err()?; + return Ok(outcome); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy reservation_token must keep its original admitted_at".to_string(), + )); + } + } + + let previous_reserved_cost_units = existing + .as_ref() + .map(|reservation| reservation.reserved_cost_units) + .unwrap_or(0); + let target_reserved_cost_units = + previous_reserved_cost_units.max(input.reserved_cost_units); + for (window_index, window) in input.windows.iter().enumerate() { + let window_start = + usage_policy_cost_i64(window.starts_at_unix_secs, "usage policy window start")?; + let window_end = + usage_policy_cost_i64(window.ends_at_unix_secs, "usage policy window end")?; + let used_cost_units = sqlx::query_scalar::<_, i64>( + r#" +SELECT CAST(COALESCE(SUM( + CASE + WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) + WHEN state = 'reserved' AND reservation_expires_at > ? THEN reserved_cost_units + ELSE 0 + END +), 0) AS SIGNED) +FROM usage_cost_reservations +WHERE subject_id = ? + AND admitted_at >= ? + AND admitted_at < ? + AND reservation_token <> ? + "#, + ) + .bind(admitted_at) + .bind(&input.subject_id) + .bind(window_start) + .bind(window_end) + .bind(&input.reservation_token) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let used_cost_units = + usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?; + if used_cost_units + .checked_add(target_reserved_cost_units) + .is_none_or(|total| total > window.limit_cost_units) + { + let outcome = ReserveUsagePolicyCostOutcome::Rejected { + window_index, + limit_cost_units: window.limit_cost_units, + used_cost_units, + }; + tx.commit().await.map_sql_err()?; + return Ok(outcome); + } + } + + let admitted_at = existing + .as_ref() + .map(|reservation| { + usage_policy_cost_i64( + reservation.admitted_at_unix_secs, + "usage policy admitted_at", + ) + }) + .transpose()? + .unwrap_or(admitted_at); + sqlx::query( + r#" +INSERT INTO usage_cost_reservations ( + request_id, subject_id, reservation_token, admitted_at, + reserved_cost_units, actual_cost_units, + state, reservation_expires_at, retain_until, finalized_at, created_at, updated_at +) +VALUES (?, ?, ?, ?, ?, NULL, 'reserved', ?, ?, NULL, ?, ?) +ON DUPLICATE KEY UPDATE + reserved_cost_units = GREATEST(reserved_cost_units, VALUES(reserved_cost_units)), + reservation_expires_at = GREATEST( + reservation_expires_at, + VALUES(reservation_expires_at) + ), + retain_until = GREATEST(retain_until, VALUES(retain_until)), + updated_at = VALUES(updated_at) + "#, + ) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.reservation_token) + .bind(admitted_at) + .bind(usage_policy_cost_i64( + target_reserved_cost_units, + "usage policy reserved_cost_units", + )?) + .bind(reservation_expires_at) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy retain_until", + )?) + .bind(updated_at) + .bind(updated_at) + .execute(&mut *tx) + .await + .map_sql_err()?; + + tx.commit().await.map_sql_err()?; + Ok(ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: target_reserved_cost_units, + additional_reserved_cost_units: target_reserved_cost_units + .saturating_sub(previous_reserved_cost_units), + }) + } + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + input.validate()?; + let actual_cost_units = + usage_policy_cost_i64(input.actual_cost_units, "usage policy actual_cost_units")?; + let finalized_at = + usage_policy_cost_i64(input.finalized_at_unix_secs, "usage policy finalized_at")?; + let updated_at = now_unix_secs()?; + + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.commit().await.map_sql_err()?; + return Ok(None); + }; + let mut reservation = usage_policy_cost_reservation_from_mysql_row(&row)?; + if reservation.request_id != input.request_id || reservation.subject_id != input.subject_id + { + // The token selects the row; audit identity must still match before the reservation + // can be finalized. + tx.commit().await.map_sql_err()?; + return Ok(None); + } + if reservation.state == UsagePolicyCostReservationState::Reserved { + sqlx::query( + r#" +UPDATE usage_cost_reservations +SET state = ?, + actual_cost_units = ?, + finalized_at = ?, + updated_at = ? +WHERE reservation_token = ? + AND request_id = ? + AND subject_id = ? + AND state = 'reserved' + "#, + ) + .bind(input.terminal_state.as_str()) + .bind(actual_cost_units) + .bind(finalized_at) + .bind(updated_at) + .bind(&input.reservation_token) + .bind(&input.request_id) + .bind(&input.subject_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + reservation.state = input.terminal_state; + reservation.actual_cost_units = Some(input.actual_cost_units); + reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs); + } + + tx.commit().await.map_sql_err()?; + Ok(Some(reservation)) + } + + async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_cost_reservations +WHERE retain_until <= ? +ORDER BY retain_until, reservation_token +LIMIT ? + "#, + ) + .bind(now) + .bind(limit) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() as usize) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -438,7 +1090,7 @@ LIMIT 1 let wallet_row = if let Some(api_key_id) = api_key_id { sqlx::query( r#" -SELECT id, balance, gift_balance, limit_mode +SELECT id, balance, gift_balance, total_consumed, limit_mode FROM wallets WHERE api_key_id = ? LIMIT 1 @@ -459,7 +1111,7 @@ FOR UPDATE if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) { sqlx::query( r#" -SELECT id, balance, gift_balance, limit_mode +SELECT id, balance, gift_balance, total_consumed, limit_mode FROM wallets WHERE user_id = ? LIMIT 1 @@ -480,14 +1132,20 @@ FOR UPDATE let wallet_can_overdraft = wallet_row.is_some(); let wallet_available_usd = match wallet_row.as_ref() { Some(row) => { + let recharge_balance: f64 = row.try_get("balance").map_sql_err()?; + let gift_balance: f64 = row.try_get("gift_balance").map_sql_err()?; + let total_consumed: f64 = row.try_get("total_consumed").map_sql_err()?; + validate_wallet_settlement_values( + recharge_balance, + gift_balance, + total_consumed, + 0.0, + )?; let limit_mode: String = row.try_get("limit_mode").map_sql_err()?; if limit_mode.eq_ignore_ascii_case("unlimited") { None } else { - Some(finite_wallet_available_usd( - row.try_get("balance").map_sql_err()?, - row.try_get("gift_balance").map_sql_err()?, - )) + Some(finite_wallet_available_usd(recharge_balance, gift_balance)) } } None => Some(0.0), @@ -566,6 +1224,7 @@ FOR UPDATE let wallet_id: String = wallet_row.try_get("id").map_sql_err()?; let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?; let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?; + let total_consumed: f64 = wallet_row.try_get("total_consumed").map_sql_err()?; let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?; let before_total = before_recharge + before_gift; let mut after_recharge = before_recharge; @@ -579,6 +1238,13 @@ FOR UPDATE (after_recharge, after_gift) = debit_plan.after_balances(before_recharge, before_gift); } + let total_consumed_after = total_consumed + wallet_debit_cost_usd; + validate_wallet_settlement_values( + after_recharge, + after_gift, + total_consumed_after, + 0.0, + )?; if final_billing_status == "settled" { sqlx::query( r#" @@ -586,14 +1252,14 @@ UPDATE wallets SET balance = ?, gift_balance = ?, - total_consumed = COALESCE(total_consumed, 0) + ?, + total_consumed = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) .bind(after_gift) - .bind(wallet_debit_cost_usd) + .bind(total_consumed_after) .bind(updated_at) .bind(&wallet_id) .execute(&mut *tx) @@ -692,10 +1358,11 @@ WHERE id = ? #[cfg(test)] mod tests { - use super::MysqlSettlementRepository; + use super::{MysqlSettlementRepository, USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL}; use crate::run_migrations; use aether_data_contracts::repository::settlement::{ - SettlementWriteRepository, UsageSettlementInput, + ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome, + SettlementWriteRepository, UsagePolicyRequestWindow, UsageSettlementInput, }; #[tokio::test] @@ -709,6 +1376,89 @@ mod tests { let _repository = MysqlSettlementRepository::new(pool); } + #[test] + fn request_admission_transactions_use_read_committed() { + assert_eq!( + USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL, + "SET TRANSACTION ISOLATION LEVEL READ COMMITTED" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn cross_subject_same_token_is_allowed_once_without_deadlock_when_url_is_set() { + let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") + .ok() + .filter(|value| !value.trim().is_empty()) + else { + eprintln!( + "skipping mysql request admission race test because AETHER_TEST_MYSQL_URL is unset" + ); + return; + }; + let pool = sqlx::mysql::MySqlPoolOptions::new() + .max_connections(2) + .connect(&database_url) + .await + .expect("mysql pool should connect"); + run_migrations(&pool) + .await + .expect("mysql migrations should run"); + cleanup_request_admission_rows(&pool).await; + sqlx::query( + r#" +INSERT INTO users (id, username, auth_source, created_at, updated_at) +VALUES + ('admission-race-user-1', 'admission-race-user-1', 'local', 1, 1), + ('admission-race-user-2', 'admission-race-user-2', 'local', 1, 1) + "#, + ) + .execute(&pool) + .await + .expect("race users should seed"); + + let repository = MysqlSettlementRepository::new(pool.clone()); + let reserve = |request_id: &str, subject_id: &str| ReserveUsagePolicyRequestInput { + request_id: request_id.to_string(), + subject_id: subject_id.to_string(), + event_token: "admission-race-token".to_string(), + admitted_at_unix_secs: 100, + retain_until_unix_secs: 1_000, + windows: vec![UsagePolicyRequestWindow { + starts_at_unix_secs: 0, + ends_at_unix_secs: 1_000, + limit_requests: 10, + }], + }; + let (first, second) = tokio::join!( + repository.reserve_usage_policy_request(reserve( + "admission-race-request-1", + "admission-race-user-1" + )), + repository.reserve_usage_policy_request(reserve( + "admission-race-request-2", + "admission-race-user-2" + )) + ); + let mut outcomes = vec![ + first.expect("first reserve should not deadlock"), + second.expect("second reserve should not deadlock"), + ]; + outcomes.sort_by_key(|outcome| match outcome { + ReserveUsagePolicyRequestOutcome::Allowed => 0, + ReserveUsagePolicyRequestOutcome::Conflict => 1, + _ => 2, + }); + assert_eq!( + outcomes, + vec![ + ReserveUsagePolicyRequestOutcome::Allowed, + ReserveUsagePolicyRequestOutcome::Conflict, + ] + ); + + cleanup_request_admission_rows(&pool).await; + } + #[tokio::test] async fn mysql_repository_settles_once_and_enqueues_provider_delta_when_url_is_set() { let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") @@ -836,4 +1586,19 @@ WHERE request_id = 'settlement-request-1' .expect("settlement cleanup should succeed"); } } + + async fn cleanup_request_admission_rows(pool: &sqlx::MySqlPool) { + sqlx::query( + "DELETE FROM usage_request_admissions WHERE event_token = 'admission-race-token'", + ) + .execute(pool) + .await + .expect("admission race row cleanup should succeed"); + sqlx::query( + "DELETE FROM users WHERE id IN ('admission-race-user-1', 'admission-race-user-2')", + ) + .execute(pool) + .await + .expect("admission race user cleanup should succeed"); + } } diff --git a/crates/aether-data/adapters/mysql/src/usage.rs b/crates/aether-data/adapters/mysql/src/usage.rs index 20869e968..2b851f1bf 100644 --- a/crates/aether-data/adapters/mysql/src/usage.rs +++ b/crates/aether-data/adapters/mysql/src/usage.rs @@ -6,7 +6,9 @@ use async_trait::async_trait; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::usage::{ - strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure, + sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence, + sanitize_usage_request_metadata, usage_can_recover_terminal_failure, + usage_error_category_for_status_code, usage_lifecycle_update_allowed, usage_request_metadata_client_family, PendingUsageCleanupSummary, StoredRequestUsageAudit, StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardSummary, StoredUsageUserTotals, UpsertUsageRecord, UsageCleanupExecutionMode, UsageCleanupPreviewCounts, @@ -437,7 +439,6 @@ ON DUPLICATE KEY UPDATE const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#" SELECT `usage`.request_id, - `usage`.status, COALESCE(usage_settlement_snapshots.billing_status, `usage`.billing_status) AS billing_status FROM `usage` LEFT JOIN usage_settlement_snapshots @@ -1046,11 +1047,29 @@ impl UsageWriteRepository for MysqlUsageWriteRepository { &self, usage: UpsertUsageRecord, ) -> Result { - let mut usage = strip_deprecated_usage_display_fields(usage); usage.validate()?; - let prepared_capture = http_capture::prepare_usage_http_capture(&mut usage)?; + // Auxiliary tables may receive only clear tombstones, never request or response content. + let capture_usage = usage.clone(); + let mut usage = sanitize_usage_for_persistence(usage); + usage.validate()?; let mut tx = self.pool.begin().await.map_sql_err()?; let existing = counters::lock_and_load_usage(&mut tx, &usage.request_id).await?; + if let Some(existing) = existing.as_ref() { + if !usage_lifecycle_update_allowed( + &existing.status, + &existing.billing_status, + existing.updated_at_unix_secs, + existing.finalized_at_unix_secs, + &usage.status, + &usage.billing_status, + usage.updated_at_unix_secs, + usage.finalized_at_unix_secs, + ) { + let existing = existing.clone(); + tx.rollback().await.map_sql_err()?; + return http_capture::hydrate_usage_body_refs(&self.pool, existing).await; + } + } let recovers_terminal_failure = existing.as_ref().is_some_and(|existing| { usage_can_recover_terminal_failure( &existing.status, @@ -1069,13 +1088,18 @@ impl UsageWriteRepository for MysqlUsageWriteRepository { } } + let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage); + let prepared_capture = http_capture::prepare_usage_http_capture(&mut capture_usage)?; let capture_update_allowed = recovers_terminal_failure || http_capture::capture_update_allowed(existing.as_ref(), &usage.status); if capture_update_allowed { - http_capture::apply_previous_metadata_tombstones(&mut usage, existing.as_ref()); + http_capture::apply_previous_metadata_tombstones(&mut capture_usage, existing.as_ref()); + usage.request_metadata = + sanitize_usage_request_metadata(capture_usage.request_metadata.clone()); } let prepared_snapshots = capture_update_allowed - .then(|| snapshots::from_usage(&usage)) + // The control projection preserves safe typed routing and allow-listed billing facts. + .then(|| snapshots::from_usage(&capture_usage)) .transpose()?; bind_upsert(sqlx::query(UPSERT_USAGE_SQL), &usage)? .execute(&mut *tx) @@ -1230,7 +1254,7 @@ SET provider_api_keys.request_count = aggregated.request_count, &self, cutoff_unix_secs: u64, now_unix_secs: u64, - timeout_minutes: u64, + _timeout_minutes: u64, batch_size: usize, ) -> Result { if batch_size == 0 { @@ -1264,7 +1288,6 @@ SET provider_api_keys.request_count = aggregated.request_count, .map(|row| { Ok(StalePendingUsageRow { request_id: row.try_get("request_id").map_sql_err()?, - status: row.try_get("status").map_sql_err()?, billing_status: row.try_get("billing_status").map_sql_err()?, }) }) @@ -1280,7 +1303,8 @@ SET provider_api_keys.request_count = aggregated.request_count, UPDATE `usage` SET status = 'completed', status_code = 200, - error_message = NULL + error_message = NULL, + error_category = NULL WHERE request_id = ? "#, ) @@ -1308,11 +1332,8 @@ WHERE request_id = ? let candidate_info = latest_failed_candidate_mysql(&mut tx, &row.request_id).await?; - let (status_code, error_message) = resolve_stale_pending_failure( - candidate_info.as_ref(), - &row.status, - timeout_minutes, - ); + let status_code = resolve_stale_pending_status_code(candidate_info.as_ref()); + let error_category = usage_error_category_for_status_code(status_code); let status_code_i64 = i64::from(status_code); if row.billing_status == "pending" { sqlx::query( @@ -1320,7 +1341,8 @@ WHERE request_id = ? UPDATE `usage` SET status = 'failed', status_code = ?, - error_message = ?, + error_message = NULL, + error_category = ?, billing_status = 'void', finalized_at = ?, total_cost_usd = 0, @@ -1329,7 +1351,7 @@ WHERE request_id = ? "#, ) .bind(status_code_i64) - .bind(&error_message) + .bind(error_category) .bind(to_i64(now_unix_secs, "usage finalized_at")?) .bind(&row.request_id) .execute(&mut *tx) @@ -1347,12 +1369,13 @@ WHERE request_id = ? UPDATE `usage` SET status = 'failed', status_code = ?, - error_message = ? + error_message = NULL, + error_category = ? WHERE request_id = ? "#, ) .bind(status_code_i64) - .bind(&error_message) + .bind(error_category) .bind(&row.request_id) .execute(&mut *tx) .await @@ -1364,7 +1387,8 @@ WHERE request_id = ? UPDATE request_candidates SET status = 'failed', finished_at = ?, - error_message = '请求超时(服务器可能已重启)' + error_type = 'internal', + error_message = NULL WHERE request_id = ? AND status IN ('pending', 'streaming') "#, @@ -1451,7 +1475,6 @@ WHERE request_id = ? struct StalePendingUsageRow { request_id: String, - status: String, billing_status: String, } @@ -1561,29 +1584,14 @@ ON DUPLICATE KEY UPDATE Ok(()) } -fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String { - format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成") -} - struct FailedCandidateCleanupInfo { status_code: Option, - error_message: Option, } -fn resolve_stale_pending_failure( - candidate: Option<&FailedCandidateCleanupInfo>, - status: &str, - timeout_minutes: u64, -) -> (u16, String) { - match candidate { - Some(info) => ( - info.status_code.unwrap_or(502), - info.error_message - .clone() - .unwrap_or_else(|| stale_pending_error_message(status, timeout_minutes)), - ), - None => (504, stale_pending_error_message(status, timeout_minutes)), - } +fn resolve_stale_pending_status_code(candidate: Option<&FailedCandidateCleanupInfo>) -> u16 { + candidate + .and_then(|info| info.status_code) + .unwrap_or(if candidate.is_some() { 502 } else { 504 }) } async fn latest_failed_candidate_mysql( @@ -1592,7 +1600,7 @@ async fn latest_failed_candidate_mysql( ) -> Result, DataLayerError> { let row = sqlx::query( r#" -SELECT status_code, error_message +SELECT status_code FROM request_candidates WHERE request_id = ? AND status IN ('failed', 'cancelled') @@ -1615,15 +1623,7 @@ LIMIT 1 .try_get::, _>("status_code") .map_sql_err()? .and_then(|value| u16::try_from(value).ok()); - let error_message = row - .try_get::, _>("error_message") - .map_sql_err()? - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - Ok(Some(FailedCandidateCleanupInfo { - status_code, - error_message, - })) + Ok(Some(FailedCandidateCleanupInfo { status_code })) } fn bind_upsert<'q>( diff --git a/crates/aether-data/adapters/mysql/src/usage/cleanup.rs b/crates/aether-data/adapters/mysql/src/usage/cleanup.rs index fc98ee971..81aea0694 100644 --- a/crates/aether-data/adapters/mysql/src/usage/cleanup.rs +++ b/crates/aether-data/adapters/mysql/src/usage/cleanup.rs @@ -1,11 +1,8 @@ -use std::io::Write; - use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, UsageBodyField, UsageCleanupExecutionMode, - UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow, + UsageCleanupExecutionMode, UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, + UsageCleanupWindow, }; use chrono::{DateTime, Utc}; -use flate2::{write::GzEncoder, Compression}; use serde_json::Value; use sqlx::Row; use tracing::warn; @@ -66,17 +63,6 @@ OR EXISTS ( ) "#; -const INLINE_OR_COMPRESSED_BODY_PREDICATE: &str = r#" -request_body IS NOT NULL -OR response_body IS NOT NULL -OR provider_request_body IS NOT NULL -OR client_response_body IS NOT NULL -OR request_body_compressed IS NOT NULL -OR response_body_compressed IS NOT NULL -OR provider_request_body_compressed IS NOT NULL -OR client_response_body_compressed IS NOT NULL -"#; - const HEADER_PREDICATE: &str = r#" request_headers IS NOT NULL OR response_headers IS NOT NULL @@ -116,6 +102,20 @@ OR request_body_compressed IS NOT NULL OR response_body_compressed IS NOT NULL OR provider_request_body_compressed IS NOT NULL OR client_response_body_compressed IS NOT NULL +OR EXISTS ( + SELECT 1 FROM usage_body_blobs + WHERE usage_body_blobs.request_id = `usage`.request_id +) +OR EXISTS ( + SELECT 1 FROM usage_http_audits + WHERE usage_http_audits.request_id = `usage`.request_id + AND ( + usage_http_audits.request_body_ref IS NOT NULL + OR usage_http_audits.provider_request_body_ref IS NOT NULL + OR usage_http_audits.response_body_ref IS NOT NULL + OR usage_http_audits.client_response_body_ref IS NOT NULL + ) +) OR ( request_metadata IS NOT NULL AND JSON_VALID(request_metadata) @@ -136,43 +136,6 @@ struct CleanupRow { request_id: String, } -#[derive(Debug)] -struct BodyRow { - id: String, - request_id: String, - request_body: Option, - request_body_compressed: Option>, - provider_request_body: Option, - provider_request_body_compressed: Option>, - response_body: Option, - response_body_compressed: Option>, - client_response_body: Option, - client_response_body_compressed: Option>, -} - -#[derive(Debug, Default)] -struct DetachedRefs { - request_body_ref: Option, - provider_request_body_ref: Option, - response_body_ref: Option, - client_response_body_ref: Option, -} - -impl DetachedRefs { - fn any_present(&self) -> bool { - self.request_body_ref.is_some() - || self.provider_request_body_ref.is_some() - || self.response_body_ref.is_some() - || self.client_response_body_ref.is_some() - } -} - -struct DetachedBlob { - body_ref: String, - body_field: &'static str, - payload_gzip: Vec, -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum BodyCleanupKind { Raw, @@ -188,10 +151,6 @@ impl BodyCleanupKind { Self::All => ALL_BODY_PREDICATE, } } - - fn clears_detached(self) -> bool { - self != Self::Raw - } } pub(crate) async fn cleanup_usage( @@ -263,12 +222,19 @@ pub(crate) async fn cleanup_usage( }; let detail_newer_than = detail_body_newer_than(window, targets); let legacy_body_refs_migrated = if targets.detail_body { - migrate_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await? + purge_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await? } else { 0 }; let body_externalized = if targets.detail_body { - externalize_detail_bodies(pool, window.detail_cutoff, detail_newer_than, batch_size).await? + cleanup_body_fields( + pool, + window.detail_cutoff, + detail_newer_than, + batch_size, + BodyCleanupKind::All, + ) + .await? } else { 0 }; @@ -291,6 +257,8 @@ pub(crate) async fn cleanup_usage( header_cleaned, keys_cleaned, records_deleted, + cost_reservations_deleted: 0, + request_admissions_deleted: 0, }) } @@ -611,14 +579,13 @@ WHERE id = ? .execute(&mut *tx) .await .map_sql_err()?; - if kind.clears_detached() { - sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") - .bind(&row.request_id) - .execute(&mut *tx) - .await - .map_sql_err()?; - sqlx::query( - r#" + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(&row.request_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" UPDATE usage_http_audits SET request_body_ref = NULL, provider_request_body_ref = NULL, @@ -628,13 +595,12 @@ SET request_body_ref = NULL, updated_at = UNIX_TIMESTAMP() WHERE request_id = ? "#, - ) - .bind(&row.request_id) - .execute(&mut *tx) - .await - .map_sql_err()?; - delete_empty_http_audit(&mut tx, &row.request_id).await?; - } + ) + .bind(&row.request_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + delete_empty_http_audit(&mut tx, &row.request_id).await?; } tx.commit().await.map_sql_err()?; total = total.saturating_add(row_count); @@ -670,14 +636,14 @@ WHERE request_id = ? Ok(()) } -async fn migrate_legacy_body_refs( +async fn purge_legacy_body_refs( pool: &MysqlPool, cutoff: DateTime, newer_than: Option>, batch_size: usize, ) -> Result { if invalid_window(cutoff, newer_than) { - warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref migration skipped due to invalid window"); + warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref purge skipped due to invalid window"); return Ok(0); } let mut total = 0usize; @@ -695,7 +661,7 @@ async fn migrate_legacy_body_refs( } let row_count = rows.len(); let mut tx = pool.begin().await.map_sql_err()?; - let mut migrated = 0usize; + let mut purged = 0usize; for row in rows { let metadata: Option = sqlx::query_scalar("SELECT request_metadata FROM `usage` WHERE id = ? LIMIT 1") @@ -704,14 +670,9 @@ async fn migrate_legacy_body_refs( .await .map_sql_err()? .flatten(); - let Some((refs, metadata)) = - legacy_body_ref_plan(&row.request_id, metadata.as_deref())? - else { + let Some(metadata) = legacy_body_ref_purge_plan(metadata.as_deref())? else { continue; }; - if refs.any_present() { - upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?; - } let updated = sqlx::query( r#" UPDATE `usage` @@ -726,23 +687,23 @@ WHERE id = ? .await .map_sql_err()? .rows_affected(); + purge_detached_body_capture(&mut tx, &row.request_id).await?; if updated > 0 { - migrated += 1; + purged += 1; } } tx.commit().await.map_sql_err()?; - total = total.saturating_add(migrated); - if row_count < batch_size || migrated == 0 { + total = total.saturating_add(purged); + if row_count < batch_size || purged == 0 { break; } } Ok(total) } -fn legacy_body_ref_plan( - request_id: &str, +fn legacy_body_ref_purge_plan( metadata: Option<&str>, -) -> Result)>, DataLayerError> { +) -> Result>, DataLayerError> { let Some(metadata) = metadata else { return Ok(None); }; @@ -752,30 +713,16 @@ fn legacy_body_ref_plan( let Value::Object(mut object) = value else { return Ok(None); }; - let mut refs = DetachedRefs::default(); let mut removed = false; - for field in [ - UsageBodyField::RequestBody, - UsageBodyField::ProviderRequestBody, - UsageBodyField::ResponseBody, - UsageBodyField::ClientResponseBody, + for key in [ + "request_body_ref", + "provider_request_body_ref", + "response_body_ref", + "client_response_body_ref", ] { - let Some(value) = object.remove(field.as_ref_key()) else { - continue; - }; - removed = true; - let parsed = value - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| { - usage_body_ref(&parsed_request_id, parsed_field) - }); - set_ref(&mut refs, field, parsed); + if object.remove(key).is_some() { + removed = true; + } } if !removed { return Ok(None); @@ -791,277 +738,35 @@ fn legacy_body_ref_plan( })?, ) }; - Ok(Some((refs, metadata))) + Ok(Some(metadata)) } -async fn externalize_detail_bodies( - pool: &MysqlPool, - cutoff: DateTime, - newer_than: Option>, - batch_size: usize, -) -> Result { - if invalid_window(cutoff, newer_than) { - warn!(%cutoff, ?newer_than, "MySQL usage body externalization skipped due to invalid window"); - return Ok(0); - } - let batch_size = batch_size.clamp(1, 25); - let mut total = 0usize; - loop { - let rows = fetch_body_rows(pool, cutoff, newer_than, batch_size).await?; - if rows.is_empty() { - break; - } - let row_count = rows.len(); - let mut externalized = 0usize; - for row in rows { - let (blobs, refs) = build_detached_bodies(&row)?; - let mut tx = pool.begin().await.map_sql_err()?; - for blob in blobs { - sqlx::query( - r#" -INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip) -VALUES (?, ?, ?, ?) -ON DUPLICATE KEY UPDATE - request_id = VALUES(request_id), - body_field = VALUES(body_field), - payload_gzip = VALUES(payload_gzip), - updated_at = UNIX_TIMESTAMP() -"#, - ) - .bind(blob.body_ref) - .bind(&row.request_id) - .bind(blob.body_field) - .bind(blob.payload_gzip) - .execute(&mut *tx) - .await - .map_sql_err()?; - } - if refs.any_present() { - upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?; - } - let updated = sqlx::query( - r#" -UPDATE `usage` -SET request_body = NULL, - response_body = NULL, - provider_request_body = NULL, - client_response_body = NULL, - request_body_compressed = NULL, - response_body_compressed = NULL, - provider_request_body_compressed = NULL, - client_response_body_compressed = NULL -WHERE id = ? -"#, - ) - .bind(row.id) - .execute(&mut *tx) - .await - .map_sql_err()? - .rows_affected(); - tx.commit().await.map_sql_err()?; - if updated > 0 { - externalized += 1; - } - } - total = total.saturating_add(externalized); - if row_count < batch_size || externalized == 0 { - break; - } - } - Ok(total) -} - -async fn fetch_body_rows( - pool: &MysqlPool, - cutoff: DateTime, - newer_than: Option>, - batch_size: usize, -) -> Result, DataLayerError> { - let newer_than = newer_than.map(|value| value.timestamp()); - let sql = format!( - r#" -SELECT id, - request_id, - CAST(request_body AS CHAR) AS request_body, - request_body_compressed, - CAST(provider_request_body AS CHAR) AS provider_request_body, - provider_request_body_compressed, - CAST(response_body AS CHAR) AS response_body, - response_body_compressed, - CAST(client_response_body AS CHAR) AS client_response_body, - client_response_body_compressed -FROM `usage` -WHERE created_at_unix_ms < ? - AND (? IS NULL OR created_at_unix_ms >= ?) - AND ({INLINE_OR_COMPRESSED_BODY_PREDICATE}) -ORDER BY created_at_unix_ms ASC, id ASC -LIMIT ? -"# - ); - sqlx::query(&sql) - .bind(cutoff.timestamp()) - .bind(newer_than) - .bind(newer_than) - .bind(i64::try_from(batch_size).unwrap_or(i64::MAX)) - .fetch_all(pool) - .await - .map_sql_err()? - .into_iter() - .map(|row| { - Ok(BodyRow { - id: row.try_get("id").map_sql_err()?, - request_id: row.try_get("request_id").map_sql_err()?, - request_body: parse_optional_json(row.try_get("request_body").map_sql_err()?)?, - request_body_compressed: row.try_get("request_body_compressed").map_sql_err()?, - provider_request_body: parse_optional_json( - row.try_get("provider_request_body").map_sql_err()?, - )?, - provider_request_body_compressed: row - .try_get("provider_request_body_compressed") - .map_sql_err()?, - response_body: parse_optional_json(row.try_get("response_body").map_sql_err()?)?, - response_body_compressed: row.try_get("response_body_compressed").map_sql_err()?, - client_response_body: parse_optional_json( - row.try_get("client_response_body").map_sql_err()?, - )?, - client_response_body_compressed: row - .try_get("client_response_body_compressed") - .map_sql_err()?, - }) - }) - .collect() -} - -fn parse_optional_json(raw: Option) -> Result, DataLayerError> { - raw.map(|raw| { - serde_json::from_str(&raw).map_err(|err| { - DataLayerError::UnexpectedValue(format!("invalid inline usage body JSON: {err}")) - }) - }) - .transpose() -} - -fn build_detached_bodies( - row: &BodyRow, -) -> Result<(Vec, DetachedRefs), DataLayerError> { - let mut blobs = Vec::new(); - let mut refs = DetachedRefs::default(); - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::RequestBody, - row.request_body.as_ref(), - row.request_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ProviderRequestBody, - row.provider_request_body.as_ref(), - row.provider_request_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ResponseBody, - row.response_body.as_ref(), - row.response_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ClientResponseBody, - row.client_response_body.as_ref(), - row.client_response_body_compressed.as_deref(), - )?; - Ok((blobs, refs)) -} - -fn add_detached_body( - blobs: &mut Vec, - refs: &mut DetachedRefs, - request_id: &str, - field: UsageBodyField, - raw: Option<&Value>, - compressed: Option<&[u8]>, -) -> Result<(), DataLayerError> { - let payload_gzip = match raw { - Some(value) => Some(compress_json(value)?), - None => compressed.map(ToOwned::to_owned), - }; - let Some(payload_gzip) = payload_gzip else { - return Ok(()); - }; - let body_ref = usage_body_ref(request_id, field); - blobs.push(DetachedBlob { - body_ref: body_ref.clone(), - body_field: field.as_storage_field(), - payload_gzip, - }); - set_ref(refs, field, Some(body_ref)); - Ok(()) -} - -fn compress_json(value: &Value) -> Result, DataLayerError> { - let bytes = serde_json::to_vec(value).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}")) - })?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6)); - encoder.write_all(&bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}")) - })?; - encoder.finish().map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}")) - }) -} - -fn set_ref(refs: &mut DetachedRefs, field: UsageBodyField, value: Option) { - match field { - UsageBodyField::RequestBody => refs.request_body_ref = value, - UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = value, - UsageBodyField::ResponseBody => refs.response_body_ref = value, - UsageBodyField::ClientResponseBody => refs.client_response_body_ref = value, - } -} - -async fn upsert_http_audit_refs( +async fn purge_detached_body_capture( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, request_id: &str, - refs: &DetachedRefs, ) -> Result<(), DataLayerError> { + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; sqlx::query( r#" -INSERT INTO usage_http_audits ( - request_id, - request_body_ref, - provider_request_body_ref, - response_body_ref, - client_response_body_ref, - body_capture_mode -) -VALUES (?, ?, ?, ?, ?, 'ref_backed') -ON DUPLICATE KEY UPDATE - request_body_ref = COALESCE(VALUES(request_body_ref), request_body_ref), - provider_request_body_ref = COALESCE(VALUES(provider_request_body_ref), provider_request_body_ref), - response_body_ref = COALESCE(VALUES(response_body_ref), response_body_ref), - client_response_body_ref = COALESCE(VALUES(client_response_body_ref), client_response_body_ref), - body_capture_mode = 'ref_backed', - updated_at = UNIX_TIMESTAMP() +UPDATE usage_http_audits +SET request_body_ref = NULL, + provider_request_body_ref = NULL, + response_body_ref = NULL, + client_response_body_ref = NULL, + body_capture_mode = 'none', + updated_at = UNIX_TIMESTAMP() +WHERE request_id = ? "#, ) .bind(request_id) - .bind(refs.request_body_ref.as_deref()) - .bind(refs.provider_request_body_ref.as_deref()) - .bind(refs.response_body_ref.as_deref()) - .bind(refs.client_response_body_ref.as_deref()) .execute(&mut **tx) .await .map_sql_err()?; - Ok(()) + delete_empty_http_audit(tx, request_id).await } async fn cleanup_expired_api_keys( @@ -1127,29 +832,21 @@ ORDER BY expires_at ASC, id ASC #[cfg(test)] mod tests { - use std::io::Read; - - use flate2::read::GzDecoder; use serde_json::json; - use super::{compress_json, legacy_body_ref_plan}; + use super::{legacy_body_ref_purge_plan, DETAIL_BODY_PREDICATE}; #[test] - fn mysql_cleanup_legacy_body_ref_plan_preserves_unrelated_metadata() { + fn mysql_cleanup_legacy_body_ref_purge_preserves_unrelated_metadata() { let metadata = json!({ "trace": "kept", "request_body_ref": "usage://request/request-1/request_body", "response_body_ref": "usage://request/other/response_body" }) .to_string(); - let (refs, metadata) = legacy_body_ref_plan("request-1", Some(&metadata)) + let metadata = legacy_body_ref_purge_plan(Some(&metadata)) .expect("legacy plan should build") .expect("legacy refs should be present"); - assert_eq!( - refs.request_body_ref.as_deref(), - Some("usage://request/request-1/request_body") - ); - assert!(refs.response_body_ref.is_none()); assert_eq!( serde_json::from_str::( metadata.as_deref().expect("trace metadata should remain") @@ -1160,18 +857,9 @@ mod tests { } #[test] - fn mysql_cleanup_gzip_payload_round_trips() { - let value = json!({"hello": "world"}); - let payload = compress_json(&value).expect("body should compress"); - let mut decoder = GzDecoder::new(payload.as_slice()); - let mut decoded = Vec::new(); - decoder - .read_to_end(&mut decoded) - .expect("body should decompress"); - assert_eq!( - serde_json::from_slice::(&decoded) - .expect("decoded body should be JSON"), - value - ); + fn mysql_detail_cleanup_includes_detached_capture() { + assert!(DETAIL_BODY_PREDICATE.contains("usage_body_blobs")); + assert!(DETAIL_BODY_PREDICATE.contains("usage_http_audits")); + assert!(!DETAIL_BODY_PREDICATE.contains("payload_gzip")); } } diff --git a/crates/aether-data/adapters/mysql/src/usage/counters.rs b/crates/aether-data/adapters/mysql/src/usage/counters.rs index 496fadc77..be41956c1 100644 --- a/crates/aether-data/adapters/mysql/src/usage/counters.rs +++ b/crates/aether-data/adapters/mysql/src/usage/counters.rs @@ -24,6 +24,7 @@ SELECT id, kind, target_id, + target_tunnel_generation, request_count_delta, total_requests_delta, success_count_delta, @@ -49,6 +50,7 @@ struct DeltaRow { id: String, kind: String, target_id: String, + target_tunnel_generation: Option, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -71,7 +73,10 @@ struct Aggregates { provider_api_keys: BTreeMap, models: BTreeMap, provider_monthly: BTreeMap, - proxy_nodes: BTreeMap, + // Keep the node incarnation in the aggregation key. A node id can be + // reused after deletion, so a bare id would route old deltas to the new + // node. + proxy_nodes: BTreeMap<(String, String), ProxyNodeCounterDelta>, management_tokens: BTreeMap, api_key_last_used: BTreeMap, } @@ -142,16 +147,28 @@ impl Aggregates { .or_default() += row.total_cost_usd_delta; } KIND_PROXY_NODE => { - let entry = aggregates - .proxy_nodes - .entry(row.target_id.clone()) - .or_insert(ProxyNodeCounterDelta { + let Some(tunnel_generation) = row + .target_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + // Legacy rows have no identity fence. Mark them + // processed without applying them to any node. + continue; + }; + let aggregate_key = (row.target_id.clone(), tunnel_generation.clone()); + let entry = aggregates.proxy_nodes.entry(aggregate_key).or_insert( + ProxyNodeCounterDelta { node_id: row.target_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), total_requests_delta: 0, failed_requests_delta: 0, dns_failures_delta: 0, stream_errors_delta: 0, - }); + }, + ); entry.total_requests_delta += row.total_requests_delta; entry.failed_requests_delta += row.error_count_delta; entry.dns_failures_delta += row.dns_failures_delta; @@ -241,8 +258,8 @@ pub(super) async fn flush( for (target_id, delta) in &aggregates.provider_monthly { apply_provider_monthly(&mut tx, target_id, *delta).await?; } - for (target_id, delta) in &aggregates.proxy_nodes { - apply_proxy_node(&mut tx, target_id, delta).await?; + for ((target_id, tunnel_generation), delta) in &aggregates.proxy_nodes { + apply_proxy_node(&mut tx, target_id, tunnel_generation, delta).await?; } for (target_id, delta) in &aggregates.management_tokens { apply_management_token(&mut tx, target_id, delta).await?; @@ -283,9 +300,42 @@ pub(super) async fn enqueue_proxy_node( if delta.is_noop() { return Ok(false); } + let Some(expected_tunnel_generation) = delta + .expected_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .filter(|value| value.len() <= 64) + .map(ToOwned::to_owned) + else { + // A bare id is not an identity fence. Reject it instead of rebinding + // the delta to whichever incarnation currently owns that id. + return Ok(false); + }; let node_id = delta.node_id.trim().to_string(); let request_id = format!("proxy_node:{node_id}:{}", uuid::Uuid::new_v4()); let mut tx = pool.begin().await.map_sql_err()?; + // Keep the parent lookup lock-free because flush claims outbox rows before + // updating proxy_nodes. The generation is stored in the outbox row and is + // checked again by flush, so a concurrent id reuse can only discard this + // delta, never apply it to the replacement row. + let tunnel_generation: Option = sqlx::query_scalar( + "SELECT tunnel_generation FROM proxy_nodes WHERE id = ? AND BINARY tunnel_generation = BINARY ? LIMIT 1", + ) + .bind(&node_id) + .bind(&expected_tunnel_generation) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(_tunnel_generation) = tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; insert_delta( &mut tx, DeltaInsert { @@ -296,6 +346,7 @@ pub(super) async fn enqueue_proxy_node( error_count_delta: delta.failed_requests_delta, dns_failures_delta: delta.dns_failures_delta, stream_errors_delta: delta.stream_errors_delta, + target_tunnel_generation: Some(&expected_tunnel_generation), ..DeltaInsert::default() }, ) @@ -731,6 +782,7 @@ struct DeltaInsert<'a> { request_id: &'a str, kind: &'a str, target_id: &'a str, + target_tunnel_generation: Option<&'a str>, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -759,18 +811,20 @@ async fn insert_delta( sqlx::query( r#" INSERT INTO usage_counter_deltas ( - id, request_id, kind, target_id, request_count_delta, total_requests_delta, + id, request_id, kind, target_id, target_tunnel_generation, + request_count_delta, total_requests_delta, success_count_delta, error_count_delta, dns_failures_delta, stream_errors_delta, total_tokens_delta, total_cost_usd_delta, total_response_time_ms_delta, last_used_at_unix_secs, last_used_ip, candidate_last_used_at_unix_secs, removed_last_used_at_unix_secs, usage_created_at_unix_secs, created_at -) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(uuid::Uuid::new_v4().to_string()) .bind(request_id) .bind(input.kind) .bind(target_id) + .bind(input.target_tunnel_generation) .bind(input.request_count_delta) .bind(input.total_requests_delta) .bind(input.success_count_delta) @@ -814,6 +868,7 @@ fn map_row(row: &sqlx::mysql::MySqlRow) -> Result { id: row.try_get("id").map_sql_err()?, kind: row.try_get("kind").map_sql_err()?, target_id: row.try_get("target_id").map_sql_err()?, + target_tunnel_generation: row.try_get("target_tunnel_generation").map_sql_err()?, request_count_delta: row.try_get("request_count_delta").map_sql_err()?, total_requests_delta: row.try_get("total_requests_delta").map_sql_err()?, success_count_delta: row.try_get("success_count_delta").map_sql_err()?, @@ -997,9 +1052,10 @@ async fn apply_provider_monthly( async fn apply_proxy_node( tx: &mut sqlx::Transaction<'_, MySql>, target_id: &str, + tunnel_generation: &str, delta: &ProxyNodeCounterDelta, ) -> Result<(), DataLayerError> { - if target_id.trim().is_empty() || delta.is_noop() { + if target_id.trim().is_empty() || tunnel_generation.trim().is_empty() || delta.is_noop() { return Ok(()); } sqlx::query( @@ -1010,7 +1066,7 @@ SET total_requests = total_requests + GREATEST(?, 0), dns_failures = dns_failures + GREATEST(?, 0), stream_errors = stream_errors + GREATEST(?, 0), updated_at = ? -WHERE id = ? + WHERE id = ? AND BINARY tunnel_generation = BINARY ? "#, ) .bind(delta.total_requests_delta) @@ -1019,6 +1075,7 @@ WHERE id = ? .bind(delta.stream_errors_delta) .bind(current_unix_secs()) .bind(target_id) + .bind(tunnel_generation) .execute(&mut **tx) .await .map_sql_err()?; diff --git a/crates/aether-data/adapters/mysql/src/usage/http_capture.rs b/crates/aether-data/adapters/mysql/src/usage/http_capture.rs index 21a91bda5..d5b120530 100644 --- a/crates/aether-data/adapters/mysql/src/usage/http_capture.rs +++ b/crates/aether-data/adapters/mysql/src/usage/http_capture.rs @@ -1,10 +1,9 @@ -use std::io::{Read, Write}; - use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, - UsageBodyCaptureState, UsageBodyField, + canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json, + usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState, + UsageBodyField, }; -use flate2::{read::GzDecoder, write::GzEncoder, Compression}; +use flate2::read::GzDecoder; use serde_json::{Map, Value}; use sqlx::{mysql::MySqlRow, Row}; @@ -28,9 +27,7 @@ pub(crate) struct PreparedUsageHttpCapture { #[derive(Debug)] struct PreparedBody { - field: UsageBodyField, payload_gzip: Option>, - clear_existing: bool, } #[derive(Debug, Default)] @@ -58,15 +55,6 @@ struct HttpAuditStates { client_response_body_state: Option, } -impl HttpAuditStates { - fn any_present(&self) -> bool { - self.request_body_state.is_some() - || self.provider_request_body_state.is_some() - || self.response_body_state.is_some() - || self.client_response_body_state.is_some() - } -} - pub(crate) fn capture_update_allowed( previous: Option<&StoredRequestUsageAudit>, incoming_status: &str, @@ -141,26 +129,10 @@ pub(crate) fn prepare_usage_http_capture( .then_some(usage.client_response_body.as_ref()) .flatten(); - let request_body = prepare_body( - UsageBodyField::RequestBody, - request_body_value, - clear_request, - )?; - let provider_request_body = prepare_body( - UsageBodyField::ProviderRequestBody, - provider_request_body_value, - clear_provider_request, - )?; - let response_body = prepare_body( - UsageBodyField::ResponseBody, - response_body_value, - clear_response, - )?; - let client_response_body = prepare_body( - UsageBodyField::ClientResponseBody, - client_response_body_value, - clear_client_response, - )?; + let request_body = prepare_body(request_body_value)?; + let provider_request_body = prepare_body(provider_request_body_value)?; + let response_body = prepare_body(response_body_value)?; + let client_response_body = prepare_body(client_response_body_value)?; let refs = HttpAuditRefs { request_body_ref: resolved_write_ref( @@ -276,29 +248,13 @@ pub(crate) fn prepare_usage_http_capture( }) } -fn prepare_body( - field: UsageBodyField, - value: Option<&Value>, - clear_existing: bool, -) -> Result { - Ok(PreparedBody { - field, - payload_gzip: value.map(compress_json).transpose()?, - clear_existing, - }) -} - -fn compress_json(value: &Value) -> Result, DataLayerError> { - let bytes = serde_json::to_vec(value).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}")) - })?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6)); - encoder.write_all(&bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}")) - })?; - encoder.finish().map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}")) - }) +fn prepare_body(value: Option<&Value>) -> Result { + if value.is_some() { + return Err(DataLayerError::InvalidInput( + "usage body persistence is disabled".to_string(), + )); + } + Ok(PreparedBody { payload_gzip: None }) } fn resolved_write_ref( @@ -308,9 +264,7 @@ fn resolved_write_ref( has_blob: bool, ) -> Option { explicit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) .or_else(|| has_blob.then(|| usage_body_ref(request_id, field))) } @@ -378,23 +332,63 @@ pub(crate) async fn sync_usage_http_capture( request_id: &str, prepared: &PreparedUsageHttpCapture, ) -> Result<(), DataLayerError> { - for body in [ + let bodies = [ &prepared.request_body, &prepared.provider_request_body, &prepared.response_body, &prepared.client_response_body, - ] { - sync_body(tx, request_id, body).await?; + ]; + let contains_capture = prepared.request_headers.is_some() + || prepared.provider_request_headers.is_some() + || prepared.response_headers.is_some() + || prepared.client_response_headers.is_some() + || prepared.refs.any_present() + || bodies.iter().any(|body| body.payload_gzip.is_some()) + || prepared.capture_mode != "none"; + if contains_capture { + return Err(DataLayerError::InvalidInput( + "usage HTTP capture persistence is disabled".to_string(), + )); } + + sqlx::query("DELETE FROM usage_http_audits WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +UPDATE `usage` +SET request_headers = NULL, + request_body = NULL, + provider_request_headers = NULL, + provider_request_body = NULL, + response_headers = NULL, + response_body = NULL, + client_response_headers = NULL, + client_response_body = NULL, + request_body_compressed = NULL, + provider_request_body_compressed = NULL, + response_body_compressed = NULL, + client_response_body_compressed = NULL +WHERE request_id = ? +"#, + ) + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + let headers_present = prepared.request_headers.is_some() || prepared.provider_request_headers.is_some() || prepared.response_headers.is_some() || prepared.client_response_headers.is_some(); - if !headers_present - && !prepared.refs.any_present() - && !prepared.states.any_present() - && prepared.capture_mode == "none" - { + if !headers_present && !prepared.refs.any_present() { return Ok(()); } @@ -512,67 +506,6 @@ ON DUPLICATE KEY UPDATE Ok(()) } -async fn sync_body( - tx: &mut sqlx::Transaction<'_, sqlx::MySql>, - request_id: &str, - body: &PreparedBody, -) -> Result<(), DataLayerError> { - let body_ref = usage_body_ref(request_id, body.field); - if body.clear_existing || body.payload_gzip.is_some() { - sqlx::query(clear_legacy_body_sql(body.field)) - .bind(request_id) - .execute(&mut **tx) - .await - .map_sql_err()?; - } - if body.clear_existing { - sqlx::query("DELETE FROM usage_body_blobs WHERE body_ref = ?") - .bind(body_ref) - .execute(&mut **tx) - .await - .map_sql_err()?; - return Ok(()); - } - if let Some(payload_gzip) = body.payload_gzip.as_deref() { - sqlx::query( - r#" -INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip) -VALUES (?, ?, ?, ?) -ON DUPLICATE KEY UPDATE - request_id = VALUES(request_id), - body_field = VALUES(body_field), - payload_gzip = VALUES(payload_gzip), - updated_at = UNIX_TIMESTAMP() -"#, - ) - .bind(body_ref) - .bind(request_id) - .bind(body.field.as_storage_field()) - .bind(payload_gzip) - .execute(&mut **tx) - .await - .map_sql_err()?; - } - Ok(()) -} - -fn clear_legacy_body_sql(field: UsageBodyField) -> &'static str { - match field { - UsageBodyField::RequestBody => { - "UPDATE `usage` SET request_body = NULL, request_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ProviderRequestBody => { - "UPDATE `usage` SET provider_request_body = NULL, provider_request_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ResponseBody => { - "UPDATE `usage` SET response_body = NULL, response_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ClientResponseBody => { - "UPDATE `usage` SET client_response_body = NULL, client_response_body_compressed = NULL WHERE request_id = ?" - } - } -} - pub(crate) fn hydrate_usage_row( row: &MySqlRow, usage: &mut StoredRequestUsageAudit, @@ -690,8 +623,7 @@ fn resolved_read_ref( has_compressed: bool, ) -> Option { audit_ref - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) + .and_then(|body_ref| canonical_usage_body_ref_for(&body_ref, request_id, field)) .or_else(|| has_compressed.then(|| usage_body_ref(request_id, field))) .or_else(|| metadata_body_ref(metadata, request_id, field)) } @@ -704,13 +636,7 @@ fn metadata_body_ref( metadata .and_then(|metadata| metadata.get(field.as_ref_key())) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field)) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) } fn optional_state( @@ -752,7 +678,11 @@ pub(crate) async fn hydrate_usage_body_refs( let Some(body_ref) = usage.body_ref(field) else { continue; }; - let value = resolve_body_ref(pool, body_ref).await?; + let Some(body_ref) = canonical_usage_body_ref_for(body_ref, &usage.request_id, field) + else { + continue; + }; + let value = resolve_body_ref(pool, &body_ref).await?; match field { UsageBodyField::RequestBody => usage.request_body = value, UsageBodyField::ProviderRequestBody => usage.provider_request_body = value, @@ -767,19 +697,22 @@ pub(crate) async fn resolve_body_ref( pool: &MysqlPool, body_ref: &str, ) -> Result, DataLayerError> { + let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { + return Ok(None); + }; + let canonical_ref = usage_body_ref(&request_id, field); if let Some(payload_gzip) = sqlx::query_scalar::<_, Vec>( - "SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? LIMIT 1", + "SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? AND request_id = ? AND body_field = ? LIMIT 1", ) - .bind(body_ref) + .bind(&canonical_ref) + .bind(&request_id) + .bind(field.as_storage_field()) .fetch_optional(pool) .await .map_sql_err()? { return inflate_json(&payload_gzip).map(Some); } - let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { - return Ok(None); - }; let (inline_column, compressed_column) = usage_body_sql_columns(field); let row = sqlx::query(&format!( "SELECT CAST({inline_column} AS CHAR) AS inline_body, {compressed_column} AS compressed_body FROM `usage` WHERE request_id = ? LIMIT 1" @@ -819,11 +752,7 @@ fn usage_body_sql_columns(field: UsageBodyField) -> (&'static str, &'static str) } fn inflate_json(bytes: &[u8]) -> Result { - let mut decoder = GzDecoder::new(bytes); - let mut decoded = Vec::new(); - decoder.read_to_end(&mut decoded).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to decompress usage body: {err}")) - })?; + let decoded = read_decompressed_usage_json(GzDecoder::new(bytes))?; serde_json::from_slice(&decoded).map_err(|err| { DataLayerError::UnexpectedValue(format!("failed to decode usage body JSON: {err}")) }) diff --git a/crates/aether-data/adapters/mysql/src/usage/tests.rs b/crates/aether-data/adapters/mysql/src/usage/tests.rs index 4d849797e..c65f027ae 100644 --- a/crates/aether-data/adapters/mysql/src/usage/tests.rs +++ b/crates/aether-data/adapters/mysql/src/usage/tests.rs @@ -152,6 +152,124 @@ fn mysql_usage_upsert_guards_candidate_identity_metadata_and_routing_from_late_l .contains("OR (status = 'streaming' AND VALUES(status) = 'pending')")); } +#[tokio::test] +async fn mysql_stale_terminal_event_is_a_full_transaction_noop_when_url_is_set() { + let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") + .ok() + .filter(|value| !value.trim().is_empty()) + else { + eprintln!("skipping MySQL stale terminal test because AETHER_TEST_MYSQL_URL is unset"); + return; + }; + + let pool = sqlx::mysql::MySqlPoolOptions::new() + .max_connections(1) + .connect(&database_url) + .await + .expect("mysql test pool should connect"); + run_migrations(&pool) + .await + .expect("mysql migrations should run"); + let suffix = unique_suffix(); + let user_id = format!("stale-user-{suffix}"); + let api_key_id = format!("stale-api-key-{suffix}"); + let provider_id = format!("stale-provider-{suffix}"); + let provider_key_id = format!("stale-provider-key-{suffix}"); + let request_id = format!("stale-request-{suffix}"); + seed_stats_targets(&pool, &user_id, &api_key_id, &provider_id, &provider_key_id).await; + let repository = MysqlUsageWriteRepository::new(pool.clone()); + + let mut newer = sample_usage( + &request_id, + &user_id, + &api_key_id, + &provider_id, + &provider_key_id, + "completed", + "pending", + 2_000, + ); + newer.candidate_id = Some("candidate-new".to_string()); + newer.route_kind = Some("route-new".to_string()); + repository + .upsert(newer) + .await + .expect("newer terminal usage should upsert"); + + let counter_rows_before: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?") + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("counter rows should count"); + let routing_before: (Option, Option) = sqlx::query_as( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?", + ) + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("routing snapshot should load"); + let settlement_before: (String, Option) = sqlx::query_as( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?", + ) + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("settlement snapshot should load"); + + let mut stale = sample_usage( + &request_id, + &user_id, + &api_key_id, + &provider_id, + &provider_key_id, + "failed", + "void", + 1_999, + ); + stale.status_code = Some(503); + stale.total_cost_usd = Some(99.0); + stale.actual_total_cost_usd = Some(98.0); + stale.candidate_id = Some("candidate-stale".to_string()); + stale.route_kind = Some("route-stale".to_string()); + let stored = repository + .upsert(stale) + .await + .expect("stale terminal usage should be ignored"); + + assert_eq!(stored.status, "completed"); + assert_eq!(stored.billing_status, "pending"); + assert_eq!(stored.status_code, Some(200)); + assert_eq!(stored.total_cost_usd, 0.5); + assert_eq!(stored.routing_candidate_id(), Some("candidate-new")); + assert_eq!(stored.routing_route_kind(), Some("route-new")); + assert_eq!(stored.updated_at_unix_secs, 2_000); + + let counter_rows_after: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?") + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("counter rows should count"); + let routing_after: (Option, Option) = sqlx::query_as( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?", + ) + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("routing snapshot should load"); + let settlement_after: (String, Option) = sqlx::query_as( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?", + ) + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("settlement snapshot should load"); + assert_eq!(counter_rows_after, counter_rows_before); + assert_eq!(routing_after, routing_before); + assert_eq!(settlement_after, settlement_before); +} + #[tokio::test] async fn mysql_usage_write_repository_upserts_and_flushes_counters_when_url_is_set() { let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") @@ -470,7 +588,7 @@ async fn mysql_concurrent_same_request_upserts_enqueue_counters_once_when_url_is } #[tokio::test] -async fn mysql_usage_http_capture_round_trips_when_url_is_set() { +async fn mysql_usage_http_capture_is_not_persisted_when_url_is_set() { let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") .ok() .filter(|value| !value.trim().is_empty()) @@ -516,29 +634,17 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() { .upsert(rich) .await .expect("MySQL canonical capture should upsert"); - assert_eq!( - stored.request_headers, - Some(serde_json::json!({"x-client": "one"})) - ); - assert_eq!( - stored.request_body, - Some(serde_json::json!({"request": true})) - ); - assert_eq!( - stored.request_body_state, - Some(UsageBodyCaptureState::Reference) - ); - assert_eq!( - stored.request_body_ref.as_deref(), - Some(format!("usage://request/{request_id}/request_body").as_str()) - ); + assert!(stored.request_headers.is_none()); + assert!(stored.request_body.is_none()); + assert!(stored.request_body_state.is_none()); + assert!(stored.request_body_ref.is_none()); let blob_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?") .bind(&request_id) .fetch_one(&pool) .await .expect("MySQL canonical blobs should count"); - assert_eq!(blob_count, 2); + assert_eq!(blob_count, 0); let legacy_body: Option = sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?") .bind(&request_id) @@ -561,8 +667,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() { .upsert(sparse) .await .expect("MySQL sparse capture should upsert"); - assert_eq!(sparse_stored.request_headers, stored.request_headers); - assert_eq!(sparse_stored.request_body, stored.request_body); + assert!(sparse_stored.request_headers.is_none()); + assert!(sparse_stored.request_body.is_none()); let mut clear = sample_usage( &request_id, @@ -582,11 +688,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() { .expect("MySQL explicit none should clear"); assert!(cleared.request_body.is_none()); assert!(cleared.request_body_ref.is_none()); - assert_eq!( - cleared.request_body_state, - Some(UsageBodyCaptureState::None) - ); - assert_eq!(cleared.provider_request_body, stored.provider_request_body); + assert!(cleared.request_body_state.is_none()); + assert!(cleared.provider_request_body.is_none()); } #[tokio::test] @@ -992,16 +1095,27 @@ async fn mysql_usage_cleanup_executes_when_url_is_set() { .await .expect("MySQL detail cleanup should succeed"); assert!(summary.body_externalized >= 1); - let body_ref: String = - sqlx::query_scalar("SELECT request_body_ref FROM usage_http_audits WHERE request_id = ?") + let stored_body: Option = + sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?") .bind(&request_id) .fetch_one(&pool) .await - .expect("externalized body ref should load"); - assert_eq!( - body_ref, - format!("usage://request/{request_id}/request_body") - ); + .expect("purged body should load"); + assert!(stored_body.is_none()); + let body_blobs: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?") + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("purged body blobs should count"); + assert_eq!(body_blobs, 0); + let body_audits: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM usage_http_audits WHERE request_id = ?") + .bind(&request_id) + .fetch_one(&pool) + .await + .expect("purged body refs should count"); + assert_eq!(body_audits, 0); let headers_only = UsageCleanupTargets { detail_body: false, diff --git a/crates/aether-data/adapters/mysql/src/users.rs b/crates/aether-data/adapters/mysql/src/users.rs index e36a3e8a5..6d039f3e4 100644 --- a/crates/aether-data/adapters/mysql/src/users.rs +++ b/crates/aether-data/adapters/mysql/src/users.rs @@ -3,11 +3,14 @@ use chrono::{DateTime, TimeZone, Utc}; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::users::{ - normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, + is_valid_bcrypt_hash, last_oauth_unbind_denial, normalize_user_group_name, + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy, - UserExportSummary, UserReadRepository, + UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED, + LAST_ACTIVE_ADMIN_UPDATE_DENIED, }; use aether_data_contracts::DataLayerError; @@ -25,6 +28,92 @@ SELECT FROM users "#; +const MYSQL_LOCK_ACTIVE_ADMINS_SQL: &str = r#" +SELECT id +FROM users +WHERE LOWER(role) = 'admin' + AND is_active = 1 + AND is_deleted = 0 +ORDER BY id +FOR UPDATE +"#; + +const MYSQL_DELETE_USER_IF_WALLET_ABSENT_SQL: &str = r#" +DELETE FROM users +WHERE id = ? + AND NOT EXISTS ( + SELECT 1 + FROM wallets AS wallet + WHERE wallet.user_id = ? + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = ? + ) + ) +"#; + +const MYSQL_DELETE_USER_API_KEYS_SQL: &str = "DELETE FROM api_keys WHERE user_id = ?"; + +const MYSQL_DELETE_USER_DEPENDENTS_SQL: &[&str] = &[ + "DELETE FROM usage_request_admissions WHERE subject_id = ?", + "DELETE FROM usage_cost_reservations WHERE subject_id = ?", + "DELETE FROM gemini_file_mappings WHERE user_id = ?", + "DELETE FROM api_key_provider_mappings WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + MYSQL_DELETE_USER_API_KEYS_SQL, + "DELETE FROM management_tokens WHERE user_id = ?", + "DELETE FROM user_sessions WHERE user_id = ?", + "DELETE FROM user_oauth_links WHERE user_id = ?", + "DELETE FROM user_group_members WHERE user_id = ?", + "DELETE FROM user_preferences WHERE user_id = ?", + "DELETE FROM user_invite_codes WHERE user_id = ?", + "DELETE FROM announcement_reads WHERE user_id = ?", +]; + +const MYSQL_PREPARE_USER_FACTS_FOR_DELETION_SQL: &[&str] = &[ + "UPDATE referral_rewards SET status = CASE WHEN status IN ('pending', 'failed', 'applying') THEN 'voided' ELSE status END, failure_reason = NULL, admin_note = NULL, updated_at = UNIX_TIMESTAMP() WHERE ? IN (inviter_user_id, invitee_user_id)", + "UPDATE referral_rewards SET failure_reason = NULL, admin_note = NULL, updated_at = UNIX_TIMESTAMP() WHERE admin_operator_id = ?", + "UPDATE user_referrals SET invite_code_snapshot = 'deleted-user', source_json = NULL, updated_at = UNIX_TIMESTAMP() WHERE ? IN (inviter_user_id, invitee_user_id)", + "UPDATE user_plan_entitlements SET status = CASE WHEN status = 'active' THEN 'revoked' ELSE status END, expires_at = LEAST(expires_at, UNIX_TIMESTAMP()), updated_at = UNIX_TIMESTAMP() WHERE user_id = ?", + "UPDATE wallets SET status = 'disabled', updated_at = UNIX_TIMESTAMP() WHERE user_id = ?", + "UPDATE wallets SET status = 'disabled', updated_at = UNIX_TIMESTAMP() WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + "UPDATE audit_logs SET description = 'deleted user event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE user_id = ?", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = ?)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE wallet_transactions SET description = NULL WHERE operator_id = ?", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order WHERE history_order.user_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.user_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?) AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE user_id = ?", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE user_id = ?", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE ? IN (requested_by, approved_by, processed_by)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = ?)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE redeem_code_batches SET description = NULL WHERE created_by = ?", +]; + +const MYSQL_ANONYMIZE_USER_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE video_tasks SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE `usage` SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE stats_user_daily SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_summary SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_model SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_api_format SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_model_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_model SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_model_provider SET username = NULL WHERE user_id = ?", +]; + +const MYSQL_ANONYMIZE_USER_API_KEY_HISTORY_SQL: &str = + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)"; + const USER_EXPORT_COLUMNS: &str = r#" SELECT id, @@ -65,6 +154,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -87,6 +177,7 @@ SELECT users.allowed_models_mode AS allowed_models_mode, users.is_active AS is_active, users.is_deleted AS is_deleted, + users.security_version AS security_version, users.created_at AS created_at, users.last_login_at AS last_login_at FROM users @@ -128,6 +219,7 @@ const USER_SESSION_COLUMNS: &str = r#" SELECT id, user_id, + security_version, client_device_id, device_label, refresh_token_hash, @@ -227,6 +319,132 @@ impl MysqlUserReadRepository { let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; rows.iter().map(map_user_group_member_row).collect() } + + async fn delete_local_auth_user_inner( + &self, + user_id: &str, + require_wallet_absent: bool, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(MYSQL_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let target_security_state = + sqlx::query("SELECT role, is_active, is_deleted FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(target_security_state) = target_security_state else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let target_role = target_security_state + .try_get::("role") + .map_sql_err()?; + let target_is_active = target_security_state + .try_get::("is_active") + .map_sql_err()?; + let target_is_deleted = target_security_state + .try_get::("is_deleted") + .map_sql_err()?; + if target_role.eq_ignore_ascii_case("admin") + && target_is_active + && !target_is_deleted + && active_admin_ids.len() <= 1 + { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_DELETE_DENIED.to_string(), + )); + } + if require_wallet_absent { + let wallet_exists: Option = sqlx::query_scalar( + r#" +SELECT 1 +FROM wallets AS wallet +WHERE wallet.user_id = ? + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = ? + ) +LIMIT 1 + "#, + ) + .bind(user_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if wallet_exists.is_some() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + } + for sql in MYSQL_PREPARE_USER_FACTS_FOR_DELETION_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + for sql in MYSQL_ANONYMIZE_USER_HISTORY_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + sqlx::query(MYSQL_ANONYMIZE_USER_API_KEY_HISTORY_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + for sql in MYSQL_DELETE_USER_DEPENDENTS_SQL { + if require_wallet_absent && *sql == MYSQL_DELETE_USER_API_KEYS_SQL { + continue; + } + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + let result = if require_wallet_absent { + sqlx::query(MYSQL_DELETE_USER_IF_WALLET_ABSENT_SQL) + .bind(user_id) + .bind(user_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + } else { + sqlx::query("DELETE FROM users WHERE id = ?") + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + }; + if require_wallet_absent && result.rows_affected() == 0 { + // A wallet may have been inserted after the initial check. Do not + // commit the history/credential mutations when the guarded delete + // loses that race. + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + if require_wallet_absent { + sqlx::query(MYSQL_DELETE_USER_API_KEYS_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + Ok(result.rows_affected() > 0) + } } #[async_trait] @@ -572,6 +790,90 @@ WHERE id = ? } } + /// Restore a group while holding its row lock. This keeps the snapshot + /// comparison and replacement atomic with respect to administrator edits. + async fn restore_user_group_if_matches( + &self, + expected: &StoredUserGroup, + restored: &StoredUserGroup, + ) -> Result { + if expected.id != restored.id || expected.id.trim().is_empty() { + return Ok(false); + } + + let mut tx = self.pool.begin().await.map_sql_err()?; + let mut builder = QueryBuilder::::new(USER_GROUP_COLUMNS); + builder + .push(" WHERE id = ") + .push_bind(&expected.id) + .push(" FOR UPDATE"); + let row = builder + .build() + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_user_group_row(&row)?; + if ¤t != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE user_groups +SET name = ?, + normalized_name = ?, + description = ?, + priority = ?, + allowed_providers = ?, + allowed_providers_mode = ?, + allowed_api_formats = ?, + allowed_api_formats_mode = ?, + allowed_models = ?, + allowed_models_mode = ?, + rate_limit = ?, + rate_limit_mode = ?, + created_at = ?, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(&restored.name) + .bind(&restored.normalized_name) + .bind(&restored.description) + .bind(restored.priority) + .bind(json_string_from_option_vec( + restored.allowed_providers.as_ref(), + )) + .bind(&restored.allowed_providers_mode) + .bind(json_string_from_option_vec( + restored.allowed_api_formats.as_ref(), + )) + .bind(&restored.allowed_api_formats_mode) + .bind(json_string_from_option_vec( + restored.allowed_models.as_ref(), + )) + .bind(&restored.allowed_models_mode) + .bind(restored.rate_limit) + .bind(&restored.rate_limit_mode) + .bind(restored.created_at.map(|value| value.timestamp())) + .bind(restored.updated_at.map(|value| value.timestamp())) + .bind(&restored.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn delete_user_group(&self, group_id: &str) -> Result { let result = sqlx::query("DELETE FROM user_groups WHERE id = ?") .bind(group_id) @@ -599,6 +901,30 @@ WHERE id = ? user_ids: &[String], ) -> Result, DataLayerError> { let mut tx = self.pool.begin().await.map_sql_err()?; + // Membership rollback serializes on the owning user row. Lock every user that can be + // removed or inserted in deterministic order so it cannot race the per-user CAS restore. + let mut locked_user_ids = normalized_ids(user_ids); + let existing_user_ids = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_group_members WHERE group_id = ? ORDER BY user_id", + ) + .bind(group_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + locked_user_ids.extend(existing_user_ids); + locked_user_ids.sort(); + locked_user_ids.dedup(); + if !locked_user_ids.is_empty() { + let mut builder = QueryBuilder::::new("SELECT id FROM users WHERE id IN ("); + { + let mut separated = builder.separated(", "); + for user_id in &locked_user_ids { + separated.push_bind(user_id); + } + } + builder.push(") ORDER BY id FOR UPDATE"); + builder.build().fetch_all(&mut *tx).await.map_sql_err()?; + } sqlx::query("DELETE FROM user_group_members WHERE group_id = ?") .bind(group_id) .execute(&mut *tx) @@ -671,6 +997,16 @@ WHERE user_group_members.user_id IN ( group_ids: &[String], ) -> Result, DataLayerError> { let mut tx = self.pool.begin().await.map_sql_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(Vec::new()); + } sqlx::query("DELETE FROM user_group_members WHERE user_id = ?") .bind(user_id) .execute(&mut *tx) @@ -692,20 +1028,106 @@ WHERE user_group_members.user_id IN ( self.list_user_groups_for_user(user_id).await } + async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + let expected = normalized_ids(expected_group_ids); + let restored = normalized_ids(restored_group_ids); + let mut tx = self.pool.begin().await.map_sql_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let current = sqlx::query_scalar::<_, String>( + "SELECT group_id FROM user_group_members WHERE user_id = ? ORDER BY group_id ASC FOR UPDATE", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if current != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + if !restored.is_empty() { + let mut builder = QueryBuilder::::new( + "SELECT COUNT(*) AS count FROM user_groups WHERE id IN (", + ); + { + let mut separated = builder.separated(", "); + for group_id in &restored { + separated.push_bind(group_id); + } + } + builder.push(")"); + let count = builder + .build() + .fetch_one(&mut *tx) + .await + .map_sql_err()? + .try_get::("count") + .map_sql_err()?; + if count != restored.len() as i64 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + } + sqlx::query("DELETE FROM user_group_members WHERE user_id = ?") + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + let now = current_unix_secs(); + for group_id in restored { + sqlx::query( + "INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)", + ) + .bind(group_id) + .bind(user_id) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn add_user_to_group( &self, group_id: &str, user_id: &str, ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } let result = sqlx::query( "INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)", ) .bind(group_id) .bind(user_id) .bind(current_unix_secs()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(result.rows_affected() > 0) } @@ -820,6 +1242,72 @@ WHERE user_group_members.user_id IN ( Ok(self.fetch_auth_rows(builder).await?.into_iter().next()) } + async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: DateTime, + _provider_enabled_snapshot: bool, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let provider_enabled: Option = sqlx::query_scalar( + "SELECT is_enabled FROM oauth_providers WHERE provider_type = ? FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_enabled != Some(true) { + tx.rollback().await.map_sql_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::ProviderUnavailable); + } + let row = sqlx::query(&format!( + "{USER_AUTH_COLUMNS_QUALIFIED} JOIN user_oauth_links ON users.id = user_oauth_links.user_id WHERE user_oauth_links.provider_type = ? AND user_oauth_links.provider_user_id = ? LIMIT 1 FOR UPDATE" + )) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::NotLinked); + }; + let mut user = map_user_auth_row(&row)?; + sqlx::query( + "UPDATE user_oauth_links SET provider_username = COALESCE(?, provider_username), provider_email = COALESCE(?, provider_email), extra_data = COALESCE(?, extra_data), last_login_at = ? WHERE provider_type = ? AND provider_user_id = ?", + ) + .bind(provider_username) + .bind(provider_email) + .bind(optional_json_string(extra_data, "user_oauth_links.extra_data")?) + .bind(touched_at.timestamp()) + .bind(provider_type) + .bind(provider_user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if let Some(verified_email) = verified_email { + let result = sqlx::query( + "UPDATE users SET email_verified = 1, updated_at = ? WHERE id = ? AND email_verified = 0 AND LOWER(TRIM(email)) = LOWER(TRIM(?))", + ) + .bind(touched_at.timestamp()) + .bind(&user.id) + .bind(verified_email) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() == 1 { + user.email_verified = true; + } + } + tx.commit().await.map_sql_err()?; + Ok(ResolveOAuthLinkedUserOutcome::Linked(user)) + } + async fn touch_oauth_link( &self, provider_type: &str, @@ -858,6 +1346,7 @@ WHERE provider_type = ? async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: DateTime, ) -> Result, DataLayerError> { @@ -869,11 +1358,12 @@ INSERT INTO users ( allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode, is_active, is_deleted, created_at, updated_at, last_login_at ) -VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?) +VALUES (?, ?, ?, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?) "#, ) .bind(&user_id) .bind(email) + .bind(email_verified) .bind(username) .bind(created_at.timestamp()) .bind(created_at.timestamp()) @@ -925,7 +1415,20 @@ VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inh Ok(total.max(0) as u64) } - async fn upsert_user_oauth_link( + async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + let exists: Option = + sqlx::query_scalar("SELECT 1 FROM user_oauth_links WHERE provider_type = ? LIMIT 1") + .bind(provider_type) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + Ok(exists.is_some()) + } + + async fn bind_user_oauth_link_if_provider_enabled( &self, user_id: &str, provider_type: &str, @@ -934,69 +1437,249 @@ VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inh provider_email: Option<&str>, extra_data: Option, linked_at: DateTime, - ) -> Result<(), DataLayerError> { + _provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result { let extra_data = optional_json_string(extra_data, "user_oauth_links.extra_data")?; - let updated = sqlx::query( - r#" -UPDATE user_oauth_links -SET provider_user_id = ?, - provider_username = ?, - provider_email = ?, - extra_data = ?, - last_login_at = ? -WHERE user_id = ? - AND provider_type = ? -"#, + let mut tx = self.pool.begin().await.map_sql_err()?; + let provider_enabled: Option = sqlx::query_scalar( + "SELECT is_enabled FROM oauth_providers WHERE provider_type = ? FOR UPDATE", ) - .bind(provider_user_id) - .bind(provider_username) - .bind(provider_email) - .bind(extra_data.as_deref()) - .bind(linked_at.timestamp()) - .bind(user_id) .bind(provider_type) - .execute(&self.pool) + .fetch_optional(&mut *tx) .await .map_sql_err()?; - if updated.rows_affected() == 0 { - sqlx::query( + if provider_enabled.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderNotFound); + } + if provider_enabled != Some(true) { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderDisabled); + } + if let Some(expectation) = session_expectation { + let session_is_current: Option = sqlx::query_scalar( r#" +SELECT 1 +FROM users +JOIN user_sessions + ON user_sessions.user_id = users.id +WHERE users.id = ? + AND users.is_active = 1 + AND users.is_deleted = 0 + AND users.security_version = ? + AND user_sessions.id = ? + AND user_sessions.security_version = ? + AND user_sessions.client_device_id = ? + AND user_sessions.revoked_at IS NULL + AND user_sessions.expires_at > GREATEST(?, UNIX_TIMESTAMP()) +FOR UPDATE +"#, + ) + .bind(user_id) + .bind(expectation.security_version) + .bind(&expectation.session_id) + .bind(expectation.security_version) + .bind(&expectation.client_device_id) + .bind(expectation.checked_at.timestamp()) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if session_is_current.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::SessionUnavailable); + } + } else { + let user_exists: Option = + sqlx::query_scalar("SELECT 1 FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::UserNotFound); + } + } + if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = ? AND provider_user_id = ? LIMIT 1 FOR UPDATE", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + { + tx.rollback().await.map_sql_err()?; + return Ok(if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }); + } + if sqlx::query_scalar::<_, i32>( + "SELECT 1 FROM user_oauth_links WHERE user_id = ? AND provider_type = ? LIMIT 1 FOR UPDATE", + ) + .bind(user_id) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + .is_some() + { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider); + } + let inserted = sqlx::query( + r#" INSERT INTO user_oauth_links ( id, user_id, provider_type, provider_user_id, provider_username, provider_email, extra_data, linked_at, last_login_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) "#, - ) - .bind(uuid::Uuid::new_v4().to_string()) - .bind(user_id) - .bind(provider_type) - .bind(provider_user_id) - .bind(provider_username) - .bind(provider_email) - .bind(extra_data.as_deref()) - .bind(linked_at.timestamp()) - .bind(linked_at.timestamp()) - .execute(&self.pool) - .await - .map_sql_err()?; + ) + .bind(uuid::Uuid::new_v4().to_string()) + .bind(user_id) + .bind(provider_type) + .bind(provider_user_id) + .bind(provider_username) + .bind(provider_email) + .bind(extra_data.as_deref()) + .bind(linked_at.timestamp()) + .bind(linked_at.timestamp()) + .execute(&mut *tx) + .await; + match inserted { + Ok(_) => { + tx.commit().await.map_sql_err()?; + Ok(BindUserOAuthLinkOutcome::Bound) + } + Err(sqlx::Error::Database(err)) if err.is_unique_violation() => { + let owner = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = ? AND provider_user_id = ? LIMIT 1", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + tx.rollback().await.map_sql_err()?; + Ok(match owner { + Some(owner) if owner == user_id => { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } + Some(_) => BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser, + None => BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider, + }) + } + Err(err) => Err(DataLayerError::sql(err)), } - Ok(()) + } + + async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: DateTime, + ) -> Result { + let result = sqlx::query( + "UPDATE users SET email_verified = 1, updated_at = ? WHERE id = ? AND email_verified = 0 AND LOWER(TRIM(email)) = LOWER(TRIM(?))", + ) + .bind(verified_at.timestamp()) + .bind(user_id) + .bind(verified_email) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) } async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { + local_password_login_allowed: bool, + _enabled_provider_types_snapshot: &[String], + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let provider_exists: Option = sqlx::query_scalar( + "SELECT provider_type FROM oauth_providers WHERE provider_type = ? FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + let user = + sqlx::query("SELECT auth_source, password_hash FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(user) = user else { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + }; + let auth_source = user.try_get::("auth_source").map_sql_err()?; + let password_hash = user + .try_get::, _>("password_hash") + .map_sql_err()?; + let provider_types = sqlx::query_scalar::<_, String>( + "SELECT provider_type FROM user_oauth_links WHERE user_id = ? FOR UPDATE", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if !provider_types.iter().any(|value| value == provider_type) { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + let enabled_provider_types = sqlx::query_scalar::<_, String>( + r#" +SELECT user_oauth_links.provider_type +FROM user_oauth_links +JOIN oauth_providers + ON oauth_providers.provider_type = user_oauth_links.provider_type +WHERE user_oauth_links.user_id = ? + AND oauth_providers.is_enabled = 1 +FOR UPDATE +"#, + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let has_remaining_enabled_oauth_link = enabled_provider_types + .iter() + .any(|value| value != provider_type); + if !has_remaining_enabled_oauth_link { + if let Some(outcome) = last_oauth_unbind_denial( + &auth_source, + password_hash.as_deref(), + local_password_login_allowed, + ) { + tx.rollback().await.map_sql_err()?; + return Ok(outcome); + } + } let result = sqlx::query("DELETE FROM user_oauth_links WHERE user_id = ? AND provider_type = ?") .bind(user_id) .bind(provider_type) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - Ok(result.rows_affected() > 0) + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + tx.commit().await.map_sql_err()?; + Ok(DeleteUserOAuthLinkOutcome::Deleted) } async fn get_or_create_ldap_auth_user( @@ -1036,14 +1719,18 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { let now = chrono::Utc::now().timestamp(); let result = sqlx::query( - "UPDATE users SET email = COALESCE(?, email), username = COALESCE(?, username), updated_at = ? WHERE id = ?", + "UPDATE users SET email = CASE WHEN ? THEN ? ELSE email END, email_verified = COALESCE(?, email_verified), username = COALESCE(?, username), updated_at = ? WHERE id = ?", ) + .bind(email_present) .bind(email) + .bind(email_verified) .bind(username) .bind(now) .bind(user_id) @@ -1056,13 +1743,175 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) self.find_user_auth_by_id(user_id).await } + async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + if expected_auth.id != restored_auth.id + || expected_export.id != expected_auth.id + || restored_export.id != restored_auth.id + { + return Ok(false); + } + let mut tx = self.pool.begin().await.map_sql_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(MYSQL_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let auth_row = sqlx::query(&format!( + "{USER_AUTH_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let export_row = sqlx::query(&format!( + "{USER_EXPORT_COLUMNS} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let (Some(auth_row), Some(export_row)) = (auth_row, export_row) else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current_auth = map_user_auth_row(&auth_row)?; + let current_export = map_user_export_row(&export_row)?; + if !current_auth.matches_restore_state(expected_auth) + || !current_export.matches_restore_state(expected_export) + || current_export.rate_limit != expected_export.rate_limit + || current_export.rate_limit_mode != expected_export.rate_limit_mode + || current_export.model_capability_settings.as_ref() + != expected_model_capability_settings + || current_export.feature_settings.as_ref() != expected_feature_settings + { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let removes_active_admin = current_auth.role.eq_ignore_ascii_case("admin") + && current_auth.is_active + && !current_auth.is_deleted + && (!restored_auth.role.eq_ignore_ascii_case("admin") || !restored_auth.is_active); + if removes_active_admin && active_admin_ids.len() <= 1 { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + let security_state_changed = expected_auth.role != restored_auth.role + || expected_auth.is_active != restored_auth.is_active; + let result = sqlx::query( + r#" +UPDATE users +SET email = ?, + email_verified = ?, + username = ?, + role = ?, + allowed_providers = ?, + allowed_providers_mode = ?, + allowed_api_formats = ?, + allowed_api_formats_mode = ?, + allowed_models = ?, + allowed_models_mode = ?, + rate_limit = ?, + rate_limit_mode = ?, + model_capability_settings = ?, + feature_settings = ?, + is_active = ?, + security_version = security_version + CASE WHEN ? THEN 1 ELSE 0 END, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(restored_auth.email.as_deref()) + .bind(restored_auth.email_verified) + .bind(&restored_auth.username) + .bind(&restored_auth.role) + .bind(optional_string_list_json( + restored_auth.allowed_providers.clone(), + "users.allowed_providers", + )?) + .bind(&restored_auth.allowed_providers_mode) + .bind(optional_string_list_json( + restored_auth.allowed_api_formats.clone(), + "users.allowed_api_formats", + )?) + .bind(&restored_auth.allowed_api_formats_mode) + .bind(optional_string_list_json( + restored_auth.allowed_models.clone(), + "users.allowed_models", + )?) + .bind(&restored_auth.allowed_models_mode) + .bind(restored_export.rate_limit) + .bind(&restored_export.rate_limit_mode) + .bind(optional_json_string( + restored_model_capability_settings.clone(), + "users.model_capability_settings", + )?) + .bind(optional_json_string( + restored_feature_settings.clone(), + "users.feature_settings", + )?) + .bind(restored_auth.is_active) + .bind(security_state_changed) + .bind(current_unix_secs()) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + if security_state_changed { + let now = current_unix_secs(); + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'user_security_state_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(now) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn update_local_auth_user_password_hash( &self, user_id: &str, password_hash: String, updated_at: DateTime, ) -> Result, DataLayerError> { - let result = sqlx::query("UPDATE users SET password_hash = ?, updated_at = ? WHERE id = ?") + let result = sqlx::query( + "UPDATE users SET password_hash = ?, security_version = security_version + 1, updated_at = ? WHERE id = ?", + ) .bind(password_hash) .bind(updated_at.timestamp()) .bind(user_id) @@ -1075,6 +1924,122 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) self.find_user_auth_by_id(user_id).await } + async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: DateTime, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE users +SET password_hash = ?, + security_version = security_version + 1, + updated_at = ? +WHERE id = ? + AND ((? IS NULL AND password_hash IS NULL) OR BINARY password_hash = BINARY ?) +"#, + ) + .bind(password_hash) + .bind(updated_at.timestamp()) + .bind(user_id) + .bind(expected_password_hash) + .bind(expected_password_hash) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let updated = sqlx::query( + "UPDATE users SET password_hash = ?, security_version = security_version + 1, updated_at = ? WHERE id = ? AND is_deleted = 0", + ) + .bind(password_hash) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'admin_password_reset', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(changed_at.timestamp()) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let updated = sqlx::query( + r#" +UPDATE users +SET password_hash = ?, security_version = security_version + 1, updated_at = ? +WHERE id = ? + AND is_active = 1 + AND is_deleted = 0 + AND ((? IS NULL AND password_hash IS NULL) OR BINARY password_hash = BINARY ?) + AND EXISTS ( + SELECT 1 FROM user_sessions + WHERE user_id = ? AND id = ? AND revoked_at IS NULL AND expires_at > ? + ) +"#, + ) + .bind(next_password_hash) + .bind(changed_at.timestamp()) + .bind(user_id) + .bind(expected_password_hash) + .bind(expected_password_hash) + .bind(user_id) + .bind(current_session_id) + .bind(changed_at.timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let revoked = sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'password_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(changed_at.timestamp()) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if revoked.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn update_local_auth_user_admin_fields( &self, user_id: &str, @@ -1089,6 +2054,45 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) rate_limit: Option, is_active: Option, ) -> Result, DataLayerError> { + let mut tx = self.pool.begin().await.map_sql_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(MYSQL_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let current_security_state = + sqlx::query("SELECT role, is_active, is_deleted FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(current_security_state) = current_security_state else { + tx.rollback().await.map_sql_err()?; + return Ok(None); + }; + let current_role = current_security_state + .try_get::("role") + .map_sql_err()?; + let current_active = current_security_state + .try_get::("is_active") + .map_sql_err()?; + let current_deleted = current_security_state + .try_get::("is_deleted") + .map_sql_err()?; + let next_role = role.as_deref().unwrap_or(current_role.as_str()); + let next_active = is_active.unwrap_or(current_active); + if current_role.eq_ignore_ascii_case("admin") + && current_active + && !current_deleted + && (!next_role.eq_ignore_ascii_case("admin") || !next_active) + && active_admin_ids.len() <= 1 + { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + let security_state_changed = + !current_role.eq_ignore_ascii_case(next_role) || current_active != next_active; let allowed_providers_mode = if allowed_providers .as_ref() .is_some_and(|values| !values.is_empty()) @@ -1131,6 +2135,7 @@ SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END, rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END, is_active = CASE WHEN ? THEN ? ELSE is_active END, + security_version = security_version + CASE WHEN ? THEN 1 ELSE 0 END, updated_at = ? WHERE id = ? "#, @@ -1164,14 +2169,45 @@ WHERE id = ? .bind(rate_limit_mode) .bind(is_active.is_some()) .bind(is_active) + .bind(security_state_changed) .bind(chrono::Utc::now().timestamp()) .bind(user_id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; if result.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; return Ok(None); } + if security_state_changed { + let revoked_at = chrono::Utc::now().timestamp(); + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'user_security_state_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(revoked_at) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; self.find_user_auth_by_id(user_id).await } @@ -1380,12 +2416,14 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?) } async fn delete_local_auth_user(&self, user_id: &str) -> Result { - let result = sqlx::query("DELETE FROM users WHERE id = ?") - .bind(user_id) - .execute(&self.pool) - .await - .map_sql_err()?; - Ok(result.rows_affected() > 0) + self.delete_local_auth_user_inner(user_id, false).await + } + + async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + self.delete_local_auth_user_inner(user_id, true).await } async fn count_active_admin_users(&self) -> Result { @@ -1505,6 +2543,19 @@ ON DUPLICATE KEY UPDATE .or(session.updated_at) .or(session.last_seen_at) .unwrap_or_else(Utc::now); + let mut tx = self.pool.begin().await.map_sql_err()?; + let user_is_eligible: Option = sqlx::query_scalar( + "SELECT id FROM users WHERE id = ? AND is_active = 1 AND is_deleted = 0 AND security_version = ? FOR UPDATE", + ) + .bind(&session.user_id) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_is_eligible.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } sqlx::query( r#" UPDATE user_sessions @@ -1517,19 +2568,20 @@ WHERE user_id = ? AND client_device_id = ? AND revoked_at IS NULL AND expires_at .bind(&session.user_id) .bind(&session.client_device_id) .bind(now.timestamp()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; sqlx::query( r#" INSERT INTO user_sessions ( - id, user_id, client_device_id, device_label, device_type, ip_address, user_agent, + id, user_id, security_version, client_device_id, device_label, device_type, ip_address, user_agent, refresh_token_hash, last_seen_at, expires_at, created_at, updated_at -) VALUES (?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) +) VALUES (?, ?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&session.id) .bind(&session.user_id) + .bind(session.security_version) .bind(&session.client_device_id) .bind(session.device_label.as_deref()) .bind(session.ip_address.as_deref()) @@ -1539,9 +2591,100 @@ INSERT INTO user_sessions ( .bind(session.expires_at.unwrap_or(now).timestamp()) .bind(session.created_at.unwrap_or(now).timestamp()) .bind(session.updated_at.unwrap_or(now).timestamp()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + let mut builder = QueryBuilder::::new(USER_SESSION_COLUMNS); + builder + .push(" WHERE user_id = ") + .push_bind(&session.user_id) + .push(" AND id = ") + .push_bind(&session.id) + .push(" LIMIT 1"); + let row = builder + .build() + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let created = row.as_ref().map(map_user_session_row).transpose()?; + tx.commit().await.map_sql_err()?; + Ok(created) + } + + async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + let now = session + .created_at + .or(session.updated_at) + .or(session.last_seen_at) + .unwrap_or_else(Utc::now); + let mut tx = self.pool.begin().await.map_sql_err()?; + let matched = sqlx::query_scalar::<_, String>( + r#" +SELECT password_hash FROM users +WHERE id = ? AND BINARY password_hash = BINARY ? AND LOWER(auth_source) = 'local' + AND is_active = 1 AND is_deleted = 0 AND security_version = ? +FOR UPDATE +"#, + ) + .bind(&session.user_id) + .bind(expected_password_hash) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if matched.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + sqlx::query("UPDATE users SET last_login_at = ? WHERE id = ?") + .bind(now.timestamp()) + .bind(&session.user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +UPDATE user_sessions +SET revoked_at = ?, revoke_reason = 'replaced_by_new_login', updated_at = ? +WHERE user_id = ? AND client_device_id = ? AND revoked_at IS NULL AND expires_at > ? +"#, + ) + .bind(now.timestamp()) + .bind(now.timestamp()) + .bind(&session.user_id) + .bind(&session.client_device_id) + .bind(now.timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +INSERT INTO user_sessions ( + id, user_id, security_version, client_device_id, device_label, device_type, ip_address, user_agent, + refresh_token_hash, last_seen_at, expires_at, created_at, updated_at +) VALUES (?, ?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) +"#, + ) + .bind(&session.id) + .bind(&session.user_id) + .bind(session.security_version) + .bind(&session.client_device_id) + .bind(session.device_label.as_deref()) + .bind(session.ip_address.as_deref()) + .bind(session.user_agent.as_deref()) + .bind(&session.refresh_token_hash) + .bind(session.last_seen_at.unwrap_or(now).timestamp()) + .bind(session.expires_at.unwrap_or(now).timestamp()) + .bind(session.created_at.unwrap_or(now).timestamp()) + .bind(session.updated_at.unwrap_or(now).timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; self.find_user_session(&session.user_id, &session.id).await } @@ -1597,7 +2740,7 @@ WHERE user_id = ? AND id = ? &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: DateTime, expires_at: DateTime, @@ -1610,10 +2753,11 @@ UPDATE user_sessions SET prev_refresh_token_hash = ?, rotated_at = ?, refresh_token_hash = ?, expires_at = ?, last_seen_at = ?, ip_address = COALESCE(?, ip_address), user_agent = COALESCE(?, user_agent), updated_at = ? -WHERE user_id = ? AND id = ? +WHERE user_id = ? AND id = ? AND BINARY refresh_token_hash = BINARY ? + AND revoked_at IS NULL AND expires_at > ? "#, ) - .bind(previous_refresh_token_hash) + .bind(expected_refresh_token_hash) .bind(rotated_at.timestamp()) .bind(next_refresh_token_hash) .bind(expires_at.timestamp()) @@ -1623,6 +2767,8 @@ WHERE user_id = ? AND id = ? .bind(rotated_at.timestamp()) .bind(user_id) .bind(session_id) + .bind(expected_refresh_token_hash) + .bind(rotated_at.timestamp()) .execute(&self.pool) .await .map_sql_err()?; @@ -1672,26 +2818,24 @@ WHERE user_id = ? AND id = ? async fn count_active_local_admin_users_with_valid_password( &self, ) -> Result { - let total: i64 = sqlx::query_scalar( + let hashes = sqlx::query_scalar::<_, String>( r#" -SELECT COUNT(*) +SELECT password_hash FROM users WHERE LOWER(role) = 'admin' AND LOWER(auth_source) = 'local' AND is_deleted = 0 AND is_active = 1 - AND CHAR_LENGTH(password_hash) = 60 - AND ( - password_hash LIKE '$2a$%' - OR password_hash LIKE '$2b$%' - OR password_hash LIKE '$2y$%' - ) + AND password_hash IS NOT NULL "#, ) - .fetch_one(&self.pool) + .fetch_all(&self.pool) .await .map_sql_err()?; - Ok(total.max(0) as u64) + Ok(hashes + .iter() + .filter(|hash| is_valid_bcrypt_hash(hash)) + .count() as u64) } } @@ -2000,6 +3144,7 @@ fn map_user_auth_row(row: &MySqlRow) -> Result Result>() + .join(" "); + assert!(normalized.contains("LOWER(role) = 'admin'")); + assert!(normalized.contains("is_active = 1")); + assert!(normalized.contains("is_deleted = 0")); + assert!(normalized.contains("ORDER BY id FOR UPDATE")); + assert!(MYSQL_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM management_tokens"))); + assert!(MYSQL_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM api_keys"))); + assert!(MYSQL_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM user_sessions"))); + assert_history_anonymization_contract(MYSQL_ANONYMIZE_USER_HISTORY_SQL); + assert!(MYSQL_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .starts_with("UPDATE stats_daily_api_key SET api_key_name = NULL")); + assert!(MYSQL_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .contains("SELECT id FROM api_keys WHERE user_id = ?")); + } + + fn assert_history_anonymization_contract(statements: &[&str]) { + const TABLES: &[&str] = &[ + "request_candidates", + "video_tasks", + "usage", + "stats_user_daily", + "stats_user_summary", + "stats_user_daily_model", + "stats_user_daily_provider", + "stats_user_daily_api_format", + "stats_user_daily_model_provider", + "stats_user_daily_cost_savings", + "stats_user_daily_cost_savings_provider", + "stats_user_daily_cost_savings_model", + "stats_user_daily_cost_savings_model_provider", + ]; + + assert_eq!(statements.len(), TABLES.len()); + for table in TABLES { + let statement = statements + .iter() + .find(|sql| { + sql.starts_with(&format!("UPDATE {table} ")) + || sql.starts_with(&format!("UPDATE `{table}` ")) + }) + .unwrap_or_else(|| panic!("missing history anonymization for {table}")); + assert!(statement.contains("username = NULL")); + assert!(statement.ends_with("WHERE user_id = ?")); + } + for table in ["request_candidates", "video_tasks", "usage"] { + let statement = statements + .iter() + .find(|sql| { + sql.starts_with(&format!("UPDATE {table} ")) + || sql.starts_with(&format!("UPDATE `{table}` ")) + }) + .expect("identity snapshot table should be covered"); + assert!(statement.contains("api_key_name = NULL")); + } + } #[tokio::test] async fn repository_builds_from_lazy_pool() { diff --git a/crates/aether-data/adapters/mysql/src/video_tasks.rs b/crates/aether-data/adapters/mysql/src/video_tasks.rs index 9996d6403..5c132689d 100644 --- a/crates/aether-data/adapters/mysql/src/video_tasks.rs +++ b/crates/aether-data/adapters/mysql/src/video_tasks.rs @@ -72,6 +72,22 @@ impl MysqlVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + async fn find_by_id_for_user( + &self, + id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{VIDEO_TASK_COLUMNS} WHERE BINARY id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1" + )) + .bind(id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + async fn find_by_short_id( &self, short_id: &str, @@ -84,13 +100,29 @@ impl MysqlVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + async fn find_by_short_id_for_user( + &self, + short_id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{VIDEO_TASK_COLUMNS} WHERE BINARY short_id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1" + )) + .bind(short_id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + async fn find_by_user_external( &self, user_id: &str, external_task_id: &str, ) -> Result, DataLayerError> { let row = sqlx::query(&format!( - "{VIDEO_TASK_COLUMNS} WHERE user_id = ? AND external_task_id = ? LIMIT 1" + "{VIDEO_TASK_COLUMNS} WHERE BINARY user_id = BINARY ? AND BINARY external_task_id = BINARY ? LIMIT 1" )) .bind(user_id) .bind(external_task_id) @@ -117,6 +149,26 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository { } } + async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + match key { + VideoTaskLookupKey::Id(id) => self.find_by_id_for_user(id, user_id).await, + VideoTaskLookupKey::ShortId(short_id) => { + self.find_by_short_id_for_user(short_id, user_id).await + } + VideoTaskLookupKey::UserExternal { + user_id: lookup_user_id, + external_task_id, + } if lookup_user_id == user_id => { + self.find_by_user_external(user_id, external_task_id).await + } + VideoTaskLookupKey::UserExternal { .. } => Ok(None), + } + } + async fn list_active(&self, limit: usize) -> Result, DataLayerError> { if limit == 0 { return Ok(Vec::new()); @@ -258,15 +310,21 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository { #[async_trait] impl VideoTaskWriteRepository for MysqlVideoTaskRepository { - async fn upsert(&self, task: UpsertVideoTask) -> Result { + async fn upsert(&self, mut task: UpsertVideoTask) -> Result { + task.sanitize_for_persistence(); let id = task.id.clone(); - bind_task(sqlx::query(UPSERT_SQL), task, true, false)? + let expected_identity = task.clone(); + bind_task(sqlx::query(upsert_sql()), task, true, false)? .execute(&self.pool) .await .map_sql_err()?; - self.find_by_id(&id) - .await? - .ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into())) + let stored = self.find_by_id(&id).await?.ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "video task {id} conflicts with persisted immutable identity" + )) + })?; + stored.ensure_immutable_identity_matches(&expected_identity)?; + Ok(stored) } async fn update_if_active( @@ -368,7 +426,76 @@ FOR UPDATE SKIP LOCKED } } -const UPSERT_SQL: &str = r#" +const IMMUTABLE_IDENTITY_MATCH_SQL: &str = r#"BINARY id <=> BINARY VALUES(id) + AND BINARY short_id <=> BINARY VALUES(short_id) + AND BINARY request_id <=> BINARY VALUES(request_id) + AND BINARY user_id <=> BINARY VALUES(user_id) + AND BINARY api_key_id <=> BINARY VALUES(api_key_id) + AND BINARY external_task_id <=> BINARY VALUES(external_task_id) + AND BINARY provider_id <=> BINARY VALUES(provider_id) + AND BINARY endpoint_id <=> BINARY VALUES(endpoint_id) + AND BINARY key_id <=> BINARY VALUES(key_id) + AND BINARY client_api_format <=> BINARY VALUES(client_api_format) + AND BINARY provider_api_format <=> BINARY VALUES(provider_api_format) + AND format_converted <=> VALUES(format_converted) + AND BINARY model <=> BINARY VALUES(model) + AND duration_seconds <=> VALUES(duration_seconds) + AND BINARY resolution <=> BINARY VALUES(resolution) + AND BINARY aspect_ratio <=> BINARY VALUES(aspect_ratio) + AND BINARY size <=> BINARY VALUES(size)"#; + +const UPSERT_UPDATE_COLUMNS: &[&str] = &[ + "short_id", + "request_id", + "user_id", + "api_key_id", + "username", + "api_key_name", + "external_task_id", + "provider_id", + "endpoint_id", + "key_id", + "client_api_format", + "provider_api_format", + "format_converted", + "model", + "prompt", + "original_request_body", + "duration_seconds", + "resolution", + "aspect_ratio", + "size", + "status", + "progress_percent", + "progress_message", + "retry_count", + "poll_interval_seconds", + "next_poll_at", + "poll_count", + "max_poll_count", + "video_url", + "error_code", + "error_message", + "request_metadata", + "submitted_at", + "completed_at", + "updated_at", +]; + +fn upsert_sql() -> &'static str { + static SQL: std::sync::OnceLock = std::sync::OnceLock::new(); + SQL.get_or_init(|| { + let guarded_updates = UPSERT_UPDATE_COLUMNS + .iter() + .map(|column| { + format!( + " {column} = IF(({IMMUTABLE_IDENTITY_MATCH_SQL}), VALUES({column}), {column})" + ) + }) + .collect::>() + .join(",\n"); + format!( + r#" INSERT INTO video_tasks ( id, short_id, request_id, user_id, api_key_id, username, api_key_name, external_task_id, provider_id, endpoint_id, key_id, client_api_format, @@ -380,43 +507,11 @@ INSERT INTO video_tasks ( ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON DUPLICATE KEY UPDATE - short_id = VALUES(short_id), - request_id = VALUES(request_id), - user_id = VALUES(user_id), - api_key_id = VALUES(api_key_id), - username = VALUES(username), - api_key_name = VALUES(api_key_name), - external_task_id = VALUES(external_task_id), - provider_id = VALUES(provider_id), - endpoint_id = VALUES(endpoint_id), - key_id = VALUES(key_id), - client_api_format = VALUES(client_api_format), - provider_api_format = VALUES(provider_api_format), - format_converted = VALUES(format_converted), - model = VALUES(model), - prompt = VALUES(prompt), - original_request_body = VALUES(original_request_body), - duration_seconds = VALUES(duration_seconds), - resolution = VALUES(resolution), - aspect_ratio = VALUES(aspect_ratio), - size = VALUES(size), - status = VALUES(status), - progress_percent = VALUES(progress_percent), - progress_message = VALUES(progress_message), - retry_count = VALUES(retry_count), - poll_interval_seconds = VALUES(poll_interval_seconds), - next_poll_at = VALUES(next_poll_at), - poll_count = VALUES(poll_count), - max_poll_count = VALUES(max_poll_count), - video_url = VALUES(video_url), - error_code = VALUES(error_code), - error_message = VALUES(error_message), - request_metadata = VALUES(request_metadata), - created_at = VALUES(created_at), - submitted_at = VALUES(submitted_at), - completed_at = VALUES(completed_at), - updated_at = VALUES(updated_at) -"#; +{guarded_updates} +"# + ) + }) +} const UPDATE_IF_ACTIVE_SQL: &str = r#" UPDATE video_tasks SET @@ -452,20 +547,38 @@ UPDATE video_tasks SET error_code = ?, error_message = ?, request_metadata = ?, - created_at = ?, + created_at = COALESCE(created_at, ?), submitted_at = ?, completed_at = ?, updated_at = ? WHERE id = ? AND status IN ('pending', 'submitted', 'queued', 'processing') + AND BINARY short_id <=> BINARY ? + AND BINARY request_id <=> BINARY ? + AND BINARY user_id <=> BINARY ? + AND BINARY api_key_id <=> BINARY ? + AND BINARY external_task_id <=> BINARY ? + AND BINARY provider_id <=> BINARY ? + AND BINARY endpoint_id <=> BINARY ? + AND BINARY key_id <=> BINARY ? + AND BINARY client_api_format <=> BINARY ? + AND BINARY provider_api_format <=> BINARY ? + AND format_converted <=> ? + AND BINARY model <=> BINARY ? + AND duration_seconds <=> ? + AND BINARY resolution <=> BINARY ? + AND BINARY aspect_ratio <=> BINARY ? + AND BINARY size <=> BINARY ? "#; fn bind_task<'q>( query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, - task: UpsertVideoTask, + mut task: UpsertVideoTask, include_insert_id: bool, include_update_id: bool, ) -> Result, DataLayerError> { + task.sanitize_for_persistence(); + let identity = task.clone(); let original_request_body = json_to_string(&task.original_request_body)?; let request_metadata = json_to_string(&task.request_metadata)?; let query = if include_insert_id { @@ -535,19 +648,45 @@ fn bind_task<'q>( "video task updated_at", )?); if include_update_id { - Ok(bound.bind(task.id)) + bind_identity_guard(bound.bind(task.id), identity) } else { Ok(bound) } } +fn bind_identity_guard<'q>( + query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, + identity: UpsertVideoTask, +) -> Result, DataLayerError> { + Ok(query + .bind(identity.short_id) + .bind(identity.request_id) + .bind(identity.user_id) + .bind(identity.api_key_id) + .bind(identity.external_task_id) + .bind(identity.provider_id) + .bind(identity.endpoint_id) + .bind(identity.key_id) + .bind(identity.client_api_format) + .bind(identity.provider_api_format) + .bind(identity.format_converted) + .bind(identity.model) + .bind(optional_u32_to_i32( + identity.duration_seconds, + "video task duration_seconds", + )?) + .bind(identity.resolution) + .bind(identity.aspect_ratio) + .bind(identity.size)) +} + fn push_filter<'args>( builder: &mut QueryBuilder<'args, MySql>, filter: &'args VideoTaskQueryFilter, created_since_unix_secs: Option, ) { if let Some(user_id) = filter.user_id.as_deref() { - push_clause(builder, "user_id = "); + push_clause(builder, "BINARY user_id = BINARY "); builder.push_bind(user_id); } if let Some(status) = filter.status { @@ -706,13 +845,86 @@ fn optional_u32_to_i32(value: Option, name: &str) -> Result, Da #[cfg(test)] mod tests { - use super::MysqlVideoTaskRepository; + use super::{ + upsert_sql, MysqlVideoTaskRepository, IMMUTABLE_IDENTITY_MATCH_SQL, UPDATE_IF_ACTIVE_SQL, + UPSERT_UPDATE_COLUMNS, + }; use crate::run_migrations; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository, }; use std::sync::Arc; + #[test] + fn mysql_write_sql_atomically_guards_immutable_identity() { + assert!(IMMUTABLE_IDENTITY_MATCH_SQL.contains("BINARY id <=> BINARY VALUES(id)")); + for column in [ + "short_id", + "request_id", + "user_id", + "api_key_id", + "external_task_id", + "provider_id", + "endpoint_id", + "key_id", + "client_api_format", + "provider_api_format", + "model", + "resolution", + "aspect_ratio", + "size", + ] { + assert!( + IMMUTABLE_IDENTITY_MATCH_SQL + .contains(&format!("BINARY {column} <=> BINARY VALUES({column})")), + "upsert identity predicate should guard {column}" + ); + } + for column in ["format_converted", "duration_seconds"] { + assert!( + IMMUTABLE_IDENTITY_MATCH_SQL.contains(&format!("{column} <=> VALUES({column})")), + "upsert identity predicate should guard {column}" + ); + } + + let upsert = upsert_sql(); + assert!(!UPSERT_UPDATE_COLUMNS.contains(&"created_at")); + for column in UPSERT_UPDATE_COLUMNS { + assert!( + upsert.contains(&format!("{column} = IF((BINARY id <=> BINARY VALUES(id)")), + "upsert assignment should be conditional for {column}" + ); + } + for column in [ + "short_id", + "request_id", + "user_id", + "api_key_id", + "external_task_id", + "provider_id", + "endpoint_id", + "key_id", + "client_api_format", + "provider_api_format", + "model", + "resolution", + "aspect_ratio", + "size", + ] { + assert!( + UPDATE_IF_ACTIVE_SQL.contains(&format!("BINARY {column} <=> BINARY ?")), + "active update should guard {column}" + ); + } + for column in ["format_converted", "duration_seconds"] { + assert!( + UPDATE_IF_ACTIVE_SQL.contains(&format!("{column} <=> ?")), + "active update should guard {column}" + ); + } + assert!(UPDATE_IF_ACTIVE_SQL.contains("created_at = COALESCE(created_at, ?)")); + } + #[tokio::test] async fn repository_builds_from_lazy_pool() { let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with( diff --git a/crates/aether-data/adapters/mysql/src/wallet.rs b/crates/aether-data/adapters/mysql/src/wallet.rs index 8b3cbf935..6418936a4 100644 --- a/crates/aether-data/adapters/mysql/src/wallet.rs +++ b/crates/aether-data/adapters/mysql/src/wallet.rs @@ -2,18 +2,39 @@ use async_trait::async_trait; use chrono::Utc; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; +use aether_data_contracts::repository::billing::{ + checked_plan_duration_days_from_snapshot, entitlements_have_replacement_selector, + entitlements_should_replace_existing, +}; use aether_data_contracts::repository::wallet::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, - redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, - AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, - AdminWalletListQuery, AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, - CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + canonicalize_payment_method, canonicalize_wallet_refund_fields, + payment_callback_amount_matches_order, payment_callback_method_matches_order, + payment_callback_provider_matches_order, payment_order_is_failed_wallet_checkout_placeholder, + payment_order_is_uncertain_wallet_checkout_placeholder, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_payment_method, + redeem_code_refundable_amount, validate_admin_redeem_code_batch_input, + validate_manual_wallet_recharge, validate_payment_order_credit_amounts, + validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements, + validate_redeem_wallet_credit, validate_wallet_recharge_order_input, + wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token, + wallet_recharge_checkout_failed_response, wallet_recharge_checkout_uncertain_response, + 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, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, + AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, + AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, + CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, + ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, @@ -22,7 +43,8 @@ use aether_data_contracts::repository::wallet::{ StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use aether_data_contracts::DataLayerError; @@ -287,7 +309,9 @@ fn push_admin_payment_order_filters<'a>( } if let Some(status) = query.status.as_deref() { builder - .push(" AND (CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ") + .push( + " AND (CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ", + ) .push_bind(now) .push(" THEN 'expired' ELSE status END) = ") .push_bind(status); @@ -490,6 +514,20 @@ impl WalletReadRepository for MysqlWalletReadRepository { ) -> Result, DataLayerError> { initialize_mysql_auth_wallet(&self.pool, Some(user_id), None, initial_gift_usd, unlimited) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_mysql_auth_wallet(&self.pool, Some(user_id), None, initial_gift_usd, unlimited) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn initialize_auth_api_key_wallet( @@ -506,6 +544,26 @@ impl WalletReadRepository for MysqlWalletReadRepository { unlimited, ) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_mysql_auth_wallet( + &self.pool, + None, + Some(api_key_id), + initial_gift_usd, + unlimited, + ) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn update_auth_user_wallet_snapshot( @@ -868,7 +926,7 @@ WHERE wallet_id = ? offset: usize, ) -> Result { let total = read_count_row( - sqlx::query("SELECT COUNT(*) AS total FROM payment_orders WHERE user_id = ?") + sqlx::query("SELECT COUNT(*) AS total FROM payment_orders WHERE user_id = ? AND order_kind = 'wallet_recharge'") .bind(user_id) .fetch_one(&self.pool) .await @@ -882,7 +940,7 @@ SELECT payment_provider, payment_channel, order_kind, product_id, product_snapshot, gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END AS status, created_at AS created_at_unix_ms, @@ -891,6 +949,7 @@ SELECT expires_at AS expires_at_unix_secs FROM payment_orders WHERE user_id = ? + AND order_kind = 'wallet_recharge' ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ? "#, @@ -959,7 +1018,7 @@ SELECT payment_provider, payment_channel, order_kind, product_id, product_snapshot, gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END AS status, created_at AS created_at_unix_ms, @@ -969,6 +1028,7 @@ SELECT FROM payment_orders WHERE user_id = ? AND id = ? + AND order_kind = 'wallet_recharge' LIMIT 1 "#, ) @@ -981,6 +1041,23 @@ LIMIT 1 row.as_ref().map(map_payment_order_row).transpose() } + async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, DataLayerError> { + let sql = payment_order_select_sql( + "WHERE user_id = ? AND order_no = ? AND order_kind = 'wallet_recharge' LIMIT 1", + ); + let row = sqlx::query(&sql) + .bind(user_id) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_payment_order_row).transpose() + } + async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, @@ -1007,6 +1084,19 @@ LIMIT 1 row.as_ref().map(map_payment_order_row).transpose() } + async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, DataLayerError> { + let sql = payment_order_select_sql("WHERE order_no = ? LIMIT 1"); + let row = sqlx::query(&sql) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_payment_order_row).transpose() + } + async fn find_wallet_refund( &self, wallet_id: &str, @@ -1127,18 +1217,404 @@ LIMIT 1 #[async_trait] impl WalletWriteRepository for MysqlWalletReadRepository { + async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: WalletLookupKey<'_>, + ) -> Result { + if wallet_id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let mut tx = self.pool.begin().await.map_sql_err()?; + let select_sql = format!( + r#" +SELECT id +FROM wallets +WHERE id = ? + AND {owner_clause} + AND balance = 0 + AND gift_balance = 0 + AND total_recharged = 0 + AND total_consumed = 0 + AND total_refunded = 0 + AND total_adjusted = 0 + AND limit_mode IN ('finite', 'unlimited') + AND currency = 'USD' + AND status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM `usage` u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1 +FOR UPDATE + "# + ); + let found = sqlx::query_scalar::<_, String>(&select_sql) + .bind(wallet_id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(found_id) = found else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let removed = sqlx::query("DELETE FROM wallets WHERE id = ?") + .bind(&found_id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + + async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if expected.id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let mut tx = self.pool.begin().await.map_sql_err()?; + let select_sql = wallet_select_sql(&format!( + r#"WHERE id = ? + AND {owner_clause} + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM `usage` u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1 +FOR UPDATE"# + )); + let row = sqlx::query(&select_sql) + .bind(&expected.id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if ¤t != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let removed = sqlx::query("DELETE FROM wallets WHERE id = ?") + .bind(&expected.id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + + async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if before.id.trim().is_empty() || after.id.trim().is_empty() { + return Ok(false); + } + if before.id != after.id { + return Err(DataLayerError::InvalidInput( + "wallet restore snapshots must reference the same wallet".to_string(), + )); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet restore requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let owner_matches = match owner { + WalletLookupKey::UserId(user_id) => { + before.user_id.as_deref() == Some(user_id) + && after.user_id.as_deref() == Some(user_id) + && before.api_key_id.is_none() + && after.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + before.api_key_id.as_deref() == Some(api_key_id) + && after.api_key_id.as_deref() == Some(api_key_id) + && before.user_id.is_none() + && after.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if !owner_matches { + return Ok(false); + } + let before_updated_at = i64::try_from(before.updated_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput( + "wallet restore timestamp is outside the supported range".to_string(), + ) + })?; + + // Keep the row lock across compare and update so a concurrent wallet mutation cannot be + // mistaken for the import's own post-state. + let mut tx = self.pool.begin().await.map_sql_err()?; + let select_sql = wallet_select_sql(&format!( + "WHERE id = ? AND {owner_clause} LIMIT 1 FOR UPDATE" + )); + let row = sqlx::query(&select_sql) + .bind(&after.id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if current != *after { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let updated = sqlx::query( + r#" +UPDATE wallets +SET balance = ?, + gift_balance = ?, + limit_mode = ?, + currency = ?, + status = ?, + total_recharged = ?, + total_consumed = ?, + total_refunded = ?, + total_adjusted = ?, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(before.balance) + .bind(before.gift_balance) + .bind(&before.limit_mode) + .bind(&before.currency) + .bind(&before.status) + .bind(before.total_recharged) + .bind(before.total_consumed) + .bind(before.total_refunded) + .bind(before.total_adjusted) + .bind(before_updated_at) + .bind(&before.id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected(); + if updated == 0 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result { + if wallet_id.trim().is_empty() || user_id.trim().is_empty() { + return Ok(false); + } + let mut tx = self.pool.begin().await.map_sql_err()?; + let found_wallet_id = sqlx::query_scalar::<_, String>( + r#" +SELECT w.id +FROM wallets AS w +WHERE w.id = ? + AND w.user_id = ? + AND w.api_key_id IS NULL + AND w.balance = 0 + AND w.gift_balance >= 0 + AND w.total_recharged = 0 + AND w.total_consumed = 0 + AND w.total_refunded = 0 + AND w.total_adjusted = w.gift_balance + AND w.limit_mode IN ('finite', 'unlimited') + AND w.currency = 'USD' + AND w.status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = w.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = w.id + ) + AND NOT EXISTS (SELECT 1 FROM `usage` u WHERE u.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = w.id + ) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = w.id + ) + AND ( + (w.gift_balance = 0 AND NOT EXISTS ( + SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = w.id + )) + OR + (w.gift_balance > 0 + AND (SELECT COUNT(*) FROM wallet_transactions t WHERE t.wallet_id = w.id) = 1 + AND EXISTS ( + SELECT 1 FROM wallet_transactions t + WHERE t.wallet_id = w.id + AND t.category = 'gift' + AND t.reason_code = 'gift_initial' + AND t.amount = w.gift_balance + AND t.balance_before = 0 + AND t.balance_after = w.gift_balance + AND t.recharge_balance_before = 0 + AND t.recharge_balance_after = 0 + AND t.gift_balance_before = 0 + AND t.gift_balance_after = w.gift_balance + AND t.link_type = 'system_task' + AND t.link_id = w.user_id + AND t.operator_id IS NULL + )) + ) +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(wallet_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(found_wallet_id) = found_wallet_id else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + sqlx::query("DELETE FROM wallet_transactions WHERE wallet_id = ?") + .bind(&found_wallet_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + let removed = sqlx::query("DELETE FROM wallets WHERE id = ? AND user_id = ?") + .bind(&found_wallet_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + async fn create_wallet_recharge_order( &self, - input: CreateWalletRechargeOrderInput, + mut input: CreateWalletRechargeOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "manual recharge amount must be finite and positive".to_string(), + )); + } + validate_wallet_recharge_order_input(&input).map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid wallet recharge numeric fields".to_string(), + )); + } + let projected_gateway_response = + project_wallet_recharge_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { DataLayerError::InvalidInput("wallet recharge expires_at overflow".to_string()) })?; - let gateway_response = - json_string(&input.gateway_response, "payment_orders.gateway_response")?; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; let mut tx = self.pool.begin().await.map_sql_err()?; + // `wallets.user_id` remains nullable to preserve deleted-user history; + // validate the live owner explicitly before creating a wallet/order. + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } + let wallet_row = sqlx::query( r#" SELECT id, status @@ -1152,14 +1628,18 @@ FOR UPDATE .fetch_optional(&mut *tx) .await .map_sql_err()?; - let (wallet_id, wallet_status) = if let Some(row) = wallet_row { - (get::(&row, "id")?, get::(&row, "status")?) + let (wallet_id, wallet_status, created_wallet) = if let Some(row) = wallet_row { + ( + get::(&row, "id")?, + get::(&row, "status")?, + false, + ) } else { - let wallet_id = input + let requested_wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - sqlx::query( + let insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, @@ -1167,24 +1647,68 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?) +ON DUPLICATE KEY UPDATE id = id "#, ) - .bind(&wallet_id) + .bind(&requested_wallet_id) .bind(&input.user_id) .bind(now) .bind(now) .execute(&mut *tx) .await .map_sql_err()?; - (wallet_id, "active".to_string()) + let Some(row) = mysql_wallet_by_user_id_for_update(&mut tx, &input.user_id).await? + else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + let wallet_id = get::(&row, "id")?; + let wallet_status = get::(&row, "status")?; + let created_wallet = + insert_result.rows_affected() > 0 && wallet_id == requested_wallet_id; + (wallet_id, wallet_status, created_wallet) }; if wallet_status != "active" { - tx.commit().await.map_sql_err()?; + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } return Ok(CreateWalletRechargeOrderOutcome::WalletInactive); } + if let Some(existing_row) = + mysql_payment_order_by_order_no_for_update(&mut tx, &input.order_no).await? + { + let existing_user_id: Option = existing_row.try_get("user_id").map_sql_err()?; + let existing_kind: String = existing_row.try_get("order_kind").map_sql_err()?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !mysql_wallet_recharge_replay_matches(&existing_row, &wallet_id, &input)? { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + let existing = map_payment_order_row(&existing_row)?; + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + let order_id = uuid::Uuid::new_v4().to_string(); - sqlx::query( + let insert_result = sqlx::query( r#" INSERT INTO payment_orders ( id, order_no, wallet_id, user_id, amount_usd, pay_amount, pay_currency, @@ -1193,6 +1717,7 @@ INSERT INTO payment_orders ( gateway_order_id, gateway_response, status, created_at, expires_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'wallet_recharge', 'pending', ?, ?, 'pending', ?, ?) +ON DUPLICATE KEY UPDATE id = id "#, ) .bind(&order_id) @@ -1214,27 +1739,420 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'wallet_recharge', 'pending', ?, .await .map_sql_err()?; - let row = mysql_payment_order_by_id(&mut tx, &order_id).await?; + let row = if insert_result.rows_affected() > 0 { + // A newly inserted row is identified by the generated id. A + // duplicate-key no-op has a different id and is resolved below. + if let Some(row) = mysql_payment_order_by_id_for_update(&mut tx, &order_id).await? { + row + } else { + let Some(existing_row) = + mysql_payment_order_by_order_no_for_update(&mut tx, &input.order_no).await? + else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge order could not be created".to_string(), + )); + }; + let existing_user_id: Option = + existing_row.try_get("user_id").map_sql_err()?; + let existing_kind: String = existing_row.try_get("order_kind").map_sql_err()?; + let existing = map_payment_order_row(&existing_row)?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !mysql_wallet_recharge_replay_matches(&existing_row, &wallet_id, &input)? { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + } else { + let Some(existing_row) = + mysql_payment_order_by_order_no_for_update(&mut tx, &input.order_no).await? + else { + if mysql_payment_order_by_gateway_order_id_for_update( + &mut tx, + &input.payment_method, + &input.gateway_order_id, + ) + .await? + .is_some() + { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment gateway order already belongs to another order".to_string(), + )); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge order could not be created".to_string(), + )); + }; + let existing_user_id: Option = existing_row.try_get("user_id").map_sql_err()?; + let existing_kind: String = existing_row.try_get("order_kind").map_sql_err()?; + let existing = map_payment_order_row(&existing_row)?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !mysql_wallet_recharge_replay_matches(&existing_row, &wallet_id, &input)? { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + }; tx.commit().await.map_sql_err()?; Ok(CreateWalletRechargeOrderOutcome::Created( map_payment_order_row(&row)?, )) } + async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() || input.gateway_order_id.trim().is_empty() { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout identifiers are required".to_string(), + )); + } + let projected_gateway_response = + match project_wallet_recharge_gateway_response(&input.gateway_response) { + Ok(value) => value, + Err(error) => return Ok(WalletMutationOutcome::Invalid(error)), + }; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; + let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(current_row) = + mysql_payment_order_by_id_for_update(&mut tx, &input.order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = get(¤t_row, "order_kind")?; + if order_kind.as_deref() != Some("wallet_recharge") { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a wallet recharge".to_string(), + )); + } + let current = map_payment_order_row(¤t_row)?; + let current_is_checkout_placeholder = + wallet_recharge_order_is_checkout_placeholder(¤t); + let current_token = current + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + let requested_token = wallet_recharge_checkout_claim_token(&projected_gateway_response); + if current_token.is_some() && current_token != requested_token { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if current.status != "pending" { + if current.gateway_order_id.as_deref() == Some(input.gateway_order_id.as_str()) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(current)); + } + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is no longer pending".to_string(), + )); + } + let now = current_unix_secs_i64(); + if current + .expires_at_unix_secs + .is_none_or(|expires_at| expires_at <= now.max(0) as u64) + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired".to_string(), + )); + } + // A newly-created row uses order_no as a temporary gateway id. Once + // the provider checkout is stored, do not let a concurrent request + // replace that checkout evidence. + if current.gateway_order_id.as_deref().is_some_and(|existing| { + existing != input.gateway_order_id.as_str() + && existing != current.order_no.as_str() + && !current_is_checkout_placeholder + }) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is already bound".to_string(), + )); + } + if let Some(row) = mysql_payment_order_by_gateway_order_id_for_update( + &mut tx, + ¤t.payment_method, + &input.gateway_order_id, + ) + .await? + { + let existing_id: String = row.try_get("id").map_sql_err()?; + if existing_id != input.order_id { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment gateway order already belongs to another order".to_string(), + )); + } + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = ?, gateway_response = ? WHERE id = ? AND status = 'pending' AND expires_at > ?", + ) + .bind(&input.gateway_order_id) + .bind(gateway_response) + .bind(&input.order_id) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired or no longer pending".to_string(), + )); + } + let updated = + map_payment_order_row(&mysql_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + + async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(row) = mysql_payment_order_by_id_for_update(&mut tx, &input.order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(false); + }; + let current = map_payment_order_row(&row)?; + let Some(replacement) = + payment_order_stripe_client_secret_cas_replacement(¤t, &input) + .map_err(DataLayerError::InvalidInput)? + else { + tx.commit().await.map_sql_err()?; + return Ok(false); + }; + let replacement = json_string(&replacement, "payment_orders.gateway_response")?; + let updated = sqlx::query("UPDATE payment_orders SET gateway_response = ? WHERE id = ?") + .bind(replacement) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(updated.rows_affected() == 1) + } + + async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout failure identifiers are required".to_string(), + )); + } + let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(row) = mysql_payment_order_by_id_for_update(&mut tx, &input.order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order = map_payment_order_row(&row)?; + if !wallet_recharge_order_is_checkout_placeholder(&order) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a checkout placeholder".to_string(), + )); + } + let current_token = order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + if current_token != Some(input.claim_token.trim()) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if order.status != "pending" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(order)); + } + let failed = if input.provider_request_may_have_succeeded { + wallet_recharge_checkout_uncertain_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs_i64().max(0) as u64, + ) + } else { + wallet_recharge_checkout_failed_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs_i64().max(0) as u64, + ) + }; + let failed = serde_json::to_string(&failed).map_err(|err| { + DataLayerError::UnexpectedValue(format!("payment_orders.gateway_response: {err}")) + })?; + let updated = sqlx::query( + "UPDATE payment_orders SET status = 'failed', gateway_response = ? WHERE id = ? AND status = 'pending'", + ) + .bind(failed) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + let updated = + map_payment_order_row(&mysql_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + + async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + let now = current_unix_secs_i64().max(0) as u64; + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + || input.expires_at_unix_secs <= now + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim identifiers are invalid".to_string(), + )); + } + if !wallet_recharge_response_is_checkout_placeholder(&input.gateway_response) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge reclaim response must be a placeholder".to_string(), + )); + } + let response = wallet_recharge_checkout_claim_response( + &input.gateway_response, + &input.claim_token, + now, + ) + .map_err(DataLayerError::InvalidInput)?; + let response = serde_json::to_string(&response).map_err(|err| { + DataLayerError::UnexpectedValue(format!("payment_orders.gateway_response: {err}")) + })?; + let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput("wallet recharge expires_at overflow".to_string()) + })?; + let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(row) = mysql_payment_order_by_id_for_update(&mut tx, &input.order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = get(&row, "order_kind")?; + let order = map_payment_order_row(&row)?; + if order_kind.as_deref() != Some("wallet_recharge") + || !wallet_recharge_order_is_reclaimable_placeholder(&order, now) + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is still in progress or already completed".to_string(), + )); + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = ?, gateway_response = ?, status = 'pending', expires_at = ? WHERE id = ?", + ) + .bind(&order.order_no) + .bind(response) + .bind(expires_at) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim lost the order race".to_string(), + )); + } + let updated = + map_payment_order_row(&mysql_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + async fn create_plan_purchase_order( &self, - input: CreatePlanPurchaseOrderInput, + mut input: CreatePlanPurchaseOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_plan_purchase_order_input(&input).map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { DataLayerError::InvalidInput("plan purchase expires_at overflow".to_string()) })?; - let gateway_response = - json_string(&input.gateway_response, "payment_orders.gateway_response")?; + let projected_gateway_response = project_wallet_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; let product_snapshot = json_string(&input.product_snapshot, "payment_orders.product_snapshot")?; let mut tx = self.pool.begin().await.map_sql_err()?; + // MySQL's baseline schema does not declare a wallet owner foreign key. + // Lock and validate the user before any automatic wallet insert. + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } + let wallet_row = sqlx::query( r#" SELECT id, status @@ -1251,11 +2169,14 @@ FOR UPDATE let (wallet_id, wallet_status) = if let Some(row) = wallet_row { (get::(&row, "id")?, get::(&row, "status")?) } else { - let wallet_id = input + let requested_wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - sqlx::query( + // Resolve both the owner uniqueness race and a preferred wallet + // identifier collision through the owner row. This keeps plan + // checkout behavior deterministic across SQL backends. + let _insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, @@ -1263,16 +2184,24 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?) +ON DUPLICATE KEY UPDATE id = id "#, ) - .bind(&wallet_id) + .bind(&requested_wallet_id) .bind(&input.user_id) .bind(now) .bind(now) .execute(&mut *tx) .await .map_sql_err()?; - (wallet_id, "active".to_string()) + let Some(row) = mysql_wallet_by_user_id_for_update(&mut tx, &input.user_id).await? + else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + (get::(&row, "id")?, get::(&row, "status")?) }; if wallet_status != "active" { tx.commit().await.map_sql_err()?; @@ -1382,9 +2311,34 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'plan_purchase', ?, ?, 'pending', &self, input: CreateWalletRefundRequestInput, ) -> Result { + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "refund amount must be finite and greater than zero".to_string(), + )); + } let now = current_unix_secs_i64(); let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(wallet_row) = sqlx::query( + r#" +SELECT id, balance +FROM wallets +WHERE id = ? + AND user_id = ? +LIMIT 1 +FOR UPDATE +"#, + ) + .bind(&input.wallet_id) + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::WalletMissing); + }; + if let Some(idempotency_key) = input.idempotency_key.as_deref() { let existing = mysql_refund_by_idempotency(&mut tx, &input.user_id, idempotency_key).await?; @@ -1396,55 +2350,51 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'plan_purchase', ?, ?, 'pending', } } - let Some(wallet_row) = sqlx::query( - r#" -SELECT id, balance -FROM wallets -WHERE id = ? -LIMIT 1 -FOR UPDATE -"#, - ) - .bind(&input.wallet_id) - .fetch_optional(&mut *tx) - .await - .map_sql_err()? - else { - tx.commit().await.map_sql_err()?; - return Ok(CreateWalletRefundRequestOutcome::WalletMissing); - }; let wallet_recharge_balance: f64 = get(&wallet_row, "balance")?; - let wallet_reserved_amount: f64 = sqlx::query_scalar( + if !wallet_recharge_balance.is_finite() { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet recharge balance is invalid".to_string(), + )); + } + let wallet_reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(SUM(amount_usd), 0) +SELECT amount_usd FROM refund_requests WHERE wallet_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(&input.wallet_id) - .fetch_one(&mut *tx) + .fetch_all(&mut *tx) .await - .map_sql_err()?; + .map_sql_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(wallet_reserved_amount) = wallet_reserved_amount else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet refund reservation is invalid".to_string(), + )); + }; if input.amount_usd > (wallet_recharge_balance - wallet_reserved_amount) { tx.commit().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableBalance); } let mut payment_order_id = None; - let mut source_type = input - .source_type - .clone() - .unwrap_or_else(|| "wallet_balance".to_string()); - let mut source_id = input.source_id.clone(); - let mut refund_mode = input - .refund_mode - .clone() - .unwrap_or_else(|| "offline_payout".to_string()); + let mut resolved_payment_method = None; if let Some(order_id) = input.payment_order_id.as_deref() { let Some(order_row) = sqlx::query( r#" -SELECT id, status, payment_method, refundable_amount_usd +SELECT id, status, payment_method, amount_usd, refunded_amount_usd, refundable_amount_usd FROM payment_orders WHERE id = ? AND wallet_id = ? @@ -1466,19 +2416,46 @@ FOR UPDATE tx.commit().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotRefundable); } - let order_reserved_amount: f64 = sqlx::query_scalar( + let order_reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(SUM(amount_usd), 0) +SELECT amount_usd FROM refund_requests WHERE payment_order_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(order_id) - .fetch_one(&mut *tx) + .fetch_all(&mut *tx) .await - .map_sql_err()?; + .map_sql_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(order_reserved_amount) = order_reserved_amount else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund reservation is invalid".to_string(), + )); + }; + let order_amount: f64 = get(&order_row, "amount_usd")?; + let refunded_amount: f64 = get(&order_row, "refunded_amount_usd")?; let refundable_amount: f64 = get(&order_row, "refundable_amount_usd")?; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_amount, + refundable_amount, + ) { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund amounts are invalid".to_string(), + )); + } if input.amount_usd > (refundable_amount - order_reserved_amount) { tx.commit().await.map_sql_err()?; return Ok( @@ -1486,14 +2463,21 @@ WHERE payment_order_id = ? ); } payment_order_id = Some(order_id.to_string()); - source_type = "payment_order".to_string(); - source_id = Some(order_id.to_string()); - if input.refund_mode.is_none() { - let payment_method: String = get(&order_row, "payment_method")?; - refund_mode = default_refund_mode_for_payment_method(&payment_method).to_string(); - } + resolved_payment_method = Some(get::(&order_row, "payment_method")?); } + let canonical = canonicalize_wallet_refund_fields( + payment_order_id.as_deref(), + input.source_type.as_deref(), + input.source_id.as_deref(), + input.refund_mode.as_deref(), + resolved_payment_method.as_deref(), + ) + .map_err(DataLayerError::InvalidInput)?; + let source_type = canonical.source_type; + let source_id = canonical.source_id; + let refund_mode = canonical.refund_mode; + let refund_id = uuid::Uuid::new_v4().to_string(); let insert = sqlx::query( r#" @@ -1523,7 +2507,11 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending_approval', ?, ?, ?, ?, ?) .await; if let Err(err) = insert { - if input.idempotency_key.is_some() { + if input.idempotency_key.is_some() + && err + .as_database_error() + .is_some_and(|database_error| database_error.is_unique_violation()) + { tx.rollback().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::DuplicateRejected); } @@ -1539,58 +2527,87 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending_approval', ?, ?, ?, ?, ?) async fn process_payment_callback( &self, - input: ProcessPaymentCallbackInput, + mut input: ProcessPaymentCallbackInput, ) -> Result { + input + .canonicalize_and_validate() + .map_err(DataLayerError::InvalidInput)?; + if input.callback_key.trim().is_empty() + || input.callback_key.chars().count() > 128 + || input.payload_hash.trim().is_empty() + || !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid payment callback numeric or identity fields".to_string(), + )); + } let now = current_unix_secs_i64(); let payload = json_string(&input.payload, "payment_callbacks.payload")?; let mut tx = self.pool.begin().await.map_sql_err()?; - let existing_callback = sqlx::query( + let candidate_callback_id = uuid::Uuid::new_v4().to_string(); + sqlx::query( r#" -SELECT id, payment_order_id, status, order_no, gateway_order_id -FROM payment_callbacks -WHERE callback_key = ? -LIMIT 1 -"#, - ) - .bind(&input.callback_key) - .fetch_optional(&mut *tx) - .await - .map_sql_err()?; - let duplicate = existing_callback.is_some(); - let callback_id = if let Some(row) = existing_callback.as_ref() { - let status: String = get(row, "status")?; - if status == "processed" { - let order_id: Option = get(row, "payment_order_id")?; - tx.commit().await.map_sql_err()?; - return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id }); - } - get(row, "id")? - } else { - let callback_id = uuid::Uuid::new_v4().to_string(); - sqlx::query( - r#" INSERT INTO payment_callbacks ( id, payment_order_id, payment_method, callback_key, order_no, gateway_order_id, payload_hash, signature_valid, status, payload, error_message, created_at, processed_at ) -VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) +VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', NULL, NULL, ?, NULL) +ON DUPLICATE KEY UPDATE id = id "#, - ) - .bind(&callback_id) - .bind(&input.payment_method) - .bind(&input.callback_key) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) - .bind(&input.payload_hash) - .bind(input.signature_valid) - .bind(&payload) - .bind(now) - .execute(&mut *tx) - .await - .map_sql_err()?; - callback_id - }; + ) + .bind(&candidate_callback_id) + .bind(&input.payment_method) + .bind(&input.callback_key) + .bind(input.order_no.as_deref()) + .bind(input.gateway_order_id.as_deref()) + .bind(&input.payload_hash) + .bind(input.signature_valid) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + let callback_row = sqlx::query( + r#" +SELECT id, payment_order_id, payment_method, payload_hash, status, order_no, gateway_order_id +FROM payment_callbacks +WHERE callback_key = ? +LIMIT 1 +FOR UPDATE +"#, + ) + .bind(&input.callback_key) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let callback_id: String = get(&callback_row, "id")?; + let duplicate = callback_id != candidate_callback_id; + let callback_order_no: Option = get(&callback_row, "order_no")?; + let callback_gateway_order_id: Option = get(&callback_row, "gateway_order_id")?; + let stored_method: String = get(&callback_row, "payment_method")?; + let stored_hash: Option = get(&callback_row, "payload_hash")?; + if !stored_method.eq_ignore_ascii_case(&input.payment_method) + || stored_hash.as_deref() != Some(input.payload_hash.as_str()) + { + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "callback key reused with different payment payload".to_string(), + }); + } + let status: String = get(&callback_row, "status")?; + if status == "processed" { + let order_id: Option = get(&callback_row, "payment_order_id")?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id }); + } if !input.signature_valid { update_mysql_payment_callback_failure( @@ -1608,20 +2625,20 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) }); } - let lookup_order_no = input.order_no.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| get(row, "order_no").ok()) - }); - let lookup_gateway_order_id = input.gateway_order_id.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| get(row, "gateway_order_id").ok()) - }); + let lookup_order_no = input.order_no.clone().or_else(|| callback_order_no.clone()); + let lookup_gateway_order_id = input + .gateway_order_id + .clone() + .or_else(|| callback_gateway_order_id.clone()); let order_row = if let Some(order_no) = lookup_order_no.as_deref() { mysql_payment_order_by_order_no_for_update(&mut tx, order_no).await? } else if let Some(gateway_order_id) = lookup_gateway_order_id.as_deref() { - mysql_payment_order_by_gateway_order_id_for_update(&mut tx, gateway_order_id).await? + mysql_payment_order_by_gateway_order_id_for_update( + &mut tx, + &input.payment_method, + gateway_order_id, + ) + .await? } else { None }; @@ -1647,19 +2664,263 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) let order_payment_method: String = get(&order_row, "payment_method")?; let order_payment_provider: Option = get(&order_row, "payment_provider")?; let order_payment_channel: Option = get(&order_row, "payment_channel")?; + let order_pay_currency: Option = get(&order_row, "pay_currency")?; + let order_gateway_order_id: Option = get(&order_row, "gateway_order_id")?; let order_kind: String = get(&order_row, "order_kind")?; let order_amount_usd: f64 = get(&order_row, "amount_usd")?; let order_pay_amount: Option = get(&order_row, "pay_amount")?; + let order_exchange_rate: Option = get(&order_row, "exchange_rate")?; let order_status: String = get(&order_row, "status")?; let expires_at_unix_secs: Option = get(&order_row, "expires_at_unix_secs")?; - - let amount_matches = if let (Some(callback_pay_amount), Some(order_pay_amount)) = - (input.pay_amount, order_pay_amount) - { - (callback_pay_amount - order_pay_amount).abs() <= 0.01 + let order_gateway_response = if order_status.eq_ignore_ascii_case("failed") { + optional_json( + get(&order_row, "gateway_response")?, + "payment_orders.gateway_response", + )? } else { - (input.amount_usd - order_amount_usd).abs() <= f64::EPSILON + None }; + let failed_checkout_recoverable = payment_order_is_failed_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + let uncertain_checkout = payment_order_is_uncertain_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + if !order_amount_usd.is_finite() + || order_amount_usd <= 0.0 + || order_pay_amount.is_some_and(|value| !value.is_finite() || value <= 0.0) + { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order amount is invalid", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order amount is invalid".to_string(), + }); + } + + // A payment order may only credit a wallet owned by the same user. + // Lock the wallet (and its API-key owner row through the join) before + // any gateway binding, entitlement, wallet, or order mutation. Legacy + // rows with an ambiguous owner shape are deliberately rejected. + let order_user_id: Option = get(&order_row, "user_id")?; + let Some(order_user_id) = order_user_id + .as_deref() + .filter(|value| !value.trim().is_empty()) + else { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order user missing", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order user missing".to_string(), + }); + }; + let Some(wallet_owner_row) = sqlx::query( + r#" +SELECT + w.user_id AS wallet_user_id, + w.api_key_id AS wallet_api_key_id, + api_keys.user_id AS api_key_user_id +FROM wallets AS w +LEFT JOIN api_keys ON api_keys.id = w.api_key_id +WHERE w.id = ? +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(&order_wallet_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + else { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "wallet not found", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet not found".to_string(), + }); + }; + let wallet_user_id: Option = get(&wallet_owner_row, "wallet_user_id")?; + let wallet_api_key_id: Option = get(&wallet_owner_row, "wallet_api_key_id")?; + let api_key_user_id: Option = get(&wallet_owner_row, "api_key_user_id")?; + let wallet_owner_matches = match ( + wallet_user_id.as_deref(), + wallet_api_key_id.as_deref(), + api_key_user_id.as_deref(), + ) { + (Some(wallet_user_id), None, _) if !wallet_user_id.trim().is_empty() => { + wallet_user_id == order_user_id + } + (None, Some(wallet_api_key_id), Some(api_key_user_id)) + if !wallet_api_key_id.trim().is_empty() && !api_key_user_id.trim().is_empty() => + { + api_key_user_id == order_user_id + } + _ => false, + }; + if !wallet_owner_matches { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order wallet owner mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order wallet owner mismatch".to_string(), + }); + } + + // The lookup identifier is not proof that the callback belongs to + // this order: order_no takes precedence over gateway_order_id. Check + // every identifier supplied by this delivery (and any persisted + // fallback from the callback row) before changing the order or + // wallet. Orders created before the gateway returns a provider + // transaction id store order_no as a placeholder; that value may be + // replaced by a verified callback, but a real id must never be + // rebound to another order. + if input + .order_no + .as_deref() + .is_some_and(|value| value != order_no) + || callback_order_no + .as_deref() + .is_some_and(|value| value != order_no) + { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order number mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order number mismatch".to_string(), + }); + } + let input_gateway_order_id = input + .gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let callback_gateway_order_id = callback_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let stored_real_gateway_order_id = order_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty() && *value != order_no); + let input_real_gateway_order_id = input_gateway_order_id.filter(|value| *value != order_no); + let callback_real_gateway_order_id = + callback_gateway_order_id.filter(|value| *value != order_no); + let effective_gateway_order_id = input_real_gateway_order_id + .or(callback_real_gateway_order_id) + .or(stored_real_gateway_order_id); + if let Some(expected_gateway_order_id) = stored_real_gateway_order_id { + if input_gateway_order_id.is_some_and(|value| value != expected_gateway_order_id) + || callback_gateway_order_id.is_some_and(|value| value != expected_gateway_order_id) + { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order mismatch".to_string(), + }); + } + } else if let (Some(input_gateway), Some(callback_gateway)) = + (input_real_gateway_order_id, callback_real_gateway_order_id) + { + if input_gateway != callback_gateway { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order identifier mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order identifier mismatch".to_string(), + }); + } + } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + let conflicting_order_id: Option = sqlx::query_scalar( + "SELECT id FROM payment_orders WHERE payment_method = ? AND gateway_order_id = ? AND id <> ? LIMIT 1", + ) + .bind(&order_payment_method) + .bind(gateway_order_id) + .bind(&order_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if conflicting_order_id.is_some() { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order belongs to another payment order", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order".to_string(), + }); + } + } + } + + let amount_matches = payment_callback_amount_matches_order( + order_amount_usd, + order_pay_amount, + order_pay_currency.as_deref(), + order_exchange_rate, + input.amount_usd, + input.pay_amount, + ); if !amount_matches { update_mysql_payment_callback_failure( &mut tx, @@ -1675,7 +2936,12 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) error: "callback amount mismatch".to_string(), }); } - if !order_payment_method.eq_ignore_ascii_case(&input.payment_method) { + if !payment_callback_method_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ) { update_mysql_payment_callback_failure( &mut tx, &callback_id, @@ -1690,31 +2956,57 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) error: "payment method mismatch".to_string(), }); } - if let Some(expected_provider) = input.payment_provider.as_deref() { - if order_payment_provider - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_provider)) - { - update_mysql_payment_callback_failure( - &mut tx, - &callback_id, - &input, - &payload, - "payment provider mismatch", - ) - .await?; - tx.commit().await.map_sql_err()?; - return Ok(ProcessPaymentCallbackOutcome::Failed { - duplicate, - error: "payment provider mismatch".to_string(), - }); - } + let payment_provider_matches = payment_callback_provider_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ); + if !payment_provider_matches { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment provider mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment provider mismatch".to_string(), + }); + } + let currency_matches = match (input.pay_currency.as_deref(), order_pay_currency.as_deref()) + { + (Some(callback), Some(order)) => order.eq_ignore_ascii_case(callback), + (None, None) => input.pay_amount.is_none() && order_pay_amount.is_none(), + _ => false, + }; + if !currency_matches { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment currency mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment currency mismatch".to_string(), + }); } if let Some(expected_channel) = input.payment_channel.as_deref() { - if order_payment_channel - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_channel)) - { + let stored_channel = order_payment_channel.as_deref().or_else(|| { + (order_payment_provider.is_none() + && ["alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(&order_payment_method))) + .then_some(order_payment_method.as_str()) + }); + if stored_channel.is_none_or(|value| !value.eq_ignore_ascii_case(expected_channel)) { update_mysql_payment_callback_failure( &mut tx, &callback_id, @@ -1748,14 +3040,16 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) wallet_id: order_wallet_id, }); } - if matches!(order_status.as_str(), "failed" | "expired" | "refunded") { + if !matches!(order_status.as_str(), "pending" | "paid") && !failed_checkout_recoverable { let error = format!("payment order is not creditable: {order_status}"); update_mysql_payment_callback_failure(&mut tx, &callback_id, &input, &payload, &error) .await?; tx.commit().await.map_sql_err()?; return Ok(ProcessPaymentCallbackOutcome::Failed { duplicate, error }); } - if order_status == "pending" && expires_at_unix_secs.is_some_and(|value| value < now) { + if (order_status == "pending" || (failed_checkout_recoverable && !uncertain_checkout)) + && expires_at_unix_secs.is_some_and(|value| value <= now) + { sqlx::query("UPDATE payment_orders SET status = 'expired' WHERE id = ?") .bind(&order_id) .execute(&mut *tx) @@ -1776,6 +3070,28 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) }); } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + if !mysql_bind_payment_gateway_order_id(&mut tx, &order_id, gateway_order_id) + .await? + { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order belongs to another payment order", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order".to_string(), + }); + } + } + } + if order_kind == "plan_purchase" { let order_user_id: Option = get(&order_row, "user_id")?; let Some(user_id) = order_user_id else { @@ -1887,7 +3203,7 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) .bind(&plan_id) .bind(&order_id) .bind(now) - .bind(plan_expires_at_unix(&snapshot, now)) + .bind(plan_expires_at_unix(&snapshot, now)?) .bind(json_string( &entitlements, "user_plan_entitlements.entitlements_snapshot", @@ -1912,9 +3228,9 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) UPDATE payment_orders SET gateway_order_id = COALESCE(?, gateway_order_id), gateway_response = ?, - pay_amount = COALESCE(?, pay_amount), - pay_currency = COALESCE(?, pay_currency), - exchange_rate = COALESCE(?, exchange_rate), + pay_amount = COALESCE(pay_amount, ?), + pay_currency = COALESCE(pay_currency, ?), + exchange_rate = COALESCE(exchange_rate, ?), status = 'credited', fulfillment_status = 'fulfilled', fulfillment_error = NULL, @@ -1924,8 +3240,11 @@ SET gateway_order_id = COALESCE(?, gateway_order_id), WHERE id = ? "#, ) - .bind(input.gateway_order_id.as_deref()) - .bind(&payload) + .bind(effective_gateway_order_id) + .bind(json_string( + &input.gateway_response_projection(&order_no, effective_gateway_order_id), + "payment_orders.gateway_response", + )?) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -1957,7 +3276,7 @@ WHERE id = ? let Some(wallet_row) = sqlx::query( r#" -SELECT id, status, balance, gift_balance +SELECT id, status, balance, gift_balance, total_recharged FROM wallets WHERE id = ? LIMIT 1 @@ -2002,6 +3321,33 @@ FOR UPDATE let before_recharge: f64 = get(&wallet_row, "balance")?; let before_gift: f64 = get(&wallet_row, "gift_balance")?; + let total_recharged: f64 = get(&wallet_row, "total_recharged")?; + // Finite recharge balances may be negative: usage settlement permits a + // finite wallet to overdraft, and a later recharge must be able to + // restore that balance. Reject only malformed values and arithmetic + // overflow here. + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order_amount_usd).is_finite() + || !(before_recharge + before_gift + order_amount_usd).is_finite() + { + update_mysql_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "wallet balance is invalid", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet balance is invalid".to_string(), + }); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order_amount_usd; let after_total = after_recharge + before_gift; @@ -2052,9 +3398,9 @@ VALUES (?, ?, 'recharge', 'topup_gateway', ?, ?, ?, ?, ?, ?, ?, 'payment_order', UPDATE payment_orders SET gateway_order_id = COALESCE(?, gateway_order_id), gateway_response = ?, - pay_amount = COALESCE(?, pay_amount), - pay_currency = COALESCE(?, pay_currency), - exchange_rate = COALESCE(?, exchange_rate), + pay_amount = COALESCE(pay_amount, ?), + pay_currency = COALESCE(pay_currency, ?), + exchange_rate = COALESCE(exchange_rate, ?), status = 'credited', paid_at = COALESCE(paid_at, ?), credited_at = ?, @@ -2062,8 +3408,11 @@ SET gateway_order_id = COALESCE(?, gateway_order_id), WHERE id = ? "#, ) - .bind(input.gateway_order_id.as_deref()) - .bind(&payload) + .bind(effective_gateway_order_id) + .bind(json_string( + &input.gateway_response_projection(&order_no, effective_gateway_order_id), + "payment_orders.gateway_response", + )?) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -2097,6 +3446,11 @@ WHERE id = ? &self, input: AdjustWalletBalanceInput, ) -> 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(), + )); + } let now = current_unix_secs_i64(); let mut tx = self.pool.begin().await.map_sql_err()?; let Some(row) = mysql_wallet_by_id_for_update(&mut tx, &input.wallet_id).await? else { @@ -2107,6 +3461,16 @@ WHERE id = ? let before_recharge: f64 = get(&row, "balance")?; let before_gift: f64 = get(&row, "gift_balance")?; let before_total = before_recharge + before_gift; + let before_total_adjusted: f64 = get(&row, "total_adjusted")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || !before_total.is_finite() + || !before_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid".to_string(), + )); + } let mut after_recharge = before_recharge; let mut after_gift = before_gift; apply_admin_balance_adjustment( @@ -2115,6 +3479,17 @@ WHERE id = ? &mut after_recharge, &mut after_gift, ); + let after_total = after_recharge + after_gift; + let after_total_adjusted = before_total_adjusted + input.amount_usd; + if !after_recharge.is_finite() + || !after_gift.is_finite() + || !after_total.is_finite() + || !after_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow during admin adjustment".to_string(), + )); + } sqlx::query( r#" @@ -2157,7 +3532,7 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, .bind(&input.wallet_id) .bind(input.amount_usd) .bind(before_total) - .bind(after_recharge + after_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -2180,7 +3555,7 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, reason_code: "adjust_admin".to_string(), amount: input.amount_usd, balance_before: before_total, - balance_after: after_recharge + after_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -2198,8 +3573,15 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, async fn create_manual_wallet_recharge( &self, - input: CreateManualWalletRechargeInput, + mut input: CreateManualWalletRechargeInput, ) -> Result, DataLayerError> { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "manual recharge amount must be finite and positive".to_string(), + )); + } let now = current_unix_secs_i64(); let mut tx = self.pool.begin().await.map_sql_err()?; let Some(wallet_row) = mysql_wallet_by_id_for_update(&mut tx, &input.wallet_id).await? @@ -2210,6 +3592,14 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, let before_recharge: f64 = get(&wallet_row, "balance")?; let before_gift: f64 = get(&wallet_row, "gift_balance")?; + let before_total_recharged: f64 = get(&wallet_row, "total_recharged")?; + let (after_recharge, after_total_recharged) = validate_manual_wallet_recharge( + input.amount_usd, + before_recharge, + before_gift, + before_total_recharged, + ) + .map_err(DataLayerError::InvalidInput)?; let user_id: Option = get(&wallet_row, "user_id")?; let order_id = uuid::Uuid::new_v4().to_string(); let gateway_response = json_string( @@ -2246,18 +3636,17 @@ VALUES (?, ?, ?, ?, ?, 0, ?, ?, 'credited', ?, ?, ?, ?) .await .map_sql_err()?; - let after_recharge = before_recharge + input.amount_usd; sqlx::query( r#" UPDATE wallets SET balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(input.amount_usd) + .bind(after_total_recharged) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) @@ -2333,6 +3722,12 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) return Ok(WalletMutationOutcome::NotFound); }; let refund = map_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if !matches!(refund.status.as_str(), "approved" | "pending_approval") { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( @@ -2349,9 +3744,28 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) }; let before_recharge: f64 = get(&wallet_row, "balance")?; let before_gift: f64 = get(&wallet_row, "gift_balance")?; - let before_total = before_recharge + before_gift; + let before_total_refunded: f64 = get(&wallet_row, "total_refunded")?; let amount_usd = refund.amount_usd; let after_recharge = before_recharge - amount_usd; + let before_total = before_recharge + before_gift; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded + amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } if after_recharge < 0.0 { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( @@ -2368,45 +3782,78 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) "payment order not found".to_string(), )); }; - let refundable_amount: f64 = get(&order_row, "refundable_amount_usd")?; - if amount_usd > refundable_amount { + let order_wallet_id: String = get(&order_row, "wallet_id")?; + let order_status: String = get(&order_row, "status")?; + if order_wallet_id != input.wallet_id || order_status != "credited" { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( - "refund amount exceeds refundable amount".to_string(), + "payment order is not refundable for this wallet".to_string(), )); } - sqlx::query( + let order_amount: f64 = get(&order_row, "amount_usd")?; + let refunded_before: f64 = get(&order_row, "refunded_amount_usd")?; + let refundable_before: f64 = get(&order_row, "refundable_amount_usd")?; + let refunded_after = refunded_before + amount_usd; + let refundable_after = refundable_before - amount_usd; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_before, + refundable_before, + ) || amount_usd > refundable_before + || !refunded_after.is_finite() + || refunded_after < 0.0 + || refunded_after > order_amount + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order_amount + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + let result = sqlx::query( r#" UPDATE payment_orders -SET refunded_amount_usd = refunded_amount_usd + ?, - refundable_amount_usd = refundable_amount_usd - ? +SET refunded_amount_usd = ?, + refundable_amount_usd = ? WHERE id = ? "#, ) - .bind(amount_usd) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .bind(payment_order_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund processing".to_string(), + )); + } } - sqlx::query( + let result = sqlx::query( r#" UPDATE wallets SET balance = ?, - total_refunded = total_refunded + ?, + total_refunded = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "wallet disappeared during refund processing".to_string(), + )); + } let wallet = map_wallet_row(&mysql_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let transaction_id = uuid::Uuid::new_v4().to_string(); @@ -2489,11 +3936,6 @@ WHERE id = ? AND wallet_id = ? input: CompleteAdminWalletRefundInput, ) -> Result, DataLayerError> { let now = current_unix_secs_i64(); - let payout_proof = input - .payout_proof - .as_ref() - .map(|value| json_string(value, "refund_requests.payout_proof")) - .transpose()?; let mut tx = self.pool.begin().await.map_sql_err()?; let Some(current_refund) = mysql_refund_by_id_and_wallet_for_update(&mut tx, &input.refund_id, &input.wallet_id) @@ -2502,24 +3944,57 @@ WHERE id = ? AND wallet_id = ? tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::NotFound); }; - let status: String = get(¤t_refund, "status")?; - if status != "processing" { + let refund = map_refund_row(¤t_refund)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let (Some(existing_id), Some(incoming_id)) = ( + refund.gateway_refund_id.as_deref(), + input.gateway_refund_id.as_deref(), + ) { + if existing_id != incoming_id { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if refund.status == "succeeded" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(refund)); + } + if refund.status != "processing" { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( "refund status must be processing before completion".to_string(), )); } + // Preserve a processing proof for ordinary replays, but allow an + // explicit successful gateway proof to upgrade it at completion. + let selected_payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| refund.payout_proof.is_none() || wallet_refund_proof_is_success(proof)) + .cloned() + .or_else(|| refund.payout_proof.clone()); + let payout_proof = selected_payout_proof + .as_ref() + .map(|value| json_string(value, "refund_requests.payout_proof")) + .transpose()?; - sqlx::query( + let refund_update = sqlx::query( r#" UPDATE refund_requests SET status = 'succeeded', - gateway_refund_id = ?, - payout_reference = ?, + gateway_refund_id = COALESCE(gateway_refund_id, ?), + payout_reference = COALESCE(payout_reference, ?), payout_proof = ?, completed_at = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' "#, ) .bind(input.gateway_refund_id.as_deref()) @@ -2532,11 +4007,107 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if refund_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during refund completion".to_string(), + )); + } let refund = map_refund_row(&mysql_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; Ok(WalletMutationOutcome::Applied(refund)) } + async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result, DataLayerError> { + if input.gateway_refund_id.trim().is_empty() || input.gateway_refund_id.len() > 128 { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier is invalid".to_string(), + )); + } + if input + .payout_proof + .as_ref() + .is_some_and(|proof| !proof.is_object()) + { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund proof must be an object".to_string(), + )); + } + let now = current_unix_secs_i64(); + let mut tx = self.pool.begin().await.map_sql_err()?; + let Some(current_row) = + mysql_refund_by_id_and_wallet_for_update(&mut tx, &input.refund_id, &input.wallet_id) + .await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let current = map_refund_row(¤t_row)?; + if !current.amount_usd.is_finite() || current.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let Some(existing_id) = current.gateway_refund_id.as_deref() { + if existing_id != input.gateway_refund_id { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if current.status == "succeeded" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(current)); + } + if current.status != "processing" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund status must be processing before gateway update".to_string(), + )); + } + // Do not overwrite durable processing evidence with an arbitrary + // replay; only a terminal success proof may replace it. + let selected_payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| current.payout_proof.is_none() || wallet_refund_proof_is_success(proof)) + .cloned() + .or_else(|| current.payout_proof.clone()); + let proof = selected_payout_proof + .as_ref() + .map(|value| json_string(value, "refund_requests.payout_proof")) + .transpose()?; + let gateway_update = sqlx::query( + r#" +UPDATE refund_requests +SET gateway_refund_id = COALESCE(gateway_refund_id, ?), + payout_proof = ?, + updated_at = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' +"#, + ) + .bind(&input.gateway_refund_id) + .bind(proof.as_deref()) + .bind(now) + .bind(&input.refund_id) + .bind(&input.wallet_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if gateway_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during gateway evidence update".to_string(), + )); + } + let updated = map_refund_row(&mysql_refund_by_id(&mut tx, &input.refund_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + async fn fail_admin_wallet_refund( &self, input: FailAdminWalletRefundInput, @@ -2558,15 +4129,21 @@ WHERE id = ? AND wallet_id = ? return Ok(WalletMutationOutcome::NotFound); }; let refund = map_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if matches!(refund.status.as_str(), "pending_approval" | "approved") { - sqlx::query( + let refund_update = sqlx::query( r#" UPDATE refund_requests SET status = 'failed', failure_reason = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(&input.reason) @@ -2576,6 +4153,11 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if refund_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during refund failure".to_string(), + )); + } let wallet = map_wallet_row(&mysql_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let refund = map_refund_row(&mysql_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; @@ -2590,6 +4172,22 @@ WHERE id = ? AND wallet_id = ? ))); } + // Only an explicitly offline payout can be released without external + // settlement evidence. An original-channel refund may still be in + // flight between the provider request and the evidence update. + if refund.gateway_refund_id.is_some() + || refund.payout_proof.is_some() + || !refund + .refund_mode + .trim() + .eq_ignore_ascii_case("offline_payout") + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "cannot fail refund while gateway settlement is processing".to_string(), + )); + } + let Some(wallet_row) = mysql_wallet_by_id_for_update(&mut tx, &input.wallet_id).await? else { tx.commit().await.map_sql_err()?; @@ -2600,26 +4198,95 @@ WHERE id = ? AND wallet_id = ? let amount_usd = refund.amount_usd; let before_recharge: f64 = get(&wallet_row, "balance")?; let before_gift: f64 = get(&wallet_row, "gift_balance")?; + let before_total_refunded: f64 = get(&wallet_row, "total_refunded")?; let before_total = before_recharge + before_gift; let after_recharge = before_recharge + amount_usd; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded - amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || before_total_refunded < amount_usd + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + || after_total_refunded < 0.0 + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid for refund recovery".to_string(), + )); + } - sqlx::query( + let mut order_amounts = None; + if let Some(payment_order_id) = refund.payment_order_id.as_deref() { + let Some(order_row) = + mysql_payment_order_by_id_for_update(&mut tx, payment_order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order not found".to_string(), + )); + }; + let order = map_payment_order_row(&order_row)?; + if order.wallet_id != input.wallet_id || order.status != "credited" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not refundable for this wallet".to_string(), + )); + } + let refunded_before = order.refunded_amount_usd; + let refundable_before = order.refundable_amount_usd; + let refunded_after = refunded_before - amount_usd; + let refundable_after = refundable_before + amount_usd; + if !payment_order_refund_amounts_are_consistent( + order.amount_usd, + refunded_before, + refundable_before, + ) || refunded_before < amount_usd + || !refunded_after.is_finite() + || refunded_after < 0.0 + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order.amount_usd + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + order_amounts = Some(( + payment_order_id.to_string(), + refunded_after, + refundable_after, + )); + } + + let wallet_result = sqlx::query( r#" UPDATE wallets SET balance = ?, - total_refunded = GREATEST(total_refunded - ?, 0), + total_refunded = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) .await .map_sql_err()?; - let wallet = map_wallet_row(&mysql_wallet_by_id(&mut tx, &input.wallet_id).await?)?; + if wallet_result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "wallet disappeared during refund recovery".to_string(), + )); + } let transaction_id = uuid::Uuid::new_v4().to_string(); sqlx::query( @@ -2636,7 +4303,7 @@ VALUES (?, ?, 'refund', 'refund_revert', ?, ?, ?, ?, ?, ?, ?, 'refund_request', .bind(&input.wallet_id) .bind(amount_usd) .bind(before_total) - .bind(after_recharge + before_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -2648,30 +4315,35 @@ VALUES (?, ?, 'refund', 'refund_revert', ?, ?, ?, ?, ?, ?, ?, 'refund_request', .await .map_sql_err()?; - if let Some(payment_order_id) = refund.payment_order_id.as_deref() { - sqlx::query( + if let Some((payment_order_id, refunded_after, refundable_after)) = order_amounts { + let result = sqlx::query( r#" UPDATE payment_orders -SET refunded_amount_usd = refunded_amount_usd - ?, - refundable_amount_usd = refundable_amount_usd + ? +SET refunded_amount_usd = ?, + refundable_amount_usd = ? WHERE id = ? "#, ) - .bind(amount_usd) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .bind(payment_order_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund recovery".to_string(), + )); + } } - sqlx::query( + let result = sqlx::query( r#" UPDATE refund_requests SET status = 'failed', failure_reason = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' "#, ) .bind(&input.reason) @@ -2681,6 +4353,12 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during recovery".to_string(), + )); + } + let wallet = map_wallet_row(&mysql_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let refund = map_refund_row(&mysql_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; Ok(WalletMutationOutcome::Applied(( @@ -2693,7 +4371,7 @@ WHERE id = ? AND wallet_id = ? reason_code: "refund_revert".to_string(), amount: amount_usd, balance_before: before_total, - balance_after: after_recharge + before_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -2830,15 +4508,31 @@ WHERE id = ? AND wallet_id = ? } if order .expires_at_unix_secs - .is_some_and(|value| value < now as u64) + .is_some_and(|value| value <= now as u64) { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( "payment order expired".to_string(), )); } - let order_kind: String = get(&order_row, "order_kind")?; + let order_payment_provider: Option = get(&order_row, "payment_provider")?; + let order_payment_channel: Option = get(&order_row, "payment_channel")?; + if validate_payment_order_credit_amounts( + &order_kind, + &order.payment_method, + order_payment_provider.as_deref(), + order_payment_channel.as_deref(), + order.amount_usd, + order.pay_amount, + ) + .is_err() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order amount is invalid".to_string(), + )); + } if order_kind == "plan_purchase" { let order_user_id: Option = get(&order_row, "user_id")?; let Some(user_id) = order_user_id else { @@ -2927,7 +4621,7 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) .bind(&plan_id) .bind(&input.order_id) .bind(now) - .bind(plan_expires_at_unix(&snapshot, now)) + .bind(plan_expires_at_unix(&snapshot, now)?) .bind(json_string( &entitlements, "user_plan_entitlements.entitlements_snapshot", @@ -3022,6 +4716,20 @@ WHERE id = ? let before_recharge: f64 = get(&wallet_row, "balance")?; let before_gift: f64 = get(&wallet_row, "gift_balance")?; + let total_recharged: f64 = get(&wallet_row, "total_recharged")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order.amount_usd).is_finite() + || !(before_recharge + before_gift + order.amount_usd).is_finite() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order.amount_usd; sqlx::query( @@ -3126,9 +4834,19 @@ WHERE id = ? &self, input: CreateAdminRedeemCodeBatchInput, ) -> Result { + validate_admin_redeem_code_batch_input(&input).map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let batch_id = uuid::Uuid::new_v4().to_string(); - let expires_at = input.expires_at_unix_secs.map(|value| value as i64); + let expires_at = input + .expires_at_unix_secs + .map(|value| { + i64::try_from(value).map_err(|_| { + DataLayerError::InvalidInput( + "redeem code batch expires_at overflow".to_string(), + ) + }) + }) + .transpose()?; let mut tx = self.pool.begin().await.map_sql_err()?; sqlx::query( @@ -3464,7 +5182,6 @@ FOR UPDATE let batch_name: String = get(&code_row, "batch_name")?; let balance_bucket: String = get(&code_row, "balance_bucket")?; let amount_usd: f64 = get(&code_row, "amount_usd")?; - let credits_recharge_balance = redeem_code_credits_recharge_balance(&balance_bucket); let wallet_row = mysql_wallet_by_user_id_for_update(&mut tx, &input.user_id).await?; let wallet_id = if let Some(row) = wallet_row.as_ref() { @@ -3478,11 +5195,16 @@ FOR UPDATE uuid::Uuid::new_v4().to_string() }; - let (before_recharge, before_gift) = if let Some(row) = wallet_row.as_ref() { - (get(row, "balance")?, get(row, "gift_balance")?) - } else { - sqlx::query( - r#" + let (before_recharge, before_gift, before_total_recharged) = + if let Some(row) = wallet_row.as_ref() { + ( + get(row, "balance")?, + get(row, "gift_balance")?, + get(row, "total_recharged")?, + ) + } else { + sqlx::query( + r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, total_recharged, total_consumed, total_refunded, total_adjusted, @@ -3490,39 +5212,37 @@ INSERT INTO wallets ( ) VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?) "#, - ) - .bind(&wallet_id) - .bind(&input.user_id) - .bind(now) - .bind(now) - .execute(&mut *tx) - .await - .map_sql_err()?; - (0.0, 0.0) - }; - let after_recharge = if credits_recharge_balance { - before_recharge + amount_usd - } else { - before_recharge - }; - let after_gift = if credits_recharge_balance { - before_gift - } else { - before_gift + amount_usd - }; + ) + .bind(&wallet_id) + .bind(&input.user_id) + .bind(now) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + (0.0, 0.0, 0.0) + }; + let (after_recharge, after_gift, after_total_recharged) = validate_redeem_wallet_credit( + &balance_bucket, + amount_usd, + before_recharge, + before_gift, + before_total_recharged, + ) + .map_err(DataLayerError::UnexpectedValue)?; sqlx::query( r#" UPDATE wallets SET balance = ?, gift_balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) .bind(after_gift) - .bind(amount_usd) + .bind(after_total_recharged) .bind(now) .bind(&wallet_id) .execute(&mut *tx) @@ -3825,34 +5545,14 @@ fn plan_purchase_limit_scope(snapshot: &serde_json::Value) -> &str { } } -fn plan_replacement_entitlement_types(snapshot: &serde_json::Value) -> Vec<&'static str> { - let entitlements = plan_entitlements_snapshot(snapshot); - let mut kinds = Vec::new(); - if entitlement_snapshot_has_type(&entitlements, "daily_quota") { - kinds.push("daily_quota"); - } - if entitlement_snapshot_has_type(&entitlements, "membership_group") { - kinds.push("membership_group"); - } - kinds -} - -fn entitlement_snapshot_has_type(snapshot: &serde_json::Value, entitlement_type: &str) -> bool { - snapshot.as_array().is_some_and(|items| { - items - .iter() - .any(|item| item.get("type").and_then(|value| value.as_str()) == Some(entitlement_type)) - }) -} - async fn replace_matching_plan_entitlements_mysql( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, user_id: &str, snapshot: &serde_json::Value, now: i64, ) -> Result<(), DataLayerError> { - let replacement_types = plan_replacement_entitlement_types(snapshot); - if replacement_types.is_empty() { + let incoming_entitlements = plan_entitlements_snapshot(snapshot); + if !entitlements_have_replacement_selector(&incoming_entitlements) { return Ok(()); } @@ -3877,9 +5577,8 @@ WHERE user_id = ? "user_plan_entitlements.entitlements_snapshot", )? .unwrap_or_else(|| serde_json::json!([])); - let should_replace = replacement_types - .iter() - .any(|kind| entitlement_snapshot_has_type(&entitlements, kind)); + let should_replace = + entitlements_should_replace_existing(&incoming_entitlements, &entitlements); if !should_replace { continue; } @@ -3907,22 +5606,18 @@ WHERE id = ? Ok(()) } -fn plan_expires_at_unix(snapshot: &serde_json::Value, starts_at_unix_secs: i64) -> i64 { - let duration_value = snapshot - .get("duration_value") - .and_then(|value| value.as_i64()) - .unwrap_or(1) - .max(1); - let days = match snapshot - .get("duration_unit") - .and_then(|value| value.as_str()) - .unwrap_or("month") - { - "day" | "custom" => duration_value, - "year" => 365 * duration_value, - _ => 30 * duration_value, - }; - starts_at_unix_secs.saturating_add(days.saturating_mul(86_400)) +fn plan_expires_at_unix( + snapshot: &serde_json::Value, + starts_at_unix_secs: i64, +) -> Result { + let days = + checked_plan_duration_days_from_snapshot(snapshot).map_err(DataLayerError::InvalidInput)?; + let seconds = days.checked_mul(86_400).ok_or_else(|| { + DataLayerError::InvalidInput("plan duration exceeds the supported range".to_string()) + })?; + starts_at_unix_secs.checked_add(seconds).ok_or_else(|| { + DataLayerError::InvalidInput("plan expiration exceeds the supported range".to_string()) + }) } async fn apply_plan_wallet_credit_mysql( @@ -3933,11 +5628,16 @@ async fn apply_plan_wallet_credit_mysql( entitlements: &serde_json::Value, now: i64, ) -> Result<(), DataLayerError> { + validate_plan_wallet_credit_entitlements(entitlements).map_err(DataLayerError::InvalidInput)?; let credits = entitlements .as_array() .into_iter() .flatten() - .filter(|item| item.get("type").and_then(|value| value.as_str()) == Some("wallet_credit")) + .filter(|item| { + item.get("type") + .and_then(|value| value.as_str()) + .is_some_and(|value| value.eq_ignore_ascii_case("wallet_credit")) + }) .filter_map(|item| { let amount = item.get("amount_usd").and_then(|value| value.as_f64())?; if amount <= 0.0 || !amount.is_finite() { @@ -3947,6 +5647,7 @@ async fn apply_plan_wallet_credit_mysql( .get("balance_bucket") .and_then(|value| value.as_str()) .unwrap_or("gift") + .trim() .to_ascii_lowercase(); Some((amount, bucket)) }) @@ -3955,7 +5656,7 @@ async fn apply_plan_wallet_credit_mysql( return Ok(()); } let Some(wallet_row) = sqlx::query( - "SELECT id, status, balance, gift_balance FROM wallets WHERE id = ? LIMIT 1 FOR UPDATE", + "SELECT id, status, balance, gift_balance, total_recharged FROM wallets WHERE id = ? LIMIT 1 FOR UPDATE", ) .bind(wallet_id) .fetch_optional(&mut **tx) @@ -3974,6 +5675,17 @@ async fn apply_plan_wallet_credit_mysql( } let mut recharge_balance: f64 = get(&wallet_row, "balance")?; let mut gift_balance: f64 = get(&wallet_row, "gift_balance")?; + let mut total_recharged: f64 = get(&wallet_row, "total_recharged")?; + if !recharge_balance.is_finite() + || !gift_balance.is_finite() + || gift_balance < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid for plan wallet_credit".to_string(), + )); + } for (amount, bucket) in credits { let before_recharge = recharge_balance; let before_gift = gift_balance; @@ -3981,23 +5693,34 @@ async fn apply_plan_wallet_credit_mysql( let credits_recharge = bucket == "recharge"; if credits_recharge { recharge_balance += amount; + total_recharged += amount; } else { gift_balance += amount; } let after_total = recharge_balance + gift_balance; + if !before_total.is_finite() + || !recharge_balance.is_finite() + || !gift_balance.is_finite() + || !total_recharged.is_finite() + || !after_total.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow for plan wallet_credit".to_string(), + )); + } sqlx::query( r#" UPDATE wallets SET balance = ?, gift_balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(recharge_balance) .bind(gift_balance) - .bind(if credits_recharge { amount } else { 0.0 }) + .bind(total_recharged) .bind(now) .bind(wallet_id) .execute(&mut **tx) @@ -4032,16 +5755,6 @@ VALUES (?, ?, 'recharge', 'plan_wallet_credit', ?, ?, ?, ?, ?, ?, ?, 'payment_or Ok(()) } -fn default_refund_mode_for_payment_method(payment_method: &str) -> &'static str { - if matches!( - payment_method, - "admin_manual" | "card_recharge" | "card_code" | "gift_code" - ) { - return "offline_payout"; - } - "original_channel" -} - fn payment_gateway_response_map( value: Option, ) -> serde_json::Map { @@ -4242,16 +5955,69 @@ async fn mysql_payment_order_by_order_no_for_update( async fn mysql_payment_order_by_gateway_order_id_for_update( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, + payment_method: &str, gateway_order_id: &str, ) -> Result, DataLayerError> { - let sql = payment_order_select_sql("WHERE gateway_order_id = ? LIMIT 1 FOR UPDATE"); + let sql = payment_order_select_sql( + "WHERE payment_method = ? AND gateway_order_id = ? LIMIT 1 FOR UPDATE", + ); sqlx::query(&sql) + .bind(payment_method) .bind(gateway_order_id) .fetch_optional(&mut **tx) .await .map_sql_err() } +async fn mysql_bind_payment_gateway_order_id( + tx: &mut sqlx::Transaction<'_, sqlx::MySql>, + order_id: &str, + gateway_order_id: &str, +) -> Result { + sqlx::query("SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + let bind_result = sqlx::query("UPDATE payment_orders SET gateway_order_id = ? WHERE id = ?") + .bind(gateway_order_id) + .bind(order_id) + .execute(&mut **tx) + .await; + match bind_result { + Ok(_) => { + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + Ok(true) + } + Err(error) + if error + .as_database_error() + .is_some_and(|database_error| database_error.is_unique_violation()) => + { + sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + Ok(false) + } + Err(error) => { + let _ = sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + let _ = sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + Err(DataLayerError::sql(error)) + } + } +} + fn refund_select_sql(where_clause: &str) -> String { format!( r#" @@ -4385,7 +6151,7 @@ async fn update_mysql_payment_callback_failure( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, callback_id: &str, input: &ProcessPaymentCallbackInput, - payload: &str, + _payload: &str, error: &str, ) -> Result<(), DataLayerError> { sqlx::query( @@ -4395,20 +6161,15 @@ SET signature_valid = ?, status = 'failed', error_message = ?, payload_hash = ?, - payload = ?, - processed_at = ?, - order_no = COALESCE(?, order_no), - gateway_order_id = COALESCE(?, gateway_order_id) + payload = NULL, + processed_at = ? WHERE id = ? "#, ) .bind(input.signature_valid) .bind(error) .bind(&input.payload_hash) - .bind(payload) .bind(current_unix_secs_i64()) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) .bind(callback_id) .execute(&mut **tx) .await @@ -4420,7 +6181,7 @@ async fn mark_mysql_payment_callback_processed( tx: &mut sqlx::Transaction<'_, sqlx::MySql>, callback_id: &str, input: &ProcessPaymentCallbackInput, - payload: &str, + _payload: &str, order_id: &str, order_no: &str, ) -> Result<(), DataLayerError> { @@ -4432,7 +6193,7 @@ SET payment_order_id = ?, status = 'processed', error_message = NULL, payload_hash = ?, - payload = ?, + payload = NULL, processed_at = ?, order_no = ?, gateway_order_id = COALESCE(?, gateway_order_id) @@ -4441,7 +6202,6 @@ WHERE id = ? ) .bind(order_id) .bind(&input.payload_hash) - .bind(payload) .bind(current_unix_secs_i64()) .bind(order_no) .bind(input.gateway_order_id.as_deref()) @@ -4521,7 +6281,20 @@ async fn initialize_mysql_auth_wallet( api_key_id: Option<&str>, initial_gift_usd: f64, unlimited: bool, -) -> Result, DataLayerError> { +) -> Result, DataLayerError> { + let owner = user_id + .or(api_key_id) + .filter(|value| !value.trim().is_empty()); + if owner.is_none() || (user_id.is_some() && api_key_id.is_some()) { + return Err(DataLayerError::InvalidInput( + "wallet owner must be exactly one non-empty user or API-key id".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "initial gift amount must be finite".to_string(), + )); + } let gift_amount = if unlimited { 0.0 } else { @@ -4543,8 +6316,80 @@ async fn initialize_mysql_auth_wallet( gift_amount, now, )?; + // Lock the owner row (when present) and keep the insert in the same + // transaction. This makes retries return the existing wallet instead of + // issuing another initial gift transaction. let mut tx = pool.begin().await.map_sql_err()?; - sqlx::query( + let owner_column = if user_id.is_some() { + "user_id" + } else { + "api_key_id" + }; + let owner_value = owner.expect("validated wallet owner"); + + // Keep the ownership lock order identical to the guarded user-deletion + // path: users first, then api_keys/wallets. This closes the window where + // a deletion can observe no wallet while a concurrent initializer inserts + // one after the user row has been removed. + if let Some(user_id) = user_id { + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + } else { + let api_key_user_id: Option = + sqlx::query_scalar("SELECT user_id FROM api_keys WHERE id = ?") + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(api_key_user_id) = api_key_user_id else { + tx.rollback().await.map_sql_err()?; + return Ok(None); + }; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? FOR UPDATE") + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + let api_key_exists: Option = + sqlx::query_scalar("SELECT id FROM api_keys WHERE id = ? AND user_id = ? FOR UPDATE") + .bind(owner_value) + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if api_key_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + } + let existing_sql = format!( + "SELECT id, user_id, api_key_id, balance, gift_balance, limit_mode, currency, status, total_recharged, total_consumed, total_refunded, total_adjusted, updated_at AS updated_at_unix_secs FROM wallets WHERE {owner_column} = ? LIMIT 1 FOR UPDATE" + ); + if let Some(row) = sqlx::query(&existing_sql) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + { + let existing = map_wallet_row(&row)?; + tx.commit().await.map_sql_err()?; + return Ok(Some((existing, false))); + } + + let insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, api_key_id, balance, gift_balance, limit_mode, currency, @@ -4552,6 +6397,7 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, ?, 0, ?, ?, 'USD', 'active', 0, 0, 0, ?, ?, ?) +ON DUPLICATE KEY UPDATE id = id "#, ) .bind(&wallet.id) @@ -4565,6 +6411,28 @@ VALUES (?, ?, ?, 0, ?, ?, 'USD', 'active', 0, 0, 0, ?, ?, ?) .execute(&mut *tx) .await .map_sql_err()?; + // Re-read under lock. A duplicate owner insert means another caller won; + // return that row and never write a second gift transaction. + let existing_sql = format!( + "SELECT id, user_id, api_key_id, balance, gift_balance, limit_mode, currency, status, total_recharged, total_consumed, total_refunded, total_adjusted, updated_at AS updated_at_unix_secs FROM wallets WHERE {owner_column} = ? LIMIT 1 FOR UPDATE" + ); + let Some(owner_row) = sqlx::query(&existing_sql) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + let owner_wallet_id: String = get(&owner_row, "id")?; + if insert_result.rows_affected() == 0 || owner_wallet_id != wallet.id { + let existing = map_wallet_row(&owner_row)?; + tx.commit().await.map_sql_err()?; + return Ok(Some((existing, false))); + } if gift_amount > 0.0 { let link_id = user_id.or(api_key_id).unwrap_or_default(); let description = if api_key_id.is_some() { @@ -4594,8 +6462,33 @@ VALUES (?, ?, 'gift', 'gift_initial', ?, 0, ?, 0, 0, 0, ?, 'system_task', ?, NUL .await .map_sql_err()?; } + let wallet = map_wallet_row(&owner_row)?; tx.commit().await.map_sql_err()?; - Ok(Some(wallet)) + Ok(Some((wallet, true))) +} + +fn mysql_wallet_recharge_replay_matches( + row: &MySqlRow, + wallet_id: &str, + input: &CreateWalletRechargeOrderInput, +) -> Result { + let existing_wallet_id: String = get(row, "wallet_id")?; + let pay_currency: Option = get(row, "pay_currency")?; + let payment_method: String = get(row, "payment_method")?; + let payment_provider: Option = get(row, "payment_provider")?; + let payment_channel: Option = get(row, "payment_channel")?; + Ok(wallet_recharge_replay_matches( + &existing_wallet_id, + get(row, "amount_usd")?, + get(row, "pay_amount")?, + pay_currency.as_deref(), + get(row, "exchange_rate")?, + &payment_method, + payment_provider.as_deref(), + payment_channel.as_deref(), + wallet_id, + input, + )) } fn map_payment_order_row(row: &MySqlRow) -> Result { @@ -4611,6 +6504,8 @@ fn map_payment_order_row(row: &MySqlRow) -> Result String { sql.split_whitespace().collect::>().join(" ") } +#[test] +fn mysql_gateway_order_uniqueness_preflights_before_persistent_changes() { + const UNIQUENESS_MIGRATION: &str = include_str!( + "../../migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql" + ); + + let executable_migration = UNIQUENESS_MIGRATION + .lines() + .filter(|line| !line.trim_start().starts_with("--")) + .collect::>() + .join("\n"); + let migration = compact_sql(&executable_migration); + let initial_cleanup = migration + .find("DROP TEMPORARY TABLE IF EXISTS aether_payment_gateway_order_uniqueness_preflight") + .expect("migration should clean up a same-session failed preflight"); + let create_preflight = migration + .find("CREATE TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight") + .expect("migration should create a non-persistent conflict guard"); + let seed_preflight = migration + .find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) VALUES (1)") + .expect("migration should seed the duplicate-key conflict guard"); + let conflict_probe = migration + .find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) SELECT 1 FROM payment_orders") + .expect("migration should reject normalized historical conflicts"); + let final_cleanup = migration + .find("DROP TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight;") + .expect("migration should remove the successful preflight guard"); + let first_persistent_update = migration + .find("UPDATE payment_orders SET payment_method") + .expect("migration should normalize payment order methods"); + let alter = migration + .find( + "MODIFY COLUMN gateway_order_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NULL", + ) + .expect("migration should enforce binary gateway identifiers"); + let unique = migration + .find("ADD UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id") + .expect("migration should create composite uniqueness"); + + assert!(migration.contains("GROUP BY LOWER(TRIM(payment_method)), CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin HAVING COUNT(*) > 1 LIMIT 1")); + assert!(migration.contains("WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method))")); + let preflight = &migration[..final_cleanup]; + assert!(!preflight.contains("UPDATE payment_orders")); + assert!(!preflight.contains("UPDATE payment_callbacks")); + assert!(!preflight.contains("ALTER TABLE")); + assert!( + initial_cleanup < create_preflight + && create_preflight < seed_preflight + && seed_preflight < conflict_probe + && conflict_probe < final_cleanup + && final_cleanup < first_persistent_update + && first_persistent_update < alter, + "the conflict probe must finish before any persistent UPDATE or ALTER" + ); + assert!( + alter < unique && !migration.contains("CREATE UNIQUE INDEX"), + "the collation and unique index must be one atomic ALTER TABLE" + ); +} + #[tokio::test] async fn mysql_wallet_read_repository_reads_wallet_contract_views() { let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL") diff --git a/crates/aether-data/adapters/postgres/migrations/20260403000000_baseline.sql b/crates/aether-data/adapters/postgres/migrations/20260403000000_baseline.sql index 7a8cc3bff..04b0d926e 100644 --- a/crates/aether-data/adapters/postgres/migrations/20260403000000_baseline.sql +++ b/crates/aether-data/adapters/postgres/migrations/20260403000000_baseline.sql @@ -706,7 +706,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( is_manual boolean DEFAULT false NOT NULL, proxy_url character varying(500), proxy_username character varying(255), - proxy_password character varying(500), + proxy_password text, created_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, updated_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, remote_config json, diff --git a/crates/aether-data/adapters/postgres/migrations/20260528010000_repair_missing_routing_profiles.sql b/crates/aether-data/adapters/postgres/migrations/20260528010000_repair_missing_routing_profiles.sql index ded28ee4e..d6b62bb2b 100644 --- a/crates/aether-data/adapters/postgres/migrations/20260528010000_repair_missing_routing_profiles.sql +++ b/crates/aether-data/adapters/postgres/migrations/20260528010000_repair_missing_routing_profiles.sql @@ -79,4 +79,4 @@ BEGIN END IF; END $$; CREATE INDEX IF NOT EXISTS routing_group_versions_group_id_idx - ON public.routing_group_versions USING btree (group_id); \ No newline at end of file + ON public.routing_group_versions USING btree (group_id); diff --git a/crates/aether-data/adapters/postgres/migrations/20260814000000_add_usage_cost_reservations.sql b/crates/aether-data/adapters/postgres/migrations/20260814000000_add_usage_cost_reservations.sql new file mode 100644 index 000000000..467b580fe --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260814000000_add_usage_cost_reservations.sql @@ -0,0 +1,43 @@ +CREATE TABLE IF NOT EXISTS public.usage_cost_reservations ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + reservation_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + reserved_cost_units bigint NOT NULL, + actual_cost_units bigint, + state character varying(20) NOT NULL, + reservation_expires_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + finalized_at timestamp with time zone, + created_at timestamp with time zone DEFAULT now() NOT NULL, + updated_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT usage_cost_reservations_pkey PRIMARY KEY (reservation_token), + CONSTRAINT usage_cost_reservations_state_check + CHECK (state IN ('reserved', 'finalized', 'released')), + CONSTRAINT usage_cost_reservations_reserved_cost_units_check + CHECK (reserved_cost_units >= 0), + CONSTRAINT usage_cost_reservations_actual_cost_units_check + CHECK (actual_cost_units IS NULL OR actual_cost_units >= 0), + CONSTRAINT usage_cost_reservations_expiry_check + CHECK (reservation_expires_at > admitted_at), + CONSTRAINT usage_cost_reservations_retention_check + CHECK (retain_until >= reservation_expires_at), + CONSTRAINT usage_cost_reservations_lifecycle_check CHECK ( + (state = 'reserved' AND actual_cost_units IS NULL AND finalized_at IS NULL) + OR (state = 'finalized' AND actual_cost_units IS NOT NULL AND finalized_at IS NOT NULL) + OR (state = 'released' AND actual_cost_units IS NOT NULL + AND actual_cost_units = 0 AND finalized_at IS NOT NULL) + ) +); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_request_id_idx + ON public.usage_cost_reservations USING btree (request_id); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_subject_admitted_at_idx + ON public.usage_cost_reservations USING btree (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_reservation_expires_at_idx + ON public.usage_cost_reservations USING btree (reservation_expires_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_retain_until_token_idx + ON public.usage_cost_reservations USING btree (retain_until, reservation_token); diff --git a/crates/aether-data/adapters/postgres/migrations/20260815000000_add_usage_request_admissions.sql b/crates/aether-data/adapters/postgres/migrations/20260815000000_add_usage_request_admissions.sql new file mode 100644 index 000000000..b64972e19 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260815000000_add_usage_request_admissions.sql @@ -0,0 +1,25 @@ +CREATE TABLE IF NOT EXISTS public.usage_request_admissions ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + event_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + state character varying(20) NOT NULL, + released_at timestamp with time zone, + created_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT usage_request_admissions_pkey PRIMARY KEY (event_token), + CONSTRAINT usage_request_admissions_retention_check + CHECK (retain_until > admitted_at), + CONSTRAINT usage_request_admissions_state_check + CHECK (state IN ('active', 'released')), + CONSTRAINT usage_request_admissions_lifecycle_check CHECK ( + (state = 'active' AND released_at IS NULL) + OR (state = 'released' AND released_at IS NOT NULL AND released_at >= admitted_at) + ) +); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_subject_admitted_at_idx + ON public.usage_request_admissions USING btree (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_retain_until_token_idx + ON public.usage_request_admissions USING btree (retain_until, event_token); diff --git a/crates/aether-data/adapters/postgres/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql b/crates/aether-data/adapters/postgres/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql new file mode 100644 index 000000000..537f5bcc2 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql @@ -0,0 +1,25 @@ +-- A feature branch shipped the two ledger migrations before their user +-- ownership relationship was enforced. Preserve those migration checksums and +-- add the cascade in a follow-up that also upgrades databases which ran that +-- branch. Rows whose user was already deleted cannot be retained safely. +DELETE FROM public.usage_cost_reservations AS reservation +WHERE NOT EXISTS ( + SELECT 1 + FROM public.users AS app_user + WHERE app_user.id = reservation.subject_id +); + +DELETE FROM public.usage_request_admissions AS admission +WHERE NOT EXISTS ( + SELECT 1 + FROM public.users AS app_user + WHERE app_user.id = admission.subject_id +); + +ALTER TABLE public.usage_cost_reservations + ADD CONSTRAINT usage_cost_reservations_subject_id_fkey + FOREIGN KEY (subject_id) REFERENCES public.users(id) ON DELETE CASCADE; + +ALTER TABLE public.usage_request_admissions + ADD CONSTRAINT usage_request_admissions_subject_id_fkey + FOREIGN KEY (subject_id) REFERENCES public.users(id) ON DELETE CASCADE; diff --git a/crates/aether-data/adapters/postgres/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql b/crates/aether-data/adapters/postgres/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql new file mode 100644 index 000000000..baae3af46 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql @@ -0,0 +1,19 @@ +-- A gateway transaction identifier may repeat across payment methods, but +-- must never identify two orders in the same method. If historical conflicts +-- exist, index creation intentionally fails without modifying financial data. +-- Diagnose with: +-- SELECT payment_method, gateway_order_id, COUNT(*) +-- FROM public.payment_orders +-- WHERE gateway_order_id IS NOT NULL +-- GROUP BY payment_method, gateway_order_id +-- HAVING COUNT(*) > 1; +UPDATE public.payment_orders +SET payment_method = lower(btrim(payment_method)) +WHERE payment_method <> lower(btrim(payment_method)); + +UPDATE public.payment_callbacks +SET payment_method = lower(btrim(payment_method)) +WHERE payment_method <> lower(btrim(payment_method)); + +CREATE UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id + ON public.payment_orders (payment_method, gateway_order_id); diff --git a/crates/aether-data/adapters/postgres/migrations/20260821130000_add_user_security_version.sql b/crates/aether-data/adapters/postgres/migrations/20260821130000_add_user_security_version.sql new file mode 100644 index 000000000..627e2dc70 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260821130000_add_user_security_version.sql @@ -0,0 +1,5 @@ +ALTER TABLE public.users + ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0; + +ALTER TABLE public.user_sessions + ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0; diff --git a/crates/aether-data/adapters/postgres/migrations/20260827040000_expand_proxy_password_ciphertext.sql b/crates/aether-data/adapters/postgres/migrations/20260827040000_expand_proxy_password_ciphertext.sql new file mode 100644 index 000000000..62762b1ca --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260827040000_expand_proxy_password_ciphertext.sql @@ -0,0 +1,3 @@ +-- Purpose-bound Fernet envelopes can exceed the former 500-character plaintext limit. +ALTER TABLE proxy_nodes + ALTER COLUMN proxy_password TYPE TEXT; diff --git a/crates/aether-data/adapters/postgres/migrations/20260827050000_anonymize_deleted_user_history.sql b/crates/aether-data/adapters/postgres/migrations/20260827050000_anonymize_deleted_user_history.sql new file mode 100644 index 000000000..a62944394 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260827050000_anonymize_deleted_user_history.sql @@ -0,0 +1,76 @@ +ALTER TABLE public.request_candidates + DROP CONSTRAINT IF EXISTS request_candidates_user_id_fkey; +ALTER TABLE public.video_tasks + DROP CONSTRAINT IF EXISTS video_tasks_user_id_fkey; +ALTER TABLE public.usage + DROP CONSTRAINT IF EXISTS usage_user_id_fkey; +ALTER TABLE public.stats_user_daily + DROP CONSTRAINT IF EXISTS stats_user_daily_user_id_fkey; +ALTER TABLE public.stats_user_summary + DROP CONSTRAINT IF EXISTS stats_user_summary_user_id_fkey; +ALTER TABLE public.stats_user_daily_model + DROP CONSTRAINT IF EXISTS stats_user_daily_model_user_id_fkey; +ALTER TABLE public.stats_user_daily_provider + DROP CONSTRAINT IF EXISTS stats_user_daily_provider_user_id_fkey; +ALTER TABLE public.stats_user_daily_api_format + DROP CONSTRAINT IF EXISTS stats_user_daily_api_format_user_id_fkey; +ALTER TABLE public.stats_user_daily_model_provider + DROP CONSTRAINT IF EXISTS stats_user_daily_model_provider_user_id_fkey; +ALTER TABLE public.stats_user_daily_cost_savings + DROP CONSTRAINT IF EXISTS stats_user_daily_cost_savings_user_id_fkey; +ALTER TABLE public.stats_user_daily_cost_savings_provider + DROP CONSTRAINT IF EXISTS stats_user_daily_cost_savings_provider_user_id_fkey; +ALTER TABLE public.stats_user_daily_cost_savings_model + DROP CONSTRAINT IF EXISTS stats_user_daily_cost_savings_model_user_id_fkey; +ALTER TABLE public.stats_user_daily_cost_savings_model_provider + DROP CONSTRAINT IF EXISTS stats_user_daily_cost_savings_model_provider_user_id_fkey; +ALTER TABLE public.stats_hourly_user_model + DROP CONSTRAINT IF EXISTS stats_hourly_user_model_user_id_fkey; +ALTER TABLE public.user_model_usage_counts + DROP CONSTRAINT IF EXISTS user_model_usage_counts_user_id_fkey; + +ALTER TABLE public.audit_logs + DROP CONSTRAINT IF EXISTS audit_logs_user_id_fkey; +ALTER TABLE public.announcements + DROP CONSTRAINT IF EXISTS announcements_author_id_fkey; +ALTER TABLE public.payment_orders + DROP CONSTRAINT IF EXISTS payment_orders_user_id_fkey; +ALTER TABLE public.proxy_nodes + DROP CONSTRAINT IF EXISTS proxy_nodes_registered_by_fkey; +ALTER TABLE public.refund_requests + DROP CONSTRAINT IF EXISTS refund_requests_user_id_fkey, + DROP CONSTRAINT IF EXISTS refund_requests_requested_by_fkey, + DROP CONSTRAINT IF EXISTS refund_requests_approved_by_fkey, + DROP CONSTRAINT IF EXISTS refund_requests_processed_by_fkey; +ALTER TABLE public.wallet_transactions + DROP CONSTRAINT IF EXISTS wallet_transactions_operator_id_fkey; +ALTER TABLE public.wallets + DROP CONSTRAINT IF EXISTS wallets_user_id_fkey, + DROP CONSTRAINT IF EXISTS wallets_api_key_id_fkey; +ALTER TABLE public.redeem_code_batches + DROP CONSTRAINT IF EXISTS redeem_code_batches_created_by_fkey; +ALTER TABLE public.redeem_codes + DROP CONSTRAINT IF EXISTS redeem_codes_redeemed_by_user_id_fkey, + DROP CONSTRAINT IF EXISTS redeem_codes_disabled_by_fkey; + +ALTER TABLE public.user_plan_entitlements + DROP CONSTRAINT IF EXISTS user_plan_entitlements_user_id_fkey; +ALTER TABLE public.entitlement_usage_ledgers + DROP CONSTRAINT IF EXISTS entitlement_usage_ledgers_user_id_fkey; +ALTER TABLE public.user_referrals + DROP CONSTRAINT IF EXISTS user_referrals_inviter_user_id_fkey, + DROP CONSTRAINT IF EXISTS user_referrals_invitee_user_id_fkey; +ALTER TABLE public.referral_rewards + DROP CONSTRAINT IF EXISTS referral_rewards_inviter_user_id_fkey, + DROP CONSTRAINT IF EXISTS referral_rewards_invitee_user_id_fkey; + +ALTER TABLE public.request_candidates + DROP CONSTRAINT IF EXISTS request_candidates_api_key_id_fkey; +ALTER TABLE public.video_tasks + DROP CONSTRAINT IF EXISTS video_tasks_api_key_id_fkey; +ALTER TABLE public.usage + DROP CONSTRAINT IF EXISTS usage_api_key_id_fkey; +ALTER TABLE public.stats_daily_api_key + DROP CONSTRAINT IF EXISTS stats_daily_api_key_api_key_id_fkey; + +-- Legacy rows are intentionally retained; current writes enforce the policy. diff --git a/crates/aether-data/adapters/postgres/migrations/20260831000000_enforce_ldap_config_singleton.sql b/crates/aether-data/adapters/postgres/migrations/20260831000000_enforce_ldap_config_singleton.sql new file mode 100644 index 000000000..9a34a8d7b --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260831000000_enforce_ldap_config_singleton.sql @@ -0,0 +1,38 @@ +-- LDAP configuration is a database-wide singleton. Preserve the row selected by the legacy +-- reader (the smallest id), remove historical duplicates, and let the database arbitrate +-- concurrent first creation. +DELETE FROM public.ldap_configs +WHERE id <> (SELECT MIN(id) FROM public.ldap_configs); + +ALTER TABLE public.ldap_configs + ADD COLUMN IF NOT EXISTS singleton_key INTEGER NOT NULL DEFAULT 1; + +UPDATE public.ldap_configs +SET singleton_key = 1 +WHERE singleton_key IS DISTINCT FROM 1; + +ALTER TABLE public.ldap_configs + ALTER COLUMN singleton_key SET DEFAULT 1, + ALTER COLUMN singleton_key SET NOT NULL; + +DO $migration$ +BEGIN + ALTER TABLE public.ldap_configs + ADD CONSTRAINT ldap_configs_singleton_key_check CHECK (singleton_key = 1); +EXCEPTION + WHEN duplicate_object THEN NULL; +END +$migration$; + +DO $migration$ +BEGIN + ALTER TABLE public.ldap_configs + ADD CONSTRAINT ldap_configs_singleton_key_key UNIQUE (singleton_key); +EXCEPTION + WHEN duplicate_object THEN NULL; + -- The empty-database snapshot already materializes this constraint. PostgreSQL + -- reports the existing constraint's backing relation as duplicate_table (42P07) + -- rather than duplicate_object, so treat that known idempotent case the same way. + WHEN duplicate_table THEN NULL; +END +$migration$; diff --git a/crates/aether-data/adapters/postgres/migrations/20260831010000_add_proxy_node_tunnel_generation.sql b/crates/aether-data/adapters/postgres/migrations/20260831010000_add_proxy_node_tunnel_generation.sql new file mode 100644 index 000000000..5c51c20c2 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260831010000_add_proxy_node_tunnel_generation.sql @@ -0,0 +1,20 @@ +ALTER TABLE public.proxy_nodes + ADD COLUMN IF NOT EXISTS tunnel_generation character varying(64); + +-- The generation is an opaque epoch marker used to reject stale mutations; it +-- is not a credential. Keep this migration independent of the optional +-- pgcrypto extension. The registration path uses a CSPRNG for new rows, +-- while this backfill combines the immutable row identity, physical tuple, +-- transaction time, and PostgreSQL's per-call PRNG to produce a distinct +-- marker for every legacy row using only core functions. +UPDATE public.proxy_nodes +SET tunnel_generation = md5( + id || ':' || + ctid::text || ':' || + clock_timestamp()::text || ':' || + random()::text +) +WHERE tunnel_generation IS NULL OR btrim(tunnel_generation) = ''; + +ALTER TABLE public.proxy_nodes + ALTER COLUMN tunnel_generation SET NOT NULL; diff --git a/crates/aether-data/adapters/postgres/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql b/crates/aether-data/adapters/postgres/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql new file mode 100644 index 000000000..3943a8a9a --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql @@ -0,0 +1,2 @@ +ALTER TABLE public.usage_counter_deltas + ADD COLUMN IF NOT EXISTS target_tunnel_generation character varying(64); diff --git a/crates/aether-data/adapters/postgres/migrations/20260901000000_expand_gemini_file_mapping_metadata.sql b/crates/aether-data/adapters/postgres/migrations/20260901000000_expand_gemini_file_mapping_metadata.sql new file mode 100644 index 000000000..5fa41a094 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260901000000_expand_gemini_file_mapping_metadata.sql @@ -0,0 +1,6 @@ +-- Match upgraded PostgreSQL installations to the portable Gemini Files +-- metadata contract already used by the logical schema and MySQL. +ALTER TABLE public.gemini_file_mappings + ALTER COLUMN file_name TYPE character varying(512), + ALTER COLUMN display_name TYPE character varying(512), + ALTER COLUMN mime_type TYPE character varying(255); diff --git a/crates/aether-data/adapters/postgres/migrations/20260903000000_add_routing_group_sort_order.sql b/crates/aether-data/adapters/postgres/migrations/20260903000000_add_routing_group_sort_order.sql new file mode 100644 index 000000000..b8a391020 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260903000000_add_routing_group_sort_order.sql @@ -0,0 +1,8 @@ +ALTER TABLE public.routing_groups + -- The empty-database snapshot may already contain the current logical + -- schema. Keep the incremental migration safe for both snapshot and + -- legacy databases. + ADD COLUMN IF NOT EXISTS sort_order bigint NOT NULL DEFAULT 0; + +CREATE INDEX IF NOT EXISTS routing_groups_enabled_sort_idx + ON public.routing_groups (enabled DESC, sort_order, name, id); diff --git a/crates/aether-data/adapters/postgres/src/auth.rs b/crates/aether-data/adapters/postgres/src/auth.rs index b36da3259..573dcdeea 100644 --- a/crates/aether-data/adapters/postgres/src/auth.rs +++ b/crates/aether-data/adapters/postgres/src/auth.rs @@ -4,9 +4,9 @@ use sqlx::{postgres::PgRow, PgPool, Row}; use aether_data_contracts::repository::auth::{ AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository, - AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, - StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, - UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, + AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, + CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, + StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; use aether_data_contracts::DataLayerError; @@ -197,6 +197,37 @@ WHERE api_keys.id = ANY($1::TEXT[]) ORDER BY api_keys.id ASC "#; +const RESTORE_API_KEY_SELECT_SQL: &str = r#" +SELECT + api_keys.user_id, + api_keys.id AS api_key_id, + api_keys.key_hash, + api_keys.key_encrypted, + api_keys.name, + api_keys.allowed_providers, + api_keys.allowed_api_formats, + api_keys.allowed_models, + api_keys.ip_rules, + api_keys.rate_limit, + api_keys.concurrent_limit, + api_keys.force_capabilities, + api_keys.feature_settings, + api_keys.is_active, + CAST(EXTRACT(EPOCH FROM api_keys.expires_at) AS BIGINT) AS expires_at_unix_secs, + api_keys.auto_delete_on_expiry, + api_keys.total_requests, + COALESCE(api_keys.total_tokens, 0)::BIGINT AS total_tokens, + COALESCE(CAST(api_keys.total_cost_usd AS DOUBLE PRECISION), 0) AS total_cost_usd, + CAST(EXTRACT(EPOCH FROM api_keys.last_used_at) AS BIGINT) AS last_used_at_unix_secs, + CAST(EXTRACT(EPOCH FROM api_keys.created_at) AS BIGINT) AS created_at_unix_secs, + CAST(EXTRACT(EPOCH FROM api_keys.updated_at) AS BIGINT) AS updated_at_unix_secs, + api_keys.is_standalone +FROM api_keys +WHERE api_keys.id = $1 +LIMIT 1 +FOR UPDATE +"#; + const LIST_EXPORT_BY_NAME_SEARCH_SQL: &str = r#" SELECT api_keys.user_id, @@ -407,15 +438,15 @@ VALUES ( $10, $11, $12, - NULL, $13, $14, $15, - FALSE, - FALSE, $16, + FALSE, + FALSE, $17, $18, + $19, NOW(), NOW() ) @@ -525,14 +556,17 @@ RETURNING const UPDATE_USER_API_KEY_BASIC_SQL: &str = r#" UPDATE api_keys SET - name = COALESCE($3, name), - rate_limit = COALESCE($4, rate_limit), - concurrent_limit = COALESCE($5, concurrent_limit), - ip_rules = CASE WHEN $6 THEN $7::jsonb ELSE ip_rules END, + key_encrypted = CASE WHEN $3 THEN $4 ELSE key_encrypted END, + name = CASE WHEN $5 THEN $6 ELSE name END, + rate_limit = CASE WHEN $7 THEN $8 ELSE rate_limit END, + concurrent_limit = CASE WHEN $9 THEN $10 ELSE concurrent_limit END, + ip_rules = CASE WHEN $11 THEN $12::jsonb ELSE ip_rules END, + feature_settings = CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END, updated_at = NOW() WHERE user_id = $1 AND id = $2 AND is_standalone = FALSE + AND ($15 = FALSE OR is_locked = FALSE) RETURNING user_id, id AS api_key_id, @@ -562,15 +596,17 @@ RETURNING const UPDATE_STANDALONE_API_KEY_BASIC_SQL: &str = r#" UPDATE api_keys SET - name = COALESCE($2, name), - rate_limit = CASE WHEN $3 THEN $4 ELSE rate_limit END, - concurrent_limit = CASE WHEN $5 THEN $6 ELSE concurrent_limit END, - allowed_providers = CASE WHEN $7 THEN $8::json ELSE allowed_providers END, - allowed_api_formats = CASE WHEN $9 THEN $10::json ELSE allowed_api_formats END, - allowed_models = CASE WHEN $11 THEN $12::json ELSE allowed_models END, - ip_rules = CASE WHEN $13 THEN $14::jsonb ELSE ip_rules END, - expires_at = CASE WHEN $15 THEN $16::timestamptz ELSE expires_at END, - auto_delete_on_expiry = CASE WHEN $17 THEN $18 ELSE auto_delete_on_expiry END, + key_encrypted = CASE WHEN $2 THEN $3 ELSE key_encrypted END, + name = CASE WHEN $4 THEN $5 ELSE name END, + force_capabilities = CASE WHEN $6 THEN $7::json ELSE force_capabilities END, + rate_limit = CASE WHEN $8 THEN $9 ELSE rate_limit END, + concurrent_limit = CASE WHEN $10 THEN $11 ELSE concurrent_limit END, + allowed_providers = CASE WHEN $12 THEN $13::json ELSE allowed_providers END, + allowed_api_formats = CASE WHEN $14 THEN $15::json ELSE allowed_api_formats END, + allowed_models = CASE WHEN $16 THEN $17::json ELSE allowed_models END, + ip_rules = CASE WHEN $18 THEN $19::jsonb ELSE ip_rules END, + expires_at = CASE WHEN $20 THEN $21::timestamptz ELSE expires_at END, + auto_delete_on_expiry = CASE WHEN $22 THEN $23 ELSE auto_delete_on_expiry END, updated_at = NOW() WHERE id = $1 AND is_standalone = TRUE @@ -608,6 +644,7 @@ SET WHERE user_id = $1 AND id = $2 AND is_standalone = FALSE + AND ($4 = FALSE OR is_locked = FALSE) RETURNING user_id, id AS api_key_id, @@ -719,6 +756,7 @@ SET WHERE user_id = $1 AND id = $2 AND is_standalone = FALSE + AND ($4 = FALSE OR is_locked = FALSE) RETURNING user_id, id AS api_key_id, @@ -753,6 +791,7 @@ SET WHERE user_id = $1 AND id = $2 AND is_standalone = FALSE + AND ($4 = FALSE OR is_locked = FALSE) RETURNING user_id, id AS api_key_id, @@ -762,6 +801,7 @@ RETURNING allowed_providers, allowed_api_formats, allowed_models, + ip_rules, rate_limit, concurrent_limit, force_capabilities, @@ -786,6 +826,7 @@ SET WHERE user_id = $1 AND id = $2 AND is_standalone = FALSE + AND ($4 = FALSE OR is_locked = FALSE) "#; const SET_STANDALONE_API_KEY_FEATURE_SETTINGS_SQL: &str = r#" @@ -797,26 +838,20 @@ WHERE id = $1 AND is_standalone = TRUE "#; -const DISABLE_WALLET_BY_API_KEY_ID_SQL: &str = r#" -UPDATE wallets -SET status = 'disabled', - updated_at = NOW() -WHERE api_key_id = $1 - AND status <> 'disabled' -"#; +const POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET api_key_name = NULL WHERE api_key_id = $1", + "UPDATE video_tasks SET api_key_name = NULL WHERE api_key_id = $1", + "UPDATE usage SET api_key_name = NULL WHERE api_key_id = $1", + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id = $1", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id = $1", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = $1)", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id = $1 AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = $1)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = $1)", +]; -const DELETE_USER_API_KEY_SQL: &str = r#" -DELETE FROM api_keys -WHERE user_id = $1 - AND id = $2 - AND is_standalone = FALSE -"#; - -const DELETE_STANDALONE_API_KEY_SQL: &str = r#" -DELETE FROM api_keys -WHERE id = $1 - AND is_standalone = TRUE -"#; +const POSTGRES_DELETE_API_KEY_DEPENDENTS_SQL: &[&str] = + &["DELETE FROM api_key_provider_mappings WHERE api_key_id = $1"]; #[derive(Debug, Clone)] pub struct SqlxAuthApiKeySnapshotReadRepository { @@ -952,10 +987,12 @@ impl SqlxAuthApiKeySnapshotReadRepository { if user_ids.is_empty() { return Ok(AuthApiKeyExportSummary::default()); } + let now_unix_secs = + datetime_from_unix_secs(now_unix_secs, "api_keys.summary_now")?.timestamp() as f64; let row = sqlx::query(SUMMARIZE_EXPORT_BY_USER_IDS_SQL) .bind(user_ids) - .bind(now_unix_secs as f64) + .bind(now_unix_secs) .fetch_one(&self.pool) .await .map_postgres_err()?; @@ -969,8 +1006,10 @@ impl SqlxAuthApiKeySnapshotReadRepository { &self, now_unix_secs: u64, ) -> Result { + let now_unix_secs = + datetime_from_unix_secs(now_unix_secs, "api_keys.summary_now")?.timestamp() as f64; let row = sqlx::query(SUMMARIZE_EXPORT_NON_STANDALONE_SQL) - .bind(now_unix_secs as f64) + .bind(now_unix_secs) .fetch_one(&self.pool) .await .map_postgres_err()?; @@ -1025,8 +1064,10 @@ impl SqlxAuthApiKeySnapshotReadRepository { &self, now_unix_secs: u64, ) -> Result { + let now_unix_secs = + datetime_from_unix_secs(now_unix_secs, "api_keys.summary_now")?.timestamp() as f64; let row = sqlx::query(SUMMARIZE_EXPORT_STANDALONE_SQL) - .bind(now_unix_secs as f64) + .bind(now_unix_secs) .fetch_one(&self.pool) .await .map_postgres_err()?; @@ -1173,12 +1214,20 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; let expires_at = record .expires_at_unix_secs - .map(|value| { - chrono::DateTime::::from_timestamp(value as i64, 0).ok_or_else(|| { - DataLayerError::UnexpectedValue(format!("invalid api_keys.expires_at: {value}")) - }) - }) + .map(|value| datetime_from_unix_secs(value, "api_keys.expires_at")) .transpose()?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + let owner_exists: Option = sqlx::query_scalar( + "SELECT id FROM users WHERE id = $1 AND is_deleted IS FALSE FOR UPDATE", + ) + .bind(&record.user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if owner_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } let row = sqlx::query(CREATE_USER_API_KEY_SQL) .bind(record.api_key_id) .bind(record.user_id) @@ -1192,16 +1241,22 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(record.rate_limit) .bind(record.concurrent_limit) .bind(record.force_capabilities) + .bind(record.feature_settings) .bind(record.is_active) .bind(expires_at) .bind(record.auto_delete_on_expiry) - .bind(record.total_requests as i64) - .bind(record.total_tokens as i64) + .bind(i64_from_u64( + record.total_requests, + "api_keys.total_requests", + )?) + .bind(i64_from_u64(record.total_tokens, "api_keys.total_tokens")?) .bind(record.total_cost_usd) - .fetch_optional(&self.pool) + .fetch_optional(&mut *tx) .await .map_postgres_err()?; - row.as_ref().map(map_auth_api_key_export_row).transpose() + let record = row.as_ref().map(map_auth_api_key_export_row).transpose()?; + tx.commit().await.map_err(postgres_error)?; + Ok(record) } async fn create_standalone_api_key( @@ -1230,12 +1285,20 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; let expires_at = record .expires_at_unix_secs - .map(|value| { - chrono::DateTime::::from_timestamp(value as i64, 0).ok_or_else(|| { - DataLayerError::UnexpectedValue(format!("invalid api_keys.expires_at: {value}")) - }) - }) + .map(|value| datetime_from_unix_secs(value, "api_keys.expires_at")) .transpose()?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + let owner_exists: Option = sqlx::query_scalar( + "SELECT id FROM users WHERE id = $1 AND is_deleted IS FALSE FOR UPDATE", + ) + .bind(&record.user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if owner_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } let row = sqlx::query(CREATE_STANDALONE_API_KEY_SQL) .bind(record.api_key_id) .bind(record.user_id) @@ -1252,13 +1315,18 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(record.is_active) .bind(expires_at) .bind(record.auto_delete_on_expiry) - .bind(record.total_requests as i64) - .bind(record.total_tokens as i64) + .bind(i64_from_u64( + record.total_requests, + "api_keys.total_requests", + )?) + .bind(i64_from_u64(record.total_tokens, "api_keys.total_tokens")?) .bind(record.total_cost_usd) - .fetch_optional(&self.pool) + .fetch_optional(&mut *tx) .await .map_postgres_err()?; - row.as_ref().map(map_auth_api_key_export_row).transpose() + let record = row.as_ref().map(map_auth_api_key_export_row).transpose()?; + tx.commit().await.map_err(postgres_error)?; + Ok(record) } async fn update_user_api_key_basic( @@ -1272,14 +1340,84 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .map(serde_json::to_value) .transpose() .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let feature_settings = record.feature_settings.clone().flatten(); let row = sqlx::query(UPDATE_USER_API_KEY_BASIC_SQL) .bind(record.user_id) .bind(record.api_key_id) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted) + .bind(record.name_present) .bind(record.name) + .bind(record.rate_limit_present) .bind(record.rate_limit) + .bind(record.concurrent_limit_present) .bind(record.concurrent_limit) .bind(record.ip_rules.is_some()) .bind(ip_rules) + .bind(record.feature_settings.is_some()) + .bind(feature_settings) + .bind(false) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_auth_api_key_export_row).transpose() + } + + async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = $1 +WHERE id = $2 + AND user_id = $3 + AND key_hash = $4 + AND is_standalone = $5 + AND key_encrypted = $6 +"#, + ) + .bind(&mutation.key_encrypted) + .bind(&mutation.api_key_id) + .bind(&mutation.user_id) + .bind(&mutation.key_hash) + .bind(mutation.is_standalone) + .bind(&mutation.expected_key_encrypted) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) + } + + async fn update_user_api_key_basic_if_unlocked( + &self, + record: UpdateUserApiKeyBasicRecord, + ) -> Result, DataLayerError> { + let ip_rules = record + .ip_rules + .clone() + .flatten() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let feature_settings = record.feature_settings.clone().flatten(); + let row = sqlx::query(UPDATE_USER_API_KEY_BASIC_SQL) + .bind(record.user_id) + .bind(record.api_key_id) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted) + .bind(record.name_present) + .bind(record.name) + .bind(record.rate_limit_present) + .bind(record.rate_limit) + .bind(record.concurrent_limit_present) + .bind(record.concurrent_limit) + .bind(record.ip_rules.is_some()) + .bind(ip_rules) + .bind(record.feature_settings.is_some()) + .bind(feature_settings) + .bind(true) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -1320,15 +1458,16 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; let expires_at = record .expires_at_unix_secs - .map(|value| { - chrono::DateTime::::from_timestamp(value as i64, 0).ok_or_else(|| { - DataLayerError::UnexpectedValue(format!("invalid api_keys.expires_at: {value}")) - }) - }) + .map(|value| datetime_from_unix_secs(value, "api_keys.expires_at")) .transpose()?; let row = sqlx::query(UPDATE_STANDALONE_API_KEY_BASIC_SQL) .bind(record.api_key_id) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted) + .bind(record.name_present) .bind(record.name) + .bind(record.force_capabilities.is_some()) + .bind(record.force_capabilities.clone().flatten()) .bind(record.rate_limit_present) .bind(record.rate_limit) .bind(record.concurrent_limit_present) @@ -1351,6 +1490,133 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { row.as_ref().map(map_auth_api_key_export_row).transpose() } + async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + if restored.api_key_id != expected.api_key_id + || restored.user_id != expected.user_id + || restored.key_hash != expected.key_hash + || restored.is_standalone != expected.is_standalone + { + return Ok(false); + } + + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query(RESTORE_API_KEY_SELECT_SQL) + .bind(&expected.api_key_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let current = map_auth_api_key_export_row(&row)?; + if current != *expected { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + + let expires_at = restored + .expires_at_unix_secs + .map(|value| datetime_from_unix_secs(value, "api_keys.expires_at")) + .transpose()?; + let last_used_at = restored + .last_used_at_unix_secs + .map(|value| datetime_from_unix_secs(value, "api_keys.last_used_at")) + .transpose()?; + let allowed_providers = restored + .allowed_providers + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let allowed_api_formats = restored + .allowed_api_formats + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let allowed_models = restored + .allowed_models + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let ip_rules = restored + .ip_rules + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = $1, + name = $2, + allowed_providers = $3::json, + allowed_api_formats = $4::json, + allowed_models = $5::json, + ip_rules = $6::jsonb, + rate_limit = $7, + concurrent_limit = $8, + force_capabilities = $9::json, + feature_settings = $10::jsonb, + is_active = $11, + expires_at = $12, + auto_delete_on_expiry = $13, + total_requests = $14, + total_tokens = $15, + total_cost_usd = $16, + last_used_at = $17, + updated_at = NOW() +WHERE id = $18 + AND user_id = $19 + AND key_hash = $20 + AND is_standalone = $21 +"#, + ) + .bind(&restored.key_encrypted) + .bind(&restored.name) + .bind(allowed_providers) + .bind(allowed_api_formats) + .bind(allowed_models) + .bind(ip_rules) + .bind(restored.rate_limit) + .bind(restored.concurrent_limit) + .bind(&restored.force_capabilities) + .bind(&restored.feature_settings) + .bind(restored.is_active) + .bind(expires_at) + .bind(restored.auto_delete_on_expiry) + .bind(i64_from_u64( + restored.total_requests, + "api_keys.total_requests", + )?) + .bind(i64_from_u64( + restored.total_tokens, + "api_keys.total_tokens", + )?) + .bind(restored.total_cost_usd) + .bind(last_used_at) + .bind(&restored.api_key_id) + .bind(&restored.user_id) + .bind(&restored.key_hash) + .bind(restored.is_standalone) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + async fn set_user_api_key_active( &self, user_id: &str, @@ -1361,6 +1627,24 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(user_id) .bind(api_key_id) .bind(is_active) + .bind(false) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_auth_api_key_export_row).transpose() + } + + async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + let row = sqlx::query(SET_USER_API_KEY_ACTIVE_SQL) + .bind(user_id) + .bind(api_key_id) + .bind(is_active) + .bind(true) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -1411,6 +1695,28 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(user_id) .bind(api_key_id) .bind(allowed_providers) + .bind(false) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_auth_api_key_export_row).transpose() + } + + async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, DataLayerError> { + let allowed_providers = allowed_providers + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let row = sqlx::query(SET_USER_API_KEY_ALLOWED_PROVIDERS_SQL) + .bind(user_id) + .bind(api_key_id) + .bind(allowed_providers) + .bind(true) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -1427,6 +1733,24 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(user_id) .bind(api_key_id) .bind(force_capabilities) + .bind(false) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_auth_api_key_export_row).transpose() + } + + async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, DataLayerError> { + let row = sqlx::query(SET_USER_API_KEY_FORCE_CAPABILITIES_SQL) + .bind(user_id) + .bind(api_key_id) + .bind(force_capabilities) + .bind(true) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -1443,6 +1767,32 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(user_id) .bind(api_key_id) .bind(feature_settings) + .bind(false) + .execute(&self.pool) + .await + .map_postgres_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + let api_key_ids = [api_key_id.to_string()]; + Ok(self + .list_export_api_keys_by_ids(&api_key_ids) + .await? + .into_iter() + .find(|record| record.user_id == user_id && !record.is_standalone)) + } + + async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, DataLayerError> { + let result = sqlx::query(SET_USER_API_KEY_FEATURE_SETTINGS_SQL) + .bind(user_id) + .bind(api_key_id) + .bind(feature_settings) + .bind(true) .execute(&self.pool) .await .map_postgres_err()?; @@ -1464,10 +1814,15 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { total_tokens: u64, total_cost_usd: f64, ) -> Result, DataLayerError> { + if !total_cost_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "api_keys.total_cost_usd is not finite".to_string(), + )); + } let row = sqlx::query(SET_API_KEY_USAGE_TOTALS_SQL) .bind(api_key_id) - .bind(total_requests as i64) - .bind(total_tokens as i64) + .bind(i64_from_u64(total_requests, "api_keys.total_requests")?) + .bind(i64_from_u64(total_tokens, "api_keys.total_tokens")?) .bind(total_cost_usd) .fetch_optional(&self.pool) .await @@ -1480,20 +1835,17 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { user_id: &str, api_key_id: &str, ) -> Result { - let mut tx = self.pool.begin().await.map_postgres_err()?; - sqlx::query(DISABLE_WALLET_BY_API_KEY_ID_SQL) - .bind(api_key_id) - .execute(&mut *tx) + self.delete_api_key(api_key_id, Some(user_id), false, false) .await - .map_postgres_err()?; - let result = sqlx::query(DELETE_USER_API_KEY_SQL) - .bind(user_id) - .bind(api_key_id) - .execute(&mut *tx) + } + + async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + self.delete_api_key(api_key_id, Some(user_id), false, true) .await - .map_postgres_err()?; - tx.commit().await.map_err(postgres_error)?; - Ok(result.rows_affected() > 0) } async fn set_standalone_api_key_feature_settings( @@ -1519,22 +1871,106 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { } async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result { - let mut tx = self.pool.begin().await.map_postgres_err()?; - sqlx::query(DISABLE_WALLET_BY_API_KEY_ID_SQL) - .bind(api_key_id) - .execute(&mut *tx) - .await - .map_postgres_err()?; - let result = sqlx::query(DELETE_STANDALONE_API_KEY_SQL) - .bind(api_key_id) - .execute(&mut *tx) - .await - .map_postgres_err()?; - tx.commit().await.map_err(postgres_error)?; - Ok(result.rows_affected() > 0) + self.delete_api_key(api_key_id, None, true, false).await } } +impl SqlxAuthApiKeySnapshotReadRepository { + async fn delete_api_key( + &self, + api_key_id: &str, + user_id: Option<&str>, + is_standalone: bool, + require_unlocked: bool, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let matching_api_key = if let Some(user_id) = user_id { + if require_unlocked { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = $1 AND user_id = $2 AND is_standalone IS FALSE AND is_locked IS FALSE FOR UPDATE", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = $1 AND user_id = $2 AND is_standalone IS FALSE FOR UPDATE", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + } + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = $1 AND is_standalone IS TRUE FOR UPDATE", + ) + .bind(api_key_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + }; + if matching_api_key.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + + sqlx::query( + "UPDATE wallets SET status = 'disabled', updated_at = NOW() WHERE api_key_id = $1 AND status <> 'disabled'", + ) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + for sql in POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + for sql in POSTGRES_DELETE_API_KEY_DEPENDENTS_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + let result = sqlx::query("DELETE FROM api_keys WHERE id = $1 AND is_standalone = $2") + .bind(api_key_id) + .bind(is_standalone) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_err(postgres_error)?; + Ok(true) + } +} + +fn i64_from_u64(value: u64, field_name: &str) -> Result { + i64::try_from(value) + .map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}"))) +} + +fn datetime_from_unix_secs( + value: u64, + field_name: &str, +) -> Result, DataLayerError> { + let unix_secs = i64_from_u64(value, field_name)?; + chrono::DateTime::::from_timestamp(unix_secs, 0).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "{field_name} is outside the supported timestamp range: {value}" + )) + }) +} + fn row_get(row: &sqlx::postgres::PgRow, column: &str) -> Result where for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type, @@ -1610,51 +2046,107 @@ fn map_auth_api_key_export_row( #[cfg(test)] mod tests { use super::{ + datetime_from_unix_secs, i64_from_u64, DataLayerError, SqlxAuthApiKeySnapshotReadRepository, CREATE_STANDALONE_API_KEY_SQL, - CREATE_USER_API_KEY_SQL, UPDATE_STANDALONE_API_KEY_BASIC_SQL, + CREATE_USER_API_KEY_SQL, POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL, + POSTGRES_DELETE_API_KEY_DEPENDENTS_SQL, UPDATE_STANDALONE_API_KEY_BASIC_SQL, UPDATE_USER_API_KEY_BASIC_SQL, }; use crate::{PostgresPoolConfig, PostgresPoolFactory}; + #[test] + fn checked_u64_conversions_reject_counter_and_timestamp_overflow() { + assert_eq!( + i64_from_u64(i64::MAX as u64, "api_keys.total_requests") + .expect("i64 maximum should fit"), + i64::MAX + ); + assert!(matches!( + i64_from_u64(u64::MAX, "api_keys.total_requests"), + Err(DataLayerError::InvalidInput(_)) + )); + assert!(matches!( + datetime_from_unix_secs(i64::MAX as u64, "api_keys.expires_at"), + Err(DataLayerError::InvalidInput(_)) + )); + } + #[test] fn create_api_key_sql_orders_expiry_before_standalone_flags() { assert!(CREATE_USER_API_KEY_SQL .contains("expires_at,\n auto_delete_on_expiry,\n is_locked,\n is_standalone,")); - assert!( - CREATE_USER_API_KEY_SQL.contains("$13,\n $14,\n $15,\n FALSE,\n FALSE,\n $16,") - ); + assert!(CREATE_USER_API_KEY_SQL + .contains("$13,\n $14,\n $15,\n $16,\n FALSE,\n FALSE,\n $17,")); assert!(CREATE_STANDALONE_API_KEY_SQL .contains("expires_at,\n auto_delete_on_expiry,\n is_locked,\n is_standalone,")); assert!(CREATE_STANDALONE_API_KEY_SQL .contains("$13,\n $14,\n $15,\n FALSE,\n TRUE,\n $16,")); } + #[test] + fn api_key_delete_sql_preserves_ids_and_removes_private_snapshots() { + for table in [ + "request_candidates", + "video_tasks", + "usage", + "stats_daily_api_key", + ] { + assert!(POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| { + sql.starts_with(&format!("UPDATE {table} ")) + && sql.contains("SET api_key_name = NULL") + && sql.ends_with("WHERE api_key_id = $1") + })); + } + assert!(POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL + .iter() + .any(|sql| sql + .starts_with("UPDATE audit_logs SET description = 'deleted API key event'"))); + assert!(POSTGRES_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| sql + .starts_with("UPDATE payment_callbacks SET payload = NULL, error_message = NULL"))); + assert_eq!( + POSTGRES_DELETE_API_KEY_DEPENDENTS_SQL, + &["DELETE FROM api_key_provider_mappings WHERE api_key_id = $1"] + ); + } + #[test] fn update_standalone_api_key_basic_sql_casts_json_case_values() { assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("concurrent_limit = CASE WHEN $5 THEN $6 ELSE concurrent_limit END")); + .contains("key_encrypted = CASE WHEN $2 THEN $3 ELSE key_encrypted END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("allowed_providers = CASE WHEN $7 THEN $8::json ELSE allowed_providers END")); + .contains("concurrent_limit = CASE WHEN $10 THEN $11 ELSE concurrent_limit END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains( - "allowed_api_formats = CASE WHEN $9 THEN $10::json ELSE allowed_api_formats END" + "allowed_providers = CASE WHEN $12 THEN $13::json ELSE allowed_providers END" + )); + assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains( + "allowed_api_formats = CASE WHEN $14 THEN $15::json ELSE allowed_api_formats END" )); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("allowed_models = CASE WHEN $11 THEN $12::json ELSE allowed_models END")); + .contains("allowed_models = CASE WHEN $16 THEN $17::json ELSE allowed_models END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("ip_rules = CASE WHEN $13 THEN $14::jsonb ELSE ip_rules END")); + .contains("ip_rules = CASE WHEN $18 THEN $19::jsonb ELSE ip_rules END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("rate_limit = CASE WHEN $3 THEN $4 ELSE rate_limit END")); + .contains("rate_limit = CASE WHEN $8 THEN $9 ELSE rate_limit END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL - .contains("expires_at = CASE WHEN $15 THEN $16::timestamptz ELSE expires_at END")); + .contains("expires_at = CASE WHEN $20 THEN $21::timestamptz ELSE expires_at END")); assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains( - "auto_delete_on_expiry = CASE WHEN $17 THEN $18 ELSE auto_delete_on_expiry END" + "auto_delete_on_expiry = CASE WHEN $22 THEN $23 ELSE auto_delete_on_expiry END" + )); + assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains( + "force_capabilities = CASE WHEN $6 THEN $7::json ELSE force_capabilities END" )); } #[test] - fn update_user_api_key_basic_sql_casts_ip_rules_as_jsonb() { + fn update_user_api_key_basic_sql_casts_json_patches_and_fences_locked_keys() { assert!(UPDATE_USER_API_KEY_BASIC_SQL - .contains("ip_rules = CASE WHEN $6 THEN $7::jsonb ELSE ip_rules END")); + .contains("key_encrypted = CASE WHEN $3 THEN $4 ELSE key_encrypted END")); + assert!(UPDATE_USER_API_KEY_BASIC_SQL + .contains("ip_rules = CASE WHEN $11 THEN $12::jsonb ELSE ip_rules END")); + assert!(UPDATE_USER_API_KEY_BASIC_SQL.contains( + "feature_settings = CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END" + )); + assert!(UPDATE_USER_API_KEY_BASIC_SQL.contains("AND ($15 = FALSE OR is_locked = FALSE)")); } #[tokio::test] diff --git a/crates/aether-data/adapters/postgres/src/auth_modules.rs b/crates/aether-data/adapters/postgres/src/auth_modules.rs index 04e8fc176..80974eb3f 100644 --- a/crates/aether-data/adapters/postgres/src/auth_modules.rs +++ b/crates/aether-data/adapters/postgres/src/auth_modules.rs @@ -3,7 +3,7 @@ use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row}; use aether_data_contracts::repository::auth_modules::*; use aether_data_contracts::DataLayerError; -use aether_data_query::{push_eq, push_limit, WhereClause}; +use aether_data_query::{push_eq, WhereClause}; use crate::error::SqlxResultExt; @@ -34,7 +34,37 @@ SELECT FROM ldap_configs "#; -const UPDATE_LDAP_CONFIG_SQL: &str = r#" +const UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL: &str = r#" +UPDATE ldap_configs +SET + server_url = $1, + bind_dn = $2, + base_dn = $3, + user_search_filter = $4, + username_attr = $5, + email_attr = $6, + display_name_attr = $7, + is_enabled = $8, + is_exclusive = $9, + use_starttls = $10, + connect_timeout = $11, + updated_at = NOW() +WHERE singleton_key = 1 + AND server_url IS NOT DISTINCT FROM $12 + AND bind_dn IS NOT DISTINCT FROM $13 + AND bind_password_encrypted IS NOT DISTINCT FROM $14 + AND base_dn IS NOT DISTINCT FROM $15 + AND user_search_filter IS NOT DISTINCT FROM $16 + AND username_attr IS NOT DISTINCT FROM $17 + AND email_attr IS NOT DISTINCT FROM $18 + AND display_name_attr IS NOT DISTINCT FROM $19 + AND is_enabled IS NOT DISTINCT FROM $20 + AND is_exclusive IS NOT DISTINCT FROM $21 + AND use_starttls IS NOT DISTINCT FROM $22 + AND connect_timeout IS NOT DISTINCT FROM $23 +"#; + +const UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL: &str = r#" UPDATE ldap_configs SET server_url = $1, @@ -50,29 +80,24 @@ SET use_starttls = $11, connect_timeout = $12, updated_at = NOW() -WHERE id = ( - SELECT id - FROM ldap_configs - ORDER BY id ASC - LIMIT 1 -) -RETURNING - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter, - username_attr, - email_attr, - display_name_attr, - is_enabled, - is_exclusive, - use_starttls, - connect_timeout +WHERE singleton_key = 1 + AND server_url IS NOT DISTINCT FROM $13 + AND bind_dn IS NOT DISTINCT FROM $14 + AND bind_password_encrypted IS NOT DISTINCT FROM $15 + AND base_dn IS NOT DISTINCT FROM $16 + AND user_search_filter IS NOT DISTINCT FROM $17 + AND username_attr IS NOT DISTINCT FROM $18 + AND email_attr IS NOT DISTINCT FROM $19 + AND display_name_attr IS NOT DISTINCT FROM $20 + AND is_enabled IS NOT DISTINCT FROM $21 + AND is_exclusive IS NOT DISTINCT FROM $22 + AND use_starttls IS NOT DISTINCT FROM $23 + AND connect_timeout IS NOT DISTINCT FROM $24 "#; const INSERT_LDAP_CONFIG_SQL: &str = r#" INSERT INTO ldap_configs ( + singleton_key, server_url, bind_dn, bind_password_encrypted, @@ -89,6 +114,7 @@ INSERT INTO ldap_configs ( updated_at ) VALUES ( + 1, $1, $2, $3, @@ -104,19 +130,6 @@ VALUES ( NOW(), NOW() ) -RETURNING - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter, - username_attr, - email_attr, - display_name_attr, - is_enabled, - is_exclusive, - use_starttls, - connect_timeout "#; #[derive(Debug, Clone)] @@ -154,8 +167,7 @@ async fn list_enabled_oauth_providers( async fn get_ldap_config(pool: &PgPool) -> Result, DataLayerError> { let mut builder = QueryBuilder::::new(LDAP_CONFIG_COLUMNS); - builder.push(" ORDER BY id ASC"); - push_limit(&mut builder, 1); + builder.push(" WHERE singleton_key = 1"); let row = builder .build() .fetch_optional(pool) @@ -192,48 +204,203 @@ impl AuthModuleReadRepository for SqlxAuthModuleRepository { #[async_trait] impl AuthModuleWriteRepository for SqlxAuthModuleRepository { - async fn upsert_ldap_config( + async fn compare_and_swap_ldap_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, DataLayerError> { - let updated = sqlx::query(UPDATE_LDAP_CONFIG_SQL) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .fetch_optional(&self.pool) - .await - .map_postgres_err()?; - if let Some(row) = updated.as_ref() { - return map_ldap_row(row).map(Some); - } + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result { + let persisted = + ldap_config_after_password_update(expected, replacement, bind_password_update)?; - let inserted = sqlx::query(INSERT_LDAP_CONFIG_SQL) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .fetch_optional(&self.pool) - .await - .map_postgres_err()?; - inserted.as_ref().map(map_ldap_row).transpose() + let Some(expected) = expected else { + let insert = sqlx::query(INSERT_LDAP_CONFIG_SQL) + .bind(&persisted.server_url) + .bind(&persisted.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&persisted.base_dn) + .bind(persisted.user_search_filter.as_deref()) + .bind(persisted.username_attr.as_deref()) + .bind(persisted.email_attr.as_deref()) + .bind(persisted.display_name_attr.as_deref()) + .bind(persisted.is_enabled) + .bind(persisted.is_exclusive) + .bind(persisted.use_starttls) + .bind(persisted.connect_timeout) + .execute(&self.pool) + .await; + return match insert { + Ok(result) if result.rows_affected() == 1 => { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } + Ok(_) => Ok(CompareAndSwapLdapConfigResult::Conflict), + Err(error) + if error + .as_database_error() + .is_some_and(|error| error.is_unique_violation()) => + { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } + Err(error) => Err(crate::error::postgres_error(error)), + }; + }; + + let rows_affected = match bind_password_update { + LdapBindPasswordUpdate::Preserve => { + sqlx::query(UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected() + } + LdapBindPasswordUpdate::Set(_) | LdapBindPasswordUpdate::Clear => { + sqlx::query(UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected() + } + }; + if rows_affected == 1 { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } else { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } } + + async fn delete_ldap_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result { + let rows_affected = sqlx::query( + r#" +DELETE FROM ldap_configs +WHERE singleton_key = 1 + AND server_url IS NOT DISTINCT FROM $1 + AND bind_dn IS NOT DISTINCT FROM $2 + AND bind_password_encrypted IS NOT DISTINCT FROM $3 + AND base_dn IS NOT DISTINCT FROM $4 + AND user_search_filter IS NOT DISTINCT FROM $5 + AND username_attr IS NOT DISTINCT FROM $6 + AND email_attr IS NOT DISTINCT FROM $7 + AND display_name_attr IS NOT DISTINCT FROM $8 + AND is_enabled IS NOT DISTINCT FROM $9 + AND is_exclusive IS NOT DISTINCT FROM $10 + AND use_starttls IS NOT DISTINCT FROM $11 + AND connect_timeout IS NOT DISTINCT FROM $12 +"#, + ) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + + async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + let rows_affected = sqlx::query( + r#" +UPDATE ldap_configs +SET bind_password_encrypted = $1, updated_at = NOW() +WHERE singleton_key = 1 + AND bind_password_encrypted = $2 +"#, + ) + .bind(replacement) + .bind(expected) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } +} + +fn ldap_config_after_password_update( + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, +) -> Result { + let bind_password_encrypted = match bind_password_update { + LdapBindPasswordUpdate::Preserve => expected + .ok_or_else(|| { + DataLayerError::InvalidConfiguration( + "LDAP bind password cannot be preserved while creating the singleton" + .to_string(), + ) + })? + .bind_password_encrypted + .clone(), + LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()), + LdapBindPasswordUpdate::Clear => None, + }; + Ok(StoredLdapModuleConfig { + bind_password_encrypted, + ..replacement.clone() + }) } fn map_oauth_row(row: &PgRow) -> Result { diff --git a/crates/aether-data/adapters/postgres/src/background_tasks.rs b/crates/aether-data/adapters/postgres/src/background_tasks.rs index 0ef8ead47..e7e1fb8d9 100644 --- a/crates/aether-data/adapters/postgres/src/background_tasks.rs +++ b/crates/aether-data/adapters/postgres/src/background_tasks.rs @@ -225,8 +225,9 @@ impl BackgroundTaskReadRepository for SqlxBackgroundTaskRepository { impl BackgroundTaskWriteRepository for SqlxBackgroundTaskRepository { async fn upsert_run( &self, - run: UpsertBackgroundTaskRun, + mut run: UpsertBackgroundTaskRun, ) -> Result { + run.sanitize_for_persistence(); run.validate()?; sqlx::query( r#" @@ -327,8 +328,9 @@ ON CONFLICT(id) DO UPDATE SET async fn upsert_event( &self, - event: UpsertBackgroundTaskEvent, + mut event: UpsertBackgroundTaskEvent, ) -> Result { + event.sanitize_for_persistence(); event.validate()?; sqlx::query( r#" @@ -383,7 +385,7 @@ fn map_run_row(row: &PgRow) -> Result { row.try_get("finished_at_unix_secs").map_postgres_err()?; let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_postgres_err()?; - Ok(StoredBackgroundTaskRun { + let mut run = StoredBackgroundTaskRun { id: row.try_get("id").map_postgres_err()?, task_key: row.try_get("task_key").map_postgres_err()?, kind: BackgroundTaskKind::from_database(&kind)?, @@ -403,19 +405,23 @@ fn map_run_row(row: &PgRow) -> Result { started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()), finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()), updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(), - }) + }; + run.sanitize_persisted_data(); + Ok(run) } fn map_event_row(row: &PgRow) -> Result { let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_postgres_err()?; - Ok(StoredBackgroundTaskEvent { + let mut event = StoredBackgroundTaskEvent { id: row.try_get("id").map_postgres_err()?, run_id: row.try_get("run_id").map_postgres_err()?, event_type: row.try_get("event_type").map_postgres_err()?, message: row.try_get("message").map_postgres_err()?, payload_json: row.try_get("payload_json").map_postgres_err()?, created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(), - }) + }; + event.sanitize_persisted_data(); + Ok(event) } fn i64_from_usize(value: usize, label: &str) -> Result { diff --git a/crates/aether-data/adapters/postgres/src/billing.rs b/crates/aether-data/adapters/postgres/src/billing.rs index d598fde12..8cc7455ef 100644 --- a/crates/aether-data/adapters/postgres/src/billing.rs +++ b/crates/aether-data/adapters/postgres/src/billing.rs @@ -4,8 +4,9 @@ use sqlx::{PgPool, Row}; use aether_data_contracts::repository::billing::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; use aether_data_contracts::DataLayerError; @@ -718,6 +719,120 @@ LIMIT 1 row.as_ref().map(map_payment_gateway_config_row).transpose() } + async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE payment_gateway_configs +SET merchant_key_encrypted = $3 +WHERE provider = $1 + AND merchant_key_encrypted = $2 + "#, + ) + .bind(update.provider.trim().to_ascii_lowercase()) + .bind(&update.expected_merchant_key_encrypted) + .bind(&update.merchant_key_encrypted) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_swap_payment_gateway_config( + &self, + mutation: &PaymentGatewayConfigCasWriteInput, + ) -> Result, DataLayerError> { + let input = &mutation.input; + let provider = input.provider.trim().to_ascii_lowercase(); + let row = if mutation.expected_existing { + sqlx::query( + r#" +UPDATE payment_gateway_configs +SET + enabled = $2, + endpoint_url = $3, + callback_base_url = $4, + merchant_id = $5, + merchant_key_encrypted = CASE + WHEN $11::BOOL THEN payment_gateway_configs.merchant_key_encrypted + ELSE $6 + END, + pay_currency = $7, + usd_exchange_rate = $8, + min_recharge_usd = $9, + channels_json = $10, + updated_at = NOW() +WHERE provider = $1 + AND merchant_key_encrypted IS NOT DISTINCT FROM $12 +RETURNING + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, + CAST(usd_exchange_rate AS DOUBLE PRECISION) AS usd_exchange_rate, + CAST(min_recharge_usd AS DOUBLE PRECISION) AS min_recharge_usd, + channels_json, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_secs, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs + "#, + ) + .bind(&provider) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(&input.channels_json) + .bind(input.preserve_existing_secret) + .bind(mutation.expected_merchant_key_encrypted.as_deref()) + .fetch_optional(&self.pool) + .await + .map_postgres_err()? + } else { + sqlx::query( + r#" +INSERT INTO payment_gateway_configs ( + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd, + channels_json, created_at, updated_at +) +VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, NOW(), NOW()) +ON CONFLICT (provider) DO NOTHING +RETURNING + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, + CAST(usd_exchange_rate AS DOUBLE PRECISION) AS usd_exchange_rate, + CAST(min_recharge_usd AS DOUBLE PRECISION) AS min_recharge_usd, + channels_json, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_secs, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs + "#, + ) + .bind(&provider) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(&input.channels_json) + .fetch_optional(&self.pool) + .await + .map_postgres_err()? + }; + match row.as_ref() { + Some(row) => Ok(AdminBillingMutationOutcome::Applied( + map_payment_gateway_config_row(row)?, + )), + None => Ok(AdminBillingMutationOutcome::NotFound), + } + } + async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -999,19 +1114,54 @@ ORDER BY expires_at ASC, created_at ASC )) } + async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE user_plan_entitlements +SET status = 'revoked', + expires_at = LEAST(expires_at, NOW()), + updated_at = NOW() +WHERE id = $1 + AND user_id = $2 + AND status = 'active' + AND expires_at > NOW() + "#, + ) + .bind(entitlement_id) + .bind(user_id) + .execute(&self.pool) + .await + .map_postgres_err()?; + if result.rows_affected() == 0 { + Ok(AdminBillingMutationOutcome::NotFound) + } else { + Ok(AdminBillingMutationOutcome::Applied(())) + } + } + async fn find_user_daily_quota_availability( &self, user_id: &str, ) -> Result, DataLayerError> { let rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = $1 - AND status = 'active' - AND starts_at <= NOW() - AND expires_at > NOW() -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = $1 + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= NOW() + AND user_plan_entitlements.expires_at > NOW() +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC "#, ) .bind(user_id) @@ -1024,9 +1174,12 @@ ORDER BY expires_at ASC, created_at ASC, id ASC let entitlement_id: String = row.try_get("id").map_postgres_err()?; let entitlements: serde_json::Value = row.try_get("entitlements_snapshot").map_postgres_err()?; + let plan_entitlements: serde_json::Value = + row.try_get("plan_entitlements_json").map_postgres_err()?; grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -1160,6 +1313,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -1185,15 +1339,27 @@ fn daily_quota_grants_from_entitlement( .and_then(serde_json::Value::as_str), now, )?, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + fn map_payment_gateway_config_row( row: &sqlx::postgres::PgRow, ) -> Result { diff --git a/crates/aether-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index 64a55e9ef..0b1650263 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -1255,11 +1255,11 @@ fn map_candidate_selection_row( key_name: row.try_get("key_name").map_postgres_err()?, key_auth_type: row.try_get("key_auth_type").map_postgres_err()?, key_is_active: row.try_get("key_is_active").map_postgres_err()?, - key_api_formats: parse_string_list( + key_api_formats: parse_key_policy_string_list( row.try_get("key_api_formats").map_postgres_err()?, "provider_api_keys.api_formats", )?, - key_allowed_models: parse_string_list( + key_allowed_models: parse_key_policy_string_list( row.try_get("key_allowed_models").map_postgres_err()?, "provider_api_keys.allowed_models", )?, @@ -1301,6 +1301,79 @@ fn parse_string_list( parse_string_list_value(&value, field_name) } +fn parse_key_policy_string_list( + value: Option, + field_name: &str, +) -> Result>, DataLayerError> { + let Some(value) = value else { + return Ok(None); + }; + parse_key_policy_string_list_value(&value, field_name) +} + +fn parse_key_policy_string_list_value( + value: &serde_json::Value, + field_name: &str, +) -> Result>, DataLayerError> { + match value { + serde_json::Value::Null => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains JSON null; use SQL NULL for an unset policy" + ))), + serde_json::Value::Array(array) => { + parse_key_policy_string_list_array(array, field_name).map(Some) + } + serde_json::Value::String(raw) => parse_embedded_key_policy_string_list(raw, field_name), + _ => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} is not a JSON array" + ))), + } +} + +fn parse_embedded_key_policy_string_list( + raw: &str, + field_name: &str, +) -> Result>, DataLayerError> { + let raw = raw.trim(); + if raw.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty string" + ))); + } + if raw.eq_ignore_ascii_case("null") { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains stringified JSON null; use SQL NULL for an unset policy" + ))); + } + + if let Ok(decoded) = serde_json::from_str::(raw) { + return parse_key_policy_string_list_value(&decoded, field_name); + } + + Ok(Some(vec![raw.to_string()])) +} + +fn parse_key_policy_string_list_array( + array: &[serde_json::Value], + field_name: &str, +) -> Result, DataLayerError> { + let mut items = Vec::with_capacity(array.len()); + for item in array { + let Some(item) = item.as_str() else { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains a non-string item" + ))); + }; + let item = item.trim(); + if item.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty item" + ))); + } + items.push(item.to_string()); + } + Ok(items) +} + fn parse_string_list_value( value: &serde_json::Value, field_name: &str, @@ -1505,9 +1578,9 @@ mod tests { use serde_json::json; use super::{ - parse_provider_model_mappings, parse_string_list, pool_key_candidate_selection_sql, - requested_model_selection_page_sql, requested_model_selection_sql, - SqlxMinimalCandidateSelectionReadRepository, + parse_key_policy_string_list, parse_provider_model_mappings, parse_string_list, + pool_key_candidate_selection_sql, requested_model_selection_page_sql, + requested_model_selection_sql, SqlxMinimalCandidateSelectionReadRepository, LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL, LIST_FOR_EXACT_API_FORMAT_SQL, LIST_POOL_KEYS_FOR_GROUP_SQL, PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_MARKER, PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_SQL, @@ -1747,6 +1820,29 @@ mod tests { assert_eq!(parsed, Some(vec!["gpt-5.2".to_string()])); } + #[test] + fn malformed_key_policy_never_degrades_to_unrestricted() { + for value in [ + json!(null), + json!("null"), + json!(""), + json!(["openai:chat", null]), + ] { + assert!( + parse_key_policy_string_list(Some(value), "provider_api_keys.api_formats",) + .is_err() + ); + } + assert_eq!( + parse_key_policy_string_list( + Some(json!(["openai:chat"])), + "provider_api_keys.api_formats", + ) + .expect("valid key policy should parse"), + Some(vec!["openai:chat".to_string()]) + ); + } + #[test] fn parse_provider_model_mappings_accepts_stringified_array() { let parsed = parse_provider_model_mappings(Some(json!( diff --git a/crates/aether-data/adapters/postgres/src/candidates.rs b/crates/aether-data/adapters/postgres/src/candidates.rs index 7f183a0f0..3604fe2ef 100644 --- a/crates/aether-data/adapters/postgres/src/candidates.rs +++ b/crates/aether-data/adapters/postgres/src/candidates.rs @@ -1,3 +1,5 @@ +use std::sync::LazyLock; + use async_trait::async_trait; use futures_util::{future::BoxFuture, stream::TryStream, TryStreamExt}; use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row}; @@ -6,7 +8,8 @@ use uuid::Uuid; use aether_data_contracts::repository::candidates::{ PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate, - UpsertRequestCandidateRecord, + UpsertRequestCandidateRecord, REQUEST_CANDIDATE_ERROR_TYPES, + REQUEST_CANDIDATE_ERROR_TYPE_ALIASES, REQUEST_CANDIDATE_SKIP_REASONS, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_in, push_limit, WhereClause}; @@ -64,7 +67,7 @@ GROUP BY FLOOR(EXTRACT(EPOCH FROM (created_at - TO_TIMESTAMP($2))) / $4)::BIGINT "#; -const UPSERT_SQL: &str = r#" +const UPSERT_SQL_TEMPLATE: &str = r#" INSERT INTO request_candidates ( id, request_id, @@ -126,16 +129,16 @@ VALUES ( ) ON CONFLICT (request_id, candidate_index, retry_index) DO UPDATE SET - user_id = COALESCE(EXCLUDED.user_id, request_candidates.user_id), - api_key_id = COALESCE(EXCLUDED.api_key_id, request_candidates.api_key_id), - username = COALESCE(EXCLUDED.username, request_candidates.username), - api_key_name = COALESCE(EXCLUDED.api_key_name, request_candidates.api_key_name), - provider_id = COALESCE(EXCLUDED.provider_id, request_candidates.provider_id), - endpoint_id = COALESCE(EXCLUDED.endpoint_id, request_candidates.endpoint_id), - key_id = COALESCE(EXCLUDED.key_id, request_candidates.key_id), + user_id = COALESCE(request_candidates.user_id, EXCLUDED.user_id), + api_key_id = COALESCE(request_candidates.api_key_id, EXCLUDED.api_key_id), + username = COALESCE(request_candidates.username, EXCLUDED.username), + api_key_name = COALESCE(request_candidates.api_key_name, EXCLUDED.api_key_name), + provider_id = COALESCE(request_candidates.provider_id, EXCLUDED.provider_id), + endpoint_id = COALESCE(request_candidates.endpoint_id, EXCLUDED.endpoint_id), + key_id = COALESCE(request_candidates.key_id, EXCLUDED.key_id), status = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status @@ -143,11 +146,11 @@ DO UPDATE SET THEN request_candidates.status ELSE EXCLUDED.status END, - skip_reason = COALESCE(EXCLUDED.skip_reason, request_candidates.skip_reason), + skip_reason = COALESCE(EXCLUDED.skip_reason, __AETHER_SANITIZED_LEGACY_SKIP_REASON__), is_cached = COALESCE($14, request_candidates.is_cached), status_code = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status_code WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status_code @@ -156,28 +159,23 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.status_code, request_candidates.status_code) END, error_type = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_type - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_type - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_type - ELSE COALESCE(EXCLUDED.error_type, request_candidates.error_type) - END, - error_message = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_message - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_message - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_message - ELSE COALESCE(EXCLUDED.error_message, request_candidates.error_message) + WHEN ( + request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') + AND EXCLUDED.status <> request_candidates.status + ) OR ( + request_candidates.status = 'pending' + AND EXCLUDED.status IN ('available', 'unused') + ) OR ( + request_candidates.status = 'streaming' + AND EXCLUDED.status IN ('available', 'unused', 'pending') + ) OR EXCLUDED.error_type IS NULL + THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__ + ELSE EXCLUDED.error_type END, + error_message = NULL, latency_ms = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.latency_ms WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.latency_ms @@ -186,41 +184,17 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms) END, concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests), - extra_data = CASE - WHEN request_candidates.extra_data IS NULL THEN EXCLUDED.extra_data - WHEN EXCLUDED.extra_data IS NULL THEN regexp_replace( - request_candidates.extra_data::text, - $aether_nul$(? request_candidates.status THEN request_candidates.finished_at WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.finished_at @@ -255,19 +229,19 @@ RETURNING CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms "#; -const UPSERT_CONFLICT_SQL: &str = r#" +const UPSERT_CONFLICT_SQL_TEMPLATE: &str = r#" ON CONFLICT (request_id, candidate_index, retry_index) DO UPDATE SET - user_id = COALESCE(EXCLUDED.user_id, request_candidates.user_id), - api_key_id = COALESCE(EXCLUDED.api_key_id, request_candidates.api_key_id), - username = COALESCE(EXCLUDED.username, request_candidates.username), - api_key_name = COALESCE(EXCLUDED.api_key_name, request_candidates.api_key_name), - provider_id = COALESCE(EXCLUDED.provider_id, request_candidates.provider_id), - endpoint_id = COALESCE(EXCLUDED.endpoint_id, request_candidates.endpoint_id), - key_id = COALESCE(EXCLUDED.key_id, request_candidates.key_id), + user_id = COALESCE(request_candidates.user_id, EXCLUDED.user_id), + api_key_id = COALESCE(request_candidates.api_key_id, EXCLUDED.api_key_id), + username = COALESCE(request_candidates.username, EXCLUDED.username), + api_key_name = COALESCE(request_candidates.api_key_name, EXCLUDED.api_key_name), + provider_id = COALESCE(request_candidates.provider_id, EXCLUDED.provider_id), + endpoint_id = COALESCE(request_candidates.endpoint_id, EXCLUDED.endpoint_id), + key_id = COALESCE(request_candidates.key_id, EXCLUDED.key_id), status = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status @@ -275,11 +249,11 @@ DO UPDATE SET THEN request_candidates.status ELSE EXCLUDED.status END, - skip_reason = COALESCE(EXCLUDED.skip_reason, request_candidates.skip_reason), + skip_reason = COALESCE(EXCLUDED.skip_reason, __AETHER_SANITIZED_LEGACY_SKIP_REASON__), is_cached = COALESCE(EXCLUDED.is_cached, request_candidates.is_cached), status_code = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status_code WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status_code @@ -288,28 +262,23 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.status_code, request_candidates.status_code) END, error_type = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_type - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_type - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_type - ELSE COALESCE(EXCLUDED.error_type, request_candidates.error_type) - END, - error_message = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_message - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_message - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_message - ELSE COALESCE(EXCLUDED.error_message, request_candidates.error_message) + WHEN ( + request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') + AND EXCLUDED.status <> request_candidates.status + ) OR ( + request_candidates.status = 'pending' + AND EXCLUDED.status IN ('available', 'unused') + ) OR ( + request_candidates.status = 'streaming' + AND EXCLUDED.status IN ('available', 'unused', 'pending') + ) OR EXCLUDED.error_type IS NULL + THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__ + ELSE EXCLUDED.error_type END, + error_message = NULL, latency_ms = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.latency_ms WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.latency_ms @@ -318,41 +287,17 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms) END, concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests), - extra_data = CASE - WHEN request_candidates.extra_data IS NULL THEN EXCLUDED.extra_data - WHEN EXCLUDED.extra_data IS NULL THEN regexp_replace( - request_candidates.extra_data::text, - $aether_nul$(? request_candidates.status THEN request_candidates.finished_at WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.finished_at @@ -362,19 +307,19 @@ DO UPDATE SET END "#; -const UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL: &str = r#" +const UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL_TEMPLATE: &str = r#" ON CONFLICT (request_id, candidate_index, retry_index) DO UPDATE SET - user_id = COALESCE(EXCLUDED.user_id, request_candidates.user_id), - api_key_id = COALESCE(EXCLUDED.api_key_id, request_candidates.api_key_id), - username = COALESCE(EXCLUDED.username, request_candidates.username), - api_key_name = COALESCE(EXCLUDED.api_key_name, request_candidates.api_key_name), - provider_id = COALESCE(EXCLUDED.provider_id, request_candidates.provider_id), - endpoint_id = COALESCE(EXCLUDED.endpoint_id, request_candidates.endpoint_id), - key_id = COALESCE(EXCLUDED.key_id, request_candidates.key_id), + user_id = COALESCE(request_candidates.user_id, EXCLUDED.user_id), + api_key_id = COALESCE(request_candidates.api_key_id, EXCLUDED.api_key_id), + username = COALESCE(request_candidates.username, EXCLUDED.username), + api_key_name = COALESCE(request_candidates.api_key_name, EXCLUDED.api_key_name), + provider_id = COALESCE(request_candidates.provider_id, EXCLUDED.provider_id), + endpoint_id = COALESCE(request_candidates.endpoint_id, EXCLUDED.endpoint_id), + key_id = COALESCE(request_candidates.key_id, EXCLUDED.key_id), status = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status @@ -382,11 +327,11 @@ DO UPDATE SET THEN request_candidates.status ELSE EXCLUDED.status END, - skip_reason = COALESCE(EXCLUDED.skip_reason, request_candidates.skip_reason), + skip_reason = COALESCE(EXCLUDED.skip_reason, __AETHER_SANITIZED_LEGACY_SKIP_REASON__), is_cached = request_candidates.is_cached, status_code = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.status_code WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.status_code @@ -395,28 +340,23 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.status_code, request_candidates.status_code) END, error_type = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_type - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_type - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_type - ELSE COALESCE(EXCLUDED.error_type, request_candidates.error_type) - END, - error_message = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_message - WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') - THEN request_candidates.error_message - WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_message - ELSE COALESCE(EXCLUDED.error_message, request_candidates.error_message) + WHEN ( + request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') + AND EXCLUDED.status <> request_candidates.status + ) OR ( + request_candidates.status = 'pending' + AND EXCLUDED.status IN ('available', 'unused') + ) OR ( + request_candidates.status = 'streaming' + AND EXCLUDED.status IN ('available', 'unused', 'pending') + ) OR EXCLUDED.error_type IS NULL + THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__ + ELSE EXCLUDED.error_type END, + error_message = NULL, latency_ms = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming') + AND EXCLUDED.status <> request_candidates.status THEN request_candidates.latency_ms WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.latency_ms @@ -425,41 +365,17 @@ DO UPDATE SET ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms) END, concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests), - extra_data = CASE - WHEN request_candidates.extra_data IS NULL THEN EXCLUDED.extra_data - WHEN EXCLUDED.extra_data IS NULL THEN regexp_replace( - request_candidates.extra_data::text, - $aether_nul$(? request_candidates.status THEN request_candidates.finished_at WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') THEN request_candidates.finished_at @@ -498,6 +414,64 @@ INSERT INTO request_candidates ( ) "#; +const LEGACY_SKIP_REASON_PLACEHOLDER: &str = "__AETHER_SANITIZED_LEGACY_SKIP_REASON__"; +const LEGACY_ERROR_TYPE_PLACEHOLDER: &str = "__AETHER_SANITIZED_LEGACY_ERROR_TYPE__"; + +static UPSERT_SQL: LazyLock = + LazyLock::new(|| postgres_candidate_upsert_sql(UPSERT_SQL_TEMPLATE)); +static UPSERT_CONFLICT_SQL: LazyLock = + LazyLock::new(|| postgres_candidate_upsert_sql(UPSERT_CONFLICT_SQL_TEMPLATE)); +static UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL: LazyLock = + LazyLock::new(|| postgres_candidate_upsert_sql(UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL_TEMPLATE)); + +fn postgres_candidate_upsert_sql(template: &str) -> String { + template + .replace( + LEGACY_SKIP_REASON_PLACEHOLDER, + postgres_sanitized_legacy_diagnostic_sql( + "request_candidates.skip_reason", + REQUEST_CANDIDATE_SKIP_REASONS, + &[], + "unclassified_skip", + ) + .as_str(), + ) + .replace( + LEGACY_ERROR_TYPE_PLACEHOLDER, + postgres_sanitized_legacy_diagnostic_sql( + "request_candidates.error_type", + REQUEST_CANDIDATE_ERROR_TYPES, + REQUEST_CANDIDATE_ERROR_TYPE_ALIASES, + "unclassified_error", + ) + .as_str(), + ) +} + +fn postgres_sanitized_legacy_diagnostic_sql( + column: &str, + allowed: &[&str], + aliases: &[(&str, &str)], + fallback: &str, +) -> String { + let normalized = format!("LOWER(BTRIM({column}))"); + let mut sql = format!("(CASE WHEN {column} IS NULL THEN NULL"); + for (alias, canonical) in aliases { + sql.push_str(format!(" WHEN {normalized} = '{alias}' THEN '{canonical}'").as_str()); + } + sql.push_str(format!(" WHEN {normalized} IN (").as_str()); + for (index, value) in allowed.iter().enumerate() { + if index > 0 { + sql.push_str(", "); + } + sql.push('\''); + sql.push_str(value); + sql.push('\''); + } + sql.push_str(format!(") THEN {normalized} ELSE '{fallback}' END)").as_str()); + sql +} + const MAX_POSTGRES_REQUEST_CANDIDATE_UPSERT_ROWS: usize = 1_000; const DELETE_CREATED_BEFORE_SQL: &str = r#" @@ -782,7 +756,7 @@ impl SqlxRequestCandidateReadRepository { self.tx_runner .run_read_write(|tx| { Box::pin(async move { - let row = sqlx::query(UPSERT_SQL) + let row = sqlx::query(UPSERT_SQL.as_str()) .bind(if candidate.id.trim().is_empty() { Uuid::new_v4().to_string() } else { @@ -1048,9 +1022,9 @@ fn split_request_candidate_upsert_batches( fn upsert_many_conflict_sql(overwrite_is_cached: bool) -> &'static str { if overwrite_is_cached { - UPSERT_CONFLICT_SQL + UPSERT_CONFLICT_SQL.as_str() } else { - UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL + UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL.as_str() } } @@ -1230,6 +1204,7 @@ fn to_i32_u64(value: u64) -> Result { } fn sanitize_request_candidate_for_postgres(candidate: &mut UpsertRequestCandidateRecord) -> usize { + candidate.sanitize_for_persistence(); let mut replacements = 0usize; for value in [ &mut candidate.username, @@ -1310,7 +1285,8 @@ fn replace_nul_characters(value: &mut String) -> usize { mod tests { use super::{ sanitize_request_candidate_for_postgres, SqlxRequestCandidateReadRepository, - UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL, UPSERT_CONFLICT_SQL, UPSERT_SQL, + UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL, UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL_TEMPLATE, + UPSERT_CONFLICT_SQL, UPSERT_CONFLICT_SQL_TEMPLATE, UPSERT_SQL, UPSERT_SQL_TEMPLATE, }; use crate::error::SqlxResultExt; use crate::{PostgresPoolConfig, PostgresPoolFactory}; @@ -1321,29 +1297,43 @@ mod tests { #[test] fn upsert_sql_does_not_default_missing_or_epoch_created_at_to_epoch() { - assert!(!UPSERT_SQL.contains("COALESCE($22, 0)")); - assert!(UPSERT_SQL.contains("WHEN $22 IS NOT NULL AND $22 > 1000.0")); - assert!(UPSERT_SQL.contains("TO_TIMESTAMP($22 / 1000.0)")); - assert!(UPSERT_SQL.contains("TO_TIMESTAMP($23 / 1000.0)")); - assert!(UPSERT_SQL.contains("TO_TIMESTAMP($24 / 1000.0)")); - assert!(UPSERT_SQL.contains("NOW()")); - assert!(UPSERT_SQL.contains("request_candidates.created_at <= TO_TIMESTAMP(1)")); - assert!(UPSERT_SQL.contains("THEN EXCLUDED.created_at")); + assert!(!UPSERT_SQL_TEMPLATE.contains("COALESCE($22, 0)")); + assert!(UPSERT_SQL_TEMPLATE.contains("WHEN $22 IS NOT NULL AND $22 > 1000.0")); + assert!(UPSERT_SQL_TEMPLATE.contains("TO_TIMESTAMP($22 / 1000.0)")); + assert!(UPSERT_SQL_TEMPLATE.contains("TO_TIMESTAMP($23 / 1000.0)")); + assert!(UPSERT_SQL_TEMPLATE.contains("TO_TIMESTAMP($24 / 1000.0)")); + assert!(UPSERT_SQL_TEMPLATE.contains("NOW()")); + assert!(UPSERT_SQL_TEMPLATE.contains("request_candidates.created_at <= TO_TIMESTAMP(1)")); + assert!(UPSERT_SQL_TEMPLATE.contains("THEN EXCLUDED.created_at")); } #[test] - fn upsert_sql_keeps_candidate_lifecycle_monotonic_when_events_arrive_late() { + fn upsert_sql_preserves_first_identity_and_terminal_fact_when_events_arrive_late() { for sql in [ - UPSERT_SQL, - UPSERT_CONFLICT_SQL, - UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL, + UPSERT_SQL_TEMPLATE, + UPSERT_CONFLICT_SQL_TEMPLATE, + UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL_TEMPLATE, ] { + for column in [ + "user_id", + "api_key_id", + "username", + "api_key_name", + "provider_id", + "endpoint_id", + "key_id", + ] { + let expected = + format!("{column} = COALESCE(request_candidates.{column}, EXCLUDED.{column})"); + assert!( + sql.contains(&expected), + "missing immutable identity merge: {expected}" + ); + } assert!(sql.contains( "request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')" )); - assert!( - sql.contains("EXCLUDED.status IN ('available', 'unused', 'pending', 'streaming')") - ); + assert!(sql.contains("EXCLUDED.status <> request_candidates.status")); assert!(sql.contains( "request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending')" )); @@ -1356,7 +1346,7 @@ mod tests { } #[test] - fn postgres_candidate_sanitizer_replaces_nul_in_text_and_nested_json() { + fn postgres_candidate_sanitizer_discards_unapproved_diagnostics_before_nul_repair() { let mut extra_data = Map::new(); extra_data.insert( "bad\0key".to_string(), @@ -1391,32 +1381,33 @@ mod tests { finished_at_unix_ms: Some(2), }; - assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 9); - assert_eq!(candidate.username.as_deref(), Some("user�name")); - assert_eq!(candidate.api_key_name.as_deref(), Some("key�name")); - assert_eq!(candidate.skip_reason.as_deref(), Some("skip�reason")); - assert_eq!(candidate.error_type.as_deref(), Some("upstream�error")); - assert_eq!(candidate.error_message.as_deref(), Some("bad�message")); - assert_eq!( - candidate.extra_data, - Some(json!({"bad�key": {"nested": ["bad�value", {"literal": "\\u0000"}]}})) - ); - assert_eq!( - candidate.required_capabilities, - Some(json!({"cap�key": "cap�value"})) - ); + assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 0); + assert!(candidate.username.is_none()); + assert!(candidate.api_key_name.is_none()); + assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip")); + assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error")); + assert!(candidate.error_message.is_none()); + assert!(candidate.extra_data.is_none()); + assert!(candidate.required_capabilities.is_none()); } #[test] - fn every_postgres_candidate_conflict_path_repairs_legacy_json_nul_escapes() { + fn every_postgres_candidate_conflict_path_discards_legacy_diagnostics() { for sql in [ - UPSERT_SQL, - UPSERT_CONFLICT_SQL, - UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL, + UPSERT_SQL.as_str(), + UPSERT_CONFLICT_SQL.as_str(), + UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL.as_str(), ] { - assert!(sql.contains("regexp_replace(")); - assert!(sql.contains(r"(?, _>(&raw, "error_message") + .expect("error_message should decode") + .is_none() + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "skip_reason") + .expect("skip_reason should decode") + .as_deref(), + Some("unclassified_skip") + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "error_type") + .expect("error_type should decode") + .as_deref(), + Some("unclassified_error") + ); + assert!( + sqlx::Row::try_get::, _>(&raw, "extra_data") + .expect("extra_data should decode") + .is_none() + ); + assert!(sqlx::Row::try_get::, _>( + &raw, + "required_capabilities" + ) + .expect("required_capabilities should decode") + .is_none()); + let rows = repository .list_by_request_id(request_id) .await .expect("sanitized candidate should be readable"); assert_eq!(rows.len(), 1); assert_eq!(rows[0].status, RequestCandidateStatus::Success); - assert_eq!(rows[0].error_message.as_deref(), Some("bad�message")); - assert_eq!( - rows[0].extra_data, - Some(json!({ - "old�key": "old�value", - "literal": "\\u0000", - "adjacent": "��", - "new": true, - "nested": "new�value" - })) - ); - assert_eq!( - rows[0].required_capabilities, - Some(json!({"cap�key": "cap�value"})) - ); + assert!(rows[0].error_message.is_none()); + assert!(rows[0].extra_data.is_none()); + assert!(rows[0].required_capabilities.is_none()); } assert_eq!( repository diff --git a/crates/aether-data/adapters/postgres/src/gemini_file_mappings.rs b/crates/aether-data/adapters/postgres/src/gemini_file_mappings.rs index 8852fcbfb..057cd2360 100644 --- a/crates/aether-data/adapters/postgres/src/gemini_file_mappings.rs +++ b/crates/aether-data/adapters/postgres/src/gemini_file_mappings.rs @@ -86,6 +86,77 @@ WHERE file_name = $1 } } + async fn find_active_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms, + EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE file_name = $1 + AND user_id = $2 + AND expires_at > TO_TIMESTAMP($3::double precision) +"#, + ) + .bind(file_name) + .bind(user_id) + .bind(now_unix_secs as f64) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + + row.as_ref().map(Self::map_row).transpose() + } + + async fn find_active_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms, + EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE file_name = $1 + AND key_id = $2 + AND user_id = $3 + AND expires_at > TO_TIMESTAMP($4::double precision) +"#, + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .bind(now_unix_secs as f64) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + + row.as_ref().map(Self::map_row).transpose() + } + async fn list_mappings( &self, query: &GeminiFileMappingListQuery, @@ -223,6 +294,61 @@ RETURNING Self::map_row(&row) } + async fn upsert_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, DataLayerError> { + record.validate()?; + let row = sqlx::query( + r#" +INSERT INTO gemini_file_mappings ( + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at, + expires_at +) +VALUES ($1,$2,$3,$4,$5,$6,$7,NOW(),TO_TIMESTAMP($8::double precision)) +ON CONFLICT (file_name) +DO UPDATE +SET + display_name = EXCLUDED.display_name, + mime_type = EXCLUDED.mime_type, + source_hash = EXCLUDED.source_hash, + expires_at = EXCLUDED.expires_at +WHERE gemini_file_mappings.key_id = EXCLUDED.key_id + AND gemini_file_mappings.user_id IS NOT DISTINCT FROM EXCLUDED.user_id +RETURNING + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms, + EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs +"#, + ) + .bind(record.id) + .bind(record.file_name) + .bind(record.key_id) + .bind(record.user_id) + .bind(record.display_name) + .bind(record.mime_type) + .bind(record.source_hash) + .bind(record.expires_at_unix_secs as f64) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + + row.as_ref().map(Self::map_row).transpose() + } + async fn delete_by_file_name(&self, file_name: &str) -> Result { let result = sqlx::query( r#" @@ -238,6 +364,48 @@ WHERE file_name = $1 Ok(result.rows_affected() > 0) } + async fn delete_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + let result = sqlx::query( + r#" +DELETE FROM gemini_file_mappings +WHERE file_name = $1 AND user_id = $2 +"#, + ) + .bind(file_name) + .bind(user_id) + .execute(&self.pool) + .await + .map_postgres_err()?; + + Ok(result.rows_affected() > 0) + } + + async fn delete_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + let result = sqlx::query( + r#" +DELETE FROM gemini_file_mappings +WHERE file_name = $1 AND key_id = $2 AND user_id = $3 +"#, + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .execute(&self.pool) + .await + .map_postgres_err()?; + + Ok(result.rows_affected() > 0) + } + async fn delete_by_id( &self, mapping_id: &str, @@ -325,6 +493,11 @@ fn apply_list_filters( where_clause: &mut WhereClause, query: &GeminiFileMappingListQuery, ) { + if let Some(user_id) = query.user_id.as_deref() { + where_clause.push_next(builder); + builder.push("user_id = "); + builder.push_bind(user_id.to_string()); + } if !query.include_expired { where_clause.push_next(builder); builder.push("expires_at > TO_TIMESTAMP("); diff --git a/crates/aether-data/adapters/postgres/src/management_tokens.rs b/crates/aether-data/adapters/postgres/src/management_tokens.rs index a43bfb591..9e19b5003 100644 --- a/crates/aether-data/adapters/postgres/src/management_tokens.rs +++ b/crates/aether-data/adapters/postgres/src/management_tokens.rs @@ -2,10 +2,10 @@ use async_trait::async_trait; use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row}; use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser, - UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, UpdateManagementTokenRecord, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause}; @@ -39,6 +39,7 @@ JOIN users u ON u.id = mt.user_id const DELETE_MANAGEMENT_TOKEN_SQL: &str = r#" DELETE FROM management_tokens WHERE id = $1 + AND ($2::text IS NULL OR user_id = $2) "#; const MANAGEMENT_TOKEN_JSON_COLUMN_TYPES_SQL: &str = r#" @@ -97,23 +98,33 @@ RETURNING const UPDATE_MANAGEMENT_TOKEN_SQL_PREFIX: &str = r#" UPDATE management_tokens -SET name = $2, - description = $3, - allowed_ips = +SET name = COALESCE($2, name), + description = CASE + WHEN $3 THEN NULL + ELSE COALESCE($4, description) + END, + allowed_ips = CASE + WHEN $5 THEN NULL + ELSE COALESCE( "#; const UPDATE_MANAGEMENT_TOKEN_SQL_MIDDLE: &str = r#", - permissions = + allowed_ips) + END, + permissions = COALESCE( "#; const UPDATE_MANAGEMENT_TOKEN_SQL_SUFFIX: &str = r#", + permissions), expires_at = CASE - WHEN $6::bigint IS NULL THEN NULL - ELSE to_timestamp($6::double precision) + WHEN $8 THEN NULL + WHEN $9::bigint IS NOT NULL THEN to_timestamp($9::double precision) + ELSE expires_at END, - is_active = $7, + is_active = COALESCE($10, is_active), updated_at = NOW() WHERE id = $1 + AND ($11::text IS NULL OR user_id = $11) RETURNING id, user_id, @@ -136,6 +147,7 @@ UPDATE management_tokens SET is_active = $2, updated_at = NOW() WHERE id = $1 + AND ($3::text IS NULL OR user_id = $3) RETURNING id, user_id, @@ -153,12 +165,56 @@ RETURNING EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs "#; +const LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL: &str = r#" +SELECT id +FROM users +WHERE id = $1 + AND is_active IS TRUE + AND is_deleted IS FALSE + AND LOWER(role::text) = 'admin' + AND security_version = $2 +FOR UPDATE +"#; + +const LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL: &str = r#" +SELECT + id, + user_id, + token_hash, + name, + description, + token_prefix, + allowed_ips, + permissions, + EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs, + EXTRACT(EPOCH FROM last_used_at)::bigint AS last_used_at_unix_secs, + last_used_ip, + COALESCE(usage_count, 0) AS usage_count, + is_active, + EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms, + EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs +FROM management_tokens +WHERE id = $1 +FOR UPDATE +"#; + +const ACTIVATE_LOCKED_MANAGEMENT_TOKEN_SQL: &str = r#" +UPDATE management_tokens +SET is_active = TRUE, + updated_at = NOW() +WHERE id = $1 + AND token_hash = $2 + AND is_active = FALSE + AND (expires_at IS NULL OR expires_at > to_timestamp($3::double precision)) +"#; + const REGENERATE_MANAGEMENT_TOKEN_SECRET_SQL: &str = r#" UPDATE management_tokens SET token_hash = $2, token_prefix = $3, updated_at = NOW() WHERE id = $1 + AND ($4::text IS NULL OR user_id = $4) RETURNING id, user_id, @@ -271,6 +327,85 @@ impl SqlxManagementTokenRepository { )), } } + + async fn update_management_token_scoped( + &self, + record: &UpdateManagementTokenRecord, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + record.validate()?; + let json_column_types = self.json_column_types().await?; + let sql = update_management_token_sql(json_column_types); + let allowed_ips = json_to_string(record.allowed_ips.as_ref())?; + let permissions = json_to_string(record.permissions.as_ref())?; + let row = sqlx::query(sql.as_str()) + .bind(&record.token_id) + .bind(record.name.as_deref()) + .bind(record.clear_description) + .bind(record.description.as_deref()) + .bind(record.clear_allowed_ips) + .bind(allowed_ips) + .bind(permissions) + .bind(record.clear_expires_at) + .bind( + record + .expires_at_unix_secs + .and_then(|value| i64::try_from(value).ok()), + ) + .bind(record.is_active) + .bind(expected_user_id) + .fetch_optional(&self.pool) + .await + .map_err(|err| map_management_token_write_error(err, record.name.as_deref()))?; + row.as_ref().map(map_token_row).transpose() + } + + async fn delete_management_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> Result { + let result = sqlx::query(DELETE_MANAGEMENT_TOKEN_SQL) + .bind(token_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() > 0) + } + + async fn set_management_token_active_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + is_active: bool, + ) -> Result, DataLayerError> { + let row = sqlx::query(SET_MANAGEMENT_TOKEN_ACTIVE_SQL) + .bind(token_id) + .bind(is_active) + .bind(expected_user_id) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_token_row).transpose() + } + + async fn regenerate_management_token_secret_scoped( + &self, + mutation: &RegenerateManagementTokenSecret, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + mutation.validate()?; + let row = sqlx::query(REGENERATE_MANAGEMENT_TOKEN_SECRET_SQL) + .bind(&mutation.token_id) + .bind(&mutation.token_hash) + .bind(mutation.token_prefix.as_deref()) + .bind(expected_user_id) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_token_row).transpose() + } } fn create_management_token_sql(types: ManagementTokenJsonColumnTypes) -> String { @@ -285,7 +420,7 @@ fn create_management_token_sql(types: ManagementTokenJsonColumnTypes) -> String fn update_management_token_sql(types: ManagementTokenJsonColumnTypes) -> String { format!( - "{} $4::text::{}{} $5::text::{}{}", + "{} $6::text::{}{} $7::text::{}{}", UPDATE_MANAGEMENT_TOKEN_SQL_PREFIX, types.allowed_ips.sql_type(), UPDATE_MANAGEMENT_TOKEN_SQL_MIDDLE, @@ -423,70 +558,29 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository { &self, record: &UpdateManagementTokenRecord, ) -> Result, DataLayerError> { - record.validate()?; - let Some(current) = self - .get_management_token_with_user(&record.token_id) - .await? - else { - return Ok(None); - }; - let json_column_types = self.json_column_types().await?; - let sql = update_management_token_sql(json_column_types); - let name = record - .name - .as_deref() - .unwrap_or(current.token.name.as_str()); - let description = if record.clear_description { - None - } else { - record - .description - .as_deref() - .or(current.token.description.as_deref()) - }; - let allowed_ips = if record.clear_allowed_ips { - None - } else { - record - .allowed_ips - .as_ref() - .or(current.token.allowed_ips.as_ref()) - }; - let permissions = record - .permissions - .as_ref() - .or(current.token.permissions.as_ref()); - let expires_at_unix_secs = if record.clear_expires_at { - None - } else { - record - .expires_at_unix_secs - .or(current.token.expires_at_unix_secs) - }; - let is_active = record.is_active.unwrap_or(current.token.is_active); - let allowed_ips = json_to_string(allowed_ips)?; - let permissions = json_to_string(permissions)?; - let row = sqlx::query(sql.as_str()) - .bind(&record.token_id) - .bind(name) - .bind(description) - .bind(allowed_ips) - .bind(permissions) - .bind(expires_at_unix_secs.and_then(|value| i64::try_from(value).ok())) - .bind(is_active) - .fetch_optional(&self.pool) + self.update_management_token_scoped(record, None).await + } + + async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, DataLayerError> { + self.update_management_token_scoped(record, Some(user_id)) .await - .map_err(|err| map_management_token_write_error(err, record.name.as_deref()))?; - row.as_ref().map(map_token_row).transpose() } async fn delete_management_token(&self, token_id: &str) -> Result { - let result = sqlx::query(DELETE_MANAGEMENT_TOKEN_SQL) - .bind(token_id) - .execute(&self.pool) + self.delete_management_token_scoped(token_id, None).await + } + + async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + self.delete_management_token_scoped(token_id, Some(user_id)) .await - .map_postgres_err()?; - Ok(result.rows_affected() > 0) } async fn set_management_token_active( @@ -494,28 +588,125 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository { token_id: &str, is_active: bool, ) -> Result, DataLayerError> { - let row = sqlx::query(SET_MANAGEMENT_TOKEN_ACTIVE_SQL) - .bind(token_id) - .bind(is_active) - .fetch_optional(&self.pool) + self.set_management_token_active_scoped(token_id, None, is_active) + .await + } + + async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + self.set_management_token_active_scoped(token_id, Some(user_id), is_active) + .await + } + + async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + let eligible_user = + sqlx::query_scalar::<_, String>(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL) + .bind(&mutation.expected_token.user_id) + .bind(mutation.expected_user_security_version) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if eligible_user.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) .await .map_postgres_err()?; - row.as_ref().map(map_token_row).transpose() + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_postgres_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + + let result = sqlx::query(ACTIVATE_LOCKED_MANAGEMENT_TOKEN_SQL) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .bind(i64::try_from(mutation.now_unix_secs).unwrap_or(i64::MAX)) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + + async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_postgres_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let result = sqlx::query( + "DELETE FROM management_tokens WHERE id = $1 AND token_hash = $2 AND is_active = FALSE", + ) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) } async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, ) -> Result, DataLayerError> { - mutation.validate()?; - let row = sqlx::query(REGENERATE_MANAGEMENT_TOKEN_SECRET_SQL) - .bind(&mutation.token_id) - .bind(&mutation.token_hash) - .bind(mutation.token_prefix.as_deref()) - .fetch_optional(&self.pool) + self.regenerate_management_token_secret_scoped(mutation, None) + .await + } + + async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, DataLayerError> { + self.regenerate_management_token_secret_scoped(mutation, Some(user_id)) .await - .map_postgres_err()?; - row.as_ref().map(map_token_row).transpose() } async fn record_management_token_usage( @@ -533,8 +724,18 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository { } } -fn optional_unix_secs(value: Option) -> Option { - value.and_then(|value| u64::try_from(value).ok()) +fn non_negative_u64(value: i64, field_name: &str) -> Result { + u64::try_from(value).map_err(|_| { + DataLayerError::UnexpectedValue(format!( + "management_tokens.{field_name} must not be negative" + )) + }) +} + +fn optional_unix_secs(value: Option, field_name: &str) -> Result, DataLayerError> { + value + .map(|value| non_negative_u64(value, field_name)) + .transpose() } fn json_to_string(value: Option<&serde_json::Value>) -> Result, DataLayerError> { @@ -588,15 +789,30 @@ fn map_token_row(row: &PgRow) -> Result { ) .with_permissions(row.try_get("permissions").map_postgres_err()?) .with_runtime_fields( - optional_unix_secs(row.try_get("expires_at_unix_secs").map_postgres_err()?), - optional_unix_secs(row.try_get("last_used_at_unix_secs").map_postgres_err()?), + optional_unix_secs( + row.try_get("expires_at_unix_secs").map_postgres_err()?, + "expires_at", + )?, + optional_unix_secs( + row.try_get("last_used_at_unix_secs").map_postgres_err()?, + "last_used_at", + )?, row.try_get("last_used_ip").map_postgres_err()?, - u64::try_from(row.try_get::("usage_count").map_postgres_err()?).unwrap_or(0), + non_negative_u64( + row.try_get::("usage_count").map_postgres_err()?, + "usage_count", + )?, row.try_get("is_active").map_postgres_err()?, ) .with_timestamps( - optional_unix_secs(row.try_get("created_at_unix_ms").map_postgres_err()?), - optional_unix_secs(row.try_get("updated_at_unix_secs").map_postgres_err()?), + optional_unix_secs( + row.try_get("created_at_unix_ms").map_postgres_err()?, + "created_at", + )?, + optional_unix_secs( + row.try_get("updated_at_unix_secs").map_postgres_err()?, + "updated_at", + )?, )) } @@ -619,8 +835,10 @@ fn map_token_with_user_row(row: &PgRow) -> Result Result { + ldap_exclusive: bool, + force_disable: bool, + _locked_users_snapshot: usize, + ) -> Result { record.validate()?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + let existing_enabled: Option = if record.is_enabled || force_disable { + None + } else { + // All provider status changes serialize in provider_type order before an + // enabled-link count is used to authorize a disable. + sqlx::query_scalar::<_, String>( + "SELECT provider_type FROM oauth_providers ORDER BY provider_type FOR UPDATE", + ) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = $1") + .bind(&record.provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + }; + if existing_enabled == Some(true) { + let affected_count: i64 = + sqlx::query_scalar(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL) + .bind(&record.provider_type) + .bind(ldap_exclusive) + .fetch_one(&mut *tx) + .await + .map_postgres_err()?; + let affected_count = usize::try_from(affected_count).map_err(|_| { + DataLayerError::UnexpectedValue( + "oauth_providers.locked_user_count is negative".to_string(), + ) + })?; + if affected_count > 0 { + tx.rollback().await.map_postgres_err()?; + return Ok( + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count, + }, + ); + } + } let row = sqlx::query(UPSERT_OAUTH_PROVIDER_CONFIG_SQL) .bind(&record.provider_type) .bind(&record.display_name) @@ -239,22 +294,57 @@ impl OAuthProviderWriteRepository for SqlxOAuthProviderRepository { .bind(record.extra_config.as_ref()) .bind(record.icon_url.as_deref()) .bind(record.is_enabled) - .fetch_one(&self.pool) + .fetch_one(&mut *tx) .await .map_postgres_err()?; - map_oauth_provider_row(&row) + let provider = map_oauth_provider_row(&row)?; + tx.commit().await.map_postgres_err()?; + Ok(UpsertOAuthProviderConfigOutcome::Upserted(provider)) } - async fn delete_oauth_provider_config( + async fn compare_and_swap_oauth_provider_client_secret( &self, provider_type: &str, + expected: &str, + replacement: &str, ) -> Result { - let result = sqlx::query(DELETE_OAUTH_PROVIDER_CONFIG_SQL) + let result = sqlx::query(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL) .bind(provider_type) + .bind(expected) + .bind(replacement) .execute(&self.pool) .await .map_postgres_err()?; - Ok(result.rows_affected() > 0) + Ok(result.rows_affected() == 1) + } + + async fn delete_oauth_provider_config_if_unlinked( + &self, + provider_type: &str, + has_links_snapshot: bool, + ) -> Result { + if has_links_snapshot { + return Ok(false); + } + let mut tx = self.pool.begin().await.map_postgres_err()?; + let provider_exists: Option = sqlx::query_scalar( + "SELECT provider_type FROM oauth_providers WHERE provider_type = $1 FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if provider_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let result = sqlx::query(DELETE_OAUTH_PROVIDER_CONFIG_SQL) + .bind(provider_type) + .execute(&mut *tx) + .await + .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; + Ok(result.rows_affected() == 1) } } @@ -349,7 +439,10 @@ fn map_oauth_provider_row(row: &PgRow) -> Result Result { config.validate()?; - - let ssl_mode = if config.require_ssl { - PgSslMode::Require - } else { - PgSslMode::Prefer - }; - let options = PgConnectOptions::from_str(config.database_url.trim()).map_err(|err| { DataLayerError::InvalidConfiguration(format!("invalid postgres database_url: {err}")) })?; + // Preserve an explicit verification mode from the URL. `require_ssl` is + // a minimum transport guarantee, so it may upgrade Disable/Allow/Prefer + // to Require but must never silently weaken VerifyCa/VerifyFull. + let ssl_mode = if config.require_ssl + && !matches!( + options.get_ssl_mode(), + PgSslMode::VerifyCa | PgSslMode::VerifyFull + ) { + PgSslMode::Require + } else { + options.get_ssl_mode() + }; + Ok(options .ssl_mode(ssl_mode) .statement_cache_capacity(config.statement_cache_capacity)) @@ -53,8 +59,43 @@ impl PostgresPoolFactory { #[cfg(test)] mod tests { - use super::PostgresPoolFactory; + use super::{connect_options, PostgresPoolFactory}; use crate::PostgresPoolConfig; + use sqlx::postgres::PgSslMode; + + fn ssl_mode(url: &str, require_ssl: bool) -> PgSslMode { + connect_options(&PostgresPoolConfig { + database_url: url.to_string(), + require_ssl, + ..PostgresPoolConfig::default() + }) + .expect("postgres options should parse") + .get_ssl_mode() + } + + #[test] + fn preserves_explicit_postgres_verification_modes() { + assert!(matches!( + ssl_mode("postgres://localhost/aether?sslmode=verify-full", false), + PgSslMode::VerifyFull + )); + assert!(matches!( + ssl_mode("postgres://localhost/aether?sslmode=verify-ca", true), + PgSslMode::VerifyCa + )); + } + + #[test] + fn require_ssl_only_upgrades_weak_postgres_modes() { + for mode in ["disable", "allow", "prefer"] { + let url = format!("postgres://localhost/aether?sslmode={mode}"); + assert!(matches!(ssl_mode(&url, true), PgSslMode::Require)); + } + assert!(matches!( + ssl_mode("postgres://localhost/aether", false), + PgSslMode::Prefer + )); + } #[tokio::test] async fn factory_builds_lazy_pool_from_valid_config() { diff --git a/crates/aether-data/adapters/postgres/src/provider_catalog.rs b/crates/aether-data/adapters/postgres/src/provider_catalog.rs index 85704aef4..fa19b83ad 100644 --- a/crates/aether-data/adapters/postgres/src/provider_catalog.rs +++ b/crates/aether-data/adapters/postgres/src/provider_catalog.rs @@ -10,9 +10,11 @@ use sqlx::{ use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate, - ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate, + ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, + ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, @@ -273,6 +275,7 @@ SELECT is_active, api_formats, NULL::jsonb AS auth_type_by_format, + NULL::jsonb AS allow_auth_channel_mismatch_formats, 'summary' AS api_key, CASE WHEN auth_config IS NULL THEN NULL @@ -398,7 +401,18 @@ WHERE id = $1 AND ($6::text IS NULL OR auth_config IS NOT DISTINCT FROM $6) "#; -const KEY_RUNTIME_METADATA_CAS_SQL: &str = r#" +const KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL: &str = r#" +SELECT + jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' + AS metadata_is_object, + COALESCE(upstream_metadata, '{}'::jsonb) ? $2 AS namespace_exists, + COALESCE(upstream_metadata, '{}'::jsonb) -> $2 AS namespace_value +FROM provider_api_keys +WHERE id = $1 +FOR UPDATE +"#; + +const KEY_RUNTIME_METADATA_UPDATE_SQL: &str = r#" UPDATE provider_api_keys SET upstream_metadata = COALESCE(upstream_metadata, '{}'::jsonb) @@ -410,10 +424,58 @@ SET END WHERE id = $1 AND jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' - AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $2) - IS NOT DISTINCT FROM $6::jsonb "#; +fn runtime_metadata_namespace_matches( + metadata_is_object: bool, + namespace_exists: bool, + current: Option<&serde_json::Value>, + expected: Option<&serde_json::Value>, +) -> bool { + metadata_is_object + && match expected { + Some(expected) => namespace_exists && current == Some(expected), + None => !namespace_exists, + } +} + +async fn lock_runtime_metadata_namespace_matches( + tx: &mut sqlx::Transaction<'_, Postgres>, + key_id: &str, + namespace: &str, + expected: Option<&serde_json::Value>, +) -> Result { + let Some(row) = sqlx::query(KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL) + .bind(key_id) + .bind(namespace) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(false); + }; + let metadata_is_object = row + .try_get::("metadata_is_object") + .map_postgres_err()?; + let namespace_exists = row + .try_get::("namespace_exists") + .map_postgres_err()?; + let current = row + .try_get::, _>("namespace_value") + .map_postgres_err()?; + + // PostgreSQL jsonb retains decimal lexemes that serde_json's default + // Number representation rounds to f64. Re-read and compare while holding + // the row lock instead of binding that rounded value back into a jsonb + // equality predicate, which would report a false CAS conflict. + Ok(runtime_metadata_namespace_matches( + metadata_is_object, + namespace_exists, + current.as_ref(), + expected, + )) +} + fn validate_key_for_update(key: &StoredProviderCatalogKey) -> Result<(), DataLayerError> { if key.id.trim().is_empty() { return Err(DataLayerError::InvalidInput( @@ -907,56 +969,11 @@ impl SqlxProviderCatalogReadRepository { .await } - pub async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - if key_id.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "provider catalog key_id is empty".to_string(), - )); - } - if encrypted_api_key.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "provider catalog oauth api_key is empty".to_string(), - )); - } - - let rows_affected = sqlx::query( - r#" -UPDATE provider_api_keys -SET - api_key = $2, - auth_config = $3, - expires_at = CASE - WHEN $4::double precision IS NULL THEN NULL - ELSE TO_TIMESTAMP($4::double precision) - END, - updated_at = NOW() -WHERE id = $1 -"#, - ) - .bind(key_id) - .bind(encrypted_api_key) - .bind(encrypted_auth_config) - .bind(expires_at_unix_secs.map(|value| value as f64)) - .execute(&self.pool) - .await - .map_postgres_err()? - .rows_affected(); - - Ok(rows_affected > 0) - } - pub async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { if key_id.trim().is_empty() { @@ -973,10 +990,9 @@ SET ELSE TO_TIMESTAMP($2::double precision) END, oauth_invalid_reason = $3, - auth_config = COALESCE($4, auth_config), updated_at = CASE - WHEN $5::double precision IS NULL THEN NOW() - ELSE TO_TIMESTAMP($5::double precision) + WHEN $4::double precision IS NULL THEN NOW() + ELSE TO_TIMESTAMP($4::double precision) END WHERE id = $1 "#, @@ -984,7 +1000,6 @@ WHERE id = $1 .bind(key_id) .bind(oauth_invalid_at_unix_secs.map(|value| value as f64)) .bind(oauth_invalid_reason) - .bind(encrypted_auth_config_update) .bind(updated_at_unix_secs.map(|value| value as f64)) .execute(&self.pool) .await @@ -1041,6 +1056,20 @@ WHERE id = $1 .to_string(), )); } + let mut tx = self.pool.begin().await.map_postgres_err()?; + if let Some(expected) = update.expected_upstream_metadata_namespace.as_ref() { + let matches = lock_runtime_metadata_namespace_matches( + &mut tx, + &update.key_id, + &expected.namespace, + expected.expected_value.as_ref(), + ) + .await?; + if !matches { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + } let rows_affected = sqlx::query( r#" UPDATE provider_api_keys @@ -1089,14 +1118,6 @@ WHERE id = $1 AND providers.provider_type = $18 ) ) - AND ( - $19::boolean IS FALSE - OR ( - jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' - AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $20) - IS NOT DISTINCT FROM $21::jsonb - ) - ) "#, ) .bind(&update.key_id) @@ -1142,24 +1163,16 @@ WHERE id = $1 .as_ref() .map(|expected| expected.provider_type.as_str()), ) - .bind(update.expected_upstream_metadata_namespace.is_some()) - .bind( - update - .expected_upstream_metadata_namespace - .as_ref() - .map(|expected| expected.namespace.as_str()), - ) - .bind( - update - .expected_upstream_metadata_namespace - .as_ref() - .and_then(|expected| expected.expected_value.as_ref()), - ) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()? .rows_affected(); - Ok(rows_affected > 0) + if rows_affected == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) } pub async fn create_provider( @@ -1455,6 +1468,60 @@ WHERE id = $1 }) } + pub async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + if update.provider_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog provider_id is empty".to_string(), + )); + } + let rows_affected = sqlx::query( + r#" +UPDATE providers +SET config = $3, updated_at = NOW() +WHERE id = $1 + AND config::jsonb IS NOT DISTINCT FROM $2::jsonb +"#, + ) + .bind(&update.provider_id) + .bind(&update.expected_config) + .bind(&update.config) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + + pub async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + if update.record_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog provider_id is empty".to_string(), + )); + } + let rows_affected = sqlx::query( + r#" +UPDATE providers +SET proxy = $3, updated_at = NOW() +WHERE id = $1 + AND proxy::jsonb IS NOT DISTINCT FROM $2::jsonb +"#, + ) + .bind(&update.record_id) + .bind(&update.expected_proxy) + .bind(&update.proxy) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + pub async fn delete_provider(&self, provider_id: &str) -> Result { if provider_id.trim().is_empty() { return Err(DataLayerError::InvalidInput( @@ -2112,6 +2179,33 @@ WHERE id = $1 }) } + pub async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + if update.record_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog endpoint_id is empty".to_string(), + )); + } + let rows_affected = sqlx::query( + r#" +UPDATE provider_endpoints +SET proxy = $3, updated_at = NOW() +WHERE id = $1 + AND proxy::jsonb IS NOT DISTINCT FROM $2::jsonb +"#, + ) + .bind(&update.record_id) + .bind(&update.expected_proxy) + .bind(&update.proxy) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + pub async fn delete_endpoint(&self, endpoint_id: &str) -> Result { if endpoint_id.trim().is_empty() { return Err(DataLayerError::InvalidInput( @@ -2164,6 +2258,65 @@ WHERE id = $1 }) } + pub async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + if update.record_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog key_id is empty".to_string(), + )); + } + let rows_affected = sqlx::query( + r#" +UPDATE provider_api_keys +SET proxy = $3, updated_at = NOW() +WHERE id = $1 + AND proxy::jsonb IS NOT DISTINCT FROM $2::jsonb +"#, + ) + .bind(&update.record_id) + .bind(&update.expected_proxy) + .bind(&update.proxy) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + + pub async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + if update.key_id.trim().is_empty() || update.expected_provider_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog key credential CAS requires key_id and provider_id".to_string(), + )); + } + let rows_affected = sqlx::query( + r#" +UPDATE provider_api_keys +SET api_key = $5, encrypted_key = NULL, auth_config = $6 +WHERE id = $1 + AND provider_id = $2 + AND COALESCE(api_key, encrypted_key) IS NOT DISTINCT FROM $3 + AND auth_config IS NOT DISTINCT FROM $4 +"#, + ) + .bind(&update.key_id) + .bind(&update.expected_provider_id) + .bind(update.expected_encrypted_api_key.as_deref()) + .bind(update.expected_encrypted_auth_config.as_deref()) + .bind(update.encrypted_api_key.as_deref()) + .bind(update.encrypted_auth_config.as_deref()) + .execute(&self.pool) + .await + .map_postgres_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + pub async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -2702,18 +2855,34 @@ WHERE id = $1 update: &ProviderCatalogKeyRuntimeMetadataUpdate, ) -> Result { validate_runtime_metadata_update(update)?; - let rows_affected = sqlx::query(KEY_RUNTIME_METADATA_CAS_SQL) + let mut tx = self.pool.begin().await.map_postgres_err()?; + if !lock_runtime_metadata_namespace_matches( + &mut tx, + &update.key_id, + &update.namespace, + update.expected_upstream_metadata_value.as_ref(), + ) + .await? + { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let rows_affected = sqlx::query(KEY_RUNTIME_METADATA_UPDATE_SQL) .bind(&update.key_id) .bind(&update.namespace) .bind(&update.upstream_metadata_value) .bind(&update.status_snapshot_patch) .bind(update.updated_at_unix_secs.map(|value| value as f64)) - .bind(update.expected_upstream_metadata_value.as_ref()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()? .rows_affected(); - Ok(rows_affected > 0) + if rows_affected == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) } pub async fn update_key_status_snapshot( @@ -2861,6 +3030,20 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository { Self::update_provider(self, provider).await } + async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_config(self, update).await + } + + async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_proxy(self, update).await + } + async fn delete_provider(&self, provider_id: &str) -> Result { Self::delete_provider(self, provider_id).await } @@ -2896,6 +3079,13 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository { self.update_endpoint(endpoint).await } + async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_endpoint_proxy(self, update).await + } + async fn delete_endpoint(&self, endpoint_id: &str) -> Result { Self::delete_endpoint(self, endpoint_id).await } @@ -2914,6 +3104,20 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository { Self::update_key(self, key).await } + async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_key_proxy(self, update).await + } + + async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + Self::compare_and_swap_key_credentials(self, update).await + } + async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -3008,29 +3212,11 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository { Self::clear_key_oauth_invalid_marker(self, key_id).await } - async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - Self::update_key_oauth_credentials( - self, - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } - async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { Self::update_key_oauth_runtime_state( @@ -3038,7 +3224,6 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -3163,6 +3348,16 @@ where row.try_get(column).map_postgres_err() } +fn optional_u64(value: Option, field_name: &str) -> Result, DataLayerError> { + value + .map(|value| { + u64::try_from(value).map_err(|_| { + DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}")) + }) + }) + .transpose() +} + fn map_provider_row(row: &PgRow) -> Result { let quota_reset_day = row_get::>(row, "quota_reset_day")? .map(|value| { @@ -3203,8 +3398,14 @@ fn map_provider_row(row: &PgRow) -> Result>(row, "quota_last_reset_at_unix_secs")?.map(|value| value as u64), - row_get::>(row, "quota_expires_at_unix_secs")?.map(|value| value as u64), + optional_u64( + row_get(row, "quota_last_reset_at_unix_secs")?, + "providers.quota_last_reset_at", + )?, + optional_u64( + row_get(row, "quota_expires_at_unix_secs")?, + "providers.quota_expires_at", + )?, ) .with_routing_fields(row_get(row, "provider_priority")?) .with_transport_fields( @@ -3321,6 +3522,10 @@ fn map_key_maintenance_summary_row( } fn map_key_row(row: &PgRow) -> Result { + let expires_at_unix_secs = optional_u64( + row_get(row, "expires_at_unix_secs")?, + "provider_api_keys.expires_at", + )?; let rpm_limit = row_get::>(row, "rpm_limit")? .map(|value| { u32::try_from(value).map_err(|_| { @@ -3471,6 +3676,9 @@ fn map_key_row(row: &PgRow) -> Result }) }) .transpose()?; + let auth_type_by_format: Option = row_get(row, "auth_type_by_format")?; + let allow_auth_channel_mismatch_formats: Option = + row_get(row, "allow_auth_channel_mismatch_formats")?; StoredProviderCatalogKey::new( row_get(row, "id")?, @@ -3487,8 +3695,7 @@ fn map_key_row(row: &PgRow) -> Result row_get(row, "rate_multipliers")?, row_get(row, "global_priority_by_format")?, row_get(row, "allowed_models")?, - row_get::>(row, "expires_at_unix_secs")? - .and_then(|value| u64::try_from(value).ok()), + expires_at_unix_secs, row_get(row, "proxy")?, row_get(row, "fingerprint")?, ) @@ -3515,9 +3722,6 @@ fn map_key_row(row: &PgRow) -> Result row.try_get("circuit_breaker_by_format").ok(), ); key.note = row.try_get("note").ok(); - key.auth_type_by_format = row.try_get("auth_type_by_format").ok(); - key.allow_auth_channel_mismatch_formats = - row.try_get("allow_auth_channel_mismatch_formats").ok(); key.internal_priority = row.try_get("internal_priority").unwrap_or(50); key.cache_ttl_minutes = row.try_get("cache_ttl_minutes").unwrap_or(5); key.max_probe_interval_minutes = row.try_get("max_probe_interval_minutes").unwrap_or(32); @@ -3540,11 +3744,20 @@ fn map_key_row(row: &PgRow) -> Result key.updated_at_unix_secs = updated_at_unix_secs; key }) + .and_then(|key| { + key.with_auth_channel_policy_fields( + auth_type_by_format, + allow_auth_channel_mismatch_formats, + ) + }) } #[cfg(test)] mod tests { - use super::SqlxProviderCatalogReadRepository; + use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate; + use serde_json::json; + + use super::{optional_u64, SqlxProviderCatalogReadRepository}; use crate::{PostgresPoolConfig, PostgresPoolFactory}; #[tokio::test] @@ -3566,6 +3779,16 @@ mod tests { let _ = repository.pool(); } + #[test] + fn provider_catalog_negative_security_timestamps_fail_closed() { + assert!(optional_u64(Some(-1), "provider_api_keys.expires_at").is_err()); + assert_eq!( + optional_u64(None, "provider_api_keys.expires_at") + .expect("SQL NULL should remain optional"), + None + ); + } + #[test] fn key_queries_include_usage_totals() { for sql in [ @@ -3634,7 +3857,7 @@ mod tests { ); assert!(source.contains("QueryBuilder::::new(select_prefix_for_in(")); assert!(source.contains(".bind(&key.allow_auth_channel_mismatch_formats)")); - assert!(source.contains("row.try_get(\"allow_auth_channel_mismatch_formats\").ok()")); + assert!(source.contains("row_get(row, \"allow_auth_channel_mismatch_formats\")?")); } #[test] @@ -3662,13 +3885,141 @@ mod tests { } #[test] - fn runtime_metadata_cas_compares_only_the_requested_namespace() { - let sql = super::KEY_RUNTIME_METADATA_CAS_SQL.to_ascii_lowercase(); - assert!(sql.contains("upstream_metadata, '{}'::jsonb) -> $2")); - assert!(sql.contains("jsonb_typeof(coalesce(upstream_metadata, '{}'::jsonb)) = 'object'")); - assert!(sql.contains("is not distinct from $6::jsonb")); - assert!(sql.contains("status_snapshot::jsonb")); - assert!(!sql.contains("is_active")); + fn runtime_metadata_cas_locks_only_the_requested_namespace() { + let lock_sql = super::KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL.to_ascii_lowercase(); + let update_sql = super::KEY_RUNTIME_METADATA_UPDATE_SQL.to_ascii_lowercase(); + + assert!(lock_sql.contains("upstream_metadata, '{}'::jsonb) -> $2")); + assert!(lock_sql.contains("upstream_metadata, '{}'::jsonb) ? $2")); + assert!(lock_sql.contains("for update")); + assert!(update_sql + .contains("jsonb_typeof(coalesce(upstream_metadata, '{}'::jsonb)) = 'object'")); + assert!(update_sql.contains("status_snapshot::jsonb")); + assert!(!update_sql.contains("is_active")); + } + + #[test] + fn runtime_metadata_namespace_cas_distinguishes_missing_from_json_null() { + assert!(super::runtime_metadata_namespace_matches( + true, false, None, None, + )); + assert!(!super::runtime_metadata_namespace_matches( + true, + true, + Some(&serde_json::Value::Null), + None, + )); + assert!(super::runtime_metadata_namespace_matches( + true, + true, + Some(&serde_json::Value::Null), + Some(&serde_json::Value::Null), + )); + assert!(!super::runtime_metadata_namespace_matches( + false, false, None, None, + )); + } + + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] + async fn live_runtime_metadata_cas_handles_high_precision_jsonb_numbers() { + let database_url = std::env::var("AETHER_TEST_DATABASE_URL") + .expect("AETHER_TEST_DATABASE_URL must point at the test database"); + let factory = PostgresPoolFactory::new(PostgresPoolConfig { + database_url, + min_connections: 1, + max_connections: 2, + acquire_timeout_ms: 10_000, + idle_timeout_ms: 30_000, + max_lifetime_ms: 60_000, + statement_cache_capacity: 64, + require_ssl: false, + }) + .expect("factory should build"); + let repository = SqlxProviderCatalogReadRepository::new( + factory.connect_lazy().expect("lazy pool should build"), + ); + crate::run_migrations(repository.pool()) + .await + .expect("test database migrations should succeed"); + + let suffix = uuid::Uuid::new_v4().simple().to_string(); + let provider_id = uuid::Uuid::new_v4().to_string(); + let key_id = uuid::Uuid::new_v4().to_string(); + let provider_name = format!("provider-metadata-cas-{suffix}"); + let key_name = format!("key-metadata-cas-{suffix}"); + sqlx::query( + "INSERT INTO providers (id, name, provider_type) VALUES ($1, $2, 'antigravity')", + ) + .bind(&provider_id) + .bind(&provider_name) + .execute(repository.pool()) + .await + .expect("provider fixture should insert"); + sqlx::query( + r#" +INSERT INTO provider_api_keys ( + id, name, provider_id, total_tokens, total_cost_usd, upstream_metadata +) +VALUES ($1, $2, $3, 0, 0, $4::jsonb) +"#, + ) + .bind(&key_id) + .bind(&key_name) + .bind(&provider_id) + .bind(r#"{"antigravity":{"used_percent":0.123456789012345678901234567890}}"#) + .execute(repository.pool()) + .await + .expect("provider key fixture should insert"); + + let observed = sqlx::query_scalar::<_, serde_json::Value>( + "SELECT upstream_metadata -> 'antigravity' FROM provider_api_keys WHERE id = $1", + ) + .bind(&key_id) + .fetch_one(repository.pool()) + .await + .expect("metadata namespace should load"); + assert_ne!( + serde_json::to_string(&observed).expect("metadata should serialize"), + r#"{"used_percent":0.123456789012345678901234567890}"#, + "the fixture must exercise precision loss in serde_json's default number representation", + ); + + let updated = repository + .update_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate { + key_id: key_id.clone(), + namespace: "antigravity".to_string(), + expected_upstream_metadata_value: Some(observed), + upstream_metadata_value: json!({"used_percent": 12.5}), + status_snapshot_patch: json!({"quota": {"used_percent": 12.5}}), + updated_at_unix_secs: Some(1_700_000_000), + }) + .await + .expect("runtime metadata CAS should execute"); + assert!( + updated, + "matching metadata must not report a false CAS conflict" + ); + + let stored = sqlx::query_scalar::<_, serde_json::Value>( + "SELECT upstream_metadata -> 'antigravity' FROM provider_api_keys WHERE id = $1", + ) + .bind(&key_id) + .fetch_one(repository.pool()) + .await + .expect("updated metadata namespace should load"); + assert_eq!(stored, json!({"used_percent": 12.5})); + + sqlx::query("DELETE FROM provider_api_keys WHERE id = $1") + .bind(&key_id) + .execute(repository.pool()) + .await + .expect("provider key fixture should delete"); + sqlx::query("DELETE FROM providers WHERE id = $1") + .bind(&provider_id) + .execute(repository.pool()) + .await + .expect("provider fixture should delete"); } #[test] @@ -3699,6 +4050,13 @@ mod tests { assert!(sql.contains("auth_config is not distinct from $6")); } + #[test] + fn credential_cas_migrates_legacy_encrypted_key_with_null_safe_fence() { + let source = include_str!("provider_catalog.rs"); + assert!(source.contains("SET api_key = $5, encrypted_key = NULL, auth_config = $6")); + assert!(source.contains("AND COALESCE(api_key, encrypted_key) IS NOT DISTINCT FROM $3")); + } + #[test] fn admin_credential_cas_has_atomic_rotation_guards() { let source = include_str!("provider_catalog.rs"); diff --git a/crates/aether-data/adapters/postgres/src/proxy_nodes.rs b/crates/aether-data/adapters/postgres/src/proxy_nodes.rs index e19bf919f..681c5b7b0 100644 --- a/crates/aether-data/adapters/postgres/src/proxy_nodes.rs +++ b/crates/aether-data/adapters/postgres/src/proxy_nodes.rs @@ -5,14 +5,14 @@ use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row}; use aether_data_contracts::repository::proxy_nodes::{ bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, - normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security, - reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation, - ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary, - ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation, - ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation, - ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent, - StoredProxyNodeMetricsBucket, TunnelErrorEventRecord, TunnelMetricsSample, - PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, + merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata, + normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, + ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, + ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository, + ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, + ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, + StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelErrorEventRecord, + TunnelMetricsSample, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_limit, WhereClause}; @@ -44,6 +44,7 @@ fn log_reported_tunnel_error_event( const FIND_PROXY_NODE_SQL: &str = r#" SELECT id, + tunnel_generation, name, ip, port, @@ -118,17 +119,58 @@ SET heartbeat_interval = COALESCE($2, heartbeat_interval), active_connections = COALESCE($3, active_connections), avg_latency_ms = COALESCE($4, avg_latency_ms), - proxy_metadata = COALESCE($5::json, proxy_metadata), - total_requests = total_requests + GREATEST(COALESCE($6, 0), 0), - failed_requests = failed_requests + GREATEST(COALESCE($7, 0), 0), - dns_failures = dns_failures + GREATEST(COALESCE($8, 0), 0), - stream_errors = stream_errors + GREATEST(COALESCE($9, 0), 0) + total_requests = total_requests + GREATEST(COALESCE($5, 0), 0), + failed_requests = failed_requests + GREATEST(COALESCE($6, 0), 0), + dns_failures = dns_failures + GREATEST(COALESCE($7, 0), 0), + stream_errors = stream_errors + GREATEST(COALESCE($8, 0), 0) WHERE id = $1 + AND tunnel_mode = TRUE + AND tunnel_generation = $9 +"#; + +const CAS_HEARTBEAT_PROXY_METADATA_SQL: &str = r#" +UPDATE proxy_nodes +SET proxy_metadata = $2::json, updated_at = NOW() +WHERE id = $1 + AND tunnel_generation = $3 + AND proxy_metadata::jsonb IS NOT DISTINCT FROM $4::jsonb +"#; + +const UPDATE_TUNNEL_STATUS_SQL: &str = r#" +UPDATE proxy_nodes +SET + tunnel_connected = $2, + active_connections = CASE + WHEN $2 THEN active_connections + ELSE 0 + END, + tunnel_connected_at = CASE + WHEN $3::double precision IS NULL THEN NOW() + ELSE TO_TIMESTAMP($3::double precision) + END, + status = CASE + WHEN $2 THEN 'online'::proxynodestatus + ELSE 'offline'::proxynodestatus + END, + updated_at = CASE + WHEN $3::double precision IS NULL THEN NOW() + ELSE TO_TIMESTAMP($3::double precision) + END +WHERE id = $1 + AND tunnel_generation = $4 + AND ( + tunnel_connected_at IS NULL + OR tunnel_connected_at <= CASE + WHEN $3::double precision IS NULL THEN NOW() + ELSE TO_TIMESTAMP($3::double precision) + END + ) "#; const FIND_EXISTING_TUNNEL_NODE_SQL: &str = r#" SELECT id, + tunnel_generation, name, ip, port, @@ -184,7 +226,8 @@ INSERT INTO proxy_nodes ( estimated_max_concurrency, tunnel_mode, tunnel_connected, - proxy_metadata + proxy_metadata, + tunnel_generation ) VALUES ( $1, @@ -203,7 +246,8 @@ VALUES ( $12, $13, FALSE, - $14::json + $14::json, + $15 ) "#; @@ -253,7 +297,8 @@ INSERT INTO proxy_nodes ( total_requests, tunnel_mode, tunnel_connected, - config_version + config_version, + tunnel_generation ) VALUES ( $1, @@ -273,7 +318,8 @@ VALUES ( 0, FALSE, FALSE, - 0 + 0, + $10 ) "#; @@ -311,6 +357,7 @@ SET updated_at = NOW() WHERE id = $1 AND is_manual = TRUE + AND tunnel_generation = $9 "#; const RECORD_PROXY_NODE_TRAFFIC_SQL: &str = r#" @@ -323,6 +370,7 @@ SET updated_at = NOW() WHERE id = $1 AND is_manual = TRUE + AND tunnel_generation = $6 "#; const UNREGISTER_PROXY_NODE_SQL: &str = r#" @@ -333,11 +381,56 @@ SET tunnel_connected_at = NOW(), updated_at = NOW() WHERE id = $1 + AND tunnel_generation = $2 "#; const DELETE_PROXY_NODE_SQL: &str = r#" DELETE FROM proxy_nodes WHERE id = $1 + AND tunnel_generation = $2 +"#; + +// Run after the parent delete commits so delete never waits on an outbox row +// already claimed by the flusher (which acquires locks in the opposite order). +const RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL: &str = r#" +DELETE FROM usage_counter_deltas +WHERE kind = 'proxy_node' + AND target_id = $1 + AND target_tunnel_generation = $2 + AND processed_at IS NULL +"#; + +const DELETE_PROXY_NODE_EVENTS_SQL: &str = r#" +DELETE FROM proxy_node_events +WHERE node_id = $1 + AND EXISTS ( + SELECT 1 + FROM proxy_nodes + WHERE id = $1 + AND tunnel_generation = $2 + ) +"#; + +const DELETE_PROXY_NODE_METRICS_1M_SQL: &str = r#" +DELETE FROM proxy_node_metrics_1m +WHERE node_id = $1 + AND EXISTS ( + SELECT 1 + FROM proxy_nodes + WHERE id = $1 + AND tunnel_generation = $2 + ) +"#; + +const DELETE_PROXY_NODE_METRICS_1H_SQL: &str = r#" +DELETE FROM proxy_node_metrics_1h +WHERE node_id = $1 + AND EXISTS ( + SELECT 1 + FROM proxy_nodes + WHERE id = $1 + AND tunnel_generation = $2 + ) "#; const UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL: &str = r#" @@ -348,6 +441,9 @@ SET config_version = config_version + 1, updated_at = NOW() WHERE id = $1 + AND tunnel_generation = $4 + AND config_version = $5 + AND is_manual = FALSE "#; const RESET_STALE_TUNNEL_STATUSES_SQL: &str = r#" @@ -372,11 +468,12 @@ SET updated_at = NOW() WHERE id = $4 AND is_manual = TRUE + AND tunnel_generation = $5 "#; const INSERT_PROXY_NODE_EVENT_SQL: &str = r#" INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at) -VALUES ( +SELECT $1, $2, $3, @@ -385,7 +482,8 @@ VALUES ( WHEN $5::double precision IS NULL THEN NOW() ELSE TO_TIMESTAMP($5::double precision) END -) +FROM proxy_nodes +WHERE id = $1 AND ($6::text IS NULL OR tunnel_generation = $6) "#; const UPSERT_PROXY_NODE_METRICS_1M_SQL: &str = r#" @@ -406,7 +504,9 @@ INSERT INTO proxy_node_metrics_1m ( ws_in_frames_delta, ws_out_frames_delta ) -VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15) +SELECT $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15 +FROM proxy_nodes +WHERE id = $1 AND ($16::text IS NULL OR tunnel_generation = $16) ON CONFLICT (node_id, bucket_start_unix_secs) DO UPDATE SET samples = proxy_node_metrics_1m.samples + EXCLUDED.samples, uptime_samples = proxy_node_metrics_1m.uptime_samples + EXCLUDED.uptime_samples, @@ -441,7 +541,9 @@ INSERT INTO proxy_node_metrics_1h ( ws_in_frames_delta, ws_out_frames_delta ) -VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15) +SELECT $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15 +FROM proxy_nodes +WHERE id = $1 AND ($16::text IS NULL OR tunnel_generation = $16) ON CONFLICT (node_id, bucket_start_unix_secs) DO UPDATE SET samples = proxy_node_metrics_1h.samples + EXCLUDED.samples, uptime_samples = proxy_node_metrics_1h.uptime_samples + EXCLUDED.uptime_samples, @@ -571,6 +673,12 @@ impl SqlxProxyNodeRepository { } fn row_to_stored(row: &PgRow) -> Result { + let tunnel_generation: String = row.try_get("tunnel_generation").map_postgres_err()?; + if tunnel_generation.trim().is_empty() { + return Err(DataLayerError::UnexpectedValue( + "proxy_nodes.tunnel_generation must not be empty".to_string(), + )); + } Ok(StoredProxyNode::new( row.try_get("id").map_postgres_err()?, row.try_get("name").map_postgres_err()?, @@ -588,6 +696,7 @@ impl SqlxProxyNodeRepository { row.try_get("tunnel_connected").map_postgres_err()?, row.try_get("config_version").map_postgres_err()?, )? + .with_tunnel_generation(tunnel_generation) .with_manual_proxy_fields( row.try_get("proxy_url").map_postgres_err()?, row.try_get("proxy_username").map_postgres_err()?, @@ -676,6 +785,7 @@ impl SqlxProxyNodeRepository { async fn insert_event( &self, node_id: &str, + expected_tunnel_generation: Option<&str>, event_type: &str, detail: Option<&str>, event_metadata: Option<&serde_json::Value>, @@ -687,6 +797,7 @@ impl SqlxProxyNodeRepository { .bind(detail) .bind(event_metadata) .bind(created_at_unix_secs.map(|value| value as f64)) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_postgres_err()?; @@ -697,6 +808,7 @@ impl SqlxProxyNodeRepository { &self, step: ProxyNodeMetricsStep, node_id: &str, + expected_tunnel_generation: Option<&str>, bucket_start: u64, sample: &TunnelMetricsSample, ) -> Result<(), DataLayerError> { @@ -720,6 +832,7 @@ impl SqlxProxyNodeRepository { .bind(sample.ws_out_bytes_delta) .bind(sample.ws_in_frames_delta) .bind(sample.ws_out_frames_delta) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_postgres_err()?; @@ -1007,6 +1120,56 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { Ok(result.rows_affected() as usize) } + async fn compare_and_set_proxy_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_password = $2, updated_at = NOW() +WHERE id = $1 AND proxy_password = $3 +"#, + ) + .bind(node_id) + .bind(replacement) + .bind(expected) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_set_proxy_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result { + let expected = serde_json::to_string(expected).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let replacement = serde_json::to_string(replacement).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_metadata = $2::json, updated_at = NOW() +WHERE id = $1 AND proxy_metadata::jsonb = $3::jsonb +"#, + ) + .bind(node_id) + .bind(replacement) + .bind(expected) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) + } + async fn create_manual_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -1028,7 +1191,12 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { )); } - let node_id = uuid::Uuid::new_v4().to_string(); + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + let tunnel_generation = uuid::Uuid::new_v4().to_string(); + if let Some((ip, port)) = proxy_node_id_owner_locked(&mut tx, &node_id).await? { + return Err(proxy_node_id_in_use_error(&node_id, &ip, port)); + } sqlx::query(INSERT_MANUAL_PROXY_NODE_SQL) .bind(&node_id) .bind(&mutation.name) @@ -1039,6 +1207,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { .bind(mutation.proxy_username.as_deref()) .bind(mutation.proxy_password.as_deref()) .bind(mutation.registered_by.as_deref()) + .bind(&tunnel_generation) .execute(&mut *tx) .await .map_postgres_err()?; @@ -1086,7 +1255,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { )); } - sqlx::query(UPDATE_MANUAL_PROXY_NODE_SQL) + let result = sqlx::query(UPDATE_MANUAL_PROXY_NODE_SQL) .bind(&mutation.node_id) .bind(mutation.name.as_deref()) .bind(mutation.ip.as_deref()) @@ -1095,10 +1264,16 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { .bind(mutation.proxy_url.as_deref()) .bind(mutation.proxy_username.as_deref()) .bind(mutation.proxy_password.as_deref()) + .bind(&existing.tunnel_generation) .execute(&mut *tx) .await .map_postgres_err()?; + if result.rows_affected() == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + tx.commit().await.map_err(postgres_error)?; self.find_proxy_node(&mutation.node_id).await } @@ -1107,9 +1282,12 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { &self, mutation: &ProxyNodeRegistrationMutation, ) -> Result { - let normalized_proxy_metadata = normalize_proxy_metadata( - mutation.proxy_metadata.as_ref(), - mutation.proxy_version.as_deref(), + let normalized_proxy_metadata = merge_proxy_metadata_for_registration( + None, + normalize_proxy_metadata( + mutation.proxy_metadata.as_ref(), + mutation.proxy_version.as_deref(), + ), ); let lock_key = Self::registration_lock_key(&mutation.ip, mutation.port); let mut tx = self.pool.begin().await.map_postgres_err()?; @@ -1129,6 +1307,18 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { let node_id = if let Some(row) = existing.as_ref() { let existing = Self::row_to_stored(row)?; + if let Some(requested_id) = requested_proxy_node_id(mutation.node_id.as_deref())? { + if requested_id != existing.id { + return Err(proxy_node_registration_identity_error( + &requested_id, + &existing.id, + )); + } + } + let registration_proxy_metadata = merge_proxy_metadata_for_registration( + existing.proxy_metadata.as_ref(), + normalized_proxy_metadata, + ); sqlx::query(UPDATE_PROXY_NODE_REGISTRATION_SQL) .bind(&existing.id) .bind(&mutation.name) @@ -1143,13 +1333,18 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { .bind(mutation.hardware_info.as_ref()) .bind(mutation.estimated_max_concurrency) .bind(mutation.tunnel_mode) - .bind(normalized_proxy_metadata.as_ref()) + .bind(registration_proxy_metadata.as_ref()) .execute(&mut *tx) .await .map_postgres_err()?; existing.id } else { - let node_id = uuid::Uuid::new_v4().to_string(); + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + let tunnel_generation = uuid::Uuid::new_v4().to_string(); + if let Some((ip, port)) = proxy_node_id_owner_locked(&mut tx, &node_id).await? { + return Err(proxy_node_id_in_use_error(&node_id, &ip, port)); + } sqlx::query(INSERT_PROXY_NODE_SQL) .bind(&node_id) .bind(&mutation.name) @@ -1165,6 +1360,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { .bind(mutation.estimated_max_concurrency) .bind(mutation.tunnel_mode) .bind(normalized_proxy_metadata.as_ref()) + .bind(&tunnel_generation) .execute(&mut *tx) .await .map_postgres_err()?; @@ -1185,6 +1381,13 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { let Some(existing) = existing else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != existing.tunnel_generation) + { + return Ok(None); + } if !existing.tunnel_mode { return Err(DataLayerError::InvalidInput( "non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode" @@ -1192,47 +1395,106 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { )); } - let normalized_proxy_metadata = normalize_proxy_metadata( + let tunnel_generation = existing.tunnel_generation.clone(); + let has_proxy_metadata_update = normalize_heartbeat_proxy_metadata( + None, mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), - ); - let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security( - existing.proxy_metadata.as_ref(), - normalized_proxy_metadata, - ); + ) + .is_some(); - sqlx::query(APPLY_HEARTBEAT_SQL) + let result = sqlx::query(APPLY_HEARTBEAT_SQL) .bind(&mutation.node_id) .bind(mutation.heartbeat_interval) .bind(mutation.active_connections) .bind(mutation.avg_latency_ms) - .bind(normalized_proxy_metadata) .bind(mutation.total_requests_delta) .bind(mutation.failed_requests_delta) .bind(mutation.dns_failures_delta) .bind(mutation.stream_errors_delta) + .bind(&tunnel_generation) .execute(&self.pool) .await .map_postgres_err()?; - - let updated = self.find_proxy_node(&mutation.node_id).await?; - let Some(updated) = updated else { + if result.rows_affected() == 0 { return Ok(None); + } + + let mut updated = None; + let mut tunnel_metrics_sample = None; + if has_proxy_metadata_update { + for _ in 0..8 { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + let Some(replacement) = normalize_heartbeat_proxy_metadata( + current.proxy_metadata.as_ref(), + mutation.proxy_metadata.as_ref(), + mutation.proxy_version.as_deref(), + ) else { + break; + }; + if current.proxy_metadata.as_ref() == Some(&replacement) { + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + Some(&replacement), + current.active_connections, + current.tunnel_connected, + ); + updated = Some(current); + break; + } + + let result = sqlx::query(CAS_HEARTBEAT_PROXY_METADATA_SQL) + .bind(&mutation.node_id) + .bind(&replacement) + .bind(&tunnel_generation) + .bind(current.proxy_metadata.as_ref()) + .execute(&self.pool) + .await + .map_postgres_err()?; + if result.rows_affected() == 0 { + continue; + } + let Some(after_cas) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if after_cas.tunnel_generation != tunnel_generation { + return Ok(None); + } + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + after_cas.proxy_metadata.as_ref(), + after_cas.active_connections, + after_cas.tunnel_connected, + ); + updated = Some(after_cas); + break; + } + } + let updated = if let Some(updated) = updated { + updated + } else { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + current }; let now_unix_secs = updated .last_heartbeat_at_unix_secs .unwrap_or_else(|| chrono::Utc::now().timestamp().max(0) as u64); - let tunnel_metrics_sample = build_tunnel_metrics_sample( - existing.proxy_metadata.as_ref(), - updated.proxy_metadata.as_ref(), - updated.active_connections, - updated.tunnel_connected, - ); if let Some(sample) = tunnel_metrics_sample.as_ref() { self.upsert_metrics_bucket( ProxyNodeMetricsStep::OneMinute, &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute), sample, ) @@ -1240,6 +1502,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { self.upsert_metrics_bucket( ProxyNodeMetricsStep::OneHour, &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour), sample, ) @@ -1261,6 +1524,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { }); self.insert_event( &updated.id, + Some(tunnel_generation.as_str()), PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, Some(detail.as_str()), Some(&event_metadata), @@ -1282,6 +1546,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { return self .update_remote_config(&ProxyNodeRemoteConfigMutation { node_id: mutation.node_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), node_name: None, allowed_ports: None, log_level: None, @@ -1299,17 +1564,38 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { &self, mutation: &ProxyNodeTrafficMutation, ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query(&format!("{FIND_PROXY_NODE_SQL} FOR UPDATE")) + .bind(&mutation.node_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let generation = Self::row_to_stored(&row)?.tunnel_generation; + let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + if expected_generation != generation { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } let result = sqlx::query(RECORD_PROXY_NODE_TRAFFIC_SQL) .bind(&mutation.node_id) .bind(mutation.total_requests_delta) .bind(mutation.failed_requests_delta) .bind(mutation.dns_failures_delta) .bind(mutation.stream_errors_delta) - .execute(&self.pool) + .bind(expected_generation) + .execute(&mut *tx) .await .map_postgres_err()?; - - Ok(result.rows_affected() > 0) + let applied = result.rows_affected() > 0; + tx.commit().await.map_err(postgres_error)?; + Ok(applied) } async fn update_tunnel_status( @@ -1320,6 +1606,14 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { let Some(existing) = existing else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != existing.tunnel_generation) + { + return Ok(None); + } + let tunnel_generation = existing.tunnel_generation.clone(); let observed_at_unix_secs = mutation.observed_at_unix_secs; let event_type = if mutation.connected { @@ -1335,100 +1629,136 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository { }); let mut tx = self.pool.begin().await.map_postgres_err()?; - - if existing - .tunnel_connected_at_unix_secs - .zip(observed_at_unix_secs) - .is_some_and(|(last_transition, observed_at)| observed_at < last_transition) - { - sqlx::query(INSERT_PROXY_NODE_EVENT_SQL) - .bind(&mutation.node_id) - .bind(event_type) - .bind(format!("[stale_ignored] {event_detail}")) - .bind(None::) - .bind(None::) - .execute(&mut *tx) - .await - .map_postgres_err()?; - tx.commit().await.map_err(postgres_error)?; - return self.find_proxy_node(&mutation.node_id).await; - } - - sqlx::query( - r#" -UPDATE proxy_nodes -SET - tunnel_connected = $2, - active_connections = CASE - WHEN $2 THEN active_connections - ELSE 0 - END, - tunnel_connected_at = CASE - WHEN $3::double precision IS NULL THEN NOW() - ELSE TO_TIMESTAMP($3::double precision) - END, - status = CASE - WHEN $2 THEN 'online'::proxynodestatus - ELSE 'offline'::proxynodestatus - END, - updated_at = CASE - WHEN $3::double precision IS NULL THEN NOW() - ELSE TO_TIMESTAMP($3::double precision) - END -WHERE id = $1 -"#, - ) - .bind(&mutation.node_id) - .bind(mutation.connected) - .bind(observed_at_unix_secs.map(|value| value as f64)) - .execute(&mut *tx) - .await - .map_postgres_err()?; + let result = sqlx::query(UPDATE_TUNNEL_STATUS_SQL) + .bind(&mutation.node_id) + .bind(mutation.connected) + .bind(observed_at_unix_secs.map(|value| value as f64)) + .bind(&tunnel_generation) + .execute(&mut *tx) + .await + .map_postgres_err()?; + let stale = result.rows_affected() == 0; + let persisted_detail = if stale { + format!("[stale_ignored] {event_detail}") + } else { + event_detail + }; sqlx::query(INSERT_PROXY_NODE_EVENT_SQL) .bind(&mutation.node_id) .bind(event_type) - .bind(event_detail) + .bind(persisted_detail) .bind(None::) - .bind(observed_at_unix_secs.map(|value| value as f64)) + .bind( + (!stale) + .then_some(observed_at_unix_secs) + .flatten() + .map(|value| value as f64), + ) + .bind(&tunnel_generation) .execute(&mut *tx) .await .map_postgres_err()?; tx.commit().await.map_err(postgres_error)?; - self.find_proxy_node(&mutation.node_id).await + let current = self.find_proxy_node(&mutation.node_id).await?; + Ok(current.filter(|node| node.tunnel_generation == tunnel_generation)) } async fn unregister_node( &self, node_id: &str, ) -> Result, DataLayerError> { - let existing = self.find_proxy_node(node_id).await?; - let Some(existing) = existing else { - return Ok(None); - }; - - sqlx::query(UNREGISTER_PROXY_NODE_SQL) + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query(&format!("{FIND_PROXY_NODE_SQL} FOR UPDATE")) .bind(node_id) - .execute(&self.pool) + .fetch_optional(&mut *tx) .await .map_postgres_err()?; - - self.find_proxy_node(&existing.id).await + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + }; + let generation = Self::row_to_stored(&row)?.tunnel_generation; + sqlx::query(UNREGISTER_PROXY_NODE_SQL) + .bind(node_id) + .bind(&generation) + .execute(&mut *tx) + .await + .map_postgres_err()?; + let updated_row = sqlx::query(&format!("{FIND_PROXY_NODE_SQL} FOR UPDATE")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_postgres_err()?; + let updated = Self::row_to_stored(&updated_row)?; + tx.commit().await.map_err(postgres_error)?; + Ok(Some(updated)) } async fn delete_node(&self, node_id: &str) -> Result, DataLayerError> { - let existing = self.find_proxy_node(node_id).await?; - let Some(existing) = existing else { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query(&format!("{FIND_PROXY_NODE_SQL} FOR UPDATE")) + .bind(node_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; return Ok(None); }; + // Reuse the canonical projection so the value returned to callers is + // identical to find_proxy_node while the row remains locked. + let existing = Self::row_to_stored(&row)?; + let generation = existing.tunnel_generation.as_str(); - sqlx::query(DELETE_PROXY_NODE_SQL) + // Keep child cleanup in this transaction and bind it to the locked + // generation. The FK cascade is still a final backstop for deployments + // that have the PostgreSQL constraints enabled. + sqlx::query(DELETE_PROXY_NODE_EVENTS_SQL) .bind(node_id) - .execute(&self.pool) + .bind(generation) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query(DELETE_PROXY_NODE_METRICS_1M_SQL) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query(DELETE_PROXY_NODE_METRICS_1H_SQL) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) .await .map_postgres_err()?; + let deleted = sqlx::query(DELETE_PROXY_NODE_SQL) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if deleted.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + tx.commit().await.map_err(postgres_error)?; + if let Err(error) = sqlx::query(RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL) + .bind(node_id) + .bind(generation) + .execute(&self.pool) + .await + .map_postgres_err() + { + tracing::warn!( + node_id = %node_id, + tunnel_generation = %generation, + error = ?error, + "failed to retire deleted proxy node counter rows" + ); + } Ok(Some(existing)) } @@ -1436,27 +1766,46 @@ WHERE id = $1 &self, mutation: &ProxyNodeRemoteConfigMutation, ) -> Result, DataLayerError> { - let existing = self.find_proxy_node(&mutation.node_id).await?; - let Some(existing) = existing else { - return Ok(None); - }; - if existing.is_manual { - return Err(DataLayerError::InvalidInput( - "手动节点不支持远程配置下发".to_string(), - )); + for _ in 0..8 { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != existing.tunnel_generation) + { + return Ok(None); + } + if existing.is_manual { + return Err(DataLayerError::InvalidInput( + "手动节点不支持远程配置下发".to_string(), + )); + } + + let tunnel_generation = existing.tunnel_generation.clone(); + let remote_config = + Self::normalize_remote_config(mutation, existing.remote_config.as_ref()); + let result = sqlx::query(UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL) + .bind(&mutation.node_id) + .bind(mutation.node_name.as_deref()) + .bind(remote_config.as_ref()) + .bind(&tunnel_generation) + .bind(existing.config_version) + .execute(&self.pool) + .await + .map_postgres_err()?; + if result.rows_affected() == 0 { + continue; + } + + let current = self.find_proxy_node(&mutation.node_id).await?; + return Ok(current.filter(|node| node.tunnel_generation == tunnel_generation)); } - let remote_config = - Self::normalize_remote_config(mutation, existing.remote_config.as_ref()); - sqlx::query(UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL) - .bind(&mutation.node_id) - .bind(mutation.node_name.as_deref()) - .bind(remote_config.as_ref()) - .execute(&self.pool) - .await - .map_postgres_err()?; - - self.find_proxy_node(&mutation.node_id).await + Err(DataLayerError::UnexpectedValue( + "proxy node remote config changed during every CAS retry".to_string(), + )) } async fn increment_manual_node_requests( @@ -1466,14 +1815,27 @@ WHERE id = $1 failed_delta: i64, latency_ms: Option, ) -> Result<(), DataLayerError> { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query(&format!("{FIND_PROXY_NODE_SQL} FOR UPDATE")) + .bind(node_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(()); + }; + let generation = Self::row_to_stored(&row)?.tunnel_generation; sqlx::query(INCREMENT_MANUAL_PROXY_NODE_REQUESTS_SQL) .bind(total_delta) .bind(failed_delta) .bind(latency_ms) .bind(node_id) - .execute(&self.pool) + .bind(&generation) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_err(postgres_error)?; Ok(()) } @@ -1535,18 +1897,63 @@ WHERE metrics.node_id = expired.node_id } } +fn requested_proxy_node_id(value: Option<&str>) -> Result, DataLayerError> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_empty() || value.trim() != value { + return Err(DataLayerError::InvalidInput( + "proxy node id must be non-empty and unpadded".to_string(), + )); + } + Ok(Some(value.to_string())) +} + +async fn proxy_node_id_owner_locked( + tx: &mut sqlx::Transaction<'_, Postgres>, + node_id: &str, +) -> Result, DataLayerError> { + let row = sqlx::query("SELECT ip, port FROM proxy_nodes WHERE id = $1 FOR UPDATE") + .bind(node_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(None); + }; + Ok(Some(( + row.try_get::("ip").map_postgres_err()?, + row.try_get::("port").map_postgres_err()?, + ))) +} + +fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node registration identity changed: requested {requested_id}, existing {existing_id}" + )) +} + +fn proxy_node_id_in_use_error(node_id: &str, ip: &str, port: i32) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node id is already in use: {node_id} ({ip}:{port})" + )) +} + #[cfg(test)] mod tests { #[test] fn proxy_node_sql_uses_json_casts_for_json_columns() { - assert!(super::APPLY_HEARTBEAT_SQL - .contains("proxy_metadata = COALESCE($5::json, proxy_metadata)")); + assert!(!super::APPLY_HEARTBEAT_SQL.contains("proxy_metadata")); + assert!(super::CAS_HEARTBEAT_PROXY_METADATA_SQL.contains("proxy_metadata = $2::json")); + assert!(super::CAS_HEARTBEAT_PROXY_METADATA_SQL + .contains("proxy_metadata::jsonb IS NOT DISTINCT FROM $4::jsonb")); assert!(super::INSERT_PROXY_NODE_SQL - .contains("\n $11::json,\n $12,\n $13,\n FALSE,\n $14::json\n")); + .contains("\n $11::json,\n $12,\n $13,\n FALSE,\n $14::json,\n $15\n")); assert!(super::UPDATE_PROXY_NODE_REGISTRATION_SQL .contains("hardware_info = COALESCE($11::json, hardware_info)")); assert!(super::UPDATE_PROXY_NODE_REGISTRATION_SQL .contains("proxy_metadata = COALESCE($14::json, proxy_metadata)")); + assert!(!super::UPDATE_PROXY_NODE_REGISTRATION_SQL.contains("tunnel_generation")); assert!(super::UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL.contains("remote_config = $3::json")); } @@ -1557,4 +1964,15 @@ mod tests { assert!(!super::UPDATE_PROXY_NODE_REGISTRATION_SQL.contains("::jsonb")); assert!(!super::UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL.contains("::jsonb")); } + + #[test] + fn tunnel_status_sql_rejects_out_of_order_transitions_atomically() { + assert!(super::UPDATE_TUNNEL_STATUS_SQL.contains("tunnel_connected_at <= CASE")); + assert!(super::UPDATE_TUNNEL_STATUS_SQL.contains("tunnel_generation = $4")); + assert!(super::APPLY_HEARTBEAT_SQL.contains("total_requests = total_requests + GREATEST")); + assert!(super::APPLY_HEARTBEAT_SQL.contains("tunnel_generation = $9")); + assert!(!super::APPLY_HEARTBEAT_SQL.contains("remote_config =")); + assert!(!super::APPLY_HEARTBEAT_SQL.contains("config_version =")); + assert!(super::UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL.contains("config_version = $5")); + } } diff --git a/crates/aether-data/adapters/postgres/src/routing_profiles.rs b/crates/aether-data/adapters/postgres/src/routing_profiles.rs index e34003efc..ab6a8b079 100644 --- a/crates/aether-data/adapters/postgres/src/routing_profiles.rs +++ b/crates/aether-data/adapters/postgres/src/routing_profiles.rs @@ -22,6 +22,7 @@ SELECT description, enabled, is_system_default, + sort_order, config_json, version, created_at, @@ -68,7 +69,9 @@ impl PostgresRoutingGroupRepository { #[async_trait] impl RoutingGroupReadRepository for PostgresRoutingGroupRepository { async fn list_routing_groups(&self) -> Result, DataLayerError> { - let sql = format!("{ROUTING_GROUP_SELECT} ORDER BY name ASC, id ASC"); + let sql = format!( + "{ROUTING_GROUP_SELECT} ORDER BY enabled DESC, sort_order ASC, name ASC, id ASC" + ); let mut rows = sqlx::query(&sql).fetch(&self.pool); let mut groups = Vec::new(); while let Some(row) = rows.try_next().await.map_postgres_err()? { @@ -178,10 +181,10 @@ impl RoutingGroupWriteRepository for PostgresRoutingGroupRepository { sqlx::query( r#" INSERT INTO routing_groups ( - id, name, description, enabled, is_system_default, config_json, + id, name, description, enabled, is_system_default, sort_order, config_json, version, created_at, updated_at, published_at ) -VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) +VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11) "#, ) .bind(&group.id) @@ -189,6 +192,7 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(&group.config_json) .bind(group.version) .bind(group.created_at) @@ -238,10 +242,11 @@ SET name = $2, description = $3, enabled = $4, is_system_default = $5, - config_json = $6, - version = $7, - updated_at = $8, - published_at = $9 + sort_order = $6, + config_json = $7, + version = $8, + updated_at = $9, + published_at = $10 WHERE id = $1 "#, ) @@ -250,6 +255,7 @@ WHERE id = $1 .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(&group.config_json) .bind(group.version) .bind(group.updated_at) @@ -441,6 +447,7 @@ fn map_group_row(row: &PgRow) -> Result { description: row.try_get("description").map_postgres_err()?, enabled: row.try_get("enabled").map_postgres_err()?, is_system_default: row.try_get("is_system_default").map_postgres_err()?, + sort_order: row.try_get("sort_order").map_postgres_err()?, config_json: row.try_get("config_json").map_postgres_err()?, version: row.try_get("version").map_postgres_err()?, created_at: row.try_get("created_at").map_postgres_err()?, diff --git a/crates/aether-data/adapters/postgres/src/settlement.rs b/crates/aether-data/adapters/postgres/src/settlement.rs index f594c0413..cd457c96f 100644 --- a/crates/aether-data/adapters/postgres/src/settlement.rs +++ b/crates/aether-data/adapters/postgres/src/settlement.rs @@ -3,8 +3,12 @@ use sqlx::{PgPool, Row}; use aether_data_contracts::repository::settlement::{ finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd, - settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement, - UsageSettlementInput, SETTLEMENT_EPSILON_USD, + settlement_billing_status_for_usage_status, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation, + StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState, + UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD, }; use aether_data_contracts::DataLayerError; @@ -155,6 +159,144 @@ impl SqlxSettlementRepository { } } +fn usage_policy_cost_i64(value: u64, field: &str) -> Result { + i64::try_from(value) + .map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds the integer range"))) +} + +fn usage_policy_cost_u64(value: i64, field: &str) -> Result { + u64::try_from(value) + .map_err(|_| DataLayerError::UnexpectedValue(format!("{field} must not be negative"))) +} + +fn usage_policy_request_admission_from_postgres_row( + row: &sqlx::postgres::PgRow, +) -> Result { + let state: String = row.try_get("state").map_postgres_err()?; + Ok(StoredUsagePolicyRequestAdmission { + request_id: row.try_get("request_id").map_postgres_err()?, + subject_id: row.try_get("subject_id").map_postgres_err()?, + event_token: row.try_get("event_token").map_postgres_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at_unix_secs").map_postgres_err()?, + "usage policy request admitted_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until_unix_secs").map_postgres_err()?, + "usage policy request retain_until", + )?, + state: UsagePolicyRequestAdmissionState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy request admission state {state}" + )) + })?, + released_at_unix_secs: row + .try_get::, _>("released_at_unix_secs") + .map_postgres_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy request released_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_REQUEST_ADMISSION_POSTGRES_SQL: &str = r#" +SELECT + request_id, + subject_id, + event_token, + CAST(EXTRACT(EPOCH FROM admitted_at) AS BIGINT) AS admitted_at_unix_secs, + CAST(EXTRACT(EPOCH FROM retain_until) AS BIGINT) AS retain_until_unix_secs, + state, + CAST(EXTRACT(EPOCH FROM released_at) AS BIGINT) AS released_at_unix_secs +FROM usage_request_admissions +WHERE event_token = $1 +FOR UPDATE +"#; + +fn usage_policy_cost_reservation_from_postgres_row( + row: &sqlx::postgres::PgRow, +) -> Result { + let state: String = row.try_get("state").map_postgres_err()?; + Ok(StoredUsagePolicyCostReservation { + request_id: row.try_get("request_id").map_postgres_err()?, + subject_id: row.try_get("subject_id").map_postgres_err()?, + reservation_token: row.try_get("reservation_token").map_postgres_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at_unix_secs").map_postgres_err()?, + "usage policy admitted_at", + )?, + reserved_cost_units: usage_policy_cost_u64( + row.try_get("reserved_cost_units").map_postgres_err()?, + "usage policy reserved_cost_units", + )?, + actual_cost_units: row + .try_get::, _>("actual_cost_units") + .map_postgres_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy actual_cost_units")) + .transpose()?, + state: UsagePolicyCostReservationState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy reservation state {state}" + )) + })?, + reservation_expires_at_unix_secs: usage_policy_cost_u64( + row.try_get("reservation_expires_at_unix_secs") + .map_postgres_err()?, + "usage policy reservation_expires_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until_unix_secs").map_postgres_err()?, + "usage policy retain_until", + )?, + finalized_at_unix_secs: row + .try_get::, _>("finalized_at_unix_secs") + .map_postgres_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy finalized_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_COST_RESERVATION_POSTGRES_SQL: &str = r#" +SELECT + request_id, + subject_id, + reservation_token, + CAST(EXTRACT(EPOCH FROM admitted_at) AS BIGINT) AS admitted_at_unix_secs, + reserved_cost_units, + actual_cost_units, + state, + CAST(EXTRACT(EPOCH FROM reservation_expires_at) AS BIGINT) + AS reservation_expires_at_unix_secs, + CAST(EXTRACT(EPOCH FROM retain_until) AS BIGINT) AS retain_until_unix_secs, + CAST(EXTRACT(EPOCH FROM finalized_at) AS BIGINT) AS finalized_at_unix_secs +FROM usage_cost_reservations +WHERE reservation_token = $1 +FOR UPDATE +"#; + +async fn lock_usage_policy_subject_postgres( + tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, + subject_id: &str, +) -> Result { + let exists = sqlx::query_scalar::<_, String>( + r#" +SELECT id +FROM users +WHERE id = $1 +FOR UPDATE + "#, + ) + .bind(subject_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + .is_some(); + Ok(exists) +} + +fn usage_policy_subject_missing() -> DataLayerError { + DataLayerError::InvalidInput("usage policy subject does not exist".to_string()) +} + fn settlement_from_row( row: &sqlx::postgres::PgRow, ) -> Result { @@ -272,6 +414,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -298,15 +441,27 @@ fn daily_quota_grants_from_entitlement( entitlement_id: entitlement_id.to_string(), daily_quota_usd, usage_date, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + async fn consume_daily_quota_postgres( tx: &mut crate::PostgresTransaction, user_id: &str, @@ -315,19 +470,30 @@ async fn consume_daily_quota_postgres( wallet_available_usd: Option, wallet_can_overdraft: bool, ) -> Result { - if total_cost_usd <= 0.0 { + if !total_cost_usd.is_finite() || total_cost_usd < 0.0 { + return Err(DataLayerError::InvalidInput( + "daily quota settlement cost must be finite and non-negative".to_string(), + )); + } + if total_cost_usd == 0.0 { return Ok(DailyQuotaDebitResult::default()); } let now = chrono::Utc::now(); let entitlement_rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = $1 - AND status = 'active' - AND starts_at <= NOW() - AND expires_at > NOW() -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = $1 + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= NOW() + AND user_plan_entitlements.expires_at > NOW() +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC FOR UPDATE "#, ) @@ -340,9 +506,12 @@ FOR UPDATE let entitlement_id: String = row.try_get("id").map_postgres_err()?; let entitlements: serde_json::Value = row.try_get("entitlements_snapshot").map_postgres_err()?; + let plan_entitlements: serde_json::Value = + row.try_get("plan_entitlements_json").map_postgres_err()?; grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -369,8 +538,18 @@ WHERE user_entitlement_id = $1 .await .map_postgres_err()? .unwrap_or(0.0); + if !used.is_finite() || used < 0.0 { + return Err(DataLayerError::UnexpectedValue( + "daily quota usage ledger total is invalid".to_string(), + )); + } let remaining = (grant.daily_quota_usd - used).max(0.0); total_remaining += remaining; + if !total_remaining.is_finite() { + return Err(DataLayerError::UnexpectedValue( + "daily quota remaining total overflowed".to_string(), + )); + } grants_with_remaining.push((grant, remaining)); } @@ -421,6 +600,496 @@ ON CONFLICT (user_entitlement_id, request_id) DO NOTHING #[async_trait] impl SettlementWriteRepository for SqlxSettlementRepository { + async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result { + input.validate()?; + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + if !lock_usage_policy_subject_postgres(tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + let existing_row = + sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_POSTGRES_SQL) + .bind(&input.event_token) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if let Some(row) = existing_row.as_ref() { + let existing = usage_policy_request_admission_from_postgres_row(row)?; + if existing.request_id != input.request_id + || existing.subject_id != input.subject_id + { + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at" + .to_string(), + )); + } + sqlx::query( + r#" +UPDATE usage_request_admissions +SET retain_until = GREATEST(retain_until, TO_TIMESTAMP($2::double precision)) +WHERE event_token = $1 + "#, + ) + .bind(&input.event_token) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + return Ok(match existing.state { + UsagePolicyRequestAdmissionState::Active => { + ReserveUsagePolicyRequestOutcome::Allowed + } + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }); + } + + for (window_index, window) in input.windows.iter().enumerate() { + let used_requests = sqlx::query_scalar::<_, i64>( + r#" +SELECT COUNT(*)::BIGINT +FROM usage_request_admissions +WHERE subject_id = $1 + AND state = 'active' + AND admitted_at >= TO_TIMESTAMP($2::double precision) + AND admitted_at < TO_TIMESTAMP($3::double precision) + "#, + ) + .bind(&input.subject_id) + .bind(usage_policy_cost_i64( + window.starts_at_unix_secs, + "usage policy request window start", + )?) + .bind(usage_policy_cost_i64( + window.ends_at_unix_secs, + "usage policy request window end", + )?) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let used_requests = usage_policy_cost_u64( + used_requests, + "usage policy request used_requests", + )?; + if used_requests >= window.limit_requests { + return Ok(ReserveUsagePolicyRequestOutcome::Rejected { + window_index, + limit_requests: window.limit_requests, + used_requests, + }); + } + } + + let insert_result = sqlx::query( + r#" +INSERT INTO usage_request_admissions ( + request_id, subject_id, event_token, admitted_at, retain_until, + state, released_at, created_at +) VALUES ( + $1, $2, $3, TO_TIMESTAMP($4::double precision), + TO_TIMESTAMP($5::double precision), 'active', NULL, NOW() +) +ON CONFLICT (event_token) DO NOTHING + "#, + ) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.event_token) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy request admitted_at", + )?) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + if insert_result.rows_affected() == 1 { + return Ok(ReserveUsagePolicyRequestOutcome::Allowed); + } + + // A token can race across different subjects, which hold different subject + // locks. The unique key resolves that race; classify it explicitly here. + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_POSTGRES_SQL) + .bind(&input.event_token) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let existing = usage_policy_request_admission_from_postgres_row(&row)?; + if existing.request_id != input.request_id + || existing.subject_id != input.subject_id + { + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at" + .to_string(), + )); + } + Ok(match existing.state { + UsagePolicyRequestAdmissionState::Active => { + ReserveUsagePolicyRequestOutcome::Allowed + } + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }) + }) + }) + .await + } + + async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, DataLayerError> { + input.validate()?; + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + if !lock_usage_policy_subject_postgres(tx, &input.subject_id).await? { + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_POSTGRES_SQL) + .bind(&input.event_token) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(None); + }; + let mut admission = usage_policy_request_admission_from_postgres_row(&row)?; + if admission.request_id != input.request_id + || admission.subject_id != input.subject_id + { + return Ok(None); + } + if input.released_at_unix_secs < admission.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy released_at must not precede admitted_at".to_string(), + )); + } + if admission.state == UsagePolicyRequestAdmissionState::Active { + sqlx::query( + r#" +UPDATE usage_request_admissions +SET state = 'released', released_at = TO_TIMESTAMP($2::double precision) +WHERE event_token = $1 AND state = 'active' + "#, + ) + .bind(&input.event_token) + .bind(usage_policy_cost_i64( + input.released_at_unix_secs, + "usage policy request released_at", + )?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + admission.state = UsagePolicyRequestAdmissionState::Released; + admission.released_at_unix_secs = Some(input.released_at_unix_secs); + } + Ok(Some(admission)) + }) + }) + .await + } + + async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy request cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_request_admissions +WHERE retain_until <= TO_TIMESTAMP($1::double precision) + AND event_token IN ( + SELECT event_token + FROM usage_request_admissions + WHERE retain_until <= TO_TIMESTAMP($1::double precision) + ORDER BY retain_until, event_token + LIMIT $2 +) + "#, + ) + .bind(now) + .bind(limit) + .execute(self.tx_runner.pool()) + .await + .map_postgres_err()?; + Ok(result.rows_affected() as usize) + } + + async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result { + input.validate()?; + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + if !lock_usage_policy_subject_postgres(tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + let existing_row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_POSTGRES_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let existing = existing_row + .as_ref() + .map(usage_policy_cost_reservation_from_postgres_row) + .transpose()?; + if let Some(existing) = existing.as_ref() { + if existing.request_id != input.request_id + || existing.subject_id != input.subject_id + { + return Ok(ReserveUsagePolicyCostOutcome::Conflict); + } + if existing.state != UsagePolicyCostReservationState::Reserved { + return Ok(ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: existing.state, + }); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy reservation_token must keep its original admitted_at" + .to_string(), + )); + } + } + + let previous_reserved_cost_units = existing + .as_ref() + .map(|reservation| reservation.reserved_cost_units) + .unwrap_or(0); + let target_reserved_cost_units = + previous_reserved_cost_units.max(input.reserved_cost_units); + for (window_index, window) in input.windows.iter().enumerate() { + let used_cost_units = sqlx::query_scalar::<_, i64>( + r#" +SELECT COALESCE(SUM( + CASE + WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) + WHEN state = 'reserved' AND reservation_expires_at > TO_TIMESTAMP($4::double precision) + THEN reserved_cost_units + ELSE 0 + END +), 0)::BIGINT +FROM usage_cost_reservations +WHERE subject_id = $1 + AND admitted_at >= TO_TIMESTAMP($2::double precision) + AND admitted_at < TO_TIMESTAMP($3::double precision) + AND reservation_token <> $5 + "#, + ) + .bind(&input.subject_id) + .bind(usage_policy_cost_i64( + window.starts_at_unix_secs, + "usage policy window start", + )?) + .bind(usage_policy_cost_i64( + window.ends_at_unix_secs, + "usage policy window end", + )?) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy admitted_at", + )?) + .bind(&input.reservation_token) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let used_cost_units = + usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?; + if used_cost_units + .checked_add(target_reserved_cost_units) + .is_none_or(|total| total > window.limit_cost_units) + { + return Ok(ReserveUsagePolicyCostOutcome::Rejected { + window_index, + limit_cost_units: window.limit_cost_units, + used_cost_units, + }); + } + } + + sqlx::query( + r#" +INSERT INTO usage_cost_reservations ( + request_id, subject_id, reservation_token, admitted_at, + reserved_cost_units, actual_cost_units, + state, reservation_expires_at, retain_until, finalized_at, created_at, updated_at +) VALUES ( + $1, $2, $3, TO_TIMESTAMP($4::double precision), $5, NULL, + 'reserved', TO_TIMESTAMP($6::double precision), TO_TIMESTAMP($7::double precision), + NULL, NOW(), NOW() +) +ON CONFLICT (reservation_token) DO UPDATE SET + reserved_cost_units = GREATEST( + usage_cost_reservations.reserved_cost_units, + EXCLUDED.reserved_cost_units + ), + reservation_expires_at = GREATEST( + usage_cost_reservations.reservation_expires_at, + EXCLUDED.reservation_expires_at + ), + retain_until = GREATEST( + usage_cost_reservations.retain_until, + EXCLUDED.retain_until + ), + updated_at = NOW() + "#, + ) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.reservation_token) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy admitted_at", + )?) + .bind(usage_policy_cost_i64( + target_reserved_cost_units, + "usage policy reserved_cost_units", + )?) + .bind(usage_policy_cost_i64( + input.reservation_expires_at_unix_secs, + "usage policy reservation_expires_at", + )?) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy retain_until", + )?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + + Ok(ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: target_reserved_cost_units, + additional_reserved_cost_units: target_reserved_cost_units + .saturating_sub(previous_reserved_cost_units), + }) + }) + }) + .await + } + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + input.validate()?; + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + if !lock_usage_policy_subject_postgres(tx, &input.subject_id).await? { + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_POSTGRES_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(None); + }; + let mut reservation = usage_policy_cost_reservation_from_postgres_row(&row)?; + if reservation.request_id != input.request_id + || reservation.subject_id != input.subject_id + { + // The token selects the row; audit identity must still match before the + // reservation can be finalized. + return Ok(None); + } + if reservation.state == UsagePolicyCostReservationState::Reserved { + sqlx::query( + r#" +UPDATE usage_cost_reservations +SET state = $4, + actual_cost_units = $5, + finalized_at = TO_TIMESTAMP($6::double precision), + updated_at = NOW() +WHERE reservation_token = $1 + AND request_id = $2 + AND subject_id = $3 + AND state = 'reserved' + "#, + ) + .bind(&input.reservation_token) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(input.terminal_state.as_str()) + .bind(usage_policy_cost_i64( + input.actual_cost_units, + "usage policy actual_cost_units", + )?) + .bind(usage_policy_cost_i64( + input.finalized_at_unix_secs, + "usage policy finalized_at", + )?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + reservation.state = input.terminal_state; + reservation.actual_cost_units = Some(input.actual_cost_units); + reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs); + } + Ok(Some(reservation)) + }) + }) + .await + } + + async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_cost_reservations +WHERE retain_until <= TO_TIMESTAMP($1::double precision) + AND reservation_token IN ( + SELECT reservation_token + FROM usage_cost_reservations + WHERE retain_until <= TO_TIMESTAMP($1::double precision) + ORDER BY retain_until, reservation_token + LIMIT $2 +) + "#, + ) + .bind(now) + .bind(limit) + .execute(self.tx_runner.pool()) + .await + .map_postgres_err()?; + Ok(result.rows_affected() as usize) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -507,6 +1176,7 @@ SELECT id, CAST(balance AS DOUBLE PRECISION) AS balance, CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(total_consumed AS DOUBLE PRECISION) AS total_consumed, limit_mode FROM wallets WHERE api_key_id = $1 @@ -534,6 +1204,7 @@ SELECT id, CAST(balance AS DOUBLE PRECISION) AS balance, CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(total_consumed AS DOUBLE PRECISION) AS total_consumed, limit_mode FROM wallets WHERE user_id = $1 @@ -555,14 +1226,26 @@ LIMIT 1 let wallet_can_overdraft = wallet_row.is_some(); let wallet_available_usd = match wallet_row.as_ref() { Some(row) => { + let recharge_balance: f64 = + row.try_get("balance").map_postgres_err()?; + let gift_balance: f64 = + row.try_get("gift_balance").map_postgres_err()?; + let total_consumed: f64 = + row.try_get("total_consumed").map_postgres_err()?; + validate_wallet_settlement_values( + recharge_balance, + gift_balance, + total_consumed, + 0.0, + )?; let limit_mode: String = row.try_get("limit_mode").map_postgres_err()?; if limit_mode.eq_ignore_ascii_case("unlimited") { None } else { Some(finite_wallet_available_usd( - row.try_get("balance").map_postgres_err()?, - row.try_get("gift_balance").map_postgres_err()?, + recharge_balance, + gift_balance, )) } } @@ -630,6 +1313,8 @@ LIMIT 1 wallet_row.try_get("balance").map_postgres_err()?; let before_gift: f64 = wallet_row.try_get("gift_balance").map_postgres_err()?; + let total_consumed: f64 = + wallet_row.try_get("total_consumed").map_postgres_err()?; let limit_mode: String = wallet_row.try_get("limit_mode").map_postgres_err()?; let before_total = before_recharge + before_gift; @@ -644,6 +1329,13 @@ LIMIT 1 (after_recharge, after_gift) = debit_plan.after_balances(before_recharge, before_gift); } + let total_consumed_after = total_consumed + wallet_debit_cost_usd; + validate_wallet_settlement_values( + after_recharge, + after_gift, + total_consumed_after, + 0.0, + )?; if final_billing_status == "settled" { sqlx::query( r#" @@ -651,7 +1343,7 @@ UPDATE wallets SET balance = $2, gift_balance = $3, - total_consumed = CAST(total_consumed AS DOUBLE PRECISION) + $4, + total_consumed = $4, updated_at = NOW() WHERE id = $1 "#, @@ -659,7 +1351,7 @@ WHERE id = $1 .bind(&wallet_id) .bind(after_recharge) .bind(after_gift) - .bind(wallet_debit_cost_usd) + .bind(total_consumed_after) .execute(&mut **tx) .await .map_postgres_err()?; diff --git a/crates/aether-data/adapters/postgres/src/usage/cleanup.rs b/crates/aether-data/adapters/postgres/src/usage/cleanup.rs index ff131617b..509164877 100644 --- a/crates/aether-data/adapters/postgres/src/usage/cleanup.rs +++ b/crates/aether-data/adapters/postgres/src/usage/cleanup.rs @@ -1,11 +1,8 @@ -use std::io::Write; - use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, UsageBodyField, UsageCleanupExecutionMode, - UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow, + UsageCleanupExecutionMode, UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, + UsageCleanupWindow, }; use chrono::{DateTime, Utc}; -use flate2::{write::GzEncoder, Compression}; use futures_util::TryStreamExt; use serde_json::Value; use sqlx::Row; @@ -214,8 +211,7 @@ SET request_body_ref = NULL, WHERE request_id = ANY($1) "#; const SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL: &str = r#" -SELECT - id +SELECT id, request_id FROM usage WHERE created_at < $1 AND ($2::timestamptz IS NULL OR created_at >= $2) @@ -228,26 +224,26 @@ WHERE created_at < $1 OR provider_request_body_compressed IS NOT NULL OR client_response_body IS NOT NULL OR client_response_body_compressed IS NOT NULL + OR EXISTS ( + SELECT 1 + FROM usage_body_blobs + WHERE usage_body_blobs.request_id = usage.request_id + ) + OR EXISTS ( + SELECT 1 + FROM usage_http_audits + WHERE usage_http_audits.request_id = usage.request_id + AND ( + usage_http_audits.request_body_ref IS NOT NULL + OR usage_http_audits.provider_request_body_ref IS NOT NULL + OR usage_http_audits.response_body_ref IS NOT NULL + OR usage_http_audits.client_response_body_ref IS NOT NULL + ) + ) ) ORDER BY created_at ASC, id ASC LIMIT $3 "#; -const SELECT_USAGE_BODY_COMPRESSION_ROW_SQL: &str = r#" -SELECT - id, - request_id, - request_body, - request_body_compressed, - response_body, - response_body_compressed, - provider_request_body, - provider_request_body_compressed, - client_response_body, - client_response_body_compressed -FROM usage -WHERE id = $1 -LIMIT 1 -"#; const SELECT_EXPIRED_ACTIVE_API_KEYS_SQL: &str = r#" SELECT id, auto_delete_on_expiry FROM api_keys @@ -274,30 +270,6 @@ WHERE id = $1 AND is_active IS TRUE "#; -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct UsageDetachedBodyBlobWrite { - pub body_ref: String, - pub body_field: &'static str, - pub payload_gzip: Vec, -} - -#[derive(Debug, Clone, Default, PartialEq, Eq)] -pub struct UsageDetachedBodyRefs { - pub request_body_ref: Option, - pub provider_request_body_ref: Option, - pub response_body_ref: Option, - pub client_response_body_ref: Option, -} - -impl UsageDetachedBodyRefs { - pub fn any_present(&self) -> bool { - self.request_body_ref.is_some() - || self.provider_request_body_ref.is_some() - || self.response_body_ref.is_some() - || self.client_response_body_ref.is_some() - } -} - #[derive(Debug, Clone, PartialEq)] pub struct UsageLegacyBodyRefMetadataRow { pub id: String, @@ -305,30 +277,9 @@ pub struct UsageLegacyBodyRefMetadataRow { pub request_metadata: Option, } -#[derive(Debug, Clone, Default, PartialEq)] -pub struct UsageLegacyBodyRefMigrationPlan { - pub refs: UsageDetachedBodyRefs, - pub request_metadata: Option, -} - #[derive(Debug, Clone, PartialEq)] -pub struct UsageBodyCompressionRow { - pub id: String, - pub request_id: String, - pub request_body: Option, - pub request_body_compressed: Option>, - pub response_body: Option, - pub response_body_compressed: Option>, - pub provider_request_body: Option, - pub provider_request_body_compressed: Option>, - pub client_response_body: Option, - pub client_response_body_compressed: Option>, -} - -#[derive(Debug, Clone, Default, PartialEq, Eq)] -pub struct UsageBodyExternalizationPlan { - pub blobs: Vec, - pub refs: UsageDetachedBodyRefs, +pub struct UsageLegacyBodyRefPurgePlan { + pub request_metadata: Option, } #[derive(Debug, Clone, PartialEq)] @@ -343,57 +294,23 @@ struct ExpiredApiKeyRow<'a> { auto_delete_on_expiry: Option, } -pub fn compress_usage_json_value(value: &Value) -> Result, DataLayerError> { - let bytes = serde_json::to_vec(value).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to serialize usage json for gzip: {err}")) - })?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6)); - encoder.write_all(&bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to gzip usage json: {err}")) - })?; - encoder.finish().map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to finish gzipped usage json: {err}")) - }) -} - -pub fn migrate_legacy_body_ref_metadata_plan( - request_id: &str, +pub fn purge_legacy_body_ref_metadata_plan( request_metadata: Option, -) -> Option { +) -> Option { let mut metadata = match request_metadata { Some(Value::Object(object)) => object, _ => return None, }; - let mut refs = UsageDetachedBodyRefs::default(); let mut removed_any = false; - for field in [ - UsageBodyField::RequestBody, - UsageBodyField::ProviderRequestBody, - UsageBodyField::ResponseBody, - UsageBodyField::ClientResponseBody, + for key in [ + "request_body_ref", + "provider_request_body_ref", + "response_body_ref", + "client_response_body_ref", ] { - let key = field.as_ref_key(); - let Some(value) = metadata.remove(key) else { - continue; - }; - removed_any = true; - let parsed = value - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| { - usage_body_ref(&parsed_request_id, parsed_field) - }); - match field { - UsageBodyField::RequestBody => refs.request_body_ref = parsed, - UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = parsed, - UsageBodyField::ResponseBody => refs.response_body_ref = parsed, - UsageBodyField::ClientResponseBody => refs.client_response_body_ref = parsed, + if metadata.remove(key).is_some() { + removed_any = true; } } @@ -401,47 +318,11 @@ pub fn migrate_legacy_body_ref_metadata_plan( return None; } - Some(UsageLegacyBodyRefMigrationPlan { - refs, + Some(UsageLegacyBodyRefPurgePlan { request_metadata: (!metadata.is_empty()).then_some(Value::Object(metadata)), }) } -pub fn build_usage_body_externalization( - row: &UsageBodyCompressionRow, -) -> Result { - let mut plan = UsageBodyExternalizationPlan::default(); - maybe_externalize_usage_body_field( - &mut plan, - &row.request_id, - UsageBodyField::RequestBody, - row.request_body.as_ref(), - row.request_body_compressed.as_deref(), - )?; - maybe_externalize_usage_body_field( - &mut plan, - &row.request_id, - UsageBodyField::ProviderRequestBody, - row.provider_request_body.as_ref(), - row.provider_request_body_compressed.as_deref(), - )?; - maybe_externalize_usage_body_field( - &mut plan, - &row.request_id, - UsageBodyField::ResponseBody, - row.response_body.as_ref(), - row.response_body_compressed.as_deref(), - )?; - maybe_externalize_usage_body_field( - &mut plan, - &row.request_id, - UsageBodyField::ClientResponseBody, - row.client_response_body.as_ref(), - row.client_response_body_compressed.as_deref(), - )?; - Ok(plan) -} - impl SqlxUsageReadRepository { pub async fn cleanup_usage( &self, @@ -477,6 +358,8 @@ impl SqlxUsageReadRepository { header_cleaned: 0, keys_cleaned: 0, records_deleted: 0, + cost_reservations_deleted: 0, + request_admissions_deleted: 0, }); } @@ -509,7 +392,7 @@ impl SqlxUsageReadRepository { }; let detail_body_newer_than = detail_body_newer_than(window, targets); let legacy_body_refs_migrated = if targets.detail_body { - migrate_legacy_usage_body_ref_metadata( + purge_legacy_usage_body_ref_metadata( &self.pool, window.detail_cutoff, batch_size, @@ -520,7 +403,7 @@ impl SqlxUsageReadRepository { 0 }; let body_externalized = if targets.detail_body { - compress_usage_body_fields( + purge_usage_detail_body_fields( &self.pool, window.detail_cutoff, batch_size, @@ -549,6 +432,8 @@ impl SqlxUsageReadRepository { header_cleaned, keys_cleaned, records_deleted, + cost_reservations_deleted: 0, + request_admissions_deleted: 0, }) } } @@ -646,12 +531,33 @@ async fn cleanup_usage_raw_body_fields( break; } let ids = rows.iter().map(|row| row.id.clone()).collect::>(); + let request_ids = rows + .iter() + .map(|row| row.request_id.clone()) + .collect::>(); + let mut tx = pool.begin().await.map_err(postgres_error)?; let cleaned = sqlx::query(CLEAR_USAGE_RAW_BODY_FIELDS_SQL) .bind(ids) - .execute(pool) + .execute(&mut *tx) .await .map_err(postgres_error)? .rows_affected(); + sqlx::query(DELETE_USAGE_BODY_BLOBS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(CLEAR_USAGE_HTTP_AUDIT_BODY_REFS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL) + .bind(request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + tx.commit().await.map_err(postgres_error)?; let cleaned = usize::try_from(cleaned).unwrap_or(usize::MAX); total_cleaned += cleaned; if rows.len() < batch_size { @@ -973,7 +879,7 @@ async fn delete_old_usage_records( Ok(total_deleted) } -async fn migrate_legacy_usage_body_ref_metadata( +async fn purge_legacy_usage_body_ref_metadata( pool: &PostgresPool, cutoff_time: DateTime, batch_size: usize, @@ -983,7 +889,7 @@ async fn migrate_legacy_usage_body_ref_metadata( warn!( cutoff_time = %cutoff_time, newer_than = ?newer_than, - "usage cleanup legacy body-ref migration skipped due to invalid window" + "usage cleanup legacy body-ref purge skipped due to invalid window" ); return Ok(0); } @@ -1014,26 +920,12 @@ async fn migrate_legacy_usage_body_ref_metadata( break; } - let mut batch_migrated = 0usize; + let mut batch_purged = 0usize; for row in rows { - let Some(plan) = - migrate_legacy_body_ref_metadata_plan(&row.request_id, row.request_metadata) - else { + let Some(plan) = purge_legacy_body_ref_metadata_plan(row.request_metadata) else { continue; }; let mut tx = pool.begin().await.map_err(postgres_error)?; - if plan.refs.any_present() { - sqlx::query(UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL) - .bind(&row.request_id) - .bind(plan.refs.request_body_ref.as_deref()) - .bind(plan.refs.provider_request_body_ref.as_deref()) - .bind(plan.refs.response_body_ref.as_deref()) - .bind(plan.refs.client_response_body_ref.as_deref()) - .bind("ref_backed") - .execute(&mut *tx) - .await - .map_err(postgres_error)?; - } let updated = sqlx::query(UPDATE_USAGE_REQUEST_METADATA_SQL) .bind(&row.id) .bind(plan.request_metadata) @@ -1041,14 +933,30 @@ async fn migrate_legacy_usage_body_ref_metadata( .await .map_err(postgres_error)? .rows_affected(); + let request_ids = vec![row.request_id]; + sqlx::query(DELETE_USAGE_BODY_BLOBS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(CLEAR_USAGE_HTTP_AUDIT_BODY_REFS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL) + .bind(request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; tx.commit().await.map_err(postgres_error)?; if updated > 0 { - batch_migrated += 1; + batch_purged += 1; } } - total_migrated += batch_migrated; - if batch_migrated == 0 || batch_migrated < batch_size { + total_migrated += batch_purged; + if batch_purged == 0 || batch_purged < batch_size { break; } } @@ -1191,7 +1099,7 @@ async fn cleanup_usage_stale_body_fields( Ok(total_cleaned) } -async fn compress_usage_body_fields( +async fn purge_usage_detail_body_fields( pool: &PostgresPool, cutoff_time: DateTime, batch_size: usize, @@ -1201,129 +1109,65 @@ async fn compress_usage_body_fields( warn!( cutoff_time = %cutoff_time, newer_than = ?newer_than, - "usage cleanup body compression skipped due to invalid window" + "usage cleanup detail body purge skipped due to invalid window" ); return Ok(0); } - let mut total_compressed = 0usize; - let mut no_progress_count = 0usize; - let batch_size = batch_size.clamp(1, 25); + let mut total_purged = 0usize; loop { let mut stream = sqlx::query(SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL) .bind(cutoff_time) .bind(newer_than) .bind(i64::try_from(batch_size).unwrap_or(i64::MAX)) .fetch(pool); - let mut ids = Vec::new(); + let mut rows = Vec::new(); while let Some(row) = stream.try_next().await.map_err(postgres_error)? { - ids.push(row.try_get::("id").map_err(postgres_error)?); - } - if ids.is_empty() { - break; - } - - let mut batch_success = 0usize; - for id in ids { - let row = sqlx::query(SELECT_USAGE_BODY_COMPRESSION_ROW_SQL) - .bind(&id) - .fetch_optional(pool) - .await - .map_err(postgres_error)?; - let Some(row) = row else { - continue; - }; - let row = UsageBodyCompressionRow { + rows.push(UsageBodyCleanupRow { id: row.try_get::("id").map_err(postgres_error)?, request_id: row .try_get::("request_id") .map_err(postgres_error)?, - request_body: row - .try_get::, _>("request_body") - .map_err(postgres_error)?, - request_body_compressed: row - .try_get::>, _>("request_body_compressed") - .map_err(postgres_error)?, - response_body: row - .try_get::, _>("response_body") - .map_err(postgres_error)?, - response_body_compressed: row - .try_get::>, _>("response_body_compressed") - .map_err(postgres_error)?, - provider_request_body: row - .try_get::, _>("provider_request_body") - .map_err(postgres_error)?, - provider_request_body_compressed: row - .try_get::>, _>("provider_request_body_compressed") - .map_err(postgres_error)?, - client_response_body: row - .try_get::, _>("client_response_body") - .map_err(postgres_error)?, - client_response_body_compressed: row - .try_get::>, _>("client_response_body_compressed") - .map_err(postgres_error)?, - }; - let detached = build_usage_body_externalization(&row)?; - if detached.refs.any_present() { - let mut tx = pool.begin().await.map_err(postgres_error)?; - for blob in &detached.blobs { - sqlx::query(super::UPSERT_USAGE_BODY_BLOB_SQL) - .bind(&blob.body_ref) - .bind(&row.request_id) - .bind(blob.body_field) - .bind(&blob.payload_gzip) - .execute(&mut *tx) - .await - .map_err(postgres_error)?; - } - sqlx::query(UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL) - .bind(&row.request_id) - .bind(detached.refs.request_body_ref.as_deref()) - .bind(detached.refs.provider_request_body_ref.as_deref()) - .bind(detached.refs.response_body_ref.as_deref()) - .bind(detached.refs.client_response_body_ref.as_deref()) - .bind("ref_backed") - .execute(&mut *tx) - .await - .map_err(postgres_error)?; - let updated = sqlx::query(UPDATE_USAGE_BODY_COMPRESSION_SQL) - .bind(&row.id) - .execute(&mut *tx) - .await - .map_err(postgres_error)? - .rows_affected(); - tx.commit().await.map_err(postgres_error)?; - if updated > 0 { - batch_success += 1; - } - continue; - } - - let updated = sqlx::query(UPDATE_USAGE_BODY_COMPRESSION_SQL) - .bind(&row.id) - .execute(pool) - .await - .map_err(postgres_error)? - .rows_affected(); - if updated > 0 { - batch_success += 1; - } + }); } - - if batch_success == 0 { - no_progress_count += 1; - if no_progress_count >= 3 { - warn!( - "usage cleanup body compression stopped after repeated zero-progress batches" - ); - break; - } - } else { - no_progress_count = 0; + if rows.is_empty() { + break; + } + let row_count = rows.len(); + let ids = rows.iter().map(|row| row.id.clone()).collect::>(); + let request_ids = rows + .iter() + .map(|row| row.request_id.clone()) + .collect::>(); + let mut tx = pool.begin().await.map_err(postgres_error)?; + let updated = sqlx::query(CLEAR_USAGE_BODY_FIELDS_SQL) + .bind(ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)? + .rows_affected(); + sqlx::query(DELETE_USAGE_BODY_BLOBS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(CLEAR_USAGE_HTTP_AUDIT_BODY_REFS_SQL) + .bind(&request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query(DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL) + .bind(request_ids) + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + tx.commit().await.map_err(postgres_error)?; + total_purged = total_purged.saturating_add(usize::try_from(updated).unwrap_or(usize::MAX)); + if row_count < batch_size { + break; } - total_compressed += batch_success; } - Ok(total_compressed) + Ok(total_purged) } async fn cleanup_expired_api_keys( @@ -1374,74 +1218,6 @@ async fn cleanup_expired_api_keys( Ok(cleaned) } -fn maybe_externalize_usage_body_field( - plan: &mut UsageBodyExternalizationPlan, - request_id: &str, - field: UsageBodyField, - inline_body: Option<&Value>, - compressed_body: Option<&[u8]>, -) -> Result<(), DataLayerError> { - let Some(payload_gzip) = (match inline_body { - Some(value) => Some(compress_usage_json_value(value)?), - None => compressed_body.map(|value| value.to_vec()), - }) else { - return Ok(()); - }; - let body_ref = usage_body_ref(request_id, field); - plan.blobs.push(UsageDetachedBodyBlobWrite { - body_ref: body_ref.clone(), - body_field: field.as_storage_field(), - payload_gzip, - }); - match field { - UsageBodyField::RequestBody => plan.refs.request_body_ref = Some(body_ref), - UsageBodyField::ProviderRequestBody => plan.refs.provider_request_body_ref = Some(body_ref), - UsageBodyField::ResponseBody => plan.refs.response_body_ref = Some(body_ref), - UsageBodyField::ClientResponseBody => plan.refs.client_response_body_ref = Some(body_ref), - } - Ok(()) -} - -const UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL: &str = r#" -INSERT INTO usage_http_audits ( - request_id, - request_body_ref, - provider_request_body_ref, - response_body_ref, - client_response_body_ref, - body_capture_mode -) -VALUES ( - $1, - $2, - $3, - $4, - $5, - $6 -) -ON CONFLICT (request_id) -DO UPDATE SET - request_body_ref = COALESCE(EXCLUDED.request_body_ref, usage_http_audits.request_body_ref), - provider_request_body_ref = COALESCE( - EXCLUDED.provider_request_body_ref, - usage_http_audits.provider_request_body_ref - ), - response_body_ref = COALESCE(EXCLUDED.response_body_ref, usage_http_audits.response_body_ref), - client_response_body_ref = COALESCE( - EXCLUDED.client_response_body_ref, - usage_http_audits.client_response_body_ref - ), - body_capture_mode = CASE - WHEN EXCLUDED.request_body_ref IS NOT NULL - OR EXCLUDED.provider_request_body_ref IS NOT NULL - OR EXCLUDED.response_body_ref IS NOT NULL - OR EXCLUDED.client_response_body_ref IS NOT NULL - THEN EXCLUDED.body_capture_mode - ELSE usage_http_audits.body_capture_mode - END, - updated_at = NOW() -"#; - const UPDATE_USAGE_REQUEST_METADATA_SQL: &str = r#" UPDATE usage SET request_metadata = $2::json, @@ -1449,41 +1225,15 @@ SET request_metadata = $2::json, WHERE id = $1 "#; -const UPDATE_USAGE_BODY_COMPRESSION_SQL: &str = r#" -UPDATE usage -SET request_body = NULL, - response_body = NULL, - provider_request_body = NULL, - client_response_body = NULL, - request_body_compressed = NULL, - response_body_compressed = NULL, - provider_request_body_compressed = NULL, - client_response_body_compressed = NULL -WHERE id = $1 -"#; - #[cfg(test)] mod tests { - use std::io::Read; - - use flate2::read::GzDecoder; use serde_json::json; use super::{ - build_usage_body_externalization, compress_usage_json_value, - migrate_legacy_body_ref_metadata_plan, UsageBodyCompressionRow, + purge_legacy_body_ref_metadata_plan, SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL, SELECT_USAGE_LEGACY_BODY_REF_METADATA_BATCH_SQL, }; - fn inflate_json(bytes: &[u8]) -> serde_json::Value { - let mut decoder = GzDecoder::new(bytes); - let mut decoded = Vec::new(); - decoder - .read_to_end(&mut decoded) - .expect("gzip should decode"); - serde_json::from_slice(&decoded).expect("json should decode") - } - #[test] fn legacy_body_ref_cleanup_index_is_embedded_and_matches_batch_predicate() { const MIGRATION_VERSION: i64 = 20_260_715_000_000; @@ -1532,88 +1282,21 @@ mod tests { } #[test] - fn usage_body_externalization_moves_inline_json_into_ref_backed_blobs() { - let row = UsageBodyCompressionRow { - id: "usage-1".to_string(), - request_id: "req-1".to_string(), - request_body: Some(json!({"hello": "world"})), - request_body_compressed: None, - response_body: None, - response_body_compressed: None, - provider_request_body: Some(json!({"provider": true})), - provider_request_body_compressed: None, - client_response_body: None, - client_response_body_compressed: None, - }; - - let plan = build_usage_body_externalization(&row).expect("plan should build"); - - assert_eq!(plan.blobs.len(), 2); - assert_eq!( - plan.refs.request_body_ref.as_deref(), - Some("usage://request/req-1/request_body") - ); - assert_eq!( - plan.refs.provider_request_body_ref.as_deref(), - Some("usage://request/req-1/provider_request_body") - ); - assert_eq!( - inflate_json(&plan.blobs[0].payload_gzip), - json!({"hello": "world"}) - ); - assert_eq!( - inflate_json(&plan.blobs[1].payload_gzip), - json!({"provider": true}) - ); + fn detail_body_cleanup_selects_detached_capture_for_deletion() { + assert!(SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL.contains("usage_body_blobs")); + assert!(SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL.contains("usage_http_audits")); + assert!(!SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL.contains("payload_gzip")); } #[test] - fn usage_body_externalization_reuses_existing_compressed_payloads() { - let compressed = compress_usage_json_value(&json!({"legacy": true})) - .expect("compressed payload should build"); - let row = UsageBodyCompressionRow { - id: "usage-1".to_string(), - request_id: "req-legacy".to_string(), - request_body: None, - request_body_compressed: Some(compressed.clone()), - response_body: None, - response_body_compressed: None, - provider_request_body: None, - provider_request_body_compressed: None, - client_response_body: None, - client_response_body_compressed: None, - }; - - let plan = build_usage_body_externalization(&row).expect("plan should build"); - - assert_eq!(plan.blobs.len(), 1); - assert_eq!(plan.blobs[0].payload_gzip, compressed); - assert_eq!( - plan.refs.request_body_ref.as_deref(), - Some("usage://request/req-legacy/request_body") - ); - } - - #[test] - fn legacy_body_ref_metadata_migration_moves_matching_refs_and_strips_keys() { - let plan = migrate_legacy_body_ref_metadata_plan( - "req-1", - Some(json!({ - "trace_id": "trace-1", - "request_body_ref": "usage://request/req-1/request_body", - "response_body_ref": "usage://request/req-1/response_body" - })), - ) + fn legacy_body_ref_metadata_purge_strips_all_ref_keys() { + let plan = purge_legacy_body_ref_metadata_plan(Some(json!({ + "trace_id": "trace-1", + "request_body_ref": "usage://request/req-1/request_body", + "response_body_ref": "usage://request/req-1/response_body" + }))) .expect("migration plan should exist"); - assert_eq!( - plan.refs.request_body_ref.as_deref(), - Some("usage://request/req-1/request_body") - ); - assert_eq!( - plan.refs.response_body_ref.as_deref(), - Some("usage://request/req-1/response_body") - ); assert_eq!( plan.request_metadata, Some(json!({ @@ -1623,18 +1306,14 @@ mod tests { } #[test] - fn legacy_body_ref_metadata_migration_strips_invalid_and_cross_request_refs() { - let plan = migrate_legacy_body_ref_metadata_plan( - "req-1", - Some(json!({ - "request_body_ref": "blob://legacy-request", - "provider_request_body_ref": "usage://request/req-other/provider_request_body", - "candidate_index": 2 - })), - ) + fn legacy_body_ref_metadata_purge_does_not_preserve_untrusted_refs() { + let plan = purge_legacy_body_ref_metadata_plan(Some(json!({ + "request_body_ref": "blob://legacy-request", + "provider_request_body_ref": "usage://request/req-other/provider_request_body", + "candidate_index": 2 + }))) .expect("migration plan should exist"); - assert!(!plan.refs.any_present()); assert_eq!( plan.request_metadata, Some(json!({ diff --git a/crates/aether-data/adapters/postgres/src/usage/mod.rs b/crates/aether-data/adapters/postgres/src/usage/mod.rs index bda75d59c..0e50692de 100644 --- a/crates/aether-data/adapters/postgres/src/usage/mod.rs +++ b/crates/aether-data/adapters/postgres/src/usage/mod.rs @@ -1,8 +1,9 @@ use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, - ProxyNodeCounterDelta, StoredUsageAuditAggregation, StoredUsageAuditSummary, - StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, - StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, + canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json, + usage_body_ref, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, ProxyNodeCounterDelta, + StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, + StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow, + StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow, @@ -33,7 +34,7 @@ use sqlx::{ PgPool, Postgres, QueryBuilder, Row, }; use std::collections::{BTreeMap, BTreeSet}; -use std::io::{Read, Write}; +use std::io::Write; use uuid::Uuid; use crate::{ @@ -41,11 +42,12 @@ use crate::{ PostgresTransactionRunner, }; use aether_data_contracts::repository::usage::{ - api_key_usage_contribution, incoming_usage_can_recover_terminal_failure, - model_usage_contribution, provider_api_key_usage_contribution, - strip_deprecated_usage_display_fields, ApiKeyUsageDelta, ModelUsageDelta, - PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, - ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, + api_key_usage_contribution, model_usage_contribution, provider_api_key_usage_contribution, + sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence, + sanitize_usage_request_metadata, usage_can_recover_terminal_failure, + usage_error_category_for_status_code, usage_lifecycle_update_allowed, ApiKeyUsageDelta, + ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, + ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredRequestUsageAudit, StoredUsageDailySummary, UpsertUsageRecord, UsageAuditListQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, @@ -61,9 +63,7 @@ pub mod cleanup; // newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits. const MAX_INLINE_USAGE_BODY_BYTES: usize = 0; const MAX_SUPPORTED_UNIX_SECS: u64 = 253_402_300_799; -const FIND_USAGE_BODY_BLOB_BY_REF_SQL: &str = - r#"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = $1 LIMIT 1"#; -const UPSERT_USAGE_BODY_BLOB_SQL: &str = include_str!("queries/upsert_usage_body_blob_sql.sql"); +const FIND_USAGE_BODY_BLOB_BY_REF_SQL: &str = r#"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = $1 AND request_id = $2 AND body_field = $3 LIMIT 1"#; const DELETE_USAGE_BODY_BLOB_SQL: &str = include_str!("queries/delete_usage_body_blob_sql.sql"); #[derive(Debug, Clone, PartialEq, Eq, Default)] @@ -1372,9 +1372,10 @@ SELECT FROM "usage" AS u WHERE u.request_id = ANY($1) "#; -const UPSERT_USAGE_HTTP_AUDIT_SQL: &str = include_str!("queries/upsert_usage_http_audit_sql.sql"); const UPSERT_USAGE_ROUTING_SNAPSHOT_SQL: &str = include_str!("queries/upsert_usage_routing_snapshot_sql.sql"); +#[cfg(test)] +const UPSERT_USAGE_HTTP_AUDIT_SQL: &str = include_str!("queries/upsert_usage_http_audit_sql.sql"); const UPSERT_USAGE_SETTLEMENT_PRICING_SNAPSHOT_SQL: &str = include_str!("queries/upsert_usage_settlement_pricing_snapshot_sql.sql"); @@ -1418,6 +1419,7 @@ INSERT INTO usage_counter_deltas ( request_id, kind, target_id, + target_tunnel_generation, request_count_delta, total_requests_delta, success_count_delta, @@ -1433,7 +1435,7 @@ INSERT INTO usage_counter_deltas ( removed_last_used_at_unix_secs, usage_created_at_unix_secs ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18 + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19 ) "#; const INSERT_USAGE_COUNTER_DELTAS_PREFIX_SQL: &str = r#" @@ -1442,6 +1444,7 @@ INSERT INTO usage_counter_deltas ( request_id, kind, target_id, + target_tunnel_generation, request_count_delta, total_requests_delta, success_count_delta, @@ -1457,7 +1460,7 @@ INSERT INTO usage_counter_deltas ( removed_last_used_at_unix_secs, usage_created_at_unix_secs ) "#; -const USAGE_COUNTER_DELTA_INSERT_BINDS_PER_ROW: usize = 18; +const USAGE_COUNTER_DELTA_INSERT_BINDS_PER_ROW: usize = 19; const USAGE_COUNTER_DELTA_INSERT_BATCH_SIZE: usize = u16::MAX as usize / USAGE_COUNTER_DELTA_INSERT_BINDS_PER_ROW; @@ -1474,6 +1477,7 @@ SELECT delta.id, delta.kind, delta.target_id, + delta.target_tunnel_generation, delta.request_count_delta, delta.total_requests_delta, delta.success_count_delta, @@ -1586,6 +1590,7 @@ SET stream_errors = stream_errors + GREATEST($5::bigint, 0), updated_at = NOW() WHERE id = $1 + AND tunnel_generation = $6 "#; const APPLY_MANAGEMENT_TOKEN_COUNTER_DELTA_SQL: &str = r#" @@ -1937,7 +1942,6 @@ INSERT INTO "usage" ( const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#" SELECT usage.request_id, - usage.status, COALESCE(usage_settlement_snapshots.billing_status, usage.billing_status) AS billing_status FROM usage LEFT JOIN usage_settlement_snapshots @@ -1966,15 +1970,15 @@ const UPDATE_RECOVERED_STALE_USAGE_SQL: &str = r#" UPDATE usage SET status = 'completed', status_code = 200, - error_message = NULL + error_message = NULL, + error_category = NULL WHERE request_id = $1 "#; const SELECT_LATEST_FAILED_CANDIDATE_FOR_STALE_REQUESTS_SQL: &str = r#" SELECT DISTINCT ON (request_id) request_id, - status_code, - error_message + status_code FROM request_candidates WHERE request_id = ANY($1) AND status IN ('failed', 'cancelled') @@ -1988,8 +1992,9 @@ ORDER BY request_id, const UPDATE_FAILED_STALE_USAGE_SQL: &str = r#" UPDATE usage SET status = 'failed', - status_code = $3, - error_message = $2 + status_code = $2, + error_message = NULL, + error_category = $3 WHERE request_id = $1 "#; @@ -1997,10 +2002,11 @@ const UPDATE_FAILED_VOID_STALE_USAGE_SQL: &str = r#" WITH updated_usage AS ( UPDATE usage SET status = 'failed', - status_code = $4, - error_message = $2, + status_code = $3, + error_message = NULL, + error_category = $4, billing_status = 'void', - finalized_at = $3, + finalized_at = $2, total_cost_usd = 0, request_cost_usd = 0, actual_total_cost_usd = 0, @@ -2013,7 +2019,7 @@ INSERT INTO usage_settlement_snapshots ( billing_status, finalized_at ) -SELECT request_id, 'void', $3 +SELECT request_id, 'void', $2 FROM updated_usage ON CONFLICT (request_id) DO UPDATE SET @@ -2037,7 +2043,8 @@ const UPDATE_FAILED_PENDING_CANDIDATES_SQL: &str = r#" UPDATE request_candidates SET status = 'failed', finished_at = $2, - error_message = '请求超时(服务器可能已重启)' + error_type = 'internal', + error_message = NULL WHERE request_id = $1 AND status IN ('pending', 'streaming') "#; @@ -2097,8 +2104,12 @@ impl PreparedPendingUsage { )); } - let usage = strip_deprecated_usage_display_fields(usage); - let prepared = prepare_usage_upsert_context(&usage)?; + // Keep the capture input separate from the accounting row. The persistence sanitizer + // intentionally removes HTTP bodies/headers/states, but the pending batch still needs + // those values to populate the canonical audit/blob tables. + let capture_usage = usage.clone(); + let usage = sanitize_usage_for_persistence(usage); + let prepared = prepare_usage_upsert_context(&capture_usage)?; let input_tokens = usage .input_tokens .map(to_i32) @@ -2215,7 +2226,7 @@ impl PreparedFirstByteUsage { )); } - let usage = strip_deprecated_usage_display_fields(usage); + let usage = sanitize_usage_for_persistence(usage); let request_metadata_json = json_bind_text(usage.request_metadata.as_ref())?; let response_time_ms = usage.response_time_ms.map(to_i32).transpose()?; let first_byte_time_ms = usage.first_byte_time_ms.map(to_i32).transpose()?; @@ -2251,6 +2262,7 @@ fn partition_first_byte_usages( let mut batch_rows = Vec::new(); let mut fallback_rows = Vec::new(); for (sequence, usage) in usages.into_iter().enumerate() { + let original_usage = usage.clone(); let prepared = PreparedFirstByteUsage::try_from_usage(usage)?; if request_id_counts .get(&prepared.usage.request_id) @@ -2260,7 +2272,10 @@ fn partition_first_byte_usages( { batch_rows.push(prepared); } else { - fallback_rows.push((sequence, prepared.usage)); + // Duplicate rows are replayed through the canonical upsert, which owns the + // persistence sanitizer. Keep the original metadata here so replay semantics remain + // lossless up to that boundary. + fallback_rows.push((sequence, original_usage)); } } Ok((batch_rows, fallback_rows)) @@ -2848,8 +2863,14 @@ ORDER BY request_count DESC, "usage".provider_name ASC } pub async fn resolve_body_ref(&self, body_ref: &str) -> Result, DataLayerError> { + let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { + return Ok(None); + }; + let canonical_ref = usage_body_ref(&request_id, field); let blob_row = sqlx::query(FIND_USAGE_BODY_BLOB_BY_REF_SQL) - .bind(body_ref) + .bind(&canonical_ref) + .bind(&request_id) + .bind(field.as_storage_field()) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -2859,9 +2880,6 @@ ORDER BY request_count DESC, "usage".provider_name ASC .map_postgres_err()?; return inflate_usage_json_value(&payload_gzip).map(Some); } - let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { - return Ok(None); - }; let (inline_column, compressed_column) = usage_body_sql_columns(field); let row = sqlx::query(&format!( "SELECT {inline_column} AS inline_body, {compressed_column} AS compressed_body FROM \"usage\" WHERE request_id = $1 LIMIT 1" @@ -2908,8 +2926,10 @@ ORDER BY request_count DESC, "usage".provider_name ASC usage: &StoredRequestUsageAudit, field: UsageBodyField, ) -> Result, DataLayerError> { - let body_ref = usage.body_ref(field); - match body_ref { + let body_ref = usage + .body_ref(field) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, &usage.request_id, field)); + match body_ref.as_deref() { Some(body_ref) => self.resolve_body_ref(body_ref).await, None => Ok(None), } @@ -8367,40 +8387,41 @@ ORDER BY "usage".user_id ASC usage: UpsertUsageRecord, ) -> Result { usage.validate()?; - let usage = strip_deprecated_usage_display_fields(usage); - let prepared = prepare_usage_upsert_context(&usage)?; + // `usage` is the sanitized accounting projection; prepare the auxiliary capture and + // snapshots from the original event so typed `none` markers can clear prior facts. + let capture_usage = usage.clone(); + let usage = sanitize_usage_for_persistence(usage); self.tx_runner .run_read_write(|tx| { - let PreparedUsageUpsert { - request_headers_json, - provider_request_headers_json, - response_headers_json, - client_response_headers_json, - request_body_storage, - provider_request_body_storage, - response_body_storage, - client_response_body_storage, - http_audit_refs, - http_audit_states, - http_audit_capture_mode, - routing_snapshot, - settlement_pricing_snapshot, - mut request_metadata_value, - mut request_metadata_json, - replace_client_request_body_facts, - replace_provider_request_body_facts, - clear_request_body, - clear_provider_request_body, - clear_response_body, - clear_client_response_body, - } = prepared; Box::pin(async move { lock_usage_request_id_in_tx(tx, &usage.request_id).await?; - if incoming_usage_can_recover_terminal_failure( - usage.status.as_str(), - usage.billing_status.as_str(), - ) { + let previous_usage = + find_usage_by_request_id_in_tx(tx, &usage.request_id).await?; + if let Some(previous) = previous_usage.as_ref() { + if !usage_lifecycle_update_allowed( + &previous.status, + &previous.billing_status, + previous.updated_at_unix_secs, + previous.finalized_at_unix_secs, + &usage.status, + &usage.billing_status, + usage.updated_at_unix_secs, + usage.finalized_at_unix_secs, + ) { + return Ok(previous.clone()); + } + } + let recovers_terminal_failure = + previous_usage.as_ref().is_some_and(|previous| { + usage_can_recover_terminal_failure( + &previous.status, + &previous.billing_status, + &usage.status, + &usage.billing_status, + ) + }); + if recovers_terminal_failure { sqlx::query(RESET_STALE_VOID_USAGE_SQL) .bind(&usage.request_id) .execute(&mut **tx) @@ -8413,14 +8434,36 @@ ORDER BY "usage".user_id ASC .map_postgres_err()?; } - let previous_usage = - find_usage_by_request_id_in_tx(tx, &usage.request_id).await?; - let capture_update_allowed = usage_capture_update_allowed( - previous_usage - .as_ref() - .map(|stored| (stored.status.as_str(), stored.billing_status.as_str())), - usage.status.as_str(), - ); + let PreparedUsageUpsert { + request_headers_json, + provider_request_headers_json, + response_headers_json, + client_response_headers_json, + request_body_storage, + provider_request_body_storage, + response_body_storage, + client_response_body_storage, + http_audit_refs, + http_audit_states, + http_audit_capture_mode, + routing_snapshot, + settlement_pricing_snapshot, + mut request_metadata_value, + mut request_metadata_json, + replace_client_request_body_facts, + replace_provider_request_body_facts, + clear_request_body, + clear_provider_request_body, + clear_response_body, + clear_client_response_body, + } = prepare_usage_upsert_context(&capture_usage)?; + let capture_update_allowed = recovers_terminal_failure + || usage_capture_update_allowed( + previous_usage.as_ref().map(|stored| { + (stored.status.as_str(), stored.billing_status.as_str()) + }), + usage.status.as_str(), + ); let replace_terminal_snapshots = matches!(usage.status.as_str(), "completed" | "failed" | "cancelled"); if capture_update_allowed @@ -8432,7 +8475,10 @@ ORDER BY "usage".user_id ASC let previous_metadata = previous_usage .as_ref() .and_then(|stored| stored.request_metadata.as_ref()); - request_metadata_value = Some(if replace_terminal_snapshots { + let preserve_empty_tombstone = !replace_terminal_snapshots + && (replace_client_request_body_facts + || replace_provider_request_body_facts); + let previous_metadata = if replace_terminal_snapshots { retain_previous_request_audit_metadata( previous_metadata, !replace_client_request_body_facts, @@ -8443,7 +8489,14 @@ ORDER BY "usage".user_id ASC replace_client_request_body_facts, replace_provider_request_body_facts, ) - }); + }; + request_metadata_value = Some( + project_usage_request_metadata( + Some(previous_metadata), + preserve_empty_tombstone, + ) + .unwrap_or_else(|| Value::Object(Map::new())), + ); request_metadata_json = json_bind_text(request_metadata_value.as_ref())?; } let _row = sqlx::query(UPSERT_SQL) @@ -8834,6 +8887,7 @@ ORDER BY "usage".user_id ASC let mut batch_rows = Vec::<(usize, PreparedPendingUsage)>::new(); let mut fallback_rows = Vec::<(usize, UpsertUsageRecord)>::new(); for (sequence, usage) in usages.into_iter().enumerate() { + let original_usage = usage.clone(); let prepared = PreparedPendingUsage::try_from_usage(usage)?; if request_id_counts .get(&prepared.usage.request_id) @@ -8843,7 +8897,9 @@ ORDER BY "usage".user_id ASC { batch_rows.push((sequence, prepared)); } else { - fallback_rows.push((sequence, prepared.usage)); + // Preserve capture markers for the canonical fallback; that path performs the + // sanitized bind only after preparing the auxiliary audit/blob state. + fallback_rows.push((sequence, original_usage)); } } @@ -9448,7 +9504,7 @@ removed_last_used_at_unix_secs, usage_created_at_unix_secs )); } - let usage = strip_deprecated_usage_display_fields(usage); + let usage = sanitize_usage_for_persistence(usage); let request_metadata_json = json_bind_text(usage.request_metadata.as_ref())?; let response_time_ms = usage.response_time_ms.map(to_i32).transpose()?; let first_byte_time_ms = usage.first_byte_time_ms.map(to_i32).transpose()?; @@ -9675,9 +9731,13 @@ DO UPDATE SET COALESCE(NULLIF(EXCLUDED.updated_at_unix_secs, 0), 0), CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) ) -WHERE "usage".billing_status = 'pending' + WHERE "usage".billing_status = 'pending' AND "usage".status IN ('pending', 'streaming') AND "usage".finalized_at IS NULL + AND EXCLUDED.updated_at_unix_secs >= COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) RETURNING request_id, provider_api_key_id, @@ -9736,8 +9796,14 @@ RETURNING apply_provider_monthly_usage_delta_in_tx(tx, provider_id.as_str(), *delta) .await?; } - for (node_id, delta) in &aggregates.proxy_nodes { - apply_proxy_node_counter_delta_in_tx(tx, node_id.as_str(), delta).await?; + for ((node_id, tunnel_generation), delta) in &aggregates.proxy_nodes { + apply_proxy_node_counter_delta_in_tx( + tx, + node_id.as_str(), + tunnel_generation.as_str(), + delta, + ) + .await?; } for (token_id, delta) in &aggregates.management_tokens { apply_management_token_counter_delta_in_tx(tx, token_id.as_str(), delta) @@ -9772,17 +9838,50 @@ RETURNING if delta.is_noop() { return Ok(false); } + let Some(expected_tunnel_generation) = delta + .expected_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .filter(|value| value.len() <= 64) + .map(ToOwned::to_owned) + else { + // Do not turn a bare node id into an implicit current-incarnation + // binding. Missing fences are rejected so stale plans fail closed. + return Ok(false); + }; let node_id = delta.node_id.trim().to_string(); let request_id = format!("proxy_node:{node_id}:{}", Uuid::new_v4()); self.tx_runner .run_read_write(|tx| { Box::pin(async move { + // Keep the parent lookup lock-free because flush claims + // outbox rows before updating proxy_nodes. The generation + // is persisted in the outbox row and checked again by the + // flush UPDATE, so id reuse can only retire this delta. + let tunnel_generation: Option = sqlx::query_scalar( + "SELECT tunnel_generation FROM proxy_nodes WHERE id = $1 AND tunnel_generation = $2 LIMIT 1", + ) + .bind(&node_id) + .bind(&expected_tunnel_generation) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(_tunnel_generation) = tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + return Ok(false); + }; insert_usage_counter_delta_in_tx( tx, UsageCounterDeltaInsert { request_id: &request_id, kind: USAGE_COUNTER_KIND_PROXY_NODE, target_id: &node_id, + target_tunnel_generation: Some(&expected_tunnel_generation), request_count_delta: 0, total_requests_delta: delta.total_requests_delta, success_count_delta: 0, @@ -9833,6 +9932,7 @@ RETURNING request_id: &request_id, kind: USAGE_COUNTER_KIND_MANAGEMENT_TOKEN, target_id: &token_id, + target_tunnel_generation: None, request_count_delta: delta.usage_count_delta, total_requests_delta: 0, success_count_delta: 0, @@ -9874,6 +9974,7 @@ RETURNING request_id: &request_id, kind: USAGE_COUNTER_KIND_API_KEY_LAST_USED, target_id: &api_key_id, + target_tunnel_generation: None, request_count_delta: 0, total_requests_delta: 0, success_count_delta: 0, @@ -9910,7 +10011,7 @@ RETURNING &self, cutoff_unix_secs: u64, now_unix_secs: u64, - timeout_minutes: u64, + _timeout_minutes: u64, batch_size: usize, ) -> Result { if batch_size == 0 { @@ -9963,7 +10064,6 @@ RETURNING .map(|row| { Ok(StalePendingUsageRow { request_id: row.try_get("request_id").map_postgres_err()?, - status: row.try_get("status").map_postgres_err()?, billing_status: row.try_get("billing_status").map_postgres_err()?, }) }) @@ -9996,18 +10096,7 @@ RETURNING .try_get::, _>("status_code") .map_postgres_err()? .and_then(|value| u16::try_from(value).ok()); - let error_message = row - .try_get::, _>("error_message") - .map_postgres_err()? - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - failed_map.insert( - request_id, - FailedCandidateCleanupInfo { - status_code, - error_message, - }, - ); + failed_map.insert(request_id, FailedCandidateCleanupInfo { status_code }); } (completed, failed_map) }; @@ -10030,23 +10119,23 @@ RETURNING } let candidate_info = failed_candidate_info.get(&row.request_id); - let (status_code, error_message) = - resolve_stale_pending_failure(candidate_info, &row.status, timeout_minutes); + let status_code = resolve_stale_pending_status_code(candidate_info); + let error_category = usage_error_category_for_status_code(status_code); let status_code_i32 = i32::from(status_code); if row.billing_status == "pending" { sqlx::query(UPDATE_FAILED_VOID_STALE_USAGE_SQL) .bind(&row.request_id) - .bind(&error_message) .bind(now) .bind(status_code_i32) + .bind(error_category) .execute(&mut *tx) .await .map_postgres_err()?; } else { sqlx::query(UPDATE_FAILED_STALE_USAGE_SQL) .bind(&row.request_id) - .bind(&error_message) .bind(status_code_i32) + .bind(error_category) .execute(&mut *tx) .await .map_postgres_err()?; @@ -10572,33 +10661,17 @@ impl UsageWriteRepository for SqlxUsageReadRepository { struct StalePendingUsageRow { request_id: String, - status: String, billing_status: String, } struct FailedCandidateCleanupInfo { status_code: Option, - error_message: Option, } -fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String { - format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成") -} - -fn resolve_stale_pending_failure( - candidate: Option<&FailedCandidateCleanupInfo>, - status: &str, - timeout_minutes: u64, -) -> (u16, String) { - match candidate { - Some(info) => ( - info.status_code.unwrap_or(502), - info.error_message - .clone() - .unwrap_or_else(|| stale_pending_error_message(status, timeout_minutes)), - ), - None => (504, stale_pending_error_message(status, timeout_minutes)), - } +fn resolve_stale_pending_status_code(candidate: Option<&FailedCandidateCleanupInfo>) -> u16 { + candidate + .and_then(|info| info.status_code) + .unwrap_or(if candidate.is_some() { 502 } else { 504 }) } async fn find_usage_by_request_id_in_tx( @@ -10768,6 +10841,7 @@ fn prepare_first_byte_provider_contribution_transitions( request_id: &transition.request_id, kind: USAGE_COUNTER_KIND_PROVIDER_API_KEY, target_id: &transition.key_id, + target_tunnel_generation: None, request_count_delta: transition.delta.request_count, total_requests_delta: 0, success_count_delta: transition.delta.success_count, @@ -10830,6 +10904,7 @@ struct UsageCounterDeltaRow { id: String, kind: String, target_id: String, + target_tunnel_generation: Option, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -10852,7 +10927,7 @@ struct UsageCounterDeltaAggregates { provider_api_keys: BTreeMap, models: BTreeMap, provider_monthly: BTreeMap, - proxy_nodes: BTreeMap, + proxy_nodes: BTreeMap<(String, String), ProxyNodeCounterDelta>, management_tokens: BTreeMap, api_key_last_used: BTreeMap, } @@ -10921,16 +10996,28 @@ impl UsageCounterDeltaAggregates { *entry += row.total_cost_usd_delta; } USAGE_COUNTER_KIND_PROXY_NODE => { - let entry = aggregates - .proxy_nodes - .entry(row.target_id.clone()) - .or_insert(ProxyNodeCounterDelta { + let Some(tunnel_generation) = row + .target_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + // Legacy rows without a generation fence are retired + // without applying them to a reused node id. + continue; + }; + let aggregate_key = (row.target_id.clone(), tunnel_generation.clone()); + let entry = aggregates.proxy_nodes.entry(aggregate_key).or_insert( + ProxyNodeCounterDelta { node_id: row.target_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), total_requests_delta: 0, failed_requests_delta: 0, dns_failures_delta: 0, stream_errors_delta: 0, - }); + }, + ); entry.total_requests_delta += row.total_requests_delta; entry.failed_requests_delta += row.error_count_delta; entry.dns_failures_delta += row.dns_failures_delta; @@ -11055,6 +11142,7 @@ async fn enqueue_api_key_usage_delta_in_tx( request_id, kind: USAGE_COUNTER_KIND_API_KEY, target_id: api_key_id, + target_tunnel_generation: None, request_count_delta: 0, total_requests_delta: delta.total_requests, success_count_delta: 0, @@ -11089,6 +11177,7 @@ async fn enqueue_model_usage_delta_in_tx( request_id, kind: USAGE_COUNTER_KIND_MODEL, target_id: model, + target_tunnel_generation: None, request_count_delta: delta.request_count, total_requests_delta: 0, success_count_delta: 0, @@ -11128,6 +11217,7 @@ async fn enqueue_provider_api_key_usage_delta_in_tx( request_id, kind: USAGE_COUNTER_KIND_PROVIDER_API_KEY, target_id: key_id, + target_tunnel_generation: None, request_count_delta: delta.request_count, total_requests_delta: 0, success_count_delta: delta.success_count, @@ -11151,6 +11241,7 @@ struct UsageCounterDeltaInsert<'a> { request_id: &'a str, kind: &'a str, target_id: &'a str, + target_tunnel_generation: Option<&'a str>, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -11173,6 +11264,7 @@ struct PreparedUsageCounterDeltaInsert { request_id: String, kind: String, target_id: String, + target_tunnel_generation: Option, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -11224,6 +11316,11 @@ fn prepare_usage_counter_delta_insert( request_id: request_id.to_string(), kind: input.kind.to_string(), target_id: target_id.to_string(), + target_tunnel_generation: input + .target_tunnel_generation + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), request_count_delta: input.request_count_delta, total_requests_delta: input.total_requests_delta, success_count_delta: input.success_count_delta, @@ -11258,6 +11355,7 @@ async fn insert_usage_counter_delta_in_tx( .bind(input.request_id) .bind(input.kind) .bind(input.target_id) + .bind(input.target_tunnel_generation) .bind(input.request_count_delta) .bind(input.total_requests_delta) .bind(input.success_count_delta) @@ -11290,6 +11388,7 @@ async fn insert_usage_counter_deltas_batch_in_tx( .push_bind(input.request_id.clone()) .push_bind(input.kind.clone()) .push_bind(input.target_id.clone()) + .push_bind(input.target_tunnel_generation.clone()) .push_bind(input.request_count_delta) .push_bind(input.total_requests_delta) .push_bind(input.success_count_delta) @@ -11340,6 +11439,9 @@ fn map_usage_counter_delta_row(row: &PgRow) -> Result("id").map_postgres_err()?, kind: row.try_get::("kind").map_postgres_err()?, target_id: row.try_get::("target_id").map_postgres_err()?, + target_tunnel_generation: row + .try_get::, _>("target_tunnel_generation") + .map_postgres_err()?, request_count_delta: row .try_get::("request_count_delta") .map_postgres_err()?, @@ -11565,9 +11667,10 @@ async fn apply_provider_monthly_usage_delta_in_tx( async fn apply_proxy_node_counter_delta_in_tx( tx: &mut sqlx::Transaction<'_, Postgres>, node_id: &str, + tunnel_generation: &str, delta: &ProxyNodeCounterDelta, ) -> Result<(), DataLayerError> { - if delta.is_noop() || node_id.trim().is_empty() { + if delta.is_noop() || node_id.trim().is_empty() || tunnel_generation.trim().is_empty() { return Ok(()); } @@ -11577,6 +11680,7 @@ async fn apply_proxy_node_counter_delta_in_tx( .bind(delta.failed_requests_delta) .bind(delta.dns_failures_delta) .bind(delta.stream_errors_delta) + .bind(tunnel_generation) .execute(&mut **tx) .await .map_postgres_err()?; @@ -12243,9 +12347,23 @@ fn json_bind_text(value: Option<&Value>) -> Result, DataLayerErro .transpose() } +fn project_usage_request_metadata( + value: Option, + preserve_empty_tombstone: bool, +) -> Option { + let projected = sanitize_usage_request_metadata(value); + if preserve_empty_tombstone && projected.is_none() { + Some(Value::Object(Map::new())) + } else { + projected + } +} + fn prepare_usage_upsert_context( usage: &UpsertUsageRecord, ) -> Result { + let usage = sanitize_usage_capture_controls_for_persistence(usage.clone()); + let usage = &usage; let replace_client_request_body_facts = request_body_capture_replaces_derived_facts( usage.request_body.as_ref(), usage.request_body_state, @@ -12382,6 +12500,14 @@ fn prepare_usage_upsert_context( clear_provider_request_body, )); } + // The raw event is used above to build the capture and billing snapshots, but only the + // allow-listed metadata may reach the accounting row. Keep an explicit empty object when a + // body `none` marker cleared the last derived fact: PostgreSQL's sparse upsert uses COALESCE + // and would otherwise resurrect the previous candidate's provider metadata from NULL. + let preserve_empty_metadata_tombstone = + (clear_request_body || clear_provider_request_body) && request_metadata_value.is_some(); + request_metadata_value = + project_usage_request_metadata(request_metadata_value, preserve_empty_metadata_tombstone); let http_audit_capture_mode = usage_http_audit_capture_mode( &http_audit_refs, [ @@ -12460,14 +12586,10 @@ fn resolved_read_usage_body_ref( http_audit_ref: Option<&str>, ) -> Option { explicit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) .or_else(|| { http_audit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) }) .or_else(|| has_compressed_storage.then(|| usage_body_ref(request_id, field))) .or_else(|| metadata_usage_body_ref_value(metadata, request_id, field)) @@ -12481,15 +12603,11 @@ fn resolved_write_usage_body_ref( http_audit_ref: Option<&str>, ) -> Option { explicit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) .or_else(|| has_compressed_storage.then(|| usage_body_ref(request_id, field))) .or_else(|| { http_audit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) }) } @@ -12511,11 +12629,7 @@ fn metadata_usage_body_ref_value( field: UsageBodyField, ) -> Option { metadata_ref_value(metadata, field.as_ref_key()) - .and_then(|value| parse_usage_body_ref(&value)) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field)) + .and_then(|body_ref| canonical_usage_body_ref_for(&body_ref, request_id, field)) } fn metadata_number_value( @@ -13076,11 +13190,7 @@ fn usage_json_column( } fn inflate_usage_json_value(bytes: &[u8]) -> Result { - let mut decoder = GzDecoder::new(bytes); - let mut json_bytes = Vec::new(); - decoder.read_to_end(&mut json_bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to decompress usage json: {err}")) - })?; + let json_bytes = read_decompressed_usage_json(GzDecoder::new(bytes))?; serde_json::from_slice(&json_bytes).map_err(|err| { DataLayerError::UnexpectedValue(format!("failed to parse decompressed usage json: {err}")) }) @@ -13252,42 +13362,19 @@ async fn sync_usage_body_blob_storage<'e, E>( executor: E, request_id: &str, field: UsageBodyField, - value: Option<&Value>, - storage: &UsageBodyStorage, - clear_existing: bool, + _value: Option<&Value>, + _storage: &UsageBodyStorage, + _clear_existing: bool, ) -> Result<(), DataLayerError> where E: sqlx::Executor<'e, Database = Postgres>, { let body_ref = usage_body_ref(request_id, field); - if clear_existing { - sqlx::query(DELETE_USAGE_BODY_BLOB_SQL) - .bind(&body_ref) - .execute(executor) - .await - .map_postgres_err()?; - return Ok(()); - } - if let Some(payload_gzip) = storage.detached_blob_bytes.as_ref() { - sqlx::query(UPSERT_USAGE_BODY_BLOB_SQL) - .bind(&body_ref) - .bind(request_id) - .bind(field.as_storage_field()) - .bind(payload_gzip) - .execute(executor) - .await - .map_postgres_err()?; - return Ok(()); - } - - if value.is_some() { - sqlx::query(DELETE_USAGE_BODY_BLOB_SQL) - .bind(&body_ref) - .execute(executor) - .await - .map_postgres_err()?; - } - + sqlx::query(DELETE_USAGE_BODY_BLOB_SQL) + .bind(&body_ref) + .execute(executor) + .await + .map_postgres_err()?; Ok(()) } @@ -13296,46 +13383,43 @@ async fn sync_usage_http_audit_storage<'e, E>( request_id: &str, headers: &UsageHttpAuditHeaders<'_>, refs: &UsageHttpAuditRefs, - states: &UsageHttpAuditStates, + _states: &UsageHttpAuditStates, body_capture_mode: &str, ) -> Result<(), DataLayerError> where E: sqlx::Executor<'e, Database = Postgres>, { - if !headers.any_present() - && !refs.any_present() - && !states.any_present() - && body_capture_mode == "none" - { - return Ok(()); + if headers.any_present() || refs.any_present() || body_capture_mode != "none" { + return Err(DataLayerError::InvalidInput( + "usage HTTP capture persistence is disabled".to_string(), + )); } - sqlx::query(UPSERT_USAGE_HTTP_AUDIT_SQL) - .bind(request_id) - .bind(headers.request_headers_json) - .bind(headers.provider_request_headers_json) - .bind(headers.response_headers_json) - .bind(headers.client_response_headers_json) - .bind(refs.request_body_ref.as_deref()) - .bind(refs.provider_request_body_ref.as_deref()) - .bind(refs.response_body_ref.as_deref()) - .bind(refs.client_response_body_ref.as_deref()) - .bind(usage_body_capture_state_bind_text( - states.request_body_state, - )) - .bind(usage_body_capture_state_bind_text( - states.provider_request_body_state, - )) - .bind(usage_body_capture_state_bind_text( - states.response_body_state, - )) - .bind(usage_body_capture_state_bind_text( - states.client_response_body_state, - )) - .bind(body_capture_mode) - .execute(executor) - .await - .map_postgres_err()?; + sqlx::query( + r#" +WITH deleted_audit AS ( + DELETE FROM usage_http_audits WHERE request_id = $1 +) +UPDATE usage +SET request_headers = NULL, + request_body = NULL, + provider_request_headers = NULL, + provider_request_body = NULL, + response_headers = NULL, + response_body = NULL, + client_response_headers = NULL, + client_response_body = NULL, + request_body_compressed = NULL, + provider_request_body_compressed = NULL, + response_body_compressed = NULL, + client_response_body_compressed = NULL +WHERE request_id = $1 +"#, + ) + .bind(request_id) + .execute(executor) + .await + .map_postgres_err()?; Ok(()) } 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 2d408bbf5..a5252794e 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 @@ -221,11 +221,13 @@ SELECT usage_settlement_snapshots.billing_rule_id AS settlement_billing_rule_id, usage_settlement_snapshots.billing_rule_version AS settlement_billing_rule_version, CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) AS created_at_unix_ms, - GREATEST( - COALESCE(NULLIF("usage".updated_at_unix_secs, 0), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), - CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + GREATEST( + COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), + COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) ) AS updated_at_unix_secs, CAST( EXTRACT( diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/find_by_request_id_sql.sql b/crates/aether-data/adapters/postgres/src/usage/queries/find_by_request_id_sql.sql index 14147c074..cbddf8d47 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/find_by_request_id_sql.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/find_by_request_id_sql.sql @@ -225,11 +225,13 @@ SELECT usage_settlement_snapshots.billing_rule_id AS settlement_billing_rule_id, usage_settlement_snapshots.billing_rule_version AS settlement_billing_rule_version, CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) AS created_at_unix_ms, - GREATEST( - COALESCE(NULLIF("usage".updated_at_unix_secs, 0), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), - CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + GREATEST( + COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), + COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) ) AS updated_at_unix_secs, CAST( EXTRACT( 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 f4515ac2d..d0909bcd2 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 @@ -327,11 +327,13 @@ SELECT usage_settlement_snapshots.billing_rule_id AS settlement_billing_rule_id, usage_settlement_snapshots.billing_rule_version AS settlement_billing_rule_version, CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) AS created_at_unix_ms, - GREATEST( - COALESCE(NULLIF("usage".updated_at_unix_secs, 0), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), - CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + GREATEST( + COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), + COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) ) AS updated_at_unix_secs, CAST( EXTRACT( 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 f4515ac2d..d0909bcd2 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 @@ -327,11 +327,13 @@ SELECT usage_settlement_snapshots.billing_rule_id AS settlement_billing_rule_id, usage_settlement_snapshots.billing_rule_version AS settlement_billing_rule_version, CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) AS created_at_unix_ms, - GREATEST( - COALESCE(NULLIF("usage".updated_at_unix_secs, 0), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), - COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), - CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + GREATEST( + COALESCE(CAST(EXTRACT(EPOCH FROM usage_settlement_snapshots.finalized_at) AS BIGINT), 0), + COALESCE(CAST(EXTRACT(EPOCH FROM "usage".finalized_at) AS BIGINT), 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) ) AS updated_at_unix_secs, CAST( EXTRACT( diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/upsert_first_byte_sql.sql b/crates/aether-data/adapters/postgres/src/usage/queries/upsert_first_byte_sql.sql index 5c44c6389..e30ac97ec 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/upsert_first_byte_sql.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/upsert_first_byte_sql.sql @@ -119,6 +119,10 @@ DO UPDATE SET WHERE "usage".billing_status = 'pending' AND "usage".status IN ('pending', 'streaming') AND "usage".finalized_at IS NULL + AND EXCLUDED.updated_at_unix_secs >= COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + ) RETURNING request_id, provider_api_key_id, diff --git a/crates/aether-data/adapters/postgres/src/usage/tests.rs b/crates/aether-data/adapters/postgres/src/usage/tests.rs index a6fab483d..7dc47bd80 100644 --- a/crates/aether-data/adapters/postgres/src/usage/tests.rs +++ b/crates/aether-data/adapters/postgres/src/usage/tests.rs @@ -187,6 +187,156 @@ async fn pending_batch_is_opt_in_and_rejects_non_pending_before_connecting() { .contains("pending usage batch requires pending status")); } +#[tokio::test] +#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] +async fn live_stale_terminal_event_is_a_full_transaction_noop() { + let database_url = std::env::var("AETHER_TEST_DATABASE_URL") + .expect("AETHER_TEST_DATABASE_URL must point at the test database"); + let factory = PostgresPoolFactory::new(PostgresPoolConfig { + database_url, + min_connections: 1, + max_connections: 2, + acquire_timeout_ms: 10_000, + idle_timeout_ms: 30_000, + max_lifetime_ms: 60_000, + statement_cache_capacity: 64, + require_ssl: false, + }) + .expect("factory should build"); + let repository = + SqlxUsageReadRepository::new(factory.connect_lazy().expect("lazy pool should build")); + crate::run_migrations(repository.pool()) + .await + .expect("test database migrations should succeed"); + + let suffix = uuid::Uuid::new_v4().simple().to_string(); + let request_id = format!("req-stale-terminal-{suffix}"); + let provider_name = format!("stale-provider-{suffix}"); + let now_unix_secs = Utc::now().timestamp().max(2) as u64; + let mut newer = fast_clear_usage_record( + &request_id, + &provider_name, + now_unix_secs, + true, + UsageBodyCaptureState::None, + None, + ); + newer.candidate_id = Some("candidate-new".to_string()); + newer.route_kind = Some("route-new".to_string()); + newer.total_cost_usd = Some(0.5); + newer.actual_total_cost_usd = Some(0.4); + repository + .upsert(newer) + .await + .expect("newer terminal usage should upsert"); + + let counter_rows_before: i64 = sqlx::query_scalar( + "SELECT COUNT(*)::BIGINT FROM usage_counter_deltas WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("counter rows should count"); + let routing_before = sqlx::query( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("routing snapshot should load"); + let routing_before = ( + routing_before + .try_get::, _>("candidate_id") + .unwrap(), + routing_before + .try_get::, _>("route_kind") + .unwrap(), + ); + let settlement_before = sqlx::query( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("settlement snapshot should load"); + let settlement_before = ( + settlement_before + .try_get::("billing_status") + .unwrap(), + settlement_before + .try_get::, _>("billing_total_cost_usd") + .unwrap(), + ); + + let mut stale = fast_clear_usage_record( + &request_id, + &provider_name, + now_unix_secs - 2, + true, + UsageBodyCaptureState::None, + None, + ); + stale.status = "failed".to_string(); + stale.billing_status = "void".to_string(); + stale.status_code = Some(503); + stale.total_cost_usd = Some(99.0); + stale.actual_total_cost_usd = Some(98.0); + stale.candidate_id = Some("candidate-stale".to_string()); + stale.route_kind = Some("route-stale".to_string()); + let stored = repository + .upsert(stale) + .await + .expect("stale terminal usage should be ignored"); + + assert_eq!(stored.status, "completed"); + assert_eq!(stored.billing_status, "pending"); + assert_eq!(stored.status_code, Some(200)); + assert_eq!(stored.total_cost_usd, 0.5); + assert_eq!(stored.routing_candidate_id(), Some("candidate-new")); + assert_eq!(stored.routing_route_kind(), Some("route-new")); + + let counter_rows_after: i64 = sqlx::query_scalar( + "SELECT COUNT(*)::BIGINT FROM usage_counter_deltas WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("counter rows should count"); + let routing_after = sqlx::query( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("routing snapshot should load"); + let routing_after = ( + routing_after + .try_get::, _>("candidate_id") + .unwrap(), + routing_after + .try_get::, _>("route_kind") + .unwrap(), + ); + let settlement_after = sqlx::query( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1", + ) + .bind(&request_id) + .fetch_one(repository.pool()) + .await + .expect("settlement snapshot should load"); + let settlement_after = ( + settlement_after + .try_get::("billing_status") + .unwrap(), + settlement_after + .try_get::, _>("billing_total_cost_usd") + .unwrap(), + ); + assert_eq!(counter_rows_after, counter_rows_before); + assert_eq!(routing_after, routing_before); + assert_eq!(settlement_after, settlement_before); +} + #[tokio::test] #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conflicts() { @@ -385,37 +535,14 @@ async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conf "late pending must not clear the first-byte observation" ); - let http = sqlx::query( - "SELECT request_headers, provider_request_headers, response_headers, client_response_headers, request_body_ref, provider_request_body_ref, response_body_ref, client_response_body_ref, request_body_state, provider_request_body_state, response_body_state, client_response_body_state FROM usage_http_audits WHERE request_id = $1", + let http_count = sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*)::BIGINT FROM usage_http_audits WHERE request_id = $1", ) .bind(&rich_request_id) .fetch_one(repository.pool()) .await - .expect("rich HTTP audit should exist"); - assert_eq!( - http.try_get::("request_headers") - .unwrap(), - json!({"x-request": "request-value"}) - ); - for field in [ - "request_body_ref", - "provider_request_body_ref", - "response_body_ref", - "client_response_body_ref", - ] { - assert!(http.try_get::, _>(field).unwrap().is_some()); - } - for field in [ - "request_body_state", - "provider_request_body_state", - "response_body_state", - "client_response_body_state", - ] { - assert_eq!( - http.try_get::, _>(field).unwrap().as_deref(), - Some("reference") - ); - } + .expect("HTTP audit count should be readable"); + assert_eq!(http_count, 0); let blob_count = sqlx::query_scalar::<_, i64>( "SELECT COUNT(*)::BIGINT FROM usage_body_blobs WHERE request_id = $1", ) @@ -423,7 +550,7 @@ async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conf .fetch_one(repository.pool()) .await .expect("body blob count should be readable"); - assert_eq!(blob_count, 4); + assert_eq!(blob_count, 0); let routing = sqlx::query( "SELECT candidate_id, candidate_index, selected_provider_api_key_id FROM usage_routing_snapshots WHERE request_id = $1", @@ -2264,6 +2391,22 @@ fn usage_sql_does_not_require_updated_at_column() { assert!(!super::UPSERT_SQL.contains("updated_at = CASE")); } +#[test] +fn usage_sql_preserves_nonzero_lifecycle_updated_revision() { + for sql in [ + super::FIND_BY_REQUEST_ID_SQL, + super::FIND_BY_ID_SQL, + super::LIST_USAGE_AUDITS_PREFIX, + super::LIST_RECENT_USAGE_AUDITS_PREFIX, + ] { + assert!(sql + .contains("COALESCE(\n NULLIF(\"usage\".updated_at_unix_secs, 0),\n GREATEST(")); + assert!( + !sql.contains("GREATEST(\n COALESCE(NULLIF(\"usage\".updated_at_unix_secs, 0), 0),") + ); + } +} + #[test] fn usage_sql_summarizes_tokens_by_api_key_ids_in_database() { let sql = super::SUMMARIZE_TOTAL_TOKENS_BY_API_KEY_IDS_SQL; @@ -3586,6 +3729,7 @@ fn usage_sql_clears_stale_failure_fields_for_non_failed_status_updates() { fn stale_cleanup_failed_candidate_sql_orders_by_effective_timestamp() { let sql = super::SELECT_LATEST_FAILED_CANDIDATE_FOR_STALE_REQUESTS_SQL; assert!(sql.contains("COALESCE(finished_at, started_at, created_at) DESC")); + assert!(!sql.contains("error_message")); assert!(!sql.contains("finished_at DESC NULLS LAST")); assert!(!sql.contains("started_at DESC NULLS LAST")); } @@ -3609,6 +3753,10 @@ fn usage_sql_does_not_allow_streaming_to_regress_back_to_pending() { #[test] fn first_byte_upsert_sql_is_single_row_guarded_and_preserves_existing_metadata() { let sql = normalize_newlines(super::UPSERT_FIRST_BYTE_SQL); + let revision_guard = r#"AND EXCLUDED.updated_at_unix_secs >= COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + CAST(EXTRACT(EPOCH FROM "usage".created_at) AS BIGINT) + )"#; assert_eq!(sql.matches("INSERT INTO").count(), 1); assert!(!sql.contains("usage_http_audits")); assert!(!sql.contains("usage_routing_snapshots")); @@ -3619,6 +3767,8 @@ fn first_byte_upsert_sql_is_single_row_guarded_and_preserves_existing_metadata() assert!(sql.contains("WHERE \"usage\".billing_status = 'pending'")); assert!(sql.contains("\"usage\".status IN ('pending', 'streaming')")); assert!(sql.contains("\"usage\".finalized_at IS NULL")); + assert!(sql.contains(revision_guard)); + assert!(normalize_newlines(include_str!("mod.rs")).contains(revision_guard)); assert!(sql.contains("$22::json->>'upstream_is_stream'")); assert!(sql.contains("\"usage\".upstream_is_stream")); @@ -3788,6 +3938,7 @@ fn first_byte_provider_counter_batch_prepares_all_columns_before_query_building( request_id: " req-counter-prepared ", kind: "provider_api_key", target_id: " key-counter-prepared ", + target_tunnel_generation: None, request_count_delta: 1, total_requests_delta: 2, success_count_delta: 3, @@ -3841,6 +3992,7 @@ fn first_byte_provider_counter_batch_prepares_all_columns_before_query_building( request_id: "req-counter-out-of-range", kind: "provider_api_key", target_id: "key-counter-out-of-range", + target_tunnel_generation: None, request_count_delta: 0, total_requests_delta: 0, success_count_delta: 0, @@ -4284,6 +4436,39 @@ fn resolved_read_usage_body_ref_prefers_typed_then_http_audit_then_compressed_th ), Some("usage://request/req-123/client_response_body".to_string()) ); + assert_eq!( + resolved_read_usage_body_ref( + Some("usage://request/req-other/request_body"), + None, + "req-123", + UsageBodyField::RequestBody, + false, + Some("usage://request/req-123/request_body"), + ), + Some(usage_body_ref("req-123", UsageBodyField::RequestBody)) + ); + assert_eq!( + resolved_read_usage_body_ref( + None, + None, + "req-123", + UsageBodyField::RequestBody, + false, + Some("usage://request/req-other/request_body"), + ), + None + ); + assert_eq!( + resolved_read_usage_body_ref( + None, + None, + "req-123", + UsageBodyField::RequestBody, + false, + Some("usage://request/req-123/response_body"), + ), + None + ); } #[test] @@ -4322,6 +4507,26 @@ fn resolved_write_usage_body_ref_ignores_metadata_compatibility_keys() { ), Some("usage://request/req-123/client_response_body".to_string()) ); + assert_eq!( + resolved_write_usage_body_ref( + Some("usage://request/req-other/request_body"), + "req-123", + UsageBodyField::RequestBody, + false, + Some("usage://request/req-123/request_body"), + ), + Some(usage_body_ref("req-123", UsageBodyField::RequestBody)) + ); + assert_eq!( + resolved_write_usage_body_ref( + Some("usage://request/req-123/response_body"), + "req-123", + UsageBodyField::RequestBody, + false, + Some("usage://request/req-other/request_body"), + ), + None + ); } #[test] diff --git a/crates/aether-data/adapters/postgres/src/users.rs b/crates/aether-data/adapters/postgres/src/users.rs index e1d5ed2fa..f11afc6f3 100644 --- a/crates/aether-data/adapters/postgres/src/users.rs +++ b/crates/aether-data/adapters/postgres/src/users.rs @@ -3,11 +3,14 @@ use futures_util::TryStreamExt; use sqlx::{PgPool, Postgres, QueryBuilder, Row}; use aether_data_contracts::repository::users::{ - normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, + is_valid_bcrypt_hash, last_oauth_unbind_denial, normalize_user_group_name, + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy, - UserExportSummary, UserReadRepository, + UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED, + LAST_ACTIVE_ADMIN_UPDATE_DENIED, }; use aether_data_contracts::DataLayerError; @@ -26,6 +29,92 @@ WHERE id = ANY($1::text[]) ORDER BY id ASC "#; +const POSTGRES_LOCK_ACTIVE_ADMINS_SQL: &str = r#" +SELECT id +FROM users +WHERE role = 'admin'::userrole + AND is_active IS TRUE + AND is_deleted IS FALSE +ORDER BY id +FOR UPDATE +"#; + +const POSTGRES_DELETE_USER_IF_WALLET_ABSENT_SQL: &str = r#" +DELETE FROM users +WHERE id = $1 + AND NOT EXISTS ( + SELECT 1 + FROM wallets AS wallet + WHERE wallet.user_id = $2 + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = $3 + ) + ) +"#; + +const POSTGRES_DELETE_USER_API_KEYS_SQL: &str = "DELETE FROM api_keys WHERE user_id = $1"; + +const POSTGRES_DELETE_USER_DEPENDENTS_SQL: &[&str] = &[ + "DELETE FROM usage_request_admissions WHERE subject_id = $1", + "DELETE FROM usage_cost_reservations WHERE subject_id = $1", + "DELETE FROM gemini_file_mappings WHERE user_id = $1", + "DELETE FROM api_key_provider_mappings WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1)", + POSTGRES_DELETE_USER_API_KEYS_SQL, + "DELETE FROM management_tokens WHERE user_id = $1", + "DELETE FROM user_sessions WHERE user_id = $1", + "DELETE FROM user_oauth_links WHERE user_id = $1", + "DELETE FROM user_group_members WHERE user_id = $1", + "DELETE FROM user_preferences WHERE user_id = $1", + "DELETE FROM user_invite_codes WHERE user_id = $1", + "DELETE FROM announcement_reads WHERE user_id = $1", +]; + +const POSTGRES_PREPARE_USER_FACTS_FOR_DELETION_SQL: &[&str] = &[ + "UPDATE referral_rewards SET status = CASE WHEN status IN ('pending', 'failed', 'applying') THEN 'voided' ELSE status END, failure_reason = NULL, admin_note = NULL, updated_at = NOW() WHERE $1 IN (inviter_user_id, invitee_user_id)", + "UPDATE referral_rewards SET failure_reason = NULL, admin_note = NULL, updated_at = NOW() WHERE admin_operator_id = $1", + "UPDATE user_referrals SET invite_code_snapshot = 'deleted-user', source_json = NULL, updated_at = NOW() WHERE $1 IN (inviter_user_id, invitee_user_id)", + "UPDATE user_plan_entitlements SET status = CASE WHEN status = 'active' THEN 'revoked' ELSE status END, expires_at = LEAST(expires_at, NOW()), updated_at = NOW() WHERE user_id = $1", + "UPDATE wallets SET status = 'disabled', updated_at = NOW() WHERE user_id = $1", + "UPDATE wallets SET status = 'disabled', updated_at = NOW() WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1)", + "UPDATE audit_logs SET description = 'deleted user event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE user_id = $1", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = $1)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1))", + "UPDATE wallet_transactions SET description = NULL WHERE operator_id = $1", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order WHERE history_order.user_id = $1 AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.user_id = $1 AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1) AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE user_id = $1", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1))", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE user_id = $1", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE $1 IN (requested_by, approved_by, processed_by)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = $1)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1))", + "UPDATE redeem_code_batches SET description = NULL WHERE created_by = $1", +]; + +const POSTGRES_ANONYMIZE_USER_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET username = NULL, api_key_name = NULL WHERE user_id = $1", + "UPDATE video_tasks SET username = NULL, api_key_name = NULL WHERE user_id = $1", + "UPDATE usage SET username = NULL, api_key_name = NULL WHERE user_id = $1", + "UPDATE stats_user_daily SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_summary SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_model SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_provider SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_api_format SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_model_provider SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_cost_savings SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_cost_savings_provider SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_cost_savings_model SET username = NULL WHERE user_id = $1", + "UPDATE stats_user_daily_cost_savings_model_provider SET username = NULL WHERE user_id = $1", +]; + +const POSTGRES_ANONYMIZE_USER_API_KEY_HISTORY_SQL: &str = + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = $1)"; + const LIST_USERS_BY_USERNAME_SEARCH_SQL: &str = r#" SELECT id, @@ -131,14 +220,14 @@ WHERE role = 'admin'::userrole AND is_active IS TRUE "#; -const COUNT_ACTIVE_LOCAL_ADMIN_USERS_WITH_VALID_PASSWORD_SQL: &str = r#" -SELECT COUNT(*)::BIGINT AS total +const LIST_ACTIVE_LOCAL_ADMIN_PASSWORD_HASHES_SQL: &str = r#" +SELECT password_hash FROM users WHERE role = 'admin'::userrole AND auth_source = 'local'::authsource AND is_deleted IS FALSE AND is_active IS TRUE - AND password_hash ~ '^\$2[aby]\$\d{2}\$.{53}$' + AND password_hash IS NOT NULL "#; const FIND_EXPORT_USER_BY_ID_SQL: &str = r#" @@ -184,6 +273,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -208,6 +298,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -232,6 +323,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -256,6 +348,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -280,6 +373,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -305,6 +399,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -345,6 +440,7 @@ SELECT users.allowed_models_mode, users.is_active, users.is_deleted, + users.security_version, users.created_at, users.last_login_at FROM user_oauth_links @@ -405,12 +501,7 @@ INSERT INTO user_oauth_links ( last_login_at ) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $8) -ON CONFLICT (user_id, provider_type) DO UPDATE -SET provider_user_id = EXCLUDED.provider_user_id, - provider_username = EXCLUDED.provider_username, - provider_email = EXCLUDED.provider_email, - extra_data = EXCLUDED.extra_data, - last_login_at = EXCLUDED.last_login_at +ON CONFLICT DO NOTHING "#; const TOUCH_AUTH_USER_LAST_LOGIN_SQL: &str = r#" @@ -514,7 +605,7 @@ LEFT JOIN providers p const FIND_USER_SESSION_SQL: &str = r#" SELECT - id, user_id, client_device_id, device_label, refresh_token_hash, + id, user_id, security_version, client_device_id, device_label, refresh_token_hash, prev_refresh_token_hash, rotated_at, last_seen_at, expires_at, revoked_at, revoke_reason, ip_address, user_agent, created_at, updated_at FROM user_sessions @@ -524,7 +615,7 @@ LIMIT 1 const LIST_USER_SESSIONS_SQL: &str = r#" SELECT - id, user_id, client_device_id, device_label, refresh_token_hash, + id, user_id, security_version, client_device_id, device_label, refresh_token_hash, prev_refresh_token_hash, rotated_at, last_seen_at, expires_at, revoked_at, revoke_reason, ip_address, user_agent, created_at, updated_at FROM user_sessions @@ -545,12 +636,12 @@ WHERE user_id = $1 const CREATE_USER_SESSION_SQL: &str = r#" INSERT INTO user_sessions ( - id, user_id, client_device_id, device_label, device_type, ip_address, - user_agent, refresh_token_hash, last_seen_at, expires_at, created_at, updated_at + id, user_id, security_version, client_device_id, device_label, device_type, + ip_address, user_agent, refresh_token_hash, last_seen_at, expires_at, created_at, updated_at ) -VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12) +VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13) RETURNING - id, user_id, client_device_id, device_label, refresh_token_hash, + id, user_id, security_version, client_device_id, device_label, refresh_token_hash, prev_refresh_token_hash, rotated_at, last_seen_at, expires_at, revoked_at, revoke_reason, ip_address, user_agent, created_at, updated_at "#; @@ -580,7 +671,11 @@ SET prev_refresh_token_hash = $3, ip_address = COALESCE($7, ip_address), user_agent = COALESCE($8, user_agent), updated_at = $4 -WHERE user_id = $1 AND id = $2 +WHERE user_id = $1 + AND id = $2 + AND refresh_token_hash = $3 + AND revoked_at IS NULL + AND expires_at > $4 "#; const REVOKE_USER_SESSION_SQL: &str = r#" @@ -827,6 +922,94 @@ WHERE id = $1 } } + /// Restore a group under a row lock so the snapshot comparison and write + /// cannot be separated by a concurrent administrator update. + pub async fn restore_user_group_if_matches( + &self, + expected: &StoredUserGroup, + restored: &StoredUserGroup, + ) -> Result { + if expected.id != restored.id || expected.id.trim().is_empty() { + return Ok(false); + } + + let mut tx = self.pool.begin().await.map_postgres_err()?; + let mut builder = QueryBuilder::::new(USER_GROUP_COLUMNS); + builder + .push(" WHERE id = ") + .push_bind(&expected.id) + .push(" FOR UPDATE"); + let row = builder + .build() + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let current = map_user_group_row(&row)?; + if ¤t != expected { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE user_groups +SET name = $2, + normalized_name = $3, + description = $4, + priority = $5, + allowed_providers = $6::json, + allowed_providers_mode = $7, + allowed_api_formats = $8::json, + allowed_api_formats_mode = $9, + allowed_models = $10::json, + allowed_models_mode = $11, + rate_limit = $12, + rate_limit_mode = $13, + created_at = $14, + updated_at = $15 +WHERE id = $1 +"#, + ) + .bind(&restored.id) + .bind(&restored.name) + .bind(&restored.normalized_name) + .bind(&restored.description) + .bind(restored.priority) + .bind( + restored + .allowed_providers + .clone() + .map(serde_json::Value::from), + ) + .bind(&restored.allowed_providers_mode) + .bind( + restored + .allowed_api_formats + .clone() + .map(serde_json::Value::from), + ) + .bind(&restored.allowed_api_formats_mode) + .bind(restored.allowed_models.clone().map(serde_json::Value::from)) + .bind(&restored.allowed_models_mode) + .bind(restored.rate_limit) + .bind(&restored.rate_limit_mode) + .bind(restored.created_at) + .bind(restored.updated_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + pub async fn delete_user_group(&self, group_id: &str) -> Result { let result = sqlx::query("DELETE FROM user_groups WHERE id = $1") .bind(group_id) @@ -854,6 +1037,26 @@ WHERE id = $1 user_ids: &[String], ) -> Result, DataLayerError> { let mut tx = self.pool.begin().await.map_postgres_err()?; + // Serialize membership replacement with the per-user CAS path by locking all affected + // users in deterministic order before deleting or inserting membership rows. + let mut locked_user_ids = normalized_ids(user_ids); + let existing_user_ids = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_group_members WHERE group_id = $1 ORDER BY user_id", + ) + .bind(group_id) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + locked_user_ids.extend(existing_user_ids); + locked_user_ids.sort(); + locked_user_ids.dedup(); + if !locked_user_ids.is_empty() { + sqlx::query("SELECT id FROM users WHERE id = ANY($1::text[]) ORDER BY id FOR UPDATE") + .bind(&locked_user_ids) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + } sqlx::query("DELETE FROM user_group_members WHERE group_id = $1") .bind(group_id) .execute(&mut *tx) @@ -927,6 +1130,16 @@ WHERE user_group_members.user_id IN ( group_ids: &[String], ) -> Result, DataLayerError> { let mut tx = self.pool.begin().await.map_postgres_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(Vec::new()); + } sqlx::query("DELETE FROM user_group_members WHERE user_id = $1") .bind(user_id) .execute(&mut *tx) @@ -946,19 +1159,92 @@ WHERE user_group_members.user_id IN ( self.list_user_groups_for_user(user_id).await } + pub async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + let expected = normalized_ids(expected_group_ids); + let restored = normalized_ids(restored_group_ids); + let mut tx = self.pool.begin().await.map_postgres_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let current = sqlx::query_scalar::<_, String>( + "SELECT group_id FROM user_group_members WHERE user_id = $1 ORDER BY group_id ASC FOR UPDATE", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + if current != expected { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + if !restored.is_empty() { + let count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM user_groups WHERE id = ANY($1::text[])") + .bind(&restored) + .fetch_one(&mut *tx) + .await + .map_postgres_err()?; + if count != restored.len() as i64 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + } + sqlx::query("DELETE FROM user_group_members WHERE user_id = $1") + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + for group_id in restored { + sqlx::query( + "INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING", + ) + .bind(group_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + pub async fn add_user_to_group( &self, group_id: &str, user_id: &str, ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } let result = sqlx::query( "INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING", ) .bind(group_id) .bind(user_id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(result.rows_affected() > 0) } @@ -1211,6 +1497,70 @@ WHERE user_group_members.user_id IN ( row.as_ref().map(map_user_auth_row).transpose() } + #[allow(clippy::too_many_arguments)] + pub async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: chrono::DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let provider_enabled: Option = sqlx::query_scalar( + "SELECT is_enabled FROM oauth_providers WHERE provider_type = $1 FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if provider_enabled != Some(true) { + tx.rollback().await.map_postgres_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::ProviderUnavailable); + } + let row = sqlx::query(&format!( + "{FIND_OAUTH_LINKED_USER_SQL} FOR UPDATE OF users, user_oauth_links" + )) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::NotLinked); + }; + let mut user = map_user_auth_row(&row)?; + sqlx::query(TOUCH_OAUTH_LINK_SQL) + .bind(provider_type) + .bind(provider_user_id) + .bind(provider_username) + .bind(provider_email) + .bind(extra_data) + .bind(touched_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if let Some(verified_email) = verified_email { + let result = sqlx::query( + "UPDATE users SET email_verified = TRUE, updated_at = $3 WHERE id = $1 AND email_verified IS FALSE AND LOWER(TRIM(email)) = LOWER(TRIM($2))", + ) + .bind(&user.id) + .bind(verified_email) + .bind(touched_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() == 1 { + user.email_verified = true; + } + } + tx.commit().await.map_postgres_err()?; + Ok(ResolveOAuthLinkedUserOutcome::Linked(user)) + } + pub async fn touch_oauth_link( &self, provider_type: &str, @@ -1236,6 +1586,7 @@ WHERE user_group_members.user_id IN ( pub async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: chrono::DateTime, ) -> Result, DataLayerError> { @@ -1248,14 +1599,15 @@ INSERT INTO users ( is_active, is_deleted, created_at, updated_at, last_login_at ) VALUES ( - $1, $2, TRUE, $3, NULL, 'user'::userrole, 'oauth'::authsource, + $1, $2, $3, $4, NULL, 'user'::userrole, 'oauth'::authsource, 'inherit', 'inherit', 'inherit', 'inherit', - TRUE, FALSE, $4, $4, $4 + TRUE, FALSE, $5, $5, $5 ) "#, ) .bind(&user_id) .bind(email) + .bind(email_verified) .bind(username) .bind(created_at) .execute(&self.pool) @@ -1304,7 +1656,7 @@ VALUES ( } #[allow(clippy::too_many_arguments)] - pub async fn upsert_user_oauth_link( + pub async fn bind_user_oauth_link( &self, user_id: &str, provider_type: &str, @@ -1313,8 +1665,64 @@ VALUES ( provider_email: Option<&str>, extra_data: Option, linked_at: chrono::DateTime, - ) -> Result<(), DataLayerError> { - sqlx::query(UPSERT_OAUTH_LINK_SQL) + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let provider_enabled = sqlx::query_scalar::<_, bool>( + "SELECT is_enabled FROM oauth_providers WHERE provider_type = $1 FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(provider_enabled) = provider_enabled else { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderNotFound); + }; + if !provider_enabled { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderDisabled); + } + let user_exists = + sqlx::query_scalar::<_, i32>("SELECT 1 FROM users WHERE id = $1 FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + .is_some(); + if !user_exists { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::UserNotFound); + } + if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = $1 AND provider_user_id = $2 LIMIT 1", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + { + tx.rollback().await.map_postgres_err()?; + return Ok(if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }); + } + if sqlx::query_scalar::<_, i32>( + "SELECT 1 FROM user_oauth_links WHERE user_id = $1 AND provider_type = $2 LIMIT 1", + ) + .bind(user_id) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + .is_some() + { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider); + } + let inserted = sqlx::query(UPSERT_OAUTH_LINK_SQL) .bind(uuid::Uuid::new_v4().to_string()) .bind(user_id) .bind(provider_type) @@ -1323,24 +1731,143 @@ VALUES ( .bind(provider_email) .bind(extra_data) .bind(linked_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; - Ok(()) + let outcome = if inserted.rows_affected() == 1 { + BindUserOAuthLinkOutcome::Bound + } else if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = $1 AND provider_user_id = $2 LIMIT 1", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + { + if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + } + } else { + BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider + }; + tx.commit().await.map_postgres_err()?; + Ok(outcome) + } + + pub async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: chrono::DateTime, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE users +SET email_verified = TRUE, + updated_at = $3 +WHERE id = $1 + AND email_verified IS FALSE + AND LOWER(TRIM(email)) = LOWER(TRIM($2)) +"#, + ) + .bind(user_id) + .bind(verified_email) + .bind(verified_at) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) } pub async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { + local_password_login_allowed: bool, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let provider_exists: Option = sqlx::query_scalar( + "SELECT provider_type FROM oauth_providers WHERE provider_type = $1 FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if provider_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + let user = sqlx::query( + "SELECT auth_source::text AS auth_source, password_hash FROM users WHERE id = $1 FOR UPDATE", + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(user) = user else { + tx.rollback().await.map_postgres_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + }; + let auth_source = user + .try_get::("auth_source") + .map_postgres_err()?; + let password_hash = user + .try_get::, _>("password_hash") + .map_postgres_err()?; + let provider_types = sqlx::query_scalar::<_, String>( + "SELECT provider_type FROM user_oauth_links WHERE user_id = $1 FOR UPDATE", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + if !provider_types.iter().any(|value| value == provider_type) { + tx.rollback().await.map_postgres_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + let enabled_provider_types = sqlx::query_scalar::<_, String>( + r#" +SELECT user_oauth_links.provider_type +FROM user_oauth_links +JOIN oauth_providers + ON oauth_providers.provider_type = user_oauth_links.provider_type +WHERE user_oauth_links.user_id = $1 + AND oauth_providers.is_enabled IS TRUE +FOR UPDATE OF user_oauth_links +"#, + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + let has_remaining_enabled_oauth_link = enabled_provider_types + .iter() + .any(|value| value != provider_type); + if !has_remaining_enabled_oauth_link { + if let Some(outcome) = last_oauth_unbind_denial( + &auth_source, + password_hash.as_deref(), + local_password_login_allowed, + ) { + tx.rollback().await.map_postgres_err()?; + return Ok(outcome); + } + } let result = sqlx::query(DELETE_USER_OAUTH_LINK_SQL) .bind(user_id) .bind(provider_type) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; - Ok(result.rows_affected() > 0) + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + tx.commit().await.map_postgres_err()?; + Ok(DeleteUserOAuthLinkOutcome::Deleted) } pub async fn get_or_create_ldap_auth_user( @@ -1394,7 +1921,7 @@ RETURNING id, email, email_verified, username, password_hash, role::text AS role, auth_source::text AS auth_source, allowed_providers, allowed_providers_mode, allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode, - is_active, is_deleted, created_at, last_login_at + is_active, is_deleted, security_version, created_at, last_login_at "#, ) .bind(&existing.id) @@ -1446,7 +1973,7 @@ RETURNING id, email, email_verified, username, password_hash, role::text AS role, auth_source::text AS auth_source, allowed_providers, allowed_providers_mode, allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode, - is_active, is_deleted, created_at, last_login_at + is_active, is_deleted, security_version, created_at, last_login_at "#, ) .bind(uuid::Uuid::new_v4().to_string()) @@ -1493,20 +2020,25 @@ RETURNING pub async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { let result = sqlx::query( r#" UPDATE users -SET email = COALESCE($2, email), - username = COALESCE($3, username), +SET email = CASE WHEN $2 THEN $3 ELSE email END, + email_verified = COALESCE($4, email_verified), + username = COALESCE($5, username), updated_at = NOW() WHERE id = $1 "#, ) .bind(user_id) + .bind(email_present) .bind(email) + .bind(email_verified) .bind(username) .execute(&self.pool) .await @@ -1517,6 +2049,163 @@ WHERE id = $1 self.find_user_auth_by_id(user_id).await } + // This operation compares and restores four correlated snapshots in one + // transaction. Keep the explicit arguments visible at the call site so a + // future restore cannot accidentally omit one consistency boundary. + #[allow(clippy::too_many_arguments)] + pub async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + if expected_auth.id != restored_auth.id + || expected_export.id != expected_auth.id + || restored_export.id != restored_auth.id + { + return Ok(false); + } + let mut tx = self.pool.begin().await.map_postgres_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(POSTGRES_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + let auth_row = sqlx::query(&format!("{FIND_USER_AUTH_BY_ID_SQL} FOR UPDATE")) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let export_row = sqlx::query(&format!("{FIND_EXPORT_USER_BY_ID_SQL} FOR UPDATE")) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let (Some(auth_row), Some(export_row)) = (auth_row, export_row) else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let current_auth = map_user_auth_row(&auth_row)?; + let current_export = map_user_export_row(&export_row)?; + if !current_auth.matches_restore_state(expected_auth) + || !current_export.matches_restore_state(expected_export) + || current_export.rate_limit != expected_export.rate_limit + || current_export.rate_limit_mode != expected_export.rate_limit_mode + || current_export.model_capability_settings.as_ref() + != expected_model_capability_settings + || current_export.feature_settings.as_ref() != expected_feature_settings + { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let removes_active_admin = current_auth.role.eq_ignore_ascii_case("admin") + && current_auth.is_active + && !current_auth.is_deleted + && (!restored_auth.role.eq_ignore_ascii_case("admin") || !restored_auth.is_active); + if removes_active_admin && active_admin_ids.len() <= 1 { + tx.rollback().await.map_postgres_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + let security_state_changed = expected_auth.role != restored_auth.role + || expected_auth.is_active != restored_auth.is_active; + let allowed_providers = restored_auth + .allowed_providers + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let allowed_api_formats = restored_auth + .allowed_api_formats + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let allowed_models = restored_auth + .allowed_models + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?; + let result = sqlx::query( + r#" +UPDATE users +SET email = $2, + email_verified = $3, + username = $4, + role = $5::userrole, + allowed_providers = $6::json, + allowed_providers_mode = $7, + allowed_api_formats = $8::json, + allowed_api_formats_mode = $9, + allowed_models = $10::json, + allowed_models_mode = $11, + rate_limit = $12, + rate_limit_mode = $13, + model_capability_settings = $14::json, + feature_settings = $15::jsonb, + is_active = $16, + security_version = security_version + CASE WHEN $17 THEN 1 ELSE 0 END, + updated_at = NOW() +WHERE id = $1 +"#, + ) + .bind(&expected_auth.id) + .bind(&restored_auth.email) + .bind(restored_auth.email_verified) + .bind(&restored_auth.username) + .bind(&restored_auth.role) + .bind(allowed_providers) + .bind(&restored_auth.allowed_providers_mode) + .bind(allowed_api_formats) + .bind(&restored_auth.allowed_api_formats_mode) + .bind(allowed_models) + .bind(&restored_auth.allowed_models_mode) + .bind(restored_export.rate_limit) + .bind(&restored_export.rate_limit_mode) + .bind(restored_model_capability_settings.clone()) + .bind(restored_feature_settings.clone()) + .bind(restored_auth.is_active) + .bind(security_state_changed) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + if security_state_changed { + sqlx::query( + "UPDATE user_sessions SET revoked_at = NOW(), revoke_reason = 'user_security_state_changed', updated_at = NOW() WHERE user_id = $1 AND revoked_at IS NULL", + ) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = FALSE, updated_at = NOW() WHERE user_id = $1 AND is_active IS TRUE", + ) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = FALSE, updated_at = NOW() WHERE user_id = $1 AND is_active IS TRUE", + ) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + pub async fn update_local_auth_user_password_hash( &self, user_id: &str, @@ -1527,6 +2216,7 @@ WHERE id = $1 r#" UPDATE users SET password_hash = $2, + security_version = security_version + 1, updated_at = $3 WHERE id = $1 "#, @@ -1543,6 +2233,142 @@ WHERE id = $1 self.find_user_auth_by_id(user_id).await } + pub async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: chrono::DateTime, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE users +SET password_hash = $2, + security_version = security_version + 1, + updated_at = $3 +WHERE id = $1 + AND (($4::TEXT IS NULL AND password_hash IS NULL) OR password_hash = $4) +"#, + ) + .bind(user_id) + .bind(password_hash) + .bind(updated_at) + .bind(expected_password_hash) + .execute(&self.pool) + .await + .map_postgres_err()?; + Ok(result.rows_affected() == 1) + } + + pub async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let updated = sqlx::query( + "UPDATE users SET password_hash = $2, security_version = security_version + 1, updated_at = $3 WHERE id = $1 AND is_deleted IS FALSE", + ) + .bind(user_id) + .bind(password_hash) + .bind(changed_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if updated.rows_affected() != 1 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + sqlx::query( + "UPDATE user_sessions SET revoked_at = $2, revoke_reason = 'admin_password_reset', updated_at = $2 WHERE user_id = $1 AND revoked_at IS NULL", + ) + .bind(user_id) + .bind(changed_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; + Ok(true) + } + + pub async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let row = sqlx::query( + r#" +SELECT password_hash, is_active, is_deleted +FROM users +WHERE id = $1 +FOR UPDATE +"#, + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let stored_password_hash = row + .try_get::, _>("password_hash") + .map_postgres_err()?; + let is_active = row.try_get::("is_active").map_postgres_err()?; + let is_deleted = row.try_get::("is_deleted").map_postgres_err()?; + if stored_password_hash.as_deref() != expected_password_hash || !is_active || is_deleted { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let current_session_exists = sqlx::query_scalar::<_, bool>( + r#" +SELECT EXISTS ( + SELECT 1 FROM user_sessions + WHERE user_id = $1 AND id = $2 AND revoked_at IS NULL AND expires_at > $3 +) +"#, + ) + .bind(user_id) + .bind(current_session_id) + .bind(changed_at) + .fetch_one(&mut *tx) + .await + .map_postgres_err()?; + if !current_session_exists { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + sqlx::query( + "UPDATE users SET password_hash = $2, security_version = security_version + 1, updated_at = $3 WHERE id = $1", + ) + .bind(user_id) + .bind(next_password_hash) + .bind(changed_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + let revoked = sqlx::query( + "UPDATE user_sessions SET revoked_at = $2, revoke_reason = 'password_changed', updated_at = $2 WHERE user_id = $1 AND revoked_at IS NULL", + ) + .bind(user_id) + .bind(changed_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + if revoked.rows_affected() == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) + } + #[allow(clippy::too_many_arguments)] pub async fn update_local_auth_user_admin_fields( &self, @@ -1558,6 +2384,46 @@ WHERE id = $1 rate_limit: Option, is_active: Option, ) -> Result, DataLayerError> { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(POSTGRES_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) + .await + .map_postgres_err()?; + let current_security_state = sqlx::query( + "SELECT role::text AS role, is_active, is_deleted FROM users WHERE id = $1 FOR UPDATE", + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(current_security_state) = current_security_state else { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + }; + let current_role = current_security_state + .try_get::("role") + .map_postgres_err()?; + let current_active = current_security_state + .try_get::("is_active") + .map_postgres_err()?; + let current_deleted = current_security_state + .try_get::("is_deleted") + .map_postgres_err()?; + let next_role = role.as_deref().unwrap_or(current_role.as_str()); + let next_active = is_active.unwrap_or(current_active); + if current_role.eq_ignore_ascii_case("admin") + && current_active + && !current_deleted + && (!next_role.eq_ignore_ascii_case("admin") || !next_active) + && active_admin_ids.len() <= 1 + { + tx.rollback().await.map_postgres_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + let security_state_changed = + !current_role.eq_ignore_ascii_case(next_role) || current_active != next_active; let allowed_providers_mode = if allowed_providers .as_ref() .is_some_and(|values| !values.is_empty()) @@ -1630,6 +2496,7 @@ SET role = CASE WHEN $16::BOOLEAN AND $17 IS NOT NULL THEN $17 ELSE is_active END, + security_version = security_version + CASE WHEN $18::BOOLEAN THEN 1 ELSE 0 END, updated_at = NOW() WHERE id = $1 "#, @@ -1651,12 +2518,38 @@ WHERE id = $1 .bind(rate_limit_mode) .bind(is_active.is_some()) .bind(is_active) - .execute(&self.pool) + .bind(security_state_changed) + .execute(&mut *tx) .await .map_postgres_err()?; if result.rows_affected() == 0 { + tx.rollback().await.map_postgres_err()?; return Ok(None); } + if security_state_changed { + sqlx::query( + "UPDATE user_sessions SET revoked_at = NOW(), revoke_reason = 'user_security_state_changed', updated_at = NOW() WHERE user_id = $1 AND revoked_at IS NULL", + ) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = FALSE, updated_at = NOW() WHERE user_id = $1 AND is_active IS TRUE", + ) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = FALSE, updated_at = NOW() WHERE user_id = $1 AND is_active IS TRUE", + ) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + tx.commit().await.map_postgres_err()?; self.find_user_auth_by_id(user_id).await } @@ -1872,15 +2765,144 @@ VALUES ( self.find_user_auth_by_id(&user_id).await } - pub async fn delete_local_auth_user(&self, user_id: &str) -> Result { - let result = sqlx::query("DELETE FROM users WHERE id = $1") - .bind(user_id) - .execute(&self.pool) + async fn delete_local_auth_user_inner( + &self, + user_id: &str, + require_wallet_absent: bool, + ) -> Result { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let active_admin_ids = sqlx::query_scalar::<_, String>(POSTGRES_LOCK_ACTIVE_ADMINS_SQL) + .fetch_all(&mut *tx) .await .map_postgres_err()?; + let target_security_state = sqlx::query( + "SELECT role::text AS role, is_active, is_deleted FROM users WHERE id = $1 FOR UPDATE", + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(target_security_state) = target_security_state else { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + }; + let target_role = target_security_state + .try_get::("role") + .map_postgres_err()?; + let target_is_active = target_security_state + .try_get::("is_active") + .map_postgres_err()?; + let target_is_deleted = target_security_state + .try_get::("is_deleted") + .map_postgres_err()?; + if target_role.eq_ignore_ascii_case("admin") + && target_is_active + && !target_is_deleted + && active_admin_ids.len() <= 1 + { + tx.rollback().await.map_postgres_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_DELETE_DENIED.to_string(), + )); + } + if require_wallet_absent { + let wallet_exists: Option = sqlx::query_scalar( + r#" +SELECT 1 +FROM wallets AS wallet +WHERE wallet.user_id = $1 + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = $2 + ) +LIMIT 1 + "#, + ) + .bind(user_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if wallet_exists.is_some() { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + } + for sql in POSTGRES_PREPARE_USER_FACTS_FOR_DELETION_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + for sql in POSTGRES_ANONYMIZE_USER_HISTORY_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + sqlx::query(POSTGRES_ANONYMIZE_USER_API_KEY_HISTORY_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + for sql in POSTGRES_DELETE_USER_DEPENDENTS_SQL { + if require_wallet_absent && *sql == POSTGRES_DELETE_USER_API_KEYS_SQL { + continue; + } + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + let result = if require_wallet_absent { + sqlx::query(POSTGRES_DELETE_USER_IF_WALLET_ABSENT_SQL) + .bind(user_id) + .bind(user_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()? + } else { + sqlx::query("DELETE FROM users WHERE id = $1") + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()? + }; + if require_wallet_absent && result.rows_affected() == 0 { + // A wallet may have been inserted after the initial check. Do not + // commit the history/credential mutations when the guarded delete + // loses that race. + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + if require_wallet_absent { + sqlx::query(POSTGRES_DELETE_USER_API_KEYS_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } + tx.commit().await.map_postgres_err()?; Ok(result.rows_affected() > 0) } + pub async fn delete_local_auth_user(&self, user_id: &str) -> Result { + self.delete_local_auth_user_inner(user_id, false).await + } + + pub async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + self.delete_local_auth_user_inner(user_id, true).await + } + pub async fn read_user_preferences( &self, user_id: &str, @@ -1951,16 +2973,35 @@ VALUES ( .or(session.updated_at) .or(session.last_seen_at) .unwrap_or_else(chrono::Utc::now); + let mut tx = self.pool.begin().await.map_postgres_err()?; + let user_is_eligible = sqlx::query_scalar::<_, String>( + r#" +SELECT id FROM users +WHERE id = $1 AND is_active IS TRUE AND is_deleted IS FALSE + AND security_version = $2 +FOR UPDATE +"#, + ) + .bind(&session.user_id) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_is_eligible.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } sqlx::query(REVOKE_ACTIVE_DEVICE_SESSIONS_SQL) .bind(&session.user_id) .bind(&session.client_device_id) .bind(now) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; let row = sqlx::query(CREATE_USER_SESSION_SQL) .bind(&session.id) .bind(&session.user_id) + .bind(session.security_version) .bind(&session.client_device_id) .bind(session.device_label.as_deref()) .bind("unknown") @@ -1971,9 +3012,74 @@ VALUES ( .bind(session.expires_at.unwrap_or(now)) .bind(session.created_at.unwrap_or(now)) .bind(session.updated_at.unwrap_or(now)) - .fetch_one(&self.pool) + .fetch_one(&mut *tx) .await .map_postgres_err()?; + let session = map_user_session_row(&row)?; + tx.commit().await.map_postgres_err()?; + Ok(Some(session)) + } + + pub async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + let now = session + .created_at + .or(session.updated_at) + .or(session.last_seen_at) + .unwrap_or_else(chrono::Utc::now); + let mut tx = self.pool.begin().await.map_postgres_err()?; + let matched = sqlx::query_scalar::<_, String>( + r#" +SELECT password_hash FROM users +WHERE id = $1 AND password_hash = $2 AND auth_source::text = 'local' + AND is_active IS TRUE AND is_deleted IS FALSE AND security_version = $3 +FOR UPDATE +"#, + ) + .bind(&session.user_id) + .bind(expected_password_hash) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if matched.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + sqlx::query("UPDATE users SET last_login_at = $2 WHERE id = $1") + .bind(&session.user_id) + .bind(now) + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query(REVOKE_ACTIVE_DEVICE_SESSIONS_SQL) + .bind(&session.user_id) + .bind(&session.client_device_id) + .bind(now) + .execute(&mut *tx) + .await + .map_postgres_err()?; + let row = sqlx::query(CREATE_USER_SESSION_SQL) + .bind(&session.id) + .bind(&session.user_id) + .bind(session.security_version) + .bind(&session.client_device_id) + .bind(session.device_label.as_deref()) + .bind("unknown") + .bind(session.ip_address.as_deref()) + .bind(session.user_agent.as_deref()) + .bind(&session.refresh_token_hash) + .bind(session.last_seen_at.unwrap_or(now)) + .bind(session.expires_at.unwrap_or(now)) + .bind(session.created_at.unwrap_or(now)) + .bind(session.updated_at.unwrap_or(now)) + .fetch_one(&mut *tx) + .await + .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(Some(map_user_session_row(&row)?)) } @@ -2020,7 +3126,7 @@ VALUES ( &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: chrono::DateTime, expires_at: chrono::DateTime, @@ -2030,7 +3136,7 @@ VALUES ( let result = sqlx::query(ROTATE_USER_SESSION_REFRESH_SQL) .bind(user_id) .bind(session_id) - .bind(previous_refresh_token_hash) + .bind(expected_refresh_token_hash) .bind(rotated_at) .bind(next_refresh_token_hash) .bind(expires_at) @@ -2079,11 +3185,14 @@ VALUES ( pub async fn count_active_local_admin_users_with_valid_password( &self, ) -> Result { - let total: i64 = sqlx::query_scalar(COUNT_ACTIVE_LOCAL_ADMIN_USERS_WITH_VALID_PASSWORD_SQL) - .fetch_one(&self.pool) + let hashes = sqlx::query_scalar::<_, String>(LIST_ACTIVE_LOCAL_ADMIN_PASSWORD_HASHES_SQL) + .fetch_all(&self.pool) .await .map_postgres_err()?; - Ok(total.max(0) as u64) + Ok(hashes + .iter() + .filter(|hash| is_valid_bcrypt_hash(hash)) + .count() as u64) } } @@ -2134,6 +3243,9 @@ fn map_user_session_row( row.try_get("created_at").map_postgres_err()?, row.try_get("updated_at").map_postgres_err()?, ) + .and_then(|record| { + record.with_security_version(row.try_get("security_version").map_postgres_err()?) + }) } fn normalize_optional_json_value(value: Option) -> Option { @@ -2258,6 +3370,9 @@ fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result Result { + self.restore_user_group_if_matches(expected, restored).await + } + async fn delete_user_group(&self, group_id: &str) -> Result { self.delete_user_group(group_id).await } @@ -2464,6 +3587,16 @@ impl UserReadRepository for SqlxUserReadRepository { self.replace_user_groups_for_user(user_id, group_ids).await } + async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + self.restore_user_groups_if_matches(user_id, expected_group_ids, restored_group_ids) + .await + } + async fn add_user_to_group( &self, group_id: &str, @@ -2530,6 +3663,29 @@ impl UserReadRepository for SqlxUserReadRepository { .await } + async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: chrono::DateTime, + _provider_enabled_snapshot: bool, + ) -> Result { + self.resolve_enabled_oauth_linked_user( + provider_type, + provider_user_id, + provider_username, + provider_email, + extra_data, + verified_email, + touched_at, + ) + .await + } + async fn touch_oauth_link( &self, provider_type: &str, @@ -2553,10 +3709,11 @@ impl UserReadRepository for SqlxUserReadRepository { async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: chrono::DateTime, ) -> Result, DataLayerError> { - self.create_oauth_auth_user(email, username, created_at) + self.create_oauth_auth_user(email, email_verified, username, created_at) .await } @@ -2582,7 +3739,20 @@ impl UserReadRepository for SqlxUserReadRepository { self.count_user_oauth_links(user_id).await } - async fn upsert_user_oauth_link( + async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + let exists: Option = + sqlx::query_scalar("SELECT 1 FROM user_oauth_links WHERE provider_type = $1 LIMIT 1") + .bind(provider_type) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + Ok(exists.is_some()) + } + + async fn bind_user_oauth_link_if_provider_enabled( &self, user_id: &str, provider_type: &str, @@ -2591,8 +3761,116 @@ impl UserReadRepository for SqlxUserReadRepository { provider_email: Option<&str>, extra_data: Option, linked_at: chrono::DateTime, - ) -> Result<(), DataLayerError> { - self.upsert_user_oauth_link( + _provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result { + if let Some(expectation) = session_expectation { + let mut tx = self.pool.begin().await.map_postgres_err()?; + let provider_enabled = sqlx::query_scalar::<_, bool>( + "SELECT is_enabled FROM oauth_providers WHERE provider_type = $1 FOR UPDATE", + ) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if provider_enabled != Some(true) { + tx.rollback().await.map_postgres_err()?; + return Ok(if provider_enabled.is_none() { + BindUserOAuthLinkOutcome::ProviderNotFound + } else { + BindUserOAuthLinkOutcome::ProviderDisabled + }); + } + let session_is_current: Option = sqlx::query_scalar( + r#" +SELECT 1 +FROM users +JOIN user_sessions ON user_sessions.user_id = users.id +WHERE users.id = $1 + AND users.is_active IS TRUE + AND users.is_deleted IS FALSE + AND users.security_version = $2 + AND user_sessions.id = $3 + AND user_sessions.security_version = $2 + AND user_sessions.client_device_id = $4 + AND user_sessions.revoked_at IS NULL + AND user_sessions.expires_at > GREATEST($5, NOW()) +FOR UPDATE OF users, user_sessions +"#, + ) + .bind(user_id) + .bind(expectation.security_version) + .bind(&expectation.session_id) + .bind(&expectation.client_device_id) + .bind(expectation.checked_at) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if session_is_current.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::SessionUnavailable); + } + if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = $1 AND provider_user_id = $2 LIMIT 1 FOR UPDATE", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? { + tx.rollback().await.map_postgres_err()?; + return Ok(if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }); + } + if sqlx::query_scalar::<_, i32>( + "SELECT 1 FROM user_oauth_links WHERE user_id = $1 AND provider_type = $2 LIMIT 1 FOR UPDATE", + ) + .bind(user_id) + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + .is_some() { + tx.rollback().await.map_postgres_err()?; + return Ok(BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider); + } + let inserted = sqlx::query(UPSERT_OAUTH_LINK_SQL) + .bind(uuid::Uuid::new_v4().to_string()) + .bind(user_id) + .bind(provider_type) + .bind(provider_user_id) + .bind(provider_username) + .bind(provider_email) + .bind(extra_data) + .bind(linked_at) + .execute(&mut *tx) + .await + .map_postgres_err()?; + let outcome = if inserted.rows_affected() == 1 { + BindUserOAuthLinkOutcome::Bound + } else if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = $1 AND provider_user_id = $2 LIMIT 1", + ) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? { + if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + } + } else { + BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider + }; + tx.commit().await.map_postgres_err()?; + return Ok(outcome); + } + self.bind_user_oauth_link( user_id, provider_type, provider_user_id, @@ -2604,12 +3882,25 @@ impl UserReadRepository for SqlxUserReadRepository { .await } + async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: chrono::DateTime, + ) -> Result { + self.upgrade_oauth_email_verification_if_matches(user_id, verified_email, verified_at) + .await + } + async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { - self.delete_user_oauth_link(user_id, provider_type).await + local_password_login_allowed: bool, + _enabled_provider_types_snapshot: &[String], + ) -> Result { + self.delete_user_oauth_link(user_id, provider_type, local_password_login_allowed) + .await } async fn get_or_create_ldap_auth_user( @@ -2635,13 +3926,39 @@ impl UserReadRepository for SqlxUserReadRepository { async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { - self.update_local_auth_user_profile(user_id, email, username) + self.update_local_auth_user_profile(user_id, email_present, email, email_verified, username) .await } + async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + self.restore_local_auth_user_state_if_matches( + expected_auth, + restored_auth, + expected_export, + restored_export, + expected_model_capability_settings, + restored_model_capability_settings, + expected_feature_settings, + restored_feature_settings, + ) + .await + } + async fn update_local_auth_user_password_hash( &self, user_id: &str, @@ -2652,6 +3969,50 @@ impl UserReadRepository for SqlxUserReadRepository { .await } + async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: chrono::DateTime, + ) -> Result { + self.restore_local_auth_user_password_hash_if_matches( + user_id, + expected_password_hash, + password_hash, + updated_at, + ) + .await + } + + async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + self.reset_local_auth_user_password_and_revoke_sessions(user_id, password_hash, changed_at) + .await + } + + async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + self.change_local_auth_password_and_revoke_sessions( + user_id, + current_session_id, + expected_password_hash, + next_password_hash, + changed_at, + ) + .await + } + async fn update_local_auth_user_admin_fields( &self, user_id: &str, @@ -2758,6 +4119,13 @@ impl UserReadRepository for SqlxUserReadRepository { self.delete_local_auth_user(user_id).await } + async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + self.delete_local_auth_user_if_wallet_absent(user_id).await + } + async fn read_user_preferences( &self, user_id: &str, @@ -2794,6 +4162,15 @@ impl UserReadRepository for SqlxUserReadRepository { self.create_user_session(session).await } + async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + self.create_user_session_if_password_matches(session, expected_password_hash) + .await + } + async fn touch_user_session( &self, user_id: &str, @@ -2821,7 +4198,7 @@ impl UserReadRepository for SqlxUserReadRepository { &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: chrono::DateTime, expires_at: chrono::DateTime, @@ -2831,7 +4208,7 @@ impl UserReadRepository for SqlxUserReadRepository { self.rotate_user_session_refresh_token( user_id, session_id, - previous_refresh_token_hash, + expected_refresh_token_hash, next_refresh_token_hash, rotated_at, expires_at, @@ -2873,3 +4250,72 @@ impl UserReadRepository for SqlxUserReadRepository { .await } } + +#[cfg(test)] +mod admin_invariant_tests { + use super::{ + POSTGRES_ANONYMIZE_USER_API_KEY_HISTORY_SQL, POSTGRES_ANONYMIZE_USER_HISTORY_SQL, + POSTGRES_DELETE_USER_DEPENDENTS_SQL, POSTGRES_LOCK_ACTIVE_ADMINS_SQL, + }; + + #[test] + fn active_admin_mutations_use_a_deterministic_postgres_row_lock() { + let normalized = POSTGRES_LOCK_ACTIVE_ADMINS_SQL + .split_whitespace() + .collect::>() + .join(" "); + assert!(normalized.contains("role = 'admin'::userrole")); + assert!(normalized.contains("is_active IS TRUE")); + assert!(normalized.contains("is_deleted IS FALSE")); + assert!(normalized.contains("ORDER BY id FOR UPDATE")); + assert!(POSTGRES_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM management_tokens"))); + assert!(POSTGRES_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM api_keys"))); + assert!(POSTGRES_DELETE_USER_DEPENDENTS_SQL + .iter() + .any(|sql| sql.starts_with("DELETE FROM user_sessions"))); + assert_history_anonymization_contract(POSTGRES_ANONYMIZE_USER_HISTORY_SQL); + assert!(POSTGRES_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .starts_with("UPDATE stats_daily_api_key SET api_key_name = NULL")); + assert!(POSTGRES_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .contains("SELECT id FROM api_keys WHERE user_id = $1")); + } + + fn assert_history_anonymization_contract(statements: &[&str]) { + const TABLES: &[&str] = &[ + "request_candidates", + "video_tasks", + "usage", + "stats_user_daily", + "stats_user_summary", + "stats_user_daily_model", + "stats_user_daily_provider", + "stats_user_daily_api_format", + "stats_user_daily_model_provider", + "stats_user_daily_cost_savings", + "stats_user_daily_cost_savings_provider", + "stats_user_daily_cost_savings_model", + "stats_user_daily_cost_savings_model_provider", + ]; + + assert_eq!(statements.len(), TABLES.len()); + for table in TABLES { + let statement = statements + .iter() + .find(|sql| sql.starts_with(&format!("UPDATE {table} "))) + .unwrap_or_else(|| panic!("missing history anonymization for {table}")); + assert!(statement.contains("username = NULL")); + assert!(statement.ends_with("WHERE user_id = $1")); + } + for table in ["request_candidates", "video_tasks", "usage"] { + let statement = statements + .iter() + .find(|sql| sql.starts_with(&format!("UPDATE {table} "))) + .expect("identity snapshot table should be covered"); + assert!(statement.contains("api_key_name = NULL")); + } + } +} diff --git a/crates/aether-data/adapters/postgres/src/video_tasks.rs b/crates/aether-data/adapters/postgres/src/video_tasks.rs index 8b2fe5788..38905f92b 100644 --- a/crates/aether-data/adapters/postgres/src/video_tasks.rs +++ b/crates/aether-data/adapters/postgres/src/video_tasks.rs @@ -101,10 +101,18 @@ fn find_by_id_sql() -> String { select_video_task_sql("WHERE id = $1\nLIMIT 1") } +fn find_by_id_for_user_sql() -> String { + select_video_task_sql("WHERE id = $1 AND user_id = $2\nLIMIT 1") +} + fn find_by_short_id_sql() -> String { select_video_task_sql("WHERE short_id = $1\nLIMIT 1") } +fn find_by_short_id_for_user_sql() -> String { + select_video_task_sql("WHERE short_id = $1 AND user_id = $2\nLIMIT 1") +} + fn find_by_user_external_sql() -> String { select_video_task_sql("WHERE user_id = $1 AND external_task_id = $2\nLIMIT 1") } @@ -302,10 +310,26 @@ ON CONFLICT (id) DO UPDATE SET error_code = EXCLUDED.error_code, error_message = EXCLUDED.error_message, request_metadata = EXCLUDED.request_metadata, - created_at = EXCLUDED.created_at, + created_at = COALESCE(video_tasks.created_at, EXCLUDED.created_at), submitted_at = EXCLUDED.submitted_at, completed_at = EXCLUDED.completed_at, updated_at = EXCLUDED.updated_at +WHERE video_tasks.short_id IS NOT DISTINCT FROM EXCLUDED.short_id + AND video_tasks.request_id IS NOT DISTINCT FROM EXCLUDED.request_id + AND video_tasks.user_id IS NOT DISTINCT FROM EXCLUDED.user_id + AND video_tasks.api_key_id IS NOT DISTINCT FROM EXCLUDED.api_key_id + AND video_tasks.external_task_id IS NOT DISTINCT FROM EXCLUDED.external_task_id + AND video_tasks.provider_id IS NOT DISTINCT FROM EXCLUDED.provider_id + AND video_tasks.endpoint_id IS NOT DISTINCT FROM EXCLUDED.endpoint_id + AND video_tasks.key_id IS NOT DISTINCT FROM EXCLUDED.key_id + AND video_tasks.client_api_format IS NOT DISTINCT FROM EXCLUDED.client_api_format + AND video_tasks.provider_api_format IS NOT DISTINCT FROM EXCLUDED.provider_api_format + AND video_tasks.format_converted IS NOT DISTINCT FROM EXCLUDED.format_converted + AND video_tasks.model IS NOT DISTINCT FROM EXCLUDED.model + AND video_tasks.duration_seconds IS NOT DISTINCT FROM EXCLUDED.duration_seconds + AND video_tasks.resolution IS NOT DISTINCT FROM EXCLUDED.resolution + AND video_tasks.aspect_ratio IS NOT DISTINCT FROM EXCLUDED.aspect_ratio + AND video_tasks.size IS NOT DISTINCT FROM EXCLUDED.size RETURNING {columns} " @@ -348,12 +372,28 @@ fn update_if_active_sql() -> String { error_code = $31, error_message = $32, request_metadata = $33, - created_at = TO_TIMESTAMP($34), + created_at = COALESCE(created_at, TO_TIMESTAMP($34)), submitted_at = TO_TIMESTAMP($35), completed_at = TO_TIMESTAMP($36), updated_at = TO_TIMESTAMP($37) WHERE id = $1 AND status = ANY($38) + AND short_id IS NOT DISTINCT FROM $2 + AND request_id IS NOT DISTINCT FROM $3 + AND user_id IS NOT DISTINCT FROM $4 + AND api_key_id IS NOT DISTINCT FROM $5 + AND external_task_id IS NOT DISTINCT FROM $8 + AND provider_id IS NOT DISTINCT FROM $9 + AND endpoint_id IS NOT DISTINCT FROM $10 + AND key_id IS NOT DISTINCT FROM $11 + AND client_api_format IS NOT DISTINCT FROM $12 + AND provider_api_format IS NOT DISTINCT FROM $13 + AND format_converted IS NOT DISTINCT FROM $14 + AND model IS NOT DISTINCT FROM $15 + AND duration_seconds IS NOT DISTINCT FROM $18 + AND resolution IS NOT DISTINCT FROM $19 + AND aspect_ratio IS NOT DISTINCT FROM $20 + AND size IS NOT DISTINCT FROM $21 RETURNING {columns} " @@ -390,6 +430,26 @@ impl SqlxVideoTaskRepository { } } + pub async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + match key { + VideoTaskLookupKey::Id(id) => self.find_by_id_for_user(id, user_id).await, + VideoTaskLookupKey::ShortId(short_id) => { + self.find_by_short_id_for_user(short_id, user_id).await + } + VideoTaskLookupKey::UserExternal { + user_id: lookup_user_id, + external_task_id, + } if lookup_user_id == user_id => { + self.find_by_user_external(user_id, external_task_id).await + } + VideoTaskLookupKey::UserExternal { .. } => Ok(None), + } + } + pub async fn find_by_id(&self, id: &str) -> Result, DataLayerError> { let sql = find_by_id_sql(); let row = sqlx::query(&sql) @@ -400,6 +460,21 @@ impl SqlxVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + pub async fn find_by_id_for_user( + &self, + id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let sql = find_by_id_for_user_sql(); + let row = sqlx::query(&sql) + .bind(id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + pub async fn find_by_short_id( &self, short_id: &str, @@ -413,6 +488,21 @@ impl SqlxVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + pub async fn find_by_short_id_for_user( + &self, + short_id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let sql = find_by_short_id_for_user_sql(); + let row = sqlx::query(&sql) + .bind(short_id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + pub async fn find_by_user_external( &self, user_id: &str, @@ -661,7 +751,12 @@ impl SqlxVideoTaskRepository { .map_err(|_| DataLayerError::UnexpectedValue(format!("invalid count result: {total}"))) } - pub async fn upsert(&self, task: UpsertVideoTask) -> Result { + pub async fn upsert( + &self, + mut task: UpsertVideoTask, + ) -> Result { + task.sanitize_for_persistence(); + let task_id = task.id.clone(); let sql = upsert_sql(); let row = sqlx::query(&sql) .bind(task.id) @@ -720,17 +815,23 @@ impl SqlxVideoTaskRepository { .bind(task.submitted_at_unix_secs.map(|value| value as f64)) .bind(task.completed_at_unix_secs.map(|value| value as f64)) .bind(task.updated_at_unix_secs as f64) - .fetch_one(&self.pool) + .fetch_optional(&self.pool) .await - .map_postgres_err()?; + .map_postgres_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "video task {task_id} conflicts with persisted immutable identity" + )) + })?; map_video_task_row(&row) } pub async fn update_if_active( &self, - task: UpsertVideoTask, + mut task: UpsertVideoTask, ) -> Result, DataLayerError> { + task.sanitize_for_persistence(); let sql = update_if_active_sql(); let row = sqlx::query(&sql) .bind(task.id) @@ -839,6 +940,14 @@ impl VideoTaskReadRepository for SqlxVideoTaskRepository { Self::find(self, key).await } + async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + Self::find_for_user(self, key, user_id).await + } + async fn list_active(&self, limit: usize) -> Result, DataLayerError> { Self::list_active(self, limit).await } @@ -1067,7 +1176,7 @@ fn map_video_task_row(row: &PgRow) -> Result { #[cfg(test)] mod tests { - use super::SqlxVideoTaskRepository; + use super::{update_if_active_sql, upsert_sql, SqlxVideoTaskRepository}; use crate::{PostgresPoolConfig, PostgresPoolFactory}; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository, @@ -1090,6 +1199,47 @@ mod tests { factory.connect_lazy().expect("pool should build") } + #[test] + fn write_sql_guards_every_immutable_identity_field() { + let upsert = upsert_sql(); + let update = update_if_active_sql(); + let immutable_columns = [ + ("short_id", "$2"), + ("request_id", "$3"), + ("user_id", "$4"), + ("api_key_id", "$5"), + ("external_task_id", "$8"), + ("provider_id", "$9"), + ("endpoint_id", "$10"), + ("key_id", "$11"), + ("client_api_format", "$12"), + ("provider_api_format", "$13"), + ("format_converted", "$14"), + ("model", "$15"), + ("duration_seconds", "$18"), + ("resolution", "$19"), + ("aspect_ratio", "$20"), + ("size", "$21"), + ]; + + for (column, parameter) in immutable_columns { + assert!( + upsert.contains(&format!( + "video_tasks.{column} IS NOT DISTINCT FROM EXCLUDED.{column}" + )), + "upsert should guard {column}" + ); + assert!( + update.contains(&format!("{column} IS NOT DISTINCT FROM {parameter}")), + "active update should guard {column}" + ); + } + assert!( + upsert.contains("created_at = COALESCE(video_tasks.created_at, EXCLUDED.created_at)") + ); + assert!(update.contains("created_at = COALESCE(created_at, TO_TIMESTAMP($34))")); + } + #[tokio::test] async fn repository_constructs_from_lazy_pool() { let repository = SqlxVideoTaskRepository::new(build_pool()); diff --git a/crates/aether-data/adapters/postgres/src/wallet.rs b/crates/aether-data/adapters/postgres/src/wallet.rs index e3870975a..91e154cc5 100644 --- a/crates/aether-data/adapters/postgres/src/wallet.rs +++ b/crates/aether-data/adapters/postgres/src/wallet.rs @@ -4,18 +4,38 @@ use futures_util::{stream::TryStream, TryStreamExt}; use sqlx::{postgres::PgRow, PgPool, Row}; use uuid::Uuid; +use aether_data_contracts::repository::billing::{ + checked_plan_duration_days_from_snapshot, entitlements_have_replacement_selector, + entitlements_should_replace_existing, +}; use aether_data_contracts::repository::wallet::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, - redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, - AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, - AdminWalletListQuery, AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, - CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + canonicalize_payment_method, canonicalize_wallet_refund_fields, + payment_callback_amount_matches_order, payment_callback_method_matches_order, + payment_callback_provider_matches_order, payment_order_is_failed_wallet_checkout_placeholder, + payment_order_is_uncertain_wallet_checkout_placeholder, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_payment_method, + redeem_code_refundable_amount, validate_admin_redeem_code_batch_input, + validate_manual_wallet_recharge, validate_payment_order_credit_amounts, + validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements, + validate_redeem_wallet_credit, validate_wallet_recharge_order_input, + wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token, + 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, + AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, + CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, + ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, @@ -24,7 +44,8 @@ use aether_data_contracts::repository::wallet::{ StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use aether_data_contracts::DataLayerError; @@ -418,7 +439,7 @@ WHERE ($1::TEXT IS NULL OR payment_method = $1) $2::TEXT IS NULL OR ( CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < NOW() THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= NOW() THEN 'expired' ELSE status END ) = $2 @@ -456,7 +477,7 @@ WHERE ($1::TEXT IS NULL OR payment_method = $1) $2::TEXT IS NULL OR ( CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < NOW() THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= NOW() THEN 'expired' ELSE status END ) = $2 @@ -500,6 +521,7 @@ const COUNT_WALLET_PAYMENT_ORDERS_BY_USER_SQL: &str = r#" SELECT COUNT(*) AS total FROM payment_orders WHERE user_id = $1 + AND order_kind = 'wallet_recharge' "#; const COUNT_PENDING_PAYMENT_ORDERS_BY_USER_SQL: &str = r#" @@ -530,7 +552,7 @@ SELECT gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < now() THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= now() THEN 'expired' ELSE status END AS status, CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, @@ -539,6 +561,7 @@ SELECT CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs FROM payment_orders WHERE user_id = $1 + AND order_kind = 'wallet_recharge' ORDER BY created_at DESC OFFSET $2 LIMIT $3 @@ -565,7 +588,7 @@ SELECT gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < now() THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= now() THEN 'expired' ELSE status END AS status, CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, @@ -575,9 +598,78 @@ SELECT FROM payment_orders WHERE user_id = $1 AND id = $2 + AND order_kind = 'wallet_recharge' LIMIT 1 "#; +const FIND_WALLET_RECHARGE_ORDER_BY_ORDER_NO_SQL: &str = r#" +SELECT + id, + order_no, + wallet_id, + user_id, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST(pay_amount AS DOUBLE PRECISION) AS pay_amount, + pay_currency, + CAST(exchange_rate AS DOUBLE PRECISION) AS exchange_rate, + CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd, + CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd, + payment_method, + payment_provider, + payment_channel, + order_kind, + product_id, + product_snapshot, + gateway_order_id, + gateway_response, + CASE + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= now() THEN 'expired' + ELSE status + END AS status, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM paid_at) AS BIGINT) AS paid_at_unix_secs, + CAST(EXTRACT(EPOCH FROM credited_at) AS BIGINT) AS credited_at_unix_secs, + CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs +FROM payment_orders +WHERE user_id = $1 + AND order_no = $2 + AND order_kind = 'wallet_recharge' +LIMIT 1 +"#; + +// A recharge retry can discover an existing order after it has provisioned a +// wallet for a user that did not have one yet. Remove that wallet only while +// it is still an untouched, unreferenced provisional row; any unexpected +// reference keeps it durable rather than risking data loss. +const DELETE_PROVISIONAL_RECHARGE_WALLET_SQL: &str = r#" +DELETE FROM wallets +WHERE id = $1 + AND user_id = $2 + AND api_key_id IS NULL + AND balance = 0 + AND gift_balance = 0 + AND total_recharged = 0 + AND total_consumed = 0 + AND total_refunded = 0 + AND total_adjusted = 0 + AND limit_mode IN ('finite', 'unlimited') + AND currency = 'USD' + AND status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM usage u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +"#; + const FIND_PENDING_PLAN_PURCHASE_ORDER_BY_USER_SQL: &str = r#" SELECT id, @@ -613,6 +705,36 @@ ORDER BY created_at DESC LIMIT 1 "#; +const FIND_PAYMENT_ORDER_BY_ORDER_NO_SQL: &str = r#" +SELECT + id, + order_no, + wallet_id, + user_id, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST(pay_amount AS DOUBLE PRECISION) AS pay_amount, + pay_currency, + CAST(exchange_rate AS DOUBLE PRECISION) AS exchange_rate, + CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd, + CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd, + payment_method, + payment_provider, + payment_channel, + order_kind, + product_id, + product_snapshot, + gateway_order_id, + gateway_response, + status, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM paid_at) AS BIGINT) AS paid_at_unix_secs, + CAST(EXTRACT(EPOCH FROM credited_at) AS BIGINT) AS credited_at_unix_secs, + CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs +FROM payment_orders +WHERE order_no = $1 +LIMIT 1 +"#; + const FIND_WALLET_REFUND_SQL: &str = r#" SELECT id, @@ -761,6 +883,26 @@ impl WalletReadRepository for SqlxWalletRepository { unlimited, ) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_postgres_auth_wallet( + &self.pool, + Some(user_id), + None, + initial_gift_usd, + unlimited, + ) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn initialize_auth_api_key_wallet( @@ -777,6 +919,26 @@ impl WalletReadRepository for SqlxWalletRepository { unlimited, ) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_postgres_auth_wallet( + &self.pool, + None, + Some(api_key_id), + initial_gift_usd, + unlimited, + ) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn update_auth_user_wallet_snapshot( @@ -1148,6 +1310,20 @@ impl WalletReadRepository for SqlxWalletRepository { row.as_ref().map(map_admin_payment_order_row).transpose() } + async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(FIND_WALLET_RECHARGE_ORDER_BY_ORDER_NO_SQL) + .bind(user_id) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_admin_payment_order_row).transpose() + } + async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, @@ -1162,6 +1338,18 @@ impl WalletReadRepository for SqlxWalletRepository { row.as_ref().map(map_admin_payment_order_row).transpose() } + async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(FIND_PAYMENT_ORDER_BY_ORDER_NO_SQL) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + row.as_ref().map(map_admin_payment_order_row).transpose() + } + async fn find_wallet_refund( &self, wallet_id: &str, @@ -1373,14 +1561,457 @@ LIMIT $4 #[async_trait] impl WalletWriteRepository for SqlxWalletRepository { + async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: WalletLookupKey<'_>, + ) -> Result { + if wallet_id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = $2 AND api_key_id IS NULL", user_id.to_string()) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => ( + "api_key_id = $2 AND user_id IS NULL", + api_key_id.to_string(), + ), + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let wallet_id = wallet_id.to_string(); + self.tx_runner + .run_read_write(|tx| { + let owner_clause = owner_clause.to_string(); + let owner_id = owner_id.clone(); + let wallet_id = wallet_id.clone(); + Box::pin(async move { + let select_sql = format!( + r#" +SELECT id +FROM wallets +WHERE id = $1 + AND {owner_clause} + AND balance = 0 + AND gift_balance = 0 + AND total_recharged = 0 + AND total_consumed = 0 + AND total_refunded = 0 + AND total_adjusted = 0 + AND limit_mode IN ('finite', 'unlimited') + AND currency = 'USD' + AND status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM usage u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1 +FOR UPDATE + "# + ); + let found = sqlx::query_scalar::<_, String>(&select_sql) + .bind(&wallet_id) + .bind(&owner_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(found_id) = found else { + return Ok(false); + }; + let removed = sqlx::query("DELETE FROM wallets WHERE id = $1") + .bind(&found_id) + .execute(&mut **tx) + .await + .map_postgres_err()? + .rows_affected() + > 0; + Ok(removed) + }) + }) + .await + } + + async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if expected.id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = $2 AND api_key_id IS NULL", user_id.to_string()) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => ( + "api_key_id = $2 AND user_id IS NULL", + api_key_id.to_string(), + ), + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let expected = expected.clone(); + self.tx_runner + .run_read_write(|tx| { + let owner_clause = owner_clause.to_string(); + let owner_id = owner_id.clone(); + Box::pin(async move { + let select_sql = format!( + r#" +SELECT + id, + user_id, + api_key_id, + CAST(balance AS DOUBLE PRECISION) AS balance, + CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + limit_mode, + currency, + status, + 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(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs +FROM wallets +WHERE id = $1 + AND {owner_clause} + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM usage u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1 +FOR UPDATE + "# + ); + let row = sqlx::query(&select_sql) + .bind(&expected.id) + .bind(&owner_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if current != expected { + return Ok(false); + } + let removed = sqlx::query("DELETE FROM wallets WHERE id = $1") + .bind(&expected.id) + .execute(&mut **tx) + .await + .map_postgres_err()? + .rows_affected() + > 0; + Ok(removed) + }) + }) + .await + } + + async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if before.id.trim().is_empty() || after.id.trim().is_empty() { + return Ok(false); + } + if before.id != after.id { + return Err(DataLayerError::InvalidInput( + "wallet restore snapshots must reference the same wallet".to_string(), + )); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = $2 AND api_key_id IS NULL", user_id.to_string()) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => ( + "api_key_id = $2 AND user_id IS NULL", + api_key_id.to_string(), + ), + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet restore requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let owner_matches = match owner { + WalletLookupKey::UserId(user_id) => { + before.user_id.as_deref() == Some(user_id) + && after.user_id.as_deref() == Some(user_id) + && before.api_key_id.is_none() + && after.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + before.api_key_id.as_deref() == Some(api_key_id) + && after.api_key_id.as_deref() == Some(api_key_id) + && before.user_id.is_none() + && after.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if !owner_matches { + return Ok(false); + } + let before_updated_at = i64::try_from(before.updated_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput( + "wallet restore timestamp is outside the supported range".to_string(), + ) + })?; + let before = before.clone(); + let after = after.clone(); + self.tx_runner + .run_read_write(|tx| { + let owner_clause = owner_clause.to_string(); + let owner_id = owner_id.clone(); + let before = before.clone(); + let after = after.clone(); + Box::pin(async move { + let select_sql = format!( + r#" +SELECT + id, + user_id, + api_key_id, + CAST(balance AS DOUBLE PRECISION) AS balance, + CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + limit_mode, + currency, + status, + 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(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs +FROM wallets +WHERE id = $1 + AND {owner_clause} +LIMIT 1 +FOR UPDATE + "# + ); + let row = sqlx::query(&select_sql) + .bind(&after.id) + .bind(&owner_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if current != after { + return Ok(false); + } + let updated = sqlx::query( + r#" +UPDATE wallets +SET balance = $2, + gift_balance = $3, + limit_mode = $4, + currency = $5, + status = $6, + total_recharged = $7, + total_consumed = $8, + total_refunded = $9, + total_adjusted = $10, + updated_at = to_timestamp($11::DOUBLE PRECISION) +WHERE id = $1 +"#, + ) + .bind(&before.id) + .bind(before.balance) + .bind(before.gift_balance) + .bind(&before.limit_mode) + .bind(&before.currency) + .bind(&before.status) + .bind(before.total_recharged) + .bind(before.total_consumed) + .bind(before.total_refunded) + .bind(before.total_adjusted) + .bind(before_updated_at) + .execute(&mut **tx) + .await + .map_postgres_err()? + .rows_affected(); + Ok(updated > 0) + }) + }) + .await + } + + async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result { + if wallet_id.trim().is_empty() || user_id.trim().is_empty() { + return Ok(false); + } + let wallet_id = wallet_id.to_string(); + let user_id = user_id.to_string(); + self.tx_runner + .run_read_write(|tx| { + let wallet_id = wallet_id.clone(); + Box::pin(async move { + let found_wallet_id = sqlx::query_scalar::<_, String>( + r#" +SELECT w.id +FROM wallets AS w +WHERE w.id = $1 + AND w.user_id = $2 + AND w.api_key_id IS NULL + AND w.balance = 0 + AND w.gift_balance >= 0 + AND w.total_recharged = 0 + AND w.total_consumed = 0 + AND w.total_refunded = 0 + AND w.total_adjusted = w.gift_balance + AND w.limit_mode IN ('finite', 'unlimited') + AND w.currency = 'USD' + AND w.status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = w.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = w.id + ) + AND NOT EXISTS (SELECT 1 FROM usage u WHERE u.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = w.id + ) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = w.id + ) + AND ( + (w.gift_balance = 0 AND NOT EXISTS ( + SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = w.id + )) + OR + (w.gift_balance > 0 + AND (SELECT COUNT(*) FROM wallet_transactions t WHERE t.wallet_id = w.id) = 1 + AND EXISTS ( + SELECT 1 FROM wallet_transactions t + WHERE t.wallet_id = w.id + AND t.category = 'gift' + AND t.reason_code = 'gift_initial' + AND t.amount = w.gift_balance + AND t.balance_before = 0 + AND t.balance_after = w.gift_balance + AND t.recharge_balance_before = 0 + AND t.recharge_balance_after = 0 + AND t.gift_balance_before = 0 + AND t.gift_balance_after = w.gift_balance + AND t.link_type = 'system_task' + AND t.link_id = w.user_id + AND t.operator_id IS NULL + )) + ) +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(&wallet_id) + .bind(&user_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + let Some(found_wallet_id) = found_wallet_id else { + return Ok(false); + }; + sqlx::query("DELETE FROM wallet_transactions WHERE wallet_id = $1") + .bind(&found_wallet_id) + .execute(&mut **tx) + .await + .map_postgres_err()?; + let removed = sqlx::query("DELETE FROM wallets WHERE id = $1 AND user_id = $2") + .bind(&found_wallet_id) + .bind(&user_id) + .execute(&mut **tx) + .await + .map_postgres_err()? + .rows_affected() + > 0; + Ok(removed) + }) + }) + .await + } + async fn create_wallet_recharge_order( &self, - input: CreateWalletRechargeOrderInput, + mut input: CreateWalletRechargeOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "manual recharge amount must be finite and positive".to_string(), + )); + } + validate_wallet_recharge_order_input(&input).map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid wallet recharge numeric fields".to_string(), + )); + } + let projected_gateway_response = + project_wallet_recharge_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; self.tx_runner .run_read_write(|tx| { Box::pin(async move { - let wallet_row = match sqlx::query( + // `wallets.user_id` is nullable so deleted-user history + // can be retained. Check the live owner before any + // automatic wallet or payment-order insert. + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM public.users WHERE id = $1 FOR UPDATE") + .bind(&input.user_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } + + let (wallet_row, created_wallet) = match sqlx::query( r#" SELECT id, status FROM wallets @@ -1394,13 +2025,13 @@ FOR UPDATE .await .map_postgres_err()? { - Some(row) => row, + Some(row) => (row, false), None => { let wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| Uuid::new_v4().to_string()); - sqlx::query( + let inserted_wallet = sqlx::query( r#" INSERT INTO wallets ( id, @@ -1432,16 +2063,34 @@ VALUES ( NOW(), NOW() ) -ON CONFLICT (user_id) DO UPDATE -SET updated_at = wallets.updated_at +ON CONFLICT DO NOTHING RETURNING id, status "#, ) .bind(&wallet_id) .bind(&input.user_id) - .fetch_one(&mut **tx) + .fetch_optional(&mut **tx) .await - .map_postgres_err()? + .map_postgres_err()?; + match inserted_wallet { + Some(row) => (row, true), + None => { + let Some(row) = sqlx::query( + "SELECT id, status FROM wallets WHERE user_id = $1 LIMIT 1 FOR UPDATE", + ) + .bind(&input.user_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner" + .to_string(), + )); + }; + (row, false) + } + } } }; let wallet_id: String = row_get(&wallet_row, "id")?; @@ -1450,12 +2099,60 @@ RETURNING id, status return Ok(CreateWalletRechargeOrderOutcome::WalletInactive); } + // `order_no` is globally unique. Inspect the existing row + // before INSERT so callers get a deterministic idempotent + // replay only for their own wallet-recharge order; a + // collision with another user or order kind is invalid. + if let Some(existing_row) = sqlx::query( + "SELECT id, user_id, order_kind FROM payment_orders WHERE order_no = $1 FOR UPDATE", + ) + .bind(&input.order_no) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + { + let existing_user_id: Option = row_get(&existing_row, "user_id")?; + let existing_kind: String = row_get(&existing_row, "order_kind")?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + let existing = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(row_get::(&existing_row, "id")?) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + if !postgres_wallet_recharge_replay_matches( + &existing, + &wallet_id, + &input, + )? { + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if created_wallet { + sqlx::query(DELETE_PROVISIONAL_RECHARGE_WALLET_SQL) + .bind(&wallet_id) + .bind(&input.user_id) + .execute(&mut **tx) + .await + .map_postgres_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing( + map_admin_payment_order_row(&existing)?, + )); + } + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { DataLayerError::InvalidInput( "wallet recharge expires_at overflow".to_string(), ) })?; - let row = sqlx::query( + let insert_result = sqlx::query( r#" INSERT INTO payment_orders ( id, @@ -1501,6 +2198,7 @@ VALUES ( NOW(), to_timestamp($14) ) +ON CONFLICT DO NOTHING RETURNING id, order_no, @@ -1539,11 +2237,73 @@ RETURNING .bind(input.payment_provider.as_deref()) .bind(input.payment_channel.as_deref()) .bind(&input.gateway_order_id) - .bind(&input.gateway_response) + .bind(&projected_gateway_response) .bind(expires_at) - .fetch_one(&mut **tx) + .fetch_optional(&mut **tx) .await .map_postgres_err()?; + let Some(row) = insert_result else { + let existing_row = sqlx::query( + "SELECT id, user_id, order_kind FROM payment_orders WHERE order_no = $1 FOR UPDATE", + ) + .bind(&input.order_no) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if let Some(existing_row) = existing_row { + let existing_user_id: Option = + row_get(&existing_row, "user_id")?; + let existing_kind: String = row_get(&existing_row, "order_kind")?; + let existing = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(row_get::(&existing_row, "id")?) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !postgres_wallet_recharge_replay_matches( + &existing, + &wallet_id, + &input, + )? { + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if created_wallet { + sqlx::query(DELETE_PROVISIONAL_RECHARGE_WALLET_SQL) + .bind(&wallet_id) + .bind(&input.user_id) + .execute(&mut **tx) + .await + .map_postgres_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing( + map_admin_payment_order_row(&existing)?, + )); + } + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + let gateway_conflict = sqlx::query( + "SELECT 1 FROM payment_orders WHERE payment_method = $1 AND gateway_order_id = $2 LIMIT 1", + ) + .bind(&input.payment_method) + .bind(&input.gateway_order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if gateway_conflict.is_some() { + return Err(DataLayerError::InvalidInput( + "payment gateway order already belongs to another order".to_string(), + )); + } + return Err(DataLayerError::InvalidInput( + "wallet recharge order could not be created".to_string(), + )); + }; Ok(CreateWalletRechargeOrderOutcome::Created( map_admin_payment_order_row(&row)?, )) @@ -1552,13 +2312,369 @@ RETURNING .await } - async fn create_plan_purchase_order( + async fn update_wallet_recharge_checkout( &self, - input: CreatePlanPurchaseOrderInput, - ) -> Result { + input: UpdateWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() || input.gateway_order_id.trim().is_empty() { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout identifiers are required".to_string(), + )); + } + let projected_gateway_response = + match project_wallet_recharge_gateway_response(&input.gateway_response) { + Ok(value) => value, + Err(error) => return Ok(WalletMutationOutcome::Invalid(error)), + }; self.tx_runner .run_read_write(|tx| { Box::pin(async move { + let Some(current_row) = sqlx::query( + "SELECT id, order_no, payment_method, payment_provider, payment_channel, order_kind, gateway_order_id, gateway_response, status, COALESCE(expires_at > NOW(), FALSE) AS checkout_live FROM payment_orders WHERE id = $1 FOR UPDATE", + ) + .bind(&input.order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = row_get(¤t_row, "order_kind")?; + if order_kind.as_deref() != Some("wallet_recharge") { + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a wallet recharge".to_string(), + )); + } + let current_gateway_response: Option = + row_get(¤t_row, "gateway_response")?; + let current_is_checkout_placeholder = current_gateway_response + .as_ref() + .is_some_and(wallet_recharge_response_is_checkout_placeholder); + let current_token = current_gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + let requested_token = + wallet_recharge_checkout_claim_token(&projected_gateway_response); + if current_token.is_some() && current_token != requested_token { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + let current_status: String = row_get(¤t_row, "status")?; + let current_gateway: Option = row_get(¤t_row, "gateway_order_id")?; + if current_status != "pending" { + if current_gateway.as_deref() == Some(input.gateway_order_id.as_str()) { + let row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&input.order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + return Ok(WalletMutationOutcome::Applied( + map_admin_payment_order_row(&row)?, + )); + } + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is no longer pending".to_string(), + )); + } + let checkout_live: bool = row_get(¤t_row, "checkout_live")?; + if !checkout_live { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired".to_string(), + )); + } + let order_no: String = row_get(¤t_row, "order_no")?; + if current_gateway.as_deref().is_some_and(|existing| { + existing != input.gateway_order_id.as_str() + && existing != order_no.as_str() + && !current_is_checkout_placeholder + }) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is already bound".to_string(), + )); + } + let payment_method: String = row_get(¤t_row, "payment_method")?; + let conflict = sqlx::query( + "SELECT id FROM payment_orders WHERE payment_method = $1 AND gateway_order_id = $2 AND id <> $3 LIMIT 1", + ) + .bind(&payment_method) + .bind(&input.gateway_order_id) + .bind(&input.order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if conflict.is_some() { + return Ok(WalletMutationOutcome::Invalid( + "payment gateway order already belongs to another order".to_string(), + )); + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = $2, gateway_response = $3 WHERE id = $1 AND status = 'pending' AND expires_at > NOW()", + ) + .bind(&input.order_id) + .bind(&input.gateway_order_id) + .bind(&projected_gateway_response) + .execute(&mut **tx) + .await + .map_postgres_err()?; + if updated.rows_affected() == 0 { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired or no longer pending".to_string(), + )); + } + let row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&input.order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + Ok(WalletMutationOutcome::Applied(map_admin_payment_order_row(&row)?)) + }) + }) + .await + } + + async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result { + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + let locked: Option = sqlx::query_scalar( + "SELECT id FROM payment_orders WHERE id = $1 FOR UPDATE", + ) + .bind(&input.order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if locked.is_none() { + return Ok(false); + } + let row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&input.order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let current = map_admin_payment_order_row(&row)?; + let Some(replacement) = + payment_order_stripe_client_secret_cas_replacement(¤t, &input) + .map_err(DataLayerError::InvalidInput)? + else { + return Ok(false); + }; + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_response = $2 WHERE id = $1", + ) + .bind(&input.order_id) + .bind(replacement) + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(updated.rows_affected() == 1) + }) + }) + .await + } + + async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout failure identifiers are required".to_string(), + )); + } + let order_id = input.order_id.clone(); + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + let Some(row) = sqlx::query( + "SELECT id, order_kind, gateway_response, status FROM payment_orders WHERE id = $1 FOR UPDATE", + ) + .bind(&order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = row_get(&row, "order_kind")?; + let gateway_response: Option = + row_get(&row, "gateway_response")?; + let status: String = row_get(&row, "status")?; + if order_kind.as_deref() != Some("wallet_recharge") + || !gateway_response + .as_ref() + .is_some_and(wallet_recharge_response_is_checkout_placeholder) + { + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a checkout placeholder".to_string(), + )); + } + let current_token = gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + if current_token != Some(input.claim_token.trim()) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + let full_row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let order = map_admin_payment_order_row(&full_row)?; + if status != "pending" { + return Ok(WalletMutationOutcome::Applied(order)); + } + let failed = if input.provider_request_may_have_succeeded { + wallet_recharge_checkout_uncertain_response( + gateway_response.as_ref(), + &input.reason, + Utc::now().timestamp().max(0) as u64, + ) + } else { + wallet_recharge_checkout_failed_response( + gateway_response.as_ref(), + &input.reason, + Utc::now().timestamp().max(0) as u64, + ) + }; + let updated = sqlx::query( + "UPDATE payment_orders SET status = 'failed', gateway_response = $2 WHERE id = $1 AND status = 'pending'", + ) + .bind(&order_id) + .bind(failed) + .execute(&mut **tx) + .await + .map_postgres_err()?; + if updated.rows_affected() == 0 { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + let row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + Ok(WalletMutationOutcome::Applied(map_admin_payment_order_row(&row)?)) + }) + }) + .await + } + + async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + let now = Utc::now().timestamp().max(0) as u64; + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + || input.expires_at_unix_secs <= now + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim identifiers are invalid".to_string(), + )); + } + if !wallet_recharge_response_is_checkout_placeholder(&input.gateway_response) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge reclaim response must be a placeholder".to_string(), + )); + } + let response = wallet_recharge_checkout_claim_response( + &input.gateway_response, + &input.claim_token, + now, + ) + .map_err(DataLayerError::InvalidInput)?; + let order_id = input.order_id.clone(); + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + let Some(row) = sqlx::query( + "SELECT id, order_no, order_kind, gateway_response, status, CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs FROM payment_orders WHERE id = $1 FOR UPDATE", + ) + .bind(&order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = row_get(&row, "order_kind")?; + let full_row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let order = map_admin_payment_order_row(&full_row)?; + if order_kind.as_deref() != Some("wallet_recharge") + || !wallet_recharge_order_is_reclaimable_placeholder(&order, now) + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is still in progress or already completed".to_string(), + )); + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = $2, gateway_response = $3, status = 'pending', expires_at = to_timestamp($4) WHERE id = $1", + ) + .bind(&order_id) + .bind(&order.order_no) + .bind(response) + .bind(i64::try_from(input.expires_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput("wallet recharge expires_at overflow".to_string()) + })?) + .execute(&mut **tx) + .await + .map_postgres_err()?; + if updated.rows_affected() == 0 { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim lost the order race".to_string(), + )); + } + let row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL) + .bind(&order_id) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + Ok(WalletMutationOutcome::Applied(map_admin_payment_order_row(&row)?)) + }) + }) + .await + } + + async fn create_plan_purchase_order( + &self, + mut input: CreatePlanPurchaseOrderInput, + ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_plan_purchase_order_input(&input).map_err(DataLayerError::InvalidInput)?; + let projected_gateway_response = project_wallet_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + // Validate and lock the order owner before the automatic + // wallet upsert. This keeps a missing-user checkout from + // creating an orphan financial row. + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(&input.user_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } + let wallet_row = match sqlx::query( r#" SELECT id, status @@ -1579,7 +2695,7 @@ FOR UPDATE .preferred_wallet_id .clone() .unwrap_or_else(|| Uuid::new_v4().to_string()); - sqlx::query( + let inserted_wallet = sqlx::query( r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, @@ -1587,16 +2703,34 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES ($1, $2, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, NOW(), NOW()) -ON CONFLICT (user_id) DO UPDATE -SET updated_at = wallets.updated_at +ON CONFLICT DO NOTHING RETURNING id, status "#, ) .bind(&wallet_id) .bind(&input.user_id) - .fetch_one(&mut **tx) + .fetch_optional(&mut **tx) .await - .map_postgres_err()? + .map_postgres_err()?; + match inserted_wallet { + Some(row) => row, + None => { + let Some(row) = sqlx::query( + "SELECT id, status FROM wallets WHERE user_id = $1 LIMIT 1 FOR UPDATE", + ) + .bind(&input.user_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner" + .to_string(), + )); + }; + row + } + } } }; let wallet_id: String = row_get(&wallet_row, "id")?; @@ -1737,7 +2871,7 @@ RETURNING .bind(&input.product_id) .bind(&input.product_snapshot) .bind(&input.gateway_order_id) - .bind(&input.gateway_response) + .bind(&projected_gateway_response) .bind(expires_at) .fetch_one(&mut **tx) .await @@ -1754,6 +2888,11 @@ RETURNING &self, input: CreateWalletRefundRequestInput, ) -> Result { + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "refund amount must be finite and greater than zero".to_string(), + )); + } self.tx_runner .run_read_write(|tx| { Box::pin(async move { @@ -1764,11 +2903,13 @@ SELECT CAST(balance AS DOUBLE PRECISION) AS balance FROM wallets WHERE id = $1 + AND user_id = $2 LIMIT 1 FOR UPDATE "#, ) .bind(&input.wallet_id) + .bind(&input.user_id) .fetch_optional(&mut **tx) .await .map_postgres_err()? @@ -1776,19 +2917,37 @@ FOR UPDATE return Ok(CreateWalletRefundRequestOutcome::WalletMissing); }; let wallet_recharge_balance: f64 = row_get(&locked_wallet_row, "balance")?; - let wallet_reserved_row = sqlx::query( + if !wallet_recharge_balance.is_finite() { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet recharge balance is invalid".to_string(), + )); + } + let wallet_reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(CAST(SUM(amount_usd) AS DOUBLE PRECISION), 0) AS total +SELECT CAST(amount_usd AS DOUBLE PRECISION) FROM refund_requests WHERE wallet_id = $1 AND status IN ('pending_approval', 'approved') "#, ) .bind(&input.wallet_id) - .fetch_one(&mut **tx) + .fetch_all(&mut **tx) .await - .map_postgres_err()?; - let wallet_reserved_amount: f64 = row_get(&wallet_reserved_row, "total")?; + .map_postgres_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(wallet_reserved_amount) = wallet_reserved_amount else { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet refund reservation is invalid".to_string(), + )); + }; if input.amount_usd > (wallet_recharge_balance - wallet_reserved_amount) { return Ok( CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableBalance, @@ -1796,15 +2955,7 @@ WHERE wallet_id = $1 } let mut payment_order_id = None; - let mut source_type = input - .source_type - .clone() - .unwrap_or_else(|| "wallet_balance".to_string()); - let mut source_id = input.source_id.clone(); - let mut refund_mode = input - .refund_mode - .clone() - .unwrap_or_else(|| "offline_payout".to_string()); + let mut resolved_payment_method = None; if let Some(order_id) = input.payment_order_id.as_deref() { let Some(order_row) = sqlx::query( r#" @@ -1812,6 +2963,8 @@ SELECT id, status, payment_method, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd, CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd FROM payment_orders WHERE id = $1 @@ -1832,36 +2985,67 @@ FOR UPDATE if status != "credited" { return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotRefundable); } - let order_reserved_row = sqlx::query( + let reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(CAST(SUM(amount_usd) AS DOUBLE PRECISION), 0) AS total +SELECT CAST(amount_usd AS DOUBLE PRECISION) FROM refund_requests WHERE payment_order_id = $1 AND status IN ('pending_approval', 'approved') "#, ) .bind(order_id) - .fetch_one(&mut **tx) + .fetch_all(&mut **tx) .await - .map_postgres_err()?; + .map_postgres_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(reserved_amount) = reserved_amount else { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund reservation is invalid".to_string(), + )); + }; + let order_amount: f64 = row_get(&order_row, "amount_usd")?; + let refunded_amount: f64 = + row_get(&order_row, "refunded_amount_usd")?; let refundable_amount: f64 = row_get(&order_row, "refundable_amount_usd")?; - let reserved_amount: f64 = row_get(&order_reserved_row, "total")?; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_amount, + refundable_amount, + ) { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund amounts are invalid".to_string(), + )); + } if input.amount_usd > (refundable_amount - reserved_amount) { return Ok( CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableOrderAmount, ); } payment_order_id = Some(order_id.to_string()); - source_type = "payment_order".to_string(); - source_id = Some(order_id.to_string()); - if input.refund_mode.is_none() { - let payment_method: String = row_get(&order_row, "payment_method")?; - refund_mode = - default_refund_mode_for_payment_method(&payment_method).to_string(); - } + resolved_payment_method = Some(row_get::(&order_row, "payment_method")?); } + let canonical = canonicalize_wallet_refund_fields( + payment_order_id.as_deref(), + input.source_type.as_deref(), + input.source_id.as_deref(), + input.refund_mode.as_deref(), + resolved_payment_method.as_deref(), + ) + .map_err(DataLayerError::InvalidInput)?; + let source_type = canonical.source_type; + let source_id = canonical.source_id; + let refund_mode = canonical.refund_mode; + let insert_result = sqlx::query( r#" INSERT INTO refund_requests ( @@ -1898,6 +3082,7 @@ VALUES ( NOW(), NOW() ) +ON CONFLICT (idempotency_key) DO NOTHING RETURNING id, refund_no, @@ -1936,15 +3121,14 @@ RETURNING .bind(input.reason.as_deref()) .bind(&input.user_id) .bind(input.idempotency_key.as_deref()) - .fetch_one(&mut **tx) - .await; + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; match insert_result { - Ok(row) => Ok(CreateWalletRefundRequestOutcome::Created( + Some(row) => Ok(CreateWalletRefundRequestOutcome::Created( map_admin_wallet_refund_row(&row)?, )), - Err(sqlx::Error::Database(err)) - if err.code().as_deref() == Some("23505") => - { + None => { if let Some(idempotency_key) = input.idempotency_key.as_deref() { let existing = sqlx::query( r#" @@ -1991,7 +3175,6 @@ LIMIT 1 } Ok(CreateWalletRefundRequestOutcome::DuplicateRejected) } - Err(err) => Err(postgres_error(err)), } }) }) @@ -2000,37 +3183,33 @@ LIMIT 1 async fn process_payment_callback( &self, - input: ProcessPaymentCallbackInput, + mut input: ProcessPaymentCallbackInput, ) -> Result { + input + .canonicalize_and_validate() + .map_err(DataLayerError::InvalidInput)?; + if input.callback_key.trim().is_empty() + || input.callback_key.chars().count() > 128 + || input.payload_hash.trim().is_empty() + || !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid payment callback numeric or identity fields".to_string(), + )); + } self.tx_runner .run_read_write(|tx| { Box::pin(async move { - let existing_callback = sqlx::query( + let candidate_callback_id = Uuid::new_v4().to_string(); + sqlx::query( r#" -SELECT id, payment_order_id, status, order_no, gateway_order_id -FROM payment_callbacks -WHERE callback_key = $1 -LIMIT 1 - "#, - ) - .bind(&input.callback_key) - .fetch_optional(&mut **tx) - .await - .map_postgres_err()?; - - let duplicate = existing_callback.is_some(); - let callback_id = if let Some(row) = existing_callback.as_ref() { - let status: String = row_get(row, "status")?; - if status == "processed" { - return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { - order_id: row_get(row, "payment_order_id")?, - }); - } - row_get(row, "id")? - } else { - let callback_id = Uuid::new_v4().to_string(); - sqlx::query( - r#" INSERT INTO payment_callbacks ( id, payment_order_id, @@ -2056,26 +3235,59 @@ VALUES ( $6, $7, 'received', - $8, + NULL, NULL, NOW(), NULL ) - "#, - ) - .bind(&callback_id) - .bind(&input.payment_method) - .bind(&input.callback_key) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) - .bind(&input.payload_hash) - .bind(input.signature_valid) - .bind(&input.payload) - .execute(&mut **tx) - .await - .map_postgres_err()?; - callback_id - }; +ON CONFLICT (callback_key) DO NOTHING + "#, + ) + .bind(&candidate_callback_id) + .bind(&input.payment_method) + .bind(&input.callback_key) + .bind(input.order_no.as_deref()) + .bind(input.gateway_order_id.as_deref()) + .bind(&input.payload_hash) + .bind(input.signature_valid) + .execute(&mut **tx) + .await + .map_postgres_err()?; + + let callback_row = sqlx::query( + r#" +SELECT id, payment_order_id, payment_method, payload_hash, status, order_no, gateway_order_id +FROM payment_callbacks +WHERE callback_key = $1 +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(&input.callback_key) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let callback_id: String = row_get(&callback_row, "id")?; + let duplicate = callback_id != candidate_callback_id; + let callback_order_no: Option = row_get(&callback_row, "order_no")?; + let callback_gateway_order_id: Option = + row_get(&callback_row, "gateway_order_id")?; + let stored_method: String = row_get(&callback_row, "payment_method")?; + let stored_hash: Option = row_get(&callback_row, "payload_hash")?; + if !stored_method.eq_ignore_ascii_case(&input.payment_method) + || stored_hash.as_deref() != Some(input.payload_hash.as_str()) + { + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate: true, + error: "callback key reused with different payment payload".to_string(), + }); + } + let status: String = row_get(&callback_row, "status")?; + if status == "processed" { + return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { + order_id: row_get(&callback_row, "payment_order_id")?, + }); + } if !input.signature_valid { update_payment_callback_failure( @@ -2091,16 +3303,12 @@ VALUES ( }); } - let lookup_order_no = input.order_no.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| row.try_get("order_no").ok()) - }); - let lookup_gateway_order_id = input.gateway_order_id.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| row.try_get("gateway_order_id").ok()) - }); + let lookup_order_no = + input.order_no.clone().or_else(|| callback_order_no.clone()); + let lookup_gateway_order_id = input + .gateway_order_id + .clone() + .or_else(|| callback_gateway_order_id.clone()); let order_row = if let Some(order_no) = lookup_order_no.as_deref() { sqlx::query( @@ -2167,11 +3375,13 @@ SELECT CAST(EXTRACT(EPOCH FROM credited_at) AS BIGINT) AS credited_at_unix_secs, CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs FROM payment_orders -WHERE gateway_order_id = $1 +WHERE payment_method = $1 + AND gateway_order_id = $2 LIMIT 1 FOR UPDATE "#, ) + .bind(&input.payment_method) .bind(gateway_order_id) .fetch_optional(&mut **tx) .await @@ -2202,21 +3412,263 @@ FOR UPDATE row_get(&order_row, "payment_provider")?; let order_payment_channel: Option = row_get(&order_row, "payment_channel")?; + let order_pay_currency: Option = row_get(&order_row, "pay_currency")?; + let order_gateway_order_id: Option = + row_get(&order_row, "gateway_order_id")?; let order_kind: String = row_get(&order_row, "order_kind")?; let order_amount_usd: f64 = row_get(&order_row, "amount_usd")?; let order_pay_amount: Option = row_get(&order_row, "pay_amount")?; + let order_exchange_rate: Option = + row_get(&order_row, "exchange_rate")?; let order_status: String = row_get(&order_row, "status")?; let expires_at_unix_secs: Option = row_get(&order_row, "expires_at_unix_secs")?; - - let amount_matches = - if let (Some(callback_pay_amount), Some(order_pay_amount)) = - (input.pay_amount, order_pay_amount) - { - (callback_pay_amount - order_pay_amount).abs() <= 0.01 + let order_gateway_response: Option = + if order_status.eq_ignore_ascii_case("failed") { + row_get(&order_row, "gateway_response")? } else { - (input.amount_usd - order_amount_usd).abs() <= f64::EPSILON + None }; + let failed_checkout_recoverable = + payment_order_is_failed_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + let uncertain_checkout = + payment_order_is_uncertain_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + if !order_amount_usd.is_finite() + || order_amount_usd <= 0.0 + || order_pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment order amount is invalid", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order amount is invalid".to_string(), + }); + } + + // A payment order may only credit a wallet owned by the + // same live user. Lock the wallet and joined API-key row + // before any gateway binding, entitlement, wallet, or + // order mutation. Reject legacy rows with an ambiguous + // owner shape instead of guessing an owner. + let order_user_id: Option = row_get(&order_row, "user_id")?; + let Some(order_user_id) = order_user_id + .as_deref() + .filter(|value| !value.trim().is_empty()) + else { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment order user missing", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order user missing".to_string(), + }); + }; + let Some(wallet_owner_row) = sqlx::query( + r#" +SELECT + w.user_id AS wallet_user_id, + w.api_key_id AS wallet_api_key_id, + api_keys.user_id AS api_key_user_id +FROM wallets AS w +LEFT JOIN api_keys ON api_keys.id = w.api_key_id +WHERE w.id = $1 +LIMIT 1 +FOR UPDATE + "#, + ) + .bind(&order_wallet_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "wallet not found", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet not found".to_string(), + }); + }; + let wallet_user_id: Option = + row_get(&wallet_owner_row, "wallet_user_id")?; + let wallet_api_key_id: Option = + row_get(&wallet_owner_row, "wallet_api_key_id")?; + let api_key_user_id: Option = + row_get(&wallet_owner_row, "api_key_user_id")?; + let wallet_owner_matches = match ( + wallet_user_id.as_deref(), + wallet_api_key_id.as_deref(), + api_key_user_id.as_deref(), + ) { + (Some(wallet_user_id), None, _) + if !wallet_user_id.trim().is_empty() => + { + wallet_user_id == order_user_id + } + (None, Some(wallet_api_key_id), Some(api_key_user_id)) + if !wallet_api_key_id.trim().is_empty() + && !api_key_user_id.trim().is_empty() => + { + api_key_user_id == order_user_id + } + _ => false, + }; + if !wallet_owner_matches { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment order wallet owner mismatch", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order wallet owner mismatch".to_string(), + }); + } + + // The lookup identifier is not proof that the callback + // belongs to this order: order_no takes precedence over + // gateway_order_id. Check every identifier supplied by + // this delivery (and any persisted fallback from the + // callback row) before changing the order or wallet. + // Orders created before the gateway returns a provider + // transaction id store order_no as a placeholder; that + // value may be replaced by a verified callback, but a + // real id must never be rebound to another order. + if input + .order_no + .as_deref() + .is_some_and(|value| value != order_no) + || callback_order_no + .as_deref() + .is_some_and(|value| value != order_no) + { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment order number mismatch", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order number mismatch".to_string(), + }); + } + let input_gateway_order_id = input + .gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let callback_gateway_order_id = callback_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let stored_real_gateway_order_id = order_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty() && *value != order_no); + let input_real_gateway_order_id = + input_gateway_order_id.filter(|value| *value != order_no); + let callback_real_gateway_order_id = + callback_gateway_order_id.filter(|value| *value != order_no); + let effective_gateway_order_id = input_real_gateway_order_id + .or(callback_real_gateway_order_id) + .or(stored_real_gateway_order_id); + if let Some(expected_gateway_order_id) = stored_real_gateway_order_id { + if input_gateway_order_id + .is_some_and(|value| value != expected_gateway_order_id) + || callback_gateway_order_id + .is_some_and(|value| value != expected_gateway_order_id) + { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment gateway order mismatch", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order mismatch".to_string(), + }); + } + } else if let (Some(input_gateway), Some(callback_gateway)) = + (input_real_gateway_order_id, callback_real_gateway_order_id) + { + if input_gateway != callback_gateway { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment gateway order identifier mismatch", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order identifier mismatch".to_string(), + }); + } + } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + let conflicting_order_id: Option = sqlx::query_scalar( + "SELECT id FROM payment_orders WHERE payment_method = $1 AND gateway_order_id = $2 AND id <> $3 LIMIT 1", + ) + .bind(&order_payment_method) + .bind(gateway_order_id) + .bind(&order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()?; + if conflicting_order_id.is_some() { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment gateway order belongs to another payment order", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order" + .to_string(), + }); + } + } + } + + let amount_matches = payment_callback_amount_matches_order( + order_amount_usd, + order_pay_amount, + order_pay_currency.as_deref(), + order_exchange_rate, + input.amount_usd, + input.pay_amount, + ); if !amount_matches { update_payment_callback_failure( tx, @@ -2230,7 +3682,12 @@ FOR UPDATE error: "callback amount mismatch".to_string(), }); } - if !order_payment_method.eq_ignore_ascii_case(&input.payment_method) { + if !payment_callback_method_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ) { update_payment_callback_failure( tx, &callback_id, @@ -2243,11 +3700,13 @@ FOR UPDATE error: "payment method mismatch".to_string(), }); } - if let Some(expected_provider) = input.payment_provider.as_deref() { - if order_payment_provider - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_provider)) - { + let payment_provider_matches = payment_callback_provider_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ); + if !payment_provider_matches { update_payment_callback_failure( tx, &callback_id, @@ -2259,12 +3718,38 @@ FOR UPDATE duplicate, error: "payment provider mismatch".to_string(), }); - } + } + let currency_matches = + match (input.pay_currency.as_deref(), order_pay_currency.as_deref()) { + (Some(callback), Some(order)) => { + order.eq_ignore_ascii_case(callback) + } + (None, None) => input.pay_amount.is_none() && order_pay_amount.is_none(), + _ => false, + }; + if !currency_matches { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment currency mismatch", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment currency mismatch".to_string(), + }); } if let Some(expected_channel) = input.payment_channel.as_deref() { - if order_payment_channel - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_channel)) + let stored_channel = order_payment_channel.as_deref().or_else(|| { + (order_payment_provider.is_none() + && ["alipay", "wxpay"].iter().any(|method| { + method.eq_ignore_ascii_case(&order_payment_method) + })) + .then_some(order_payment_method.as_str()) + }); + if stored_channel + .is_none_or(|value| !value.eq_ignore_ascii_case(expected_channel)) { update_payment_callback_failure( tx, @@ -2295,14 +3780,18 @@ FOR UPDATE wallet_id: order_wallet_id, }); } - if matches!(order_status.as_str(), "failed" | "expired" | "refunded") { + if !matches!(order_status.as_str(), "pending" | "paid") + && !failed_checkout_recoverable + { let error = format!("payment order is not creditable: {order_status}"); update_payment_callback_failure(tx, &callback_id, &input, &error).await?; return Ok(ProcessPaymentCallbackOutcome::Failed { duplicate, error }); } - if order_status == "pending" { + if order_status == "pending" + || (failed_checkout_recoverable && !uncertain_checkout) + { let now = Utc::now().timestamp(); - if expires_at_unix_secs.is_some_and(|value| value < now) { + if expires_at_unix_secs.is_some_and(|value| value <= now) { sqlx::query( "UPDATE payment_orders SET status = 'expired' WHERE id = $1", ) @@ -2324,6 +3813,31 @@ FOR UPDATE } } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + if !postgres_bind_payment_gateway_order_id( + tx, + &order_id, + gateway_order_id, + ) + .await? + { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "payment gateway order belongs to another payment order", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order" + .to_string(), + }); + } + } + } + if order_kind == "plan_purchase" { let product_id: Option = row_get(&order_row, "product_id")?; let product_snapshot: Option = @@ -2352,7 +3866,7 @@ FOR UPDATE }); let entitlements = plan_entitlements_snapshot(&snapshot); let now = Utc::now(); - let expires_at = plan_expires_at(&snapshot, now); + let expires_at = plan_expires_at(&snapshot, now)?; let existing_entitlement_id = sqlx::query_scalar::<_, String>( r#" SELECT id @@ -2457,9 +3971,9 @@ VALUES ($1, $2, $3, $4, 'active', $5, $6, $7, NOW(), NOW()) UPDATE payment_orders SET gateway_order_id = COALESCE($2, gateway_order_id), gateway_response = $3, - pay_amount = COALESCE($4, pay_amount), - pay_currency = COALESCE($5, pay_currency), - exchange_rate = COALESCE($6, exchange_rate), + pay_amount = COALESCE(pay_amount, $4), + pay_currency = COALESCE(pay_currency, $5), + exchange_rate = COALESCE(exchange_rate, $6), status = 'credited', fulfillment_status = 'fulfilled', fulfillment_error = NULL, @@ -2494,8 +4008,11 @@ RETURNING "#, ) .bind(&order_id) - .bind(input.gateway_order_id.as_deref()) - .bind(&input.payload) + .bind(effective_gateway_order_id) + .bind(input.gateway_response_projection( + &order_no, + effective_gateway_order_id, + )) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -2525,7 +4042,8 @@ SELECT id, status, CAST(balance AS DOUBLE PRECISION) AS balance, - CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance + CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(total_recharged AS DOUBLE PRECISION) AS total_recharged FROM wallets WHERE id = $1 LIMIT 1 @@ -2566,6 +4084,30 @@ FOR UPDATE let before_recharge: f64 = row_get(&wallet_row, "balance")?; let before_gift: f64 = row_get(&wallet_row, "gift_balance")?; + let total_recharged: f64 = row_get(&wallet_row, "total_recharged")?; + // Finite recharge balances may be negative: usage settlement permits a + // finite wallet to overdraft, and a later recharge must be able to restore + // that balance. Reject only malformed values and arithmetic overflow here. + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order_amount_usd).is_finite() + || !(before_recharge + before_gift + order_amount_usd).is_finite() + { + update_payment_callback_failure( + tx, + &callback_id, + &input, + "wallet balance is invalid", + ) + .await?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet balance is invalid".to_string(), + }); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order_amount_usd; let after_total = after_recharge + before_gift; @@ -2646,9 +4188,9 @@ VALUES ( UPDATE payment_orders SET gateway_order_id = COALESCE($2, gateway_order_id), gateway_response = $3, - pay_amount = COALESCE($4, pay_amount), - pay_currency = COALESCE($5, pay_currency), - exchange_rate = COALESCE($6, exchange_rate), + pay_amount = COALESCE(pay_amount, $4), + pay_currency = COALESCE(pay_currency, $5), + exchange_rate = COALESCE(exchange_rate, $6), status = 'credited', paid_at = COALESCE(paid_at, NOW()), credited_at = NOW(), @@ -2681,8 +4223,11 @@ RETURNING "#, ) .bind(&order_id) - .bind(input.gateway_order_id.as_deref()) - .bind(&input.payload) + .bind(effective_gateway_order_id) + .bind(input.gateway_response_projection( + &order_no, + effective_gateway_order_id, + )) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -2707,6 +4252,11 @@ RETURNING &self, input: AdjustWalletBalanceInput, ) -> 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(), + )); + } self.tx_runner .run_read_write(|tx| { Box::pin(async move { @@ -2741,6 +4291,16 @@ FOR UPDATE let before_recharge: f64 = row_get(&row, "balance")?; let before_gift: f64 = row_get(&row, "gift_balance")?; let before_total = before_recharge + before_gift; + let before_total_adjusted: f64 = row_get(&row, "total_adjusted")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || !before_total.is_finite() + || !before_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid".to_string(), + )); + } let mut after_recharge = before_recharge; let mut after_gift = before_gift; @@ -2772,6 +4332,17 @@ FOR UPDATE after_recharge -= remaining; } } + let after_total = after_recharge + after_gift; + let after_total_adjusted = before_total_adjusted + input.amount_usd; + if !after_recharge.is_finite() + || !after_gift.is_finite() + || !after_total.is_finite() + || !after_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow during admin adjustment".to_string(), + )); + } let wallet_row = sqlx::query( r#" @@ -2859,7 +4430,7 @@ VALUES ( .bind(&input.wallet_id) .bind(input.amount_usd) .bind(before_total) - .bind(after_recharge + after_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -2880,7 +4451,7 @@ VALUES ( reason_code: "adjust_admin".to_string(), amount: input.amount_usd, balance_before: before_total, - balance_after: after_recharge + after_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -2901,8 +4472,15 @@ VALUES ( async fn create_manual_wallet_recharge( &self, - input: CreateManualWalletRechargeInput, + mut input: CreateManualWalletRechargeInput, ) -> Result, DataLayerError> { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "manual recharge amount must be finite and positive".to_string(), + )); + } self.tx_runner .run_read_write(|tx| { Box::pin(async move { @@ -2936,6 +4514,14 @@ FOR UPDATE let before_recharge: f64 = row_get(&wallet_row, "balance")?; let before_gift: f64 = row_get(&wallet_row, "gift_balance")?; + let before_total_recharged: f64 = row_get(&wallet_row, "total_recharged")?; + let (after_recharge, after_total_recharged) = validate_manual_wallet_recharge( + input.amount_usd, + before_recharge, + before_gift, + before_total_recharged, + ) + .map_err(DataLayerError::InvalidInput)?; let user_id: Option = row_get(&wallet_row, "user_id")?; let gateway_response = serde_json::json!({ "source": "manual", @@ -2989,13 +4575,12 @@ VALUES ( .await .map_postgres_err()?; - let after_recharge = before_recharge + input.amount_usd; let wallet_row = sqlx::query( r#" UPDATE wallets SET balance = $2, - total_recharged = total_recharged + $3, + total_recharged = $3, updated_at = NOW() WHERE id = $1 RETURNING @@ -3016,7 +4601,7 @@ RETURNING ) .bind(&input.wallet_id) .bind(after_recharge) - .bind(input.amount_usd) + .bind(after_total_recharged) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -3182,6 +4767,11 @@ FOR UPDATE return Ok(WalletMutationOutcome::NotFound); }; let refund = map_admin_wallet_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if !matches!(refund.status.as_str(), "approved" | "pending_approval") { return Ok(WalletMutationOutcome::Invalid( "refund status is not approvable".to_string(), @@ -3220,9 +4810,27 @@ FOR UPDATE }; let before_recharge: f64 = row_get(&wallet_row, "balance")?; let before_gift: f64 = row_get(&wallet_row, "gift_balance")?; - let before_total = before_recharge + before_gift; + let before_total_refunded: f64 = row_get(&wallet_row, "total_refunded")?; let amount_usd = refund.amount_usd; let after_recharge = before_recharge - amount_usd; + let before_total = before_recharge + before_gift; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded + amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + { + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } if after_recharge < 0.0 { return Ok(WalletMutationOutcome::Invalid( "refund amount exceeds refundable recharge balance".to_string(), @@ -3234,6 +4842,9 @@ FOR UPDATE r#" SELECT id, + wallet_id, + status, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd, CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd FROM payment_orders @@ -3250,26 +4861,54 @@ FOR UPDATE "payment order not found".to_string(), )); }; - let refundable_amount: f64 = row_get(&order_row, "refundable_amount_usd")?; - if amount_usd > refundable_amount { + let order_wallet_id: String = row_get(&order_row, "wallet_id")?; + let order_status: String = row_get(&order_row, "status")?; + if order_wallet_id != input.wallet_id || order_status != "credited" { return Ok(WalletMutationOutcome::Invalid( - "refund amount exceeds refundable amount".to_string(), + "payment order is not refundable for this wallet".to_string(), )); } - sqlx::query( + let order_amount: f64 = row_get(&order_row, "amount_usd")?; + let refunded_before: f64 = row_get(&order_row, "refunded_amount_usd")?; + let refundable_before: f64 = row_get(&order_row, "refundable_amount_usd")?; + let refunded_after = refunded_before + amount_usd; + let refundable_after = refundable_before - amount_usd; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_before, + refundable_before, + ) || amount_usd > refundable_before + || !refunded_after.is_finite() + || refunded_after < 0.0 + || refunded_after > order_amount + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order_amount + { + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + let result = sqlx::query( r#" UPDATE payment_orders SET - refunded_amount_usd = refunded_amount_usd + $2, - refundable_amount_usd = refundable_amount_usd - $2 + refunded_amount_usd = $2, + refundable_amount_usd = $3 WHERE id = $1 "#, ) .bind(payment_order_id) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .execute(&mut **tx) .await .map_postgres_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund processing".to_string(), + )); + } } let wallet_row = sqlx::query( @@ -3277,7 +4916,7 @@ WHERE id = $1 UPDATE wallets SET balance = $2, - total_refunded = total_refunded + $3, + total_refunded = $3, updated_at = NOW() WHERE id = $1 RETURNING @@ -3298,7 +4937,7 @@ RETURNING ) .bind(&input.wallet_id) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -3441,7 +5080,30 @@ RETURNING Box::pin(async move { let Some(current_refund) = sqlx::query( r#" -SELECT status +SELECT + id, + refund_no, + wallet_id, + user_id, + payment_order_id, + source_type, + source_id, + refund_mode, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, + reason, + failure_reason, + gateway_refund_id, + payout_method, + payout_reference, + payout_proof, + requested_by, + approved_by, + processed_by, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs, + CAST(EXTRACT(EPOCH FROM processed_at) AS BIGINT) AS processed_at_unix_secs, + CAST(EXTRACT(EPOCH FROM completed_at) AS BIGINT) AS completed_at_unix_secs FROM refund_requests WHERE id = $1 AND wallet_id = $2 FOR UPDATE @@ -3455,20 +5117,50 @@ FOR UPDATE else { return Ok(WalletMutationOutcome::NotFound); }; - let status: String = row_get(¤t_refund, "status")?; - if status != "processing" { + let refund = map_admin_wallet_refund_row(¤t_refund)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let (Some(existing_id), Some(incoming_id)) = ( + refund.gateway_refund_id.as_deref(), + input.gateway_refund_id.as_deref(), + ) { + if existing_id != incoming_id { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence" + .to_string(), + )); + } + } + if refund.status == "succeeded" { + return Ok(WalletMutationOutcome::Applied(refund)); + } + if refund.status != "processing" { return Ok(WalletMutationOutcome::Invalid( "refund status must be processing before completion".to_string(), )); } + // A processing proof is durable evidence of the provider's last + // response. Preserve it for ordinary replays, but replace it when the + // same gateway refund reaches an explicit successful terminal state. + let payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| { + refund.payout_proof.is_none() || wallet_refund_proof_is_success(proof) + }) + .cloned() + .or_else(|| refund.payout_proof.clone()); let refund_row = sqlx::query( r#" UPDATE refund_requests SET status = 'succeeded', - gateway_refund_id = $3, - payout_reference = $4, + gateway_refund_id = COALESCE($3, gateway_refund_id), + payout_reference = COALESCE($4, payout_reference), payout_proof = $5, completed_at = NOW(), updated_at = NOW() @@ -3503,7 +5195,7 @@ RETURNING .bind(&input.wallet_id) .bind(input.gateway_refund_id.as_deref()) .bind(input.payout_reference.as_deref()) - .bind(&input.payout_proof) + .bind(&payout_proof) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -3515,6 +5207,146 @@ RETURNING .await } + async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result, DataLayerError> { + if input.gateway_refund_id.trim().is_empty() || input.gateway_refund_id.len() > 128 { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier is invalid".to_string(), + )); + } + if input + .payout_proof + .as_ref() + .is_some_and(|proof| !proof.is_object()) + { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund proof must be an object".to_string(), + )); + } + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + let Some(current_row) = sqlx::query( + r#" +SELECT + id, + refund_no, + wallet_id, + user_id, + payment_order_id, + source_type, + source_id, + refund_mode, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, + reason, + failure_reason, + gateway_refund_id, + payout_method, + payout_reference, + payout_proof, + requested_by, + approved_by, + processed_by, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs, + CAST(EXTRACT(EPOCH FROM processed_at) AS BIGINT) AS processed_at_unix_secs, + CAST(EXTRACT(EPOCH FROM completed_at) AS BIGINT) AS completed_at_unix_secs +FROM refund_requests +WHERE id = $1 AND wallet_id = $2 +FOR UPDATE + "#, + ) + .bind(&input.refund_id) + .bind(&input.wallet_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(WalletMutationOutcome::NotFound); + }; + let current = map_admin_wallet_refund_row(¤t_row)?; + if !current.amount_usd.is_finite() || current.amount_usd <= 0.0 { + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let Some(existing_id) = current.gateway_refund_id.as_deref() { + if existing_id != input.gateway_refund_id { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence" + .to_string(), + )); + } + } + if current.status == "succeeded" { + return Ok(WalletMutationOutcome::Applied(current)); + } + if current.status != "processing" { + return Ok(WalletMutationOutcome::Invalid( + "refund status must be processing before gateway update".to_string(), + )); + } + // Do not let a retry overwrite a pending proof with arbitrary data. A + // provider's explicit success response is the one permitted upgrade. + let payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| { + current.payout_proof.is_none() || wallet_refund_proof_is_success(proof) + }) + .cloned() + .or_else(|| current.payout_proof.clone()); + let row = sqlx::query( + r#" +UPDATE refund_requests +SET gateway_refund_id = COALESCE(gateway_refund_id, $3), + payout_proof = $4, + updated_at = NOW() +WHERE id = $1 AND wallet_id = $2 AND status = 'processing' +RETURNING + id, + refund_no, + wallet_id, + user_id, + payment_order_id, + source_type, + source_id, + refund_mode, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, + reason, + failure_reason, + gateway_refund_id, + payout_method, + payout_reference, + payout_proof, + requested_by, + approved_by, + processed_by, + CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs, + CAST(EXTRACT(EPOCH FROM processed_at) AS BIGINT) AS processed_at_unix_secs, + CAST(EXTRACT(EPOCH FROM completed_at) AS BIGINT) AS completed_at_unix_secs + "#, + ) + .bind(&input.refund_id) + .bind(&input.wallet_id) + .bind(&input.gateway_refund_id) + .bind(&payout_proof) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + Ok(WalletMutationOutcome::Applied(map_admin_wallet_refund_row( + &row, + )?)) + }) + }) + .await + } + async fn fail_admin_wallet_refund( &self, input: FailAdminWalletRefundInput, @@ -3569,6 +5401,11 @@ FOR UPDATE return Ok(WalletMutationOutcome::NotFound); }; let refund = map_admin_wallet_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if matches!(refund.status.as_str(), "pending_approval" | "approved") { let refund_row = sqlx::query( @@ -3649,6 +5486,22 @@ WHERE id = $1 ))); } + // Only an explicitly offline payout can be released without + // external settlement evidence. An original-channel refund + // may still be in flight between the provider request and + // the evidence update. + if refund.gateway_refund_id.is_some() + || refund.payout_proof.is_some() + || !refund + .refund_mode + .trim() + .eq_ignore_ascii_case("offline_payout") + { + return Ok(WalletMutationOutcome::Invalid( + "cannot fail refund while gateway settlement is processing".to_string(), + )); + } + let Some(wallet_row) = sqlx::query( r#" SELECT @@ -3682,15 +5535,93 @@ FOR UPDATE let amount_usd = refund.amount_usd; let before_recharge: f64 = row_get(&wallet_row, "balance")?; let before_gift: f64 = row_get(&wallet_row, "gift_balance")?; + let before_total_refunded: f64 = row_get(&wallet_row, "total_refunded")?; let before_total = before_recharge + before_gift; let after_recharge = before_recharge + amount_usd; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded - amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || before_total_refunded < amount_usd + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + || after_total_refunded < 0.0 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid for refund recovery".to_string(), + )); + } + + let mut order_amounts = None; + if let Some(payment_order_id) = refund.payment_order_id.as_deref() { + let Some(order_row) = sqlx::query( + r#" +SELECT + wallet_id, + status, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd, + CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd +FROM payment_orders +WHERE id = $1 +FOR UPDATE + "#, + ) + .bind(payment_order_id) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(WalletMutationOutcome::Invalid( + "payment order not found".to_string(), + )); + }; + let order_wallet_id: String = row_get(&order_row, "wallet_id")?; + let order_status: String = row_get(&order_row, "status")?; + if order_wallet_id != input.wallet_id || order_status != "credited" { + return Ok(WalletMutationOutcome::Invalid( + "payment order is not refundable for this wallet".to_string(), + )); + } + let order_amount: f64 = row_get(&order_row, "amount_usd")?; + let refunded_before: f64 = row_get(&order_row, "refunded_amount_usd")?; + let refundable_before: f64 = row_get(&order_row, "refundable_amount_usd")?; + let refunded_after = refunded_before - amount_usd; + let refundable_after = refundable_before + amount_usd; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_before, + refundable_before, + ) || refunded_before < amount_usd + || !refunded_after.is_finite() + || refunded_after < 0.0 + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order_amount + { + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + order_amounts = Some(( + payment_order_id.to_string(), + refunded_after, + refundable_after, + )); + } let wallet_row = sqlx::query( r#" UPDATE wallets SET balance = $2, - total_refunded = GREATEST(total_refunded - $3, 0), + total_refunded = $3, updated_at = NOW() WHERE id = $1 RETURNING @@ -3711,7 +5642,7 @@ RETURNING ) .bind(&input.wallet_id) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -3763,7 +5694,7 @@ VALUES ( .bind(&input.wallet_id) .bind(amount_usd) .bind(before_total) - .bind(after_recharge + before_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -3774,21 +5705,29 @@ VALUES ( .await .map_postgres_err()?; - if let Some(payment_order_id) = refund.payment_order_id.as_deref() { - let _ = sqlx::query( + if let Some((payment_order_id, refunded_after, refundable_after)) = + order_amounts + { + let result = sqlx::query( r#" UPDATE payment_orders SET - refunded_amount_usd = refunded_amount_usd - $2, - refundable_amount_usd = refundable_amount_usd + $2 + refunded_amount_usd = $2, + refundable_amount_usd = $3 WHERE id = $1 "#, ) .bind(payment_order_id) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .execute(&mut **tx) .await .map_postgres_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund recovery".to_string(), + )); + } } let refund_row = sqlx::query( @@ -3798,7 +5737,7 @@ SET status = 'failed', failure_reason = $3, updated_at = NOW() -WHERE id = $1 AND wallet_id = $2 +WHERE id = $1 AND wallet_id = $2 AND status = 'processing' RETURNING id, refund_no, @@ -3828,9 +5767,14 @@ RETURNING .bind(&input.refund_id) .bind(&input.wallet_id) .bind(&input.reason) - .fetch_one(&mut **tx) + .fetch_optional(&mut **tx) .await - .map_postgres_err()?; + .map_postgres_err()? + .ok_or_else(|| { + DataLayerError::UnexpectedValue( + "refund status changed during recovery".to_string(), + ) + })?; Ok(WalletMutationOutcome::Applied(( wallet, map_admin_wallet_refund_row(&refund_row)?, @@ -3841,7 +5785,7 @@ RETURNING reason_code: "refund_revert".to_string(), amount: amount_usd, balance_before: before_total, - balance_after: after_recharge + before_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -4128,14 +6072,31 @@ FOR UPDATE } if order .expires_at_unix_secs - .is_some_and(|value| value < Utc::now().timestamp().max(0) as u64) + .is_some_and(|value| value <= Utc::now().timestamp().max(0) as u64) { return Ok(WalletMutationOutcome::Invalid( "payment order expired".to_string(), )); } - let order_kind: String = row_get(&order_row, "order_kind")?; + let order_payment_provider: Option = + row_get(&order_row, "payment_provider")?; + let order_payment_channel: Option = + row_get(&order_row, "payment_channel")?; + if validate_payment_order_credit_amounts( + &order_kind, + &order.payment_method, + order_payment_provider.as_deref(), + order_payment_channel.as_deref(), + order.amount_usd, + order.pay_amount, + ) + .is_err() + { + return Ok(WalletMutationOutcome::Invalid( + "payment order amount is invalid".to_string(), + )); + } if order_kind == "plan_purchase" { let order_user_id: Option = row_get(&order_row, "user_id")?; let Some(user_id) = order_user_id else { @@ -4156,7 +6117,7 @@ FOR UPDATE }); let entitlements = plan_entitlements_snapshot(&snapshot); let now = Utc::now(); - let expires_at = plan_expires_at(&snapshot, now); + let expires_at = plan_expires_at(&snapshot, now)?; let existing_entitlement_id = sqlx::query_scalar::<_, String>( r#" SELECT id @@ -4334,7 +6295,8 @@ SELECT id, status, CAST(balance AS DOUBLE PRECISION) AS balance, - CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance + CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(total_recharged AS DOUBLE PRECISION) AS total_recharged FROM wallets WHERE id = $1 FOR UPDATE @@ -4358,6 +6320,19 @@ FOR UPDATE let before_recharge: f64 = row_get(&wallet_row, "balance")?; let before_gift: f64 = row_get(&wallet_row, "gift_balance")?; + let total_recharged: f64 = row_get(&wallet_row, "total_recharged")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order.amount_usd).is_finite() + || !(before_recharge + before_gift + order.amount_usd).is_finite() + { + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order.amount_usd; let now_unix_secs = Utc::now().timestamp().max(0) as u64; @@ -4518,6 +6493,7 @@ RETURNING &self, input: CreateAdminRedeemCodeBatchInput, ) -> Result { + validate_admin_redeem_code_batch_input(&input).map_err(DataLayerError::InvalidInput)?; self.tx_runner .run_read_write(|tx| { Box::pin(async move { @@ -5113,9 +7089,6 @@ FOR UPDATE OF codes, batches row_get(&code_row, "batch_expires_at_unix_secs")?, "redeem_code_batches.expires_at", )?; - let credits_recharge_balance = - redeem_code_credits_recharge_balance(&balance_bucket); - if code_status == "disabled" { return Ok(RedeemWalletCodeOutcome::CodeDisabled); } @@ -5224,16 +7197,15 @@ RETURNING let before_recharge = wallet_snapshot.balance; let before_gift = wallet_snapshot.gift_balance; let before_total = before_recharge + before_gift; - let after_recharge = if credits_recharge_balance { - before_recharge + amount_usd - } else { - before_recharge - }; - let after_gift = if credits_recharge_balance { - before_gift - } else { - before_gift + amount_usd - }; + let (after_recharge, after_gift, after_total_recharged) = + validate_redeem_wallet_credit( + &balance_bucket, + amount_usd, + before_recharge, + before_gift, + wallet_snapshot.total_recharged, + ) + .map_err(DataLayerError::UnexpectedValue)?; let wallet_row = sqlx::query( r#" @@ -5241,7 +7213,7 @@ UPDATE wallets SET balance = $2, gift_balance = $3, - total_recharged = total_recharged + $4, + total_recharged = $4, updated_at = NOW() WHERE id = $1 RETURNING @@ -5263,7 +7235,7 @@ RETURNING .bind(&wallet_snapshot.id) .bind(after_recharge) .bind(after_gift) - .bind(amount_usd) + .bind(after_total_recharged) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -5453,16 +7425,6 @@ fn as_i64(value: usize, field: &str) -> Result { .map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {field}: {value}"))) } -fn default_refund_mode_for_payment_method(payment_method: &str) -> &'static str { - if matches!( - payment_method, - "admin_manual" | "card_recharge" | "card_code" | "gift_code" - ) { - return "offline_payout"; - } - "original_channel" -} - fn payment_gateway_response_map( value: Option, ) -> serde_json::Map { @@ -5550,34 +7512,14 @@ fn plan_purchase_limit_scope(snapshot: &serde_json::Value) -> &str { } } -fn plan_replacement_entitlement_types(snapshot: &serde_json::Value) -> Vec<&'static str> { - let entitlements = plan_entitlements_snapshot(snapshot); - let mut kinds = Vec::new(); - if entitlement_snapshot_has_type(&entitlements, "daily_quota") { - kinds.push("daily_quota"); - } - if entitlement_snapshot_has_type(&entitlements, "membership_group") { - kinds.push("membership_group"); - } - kinds -} - -fn entitlement_snapshot_has_type(snapshot: &serde_json::Value, entitlement_type: &str) -> bool { - snapshot.as_array().is_some_and(|items| { - items - .iter() - .any(|item| item.get("type").and_then(|value| value.as_str()) == Some(entitlement_type)) - }) -} - async fn replace_matching_plan_entitlements_postgres( tx: &mut crate::PostgresTransaction, user_id: &str, snapshot: &serde_json::Value, now: chrono::DateTime, ) -> Result<(), DataLayerError> { - let replacement_types = plan_replacement_entitlement_types(snapshot); - if replacement_types.is_empty() { + let incoming_entitlements = plan_entitlements_snapshot(snapshot); + if !entitlements_have_replacement_selector(&incoming_entitlements) { return Ok(()); } @@ -5598,9 +7540,8 @@ WHERE user_id = $1 for row in rows { let entitlements: serde_json::Value = row_get(&row, "entitlements_snapshot")?; - let should_replace = replacement_types - .iter() - .any(|kind| entitlement_snapshot_has_type(&entitlements, kind)); + let should_replace = + entitlements_should_replace_existing(&incoming_entitlements, &entitlements); if !should_replace { continue; } @@ -5628,22 +7569,15 @@ WHERE id = $2 fn plan_expires_at( snapshot: &serde_json::Value, starts_at: chrono::DateTime, -) -> chrono::DateTime { - let duration_value = snapshot - .get("duration_value") - .and_then(|value| value.as_i64()) - .unwrap_or(1) - .max(1); - match snapshot - .get("duration_unit") - .and_then(|value| value.as_str()) - .unwrap_or("month") - { - "day" => starts_at + chrono::Duration::days(duration_value), - "year" => starts_at + chrono::Duration::days(365 * duration_value), - "custom" => starts_at + chrono::Duration::days(duration_value), - _ => starts_at + chrono::Duration::days(30 * duration_value), - } +) -> Result, DataLayerError> { + let days = + checked_plan_duration_days_from_snapshot(snapshot).map_err(DataLayerError::InvalidInput)?; + let duration = chrono::TimeDelta::try_days(days).ok_or_else(|| { + DataLayerError::InvalidInput("plan duration exceeds the supported range".to_string()) + })?; + starts_at.checked_add_signed(duration).ok_or_else(|| { + DataLayerError::InvalidInput("plan expiration exceeds the supported range".to_string()) + }) } async fn apply_plan_wallet_credit_postgres( @@ -5653,11 +7587,16 @@ async fn apply_plan_wallet_credit_postgres( payment_method: &str, entitlements: &serde_json::Value, ) -> Result<(), DataLayerError> { + validate_plan_wallet_credit_entitlements(entitlements).map_err(DataLayerError::InvalidInput)?; let credits = entitlements .as_array() .into_iter() .flatten() - .filter(|item| item.get("type").and_then(|value| value.as_str()) == Some("wallet_credit")) + .filter(|item| { + item.get("type") + .and_then(|value| value.as_str()) + .is_some_and(|value| value.eq_ignore_ascii_case("wallet_credit")) + }) .filter_map(|item| { let amount = item.get("amount_usd").and_then(|value| value.as_f64())?; if amount <= 0.0 || !amount.is_finite() { @@ -5667,6 +7606,7 @@ async fn apply_plan_wallet_credit_postgres( .get("balance_bucket") .and_then(|value| value.as_str()) .unwrap_or("gift") + .trim() .to_ascii_lowercase(); Some((amount, bucket)) }) @@ -5681,7 +7621,8 @@ SELECT id, status, CAST(balance AS DOUBLE PRECISION) AS balance, - CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance + CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(total_recharged AS DOUBLE PRECISION) AS total_recharged FROM wallets WHERE id = $1 LIMIT 1 @@ -5705,6 +7646,17 @@ FOR UPDATE } let mut recharge_balance: f64 = row_get(&wallet_row, "balance")?; let mut gift_balance: f64 = row_get(&wallet_row, "gift_balance")?; + let mut total_recharged: f64 = row_get(&wallet_row, "total_recharged")?; + if !recharge_balance.is_finite() + || !gift_balance.is_finite() + || gift_balance < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid for plan wallet_credit".to_string(), + )); + } for (amount, bucket) in credits { let before_recharge = recharge_balance; let before_gift = gift_balance; @@ -5712,16 +7664,27 @@ FOR UPDATE let credits_recharge = bucket == "recharge"; if credits_recharge { recharge_balance += amount; + total_recharged += amount; } else { gift_balance += amount; } let after_total = recharge_balance + gift_balance; + if !before_total.is_finite() + || !recharge_balance.is_finite() + || !gift_balance.is_finite() + || !total_recharged.is_finite() + || !after_total.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow for plan wallet_credit".to_string(), + )); + } sqlx::query( r#" UPDATE wallets SET balance = $2, gift_balance = $3, - total_recharged = total_recharged + $4, + total_recharged = $4, updated_at = NOW() WHERE id = $1 "#, @@ -5729,7 +7692,7 @@ WHERE id = $1 .bind(wallet_id) .bind(recharge_balance) .bind(gift_balance) - .bind(if credits_recharge { amount } else { 0.0 }) + .bind(total_recharged) .execute(&mut **tx) .await .map_postgres_err()?; @@ -5777,10 +7740,8 @@ SET signature_valid = $2, status = 'failed', error_message = $3, payload_hash = $4, - payload = $5, - processed_at = NOW(), - order_no = COALESCE($6, order_no), - gateway_order_id = COALESCE($7, gateway_order_id) + payload = NULL, + processed_at = NOW() WHERE id = $1 "#, ) @@ -5788,15 +7749,61 @@ WHERE id = $1 .bind(input.signature_valid) .bind(error) .bind(&input.payload_hash) - .bind(&input.payload) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) .execute(&mut **tx) .await .map_postgres_err()?; Ok(()) } +async fn postgres_bind_payment_gateway_order_id( + tx: &mut crate::PostgresTransaction, + order_id: &str, + gateway_order_id: &str, +) -> Result { + sqlx::query("SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_postgres_err()?; + let bind_result = sqlx::query("UPDATE payment_orders SET gateway_order_id = $2 WHERE id = $1") + .bind(order_id) + .bind(gateway_order_id) + .execute(&mut **tx) + .await; + match bind_result { + Ok(_) => { + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(true) + } + Err(error) + if error + .as_database_error() + .is_some_and(|database_error| database_error.is_unique_violation()) => + { + sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_postgres_err()?; + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(false) + } + Err(error) => { + let _ = sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + let _ = sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + Err(postgres_error(error)) + } + } +} + async fn mark_payment_callback_processed( tx: &mut crate::PostgresTransaction, callback_id: &str, @@ -5812,17 +7819,16 @@ SET payment_order_id = $2, status = 'processed', error_message = NULL, payload_hash = $3, - payload = $4, + payload = NULL, processed_at = NOW(), - order_no = $5, - gateway_order_id = COALESCE($6, gateway_order_id) + order_no = $4, + gateway_order_id = COALESCE($5, gateway_order_id) WHERE id = $1 "#, ) .bind(callback_id) .bind(order_id) .bind(&input.payload_hash) - .bind(&input.payload) .bind(order_no) .bind(input.gateway_order_id.as_deref()) .execute(&mut **tx) @@ -6164,6 +8170,30 @@ fn map_admin_redeem_code_row(row: &PgRow) -> Result Result { + let existing_wallet_id: String = row_get(row, "wallet_id")?; + let pay_currency: Option = row_get(row, "pay_currency")?; + let payment_method: String = row_get(row, "payment_method")?; + let payment_provider: Option = row_get(row, "payment_provider")?; + let payment_channel: Option = row_get(row, "payment_channel")?; + Ok(wallet_recharge_replay_matches( + &existing_wallet_id, + row_get(row, "amount_usd")?, + row_get(row, "pay_amount")?, + pay_currency.as_deref(), + row_get(row, "exchange_rate")?, + &payment_method, + payment_provider.as_deref(), + payment_channel.as_deref(), + wallet_id, + input, + )) +} + fn map_admin_payment_order_row(row: &PgRow) -> Result { Ok(StoredAdminPaymentOrder { id: row_get(row, "id")?, @@ -6177,6 +8207,8 @@ fn map_admin_payment_order_row(row: &PgRow) -> Result, initial_gift_usd: f64, unlimited: bool, -) -> Result, DataLayerError> { +) -> Result, DataLayerError> { + let owner = user_id + .or(api_key_id) + .filter(|value| !value.trim().is_empty()); + if owner.is_none() || (user_id.is_some() && api_key_id.is_some()) { + return Err(DataLayerError::InvalidInput( + "wallet owner must be exactly one non-empty user or API-key id".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "initial gift amount must be finite".to_string(), + )); + } let gift_amount = if unlimited { 0.0 } else { initial_gift_usd.max(0.0) }; + let owner_column = if user_id.is_some() { + "user_id" + } else { + "api_key_id" + }; + let owner_value = owner.expect("validated wallet owner"); let mut tx = pool.begin().await.map_postgres_err()?; + + // Keep the ownership lock order identical to the guarded user-deletion + // path: users first, then api_keys/wallets. This closes the window where + // a deletion can observe no wallet while a concurrent initializer inserts + // one after the user row has been removed. + if let Some(user_id) = user_id { + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + } else { + let api_key_user_id: Option = + sqlx::query_scalar("SELECT user_id FROM api_keys WHERE id = $1") + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(api_key_user_id) = api_key_user_id else { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + }; + let user_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE") + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if user_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + let api_key_exists: Option = + sqlx::query_scalar("SELECT id FROM api_keys WHERE id = $1 AND user_id = $2 FOR UPDATE") + .bind(owner_value) + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + if api_key_exists.is_none() { + tx.rollback().await.map_postgres_err()?; + return Ok(None); + } + } + let existing_sql = format!( + "SELECT id, user_id, api_key_id, CAST(balance AS DOUBLE PRECISION) AS balance, CAST(gift_balance AS DOUBLE PRECISION) AS gift_balance, limit_mode, currency, status, 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(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs FROM wallets WHERE {owner_column} = $1 LIMIT 1 FOR UPDATE" + ); + if let Some(row) = sqlx::query(&existing_sql) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + { + let existing = map_wallet_row(&row)?; + tx.commit().await.map_err(postgres_error)?; + return Ok(Some((existing, false))); + } + let wallet_row = sqlx::query( r#" INSERT INTO wallets ( @@ -6297,6 +8411,7 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES ($1, $2, $3, 0, $4, $5, 'USD', 'active', 0, 0, 0, $6, NOW(), NOW()) +ON CONFLICT DO NOTHING RETURNING id, user_id, @@ -6319,9 +8434,28 @@ RETURNING .bind(gift_amount) .bind(if unlimited { "unlimited" } else { "finite" }) .bind(gift_amount) - .fetch_one(&mut *tx) + .fetch_optional(&mut *tx) .await .map_postgres_err()?; + let Some(wallet_row) = wallet_row else { + // Another initializer may have committed the owner row while this + // transaction waited on the unique index. Read it under lock and + // return it without creating a duplicate gift transaction. + let Some(row) = sqlx::query(&existing_sql) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()? + else { + tx.rollback().await.map_err(postgres_error)?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + let existing = map_wallet_row(&row)?; + tx.commit().await.map_err(postgres_error)?; + return Ok(Some((existing, false))); + }; let wallet = map_wallet_row(&wallet_row)?; if gift_amount > 0.0 { let link_id = user_id.or(api_key_id).unwrap_or_default(); @@ -6351,7 +8485,7 @@ VALUES ($1, $2, 'gift', 'gift_initial', $3, 0, $3, 0, 0, 0, $3, 'system_task', $ .map_postgres_err()?; } tx.commit().await.map_err(postgres_error)?; - Ok(Some(wallet)) + Ok(Some((wallet, true))) } #[cfg(test)] diff --git a/crates/aether-data/adapters/sqlite/migrations/20260814000000_add_usage_cost_reservations.sql b/crates/aether-data/adapters/sqlite/migrations/20260814000000_add_usage_cost_reservations.sql new file mode 100644 index 000000000..9253bdeac --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260814000000_add_usage_cost_reservations.sql @@ -0,0 +1,34 @@ +CREATE TABLE IF NOT EXISTS usage_cost_reservations ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + reservation_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + reserved_cost_units INTEGER NOT NULL CHECK (reserved_cost_units >= 0), + actual_cost_units INTEGER CHECK (actual_cost_units IS NULL OR actual_cost_units >= 0), + state TEXT NOT NULL CHECK (state IN ('reserved', 'finalized', 'released')), + reservation_expires_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + finalized_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + CHECK (reservation_expires_at > admitted_at), + CHECK (retain_until >= reservation_expires_at), + CHECK ( + (state = 'reserved' AND actual_cost_units IS NULL AND finalized_at IS NULL) + OR (state = 'finalized' AND actual_cost_units IS NOT NULL AND finalized_at IS NOT NULL) + OR (state = 'released' AND actual_cost_units IS NOT NULL + AND actual_cost_units = 0 AND finalized_at IS NOT NULL) + ) +); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_request_id_idx + ON usage_cost_reservations (request_id); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_subject_admitted_at_idx + ON usage_cost_reservations (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_reservation_expires_at_idx + ON usage_cost_reservations (reservation_expires_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_retain_until_token_idx + ON usage_cost_reservations (retain_until, reservation_token); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260815000000_add_usage_request_admissions.sql b/crates/aether-data/adapters/sqlite/migrations/20260815000000_add_usage_request_admissions.sql new file mode 100644 index 000000000..32f156149 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260815000000_add_usage_request_admissions.sql @@ -0,0 +1,21 @@ +CREATE TABLE IF NOT EXISTS usage_request_admissions ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + event_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + state TEXT NOT NULL CHECK (state IN ('active', 'released')), + released_at INTEGER, + created_at INTEGER NOT NULL, + CHECK (retain_until > admitted_at), + CHECK ( + (state = 'active' AND released_at IS NULL) + OR (state = 'released' AND released_at IS NOT NULL AND released_at >= admitted_at) + ) +); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_subject_admitted_at_idx + ON usage_request_admissions (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_retain_until_token_idx + ON usage_request_admissions (retain_until, event_token); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql b/crates/aether-data/adapters/sqlite/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql new file mode 100644 index 000000000..14066f283 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260816000000_add_usage_policy_user_foreign_keys.sql @@ -0,0 +1,117 @@ +-- SQLite cannot add a foreign key with ALTER TABLE. Rebuild both ledgers while +-- preserving valid rows and dropping only records whose owning user was +-- already deleted before the relationship became enforceable. +CREATE TABLE usage_cost_reservations_with_user_fk ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + reservation_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + reserved_cost_units INTEGER NOT NULL CHECK (reserved_cost_units >= 0), + actual_cost_units INTEGER CHECK (actual_cost_units IS NULL OR actual_cost_units >= 0), + state TEXT NOT NULL CHECK (state IN ('reserved', 'finalized', 'released')), + reservation_expires_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + finalized_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + CONSTRAINT usage_cost_reservations_subject_id_fkey + FOREIGN KEY (subject_id) REFERENCES users (id) ON DELETE CASCADE, + CHECK (reservation_expires_at > admitted_at), + CHECK (retain_until >= reservation_expires_at), + CHECK ( + (state = 'reserved' AND actual_cost_units IS NULL AND finalized_at IS NULL) + OR (state = 'finalized' AND actual_cost_units IS NOT NULL AND finalized_at IS NOT NULL) + OR (state = 'released' AND actual_cost_units IS NOT NULL + AND actual_cost_units = 0 AND finalized_at IS NOT NULL) + ) +); + +INSERT INTO usage_cost_reservations_with_user_fk ( + request_id, + subject_id, + reservation_token, + admitted_at, + reserved_cost_units, + actual_cost_units, + state, + reservation_expires_at, + retain_until, + finalized_at, + created_at, + updated_at +) +SELECT + reservation.request_id, + reservation.subject_id, + reservation.reservation_token, + reservation.admitted_at, + reservation.reserved_cost_units, + reservation.actual_cost_units, + reservation.state, + reservation.reservation_expires_at, + reservation.retain_until, + reservation.finalized_at, + reservation.created_at, + reservation.updated_at +FROM usage_cost_reservations AS reservation +INNER JOIN users AS app_user ON app_user.id = reservation.subject_id; + +DROP TABLE usage_cost_reservations; +ALTER TABLE usage_cost_reservations_with_user_fk RENAME TO usage_cost_reservations; + +CREATE INDEX usage_cost_reservations_request_id_idx + ON usage_cost_reservations (request_id); +CREATE INDEX usage_cost_reservations_subject_admitted_at_idx + ON usage_cost_reservations (subject_id, admitted_at); +CREATE INDEX usage_cost_reservations_reservation_expires_at_idx + ON usage_cost_reservations (reservation_expires_at); +CREATE INDEX usage_cost_reservations_retain_until_token_idx + ON usage_cost_reservations (retain_until, reservation_token); + +CREATE TABLE usage_request_admissions_with_user_fk ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + event_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + state TEXT NOT NULL CHECK (state IN ('active', 'released')), + released_at INTEGER, + created_at INTEGER NOT NULL, + CONSTRAINT usage_request_admissions_subject_id_fkey + FOREIGN KEY (subject_id) REFERENCES users (id) ON DELETE CASCADE, + CHECK (retain_until > admitted_at), + CHECK ( + (state = 'active' AND released_at IS NULL) + OR (state = 'released' AND released_at IS NOT NULL AND released_at >= admitted_at) + ) +); + +INSERT INTO usage_request_admissions_with_user_fk ( + request_id, + subject_id, + event_token, + admitted_at, + retain_until, + state, + released_at, + created_at +) +SELECT + admission.request_id, + admission.subject_id, + admission.event_token, + admission.admitted_at, + admission.retain_until, + admission.state, + admission.released_at, + admission.created_at +FROM usage_request_admissions AS admission +INNER JOIN users AS app_user ON app_user.id = admission.subject_id; + +DROP TABLE usage_request_admissions; +ALTER TABLE usage_request_admissions_with_user_fk RENAME TO usage_request_admissions; + +CREATE INDEX usage_request_admissions_subject_admitted_at_idx + ON usage_request_admissions (subject_id, admitted_at); +CREATE INDEX usage_request_admissions_retain_until_token_idx + ON usage_request_admissions (retain_until, event_token); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql b/crates/aether-data/adapters/sqlite/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql new file mode 100644 index 000000000..4c6b8cdba --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql @@ -0,0 +1,19 @@ +-- A gateway transaction identifier may repeat across payment methods, but +-- must never identify two orders in the same method. If historical conflicts +-- exist, index creation intentionally fails without modifying financial data. +-- Diagnose with: +-- SELECT payment_method, gateway_order_id, COUNT(*) +-- FROM payment_orders +-- WHERE gateway_order_id IS NOT NULL +-- GROUP BY payment_method, gateway_order_id +-- HAVING COUNT(*) > 1; +UPDATE payment_orders +SET payment_method = lower(trim(payment_method)) +WHERE payment_method <> lower(trim(payment_method)); + +UPDATE payment_callbacks +SET payment_method = lower(trim(payment_method)) +WHERE payment_method <> lower(trim(payment_method)); + +CREATE UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id + ON payment_orders (payment_method, gateway_order_id); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260821130000_add_user_security_version.sql b/crates/aether-data/adapters/sqlite/migrations/20260821130000_add_user_security_version.sql new file mode 100644 index 000000000..5a97e0460 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260821130000_add_user_security_version.sql @@ -0,0 +1,5 @@ +ALTER TABLE users + ADD COLUMN security_version INTEGER NOT NULL DEFAULT 0; + +ALTER TABLE user_sessions + ADD COLUMN security_version INTEGER NOT NULL DEFAULT 0; diff --git a/crates/aether-data/adapters/sqlite/migrations/20260827050000_anonymize_deleted_user_history.sql b/crates/aether-data/adapters/sqlite/migrations/20260827050000_anonymize_deleted_user_history.sql new file mode 100644 index 000000000..05d81afcc --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260827050000_anonymize_deleted_user_history.sql @@ -0,0 +1,157 @@ +-- Retain legacy row contents. This migration only removes user foreign keys +-- required by the current account-deletion flow. +ALTER TABLE user_plan_entitlements RENAME TO _aether_user_plan_entitlements_with_user_fk; +DROP INDEX IF EXISTS idx_user_plan_entitlements_user_active; +DROP INDEX IF EXISTS idx_user_plan_entitlements_order; + +CREATE TABLE user_plan_entitlements ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + plan_id TEXT NOT NULL, + payment_order_id TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'active', + starts_at INTEGER NOT NULL, + expires_at INTEGER NOT NULL, + entitlements_snapshot TEXT NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + FOREIGN KEY(plan_id) REFERENCES billing_plans(id) ON DELETE RESTRICT, + FOREIGN KEY(payment_order_id) REFERENCES payment_orders(id) ON DELETE RESTRICT +); + +INSERT INTO user_plan_entitlements ( + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +) +SELECT + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +FROM _aether_user_plan_entitlements_with_user_fk; + +CREATE INDEX idx_user_plan_entitlements_user_active + ON user_plan_entitlements (user_id, status, expires_at); +CREATE INDEX idx_user_plan_entitlements_order + ON user_plan_entitlements (payment_order_id); + +ALTER TABLE entitlement_usage_ledgers RENAME TO _aether_entitlement_usage_ledgers_with_user_fk; +DROP INDEX IF EXISTS idx_entitlement_usage_user_date; +DROP INDEX IF EXISTS idx_entitlement_usage_entitlement_date; + +CREATE TABLE entitlement_usage_ledgers ( + id TEXT PRIMARY KEY, + user_entitlement_id TEXT NOT NULL, + user_id TEXT NOT NULL, + request_id TEXT NOT NULL, + amount_usd REAL NOT NULL, + balance_before REAL NOT NULL, + balance_after REAL NOT NULL, + usage_date TEXT NOT NULL, + created_at INTEGER NOT NULL, + UNIQUE (user_entitlement_id, request_id), + FOREIGN KEY(user_entitlement_id) REFERENCES user_plan_entitlements(id) ON DELETE CASCADE +); + +INSERT INTO entitlement_usage_ledgers ( + id, user_entitlement_id, user_id, request_id, amount_usd, + balance_before, balance_after, usage_date, created_at +) +SELECT + id, user_entitlement_id, user_id, request_id, amount_usd, + balance_before, balance_after, usage_date, created_at +FROM _aether_entitlement_usage_ledgers_with_user_fk; + +CREATE INDEX idx_entitlement_usage_user_date + ON entitlement_usage_ledgers (user_id, usage_date); +CREATE INDEX idx_entitlement_usage_entitlement_date + ON entitlement_usage_ledgers (user_entitlement_id, usage_date); + +DROP TABLE _aether_entitlement_usage_ledgers_with_user_fk; +DROP TABLE _aether_user_plan_entitlements_with_user_fk; + +ALTER TABLE user_referrals RENAME TO _aether_user_referrals_with_user_fk; +DROP INDEX IF EXISTS idx_user_referrals_inviter; +DROP INDEX IF EXISTS idx_user_referrals_created; +DROP INDEX IF EXISTS idx_user_referrals_invite_code; + +CREATE TABLE user_referrals ( + id TEXT PRIMARY KEY, + inviter_user_id TEXT NOT NULL, + invitee_user_id TEXT NOT NULL UNIQUE, + invite_code_snapshot TEXT NOT NULL, + source_json TEXT, + first_paid_order_id TEXT, + first_paid_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + FOREIGN KEY(first_paid_order_id) REFERENCES payment_orders(id) ON DELETE SET NULL +); + +INSERT INTO user_referrals ( + id, inviter_user_id, invitee_user_id, invite_code_snapshot, source_json, + first_paid_order_id, first_paid_at, created_at, updated_at +) +SELECT + id, inviter_user_id, invitee_user_id, invite_code_snapshot, source_json, + first_paid_order_id, first_paid_at, created_at, updated_at +FROM _aether_user_referrals_with_user_fk; + +CREATE INDEX idx_user_referrals_inviter + ON user_referrals (inviter_user_id, created_at); +CREATE INDEX idx_user_referrals_created + ON user_referrals (created_at); +CREATE INDEX idx_user_referrals_invite_code + ON user_referrals (invite_code_snapshot); + +ALTER TABLE referral_rewards RENAME TO _aether_referral_rewards_with_user_fk; +DROP INDEX IF EXISTS idx_referral_rewards_inviter_status; +DROP INDEX IF EXISTS idx_referral_rewards_inviter_created; +DROP INDEX IF EXISTS idx_referral_rewards_created; +DROP INDEX IF EXISTS idx_referral_rewards_source_order; + +CREATE TABLE referral_rewards ( + id TEXT PRIMARY KEY, + referral_id TEXT NOT NULL, + inviter_user_id TEXT NOT NULL, + invitee_user_id TEXT NOT NULL, + reward_type TEXT NOT NULL, + trigger_point TEXT NOT NULL, + source_order_id TEXT, + idempotency_key TEXT NOT NULL UNIQUE, + amount_usd REAL NOT NULL, + status TEXT NOT NULL DEFAULT 'pending', + wallet_transaction_id TEXT, + reversed_amount_usd REAL NOT NULL DEFAULT 0, + pending_reversal_amount_usd REAL NOT NULL DEFAULT 0, + failure_reason TEXT, + admin_operator_id TEXT, + admin_note TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + FOREIGN KEY(referral_id) REFERENCES user_referrals(id) ON DELETE CASCADE, + FOREIGN KEY(source_order_id) REFERENCES payment_orders(id) ON DELETE SET NULL +); + +INSERT INTO referral_rewards ( + id, referral_id, inviter_user_id, invitee_user_id, reward_type, + trigger_point, source_order_id, idempotency_key, amount_usd, status, + wallet_transaction_id, reversed_amount_usd, pending_reversal_amount_usd, + failure_reason, admin_operator_id, admin_note, created_at, updated_at +) +SELECT + id, referral_id, inviter_user_id, invitee_user_id, reward_type, + trigger_point, source_order_id, idempotency_key, amount_usd, status, + wallet_transaction_id, reversed_amount_usd, pending_reversal_amount_usd, + failure_reason, admin_operator_id, admin_note, created_at, updated_at +FROM _aether_referral_rewards_with_user_fk; + +CREATE INDEX idx_referral_rewards_inviter_status + ON referral_rewards (inviter_user_id, status, created_at); +CREATE INDEX idx_referral_rewards_inviter_created + ON referral_rewards (inviter_user_id, created_at); +CREATE INDEX idx_referral_rewards_created + ON referral_rewards (created_at); +CREATE INDEX idx_referral_rewards_source_order + ON referral_rewards (source_order_id); + +DROP TABLE _aether_referral_rewards_with_user_fk; +DROP TABLE _aether_user_referrals_with_user_fk; diff --git a/crates/aether-data/adapters/sqlite/migrations/20260831000000_enforce_ldap_config_singleton.sql b/crates/aether-data/adapters/sqlite/migrations/20260831000000_enforce_ldap_config_singleton.sql new file mode 100644 index 000000000..a2333a451 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260831000000_enforce_ldap_config_singleton.sql @@ -0,0 +1,11 @@ +-- LDAP configuration is a database-wide singleton. Preserve the row selected by the legacy +-- reader (the smallest id), remove historical duplicates, and let the database arbitrate +-- concurrent first creation. +DELETE FROM ldap_configs +WHERE id <> (SELECT MIN(id) FROM ldap_configs); + +ALTER TABLE ldap_configs +ADD COLUMN singleton_key INTEGER NOT NULL DEFAULT 1 CHECK (singleton_key = 1); + +CREATE UNIQUE INDEX ldap_configs_singleton_key_key +ON ldap_configs (singleton_key); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260831010000_add_proxy_node_tunnel_generation.sql b/crates/aether-data/adapters/sqlite/migrations/20260831010000_add_proxy_node_tunnel_generation.sql new file mode 100644 index 000000000..78f5333a9 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260831010000_add_proxy_node_tunnel_generation.sql @@ -0,0 +1,20 @@ +ALTER TABLE proxy_nodes + ADD COLUMN tunnel_generation TEXT NOT NULL DEFAULT ''; + +UPDATE proxy_nodes +SET tunnel_generation = lower(hex(randomblob(16))) +WHERE tunnel_generation = ''; + +-- SQLite only permits a constant default when adding a NOT NULL column. Keep +-- the upgrade compatible with legacy rows, then replace the temporary empty +-- default for future legacy writers with a per-row random generation. This +-- also covers importers that omit the newly added column. +CREATE TRIGGER IF NOT EXISTS proxy_nodes_fill_tunnel_generation +AFTER INSERT ON proxy_nodes +WHEN NEW.tunnel_generation IS NULL OR trim(NEW.tunnel_generation) = '' +BEGIN + UPDATE proxy_nodes + SET tunnel_generation = lower(hex(randomblob(16))) + WHERE id = NEW.id + AND (tunnel_generation IS NULL OR trim(tunnel_generation) = ''); +END; diff --git a/crates/aether-data/adapters/sqlite/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql b/crates/aether-data/adapters/sqlite/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql new file mode 100644 index 000000000..c1475a9f9 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260831020000_enforce_proxy_node_endpoint_uniqueness.sql @@ -0,0 +1,9 @@ +-- A proxy endpoint has one stable node identity across manual and tunnel +-- registrations. Index creation intentionally fails if historical duplicates +-- exist; operators must resolve the conflicting identities explicitly. +-- Diagnose with: +-- SELECT ip, port, COUNT(*) +-- FROM proxy_nodes +-- GROUP BY ip, port +-- HAVING COUNT(*) > 1; +CREATE UNIQUE INDEX uq_proxy_node_ip_port ON proxy_nodes (ip, port); diff --git a/crates/aether-data/adapters/sqlite/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql b/crates/aether-data/adapters/sqlite/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql new file mode 100644 index 000000000..1db6a846d --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260831030000_add_usage_counter_delta_tunnel_generation.sql @@ -0,0 +1,2 @@ +ALTER TABLE usage_counter_deltas + ADD COLUMN target_tunnel_generation TEXT; diff --git a/crates/aether-data/adapters/sqlite/migrations/20260903000000_add_routing_group_sort_order.sql b/crates/aether-data/adapters/sqlite/migrations/20260903000000_add_routing_group_sort_order.sql new file mode 100644 index 000000000..7cfdc1c1e --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260903000000_add_routing_group_sort_order.sql @@ -0,0 +1,4 @@ +ALTER TABLE routing_groups ADD COLUMN sort_order INTEGER NOT NULL DEFAULT 0; + +CREATE INDEX IF NOT EXISTS routing_groups_enabled_sort_idx + ON routing_groups (enabled, sort_order, name, id); diff --git a/crates/aether-data/adapters/sqlite/src/auth.rs b/crates/aether-data/adapters/sqlite/src/auth.rs index be4c12471..c848067b7 100644 --- a/crates/aether-data/adapters/sqlite/src/auth.rs +++ b/crates/aether-data/adapters/sqlite/src/auth.rs @@ -3,9 +3,9 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::auth::{ AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository, - AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, - StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, - UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, + AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, + CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, + StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; use aether_data_contracts::DataLayerError; @@ -69,6 +69,21 @@ SELECT FROM api_keys "#; +const SQLITE_ANONYMIZE_API_KEY_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE video_tasks SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE usage SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id = ?", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id = ?", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)", +]; + +const SQLITE_DELETE_API_KEY_DEPENDENTS_SQL: &[&str] = + &["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"]; + #[derive(Debug, Clone)] pub struct SqliteAuthApiKeyReadRepository { pool: SqlitePool, @@ -111,6 +126,21 @@ impl SqliteAuthApiKeyReadRepository { record: CreateApiKeyInsertRecord, ) -> Result, DataLayerError> { let now = current_unix_secs(); + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let owner_exists: Option = + sqlx::query_scalar("SELECT id FROM users WHERE id = ? AND is_deleted = 0") + .bind(&record.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if owner_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } sqlx::query( r#" INSERT INTO api_keys ( @@ -120,7 +150,7 @@ INSERT INTO api_keys ( total_requests, total_tokens, total_cost_usd, is_standalone, created_at, updated_at ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&record.api_key_id) @@ -150,6 +180,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) &record.force_capabilities, "api_keys.force_capabilities", )?) + .bind(optional_json_to_string( + &record.feature_settings, + "api_keys.feature_settings", + )?) .bind(record.is_active) .bind(optional_i64_from_u64( record.expires_at_unix_secs, @@ -165,11 +199,25 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(record.is_standalone) .bind(now as i64) .bind(now as i64) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - - self.reload_export_by_id(&record.api_key_id).await + let reload_sql = format!("{EXPORT_COLUMNS}\nWHERE api_keys.id = ?\nLIMIT 1"); + let row = sqlx::query(&reload_sql) + .bind(&record.api_key_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::UnexpectedValue(format!( + "created api_keys row is missing: {}", + record.api_key_id + ))); + }; + let created = map_auth_api_key_export_row(&row)?; + tx.commit().await.map_sql_err()?; + Ok(Some(created)) } } @@ -186,6 +234,7 @@ struct CreateApiKeyInsertRecord { rate_limit: Option, concurrent_limit: Option, force_capabilities: Option, + feature_settings: Option, is_active: bool, expires_at_unix_secs: Option, auto_delete_on_expiry: bool, @@ -337,6 +386,7 @@ impl AuthApiKeyReadRepository for SqliteAuthApiKeyReadRepository { if user_ids.is_empty() { return Ok(AuthApiKeyExportSummary::default()); } + let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let mut builder = QueryBuilder::::new( r#" @@ -345,7 +395,7 @@ SELECT SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >= "#, ); - builder.push_bind(now_unix_secs as i64); + builder.push_bind(now_unix_secs); builder.push( r#") THEN 1 ELSE 0 END) AS active FROM api_keys @@ -429,6 +479,7 @@ WHERE id = ? rate_limit: Some(record.rate_limit), concurrent_limit: record.concurrent_limit, force_capabilities: record.force_capabilities, + feature_settings: record.feature_settings, is_active: record.is_active, expires_at_unix_secs: record.expires_at_unix_secs, auto_delete_on_expiry: record.auto_delete_on_expiry, @@ -457,6 +508,7 @@ WHERE id = ? rate_limit: record.rate_limit, concurrent_limit: record.concurrent_limit, force_capabilities: record.force_capabilities, + feature_settings: None, is_active: record.is_active, expires_at_unix_secs: record.expires_at_unix_secs, auto_delete_on_expiry: record.auto_delete_on_expiry, @@ -472,35 +524,41 @@ WHERE id = ? &self, record: UpdateUserApiKeyBasicRecord, ) -> Result, DataLayerError> { - let now = current_unix_secs() as i64; - sqlx::query( + self.update_user_api_key_basic_scoped(record, false).await + } + + async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + let result = sqlx::query( r#" UPDATE api_keys -SET name = COALESCE(?, name), - rate_limit = COALESCE(?, rate_limit), - concurrent_limit = COALESCE(?, concurrent_limit), - ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END, - updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 +SET key_encrypted = ? +WHERE CAST(id AS BLOB) = CAST(? AS BLOB) + AND CAST(user_id AS BLOB) = CAST(? AS BLOB) + AND CAST(key_hash AS BLOB) = CAST(? AS BLOB) + AND is_standalone = ? + AND CAST(key_encrypted AS BLOB) = CAST(? AS BLOB) "#, ) - .bind(record.name.as_deref()) - .bind(record.rate_limit) - .bind(record.concurrent_limit) - .bind(record.ip_rules.is_some()) - .bind(json_string_from_nested_string_list( - &record.ip_rules, - "api_keys.ip_rules", - )?) - .bind(now) - .bind(&record.api_key_id) - .bind(&record.user_id) + .bind(&mutation.key_encrypted) + .bind(&mutation.api_key_id) + .bind(&mutation.user_id) + .bind(&mutation.key_hash) + .bind(mutation.is_standalone) + .bind(&mutation.expected_key_encrypted) .execute(&self.pool) .await .map_sql_err()?; - self.reload_export_by_id(&record.api_key_id).await + Ok(result.rows_affected() == 1) + } + + async fn update_user_api_key_basic_if_unlocked( + &self, + record: UpdateUserApiKeyBasicRecord, + ) -> Result, DataLayerError> { + self.update_user_api_key_basic_scoped(record, true).await } async fn update_standalone_api_key_basic( @@ -511,7 +569,9 @@ WHERE id = ? sqlx::query( r#" UPDATE api_keys -SET name = COALESCE(?, name), +SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END, + name = CASE WHEN ? THEN ? ELSE name END, + force_capabilities = CASE WHEN ? THEN ? ELSE force_capabilities END, rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END, allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END, @@ -525,7 +585,15 @@ WHERE id = ? AND is_standalone = 1 "#, ) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted.as_deref()) + .bind(record.name_present) .bind(record.name.as_deref()) + .bind(record.force_capabilities.is_some()) + .bind(optional_json_to_string( + &record.force_capabilities.clone().flatten(), + "api_keys.force_capabilities", + )?) .bind(record.rate_limit_present) .bind(record.rate_limit) .bind(record.concurrent_limit_present) @@ -565,13 +633,147 @@ WHERE id = ? self.reload_export_by_id(&record.api_key_id).await } + async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + if restored.api_key_id != expected.api_key_id + || restored.user_id != expected.user_id + || restored.key_hash != expected.key_hash + || restored.is_standalone != expected.is_standalone + { + return Ok(false); + } + + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let select_sql = format!("{EXPORT_COLUMNS} WHERE api_keys.id = ? LIMIT 1"); + let row = sqlx::query(&select_sql) + .bind(&expected.api_key_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_auth_api_key_export_row(&row)?; + if current != *expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = ?, + name = ?, + allowed_providers = ?, + allowed_api_formats = ?, + allowed_models = ?, + ip_rules = ?, + rate_limit = ?, + concurrent_limit = ?, + force_capabilities = ?, + feature_settings = ?, + is_active = ?, + expires_at = ?, + auto_delete_on_expiry = ?, + total_requests = ?, + total_tokens = ?, + total_cost_usd = ?, + last_used_at = ?, + updated_at = ? +WHERE id = ? + AND user_id = ? + AND key_hash = ? + AND is_standalone = ? +"#, + ) + .bind(restored.key_encrypted.as_deref()) + .bind(restored.name.as_deref()) + .bind(json_string_from_string_list( + restored.allowed_providers.as_ref(), + "api_keys.allowed_providers", + )?) + .bind(json_string_from_string_list( + restored.allowed_api_formats.as_ref(), + "api_keys.allowed_api_formats", + )?) + .bind(json_string_from_string_list( + restored.allowed_models.as_ref(), + "api_keys.allowed_models", + )?) + .bind(json_string_from_string_list( + restored.ip_rules.as_ref(), + "api_keys.ip_rules", + )?) + .bind(restored.rate_limit) + .bind(restored.concurrent_limit) + .bind(optional_json_to_string( + &restored.force_capabilities, + "api_keys.force_capabilities", + )?) + .bind(optional_json_to_string( + &restored.feature_settings, + "api_keys.feature_settings", + )?) + .bind(restored.is_active) + .bind(optional_i64_from_u64( + restored.expires_at_unix_secs, + "api_keys.expires_at", + )?) + .bind(restored.auto_delete_on_expiry) + .bind(i64_from_u64( + restored.total_requests, + "api_keys.total_requests", + )?) + .bind(i64_from_u64( + restored.total_tokens, + "api_keys.total_tokens", + )?) + .bind(restored.total_cost_usd) + .bind(optional_i64_from_u64( + restored.last_used_at_unix_secs, + "api_keys.last_used_at", + )?) + .bind(current_unix_secs() as i64) + .bind(&restored.api_key_id) + .bind(&restored.user_id) + .bind(&restored.key_hash) + .bind(restored.is_standalone) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn set_user_api_key_active( &self, user_id: &str, api_key_id: &str, is_active: bool, ) -> Result, DataLayerError> { - self.set_active(api_key_id, Some(user_id), is_active, false) + self.set_active(api_key_id, Some(user_id), is_active, false, false) + .await + } + + async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + self.set_active(api_key_id, Some(user_id), is_active, false, true) .await } @@ -580,7 +782,8 @@ WHERE id = ? api_key_id: &str, is_active: bool, ) -> Result, DataLayerError> { - self.set_active(api_key_id, None, is_active, true).await + self.set_active(api_key_id, None, is_active, true, false) + .await } async fn set_user_api_key_locked( @@ -615,26 +818,23 @@ WHERE id = ? api_key_id: &str, allowed_providers: Option>, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET allowed_providers = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, + self.set_user_api_key_allowed_providers_scoped( + user_id, + api_key_id, + allowed_providers, + false, ) - .bind(json_string_from_string_list( - allowed_providers.as_ref(), - "api_keys.allowed_providers", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, DataLayerError> { + self.set_user_api_key_allowed_providers_scoped(user_id, api_key_id, allowed_providers, true) + .await } async fn set_user_api_key_force_capabilities( @@ -643,26 +843,28 @@ WHERE id = ? api_key_id: &str, force_capabilities: Option, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET force_capabilities = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, + self.set_user_api_key_force_capabilities_scoped( + user_id, + api_key_id, + force_capabilities, + false, + ) + .await + } + + async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, DataLayerError> { + self.set_user_api_key_force_capabilities_scoped( + user_id, + api_key_id, + force_capabilities, + true, ) - .bind(optional_json_to_string( - &force_capabilities, - "api_keys.force_capabilities", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await } async fn set_user_api_key_feature_settings( @@ -671,26 +873,18 @@ WHERE id = ? api_key_id: &str, feature_settings: Option, ) -> Result, DataLayerError> { - sqlx::query( - r#" -UPDATE api_keys -SET feature_settings = ?, updated_at = ? -WHERE id = ? - AND user_id = ? - AND is_standalone = 0 -"#, - ) - .bind(optional_json_to_string( - &feature_settings, - "api_keys.feature_settings", - )?) - .bind(current_unix_secs() as i64) - .bind(api_key_id) - .bind(user_id) - .execute(&self.pool) - .await - .map_sql_err()?; - self.reload_export_by_id(api_key_id).await + self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, false) + .await + } + + async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, DataLayerError> { + self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, true) + .await } async fn set_api_key_usage_totals( @@ -700,6 +894,11 @@ WHERE id = ? total_tokens: u64, total_cost_usd: f64, ) -> Result, DataLayerError> { + if !total_cost_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "api_keys.total_cost_usd is not finite".to_string(), + )); + } sqlx::query( r#" UPDATE api_keys @@ -710,8 +909,8 @@ SET total_requests = ?, WHERE id = ? "#, ) - .bind(total_requests as i64) - .bind(total_tokens as i64) + .bind(i64_from_u64(total_requests, "api_keys.total_requests")?) + .bind(i64_from_u64(total_tokens, "api_keys.total_tokens")?) .bind(total_cost_usd) .bind(current_unix_secs() as i64) .bind(api_key_id) @@ -726,11 +925,21 @@ WHERE id = ? user_id: &str, api_key_id: &str, ) -> Result { - self.delete_api_key(api_key_id, Some(user_id), false).await + self.delete_api_key(api_key_id, Some(user_id), false, false) + .await + } + + async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + self.delete_api_key(api_key_id, Some(user_id), false, true) + .await } async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result { - self.delete_api_key(api_key_id, None, true).await + self.delete_api_key(api_key_id, None, true, false).await } async fn set_standalone_api_key_feature_settings( @@ -760,12 +969,65 @@ WHERE id = ? } impl SqliteAuthApiKeyReadRepository { + async fn update_user_api_key_basic_scoped( + &self, + record: UpdateUserApiKeyBasicRecord, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END, + name = CASE WHEN ? THEN ? ELSE name END, + rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, + concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END, + ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END, + feature_settings = CASE WHEN ? THEN ? ELSE feature_settings END, + updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(record.key_encrypted_present) + .bind(record.key_encrypted.as_deref()) + .bind(record.name_present) + .bind(record.name.as_deref()) + .bind(record.rate_limit_present) + .bind(record.rate_limit) + .bind(record.concurrent_limit_present) + .bind(record.concurrent_limit) + .bind(record.ip_rules.is_some()) + .bind(json_string_from_nested_string_list( + &record.ip_rules, + "api_keys.ip_rules", + )?) + .bind(record.feature_settings.is_some()) + .bind(optional_json_to_string( + &record.feature_settings.clone().flatten(), + "api_keys.feature_settings", + )?) + .bind(current_unix_secs() as i64) + .bind(&record.api_key_id) + .bind(&record.user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(&record.api_key_id).await + } + async fn set_active( &self, api_key_id: &str, user_id: Option<&str>, is_active: bool, is_standalone: bool, + require_unlocked: bool, ) -> Result, DataLayerError> { let mut builder = QueryBuilder::::new("UPDATE api_keys SET is_active = "); builder @@ -779,7 +1041,115 @@ impl SqliteAuthApiKeyReadRepository { if let Some(user_id) = user_id { builder.push(" AND user_id = ").push_bind(user_id); } - builder.build().execute(&self.pool).await.map_sql_err()?; + if require_unlocked { + builder.push(" AND is_locked = ").push_bind(false); + } + let result = builder.build().execute(&self.pool).await.map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_allowed_providers_scoped( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET allowed_providers = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(json_string_from_string_list( + allowed_providers.as_ref(), + "api_keys.allowed_providers", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_force_capabilities_scoped( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET force_capabilities = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(optional_json_to_string( + &force_capabilities, + "api_keys.force_capabilities", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.reload_export_by_id(api_key_id).await + } + + async fn set_user_api_key_feature_settings_scoped( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + require_unlocked: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + r#" +UPDATE api_keys +SET feature_settings = ?, updated_at = ? +WHERE id = ? + AND user_id = ? + AND is_standalone = 0 + AND (? = 0 OR is_locked = 0) +"#, + ) + .bind(optional_json_to_string( + &feature_settings, + "api_keys.feature_settings", + )?) + .bind(current_unix_secs() as i64) + .bind(api_key_id) + .bind(user_id) + .bind(require_unlocked) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } self.reload_export_by_id(api_key_id).await } @@ -788,22 +1158,80 @@ impl SqliteAuthApiKeyReadRepository { api_key_id: &str, user_id: Option<&str>, is_standalone: bool, + require_unlocked: bool, ) -> Result { - let mut builder = QueryBuilder::::new("DELETE FROM api_keys WHERE id = "); - builder - .push_bind(api_key_id) - .push(" AND is_standalone = ") - .push_bind(is_standalone); - if let Some(user_id) = user_id { - builder.push(" AND user_id = ").push_bind(user_id); - } - let rows_affected = builder - .build() - .execute(&self.pool) + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let matching_api_key = if let Some(user_id) = user_id { + if require_unlocked { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 AND is_locked = 0", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0", + ) + .bind(api_key_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + } + } else { + sqlx::query_scalar::<_, String>( + "SELECT id FROM api_keys WHERE id = ? AND is_standalone = 1", + ) + .bind(api_key_id) + .fetch_optional(&mut *tx) .await .map_sql_err()? - .rows_affected(); - Ok(rows_affected > 0) + }; + if matching_api_key.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + sqlx::query( + "UPDATE wallets SET status = 'disabled', updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE api_key_id = ? AND status <> 'disabled'", + ) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + for sql in SQLITE_ANONYMIZE_API_KEY_HISTORY_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + for sql in SQLITE_DELETE_API_KEY_DEPENDENTS_SQL { + sqlx::query(sql) + .bind(api_key_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + let result = sqlx::query("DELETE FROM api_keys WHERE id = ? AND is_standalone = ?") + .bind(api_key_id) + .bind(is_standalone) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) } } @@ -827,6 +1255,7 @@ async fn summarize_api_keys( is_standalone: bool, now_unix_secs: u64, ) -> Result { + let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let row = sqlx::query( r#" SELECT @@ -836,7 +1265,7 @@ FROM api_keys WHERE is_standalone = ? "#, ) - .bind(now_unix_secs as i64) + .bind(now_unix_secs) .bind(is_standalone) .fetch_one(pool) .await @@ -1045,6 +1474,7 @@ mod tests { UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; use serde_json::json; + use sqlx::Row; #[tokio::test] async fn sqlite_repository_reads_auth_api_key_contract_views() { @@ -1131,6 +1561,31 @@ mod tests { seed_auth_user(&pool).await; let repository = SqliteAuthApiKeyReadRepository::new(pool); + let missing_owner_key = repository + .create_user_api_key(CreateUserApiKeyRecord { + user_id: "missing-user".to_string(), + api_key_id: "key-missing-owner".to_string(), + key_hash: "hash-missing-owner".to_string(), + key_encrypted: Some("enc-missing-owner".to_string()), + name: Some("Missing Owner".to_string()), + allowed_providers: None, + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + rate_limit: 100, + concurrent_limit: None, + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + }) + .await + .expect("missing-owner creation should resolve"); + assert!(missing_owner_key.is_none()); let user_key = repository .create_user_api_key(CreateUserApiKeyRecord { user_id: "user-1".to_string(), @@ -1145,6 +1600,7 @@ mod tests { rate_limit: 100, concurrent_limit: Some(5), force_capabilities: Some(json!({"cache": true})), + feature_settings: Some(json!({"chat_pii_redaction": {"enabled": true}})), is_active: true, expires_at_unix_secs: Some(2_000_000_000), auto_delete_on_expiry: false, @@ -1157,21 +1613,84 @@ mod tests { .expect("user key should reload"); assert_eq!(user_key.allowed_models, Some(vec!["gpt-4.1".to_string()])); assert_eq!(user_key.total_tokens, 42); + assert_eq!( + user_key.feature_settings, + Some(json!({"chat_pii_redaction": {"enabled": true}})) + ); let updated_user_key = repository .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { user_id: "user-1".to_string(), api_key_id: "key-created-user".to_string(), + key_encrypted: Some("enc-user-rotated".to_string()), + key_encrypted_present: true, name: Some("Updated User".to_string()), + name_present: true, rate_limit: Some(150), + rate_limit_present: true, concurrent_limit: Some(6), + concurrent_limit_present: true, ip_rules: Some(Some(vec!["10.0.0.0/24".to_string()])), + feature_settings: Some(Some(json!({"compact": true}))), }) .await .expect("user key should update") .expect("user key should reload"); assert_eq!(updated_user_key.name, Some("Updated User".to_string())); + assert_eq!( + updated_user_key.key_encrypted.as_deref(), + Some("enc-user-rotated") + ); assert_eq!(updated_user_key.concurrent_limit, Some(6)); + assert_eq!( + updated_user_key.feature_settings, + Some(json!({"compact": true})) + ); + + // Rollback can explicitly restore nullable values, while a present zero remains a + // meaningful rate limit rather than being treated as an omitted field. + let cleared_user_key = repository + .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-created-user".to_string(), + key_encrypted: None, + key_encrypted_present: true, + name: None, + name_present: true, + rate_limit: None, + rate_limit_present: true, + concurrent_limit: None, + concurrent_limit_present: true, + ip_rules: None, + feature_settings: None, + }) + .await + .expect("nullable values should clear") + .expect("user key should remain"); + assert!(cleared_user_key.key_encrypted.is_none()); + assert!(cleared_user_key.name.is_none()); + assert!(cleared_user_key.rate_limit.is_none()); + assert!(cleared_user_key.concurrent_limit.is_none()); + + let zero_rate_limit = repository + .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-created-user".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: None, + name_present: false, + rate_limit: Some(0), + rate_limit_present: true, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: None, + }) + .await + .expect("zero rate limit should persist") + .expect("user key should remain"); + assert_eq!(zero_rate_limit.rate_limit, Some(0)); assert!(repository .set_user_api_key_locked("user-1", "key-created-user", true) @@ -1184,6 +1703,78 @@ mod tests { .expect("snapshot should exist"); assert!(snapshot.api_key_is_locked); + assert!(repository + .update_user_api_key_basic_if_unlocked(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-created-user".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: Some("must-not-change".to_string()), + name_present: true, + rate_limit: None, + rate_limit_present: false, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: Some(Some(json!({"must_not_change": true}))), + }) + .await + .expect("locked basic update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_active_if_unlocked("user-1", "key-created-user", false) + .await + .expect("locked status update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_allowed_providers_if_unlocked( + "user-1", + "key-created-user", + Some(vec!["must-not-change".to_string()]), + ) + .await + .expect("locked provider update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_force_capabilities_if_unlocked( + "user-1", + "key-created-user", + Some(json!({"must_not_change": true})), + ) + .await + .expect("locked capability update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_feature_settings_if_unlocked( + "user-1", + "key-created-user", + Some(json!({"must_not_change": true})), + ) + .await + .expect("locked feature update should resolve") + .is_none()); + assert!(!repository + .delete_user_api_key_if_unlocked("user-1", "key-created-user") + .await + .expect("locked delete should resolve")); + + let unchanged = repository + .reload_export_by_id("key-created-user") + .await + .expect("locked key should reload") + .expect("locked key should still exist"); + assert_ne!(unchanged.name.as_deref(), Some("must-not-change")); + assert!(unchanged.is_active); + assert_ne!( + unchanged.allowed_providers, + Some(vec!["must-not-change".to_string()]) + ); + assert_ne!( + unchanged.force_capabilities, + Some(json!({"must_not_change": true})) + ); + assert_eq!(unchanged.feature_settings, Some(json!({"compact": true}))); + let active_user_key = repository .set_user_api_key_active("user-1", "key-created-user", false) .await @@ -1191,6 +1782,12 @@ mod tests { .expect("user key should reload"); assert!(!active_user_key.is_active); + assert!(repository + .set_user_api_key_active("wrong-owner", "key-created-user", true) + .await + .expect("wrong-owner status update should resolve") + .is_none()); + let provider_updated = repository .set_user_api_key_allowed_providers( "user-1", @@ -1260,7 +1857,11 @@ mod tests { let standalone = repository .update_standalone_api_key_basic(UpdateStandaloneApiKeyBasicRecord { api_key_id: "key-created-standalone".to_string(), + key_encrypted: Some("enc-standalone-rotated".to_string()), + key_encrypted_present: true, name: Some("Updated Standalone".to_string()), + name_present: true, + force_capabilities: None, rate_limit_present: true, rate_limit: Some(20), concurrent_limit_present: true, @@ -1278,6 +1879,10 @@ mod tests { .expect("standalone key should update") .expect("standalone key should reload"); assert_eq!(standalone.name, Some("Updated Standalone".to_string())); + assert_eq!( + standalone.key_encrypted.as_deref(), + Some("enc-standalone-rotated") + ); assert_eq!(standalone.allowed_providers, None); assert_eq!( standalone.allowed_api_formats, @@ -1303,6 +1908,250 @@ mod tests { .expect("user key should delete")); } + #[tokio::test] + async fn sqlite_api_key_delete_is_owner_scoped_and_preserves_anonymized_facts() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users ( + id, username, role, auth_source, is_active, is_deleted, created_at, updated_at +) VALUES + ('key-owner', 'key-owner', 'user', 'local', 1, 0, 1, 1), + ('other-owner', 'other-owner', 'user', 'local', 1, 0, 1, 1); + +INSERT INTO api_keys ( + id, user_id, key_hash, name, is_standalone, created_at, updated_at +) VALUES ( + 'key-to-delete', 'key-owner', 'key-to-delete-hash', 'private key name', 0, 1, 1 +); + +INSERT INTO wallets ( + id, api_key_id, balance, gift_balance, status, created_at, updated_at +) VALUES ( + 'key-wallet', 'key-to-delete', 12, 3, 'active', 1, 1 +); + +INSERT INTO request_candidates ( + id, request_id, user_id, api_key_id, username, api_key_name, + candidate_index, status, created_at +) VALUES ( + 'key-history-row', 'key-candidate-request', 'key-owner', 'key-to-delete', + 'key-owner', 'private key name', 0, 'success', 1 +); + +INSERT INTO video_tasks ( + id, request_id, user_id, api_key_id, username, api_key_name, created_at, updated_at +) VALUES ( + 'key-history-row', 'key-video-request', 'key-owner', 'key-to-delete', + 'key-owner', 'private key name', 1, 1 +); + +INSERT INTO usage ( + request_id, id, user_id, api_key_id, username, api_key_name +) VALUES ( + 'key-usage-request', 'key-history-row', 'key-owner', 'key-to-delete', + 'key-owner', 'private key name' +); + +INSERT INTO stats_daily_api_key ( + id, api_key_id, date, api_key_name, created_at, updated_at +) VALUES ( + 'key-history-row', 'key-to-delete', 1, 'private key name', 1, 1 +); + +INSERT INTO audit_logs ( + id, event_type, api_key_id, description, ip_address, user_agent, + event_metadata, error_message, created_at +) VALUES ( + 'key-audit', 'key_event', 'key-to-delete', 'private description', + '192.0.2.20', 'private agent', '{"private":true}', 'private error', 1 +); + +INSERT INTO api_key_provider_mappings ( + id, api_key_id, provider_id, created_at, updated_at +) VALUES ( + 'key-mapping', 'key-to-delete', 'provider-1', 1, 1 +); + +INSERT INTO payment_orders ( + id, order_no, wallet_id, amount_usd, payment_method, + gateway_response, status, created_at +) VALUES ( + 'key-order', 'key-order-no', 'key-wallet', 12, 'test', + '{"customer_email":"key@example.com"}', 'credited', 1 +); + +INSERT INTO payment_callbacks ( + id, payment_order_id, payment_method, callback_key, order_no, + payload_hash, signature_valid, status, payload, error_message, created_at +) VALUES ( + 'key-callback', 'key-order', 'test', 'key-callback-key', 'key-order-no', + 'key-payload-hash', 1, 'processed', + '{"customer_email":"key@example.com"}', 'private callback error', 1 +); +"#, + ) + .execute(&pool) + .await + .expect("API key deletion fixtures should insert"); + + let repository = SqliteAuthApiKeyReadRepository::new(pool.clone()); + assert!(!repository + .delete_user_api_key("other-owner", "key-to-delete") + .await + .expect("wrong-owner delete should resolve")); + assert_eq!( + sqlx::query_scalar::<_, String>("SELECT status FROM wallets WHERE id = 'key-wallet'",) + .fetch_one(&pool) + .await + .expect("wallet status should load"), + "active" + ); + assert_eq!( + sqlx::query_scalar::<_, String>( + "SELECT api_key_name FROM request_candidates WHERE id = 'key-history-row'", + ) + .fetch_one(&pool) + .await + .expect("candidate key name should load"), + "private key name" + ); + + assert!(repository + .delete_user_api_key("key-owner", "key-to-delete") + .await + .expect("owner-scoped delete should succeed")); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM api_keys WHERE id = 'key-to-delete'", + ) + .fetch_one(&pool) + .await + .expect("API key count should load"), + 0 + ); + + let wallet = sqlx::query("SELECT api_key_id, status FROM wallets WHERE id = 'key-wallet'") + .fetch_one(&pool) + .await + .expect("wallet fact should remain"); + assert_eq!( + wallet + .try_get::, _>("api_key_id") + .expect("wallet API key id should decode") + .as_deref(), + Some("key-to-delete") + ); + assert_eq!( + wallet + .try_get::("status") + .expect("wallet status should decode"), + "disabled" + ); + + for table in [ + "request_candidates", + "video_tasks", + "usage", + "stats_daily_api_key", + ] { + let row = sqlx::query(&format!( + "SELECT api_key_id, api_key_name FROM {table} WHERE id = 'key-history-row'", + )) + .fetch_one(&pool) + .await + .unwrap_or_else(|error| panic!("{table} fact should remain: {error}")); + assert_eq!( + row.try_get::("api_key_id") + .expect("history API key id should decode"), + "key-to-delete", + "{table} API key id must remain stable" + ); + assert_eq!( + row.try_get::, _>("api_key_name") + .expect("history API key name should decode"), + None, + "{table} API key name must be removed" + ); + } + + let audit = sqlx::query( + "SELECT api_key_id, description, ip_address, user_agent, event_metadata, error_message FROM audit_logs WHERE id = 'key-audit'", + ) + .fetch_one(&pool) + .await + .expect("audit fact should remain"); + assert_eq!( + audit + .try_get::, _>("api_key_id") + .expect("audit API key id should decode") + .as_deref(), + Some("key-to-delete") + ); + assert_eq!( + audit + .try_get::("description") + .expect("audit description should decode"), + "deleted API key event" + ); + for column in [ + "ip_address", + "user_agent", + "event_metadata", + "error_message", + ] { + assert_eq!( + audit + .try_get::, _>(column) + .unwrap_or_else(|error| panic!("audit {column} should decode: {error}")), + None, + "audit {column} must be removed" + ); + } + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM api_key_provider_mappings WHERE id = 'key-mapping'", + ) + .fetch_one(&pool) + .await + .expect("mapping count should load"), + 0 + ); + + let order_gateway_response: Option = sqlx::query_scalar( + "SELECT gateway_response FROM payment_orders WHERE id = 'key-order'", + ) + .fetch_one(&pool) + .await + .expect("payment order should remain"); + assert_eq!(order_gateway_response, None); + let callback = sqlx::query( + "SELECT payload, error_message FROM payment_callbacks WHERE id = 'key-callback'", + ) + .fetch_one(&pool) + .await + .expect("payment callback should remain"); + assert_eq!( + callback + .try_get::, _>("payload") + .expect("callback payload should decode"), + None + ); + assert_eq!( + callback + .try_get::, _>("error_message") + .expect("callback error should decode"), + None + ); + } + async fn seed_auth_api_key_rows(pool: &sqlx::SqlitePool) { seed_auth_user(pool).await; sqlx::query( diff --git a/crates/aether-data/adapters/sqlite/src/auth_modules.rs b/crates/aether-data/adapters/sqlite/src/auth_modules.rs index a4243aa7f..2c1261888 100644 --- a/crates/aether-data/adapters/sqlite/src/auth_modules.rs +++ b/crates/aether-data/adapters/sqlite/src/auth_modules.rs @@ -3,7 +3,7 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::auth_modules::*; use aether_data_contracts::DataLayerError; -use aether_data_query::{push_eq, push_limit, WhereClause}; +use aether_data_query::{push_eq, WhereClause}; use crate::error::SqlResultExt; use crate::SqlitePool; @@ -35,6 +35,87 @@ SELECT FROM ldap_configs "#; +const UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL: &str = r#" +UPDATE ldap_configs +SET + server_url = ?, + bind_dn = ?, + base_dn = ?, + user_search_filter = ?, + username_attr = ?, + email_attr = ?, + display_name_attr = ?, + is_enabled = ?, + is_exclusive = ?, + use_starttls = ?, + connect_timeout = ?, + updated_at = MAX(updated_at + 1, ?) +WHERE singleton_key = 1 + AND server_url IS ? + AND bind_dn IS ? + AND bind_password_encrypted IS ? + AND base_dn IS ? + AND user_search_filter IS ? + AND username_attr IS ? + AND email_attr IS ? + AND display_name_attr IS ? + AND is_enabled IS ? + AND is_exclusive IS ? + AND use_starttls IS ? + AND connect_timeout IS ? +"#; + +const UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL: &str = r#" +UPDATE ldap_configs +SET + server_url = ?, + bind_dn = ?, + bind_password_encrypted = ?, + base_dn = ?, + user_search_filter = ?, + username_attr = ?, + email_attr = ?, + display_name_attr = ?, + is_enabled = ?, + is_exclusive = ?, + use_starttls = ?, + connect_timeout = ?, + updated_at = MAX(updated_at + 1, ?) +WHERE singleton_key = 1 + AND server_url IS ? + AND bind_dn IS ? + AND bind_password_encrypted IS ? + AND base_dn IS ? + AND user_search_filter IS ? + AND username_attr IS ? + AND email_attr IS ? + AND display_name_attr IS ? + AND is_enabled IS ? + AND is_exclusive IS ? + AND use_starttls IS ? + AND connect_timeout IS ? +"#; + +const INSERT_LDAP_CONFIG_SQL: &str = r#" +INSERT INTO ldap_configs ( + singleton_key, + server_url, + bind_dn, + bind_password_encrypted, + base_dn, + user_search_filter, + username_attr, + email_attr, + display_name_attr, + is_enabled, + is_exclusive, + use_starttls, + connect_timeout, + created_at, + updated_at +) VALUES (1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +"#; + #[derive(Debug, Clone)] pub struct SqliteAuthModuleReadRepository { pool: SqlitePool, @@ -72,8 +153,7 @@ async fn get_ldap_config( pool: &SqlitePool, ) -> Result, DataLayerError> { let mut builder = QueryBuilder::::new(LDAP_CONFIG_COLUMNS); - builder.push(" ORDER BY id ASC"); - push_limit(&mut builder, 1); + builder.push(" WHERE singleton_key = 1"); let row = builder.build().fetch_optional(pool).await.map_sql_err()?; row.as_ref().map(map_ldap_row).transpose() } @@ -106,95 +186,208 @@ impl AuthModuleReadRepository for SqliteAuthModuleRepository { #[async_trait] impl AuthModuleWriteRepository for SqliteAuthModuleRepository { - async fn upsert_ldap_config( + async fn compare_and_swap_ldap_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, DataLayerError> { + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result { + let persisted = + ldap_config_after_password_update(expected, replacement, bind_password_update)?; let now = now_unix_secs(); - let updated = sqlx::query( + let Some(expected) = expected else { + let insert = sqlx::query(INSERT_LDAP_CONFIG_SQL) + .bind(&persisted.server_url) + .bind(&persisted.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&persisted.base_dn) + .bind(persisted.user_search_filter.as_deref()) + .bind(persisted.username_attr.as_deref()) + .bind(persisted.email_attr.as_deref()) + .bind(persisted.display_name_attr.as_deref()) + .bind(persisted.is_enabled) + .bind(persisted.is_exclusive) + .bind(persisted.use_starttls) + .bind(persisted.connect_timeout) + .bind(now as i64) + .bind(now as i64) + .execute(&self.pool) + .await; + return match insert { + Ok(result) if result.rows_affected() == 1 => { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } + Ok(_) => Ok(CompareAndSwapLdapConfigResult::Conflict), + Err(error) + if error + .as_database_error() + .is_some_and(|error| error.is_unique_violation()) => + { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } + Err(error) => Err(DataLayerError::sql(error)), + }; + }; + + let rows_affected = match bind_password_update { + LdapBindPasswordUpdate::Preserve => { + sqlx::query(UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected() + } + LdapBindPasswordUpdate::Set(_) | LdapBindPasswordUpdate::Clear => { + sqlx::query(UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL) + .bind(&replacement.server_url) + .bind(&replacement.bind_dn) + .bind(persisted.bind_password_encrypted.as_deref()) + .bind(&replacement.base_dn) + .bind(replacement.user_search_filter.as_deref()) + .bind(replacement.username_attr.as_deref()) + .bind(replacement.email_attr.as_deref()) + .bind(replacement.display_name_attr.as_deref()) + .bind(replacement.is_enabled) + .bind(replacement.is_exclusive) + .bind(replacement.use_starttls) + .bind(replacement.connect_timeout) + .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected() + } + }; + if rows_affected == 1 { + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } else { + Ok(CompareAndSwapLdapConfigResult::Conflict) + } + } + + async fn delete_ldap_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result { + let rows_affected = sqlx::query( r#" -UPDATE ldap_configs -SET - server_url = ?, - bind_dn = ?, - bind_password_encrypted = ?, - base_dn = ?, - user_search_filter = ?, - username_attr = ?, - email_attr = ?, - display_name_attr = ?, - is_enabled = ?, - is_exclusive = ?, - use_starttls = ?, - connect_timeout = ?, - updated_at = ? -WHERE id = ( - SELECT id - FROM ldap_configs - ORDER BY id ASC - LIMIT 1 -) +DELETE FROM ldap_configs +WHERE singleton_key = 1 + AND server_url IS ? + AND bind_dn IS ? + AND bind_password_encrypted IS ? + AND base_dn IS ? + AND user_search_filter IS ? + AND username_attr IS ? + AND email_attr IS ? + AND display_name_attr IS ? + AND is_enabled IS ? + AND is_exclusive IS ? + AND use_starttls IS ? + AND connect_timeout IS ? "#, ) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .bind(now as i64) + .bind(&expected.server_url) + .bind(&expected.bind_dn) + .bind(expected.bind_password_encrypted.as_deref()) + .bind(&expected.base_dn) + .bind(expected.user_search_filter.as_deref()) + .bind(expected.username_attr.as_deref()) + .bind(expected.email_attr.as_deref()) + .bind(expected.display_name_attr.as_deref()) + .bind(expected.is_enabled) + .bind(expected.is_exclusive) + .bind(expected.use_starttls) + .bind(expected.connect_timeout) .execute(&self.pool) .await - .map_sql_err()?; - - if updated.rows_affected() == 0 { - sqlx::query( - r#" -INSERT INTO ldap_configs ( - server_url, - bind_dn, - bind_password_encrypted, - base_dn, - user_search_filter, - username_attr, - email_attr, - display_name_attr, - is_enabled, - is_exclusive, - use_starttls, - connect_timeout, - created_at, - updated_at -) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) -"#, - ) - .bind(&config.server_url) - .bind(&config.bind_dn) - .bind(config.bind_password_encrypted.as_deref()) - .bind(&config.base_dn) - .bind(config.user_search_filter.as_deref()) - .bind(config.username_attr.as_deref()) - .bind(config.email_attr.as_deref()) - .bind(config.display_name_attr.as_deref()) - .bind(config.is_enabled) - .bind(config.is_exclusive) - .bind(config.use_starttls) - .bind(config.connect_timeout) - .bind(now as i64) - .bind(now as i64) - .execute(&self.pool) - .await - .map_sql_err()?; - } - - self.get_ldap_config().await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) } + + async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + let rows_affected = sqlx::query( + r#" +UPDATE ldap_configs +SET bind_password_encrypted = ?, updated_at = MAX(updated_at + 1, ?) +WHERE singleton_key = 1 + AND bind_password_encrypted = ? +"#, + ) + .bind(replacement) + .bind(now_unix_secs() as i64) + .bind(expected) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } +} + +fn ldap_config_after_password_update( + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, +) -> Result { + let bind_password_encrypted = match bind_password_update { + LdapBindPasswordUpdate::Preserve => expected + .ok_or_else(|| { + DataLayerError::InvalidConfiguration( + "LDAP bind password cannot be preserved while creating the singleton" + .to_string(), + ) + })? + .bind_password_encrypted + .clone(), + LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()), + LdapBindPasswordUpdate::Clear => None, + }; + Ok(StoredLdapModuleConfig { + bind_password_encrypted, + ..replacement.clone() + }) } fn now_unix_secs() -> u64 { @@ -232,7 +425,8 @@ fn map_ldap_row(row: &SqliteRow) -> Result Result { + run.sanitize_for_persistence(); run.validate()?; sqlx::query( r#" @@ -307,8 +308,9 @@ ON CONFLICT(id) DO UPDATE SET async fn upsert_event( &self, - event: UpsertBackgroundTaskEvent, + mut event: UpsertBackgroundTaskEvent, ) -> Result { + event.sanitize_for_persistence(); event.validate()?; sqlx::query( r#" @@ -361,7 +363,7 @@ fn map_run_row(row: &SqliteRow) -> Result = row.try_get("finished_at_unix_secs").map_sql_err()?; let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?; - Ok(StoredBackgroundTaskRun { + let mut run = StoredBackgroundTaskRun { id: row.try_get("id").map_sql_err()?, task_key: row.try_get("task_key").map_sql_err()?, kind: BackgroundTaskKind::from_database(&kind)?, @@ -381,19 +383,23 @@ fn map_run_row(row: &SqliteRow) -> Result Result { let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?; - Ok(StoredBackgroundTaskEvent { + let mut event = StoredBackgroundTaskEvent { id: row.try_get("id").map_sql_err()?, run_id: row.try_get("run_id").map_sql_err()?, event_type: row.try_get("event_type").map_sql_err()?, message: row.try_get("message").map_sql_err()?, payload_json: parse_optional_json(row.try_get("payload_json").map_sql_err()?)?, created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(), - }) + }; + event.sanitize_persisted_data(); + Ok(event) } fn parse_optional_json(value: Option) -> Result, DataLayerError> { @@ -451,7 +457,11 @@ mod tests { owner_instance: None, progress_percent: 0, progress_message: None, - payload_json: Some(json!({"partition": 7})), + payload_json: Some(json!({ + "partition": 7, + "refresh_token": "sensitive-refresh-token", + "nested": {"authorization": "Bearer sensitive"} + })), result_json: None, error_message: None, cancel_requested: false, @@ -471,7 +481,10 @@ mod tests { run_id: run.id.clone(), event_type: "queued".to_string(), message: "task queued".to_string(), - payload_json: Some(json!({"attempt": 0})), + payload_json: Some(json!({ + "error_code": "provider_delete_failed", + "error": "sensitive provider detail" + })), created_at_unix_secs: 11, }) .await @@ -496,7 +509,11 @@ mod tests { .await .expect("background task events should list"); assert_eq!(events.len(), 1); - assert_eq!(events[0].payload_json, Some(json!({"attempt": 0}))); + assert_eq!(events[0].message, "queued"); + assert_eq!( + events[0].payload_json, + Some(json!({"error_code": "provider_delete_failed"})) + ); assert!(repository .request_cancel("run-1", 20) diff --git a/crates/aether-data/adapters/sqlite/src/billing.rs b/crates/aether-data/adapters/sqlite/src/billing.rs index f5242fb84..e13c08abf 100644 --- a/crates/aether-data/adapters/sqlite/src/billing.rs +++ b/crates/aether-data/adapters/sqlite/src/billing.rs @@ -4,8 +4,9 @@ use sqlx::{sqlite::SqliteRow, Row}; use aether_data_contracts::repository::billing::{ AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; use aether_data_contracts::DataLayerError; @@ -637,6 +638,126 @@ LIMIT 1 .transpose() } + async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE payment_gateway_configs +SET merchant_key_encrypted = ? +WHERE provider = ? + AND merchant_key_encrypted IS ? + "#, + ) + .bind(&update.merchant_key_encrypted) + .bind(update.provider.trim().to_ascii_lowercase()) + .bind(&update.expected_merchant_key_encrypted) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_swap_payment_gateway_config( + &self, + mutation: &PaymentGatewayConfigCasWriteInput, + ) -> Result, DataLayerError> { + let input = &mutation.input; + let provider = input.provider.trim().to_ascii_lowercase(); + let now = current_unix_secs_i64(); + let mut tx = self.pool.begin().await.map_sql_err()?; + let result = if mutation.expected_existing { + sqlx::query( + r#" +UPDATE payment_gateway_configs +SET + enabled = ?, + endpoint_url = ?, + callback_base_url = ?, + merchant_id = ?, + merchant_key_encrypted = CASE + WHEN ? THEN merchant_key_encrypted + ELSE ? + END, + pay_currency = ?, + usd_exchange_rate = ?, + min_recharge_usd = ?, + channels_json = ?, + updated_at = ? +WHERE provider = ? + AND merchant_key_encrypted IS ? + "#, + ) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.preserve_existing_secret) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(json_to_string(&input.channels_json)?) + .bind(now) + .bind(&provider) + .bind(mutation.expected_merchant_key_encrypted.as_deref()) + .execute(&mut *tx) + .await + .map_sql_err()? + } else { + sqlx::query( + r#" +INSERT INTO payment_gateway_configs ( + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd, + channels_json, created_at, updated_at +) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +ON CONFLICT(provider) DO NOTHING + "#, + ) + .bind(&provider) + .bind(input.enabled) + .bind(&input.endpoint_url) + .bind(input.callback_base_url.as_deref()) + .bind(&input.merchant_id) + .bind(input.merchant_key_encrypted.as_deref()) + .bind(&input.pay_currency) + .bind(input.usd_exchange_rate) + .bind(input.min_recharge_usd) + .bind(json_to_string(&input.channels_json)?) + .bind(now) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()? + }; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(AdminBillingMutationOutcome::NotFound); + } + + let row = sqlx::query( + r#" +SELECT + provider, enabled, endpoint_url, callback_base_url, merchant_id, + merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd, + channels_json, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs +FROM payment_gateway_configs +WHERE provider = ? +LIMIT 1 + "#, + ) + .bind(&provider) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let record = map_payment_gateway_config_sqlite(&row)?; + tx.commit().await.map_sql_err()?; + Ok(AdminBillingMutationOutcome::Applied(record)) + } + async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -920,6 +1041,40 @@ ORDER BY expires_at ASC, created_at ASC )) } + async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, DataLayerError> { + let now = current_unix_secs_i64(); + let result = sqlx::query( + r#" +UPDATE user_plan_entitlements +SET status = 'revoked', + expires_at = CASE WHEN expires_at > ? THEN ? ELSE expires_at END, + updated_at = ? +WHERE id = ? + AND user_id = ? + AND status = 'active' + AND expires_at > ? + "#, + ) + .bind(now) + .bind(now) + .bind(now) + .bind(entitlement_id) + .bind(user_id) + .bind(now) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + Ok(AdminBillingMutationOutcome::NotFound) + } else { + Ok(AdminBillingMutationOutcome::Applied(())) + } + } + async fn find_user_daily_quota_availability( &self, user_id: &str, @@ -927,13 +1082,19 @@ ORDER BY expires_at ASC, created_at ASC let now_unix_secs = current_unix_secs_i64(); let rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = ? - AND status = 'active' - AND starts_at <= ? - AND expires_at > ? -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = ? + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= ? + AND user_plan_entitlements.expires_at > ? +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC "#, ) .bind(user_id) @@ -948,9 +1109,13 @@ ORDER BY expires_at ASC, created_at ASC, id ASC let entitlement_id: String = row.try_get("id").map_sql_err()?; let entitlements = parse_json(row.try_get("entitlements_snapshot").ok().flatten())? .unwrap_or_else(|| serde_json::json!([])); + let plan_entitlements = + parse_json(row.try_get("plan_entitlements_json").ok().flatten())? + .unwrap_or_else(|| serde_json::json!([])); grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -1174,6 +1339,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -1199,15 +1365,27 @@ fn daily_quota_grants_from_entitlement( .and_then(serde_json::Value::as_str), now, )?, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + fn read_count_sqlite(row: &SqliteRow) -> Result { Ok(row.try_get::("total").map_sql_err()?.max(0) as u64) } @@ -1405,7 +1583,8 @@ mod tests { use crate::run_migrations; use aether_data_contracts::repository::billing::{ AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingRuleWriteInput, - BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigWriteInput, + BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigCasWriteInput, + PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, }; #[tokio::test] @@ -1558,6 +1737,106 @@ mod tests { assert_eq!(preset.errors, Vec::::new()); } + #[tokio::test] + async fn sqlite_repository_revokes_active_user_plan_entitlement() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let now = super::current_unix_secs_i64(); + sqlx::query( + r#" +INSERT INTO users ( + id, username, email, role, auth_source, password_hash, is_active, + is_deleted, created_at, updated_at +) VALUES ( + 'user-revoke', 'revoke-user', 'revoke@example.com', 'user', 'local', + 'hash', 1, 0, 1, 1 +); +INSERT INTO wallets ( + id, user_id, balance, gift_balance, limit_mode, created_at, updated_at +) VALUES ( + 'wallet-revoke', 'user-revoke', 5.0, 0.0, 'finite', 1, 1 +); +INSERT INTO billing_plans ( + id, title, price_amount, price_currency, duration_unit, + duration_value, entitlements_json, created_at, updated_at +) VALUES ( + 'plan-revoke', 'Revocable Plan', 0.0, 'USD', 'month', 1, + '[{"type":"daily_quota","daily_quota_usd":10.0,"allow_wallet_overage":true}]', + 1, 1 +); +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd, + refundable_amount_usd, payment_method, gateway_response, status, created_at +) VALUES ( + 'order-revoke', 'order-revoke', 'wallet-revoke', 'user-revoke', 0.0, 0.0, + 0.0, 'admin_manual', '{}', 'credited', 1 +); +INSERT INTO user_plan_entitlements ( + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +) VALUES ( + 'entitlement-revoke', 'user-revoke', 'plan-revoke', 'order-revoke', + 'active', ?, ?, + '[{"type":"daily_quota","daily_quota_usd":10.0,"allow_wallet_overage":false}]', + ?, ? +); +"#, + ) + .bind(now - 60) + .bind(now + 3600) + .bind(now - 60) + .bind(now - 60) + .execute(&pool) + .await + .expect("revocable entitlement should seed"); + let repository = SqliteBillingReadRepository::new(pool.clone()); + + let quota = repository + .find_user_daily_quota_availability("user-revoke") + .await + .expect("quota should load") + .expect("quota should be available"); + assert!(quota.has_active_daily_quota); + assert!(quota.allow_wallet_overage); + + let wrong_user = repository + .revoke_user_plan_entitlement("other-user", "entitlement-revoke") + .await + .expect("ownership check should run"); + assert_eq!(wrong_user, AdminBillingMutationOutcome::NotFound); + + let outcome = repository + .revoke_user_plan_entitlement("user-revoke", "entitlement-revoke") + .await + .expect("entitlement revoke should run"); + assert_eq!(outcome, AdminBillingMutationOutcome::Applied(())); + let active = repository + .list_user_plan_entitlements("user-revoke") + .await + .expect("entitlements should load") + .expect("entitlements should be available"); + assert!(active.is_empty()); + let quota = repository + .find_user_daily_quota_availability("user-revoke") + .await + .expect("quota should load") + .expect("quota should be available"); + assert!(!quota.has_active_daily_quota); + let status: String = sqlx::query_scalar( + "SELECT status FROM user_plan_entitlements WHERE id = 'entitlement-revoke'", + ) + .fetch_one(&pool) + .await + .expect("entitlement status should load"); + assert_eq!(status, "revoked"); + } + #[tokio::test] async fn sqlite_repository_deletes_unused_billing_plans_only() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1719,6 +1998,96 @@ VALUES ('order-1', 'order-no-1', 'wallet-1', 0, 'epay', 'plan_purchase', ); } + #[tokio::test] + async fn sqlite_gateway_cas_is_create_only_and_secret_exact() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteBillingReadRepository::new(pool); + let input = PaymentGatewayConfigWriteInput { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: "https://api.stripe.com".to_string(), + callback_base_url: None, + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some("legacy-ciphertext".to_string()), + preserve_existing_secret: false, + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: json!({"channels": []}), + }; + let create = PaymentGatewayConfigCasWriteInput { + input: input.clone(), + expected_existing: false, + expected_merchant_key_encrypted: None, + }; + assert!(matches!( + repository + .compare_and_swap_payment_gateway_config(&create) + .await + .expect("create should run"), + AdminBillingMutationOutcome::Applied(_) + )); + assert_eq!( + repository + .compare_and_swap_payment_gateway_config(&create) + .await + .expect("conflicting create should run"), + AdminBillingMutationOutcome::NotFound + ); + + let before = repository + .find_payment_gateway_config("stripe") + .await + .expect("lookup should run") + .expect("config should exist"); + assert!(!repository + .compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate { + provider: "stripe".to_string(), + expected_merchant_key_encrypted: "LEGACY-ciphertext".to_string(), + merchant_key_encrypted: "v2-ciphertext".to_string(), + }) + .await + .expect("case-mismatched CAS should run")); + assert!(repository + .compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate { + provider: "stripe".to_string(), + expected_merchant_key_encrypted: "legacy-ciphertext".to_string(), + merchant_key_encrypted: "v2-ciphertext".to_string(), + }) + .await + .expect("exact CAS should run")); + let after = repository + .find_payment_gateway_config("stripe") + .await + .expect("lookup should run") + .expect("config should exist"); + assert_eq!(after.updated_at_unix_secs, before.updated_at_unix_secs); + assert_eq!( + after.merchant_key_encrypted.as_deref(), + Some("v2-ciphertext") + ); + + let stale_update = PaymentGatewayConfigCasWriteInput { + input, + expected_existing: true, + expected_merchant_key_encrypted: Some("legacy-ciphertext".to_string()), + }; + assert_eq!( + repository + .compare_and_swap_payment_gateway_config(&stale_update) + .await + .expect("stale update should run"), + AdminBillingMutationOutcome::NotFound + ); + } + async fn seed_billing_context(pool: &sqlx::SqlitePool) { sqlx::query( r#" diff --git a/crates/aether-data/adapters/sqlite/src/candidate_selection.rs b/crates/aether-data/adapters/sqlite/src/candidate_selection.rs index 54b5dbc64..86d1f6615 100644 --- a/crates/aether-data/adapters/sqlite/src/candidate_selection.rs +++ b/crates/aether-data/adapters/sqlite/src/candidate_selection.rs @@ -1070,12 +1070,12 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result, + field_name: &str, +) -> Result>, DataLayerError> { + let Some(raw) = raw else { + return Ok(None); + }; + let value = serde_json::from_str::(&raw).map_err(|err| { + DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}")) + })?; + parse_key_policy_string_list_value(&value, field_name) +} + +fn parse_key_policy_string_list_value( + value: &serde_json::Value, + field_name: &str, +) -> Result>, DataLayerError> { + match value { + serde_json::Value::Null => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains JSON null; use SQL NULL for an unset policy" + ))), + serde_json::Value::Array(array) => { + parse_key_policy_string_list_array(array, field_name).map(Some) + } + serde_json::Value::String(raw) => parse_embedded_key_policy_string_list(raw, field_name), + _ => Err(DataLayerError::UnexpectedValue(format!( + "{field_name} is not a JSON array" + ))), + } +} + +fn parse_embedded_key_policy_string_list( + raw: &str, + field_name: &str, +) -> Result>, DataLayerError> { + let raw = raw.trim(); + if raw.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty string" + ))); + } + if raw.eq_ignore_ascii_case("null") { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains stringified JSON null; use SQL NULL for an unset policy" + ))); + } + + if let Ok(decoded) = serde_json::from_str::(raw) { + return parse_key_policy_string_list_value(&decoded, field_name); + } + + Ok(Some(vec![raw.to_string()])) +} + +fn parse_key_policy_string_list_array( + array: &[serde_json::Value], + field_name: &str, +) -> Result, DataLayerError> { + let mut items = Vec::with_capacity(array.len()); + for item in array { + let Some(item) = item.as_str() else { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains a non-string item" + ))); + }; + let item = item.trim(); + if item.is_empty() { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty item" + ))); + } + items.push(item.to_string()); + } + Ok(items) +} + fn parse_string_list_value( value: &serde_json::Value, field_name: &str, @@ -1338,9 +1414,10 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec { #[cfg(test)] mod tests { use super::{ - provider_model_mapping_api_format_covers, push_key_auth_channel_sql_filter, - push_pool_key_order, vertex_key_auth_channel_matches, ExactPageAccumulator, - SqliteMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT, + parse_stored_key_policy_string_list, provider_model_mapping_api_format_covers, + push_key_auth_channel_sql_filter, push_pool_key_order, vertex_key_auth_channel_matches, + ExactPageAccumulator, SqliteMinimalCandidateSelectionReadRepository, + REQUESTED_MODEL_RAW_SCAN_LIMIT, }; use crate::run_migrations; use aether_data_contracts::repository::candidate_selection::{ @@ -1378,6 +1455,25 @@ mod tests { assert!(vertex_clause.contains("gemini:embedding")); } + #[test] + fn malformed_key_policy_never_degrades_to_unrestricted() { + for raw in ["null", "\"null\"", "\"\"", "[\"openai:chat\",null]"] { + assert!(parse_stored_key_policy_string_list( + Some(raw.to_string()), + "provider_api_keys.api_formats", + ) + .is_err()); + } + assert_eq!( + parse_stored_key_policy_string_list( + Some("[\"openai:chat\"]".to_string()), + "provider_api_keys.api_formats", + ) + .expect("valid key policy should parse"), + Some(vec!["openai:chat".to_string()]) + ); + } + #[test] fn codex_auth_sql_allows_live_for_oauth_keys() { let mut builder = sqlx::QueryBuilder::::new("SELECT 1 WHERE 1 = 1"); diff --git a/crates/aether-data/adapters/sqlite/src/candidates.rs b/crates/aether-data/adapters/sqlite/src/candidates.rs index 065ea216f..c5b9bbc1f 100644 --- a/crates/aether-data/adapters/sqlite/src/candidates.rs +++ b/crates/aether-data/adapters/sqlite/src/candidates.rs @@ -220,8 +220,9 @@ impl RequestCandidateReadRepository for SqliteRequestCandidateRepository { impl RequestCandidateWriteRepository for SqliteRequestCandidateRepository { async fn upsert( &self, - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, ) -> Result { + candidate.sanitize_for_persistence(); candidate.validate()?; let mut tx = self.pool.begin().await.map_sql_err()?; match upsert_candidate_in_transaction(&mut tx, candidate).await { @@ -238,12 +239,13 @@ impl RequestCandidateWriteRepository for SqliteRequestCandidateRepository { async fn upsert_many( &self, - candidates: Vec, + mut candidates: Vec, ) -> Result { if candidates.is_empty() { return Ok(0); } - for candidate in &candidates { + for candidate in &mut candidates { + candidate.sanitize_for_persistence(); candidate.validate()?; } @@ -428,30 +430,8 @@ ON CONFLICT(request_id, candidate_index, retry_index) DO UPDATE SET THEN request_candidates.status_code ELSE COALESCE(excluded.status_code, request_candidates.status_code) END, - error_type = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND excluded.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_type - WHEN request_candidates.status = 'pending' - AND excluded.status IN ('available', 'unused') - THEN request_candidates.error_type - WHEN request_candidates.status = 'streaming' - AND excluded.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_type - ELSE COALESCE(excluded.error_type, request_candidates.error_type) - END, - error_message = CASE - WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') - AND excluded.status IN ('available', 'unused', 'pending', 'streaming') - THEN request_candidates.error_message - WHEN request_candidates.status = 'pending' - AND excluded.status IN ('available', 'unused') - THEN request_candidates.error_message - WHEN request_candidates.status = 'streaming' - AND excluded.status IN ('available', 'unused', 'pending') - THEN request_candidates.error_message - ELSE COALESCE(excluded.error_message, request_candidates.error_message) - END, + error_type = excluded.error_type, + error_message = NULL, latency_ms = CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') AND excluded.status IN ('available', 'unused', 'pending', 'streaming') @@ -523,9 +503,10 @@ ON CONFLICT(request_id, candidate_index, retry_index) DO UPDATE SET } fn merge_candidate( - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, existing: Option, ) -> Result { + candidate.sanitize_for_persistence(); let preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| { request_candidate_lifecycle_would_regress(value.status, candidate.status) }); @@ -537,15 +518,11 @@ fn merge_candidate( } else { candidate.status }; - let created_at_unix_ms = candidate - .created_at_unix_ms + let created_at_unix_ms = existing + .as_ref() + .map(|value| value.created_at_unix_ms) .filter(|value| *value > 1000) - .or_else(|| { - existing - .as_ref() - .map(|value| value.created_at_unix_ms) - .filter(|value| *value > 1000) - }) + .or_else(|| candidate.created_at_unix_ms.filter(|value| *value > 1000)) .or(candidate.started_at_unix_ms) .or(candidate.finished_at_unix_ms) .unwrap_or_else(current_unix_ms); @@ -560,35 +537,36 @@ fn merge_candidate( StoredRequestCandidate::new( id, candidate.request_id, - candidate - .user_id - .or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())), - candidate - .api_key_id - .or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())), - candidate - .username - .or_else(|| existing.as_ref().and_then(|value| value.username.clone())), - candidate.api_key_name.or_else(|| { - existing - .as_ref() - .and_then(|value| value.api_key_name.clone()) - }), + existing + .as_ref() + .and_then(|value| value.user_id.clone()) + .or(candidate.user_id), + existing + .as_ref() + .and_then(|value| value.api_key_id.clone()) + .or(candidate.api_key_id), + existing + .as_ref() + .and_then(|value| value.username.clone()) + .or(candidate.username), + existing + .as_ref() + .and_then(|value| value.api_key_name.clone()) + .or(candidate.api_key_name), to_i32(candidate.candidate_index)?, to_i32(candidate.retry_index)?, - candidate.provider_id.or_else(|| { - existing - .as_ref() - .and_then(|value| value.provider_id.clone()) - }), - candidate.endpoint_id.or_else(|| { - existing - .as_ref() - .and_then(|value| value.endpoint_id.clone()) - }), - candidate - .key_id - .or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())), + existing + .as_ref() + .and_then(|value| value.provider_id.clone()) + .or(candidate.provider_id), + existing + .as_ref() + .and_then(|value| value.endpoint_id.clone()) + .or(candidate.endpoint_id), + existing + .as_ref() + .and_then(|value| value.key_id.clone()) + .or(candidate.key_id), merged_status, candidate.skip_reason.or_else(|| { existing @@ -616,17 +594,7 @@ fn merge_candidate( .error_type .or_else(|| existing.as_ref().and_then(|value| value.error_type.clone())) }, - if preserve_existing_lifecycle { - existing - .as_ref() - .and_then(|value| value.error_message.clone()) - } else { - candidate.error_message.or_else(|| { - existing - .as_ref() - .and_then(|value| value.error_message.clone()) - }) - }, + None, if preserve_existing_lifecycle { match existing.as_ref().and_then(|value| value.latency_ms) { Some(value) => Some(to_i32_u64(value)?), @@ -656,9 +624,10 @@ fn merge_candidate( .and_then(|value| value.required_capabilities.clone()) }), u64_to_i64(created_at_unix_ms, "request candidate created_at")?, - candidate - .started_at_unix_ms - .or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms)) + existing + .as_ref() + .and_then(|value| value.started_at_unix_ms) + .or(candidate.started_at_unix_ms) .map(|value| u64_to_i64(value, "request candidate started_at")) .transpose()?, if preserve_existing_lifecycle { @@ -889,45 +858,147 @@ mod tests { .await .expect("sqlite migrations should run"); - let repository = SqliteRequestCandidateRepository::new(pool); + let repository = SqliteRequestCandidateRepository::new(pool.clone()); let created = repository .upsert(sample_upsert( "candidate-1", RequestCandidateStatus::Pending, - Some(json!({"a": 1})), + Some(json!({"gateway_execution_runtime": true})), 1_000_000, )) .await .expect("candidate should insert"); assert_eq!(created.request_id, "request-1"); + sqlx::query( + "UPDATE request_candidates SET skip_reason = ?, error_type = ?, error_message = ?, extra_data = ?, required_capabilities = ? WHERE request_id = ?", + ) + .bind("legacy skip reason with tenant-secret") + .bind("legacy_error_type_with_token") + .bind("Bearer legacy-secret") + .bind(r#"{"gateway_execution_runtime":true,"request_body":{"password":"secret"}}"#) + .bind(r#"{"streaming":true,"internal_capability":"secret"}"#) + .bind("request-1") + .execute(&pool) + .await + .expect("legacy diagnostics should be injected for the conflict test"); let updated = repository .upsert(sample_upsert( "candidate-replacement", RequestCandidateStatus::Success, - Some(json!({"b": 2})), + Some(json!({"stream_completed": true})), 1_000_500, )) .await .expect("candidate should update"); assert_eq!(updated.id, "candidate-1"); - assert_eq!(updated.extra_data, Some(json!({"a": 1, "b": 2}))); - - let late_streaming = repository - .upsert(sample_upsert( - "candidate-late-streaming", - RequestCandidateStatus::Streaming, - Some(json!({"late": true})), - 1_000_250, - )) - .await - .expect("late streaming candidate should not regress terminal status"); - assert_eq!(late_streaming.id, "candidate-1"); - assert_eq!(late_streaming.status, RequestCandidateStatus::Success); - assert_eq!(late_streaming.finished_at_unix_ms, Some(1_000_502)); assert_eq!( - late_streaming.extra_data, - Some(json!({"a": 1, "b": 2, "late": true})) + updated.extra_data, + Some(json!({ + "gateway_execution_runtime": true, + "stream_completed": true + })) + ); + let raw = sqlx::query( + "SELECT skip_reason, error_type, error_message, extra_data, required_capabilities FROM request_candidates WHERE request_id = ?", + ) + .bind("request-1") + .fetch_one(&pool) + .await + .expect("raw candidate diagnostics should load"); + assert!( + sqlx::Row::try_get::, _>(&raw, "error_message") + .expect("error_message should decode") + .is_none() + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "skip_reason") + .expect("skip_reason should decode") + .as_deref(), + Some("unclassified_skip") + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "error_type") + .expect("error_type should decode") + .as_deref(), + Some("unclassified_error") + ); + let raw_extra = sqlx::Row::try_get::, _>(&raw, "extra_data") + .expect("extra_data should decode") + .and_then(|value| serde_json::from_str::(&value).ok()); + assert_eq!(raw_extra, updated.extra_data); + let raw_capabilities = + sqlx::Row::try_get::, _>(&raw, "required_capabilities") + .expect("required_capabilities should decode") + .and_then(|value| serde_json::from_str::(&value).ok()); + assert_eq!(raw_capabilities, Some(json!({"streaming": true}))); + + sqlx::query( + "UPDATE request_candidates SET skip_reason = ?, error_type = ? WHERE request_id = ?", + ) + .bind("pool_cooldown") + .bind(" FirstByteTimeout ") + .bind("request-1") + .execute(&pool) + .await + .expect("known legacy diagnostics should be injected for the regression test"); + let mut late = sample_upsert( + "candidate-late-terminal", + RequestCandidateStatus::Failed, + Some(json!({"cache_1h": true})), + 1_000_250, + ); + late.user_id = Some("attacker-user".to_string()); + late.api_key_id = Some("attacker-api-key".to_string()); + late.username = Some("mallory".to_string()); + late.api_key_name = Some("attacker-key".to_string()); + late.provider_id = Some("attacker-provider".to_string()); + late.endpoint_id = Some("attacker-endpoint".to_string()); + late.key_id = Some("attacker-provider-key".to_string()); + late.error_type = Some("upstream5xx".to_string()); + let late_terminal = repository + .upsert(late) + .await + .expect("late terminal candidate should not replace the first terminal fact"); + assert_eq!(late_terminal.id, "candidate-1"); + assert_eq!(late_terminal.status, RequestCandidateStatus::Success); + assert_eq!(late_terminal.user_id.as_deref(), Some("user-1")); + assert_eq!(late_terminal.api_key_id.as_deref(), Some("key-1")); + assert_eq!(late_terminal.provider_id.as_deref(), Some("provider-1")); + assert_eq!(late_terminal.endpoint_id.as_deref(), Some("endpoint-1")); + assert_eq!(late_terminal.key_id.as_deref(), Some("provider-key-1")); + assert_eq!(late_terminal.skip_reason.as_deref(), Some("pool_cooldown")); + assert_eq!( + late_terminal.error_type.as_deref(), + Some("first_byte_timeout") + ); + assert_eq!(late_terminal.finished_at_unix_ms, Some(1_000_502)); + assert_eq!( + late_terminal.extra_data, + Some(json!({ + "cache_1h": true, + "gateway_execution_runtime": true, + "stream_completed": true + })) + ); + let raw = sqlx::query( + "SELECT skip_reason, error_type FROM request_candidates WHERE request_id = ?", + ) + .bind("request-1") + .fetch_one(&pool) + .await + .expect("raw candidate classifications should load"); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "skip_reason") + .expect("skip_reason should decode") + .as_deref(), + Some("pool_cooldown") + ); + assert_eq!( + sqlx::Row::try_get::, _>(&raw, "error_type") + .expect("error_type should decode") + .as_deref(), + Some("first_byte_timeout") ); assert_eq!( @@ -993,7 +1064,7 @@ mod tests { let mut initial = sample_upsert( "initial", RequestCandidateStatus::Pending, - Some(json!({"initial": true})), + Some(json!({"gateway_execution_runtime": true})), 3_000_000, ); initial.request_id = request_id.clone(); @@ -1015,12 +1086,23 @@ mod tests { } else { RequestCandidateStatus::Streaming }; - let mut extra_data = serde_json::Map::new(); - extra_data.insert(format!("writer_{writer}"), json!(writer)); + let extra_data = match writer { + 0 => json!({"stream_completed": true}), + 1 => json!({"cache_1h": true}), + 2 => json!({"first_byte_time_ms": 2}), + 3 => json!({"pool_key_index": 3}), + 4 => json!({"priority_slot": 4}), + 5 => json!({"ranking_index": 5}), + 6 => json!({"phase": "provider_request"}), + 7 => json!({"provider_api_format": "openai:responses"}), + 8 => json!({"client_api_format": "openai:chat"}), + 9 => json!({"execution_strategy": "local_cross_format"}), + _ => unreachable!("writer index is bounded by WRITERS"), + }; let mut candidate = sample_upsert( format!("writer-{writer}").as_str(), status, - Some(serde_json::Value::Object(extra_data)), + Some(extra_data), 3_100_000 + u64::try_from(writer).expect("writer index should fit") * 10, ); candidate.request_id = request_id; @@ -1048,18 +1130,22 @@ mod tests { assert_eq!(candidate.status, RequestCandidateStatus::Success); assert_eq!(candidate.latency_ms, Some(123)); assert_eq!(candidate.finished_at_unix_ms, Some(3_100_002)); - let extra_data = candidate - .extra_data - .as_ref() - .and_then(serde_json::Value::as_object) - .expect("merged extra data should be an object"); - assert_eq!(extra_data.get("initial"), Some(&json!(true))); - for writer in 0..WRITERS { - assert_eq!( - extra_data.get(format!("writer_{writer}").as_str()), - Some(&json!(writer)) - ); - } + assert_eq!( + candidate.extra_data, + Some(json!({ + "cache_1h": true, + "client_api_format": "openai:chat", + "execution_strategy": "local_cross_format", + "first_byte_time_ms": 2, + "gateway_execution_runtime": true, + "phase": "provider_request", + "pool_key_index": 3, + "priority_slot": 4, + "provider_api_format": "openai:responses", + "ranking_index": 5, + "stream_completed": true + })) + ); drop(repository); pool.close().await; @@ -1084,14 +1170,14 @@ mod tests { let mut pending = sample_upsert( "batch-first", RequestCandidateStatus::Pending, - Some(json!({"pending": true})), + Some(json!({"gateway_execution_runtime": true})), 4_000_000, ); pending.request_id = request_id.to_string(); let mut streaming = sample_upsert( "batch-second", RequestCandidateStatus::Streaming, - Some(json!({"streaming": true})), + Some(json!({"stream_completed": true})), 4_000_100, ); streaming.request_id = request_id.to_string(); @@ -1099,7 +1185,7 @@ mod tests { let mut success = sample_upsert( "batch-third", RequestCandidateStatus::Success, - Some(json!({"success": true})), + Some(json!({"cache_1h": true})), 4_000_200, ); success.request_id = request_id.to_string(); @@ -1107,7 +1193,7 @@ mod tests { let mut late_pending = sample_upsert( "batch-fourth", RequestCandidateStatus::Pending, - Some(json!({"late": true})), + Some(json!({"first_byte_time_ms": 42})), 4_000_300, ); late_pending.request_id = request_id.to_string(); @@ -1136,10 +1222,10 @@ mod tests { assert_eq!( candidate.extra_data, Some(json!({ - "pending": true, - "streaming": true, - "success": true, - "late": true + "cache_1h": true, + "first_byte_time_ms": 42, + "gateway_execution_runtime": true, + "stream_completed": true })) ); diff --git a/crates/aether-data/adapters/sqlite/src/gemini_file_mappings.rs b/crates/aether-data/adapters/sqlite/src/gemini_file_mappings.rs index 9f787c5bc..71a9a2fee 100644 --- a/crates/aether-data/adapters/sqlite/src/gemini_file_mappings.rs +++ b/crates/aether-data/adapters/sqlite/src/gemini_file_mappings.rs @@ -63,6 +63,80 @@ LIMIT 1 row.as_ref().map(map_row).transpose() } + async fn find_active_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at AS created_at_unix_ms, + expires_at AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE file_name = ? AND user_id = ? AND expires_at > ? +LIMIT 1 +"#, + ) + .bind(file_name) + .bind(user_id) + .bind(i64_from_u64( + now_unix_secs, + "gemini_file_mappings.owner_read_now", + )?) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + + row.as_ref().map(map_row).transpose() + } + + async fn find_active_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT + id, + file_name, + key_id, + user_id, + display_name, + mime_type, + source_hash, + created_at AS created_at_unix_ms, + expires_at AS expires_at_unix_secs +FROM gemini_file_mappings +WHERE file_name = ? AND key_id = ? AND user_id = ? AND expires_at > ? +LIMIT 1 +"#, + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .bind(i64_from_u64( + now_unix_secs, + "gemini_file_mappings.owner_read_now", + )?) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + + row.as_ref().map(map_row).transpose() + } + async fn list_mappings( &self, query: &GeminiFileMappingListQuery, @@ -185,6 +259,49 @@ ON CONFLICT(file_name) DO UPDATE SET self.reload_by_file_name(&record.file_name).await } + async fn upsert_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, DataLayerError> { + record.validate()?; + let rows_affected = sqlx::query( + r#" +INSERT INTO gemini_file_mappings ( + id, file_name, key_id, user_id, display_name, mime_type, source_hash, + created_at, expires_at +) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) +ON CONFLICT(file_name) DO UPDATE SET + display_name = excluded.display_name, + mime_type = excluded.mime_type, + source_hash = excluded.source_hash, + expires_at = excluded.expires_at +WHERE gemini_file_mappings.key_id = excluded.key_id + AND gemini_file_mappings.user_id IS excluded.user_id +"#, + ) + .bind(&record.id) + .bind(&record.file_name) + .bind(&record.key_id) + .bind(&record.user_id) + .bind(&record.display_name) + .bind(&record.mime_type) + .bind(&record.source_hash) + .bind(current_unix_secs() as i64) + .bind(i64_from_u64( + record.expires_at_unix_secs, + "gemini_file_mappings.expires_at", + )?) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + if rows_affected == 0 { + return Ok(None); + } + self.reload_by_file_name(&record.file_name).await.map(Some) + } + async fn delete_by_file_name(&self, file_name: &str) -> Result { let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?") .bind(file_name) @@ -195,6 +312,41 @@ ON CONFLICT(file_name) DO UPDATE SET Ok(rows_affected > 0) } + async fn delete_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + let rows_affected = + sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ? AND user_id = ?") + .bind(file_name) + .bind(user_id) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected > 0) + } + + async fn delete_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + let rows_affected = sqlx::query( + "DELETE FROM gemini_file_mappings WHERE file_name = ? AND key_id = ? AND user_id = ?", + ) + .bind(file_name) + .bind(key_id) + .bind(user_id) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected > 0) + } + async fn delete_by_id( &self, mapping_id: &str, @@ -282,6 +434,11 @@ fn apply_list_filters( where_clause: &mut WhereClause, query: &GeminiFileMappingListQuery, ) { + if let Some(user_id) = query.user_id.as_deref() { + where_clause.push_next(builder); + builder.push("user_id = "); + builder.push_bind(user_id.to_string()); + } if !query.include_expired { where_clause.push_next(builder); builder.push("expires_at > "); @@ -394,8 +551,77 @@ mod tests { assert_eq!(updated.id, "mapping-1"); assert_eq!(updated.key_id, "key-2"); + let guarded_reassignment = repository + .upsert_if_owner_matches(UpsertGeminiFileMappingRecord { + id: "mapping-attacker".to_string(), + file_name: "files/example.png".to_string(), + key_id: "key-attacker".to_string(), + user_id: Some("user-attacker".to_string()), + display_name: Some("Attacker".to_string()), + mime_type: Some("application/octet-stream".to_string()), + source_hash: Some("hash-attacker".to_string()), + expires_at_unix_secs: 900, + }) + .await + .expect("guarded reassignment should run"); + assert!(guarded_reassignment.is_none()); + let after_reassignment = repository + .find_by_file_name("files/example.png") + .await + .expect("mapping should read") + .expect("mapping should remain"); + assert_eq!(after_reassignment.key_id, "key-2"); + assert_eq!(after_reassignment.user_id.as_deref(), Some("user-2")); + + let guarded_refresh = repository + .upsert_if_owner_matches(UpsertGeminiFileMappingRecord { + id: "mapping-refresh".to_string(), + file_name: "files/example.png".to_string(), + key_id: "key-2".to_string(), + user_id: Some("user-2".to_string()), + display_name: Some("Updated".to_string()), + mime_type: Some("image/jpeg".to_string()), + source_hash: Some("hash-refreshed".to_string()), + expires_at_unix_secs: 500, + }) + .await + .expect("same-owner refresh should run"); + assert!(guarded_refresh.is_some()); + + assert!(repository + .find_active_by_file_name_for_user("files/example.png", "user-2", 400) + .await + .expect("owner-scoped read should run") + .is_some()); + assert!(repository + .find_active_by_file_name_for_user("files/example.png", "user-1", 400) + .await + .expect("foreign owner read should run") + .is_none()); + assert!(repository + .find_active_by_file_name_for_owner("files/example.png", "key-2", "user-2", 400,) + .await + .expect("provider owner read should run") + .is_some()); + assert!(repository + .find_active_by_file_name_for_owner("files/example.png", "key-attacker", "user-2", 400,) + .await + .expect("foreign provider read should run") + .is_none()); + assert!(repository + .find_active_by_file_name_for_user("files/example.png", "user-2", 500) + .await + .expect("expired owner read should run") + .is_none()); + + assert!(!repository + .delete_by_file_name_for_owner("files/example.png", "key-attacker", "user-2") + .await + .expect("wrong-key owner delete should run")); + let page = repository .list_mappings(&GeminiFileMappingListQuery { + user_id: Some("user-2".to_string()), include_expired: false, search: Some("updated".to_string()), offset: 0, @@ -407,6 +633,16 @@ mod tests { assert_eq!(page.total, 1); assert_eq!(page.items[0].file_name, "files/example.png"); + assert!(!repository + .delete_by_file_name_for_user("files/example.png", "user-1") + .await + .expect("non-owner delete should run")); + assert!(repository + .find_by_file_name("files/example.png") + .await + .expect("mapping should remain") + .is_some()); + let stats = repository .summarize_mappings(400) .await diff --git a/crates/aether-data/adapters/sqlite/src/management_tokens.rs b/crates/aether-data/adapters/sqlite/src/management_tokens.rs index 89e9f6e9b..e5b212ea8 100644 --- a/crates/aether-data/adapters/sqlite/src/management_tokens.rs +++ b/crates/aether-data/adapters/sqlite/src/management_tokens.rs @@ -2,10 +2,10 @@ use async_trait::async_trait; use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser, - UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, UpdateManagementTokenRecord, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause}; @@ -38,8 +38,147 @@ impl SqliteManagementTokenRepository { .map_sql_err()?; row.as_ref().map(map_token_row).transpose() } + + async fn get_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + let mut builder = QueryBuilder::::new(TOKEN_COLUMNS); + let mut where_clause = WhereClause::new(); + push_eq(&mut builder, &mut where_clause, "id", token_id.to_string()); + push_optional_eq( + &mut builder, + &mut where_clause, + "user_id", + expected_user_id.map(ToOwned::to_owned), + ); + push_limit(&mut builder, 1); + let row = builder + .build() + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_token_row).transpose() + } + + async fn update_management_token_scoped( + &self, + record: &UpdateManagementTokenRecord, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + record.validate()?; + let allowed_ips = json_to_string(record.allowed_ips.as_ref())?; + let permissions = json_to_string(record.permissions.as_ref())?; + let now = now_unix_secs(); + + sqlx::query(UPDATE_MANAGEMENT_TOKEN_SQL) + .bind(record.name.as_deref()) + .bind(record.clear_description) + .bind(record.description.as_deref()) + .bind(record.clear_allowed_ips) + .bind(allowed_ips) + .bind(permissions) + .bind(record.clear_expires_at) + .bind( + record + .expires_at_unix_secs + .and_then(|value| i64::try_from(value).ok()), + ) + .bind(record.is_active) + .bind(now as i64) + .bind(&record.token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_err(|err| map_sqlite_write_error(err, record.name.as_deref()))?; + self.get_token_scoped(&record.token_id, expected_user_id) + .await + } + + async fn delete_management_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> Result { + let result = sqlx::query( + "DELETE FROM management_tokens WHERE id = ? AND (? IS NULL OR user_id = ?)", + ) + .bind(token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() > 0) + } + + async fn set_management_token_active_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + is_active: bool, + ) -> Result, DataLayerError> { + let result = sqlx::query( + "UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ? AND (? IS NULL OR user_id = ?)", + ) + .bind(is_active) + .bind(now_unix_secs() as i64) + .bind(token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.get_token_scoped(token_id, expected_user_id).await + } + + async fn regenerate_management_token_secret_scoped( + &self, + mutation: &RegenerateManagementTokenSecret, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + mutation.validate()?; + let result = sqlx::query( + r#" +UPDATE management_tokens +SET token_hash = ?, token_prefix = ?, updated_at = ? +WHERE id = ? AND (? IS NULL OR user_id = ?) +"#, + ) + .bind(&mutation.token_hash) + .bind(mutation.token_prefix.as_deref()) + .bind(now_unix_secs() as i64) + .bind(&mutation.token_id) + .bind(expected_user_id) + .bind(expected_user_id) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); + } + self.get_token_scoped(&mutation.token_id, expected_user_id) + .await + } } +const UPDATE_MANAGEMENT_TOKEN_SQL: &str = r#" +UPDATE management_tokens +SET name = COALESCE(?, name), + description = CASE WHEN ? THEN NULL ELSE COALESCE(?, description) END, + allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END, + permissions = COALESCE(?, permissions), + expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END, + is_active = COALESCE(?, is_active), + updated_at = ? +WHERE id = ? AND (? IS NULL OR user_id = ?) +"#; + const TOKEN_COLUMNS: &str = r#" SELECT id, @@ -59,6 +198,37 @@ SELECT FROM management_tokens "#; +const LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL: &str = r#" +SELECT id +FROM users +WHERE id = ? + AND is_active = 1 + AND is_deleted = 0 + AND LOWER(role) = 'admin' + AND security_version = ? +"#; + +const LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL: &str = r#" +SELECT + id, + user_id, + token_hash, + name, + description, + token_prefix, + allowed_ips, + permissions, + expires_at AS expires_at_unix_secs, + last_used_at AS last_used_at_unix_secs, + last_used_ip, + COALESCE(usage_count, 0) AS usage_count, + is_active, + created_at AS created_at_unix_ms, + updated_at AS updated_at_unix_secs +FROM management_tokens +WHERE id = ? +"#; + const TOKEN_WITH_USER_COLUMNS: &str = r#" SELECT mt.id, @@ -220,71 +390,29 @@ INSERT INTO management_tokens ( &self, record: &UpdateManagementTokenRecord, ) -> Result, DataLayerError> { - record.validate()?; - let current = self.get_token(&record.token_id).await?; - let Some(current) = current else { - return Ok(None); - }; - let name = record.name.as_deref().unwrap_or(¤t.name); - let description = if record.clear_description { - None - } else { - record - .description - .as_deref() - .or(current.description.as_deref()) - }; - let allowed_ips = if record.clear_allowed_ips { - None - } else { - record.allowed_ips.as_ref().or(current.allowed_ips.as_ref()) - }; - let permissions = record.permissions.as_ref().or(current.permissions.as_ref()); - let expires_at = if record.clear_expires_at { - None - } else { - record.expires_at_unix_secs.or(current.expires_at_unix_secs) - }; - let is_active = record.is_active.unwrap_or(current.is_active); - let now = now_unix_secs(); + self.update_management_token_scoped(record, None).await + } - let result = sqlx::query( - r#" -UPDATE management_tokens -SET name = ?, - description = ?, - allowed_ips = ?, - permissions = ?, - expires_at = ?, - is_active = ?, - updated_at = ? -WHERE id = ? -"#, - ) - .bind(name) - .bind(description) - .bind(json_to_string(allowed_ips)?) - .bind(json_to_string(permissions)?) - .bind(expires_at.and_then(|value| i64::try_from(value).ok())) - .bind(is_active) - .bind(now as i64) - .bind(&record.token_id) - .execute(&self.pool) - .await - .map_err(|err| map_sqlite_write_error(err, record.name.as_deref()))?; - if result.rows_affected() == 0 { - return Ok(None); - } - self.get_token(&record.token_id).await + async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, DataLayerError> { + self.update_management_token_scoped(record, Some(user_id)) + .await } async fn delete_management_token(&self, token_id: &str) -> Result { - let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?") - .bind(token_id) - .execute(&self.pool) + self.delete_management_token_scoped(token_id, None).await + } + + async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + self.delete_management_token_scoped(token_id, Some(user_id)) .await - .map_sql_err()?; - Ok(result.rows_affected() > 0) } async fn set_management_token_active( @@ -292,43 +420,145 @@ WHERE id = ? token_id: &str, is_active: bool, ) -> Result, DataLayerError> { - let result = - sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?") - .bind(is_active) - .bind(now_unix_secs() as i64) - .bind(token_id) - .execute(&self.pool) + self.set_management_token_active_scoped(token_id, None, is_active) + .await + } + + async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + self.set_management_token_active_scoped(token_id, Some(user_id), is_active) + .await + } + + async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + // BEGIN IMMEDIATE prevents a concurrent user downgrade or token mutation between the + // canonical pre-state reads and the activation write. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let eligible_user = + sqlx::query_scalar::<_, String>(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL) + .bind(&mutation.expected_token.user_id) + .bind(mutation.expected_user_security_version) + .fetch_optional(&mut *tx) .await .map_sql_err()?; - if result.rows_affected() == 0 { - return Ok(None); + if eligible_user.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); } - self.get_token(token_id).await + + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_sql_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE management_tokens +SET is_active = 1, updated_at = ? +WHERE id = ? + AND token_hash = ? + AND is_active = 0 + AND (expires_at IS NULL OR expires_at > ?) +"#, + ) + .bind(now_unix_secs() as i64) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .bind(i64::try_from(mutation.now_unix_secs).unwrap_or(i64::MAX)) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL) + .bind(&mutation.expected_token.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let snapshot_matches = match locked.as_ref() { + Some(row) => { + let token_hash: String = row.try_get("token_hash").map_sql_err()?; + let token = map_token_row(row)?; + mutation.matches_locked_token_snapshot(&token, &token_hash) + } + None => false, + }; + if !snapshot_matches { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let result = sqlx::query( + "DELETE FROM management_tokens WHERE id = ? AND token_hash = ? AND is_active = 0", + ) + .bind(&mutation.expected_token.id) + .bind(&mutation.token_hash) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) } async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, ) -> Result, DataLayerError> { - mutation.validate()?; - let result = sqlx::query( - r#" -UPDATE management_tokens -SET token_hash = ?, token_prefix = ?, updated_at = ? -WHERE id = ? -"#, - ) - .bind(&mutation.token_hash) - .bind(mutation.token_prefix.as_deref()) - .bind(now_unix_secs() as i64) - .bind(&mutation.token_id) - .execute(&self.pool) - .await - .map_sql_err()?; - if result.rows_affected() == 0 { - return Ok(None); - } - self.get_token(&mutation.token_id).await + self.regenerate_management_token_secret_scoped(mutation, None) + .await + } + + async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, DataLayerError> { + self.regenerate_management_token_secret_scoped(mutation, Some(user_id)) + .await } async fn record_management_token_usage( @@ -365,8 +595,18 @@ fn now_unix_secs() -> u64 { chrono::Utc::now().timestamp().max(0) as u64 } -fn optional_unix_secs(value: Option) -> Option { - value.and_then(|value| u64::try_from(value).ok()) +fn non_negative_u64(value: i64, field_name: &str) -> Result { + u64::try_from(value).map_err(|_| { + DataLayerError::UnexpectedValue(format!( + "management_tokens.{field_name} must not be negative" + )) + }) +} + +fn optional_unix_secs(value: Option, field_name: &str) -> Result, DataLayerError> { + value + .map(|value| non_negative_u64(value, field_name)) + .transpose() } fn json_to_string(value: Option<&serde_json::Value>) -> Result, DataLayerError> { @@ -418,15 +658,30 @@ fn map_token_row(row: &SqliteRow) -> Result("usage_count").map_sql_err()?).unwrap_or(0), + non_negative_u64( + row.try_get::("usage_count").map_sql_err()?, + "usage_count", + )?, row.try_get("is_active").map_sql_err()?, ) .with_timestamps( - optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?), - optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?), + optional_unix_secs( + row.try_get("created_at_unix_ms").map_sql_err()?, + "created_at", + )?, + optional_unix_secs( + row.try_get("updated_at_unix_secs").map_sql_err()?, + "updated_at", + )?, )) } @@ -452,14 +707,41 @@ fn map_token_with_user_row( #[cfg(test)] mod tests { - use super::SqliteManagementTokenRepository; + use super::{ + non_negative_u64, optional_unix_secs, SqliteManagementTokenRepository, + UPDATE_MANAGEMENT_TOKEN_SQL, + }; use crate::run_migrations; use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, - StoredManagementTokenUserSummary, UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, + RegenerateManagementTokenSecret, StoredManagementTokenUserSummary, + UpdateManagementTokenRecord, }; + #[test] + fn sqlite_management_token_mapping_rejects_negative_integer_state() { + assert!(optional_unix_secs(Some(-1), "expires_at").is_err()); + assert_eq!( + optional_unix_secs(None, "expires_at").expect("SQL NULL should remain optional"), + None + ); + assert!(non_negative_u64(-1, "usage_count").is_err()); + } + + #[test] + fn sqlite_management_token_updates_patch_only_explicit_fields() { + for clause in [ + "name = COALESCE(?, name)", + "allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END", + "permissions = COALESCE(?, permissions)", + "expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END", + "is_active = COALESCE(?, is_active)", + ] { + assert!(UPDATE_MANAGEMENT_TOKEN_SQL.contains(clause)); + } + } + #[tokio::test] async fn sqlite_repository_round_trips_management_tokens() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -529,6 +811,103 @@ VALUES ('user-1', 'user-1@example.com', 'user-1', 'admin', 1, 1, 1) .expect("token should exist"); assert_eq!(by_hash.user.username, "user-1"); + let pending_install = repository + .create_management_token(&CreateManagementTokenRecord { + id: "token-install".to_string(), + user_id: "user-1".to_string(), + user: by_hash.user.clone(), + token_hash: "hash-install".to_string(), + token_prefix: Some("ae_install".to_string()), + name: "pending install".to_string(), + description: None, + allowed_ips: Some(serde_json::json!(["127.0.0.1"])), + permissions: Some(serde_json::json!(["admin:proxy_nodes:write"])), + expires_at_unix_secs: Some(1_800_000_000), + is_active: false, + }) + .await + .expect("pending install token should create"); + let activation = ActivateManagementTokenIfMatches { + expected_token: pending_install, + token_hash: "hash-install".to_string(), + expected_user_security_version: 0, + now_unix_secs: 1_700_000_000, + }; + let mut wrong_secret = activation.clone(); + wrong_secret.token_hash = "hash-other".to_string(); + assert!(!repository + .activate_management_token_if_matches(&wrong_secret) + .await + .expect("wrong-secret activation should execute")); + let mut stale_snapshot = activation.clone(); + stale_snapshot.expected_token.description = Some("changed snapshot".to_string()); + assert!(!repository + .activate_management_token_if_matches(&stale_snapshot) + .await + .expect("stale-snapshot activation should execute")); + assert!(repository + .activate_management_token_if_matches(&activation) + .await + .expect("matching activation should execute")); + assert!(!repository + .activate_management_token_if_matches(&activation) + .await + .expect("already-active token must not activate again")); + + let cross_owner_update = UpdateManagementTokenRecord { + token_id: "token-1".to_string(), + name: Some("hijacked".to_string()), + description: None, + clear_description: false, + allowed_ips: None, + clear_allowed_ips: false, + permissions: None, + expires_at_unix_secs: None, + clear_expires_at: false, + is_active: None, + }; + assert!(repository + .update_management_token_for_user(&cross_owner_update, "user-2") + .await + .expect("owner-scoped update should execute") + .is_none()); + assert!(repository + .set_management_token_active_for_user("token-1", "user-2", false) + .await + .expect("owner-scoped toggle should execute") + .is_none()); + assert!(repository + .regenerate_management_token_secret_for_user( + &RegenerateManagementTokenSecret { + token_id: "token-1".to_string(), + token_hash: "hash-hijacked".to_string(), + token_prefix: Some("ae_hijacked".to_string()), + }, + "user-2", + ) + .await + .expect("owner-scoped regeneration should execute") + .is_none()); + assert!(!repository + .delete_management_token_for_user("token-1", "user-2") + .await + .expect("owner-scoped delete should execute")); + assert_eq!( + repository + .get_management_token_with_user_by_hash("hash-1") + .await + .expect("original hash lookup should succeed") + .expect("token should remain") + .token + .name, + "primary" + ); + assert!(repository + .get_management_token_with_user_by_hash("hash-hijacked") + .await + .expect("replacement hash lookup should succeed") + .is_none()); + let updated = repository .update_management_token(&UpdateManagementTokenRecord { token_id: "token-1".to_string(), @@ -589,5 +968,145 @@ VALUES ('user-1', 'user-1@example.com', 'user-1', 'admin', 1, 1, 1) .delete_management_token("token-1") .await .expect("delete should succeed")); + assert!(repository + .delete_management_token("token-install") + .await + .expect("install token delete should succeed")); + } + + #[tokio::test] + async fn sqlite_install_activation_rejects_changed_admin_identity_snapshot() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO users ( + id, email, username, role, is_active, is_deleted, security_version, created_at, updated_at +) +VALUES ('admin-install', 'admin@example.com', 'admin-install', 'admin', 1, 0, 7, 1, 1) +"#, + ) + .execute(&pool) + .await + .expect("admin user should insert"); + + let repository = SqliteManagementTokenRepository::new(pool.clone()); + let user = StoredManagementTokenUserSummary::new( + "admin-install".to_string(), + Some("admin@example.com".to_string()), + "admin-install".to_string(), + "admin".to_string(), + ) + .expect("user summary should build"); + + async fn create_pending( + repository: &SqliteManagementTokenRepository, + user: &StoredManagementTokenUserSummary, + id: &str, + ) -> aether_data_contracts::repository::management_tokens::StoredManagementToken { + repository + .create_management_token(&CreateManagementTokenRecord { + id: id.to_string(), + user_id: user.id.clone(), + user: user.clone(), + token_hash: format!("hash-{id}"), + token_prefix: Some("ae_install".to_string()), + name: id.to_string(), + description: Some("one-time tunnel install".to_string()), + allowed_ips: Some(serde_json::json!(["127.0.0.1"])), + permissions: Some(serde_json::json!(["admin:proxy_nodes:write"])), + expires_at_unix_secs: Some(1_800_000_000), + is_active: false, + }) + .await + .expect("pending install token should create") + } + + fn activation( + token: aether_data_contracts::repository::management_tokens::StoredManagementToken, + security_version: i64, + ) -> ActivateManagementTokenIfMatches { + ActivateManagementTokenIfMatches { + token_hash: format!("hash-{}", token.id), + expected_token: token, + expected_user_security_version: security_version, + now_unix_secs: 1_700_000_000, + } + } + + let role_activation = + activation(create_pending(&repository, &user, "install-role").await, 7); + sqlx::query("UPDATE users SET role = 'user' WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("role downgrade should update"); + assert!(!repository + .activate_management_token_if_matches(&role_activation) + .await + .expect("role-mismatched activation should execute")); + sqlx::query("UPDATE users SET role = 'admin' WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("admin role should restore"); + + let version_activation = activation( + create_pending(&repository, &user, "install-version").await, + 7, + ); + sqlx::query("UPDATE users SET security_version = 8 WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("security version should update"); + assert!(!repository + .activate_management_token_if_matches(&version_activation) + .await + .expect("version-mismatched activation should execute")); + + let inactive_activation = activation( + create_pending(&repository, &user, "install-inactive").await, + 8, + ); + sqlx::query("UPDATE users SET is_active = 0 WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("administrator should deactivate"); + assert!(!repository + .activate_management_token_if_matches(&inactive_activation) + .await + .expect("inactive-admin activation should execute")); + sqlx::query("UPDATE users SET is_active = 1 WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("administrator should reactivate"); + + let deleted_activation = activation( + create_pending(&repository, &user, "install-deleted").await, + 8, + ); + sqlx::query("UPDATE users SET is_deleted = 1 WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("administrator should be soft deleted"); + assert!(!repository + .activate_management_token_if_matches(&deleted_activation) + .await + .expect("deleted-admin activation should execute")); + sqlx::query("UPDATE users SET is_deleted = 0 WHERE id = 'admin-install'") + .execute(&pool) + .await + .expect("administrator deletion flag should restore"); + + let valid_activation = + activation(create_pending(&repository, &user, "install-valid").await, 8); + assert!(repository + .activate_management_token_if_matches(&valid_activation) + .await + .expect("matching administrator activation should execute")); } } diff --git a/crates/aether-data/adapters/sqlite/src/migrations.rs b/crates/aether-data/adapters/sqlite/src/migrations.rs index 89c07fd66..d19bb3950 100644 --- a/crates/aether-data/adapters/sqlite/src/migrations.rs +++ b/crates/aether-data/adapters/sqlite/src/migrations.rs @@ -110,6 +110,54 @@ mod tests { .is_empty()); } + #[tokio::test] + async fn legacy_proxy_node_inserts_receive_non_empty_tunnel_generations() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("in-memory sqlite pool"); + + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + // Simulate an older writer that does not know about tunnel_generation. + sqlx::query( + "INSERT INTO proxy_nodes (id, name, ip, port, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("legacy-node-a") + .bind("legacy node a") + .bind("127.0.0.1") + .bind(18080_i32) + .bind(1_i64) + .bind(1_i64) + .execute(&pool) + .await + .expect("legacy proxy node insert should succeed"); + sqlx::query( + "INSERT INTO proxy_nodes (id, name, ip, port, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("legacy-node-b") + .bind("legacy node b") + .bind("127.0.0.2") + .bind(18081_i32) + .bind(1_i64) + .bind(1_i64) + .execute(&pool) + .await + .expect("second legacy proxy node insert should succeed"); + + let generations = sqlx::query_scalar::<_, String>( + "SELECT tunnel_generation FROM proxy_nodes ORDER BY id", + ) + .fetch_all(&pool) + .await + .expect("proxy node generations should load"); + assert_eq!(generations.len(), 2); + assert!(generations.iter().all(|generation| !generation.is_empty())); + assert_ne!(generations[0], generations[1]); + } + #[test] fn rejects_applied_migration_versions_unknown_to_this_binary() { let version = MIGRATOR @@ -678,6 +726,85 @@ ORDER BY id ); } + #[tokio::test] + async fn ldap_singleton_migration_keeps_legacy_min_id_and_enforces_one_fixed_row() { + const LDAP_SINGLETON_VERSION: i64 = 20260831000000; + + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("in-memory sqlite pool"); + for migration in MIGRATOR + .iter() + .filter(|migration| migration.version < LDAP_SINGLETON_VERSION) + { + sqlx::raw_sql(migration.sql.as_ref()) + .execute(&pool) + .await + .unwrap_or_else(|err| panic!("migration {} should run: {err}", migration.version)); + } + + sqlx::query( + r#" +INSERT INTO ldap_configs ( + server_url, bind_dn, bind_password_encrypted, base_dn, created_at, updated_at +) VALUES + ('ldaps://first.example.com', 'cn=first', 'first-ciphertext', 'dc=first', 1, 1), + ('ldaps://second.example.com', 'cn=second', 'second-ciphertext', 'dc=second', 2, 2) +"#, + ) + .execute(&pool) + .await + .expect("legacy duplicate LDAP rows should seed"); + + let migration = MIGRATOR + .iter() + .find(|migration| migration.version == LDAP_SINGLETON_VERSION) + .expect("LDAP singleton migration should be embedded"); + sqlx::raw_sql(migration.sql.as_ref()) + .execute(&pool) + .await + .expect("LDAP singleton migration should run"); + + let surviving = sqlx::query_as::<_, (String, Option, i64)>( + "SELECT server_url, bind_password_encrypted, singleton_key FROM ldap_configs", + ) + .fetch_all(&pool) + .await + .expect("singleton LDAP row should load"); + assert_eq!( + surviving, + vec![( + "ldaps://first.example.com".to_string(), + Some("first-ciphertext".to_string()), + 1, + )] + ); + + let duplicate = sqlx::query( + r#" +INSERT INTO ldap_configs ( + server_url, bind_dn, bind_password_encrypted, base_dn, created_at, updated_at +) VALUES ('ldaps://third.example.com', 'cn=third', 'third-ciphertext', 'dc=third', 3, 3) +"#, + ) + .execute(&pool) + .await; + assert!( + duplicate.is_err(), + "a second singleton row must be rejected" + ); + + let invalid_key = sqlx::query("UPDATE ldap_configs SET singleton_key = 2") + .execute(&pool) + .await; + assert!( + invalid_key.is_err(), + "the singleton discriminator must remain fixed at one" + ); + } + #[tokio::test] async fn codex_live_permission_migration_is_scoped_and_idempotent() { const MIGRATION_VERSION: i64 = 20260821000000; diff --git a/crates/aether-data/adapters/sqlite/src/oauth_providers.rs b/crates/aether-data/adapters/sqlite/src/oauth_providers.rs index a4437363c..a49ccf36a 100644 --- a/crates/aether-data/adapters/sqlite/src/oauth_providers.rs +++ b/crates/aether-data/adapters/sqlite/src/oauth_providers.rs @@ -3,7 +3,7 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::oauth_providers::{ OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig, - UpsertOAuthProviderConfigRecord, + UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord, }; use aether_data_contracts::DataLayerError; use aether_data_query::{push_eq, push_limit, WhereClause}; @@ -101,6 +101,13 @@ WHERE users.is_active = 1 ) "#; +const COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL: &str = r#" +UPDATE oauth_providers +SET client_secret_encrypted = ? +WHERE provider_type = ? + AND client_secret_encrypted = ? +"#; + #[async_trait] impl OAuthProviderReadRepository for SqliteOAuthProviderRepository { async fn list_oauth_provider_configs( @@ -143,12 +150,51 @@ impl OAuthProviderReadRepository for SqliteOAuthProviderRepository { #[async_trait] impl OAuthProviderWriteRepository for SqliteOAuthProviderRepository { - async fn upsert_oauth_provider_config( + async fn upsert_oauth_provider_config_guarded( &self, record: &UpsertOAuthProviderConfigRecord, - ) -> Result { + ldap_exclusive: bool, + force_disable: bool, + _locked_users_snapshot: usize, + ) -> Result { record.validate()?; let now = now_unix_secs(); + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let existing_enabled: Option = + sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?") + .bind(&record.provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if !force_disable && !record.is_enabled && existing_enabled == Some(true) { + let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL) + .bind(&record.provider_type) + .bind(&record.provider_type) + .bind(ldap_exclusive) + .bind(&record.provider_type) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let affected_count = + usize::try_from(row.try_get::("locked_count").map_sql_err()?.max(0)) + .map_err(|_| { + DataLayerError::UnexpectedValue( + "oauth_providers.locked_user_count overflowed".to_string(), + ) + })?; + if affected_count > 0 { + tx.rollback().await.map_sql_err()?; + return Ok( + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count, + }, + ); + } + } sqlx::query( r#" INSERT INTO oauth_providers ( @@ -213,27 +259,70 @@ ON CONFLICT(provider_type) DO UPDATE SET .bind(now as i64) .bind(record.client_secret_encrypted.mode_name()) .bind(record.client_secret_encrypted.value()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - self.get_provider(&record.provider_type) - .await? - .ok_or_else(|| { - DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string()) - }) + let row = sqlx::query(&format!( + "{OAUTH_PROVIDER_COLUMNS} WHERE provider_type = ? LIMIT 1" + )) + .bind(&record.provider_type) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let provider = map_oauth_provider_row(&row)?; + tx.commit().await.map_sql_err()?; + Ok(UpsertOAuthProviderConfigOutcome::Upserted(provider)) } - async fn delete_oauth_provider_config( + async fn compare_and_swap_oauth_provider_client_secret( &self, provider_type: &str, + expected: &str, + replacement: &str, ) -> Result { - let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?") + let result = sqlx::query(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL) + .bind(replacement) .bind(provider_type) + .bind(expected) .execute(&self.pool) .await .map_sql_err()?; - Ok(result.rows_affected() > 0) + Ok(result.rows_affected() == 1) + } + + async fn delete_oauth_provider_config_if_unlinked( + &self, + provider_type: &str, + has_links_snapshot: bool, + ) -> Result { + if has_links_snapshot { + return Ok(false); + } + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let provider_exists: Option = + sqlx::query_scalar("SELECT provider_type FROM oauth_providers WHERE provider_type = ?") + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let result = sqlx::query( + "DELETE FROM oauth_providers WHERE provider_type = ? AND NOT EXISTS (SELECT 1 FROM user_oauth_links WHERE user_oauth_links.provider_type = oauth_providers.provider_type)", + ) + .bind(provider_type) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(result.rows_affected() == 1) } } @@ -370,7 +459,7 @@ mod tests { use crate::run_migrations; use aether_data_contracts::repository::oauth_providers::{ EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository, - UpsertOAuthProviderConfigRecord, + UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord, }; fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord { @@ -379,8 +468,10 @@ mod tests { display_name: format!("{provider_type} display"), client_id: format!("{provider_type}-client"), client_secret_encrypted: EncryptedSecretUpdate::Preserve, - authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")), - token_url_override: Some(format!("https://{provider_type}.example.com/token")), + authorization_url_override: Some( + "https://connect.linux.do/oauth2/authorize".to_string(), + ), + token_url_override: Some("https://connect.linux.do/oauth2/token".to_string()), userinfo_url_override: None, scopes: Some(vec!["openid".to_string(), "profile".to_string()]), redirect_uri: format!("https://{provider_type}.example.com/redirect"), @@ -407,7 +498,7 @@ mod tests { let created = repository .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()), - ..sample_upsert("github") + ..sample_upsert("linuxdo") }) .await .expect("provider should upsert"); @@ -420,13 +511,38 @@ mod tests { let updated = repository .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { client_secret_encrypted: EncryptedSecretUpdate::Preserve, - display_name: "GitHub".to_string(), - ..sample_upsert("github") + display_name: "Linux.do".to_string(), + ..sample_upsert("linuxdo") }) .await .expect("provider should update"); - assert_eq!(updated.display_name, "GitHub"); + assert_eq!(updated.display_name, "Linux.do"); assert_eq!(updated.client_secret_encrypted.as_deref(), Some("secret-1")); + assert!( + repository + .compare_and_swap_oauth_provider_client_secret( + "linuxdo", + "secret-1", + "record-bound-v2", + ) + .await + .expect("client secret CAS should execute") + ); + let migrated = repository + .get_oauth_provider_config("linuxdo") + .await + .expect("provider should fetch") + .expect("provider should exist"); + assert_eq!(migrated.display_name, "Linux.do"); + assert_eq!(migrated.updated_at_unix_secs, updated.updated_at_unix_secs); + assert_eq!( + migrated.client_secret_encrypted.as_deref(), + Some("record-bound-v2") + ); + assert!(!repository + .compare_and_swap_oauth_provider_client_secret("linuxdo", "secret-1", "must-not-win",) + .await + .expect("stale client secret CAS should execute")); let listed = repository .list_oauth_provider_configs() @@ -435,7 +551,7 @@ mod tests { assert_eq!(listed.len(), 1); let fetched = repository - .get_oauth_provider_config("github") + .get_oauth_provider_config("linuxdo") .await .expect("provider should fetch") .expect("provider should exist"); @@ -461,8 +577,8 @@ INSERT INTO users ( INSERT INTO user_oauth_links ( id, user_id, provider_type, provider_user_id, linked_at ) VALUES - ('link-1', 'user-oauth', 'github', 'gh-1', 1), - ('link-2', 'user-local', 'github', 'gh-2', 1) + ('link-1', 'user-oauth', 'linuxdo', 'linuxdo-1', 1), + ('link-2', 'user-local', 'linuxdo', 'linuxdo-2', 1) "#, ) .execute(&pool) @@ -470,22 +586,68 @@ INSERT INTO user_oauth_links ( .expect("oauth links should seed"); assert_eq!( repository - .count_locked_users_if_provider_disabled("github", false) + .count_locked_users_if_provider_disabled("linuxdo", false) .await .expect("locked users should count"), 1 ); assert_eq!( repository - .count_locked_users_if_provider_disabled("github", true) + .count_locked_users_if_provider_disabled("linuxdo", true) .await .expect("locked users should count"), 2 ); - assert!(repository - .delete_oauth_provider_config("github") + assert_eq!( + repository + .upsert_oauth_provider_config_guarded( + &UpsertOAuthProviderConfigRecord { + is_enabled: false, + ..sample_upsert("linuxdo") + }, + false, + false, + 0, + ) + .await + .expect("guarded disable should resolve"), + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { affected_count: 1 } + ); + let forced = repository + .upsert_oauth_provider_config_guarded( + &UpsertOAuthProviderConfigRecord { + is_enabled: false, + ..sample_upsert("linuxdo") + }, + false, + true, + 0, + ) .await - .expect("provider should delete")); + .expect("forced provider disable should succeed"); + assert!(matches!( + forced, + UpsertOAuthProviderConfigOutcome::Upserted(provider) if !provider.is_enabled + )); + + assert!(!repository + .delete_oauth_provider_config_if_unlinked("linuxdo", false) + .await + .expect("linked provider deletion should resolve")); + assert!(repository + .get_oauth_provider_config("linuxdo") + .await + .expect("provider should fetch") + .is_some()); + + sqlx::query("DELETE FROM user_oauth_links WHERE provider_type = 'linuxdo'") + .execute(&pool) + .await + .expect("links should delete"); + assert!(repository + .delete_oauth_provider_config_if_unlinked("linuxdo", false) + .await + .expect("unlinked provider should delete")); } } diff --git a/crates/aether-data/adapters/sqlite/src/provider_catalog.rs b/crates/aether-data/adapters/sqlite/src/provider_catalog.rs index 567f28362..11be74256 100644 --- a/crates/aether-data/adapters/sqlite/src/provider_catalog.rs +++ b/crates/aether-data/adapters/sqlite/src/provider_catalog.rs @@ -9,9 +9,11 @@ use sqlx::{ use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate, - ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate, + ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, + ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, @@ -731,6 +733,48 @@ WHERE id = ? self.reload_provider(&provider.id, "updated").await } + pub async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + validate_non_empty(&update.provider_id, "provider catalog provider_id")?; + let expected_config = + optional_json_to_string(&update.expected_config, "providers.expected_config")?; + let config = optional_json_to_string(&update.config, "providers.config")?; + let rows_affected = sqlx::query( + r#" +UPDATE providers +SET config = ?, updated_at = ? +WHERE id = ? + AND config IS ? +"#, + ) + .bind(config) + .bind(current_unix_secs() as i64) + .bind(&update.provider_id) + .bind(expected_config) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + + pub async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog provider_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM providers WHERE id = ?", + "UPDATE providers SET proxy = ?, updated_at = ? WHERE id = ? AND proxy IS ?", + update, + "providers.proxy", + ) + .await + } + pub async fn delete_provider(&self, provider_id: &str) -> Result { validate_non_empty(provider_id, "provider catalog provider_id")?; let rows_affected = sqlx::query("DELETE FROM providers WHERE id = ?") @@ -935,6 +979,21 @@ WHERE id = ? self.reload_endpoint(&endpoint.id, "updated").await } + pub async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog endpoint_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM provider_endpoints WHERE id = ?", + "UPDATE provider_endpoints SET proxy = ?, updated_at = ? WHERE id = ? AND proxy IS ?", + update, + "provider_endpoints.proxy", + ) + .await + } + pub async fn delete_endpoint(&self, endpoint_id: &str) -> Result { validate_non_empty(endpoint_id, "provider catalog endpoint_id")?; let rows_affected = sqlx::query("DELETE FROM provider_endpoints WHERE id = ?") @@ -1117,6 +1176,46 @@ WHERE id = ? self.reload_key(&key.id, "updated").await } + pub async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + validate_non_empty(&update.record_id, "provider catalog key_id")?; + compare_and_swap_proxy_json( + &self.pool, + "SELECT proxy FROM provider_api_keys WHERE id = ?", + "UPDATE provider_api_keys SET proxy = ?, updated_at = ? WHERE id = ? AND proxy IS ?", + update, + "provider_api_keys.proxy", + ) + .await + } + + pub async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + validate_non_empty(&update.key_id, "provider catalog key_id")?; + validate_non_empty( + &update.expected_provider_id, + "provider catalog expected provider_id", + )?; + let rows_affected = sqlx::query( + "UPDATE provider_api_keys SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND provider_id = ? AND COALESCE(api_key, encrypted_key) IS ? AND auth_config IS ?", + ) + .bind(update.encrypted_api_key.as_deref()) + .bind(update.encrypted_auth_config.as_deref()) + .bind(&update.key_id) + .bind(&update.expected_provider_id) + .bind(update.expected_encrypted_api_key.as_deref()) + .bind(update.expected_encrypted_auth_config.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) + } + pub async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -1543,51 +1642,18 @@ WHERE id = ? Ok(rows_affected > 0) } - pub async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - validate_non_empty(key_id, "provider catalog key_id")?; - validate_non_empty(encrypted_api_key, "provider catalog oauth api_key")?; - let rows_affected = sqlx::query( - r#" -UPDATE provider_api_keys -SET api_key = ?, auth_config = ?, expires_at = ?, updated_at = ? -WHERE id = ? -"#, - ) - .bind(encrypted_api_key) - .bind(encrypted_auth_config) - .bind(optional_i64_from_u64( - expires_at_unix_secs, - "provider_api_keys.expires_at", - )?) - .bind(current_unix_secs() as i64) - .bind(key_id) - .execute(&self.pool) - .await - .map_sql_err()? - .rows_affected(); - Ok(rows_affected > 0) - } - pub async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { validate_non_empty(key_id, "provider catalog key_id")?; let rows_affected = sqlx::query( r#" UPDATE provider_api_keys -SET oauth_invalid_at = ?, oauth_invalid_reason = ?, - auth_config = COALESCE(?, auth_config), updated_at = ? +SET oauth_invalid_at = ?, oauth_invalid_reason = ?, updated_at = ? WHERE id = ? "#, ) @@ -1596,7 +1662,6 @@ WHERE id = ? "provider_api_keys.oauth_invalid_at", )?) .bind(oauth_invalid_reason) - .bind(encrypted_auth_config_update) .bind(updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64) .bind(key_id) .execute(&self.pool) @@ -2249,6 +2314,20 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository { Self::update_provider(self, provider).await } + async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_config(self, update).await + } + + async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_provider_proxy(self, update).await + } + async fn delete_provider(&self, provider_id: &str) -> Result { Self::delete_provider(self, provider_id).await } @@ -2284,6 +2363,13 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository { Self::update_endpoint(self, endpoint).await } + async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_endpoint_proxy(self, update).await + } + async fn delete_endpoint(&self, endpoint_id: &str) -> Result { Self::delete_endpoint(self, endpoint_id).await } @@ -2302,6 +2388,20 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository { Self::update_key(self, key).await } + async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Self::compare_and_swap_key_proxy(self, update).await + } + + async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + Self::compare_and_swap_key_credentials(self, update).await + } + async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -2396,29 +2496,11 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository { Self::clear_key_oauth_invalid_marker(self, key_id).await } - async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - Self::update_key_oauth_credentials( - self, - key_id, - encrypted_api_key, - encrypted_auth_config, - expires_at_unix_secs, - ) - .await - } - async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { Self::update_key_oauth_runtime_state( @@ -2426,7 +2508,6 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository { key_id, oauth_invalid_at_unix_secs, oauth_invalid_reason, - encrypted_auth_config_update, updated_at_unix_secs, ) .await @@ -2752,6 +2833,44 @@ fn optional_json_to_string( optional_json_ref_to_string(value.as_ref(), field_name) } +async fn compare_and_swap_proxy_json( + pool: &SqlitePool, + select_sql: &'static str, + update_sql: &'static str, + update: &ProviderCatalogProxyCasUpdate, + field_name: &'static str, +) -> Result { + // SQLite stores these JSON fields as TEXT. Legacy Python rows commonly include + // insignificant whitespace, so comparing a serde_json re-serialization directly would + // make lazy credential migration conflict forever. Compare semantic JSON first and use + // the exact observed bytes as the atomic write fence. + // Outer None means the row does not exist; inner None is an existing SQL NULL proxy. + let observed_raw: Option> = sqlx::query_scalar::<_, Option>(select_sql) + .bind(&update.record_id) + .fetch_optional(pool) + .await + .map_sql_err()?; + let Some(observed_raw) = observed_raw else { + return Ok(false); + }; + let observed = optional_json_from_string(observed_raw.clone(), field_name)?; + if observed != update.expected_proxy { + return Ok(false); + } + + let replacement = optional_json_to_string(&update.proxy, field_name)?; + let rows_affected = sqlx::query(update_sql) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(&update.record_id) + .bind(observed_raw) + .execute(pool) + .await + .map_sql_err()? + .rows_affected(); + Ok(rows_affected == 1) +} + fn key_insert_sql() -> &'static str { r#" INSERT INTO provider_api_keys ( @@ -3306,11 +3425,11 @@ fn map_key_row(row: &SqliteRow) -> Result Result(key) + key.with_auth_channel_policy_fields( + auth_type_by_format, + allow_auth_channel_mismatch_formats, + ) })? } @@ -3390,12 +3512,222 @@ mod tests { ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, - ProviderCatalogUpstreamMetadataNamespaceExpectation, + ProviderCatalogProxyCasUpdate, ProviderCatalogUpstreamMetadataNamespaceExpectation, ProviderCatalogUpstreamMetadataNamespaceUpdate, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use serde_json::json; + #[test] + fn credential_cas_migrates_legacy_encrypted_key_with_null_safe_fence() { + let source = include_str!("provider_catalog.rs"); + assert!(source.contains( + "UPDATE provider_api_keys SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND provider_id = ? AND COALESCE(api_key, encrypted_key) IS ? AND auth_config IS ?" + )); + } + + #[tokio::test] + async fn sqlite_proxy_cas_accepts_python_spaced_json_and_distinguishes_null_from_missing_rows() + { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteProviderCatalogReadRepository::new(pool.clone()); + + let expected_proxy = json!({ + "url": "http://proxy.example.test:8080/", + "username": "alice" + }); + let replacement_proxy = json!({ + "url": "http://proxy.example.test:8080/", + "username": "aether-runtime-secret:v1:sealed" + }); + let mut spaced_provider = StoredProviderCatalogProvider::new( + "proxy-cas-spaced-provider".to_string(), + "Proxy CAS Spaced Provider".to_string(), + None, + "custom".to_string(), + ) + .expect("provider should build"); + spaced_provider.proxy = Some(expected_proxy.clone()); + repository + .create_provider(&spaced_provider, None) + .await + .expect("provider should create"); + + let legacy_python_json = + r#"{"url": "http://proxy.example.test:8080/", "username": "alice"}"#; + sqlx::query("UPDATE providers SET proxy = ? WHERE id = ?") + .bind(legacy_python_json) + .bind(&spaced_provider.id) + .execute(&pool) + .await + .expect("legacy proxy JSON should be seeded"); + + let update = ProviderCatalogProxyCasUpdate { + record_id: spaced_provider.id.clone(), + expected_proxy: Some(expected_proxy), + proxy: Some(replacement_proxy.clone()), + }; + assert!(repository + .compare_and_swap_provider_proxy(&update) + .await + .expect("semantic proxy CAS should run")); + assert!(!repository + .compare_and_swap_provider_proxy(&update) + .await + .expect("stale semantic proxy CAS should run")); + let stored_raw: Option = + sqlx::query_scalar("SELECT proxy FROM providers WHERE id = ?") + .bind(&spaced_provider.id) + .fetch_one(&pool) + .await + .expect("updated proxy should load"); + assert_eq!( + stored_raw + .as_deref() + .map(serde_json::from_str::) + .transpose() + .expect("stored proxy should remain valid JSON"), + Some(replacement_proxy.clone()) + ); + assert_ne!(stored_raw.as_deref(), Some(legacy_python_json)); + + let endpoint = StoredProviderCatalogEndpoint::new( + "proxy-cas-spaced-endpoint".to_string(), + spaced_provider.id.clone(), + "openai:chat".to_string(), + Some("openai".to_string()), + Some("chat".to_string()), + true, + ) + .expect("endpoint should build") + .with_transport_fields( + "https://api.example.test/v1".to_string(), + None, + None, + None, + None, + None, + None, + Some(json!({ + "url": "http://proxy.example.test:8080/", + "username": "alice" + })), + ) + .expect("endpoint transport should build"); + repository + .create_endpoint(&endpoint) + .await + .expect("endpoint should create"); + sqlx::query("UPDATE provider_endpoints SET proxy = ? WHERE id = ?") + .bind(legacy_python_json) + .bind(&endpoint.id) + .execute(&pool) + .await + .expect("legacy endpoint proxy JSON should be seeded"); + let endpoint_update = ProviderCatalogProxyCasUpdate { + record_id: endpoint.id.clone(), + expected_proxy: Some(json!({ + "url": "http://proxy.example.test:8080/", + "username": "alice" + })), + proxy: Some(replacement_proxy.clone()), + }; + assert!(repository + .compare_and_swap_endpoint_proxy(&endpoint_update) + .await + .expect("semantic endpoint proxy CAS should run")); + assert!(!repository + .compare_and_swap_endpoint_proxy(&endpoint_update) + .await + .expect("stale endpoint proxy CAS should run")); + + let key = StoredProviderCatalogKey::new( + "proxy-cas-spaced-key".to_string(), + spaced_provider.id.clone(), + "Proxy CAS Key".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + None, + None::, + None, + None, + None, + None, + None, + Some(json!({ + "url": "http://proxy.example.test:8080/", + "username": "alice" + })), + None, + ) + .expect("key transport should build"); + repository + .create_key(&key) + .await + .expect("key should create"); + sqlx::query("UPDATE provider_api_keys SET proxy = ? WHERE id = ?") + .bind(legacy_python_json) + .bind(&key.id) + .execute(&pool) + .await + .expect("legacy key proxy JSON should be seeded"); + let key_update = ProviderCatalogProxyCasUpdate { + record_id: key.id.clone(), + expected_proxy: Some(json!({ + "url": "http://proxy.example.test:8080/", + "username": "alice" + })), + proxy: Some(replacement_proxy), + }; + assert!(repository + .compare_and_swap_key_proxy(&key_update) + .await + .expect("semantic key proxy CAS should run")); + assert!(!repository + .compare_and_swap_key_proxy(&key_update) + .await + .expect("stale key proxy CAS should run")); + + let null_provider = StoredProviderCatalogProvider::new( + "proxy-cas-null-provider".to_string(), + "Proxy CAS Null Provider".to_string(), + None, + "custom".to_string(), + ) + .expect("provider should build"); + repository + .create_provider(&null_provider, None) + .await + .expect("null-proxy provider should create"); + assert!(repository + .compare_and_swap_provider_proxy(&ProviderCatalogProxyCasUpdate { + record_id: null_provider.id, + expected_proxy: None, + proxy: Some(json!({"url": "http://proxy.example.test:8081/"})), + }) + .await + .expect("existing NULL proxy should be distinguishable from a missing row")); + assert!(!repository + .compare_and_swap_provider_proxy(&ProviderCatalogProxyCasUpdate { + record_id: "proxy-cas-missing-provider".to_string(), + expected_proxy: None, + proxy: Some(json!({"url": "http://proxy.example.test:8082/"})), + }) + .await + .expect("missing-row proxy CAS should run")); + } + #[tokio::test] async fn sqlite_admin_credential_cas_rotates_codex_namespace_atomically() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -4682,15 +5014,6 @@ mod tests { ) .await .expect("model fetch success should update atomically")); - assert!(repository - .update_key_oauth_credentials( - "key-write-1", - "enc-key-2", - Some("enc-auth-2"), - Some(1_750_000_000), - ) - .await - .expect("oauth credentials should update")); assert!(repository .update_key_health_state( "key-write-1", @@ -4711,10 +5034,6 @@ mod tests { .expect("key should reload") .pop() .expect("key should exist"); - assert_eq!( - reloaded_key.encrypted_api_key, - Some("enc-key-2".to_string()) - ); assert_eq!( reloaded_key.upstream_metadata, Some(json!({ diff --git a/crates/aether-data/adapters/sqlite/src/proxy_nodes.rs b/crates/aether-data/adapters/sqlite/src/proxy_nodes.rs index 25d4b4443..ee4d29d59 100644 --- a/crates/aether-data/adapters/sqlite/src/proxy_nodes.rs +++ b/crates/aether-data/adapters/sqlite/src/proxy_nodes.rs @@ -3,7 +3,8 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::proxy_nodes::{ bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, - normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security, + merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata, + normalize_proxy_metadata, proxy_metadata_has_explicit_tunnel_security, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation, @@ -18,6 +19,8 @@ use aether_data_query::{push_eq, push_limit, WhereClause}; use crate::error::SqlResultExt; use crate::SqlitePool; +const PROXY_NODE_REGISTRATION_CAS_RETRIES: usize = 8; + fn log_reported_tunnel_error_event( node_id: &str, event: &TunnelErrorEventRecord, @@ -50,19 +53,22 @@ impl SqliteProxyNodeReadRepository { Self { pool } } - async fn upsert_node(&self, node: &StoredProxyNode) -> Result<(), DataLayerError> { + async fn write_node( + &self, + node: &StoredProxyNode, + update_existing: bool, + ) -> Result<(), DataLayerError> { let now = current_unix_secs(); - sqlx::query( - r#" + let upsert_sql = r#" INSERT INTO proxy_nodes ( - id, name, ip, port, region, status, registered_by, last_heartbeat_at, + id, tunnel_generation, name, ip, port, region, status, registered_by, last_heartbeat_at, heartbeat_interval, active_connections, total_requests, avg_latency_ms, is_manual, proxy_url, proxy_username, proxy_password, created_at, updated_at, remote_config, config_version, hardware_info, estimated_max_concurrency, tunnel_mode, tunnel_connected, tunnel_connected_at, failed_requests, dns_failures, stream_errors, proxy_metadata ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET name = excluded.name, ip = excluded.ip, @@ -91,58 +97,115 @@ ON CONFLICT(id) DO UPDATE SET dns_failures = excluded.dns_failures, stream_errors = excluded.stream_errors, proxy_metadata = excluded.proxy_metadata -"#, - ) - .bind(&node.id) - .bind(&node.name) - .bind(&node.ip) - .bind(node.port) - .bind(&node.region) - .bind(&node.status) - .bind(&node.registered_by) - .bind(optional_i64_from_u64( - node.last_heartbeat_at_unix_secs, - "proxy_nodes.last_heartbeat_at", - )?) - .bind(node.heartbeat_interval) - .bind(node.active_connections) - .bind(node.total_requests) - .bind(node.avg_latency_ms) - .bind(node.is_manual) - .bind(&node.proxy_url) - .bind(&node.proxy_username) - .bind(&node.proxy_password) - .bind(node.created_at_unix_ms.unwrap_or(now) as i64) - .bind(node.updated_at_unix_secs.unwrap_or(now) as i64) - .bind(optional_json_to_string( - &node.remote_config, - "proxy_nodes.remote_config", - )?) - .bind(node.config_version) - .bind(optional_json_to_string( - &node.hardware_info, - "proxy_nodes.hardware_info", - )?) - .bind(node.estimated_max_concurrency) - .bind(node.tunnel_mode) - .bind(node.tunnel_connected) - .bind(optional_i64_from_u64( - node.tunnel_connected_at_unix_secs, - "proxy_nodes.tunnel_connected_at", - )?) - .bind(node.failed_requests) - .bind(node.dns_failures) - .bind(node.stream_errors) - .bind(optional_json_to_string( - &node.proxy_metadata, - "proxy_nodes.proxy_metadata", - )?) - .execute(&self.pool) - .await - .map_sql_err()?; +"#; + let sql = if update_existing { + upsert_sql + } else { + upsert_sql + .split_once("\nON CONFLICT(id) DO UPDATE SET") + .map(|(insert_sql, _)| insert_sql) + .expect("proxy node upsert SQL should contain its conflict clause") + }; + sqlx::query(sql) + .bind(&node.id) + .bind(&node.tunnel_generation) + .bind(&node.name) + .bind(&node.ip) + .bind(node.port) + .bind(&node.region) + .bind(&node.status) + .bind(&node.registered_by) + .bind(optional_i64_from_u64( + node.last_heartbeat_at_unix_secs, + "proxy_nodes.last_heartbeat_at", + )?) + .bind(node.heartbeat_interval) + .bind(node.active_connections) + .bind(node.total_requests) + .bind(node.avg_latency_ms) + .bind(node.is_manual) + .bind(&node.proxy_url) + .bind(&node.proxy_username) + .bind(&node.proxy_password) + .bind(node.created_at_unix_ms.unwrap_or(now) as i64) + .bind(node.updated_at_unix_secs.unwrap_or(now) as i64) + .bind(optional_json_to_string( + &node.remote_config, + "proxy_nodes.remote_config", + )?) + .bind(node.config_version) + .bind(optional_json_to_string( + &node.hardware_info, + "proxy_nodes.hardware_info", + )?) + .bind(node.estimated_max_concurrency) + .bind(node.tunnel_mode) + .bind(node.tunnel_connected) + .bind(optional_i64_from_u64( + node.tunnel_connected_at_unix_secs, + "proxy_nodes.tunnel_connected_at", + )?) + .bind(node.failed_requests) + .bind(node.dns_failures) + .bind(node.stream_errors) + .bind(optional_json_to_string( + &node.proxy_metadata, + "proxy_nodes.proxy_metadata", + )?) + .execute(&self.pool) + .await + .map_sql_err()?; Ok(()) } + async fn insert_node(&self, node: &StoredProxyNode) -> Result<(), DataLayerError> { + self.write_node(node, false).await + } + + async fn update_existing_registration_if_unchanged( + &self, + mutation: &ProxyNodeRegistrationMutation, + existing: &StoredProxyNode, + replacement_proxy_metadata: Option<&serde_json::Value>, + now: u64, + ) -> Result { + let hardware_info = + optional_json_to_string(&mutation.hardware_info, "proxy_nodes.hardware_info")?; + let replacement_proxy_metadata = optional_json_to_string( + &replacement_proxy_metadata.cloned(), + "proxy_nodes.proxy_metadata", + )?; + let expected_proxy_metadata = + optional_json_to_string(&existing.proxy_metadata, "proxy_nodes.proxy_metadata")?; + let result = sqlx::query(UPDATE_PROXY_NODE_REGISTRATION_SQL) + .bind(&mutation.name) + .bind(&mutation.ip) + .bind(mutation.port) + .bind(mutation.region.as_deref()) + .bind(mutation.registered_by.as_deref()) + .bind(now as i64) + .bind(mutation.heartbeat_interval) + .bind(mutation.active_connections) + .bind(mutation.total_requests) + .bind(mutation.avg_latency_ms) + .bind(hardware_info) + .bind(mutation.estimated_max_concurrency) + .bind(mutation.tunnel_mode) + .bind(replacement_proxy_metadata) + .bind(now as i64) + .bind(&existing.id) + .bind(&existing.tunnel_generation) + .bind(&existing.ip) + .bind(existing.port) + .bind(expected_proxy_metadata.as_deref()) + .bind(expected_proxy_metadata.as_deref()) + .bind(expected_proxy_metadata.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + async fn find_duplicate_proxy_node( &self, ip: &str, @@ -172,9 +235,26 @@ ON CONFLICT(id) DO UPDATE SET row.as_ref().map(map_proxy_node_row).transpose() } + async fn find_registered_proxy_node_by_endpoint( + &self, + ip: &str, + port: i32, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? AND is_manual = 0 ORDER BY created_at ASC, id ASC LIMIT 1" + )) + .bind(ip) + .bind(port) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_proxy_node_row).transpose() + } + async fn insert_event( &self, node_id: &str, + expected_tunnel_generation: Option<&str>, event_type: &str, detail: Option<&str>, event_metadata: Option<&serde_json::Value>, @@ -183,10 +263,11 @@ ON CONFLICT(id) DO UPDATE SET sqlx::query( r#" INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at) -VALUES (?, ?, ?, ?, ?) +SELECT id, ?, ?, ?, ? +FROM proxy_nodes +WHERE id = ? AND (? IS NULL OR tunnel_generation = ?) "#, ) - .bind(node_id) .bind(event_type) .bind(detail) .bind(optional_json_to_string( @@ -194,6 +275,9 @@ VALUES (?, ?, ?, ?, ?) "proxy_node_events.event_metadata", )?) .bind(created_at_unix_secs.unwrap_or_else(current_unix_secs) as i64) + .bind(node_id) + .bind(expected_tunnel_generation) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_sql_err()?; @@ -204,6 +288,7 @@ VALUES (?, ?, ?, ?, ?) &self, table: &str, node_id: &str, + expected_tunnel_generation: Option<&str>, bucket_start: u64, sample: &TunnelMetricsSample, ) -> Result<(), DataLayerError> { @@ -226,7 +311,9 @@ INSERT INTO {table} ( ws_in_frames_delta, ws_out_frames_delta ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +SELECT ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? +FROM proxy_nodes +WHERE id = ? AND (? IS NULL OR tunnel_generation = ?) ON CONFLICT(node_id, bucket_start_unix_secs) DO UPDATE SET samples = {table}.samples + excluded.samples, uptime_samples = {table}.uptime_samples + excluded.uptime_samples, @@ -258,6 +345,9 @@ ON CONFLICT(node_id, bucket_start_unix_secs) DO UPDATE SET .bind(sample.ws_out_bytes_delta) .bind(sample.ws_in_frames_delta) .bind(sample.ws_out_frames_delta) + .bind(node_id) + .bind(expected_tunnel_generation) + .bind(expected_tunnel_generation) .execute(&self.pool) .await .map_sql_err()?; @@ -331,6 +421,7 @@ ON CONFLICT(node_id, bucket_start_unix_secs) DO UPDATE SET const PROXY_NODE_COLUMNS: &str = r#" SELECT id, + tunnel_generation, name, ip, port, @@ -373,6 +464,139 @@ SELECT FROM proxy_node_events "#; +const APPLY_HEARTBEAT_SQL: &str = r#" +UPDATE proxy_nodes +SET last_heartbeat_at = ?, + tunnel_connected_at = CASE + WHEN status <> 'online' OR tunnel_connected = 0 THEN ? + ELSE tunnel_connected_at + END, + updated_at = CASE + WHEN status <> 'online' OR tunnel_connected = 0 THEN ? + ELSE updated_at + END, + status = 'online', + tunnel_connected = 1, + heartbeat_interval = COALESCE(?, heartbeat_interval), + active_connections = COALESCE(?, active_connections), + avg_latency_ms = COALESCE(?, avg_latency_ms), + total_requests = total_requests + MAX(COALESCE(?, 0), 0), + failed_requests = failed_requests + MAX(COALESCE(?, 0), 0), + dns_failures = dns_failures + MAX(COALESCE(?, 0), 0), + stream_errors = stream_errors + MAX(COALESCE(?, 0), 0) +WHERE id = ? + AND tunnel_mode = 1 + AND tunnel_generation = ? +"#; + +const CAS_HEARTBEAT_PROXY_METADATA_SQL: &str = r#" +UPDATE proxy_nodes +SET proxy_metadata = ?, updated_at = ? +WHERE id = ? AND tunnel_generation = ? + AND ( + (proxy_metadata IS NULL AND ? IS NULL) + OR ( + proxy_metadata IS NOT NULL AND ? IS NOT NULL + AND json(proxy_metadata) = json(?) + ) + ) +"#; + +const UPDATE_TUNNEL_STATUS_SQL: &str = r#" +UPDATE proxy_nodes +SET tunnel_connected = ?, + active_connections = CASE WHEN ? THEN active_connections ELSE 0 END, + tunnel_connected_at = ?, + status = CASE WHEN ? THEN 'online' ELSE 'offline' END, + updated_at = ? +WHERE id = ? + AND tunnel_generation = ? + AND (tunnel_connected_at IS NULL OR tunnel_connected_at <= ?) +"#; + +const UPDATE_MANUAL_PROXY_NODE_SQL: &str = r#" +UPDATE proxy_nodes +SET name = COALESCE(?, name), + ip = COALESCE(?, ip), + port = COALESCE(?, port), + region = COALESCE(?, region), + proxy_url = COALESCE(?, proxy_url), + proxy_username = COALESCE(?, proxy_username), + proxy_password = COALESCE(?, proxy_password), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND tunnel_generation = ? +"#; + +const UPDATE_PROXY_NODE_REGISTRATION_SQL: &str = r#" +UPDATE proxy_nodes +SET name = ?, ip = ?, port = ?, region = ?, registered_by = ?, + last_heartbeat_at = ?, heartbeat_interval = ?, + active_connections = COALESCE(?, active_connections), + total_requests = COALESCE(?, total_requests), + avg_latency_ms = COALESCE(?, avg_latency_ms), + hardware_info = COALESCE(?, hardware_info), + estimated_max_concurrency = COALESCE(?, estimated_max_concurrency), + tunnel_mode = ?, proxy_metadata = COALESCE(?, proxy_metadata), updated_at = ? +WHERE id = ? AND tunnel_generation = ? + AND is_manual = 0 AND ip = ? AND port = ? + AND ( + (proxy_metadata IS NULL AND ? IS NULL) + OR ( + proxy_metadata IS NOT NULL AND ? IS NOT NULL + AND json(proxy_metadata) = json(?) + ) + ) +"#; + +const UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL: &str = r#" +UPDATE proxy_nodes +SET name = COALESCE(?, name), remote_config = ?, + config_version = config_version + 1, updated_at = ? +WHERE id = ? AND tunnel_generation = ? AND config_version = ? + AND is_manual = 0 +"#; + +const RECORD_PROXY_NODE_TRAFFIC_SQL: &str = r#" +UPDATE proxy_nodes +SET total_requests = total_requests + MAX(?, 0), + failed_requests = failed_requests + MAX(?, 0), + dns_failures = dns_failures + MAX(?, 0), + stream_errors = stream_errors + MAX(?, 0), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND tunnel_generation = ? +"#; + +const INCREMENT_MANUAL_PROXY_NODE_REQUESTS_SQL: &str = r#" +UPDATE proxy_nodes +SET total_requests = total_requests + MAX(?, 0), + failed_requests = failed_requests + MAX(?, 0), + avg_latency_ms = COALESCE(?, avg_latency_ms), + updated_at = ? +WHERE id = ? AND is_manual = 1 + AND tunnel_generation = ? +"#; + +const UNREGISTER_PROXY_NODE_SQL: &str = r#" +UPDATE proxy_nodes +SET status = 'offline', tunnel_connected = 0, active_connections = 0, + tunnel_connected_at = ?, updated_at = ? +WHERE id = ? + AND tunnel_generation = ? +"#; + +// Retire counter rows after the parent delete commits. Keeping this statement +// outside the parent-locking transaction avoids the outbox -> proxy_nodes +// versus proxy_nodes -> outbox lock inversion with the counter flusher. +const RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL: &str = r#" +DELETE FROM usage_counter_deltas +WHERE kind = 'proxy_node' + AND target_id = ? + AND target_tunnel_generation = ? + AND processed_at IS NULL +"#; + #[async_trait] impl ProxyNodeReadRepository for SqliteProxyNodeReadRepository { async fn list_proxy_nodes(&self) -> Result, DataLayerError> { @@ -573,6 +797,60 @@ WHERE is_manual = 0 Ok(result.rows_affected() as usize) } + async fn compare_and_set_proxy_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_password = ?, updated_at = ? +WHERE id = ? AND proxy_password = ? +"#, + ) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(expected) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn compare_and_set_proxy_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result { + let expected = serde_json::to_string(expected).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let replacement = serde_json::to_string(replacement).map_err(|err| { + DataLayerError::InvalidInput(format!("proxy_nodes.proxy_metadata is invalid: {err}")) + })?; + let result = sqlx::query( + r#" +UPDATE proxy_nodes +SET proxy_metadata = ?, updated_at = ? +WHERE id = ? + AND proxy_metadata IS NOT NULL + AND json(proxy_metadata) = json(?) +"#, + ) + .bind(replacement) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(expected) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + async fn create_manual_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -584,9 +862,14 @@ WHERE is_manual = 0 return Err(duplicate_proxy_node_error(&existing)); } + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + if let Some(existing) = self.find_proxy_node(&node_id).await? { + return Err(proxy_node_id_in_use_error(&existing)); + } let now = Some(current_unix_secs()); let node = StoredProxyNode::new( - uuid::Uuid::new_v4().to_string(), + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -621,7 +904,18 @@ WHERE is_manual = 0 now, ); - self.upsert_node(&node).await?; + if let Err(error) = self.insert_node(&node).await { + if let Some(duplicate) = self + .find_duplicate_proxy_node(&mutation.ip, mutation.port, None) + .await? + { + return Err(duplicate_proxy_node_error(&duplicate)); + } + if let Some(owner) = self.find_proxy_node(&node.id).await? { + return Err(proxy_node_id_in_use_error(&owner)); + } + return Err(error); + } Ok(node) } @@ -629,17 +923,17 @@ WHERE is_manual = 0 &self, mutation: &ProxyNodeManualUpdateMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; - if !node.is_manual { + if !existing.is_manual { return Err(DataLayerError::InvalidInput( "只能编辑手动添加的代理节点".to_string(), )); } - let next_ip = mutation.ip.as_deref().unwrap_or(node.ip.as_str()); - let next_port = mutation.port.unwrap_or(node.port); + let next_ip = mutation.ip.as_deref().unwrap_or(existing.ip.as_str()); + let next_port = mutation.port.unwrap_or(existing.port); if let Some(existing) = self .find_duplicate_proxy_node(next_ip, next_port, Some(&mutation.node_id)) .await? @@ -647,56 +941,62 @@ WHERE is_manual = 0 return Err(duplicate_proxy_node_error(&existing)); } - if let Some(name) = mutation.name.as_ref() { - node.name = name.clone(); + let result = sqlx::query(UPDATE_MANUAL_PROXY_NODE_SQL) + .bind(mutation.name.as_deref()) + .bind(mutation.ip.as_deref()) + .bind(mutation.port) + .bind(mutation.region.as_deref()) + .bind(mutation.proxy_url.as_deref()) + .bind(mutation.proxy_username.as_deref()) + .bind(mutation.proxy_password.as_deref()) + .bind(current_unix_secs() as i64) + .bind(&mutation.node_id) + .bind(&existing.tunnel_generation) + .execute(&self.pool) + .await; + let result = match result { + Ok(result) => result, + Err(error) => { + if let Some(duplicate) = self + .find_duplicate_proxy_node(next_ip, next_port, Some(&mutation.node_id)) + .await? + { + return Err(duplicate_proxy_node_error(&duplicate)); + } + return Err(DataLayerError::sql(error)); + } + }; + if result.rows_affected() == 0 { + return Ok(None); } - if let Some(ip) = mutation.ip.as_ref() { - node.ip = ip.clone(); - } - if let Some(port) = mutation.port { - node.port = port; - } - if let Some(region) = mutation.region.as_ref() { - node.region = Some(region.clone()); - } - if let Some(proxy_url) = mutation.proxy_url.as_ref() { - node.proxy_url = Some(proxy_url.clone()); - } - if let Some(proxy_username) = mutation.proxy_username.as_ref() { - node.proxy_username = Some(proxy_username.clone()); - } - if let Some(proxy_password) = mutation.proxy_password.as_ref() { - node.proxy_password = Some(proxy_password.clone()); - } - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(Some(node)) + self.find_proxy_node(&mutation.node_id).await } async fn register_node( &self, mutation: &ProxyNodeRegistrationMutation, ) -> Result { - let now = Some(current_unix_secs()); + let requested_id = requested_proxy_node_id(mutation.node_id.as_deref())?; let normalized_proxy_metadata = normalize_proxy_metadata( mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), ); + let rotates_tunnel_security = + proxy_metadata_has_explicit_tunnel_security(normalized_proxy_metadata.as_ref()); - let existing = sqlx::query(&format!( - "{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? AND is_manual = 0 ORDER BY created_at ASC, id ASC LIMIT 1" - )) - .bind(&mutation.ip) - .bind(mutation.port) - .fetch_optional(&self.pool) - .await - .map_sql_err()?; - - let mut node = if let Some(row) = existing.as_ref() { - map_proxy_node_row(row)? - } else { - StoredProxyNode::new( - uuid::Uuid::new_v4().to_string(), + let Some(initial_existing) = self + .find_registered_proxy_node_by_endpoint(&mutation.ip, mutation.port) + .await? + else { + let node_id = requested_id + .clone() + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + if let Some(existing) = self.find_proxy_node(&node_id).await? { + return Err(proxy_node_id_in_use_error(&existing)); + } + let now = Some(current_unix_secs()); + let node = StoredProxyNode::new( + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -717,143 +1017,243 @@ WHERE is_manual = 0 mutation.registered_by.clone(), now, mutation.avg_latency_ms, - normalized_proxy_metadata.clone(), + merge_proxy_metadata_for_registration(None, normalized_proxy_metadata.clone()), mutation.hardware_info.clone(), mutation.estimated_max_concurrency, None, None, now, now, - ) + ); + if let Err(error) = self.insert_node(&node).await { + if let Some(winner) = self + .find_duplicate_proxy_node(&mutation.ip, mutation.port, None) + .await? + { + if winner.is_manual { + return Err(duplicate_proxy_node_error(&winner)); + } + if let Some(requested_id) = requested_id.as_deref() { + if requested_id != winner.id { + return Err(proxy_node_registration_identity_error( + requested_id, + &winner.id, + )); + } + } + return Ok(winner); + } + if let Some(owner) = self.find_proxy_node(&node.id).await? { + return Err(proxy_node_id_in_use_error(&owner)); + } + return Err(error); + } + return Ok(node); }; - node.name = mutation.name.clone(); - node.ip = mutation.ip.clone(); - node.port = mutation.port; - node.region = mutation.region.clone(); - node.registered_by = mutation.registered_by.clone(); - node.last_heartbeat_at_unix_secs = now; - node.heartbeat_interval = mutation.heartbeat_interval; - node.tunnel_mode = mutation.tunnel_mode; - if let Some(active_connections) = mutation.active_connections { - node.active_connections = active_connections; + if let Some(requested_id) = requested_id.as_deref() { + if requested_id != initial_existing.id { + return Err(proxy_node_registration_identity_error( + requested_id, + &initial_existing.id, + )); + } } - if let Some(total_requests) = mutation.total_requests { - node.total_requests = total_requests; + + let pinned_id = initial_existing.id.clone(); + let pinned_generation = initial_existing.tunnel_generation.clone(); + let mut existing = initial_existing; + for attempt in 0..PROXY_NODE_REGISTRATION_CAS_RETRIES { + if attempt != 0 { + existing = self + .find_registered_proxy_node_by_endpoint(&mutation.ip, mutation.port) + .await? + .ok_or_else(proxy_node_registration_changed_error)?; + } + if existing.id != pinned_id || existing.tunnel_generation != pinned_generation { + return Err(proxy_node_registration_changed_error()); + } + + let replacement_proxy_metadata = merge_proxy_metadata_for_registration( + existing.proxy_metadata.as_ref(), + normalized_proxy_metadata.clone(), + ); + let now = current_unix_secs(); + if self + .update_existing_registration_if_unchanged( + mutation, + &existing, + replacement_proxy_metadata.as_ref(), + now, + ) + .await? + { + return self + .find_proxy_node(&pinned_id) + .await? + .filter(|current| current.tunnel_generation == pinned_generation) + .ok_or_else(proxy_node_registration_changed_error); + } + if rotates_tunnel_security { + return Err(DataLayerError::UnexpectedValue( + "proxy node changed during explicit tunnel security rotation".to_string(), + )); + } } - if let Some(avg_latency_ms) = mutation.avg_latency_ms { - node.avg_latency_ms = Some(avg_latency_ms); - } - if let Some(hardware_info) = mutation.hardware_info.as_ref() { - node.hardware_info = Some(hardware_info.clone()); - } - if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency { - node.estimated_max_concurrency = Some(estimated_max_concurrency); - } - if let Some(proxy_metadata) = normalized_proxy_metadata { - node.proxy_metadata = Some(proxy_metadata); - } - if node.created_at_unix_ms.is_none() { - node.created_at_unix_ms = now; - } - node.updated_at_unix_secs = now; - self.upsert_node(&node).await?; - Ok(node) + + Err(DataLayerError::UnexpectedValue( + "proxy node registration changed during every CAS retry".to_string(), + )) } async fn apply_heartbeat( &self, mutation: &ProxyNodeHeartbeatMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; - if !node.tunnel_mode { + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != existing.tunnel_generation) + { + return Ok(None); + } + if !existing.tunnel_mode { return Err(DataLayerError::InvalidInput( "non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode" .to_string(), )); } - let previous_proxy_metadata = node.proxy_metadata.clone(); + let tunnel_generation = existing.tunnel_generation.clone(); let now_unix_secs = current_unix_secs(); - let now = Some(now_unix_secs); - node.last_heartbeat_at_unix_secs = now; - if node.status != "online" || !node.tunnel_connected { - node.status = "online".to_string(); - node.tunnel_connected = true; - node.tunnel_connected_at_unix_secs = now; - node.updated_at_unix_secs = now; - } - if let Some(value) = mutation.heartbeat_interval { - node.heartbeat_interval = value; - } - if let Some(value) = mutation.active_connections { - node.active_connections = value; - } - if let Some(value) = mutation.avg_latency_ms { - node.avg_latency_ms = Some(value); - } - let normalized_proxy_metadata = normalize_proxy_metadata( + let now = i64::try_from(now_unix_secs).unwrap_or(i64::MAX); + let has_proxy_metadata_update = normalize_heartbeat_proxy_metadata( + None, mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), - ); - let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security( - previous_proxy_metadata.as_ref(), - normalized_proxy_metadata, - ); - if let Some(value) = normalized_proxy_metadata { - node.proxy_metadata = Some(value); - } - if let Some(value) = mutation.total_requests_delta.filter(|value| *value > 0) { - node.total_requests += value; - } - if let Some(value) = mutation.failed_requests_delta.filter(|value| *value > 0) { - node.failed_requests += value; - } - if let Some(value) = mutation.dns_failures_delta.filter(|value| *value > 0) { - node.dns_failures += value; - } - if let Some(value) = mutation.stream_errors_delta.filter(|value| *value > 0) { - node.stream_errors += value; - } - let reconciled_remote_config = reconcile_remote_config_after_heartbeat( - node.remote_config.as_ref(), - mutation.proxy_version.as_deref(), - ); - if reconciled_remote_config != node.remote_config { - node.remote_config = reconciled_remote_config; - node.config_version = node.config_version.saturating_add(1); - node.updated_at_unix_secs = now; + ) + .is_some(); + + let result = sqlx::query(APPLY_HEARTBEAT_SQL) + .bind(now) + .bind(now) + .bind(now) + .bind(mutation.heartbeat_interval) + .bind(mutation.active_connections) + .bind(mutation.avg_latency_ms) + .bind(mutation.total_requests_delta) + .bind(mutation.failed_requests_delta) + .bind(mutation.dns_failures_delta) + .bind(mutation.stream_errors_delta) + .bind(&mutation.node_id) + .bind(&tunnel_generation) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + return Ok(None); } - let tunnel_metrics_sample = build_tunnel_metrics_sample( - previous_proxy_metadata.as_ref(), - node.proxy_metadata.as_ref(), - node.active_connections, - node.tunnel_connected, - ); + let mut updated = None; + let mut tunnel_metrics_sample = None; + if has_proxy_metadata_update { + for _ in 0..8 { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + let Some(replacement) = normalize_heartbeat_proxy_metadata( + current.proxy_metadata.as_ref(), + mutation.proxy_metadata.as_ref(), + mutation.proxy_version.as_deref(), + ) else { + break; + }; + if current.proxy_metadata.as_ref() == Some(&replacement) { + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + Some(&replacement), + current.active_connections, + current.tunnel_connected, + ); + updated = Some(current); + break; + } - self.upsert_node(&node).await?; + let expected = + optional_json_to_string(¤t.proxy_metadata, "proxy_nodes.proxy_metadata")?; + let replacement_json = serde_json::to_string(&replacement).map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "proxy_nodes.proxy_metadata contains unserializable JSON: {error}" + )) + })?; + let result = sqlx::query(CAS_HEARTBEAT_PROXY_METADATA_SQL) + .bind(replacement_json) + .bind(now) + .bind(&mutation.node_id) + .bind(&tunnel_generation) + .bind(expected.as_deref()) + .bind(expected.as_deref()) + .bind(expected.as_deref()) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + continue; + } + let Some(after_cas) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if after_cas.tunnel_generation != tunnel_generation { + return Ok(None); + } + tunnel_metrics_sample = build_tunnel_metrics_sample( + current.proxy_metadata.as_ref(), + after_cas.proxy_metadata.as_ref(), + after_cas.active_connections, + after_cas.tunnel_connected, + ); + updated = Some(after_cas); + break; + } + } + let updated = if let Some(updated) = updated { + updated + } else { + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if current.tunnel_generation != tunnel_generation { + return Ok(None); + } + current + }; if let Some(sample) = tunnel_metrics_sample.as_ref() { self.upsert_metrics_bucket( "proxy_node_metrics_1m", - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute), sample, ) .await?; self.upsert_metrics_bucket( "proxy_node_metrics_1h", - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour), sample, ) .await?; for error in &sample.recent_error_events { - log_reported_tunnel_error_event(&node.id, error, now_unix_secs); + log_reported_tunnel_error_event(&updated.id, error, now_unix_secs); let detail = build_tunnel_error_event_detail(error); let event_metadata = serde_json::json!({ "source": "heartbeat", @@ -867,7 +1267,8 @@ WHERE is_manual = 0 "timestamp_unix_ms": error.timestamp_unix_ms, }); self.insert_event( - &node.id, + &updated.id, + Some(tunnel_generation.as_str()), PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, Some(detail.as_str()), Some(&event_metadata), @@ -881,35 +1282,86 @@ WHERE is_manual = 0 } } - Ok(Some(node)) + if reconcile_remote_config_after_heartbeat( + updated.remote_config.as_ref(), + mutation.proxy_version.as_deref(), + ) != updated.remote_config + { + return self + .update_remote_config(&ProxyNodeRemoteConfigMutation { + node_id: mutation.node_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), + node_name: None, + allowed_ports: None, + log_level: None, + heartbeat_interval: None, + scheduling_state: None, + upgrade_to: Some(None), + }) + .await; + } + + Ok(Some(updated)) } async fn record_traffic( &self, mutation: &ProxyNodeTrafficMutation, ) -> Result { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { - return Ok(false); - }; - if !node.is_manual { + let mut tx = self.pool.begin().await.map_sql_err()?; + let lock = sqlx::query("UPDATE proxy_nodes SET id = id WHERE id = ?") + .bind(&mutation.node_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if lock.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; return Ok(false); } - node.total_requests += mutation.total_requests_delta.max(0); - node.failed_requests += mutation.failed_requests_delta.max(0); - node.dns_failures += mutation.dns_failures_delta.max(0); - node.stream_errors += mutation.stream_errors_delta.max(0); - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(true) + let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(&mutation.node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let generation = map_proxy_node_row(&row)?.tunnel_generation; + let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + if expected_generation != generation { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let result = sqlx::query(RECORD_PROXY_NODE_TRAFFIC_SQL) + .bind(mutation.total_requests_delta) + .bind(mutation.failed_requests_delta) + .bind(mutation.dns_failures_delta) + .bind(mutation.stream_errors_delta) + .bind(current_unix_secs() as i64) + .bind(&mutation.node_id) + .bind(expected_generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + let applied = result.rows_affected() > 0; + tx.commit().await.map_sql_err()?; + Ok(applied) } async fn update_tunnel_status( &self, mutation: &ProxyNodeTunnelStatusMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { + let Some(node) = self.find_proxy_node(&mutation.node_id).await? else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } let event_time = mutation .observed_at_unix_secs @@ -926,108 +1378,225 @@ WHERE is_manual = 0 ) }); - if node - .tunnel_connected_at_unix_secs - .is_some_and(|last_transition| event_time < last_transition) - { - self.insert_event( - &mutation.node_id, - event_type, - Some(&format!("[stale_ignored] {event_detail}")), - None, - Some(current_unix_secs()), - ) - .await?; - return Ok(Some(node)); - } - - node.tunnel_connected = mutation.connected; - node.tunnel_connected_at_unix_secs = Some(event_time); - node.status = if mutation.connected { - "online".to_string() - } else { - "offline".to_string() + let event_time_i64 = i64::try_from(event_time).unwrap_or(i64::MAX); + let result = sqlx::query(UPDATE_TUNNEL_STATUS_SQL) + .bind(mutation.connected) + .bind(mutation.connected) + .bind(event_time_i64) + .bind(mutation.connected) + .bind(event_time_i64) + .bind(&mutation.node_id) + .bind(&node.tunnel_generation) + .bind(event_time_i64) + .execute(&self.pool) + .await + .map_sql_err()?; + let Some(current) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); }; - if !mutation.connected { - node.active_connections = 0; + if current.tunnel_generation != node.tunnel_generation { + return Ok(None); } - node.updated_at_unix_secs = Some(event_time); - self.upsert_node(&node).await?; + let stale = result.rows_affected() == 0 + && current + .tunnel_connected_at_unix_secs + .is_some_and(|last_transition| event_time < last_transition); + let persisted_detail = if stale { + format!("[stale_ignored] {event_detail}") + } else { + event_detail + }; self.insert_event( &mutation.node_id, + Some(node.tunnel_generation.as_str()), event_type, - Some(&event_detail), + Some(&persisted_detail), None, - Some(event_time), + Some(if stale { + current_unix_secs() + } else { + event_time + }), ) .await?; - Ok(Some(node)) + Ok(Some(current)) } async fn unregister_node( &self, node_id: &str, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(node_id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + let lock = sqlx::query("UPDATE proxy_nodes SET id = id WHERE id = ?") + .bind(node_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if lock.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; return Ok(None); - }; - let now = Some(current_unix_secs()); - node.status = "offline".to_string(); - node.tunnel_connected = false; - node.active_connections = 0; - node.tunnel_connected_at_unix_secs = now; - node.updated_at_unix_secs = now; - self.upsert_node(&node).await?; - Ok(Some(node)) + } + let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let generation = map_proxy_node_row(&row)?.tunnel_generation; + let now = current_unix_secs() as i64; + sqlx::query(UNREGISTER_PROXY_NODE_SQL) + .bind(now) + .bind(now) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + let updated = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err() + .and_then(|row| map_proxy_node_row(&row)); + let updated = updated?; + tx.commit().await.map_sql_err()?; + Ok(Some(updated)) } async fn delete_node(&self, node_id: &str) -> Result, DataLayerError> { - let existing = self.find_proxy_node(node_id).await?; - if existing.is_some() { - sqlx::query("DELETE FROM proxy_node_events WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_node_metrics_1m WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_node_metrics_1h WHERE node_id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; - sqlx::query("DELETE FROM proxy_nodes WHERE id = ?") - .bind(node_id) - .execute(&self.pool) - .await - .map_sql_err()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + + // SQLite has no SELECT ... FOR UPDATE. A matched no-op UPDATE acquires the + // connection's write lock before we read the generation, so a concurrent + // unregister/re-register cannot interleave with the cleanup below. + let lock = sqlx::query("UPDATE proxy_nodes SET id = id WHERE id = ?") + .bind(node_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if lock.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; + return Ok(None); } - Ok(existing) + + let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let existing = map_proxy_node_row(&row)?; + let generation = existing.tunnel_generation.as_str(); + + // Child tables do not carry the generation themselves. Keep the parent + // identity predicate on every delete so this remains correct if schema + // constraints differ between installations. + sqlx::query( + "DELETE FROM proxy_node_events WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes WHERE id = ? AND tunnel_generation = ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "DELETE FROM proxy_node_metrics_1m WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes WHERE id = ? AND tunnel_generation = ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "DELETE FROM proxy_node_metrics_1h WHERE node_id = ? AND EXISTS (SELECT 1 FROM proxy_nodes WHERE id = ? AND tunnel_generation = ?)", + ) + .bind(node_id) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + + let deleted = sqlx::query("DELETE FROM proxy_nodes WHERE id = ? AND tunnel_generation = ?") + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + if deleted.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + tx.commit().await.map_sql_err()?; + // The parent lock is released before this cleanup. If a flusher already + // claimed one of the rows it can finish (and fail the generation + // predicate) without forming a lock cycle with the delete transaction. + if let Err(error) = sqlx::query(RETIRE_PROXY_NODE_PENDING_COUNTERS_SQL) + .bind(node_id) + .bind(generation) + .execute(&self.pool) + .await + .map_sql_err() + { + tracing::warn!( + node_id = %node_id, + tunnel_generation = %generation, + error = ?error, + "failed to retire deleted proxy node counter rows" + ); + } + Ok(Some(existing)) } async fn update_remote_config( &self, mutation: &ProxyNodeRemoteConfigMutation, ) -> Result, DataLayerError> { - let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else { - return Ok(None); - }; - if node.is_manual { - return Err(DataLayerError::InvalidInput( - "手动节点不支持远程配置下发".to_string(), - )); + for _ in 0..8 { + let Some(node) = self.find_proxy_node(&mutation.node_id).await? else { + return Ok(None); + }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } + if node.is_manual { + return Err(DataLayerError::InvalidInput( + "手动节点不支持远程配置下发".to_string(), + )); + } + + let remote_config = + Self::normalize_remote_config(mutation, node.remote_config.as_ref()); + let remote_config = + optional_json_to_string(&remote_config, "proxy_nodes.remote_config")?; + let now = current_unix_secs() as i64; + let result = sqlx::query(UPDATE_PROXY_NODE_REMOTE_CONFIG_SQL) + .bind(mutation.node_name.as_deref()) + .bind(remote_config) + .bind(now) + .bind(&mutation.node_id) + .bind(&node.tunnel_generation) + .bind(node.config_version) + .execute(&self.pool) + .await + .map_sql_err()?; + if result.rows_affected() == 0 { + continue; + } + + let current = self.find_proxy_node(&mutation.node_id).await?; + return Ok( + current.filter(|current| current.tunnel_generation == node.tunnel_generation) + ); } - if let Some(node_name) = mutation.node_name.as_ref() { - node.name = node_name.clone(); - } - node.remote_config = Self::normalize_remote_config(mutation, node.remote_config.as_ref()); - node.config_version = node.config_version.saturating_add(1); - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await?; - Ok(Some(node)) + + Err(DataLayerError::UnexpectedValue( + "proxy node remote config changed during every CAS retry".to_string(), + )) } async fn increment_manual_node_requests( @@ -1037,23 +1606,34 @@ WHERE is_manual = 0 failed_delta: i64, latency_ms: Option, ) -> Result<(), DataLayerError> { - let Some(mut node) = self.find_proxy_node(node_id).await? else { - return Ok(()); - }; - if !node.is_manual { + let mut tx = self.pool.begin().await.map_sql_err()?; + let lock = sqlx::query("UPDATE proxy_nodes SET id = id WHERE id = ?") + .bind(node_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if lock.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; return Ok(()); } - if total_delta > 0 { - node.total_requests += total_delta; - } - if failed_delta > 0 { - node.failed_requests += failed_delta; - } - if let Some(ms) = latency_ms { - node.avg_latency_ms = Some(ms as f64); - } - node.updated_at_unix_secs = Some(current_unix_secs()); - self.upsert_node(&node).await + let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(node_id) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let generation = map_proxy_node_row(&row)?.tunnel_generation; + sqlx::query(INCREMENT_MANUAL_PROXY_NODE_REQUESTS_SQL) + .bind(total_delta) + .bind(failed_delta) + .bind(latency_ms.map(|value| value as f64)) + .bind(current_unix_secs() as i64) + .bind(node_id) + .bind(generation) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(()) } async fn cleanup_proxy_node_metrics( @@ -1152,6 +1732,37 @@ fn duplicate_proxy_node_error(node: &StoredProxyNode) -> DataLayerError { )) } +fn requested_proxy_node_id(value: Option<&str>) -> Result, DataLayerError> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_empty() || value.trim() != value { + return Err(DataLayerError::InvalidInput( + "proxy node id must be non-empty and unpadded".to_string(), + )); + } + Ok(Some(value.to_string())) +} + +fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node registration identity changed: requested {requested_id}, existing {existing_id}" + )) +} + +fn proxy_node_registration_changed_error() -> DataLayerError { + DataLayerError::UnexpectedValue( + "registered proxy node identity changed during registration".to_string(), + ) +} + +fn proxy_node_id_in_use_error(node: &StoredProxyNode) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node id is already in use: {} ({}:{})", + node.id, node.ip, node.port + )) +} + fn optional_json_from_string( value: Option, field_name: &str, @@ -1168,6 +1779,12 @@ fn optional_json_from_string( } fn map_proxy_node_row(row: &SqliteRow) -> Result { + let tunnel_generation: String = row.try_get("tunnel_generation").map_sql_err()?; + if tunnel_generation.trim().is_empty() { + return Err(DataLayerError::UnexpectedValue( + "proxy_nodes.tunnel_generation must not be empty".to_string(), + )); + } Ok(StoredProxyNode::new( row.try_get("id").map_sql_err()?, row.try_get("name").map_sql_err()?, @@ -1185,6 +1802,7 @@ fn map_proxy_node_row(row: &SqliteRow) -> Result success_count += 1, + Err(error) if error.to_string().contains("identity changed") => { + identity_error_count += 1 + } + Err(error) => panic!("unexpected concurrent registration error: {error}"), + } + } + assert_eq!(success_count, 1); + assert_eq!(identity_error_count, 1); + let endpoint_rows: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM proxy_nodes WHERE ip = ? AND port = ?") + .bind("127.0.0.42") + .bind(7042_i32) + .fetch_one(&pool) + .await + .expect("endpoint identity count should read"); + assert_eq!(endpoint_rows, 1); + + drop(repository); + drop(pool); + let _ = std::fs::remove_file(database_path); + } #[tokio::test] async fn sqlite_repository_reads_proxy_nodes_and_events() { @@ -1303,12 +2263,12 @@ mod tests { sqlx::query( r#" INSERT INTO proxy_nodes ( - id, name, ip, port, status, heartbeat_interval, active_connections, + id, tunnel_generation, name, ip, port, status, heartbeat_interval, active_connections, total_requests, failed_requests, dns_failures, stream_errors, tunnel_mode, tunnel_connected, config_version, proxy_metadata, hardware_info, remote_config, created_at, updated_at ) VALUES ( - 'node-1', 'Node 1', '127.0.0.1', 8080, 'online', 30, 1, + 'node-1', 'test-generation-node-1', 'Node 1', '127.0.0.1', 8080, 'online', 30, 1, 10, 2, 1, 0, 1, 1, 3, '{"version":"1.0.0"}', '{"cpu":"m1"}', '{"log_level":"debug"}', 1, 2 ) @@ -1368,6 +2328,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let repository = SqliteProxyNodeReadRepository::new(pool); let manual = repository .create_manual_node(&ProxyNodeManualCreateMutation { + node_id: Some("manual-1-fixed-id".to_string()), name: "manual-1".to_string(), ip: "127.0.0.2".to_string(), port: 8081, @@ -1380,7 +2341,28 @@ VALUES ('node-1', 'registered', 'ok', 3) .await .expect("manual node should create"); assert!(manual.is_manual); + assert_eq!(manual.id, "manual-1-fixed-id"); assert_eq!(manual.status, "online"); + assert!(repository + .create_manual_node(&ProxyNodeManualCreateMutation { + node_id: Some("manual-1-fixed-id".to_string()), + name: "manual-id-collision".to_string(), + ip: "127.0.0.3".to_string(), + port: 8082, + region: None, + proxy_url: "http://127.0.0.3:8082".to_string(), + proxy_username: Some("attacker".to_string()), + proxy_password: Some("replacement-pass".to_string()), + registered_by: None, + }) + .await + .is_err()); + let after_collision = repository + .find_proxy_node("manual-1-fixed-id") + .await + .expect("manual node should reload") + .expect("manual node should remain"); + assert_eq!(after_collision.proxy_password.as_deref(), Some("pass")); let manual = repository .update_manual_node(&ProxyNodeManualUpdateMutation { @@ -1402,6 +2384,7 @@ VALUES ('node-1', 'registered', 'ok', 3) assert!(repository .record_traffic(&ProxyNodeTrafficMutation { node_id: manual.id.clone(), + expected_tunnel_generation: Some(manual.tunnel_generation.clone()), total_requests_delta: 5, failed_requests_delta: 1, dns_failures_delta: 1, @@ -1424,6 +2407,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let registered = repository .register_node(&ProxyNodeRegistrationMutation { + node_id: Some("tunnel-1-fixed-id".to_string()), name: "tunnel-1".to_string(), ip: "10.0.0.1".to_string(), port: 7000, @@ -1434,7 +2418,13 @@ VALUES ('node-1', 'registered', 'ok', 3) avg_latency_ms: Some(12.5), hardware_info: Some(json!({"cpu":"m1"})), estimated_max_concurrency: Some(100), - proxy_metadata: Some(json!({"arch":"arm64"})), + proxy_metadata: Some(json!({ + "arch":"arm64", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:trusted-registration-key" + } + })), proxy_version: Some("1.0.0".to_string()), registered_by: Some("proxy".to_string()), tunnel_mode: true, @@ -1442,6 +2432,7 @@ VALUES ('node-1', 'registered', 'ok', 3) .await .expect("tunnel node should register"); assert!(!registered.is_manual); + assert_eq!(registered.id, "tunnel-1-fixed-id"); assert!(registered.tunnel_mode); assert_eq!( registered @@ -1455,6 +2446,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let configured = repository .update_remote_config(&ProxyNodeRemoteConfigMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, node_name: Some("tunnel-renamed".to_string()), allowed_ports: Some(vec![443, 8443]), log_level: Some("debug".to_string()), @@ -1471,6 +2463,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let heartbeat = repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, heartbeat_interval: Some(45), active_connections: Some(4), total_requests_delta: Some(6), @@ -1478,7 +2471,13 @@ VALUES ('node-1', 'registered', 'ok', 3) failed_requests_delta: Some(1), dns_failures_delta: Some(0), stream_errors_delta: Some(2), - proxy_metadata: Some(json!({"arch":"arm64"})), + proxy_metadata: Some(json!({ + "arch":"arm64", + "tunnel_security": { + "mode": "disabled", + "encryption_key": "heartbeat-attacker-controlled" + } + })), proxy_version: Some("2.0.0".to_string()), }) .await @@ -1489,6 +2488,14 @@ VALUES ('node-1', 'registered', 'ok', 3) assert_eq!(heartbeat.active_connections, 4); assert_eq!(heartbeat.total_requests, 16); assert_eq!(heartbeat.config_version, 2); + assert_eq!( + heartbeat + .proxy_metadata + .as_ref() + .and_then(|value| value.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:trusted-registration-key") + ); assert!(heartbeat .remote_config .as_ref() @@ -1498,6 +2505,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let stale = repository .update_tunnel_status(&ProxyNodeTunnelStatusMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, connected: false, conn_count: 0, detail: None, @@ -1511,6 +2519,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let disconnected = repository .update_tunnel_status(&ProxyNodeTunnelStatusMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, connected: false, conn_count: 0, detail: Some("closed".to_string()), @@ -1571,6 +2580,7 @@ VALUES ('node-1', 'registered', 'ok', 3) let repository = SqliteProxyNodeReadRepository::new(pool); let registered = repository .register_node(&ProxyNodeRegistrationMutation { + node_id: None, name: "tunnel-1".to_string(), ip: "10.0.0.1".to_string(), port: 7000, @@ -1592,6 +2602,7 @@ VALUES ('node-1', 'registered', 'ok', 3) repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, heartbeat_interval: Some(30), active_connections: Some(0), total_requests_delta: None, @@ -1619,6 +2630,7 @@ VALUES ('node-1', 'registered', 'ok', 3) repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, heartbeat_interval: Some(30), active_connections: Some(5), total_requests_delta: None, @@ -1747,4 +2759,173 @@ VALUES ('node-1', 'registered', 'ok', 3) assert_eq!(cleanup.deleted_1m_rows, 1); assert_eq!(cleanup.deleted_1h_rows, 1); } + + #[tokio::test] + async fn sqlite_registration_preserves_omitted_security_and_allows_rotation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteProxyNodeReadRepository::new(pool); + + let first = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some("registration-security-node".to_string()), + name: "registration-security-node".to_string(), + ip: "127.0.0.70".to_string(), + port: 7070, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "version": "1.0.0", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + })), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("first registration should succeed"); + + let refreshed = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some(first.id.clone()), + name: "registration-security-node-refreshed".to_string(), + ip: "127.0.0.70".to_string(), + port: 7070, + region: None, + heartbeat_interval: 45, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({"runtime": "refreshed"})), + proxy_version: Some("2.0.0".to_string()), + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("metadata-only re-registration should succeed"); + assert_eq!(refreshed.id, first.id); + assert_eq!( + refreshed + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old") + ); + assert_eq!( + refreshed + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.get("runtime")), + Some(&json!("refreshed")) + ); + + let rotated = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some(first.id.clone()), + name: "registration-security-node-rotated".to_string(), + ip: "127.0.0.70".to_string(), + port: 7070, + region: None, + heartbeat_interval: 45, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new" + } + })), + proxy_version: Some("2.1.0".to_string()), + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("explicit security rotation should succeed"); + assert_eq!( + rotated + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new") + ); + + let stale_refresh = ProxyNodeRegistrationMutation { + node_id: Some(first.id.clone()), + name: "registration-security-node-stale".to_string(), + ip: "127.0.0.70".to_string(), + port: 7070, + region: None, + heartbeat_interval: 45, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({"runtime": "stale-writer"})), + proxy_version: Some("2.2.0".to_string()), + registered_by: None, + tunnel_mode: true, + }; + let stale_replacement = merge_proxy_metadata_for_registration( + refreshed.proxy_metadata.as_ref(), + normalize_proxy_metadata( + stale_refresh.proxy_metadata.as_ref(), + stale_refresh.proxy_version.as_deref(), + ), + ); + assert!(!repository + .update_existing_registration_if_unchanged( + &stale_refresh, + &refreshed, + stale_replacement.as_ref(), + super::current_unix_secs(), + ) + .await + .expect("stale registration CAS should execute")); + + let committed_refresh = repository + .register_node(&ProxyNodeRegistrationMutation { + name: "registration-security-node-committed".to_string(), + proxy_metadata: Some(json!({"runtime": "committed-after-rotation"})), + ..stale_refresh + }) + .await + .expect("metadata refresh should retry from current security state"); + assert_eq!( + committed_refresh + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(serde_json::Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new") + ); + assert_eq!( + committed_refresh + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.get("runtime")), + Some(&json!("committed-after-rotation")) + ); + } } diff --git a/crates/aether-data/adapters/sqlite/src/routing_profiles.rs b/crates/aether-data/adapters/sqlite/src/routing_profiles.rs index 22bcdf8ac..e1fa9778b 100644 --- a/crates/aether-data/adapters/sqlite/src/routing_profiles.rs +++ b/crates/aether-data/adapters/sqlite/src/routing_profiles.rs @@ -15,6 +15,7 @@ SELECT description, enabled, is_system_default, + sort_order, config_json, version, created_at, @@ -61,10 +62,12 @@ impl SqliteRoutingGroupRepository { #[async_trait] impl RoutingGroupReadRepository for SqliteRoutingGroupRepository { async fn list_routing_groups(&self) -> Result, DataLayerError> { - let rows = sqlx::query(&format!("{ROUTING_GROUP_SELECT} ORDER BY name ASC, id ASC")) - .fetch_all(&self.pool) - .await - .map_sql_err()?; + let rows = sqlx::query(&format!( + "{ROUTING_GROUP_SELECT} ORDER BY enabled DESC, sort_order ASC, name ASC, id ASC" + )) + .fetch_all(&self.pool) + .await + .map_sql_err()?; rows.iter().map(map_group_row).collect() } @@ -168,10 +171,10 @@ impl RoutingGroupWriteRepository for SqliteRoutingGroupRepository { sqlx::query( r#" INSERT INTO routing_groups ( - id, name, description, enabled, is_system_default, config_json, + id, name, description, enabled, is_system_default, sort_order, config_json, version, created_at, updated_at, published_at ) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&group.id) @@ -179,6 +182,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(json_to_string( &group.config_json, "routing_groups.config_json", @@ -229,6 +233,7 @@ SET name = ?, description = ?, enabled = ?, is_system_default = ?, + sort_order = ?, config_json = ?, version = ?, updated_at = ?, @@ -240,6 +245,7 @@ WHERE id = ? .bind(&group.description) .bind(group.enabled) .bind(group.is_system_default) + .bind(group.sort_order) .bind(json_to_string( &group.config_json, "routing_groups.config_json", @@ -438,6 +444,7 @@ fn map_group_row(row: &SqliteRow) -> Result description: row.try_get("description").map_sql_err()?, enabled: row.try_get("enabled").map_sql_err()?, is_system_default: row.try_get("is_system_default").map_sql_err()?, + sort_order: row.try_get("sort_order").map_sql_err()?, config_json: json_from_string( row.try_get("config_json").map_sql_err()?, "routing_groups.config_json", @@ -514,6 +521,7 @@ mod tests { description: Some("initial".to_string()), enabled: true, is_system_default: true, + sort_order: 0, config_json: json!({"allowed_models": ["gpt-*"]}), version: 1, created_at: 10, @@ -790,6 +798,7 @@ SET is_default = 1, description: None, enabled: true, is_system_default, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, diff --git a/crates/aether-data/adapters/sqlite/src/settlement.rs b/crates/aether-data/adapters/sqlite/src/settlement.rs index ef77b96d5..6be6bfa54 100644 --- a/crates/aether-data/adapters/sqlite/src/settlement.rs +++ b/crates/aether-data/adapters/sqlite/src/settlement.rs @@ -3,8 +3,12 @@ use sqlx::{sqlite::SqliteRow, Row}; use aether_data_contracts::repository::settlement::{ finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd, - settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement, - UsageSettlementInput, SETTLEMENT_EPSILON_USD, + settlement_billing_status_for_usage_status, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation, + StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState, + UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD, }; use aether_data_contracts::DataLayerError; @@ -127,6 +131,128 @@ impl SqliteSettlementRepository { } } +fn usage_policy_cost_i64(value: u64, field: &str) -> Result { + i64::try_from(value) + .map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds the integer range"))) +} + +fn usage_policy_cost_u64(value: i64, field: &str) -> Result { + u64::try_from(value) + .map_err(|_| DataLayerError::UnexpectedValue(format!("{field} must not be negative"))) +} + +fn usage_policy_request_admission_from_sqlite_row( + row: &SqliteRow, +) -> Result { + let state: String = row.try_get("state").map_sql_err()?; + Ok(StoredUsagePolicyRequestAdmission { + request_id: row.try_get("request_id").map_sql_err()?, + subject_id: row.try_get("subject_id").map_sql_err()?, + event_token: row.try_get("event_token").map_sql_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at_unix_secs").map_sql_err()?, + "usage policy request admitted_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until_unix_secs").map_sql_err()?, + "usage policy request retain_until", + )?, + state: UsagePolicyRequestAdmissionState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy request admission state {state}" + )) + })?, + released_at_unix_secs: row + .try_get::, _>("released_at_unix_secs") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy request released_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL: &str = r#" +SELECT request_id, subject_id, event_token, + admitted_at AS admitted_at_unix_secs, + retain_until AS retain_until_unix_secs, + state, released_at AS released_at_unix_secs +FROM usage_request_admissions +WHERE event_token = ? +"#; + +const INSERT_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL: &str = r#" +INSERT INTO usage_request_admissions ( + request_id, subject_id, event_token, admitted_at, retain_until, + state, released_at, created_at +) VALUES (?, ?, ?, ?, ?, 'active', NULL, ?) +ON CONFLICT(event_token) DO NOTHING +"#; + +fn usage_policy_cost_reservation_from_sqlite_row( + row: &SqliteRow, +) -> Result { + let state: String = row.try_get("state").map_sql_err()?; + Ok(StoredUsagePolicyCostReservation { + request_id: row.try_get("request_id").map_sql_err()?, + subject_id: row.try_get("subject_id").map_sql_err()?, + reservation_token: row.try_get("reservation_token").map_sql_err()?, + admitted_at_unix_secs: usage_policy_cost_u64( + row.try_get("admitted_at").map_sql_err()?, + "usage policy admitted_at", + )?, + reserved_cost_units: usage_policy_cost_u64( + row.try_get("reserved_cost_units").map_sql_err()?, + "usage policy reserved_cost_units", + )?, + actual_cost_units: row + .try_get::, _>("actual_cost_units") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy actual_cost_units")) + .transpose()?, + state: UsagePolicyCostReservationState::parse(&state).ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "unknown usage policy reservation state {state}" + )) + })?, + reservation_expires_at_unix_secs: usage_policy_cost_u64( + row.try_get("reservation_expires_at").map_sql_err()?, + "usage policy reservation_expires_at", + )?, + retain_until_unix_secs: usage_policy_cost_u64( + row.try_get("retain_until").map_sql_err()?, + "usage policy retain_until", + )?, + finalized_at_unix_secs: row + .try_get::, _>("finalized_at") + .map_sql_err()? + .map(|value| usage_policy_cost_u64(value, "usage policy finalized_at")) + .transpose()?, + }) +} + +const FIND_USAGE_POLICY_COST_RESERVATION_SQLITE_SQL: &str = r#" +SELECT request_id, subject_id, reservation_token, admitted_at, + reserved_cost_units, actual_cost_units, state, + reservation_expires_at, retain_until, finalized_at +FROM usage_cost_reservations +WHERE reservation_token = ? +"#; + +async fn lock_usage_policy_subject_sqlite( + tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, + subject_id: &str, +) -> Result { + let result = sqlx::query("UPDATE users SET updated_at = updated_at WHERE id = ?") + .bind(subject_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + Ok(result.rows_affected() > 0) +} + +fn usage_policy_subject_missing() -> DataLayerError { + DataLayerError::InvalidInput("usage policy subject does not exist".to_string()) +} + fn settlement_from_row(row: &SqliteRow) -> Result { Ok(StoredUsageSettlement { request_id: row.try_get("request_id").map_sql_err()?, @@ -219,6 +345,7 @@ fn daily_quota_usage_date( fn daily_quota_grants_from_entitlement( entitlement_id: &str, entitlements: &serde_json::Value, + current_allow_wallet_overage: Option, now: chrono::DateTime, ) -> Result, DataLayerError> { let mut grants = Vec::new(); @@ -244,15 +371,27 @@ fn daily_quota_grants_from_entitlement( .and_then(serde_json::Value::as_str), now, )?, - allow_wallet_overage: item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false), + allow_wallet_overage: current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }), }); } Ok(grants) } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + async fn consume_daily_quota_sqlite( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, user_id: &str, @@ -262,18 +401,29 @@ async fn consume_daily_quota_sqlite( wallet_can_overdraft: bool, now_unix_secs: i64, ) -> Result { - if total_cost_usd <= 0.0 { + if !total_cost_usd.is_finite() || total_cost_usd < 0.0 { + return Err(DataLayerError::InvalidInput( + "daily quota settlement cost must be finite and non-negative".to_string(), + )); + } + if total_cost_usd == 0.0 { return Ok(DailyQuotaDebitResult::default()); } let rows = sqlx::query( r#" -SELECT id, entitlements_snapshot +SELECT + user_plan_entitlements.id, + user_plan_entitlements.entitlements_snapshot, + billing_plans.entitlements_json AS plan_entitlements_json FROM user_plan_entitlements -WHERE user_id = ? - AND status = 'active' - AND starts_at <= ? - AND expires_at > ? -ORDER BY expires_at ASC, created_at ASC, id ASC +JOIN billing_plans ON billing_plans.id = user_plan_entitlements.plan_id +WHERE user_plan_entitlements.user_id = ? + AND user_plan_entitlements.status = 'active' + AND user_plan_entitlements.starts_at <= ? + AND user_plan_entitlements.expires_at > ? +ORDER BY user_plan_entitlements.expires_at ASC, + user_plan_entitlements.created_at ASC, + user_plan_entitlements.id ASC "#, ) .bind(user_id) @@ -293,9 +443,17 @@ ORDER BY expires_at ASC, created_at ASC, id ASC "user_plan_entitlements.entitlements_snapshot invalid json: {err}" )) })?; + let plan_entitlements_raw: String = row.try_get("plan_entitlements_json").map_sql_err()?; + let plan_entitlements = serde_json::from_str::(&plan_entitlements_raw) + .map_err(|err| { + DataLayerError::UnexpectedValue(format!( + "billing_plans.entitlements_json invalid json: {err}" + )) + })?; grants.extend(daily_quota_grants_from_entitlement( &entitlement_id, &entitlements, + daily_quota_wallet_overage_policy(&plan_entitlements), now, )?); } @@ -321,8 +479,18 @@ WHERE user_entitlement_id = ? .fetch_one(&mut **tx) .await .map_sql_err()?; + if !used.is_finite() || used < 0.0 { + return Err(DataLayerError::UnexpectedValue( + "daily quota usage ledger total is invalid".to_string(), + )); + } let remaining = (grant.daily_quota_usd - used).max(0.0); total_remaining += remaining; + if !total_remaining.is_finite() { + return Err(DataLayerError::UnexpectedValue( + "daily quota remaining total overflowed".to_string(), + )); + } grants_with_remaining.push((grant, remaining)); } let insufficient = (!allow_wallet_overage && total_remaining + 0.000_000_01 < total_cost_usd) @@ -372,6 +540,469 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) #[async_trait] impl SettlementWriteRepository for SqliteSettlementRepository { + async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result { + input.validate()?; + let now = now_unix_secs()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + // This no-op update acquires SQLite's single writer slot before any admission reads. + if !lock_usage_policy_subject_sqlite(&mut tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + let existing_row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL) + .bind(&input.event_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if let Some(row) = existing_row.as_ref() { + let existing = usage_policy_request_admission_from_sqlite_row(row)?; + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at".to_string(), + )); + } + sqlx::query( + "UPDATE usage_request_admissions SET retain_until = MAX(retain_until, ?) WHERE event_token = ?", + ) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + let outcome = match existing.state { + UsagePolicyRequestAdmissionState::Active => { + ReserveUsagePolicyRequestOutcome::Allowed + } + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }; + tx.commit().await.map_sql_err()?; + return Ok(outcome); + } + + for (window_index, window) in input.windows.iter().enumerate() { + let used_requests = sqlx::query_scalar::<_, i64>( + r#" +SELECT COUNT(*) +FROM usage_request_admissions +WHERE subject_id = ? + AND state = 'active' + AND admitted_at >= ? + AND admitted_at < ? + "#, + ) + .bind(&input.subject_id) + .bind(usage_policy_cost_i64( + window.starts_at_unix_secs, + "usage policy request window start", + )?) + .bind(usage_policy_cost_i64( + window.ends_at_unix_secs, + "usage policy request window end", + )?) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let used_requests = + usage_policy_cost_u64(used_requests, "usage policy request used_requests")?; + if used_requests >= window.limit_requests { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Rejected { + window_index, + limit_requests: window.limit_requests, + used_requests, + }); + } + } + + let insert_result = sqlx::query(INSERT_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.event_token) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy request admitted_at", + )?) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + if insert_result.rows_affected() == 1 { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Allowed); + } + + // The writer lock above normally makes this branch unreachable for concurrent reserves, + // but classify the unique-token race explicitly so future lock changes cannot surface a + // raw SQLite constraint error or accidentally reactivate a released tombstone. + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL) + .bind(&input.event_token) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let existing = usage_policy_request_admission_from_sqlite_row(&row)?; + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at".to_string(), + )); + } + sqlx::query( + "UPDATE usage_request_admissions SET retain_until = MAX(retain_until, ?) WHERE event_token = ?", + ) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy request retain_until", + )?) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + let outcome = match existing.state { + UsagePolicyRequestAdmissionState::Active => ReserveUsagePolicyRequestOutcome::Allowed, + UsagePolicyRequestAdmissionState::Released => { + ReserveUsagePolicyRequestOutcome::AlreadyReleased + } + }; + tx.commit().await.map_sql_err()?; + Ok(outcome) + } + + async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, DataLayerError> { + input.validate()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_sqlite(&mut tx, &input.subject_id).await? { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL) + .bind(&input.event_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.commit().await.map_sql_err()?; + return Ok(None); + }; + let mut admission = usage_policy_request_admission_from_sqlite_row(&row)?; + if admission.request_id != input.request_id || admission.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + if input.released_at_unix_secs < admission.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy released_at must not precede admitted_at".to_string(), + )); + } + if admission.state == UsagePolicyRequestAdmissionState::Active { + sqlx::query( + "UPDATE usage_request_admissions SET state = 'released', released_at = ? WHERE event_token = ? AND state = 'active'", + ) + .bind(usage_policy_cost_i64( + input.released_at_unix_secs, + "usage policy request released_at", + )?) + .bind(&input.event_token) + .execute(&mut *tx) + .await + .map_sql_err()?; + admission.state = UsagePolicyRequestAdmissionState::Released; + admission.released_at_unix_secs = Some(input.released_at_unix_secs); + } + tx.commit().await.map_sql_err()?; + Ok(Some(admission)) + } + + async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy request cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_request_admissions +WHERE rowid IN ( + SELECT rowid + FROM usage_request_admissions + WHERE retain_until <= ? + ORDER BY retain_until, event_token + LIMIT ? +) + "#, + ) + .bind(now) + .bind(limit) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() as usize) + } + + async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result { + input.validate()?; + let now = now_unix_secs()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_sqlite(&mut tx, &input.subject_id).await? { + return Err(usage_policy_subject_missing()); + } + let existing_row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_SQLITE_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let existing = existing_row + .as_ref() + .map(usage_policy_cost_reservation_from_sqlite_row) + .transpose()?; + if let Some(existing) = existing.as_ref() { + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyCostOutcome::Conflict); + } + if existing.state != UsagePolicyCostReservationState::Reserved { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: existing.state, + }); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy reservation_token must keep its original admitted_at".to_string(), + )); + } + } + + let previous_reserved_cost_units = existing + .as_ref() + .map(|reservation| reservation.reserved_cost_units) + .unwrap_or(0); + let target_reserved_cost_units = + previous_reserved_cost_units.max(input.reserved_cost_units); + for (window_index, window) in input.windows.iter().enumerate() { + let used_cost_units = sqlx::query_scalar::<_, i64>( + r#" +SELECT COALESCE(SUM( + CASE + WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) + WHEN state = 'reserved' AND reservation_expires_at > ? THEN reserved_cost_units + ELSE 0 + END +), 0) +FROM usage_cost_reservations +WHERE subject_id = ? + AND admitted_at >= ? + AND admitted_at < ? + AND reservation_token <> ? + "#, + ) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy admitted_at", + )?) + .bind(&input.subject_id) + .bind(usage_policy_cost_i64( + window.starts_at_unix_secs, + "usage policy window start", + )?) + .bind(usage_policy_cost_i64( + window.ends_at_unix_secs, + "usage policy window end", + )?) + .bind(&input.reservation_token) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + let used_cost_units = + usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?; + if used_cost_units + .checked_add(target_reserved_cost_units) + .is_none_or(|total| total > window.limit_cost_units) + { + tx.commit().await.map_sql_err()?; + return Ok(ReserveUsagePolicyCostOutcome::Rejected { + window_index, + limit_cost_units: window.limit_cost_units, + used_cost_units, + }); + } + } + + sqlx::query( + r#" +INSERT INTO usage_cost_reservations ( + request_id, subject_id, reservation_token, admitted_at, + reserved_cost_units, actual_cost_units, + state, reservation_expires_at, retain_until, finalized_at, created_at, updated_at +) VALUES (?, ?, ?, ?, ?, NULL, 'reserved', ?, ?, NULL, ?, ?) +ON CONFLICT (reservation_token) DO UPDATE SET + reserved_cost_units = MAX( + usage_cost_reservations.reserved_cost_units, + excluded.reserved_cost_units + ), + reservation_expires_at = MAX( + usage_cost_reservations.reservation_expires_at, + excluded.reservation_expires_at + ), + retain_until = MAX( + usage_cost_reservations.retain_until, + excluded.retain_until + ), + updated_at = excluded.updated_at + "#, + ) + .bind(&input.request_id) + .bind(&input.subject_id) + .bind(&input.reservation_token) + .bind(usage_policy_cost_i64( + input.admitted_at_unix_secs, + "usage policy admitted_at", + )?) + .bind(usage_policy_cost_i64( + target_reserved_cost_units, + "usage policy reserved_cost_units", + )?) + .bind(usage_policy_cost_i64( + input.reservation_expires_at_unix_secs, + "usage policy reservation_expires_at", + )?) + .bind(usage_policy_cost_i64( + input.retain_until_unix_secs, + "usage policy retain_until", + )?) + .bind(now) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: target_reserved_cost_units, + additional_reserved_cost_units: target_reserved_cost_units + .saturating_sub(previous_reserved_cost_units), + }) + } + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + input.validate()?; + let now = now_unix_secs()?; + let mut tx = self.pool.begin().await.map_sql_err()?; + if !lock_usage_policy_subject_sqlite(&mut tx, &input.subject_id).await? { + tx.commit().await.map_sql_err()?; + return Ok(None); + } + let row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_SQLITE_SQL) + .bind(&input.reservation_token) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.commit().await.map_sql_err()?; + return Ok(None); + }; + let mut reservation = usage_policy_cost_reservation_from_sqlite_row(&row)?; + if reservation.request_id != input.request_id || reservation.subject_id != input.subject_id + { + // The token selects the row; audit identity must still match before the reservation + // can be finalized. + tx.commit().await.map_sql_err()?; + return Ok(None); + } + if reservation.state == UsagePolicyCostReservationState::Reserved { + sqlx::query( + r#" +UPDATE usage_cost_reservations +SET state = ?, actual_cost_units = ?, finalized_at = ?, updated_at = ? +WHERE reservation_token = ? + AND request_id = ? + AND subject_id = ? + AND state = 'reserved' + "#, + ) + .bind(input.terminal_state.as_str()) + .bind(usage_policy_cost_i64( + input.actual_cost_units, + "usage policy actual_cost_units", + )?) + .bind(usage_policy_cost_i64( + input.finalized_at_unix_secs, + "usage policy finalized_at", + )?) + .bind(now) + .bind(&input.reservation_token) + .bind(&input.request_id) + .bind(&input.subject_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + reservation.state = input.terminal_state; + reservation.actual_cost_units = Some(input.actual_cost_units); + reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs); + } + tx.commit().await.map_sql_err()?; + Ok(Some(reservation)) + } + + async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let now = usage_policy_cost_i64(now_unix_secs, "usage policy cleanup timestamp")?; + let limit = i64::try_from(batch_size).unwrap_or(i64::MAX); + let result = sqlx::query( + r#" +DELETE FROM usage_cost_reservations +WHERE rowid IN ( + SELECT rowid + FROM usage_cost_reservations + WHERE retain_until <= ? + ORDER BY retain_until, reservation_token + LIMIT ? +) + "#, + ) + .bind(now) + .bind(limit) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() as usize) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -457,7 +1088,7 @@ LIMIT 1 let wallet_row = if let Some(api_key_id) = api_key_id { sqlx::query( r#" -SELECT id, balance, gift_balance, limit_mode +SELECT id, balance, gift_balance, total_consumed, limit_mode FROM wallets WHERE api_key_id = ? LIMIT 1 @@ -477,7 +1108,7 @@ LIMIT 1 if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) { sqlx::query( r#" -SELECT id, balance, gift_balance, limit_mode +SELECT id, balance, gift_balance, total_consumed, limit_mode FROM wallets WHERE user_id = ? LIMIT 1 @@ -497,14 +1128,20 @@ LIMIT 1 let wallet_can_overdraft = wallet_row.is_some(); let wallet_available_usd = match wallet_row.as_ref() { Some(row) => { + let recharge_balance = sqlite_real(row, "balance")?; + let gift_balance = sqlite_real(row, "gift_balance")?; + let total_consumed = sqlite_real(row, "total_consumed")?; + validate_wallet_settlement_values( + recharge_balance, + gift_balance, + total_consumed, + 0.0, + )?; let limit_mode: String = row.try_get("limit_mode").map_sql_err()?; if limit_mode.eq_ignore_ascii_case("unlimited") { None } else { - Some(finite_wallet_available_usd( - sqlite_real(row, "balance")?, - sqlite_real(row, "gift_balance")?, - )) + Some(finite_wallet_available_usd(recharge_balance, gift_balance)) } } None => Some(0.0), @@ -583,6 +1220,7 @@ LIMIT 1 let wallet_id: String = wallet_row.try_get("id").map_sql_err()?; let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; + let total_consumed = sqlite_real(&wallet_row, "total_consumed")?; let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?; let before_total = before_recharge + before_gift; let mut after_recharge = before_recharge; @@ -596,6 +1234,13 @@ LIMIT 1 (after_recharge, after_gift) = debit_plan.after_balances(before_recharge, before_gift); } + let total_consumed_after = total_consumed + wallet_debit_cost_usd; + validate_wallet_settlement_values( + after_recharge, + after_gift, + total_consumed_after, + 0.0, + )?; if final_billing_status == "settled" { sqlx::query( r#" @@ -603,14 +1248,14 @@ UPDATE wallets SET balance = ?, gift_balance = ?, - total_consumed = COALESCE(total_consumed, 0) + ?, + total_consumed = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) .bind(after_gift) - .bind(wallet_debit_cost_usd) + .bind(total_consumed_after) .bind(updated_at) .bind(&wallet_id) .execute(&mut *tx) @@ -709,11 +1354,15 @@ WHERE id = ? #[cfg(test)] mod tests { - use super::SqliteSettlementRepository; - use crate::run_migrations; + use super::{SqliteSettlementRepository, INSERT_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL}; + use crate::{run_migrations, SqliteUserReadRepository}; use aether_data_contracts::repository::settlement::{ - SettlementWriteRepository, UsageSettlementInput, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput, SettlementWriteRepository, + UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestWindow, + UsageSettlementInput, }; + use aether_data_contracts::repository::users::UserReadRepository; use sqlx::Row; use std::time::Duration; @@ -824,6 +1473,219 @@ WHERE request_id = 'request-1' ); } + #[tokio::test] + async fn sqlite_settlement_rejects_corrupt_wallet_before_financial_mutation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_settlement_rows(&pool).await; + sqlx::query("UPDATE wallets SET balance = ? WHERE id = 'wallet-1'") + .bind(f64::INFINITY) + .execute(&pool) + .await + .expect("corrupt wallet fixture should update"); + + let result = SqliteSettlementRepository::new(pool.clone()) + .settle_usage(UsageSettlementInput { + request_id: "request-1".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: None, + api_key_is_standalone: false, + provider_id: Some("provider-1".to_string()), + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 3.0, + actual_total_cost_usd: 6.0, + finalized_at_unix_secs: Some(1_234), + }) + .await; + assert!(result.is_err()); + + let billing_status: String = + sqlx::query_scalar("SELECT billing_status FROM usage WHERE request_id = 'request-1'") + .fetch_one(&pool) + .await + .expect("usage should load"); + assert_eq!(billing_status, "pending"); + let settlement_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_settlement_snapshots WHERE request_id = 'request-1'", + ) + .fetch_one(&pool) + .await + .expect("settlement snapshots should count"); + assert_eq!(settlement_count, 0); + let delta_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = 'request-1'", + ) + .fetch_one(&pool) + .await + .expect("usage deltas should count"); + assert_eq!(delta_count, 0); + } + + #[tokio::test] + async fn request_admission_insert_defensively_preserves_the_existing_token() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO users (id, username, auth_source, created_at, updated_at) +VALUES + ('defensive-user-1', 'defensive-user-1', 'local', 1, 1), + ('defensive-user-2', 'defensive-user-2', 'local', 1, 1) + "#, + ) + .execute(&pool) + .await + .expect("usage policy subjects should insert"); + + let insert = |request_id: &'static str, subject_id: &'static str| { + sqlx::query(INSERT_USAGE_POLICY_REQUEST_ADMISSION_SQLITE_SQL) + .bind(request_id) + .bind(subject_id) + .bind("defensive-event-token") + .bind(100_i64) + .bind(200_i64) + .bind(100_i64) + }; + assert_eq!( + insert("defensive-request-1", "defensive-user-1") + .execute(&pool) + .await + .expect("initial admission should insert") + .rows_affected(), + 1 + ); + assert_eq!( + insert("defensive-request-2", "defensive-user-2") + .execute(&pool) + .await + .expect("duplicate token should be ignored") + .rows_affected(), + 0 + ); + let stored: (String, String, i64) = sqlx::query_as( + "SELECT request_id, subject_id, admitted_at FROM usage_request_admissions WHERE event_token = 'defensive-event-token'", + ) + .fetch_one(&pool) + .await + .expect("original admission should remain"); + assert_eq!( + stored, + ( + "defensive-request-1".to_string(), + "defensive-user-1".to_string(), + 100, + ) + ); + } + + #[tokio::test] + async fn deleting_user_cascades_usage_policy_ledgers_and_terminal_calls_are_noops() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO users (id, username, auth_source, created_at, updated_at) +VALUES ('usage-policy-delete-user', 'usage-policy-delete-user', 'local', 1, 1) + "#, + ) + .execute(&pool) + .await + .expect("usage policy user should insert"); + + let repository = SqliteSettlementRepository::new(pool.clone()); + repository + .reserve_usage_policy_request(ReserveUsagePolicyRequestInput { + request_id: "usage-policy-delete-request".to_string(), + subject_id: "usage-policy-delete-user".to_string(), + event_token: "usage-policy-delete-event".to_string(), + admitted_at_unix_secs: 100, + retain_until_unix_secs: 200, + windows: vec![UsagePolicyRequestWindow { + starts_at_unix_secs: 50, + ends_at_unix_secs: 200, + limit_requests: 10, + }], + }) + .await + .expect("request admission should reserve"); + repository + .reserve_usage_policy_cost(ReserveUsagePolicyCostInput { + request_id: "usage-policy-delete-request".to_string(), + subject_id: "usage-policy-delete-user".to_string(), + reservation_token: "usage-policy-delete-reservation".to_string(), + admitted_at_unix_secs: 100, + reserved_cost_units: 1, + reservation_expires_at_unix_secs: 150, + retain_until_unix_secs: 200, + windows: vec![UsagePolicyCostWindow { + window_id: "usage-policy-delete-window".to_string(), + starts_at_unix_secs: 50, + ends_at_unix_secs: 200, + limit_cost_units: 10, + }], + }) + .await + .expect("cost reservation should reserve"); + + assert!(SqliteUserReadRepository::new(pool.clone()) + .delete_local_auth_user("usage-policy-delete-user") + .await + .expect("user deletion should succeed")); + let ledger_count: i64 = sqlx::query_scalar( + r#" +SELECT + (SELECT COUNT(*) FROM usage_request_admissions) + + (SELECT COUNT(*) FROM usage_cost_reservations) + "#, + ) + .fetch_one(&pool) + .await + .expect("usage policy ledgers should count"); + assert_eq!(ledger_count, 0); + + assert!(repository + .release_usage_policy_request_admission(ReleaseUsagePolicyRequestAdmissionInput { + request_id: "usage-policy-delete-request".to_string(), + subject_id: "usage-policy-delete-user".to_string(), + event_token: "usage-policy-delete-event".to_string(), + released_at_unix_secs: 150, + },) + .await + .expect("post-delete release should be a no-op") + .is_none()); + assert!(repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "usage-policy-delete-request".to_string(), + subject_id: "usage-policy-delete-user".to_string(), + reservation_token: "usage-policy-delete-reservation".to_string(), + actual_cost_units: 1, + terminal_state: UsagePolicyCostReservationState::Finalized, + finalized_at_unix_secs: 150, + }) + .await + .expect("post-delete reconciliation should be a no-op") + .is_none()); + } + #[tokio::test] async fn sqlite_repository_voids_failed_usage_without_wallet_mutation() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1012,6 +1874,58 @@ WHERE request_id = 'request-1' assert_eq!(quota_used, 10.0); } + #[tokio::test] + async fn sqlite_repository_uses_current_plan_wallet_overage_policy() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_quota_covered_settlement_rows(&pool).await; + sqlx::query( + r#" +UPDATE wallets SET balance = 5.0 WHERE id = 'wallet-quota'; +UPDATE billing_plans +SET entitlements_json = '[{"type":"daily_quota","daily_quota_usd":10.0,"reset_timezone":"Asia/Shanghai","allow_wallet_overage":true}]' +WHERE id = 'plan-quota'; +"#, + ) + .execute(&pool) + .await + .expect("plan overage policy should update"); + + let repository = SqliteSettlementRepository::new(pool.clone()); + let settlement = repository + .settle_usage(UsageSettlementInput { + request_id: "request-quota-overrun".to_string(), + user_id: Some("user-quota".to_string()), + api_key_id: Some("key-quota".to_string()), + api_key_is_standalone: false, + provider_id: None, + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 12.0, + actual_total_cost_usd: 12.0, + finalized_at_unix_secs: Some(1_261), + }) + .await + .expect("settlement should run") + .expect("usage should exist"); + + assert_eq!(settlement.billing_status, "settled"); + assert_eq!(settlement.wallet_balance_after, Some(3.0)); + let quota_used: f64 = sqlx::query_scalar( + "SELECT CAST(COALESCE(SUM(amount_usd), 0) AS REAL) FROM entitlement_usage_ledgers WHERE request_id = 'request-quota-overrun'", + ) + .fetch_one(&pool) + .await + .expect("quota ledger should load"); + assert_eq!(quota_used, 10.0); + } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn sqlite_repository_exhausts_strict_quota_across_concurrent_requests() { let database_path = std::env::temp_dir().join(format!( diff --git a/crates/aether-data/adapters/sqlite/src/usage.rs b/crates/aether-data/adapters/sqlite/src/usage.rs index 64b4d3534..c1e5aeba5 100644 --- a/crates/aether-data/adapters/sqlite/src/usage.rs +++ b/crates/aether-data/adapters/sqlite/src/usage.rs @@ -1,5 +1,4 @@ use std::collections::{BTreeMap, HashSet}; -use std::io::Read; use std::time::{SystemTime, UNIX_EPOCH}; use aether_ai_formats::UPSTREAM_IS_STREAM_KEY; @@ -11,9 +10,11 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use crate::error::SqlResultExt; use crate::{sqlite_optional_real, sqlite_real, SqlitePool}; use aether_data_contracts::repository::usage::{ - strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure, - usage_request_metadata_client_family, PendingUsageCleanupSummary, - ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, + read_decompressed_usage_json, sanitize_usage_capture_controls_for_persistence, + sanitize_usage_for_persistence, sanitize_usage_request_metadata, + usage_can_recover_terminal_failure, usage_error_category_for_status_code, + usage_lifecycle_update_allowed, usage_request_metadata_client_family, + PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow, @@ -542,12 +543,15 @@ const UPSERT_FIRST_BYTE_BATCH_UPDATE_SUFFIX_SQL: &str = r#", WHERE "usage".billing_status = 'pending' AND "usage".status IN ('pending', 'streaming') AND "usage".finalized_at IS NULL + AND excluded.updated_at_unix_secs >= COALESCE( + NULLIF("usage".updated_at_unix_secs, 0), + COALESCE("usage".created_at_unix_ms, 0) + ) "#; const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#" SELECT "usage".request_id, - "usage".status, COALESCE(usage_settlement_snapshots.billing_status, "usage".billing_status) AS billing_status FROM "usage" LEFT JOIN usage_settlement_snapshots @@ -1362,11 +1366,7 @@ fn sqlite_usage_body_sql_columns(field: UsageBodyField) -> (&'static str, &'stat } fn inflate_usage_json_value(bytes: &[u8]) -> Result { - let mut decoder = GzDecoder::new(bytes); - let mut json_bytes = Vec::new(); - decoder.read_to_end(&mut json_bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to decompress usage json: {err}")) - })?; + let json_bytes = read_decompressed_usage_json(GzDecoder::new(bytes))?; serde_json::from_slice(&json_bytes).map_err(|err| { DataLayerError::UnexpectedValue(format!("failed to parse decompressed usage json: {err}")) }) @@ -1403,7 +1403,7 @@ impl PreparedFirstByteUsage { )); } - let usage = strip_deprecated_usage_display_fields(usage); + let usage = sanitize_usage_for_persistence(usage); let request_metadata_json = usage .request_metadata .as_ref() @@ -4155,10 +4155,23 @@ fn usage_current_unix_secs() -> u64 { .unwrap_or_default() } -fn first_byte_transition_allowed(existing: &StoredRequestUsageAudit) -> bool { +fn first_byte_transition_allowed( + existing: &StoredRequestUsageAudit, + incoming: &UpsertUsageRecord, +) -> bool { existing.billing_status == "pending" && matches!(existing.status.as_str(), "pending" | "streaming") && existing.finalized_at_unix_secs.is_none() + && usage_lifecycle_update_allowed( + &existing.status, + &existing.billing_status, + existing.updated_at_unix_secs, + existing.finalized_at_unix_secs, + &incoming.status, + &incoming.billing_status, + incoming.updated_at_unix_secs, + incoming.finalized_at_unix_secs, + ) } impl SqliteUsageWriteRepository { @@ -4193,10 +4206,26 @@ impl SqliteUsageWriteRepository { tx: &mut sqlx::Transaction<'_, Sqlite>, usage: UpsertUsageRecord, ) -> Result<(), DataLayerError> { - let mut usage = strip_deprecated_usage_display_fields(usage); usage.validate()?; - let prepared_capture = http_capture::prepare_usage_http_capture(&mut usage)?; + // Auxiliary tables may receive only clear tombstones, never request or response content. + let capture_usage = usage.clone(); + let mut usage = sanitize_usage_for_persistence(usage); + usage.validate()?; let existing = counters::lock_and_load_usage(tx, &usage.request_id).await?; + if existing.as_ref().is_some_and(|existing| { + !usage_lifecycle_update_allowed( + &existing.status, + &existing.billing_status, + existing.updated_at_unix_secs, + existing.finalized_at_unix_secs, + &usage.status, + &usage.billing_status, + usage.updated_at_unix_secs, + usage.finalized_at_unix_secs, + ) + }) { + return Ok(()); + } let recovers_terminal_failure = existing.as_ref().is_some_and(|existing| { usage_can_recover_terminal_failure( &existing.status, @@ -4212,13 +4241,18 @@ impl SqliteUsageWriteRepository { return Ok(()); } + let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage); + let prepared_capture = http_capture::prepare_usage_http_capture(&mut capture_usage)?; let capture_update_allowed = recovers_terminal_failure || http_capture::capture_update_allowed(existing.as_ref(), &usage.status); if capture_update_allowed { - http_capture::apply_previous_metadata_tombstones(&mut usage, existing.as_ref()); + http_capture::apply_previous_metadata_tombstones(&mut capture_usage, existing.as_ref()); + usage.request_metadata = + sanitize_usage_request_metadata(capture_usage.request_metadata.clone()); } let prepared_snapshots = capture_update_allowed - .then(|| snapshots::from_usage(&usage)) + // The control projection preserves safe typed routing and allow-listed billing facts. + .then(|| snapshots::from_usage(&capture_usage)) .transpose()?; bind_upsert(sqlx::query(UPSERT_USAGE_SQL), &usage)? .execute(&mut **tx) @@ -4351,7 +4385,7 @@ impl SqliteUsageWriteRepository { before .get(&row.usage.request_id) .and_then(Option::as_ref) - .is_none_or(first_byte_transition_allowed) + .is_none_or(|existing| first_byte_transition_allowed(existing, &row.usage)) }) .collect::>(); for preserve_existing_format_conversion in [true, false] { @@ -4381,7 +4415,7 @@ impl SqliteUsageWriteRepository { let existing = counters::lock_and_load_usage(&mut tx, &row.usage.request_id).await?; if existing .as_ref() - .is_some_and(|existing| !first_byte_transition_allowed(existing)) + .is_some_and(|existing| !first_byte_transition_allowed(existing, &row.usage)) { continue; } @@ -4628,7 +4662,7 @@ WHERE id = ? &self, cutoff_unix_secs: u64, now_unix_secs: u64, - timeout_minutes: u64, + _timeout_minutes: u64, batch_size: usize, ) -> Result { if batch_size == 0 { @@ -4662,7 +4696,6 @@ WHERE id = ? .map(|row| { Ok(StalePendingUsageRow { request_id: row.try_get("request_id").map_sql_err()?, - status: row.try_get("status").map_sql_err()?, billing_status: row.try_get("billing_status").map_sql_err()?, }) }) @@ -4678,7 +4711,8 @@ WHERE id = ? UPDATE "usage" SET status = 'completed', status_code = 200, - error_message = NULL + error_message = NULL, + error_category = NULL WHERE request_id = ? "#, ) @@ -4706,11 +4740,8 @@ WHERE request_id = ? let candidate_info = latest_failed_candidate_sqlite(&mut tx, &row.request_id).await?; - let (status_code, error_message) = resolve_stale_pending_failure( - candidate_info.as_ref(), - &row.status, - timeout_minutes, - ); + let status_code = resolve_stale_pending_status_code(candidate_info.as_ref()); + let error_category = usage_error_category_for_status_code(status_code); let status_code_i64 = i64::from(status_code); if row.billing_status == "pending" { sqlx::query( @@ -4718,7 +4749,8 @@ WHERE request_id = ? UPDATE "usage" SET status = 'failed', status_code = ?, - error_message = ?, + error_message = NULL, + error_category = ?, billing_status = 'void', finalized_at = ?, total_cost_usd = 0.0, @@ -4727,7 +4759,7 @@ WHERE request_id = ? "#, ) .bind(status_code_i64) - .bind(&error_message) + .bind(error_category) .bind(to_i64(now_unix_secs, "usage finalized_at")?) .bind(&row.request_id) .execute(&mut *tx) @@ -4745,12 +4777,13 @@ WHERE request_id = ? UPDATE "usage" SET status = 'failed', status_code = ?, - error_message = ? + error_message = NULL, + error_category = ? WHERE request_id = ? "#, ) .bind(status_code_i64) - .bind(&error_message) + .bind(error_category) .bind(&row.request_id) .execute(&mut *tx) .await @@ -4762,7 +4795,8 @@ WHERE request_id = ? UPDATE request_candidates SET status = 'failed', finished_at = ?, - error_message = '请求超时(服务器可能已重启)' + error_type = 'internal', + error_message = NULL WHERE request_id = ? AND status IN ('pending', 'streaming') "#, @@ -4849,7 +4883,6 @@ WHERE request_id = ? struct StalePendingUsageRow { request_id: String, - status: String, billing_status: String, } @@ -4960,29 +4993,14 @@ DO UPDATE SET Ok(()) } -fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String { - format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成") -} - struct FailedCandidateCleanupInfo { status_code: Option, - error_message: Option, } -fn resolve_stale_pending_failure( - candidate: Option<&FailedCandidateCleanupInfo>, - status: &str, - timeout_minutes: u64, -) -> (u16, String) { - match candidate { - Some(info) => ( - info.status_code.unwrap_or(502), - info.error_message - .clone() - .unwrap_or_else(|| stale_pending_error_message(status, timeout_minutes)), - ), - None => (504, stale_pending_error_message(status, timeout_minutes)), - } +fn resolve_stale_pending_status_code(candidate: Option<&FailedCandidateCleanupInfo>) -> u16 { + candidate + .and_then(|info| info.status_code) + .unwrap_or(if candidate.is_some() { 502 } else { 504 }) } async fn latest_failed_candidate_sqlite( @@ -4991,7 +5009,7 @@ async fn latest_failed_candidate_sqlite( ) -> Result, DataLayerError> { let row = sqlx::query( r#" -SELECT status_code, error_message +SELECT status_code FROM request_candidates WHERE request_id = ? AND status IN ('failed', 'cancelled') @@ -5014,15 +5032,7 @@ LIMIT 1 .try_get::, _>("status_code") .map_sql_err()? .and_then(|value| u16::try_from(value).ok()); - let error_message = row - .try_get::, _>("error_message") - .map_sql_err()? - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - Ok(Some(FailedCandidateCleanupInfo { - status_code, - error_message, - })) + Ok(Some(FailedCandidateCleanupInfo { status_code })) } fn bind_upsert<'q>( diff --git a/crates/aether-data/adapters/sqlite/src/usage/cleanup.rs b/crates/aether-data/adapters/sqlite/src/usage/cleanup.rs index b3299f545..c81d4bc8d 100644 --- a/crates/aether-data/adapters/sqlite/src/usage/cleanup.rs +++ b/crates/aether-data/adapters/sqlite/src/usage/cleanup.rs @@ -1,11 +1,8 @@ -use std::io::Write; - use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, UsageBodyField, UsageCleanupExecutionMode, - UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow, + UsageCleanupExecutionMode, UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, + UsageCleanupWindow, }; use chrono::{DateTime, Utc}; -use flate2::{write::GzEncoder, Compression}; use serde_json::Value; use sqlx::Row; use tracing::warn; @@ -66,17 +63,6 @@ OR EXISTS ( ) "#; -const INLINE_OR_COMPRESSED_BODY_PREDICATE: &str = r#" -request_body IS NOT NULL -OR response_body IS NOT NULL -OR provider_request_body IS NOT NULL -OR client_response_body IS NOT NULL -OR request_body_compressed IS NOT NULL -OR response_body_compressed IS NOT NULL -OR provider_request_body_compressed IS NOT NULL -OR client_response_body_compressed IS NOT NULL -"#; - const HEADER_PREDICATE: &str = r#" request_headers IS NOT NULL OR response_headers IS NOT NULL @@ -114,6 +100,20 @@ OR request_body_compressed IS NOT NULL OR response_body_compressed IS NOT NULL OR provider_request_body_compressed IS NOT NULL OR client_response_body_compressed IS NOT NULL +OR EXISTS ( + SELECT 1 FROM usage_body_blobs + WHERE usage_body_blobs.request_id = "usage".request_id +) +OR EXISTS ( + SELECT 1 FROM usage_http_audits + WHERE usage_http_audits.request_id = "usage".request_id + AND ( + usage_http_audits.request_body_ref IS NOT NULL + OR usage_http_audits.provider_request_body_ref IS NOT NULL + OR usage_http_audits.response_body_ref IS NOT NULL + OR usage_http_audits.client_response_body_ref IS NOT NULL + ) +) OR ( request_metadata IS NOT NULL AND json_valid(request_metadata) @@ -132,43 +132,6 @@ struct CleanupRow { request_id: String, } -#[derive(Debug)] -struct BodyRow { - id: String, - request_id: String, - request_body: Option, - request_body_compressed: Option>, - provider_request_body: Option, - provider_request_body_compressed: Option>, - response_body: Option, - response_body_compressed: Option>, - client_response_body: Option, - client_response_body_compressed: Option>, -} - -#[derive(Debug, Default)] -struct DetachedRefs { - request_body_ref: Option, - provider_request_body_ref: Option, - response_body_ref: Option, - client_response_body_ref: Option, -} - -impl DetachedRefs { - fn any_present(&self) -> bool { - self.request_body_ref.is_some() - || self.provider_request_body_ref.is_some() - || self.response_body_ref.is_some() - || self.client_response_body_ref.is_some() - } -} - -struct DetachedBlob { - body_ref: String, - body_field: &'static str, - payload_gzip: Vec, -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum BodyCleanupKind { Raw, @@ -184,10 +147,6 @@ impl BodyCleanupKind { Self::All => ALL_BODY_PREDICATE, } } - - fn clears_detached(self) -> bool { - self != Self::Raw - } } pub(crate) async fn cleanup_usage( @@ -259,12 +218,19 @@ pub(crate) async fn cleanup_usage( }; let detail_newer_than = detail_body_newer_than(window, targets); let legacy_body_refs_migrated = if targets.detail_body { - migrate_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await? + purge_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await? } else { 0 }; let body_externalized = if targets.detail_body { - externalize_detail_bodies(pool, window.detail_cutoff, detail_newer_than, batch_size).await? + cleanup_body_fields( + pool, + window.detail_cutoff, + detail_newer_than, + batch_size, + BodyCleanupKind::All, + ) + .await? } else { 0 }; @@ -287,6 +253,8 @@ pub(crate) async fn cleanup_usage( header_cleaned, keys_cleaned, records_deleted, + cost_reservations_deleted: 0, + request_admissions_deleted: 0, }) } @@ -621,14 +589,13 @@ WHERE id = ? .map_sql_err()?; } - if kind.clears_detached() { - sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") - .bind(&row.request_id) - .execute(&mut *tx) - .await - .map_sql_err()?; - sqlx::query( - r#" + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(&row.request_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" UPDATE usage_http_audits SET request_body_ref = NULL, provider_request_body_ref = NULL, @@ -638,13 +605,12 @@ SET request_body_ref = NULL, updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE request_id = ? "#, - ) - .bind(&row.request_id) - .execute(&mut *tx) - .await - .map_sql_err()?; - delete_empty_http_audit(&mut tx, &row.request_id).await?; - } + ) + .bind(&row.request_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + delete_empty_http_audit(&mut tx, &row.request_id).await?; } tx.commit().await.map_sql_err()?; total = total.saturating_add(row_count); @@ -680,14 +646,14 @@ WHERE request_id = ? Ok(()) } -async fn migrate_legacy_body_refs( +async fn purge_legacy_body_refs( pool: &SqlitePool, cutoff: DateTime, newer_than: Option>, batch_size: usize, ) -> Result { if invalid_window(cutoff, newer_than) { - warn!(%cutoff, ?newer_than, "SQLite usage legacy body-ref migration skipped due to invalid window"); + warn!(%cutoff, ?newer_than, "SQLite usage legacy body-ref purge skipped due to invalid window"); return Ok(0); } let mut total = 0usize; @@ -705,7 +671,7 @@ async fn migrate_legacy_body_refs( } let row_count = rows.len(); let mut tx = pool.begin().await.map_sql_err()?; - let mut migrated = 0usize; + let mut purged = 0usize; for row in rows { let metadata: Option = sqlx::query_scalar("SELECT request_metadata FROM \"usage\" WHERE id = ? LIMIT 1") @@ -714,14 +680,9 @@ async fn migrate_legacy_body_refs( .await .map_sql_err()? .flatten(); - let Some((refs, metadata)) = - legacy_body_ref_plan(&row.request_id, metadata.as_deref())? - else { + let Some(metadata) = legacy_body_ref_purge_plan(metadata.as_deref())? else { continue; }; - if refs.any_present() { - upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?; - } let updated = sqlx::query( r#" UPDATE "usage" @@ -736,23 +697,23 @@ WHERE id = ? .await .map_sql_err()? .rows_affected(); + purge_detached_body_capture(&mut tx, &row.request_id).await?; if updated > 0 { - migrated += 1; + purged += 1; } } tx.commit().await.map_sql_err()?; - total = total.saturating_add(migrated); - if row_count < batch_size || migrated == 0 { + total = total.saturating_add(purged); + if row_count < batch_size || purged == 0 { break; } } Ok(total) } -fn legacy_body_ref_plan( - request_id: &str, +fn legacy_body_ref_purge_plan( metadata: Option<&str>, -) -> Result)>, DataLayerError> { +) -> Result>, DataLayerError> { let Some(metadata) = metadata else { return Ok(None); }; @@ -762,30 +723,16 @@ fn legacy_body_ref_plan( let Value::Object(mut object) = value else { return Ok(None); }; - let mut refs = DetachedRefs::default(); let mut removed = false; - for field in [ - UsageBodyField::RequestBody, - UsageBodyField::ProviderRequestBody, - UsageBodyField::ResponseBody, - UsageBodyField::ClientResponseBody, + for key in [ + "request_body_ref", + "provider_request_body_ref", + "response_body_ref", + "client_response_body_ref", ] { - let Some(value) = object.remove(field.as_ref_key()) else { - continue; - }; - removed = true; - let parsed = value - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| { - usage_body_ref(&parsed_request_id, parsed_field) - }); - set_ref(&mut refs, field, parsed); + if object.remove(key).is_some() { + removed = true; + } } if !removed { return Ok(None); @@ -801,283 +748,35 @@ fn legacy_body_ref_plan( })?, ) }; - Ok(Some((refs, metadata))) + Ok(Some(metadata)) } -async fn externalize_detail_bodies( - pool: &SqlitePool, - cutoff: DateTime, - newer_than: Option>, - batch_size: usize, -) -> Result { - if invalid_window(cutoff, newer_than) { - warn!(%cutoff, ?newer_than, "SQLite usage body externalization skipped due to invalid window"); - return Ok(0); - } - let batch_size = batch_size.clamp(1, 25); - let mut total = 0usize; - loop { - let rows = fetch_body_rows(pool, cutoff, newer_than, batch_size).await?; - if rows.is_empty() { - break; - } - let row_count = rows.len(); - let mut externalized = 0usize; - for row in rows { - let (blobs, refs) = build_detached_bodies(&row)?; - let mut tx = pool.begin().await.map_sql_err()?; - for blob in blobs { - sqlx::query( - r#" -INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip) -VALUES (?, ?, ?, ?) -ON CONFLICT(body_ref) DO UPDATE SET - request_id = excluded.request_id, - body_field = excluded.body_field, - payload_gzip = excluded.payload_gzip, - updated_at = CAST(strftime('%s', 'now') AS INTEGER) -"#, - ) - .bind(blob.body_ref) - .bind(&row.request_id) - .bind(blob.body_field) - .bind(blob.payload_gzip) - .execute(&mut *tx) - .await - .map_sql_err()?; - } - if refs.any_present() { - upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?; - } - let updated = sqlx::query( - r#" -UPDATE "usage" -SET request_body = NULL, - response_body = NULL, - provider_request_body = NULL, - client_response_body = NULL, - request_body_compressed = NULL, - response_body_compressed = NULL, - provider_request_body_compressed = NULL, - client_response_body_compressed = NULL -WHERE id = ? -"#, - ) - .bind(row.id) - .execute(&mut *tx) - .await - .map_sql_err()? - .rows_affected(); - tx.commit().await.map_sql_err()?; - if updated > 0 { - externalized += 1; - } - } - total = total.saturating_add(externalized); - if row_count < batch_size || externalized == 0 { - break; - } - } - Ok(total) -} - -async fn fetch_body_rows( - pool: &SqlitePool, - cutoff: DateTime, - newer_than: Option>, - batch_size: usize, -) -> Result, DataLayerError> { - let newer_than = newer_than.map(|value| value.timestamp()); - let sql = format!( - r#" -SELECT id, - request_id, - request_body, - request_body_compressed, - provider_request_body, - provider_request_body_compressed, - response_body, - response_body_compressed, - client_response_body, - client_response_body_compressed -FROM "usage" -WHERE created_at_unix_ms < ? - AND (? IS NULL OR created_at_unix_ms >= ?) - AND ({INLINE_OR_COMPRESSED_BODY_PREDICATE}) -ORDER BY created_at_unix_ms ASC, id ASC -LIMIT ? -"# - ); - sqlx::query(&sql) - .bind(cutoff.timestamp()) - .bind(newer_than) - .bind(newer_than) - .bind(i64::try_from(batch_size).unwrap_or(i64::MAX)) - .fetch_all(pool) - .await - .map_sql_err()? - .into_iter() - .map(|row| { - Ok(BodyRow { - id: row.try_get("id").map_sql_err()?, - request_id: row.try_get("request_id").map_sql_err()?, - request_body: parse_optional_json(row.try_get("request_body").map_sql_err()?)?, - request_body_compressed: row.try_get("request_body_compressed").map_sql_err()?, - provider_request_body: parse_optional_json( - row.try_get("provider_request_body").map_sql_err()?, - )?, - provider_request_body_compressed: row - .try_get("provider_request_body_compressed") - .map_sql_err()?, - response_body: parse_optional_json(row.try_get("response_body").map_sql_err()?)?, - response_body_compressed: row.try_get("response_body_compressed").map_sql_err()?, - client_response_body: parse_optional_json( - row.try_get("client_response_body").map_sql_err()?, - )?, - client_response_body_compressed: row - .try_get("client_response_body_compressed") - .map_sql_err()?, - }) - }) - .collect() -} - -fn parse_optional_json(raw: Option) -> Result, DataLayerError> { - raw.map(|raw| { - serde_json::from_str(&raw).map_err(|err| { - DataLayerError::UnexpectedValue(format!("invalid inline usage body JSON: {err}")) - }) - }) - .transpose() -} - -fn build_detached_bodies( - row: &BodyRow, -) -> Result<(Vec, DetachedRefs), DataLayerError> { - let mut blobs = Vec::new(); - let mut refs = DetachedRefs::default(); - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::RequestBody, - row.request_body.as_ref(), - row.request_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ProviderRequestBody, - row.provider_request_body.as_ref(), - row.provider_request_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ResponseBody, - row.response_body.as_ref(), - row.response_body_compressed.as_deref(), - )?; - add_detached_body( - &mut blobs, - &mut refs, - &row.request_id, - UsageBodyField::ClientResponseBody, - row.client_response_body.as_ref(), - row.client_response_body_compressed.as_deref(), - )?; - Ok((blobs, refs)) -} - -fn add_detached_body( - blobs: &mut Vec, - refs: &mut DetachedRefs, - request_id: &str, - field: UsageBodyField, - raw: Option<&Value>, - compressed: Option<&[u8]>, -) -> Result<(), DataLayerError> { - let payload_gzip = match raw { - Some(value) => Some(compress_json(value)?), - None => compressed.map(ToOwned::to_owned), - }; - let Some(payload_gzip) = payload_gzip else { - return Ok(()); - }; - let body_ref = usage_body_ref(request_id, field); - blobs.push(DetachedBlob { - body_ref: body_ref.clone(), - body_field: field.as_storage_field(), - payload_gzip, - }); - set_ref(refs, field, Some(body_ref)); - Ok(()) -} - -fn compress_json(value: &Value) -> Result, DataLayerError> { - let bytes = serde_json::to_vec(value).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}")) - })?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6)); - encoder.write_all(&bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}")) - })?; - encoder.finish().map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}")) - }) -} - -fn set_ref(refs: &mut DetachedRefs, field: UsageBodyField, value: Option) { - match field { - UsageBodyField::RequestBody => refs.request_body_ref = value, - UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = value, - UsageBodyField::ResponseBody => refs.response_body_ref = value, - UsageBodyField::ClientResponseBody => refs.client_response_body_ref = value, - } -} - -async fn upsert_http_audit_refs( +async fn purge_detached_body_capture( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, request_id: &str, - refs: &DetachedRefs, ) -> Result<(), DataLayerError> { + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; sqlx::query( r#" -INSERT INTO usage_http_audits ( - request_id, - request_body_ref, - provider_request_body_ref, - response_body_ref, - client_response_body_ref, - body_capture_mode -) -VALUES (?, ?, ?, ?, ?, 'ref_backed') -ON CONFLICT(request_id) DO UPDATE SET - request_body_ref = COALESCE(excluded.request_body_ref, usage_http_audits.request_body_ref), - provider_request_body_ref = COALESCE( - excluded.provider_request_body_ref, - usage_http_audits.provider_request_body_ref - ), - response_body_ref = COALESCE(excluded.response_body_ref, usage_http_audits.response_body_ref), - client_response_body_ref = COALESCE( - excluded.client_response_body_ref, - usage_http_audits.client_response_body_ref - ), - body_capture_mode = 'ref_backed', +UPDATE usage_http_audits +SET request_body_ref = NULL, + provider_request_body_ref = NULL, + response_body_ref = NULL, + client_response_body_ref = NULL, + body_capture_mode = 'none', updated_at = CAST(strftime('%s', 'now') AS INTEGER) +WHERE request_id = ? "#, ) .bind(request_id) - .bind(refs.request_body_ref.as_deref()) - .bind(refs.provider_request_body_ref.as_deref()) - .bind(refs.response_body_ref.as_deref()) - .bind(refs.client_response_body_ref.as_deref()) .execute(&mut **tx) .await .map_sql_err()?; - Ok(()) + delete_empty_http_audit(tx, request_id).await } async fn cleanup_expired_api_keys( diff --git a/crates/aether-data/adapters/sqlite/src/usage/counters.rs b/crates/aether-data/adapters/sqlite/src/usage/counters.rs index 4e265b5f8..b76b936d4 100644 --- a/crates/aether-data/adapters/sqlite/src/usage/counters.rs +++ b/crates/aether-data/adapters/sqlite/src/usage/counters.rs @@ -25,6 +25,7 @@ SELECT id, kind, target_id, + target_tunnel_generation, request_count_delta, total_requests_delta, success_count_delta, @@ -49,6 +50,7 @@ struct DeltaRow { id: String, kind: String, target_id: String, + target_tunnel_generation: Option, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -71,7 +73,9 @@ struct Aggregates { provider_api_keys: BTreeMap, models: BTreeMap, provider_monthly: BTreeMap, - proxy_nodes: BTreeMap, + // Generation is part of the key: rows from an old incarnation must never + // be coalesced with rows for a newly registered node using the same id. + proxy_nodes: BTreeMap<(String, String), ProxyNodeCounterDelta>, management_tokens: BTreeMap, api_key_last_used: BTreeMap, } @@ -142,16 +146,28 @@ impl Aggregates { .or_default() += row.total_cost_usd_delta; } KIND_PROXY_NODE => { - let entry = aggregates - .proxy_nodes - .entry(row.target_id.clone()) - .or_insert(ProxyNodeCounterDelta { + let Some(tunnel_generation) = row + .target_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + // Legacy rows without an identity fence are intentionally + // discarded when marked processed below. + continue; + }; + let aggregate_key = (row.target_id.clone(), tunnel_generation.clone()); + let entry = aggregates.proxy_nodes.entry(aggregate_key).or_insert( + ProxyNodeCounterDelta { node_id: row.target_id.clone(), + expected_tunnel_generation: Some(tunnel_generation), total_requests_delta: 0, failed_requests_delta: 0, dns_failures_delta: 0, stream_errors_delta: 0, - }); + }, + ); entry.total_requests_delta += row.total_requests_delta; entry.failed_requests_delta += row.error_count_delta; entry.dns_failures_delta += row.dns_failures_delta; @@ -247,8 +263,8 @@ pub(super) async fn flush( for (target_id, delta) in &aggregates.provider_monthly { apply_provider_monthly(&mut tx, target_id, *delta).await?; } - for (target_id, delta) in &aggregates.proxy_nodes { - apply_proxy_node(&mut tx, target_id, delta).await?; + for ((target_id, tunnel_generation), delta) in &aggregates.proxy_nodes { + apply_proxy_node(&mut tx, target_id, tunnel_generation, delta).await?; } for (target_id, delta) in &aggregates.management_tokens { apply_management_token(&mut tx, target_id, delta).await?; @@ -289,9 +305,43 @@ pub(super) async fn enqueue_proxy_node( if delta.is_noop() { return Ok(false); } + let Some(expected_tunnel_generation) = delta + .expected_tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .filter(|value| value.len() <= 64) + .map(ToOwned::to_owned) + else { + // Never infer an incarnation from a bare node id. A missing fence can + // otherwise make a stale request update a node recreated under that id. + return Ok(false); + }; let node_id = delta.node_id.trim().to_string(); let request_id = format!("proxy_node:{node_id}:{}", uuid::Uuid::new_v4()); let mut tx = pool.begin().await.map_sql_err()?; + // Read the generation as a predicate without taking an exclusive parent + // lock. Flush claims outbox rows before touching proxy_nodes; keeping this + // read lock-free avoids the inverse parent->outbox lock order. The value is + // persisted in the outbox row and the flush UPDATE repeats the generation + // predicate, so a delete/re-register race can only retire the delta. + let tunnel_generation: Option = sqlx::query_scalar( + "SELECT tunnel_generation FROM proxy_nodes WHERE id = ? AND tunnel_generation = ? LIMIT 1", + ) + .bind(&node_id) + .bind(&expected_tunnel_generation) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(_tunnel_generation) = tunnel_generation + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; insert_delta( &mut tx, DeltaInsert { @@ -302,6 +352,7 @@ pub(super) async fn enqueue_proxy_node( error_count_delta: delta.failed_requests_delta, dns_failures_delta: delta.dns_failures_delta, stream_errors_delta: delta.stream_errors_delta, + target_tunnel_generation: Some(&expected_tunnel_generation), ..DeltaInsert::default() }, ) @@ -734,6 +785,7 @@ struct DeltaInsert<'a> { request_id: &'a str, kind: &'a str, target_id: &'a str, + target_tunnel_generation: Option<&'a str>, request_count_delta: i64, total_requests_delta: i64, success_count_delta: i64, @@ -763,11 +815,12 @@ async fn insert_delta( r#" INSERT INTO usage_counter_deltas ( id, request_id, kind, target_id, request_count_delta, total_requests_delta, + target_tunnel_generation, success_count_delta, error_count_delta, dns_failures_delta, stream_errors_delta, total_tokens_delta, total_cost_usd_delta, total_response_time_ms_delta, last_used_at_unix_secs, last_used_ip, candidate_last_used_at_unix_secs, removed_last_used_at_unix_secs, usage_created_at_unix_secs, created_at -) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(uuid::Uuid::new_v4().to_string()) @@ -776,6 +829,7 @@ INSERT INTO usage_counter_deltas ( .bind(target_id) .bind(input.request_count_delta) .bind(input.total_requests_delta) + .bind(input.target_tunnel_generation) .bind(input.success_count_delta) .bind(input.error_count_delta) .bind(input.dns_failures_delta) @@ -817,6 +871,7 @@ fn map_row(row: &sqlx::sqlite::SqliteRow) -> Result { id: row.try_get("id").map_sql_err()?, kind: row.try_get("kind").map_sql_err()?, target_id: row.try_get("target_id").map_sql_err()?, + target_tunnel_generation: row.try_get("target_tunnel_generation").map_sql_err()?, request_count_delta: row.try_get("request_count_delta").map_sql_err()?, total_requests_delta: row.try_get("total_requests_delta").map_sql_err()?, success_count_delta: row.try_get("success_count_delta").map_sql_err()?, @@ -1000,9 +1055,10 @@ async fn apply_provider_monthly( async fn apply_proxy_node( tx: &mut sqlx::Transaction<'_, Sqlite>, target_id: &str, + tunnel_generation: &str, delta: &ProxyNodeCounterDelta, ) -> Result<(), DataLayerError> { - if target_id.trim().is_empty() || delta.is_noop() { + if target_id.trim().is_empty() || tunnel_generation.trim().is_empty() || delta.is_noop() { return Ok(()); } sqlx::query( @@ -1013,7 +1069,7 @@ SET total_requests = total_requests + MAX(?, 0), dns_failures = dns_failures + MAX(?, 0), stream_errors = stream_errors + MAX(?, 0), updated_at = ? -WHERE id = ? +WHERE id = ? AND tunnel_generation = ? "#, ) .bind(delta.total_requests_delta) @@ -1022,6 +1078,7 @@ WHERE id = ? .bind(delta.stream_errors_delta) .bind(current_unix_secs()) .bind(target_id) + .bind(tunnel_generation) .execute(&mut **tx) .await .map_sql_err()?; @@ -1157,13 +1214,18 @@ fn optional_nonnegative_u64(value: Option) -> Option { #[cfg(test)] mod tests { + use std::{sync::Arc, time::Duration}; + use super::{ cleanup_processed, enqueue_api_key_last_used, enqueue_management_token, enqueue_proxy_node, flush, read_health, read_pending_health, }; + use crate::proxy_nodes::SqliteProxyNodeReadRepository; + use aether_data_contracts::repository::proxy_nodes::ProxyNodeWriteRepository; use aether_data_contracts::repository::usage::{ ApiKeyLastUsedDelta, ManagementTokenCounterDelta, ProxyNodeCounterDelta, }; + use tokio::sync::Barrier; #[tokio::test] async fn auxiliary_counters_flush_report_health_and_cleanup() { @@ -1186,8 +1248,8 @@ INSERT INTO management_tokens ( ) VALUES ( 'counter-token', 'counter-user', 'counter token', 'counter-token-hash', 1, 1 ); -INSERT INTO proxy_nodes (id, name, ip, port, created_at, updated_at) -VALUES ('counter-node', 'counter node', '127.0.0.1', 8080, 1, 1); +INSERT INTO proxy_nodes (id, tunnel_generation, name, ip, port, created_at, updated_at) +VALUES ('counter-node', 'test-generation-counter-node', 'counter node', '127.0.0.1', 8080, 1, 1); "#, ) .execute(&pool) @@ -1198,6 +1260,7 @@ VALUES ('counter-node', 'counter node', '127.0.0.1', 8080, 1, 1); &pool, ProxyNodeCounterDelta { node_id: "counter-node".to_string(), + expected_tunnel_generation: Some("test-generation-counter-node".to_string()), total_requests_delta: 3, failed_requests_delta: 1, dns_failures_delta: 2, @@ -1280,4 +1343,261 @@ VALUES ('counter-node', 'counter node', '127.0.0.1', 8080, 1, 1); 0 ); } + + #[tokio::test] + async fn proxy_counter_enqueue_requires_the_expected_generation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + crate::run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO proxy_nodes (id, tunnel_generation, name, ip, port, created_at, updated_at) +VALUES ('generation-node', 'generation-a', 'generation node', '127.0.0.1', 8080, 1, 1) +"#, + ) + .execute(&pool) + .await + .expect("proxy node should seed"); + + assert!(!enqueue_proxy_node( + &pool, + ProxyNodeCounterDelta { + node_id: "generation-node".to_string(), + expected_tunnel_generation: None, + total_requests_delta: 1, + failed_requests_delta: 0, + dns_failures_delta: 0, + stream_errors_delta: 0, + }, + ) + .await + .expect("missing generation should be handled")); + assert!(!enqueue_proxy_node( + &pool, + ProxyNodeCounterDelta { + node_id: "generation-node".to_string(), + expected_tunnel_generation: Some("generation-b".to_string()), + total_requests_delta: 1, + failed_requests_delta: 0, + dns_failures_delta: 0, + stream_errors_delta: 0, + }, + ) + .await + .expect("mismatched generation should be handled")); + let pending: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_counter_deltas WHERE processed_at IS NULL", + ) + .fetch_one(&pool) + .await + .expect("pending rows should load"); + assert_eq!(pending, 0); + + assert!(enqueue_proxy_node( + &pool, + ProxyNodeCounterDelta { + node_id: "generation-node".to_string(), + expected_tunnel_generation: Some("generation-a".to_string()), + total_requests_delta: 1, + failed_requests_delta: 0, + dns_failures_delta: 0, + stream_errors_delta: 0, + }, + ) + .await + .expect("matching generation should enqueue")); + } + + #[tokio::test] + async fn proxy_counter_flush_does_not_cross_generation_reuse() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + crate::run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO proxy_nodes (id, tunnel_generation, name, ip, port, created_at, updated_at) +VALUES ('reused-node', 'generation-a', 'reused node', '127.0.0.1', 8080, 1, 1) +"#, + ) + .execute(&pool) + .await + .expect("first proxy incarnation should seed"); + assert!(enqueue_proxy_node( + &pool, + ProxyNodeCounterDelta { + node_id: "reused-node".to_string(), + expected_tunnel_generation: Some("generation-a".to_string()), + total_requests_delta: 7, + failed_requests_delta: 3, + dns_failures_delta: 2, + stream_errors_delta: 1, + }, + ) + .await + .expect("first incarnation delta should enqueue")); + + // Simulate deletion followed by id reuse while the old outbox row is + // still pending. The generation predicate must retire the row without + // applying it to the replacement incarnation. + sqlx::query("DELETE FROM proxy_nodes WHERE id = 'reused-node'") + .execute(&pool) + .await + .expect("first incarnation should delete"); + sqlx::query( + r#" +INSERT INTO proxy_nodes (id, tunnel_generation, name, ip, port, created_at, updated_at) +VALUES ('reused-node', 'generation-b', 'reused node', '127.0.0.1', 8080, 2, 2) +"#, + ) + .execute(&pool) + .await + .expect("replacement incarnation should seed"); + + let summary = flush(&pool, 100) + .await + .expect("counter flush should succeed"); + assert_eq!(summary.rows_claimed, 1); + let counters: (i64, i64, i64, i64) = sqlx::query_as( + "SELECT total_requests, failed_requests, dns_failures, stream_errors FROM proxy_nodes WHERE id = 'reused-node'", + ) + .fetch_one(&pool) + .await + .expect("replacement counters should load"); + assert_eq!(counters, (0, 0, 0, 0)); + let pending: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_counter_deltas WHERE processed_at IS NULL", + ) + .fetch_one(&pool) + .await + .expect("pending rows should load"); + assert_eq!(pending, 0); + } + + #[tokio::test] + async fn proxy_counter_flush_and_delete_do_not_deadlock_or_cross_generation() { + let database_path = std::env::temp_dir().join(format!( + "aether-proxy-counter-delete-flush-{}.sqlite", + uuid::Uuid::new_v4().simple() + )); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .busy_timeout(Duration::from_secs(10)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(8) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + crate::run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let node_id = "delete-flush-node"; + let generation_a = "delete-flush-generation-a"; + sqlx::query( + r#" +INSERT INTO proxy_nodes ( + id, tunnel_generation, name, ip, port, is_manual, proxy_url, created_at, updated_at +) +VALUES (?, ?, 'delete/flush node', '127.0.0.91', 8091, 1, 'http://127.0.0.91:8091', 1, 1) +"#, + ) + .bind(node_id) + .bind(generation_a) + .execute(&pool) + .await + .expect("proxy node should seed"); + assert!(enqueue_proxy_node( + &pool, + ProxyNodeCounterDelta { + node_id: node_id.to_string(), + expected_tunnel_generation: Some(generation_a.to_string()), + total_requests_delta: 11, + failed_requests_delta: 3, + dns_failures_delta: 2, + stream_errors_delta: 1, + }, + ) + .await + .expect("counter delta should enqueue")); + + let repository = SqliteProxyNodeReadRepository::new(pool.clone()); + let barrier = Arc::new(Barrier::new(2)); + let flush_barrier = Arc::clone(&barrier); + let flush_pool = pool.clone(); + let flush_task = tokio::spawn(async move { + flush_barrier.wait().await; + tokio::time::timeout(Duration::from_secs(15), flush(&flush_pool, 100)) + .await + .expect("counter flush should not deadlock") + .expect("counter flush should succeed") + }); + let delete_barrier = Arc::clone(&barrier); + let delete_repository = repository.clone(); + let delete_task = tokio::spawn(async move { + delete_barrier.wait().await; + tokio::time::timeout( + Duration::from_secs(15), + delete_repository.delete_node(node_id), + ) + .await + .expect("proxy delete should not deadlock") + .expect("proxy delete should succeed") + }); + let (flush_summary, deleted) = tokio::join!(flush_task, delete_task); + let flush_summary = flush_summary.expect("flush task should join"); + let deleted = deleted.expect("delete task should join"); + assert!(deleted.is_some(), "the original node should be deleted"); + assert!(flush_summary.rows_claimed <= 1); + + // Reuse the id immediately. Any row that lost the race with the + // post-commit cleanup is still safe because flush checks generation. + let generation_b = "delete-flush-generation-b"; + sqlx::query( + r#" +INSERT INTO proxy_nodes ( + id, tunnel_generation, name, ip, port, is_manual, proxy_url, created_at, updated_at +) +VALUES (?, ?, 'replacement node', '127.0.0.92', 8092, 1, 'http://127.0.0.92:8092', 2, 2) +"#, + ) + .bind(node_id) + .bind(generation_b) + .execute(&pool) + .await + .expect("replacement node should seed"); + flush(&pool, 100) + .await + .expect("stale counter flush should succeed"); + + let counters: (i64, i64, i64, i64) = sqlx::query_as( + "SELECT total_requests, failed_requests, dns_failures, stream_errors FROM proxy_nodes WHERE id = ?", + ) + .bind(node_id) + .fetch_one(&pool) + .await + .expect("replacement counters should load"); + assert_eq!(counters, (0, 0, 0, 0)); + let pending: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_counter_deltas WHERE processed_at IS NULL", + ) + .fetch_one(&pool) + .await + .expect("pending rows should load"); + assert_eq!(pending, 0); + + pool.close().await; + let _ = std::fs::remove_file(database_path); + } } diff --git a/crates/aether-data/adapters/sqlite/src/usage/http_capture.rs b/crates/aether-data/adapters/sqlite/src/usage/http_capture.rs index 9e66588e7..44b6a4580 100644 --- a/crates/aether-data/adapters/sqlite/src/usage/http_capture.rs +++ b/crates/aether-data/adapters/sqlite/src/usage/http_capture.rs @@ -1,10 +1,7 @@ -use std::io::Write; - use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, - UsageBodyCaptureState, UsageBodyField, + canonical_usage_body_ref_for, parse_usage_body_ref, usage_body_ref, StoredRequestUsageAudit, + UpsertUsageRecord, UsageBodyCaptureState, UsageBodyField, }; -use flate2::{write::GzEncoder, Compression}; use serde_json::{Map, Value}; use sqlx::{sqlite::SqliteRow, Row}; @@ -28,9 +25,7 @@ pub(crate) struct PreparedUsageHttpCapture { #[derive(Debug)] struct PreparedBody { - field: UsageBodyField, payload_gzip: Option>, - clear_existing: bool, } #[derive(Debug, Default)] @@ -58,15 +53,6 @@ struct HttpAuditStates { client_response_body_state: Option, } -impl HttpAuditStates { - fn any_present(&self) -> bool { - self.request_body_state.is_some() - || self.provider_request_body_state.is_some() - || self.response_body_state.is_some() - || self.client_response_body_state.is_some() - } -} - pub(crate) fn capture_update_allowed( previous: Option<&StoredRequestUsageAudit>, incoming_status: &str, @@ -141,26 +127,10 @@ pub(crate) fn prepare_usage_http_capture( .then_some(usage.client_response_body.as_ref()) .flatten(); - let request_body = prepare_body( - UsageBodyField::RequestBody, - request_body_value, - clear_request, - )?; - let provider_request_body = prepare_body( - UsageBodyField::ProviderRequestBody, - provider_request_body_value, - clear_provider_request, - )?; - let response_body = prepare_body( - UsageBodyField::ResponseBody, - response_body_value, - clear_response, - )?; - let client_response_body = prepare_body( - UsageBodyField::ClientResponseBody, - client_response_body_value, - clear_client_response, - )?; + let request_body = prepare_body(request_body_value)?; + let provider_request_body = prepare_body(provider_request_body_value)?; + let response_body = prepare_body(response_body_value)?; + let client_response_body = prepare_body(client_response_body_value)?; let refs = HttpAuditRefs { request_body_ref: resolved_write_ref( @@ -276,30 +246,13 @@ pub(crate) fn prepare_usage_http_capture( }) } -fn prepare_body( - field: UsageBodyField, - value: Option<&Value>, - clear_existing: bool, -) -> Result { - let payload_gzip = value.map(compress_json).transpose()?; - Ok(PreparedBody { - field, - payload_gzip, - clear_existing, - }) -} - -fn compress_json(value: &Value) -> Result, DataLayerError> { - let bytes = serde_json::to_vec(value).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}")) - })?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6)); - encoder.write_all(&bytes).map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}")) - })?; - encoder.finish().map_err(|err| { - DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}")) - }) +fn prepare_body(value: Option<&Value>) -> Result { + if value.is_some() { + return Err(DataLayerError::InvalidInput( + "usage body persistence is disabled".to_string(), + )); + } + Ok(PreparedBody { payload_gzip: None }) } fn resolved_write_ref( @@ -309,9 +262,7 @@ fn resolved_write_ref( has_blob: bool, ) -> Option { explicit_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) .or_else(|| has_blob.then(|| usage_body_ref(request_id, field))) } @@ -379,24 +330,63 @@ pub(crate) async fn sync_usage_http_capture( request_id: &str, prepared: &PreparedUsageHttpCapture, ) -> Result<(), DataLayerError> { - for body in [ + let bodies = [ &prepared.request_body, &prepared.provider_request_body, &prepared.response_body, &prepared.client_response_body, - ] { - sync_body(tx, request_id, body).await?; + ]; + let contains_capture = prepared.request_headers.is_some() + || prepared.provider_request_headers.is_some() + || prepared.response_headers.is_some() + || prepared.client_response_headers.is_some() + || prepared.refs.any_present() + || bodies.iter().any(|body| body.payload_gzip.is_some()) + || prepared.capture_mode != "none"; + if contains_capture { + return Err(DataLayerError::InvalidInput( + "usage HTTP capture persistence is disabled".to_string(), + )); } + sqlx::query("DELETE FROM usage_http_audits WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?") + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +UPDATE "usage" +SET request_headers = NULL, + request_body = NULL, + provider_request_headers = NULL, + provider_request_body = NULL, + response_headers = NULL, + response_body = NULL, + client_response_headers = NULL, + client_response_body = NULL, + request_body_compressed = NULL, + provider_request_body_compressed = NULL, + response_body_compressed = NULL, + client_response_body_compressed = NULL +WHERE request_id = ? +"#, + ) + .bind(request_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + let headers_present = prepared.request_headers.is_some() || prepared.provider_request_headers.is_some() || prepared.response_headers.is_some() || prepared.client_response_headers.is_some(); - if !headers_present - && !prepared.refs.any_present() - && !prepared.states.any_present() - && prepared.capture_mode == "none" - { + if !headers_present && !prepared.refs.any_present() { return Ok(()); } @@ -526,67 +516,6 @@ ON CONFLICT(request_id) DO UPDATE SET Ok(()) } -async fn sync_body( - tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, - request_id: &str, - body: &PreparedBody, -) -> Result<(), DataLayerError> { - let body_ref = usage_body_ref(request_id, body.field); - if body.clear_existing || body.payload_gzip.is_some() { - sqlx::query(clear_legacy_body_sql(body.field)) - .bind(request_id) - .execute(&mut **tx) - .await - .map_sql_err()?; - } - if body.clear_existing { - sqlx::query("DELETE FROM usage_body_blobs WHERE body_ref = ?") - .bind(body_ref) - .execute(&mut **tx) - .await - .map_sql_err()?; - return Ok(()); - } - if let Some(payload_gzip) = body.payload_gzip.as_deref() { - sqlx::query( - r#" -INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip) -VALUES (?, ?, ?, ?) -ON CONFLICT(body_ref) DO UPDATE SET - request_id = excluded.request_id, - body_field = excluded.body_field, - payload_gzip = excluded.payload_gzip, - updated_at = CAST(strftime('%s', 'now') AS INTEGER) -"#, - ) - .bind(body_ref) - .bind(request_id) - .bind(body.field.as_storage_field()) - .bind(payload_gzip) - .execute(&mut **tx) - .await - .map_sql_err()?; - } - Ok(()) -} - -fn clear_legacy_body_sql(field: UsageBodyField) -> &'static str { - match field { - UsageBodyField::RequestBody => { - "UPDATE \"usage\" SET request_body = NULL, request_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ProviderRequestBody => { - "UPDATE \"usage\" SET provider_request_body = NULL, provider_request_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ResponseBody => { - "UPDATE \"usage\" SET response_body = NULL, response_body_compressed = NULL WHERE request_id = ?" - } - UsageBodyField::ClientResponseBody => { - "UPDATE \"usage\" SET client_response_body = NULL, client_response_body_compressed = NULL WHERE request_id = ?" - } - } -} - pub(crate) fn hydrate_usage_row( row: &SqliteRow, usage: &mut StoredRequestUsageAudit, @@ -702,8 +631,7 @@ fn resolved_read_ref( has_compressed: bool, ) -> Option { audit_ref - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) + .and_then(|body_ref| canonical_usage_body_ref_for(&body_ref, request_id, field)) .or_else(|| has_compressed.then(|| usage_body_ref(request_id, field))) .or_else(|| metadata_body_ref(metadata, request_id, field)) } @@ -716,13 +644,7 @@ fn metadata_body_ref( metadata .and_then(|metadata| metadata.get(field.as_ref_key())) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field)) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) } fn optional_state( @@ -764,7 +686,11 @@ pub(crate) async fn hydrate_usage_body_refs( let Some(body_ref) = usage.body_ref(field) else { continue; }; - let value = resolve_body_ref(pool, body_ref).await?; + let Some(body_ref) = canonical_usage_body_ref_for(body_ref, &usage.request_id, field) + else { + continue; + }; + let value = resolve_body_ref(pool, &body_ref).await?; match field { UsageBodyField::RequestBody => usage.request_body = value, UsageBodyField::ProviderRequestBody => usage.provider_request_body = value, @@ -779,19 +705,22 @@ pub(crate) async fn resolve_body_ref( pool: &SqlitePool, body_ref: &str, ) -> Result, DataLayerError> { + let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { + return Ok(None); + }; + let canonical_ref = usage_body_ref(&request_id, field); if let Some(payload_gzip) = sqlx::query_scalar::<_, Vec>( - "SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? LIMIT 1", + "SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? AND request_id = ? AND body_field = ? LIMIT 1", ) - .bind(body_ref) + .bind(&canonical_ref) + .bind(&request_id) + .bind(field.as_storage_field()) .fetch_optional(pool) .await .map_sql_err()? { return super::inflate_usage_json_value(&payload_gzip).map(Some); } - let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { - return Ok(None); - }; let (inline_column, compressed_column) = super::sqlite_usage_body_sql_columns(field); let row = sqlx::query(&format!( "SELECT {inline_column} AS inline_body, {compressed_column} AS compressed_body FROM \"usage\" WHERE request_id = ? LIMIT 1" diff --git a/crates/aether-data/adapters/sqlite/src/usage/tests.rs b/crates/aether-data/adapters/sqlite/src/usage/tests.rs index 425da833d..377d253bd 100644 --- a/crates/aether-data/adapters/sqlite/src/usage/tests.rs +++ b/crates/aether-data/adapters/sqlite/src/usage/tests.rs @@ -10,6 +10,16 @@ use aether_data_contracts::repository::usage::{ UsageWriteRepository, }; use chrono::{DateTime, Utc}; +use flate2::{write::GzEncoder, Compression}; +use std::io::Write; + +fn gzip_json_for_test(value: &serde_json::Value) -> Vec { + let mut encoder = GzEncoder::new(Vec::new(), Compression::fast()); + encoder + .write_all(&serde_json::to_vec(value).expect("test JSON should serialize")) + .expect("test JSON should compress"); + encoder.finish().expect("test gzip should finish") +} #[test] fn sqlite_usage_upsert_guards_candidate_identity_metadata_and_routing_from_late_lifecycle() { @@ -54,6 +64,13 @@ fn sqlite_usage_upsert_guards_candidate_identity_metadata_and_routing_from_late_ .contains("OR (\"usage\".status = 'streaming' AND excluded.status = 'pending')")); } +#[test] +fn sqlite_first_byte_upsert_rejects_older_revision_in_sql() { + assert!(super::UPSERT_FIRST_BYTE_BATCH_UPDATE_SUFFIX_SQL.contains( + "excluded.updated_at_unix_secs >= COALESCE(\n NULLIF(\"usage\".updated_at_unix_secs, 0)," + )); +} + #[tokio::test] async fn sqlite_provider_performance_can_skip_timeline() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -364,7 +381,7 @@ WHERE request_id = 'rebuild-completed'; } #[tokio::test] -async fn sqlite_usage_http_capture_round_trips_and_preserves_sparse_updates() { +async fn sqlite_usage_http_capture_is_not_persisted() { let pool = sqlx::sqlite::SqlitePoolOptions::new() .max_connections(1) .connect("sqlite::memory:") @@ -398,31 +415,13 @@ async fn sqlite_usage_http_capture_round_trips_and_preserves_sparse_updates() { .upsert(rich) .await .expect("canonical capture should upsert"); - assert_eq!( - stored.request_headers, - Some(serde_json::json!({"x-client": "one"})) - ); - assert_eq!(stored.request_body, Some(serde_json::json!({"request": 1}))); - assert_eq!( - stored.provider_request_body, - Some(serde_json::json!({"provider_request": 2})) - ); - assert_eq!( - stored.response_body, - Some(serde_json::json!({"response": 3})) - ); - assert_eq!( - stored.client_response_body, - Some(serde_json::json!({"client_response": 4})) - ); - assert_eq!( - stored.request_body_state, - Some(UsageBodyCaptureState::Reference) - ); - assert_eq!( - stored.request_body_ref.as_deref(), - Some("usage://request/canonical-capture/request_body") - ); + assert!(stored.request_headers.is_none()); + assert!(stored.request_body.is_none()); + assert!(stored.provider_request_body.is_none()); + assert!(stored.response_body.is_none()); + assert!(stored.client_response_body.is_none()); + assert!(stored.request_body_state.is_none()); + assert!(stored.request_body_ref.is_none()); assert_eq!( stored.request_metadata.as_ref().unwrap()["trace_id"], "canonical-trace" @@ -441,41 +440,36 @@ async fn sqlite_usage_http_capture_round_trips_and_preserves_sparse_updates() { .await .expect("legacy columns should load"); assert_eq!(legacy_columns, (None, None, None)); - let audit: (String, String, String) = sqlx::query_as( - "SELECT request_headers, request_body_ref, request_body_state FROM usage_http_audits WHERE request_id = 'canonical-capture'", + let audit_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_http_audits WHERE request_id = 'canonical-capture'", ) .fetch_one(&pool) .await - .expect("canonical audit should load"); - assert_eq!( - serde_json::from_str::(&audit.0).expect("header JSON should decode"), - serde_json::json!({"x-client": "one"}) - ); - assert_eq!(audit.1, "usage://request/canonical-capture/request_body"); - assert_eq!(audit.2, "reference"); + .expect("canonical audits should count"); + assert_eq!(audit_count, 0); let blob_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = 'canonical-capture'", ) .fetch_one(&pool) .await .expect("canonical blobs should count"); - assert_eq!(blob_count, 4); + assert_eq!(blob_count, 0); let sparse = sample_usage("canonical-capture", "streaming", "pending", 1_001); let sparse_stored = writer .upsert(sparse) .await .expect("sparse lifecycle update should upsert"); - assert_eq!(sparse_stored.request_headers, stored.request_headers); - assert_eq!(sparse_stored.request_body, stored.request_body); - assert_eq!(sparse_stored.response_body, stored.response_body); + assert!(sparse_stored.request_headers.is_none()); + assert!(sparse_stored.request_body.is_none()); + assert!(sparse_stored.response_body.is_none()); let sparse_blob_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = 'canonical-capture'", ) .fetch_one(&pool) .await .expect("preserved blobs should count"); - assert_eq!(sparse_blob_count, 4); + assert_eq!(sparse_blob_count, 0); let mut clear = sample_usage("canonical-capture", "streaming", "pending", 1_002); clear.request_body = Some(serde_json::json!({"residual": true})); @@ -487,35 +481,29 @@ async fn sqlite_usage_http_capture_round_trips_and_preserves_sparse_updates() { .expect("explicit none capture should clear"); assert!(cleared.request_body.is_none()); assert!(cleared.request_body_ref.is_none()); - assert_eq!( - cleared.request_body_state, - Some(UsageBodyCaptureState::None) - ); - assert_eq!(cleared.provider_request_body, stored.provider_request_body); + assert!(cleared.request_body_state.is_none()); + assert!(cleared.provider_request_body.is_none()); let cleared_blob_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = 'canonical-capture'", ) .fetch_one(&pool) .await .expect("remaining blobs should count"); - assert_eq!(cleared_blob_count, 3); + assert_eq!(cleared_blob_count, 0); let reader = SqliteUsageReadRepository::new(pool.clone()); let resolved = reader .resolve_body_ref("usage://request/canonical-capture/provider_request_body") .await .expect("body ref should resolve"); - assert_eq!(resolved, stored.provider_request_body); + assert!(resolved.is_none()); let loaded = reader .find_by_request_id("canonical-capture") .await .expect("canonical usage should load") .expect("canonical usage should exist"); - assert_eq!( - loaded.provider_request_headers, - stored.provider_request_headers - ); - assert_eq!(loaded.provider_request_body, stored.provider_request_body); + assert!(loaded.provider_request_headers.is_none()); + assert!(loaded.provider_request_body.is_none()); } #[tokio::test] @@ -536,25 +524,18 @@ async fn sqlite_usage_http_read_falls_back_to_legacy_inline_and_compressed_colum .upsert(captured) .await .expect("temporary canonical body should upsert"); - let payload: Vec = sqlx::query_scalar( - "SELECT payload_gzip FROM usage_body_blobs WHERE request_id = 'legacy-capture' AND body_field = 'request_body'", - ) - .fetch_one(&pool) - .await - .expect("temporary gzip should load"); sqlx::query( r#" DELETE FROM usage_http_audits WHERE request_id = 'legacy-capture'; DELETE FROM usage_body_blobs WHERE request_id = 'legacy-capture'; UPDATE "usage" SET request_headers = '{"legacy":true}', - request_body_compressed = ?, + request_body = '{"compressed":true}', response_body = '{"inline":true}', request_metadata = '{"request_body_ref":"usage://request/legacy-capture/request_body"}' WHERE request_id = 'legacy-capture'; "#, ) - .bind(payload) .execute(&pool) .await .expect("legacy capture should seed"); @@ -592,10 +573,7 @@ WHERE request_id = 'legacy-capture'; .expect("explicit none should clear legacy fallback storage"); assert!(cleared.request_body.is_none()); assert!(cleared.request_body_ref.is_none()); - assert_eq!( - cleared.request_body_state, - Some(UsageBodyCaptureState::None) - ); + assert!(cleared.request_body_state.is_none()); assert!(cleared .request_metadata .as_ref() @@ -610,6 +588,78 @@ WHERE request_id = 'legacy-capture'; assert!(compressed_after_clear.is_none()); } +#[tokio::test] +async fn sqlite_usage_body_refs_enforce_request_and_field_ownership() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_stats_targets(&pool).await; + + let writer = SqliteUsageWriteRepository::new(pool.clone()); + for (request_id, updated_at) in [("ref-target", 3_000), ("ref-owner", 3_001)] { + writer + .upsert(sample_usage(request_id, "completed", "settled", updated_at)) + .await + .expect("usage should seed"); + } + + let mismatched_payload = gzip_json_for_test(&serde_json::json!({ + "secret": "belongs to ref-owner response" + })); + sqlx::query( + "INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip) VALUES (?, ?, ?, ?)", + ) + .bind("usage://request/ref-target/request_body") + .bind("ref-owner") + .bind("response_body") + .bind(mismatched_payload) + .execute(&pool) + .await + .expect("mismatched legacy blob should seed"); + + let reader = SqliteUsageReadRepository::new(pool.clone()); + assert_eq!( + reader + .resolve_body_ref("usage://request/ref-target/request_body") + .await + .expect("mismatched blob lookup should remain safe"), + None + ); + + sqlx::query("DELETE FROM usage_body_blobs") + .execute(&pool) + .await + .expect("mismatched blob should clear"); + sqlx::query("UPDATE \"usage\" SET request_body = ? WHERE request_id = ?") + .bind(r#"{"secret":"belongs to ref-owner request"}"#) + .bind("ref-owner") + .execute(&pool) + .await + .expect("legacy owner body should seed"); + sqlx::query( + "INSERT INTO usage_http_audits (request_id, request_body_ref, body_capture_mode) VALUES (?, ?, ?)", + ) + .bind("ref-target") + .bind("usage://request/ref-owner/request_body") + .bind("ref_backed") + .execute(&pool) + .await + .expect("cross-request audit ref should seed"); + + let target = reader + .find_by_request_id("ref-target") + .await + .expect("target usage should load") + .expect("target usage should exist"); + assert!(target.request_body_ref.is_none()); + assert!(target.request_body.is_none()); +} + #[tokio::test] async fn sqlite_usage_canonical_snapshots_round_trip_preserve_sparse_and_clear_terminal() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1022,16 +1072,13 @@ VALUES ( .expect("stale blobs should count"); assert_eq!(stale_blobs, 0); - let detail_blob: Vec = sqlx::query_scalar( - "SELECT payload_gzip FROM usage_body_blobs WHERE request_id = 'cleanup-detail'", + let detail_blobs: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = 'cleanup-detail'", ) .fetch_one(&pool) .await - .expect("externalized body should load"); - assert_eq!( - super::inflate_usage_json_value(&detail_blob).expect("body gzip should decode"), - serde_json::json!({"detail": true}) - ); + .expect("purged body blobs should count"); + assert_eq!(detail_blobs, 0); let detail_inline: Option = sqlx::query_scalar( "SELECT request_body FROM \"usage\" WHERE request_id = 'cleanup-detail'", ) @@ -1050,16 +1097,13 @@ VALUES ( serde_json::from_str::(&legacy_metadata).expect("valid metadata"), serde_json::json!({"trace": "kept"}) ); - let legacy_ref: String = sqlx::query_scalar( - "SELECT request_body_ref FROM usage_http_audits WHERE request_id = 'cleanup-legacy'", + let legacy_audits: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_http_audits WHERE request_id = 'cleanup-legacy'", ) .fetch_one(&pool) .await - .expect("legacy ref should migrate"); - assert_eq!( - legacy_ref, - "usage://request/cleanup-legacy/request_body".to_string() - ); + .expect("legacy audit refs should count"); + assert_eq!(legacy_audits, 0); let disabled_key: i64 = sqlx::query_scalar("SELECT is_active FROM api_keys WHERE id = 'cleanup-disable-key'") @@ -1182,6 +1226,20 @@ VALUES ( .await .expect("audit headers should remain"); assert!(audit_headers.is_some()); + let audit_body_ref: Option = sqlx::query_scalar( + "SELECT request_body_ref FROM usage_http_audits WHERE request_id = 'cleanup-before-now'", + ) + .fetch_one(&pool) + .await + .expect("audit body ref should load"); + assert!(audit_body_ref.is_none()); + let body_blobs: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = 'cleanup-before-now'", + ) + .fetch_one(&pool) + .await + .expect("body blobs should count"); + assert_eq!(body_blobs, 0); } #[tokio::test] @@ -1211,6 +1269,89 @@ async fn sqlite_usage_write_repository_does_not_regress_void_usage() { assert_eq!(existing.updated_at_unix_secs, 1_000); } +#[tokio::test] +async fn sqlite_stale_terminal_event_is_a_full_transaction_noop() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_stats_targets(&pool).await; + + let repository = SqliteUsageWriteRepository::new(pool.clone()); + let mut newer = sample_usage("request-stale-terminal", "completed", "pending", 2_000); + newer.candidate_id = Some("candidate-new".to_string()); + newer.route_kind = Some("route-new".to_string()); + repository + .upsert(newer) + .await + .expect("newer terminal usage should upsert"); + + let counter_rows_before: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas") + .fetch_one(&pool) + .await + .expect("counter rows should count"); + let routing_before: (Option, Option) = sqlx::query_as( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?", + ) + .bind("request-stale-terminal") + .fetch_one(&pool) + .await + .expect("routing snapshot should load"); + let settlement_before: (String, Option) = sqlx::query_as( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?", + ) + .bind("request-stale-terminal") + .fetch_one(&pool) + .await + .expect("settlement snapshot should load"); + + let mut stale = sample_usage("request-stale-terminal", "failed", "void", 1_999); + stale.status_code = Some(503); + stale.total_cost_usd = Some(99.0); + stale.actual_total_cost_usd = Some(98.0); + stale.candidate_id = Some("candidate-stale".to_string()); + stale.route_kind = Some("route-stale".to_string()); + let stored = repository + .upsert(stale) + .await + .expect("stale terminal usage should be ignored"); + + assert_eq!(stored.status, "completed"); + assert_eq!(stored.billing_status, "pending"); + assert_eq!(stored.status_code, Some(200)); + assert_eq!(stored.total_cost_usd, 0.5); + assert_eq!(stored.routing_candidate_id(), Some("candidate-new")); + assert_eq!(stored.routing_route_kind(), Some("route-new")); + assert_eq!(stored.updated_at_unix_secs, 2_000); + + let counter_rows_after: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas") + .fetch_one(&pool) + .await + .expect("counter rows should count"); + let routing_after: (Option, Option) = sqlx::query_as( + "SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?", + ) + .bind("request-stale-terminal") + .fetch_one(&pool) + .await + .expect("routing snapshot should load"); + let settlement_after: (String, Option) = sqlx::query_as( + "SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?", + ) + .bind("request-stale-terminal") + .fetch_one(&pool) + .await + .expect("settlement snapshot should load"); + + assert_eq!(counter_rows_after, counter_rows_before); + assert_eq!(routing_after, routing_before); + assert_eq!(settlement_after, settlement_before); +} + #[tokio::test] async fn sqlite_usage_write_repository_does_not_reopen_void_failure_from_late_streaming() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1621,8 +1762,8 @@ async fn sqlite_usage_write_repository_cleanup_uses_failed_candidate_status_when .await .expect("pending usage should upsert"); - // request-upstream-reset has a failed candidate carrying a concrete 502 status - // and a connection-reset message — cleanup should use them instead of 504. + // request-upstream-reset has a failed candidate carrying a concrete 502 status. + // Cleanup keeps the status but must not copy the candidate's diagnostic text. // request-stuck has only a still-pending candidate, so cleanup should fall back to 504. sqlx::query( r#" @@ -1654,6 +1795,33 @@ INSERT INTO request_candidates ( assert_eq!(summary.recovered, 0); assert_eq!(summary.failed, 2); + let raw_messages = sqlx::query_as::<_, (String, Option, Option)>( + r#" +SELECT request_id, error_message, error_category +FROM "usage" +WHERE request_id IN ('request-upstream-reset', 'request-stuck') +ORDER BY request_id +"#, + ) + .fetch_all(&pool) + .await + .expect("raw stale usage diagnostics should load"); + assert_eq!( + raw_messages, + vec![ + ( + "request-stuck".to_string(), + None, + Some("server_error".to_string()) + ), + ( + "request-upstream-reset".to_string(), + None, + Some("server_error".to_string()) + ), + ] + ); + let reset = repository .find_by_request_id("request-upstream-reset") .await @@ -1661,10 +1829,8 @@ INSERT INTO request_candidates ( .expect("upstream-reset usage should exist"); assert_eq!(reset.status, "failed"); assert_eq!(reset.status_code, Some(502)); - assert_eq!( - reset.error_message.as_deref(), - Some("upstream connection reset by peer") - ); + assert_eq!(reset.error_message, None); + assert_eq!(reset.error_category.as_deref(), Some("server_error")); let stuck = repository .find_by_request_id("request-stuck") @@ -1673,10 +1839,59 @@ INSERT INTO request_candidates ( .expect("stuck usage should exist"); assert_eq!(stuck.status, "failed"); assert_eq!(stuck.status_code, Some(504)); - assert!(stuck - .error_message - .as_deref() - .is_some_and(|message| message.contains("超过 5 分钟未完成"))); + assert_eq!(stuck.error_message, None); + assert_eq!(stuck.error_category.as_deref(), Some("server_error")); +} + +#[tokio::test] +async fn sqlite_stale_cleanup_derives_category_from_candidate_status_code() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_stats_targets(&pool).await; + + let repository = SqliteUsageWriteRepository::new(pool.clone()); + repository + .upsert(sample_usage( + "request-client-error-cleanup", + "pending", + "pending", + 1, + )) + .await + .expect("pending usage should upsert"); + sqlx::query( + r#" +INSERT INTO request_candidates ( + id, request_id, candidate_index, retry_index, status, status_code, + is_cached, created_at, started_at, finished_at +) VALUES ('candidate-client-error-cleanup', 'request-client-error-cleanup', 0, 0, + 'failed', 429, 0, 1, 2, 3) +"#, + ) + .execute(&pool) + .await + .expect("failed candidate should seed"); + + let summary = repository + .cleanup_stale_pending_requests(2, 10, 5, 10) + .await + .expect("cleanup should run"); + assert_eq!(summary.failed, 1); + + let stored = repository + .find_by_request_id("request-client-error-cleanup") + .await + .expect("usage should load") + .expect("usage should exist"); + assert_eq!(stored.status_code, Some(429)); + assert_eq!(stored.error_category.as_deref(), Some("client_error")); + assert_eq!(stored.error_message, None); } #[tokio::test] @@ -2095,10 +2310,8 @@ async fn sqlite_first_byte_fast_path_preserves_lifecycle_state_and_counters() { duplicate.request_metadata.as_ref().unwrap()["trace_id"], "pending-first-byte-duplicate" ); - assert_eq!( - duplicate.request_body, - Some(serde_json::json!({"prompt": "first-byte-duplicate"})) - ); + assert!(duplicate.request_body.is_none()); + assert!(duplicate.request_body_ref.is_none()); let unique = repository .find_by_request_id("first-byte-unique") @@ -2152,6 +2365,52 @@ WHERE request_id = 'first-byte-missing' assert_eq!(missing_counter_delta, 1); } +#[tokio::test] +async fn sqlite_first_byte_fast_path_rejects_stale_revision() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_stats_targets(&pool).await; + let repository = SqliteUsageWriteRepository::new(pool.clone()); + + let mut current = sample_usage("first-byte-stale", "streaming", "pending", 2_000); + current.finalized_at_unix_secs = None; + current.provider_name = "current-provider".to_string(); + current.model = "current-model".to_string(); + current.response_time_ms = Some(50); + repository + .upsert(current) + .await + .expect("current streaming usage should seed"); + + let mut stale = sample_usage("first-byte-stale", "streaming", "pending", 1_999); + stale.finalized_at_unix_secs = None; + stale.provider_name = "stale-provider".to_string(); + stale.model = "stale-model".to_string(); + stale.status_code = Some(503); + stale.response_time_ms = Some(999); + repository + .upsert_first_byte(stale) + .await + .expect("stale first-byte usage should be ignored"); + + let stored = repository + .find_by_request_id("first-byte-stale") + .await + .expect("streaming usage should load") + .expect("streaming usage should exist"); + assert_eq!(stored.updated_at_unix_secs, 2_000); + assert_eq!(stored.provider_name, "current-provider"); + assert_eq!(stored.model, "current-model"); + assert_eq!(stored.status_code, Some(200)); + assert_eq!(stored.response_time_ms, Some(50)); +} + #[tokio::test] async fn sqlite_pending_batch_is_atomic_and_persists_auxiliary_state() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -2177,8 +2436,8 @@ async fn sqlite_pending_batch_is_atomic_and_persists_auxiliary_state() { sqlx::query( r#" -CREATE TRIGGER reject_second_pending_audit -BEFORE INSERT ON usage_http_audits + CREATE TRIGGER reject_second_pending_routing_snapshot + BEFORE INSERT ON usage_routing_snapshots WHEN NEW.request_id = 'pending-batch-second' BEGIN SELECT RAISE(ABORT, 'reject pending batch test row'); @@ -2207,7 +2466,7 @@ END assert_eq!(rolled_back_usage, 0); assert_eq!(rolled_back_deltas, 0); - sqlx::query("DROP TRIGGER reject_second_pending_audit") + sqlx::query("DROP TRIGGER reject_second_pending_routing_snapshot") .execute(&pool) .await .expect("rollback trigger should drop"); @@ -2229,7 +2488,7 @@ SELECT .fetch_one(&pool) .await .expect("pending batch auxiliary rows should count"); - assert_eq!(committed, (2, 2, 1, 2, 2)); + assert_eq!(committed, (2, 0, 0, 2, 2)); let provider_deltas: i64 = sqlx::query_scalar( r#" diff --git a/crates/aether-data/adapters/sqlite/src/users.rs b/crates/aether-data/adapters/sqlite/src/users.rs index 00fc6b42f..f7c0eaf0b 100644 --- a/crates/aether-data/adapters/sqlite/src/users.rs +++ b/crates/aether-data/adapters/sqlite/src/users.rs @@ -3,11 +3,14 @@ use chrono::{DateTime, TimeZone, Utc}; use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::users::{ - normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, + is_valid_bcrypt_hash, last_oauth_unbind_denial, normalize_user_group_name, + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy, - UserExportSummary, UserReadRepository, + UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED, + LAST_ACTIVE_ADMIN_UPDATE_DENIED, }; use aether_data_contracts::DataLayerError; @@ -25,6 +28,130 @@ SELECT FROM users "#; +const SQLITE_ACTIVE_ADMIN_UPDATE_GUARD: &str = r#" + AND ( + ? = 0 + OR COALESCE(LOWER(role), '') != 'admin' + OR is_active = 0 + OR is_deleted != 0 + OR EXISTS ( + SELECT 1 + FROM users AS other_admin + WHERE other_admin.id != users.id + AND LOWER(other_admin.role) = 'admin' + AND other_admin.is_active = 1 + AND other_admin.is_deleted = 0 + ) + ) +"#; + +const SQLITE_DELETE_USER_SQL: &str = r#" +DELETE FROM users +WHERE id = ? + AND ( + COALESCE(LOWER(role), '') != 'admin' + OR is_active = 0 + OR is_deleted != 0 + OR EXISTS ( + SELECT 1 + FROM users AS other_admin + WHERE other_admin.id != users.id + AND LOWER(other_admin.role) = 'admin' + AND other_admin.is_active = 1 + AND other_admin.is_deleted = 0 + ) + ) +"#; + +const SQLITE_DELETE_USER_IF_WALLET_ABSENT_SQL: &str = r#" +DELETE FROM users +WHERE id = ? + AND NOT EXISTS ( + SELECT 1 + FROM wallets AS wallet + WHERE wallet.user_id = ? + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = ? + ) + ) + AND ( + COALESCE(LOWER(role), '') != 'admin' + OR is_active = 0 + OR is_deleted != 0 + OR EXISTS ( + SELECT 1 + FROM users AS other_admin + WHERE other_admin.id != users.id + AND LOWER(other_admin.role) = 'admin' + AND other_admin.is_active = 1 + AND other_admin.is_deleted = 0 + ) + ) +"#; + +const SQLITE_DELETE_USER_API_KEYS_SQL: &str = "DELETE FROM api_keys WHERE user_id = ?"; + +const SQLITE_DELETE_USER_DEPENDENTS_SQL: &[&str] = &[ + "DELETE FROM usage_request_admissions WHERE subject_id = ?", + "DELETE FROM usage_cost_reservations WHERE subject_id = ?", + "DELETE FROM gemini_file_mappings WHERE user_id = ?", + "DELETE FROM api_key_provider_mappings WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + SQLITE_DELETE_USER_API_KEYS_SQL, + "DELETE FROM management_tokens WHERE user_id = ?", + "DELETE FROM user_sessions WHERE user_id = ?", + "DELETE FROM user_oauth_links WHERE user_id = ?", + "DELETE FROM user_group_members WHERE user_id = ?", + "DELETE FROM user_preferences WHERE user_id = ?", + "DELETE FROM user_invite_codes WHERE user_id = ?", + "DELETE FROM announcement_reads WHERE user_id = ?", +]; + +const SQLITE_PREPARE_USER_FACTS_FOR_DELETION_SQL: &[&str] = &[ + "UPDATE referral_rewards SET status = CASE WHEN status IN ('pending', 'failed', 'applying') THEN 'voided' ELSE status END, failure_reason = NULL, admin_note = NULL, updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE ? IN (inviter_user_id, invitee_user_id)", + "UPDATE referral_rewards SET failure_reason = NULL, admin_note = NULL, updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE admin_operator_id = ?", + "UPDATE user_referrals SET invite_code_snapshot = 'deleted-user', source_json = NULL, updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE ? IN (inviter_user_id, invitee_user_id)", + "UPDATE user_plan_entitlements SET status = CASE WHEN status = 'active' THEN 'revoked' ELSE status END, expires_at = MIN(expires_at, CAST(strftime('%s', 'now') AS INTEGER)), updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE user_id = ?", + "UPDATE wallets SET status = 'disabled', updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE user_id = ?", + "UPDATE wallets SET status = 'disabled', updated_at = CAST(strftime('%s', 'now') AS INTEGER) WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + "UPDATE audit_logs SET description = 'deleted user event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE user_id = ?", + "UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = ?)", + "UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE wallet_transactions SET description = NULL WHERE operator_id = ?", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order WHERE history_order.user_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.user_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?) AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))", + "UPDATE payment_orders SET gateway_response = NULL WHERE user_id = ?", + "UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE user_id = ?", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE ? IN (requested_by, approved_by, processed_by)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE user_id = ?)", + "UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?))", + "UPDATE redeem_code_batches SET description = NULL WHERE created_by = ?", +]; + +const SQLITE_ANONYMIZE_USER_HISTORY_SQL: &[&str] = &[ + "UPDATE request_candidates SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE video_tasks SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE usage SET username = NULL, api_key_name = NULL WHERE user_id = ?", + "UPDATE stats_user_daily SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_summary SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_model SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_api_format SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_model_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_provider SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_model SET username = NULL WHERE user_id = ?", + "UPDATE stats_user_daily_cost_savings_model_provider SET username = NULL WHERE user_id = ?", +]; + +const SQLITE_ANONYMIZE_USER_API_KEY_HISTORY_SQL: &str = + "UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id IN (SELECT id FROM api_keys WHERE user_id = ?)"; + const USER_EXPORT_COLUMNS: &str = r#" SELECT id, @@ -65,6 +192,7 @@ SELECT allowed_models_mode, is_active, is_deleted, + security_version, created_at, last_login_at FROM users @@ -87,6 +215,7 @@ SELECT users.allowed_models_mode AS allowed_models_mode, users.is_active AS is_active, users.is_deleted AS is_deleted, + users.security_version AS security_version, users.created_at AS created_at, users.last_login_at AS last_login_at FROM users @@ -128,6 +257,7 @@ const USER_SESSION_COLUMNS: &str = r#" SELECT id, user_id, + security_version, client_device_id, device_label, refresh_token_hash, @@ -227,6 +357,112 @@ impl SqliteUserReadRepository { let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; rows.iter().map(map_user_group_member_row).collect() } + + async fn delete_local_auth_user_inner( + &self, + user_id: &str, + require_wallet_absent: bool, + ) -> Result { + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + if require_wallet_absent { + let wallet_exists: Option = sqlx::query_scalar( + r#" +SELECT 1 +FROM wallets AS wallet +WHERE wallet.user_id = ? + OR EXISTS ( + SELECT 1 + FROM api_keys AS api_key + WHERE api_key.id = wallet.api_key_id + AND api_key.user_id = ? + ) +LIMIT 1 + "#, + ) + .bind(user_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if wallet_exists.is_some() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + } + for sql in SQLITE_PREPARE_USER_FACTS_FOR_DELETION_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + for sql in SQLITE_ANONYMIZE_USER_HISTORY_SQL { + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + sqlx::query(SQLITE_ANONYMIZE_USER_API_KEY_HISTORY_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + for sql in SQLITE_DELETE_USER_DEPENDENTS_SQL { + if require_wallet_absent && *sql == SQLITE_DELETE_USER_API_KEYS_SQL { + continue; + } + sqlx::query(sql) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + let result = if require_wallet_absent { + sqlx::query(SQLITE_DELETE_USER_IF_WALLET_ABSENT_SQL) + .bind(user_id) + .bind(user_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + } else { + sqlx::query(SQLITE_DELETE_USER_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + }; + if result.rows_affected() > 0 { + if require_wallet_absent { + sqlx::query(SQLITE_DELETE_USER_API_KEYS_SQL) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + return Ok(true); + } + let blocked_active_admin: Option = sqlx::query_scalar( + "SELECT 1 FROM users WHERE id = ? AND LOWER(role) = 'admin' AND is_active = 1 AND is_deleted = 0 LIMIT 1", + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + tx.rollback().await.map_sql_err()?; + if blocked_active_admin.is_some() { + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_DELETE_DENIED.to_string(), + )); + } + Ok(false) + } } #[async_trait] @@ -572,6 +808,92 @@ WHERE id = ? } } + /// BEGIN IMMEDIATE serializes writers while the complete snapshot is + /// compared and restored, preventing a rollback from overwriting a newer + /// administrator update. + async fn restore_user_group_if_matches( + &self, + expected: &StoredUserGroup, + restored: &StoredUserGroup, + ) -> Result { + if expected.id != restored.id || expected.id.trim().is_empty() { + return Ok(false); + } + + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let mut builder = QueryBuilder::::new(USER_GROUP_COLUMNS); + builder.push(" WHERE id = ").push_bind(&expected.id); + let row = builder + .build() + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_user_group_row(&row)?; + if ¤t != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let result = sqlx::query( + r#" +UPDATE user_groups +SET name = ?, + normalized_name = ?, + description = ?, + priority = ?, + allowed_providers = ?, + allowed_providers_mode = ?, + allowed_api_formats = ?, + allowed_api_formats_mode = ?, + allowed_models = ?, + allowed_models_mode = ?, + rate_limit = ?, + rate_limit_mode = ?, + created_at = ?, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(&restored.name) + .bind(&restored.normalized_name) + .bind(&restored.description) + .bind(restored.priority) + .bind(json_string_from_option_vec( + restored.allowed_providers.as_ref(), + )) + .bind(&restored.allowed_providers_mode) + .bind(json_string_from_option_vec( + restored.allowed_api_formats.as_ref(), + )) + .bind(&restored.allowed_api_formats_mode) + .bind(json_string_from_option_vec( + restored.allowed_models.as_ref(), + )) + .bind(&restored.allowed_models_mode) + .bind(restored.rate_limit) + .bind(&restored.rate_limit_mode) + .bind(restored.created_at.map(|value| value.timestamp())) + .bind(restored.updated_at.map(|value| value.timestamp())) + .bind(&restored.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn delete_user_group(&self, group_id: &str) -> Result { let result = sqlx::query("DELETE FROM user_groups WHERE id = ?") .bind(group_id) @@ -598,7 +920,11 @@ WHERE id = ? group_id: &str, user_ids: &[String], ) -> Result, DataLayerError> { - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; sqlx::query("DELETE FROM user_group_members WHERE group_id = ?") .bind(group_id) .execute(&mut *tx) @@ -670,7 +996,20 @@ WHERE user_group_members.user_id IN ( user_id: &str, group_ids: &[String], ) -> Result, DataLayerError> { - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(Vec::new()); + } sqlx::query("DELETE FROM user_group_members WHERE user_id = ?") .bind(user_id) .execute(&mut *tx) @@ -692,20 +1031,112 @@ WHERE user_group_members.user_id IN ( self.list_user_groups_for_user(user_id).await } + async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + let expected = normalized_ids(expected_group_ids); + let restored = normalized_ids(restored_group_ids); + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let current = sqlx::query_scalar::<_, String>( + "SELECT group_id FROM user_group_members WHERE user_id = ? ORDER BY group_id ASC", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if current != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + if !restored.is_empty() { + let mut builder = QueryBuilder::::new( + "SELECT COUNT(*) AS count FROM user_groups WHERE id IN (", + ); + { + let mut separated = builder.separated(", "); + for group_id in &restored { + separated.push_bind(group_id); + } + } + builder.push(")"); + let count = builder + .build() + .fetch_one(&mut *tx) + .await + .map_sql_err()? + .try_get::("count") + .map_sql_err()?; + if count != restored.len() as i64 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + } + sqlx::query("DELETE FROM user_group_members WHERE user_id = ?") + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + let now = current_unix_secs(); + for group_id in restored { + sqlx::query( + "INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)", + ) + .bind(group_id) + .bind(user_id) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn add_user_to_group( &self, group_id: &str, user_id: &str, ) -> Result { + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } let result = sqlx::query( "INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)", ) .bind(group_id) .bind(user_id) .bind(current_unix_secs()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(result.rows_affected() > 0) } @@ -820,6 +1251,75 @@ WHERE user_group_members.user_id IN ( Ok(self.fetch_auth_rows(builder).await?.into_iter().next()) } + async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: DateTime, + _provider_enabled_snapshot: bool, + ) -> Result { + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let provider_enabled: Option = + sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?") + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_enabled != Some(true) { + tx.rollback().await.map_sql_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::ProviderUnavailable); + } + let row = sqlx::query(&format!( + "{USER_AUTH_COLUMNS_QUALIFIED} JOIN user_oauth_links ON users.id = user_oauth_links.user_id WHERE user_oauth_links.provider_type = ? AND user_oauth_links.provider_user_id = ? LIMIT 1" + )) + .bind(provider_type) + .bind(provider_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(ResolveOAuthLinkedUserOutcome::NotLinked); + }; + let mut user = map_user_auth_row(&row)?; + sqlx::query( + "UPDATE user_oauth_links SET provider_username = COALESCE(?, provider_username), provider_email = COALESCE(?, provider_email), extra_data = COALESCE(?, extra_data), last_login_at = ? WHERE provider_type = ? AND provider_user_id = ?", + ) + .bind(provider_username) + .bind(provider_email) + .bind(optional_json_string(extra_data, "user_oauth_links.extra_data")?) + .bind(touched_at.timestamp()) + .bind(provider_type) + .bind(provider_user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if let Some(verified_email) = verified_email { + let result = sqlx::query( + "UPDATE users SET email_verified = 1, updated_at = ? WHERE id = ? AND email_verified = 0 AND LOWER(TRIM(email)) = LOWER(TRIM(?))", + ) + .bind(touched_at.timestamp()) + .bind(&user.id) + .bind(verified_email) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() == 1 { + user.email_verified = true; + } + } + tx.commit().await.map_sql_err()?; + Ok(ResolveOAuthLinkedUserOutcome::Linked(user)) + } + async fn touch_oauth_link( &self, provider_type: &str, @@ -858,6 +1358,7 @@ WHERE provider_type = ? async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: DateTime, ) -> Result, DataLayerError> { @@ -869,11 +1370,12 @@ INSERT INTO users ( allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode, is_active, is_deleted, created_at, updated_at, last_login_at ) -VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?) +VALUES (?, ?, ?, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?) "#, ) .bind(&user_id) .bind(email) + .bind(email_verified) .bind(username) .bind(created_at.timestamp()) .bind(created_at.timestamp()) @@ -925,7 +1427,20 @@ VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inh Ok(total.max(0) as u64) } - async fn upsert_user_oauth_link( + async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + let exists: Option = + sqlx::query_scalar("SELECT 1 FROM user_oauth_links WHERE provider_type = ? LIMIT 1") + .bind(provider_type) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + Ok(exists.is_some()) + } + + async fn bind_user_oauth_link_if_provider_enabled( &self, user_id: &str, provider_type: &str, @@ -934,69 +1449,219 @@ VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inh provider_email: Option<&str>, extra_data: Option, linked_at: DateTime, - ) -> Result<(), DataLayerError> { + _provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result { let extra_data = optional_json_string(extra_data, "user_oauth_links.extra_data")?; - let updated = sqlx::query( - r#" -UPDATE user_oauth_links -SET provider_user_id = ?, - provider_username = ?, - provider_email = ?, - extra_data = ?, - last_login_at = ? -WHERE user_id = ? - AND provider_type = ? + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let provider_enabled: Option = + sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?") + .bind(provider_type) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if provider_enabled.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderNotFound); + } + if provider_enabled != Some(true) { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::ProviderDisabled); + } + if let Some(expectation) = session_expectation { + let session_is_current: Option = sqlx::query_scalar( + r#" +SELECT 1 +FROM users +JOIN user_sessions + ON user_sessions.user_id = users.id +WHERE users.id = ? + AND users.is_active = 1 + AND users.is_deleted = 0 + AND users.security_version = ? + AND user_sessions.id = ? + AND user_sessions.security_version = ? + AND user_sessions.client_device_id = ? + AND user_sessions.revoked_at IS NULL + AND user_sessions.expires_at > MAX(?, CAST(strftime('%s', 'now') AS INTEGER)) "#, + ) + .bind(user_id) + .bind(expectation.security_version) + .bind(&expectation.session_id) + .bind(expectation.security_version) + .bind(&expectation.client_device_id) + .bind(expectation.checked_at.timestamp()) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if session_is_current.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::SessionUnavailable); + } + } else { + let user_exists: Option = sqlx::query_scalar("SELECT 1 FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::UserNotFound); + } + } + if let Some(owner) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = ? AND provider_user_id = ? LIMIT 1", ) + .bind(provider_type) .bind(provider_user_id) - .bind(provider_username) - .bind(provider_email) - .bind(extra_data.as_deref()) - .bind(linked_at.timestamp()) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + { + tx.rollback().await.map_sql_err()?; + return Ok(if owner == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }); + } + if sqlx::query_scalar::<_, i32>( + "SELECT 1 FROM user_oauth_links WHERE user_id = ? AND provider_type = ? LIMIT 1", + ) .bind(user_id) .bind(provider_type) - .execute(&self.pool) + .fetch_optional(&mut *tx) .await - .map_sql_err()?; - if updated.rows_affected() == 0 { - sqlx::query( - r#" + .map_sql_err()? + .is_some() + { + tx.rollback().await.map_sql_err()?; + return Ok(BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider); + } + sqlx::query( + r#" INSERT INTO user_oauth_links ( id, user_id, provider_type, provider_user_id, provider_username, provider_email, extra_data, linked_at, last_login_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) "#, - ) - .bind(uuid::Uuid::new_v4().to_string()) - .bind(user_id) - .bind(provider_type) - .bind(provider_user_id) - .bind(provider_username) - .bind(provider_email) - .bind(extra_data.as_deref()) - .bind(linked_at.timestamp()) - .bind(linked_at.timestamp()) - .execute(&self.pool) - .await - .map_sql_err()?; - } - Ok(()) + ) + .bind(uuid::Uuid::new_v4().to_string()) + .bind(user_id) + .bind(provider_type) + .bind(provider_user_id) + .bind(provider_username) + .bind(provider_email) + .bind(extra_data.as_deref()) + .bind(linked_at.timestamp()) + .bind(linked_at.timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(BindUserOAuthLinkOutcome::Bound) + } + + async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: DateTime, + ) -> Result { + let result = sqlx::query( + "UPDATE users SET email_verified = 1, updated_at = ? WHERE id = ? AND email_verified = 0 AND LOWER(TRIM(email)) = LOWER(TRIM(?))", + ) + .bind(verified_at.timestamp()) + .bind(user_id) + .bind(verified_email) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) } async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { + local_password_login_allowed: bool, + _enabled_provider_types_snapshot: &[String], + ) -> Result { + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let user = sqlx::query("SELECT auth_source, password_hash FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(user) = user else { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + }; + let auth_source = user.try_get::("auth_source").map_sql_err()?; + let password_hash = user + .try_get::, _>("password_hash") + .map_sql_err()?; + let provider_types = sqlx::query_scalar::<_, String>( + "SELECT provider_type FROM user_oauth_links WHERE user_id = ?", + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if !provider_types.iter().any(|value| value == provider_type) { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + let enabled_provider_types = sqlx::query_scalar::<_, String>( + r#" +SELECT user_oauth_links.provider_type +FROM user_oauth_links +JOIN oauth_providers + ON oauth_providers.provider_type = user_oauth_links.provider_type +WHERE user_oauth_links.user_id = ? + AND oauth_providers.is_enabled = 1 +"#, + ) + .bind(user_id) + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let has_remaining_enabled_oauth_link = enabled_provider_types + .iter() + .any(|value| value != provider_type); + if !has_remaining_enabled_oauth_link { + if let Some(outcome) = last_oauth_unbind_denial( + &auth_source, + password_hash.as_deref(), + local_password_login_allowed, + ) { + tx.rollback().await.map_sql_err()?; + return Ok(outcome); + } + } let result = sqlx::query("DELETE FROM user_oauth_links WHERE user_id = ? AND provider_type = ?") .bind(user_id) .bind(provider_type) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; - Ok(result.rows_affected() > 0) + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + tx.commit().await.map_sql_err()?; + Ok(DeleteUserOAuthLinkOutcome::Deleted) } async fn get_or_create_ldap_auth_user( @@ -1036,14 +1701,18 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { let now = chrono::Utc::now().timestamp(); let result = sqlx::query( - "UPDATE users SET email = COALESCE(?, email), username = COALESCE(?, username), updated_at = ? WHERE id = ?", + "UPDATE users SET email = CASE WHEN ? THEN ? ELSE email END, email_verified = COALESCE(?, email_verified), username = COALESCE(?, username), updated_at = ? WHERE id = ?", ) + .bind(email_present) .bind(email) + .bind(email_verified) .bind(username) .bind(now) .bind(user_id) @@ -1056,13 +1725,181 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) self.find_user_auth_by_id(user_id).await } + async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + if expected_auth.id != restored_auth.id + || expected_export.id != expected_auth.id + || restored_export.id != restored_auth.id + { + return Ok(false); + } + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let auth_row = sqlx::query(&format!("{USER_AUTH_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let export_row = sqlx::query(&format!("{USER_EXPORT_COLUMNS} WHERE id = ? LIMIT 1")) + .bind(&expected_auth.id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let (Some(auth_row), Some(export_row)) = (auth_row, export_row) else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current_auth = map_user_auth_row(&auth_row)?; + let current_export = map_user_export_row(&export_row)?; + if !current_auth.matches_restore_state(expected_auth) + || !current_export.matches_restore_state(expected_export) + || current_export.rate_limit != expected_export.rate_limit + || current_export.rate_limit_mode != expected_export.rate_limit_mode + || current_export.model_capability_settings.as_ref() + != expected_model_capability_settings + || current_export.feature_settings.as_ref() != expected_feature_settings + { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + + let removes_active_admin = current_auth.role.eq_ignore_ascii_case("admin") + && current_auth.is_active + && !current_auth.is_deleted + && (!restored_auth.role.eq_ignore_ascii_case("admin") || !restored_auth.is_active); + if removes_active_admin { + let active_admin_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM users WHERE LOWER(role) = 'admin' AND is_active = 1 AND is_deleted = 0", + ) + .fetch_one(&mut *tx) + .await + .map_sql_err()?; + if active_admin_count <= 1 { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + } + + let security_state_changed = expected_auth.role != restored_auth.role + || expected_auth.is_active != restored_auth.is_active; + let result = sqlx::query( + r#" +UPDATE users +SET email = ?, + email_verified = ?, + username = ?, + role = ?, + allowed_providers = ?, + allowed_providers_mode = ?, + allowed_api_formats = ?, + allowed_api_formats_mode = ?, + allowed_models = ?, + allowed_models_mode = ?, + rate_limit = ?, + rate_limit_mode = ?, + model_capability_settings = ?, + feature_settings = ?, + is_active = ?, + security_version = security_version + CASE WHEN ? THEN 1 ELSE 0 END, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(restored_auth.email.as_deref()) + .bind(restored_auth.email_verified) + .bind(&restored_auth.username) + .bind(&restored_auth.role) + .bind(optional_string_list_json( + restored_auth.allowed_providers.clone(), + "users.allowed_providers", + )?) + .bind(&restored_auth.allowed_providers_mode) + .bind(optional_string_list_json( + restored_auth.allowed_api_formats.clone(), + "users.allowed_api_formats", + )?) + .bind(&restored_auth.allowed_api_formats_mode) + .bind(optional_string_list_json( + restored_auth.allowed_models.clone(), + "users.allowed_models", + )?) + .bind(&restored_auth.allowed_models_mode) + .bind(restored_export.rate_limit) + .bind(&restored_export.rate_limit_mode) + .bind(optional_json_string( + restored_model_capability_settings.clone(), + "users.model_capability_settings", + )?) + .bind(optional_json_string( + restored_feature_settings.clone(), + "users.feature_settings", + )?) + .bind(restored_auth.is_active) + .bind(security_state_changed) + .bind(current_unix_secs()) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if result.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + if security_state_changed { + let now = current_unix_secs(); + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'user_security_state_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(now) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(now) + .bind(&expected_auth.id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn update_local_auth_user_password_hash( &self, user_id: &str, password_hash: String, updated_at: DateTime, ) -> Result, DataLayerError> { - let result = sqlx::query("UPDATE users SET password_hash = ?, updated_at = ? WHERE id = ?") + let result = sqlx::query( + "UPDATE users SET password_hash = ?, security_version = security_version + 1, updated_at = ? WHERE id = ?", + ) .bind(password_hash) .bind(updated_at.timestamp()) .bind(user_id) @@ -1075,6 +1912,122 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) self.find_user_auth_by_id(user_id).await } + async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: DateTime, + ) -> Result { + let result = sqlx::query( + r#" +UPDATE users +SET password_hash = ?, + security_version = security_version + 1, + updated_at = ? +WHERE id = ? + AND ((? IS NULL AND password_hash IS NULL) OR password_hash = ?) +"#, + ) + .bind(password_hash) + .bind(updated_at.timestamp()) + .bind(user_id) + .bind(expected_password_hash) + .bind(expected_password_hash) + .execute(&self.pool) + .await + .map_sql_err()?; + Ok(result.rows_affected() == 1) + } + + async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let updated = sqlx::query( + "UPDATE users SET password_hash = ?, security_version = security_version + 1, updated_at = ? WHERE id = ? AND is_deleted = 0", + ) + .bind(password_hash) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'admin_password_reset', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(changed_at.timestamp()) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: DateTime, + ) -> Result { + let mut tx = self.pool.begin().await.map_sql_err()?; + let updated = sqlx::query( + r#" +UPDATE users +SET password_hash = ?, security_version = security_version + 1, updated_at = ? +WHERE id = ? + AND is_active = 1 + AND is_deleted = 0 + AND ((? IS NULL AND password_hash IS NULL) OR password_hash = ?) + AND EXISTS ( + SELECT 1 FROM user_sessions + WHERE user_id = ? AND id = ? AND revoked_at IS NULL AND expires_at > ? + ) +"#, + ) + .bind(next_password_hash) + .bind(changed_at.timestamp()) + .bind(user_id) + .bind(expected_password_hash) + .bind(expected_password_hash) + .bind(user_id) + .bind(current_session_id) + .bind(changed_at.timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() != 1 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let revoked = sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'password_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(changed_at.timestamp()) + .bind(changed_at.timestamp()) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if revoked.rows_affected() == 0 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + async fn update_local_auth_user_admin_fields( &self, user_id: &str, @@ -1089,6 +2042,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) rate_limit: Option, is_active: Option, ) -> Result, DataLayerError> { + let removes_active_admin = role + .as_deref() + .is_some_and(|value| !value.eq_ignore_ascii_case("admin")) + || is_active == Some(false); let allowed_providers_mode = if allowed_providers .as_ref() .is_some_and(|values| !values.is_empty()) @@ -1118,7 +2075,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) } else { "system" }; - let result = sqlx::query( + let update_sql = format!( r#" UPDATE users SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END, @@ -1131,47 +2088,109 @@ SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END, rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END, rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END, is_active = CASE WHEN ? THEN ? ELSE is_active END, + security_version = security_version + CASE WHEN ? THEN 1 ELSE 0 END, updated_at = ? WHERE id = ? +{SQLITE_ACTIVE_ADMIN_UPDATE_GUARD} "#, - ) - .bind(role.is_some()) - .bind(role) - .bind(allowed_providers_present) - .bind(optional_string_list_json( - allowed_providers, - "users.allowed_providers", - )?) - .bind(allowed_providers_present) - .bind(allowed_providers_mode) - .bind(allowed_api_formats_present) - .bind(optional_string_list_json( - allowed_api_formats, - "users.allowed_api_formats", - )?) - .bind(allowed_api_formats_present) - .bind(allowed_api_formats_mode) - .bind(allowed_models_present) - .bind(optional_string_list_json( - allowed_models, - "users.allowed_models", - )?) - .bind(allowed_models_present) - .bind(allowed_models_mode) - .bind(rate_limit_present) - .bind(rate_limit) - .bind(rate_limit_present) - .bind(rate_limit_mode) - .bind(is_active.is_some()) - .bind(is_active) - .bind(chrono::Utc::now().timestamp()) - .bind(user_id) - .execute(&self.pool) - .await - .map_sql_err()?; + ); + let mut tx = self.pool.begin().await.map_sql_err()?; + let current_security_state = sqlx::query("SELECT role, is_active FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let security_state_changed = current_security_state.as_ref().is_some_and(|row| { + role.as_deref().is_some_and(|next_role| { + row.try_get::("role") + .is_ok_and(|current_role| !current_role.eq_ignore_ascii_case(next_role)) + }) || is_active.is_some_and(|next_active| { + row.try_get::("is_active") + .is_ok_and(|current_active| current_active != next_active) + }) + }); + let result = sqlx::query(&update_sql) + .bind(role.is_some()) + .bind(role) + .bind(allowed_providers_present) + .bind(optional_string_list_json( + allowed_providers, + "users.allowed_providers", + )?) + .bind(allowed_providers_present) + .bind(allowed_providers_mode) + .bind(allowed_api_formats_present) + .bind(optional_string_list_json( + allowed_api_formats, + "users.allowed_api_formats", + )?) + .bind(allowed_api_formats_present) + .bind(allowed_api_formats_mode) + .bind(allowed_models_present) + .bind(optional_string_list_json( + allowed_models, + "users.allowed_models", + )?) + .bind(allowed_models_present) + .bind(allowed_models_mode) + .bind(rate_limit_present) + .bind(rate_limit) + .bind(rate_limit_present) + .bind(rate_limit_mode) + .bind(is_active.is_some()) + .bind(is_active) + .bind(security_state_changed) + .bind(chrono::Utc::now().timestamp()) + .bind(user_id) + .bind(removes_active_admin) + .execute(&mut *tx) + .await + .map_sql_err()?; if result.rows_affected() == 0 { + let blocked_active_admin: Option = sqlx::query_scalar( + "SELECT 1 FROM users WHERE id = ? AND LOWER(role) = 'admin' AND is_active = 1 AND is_deleted = 0 LIMIT 1", + ) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + tx.rollback().await.map_sql_err()?; + if removes_active_admin && blocked_active_admin.is_some() { + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } return Ok(None); } + if security_state_changed { + let revoked_at = chrono::Utc::now().timestamp(); + sqlx::query( + "UPDATE user_sessions SET revoked_at = ?, revoke_reason = 'user_security_state_changed', updated_at = ? WHERE user_id = ? AND revoked_at IS NULL", + ) + .bind(revoked_at) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE api_keys SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + "UPDATE management_tokens SET is_active = 0, updated_at = ? WHERE user_id = ? AND is_active = 1", + ) + .bind(revoked_at) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } + tx.commit().await.map_sql_err()?; self.find_user_auth_by_id(user_id).await } @@ -1380,12 +2399,14 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?) } async fn delete_local_auth_user(&self, user_id: &str) -> Result { - let result = sqlx::query("DELETE FROM users WHERE id = ?") - .bind(user_id) - .execute(&self.pool) - .await - .map_sql_err()?; - Ok(result.rows_affected() > 0) + self.delete_local_auth_user_inner(user_id, false).await + } + + async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + self.delete_local_auth_user_inner(user_id, true).await } async fn count_active_admin_users(&self) -> Result { @@ -1505,6 +2526,23 @@ ON CONFLICT(user_id) DO UPDATE SET .or(session.updated_at) .or(session.last_seen_at) .unwrap_or_else(Utc::now); + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let user_is_eligible: Option = sqlx::query_scalar( + "SELECT security_version FROM users WHERE id = ? AND is_active = 1 AND is_deleted = 0 AND security_version = ?", + ) + .bind(&session.user_id) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_is_eligible.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } sqlx::query( r#" UPDATE user_sessions @@ -1517,19 +2555,20 @@ WHERE user_id = ? AND client_device_id = ? AND revoked_at IS NULL AND expires_at .bind(&session.user_id) .bind(&session.client_device_id) .bind(now.timestamp()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; sqlx::query( r#" INSERT INTO user_sessions ( - id, user_id, client_device_id, device_label, device_type, ip_address, user_agent, + id, user_id, security_version, client_device_id, device_label, device_type, ip_address, user_agent, refresh_token_hash, last_seen_at, expires_at, created_at, updated_at -) VALUES (?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) +) VALUES (?, ?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&session.id) .bind(&session.user_id) + .bind(session.security_version) .bind(&session.client_device_id) .bind(session.device_label.as_deref()) .bind(session.ip_address.as_deref()) @@ -1539,9 +2578,99 @@ INSERT INTO user_sessions ( .bind(session.expires_at.unwrap_or(now).timestamp()) .bind(session.created_at.unwrap_or(now).timestamp()) .bind(session.updated_at.unwrap_or(now).timestamp()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + let mut builder = QueryBuilder::::new(USER_SESSION_COLUMNS); + builder + .push(" WHERE user_id = ") + .push_bind(&session.user_id) + .push(" AND id = ") + .push_bind(&session.id) + .push(" LIMIT 1"); + let row = builder + .build() + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let created = row.as_ref().map(map_user_session_row).transpose()?; + tx.commit().await.map_sql_err()?; + Ok(created) + } + + async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + let now = session + .created_at + .or(session.updated_at) + .or(session.last_seen_at) + .unwrap_or_else(Utc::now); + let mut tx = self.pool.begin().await.map_sql_err()?; + let matched = sqlx::query_scalar::<_, String>( + r#" +SELECT password_hash FROM users +WHERE id = ? AND password_hash = ? AND LOWER(auth_source) = 'local' + AND is_active = 1 AND is_deleted = 0 AND security_version = ? +"#, + ) + .bind(&session.user_id) + .bind(expected_password_hash) + .bind(session.security_version) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if matched.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + sqlx::query("UPDATE users SET last_login_at = ? WHERE id = ?") + .bind(now.timestamp()) + .bind(&session.user_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +UPDATE user_sessions +SET revoked_at = ?, revoke_reason = 'replaced_by_new_login', updated_at = ? +WHERE user_id = ? AND client_device_id = ? AND revoked_at IS NULL AND expires_at > ? +"#, + ) + .bind(now.timestamp()) + .bind(now.timestamp()) + .bind(&session.user_id) + .bind(&session.client_device_id) + .bind(now.timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + sqlx::query( + r#" +INSERT INTO user_sessions ( + id, user_id, security_version, client_device_id, device_label, device_type, ip_address, user_agent, + refresh_token_hash, last_seen_at, expires_at, created_at, updated_at +) VALUES (?, ?, ?, ?, ?, 'unknown', ?, ?, ?, ?, ?, ?, ?) +"#, + ) + .bind(&session.id) + .bind(&session.user_id) + .bind(session.security_version) + .bind(&session.client_device_id) + .bind(session.device_label.as_deref()) + .bind(session.ip_address.as_deref()) + .bind(session.user_agent.as_deref()) + .bind(&session.refresh_token_hash) + .bind(session.last_seen_at.unwrap_or(now).timestamp()) + .bind(session.expires_at.unwrap_or(now).timestamp()) + .bind(session.created_at.unwrap_or(now).timestamp()) + .bind(session.updated_at.unwrap_or(now).timestamp()) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; self.find_user_session(&session.user_id, &session.id).await } @@ -1601,7 +2730,7 @@ WHERE user_id = ? AND id = ? &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: DateTime, expires_at: DateTime, @@ -1614,10 +2743,11 @@ UPDATE user_sessions SET prev_refresh_token_hash = ?, rotated_at = ?, refresh_token_hash = ?, expires_at = ?, last_seen_at = ?, ip_address = COALESCE(?, ip_address), user_agent = COALESCE(?, user_agent), updated_at = ? -WHERE user_id = ? AND id = ? +WHERE user_id = ? AND id = ? AND refresh_token_hash = ? + AND revoked_at IS NULL AND expires_at > ? "#, ) - .bind(previous_refresh_token_hash) + .bind(expected_refresh_token_hash) .bind(rotated_at.timestamp()) .bind(next_refresh_token_hash) .bind(expires_at.timestamp()) @@ -1627,6 +2757,8 @@ WHERE user_id = ? AND id = ? .bind(rotated_at.timestamp()) .bind(user_id) .bind(session_id) + .bind(expected_refresh_token_hash) + .bind(rotated_at.timestamp()) .execute(&self.pool) .await .map_sql_err()?; @@ -1676,26 +2808,24 @@ WHERE user_id = ? AND id = ? async fn count_active_local_admin_users_with_valid_password( &self, ) -> Result { - let total: i64 = sqlx::query_scalar( + let hashes = sqlx::query_scalar::<_, String>( r#" -SELECT COUNT(*) +SELECT password_hash FROM users WHERE LOWER(role) = 'admin' AND LOWER(auth_source) = 'local' AND is_deleted = 0 AND is_active = 1 - AND LENGTH(password_hash) = 60 - AND ( - password_hash LIKE '$2a$%' - OR password_hash LIKE '$2b$%' - OR password_hash LIKE '$2y$%' - ) + AND password_hash IS NOT NULL "#, ) - .fetch_one(&self.pool) + .fetch_all(&self.pool) .await .map_sql_err()?; - Ok(total.max(0) as u64) + Ok(hashes + .iter() + .filter(|hash| is_valid_bcrypt_hash(hash)) + .count() as u64) } } @@ -2004,6 +3134,7 @@ fn map_user_auth_row(row: &SqliteRow) -> Result Result, + ) -> StoredUserSessionRecord { + StoredUserSessionRecord::new( + id.to_string(), + user_id.to_string(), + client_device_id.to_string(), + None, + StoredUserSessionRecord::hash_refresh_token(refresh_token), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("session should build") + } + + #[test] + fn hard_delete_anonymizes_every_sqlite_history_snapshot() { + assert_eq!( + SQLITE_ANONYMIZE_USER_HISTORY_SQL.len(), + USER_HISTORY_TABLES.len() + ); + for table in USER_HISTORY_TABLES { + let statement = SQLITE_ANONYMIZE_USER_HISTORY_SQL + .iter() + .find(|sql| sql.starts_with(&format!("UPDATE {table} "))) + .unwrap_or_else(|| panic!("missing history anonymization for {table}")); + assert!(statement.contains("username = NULL")); + assert!(statement.ends_with("WHERE user_id = ?")); + } + for table in ["request_candidates", "video_tasks", "usage"] { + let statement = SQLITE_ANONYMIZE_USER_HISTORY_SQL + .iter() + .find(|sql| sql.starts_with(&format!("UPDATE {table} "))) + .expect("identity snapshot table should be covered"); + assert!(statement.contains("api_key_name = NULL")); + } + assert!(SQLITE_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .starts_with("UPDATE stats_daily_api_key SET api_key_name = NULL")); + assert!(SQLITE_ANONYMIZE_USER_API_KEY_HISTORY_SQL + .contains("SELECT id FROM api_keys WHERE user_id = ?")); + } + + #[tokio::test] + async fn sqlite_hard_delete_preserves_history_ids_and_anonymizes_snapshots() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users ( + id, username, password_hash, role, auth_source, + is_active, is_deleted, created_at, updated_at +) VALUES ('history-user', 'history-name', 'history-hash', 'user', 'local', 1, 0, 1, 1); + +INSERT INTO request_candidates ( + id, request_id, user_id, api_key_id, username, api_key_name, + candidate_index, status, created_at +) VALUES ( + 'history-row', 'history-request-candidate', 'history-user', 'history-key', + 'history-name', 'history-key-name', 0, 'success', 1 +); + +INSERT INTO video_tasks ( + id, request_id, user_id, api_key_id, username, api_key_name, created_at, updated_at +) VALUES ( + 'history-row', 'history-video-request', 'history-user', 'history-key', + 'history-name', 'history-key-name', 1, 1 +); + +INSERT INTO usage ( + request_id, id, user_id, api_key_id, username, api_key_name +) VALUES ( + 'history-usage-request', 'history-row', 'history-user', 'history-key', + 'history-name', 'history-key-name' +); + +INSERT INTO stats_user_daily (id, user_id, date, username, created_at, updated_at) +VALUES ('history-row', 'history-user', 1, 'history-name', 1, 1); +INSERT INTO stats_user_summary (id, user_id, username, cutoff_date, created_at, updated_at) +VALUES ('history-row', 'history-user', 'history-name', 1, 1, 1); +INSERT INTO stats_user_daily_model (id, user_id, username, date, model, created_at, updated_at) +VALUES ('history-row', 'history-user', 'history-name', 1, 'history-model', 1, 1); +INSERT INTO stats_user_daily_provider ( + id, user_id, username, date, provider_name, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-name', 1, 'history-provider', 1, 1); +INSERT INTO stats_user_daily_api_format ( + id, user_id, username, date, api_format, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-name', 1, 'history-format', 1, 1); +INSERT INTO stats_user_daily_model_provider ( + id, user_id, username, date, model, provider_name, created_at, updated_at +) VALUES ( + 'history-row', 'history-user', 'history-name', 1, + 'history-model', 'history-provider', 1, 1 +); +INSERT INTO stats_user_daily_cost_savings ( + id, user_id, username, date, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-name', 1, 1, 1); +INSERT INTO stats_user_daily_cost_savings_provider ( + id, user_id, username, date, provider_name, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-name', 1, 'history-provider', 1, 1); +INSERT INTO stats_user_daily_cost_savings_model ( + id, user_id, username, date, model, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-name', 1, 'history-model', 1, 1); +INSERT INTO stats_user_daily_cost_savings_model_provider ( + id, user_id, username, date, model, provider_name, created_at, updated_at +) VALUES ( + 'history-row', 'history-user', 'history-name', 1, + 'history-model', 'history-provider', 1, 1 +); +INSERT INTO stats_hourly_user ( + id, hour_utc, user_id, created_at, updated_at +) VALUES ('history-row', 1, 'history-user', 1, 1); +INSERT INTO stats_hourly_user_model ( + id, hour_utc, user_id, model, created_at, updated_at +) VALUES ('history-row', 1, 'history-user', 'history-model', 1, 1); +INSERT INTO user_model_usage_counts ( + id, user_id, model, created_at, updated_at +) VALUES ('history-row', 'history-user', 'history-model', 1, 1); +INSERT INTO stats_daily_api_key ( + id, api_key_id, date, api_key_name, created_at, updated_at +) VALUES ('history-row', 'history-key', 1, 'history-key-name', 1, 1); +INSERT INTO api_keys ( + id, user_id, name, key_hash, created_at, updated_at +) VALUES ('history-key', 'history-user', 'history-key-name', 'history-key-hash', 1, 1); + +INSERT INTO users ( + id, username, password_hash, role, auth_source, + is_active, is_deleted, created_at, updated_at +) VALUES ('history-inviter', 'history-inviter', 'history-hash', 'user', 'local', 1, 0, 1, 1); + +INSERT INTO wallets ( + id, user_id, balance, gift_balance, status, created_at, updated_at +) VALUES ('history-wallet', 'history-user', 12, 3, 'active', 1, 1); + +INSERT INTO wallet_transactions ( + id, wallet_id, category, reason_code, amount, + balance_before, balance_after, + recharge_balance_before, recharge_balance_after, + gift_balance_before, gift_balance_after, + operator_id, description, created_at +) VALUES ( + 'history-wallet-tx', 'history-wallet', 'adjust', 'manual', 1, + 14, 15, 11, 12, 3, 3, + 'history-user', 'private wallet note', 1 +); + +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, payment_method, + gateway_response, status, created_at +) VALUES ( + 'history-order', 'history-order-no', 'history-wallet', 'history-user', 12, + 'test', '{"customer_email":"history@example.com"}', 'credited', 1 +); + +INSERT INTO payment_callbacks ( + id, payment_order_id, payment_method, callback_key, order_no, + payload_hash, signature_valid, status, payload, error_message, created_at +) VALUES ( + 'history-callback', 'history-order', 'test', 'history-callback-key', + 'history-order-no', 'history-payload-hash', 1, 'processed', + '{"customer_email":"history@example.com"}', 'private callback error', 1 +); + +INSERT INTO billing_plans ( + id, title, price_amount, duration_unit, duration_value, + entitlements_json, created_at, updated_at +) VALUES ('history-plan', 'history plan', 12, 'day', 30, '[]', 1, 1); + +INSERT INTO user_plan_entitlements ( + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +) VALUES ( + 'history-entitlement', 'history-user', 'history-plan', 'history-order', + 'active', 1, 4102444800, '[]', 1, 1 +); + +INSERT INTO entitlement_usage_ledgers ( + id, user_entitlement_id, user_id, request_id, amount_usd, + balance_before, balance_after, usage_date, created_at +) VALUES ( + 'history-entitlement-ledger', 'history-entitlement', 'history-user', + 'history-entitlement-request', 1, 12, 11, '2026-08-27', 1 +); + +INSERT INTO user_referrals ( + id, inviter_user_id, invitee_user_id, invite_code_snapshot, source_json, + first_paid_order_id, first_paid_at, created_at, updated_at +) VALUES ( + 'history-referral', 'history-inviter', 'history-user', 'PRIVATE-CODE', + '{"ip":"192.0.2.1"}', 'history-order', 1, 1, 1 +); + +INSERT INTO referral_rewards ( + id, referral_id, inviter_user_id, invitee_user_id, reward_type, + trigger_point, source_order_id, idempotency_key, amount_usd, status, + failure_reason, admin_operator_id, admin_note, created_at, updated_at +) VALUES ( + 'history-reward', 'history-referral', 'history-inviter', 'history-user', + 'percent', 'paid_order', 'history-order', 'history-reward-key', 1, 'failed', + 'private failure', 'history-user', 'private admin note', 1, 1 +); + +INSERT INTO audit_logs ( + id, event_type, user_id, description, ip_address, user_agent, + event_metadata, error_message, created_at +) VALUES ( + 'history-audit', 'history_event', 'history-user', 'private description', + '192.0.2.2', 'private agent', '{"private":true}', 'private error', 1 +); +"#, + ) + .execute(&pool) + .await + .expect("history fixtures should insert"); + + let repository = SqliteUserReadRepository::new(pool.clone()); + assert!(repository + .delete_local_auth_user("history-user") + .await + .expect("hard delete should succeed")); + + for table in USER_HISTORY_TABLES { + let row = sqlx::query(&format!( + "SELECT user_id, username FROM {table} WHERE id = 'history-row'" + )) + .fetch_one(&pool) + .await + .unwrap_or_else(|error| panic!("{table} history should remain: {error}")); + assert_eq!( + row.try_get::, _>("user_id") + .expect("user_id should decode") + .as_deref(), + Some("history-user"), + "{table} user_id must remain stable" + ); + assert_eq!( + row.try_get::, _>("username") + .expect("username should decode"), + None, + "{table} username snapshot must be removed" + ); + } + for table in ["request_candidates", "video_tasks", "usage"] { + let row = sqlx::query(&format!( + "SELECT api_key_id, api_key_name FROM {table} WHERE id = 'history-row'" + )) + .fetch_one(&pool) + .await + .unwrap_or_else(|error| panic!("{table} identity snapshot should remain: {error}")); + assert_eq!( + row.try_get::, _>("api_key_id") + .expect("api_key_id should decode") + .as_deref(), + Some("history-key"), + "{table} api_key_id must remain stable" + ); + assert_eq!( + row.try_get::, _>("api_key_name") + .expect("api_key_name should decode"), + None, + "{table} API key name snapshot must be removed" + ); + } + for table in STABLE_USER_ID_ONLY_TABLES { + let user_id: String = sqlx::query_scalar(&format!( + "SELECT user_id FROM {table} WHERE id = 'history-row'" + )) + .fetch_one(&pool) + .await + .unwrap_or_else(|error| panic!("{table} fact should remain: {error}")); + assert_eq!( + user_id, "history-user", + "{table} user_id must remain stable" + ); + } + let api_key_fact = sqlx::query( + "SELECT api_key_id, api_key_name FROM stats_daily_api_key WHERE id = 'history-row'", + ) + .fetch_one(&pool) + .await + .expect("API key aggregate should remain"); + assert_eq!( + api_key_fact + .try_get::("api_key_id") + .expect("aggregate api_key_id should decode"), + "history-key" + ); + assert_eq!( + api_key_fact + .try_get::, _>("api_key_name") + .expect("aggregate api_key_name should decode"), + None + ); + let wallet_fact = + sqlx::query("SELECT user_id, status FROM wallets WHERE id = 'history-wallet'") + .fetch_one(&pool) + .await + .expect("wallet fact should remain"); + assert_eq!( + wallet_fact + .try_get::, _>("user_id") + .expect("wallet user_id should decode") + .as_deref(), + Some("history-user") + ); + assert_eq!( + wallet_fact + .try_get::("status") + .expect("wallet status should decode"), + "disabled" + ); + let order_fact = sqlx::query( + "SELECT user_id, gateway_response FROM payment_orders WHERE id = 'history-order'", + ) + .fetch_one(&pool) + .await + .expect("payment order fact should remain"); + assert_eq!( + order_fact + .try_get::, _>("user_id") + .expect("order user_id should decode") + .as_deref(), + Some("history-user") + ); + assert_eq!( + order_fact + .try_get::, _>("gateway_response") + .expect("gateway response should decode"), + None + ); + let callback_fact = sqlx::query( + "SELECT payment_order_id, order_no, payload, error_message FROM payment_callbacks WHERE id = 'history-callback'", + ) + .fetch_one(&pool) + .await + .expect("payment callback fact should remain"); + assert_eq!( + callback_fact + .try_get::, _>("payment_order_id") + .expect("callback order id should decode") + .as_deref(), + Some("history-order") + ); + assert_eq!( + callback_fact + .try_get::, _>("order_no") + .expect("callback order number should decode") + .as_deref(), + Some("history-order-no") + ); + assert_eq!( + callback_fact + .try_get::, _>("payload") + .expect("callback payload should decode"), + None + ); + assert_eq!( + callback_fact + .try_get::, _>("error_message") + .expect("callback error should decode"), + None + ); + let entitlement_fact = sqlx::query( + "SELECT user_id, status FROM user_plan_entitlements WHERE id = 'history-entitlement'", + ) + .fetch_one(&pool) + .await + .expect("entitlement fact should remain"); + assert_eq!( + entitlement_fact + .try_get::("user_id") + .expect("entitlement user_id should decode"), + "history-user" + ); + assert_eq!( + entitlement_fact + .try_get::("status") + .expect("entitlement status should decode"), + "revoked" + ); + let ledger_user_id: String = sqlx::query_scalar( + "SELECT user_id FROM entitlement_usage_ledgers WHERE id = 'history-entitlement-ledger'", + ) + .fetch_one(&pool) + .await + .expect("entitlement ledger should remain"); + assert_eq!(ledger_user_id, "history-user"); + let referral_fact = sqlx::query( + "SELECT invitee_user_id, invite_code_snapshot, source_json FROM user_referrals WHERE id = 'history-referral'", + ) + .fetch_one(&pool) + .await + .expect("referral fact should remain"); + assert_eq!( + referral_fact + .try_get::("invitee_user_id") + .expect("invitee user_id should decode"), + "history-user" + ); + assert_eq!( + referral_fact + .try_get::("invite_code_snapshot") + .expect("invite code snapshot should decode"), + "deleted-user" + ); + assert_eq!( + referral_fact + .try_get::, _>("source_json") + .expect("referral source should decode"), + None + ); + let reward_fact = sqlx::query( + "SELECT invitee_user_id, status, failure_reason, admin_note FROM referral_rewards WHERE id = 'history-reward'", + ) + .fetch_one(&pool) + .await + .expect("referral reward fact should remain"); + assert_eq!( + reward_fact + .try_get::("invitee_user_id") + .expect("reward invitee should decode"), + "history-user" + ); + assert_eq!( + reward_fact + .try_get::("status") + .expect("reward status should decode"), + "voided" + ); + assert_eq!( + reward_fact + .try_get::, _>("failure_reason") + .expect("reward failure reason should decode"), + None + ); + assert_eq!( + reward_fact + .try_get::, _>("admin_note") + .expect("reward admin note should decode"), + None + ); + let audit_fact = sqlx::query( + "SELECT user_id, ip_address, user_agent, event_metadata, error_message FROM audit_logs WHERE id = 'history-audit'", + ) + .fetch_one(&pool) + .await + .expect("audit fact should remain"); + assert_eq!( + audit_fact + .try_get::, _>("user_id") + .expect("audit user_id should decode") + .as_deref(), + Some("history-user") + ); + for column in [ + "ip_address", + "user_agent", + "event_metadata", + "error_message", + ] { + assert_eq!( + audit_fact + .try_get::, _>(column) + .unwrap_or_else(|error| panic!("audit {column} should decode: {error}")), + None, + "audit {column} must be removed" + ); + } + let transaction_description: Option = sqlx::query_scalar( + "SELECT description FROM wallet_transactions WHERE id = 'history-wallet-tx'", + ) + .fetch_one(&pool) + .await + .expect("wallet transaction should remain"); + assert_eq!(transaction_description, None); + let user_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM users WHERE id = 'history-user'") + .fetch_one(&pool) + .await + .expect("deleted user count should load"); + assert_eq!(user_count, 0); + } + + #[tokio::test] + async fn sqlite_atomic_user_delete_requires_wallet_absence() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + "INSERT INTO users (id, username, email, role, auth_source, is_active, is_deleted, created_at, updated_at) VALUES ('atomic-user', 'atomic-user', 'atomic@example.com', 'user', 'local', 1, 0, 1, 1)", + ) + .execute(&pool) + .await + .expect("user should seed"); + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, created_at, updated_at) VALUES ('atomic-wallet', 'atomic-user', 0, 0, 'active', 1, 1)", + ) + .execute(&pool) + .await + .expect("wallet should seed"); + + let repository = SqliteUserReadRepository::new(pool.clone()); + assert!(!repository + .delete_local_auth_user_if_wallet_absent("atomic-user") + .await + .expect("wallet guard should resolve")); + assert!(repository + .find_user_auth_by_id("atomic-user") + .await + .expect("user lookup should succeed") + .is_some()); + + sqlx::query("DELETE FROM wallets WHERE id = 'atomic-wallet'") + .execute(&pool) + .await + .expect("wallet should remove"); + assert!(repository + .delete_local_auth_user_if_wallet_absent("atomic-user") + .await + .expect("wallet-free user should delete")); + assert!(repository + .find_user_auth_by_id("atomic-user") + .await + .expect("user lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn sqlite_atomic_user_delete_detects_api_key_wallet() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + "INSERT INTO users (id, username, email, role, auth_source, is_active, is_deleted, created_at, updated_at) VALUES ('api-wallet-user', 'api-wallet-user', 'api-wallet@example.com', 'user', 'local', 1, 0, 1, 1)", + ) + .execute(&pool) + .await + .expect("user should seed"); + sqlx::query( + "INSERT INTO api_keys (id, user_id, key_hash, created_at, updated_at) VALUES ('api-wallet-key', 'api-wallet-user', 'api-wallet-key-hash', 1, 1)", + ) + .execute(&pool) + .await + .expect("api key should seed"); + sqlx::query( + "INSERT INTO wallets (id, api_key_id, balance, gift_balance, status, created_at, updated_at) VALUES ('api-wallet', 'api-wallet-key', 25, 0, 'active', 1, 1)", + ) + .execute(&pool) + .await + .expect("api key wallet should seed"); + + let repository = SqliteUserReadRepository::new(pool.clone()); + assert!(!repository + .delete_local_auth_user_if_wallet_absent("api-wallet-user") + .await + .expect("api key wallet guard should resolve")); + + assert!(repository + .find_user_auth_by_id("api-wallet-user") + .await + .expect("user lookup should succeed") + .is_some()); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM api_keys WHERE id = 'api-wallet-key'", + ) + .fetch_one(&pool) + .await + .expect("api key count should query"), + 1 + ); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallets WHERE id = 'api-wallet' AND api_key_id = 'api-wallet-key'", + ) + .fetch_one(&pool) + .await + .expect("wallet count should query"), + 1 + ); + + sqlx::query("DELETE FROM wallets WHERE id = 'api-wallet'") + .execute(&pool) + .await + .expect("api key wallet should remove"); + assert!(repository + .delete_local_auth_user_if_wallet_absent("api-wallet-user") + .await + .expect("wallet-free api key owner should delete")); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM api_keys WHERE id = 'api-wallet-key'", + ) + .fetch_one(&pool) + .await + .expect("api key count should query after delete"), + 0 + ); + } + + #[tokio::test] + async fn sqlite_atomically_preserves_last_active_admin_and_revokes_on_security_change() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_sqlite_admin(&pool, "admin-1", "admin_one").await; + sqlx::query( + "INSERT INTO management_tokens (id, user_id, name, token_hash, created_at, updated_at) VALUES ('token-admin-1', 'admin-1', 'admin token', 'token-hash-admin-1', 1, 1)", + ) + .execute(&pool) + .await + .expect("management token should insert"); + let repository = SqliteUserReadRepository::new(pool.clone()); + + let update_error = repository + .update_local_auth_user_admin_fields( + "admin-1", + Some("audit_admin".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + None, + ) + .await + .expect_err("last active admin demotion must be rejected"); + assert!(is_last_active_admin_update_denied(&update_error)); + assert_eq!( + repository + .find_user_auth_by_id("admin-1") + .await + .expect("admin lookup should succeed") + .expect("admin should remain") + .role, + "admin" + ); + + let delete_error = repository + .delete_local_auth_user("admin-1") + .await + .expect_err("last active admin delete must be rejected"); + assert!(is_last_active_admin_delete_denied(&delete_error)); + let token_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM management_tokens WHERE user_id = 'admin-1'") + .fetch_one(&pool) + .await + .expect("management token count should load"); + assert_eq!( + token_count, 1, + "rejected delete must roll back credential cleanup" + ); + + seed_sqlite_admin(&pool, "admin-2", "admin_two").await; + let now = chrono::Utc::now(); + let session = test_user_session("session-admin-1", "admin-1", "device-1", "refresh", now); + repository + .create_user_session(&session) + .await + .expect("admin session should create") + .expect("admin session should exist"); + sqlx::query( + "INSERT INTO api_keys (id, user_id, key_hash, created_at, updated_at) VALUES ('key-admin-1', 'admin-1', 'key-hash-admin-1', 1, 1)", + ) + .execute(&pool) + .await + .expect("api key should insert"); + sqlx::query( + "INSERT INTO api_key_provider_mappings (id, api_key_id, provider_id, created_at, updated_at) VALUES ('mapping-admin-1', 'key-admin-1', 'provider-1', 1, 1)", + ) + .execute(&pool) + .await + .expect("api key mapping should insert"); + sqlx::query( + "INSERT INTO user_oauth_links (id, user_id, provider_type, provider_user_id, linked_at) VALUES ('oauth-admin-1', 'admin-1', 'test', 'subject-admin-1', 1)", + ) + .execute(&pool) + .await + .expect("oauth link should insert"); + sqlx::query( + "INSERT INTO user_group_members (group_id, user_id, created_at) VALUES ('00000000-0000-0000-0000-000000000001', 'admin-1', 1)", + ) + .execute(&pool) + .await + .expect("group membership should insert"); + sqlx::query( + "INSERT INTO user_preferences (id, user_id, created_at, updated_at) VALUES ('preferences-admin-1', 'admin-1', 1, 1)", + ) + .execute(&pool) + .await + .expect("preferences should insert"); + sqlx::query( + "INSERT INTO announcements (id, title, content, created_at, updated_at) VALUES ('announcement-1', 'notice', 'content', 1, 1)", + ) + .execute(&pool) + .await + .expect("announcement should insert"); + sqlx::query( + "INSERT INTO announcement_reads (id, user_id, announcement_id, read_at) VALUES ('read-admin-1', 'admin-1', 'announcement-1', 1)", + ) + .execute(&pool) + .await + .expect("announcement read should insert"); + let updated = repository + .update_local_auth_user_admin_fields( + "admin-1", + Some("audit_admin".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + None, + ) + .await + .expect("demotion with another active admin should succeed") + .expect("admin should exist"); + assert_eq!(updated.role, "audit_admin"); + let revoked = repository + .find_user_session("admin-1", "session-admin-1") + .await + .expect("session lookup should succeed") + .expect("session should remain as audit record"); + assert!(revoked.revoked_at.is_some()); + assert_eq!( + revoked.revoke_reason.as_deref(), + Some("user_security_state_changed") + ); + assert!(repository + .delete_local_auth_user("admin-1") + .await + .expect("non-full-admin delete should succeed")); + for (table, predicate) in [ + ("api_key_provider_mappings", "api_key_id = 'key-admin-1'"), + ("api_keys", "user_id = 'admin-1'"), + ("management_tokens", "user_id = 'admin-1'"), + ("user_sessions", "user_id = 'admin-1'"), + ("user_oauth_links", "user_id = 'admin-1'"), + ("user_group_members", "user_id = 'admin-1'"), + ("user_preferences", "user_id = 'admin-1'"), + ("announcement_reads", "user_id = 'admin-1'"), + ] { + let count: i64 = + sqlx::query_scalar(&format!("SELECT COUNT(*) FROM {table} WHERE {predicate}")) + .fetch_one(&pool) + .await + .expect("dependent row count should load"); + assert_eq!(count, 0, "{table} credentials must be removed"); + } + } + + async fn seed_active_session_test_user(pool: &crate::SqlitePool, user_id: &str) { + sqlx::query( + r#" +INSERT INTO users ( + id, email, email_verified, username, password_hash, role, auth_source, + is_active, is_deleted, created_at, updated_at +) VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 1, 0, ?, ?) +"#, + ) + .bind(user_id) + .bind(format!("{user_id}@example.com")) + .bind(user_id) + .bind(chrono::Utc::now().timestamp()) + .bind(chrono::Utc::now().timestamp()) + .execute(pool) + .await + .expect("session test user should insert"); + } #[tokio::test] async fn sqlite_repository_reads_user_contract_views() { @@ -2168,7 +4154,7 @@ INSERT INTO users ( .execute(&pool) .await .expect("seed users should insert"); - let valid_hash = format!("$2b$12${}", "a".repeat(53)); + let valid_hash = "$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string(); sqlx::query( r#" INSERT INTO users ( @@ -2186,7 +4172,7 @@ INSERT INTO users ( .await .expect("valid local admin should insert"); - let repository = SqliteUserReadRepository::new(pool); + let repository = SqliteUserReadRepository::new(pool.clone()); let summaries = repository .list_users_by_ids(&["user-1".to_string(), "admin-1".to_string()]) .await @@ -2258,7 +4244,9 @@ INSERT INTO users ( let profile_updated = repository .update_local_auth_user_profile( "user-1", + true, Some("user-1b@example.com".to_string()), + Some(true), Some("alice-b".to_string()), ) .await @@ -2268,6 +4256,7 @@ INSERT INTO users ( profile_updated.email.as_deref(), Some("user-1b@example.com") ); + assert!(profile_updated.email_verified); assert_eq!(profile_updated.username, "alice-b"); let password_updated = repository .update_local_auth_user_password_hash( @@ -2366,6 +4355,13 @@ INSERT INTO users ( .await .expect("email lookup should load") .is_none()); + let cleared_profile = repository + .update_local_auth_user_profile("user-1", true, None, Some(false), None) + .await + .expect("nullable email update should succeed") + .expect("profile should remain"); + assert!(cleared_profile.email.is_none()); + assert!(!cleared_profile.email_verified); assert_eq!( repository .count_active_admin_users() @@ -2425,7 +4421,9 @@ INSERT INTO users ( Some(now), Some(now), ) - .expect("session should build"); + .expect("session should build") + .with_security_version(1) + .expect("session security version should be valid"); assert_eq!( repository .create_user_session(&session) @@ -2479,6 +4477,318 @@ INSERT INTO users ( .is_none()); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn non_password_session_replacement_is_atomic_under_concurrency() { + let database_path = std::env::temp_dir().join(format!( + "aether-sqlite-session-replacement-race-{}.db", + uuid::Uuid::new_v4() + )); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .foreign_keys(true) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .busy_timeout(Duration::from_secs(30)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_active_session_test_user(&pool, "session-race-user").await; + + let repository = SqliteUserReadRepository::new(pool.clone()); + let now = chrono::Utc::now(); + let first = test_user_session( + "session-race-first", + "session-race-user", + "shared-device", + "refresh-first", + now, + ); + let second = test_user_session( + "session-race-second", + "session-race-user", + "shared-device", + "refresh-second", + now, + ); + let barrier = Arc::new(tokio::sync::Barrier::new(2)); + let first_repository = repository.clone(); + let first_barrier = Arc::clone(&barrier); + let first_create = tokio::spawn(async move { + first_barrier.wait().await; + first_repository.create_user_session(&first).await + }); + let second_repository = repository.clone(); + let second_barrier = Arc::clone(&barrier); + let second_create = tokio::spawn(async move { + second_barrier.wait().await; + second_repository.create_user_session(&second).await + }); + + assert!(first_create + .await + .expect("first login task should join") + .expect("first login should succeed") + .is_some()); + assert!(second_create + .await + .expect("second login task should join") + .expect("second login should succeed") + .is_some()); + let active = repository + .list_user_sessions("session-race-user") + .await + .expect("active sessions should list"); + assert_eq!(active.len(), 1); + assert_eq!(active[0].client_device_id, "shared-device"); + + pool.close().await; + let _ = std::fs::remove_file(database_path); + } + + #[tokio::test] + async fn failed_non_password_session_insert_rolls_back_device_revocation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_active_session_test_user(&pool, "session-rollback-user").await; + let repository = SqliteUserReadRepository::new(pool); + let now = chrono::Utc::now(); + let current = test_user_session( + "session-duplicate-id", + "session-rollback-user", + "shared-device", + "refresh-current", + now, + ); + repository + .create_user_session(¤t) + .await + .expect("initial session should create") + .expect("initial session should exist"); + + let duplicate = test_user_session( + "session-duplicate-id", + "session-rollback-user", + "shared-device", + "refresh-duplicate", + now + chrono::Duration::seconds(1), + ); + assert!(repository.create_user_session(&duplicate).await.is_err()); + + let active = repository + .list_user_sessions("session-rollback-user") + .await + .expect("active sessions should list"); + assert_eq!(active.len(), 1); + assert_eq!(active[0].refresh_token_hash, current.refresh_token_hash); + assert!(!active[0].is_revoked()); + } + + #[tokio::test] + async fn security_state_changes_revoke_sessions_without_reactivation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_active_session_test_user(&pool, "sqlite-security-state-user").await; + let repository = SqliteUserReadRepository::new(pool); + let now = chrono::Utc::now(); + let session = test_user_session( + "sqlite-security-state-session", + "sqlite-security-state-user", + "sqlite-security-state-device", + "sqlite-security-state-refresh", + now, + ); + repository + .create_user_session(&session) + .await + .expect("initial session should create") + .expect("initial session should exist"); + + repository + .update_local_auth_user_admin_fields( + "sqlite-security-state-user", + None, + false, + None, + false, + None, + false, + None, + false, + None, + Some(false), + ) + .await + .expect("disable should succeed") + .expect("user should exist"); + let revoked = repository + .find_user_session( + "sqlite-security-state-user", + "sqlite-security-state-session", + ) + .await + .expect("revoked session should load") + .expect("revoked session should remain stored"); + assert!(revoked.is_revoked()); + assert_eq!( + revoked.revoke_reason.as_deref(), + Some("user_security_state_changed") + ); + + repository + .update_local_auth_user_admin_fields( + "sqlite-security-state-user", + None, + false, + None, + false, + None, + false, + None, + false, + None, + Some(true), + ) + .await + .expect("reactivation should succeed") + .expect("user should exist"); + assert!(repository + .list_user_sessions("sqlite-security-state-user") + .await + .expect("sessions should list") + .is_empty()); + + let replacement = test_user_session( + "sqlite-security-state-replacement", + "sqlite-security-state-user", + "sqlite-security-state-device", + "sqlite-security-state-replacement-refresh", + now + chrono::Duration::seconds(1), + ); + let replacement = replacement + .with_security_version( + repository + .find_user_auth_by_id("sqlite-security-state-user") + .await + .expect("user lookup should succeed") + .expect("user should exist") + .security_version, + ) + .expect("security version should be valid"); + let created = repository + .create_user_session(&replacement) + .await + .expect("replacement session should create") + .expect("replacement session should exist"); + assert_eq!(created.security_version, replacement.security_version); + let persisted = repository + .find_user_session( + "sqlite-security-state-user", + "sqlite-security-state-replacement", + ) + .await + .expect("replacement session should load") + .expect("replacement session should remain stored"); + assert_eq!(persisted.security_version, replacement.security_version); + let active = repository + .list_user_sessions("sqlite-security-state-user") + .await + .expect("replacement session should list"); + assert_eq!(active.len(), 1); + assert_eq!(active[0].security_version, replacement.security_version); + repository + .update_local_auth_user_admin_fields( + "sqlite-security-state-user", + Some("audit_admin".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + None, + ) + .await + .expect("role update should succeed") + .expect("user should exist"); + assert!(repository + .list_user_sessions("sqlite-security-state-user") + .await + .expect("sessions should list") + .is_empty()); + } + + #[tokio::test] + async fn unchanged_security_state_preserves_sqlite_sessions() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_active_session_test_user(&pool, "sqlite-unchanged-security-user").await; + let repository = SqliteUserReadRepository::new(pool); + let now = chrono::Utc::now(); + let session = test_user_session( + "sqlite-unchanged-security-session", + "sqlite-unchanged-security-user", + "sqlite-unchanged-security-device", + "sqlite-unchanged-security-refresh", + now, + ); + repository + .create_user_session(&session) + .await + .expect("initial session should create") + .expect("initial session should exist"); + + repository + .update_local_auth_user_admin_fields( + "sqlite-unchanged-security-user", + Some("USER".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + Some(true), + ) + .await + .expect("idempotent security update should succeed") + .expect("user should exist"); + assert_eq!( + repository + .list_user_sessions("sqlite-unchanged-security-user") + .await + .expect("sessions should list") + .len(), + 1 + ); + } + #[tokio::test] async fn sqlite_repository_manages_oauth_users_and_links() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -2494,7 +4804,9 @@ INSERT INTO users ( INSERT INTO oauth_providers ( provider_type, display_name, client_id, redirect_uri, frontend_callback_url, is_enabled, created_at, updated_at -) VALUES ('linuxdo', 'Linux.do', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1) +) VALUES + ('linuxdo', 'Linux.do', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1), + ('github', 'GitHub', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1) "#, ) .execute(&pool) @@ -2506,6 +4818,7 @@ INSERT INTO oauth_providers ( let user = repository .create_oauth_auth_user( Some("OAuth@Example.com".to_string()), + false, "oauth_user".to_string(), now, ) @@ -2513,6 +4826,7 @@ INSERT INTO oauth_providers ( .expect("oauth user should create") .expect("oauth user should exist"); assert_eq!(user.auth_source, "oauth"); + assert!(!user.email_verified); assert_eq!( repository .find_active_user_auth_by_email_ci("oauth@example.com") @@ -2522,18 +4836,38 @@ INSERT INTO oauth_providers ( Some(user.id.clone()) ); - repository - .upsert_user_oauth_link( - &user.id, - "linuxdo", - "subject-1", - Some("alice"), - Some("alice@example.com"), - Some(serde_json::json!({"sub": "subject-1"})), - now, - ) + assert!(!repository + .upgrade_oauth_email_verification_if_matches(&user.id, "different@example.com", now,) .await - .expect("oauth link should upsert"); + .expect("mismatched verification should resolve")); + assert!(repository + .upgrade_oauth_email_verification_if_matches(&user.id, "oauth@example.com", now) + .await + .expect("matching verification should resolve")); + assert!( + repository + .find_user_auth_by_id(&user.id) + .await + .expect("user should load") + .expect("user should exist") + .email_verified + ); + + assert_eq!( + repository + .bind_user_oauth_link( + &user.id, + "linuxdo", + "subject-1", + Some("alice"), + Some("alice@example.com"), + Some(serde_json::json!({"sub": "subject-1"})), + now, + ) + .await + .expect("oauth link should bind"), + BindUserOAuthLinkOutcome::Bound + ); assert_eq!( repository .find_oauth_link_owner("linuxdo", "subject-1") @@ -2573,16 +4907,430 @@ INSERT INTO oauth_providers ( .expect("link count should load"), 1 ); - assert!(repository - .delete_user_oauth_link(&user.id, "linuxdo") + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &[]) + .await + .expect("last link deletion should resolve"), + DeleteUserOAuthLinkOutcome::LastOAuthBinding + ); + repository + .bind_user_oauth_link( + &user.id, + "github", + "subject-2", + Some("alice"), + Some("alice@example.com"), + None, + now, + ) .await - .expect("link should delete")); + .expect("second link should upsert"); + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &[]) + .await + .expect("link should delete"), + DeleteUserOAuthLinkOutcome::Deleted + ); assert_eq!( repository .count_user_oauth_links(&user.id) .await .expect("link count should load"), - 0 + 1 ); } + + #[tokio::test] + async fn concurrent_sqlite_oauth_unbinds_preserve_one_login_method() { + let database_path = std::env::temp_dir().join(format!( + "aether-oauth-unbind-{}.sqlite", + uuid::Uuid::new_v4() + )); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .foreign_keys(true) + .busy_timeout(Duration::from_secs(5)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES + ('linuxdo', 'Linux.do', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1), + ('github', 'GitHub', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1) +"#, + ) + .execute(&pool) + .await + .expect("providers should insert"); + + let repository = Arc::new(SqliteUserReadRepository::new(pool.clone())); + let now = chrono::Utc::now(); + let user = repository + .create_oauth_auth_user( + Some("concurrent-oauth@example.com".to_string()), + true, + "concurrent-oauth".to_string(), + now, + ) + .await + .expect("oauth user should create") + .expect("oauth user should exist"); + for (provider_type, subject) in [("linuxdo", "subject-1"), ("github", "subject-2")] { + repository + .bind_user_oauth_link(&user.id, provider_type, subject, None, None, None, now) + .await + .expect("oauth link should upsert"); + } + + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let first_repository = Arc::clone(&repository); + let first_barrier = Arc::clone(&barrier); + let first_user_id = user.id.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository + .delete_user_oauth_link(&first_user_id, "linuxdo", false, &[]) + .await + .expect("first unlink should resolve") + }); + let second_repository = Arc::clone(&repository); + let second_barrier = Arc::clone(&barrier); + let second_user_id = user.id.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .delete_user_oauth_link(&second_user_id, "github", false, &[]) + .await + .expect("second unlink should resolve") + }); + barrier.wait().await; + let outcomes = [ + first.await.expect("first unlink task should join"), + second.await.expect("second unlink task should join"), + ]; + + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == DeleteUserOAuthLinkOutcome::Deleted) + .count(), + 1 + ); + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == DeleteUserOAuthLinkOutcome::LastOAuthBinding) + .count(), + 1 + ); + assert_eq!( + repository + .count_user_oauth_links(&user.id) + .await + .expect("remaining links should count"), + 1 + ); + + drop(repository); + pool.close().await; + let _ = std::fs::remove_file(database_path); + } + + #[tokio::test] + async fn concurrent_sqlite_oauth_binds_preserve_single_identity_owner() { + let database_path = + std::env::temp_dir().join(format!("aether-oauth-bind-{}.sqlite", uuid::Uuid::new_v4())); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .foreign_keys(true) + .busy_timeout(Duration::from_secs(5)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://example.test/callback', + 'https://example.test/app', 1, 1, 1 +) +"#, + ) + .execute(&pool) + .await + .expect("provider should insert"); + + let repository = Arc::new(SqliteUserReadRepository::new(pool.clone())); + let now = chrono::Utc::now(); + let first_user = repository + .create_oauth_auth_user(None, false, "bind-first".to_string(), now) + .await + .expect("first user should create") + .expect("first user should exist"); + let second_user = repository + .create_oauth_auth_user(None, false, "bind-second".to_string(), now) + .await + .expect("second user should create") + .expect("second user should exist"); + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + + let first_repository = Arc::clone(&repository); + let first_barrier = Arc::clone(&barrier); + let first_id = first_user.id.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository + .bind_user_oauth_link( + &first_id, + "linuxdo", + "shared-subject", + None, + None, + None, + now, + ) + .await + .expect("first bind should resolve") + }); + let second_repository = Arc::clone(&repository); + let second_barrier = Arc::clone(&barrier); + let second_id = second_user.id.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .bind_user_oauth_link( + &second_id, + "linuxdo", + "shared-subject", + None, + None, + None, + now, + ) + .await + .expect("second bind should resolve") + }); + barrier.wait().await; + let outcomes = [ + first.await.expect("first bind task should join"), + second.await.expect("second bind task should join"), + ]; + + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == BindUserOAuthLinkOutcome::Bound) + .count(), + 1 + ); + assert_eq!( + outcomes + .iter() + .filter(|outcome| { + **outcome == BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }) + .count(), + 1 + ); + let link_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM user_oauth_links WHERE provider_type = 'linuxdo' AND provider_user_id = 'shared-subject'", + ) + .fetch_one(&pool) + .await + .expect("link count should load"); + assert_eq!(link_count, 1); + + drop(repository); + pool.close().await; + let _ = std::fs::remove_file(database_path); + } + + #[tokio::test] + async fn sqlite_oauth_bind_rejects_disabled_provider() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://example.test/callback', + 'https://example.test/app', 0, 1, 1 +) +"#, + ) + .execute(&pool) + .await + .expect("disabled provider should insert"); + let repository = SqliteUserReadRepository::new(pool); + let now = chrono::Utc::now(); + let user = repository + .create_oauth_auth_user(None, false, "disabled-provider-user".to_string(), now) + .await + .expect("user should create") + .expect("user should exist"); + + assert_eq!( + repository + .bind_user_oauth_link(&user.id, "linuxdo", "subject", None, None, None, now) + .await + .expect("bind should resolve"), + BindUserOAuthLinkOutcome::ProviderDisabled + ); + assert!(!repository + .has_user_oauth_provider_link(&user.id, "linuxdo") + .await + .expect("link lookup should succeed")); + } + + #[tokio::test] + async fn sqlite_oauth_unbind_only_counts_enabled_provider_links() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES + ('linuxdo', 'Linux.do', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1), + ('github', 'GitHub', 'client', 'https://example.test/callback', 'https://example.test/app', 1, 1, 1) +"#, + ) + .execute(&pool) + .await + .expect("providers should insert"); + let repository = SqliteUserReadRepository::new(pool.clone()); + let now = chrono::Utc::now(); + let user = repository + .create_oauth_auth_user( + Some("enabled-link@example.com".to_string()), + true, + "enabled-link".to_string(), + now, + ) + .await + .expect("oauth user should create") + .expect("oauth user should exist"); + for (provider_type, subject) in [("linuxdo", "subject-1"), ("github", "subject-2")] { + assert_eq!( + repository + .bind_user_oauth_link(&user.id, provider_type, subject, None, None, None, now) + .await + .expect("oauth link should bind"), + BindUserOAuthLinkOutcome::Bound + ); + } + sqlx::query("UPDATE oauth_providers SET is_enabled = 0 WHERE provider_type = 'github'") + .execute(&pool) + .await + .expect("provider should disable after its link is created"); + + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &[]) + .await + .expect("enabled link deletion should resolve"), + DeleteUserOAuthLinkOutcome::LastOAuthBinding + ); + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "github", false, &[]) + .await + .expect("disabled link deletion should resolve"), + DeleteUserOAuthLinkOutcome::Deleted + ); + assert!(repository + .has_user_oauth_provider_link(&user.id, "linuxdo") + .await + .expect("enabled provider link lookup should work")); + } + + #[tokio::test] + async fn sqlite_oauth_unbind_respects_ldap_exclusive_local_login_policy() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://example.test/callback', + 'https://example.test/app', 1, 1, 1 +) +"#, + ) + .execute(&pool) + .await + .expect("provider should insert"); + let repository = SqliteUserReadRepository::new(pool); + let now = chrono::Utc::now(); + let valid_hash = "$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string(); + let user = repository + .create_local_auth_user( + Some("ldap-exclusive-local@example.com".to_string()), + true, + "ldap-exclusive-local".to_string(), + valid_hash, + ) + .await + .expect("local user should create") + .expect("local user should exist"); + repository + .bind_user_oauth_link(&user.id, "linuxdo", "subject-1", None, None, None, now) + .await + .expect("oauth link should upsert"); + + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &[]) + .await + .expect("unlink should resolve"), + DeleteUserOAuthLinkOutcome::LastLoginMethod + ); + assert!(repository + .has_user_oauth_provider_link(&user.id, "linuxdo") + .await + .expect("oauth link lookup should work")); + } } diff --git a/crates/aether-data/adapters/sqlite/src/video_tasks.rs b/crates/aether-data/adapters/sqlite/src/video_tasks.rs index a4095fc54..a00c0f85d 100644 --- a/crates/aether-data/adapters/sqlite/src/video_tasks.rs +++ b/crates/aether-data/adapters/sqlite/src/video_tasks.rs @@ -72,6 +72,22 @@ impl SqliteVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + async fn find_by_id_for_user( + &self, + id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{VIDEO_TASK_COLUMNS} WHERE id = ? AND user_id = ? LIMIT 1" + )) + .bind(id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + async fn find_by_short_id( &self, short_id: &str, @@ -84,6 +100,22 @@ impl SqliteVideoTaskRepository { row.as_ref().map(map_video_task_row).transpose() } + async fn find_by_short_id_for_user( + &self, + short_id: &str, + user_id: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query(&format!( + "{VIDEO_TASK_COLUMNS} WHERE short_id = ? AND user_id = ? LIMIT 1" + )) + .bind(short_id) + .bind(user_id) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_video_task_row).transpose() + } + async fn find_by_user_external( &self, user_id: &str, @@ -117,6 +149,26 @@ impl VideoTaskReadRepository for SqliteVideoTaskRepository { } } + async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + match key { + VideoTaskLookupKey::Id(id) => self.find_by_id_for_user(id, user_id).await, + VideoTaskLookupKey::ShortId(short_id) => { + self.find_by_short_id_for_user(short_id, user_id).await + } + VideoTaskLookupKey::UserExternal { + user_id: lookup_user_id, + external_task_id, + } if lookup_user_id == user_id => { + self.find_by_user_external(user_id, external_task_id).await + } + VideoTaskLookupKey::UserExternal { .. } => Ok(None), + } + } + async fn list_active(&self, limit: usize) -> Result, DataLayerError> { if limit == 0 { return Ok(Vec::new()); @@ -258,15 +310,21 @@ impl VideoTaskReadRepository for SqliteVideoTaskRepository { #[async_trait] impl VideoTaskWriteRepository for SqliteVideoTaskRepository { - async fn upsert(&self, task: UpsertVideoTask) -> Result { + async fn upsert(&self, mut task: UpsertVideoTask) -> Result { + task.sanitize_for_persistence(); let id = task.id.clone(); + let expected_identity = task.clone(); bind_task(sqlx::query(UPSERT_SQL), task, true, false)? .execute(&self.pool) .await .map_sql_err()?; - self.find_by_id(&id) - .await? - .ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into())) + let stored = self.find_by_id(&id).await?.ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "video task {id} conflicts with persisted immutable identity" + )) + })?; + stored.ensure_immutable_identity_matches(&expected_identity)?; + Ok(stored) } async fn update_if_active( @@ -405,10 +463,26 @@ ON CONFLICT(id) DO UPDATE SET error_code = excluded.error_code, error_message = excluded.error_message, request_metadata = excluded.request_metadata, - created_at = excluded.created_at, + created_at = COALESCE(video_tasks.created_at, excluded.created_at), submitted_at = excluded.submitted_at, completed_at = excluded.completed_at, updated_at = excluded.updated_at +WHERE video_tasks.short_id IS excluded.short_id + AND video_tasks.request_id IS excluded.request_id + AND video_tasks.user_id IS excluded.user_id + AND video_tasks.api_key_id IS excluded.api_key_id + AND video_tasks.external_task_id IS excluded.external_task_id + AND video_tasks.provider_id IS excluded.provider_id + AND video_tasks.endpoint_id IS excluded.endpoint_id + AND video_tasks.key_id IS excluded.key_id + AND video_tasks.client_api_format IS excluded.client_api_format + AND video_tasks.provider_api_format IS excluded.provider_api_format + AND video_tasks.format_converted IS excluded.format_converted + AND video_tasks.model IS excluded.model + AND video_tasks.duration_seconds IS excluded.duration_seconds + AND video_tasks.resolution IS excluded.resolution + AND video_tasks.aspect_ratio IS excluded.aspect_ratio + AND video_tasks.size IS excluded.size "#; const UPDATE_IF_ACTIVE_SQL: &str = r#" @@ -445,20 +519,38 @@ UPDATE video_tasks SET error_code = ?, error_message = ?, request_metadata = ?, - created_at = ?, + created_at = COALESCE(created_at, ?), submitted_at = ?, completed_at = ?, updated_at = ? WHERE id = ? AND status IN ('pending', 'submitted', 'queued', 'processing') + AND short_id IS ? + AND request_id IS ? + AND user_id IS ? + AND api_key_id IS ? + AND external_task_id IS ? + AND provider_id IS ? + AND endpoint_id IS ? + AND key_id IS ? + AND client_api_format IS ? + AND provider_api_format IS ? + AND format_converted IS ? + AND model IS ? + AND duration_seconds IS ? + AND resolution IS ? + AND aspect_ratio IS ? + AND size IS ? "#; fn bind_task<'q>( query: sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>, - task: UpsertVideoTask, + mut task: UpsertVideoTask, include_insert_id: bool, include_update_id: bool, ) -> Result>, DataLayerError> { + task.sanitize_for_persistence(); + let identity = task.clone(); let original_request_body = json_to_string(&task.original_request_body)?; let request_metadata = json_to_string(&task.request_metadata)?; let query = if include_insert_id { @@ -528,12 +620,38 @@ fn bind_task<'q>( "video task updated_at", )?); if include_update_id { - Ok(bound.bind(task.id)) + bind_identity_guard(bound.bind(task.id), identity) } else { Ok(bound) } } +fn bind_identity_guard<'q>( + query: sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>, + identity: UpsertVideoTask, +) -> Result>, DataLayerError> { + Ok(query + .bind(identity.short_id) + .bind(identity.request_id) + .bind(identity.user_id) + .bind(identity.api_key_id) + .bind(identity.external_task_id) + .bind(identity.provider_id) + .bind(identity.endpoint_id) + .bind(identity.key_id) + .bind(identity.client_api_format) + .bind(identity.provider_api_format) + .bind(identity.format_converted) + .bind(identity.model) + .bind(optional_u32_to_i32( + identity.duration_seconds, + "video task duration_seconds", + )?) + .bind(identity.resolution) + .bind(identity.aspect_ratio) + .bind(identity.size)) +} + fn push_filter<'args>( builder: &mut QueryBuilder<'args, Sqlite>, filter: &'args VideoTaskQueryFilter, @@ -699,7 +817,7 @@ fn optional_u32_to_i32(value: Option, name: &str) -> Result, Da #[cfg(test)] mod tests { - use super::SqliteVideoTaskRepository; + use super::{SqliteVideoTaskRepository, UPDATE_IF_ACTIVE_SQL, UPSERT_SQL}; use crate::run_migrations; use aether_data_contracts::repository::video_tasks::{ UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository, @@ -708,6 +826,40 @@ mod tests { use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions}; use std::{sync::Arc, time::Duration}; + #[test] + fn sqlite_write_sql_guards_every_immutable_identity_field() { + for column in [ + "short_id", + "request_id", + "user_id", + "api_key_id", + "external_task_id", + "provider_id", + "endpoint_id", + "key_id", + "client_api_format", + "provider_api_format", + "format_converted", + "model", + "duration_seconds", + "resolution", + "aspect_ratio", + "size", + ] { + assert!( + UPSERT_SQL.contains(&format!("video_tasks.{column} IS excluded.{column}")), + "upsert should guard {column}" + ); + assert!( + UPDATE_IF_ACTIVE_SQL.contains(&format!("{column} IS ?")), + "active update should guard {column}" + ); + } + assert!(UPSERT_SQL + .contains("created_at = COALESCE(video_tasks.created_at, excluded.created_at)")); + assert!(UPDATE_IF_ACTIVE_SQL.contains("created_at = COALESCE(created_at, ?)")); + } + #[tokio::test] async fn sqlite_repository_writes_and_reads_video_tasks() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -750,6 +902,26 @@ mod tests { .await .expect("user/external lookup should load") .is_some()); + assert!(repository + .find_for_user(VideoTaskLookupKey::Id("task-1"), "user-1") + .await + .expect("owner id lookup should load") + .is_some()); + assert!(repository + .find_for_user(VideoTaskLookupKey::Id("task-1"), "user-2") + .await + .expect("foreign id lookup should run") + .is_none()); + assert!(repository + .find_for_user(VideoTaskLookupKey::ShortId("short-task-1"), "user-1") + .await + .expect("owner short id lookup should load") + .is_some()); + assert!(repository + .find_for_user(VideoTaskLookupKey::ShortId("short-task-1"), "user-2") + .await + .expect("foreign short id lookup should run") + .is_none()); let due = repository .list_due(100, 10) @@ -795,6 +967,8 @@ mod tests { .update_if_active(UpsertVideoTask { status: VideoTaskStatus::Processing, progress_percent: 50, + created_at_unix_ms: 90, + submitted_at_unix_secs: Some(90), ..sample_task("task-1", VideoTaskStatus::Processing, 150) }) .await @@ -803,6 +977,86 @@ mod tests { assert_eq!(updated.progress_percent, 50); } + #[tokio::test] + async fn sqlite_rejects_identity_conflicts_without_modifying_task_state() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let repository = SqliteVideoTaskRepository::new(pool); + let original = sample_task("task-owned", VideoTaskStatus::Submitted, 100); + repository + .upsert(original.clone()) + .await + .expect("original task should insert"); + + let conflict = repository + .upsert(UpsertVideoTask { + user_id: Some("attacker".to_string()), + key_id: Some("attacker-key".to_string()), + status: VideoTaskStatus::Completed, + progress_percent: 100, + completed_at_unix_secs: Some(200), + updated_at_unix_secs: 200, + ..original.clone() + }) + .await + .expect_err("conflicting owner should be rejected"); + assert!(conflict.to_string().contains("immutable field user_id")); + + let after_upsert = repository + .find(VideoTaskLookupKey::Id("task-owned")) + .await + .expect("task lookup should succeed") + .expect("original task should remain"); + assert_eq!(after_upsert.user_id.as_deref(), Some("user-1")); + assert_eq!(after_upsert.key_id.as_deref(), Some("provider-key-1")); + assert_eq!(after_upsert.status, VideoTaskStatus::Submitted); + assert_eq!(after_upsert.progress_percent, 0); + assert_eq!(after_upsert.completed_at_unix_secs, None); + assert_eq!(after_upsert.updated_at_unix_secs, 100); + + let active_conflict = repository + .update_if_active(UpsertVideoTask { + request_id: "attacker-request".to_string(), + status: VideoTaskStatus::Failed, + updated_at_unix_secs: 300, + ..original.clone() + }) + .await + .expect("guarded active update should execute"); + assert!(active_conflict.is_none()); + let after_active_conflict = repository + .find(VideoTaskLookupKey::Id("task-owned")) + .await + .expect("task lookup should succeed") + .expect("original task should remain"); + assert_eq!(after_active_conflict.request_id, "request-task-owned"); + assert_eq!(after_active_conflict.status, VideoTaskStatus::Submitted); + assert_eq!(after_active_conflict.updated_at_unix_secs, 100); + + let updated = repository + .upsert(UpsertVideoTask { + status: VideoTaskStatus::Processing, + progress_percent: 50, + poll_count: 2, + created_at_unix_ms: 999, + updated_at_unix_secs: 200, + ..original + }) + .await + .expect("same owner state update should succeed"); + assert_eq!(updated.status, VideoTaskStatus::Processing); + assert_eq!(updated.progress_percent, 50); + assert_eq!(updated.poll_count, 2); + assert_eq!(updated.created_at_unix_ms, 90); + } + #[tokio::test] async fn sqlite_claim_due_does_not_return_one_task_to_multiple_workers() { const WORKERS: usize = 8; diff --git a/crates/aether-data/adapters/sqlite/src/wallet.rs b/crates/aether-data/adapters/sqlite/src/wallet.rs index 53cd46e27..e1c8c34fc 100644 --- a/crates/aether-data/adapters/sqlite/src/wallet.rs +++ b/crates/aether-data/adapters/sqlite/src/wallet.rs @@ -1,19 +1,41 @@ use crate::error::SqlResultExt; use crate::{sqlite_optional_real, sqlite_real, SqlitePool}; -use aether_data_contracts::repository::wallet::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, redeem_code_refundable_amount, +use aether_data_contracts::repository::billing::{ + checked_plan_duration_days_from_snapshot, entitlements_have_replacement_selector, + entitlements_should_replace_existing, }; use aether_data_contracts::repository::wallet::{ + canonicalize_payment_method, canonicalize_wallet_refund_fields, + payment_callback_amount_matches_order, payment_callback_method_matches_order, + payment_callback_provider_matches_order, payment_order_is_failed_wallet_checkout_placeholder, + payment_order_is_uncertain_wallet_checkout_placeholder, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_payment_method, + redeem_code_refundable_amount, validate_admin_redeem_code_batch_input, + validate_manual_wallet_recharge, validate_payment_order_credit_amounts, + validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements, + validate_redeem_wallet_credit, validate_wallet_recharge_order_input, + wallet_recharge_replay_matches, +}; +use aether_data_contracts::repository::wallet::{ + wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token, + wallet_recharge_checkout_failed_response, wallet_recharge_checkout_uncertain_response, + wallet_recharge_order_is_checkout_placeholder, + wallet_recharge_order_is_reclaimable_placeholder, + wallet_recharge_response_is_checkout_placeholder, wallet_refund_proof_is_success, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, - AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, - CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, + CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, + ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, @@ -22,7 +44,8 @@ use aether_data_contracts::repository::wallet::{ StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use aether_data_contracts::DataLayerError; @@ -107,6 +130,20 @@ impl WalletReadRepository for SqliteWalletReadRepository { ) -> Result, DataLayerError> { initialize_sqlite_auth_wallet(&self.pool, Some(user_id), None, initial_gift_usd, unlimited) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_sqlite_auth_wallet(&self.pool, Some(user_id), None, initial_gift_usd, unlimited) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn initialize_auth_api_key_wallet( @@ -123,6 +160,26 @@ impl WalletReadRepository for SqliteWalletReadRepository { unlimited, ) .await + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + initialize_sqlite_auth_wallet( + &self.pool, + None, + Some(api_key_id), + initial_gift_usd, + unlimited, + ) + .await + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn update_auth_user_wallet_snapshot( @@ -590,7 +647,7 @@ WHERE (? IS NULL OR payment_method = ?) ? IS NULL OR ( CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END ) = ? @@ -623,7 +680,7 @@ WHERE (? IS NULL OR payment_method = ?) ? IS NULL OR ( CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END ) = ? @@ -669,7 +726,7 @@ LIMIT ? OFFSET ? offset: usize, ) -> Result { let total = read_count_row( - sqlx::query("SELECT COUNT(*) AS total FROM payment_orders WHERE user_id = ?") + sqlx::query("SELECT COUNT(*) AS total FROM payment_orders WHERE user_id = ? AND order_kind = 'wallet_recharge'") .bind(user_id) .fetch_one(&self.pool) .await @@ -684,7 +741,7 @@ SELECT payment_provider, payment_channel, order_kind, product_id, product_snapshot, gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END AS status, created_at AS created_at_unix_ms, @@ -693,6 +750,7 @@ SELECT expires_at AS expires_at_unix_secs FROM payment_orders WHERE user_id = ? + AND order_kind = 'wallet_recharge' ORDER BY created_at DESC LIMIT ? OFFSET ? "#, @@ -761,7 +819,7 @@ SELECT payment_provider, payment_channel, order_kind, product_id, product_snapshot, gateway_order_id, gateway_response, CASE - WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired' + WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired' ELSE status END AS status, created_at AS created_at_unix_ms, @@ -771,6 +829,7 @@ SELECT FROM payment_orders WHERE user_id = ? AND id = ? + AND order_kind = 'wallet_recharge' LIMIT 1 "#, ) @@ -783,6 +842,23 @@ LIMIT 1 row.as_ref().map(map_payment_order_row).transpose() } + async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, DataLayerError> { + let sql = payment_order_select_sql( + "WHERE user_id = ? AND order_no = ? AND order_kind = 'wallet_recharge' LIMIT 1", + ); + let row = sqlx::query(&sql) + .bind(user_id) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_payment_order_row).transpose() + } + async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, @@ -809,6 +885,19 @@ LIMIT 1 row.as_ref().map(map_payment_order_row).transpose() } + async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, DataLayerError> { + let sql = payment_order_select_sql("WHERE order_no = ? LIMIT 1"); + let row = sqlx::query(&sql) + .bind(order_no) + .fetch_optional(&self.pool) + .await + .map_sql_err()?; + row.as_ref().map(map_payment_order_row).transpose() + } + async fn find_wallet_refund( &self, wallet_id: &str, @@ -1005,17 +1094,422 @@ WHERE batch_id = ? #[async_trait] impl WalletWriteRepository for SqliteWalletReadRepository { + async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: WalletLookupKey<'_>, + ) -> Result { + if wallet_id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + // The reference predicates and the delete must run while holding SQLite's single + // writer slot. A deferred transaction could observe an unreferenced wallet, then let a + // concurrent writer attach a financial row before the delete is issued. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let select_sql = format!( + r#" +SELECT id +FROM wallets +WHERE id = ? + AND {owner_clause} + AND balance = 0 + AND gift_balance = 0 + AND total_recharged = 0 + AND total_consumed = 0 + AND total_refunded = 0 + AND total_adjusted = 0 + AND limit_mode IN ('finite', 'unlimited') + AND currency = 'USD' + AND status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM "usage" u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1 + "# + ); + let found = sqlx::query_scalar::<_, String>(&select_sql) + .bind(wallet_id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(found_id) = found else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let removed = sqlx::query("DELETE FROM wallets WHERE id = ?") + .bind(&found_id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + + async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if expected.id.trim().is_empty() { + return Ok(false); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + // Serialize the snapshot check, reference check, and delete behind SQLite's writer lock. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let select_sql = wallet_select_sql(&format!( + r#"WHERE id = ? + AND {owner_clause} + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = wallets.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM "usage" u WHERE u.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = wallets.id + ) + AND NOT EXISTS (SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = wallets.id) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = wallets.id + ) +LIMIT 1"# + )); + let row = sqlx::query(&select_sql) + .bind(&expected.id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if ¤t != expected { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let removed = sqlx::query("DELETE FROM wallets WHERE id = ?") + .bind(&expected.id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + + async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if before.id.trim().is_empty() || after.id.trim().is_empty() { + return Ok(false); + } + if before.id != after.id { + return Err(DataLayerError::InvalidInput( + "wallet restore snapshots must reference the same wallet".to_string(), + )); + } + let (owner_clause, owner_id) = match owner { + WalletLookupKey::UserId(user_id) if !user_id.trim().is_empty() => { + ("user_id = ? AND api_key_id IS NULL", user_id) + } + WalletLookupKey::ApiKeyId(api_key_id) if !api_key_id.trim().is_empty() => { + ("api_key_id = ? AND user_id IS NULL", api_key_id) + } + WalletLookupKey::WalletId(_) => { + return Err(DataLayerError::InvalidInput( + "wallet restore requires an explicit user or API-key owner".to_string(), + )) + } + _ => return Ok(false), + }; + let owner_matches = match owner { + WalletLookupKey::UserId(user_id) => { + before.user_id.as_deref() == Some(user_id) + && after.user_id.as_deref() == Some(user_id) + && before.api_key_id.is_none() + && after.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + before.api_key_id.as_deref() == Some(api_key_id) + && after.api_key_id.as_deref() == Some(api_key_id) + && before.user_id.is_none() + && after.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if !owner_matches { + return Ok(false); + } + let before_updated_at = i64::try_from(before.updated_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput( + "wallet restore timestamp is outside the supported range".to_string(), + ) + })?; + + // BEGIN IMMEDIATE serializes the snapshot check and replacement with all SQLite wallet + // writers. No caller can change the row after it is read but before it is restored. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let select_sql = wallet_select_sql(&format!("WHERE id = ? AND {owner_clause} LIMIT 1")); + let row = sqlx::query(&select_sql) + .bind(&after.id) + .bind(owner_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(row) = row else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + let current = map_wallet_row(&row)?; + if current != *after { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + let updated = sqlx::query( + r#" +UPDATE wallets +SET balance = ?, + gift_balance = ?, + limit_mode = ?, + currency = ?, + status = ?, + total_recharged = ?, + total_consumed = ?, + total_refunded = ?, + total_adjusted = ?, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(before.balance) + .bind(before.gift_balance) + .bind(&before.limit_mode) + .bind(&before.currency) + .bind(&before.status) + .bind(before.total_recharged) + .bind(before.total_consumed) + .bind(before.total_refunded) + .bind(before.total_adjusted) + .bind(before_updated_at) + .bind(&before.id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected(); + if updated == 0 { + tx.rollback().await.map_sql_err()?; + return Ok(false); + } + tx.commit().await.map_sql_err()?; + Ok(true) + } + + async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result { + if wallet_id.trim().is_empty() || user_id.trim().is_empty() { + return Ok(false); + } + // Keep the eligibility read and the compensating delete atomic with respect to any new + // wallet/order/ledger writes (SQLite transactions are deferred by default). + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let found_wallet_id = sqlx::query_scalar::<_, String>( + r#" +SELECT w.id +FROM wallets AS w +WHERE w.id = ? + AND w.user_id = ? + AND w.api_key_id IS NULL + AND w.balance = 0 + AND w.gift_balance >= 0 + AND w.total_recharged = 0 + AND w.total_consumed = 0 + AND w.total_refunded = 0 + AND w.total_adjusted = w.gift_balance + AND w.limit_mode IN ('finite', 'unlimited') + AND w.currency = 'USD' + AND w.status = 'active' + AND NOT EXISTS (SELECT 1 FROM payment_orders p WHERE p.wallet_id = w.id) + AND NOT EXISTS (SELECT 1 FROM refund_requests r WHERE r.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM wallet_daily_usage_ledgers d WHERE d.wallet_id = w.id + ) + AND NOT EXISTS (SELECT 1 FROM "usage" u WHERE u.wallet_id = w.id) + AND NOT EXISTS ( + SELECT 1 FROM usage_settlement_snapshots s WHERE s.wallet_id = w.id + ) + AND NOT EXISTS ( + SELECT 1 FROM redeem_codes c WHERE c.redeemed_wallet_id = w.id + ) + AND ( + (w.gift_balance = 0 AND NOT EXISTS ( + SELECT 1 FROM wallet_transactions t WHERE t.wallet_id = w.id + )) + OR + (w.gift_balance > 0 + AND (SELECT COUNT(*) FROM wallet_transactions t WHERE t.wallet_id = w.id) = 1 + AND EXISTS ( + SELECT 1 FROM wallet_transactions t + WHERE t.wallet_id = w.id + AND t.category = 'gift' + AND t.reason_code = 'gift_initial' + AND t.amount = w.gift_balance + AND t.balance_before = 0 + AND t.balance_after = w.gift_balance + AND t.recharge_balance_before = 0 + AND t.recharge_balance_after = 0 + AND t.gift_balance_before = 0 + AND t.gift_balance_after = w.gift_balance + AND t.link_type = 'system_task' + AND t.link_id = w.user_id + AND t.operator_id IS NULL + )) + ) +LIMIT 1 + "#, + ) + .bind(wallet_id) + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(found_wallet_id) = found_wallet_id else { + tx.rollback().await.map_sql_err()?; + return Ok(false); + }; + sqlx::query("DELETE FROM wallet_transactions WHERE wallet_id = ?") + .bind(&found_wallet_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + let removed = sqlx::query("DELETE FROM wallets WHERE id = ? AND user_id = ?") + .bind(&found_wallet_id) + .bind(user_id) + .execute(&mut *tx) + .await + .map_sql_err()? + .rows_affected() + > 0; + tx.commit().await.map_sql_err()?; + Ok(removed) + } + async fn create_wallet_recharge_order( &self, - input: CreateWalletRechargeOrderInput, + mut input: CreateWalletRechargeOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_wallet_recharge_order_input(&input).map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid wallet recharge numeric fields".to_string(), + )); + } + let projected_gateway_response = + project_wallet_recharge_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { DataLayerError::InvalidInput("wallet recharge expires_at overflow".to_string()) })?; - let gateway_response = - json_string(&input.gateway_response, "payment_orders.gateway_response")?; - let mut tx = self.pool.begin().await.map_sql_err()?; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; + // Serialize wallet creation and order idempotency checks with the insert. A deferred + // transaction can let concurrent callers race through the read-before-write section. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + + // The wallet schema intentionally keeps `user_id` nullable for + // deleted-user history, so it cannot protect this creation path with + // a mandatory foreign key. Validate the owner while holding the + // writer transaction before creating either the wallet or order. + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } let wallet_row = sqlx::query( r#" @@ -1029,14 +1523,18 @@ LIMIT 1 .fetch_optional(&mut *tx) .await .map_sql_err()?; - let (wallet_id, wallet_status) = if let Some(row) = wallet_row { - (get::(&row, "id")?, get::(&row, "status")?) + let (wallet_id, wallet_status, created_wallet) = if let Some(row) = wallet_row { + ( + get::(&row, "id")?, + get::(&row, "status")?, + false, + ) } else { - let wallet_id = input + let requested_wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - sqlx::query( + let insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, @@ -1044,24 +1542,71 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?) +ON CONFLICT DO NOTHING "#, ) - .bind(&wallet_id) + .bind(&requested_wallet_id) .bind(&input.user_id) .bind(now) .bind(now) .execute(&mut *tx) .await .map_sql_err()?; - (wallet_id, "active".to_string()) + if insert_result.rows_affected() == 0 { + let Some(row) = sqlite_wallet_by_user_id(&mut tx, &input.user_id).await? else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + ( + get::(&row, "id")?, + get::(&row, "status")?, + false, + ) + } else { + (requested_wallet_id, "active".to_string(), true) + } }; if wallet_status != "active" { - tx.commit().await.map_sql_err()?; + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } return Ok(CreateWalletRechargeOrderOutcome::WalletInactive); } + if let Some(existing_row) = + sqlite_payment_order_by_order_no(&mut tx, &input.order_no).await? + { + let existing_user_id: Option = existing_row.try_get("user_id").map_sql_err()?; + let existing_kind: String = existing_row.try_get("order_kind").map_sql_err()?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !sqlite_wallet_recharge_replay_matches(&existing_row, &wallet_id, &input)? { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + let existing = map_payment_order_row(&existing_row)?; + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + let order_id = uuid::Uuid::new_v4().to_string(); - sqlx::query( + let insert_result = sqlx::query( r#" INSERT INTO payment_orders ( id, order_no, wallet_id, user_id, amount_usd, pay_amount, pay_currency, @@ -1070,6 +1615,7 @@ INSERT INTO payment_orders ( gateway_order_id, gateway_response, status, created_at, expires_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'wallet_recharge', 'pending', ?, ?, 'pending', ?, ?) +ON CONFLICT DO NOTHING "#, ) .bind(&order_id) @@ -1091,6 +1637,54 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'wallet_recharge', 'pending', ?, .await .map_sql_err()?; + if insert_result.rows_affected() == 0 { + if let Some(existing_row) = + sqlite_payment_order_by_order_no(&mut tx, &input.order_no).await? + { + let existing_user_id: Option = + existing_row.try_get("user_id").map_sql_err()?; + let existing_kind: String = existing_row.try_get("order_kind").map_sql_err()?; + let existing = map_payment_order_row(&existing_row)?; + if existing_user_id.as_deref() == Some(input.user_id.as_str()) + && existing_kind == "wallet_recharge" + { + if !sqlite_wallet_recharge_replay_matches(&existing_row, &wallet_id, &input)? { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if created_wallet { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + if sqlite_payment_order_by_gateway_order_id( + &mut tx, + &input.payment_method, + &input.gateway_order_id, + ) + .await? + .is_some() + { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "payment gateway order already belongs to another order".to_string(), + )); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet recharge order could not be created".to_string(), + )); + } + let row = sqlite_payment_order_by_id(&mut tx, &order_id).await?; tx.commit().await.map_sql_err()?; Ok(CreateWalletRechargeOrderOutcome::Created( @@ -1098,19 +1692,342 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'wallet_recharge', 'pending', ?, )) } + async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() || input.gateway_order_id.trim().is_empty() { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout identifiers are required".to_string(), + )); + } + let projected_gateway_response = + match project_wallet_recharge_gateway_response(&input.gateway_response) { + Ok(value) => value, + Err(error) => return Ok(WalletMutationOutcome::Invalid(error)), + }; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let Some(current_row) = + sqlite_payment_order_by_id_optional(&mut tx, &input.order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = get(¤t_row, "order_kind")?; + if order_kind.as_deref() != Some("wallet_recharge") { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a wallet recharge".to_string(), + )); + } + let current = map_payment_order_row(¤t_row)?; + let current_is_checkout_placeholder = + wallet_recharge_order_is_checkout_placeholder(¤t); + let current_token = current + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + let requested_token = wallet_recharge_checkout_claim_token(&projected_gateway_response); + if current_token.is_some() && current_token != requested_token { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if current.status != "pending" { + if current.gateway_order_id.as_deref() == Some(input.gateway_order_id.as_str()) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(current)); + } + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is no longer pending".to_string(), + )); + } + let now = current_unix_secs_i64(); + if current + .expires_at_unix_secs + .is_none_or(|expires_at| expires_at <= now.max(0) as u64) + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired".to_string(), + )); + } + // A newly-created row uses order_no as a temporary gateway id. Once + // the provider checkout is stored, do not let a concurrent request + // replace that checkout evidence. + if current.gateway_order_id.as_deref().is_some_and(|existing| { + existing != input.gateway_order_id.as_str() + && existing != current.order_no.as_str() + && !current_is_checkout_placeholder + }) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is already bound".to_string(), + )); + } + let conflict = sqlite_payment_order_by_gateway_order_id( + &mut tx, + ¤t.payment_method, + &input.gateway_order_id, + ) + .await? + .is_some_and(|row| { + row.try_get::("id").ok().as_deref() != Some(input.order_id.as_str()) + }); + if conflict { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment gateway order already belongs to another order".to_string(), + )); + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = ?, gateway_response = ? WHERE id = ? AND status = 'pending' AND expires_at > ?", + ) + .bind(&input.gateway_order_id) + .bind(gateway_response) + .bind(&input.order_id) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired or no longer pending".to_string(), + )); + } + let updated = + map_payment_order_row(&sqlite_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + + async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result { + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let Some(row) = sqlite_payment_order_by_id_optional(&mut tx, &input.order_id).await? else { + tx.commit().await.map_sql_err()?; + return Ok(false); + }; + let current = map_payment_order_row(&row)?; + let Some(replacement) = + payment_order_stripe_client_secret_cas_replacement(¤t, &input) + .map_err(DataLayerError::InvalidInput)? + else { + tx.commit().await.map_sql_err()?; + return Ok(false); + }; + let replacement = json_string(&replacement, "payment_orders.gateway_response")?; + let updated = sqlx::query("UPDATE payment_orders SET gateway_response = ? WHERE id = ?") + .bind(replacement) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + tx.commit().await.map_sql_err()?; + Ok(updated.rows_affected() == 1) + } + + async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout failure identifiers are required".to_string(), + )); + } + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let Some(row) = sqlite_payment_order_by_id_optional(&mut tx, &input.order_id).await? else { + tx.rollback().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order = map_payment_order_row(&row)?; + if !wallet_recharge_order_is_checkout_placeholder(&order) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a checkout placeholder".to_string(), + )); + } + let current_token = order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + if current_token != Some(input.claim_token.trim()) { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if order.status != "pending" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(order)); + } + let failed = if input.provider_request_may_have_succeeded { + wallet_recharge_checkout_uncertain_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs_i64().max(0) as u64, + ) + } else { + wallet_recharge_checkout_failed_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs_i64().max(0) as u64, + ) + }; + let failed = json_string(&failed, "payment_orders.gateway_response")?; + let updated = sqlx::query( + "UPDATE payment_orders SET status = 'failed', gateway_response = ? WHERE id = ? AND status = 'pending'", + ) + .bind(failed) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + let updated = + map_payment_order_row(&sqlite_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + + async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + let now = current_unix_secs_i64().max(0) as u64; + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + || input.expires_at_unix_secs <= now + || input.expires_at_unix_secs > i64::MAX as u64 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim identifiers are invalid".to_string(), + )); + } + if !wallet_recharge_response_is_checkout_placeholder(&input.gateway_response) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge reclaim response must be a placeholder".to_string(), + )); + } + let response = wallet_recharge_checkout_claim_response( + &input.gateway_response, + &input.claim_token, + now, + ) + .map_err(DataLayerError::InvalidInput)?; + let response = json_string(&response, "payment_orders.gateway_response")?; + let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { + DataLayerError::InvalidInput("wallet recharge expires_at overflow".to_string()) + })?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let Some(row) = sqlite_payment_order_by_id_optional(&mut tx, &input.order_id).await? else { + tx.rollback().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let order_kind: Option = get(&row, "order_kind")?; + let order = map_payment_order_row(&row)?; + if order_kind.as_deref() != Some("wallet_recharge") + || !wallet_recharge_order_is_reclaimable_placeholder(&order, now) + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is still in progress or already completed".to_string(), + )); + } + let updated = sqlx::query( + "UPDATE payment_orders SET gateway_order_id = ?, gateway_response = ?, status = 'pending', expires_at = ? WHERE id = ?", + ) + .bind(&order.order_no) + .bind(response) + .bind(expires_at) + .bind(&input.order_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if updated.rows_affected() == 0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim lost the order race".to_string(), + )); + } + let updated = + map_payment_order_row(&sqlite_payment_order_by_id(&mut tx, &input.order_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + async fn create_plan_purchase_order( &self, - input: CreatePlanPurchaseOrderInput, + mut input: CreatePlanPurchaseOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_plan_purchase_order_input(&input).map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let expires_at = i64::try_from(input.expires_at_unix_secs).map_err(|_| { DataLayerError::InvalidInput("plan purchase expires_at overflow".to_string()) })?; - let gateway_response = - json_string(&input.gateway_response, "payment_orders.gateway_response")?; + let projected_gateway_response = project_wallet_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; + let gateway_response = json_string( + &projected_gateway_response, + "payment_orders.gateway_response", + )?; let product_snapshot = json_string(&input.product_snapshot, "payment_orders.product_snapshot")?; - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + + // The baseline SQLite schema does not enforce the wallet -> user + // relationship. Validate the order owner before creating an automatic + // wallet, so an invalid checkout cannot leave a financial row behind. + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput("user not found".to_string())); + } let wallet_row = sqlx::query( r#" @@ -1127,11 +2044,16 @@ LIMIT 1 let (wallet_id, wallet_status) = if let Some(row) = wallet_row { (get::(&row, "id")?, get::(&row, "status")?) } else { - let wallet_id = input + let requested_wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - sqlx::query( + // `user_id` is unique in the portable schema, but the explicit + // read above can still race with another initializer on a + // connection that started before it committed. Treat either a + // user or wallet-id conflict as a no-op, then resolve the winner + // by owner instead of leaking a database constraint error. + let insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, @@ -1139,16 +2061,27 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?) +ON CONFLICT DO NOTHING "#, ) - .bind(&wallet_id) + .bind(&requested_wallet_id) .bind(&input.user_id) .bind(now) .bind(now) .execute(&mut *tx) .await .map_sql_err()?; - (wallet_id, "active".to_string()) + if insert_result.rows_affected() == 0 { + let Some(row) = sqlite_wallet_by_user_id(&mut tx, &input.user_id).await? else { + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + }; + (get::(&row, "id")?, get::(&row, "status")?) + } else { + (requested_wallet_id, "active".to_string()) + } }; if wallet_status != "active" { tx.commit().await.map_sql_err()?; @@ -1258,8 +2191,39 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'plan_purchase', ?, ?, 'pending', &self, input: CreateWalletRefundRequestInput, ) -> Result { + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "refund amount must be finite and greater than zero".to_string(), + )); + } let now = current_unix_secs_i64(); - let mut tx = self.pool.begin().await.map_sql_err()?; + // Hold SQLite's single writer slot while checking existing reservations and inserting + // the new request. A deferred transaction would allow two callers to both observe the + // same available balance before either reservation is written. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + + let Some(wallet_row) = sqlx::query( + r#" +SELECT id, balance +FROM wallets +WHERE id = ? + AND user_id = ? +LIMIT 1 +"#, + ) + .bind(&input.wallet_id) + .bind(&input.user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::WalletMissing); + }; if let Some(idempotency_key) = input.idempotency_key.as_deref() { let existing = @@ -1272,54 +2236,51 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, ?, ?, 'plan_purchase', ?, ?, 'pending', } } - let Some(wallet_row) = sqlx::query( - r#" -SELECT id, balance -FROM wallets -WHERE id = ? -LIMIT 1 -"#, - ) - .bind(&input.wallet_id) - .fetch_optional(&mut *tx) - .await - .map_sql_err()? - else { - tx.commit().await.map_sql_err()?; - return Ok(CreateWalletRefundRequestOutcome::WalletMissing); - }; let wallet_recharge_balance = sqlite_real(&wallet_row, "balance")?; - let wallet_reserved_amount: f64 = sqlx::query_scalar( + if !wallet_recharge_balance.is_finite() { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet recharge balance is invalid".to_string(), + )); + } + let wallet_reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(SUM(amount_usd), 0.0) +SELECT amount_usd FROM refund_requests WHERE wallet_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(&input.wallet_id) - .fetch_one(&mut *tx) + .fetch_all(&mut *tx) .await - .map_sql_err()?; + .map_sql_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(wallet_reserved_amount) = wallet_reserved_amount else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet refund reservation is invalid".to_string(), + )); + }; if input.amount_usd > (wallet_recharge_balance - wallet_reserved_amount) { tx.commit().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableBalance); } let mut payment_order_id = None; - let mut source_type = input - .source_type - .clone() - .unwrap_or_else(|| "wallet_balance".to_string()); - let mut source_id = input.source_id.clone(); - let mut refund_mode = input - .refund_mode - .clone() - .unwrap_or_else(|| "offline_payout".to_string()); + let mut resolved_payment_method = None; if let Some(order_id) = input.payment_order_id.as_deref() { let Some(order_row) = sqlx::query( r#" -SELECT id, status, payment_method, refundable_amount_usd +SELECT id, status, payment_method, amount_usd, refunded_amount_usd, refundable_amount_usd FROM payment_orders WHERE id = ? AND wallet_id = ? @@ -1340,19 +2301,46 @@ LIMIT 1 tx.commit().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotRefundable); } - let order_reserved_amount: f64 = sqlx::query_scalar( + let order_reserved_amount = sqlx::query_scalar::<_, Option>( r#" -SELECT COALESCE(SUM(amount_usd), 0.0) +SELECT amount_usd FROM refund_requests WHERE payment_order_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(order_id) - .fetch_one(&mut *tx) + .fetch_all(&mut *tx) .await - .map_sql_err()?; + .map_sql_err()? + .into_iter() + .try_fold(0.0_f64, |total, amount| { + let amount = amount?; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(order_reserved_amount) = order_reserved_amount else { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund reservation is invalid".to_string(), + )); + }; + let order_amount = sqlite_real(&order_row, "amount_usd")?; + let refunded_amount = sqlite_real(&order_row, "refunded_amount_usd")?; let refundable_amount = sqlite_real(&order_row, "refundable_amount_usd")?; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_amount, + refundable_amount, + ) { + tx.commit().await.map_sql_err()?; + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund amounts are invalid".to_string(), + )); + } if input.amount_usd > (refundable_amount - order_reserved_amount) { tx.commit().await.map_sql_err()?; return Ok( @@ -1360,14 +2348,21 @@ WHERE payment_order_id = ? ); } payment_order_id = Some(order_id.to_string()); - source_type = "payment_order".to_string(); - source_id = Some(order_id.to_string()); - if input.refund_mode.is_none() { - let payment_method: String = get(&order_row, "payment_method")?; - refund_mode = default_refund_mode_for_payment_method(&payment_method).to_string(); - } + resolved_payment_method = Some(get::(&order_row, "payment_method")?); } + let canonical = canonicalize_wallet_refund_fields( + payment_order_id.as_deref(), + input.source_type.as_deref(), + input.source_id.as_deref(), + input.refund_mode.as_deref(), + resolved_payment_method.as_deref(), + ) + .map_err(DataLayerError::InvalidInput)?; + let source_type = canonical.source_type; + let source_id = canonical.source_id; + let refund_mode = canonical.refund_mode; + let refund_id = uuid::Uuid::new_v4().to_string(); let insert = sqlx::query( r#" @@ -1397,7 +2392,11 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending_approval', ?, ?, ?, ?, ?) .await; if let Err(err) = insert { - if input.idempotency_key.is_some() { + if input.idempotency_key.is_some() + && err + .as_database_error() + .is_some_and(|database_error| database_error.is_unique_violation()) + { tx.rollback().await.map_sql_err()?; return Ok(CreateWalletRefundRequestOutcome::DuplicateRejected); } @@ -1413,58 +2412,100 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending_approval', ?, ?, ?, ?, ?) async fn process_payment_callback( &self, - input: ProcessPaymentCallbackInput, + mut input: ProcessPaymentCallbackInput, ) -> Result { + input + .canonicalize_and_validate() + .map_err(DataLayerError::InvalidInput)?; + if input.callback_key.trim().is_empty() + || input.callback_key.chars().count() > 128 + || input.payload_hash.trim().is_empty() + || !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err(DataLayerError::InvalidInput( + "invalid payment callback numeric or identity fields".to_string(), + )); + } let now = current_unix_secs_i64(); let payload = json_string(&input.payload, "payment_callbacks.payload")?; - let mut tx = self.pool.begin().await.map_sql_err()?; + // Callback processing reads and then mutates the callback, order, and + // wallet rows. Acquire SQLite's writer lock before the first read so + // a deferred transaction cannot observe `pending` and later fail while + // upgrading after another callback or admin credit has committed. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; - let existing_callback = sqlx::query( + // Register the callback atomically. A preceding SELECT allows two + // concurrent first deliveries to both observe a missing key and then + // race on the unique constraint. Insert first and re-read by the + // key so every caller uses the row that actually won the race. + let candidate_callback_id = uuid::Uuid::new_v4().to_string(); + sqlx::query( r#" -SELECT id, payment_order_id, status, order_no, gateway_order_id +INSERT INTO payment_callbacks ( + id, payment_order_id, payment_method, callback_key, order_no, gateway_order_id, + payload_hash, signature_valid, status, payload, error_message, created_at, processed_at +) +VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', NULL, NULL, ?, NULL) +ON CONFLICT(callback_key) DO NOTHING +"#, + ) + .bind(&candidate_callback_id) + .bind(&input.payment_method) + .bind(&input.callback_key) + .bind(input.order_no.as_deref()) + .bind(input.gateway_order_id.as_deref()) + .bind(&input.payload_hash) + .bind(sqlite_bool(input.signature_valid)) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + + let callback_row = sqlx::query( + r#" +SELECT id, payment_order_id, payment_method, payload_hash, status, order_no, gateway_order_id FROM payment_callbacks WHERE callback_key = ? LIMIT 1 "#, ) .bind(&input.callback_key) - .fetch_optional(&mut *tx) + .fetch_one(&mut *tx) .await .map_sql_err()?; - let duplicate = existing_callback.is_some(); - let callback_id = if let Some(row) = existing_callback.as_ref() { - let status: String = get(row, "status")?; - if status == "processed" { - let order_id: Option = get(row, "payment_order_id")?; - tx.commit().await.map_sql_err()?; - return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id }); - } - get(row, "id")? - } else { - let callback_id = uuid::Uuid::new_v4().to_string(); - sqlx::query( - r#" -INSERT INTO payment_callbacks ( - id, payment_order_id, payment_method, callback_key, order_no, gateway_order_id, - payload_hash, signature_valid, status, payload, error_message, created_at, processed_at -) -VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) -"#, - ) - .bind(&callback_id) - .bind(&input.payment_method) - .bind(&input.callback_key) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) - .bind(&input.payload_hash) - .bind(sqlite_bool(input.signature_valid)) - .bind(&payload) - .bind(now) - .execute(&mut *tx) - .await - .map_sql_err()?; - callback_id - }; + let callback_id: String = get(&callback_row, "id")?; + let duplicate = callback_id != candidate_callback_id; + let callback_order_no: Option = get(&callback_row, "order_no")?; + let callback_gateway_order_id: Option = get(&callback_row, "gateway_order_id")?; + + let stored_method: String = get(&callback_row, "payment_method")?; + let stored_hash: Option = get(&callback_row, "payload_hash")?; + if !stored_method.eq_ignore_ascii_case(&input.payment_method) + || stored_hash.as_deref() != Some(input.payload_hash.as_str()) + { + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate: true, + error: "callback key reused with different payment payload".to_string(), + }); + } + let status: String = get(&callback_row, "status")?; + if status == "processed" { + let order_id: Option = get(&callback_row, "payment_order_id")?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::DuplicateProcessed { order_id }); + } if !input.signature_valid { update_sqlite_payment_callback_failure( @@ -1482,20 +2523,20 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) }); } - let lookup_order_no = input.order_no.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| get(row, "order_no").ok()) - }); - let lookup_gateway_order_id = input.gateway_order_id.clone().or_else(|| { - existing_callback - .as_ref() - .and_then(|row| get(row, "gateway_order_id").ok()) - }); + let lookup_order_no = input.order_no.clone().or_else(|| callback_order_no.clone()); + let lookup_gateway_order_id = input + .gateway_order_id + .clone() + .or_else(|| callback_gateway_order_id.clone()); let order_row = if let Some(order_no) = lookup_order_no.as_deref() { sqlite_payment_order_by_order_no(&mut tx, order_no).await? } else if let Some(gateway_order_id) = lookup_gateway_order_id.as_deref() { - sqlite_payment_order_by_gateway_order_id(&mut tx, gateway_order_id).await? + sqlite_payment_order_by_gateway_order_id( + &mut tx, + &input.payment_method, + gateway_order_id, + ) + .await? } else { None }; @@ -1521,19 +2562,262 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) let order_payment_method: String = get(&order_row, "payment_method")?; let order_payment_provider: Option = get(&order_row, "payment_provider")?; let order_payment_channel: Option = get(&order_row, "payment_channel")?; + let order_pay_currency: Option = get(&order_row, "pay_currency")?; + let order_gateway_order_id: Option = get(&order_row, "gateway_order_id")?; let order_kind: String = get(&order_row, "order_kind")?; let order_amount_usd = sqlite_real(&order_row, "amount_usd")?; let order_pay_amount = sqlite_optional_real(&order_row, "pay_amount")?; + let order_exchange_rate = sqlite_optional_real(&order_row, "exchange_rate")?; let order_status: String = get(&order_row, "status")?; let expires_at_unix_secs: Option = get(&order_row, "expires_at_unix_secs")?; - - let amount_matches = if let (Some(callback_pay_amount), Some(order_pay_amount)) = - (input.pay_amount, order_pay_amount) - { - (callback_pay_amount - order_pay_amount).abs() <= 0.01 + let order_gateway_response = if order_status.eq_ignore_ascii_case("failed") { + optional_json( + get(&order_row, "gateway_response")?, + "payment_orders.gateway_response", + )? } else { - (input.amount_usd - order_amount_usd).abs() <= f64::EPSILON + None }; + let failed_checkout_recoverable = payment_order_is_failed_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + let uncertain_checkout = payment_order_is_uncertain_wallet_checkout_placeholder( + &order_status, + &order_kind, + order_gateway_response.as_ref(), + ); + if !order_amount_usd.is_finite() + || order_amount_usd <= 0.0 + || order_pay_amount.is_some_and(|value| !value.is_finite() || value <= 0.0) + { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order amount is invalid", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order amount is invalid".to_string(), + }); + } + + // A payment order must credit the wallet that belongs to the same + // live user. Do this before binding a gateway id or changing any + // entitlement, wallet, or order state. Legacy rows may violate the + // wallet owner XOR invariant, so reject every ambiguous shape here. + let order_user_id: Option = get(&order_row, "user_id")?; + let Some(order_user_id) = order_user_id + .as_deref() + .filter(|value| !value.trim().is_empty()) + else { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order user missing", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order user missing".to_string(), + }); + }; + let Some(wallet_owner_row) = sqlx::query( + r#" +SELECT + w.user_id AS wallet_user_id, + w.api_key_id AS wallet_api_key_id, + api_keys.user_id AS api_key_user_id +FROM wallets AS w +LEFT JOIN api_keys ON api_keys.id = w.api_key_id +WHERE w.id = ? +LIMIT 1 + "#, + ) + .bind(&order_wallet_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()? + else { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "wallet not found", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet not found".to_string(), + }); + }; + let wallet_user_id: Option = get(&wallet_owner_row, "wallet_user_id")?; + let wallet_api_key_id: Option = get(&wallet_owner_row, "wallet_api_key_id")?; + let api_key_user_id: Option = get(&wallet_owner_row, "api_key_user_id")?; + let wallet_owner_matches = match ( + wallet_user_id.as_deref(), + wallet_api_key_id.as_deref(), + api_key_user_id.as_deref(), + ) { + (Some(wallet_user_id), None, _) if !wallet_user_id.trim().is_empty() => { + wallet_user_id == order_user_id + } + (None, Some(wallet_api_key_id), Some(api_key_user_id)) + if !wallet_api_key_id.trim().is_empty() && !api_key_user_id.trim().is_empty() => + { + api_key_user_id == order_user_id + } + _ => false, + }; + if !wallet_owner_matches { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order wallet owner mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order wallet owner mismatch".to_string(), + }); + } + + // The lookup identifier is not proof that the callback belongs to + // this order: order_no takes precedence over gateway_order_id. Check + // every identifier supplied by this delivery (and any persisted + // fallback from the callback row) before changing the order or + // wallet. Orders created before the gateway returns a provider + // transaction id store order_no as a placeholder; that value may be + // replaced by a verified callback, but a real id must never be + // rebound to another order. + if input + .order_no + .as_deref() + .is_some_and(|value| value != order_no) + || callback_order_no + .as_deref() + .is_some_and(|value| value != order_no) + { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment order number mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment order number mismatch".to_string(), + }); + } + let input_gateway_order_id = input + .gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let callback_gateway_order_id = callback_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()); + let stored_real_gateway_order_id = order_gateway_order_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty() && *value != order_no); + let input_real_gateway_order_id = input_gateway_order_id.filter(|value| *value != order_no); + let callback_real_gateway_order_id = + callback_gateway_order_id.filter(|value| *value != order_no); + let effective_gateway_order_id = input_real_gateway_order_id + .or(callback_real_gateway_order_id) + .or(stored_real_gateway_order_id); + if let Some(expected_gateway_order_id) = stored_real_gateway_order_id { + if input_gateway_order_id.is_some_and(|value| value != expected_gateway_order_id) + || callback_gateway_order_id.is_some_and(|value| value != expected_gateway_order_id) + { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order mismatch".to_string(), + }); + } + } else if let (Some(input_gateway), Some(callback_gateway)) = + (input_real_gateway_order_id, callback_real_gateway_order_id) + { + if input_gateway != callback_gateway { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order identifier mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order identifier mismatch".to_string(), + }); + } + } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + let conflicting_order_id: Option = sqlx::query_scalar( + "SELECT id FROM payment_orders WHERE payment_method = ? AND gateway_order_id = ? AND id <> ? LIMIT 1", + ) + .bind(&order_payment_method) + .bind(gateway_order_id) + .bind(&order_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if conflicting_order_id.is_some() { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order belongs to another payment order", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order".to_string(), + }); + } + } + } + + let amount_matches = payment_callback_amount_matches_order( + order_amount_usd, + order_pay_amount, + order_pay_currency.as_deref(), + order_exchange_rate, + input.amount_usd, + input.pay_amount, + ); if !amount_matches { update_sqlite_payment_callback_failure( &mut tx, @@ -1549,7 +2833,12 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) error: "callback amount mismatch".to_string(), }); } - if !order_payment_method.eq_ignore_ascii_case(&input.payment_method) { + if !payment_callback_method_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ) { update_sqlite_payment_callback_failure( &mut tx, &callback_id, @@ -1564,31 +2853,57 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) error: "payment method mismatch".to_string(), }); } - if let Some(expected_provider) = input.payment_provider.as_deref() { - if order_payment_provider - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_provider)) - { - update_sqlite_payment_callback_failure( - &mut tx, - &callback_id, - &input, - &payload, - "payment provider mismatch", - ) - .await?; - tx.commit().await.map_sql_err()?; - return Ok(ProcessPaymentCallbackOutcome::Failed { - duplicate, - error: "payment provider mismatch".to_string(), - }); - } + let payment_provider_matches = payment_callback_provider_matches_order( + &order_payment_method, + order_payment_provider.as_deref(), + &input.payment_method, + input.payment_provider.as_deref(), + ); + if !payment_provider_matches { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment provider mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment provider mismatch".to_string(), + }); + } + let currency_matches = match (input.pay_currency.as_deref(), order_pay_currency.as_deref()) + { + (Some(callback), Some(order)) => order.eq_ignore_ascii_case(callback), + (None, None) => input.pay_amount.is_none() && order_pay_amount.is_none(), + _ => false, + }; + if !currency_matches { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment currency mismatch", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment currency mismatch".to_string(), + }); } if let Some(expected_channel) = input.payment_channel.as_deref() { - if order_payment_channel - .as_deref() - .is_some_and(|value| !value.eq_ignore_ascii_case(expected_channel)) - { + let stored_channel = order_payment_channel.as_deref().or_else(|| { + (order_payment_provider.is_none() + && ["alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(&order_payment_method))) + .then_some(order_payment_method.as_str()) + }); + if !stored_channel.is_some_and(|value| value.eq_ignore_ascii_case(expected_channel)) { update_sqlite_payment_callback_failure( &mut tx, &callback_id, @@ -1622,14 +2937,16 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) wallet_id: order_wallet_id, }); } - if matches!(order_status.as_str(), "failed" | "expired" | "refunded") { + if !matches!(order_status.as_str(), "pending" | "paid") && !failed_checkout_recoverable { let error = format!("payment order is not creditable: {order_status}"); update_sqlite_payment_callback_failure(&mut tx, &callback_id, &input, &payload, &error) .await?; tx.commit().await.map_sql_err()?; return Ok(ProcessPaymentCallbackOutcome::Failed { duplicate, error }); } - if order_status == "pending" && expires_at_unix_secs.is_some_and(|value| value < now) { + if (order_status == "pending" || (failed_checkout_recoverable && !uncertain_checkout)) + && expires_at_unix_secs.is_some_and(|value| value <= now) + { sqlx::query("UPDATE payment_orders SET status = 'expired' WHERE id = ?") .bind(&order_id) .execute(&mut *tx) @@ -1650,6 +2967,28 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL) }); } + if stored_real_gateway_order_id.is_none() { + if let Some(gateway_order_id) = effective_gateway_order_id { + if !sqlite_bind_payment_gateway_order_id(&mut tx, &order_id, gateway_order_id) + .await? + { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "payment gateway order belongs to another payment order", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "payment gateway order belongs to another payment order".to_string(), + }); + } + } + } + if order_kind == "plan_purchase" { let order_user_id: Option = get(&order_row, "user_id")?; let Some(user_id) = order_user_id else { @@ -1757,7 +3096,7 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) .bind(&plan_id) .bind(&order_id) .bind(now) - .bind(plan_expires_at_unix(&snapshot, now)) + .bind(plan_expires_at_unix(&snapshot, now)?) .bind(json_string( &entitlements, "user_plan_entitlements.entitlements_snapshot", @@ -1782,9 +3121,9 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) UPDATE payment_orders SET gateway_order_id = COALESCE(?, gateway_order_id), gateway_response = ?, - pay_amount = COALESCE(?, pay_amount), - pay_currency = COALESCE(?, pay_currency), - exchange_rate = COALESCE(?, exchange_rate), + pay_amount = COALESCE(pay_amount, ?), + pay_currency = COALESCE(pay_currency, ?), + exchange_rate = COALESCE(exchange_rate, ?), status = 'credited', fulfillment_status = 'fulfilled', fulfillment_error = NULL, @@ -1794,8 +3133,11 @@ SET gateway_order_id = COALESCE(?, gateway_order_id), WHERE id = ? "#, ) - .bind(input.gateway_order_id.as_deref()) - .bind(&payload) + .bind(effective_gateway_order_id) + .bind(json_string( + &input.gateway_response_projection(&order_no, effective_gateway_order_id), + "payment_orders.gateway_response", + )?) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -1827,7 +3169,7 @@ WHERE id = ? let Some(wallet_row) = sqlx::query( r#" -SELECT id, status, balance, gift_balance +SELECT id, status, balance, gift_balance, total_recharged FROM wallets WHERE id = ? LIMIT 1 @@ -1871,6 +3213,33 @@ LIMIT 1 let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; + let total_recharged = sqlite_real(&wallet_row, "total_recharged")?; + // Finite recharge balances may be negative: usage settlement permits a + // finite wallet to overdraft, and a later recharge must be able to + // restore that balance. Reject only malformed values and arithmetic + // overflow here. + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order_amount_usd).is_finite() + || !(before_recharge + before_gift + order_amount_usd).is_finite() + { + update_sqlite_payment_callback_failure( + &mut tx, + &callback_id, + &input, + &payload, + "wallet balance is invalid", + ) + .await?; + tx.commit().await.map_sql_err()?; + return Ok(ProcessPaymentCallbackOutcome::Failed { + duplicate, + error: "wallet balance is invalid".to_string(), + }); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order_amount_usd; let after_total = after_recharge + before_gift; @@ -1921,9 +3290,9 @@ VALUES (?, ?, 'recharge', 'topup_gateway', ?, ?, ?, ?, ?, ?, ?, 'payment_order', UPDATE payment_orders SET gateway_order_id = COALESCE(?, gateway_order_id), gateway_response = ?, - pay_amount = COALESCE(?, pay_amount), - pay_currency = COALESCE(?, pay_currency), - exchange_rate = COALESCE(?, exchange_rate), + pay_amount = COALESCE(pay_amount, ?), + pay_currency = COALESCE(pay_currency, ?), + exchange_rate = COALESCE(exchange_rate, ?), status = 'credited', paid_at = COALESCE(paid_at, ?), credited_at = ?, @@ -1931,8 +3300,11 @@ SET gateway_order_id = COALESCE(?, gateway_order_id), WHERE id = ? "#, ) - .bind(input.gateway_order_id.as_deref()) - .bind(&payload) + .bind(effective_gateway_order_id) + .bind(json_string( + &input.gateway_response_projection(&order_no, effective_gateway_order_id), + "payment_orders.gateway_response", + )?) .bind(input.pay_amount) .bind(input.pay_currency.as_deref()) .bind(input.exchange_rate) @@ -1966,8 +3338,20 @@ WHERE id = ? &self, input: AdjustWalletBalanceInput, ) -> 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(), + )); + } let now = current_unix_secs_i64(); - let mut tx = self.pool.begin().await.map_sql_err()?; + // Admin credit follows the same read/validate/write sequence as a + // provider callback. Serialize it at transaction start to prevent two + // writers from both crediting a pending order. + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; let Some(row) = sqlite_wallet_by_id_optional(&mut tx, &input.wallet_id).await? else { tx.commit().await.map_sql_err()?; return Ok(None); @@ -1976,6 +3360,16 @@ WHERE id = ? let before_recharge = sqlite_real(&row, "balance")?; let before_gift = sqlite_real(&row, "gift_balance")?; let before_total = before_recharge + before_gift; + let before_total_adjusted = sqlite_real(&row, "total_adjusted")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || !before_total.is_finite() + || !before_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid".to_string(), + )); + } let mut after_recharge = before_recharge; let mut after_gift = before_gift; apply_admin_balance_adjustment( @@ -1984,6 +3378,17 @@ WHERE id = ? &mut after_recharge, &mut after_gift, ); + let after_total = after_recharge + after_gift; + let after_total_adjusted = before_total_adjusted + input.amount_usd; + if !after_recharge.is_finite() + || !after_gift.is_finite() + || !after_total.is_finite() + || !after_total_adjusted.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow during admin adjustment".to_string(), + )); + } sqlx::query( r#" @@ -2026,7 +3431,7 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, .bind(&input.wallet_id) .bind(input.amount_usd) .bind(before_total) - .bind(after_recharge + after_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -2049,7 +3454,7 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, reason_code: "adjust_admin".to_string(), amount: input.amount_usd, balance_before: before_total, - balance_after: after_recharge + after_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -2067,10 +3472,21 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, async fn create_manual_wallet_recharge( &self, - input: CreateManualWalletRechargeInput, + mut input: CreateManualWalletRechargeInput, ) -> Result, DataLayerError> { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "manual recharge amount must be finite and positive".to_string(), + )); + } let now = current_unix_secs_i64(); - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; let Some(wallet_row) = sqlite_wallet_by_id_optional(&mut tx, &input.wallet_id).await? else { tx.commit().await.map_sql_err()?; @@ -2079,6 +3495,14 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?, let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; + let before_total_recharged = sqlite_real(&wallet_row, "total_recharged")?; + let (after_recharge, after_total_recharged) = validate_manual_wallet_recharge( + input.amount_usd, + before_recharge, + before_gift, + before_total_recharged, + ) + .map_err(DataLayerError::InvalidInput)?; let user_id: Option = get(&wallet_row, "user_id")?; let order_id = uuid::Uuid::new_v4().to_string(); let gateway_response = json_string( @@ -2115,18 +3539,17 @@ VALUES (?, ?, ?, ?, ?, 0, ?, ?, 'credited', ?, ?, ?, ?) .await .map_sql_err()?; - let after_recharge = before_recharge + input.amount_usd; sqlx::query( r#" UPDATE wallets SET balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(input.amount_usd) + .bind(after_total_recharged) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) @@ -2193,7 +3616,11 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) DataLayerError, > { let now = current_unix_secs_i64(); - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; let Some(refund_row) = sqlite_refund_by_id_and_wallet(&mut tx, &input.refund_id, &input.wallet_id).await? else { @@ -2201,6 +3628,12 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) return Ok(WalletMutationOutcome::NotFound); }; let refund = map_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if !matches!(refund.status.as_str(), "approved" | "pending_approval") { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( @@ -2217,9 +3650,28 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) }; let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; - let before_total = before_recharge + before_gift; + let before_total_refunded = sqlite_real(&wallet_row, "total_refunded")?; let amount_usd = refund.amount_usd; let after_recharge = before_recharge - amount_usd; + let before_total = before_recharge + before_gift; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded + amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } if after_recharge < 0.0 { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( @@ -2236,45 +3688,78 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?) "payment order not found".to_string(), )); }; - let refundable_amount = sqlite_real(&order_row, "refundable_amount_usd")?; - if amount_usd > refundable_amount { + let order_wallet_id: String = get(&order_row, "wallet_id")?; + let order_status: String = get(&order_row, "status")?; + if order_wallet_id != input.wallet_id || order_status != "credited" { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( - "refund amount exceeds refundable amount".to_string(), + "payment order is not refundable for this wallet".to_string(), )); } - sqlx::query( + let order_amount = sqlite_real(&order_row, "amount_usd")?; + let refunded_before = sqlite_real(&order_row, "refunded_amount_usd")?; + let refundable_before = sqlite_real(&order_row, "refundable_amount_usd")?; + let refunded_after = refunded_before + amount_usd; + let refundable_after = refundable_before - amount_usd; + if !payment_order_refund_amounts_are_consistent( + order_amount, + refunded_before, + refundable_before, + ) || amount_usd > refundable_before + || !refunded_after.is_finite() + || refunded_after < 0.0 + || refunded_after > order_amount + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order_amount + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + let result = sqlx::query( r#" UPDATE payment_orders -SET refunded_amount_usd = refunded_amount_usd + ?, - refundable_amount_usd = refundable_amount_usd - ? +SET refunded_amount_usd = ?, + refundable_amount_usd = ? WHERE id = ? "#, ) - .bind(amount_usd) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .bind(payment_order_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund processing".to_string(), + )); + } } - sqlx::query( + let result = sqlx::query( r#" UPDATE wallets SET balance = ?, - total_refunded = total_refunded + ?, + total_refunded = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "wallet disappeared during refund processing".to_string(), + )); + } let wallet = map_wallet_row(&sqlite_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let transaction_id = uuid::Uuid::new_v4().to_string(); @@ -2304,15 +3789,17 @@ VALUES (?, ?, 'refund', 'refund_out', ?, ?, ?, ?, ?, ?, ?, 'refund_request', ?, .await .map_sql_err()?; - sqlx::query( + let refund_update = sqlx::query( r#" -UPDATE refund_requests + UPDATE refund_requests SET status = 'processing', approved_by = ?, processed_by = ?, processed_at = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? + AND wallet_id = ? + AND status IN ('approved', 'pending_approval') "#, ) .bind(input.operator_id.as_deref()) @@ -2324,6 +3811,11 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if refund_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during refund processing".to_string(), + )); + } let refund = map_refund_row(&sqlite_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; Ok(WalletMutationOutcome::Applied(( @@ -2357,36 +3849,68 @@ WHERE id = ? AND wallet_id = ? input: CompleteAdminWalletRefundInput, ) -> Result, DataLayerError> { let now = current_unix_secs_i64(); - let payout_proof = input - .payout_proof - .as_ref() - .map(|value| json_string(value, "refund_requests.payout_proof")) - .transpose()?; - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; let Some(current_refund) = sqlite_refund_by_id_and_wallet(&mut tx, &input.refund_id, &input.wallet_id).await? else { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::NotFound); }; - let status: String = get(¤t_refund, "status")?; - if status != "processing" { + let refund = map_refund_row(¤t_refund)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let (Some(existing_id), Some(incoming_id)) = ( + refund.gateway_refund_id.as_deref(), + input.gateway_refund_id.as_deref(), + ) { + if existing_id != incoming_id { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if refund.status == "succeeded" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(refund)); + } + if refund.status != "processing" { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( "refund status must be processing before completion".to_string(), )); } + // Preserve a processing proof for ordinary replays, but allow an + // explicit successful gateway proof to upgrade it at completion. + let selected_payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| refund.payout_proof.is_none() || wallet_refund_proof_is_success(proof)) + .cloned() + .or_else(|| refund.payout_proof.clone()); + let payout_proof = selected_payout_proof + .as_ref() + .map(|value| json_string(value, "refund_requests.payout_proof")) + .transpose()?; - sqlx::query( + let refund_update = sqlx::query( r#" UPDATE refund_requests SET status = 'succeeded', - gateway_refund_id = ?, - payout_reference = ?, + gateway_refund_id = COALESCE(gateway_refund_id, ?), + payout_reference = COALESCE(payout_reference, ?), payout_proof = ?, completed_at = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' "#, ) .bind(input.gateway_refund_id.as_deref()) @@ -2399,11 +3923,110 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if refund_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during refund completion".to_string(), + )); + } let refund = map_refund_row(&sqlite_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; Ok(WalletMutationOutcome::Applied(refund)) } + async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result, DataLayerError> { + if input.gateway_refund_id.trim().is_empty() || input.gateway_refund_id.len() > 128 { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier is invalid".to_string(), + )); + } + if input + .payout_proof + .as_ref() + .is_some_and(|proof| !proof.is_object()) + { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund proof must be an object".to_string(), + )); + } + let now = current_unix_secs_i64(); + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; + let Some(current_row) = + sqlite_refund_by_id_and_wallet(&mut tx, &input.refund_id, &input.wallet_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::NotFound); + }; + let current = map_refund_row(¤t_row)?; + if !current.amount_usd.is_finite() || current.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let Some(existing_id) = current.gateway_refund_id.as_deref() { + if existing_id != input.gateway_refund_id { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if current.status == "succeeded" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Applied(current)); + } + if current.status != "processing" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund status must be processing before gateway update".to_string(), + )); + } + // Do not overwrite durable processing evidence with an arbitrary + // replay; only a terminal success proof may replace it. + let selected_payout_proof = input + .payout_proof + .as_ref() + .filter(|proof| current.payout_proof.is_none() || wallet_refund_proof_is_success(proof)) + .cloned() + .or_else(|| current.payout_proof.clone()); + let proof = selected_payout_proof + .as_ref() + .map(|value| json_string(value, "refund_requests.payout_proof")) + .transpose()?; + let gateway_update = sqlx::query( + r#" +UPDATE refund_requests +SET gateway_refund_id = COALESCE(gateway_refund_id, ?), + payout_proof = ?, + updated_at = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' +"#, + ) + .bind(&input.gateway_refund_id) + .bind(proof.as_deref()) + .bind(now) + .bind(&input.refund_id) + .bind(&input.wallet_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + if gateway_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during gateway evidence update".to_string(), + )); + } + let updated = map_refund_row(&sqlite_refund_by_id(&mut tx, &input.refund_id).await?)?; + tx.commit().await.map_sql_err()?; + Ok(WalletMutationOutcome::Applied(updated)) + } + async fn fail_admin_wallet_refund( &self, input: FailAdminWalletRefundInput, @@ -2416,7 +4039,11 @@ WHERE id = ? AND wallet_id = ? DataLayerError, > { let now = current_unix_secs_i64(); - let mut tx = self.pool.begin().await.map_sql_err()?; + let mut tx = self + .pool + .begin_with("BEGIN IMMEDIATE") + .await + .map_sql_err()?; let Some(refund_row) = sqlite_refund_by_id_and_wallet(&mut tx, &input.refund_id, &input.wallet_id).await? else { @@ -2424,15 +4051,21 @@ WHERE id = ? AND wallet_id = ? return Ok(WalletMutationOutcome::NotFound); }; let refund = map_refund_row(&refund_row)?; + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } if matches!(refund.status.as_str(), "pending_approval" | "approved") { - sqlx::query( + let refund_update = sqlx::query( r#" UPDATE refund_requests SET status = 'failed', failure_reason = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status IN ('pending_approval', 'approved') "#, ) .bind(&input.reason) @@ -2442,6 +4075,11 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if refund_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during refund failure".to_string(), + )); + } let wallet = map_wallet_row(&sqlite_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let refund = map_refund_row(&sqlite_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; @@ -2456,6 +4094,22 @@ WHERE id = ? AND wallet_id = ? ))); } + // Only an explicitly offline payout can be released without external + // settlement evidence. An original-channel refund may still be in + // flight between the provider request and the evidence update. + if refund.gateway_refund_id.is_some() + || refund.payout_proof.is_some() + || !refund + .refund_mode + .trim() + .eq_ignore_ascii_case("offline_payout") + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "cannot fail refund while gateway settlement is processing".to_string(), + )); + } + let Some(wallet_row) = sqlite_wallet_by_id_optional(&mut tx, &input.wallet_id).await? else { tx.commit().await.map_sql_err()?; @@ -2466,26 +4120,95 @@ WHERE id = ? AND wallet_id = ? let amount_usd = refund.amount_usd; let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; + let before_total_refunded = sqlite_real(&wallet_row, "total_refunded")?; let before_total = before_recharge + before_gift; let after_recharge = before_recharge + amount_usd; + let after_total = after_recharge + before_gift; + let after_total_refunded = before_total_refunded - amount_usd; + if !before_recharge.is_finite() + || before_recharge < 0.0 + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_refunded.is_finite() + || before_total_refunded < 0.0 + || before_total_refunded < amount_usd + || !before_total.is_finite() + || !after_recharge.is_finite() + || !after_total.is_finite() + || !after_total_refunded.is_finite() + || after_total_refunded < 0.0 + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid for refund recovery".to_string(), + )); + } - sqlx::query( + let mut order_amounts = None; + if let Some(payment_order_id) = refund.payment_order_id.as_deref() { + let Some(order_row) = + sqlite_payment_order_by_id_optional(&mut tx, payment_order_id).await? + else { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order not found".to_string(), + )); + }; + let order = map_payment_order_row(&order_row)?; + if order.wallet_id != input.wallet_id || order.status != "credited" { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order is not refundable for this wallet".to_string(), + )); + } + let refunded_before = order.refunded_amount_usd; + let refundable_before = order.refundable_amount_usd; + let refunded_after = refunded_before - amount_usd; + let refundable_after = refundable_before + amount_usd; + if !payment_order_refund_amounts_are_consistent( + order.amount_usd, + refunded_before, + refundable_before, + ) || refunded_before < amount_usd + || !refunded_after.is_finite() + || refunded_after < 0.0 + || !refundable_after.is_finite() + || refundable_after < 0.0 + || refundable_after > order.amount_usd + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order refund amounts are invalid".to_string(), + )); + } + order_amounts = Some(( + payment_order_id.to_string(), + refunded_after, + refundable_after, + )); + } + + let wallet_update = sqlx::query( r#" UPDATE wallets SET balance = ?, - total_refunded = MAX(total_refunded - ?, 0), + total_refunded = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) - .bind(amount_usd) + .bind(after_total_refunded) .bind(now) .bind(&input.wallet_id) .execute(&mut *tx) .await .map_sql_err()?; - let wallet = map_wallet_row(&sqlite_wallet_by_id(&mut tx, &input.wallet_id).await?)?; + if wallet_update.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "wallet disappeared during refund recovery".to_string(), + )); + } let transaction_id = uuid::Uuid::new_v4().to_string(); sqlx::query( @@ -2502,7 +4225,7 @@ VALUES (?, ?, 'refund', 'refund_revert', ?, ?, ?, ?, ?, ?, ?, 'refund_request', .bind(&input.wallet_id) .bind(amount_usd) .bind(before_total) - .bind(after_recharge + before_gift) + .bind(after_total) .bind(before_recharge) .bind(after_recharge) .bind(before_gift) @@ -2514,30 +4237,35 @@ VALUES (?, ?, 'refund', 'refund_revert', ?, ?, ?, ?, ?, ?, ?, 'refund_request', .await .map_sql_err()?; - if let Some(payment_order_id) = refund.payment_order_id.as_deref() { - sqlx::query( + if let Some((payment_order_id, refunded_after, refundable_after)) = order_amounts { + let result = sqlx::query( r#" UPDATE payment_orders -SET refunded_amount_usd = refunded_amount_usd - ?, - refundable_amount_usd = refundable_amount_usd + ? +SET refunded_amount_usd = ?, + refundable_amount_usd = ? WHERE id = ? "#, ) - .bind(amount_usd) - .bind(amount_usd) + .bind(refunded_after) + .bind(refundable_after) .bind(payment_order_id) .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "payment order disappeared during refund recovery".to_string(), + )); + } } - sqlx::query( + let result = sqlx::query( r#" UPDATE refund_requests SET status = 'failed', failure_reason = ?, updated_at = ? -WHERE id = ? AND wallet_id = ? +WHERE id = ? AND wallet_id = ? AND status = 'processing' "#, ) .bind(&input.reason) @@ -2547,6 +4275,12 @@ WHERE id = ? AND wallet_id = ? .execute(&mut *tx) .await .map_sql_err()?; + if result.rows_affected() != 1 { + return Err(DataLayerError::UnexpectedValue( + "refund status changed during recovery".to_string(), + )); + } + let wallet = map_wallet_row(&sqlite_wallet_by_id(&mut tx, &input.wallet_id).await?)?; let refund = map_refund_row(&sqlite_refund_by_id(&mut tx, &input.refund_id).await?)?; tx.commit().await.map_sql_err()?; Ok(WalletMutationOutcome::Applied(( @@ -2559,7 +4293,7 @@ WHERE id = ? AND wallet_id = ? reason_code: "refund_revert".to_string(), amount: amount_usd, balance_before: before_total, - balance_after: after_recharge + before_gift, + balance_after: after_total, recharge_balance_before: before_recharge, recharge_balance_after: after_recharge, gift_balance_before: before_gift, @@ -2695,15 +4429,31 @@ WHERE id = ? AND wallet_id = ? } if order .expires_at_unix_secs - .is_some_and(|value| value < now as u64) + .is_some_and(|value| value <= now as u64) { tx.commit().await.map_sql_err()?; return Ok(WalletMutationOutcome::Invalid( "payment order expired".to_string(), )); } - let order_kind: String = get(&order_row, "order_kind")?; + let order_payment_provider: Option = get(&order_row, "payment_provider")?; + let order_payment_channel: Option = get(&order_row, "payment_channel")?; + if validate_payment_order_credit_amounts( + &order_kind, + &order.payment_method, + order_payment_provider.as_deref(), + order_payment_channel.as_deref(), + order.amount_usd, + order.pay_amount, + ) + .is_err() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "payment order amount is invalid".to_string(), + )); + } if order_kind == "plan_purchase" { let order_user_id: Option = get(&order_row, "user_id")?; let Some(user_id) = order_user_id else { @@ -2793,7 +4543,7 @@ VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) .bind(&plan_id) .bind(&input.order_id) .bind(now) - .bind(plan_expires_at_unix(&snapshot, now)) + .bind(plan_expires_at_unix(&snapshot, now)?) .bind(json_string( &entitlements, "user_plan_entitlements.entitlements_snapshot", @@ -2889,6 +4639,20 @@ WHERE id = ? let before_recharge = sqlite_real(&wallet_row, "balance")?; let before_gift = sqlite_real(&wallet_row, "gift_balance")?; + let total_recharged = sqlite_real(&wallet_row, "total_recharged")?; + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + || !(total_recharged + order.amount_usd).is_finite() + || !(before_recharge + before_gift + order.amount_usd).is_finite() + { + tx.commit().await.map_sql_err()?; + return Ok(WalletMutationOutcome::Invalid( + "wallet balance is invalid".to_string(), + )); + } let before_total = before_recharge + before_gift; let after_recharge = before_recharge + order.amount_usd; sqlx::query( @@ -2993,9 +4757,19 @@ WHERE id = ? &self, input: CreateAdminRedeemCodeBatchInput, ) -> Result { + validate_admin_redeem_code_batch_input(&input).map_err(DataLayerError::InvalidInput)?; let now = current_unix_secs_i64(); let batch_id = uuid::Uuid::new_v4().to_string(); - let expires_at = input.expires_at_unix_secs.map(|value| value as i64); + let expires_at = input + .expires_at_unix_secs + .map(|value| { + i64::try_from(value).map_err(|_| { + DataLayerError::InvalidInput( + "redeem code batch expires_at overflow".to_string(), + ) + }) + }) + .transpose()?; let mut tx = self.pool.begin().await.map_sql_err()?; sqlx::query( @@ -3327,7 +5101,6 @@ LIMIT 1 let batch_name: String = get(&code_row, "batch_name")?; let balance_bucket: String = get(&code_row, "balance_bucket")?; let amount_usd = sqlite_real(&code_row, "amount_usd")?; - let credits_recharge_balance = redeem_code_credits_recharge_balance(&balance_bucket); let wallet_row = sqlite_wallet_by_user_id(&mut tx, &input.user_id).await?; let wallet_id = if let Some(row) = wallet_row.as_ref() { @@ -3341,14 +5114,16 @@ LIMIT 1 uuid::Uuid::new_v4().to_string() }; - let (before_recharge, before_gift) = if let Some(row) = wallet_row.as_ref() { - ( - sqlite_real(row, "balance")?, - sqlite_real(row, "gift_balance")?, - ) - } else { - sqlx::query( - r#" + let (before_recharge, before_gift, before_total_recharged) = + if let Some(row) = wallet_row.as_ref() { + ( + sqlite_real(row, "balance")?, + sqlite_real(row, "gift_balance")?, + sqlite_real(row, "total_recharged")?, + ) + } else { + sqlx::query( + r#" INSERT INTO wallets ( id, user_id, balance, gift_balance, limit_mode, currency, status, total_recharged, total_consumed, total_refunded, total_adjusted, @@ -3356,39 +5131,37 @@ INSERT INTO wallets ( ) VALUES (?, ?, 0.0, 0.0, 'finite', 'USD', 'active', 0.0, 0.0, 0.0, 0.0, ?, ?) "#, - ) - .bind(&wallet_id) - .bind(&input.user_id) - .bind(now) - .bind(now) - .execute(&mut *tx) - .await - .map_sql_err()?; - (0.0, 0.0) - }; - let after_recharge = if credits_recharge_balance { - before_recharge + amount_usd - } else { - before_recharge - }; - let after_gift = if credits_recharge_balance { - before_gift - } else { - before_gift + amount_usd - }; + ) + .bind(&wallet_id) + .bind(&input.user_id) + .bind(now) + .bind(now) + .execute(&mut *tx) + .await + .map_sql_err()?; + (0.0, 0.0, 0.0) + }; + let (after_recharge, after_gift, after_total_recharged) = validate_redeem_wallet_credit( + &balance_bucket, + amount_usd, + before_recharge, + before_gift, + before_total_recharged, + ) + .map_err(DataLayerError::UnexpectedValue)?; sqlx::query( r#" UPDATE wallets SET balance = ?, gift_balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(after_recharge) .bind(after_gift) - .bind(amount_usd) + .bind(after_total_recharged) .bind(now) .bind(&wallet_id) .execute(&mut *tx) @@ -3699,34 +5472,14 @@ fn plan_purchase_limit_scope(snapshot: &serde_json::Value) -> &str { } } -fn plan_replacement_entitlement_types(snapshot: &serde_json::Value) -> Vec<&'static str> { - let entitlements = plan_entitlements_snapshot(snapshot); - let mut kinds = Vec::new(); - if entitlement_snapshot_has_type(&entitlements, "daily_quota") { - kinds.push("daily_quota"); - } - if entitlement_snapshot_has_type(&entitlements, "membership_group") { - kinds.push("membership_group"); - } - kinds -} - -fn entitlement_snapshot_has_type(snapshot: &serde_json::Value, entitlement_type: &str) -> bool { - snapshot.as_array().is_some_and(|items| { - items - .iter() - .any(|item| item.get("type").and_then(|value| value.as_str()) == Some(entitlement_type)) - }) -} - async fn replace_matching_plan_entitlements_sqlite( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, user_id: &str, snapshot: &serde_json::Value, now: i64, ) -> Result<(), DataLayerError> { - let replacement_types = plan_replacement_entitlement_types(snapshot); - if replacement_types.is_empty() { + let incoming_entitlements = plan_entitlements_snapshot(snapshot); + if !entitlements_have_replacement_selector(&incoming_entitlements) { return Ok(()); } @@ -3751,9 +5504,8 @@ WHERE user_id = ? "user_plan_entitlements.entitlements_snapshot", )? .unwrap_or_else(|| serde_json::json!([])); - let should_replace = replacement_types - .iter() - .any(|kind| entitlement_snapshot_has_type(&entitlements, kind)); + let should_replace = + entitlements_should_replace_existing(&incoming_entitlements, &entitlements); if !should_replace { continue; } @@ -3781,22 +5533,18 @@ WHERE id = ? Ok(()) } -fn plan_expires_at_unix(snapshot: &serde_json::Value, starts_at_unix_secs: i64) -> i64 { - let duration_value = snapshot - .get("duration_value") - .and_then(|value| value.as_i64()) - .unwrap_or(1) - .max(1); - let days = match snapshot - .get("duration_unit") - .and_then(|value| value.as_str()) - .unwrap_or("month") - { - "day" | "custom" => duration_value, - "year" => 365 * duration_value, - _ => 30 * duration_value, - }; - starts_at_unix_secs.saturating_add(days.saturating_mul(86_400)) +fn plan_expires_at_unix( + snapshot: &serde_json::Value, + starts_at_unix_secs: i64, +) -> Result { + let days = + checked_plan_duration_days_from_snapshot(snapshot).map_err(DataLayerError::InvalidInput)?; + let seconds = days.checked_mul(86_400).ok_or_else(|| { + DataLayerError::InvalidInput("plan duration exceeds the supported range".to_string()) + })?; + starts_at_unix_secs.checked_add(seconds).ok_or_else(|| { + DataLayerError::InvalidInput("plan expiration exceeds the supported range".to_string()) + }) } async fn apply_plan_wallet_credit_sqlite( @@ -3807,11 +5555,16 @@ async fn apply_plan_wallet_credit_sqlite( entitlements: &serde_json::Value, now: i64, ) -> Result<(), DataLayerError> { + validate_plan_wallet_credit_entitlements(entitlements).map_err(DataLayerError::InvalidInput)?; let credits = entitlements .as_array() .into_iter() .flatten() - .filter(|item| item.get("type").and_then(|value| value.as_str()) == Some("wallet_credit")) + .filter(|item| { + item.get("type") + .and_then(|value| value.as_str()) + .is_some_and(|value| value.eq_ignore_ascii_case("wallet_credit")) + }) .filter_map(|item| { let amount = item.get("amount_usd").and_then(|value| value.as_f64())?; if amount <= 0.0 || !amount.is_finite() { @@ -3821,6 +5574,7 @@ async fn apply_plan_wallet_credit_sqlite( .get("balance_bucket") .and_then(|value| value.as_str()) .unwrap_or("gift") + .trim() .to_ascii_lowercase(); Some((amount, bucket)) }) @@ -3829,7 +5583,7 @@ async fn apply_plan_wallet_credit_sqlite( return Ok(()); } let Some(wallet_row) = - sqlx::query("SELECT id, status, balance, gift_balance FROM wallets WHERE id = ? LIMIT 1") + sqlx::query("SELECT id, status, balance, gift_balance, total_recharged FROM wallets WHERE id = ? LIMIT 1") .bind(wallet_id) .fetch_optional(&mut **tx) .await @@ -3847,6 +5601,17 @@ async fn apply_plan_wallet_credit_sqlite( } let mut recharge_balance = sqlite_real(&wallet_row, "balance")?; let mut gift_balance = sqlite_real(&wallet_row, "gift_balance")?; + let mut total_recharged = sqlite_real(&wallet_row, "total_recharged")?; + if !recharge_balance.is_finite() + || !gift_balance.is_finite() + || gift_balance < 0.0 + || !total_recharged.is_finite() + || total_recharged < 0.0 + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance is invalid for plan wallet_credit".to_string(), + )); + } for (amount, bucket) in credits { let before_recharge = recharge_balance; let before_gift = gift_balance; @@ -3854,23 +5619,34 @@ async fn apply_plan_wallet_credit_sqlite( let credits_recharge = bucket == "recharge"; if credits_recharge { recharge_balance += amount; + total_recharged += amount; } else { gift_balance += amount; } let after_total = recharge_balance + gift_balance; + if !before_total.is_finite() + || !recharge_balance.is_finite() + || !gift_balance.is_finite() + || !total_recharged.is_finite() + || !after_total.is_finite() + { + return Err(DataLayerError::UnexpectedValue( + "wallet balance overflow for plan wallet_credit".to_string(), + )); + } sqlx::query( r#" UPDATE wallets SET balance = ?, gift_balance = ?, - total_recharged = total_recharged + ?, + total_recharged = ?, updated_at = ? WHERE id = ? "#, ) .bind(recharge_balance) .bind(gift_balance) - .bind(if credits_recharge { amount } else { 0.0 }) + .bind(total_recharged) .bind(now) .bind(wallet_id) .execute(&mut **tx) @@ -3905,16 +5681,6 @@ VALUES (?, ?, 'recharge', 'plan_wallet_credit', ?, ?, ?, ?, ?, ?, ?, 'payment_or Ok(()) } -fn default_refund_mode_for_payment_method(payment_method: &str) -> &'static str { - if matches!( - payment_method, - "admin_manual" | "card_recharge" | "card_code" | "gift_code" - ) { - return "offline_payout"; - } - "original_channel" -} - fn payment_gateway_response_map( value: Option, ) -> serde_json::Map { @@ -4115,16 +5881,67 @@ async fn sqlite_payment_order_by_order_no( async fn sqlite_payment_order_by_gateway_order_id( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, + payment_method: &str, gateway_order_id: &str, ) -> Result, DataLayerError> { - let sql = payment_order_select_sql("WHERE gateway_order_id = ? LIMIT 1"); + let sql = payment_order_select_sql("WHERE payment_method = ? AND gateway_order_id = ? LIMIT 1"); sqlx::query(&sql) + .bind(payment_method) .bind(gateway_order_id) .fetch_optional(&mut **tx) .await .map_sql_err() } +async fn sqlite_bind_payment_gateway_order_id( + tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, + order_id: &str, + gateway_order_id: &str, +) -> Result { + sqlx::query("SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + let bind_result = sqlx::query("UPDATE payment_orders SET gateway_order_id = ? WHERE id = ?") + .bind(gateway_order_id) + .bind(order_id) + .execute(&mut **tx) + .await; + match bind_result { + Ok(_) => { + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + Ok(true) + } + Err(error) + if error + .as_database_error() + .is_some_and(|database_error| database_error.is_unique_violation()) => + { + sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await + .map_sql_err()?; + Ok(false) + } + Err(error) => { + let _ = sqlx::query("ROLLBACK TO SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + let _ = sqlx::query("RELEASE SAVEPOINT payment_gateway_order_binding") + .execute(&mut **tx) + .await; + Err(DataLayerError::sql(error)) + } + } +} + fn refund_select_sql(where_clause: &str) -> String { format!( r#" @@ -4258,7 +6075,7 @@ async fn update_sqlite_payment_callback_failure( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, callback_id: &str, input: &ProcessPaymentCallbackInput, - payload: &str, + _payload: &str, error: &str, ) -> Result<(), DataLayerError> { sqlx::query( @@ -4268,20 +6085,15 @@ SET signature_valid = ?, status = 'failed', error_message = ?, payload_hash = ?, - payload = ?, - processed_at = ?, - order_no = COALESCE(?, order_no), - gateway_order_id = COALESCE(?, gateway_order_id) + payload = NULL, + processed_at = ? WHERE id = ? "#, ) .bind(sqlite_bool(input.signature_valid)) .bind(error) .bind(&input.payload_hash) - .bind(payload) .bind(current_unix_secs_i64()) - .bind(input.order_no.as_deref()) - .bind(input.gateway_order_id.as_deref()) .bind(callback_id) .execute(&mut **tx) .await @@ -4293,7 +6105,7 @@ async fn mark_sqlite_payment_callback_processed( tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, callback_id: &str, input: &ProcessPaymentCallbackInput, - payload: &str, + _payload: &str, order_id: &str, order_no: &str, ) -> Result<(), DataLayerError> { @@ -4305,7 +6117,7 @@ SET payment_order_id = ?, status = 'processed', error_message = NULL, payload_hash = ?, - payload = ?, + payload = NULL, processed_at = ?, order_no = ?, gateway_order_id = COALESCE(?, gateway_order_id) @@ -4314,7 +6126,6 @@ WHERE id = ? ) .bind(order_id) .bind(&input.payload_hash) - .bind(payload) .bind(current_unix_secs_i64()) .bind(order_no) .bind(input.gateway_order_id.as_deref()) @@ -4394,7 +6205,20 @@ async fn initialize_sqlite_auth_wallet( api_key_id: Option<&str>, initial_gift_usd: f64, unlimited: bool, -) -> Result, DataLayerError> { +) -> Result, DataLayerError> { + let owner = user_id + .or(api_key_id) + .filter(|value| !value.trim().is_empty()); + if owner.is_none() || (user_id.is_some() && api_key_id.is_some()) { + return Err(DataLayerError::InvalidInput( + "wallet owner must be exactly one non-empty user or API-key id".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "initial gift amount must be finite".to_string(), + )); + } let gift_amount = if unlimited { 0.0 } else { @@ -4416,8 +6240,77 @@ async fn initialize_sqlite_auth_wallet( gift_amount, now, )?; - let mut tx = pool.begin().await.map_sql_err()?; - sqlx::query( + // Keep the owner lookup and insert in one writer transaction so concurrent + // initialization retries cannot mint duplicate wallets or gift entries. + let mut tx = pool.begin_with("BEGIN IMMEDIATE").await.map_sql_err()?; + let owner_column = if user_id.is_some() { + "user_id" + } else { + "api_key_id" + }; + let owner_value = owner.expect("validated wallet owner"); + + // SQLite's baseline schema does not declare owner foreign keys on wallets. + // Validate the owner while holding the writer transaction so a wallet (and + // its initial gift journal entry) can never be created for a missing auth + // record. + if let Some(user_id) = user_id { + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + } else { + let api_key_user_id: Option = + sqlx::query_scalar("SELECT user_id FROM api_keys WHERE id = ?") + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(api_key_user_id) = api_key_user_id else { + tx.rollback().await.map_sql_err()?; + return Ok(None); + }; + let user_exists: Option = sqlx::query_scalar("SELECT id FROM users WHERE id = ?") + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if user_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + let api_key_exists: Option = + sqlx::query_scalar("SELECT id FROM api_keys WHERE id = ? AND user_id = ?") + .bind(owner_value) + .bind(&api_key_user_id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if api_key_exists.is_none() { + tx.rollback().await.map_sql_err()?; + return Ok(None); + } + } + + let existing_row = sqlx::query(&wallet_select_sql(&format!( + "WHERE {owner_column} = ? LIMIT 1" + ))) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if let Some(row) = existing_row { + let existing = map_wallet_row(&row)?; + tx.commit().await.map_sql_err()?; + return Ok(Some((existing, false))); + } + + let insert_result = sqlx::query( r#" INSERT INTO wallets ( id, user_id, api_key_id, balance, gift_balance, limit_mode, currency, @@ -4425,6 +6318,7 @@ INSERT INTO wallets ( created_at, updated_at ) VALUES (?, ?, ?, 0, ?, ?, 'USD', 'active', 0, 0, 0, ?, ?, ?) +ON CONFLICT DO NOTHING "#, ) .bind(&wallet.id) @@ -4438,6 +6332,26 @@ VALUES (?, ?, ?, 0, ?, ?, 'USD', 'active', 0, 0, 0, ?, ?, ?) .execute(&mut *tx) .await .map_sql_err()?; + if insert_result.rows_affected() == 0 { + // Another initializer may have won the owner race. If no owner row is + // visible, the generated wallet id collided with a different owner. + let row = sqlx::query(&wallet_select_sql(&format!( + "WHERE {owner_column} = ? LIMIT 1" + ))) + .bind(owner_value) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + if let Some(row) = row { + let existing = map_wallet_row(&row)?; + tx.commit().await.map_sql_err()?; + return Ok(Some((existing, false))); + } + tx.rollback().await.map_sql_err()?; + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + } if gift_amount > 0.0 { let link_id = user_id.or(api_key_id).unwrap_or_default(); let description = if api_key_id.is_some() { @@ -4467,8 +6381,34 @@ VALUES (?, ?, 'gift', 'gift_initial', ?, 0, ?, 0, 0, 0, ?, 'system_task', ?, NUL .await .map_sql_err()?; } + let row = sqlite_wallet_by_id(&mut tx, &wallet.id).await?; + let wallet = map_wallet_row(&row)?; tx.commit().await.map_sql_err()?; - Ok(Some(wallet)) + Ok(Some((wallet, true))) +} + +fn sqlite_wallet_recharge_replay_matches( + row: &SqliteRow, + wallet_id: &str, + input: &CreateWalletRechargeOrderInput, +) -> Result { + let existing_wallet_id: String = get(row, "wallet_id")?; + let pay_currency: Option = get(row, "pay_currency")?; + let payment_method: String = get(row, "payment_method")?; + let payment_provider: Option = get(row, "payment_provider")?; + let payment_channel: Option = get(row, "payment_channel")?; + Ok(wallet_recharge_replay_matches( + &existing_wallet_id, + sqlite_real(row, "amount_usd")?, + sqlite_optional_real(row, "pay_amount")?, + pay_currency.as_deref(), + sqlite_optional_real(row, "exchange_rate")?, + &payment_method, + payment_provider.as_deref(), + payment_channel.as_deref(), + wallet_id, + input, + )) } fn map_payment_order_row(row: &SqliteRow) -> Result { @@ -4484,6 +6424,8 @@ fn map_payment_order_row(row: &SqliteRow) -> Result order, + other => panic!("unexpected order creation outcome: {other:?}"), + }; + let observed = order + .gateway_response + .clone() + .expect("created order should contain a response"); + let replacement = concat!( + "aether-payment-order-stripe-client-secret-v2:", + "aether-runtime-secret-v1:gAAAAABsqlite-replacement" + ); + let input = CompareAndSwapPaymentOrderStripeClientSecretInput { + order_id: order.id.clone(), + order_no: order.order_no.clone(), + wallet_id: order.wallet_id.clone(), + user_id: order.user_id.clone(), + payment_method: order.payment_method.clone(), + payment_provider: order.payment_provider.clone(), + order_kind: order.order_kind.clone(), + gateway_order_id: order.gateway_order_id.clone(), + expected_status: order.status.clone(), + expected_expires_at_unix_secs: order.expires_at_unix_secs, + expected_gateway_response: observed, + expected_client_secret_encrypted: legacy.to_string(), + replacement_client_secret_encrypted: replacement.to_string(), + }; + + let mut foreign = input.clone(); + foreign.user_id = Some("stripe-cas-foreign-user".to_string()); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(foreign) + .await + .expect("identity mismatch should be a normal CAS miss")); + assert!(repository + .compare_and_swap_payment_order_stripe_client_secret(input.clone()) + .await + .expect("exact CAS should succeed")); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(input) + .await + .expect("stale CAS must not replace the new value")); + + let stored: String = + sqlx::query_scalar("SELECT gateway_response FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("stored gateway response should query"); + let stored: serde_json::Value = + serde_json::from_str(&stored).expect("stored response should be valid JSON"); + assert_eq!( + stored["_stripe_client_secret_encrypted"].as_str(), + Some(replacement) + ); + assert_eq!(stored["publishable_key"], "pk_test_public"); +} + +#[tokio::test] +async fn sqlite_admin_balance_adjustment_rejects_invalid_numbers_without_writes() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "invalid-adjustment-user").await; + let wallet = repository + .initialize_auth_user_wallet("invalid-adjustment-user", 0.0, false) + .await + .expect("wallet initialization should run") + .expect("wallet should exist"); + + for amount_usd in [0.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let error = repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet.id.clone(), + amount_usd, + balance_type: "recharge".to_string(), + operator_id: Some("admin-1".to_string()), + description: None, + }) + .await + .expect_err("invalid adjustment should fail before writing"); + assert!(matches!(error, DataLayerError::InvalidInput(_))); + } + + sqlx::query("UPDATE wallets SET balance = ? WHERE id = ?") + .bind(f64::MAX) + .bind(&wallet.id) + .execute(&pool) + .await + .expect("overflow fixture should update"); + let error = repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet.id.clone(), + amount_usd: f64::MAX, + balance_type: "recharge".to_string(), + operator_id: Some("admin-1".to_string()), + description: None, + }) + .await + .expect_err("overflowing adjustment should fail before writing"); + assert!(matches!(error, DataLayerError::UnexpectedValue(_))); + + let stored_balance = sqlx::query_scalar::<_, f64>("SELECT balance FROM wallets WHERE id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("wallet balance should query"); + assert_eq!(stored_balance, f64::MAX); + let adjustment_count = sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ? AND category = 'adjust'", + ) + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("adjustment ledger count should query"); + assert_eq!(adjustment_count, 0); +} + +#[tokio::test] +async fn sqlite_manual_wallet_recharge_rejects_invalid_numbers_without_writes() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + ensure_test_user(&pool, "invalid-manual-recharge-user").await; + let repository = SqliteWalletReadRepository::new(pool.clone()); + let wallet = repository + .initialize_auth_user_wallet("invalid-manual-recharge-user", 0.0, false) + .await + .expect("wallet initialization should run") + .expect("wallet should exist"); + let initial_order_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM payment_orders WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("payment order count should query"); + let initial_transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("wallet transaction count should query"); + + for (index, amount_usd) in [0.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] + .into_iter() + .enumerate() + { + let error = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: wallet.id.clone(), + amount_usd, + payment_method: "admin_manual".to_string(), + operator_id: Some("admin-invalid-recharge".to_string()), + description: None, + order_no: format!("invalid-manual-recharge-{index}"), + }) + .await + .expect_err("invalid manual recharge should fail before writing"); + assert!(matches!(error, DataLayerError::InvalidInput(_))); + } + + let unchanged = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet should query") + .expect("wallet should remain present"); + assert_eq!(unchanged, wallet); + let order_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM payment_orders WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("payment order count should query"); + let transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("wallet transaction count should query"); + assert_eq!(order_count, initial_order_count); + assert_eq!(transaction_count, initial_transaction_count); + + sqlx::query("UPDATE wallets SET balance = ?, total_recharged = ? WHERE id = ?") + .bind(f64::MAX) + .bind(f64::MAX) + .bind(&wallet.id) + .execute(&pool) + .await + .expect("overflow fixture should update"); + let overflow_fixture = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("overflow fixture should query") + .expect("wallet should remain present"); + let error = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: wallet.id.clone(), + amount_usd: f64::MAX, + payment_method: "admin_manual".to_string(), + operator_id: Some("admin-overflow-recharge".to_string()), + description: None, + order_no: "overflow-manual-recharge".to_string(), + }) + .await + .expect_err("overflowing manual recharge should fail before writing"); + assert!(matches!(error, DataLayerError::InvalidInput(_))); + + let unchanged = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet should query") + .expect("wallet should remain present"); + assert_eq!(unchanged, overflow_fixture); + let order_count_after_overflow: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM payment_orders WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("payment order count should query"); + let transaction_count_after_overflow: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?") + .bind(&wallet.id) + .fetch_one(&pool) + .await + .expect("wallet transaction count should query"); + assert_eq!(order_count_after_overflow, initial_order_count); + assert_eq!(transaction_count_after_overflow, initial_transaction_count); +} + +#[tokio::test] +async fn sqlite_wallet_initialization_rejects_missing_owners_without_writes() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + assert!(repository + .initialize_auth_user_wallet("missing-wallet-user", 5.0, false) + .await + .expect("missing user initialization should resolve") + .is_none()); + assert!(repository + .initialize_auth_api_key_wallet("missing-wallet-api-key", 5.0, false) + .await + .expect("missing api key initialization should resolve") + .is_none()); + + let wallet_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM wallets") + .fetch_one(&pool) + .await + .expect("wallet count should query"); + assert_eq!(wallet_count, 0); + let transaction_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions") + .fetch_one(&pool) + .await + .expect("transaction count should query"); + assert_eq!(transaction_count, 0); + + ensure_test_user(&pool, "wallet-owner-user").await; + sqlx::query( + "INSERT INTO api_keys (id, user_id, key_hash, created_at, updated_at) VALUES (?, ?, ?, ?, ?)", + ) + .bind("wallet-owner-api-key") + .bind("wallet-owner-user") + .bind("wallet-owner-api-key-hash") + .bind(1_i64) + .bind(1_i64) + .execute(&pool) + .await + .expect("api key should seed"); + let initialized = repository + .initialize_auth_api_key_wallet("wallet-owner-api-key", 5.0, false) + .await + .expect("valid api key initialization should resolve") + .expect("valid api key wallet should be created"); + assert_eq!( + initialized.api_key_id.as_deref(), + Some("wallet-owner-api-key") + ); + assert_eq!(initialized.gift_balance, 5.0); +} #[tokio::test] async fn sqlite_wallet_read_repository_reads_wallet_contract_views() { @@ -103,6 +466,1144 @@ async fn sqlite_wallet_read_repository_reads_wallet_contract_views() { assert_eq!(daily.total_requests, 2); } +#[tokio::test] +async fn sqlite_provisional_wallet_cleanup_is_activity_guarded() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + ensure_test_user(&pool, "provisional-user").await; + let provisional_wallet = repository + .initialize_auth_user_wallet("provisional-user", 10.0, false) + .await + .expect("wallet initialization should succeed") + .expect("provisional wallet should exist"); + assert!(repository + .delete_provisional_auth_user_wallet(&provisional_wallet.id, "provisional-user") + .await + .expect("provisional cleanup should succeed")); + assert!(repository + .find(WalletLookupKey::UserId("provisional-user")) + .await + .expect("wallet lookup should succeed") + .is_none()); + + ensure_test_user(&pool, "active-user").await; + repository + .initialize_auth_user_wallet("active-user", 10.0, false) + .await + .expect("wallet initialization should succeed"); + let active_wallet = repository + .find(WalletLookupKey::UserId("active-user")) + .await + .expect("wallet lookup should succeed") + .expect("active wallet should exist"); + sqlx::query( + r#"INSERT INTO wallet_daily_usage_ledgers ( + id, wallet_id, billing_date, billing_timezone, total_cost_usd, + total_requests, input_tokens, output_tokens, cache_creation_tokens, + cache_read_tokens, aggregated_at, created_at, updated_at + ) VALUES (?, ?, '2000-01-01', 'UTC', 1.0, 1, 1, 1, 0, 0, 1, 1, 1)"#, + ) + .bind("active-daily") + .bind(&active_wallet.id) + .execute(&pool) + .await + .expect("activity row should insert"); + assert!(!repository + .delete_provisional_auth_user_wallet(&active_wallet.id, "active-user") + .await + .expect("provisional cleanup should succeed")); + assert!(repository + .find(WalletLookupKey::UserId("active-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); +} + +#[tokio::test] +async fn sqlite_wallet_compensation_delete_is_owner_and_reference_guarded() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + ensure_test_user(&pool, "funded-compensation-user").await; + let funded_wallet = repository + .initialize_auth_user_wallet("funded-compensation-user", 0.0, false) + .await + .expect("wallet initialization should succeed") + .expect("wallet should exist"); + sqlx::query( + "UPDATE wallets SET balance = ?, total_recharged = ?, total_adjusted = ? WHERE id = ?", + ) + .bind(1.0) + .bind(1.0) + .bind(1.0) + .bind(&funded_wallet.id) + .execute(&pool) + .await + .expect("funded wallet update should succeed"); + assert!(!repository + .delete_wallet_if_unreferenced( + &funded_wallet.id, + WalletLookupKey::UserId("funded-compensation-user"), + ) + .await + .expect("funded wallet must not be deleted")); + assert!(repository + .find(WalletLookupKey::UserId("funded-compensation-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + + ensure_test_user(&pool, "compensation-user").await; + let wallet = repository + .initialize_auth_user_wallet("compensation-user", 0.0, false) + .await + .expect("wallet initialization should succeed") + .expect("wallet should exist"); + let wallet_id = wallet.id.clone(); + + // A caller with a different owner must not be able to reclaim this wallet, even when the + // wallet id is known. + assert!(!repository + .delete_wallet_if_unreferenced(&wallet_id, WalletLookupKey::ApiKeyId("different-api-key"),) + .await + .expect("owner mismatch should be handled cleanly")); + + sqlx::query( + r#" +INSERT INTO wallet_daily_usage_ledgers ( + id, wallet_id, billing_date, billing_timezone, total_cost_usd, + total_requests, input_tokens, output_tokens, cache_creation_tokens, + cache_read_tokens, aggregated_at, created_at, updated_at +) VALUES (?, ?, '2000-01-01', 'UTC', 0, 1, 0, 0, 0, 0, 1, 1, 1) +"#, + ) + .bind("compensation-daily") + .bind(&wallet_id) + .execute(&pool) + .await + .expect("usage reference should insert"); + + assert!(!repository + .delete_wallet_if_unreferenced(&wallet_id, WalletLookupKey::UserId("compensation-user"),) + .await + .expect("referenced wallet should be retained")); + assert!(repository + .find(WalletLookupKey::UserId("compensation-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + + sqlx::query("DELETE FROM wallet_daily_usage_ledgers WHERE id = ?") + .bind("compensation-daily") + .execute(&pool) + .await + .expect("usage reference should delete"); + assert!(repository + .delete_wallet_if_unreferenced(&wallet_id, WalletLookupKey::UserId("compensation-user"),) + .await + .expect("unreferenced wallet should be deleted")); + assert!(repository + .find(WalletLookupKey::UserId("compensation-user")) + .await + .expect("wallet lookup should succeed") + .is_none()); +} + +#[tokio::test] +async fn sqlite_wallet_snapshot_compensation_deletes_funded_match_only() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + ensure_test_user(&pool, "snapshot-compensation-user").await; + let wallet = repository + .initialize_auth_user_wallet("snapshot-compensation-user", 0.0, false) + .await + .expect("wallet initialization should succeed") + .expect("wallet should exist"); + sqlx::query( + "UPDATE wallets SET balance = ?, total_recharged = ?, total_adjusted = ? WHERE id = ?", + ) + .bind(12.5) + .bind(12.5) + .bind(0.0) + .bind(&wallet.id) + .execute(&pool) + .await + .expect("wallet funding fixture should update"); + let expected = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet lookup should succeed") + .expect("funded wallet should exist"); + assert!(repository + .delete_wallet_if_snapshot_matches_and_unreferenced( + &expected, + WalletLookupKey::UserId("snapshot-compensation-user"), + ) + .await + .expect("matching snapshot delete should succeed")); + + ensure_test_user(&pool, "snapshot-compensation-changed").await; + let wallet = repository + .initialize_auth_user_wallet("snapshot-compensation-changed", 0.0, false) + .await + .expect("wallet initialization should succeed") + .expect("wallet should exist"); + let expected = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet lookup should succeed") + .expect("wallet should exist"); + sqlx::query("UPDATE wallets SET balance = 1.0 WHERE id = ?") + .bind(&wallet.id) + .execute(&pool) + .await + .expect("concurrent wallet change fixture should update"); + assert!(!repository + .delete_wallet_if_snapshot_matches_and_unreferenced( + &expected, + WalletLookupKey::UserId("snapshot-compensation-changed"), + ) + .await + .expect("changed snapshot should be retained")); + assert!(repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet lookup should succeed") + .is_some()); + + ensure_test_user(&pool, "snapshot-compensation-referenced").await; + let wallet = repository + .initialize_auth_user_wallet("snapshot-compensation-referenced", 0.0, false) + .await + .expect("wallet initialization should succeed") + .expect("wallet should exist"); + let expected = repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet lookup should succeed") + .expect("wallet should exist"); + sqlx::query( + r#" +INSERT INTO wallet_daily_usage_ledgers ( + id, wallet_id, billing_date, billing_timezone, total_cost_usd, + total_requests, input_tokens, output_tokens, cache_creation_tokens, + cache_read_tokens, aggregated_at, created_at, updated_at +) VALUES (?, ?, '2000-01-01', 'UTC', 0, 1, 0, 0, 0, 0, 1, 1, 1) +"#, + ) + .bind("snapshot-compensation-reference") + .bind(&wallet.id) + .execute(&pool) + .await + .expect("usage reference should insert"); + assert!(!repository + .delete_wallet_if_snapshot_matches_and_unreferenced( + &expected, + WalletLookupKey::UserId("snapshot-compensation-referenced"), + ) + .await + .expect("referenced snapshot should be retained")); + assert!(repository + .find(WalletLookupKey::WalletId(&wallet.id)) + .await + .expect("wallet lookup should succeed") + .is_some()); +} + +#[tokio::test] +async fn sqlite_wallet_refund_rejects_invalid_amounts_and_foreign_wallets() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + ensure_test_user(&pool, "refund-security-user").await; + let wallet = repository + .initialize_auth_user_wallet("refund-security-user", 0.0, false) + .await + .expect("setup wallet initialization should run") + .expect("setup wallet should exist"); + let (wallet, _) = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: wallet.id, + amount_usd: 10.0, + payment_method: "admin_manual".to_string(), + operator_id: Some("admin-1".to_string()), + description: Some("refund security setup".to_string()), + order_no: "refund-security-recharge".to_string(), + }) + .await + .expect("manual recharge should run") + .expect("setup wallet should exist"); + + for (index, amount_usd) in [0.0, -1.0, f64::NAN, f64::INFINITY].into_iter().enumerate() { + let outcome = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: wallet.id.clone(), + user_id: wallet + .user_id + .clone() + .expect("setup wallet should have an owner"), + amount_usd, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: None, + idempotency_key: Some(format!("refund-security-invalid-{index}")), + refund_no: format!("refund-security-invalid-{index}"), + }) + .await + .expect("invalid refund should be rejected cleanly"); + assert!(matches!( + outcome, + CreateWalletRefundRequestOutcome::InvalidInput(_) + )); + } + + let foreign_outcome = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: wallet.id, + user_id: "different-user".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: None, + idempotency_key: Some("refund-security-foreign-wallet".to_string()), + refund_no: "refund-security-foreign-wallet".to_string(), + }) + .await + .expect("foreign wallet refund should be rejected cleanly"); + assert!(matches!( + foreign_outcome, + CreateWalletRefundRequestOutcome::WalletMissing + )); + + let refund_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM refund_requests") + .fetch_one(&pool) + .await + .expect("refund count should query"); + assert_eq!(refund_count, 0); +} + +#[tokio::test] +async fn sqlite_wallet_refund_rejects_invalid_reserved_amounts() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + ensure_test_user(&pool, "refund-reservation-user").await; + let wallet = repository + .initialize_auth_user_wallet("refund-reservation-user", 0.0, false) + .await + .expect("reservation wallet initialization should run") + .expect("reservation wallet should exist"); + let (wallet, order) = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: wallet.id, + amount_usd: 10.0, + payment_method: "admin_manual".to_string(), + operator_id: Some("reservation-admin".to_string()), + description: Some("reservation setup".to_string()), + order_no: "reservation-order".to_string(), + }) + .await + .expect("reservation recharge should run") + .expect("reservation wallet should still exist"); + + for (id, status, amount_usd) in [ + ("reservation-valid", "pending_approval", 2.0), + ("reservation-negative", "approved", -100.0), + ("reservation-infinite", "pending_approval", f64::INFINITY), + ] { + sqlx::query( + r#" +INSERT INTO refund_requests ( + id, refund_no, wallet_id, user_id, payment_order_id, source_type, + refund_mode, amount_usd, status, created_at, updated_at +) +VALUES (?, ?, ?, ?, ?, 'payment_order', 'offline_payout', ?, ?, 1, 1) +"#, + ) + .bind(id) + .bind(format!("{id}-no")) + .bind(&wallet.id) + .bind("refund-reservation-user") + .bind(&order.id) + .bind(amount_usd) + .bind(status) + .execute(&pool) + .await + .expect("corrupt reservation row should insert"); + } + + let outcome = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: wallet.id, + user_id: "refund-reservation-user".to_string(), + amount_usd: 8.0, + payment_order_id: Some(order.id), + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("reservation regression".to_string()), + idempotency_key: Some("reservation-regression-idempotency".to_string()), + refund_no: "reservation-regression-refund".to_string(), + }) + .await + .expect("reservation request should run"); + assert!(matches!( + outcome, + CreateWalletRefundRequestOutcome::InvalidInput(_) + )); + let refund_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM refund_requests WHERE idempotency_key = ?") + .bind("reservation-regression-idempotency") + .fetch_one(&pool) + .await + .expect("refund count should query"); + assert_eq!(refund_count, 0); +} + +#[derive(Debug, PartialEq)] +struct RefundMutationSnapshot { + wallet_balance: f64, + wallet_total_refunded: f64, + order_refunded_amount: f64, + order_refundable_amount: f64, + refund_status: String, + refund_failure_reason: Option, + refund_gateway_id: Option, + refund_payout_reference: Option, + refund_payout_proof: Option, + wallet_transaction_count: i64, +} + +struct RefundMutationFixture { + wallet_id: String, + payment_order_id: String, + refund_id: String, +} + +async fn create_refund_mutation_fixture( + repository: &SqliteWalletReadRepository, + case_name: &str, +) -> RefundMutationFixture { + let user_id = format!("refund-corruption-{case_name}"); + ensure_test_user(repository.pool(), &user_id).await; + let wallet = repository + .initialize_auth_user_wallet(&user_id, 0.0, false) + .await + .expect("refund corruption wallet initialization should run") + .expect("refund corruption wallet should exist"); + let (wallet, order) = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: wallet.id, + amount_usd: 10.0, + payment_method: "admin_manual".to_string(), + operator_id: Some("admin-refund-corruption".to_string()), + description: Some("refund corruption setup".to_string()), + order_no: format!("refund-corruption-order-{case_name}"), + }) + .await + .expect("refund corruption recharge should run") + .expect("refund corruption wallet should still exist"); + let refund = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: wallet.id.clone(), + user_id, + amount_usd: 2.0, + payment_order_id: Some(order.id.clone()), + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("refund corruption setup".to_string()), + idempotency_key: Some(format!("refund-corruption-idempotency-{case_name}")), + refund_no: format!("refund-corruption-refund-{case_name}"), + }) + .await + .expect("refund corruption request should run"); + let CreateWalletRefundRequestOutcome::Created(refund) = refund else { + panic!("refund corruption request should be created"); + }; + + RefundMutationFixture { + wallet_id: wallet.id, + payment_order_id: order.id, + refund_id: refund.id, + } +} + +async fn process_refund_mutation_fixture( + repository: &SqliteWalletReadRepository, + fixture: &RefundMutationFixture, +) { + let outcome = repository + .process_admin_wallet_refund(ProcessAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + operator_id: Some("admin-refund-corruption".to_string()), + }) + .await + .expect("refund corruption setup should process"); + assert!(matches!(outcome, WalletMutationOutcome::Applied(_))); +} + +async fn corrupt_refund_amount(pool: &sqlx::SqlitePool, refund_id: &str, amount_usd: f64) { + let result = sqlx::query("UPDATE refund_requests SET amount_usd = ? WHERE id = ?") + .bind(amount_usd) + .bind(refund_id) + .execute(pool) + .await + .expect("persisted refund amount should be corruptible for the regression test"); + assert_eq!(result.rows_affected(), 1); + + let stored_amount: f64 = + sqlx::query_scalar("SELECT amount_usd FROM refund_requests WHERE id = ?") + .bind(refund_id) + .fetch_one(pool) + .await + .expect("corrupted refund amount should query"); + if amount_usd.is_infinite() { + assert!(stored_amount.is_infinite()); + assert_eq!( + stored_amount.is_sign_positive(), + amount_usd.is_sign_positive() + ); + } else { + assert_eq!(stored_amount, amount_usd); + } +} + +async fn refund_mutation_snapshot( + pool: &sqlx::SqlitePool, + fixture: &RefundMutationFixture, +) -> RefundMutationSnapshot { + let (wallet_balance, wallet_total_refunded): (f64, f64) = + sqlx::query_as("SELECT balance, total_refunded FROM wallets WHERE id = ?") + .bind(&fixture.wallet_id) + .fetch_one(pool) + .await + .expect("refund corruption wallet should query"); + let (order_refunded_amount, order_refundable_amount): (f64, f64) = sqlx::query_as( + "SELECT refunded_amount_usd, refundable_amount_usd FROM payment_orders WHERE id = ?", + ) + .bind(&fixture.payment_order_id) + .fetch_one(pool) + .await + .expect("refund corruption payment order should query"); + let ( + refund_status, + refund_failure_reason, + refund_gateway_id, + refund_payout_reference, + refund_payout_proof, + ): ( + String, + Option, + Option, + Option, + Option, + ) = sqlx::query_as( + r#" +SELECT status, failure_reason, gateway_refund_id, payout_reference, payout_proof +FROM refund_requests +WHERE id = ? +"#, + ) + .bind(&fixture.refund_id) + .fetch_one(pool) + .await + .expect("refund corruption request should query"); + let wallet_transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?") + .bind(&fixture.wallet_id) + .fetch_one(pool) + .await + .expect("refund corruption transactions should count"); + + RefundMutationSnapshot { + wallet_balance, + wallet_total_refunded, + order_refunded_amount, + order_refundable_amount, + refund_status, + refund_failure_reason, + refund_gateway_id, + refund_payout_reference, + refund_payout_proof, + wallet_transaction_count, + } +} + +#[tokio::test] +async fn sqlite_process_refund_rejects_corrupt_persisted_amount_without_side_effects() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + for (index, amount_usd) in [0.0, -1.0, f64::INFINITY].into_iter().enumerate() { + let fixture = + create_refund_mutation_fixture(&repository, &format!("process-{index}")).await; + corrupt_refund_amount(&pool, &fixture.refund_id, amount_usd).await; + let before = refund_mutation_snapshot(&pool, &fixture).await; + + let outcome = repository + .process_admin_wallet_refund(ProcessAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + operator_id: Some("admin-refund-corruption".to_string()), + }) + .await + .expect("corrupted refund process should be rejected cleanly"); + + assert!(matches!(outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, &fixture).await, before); + } +} + +#[tokio::test] +async fn sqlite_complete_refund_rejects_corrupt_persisted_amount_without_side_effects() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + for (index, amount_usd) in [0.0, -1.0, f64::INFINITY].into_iter().enumerate() { + let fixture = + create_refund_mutation_fixture(&repository, &format!("complete-{index}")).await; + process_refund_mutation_fixture(&repository, &fixture).await; + corrupt_refund_amount(&pool, &fixture.refund_id, amount_usd).await; + let before = refund_mutation_snapshot(&pool, &fixture).await; + + let outcome = repository + .complete_admin_wallet_refund(CompleteAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: Some("must-not-be-stored".to_string()), + payout_reference: Some("must-not-be-stored".to_string()), + payout_proof: Some(json!({ "proof": "must-not-be-stored" })), + }) + .await + .expect("corrupted refund completion should be rejected cleanly"); + + assert!(matches!(outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, &fixture).await, before); + } +} + +#[tokio::test] +async fn sqlite_fail_refund_rejects_corrupt_persisted_amount_without_side_effects() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + for (index, amount_usd) in [0.0, -1.0, f64::INFINITY].into_iter().enumerate() { + let fixture = create_refund_mutation_fixture(&repository, &format!("fail-{index}")).await; + process_refund_mutation_fixture(&repository, &fixture).await; + corrupt_refund_amount(&pool, &fixture.refund_id, amount_usd).await; + let before = refund_mutation_snapshot(&pool, &fixture).await; + + let outcome = repository + .fail_admin_wallet_refund(FailAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + reason: "must not be stored".to_string(), + operator_id: Some("admin-refund-corruption".to_string()), + }) + .await + .expect("corrupted refund failure should be rejected cleanly"); + + assert!(matches!(outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, &fixture).await, before); + } +} + +#[tokio::test] +async fn sqlite_fail_refund_rejects_negative_recharge_balance_without_side_effects() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let fixture = create_refund_mutation_fixture(&repository, "fail-negative-balance").await; + process_refund_mutation_fixture(&repository, &fixture).await; + + sqlx::query("UPDATE wallets SET balance = ? WHERE id = ?") + .bind(-1.0_f64) + .bind(&fixture.wallet_id) + .execute(&pool) + .await + .expect("negative wallet balance should be seedable for the regression test"); + let before = refund_mutation_snapshot(&pool, &fixture).await; + + let outcome = repository + .fail_admin_wallet_refund(FailAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + reason: "must not recover a corrupt wallet".to_string(), + operator_id: Some("admin-refund-corruption".to_string()), + }) + .await + .expect("negative wallet balance failure should resolve"); + + assert!(matches!(outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, &fixture).await, before); +} + +#[tokio::test] +async fn sqlite_pending_gateway_refund_evidence_is_durable_and_cannot_be_reverted() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let fixture = create_refund_mutation_fixture(&repository, "pending-gateway").await; + process_refund_mutation_fixture(&repository, &fixture).await; + + let proof = json!({ + "gateway": "wxpay", + "status": "processing", + "refund_no": "provider-refund-1" + }); + let recorded = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-1".to_string(), + payout_proof: Some(proof.clone()), + }) + .await + .expect("gateway evidence update should run"); + let WalletMutationOutcome::Applied(recorded_refund) = recorded else { + panic!("gateway evidence should be recorded"); + }; + assert_eq!(recorded_refund.status, "processing"); + assert_eq!( + recorded_refund.gateway_refund_id.as_deref(), + Some("provider-refund-1") + ); + assert_eq!(recorded_refund.payout_proof, Some(proof.clone())); + + let replay = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-1".to_string(), + payout_proof: Some(json!({ "status": "different" })), + }) + .await + .expect("same gateway evidence replay should run"); + let WalletMutationOutcome::Applied(replayed_refund) = replay else { + panic!("same gateway evidence replay should be accepted"); + }; + assert_eq!(replayed_refund.payout_proof, Some(proof)); + + let conflict = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-attacker".to_string(), + payout_proof: None, + }) + .await + .expect("conflicting gateway evidence should resolve"); + assert!(matches!(conflict, WalletMutationOutcome::Invalid(_))); + + let before_fail = refund_mutation_snapshot(&pool, &fixture).await; + let fail = repository + .fail_admin_wallet_refund(FailAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + reason: "provider still processing".to_string(), + operator_id: Some("admin-refund-pending".to_string()), + }) + .await + .expect("processing refund failure should resolve"); + assert!(matches!(fail, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, &fixture).await, before_fail); + + let completed = repository + .complete_admin_wallet_refund(CompleteAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: None, + payout_reference: None, + payout_proof: None, + }) + .await + .expect("completion should preserve provider evidence"); + let WalletMutationOutcome::Applied(completed_refund) = completed else { + panic!("processing refund should complete"); + }; + assert_eq!(completed_refund.status, "succeeded"); + assert_eq!( + completed_refund.gateway_refund_id.as_deref(), + Some("provider-refund-1") + ); + assert_eq!( + completed_refund.payout_proof, + Some(json!({ + "gateway": "wxpay", + "status": "processing", + "refund_no": "provider-refund-1" + })) + ); + + let terminal_update_conflict = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-attacker".to_string(), + payout_proof: None, + }) + .await + .expect("terminal gateway evidence conflict should resolve"); + assert!(matches!( + terminal_update_conflict, + WalletMutationOutcome::Invalid(_) + )); + + let terminal_complete_conflict = repository + .complete_admin_wallet_refund(CompleteAdminWalletRefundInput { + wallet_id: fixture.wallet_id, + refund_id: fixture.refund_id, + gateway_refund_id: Some("provider-refund-attacker".to_string()), + payout_reference: None, + payout_proof: None, + }) + .await + .expect("terminal completion conflict should resolve"); + assert!(matches!( + terminal_complete_conflict, + WalletMutationOutcome::Invalid(_) + )); +} + +#[tokio::test] +async fn sqlite_success_gateway_refund_proof_upgrades_processing_evidence() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let fixture = create_refund_mutation_fixture(&repository, "success-proof-upgrade").await; + process_refund_mutation_fixture(&repository, &fixture).await; + + let processing_proof = json!({ + "gateway": "wxpay", + "id": "provider-refund-upgrade", + "status": "processing" + }); + let recorded = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-upgrade".to_string(), + payout_proof: Some(processing_proof.clone()), + }) + .await + .expect("processing evidence should persist"); + assert!(matches!(recorded, WalletMutationOutcome::Applied(_))); + + let replay = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-upgrade".to_string(), + payout_proof: Some(json!({ + "gateway": "wxpay", + "id": "provider-refund-upgrade", + "status": "processing", + "attempt": 2 + })), + }) + .await + .expect("processing replay should resolve"); + let WalletMutationOutcome::Applied(replayed) = replay else { + panic!("processing replay should be accepted"); + }; + assert_eq!(replayed.payout_proof, Some(processing_proof)); + + let success_proof = json!({ + "gateway": "wxpay", + "id": "provider-refund-upgrade", + "status": "success", + "processed_at": "2026-08-29T00:00:00Z" + }); + let upgraded = repository + .update_admin_wallet_refund_gateway(UpdateAdminWalletRefundGatewayInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + gateway_refund_id: "provider-refund-upgrade".to_string(), + payout_proof: Some(success_proof.clone()), + }) + .await + .expect("success evidence should upgrade processing proof"); + let WalletMutationOutcome::Applied(upgraded) = upgraded else { + panic!("success evidence should be accepted"); + }; + assert_eq!(upgraded.payout_proof, Some(success_proof.clone())); + + let completed = repository + .complete_admin_wallet_refund(CompleteAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id, + gateway_refund_id: Some("provider-refund-upgrade".to_string()), + payout_reference: None, + payout_proof: None, + }) + .await + .expect("refund should complete"); + let WalletMutationOutcome::Applied(completed) = completed else { + panic!("refund should complete"); + }; + assert_eq!(completed.status, "succeeded"); + assert_eq!(completed.payout_proof, Some(success_proof)); +} + +#[tokio::test] +async fn sqlite_offline_processing_refund_failure_releases_reservation() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let fixture = create_refund_mutation_fixture(&repository, "offline-failure").await; + process_refund_mutation_fixture(&repository, &fixture).await; + let before = refund_mutation_snapshot(&pool, &fixture).await; + + let outcome = repository + .fail_admin_wallet_refund(FailAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + reason: "offline payout was not sent".to_string(), + operator_id: Some("admin-offline-failure".to_string()), + }) + .await + .expect("offline processing refund failure should resolve"); + let WalletMutationOutcome::Applied((wallet, refund, transaction)) = outcome else { + panic!("offline processing refund should be released"); + }; + let transaction = transaction.expect("refund recovery transaction should be recorded"); + assert_eq!(refund.status, "failed"); + assert_eq!( + refund.failure_reason.as_deref(), + Some("offline payout was not sent") + ); + assert_eq!(transaction.reason_code, "refund_revert"); + assert_eq!(wallet.balance, 10.0); + assert_eq!(wallet.total_refunded, 0.0); + + let after = refund_mutation_snapshot(&pool, &fixture).await; + assert_eq!(after.wallet_balance, before.wallet_balance + 2.0); + assert_eq!(after.wallet_total_refunded, 0.0); + assert_eq!(after.order_refunded_amount, 0.0); + assert_eq!(after.order_refundable_amount, 10.0); + assert_eq!(after.refund_status, "failed"); + assert_eq!( + after.wallet_transaction_count, + before.wallet_transaction_count + 1 + ); +} + +#[tokio::test] +async fn sqlite_processing_refund_requires_offline_mode_without_gateway_evidence() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + let original_channel = create_refund_mutation_fixture(&repository, "original-channel").await; + process_refund_mutation_fixture(&repository, &original_channel).await; + sqlx::query("UPDATE refund_requests SET refund_mode = 'original_channel' WHERE id = ?") + .bind(&original_channel.refund_id) + .execute(&pool) + .await + .expect("refund mode should update for the regression test"); + + let proof_only = create_refund_mutation_fixture(&repository, "proof-only").await; + process_refund_mutation_fixture(&repository, &proof_only).await; + sqlx::query("UPDATE refund_requests SET payout_proof = ? WHERE id = ?") + .bind(r#"{"gateway":"manual-settlement"}"#) + .bind(&proof_only.refund_id) + .execute(&pool) + .await + .expect("proof-only evidence should update for the regression test"); + + for fixture in [&original_channel, &proof_only] { + let before = refund_mutation_snapshot(&pool, fixture).await; + let outcome = repository + .fail_admin_wallet_refund(FailAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + reason: "must remain reserved".to_string(), + operator_id: Some("admin-preserve-reservation".to_string()), + }) + .await + .expect("protected processing refund failure should resolve"); + assert!(matches!(outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!(refund_mutation_snapshot(&pool, fixture).await, before); + } +} + +#[tokio::test] +async fn sqlite_refund_rejects_foreign_or_uncredited_payment_order_without_side_effects() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let fixture = create_refund_mutation_fixture(&repository, "order-integrity").await; + + ensure_test_user(&pool, "refund-order-integrity-other").await; + let other_wallet = repository + .initialize_auth_user_wallet("refund-order-integrity-other", 0.0, false) + .await + .expect("other wallet initialization should run") + .expect("other wallet should exist"); + let (_, other_order) = repository + .create_manual_wallet_recharge(CreateManualWalletRechargeInput { + wallet_id: other_wallet.id, + amount_usd: 10.0, + payment_method: "admin_manual".to_string(), + operator_id: Some("admin-order-integrity".to_string()), + description: Some("order integrity setup".to_string()), + order_no: "refund-order-integrity-other-order".to_string(), + }) + .await + .expect("other recharge should run") + .expect("other recharge should create an order"); + + sqlx::query("UPDATE refund_requests SET payment_order_id = ? WHERE id = ?") + .bind(&other_order.id) + .bind(&fixture.refund_id) + .execute(&pool) + .await + .expect("foreign payment order should be assignable for regression setup"); + let before_foreign = refund_mutation_snapshot(&pool, &fixture).await; + let foreign_outcome = repository + .process_admin_wallet_refund(ProcessAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + operator_id: Some("admin-order-integrity".to_string()), + }) + .await + .expect("foreign payment order refund should resolve"); + assert!(matches!(foreign_outcome, WalletMutationOutcome::Invalid(_))); + assert_eq!( + refund_mutation_snapshot(&pool, &fixture).await, + before_foreign + ); + + sqlx::query("UPDATE refund_requests SET payment_order_id = ? WHERE id = ?") + .bind(&fixture.payment_order_id) + .bind(&fixture.refund_id) + .execute(&pool) + .await + .expect("original payment order should be restored for regression setup"); + sqlx::query("UPDATE payment_orders SET status = 'pending' WHERE id = ?") + .bind(&fixture.payment_order_id) + .execute(&pool) + .await + .expect("payment order status should be corruptible for regression setup"); + let before_uncredited = refund_mutation_snapshot(&pool, &fixture).await; + let uncredited_outcome = repository + .process_admin_wallet_refund(ProcessAdminWalletRefundInput { + wallet_id: fixture.wallet_id.clone(), + refund_id: fixture.refund_id.clone(), + operator_id: Some("admin-order-integrity".to_string()), + }) + .await + .expect("uncredited payment order refund should resolve"); + assert!(matches!( + uncredited_outcome, + WalletMutationOutcome::Invalid(_) + )); + assert_eq!( + refund_mutation_snapshot(&pool, &fixture).await, + before_uncredited + ); +} + #[tokio::test] async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_refund() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -114,7 +1615,17 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref .await .expect("sqlite migrations should run"); - let repository = SqliteWalletReadRepository::new(pool); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users( + &pool, + &[ + "user-write-1", + "user-credit-1", + "user-expire-1", + "user-fail-1", + ], + ) + .await; let order = match repository .create_wallet_recharge_order(CreateWalletRechargeOrderInput { preferred_wallet_id: Some("wallet-write-1".to_string()), @@ -124,8 +1635,8 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref pay_currency: Some("USD".to_string()), exchange_rate: Some(1.0), payment_method: "alipay".to_string(), - payment_provider: None, - payment_channel: None, + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), gateway_order_id: "gateway-order-write-1".to_string(), gateway_response: json!({ "checkout": true }), order_no: "order-no-write-1".to_string(), @@ -138,14 +1649,17 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new wallet order should not already exist") + } }; assert_eq!(order.status, "pending"); let callback = repository .process_payment_callback(ProcessPaymentCallbackInput { payment_method: "alipay".to_string(), - payment_provider: None, - payment_channel: None, + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), callback_key: "callback-write-1".to_string(), order_no: Some("order-no-write-1".to_string()), gateway_order_id: Some("gateway-order-write-1".to_string()), @@ -154,7 +1668,12 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref pay_currency: Some("USD".to_string()), exchange_rate: Some(1.0), payload_hash: "payload-hash-write-1".to_string(), - payload: json!({ "status": "paid" }), + payload: json!({ + "status": "paid", + "client_secret": "pi_1_secret_replayable", + "customer": {"email": "payer@example.com"}, + "authorization": "Bearer upstream-secret", + }), signature_valid: true, }) .await @@ -167,6 +1686,55 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref }; assert_eq!(wallet_id, "wallet-write-1"); assert_eq!(order.status, "credited"); + let stored_callback_payload: Option = + sqlx::query_scalar("SELECT payload FROM payment_callbacks WHERE callback_key = ?") + .bind("callback-write-1") + .fetch_one(&pool) + .await + .expect("callback payload should query"); + assert_eq!(stored_callback_payload, None); + let stored_gateway_response: String = + sqlx::query_scalar("SELECT gateway_response FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("gateway response should query"); + let stored_gateway_response: serde_json::Value = + serde_json::from_str(&stored_gateway_response).expect("gateway response should be JSON"); + assert_eq!(stored_gateway_response["gateway"], "alipay"); + assert_eq!(stored_gateway_response["payment_provider"], "alipay"); + assert_eq!(stored_gateway_response["payment_channel"], "alipay"); + assert_eq!(stored_gateway_response["order_no"], "order-no-write-1"); + assert_eq!(stored_gateway_response["amount_usd"], 12.5); + assert!(stored_gateway_response.get("status").is_none()); + let stored_settlement_binding: (String, Option, Option) = sqlx::query_as( + "SELECT payment_method, payment_provider, payment_channel FROM payment_orders WHERE id = ?", + ) + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("settlement binding should query"); + assert_eq!( + stored_settlement_binding, + ( + "alipay".to_string(), + Some("alipay".to_string()), + Some("alipay".to_string()), + ) + ); + let encoded_gateway_response = stored_gateway_response.to_string(); + for forbidden in [ + "client_secret", + "replayable", + "payer@example.com", + "authorization", + "upstream-secret", + ] { + assert!( + !encoded_gateway_response.contains(forbidden), + "persisted {forbidden}" + ); + } let wallet = repository .find(WalletLookupKey::UserId("user-write-1")) @@ -248,6 +1816,34 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref completed_refund.payout_proof.as_ref().unwrap()["proof"], "ok" ); + let repeated_completion = repository + .complete_admin_wallet_refund(CompleteAdminWalletRefundInput { + wallet_id: wallet.id.clone(), + refund_id: refund.id.clone(), + gateway_refund_id: Some("gateway-refund-attacker".to_string()), + payout_reference: Some("payout-ref-attacker".to_string()), + payout_proof: Some(json!({ "proof": "attacker" })), + }) + .await + .expect("completed refund replay should resolve"); + assert!(matches!( + repeated_completion, + WalletMutationOutcome::Invalid(_) + )); + let repeated_refund = repository + .find_wallet_refund(&wallet.id, &refund.id) + .await + .expect("completed refund should still load") + .expect("completed refund should still exist"); + assert_eq!( + repeated_refund.gateway_refund_id.as_deref(), + Some("gateway-refund-write-1") + ); + assert_eq!( + repeated_refund.payout_reference.as_deref(), + Some("payout-ref-write-1") + ); + assert_eq!(repeated_refund.payout_proof, completed_refund.payout_proof); let refund_to_fail = repository .create_wallet_refund_request(CreateWalletRefundRequestInput { @@ -267,18 +1863,15 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref let CreateWalletRefundRequestOutcome::Created(refund_to_fail) = refund_to_fail else { panic!("second refund request should be created"); }; - let processed_to_fail = repository - .process_admin_wallet_refund(ProcessAdminWalletRefundInput { + let before_fail = refund_mutation_snapshot( + &pool, + &RefundMutationFixture { wallet_id: wallet.id.clone(), + payment_order_id: order.id.clone(), refund_id: refund_to_fail.id.clone(), - operator_id: Some("admin-1".to_string()), - }) - .await - .expect("second refund should process"); - assert!(matches!( - processed_to_fail, - WalletMutationOutcome::Applied(_) - )); + }, + ) + .await; let failed = repository .fail_admin_wallet_refund(FailAdminWalletRefundInput { wallet_id: wallet.id.clone(), @@ -287,16 +1880,31 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref operator_id: Some("admin-1".to_string()), }) .await - .expect("second refund should fail"); - let WalletMutationOutcome::Applied((wallet, failed_refund, revert_transaction)) = failed else { - panic!("second refund should fail with wallet restoration"); + .expect("second refund failure should resolve"); + assert!(matches!(&failed, WalletMutationOutcome::Applied(_))); + let failed_refund = match failed { + WalletMutationOutcome::Applied((_, refund, transaction)) => { + assert!(transaction.is_none()); + refund + } + _ => unreachable!(), }; - assert_eq!(wallet.balance, 8.5); assert_eq!(failed_refund.status, "failed"); + let after_fail = refund_mutation_snapshot( + &pool, + &RefundMutationFixture { + wallet_id: wallet.id.clone(), + payment_order_id: order.id.clone(), + refund_id: refund_to_fail.id.clone(), + }, + ) + .await; + assert_eq!(after_fail.wallet_balance, before_fail.wallet_balance); assert_eq!( - revert_transaction.as_ref().unwrap().reason_code, - "refund_revert" + after_fail.wallet_total_refunded, + before_fail.wallet_total_refunded ); + assert_eq!(after_fail.refund_status, "failed"); let batch = repository .create_admin_redeem_code_batch(CreateAdminRedeemCodeBatchInput { @@ -440,6 +2048,9 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new credit wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new credit order should not already exist") + } }; let credited = repository .credit_admin_payment_order(CreditAdminPaymentOrderInput { @@ -502,6 +2113,9 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new expire wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new expiring order should not already exist") + } }; let expired = repository .expire_admin_payment_order(&expiring_order.id) @@ -540,6 +2154,9 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new fail wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new failing order should not already exist") + } }; let failed_order = repository .fail_admin_payment_order(&failing_order.id) @@ -551,6 +2168,2027 @@ async fn sqlite_wallet_write_repository_handles_public_recharge_callback_and_ref assert_eq!(failed_order.status, "failed"); } +#[tokio::test] +async fn sqlite_payment_callback_rejects_gateway_identifier_mismatch_without_crediting() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-identifier-mismatch").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-identifier-mismatch".to_string()), + user_id: "user-callback-identifier-mismatch".to_string(), + amount_usd: 15.0, + pay_amount: Some(15.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-order-original".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-no-identifier-mismatch".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => { + panic!("new wallet should be active") + } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new identifier-mismatch order should not already exist") + } + }; + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-identifier-mismatch".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-order-attacker".to_string()), + amount_usd: 15.0, + pay_amount: Some(15.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-identifier-mismatch".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("mismatched callback should resolve"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment gateway order mismatch" + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("wallet-callback-identifier-mismatch") + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (0.0, 0.0)); + + let stored_order: (String, String, Option, Option) = sqlx::query_as( + "SELECT status, gateway_order_id, paid_at, credited_at FROM payment_orders WHERE id = ?", + ) + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order should load"); + assert_eq!(stored_order.0, "pending"); + assert_eq!(stored_order.1, "gateway-order-original"); + assert_eq!(stored_order.2, None); + assert_eq!(stored_order.3, None); + + let callback: (String, Option) = sqlx::query_as( + "SELECT status, error_message FROM payment_callbacks WHERE callback_key = ?", + ) + .bind("callback-identifier-mismatch") + .fetch_one(&pool) + .await + .expect("callback should load"); + assert_eq!(callback.0, "failed"); + assert_eq!( + callback.1.as_deref(), + Some("payment gateway order mismatch") + ); +} + +#[tokio::test] +async fn sqlite_payment_callback_rejects_wallet_owner_mismatch_without_crediting() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users(&pool, &["callback-order-owner", "callback-wallet-owner"]).await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-owner-mismatch".to_string()), + user_id: "callback-wallet-owner".to_string(), + amount_usd: 8.0, + pay_amount: Some(8.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-owner-mismatch".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-owner-mismatch".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + // Simulate a corrupted/imported order that points at a different user + // than the wallet selected at checkout. + sqlx::query("UPDATE payment_orders SET user_id = ? WHERE id = ?") + .bind("callback-order-owner") + .bind(&order.id) + .execute(&pool) + .await + .expect("order owner should update for regression setup"); + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-owner-mismatch".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-owner-mismatch".to_string()), + amount_usd: 8.0, + pay_amount: Some(8.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-owner-mismatch".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("owner-mismatch callback should resolve"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment order wallet owner mismatch" + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind(&order.wallet_id) + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (0.0, 0.0)); + let stored_order: (String, Option, Option) = + sqlx::query_as("SELECT status, paid_at, credited_at FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order should load"); + assert_eq!(stored_order, ("pending".to_string(), None, None)); + let callback: (String, Option) = sqlx::query_as( + "SELECT status, error_message FROM payment_callbacks WHERE callback_key = ?", + ) + .bind("callback-owner-mismatch") + .fetch_one(&pool) + .await + .expect("callback should load"); + assert_eq!(callback.0, "failed"); + assert_eq!( + callback.1.as_deref(), + Some("payment order wallet owner mismatch") + ); +} + +#[tokio::test] +async fn sqlite_payment_callback_credits_overdrawn_recharge_balance() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-overdrawn").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-overdrawn".to_string()), + user_id: "user-callback-overdrawn".to_string(), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-callback-overdrawn".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-callback-overdrawn".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + sqlx::query("UPDATE wallets SET balance = -3.0, total_recharged = 0.0 WHERE id = ?") + .bind(&order.wallet_id) + .execute(&pool) + .await + .expect("wallet should be made overdrawn"); + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-overdrawn".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-callback-overdrawn".to_string()), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-callback-overdrawn".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("payment callback should process"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + + let wallet = repository + .find(WalletLookupKey::UserId("user-callback-overdrawn")) + .await + .expect("wallet should query") + .expect("wallet should exist"); + assert_eq!(wallet.balance, 2.0); + assert_eq!(wallet.total_recharged, 5.0); + + let callback: (String, Option) = sqlx::query_as( + "SELECT status, error_message FROM payment_callbacks WHERE callback_key = ?", + ) + .bind("callback-overdrawn") + .fetch_one(&pool) + .await + .expect("callback should load"); + assert_eq!(callback.0, "processed"); + assert_eq!(callback.1, None); +} + +#[tokio::test] +async fn sqlite_manual_credit_rejects_invalid_order_and_wallet_values() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-manual-credit-invalid").await; + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-manual-credit-invalid".to_string()), + user_id: "user-manual-credit-invalid".to_string(), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "manual_gateway".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "gateway-manual-credit-invalid".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-manual-credit-invalid".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("credit order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + sqlx::query("UPDATE payment_orders SET amount_usd = ? WHERE id = ?") + .bind(-5.0) + .bind(&order.id) + .execute(&pool) + .await + .expect("invalid order fixture should update"); + let invalid_order = repository + .credit_admin_payment_order(CreditAdminPaymentOrderInput { + order_id: order.id.clone(), + gateway_order_id: None, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + gateway_response_patch: None, + operator_id: Some("admin-invalid".to_string()), + }) + .await + .expect("invalid order credit should resolve"); + assert!(matches!( + invalid_order, + WalletMutationOutcome::Invalid(ref error) if error == "payment order amount is invalid" + )); + let order_status: String = sqlx::query_scalar("SELECT status FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order status should query"); + assert_eq!(order_status, "pending"); + + sqlx::query("UPDATE payment_orders SET amount_usd = ? WHERE id = ?") + .bind(5.0) + .bind(&order.id) + .execute(&pool) + .await + .expect("order fixture should restore"); + sqlx::query("UPDATE wallets SET gift_balance = ? WHERE id = ?") + .bind(-1.0) + .bind(&order.wallet_id) + .execute(&pool) + .await + .expect("invalid wallet fixture should update"); + let invalid_wallet = repository + .credit_admin_payment_order(CreditAdminPaymentOrderInput { + order_id: order.id.clone(), + gateway_order_id: None, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + gateway_response_patch: None, + operator_id: Some("admin-invalid".to_string()), + }) + .await + .expect("invalid wallet credit should resolve"); + assert!(matches!( + invalid_wallet, + WalletMutationOutcome::Invalid(ref error) if error == "wallet balance is invalid" + )); + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, gift_balance FROM wallets WHERE id = ?") + .bind(&order.wallet_id) + .fetch_one(&pool) + .await + .expect("wallet should query"); + assert_eq!(wallet, (0.0, -1.0)); + let order_status: String = sqlx::query_scalar("SELECT status FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order status should query"); + assert_eq!(order_status, "pending"); +} + +#[tokio::test] +async fn sqlite_payment_callback_rejects_unknown_order_status_without_crediting() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-invalid-state").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-invalid-state".to_string()), + user_id: "user-callback-invalid-state".to_string(), + amount_usd: 11.0, + pay_amount: Some(11.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-invalid-state".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-invalid-state".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + sqlx::query("UPDATE payment_orders SET status = 'cancelled' WHERE id = ?") + .bind(&order.id) + .execute(&pool) + .await + .expect("test order status should update"); + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-invalid-state".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-invalid-state".to_string()), + amount_usd: 11.0, + pay_amount: Some(11.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-invalid-state".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("invalid-state callback should resolve"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment order is not creditable: cancelled" + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("wallet-callback-invalid-state") + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (0.0, 0.0)); + let stored_order: (String, Option, Option) = + sqlx::query_as("SELECT status, paid_at, credited_at FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order should load"); + assert_eq!(stored_order, ("cancelled".to_string(), None, None)); +} + +#[tokio::test] +async fn sqlite_payment_callback_recovers_failed_checkout_placeholder_without_losing_credit() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-failed-checkout").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-failed-checkout".to_string()), + user_id: "user-callback-failed-checkout".to_string(), + amount_usd: 12.0, + pay_amount: Some(12.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "order-callback-failed-checkout".to_string(), + gateway_response: json!({ + "gateway": "alipay", + "gateway_order_id": "order-callback-failed-checkout", + "order_kind": "wallet_recharge", + "payment_channel": "alipay", + "pay_amount": 12.0, + "pay_currency": "USD", + "integration_status": "checkout_pending", + "checkout_claim_token": "claim-failed-checkout", + "checkout_claimed_at_unix_secs": 1, + }), + order_no: "order-callback-failed-checkout".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + let claim_token = order + .gateway_response + .as_ref() + .and_then(|response| response.get("checkout_claim_token")) + .and_then(serde_json::Value::as_str) + .expect("checkout claim token should persist") + .to_string(); + + let failed = repository + .fail_wallet_recharge_checkout(FailWalletRechargeCheckoutInput { + order_id: order.id.clone(), + claim_token, + reason: "checkout response timed out after provider acceptance".to_string(), + provider_request_may_have_succeeded: true, + }) + .await + .expect("checkout failure should resolve"); + let WalletMutationOutcome::Applied(failed) = failed else { + panic!("checkout placeholder should become failed"); + }; + assert_eq!(failed.status, "failed"); + assert_eq!( + failed + .gateway_response + .as_ref() + .and_then(|response| response.get("integration_status")) + .and_then(serde_json::Value::as_str), + Some("checkout_uncertain") + ); + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-failed-checkout-recovery".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("provider-order-failed-checkout".to_string()), + amount_usd: 12.0, + pay_amount: Some(12.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-failed-checkout-recovery".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("provider callback should process"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind(&order.wallet_id) + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (12.0, 12.0)); + + let stored_order: (String, Option) = + sqlx::query_as("SELECT status, gateway_order_id FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(&pool) + .await + .expect("order should load"); + assert_eq!( + stored_order, + ( + "credited".to_string(), + Some("provider-order-failed-checkout".to_string()) + ) + ); + + let callback_status: String = + sqlx::query_scalar("SELECT status FROM payment_callbacks WHERE callback_key = ?") + .bind("callback-failed-checkout-recovery") + .fetch_one(&pool) + .await + .expect("callback should load"); + assert_eq!(callback_status, "processed"); +} + +#[tokio::test] +async fn sqlite_payment_callback_rejects_corrupt_stored_credit_amount() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-corrupt-amount").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-corrupt-amount".to_string()), + user_id: "user-callback-corrupt-amount".to_string(), + amount_usd: 11.0, + pay_amount: Some(11.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-corrupt-amount".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-corrupt-amount".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + sqlx::query("UPDATE payment_orders SET amount_usd = -11 WHERE id = ?") + .bind(&order.id) + .execute(&pool) + .await + .expect("test order amount should update"); + + let outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-corrupt-amount".to_string(), + order_no: Some(order.order_no), + gateway_order_id: Some("gateway-corrupt-amount".to_string()), + amount_usd: 11.0, + pay_amount: Some(11.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-corrupt-amount".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("corrupt-amount callback should resolve"); + assert!(matches!( + outcome, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment order amount is invalid" + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("wallet-callback-corrupt-amount") + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (0.0, 0.0)); +} + +#[tokio::test] +async fn sqlite_payment_callback_reconstructs_legacy_provider_amount_from_order_terms() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users( + &pool, + &["legacy-cny-callback-user", "legacy-usd-callback-user"], + ) + .await; + + let cny_order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("legacy-cny-callback-wallet".to_string()), + user_id: "legacy-cny-callback-user".to_string(), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "legacy-cny-callback-gateway".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "legacy-cny-callback-order".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("CNY recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + sqlx::query("UPDATE payment_orders SET pay_amount = NULL WHERE id = ?") + .bind(&cny_order.id) + .execute(&pool) + .await + .expect("legacy CNY order should drop provider amount"); + + let cny_outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "legacy-cny-callback".to_string(), + order_no: Some(cny_order.order_no), + gateway_order_id: Some("legacy-cny-callback-gateway".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payload_hash: "legacy-cny-callback-payload".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("legacy CNY callback should resolve"); + assert!(matches!( + cny_outcome, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + let cny_wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("legacy-cny-callback-wallet") + .fetch_one(&pool) + .await + .expect("legacy CNY wallet should load"); + assert_eq!(cny_wallet, (10.0, 10.0)); + + let usd_order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("legacy-usd-callback-wallet".to_string()), + user_id: "legacy-usd-callback-user".to_string(), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + // Old rows could retain the historical CNY default for USD. + exchange_rate: Some(7.2), + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "legacy-usd-callback-gateway".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "legacy-usd-callback-order".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("USD recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + sqlx::query("UPDATE payment_orders SET pay_amount = NULL WHERE id = ?") + .bind(&usd_order.id) + .execute(&pool) + .await + .expect("legacy USD order should drop provider amount"); + + let wrong_usd_outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + callback_key: "legacy-usd-callback-wrong".to_string(), + order_no: Some(usd_order.order_no.clone()), + gateway_order_id: Some("legacy-usd-callback-gateway".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(7.2), + payload_hash: "legacy-usd-callback-wrong-payload".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("wrong USD callback should resolve"); + assert!(matches!( + wrong_usd_outcome, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "callback amount mismatch" + )); + let usd_wallet_before: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("legacy-usd-callback-wallet") + .fetch_one(&pool) + .await + .expect("legacy USD wallet should load before valid callback"); + assert_eq!(usd_wallet_before, (0.0, 0.0)); + + let valid_usd_outcome = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + callback_key: "legacy-usd-callback-valid".to_string(), + order_no: Some(usd_order.order_no), + gateway_order_id: Some("legacy-usd-callback-gateway".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(7.2), + payload_hash: "legacy-usd-callback-valid-payload".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("valid USD callback should resolve"); + assert!(matches!( + valid_usd_outcome, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + let usd_wallet_after: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("legacy-usd-callback-wallet") + .fetch_one(&pool) + .await + .expect("legacy USD wallet should load after valid callback"); + assert_eq!(usd_wallet_after, (10.0, 10.0)); +} + +#[tokio::test] +async fn sqlite_payment_callback_requires_provider_namespace_to_match_exactly() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-provider-boundary").await; + + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-provider-boundary".to_string()), + user_id: "user-provider-boundary".to_string(), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-provider-boundary".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-provider-boundary".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + let error = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: None, + payment_channel: None, + callback_key: "callback-provider-boundary".to_string(), + order_no: Some(order.order_no), + gateway_order_id: Some("gateway-provider-boundary".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payload_hash: "payload-provider-boundary".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect_err("provider-less official callback must be rejected at repository boundary"); + assert!(matches!( + error, + aether_data_contracts::DataLayerError::InvalidInput(ref detail) + if detail == "official payment callback provider binding mismatch" + )); + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("wallet-provider-boundary") + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (0.0, 0.0)); +} + +#[tokio::test] +async fn sqlite_legacy_epay_channel_order_without_provider_is_compatible() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-legacy-epay").await; + + // Rows written before payment_provider/payment_channel were introduced + // stored the selected EPay channel as payment_method and left both new + // columns NULL. + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-legacy-epay".to_string()), + user_id: "user-legacy-epay".to_string(), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payment_method: "alipay".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "gateway-legacy-epay".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-legacy-epay".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("legacy order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + let wrong_channel = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("wxpay".to_string()), + callback_key: "callback-legacy-epay-wrong-channel".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-legacy-epay".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payload_hash: "payload-legacy-epay-wrong-channel".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("wrong-channel callback should resolve"); + assert!(matches!( + wrong_channel, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment channel mismatch" + )); + + let applied = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-legacy-epay-success".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-legacy-epay".to_string()), + amount_usd: 10.0, + pay_amount: Some(72.0), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payload_hash: "payload-legacy-epay-success".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("legacy EPay callback should process"); + assert!(matches!( + applied, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind(&order.wallet_id) + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (10.0, 10.0)); +} + +#[tokio::test] +async fn sqlite_payment_callback_validates_placeholder_gateway_binding_and_conflicts() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users( + &pool, + &[ + "user-callback-placeholder-a", + "user-callback-placeholder-b", + "user-callback-placeholder-c", + ], + ) + .await; + + // Order B already owns the real provider transaction id. + let order_b = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-placeholder-b".to_string()), + user_id: "user-callback-placeholder-b".to_string(), + amount_usd: 7.0, + pay_amount: Some(7.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-b".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-b".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("order B should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + // Order A stores its merchant order number as a provider-id placeholder. + let order_a = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-placeholder-a".to_string()), + user_id: "user-callback-placeholder-a".to_string(), + amount_usd: 7.0, + pay_amount: Some(7.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "order-a".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-a".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("order A should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + let conflict = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-placeholder-conflict".to_string(), + order_no: Some(order_a.order_no.clone()), + gateway_order_id: Some("gateway-b".to_string()), + amount_usd: 7.0, + pay_amount: Some(7.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-placeholder-conflict".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("placeholder conflict callback should resolve"); + assert!(matches!( + conflict, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment gateway order belongs to another payment order" + )); + let order_a_state: (String, String) = + sqlx::query_as("SELECT status, gateway_order_id FROM payment_orders WHERE id = ?") + .bind(&order_a.id) + .fetch_one(&pool) + .await + .expect("order A should load"); + assert_eq!( + order_a_state, + ("pending".to_string(), "order-a".to_string()) + ); + let order_b_state: (String, f64) = + sqlx::query_as("SELECT status, amount_usd FROM payment_orders WHERE id = ?") + .bind(&order_b.id) + .fetch_one(&pool) + .await + .expect("order B should load"); + assert_eq!(order_b_state.0, "pending"); + assert_eq!(order_b_state.1, 7.0); + + // A fresh provider id that is not owned by another order is bound and + // credited atomically, replacing the placeholder. + let order_c = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-placeholder-c".to_string()), + user_id: "user-callback-placeholder-c".to_string(), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "order-c".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-c".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("order C should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + let applied = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-placeholder-success".to_string(), + order_no: Some(order_c.order_no.clone()), + gateway_order_id: Some("gateway-c".to_string()), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-placeholder-success".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("placeholder binding callback should apply"); + assert!(matches!( + applied, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + let order_c_state: (String, String) = + sqlx::query_as("SELECT status, gateway_order_id FROM payment_orders WHERE id = ?") + .bind(&order_c.id) + .fetch_one(&pool) + .await + .expect("order C should load"); + assert_eq!( + order_c_state, + ("credited".to_string(), "gateway-c".to_string()) + ); + let wallet_c_balance: f64 = sqlx::query_scalar("SELECT balance FROM wallets WHERE id = ?") + .bind("wallet-callback-placeholder-c") + .fetch_one(&pool) + .await + .expect("wallet C should load"); + assert_eq!(wallet_c_balance, 5.0); +} + +#[tokio::test] +async fn sqlite_payment_gateway_order_identifier_is_unique_within_payment_method() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users( + &pool, + &[ + "user-gateway-unique-first", + "user-gateway-unique-second", + "user-gateway-other-method", + "user-gateway-case-distinct", + ], + ) + .await; + + let first = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-gateway-unique-first".to_string()), + user_id: "user-gateway-unique-first".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: " EPAY ".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "shared-provider-transaction".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-gateway-unique-first".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("first order should create"); + assert!(matches!( + first, + CreateWalletRechargeOrderOutcome::Created(ref order) if order.payment_method == "epay" + )); + + let duplicate = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-gateway-unique-second".to_string()), + user_id: "user-gateway-unique-second".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "shared-provider-transaction".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-gateway-unique-second".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await; + assert!(duplicate.is_err()); + assert!(repository + .find(WalletLookupKey::UserId("user-gateway-unique-second")) + .await + .expect("conflicting order must not leave a wallet") + .is_none()); + + let other_method = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-gateway-other-method".to_string()), + user_id: "user-gateway-other-method".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "stripe".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "shared-provider-transaction".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-gateway-other-method".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("another payment method may reuse the identifier"); + assert!(matches!( + other_method, + CreateWalletRechargeOrderOutcome::Created(_) + )); + + let case_distinct_identifier = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-gateway-case-distinct".to_string()), + user_id: "user-gateway-case-distinct".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "Shared-Provider-Transaction".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-gateway-case-distinct".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("opaque identifiers that differ by case may coexist"); + assert!(matches!( + case_distinct_identifier, + CreateWalletRechargeOrderOutcome::Created(_) + )); +} + +#[tokio::test] +async fn sqlite_payment_callback_failure_does_not_rebind_existing_identifiers() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-failure-preserve").await; + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-failure-preserve".to_string()), + user_id: "user-callback-failure-preserve".to_string(), + amount_usd: 6.0, + pay_amount: Some(6.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-preserve".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-failure-preserve".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => panic!("wallet should be active"), + CreateWalletRechargeOrderOutcome::Existing(_) => panic!("order should be new"), + }; + + let first_failure = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-failure-preserve".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-preserve".to_string()), + amount_usd: 6.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-failure-preserve".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("first failure should resolve"); + assert!(matches!( + first_failure, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "callback amount mismatch" + )); + + let second_failure = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-failure-preserve".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-attacker-retry".to_string()), + amount_usd: 6.0, + pay_amount: Some(6.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-failure-preserve".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("second failure should resolve"); + assert!(matches!( + second_failure, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment gateway order mismatch" + )); + + let callback: (String, String, Option) = sqlx::query_as( + "SELECT gateway_order_id, status, error_message FROM payment_callbacks WHERE callback_key = ?", + ) + .bind("callback-failure-preserve") + .fetch_one(&pool) + .await + .expect("callback should load"); + assert_eq!(callback.0, "gateway-preserve"); + assert_eq!(callback.1, "failed"); + assert_eq!( + callback.2.as_deref(), + Some("payment gateway order mismatch") + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sqlite_payment_callback_registration_is_atomic_under_concurrency() { + let database_path = std::env::temp_dir().join(format!( + "aether-sqlite-payment-callback-race-{}.db", + uuid::Uuid::new_v4() + )); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .foreign_keys(true) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .busy_timeout(Duration::from_secs(30)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "user-callback-race").await; + let order = match repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-callback-race".to_string()), + user_id: "user-callback-race".to_string(), + amount_usd: 9.0, + pay_amount: Some(9.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-callback-race".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-callback-race".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created") + { + CreateWalletRechargeOrderOutcome::Created(order) => order, + CreateWalletRechargeOrderOutcome::WalletInactive => { + panic!("new wallet should be active") + } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new callback-race order should not already exist") + } + }; + + let callback_input = ProcessPaymentCallbackInput { + payment_method: "alipay".to_string(), + payment_provider: Some("alipay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-key-race".to_string(), + order_no: Some(order.order_no.clone()), + gateway_order_id: Some("gateway-callback-race".to_string()), + amount_usd: 9.0, + pay_amount: Some(9.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "payload-hash-race".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }; + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let first_repository = repository.clone(); + let first_barrier = barrier.clone(); + let first_input = callback_input.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository.process_payment_callback(first_input).await + }); + let second_repository = repository.clone(); + let second_barrier = barrier.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .process_payment_callback(callback_input) + .await + }); + barrier.wait().await; + + let outcomes = [ + first + .await + .expect("first callback task should join") + .expect("first callback should resolve"), + second + .await + .expect("second callback task should join") + .expect("second callback should resolve"), + ]; + let mut applied = 0; + let mut duplicate = 0; + for outcome in outcomes { + match outcome { + ProcessPaymentCallbackOutcome::Applied { + duplicate: false, .. + } => applied += 1, + ProcessPaymentCallbackOutcome::DuplicateProcessed { .. } + | ProcessPaymentCallbackOutcome::Applied { + duplicate: true, .. + } + | ProcessPaymentCallbackOutcome::AlreadyCredited { + duplicate: true, .. + } => duplicate += 1, + other => panic!("unexpected concurrent callback outcome: {other:?}"), + } + } + assert_eq!(applied, 1, "exactly one callback should apply the credit"); + assert_eq!(duplicate, 1, "the other callback should be a duplicate"); + + let wallet: (f64, f64) = + sqlx::query_as("SELECT balance, total_recharged FROM wallets WHERE id = ?") + .bind("wallet-callback-race") + .fetch_one(&pool) + .await + .expect("wallet should load"); + assert_eq!(wallet, (9.0, 9.0)); + let callback_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM payment_callbacks WHERE callback_key = ?") + .bind("callback-key-race") + .fetch_one(&pool) + .await + .expect("callback count should query"); + assert_eq!(callback_count, 1); + + pool.close().await; + let _ = std::fs::remove_file(&database_path); + let _ = std::fs::remove_file(format!("{}-wal", database_path.display())); + let _ = std::fs::remove_file(format!("{}-shm", database_path.display())); +} + +#[tokio::test] +async fn sqlite_payment_callback_binds_settlement_amount_before_usd_conversion() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let repository = SqliteWalletReadRepository::new(pool); + ensure_test_user(repository.pool(), "user-fee-callback").await; + let order = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-fee-callback".to_string()), + user_id: "user-fee-callback".to_string(), + amount_usd: 10.0, + pay_amount: Some(73.5), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-fee-callback".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-fee-callback".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created"); + assert!(matches!( + order, + CreateWalletRechargeOrderOutcome::Created(_) + )); + + let mismatched = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-fee-mismatch".to_string(), + order_no: Some("order-fee-callback".to_string()), + gateway_order_id: Some("gateway-fee-callback".to_string()), + amount_usd: 10.0, + pay_amount: Some(73.49), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.0), + payload_hash: "payload-fee-mismatch".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("mismatched callback should resolve"); + assert!(matches!( + mismatched, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "callback amount mismatch" + )); + + let wrong_currency = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-fee-wrong-currency".to_string(), + order_no: Some("order-fee-callback".to_string()), + gateway_order_id: Some("gateway-fee-callback".to_string()), + amount_usd: 10.0, + pay_amount: Some(73.5), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(7.0), + payload_hash: "payload-fee-wrong-currency".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("wrong-currency callback should resolve"); + assert!(matches!( + wrong_currency, + ProcessPaymentCallbackOutcome::Failed { ref error, .. } + if error == "payment currency mismatch" + )); + + let applied = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + callback_key: "callback-fee-applied".to_string(), + order_no: Some("order-fee-callback".to_string()), + gateway_order_id: Some("gateway-fee-callback".to_string()), + // Gateway callbacks derive USD from the fee-inclusive settlement + // amount. The stored order's USD amount remains the net credit. + amount_usd: 10.5, + pay_amount: Some(73.5000005), + pay_currency: Some("cny".to_string()), + // A callback may carry a conflicting provider-side rate. The + // order's checkout-time rate is the settlement proof and must + // remain unchanged. + exchange_rate: Some(99.0), + payload_hash: "payload-fee-applied".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("fee-inclusive callback should process"); + assert!(matches!( + applied, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + + let wallet = repository + .find(WalletLookupKey::UserId("user-fee-callback")) + .await + .expect("wallet should query") + .expect("wallet should exist"); + assert_eq!(wallet.balance, 10.0); + assert_eq!(wallet.total_recharged, 10.0); + let persisted_terms: (Option, Option, Option) = sqlx::query_as( + "SELECT pay_amount, pay_currency, exchange_rate FROM payment_orders WHERE order_no = ?", + ) + .bind("order-fee-callback") + .fetch_one(repository.pool()) + .await + .expect("settlement terms should query"); + assert_eq!( + persisted_terms, + (Some(73.5), Some("CNY".to_string()), Some(7.0)) + ); + + repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some(wallet.id), + user_id: "user-fee-callback".to_string(), + amount_usd: 5.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + payment_method: "manual".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "gateway-usd-fallback".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-usd-fallback".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("fallback order should be created"); + let fallback = repository + .process_payment_callback(ProcessPaymentCallbackInput { + payment_method: "manual".to_string(), + payment_provider: None, + payment_channel: None, + callback_key: "callback-usd-fallback".to_string(), + order_no: Some("order-usd-fallback".to_string()), + // The order already has a verified gateway id. A generic retry + // may identify it by order number without repeating that id. + gateway_order_id: None, + amount_usd: 5.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + payload_hash: "payload-usd-fallback".to_string(), + payload: json!({ "status": "paid" }), + signature_valid: true, + }) + .await + .expect("USD fallback callback should process"); + assert!(matches!( + fallback, + ProcessPaymentCallbackOutcome::Applied { .. } + )); + let persisted_gateway: (Option, String) = sqlx::query_as( + "SELECT gateway_order_id, gateway_response FROM payment_orders WHERE order_no = ?", + ) + .bind("order-usd-fallback") + .fetch_one(repository.pool()) + .await + .expect("fallback order should remain queryable"); + assert_eq!(persisted_gateway.0.as_deref(), Some("gateway-usd-fallback")); + assert_eq!( + serde_json::from_str::(&persisted_gateway.1) + .expect("gateway response should be JSON") + .get("gateway_order_id"), + Some(&json!("gateway-usd-fallback")) + ); +} + +#[tokio::test] +async fn sqlite_plan_purchase_rejects_missing_user_without_creating_wallet_or_order() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + let input = CreatePlanPurchaseOrderInput { + preferred_wallet_id: Some("missing-user-wallet".to_string()), + user_id: "missing-plan-user".to_string(), + amount_usd: 1.0, + pay_amount: 1.0, + pay_currency: "USD".to_string(), + exchange_rate: 1.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "missing-user-gateway".to_string(), + gateway_response: json!({"checkout": true}), + order_no: "missing-user-order".to_string(), + product_id: "missing-user-plan".to_string(), + product_snapshot: json!({ + "id": "missing-user-plan", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "unlimited", + "entitlements": [] + }), + expires_at_unix_secs: 4_102_444_800, + }; + let error = repository + .create_plan_purchase_order(input) + .await + .expect_err("missing user must be rejected"); + assert!(matches!(error, DataLayerError::InvalidInput(message) if message == "user not found")); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallets WHERE id = 'missing-user-wallet'", + ) + .fetch_one(&pool) + .await + .expect("wallet count should query"), + 0 + ); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM payment_orders WHERE order_no = 'missing-user-order'", + ) + .fetch_one(&pool) + .await + .expect("payment order count should query"), + 0 + ); +} + +#[tokio::test] +async fn sqlite_plan_purchase_rejects_preferred_wallet_id_owned_by_another_user() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_users(&pool, &["plan-wallet-owner", "plan-wallet-conflict"]).await; + let existing_wallet = repository + .initialize_auth_user_wallet("plan-wallet-owner", 0.0, false) + .await + .expect("owner wallet initialization should run") + .expect("owner wallet should exist"); + + let result = repository + .create_plan_purchase_order(CreatePlanPurchaseOrderInput { + preferred_wallet_id: Some(existing_wallet.id.clone()), + user_id: "plan-wallet-conflict".to_string(), + amount_usd: 1.0, + pay_amount: 1.0, + pay_currency: "USD".to_string(), + exchange_rate: 1.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "gateway-plan-wallet-conflict".to_string(), + gateway_response: json!({"checkout": true}), + order_no: "order-plan-wallet-conflict".to_string(), + product_id: "plan-wallet-conflict-product".to_string(), + product_snapshot: json!({ + "id": "plan-wallet-conflict-product", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "unlimited", + "entitlements": [] + }), + expires_at_unix_secs: 4_102_444_800, + }) + .await; + + assert!(matches!( + result, + Err(DataLayerError::InvalidInput(message)) + if message == "wallet identifier already belongs to another owner" + )); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallets WHERE user_id = 'plan-wallet-conflict'", + ) + .fetch_one(&pool) + .await + .expect("conflicting user wallet count should query"), + 0 + ); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM payment_orders WHERE order_no = 'order-plan-wallet-conflict'", + ) + .fetch_one(&pool) + .await + .expect("conflicting plan order count should query"), + 0 + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sqlite_plan_purchase_initializes_one_wallet_under_concurrency() { + let database_path = std::env::temp_dir().join(format!( + "aether-sqlite-plan-wallet-race-{}.db", + uuid::Uuid::new_v4() + )); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&database_path) + .create_if_missing(true) + .foreign_keys(true) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .busy_timeout(Duration::from_secs(30)); + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + ensure_test_user(&pool, "plan-wallet-race-user").await; + + let plan_snapshot = json!({ + "id": "plan-wallet-race-product", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "unlimited", + "entitlements": [] + }); + let first_input = CreatePlanPurchaseOrderInput { + preferred_wallet_id: Some("plan-wallet-race-first".to_string()), + user_id: "plan-wallet-race-user".to_string(), + amount_usd: 1.0, + pay_amount: 1.0, + pay_currency: "USD".to_string(), + exchange_rate: 1.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "gateway-plan-wallet-race-first".to_string(), + gateway_response: json!({"checkout": true}), + order_no: "order-plan-wallet-race-first".to_string(), + product_id: "plan-wallet-race-product".to_string(), + product_snapshot: plan_snapshot.clone(), + expires_at_unix_secs: 4_102_444_800, + }; + let second_input = CreatePlanPurchaseOrderInput { + preferred_wallet_id: Some("plan-wallet-race-second".to_string()), + gateway_order_id: "gateway-plan-wallet-race-second".to_string(), + order_no: "order-plan-wallet-race-second".to_string(), + ..first_input.clone() + }; + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let first_repository = repository.clone(); + let first_barrier = barrier.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository + .create_plan_purchase_order(first_input) + .await + }); + let second_repository = repository.clone(); + let second_barrier = barrier.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .create_plan_purchase_order(second_input) + .await + }); + barrier.wait().await; + + let first = first + .await + .expect("first plan task should join") + .expect("first plan task should resolve"); + let second = second + .await + .expect("second plan task should join") + .expect("second plan task should resolve"); + let first_wallet_id = match first { + CreatePlanPurchaseOrderOutcome::Created(order) => order.wallet_id, + other => panic!("first concurrent plan should be created, got {other:?}"), + }; + let second_wallet_id = match second { + CreatePlanPurchaseOrderOutcome::Created(order) => order.wallet_id, + other => panic!("second concurrent plan should be created, got {other:?}"), + }; + assert_eq!(first_wallet_id, second_wallet_id); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallets WHERE user_id = 'plan-wallet-race-user'", + ) + .fetch_one(&pool) + .await + .expect("wallet count should query"), + 1 + ); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM payment_orders WHERE user_id = 'plan-wallet-race-user' AND order_kind = 'plan_purchase'", + ) + .fetch_one(&pool) + .await + .expect("plan order count should query"), + 2 + ); + + pool.close().await; + let _ = std::fs::remove_file(database_path); +} + +#[tokio::test] +async fn sqlite_wallet_recharge_rejects_missing_user_without_creating_wallet_or_order() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool.clone()); + + let error = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("missing-recharge-wallet".to_string()), + user_id: "missing-recharge-user".to_string(), + amount_usd: 2.0, + pay_amount: Some(2.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "missing-recharge-gateway".to_string(), + gateway_response: json!({"checkout": true}), + order_no: "missing-recharge-order".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect_err("missing user must be rejected"); + assert!(matches!( + error, + DataLayerError::InvalidInput(message) if message == "user not found" + )); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM wallets WHERE id = 'missing-recharge-wallet'", + ) + .fetch_one(&pool) + .await + .expect("wallet count should query"), + 0 + ); + assert_eq!( + sqlx::query_scalar::<_, i64>( + "SELECT COUNT(*) FROM payment_orders WHERE order_no = 'missing-recharge-order'", + ) + .fetch_one(&pool) + .await + .expect("payment order count should query"), + 0 + ); +} + #[tokio::test] async fn sqlite_plan_purchase_blocks_duplicate_pending_active_period_order_and_manual_credit_fulfills( ) { @@ -600,6 +4238,9 @@ async fn sqlite_plan_purchase_blocks_duplicate_pending_active_period_order_and_m CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new bootstrap order should not already exist") + } }; let plan_snapshot = json!({ @@ -725,6 +4366,100 @@ INSERT INTO billing_plans ( .await .expect("wallet balance should query"); assert_eq!(wallet_balance, 0.0); + + // A malformed wallet_credit must abort fulfillment rather than silently + // activating the plan without delivering its promised balance. + let malformed_snapshot = json!({ + "id": "active-period-plan", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "unlimited", + "entitlements": [{ + "type": "wallet_credit", + "amount_usd": 5.0, + "balance_bucket": "not-a-wallet-bucket" + }] + }); + let valid_legacy_snapshot = json!({ + "id": "active-period-plan", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "unlimited", + "entitlements": [{ + "type": "wallet_credit", + "amount_usd": 5.0, + "balance_bucket": "gift" + }] + }); + let malformed_order = match repository + .create_plan_purchase_order(CreatePlanPurchaseOrderInput { + preferred_wallet_id: None, + user_id: "user-active-period-1".to_string(), + amount_usd: 2.0, + pay_amount: 2.0, + pay_currency: "USD".to_string(), + exchange_rate: 1.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "gateway-malformed-wallet-credit".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-malformed-wallet-credit".to_string(), + product_id: "active-period-plan".to_string(), + // Create through the validated boundary first; the malformed + // snapshot is installed below to simulate a legacy/corrupt row + // that predates the boundary validator. + product_snapshot: valid_legacy_snapshot, + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("malformed plan order should be persisted for fulfillment test") + { + CreatePlanPurchaseOrderOutcome::Created(order) => order, + other => panic!("malformed plan order should be created, got {other:?}"), + }; + sqlx::query("UPDATE payment_orders SET product_snapshot = ? WHERE id = ?") + .bind(malformed_snapshot.to_string()) + .bind(&malformed_order.id) + .execute(repository.pool()) + .await + .expect("malformed legacy snapshot should be installed"); + let credit_result = repository + .credit_admin_payment_order(CreditAdminPaymentOrderInput { + order_id: malformed_order.id.clone(), + gateway_order_id: Some("gateway-malformed-wallet-credit-paid".to_string()), + pay_amount: Some(2.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + gateway_response_patch: Some(json!({ "settled": true })), + operator_id: Some("admin-1".to_string()), + }) + .await; + assert!(matches!( + credit_result, + Err(DataLayerError::InvalidInput(_)) + )); + let malformed_status: String = + sqlx::query_scalar("SELECT status FROM payment_orders WHERE id = ?") + .bind(&malformed_order.id) + .fetch_one(repository.pool()) + .await + .expect("malformed order status should remain queryable"); + assert_eq!(malformed_status, "pending"); + let malformed_entitlements: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM user_plan_entitlements WHERE payment_order_id = ?", + ) + .bind(&malformed_order.id) + .fetch_one(repository.pool()) + .await + .expect("malformed entitlement count should query"); + assert_eq!(malformed_entitlements, 0); + let wallet_balance_after: f64 = sqlx::query_scalar("SELECT balance FROM wallets WHERE id = ?") + .bind("wallet-active-period-1") + .fetch_one(repository.pool()) + .await + .expect("wallet balance after rejected credit should query"); + assert_eq!(wallet_balance_after, 0.0); } #[tokio::test] @@ -775,6 +4510,9 @@ async fn sqlite_finds_reusable_pending_plan_purchase_order() { CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new pending-plan order should not already exist") + } }; let plan_snapshot = json!({ @@ -944,6 +4682,9 @@ async fn sqlite_plan_purchase_replaces_same_class_entitlements_on_manual_credit( CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new upgrade order should not already exist") + } }; let low_snapshot = json!({ @@ -1092,6 +4833,150 @@ INSERT INTO billing_plans ( assert_eq!(active_high_count, 1); } +#[tokio::test] +async fn sqlite_plan_replacement_stacks_usage_policies_unless_groups_match() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let repository = SqliteWalletReadRepository::new(pool); + sqlx::query("PRAGMA foreign_keys = OFF") + .execute(repository.pool()) + .await + .expect("foreign keys should be disabled for isolated replacement fixtures"); + + let fixtures = [ + ( + "usage-plain", + json!([{"type": "usage_policy", "policy_id": "weekly", "rules": []}]), + ), + ( + "usage-pro", + json!([{ + "type": "usage_policy", + "replacement_group": "pro-tier", + "rules": [] + }]), + ), + ( + "usage-team", + json!([{ + "type": "usage_policy", + "replacement_group": "team-tier", + "rules": [] + }]), + ), + ( + "daily-legacy", + json!([{"type": "daily_quota", "daily_quota_usd": 10.0}]), + ), + ]; + for (id, entitlements) in fixtures { + sqlx::query( + r#" +INSERT INTO user_plan_entitlements ( + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +) VALUES (?, 'replacement-user', ?, ?, 'active', 1, 4102444800, ?, 1, 1) + "#, + ) + .bind(id) + .bind(format!("plan-{id}")) + .bind(format!("order-{id}")) + .bind(entitlements.to_string()) + .execute(repository.pool()) + .await + .expect("entitlement fixture should seed"); + } + + let mut tx = repository.pool().begin().await.expect("tx should start"); + replace_matching_plan_entitlements_sqlite( + &mut tx, + "replacement-user", + &json!({ + "entitlements": [{"type": "usage_policy", "policy_id": "five-hour", "rules": []}] + }), + 100, + ) + .await + .expect("ungrouped usage policy replacement should run"); + tx.commit().await.expect("tx should commit"); + let active_count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM user_plan_entitlements WHERE user_id = ? AND status = 'active'", + ) + .bind("replacement-user") + .fetch_one(repository.pool()) + .await + .expect("active entitlement count should query"); + assert_eq!(active_count, 4); + + let mut tx = repository.pool().begin().await.expect("tx should start"); + replace_matching_plan_entitlements_sqlite( + &mut tx, + "replacement-user", + &json!({ + "entitlements": [{ + "type": "usage_policy", + "replacement_group": "pro-tier", + "rules": [] + }] + }), + 200, + ) + .await + .expect("grouped usage policy replacement should run"); + tx.commit().await.expect("tx should commit"); + + let statuses = sqlx::query_as::<_, (String, String)>( + "SELECT id, status FROM user_plan_entitlements ORDER BY id", + ) + .fetch_all(repository.pool()) + .await + .expect("entitlement statuses should query") + .into_iter() + .collect::>(); + assert_eq!( + statuses.get("usage-pro").map(String::as_str), + Some("replaced") + ); + assert_eq!( + statuses.get("usage-plain").map(String::as_str), + Some("active") + ); + assert_eq!( + statuses.get("usage-team").map(String::as_str), + Some("active") + ); + assert_eq!( + statuses.get("daily-legacy").map(String::as_str), + Some("active") + ); + + let mut tx = repository.pool().begin().await.expect("tx should start"); + replace_matching_plan_entitlements_sqlite( + &mut tx, + "replacement-user", + &json!({ + "entitlements": [{"type": "daily_quota", "daily_quota_usd": 50.0}] + }), + 300, + ) + .await + .expect("legacy daily quota replacement should run"); + tx.commit().await.expect("tx should commit"); + let daily_status: String = + sqlx::query_scalar("SELECT status FROM user_plan_entitlements WHERE id = 'daily-legacy'") + .fetch_one(repository.pool()) + .await + .expect("daily entitlement status should query"); + assert_eq!(daily_status, "replaced"); +} + #[tokio::test] async fn sqlite_plan_purchase_respects_lifetime_purchase_limit() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1140,6 +5025,9 @@ async fn sqlite_plan_purchase_respects_lifetime_purchase_limit() { CreateWalletRechargeOrderOutcome::WalletInactive => { panic!("new wallet should be active") } + CreateWalletRechargeOrderOutcome::Existing(_) => { + panic!("new lifetime order should not already exist") + } }; let plan_snapshot = json!({ @@ -1215,9 +5103,9 @@ INSERT INTO billing_plans ( order_no: Some("order-plan-lifetime-1".to_string()), gateway_order_id: Some("gateway-plan-lifetime-1".to_string()), amount_usd: 1.0, - pay_amount: Some(7.2), - pay_currency: Some("CNY".to_string()), - exchange_rate: Some(7.2), + pay_amount: Some(7.2000005), + pay_currency: Some("cny".to_string()), + exchange_rate: Some(99.0), payload_hash: "payload-plan-lifetime-1".to_string(), payload: json!({ "trade_status": "TRADE_SUCCESS" }), signature_valid: true, @@ -1228,6 +5116,18 @@ INSERT INTO billing_plans ( callback, ProcessPaymentCallbackOutcome::Applied { .. } )); + let persisted_plan_terms: (f64, String, f64) = sqlx::query_as( + "SELECT pay_amount, pay_currency, exchange_rate FROM payment_orders WHERE order_no = ?", + ) + .bind("order-plan-lifetime-1") + .fetch_one(repository.pool()) + .await + .expect("plan settlement terms should query"); + assert_eq!( + persisted_plan_terms, + (7.2, "CNY".to_string(), 7.2), + "callback data must not overwrite checkout-time settlement terms" + ); let entitlement_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM user_plan_entitlements WHERE user_id = ? AND plan_id = ?", @@ -1372,6 +5272,79 @@ impl SqliteWalletReadRepository { } } +#[tokio::test] +async fn sqlite_recharge_checkout_update_rejects_expired_order() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteWalletReadRepository::new(pool); + + sqlx::query( + "INSERT INTO users (id, username, email, auth_source, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("user-expired-checkout") + .bind("Expired Checkout") + .bind("expired-checkout@example.com") + .bind("local") + .bind(1_i64) + .bind(1_i64) + .execute(repository.pool()) + .await + .expect("user should seed"); + + let created = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-expired-checkout".to_string()), + user_id: "user-expired-checkout".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "order-expired-checkout".to_string(), + gateway_response: json!({ + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending" + }), + order_no: "order-expired-checkout".to_string(), + expires_at_unix_secs: 1, + }) + .await + .expect("expired recharge order should be creatable for regression setup"); + let CreateWalletRechargeOrderOutcome::Created(order) = created else { + panic!("expected a newly created recharge order"); + }; + + let result = repository + .update_wallet_recharge_checkout(UpdateWalletRechargeCheckoutInput { + order_id: order.id.clone(), + gateway_order_id: "provider-expired-checkout".to_string(), + gateway_response: json!({ + "order_kind": "wallet_recharge", + "payment_url": "https://pay.example.test/expired" + }), + }) + .await + .expect("expired checkout update should resolve"); + assert!(matches!(result, WalletMutationOutcome::Invalid(_))); + + let persisted: (Option, String) = + sqlx::query_as("SELECT gateway_order_id, status FROM payment_orders WHERE id = ?") + .bind(&order.id) + .fetch_one(repository.pool()) + .await + .expect("expired recharge order should remain queryable"); + assert_eq!(persisted.0.as_deref(), Some("order-expired-checkout")); + assert_eq!(persisted.1, "pending"); +} + async fn seed_rows(pool: &sqlx::SqlitePool) { sqlx::query( r#" diff --git a/crates/aether-data/contracts/Cargo.toml b/crates/aether-data/contracts/Cargo.toml index a9607a8c8..107a663d7 100644 --- a/crates/aether-data/contracts/Cargo.toml +++ b/crates/aether-data/contracts/Cargo.toml @@ -7,14 +7,20 @@ repository.workspace = true description = "Shared data contracts and repository traits for Aether Rust services" [dependencies] +aether-contracts.workspace = true aether-ai-formats.workspace = true aether-routing-core.workspace = true async-trait.workspace = true +base64.workspace = true +bcrypt.workspace = true chrono.workspace = true +chrono-tz.workspace = true serde.workspace = true serde_json.workspace = true sha2.workspace = true thiserror.workspace = true +url.workspace = true +uuid.workspace = true [dev-dependencies] tokio.workspace = true diff --git a/crates/aether-data/contracts/src/repository/auth.rs b/crates/aether-data/contracts/src/repository/auth.rs index df5b2aff1..3fedc5c0b 100644 --- a/crates/aether-data/contracts/src/repository/auth.rs +++ b/crates/aether-data/contracts/src/repository/auth.rs @@ -1,5 +1,9 @@ use async_trait::async_trait; +fn redacted_optional_secret(value: &Option) -> Option<&'static str> { + value.as_ref().map(|_| "[REDACTED]") +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredAuthApiKeySnapshot { pub user_id: String, @@ -121,7 +125,7 @@ impl StoredAuthApiKeySnapshot { return false; } if let Some(expires_at_unix_secs) = self.api_key_expires_at_unix_secs { - if expires_at_unix_secs < now_unix_secs { + if expires_at_unix_secs <= now_unix_secs { return false; } } @@ -350,7 +354,7 @@ pub async fn read_resolved_auth_api_key_snapshot_by_user_api_key_ids( .await } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredAuthApiKeyExportRecord { pub user_id: String, pub api_key_id: String, @@ -377,6 +381,22 @@ pub struct StoredAuthApiKeyExportRecord { pub is_standalone: bool, } +impl std::fmt::Debug for StoredAuthApiKeyExportRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredAuthApiKeyExportRecord") + .field("user_id", &self.user_id) + .field("api_key_id", &self.api_key_id) + .field("key_hash", &self.key_hash) + .field( + "key_encrypted", + &redacted_optional_secret(&self.key_encrypted), + ) + .field("is_standalone", &self.is_standalone) + .finish_non_exhaustive() + } +} + impl StoredAuthApiKeyExportRecord { #[allow(clippy::too_many_arguments)] pub fn new( @@ -497,7 +517,7 @@ pub struct StandaloneApiKeyExportListQuery { pub is_active: Option, } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct CreateUserApiKeyRecord { pub user_id: String, pub api_key_id: String, @@ -511,6 +531,7 @@ pub struct CreateUserApiKeyRecord { pub rate_limit: i32, pub concurrent_limit: Option, pub force_capabilities: Option, + pub feature_settings: Option, pub is_active: bool, pub expires_at_unix_secs: Option, pub auto_delete_on_expiry: bool, @@ -519,17 +540,61 @@ pub struct CreateUserApiKeyRecord { pub total_cost_usd: f64, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for CreateUserApiKeyRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreateUserApiKeyRecord") + .field("user_id", &self.user_id) + .field("api_key_id", &self.api_key_id) + .field("key_hash", &self.key_hash) + .field( + "key_encrypted", + &redacted_optional_secret(&self.key_encrypted), + ) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct UpdateUserApiKeyBasicRecord { pub user_id: String, pub api_key_id: String, + pub key_encrypted: Option, + /// Whether `key_encrypted` is an explicit replacement, including an explicit `NULL`. + /// Ordinary callers should leave this false to retain the existing value. + pub key_encrypted_present: bool, pub name: Option, + /// Whether `name` is an explicit replacement, including an explicit `NULL`. + pub name_present: bool, pub rate_limit: Option, + /// Whether `rate_limit` is an explicit replacement, including an explicit `NULL`. + pub rate_limit_present: bool, pub concurrent_limit: Option, + /// Whether `concurrent_limit` is an explicit replacement, including an explicit `NULL`. + pub concurrent_limit_present: bool, pub ip_rules: Option>>, + /// `Some(Some(value))` replaces the settings; `Some(None)` clears them; `None` leaves them + /// unchanged. Keeping this patch in the basic mutation record lets repositories apply the + /// complete user-key update in one atomic write. + pub feature_settings: Option>, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for UpdateUserApiKeyBasicRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("UpdateUserApiKeyBasicRecord") + .field("user_id", &self.user_id) + .field("api_key_id", &self.api_key_id) + .field( + "key_encrypted", + &redacted_optional_secret(&self.key_encrypted), + ) + .field("key_encrypted_present", &self.key_encrypted_present) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq)] pub struct CreateStandaloneApiKeyRecord { pub user_id: String, pub api_key_id: String, @@ -551,10 +616,35 @@ pub struct CreateStandaloneApiKeyRecord { pub total_cost_usd: f64, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for CreateStandaloneApiKeyRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreateStandaloneApiKeyRecord") + .field("user_id", &self.user_id) + .field("api_key_id", &self.api_key_id) + .field("key_hash", &self.key_hash) + .field( + "key_encrypted", + &redacted_optional_secret(&self.key_encrypted), + ) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct UpdateStandaloneApiKeyBasicRecord { pub api_key_id: String, + pub key_encrypted: Option, + /// Whether `key_encrypted` is an explicit replacement, including an explicit `NULL`. + /// Ordinary callers should leave this false to retain the existing value. + pub key_encrypted_present: bool, pub name: Option, + /// Whether `name` is an explicit replacement, including an explicit `NULL`. + /// Ordinary callers should leave this false to retain the existing value. + pub name_present: bool, + /// `Some(Some(value))` sets the capability map; `Some(None)` clears it; `None` leaves it + /// unchanged. + pub force_capabilities: Option>, pub rate_limit_present: bool, pub rate_limit: Option, pub concurrent_limit_present: bool, @@ -569,6 +659,48 @@ pub struct UpdateStandaloneApiKeyBasicRecord { pub auto_delete_on_expiry: bool, } +impl std::fmt::Debug for UpdateStandaloneApiKeyBasicRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("UpdateStandaloneApiKeyBasicRecord") + .field("api_key_id", &self.api_key_id) + .field( + "key_encrypted", + &redacted_optional_secret(&self.key_encrypted), + ) + .field("key_encrypted_present", &self.key_encrypted_present) + .finish_non_exhaustive() + } +} + +/// Replace only the recoverable API-key ciphertext when the complete immutable identity and the +/// exact ciphertext observed by the caller still match. This is intentionally narrower than the +/// ordinary admin update records so lazy envelope migration cannot overwrite a concurrent secret +/// restore or move ciphertext between owners, scopes, hashes, or key IDs. +#[derive(Clone, PartialEq, Eq)] +pub struct CompareAndSwapAuthApiKeyCiphertext { + pub user_id: String, + pub api_key_id: String, + pub key_hash: String, + pub is_standalone: bool, + pub expected_key_encrypted: String, + pub key_encrypted: String, +} + +impl std::fmt::Debug for CompareAndSwapAuthApiKeyCiphertext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CompareAndSwapAuthApiKeyCiphertext") + .field("user_id", &self.user_id) + .field("api_key_id", &self.api_key_id) + .field("key_hash", &self.key_hash) + .field("is_standalone", &self.is_standalone) + .field("expected_key_encrypted", &"[REDACTED]") + .field("key_encrypted", &"[REDACTED]") + .finish() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AuthApiKeyLookupKey<'a> { KeyHash(&'a str), @@ -650,6 +782,20 @@ pub trait AuthApiKeyReadRepository: Send + Sync { pub trait AuthApiKeyWriteRepository: Send + Sync { async fn touch_last_used_at(&self, api_key_id: &str) -> Result; + /// Synchronize an authoritative user snapshot for repositories used by gateway tests. + /// + /// Production database repositories validate the API-key owner in the same transaction as + /// key creation and deliberately keep the default no-op implementation. Test repositories + /// may override this hook, but must derive owner state exclusively from `StoredUserAuthRecord` + /// rather than from an API-key mutation request. + async fn synchronize_user_api_key_owner_for_tests( + &self, + user: &crate::repository::users::StoredUserAuthRecord, + ) -> Result<(), crate::DataLayerError> { + let _ = user; + Ok(()) + } + async fn create_user_api_key( &self, record: CreateUserApiKeyRecord, @@ -660,16 +806,53 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { record: CreateStandaloneApiKeyRecord, ) -> Result, crate::DataLayerError>; + async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + let _ = mutation; + Err(crate::DataLayerError::InvalidInput( + "atomic API-key ciphertext migration is not available".to_string(), + )) + } + async fn update_user_api_key_basic( &self, record: UpdateUserApiKeyBasicRecord, ) -> Result, crate::DataLayerError>; + /// Apply a user-owned key update only while the key remains unlocked. + /// Implementations must test ownership, non-standalone status, and + /// `is_locked = false` in the same atomic write as the mutation. + async fn update_user_api_key_basic_if_unlocked( + &self, + record: UpdateUserApiKeyBasicRecord, + ) -> Result, crate::DataLayerError> { + let _ = record; + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key update is not available".to_string(), + )) + } + async fn update_standalone_api_key_basic( &self, record: UpdateStandaloneApiKeyBasicRecord, ) -> Result, crate::DataLayerError>; + /// Restore an API-key export row only when its complete exported post-state still matches + /// `expected`. Implementations must perform the compare and update atomically so a failed + /// import cannot overwrite a concurrent administrator change. + async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + let _ = (expected, restored); + Err(crate::DataLayerError::InvalidInput( + "atomic API key restore is not available".to_string(), + )) + } + async fn set_user_api_key_active( &self, user_id: &str, @@ -677,6 +860,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { is_active: bool, ) -> Result, crate::DataLayerError>; + async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, api_key_id, is_active); + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key status update is not available".to_string(), + )) + } + async fn set_standalone_api_key_active( &self, api_key_id: &str, @@ -697,6 +892,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { allowed_providers: Option>, ) -> Result, crate::DataLayerError>; + async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, api_key_id, allowed_providers); + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key provider update is not available".to_string(), + )) + } + async fn set_user_api_key_force_capabilities( &self, user_id: &str, @@ -704,6 +911,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { force_capabilities: Option, ) -> Result, crate::DataLayerError>; + async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, api_key_id, force_capabilities); + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key capability update is not available".to_string(), + )) + } + async fn set_user_api_key_feature_settings( &self, user_id: &str, @@ -711,6 +930,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { feature_settings: Option, ) -> Result, crate::DataLayerError>; + async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, api_key_id, feature_settings); + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key feature update is not available".to_string(), + )) + } + async fn set_api_key_usage_totals( &self, api_key_id: &str, @@ -725,6 +956,17 @@ pub trait AuthApiKeyWriteRepository: Send + Sync { api_key_id: &str, ) -> Result; + async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + let _ = (user_id, api_key_id); + Err(crate::DataLayerError::InvalidInput( + "atomic unlocked user API key deletion is not available".to_string(), + )) + } + async fn delete_standalone_api_key( &self, api_key_id: &str, @@ -762,7 +1004,9 @@ fn parse_string_list_value( field_name: &str, ) -> Result>, crate::DataLayerError> { match value { - serde_json::Value::Null => Ok(None), + serde_json::Value::Null => Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains JSON null; use SQL NULL for an unset policy" + ))), serde_json::Value::Array(array) => parse_string_list_array(array, field_name).map(Some), serde_json::Value::String(raw) => parse_embedded_string_list(raw, field_name), _ => Err(crate::DataLayerError::UnexpectedValue(format!( @@ -776,8 +1020,15 @@ fn parse_embedded_string_list( field_name: &str, ) -> Result>, crate::DataLayerError> { let raw = raw.trim(); - if raw.is_empty() || raw.eq_ignore_ascii_case("null") { - return Ok(None); + if raw.is_empty() { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty string" + ))); + } + if raw.eq_ignore_ascii_case("null") { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains stringified JSON null; use SQL NULL for an unset policy" + ))); } if let Ok(decoded) = serde_json::from_str::(raw) { @@ -799,9 +1050,12 @@ fn parse_string_list_array( ))); }; let item = item.trim(); - if !item.is_empty() { - items.push(item.to_string()); + if item.is_empty() { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty item" + ))); } + items.push(item.to_string()); } Ok(items) } @@ -826,10 +1080,80 @@ mod tests { use super::{ read_resolved_auth_api_key_snapshot_by_key_hash, read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyLookupKey, - ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeyExportRecord, - StoredAuthApiKeySnapshot, + CompareAndSwapAuthApiKeyCiphertext, ResolvedAuthApiKeySnapshot, + ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, }; + #[test] + fn api_key_record_debug_output_redacts_recoverable_ciphertext() { + let ciphertext = "debug-secret-api-key-ciphertext"; + let replacement = "debug-secret-api-key-replacement"; + let record = StoredAuthApiKeyExportRecord::new( + "user-1".to_string(), + "key-1".to_string(), + "key-hash".to_string(), + Some(ciphertext.to_string()), + Some("test key".to_string()), + None, + None, + None, + None, + None, + None, + true, + None, + false, + 0, + 0, + 0.0, + false, + ) + .expect("API key export record should build"); + let mutation = CompareAndSwapAuthApiKeyCiphertext { + user_id: "user-1".to_string(), + api_key_id: "key-1".to_string(), + key_hash: "key-hash".to_string(), + is_standalone: false, + expected_key_encrypted: ciphertext.to_string(), + key_encrypted: replacement.to_string(), + }; + + for rendered in [format!("{record:?}"), format!("{mutation:?}")] { + assert!(!rendered.contains(ciphertext)); + assert!(!rendered.contains(replacement)); + assert!(rendered.contains("[REDACTED]")); + } + } + + #[test] + fn stored_security_lists_distinguish_sql_null_from_malformed_json_null() { + assert_eq!( + super::parse_string_list(None, "api_keys.allowed_providers") + .expect("SQL NULL should remain an unset policy"), + None + ); + assert!(super::parse_string_list( + Some(serde_json::Value::Null), + "api_keys.allowed_providers" + ) + .is_err()); + assert!(super::parse_string_list( + Some(serde_json::json!("null")), + "api_keys.allowed_providers" + ) + .is_err()); + assert!(super::parse_string_list( + Some(serde_json::json!([" "])), + "api_keys.allowed_providers" + ) + .is_err()); + assert_eq!( + super::parse_string_list(Some(serde_json::json!([])), "api_keys.allowed_providers") + .expect("an intentional empty policy should preserve its existing semantics"), + Some(Vec::new()) + ); + } + #[test] fn api_format_policy_intersection_preserves_companion_scope() { assert_eq!( @@ -988,6 +1312,8 @@ mod tests { ) .expect("snapshot should build"); + assert!(snapshot.is_currently_usable(99)); + assert!(!snapshot.is_currently_usable(100)); assert!(!snapshot.is_currently_usable(101)); } diff --git a/crates/aether-data/contracts/src/repository/auth_modules.rs b/crates/aether-data/contracts/src/repository/auth_modules.rs index 8755c07d5..dfa1b34fb 100644 --- a/crates/aether-data/contracts/src/repository/auth_modules.rs +++ b/crates/aether-data/contracts/src/repository/auth_modules.rs @@ -1,6 +1,10 @@ use async_trait::async_trait; -#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +fn redacted_optional_secret(value: &Option) -> Option<&'static str> { + value.as_ref().map(|_| "[REDACTED]") +} + +#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredOAuthProviderModuleConfig { pub provider_type: String, pub display_name: String, @@ -9,6 +13,22 @@ pub struct StoredOAuthProviderModuleConfig { pub redirect_uri: String, } +impl std::fmt::Debug for StoredOAuthProviderModuleConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredOAuthProviderModuleConfig") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("client_id", &self.client_id) + .field( + "client_secret_encrypted", + &redacted_optional_secret(&self.client_secret_encrypted), + ) + .field("redirect_uri", &self.redirect_uri) + .finish() + } +} + impl StoredOAuthProviderModuleConfig { pub fn new( provider_type: String, @@ -37,7 +57,7 @@ impl StoredOAuthProviderModuleConfig { } } -#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredLdapModuleConfig { pub server_url: String, pub bind_dn: String, @@ -53,6 +73,62 @@ pub struct StoredLdapModuleConfig { pub connect_timeout: Option, } +impl std::fmt::Debug for StoredLdapModuleConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredLdapModuleConfig") + .field("server_url", &self.server_url) + .field("bind_dn", &self.bind_dn) + .field( + "bind_password_encrypted", + &redacted_optional_secret(&self.bind_password_encrypted), + ) + .field("base_dn", &self.base_dn) + .field("user_search_filter", &self.user_search_filter) + .field("username_attr", &self.username_attr) + .field("email_attr", &self.email_attr) + .field("display_name_attr", &self.display_name_attr) + .field("is_enabled", &self.is_enabled) + .field("is_exclusive", &self.is_exclusive) + .field("use_starttls", &self.use_starttls) + .field("connect_timeout", &self.connect_timeout) + .finish() + } +} + +/// Explicit mutation semantics for the LDAP bind password. +/// +/// The password is deliberately kept separate from [`StoredLdapModuleConfig`] updates so a +/// caller that only changes non-secret fields cannot accidentally write a stale ciphertext back +/// to storage. +#[derive(Clone, PartialEq, Eq)] +pub enum LdapBindPasswordUpdate { + Preserve, + Set(String), + Clear, +} + +impl std::fmt::Debug for LdapBindPasswordUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Preserve => formatter.write_str("Preserve"), + Self::Set(_) => formatter.write_str("Set([REDACTED])"), + Self::Clear => formatter.write_str("Clear"), + } + } +} + +// The successful branch intentionally returns the complete persisted +// configuration so callers can continue with the exact CAS snapshot. Boxing +// it would change this public repository contract and add needless allocation +// on the normal (successful) path. +#[allow(clippy::large_enum_variant)] +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CompareAndSwapLdapConfigResult { + Applied(StoredLdapModuleConfig), + Conflict, +} + #[async_trait] pub trait AuthModuleReadRepository: Send + Sync { async fn list_enabled_oauth_providers( @@ -66,8 +142,80 @@ pub trait AuthModuleReadRepository: Send + Sync { #[async_trait] pub trait AuthModuleWriteRepository: Send + Sync { - async fn upsert_ldap_config( + /// Atomically create or replace the singleton LDAP configuration. + /// + /// `expected` is the complete snapshot observed by the caller. `None` means the caller + /// expects the singleton not to exist. Implementations must compare every persisted config + /// field, including the encrypted password, before applying the replacement. The password + /// field in `replacement` is never authoritative; only `bind_password_update` controls the + /// stored secret. + async fn compare_and_swap_ldap_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, crate::DataLayerError>; + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result; + + /// Delete the singleton LDAP configuration only when every persisted field still matches the + /// supplied snapshot. This is used by aggregate-import compensation and must not remove a + /// configuration that another operation changed after it was created. + async fn delete_ldap_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result; + + async fn compare_and_swap_ldap_bind_password( + &self, + _expected: &str, + _replacement: &str, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "LDAP bind password compare-and-swap is not supported by this repository".to_string(), + )) + } +} + +#[cfg(test)] +mod tests { + use super::{LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig}; + + #[test] + fn auth_module_debug_output_redacts_encrypted_secrets() { + let oauth_secret = "debug-secret-oauth-ciphertext"; + let oauth = StoredOAuthProviderModuleConfig::new( + "linuxdo".to_string(), + "Linux.do".to_string(), + "client-id".to_string(), + Some(oauth_secret.to_string()), + "https://example.com/callback".to_string(), + ) + .expect("OAuth module config should build"); + let ldap_secret = "debug-secret-ldap-ciphertext"; + let ldap = StoredLdapModuleConfig { + server_url: "ldaps://ldap.example.com".to_string(), + bind_dn: "cn=admin,dc=example,dc=com".to_string(), + bind_password_encrypted: Some(ldap_secret.to_string()), + base_dn: "dc=example,dc=com".to_string(), + user_search_filter: None, + username_attr: None, + email_attr: None, + display_name_attr: None, + is_enabled: true, + is_exclusive: false, + use_starttls: false, + connect_timeout: Some(5), + }; + + for (rendered, secret) in [ + (format!("{oauth:?}"), oauth_secret), + (format!("{ldap:?}"), ldap_secret), + ( + format!("{:?}", LdapBindPasswordUpdate::Set(ldap_secret.to_string())), + ldap_secret, + ), + ] { + assert!(!rendered.contains(secret)); + assert!(rendered.contains("[REDACTED]")); + } + } } diff --git a/crates/aether-data/contracts/src/repository/background_tasks/types.rs b/crates/aether-data/contracts/src/repository/background_tasks/types.rs index 2fcff266d..c07681c80 100644 --- a/crates/aether-data/contracts/src/repository/background_tasks/types.rs +++ b/crates/aether-data/contracts/src/repository/background_tasks/types.rs @@ -3,6 +3,45 @@ use std::collections::BTreeMap; use async_trait::async_trait; use serde_json::Value; +const BACKGROUND_TASK_DEFAULT_ERROR_CODE: &str = "background_task_failed"; +const BACKGROUND_TASK_UNCLASSIFIED_EVENT: &str = "unclassified_event"; +const MAX_BACKGROUND_TASK_METADATA_FIELDS: usize = 48; +const SAFE_BACKGROUND_TASK_METADATA_FIELDS: &[&str] = &[ + "automatic_deletions", + "bytes", + "compression", + "created_count", + "deleted_endpoints", + "deleted_keys", + "encryption", + "error_code", + "export_version", + "exported_at", + "failed", + "import_kind", + "legacy_encrypted_copies_created", + "legacy_encrypted_copies_verified", + "legacy_plaintext_objects_deleted", + "legacy_plaintext_objects_retained", + "object_cleanup_mode", + "partition", + "provider_type", + "replaced_count", + "retention_cleanup_candidates", + "scheduled_slot", + "scope", + "sha256", + "stage", + "status", + "success", + "total", + "total_endpoints", + "total_keys", + "trigger", + "versioned_storage_cleanup_notice", + "versioned_storage_cleanup_required", +]; + #[derive( Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize, )] @@ -101,6 +140,18 @@ pub struct StoredBackgroundTaskRun { pub updated_at_unix_secs: u64, } +impl StoredBackgroundTaskRun { + pub fn sanitize_persisted_data(&mut self) { + self.owner_instance = None; + self.created_by = sanitize_background_task_actor(self.created_by.take()); + self.progress_message = None; + self.payload_json = sanitize_background_task_metadata(self.payload_json.take()); + self.result_json = sanitize_background_task_metadata(self.result_json.take()); + self.error_message = + sanitize_background_task_error_code(self.status, self.error_message.take()); + } +} + #[derive(Debug, Clone, PartialEq)] pub struct UpsertBackgroundTaskRun { pub id: String, @@ -125,6 +176,16 @@ pub struct UpsertBackgroundTaskRun { } impl UpsertBackgroundTaskRun { + pub fn sanitize_for_persistence(&mut self) { + self.owner_instance = None; + self.created_by = sanitize_background_task_actor(self.created_by.take()); + self.progress_message = None; + self.payload_json = sanitize_background_task_metadata(self.payload_json.take()); + self.result_json = sanitize_background_task_metadata(self.result_json.take()); + self.error_message = + sanitize_background_task_error_code(self.status, self.error_message.take()); + } + pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.id.trim().is_empty() || self.task_key.trim().is_empty() @@ -143,7 +204,8 @@ impl UpsertBackgroundTaskRun { Ok(()) } - pub fn into_stored(self) -> StoredBackgroundTaskRun { + pub fn into_stored(mut self) -> StoredBackgroundTaskRun { + self.sanitize_for_persistence(); StoredBackgroundTaskRun { id: self.id, task_key: self.task_key, @@ -178,6 +240,14 @@ pub struct StoredBackgroundTaskEvent { pub created_at_unix_secs: u64, } +impl StoredBackgroundTaskEvent { + pub fn sanitize_persisted_data(&mut self) { + self.event_type = sanitize_background_task_event_type(&self.event_type); + self.message = self.event_type.clone(); + self.payload_json = sanitize_background_task_metadata(self.payload_json.take()); + } +} + #[derive(Debug, Clone, PartialEq)] pub struct UpsertBackgroundTaskEvent { pub id: String, @@ -189,6 +259,12 @@ pub struct UpsertBackgroundTaskEvent { } impl UpsertBackgroundTaskEvent { + pub fn sanitize_for_persistence(&mut self) { + self.event_type = sanitize_background_task_event_type(&self.event_type); + self.message = self.event_type.clone(); + self.payload_json = sanitize_background_task_metadata(self.payload_json.take()); + } + pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.id.trim().is_empty() || self.run_id.trim().is_empty() @@ -202,7 +278,8 @@ impl UpsertBackgroundTaskEvent { Ok(()) } - pub fn into_stored(self) -> StoredBackgroundTaskEvent { + pub fn into_stored(mut self) -> StoredBackgroundTaskEvent { + self.sanitize_for_persistence(); StoredBackgroundTaskEvent { id: self.id, run_id: self.run_id, @@ -214,6 +291,199 @@ impl UpsertBackgroundTaskEvent { } } +fn sanitize_background_task_error_code( + status: BackgroundTaskStatus, + value: Option, +) -> Option { + if status != BackgroundTaskStatus::Failed { + return None; + } + let value = value?.trim().to_ascii_lowercase(); + let code = match value.as_str() { + "background_task_failed" + | "background_task_panicked" + | "provider_delete_failed" + | "provider_oauth_batch_import_failed" + | "s3_backup_failed" + | "s3_backup_slot_record_failed" => value, + _ => BACKGROUND_TASK_DEFAULT_ERROR_CODE.to_string(), + }; + Some(code) +} + +fn sanitize_background_task_event_type(value: &str) -> String { + let value = value.trim().to_ascii_lowercase(); + match value.as_str() { + "cancel_requested" | "failed" | "queued" | "running" | "skipped" | "succeeded" + | "worker_boot" => value, + _ => BACKGROUND_TASK_UNCLASSIFIED_EVENT.to_string(), + } +} + +fn sanitize_background_task_actor(value: Option) -> Option { + let value = value?.trim().to_ascii_lowercase(); + matches!(value.as_str(), "admin" | "scheduler" | "system").then_some(value) +} + +fn sanitize_background_task_metadata(value: Option) -> Option { + let Value::Object(object) = value? else { + return None; + }; + let mut sanitized = serde_json::Map::new(); + for (key, value) in object.into_iter().take(MAX_BACKGROUND_TASK_METADATA_FIELDS) { + let normalized_key = key.trim().to_ascii_lowercase(); + if !SAFE_BACKGROUND_TASK_METADATA_FIELDS.contains(&normalized_key.as_str()) { + continue; + } + let Some(value) = sanitize_background_task_metadata_value(&normalized_key, value) else { + continue; + }; + sanitized.insert(normalized_key, value); + } + (!sanitized.is_empty()).then_some(Value::Object(sanitized)) +} + +fn sanitize_background_task_metadata_value(key: &str, value: Value) -> Option { + match key { + "automatic_deletions" + | "bytes" + | "created_count" + | "deleted_endpoints" + | "deleted_keys" + | "failed" + | "legacy_encrypted_copies_created" + | "legacy_encrypted_copies_verified" + | "legacy_plaintext_objects_deleted" + | "legacy_plaintext_objects_retained" + | "partition" + | "replaced_count" + | "retention_cleanup_candidates" + | "success" + | "total" + | "total_endpoints" + | "total_keys" => value.is_u64().then_some(value), + "versioned_storage_cleanup_required" => value.is_boolean().then_some(value), + "error_code" => value + .as_str() + .and_then(sanitize_background_task_metadata_error_code) + .map(Value::String), + "scope" => sanitize_background_task_metadata_enum(value, &["config", "data", "users"]), + "compression" => sanitize_background_task_metadata_enum(value, &["zstd"]), + "encryption" => sanitize_background_task_metadata_enum(value, &["aes-256-gcm-v2"]), + "trigger" => sanitize_background_task_metadata_enum(value, &["manual", "scheduled"]), + "import_kind" => sanitize_background_task_metadata_enum( + value, + &["agent_identity", "cookie_authorize", "oauth_batch"], + ), + "provider_type" => sanitize_background_task_metadata_enum( + value, + &[ + "antigravity", + "chatgpt_web", + "claude_code", + "codex", + "gemini_cli", + "kiro", + "windsurf", + ], + ), + "status" => sanitize_background_task_metadata_enum( + value, + &[ + "cancelled", + "completed", + "failed", + "pending", + "processing", + "queued", + "running", + "skipped", + "succeeded", + ], + ), + "stage" => sanitize_background_task_metadata_enum( + value, + &[ + "completed", + "deleting_endpoints", + "deleting_keys", + "deleting_models", + "deleting_provider", + "disabling", + "failed", + "preparing", + "queued", + "skipped", + ], + ), + "object_cleanup_mode" => sanitize_background_task_metadata_enum( + value, + &["legacy_plaintext_deleted_after_verified_encryption"], + ), + "sha256" => sanitize_background_task_metadata_hex(value, 64, 64), + "export_version" => sanitize_background_task_metadata_version(value), + "exported_at" => sanitize_background_task_metadata_rfc3339(value), + "scheduled_slot" => sanitize_background_task_metadata_scheduled_slot(value), + "versioned_storage_cleanup_notice" => sanitize_background_task_metadata_enum( + value, + &["legacy_plaintext_versions_require_external_cleanup"], + ), + _ => None, + } +} + +fn sanitize_background_task_metadata_enum(value: Value, allowed: &[&str]) -> Option { + let value = value.as_str()?.trim().to_ascii_lowercase(); + allowed + .contains(&value.as_str()) + .then_some(Value::String(value)) +} + +fn sanitize_background_task_metadata_hex(value: Value, min: usize, max: usize) -> Option { + let value = value.as_str()?.trim().to_ascii_lowercase(); + ((min..=max).contains(&value.len()) && value.bytes().all(|byte| byte.is_ascii_hexdigit())) + .then_some(Value::String(value)) +} + +fn sanitize_background_task_metadata_version(value: Value) -> Option { + let value = value.as_str()?.trim(); + (!value.is_empty() + && value.len() <= 32 + && value + .bytes() + .all(|byte| byte.is_ascii_digit() || matches!(byte, b'.' | b'-' | b'_'))) + .then_some(Value::String(value.to_string())) +} + +fn sanitize_background_task_metadata_rfc3339(value: Value) -> Option { + let value = value.as_str()?.trim(); + (value.len() <= 64 && chrono::DateTime::parse_from_rfc3339(value).is_ok()) + .then_some(Value::String(value.to_string())) +} + +fn sanitize_background_task_metadata_scheduled_slot(value: Value) -> Option { + let value = value.as_str()?.trim(); + let (unit, timestamp) = value.split_once(':')?; + (matches!(unit, "hours" | "days" | "weeks" | "months") + && value.len() <= 80 + && chrono::DateTime::parse_from_rfc3339(timestamp).is_ok()) + .then_some(Value::String(value.to_string())) +} + +fn sanitize_background_task_metadata_error_code(value: &str) -> Option { + let value = value.trim().to_ascii_lowercase(); + Some(match value.as_str() { + "background_task_failed" + | "background_task_panicked" + | "provider_delete_failed" + | "provider_oauth_batch_import_failed" + | "s3_backup_failed" + | "s3_backup_slot_record_failed" => value, + _ if value.is_empty() => return None, + _ => BACKGROUND_TASK_DEFAULT_ERROR_CODE.to_string(), + }) +} + #[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)] pub struct BackgroundTaskListQuery { pub task_key_substring: Option, @@ -288,3 +558,152 @@ impl BackgroundTaskRepository for T where T: BackgroundTaskReadRepository + BackgroundTaskWriteRepository + Send + Sync { } + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + + fn background_task_run(status: BackgroundTaskStatus) -> UpsertBackgroundTaskRun { + UpsertBackgroundTaskRun { + id: "run-1".to_string(), + task_key: "security-review".to_string(), + kind: BackgroundTaskKind::OnDemand, + trigger: "manual".to_string(), + status, + attempt: 1, + max_attempts: 3, + owner_instance: None, + progress_percent: 50, + progress_message: None, + payload_json: None, + result_json: None, + error_message: None, + cancel_requested: false, + created_by: None, + created_at_unix_secs: 1, + started_at_unix_secs: Some(2), + finished_at_unix_secs: None, + updated_at_unix_secs: 3, + } + } + + #[test] + fn run_sanitization_removes_sensitive_and_nested_metadata() { + let mut run = background_task_run(BackgroundTaskStatus::Running); + run.owner_instance = Some("gateway-a".to_string()); + run.created_by = Some("admin@example.com".to_string()); + run.progress_message = Some("Bearer secret-access-token".to_string()); + run.payload_json = Some(json!({ + "provider_id": "provider-1", + "gateway_instance_id": "gateway-a", + "bucket": "safe-bucket", + "partition": 7, + "success": 1, + "access_token": "secret-access-token", + "authorization": "Bearer secret-access-token", + "password": "secret-password", + "error": "upstream detail containing secret", + "detail": "private diagnostic", + "provider_id ": "Bearer secret-access-token", + "gateway_instance_id ": "gateway-a; Authorization: secret", + "bucket ": "secret/bucket", + "nested": {"refresh_token": "secret-refresh-token"}, + "unknown": "must not be persisted" + })); + + let stored = run.into_stored(); + + assert_eq!(stored.progress_message, None); + assert_eq!(stored.owner_instance, None); + assert_eq!(stored.created_by, None); + assert_eq!( + stored.payload_json, + Some(json!({ + "partition": 7, + "success": 1 + })) + ); + } + + #[test] + fn run_sanitization_classifies_errors_and_clears_non_failure_errors() { + let mut failed = background_task_run(BackgroundTaskStatus::Failed); + failed.error_message = Some("upstream response included a credential".to_string()); + failed.result_json = Some(json!({ + "error_code": "raw-provider-error", + "failed": 2, + "token": "secret" + })); + + let failed = failed.into_stored(); + assert_eq!( + failed.error_message.as_deref(), + Some(BACKGROUND_TASK_DEFAULT_ERROR_CODE) + ); + assert_eq!( + failed.result_json, + Some(json!({ + "error_code": BACKGROUND_TASK_DEFAULT_ERROR_CODE, + "failed": 2 + })) + ); + + let mut succeeded = background_task_run(BackgroundTaskStatus::Succeeded); + succeeded.error_message = Some("provider_delete_failed".to_string()); + assert_eq!(succeeded.into_stored().error_message, None); + } + + #[test] + fn historical_run_sanitization_applies_the_same_read_boundary() { + let mut stored = background_task_run(BackgroundTaskStatus::Failed).into_stored(); + stored.progress_message = Some("legacy diagnostic with token".to_string()); + stored.payload_json = Some(json!({ + "scope": "data", + "refresh_token": "legacy-secret", + "nested": {"password": "legacy-password"} + })); + stored.error_message = Some("legacy upstream error: legacy-secret".to_string()); + + stored.sanitize_persisted_data(); + + assert_eq!(stored.progress_message, None); + assert_eq!(stored.payload_json, Some(json!({"scope": "data"}))); + assert_eq!( + stored.error_message.as_deref(), + Some(BACKGROUND_TASK_DEFAULT_ERROR_CODE) + ); + } + + #[test] + fn event_sanitization_canonicalizes_type_message_and_payload() { + let event = UpsertBackgroundTaskEvent { + id: "event-1".to_string(), + run_id: "run-1".to_string(), + event_type: "provider returned secret-token".to_string(), + message: "Authorization: Bearer secret-token".to_string(), + payload_json: Some(json!({ + "stage": "finalize", + "bytes": 42, + "error_code": "upstream said secret-token", + "error": "secret-token", + "detail": "credential detail", + "token": "secret-token", + "nested": {"authorization": "Bearer secret-token"} + })), + created_at_unix_secs: 4, + } + .into_stored(); + + assert_eq!(event.event_type, BACKGROUND_TASK_UNCLASSIFIED_EVENT); + assert_eq!(event.message, BACKGROUND_TASK_UNCLASSIFIED_EVENT); + assert_eq!( + event.payload_json, + Some(json!({ + "bytes": 42, + "error_code": BACKGROUND_TASK_DEFAULT_ERROR_CODE + })) + ); + } +} diff --git a/crates/aether-data/contracts/src/repository/billing/mod.rs b/crates/aether-data/contracts/src/repository/billing/mod.rs index 6753b7931..68fe35886 100644 --- a/crates/aether-data/contracts/src/repository/billing/mod.rs +++ b/crates/aether-data/contracts/src/repository/billing/mod.rs @@ -1,9 +1,27 @@ +mod replacement; mod types; +mod usage_policy; +pub use replacement::{ + entitlements_have_replacement_selector, entitlements_should_replace_existing, + validate_entitlement_replacement_groups, EntitlementReplacementGroupValidationError, + ENTITLEMENT_REPLACEMENT_GROUP_FIELD, MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH, +}; pub use types::{ + checked_plan_duration_days, checked_plan_duration_days_from_snapshot, AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, - BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, }; +pub use usage_policy::{ + nonnegative_usd_to_usage_policy_cost_units, parse_usage_policy_entitlements, + usd_to_usage_policy_cost_units, UsagePolicyEnforcement, UsagePolicyEntitlement, + UsagePolicyEntitlementType, UsagePolicyMetric, UsagePolicyParseError, UsagePolicyRule, + UsagePolicyValidationError, UsagePolicyWindow, MAX_USAGE_POLICY_ENTITLEMENTS, + MAX_USAGE_POLICY_EXACT_INTEGER, MAX_USAGE_POLICY_ROLLING_WINDOW_SECONDS, + MAX_USAGE_POLICY_RULES, MAX_USAGE_POLICY_TEXT_LENGTH, MAX_USAGE_POLICY_TOTAL_RULES, + USAGE_POLICY_COST_UNITS_PER_USD, USAGE_POLICY_ENTITLEMENT_TYPE, +}; diff --git a/crates/aether-data/contracts/src/repository/billing/replacement.rs b/crates/aether-data/contracts/src/repository/billing/replacement.rs new file mode 100644 index 000000000..a9d31602d --- /dev/null +++ b/crates/aether-data/contracts/src/repository/billing/replacement.rs @@ -0,0 +1,200 @@ +use std::collections::HashSet; + +use serde_json::Value; + +pub const ENTITLEMENT_REPLACEMENT_GROUP_FIELD: &str = "replacement_group"; +pub const MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH: usize = 128; + +const LEGACY_REPLACEMENT_ENTITLEMENT_TYPES: [&str; 2] = ["daily_quota", "membership_group"]; + +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum EntitlementReplacementGroupValidationError { + #[error("entitlements must be an array")] + EntitlementsMustBeArray, + #[error("entitlements[{index}].replacement_group must be a string")] + InvalidType { index: usize }, + #[error("entitlements[{index}].replacement_group must not be empty")] + Empty { index: usize }, + #[error("entitlements[{index}].replacement_group exceeds maximum length {max_len}")] + TooLong { index: usize, max_len: usize }, +} + +pub fn validate_entitlement_replacement_groups( + entitlements: &Value, +) -> Result<(), EntitlementReplacementGroupValidationError> { + let items = entitlements + .as_array() + .ok_or(EntitlementReplacementGroupValidationError::EntitlementsMustBeArray)?; + + for (index, item) in items.iter().enumerate() { + let Some(group) = item.get(ENTITLEMENT_REPLACEMENT_GROUP_FIELD) else { + continue; + }; + let group = group + .as_str() + .ok_or(EntitlementReplacementGroupValidationError::InvalidType { index })? + .trim(); + if group.is_empty() { + return Err(EntitlementReplacementGroupValidationError::Empty { index }); + } + if group.chars().count() > MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH { + return Err(EntitlementReplacementGroupValidationError::TooLong { + index, + max_len: MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH, + }); + } + } + + Ok(()) +} + +pub fn entitlements_have_replacement_selector(entitlements: &Value) -> bool { + let Some(items) = entitlements.as_array() else { + return false; + }; + + items.iter().any(|item| { + let entitlement_type = item.get("type").and_then(Value::as_str); + LEGACY_REPLACEMENT_ENTITLEMENT_TYPES.contains(&entitlement_type.unwrap_or_default()) + || replacement_group(item).is_some() + }) +} + +pub fn entitlements_should_replace_existing(incoming: &Value, existing: &Value) -> bool { + let (Some(incoming), Some(existing)) = (incoming.as_array(), existing.as_array()) else { + return false; + }; + + if LEGACY_REPLACEMENT_ENTITLEMENT_TYPES.iter().any(|kind| { + entitlement_items_have_type(incoming, kind) && entitlement_items_have_type(existing, kind) + }) { + return true; + } + + let incoming_groups = incoming + .iter() + .filter_map(replacement_group) + .collect::>(); + !incoming_groups.is_empty() + && existing + .iter() + .filter_map(replacement_group) + .any(|group| incoming_groups.contains(group)) +} + +fn entitlement_items_have_type(items: &[Value], entitlement_type: &str) -> bool { + items + .iter() + .any(|item| item.get("type").and_then(Value::as_str) == Some(entitlement_type)) +} + +fn replacement_group(item: &Value) -> Option<&str> { + item.get(ENTITLEMENT_REPLACEMENT_GROUP_FIELD) + .and_then(Value::as_str) + .map(str::trim) + .filter(|group| !group.is_empty()) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn legacy_daily_quota_and_membership_groups_remain_mutually_exclusive() { + assert!(entitlements_should_replace_existing( + &json!([{"type": "daily_quota", "daily_quota_usd": 20}]), + &json!([ + {"type": "daily_quota", "daily_quota_usd": 10}, + {"type": "usage_policy", "rules": []} + ]), + )); + assert!(entitlements_should_replace_existing( + &json!([{"type": "membership_group", "grant_user_groups": ["pro"]}]), + &json!([{"type": "membership_group", "grant_user_groups": ["basic"]}]), + )); + } + + #[test] + fn usage_policies_stack_by_default() { + let incoming = json!([{"type": "usage_policy", "policy_id": "weekly", "rules": []}]); + let existing = json!([{"type": "usage_policy", "policy_id": "five-hour", "rules": []}]); + + assert!(!entitlements_have_replacement_selector(&incoming)); + assert!(!entitlements_should_replace_existing(&incoming, &existing)); + } + + #[test] + fn matching_explicit_groups_replace_the_whole_package() { + let incoming = json!([{ + "type": "usage_policy", + "replacement_group": "pro-tier", + "rules": [] + }]); + let existing = json!([ + {"type": "wallet_credit", "amount_usd": 10}, + { + "type": "usage_policy", + "replacement_group": "pro-tier", + "rules": [] + } + ]); + + assert!(entitlements_have_replacement_selector(&incoming)); + assert!(entitlements_should_replace_existing(&incoming, &existing)); + assert!(!entitlements_should_replace_existing( + &incoming, + &json!([{ + "type": "usage_policy", + "replacement_group": "team-tier", + "rules": [] + }]), + )); + } + + #[test] + fn explicit_groups_can_span_entitlement_types_and_ignore_outer_whitespace() { + assert!(entitlements_should_replace_existing( + &json!([{ + "type": "usage_policy", + "replacement_group": " traffic-tier ", + "rules": [] + }]), + &json!([{ + "type": "wallet_credit", + "replacement_group": "traffic-tier", + "amount_usd": 10 + }]), + )); + } + + #[test] + fn validates_explicit_group_shape_and_bounds() { + assert!(validate_entitlement_replacement_groups(&json!([{ + "type": "usage_policy", + "replacement_group": "pro-tier" + }])) + .is_ok()); + assert_eq!( + validate_entitlement_replacement_groups(&json!([{ + "type": "usage_policy", + "replacement_group": " " + }])), + Err(EntitlementReplacementGroupValidationError::Empty { index: 0 }) + ); + assert_eq!( + validate_entitlement_replacement_groups(&json!([{ + "type": "usage_policy", + "replacement_group": 42 + }])), + Err(EntitlementReplacementGroupValidationError::InvalidType { index: 0 }) + ); + assert!(matches!( + validate_entitlement_replacement_groups(&json!([{ + "type": "usage_policy", + "replacement_group": "x".repeat(MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH + 1) + }])), + Err(EntitlementReplacementGroupValidationError::TooLong { index: 0, .. }) + )); + } +} diff --git a/crates/aether-data/contracts/src/repository/billing/types.rs b/crates/aether-data/contracts/src/repository/billing/types.rs index 2ec998899..82b891e1b 100644 --- a/crates/aether-data/contracts/src/repository/billing/types.rs +++ b/crates/aether-data/contracts/src/repository/billing/types.rs @@ -150,7 +150,7 @@ pub enum AdminBillingMutationOutcome { Unavailable, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct PaymentGatewayConfigRecord { pub provider: String, pub enabled: bool, @@ -166,7 +166,30 @@ pub struct PaymentGatewayConfigRecord { pub updated_at_unix_secs: u64, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for PaymentGatewayConfigRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PaymentGatewayConfigRecord") + .field("provider", &self.provider) + .field("enabled", &self.enabled) + .field("endpoint_url", &"[REDACTED]") + .field( + "callback_base_url", + &self.callback_base_url.as_ref().map(|_| "[REDACTED]"), + ) + .field("merchant_id", &self.merchant_id) + .field( + "merchant_key_encrypted", + &self.merchant_key_encrypted.as_ref().map(|_| "[REDACTED]"), + ) + .field("channels_json", &"[REDACTED]") + .field("created_at_unix_secs", &self.created_at_unix_secs) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq)] pub struct PaymentGatewayConfigWriteInput { pub provider: String, pub enabled: bool, @@ -181,6 +204,113 @@ pub struct PaymentGatewayConfigWriteInput { pub channels_json: Value, } +impl std::fmt::Debug for PaymentGatewayConfigWriteInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PaymentGatewayConfigWriteInput") + .field("provider", &self.provider) + .field("enabled", &self.enabled) + .field("endpoint_url", &"[REDACTED]") + .field( + "callback_base_url", + &self.callback_base_url.as_ref().map(|_| "[REDACTED]"), + ) + .field("merchant_id", &self.merchant_id) + .field( + "merchant_key_encrypted", + &self.merchant_key_encrypted.as_ref().map(|_| "[REDACTED]"), + ) + .field("preserve_existing_secret", &self.preserve_existing_secret) + .field("channels_json", &"[REDACTED]") + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq)] +pub struct PaymentGatewaySecretCasUpdate { + pub provider: String, + pub expected_merchant_key_encrypted: String, + pub merchant_key_encrypted: String, +} + +impl std::fmt::Debug for PaymentGatewaySecretCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PaymentGatewaySecretCasUpdate") + .field("provider", &self.provider) + .field("expected_merchant_key_encrypted", &"[REDACTED]") + .field("merchant_key_encrypted", &"[REDACTED]") + .finish() + } +} + +#[derive(Clone, PartialEq)] +pub struct PaymentGatewayConfigCasWriteInput { + pub input: PaymentGatewayConfigWriteInput, + pub expected_existing: bool, + pub expected_merchant_key_encrypted: Option, +} + +impl std::fmt::Debug for PaymentGatewayConfigCasWriteInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PaymentGatewayConfigCasWriteInput") + .field("input", &self.input) + .field("expected_existing", &self.expected_existing) + .field( + "expected_merchant_key_encrypted", + &self + .expected_merchant_key_encrypted + .as_ref() + .map(|_| "[REDACTED]"), + ) + .finish() + } +} + +#[cfg(test)] +mod payment_gateway_debug_tests { + use super::{PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigWriteInput}; + + #[test] + fn payment_gateway_config_debug_output_redacts_credential_material() { + let input = PaymentGatewayConfigCasWriteInput { + input: PaymentGatewayConfigWriteInput { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: "https://endpoint.example/?key=endpoint-canary".to_string(), + callback_base_url: Some( + "https://callback.example/?token=callback-canary".to_string(), + ), + merchant_id: "merchant".to_string(), + merchant_key_encrypted: Some("merchant-key-canary".to_string()), + preserve_existing_secret: false, + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: serde_json::json!({"secret": "channels-canary"}), + }, + expected_existing: true, + expected_merchant_key_encrypted: Some("expected-merchant-key-canary".to_string()), + }; + + let debug = format!("{input:?}"); + assert!(debug.contains("[REDACTED]")); + for secret in [ + "endpoint-canary", + "callback-canary", + "merchant-key-canary", + "channels-canary", + "expected-merchant-key-canary", + ] { + assert!( + !debug.contains(secret), + "debug output leaked {secret}: {debug}" + ); + } + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct BillingPlanRecord { pub id: String, @@ -214,6 +344,43 @@ pub struct BillingPlanWriteInput { pub entitlements_json: Value, } +/// Convert a plan duration into whole days without allowing integer or +/// `chrono::TimeDelta` overflow. Plan snapshots are persisted and may later be +/// fulfilled by any database adapter, so the accepted range must be portable +/// across all of them. +pub fn checked_plan_duration_days(duration_unit: &str, duration_value: i64) -> Result { + if duration_value <= 0 { + return Err("plan duration_value must be positive".to_string()); + } + let days = match duration_unit.trim() { + "day" | "custom" => Some(duration_value), + "month" => duration_value.checked_mul(30), + "year" => duration_value.checked_mul(365), + _ => return Err("plan duration_unit is invalid".to_string()), + } + .ok_or_else(|| "plan duration exceeds the supported range".to_string())?; + chrono::TimeDelta::try_days(days) + .ok_or_else(|| "plan duration exceeds the supported range".to_string())?; + Ok(days) +} + +/// Read a persisted plan snapshot using the historical month/one defaults, +/// while rejecting malformed or unrepresentable explicit values. +pub fn checked_plan_duration_days_from_snapshot(snapshot: &Value) -> Result { + let duration_unit = match snapshot.get("duration_unit") { + None => "month", + Some(Value::String(value)) => value.as_str(), + Some(_) => return Err("product_snapshot.duration_unit is invalid".to_string()), + }; + let duration_value = match snapshot.get("duration_value") { + None => 1, + Some(value) => value + .as_i64() + .ok_or_else(|| "product_snapshot.duration_value must be an integer".to_string())?, + }; + checked_plan_duration_days(duration_unit, duration_value) +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct UserPlanEntitlementRecord { pub id: String, @@ -369,6 +536,37 @@ pub trait BillingReadRepository: Send + Sync { Ok(None) } + /// Re-read a gateway configuration from the authoritative backing store. + /// Implementations without a read cache may delegate to the normal read. + async fn find_payment_gateway_config_strong( + &self, + provider: &str, + ) -> Result, crate::DataLayerError> { + self.find_payment_gateway_config(provider).await + } + + /// Replace only the encrypted merchant secret when the exact previously + /// observed ciphertext is still stored. Timestamps and all other fields + /// must remain unchanged. + async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + let _ = update; + Ok(false) + } + + /// Create a configuration only when absent, or update it only when the + /// exact nullable merchant-secret fence still matches. + async fn compare_and_swap_payment_gateway_config( + &self, + input: &PaymentGatewayConfigCasWriteInput, + ) -> Result, crate::DataLayerError> + { + let _ = input; + Ok(AdminBillingMutationOutcome::Unavailable) + } + async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -436,6 +634,15 @@ pub trait BillingReadRepository: Send + Sync { Ok(None) } + async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, entitlement_id); + Ok(AdminBillingMutationOutcome::Unavailable) + } + async fn find_user_daily_quota_availability( &self, user_id: &str, diff --git a/crates/aether-data/contracts/src/repository/billing/usage_policy.rs b/crates/aether-data/contracts/src/repository/billing/usage_policy.rs new file mode 100644 index 000000000..f93dac8da --- /dev/null +++ b/crates/aether-data/contracts/src/repository/billing/usage_policy.rs @@ -0,0 +1,1059 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +pub const USAGE_POLICY_ENTITLEMENT_TYPE: &str = "usage_policy"; +pub const MAX_USAGE_POLICY_RULES: usize = 32; +pub const MAX_USAGE_POLICY_ENTITLEMENTS: usize = 16; +pub const MAX_USAGE_POLICY_TOTAL_RULES: usize = 64; +pub const MAX_USAGE_POLICY_ROLLING_WINDOW_SECONDS: u64 = 30 * 24 * 60 * 60; +pub const MAX_USAGE_POLICY_TEXT_LENGTH: usize = 128; +pub const MAX_USAGE_POLICY_EXACT_INTEGER: u64 = (1_u64 << 53) - 1; +pub const USAGE_POLICY_COST_UNITS_PER_USD: u64 = 100_000_000; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum UsagePolicyEntitlementType { + UsagePolicy, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct UsagePolicyEntitlement { + #[serde(rename = "type")] + pub entitlement_type: UsagePolicyEntitlementType, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub policy_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub replacement_group: Option, + pub rules: Vec, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct UsagePolicyRule { + pub metric: UsagePolicyMetric, + pub window: UsagePolicyWindow, + pub limit: f64, + #[serde(default)] + pub enforcement: UsagePolicyEnforcement, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum UsagePolicyMetric { + RequestCount, + Concurrency, + ActualCostUsd, +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum UsagePolicyEnforcement { + #[default] + HardCap, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)] +pub enum UsagePolicyWindow { + Rolling { + seconds: u64, + }, + CalendarDay { + #[serde(default, skip_serializing_if = "Option::is_none")] + timezone: Option, + }, + CalendarWeek { + #[serde(default, skip_serializing_if = "Option::is_none")] + timezone: Option, + #[serde(default = "default_week_start")] + week_start: u8, + }, + CalendarMonth { + #[serde(default, skip_serializing_if = "Option::is_none")] + timezone: Option, + }, + SubscriptionPeriod, + Concurrent, +} + +const fn default_week_start() -> u8 { + 1 +} + +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum UsagePolicyValidationError { + #[error("{field} must not be empty")] + EmptyText { field: String }, + #[error("{field} exceeds maximum length {max_len}")] + TextTooLong { field: String, max_len: usize }, + #[error("usage_policy.rules must not be empty")] + EmptyRules, + #[error("usage_policy.rules must contain at most {max_rules} rules")] + TooManyRules { max_rules: usize }, + #[error("entitlements must contain at most {max_policies} usage_policy entries")] + TooManyPolicies { max_policies: usize }, + #[error("usage_policy entries must contain at most {max_rules} rules in total")] + TooManyTotalRules { max_rules: usize }, + #[error("usage_policy.rules[{index}].limit must be finite and positive")] + InvalidLimit { index: usize }, + #[error("usage_policy.rules[{index}].limit must be a positive exact integer")] + InvalidIntegerLimit { index: usize }, + #[error("usage_policy.rules[{index}].limit is outside the supported cost range")] + InvalidCostLimit { index: usize }, + #[error("usage_policy.rules[{index}].window.seconds must be positive")] + ZeroRollingWindow { index: usize }, + #[error("usage_policy.rules[{index}].window.seconds must not exceed {max_seconds} seconds")] + RollingWindowTooLong { index: usize, max_seconds: u64 }, + #[error("usage_policy.rules[{index}].window.timezone is not a valid IANA timezone")] + InvalidTimezone { index: usize }, + #[error("usage_policy.rules[{index}].window.week_start must be between 1 and 7")] + InvalidWeekStart { index: usize }, + #[error("usage_policy.rules[{index}] request_count requires a time-based window")] + RequestCountRequiresTimeWindow { index: usize }, + #[error("usage_policy.rules[{index}] concurrency requires a concurrent window")] + ConcurrencyRequiresConcurrentWindow { index: usize }, + #[error("usage_policy.rules[{index}] actual_cost_usd requires a time-based window")] + ActualCostRequiresTimeWindow { index: usize }, +} + +#[derive(Debug, thiserror::Error)] +pub enum UsagePolicyParseError { + #[error("entitlements must be an array")] + EntitlementsMustBeArray, + #[error("usage_policy entitlement at index {index} has invalid JSON shape: {source}")] + InvalidShape { + index: usize, + #[source] + source: serde_json::Error, + }, + #[error("usage_policy entitlement at index {index} is invalid: {source}")] + InvalidPolicy { + index: usize, + #[source] + source: UsagePolicyValidationError, + }, +} + +impl UsagePolicyEntitlement { + pub fn validate(&self) -> Result<(), UsagePolicyValidationError> { + validate_optional_text( + self.policy_id.as_deref(), + "usage_policy.policy_id", + MAX_USAGE_POLICY_TEXT_LENGTH, + )?; + validate_optional_text( + self.name.as_deref(), + "usage_policy.name", + MAX_USAGE_POLICY_TEXT_LENGTH, + )?; + validate_optional_text( + self.replacement_group.as_deref(), + "usage_policy.replacement_group", + MAX_USAGE_POLICY_TEXT_LENGTH, + )?; + + if self.rules.is_empty() { + return Err(UsagePolicyValidationError::EmptyRules); + } + if self.rules.len() > MAX_USAGE_POLICY_RULES { + return Err(UsagePolicyValidationError::TooManyRules { + max_rules: MAX_USAGE_POLICY_RULES, + }); + } + + for (index, rule) in self.rules.iter().enumerate() { + rule.validate(index)?; + } + Ok(()) + } +} + +impl UsagePolicyRule { + fn validate(&self, index: usize) -> Result<(), UsagePolicyValidationError> { + if !self.limit.is_finite() || self.limit <= 0.0 { + return Err(UsagePolicyValidationError::InvalidLimit { index }); + } + match self.metric { + UsagePolicyMetric::RequestCount | UsagePolicyMetric::Concurrency => { + if self.request_limit().is_none() { + return Err(UsagePolicyValidationError::InvalidIntegerLimit { index }); + } + } + UsagePolicyMetric::ActualCostUsd => { + if self.cost_limit_units().is_none() { + return Err(UsagePolicyValidationError::InvalidCostLimit { index }); + } + } + } + + match &self.window { + UsagePolicyWindow::Rolling { seconds: 0 } => { + return Err(UsagePolicyValidationError::ZeroRollingWindow { index }); + } + UsagePolicyWindow::Rolling { seconds } + if *seconds > MAX_USAGE_POLICY_ROLLING_WINDOW_SECONDS => + { + return Err(UsagePolicyValidationError::RollingWindowTooLong { + index, + max_seconds: MAX_USAGE_POLICY_ROLLING_WINDOW_SECONDS, + }); + } + UsagePolicyWindow::CalendarDay { timezone } + | UsagePolicyWindow::CalendarMonth { timezone } + | UsagePolicyWindow::CalendarWeek { timezone, .. } => { + validate_timezone(timezone.as_deref(), index)?; + } + UsagePolicyWindow::Rolling { .. } + | UsagePolicyWindow::SubscriptionPeriod + | UsagePolicyWindow::Concurrent => {} + } + + if let UsagePolicyWindow::CalendarWeek { week_start, .. } = &self.window { + if !(1..=7).contains(week_start) { + return Err(UsagePolicyValidationError::InvalidWeekStart { index }); + } + } + + match (&self.metric, &self.window) { + (UsagePolicyMetric::RequestCount, UsagePolicyWindow::Concurrent) => { + Err(UsagePolicyValidationError::RequestCountRequiresTimeWindow { index }) + } + (UsagePolicyMetric::Concurrency, UsagePolicyWindow::Concurrent) => Ok(()), + (UsagePolicyMetric::Concurrency, _) => { + Err(UsagePolicyValidationError::ConcurrencyRequiresConcurrentWindow { index }) + } + (UsagePolicyMetric::RequestCount, _) => Ok(()), + (UsagePolicyMetric::ActualCostUsd, UsagePolicyWindow::Concurrent) => { + Err(UsagePolicyValidationError::ActualCostRequiresTimeWindow { index }) + } + (UsagePolicyMetric::ActualCostUsd, _) => Ok(()), + } + } + + pub fn request_limit(&self) -> Option { + if !matches!( + self.metric, + UsagePolicyMetric::RequestCount | UsagePolicyMetric::Concurrency + ) || !self.limit.is_finite() + || self.limit <= 0.0 + || self.limit.fract() != 0.0 + || self.limit > MAX_USAGE_POLICY_EXACT_INTEGER as f64 + { + return None; + } + Some(self.limit as u64) + } + + pub fn cost_limit_units(&self) -> Option { + if self.metric != UsagePolicyMetric::ActualCostUsd { + return None; + } + usd_to_usage_policy_cost_units(self.limit) + } +} + +pub fn usd_to_usage_policy_cost_units(value: f64) -> Option { + if !value.is_finite() || value <= 0.0 { + return None; + } + let scaled = value * USAGE_POLICY_COST_UNITS_PER_USD as f64; + if !scaled.is_finite() || scaled < 0.5 || scaled > i64::MAX as f64 { + return None; + } + let rounded = scaled.round(); + let units = rounded as u64; + (units <= i64::MAX as u64).then_some(units) +} + +pub fn nonnegative_usd_to_usage_policy_cost_units(value: f64) -> Option { + if !value.is_finite() || value < 0.0 { + return None; + } + let scaled = value * USAGE_POLICY_COST_UNITS_PER_USD as f64; + if !scaled.is_finite() || scaled > i64::MAX as f64 { + return None; + } + let units = scaled.round() as u64; + (units <= i64::MAX as u64).then_some(units) +} + +pub fn parse_usage_policy_entitlements( + entitlements: &Value, +) -> Result, UsagePolicyParseError> { + let items = entitlements + .as_array() + .ok_or(UsagePolicyParseError::EntitlementsMustBeArray)?; + let mut policies = Vec::new(); + let mut total_rules = 0_usize; + + for (index, item) in items.iter().enumerate() { + if item.get("type").and_then(Value::as_str) != Some(USAGE_POLICY_ENTITLEMENT_TYPE) { + continue; + } + + let policy = serde_json::from_value::(item.clone()) + .map_err(|source| UsagePolicyParseError::InvalidShape { index, source })?; + policy + .validate() + .map_err(|source| UsagePolicyParseError::InvalidPolicy { index, source })?; + if policies.len() >= MAX_USAGE_POLICY_ENTITLEMENTS { + return Err(UsagePolicyParseError::InvalidPolicy { + index, + source: UsagePolicyValidationError::TooManyPolicies { + max_policies: MAX_USAGE_POLICY_ENTITLEMENTS, + }, + }); + } + total_rules = total_rules.saturating_add(policy.rules.len()); + if total_rules > MAX_USAGE_POLICY_TOTAL_RULES { + return Err(UsagePolicyParseError::InvalidPolicy { + index, + source: UsagePolicyValidationError::TooManyTotalRules { + max_rules: MAX_USAGE_POLICY_TOTAL_RULES, + }, + }); + } + policies.push(policy); + } + + Ok(policies) +} + +fn validate_optional_text( + value: Option<&str>, + field: &str, + max_len: usize, +) -> Result<(), UsagePolicyValidationError> { + let Some(value) = value else { + return Ok(()); + }; + let value = value.trim(); + if value.is_empty() { + return Err(UsagePolicyValidationError::EmptyText { + field: field.to_string(), + }); + } + if value.chars().count() > max_len { + return Err(UsagePolicyValidationError::TextTooLong { + field: field.to_string(), + max_len, + }); + } + Ok(()) +} + +fn validate_timezone( + timezone: Option<&str>, + index: usize, +) -> Result<(), UsagePolicyValidationError> { + let Some(timezone) = timezone else { + return Ok(()); + }; + let timezone = timezone.trim(); + if timezone.is_empty() || timezone.parse::().is_err() { + return Err(UsagePolicyValidationError::InvalidTimezone { index }); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn parses_flexible_combinations_and_ignores_legacy_entitlements() { + let entitlements = json!([ + { + "type": "daily_quota", + "daily_quota_usd": 25.0 + }, + { + "type": "usage_policy", + "policy_id": "standard-traffic", + "name": "Standard traffic limits", + "replacement_group": "pro-tier", + "rules": [ + { + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 18000 }, + "limit": 500 + }, + { + "metric": "request_count", + "window": { + "kind": "calendar_week", + "timezone": "Asia/Shanghai", + "week_start": 1 + }, + "limit": 10000, + "enforcement": "hard_cap" + }, + { + "metric": "concurrency", + "window": { "kind": "concurrent" }, + "limit": 4 + } + ] + } + ]); + + let policies = parse_usage_policy_entitlements(&entitlements).unwrap(); + + assert_eq!(policies.len(), 1); + assert_eq!(policies[0].policy_id.as_deref(), Some("standard-traffic")); + assert_eq!(policies[0].replacement_group.as_deref(), Some("pro-tier")); + assert_eq!(policies[0].rules.len(), 3); + assert_eq!( + policies[0].rules[0].enforcement, + UsagePolicyEnforcement::HardCap + ); + assert_eq!(policies[0].rules[2].window, UsagePolicyWindow::Concurrent); + } + + #[test] + fn supports_single_weekly_rule_and_subscription_period() { + let entitlements = json!([ + { + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_week" }, + "limit": 1000 + }] + }, + { + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "subscription_period" }, + "limit": 10000 + }] + } + ]); + + let policies = parse_usage_policy_entitlements(&entitlements).unwrap(); + + assert_eq!(policies.len(), 2); + assert_eq!( + policies[0].rules[0].window, + UsagePolicyWindow::CalendarWeek { + timezone: None, + week_start: 1, + } + ); + assert_eq!( + policies[1].rules[0].window, + UsagePolicyWindow::SubscriptionPeriod + ); + } + + #[test] + fn serializes_the_public_json_contract() { + let policy = UsagePolicyEntitlement { + entitlement_type: UsagePolicyEntitlementType::UsagePolicy, + policy_id: None, + name: None, + replacement_group: None, + rules: vec![UsagePolicyRule { + metric: UsagePolicyMetric::RequestCount, + window: UsagePolicyWindow::Rolling { seconds: 60 }, + limit: 120.0, + enforcement: UsagePolicyEnforcement::HardCap, + }], + }; + + assert_eq!( + serde_json::to_value(policy).unwrap(), + json!({ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 120.0, + "enforcement": "hard_cap" + }] + }) + ); + } + + #[test] + fn rejects_invalid_limits_and_zero_rolling_windows() { + let zero_limit = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 0 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&zero_limit), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidLimit { index: 0 }, + .. + }) + )); + + let zero_window = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 0 }, + "limit": 1 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&zero_window), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::ZeroRollingWindow { index: 0 }, + .. + }) + )); + + let excessive_window = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 2_592_001 }, + "limit": 1 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&excessive_window), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::RollingWindowTooLong { + index: 0, + max_seconds: 2_592_000 + }, + .. + }) + )); + } + + #[test] + fn validates_metric_specific_limits_and_cost_conversion() { + let cost_policy = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "actual_cost_usd", + "window": { "kind": "rolling", "seconds": 18000 }, + "limit": 12.34567891 + }] + }]); + let policies = parse_usage_policy_entitlements(&cost_policy).unwrap(); + assert_eq!(policies[0].rules[0].cost_limit_units(), Some(1_234_567_891)); + assert_eq!(policies[0].rules[0].request_limit(), None); + + let fractional_requests = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 1.5 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&fractional_requests), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidIntegerLimit { index: 0 }, + .. + }) + )); + + let sub_unit_cost = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "actual_cost_usd", + "window": { "kind": "calendar_day" }, + "limit": 0.000000001 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&sub_unit_cost), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidCostLimit { index: 0 }, + .. + }) + )); + } + + #[test] + fn rejects_metric_and_window_mismatches() { + let concurrency_with_rolling_window = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "concurrency", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 3 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&concurrency_with_rolling_window), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::ConcurrencyRequiresConcurrentWindow { + index: 0 + }, + .. + }) + )); + + let cost_with_concurrent_window = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "actual_cost_usd", + "window": { "kind": "concurrent" }, + "limit": 1 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&cost_with_concurrent_window), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::ActualCostRequiresTimeWindow { index: 0 }, + .. + }) + )); + + let request_count_with_concurrent_window = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "concurrent" }, + "limit": 3 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&request_count_with_concurrent_window), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::RequestCountRequiresTimeWindow { index: 0 }, + .. + }) + )); + } + + #[test] + fn rejects_invalid_calendar_configuration() { + let invalid_timezone = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_day", "timezone": "Mars/Olympus" }, + "limit": 100 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&invalid_timezone), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidTimezone { index: 0 }, + .. + }) + )); + + let invalid_week_start = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_week", "week_start": 0 }, + "limit": 100 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&invalid_week_start), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidWeekStart { index: 0 }, + .. + }) + )); + } + + #[test] + fn rejects_empty_rules_unknown_fields_and_non_array_roots() { + let empty_rules = json!([{"type": "usage_policy", "rules": []}]); + assert!(matches!( + parse_usage_policy_entitlements(&empty_rules), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::EmptyRules, + .. + }) + )); + + let typo = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 100, + "enforcment": "hard_cap" + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&typo), + Err(UsagePolicyParseError::InvalidShape { .. }) + )); + + assert!(matches!( + parse_usage_policy_entitlements(&json!({})), + Err(UsagePolicyParseError::EntitlementsMustBeArray) + )); + } + + #[test] + fn window_discriminator_requires_kind_instead_of_type() { + for window in [ + json!({"type": "rolling", "seconds": 60}), + json!({"seconds": 60}), + ] { + let entitlements = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": window, + "limit": 100 + }] + }]); + + assert!(matches!( + parse_usage_policy_entitlements(&entitlements), + Err(UsagePolicyParseError::InvalidShape { .. }) + )); + } + } + + #[test] + fn enforces_collection_and_metadata_bounds() { + let rule = json!({ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 1 }, + "limit": 1 + }); + let maximum_rules = json!([{ + "type": "usage_policy", + "rules": vec![rule.clone(); MAX_USAGE_POLICY_RULES] + }]); + assert_eq!( + parse_usage_policy_entitlements(&maximum_rules) + .unwrap() + .first() + .unwrap() + .rules + .len(), + MAX_USAGE_POLICY_RULES + ); + + let too_many_rules = json!([{ + "type": "usage_policy", + "rules": vec![rule; MAX_USAGE_POLICY_RULES + 1] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&too_many_rules), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::TooManyRules { + max_rules: MAX_USAGE_POLICY_RULES + }, + .. + }) + )); + + let empty_name = json!([{ + "type": "usage_policy", + "name": " ", + "rules": [{ + "metric": "concurrency", + "window": { "kind": "concurrent" }, + "limit": 1 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&empty_name), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::EmptyText { .. }, + .. + }) + )); + } + + #[test] + fn bounds_policy_count_and_total_rule_count() { + let one_rule_policy = json!({ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 1 + }] + }); + let too_many_policies = Value::Array( + (0..=MAX_USAGE_POLICY_ENTITLEMENTS) + .map(|_| one_rule_policy.clone()) + .collect(), + ); + assert!(matches!( + parse_usage_policy_entitlements(&too_many_policies), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::TooManyPolicies { .. }, + .. + }) + )); + + let max_rules_policy = json!({ + "type": "usage_policy", + "rules": vec![one_rule_policy["rules"][0].clone(); MAX_USAGE_POLICY_RULES] + }); + let too_many_total_rules = + json!([max_rules_policy.clone(), max_rules_policy, one_rule_policy]); + assert!(matches!( + parse_usage_policy_entitlements(&too_many_total_rules), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::TooManyTotalRules { .. }, + .. + }) + )); + } + + #[test] + fn supports_week_only_and_five_hour_plus_week_combinations() { + let entitlements = json!([ + { + "type": "usage_policy", + "policy_id": "weekly-only", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_week", "timezone": "Asia/Shanghai" }, + "limit": 10_000 + }] + }, + { + "type": "usage_policy", + "policy_id": "burst-and-weekly", + "rules": [ + { + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 18_000 }, + "limit": 500 + }, + { + "metric": "request_count", + "window": { "kind": "calendar_week", "timezone": "Asia/Shanghai" }, + "limit": 20_000 + } + ] + } + ]); + + let policies = parse_usage_policy_entitlements(&entitlements).expect("policy combinations"); + + assert_eq!(policies.len(), 2); + assert_eq!(policies[0].rules.len(), 1); + assert_eq!(policies[1].rules.len(), 2); + assert!(matches!( + policies[1].rules[0].window, + UsagePolicyWindow::Rolling { seconds: 18_000 } + )); + assert!(matches!( + policies[1].rules[1].window, + UsagePolicyWindow::CalendarWeek { .. } + )); + } + + #[test] + fn supports_qps_rpm_concurrency_and_cost_in_one_policy() { + let entitlements = json!([{ + "type": "usage_policy", + "policy_id": "full-traffic-policy", + "rules": [ + { + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 1 }, + "limit": 2 + }, + { + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 60 + }, + { + "metric": "concurrency", + "window": { "kind": "concurrent" }, + "limit": 8 + }, + { + "metric": "actual_cost_usd", + "window": { "kind": "calendar_month", "timezone": "UTC" }, + "limit": 125.50 + } + ] + }]); + + let policy = &parse_usage_policy_entitlements(&entitlements).expect("combined policy")[0]; + + assert_eq!(policy.rules.len(), 4); + assert_eq!(policy.rules[0].request_limit(), Some(2)); + assert_eq!(policy.rules[1].request_limit(), Some(60)); + assert_eq!(policy.rules[2].request_limit(), Some(8)); + assert_eq!(policy.rules[3].cost_limit_units(), Some(12_550_000_000)); + assert!(matches!( + policy.rules[2].window, + UsagePolicyWindow::Concurrent + )); + } + + #[test] + fn keeps_multiple_policy_entitlements_independent() { + let entitlements = json!([ + { + "type": "usage_policy", + "policy_id": "api-traffic", + "rules": [{ + "metric": "request_count", + "window": { "kind": "rolling", "seconds": 60 }, + "limit": 100 + }] + }, + { + "type": "usage_policy", + "policy_id": "model-cost", + "rules": [{ + "metric": "actual_cost_usd", + "window": { "kind": "rolling", "seconds": 18_000 }, + "limit": 5 + }] + }, + { + "type": "usage_policy", + "policy_id": "subscription-lifetime", + "rules": [{ + "metric": "request_count", + "window": { "kind": "subscription_period" }, + "limit": 10_000 + }] + } + ]); + + let policies = + parse_usage_policy_entitlements(&entitlements).expect("independent policies"); + + assert_eq!( + policies + .iter() + .map(|policy| policy.rules.len()) + .sum::(), + 3 + ); + assert_eq!( + policies + .iter() + .map(|policy| policy.policy_id.as_deref()) + .collect::>(), + vec![ + Some("api-traffic"), + Some("model-cost"), + Some("subscription-lifetime") + ] + ); + } + + #[test] + fn validates_dst_timezones_and_week_start_boundaries() { + let dst_policies = json!([ + { + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_day", "timezone": "America/New_York" }, + "limit": 100 + }] + }, + { + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { + "kind": "calendar_week", + "timezone": "Europe/Berlin", + "week_start": 7 + }, + "limit": 100 + }] + }, + { + "type": "usage_policy", + "rules": [{ + "metric": "actual_cost_usd", + "window": { "kind": "calendar_month", "timezone": "Pacific/Apia" }, + "limit": 1 + }] + } + ]); + + let policies = parse_usage_policy_entitlements(&dst_policies).expect("DST timezones"); + assert_eq!(policies.len(), 3); + assert!(matches!( + policies[1].rules[0].window, + UsagePolicyWindow::CalendarWeek { week_start: 7, .. } + )); + + for (timezone, expected) in [ + ("America/New_York", true), + ("Europe/Berlin", true), + ("Pacific/Apia", true), + ("UTC", true), + // Text fields are normalized by trimming surrounding whitespace before parsing. + ("America/New_York ", true), + // chrono-tz intentionally accepts the POSIX-compatible EST alias. + ("EST", true), + ] { + let value = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_day", "timezone": timezone }, + "limit": 1 + }] + }]); + assert_eq!( + parse_usage_policy_entitlements(&value).is_ok(), + expected, + "{timezone}" + ); + } + } + + #[test] + fn accepts_week_start_one_and_rejects_values_outside_one_through_seven() { + for week_start in 1..=7 { + let value = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_week", "week_start": week_start }, + "limit": 1 + }] + }]); + assert!( + parse_usage_policy_entitlements(&value).is_ok(), + "week_start={week_start}" + ); + } + + for week_start in [0, 8, u8::MAX] { + let value = json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": { "kind": "calendar_week", "week_start": week_start }, + "limit": 1 + }] + }]); + assert!(matches!( + parse_usage_policy_entitlements(&value), + Err(UsagePolicyParseError::InvalidPolicy { + source: UsagePolicyValidationError::InvalidWeekStart { index: 0 }, + .. + }) + )); + } + } +} diff --git a/crates/aether-data/contracts/src/repository/candidates/mod.rs b/crates/aether-data/contracts/src/repository/candidates/mod.rs index 634035623..a1aae8964 100644 --- a/crates/aether-data/contracts/src/repository/candidates/mod.rs +++ b/crates/aether-data/contracts/src/repository/candidates/mod.rs @@ -2,9 +2,12 @@ mod types; pub use types::{ build_decision_trace, derive_request_candidate_final_status, - request_candidate_lifecycle_would_regress, DecisionTrace, DecisionTraceCandidate, - PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateFinalStatus, - RequestCandidateReadRepository, RequestCandidateRepository, RequestCandidateStatus, - RequestCandidateTrace, RequestCandidateWriteRepository, StoredRequestCandidate, - UpsertRequestCandidateRecord, + request_candidate_lifecycle_would_regress, sanitize_request_candidate_api_formats, + sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data, + sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason, + DecisionTrace, DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket, + RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository, + RequestCandidateStatus, RequestCandidateTrace, RequestCandidateWriteRepository, + StoredRequestCandidate, UpsertRequestCandidateRecord, REQUEST_CANDIDATE_ERROR_TYPES, + REQUEST_CANDIDATE_ERROR_TYPE_ALIASES, REQUEST_CANDIDATE_SKIP_REASONS, }; diff --git a/crates/aether-data/contracts/src/repository/candidates/types.rs b/crates/aether-data/contracts/src/repository/candidates/types.rs index 461e521ac..58bf33b6f 100644 --- a/crates/aether-data/contracts/src/repository/candidates/types.rs +++ b/crates/aether-data/contracts/src/repository/candidates/types.rs @@ -6,6 +6,156 @@ use crate::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; +const UNCLASSIFIED_CANDIDATE_ERROR_TYPE: &str = "unclassified_error"; +const UNCLASSIFIED_CANDIDATE_SKIP_REASON: &str = "unclassified_skip"; + +macro_rules! define_candidate_diagnostic_categories { + ($constant:ident, $predicate:ident, [$($value:literal),+ $(,)?]) => { + pub const $constant: &[&str] = &[$($value),+]; + + fn $predicate(value: &str) -> bool { + matches!(value, $($value)|+) + } + }; +} + +define_candidate_diagnostic_categories!( + REQUEST_CANDIDATE_SKIP_REASONS, + is_known_request_candidate_skip_reason, + [ + "account_quota_exhausted", + "api_key_concurrency_limit_reached", + "auth_api_key_concurrency_limit_reached", + "auth_channel_mismatch", + "auth_snapshot_missing", + "endpoint_api_format_changed", + "endpoint_inactive", + "format_conversion_disabled", + "gemini_file_mapping_mismatch", + "key_api_format_disabled", + "key_circuit_open", + "key_health_score_zero", + "key_inactive", + "key_model_disabled", + "key_model_not_allowed", + "key_rpm_exhausted", + "mapped_model_missing", + "oauth_invalid", + "pool_account_blocked", + "pool_account_exhausted", + "pool_active_probe_sealed", + "pool_cooldown", + "pool_cost_limit_reached", + "pool_group_exhausted", + "pool_key_lease_busy", + "pool_score_member_missing", + "provider_concurrency_limit_reached", + "provider_inactive", + "provider_key_concurrency_limit_reached", + "provider_quota_blocked", + "provider_request_body_build_failed", + "provider_request_body_missing", + "routing_profile_disallowed_key", + "routing_profile_disallowed_provider", + "transport_api_format_mismatch", + "transport_api_format_unsupported", + "transport_auth_unavailable", + "transport_body_rules_apply_failed", + "transport_body_rules_unsupported", + "transport_body_rules_unsupported_for_binary_upload", + "transport_custom_path_unsupported", + "transport_endpoint_kind_unsupported", + "transport_header_rules_apply_failed", + "transport_header_rules_unsupported", + "transport_oauth_resolution_unsupported", + "transport_operation_unsupported", + "transport_profile_unsupported", + "transport_provider_type_unsupported", + "transport_proxy_or_profile_unsupported", + "transport_proxy_unsupported", + "transport_snapshot_missing", + "transport_unsupported", + "upstream_url_missing", + ] +); + +pub const REQUEST_CANDIDATE_ERROR_TYPE_ALIASES: &[(&str, &str)] = &[ + ("connecttimeout", "connect_timeout"), + ("firstbytetimeout", "first_byte_timeout"), + ("protocolerror", "protocol_error"), + ("proxyerror", "proxy_error"), + ("readtimeout", "read_timeout"), + ("tlserror", "tls_error"), +]; + +define_candidate_diagnostic_categories!( + REQUEST_CANDIDATE_ERROR_TYPES, + is_known_request_candidate_error_type, + [ + "api_error", + "authentication_error", + "all_candidates_skipped", + "cancelled", + "candidate_list_empty", + "chatgpt_web_image_execution_unavailable", + "client_delivery_failed", + "connect_timeout", + "control_fallback", + "downstream_disconnect", + "execution_runtime_http_error", + "execution_runtime_stream_chunk_decode_error", + "execution_runtime_stream_frame_decode_error", + "execution_runtime_stream_non_success_status", + "execution_runtime_stream_read_error", + "execution_runtime_stream_rewrite_error", + "execution_runtime_stream_rewrite_flush_error", + "execution_runtime_sync_json_stream_bridge_error", + "execution_runtime_unavailable", + "first_byte_timeout", + "gateway_admission_failed", + "gateway_admission_timeout", + "grok_execution_unavailable", + "grok_upstream_error", + "image_sync_total_timeout", + "internal", + "invalid_provider_success_response", + "invalid_request_error", + "kiro_web_search_mcp_unavailable", + "local_stream_candidate_watchdog_timeout", + "local_stream_attempt_cancelled", + "local_sync_attempt_aborted", + "local_sync_attempt_cancelled", + "not_found_error", + "no_local_stream_plans", + "no_local_sync_plans", + "overloaded_error", + "permission_error", + "plan_usage_limit_exceeded", + "provider_request_body_build_failed", + "provider_request_body_missing", + "protocol_error", + "proxy_error", + "rate_limit_error", + "read_timeout", + "resource_exhausted", + "retryable_upstream_status", + "server_error", + "stream_http_error", + "stream_missing_terminal_event", + "stream_terminal_error", + "success_failover_pattern", + "tls_error", + "upstream4xx", + "upstream5xx", + "upstream_error", + "upstream_response_decode_failed", + "upstream_response_too_large", + "upstream_url_missing", + "websocket_cancelled", + "windsurf_native_execution_unavailable", + ] +); + #[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] #[serde(rename_all = "snake_case")] pub enum RequestCandidateStatus { @@ -74,6 +224,17 @@ pub struct StoredRequestCandidate { } impl StoredRequestCandidate { + pub fn sanitize_sensitive_diagnostics(&mut self) { + self.username = None; + self.api_key_name = None; + self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take()); + self.error_type = sanitize_request_candidate_error_type(self.error_type.take()); + self.error_message = None; + self.extra_data = sanitize_request_candidate_extra_data(self.extra_data.take()); + self.required_capabilities = + sanitize_request_candidate_required_capabilities(self.required_capabilities.take()); + } + #[allow(clippy::too_many_arguments)] pub fn new( id: String, @@ -92,7 +253,7 @@ impl StoredRequestCandidate { is_cached: bool, status_code: Option, error_type: Option, - error_message: Option, + _error_message: Option, latency_ms: Option, concurrent_requests: Option, extra_data: Option, @@ -162,7 +323,7 @@ impl StoredRequestCandidate { }) .transpose()?; - Ok(Self { + let mut candidate = Self { id, request_id, user_id, @@ -179,7 +340,7 @@ impl StoredRequestCandidate { is_cached, status_code, error_type, - error_message, + error_message: _error_message, latency_ms, concurrent_requests, extra_data, @@ -187,7 +348,9 @@ impl StoredRequestCandidate { created_at_unix_ms, started_at_unix_ms, finished_at_unix_ms, - }) + }; + candidate.sanitize_sensitive_diagnostics(); + Ok(candidate) } } @@ -211,11 +374,20 @@ pub struct RequestCandidateTrace { } impl RequestCandidateTrace { + pub fn sanitize_sensitive_diagnostics(&mut self) { + for candidate in &mut self.candidates { + candidate.sanitize_sensitive_diagnostics(); + } + } + pub fn from_candidates( request_id: impl Into, - all_candidates: Vec, + mut all_candidates: Vec, attempted_only: bool, ) -> Option { + for candidate in &mut all_candidates { + candidate.sanitize_sensitive_diagnostics(); + } if all_candidates.is_empty() { return None; } @@ -337,6 +509,22 @@ pub struct DecisionTraceCandidate { pub provider_key_is_active: Option, } +impl DecisionTraceCandidate { + pub fn sanitize_sensitive_diagnostics(&mut self) { + self.candidate.sanitize_sensitive_diagnostics(); + self.provider_website = self + .provider_website + .take() + .and_then(|value| sanitize_candidate_url(&value)); + self.endpoint_format_acceptance_config = None; + self.provider_key_api_formats = + sanitize_request_candidate_api_formats(self.provider_key_api_formats.take()); + self.provider_key_global_priority_by_format = None; + self.provider_key_capabilities = + sanitize_request_candidate_required_capabilities(self.provider_key_capabilities.take()); + } +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct DecisionTrace { pub request_id: String, @@ -346,6 +534,14 @@ pub struct DecisionTrace { pub candidates: Vec, } +impl DecisionTrace { + pub fn sanitize_sensitive_diagnostics(&mut self) { + for item in &mut self.candidates { + item.sanitize_sensitive_diagnostics(); + } + } +} + pub fn build_decision_trace( trace: RequestCandidateTrace, providers: Vec, @@ -365,7 +561,7 @@ pub fn build_decision_trace( .map(|item| (item.id.clone(), item)) .collect::>(); - DecisionTrace { + let mut trace = DecisionTrace { request_id: trace.request_id, total_candidates: trace.total_candidates, final_status: trace.final_status, @@ -377,7 +573,9 @@ pub fn build_decision_trace( enrich_decision_trace_candidate(candidate, &provider_map, &endpoint_map, &key_map) }) .collect(), - } + }; + trace.sanitize_sensitive_diagnostics(); + trace } fn enrich_decision_trace_candidate( @@ -524,6 +722,17 @@ pub struct UpsertRequestCandidateRecord { } impl UpsertRequestCandidateRecord { + pub fn sanitize_for_persistence(&mut self) { + self.username = None; + self.api_key_name = None; + self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take()); + self.error_type = sanitize_request_candidate_error_type(self.error_type.take()); + self.error_message = None; + self.extra_data = sanitize_request_candidate_extra_data(self.extra_data.take()); + self.required_capabilities = + sanitize_request_candidate_required_capabilities(self.required_capabilities.take()); + } + pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.id.trim().is_empty() { return Err(crate::DataLayerError::InvalidInput( @@ -554,6 +763,955 @@ impl UpsertRequestCandidateRecord { } } +pub fn sanitize_request_candidate_skip_reason(value: Option) -> Option { + let value = value?; + let normalized = value.trim().to_ascii_lowercase(); + let safe = if is_known_request_candidate_skip_reason(normalized.as_str()) { + normalized + } else { + UNCLASSIFIED_CANDIDATE_SKIP_REASON.to_string() + }; + Some(safe) +} + +pub fn sanitize_request_candidate_error_type(value: Option) -> Option { + let value = value?; + let normalized = value.trim().to_ascii_lowercase(); + let safe = REQUEST_CANDIDATE_ERROR_TYPE_ALIASES + .iter() + .find_map(|(alias, canonical)| (normalized == *alias).then_some(*canonical)) + .map(str::to_string) + .or_else(|| { + is_known_request_candidate_error_type(normalized.as_str()).then_some(normalized) + }) + .unwrap_or_else(|| UNCLASSIFIED_CANDIDATE_ERROR_TYPE.to_string()); + Some(safe) +} + +pub fn sanitize_request_candidate_extra_data( + extra_data: Option, +) -> Option { + let serde_json::Value::Object(object) = extra_data? else { + return None; + }; + let mut sanitized = serde_json::Map::new(); + + for field in ["gateway_execution_runtime", "stream_completed", "cache_1h"] { + insert_candidate_bool(&object, &mut sanitized, field); + } + for field in ["first_byte_time_ms", "pool_key_index"] { + insert_candidate_u64(&object, &mut sanitized, field); + } + insert_candidate_i64(&object, &mut sanitized, "priority_slot"); + insert_candidate_u64(&object, &mut sanitized, "ranking_index"); + + insert_candidate_known_string(&object, &mut sanitized, "phase", sanitize_candidate_phase); + for field in [ + "client_api_format", + "provider_api_format", + "client_contract", + "provider_contract", + ] { + insert_candidate_known_string( + &object, + &mut sanitized, + field, + sanitize_candidate_api_format, + ); + } + insert_candidate_known_string( + &object, + &mut sanitized, + "execution_strategy", + sanitize_candidate_execution_strategy, + ); + insert_candidate_known_string( + &object, + &mut sanitized, + "conversion_mode", + sanitize_candidate_conversion_mode, + ); + insert_candidate_known_string( + &object, + &mut sanitized, + "ranking_mode", + sanitize_candidate_ranking_mode, + ); + insert_candidate_known_string( + &object, + &mut sanitized, + "priority_mode", + sanitize_candidate_priority_mode, + ); + insert_candidate_known_string( + &object, + &mut sanitized, + "promoted_by", + sanitize_candidate_promotion_reason, + ); + insert_candidate_known_string( + &object, + &mut sanitized, + "demoted_by", + sanitize_candidate_demotion_reason, + ); + insert_candidate_known_string(&object, &mut sanitized, "source", sanitize_candidate_source); + insert_candidate_known_string( + &object, + &mut sanitized, + "execution_path", + sanitize_candidate_execution_path, + ); + + if let Some(url) = object + .get("upstream_url") + .and_then(serde_json::Value::as_str) + .and_then(sanitize_candidate_url) + { + sanitized.insert("upstream_url".to_string(), serde_json::Value::String(url)); + } + if let Some(summary) = object + .get("header_rules") + .and_then(sanitize_candidate_rules_summary) + { + sanitized.insert("header_rules".to_string(), summary); + } + if let Some(summary) = object + .get("body_rules") + .and_then(sanitize_candidate_rules_summary) + { + sanitized.insert("body_rules".to_string(), summary); + } + if let Some(summary) = object.get("proxy").and_then(sanitize_candidate_proxy) { + sanitized.insert("proxy".to_string(), summary); + } + if let Some(summary) = object + .get("error_flow") + .and_then(sanitize_candidate_error_flow) + { + sanitized.insert("error_flow".to_string(), summary); + } + if let Some(summary) = object + .get("routing_trace") + .and_then(sanitize_candidate_routing_trace) + { + sanitized.insert("routing_trace".to_string(), summary); + } + + if let Some(summary) = object + .get("upstream_response") + .and_then(sanitize_candidate_upstream_response) + { + sanitized.insert("upstream_response".to_string(), summary); + } + if let Some(progress) = object + .get("image_progress") + .and_then(sanitize_candidate_image_progress) + { + sanitized.insert("image_progress".to_string(), progress); + } + if let Some(exhaustion) = object + .get("pool_group_exhaustion") + .and_then(sanitize_candidate_pool_group_exhaustion) + { + sanitized.insert("pool_group_exhaustion".to_string(), exhaustion); + } + + (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)) +} + +pub fn sanitize_request_candidate_required_capabilities( + required_capabilities: Option, +) -> Option { + let serde_json::Value::Object(object) = required_capabilities? else { + return None; + }; + let mut sanitized = serde_json::Map::new(); + for capability in [ + "cache_1h", + "context_1m", + "gemini_files", + "streaming", + "vision", + ] { + let Some(enabled) = object + .get(capability) + .or_else(|| { + object + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case(capability)) + .map(|(_, value)| value) + }) + .and_then(sanitize_candidate_capability_value) + else { + continue; + }; + sanitized.insert(capability.to_string(), serde_json::Value::Bool(enabled)); + } + (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)) +} + +pub fn sanitize_request_candidate_api_formats( + api_formats: Option, +) -> Option { + let serde_json::Value::Array(values) = api_formats? else { + return None; + }; + let mut sanitized = Vec::new(); + for format in values + .iter() + .filter_map(serde_json::Value::as_str) + .filter_map(sanitize_candidate_api_format) + { + if sanitized + .iter() + .any(|existing: &serde_json::Value| existing.as_str() == Some(format)) + { + continue; + } + sanitized.push(serde_json::Value::String(format.to_string())); + } + (!sanitized.is_empty()).then_some(serde_json::Value::Array(sanitized)) +} + +fn sanitize_candidate_capability_value(value: &serde_json::Value) -> Option { + match value { + serde_json::Value::Bool(value) => Some(*value), + serde_json::Value::String(value) => match value.trim().to_ascii_lowercase().as_str() { + "true" => Some(true), + "false" => Some(false), + _ => None, + }, + serde_json::Value::Number(value) => value + .as_i64() + .map(|value| value > 0) + .or_else(|| value.as_u64().map(|value| value > 0)) + .or_else(|| value.as_f64().map(|value| value > 0.0)), + _ => None, + } +} + +fn sanitize_candidate_url(value: &str) -> Option { + let mut url = url::Url::parse(value.trim()).ok()?; + if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() { + return None; + } + + url.set_username("").ok()?; + url.set_password(None).ok()?; + url.set_path("/"); + url.set_query(None); + url.set_fragment(None); + Some(url.into()) +} + +fn sanitize_candidate_rules_summary(value: &serde_json::Value) -> Option { + if let Some(summary) = value.as_object() { + let mut sanitized = serde_json::Map::new(); + for field in ["count", "enabled_count", "conditional_count"] { + insert_candidate_u64(summary, &mut sanitized, field); + } + if let Some(action_counts) = summary + .get("action_counts") + .and_then(sanitize_candidate_rule_action_counts) + { + sanitized.insert("action_counts".to_string(), action_counts); + } + return (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)); + } + + let rules = value.as_array()?; + let mut summary = serde_json::Map::new(); + summary.insert( + "count".to_string(), + serde_json::Value::Number((rules.len() as u64).into()), + ); + + let mut enabled_count = 0_u64; + let mut conditional_count = 0_u64; + let mut action_counts = serde_json::Map::new(); + for rule in rules.iter().filter_map(serde_json::Value::as_object) { + if rule.get("enabled").and_then(serde_json::Value::as_bool) == Some(false) { + continue; + } + enabled_count = enabled_count.saturating_add(1); + if rule.get("condition").is_some_and(|value| !value.is_null()) { + conditional_count = conditional_count.saturating_add(1); + } + let Some(action) = rule + .get("action") + .or_else(|| rule.get("op")) + .and_then(serde_json::Value::as_str) + .and_then(sanitize_candidate_rule_action) + else { + continue; + }; + let count = action_counts + .get(action) + .and_then(serde_json::Value::as_u64) + .unwrap_or_default() + .saturating_add(1); + action_counts.insert(action.to_string(), serde_json::Value::Number(count.into())); + } + summary.insert( + "enabled_count".to_string(), + serde_json::Value::Number(enabled_count.into()), + ); + summary.insert( + "conditional_count".to_string(), + serde_json::Value::Number(conditional_count.into()), + ); + if !action_counts.is_empty() { + summary.insert( + "action_counts".to_string(), + serde_json::Value::Object(action_counts), + ); + } + Some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_rule_action_counts(value: &serde_json::Value) -> Option { + let counts = value.as_object()?; + let mut sanitized = serde_json::Map::new(); + for action in [ + "add", + "append", + "drop", + "insert", + "regex_replace", + "remove", + "rename", + "replace", + "set", + ] { + insert_candidate_u64(counts, &mut sanitized, action); + } + (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)) +} + +fn sanitize_candidate_rule_action(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "add" => Some("add"), + "append" => Some("append"), + "drop" => Some("drop"), + "insert" => Some("insert"), + "regex_replace" => Some("regex_replace"), + "remove" => Some("remove"), + "rename" => Some("rename"), + "replace" => Some("replace"), + "set" => Some("set"), + _ => None, + } +} + +fn sanitize_candidate_proxy(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + insert_candidate_known_string(object, &mut summary, "mode", sanitize_candidate_proxy_mode); + insert_candidate_known_string( + object, + &mut summary, + "source", + sanitize_candidate_proxy_source, + ); + if let Some(url) = object + .get("url") + .and_then(serde_json::Value::as_str) + .and_then(sanitize_candidate_url) + { + summary.insert("url".to_string(), serde_json::Value::String(url)); + } + for field in ["ttfb_ms", "connection_acquire_ms", "response_wait_ms"] { + insert_candidate_u64(object, &mut summary, field); + } + if let Some(timing) = object + .get("timing") + .and_then(sanitize_candidate_proxy_timing) + { + summary.insert("timing".to_string(), timing); + } + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_proxy_timing(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + for field in [ + "connection_acquire_ms", + "connection_ms", + "response_wait_ms", + "ttfb_ms", + "upstream_processing_ms", + ] { + insert_candidate_u64(object, &mut summary, field); + } + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_proxy_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "direct" => Some("direct"), + "manual" => Some("manual"), + "node" => Some("node"), + "system" => Some("system"), + "tunnel" => Some("tunnel"), + _ => None, + } +} + +fn sanitize_candidate_proxy_source(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "endpoint" => Some("endpoint"), + "key" => Some("key"), + "provider" => Some("provider"), + "system" => Some("system"), + "tunnel_affinity" => Some("tunnel_affinity"), + _ => None, + } +} + +fn sanitize_candidate_error_flow(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + insert_candidate_known_string( + object, + &mut summary, + "stage", + sanitize_candidate_error_stage, + ); + insert_candidate_known_string( + object, + &mut summary, + "source", + sanitize_candidate_error_source, + ); + insert_candidate_known_string( + object, + &mut summary, + "classification", + sanitize_candidate_error_classification, + ); + insert_candidate_known_string( + object, + &mut summary, + "decision", + sanitize_candidate_error_decision, + ); + insert_candidate_known_string( + object, + &mut summary, + "propagation", + sanitize_candidate_error_propagation, + ); + for field in ["retryable", "safe_to_expose", "safe_to_expose_upstream"] { + insert_candidate_bool(object, &mut summary, field); + } + if let Some(status_code) = object + .get("status_code") + .and_then(serde_json::Value::as_u64) + .filter(|status_code| *status_code <= u64::from(u16::MAX)) + { + summary.insert( + "status_code".to_string(), + serde_json::Value::Number(status_code.into()), + ); + } + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_error_stage(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "candidate" => Some("candidate"), + "client" => Some("client"), + "gateway" => Some("gateway"), + "request" => Some("request"), + "upstream" => Some("upstream"), + _ => None, + } +} + +fn sanitize_candidate_error_source(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "client" => Some("client"), + "client_response" => Some("client_response"), + "gateway" => Some("gateway"), + "request" => Some("request"), + "summary" => Some("summary"), + "upstream" => Some("upstream"), + "upstream_response" => Some("upstream_response"), + _ => None, + } +} + +fn sanitize_candidate_error_classification(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "retry_status_code" => Some("retry_status_code"), + "retry_success_pattern" => Some("retry_success_pattern"), + "retry_transport_error" => Some("retry_transport_error"), + "retry_upstream_failure" => Some("retry_upstream_failure"), + "stop_cyber_policy" => Some("stop_cyber_policy"), + "stop_error_pattern" => Some("stop_error_pattern"), + "stop_execution_error" => Some("stop_execution_error"), + "stop_status_code" => Some("stop_status_code"), + "stop_transport_error" => Some("stop_transport_error"), + "use_default" => Some("use_default"), + _ => None, + } +} + +fn sanitize_candidate_error_decision(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "retry_next_candidate" => Some("retry_next_candidate"), + "stop_local_failover" => Some("stop_local_failover"), + "use_default" => Some("use_default"), + _ => None, + } +} + +fn sanitize_candidate_error_propagation(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "captured" => Some("captured"), + "converted" => Some("converted"), + "local" => Some("local"), + "none" => Some("none"), + "passthrough" => Some("passthrough"), + "suppressed" => Some("suppressed"), + _ => None, + } +} + +fn sanitize_candidate_routing_trace(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + insert_candidate_i64(object, &mut summary, "group_version"); + insert_candidate_known_string( + object, + &mut summary, + "selection_source", + sanitize_candidate_routing_selection_source, + ); + insert_candidate_known_string( + object, + &mut summary, + "client_api_format", + sanitize_candidate_api_format, + ); + for (field, output) in [ + ("selected_rules", "selected_rule_count"), + ("global_candidates", "global_candidate_count"), + ("pool_expansion", "pool_expansion_count"), + ] { + if let Some(count) = object + .get(field) + .and_then(serde_json::Value::as_array) + .map(Vec::len) + .map(|value| value as u64) + .or_else(|| object.get(output).and_then(serde_json::Value::as_u64)) + { + summary.insert(output.to_string(), serde_json::Value::Number(count.into())); + } + } + for field in [ + "client_request_patch_summary", + "provider_request_patch_summary", + ] { + if let Some(patch) = object + .get(field) + .and_then(sanitize_candidate_routing_patch_summary) + { + summary.insert(field.to_string(), patch); + } + } + if let Some(facts) = object + .get("runtime_facts") + .and_then(sanitize_candidate_routing_runtime_facts) + { + summary.insert("runtime_facts".to_string(), facts); + } + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_routing_patch_summary( + value: &serde_json::Value, +) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + for (field, output) in [ + ("body_paths", "body_patch_count"), + ("header_names", "header_patch_count"), + ] { + if let Some(count) = object + .get(field) + .and_then(serde_json::Value::as_array) + .map(Vec::len) + .map(|value| value as u64) + .or_else(|| object.get(output).and_then(serde_json::Value::as_u64)) + { + summary.insert(output.to_string(), serde_json::Value::Number(count.into())); + } + } + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_routing_runtime_facts( + value: &serde_json::Value, +) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + insert_candidate_bool(object, &mut summary, "cache_affinity_hit"); + insert_candidate_known_string( + object, + &mut summary, + "scheduler_mode", + sanitize_candidate_routing_scheduler_mode, + ); + insert_candidate_known_string( + object, + &mut summary, + "priority_mode", + sanitize_candidate_routing_priority_mode, + ); + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_routing_selection_source(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "admin_dry_run" => Some("admin_dry_run"), + "api_key_default" => Some("api_key_default"), + "explicit" => Some("explicit"), + "explicit_header" => Some("explicit_header"), + "system_default" => Some("system_default"), + "user_default" => Some("user_default"), + "user_group_default" => Some("user_group_default"), + _ => None, + } +} + +fn sanitize_candidate_routing_scheduler_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "cache_affinity" | "cacheaffinity" => Some("cache_affinity"), + "fixed_order" | "fixedorder" => Some("fixed_order"), + "load_balance" | "loadbalance" => Some("load_balance"), + _ => None, + } +} + +fn sanitize_candidate_routing_priority_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "global_key" | "globalkey" => Some("global_key"), + "provider" => Some("provider"), + _ => None, + } +} + +fn insert_candidate_bool( + source: &serde_json::Map, + target: &mut serde_json::Map, + field: &str, +) { + if let Some(value) = source.get(field).and_then(serde_json::Value::as_bool) { + target.insert(field.to_string(), serde_json::Value::Bool(value)); + } +} + +fn insert_candidate_u64( + source: &serde_json::Map, + target: &mut serde_json::Map, + field: &str, +) { + if let Some(value) = source.get(field).and_then(serde_json::Value::as_u64) { + target.insert(field.to_string(), serde_json::Value::Number(value.into())); + } +} + +fn insert_candidate_i64( + source: &serde_json::Map, + target: &mut serde_json::Map, + field: &str, +) { + if let Some(value) = source.get(field).and_then(serde_json::Value::as_i64) { + target.insert(field.to_string(), serde_json::Value::Number(value.into())); + } +} + +fn insert_candidate_known_string( + source: &serde_json::Map, + target: &mut serde_json::Map, + field: &str, + sanitize: fn(&str) -> Option<&'static str>, +) { + if let Some(value) = source + .get(field) + .and_then(serde_json::Value::as_str) + .and_then(sanitize) + { + target.insert( + field.to_string(), + serde_json::Value::String(value.to_string()), + ); + } +} + +fn sanitize_candidate_phase(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "3c_trial" => Some("3c_trial"), + "provider_request" => Some("provider_request"), + _ => None, + } +} + +fn sanitize_candidate_api_format(value: &str) -> Option<&'static str> { + match aether_ai_formats::normalize_api_format_alias(value).as_str() { + "openai:chat" => Some("openai:chat"), + "openai:responses" => Some("openai:responses"), + "openai:responses:compact" => Some("openai:responses:compact"), + "openai:search" => Some("openai:search"), + "openai:embedding" => Some("openai:embedding"), + "openai:rerank" => Some("openai:rerank"), + "openai:image" => Some("openai:image"), + "openai:video" => Some("openai:video"), + "claude:messages" | "anthropic:messages" => Some("claude:messages"), + "gemini:generate_content" => Some("gemini:generate_content"), + "gemini:interactions" => Some("gemini:interactions"), + "gemini:embedding" => Some("gemini:embedding"), + "gemini:files" => Some("gemini:files"), + "gemini:video" => Some("gemini:video"), + "jina:embedding" => Some("jina:embedding"), + "jina:rerank" => Some("jina:rerank"), + "doubao:embedding" => Some("doubao:embedding"), + "aliyun:multimodal_embedding" => Some("aliyun:multimodal_embedding"), + _ => None, + } +} + +fn sanitize_candidate_execution_strategy(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "gateway_affinity_forward" => Some("gateway_affinity_forward"), + "raw_public_proxy" => Some("raw_public_proxy"), + "local_same_format" => Some("local_same_format"), + "local_cross_format" => Some("local_cross_format"), + _ => None, + } +} + +fn sanitize_candidate_conversion_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "none" => Some("none"), + "request_only" => Some("request_only"), + "response_only" => Some("response_only"), + "bidirectional" => Some("bidirectional"), + _ => None, + } +} + +fn sanitize_candidate_ranking_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "fixedorder" | "fixed_order" => Some("FixedOrder"), + "cacheaffinity" | "cache_affinity" => Some("CacheAffinity"), + "loadbalance" | "load_balance" => Some("LoadBalance"), + _ => None, + } +} + +fn sanitize_candidate_priority_mode(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "provider" => Some("Provider"), + "globalkey" | "global_key" => Some("GlobalKey"), + _ => None, + } +} + +fn sanitize_candidate_promotion_reason(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "cached_affinity" => Some("cached_affinity"), + "local_tunnel" => Some("local_tunnel"), + _ => None, + } +} + +fn sanitize_candidate_demotion_reason(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "cross_format" => Some("cross_format"), + _ => None, + } +} + +fn sanitize_candidate_source(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "execution_runtime" => Some("execution_runtime"), + "upstream_response" => Some("upstream_response"), + "usage_routing_snapshot" => Some("usage_routing_snapshot"), + _ => None, + } +} + +fn sanitize_candidate_execution_path(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "public_proxy_passthrough" => Some("public_proxy_passthrough"), + "local_proxy_passthrough_removed" => Some("local_proxy_passthrough_removed"), + "execution_runtime_sync" => Some("execution_runtime_sync"), + "execution_runtime_stream" => Some("execution_runtime_stream"), + "control_execute_sync" => Some("control_execute_sync"), + "control_execute_stream" => Some("control_execute_stream"), + "local_execution_runtime_miss" => Some("local_execution_runtime_miss"), + "local_execution_planning_timeout" => Some("local_execution_planning_timeout"), + "local_api_key_concurrency_limited" => Some("local_api_key_concurrency_limited"), + "local_auth_denied" => Some("local_auth_denied"), + "local_rate_limited" => Some("local_rate_limited"), + "local_invalid_request" => Some("local_invalid_request"), + "local_route_not_found" => Some("local_route_not_found"), + "local_overloaded" => Some("local_overloaded"), + "distributed_overloaded" => Some("distributed_overloaded"), + "local_ai_public" => Some("local_ai_public"), + "local_execution_loop_detected" => Some("local_execution_loop_detected"), + "tunnel_affinity_forward" => Some("tunnel_affinity_forward"), + "responses_websocket_bridge" => Some("responses_websocket_bridge"), + _ => None, + } +} + +fn sanitize_candidate_upstream_response(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut summary = serde_json::Map::new(); + insert_candidate_known_string(object, &mut summary, "source", sanitize_candidate_source); + if let Some(status_code) = object + .get("status_code") + .and_then(serde_json::Value::as_u64) + .filter(|status_code| *status_code <= u64::from(u16::MAX)) + { + summary.insert( + "status_code".to_string(), + serde_json::Value::Number(status_code.into()), + ); + } + insert_candidate_known_string( + object, + &mut summary, + "body_state", + sanitize_candidate_body_state, + ); + (!summary.is_empty()).then_some(serde_json::Value::Object(summary)) +} + +fn sanitize_candidate_body_state(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "none" => Some("none"), + "inline" => Some("inline"), + "reference" => Some("reference"), + "truncated" => Some("truncated"), + "disabled" => Some("disabled"), + "unavailable" => Some("unavailable"), + _ => None, + } +} + +fn sanitize_candidate_image_progress(value: &serde_json::Value) -> Option { + let object = value.as_object()?; + let mut progress = serde_json::Map::new(); + insert_candidate_known_string( + object, + &mut progress, + "phase", + sanitize_candidate_image_progress_phase, + ); + insert_candidate_known_string( + object, + &mut progress, + "last_upstream_event", + sanitize_candidate_image_upstream_event, + ); + insert_candidate_known_string( + object, + &mut progress, + "last_client_visible_event", + sanitize_candidate_image_client_event, + ); + for field in [ + "upstream_ttfb_ms", + "upstream_sse_frame_count", + "last_upstream_frame_at_unix_ms", + "partial_image_count", + "downstream_heartbeat_count", + "last_downstream_heartbeat_at_unix_ms", + "downstream_heartbeat_interval_ms", + ] { + insert_candidate_u64(object, &mut progress, field); + } + (!progress.is_empty()).then_some(serde_json::Value::Object(progress)) +} + +fn sanitize_candidate_image_progress_phase(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "upstream_connecting" => Some("upstream_connecting"), + "upstream_streaming" => Some("upstream_streaming"), + "upstream_completed" => Some("upstream_completed"), + "failed" => Some("failed"), + _ => None, + } +} + +fn sanitize_candidate_image_upstream_event(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "response.image_generation_call.partial_image" => { + Some("response.image_generation_call.partial_image") + } + "response.completed" => Some("response.completed"), + "response.failed" => Some("response.failed"), + "response.error" => Some("response.error"), + "error" => Some("error"), + "done" => Some("done"), + _ => None, + } +} + +fn sanitize_candidate_image_client_event(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "image_generation.partial_image" => Some("image_generation.partial_image"), + "image_generation.completed" => Some("image_generation.completed"), + "image_generation.failed" => Some("image_generation.failed"), + _ => None, + } +} + +fn sanitize_candidate_pool_group_exhaustion( + value: &serde_json::Value, +) -> Option { + let object = value.as_object()?; + let mut exhaustion = serde_json::Map::new(); + for field in ["scanned_keys", "budget_scanned_keys"] { + insert_candidate_u64(object, &mut exhaustion, field); + } + + let mut sanitized_counts = serde_json::Map::new(); + if let Some(counts) = object + .get("skip_reason_counts") + .and_then(serde_json::Value::as_object) + { + for (reason, count) in counts { + let Some(count) = count.as_u64() else { + continue; + }; + let Some(reason) = sanitize_request_candidate_skip_reason(Some(reason.clone())) else { + continue; + }; + let combined = sanitized_counts + .get(&reason) + .and_then(serde_json::Value::as_u64) + .unwrap_or_default() + .saturating_add(count); + sanitized_counts.insert(reason, serde_json::Value::Number(combined.into())); + } + } + if !sanitized_counts.is_empty() { + exhaustion.insert( + "skip_reason_counts".to_string(), + serde_json::Value::Object(sanitized_counts), + ); + } + + (!exhaustion.is_empty()).then_some(serde_json::Value::Object(exhaustion)) +} + #[async_trait] pub trait RequestCandidateWriteRepository: Send + Sync { async fn upsert( @@ -594,23 +1752,20 @@ pub fn request_candidate_lifecycle_would_regress( existing: RequestCandidateStatus, incoming: RequestCandidateStatus, ) -> bool { - matches!( + let existing_is_terminal = matches!( existing, RequestCandidateStatus::Success | RequestCandidateStatus::Failed | RequestCandidateStatus::Cancelled | RequestCandidateStatus::Skipped - ) && matches!( - incoming, - RequestCandidateStatus::Available - | RequestCandidateStatus::Unused - | RequestCandidateStatus::Pending - | RequestCandidateStatus::Streaming - ) || existing == RequestCandidateStatus::Pending - && matches!( - incoming, - RequestCandidateStatus::Available | RequestCandidateStatus::Unused - ) + ); + + existing_is_terminal && incoming != existing + || existing == RequestCandidateStatus::Pending + && matches!( + incoming, + RequestCandidateStatus::Available | RequestCandidateStatus::Unused + ) || existing == RequestCandidateStatus::Streaming && matches!( incoming, @@ -622,10 +1777,15 @@ pub fn request_candidate_lifecycle_would_regress( #[cfg(test)] mod tests { + use serde_json::json; + use super::{ derive_request_candidate_final_status, request_candidate_lifecycle_would_regress, + sanitize_request_candidate_error_type, sanitize_request_candidate_skip_reason, RequestCandidateFinalStatus, RequestCandidateStatus, StoredRequestCandidate, - UpsertRequestCandidateRecord, + UpsertRequestCandidateRecord, REQUEST_CANDIDATE_ERROR_TYPES, + REQUEST_CANDIDATE_ERROR_TYPE_ALIASES, REQUEST_CANDIDATE_SKIP_REASONS, + UNCLASSIFIED_CANDIDATE_ERROR_TYPE, UNCLASSIFIED_CANDIDATE_SKIP_REASON, }; fn candidate( @@ -729,6 +1889,29 @@ mod tests { } } + #[test] + fn terminal_candidate_cannot_be_rewritten_to_a_different_terminal_fact() { + for existing in [ + RequestCandidateStatus::Success, + RequestCandidateStatus::Failed, + RequestCandidateStatus::Cancelled, + RequestCandidateStatus::Skipped, + ] { + for incoming in [ + RequestCandidateStatus::Success, + RequestCandidateStatus::Failed, + RequestCandidateStatus::Cancelled, + RequestCandidateStatus::Skipped, + ] { + assert_eq!( + request_candidate_lifecycle_would_regress(existing, incoming), + existing != incoming, + "first terminal fact must win: {existing:?} -> {incoming:?}", + ); + } + } + } + #[test] fn candidate_upsert_rejects_nul_in_persistence_identity_fields() { let mut record = UpsertRequestCandidateRecord { @@ -764,4 +1947,457 @@ mod tests { record.key_id = Some("key\0poison".to_string()); assert!(record.validate().is_err()); } + + #[test] + fn candidate_persistence_removes_credentials_and_raw_payloads() { + let mut record = UpsertRequestCandidateRecord { + id: "candidate-1".to_string(), + request_id: "request-1".to_string(), + user_id: None, + api_key_id: None, + username: Some("alice-sensitive-display".to_string()), + api_key_name: Some("production-sensitive-label".to_string()), + candidate_index: 0, + retry_index: 0, + provider_id: None, + endpoint_id: None, + key_id: None, + status: RequestCandidateStatus::Failed, + skip_reason: None, + is_cached: None, + status_code: Some(401), + error_type: Some("upstream_error".to_string()), + error_message: Some("unauthorized".to_string()), + latency_ms: None, + concurrent_requests: None, + extra_data: Some(json!({ + "upstream_url": "https://user:pass@example.com/v1/models?key=vertex-secret#fragment", + "request_path_and_query": "/v1/models?key=client-secret&alt=sse", + "key_name": "credential-label-secret", + "header_rules": [{ + "id": "header-rule-secret", + "action": "set", + "name": "x-auth", + "value": "secret", + "condition": {"path": "$.tenant_secret"} + }], + "body_rules": [{ + "id": "body-rule-secret", + "action": "replace", + "path": "$.api_key", + "pattern": "secret-pattern", + "replacement": "secret" + }], + "proxy": { + "mode": "manual", + "source": "endpoint", + "url": "https://user:pass@proxy.example/private-secret?token=secret", + "ttfb_ms": 17 + }, + "routing_trace": { + "selection_source": "system_default", + "selected_rules": ["tenant-secret"], + "global_candidates": [{"key_id": "secret-key"}], + "client_request_patch_summary": { + "body_paths": ["$.secret"], + "header_names": ["authorization"] + } + }, + "error_flow": { + "stage": "upstream", + "source": "upstream_response", + "classification": "retry_status_code", + "decision": "retry_next_candidate", + "propagation": "captured", + "retryable": true, + "status_code": 401, + "message": "token vertex-secret rejected" + }, + "unknown": {"message": "secret"}, + "free_text": "Bearer secret", + "gateway_execution_runtime": true, + "client_api_format": "OPENAI:RESPONSES", + "provider_api_format": "claude:messages", + "execution_strategy": "local_cross_format", + "conversion_mode": "bidirectional", + "ranking_mode": "CacheAffinity", + "priority_mode": "Provider", + "ranking_index": 2, + "priority_slot": 7, + "promoted_by": "cached_affinity", + "demoted_by": "cross_format", + "upstream_response": { + "source": "upstream_response", + "status_code": 401, + "headers": {"set-cookie": "session=secret"}, + "body": {"error": {"message": "token vertex-secret rejected"}}, + "body_state": "inline" + }, + "image_progress": { + "phase": "upstream_streaming", + "upstream_ttfb_ms": 20, + "upstream_sse_frame_count": 3, + "last_upstream_event": "response.image_generation_call.partial_image", + "last_upstream_frame_at_unix_ms": 1_700_000_000_100_u64, + "partial_image_count": 2, + "last_client_visible_event": "image_generation.partial_image", + "downstream_heartbeat_count": 4, + "last_downstream_heartbeat_at_unix_ms": 1_700_000_000_200_u64, + "downstream_heartbeat_interval_ms": 1_000, + "message": "secret" + }, + "pool_group_exhaustion": { + "scanned_keys": 3, + "budget_scanned_keys": 2, + "skip_reason_counts": { + "pool_cooldown": 2, + "secret reason": 1 + }, + "message": "secret" + } + })), + required_capabilities: Some(json!({ + "cache_1h": "TRUE", + "context_1m": 1, + "gemini_files": 0, + "streaming": "false", + "vision": true, + "tenant_secret_capability": "Bearer secret", + "billing": {"account": "secret-account"} + })), + created_at_unix_ms: None, + started_at_unix_ms: None, + finished_at_unix_ms: None, + }; + + record.sanitize_for_persistence(); + let sanitized_once = record.clone(); + record.sanitize_for_persistence(); + assert_eq!( + record, sanitized_once, + "candidate sanitization must be idempotent" + ); + assert!(record.username.is_none()); + assert!(record.api_key_name.is_none()); + assert!(record.error_message.is_none()); + let extra = record + .extra_data + .as_ref() + .expect("safe candidate data should remain"); + for field in ["request_path_and_query", "key_name", "unknown", "free_text"] { + assert!(extra.get(field).is_none(), "{field} must not be persisted"); + } + assert_eq!(extra["upstream_url"], "https://example.com/"); + assert_eq!(extra["header_rules"]["count"], 1); + assert_eq!(extra["header_rules"]["enabled_count"], 1); + assert_eq!(extra["header_rules"]["conditional_count"], 1); + assert_eq!(extra["header_rules"]["action_counts"]["set"], 1); + assert_eq!(extra["body_rules"]["count"], 1); + assert_eq!(extra["body_rules"]["action_counts"]["replace"], 1); + assert_eq!(extra["proxy"]["mode"], "manual"); + assert_eq!(extra["proxy"]["source"], "endpoint"); + assert_eq!(extra["proxy"]["url"], "https://proxy.example/"); + assert_eq!(extra["proxy"]["ttfb_ms"], 17); + assert_eq!(extra["routing_trace"]["selection_source"], "system_default"); + assert_eq!(extra["routing_trace"]["selected_rule_count"], 1); + assert_eq!(extra["routing_trace"]["global_candidate_count"], 1); + assert_eq!( + extra["routing_trace"]["client_request_patch_summary"]["body_patch_count"], + 1 + ); + assert_eq!( + extra["routing_trace"]["client_request_patch_summary"]["header_patch_count"], + 1 + ); + assert_eq!(extra["error_flow"]["stage"], "upstream"); + assert_eq!(extra["error_flow"]["retryable"], true); + assert_eq!(extra["error_flow"]["status_code"], 401); + assert!(extra["error_flow"].get("message").is_none()); + assert_eq!(extra["gateway_execution_runtime"], true); + assert_eq!(extra["client_api_format"], "openai:responses"); + assert_eq!(extra["provider_api_format"], "claude:messages"); + assert_eq!(extra["execution_strategy"], "local_cross_format"); + assert_eq!(extra["conversion_mode"], "bidirectional"); + assert_eq!(extra["ranking_mode"], "CacheAffinity"); + assert_eq!(extra["priority_mode"], "Provider"); + assert_eq!(extra["ranking_index"], 2); + assert_eq!(extra["priority_slot"], 7); + assert_eq!(extra["promoted_by"], "cached_affinity"); + assert_eq!(extra["demoted_by"], "cross_format"); + assert_eq!(extra["upstream_response"]["source"], "upstream_response"); + assert_eq!(extra["upstream_response"]["status_code"], 401); + assert_eq!(extra["upstream_response"]["body_state"], "inline"); + assert!(extra["upstream_response"].get("headers").is_none()); + assert!(extra["upstream_response"].get("body").is_none()); + assert_eq!( + extra["image_progress"]["last_client_visible_event"], + "image_generation.partial_image" + ); + assert_eq!(extra["image_progress"]["phase"], "upstream_streaming"); + assert_eq!(extra["image_progress"]["upstream_ttfb_ms"], 20); + assert_eq!(extra["image_progress"]["upstream_sse_frame_count"], 3); + assert_eq!( + extra["image_progress"]["last_upstream_event"], + "response.image_generation_call.partial_image" + ); + assert_eq!( + extra["image_progress"]["last_upstream_frame_at_unix_ms"], + 1_700_000_000_100_u64 + ); + assert_eq!(extra["image_progress"]["partial_image_count"], 2); + assert_eq!(extra["image_progress"]["downstream_heartbeat_count"], 4); + assert_eq!( + extra["image_progress"]["last_downstream_heartbeat_at_unix_ms"], + 1_700_000_000_200_u64 + ); + assert_eq!( + extra["image_progress"]["downstream_heartbeat_interval_ms"], + 1_000 + ); + assert!(extra["image_progress"].get("message").is_none()); + assert_eq!(extra["pool_group_exhaustion"]["scanned_keys"], 3); + assert_eq!( + extra["pool_group_exhaustion"]["skip_reason_counts"]["pool_cooldown"], + 2 + ); + assert_eq!( + extra["pool_group_exhaustion"]["skip_reason_counts"]["unclassified_skip"], + 1 + ); + assert!(extra["pool_group_exhaustion"].get("message").is_none()); + assert_eq!( + record.required_capabilities, + Some(json!({ + "cache_1h": true, + "context_1m": true, + "gemini_files": false, + "streaming": false, + "vision": true + })) + ); + + let serialized = serde_json::to_string(&record).expect("candidate should serialize"); + for sensitive in [ + "vertex-secret", + "client-secret", + "credential-label-secret", + "header-rule-secret", + "body-rule-secret", + "x-auth", + "$.api_key", + "secret-pattern", + "tenant-secret", + "secret-key", + "authorization", + "tenant_secret_capability", + "secret-account", + "alice-sensitive-display", + "production-sensitive-label", + ] { + assert!( + !serialized.contains(sensitive), + "candidate must not retain {sensitive}" + ); + } + } + + #[test] + fn candidate_persistence_keeps_only_known_diagnostic_categories() { + let mut record = UpsertRequestCandidateRecord { + id: "candidate-1".to_string(), + request_id: "request-1".to_string(), + user_id: None, + api_key_id: None, + username: None, + api_key_name: None, + candidate_index: 0, + retry_index: 0, + provider_id: None, + endpoint_id: None, + key_id: None, + status: RequestCandidateStatus::Failed, + skip_reason: Some("Provider auth failed with token sk-secret".to_string()), + is_cached: None, + status_code: Some(401), + error_type: Some("sk_secret_looks_like_a_category".to_string()), + error_message: None, + latency_ms: None, + concurrent_requests: None, + extra_data: None, + required_capabilities: None, + created_at_unix_ms: None, + started_at_unix_ms: None, + finished_at_unix_ms: None, + }; + + record.sanitize_for_persistence(); + assert_eq!( + record.skip_reason.as_deref(), + Some(UNCLASSIFIED_CANDIDATE_SKIP_REASON) + ); + assert_eq!( + record.error_type.as_deref(), + Some(UNCLASSIFIED_CANDIDATE_ERROR_TYPE) + ); + + record.skip_reason = Some(" Pool_Cooldown ".to_string()); + record.error_type = Some(" FirstByteTimeout ".to_string()); + record.sanitize_for_persistence(); + assert_eq!(record.skip_reason.as_deref(), Some("pool_cooldown")); + assert_eq!(record.error_type.as_deref(), Some("first_byte_timeout")); + + record.skip_reason = Some("Pool_Score_Member_Missing".to_string()); + record.sanitize_for_persistence(); + assert_eq!( + record.skip_reason.as_deref(), + Some("pool_score_member_missing") + ); + + record.error_type = Some("Upstream5xx".to_string()); + record.sanitize_for_persistence(); + assert_eq!(record.error_type.as_deref(), Some("upstream5xx")); + + for reason in REQUEST_CANDIDATE_SKIP_REASONS { + assert_eq!( + sanitize_request_candidate_skip_reason(Some((*reason).to_string())).as_deref(), + Some(*reason) + ); + } + for error_type in REQUEST_CANDIDATE_ERROR_TYPES { + assert_eq!( + sanitize_request_candidate_error_type(Some((*error_type).to_string())).as_deref(), + Some(*error_type) + ); + } + for (alias, canonical) in REQUEST_CANDIDATE_ERROR_TYPE_ALIASES { + assert_eq!( + sanitize_request_candidate_error_type(Some((*alias).to_string())).as_deref(), + Some(*canonical) + ); + } + } + + #[test] + fn candidate_database_read_sanitizes_legacy_diagnostic_text() { + let candidate = StoredRequestCandidate::new( + "candidate-1".to_string(), + "request-1".to_string(), + None, + None, + Some("legacy-sensitive-user".to_string()), + Some("legacy-sensitive-key-label".to_string()), + 0, + 0, + None, + None, + None, + RequestCandidateStatus::Failed, + Some("legacy secret in skip reason".to_string()), + false, + Some(500), + Some("legacy_secret_code".to_string()), + Some("legacy secret message".to_string()), + None, + None, + None, + None, + 1, + None, + Some(2), + ) + .expect("candidate should build"); + + assert_eq!( + candidate.skip_reason.as_deref(), + Some(UNCLASSIFIED_CANDIDATE_SKIP_REASON) + ); + assert_eq!( + candidate.error_type.as_deref(), + Some(UNCLASSIFIED_CANDIDATE_ERROR_TYPE) + ); + assert!(candidate.error_message.is_none()); + assert!(candidate.username.is_none()); + assert!(candidate.api_key_name.is_none()); + } + + #[test] + fn decision_trace_candidate_sanitizes_catalog_enrichment_and_is_idempotent() { + let mut stored = candidate("candidate-1", RequestCandidateStatus::Failed, Some(500)); + stored.error_message = Some("Bearer candidate-secret".to_string()); + stored.extra_data = Some(json!({ + "gateway_execution_runtime": true, + "request_body": {"password": "candidate-secret"} + })); + stored.required_capabilities = Some(json!({ + "VISION": 1, + "tenant_secret": "candidate-secret" + })); + let mut item = super::DecisionTraceCandidate { + candidate: stored, + provider_name: Some("Provider".to_string()), + provider_website: Some( + "https://user:pass@example.com/private/tenant-secret?token=secret#fragment" + .to_string(), + ), + provider_type: Some("custom".to_string()), + provider_priority: Some(1), + provider_keep_priority_on_conversion: Some(false), + provider_enable_format_conversion: Some(true), + endpoint_api_format: Some("openai:chat".to_string()), + endpoint_api_family: Some("openai".to_string()), + endpoint_kind: Some("chat".to_string()), + endpoint_format_acceptance_config: Some(json!({ + "secret_pattern": "tenant-secret" + })), + provider_key_name: Some("prod".to_string()), + provider_key_auth_type: Some("api_key".to_string()), + provider_key_api_formats: Some(json!([ + "OPENAI:CHAT", + "anthropic:messages", + "tenant-secret-format" + ])), + provider_key_internal_priority: Some(5), + provider_key_global_priority_by_format: Some(json!({ + "tenant-secret-format": 1 + })), + provider_key_capabilities: Some(json!({ + "cache_1h": "TRUE", + "tenant_secret": "candidate-secret" + })), + provider_key_is_active: Some(true), + }; + + item.sanitize_sensitive_diagnostics(); + let sanitized_once = item.clone(); + item.sanitize_sensitive_diagnostics(); + + assert_eq!(item, sanitized_once); + assert_eq!( + item.provider_website.as_deref(), + Some("https://example.com/") + ); + assert!(item.endpoint_format_acceptance_config.is_none()); + assert_eq!( + item.provider_key_api_formats, + Some(json!(["openai:chat", "claude:messages"])) + ); + assert!(item.provider_key_global_priority_by_format.is_none()); + assert_eq!( + item.provider_key_capabilities, + Some(json!({"cache_1h": true})) + ); + assert!(item.candidate.error_message.is_none()); + assert_eq!( + item.candidate.extra_data, + Some(json!({"gateway_execution_runtime": true})) + ); + assert_eq!( + item.candidate.required_capabilities, + Some(json!({"vision": true})) + ); + let serialized = serde_json::to_string(&item).expect("trace candidate should serialize"); + assert!(!serialized.contains("candidate-secret")); + assert!(!serialized.contains("tenant-secret")); + assert!(!serialized.contains("user:pass")); + } } diff --git a/crates/aether-data/contracts/src/repository/gemini_file_mappings.rs b/crates/aether-data/contracts/src/repository/gemini_file_mappings.rs index f4a3caeff..a797180fc 100644 --- a/crates/aether-data/contracts/src/repository/gemini_file_mappings.rs +++ b/crates/aether-data/contracts/src/repository/gemini_file_mappings.rs @@ -1,7 +1,12 @@ use async_trait::async_trait; +pub const GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS: usize = 512; +pub const GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS: usize = 512; +pub const GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS: usize = 255; + #[derive(Debug, Clone, PartialEq, Eq, Default)] pub struct GeminiFileMappingListQuery { + pub user_id: Option, pub include_expired: bool, pub search: Option, pub offset: usize, @@ -103,6 +108,25 @@ impl UpsertGeminiFileMappingRecord { "gemini_file_mappings.file_name is empty".to_string(), )); } + validate_text_length( + &self.file_name, + "file_name", + GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS, + )?; + if let Some(display_name) = self.display_name.as_deref() { + validate_text_length( + display_name, + "display_name", + GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, + )?; + } + if let Some(mime_type) = self.mime_type.as_deref() { + validate_text_length( + mime_type, + "mime_type", + GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS, + )?; + } if self.key_id.trim().is_empty() { return Err(crate::DataLayerError::InvalidInput( "gemini_file_mappings.key_id is empty".to_string(), @@ -117,6 +141,19 @@ impl UpsertGeminiFileMappingRecord { } } +fn validate_text_length( + value: &str, + field: &str, + max_chars: usize, +) -> Result<(), crate::DataLayerError> { + if value.chars().nth(max_chars).is_some() { + return Err(crate::DataLayerError::InvalidInput(format!( + "gemini_file_mappings.{field} exceeds maximum length {max_chars}" + ))); + } + Ok(()) +} + #[async_trait] pub trait GeminiFileMappingReadRepository: Send + Sync { async fn find_by_file_name( @@ -124,6 +161,29 @@ pub trait GeminiFileMappingReadRepository: Send + Sync { file_name: &str, ) -> Result, crate::DataLayerError>; + /// Return an unexpired mapping only when it belongs to `user_id`. + /// + /// Repository implementations must apply the file name, user and expiry + /// predicates in the same read operation. Public callers must not emulate + /// this with an unrestricted lookup followed by an ownership check. + async fn find_active_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, crate::DataLayerError>; + + /// Return an unexpired mapping only when both user and provider key match. + /// This is the routing lookup used before forwarding file object requests + /// to an upstream provider credential. + async fn find_active_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, crate::DataLayerError>; + async fn list_mappings( &self, query: &GeminiFileMappingListQuery, @@ -141,8 +201,32 @@ pub trait GeminiFileMappingWriteRepository: Send + Sync { &self, record: UpsertGeminiFileMappingRecord, ) -> Result; + + /// Insert a new mapping or refresh it only when the persisted owner is the + /// same provider key and user. The ownership check and write must be one + /// atomic repository operation so callers cannot be bypassed with a + /// check-then-write race. + async fn upsert_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, crate::DataLayerError>; + async fn delete_by_file_name(&self, file_name: &str) -> Result; + async fn delete_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result; + + /// Delete only when both persisted ownership dimensions still match. + async fn delete_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result; + async fn delete_by_id( &self, mapping_id: &str, @@ -163,3 +247,64 @@ impl GeminiFileMappingRepository for T where T: GeminiFileMappingReadRepository + GeminiFileMappingWriteRepository { } + +#[cfg(test)] +mod tests { + use super::{ + UpsertGeminiFileMappingRecord, GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, + GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS, + }; + use crate::DataLayerError; + + fn record() -> UpsertGeminiFileMappingRecord { + UpsertGeminiFileMappingRecord { + id: "mapping-1".to_string(), + file_name: "files/example".to_string(), + key_id: "key-1".to_string(), + user_id: Some("user-1".to_string()), + display_name: Some("example".to_string()), + mime_type: Some("application/octet-stream".to_string()), + source_hash: None, + expires_at_unix_secs: 1, + } + } + + #[test] + fn mapping_metadata_accepts_schema_limits() { + let mut record = record(); + record.file_name = "f".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS); + record.display_name = Some("d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS)); + record.mime_type = Some("m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS)); + + record.validate().expect("schema limits should validate"); + } + + #[test] + fn mapping_metadata_rejects_values_beyond_schema_limits() { + for (field, record) in [ + ("file_name", { + let mut record = record(); + record.file_name = "f".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS + 1); + record + }), + ("display_name", { + let mut record = record(); + record.display_name = + Some("d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS + 1)); + record + }), + ("mime_type", { + let mut record = record(); + record.mime_type = Some("m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS + 1)); + record + }), + ] { + let error = record + .validate() + .expect_err("oversized mapping metadata should fail"); + assert!( + matches!(error, DataLayerError::InvalidInput(message) if message.contains(field)) + ); + } + } +} diff --git a/crates/aether-data/contracts/src/repository/management_tokens.rs b/crates/aether-data/contracts/src/repository/management_tokens.rs index 9f8e2474a..402dd7fbc 100644 --- a/crates/aether-data/contracts/src/repository/management_tokens.rs +++ b/crates/aether-data/contracts/src/repository/management_tokens.rs @@ -163,7 +163,7 @@ pub struct ManagementTokenListQuery { pub limit: usize, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct CreateManagementTokenRecord { pub id: String, pub user_id: String, @@ -178,6 +178,46 @@ pub struct CreateManagementTokenRecord { pub is_active: bool, } +impl std::fmt::Debug for CreateManagementTokenRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreateManagementTokenRecord") + .field("id", &self.id) + .field("user_id", &self.user_id) + .field("token_hash", &"[REDACTED]") + .field( + "token_prefix", + &self.token_prefix.as_ref().map(|_| "[REDACTED]"), + ) + .field("name", &self.name) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .field("is_active", &self.is_active) + .finish_non_exhaustive() + } +} + +fn validate_management_token_unix_secs_storage_range( + value: u64, + field_name: &str, +) -> Result<(), crate::DataLayerError> { + if i64::try_from(value).is_err() { + return Err(crate::DataLayerError::InvalidInput(format!( + "{field_name} exceeds supported storage range" + ))); + } + Ok(()) +} + +fn validate_optional_management_token_unix_secs_storage_range( + value: Option, + field_name: &str, +) -> Result<(), crate::DataLayerError> { + match value { + Some(value) => validate_management_token_unix_secs_storage_range(value, field_name), + None => Ok(()), + } +} + impl CreateManagementTokenRecord { pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.id.trim().is_empty() { @@ -223,6 +263,10 @@ impl CreateManagementTokenRecord { } } validate_management_token_permissions(self.permissions.as_ref())?; + validate_optional_management_token_unix_secs_storage_range( + self.expires_at_unix_secs, + "expires_at_unix_secs", + )?; Ok(()) } } @@ -248,6 +292,21 @@ impl UpdateManagementTokenRecord { "token_id is required".to_string(), )); } + if self.clear_description && self.description.is_some() { + return Err(crate::DataLayerError::InvalidInput( + "description and clear_description are mutually exclusive".to_string(), + )); + } + if self.clear_allowed_ips && self.allowed_ips.is_some() { + return Err(crate::DataLayerError::InvalidInput( + "allowed_ips and clear_allowed_ips are mutually exclusive".to_string(), + )); + } + if self.clear_expires_at && self.expires_at_unix_secs.is_some() { + return Err(crate::DataLayerError::InvalidInput( + "expires_at_unix_secs and clear_expires_at are mutually exclusive".to_string(), + )); + } if let Some(name) = &self.name { if name.trim().is_empty() { return Err(crate::DataLayerError::InvalidInput( @@ -273,6 +332,10 @@ impl UpdateManagementTokenRecord { } } validate_management_token_permissions(self.permissions.as_ref())?; + validate_optional_management_token_unix_secs_storage_range( + self.expires_at_unix_secs, + "expires_at_unix_secs", + )?; Ok(()) } } @@ -301,13 +364,127 @@ fn validate_management_token_permissions( Ok(()) } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct RegenerateManagementTokenSecret { pub token_id: String, pub token_hash: String, pub token_prefix: Option, } +impl std::fmt::Debug for RegenerateManagementTokenSecret { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("RegenerateManagementTokenSecret") + .field("token_id", &self.token_id) + .field("token_hash", &"[REDACTED]") + .field( + "token_prefix", + &self.token_prefix.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct ActivateManagementTokenIfMatches { + pub expected_token: StoredManagementToken, + pub token_hash: String, + pub expected_user_security_version: i64, + pub now_unix_secs: u64, +} + +impl std::fmt::Debug for ActivateManagementTokenIfMatches { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ActivateManagementTokenIfMatches") + .field("expected_token", &self.expected_token) + .field("token_hash", &"[REDACTED]") + .field( + "expected_user_security_version", + &self.expected_user_security_version, + ) + .field("now_unix_secs", &self.now_unix_secs) + .finish() + } +} + +impl ActivateManagementTokenIfMatches { + pub fn validate(&self) -> Result<(), crate::DataLayerError> { + if self.expected_token.id.trim().is_empty() { + return Err(crate::DataLayerError::InvalidInput( + "token_id is required".to_string(), + )); + } + if self.expected_token.user_id.trim().is_empty() { + return Err(crate::DataLayerError::InvalidInput( + "user_id is required".to_string(), + )); + } + if self.expected_token.name.trim().is_empty() { + return Err(crate::DataLayerError::InvalidInput( + "token name is required".to_string(), + )); + } + if self.expected_token.is_active { + return Err(crate::DataLayerError::InvalidInput( + "activation snapshot must be inactive".to_string(), + )); + } + if self.token_hash.trim().is_empty() { + return Err(crate::DataLayerError::InvalidInput( + "token_hash is required".to_string(), + )); + } + if self.expected_user_security_version < 0 { + return Err(crate::DataLayerError::InvalidInput( + "expected_user_security_version must not be negative".to_string(), + )); + } + validate_management_token_unix_secs_storage_range(self.now_unix_secs, "now_unix_secs")?; + validate_optional_management_token_unix_secs_storage_range( + self.expected_token.expires_at_unix_secs, + "expires_at_unix_secs", + )?; + validate_optional_management_token_unix_secs_storage_range( + self.expected_token.last_used_at_unix_secs, + "last_used_at_unix_secs", + )?; + validate_optional_management_token_unix_secs_storage_range( + self.expected_token.created_at_unix_ms, + "created_at_unix_ms", + )?; + validate_optional_management_token_unix_secs_storage_range( + self.expected_token.updated_at_unix_secs, + "updated_at_unix_secs", + )?; + validate_management_token_unix_secs_storage_range( + self.expected_token.usage_count, + "usage_count", + )?; + if let Some(allowed_ips) = &self.expected_token.allowed_ips { + let Some(items) = allowed_ips.as_array() else { + return Err(crate::DataLayerError::InvalidInput( + "IP 限制规则必须是数组".to_string(), + )); + }; + if items.is_empty() || items.iter().any(|value| value.as_str().is_none()) { + return Err(crate::DataLayerError::InvalidInput( + "IP 限制规则只能是非空字符串数组".to_string(), + )); + } + } + validate_management_token_permissions(self.expected_token.permissions.as_ref()) + } + + pub fn matches_locked_token_snapshot( + &self, + current: &StoredManagementToken, + current_token_hash: &str, + ) -> bool { + current_token_hash == self.token_hash && current == &self.expected_token + } +} + impl RegenerateManagementTokenSecret { pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.token_id.trim().is_empty() { @@ -360,22 +537,277 @@ pub trait ManagementTokenWriteRepository: Send + Sync { record: &UpdateManagementTokenRecord, ) -> Result, crate::DataLayerError>; + /// Update a token only while it still belongs to `user_id`. + /// + /// Self-service callers must use this owner-scoped mutation instead of + /// relying on a preceding read-side ownership check. + async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, crate::DataLayerError>; + async fn delete_management_token(&self, token_id: &str) -> Result; + async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result; + async fn set_management_token_active( &self, token_id: &str, is_active: bool, ) -> Result, crate::DataLayerError>; + async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, crate::DataLayerError>; + + /// Atomically activate an inactive one-time-install token only while all + /// security-relevant fields still match the session that issued it. + async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result; + + /// Delete an unclaimed one-time-install token only while it still matches + /// the session that created it and remains inactive. + async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result; + async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, ) -> Result, crate::DataLayerError>; + async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, crate::DataLayerError>; + async fn record_management_token_usage( &self, token_id: &str, last_used_ip: Option<&str>, ) -> Result, crate::DataLayerError>; } + +#[cfg(test)] +mod tests { + use super::{ + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, + RegenerateManagementTokenSecret, StoredManagementTokenUserSummary, + UpdateManagementTokenRecord, + }; + + fn activation() -> ActivateManagementTokenIfMatches { + ActivateManagementTokenIfMatches { + expected_token: super::StoredManagementToken::new( + "token-1".to_string(), + "user-1".to_string(), + "install token".to_string(), + ) + .expect("token should build") + .with_display_fields(None, Some("ae_install".to_string()), None) + .with_permissions(Some(serde_json::json!(["admin:proxy_nodes:write"]))) + .with_runtime_fields(Some(2_000_000_000), None, None, 0, false) + .with_timestamps(Some(1_800_000_000), Some(1_800_000_000)), + token_hash: "hash-1".to_string(), + expected_user_security_version: 7, + now_unix_secs: 1_900_000_000, + } + } + + #[test] + fn activation_rejects_timestamps_outside_sql_storage_range() { + let mut mutation = activation(); + mutation.expected_token.expires_at_unix_secs = Some(i64::MAX as u64 + 1); + assert!(mutation.validate().is_err()); + + let mut mutation = activation(); + mutation.now_unix_secs = i64::MAX as u64 + 1; + assert!(mutation.validate().is_err()); + } + + #[test] + fn activation_requires_non_negative_user_security_version() { + let mut mutation = activation(); + mutation.expected_user_security_version = -1; + assert!(mutation.validate().is_err()); + } + + #[test] + fn activation_matches_the_complete_canonical_token_snapshot() { + let mutation = activation(); + assert!( + mutation.matches_locked_token_snapshot(&mutation.expected_token, &mutation.token_hash,) + ); + + let mut changed = mutation.expected_token.clone(); + changed.description = Some("changed after session creation".to_string()); + assert!(!mutation.matches_locked_token_snapshot(&changed, &mutation.token_hash)); + + let mut changed = mutation.expected_token.clone(); + changed.updated_at_unix_secs = changed.updated_at_unix_secs.map(|value| value + 1); + assert!(!mutation.matches_locked_token_snapshot(&changed, &mutation.token_hash)); + + assert!( + !mutation.matches_locked_token_snapshot(&mutation.expected_token, "different-hash",) + ); + } + + #[test] + fn management_token_secret_debug_output_is_redacted() { + let activation = ActivateManagementTokenIfMatches { + token_hash: "activation-token-hash-canary".to_string(), + ..activation() + }; + let activation_debug = format!("{activation:?}"); + assert!(activation_debug.contains("[REDACTED]")); + assert!(!activation_debug.contains("activation-token-hash-canary")); + + let regenerate = RegenerateManagementTokenSecret { + token_id: "token-1".to_string(), + token_hash: "regenerate-token-hash-canary".to_string(), + token_prefix: Some("regenerate-token-prefix-canary".to_string()), + }; + let regenerate_debug = format!("{regenerate:?}"); + assert!(regenerate_debug.contains("[REDACTED]")); + assert!(!regenerate_debug.contains("regenerate-token-hash-canary")); + assert!(!regenerate_debug.contains("regenerate-token-prefix-canary")); + + let user = StoredManagementTokenUserSummary::new( + "user-1".to_string(), + None, + "admin".to_string(), + "admin".to_string(), + ) + .expect("user summary should build"); + let create = CreateManagementTokenRecord { + id: "token-1".to_string(), + user_id: user.id.clone(), + user, + token_hash: "create-token-hash-canary".to_string(), + token_prefix: Some("create-token-prefix-canary".to_string()), + name: "token".to_string(), + description: None, + allowed_ips: None, + permissions: None, + expires_at_unix_secs: None, + is_active: true, + }; + let create_debug = format!("{create:?}"); + assert!(create_debug.contains("[REDACTED]")); + assert!(!create_debug.contains("create-token-hash-canary")); + assert!(!create_debug.contains("create-token-prefix-canary")); + } + + #[test] + fn activation_rejects_json_null_without_conflating_it_with_sql_null() { + let mutation = activation(); + + let mut json_null_allowed_ips = mutation.expected_token.clone(); + json_null_allowed_ips.allowed_ips = Some(serde_json::Value::Null); + assert!( + !mutation.matches_locked_token_snapshot(&json_null_allowed_ips, &mutation.token_hash) + ); + let mut invalid = mutation.clone(); + invalid.expected_token = json_null_allowed_ips; + assert!(invalid.validate().is_err()); + + let mut json_null_permissions = mutation.expected_token.clone(); + json_null_permissions.permissions = Some(serde_json::Value::Null); + assert!( + !mutation.matches_locked_token_snapshot(&json_null_permissions, &mutation.token_hash) + ); + let mut invalid = mutation; + invalid.expected_token = json_null_permissions; + assert!(invalid.validate().is_err()); + } + + #[test] + fn create_and_update_reject_expiry_outside_sql_storage_range() { + let unsupported_expiry = i64::MAX as u64 + 1; + let user = StoredManagementTokenUserSummary::new( + "user-1".to_string(), + None, + "admin".to_string(), + "admin".to_string(), + ) + .expect("user summary should build"); + let create = CreateManagementTokenRecord { + id: "token-1".to_string(), + user_id: user.id.clone(), + user, + token_hash: "hash-1".to_string(), + token_prefix: Some("ae_1234".to_string()), + name: "token".to_string(), + description: None, + allowed_ips: None, + permissions: Some(serde_json::json!(["admin:proxy_nodes:write"])), + expires_at_unix_secs: Some(unsupported_expiry), + is_active: false, + }; + assert!(create.validate().is_err()); + + let update = UpdateManagementTokenRecord { + token_id: "token-1".to_string(), + name: None, + description: None, + clear_description: false, + allowed_ips: None, + clear_allowed_ips: false, + permissions: None, + expires_at_unix_secs: Some(unsupported_expiry), + clear_expires_at: false, + is_active: None, + }; + assert!(update.validate().is_err()); + } + + #[test] + fn management_token_update_rejects_ambiguous_set_and_clear_operations() { + let base = UpdateManagementTokenRecord { + token_id: "token-1".to_string(), + name: None, + description: None, + clear_description: false, + allowed_ips: None, + clear_allowed_ips: false, + permissions: None, + expires_at_unix_secs: None, + clear_expires_at: false, + is_active: None, + }; + + assert!(UpdateManagementTokenRecord { + description: Some("description".to_string()), + clear_description: true, + ..base.clone() + } + .validate() + .is_err()); + assert!(UpdateManagementTokenRecord { + allowed_ips: Some(serde_json::json!(["127.0.0.1"])), + clear_allowed_ips: true, + ..base.clone() + } + .validate() + .is_err()); + assert!(UpdateManagementTokenRecord { + expires_at_unix_secs: Some(100), + clear_expires_at: true, + ..base + } + .validate() + .is_err()); + } +} diff --git a/crates/aether-data/contracts/src/repository/oauth_providers.rs b/crates/aether-data/contracts/src/repository/oauth_providers.rs index 3e2b9897d..7deec281c 100644 --- a/crates/aether-data/contracts/src/repository/oauth_providers.rs +++ b/crates/aether-data/contracts/src/repository/oauth_providers.rs @@ -1,6 +1,213 @@ use async_trait::async_trait; +use std::net::IpAddr; +use url::{Host, Url}; -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub fn validate_oauth_redirect_uri(value: &str) -> Result<(), String> { + let parsed = + Url::parse(value).map_err(|_| "redirect_uri must be an absolute URL".to_string())?; + let Some(host) = parsed.host() else { + return Err("redirect_uri must be an absolute URL".to_string()); + }; + let is_loopback = match host { + Host::Domain(domain) => domain.eq_ignore_ascii_case("localhost"), + Host::Ipv4(address) => address.is_loopback(), + Host::Ipv6(address) => address.is_loopback(), + }; + if parsed.scheme() != "https" && !(parsed.scheme() == "http" && is_loopback) { + return Err( + "redirect_uri must use https, except for localhost or loopback IPs".to_string(), + ); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err("redirect_uri must not contain URL credentials".to_string()); + } + if parsed.fragment().is_some() { + return Err("redirect_uri must not contain a fragment".to_string()); + } + Ok(()) +} + +pub fn validate_oauth_frontend_callback_url(value: &str) -> Result<(), String> { + let parsed = Url::parse(value) + .map_err(|_| "frontend_callback_url must be an absolute URL".to_string())?; + let Some(host) = parsed.host() else { + return Err("frontend_callback_url must be an absolute URL".to_string()); + }; + let is_loopback = match host { + Host::Domain(domain) => domain.eq_ignore_ascii_case("localhost"), + Host::Ipv4(address) => address.is_loopback(), + Host::Ipv6(address) => address.is_loopback(), + }; + if parsed.scheme() != "https" && !(parsed.scheme() == "http" && is_loopback) { + return Err( + "frontend_callback_url must use https, except for localhost or loopback IPs" + .to_string(), + ); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err("frontend_callback_url must not contain URL credentials".to_string()); + } + if parsed.query().is_some() || parsed.fragment().is_some() { + return Err("frontend_callback_url must not contain a query or fragment".to_string()); + } + if !parsed + .path() + .trim_end_matches('/') + .ends_with("/auth/callback") + { + return Err("frontend_callback_url path must end with /auth/callback".to_string()); + } + Ok(()) +} + +pub fn validate_oauth_provider_endpoint_config( + provider_type: &str, + authorization_url_override: Option<&str>, + token_url_override: Option<&str>, + userinfo_url_override: Option<&str>, + extra_config: Option<&serde_json::Value>, +) -> Result<(), String> { + let provider_type = provider_type.trim().to_ascii_lowercase(); + let mut provider_chars = provider_type.chars(); + if !(3..=64).contains(&provider_type.len()) + || !provider_chars + .next() + .is_some_and(|character| character.is_ascii_lowercase()) + || !provider_chars.all(|character| { + character.is_ascii_lowercase() + || character.is_ascii_digit() + || matches!(character, '_' | '-') + }) + { + return Err("provider_type contains invalid characters".to_string()); + } + let is_custom = provider_type == "custom_oidc" + || provider_type.starts_with("custom_oidc_") + || provider_type.starts_with("custom_") + || provider_type.starts_with("oidc_"); + let allowed_domains = if provider_type == "linuxdo" { + vec![ + "linux.do".to_string(), + "connect.linux.do".to_string(), + "connect.linuxdo.org".to_string(), + ] + } else if is_custom { + oauth_custom_allowed_domains(extra_config)? + } else { + return Err("unsupported identity OAuth provider_type".to_string()); + }; + + for (field, value) in [ + ("authorization_url_override", authorization_url_override), + ("token_url_override", token_url_override), + ("userinfo_url_override", userinfo_url_override), + ] { + let value = value.map(str::trim).filter(|value| !value.is_empty()); + if is_custom && value.is_none() { + return Err(format!("custom OIDC providers must configure {field}")); + } + if let Some(value) = value { + validate_oauth_endpoint_url(field, value, &allowed_domains)?; + } + } + Ok(()) +} + +fn oauth_custom_allowed_domains( + extra_config: Option<&serde_json::Value>, +) -> Result, String> { + let values = extra_config + .and_then(serde_json::Value::as_object) + .and_then(|object| { + object + .get("allowed_domains") + .or_else(|| object.get("oauth_allowed_domains")) + }) + .and_then(serde_json::Value::as_array) + .ok_or_else(|| { + "custom OIDC providers must configure extra_config.allowed_domains".to_string() + })?; + let mut domains = Vec::with_capacity(values.len()); + for value in values { + let domain = value + .as_str() + .map(str::trim) + .map(|value| value.trim_end_matches('.')) + .filter(|value| !value.is_empty()) + .ok_or_else(|| "OAuth allowed_domains must contain only host names".to_string())?; + if domain.contains('/') + || domain.contains('\\') + || domain.contains('@') + || domain.contains(':') + || domain.contains(char::is_whitespace) + || domain.parse::().is_ok() + { + return Err( + "OAuth allowed_domains must contain DNS host names, not IP literals".to_string(), + ); + } + domains.push(domain.to_ascii_lowercase()); + } + if domains.is_empty() { + return Err("custom OIDC providers must configure allowed domains".to_string()); + } + Ok(domains) +} + +fn validate_oauth_endpoint_url( + field: &str, + value: &str, + allowed_domains: &[String], +) -> Result<(), String> { + let parsed = Url::parse(value).map_err(|_| format!("{field} must be an absolute URL"))?; + if parsed.scheme() != "https" || parsed.host_str().is_none() { + return Err(format!("{field} must be an absolute https URL")); + } + if matches!(parsed.host(), Some(Host::Ipv4(_)) | Some(Host::Ipv6(_))) { + return Err(format!( + "{field} must use a DNS host name, not an IP literal" + )); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err(format!("{field} must not contain URL credentials")); + } + if parsed.fragment().is_some() { + return Err(format!("{field} must not contain a fragment")); + } + if field == "authorization_url_override" { + for (name, _) in parsed.query_pairs() { + if matches!( + name.to_ascii_lowercase().as_str(), + "response_type" + | "client_id" + | "redirect_uri" + | "state" + | "scope" + | "code_challenge" + | "code_challenge_method" + ) { + return Err(format!( + "{field} must not predefine OAuth authorization parameters" + )); + } + } + } + if !allowed_domains.is_empty() { + let host = parsed + .host_str() + .map(|value| value.trim_end_matches('.').to_ascii_lowercase()) + .unwrap_or_default(); + if !allowed_domains + .iter() + .any(|domain| host == *domain || host.ends_with(&format!(".{domain}"))) + { + return Err(format!("{field} host is not in the provider allowlist")); + } + } + Ok(()) +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredOAuthProviderConfig { pub provider_type: String, pub display_name: String, @@ -20,6 +227,51 @@ pub struct StoredOAuthProviderConfig { pub updated_at_unix_secs: Option, } +impl std::fmt::Debug for StoredOAuthProviderConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredOAuthProviderConfig") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("client_id", &self.client_id) + .field( + "client_secret_encrypted", + &self.client_secret_encrypted.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "authorization_url_override", + &self + .authorization_url_override + .as_ref() + .map(|_| "[REDACTED]"), + ) + .field( + "token_url_override", + &self.token_url_override.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "userinfo_url_override", + &self.userinfo_url_override.as_ref().map(|_| "[REDACTED]"), + ) + .field("scopes", &self.scopes) + .field("redirect_uri", &self.redirect_uri) + .field("frontend_callback_url", &self.frontend_callback_url) + .field( + "attribute_mapping", + &self.attribute_mapping.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "extra_config", + &self.extra_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("icon_url", &self.icon_url) + .field("is_enabled", &self.is_enabled) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish() + } +} + impl StoredOAuthProviderConfig { pub fn new( provider_type: String, @@ -110,7 +362,7 @@ impl StoredOAuthProviderConfig { } } -#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)] +#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)] pub enum EncryptedSecretUpdate { #[default] Preserve, @@ -118,6 +370,16 @@ pub enum EncryptedSecretUpdate { Set(String), } +impl std::fmt::Debug for EncryptedSecretUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Preserve => formatter.write_str("Preserve"), + Self::Clear => formatter.write_str("Clear"), + Self::Set(_) => formatter.write_str("Set([REDACTED])"), + } + } +} + impl EncryptedSecretUpdate { pub fn mode_name(&self) -> &'static str { match self { @@ -135,7 +397,7 @@ impl EncryptedSecretUpdate { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct UpsertOAuthProviderConfigRecord { pub provider_type: String, pub display_name: String, @@ -153,6 +415,56 @@ pub struct UpsertOAuthProviderConfigRecord { pub is_enabled: bool, } +impl std::fmt::Debug for UpsertOAuthProviderConfigRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("UpsertOAuthProviderConfigRecord") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("client_id", &self.client_id) + .field("client_secret_encrypted", &self.client_secret_encrypted) + .field( + "authorization_url_override", + &self + .authorization_url_override + .as_ref() + .map(|_| "[REDACTED]"), + ) + .field( + "token_url_override", + &self.token_url_override.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "userinfo_url_override", + &self.userinfo_url_override.as_ref().map(|_| "[REDACTED]"), + ) + .field("scopes", &self.scopes) + .field("redirect_uri", &"[REDACTED]") + .field("frontend_callback_url", &"[REDACTED]") + .field( + "attribute_mapping", + &self.attribute_mapping.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "extra_config", + &self.extra_config.as_ref().map(|_| "[REDACTED]"), + ) + .field("icon_url", &self.icon_url) + .field("is_enabled", &self.is_enabled) + .finish() + } +} + +// Keep the returned provider value inline: this outcome is part of the +// repository API and the successful value is consumed immediately by callers. +// Boxing would be an API/ownership change for no security benefit. +#[allow(clippy::large_enum_variant)] +#[derive(Debug, Clone, PartialEq)] +pub enum UpsertOAuthProviderConfigOutcome { + Upserted(StoredOAuthProviderConfig), + DisableRequiresConfirmation { affected_count: usize }, +} + impl UpsertOAuthProviderConfigRecord { pub fn validate(&self) -> Result<(), crate::DataLayerError> { if self.provider_type.trim().is_empty() { @@ -175,11 +487,23 @@ impl UpsertOAuthProviderConfigRecord { "redirect_uri is required".to_string(), )); } + validate_oauth_redirect_uri(self.redirect_uri.trim()) + .map_err(crate::DataLayerError::InvalidInput)?; if self.frontend_callback_url.trim().is_empty() { return Err(crate::DataLayerError::InvalidInput( "frontend_callback_url is required".to_string(), )); } + validate_oauth_frontend_callback_url(self.frontend_callback_url.trim()) + .map_err(crate::DataLayerError::InvalidInput)?; + validate_oauth_provider_endpoint_config( + &self.provider_type, + self.authorization_url_override.as_deref(), + self.token_url_override.as_deref(), + self.userinfo_url_override.as_deref(), + self.extra_config.as_ref(), + ) + .map_err(crate::DataLayerError::InvalidInput)?; if let Some(scopes) = &self.scopes { for scope in scopes { if scope.trim().is_empty() { @@ -193,6 +517,190 @@ impl UpsertOAuthProviderConfigRecord { } } +// Validation tests are kept beside the validation implementation so changes +// to endpoint policy are reviewed together. The repository traits below are +// intentionally declared after this focused test module for API readability. +#[allow(clippy::items_after_test_module)] +#[cfg(test)] +mod tests { + use super::{ + validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config, + validate_oauth_redirect_uri, EncryptedSecretUpdate, StoredOAuthProviderConfig, + }; + + #[test] + fn oauth_provider_debug_output_redacts_encrypted_client_secrets() { + let secret = "debug-secret-oauth-provider-ciphertext"; + let provider = StoredOAuthProviderConfig::new( + "linuxdo".to_string(), + "Linux.do".to_string(), + "client-id".to_string(), + "https://gateway.example/api/oauth/linuxdo/callback".to_string(), + "https://frontend.example/auth/callback".to_string(), + ) + .expect("provider should build") + .with_config_fields( + Some(secret.to_string()), + None, + None, + None, + None, + None, + None, + None, + true, + ); + + for rendered in [ + format!("{provider:?}"), + format!("{:?}", EncryptedSecretUpdate::Set(secret.to_string())), + ] { + assert!(!rendered.contains(secret)); + assert!(rendered.contains("[REDACTED]")); + } + } + + #[test] + fn oauth_redirect_uri_requires_absolute_http_url_without_credentials() { + assert!( + validate_oauth_redirect_uri("https://gateway.example/api/oauth/custom/callback") + .is_ok() + ); + assert!( + validate_oauth_redirect_uri("http://localhost:8080/api/oauth/custom/callback").is_ok() + ); + for value in [ + "http://gateway.example/api/oauth/custom/callback", + "/api/oauth/custom/callback", + "javascript:alert(1)", + "https://user:password@gateway.example/api/oauth/custom/callback", + ] { + assert!( + validate_oauth_redirect_uri(value).is_err(), + "accepted {value}" + ); + } + } + + #[test] + fn oauth_provider_endpoints_require_https_and_the_configured_domain() { + let extra = serde_json::json!({"allowed_domains": ["idp.example"]}); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://idp.example/oauth/authorize"), + Some("https://idp.example/oauth/token"), + Some("https://accounts.idp.example/oauth/userinfo"), + Some(&extra), + ) + .is_ok()); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://idp.example/oauth/authorize"), + Some("https://attacker.example/oauth/token"), + Some("https://idp.example/oauth/userinfo"), + Some(&extra), + ) + .is_err()); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://127.0.0.1/oauth/authorize"), + Some("https://idp.example/oauth/token"), + Some("https://idp.example/oauth/userinfo"), + Some(&serde_json::json!({"allowed_domains": ["127.0.0.1"]})), + ) + .is_err()); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://idp.example/oauth/authorize?client_id=attacker"), + Some("https://idp.example/oauth/token"), + Some("https://idp.example/oauth/userinfo"), + Some(&extra), + ) + .is_err()); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://idp.example/oauth/authorize"), + Some("http://idp.example/oauth/token"), + Some("https://idp.example/oauth/userinfo"), + Some(&extra), + ) + .is_err()); + } + + #[test] + fn oauth_provider_endpoints_reject_ip_literals_and_predefined_authorization_parameters() { + for host in ["127.0.0.1", "[::1]"] { + let extra = serde_json::json!({"allowed_domains": [host]}); + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some(&format!("https://{host}/oauth/authorize")), + Some(&format!("https://{host}/oauth/token")), + Some(&format!("https://{host}/oauth/userinfo")), + Some(&extra), + ) + .is_err()); + } + + let extra = serde_json::json!({"allowed_domains": ["idp.example"]}); + for name in [ + "response_type", + "client_id", + "redirect_uri", + "state", + "scope", + "code_challenge", + "code_challenge_method", + ] { + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some(&format!( + "https://idp.example/oauth/authorize?{name}=attacker" + )), + Some("https://idp.example/oauth/token?tenant=workforce"), + Some("https://idp.example/oauth/userinfo?schema=current"), + Some(&extra), + ) + .is_err()); + } + + assert!(validate_oauth_provider_endpoint_config( + "custom_oidc_work", + Some("https://idp.example/oauth/authorize?tenant=workforce"), + Some("https://idp.example/oauth/token?tenant=workforce"), + Some("https://idp.example/oauth/userinfo?schema=current"), + Some(&extra), + ) + .is_ok()); + } + + #[test] + fn oauth_frontend_callback_rejects_token_exfiltration_targets() { + for value in [ + "https://frontend.example/auth/callback", + "http://localhost:5173/auth/callback", + "http://127.0.0.1:5173/auth/callback", + "http://[::1]:5173/auth/callback", + ] { + assert!( + validate_oauth_frontend_callback_url(value).is_ok(), + "rejected {value}" + ); + } + for value in [ + "http://attacker.example/auth/callback", + "https://user:password@frontend.example/auth/callback", + "https://frontend.example/auth/callback?next=https://attacker.example", + "https://frontend.example/auth/callback#access_token=stolen", + "https://frontend.example/not-the-callback", + ] { + assert!( + validate_oauth_frontend_callback_url(value).is_err(), + "accepted {value}" + ); + } + } +} + #[async_trait] pub trait OAuthProviderReadRepository: Send + Sync { async fn list_oauth_provider_configs( @@ -216,11 +724,42 @@ pub trait OAuthProviderWriteRepository: Send + Sync { async fn upsert_oauth_provider_config( &self, record: &UpsertOAuthProviderConfigRecord, - ) -> Result; + ) -> Result { + match self + .upsert_oauth_provider_config_guarded(record, false, false, 0) + .await? + { + UpsertOAuthProviderConfigOutcome::Upserted(provider) => Ok(provider), + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { affected_count } => { + Err(crate::DataLayerError::InvalidInput(format!( + "disabling OAuth provider requires confirmation for {affected_count} affected users" + ))) + } + } + } - async fn delete_oauth_provider_config( + async fn upsert_oauth_provider_config_guarded( + &self, + record: &UpsertOAuthProviderConfigRecord, + ldap_exclusive: bool, + force_disable: bool, + locked_users_snapshot: usize, + ) -> Result; + + /// Replace only the stored client secret when the provider and exact previously observed + /// ciphertext still match. Implementations must not modify `updated_at` or any non-secret + /// provider field; this is used by lazy record-bound ciphertext migration. + async fn compare_and_swap_oauth_provider_client_secret( &self, provider_type: &str, + expected: &str, + replacement: &str, + ) -> Result; + + async fn delete_oauth_provider_config_if_unlinked( + &self, + provider_type: &str, + has_links_snapshot: bool, ) -> Result; } diff --git a/crates/aether-data/contracts/src/repository/provider_catalog/mod.rs b/crates/aether-data/contracts/src/repository/provider_catalog/mod.rs index 7d35e5e05..33c248ff4 100644 --- a/crates/aether-data/contracts/src/repository/provider_catalog/mod.rs +++ b/crates/aether-data/contracts/src/repository/provider_catalog/mod.rs @@ -4,11 +4,12 @@ mod types; pub use snapshot::ProviderCatalogSnapshot; pub use types::{ ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, - ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate, - ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, + ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate, + ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, - ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, + ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate, + ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceExpectation, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, diff --git a/crates/aether-data/contracts/src/repository/provider_catalog/types.rs b/crates/aether-data/contracts/src/repository/provider_catalog/types.rs index 5e1f277de..37f65824d 100644 --- a/crates/aether-data/contracts/src/repository/provider_catalog/types.rs +++ b/crates/aether-data/contracts/src/repository/provider_catalog/types.rs @@ -1,11 +1,27 @@ use async_trait::async_trait; -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +const REDACTED_DEBUG_VALUE: &str = "[REDACTED]"; + +fn redacted_debug_option(value: &Option) -> Option<&'static str> { + value.as_ref().map(|_| REDACTED_DEBUG_VALUE) +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogUpstreamMetadataNamespaceUpdate { pub namespace: String, pub value: serde_json::Value, } +impl std::fmt::Debug for ProviderCatalogUpstreamMetadataNamespaceUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogUpstreamMetadataNamespaceUpdate") + .field("namespace", &self.namespace) + .field("value", &REDACTED_DEBUG_VALUE) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyAdaptiveState { pub learned_rpm_limit: Option, @@ -28,7 +44,7 @@ impl ProviderCatalogKeyAdaptiveState { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyAdaptiveStateUpdate { pub key_id: String, /// Optional auth_config fence for request-owned adaptive feedback. @@ -41,7 +57,24 @@ pub struct ProviderCatalogKeyAdaptiveStateUpdate { pub updated_at_unix_secs: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for ProviderCatalogKeyAdaptiveStateUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyAdaptiveStateUpdate") + .field("key_id", &self.key_id) + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field("expected", &self.expected) + .field("next", &self.next) + .field("status_snapshot_patch", &REDACTED_DEBUG_VALUE) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyRuntimeMetadataUpdate { pub key_id: String, pub namespace: String, @@ -59,7 +92,24 @@ pub struct ProviderCatalogKeyRuntimeMetadataUpdate { pub updated_at_unix_secs: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for ProviderCatalogKeyRuntimeMetadataUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyRuntimeMetadataUpdate") + .field("key_id", &self.key_id) + .field("namespace", &self.namespace) + .field( + "expected_upstream_metadata_value", + &redacted_debug_option(&self.expected_upstream_metadata_value), + ) + .field("upstream_metadata_value", &REDACTED_DEBUG_VALUE) + .field("status_snapshot_patch", &REDACTED_DEBUG_VALUE) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyStatusSnapshotUpdate { pub key_id: String, /// Top-level status fields owned by the caller. @@ -67,10 +117,21 @@ pub struct ProviderCatalogKeyStatusSnapshotUpdate { pub updated_at_unix_secs: Option, } +impl std::fmt::Debug for ProviderCatalogKeyStatusSnapshotUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyStatusSnapshotUpdate") + .field("key_id", &self.key_id) + .field("status_snapshot_patch", &REDACTED_DEBUG_VALUE) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish() + } +} + /// Credential context observed before an OAuth refresh started. Repositories /// compare every field atomically with the runtime-state update so an /// administrator replacement cannot be overwritten by an older refresh. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyOAuthCredentialFence { /// Exact nullable ciphertext stored in `provider_api_keys.api_key`. pub encrypted_api_key: Option, @@ -79,10 +140,25 @@ pub struct ProviderCatalogKeyOAuthCredentialFence { pub provider_type: String, } +impl std::fmt::Debug for ProviderCatalogKeyOAuthCredentialFence { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyOAuthCredentialFence") + .field( + "encrypted_api_key", + &redacted_debug_option(&self.encrypted_api_key), + ) + .field("auth_type", &self.auth_type) + .field("provider_id", &self.provider_id) + .field("provider_type", &self.provider_type) + .finish() + } +} + /// Administrator-owned key replacement fenced by the exact credential state /// observed while the edit was prepared. This prevents an older admin request /// from restoring credentials that a concurrent request already replaced. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyAdminCasUpdate { pub expected_encrypted_auth_config: Option, pub expected_credential: ProviderCatalogKeyOAuthCredentialFence, @@ -102,9 +178,28 @@ pub struct ProviderCatalogKeyAdminCasUpdate { pub reset_oauth_runtime: bool, } +impl std::fmt::Debug for ProviderCatalogKeyAdminCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyAdminCasUpdate") + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field("expected_credential", &self.expected_credential) + .field("key", &self.key) + .field( + "codex_rotation", + &redacted_debug_option(&self.codex_rotation), + ) + .field("reset_oauth_runtime", &self.reset_oauth_runtime) + .finish() + } +} + /// Atomic key deletion fenced by the exact OAuth credential generation that /// produced the terminal failure and, when supplied, one metadata namespace. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyOAuthCredentialCasDelete { pub key_id: String, pub expected_encrypted_auth_config: Option, @@ -114,22 +209,53 @@ pub struct ProviderCatalogKeyOAuthCredentialCasDelete { Option, } +impl std::fmt::Debug for ProviderCatalogKeyOAuthCredentialCasDelete { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyOAuthCredentialCasDelete") + .field("key_id", &self.key_id) + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field("expected_credential", &self.expected_credential) + .field( + "expected_upstream_metadata_namespace", + &self.expected_upstream_metadata_namespace, + ) + .finish() + } +} + /// Optional single-namespace metadata fence for an OAuth runtime CAS. /// /// The outer option on the owning update controls whether the namespace is /// compared. Within an expectation, `None` requires the namespace to be absent. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogUpstreamMetadataNamespaceExpectation { pub namespace: String, #[serde(default, skip_serializing_if = "Option::is_none")] pub expected_value: Option, } +impl std::fmt::Debug for ProviderCatalogUpstreamMetadataNamespaceExpectation { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogUpstreamMetadataNamespaceExpectation") + .field("namespace", &self.namespace) + .field( + "expected_value", + &redacted_debug_option(&self.expected_value), + ) + .finish() + } +} + /// Agent/runtime-owned OAuth state update fenced by the exact encrypted /// auth_config and, when supplied, credential context observed before the /// refresh started. Repositories must update only these fields and return /// `false` when an expected value changed. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate { pub key_id: String, pub expected_encrypted_auth_config: Option, @@ -162,7 +288,53 @@ pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate { pub updated_at_unix_secs: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for ProviderCatalogKeyOAuthRuntimeStateCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyOAuthRuntimeStateCasUpdate") + .field("key_id", &self.key_id) + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field("expected_credential", &self.expected_credential) + .field( + "expected_upstream_metadata_namespace", + &self.expected_upstream_metadata_namespace, + ) + .field("encrypted_auth_config", &REDACTED_DEBUG_VALUE) + .field( + "encrypted_api_key_update", + &redacted_debug_option(&self.encrypted_api_key_update), + ) + .field( + "expires_at_unix_secs_update", + &self.expires_at_unix_secs_update, + ) + .field( + "oauth_invalid_at_unix_secs", + &self.oauth_invalid_at_unix_secs, + ) + .field( + "oauth_invalid_reason", + &redacted_debug_option(&self.oauth_invalid_reason), + ) + .field( + "upstream_metadata_patch", + &redacted_debug_option(&self.upstream_metadata_patch), + ) + .field( + "upstream_metadata_namespace_to_remove", + &self.upstream_metadata_namespace_to_remove, + ) + .field("status_snapshot_patch", &REDACTED_DEBUG_VALUE) + .field("reset_error_count", &self.reset_error_count) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProviderCatalogKeyHealthStateUpdate { pub key_id: String, /// Optional auth_config fence for lifecycle-owned health recovery. @@ -174,7 +346,108 @@ pub struct ProviderCatalogKeyHealthStateUpdate { pub circuit_breaker_by_format: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for ProviderCatalogKeyHealthStateUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyHealthStateUpdate") + .field("key_id", &self.key_id) + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field("expected_health_by_format", &self.expected_health_by_format) + .field( + "expected_circuit_breaker_by_format", + &self.expected_circuit_breaker_by_format, + ) + .field("health_by_format", &self.health_by_format) + .field("circuit_breaker_by_format", &self.circuit_breaker_by_format) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct ProviderCatalogProviderConfigCasUpdate { + pub provider_id: String, + pub expected_config: Option, + pub config: Option, +} + +impl std::fmt::Debug for ProviderCatalogProviderConfigCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogProviderConfigCasUpdate") + .field("provider_id", &self.provider_id) + .field( + "expected_config", + &redacted_debug_option(&self.expected_config), + ) + .field("config", &redacted_debug_option(&self.config)) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct ProviderCatalogProxyCasUpdate { + pub record_id: String, + pub expected_proxy: Option, + pub proxy: Option, +} + +impl std::fmt::Debug for ProviderCatalogProxyCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogProxyCasUpdate") + .field("record_id", &self.record_id) + .field( + "expected_proxy", + &redacted_debug_option(&self.expected_proxy), + ) + .field("proxy", &redacted_debug_option(&self.proxy)) + .finish() + } +} + +/// Secret-only migration fenced by the complete catalog-key credential +/// identity. The provider fence prevents a legacy credential from being +/// re-encrypted for an obsolete provider after a concurrent key move. +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct ProviderCatalogKeyCredentialsCasUpdate { + pub key_id: String, + pub expected_provider_id: String, + pub expected_encrypted_api_key: Option, + pub expected_encrypted_auth_config: Option, + pub encrypted_api_key: Option, + pub encrypted_auth_config: Option, +} + +impl std::fmt::Debug for ProviderCatalogKeyCredentialsCasUpdate { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderCatalogKeyCredentialsCasUpdate") + .field("key_id", &self.key_id) + .field("expected_provider_id", &self.expected_provider_id) + .field( + "expected_encrypted_api_key", + &redacted_debug_option(&self.expected_encrypted_api_key), + ) + .field( + "expected_encrypted_auth_config", + &redacted_debug_option(&self.expected_encrypted_auth_config), + ) + .field( + "encrypted_api_key", + &redacted_debug_option(&self.encrypted_api_key), + ) + .field( + "encrypted_auth_config", + &redacted_debug_option(&self.encrypted_auth_config), + ) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredProviderCatalogProvider { pub id: String, pub name: String, @@ -201,6 +474,23 @@ pub struct StoredProviderCatalogProvider { pub updated_at_unix_secs: Option, } +impl std::fmt::Debug for StoredProviderCatalogProvider { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredProviderCatalogProvider") + .field("id", &self.id) + .field("name", &self.name) + .field("provider_type", &self.provider_type) + .field("billing_type", &self.billing_type) + .field("is_active", &self.is_active) + .field("proxy", &redacted_debug_option(&self.proxy)) + .field("config", &redacted_debug_option(&self.config)) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish_non_exhaustive() + } +} + impl StoredProviderCatalogProvider { pub fn new( id: String, @@ -311,7 +601,7 @@ impl StoredProviderCatalogProvider { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredProviderCatalogEndpoint { pub id: String, pub provider_id: String, @@ -332,6 +622,27 @@ pub struct StoredProviderCatalogEndpoint { pub updated_at_unix_secs: Option, } +impl std::fmt::Debug for StoredProviderCatalogEndpoint { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredProviderCatalogEndpoint") + .field("id", &self.id) + .field("provider_id", &self.provider_id) + .field("api_format", &self.api_format) + .field("api_family", &self.api_family) + .field("endpoint_kind", &self.endpoint_kind) + .field("is_active", &self.is_active) + .field("base_url", &REDACTED_DEBUG_VALUE) + .field("header_rules", &redacted_debug_option(&self.header_rules)) + .field("body_rules", &redacted_debug_option(&self.body_rules)) + .field("config", &redacted_debug_option(&self.config)) + .field("proxy", &redacted_debug_option(&self.proxy)) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish_non_exhaustive() + } +} + impl StoredProviderCatalogEndpoint { pub fn new( id: String, @@ -413,7 +724,7 @@ impl StoredProviderCatalogEndpoint { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredProviderCatalogKey { pub id: String, pub provider_id: String, @@ -470,7 +781,44 @@ pub struct StoredProviderCatalogKey { pub circuit_breaker_by_format: Option, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for StoredProviderCatalogKey { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredProviderCatalogKey") + .field("id", &self.id) + .field("provider_id", &self.provider_id) + .field("name", &self.name) + .field("auth_type", &self.auth_type) + .field("is_active", &self.is_active) + .field( + "encrypted_api_key", + &redacted_debug_option(&self.encrypted_api_key), + ) + .field( + "encrypted_auth_config", + &redacted_debug_option(&self.encrypted_auth_config), + ) + .field("proxy", &redacted_debug_option(&self.proxy)) + .field("fingerprint", &redacted_debug_option(&self.fingerprint)) + .field( + "upstream_metadata", + &redacted_debug_option(&self.upstream_metadata), + ) + .field( + "oauth_invalid_reason", + &redacted_debug_option(&self.oauth_invalid_reason), + ) + .field( + "status_snapshot", + &redacted_debug_option(&self.status_snapshot), + ) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .field("updated_at_unix_secs", &self.updated_at_unix_secs) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq)] pub struct StoredProviderCatalogKeyMaintenanceSummary { pub id: String, pub provider_id: String, @@ -478,6 +826,21 @@ pub struct StoredProviderCatalogKeyMaintenanceSummary { pub upstream_metadata: Option, } +impl std::fmt::Debug for StoredProviderCatalogKeyMaintenanceSummary { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredProviderCatalogKeyMaintenanceSummary") + .field("id", &self.id) + .field("provider_id", &self.provider_id) + .field("is_active", &self.is_active) + .field( + "upstream_metadata", + &redacted_debug_option(&self.upstream_metadata), + ) + .finish() + } +} + impl StoredProviderCatalogKey { pub fn new( id: String, @@ -590,6 +953,18 @@ impl StoredProviderCatalogKey { Ok(self) } + pub fn with_auth_channel_policy_fields( + mut self, + auth_type_by_format: Option, + allow_auth_channel_mismatch_formats: Option, + ) -> Result { + validate_auth_type_by_format(auth_type_by_format.as_ref())?; + validate_auth_channel_mismatch_formats(allow_auth_channel_mismatch_formats.as_ref())?; + self.auth_type_by_format = auth_type_by_format; + self.allow_auth_channel_mismatch_formats = allow_auth_channel_mismatch_formats; + Ok(self) + } + #[allow(clippy::too_many_arguments)] pub fn with_rate_limit_fields( mut self, @@ -646,6 +1021,58 @@ impl StoredProviderCatalogKey { } } +fn validate_auth_type_by_format( + value: Option<&serde_json::Value>, +) -> Result<(), crate::DataLayerError> { + let Some(value) = value else { + return Ok(()); + }; + let Some(entries) = value.as_object() else { + return Err(crate::DataLayerError::UnexpectedValue( + "provider_api_keys.auth_type_by_format must be a JSON object".to_string(), + )); + }; + for (api_format, auth_type) in entries { + let valid_api_format = !api_format.trim().is_empty(); + let valid_auth_type = auth_type.as_str().is_some_and(|value| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "api_key" | "bearer" + ) + }); + if !valid_api_format || !valid_auth_type { + return Err(crate::DataLayerError::UnexpectedValue( + "provider_api_keys.auth_type_by_format contains an invalid entry".to_string(), + )); + } + } + Ok(()) +} + +fn validate_auth_channel_mismatch_formats( + value: Option<&serde_json::Value>, +) -> Result<(), crate::DataLayerError> { + let Some(value) = value else { + return Ok(()); + }; + let Some(items) = value.as_array() else { + return Err(crate::DataLayerError::UnexpectedValue( + "provider_api_keys.allow_auth_channel_mismatch_formats must be a JSON array" + .to_string(), + )); + }; + if items + .iter() + .any(|item| item.as_str().is_none_or(|value| value.trim().is_empty())) + { + return Err(crate::DataLayerError::UnexpectedValue( + "provider_api_keys.allow_auth_channel_mismatch_formats contains an invalid entry" + .to_string(), + )); + } + Ok(()) +} + impl From<&StoredProviderCatalogKey> for ProviderCatalogKeyAdaptiveState { fn from(key: &StoredProviderCatalogKey) -> Self { Self { @@ -664,7 +1091,60 @@ impl From<&StoredProviderCatalogKey> for ProviderCatalogKeyAdaptiveState { #[cfg(test)] mod transport_tests { - use super::StoredProviderCatalogKey; + use super::{ + ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyOAuthCredentialFence, + ProviderCatalogKeyOAuthRuntimeStateCasUpdate, + ProviderCatalogUpstreamMetadataNamespaceExpectation, StoredProviderCatalogEndpoint, + StoredProviderCatalogKey, StoredProviderCatalogProvider, + }; + + fn assert_debug_redacts(value: &T, secrets: &[&str]) { + let debug = format!("{value:?}"); + assert!(debug.contains("[REDACTED]"), "debug output: {debug}"); + for secret in secrets { + assert!( + !debug.contains(secret), + "debug output leaked {secret}: {debug}" + ); + } + } + + fn sample_key() -> StoredProviderCatalogKey { + StoredProviderCatalogKey::new( + "key-policy".to_string(), + "provider-policy".to_string(), + "policy".to_string(), + "api_key".to_string(), + None, + true, + ) + .expect("key should build") + } + + #[test] + fn provider_catalog_key_auth_channel_policy_rejects_malformed_stored_json() { + assert!(sample_key() + .with_auth_channel_policy_fields( + Some(serde_json::json!({"openai:chat": "bearer"})), + Some(serde_json::json!([])), + ) + .is_ok()); + assert!(sample_key() + .with_auth_channel_policy_fields(Some(serde_json::Value::Null), None) + .is_err()); + assert!(sample_key() + .with_auth_channel_policy_fields( + Some(serde_json::json!({"openai:chat": "oauth"})), + None, + ) + .is_err()); + assert!(sample_key() + .with_auth_channel_policy_fields(None, Some(serde_json::Value::Null)) + .is_err()); + assert!(sample_key() + .with_auth_channel_policy_fields(None, Some(serde_json::json!([""]))) + .is_err()); + } #[test] fn provider_catalog_key_defaults_concurrent_limit_to_none() { @@ -707,6 +1187,152 @@ mod transport_tests { assert_eq!(key.rpm_limit, Some(120)); assert_eq!(key.concurrent_limit, Some(3)); } + + #[test] + fn provider_catalog_debug_output_redacts_credentials_and_transport_metadata() { + let mut key = sample_key(); + key.encrypted_api_key = Some("catalog-api-key-ciphertext-canary".to_string()); + key.encrypted_auth_config = Some("catalog-auth-config-ciphertext-canary".to_string()); + key.proxy = Some(serde_json::json!({"password": "catalog-proxy-canary"})); + key.fingerprint = Some(serde_json::json!({"device": "catalog-device-canary"})); + key.upstream_metadata = Some(serde_json::json!({"token": "catalog-metadata-canary"})); + key.oauth_invalid_reason = Some("catalog-oauth-reason-canary".to_string()); + key.status_snapshot = Some(serde_json::json!({"raw": "catalog-status-canary"})); + assert_debug_redacts( + &key, + &[ + "catalog-api-key-ciphertext-canary", + "catalog-auth-config-ciphertext-canary", + "catalog-proxy-canary", + "catalog-device-canary", + "catalog-metadata-canary", + "catalog-oauth-reason-canary", + "catalog-status-canary", + ], + ); + + let provider = StoredProviderCatalogProvider::new( + "provider-debug".to_string(), + "debug".to_string(), + None, + "openai".to_string(), + ) + .expect("provider should build") + .with_transport_fields( + true, + false, + false, + None, + None, + Some(serde_json::json!({"password": "provider-proxy-canary"})), + None, + None, + Some(serde_json::json!({"secret": "provider-config-canary"})), + ); + assert_debug_redacts( + &provider, + &["provider-proxy-canary", "provider-config-canary"], + ); + + let endpoint = StoredProviderCatalogEndpoint::new( + "endpoint-debug".to_string(), + "provider-debug".to_string(), + "openai:chat".to_string(), + None, + None, + true, + ) + .expect("endpoint should build") + .with_transport_fields( + "https://endpoint-user:endpoint-password-canary@example.com/endpoint-token-canary" + .to_string(), + Some(serde_json::json!({"Authorization": "endpoint-header-canary"})), + Some(serde_json::json!({"credential": "endpoint-body-canary"})), + None, + None, + Some(serde_json::json!({"secret": "endpoint-config-canary"})), + None, + Some(serde_json::json!({"password": "endpoint-proxy-canary"})), + ) + .expect("endpoint should accept transport fields"); + assert_debug_redacts( + &endpoint, + &[ + "endpoint-password-canary", + "endpoint-token-canary", + "endpoint-header-canary", + "endpoint-body-canary", + "endpoint-config-canary", + "endpoint-proxy-canary", + ], + ); + } + + #[test] + fn provider_catalog_cas_debug_output_redacts_credential_fences() { + let fence = ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: Some("fence-api-key-canary".to_string()), + auth_type: "oauth".to_string(), + provider_id: "provider-debug".to_string(), + provider_type: "codex".to_string(), + }; + assert_debug_redacts(&fence, &["fence-api-key-canary"]); + + let credentials = ProviderCatalogKeyCredentialsCasUpdate { + key_id: "key-debug".to_string(), + expected_provider_id: "provider-debug".to_string(), + expected_encrypted_api_key: Some("expected-api-key-canary".to_string()), + expected_encrypted_auth_config: Some("expected-auth-config-canary".to_string()), + encrypted_api_key: Some("replacement-api-key-canary".to_string()), + encrypted_auth_config: Some("replacement-auth-config-canary".to_string()), + }; + assert_debug_redacts( + &credentials, + &[ + "expected-api-key-canary", + "expected-auth-config-canary", + "replacement-api-key-canary", + "replacement-auth-config-canary", + ], + ); + + let runtime = ProviderCatalogKeyOAuthRuntimeStateCasUpdate { + key_id: "key-debug".to_string(), + expected_encrypted_auth_config: Some("runtime-expected-auth-canary".to_string()), + expected_credential: Some(fence), + expected_upstream_metadata_namespace: Some( + ProviderCatalogUpstreamMetadataNamespaceExpectation { + namespace: "oauth".to_string(), + expected_value: Some(serde_json::json!({"token": "runtime-fence-canary"})), + }, + ), + encrypted_auth_config: "runtime-auth-config-canary".to_string(), + encrypted_api_key_update: Some("runtime-api-key-canary".to_string()), + expires_at_unix_secs_update: Some(Some(123)), + oauth_invalid_at_unix_secs: Some(124), + oauth_invalid_reason: Some("runtime-provider-error-canary".to_string()), + upstream_metadata_patch: Some(serde_json::json!({ + "refresh_token": "runtime-metadata-canary" + })), + upstream_metadata_namespace_to_remove: None, + status_snapshot_patch: serde_json::json!({"raw": "runtime-status-canary"}), + reset_error_count: true, + updated_at_unix_secs: Some(125), + }; + assert_debug_redacts( + &runtime, + &[ + "runtime-expected-auth-canary", + "fence-api-key-canary", + "runtime-fence-canary", + "runtime-auth-config-canary", + "runtime-api-key-canary", + "runtime-provider-error-canary", + "runtime-metadata-canary", + "runtime-status-canary", + ], + ); + } } #[derive(Debug, Clone, PartialEq, Eq, Default)] @@ -848,6 +1474,26 @@ pub trait ProviderCatalogWriteRepository: Send + Sync { provider: &StoredProviderCatalogProvider, ) -> Result; + async fn compare_and_swap_provider_config( + &self, + _update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "provider catalog config compare-and-swap is not supported by this repository" + .to_string(), + )) + } + + async fn compare_and_swap_provider_proxy( + &self, + _update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "provider catalog provider proxy compare-and-swap is not supported by this repository" + .to_string(), + )) + } + async fn delete_provider(&self, provider_id: &str) -> Result; async fn cleanup_deleted_provider_refs( @@ -868,6 +1514,16 @@ pub trait ProviderCatalogWriteRepository: Send + Sync { endpoint: &StoredProviderCatalogEndpoint, ) -> Result; + async fn compare_and_swap_endpoint_proxy( + &self, + _update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "provider catalog endpoint proxy compare-and-swap is not supported by this repository" + .to_string(), + )) + } + async fn delete_endpoint(&self, endpoint_id: &str) -> Result; async fn create_key( @@ -880,6 +1536,26 @@ pub trait ProviderCatalogWriteRepository: Send + Sync { key: &StoredProviderCatalogKey, ) -> Result; + async fn compare_and_swap_key_proxy( + &self, + _update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "provider catalog key proxy compare-and-swap is not supported by this repository" + .to_string(), + )) + } + + async fn compare_and_swap_key_credentials( + &self, + _update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + Err(crate::DataLayerError::InvalidConfiguration( + "provider catalog key credential compare-and-swap is not supported by this repository" + .to_string(), + )) + } + /// Compare-and-swap administrator-owned key configuration. Credential /// rotation, Codex namespace replacement, and quota invalidation must be /// committed atomically with the configuration update. @@ -947,20 +1623,11 @@ pub trait ProviderCatalogWriteRepository: Send + Sync { key_id: &str, ) -> Result; - async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result; - async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result; diff --git a/crates/aether-data/contracts/src/repository/proxy_nodes.rs b/crates/aether-data/contracts/src/repository/proxy_nodes.rs index ec924a707..d22059cca 100644 --- a/crates/aether-data/contracts/src/repository/proxy_nodes.rs +++ b/crates/aether-data/contracts/src/repository/proxy_nodes.rs @@ -1,9 +1,15 @@ +use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED; use async_trait::async_trait; use serde_json::Value; -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +const PROXY_NODE_BOUND_TUNNEL_SECRET_PREFIX: &str = + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:"; + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredProxyNode { pub id: String, + #[serde(default = "new_proxy_node_tunnel_generation")] + pub tunnel_generation: String, pub name: String, pub ip: String, pub port: i32, @@ -34,6 +40,42 @@ pub struct StoredProxyNode { pub updated_at_unix_secs: Option, } +impl std::fmt::Debug for StoredProxyNode { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredProxyNode") + .field("id", &self.id) + .field("tunnel_generation", &self.tunnel_generation) + .field("name", &self.name) + .field("ip", &self.ip) + .field("port", &self.port) + .field("region", &self.region) + .field("is_manual", &self.is_manual) + .field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]")) + .field("proxy_username", &self.proxy_username) + .field( + "proxy_password", + &self.proxy_password.as_ref().map(|_| "[REDACTED]"), + ) + .field("status", &self.status) + .field("tunnel_mode", &self.tunnel_mode) + .field("tunnel_connected", &self.tunnel_connected) + .field( + "proxy_metadata", + &self.proxy_metadata.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "hardware_info", + &self.hardware_info.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "remote_config", + &self.remote_config.as_ref().map(|_| "[REDACTED]"), + ) + .finish_non_exhaustive() + } +} + impl StoredProxyNode { #[allow(clippy::too_many_arguments)] pub fn new( @@ -76,6 +118,7 @@ impl StoredProxyNode { Ok(Self { id, + tunnel_generation: new_proxy_node_tunnel_generation(), name, ip, port, @@ -147,11 +190,22 @@ impl StoredProxyNode { self.proxy_password = proxy_password; self } + + pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self { + self.tunnel_generation = tunnel_generation; + self + } +} + +pub fn new_proxy_node_tunnel_generation() -> String { + uuid::Uuid::new_v4().to_string() } #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeHeartbeatMutation { pub node_id: String, + #[serde(default)] + pub expected_tunnel_generation: Option, pub heartbeat_interval: Option, pub active_connections: Option, pub total_requests_delta: Option, @@ -166,6 +220,11 @@ pub struct ProxyNodeHeartbeatMutation { #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeTrafficMutation { pub node_id: String, + /// Incarnation fence captured when the request plan selected this node. + /// Missing fences are rejected by the gateway path so a stale plan cannot + /// update a node recreated under the same id. + #[serde(default)] + pub expected_tunnel_generation: Option, pub total_requests_delta: i64, pub failed_requests_delta: i64, pub dns_failures_delta: i64, @@ -174,6 +233,11 @@ pub struct ProxyNodeTrafficMutation { #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeRegistrationMutation { + /// The stable identity selected by the caller before any secret is + /// protected. Re-registration of an existing endpoint must use its + /// existing id; repositories reject attempts to replace it. + #[serde(default)] + pub node_id: Option, pub name: String, pub ip: String, pub port: i32, @@ -190,8 +254,13 @@ pub struct ProxyNodeRegistrationMutation { pub tunnel_mode: bool, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeManualCreateMutation { + /// Optional caller-selected id used to bind credentials before the row is + /// inserted. Repositories generate one only for legacy callers that do + /// not provide it. + #[serde(default)] + pub node_id: Option, pub name: String, pub ip: String, pub port: i32, @@ -202,7 +271,27 @@ pub struct ProxyNodeManualCreateMutation { pub registered_by: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for ProxyNodeManualCreateMutation { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProxyNodeManualCreateMutation") + .field("node_id", &self.node_id) + .field("name", &self.name) + .field("ip", &self.ip) + .field("port", &self.port) + .field("region", &self.region) + .field("proxy_url", &"[REDACTED]") + .field("proxy_username", &self.proxy_username) + .field( + "proxy_password", + &self.proxy_password.as_ref().map(|_| "[REDACTED]"), + ) + .field("registered_by", &self.registered_by) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeManualUpdateMutation { pub node_id: String, pub name: Option, @@ -214,9 +303,30 @@ pub struct ProxyNodeManualUpdateMutation { pub proxy_password: Option, } +impl std::fmt::Debug for ProxyNodeManualUpdateMutation { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProxyNodeManualUpdateMutation") + .field("node_id", &self.node_id) + .field("name", &self.name) + .field("ip", &self.ip) + .field("port", &self.port) + .field("region", &self.region) + .field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]")) + .field("proxy_username", &self.proxy_username) + .field( + "proxy_password", + &self.proxy_password.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeTunnelStatusMutation { pub node_id: String, + #[serde(default)] + pub expected_tunnel_generation: Option, pub connected: bool, pub conn_count: i32, pub detail: Option, @@ -226,6 +336,8 @@ pub struct ProxyNodeTunnelStatusMutation { #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeRemoteConfigMutation { pub node_id: String, + #[serde(default)] + pub expected_tunnel_generation: Option, pub node_name: Option, pub allowed_ports: Option>, pub log_level: Option, @@ -463,6 +575,31 @@ pub fn normalize_proxy_metadata( } } +pub fn normalize_heartbeat_proxy_metadata( + previous_proxy_metadata: Option<&Value>, + proxy_metadata: Option<&Value>, + proxy_version: Option<&str>, +) -> Option { + let Some(Value::Object(mut normalized)) = + normalize_proxy_metadata(proxy_metadata, proxy_version) + else { + return None; + }; + + // Tunnel security is control-plane state. A heartbeat may refresh runtime + // metadata, but it must never introduce or replace this trusted field. + normalized.remove("tunnel_security"); + let merged = preserve_proxy_metadata_tunnel_security( + previous_proxy_metadata, + Some(Value::Object(normalized)), + ); + merged.filter(|value| { + value + .as_object() + .is_some_and(|metadata| !metadata.is_empty()) + }) +} + pub fn preserve_proxy_metadata_tunnel_security( previous_proxy_metadata: Option<&Value>, next_proxy_metadata: Option, @@ -477,9 +614,7 @@ pub fn preserve_proxy_metadata_tunnel_security( match next_proxy_metadata { Some(Value::Object(mut metadata)) => { - metadata - .entry("tunnel_security".to_string()) - .or_insert(tunnel_security); + metadata.insert("tunnel_security".to_string(), tunnel_security); Some(Value::Object(metadata)) } Some(value) => Some(value), @@ -491,6 +626,70 @@ pub fn preserve_proxy_metadata_tunnel_security( } } +/// Merge metadata received during a trusted registration/re-registration. +/// +/// Registration is the control-plane path that may rotate a tunnel PSK. A +/// registration payload that omits `tunnel_security` is therefore a partial +/// metadata refresh and must not clear the previously trusted security state. +/// Only a non-empty, gateway-bound v2 ciphertext proves that the registration +/// passed through the gateway credential-binding path, and that ciphertext is +/// accepted only with the required non-TLS security mode. Mode-only, +/// plaintext, malformed, null, scalar, empty, and disabled security values are +/// treated as omission and cannot clear a previously trusted object. +pub fn merge_proxy_metadata_for_registration( + previous_proxy_metadata: Option<&Value>, + next_proxy_metadata: Option, +) -> Option { + let Some(next_proxy_metadata) = next_proxy_metadata else { + return previous_proxy_metadata.cloned(); + }; + let incoming_security_is_explicit = + proxy_metadata_has_explicit_tunnel_security(Some(&next_proxy_metadata)); + + let Value::Object(mut metadata) = next_proxy_metadata else { + // `normalize_proxy_metadata` normally prevents this branch. Keep a + // malformed replacement from erasing trusted control-plane state. + return previous_proxy_metadata.cloned(); + }; + + if incoming_security_is_explicit { + return Some(Value::Object(metadata)); + } + + // Null, scalar, and empty security objects are not valid replacements. + // Remove them before restoring the previous trusted object so malformed + // input cannot mask or downgrade the registered security policy. + metadata.remove("tunnel_security"); + if let Some(previous_security) = previous_proxy_metadata + .and_then(|value| value.get("tunnel_security")) + .filter(|value| value.is_object()) + .cloned() + { + metadata.insert("tunnel_security".to_string(), previous_security); + } + + (!metadata.is_empty()).then_some(Value::Object(metadata)) +} + +/// Return whether metadata contains a complete gateway-bound tunnel security +/// replacement that a trusted registration may persist. +pub fn proxy_metadata_has_explicit_tunnel_security(proxy_metadata: Option<&Value>) -> bool { + proxy_metadata + .and_then(Value::as_object) + .and_then(|metadata| metadata.get("tunnel_security")) + .and_then(Value::as_object) + .is_some_and(|security| { + security.get("mode").and_then(Value::as_str) == Some(TUNNEL_SECURITY_NON_TLS_REQUIRED) + && security + .get("encryption_key_encrypted") + .and_then(Value::as_str) + .and_then(|encrypted| { + encrypted.strip_prefix(PROXY_NODE_BOUND_TUNNEL_SECRET_PREFIX) + }) + .is_some_and(|ciphertext| !ciphertext.is_empty()) + }) +} + fn extract_tunnel_metrics_counters( proxy_metadata: Option<&Value>, ) -> Option { @@ -712,6 +911,20 @@ pub trait ProxyNodeReadRepository: Send + Sync { pub trait ProxyNodeWriteRepository: Send + Sync { async fn reset_stale_tunnel_statuses(&self) -> Result; + async fn compare_and_set_proxy_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result; + + async fn compare_and_set_proxy_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result; + async fn create_manual_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -779,12 +992,77 @@ mod tests { use super::{ bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, + merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata, normalize_proxy_node_scheduling_state, preserve_proxy_metadata_tunnel_security, proxy_node_accepts_new_tunnels, proxy_reported_version, reconcile_remote_config_after_heartbeat, remote_config_scheduling_state, - remote_config_upgrade_target, ProxyNodeMetricsStep, StoredProxyNode, + remote_config_upgrade_target, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, + ProxyNodeMetricsStep, StoredProxyNode, }; + #[test] + fn proxy_node_debug_output_redacts_credentials_and_untrusted_metadata() { + let password = "debug-secret-proxy-password"; + let proxy_url = "https://user:debug-secret-url@example.com"; + let metadata_secret = "debug-secret-proxy-metadata"; + let mut stored = StoredProxyNode::new( + "node-1".to_string(), + "node".to_string(), + "127.0.0.1".to_string(), + 8080, + true, + "online".to_string(), + 30, + 0, + 0, + 0, + 0, + 0, + false, + false, + 1, + ) + .expect("proxy node should build") + .with_manual_proxy_fields( + Some(proxy_url.to_string()), + Some("user".to_string()), + Some(password.to_string()), + ); + stored.proxy_metadata = Some(json!({"secret": metadata_secret})); + let create = ProxyNodeManualCreateMutation { + node_id: Some("node-1".to_string()), + name: "node".to_string(), + ip: "127.0.0.1".to_string(), + port: 8080, + region: None, + proxy_url: proxy_url.to_string(), + proxy_username: Some("user".to_string()), + proxy_password: Some(password.to_string()), + registered_by: None, + }; + let update = ProxyNodeManualUpdateMutation { + node_id: "node-1".to_string(), + name: None, + ip: None, + port: None, + region: None, + proxy_url: Some(proxy_url.to_string()), + proxy_username: Some("user".to_string()), + proxy_password: Some(password.to_string()), + }; + + for rendered in [ + format!("{stored:?}"), + format!("{create:?}"), + format!("{update:?}"), + ] { + for secret in [password, proxy_url, metadata_secret] { + assert!(!rendered.contains(secret), "Debug output leaked {secret}"); + } + assert!(rendered.contains("[REDACTED]")); + } + } + #[test] fn normalizes_reported_versions_and_clears_completed_upgrade_targets() { let remote_config = json!({ @@ -985,7 +1263,11 @@ mod tests { }); let next = json!({ "version": "1.0.1", - "tunnel_metrics": {"connect_successes": 1} + "tunnel_metrics": {"connect_successes": 1}, + "tunnel_security": { + "mode": "disabled", + "encryption_key": "attacker-controlled" + } }); let merged = preserve_proxy_metadata_tunnel_security(Some(&previous), Some(next)) @@ -1008,6 +1290,202 @@ mod tests { ); } + #[test] + fn registration_metadata_preserves_omitted_tunnel_security() { + let previous = json!({ + "version": "1.0.0", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + }); + let next = json!({ + "version": "1.1.0", + "tunnel_metrics": {"connect_successes": 2} + }); + + let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next)) + .expect("registration metadata should remain present"); + assert_eq!(merged.get("version"), Some(&json!("1.1.0"))); + assert_eq!( + merged.pointer("/tunnel_security/encryption_key_encrypted"), + Some(&json!( + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + )) + ); + assert_eq!( + merged.pointer("/tunnel_metrics/connect_successes"), + Some(&json!(2)) + ); + } + + #[test] + fn registration_metadata_accepts_explicit_tunnel_security_rotation() { + let previous = json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + }); + let next = json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new" + } + }); + + let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next)) + .expect("rotated registration metadata should remain present"); + assert_eq!( + merged.pointer("/tunnel_security/encryption_key_encrypted"), + Some(&json!( + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new" + )) + ); + } + + #[test] + fn registration_metadata_rejects_invalid_security_replacement() { + let previous = json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + }); + let next = json!({ + "version": "1.2.0", + "tunnel_security": null + }); + + let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next)) + .expect("previous security should be retained"); + assert_eq!(merged.get("version"), Some(&json!("1.2.0"))); + assert_eq!( + merged.pointer("/tunnel_security/encryption_key_encrypted"), + Some(&json!( + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + )) + ); + } + + #[test] + fn registration_metadata_rejects_mode_only_security_downgrade() { + let previous = json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + }); + let next = json!({ + "version": "1.3.0", + "tunnel_security": {"mode": "disabled"} + }); + + let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next)) + .expect("mode-only security must not replace the registered credential"); + assert_eq!(merged.get("version"), Some(&json!("1.3.0"))); + assert_eq!( + merged.pointer("/tunnel_security/encryption_key_encrypted"), + Some(&json!( + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + )) + ); + assert_eq!( + merged.pointer("/tunnel_security/mode"), + Some(&json!("non_tls_required")) + ); + } + + #[test] + fn registration_metadata_rejects_bound_ciphertext_with_disabled_mode() { + let previous = json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + }); + let next = json!({ + "version": "1.3.1", + "tunnel_security": { + "mode": "disabled", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-attacker" + } + }); + + let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next)) + .expect("disabled security must not replace the registered credential"); + assert_eq!(merged.get("version"), Some(&json!("1.3.1"))); + assert_eq!( + merged.pointer("/tunnel_security/encryption_key_encrypted"), + Some(&json!( + "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + )) + ); + assert_eq!( + merged.pointer("/tunnel_security/mode"), + Some(&json!("non_tls_required")) + ); + } + + #[test] + fn new_registration_drops_invalid_tunnel_security_fields() { + for invalid_security in [ + json!(null), + json!("disabled"), + json!({}), + json!({"mode": "disabled"}), + json!({"mode": "disabled", "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-attacker"}), + json!({"mode": "non_tls_required", "encryption_key_encrypted": ""}), + json!({"mode": "non_tls_required", "encryption_key_encrypted": "not-gateway-bound"}), + ] { + let next = json!({ + "version": "1.4.0", + "tunnel_security": invalid_security + }); + let merged = merge_proxy_metadata_for_registration(None, Some(next)) + .expect("valid non-security metadata should remain"); + assert_eq!(merged.get("version"), Some(&json!("1.4.0"))); + assert!(merged.get("tunnel_security").is_none()); + } + + assert_eq!( + merge_proxy_metadata_for_registration( + None, + Some(json!({"tunnel_security": {"mode": "disabled"}})), + ), + None + ); + } + + #[test] + fn heartbeat_metadata_cannot_introduce_tunnel_security() { + let injected = json!({ + "tunnel_security": { + "mode": "disabled", + "encryption_key": "attacker-controlled" + } + }); + assert_eq!( + normalize_heartbeat_proxy_metadata(None, Some(&injected), None), + None + ); + + let injected_with_runtime_metadata = json!({ + "version": "1.2.3", + "arch": "arm64", + "tunnel_security": { + "mode": "disabled", + "encryption_key": "attacker-controlled" + } + }); + let normalized = + normalize_heartbeat_proxy_metadata(None, Some(&injected_with_runtime_metadata), None) + .expect("runtime metadata should remain present"); + assert_eq!(normalized.get("version"), Some(&json!("1.2.3"))); + assert_eq!(normalized.get("arch"), Some(&json!("arm64"))); + assert!(normalized.get("tunnel_security").is_none()); + } + #[test] fn maps_timestamps_to_metric_buckets() { assert_eq!( diff --git a/crates/aether-data/contracts/src/repository/routing_profiles/types.rs b/crates/aether-data/contracts/src/repository/routing_profiles/types.rs index 09c0620d3..d949d7836 100644 --- a/crates/aether-data/contracts/src/repository/routing_profiles/types.rs +++ b/crates/aether-data/contracts/src/repository/routing_profiles/types.rs @@ -10,6 +10,9 @@ pub struct StoredRoutingGroup { pub description: Option, pub enabled: bool, pub is_system_default: bool, + /// Stable administrator-defined display order. This is intentionally not + /// consulted by request routing or candidate selection. + pub sort_order: i64, pub config_json: Value, pub version: i64, pub created_at: i64, @@ -28,6 +31,7 @@ impl StoredRoutingGroup { description: record.description, enabled: record.enabled, is_system_default: record.is_system_default, + sort_order: record.sort_order.max(0), config_json: record.config_json, version: record.version.max(1), created_at: record.created_at, @@ -44,6 +48,7 @@ pub struct CreateRoutingGroupRecord { pub description: Option, pub enabled: bool, pub is_system_default: bool, + pub sort_order: i64, pub config_json: Value, pub version: i64, pub created_at: i64, @@ -57,6 +62,7 @@ pub struct UpdateRoutingGroupRecord { pub description: Option>, pub enabled: Option, pub is_system_default: Option, + pub sort_order: Option, pub config_json: Option, pub version: Option, pub updated_at: i64, @@ -255,6 +261,9 @@ pub fn apply_group_patch( if let Some(is_system_default) = patch.is_system_default { group.is_system_default = is_system_default; } + if let Some(sort_order) = patch.sort_order { + group.sort_order = sort_order.max(0); + } if let Some(config_json) = patch.config_json { if !config_json.is_object() { return Err(crate::DataLayerError::InvalidInput( diff --git a/crates/aether-data/contracts/src/repository/settlement/mod.rs b/crates/aether-data/contracts/src/repository/settlement/mod.rs index 5fd346872..420d56065 100644 --- a/crates/aether-data/contracts/src/repository/settlement/mod.rs +++ b/crates/aether-data/contracts/src/repository/settlement/mod.rs @@ -2,6 +2,11 @@ mod types; pub use types::{ finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd, - settlement_billing_status_for_usage_status, SettlementRepository, SettlementWriteRepository, - StoredUsageSettlement, UsageSettlementInput, WalletDebitPlan, SETTLEMENT_EPSILON_USD, + settlement_billing_status_for_usage_status, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementRepository, SettlementWriteRepository, + StoredUsagePolicyCostReservation, StoredUsagePolicyRequestAdmission, StoredUsageSettlement, + UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestAdmissionState, + UsagePolicyRequestWindow, UsageSettlementInput, WalletDebitPlan, SETTLEMENT_EPSILON_USD, }; diff --git a/crates/aether-data/contracts/src/repository/settlement/types.rs b/crates/aether-data/contracts/src/repository/settlement/types.rs index cd5521e81..4d9c6a0e6 100644 --- a/crates/aether-data/contracts/src/repository/settlement/types.rs +++ b/crates/aether-data/contracts/src/repository/settlement/types.rs @@ -1,5 +1,339 @@ use async_trait::async_trait; +use crate::repository::billing::MAX_USAGE_POLICY_TOTAL_RULES; + +const MAX_USAGE_POLICY_LEDGER_ID_BYTES: usize = 128; + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct UsagePolicyRequestWindow { + pub starts_at_unix_secs: u64, + pub ends_at_unix_secs: u64, + pub limit_requests: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct ReserveUsagePolicyRequestInput { + pub request_id: String, + pub subject_id: String, + pub event_token: String, + pub admitted_at_unix_secs: u64, + /// Exclusive timestamp after which this admission cannot affect any supplied window and its + /// idempotency tombstone may be deleted safely. + pub retain_until_unix_secs: u64, + pub windows: Vec, +} + +impl ReserveUsagePolicyRequestInput { + pub fn validate(&self) -> Result<(), crate::DataLayerError> { + validate_bounded_id(&self.request_id, "usage policy request_id")?; + validate_bounded_id(&self.subject_id, "usage policy subject_id")?; + validate_bounded_id(&self.event_token, "usage policy event_token")?; + validate_database_u64(self.admitted_at_unix_secs, "usage policy admitted_at")?; + validate_database_u64(self.retain_until_unix_secs, "usage policy retain_until")?; + if self.windows.is_empty() || self.windows.len() > MAX_USAGE_POLICY_TOTAL_RULES { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy request admission requires 1 to {MAX_USAGE_POLICY_TOTAL_RULES} windows" + ))); + } + for (index, window) in self.windows.iter().enumerate() { + validate_database_u64( + window.starts_at_unix_secs, + &format!("usage policy request window {index} start"), + )?; + validate_database_u64( + window.ends_at_unix_secs, + &format!("usage policy request window {index} end"), + )?; + validate_database_u64( + window.limit_requests, + &format!("usage policy request window {index} limit"), + )?; + if window.limit_requests == 0 + || window.starts_at_unix_secs >= window.ends_at_unix_secs + || self.admitted_at_unix_secs < window.starts_at_unix_secs + || self.admitted_at_unix_secs >= window.ends_at_unix_secs + { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy request window {index} does not contain admission or has invalid bounds" + ))); + } + if self.retain_until_unix_secs < window.ends_at_unix_secs { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy request retain_until precedes window {index} end" + ))); + } + } + Ok(()) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum UsagePolicyRequestAdmissionState { + Active, + Released, +} + +impl UsagePolicyRequestAdmissionState { + pub const fn as_str(self) -> &'static str { + match self { + Self::Active => "active", + Self::Released => "released", + } + } + + pub fn parse(value: &str) -> Option { + match value { + "active" => Some(Self::Active), + "released" => Some(Self::Released), + _ => None, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(tag = "status", rename_all = "snake_case")] +pub enum ReserveUsagePolicyRequestOutcome { + Allowed, + Rejected { + window_index: usize, + limit_requests: u64, + used_requests: u64, + }, + AlreadyReleased, + Conflict, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct ReleaseUsagePolicyRequestAdmissionInput { + pub request_id: String, + pub subject_id: String, + pub event_token: String, + pub released_at_unix_secs: u64, +} + +impl ReleaseUsagePolicyRequestAdmissionInput { + pub fn validate(&self) -> Result<(), crate::DataLayerError> { + validate_bounded_id(&self.request_id, "usage policy request_id")?; + validate_bounded_id(&self.subject_id, "usage policy subject_id")?; + validate_bounded_id(&self.event_token, "usage policy event_token")?; + validate_database_u64(self.released_at_unix_secs, "usage policy released_at") + } +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct StoredUsagePolicyRequestAdmission { + pub request_id: String, + pub subject_id: String, + pub event_token: String, + pub admitted_at_unix_secs: u64, + pub retain_until_unix_secs: u64, + pub state: UsagePolicyRequestAdmissionState, + pub released_at_unix_secs: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct UsagePolicyCostWindow { + pub window_id: String, + pub starts_at_unix_secs: u64, + pub ends_at_unix_secs: u64, + pub limit_cost_units: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct ReserveUsagePolicyCostInput { + pub request_id: String, + pub subject_id: String, + pub reservation_token: String, + pub admitted_at_unix_secs: u64, + pub reserved_cost_units: u64, + pub reservation_expires_at_unix_secs: u64, + /// Exclusive timestamp after which this reservation can no longer affect any future window + /// and its idempotency tombstone may be deleted safely. + pub retain_until_unix_secs: u64, + pub windows: Vec, +} + +impl ReserveUsagePolicyCostInput { + pub fn validate(&self) -> Result<(), crate::DataLayerError> { + validate_bounded_id(&self.request_id, "usage policy request_id")?; + validate_bounded_id(&self.subject_id, "usage policy subject_id")?; + validate_bounded_id(&self.reservation_token, "usage policy reservation_token")?; + validate_cost_units(self.reserved_cost_units, "reserved_cost_units")?; + if self.reservation_expires_at_unix_secs <= self.admitted_at_unix_secs { + return Err(crate::DataLayerError::InvalidInput( + "usage policy reservation must expire after admission".to_string(), + )); + } + if self.retain_until_unix_secs < self.reservation_expires_at_unix_secs { + return Err(crate::DataLayerError::InvalidInput( + "usage policy retain_until must not precede reservation expiry".to_string(), + )); + } + if self.windows.is_empty() || self.windows.len() > MAX_USAGE_POLICY_TOTAL_RULES { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy reservation requires 1 to {MAX_USAGE_POLICY_TOTAL_RULES} windows" + ))); + } + for (index, window) in self.windows.iter().enumerate() { + validate_non_empty_id( + &window.window_id, + &format!("usage policy window {index} id"), + )?; + validate_cost_units( + window.limit_cost_units, + &format!("usage policy window {index} limit_cost_units"), + )?; + if window.limit_cost_units == 0 + || window.starts_at_unix_secs >= window.ends_at_unix_secs + || self.admitted_at_unix_secs < window.starts_at_unix_secs + || self.admitted_at_unix_secs >= window.ends_at_unix_secs + { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy window {index} does not contain admission or has invalid bounds" + ))); + } + if self.windows[..index] + .iter() + .any(|previous| previous.window_id == window.window_id) + { + return Err(crate::DataLayerError::InvalidInput(format!( + "usage policy window {index} duplicates window_id {}", + window.window_id + ))); + } + } + Ok(()) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum UsagePolicyCostReservationState { + Reserved, + Finalized, + Released, +} + +impl UsagePolicyCostReservationState { + pub const fn as_str(self) -> &'static str { + match self { + Self::Reserved => "reserved", + Self::Finalized => "finalized", + Self::Released => "released", + } + } + + pub fn parse(value: &str) -> Option { + match value { + "reserved" => Some(Self::Reserved), + "finalized" => Some(Self::Finalized), + "released" => Some(Self::Released), + _ => None, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(tag = "status", rename_all = "snake_case")] +pub enum ReserveUsagePolicyCostOutcome { + Allowed { + reserved_cost_units: u64, + additional_reserved_cost_units: u64, + }, + Rejected { + window_index: usize, + limit_cost_units: u64, + used_cost_units: u64, + }, + AlreadyTerminal { + state: UsagePolicyCostReservationState, + }, + Conflict, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct ReconcileUsagePolicyCostInput { + pub request_id: String, + pub subject_id: String, + pub reservation_token: String, + pub actual_cost_units: u64, + pub terminal_state: UsagePolicyCostReservationState, + pub finalized_at_unix_secs: u64, +} + +impl ReconcileUsagePolicyCostInput { + pub fn validate(&self) -> Result<(), crate::DataLayerError> { + validate_bounded_id(&self.request_id, "usage policy request_id")?; + validate_bounded_id(&self.subject_id, "usage policy subject_id")?; + validate_bounded_id(&self.reservation_token, "usage policy reservation_token")?; + validate_cost_units(self.actual_cost_units, "actual_cost_units")?; + match self.terminal_state { + UsagePolicyCostReservationState::Reserved => Err(crate::DataLayerError::InvalidInput( + "usage policy reconciliation requires a terminal state".to_string(), + )), + UsagePolicyCostReservationState::Released if self.actual_cost_units != 0 => { + Err(crate::DataLayerError::InvalidInput( + "released usage policy reservations must have zero actual cost".to_string(), + )) + } + UsagePolicyCostReservationState::Finalized + | UsagePolicyCostReservationState::Released => Ok(()), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct StoredUsagePolicyCostReservation { + pub request_id: String, + pub subject_id: String, + pub reservation_token: String, + pub admitted_at_unix_secs: u64, + pub reserved_cost_units: u64, + pub actual_cost_units: Option, + pub state: UsagePolicyCostReservationState, + pub reservation_expires_at_unix_secs: u64, + pub retain_until_unix_secs: u64, + pub finalized_at_unix_secs: Option, +} + +fn validate_non_empty_id(value: &str, field: &str) -> Result<(), crate::DataLayerError> { + if value.trim().is_empty() { + return Err(crate::DataLayerError::InvalidInput(format!( + "{field} must not be empty" + ))); + } + Ok(()) +} + +fn validate_bounded_id(value: &str, field: &str) -> Result<(), crate::DataLayerError> { + validate_non_empty_id(value, field)?; + if value.len() > MAX_USAGE_POLICY_LEDGER_ID_BYTES { + return Err(crate::DataLayerError::InvalidInput(format!( + "{field} exceeds {MAX_USAGE_POLICY_LEDGER_ID_BYTES} bytes" + ))); + } + Ok(()) +} + +fn validate_database_u64(value: u64, field: &str) -> Result<(), crate::DataLayerError> { + if value > i64::MAX as u64 { + return Err(crate::DataLayerError::InvalidInput(format!( + "{field} exceeds the database integer range" + ))); + } + Ok(()) +} + +fn validate_cost_units(value: u64, field: &str) -> Result<(), crate::DataLayerError> { + if value > i64::MAX as u64 { + return Err(crate::DataLayerError::InvalidInput(format!( + "{field} exceeds the database integer range" + ))); + } + Ok(()) +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct UsageSettlementInput { pub request_id: String, @@ -53,6 +387,44 @@ pub struct StoredUsageSettlement { #[async_trait] pub trait SettlementWriteRepository: Send + Sync { + async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result; + + async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, crate::DataLayerError>; + + async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + let _ = (now_unix_secs, batch_size); + Ok(0) + } + + async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result; + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, crate::DataLayerError>; + + async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + let _ = (now_unix_secs, batch_size); + Ok(0) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -85,6 +457,35 @@ pub fn finite_wallet_available_usd(recharge_balance: f64, gift_balance: f64) -> recharge_balance.max(0.0) + gift_balance.max(0.0) } +/// Reject corrupted or overflowing financial values before usage settlement mutates a wallet. +/// Recharge balances may legitimately be negative after an admitted request settles, but gift +/// balances and cumulative consumption may not be negative. Derived totals must remain finite so +/// persisted `NaN`/infinity values cannot turn a finite wallet into an implicit unlimited wallet. +pub fn validate_wallet_settlement_values( + recharge_balance: f64, + gift_balance: f64, + total_consumed: f64, + additional_consumed: f64, +) -> Result<(), crate::DataLayerError> { + let balance_total = recharge_balance + gift_balance; + let consumed_after = total_consumed + additional_consumed; + if !recharge_balance.is_finite() + || !gift_balance.is_finite() + || gift_balance < 0.0 + || !balance_total.is_finite() + || !total_consumed.is_finite() + || total_consumed < 0.0 + || !additional_consumed.is_finite() + || additional_consumed < 0.0 + || !consumed_after.is_finite() + { + return Err(crate::DataLayerError::UnexpectedValue( + "wallet financial state is invalid for usage settlement".to_string(), + )); + } + Ok(()) +} + pub fn plan_finite_wallet_debit( recharge_balance: f64, gift_balance: f64, @@ -115,7 +516,12 @@ pub fn settlement_billable_cost_usd(input: &UsageSettlementInput) -> f64 { #[cfg(test)] mod tests { - use super::UsageSettlementInput; + use super::{ + validate_wallet_settlement_values, ReconcileUsagePolicyCostInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput, + UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestWindow, + UsageSettlementInput, + }; #[test] fn rejects_invalid_settlement_input() { @@ -133,4 +539,101 @@ mod tests { }; assert!(input.validate().is_err()); } + + #[test] + fn wallet_settlement_values_reject_corruption_and_overflow() { + assert!(validate_wallet_settlement_values(-3.0, 0.0, 12.0, 1.0).is_ok()); + for (recharge, gift, consumed, additional) in [ + (f64::NAN, 0.0, 0.0, 1.0), + (f64::INFINITY, 0.0, 0.0, 1.0), + (0.0, f64::NAN, 0.0, 1.0), + (0.0, -0.01, 0.0, 1.0), + (0.0, 0.0, f64::INFINITY, 1.0), + (0.0, 0.0, -0.01, 1.0), + (0.0, 0.0, 0.0, f64::NAN), + (f64::MAX, f64::MAX, 0.0, 1.0), + (0.0, 0.0, f64::MAX, f64::MAX), + ] { + assert!( + validate_wallet_settlement_values(recharge, gift, consumed, additional).is_err() + ); + } + } + + #[test] + fn validates_request_admission_window_and_retention_bounds() { + let valid = ReserveUsagePolicyRequestInput { + request_id: "request-1".to_string(), + subject_id: "user-1".to_string(), + event_token: "event-1".to_string(), + admitted_at_unix_secs: 100, + retain_until_unix_secs: 200, + windows: vec![UsagePolicyRequestWindow { + starts_at_unix_secs: 50, + ends_at_unix_secs: 200, + limit_requests: 10, + }], + }; + assert!(valid.validate().is_ok()); + + let mut invalid_retention = valid.clone(); + invalid_retention.retain_until_unix_secs = 199; + assert!(invalid_retention.validate().is_err()); + + let mut invalid_window = valid; + invalid_window.windows[0].starts_at_unix_secs = 101; + assert!(invalid_window.validate().is_err()); + } + + #[test] + fn usage_policy_cost_ids_match_sql_column_bounds() { + let bounded = "x".repeat(128); + let reserve = ReserveUsagePolicyCostInput { + request_id: bounded.clone(), + subject_id: bounded.clone(), + reservation_token: bounded.clone(), + admitted_at_unix_secs: 100, + reserved_cost_units: 1, + reservation_expires_at_unix_secs: 150, + retain_until_unix_secs: 200, + windows: vec![UsagePolicyCostWindow { + window_id: "window-1".to_string(), + starts_at_unix_secs: 50, + ends_at_unix_secs: 200, + limit_cost_units: 10, + }], + }; + assert!(reserve.validate().is_ok()); + + for field in ["request_id", "subject_id", "reservation_token"] { + let mut too_long = reserve.clone(); + match field { + "request_id" => too_long.request_id.push('x'), + "subject_id" => too_long.subject_id.push('x'), + "reservation_token" => too_long.reservation_token.push('x'), + _ => unreachable!(), + } + assert!(too_long.validate().is_err(), "{field} must be bounded"); + } + + let reconcile = ReconcileUsagePolicyCostInput { + request_id: bounded.clone(), + subject_id: bounded.clone(), + reservation_token: bounded, + actual_cost_units: 1, + terminal_state: UsagePolicyCostReservationState::Finalized, + finalized_at_unix_secs: 200, + }; + assert!(reconcile.validate().is_ok()); + for field in ["request_id", "subject_id", "reservation_token"] { + let mut too_long = reconcile.clone(); + match field { + "request_id" => too_long.request_id.push('x'), + "subject_id" => too_long.subject_id.push('x'), + "reservation_token" => too_long.reservation_token.push('x'), + _ => unreachable!(), + } + assert!(too_long.validate().is_err(), "{field} must be bounded"); + } + } } diff --git a/crates/aether-data/contracts/src/repository/usage/compression.rs b/crates/aether-data/contracts/src/repository/usage/compression.rs new file mode 100644 index 000000000..855af1a20 --- /dev/null +++ b/crates/aether-data/contracts/src/repository/usage/compression.rs @@ -0,0 +1,58 @@ +use std::io::Read; + +use crate::DataLayerError; + +/// Hard ceiling for usage JSON after decompression. +/// +/// Usage bodies may contain large model responses, so this stays aligned with the gateway's +/// largest routinely buffered response while still bounding gzip expansion from stored data. +pub const MAX_DECOMPRESSED_USAGE_JSON_BYTES: usize = 64 * 1024 * 1024; + +pub fn read_decompressed_usage_json(reader: impl Read) -> Result, DataLayerError> { + read_decompressed_usage_json_with_limit(reader, MAX_DECOMPRESSED_USAGE_JSON_BYTES) +} + +fn read_decompressed_usage_json_with_limit( + reader: impl Read, + limit_bytes: usize, +) -> Result, DataLayerError> { + let read_limit = u64::try_from(limit_bytes) + .unwrap_or(u64::MAX) + .saturating_add(1); + let mut limited = reader.take(read_limit); + let mut decoded = Vec::new(); + limited.read_to_end(&mut decoded).map_err(|err| { + DataLayerError::UnexpectedValue(format!("failed to decompress usage json: {err}")) + })?; + if decoded.len() > limit_bytes { + return Err(DataLayerError::UnexpectedValue(format!( + "decompressed usage json exceeds {limit_bytes} bytes" + ))); + } + Ok(decoded) +} + +#[cfg(test)] +mod tests { + use std::io::Cursor; + + use super::read_decompressed_usage_json_with_limit; + + #[test] + fn decompressed_usage_json_reader_accepts_exact_limit() { + let decoded = read_decompressed_usage_json_with_limit(Cursor::new(b"1234"), 4) + .expect("payload at the hard limit should decode"); + + assert_eq!(decoded, b"1234"); + } + + #[test] + fn decompressed_usage_json_reader_rejects_limit_plus_one() { + let error = read_decompressed_usage_json_with_limit(Cursor::new(b"12345"), 4) + .expect_err("payload over the hard limit should fail"); + + assert!(error + .to_string() + .contains("decompressed usage json exceeds 4 bytes")); + } +} diff --git a/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs new file mode 100644 index 000000000..d9463ca9d --- /dev/null +++ b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs @@ -0,0 +1,1541 @@ +use std::net::IpAddr; + +use aether_ai_formats::api::{ + sanitize_request_path, sanitize_request_path_and_query, sanitize_request_query_string, +}; +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, + 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, +}; + +const UPSTREAM_IS_STREAM_KEY: &str = "upstream_is_stream"; +const PLAN_USAGE_RESERVATION_TOKEN_KEY: &str = "plan_usage_reservation_token"; +const BODY_SIZE_BASIS: &str = "serialized gateway request bodies after normalization"; + +/// Projects request metadata onto the persistence contract. Unknown fields and malformed values +/// are discarded instead of being recursively copied into an audit row. +pub fn sanitize_usage_request_metadata(value: Option) -> Option { + let Value::Object(object) = value? else { + return None; + }; + sanitize_usage_request_metadata_object(&object) +} + +pub fn sanitize_usage_request_metadata_ref(value: Option<&Value>) -> Option { + sanitize_usage_request_metadata_object(value?.as_object()?) +} + +pub fn sanitize_usage_request_metadata_object(source: &Map) -> Option { + let mut target = Map::new(); + + insert_token(source, &mut target, "trace_id", 128); + insert_ip_address(source, &mut target, "client_ip"); + insert_client_family(source, &mut target); + for key in [ + "client_requested_stream", + UPSTREAM_IS_STREAM_KEY, + "api_key_is_standalone", + WEBSOCKET_MODE_METADATA_KEY, + PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, + "transport_error", + "is_free_tier", + USAGE_AVAILABLE_METADATA_KEY, + USAGE_PRICING_AVAILABLE_METADATA_KEY, + ] { + insert_bool(source, &mut target, key); + } + insert_known_string( + source, + &mut target, + WEBSOCKET_TRANSPORT_METADATA_KEY, + sanitize_websocket_transport, + ); + if let Some(session) = source + .get(LIVE_SESSION_METADATA_KEY) + .and_then(project_live_session) + { + target.insert(LIVE_SESSION_METADATA_KEY.to_string(), session); + } + if let Some(session) = source + .get(REALTIME_SESSION_METADATA_KEY) + .and_then(project_realtime_session) + { + target.insert(REALTIME_SESSION_METADATA_KEY.to_string(), session); + } + insert_uuid(source, &mut target, PLAN_USAGE_RESERVATION_TOKEN_KEY); + insert_request_paths(source, &mut target); + for key in [ + REQUESTED_REASONING_EFFORT_METADATA_KEY, + PROVIDER_REASONING_EFFORT_METADATA_KEY, + ] { + insert_known_string(source, &mut target, key, sanitize_reasoning_effort); + } + for key in [ + PROVIDER_SERVICE_TIER_METADATA_KEY, + PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, + ] { + insert_known_string(source, &mut target, key, sanitize_service_tier); + } + insert_bounded_u64( + source, + &mut target, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, + 10_080, + ); + for key in [ + "provider_request_body_base64_bytes", + "provider_response_body_base64_bytes", + "client_response_body_base64_bytes", + "end_to_end_time_ms", + "end_to_end_first_byte_time_ms", + ] { + insert_u64(source, &mut target, key); + } + insert_bounded_u64(source, &mut target, "client_response_status_code", 599); + if let Some(body_size) = source.get("body_size").and_then(project_body_size) { + target.insert("body_size".to_string(), body_size); + } + insert_transport_error_type(source, &mut target); + + for key in ["model_id", "global_model_id", "global_model_name"] { + insert_model_token(source, &mut target, key); + } + if let Some(dimensions) = source.get("dimensions").and_then(project_dimensions) { + target.insert("dimensions".to_string(), dimensions); + } + if let Some(dimensions) = source + .get("billing_dimensions") + .and_then(project_dimensions) + { + target.insert("billing_dimensions".to_string(), dimensions); + } + insert_routing_skip_reason(source, &mut target); + if let Some(diagnostic) = source + .get(ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY) + .and_then(project_routing_failure_diagnostic) + { + target.insert( + ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY.to_string(), + diagnostic, + ); + } + + for key in [ + "rate_multiplier", + "input_price_per_1m", + "output_price_per_1m", + "cache_creation_price_per_1m", + "cache_read_price_per_1m", + "price_per_request", + ] { + insert_nonnegative_number(source, &mut target, key); + } + + let billing_snapshot = source + .get("billing_snapshot") + .and_then(project_billing_snapshot); + insert_schema_version_with_fallback( + source, + &mut target, + "billing_snapshot_schema_version", + billing_snapshot.as_ref(), + ); + insert_billing_status_with_fallback( + source, + &mut target, + "billing_snapshot_status", + billing_snapshot.as_ref(), + ); + if let Some(snapshot) = billing_snapshot { + target.insert("billing_snapshot".to_string(), snapshot); + } + + let settlement_snapshot = source + .get("settlement_snapshot") + .and_then(project_settlement_snapshot); + insert_schema_version_with_fallback( + source, + &mut target, + "settlement_snapshot_schema_version", + settlement_snapshot.as_ref(), + ); + if let Some(snapshot) = settlement_snapshot { + target.insert("settlement_snapshot".to_string(), snapshot); + } + + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn insert_request_paths(source: &Map, target: &mut Map) { + let path = source + .get("request_path") + .and_then(Value::as_str) + .and_then(sanitize_request_path); + let query = source + .get("request_query_string") + .and_then(Value::as_str) + .and_then(sanitize_request_query_string); + let combined = path + .as_deref() + .and_then(|path| sanitize_request_path_and_query(path, query.as_deref())) + .or_else(|| { + source + .get("request_path_and_query") + .and_then(Value::as_str) + .and_then(|value| sanitize_request_path_and_query(value, None)) + }); + + insert_owned_string(target, "request_path", path); + insert_owned_string(target, "request_query_string", query); + insert_owned_string(target, "request_path_and_query", combined); +} + +fn insert_client_family(source: &Map, target: &mut Map) { + let value = source + .get("client_family") + .and_then(Value::as_str) + .or_else(|| { + source + .get("client_session_affinity") + .and_then(Value::as_object) + .and_then(|affinity| affinity.get("client_family")) + .and_then(Value::as_str) + }) + .and_then(sanitize_client_family) + .or_else(|| { + source + .get("user_agent") + .and_then(Value::as_str) + .and_then(infer_client_family_from_user_agent) + }); + insert_owned_string(target, "client_family", value); +} + +fn sanitize_client_family(value: &str) -> Option { + known_lowercase( + value, + &[ + "aider", + "anthropic_js_sdk", + "anthropic_python_sdk", + "cherrystudio", + "claude_code", + "cline", + "codex", + "codex_vscode", + "continue", + "cursor", + "gemini_cli", + "generic", + "kilocode", + "langchain", + "llamaindex", + "openai_js_sdk", + "openai_python_sdk", + "opencode", + "openui", + "qwen_code", + "roo_code", + "sdk", + "unknown", + "windsurf", + ], + ) +} + +fn infer_client_family_from_user_agent(value: &str) -> Option { + let normalized = value.trim().to_ascii_lowercase(); + let family = if normalized.starts_with("codex_vscode") { + "codex_vscode" + } else if normalized.starts_with("codex") { + "codex" + } else if normalized.contains("claude-code") || normalized.contains("claude_code") { + "claude_code" + } else if normalized.contains("opencode") { + "opencode" + } else if normalized.contains("geminicli") || normalized.contains("gemini-cli") { + "gemini_cli" + } else if normalized.contains("qwencode") { + "qwen_code" + } else if normalized.contains("roo-code") || normalized.contains("roocode") { + "roo_code" + } else if normalized.contains("kilo-code") || normalized.contains("kilocode") { + "kilocode" + } else if normalized.contains("cherrystudio") || normalized.contains("cherry-studio") { + "cherrystudio" + } else if normalized.contains("openui-agent-manager") || normalized.contains("openui") { + "openui" + } else if normalized.contains("cursor") { + "cursor" + } else if normalized.contains("windsurf") { + "windsurf" + } else if normalized.contains("continue") { + "continue" + } else if normalized.contains("cline") { + "cline" + } else if normalized.contains("aider") { + "aider" + } else if normalized.contains("langchain") { + "langchain" + } else if normalized.contains("llamaindex") || normalized.contains("llama-index") { + "llamaindex" + } else if normalized.starts_with("openai/js") { + "openai_js_sdk" + } else if normalized.starts_with("openai/python") { + "openai_python_sdk" + } else if normalized.starts_with("anthropic/js") + || normalized.contains("anthropic-sdk-typescript") + { + "anthropic_js_sdk" + } else if normalized.starts_with("anthropic/python") + || normalized.contains("anthropic-sdk-python") + { + "anthropic_python_sdk" + } else if normalized.contains("/js ") || normalized.contains("/python ") { + "sdk" + } else { + return None; + }; + Some(family.to_string()) +} + +fn sanitize_websocket_transport(value: &str) -> Option { + known_lowercase( + value, + &[ + "codex_live_direct", + "codex_live_sideband", + "openai_realtime", + "openai_responses", + "responses", + ], + ) +} + +fn sanitize_reasoning_effort(value: &str) -> Option { + known_lowercase( + value, + &["none", "minimal", "low", "medium", "high", "xhigh", "max"], + ) +} + +fn sanitize_service_tier(value: &str) -> Option { + known_lowercase( + value, + &[ + "auto", + "batch", + "default", + "expedited", + "fast", + "flex", + "free_tier", + "priority", + "standard", + ], + ) +} + +fn known_lowercase(value: &str, allowed: &[&str]) -> Option { + let value = value.trim(); + if value.len() > 128 { + return None; + } + let normalized = value.to_ascii_lowercase(); + allowed.contains(&normalized.as_str()).then_some(normalized) +} + +fn project_live_session(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + insert_schema_version(source, &mut target, "schema_version"); + insert_known_string( + source, + &mut target, + "transport", + sanitize_live_session_transport, + ); + insert_known_string(source, &mut target, "mode", sanitize_live_session_mode); + insert_known_string(source, &mut target, "state", sanitize_session_state); + insert_known_string( + source, + &mut target, + "termination", + sanitize_session_termination, + ); + for key in [ + "elapsed_ms", + "client_frames", + "client_bytes", + "upstream_frames", + "upstream_bytes", + "first_upstream_frame_ms", + ] { + insert_u64(source, &mut target, key); + } + insert_known_string( + source, + &mut target, + "usage_state", + sanitize_live_usage_state, + ); + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_realtime_session(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + insert_schema_version(source, &mut target, "schema_version"); + insert_known_string( + source, + &mut target, + "transport", + sanitize_realtime_session_transport, + ); + insert_known_string(source, &mut target, "state", sanitize_session_state); + insert_known_string( + source, + &mut target, + "termination", + sanitize_session_termination, + ); + for key in [ + "elapsed_ms", + "client_frames", + "client_bytes", + "upstream_frames", + "upstream_bytes", + "first_upstream_frame_ms", + "usage_response_count", + "cached_input_tokens", + "input_audio_tokens", + "output_audio_tokens", + ] { + insert_u64(source, &mut target, key); + } + insert_known_string( + source, + &mut target, + "usage_state", + sanitize_realtime_usage_state, + ); + insert_known_string( + source, + &mut target, + "pricing_state", + sanitize_realtime_pricing_state, + ); + insert_known_string(source, &mut target, "usage_scope", sanitize_usage_scope); + insert_bool(source, &mut target, "input_transcription_usage_included"); + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn sanitize_live_session_transport(value: &str) -> Option { + known_lowercase(value, &["sideband", "webrtc", "websocket"]) +} + +fn sanitize_realtime_session_transport(value: &str) -> Option { + known_lowercase(value, &["websocket"]) +} + +fn sanitize_live_session_mode(value: &str) -> Option { + known_lowercase(value, &["call_create", "direct", "sideband"]) +} + +fn sanitize_session_state(value: &str) -> Option { + known_lowercase(value, &["cancelled", "closed", "failed"]) +} + +fn sanitize_live_usage_state(value: &str) -> Option { + known_lowercase(value, &["unavailable"]) +} + +fn sanitize_realtime_usage_state(value: &str) -> Option { + known_lowercase(value, &["authoritative", "unavailable"]) +} + +fn sanitize_realtime_pricing_state(value: &str) -> Option { + known_lowercase( + value, + &[ + "compatible_text_usage", + "unsupported_audio_breakdown", + "usage_unavailable", + ], + ) +} + +fn sanitize_usage_scope(value: &str) -> Option { + known_lowercase(value, &["response_done"]) +} + +fn sanitize_session_termination(value: &str) -> Option { + let value = value.trim(); + if value.len() > 96 { + return Some("other".to_string()); + } + let normalized = value.to_ascii_lowercase(); + let allowed = [ + "admission_failed", + "admission_plan_unavailable", + "admission_planning_timeout", + "admission_timeout", + "auth_context_missing", + "authentication_required", + "balance_capacity_check_failed", + "balance_capacity_rejected", + "balance_rejected", + "binding_failed", + "binding_changed", + "call_created", + "candidate_unavailable", + "client_close_frame", + "client_closed", + "client_read_failed", + "client_write_failed", + "connection_admission_lost", + "connection_duration_limit", + "control_unavailable", + "codex_live_architecture_invalid", + "codex_live_boundary_invalid", + "codex_live_body_too_large", + "codex_live_call_id_invalid", + "codex_live_call_location_invalid", + "codex_live_expected_session_update", + "codex_live_initial_client_read_failed", + "codex_live_initial_event_invalid", + "codex_live_initial_event_must_be_text", + "codex_live_initial_session_update_timeout", + "codex_live_intent_invalid", + "codex_live_media_type_unsupported", + "codex_live_model_invalid", + "codex_live_model_query_invalid", + "codex_live_multipart_invalid", + "codex_live_multipart_part_duplicate", + "codex_live_multipart_part_unexpected", + "codex_live_oauth_direct_unsupported", + "codex_live_oauth_upstream_unsupported", + "codex_live_sdp_invalid", + "codex_live_sdp_missing", + "codex_live_sdp_too_large", + "codex_live_session_invalid", + "codex_live_session_missing", + "codex_live_session_too_large", + "codex_live_upstream_url_invalid", + "codex_live_upstream_url_missing", + "downstream_response_build_failed", + "explicit_failure", + "finite_balance_unsupported", + "initial_upstream_write_failed", + "last_admin_delete_denied", + "last_admin_update_denied", + "location_invalid", + "location_missing", + "model_invalid", + "model_missing", + "multipart_parse_failed", + "planning_failed", + "pool_key_lease_lost", + "pool_lease_lost", + "provider_body_build_failed", + "provider_plan_build_failed", + "provider_plan_unavailable", + "relay_cancelled", + "request_rejected", + "request_body_missing", + "request_body_too_large", + "request_future_cancelled", + "response_body_unavailable", + "route_unavailable", + "session_close_drain_timeout", + "sideband_attachment_conflict", + "sideband_attachment_lease_lost", + "sideband_attachment_lease_renewal_failed", + "sideband_attachment_timeout", + "sideband_attachment_unavailable", + "sideband_binding_changed", + "sideband_binding_disabled", + "sideband_binding_expired", + "sideband_binding_lookup_timeout", + "sideband_binding_missing", + "sideband_binding_unavailable", + "upstream_close_frame", + "upstream_closed", + "upstream_connect_failed", + "upstream_error_body_unavailable", + "upstream_execute_failed", + "upstream_read_failed", + "upstream_rejected", + "upstream_url_invalid", + "upstream_write_failed", + "usage_settlement_unavailable", + ]; + if allowed.contains(&normalized.as_str()) { + Some(normalized) + } else { + Some("other".to_string()) + } +} + +fn insert_transport_error_type(source: &Map, target: &mut Map) { + let Some(value) = source + .get("transport_error_type") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return; + }; + let normalized = value.to_ascii_lowercase(); + let value = match normalized.as_str() { + "chatgpt_web_image_execution_unavailable" + | "connect_timeout" + | "execution_runtime_unavailable" + | "first_byte_timeout" + | "gateway_admission_timeout" + | "grok_execution_unavailable" + | "kiro_web_search_mcp_unavailable" + | "local_stream_candidate_watchdog_timeout" + | "protocol_error" + | "proxy_error" + | "read_timeout" + | "tls_error" + | "upstream_transport_error" + | "windsurf_native_execution_unavailable" => normalized, + _ => "other_transport_error".to_string(), + }; + target.insert("transport_error_type".to_string(), Value::String(value)); +} + +fn insert_routing_skip_reason(source: &Map, target: &mut Map) { + let value = source + .get(ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + if let Some(value) = sanitize_request_candidate_skip_reason(value) { + target.insert( + ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY.to_string(), + Value::String(value), + ); + } +} + +fn project_routing_failure_diagnostic(value: &Value) -> Option { + let source = value.as_object()?; + let kind = source + .get("kind") + .and_then(Value::as_str) + .and_then(|value| { + known_lowercase( + value, + &[ + "body_rules", + "envelope_build", + "header_rules", + "request_body_build", + "request_conversion", + "transport_auth", + "url_build", + ], + ) + })?; + let mut target = Map::from_iter([("kind".to_string(), Value::String(kind))]); + if let Some(path) = source + .get("path") + .and_then(Value::as_str) + .and_then(sanitize_json_path) + { + target.insert("path".to_string(), Value::String(path)); + } + Some(Value::Object(target)) +} + +fn sanitize_json_path(value: &str) -> Option { + let value = value.trim(); + if value.is_empty() + || value.len() > 256 + || !value.starts_with('$') + || !value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"$._[]:-".contains(&byte)) + { + return None; + } + Some(value.to_string()) +} + +fn project_body_size(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + if source + .get("basis") + .and_then(Value::as_str) + .map(str::trim) + .is_some_and(|value| value == BODY_SIZE_BASIS) + { + target.insert( + "basis".to_string(), + Value::String(BODY_SIZE_BASIS.to_string()), + ); + } + for key in [ + "client_request_body", + "provider_request_body", + "provider_over_client", + ] { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.len() <= 32 + && value.bytes().all(|byte| { + byte.is_ascii_digit() + || byte == b'.' + || byte == b' ' + || matches!(byte, b'B' | b'K' | b'M' | b'G' | b'x') + }) + }) + else { + continue; + }; + target.insert(key.to_string(), Value::String(value.to_string())); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_dimensions(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + for key in [ + "input_tokens", + "effective_input_tokens", + "output_tokens", + "total_tokens", + "reasoning_tokens", + "cache_creation_tokens", + "cache_creation_uncategorized_tokens", + "cache_creation_ephemeral_5m_tokens", + "cache_creation_ephemeral_1h_tokens", + "cache_read_tokens", + "request_count", + "image_count", + "image_count_unmetered", + "total_input_context", + "cache_ttl_minutes", + "cache_creation_ephemeral_5m_ttl_minutes", + "cache_creation_ephemeral_1h_ttl_minutes", + "image_pixels", + "windsurf_generator_entry_count", + ] { + insert_u64(source, &mut target, key); + } + for key in ["cache_storage_token_hours", "image_output_price_per_image"] { + insert_nonnegative_number(source, &mut target, key); + } + for key in [ + "image_output_pricing_enabled", + "image_output_matrix_enabled", + "image_output_range_enabled", + ] { + insert_bool(source, &mut target, key); + } + insert_known_string( + source, + &mut target, + "effective_task_type", + sanitize_task_type, + ); + for key in [ + "requested_processing_tier", + "actual_processing_tier", + "billing_processing_tier", + ] { + insert_nullable_known_string(source, &mut target, key, sanitize_service_tier); + } + insert_known_string( + source, + &mut target, + "image_output_pricing_mode", + sanitize_image_pricing_mode, + ); + insert_known_string(source, &mut target, "image_quality", sanitize_image_quality); + insert_known_string( + source, + &mut target, + "image_output_format", + sanitize_image_output_format, + ); + for key in ["image_size", "image_price_key", "image_output_price_bucket"] { + insert_dimension_token(source, &mut target, key); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn sanitize_task_type(value: &str) -> Option { + known_lowercase( + value, + &["chat", "embedding", "image", "rerank", "search", "video"], + ) +} + +fn sanitize_image_pricing_mode(value: &str) -> Option { + known_lowercase(value, &["matrix", "none", "per_image", "pixel_tiers"]) +} + +fn sanitize_image_quality(value: &str) -> Option { + known_lowercase(value, &["auto", "low", "medium", "high", "standard", "hd"]) +} + +fn sanitize_image_output_format(value: &str) -> Option { + known_lowercase(value, &["jpeg", "jpg", "png", "webp"]) +} + +fn project_billing_snapshot(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + insert_schema_version(source, &mut target, "schema_version"); + insert_token(source, &mut target, "rule_id", 128); + insert_billing_status(source, &mut target, "status"); + if let Some(value) = source + .get("resolved_dimensions") + .and_then(project_dimensions) + { + target.insert("resolved_dimensions".to_string(), value); + } + if let Some(value) = source + .get("resolved_variables") + .and_then(project_resolved_variables) + { + target.insert("resolved_variables".to_string(), value); + } + if let Some(value) = source + .get("cost_breakdown") + .and_then(project_cost_breakdown) + { + target.insert("cost_breakdown".to_string(), value); + } + insert_nonnegative_number(source, &mut target, "total_cost"); + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_settlement_snapshot(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + insert_schema_version(source, &mut target, "schema_version"); + insert_billing_status(source, &mut target, "status"); + for key in ["total_cost", "actual_total_cost"] { + insert_nonnegative_number(source, &mut target, key); + } + if let Some(value) = source + .get("pricing_snapshot") + .and_then(project_pricing_snapshot) + { + target.insert("pricing_snapshot".to_string(), value); + } + if let Some(value) = source + .get("billing_plan_snapshot") + .and_then(project_billing_plan_snapshot) + { + target.insert("billing_plan_snapshot".to_string(), value); + } + if let Some(value) = source + .get("resolved_dimensions") + .and_then(project_dimensions) + { + target.insert("resolved_dimensions".to_string(), value); + } + if let Some(value) = source + .get("resolved_variables") + .and_then(project_resolved_variables) + { + target.insert("resolved_variables".to_string(), value); + } + if let Some(value) = source + .get("cost_breakdown") + .and_then(project_cost_breakdown) + { + target.insert("cost_breakdown".to_string(), value); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_pricing_snapshot(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + for key in [ + "requested_processing_tier", + "actual_processing_tier", + "billing_processing_tier", + ] { + insert_nullable_known_string(source, &mut target, key, sanitize_service_tier); + } + for key in [ + "pricing_source", + "tiered_pricing_source", + "price_per_request_source", + ] { + insert_nullable_known_string(source, &mut target, key, sanitize_pricing_source); + } + for key in [ + "processing_tier_price_multiplier", + "price_per_request", + "rate_multiplier", + ] { + insert_nonnegative_number(source, &mut target, key); + } + insert_bool(source, &mut target, "is_free_tier"); + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn sanitize_pricing_source(value: &str) -> Option { + known_lowercase( + value, + &["global_default", "mixed", "provider_override", "unpriced"], + ) +} + +fn project_billing_plan_snapshot(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + insert_token(source, &mut target, "rule_id", 128); + if let Some(value) = source.get("rule_version").and_then(safe_version_value) { + target.insert("rule_version".to_string(), Value::String(value)); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_resolved_variables(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + for key in [ + "input_price_per_1m", + "output_price_per_1m", + "cache_creation_price_per_1m", + "cache_creation_ephemeral_5m_price_per_1m", + "cache_creation_ephemeral_1h_price_per_1m", + "cache_read_price_per_1m", + "price_per_request", + "image_output_price_per_image", + ] { + insert_nonnegative_number(source, &mut target, key); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn project_cost_breakdown(value: &Value) -> Option { + let source = value.as_object()?; + let mut target = Map::new(); + for key in [ + "input_cost", + "output_cost", + "cache_creation_uncategorized_cost", + "cache_creation_ephemeral_5m_cost", + "cache_creation_ephemeral_1h_cost", + "cache_creation_cost", + "cache_read_cost", + "image_output_cost", + "request_cost", + ] { + insert_nonnegative_number(source, &mut target, key); + } + (!target.is_empty()).then_some(Value::Object(target)) +} + +fn insert_schema_version_with_fallback( + source: &Map, + target: &mut Map, + key: &str, + snapshot: Option<&Value>, +) { + let value = source + .get(key) + .and_then(Value::as_str) + .and_then(sanitize_schema_version) + .or_else(|| { + snapshot + .and_then(Value::as_object) + .and_then(|snapshot| snapshot.get("schema_version")) + .and_then(Value::as_str) + .and_then(sanitize_schema_version) + }); + insert_owned_string(target, key, value); +} + +fn insert_billing_status_with_fallback( + source: &Map, + target: &mut Map, + key: &str, + snapshot: Option<&Value>, +) { + let value = source + .get(key) + .and_then(Value::as_str) + .and_then(sanitize_billing_status) + .or_else(|| { + snapshot + .and_then(Value::as_object) + .and_then(|snapshot| snapshot.get("status")) + .and_then(Value::as_str) + .and_then(sanitize_billing_status) + }); + insert_owned_string(target, key, value); +} + +fn insert_schema_version(source: &Map, target: &mut Map, key: &str) { + insert_known_string(source, target, key, sanitize_schema_version); +} + +fn sanitize_schema_version(value: &str) -> Option { + let value = value.trim(); + if value.len() > 32 { + return None; + } + let value = value.to_ascii_lowercase(); + (!value.is_empty() + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"._-".contains(&byte))) + .then_some(value) +} + +fn insert_billing_status(source: &Map, target: &mut Map, key: &str) { + insert_known_string(source, target, key, sanitize_billing_status); +} + +fn sanitize_billing_status(value: &str) -> Option { + known_lowercase( + value, + &[ + "complete", + "incomplete", + "legacy", + "no_rule", + "pending", + "resolved", + "void", + ], + ) +} + +fn insert_ip_address(source: &Map, target: &mut Map, key: &str) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .and_then(|value| value.parse::().ok()) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_string())); +} + +fn insert_uuid(source: &Map, target: &mut Map, key: &str) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| is_canonical_uuid(value)) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_ascii_lowercase())); +} + +fn is_canonical_uuid(value: &str) -> bool { + value.len() == 36 + && value.bytes().enumerate().all(|(index, byte)| { + if matches!(index, 8 | 13 | 18 | 23) { + byte == b'-' + } else { + byte.is_ascii_hexdigit() + } + }) +} + +fn insert_model_token(source: &Map, target: &mut Map, key: &str) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.len() <= 256 + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"._:/@+-".contains(&byte)) + }) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_string())); +} + +fn insert_token( + source: &Map, + target: &mut Map, + key: &str, + max_len: usize, +) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.len() <= max_len + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"._:-".contains(&byte)) + }) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_string())); +} + +fn insert_dimension_token(source: &Map, target: &mut Map, key: &str) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.len() <= 64 + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"._:<>=".contains(&byte)) + }) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_ascii_lowercase())); +} + +fn insert_known_string( + source: &Map, + target: &mut Map, + key: &str, + sanitize: fn(&str) -> Option, +) { + let value = source.get(key).and_then(Value::as_str).and_then(sanitize); + insert_owned_string(target, key, value); +} + +fn insert_nullable_known_string( + source: &Map, + target: &mut Map, + key: &str, + sanitize: fn(&str) -> Option, +) { + match source.get(key) { + Some(Value::Null) => { + target.insert(key.to_string(), Value::Null); + } + Some(Value::String(value)) => { + if let Some(value) = sanitize(value) { + target.insert(key.to_string(), Value::String(value)); + } + } + _ => {} + } +} + +fn insert_owned_string(target: &mut Map, key: &str, value: Option) { + if let Some(value) = value { + target.insert(key.to_string(), Value::String(value)); + } +} + +fn insert_bool(source: &Map, target: &mut Map, key: &str) { + if let Some(value) = source.get(key).and_then(Value::as_bool) { + target.insert(key.to_string(), Value::Bool(value)); + } +} + +fn insert_u64(source: &Map, target: &mut Map, key: &str) { + if let Some(value) = source.get(key).and_then(Value::as_u64) { + target.insert(key.to_string(), Value::Number(value.into())); + } +} + +fn insert_bounded_u64( + source: &Map, + target: &mut Map, + key: &str, + max: u64, +) { + if let Some(value) = source + .get(key) + .and_then(Value::as_u64) + .filter(|value| *value <= max) + { + target.insert(key.to_string(), Value::Number(value.into())); + } +} + +fn insert_nonnegative_number( + source: &Map, + target: &mut Map, + key: &str, +) { + let Some(value) = source.get(key).filter(|value| { + value + .as_f64() + .is_some_and(|value| value.is_finite() && value >= 0.0) + }) else { + return; + }; + target.insert(key.to_string(), value.clone()); +} + +fn safe_version_value(value: &Value) -> Option { + match value { + Value::String(value) => sanitize_schema_version(value), + Value::Number(value) => sanitize_schema_version(&value.to_string()), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{sanitize_usage_request_metadata, sanitize_usage_request_metadata_ref}; + + #[test] + fn persistence_projection_drops_credentials_and_free_diagnostics() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "trace_id": "trace-1", + "client_ip": "203.0.113.8", + "user_agent": "Bearer browser-secret", + "client_session_affinity": { + "client_family": "codex", + "session_key": "tenant/session-secret" + }, + "proxy": {"url": "https://user:pass@proxy.example"}, + "tls_fingerprint": {"ja3": "fingerprint"}, + "stage_timings_ms": {"secret_operation": 12}, + "db_timings_ms": {"query": "SELECT credential"}, + "scheduling_audit": {"key_id": "secret-key"}, + "billing_rule_snapshot": {"expression": "tenant_secret * input_tokens"}, + "routing_candidate_skip_reason": "Authorization: Bearer secret", + "routing_failure_diagnostic": { + "kind": "request_body_build", + "path": "$.reasoning.summary", + "message": "Bearer secret", + "source": "https://user:pass@example.com", + "client_api_format": "secret", + "provider_api_format": "secret" + }, + "unknown": {"authorization": "Bearer secret"} + }))) + .expect("safe metadata should remain"); + + assert_eq!(metadata["trace_id"], "trace-1"); + assert_eq!(metadata["client_ip"], "203.0.113.8"); + assert_eq!(metadata["client_family"], "codex"); + assert_eq!( + metadata["routing_candidate_skip_reason"], + "unclassified_skip" + ); + assert_eq!( + metadata["routing_failure_diagnostic"], + json!({"kind": "request_body_build", "path": "$.reasoning.summary"}) + ); + for key in [ + "user_agent", + "client_session_affinity", + "proxy", + "tls_fingerprint", + "stage_timings_ms", + "db_timings_ms", + "scheduling_audit", + "billing_rule_snapshot", + "unknown", + ] { + assert!(metadata.get(key).is_none(), "{key} must not be persisted"); + } + } + + #[test] + fn persistence_projection_keeps_only_bounded_settlement_facts() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "billing_snapshot": { + "schema_version": "2.0", + "rule_id": "__default__", + "rule_name": "tenant secret pricing", + "scope": "secret scope", + "expression": "secret_variable * input_tokens", + "resolved_dimensions": { + "input_tokens": 100, + "image_quality": "high", + "tenant_id": "tenant-secret" + }, + "resolved_variables": { + "input_price_per_1m": 3.0, + "api_key": "secret" + }, + "cost_breakdown": { + "input_cost": 0.0003, + "tenant_secret_cost": 99 + }, + "total_cost": 0.0003, + "status": "complete", + "tier_info": {"catalog": "secret"} + }, + "settlement_snapshot": { + "schema_version": "3.0", + "pricing_snapshot": { + "provider_api_key_id": "key-secret", + "tiered_pricing": {"tenant": "secret"}, + "billing_processing_tier": "priority", + "pricing_source": "provider_override", + "price_per_request": 0.02, + "rate_multiplier": 1.25, + "is_free_tier": false + }, + "billing_plan_snapshot": { + "rule_id": "rule-1", + "rule_version": "7", + "rule_name": "secret rule", + "expression": "secret" + }, + "resolved_dimensions": {"input_tokens": 100, "secret": "value"}, + "resolved_variables": {"input_price_per_1m": 3.0, "secret": 4}, + "cost_breakdown": {"input_cost": 0.0003, "secret_cost": 4}, + "total_cost": 0.0003, + "actual_total_cost": 0.000375, + "status": "complete", + "calculated_at": "secret timestamp" + }, + "billing_dimensions": { + "input_tokens": 100, + "image_size": "1024x1024", + "secret_dimension": "secret" + } + }))) + .expect("settlement facts should remain"); + + assert_eq!( + metadata["plan_usage_reservation_token"], + "550e8400-e29b-41d4-a716-446655440000" + ); + assert_eq!(metadata["billing_snapshot_schema_version"], "2.0"); + assert_eq!(metadata["billing_snapshot_status"], "complete"); + assert_eq!(metadata["settlement_snapshot_schema_version"], "3.0"); + assert_eq!( + metadata.pointer("/billing_snapshot/resolved_variables/input_price_per_1m"), + Some(&json!(3.0)) + ); + assert_eq!( + metadata.pointer("/settlement_snapshot/billing_plan_snapshot/rule_version"), + Some(&json!("7")) + ); + assert_eq!( + metadata.pointer("/settlement_snapshot/pricing_snapshot/pricing_source"), + Some(&json!("provider_override")) + ); + assert!(metadata.pointer("/billing_snapshot/expression").is_none()); + assert!(metadata + .pointer("/settlement_snapshot/pricing_snapshot/provider_api_key_id") + .is_none()); + assert!(metadata + .pointer("/settlement_snapshot/pricing_snapshot/tiered_pricing") + .is_none()); + assert!(metadata + .pointer("/billing_dimensions/secret_dimension") + .is_none()); + } + + #[test] + fn persistence_projection_rejects_malformed_identifiers_and_paths() { + assert!(sanitize_usage_request_metadata(Some(json!({ + "client_ip": "127.0.0.1, 10.0.0.1", + "trace_id": "Bearer secret", + "plan_usage_reservation_token": "server-token", + "request_path": "/install/sensitive-code", + "request_query_string": "key=secret&alt=sse", + "routing_failure_diagnostic": { + "kind": "request_body_build", + "path": "$['api_key=secret']" + } + }))) + .is_some_and(|metadata| { + metadata.get("client_ip").is_none() + && metadata.get("trace_id").is_none() + && metadata.get("plan_usage_reservation_token").is_none() + && metadata["request_path"] == "/install/[redacted]" + && metadata["request_query_string"] == "alt=sse" + && metadata + .pointer("/routing_failure_diagnostic/path") + .is_none() + })); + } + + #[test] + fn persistence_projection_keeps_only_known_body_size_basis() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "body_size": { + "basis": " serialized gateway request bodies after normalization ", + "client_request_body": "1 KB", + "provider_request_body": "4 KB", + "untrusted": "drop-me" + } + }))) + .expect("body size metadata should remain"); + + assert_eq!( + metadata, + json!({ + "body_size": { + "basis": "serialized gateway request bodies after normalization", + "client_request_body": "1 KB", + "provider_request_body": "4 KB" + } + }) + ); + assert!(sanitize_usage_request_metadata(Some(json!({ + "body_size": {"basis": "untrusted basis"} + }))) + .is_none()); + } + + #[test] + fn borrowed_and_owned_projection_match() { + let value = json!({ + "trace_id": "trace-1", + "client_ip": "2001:db8::1", + "billing_dimensions": {"input_tokens": 5} + }); + assert_eq!( + sanitize_usage_request_metadata_ref(Some(&value)), + sanitize_usage_request_metadata(Some(value)) + ); + } + + #[test] + fn user_agent_is_reduced_to_a_controlled_client_family() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "user_agent": "codex_vscode/0.131.0-alpha.9 (Windows; x86_64; tenant=secret)" + }))) + .expect("recognized client family should remain"); + assert_eq!(metadata, json!({"client_family": "codex_vscode"})); + + assert!(sanitize_usage_request_metadata(Some(json!({ + "user_agent": "private-client/1.0 account-secret" + }))) + .is_none()); + } + + #[test] + fn persistence_projection_keeps_bounded_websocket_session_facts() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "websocket_mode": true, + "websocket_transport": "CODEX_LIVE_DIRECT", + "usage_available": false, + "usage_pricing_available": false, + "live_session": { + "schema_version": "1", + "transport": "websocket", + "mode": "direct", + "state": "cancelled", + "termination": "client_close_frame", + "elapsed_ms": 1200, + "client_frames": 4, + "client_bytes": 128, + "upstream_frames": 3, + "upstream_bytes": 96, + "first_upstream_frame_ms": 42, + "usage_state": "unavailable", + "authorization": "Bearer secret", + "nested": {"secret": "drop-me"} + }, + "realtime_session": { + "schema_version": "1", + "transport": "websocket", + "state": "failed", + "termination": "Bearer secret", + "elapsed_ms": 900, + "client_frames": 2, + "client_bytes": 64, + "upstream_frames": 1, + "upstream_bytes": 32, + "usage_state": "authoritative", + "pricing_state": "unsupported_audio_breakdown", + "usage_scope": "response_done", + "input_transcription_usage_included": false, + "usage_response_count": 1, + "cached_input_tokens": 5, + "input_audio_tokens": 6, + "output_audio_tokens": 7, + "authorization": "Bearer secret", + "nested": {"secret": "drop-me"} + } + }))) + .expect("bounded session metadata should remain"); + + assert_eq!(metadata["websocket_mode"], true); + assert_eq!(metadata["websocket_transport"], "codex_live_direct"); + assert_eq!(metadata["usage_available"], false); + assert_eq!(metadata["usage_pricing_available"], false); + assert_eq!( + metadata["live_session"], + json!({ + "schema_version": "1", + "transport": "websocket", + "mode": "direct", + "state": "cancelled", + "termination": "client_close_frame", + "elapsed_ms": 1200, + "client_frames": 4, + "client_bytes": 128, + "upstream_frames": 3, + "upstream_bytes": 96, + "first_upstream_frame_ms": 42, + "usage_state": "unavailable" + }) + ); + assert_eq!( + metadata["realtime_session"], + json!({ + "schema_version": "1", + "transport": "websocket", + "state": "failed", + "termination": "other", + "elapsed_ms": 900, + "client_frames": 2, + "client_bytes": 64, + "upstream_frames": 1, + "upstream_bytes": 32, + "usage_state": "authoritative", + "pricing_state": "unsupported_audio_breakdown", + "usage_scope": "response_done", + "input_transcription_usage_included": false, + "usage_response_count": 1, + "cached_input_tokens": 5, + "input_audio_tokens": 6, + "output_audio_tokens": 7 + }) + ); + assert!(metadata.pointer("/live_session/authorization").is_none()); + assert!(metadata.pointer("/live_session/nested").is_none()); + assert!(metadata + .pointer("/realtime_session/authorization") + .is_none()); + assert!(metadata.pointer("/realtime_session/nested").is_none()); + } +} diff --git a/crates/aether-data/contracts/src/repository/usage/mod.rs b/crates/aether-data/contracts/src/repository/usage/mod.rs index 00329640d..6ed29768f 100644 --- a/crates/aether-data/contracts/src/repository/usage/mod.rs +++ b/crates/aether-data/contracts/src/repository/usage/mod.rs @@ -1,9 +1,13 @@ +mod compression; +mod metadata_policy; mod policy; mod types; +pub use compression::{read_decompressed_usage_json, MAX_DECOMPRESSED_USAGE_JSON_BYTES}; +pub use metadata_policy::*; pub use policy::*; pub use types::{ - extract_provider_actual_service_tier_from_response, + 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, resolve_provider_cache_ttl_minutes, resolve_provider_service_tier_from_request_capture, @@ -34,10 +38,11 @@ pub use types::{ UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery, UsageWriteRepository, LIVE_SESSION_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, + 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, }; diff --git a/crates/aether-data/contracts/src/repository/usage/policy.rs b/crates/aether-data/contracts/src/repository/usage/policy.rs index c8780a420..ee2b46a2b 100644 --- a/crates/aether-data/contracts/src/repository/usage/policy.rs +++ b/crates/aether-data/contracts/src/repository/usage/policy.rs @@ -1,4 +1,14 @@ use super::{StoredRequestUsageAudit, UpsertUsageRecord}; +use serde_json::{Map, Value}; + +const MAX_USAGE_CANDIDATE_INDEX: u64 = i32::MAX as u64; +const MAX_USAGE_CANDIDATE_ID_LEN: usize = 128; +const MAX_USAGE_KEY_NAME_LEN: usize = 255; +const MAX_USAGE_PLANNER_KIND_LEN: usize = 64; +const MAX_USAGE_ROUTE_FAMILY_LEN: usize = 80; +const MAX_USAGE_ROUTE_KIND_LEN: usize = 80; +const MAX_USAGE_EXECUTION_PATH_LEN: usize = 80; +const MAX_USAGE_RUNTIME_MISS_REASON_LEN: usize = 120; #[derive(Debug, Clone, PartialEq, Default)] pub struct ApiKeyUsageContribution { @@ -201,12 +211,289 @@ pub fn usage_can_recover_terminal_failure( && incoming_usage_can_recover_terminal_failure(incoming_status, incoming_billing_status) } +/// Decide whether an incoming lifecycle event may replace an existing usage revision. +/// +/// `updated_at_unix_secs` is authoritative. `finalized_at_unix_secs` breaks ties when writers +/// observe multiple transitions in the same second. Equal pending revisions may still progress to +/// streaming or terminal states, while every terminal replay requires a strictly newer revision. +/// The explicit void-failure recovery remains available at an equal revision, but never for an +/// older event. +#[allow(clippy::too_many_arguments)] +pub fn usage_lifecycle_update_allowed( + existing_status: &str, + existing_billing_status: &str, + existing_updated_at_unix_secs: u64, + existing_finalized_at_unix_secs: Option, + incoming_status: &str, + incoming_billing_status: &str, + incoming_updated_at_unix_secs: u64, + incoming_finalized_at_unix_secs: Option, +) -> bool { + let existing_revision = ( + existing_updated_at_unix_secs, + existing_finalized_at_unix_secs.unwrap_or_default(), + ); + let incoming_revision = ( + incoming_updated_at_unix_secs, + incoming_finalized_at_unix_secs.unwrap_or_default(), + ); + if incoming_revision < existing_revision { + return false; + } + + let can_recover = usage_can_recover_terminal_failure( + existing_status, + existing_billing_status, + incoming_status, + incoming_billing_status, + ); + let existing_is_terminal = matches!(existing_status, "completed" | "failed" | "cancelled"); + let incoming_is_terminal = matches!(incoming_status, "completed" | "failed" | "cancelled"); + if existing_is_terminal && !incoming_is_terminal { + return false; + } + if existing_status == "streaming" && incoming_status == "pending" { + return false; + } + if incoming_revision == existing_revision && existing_is_terminal && incoming_is_terminal { + return can_recover; + } + + true +} + pub fn strip_deprecated_usage_display_fields(mut usage: UpsertUsageRecord) -> UpsertUsageRecord { usage.username = None; usage.api_key_name = None; usage } +pub fn sanitize_usage_for_persistence(mut usage: UpsertUsageRecord) -> UpsertUsageRecord { + usage = strip_deprecated_usage_display_fields(usage); + sanitize_usage_routing_fields(&mut usage, None); + usage.error_message = None; + usage.error_category = sanitize_usage_error_category(usage.error_category); + if usage.error_category.is_none() && usage.status == "failed" { + usage.error_category = usage + .status_code + .map(usage_error_category_for_status_code) + .map(str::to_string); + } + usage.request_metadata = super::sanitize_usage_request_metadata(usage.request_metadata); + usage.request_headers = None; + usage.request_body = None; + usage.request_body_ref = None; + usage.request_body_state = None; + usage.provider_request_headers = None; + usage.provider_request_body = None; + usage.provider_request_body_ref = None; + usage.provider_request_body_state = None; + usage.response_headers = None; + usage.response_body = None; + usage.response_body_ref = None; + usage.response_body_state = None; + usage.client_response_headers = None; + usage.client_response_body = None; + usage.client_response_body_ref = None; + usage.client_response_body_state = None; + usage +} + +/// Project an event onto the non-content controls accepted by auxiliary usage storage. +/// +/// Explicit `none` states are retained only as tombstones for removing historical captures. +/// Every header, body, reference, and non-clear capture state is discarded. +pub fn sanitize_usage_capture_controls_for_persistence( + mut usage: UpsertUsageRecord, +) -> UpsertUsageRecord { + // Routing facts are allowed in the transient event metadata for compatibility with older + // writers. Project only the known scalar fields into typed slots before the general metadata + // sanitizer drops unknown keys. This keeps snapshots useful without re-persisting arbitrary + // metadata (or any body/header material). + let metadata = usage + .request_metadata + .as_ref() + .and_then(Value::as_object) + .cloned(); + sanitize_usage_routing_fields(&mut usage, metadata.as_ref()); + let clear_request_body = usage.request_body_state == Some(super::UsageBodyCaptureState::None); + let clear_provider_request_body = + usage.provider_request_body_state == Some(super::UsageBodyCaptureState::None); + let clear_response_body = usage.response_body_state == Some(super::UsageBodyCaptureState::None); + let clear_client_response_body = + usage.client_response_body_state == Some(super::UsageBodyCaptureState::None); + + let mut usage = sanitize_usage_for_persistence(usage); + usage.request_body_state = clear_request_body.then_some(super::UsageBodyCaptureState::None); + usage.provider_request_body_state = + clear_provider_request_body.then_some(super::UsageBodyCaptureState::None); + usage.response_body_state = clear_response_body.then_some(super::UsageBodyCaptureState::None); + usage.client_response_body_state = + clear_client_response_body.then_some(super::UsageBodyCaptureState::None); + usage +} + +fn sanitize_usage_routing_fields( + usage: &mut UpsertUsageRecord, + metadata: Option<&Map>, +) { + usage.candidate_id = sanitize_usage_routing_string_with_metadata( + usage.candidate_id.take(), + metadata, + "candidate_id", + MAX_USAGE_CANDIDATE_ID_LEN, + false, + ); + usage.candidate_index = + sanitize_usage_routing_index_with_metadata(usage.candidate_index.take(), metadata); + usage.key_name = sanitize_usage_routing_string_with_metadata( + usage.key_name.take(), + metadata, + "key_name", + MAX_USAGE_KEY_NAME_LEN, + true, + ); + usage.planner_kind = sanitize_usage_routing_string_with_metadata( + usage.planner_kind.take(), + metadata, + "planner_kind", + MAX_USAGE_PLANNER_KIND_LEN, + false, + ); + usage.route_family = sanitize_usage_routing_string_with_metadata( + usage.route_family.take(), + metadata, + "route_family", + MAX_USAGE_ROUTE_FAMILY_LEN, + false, + ); + usage.route_kind = sanitize_usage_routing_string_with_metadata( + usage.route_kind.take(), + metadata, + "route_kind", + MAX_USAGE_ROUTE_KIND_LEN, + false, + ); + usage.execution_path = sanitize_usage_routing_string_with_metadata( + usage.execution_path.take(), + metadata, + "execution_path", + MAX_USAGE_EXECUTION_PATH_LEN, + false, + ); + usage.local_execution_runtime_miss_reason = sanitize_usage_routing_string_with_metadata( + usage.local_execution_runtime_miss_reason.take(), + metadata, + "local_execution_runtime_miss_reason", + MAX_USAGE_RUNTIME_MISS_REASON_LEN, + false, + ); +} + +fn sanitize_usage_routing_string_with_metadata( + typed: Option, + metadata: Option<&Map>, + key: &str, + max_len: usize, + allow_spaces: bool, +) -> Option { + match typed { + Some(value) => sanitize_usage_routing_string(Some(value), max_len, allow_spaces), + None => metadata_routing_string(metadata, key, max_len, allow_spaces), + } +} + +fn sanitize_usage_routing_index_with_metadata( + typed: Option, + metadata: Option<&Map>, +) -> Option { + match typed { + Some(value) => sanitize_usage_routing_index(Some(value)), + None => metadata + .and_then(|object| object.get("candidate_index")) + .and_then(|value| { + value + .as_u64() + .or_else(|| value.as_i64().and_then(|value| u64::try_from(value).ok())) + .filter(|value| *value <= MAX_USAGE_CANDIDATE_INDEX) + }), + } +} + +fn metadata_routing_string( + metadata: Option<&Map>, + key: &str, + max_len: usize, + allow_spaces: bool, +) -> Option { + metadata + .and_then(|object| object.get(key)) + .and_then(Value::as_str) + .and_then(|value| { + sanitize_usage_routing_string(Some(value.to_string()), max_len, allow_spaces) + }) +} + +fn sanitize_usage_routing_index(value: Option) -> Option { + value.filter(|value| *value <= MAX_USAGE_CANDIDATE_INDEX) +} + +fn sanitize_usage_routing_string( + value: Option, + max_len: usize, + allow_spaces: bool, +) -> Option { + let value = value?; + let value = value.trim(); + if value.is_empty() || value.len() > max_len { + return None; + } + if !value.bytes().all(|byte| { + byte.is_ascii_alphanumeric() || b"._:/@+-".contains(&byte) || (allow_spaces && byte == b' ') + }) { + return None; + } + Some(value.to_string()) +} + +pub(crate) fn sanitize_usage_error_category(value: Option) -> Option { + let value = value?.trim().to_ascii_lowercase(); + let category = match value.as_str() { + "auth" + | "cancelled" + | "client_error" + | "http_error" + | "non_success_status" + | "provider_error" + | "rate_limit" + | "redirect" + | "server_error" + | "stream_missing_terminal_event" + | "stream_terminal_error" + | "upstream_error" => value, + "" => return None, + _ => "other_error".to_string(), + }; + Some(category) +} + +/// Return the bounded error category represented by an HTTP status code. +/// +/// Stale-request cleanup may have only a candidate status code available. Do +/// not persist provider-supplied diagnostic text in that case; derive one of +/// the same fixed categories used by the usage writer instead. +pub fn usage_error_category_for_status_code(status_code: u16) -> &'static str { + if status_code >= 500 { + "server_error" + } else if status_code >= 400 { + "client_error" + } else if status_code >= 300 { + "redirect" + } else { + "non_success_status" + } +} + pub fn provider_api_key_usage_is_success( status: &str, status_code: Option, @@ -339,6 +626,361 @@ fn newer_last_used_at(before: Option, after: Option) -> Option { mod tests { use serde_json::json; + use super::{ + sanitize_usage_capture_controls_for_persistence, sanitize_usage_error_category, + sanitize_usage_for_persistence, usage_error_category_for_status_code, + usage_lifecycle_update_allowed, + }; + use crate::repository::usage::{UpsertUsageRecord, UsageBodyCaptureState}; + + fn usage_with_http_capture() -> UpsertUsageRecord { + UpsertUsageRecord { + request_id: "req-sensitive-capture".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + username: Some("alice".to_string()), + api_key_name: Some("primary".to_string()), + provider_name: "OpenAI".to_string(), + model: "gpt-5".to_string(), + target_model: None, + provider_id: Some("provider-1".to_string()), + provider_endpoint_id: None, + provider_api_key_id: None, + request_type: Some("chat".to_string()), + api_format: Some("openai:chat".to_string()), + api_family: Some("openai".to_string()), + endpoint_kind: Some("chat".to_string()), + endpoint_api_format: Some("openai:chat".to_string()), + provider_api_family: Some("openai".to_string()), + provider_endpoint_kind: Some("chat".to_string()), + has_format_conversion: Some(false), + is_stream: Some(false), + input_tokens: Some(1), + output_tokens: Some(2), + total_tokens: Some(3), + cache_creation_input_tokens: None, + cache_creation_ephemeral_5m_input_tokens: None, + cache_creation_ephemeral_1h_input_tokens: None, + cache_read_input_tokens: None, + cache_creation_cost_usd: None, + cache_read_cost_usd: None, + output_price_per_1m: None, + total_cost_usd: Some(0.01), + actual_total_cost_usd: Some(0.01), + status_code: Some(200), + error_message: Some("Bearer secret".to_string()), + error_category: Some("provider_error".to_string()), + response_time_ms: Some(10), + first_byte_time_ms: Some(5), + status: "completed".to_string(), + billing_status: "settled".to_string(), + request_headers: Some(json!({"authorization": "Bearer secret"})), + request_body: Some(json!({"prompt": "private"})), + request_body_ref: Some("usage://request/body".to_string()), + request_body_state: Some(UsageBodyCaptureState::Reference), + provider_request_headers: Some(json!({"x-api-key": "secret"})), + provider_request_body: Some(json!({"prompt": "private"})), + provider_request_body_ref: Some("usage://provider/body".to_string()), + provider_request_body_state: Some(UsageBodyCaptureState::Inline), + response_headers: Some(json!({"set-cookie": "secret"})), + response_body: Some(json!({"output": "private"})), + response_body_ref: Some("usage://response/body".to_string()), + response_body_state: Some(UsageBodyCaptureState::Truncated), + client_response_headers: Some(json!({"x-private": "secret"})), + client_response_body: Some(json!({"output": "private"})), + client_response_body_ref: Some("usage://client/body".to_string()), + client_response_body_state: Some(UsageBodyCaptureState::Disabled), + candidate_id: Some("candidate-1".to_string()), + candidate_index: Some(0), + key_name: None, + planner_kind: None, + route_family: None, + route_kind: None, + execution_path: None, + local_execution_runtime_miss_reason: None, + request_metadata: Some(json!({"client_ip": "203.0.113.8"})), + finalized_at_unix_secs: Some(2), + created_at_unix_ms: Some(1_000), + updated_at_unix_secs: 2, + } + } + + #[test] + fn usage_error_categories_are_bounded_to_controlled_values() { + assert_eq!( + sanitize_usage_error_category(Some(" Server_Error ".to_string())).as_deref(), + Some("server_error") + ); + assert_eq!( + sanitize_usage_error_category(Some("Authorization: Bearer secret".to_string())) + .as_deref(), + Some("other_error") + ); + assert_eq!(sanitize_usage_error_category(Some(" ".to_string())), None); + } + + #[test] + fn status_codes_map_to_bounded_error_categories() { + assert_eq!(usage_error_category_for_status_code(599), "server_error"); + assert_eq!(usage_error_category_for_status_code(400), "client_error"); + assert_eq!(usage_error_category_for_status_code(302), "redirect"); + assert_eq!( + usage_error_category_for_status_code(200), + "non_success_status" + ); + } + + #[test] + fn persistence_derives_missing_failed_category_from_status_code() { + let mut usage = usage_with_http_capture(); + usage.status = "failed".to_string(); + usage.status_code = Some(429); + usage.error_category = None; + + let usage = sanitize_usage_for_persistence(usage); + assert_eq!(usage.error_category.as_deref(), Some("client_error")); + } + + #[test] + fn persistence_boundary_drops_all_http_capture_material() { + let usage = sanitize_usage_for_persistence(usage_with_http_capture()); + + assert!(usage.username.is_none()); + assert!(usage.api_key_name.is_none()); + assert!(usage.error_message.is_none()); + assert!(usage.request_headers.is_none()); + assert!(usage.request_body.is_none()); + assert!(usage.request_body_ref.is_none()); + assert!(usage.request_body_state.is_none()); + assert!(usage.provider_request_headers.is_none()); + assert!(usage.provider_request_body.is_none()); + assert!(usage.provider_request_body_ref.is_none()); + assert!(usage.provider_request_body_state.is_none()); + assert!(usage.response_headers.is_none()); + assert!(usage.response_body.is_none()); + assert!(usage.response_body_ref.is_none()); + assert!(usage.response_body_state.is_none()); + assert!(usage.client_response_headers.is_none()); + assert!(usage.client_response_body.is_none()); + assert!(usage.client_response_body_ref.is_none()); + assert!(usage.client_response_body_state.is_none()); + assert_eq!(usage.error_category.as_deref(), Some("provider_error")); + assert_eq!(usage.candidate_id.as_deref(), Some("candidate-1")); + assert_eq!( + usage.request_metadata, + Some(json!({"client_ip": "203.0.113.8"})) + ); + } + + #[test] + fn auxiliary_capture_projection_keeps_only_explicit_clear_tombstones() { + let mut input = usage_with_http_capture(); + input.request_body_state = Some(UsageBodyCaptureState::None); + input.response_body_state = Some(UsageBodyCaptureState::Disabled); + + let usage = sanitize_usage_capture_controls_for_persistence(input); + + assert!(usage.request_headers.is_none()); + assert!(usage.request_body.is_none()); + assert!(usage.request_body_ref.is_none()); + assert_eq!(usage.request_body_state, Some(UsageBodyCaptureState::None)); + assert!(usage.provider_request_body_state.is_none()); + assert!(usage.response_body_state.is_none()); + assert!(usage.client_response_body_state.is_none()); + } + + #[test] + fn auxiliary_projection_promotes_only_bounded_routing_metadata() { + let mut input = usage_with_http_capture(); + input.candidate_id = None; + input.candidate_index = None; + input.key_name = None; + input.planner_kind = None; + input.route_family = None; + input.route_kind = None; + input.execution_path = None; + input.local_execution_runtime_miss_reason = None; + input.request_metadata = Some(json!({ + "trace_id": "trace", + "candidate_id": "candidate-from-metadata", + "candidate_index": 7, + "key_name": "primary key", + "planner_kind": "fallback", + "route_family": "chat", + "route_kind": "remote", + "execution_path": "execution_runtime_stream", + "local_execution_runtime_miss_reason": "runtime_busy", + "authorization": "Bearer should-not-persist", + })); + + let usage = sanitize_usage_capture_controls_for_persistence(input); + + assert_eq!( + usage.candidate_id.as_deref(), + Some("candidate-from-metadata") + ); + assert_eq!(usage.candidate_index, Some(7)); + assert_eq!(usage.key_name.as_deref(), Some("primary key")); + assert_eq!(usage.planner_kind.as_deref(), Some("fallback")); + assert_eq!(usage.route_family.as_deref(), Some("chat")); + assert_eq!(usage.route_kind.as_deref(), Some("remote")); + assert_eq!( + usage.execution_path.as_deref(), + Some("execution_runtime_stream") + ); + assert_eq!( + usage.local_execution_runtime_miss_reason.as_deref(), + Some("runtime_busy") + ); + assert_eq!(usage.request_metadata, Some(json!({"trace_id": "trace"}))); + } + + #[test] + fn routing_projection_rejects_unbounded_or_control_character_values() { + let mut input = usage_with_http_capture(); + input.candidate_id = Some("candidate\nforged".to_string()); + input.candidate_index = Some(u64::MAX); + input.key_name = Some("key\0name".to_string()); + input.planner_kind = Some("p".repeat(65)); + input.route_family = Some("route\tname".to_string()); + input.route_kind = Some("route-kind".to_string()); + input.execution_path = Some("execution-path".to_string()); + input.local_execution_runtime_miss_reason = Some("m".repeat(121)); + + let usage = sanitize_usage_for_persistence(input); + + assert!(usage.candidate_id.is_none()); + assert!(usage.candidate_index.is_none()); + assert!(usage.key_name.is_none()); + assert!(usage.planner_kind.is_none()); + assert!(usage.route_family.is_none()); + assert_eq!(usage.route_kind.as_deref(), Some("route-kind")); + assert_eq!(usage.execution_path.as_deref(), Some("execution-path")); + assert!(usage.local_execution_runtime_miss_reason.is_none()); + } + + #[test] + fn invalid_typed_routing_values_do_not_fall_back_to_metadata() { + let mut input = usage_with_http_capture(); + input.candidate_id = Some("candidate\nforged".to_string()); + input.candidate_index = Some(u64::MAX); + input.key_name = Some("key\0name".to_string()); + input.planner_kind = Some("p".repeat(65)); + input.route_family = Some("route\tname".to_string()); + input.route_kind = Some("route\nkind".to_string()); + input.execution_path = Some("execution\npath".to_string()); + input.local_execution_runtime_miss_reason = Some("m".repeat(121)); + input.request_metadata = Some(json!({ + "candidate_id": "metadata-candidate", + "candidate_index": 7, + "key_name": "metadata key", + "planner_kind": "metadata-planner", + "route_family": "metadata-family", + "route_kind": "metadata-kind", + "execution_path": "metadata-path", + "local_execution_runtime_miss_reason": "metadata-reason", + })); + + let usage = sanitize_usage_capture_controls_for_persistence(input); + + assert!(usage.candidate_id.is_none()); + assert!(usage.candidate_index.is_none()); + assert!(usage.key_name.is_none()); + assert!(usage.planner_kind.is_none()); + assert!(usage.route_family.is_none()); + assert!(usage.route_kind.is_none()); + assert!(usage.execution_path.is_none()); + assert!(usage.local_execution_runtime_miss_reason.is_none()); + } + + #[test] + fn lifecycle_order_rejects_stale_and_equal_conflicting_terminal_events() { + assert!(!usage_lifecycle_update_allowed( + "completed", + "pending", + 20, + Some(20), + "failed", + "void", + 19, + Some(19), + )); + assert!(!usage_lifecycle_update_allowed( + "completed", + "pending", + 20, + Some(20), + "failed", + "void", + 20, + Some(20), + )); + assert!(!usage_lifecycle_update_allowed( + "completed", + "pending", + 20, + Some(20), + "completed", + "settled", + 20, + Some(20), + )); + assert!(usage_lifecycle_update_allowed( + "completed", + "pending", + 20, + Some(20), + "failed", + "void", + 21, + Some(21), + )); + } + + #[test] + fn lifecycle_order_allows_same_second_progress_and_fresh_void_recovery() { + assert!(usage_lifecycle_update_allowed( + "pending", + "pending", + 20, + None, + "streaming", + "pending", + 20, + None, + )); + assert!(usage_lifecycle_update_allowed( + "streaming", + "pending", + 20, + None, + "completed", + "pending", + 20, + Some(20), + )); + assert!(usage_lifecycle_update_allowed( + "failed", + "void", + 20, + Some(20), + "completed", + "pending", + 20, + Some(20), + )); + assert!(!usage_lifecycle_update_allowed( + "failed", + "void", + 20, + Some(21), + "completed", + "pending", + 20, + Some(20), + )); + } + use super::{api_key_usage_contribution, provider_api_key_usage_contribution}; use crate::repository::usage::StoredRequestUsageAudit; diff --git a/crates/aether-data/contracts/src/repository/usage/types.rs b/crates/aether-data/contracts/src/repository/usage/types.rs index f0f1cea61..dd7e8e65d 100644 --- a/crates/aether-data/contracts/src/repository/usage/types.rs +++ b/crates/aether-data/contracts/src/repository/usage/types.rs @@ -11,6 +11,7 @@ pub const ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY: &str = "routing_candidate_ pub const ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY: &str = "routing_failure_diagnostic"; pub const WEBSOCKET_MODE_METADATA_KEY: &str = "websocket_mode"; pub const WEBSOCKET_TRANSPORT_METADATA_KEY: &str = "websocket_transport"; +pub const PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY: &str = "plan_usage_reservation_deferred"; /// Whether token/cost usage is authoritative for this audit row. /// /// The field is absent for legacy and normally-metered requests. An explicit @@ -386,7 +387,7 @@ impl StoredRequestUsageAudit { total_cost_usd: f64, actual_total_cost_usd: f64, status_code: Option, - error_message: Option, + _error_message: Option, error_category: Option, response_time_ms: Option, first_byte_time_ms: Option, @@ -421,14 +422,14 @@ impl StoredRequestUsageAudit { "usage.billing_status is empty".to_string(), )); } - if !total_cost_usd.is_finite() { + if !total_cost_usd.is_finite() || total_cost_usd < 0.0 { return Err(crate::DataLayerError::UnexpectedValue( - "usage.total_cost_usd is not finite".to_string(), + "usage.total_cost_usd must be finite and non-negative".to_string(), )); } - if !actual_total_cost_usd.is_finite() { + if !actual_total_cost_usd.is_finite() || actual_total_cost_usd < 0.0 { return Err(crate::DataLayerError::UnexpectedValue( - "usage.actual_total_cost_usd is not finite".to_string(), + "usage.actual_total_cost_usd must be finite and non-negative".to_string(), )); } @@ -468,8 +469,8 @@ impl StoredRequestUsageAudit { total_cost_usd, actual_total_cost_usd, status_code: parse_u16(status_code, "usage.status_code")?, - error_message, - error_category, + error_message: None, + error_category: super::policy::sanitize_usage_error_category(error_category), response_time_ms: parse_optional_u64(response_time_ms, "usage.response_time_ms")?, first_byte_time_ms: parse_optional_u64(first_byte_time_ms, "usage.first_byte_time_ms")?, status, @@ -1721,6 +1722,16 @@ pub fn parse_usage_body_ref(body_ref: &str) -> Option<(String, UsageBodyField)> )) } +pub fn canonical_usage_body_ref_for( + body_ref: &str, + expected_request_id: &str, + expected_field: UsageBodyField, +) -> Option { + parse_usage_body_ref(body_ref) + .filter(|(request_id, field)| request_id == expected_request_id && *field == expected_field) + .map(|(request_id, field)| usage_body_ref(&request_id, field)) +} + #[async_trait] pub trait UsageReadRepository: Send + Sync { async fn find_by_id( @@ -2043,48 +2054,58 @@ impl UpsertUsageRecord { "usage upsert model cannot be empty".to_string(), )); } - if self.status.trim().is_empty() { - return Err(crate::DataLayerError::InvalidInput( - "usage upsert status cannot be empty".to_string(), - )); + if !matches!( + self.status.as_str(), + "pending" | "streaming" | "completed" | "failed" | "cancelled" + ) { + return Err(crate::DataLayerError::InvalidInput(format!( + "invalid usage upsert status: {}", + self.status + ))); } - if self.billing_status.trim().is_empty() { - return Err(crate::DataLayerError::InvalidInput( - "usage upsert billing_status cannot be empty".to_string(), - )); + if !matches!( + self.billing_status.as_str(), + "pending" | "settled" | "void" | "insufficient_quota" + ) { + return Err(crate::DataLayerError::InvalidInput(format!( + "invalid usage upsert billing_status: {}", + self.billing_status + ))); } if let Some(value) = self.total_cost_usd { - if !value.is_finite() { + if !value.is_finite() || value < 0.0 { return Err(crate::DataLayerError::InvalidInput( - "usage upsert total_cost_usd must be finite".to_string(), + "usage upsert total_cost_usd must be finite and non-negative".to_string(), )); } } if let Some(value) = self.cache_creation_cost_usd { - if !value.is_finite() { + if !value.is_finite() || value < 0.0 { return Err(crate::DataLayerError::InvalidInput( - "usage upsert cache_creation_cost_usd must be finite".to_string(), + "usage upsert cache_creation_cost_usd must be finite and non-negative" + .to_string(), )); } } if let Some(value) = self.cache_read_cost_usd { - if !value.is_finite() { + if !value.is_finite() || value < 0.0 { return Err(crate::DataLayerError::InvalidInput( - "usage upsert cache_read_cost_usd must be finite".to_string(), + "usage upsert cache_read_cost_usd must be finite and non-negative".to_string(), )); } } if let Some(value) = self.output_price_per_1m { - if !value.is_finite() { + if !value.is_finite() || value < 0.0 { return Err(crate::DataLayerError::InvalidInput( - "usage upsert output_price_per_1m must be finite".to_string(), + "usage upsert output_price_per_1m must be finite and non-negative".to_string(), )); } } if let Some(value) = self.actual_total_cost_usd { - if !value.is_finite() { + if !value.is_finite() || value < 0.0 { return Err(crate::DataLayerError::InvalidInput( - "usage upsert actual_total_cost_usd must be finite".to_string(), + "usage upsert actual_total_cost_usd must be finite and non-negative" + .to_string(), )); } } @@ -2271,9 +2292,14 @@ pub struct UsageCounterPendingHealthSnapshot { pub pending_by_kind: std::collections::BTreeMap, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct ProxyNodeCounterDelta { pub node_id: String, + /// Incarnation fence captured when the request plan selected this node. + /// Counter writes must never silently rebind to a different incarnation + /// that reused the same node id. + #[serde(default)] + pub expected_tunnel_generation: Option, pub total_requests_delta: i64, pub failed_requests_delta: i64, pub dns_failures_delta: i64, @@ -2332,6 +2358,8 @@ pub struct UsageCleanupSummary { pub header_cleaned: usize, pub keys_cleaned: usize, pub records_deleted: usize, + pub cost_reservations_deleted: usize, + pub request_admissions_deleted: usize, } #[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] @@ -2452,12 +2480,13 @@ fn parse_timestamp(value: i64, field_name: &str) -> Result(value: &Option) -> Option<&'static str> { + value.as_ref().map(|_| "[REDACTED]") +} + +pub const LAST_ACTIVE_ADMIN_UPDATE_DENIED: &str = "last_active_admin_update_denied"; +pub const LAST_ACTIVE_ADMIN_DELETE_DENIED: &str = "last_active_admin_delete_denied"; + +pub fn is_last_active_admin_update_denied(error: &crate::DataLayerError) -> bool { + matches!(error, crate::DataLayerError::InvalidInput(message) if message == LAST_ACTIVE_ADMIN_UPDATE_DENIED) +} + +pub fn is_last_active_admin_delete_denied(error: &crate::DataLayerError) -> bool { + matches!(error, crate::DataLayerError::InvalidInput(message) if message == LAST_ACTIVE_ADMIN_DELETE_DENIED) +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredUserSummary { pub id: String, @@ -47,7 +62,7 @@ impl StoredUserSummary { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredUserAuthRecord { pub id: String, pub email: Option, @@ -64,10 +79,41 @@ pub struct StoredUserAuthRecord { pub allowed_models_mode: String, pub is_active: bool, pub is_deleted: bool, + #[serde(default)] + pub security_version: i64, pub created_at: Option>, pub last_login_at: Option>, } +impl std::fmt::Debug for StoredUserAuthRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredUserAuthRecord") + .field("id", &self.id) + .field("email", &self.email) + .field("email_verified", &self.email_verified) + .field("username", &self.username) + .field( + "password_hash", + &redacted_optional_secret(&self.password_hash), + ) + .field("role", &self.role) + .field("auth_source", &self.auth_source) + .field("allowed_providers", &self.allowed_providers) + .field("allowed_providers_mode", &self.allowed_providers_mode) + .field("allowed_api_formats", &self.allowed_api_formats) + .field("allowed_api_formats_mode", &self.allowed_api_formats_mode) + .field("allowed_models", &self.allowed_models) + .field("allowed_models_mode", &self.allowed_models_mode) + .field("is_active", &self.is_active) + .field("is_deleted", &self.is_deleted) + .field("security_version", &self.security_version) + .field("created_at", &self.created_at) + .field("last_login_at", &self.last_login_at) + .finish() + } +} + impl StoredUserAuthRecord { #[allow(clippy::too_many_arguments)] pub fn new( @@ -126,6 +172,7 @@ impl StoredUserAuthRecord { allowed_models_mode: "unrestricted".to_string(), is_active, is_deleted, + security_version: 0, created_at, last_login_at, }) @@ -149,6 +196,40 @@ impl StoredUserAuthRecord { Ok(self) } + pub fn with_security_version( + mut self, + security_version: i64, + ) -> Result { + if security_version < 0 { + return Err(crate::DataLayerError::UnexpectedValue( + "users.security_version is negative".to_string(), + )); + } + self.security_version = security_version; + Ok(self) + } + + /// Compare the user fields that an aggregate import is allowed to restore. + /// Passwords, security versions, and timestamps are intentionally excluded: + /// password restoration has its own nullable CAS operation and the remaining + /// fields are server-managed concurrency markers. + pub fn matches_restore_state(&self, expected: &Self) -> bool { + self.id == expected.id + && self.email == expected.email + && self.email_verified == expected.email_verified + && self.username == expected.username + && self.role == expected.role + && self.auth_source == expected.auth_source + && self.allowed_providers == expected.allowed_providers + && self.allowed_providers_mode == expected.allowed_providers_mode + && self.allowed_api_formats == expected.allowed_api_formats + && self.allowed_api_formats_mode == expected.allowed_api_formats_mode + && self.allowed_models == expected.allowed_models + && self.allowed_models_mode == expected.allowed_models_mode + && self.is_active == expected.is_active + && self.is_deleted == expected.is_deleted + } + fn with_legacy_policy_modes(mut self) -> Self { self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers); self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats); @@ -174,6 +255,127 @@ pub struct LdapAuthUserProvisioningOutcome { pub created: bool, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DeleteUserOAuthLinkOutcome { + Deleted, + NotFound, + LastOAuthBinding, + LastLoginMethod, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BindUserOAuthLinkOutcome { + Bound, + IdentityAlreadyBoundToUser, + IdentityBoundToAnotherUser, + UserAlreadyLinkedProvider, + UserNotFound, + SessionUnavailable, + ProviderNotFound, + ProviderDisabled, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct BindUserOAuthLinkSessionExpectation { + pub session_id: String, + pub client_device_id: String, + pub security_version: i64, + pub checked_at: DateTime, +} + +impl BindUserOAuthLinkSessionExpectation { + pub fn new( + session_id: impl Into, + client_device_id: impl Into, + security_version: i64, + checked_at: DateTime, + ) -> Result { + let session_id = session_id.into(); + let client_device_id = client_device_id.into(); + if session_id.trim().is_empty() + || session_id != session_id.trim() + || client_device_id.trim().is_empty() + || client_device_id != client_device_id.trim() + || security_version < 0 + { + return Err(crate::DataLayerError::InvalidInput( + "OAuth link session expectation is invalid".to_string(), + )); + } + Ok(Self { + session_id, + client_device_id, + security_version, + checked_at, + }) + } +} + +// Returning the full linked-user record avoids a second repository lookup and +// is the established public contract. Keep the success value inline rather +// than changing every adapter/caller to an allocated box. +#[allow(clippy::large_enum_variant)] +#[derive(Debug, Clone, PartialEq)] +pub enum ResolveOAuthLinkedUserOutcome { + Linked(StoredUserAuthRecord), + NotLinked, + ProviderUnavailable, +} + +pub fn last_oauth_unbind_denial( + auth_source: &str, + password_hash: Option<&str>, + local_password_login_allowed: bool, +) -> Option { + if auth_source.eq_ignore_ascii_case("oauth") { + return Some(DeleteUserOAuthLinkOutcome::LastOAuthBinding); + } + if !auth_source.eq_ignore_ascii_case("local") + || !local_password_login_allowed + || !password_hash.is_some_and(is_valid_bcrypt_hash) + { + return Some(DeleteUserOAuthLinkOutcome::LastLoginMethod); + } + None +} + +pub fn is_valid_bcrypt_hash(value: &str) -> bool { + use base64::Engine as _; + use std::str::FromStr; + + let bytes = value.as_bytes(); + if value.len() != 60 + || !matches!(value.get(0..4), Some("$2a$") | Some("$2b$") | Some("$2y$")) + || !bytes.get(4).is_some_and(u8::is_ascii_digit) + || !bytes.get(5).is_some_and(u8::is_ascii_digit) + || bytes.get(6) != Some(&b'$') + { + return false; + } + let Ok(parts) = bcrypt::HashParts::from_str(value) else { + return false; + }; + if !(4..=31).contains(&parts.get_cost()) { + return false; + } + let Some(payload) = value.get(7..) else { + return false; + }; + if !payload.bytes().all(is_bcrypt_base64_byte) { + return false; + } + bcrypt::BASE_64 + .decode(&payload[..22]) + .is_ok_and(|salt| salt.len() == 16) + && bcrypt::BASE_64 + .decode(&payload[22..]) + .is_ok_and(|hash| hash.len() == 23) +} + +fn is_bcrypt_base64_byte(value: u8) -> bool { + value.is_ascii_alphanumeric() || matches!(value, b'.' | b'/') +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredUserOAuthLinkSummary { pub provider_type: String, @@ -218,7 +420,7 @@ impl StoredUserOAuthLinkSummary { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredUserExportRow { pub id: String, pub email: Option, @@ -240,6 +442,35 @@ pub struct StoredUserExportRow { pub is_active: bool, } +impl std::fmt::Debug for StoredUserExportRow { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredUserExportRow") + .field("id", &self.id) + .field("email", &self.email) + .field("email_verified", &self.email_verified) + .field("username", &self.username) + .field( + "password_hash", + &redacted_optional_secret(&self.password_hash), + ) + .field("role", &self.role) + .field("auth_source", &self.auth_source) + .field("allowed_providers", &self.allowed_providers) + .field("allowed_providers_mode", &self.allowed_providers_mode) + .field("allowed_api_formats", &self.allowed_api_formats) + .field("allowed_api_formats_mode", &self.allowed_api_formats_mode) + .field("allowed_models", &self.allowed_models) + .field("allowed_models_mode", &self.allowed_models_mode) + .field("rate_limit", &self.rate_limit) + .field("rate_limit_mode", &self.rate_limit_mode) + .field("model_capability_settings", &self.model_capability_settings) + .field("feature_settings", &self.feature_settings) + .field("is_active", &self.is_active) + .finish() + } +} + impl StoredUserExportRow { #[allow(clippy::too_many_arguments)] pub fn new( @@ -324,6 +555,25 @@ impl StoredUserExportRow { Ok(self) } + /// Compare the identity and policy fields that an aggregate import may + /// restore. Passwords, rate/settings payloads, and server timestamps are + /// checked separately by the rollback operation. + pub fn matches_restore_state(&self, expected: &Self) -> bool { + self.id == expected.id + && self.email == expected.email + && self.email_verified == expected.email_verified + && self.username == expected.username + && self.role == expected.role + && self.auth_source == expected.auth_source + && self.allowed_providers == expected.allowed_providers + && self.allowed_providers_mode == expected.allowed_providers_mode + && self.allowed_api_formats == expected.allowed_api_formats + && self.allowed_api_formats_mode == expected.allowed_api_formats_mode + && self.allowed_models == expected.allowed_models + && self.allowed_models_mode == expected.allowed_models_mode + && self.is_active == expected.is_active + } + pub fn with_feature_settings(mut self, feature_settings: Option) -> Self { self.feature_settings = normalize_optional_json(feature_settings); self @@ -342,7 +592,7 @@ impl StoredUserExportRow { } } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredUserSessionRecord { pub id: String, pub user_id: String, @@ -359,6 +609,35 @@ pub struct StoredUserSessionRecord { pub user_agent: Option, pub created_at: Option>, pub updated_at: Option>, + #[serde(skip)] + pub security_version: i64, +} + +impl std::fmt::Debug for StoredUserSessionRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredUserSessionRecord") + .field("id", &self.id) + .field("user_id", &self.user_id) + .field("client_device_id", &self.client_device_id) + .field("device_label", &self.device_label) + .field("refresh_token_hash", &"[REDACTED]") + .field( + "prev_refresh_token_hash", + &redacted_optional_secret(&self.prev_refresh_token_hash), + ) + .field("rotated_at", &self.rotated_at) + .field("last_seen_at", &self.last_seen_at) + .field("expires_at", &self.expires_at) + .field("revoked_at", &self.revoked_at) + .field("revoke_reason", &self.revoke_reason) + .field("ip_address", &self.ip_address) + .field("user_agent", &self.user_agent) + .field("created_at", &self.created_at) + .field("updated_at", &self.updated_at) + .field("security_version", &self.security_version) + .finish() + } } impl StoredUserSessionRecord { @@ -420,9 +699,23 @@ impl StoredUserSessionRecord { user_agent, created_at, updated_at, + security_version: 0, }) } + pub fn with_security_version( + mut self, + security_version: i64, + ) -> Result { + if security_version < 0 { + return Err(crate::DataLayerError::UnexpectedValue( + "user_sessions.security_version is negative".to_string(), + )); + } + self.security_version = security_version; + Ok(self) + } + pub fn hash_refresh_token(token: &str) -> String { use sha2::Digest; @@ -442,8 +735,10 @@ impl StoredUserSessionRecord { let Some(rotated_at) = self.rotated_at else { return (false, false); }; + let age = now.signed_duration_since(rotated_at); if prev_hash == &token_hash - && now.signed_duration_since(rotated_at).num_seconds() <= Self::REFRESH_GRACE_SECONDS + && age >= chrono::Duration::zero() + && age <= chrono::Duration::seconds(Self::REFRESH_GRACE_SECONDS) { return (true, true); } @@ -744,6 +1039,22 @@ pub trait UserReadRepository: Send + Sync { record: UpsertUserGroupRecord, ) -> Result, crate::DataLayerError>; + /// Restore an existing user group only when its complete stored snapshot + /// still equals `expected`. The compare and replacement must be atomic so + /// an import rollback cannot overwrite a concurrent administrator update. + /// `false` means the row is missing, the identities differ, or the current + /// snapshot no longer matches the expected post-import state. + async fn restore_user_group_if_matches( + &self, + expected: &StoredUserGroup, + restored: &StoredUserGroup, + ) -> Result { + let _ = (expected, restored); + Err(crate::DataLayerError::InvalidInput( + "atomic user group restore is not available".to_string(), + )) + } + async fn delete_user_group(&self, group_id: &str) -> Result; async fn list_user_group_members( @@ -773,6 +1084,20 @@ pub trait UserReadRepository: Send + Sync { group_ids: &[String], ) -> Result, crate::DataLayerError>; + /// Restore a user's group memberships only when the current set still + /// equals the post-import set. The compare and replacement are atomic. + async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + let _ = (user_id, expected_group_ids, restored_group_ids); + Err(crate::DataLayerError::InvalidInput( + "atomic user group restore is not available".to_string(), + )) + } + async fn add_user_to_group( &self, group_id: &str, @@ -824,6 +1149,19 @@ pub trait UserReadRepository: Send + Sync { provider_user_id: &str, ) -> Result, crate::DataLayerError>; + #[allow(clippy::too_many_arguments)] + async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: DateTime, + provider_enabled_snapshot: bool, + ) -> Result; + async fn touch_oauth_link( &self, provider_type: &str, @@ -837,6 +1175,7 @@ pub trait UserReadRepository: Send + Sync { async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: DateTime, ) -> Result, crate::DataLayerError>; @@ -855,8 +1194,22 @@ pub trait UserReadRepository: Send + Sync { async fn count_user_oauth_links(&self, user_id: &str) -> Result; + async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result; + + async fn count_locked_users_if_oauth_provider_disabled( + &self, + _provider_type: &str, + _enabled_provider_types_snapshot: &[String], + _ldap_exclusive: bool, + ) -> Result { + Ok(0) + } + #[allow(clippy::too_many_arguments)] - async fn upsert_user_oauth_link( + async fn bind_user_oauth_link( &self, user_id: &str, provider_type: &str, @@ -865,13 +1218,51 @@ pub trait UserReadRepository: Send + Sync { provider_email: Option<&str>, extra_data: Option, linked_at: DateTime, - ) -> Result<(), crate::DataLayerError>; + ) -> Result { + self.bind_user_oauth_link_if_provider_enabled( + user_id, + provider_type, + provider_user_id, + provider_username, + provider_email, + extra_data, + linked_at, + true, + None, + ) + .await + } + #[allow(clippy::too_many_arguments)] + async fn bind_user_oauth_link_if_provider_enabled( + &self, + user_id: &str, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + linked_at: DateTime, + provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result; + + /// Marks an email as verified only if the user's current normalized email still matches. + async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + verified_at: DateTime, + ) -> Result; + + /// Atomically deletes the requested link only when doing so leaves a valid login method. async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result; + local_password_login_allowed: bool, + enabled_provider_types_snapshot: &[String], + ) -> Result; async fn get_or_create_ldap_auth_user( &self, @@ -891,10 +1282,44 @@ pub trait UserReadRepository: Send + Sync { async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, crate::DataLayerError>; + /// Restore all user fields represented by an aggregate import checkpoint in + /// one compare-and-write transaction. The password hash is deliberately + /// excluded and must be restored through the nullable password CAS method. + /// Implementations must not write anything when the current state differs + /// from `expected_auth` (or the exported rate/settings fields). + #[allow(clippy::too_many_arguments)] + async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&Value>, + restored_feature_settings: Option, + ) -> Result { + let _ = ( + expected_auth, + restored_auth, + expected_export, + restored_export, + expected_model_capability_settings, + restored_model_capability_settings, + expected_feature_settings, + restored_feature_settings, + ); + Err(crate::DataLayerError::InvalidInput( + "atomic user state restore is not available".to_string(), + )) + } + async fn update_local_auth_user_password_hash( &self, user_id: &str, @@ -902,6 +1327,37 @@ pub trait UserReadRepository: Send + Sync { updated_at: DateTime, ) -> Result, crate::DataLayerError>; + /// Replace a user's password hash, including an explicit `NULL`, only when the current hash + /// still equals `expected_password_hash`. The compare and write must be atomic. + async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + updated_at: DateTime, + ) -> Result { + let _ = (user_id, expected_password_hash, password_hash, updated_at); + Err(crate::DataLayerError::InvalidInput( + "atomic nullable password restore is not available".to_string(), + )) + } + + async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: DateTime, + ) -> Result; + + async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: DateTime, + ) -> Result; + #[allow(clippy::too_many_arguments)] async fn update_local_auth_user_admin_fields( &self, @@ -963,6 +1419,23 @@ pub trait UserReadRepository: Send + Sync { async fn delete_local_auth_user(&self, user_id: &str) -> Result; + /// Delete a local-auth user only when no wallet currently belongs to the user. + /// + /// Authentication provisioning compensation uses this guard after removing a + /// wallet it can prove ownership of. Implementations must evaluate the + /// wallet absence check in the same database transaction as the user delete; + /// a default implementation is deliberately fail-closed for repositories + /// that cannot provide that atomicity. + async fn delete_local_auth_user_if_wallet_absent( + &self, + user_id: &str, + ) -> Result { + let _ = user_id; + Err(crate::DataLayerError::InvalidInput( + "atomic user deletion without a wallet is not available".to_string(), + )) + } + async fn read_user_preferences( &self, user_id: &str, @@ -989,6 +1462,12 @@ pub trait UserReadRepository: Send + Sync { session: &StoredUserSessionRecord, ) -> Result, crate::DataLayerError>; + async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, crate::DataLayerError>; + async fn touch_user_session( &self, user_id: &str, @@ -1011,7 +1490,7 @@ pub trait UserReadRepository: Send + Sync { &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: DateTime, expires_at: DateTime, @@ -1104,7 +1583,9 @@ fn parse_string_list_value( field_name: &str, ) -> Result>, crate::DataLayerError> { match value { - Value::Null => Ok(None), + Value::Null => Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains JSON null; use SQL NULL for an unset policy" + ))), Value::Array(array) => parse_string_list_array(array, field_name).map(Some), Value::String(raw) => parse_embedded_string_list(raw, field_name), _ => Err(crate::DataLayerError::UnexpectedValue(format!( @@ -1118,8 +1599,15 @@ fn parse_embedded_string_list( field_name: &str, ) -> Result>, crate::DataLayerError> { let raw = raw.trim(); - if raw.is_empty() || raw.eq_ignore_ascii_case("null") { - return Ok(None); + if raw.is_empty() { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty string" + ))); + } + if raw.eq_ignore_ascii_case("null") { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains stringified JSON null; use SQL NULL for an unset policy" + ))); } if let Ok(decoded) = serde_json::from_str::(raw) { @@ -1141,9 +1629,12 @@ fn parse_string_list_array( ))); }; let item = item.trim(); - if !item.is_empty() { - items.push(item.to_string()); + if item.is_empty() { + return Err(crate::DataLayerError::UnexpectedValue(format!( + "{field_name} contains an empty item" + ))); } + items.push(item.to_string()); } Ok(items) } @@ -1154,10 +1645,54 @@ mod tests { use serde_json::Value; use super::{ - legacy_list_policy_mode, StoredUserAuthRecord, StoredUserExportRow, + is_valid_bcrypt_hash, last_oauth_unbind_denial, legacy_list_policy_mode, + DeleteUserOAuthLinkOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserPreferenceRecord, StoredUserSessionRecord, }; + #[test] + fn classifies_last_oauth_unbind_by_remaining_login_method() { + let valid_hash = + bcrypt::hash("Secret123!", bcrypt::DEFAULT_COST).expect("bcrypt fixture should hash"); + + assert_eq!( + last_oauth_unbind_denial("oauth", None, false), + Some(DeleteUserOAuthLinkOutcome::LastOAuthBinding) + ); + assert_eq!( + last_oauth_unbind_denial("local", None, true), + Some(DeleteUserOAuthLinkOutcome::LastLoginMethod) + ); + assert_eq!( + last_oauth_unbind_denial("local", Some("not-a-password-hash"), true), + Some(DeleteUserOAuthLinkOutcome::LastLoginMethod) + ); + assert_eq!( + last_oauth_unbind_denial("local", Some(&valid_hash), false), + Some(DeleteUserOAuthLinkOutcome::LastLoginMethod) + ); + assert_eq!( + last_oauth_unbind_denial("local", Some(&valid_hash), true), + None + ); + } + + #[test] + fn validates_complete_bcrypt_encoding_and_cost() { + let valid_hash = + bcrypt::hash("Secret123!", bcrypt::DEFAULT_COST).expect("bcrypt fixture should hash"); + assert!(is_valid_bcrypt_hash(&valid_hash)); + + for invalid in [ + format!("$2b$99${}", &valid_hash[7..]), + format!("$2x$12${}", &valid_hash[7..]), + format!("$2b$12${}", "!".repeat(53)), + valid_hash[..59].to_string(), + ] { + assert!(!is_valid_bcrypt_hash(&invalid), "accepted {invalid}"); + } + } + #[test] fn builds_user_export_row_with_allowed_lists() { let row = StoredUserExportRow::new( @@ -1203,7 +1738,7 @@ mod tests { "user".to_string(), "local".to_string(), Some(serde_json::json!("[\"openai\"]")), - Some(serde_json::json!("null")), + None, Some(serde_json::json!("gpt-4.1")), None, Some(Value::Null), @@ -1217,6 +1752,29 @@ mod tests { assert_eq!(row.model_capability_settings, None); } + #[test] + fn stored_user_security_lists_distinguish_sql_null_from_malformed_json_null() { + assert_eq!( + super::parse_string_list(None, "users.allowed_providers") + .expect("SQL NULL should remain an unset policy"), + None + ); + assert!( + super::parse_string_list(Some(serde_json::Value::Null), "users.allowed_providers") + .is_err() + ); + assert!(super::parse_string_list( + Some(serde_json::json!("null")), + "users.allowed_providers" + ) + .is_err()); + assert!(super::parse_string_list( + Some(serde_json::json!([" "])), + "users.allowed_providers" + ) + .is_err()); + } + #[test] fn rejects_object_allowed_providers_for_user_export_row() { let result = StoredUserExportRow::new( @@ -1266,6 +1824,73 @@ mod tests { assert_eq!(row.allowed_models, Some(vec!["gpt-4.1".to_string()])); } + #[test] + fn user_record_debug_output_redacts_password_and_refresh_hashes() { + let password_hash = "debug-secret-password-hash"; + let auth = StoredUserAuthRecord::new( + "user-debug".to_string(), + Some("debug@example.com".to_string()), + true, + "debug-user".to_string(), + Some(password_hash.to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + None, + None, + ) + .expect("auth record should build"); + let export = StoredUserExportRow::new( + "user-debug".to_string(), + Some("debug@example.com".to_string()), + true, + "debug-user".to_string(), + Some(password_hash.to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + None, + None, + true, + ) + .expect("export record should build"); + let current_hash = "debug-secret-current-refresh-hash"; + let previous_hash = "debug-secret-previous-refresh-hash"; + let session = StoredUserSessionRecord::new( + "session-debug".to_string(), + "user-debug".to_string(), + "device-debug".to_string(), + None, + current_hash.to_string(), + Some(previous_hash.to_string()), + None, + None, + None, + None, + None, + None, + None, + None, + None, + ) + .expect("session should build"); + + for rendered in [format!("{auth:?}"), format!("{export:?}")] { + assert!(!rendered.contains(password_hash)); + assert!(rendered.contains("[REDACTED]")); + } + let rendered = format!("{session:?}"); + assert!(!rendered.contains(current_hash)); + assert!(!rendered.contains(previous_hash)); + assert!(rendered.contains("[REDACTED]")); + } + #[test] fn legacy_policy_mode_treats_empty_lists_as_unrestricted() { assert_eq!(legacy_list_policy_mode(&None), "unrestricted"); @@ -1306,6 +1931,29 @@ mod tests { session.verify_refresh_token("current-token", now), (true, false) ); + + let future_rotation = StoredUserSessionRecord::new( + "session-2".to_string(), + "user-1".to_string(), + "device-1".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("current-token"), + Some(StoredUserSessionRecord::hash_refresh_token("prev-token")), + Some(now + Duration::seconds(1)), + None, + None, + None, + None, + None, + None, + None, + None, + ) + .expect("session should build"); + assert_eq!( + future_rotation.verify_refresh_token("prev-token", now), + (false, false) + ); } #[test] diff --git a/crates/aether-data/contracts/src/repository/video_tasks/types.rs b/crates/aether-data/contracts/src/repository/video_tasks/types.rs index 759994775..a2d5d2007 100644 --- a/crates/aether-data/contracts/src/repository/video_tasks/types.rs +++ b/crates/aether-data/contracts/src/repository/video_tasks/types.rs @@ -1,6 +1,8 @@ use async_trait::async_trait; use serde_json::Value; +const SAFE_VIDEO_URL_QUERY_KEYS: &[(&str, &str)] = &[("alt", "media")]; + #[derive( Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize, )] @@ -84,6 +86,13 @@ pub struct StoredVideoTask { } impl StoredVideoTask { + pub fn effective_api_format(&self) -> Option<&str> { + effective_video_task_api_format( + self.client_api_format.as_deref(), + self.provider_api_format.as_deref(), + ) + } + #[allow(clippy::too_many_arguments)] pub fn new( id: String, @@ -168,7 +177,7 @@ impl StoredVideoTask { None => None, }; - Ok(Self { + let mut task = Self { id, short_id, request_id, @@ -206,7 +215,84 @@ impl StoredVideoTask { error_message, video_url, request_metadata, - }) + }; + task.sanitize_persisted_diagnostics(); + Ok(task) + } + + fn sanitize_persisted_diagnostics(&mut self) { + self.prompt = None; + self.original_request_body = None; + self.progress_message = None; + self.error_code = sanitize_video_task_error_code(self.error_code.take()); + self.error_message = None; + self.video_url = sanitize_video_task_url( + self.client_api_format.as_deref(), + self.provider_api_format.as_deref(), + self.video_url.take(), + ); + self.request_metadata = None; + } + + /// Verifies that an update still refers to the task identity persisted for `id`. + /// + /// These fields select the owner, upstream target, request shape, or immutable + /// creation identity of a video task. Repository implementations must reject an + /// upsert that changes any of them instead of treating possession of `id` as + /// permission to replace the row. + /// + /// `created_at_unix_ms` is deliberately not compared because some snapshot + /// projections recompute it. Repositories preserve the already stored creation + /// time while accepting an otherwise matching lifecycle update. + pub fn ensure_immutable_identity_matches( + &self, + incoming: &UpsertVideoTask, + ) -> Result<(), crate::DataLayerError> { + let mismatched_field = if self.id != incoming.id { + Some("id") + } else if self.short_id != incoming.short_id { + Some("short_id") + } else if self.request_id != incoming.request_id { + Some("request_id") + } else if self.user_id != incoming.user_id { + Some("user_id") + } else if self.api_key_id != incoming.api_key_id { + Some("api_key_id") + } else if self.external_task_id != incoming.external_task_id { + Some("external_task_id") + } else if self.provider_id != incoming.provider_id { + Some("provider_id") + } else if self.endpoint_id != incoming.endpoint_id { + Some("endpoint_id") + } else if self.key_id != incoming.key_id { + Some("key_id") + } else if self.client_api_format != incoming.client_api_format { + Some("client_api_format") + } else if self.provider_api_format != incoming.provider_api_format { + Some("provider_api_format") + } else if self.format_converted != incoming.format_converted { + Some("format_converted") + } else if self.model != incoming.model { + Some("model") + } else if self.duration_seconds != incoming.duration_seconds { + Some("duration_seconds") + } else if self.resolution != incoming.resolution { + Some("resolution") + } else if self.aspect_ratio != incoming.aspect_ratio { + Some("aspect_ratio") + } else if self.size != incoming.size { + Some("size") + } else { + None + }; + + match mismatched_field { + Some(field) => Err(crate::DataLayerError::InvalidInput(format!( + "video task {} conflicts with persisted immutable field {field}", + incoming.id + ))), + None => Ok(()), + } } } @@ -252,7 +338,24 @@ pub struct UpsertVideoTask { } impl UpsertVideoTask { - pub fn into_stored(self) -> StoredVideoTask { + pub fn sanitize_for_persistence(&mut self) { + self.username = None; + self.api_key_name = None; + self.prompt = None; + self.original_request_body = None; + self.progress_message = None; + self.error_code = sanitize_video_task_error_code(self.error_code.take()); + self.error_message = None; + self.video_url = sanitize_video_task_url( + self.client_api_format.as_deref(), + self.provider_api_format.as_deref(), + self.video_url.take(), + ); + self.request_metadata = None; + } + + pub fn into_stored(mut self) -> StoredVideoTask { + self.sanitize_for_persistence(); StoredVideoTask { id: self.id, short_id: self.short_id, @@ -295,6 +398,75 @@ impl UpsertVideoTask { } } +fn sanitize_video_task_error_code(value: Option) -> Option { + let value = value?.trim().to_ascii_lowercase(); + if value.is_empty() { + return None; + } + Some(match value.as_str() { + "authentication_error" + | "cancelled" + | "content_policy_violation" + | "expired" + | "invalid_request" + | "not_found" + | "permission_denied" + | "poll_permanent_error" + | "poll_timeout" + | "provider_error" + | "rate_limit_exceeded" + | "server_error" + | "unknown" => value, + _ => "provider_error".to_string(), + }) +} + +fn sanitize_video_task_url( + client_api_format: Option<&str>, + provider_api_format: Option<&str>, + value: Option, +) -> Option { + if effective_video_task_api_format(client_api_format, provider_api_format) + != Some("gemini:video") + { + return None; + } + let mut url = url::Url::parse(value?.trim()).ok()?; + if !matches!(url.scheme(), "http" | "https") + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + { + return None; + } + + let query = url + .query_pairs() + .filter(|(key, value)| SAFE_VIDEO_URL_QUERY_KEYS.contains(&(key.as_ref(), value.as_ref()))) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect::>(); + url.set_query(None); + if !query.is_empty() { + url.query_pairs_mut().extend_pairs(query); + } + url.set_fragment(None); + Some(url.into()) +} + +fn effective_video_task_api_format<'a>( + client_api_format: Option<&'a str>, + provider_api_format: Option<&'a str>, +) -> Option<&'a str> { + provider_api_format + .map(str::trim) + .filter(|value| !value.is_empty()) + .or_else(|| { + client_api_format + .map(str::trim) + .filter(|value| !value.is_empty()) + }) +} + impl From for UpsertVideoTask { fn from(task: StoredVideoTask) -> Self { Self { @@ -376,6 +548,15 @@ pub trait VideoTaskReadRepository: Send + Sync { key: VideoTaskLookupKey<'_>, ) -> Result, crate::DataLayerError>; + /// Resolve a public task identifier only when the persisted task belongs + /// to `user_id`. The lookup key and owner predicate must be evaluated by + /// one repository operation. + async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, crate::DataLayerError>; + async fn list_active( &self, limit: usize, @@ -468,7 +649,7 @@ fn coerce_optional_unix_secs( #[cfg(test)] mod tests { - use super::{StoredVideoTask, VideoTaskStatus}; + use super::{StoredVideoTask, UpsertVideoTask, VideoTaskStatus}; #[allow(clippy::type_complexity)] fn base_new_args() -> ( @@ -615,4 +796,203 @@ mod tests { ) .is_err()); } + + #[test] + fn immutable_identity_validation_rejects_every_protected_field() { + let task = UpsertVideoTask { + id: "task-1".to_string(), + short_id: Some("short-1".to_string()), + request_id: "request-1".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("api-key-1".to_string()), + username: None, + api_key_name: None, + external_task_id: Some("external-1".to_string()), + provider_id: Some("provider-1".to_string()), + endpoint_id: Some("endpoint-1".to_string()), + key_id: Some("key-1".to_string()), + client_api_format: Some("openai:video".to_string()), + provider_api_format: Some("gemini:video".to_string()), + format_converted: true, + model: Some("video-model".to_string()), + prompt: None, + original_request_body: None, + duration_seconds: Some(4), + resolution: Some("720p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("1280x720".to_string()), + status: VideoTaskStatus::Submitted, + progress_percent: 0, + progress_message: None, + retry_count: 0, + poll_interval_seconds: 10, + next_poll_at_unix_secs: Some(10), + poll_count: 0, + max_poll_count: 360, + created_at_unix_ms: 1, + submitted_at_unix_secs: Some(1), + completed_at_unix_secs: None, + updated_at_unix_secs: 1, + error_code: None, + error_message: None, + video_url: None, + request_metadata: None, + }; + let stored = task.clone().into_stored(); + + let same_identity_update = UpsertVideoTask { + status: VideoTaskStatus::Completed, + progress_percent: 100, + created_at_unix_ms: 2, + completed_at_unix_secs: Some(2), + updated_at_unix_secs: 2, + ..task.clone() + }; + stored + .ensure_immutable_identity_matches(&same_identity_update) + .expect("mutable state changes should keep the same identity"); + + macro_rules! assert_identity_conflict { + ($field:ident, $value:expr) => {{ + let mut conflicting = task.clone(); + conflicting.$field = $value; + let error = stored + .ensure_immutable_identity_matches(&conflicting) + .expect_err(concat!(stringify!($field), " should be immutable")); + assert!( + error + .to_string() + .contains(concat!("immutable field ", stringify!($field))), + "unexpected error for {}: {error}", + stringify!($field) + ); + }}; + } + + assert_identity_conflict!(id, "task-2".to_string()); + assert_identity_conflict!(short_id, Some("short-2".to_string())); + assert_identity_conflict!(request_id, "request-2".to_string()); + assert_identity_conflict!(user_id, Some("user-2".to_string())); + assert_identity_conflict!(api_key_id, Some("api-key-2".to_string())); + assert_identity_conflict!(external_task_id, Some("external-2".to_string())); + assert_identity_conflict!(provider_id, Some("provider-2".to_string())); + assert_identity_conflict!(endpoint_id, Some("endpoint-2".to_string())); + assert_identity_conflict!(key_id, Some("key-2".to_string())); + assert_identity_conflict!(client_api_format, Some("gemini:video".to_string())); + assert_identity_conflict!(provider_api_format, Some("openai:video".to_string())); + assert_identity_conflict!(format_converted, false); + assert_identity_conflict!(model, Some("other-model".to_string())); + assert_identity_conflict!(duration_seconds, Some(8)); + assert_identity_conflict!(resolution, Some("1080p".to_string())); + assert_identity_conflict!(aspect_ratio, Some("9:16".to_string())); + assert_identity_conflict!(size, Some("1920x1080".to_string())); + } + + #[test] + fn upsert_sanitization_drops_sensitive_diagnostics() { + let mut task = UpsertVideoTask { + id: "task-1".to_string(), + short_id: None, + request_id: "request-1".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("api-key-1".to_string()), + username: Some("private-user-name".to_string()), + api_key_name: Some("private-key-name".to_string()), + external_task_id: Some("upstream-1".to_string()), + provider_id: Some("provider-1".to_string()), + endpoint_id: Some("endpoint-1".to_string()), + key_id: Some("key-1".to_string()), + client_api_format: Some("openai:video".to_string()), + provider_api_format: Some("openai:video".to_string()), + format_converted: false, + model: Some("video-model".to_string()), + prompt: Some("prompt".to_string()), + original_request_body: Some(serde_json::json!({ + "prompt": "private prompt", + "api_key": "secret" + })), + duration_seconds: Some(4), + resolution: Some("720p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("1280x720".to_string()), + status: VideoTaskStatus::Failed, + progress_percent: 100, + progress_message: Some("provider response: secret".to_string()), + retry_count: 1, + poll_interval_seconds: 10, + next_poll_at_unix_secs: None, + poll_count: 2, + max_poll_count: 360, + created_at_unix_ms: 1, + submitted_at_unix_secs: Some(1), + completed_at_unix_secs: Some(2), + updated_at_unix_secs: 2, + error_code: Some("secret provider code".to_string()), + error_message: Some("Authorization: Bearer secret".to_string()), + video_url: Some("https://cdn.example.test/video.mp4?token=secret".to_string()), + request_metadata: Some(serde_json::json!({ + "rust_local_snapshot": {"transport": {"headers": {"authorization": "secret"}}}, + "poll_raw_response": {"error": "secret"} + })), + }; + + task.sanitize_for_persistence(); + + assert_eq!(task.user_id.as_deref(), Some("user-1")); + assert_eq!(task.api_key_id.as_deref(), Some("api-key-1")); + assert_eq!(task.username, None); + assert_eq!(task.api_key_name, None); + assert_eq!(task.original_request_body, None); + assert_eq!(task.progress_message, None); + assert_eq!(task.error_message, None); + assert_eq!(task.error_code.as_deref(), Some("provider_error")); + assert_eq!(task.video_url, None); + assert_eq!(task.request_metadata, None); + assert_eq!(task.prompt, None); + assert_eq!(task.duration_seconds, Some(4)); + } + + #[test] + fn upsert_sanitization_keeps_only_noncredential_video_urls() { + let mut args = base_new_args(); + args.12 = Some("gemini:video".to_string()); + args.15 = Some("private prompt".to_string()); + args.35 = Some( + "https://cdn.example.test/video.mp4?key=secret&alt=media&signature=private#fragment" + .to_string(), + ); + let task = StoredVideoTask::new( + args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9, + args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18, + args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27, + args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36, + ) + .expect("stored task should build"); + assert_eq!(task.prompt, None); + assert_eq!( + task.video_url.as_deref(), + Some("https://cdn.example.test/video.mp4?alt=media") + ); + } + + #[test] + fn upsert_sanitization_uses_client_format_when_legacy_provider_format_is_blank() { + let mut args = base_new_args(); + args.11 = Some("gemini:video".to_string()); + args.12 = Some(" ".to_string()); + args.35 = + Some("https://cdn.example.test/video.mp4?key=secret&alt=media#fragment".to_string()); + let task = StoredVideoTask::new( + args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9, + args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18, + args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27, + args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36, + ) + .expect("legacy stored task should build"); + + assert_eq!( + task.video_url.as_deref(), + Some("https://cdn.example.test/video.mp4?alt=media") + ); + } } diff --git a/crates/aether-data/contracts/src/repository/wallet/snapshot.rs b/crates/aether-data/contracts/src/repository/wallet/snapshot.rs index ea6a2ce94..796a753be 100644 --- a/crates/aether-data/contracts/src/repository/wallet/snapshot.rs +++ b/crates/aether-data/contracts/src/repository/wallet/snapshot.rs @@ -19,6 +19,11 @@ pub struct WalletReadSeed { pub payment_callbacks: Vec, pub wallet_transactions: Vec, pub refunds: Vec, + /// `(user_id, idempotency_key, refund_id)` entries for mutable in-memory + /// repositories. Refund read records intentionally do not expose their + /// idempotency keys, so callers must seed this private write index + /// explicitly when replay behavior matters. + pub refund_idempotency: Vec<(String, String, String)>, pub redeem_batches: Vec, pub redeem_codes: Vec, } @@ -248,7 +253,7 @@ impl WalletReadSnapshot { let effective = if order.status == "pending" && order .expires_at_unix_secs - .is_some_and(|value| value < now_unix_secs) + .is_some_and(|value| value <= now_unix_secs) { "expired" } else { @@ -282,6 +287,14 @@ impl WalletReadSnapshot { .payment_orders .values() .filter(|order| order.user_id.as_deref() == Some(user_id)) + .filter(|order| { + order + .gateway_response + .as_ref() + .and_then(|value| value.get("order_kind")) + .and_then(serde_json::Value::as_str) + != Some("plan_purchase") + }) .cloned() .collect::>(); items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms)); @@ -317,7 +330,15 @@ impl WalletReadSnapshot { ) -> Option { self.payment_orders .get(order_id) - .filter(|order| order.user_id.as_deref() == Some(user_id)) + .filter(|order| { + order.user_id.as_deref() == Some(user_id) + && order + .gateway_response + .as_ref() + .and_then(|value| value.get("order_kind")) + .and_then(serde_json::Value::as_str) + != Some("plan_purchase") + }) .cloned() } @@ -441,3 +462,51 @@ fn page(items: Vec, offset: usize, limit: usize, build: impl Fn(Vec, let total = items.len() as u64; build(items.into_iter().skip(offset).take(limit).collect(), total) } + +#[cfg(test)] +mod tests { + use super::{ + AdminPaymentOrderListQuery, StoredAdminPaymentOrder, WalletReadSeed, WalletReadSnapshot, + }; + + #[test] + fn payment_orders_expire_at_the_exact_boundary() { + let order = StoredAdminPaymentOrder { + id: "order-boundary".to_string(), + order_no: "order-boundary".to_string(), + wallet_id: "wallet-boundary".to_string(), + user_id: Some("user-boundary".to_string()), + amount_usd: 1.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + order_kind: "wallet_recharge".to_string(), + gateway_order_id: None, + gateway_response: None, + status: "pending".to_string(), + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(100), + }; + let snapshot = WalletReadSnapshot::new(WalletReadSeed { + payment_orders: vec![order], + ..WalletReadSeed::default() + }); + let page = snapshot.list_admin_payment_orders( + &AdminPaymentOrderListQuery { + status: Some("expired".to_string()), + payment_method: None, + limit: 10, + offset: 0, + }, + 100, + ); + assert_eq!(page.total, 1); + assert_eq!(page.items[0].id, "order-boundary"); + } +} diff --git a/crates/aether-data/contracts/src/repository/wallet/types.rs b/crates/aether-data/contracts/src/repository/wallet/types.rs index 416f7bd45..25a7e8ea5 100644 --- a/crates/aether-data/contracts/src/repository/wallet/types.rs +++ b/crates/aether-data/contracts/src/repository/wallet/types.rs @@ -1,5 +1,11 @@ use async_trait::async_trait; +const WALLET_REDACTED_DEBUG_VALUE: &str = "[REDACTED]"; + +fn wallet_redacted_debug_option(value: &Option) -> Option<&'static str> { + value.as_ref().map(|_| WALLET_REDACTED_DEBUG_VALUE) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum WalletLookupKey<'a> { WalletId(&'a str), @@ -24,6 +30,18 @@ pub struct StoredWalletSnapshot { pub updated_at_unix_secs: u64, } +/// Result of an idempotent authentication-wallet initialization. +/// +/// Callers that may need to compensate a partially completed operation must +/// know whether the returned row was created by that operation or was already +/// present. Returning this bit from the same atomic repository operation +/// avoids the racy `find -> initialize` ownership inference used previously. +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct InitializeAuthWalletOutcome { + pub wallet: StoredWalletSnapshot, + pub created: bool, +} + impl StoredWalletSnapshot { #[allow(clippy::too_many_arguments)] pub fn new( @@ -353,7 +371,7 @@ pub struct AdminPaymentOrderListQuery { pub offset: usize, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredAdminPaymentOrder { pub id: String, pub order_no: String, @@ -366,6 +384,10 @@ pub struct StoredAdminPaymentOrder { pub refunded_amount_usd: f64, pub refundable_amount_usd: f64, pub payment_method: String, + #[serde(default)] + pub payment_provider: Option, + #[serde(default)] + pub order_kind: String, pub gateway_order_id: Option, pub gateway_response: Option, pub status: String, @@ -375,7 +397,30 @@ pub struct StoredAdminPaymentOrder { pub expires_at_unix_secs: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for StoredAdminPaymentOrder { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredAdminPaymentOrder") + .field("id", &self.id) + .field("order_no", &self.order_no) + .field("wallet_id", &self.wallet_id) + .field("user_id", &self.user_id) + .field("amount_usd", &self.amount_usd) + .field("payment_method", &self.payment_method) + .field("payment_provider", &self.payment_provider) + .field("order_kind", &self.order_kind) + .field("gateway_order_id", &self.gateway_order_id) + .field( + "gateway_response", + &wallet_redacted_debug_option(&self.gateway_response), + ) + .field("status", &self.status) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct AdminWalletPaymentOrderRecord { pub id: String, pub order_no: String, @@ -397,13 +442,34 @@ pub struct AdminWalletPaymentOrderRecord { pub expires_at_unix_secs: Option, } +impl std::fmt::Debug for AdminWalletPaymentOrderRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminWalletPaymentOrderRecord") + .field("id", &self.id) + .field("order_no", &self.order_no) + .field("wallet_id", &self.wallet_id) + .field("user_id", &self.user_id) + .field("amount_usd", &self.amount_usd) + .field("payment_method", &self.payment_method) + .field("gateway_order_id", &self.gateway_order_id) + .field("status", &self.status) + .field( + "gateway_response", + &wallet_redacted_debug_option(&self.gateway_response), + ) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .finish_non_exhaustive() + } +} + #[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)] pub struct StoredAdminPaymentOrderPage { pub items: Vec, pub total: u64, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct StoredAdminPaymentCallback { pub id: String, pub payment_order_id: Option, @@ -420,7 +486,33 @@ pub struct StoredAdminPaymentCallback { pub processed_at_unix_secs: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +impl std::fmt::Debug for StoredAdminPaymentCallback { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("StoredAdminPaymentCallback") + .field("id", &self.id) + .field("payment_order_id", &self.payment_order_id) + .field("payment_method", &self.payment_method) + .field("callback_key", &WALLET_REDACTED_DEBUG_VALUE) + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field( + "payload_hash", + &wallet_redacted_debug_option(&self.payload_hash), + ) + .field("signature_valid", &self.signature_valid) + .field("status", &self.status) + .field("payload", &wallet_redacted_debug_option(&self.payload)) + .field( + "error_message", + &wallet_redacted_debug_option(&self.error_message), + ) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .finish_non_exhaustive() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct AdminPaymentCallbackRecord { pub id: String, pub payment_order_id: Option, @@ -437,6 +529,32 @@ pub struct AdminPaymentCallbackRecord { pub processed_at_unix_secs: Option, } +impl std::fmt::Debug for AdminPaymentCallbackRecord { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdminPaymentCallbackRecord") + .field("id", &self.id) + .field("payment_order_id", &self.payment_order_id) + .field("payment_method", &self.payment_method) + .field("callback_key", &WALLET_REDACTED_DEBUG_VALUE) + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field( + "payload_hash", + &wallet_redacted_debug_option(&self.payload_hash), + ) + .field("signature_valid", &self.signature_valid) + .field("status", &self.status) + .field("payload", &wallet_redacted_debug_option(&self.payload)) + .field( + "error_message", + &wallet_redacted_debug_option(&self.error_message), + ) + .field("created_at_unix_ms", &self.created_at_unix_ms) + .finish_non_exhaustive() + } +} + #[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)] pub struct StoredAdminPaymentCallbackPage { pub items: Vec, @@ -571,6 +689,1144 @@ pub fn redeem_code_payment_method(balance_bucket: &str) -> &'static str { } } +/// Canonicalize the payment-method namespace before it reaches storage. +/// +/// Payment methods are identifiers, not display labels. Keeping one lowercase +/// representation prevents case-only aliases from bypassing gateway-order +/// uniqueness and refund-routing rules across database backends. +pub fn canonicalize_payment_method(payment_method: &str) -> Result { + let normalized = payment_method.trim().to_ascii_lowercase(); + if normalized.is_empty() { + return Err("payment method is required".to_string()); + } + if normalized.chars().count() > 64 { + return Err("payment method exceeds 64 characters".to_string()); + } + if !normalized + .chars() + .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '_' || ch == '-') + { + return Err("payment method contains invalid characters".to_string()); + } + Ok(normalized) +} + +/// Validate the relationship between a stored payment method, provider, and +/// channel. Provider integrations use an explicit namespace so a callback +/// cannot be routed to a different gateway than the one that created the +/// order. EPay is the one compatibility exception: older orders stored the +/// selected `alipay`/`wxpay` channel as the method while using `epay` as the +/// provider. +pub fn validate_payment_provider_channel_binding( + payment_method: &str, + payment_provider: Option<&str>, + payment_channel: Option<&str>, +) -> Result<(), String> { + let payment_method = canonicalize_payment_method(payment_method)?; + let payment_provider = payment_provider + .map(canonicalize_payment_method) + .transpose()?; + let payment_channel = payment_channel + .map(canonicalize_payment_method) + .transpose()?; + + let Some(payment_provider) = payment_provider else { + return Ok(()); + }; + let Some(payment_channel) = payment_channel else { + return Err("payment provider requires a payment channel".to_string()); + }; + + let valid = match payment_provider.as_str() { + "epay" => { + matches!(payment_method.as_str(), "epay" | "alipay" | "wxpay") + && (payment_method == "epay" || payment_channel == payment_method) + } + "alipay" => payment_method == "alipay" && payment_channel == "alipay", + // Keep this list aligned with the direct gateway implementations and + // the public channel resolver. Accepting an arbitrary value here + // would let a repository caller persist a channel that no callback + // verifier or checkout implementation can ever produce. + "wxpay" => { + payment_method == "wxpay" + && matches!(payment_channel.as_str(), "native" | "h5" | "jsapi") + } + "stripe" => { + payment_method == "stripe" + && matches!( + payment_channel.as_str(), + "card" | "alipay" | "wechat_pay" | "link" + ) + } + "admin" => payment_method == "admin_grant" && payment_channel == "manual", + _ => payment_method == payment_provider, + }; + if valid { + Ok(()) + } else { + Err("payment method, provider, and channel binding mismatch".to_string()) + } +} + +/// Validate the wallet-credit entries embedded in a plan entitlement snapshot. +/// +/// A malformed `wallet_credit` must fail the whole fulfillment transaction. +/// Silently filtering it out would leave the entitlement active while the +/// balance promised by the plan was never delivered. Entries for other +/// entitlement kinds are intentionally left to their existing validators. +pub fn validate_plan_wallet_credit_entitlements( + entitlements: &serde_json::Value, +) -> Result<(), String> { + let Some(items) = entitlements.as_array() else { + return Err("plan entitlements must be an array".to_string()); + }; + for (index, item) in items.iter().enumerate() { + let Some(object) = item.as_object() else { + return Err(format!( + "plan entitlement at index {index} must be an object" + )); + }; + let Some(kind) = object.get("type").and_then(serde_json::Value::as_str) else { + return Err(format!("plan entitlement at index {index} is missing type")); + }; + if !kind.eq_ignore_ascii_case("wallet_credit") { + continue; + } + let Some(amount) = object.get("amount_usd").and_then(serde_json::Value::as_f64) else { + return Err(format!( + "wallet_credit.amount_usd is missing at entitlement index {index}" + )); + }; + if !amount.is_finite() || amount <= 0.0 { + return Err(format!( + "wallet_credit.amount_usd is invalid at entitlement index {index}" + )); + } + if let Some(bucket) = object.get("balance_bucket") { + let Some(bucket) = bucket.as_str() else { + return Err(format!( + "wallet_credit.balance_bucket is invalid at entitlement index {index}" + )); + }; + if !matches!( + bucket.trim().to_ascii_lowercase().as_str(), + "recharge" | "gift" + ) { + return Err(format!( + "wallet_credit.balance_bucket is invalid at entitlement index {index}" + )); + } + } + } + Ok(()) +} + +/// Validate the immutable fields of a plan-purchase order before any wallet +/// row or payment order is created. +/// +/// The public handlers already perform most of these checks, but repository +/// methods are also called by administrative workflows and import/recovery +/// code. Keeping the checks here prevents a backend-specific caller from +/// persisting values that another adapter cannot represent (or that later +/// fulfillment would interpret differently). +pub fn validate_plan_purchase_order_input( + input: &CreatePlanPurchaseOrderInput, +) -> Result<(), String> { + fn required_identifier(value: &str, field: &str, max_len: usize) -> Result<(), String> { + let value = value.trim(); + if value.is_empty() { + return Err(format!("{field} is required")); + } + if value.chars().count() > max_len { + return Err(format!("{field} exceeds {max_len} characters")); + } + if value.chars().any(char::is_control) { + return Err(format!("{field} contains control characters")); + } + Ok(()) + } + + fn optional_identifier(value: Option<&str>, field: &str, max_len: usize) -> Result<(), String> { + let Some(value) = value else { + return Ok(()); + }; + required_identifier(value, field, max_len) + } + + required_identifier(&input.user_id, "user_id", 128)?; + required_identifier(&input.order_no, "order_no", 64)?; + required_identifier(&input.gateway_order_id, "gateway_order_id", 128)?; + required_identifier(&input.product_id, "product_id", 64)?; + required_identifier(&input.pay_currency, "pay_currency", 3)?; + if input.pay_currency.trim().chars().count() != 3 + || !input + .pay_currency + .trim() + .chars() + .all(|character| character.is_ascii_alphabetic()) + { + return Err("pay_currency must be a 3-letter currency code".to_string()); + } + + let payment_method = canonicalize_payment_method(&input.payment_method)?; + optional_identifier(input.payment_provider.as_deref(), "payment_provider", 64)?; + optional_identifier(input.payment_channel.as_deref(), "payment_channel", 64)?; + validate_payment_provider_channel_binding( + &payment_method, + input.payment_provider.as_deref(), + input.payment_channel.as_deref(), + )?; + + let is_admin_grant = payment_method == "admin_grant" + && input + .payment_provider + .as_deref() + .is_some_and(|value| value.trim().eq_ignore_ascii_case("admin")) + && input + .payment_channel + .as_deref() + .is_some_and(|value| value.trim().eq_ignore_ascii_case("manual")); + if !input.amount_usd.is_finite() || input.amount_usd < 0.0 { + return Err("amount_usd must be finite and non-negative".to_string()); + } + if !input.pay_amount.is_finite() || input.pay_amount < 0.0 { + return Err("pay_amount must be finite and non-negative".to_string()); + } + if is_admin_grant { + if !input.pay_currency.trim().eq_ignore_ascii_case("USD") { + return Err("admin_grant orders must use USD as the settlement currency".to_string()); + } + if input.amount_usd != 0.0 || input.pay_amount != 0.0 { + return Err("admin_grant orders must have zero amounts".to_string()); + } + } else if input.amount_usd <= 0.0 || input.pay_amount <= 0.0 { + return Err("paid plan orders must have positive amounts".to_string()); + } + if !input.exchange_rate.is_finite() || input.exchange_rate <= 0.0 { + return Err("exchange_rate must be finite and positive".to_string()); + } + if input.expires_at_unix_secs == 0 || input.expires_at_unix_secs > i64::MAX as u64 { + return Err("expires_at_unix_secs is outside the supported range".to_string()); + } + if !input.gateway_response.is_object() { + return Err("gateway_response must be an object".to_string()); + } + + let Some(snapshot) = input.product_snapshot.as_object() else { + return Err("product_snapshot must be an object".to_string()); + }; + let Some(snapshot_id) = snapshot.get("id").and_then(serde_json::Value::as_str) else { + return Err("product_snapshot.id is required".to_string()); + }; + if snapshot_id.trim() != input.product_id.trim() { + return Err("product_snapshot.id must match product_id".to_string()); + } + crate::repository::billing::checked_plan_duration_days_from_snapshot(&input.product_snapshot)?; + if let Some(max_active) = snapshot.get("max_active_per_user") { + if max_active.as_i64().is_none_or(|value| value <= 0) { + return Err("product_snapshot.max_active_per_user must be positive".to_string()); + } + } + if let Some(scope) = snapshot + .get("purchase_limit_scope") + .and_then(serde_json::Value::as_str) + { + if !matches!(scope, "active_period" | "lifetime" | "unlimited") { + return Err("product_snapshot.purchase_limit_scope is invalid".to_string()); + } + } + let entitlements = snapshot + .get("entitlements") + .or_else(|| snapshot.get("entitlements_json")) + .cloned() + .unwrap_or_else(|| serde_json::json!([])); + validate_plan_wallet_credit_entitlements(&entitlements) +} + +/// Validate a wallet-recharge order before a backend creates its wallet or +/// payment row. Recharge orders predate the explicit provider/channel +/// columns, so a provider may be absent and a legacy provider-bound row may +/// also omit its channel. When a channel is present, however, it must be a +/// channel that the corresponding official integration can verify. +pub fn validate_wallet_recharge_order_input( + input: &CreateWalletRechargeOrderInput, +) -> Result<(), String> { + fn required_identifier(value: &str, field: &str, max_len: usize) -> Result<(), String> { + let value = value.trim(); + if value.is_empty() { + return Err(format!("{field} is required")); + } + if value.chars().count() > max_len { + return Err(format!("{field} exceeds {max_len} characters")); + } + if value.chars().any(char::is_control) { + return Err(format!("{field} contains control characters")); + } + Ok(()) + } + + fn optional_identifier(value: Option<&str>, field: &str, max_len: usize) -> Result<(), String> { + let Some(value) = value else { + return Ok(()); + }; + required_identifier(value, field, max_len) + } + + required_identifier(&input.user_id, "user_id", 128)?; + optional_identifier( + input.preferred_wallet_id.as_deref(), + "preferred_wallet_id", + 64, + )?; + required_identifier(&input.order_no, "order_no", 128)?; + required_identifier(&input.gateway_order_id, "gateway_order_id", 128)?; + let payment_method = canonicalize_payment_method(&input.payment_method)?; + optional_identifier(input.payment_provider.as_deref(), "payment_provider", 64)?; + optional_identifier(input.payment_channel.as_deref(), "payment_channel", 64)?; + + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err("amount_usd must be finite and positive".to_string()); + } + if input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err("pay_amount must be finite and positive".to_string()); + } + if input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + { + return Err("exchange_rate must be finite and positive".to_string()); + } + if let Some(currency) = input.pay_currency.as_deref() { + let currency = currency.trim(); + if currency.len() != 3 || !currency.bytes().all(|byte| byte.is_ascii_alphabetic()) { + return Err("pay_currency must be a 3-letter currency code".to_string()); + } + } + if input.expires_at_unix_secs == 0 || input.expires_at_unix_secs > i64::MAX as u64 { + return Err("expires_at_unix_secs is outside the supported range".to_string()); + } + if !input.gateway_response.is_object() { + return Err("gateway_response must be an object".to_string()); + } + + let provider = input + .payment_provider + .as_deref() + .map(canonicalize_payment_method) + .transpose()?; + let channel = input + .payment_channel + .as_deref() + .map(canonicalize_payment_method) + .transpose()?; + if let Some(provider) = provider.as_deref() { + // Keep legacy rows that have a provider but no channel readable. A + // present channel still goes through the strict official-provider + // allowlist above, preventing unsupported channels from being stored. + if let Some(channel) = channel.as_deref() { + validate_payment_provider_channel_binding( + &payment_method, + Some(provider), + Some(channel), + )?; + } else { + let method_matches = match provider { + "epay" => matches!(payment_method.as_str(), "epay" | "alipay" | "wxpay"), + "alipay" | "wxpay" | "stripe" => payment_method == provider, + "admin" => payment_method == "admin_grant", + _ => payment_method == provider, + }; + if !method_matches { + return Err("payment method and provider binding mismatch".to_string()); + } + } + } else if channel.is_some() { + // Provider-less rows are a legacy shape from before the explicit + // provider/channel columns existed. Such rows never had a channel; + // accepting one now would persist an identity that no official + // callback verifier can bind to a gateway. + return Err("payment channel requires a payment provider".to_string()); + } + Ok(()) +} + +fn wallet_recharge_gateway_identity( + payment_method: &str, + payment_provider: Option<&str>, + payment_channel: Option<&str>, +) -> Option<(String, Option)> { + let method = canonicalize_payment_method(payment_method).ok()?; + let provider = payment_provider + .map(canonicalize_payment_method) + .transpose() + .ok()?; + let channel = payment_channel + .map(canonicalize_payment_method) + .transpose() + .ok()?; + + // Historical EPay orders used alipay/wxpay as the method and either left + // provider/channel empty or stored provider=epay. Normalize those rows to + // the same identity as today's method=epay representation. + if provider.as_deref() == Some("epay") + || (provider.is_none() && matches!(method.as_str(), "alipay" | "wxpay")) + { + let channel = channel + .or_else(|| matches!(method.as_str(), "alipay" | "wxpay").then(|| method.clone())); + return Some(("epay".to_string(), channel)); + } + + Some((provider.unwrap_or(method), channel)) +} + +#[allow(clippy::too_many_arguments)] +pub fn wallet_recharge_replay_matches( + existing_wallet_id: &str, + existing_amount_usd: f64, + existing_pay_amount: Option, + existing_pay_currency: Option<&str>, + existing_exchange_rate: Option, + existing_payment_method: &str, + existing_payment_provider: Option<&str>, + existing_payment_channel: Option<&str>, + wallet_id: &str, + input: &CreateWalletRechargeOrderInput, +) -> bool { + fn same_number(left: f64, right: f64) -> bool { + left.is_finite() && right.is_finite() && (left - right).abs() <= 0.00000001 + } + fn same_optional_number(left: Option, right: Option) -> bool { + match (left, right) { + (Some(left), Some(right)) => same_number(left, right), + (None, None) => true, + _ => false, + } + } + fn same_optional_text(left: Option<&str>, right: Option<&str>) -> bool { + match (left, right) { + (Some(left), Some(right)) => left.trim().eq_ignore_ascii_case(right.trim()), + (None, None) => true, + _ => false, + } + } + + existing_wallet_id == wallet_id + && same_number(existing_amount_usd, input.amount_usd) + && same_optional_number(existing_pay_amount, input.pay_amount) + && same_optional_text(existing_pay_currency, input.pay_currency.as_deref()) + && same_optional_number(existing_exchange_rate, input.exchange_rate) + && wallet_recharge_gateway_identity( + existing_payment_method, + existing_payment_provider, + existing_payment_channel, + ) == wallet_recharge_gateway_identity( + &input.payment_method, + input.payment_provider.as_deref(), + input.payment_channel.as_deref(), + ) +} + +/// Validate amounts read from a pending payment order before crediting it. +/// A zero-value order is valid only for the server-controlled admin grant +/// namespace; ordinary payment orders must always carry positive amounts. +pub fn validate_payment_order_credit_amounts( + order_kind: &str, + payment_method: &str, + payment_provider: Option<&str>, + payment_channel: Option<&str>, + amount_usd: f64, + pay_amount: Option, +) -> Result<(), String> { + let is_admin_grant = order_kind.eq_ignore_ascii_case("plan_purchase") + && payment_method.eq_ignore_ascii_case("admin_grant") + && payment_provider.is_some_and(|value| value.trim().eq_ignore_ascii_case("admin")) + && payment_channel.is_some_and(|value| value.trim().eq_ignore_ascii_case("manual")); + if !amount_usd.is_finite() || amount_usd < 0.0 { + return Err("payment order amount is invalid".to_string()); + } + if pay_amount.is_some_and(|value| !value.is_finite() || value < 0.0) { + return Err("payment order pay amount is invalid".to_string()); + } + if is_admin_grant { + if amount_usd != 0.0 || pay_amount.is_some_and(|value| value != 0.0) { + return Err("admin_grant payment order amounts are invalid".to_string()); + } + } else if amount_usd <= 0.0 || pay_amount.is_some_and(|value| value <= 0.0) { + return Err("payment order amount is invalid".to_string()); + } + Ok(()) +} + +/// Match the provider settlement amount against a payment order. +/// +/// `pay_amount` was nullable in the original payment-order schema. New +/// orders carry the provider amount and must compare it exactly (within the +/// storage precision). For a legacy row that has no provider amount, only a +/// deterministic amount reconstructed from the order's own USD amount, +/// currency, and exchange rate is accepted. In particular, a callback's +/// self-reported `amount_usd` is not sufficient for an official callback: +/// gateway handlers may intentionally project that field from the order +/// snapshot before reaching the repository. +pub fn payment_callback_amount_matches_order( + order_amount_usd: f64, + order_pay_amount: Option, + order_pay_currency: Option<&str>, + order_exchange_rate: Option, + callback_amount_usd: f64, + callback_pay_amount: Option, +) -> bool { + const EPSILON: f64 = 0.000001; + + fn valid_positive(value: f64) -> bool { + value.is_finite() && value > 0.0 + } + + fn rounded_major(value: f64) -> Option { + if !value.is_finite() { + return None; + } + let rounded = (value * 100.0).round() / 100.0; + valid_positive(rounded).then_some(rounded) + } + + if !valid_positive(order_amount_usd) || !valid_positive(callback_amount_usd) { + return false; + } + + match (callback_pay_amount, order_pay_amount) { + (Some(callback), Some(order)) => { + valid_positive(callback) && valid_positive(order) && (callback - order).abs() <= EPSILON + } + // A provider amount on the callback cannot be accepted against a + // legacy row unless the expected settlement can be reconstructed from + // values that were already persisted with that row. + (Some(callback), None) if valid_positive(callback) => { + let Some(currency) = order_pay_currency.map(str::trim) else { + return false; + }; + let currency = currency.to_ascii_uppercase(); + if currency.len() != 3 || !currency.bytes().all(|byte| byte.is_ascii_alphabetic()) { + return false; + } + // USD is the accounting currency. Historical rows sometimes + // stored the old default CNY rate even when the currency was USD; + // never turn $10 into 72 CNY during compatibility matching. + let exchange_rate = if currency == "USD" { + 1.0 + } else { + let Some(rate) = order_exchange_rate else { + return false; + }; + if !valid_positive(rate) { + return false; + } + rate + }; + let expected = rounded_major(order_amount_usd * exchange_rate); + expected.is_some_and(|expected| (callback - expected).abs() <= EPSILON) + } + // Keep malformed provider amounts out of the compatibility path. + (Some(_), None) => false, + // Callbacks without a provider amount are legacy/non-official input. + // Keep the old USD fallback, but never use it when the order had a + // provider amount that the callback omitted. + (None, None) => (callback_amount_usd - order_amount_usd).abs() <= EPSILON, + (None, Some(_)) => false, + } +} + +const WALLET_CHECKOUT_SAFE_KEYS: &[&str] = &[ + "gateway", + "display_name", + "gateway_order_id", + "payment_url", + "payment_params", + "submit_method", + "qr_code", + "expires_at", + "pay_amount", + "base_pay_amount", + "fee_rate", + "fee_amount", + "pay_currency", + "payment_channel", + "payment_provider", + "code_url", + "h5_url", + "jsapi", + "publishable_key", + "intent_id", + "payment_method_types", + "provider_label", + "subject", + "instructions", + "callback_url", + "return_url", + "integration_status", + "manual_credit", + // Internal, server-generated checkout-claim metadata. These fields are + // deliberately excluded by the public response projection. + "checkout_claim_token", + "checkout_claimed_at_unix_secs", + "failed_at_unix_secs", + "failure_reason", +]; +const WALLET_CHECKOUT_SAFE_EPAY_PARAM_KEYS: &[&str] = &[ + "pid", + "type", + "out_trade_no", + "notify_url", + "return_url", + "name", + "money", + "sign_type", + "sign", +]; +const STRIPE_CLIENT_SECRET_ENCRYPTED_KEY: &str = "_stripe_client_secret_encrypted"; +const PAYMENT_ORDER_STRIPE_CLIENT_SECRET_V2_PREFIX: &str = + "aether-payment-order-stripe-client-secret-v2:aether-runtime-secret-v1:"; +pub const WALLET_RECHARGE_CHECKOUT_CLAIM_LEASE_SECS: u64 = 120; + +/// Project a provider checkout response before it is persisted in +/// `payment_orders.gateway_response`. +/// +/// Provider responses are untrusted input. Keep this projection in the data +/// contract so callers that bypass the HTTP handler cannot persist credentials, +/// customer data, or arbitrary nested JSON. The raw Stripe `client_secret` is +/// deliberately excluded; the gateway may pass only its encrypted form. +/// +/// This generic projection intentionally does not assign an order kind. Plan +/// purchases and wallet recharges share the same gateway response shape, and +/// stamping a wallet discriminator here would misclassify plan orders. +pub fn project_wallet_gateway_response( + value: &serde_json::Value, +) -> Result { + let Some(object) = value.as_object() else { + return Err("wallet recharge gateway response must be an object".to_string()); + }; + let mut projected = serde_json::Map::new(); + for key in WALLET_CHECKOUT_SAFE_KEYS { + let Some(item) = object.get(*key) else { + continue; + }; + let item = match *key { + "payment_method_types" => { + let Some(values) = item.as_array() else { + continue; + }; + let values = values + .iter() + .filter_map(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(|value| serde_json::Value::String(value.to_string())) + .collect::>(); + if values.is_empty() { + continue; + } + serde_json::Value::Array(values) + } + _ if *key == "payment_params" => { + let Some(params) = item.as_object() else { + continue; + }; + let mut safe_params = serde_json::Map::new(); + for param_key in WALLET_CHECKOUT_SAFE_EPAY_PARAM_KEYS { + if let Some(param_value) = params.get(*param_key) { + if param_value.is_string() { + safe_params.insert((*param_key).to_string(), param_value.clone()); + } + } + } + if safe_params.is_empty() { + continue; + } + serde_json::Value::Object(safe_params) + } + _ if item.is_object() || item.is_array() => continue, + _ => item.clone(), + }; + projected.insert((*key).to_string(), item); + } + + if let Some(encrypted) = object + .get(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty() && value.len() <= 8192) + { + projected.insert( + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY.to_string(), + serde_json::Value::String(encrypted.to_string()), + ); + } + Ok(serde_json::Value::Object(projected)) +} + +/// Project a wallet-recharge checkout response and attach its server-controlled +/// discriminator. The in-memory repository uses this marker because it has no +/// dedicated `order_kind` column. +pub fn project_wallet_recharge_gateway_response( + value: &serde_json::Value, +) -> Result { + let mut projected = project_wallet_gateway_response(value)?; + let Some(object) = projected.as_object_mut() else { + return Err("wallet recharge gateway response must be an object".to_string()); + }; + object.insert( + "order_kind".to_string(), + serde_json::Value::String("wallet_recharge".to_string()), + ); + Ok(projected) +} + +pub fn wallet_recharge_checkout_claim_token(value: &serde_json::Value) -> Option<&str> { + value + .as_object() + .and_then(|object| object.get("checkout_claim_token")) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|token| !token.is_empty() && token.len() <= 128) +} + +pub fn wallet_recharge_checkout_claimed_at(value: &serde_json::Value) -> Option { + value + .as_object() + .and_then(|object| object.get("checkout_claimed_at_unix_secs")) + .and_then(serde_json::Value::as_u64) +} + +/// Build the server-controlled placeholder used while an external checkout +/// request is in flight. The response is projected again at this boundary so +/// callers cannot smuggle additional fields into the persisted claim. +pub fn wallet_recharge_checkout_claim_response( + value: &serde_json::Value, + claim_token: &str, + claimed_at_unix_secs: u64, +) -> Result { + let token = claim_token.trim(); + if token.is_empty() || token.len() > 128 || token.chars().any(char::is_control) { + return Err("wallet recharge checkout claim token is invalid".to_string()); + } + let mut projected = project_wallet_recharge_gateway_response(value)?; + let Some(object) = projected.as_object_mut() else { + return Err("wallet recharge gateway response must be an object".to_string()); + }; + object.insert( + "order_kind".to_string(), + serde_json::Value::String("wallet_recharge".to_string()), + ); + object.insert( + "integration_status".to_string(), + serde_json::Value::String("checkout_pending".to_string()), + ); + object.insert( + "checkout_claim_token".to_string(), + serde_json::Value::String(token.to_string()), + ); + object.insert( + "checkout_claimed_at_unix_secs".to_string(), + serde_json::Value::Number(serde_json::Number::from(claimed_at_unix_secs)), + ); + Ok(projected) +} + +pub fn wallet_recharge_checkout_failed_response( + value: Option<&serde_json::Value>, + reason: &str, + failed_at_unix_secs: u64, +) -> serde_json::Value { + wallet_recharge_checkout_failure_response(value, reason, failed_at_unix_secs, false) +} + +/// Record a checkout failure whose provider-side outcome cannot be proven. +/// +/// A transport timeout, truncated response, or persistence error can happen +/// after a gateway has accepted the payment request. Such an order remains +/// settleable by a verified callback, but must not be reclaimed for a second +/// checkout because replacing its gateway identifier could strand the first +/// payment. +pub fn wallet_recharge_checkout_uncertain_response( + value: Option<&serde_json::Value>, + reason: &str, + failed_at_unix_secs: u64, +) -> serde_json::Value { + wallet_recharge_checkout_failure_response(value, reason, failed_at_unix_secs, true) +} + +fn wallet_recharge_checkout_failure_response( + value: Option<&serde_json::Value>, + reason: &str, + failed_at_unix_secs: u64, + provider_request_may_have_succeeded: bool, +) -> serde_json::Value { + let integration_status = if provider_request_may_have_succeeded { + "checkout_uncertain" + } else { + "checkout_failed" + }; + let mut projected = value + .and_then(|value| project_wallet_recharge_gateway_response(value).ok()) + .unwrap_or_else(|| serde_json::json!({})); + let Some(object) = projected.as_object_mut() else { + return serde_json::json!({ + "order_kind": "wallet_recharge", + "integration_status": integration_status, + }); + }; + object.remove("checkout_claim_token"); + object.remove("checkout_claimed_at_unix_secs"); + object.insert( + "order_kind".to_string(), + serde_json::Value::String("wallet_recharge".to_string()), + ); + object.insert( + "integration_status".to_string(), + serde_json::Value::String(integration_status.to_string()), + ); + let reason = reason.trim(); + if !reason.is_empty() { + let bounded = reason.chars().take(512).collect::(); + object.insert( + "failure_reason".to_string(), + serde_json::Value::String(bounded), + ); + } + object.insert( + "failed_at_unix_secs".to_string(), + serde_json::Value::Number(serde_json::Number::from(failed_at_unix_secs)), + ); + projected +} + +const WALLET_RECHARGE_CHECKOUT_EVIDENCE_KEYS: &[&str] = &[ + "payment_url", + "payment_params", + "qr_code", + "code_url", + "h5_url", + "jsapi", + "client_secret", + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY, + "intent_id", +]; + +/// Whether an order still contains only the server-created checkout claim and +/// has no provider checkout evidence. This marker is intentionally strict: +/// an order with any provider URL, token, or intent must never be reclaimed. +pub fn wallet_recharge_response_is_checkout_placeholder(value: &serde_json::Value) -> bool { + let Some(object) = value.as_object() else { + return false; + }; + let integration_status = object + .get("integration_status") + .and_then(serde_json::Value::as_str); + if !matches!( + integration_status, + Some("checkout_pending" | "checkout_failed" | "checkout_uncertain") + ) { + return false; + } + if object.get("order_kind").and_then(serde_json::Value::as_str) != Some("wallet_recharge") { + return false; + } + !WALLET_RECHARGE_CHECKOUT_EVIDENCE_KEYS + .iter() + .any(|key| object.contains_key(*key)) +} + +/// Whether a failed payment order is the server-created checkout placeholder +/// that may still be settled by a verified provider callback. +/// +/// Checkout failures are persisted before the provider's eventual callback +/// can arrive (for example when the checkout response is lost after the +/// provider accepted it). Keep this exception narrowly scoped: ordinary +/// failed orders and placeholders that contain any provider checkout evidence +/// must remain non-creditable. +pub fn payment_order_is_failed_wallet_checkout_placeholder( + order_status: &str, + order_kind: &str, + gateway_response: Option<&serde_json::Value>, +) -> bool { + if !order_status.trim().eq_ignore_ascii_case("failed") + || !order_kind.trim().eq_ignore_ascii_case("wallet_recharge") + { + return false; + } + let Some(gateway_response) = gateway_response else { + return false; + }; + if !wallet_recharge_response_is_checkout_placeholder(gateway_response) { + return false; + } + gateway_response + .get("integration_status") + .and_then(serde_json::Value::as_str) + .is_some_and(|status| { + matches!( + status.trim().to_ascii_lowercase().as_str(), + "checkout_failed" | "checkout_uncertain" + ) + }) +} + +/// Whether a failed wallet checkout was marked as provider-outcome-uncertain. +/// Such orders may be settled by a verified callback even after the local +/// checkout expiry, because the provider may have accepted the request before +/// the response was lost. +pub fn payment_order_is_uncertain_wallet_checkout_placeholder( + order_status: &str, + order_kind: &str, + gateway_response: Option<&serde_json::Value>, +) -> bool { + if !payment_order_is_failed_wallet_checkout_placeholder( + order_status, + order_kind, + gateway_response, + ) { + return false; + } + gateway_response + .and_then(|value| value.get("integration_status")) + .and_then(serde_json::Value::as_str) + .is_some_and(|status| status.trim().eq_ignore_ascii_case("checkout_uncertain")) +} + +pub fn wallet_recharge_order_is_checkout_placeholder(order: &StoredAdminPaymentOrder) -> bool { + order + .gateway_response + .as_ref() + .is_some_and(wallet_recharge_response_is_checkout_placeholder) +} + +/// Convert a persisted timestamp that may use either seconds or milliseconds +/// to epoch seconds. +/// +/// The `*_unix_ms` field names predate the current adapters. The in-memory +/// repository stores milliseconds, while the SQL adapters historically +/// selected epoch seconds under the same aliases. Keep this compatibility +/// conversion at the contract boundary so public payloads and lease decisions +/// agree across backends. +pub fn stored_timestamp_unix_secs(value: u64) -> u64 { + const MILLIS_THRESHOLD: u64 = 10_000_000_000; + if value >= MILLIS_THRESHOLD { + value / 1000 + } else { + value + } +} + +/// Return the creation time in epoch seconds for checkout lease decisions. +pub fn wallet_recharge_order_created_at_unix_secs(order: &StoredAdminPaymentOrder) -> u64 { + stored_timestamp_unix_secs(order.created_at_unix_ms) +} + +/// A failed or expired placeholder may be claimed by one subsequent checkout +/// attempt. Live pending placeholders remain owned by the request that first +/// created them, which prevents concurrent provider-side duplicate orders. +pub fn wallet_recharge_order_is_reclaimable_placeholder( + order: &StoredAdminPaymentOrder, + now_unix_secs: u64, +) -> bool { + if !wallet_recharge_order_is_checkout_placeholder(order) { + return false; + } + // An uncertain provider result may still arrive as a signed callback. + // Never replace its gateway identity, even after the original claim lease + // has elapsed. + let integration_status = order + .gateway_response + .as_ref() + .and_then(|value| value.get("integration_status")) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .unwrap_or_default(); + if integration_status.eq_ignore_ascii_case("checkout_uncertain") { + return false; + } + match order.status.trim().to_ascii_lowercase().as_str() { + "failed" | "expired" => true, + "pending" => { + let claimed_at = order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claimed_at) + .unwrap_or_else(|| wallet_recharge_order_created_at_unix_secs(order)); + now_unix_secs.saturating_sub(claimed_at) >= WALLET_RECHARGE_CHECKOUT_CLAIM_LEASE_SECS + } + _ => false, + } +} + +/// Enforce the provider namespace for callbacks verified by an official +/// gateway integration. This is a repository boundary check: callers cannot +/// relabel an official payment method as a provider-less generic callback. +pub fn validate_payment_callback_provider_binding( + payment_method: &str, + payment_provider: Option<&str>, +) -> Result<(), String> { + let payment_method = canonicalize_payment_method(payment_method)?; + let payment_provider = payment_provider + .map(canonicalize_payment_method) + .transpose()?; + let provider_matches = match payment_method.as_str() { + // Legacy EPay orders stored the selected channel as the method. Keep + // that explicit, signed aggregator binding valid while rejecting a + // provider-less generic callback. + "alipay" | "wxpay" => matches!( + payment_provider.as_deref(), + Some(provider) if provider == payment_method || provider == "epay" + ), + "stripe" => payment_provider.as_deref() == Some("stripe"), + "epay" => payment_provider.as_deref() == Some("epay"), + _ => true, + }; + if !provider_matches { + return Err("official payment callback provider binding mismatch".to_string()); + } + Ok(()) +} + +/// Match a verified callback's method against the method stored on an order. +/// +/// EPay historically stored the selected channel (`alipay` or `wxpay`) in +/// `payment_method`, while the notification endpoint identifies itself as the +/// `epay` provider. Keep that legacy representation compatible, but only when +/// both sides are in the EPay namespace; the adapter still verifies the +/// explicit payment channel immediately afterwards. +pub fn payment_callback_method_matches_order( + order_method: &str, + order_provider: Option<&str>, + callback_method: &str, + callback_provider: Option<&str>, +) -> bool { + if order_method.eq_ignore_ascii_case(callback_method) { + return true; + } + let order_method_is_epay_alias = ["epay", "alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(order_method)); + let callback_method_is_epay_alias = ["epay", "alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(callback_method)); + let order_is_epay_namespace = order_provider + .map(|provider| provider.eq_ignore_ascii_case("epay")) + .unwrap_or_else(|| { + // Before payment_provider was added, EPay direct-channel orders + // persisted alipay/wxpay as payment_method and left the provider + // column NULL. Keep only those legacy channel rows compatible. + ["alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(order_method)) + }); + order_is_epay_namespace + && callback_provider.is_some_and(|provider| provider.eq_ignore_ascii_case("epay")) + && order_method_is_epay_alias + && callback_method_is_epay_alias +} + +/// Match the provider namespace stored on a payment order against a verified +/// callback. Historical EPay channel orders predate `payment_provider` and +/// therefore have a NULL provider; only their explicit alipay/wxpay method +/// aliases may be upgraded by an EPay callback. +pub fn payment_callback_provider_matches_order( + order_method: &str, + order_provider: Option<&str>, + callback_method: &str, + callback_provider: Option<&str>, +) -> bool { + match (order_provider, callback_provider) { + (None, None) => true, + (Some(order), Some(callback)) => order.eq_ignore_ascii_case(callback), + (None, Some(callback)) => { + callback.eq_ignore_ascii_case("epay") + && ["alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(order_method)) + && ["epay", "alipay", "wxpay"] + .iter() + .any(|method| method.eq_ignore_ascii_case(callback_method)) + } + _ => false, + } +} + +impl ProcessPaymentCallbackInput { + pub fn canonicalize_and_validate(&mut self) -> Result<(), String> { + self.payment_method = canonicalize_payment_method(&self.payment_method)?; + self.payment_provider = self + .payment_provider + .as_deref() + .map(canonicalize_payment_method) + .transpose()?; + self.payment_channel = self + .payment_channel + .as_deref() + .map(canonicalize_payment_method) + .transpose()?; + validate_payment_callback_provider_binding( + &self.payment_method, + self.payment_provider.as_deref(), + )?; + validate_payment_provider_channel_binding( + &self.payment_method, + self.payment_provider.as_deref(), + self.payment_channel.as_deref(), + )?; + + if self.payment_provider.is_some() + && (self.payment_channel.is_none() + || self + .order_no + .as_deref() + .is_none_or(|value| value.trim().is_empty()) + || self + .gateway_order_id + .as_deref() + .is_none_or(|value| value.trim().is_empty()) + || self.pay_amount.is_none() + || self + .pay_currency + .as_deref() + .is_none_or(|value| value.trim().len() != 3)) + { + return Err( + "official payment callback is missing settlement binding fields".to_string(), + ); + } + Ok(()) + } + + /// Return the only callback data that may be copied to a payment order. + /// + /// `payload` is provider-controlled (and may contain credentials, payment + /// capabilities, or customer data), so adapters must never persist it as + /// `payment_orders.gateway_response`. Keep this projection at the shared + /// contract boundary so every database backend applies the same allowlist, + /// including callers that bypass the HTTP handler's projection. + pub fn gateway_response_projection( + &self, + order_no: &str, + gateway_order_id: Option<&str>, + ) -> serde_json::Value { + serde_json::json!({ + "gateway": self.payment_method, + "payment_provider": self.payment_provider, + "payment_channel": self.payment_channel, + "order_no": order_no, + "gateway_order_id": gateway_order_id, + "amount_usd": self.amount_usd, + "pay_amount": self.pay_amount, + "pay_currency": self.pay_currency, + "exchange_rate": self.exchange_rate, + "signature_valid": self.signature_valid, + }) + } +} + pub fn redeem_code_refundable_amount(balance_bucket: &str, amount_usd: f64) -> f64 { if redeem_code_credits_recharge_balance(balance_bucket) { amount_usd @@ -579,6 +1835,117 @@ pub fn redeem_code_refundable_amount(balance_bucket: &str, amount_usd: f64) -> f } } +pub fn validate_admin_redeem_code_batch_input( + input: &CreateAdminRedeemCodeBatchInput, +) -> Result<(), String> { + let name = input.name.trim(); + if name.is_empty() || name.chars().count() > 120 || name.chars().any(char::is_control) { + return Err("redeem code batch name is invalid".to_string()); + } + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Err("redeem code batch amount_usd must be finite and positive".to_string()); + } + if !input.currency.trim().eq_ignore_ascii_case("USD") { + return Err("redeem code batch currency must be USD".to_string()); + } + if !matches!(input.balance_bucket.trim(), "gift" | "recharge") { + return Err("redeem code batch balance_bucket must be gift or recharge".to_string()); + } + if input.total_count == 0 || input.total_count > 5_000 { + return Err("redeem code batch total_count must be between 1 and 5000".to_string()); + } + if input + .expires_at_unix_secs + .is_some_and(|value| value == 0 || value > i64::MAX as u64) + { + return Err( + "redeem code batch expires_at_unix_secs is outside the supported range".to_string(), + ); + } + Ok(()) +} + +/// Validate a persisted redeem batch and the wallet arithmetic before any +/// credit is applied. Imported rows can bypass repository input validation, +/// so redemption must fail closed when either side contains invalid amounts. +pub fn validate_redeem_wallet_credit( + balance_bucket: &str, + amount_usd: f64, + before_recharge: f64, + before_gift: f64, + before_total_recharged: f64, +) -> Result<(f64, f64, f64), String> { + let balance_bucket = balance_bucket.trim(); + if !matches!(balance_bucket, "gift" | "recharge") { + return Err("redeem code batch balance bucket is invalid".to_string()); + } + if !amount_usd.is_finite() || amount_usd <= 0.0 { + return Err("redeem code batch amount is invalid".to_string()); + } + if !before_recharge.is_finite() + || !before_gift.is_finite() + || !before_total_recharged.is_finite() + || before_gift < 0.0 + || before_total_recharged < 0.0 + || !(before_recharge + before_gift).is_finite() + { + return Err("wallet amount is invalid".to_string()); + } + + let after_recharge = if balance_bucket == "recharge" { + before_recharge + amount_usd + } else { + before_recharge + }; + let after_gift = if balance_bucket == "recharge" { + before_gift + } else { + before_gift + amount_usd + }; + let after_total_recharged = before_total_recharged + amount_usd; + if !after_recharge.is_finite() + || !after_gift.is_finite() + || !after_total_recharged.is_finite() + || !(after_recharge + after_gift).is_finite() + { + return Err("redeem wallet credit would overflow".to_string()); + } + Ok((after_recharge, after_gift, after_total_recharged)) +} + +/// Validate a manual recharge and calculate the exact wallet values that may +/// be persisted. Manual recharges are also invoked by import and recovery +/// workflows, so the repository cannot rely on the admin HTTP validator. +pub fn validate_manual_wallet_recharge( + amount_usd: f64, + before_recharge: f64, + before_gift: f64, + before_total_recharged: f64, +) -> Result<(f64, f64), String> { + if !amount_usd.is_finite() || amount_usd <= 0.0 { + return Err("manual recharge amount must be finite and positive".to_string()); + } + if !before_recharge.is_finite() + || !before_gift.is_finite() + || before_gift < 0.0 + || !before_total_recharged.is_finite() + || before_total_recharged < 0.0 + || !(before_recharge + before_gift).is_finite() + { + return Err("wallet amount is invalid for manual recharge".to_string()); + } + + let after_recharge = before_recharge + amount_usd; + let after_total_recharged = before_total_recharged + amount_usd; + if !after_recharge.is_finite() + || !after_total_recharged.is_finite() + || !(after_recharge + before_gift).is_finite() + { + return Err("manual recharge would overflow wallet amounts".to_string()); + } + Ok((after_recharge, after_total_recharged)) +} + #[allow(clippy::large_enum_variant)] #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub enum RedeemWalletCodeOutcome { @@ -597,7 +1964,7 @@ pub enum RedeemWalletCodeOutcome { WalletInactive, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct CreateWalletRechargeOrderInput { pub preferred_wallet_id: Option, pub user_id: String, @@ -614,14 +1981,209 @@ pub struct CreateWalletRechargeOrderInput { pub expires_at_unix_secs: u64, } +impl std::fmt::Debug for CreateWalletRechargeOrderInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreateWalletRechargeOrderInput") + .field("preferred_wallet_id", &self.preferred_wallet_id) + .field("user_id", &self.user_id) + .field("amount_usd", &self.amount_usd) + .field("payment_method", &self.payment_method) + .field("payment_provider", &self.payment_provider) + .field("payment_channel", &self.payment_channel) + .field("gateway_order_id", &self.gateway_order_id) + .field("gateway_response", &WALLET_REDACTED_DEBUG_VALUE) + .field("order_no", &self.order_no) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .finish_non_exhaustive() + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] #[allow(clippy::large_enum_variant)] pub enum CreateWalletRechargeOrderOutcome { Created(StoredAdminPaymentOrder), + Existing(StoredAdminPaymentOrder), WalletInactive, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct UpdateWalletRechargeCheckoutInput { + pub order_id: String, + pub gateway_order_id: String, + pub gateway_response: serde_json::Value, +} + +impl std::fmt::Debug for UpdateWalletRechargeCheckoutInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("UpdateWalletRechargeCheckoutInput") + .field("order_id", &self.order_id) + .field("gateway_order_id", &self.gateway_order_id) + .field("gateway_response", &WALLET_REDACTED_DEBUG_VALUE) + .finish() + } +} + +/// Exact compare-and-swap for lazily migrating one Stripe client-secret field. +/// +/// Every immutable identity component and the complete observed gateway +/// response are included so a stale reader cannot overwrite another checkout +/// update or move a capability between payment orders. +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct CompareAndSwapPaymentOrderStripeClientSecretInput { + pub order_id: String, + pub order_no: String, + pub wallet_id: String, + pub user_id: Option, + pub payment_method: String, + pub payment_provider: Option, + pub order_kind: String, + pub gateway_order_id: Option, + pub expected_status: String, + pub expected_expires_at_unix_secs: Option, + pub expected_gateway_response: serde_json::Value, + pub expected_client_secret_encrypted: String, + pub replacement_client_secret_encrypted: String, +} + +impl std::fmt::Debug for CompareAndSwapPaymentOrderStripeClientSecretInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CompareAndSwapPaymentOrderStripeClientSecretInput") + .field("order_id", &self.order_id) + .field("order_no", &self.order_no) + .field("wallet_id", &self.wallet_id) + .field("user_id", &self.user_id) + .field("payment_method", &self.payment_method) + .field("payment_provider", &self.payment_provider) + .field("order_kind", &self.order_kind) + .field("gateway_order_id", &self.gateway_order_id) + .field("expected_status", &self.expected_status) + .field( + "expected_expires_at_unix_secs", + &self.expected_expires_at_unix_secs, + ) + .field("expected_gateway_response", &WALLET_REDACTED_DEBUG_VALUE) + .field( + "expected_client_secret_encrypted", + &WALLET_REDACTED_DEBUG_VALUE, + ) + .field( + "replacement_client_secret_encrypted", + &WALLET_REDACTED_DEBUG_VALUE, + ) + .finish() + } +} + +/// Validate a Stripe-secret migration against a row locked by a repository +/// implementation and build the exact replacement response. +/// +/// `Ok(None)` is a normal CAS miss. Invalid replacement values are rejected as +/// input errors; all unrelated JSON fields are preserved byte-for-value. +pub fn payment_order_stripe_client_secret_cas_replacement( + current: &StoredAdminPaymentOrder, + input: &CompareAndSwapPaymentOrderStripeClientSecretInput, +) -> Result, String> { + let replacement_ciphertext = input + .replacement_client_secret_encrypted + .strip_prefix(PAYMENT_ORDER_STRIPE_CLIENT_SECRET_V2_PREFIX); + if input.replacement_client_secret_encrypted.len() > 8192 + || replacement_ciphertext.is_none_or(|ciphertext| ciphertext.is_empty()) + || input + .replacement_client_secret_encrypted + .chars() + .any(char::is_control) + { + return Err("replacement Stripe client-secret envelope is invalid".to_string()); + } + if current.id != input.order_id + || current.order_no != input.order_no + || current.wallet_id != input.wallet_id + || current.user_id != input.user_id + || current.payment_method != input.payment_method + || current.payment_provider != input.payment_provider + || current.order_kind != input.order_kind + || current.gateway_order_id != input.gateway_order_id + || current.status != input.expected_status + || current.expires_at_unix_secs != input.expected_expires_at_unix_secs + || current.gateway_response.as_ref() != Some(&input.expected_gateway_response) + { + return Ok(None); + } + let Some(object) = input.expected_gateway_response.as_object() else { + return Ok(None); + }; + if object + .get(STRIPE_CLIENT_SECRET_ENCRYPTED_KEY) + .and_then(serde_json::Value::as_str) + != Some(input.expected_client_secret_encrypted.as_str()) + { + return Ok(None); + } + let mut replacement = input.expected_gateway_response.clone(); + let Some(object) = replacement.as_object_mut() else { + return Ok(None); + }; + object.insert( + STRIPE_CLIENT_SECRET_ENCRYPTED_KEY.to_string(), + serde_json::Value::String(input.replacement_client_secret_encrypted.clone()), + ); + Ok(Some(replacement)) +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct FailWalletRechargeCheckoutInput { + pub order_id: String, + /// The claim token assigned to the request that created/reclaimed the + /// placeholder. Requiring it prevents a slow error from invalidating a + /// newer retry that has already taken over the order. + pub claim_token: String, + pub reason: String, + /// True when the provider request may have been accepted even though the + /// gateway response was not durably observed. Such failures are marked + /// `checkout_uncertain` and cannot be reclaimed for another checkout. + #[serde(default)] + pub provider_request_may_have_succeeded: bool, +} + +impl std::fmt::Debug for FailWalletRechargeCheckoutInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("FailWalletRechargeCheckoutInput") + .field("order_id", &self.order_id) + .field("claim_token", &WALLET_REDACTED_DEBUG_VALUE) + .field("reason", &WALLET_REDACTED_DEBUG_VALUE) + .field( + "provider_request_may_have_succeeded", + &self.provider_request_may_have_succeeded, + ) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct ReclaimWalletRechargeCheckoutInput { + pub order_id: String, + pub claim_token: String, + pub gateway_response: serde_json::Value, + pub expires_at_unix_secs: u64, +} + +impl std::fmt::Debug for ReclaimWalletRechargeCheckoutInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ReclaimWalletRechargeCheckoutInput") + .field("order_id", &self.order_id) + .field("claim_token", &WALLET_REDACTED_DEBUG_VALUE) + .field("gateway_response", &WALLET_REDACTED_DEBUG_VALUE) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct CreatePlanPurchaseOrderInput { pub preferred_wallet_id: Option, pub user_id: String, @@ -640,6 +2202,25 @@ pub struct CreatePlanPurchaseOrderInput { pub expires_at_unix_secs: u64, } +impl std::fmt::Debug for CreatePlanPurchaseOrderInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreatePlanPurchaseOrderInput") + .field("preferred_wallet_id", &self.preferred_wallet_id) + .field("user_id", &self.user_id) + .field("amount_usd", &self.amount_usd) + .field("payment_method", &self.payment_method) + .field("payment_provider", &self.payment_provider) + .field("payment_channel", &self.payment_channel) + .field("gateway_order_id", &self.gateway_order_id) + .field("gateway_response", &WALLET_REDACTED_DEBUG_VALUE) + .field("order_no", &self.order_no) + .field("product_id", &self.product_id) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .finish_non_exhaustive() + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] #[allow(clippy::large_enum_variant)] pub enum CreatePlanPurchaseOrderOutcome { @@ -662,10 +2243,143 @@ pub struct CreateWalletRefundRequestInput { pub refund_no: String, } +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct CanonicalWalletRefundFields { + pub source_type: String, + pub source_id: Option, + pub refund_mode: String, +} + +/// Validate and canonicalize refund provenance fields at the data boundary. +/// The client may omit fields, but it must not be able to select a different +/// refund route or claim a different source than the server-resolved order. +pub fn canonicalize_wallet_refund_fields( + payment_order_id: Option<&str>, + source_type: Option<&str>, + source_id: Option<&str>, + refund_mode: Option<&str>, + payment_method: Option<&str>, +) -> Result { + fn normalize(value: Option<&str>) -> Option<&str> { + value.map(str::trim).filter(|value| !value.is_empty()) + } + + if let Some(order_id) = normalize(payment_order_id) { + let Some(payment_method) = normalize(payment_method) else { + return Err("payment method is required for an order refund".to_string()); + }; + let expected_mode = default_refund_mode_for_payment_method(payment_method); + if let Some(value) = normalize(source_type) { + if !value.eq_ignore_ascii_case("payment_order") { + return Err("source_type does not match payment_order".to_string()); + } + } + if let Some(value) = normalize(source_id) { + if value != order_id { + return Err("source_id does not match payment_order_id".to_string()); + } + } + if let Some(value) = normalize(refund_mode) { + if !value.eq_ignore_ascii_case(expected_mode) { + return Err("refund_mode does not match the payment method".to_string()); + } + } + return Ok(CanonicalWalletRefundFields { + source_type: "payment_order".to_string(), + source_id: Some(order_id.to_string()), + refund_mode: expected_mode.to_string(), + }); + } + + if normalize(payment_method).is_some() { + return Err("payment method is only valid for an order refund".to_string()); + } + if let Some(value) = normalize(source_type) { + if !value.eq_ignore_ascii_case("wallet_balance") { + return Err("source_type must be wallet_balance without an order".to_string()); + } + } + if normalize(source_id).is_some() { + return Err("source_id is not valid without a payment order".to_string()); + } + if let Some(value) = normalize(refund_mode) { + if !value.eq_ignore_ascii_case("offline_payout") { + return Err("refund_mode must be offline_payout without an order".to_string()); + } + } + Ok(CanonicalWalletRefundFields { + source_type: "wallet_balance".to_string(), + source_id: None, + refund_mode: "offline_payout".to_string(), + }) +} + +fn default_refund_mode_for_payment_method(payment_method: &str) -> &'static str { + if matches!( + payment_method.trim().to_ascii_lowercase().as_str(), + "admin_manual" | "card_recharge" | "card_code" | "gift_code" + ) { + return "offline_payout"; + } + "original_channel" +} + +/// Validate the durable accounting split for a refundable payment order. +/// +/// Database numeric conversions and repeated partial refunds can introduce a +/// very small rounding delta, so compare the invariant with a sub-cent +/// tolerance while still rejecting malformed or out-of-range components. +pub fn payment_order_refund_amounts_are_consistent( + amount_usd: f64, + refunded_amount_usd: f64, + refundable_amount_usd: f64, +) -> bool { + const REFUND_AMOUNT_EPSILON_USD: f64 = 0.000_001; + + amount_usd.is_finite() + && amount_usd > 0.0 + && refunded_amount_usd.is_finite() + && refunded_amount_usd >= 0.0 + && refunded_amount_usd <= amount_usd + REFUND_AMOUNT_EPSILON_USD + && refundable_amount_usd.is_finite() + && refundable_amount_usd >= 0.0 + && refundable_amount_usd <= amount_usd + REFUND_AMOUNT_EPSILON_USD + && (refunded_amount_usd + refundable_amount_usd - amount_usd).abs() + <= REFUND_AMOUNT_EPSILON_USD +} + +/// Return whether a persisted refund proof represents a terminal successful +/// gateway settlement. Gateway evidence is stored either directly (the +/// pending response) or under the `gateway_refund` projection (the successful +/// completion response), so accept both shapes while ignoring arbitrary +/// nested payloads. +pub fn wallet_refund_proof_is_success(value: &serde_json::Value) -> bool { + fn status_is_success(value: Option<&serde_json::Value>) -> bool { + value + .and_then(serde_json::Value::as_str) + .map(str::trim) + .is_some_and(|status| { + status.eq_ignore_ascii_case("success") || status.eq_ignore_ascii_case("succeeded") + }) + } + + let Some(object) = value.as_object() else { + return false; + }; + if let Some(gateway_refund) = object.get("gateway_refund") { + return gateway_refund + .as_object() + .and_then(|gateway_refund| gateway_refund.get("status")) + .is_some_and(|status| status_is_success(Some(status))); + } + status_is_success(object.get("status")) +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub enum CreateWalletRefundRequestOutcome { Created(StoredAdminWalletRefund), Duplicate(StoredAdminWalletRefund), + InvalidInput(String), WalletMissing, RefundAmountExceedsAvailableBalance, PaymentOrderNotFound, @@ -674,7 +2388,7 @@ pub enum CreateWalletRefundRequestOutcome { DuplicateRejected, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct ProcessPaymentCallbackInput { pub payment_method: String, pub payment_provider: Option, @@ -691,6 +2405,24 @@ pub struct ProcessPaymentCallbackInput { pub signature_valid: bool, } +impl std::fmt::Debug for ProcessPaymentCallbackInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProcessPaymentCallbackInput") + .field("payment_method", &self.payment_method) + .field("payment_provider", &self.payment_provider) + .field("payment_channel", &self.payment_channel) + .field("callback_key", &WALLET_REDACTED_DEBUG_VALUE) + .field("order_no", &self.order_no) + .field("gateway_order_id", &self.gateway_order_id) + .field("amount_usd", &self.amount_usd) + .field("payload_hash", &WALLET_REDACTED_DEBUG_VALUE) + .field("payload", &WALLET_REDACTED_DEBUG_VALUE) + .field("signature_valid", &self.signature_valid) + .finish_non_exhaustive() + } +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] #[allow(clippy::large_enum_variant)] pub enum ProcessPaymentCallbackOutcome { @@ -758,6 +2490,17 @@ pub struct CompleteAdminWalletRefundInput { pub payout_proof: Option, } +/// Records a provider response while the local refund remains processing. +/// This persists asynchronous gateway evidence without releasing the wallet +/// reservation or changing the refund state. +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct UpdateAdminWalletRefundGatewayInput { + pub wallet_id: String, + pub refund_id: String, + pub gateway_refund_id: String, + pub payout_proof: Option, +} + #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct FailAdminWalletRefundInput { pub wallet_id: String, @@ -766,7 +2509,7 @@ pub struct FailAdminWalletRefundInput { pub operator_id: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub struct CreditAdminPaymentOrderInput { pub order_id: String, pub gateway_order_id: Option, @@ -777,6 +2520,24 @@ pub struct CreditAdminPaymentOrderInput { pub operator_id: Option, } +impl std::fmt::Debug for CreditAdminPaymentOrderInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CreditAdminPaymentOrderInput") + .field("order_id", &self.order_id) + .field("gateway_order_id", &self.gateway_order_id) + .field("pay_amount", &self.pay_amount) + .field("pay_currency", &self.pay_currency) + .field("exchange_rate", &self.exchange_rate) + .field( + "gateway_response_patch", + &wallet_redacted_debug_option(&self.gateway_response_patch), + ) + .field("operator_id", &self.operator_id) + .finish() + } +} + #[async_trait] pub trait WalletReadRepository: Send + Sync { async fn find( @@ -803,6 +2564,22 @@ pub trait WalletReadRepository: Send + Sync { unlimited: bool, ) -> Result, crate::DataLayerError>; + /// Atomically initialize a user wallet and report whether this call won + /// the create race. Implementations that cannot provide the creation bit + /// must fail explicitly; callers use it to decide whether compensation is + /// allowed during an aggregate import rollback. + async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, initial_gift_usd, unlimited); + Err(crate::DataLayerError::InvalidInput( + "atomic user wallet initialization is not available".to_string(), + )) + } + async fn initialize_auth_api_key_wallet( &self, api_key_id: &str, @@ -810,6 +2587,20 @@ pub trait WalletReadRepository: Send + Sync { unlimited: bool, ) -> Result, crate::DataLayerError>; + /// Atomically initialize an API-key wallet and report whether this call + /// created the row. See the user-wallet variant for the rollback contract. + async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, crate::DataLayerError> { + let _ = (api_key_id, initial_gift_usd, unlimited); + Err(crate::DataLayerError::InvalidInput( + "atomic API-key wallet initialization is not available".to_string(), + )) + } + #[allow(clippy::too_many_arguments)] async fn update_auth_user_wallet_snapshot( &self, @@ -927,12 +2718,38 @@ pub trait WalletReadRepository: Send + Sync { order_id: &str, ) -> Result, crate::DataLayerError>; + /// Finds a wallet-recharge order by the stable merchant order number. + /// + /// Public recharge retries use a deterministic order number derived from + /// their idempotency key. Keeping this lookup on the repository boundary + /// lets the gateway replay the original checkout without contacting the + /// payment provider again. + async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, crate::DataLayerError> { + let _ = (user_id, order_no); + Ok(None) + } + async fn find_pending_plan_purchase_order_by_user_id( &self, user_id: &str, product_id: &str, ) -> Result, crate::DataLayerError>; + /// Finds any payment order by its globally unique merchant order number. + /// Checkout uses this narrow lookup before contacting a provider so a + /// deterministic retry key cannot collide with an older terminal order. + async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, crate::DataLayerError> { + let _ = order_no; + Ok(None) + } + async fn find_wallet_refund( &self, wallet_id: &str, @@ -964,11 +2781,95 @@ pub trait WalletReadRepository: Send + Sync { #[async_trait] pub trait WalletWriteRepository: Send + Sync { + /// Delete one wallet identified by its exact id and owner, but only when it has no + /// financial or usage references. This is reserved for compensating a wallet created by + /// an operation that has not completed; it must never be used to erase an established + /// wallet by owner lookup alone. + async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: WalletLookupKey<'_>, + ) -> Result; + + /// Delete a wallet only when its complete persisted snapshot still matches + /// the snapshot captured by the compensating operation and it has no + /// financial, usage, or redemption references. Unlike the zero-balance + /// helper above, this is intentionally suitable for rolling back an + /// imported wallet whose snapshot contains a non-zero balance. + async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result; + + /// Restore an existing wallet to its pre-import snapshot only when the + /// persisted row still exactly matches the post-import snapshot captured + /// by the same operation. Implementations must perform the compare and + /// update while holding the wallet row/lifecycle lock; a mismatch, missing + /// row, or owner mismatch returns `false` without changing the wallet. + async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result; + + /// Remove the exact user wallet created during an incomplete authentication provisioning + /// flow, but only when it is still an untouched initial wallet. Requiring both the wallet id + /// and owner prevents a failed initializer from deleting a wallet created concurrently by a + /// different operation. + async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result; + async fn create_wallet_recharge_order( &self, input: CreateWalletRechargeOrderInput, ) -> Result; + async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result, crate::DataLayerError>; + + /// Replace only the Stripe client-secret envelope when the complete + /// payment-order identity and observed gateway response still match. + async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result { + let _ = input; + Ok(false) + } + + /// Mark the currently claimed checkout placeholder as failed. The + /// expected expiry makes this conditional on the caller's claim, so a + /// delayed error cannot invalidate a newer retry. + async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result, crate::DataLayerError> { + let _ = input; + Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout failure handling is not available".to_string(), + )) + } + + /// Atomically take over a failed, expired, or timed-out checkout + /// placeholder. Implementations must lock/compare the current claim before + /// writing the new token, so at most one caller proceeds to the provider. + async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result, crate::DataLayerError> { + let _ = input; + Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim is not available".to_string(), + )) + } + async fn create_plan_purchase_order( &self, input: CreatePlanPurchaseOrderInput, @@ -1011,6 +2912,11 @@ pub trait WalletWriteRepository: Send + Sync { crate::DataLayerError, >; + async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result, crate::DataLayerError>; + async fn complete_admin_wallet_refund( &self, input: CompleteAdminWalletRefundInput, @@ -1076,11 +2982,180 @@ impl WalletRepository for T where T: WalletReadRepository + WalletWriteReposi #[cfg(test)] mod tests { use super::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, - redeem_code_refundable_amount, StoredWalletSnapshot, + canonicalize_payment_method, payment_callback_amount_matches_order, + payment_callback_method_matches_order, payment_callback_provider_matches_order, + payment_order_is_failed_wallet_checkout_placeholder, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_credits_recharge_balance, + redeem_code_payment_method, redeem_code_refundable_amount, validate_manual_wallet_recharge, + validate_payment_callback_provider_binding, validate_payment_order_credit_amounts, + validate_payment_provider_channel_binding, validate_plan_purchase_order_input, + validate_plan_wallet_credit_entitlements, validate_wallet_recharge_order_input, + wallet_recharge_order_created_at_unix_secs, wallet_refund_proof_is_success, + CompareAndSwapPaymentOrderStripeClientSecretInput, CreatePlanPurchaseOrderInput, + CreateWalletRechargeOrderInput, ProcessPaymentCallbackInput, StoredAdminPaymentOrder, + StoredWalletSnapshot, }; use crate::repository::settlement::UsageSettlementInput; + fn stripe_secret_cas_fixture() -> ( + StoredAdminPaymentOrder, + CompareAndSwapPaymentOrderStripeClientSecretInput, + ) { + let legacy = "gAAAAABlegacy-ciphertext"; + let gateway_response = serde_json::json!({ + "gateway": "stripe", + "publishable_key": "pk_test_public", + "provider_label": "Stripe", + "nested_unrelated": {"keep": true}, + "_stripe_client_secret_encrypted": legacy, + }); + let order = StoredAdminPaymentOrder { + id: "order-cas".to_string(), + order_no: "po-cas".to_string(), + wallet_id: "wallet-cas".to_string(), + user_id: Some("user-cas".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + refunded_amount_usd: 0.0, + refundable_amount_usd: 10.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + order_kind: "wallet_recharge".to_string(), + gateway_order_id: Some("pi-cas".to_string()), + gateway_response: Some(gateway_response.clone()), + status: "pending".to_string(), + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(4_102_444_800), + }; + let input = CompareAndSwapPaymentOrderStripeClientSecretInput { + order_id: order.id.clone(), + order_no: order.order_no.clone(), + wallet_id: order.wallet_id.clone(), + user_id: order.user_id.clone(), + payment_method: order.payment_method.clone(), + payment_provider: order.payment_provider.clone(), + order_kind: order.order_kind.clone(), + gateway_order_id: order.gateway_order_id.clone(), + expected_status: order.status.clone(), + expected_expires_at_unix_secs: order.expires_at_unix_secs, + expected_gateway_response: gateway_response, + expected_client_secret_encrypted: legacy.to_string(), + replacement_client_secret_encrypted: concat!( + "aether-payment-order-stripe-client-secret-v2:", + "aether-runtime-secret-v1:gAAAAABreplacement" + ) + .to_string(), + }; + (order, input) + } + + #[test] + fn payment_order_debug_output_redacts_gateway_and_stripe_secret_material() { + let (order, input) = stripe_secret_cas_fixture(); + let order_debug = format!("{order:?}"); + let input_debug = format!("{input:?}"); + + assert!(order_debug.contains("[REDACTED]")); + assert!(input_debug.contains("[REDACTED]")); + for secret in [ + "gAAAAABlegacy-ciphertext", + "gAAAAABreplacement", + "nested_unrelated", + ] { + assert!(!order_debug.contains(secret), "order debug leaked {secret}"); + assert!(!input_debug.contains(secret), "CAS debug leaked {secret}"); + } + } + + #[test] + fn stripe_secret_cas_replaces_only_the_exact_observed_field() { + let (order, input) = stripe_secret_cas_fixture(); + let replacement = payment_order_stripe_client_secret_cas_replacement(&order, &input) + .expect("valid replacement should be accepted") + .expect("complete observed row should match"); + + assert_eq!( + replacement["_stripe_client_secret_encrypted"].as_str(), + Some(input.replacement_client_secret_encrypted.as_str()) + ); + assert_eq!(replacement["publishable_key"], "pk_test_public"); + assert_eq!(replacement["provider_label"], "Stripe"); + assert_eq!( + replacement["nested_unrelated"], + serde_json::json!({"keep": true}) + ); + assert_eq!(replacement.as_object().map(|object| object.len()), Some(5)); + } + + #[test] + fn stripe_secret_cas_rejects_stale_json_ciphertext_and_identity() { + let (order, input) = stripe_secret_cas_fixture(); + + let mut stale_json = input.clone(); + stale_json.expected_gateway_response["provider_label"] = serde_json::json!("changed"); + assert_eq!( + payment_order_stripe_client_secret_cas_replacement(&order, &stale_json) + .expect("stale JSON is a normal miss"), + None + ); + + let mut stale_ciphertext = input.clone(); + stale_ciphertext.expected_client_secret_encrypted = "gAAAAABstale".to_string(); + assert_eq!( + payment_order_stripe_client_secret_cas_replacement(&order, &stale_ciphertext) + .expect("stale ciphertext is a normal miss"), + None + ); + + for foreign in [ + CompareAndSwapPaymentOrderStripeClientSecretInput { + order_no: "po-foreign".to_string(), + ..input.clone() + }, + CompareAndSwapPaymentOrderStripeClientSecretInput { + user_id: Some("user-foreign".to_string()), + ..input.clone() + }, + CompareAndSwapPaymentOrderStripeClientSecretInput { + order_kind: "plan_purchase".to_string(), + ..input.clone() + }, + CompareAndSwapPaymentOrderStripeClientSecretInput { + payment_provider: Some("STRIPE".to_string()), + ..input.clone() + }, + ] { + assert_eq!( + payment_order_stripe_client_secret_cas_replacement(&order, &foreign) + .expect("identity mismatch is a normal miss"), + None + ); + } + } + + #[test] + fn stripe_secret_cas_rejects_unknown_or_malformed_replacement_envelopes() { + let (order, input) = stripe_secret_cas_fixture(); + for replacement in [ + "aether-payment-order-stripe-client-secret-v3:unknown", + "aether-payment-order-stripe-client-secret-v2:", + "aether-payment-order-stripe-client-secret-v2:aether-runtime-secret-v2:unknown", + "aether-payment-order-stripe-client-secret-v2:aether-runtime-secret-v1:\n", + ] { + let invalid = CompareAndSwapPaymentOrderStripeClientSecretInput { + replacement_client_secret_encrypted: replacement.to_string(), + ..input.clone() + }; + assert!(payment_order_stripe_client_secret_cas_replacement(&order, &invalid).is_err()); + } + } + #[test] fn rejects_invalid_wallet_snapshot() { assert!(StoredWalletSnapshot::new( @@ -1133,4 +3208,803 @@ mod tests { assert_eq!(redeem_code_payment_method("recharge"), "card_code"); assert_eq!(redeem_code_refundable_amount("recharge", 8.5), 8.5); } + + #[test] + fn plan_wallet_credit_entitlements_fail_closed_on_malformed_entries() { + assert!( + validate_plan_wallet_credit_entitlements(&serde_json::json!([ + {"type": "wallet_credit", "amount_usd": 5.0, "balance_bucket": "gift"} + ])) + .is_ok() + ); + for invalid in [ + serde_json::json!([{"type": "wallet_credit"}]), + serde_json::json!([{"type": "wallet_credit", "amount_usd": 0.0}]), + serde_json::json!([{"type": "wallet_credit", "amount_usd": -1.0}]), + serde_json::json!([{"type": "wallet_credit", "amount_usd": 5.0, "balance_bucket": "unknown"}]), + serde_json::json!([{"type": "wallet_credit", "amount_usd": 5.0, "balance_bucket": 7}]), + serde_json::json!(["wallet_credit"]), + serde_json::json!({"type": "wallet_credit", "amount_usd": 5.0}), + ] { + assert!( + validate_plan_wallet_credit_entitlements(&invalid).is_err(), + "malformed wallet credit should be rejected: {invalid}" + ); + } + } + + #[test] + fn plan_purchase_input_rejects_invalid_money_identity_and_snapshot_fields() { + let valid = CreatePlanPurchaseOrderInput { + preferred_wallet_id: None, + user_id: "user-1".to_string(), + amount_usd: 1.0, + pay_amount: 7.2, + pay_currency: "CNY".to_string(), + exchange_rate: 7.2, + payment_method: "alipay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "gateway-1".to_string(), + gateway_response: serde_json::json!({}), + order_no: "order-1".to_string(), + product_id: "plan-1".to_string(), + product_snapshot: serde_json::json!({ + "id": "plan-1", + "duration_unit": "month", + "duration_value": 1, + "purchase_limit_scope": "active_period", + "entitlements": [{"type": "daily_quota", "daily_quota_usd": 1.0}] + }), + expires_at_unix_secs: 4_102_444_800, + }; + assert!(validate_plan_purchase_order_input(&valid).is_ok()); + + let mut invalid = valid.clone(); + invalid.amount_usd = f64::NAN; + assert!(validate_plan_purchase_order_input(&invalid).is_err()); + invalid = valid.clone(); + invalid.product_snapshot["id"] = serde_json::json!("another-plan"); + assert!(validate_plan_purchase_order_input(&invalid).is_err()); + invalid = valid.clone(); + invalid.gateway_response = serde_json::json!("provider payload"); + assert!(validate_plan_purchase_order_input(&invalid).is_err()); + + for (duration_unit, duration_value) in + [("day", i64::MAX), ("month", i64::MAX), ("year", i64::MAX)] + { + invalid = valid.clone(); + invalid.product_snapshot["duration_unit"] = serde_json::json!(duration_unit); + invalid.product_snapshot["duration_value"] = serde_json::json!(duration_value); + assert!( + validate_plan_purchase_order_input(&invalid).is_err(), + "overflowing {duration_unit} duration must be rejected" + ); + } + } + + #[test] + fn wallet_recharge_input_rejects_invalid_channels_and_numbers() { + let valid = CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: "user-1".to_string(), + amount_usd: 1.0, + pay_amount: Some(7.2), + pay_currency: Some("CNY".to_string()), + exchange_rate: Some(7.2), + payment_method: "wxpay".to_string(), + payment_provider: Some("wxpay".to_string()), + payment_channel: Some("native".to_string()), + gateway_order_id: "gateway-1".to_string(), + gateway_response: serde_json::json!({}), + order_no: "order-1".to_string(), + expires_at_unix_secs: 4_102_444_800, + }; + assert!(validate_wallet_recharge_order_input(&valid).is_ok()); + + let mut invalid = valid.clone(); + invalid.payment_channel = Some("app".to_string()); + assert!(validate_wallet_recharge_order_input(&invalid).is_err()); + invalid = valid.clone(); + invalid.pay_currency = Some("C1Y".to_string()); + assert!(validate_wallet_recharge_order_input(&invalid).is_err()); + invalid = valid.clone(); + invalid.amount_usd = f64::INFINITY; + assert!(validate_wallet_recharge_order_input(&invalid).is_err()); + + // Legacy checkout placeholders may omit the explicit channel; keep + // those rows readable while still checking the provider/method pair. + invalid = valid; + invalid.payment_channel = None; + assert!(validate_wallet_recharge_order_input(&invalid).is_ok()); + + invalid.payment_provider = None; + invalid.payment_channel = Some("native".to_string()); + assert!(validate_wallet_recharge_order_input(&invalid).is_err()); + + invalid.payment_channel = None; + assert!(validate_wallet_recharge_order_input(&invalid).is_ok()); + } + + #[test] + fn zero_value_credit_is_limited_to_admin_grants() { + assert!(validate_payment_order_credit_amounts( + "plan_purchase", + "admin_grant", + Some("admin"), + Some("manual"), + 0.0, + Some(0.0), + ) + .is_ok()); + assert!(validate_payment_order_credit_amounts( + "wallet_recharge", + "admin_grant", + Some("admin"), + Some("manual"), + 0.0, + Some(0.0), + ) + .is_err()); + assert!(validate_payment_order_credit_amounts( + "plan_purchase", + "stripe", + Some("stripe"), + Some("card"), + 0.0, + Some(0.0), + ) + .is_err()); + } + + #[test] + fn payment_method_namespace_is_trimmed_lowercase_and_bounded() { + assert_eq!( + canonicalize_payment_method(" Admin_Manual "), + Ok("admin_manual".to_string()) + ); + assert!(canonicalize_payment_method(" ").is_err()); + assert!(canonicalize_payment_method(&"X".repeat(65)).is_err()); + assert!(canonicalize_payment_method("epay/card").is_err()); + } + + #[test] + fn official_payment_callbacks_require_an_explicit_matching_provider() { + for provider in ["alipay", "wxpay", "stripe", "epay"] { + assert!(validate_payment_callback_provider_binding(provider, Some(provider)).is_ok()); + assert!(validate_payment_callback_provider_binding(provider, None).is_err()); + assert!(validate_payment_callback_provider_binding(provider, Some("manual")).is_err()); + } + assert!(validate_payment_callback_provider_binding("alipay", Some("epay")).is_ok()); + assert!(validate_payment_callback_provider_binding("wxpay", Some("epay")).is_ok()); + assert!(validate_payment_callback_provider_binding("manual", None).is_ok()); + } + + #[test] + fn payment_provider_channel_binding_rejects_cross_gateway_orders() { + assert!( + validate_payment_provider_channel_binding("stripe", Some("stripe"), Some("card")) + .is_ok() + ); + assert!( + validate_payment_provider_channel_binding("wxpay", Some("wxpay"), Some("native")) + .is_ok() + ); + assert!(validate_payment_provider_channel_binding( + "stripe", + Some("stripe"), + Some("wechat_pay") + ) + .is_ok()); + assert!( + validate_payment_provider_channel_binding("alipay", Some("epay"), Some("alipay")) + .is_ok() + ); + assert!( + validate_payment_provider_channel_binding("epay", Some("epay"), Some("wxpay")).is_ok() + ); + assert!( + validate_payment_provider_channel_binding("epay", Some("epay"), Some("qqpay")).is_ok() + ); + assert!(validate_payment_provider_channel_binding( + "admin_grant", + Some("admin"), + Some("manual") + ) + .is_ok()); + for invalid in [ + ("stripe", Some("epay"), Some("card")), + ("alipay", Some("epay"), Some("wxpay")), + ("alipay", Some("alipay"), Some("native")), + ("stripe", Some("stripe"), None), + ("wxpay", Some("wxpay"), Some("app")), + ("stripe", Some("stripe"), Some("bank_transfer")), + ("admin_grant", Some("admin"), Some("card")), + ] { + assert!( + validate_payment_provider_channel_binding(invalid.0, invalid.1, invalid.2).is_err(), + "invalid payment binding should be rejected: {:?}", + invalid + ); + } + } + + #[test] + fn epay_callback_accepts_legacy_method_aliases_only_with_epay_binding() { + assert!(payment_callback_method_matches_order( + "alipay", + Some("epay"), + "epay", + Some("epay") + )); + assert!(payment_callback_method_matches_order( + "epay", + Some("epay"), + "wxpay", + Some("epay") + )); + assert!(!payment_callback_method_matches_order( + "alipay", + Some("alipay"), + "epay", + Some("epay") + )); + assert!(!payment_callback_method_matches_order( + "manual", + Some("epay"), + "epay", + Some("epay") + )); + } + + #[test] + fn legacy_epay_channel_orders_accept_only_epay_provider_callbacks() { + assert!(payment_callback_provider_matches_order( + "alipay", + None, + "epay", + Some("epay") + )); + assert!(payment_callback_provider_matches_order( + "wxpay", + None, + "alipay", + Some("epay") + )); + assert!(!payment_callback_provider_matches_order( + "alipay", + None, + "alipay", + Some("stripe") + )); + assert!(!payment_callback_provider_matches_order( + "manual", + None, + "epay", + Some("epay") + )); + assert!(!payment_callback_provider_matches_order( + "alipay", + None, + "alipay", + Some("alipay") + )); + } + + #[test] + fn official_payment_callbacks_require_complete_settlement_bindings() { + let valid = ProcessPaymentCallbackInput { + payment_method: " Stripe ".to_string(), + payment_provider: Some("STRIPE".to_string()), + payment_channel: Some("CARD".to_string()), + callback_key: "stripe:event-1".to_string(), + order_no: Some("po_1".to_string()), + gateway_order_id: Some("pi_1".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "hash-1".to_string(), + payload: serde_json::json!({}), + signature_valid: true, + }; + + let mut canonical = valid.clone(); + canonical + .canonicalize_and_validate() + .expect("complete official callback should validate"); + assert_eq!(canonical.payment_method, "stripe"); + assert_eq!(canonical.payment_provider.as_deref(), Some("stripe")); + assert_eq!(canonical.payment_channel.as_deref(), Some("card")); + + for mut incomplete in [ + { + let mut value = valid.clone(); + value.payment_channel = None; + value + }, + { + let mut value = valid.clone(); + value.order_no = None; + value + }, + { + let mut value = valid.clone(); + value.gateway_order_id = None; + value + }, + { + let mut value = valid.clone(); + value.pay_amount = None; + value + }, + { + let mut value = valid.clone(); + value.pay_currency = None; + value + }, + ] { + assert!(incomplete.canonicalize_and_validate().is_err()); + } + } + + #[test] + fn callback_amounts_require_provider_settlement_to_match_order() { + assert!(payment_callback_amount_matches_order( + 10.0, + Some(72.0), + Some("CNY"), + Some(7.2), + 10.0, + Some(72.0), + )); + assert!(payment_callback_amount_matches_order( + 10.0, + Some(72.0), + Some("CNY"), + Some(7.2), + 10.0000005, + Some(72.0000005), + )); + // The gateway's USD presentation may include a fee or use a + // provider-side conversion. The repository credits the order's + // locked USD amount and therefore only binds the signed settlement. + assert!(payment_callback_amount_matches_order( + 10.0, + Some(72.0), + Some("CNY"), + Some(7.2), + 11.0, + Some(72.0), + )); + assert!(!payment_callback_amount_matches_order( + 10.0, + Some(72.0), + Some("CNY"), + Some(7.2), + 10.0, + Some(71.0), + )); + } + + #[test] + fn legacy_callback_amounts_are_reconstructed_only_from_order_terms() { + assert!(payment_callback_amount_matches_order( + 10.0, + None, + Some("CNY"), + Some(7.2), + 10.0, + Some(72.0), + )); + assert!(!payment_callback_amount_matches_order( + 10.0, + None, + Some("CNY"), + Some(7.2), + 10.0, + Some(72.01), + )); + + // A stale historical USD rate must not reinterpret a USD order as + // CNY. The currency itself determines the effective rate. + assert!(payment_callback_amount_matches_order( + 10.0, + None, + Some("USD"), + Some(7.2), + 10.0, + Some(10.0), + )); + assert!(!payment_callback_amount_matches_order( + 10.0, + None, + Some("USD"), + Some(7.2), + 10.0, + Some(72.0), + )); + + // An unknown currency/rate cannot be completed from provider data. + assert!(!payment_callback_amount_matches_order( + 10.0, + None, + None, + Some(7.2), + 10.0, + Some(72.0), + )); + assert!(!payment_callback_amount_matches_order( + 10.0, + None, + Some("CNY"), + None, + 10.0, + Some(72.0), + )); + } + + #[test] + fn payment_callback_gateway_response_projection_drops_provider_payload() { + let input = ProcessPaymentCallbackInput { + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + callback_key: "evt_1".to_string(), + order_no: Some("po_1".to_string()), + gateway_order_id: Some("pi_1".to_string()), + amount_usd: 10.0, + pay_amount: Some(10.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payload_hash: "hash-1".to_string(), + payload: serde_json::json!({ + "client_secret": "pi_1_secret_replayable", + "customer": {"email": "payer@example.com"}, + "authorization": "Bearer upstream-secret", + }), + signature_valid: true, + }; + + let projected = input.gateway_response_projection("po_1", Some("pi_1")); + assert_eq!( + projected, + serde_json::json!({ + "gateway": "stripe", + "payment_provider": "stripe", + "payment_channel": "card", + "order_no": "po_1", + "gateway_order_id": "pi_1", + "amount_usd": 10.0, + "pay_amount": 10.0, + "pay_currency": "USD", + "exchange_rate": 1.0, + "signature_valid": true, + }) + ); + let encoded = projected.to_string(); + for forbidden in [ + "client_secret", + "replayable", + "payer@example.com", + "authorization", + "upstream-secret", + ] { + assert!(!encoded.contains(forbidden), "persisted {forbidden}"); + } + } + + #[test] + fn wallet_recharge_gateway_projection_is_an_allowlist() { + let projected = project_wallet_recharge_gateway_response(&serde_json::json!({ + "gateway": "stripe", + "instructions": "confirm payment", + "payment_method_types": ["card", {"secret": "drop"}, ""], + "payment_params": { + "pid": "merchant", + "sign": "signed", + "nested": {"credential": "drop"} + }, + "_stripe_client_secret_encrypted": "enc:v1:secret", + "client_secret": "pi_1_secret_raw", + "customer": {"email": "payer@example.com"}, + "unknown": "drop", + })) + .expect("checkout object should project"); + + assert_eq!( + projected, + serde_json::json!({ + "gateway": "stripe", + "instructions": "confirm payment", + "payment_method_types": ["card"], + "payment_params": {"pid": "merchant", "sign": "signed"}, + "_stripe_client_secret_encrypted": "enc:v1:secret", + "order_kind": "wallet_recharge", + }) + ); + let encoded = projected.to_string(); + for forbidden in [ + "pi_1_secret_raw", + "payer@example.com", + "credential", + "unknown", + ] { + assert!(!encoded.contains(forbidden), "projected {forbidden}"); + } + assert!(project_wallet_recharge_gateway_response(&serde_json::json!(null)).is_err()); + } + + #[test] + fn plan_gateway_projection_is_an_allowlist_without_wallet_marker() { + let projected = project_wallet_gateway_response(&serde_json::json!({ + "gateway": "stripe", + "gateway_order_id": "pi_1", + "intent_id": "pi_1", + "publishable_key": "pk_test_public", + "client_secret": "pi_1_secret_raw", + "_stripe_client_secret_encrypted": "enc:v1:secret", + "order_kind": "wallet_recharge", + "product_id": "plan-secret-overwrite", + "customer": {"email": "payer@example.com"}, + "provider_private_token": "drop", + })) + .expect("plan checkout object should project"); + + assert_eq!( + projected, + serde_json::json!({ + "gateway": "stripe", + "gateway_order_id": "pi_1", + "intent_id": "pi_1", + "publishable_key": "pk_test_public", + "_stripe_client_secret_encrypted": "enc:v1:secret", + }) + ); + assert_ne!( + projected + .get("order_kind") + .and_then(serde_json::Value::as_str), + Some("wallet_recharge") + ); + let encoded = projected.to_string(); + for forbidden in [ + "pi_1_secret_raw", + "payer@example.com", + "provider_private_token", + "plan-secret-overwrite", + ] { + assert!(!encoded.contains(forbidden), "projected {forbidden}"); + } + } + + #[test] + fn failed_wallet_checkout_placeholder_is_the_only_recoverable_failed_order() { + let placeholder = serde_json::json!({ + "order_kind": "wallet_recharge", + "integration_status": "checkout_failed", + "gateway": "stripe", + "gateway_order_id": "order-1", + "failure_reason": "provider response lost", + }); + assert!(payment_order_is_failed_wallet_checkout_placeholder( + "failed", + "wallet_recharge", + Some(&placeholder), + )); + assert!(!payment_order_is_failed_wallet_checkout_placeholder( + "pending", + "wallet_recharge", + Some(&placeholder), + )); + assert!(!payment_order_is_failed_wallet_checkout_placeholder( + "failed", + "plan_purchase", + Some(&placeholder), + )); + assert!(!payment_order_is_failed_wallet_checkout_placeholder( + "failed", + "wallet_recharge", + None, + )); + + for evidence_key in [ + "payment_url", + "qr_code", + "code_url", + "h5_url", + "jsapi", + "client_secret", + "_stripe_client_secret_encrypted", + "intent_id", + ] { + let mut with_evidence = placeholder.clone(); + with_evidence[evidence_key] = serde_json::json!("provider-evidence"); + assert!( + !payment_order_is_failed_wallet_checkout_placeholder( + "failed", + "wallet_recharge", + Some(&with_evidence), + ), + "provider evidence {evidence_key} must keep a failed order non-creditable" + ); + } + } + + #[test] + fn uncertain_wallet_checkout_is_callback_settleable_but_not_reclaimable() { + let uncertain = super::wallet_recharge_checkout_uncertain_response( + Some(&serde_json::json!({ + "gateway": "stripe", + "gateway_order_id": "po_1", + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending", + "checkout_claim_token": "claim-1", + "checkout_claimed_at_unix_secs": 1, + })), + "response lost after provider acceptance", + 2, + ); + assert_eq!(uncertain["integration_status"], "checkout_uncertain"); + assert!(super::payment_order_is_failed_wallet_checkout_placeholder( + "failed", + "wallet_recharge", + Some(&uncertain), + )); + let order = StoredAdminPaymentOrder { + id: "order-uncertain".to_string(), + order_no: "po_1".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 1.0, + pay_amount: Some(1.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + order_kind: "wallet_recharge".to_string(), + gateway_order_id: Some("po_1".to_string()), + gateway_response: Some(uncertain), + status: "failed".to_string(), + created_at_unix_ms: 1, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: Some(u64::MAX), + }; + assert!(!super::wallet_recharge_order_is_reclaimable_placeholder( + &order, 10_000, + )); + } + + #[test] + fn checkout_lease_fallback_accepts_legacy_seconds_and_milliseconds() { + let base = StoredAdminPaymentOrder { + id: "order-1".to_string(), + order_no: "order-no-1".to_string(), + wallet_id: "wallet-1".to_string(), + user_id: Some("user-1".to_string()), + amount_usd: 1.0, + pay_amount: Some(1.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + refunded_amount_usd: 0.0, + refundable_amount_usd: 0.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + order_kind: "wallet_recharge".to_string(), + gateway_order_id: Some("order-no-1".to_string()), + gateway_response: None, + status: "pending".to_string(), + created_at_unix_ms: 1_700_000_000, + paid_at_unix_secs: None, + credited_at_unix_secs: None, + expires_at_unix_secs: None, + }; + assert_eq!( + wallet_recharge_order_created_at_unix_secs(&base), + 1_700_000_000 + ); + + let mut millis = base.clone(); + millis.created_at_unix_ms = 1_700_000_000_000; + assert_eq!( + wallet_recharge_order_created_at_unix_secs(&millis), + 1_700_000_000 + ); + } + + #[test] + fn canonicalizes_and_rejects_refund_provenance_tampering() { + let fields = super::canonicalize_wallet_refund_fields( + Some("order-1"), + None, + None, + None, + Some("stripe"), + ) + .expect("order fields should canonicalize"); + assert_eq!(fields.source_type, "payment_order"); + assert_eq!(fields.source_id.as_deref(), Some("order-1")); + assert_eq!(fields.refund_mode, "original_channel"); + + assert!(super::canonicalize_wallet_refund_fields( + Some("order-1"), + Some("wallet_balance"), + None, + None, + Some("stripe"), + ) + .is_err()); + assert!(super::canonicalize_wallet_refund_fields( + None, + None, + Some("forged-source"), + None, + None, + ) + .is_err()); + } + + #[test] + fn manual_wallet_recharge_requires_safe_finite_arithmetic() { + assert_eq!( + validate_manual_wallet_recharge(5.0, 10.0, 2.0, 20.0), + Ok((15.0, 25.0)) + ); + assert_eq!( + validate_manual_wallet_recharge(5.0, -10.0, 2.0, 20.0), + Ok((-5.0, 25.0)) + ); + + for amount_usd in [0.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + assert!( + validate_manual_wallet_recharge(amount_usd, 10.0, 2.0, 20.0).is_err(), + "invalid recharge amount should be rejected: {amount_usd:?}" + ); + } + assert!(validate_manual_wallet_recharge(f64::MAX, f64::MAX, 0.0, f64::MAX).is_err()); + } + + #[test] + fn payment_order_refund_amounts_require_a_consistent_split() { + assert!(payment_order_refund_amounts_are_consistent(10.0, 0.0, 10.0)); + assert!(payment_order_refund_amounts_are_consistent(10.0, 2.5, 7.5)); + assert!(payment_order_refund_amounts_are_consistent( + 10.0, 2.5, 7.5000005 + )); + for values in [ + (10.0, 0.0, 5.0), + (10.0, -1.0, 11.0), + (10.0, 11.0, 0.0), + (f64::NAN, 0.0, 0.0), + (10.0, f64::INFINITY, 0.0), + ] { + assert!( + !payment_order_refund_amounts_are_consistent(values.0, values.1, values.2), + "malformed refund split should be rejected: {values:?}" + ); + } + } + + #[test] + fn refund_proof_success_detection_only_accepts_terminal_status() { + assert!(wallet_refund_proof_is_success(&serde_json::json!({ + "status": "success" + }))); + assert!(wallet_refund_proof_is_success(&serde_json::json!({ + "gateway_refund": { "status": "succeeded" } + }))); + assert!(!wallet_refund_proof_is_success(&serde_json::json!({ + "status": "processing" + }))); + assert!(!wallet_refund_proof_is_success(&serde_json::json!({ + "gateway_refund": { "status": "processing" } + }))); + assert!(!wallet_refund_proof_is_success(&serde_json::json!({ + "status": "success", + "gateway_refund": { "status": "processing" } + }))); + } } diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql index 44f1df04e..1f3837fad 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql @@ -265,11 +265,11 @@ CREATE TABLE IF NOT EXISTS public.dimension_collectors ( CREATE TABLE IF NOT EXISTS public.gemini_file_mappings ( id character varying(36) NOT NULL, - file_name character varying(255) NOT NULL, + file_name character varying(512) NOT NULL, key_id character varying(36) NOT NULL, user_id character varying(36), - display_name character varying(255), - mime_type character varying(100), + display_name character varying(512), + mime_type character varying(255), source_hash character varying(64), created_at timestamp with time zone NOT NULL, expires_at timestamp with time zone NOT NULL @@ -305,6 +305,7 @@ CREATE TABLE IF NOT EXISTS public.global_models ( CREATE TABLE IF NOT EXISTS public.ldap_configs ( id integer NOT NULL, + singleton_key integer DEFAULT 1 NOT NULL, server_url character varying(255) NOT NULL, bind_dn text NOT NULL, bind_password_encrypted text, @@ -788,6 +789,7 @@ ALTER SEQUENCE public.proxy_node_events_id_seq OWNED BY public.proxy_node_events CREATE TABLE IF NOT EXISTS public.proxy_nodes ( id character varying(36) NOT NULL, + tunnel_generation character varying(64) NOT NULL, name character varying(100) NOT NULL, ip character varying(512) NOT NULL, port integer NOT NULL, @@ -802,7 +804,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( is_manual boolean DEFAULT false NOT NULL, proxy_url character varying(500), proxy_username character varying(255), - proxy_password character varying(500), + proxy_password text, created_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, updated_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, remote_config json, @@ -1275,6 +1277,7 @@ CREATE TABLE IF NOT EXISTS public.usage_counter_deltas ( request_id character varying(128) NOT NULL, kind character varying(64) NOT NULL, target_id text NOT NULL, + target_tunnel_generation character varying(64), request_count_delta bigint DEFAULT 0 NOT NULL, total_requests_delta bigint DEFAULT 0 NOT NULL, success_count_delta bigint DEFAULT 0 NOT NULL, @@ -1361,6 +1364,7 @@ CREATE TABLE IF NOT EXISTS public.user_preferences ( CREATE TABLE IF NOT EXISTS public.user_sessions ( id character varying(36) NOT NULL, user_id character varying(36) NOT NULL, + security_version bigint DEFAULT 0 NOT NULL, client_device_id character varying(128) NOT NULL, device_label character varying(120), device_type character varying(20) DEFAULT 'unknown'::character varying NOT NULL, @@ -1406,6 +1410,7 @@ CREATE TABLE IF NOT EXISTS public.users ( feature_settings jsonb, is_active boolean DEFAULT true NOT NULL, is_deleted boolean DEFAULT false NOT NULL, + security_version bigint DEFAULT 0 NOT NULL, created_at timestamp with time zone DEFAULT now() NOT NULL, updated_at timestamp with time zone DEFAULT now() NOT NULL, last_login_at timestamp with time zone, diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/003_constraints.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/003_constraints.sql index 8975f3bbf..06935f4f0 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/003_constraints.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/003_constraints.sql @@ -177,6 +177,36 @@ END $mig$; +-- +-- Name: ldap_configs ldap_configs_singleton_key_check; Type: CHECK CONSTRAINT; Schema: public; Owner: - +-- + +DO $mig$ BEGIN + ALTER TABLE ONLY public.ldap_configs + ADD CONSTRAINT ldap_configs_singleton_key_check CHECK (singleton_key = 1); +EXCEPTION + WHEN duplicate_object THEN NULL; + WHEN duplicate_table THEN NULL; + WHEN invalid_table_definition THEN NULL; +END $mig$; + + + +-- +-- Name: ldap_configs ldap_configs_singleton_key_key; Type: CONSTRAINT; Schema: public; Owner: - +-- + +DO $mig$ BEGIN + ALTER TABLE ONLY public.ldap_configs + ADD CONSTRAINT ldap_configs_singleton_key_key UNIQUE (singleton_key); +EXCEPTION + WHEN duplicate_object THEN NULL; + WHEN duplicate_table THEN NULL; + WHEN invalid_table_definition THEN NULL; +END $mig$; + + + -- -- Name: management_tokens management_tokens_pkey; Type: CONSTRAINT; Schema: public; Owner: - -- diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/004_indexes.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/004_indexes.sql index ecc3fb340..8c07e3379 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/004_indexes.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/004_indexes.sql @@ -101,6 +101,14 @@ CREATE INDEX IF NOT EXISTS idx_payment_orders_gateway_order_id ON public.payment +-- +-- Name: uq_payment_orders_payment_method_gateway_order_id; Type: INDEX; Schema: public; Owner: - +-- + +CREATE UNIQUE INDEX IF NOT EXISTS uq_payment_orders_payment_method_gateway_order_id ON public.payment_orders USING btree (payment_method, gateway_order_id); + + + -- -- Name: idx_payment_orders_kind_status; Type: INDEX; Schema: public; Owner: - -- diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/005_foreign_keys.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/005_foreign_keys.sql index cc8165317..c93a599b1 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/005_foreign_keys.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/005_foreign_keys.sql @@ -694,36 +694,6 @@ END $mig$; --- --- Name: user_referrals user_referrals_inviter_user_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - --- - -DO $mig$ BEGIN - ALTER TABLE ONLY public.user_referrals - ADD CONSTRAINT user_referrals_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -EXCEPTION - WHEN duplicate_object THEN NULL; - WHEN duplicate_table THEN NULL; - WHEN invalid_table_definition THEN NULL; -END $mig$; - - - --- --- Name: user_referrals user_referrals_invitee_user_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - --- - -DO $mig$ BEGIN - ALTER TABLE ONLY public.user_referrals - ADD CONSTRAINT user_referrals_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -EXCEPTION - WHEN duplicate_object THEN NULL; - WHEN duplicate_table THEN NULL; - WHEN invalid_table_definition THEN NULL; -END $mig$; - - - -- -- Name: user_referrals user_referrals_first_paid_order_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - -- @@ -754,36 +724,6 @@ END $mig$; --- --- Name: referral_rewards referral_rewards_inviter_user_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - --- - -DO $mig$ BEGIN - ALTER TABLE ONLY public.referral_rewards - ADD CONSTRAINT referral_rewards_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -EXCEPTION - WHEN duplicate_object THEN NULL; - WHEN duplicate_table THEN NULL; - WHEN invalid_table_definition THEN NULL; -END $mig$; - - - --- --- Name: referral_rewards referral_rewards_invitee_user_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - --- - -DO $mig$ BEGIN - ALTER TABLE ONLY public.referral_rewards - ADD CONSTRAINT referral_rewards_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -EXCEPTION - WHEN duplicate_object THEN NULL; - WHEN duplicate_table THEN NULL; - WHEN invalid_table_definition THEN NULL; -END $mig$; - - - -- -- Name: referral_rewards referral_rewards_source_order_id_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - -- diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/100_usage_capture.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/100_usage_capture.sql index 9e169b2c5..4d2b9f4e0 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/100_usage_capture.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/100_usage_capture.sql @@ -145,6 +145,84 @@ CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_schema_version CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_pricing_source ON public.usage_settlement_snapshots USING btree (billing_pricing_source); +CREATE TABLE IF NOT EXISTS public.usage_cost_reservations ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + reservation_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + reserved_cost_units bigint NOT NULL, + actual_cost_units bigint, + state character varying(20) NOT NULL, + reservation_expires_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + finalized_at timestamp with time zone, + created_at timestamp with time zone DEFAULT now() NOT NULL, + updated_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT usage_cost_reservations_pkey PRIMARY KEY (reservation_token), + CONSTRAINT usage_cost_reservations_subject_id_fkey + FOREIGN KEY (subject_id) + REFERENCES public.users(id) + ON DELETE CASCADE, + CONSTRAINT usage_cost_reservations_state_check + CHECK (state IN ('reserved', 'finalized', 'released')), + CONSTRAINT usage_cost_reservations_reserved_cost_units_check + CHECK (reserved_cost_units >= 0), + CONSTRAINT usage_cost_reservations_actual_cost_units_check + CHECK (actual_cost_units IS NULL OR actual_cost_units >= 0), + CONSTRAINT usage_cost_reservations_expiry_check + CHECK (reservation_expires_at > admitted_at), + CONSTRAINT usage_cost_reservations_retention_check + CHECK (retain_until >= reservation_expires_at), + CONSTRAINT usage_cost_reservations_lifecycle_check CHECK ( + (state = 'reserved' AND actual_cost_units IS NULL AND finalized_at IS NULL) + OR (state = 'finalized' AND actual_cost_units IS NOT NULL AND finalized_at IS NOT NULL) + OR (state = 'released' AND actual_cost_units IS NOT NULL + AND actual_cost_units = 0 AND finalized_at IS NOT NULL) + ) +); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_request_id_idx + ON public.usage_cost_reservations USING btree (request_id); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_subject_admitted_at_idx + ON public.usage_cost_reservations USING btree (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_reservation_expires_at_idx + ON public.usage_cost_reservations USING btree (reservation_expires_at); + +CREATE INDEX IF NOT EXISTS usage_cost_reservations_retain_until_token_idx + ON public.usage_cost_reservations USING btree (retain_until, reservation_token); + +CREATE TABLE IF NOT EXISTS public.usage_request_admissions ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + event_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + state character varying(20) NOT NULL, + released_at timestamp with time zone, + created_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT usage_request_admissions_pkey PRIMARY KEY (event_token), + CONSTRAINT usage_request_admissions_subject_id_fkey + FOREIGN KEY (subject_id) + REFERENCES public.users(id) + ON DELETE CASCADE, + CONSTRAINT usage_request_admissions_retention_check + CHECK (retain_until > admitted_at), + CONSTRAINT usage_request_admissions_state_check + CHECK (state IN ('active', 'released')), + CONSTRAINT usage_request_admissions_lifecycle_check CHECK ( + (state = 'active' AND released_at IS NULL) + OR (state = 'released' AND released_at IS NOT NULL AND released_at >= admitted_at) + ) +); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_subject_admitted_at_idx + ON public.usage_request_admissions USING btree (subject_id, admitted_at); + +CREATE INDEX IF NOT EXISTS usage_request_admissions_retain_until_token_idx + ON public.usage_request_admissions USING btree (retain_until, event_token); + CREATE INDEX IF NOT EXISTS idx_usage_settlement_dashboard_cover ON public.usage_settlement_snapshots USING btree (request_id) INCLUDE ( diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/160_routing_profiles.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/160_routing_profiles.sql index 233f8805e..0f2828628 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/160_routing_profiles.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/160_routing_profiles.sql @@ -4,6 +4,7 @@ CREATE TABLE IF NOT EXISTS public.routing_groups ( description text, enabled boolean DEFAULT true NOT NULL, is_system_default boolean DEFAULT false NOT NULL, + sort_order bigint DEFAULT 0 NOT NULL, config_json jsonb NOT NULL, version bigint DEFAULT 1 NOT NULL, created_at bigint NOT NULL, @@ -31,6 +32,8 @@ CREATE INDEX IF NOT EXISTS routing_groups_system_default_idx CREATE UNIQUE INDEX IF NOT EXISTS routing_groups_one_system_default_key ON public.routing_groups (is_system_default) WHERE is_system_default = TRUE; +CREATE INDEX IF NOT EXISTS routing_groups_enabled_sort_idx + ON public.routing_groups (enabled DESC, sort_order, name, id); CREATE TABLE IF NOT EXISTS public.routing_group_bindings ( id character varying(64) NOT NULL, @@ -85,4 +88,4 @@ BEGIN END IF; END $$; CREATE INDEX IF NOT EXISTS routing_group_versions_group_id_idx - ON public.routing_group_versions USING btree (group_id); \ No newline at end of file + ON public.routing_group_versions USING btree (group_id); diff --git a/crates/aether-data/runtime/schema/drivers/mysql/baseline/004_proxy_nodes.sql b/crates/aether-data/runtime/schema/drivers/mysql/baseline/004_proxy_nodes.sql index 5502b9925..18904c145 100644 --- a/crates/aether-data/runtime/schema/drivers/mysql/baseline/004_proxy_nodes.sql +++ b/crates/aether-data/runtime/schema/drivers/mysql/baseline/004_proxy_nodes.sql @@ -14,7 +14,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( is_manual TINYINT(1) NOT NULL DEFAULT 0, proxy_url VARCHAR(500), proxy_username VARCHAR(255), - proxy_password VARCHAR(500), + proxy_password TEXT, created_at BIGINT NOT NULL, updated_at BIGINT NOT NULL, remote_config TEXT, diff --git a/crates/aether-data/runtime/schema/drivers/postgres/baseline/001_types_and_tables.sql b/crates/aether-data/runtime/schema/drivers/postgres/baseline/001_types_and_tables.sql index 999b93d56..a1f597f3b 100644 --- a/crates/aether-data/runtime/schema/drivers/postgres/baseline/001_types_and_tables.sql +++ b/crates/aether-data/runtime/schema/drivers/postgres/baseline/001_types_and_tables.sql @@ -706,7 +706,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( is_manual boolean DEFAULT false NOT NULL, proxy_url character varying(500), proxy_username character varying(255), - proxy_password character varying(500), + proxy_password text, created_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, updated_at timestamp with time zone DEFAULT CURRENT_TIMESTAMP NOT NULL, remote_config json, diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/001_identity.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/001_identity.sql index 5b1f0a6bc..fd7a01c27 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/001_identity.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/001_identity.sql @@ -12,6 +12,7 @@ CREATE TABLE IF NOT EXISTS users ( `email_verified` TINYINT(1) NOT NULL DEFAULT 0, `is_active` TINYINT(1) NOT NULL DEFAULT 1, `is_deleted` TINYINT(1) NOT NULL DEFAULT 0, + `security_version` BIGINT NOT NULL DEFAULT 0, `allowed_models` JSON, `allowed_models_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted', `allowed_providers` JSON, @@ -193,6 +194,7 @@ CREATE TABLE IF NOT EXISTS user_preferences ( CREATE TABLE IF NOT EXISTS user_sessions ( `id` VARCHAR(64) NOT NULL, `user_id` VARCHAR(64) NOT NULL, + `security_version` BIGINT NOT NULL DEFAULT 0, `client_device_id` VARCHAR(128) NOT NULL, `device_label` VARCHAR(120), `device_type` VARCHAR(20) NOT NULL DEFAULT 'unknown', diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/002_provider_catalog.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/002_provider_catalog.sql index f9e96ec0d..f926876ae 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/002_provider_catalog.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/002_provider_catalog.sql @@ -391,6 +391,7 @@ CREATE TABLE IF NOT EXISTS routing_groups ( `description` LONGTEXT, `enabled` TINYINT(1) NOT NULL DEFAULT 1, `is_system_default` TINYINT(1) NOT NULL DEFAULT 0, + `sort_order` BIGINT NOT NULL DEFAULT 0, `config_json` JSON NOT NULL, `version` BIGINT NOT NULL DEFAULT 1, `created_at` BIGINT NOT NULL, @@ -398,7 +399,8 @@ CREATE TABLE IF NOT EXISTS routing_groups ( `published_at` BIGINT, PRIMARY KEY (`id`), UNIQUE KEY routing_groups_name_key (`name`), - KEY routing_groups_system_default_idx (`is_system_default`, `enabled`) + KEY routing_groups_system_default_idx (`is_system_default`, `enabled`), + KEY routing_groups_enabled_sort_idx (`enabled`, `sort_order`, `name`, `id`) ); CREATE TABLE IF NOT EXISTS routing_group_bindings ( diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/003_auth_config.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/003_auth_config.sql index 07495b0b1..7ddf3f6a8 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/003_auth_config.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/003_auth_config.sql @@ -45,6 +45,7 @@ CREATE TABLE IF NOT EXISTS oauth_providers ( CREATE TABLE IF NOT EXISTS ldap_configs ( `id` BIGINT NOT NULL AUTO_INCREMENT, + `singleton_key` INT NOT NULL DEFAULT 1, `server_url` VARCHAR(255) NOT NULL, `bind_dn` LONGTEXT NOT NULL, `bind_password_encrypted` LONGTEXT, @@ -59,7 +60,8 @@ CREATE TABLE IF NOT EXISTS ldap_configs ( `connect_timeout` INT NOT NULL DEFAULT 10, `created_at` BIGINT NOT NULL, `updated_at` BIGINT NOT NULL, - PRIMARY KEY (`id`) + PRIMARY KEY (`id`), + UNIQUE KEY ldap_configs_singleton_key_key (`singleton_key`) ); CREATE TABLE IF NOT EXISTS user_oauth_links ( diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/004_proxy_nodes.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/004_proxy_nodes.sql index f6b6cfefd..cd104a852 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/004_proxy_nodes.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/004_proxy_nodes.sql @@ -3,6 +3,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( `id` VARCHAR(64) NOT NULL, + `tunnel_generation` VARCHAR(64) NOT NULL, `name` VARCHAR(255) NOT NULL, `ip` VARCHAR(512) NOT NULL, `port` INT NOT NULL, @@ -17,7 +18,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( `is_manual` TINYINT(1) NOT NULL DEFAULT 0, `proxy_url` VARCHAR(500), `proxy_username` VARCHAR(255), - `proxy_password` VARCHAR(500), + `proxy_password` TEXT, `created_at` BIGINT NOT NULL, `updated_at` BIGINT NOT NULL, `remote_config` JSON, @@ -31,7 +32,8 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( `dns_failures` BIGINT NOT NULL DEFAULT 0, `stream_errors` BIGINT NOT NULL DEFAULT 0, `proxy_metadata` JSON, - PRIMARY KEY (`id`) + PRIMARY KEY (`id`), + UNIQUE KEY uq_proxy_node_ip_port (`ip`, `port`) ); CREATE TABLE IF NOT EXISTS proxy_node_events ( diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/005_wallet_billing.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/005_wallet_billing.sql index 1111eca99..6b46a309c 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/005_wallet_billing.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/005_wallet_billing.sql @@ -87,7 +87,7 @@ CREATE TABLE IF NOT EXISTS payment_orders ( `product_snapshot` JSON, `fulfillment_status` VARCHAR(64) NOT NULL DEFAULT 'pending', `fulfillment_error` LONGTEXT, - `gateway_order_id` VARCHAR(128), + `gateway_order_id` VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin, `gateway_response` JSON, `status` VARCHAR(64) NOT NULL DEFAULT 'pending', `created_at` BIGINT NOT NULL, @@ -100,6 +100,7 @@ CREATE TABLE IF NOT EXISTS payment_orders ( KEY idx_payment_orders_user_created (`user_id`, `created_at`), KEY idx_payment_orders_status (`status`), KEY idx_payment_orders_gateway_order_id (`gateway_order_id`), + UNIQUE KEY uq_payment_orders_payment_method_gateway_order_id (`payment_method`, `gateway_order_id`), KEY idx_payment_orders_kind_status (`order_kind`, `status`), KEY idx_payment_orders_product (`product_id`) ); @@ -130,8 +131,6 @@ CREATE TABLE IF NOT EXISTS user_referrals ( KEY idx_user_referrals_inviter (`inviter_user_id`, `created_at`), KEY idx_user_referrals_created (`created_at`), KEY idx_user_referrals_invite_code (`invite_code_snapshot`), - CONSTRAINT user_referrals_inviter_user_id_fkey FOREIGN KEY (`inviter_user_id`) REFERENCES users (`id`) ON DELETE CASCADE, - CONSTRAINT user_referrals_invitee_user_id_fkey FOREIGN KEY (`invitee_user_id`) REFERENCES users (`id`) ON DELETE CASCADE, CONSTRAINT user_referrals_first_paid_order_fkey FOREIGN KEY (`first_paid_order_id`) REFERENCES payment_orders (`id`) ON DELETE SET NULL ); @@ -161,8 +160,6 @@ CREATE TABLE IF NOT EXISTS referral_rewards ( KEY idx_referral_rewards_created (`created_at`), KEY idx_referral_rewards_source_order (`source_order_id`), CONSTRAINT referral_rewards_referral_id_fkey FOREIGN KEY (`referral_id`) REFERENCES user_referrals (`id`) ON DELETE CASCADE, - CONSTRAINT referral_rewards_inviter_user_id_fkey FOREIGN KEY (`inviter_user_id`) REFERENCES users (`id`) ON DELETE CASCADE, - CONSTRAINT referral_rewards_invitee_user_id_fkey FOREIGN KEY (`invitee_user_id`) REFERENCES users (`id`) ON DELETE CASCADE, CONSTRAINT referral_rewards_source_order_fkey FOREIGN KEY (`source_order_id`) REFERENCES payment_orders (`id`) ON DELETE SET NULL ); diff --git a/crates/aether-data/runtime/schema/generated/mysql/baseline/006_usage.sql b/crates/aether-data/runtime/schema/generated/mysql/baseline/006_usage.sql index 1a2b1dbb6..15f272948 100644 --- a/crates/aether-data/runtime/schema/generated/mysql/baseline/006_usage.sql +++ b/crates/aether-data/runtime/schema/generated/mysql/baseline/006_usage.sql @@ -173,6 +173,7 @@ CREATE TABLE IF NOT EXISTS usage_counter_deltas ( `request_id` VARCHAR(128) NOT NULL, `kind` VARCHAR(64) NOT NULL, `target_id` TEXT NOT NULL, + `target_tunnel_generation` VARCHAR(64), `request_count_delta` BIGINT NOT NULL DEFAULT 0, `total_requests_delta` BIGINT NOT NULL DEFAULT 0, `success_count_delta` BIGINT NOT NULL DEFAULT 0, @@ -243,3 +244,39 @@ CREATE TABLE IF NOT EXISTS usage_settlement_snapshots ( KEY ix_usage_settlement_snapshots_pricing_source (`billing_pricing_source`) ); +CREATE TABLE IF NOT EXISTS usage_cost_reservations ( + `request_id` VARCHAR(128) NOT NULL, + `subject_id` VARCHAR(128) NOT NULL, + `reservation_token` VARCHAR(128) NOT NULL, + `admitted_at` BIGINT NOT NULL, + `reserved_cost_units` BIGINT NOT NULL, + `actual_cost_units` BIGINT, + `state` VARCHAR(20) NOT NULL, + `reservation_expires_at` BIGINT NOT NULL, + `retain_until` BIGINT NOT NULL, + `finalized_at` BIGINT, + `created_at` BIGINT NOT NULL, + `updated_at` BIGINT NOT NULL, + PRIMARY KEY (`reservation_token`), + KEY usage_cost_reservations_request_id_idx (`request_id`), + KEY usage_cost_reservations_subject_admitted_at_idx (`subject_id`, `admitted_at`), + KEY usage_cost_reservations_reservation_expires_at_idx (`reservation_expires_at`), + KEY usage_cost_reservations_retain_until_token_idx (`retain_until`, `reservation_token`), + CONSTRAINT usage_cost_reservations_subject_id_fkey FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE +); + +CREATE TABLE IF NOT EXISTS usage_request_admissions ( + `request_id` VARCHAR(128) NOT NULL, + `subject_id` VARCHAR(128) NOT NULL, + `event_token` VARCHAR(128) NOT NULL, + `admitted_at` BIGINT NOT NULL, + `retain_until` BIGINT NOT NULL, + `state` VARCHAR(20) NOT NULL, + `released_at` BIGINT, + `created_at` BIGINT NOT NULL, + PRIMARY KEY (`event_token`), + KEY usage_request_admissions_subject_admitted_at_idx (`subject_id`, `admitted_at`), + KEY usage_request_admissions_retain_until_token_idx (`retain_until`, `event_token`), + CONSTRAINT usage_request_admissions_subject_id_fkey FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE +); + diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/001_identity.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/001_identity.sql index d036d8d38..0f6b785c2 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/001_identity.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/001_identity.sql @@ -12,6 +12,7 @@ CREATE TABLE IF NOT EXISTS public.users ( email_verified boolean DEFAULT false NOT NULL, is_active boolean DEFAULT true NOT NULL, is_deleted boolean DEFAULT false NOT NULL, + security_version bigint DEFAULT 0 NOT NULL, allowed_models jsonb, allowed_models_mode character varying(32) DEFAULT 'unrestricted' NOT NULL, allowed_providers jsonb, @@ -202,6 +203,7 @@ CREATE INDEX IF NOT EXISTS user_preferences_user_id_idx ON public.user_preferenc CREATE TABLE IF NOT EXISTS public.user_sessions ( id character varying(64) NOT NULL, user_id character varying(64) NOT NULL, + security_version bigint DEFAULT 0 NOT NULL, client_device_id character varying(128) NOT NULL, device_label character varying(120), device_type character varying(20) DEFAULT 'unknown' NOT NULL, diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/002_provider_catalog.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/002_provider_catalog.sql index 079698927..47008357d 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/002_provider_catalog.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/002_provider_catalog.sql @@ -404,6 +404,7 @@ CREATE TABLE IF NOT EXISTS public.routing_groups ( description text, enabled boolean DEFAULT true NOT NULL, is_system_default boolean DEFAULT false NOT NULL, + sort_order bigint DEFAULT 0 NOT NULL, config_json jsonb NOT NULL, version bigint DEFAULT 1 NOT NULL, created_at bigint NOT NULL, @@ -414,6 +415,7 @@ CREATE TABLE IF NOT EXISTS public.routing_groups ( ALTER TABLE ONLY public.routing_groups ADD CONSTRAINT routing_groups_pkey PRIMARY KEY (id); ALTER TABLE ONLY public.routing_groups ADD CONSTRAINT routing_groups_name_key UNIQUE (name); CREATE INDEX IF NOT EXISTS routing_groups_system_default_idx ON public.routing_groups USING btree (is_system_default, enabled); +CREATE INDEX IF NOT EXISTS routing_groups_enabled_sort_idx ON public.routing_groups USING btree (enabled, sort_order, name, id); CREATE TABLE IF NOT EXISTS public.routing_group_bindings ( id character varying(64) NOT NULL, diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/003_auth_config.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/003_auth_config.sql index 766771ef2..b7c0cf3bb 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/003_auth_config.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/003_auth_config.sql @@ -48,6 +48,7 @@ ALTER TABLE ONLY public.oauth_providers ADD CONSTRAINT oauth_providers_pkey PRIM CREATE TABLE IF NOT EXISTS public.ldap_configs ( id bigserial NOT NULL, + singleton_key integer DEFAULT 1 NOT NULL, server_url character varying(255) NOT NULL, bind_dn text NOT NULL, bind_password_encrypted text, @@ -65,6 +66,7 @@ CREATE TABLE IF NOT EXISTS public.ldap_configs ( ); ALTER TABLE ONLY public.ldap_configs ADD CONSTRAINT ldap_configs_pkey PRIMARY KEY (id); +ALTER TABLE ONLY public.ldap_configs ADD CONSTRAINT ldap_configs_singleton_key_key UNIQUE (singleton_key); CREATE TABLE IF NOT EXISTS public.user_oauth_links ( id character varying(64) NOT NULL, diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/004_proxy_nodes.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/004_proxy_nodes.sql index d6c0e418b..4c4dd6aa4 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/004_proxy_nodes.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/004_proxy_nodes.sql @@ -3,6 +3,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( id character varying(64) NOT NULL, + tunnel_generation character varying(64) NOT NULL, name character varying(255) NOT NULL, ip character varying(512) NOT NULL, port integer NOT NULL, @@ -17,7 +18,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( is_manual boolean DEFAULT false NOT NULL, proxy_url character varying(500), proxy_username character varying(255), - proxy_password character varying(500), + proxy_password text, created_at bigint NOT NULL, updated_at bigint NOT NULL, remote_config jsonb, @@ -34,6 +35,7 @@ CREATE TABLE IF NOT EXISTS public.proxy_nodes ( ); ALTER TABLE ONLY public.proxy_nodes ADD CONSTRAINT proxy_nodes_pkey PRIMARY KEY (id); +ALTER TABLE ONLY public.proxy_nodes ADD CONSTRAINT uq_proxy_node_ip_port UNIQUE (ip, port); CREATE TABLE IF NOT EXISTS public.proxy_node_events ( id bigserial NOT NULL, 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 657288fe7..da83615e5 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 @@ -105,6 +105,7 @@ CREATE INDEX IF NOT EXISTS idx_payment_orders_wallet_created ON public.payment_o CREATE INDEX IF NOT EXISTS idx_payment_orders_user_created ON public.payment_orders USING btree (user_id, created_at); CREATE INDEX IF NOT EXISTS idx_payment_orders_status ON public.payment_orders USING btree (status); CREATE INDEX IF NOT EXISTS idx_payment_orders_gateway_order_id ON public.payment_orders USING btree (gateway_order_id); +CREATE UNIQUE INDEX IF NOT EXISTS uq_payment_orders_payment_method_gateway_order_id ON public.payment_orders USING btree (payment_method, gateway_order_id); CREATE INDEX IF NOT EXISTS idx_payment_orders_kind_status ON public.payment_orders USING btree (order_kind, status); CREATE INDEX IF NOT EXISTS idx_payment_orders_product ON public.payment_orders USING btree (product_id); @@ -137,8 +138,6 @@ ALTER TABLE ONLY public.user_referrals ADD CONSTRAINT user_referrals_invitee_use CREATE INDEX IF NOT EXISTS idx_user_referrals_inviter ON public.user_referrals USING btree (inviter_user_id, created_at); CREATE INDEX IF NOT EXISTS idx_user_referrals_created ON public.user_referrals USING btree (created_at); CREATE INDEX IF NOT EXISTS idx_user_referrals_invite_code ON public.user_referrals USING btree (invite_code_snapshot); -ALTER TABLE ONLY public.user_referrals ADD CONSTRAINT user_referrals_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -ALTER TABLE ONLY public.user_referrals ADD CONSTRAINT user_referrals_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES public.users(id) ON DELETE CASCADE; ALTER TABLE ONLY public.user_referrals ADD CONSTRAINT user_referrals_first_paid_order_fkey FOREIGN KEY (first_paid_order_id) REFERENCES public.payment_orders(id) ON DELETE SET NULL; CREATE TABLE IF NOT EXISTS public.referral_rewards ( @@ -169,8 +168,6 @@ CREATE INDEX IF NOT EXISTS idx_referral_rewards_inviter_created ON public.referr CREATE INDEX IF NOT EXISTS idx_referral_rewards_created ON public.referral_rewards USING btree (created_at); CREATE INDEX IF NOT EXISTS idx_referral_rewards_source_order ON public.referral_rewards USING btree (source_order_id); ALTER TABLE ONLY public.referral_rewards ADD CONSTRAINT referral_rewards_referral_id_fkey FOREIGN KEY (referral_id) REFERENCES public.user_referrals(id) ON DELETE CASCADE; -ALTER TABLE ONLY public.referral_rewards ADD CONSTRAINT referral_rewards_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES public.users(id) ON DELETE CASCADE; -ALTER TABLE ONLY public.referral_rewards ADD CONSTRAINT referral_rewards_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES public.users(id) ON DELETE CASCADE; ALTER TABLE ONLY public.referral_rewards ADD CONSTRAINT referral_rewards_source_order_fkey FOREIGN KEY (source_order_id) REFERENCES public.payment_orders(id) ON DELETE SET NULL; CREATE TABLE IF NOT EXISTS public.payment_gateway_configs ( diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/006_usage.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/006_usage.sql index 3b3dbf2ac..173d7b7de 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/006_usage.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/006_usage.sql @@ -177,6 +177,7 @@ CREATE TABLE IF NOT EXISTS public.usage_counter_deltas ( request_id character varying(128) NOT NULL, kind character varying(64) NOT NULL, target_id text NOT NULL, + target_tunnel_generation character varying(64), request_count_delta bigint DEFAULT 0 NOT NULL, total_requests_delta bigint DEFAULT 0 NOT NULL, success_count_delta bigint DEFAULT 0 NOT NULL, @@ -249,3 +250,41 @@ CREATE INDEX IF NOT EXISTS usage_settlement_snapshots_wallet_id_idx ON public.us CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_schema_version ON public.usage_settlement_snapshots USING btree (settlement_snapshot_schema_version); CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_pricing_source ON public.usage_settlement_snapshots USING btree (billing_pricing_source); +CREATE TABLE IF NOT EXISTS public.usage_cost_reservations ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + reservation_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + reserved_cost_units bigint NOT NULL, + actual_cost_units bigint, + state character varying(20) NOT NULL, + reservation_expires_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + finalized_at timestamp with time zone, + created_at timestamp with time zone NOT NULL, + updated_at timestamp with time zone NOT NULL +); + +ALTER TABLE ONLY public.usage_cost_reservations ADD CONSTRAINT usage_cost_reservations_pkey PRIMARY KEY (reservation_token); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_request_id_idx ON public.usage_cost_reservations USING btree (request_id); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_subject_admitted_at_idx ON public.usage_cost_reservations USING btree (subject_id, admitted_at); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_reservation_expires_at_idx ON public.usage_cost_reservations USING btree (reservation_expires_at); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_retain_until_token_idx ON public.usage_cost_reservations USING btree (retain_until, reservation_token); +ALTER TABLE ONLY public.usage_cost_reservations ADD CONSTRAINT usage_cost_reservations_subject_id_fkey FOREIGN KEY (subject_id) REFERENCES public.users(id) ON DELETE CASCADE; + +CREATE TABLE IF NOT EXISTS public.usage_request_admissions ( + request_id character varying(128) NOT NULL, + subject_id character varying(128) NOT NULL, + event_token character varying(128) NOT NULL, + admitted_at timestamp with time zone NOT NULL, + retain_until timestamp with time zone NOT NULL, + state character varying(20) NOT NULL, + released_at timestamp with time zone, + created_at timestamp with time zone NOT NULL +); + +ALTER TABLE ONLY public.usage_request_admissions ADD CONSTRAINT usage_request_admissions_pkey PRIMARY KEY (event_token); +CREATE INDEX IF NOT EXISTS usage_request_admissions_subject_admitted_at_idx ON public.usage_request_admissions USING btree (subject_id, admitted_at); +CREATE INDEX IF NOT EXISTS usage_request_admissions_retain_until_token_idx ON public.usage_request_admissions USING btree (retain_until, event_token); +ALTER TABLE ONLY public.usage_request_admissions ADD CONSTRAINT usage_request_admissions_subject_id_fkey FOREIGN KEY (subject_id) REFERENCES public.users(id) ON DELETE CASCADE; + diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/001_identity.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/001_identity.sql index 560f5c8fe..8a9eb62e6 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/001_identity.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/001_identity.sql @@ -12,6 +12,7 @@ CREATE TABLE IF NOT EXISTS users ( email_verified INTEGER NOT NULL DEFAULT 0, is_active INTEGER NOT NULL DEFAULT 1, is_deleted INTEGER NOT NULL DEFAULT 0, + security_version INTEGER NOT NULL DEFAULT 0, allowed_models TEXT, allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted', allowed_providers TEXT, @@ -185,6 +186,7 @@ CREATE INDEX IF NOT EXISTS user_preferences_user_id_idx ON user_preferences (use CREATE TABLE IF NOT EXISTS user_sessions ( id TEXT PRIMARY KEY NOT NULL, user_id TEXT NOT NULL, + security_version INTEGER NOT NULL DEFAULT 0, client_device_id TEXT NOT NULL, device_label TEXT, device_type TEXT NOT NULL DEFAULT 'unknown', diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/002_provider_catalog.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/002_provider_catalog.sql index feaac5141..a09bebaf1 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/002_provider_catalog.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/002_provider_catalog.sql @@ -378,6 +378,7 @@ CREATE TABLE IF NOT EXISTS routing_groups ( description TEXT, enabled INTEGER NOT NULL DEFAULT 1, is_system_default INTEGER NOT NULL DEFAULT 0, + sort_order INTEGER NOT NULL DEFAULT 0, config_json TEXT NOT NULL, version INTEGER NOT NULL DEFAULT 1, created_at INTEGER NOT NULL, @@ -386,6 +387,7 @@ CREATE TABLE IF NOT EXISTS routing_groups ( UNIQUE (name) ); CREATE INDEX IF NOT EXISTS routing_groups_system_default_idx ON routing_groups (is_system_default, enabled); +CREATE INDEX IF NOT EXISTS routing_groups_enabled_sort_idx ON routing_groups (enabled, sort_order, name, id); CREATE TABLE IF NOT EXISTS routing_group_bindings ( id TEXT PRIMARY KEY NOT NULL, diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/003_auth_config.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/003_auth_config.sql index 9f2e2c7c5..7a0f3218b 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/003_auth_config.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/003_auth_config.sql @@ -42,6 +42,7 @@ CREATE TABLE IF NOT EXISTS oauth_providers ( CREATE TABLE IF NOT EXISTS ldap_configs ( id INTEGER PRIMARY KEY AUTOINCREMENT, + singleton_key INTEGER NOT NULL DEFAULT 1, server_url TEXT NOT NULL, bind_dn TEXT NOT NULL, bind_password_encrypted TEXT, @@ -55,7 +56,8 @@ CREATE TABLE IF NOT EXISTS ldap_configs ( use_starttls INTEGER NOT NULL DEFAULT 0, connect_timeout INTEGER NOT NULL DEFAULT 10, created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL + updated_at INTEGER NOT NULL, + UNIQUE (singleton_key) ); CREATE TABLE IF NOT EXISTS user_oauth_links ( diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/004_proxy_nodes.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/004_proxy_nodes.sql index 76777cd46..2e7294ba9 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/004_proxy_nodes.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/004_proxy_nodes.sql @@ -3,6 +3,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( id TEXT PRIMARY KEY NOT NULL, + tunnel_generation TEXT NOT NULL, name TEXT NOT NULL, ip TEXT NOT NULL, port INTEGER NOT NULL, @@ -30,7 +31,8 @@ CREATE TABLE IF NOT EXISTS proxy_nodes ( failed_requests INTEGER NOT NULL DEFAULT 0, dns_failures INTEGER NOT NULL DEFAULT 0, stream_errors INTEGER NOT NULL DEFAULT 0, - proxy_metadata TEXT + proxy_metadata TEXT, + UNIQUE (ip, port) ); CREATE TABLE IF NOT EXISTS proxy_node_events ( diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/005_wallet_billing.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/005_wallet_billing.sql index 3a0194c12..c1f09090b 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/005_wallet_billing.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/005_wallet_billing.sql @@ -97,6 +97,7 @@ CREATE INDEX IF NOT EXISTS idx_payment_orders_wallet_created ON payment_orders ( CREATE INDEX IF NOT EXISTS idx_payment_orders_user_created ON payment_orders (user_id, created_at); CREATE INDEX IF NOT EXISTS idx_payment_orders_status ON payment_orders (status); CREATE INDEX IF NOT EXISTS idx_payment_orders_gateway_order_id ON payment_orders (gateway_order_id); +CREATE UNIQUE INDEX IF NOT EXISTS uq_payment_orders_payment_method_gateway_order_id ON payment_orders (payment_method, gateway_order_id); CREATE INDEX IF NOT EXISTS idx_payment_orders_kind_status ON payment_orders (order_kind, status); CREATE INDEX IF NOT EXISTS idx_payment_orders_product ON payment_orders (product_id); @@ -121,8 +122,6 @@ CREATE TABLE IF NOT EXISTS user_referrals ( created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL, UNIQUE (invitee_user_id), - CONSTRAINT user_referrals_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES users (id) ON DELETE CASCADE, - CONSTRAINT user_referrals_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES users (id) ON DELETE CASCADE, CONSTRAINT user_referrals_first_paid_order_fkey FOREIGN KEY (first_paid_order_id) REFERENCES payment_orders (id) ON DELETE SET NULL ); CREATE INDEX IF NOT EXISTS idx_user_referrals_inviter ON user_referrals (inviter_user_id, created_at); @@ -150,8 +149,6 @@ CREATE TABLE IF NOT EXISTS referral_rewards ( updated_at INTEGER NOT NULL, UNIQUE (idempotency_key), CONSTRAINT referral_rewards_referral_id_fkey FOREIGN KEY (referral_id) REFERENCES user_referrals (id) ON DELETE CASCADE, - CONSTRAINT referral_rewards_inviter_user_id_fkey FOREIGN KEY (inviter_user_id) REFERENCES users (id) ON DELETE CASCADE, - CONSTRAINT referral_rewards_invitee_user_id_fkey FOREIGN KEY (invitee_user_id) REFERENCES users (id) ON DELETE CASCADE, CONSTRAINT referral_rewards_source_order_fkey FOREIGN KEY (source_order_id) REFERENCES payment_orders (id) ON DELETE SET NULL ); CREATE INDEX IF NOT EXISTS idx_referral_rewards_inviter_status ON referral_rewards (inviter_user_id, status, created_at); diff --git a/crates/aether-data/runtime/schema/generated/sqlite/baseline/006_usage.sql b/crates/aether-data/runtime/schema/generated/sqlite/baseline/006_usage.sql index f4495c760..890d6915b 100644 --- a/crates/aether-data/runtime/schema/generated/sqlite/baseline/006_usage.sql +++ b/crates/aether-data/runtime/schema/generated/sqlite/baseline/006_usage.sql @@ -169,6 +169,7 @@ CREATE TABLE IF NOT EXISTS usage_counter_deltas ( request_id TEXT NOT NULL, kind TEXT NOT NULL, target_id TEXT NOT NULL, + target_tunnel_generation TEXT, request_count_delta INTEGER NOT NULL DEFAULT 0, total_requests_delta INTEGER NOT NULL DEFAULT 0, success_count_delta INTEGER NOT NULL DEFAULT 0, @@ -237,3 +238,37 @@ CREATE INDEX IF NOT EXISTS usage_settlement_snapshots_wallet_id_idx ON usage_set CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_schema_version ON usage_settlement_snapshots (settlement_snapshot_schema_version); CREATE INDEX IF NOT EXISTS ix_usage_settlement_snapshots_pricing_source ON usage_settlement_snapshots (billing_pricing_source); +CREATE TABLE IF NOT EXISTS usage_cost_reservations ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + reservation_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + reserved_cost_units INTEGER NOT NULL, + actual_cost_units INTEGER, + state TEXT NOT NULL, + reservation_expires_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + finalized_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + CONSTRAINT usage_cost_reservations_subject_id_fkey FOREIGN KEY (subject_id) REFERENCES users (id) ON DELETE CASCADE +); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_request_id_idx ON usage_cost_reservations (request_id); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_subject_admitted_at_idx ON usage_cost_reservations (subject_id, admitted_at); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_reservation_expires_at_idx ON usage_cost_reservations (reservation_expires_at); +CREATE INDEX IF NOT EXISTS usage_cost_reservations_retain_until_token_idx ON usage_cost_reservations (retain_until, reservation_token); + +CREATE TABLE IF NOT EXISTS usage_request_admissions ( + request_id TEXT NOT NULL, + subject_id TEXT NOT NULL, + event_token TEXT PRIMARY KEY NOT NULL, + admitted_at INTEGER NOT NULL, + retain_until INTEGER NOT NULL, + state TEXT NOT NULL, + released_at INTEGER, + created_at INTEGER NOT NULL, + CONSTRAINT usage_request_admissions_subject_id_fkey FOREIGN KEY (subject_id) REFERENCES users (id) ON DELETE CASCADE +); +CREATE INDEX IF NOT EXISTS usage_request_admissions_subject_admitted_at_idx ON usage_request_admissions (subject_id, admitted_at); +CREATE INDEX IF NOT EXISTS usage_request_admissions_retain_until_token_idx ON usage_request_admissions (retain_until, event_token); + diff --git a/crates/aether-data/runtime/schema/logical/001_identity.toml b/crates/aether-data/runtime/schema/logical/001_identity.toml index e99fb74db..15ce325e6 100644 --- a/crates/aether-data/runtime/schema/logical/001_identity.toml +++ b/crates/aether-data/runtime/schema/logical/001_identity.toml @@ -59,6 +59,11 @@ name = "is_deleted" type = "bool" default = false +[[table.users.columns]] +name = "security_version" +type = "int64" +default = 0 + [[table.users.columns]] name = "allowed_models" type = "json" @@ -818,6 +823,11 @@ name = "user_id" type = "text_id" length = 64 +[[table.user_sessions.columns]] +name = "security_version" +type = "int64" +default = 0 + [[table.user_sessions.columns]] name = "client_device_id" type = "text" diff --git a/crates/aether-data/runtime/schema/logical/002_provider_catalog.toml b/crates/aether-data/runtime/schema/logical/002_provider_catalog.toml index e838f57ff..db035eff4 100644 --- a/crates/aether-data/runtime/schema/logical/002_provider_catalog.toml +++ b/crates/aether-data/runtime/schema/logical/002_provider_catalog.toml @@ -1733,6 +1733,11 @@ name = "is_system_default" type = "bool" default = false +[[table.routing_groups.columns]] +name = "sort_order" +type = "int64" +default = 0 + [[table.routing_groups.columns]] name = "config_json" type = "json" @@ -1763,6 +1768,10 @@ columns = ["name"] name = "routing_groups_system_default_idx" columns = ["is_system_default", "enabled"] +[[table.routing_groups.indexes]] +name = "routing_groups_enabled_sort_idx" +columns = ["enabled", "sort_order", "name", "id"] + [table.routing_group_bindings] domain = "provider_catalog" order = 111 diff --git a/crates/aether-data/runtime/schema/logical/003_auth_config.toml b/crates/aether-data/runtime/schema/logical/003_auth_config.toml index 2e441bf8e..731ff316a 100644 --- a/crates/aether-data/runtime/schema/logical/003_auth_config.toml +++ b/crates/aether-data/runtime/schema/logical/003_auth_config.toml @@ -166,6 +166,11 @@ name = "id" type = "int64" auto_increment = true +[[table.ldap_configs.columns]] +name = "singleton_key" +type = "int32" +default = 1 + [[table.ldap_configs.columns]] name = "server_url" type = "text" @@ -236,6 +241,10 @@ type = "unix_seconds" name = "updated_at" type = "unix_seconds" +[[table.ldap_configs.uniques]] +name = "ldap_configs_singleton_key_key" +columns = ["singleton_key"] + [table.user_oauth_links] domain = "auth_config" order = 50 diff --git a/crates/aether-data/runtime/schema/logical/004_proxy_nodes.toml b/crates/aether-data/runtime/schema/logical/004_proxy_nodes.toml index f1b9090c6..25de0cdc6 100644 --- a/crates/aether-data/runtime/schema/logical/004_proxy_nodes.toml +++ b/crates/aether-data/runtime/schema/logical/004_proxy_nodes.toml @@ -8,6 +8,11 @@ name = "id" type = "text_id" length = 64 +[[table.proxy_nodes.columns]] +name = "tunnel_generation" +type = "text_id" +length = 64 + [[table.proxy_nodes.columns]] name = "name" type = "text" @@ -85,7 +90,6 @@ nullable = true [[table.proxy_nodes.columns]] name = "proxy_password" type = "text" -length = 500 nullable = true [[table.proxy_nodes.columns]] @@ -151,6 +155,10 @@ name = "proxy_metadata" type = "json" nullable = true +[[table.proxy_nodes.uniques]] +name = "uq_proxy_node_ip_port" +columns = ["ip", "port"] + [table.proxy_node_events] domain = "proxy_nodes" order = 20 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 40f3289ae..d2f6be018 100644 --- a/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml +++ b/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml @@ -380,6 +380,9 @@ type = "text" length = 128 nullable = true + [table.payment_orders.columns.driver.mysql] +type = "VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin" + [[table.payment_orders.columns]] name = "gateway_response" type = "json" @@ -430,6 +433,11 @@ columns = ["status"] name = "idx_payment_orders_gateway_order_id" columns = ["gateway_order_id"] +[[table.payment_orders.indexes]] +name = "uq_payment_orders_payment_method_gateway_order_id" +columns = ["payment_method", "gateway_order_id"] +unique = true + [[table.payment_orders.indexes]] name = "idx_payment_orders_kind_status" columns = ["order_kind", "status"] @@ -542,20 +550,6 @@ columns = ["created_at"] name = "idx_user_referrals_invite_code" columns = ["invite_code_snapshot"] -[[table.user_referrals.foreign_keys]] -name = "user_referrals_inviter_user_id_fkey" -columns = ["inviter_user_id"] -references_table = "users" -references_columns = ["id"] -on_delete = "cascade" - -[[table.user_referrals.foreign_keys]] -name = "user_referrals_invitee_user_id_fkey" -columns = ["invitee_user_id"] -references_table = "users" -references_columns = ["id"] -on_delete = "cascade" - [[table.user_referrals.foreign_keys]] name = "user_referrals_first_paid_order_fkey" columns = ["first_paid_order_id"] @@ -686,20 +680,6 @@ references_table = "user_referrals" references_columns = ["id"] on_delete = "cascade" -[[table.referral_rewards.foreign_keys]] -name = "referral_rewards_inviter_user_id_fkey" -columns = ["inviter_user_id"] -references_table = "users" -references_columns = ["id"] -on_delete = "cascade" - -[[table.referral_rewards.foreign_keys]] -name = "referral_rewards_invitee_user_id_fkey" -columns = ["invitee_user_id"] -references_table = "users" -references_columns = ["id"] -on_delete = "cascade" - [[table.referral_rewards.foreign_keys]] name = "referral_rewards_source_order_fkey" columns = ["source_order_id"] diff --git a/crates/aether-data/runtime/schema/logical/006_usage.toml b/crates/aether-data/runtime/schema/logical/006_usage.toml index f137bae1d..d27fed4b9 100644 --- a/crates/aether-data/runtime/schema/logical/006_usage.toml +++ b/crates/aether-data/runtime/schema/logical/006_usage.toml @@ -831,6 +831,12 @@ length = 64 name = "target_id" type = "text" +[[table.usage_counter_deltas.columns]] +name = "target_tunnel_generation" +type = "text" +length = 64 +nullable = true + [[table.usage_counter_deltas.columns]] name = "request_count_delta" type = "int64" @@ -1147,3 +1153,142 @@ columns = ["settlement_snapshot_schema_version"] [[table.usage_settlement_snapshots.indexes]] name = "ix_usage_settlement_snapshots_pricing_source" columns = ["billing_pricing_source"] + +[table.usage_cost_reservations] +domain = "usage" +order = 21 +primary_key = ["reservation_token"] + +[[table.usage_cost_reservations.columns]] +name = "request_id" +type = "text" +length = 128 + +[[table.usage_cost_reservations.columns]] +name = "subject_id" +type = "text" +length = 128 + +[[table.usage_cost_reservations.columns]] +name = "reservation_token" +type = "text" +length = 128 + +[[table.usage_cost_reservations.columns]] +name = "admitted_at" +type = "timestamp" + +[[table.usage_cost_reservations.columns]] +name = "reserved_cost_units" +type = "int64" + +[[table.usage_cost_reservations.columns]] +name = "actual_cost_units" +type = "int64" +nullable = true + +[[table.usage_cost_reservations.columns]] +name = "state" +type = "text" +length = 20 + +[[table.usage_cost_reservations.columns]] +name = "reservation_expires_at" +type = "timestamp" + +[[table.usage_cost_reservations.columns]] +name = "retain_until" +type = "timestamp" + +[[table.usage_cost_reservations.columns]] +name = "finalized_at" +type = "timestamp" +nullable = true + +[[table.usage_cost_reservations.columns]] +name = "created_at" +type = "timestamp" + +[[table.usage_cost_reservations.columns]] +name = "updated_at" +type = "timestamp" + +[[table.usage_cost_reservations.indexes]] +name = "usage_cost_reservations_request_id_idx" +columns = ["request_id"] + +[[table.usage_cost_reservations.indexes]] +name = "usage_cost_reservations_subject_admitted_at_idx" +columns = ["subject_id", "admitted_at"] + +[[table.usage_cost_reservations.indexes]] +name = "usage_cost_reservations_reservation_expires_at_idx" +columns = ["reservation_expires_at"] + +[[table.usage_cost_reservations.indexes]] +name = "usage_cost_reservations_retain_until_token_idx" +columns = ["retain_until", "reservation_token"] + +[[table.usage_cost_reservations.foreign_keys]] +name = "usage_cost_reservations_subject_id_fkey" +columns = ["subject_id"] +references_table = "users" +references_columns = ["id"] +on_delete = "cascade" + +[table.usage_request_admissions] +domain = "usage" +order = 22 +primary_key = ["event_token"] + +[[table.usage_request_admissions.columns]] +name = "request_id" +type = "text" +length = 128 + +[[table.usage_request_admissions.columns]] +name = "subject_id" +type = "text" +length = 128 + +[[table.usage_request_admissions.columns]] +name = "event_token" +type = "text" +length = 128 + +[[table.usage_request_admissions.columns]] +name = "admitted_at" +type = "timestamp" + +[[table.usage_request_admissions.columns]] +name = "retain_until" +type = "timestamp" + +[[table.usage_request_admissions.columns]] +name = "state" +type = "text" +length = 20 + +[[table.usage_request_admissions.columns]] +name = "released_at" +type = "timestamp" +nullable = true + +[[table.usage_request_admissions.columns]] +name = "created_at" +type = "timestamp" + +[[table.usage_request_admissions.indexes]] +name = "usage_request_admissions_subject_admitted_at_idx" +columns = ["subject_id", "admitted_at"] + +[[table.usage_request_admissions.indexes]] +name = "usage_request_admissions_retain_until_token_idx" +columns = ["retain_until", "event_token"] + +[[table.usage_request_admissions.foreign_keys]] +name = "usage_request_admissions_subject_id_fkey" +columns = ["subject_id"] +references_table = "users" +references_columns = ["id"] +on_delete = "cascade" diff --git a/crates/aether-data/runtime/src/backend/maintenance.rs b/crates/aether-data/runtime/src/backend/maintenance.rs index 45a8aed79..fa465c6d0 100644 --- a/crates/aether-data/runtime/src/backend/maintenance.rs +++ b/crates/aether-data/runtime/src/backend/maintenance.rs @@ -202,6 +202,22 @@ impl DataBackends { } } + pub async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + match self.sql_backend() { + Some(backend) => { + backend + .compare_and_set_system_config_string_value(key, expected, replacement) + .await + } + None => Ok(false), + } + } + pub async fn list_system_config_entries( &self, ) -> Result, DataLayerError> { @@ -307,6 +323,8 @@ impl<'a> SqlBackendRef<'a> { Self::Sqlite(sqlite) => { warm_pool(sqlite.pool(), sqlite.config().pool.min_connections).await } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -321,6 +339,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.run_table_maintenance(table_names).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.run_table_maintenance(table_names).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -341,6 +361,8 @@ impl<'a> SqlBackendRef<'a> { crate::lifecycle::migrate::run_sqlite_migrations(sqlite.pool()).await?; Ok(true) } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -361,6 +383,8 @@ impl<'a> SqlBackendRef<'a> { crate::lifecycle::backfill::run_sqlite_backfills(sqlite.pool()).await?; Ok(true) } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -380,6 +404,8 @@ impl<'a> SqlBackendRef<'a> { Self::Sqlite(sqlite) => Ok(Some( crate::lifecycle::migrate::pending_sqlite_migrations(sqlite.pool()).await?, )), + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -400,6 +426,8 @@ impl<'a> SqlBackendRef<'a> { crate::lifecycle::migrate::prepare_sqlite_database_for_startup(sqlite.pool()) .await?, )), + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -419,6 +447,8 @@ impl<'a> SqlBackendRef<'a> { Self::Sqlite(sqlite) => Ok(Some( crate::lifecycle::backfill::pending_sqlite_backfills(sqlite.pool()).await?, )), + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -445,6 +475,8 @@ impl<'a> SqlBackendRef<'a> { sqlite.pool().num_idle(), sqlite.config().pool.max_connections, ), + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -459,6 +491,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.aggregate_wallet_daily_usage(input).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.aggregate_wallet_daily_usage(input).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -473,6 +507,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.aggregate_stats_hourly(input).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.aggregate_stats_hourly(input).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -487,6 +523,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.aggregate_stats_daily(input).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.aggregate_stats_daily(input).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -501,6 +539,38 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.find_system_config_value(key).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.find_system_config_value(key).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), + } + } + + async fn compare_and_set_system_config_string_value( + self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + match self { + #[cfg(feature = "postgres")] + Self::Postgres(postgres) => { + postgres + .compare_and_set_system_config_string_value(key, expected, replacement) + .await + } + #[cfg(feature = "mysql")] + Self::Mysql(mysql) => { + mysql + .compare_and_set_system_config_string_value(key, expected, replacement) + .await + } + #[cfg(feature = "sqlite")] + Self::Sqlite(sqlite) => { + sqlite + .compare_and_set_system_config_string_value(key, expected, replacement) + .await + } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -514,6 +584,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.list_system_config_entries().await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.list_system_config_entries().await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -542,6 +614,8 @@ impl<'a> SqlBackendRef<'a> { .upsert_system_config_entry(key, value, description) .await } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -553,6 +627,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.delete_system_config_value(key).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.delete_system_config_value(key).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -564,6 +640,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.read_admin_system_stats().await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.read_admin_system_stats().await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -578,6 +656,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.purge_admin_system_data(target).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.purge_admin_system_data(target).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -591,6 +671,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.export_admin_system_usage_aggregates().await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.export_admin_system_usage_aggregates().await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -635,6 +717,8 @@ impl<'a> SqlBackendRef<'a> { ) .await } + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } @@ -649,6 +733,8 @@ impl<'a> SqlBackendRef<'a> { Self::Mysql(mysql) => mysql.purge_admin_request_bodies_batch(batch_size).await, #[cfg(feature = "sqlite")] Self::Sqlite(sqlite) => sqlite.purge_admin_request_bodies_batch(batch_size).await, + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"), } } } diff --git a/crates/aether-data/runtime/src/backend/mod.rs b/crates/aether-data/runtime/src/backend/mod.rs index cc76eeb58..af170ed4b 100644 --- a/crates/aether-data/runtime/src/backend/mod.rs +++ b/crates/aether-data/runtime/src/backend/mod.rs @@ -32,9 +32,9 @@ pub use mysql::MysqlBackend; pub use postgres::PostgresBackend; pub use read::DataReadRepositories; pub use referrals::{ - ReferralAdminStats, ReferralDataState, ReferralMutationStatus, ReferralRelationshipListQuery, - ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery, - ReferralRewardRecord, ReferralUserDashboard, + ReferralAdminStats, ReferralDataState, ReferralMutationStatus, ReferralReconciliationSummary, + ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig, + ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard, }; #[cfg(feature = "sqlite")] pub use sqlite::SqliteBackend; @@ -52,6 +52,11 @@ enum SqlBackendRef<'a> { Mysql(&'a MysqlBackend), #[cfg(feature = "sqlite")] Sqlite(&'a SqliteBackend), + // Keep the reference lifetime represented when this crate is built without + // any SQL driver features. The no-driver build still exposes the + // maintenance facade, but has no concrete backend variant to carry `'a`. + #[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))] + Disabled(std::marker::PhantomData<&'a ()>), } #[derive(Debug, Clone, Default)] diff --git a/crates/aether-data/runtime/src/backend/referrals.rs b/crates/aether-data/runtime/src/backend/referrals.rs index 623c605a9..c641e4e12 100644 --- a/crates/aether-data/runtime/src/backend/referrals.rs +++ b/crates/aether-data/runtime/src/backend/referrals.rs @@ -1,7 +1,7 @@ use crate::DataLayerError; +use aether_data_contracts::repository::wallet::payment_order_refund_amounts_are_consistent; use serde::{Deserialize, Serialize}; use sqlx::Row; -use std::collections::HashSet; use super::DataBackends; @@ -16,6 +16,12 @@ impl<'a> ReferralDataState<'a> { } } +const REFERRAL_RECONCILIATION_LIMIT: usize = 200; + +// The list tests intentionally build one page larger than the historical +// in-memory fetch cap. Keep the fixture cap test-only now that production +// queries paginate directly in SQL. +#[cfg(all(test, feature = "sqlite"))] const REFERRAL_FETCH_LIMIT: usize = 5_000; #[derive(Debug, Clone, Serialize)] @@ -72,6 +78,22 @@ pub struct ReferralAdminStats { pub reversed_reward_usd: f64, } +/// Result of one bounded referral reconciliation pass. +/// +/// The pass is intentionally idempotent: rows that cannot be applied (for +/// example, because the inviter wallet is temporarily unavailable) remain in +/// their durable pending/failed state and are picked up by the next pass. +#[derive(Debug, Clone, Copy, Default, Serialize)] +pub struct ReferralReconciliationSummary { + pub order_attempted: u64, + pub order_repaired: u64, + pub reward_attempted: u64, + pub reward_applied: u64, + pub reversal_attempted: u64, + pub reversal_applied: u64, + pub deferred: u64, +} + #[derive(Debug, Clone, Default, Deserialize)] pub struct ReferralRelationshipListQuery { pub inviter: Option, @@ -108,6 +130,13 @@ pub enum ReferralMutationStatus { Unavailable, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ReferralApplyingRecovery { + Applied, + Failed, + Unchanged, +} + #[derive(Debug, Clone)] struct ReferralPaymentOrderContext { id: String, @@ -199,11 +228,6 @@ fn now_unix_secs() -> u64 { chrono::Utc::now().timestamp().max(0) as u64 } -#[cfg(any(feature = "mysql", feature = "sqlite"))] -fn now_unix_ms() -> u64 { - chrono::Utc::now().timestamp_millis().max(0) as u64 -} - fn row_unix_secs(row: &R, column: &str) -> Result where R: Row, @@ -240,47 +264,43 @@ fn generate_invite_code() -> String { ) } -fn referral_text_matches(value: Option<&str>, needle: Option<&str>) -> bool { - let Some(needle) = needle.map(str::trim).filter(|value| !value.is_empty()) else { - return true; +/// Build a case-insensitive SQL `LIKE` pattern while treating user input as a +/// literal substring. `!` is used as the escape character because it is +/// accepted consistently by PostgreSQL, MySQL, and SQLite. +fn referral_like_pattern(value: Option<&str>) -> String { + let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else { + return String::new(); }; - let Some(value) = value else { - return false; - }; - value - .to_ascii_lowercase() - .contains(&needle.to_ascii_lowercase()) + let escaped = value + .replace('!', "!!") + .replace('%', "!%") + .replace('_', "!_") + .to_ascii_lowercase(); + format!("%{escaped}%") } -fn referral_list_window(items: &[T], limit: usize, offset: usize) -> Vec { - let limit = limit.clamp(1, 200); - items.iter().skip(offset).take(limit).cloned().collect() +fn referral_page_bounds(limit: usize, offset: usize) -> (i64, i64) { + let limit = limit.clamp(1, 200) as i64; + let offset = i64::try_from(offset).unwrap_or(i64::MAX); + (limit, offset) } -fn referral_admin_stats( - relationships: &[ReferralRelationshipRecord], - rewards: &[ReferralRewardRecord], -) -> ReferralAdminStats { - ReferralAdminStats { - total_invites: relationships.len() as u64, - effective_invites: relationships - .iter() - .filter(|item| item.first_paid_order_id.is_some()) - .count() as u64, - paid_reward_usd: rewards - .iter() - .filter(|item| item.status == "applied") - .map(|item| item.amount_usd) - .sum(), - pending_reward_usd: rewards - .iter() - .filter(|item| matches!(item.status.as_str(), "pending" | "failed")) - .map(|item| item.amount_usd) - .sum(), - reversed_reward_usd: rewards.iter().map(|item| item.reversed_amount_usd).sum(), +fn referral_stats_amount(value: f64) -> f64 { + if value.is_nan() || value < 0.0 { + 0.0 + } else if value.is_infinite() { + // Database SUM over legacy rows can overflow a binary float. Keep a + // finite, monotonic public value instead of silently reporting zero. + f64::MAX + } else { + value } } +fn referral_stats_count(value: i64) -> u64 { + value.max(0) as u64 +} + fn normalize_referral_code(value: &str) -> Option { let value = value.trim().to_ascii_uppercase(); (!value.is_empty() && value.len() <= 64).then_some(value) @@ -302,6 +322,152 @@ fn referral_void_allowed(status: &str) -> bool { matches!(status, "pending" | "failed") } +fn referral_percent_rate_valid(percent_rate: f64) -> bool { + percent_rate.is_finite() && percent_rate > 0.0 && percent_rate <= 100.0 +} + +fn referral_payment_method_excluded(payment_method: &str) -> bool { + matches!( + payment_method.trim().to_ascii_lowercase().as_str(), + "manual" | "admin_manual" | "redeem_code" | "gift" + ) +} + +fn referral_refund_context_valid(context: &ReferralPaymentOrderRefundContext) -> bool { + context.refunded_amount_usd > 0.0 + && payment_order_refund_amounts_are_consistent( + context.amount_usd, + context.refunded_amount_usd, + (context.amount_usd - context.refunded_amount_usd).max(0.0), + ) +} + +fn referral_wallet_values_valid(balance: f64, gift_balance: f64) -> bool { + // Recharge balances may legitimately be negative when overdraft is + // enabled. Gift balances, however, are never allowed to go below zero. + balance.is_finite() && gift_balance.is_finite() && gift_balance >= 0.0 +} + +fn referral_amounts_match(left: f64, right: f64) -> bool { + if !left.is_finite() || !right.is_finite() { + return false; + } + // PostgreSQL NUMERIC values are decoded through f64 in this runtime. + // Preserve the eight-decimal storage tolerance while allowing a handful + // of ULPs when the running balance is large. + let scale = left.abs().max(right.abs()).max(1.0); + let tolerance = 0.00000001_f64.max(scale * f64::EPSILON * 8.0); + (left - right).abs() <= tolerance +} + +/// Validate the durable wallet snapshot written alongside a referral credit. +/// +/// Matching only `link_id` and `amount` is insufficient: a malformed or +/// manually-inserted transaction could otherwise turn an interrupted +/// `applying` reward into `applied` without ever increasing the inviter's gift +/// balance. The normal credit path writes a complete before/after snapshot, +/// so recovery can require those same invariants before trusting the fact. +// The fact validator compares the complete before/after ledger snapshot. Keep +// each value explicit so a caller cannot accidentally substitute a bucket or +// omit one of the persisted invariants. +#[allow(clippy::too_many_arguments)] +fn referral_credit_transaction_fact_valid( + reward_amount_usd: f64, + amount: f64, + balance_before: f64, + balance_after: f64, + recharge_balance_before: f64, + recharge_balance_after: f64, + gift_balance_before: f64, + gift_balance_after: f64, +) -> bool { + if !reward_amount_usd.is_finite() + || reward_amount_usd <= 0.0 + || !amount.is_finite() + || amount <= 0.0 + || !referral_amounts_match(amount, reward_amount_usd) + || !balance_before.is_finite() + || !balance_after.is_finite() + || !recharge_balance_before.is_finite() + || !recharge_balance_after.is_finite() + || !gift_balance_before.is_finite() + || !gift_balance_after.is_finite() + || gift_balance_before < 0.0 + || gift_balance_after < 0.0 + { + return false; + } + + // Referral credits affect only the gift bucket. The total balance and + // both bucket decompositions must agree with the signed transaction. + referral_amounts_match(recharge_balance_before, recharge_balance_after) + && referral_amounts_match(balance_after, balance_before + amount) + && referral_amounts_match(gift_balance_after, gift_balance_before + amount) + && referral_amounts_match( + balance_before, + recharge_balance_before + gift_balance_before, + ) + && referral_amounts_match(balance_after, recharge_balance_after + gift_balance_after) +} + +fn referral_reversal_state_valid( + reward_amount_usd: f64, + current_reversed_amount_usd: f64, + current_pending_amount_usd: f64, + actual_reverse_amount_usd: f64, + pending_after_usd: f64, +) -> bool { + if !reward_amount_usd.is_finite() + || !current_reversed_amount_usd.is_finite() + || !current_pending_amount_usd.is_finite() + || !actual_reverse_amount_usd.is_finite() + || !pending_after_usd.is_finite() + || reward_amount_usd <= 0.0 + || current_reversed_amount_usd < 0.0 + || current_pending_amount_usd < 0.0 + || actual_reverse_amount_usd < 0.0 + || pending_after_usd < 0.0 + { + return false; + } + let reversed_after_usd = current_reversed_amount_usd + actual_reverse_amount_usd; + let total_reversal_after_usd = reversed_after_usd + pending_after_usd; + reversed_after_usd.is_finite() + && total_reversal_after_usd.is_finite() + && reversed_after_usd <= reward_amount_usd + 0.00000001 + && total_reversal_after_usd <= reward_amount_usd + 0.00000001 +} + +/// Validate the durable reversal counters before calculating or persisting a +/// new debt. In particular, this must run before the wallet lookup: a missing +/// wallet is a normal retry condition, but it must not become a way to carry +/// malformed negative/overflowed counters forward indefinitely. +fn referral_reversal_inputs_valid( + reward_amount_usd: f64, + target_reversal_amount_usd: f64, + current_reversed_amount_usd: f64, + current_pending_amount_usd: f64, +) -> bool { + if !reward_amount_usd.is_finite() + || !target_reversal_amount_usd.is_finite() + || !current_reversed_amount_usd.is_finite() + || !current_pending_amount_usd.is_finite() + || reward_amount_usd <= 0.0 + || target_reversal_amount_usd < 0.0 + || current_reversed_amount_usd < 0.0 + || current_pending_amount_usd < 0.0 + { + return false; + } + + let total_reversal = current_reversed_amount_usd + current_pending_amount_usd; + let tolerance = 0.00000001_f64; + total_reversal.is_finite() + && current_reversed_amount_usd <= reward_amount_usd + tolerance + && total_reversal <= reward_amount_usd + tolerance + && target_reversal_amount_usd <= reward_amount_usd + tolerance +} + fn referral_reversal_delta( reward_amount_usd: f64, order_amount_usd: f64, @@ -311,7 +477,60 @@ fn referral_reversal_delta( ) -> f64 { let target_reversal = referral_reversal_target(reward_amount_usd, order_amount_usd, refunded_amount_usd); - (target_reversal - reversed_amount_usd - pending_reversal_amount_usd).max(0.0) + referral_reversal_due_bounded( + target_reversal, + reward_amount_usd, + reversed_amount_usd, + pending_reversal_amount_usd, + ) +} + +fn referral_reversal_due( + target_reversal_amount_usd: f64, + reversed_amount_usd: f64, + pending_reversal_amount_usd: f64, +) -> f64 { + // A pending amount is a debt, not an amount that has already been + // reversed. Keep it eligible on later passes while also accounting for a + // refund that increased the cumulative target. + (target_reversal_amount_usd - reversed_amount_usd) + .max(0.0) + .max(pending_reversal_amount_usd.max(0.0)) +} + +fn referral_reversal_due_bounded( + target_reversal_amount_usd: f64, + reward_amount_usd: f64, + reversed_amount_usd: f64, + pending_reversal_amount_usd: f64, +) -> f64 { + if !target_reversal_amount_usd.is_finite() + || !reward_amount_usd.is_finite() + || !reversed_amount_usd.is_finite() + || !pending_reversal_amount_usd.is_finite() + { + return 0.0; + } + referral_reversal_due( + target_reversal_amount_usd, + reversed_amount_usd, + pending_reversal_amount_usd, + ) + // Cap the debt at the reward's remaining principal even when a legacy row + // contains an oversized pending value. + .min((reward_amount_usd - reversed_amount_usd.max(0.0)).max(0.0)) +} + +fn referral_pending_reversal_capped( + reward_amount_usd: f64, + reversed_amount_usd: f64, + current_pending_amount_usd: f64, + due_amount_usd: f64, +) -> f64 { + let remaining_principal = (reward_amount_usd - reversed_amount_usd.max(0.0)).max(0.0); + current_pending_amount_usd + .max(due_amount_usd) + .min(remaining_principal) } fn referral_reversal_target( @@ -319,7 +538,13 @@ fn referral_reversal_target( order_amount_usd: f64, refunded_amount_usd: f64, ) -> f64 { - if reward_amount_usd <= 0.0 || order_amount_usd <= 0.0 || refunded_amount_usd <= 0.0 { + if !reward_amount_usd.is_finite() + || !order_amount_usd.is_finite() + || !refunded_amount_usd.is_finite() + || reward_amount_usd <= 0.0 + || order_amount_usd <= 0.0 + || refunded_amount_usd <= 0.0 + { return 0.0; } reward_amount_usd * (refunded_amount_usd / order_amount_usd).clamp(0.0, 1.0) @@ -404,11 +629,9 @@ WHERE id = ? let Some(invite_code) = self.ensure_referral_invite_code(user_id).await? else { return Ok(None); }; - let relationships = self - .list_referral_relationships_raw(Some(user_id), None) - .await?; - let rewards = self.list_referral_rewards_raw(Some(user_id)).await?; - let stats = referral_admin_stats(&relationships, &rewards); + // Dashboard metrics must cover the complete history; do not derive + // them from a bounded list page. + let stats = self.referral_admin_stats_global(Some(user_id)).await?; Ok(Some(ReferralUserDashboard { invite_code, total_invites: stats.total_invites, @@ -427,52 +650,13 @@ WHERE id = ? if self.backends.is_none() { return Ok(None); } - let all_relationships = self.list_referral_relationships_raw(None, None).await?; - let all_rewards = self.list_referral_rewards_raw(None).await?; - let filtered = all_relationships - .into_iter() - .filter(|item| { - referral_text_matches(item.inviter_username.as_deref(), query.inviter.as_deref()) - || referral_text_matches( - Some(item.inviter_user_id.as_str()), - query.inviter.as_deref(), - ) - }) - .filter(|item| { - referral_text_matches(item.invitee_username.as_deref(), query.invitee.as_deref()) - || referral_text_matches( - Some(item.invitee_user_id.as_str()), - query.invitee.as_deref(), - ) - }) - .filter(|item| { - referral_text_matches( - Some(item.invite_code_snapshot.as_str()), - query.invite_code.as_deref(), - ) - }) - .filter(|item| { - query - .first_paid - .map(|expected| item.first_paid_order_id.is_some() == expected) - .unwrap_or(true) - }) - .collect::>(); - let total = filtered.len() as u64; - let filtered_referral_ids = filtered - .iter() - .map(|item| item.id.as_str()) - .collect::>(); - let filtered_rewards = all_rewards - .into_iter() - .filter(|item| filtered_referral_ids.contains(item.referral_id.as_str())) - .collect::>(); - let stats = referral_admin_stats(&filtered, &filtered_rewards); - Ok(Some(( - referral_list_window(&filtered, query.limit, query.offset), - total, - stats, - ))) + // The cards in the admin view are global totals (their labels use + // "total"/"paid" rather than "filtered"). Compute them with an + // aggregate query instead of deriving them from the bounded list + // page, so pagination and filters cannot change the headline stats. + let (items, total) = self.list_referral_relationships_raw(&query).await?; + let stats = self.referral_admin_stats_global(None).await?; + Ok(Some((items, total, stats))) } pub async fn list_admin_referral_rewards( @@ -482,39 +666,193 @@ WHERE id = ? if self.backends.is_none() { return Ok(None); } - let relationships = self.list_referral_relationships_raw(None, None).await?; - let rewards = self.list_referral_rewards_raw(None).await?; - let filtered = rewards - .iter() - .filter(|item| { - referral_text_matches(item.source_order_id.as_deref(), query.order_id.as_deref()) - }) - .filter(|item| { - referral_text_matches( - Some(item.reward_type.as_str()), - query.reward_type.as_deref(), - ) - }) - .filter(|item| { - referral_text_matches(Some(item.status.as_str()), query.status.as_deref()) - }) - .cloned() - .collect::>(); - let total = filtered.len() as u64; - let filtered_referral_ids = filtered - .iter() - .map(|item| item.referral_id.as_str()) - .collect::>(); - let filtered_relationships = relationships - .into_iter() - .filter(|item| filtered_referral_ids.contains(item.id.as_str())) - .collect::>(); - let stats = referral_admin_stats(&filtered_relationships, &filtered); - Ok(Some(( - referral_list_window(&filtered, query.limit, query.offset), - total, - stats, - ))) + let (items, total) = self.list_referral_rewards_raw(&query).await?; + let stats = self.referral_admin_stats_global(None).await?; + Ok(Some((items, total, stats))) + } + + /// Read the headline referral metrics without applying the list window. + /// + /// Admin list endpoints intentionally cap their row payloads, so deriving + /// metrics from those rows would silently under-count once the history is + /// larger than the fetch limit. Keep the aggregate in the data layer and + /// use the native numeric type of each backend before normalising it to the + /// public `f64` contract. + async fn referral_admin_stats_global( + &self, + inviter_user_id: Option<&str>, + ) -> Result { + let Some(backends) = self.backends.as_ref() else { + return Ok(ReferralAdminStats::default()); + }; + + #[cfg(feature = "postgres")] + if let Some(backend) = backends.postgres() { + let row = sqlx::query( + r#" +SELECT + (SELECT COUNT(*) FROM user_referrals + WHERE ($1::TEXT IS NULL OR inviter_user_id = $1)) AS total_invites, + (SELECT COUNT(*) FROM user_referrals + WHERE ($1::TEXT IS NULL OR inviter_user_id = $1) + AND first_paid_order_id IS NOT NULL) + AS effective_invites, + CAST(COALESCE(SUM(CASE + WHEN status = 'applied' AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS DOUBLE PRECISION) AS paid_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN status IN ('pending', 'failed', 'applying') AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS DOUBLE PRECISION) AS pending_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN reversed_amount_usd > 0 THEN reversed_amount_usd ELSE 0 END), 0) + AS DOUBLE PRECISION) AS reversed_reward_usd +FROM referral_rewards +WHERE ($1::TEXT IS NULL OR inviter_user_id = $1) +"#, + ) + .bind(inviter_user_id) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + return Ok(ReferralAdminStats { + total_invites: referral_stats_count( + row.try_get::("total_invites") + .map_err(DataLayerError::postgres)?, + ), + effective_invites: referral_stats_count( + row.try_get::("effective_invites") + .map_err(DataLayerError::postgres)?, + ), + paid_reward_usd: referral_stats_amount( + row.try_get::("paid_reward_usd") + .map_err(DataLayerError::postgres)?, + ), + pending_reward_usd: referral_stats_amount( + row.try_get::("pending_reward_usd") + .map_err(DataLayerError::postgres)?, + ), + reversed_reward_usd: referral_stats_amount( + row.try_get::("reversed_reward_usd") + .map_err(DataLayerError::postgres)?, + ), + }); + } + + #[cfg(feature = "mysql")] + if let Some(backend) = backends.mysql() { + let row = sqlx::query( + r#" +SELECT + (SELECT COUNT(*) FROM user_referrals + WHERE (? IS NULL OR inviter_user_id = ?)) AS total_invites, + (SELECT COUNT(*) FROM user_referrals + WHERE (? IS NULL OR inviter_user_id = ?) + AND first_paid_order_id IS NOT NULL) + AS effective_invites, + CAST(COALESCE(SUM(CASE + WHEN status = 'applied' AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS DOUBLE) AS paid_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN status IN ('pending', 'failed', 'applying') AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS DOUBLE) AS pending_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN reversed_amount_usd > 0 THEN reversed_amount_usd ELSE 0 END), 0) + AS DOUBLE) AS reversed_reward_usd +FROM referral_rewards +WHERE (? IS NULL OR inviter_user_id = ?) +"#, + ) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return Ok(ReferralAdminStats { + total_invites: referral_stats_count( + row.try_get::("total_invites") + .map_err(DataLayerError::sql)?, + ), + effective_invites: referral_stats_count( + row.try_get::("effective_invites") + .map_err(DataLayerError::sql)?, + ), + paid_reward_usd: referral_stats_amount( + row.try_get::("paid_reward_usd") + .map_err(DataLayerError::sql)?, + ), + pending_reward_usd: referral_stats_amount( + row.try_get::("pending_reward_usd") + .map_err(DataLayerError::sql)?, + ), + reversed_reward_usd: referral_stats_amount( + row.try_get::("reversed_reward_usd") + .map_err(DataLayerError::sql)?, + ), + }); + } + + #[cfg(feature = "sqlite")] + if let Some(backend) = backends.sqlite() { + let row = sqlx::query( + r#" +SELECT + (SELECT COUNT(*) FROM user_referrals + WHERE (? IS NULL OR inviter_user_id = ?)) AS total_invites, + (SELECT COUNT(*) FROM user_referrals + WHERE (? IS NULL OR inviter_user_id = ?) + AND first_paid_order_id IS NOT NULL) + AS effective_invites, + CAST(COALESCE(SUM(CASE + WHEN status = 'applied' AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS REAL) AS paid_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN status IN ('pending', 'failed', 'applying') AND amount_usd > 0 THEN amount_usd ELSE 0 END), 0) + AS REAL) AS pending_reward_usd, + CAST(COALESCE(SUM(CASE + WHEN reversed_amount_usd > 0 THEN reversed_amount_usd ELSE 0 END), 0) + AS REAL) AS reversed_reward_usd +FROM referral_rewards +WHERE (? IS NULL OR inviter_user_id = ?) +"#, + ) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .bind(inviter_user_id) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return Ok(ReferralAdminStats { + total_invites: referral_stats_count( + row.try_get::("total_invites") + .map_err(DataLayerError::sql)?, + ), + effective_invites: referral_stats_count( + row.try_get::("effective_invites") + .map_err(DataLayerError::sql)?, + ), + paid_reward_usd: referral_stats_amount( + row.try_get::("paid_reward_usd") + .map_err(DataLayerError::sql)?, + ), + pending_reward_usd: referral_stats_amount( + row.try_get::("pending_reward_usd") + .map_err(DataLayerError::sql)?, + ), + reversed_reward_usd: referral_stats_amount( + row.try_get::("reversed_reward_usd") + .map_err(DataLayerError::sql)?, + ), + }); + } + + Ok(ReferralAdminStats::default()) } pub async fn bind_referral_invite_code( @@ -557,7 +895,7 @@ WHERE id = ? amount_usd: f64, trigger_point: &str, ) -> Result, DataLayerError> { - if amount_usd <= 0.0 { + if !amount_usd.is_finite() || amount_usd <= 0.0 { return Ok(Vec::new()); } let Some(relationship) = self @@ -594,7 +932,10 @@ WHERE id = ? let Some(context) = self.find_referral_payment_order_context(order_id).await? else { return Ok(Vec::new()); }; - if context.status != "credited" { + if context.status != "credited" + || !context.amount_usd.is_finite() + || context.amount_usd <= 0.0 + { return Ok(Vec::new()); } if !matches!( @@ -603,10 +944,7 @@ WHERE id = ? ) { return Ok(Vec::new()); } - if matches!( - context.payment_method.as_str(), - "manual" | "admin_manual" | "redeem_code" | "gift" - ) { + if referral_payment_method_excluded(&context.payment_method) { return Ok(Vec::new()); } let Some(relationship) = self @@ -615,14 +953,18 @@ WHERE id = ? else { return Ok(Vec::new()); }; - let marked_first_paid = self + let newly_marked_first_paid = self .mark_referral_first_paid_order(&relationship.id, &context.id) .await?; + // A replay of the winning order must repair a crash between marking + // first-paid and inserting its idempotent reward row. + let owns_first_paid_order = newly_marked_first_paid + || relationship.first_paid_order_id.as_deref() == Some(context.id.as_str()); let mut idempotency_keys = Vec::new(); - if config.percent_enabled && config.percent_rate > 0.0 { + if config.percent_enabled && referral_percent_rate_valid(config.percent_rate) { let amount_usd = (context.amount_usd * config.percent_rate / 100.0).max(0.0); - if amount_usd > 0.0 { + if amount_usd.is_finite() && amount_usd > 0.0 { let idempotency_key = format!("referral:{}:percent:{}", relationship.id, context.id); self.insert_referral_reward( @@ -638,9 +980,10 @@ WHERE id = ? } } if config.headcount_enabled + && config.headcount_amount_usd.is_finite() && config.headcount_amount_usd > 0.0 && config.headcount_trigger == "first_paid_order" - && marked_first_paid + && owns_first_paid_order { let idempotency_key = format!("referral:{}:headcount:first_paid_order", relationship.id); @@ -659,8 +1002,27 @@ WHERE id = ? if idempotency_keys.is_empty() { return Ok(Vec::new()); } - self.credit_pending_referral_rewards(&idempotency_keys, None, None) - .await + let mut rewards = self + .credit_pending_referral_rewards(&idempotency_keys, None, None) + .await?; + + // Payment credit and referral application use separate transactions. + // Whichever side wins a race with a refund must converge on the same + // cumulative reversal state. + if self + .find_referral_payment_order_refund_context(&context.id) + .await? + .is_some_and(|refund| refund.refunded_amount_usd > 0.0) + { + self.reverse_referral_rewards_for_order(&context.id, context.amount_usd) + .await?; + for reward in &mut rewards { + if let Some(updated) = self.find_referral_reward(&reward.id).await? { + *reward = updated; + } + } + } + Ok(rewards) } pub async fn retry_referral_reward( @@ -677,10 +1039,44 @@ WHERE id = ? "仅失败返利可以补发".to_string(), )); } + if !reward.amount_usd.is_finite() || reward.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "返利金额无效,无法补发".to_string(), + )); + } let rewards = self .credit_pending_referral_rewards(&[reward.idempotency_key], operator_id, note) .await?; - Ok(rewards.into_iter().next()) + let Some(mut updated) = rewards.into_iter().next() else { + return Ok(None); + }; + + // A manual retry can race with a payment refund. Do the same + // refund-aware reconciliation as the normal paid-order path so a + // successful retry can never leave a newly credited, already-refunded + // order permanently over-rewarded. + if updated.status == "applied" { + if let Some(order_id) = updated.source_order_id.as_deref() { + let refund_amount = self + .find_referral_payment_order_refund_context(order_id) + .await? + .and_then(|refund| { + let valid = refund.amount_usd.is_finite() + && refund.amount_usd > 0.0 + && refund.refunded_amount_usd.is_finite() + && refund.refunded_amount_usd > 0.0; + valid.then_some(refund.refunded_amount_usd) + }); + if let Some(refund_amount) = refund_amount { + self.reverse_referral_rewards_for_order(order_id, refund_amount) + .await?; + if let Some(refreshed) = self.find_referral_reward(&updated.id).await? { + updated = refreshed; + } + } + } + } + Ok(Some(updated)) } pub async fn void_referral_reward( @@ -707,7 +1103,7 @@ WHERE id = ? order_id: &str, amount_usd: f64, ) -> Result, DataLayerError> { - if amount_usd <= 0.0 { + if !amount_usd.is_finite() || amount_usd <= 0.0 { return Ok(Vec::new()); } let Some(refund_context) = self @@ -716,6 +1112,9 @@ WHERE id = ? else { return Ok(Vec::new()); }; + if !referral_refund_context_valid(&refund_context) { + return Ok(Vec::new()); + } let rewards = self .find_applied_referral_rewards_by_order(order_id) .await?; @@ -731,19 +1130,173 @@ WHERE id = ? if reversal_amount <= 0.0 { continue; } - let target_reversal = referral_reversal_target( - reward.amount_usd, - refund_context.amount_usd, - refund_context.refunded_amount_usd, - ); - self.apply_referral_reward_reversal(&reward, target_reversal) - .await?; + // The reversal transaction re-reads and locks the source payment + // order before calculating its target. The context above is only + // the caller-side eligibility check and may be stale by now. + self.apply_referral_reward_reversal(&reward).await?; if let Some(updated) = self.find_referral_reward(&reward.id).await? { reversed.push(updated); } } Ok(reversed) } + + /// Reconcile durable referral obligations left behind by an interrupted + /// payment callback or by a temporarily unavailable inviter wallet. + /// + /// The current reward configuration is accepted for API compatibility, but + /// it is deliberately not used to infer missing rows from payment history: + /// configuration has no historical snapshot, so doing that would + /// retroactively apply today's rate/mode to orders made before the feature + /// was enabled (or while a different mode was active). Only durable + /// pending/failed/applying reward rows and reversal debts are retried. + pub async fn reconcile_referral_rewards_once( + &self, + _reward_config: Option, + ) -> Result { + if self.backends.is_none() { + return Ok(ReferralReconciliationSummary::default()); + } + + let mut summary = ReferralReconciliationSummary::default(); + let mut first_error = None; + + let reward_keys = match self.list_referral_reward_retry_keys().await { + Ok(keys) => keys, + Err(error) => { + if let Some(first_error) = first_error { + return Err(first_error); + } + return Err(error); + } + }; + + // Retry rows whose reward credit transaction did not reach `applied`. + // Process each key independently so one broken wallet does not starve + // unrelated referral rewards in the same pass. + for idempotency_key in reward_keys { + summary.reward_attempted += 1; + let result = self + .credit_pending_referral_rewards(std::slice::from_ref(&idempotency_key), None, None) + .await; + match result { + Ok(updated) if updated.iter().any(|item| item.status == "applied") => { + summary.reward_applied += 1; + } + Ok(_) => summary.deferred += 1, + Err(error) => { + summary.deferred += 1; + if first_error.is_none() { + first_error = Some(error); + } + } + } + } + + // An older implementation could commit the intermediate `applying` + // state independently from the wallet credit. Resolve those rows from + // the durable wallet transaction fact, never by crediting them again. + // Rows without a matching transaction become `failed` and are only + // eligible for the normal credit path on a later pass. + let applying_reward_ids = match self.list_applying_referral_reward_ids().await { + Ok(ids) => ids, + Err(error) => { + if first_error.is_none() { + first_error = Some(error); + } + Vec::new() + } + }; + for reward_id in applying_reward_ids { + summary.reward_attempted += 1; + match self.recover_applying_referral_reward(&reward_id).await { + Ok(ReferralApplyingRecovery::Applied) => summary.reward_applied += 1, + Ok(ReferralApplyingRecovery::Failed | ReferralApplyingRecovery::Unchanged) => { + summary.deferred += 1; + } + Err(error) => { + summary.deferred += 1; + if first_error.is_none() { + first_error = Some(error); + } + } + } + } + + // Refresh the rows after reward retries. A reward that was applied in + // the first phase may itself carry an outstanding refund reversal. + let rewards = match self.list_referral_reversal_candidates().await { + Ok(rewards) => rewards, + Err(error) => { + if first_error.is_none() { + first_error = Some(error); + } + Vec::new() + } + }; + for reward in rewards.iter().take(REFERRAL_RECONCILIATION_LIMIT) { + let Some(order_id) = reward.source_order_id.as_deref() else { + summary.deferred += 1; + continue; + }; + let refund_context = match self + .find_referral_payment_order_refund_context(order_id) + .await + { + Ok(Some(context)) if referral_refund_context_valid(&context) => context, + Ok(Some(_)) => { + // Do not let a malformed historical order authorize a + // pending reversal. Pending debt is retried only after + // its source refund can be validated again. + summary.deferred += 1; + continue; + } + Ok(None) => { + summary.deferred += 1; + continue; + } + Err(error) => { + summary.deferred += 1; + if first_error.is_none() { + first_error = Some(error); + } + continue; + } + }; + let target_reversal = referral_reversal_target( + reward.amount_usd, + refund_context.amount_usd, + refund_context.refunded_amount_usd, + ); + let due = referral_reversal_due_bounded( + target_reversal, + reward.amount_usd, + reward.reversed_amount_usd, + reward.pending_reversal_amount_usd, + ); + if due <= f64::EPSILON { + continue; + } + summary.reversal_attempted += 1; + // `target_reversal` was calculated from the candidate-list + // snapshot. The transaction below obtains a fresh, locked order + // row and recalculates it before mutating either balance or debt. + match self.apply_referral_reward_reversal(reward).await { + Ok(()) => summary.reversal_applied += 1, + Err(error) => { + summary.deferred += 1; + if first_error.is_none() { + first_error = Some(error); + } + } + } + } + + if let Some(error) = first_error { + return Err(error); + } + Ok(summary) + } } impl ReferralDataState<'_> { @@ -1024,14 +1577,43 @@ VALUES (?, ?, ?, ?, ?, ?, ?) async fn list_referral_relationships_raw( &self, - inviter_user_id: Option<&str>, - invitee_user_id: Option<&str>, - ) -> Result, DataLayerError> { + query: &ReferralRelationshipListQuery, + ) -> Result<(Vec, u64), DataLayerError> { let Some(backends) = self.backends.as_ref() else { - return Ok(Vec::new()); + return Ok((Vec::new(), 0)); }; + let inviter_pattern = referral_like_pattern(query.inviter.as_deref()); + let invitee_pattern = referral_like_pattern(query.invitee.as_deref()); + let invite_code_pattern = referral_like_pattern(query.invite_code.as_deref()); + let first_paid = query + .first_paid + .map(|value| i64::from(value as u8)) + .unwrap_or(-1); + let (limit, offset) = referral_page_bounds(query.limit, query.offset); #[cfg(feature = "postgres")] if let Some(backend) = backends.postgres() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM user_referrals r +LEFT JOIN users inviter ON inviter.id = r.inviter_user_id +LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +WHERE ($1 = '' OR LOWER(COALESCE(inviter.username, '')) LIKE $1 ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE $1 ESCAPE '!') + AND ($2 = '' OR LOWER(COALESCE(invitee.username, '')) LIKE $2 ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE $2 ESCAPE '!') + AND ($3 = '' OR LOWER(r.invite_code_snapshot) LIKE $3 ESCAPE '!') + AND ($4 < 0 OR ($4 = 1 AND r.first_paid_order_id IS NOT NULL) OR ($4 = 0 AND r.first_paid_order_id IS NULL)) +"#, + ) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::postgres)?; let rows = sqlx::query( r#" SELECT @@ -1044,22 +1626,60 @@ SELECT FROM user_referrals r LEFT JOIN users inviter ON inviter.id = r.inviter_user_id LEFT JOIN users invitee ON invitee.id = r.invitee_user_id -WHERE ($1::TEXT IS NULL OR r.inviter_user_id = $1) - AND ($2::TEXT IS NULL OR r.invitee_user_id = $2) -ORDER BY r.created_at DESC -LIMIT $3 +WHERE ($1 = '' OR LOWER(COALESCE(inviter.username, '')) LIKE $1 ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE $1 ESCAPE '!') + AND ($2 = '' OR LOWER(COALESCE(invitee.username, '')) LIKE $2 ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE $2 ESCAPE '!') + AND ($3 = '' OR LOWER(r.invite_code_snapshot) LIKE $3 ESCAPE '!') + AND ($4 < 0 OR ($4 = 1 AND r.first_paid_order_id IS NOT NULL) OR ($4 = 0 AND r.first_paid_order_id IS NULL)) +ORDER BY r.created_at DESC, r.id DESC +LIMIT $5 OFFSET $6 "#, ) - .bind(inviter_user_id) - .bind(invitee_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .bind(limit) + .bind(offset) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::postgres)?; - return rows.iter().map(|row| relationship_from_row!(row)).collect(); + let items = rows + .iter() + .map(|row| relationship_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); } #[cfg(feature = "mysql")] if let Some(backend) = backends.mysql() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM user_referrals r +LEFT JOIN users inviter ON inviter.id = r.inviter_user_id +LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +WHERE (? = '' OR LOWER(COALESCE(inviter.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(COALESCE(invitee.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(r.invite_code_snapshot) LIKE ? ESCAPE '!') + AND (? < 0 OR (? = 1 AND r.first_paid_order_id IS NOT NULL) OR (? = 0 AND r.first_paid_order_id IS NULL)) +"#, + ) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .bind(first_paid) + .bind(first_paid) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::sql)?; let rows = sqlx::query( r#" SELECT @@ -1072,24 +1692,67 @@ SELECT FROM user_referrals r LEFT JOIN users inviter ON inviter.id = r.inviter_user_id LEFT JOIN users invitee ON invitee.id = r.invitee_user_id -WHERE (? IS NULL OR r.inviter_user_id = ?) - AND (? IS NULL OR r.invitee_user_id = ?) -ORDER BY r.created_at DESC -LIMIT ? +WHERE (? = '' OR LOWER(COALESCE(inviter.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(COALESCE(invitee.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(r.invite_code_snapshot) LIKE ? ESCAPE '!') + AND (? < 0 OR (? = 1 AND r.first_paid_order_id IS NOT NULL) OR (? = 0 AND r.first_paid_order_id IS NULL)) +ORDER BY r.created_at DESC, r.id DESC +LIMIT ? OFFSET ? "#, ) - .bind(inviter_user_id) - .bind(inviter_user_id) - .bind(invitee_user_id) - .bind(invitee_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .bind(first_paid) + .bind(first_paid) + .bind(limit) + .bind(offset) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::sql)?; - return rows.iter().map(|row| relationship_from_row!(row)).collect(); + let items = rows + .iter() + .map(|row| relationship_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); } #[cfg(feature = "sqlite")] if let Some(backend) = backends.sqlite() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM user_referrals r +LEFT JOIN users inviter ON inviter.id = r.inviter_user_id +LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +WHERE (? = '' OR LOWER(COALESCE(inviter.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(COALESCE(invitee.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(r.invite_code_snapshot) LIKE ? ESCAPE '!') + AND (? < 0 OR (? = 1 AND r.first_paid_order_id IS NOT NULL) OR (? = 0 AND r.first_paid_order_id IS NULL)) +"#, + ) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .bind(first_paid) + .bind(first_paid) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::sql)?; let rows = sqlx::query( r#" SELECT @@ -1102,23 +1765,37 @@ SELECT FROM user_referrals r LEFT JOIN users inviter ON inviter.id = r.inviter_user_id LEFT JOIN users invitee ON invitee.id = r.invitee_user_id -WHERE (? IS NULL OR r.inviter_user_id = ?) - AND (? IS NULL OR r.invitee_user_id = ?) -ORDER BY r.created_at DESC -LIMIT ? +WHERE (? = '' OR LOWER(COALESCE(inviter.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.inviter_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(COALESCE(invitee.username, '')) LIKE ? ESCAPE '!' OR LOWER(r.invitee_user_id) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(r.invite_code_snapshot) LIKE ? ESCAPE '!') + AND (? < 0 OR (? = 1 AND r.first_paid_order_id IS NOT NULL) OR (? = 0 AND r.first_paid_order_id IS NULL)) +ORDER BY r.created_at DESC, r.id DESC +LIMIT ? OFFSET ? "#, ) - .bind(inviter_user_id) - .bind(inviter_user_id) - .bind(invitee_user_id) - .bind(invitee_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&inviter_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invitee_pattern) + .bind(&invite_code_pattern) + .bind(&invite_code_pattern) + .bind(first_paid) + .bind(first_paid) + .bind(first_paid) + .bind(limit) + .bind(offset) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::sql)?; - return rows.iter().map(|row| relationship_from_row!(row)).collect(); + let items = rows + .iter() + .map(|row| relationship_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); } - Ok(Vec::new()) + Ok((Vec::new(), 0)) } async fn find_referral_relationship( @@ -1222,8 +1899,10 @@ SELECT r.source_json::TEXT AS source_json, EXTRACT(EPOCH FROM r.created_at)::BIGINT AS created_at_unix_secs FROM user_referrals r -LEFT JOIN users inviter ON inviter.id = r.inviter_user_id -LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +JOIN users inviter ON inviter.id = r.inviter_user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE +JOIN users invitee ON invitee.id = r.invitee_user_id + AND invitee.is_active IS TRUE AND invitee.is_deleted IS FALSE WHERE r.invitee_user_id = $1 LIMIT 1 "#, @@ -1246,8 +1925,10 @@ SELECT r.source_json AS source_json, r.created_at AS created_at_unix_secs FROM user_referrals r -LEFT JOIN users inviter ON inviter.id = r.inviter_user_id -LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +JOIN users inviter ON inviter.id = r.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +JOIN users invitee ON invitee.id = r.invitee_user_id + AND invitee.is_active = 1 AND invitee.is_deleted = 0 WHERE r.invitee_user_id = ? LIMIT 1 "#, @@ -1270,8 +1951,10 @@ SELECT r.source_json AS source_json, r.created_at AS created_at_unix_secs FROM user_referrals r -LEFT JOIN users inviter ON inviter.id = r.inviter_user_id -LEFT JOIN users invitee ON invitee.id = r.invitee_user_id +JOIN users inviter ON inviter.id = r.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +JOIN users invitee ON invitee.id = r.invitee_user_id + AND invitee.is_active = 1 AND invitee.is_deleted = 0 WHERE r.invitee_user_id = ? LIMIT 1 "#, @@ -1287,7 +1970,335 @@ LIMIT 1 async fn list_referral_rewards_raw( &self, - inviter_user_id: Option<&str>, + query: &ReferralRewardListQuery, + ) -> Result<(Vec, u64), DataLayerError> { + let Some(backends) = self.backends.as_ref() else { + return Ok((Vec::new(), 0)); + }; + let order_pattern = referral_like_pattern(query.order_id.as_deref()); + let reward_type_pattern = referral_like_pattern(query.reward_type.as_deref()); + let status_pattern = referral_like_pattern(query.status.as_deref()); + let (limit, offset) = referral_page_bounds(query.limit, query.offset); + #[cfg(feature = "postgres")] + if let Some(backend) = backends.postgres() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM referral_rewards +WHERE ($1 = '' OR LOWER(COALESCE(source_order_id, '')) LIKE $1 ESCAPE '!') + AND ($2 = '' OR LOWER(reward_type) LIKE $2 ESCAPE '!') + AND ($3 = '' OR LOWER(status) LIKE $3 ESCAPE '!') +"#, + ) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::postgres)?; + let rows = sqlx::query( + r#" +SELECT + id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, + EXTRACT(EPOCH FROM created_at)::BIGINT AS created_at_unix_secs, + EXTRACT(EPOCH FROM updated_at)::BIGINT AS updated_at_unix_secs +FROM referral_rewards +WHERE ($1 = '' OR LOWER(COALESCE(source_order_id, '')) LIKE $1 ESCAPE '!') + AND ($2 = '' OR LOWER(reward_type) LIKE $2 ESCAPE '!') + AND ($3 = '' OR LOWER(status) LIKE $3 ESCAPE '!') +ORDER BY created_at DESC, id DESC +LIMIT $4 OFFSET $5 +"#, + ) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .bind(limit) + .bind(offset) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + let items = rows + .iter() + .map(|row| reward_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); + } + #[cfg(feature = "mysql")] + if let Some(backend) = backends.mysql() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM referral_rewards +WHERE (? = '' OR LOWER(COALESCE(source_order_id, '')) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(reward_type) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(status) LIKE ? ESCAPE '!') +"#, + ) + .bind(&order_pattern) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .bind(&status_pattern) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::sql)?; + let rows = sqlx::query( + r#" +SELECT + id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, + trigger_point, CAST(amount_usd AS DOUBLE) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, + created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs +FROM referral_rewards +WHERE (? = '' OR LOWER(COALESCE(source_order_id, '')) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(reward_type) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(status) LIKE ? ESCAPE '!') +ORDER BY created_at DESC, id DESC +LIMIT ? OFFSET ? +"#, + ) + .bind(&order_pattern) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .bind(&status_pattern) + .bind(limit) + .bind(offset) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let items = rows + .iter() + .map(|row| reward_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); + } + #[cfg(feature = "sqlite")] + if let Some(backend) = backends.sqlite() { + let count = sqlx::query( + r#" +SELECT COUNT(*) AS total +FROM referral_rewards +WHERE (? = '' OR LOWER(COALESCE(source_order_id, '')) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(reward_type) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(status) LIKE ? ESCAPE '!') +"#, + ) + .bind(&order_pattern) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .bind(&status_pattern) + .fetch_one(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let total = count + .try_get::("total") + .map_err(DataLayerError::sql)?; + let rows = sqlx::query( + r#" +SELECT + id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, + created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs +FROM referral_rewards +WHERE (? = '' OR LOWER(COALESCE(source_order_id, '')) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(reward_type) LIKE ? ESCAPE '!') + AND (? = '' OR LOWER(status) LIKE ? ESCAPE '!') +ORDER BY created_at DESC, id DESC +LIMIT ? OFFSET ? +"#, + ) + .bind(&order_pattern) + .bind(&order_pattern) + .bind(&reward_type_pattern) + .bind(&reward_type_pattern) + .bind(&status_pattern) + .bind(&status_pattern) + .bind(limit) + .bind(offset) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + let items = rows + .iter() + .map(|row| reward_from_row!(row)) + .collect::, _>>()?; + return Ok((items, total.max(0) as u64)); + } + Ok((Vec::new(), 0)) + } + + async fn list_applying_referral_reward_ids(&self) -> Result, DataLayerError> { + let Some(backends) = self.backends.as_ref() else { + return Ok(Vec::new()); + }; + #[cfg(feature = "postgres")] + if let Some(backend) = backends.postgres() { + let rows = sqlx::query( + r#" +SELECT id +FROM referral_rewards +WHERE status = 'applying' +ORDER BY updated_at ASC, created_at ASC, id ASC +LIMIT $1 +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + return rows.iter().map(|row| Ok(row_string!(row, "id"))).collect(); + } + #[cfg(feature = "mysql")] + if let Some(backend) = backends.mysql() { + let rows = sqlx::query( + r#" +SELECT id +FROM referral_rewards +WHERE status = 'applying' +ORDER BY updated_at ASC, created_at ASC, id ASC +LIMIT ? +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return rows.iter().map(|row| Ok(row_string!(row, "id"))).collect(); + } + #[cfg(feature = "sqlite")] + if let Some(backend) = backends.sqlite() { + let rows = sqlx::query( + r#" +SELECT id +FROM referral_rewards +WHERE status = 'applying' +ORDER BY updated_at ASC, created_at ASC, id ASC +LIMIT ? +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return rows.iter().map(|row| Ok(row_string!(row, "id"))).collect(); + } + Ok(Vec::new()) + } + + /// Select only rewards that can be credited now. Ineligible historical + /// rows must not occupy the bounded retry page and starve valid rewards. + async fn list_referral_reward_retry_keys(&self) -> Result, DataLayerError> { + let Some(backends) = self.backends.as_ref() else { + return Ok(Vec::new()); + }; + #[cfg(feature = "postgres")] + if let Some(backend) = backends.postgres() { + let rows = sqlx::query( + r#" +SELECT rw.idempotency_key +FROM referral_rewards rw +JOIN wallets ON wallets.user_id = rw.inviter_user_id + AND wallets.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE +WHERE rw.status IN ('pending', 'failed') + AND rw.amount_usd > 0 +ORDER BY rw.created_at ASC, rw.id ASC +LIMIT $1 +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::postgres)?; + return rows + .iter() + .map(|row| Ok(row_string!(row, "idempotency_key"))) + .collect(); + } + #[cfg(feature = "mysql")] + if let Some(backend) = backends.mysql() { + let rows = sqlx::query( + r#" +SELECT rw.idempotency_key +FROM referral_rewards rw +JOIN wallets ON wallets.user_id = rw.inviter_user_id + AND wallets.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +WHERE rw.status IN ('pending', 'failed') + AND rw.amount_usd > 0 +ORDER BY rw.created_at ASC, rw.id ASC +LIMIT ? +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return rows + .iter() + .map(|row| Ok(row_string!(row, "idempotency_key"))) + .collect(); + } + #[cfg(feature = "sqlite")] + if let Some(backend) = backends.sqlite() { + let rows = sqlx::query( + r#" +SELECT rw.idempotency_key +FROM referral_rewards rw +JOIN wallets ON wallets.user_id = rw.inviter_user_id + AND wallets.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +WHERE rw.status IN ('pending', 'failed') + AND rw.amount_usd > 0 +ORDER BY rw.created_at ASC, rw.id ASC +LIMIT ? +"#, + ) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) + .fetch_all(&backend.pool_clone()) + .await + .map_err(DataLayerError::sql)?; + return rows + .iter() + .map(|row| Ok(row_string!(row, "idempotency_key"))) + .collect(); + } + Ok(Vec::new()) + } + + /// Return rewards that can have a refund reversal. Filtering against the + /// payment order here is important: a reward may be newly applied after a + /// refund has already completed, in which case its pending column is + /// still zero and a pending-only scan would miss it forever. + async fn list_referral_reversal_candidates( + &self, ) -> Result, DataLayerError> { let Some(backends) = self.backends.as_ref() else { return Ok(Vec::new()); @@ -1297,19 +2308,77 @@ LIMIT 1 let rows = sqlx::query( r#" SELECT - id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, - EXTRACT(EPOCH FROM created_at)::BIGINT AS created_at_unix_secs, - EXTRACT(EPOCH FROM updated_at)::BIGINT AS updated_at_unix_secs -FROM referral_rewards -WHERE ($1::TEXT IS NULL OR inviter_user_id = $1) -ORDER BY created_at DESC -LIMIT $2 + rw.id, rw.referral_id, rw.inviter_user_id, rw.invitee_user_id, + rw.reward_type, rw.source_order_id, rw.trigger_point, + CAST(rw.amount_usd AS DOUBLE PRECISION) AS amount_usd, + rw.status, rw.wallet_transaction_id, rw.idempotency_key, + CAST(rw.reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(rw.pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + rw.admin_operator_id, rw.admin_note, + EXTRACT(EPOCH FROM rw.created_at)::BIGINT AS created_at_unix_secs, + EXTRACT(EPOCH FROM rw.updated_at)::BIGINT AS updated_at_unix_secs +FROM referral_rewards rw +JOIN ( + SELECT + po0.id, + CAST(po0.amount_usd AS DOUBLE PRECISION) AS amount_usd, + po0.credited_at, + po0.paid_at, + po0.created_at, + CAST( + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + END AS DOUBLE PRECISION + ) AS refunded_amount_usd + FROM payment_orders po0 +) po ON po.id = rw.source_order_id +JOIN wallets wallet ON wallet.user_id = rw.inviter_user_id + AND wallet.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE +WHERE rw.status IN ('applied', 'reversed') + AND po.refunded_amount_usd > 0 + AND ( + rw.pending_reversal_amount_usd > 0.00000001 + OR ( + po.amount_usd > 0 + AND rw.amount_usd > 0 + AND rw.reversed_amount_usd + 0.00000001 < + rw.amount_usd * CASE + WHEN po.refunded_amount_usd >= po.amount_usd THEN 1.0 + ELSE po.refunded_amount_usd / po.amount_usd + END + ) + ) +ORDER BY COALESCE(po.credited_at, po.paid_at, po.created_at) ASC, + rw.created_at ASC, rw.id ASC +LIMIT $1 "#, ) - .bind(inviter_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::postgres)?; @@ -1320,19 +2389,73 @@ LIMIT $2 let rows = sqlx::query( r#" SELECT - id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, - created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs -FROM referral_rewards -WHERE (? IS NULL OR inviter_user_id = ?) -ORDER BY created_at DESC + rw.id, rw.referral_id, rw.inviter_user_id, rw.invitee_user_id, + rw.reward_type, rw.source_order_id, rw.trigger_point, rw.amount_usd, + rw.status, rw.wallet_transaction_id, rw.idempotency_key, + rw.reversed_amount_usd, rw.pending_reversal_amount_usd, + rw.admin_operator_id, rw.admin_note, + rw.created_at AS created_at_unix_secs, + rw.updated_at AS updated_at_unix_secs +FROM referral_rewards rw +JOIN ( + SELECT + po0.id, + po0.amount_usd, + po0.credited_at, + po0.paid_at, + po0.created_at, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd + FROM payment_orders po0 +) po ON po.id = rw.source_order_id +JOIN wallets wallet ON wallet.user_id = rw.inviter_user_id + AND wallet.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +WHERE rw.status IN ('applied', 'reversed') + AND po.refunded_amount_usd > 0 + AND ( + rw.pending_reversal_amount_usd > 0.00000001 + OR ( + po.amount_usd > 0 + AND rw.amount_usd > 0 + AND rw.reversed_amount_usd + 0.00000001 < + rw.amount_usd * CASE + WHEN po.refunded_amount_usd >= po.amount_usd THEN 1.0 + ELSE po.refunded_amount_usd / po.amount_usd + END + ) + ) +ORDER BY COALESCE(po.credited_at, po.paid_at, po.created_at) ASC, + rw.created_at ASC, rw.id ASC LIMIT ? "#, ) - .bind(inviter_user_id) - .bind(inviter_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::sql)?; @@ -1343,19 +2466,73 @@ LIMIT ? let rows = sqlx::query( r#" SELECT - id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, - created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs -FROM referral_rewards -WHERE (? IS NULL OR inviter_user_id = ?) -ORDER BY created_at DESC + rw.id, rw.referral_id, rw.inviter_user_id, rw.invitee_user_id, + rw.reward_type, rw.source_order_id, rw.trigger_point, rw.amount_usd, + rw.status, rw.wallet_transaction_id, rw.idempotency_key, + rw.reversed_amount_usd, rw.pending_reversal_amount_usd, + rw.admin_operator_id, rw.admin_note, + rw.created_at AS created_at_unix_secs, + rw.updated_at AS updated_at_unix_secs +FROM referral_rewards rw +JOIN ( + SELECT + po0.id, + po0.amount_usd, + po0.credited_at, + po0.paid_at, + po0.created_at, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po0.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po0.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd + FROM payment_orders po0 +) po ON po.id = rw.source_order_id +JOIN wallets wallet ON wallet.user_id = rw.inviter_user_id + AND wallet.status = 'active' +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 +WHERE rw.status IN ('applied', 'reversed') + AND po.refunded_amount_usd > 0 + AND ( + rw.pending_reversal_amount_usd > 0.00000001 + OR ( + po.amount_usd > 0 + AND rw.amount_usd > 0 + AND rw.reversed_amount_usd + 0.00000001 < + rw.amount_usd * CASE + WHEN po.refunded_amount_usd >= po.amount_usd THEN 1.0 + ELSE po.refunded_amount_usd / po.amount_usd + END + ) + ) +ORDER BY COALESCE(po.credited_at, po.paid_at, po.created_at) ASC, + rw.created_at ASC, rw.id ASC LIMIT ? "#, ) - .bind(inviter_user_id) - .bind(inviter_user_id) - .bind(REFERRAL_FETCH_LIMIT as i64) + .bind(REFERRAL_RECONCILIATION_LIMIT as i64) .fetch_all(&backend.pool_clone()) .await .map_err(DataLayerError::sql)?; @@ -1377,8 +2554,11 @@ LIMIT ? r#" SELECT id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, EXTRACT(EPOCH FROM created_at)::BIGINT AS created_at_unix_secs, EXTRACT(EPOCH FROM updated_at)::BIGINT AS updated_at_unix_secs FROM referral_rewards @@ -1398,8 +2578,11 @@ LIMIT 1 r#" SELECT id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, + trigger_point, CAST(amount_usd AS DOUBLE) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs FROM referral_rewards WHERE id = ? @@ -1418,8 +2601,11 @@ LIMIT 1 r#" SELECT id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs FROM referral_rewards WHERE id = ? @@ -1448,8 +2634,11 @@ LIMIT 1 r#" SELECT id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, EXTRACT(EPOCH FROM created_at)::BIGINT AS created_at_unix_secs, EXTRACT(EPOCH FROM updated_at)::BIGINT AS updated_at_unix_secs FROM referral_rewards @@ -1519,12 +2708,16 @@ LIMIT 1 r#" SELECT id, referral_id, inviter_user_id, invitee_user_id, reward_type, source_order_id, - trigger_point, amount_usd, status, wallet_transaction_id, idempotency_key, - reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, + trigger_point, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + status, wallet_transaction_id, idempotency_key, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd, + admin_operator_id, admin_note, EXTRACT(EPOCH FROM created_at)::BIGINT AS created_at_unix_secs, EXTRACT(EPOCH FROM updated_at)::BIGINT AS updated_at_unix_secs FROM referral_rewards -WHERE source_order_id = $1 AND status = 'applied' +WHERE source_order_id = $1 + AND status IN ('applied', 'reversed') ORDER BY created_at ASC "#, ) @@ -1544,7 +2737,8 @@ SELECT reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs FROM referral_rewards -WHERE source_order_id = ? AND status = 'applied' +WHERE source_order_id = ? + AND status IN ('applied', 'reversed') ORDER BY created_at ASC "#, ) @@ -1564,7 +2758,8 @@ SELECT reversed_amount_usd, pending_reversal_amount_usd, admin_operator_id, admin_note, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs FROM referral_rewards -WHERE source_order_id = ? AND status = 'applied' +WHERE source_order_id = ? + AND status IN ('applied', 'reversed') ORDER BY created_at ASC "#, ) @@ -1687,7 +2882,8 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, 'pending', ?, ?, ?) if let Some(backend) = backends.postgres() { let row = sqlx::query( r#" -SELECT id, user_id, amount_usd, payment_method, status, order_kind +SELECT id, user_id, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + payment_method, status, order_kind FROM payment_orders WHERE id = $1 "#, @@ -1742,9 +2938,37 @@ WHERE id = ? if let Some(backend) = backends.postgres() { let row = sqlx::query( r#" -SELECT amount_usd, refunded_amount_usd -FROM payment_orders -WHERE id = $1 +SELECT CAST(po.amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST( + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS DOUBLE PRECISION + ) AS refunded_amount_usd +FROM payment_orders po +WHERE po.id = $1 "#, ) .bind(order_id) @@ -1757,9 +2981,35 @@ WHERE id = $1 if let Some(backend) = backends.mysql() { let row = sqlx::query( r#" -SELECT amount_usd, refunded_amount_usd -FROM payment_orders -WHERE id = ? +SELECT po.amount_usd, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd +FROM payment_orders po +WHERE po.id = ? "#, ) .bind(order_id) @@ -1772,9 +3022,35 @@ WHERE id = ? if let Some(backend) = backends.sqlite() { let row = sqlx::query( r#" -SELECT amount_usd, refunded_amount_usd -FROM payment_orders -WHERE id = ? +SELECT po.amount_usd, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd +FROM payment_orders po +WHERE po.id = ? "#, ) .bind(order_id) @@ -1939,6 +3215,172 @@ WHERE id = ? AND status IN ('pending', 'failed') Ok(false) } + async fn recover_applying_referral_reward( + &self, + reward_id: &str, + ) -> Result { + #[cfg(feature = "postgres")] + if let Some(backend) = self.backends.and_then(DataBackends::postgres) { + let mut tx = backend + .pool_clone() + .begin() + .await + .map_err(DataLayerError::postgres)?; + let reward = sqlx::query( + r#" +SELECT id, inviter_user_id, CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd +FROM referral_rewards +WHERE id = $1 AND status = 'applying' +FOR UPDATE +"#, + ) + .bind(reward_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + if reward.is_none() { + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(ReferralApplyingRecovery::Unchanged); + } + + let wallet_transactions = sqlx::query( + r#" +SELECT tx.id, + CAST(tx.amount AS DOUBLE PRECISION) AS amount, + CAST(tx.balance_before AS DOUBLE PRECISION) AS balance_before, + CAST(tx.balance_after AS DOUBLE PRECISION) AS balance_after, + CAST(tx.recharge_balance_before AS DOUBLE PRECISION) AS recharge_balance_before, + CAST(tx.recharge_balance_after AS DOUBLE PRECISION) AS recharge_balance_after, + CAST(tx.gift_balance_before AS DOUBLE PRECISION) AS gift_balance_before, + CAST(tx.gift_balance_after AS DOUBLE PRECISION) AS gift_balance_after +FROM wallet_transactions tx +JOIN wallets wallet ON wallet.id = tx.wallet_id +WHERE tx.category = 'adjust' + AND tx.reason_code = 'referral_reward' + AND tx.link_type = 'referral_reward' + AND tx.link_id = $1 + AND wallet.user_id = (SELECT inviter_user_id FROM referral_rewards WHERE id = $1) + AND tx.amount > 0 +ORDER BY tx.created_at ASC, tx.id ASC +LIMIT 32 +"#, + ) + .bind(reward_id) + .fetch_all(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + let reward_amount = reward + .as_ref() + .and_then(|row| row.try_get::("amount_usd").ok()) + .unwrap_or(0.0); + let has_wallet_transaction = !wallet_transactions.is_empty(); + let valid_wallet_transaction_ids = wallet_transactions + .into_iter() + .filter_map(|row| { + let amount = row.try_get::("amount").ok()?; + let balance_before = row.try_get::("balance_before").ok()?; + let balance_after = row.try_get::("balance_after").ok()?; + let recharge_balance_before = + row.try_get::("recharge_balance_before").ok()?; + let recharge_balance_after = + row.try_get::("recharge_balance_after").ok()?; + let gift_balance_before = row.try_get::("gift_balance_before").ok()?; + let gift_balance_after = row.try_get::("gift_balance_after").ok()?; + if !referral_credit_transaction_fact_valid( + reward_amount, + amount, + balance_before, + balance_after, + recharge_balance_before, + recharge_balance_after, + gift_balance_before, + gift_balance_after, + ) { + return None; + } + row.try_get::("id").ok() + }) + .collect::>(); + // Exactly one valid transaction fact is required. If multiple + // facts match the same reward, the historical write may already + // have credited the wallet twice; silently choosing the first + // would hide that ambiguity and make the ledger unreconcilable. + let wallet_transaction_id = (valid_wallet_transaction_ids.len() == 1) + .then(|| valid_wallet_transaction_ids[0].clone()); + let recovery = if !reward_amount.is_finite() || reward_amount <= 0.0 { + // A malformed durable amount must never enter the normal + // failed-reward retry path. Leave it for operator repair, + // just like an ambiguous wallet snapshot. + ReferralApplyingRecovery::Unchanged + } else if wallet_transaction_id.is_some() { + ReferralApplyingRecovery::Applied + } else if has_wallet_transaction { + // A matching transaction whose durable snapshot is malformed + // is evidence of an ambiguous historical write. Retrying it + // as a normal failed reward could credit the inviter twice. + // Keep the row applying until an operator repairs the fact. + ReferralApplyingRecovery::Unchanged + } else { + ReferralApplyingRecovery::Failed + }; + if recovery == ReferralApplyingRecovery::Unchanged { + // `applying` rows are processed in a bounded queue. Bump the + // retry timestamp for ambiguous facts so one permanently + // malformed row cannot occupy the oldest page forever. + sqlx::query( + "UPDATE referral_rewards SET updated_at = GREATEST(updated_at + INTERVAL '1 microsecond', NOW()) WHERE id = $1 AND status = 'applying'", + ) + .bind(reward_id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(recovery); + } + let status = match recovery { + ReferralApplyingRecovery::Applied => "applied", + ReferralApplyingRecovery::Failed => "failed", + ReferralApplyingRecovery::Unchanged => unreachable!(), + }; + sqlx::query( + r#" +UPDATE referral_rewards +SET status = $2, + wallet_transaction_id = $3, + updated_at = NOW() +WHERE id = $1 AND status = 'applying' +"#, + ) + .bind(reward_id) + .bind(status) + .bind(wallet_transaction_id.as_deref()) + .execute(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(recovery); + } + #[cfg(feature = "mysql")] + if let Some(backend) = self.backends.and_then(DataBackends::mysql) { + return self + .recover_applying_referral_reward_mysql_numeric_time( + &backend.pool_clone(), + reward_id, + ) + .await; + } + #[cfg(feature = "sqlite")] + if let Some(backend) = self.backends.and_then(DataBackends::sqlite) { + return self + .recover_applying_referral_reward_sqlite_numeric_time( + &backend.pool_clone(), + reward_id, + ) + .await; + } + Ok(ReferralApplyingRecovery::Unchanged) + } + async fn credit_pending_referral_rewards( &self, idempotency_keys: &[String], @@ -1970,12 +3412,17 @@ WHERE id = ? AND status IN ('pending', 'failed') let row = sqlx::query( r#" SELECT - rw.id, rw.inviter_user_id, rw.invitee_user_id, rw.amount_usd, rw.reward_type, + rw.id, rw.inviter_user_id, rw.invitee_user_id, + CAST(rw.amount_usd AS DOUBLE PRECISION) AS amount_usd, + rw.reward_type, rw.trigger_point, wallets.id AS wallet_id FROM referral_rewards rw JOIN wallets ON wallets.user_id = rw.inviter_user_id +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE WHERE rw.idempotency_key = $1 AND rw.status IN ('pending', 'failed') + AND wallets.status = 'active' "#, ) .bind(idempotency_key) @@ -1993,8 +3440,11 @@ SELECT rw.trigger_point, wallets.id AS wallet_id FROM referral_rewards rw JOIN wallets ON wallets.user_id = rw.inviter_user_id +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 WHERE rw.idempotency_key = ? AND rw.status IN ('pending', 'failed') + AND wallets.status = 'active' "#, ) .bind(idempotency_key) @@ -2012,8 +3462,11 @@ SELECT rw.trigger_point, wallets.id AS wallet_id FROM referral_rewards rw JOIN wallets ON wallets.user_id = rw.inviter_user_id +JOIN users inviter ON inviter.id = rw.inviter_user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 WHERE rw.idempotency_key = ? AND rw.status IN ('pending', 'failed') + AND wallets.status = 'active' "#, ) .bind(idempotency_key) @@ -2031,6 +3484,11 @@ WHERE rw.idempotency_key = ? operator_id: Option<&str>, note: Option<&str>, ) -> Result<(), DataLayerError> { + if !target.amount_usd.is_finite() || target.amount_usd <= 0.0 { + return Err(DataLayerError::InvalidInput( + "referral reward amount must be finite and greater than zero".to_string(), + )); + } #[cfg(feature = "postgres")] if let Some(backend) = self.backends.and_then(DataBackends::postgres) { let mut tx = backend @@ -2061,9 +3519,14 @@ WHERE id = $1 AND status IN ('pending', 'failed') } let wallet = sqlx::query( r#" -SELECT balance, gift_balance +SELECT CAST(wallets.balance AS DOUBLE PRECISION) AS balance, + CAST(wallets.gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(wallets.total_adjusted AS DOUBLE PRECISION) AS total_adjusted FROM wallets -WHERE id = $1 +JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE +WHERE wallets.id = $1 + AND wallets.status = 'active' FOR UPDATE "#, ) @@ -2093,7 +3556,27 @@ WHERE id = $1 }; let balance = row_f64!(wallet, "balance"); let gift_before = row_f64!(wallet, "gift_balance"); + let total_adjusted_before = row_f64!(wallet, "total_adjusted"); + if !referral_wallet_values_valid(balance, gift_before) + || !total_adjusted_before.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance is invalid".to_string(), + )); + } + let total_before = balance + gift_before; let gift_after = gift_before + target.amount_usd; + let total_after = balance + gift_after; + let total_adjusted_after = total_adjusted_before + target.amount_usd; + if !gift_after.is_finite() + || !total_before.is_finite() + || !total_after.is_finite() + || !total_adjusted_after.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance overflowed".to_string(), + )); + } let tx_id = uuid::Uuid::new_v4().to_string(); let description = note .map(ToOwned::to_owned) @@ -2102,14 +3585,14 @@ WHERE id = $1 r#" UPDATE wallets SET gift_balance = $2, - total_adjusted = total_adjusted + $3, + total_adjusted = $3, updated_at = NOW() WHERE id = $1 "#, ) .bind(&target.wallet_id) .bind(gift_after) - .bind(target.amount_usd) + .bind(total_adjusted_after) .execute(&mut *tx) .await .map_err(DataLayerError::postgres)?; @@ -2129,8 +3612,8 @@ VALUES ($1, $2, 'adjust', 'referral_reward', $3, $4, $5, $6, $6, $7, $8, .bind(&tx_id) .bind(&target.wallet_id) .bind(target.amount_usd) - .bind(balance + gift_before) - .bind(balance + gift_after) + .bind(total_before) + .bind(total_after) .bind(balance) .bind(gift_before) .bind(gift_after) @@ -2189,7 +3672,6 @@ WHERE id = $1 async fn apply_referral_reward_reversal( &self, reward: &ReferralRewardRecord, - target_reversal_amount_usd: f64, ) -> Result<(), DataLayerError> { #[cfg(feature = "postgres")] if let Some(backend) = self.backends.and_then(DataBackends::postgres) { @@ -2200,7 +3682,12 @@ WHERE id = $1 .map_err(DataLayerError::postgres)?; let reward_row = sqlx::query( r#" -SELECT reversed_amount_usd, pending_reversal_amount_usd +SELECT status, + inviter_user_id, + source_order_id, + CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST(reversed_amount_usd AS DOUBLE PRECISION) AS reversed_amount_usd, + CAST(pending_reversal_amount_usd AS DOUBLE PRECISION) AS pending_reversal_amount_usd FROM referral_rewards WHERE id = $1 FOR UPDATE @@ -2214,50 +3701,197 @@ FOR UPDATE tx.commit().await.map_err(DataLayerError::postgres)?; return Ok(()); }; + let reward_status = row_string!(reward_row, "status"); + if !matches!(reward_status.as_str(), "applied" | "reversed") { + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(()); + } + let Some(source_order_id) = row_optional_string!(reward_row, "source_order_id") else { + // Registration/headcount rewards have no payment source and + // therefore can never be authorized for a refund reversal. + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(()); + }; + let inviter_user_id = row_string!(reward_row, "inviter_user_id"); + let reward_amount = row_f64!(reward_row, "amount_usd"); let current_reversed = row_f64!(reward_row, "reversed_amount_usd"); let current_pending = row_f64!(reward_row, "pending_reversal_amount_usd"); - let amount_usd = - (target_reversal_amount_usd - current_reversed - current_pending).max(0.0); + + // Keep the lock order aligned with the wallet refund path + // (wallet -> payment order). The order is re-read after its row + // lock, so a refund committed after the caller's candidate query + // cannot leave this reversal using an obsolete target amount. + let wallet = sqlx::query( + r#" +SELECT wallets.id, + CAST(wallets.balance AS DOUBLE PRECISION) AS balance, + CAST(wallets.gift_balance AS DOUBLE PRECISION) AS gift_balance, + CAST(wallets.total_adjusted AS DOUBLE PRECISION) AS total_adjusted +FROM wallets +JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active IS TRUE AND inviter.is_deleted IS FALSE +WHERE wallets.user_id = $1 + AND wallets.status = 'active' +FOR UPDATE +"#, + ) + .bind(&inviter_user_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + let order_row = sqlx::query( + r#" +SELECT CAST(po.amount_usd AS DOUBLE PRECISION) AS amount_usd, + CAST( + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS DOUBLE PRECISION + ) AS refunded_amount_usd +FROM payment_orders po +WHERE po.id = $1 +FOR UPDATE +"#, + ) + .bind(&source_order_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; + let Some(order_row) = order_row else { + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(()); + }; + let refund_context = payment_order_refund_context_from_row(order_row)?; + if !referral_refund_context_valid(&refund_context) { + tx.commit().await.map_err(DataLayerError::postgres)?; + return Ok(()); + } + let target_reversal_amount_usd = referral_reversal_target( + reward_amount, + refund_context.amount_usd, + refund_context.refunded_amount_usd, + ); + if !referral_reversal_inputs_valid( + reward_amount, + target_reversal_amount_usd, + current_reversed, + current_pending, + ) { + return Err(DataLayerError::InvalidInput( + "referral reversal state is invalid".to_string(), + )); + } + let amount_usd = referral_reversal_due_bounded( + target_reversal_amount_usd, + reward_amount, + current_reversed, + current_pending, + ); if amount_usd <= 0.0 { tx.commit().await.map_err(DataLayerError::postgres)?; return Ok(()); } - let wallet = sqlx::query( - r#" -SELECT id, balance, gift_balance -FROM wallets -WHERE user_id = $1 -FOR UPDATE -"#, - ) - .bind(&reward.inviter_user_id) - .fetch_optional(&mut *tx) - .await - .map_err(DataLayerError::postgres)?; let Some(wallet) = wallet else { + // Keep the unrecovered amount durable even when the inviter + // wallet is temporarily absent/inactive. A later + // reconciliation pass can consume it after the wallet is + // restored. + sqlx::query( + r#" +UPDATE referral_rewards +SET pending_reversal_amount_usd = $2, + status = CASE + WHEN status = 'reversed' THEN 'applied' + ELSE status + END, + updated_at = NOW() +WHERE id = $1 +"#, + ) + .bind(&reward.id) + .bind(referral_pending_reversal_capped( + reward_amount, + current_reversed, + current_pending, + amount_usd, + )) + .execute(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; tx.commit().await.map_err(DataLayerError::postgres)?; return Ok(()); }; let wallet_id = row_string!(wallet, "id"); let balance = row_f64!(wallet, "balance"); let gift_before = row_f64!(wallet, "gift_balance"); - let actual_reverse = gift_before.min(amount_usd); + let total_adjusted_before = row_f64!(wallet, "total_adjusted"); + if !referral_wallet_values_valid(balance, gift_before) + || !total_adjusted_before.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance is invalid".to_string(), + )); + } + let actual_reverse = gift_before.max(0.0).min(amount_usd); let pending_reverse = (amount_usd - actual_reverse).max(0.0); let gift_after = gift_before - actual_reverse; + let total_before = balance + gift_before; + let total_after = balance + gift_after; + let total_adjusted_after = total_adjusted_before - actual_reverse; + if !actual_reverse.is_finite() + || !pending_reverse.is_finite() + || !gift_after.is_finite() + || !total_before.is_finite() + || !total_after.is_finite() + || !total_adjusted_after.is_finite() + || !referral_reversal_state_valid( + reward_amount, + current_reversed, + current_pending, + actual_reverse, + pending_reverse, + ) + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance overflowed".to_string(), + )); + } let tx_id = uuid::Uuid::new_v4().to_string(); if actual_reverse > 0.0 { sqlx::query( r#" UPDATE wallets SET gift_balance = $2, - total_adjusted = total_adjusted - $3, + total_adjusted = $3, updated_at = NOW() WHERE id = $1 "#, ) .bind(&wallet_id) .bind(gift_after) - .bind(actual_reverse) + .bind(total_adjusted_after) .execute(&mut *tx) .await .map_err(DataLayerError::postgres)?; @@ -2277,8 +3911,8 @@ VALUES ($1, $2, 'adjust', 'referral_reward_reversal', $3, $4, $5, $6, $6, $7, $8 .bind(&tx_id) .bind(&wallet_id) .bind(-actual_reverse) - .bind(balance + gift_before) - .bind(balance + gift_after) + .bind(total_before) + .bind(total_after) .bind(balance) .bind(gift_before) .bind(gift_after) @@ -2291,11 +3925,7 @@ VALUES ($1, $2, 'adjust', 'referral_reward_reversal', $3, $4, $5, $6, $6, $7, $8 r#" UPDATE referral_rewards SET reversed_amount_usd = reversed_amount_usd + $2, - pending_reversal_amount_usd = pending_reversal_amount_usd + $3, - status = CASE - WHEN reversed_amount_usd + $2 >= amount_usd THEN 'reversed' - ELSE status - END, + pending_reversal_amount_usd = $3, updated_at = NOW() WHERE id = $1 "#, @@ -2306,6 +3936,24 @@ WHERE id = $1 .execute(&mut *tx) .await .map_err(DataLayerError::postgres)?; + sqlx::query( + r#" +UPDATE referral_rewards +SET status = CASE + WHEN pending_reversal_amount_usd > 0.00000001 AND status = 'reversed' THEN 'applied' + WHEN pending_reversal_amount_usd <= 0.00000001 + AND reversed_amount_usd >= amount_usd + AND status IN ('applied', 'reversed') THEN 'reversed' + ELSE status + END, + updated_at = NOW() +WHERE id = $1 +"#, + ) + .bind(&reward.id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::postgres)?; tx.commit().await.map_err(DataLayerError::postgres)?; return Ok(()); } @@ -2313,7 +3961,7 @@ WHERE id = $1 { // MySQL/SQLite refunds use integer timestamps in the wallet tables. return self - .apply_referral_reward_reversal_numeric_time(reward, target_reversal_amount_usd) + .apply_referral_reward_reversal_numeric_time(reward) .await; } #[cfg(not(any(feature = "mysql", feature = "sqlite")))] @@ -2321,6 +3969,163 @@ WHERE id = $1 } } +#[cfg(any(feature = "mysql", feature = "sqlite"))] +macro_rules! referral_applying_recovery_numeric_method { + ($name:ident, $pool_ty:ty) => { + async fn $name( + &self, + pool: &$pool_ty, + reward_id: &str, + ) -> Result { + let mut tx = pool.begin().await.map_err(DataLayerError::sql)?; + + // Both drivers begin deferred transactions. This harmless write + // takes the write/row lock before the transaction fact is read. + sqlx::query( + "UPDATE referral_rewards SET updated_at = updated_at WHERE id = ? AND status = 'applying'", + ) + .bind(reward_id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + let reward = sqlx::query( + "SELECT id, inviter_user_id, amount_usd FROM referral_rewards WHERE id = ? AND status = 'applying'", + ) + .bind(reward_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + if reward.is_none() { + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(ReferralApplyingRecovery::Unchanged); + } + + let wallet_transactions = sqlx::query( + r#" +SELECT tx.id, + tx.amount, + tx.balance_before, + tx.balance_after, + tx.recharge_balance_before, + tx.recharge_balance_after, + tx.gift_balance_before, + tx.gift_balance_after +FROM wallet_transactions tx +JOIN wallets wallet ON wallet.id = tx.wallet_id +WHERE tx.category = 'adjust' + AND tx.reason_code = 'referral_reward' + AND tx.link_type = 'referral_reward' + AND tx.link_id = ? + AND wallet.user_id = (SELECT inviter_user_id FROM referral_rewards WHERE id = ?) + AND tx.amount > 0 +ORDER BY tx.created_at ASC, tx.id ASC +LIMIT 32 +"#, + ) + .bind(reward_id) + .bind(reward_id) + .fetch_all(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + let reward_amount = reward + .as_ref() + .and_then(|row| row.try_get::("amount_usd").ok()) + .unwrap_or(0.0); + let has_wallet_transaction = !wallet_transactions.is_empty(); + let valid_wallet_transaction_ids = wallet_transactions + .into_iter() + .filter_map(|row| { + let amount = row.try_get::("amount").ok()?; + let balance_before = row.try_get::("balance_before").ok()?; + let balance_after = row.try_get::("balance_after").ok()?; + let recharge_balance_before = row + .try_get::("recharge_balance_before") + .ok()?; + let recharge_balance_after = row + .try_get::("recharge_balance_after") + .ok()?; + let gift_balance_before = row.try_get::("gift_balance_before").ok()?; + let gift_balance_after = row.try_get::("gift_balance_after").ok()?; + if !referral_credit_transaction_fact_valid( + reward_amount, + amount, + balance_before, + balance_after, + recharge_balance_before, + recharge_balance_after, + gift_balance_before, + gift_balance_after, + ) { + return None; + } + row.try_get::("id").ok() + }) + .collect::>(); + // Multiple valid facts for one reward indicate a possible + // duplicate credit. Do not mark the reward applied by selecting + // an arbitrary transaction. + let wallet_transaction_id = (valid_wallet_transaction_ids.len() == 1) + .then(|| valid_wallet_transaction_ids[0].clone()); + let recovery = if !reward_amount.is_finite() || reward_amount <= 0.0 { + // A malformed durable amount must never enter the normal + // failed-reward retry path. Leave it for operator repair, + // just like an ambiguous wallet snapshot. + ReferralApplyingRecovery::Unchanged + } else if wallet_transaction_id.is_some() { + ReferralApplyingRecovery::Applied + } else if has_wallet_transaction { + // A matching transaction with an invalid snapshot is + // ambiguous: retrying it as failed could credit twice. + // Leave the reward applying until the historical fact is + // repaired by an operator. + ReferralApplyingRecovery::Unchanged + } else { + ReferralApplyingRecovery::Failed + }; + if recovery == ReferralApplyingRecovery::Unchanged { + // `applying` rows are processed in a bounded queue. Bump the + // retry timestamp for ambiguous facts so one permanently + // malformed row cannot occupy the oldest page forever. + let rotated_at = now_unix_secs() as i64; + sqlx::query( + "UPDATE referral_rewards SET updated_at = CASE WHEN updated_at >= ? THEN updated_at + 1 ELSE ? END WHERE id = ? AND status = 'applying'", + ) + .bind(rotated_at) + .bind(rotated_at) + .bind(reward_id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(recovery); + } + let status = match recovery { + ReferralApplyingRecovery::Applied => "applied", + ReferralApplyingRecovery::Failed => "failed", + ReferralApplyingRecovery::Unchanged => unreachable!(), + }; + sqlx::query( + r#" +UPDATE referral_rewards +SET status = ?, + wallet_transaction_id = ?, + updated_at = ? +WHERE id = ? AND status = 'applying' +"#, + ) + .bind(status) + .bind(wallet_transaction_id.as_deref()) + .bind(now_unix_secs() as i64) + .bind(reward_id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + tx.commit().await.map_err(DataLayerError::sql)?; + Ok(recovery) + } + }; +} + #[cfg(any(feature = "mysql", feature = "sqlite"))] macro_rules! referral_credit_numeric_method { ($name:ident, $pool_ty:ty, $wallet_sql:expr) => { @@ -2382,7 +4187,27 @@ WHERE id = ? }; let balance = row_f64!(wallet, "balance"); let gift_before = row_f64!(wallet, "gift_balance"); + let total_adjusted_before = row_f64!(wallet, "total_adjusted"); + if !referral_wallet_values_valid(balance, gift_before) + || !total_adjusted_before.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance is invalid".to_string(), + )); + } + let total_before = balance + gift_before; let gift_after = gift_before + target.amount_usd; + let total_after = balance + gift_after; + let total_adjusted_after = total_adjusted_before + target.amount_usd; + if !gift_after.is_finite() + || !total_before.is_finite() + || !total_after.is_finite() + || !total_adjusted_after.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance overflowed".to_string(), + )); + } let tx_id = uuid::Uuid::new_v4().to_string(); let description = note .map(ToOwned::to_owned) @@ -2391,13 +4216,13 @@ WHERE id = ? r#" UPDATE wallets SET gift_balance = ?, - total_adjusted = total_adjusted + ?, + total_adjusted = ?, updated_at = ? WHERE id = ? "#, ) .bind(gift_after) - .bind(target.amount_usd) + .bind(total_adjusted_after) .bind(now_unix_secs() as i64) .bind(&target.wallet_id) .execute(&mut *tx) @@ -2419,8 +4244,8 @@ VALUES (?, ?, 'adjust', 'referral_reward', ?, ?, ?, ?, ?, ?, ?, .bind(&tx_id) .bind(&target.wallet_id) .bind(target.amount_usd) - .bind(balance + gift_before) - .bind(balance + gift_after) + .bind(total_before) + .bind(total_after) .bind(balance) .bind(balance) .bind(gift_before) @@ -2428,7 +4253,7 @@ VALUES (?, ?, 'adjust', 'referral_reward', ?, ?, ?, ?, ?, ?, ?, .bind(&target.id) .bind(operator_id) .bind(&description) - .bind(now_unix_ms() as i64) + .bind(now_unix_secs() as i64) .execute(&mut *tx) .await .map_err(DataLayerError::sql)?; @@ -2458,45 +4283,56 @@ WHERE id = ? } impl ReferralDataState<'_> { + #[cfg(feature = "mysql")] + referral_applying_recovery_numeric_method!( + recover_applying_referral_reward_mysql_numeric_time, + sqlx::MySqlPool + ); + #[cfg(feature = "sqlite")] + referral_applying_recovery_numeric_method!( + recover_applying_referral_reward_sqlite_numeric_time, + sqlx::SqlitePool + ); + #[cfg(feature = "mysql")] referral_credit_numeric_method!( credit_referral_reward_mysql_numeric_time, sqlx::MySqlPool, - "SELECT balance, gift_balance FROM wallets WHERE id = ? FOR UPDATE" + "SELECT wallets.balance, wallets.gift_balance, wallets.total_adjusted + FROM wallets + JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 + WHERE wallets.id = ? AND wallets.status = 'active' + FOR UPDATE" ); #[cfg(feature = "sqlite")] referral_credit_numeric_method!( credit_referral_reward_sqlite_numeric_time, sqlx::SqlitePool, - "SELECT balance, gift_balance FROM wallets WHERE id = ?" + "SELECT wallets.balance, wallets.gift_balance, wallets.total_adjusted + FROM wallets + JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 + WHERE wallets.id = ? AND wallets.status = 'active'" ); #[cfg(any(feature = "mysql", feature = "sqlite"))] async fn apply_referral_reward_reversal_numeric_time( &self, reward: &ReferralRewardRecord, - target_reversal_amount_usd: f64, ) -> Result<(), DataLayerError> { let Some(backends) = self.backends.as_ref() else { return Ok(()); }; #[cfg(feature = "mysql")] if let Some(backend) = backends.mysql() { - return apply_referral_reward_reversal_for_mysql_pool( - &backend.pool_clone(), - reward, - target_reversal_amount_usd, - ) - .await; + return apply_referral_reward_reversal_for_mysql_pool(&backend.pool_clone(), reward) + .await; } #[cfg(feature = "sqlite")] if let Some(backend) = backends.sqlite() { - return apply_referral_reward_reversal_for_sqlite_pool( - &backend.pool_clone(), - reward, - target_reversal_amount_usd, - ) - .await; + return apply_referral_reward_reversal_for_sqlite_pool(&backend.pool_clone(), reward) + .await; } Ok(()) } @@ -2556,11 +4392,10 @@ where #[cfg(any(feature = "mysql", feature = "sqlite"))] macro_rules! referral_reversal_numeric_fn { - ($name:ident, $pool_ty:ty, $wallet_sql:expr) => { + ($name:ident, $pool_ty:ty, $wallet_sql:expr, $order_sql:expr) => { async fn $name( pool: &$pool_ty, reward: &ReferralRewardRecord, - target_reversal_amount_usd: f64, ) -> Result<(), DataLayerError> { let mut tx = pool.begin().await.map_err(DataLayerError::sql)?; sqlx::query("UPDATE referral_rewards SET updated_at = updated_at WHERE id = ?") @@ -2570,7 +4405,8 @@ macro_rules! referral_reversal_numeric_fn { .map_err(DataLayerError::sql)?; let reward_row = sqlx::query( r#" -SELECT reversed_amount_usd, pending_reversal_amount_usd +SELECT status, inviter_user_id, source_order_id, amount_usd, + reversed_amount_usd, pending_reversal_amount_usd FROM referral_rewards WHERE id = ? "#, @@ -2583,41 +4419,144 @@ WHERE id = ? tx.commit().await.map_err(DataLayerError::sql)?; return Ok(()); }; + let reward_status = row_string!(reward_row, "status"); + if !matches!(reward_status.as_str(), "applied" | "reversed") { + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(()); + } + let Some(source_order_id) = row_optional_string!(reward_row, "source_order_id") else { + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(()); + }; + let inviter_user_id = row_string!(reward_row, "inviter_user_id"); + let reward_amount = row_f64!(reward_row, "amount_usd"); let current_reversed = row_f64!(reward_row, "reversed_amount_usd"); let current_pending = row_f64!(reward_row, "pending_reversal_amount_usd"); - let amount_usd = - (target_reversal_amount_usd - current_reversed - current_pending).max(0.0); + + // Match the wallet refund lock order. The payment order is read + // only after its row lock so the target reflects the cumulative + // refund that actually won the race with this transaction. + let wallet = sqlx::query($wallet_sql) + .bind(&inviter_user_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + let order_row = sqlx::query($order_sql) + .bind(&source_order_id) + .fetch_optional(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + let Some(order_row) = order_row else { + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(()); + }; + let refund_context = payment_order_refund_context_from_row(order_row)?; + if !referral_refund_context_valid(&refund_context) { + tx.commit().await.map_err(DataLayerError::sql)?; + return Ok(()); + } + let target_reversal_amount_usd = referral_reversal_target( + reward_amount, + refund_context.amount_usd, + refund_context.refunded_amount_usd, + ); + if !referral_reversal_inputs_valid( + reward_amount, + target_reversal_amount_usd, + current_reversed, + current_pending, + ) { + return Err(DataLayerError::InvalidInput( + "referral reversal state is invalid".to_string(), + )); + } + let amount_usd = referral_reversal_due_bounded( + target_reversal_amount_usd, + reward_amount, + current_reversed, + current_pending, + ); if amount_usd <= 0.0 { tx.commit().await.map_err(DataLayerError::sql)?; return Ok(()); } - let wallet = sqlx::query($wallet_sql) - .bind(&reward.inviter_user_id) - .fetch_optional(&mut *tx) + let Some(wallet) = wallet else { + // Preserve the debt when the inviter wallet is temporarily + // absent/inactive; the periodic reconciliation pass will + // retry after the wallet becomes available. + sqlx::query( + r#" +UPDATE referral_rewards +SET pending_reversal_amount_usd = ?, + status = CASE + WHEN status = 'reversed' THEN 'applied' + ELSE status + END, + updated_at = ? +WHERE id = ? +"#, + ) + .bind(referral_pending_reversal_capped( + reward_amount, + current_reversed, + current_pending, + amount_usd, + )) + .bind(now_unix_secs() as i64) + .bind(&reward.id) + .execute(&mut *tx) .await .map_err(DataLayerError::sql)?; - let Some(wallet) = wallet else { tx.commit().await.map_err(DataLayerError::sql)?; return Ok(()); }; let wallet_id = row_string!(wallet, "id"); let balance = row_f64!(wallet, "balance"); let gift_before = row_f64!(wallet, "gift_balance"); - let actual_reverse = gift_before.min(amount_usd); + let total_adjusted_before = row_f64!(wallet, "total_adjusted"); + if !referral_wallet_values_valid(balance, gift_before) + || !total_adjusted_before.is_finite() + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance is invalid".to_string(), + )); + } + let actual_reverse = gift_before.max(0.0).min(amount_usd); let pending_reverse = (amount_usd - actual_reverse).max(0.0); let gift_after = gift_before - actual_reverse; + let total_before = balance + gift_before; + let total_after = balance + gift_after; + let total_adjusted_after = total_adjusted_before - actual_reverse; + if !actual_reverse.is_finite() + || !pending_reverse.is_finite() + || !gift_after.is_finite() + || !total_before.is_finite() + || !total_after.is_finite() + || !total_adjusted_after.is_finite() + || !referral_reversal_state_valid( + reward_amount, + current_reversed, + current_pending, + actual_reverse, + pending_reverse, + ) + { + return Err(DataLayerError::InvalidInput( + "inviter wallet balance overflowed".to_string(), + )); + } if actual_reverse > 0.0 { sqlx::query( r#" UPDATE wallets SET gift_balance = ?, - total_adjusted = total_adjusted - ?, + total_adjusted = ?, updated_at = ? WHERE id = ? "#, ) .bind(gift_after) - .bind(actual_reverse) + .bind(total_adjusted_after) .bind(now_unix_secs() as i64) .bind(&wallet_id) .execute(&mut *tx) @@ -2639,14 +4578,14 @@ VALUES (?, ?, 'adjust', 'referral_reward_reversal', ?, ?, ?, ?, ?, ?, ?, .bind(uuid::Uuid::new_v4().to_string()) .bind(&wallet_id) .bind(-actual_reverse) - .bind(balance + gift_before) - .bind(balance + gift_after) + .bind(total_before) + .bind(total_after) .bind(balance) .bind(balance) .bind(gift_before) .bind(gift_after) .bind(&reward.id) - .bind(now_unix_ms() as i64) + .bind(now_unix_secs() as i64) .execute(&mut *tx) .await .map_err(DataLayerError::sql)?; @@ -2655,18 +4594,32 @@ VALUES (?, ?, 'adjust', 'referral_reward_reversal', ?, ?, ?, ?, ?, ?, ?, r#" UPDATE referral_rewards SET reversed_amount_usd = reversed_amount_usd + ?, - pending_reversal_amount_usd = pending_reversal_amount_usd + ?, - status = CASE - WHEN reversed_amount_usd + ? >= amount_usd THEN 'reversed' - ELSE status - END, + pending_reversal_amount_usd = ?, updated_at = ? WHERE id = ? "#, ) .bind(actual_reverse) .bind(pending_reverse) - .bind(actual_reverse) + .bind(now_unix_secs() as i64) + .bind(&reward.id) + .execute(&mut *tx) + .await + .map_err(DataLayerError::sql)?; + sqlx::query( + r#" +UPDATE referral_rewards +SET status = CASE + WHEN pending_reversal_amount_usd > 0.00000001 AND status = 'reversed' THEN 'applied' + WHEN pending_reversal_amount_usd <= 0.00000001 + AND reversed_amount_usd >= amount_usd + AND status IN ('applied', 'reversed') THEN 'reversed' + ELSE status + END, + updated_at = ? +WHERE id = ? +"#, + ) .bind(now_unix_secs() as i64) .bind(&reward.id) .execute(&mut *tx) @@ -2682,13 +4635,81 @@ WHERE id = ? referral_reversal_numeric_fn!( apply_referral_reward_reversal_for_mysql_pool, sqlx::MySqlPool, - "SELECT id, balance, gift_balance FROM wallets WHERE user_id = ? FOR UPDATE" + "SELECT wallets.id, wallets.balance, wallets.gift_balance, wallets.total_adjusted + FROM wallets + JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 + WHERE wallets.user_id = ? AND wallets.status = 'active' + FOR UPDATE", + "SELECT po.amount_usd, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd + FROM payment_orders po + WHERE po.id = ? + FOR UPDATE" ); #[cfg(feature = "sqlite")] referral_reversal_numeric_fn!( apply_referral_reward_reversal_for_sqlite_pool, sqlx::SqlitePool, - "SELECT id, balance, gift_balance FROM wallets WHERE user_id = ?" + "SELECT wallets.id, wallets.balance, wallets.gift_balance, wallets.total_adjusted + FROM wallets + JOIN users inviter ON inviter.id = wallets.user_id + AND inviter.is_active = 1 AND inviter.is_deleted = 0 + WHERE wallets.user_id = ? AND wallets.status = 'active'", + "SELECT po.amount_usd, + CASE + WHEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) >= + COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + THEN COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'succeeded' + ), 0.0) + ELSE COALESCE(po.refunded_amount_usd, 0.0) - COALESCE(( + SELECT SUM(rr.amount_usd) + FROM refund_requests rr + WHERE rr.payment_order_id = po.id + AND rr.status = 'processing' + ), 0.0) + END AS refunded_amount_usd + FROM payment_orders po + WHERE po.id = ?" ); #[cfg(test)] @@ -2722,10 +4743,1343 @@ mod tests { let second = referral_reversal_delta(10.0, 100.0, 50.0, 2.0, 0.0); assert!((second - 3.0).abs() < f64::EPSILON); + // A previously deferred reversal remains due until a later pass can + // consume the inviter's replenished gift balance. let repeated = referral_reversal_delta(10.0, 100.0, 50.0, 2.0, 3.0); - assert_eq!(repeated, 0.0); + assert!((repeated - 3.0).abs() < f64::EPSILON); + + let increased_target = referral_reversal_delta(10.0, 100.0, 80.0, 2.0, 3.0); + assert!((increased_target - 6.0).abs() < f64::EPSILON); let full = referral_reversal_delta(10.0, 100.0, 125.0, 5.0, 0.0); assert!((full - 5.0).abs() < f64::EPSILON); } + + #[test] + fn referral_reversal_target_rejects_non_finite_amounts() { + assert_eq!(referral_reversal_target(f64::NAN, 100.0, 10.0), 0.0); + assert_eq!(referral_reversal_target(10.0, f64::INFINITY, 10.0), 0.0); + assert_eq!( + referral_reversal_target(10.0, 100.0, f64::NEG_INFINITY), + 0.0 + ); + } + + #[test] + fn referral_reversal_due_never_exceeds_current_refund_or_reward() { + // A legacy additive pending value may be larger than the current + // target. It must not authorize an over-reversal beyond the reward. + assert_eq!(referral_reversal_due_bounded(5.0, 10.0, 0.0, 15.0), 10.0); + assert_eq!(referral_reversal_due_bounded(5.0, 10.0, 10.0, 3.0), 0.0); + // A malformed target above the reward is still capped by the reward + // remainder. + assert_eq!(referral_reversal_due_bounded(20.0, 10.0, 2.0, 3.0), 8.0); + } + + #[test] + fn referral_percent_rate_must_be_finite_and_at_most_one_hundred() { + assert!(referral_percent_rate_valid(0.01)); + assert!(referral_percent_rate_valid(100.0)); + for value in [0.0, -1.0, 100.000001, f64::NAN, f64::INFINITY] { + assert!( + !referral_percent_rate_valid(value), + "{value:?} must be rejected" + ); + } + } + + #[test] + fn referral_payment_method_exclusion_is_case_and_whitespace_insensitive() { + for method in [ + "manual", + "MANUAL", + " Admin_Manual ", + "REDEEM_CODE", + " Gift\t", + ] { + assert!( + referral_payment_method_excluded(method), + "{method:?} must be excluded" + ); + } + for method in ["stripe", "paypal", "manual_review"] { + assert!( + !referral_payment_method_excluded(method), + "{method:?} must remain eligible" + ); + } + } + + #[test] + fn referral_wallet_values_allow_overdraft_but_reject_invalid_gifts() { + assert!(referral_wallet_values_valid(-25.0, 3.0)); + assert!(!referral_wallet_values_valid(f64::NAN, 3.0)); + assert!(!referral_wallet_values_valid(1.0, f64::INFINITY)); + assert!(!referral_wallet_values_valid(1.0, -0.01)); + } + + #[test] + fn referral_refund_context_rejects_amounts_outside_order_total() { + assert!(referral_refund_context_valid( + &ReferralPaymentOrderRefundContext { + amount_usd: 100.0, + refunded_amount_usd: 25.0, + } + )); + assert!(!referral_refund_context_valid( + &ReferralPaymentOrderRefundContext { + amount_usd: 100.0, + refunded_amount_usd: 100.00001, + } + )); + assert!(!referral_refund_context_valid( + &ReferralPaymentOrderRefundContext { + amount_usd: 100.0, + refunded_amount_usd: -1.0, + } + )); + assert!(!referral_refund_context_valid( + &ReferralPaymentOrderRefundContext { + amount_usd: f64::NAN, + refunded_amount_usd: 1.0, + } + )); + } + + #[test] + fn referral_credit_fact_requires_a_consistent_gift_only_delta() { + assert!(referral_credit_transaction_fact_valid( + 3.0, 3.0, 2.0, 5.0, -1.0, -1.0, 3.0, 6.0, + )); + // Matching amount/link metadata alone must not be trusted. + assert!(!referral_credit_transaction_fact_valid( + 3.0, 3.0, 2.0, 5.0, -1.0, -1.0, 3.0, 3.0, + )); + assert!(!referral_credit_transaction_fact_valid( + 3.0, 3.0, 2.0, 5.0, -1.0, 0.0, 3.0, 6.0, + )); + assert!(!referral_credit_transaction_fact_valid( + 3.0, 3.0, 2.0, 5.0, -1.0, -1.0, -1.0, 2.0, + )); + assert!(!referral_credit_transaction_fact_valid( + 3.0, + f64::NAN, + 2.0, + 5.0, + -1.0, + -1.0, + 3.0, + 6.0, + )); + } + + #[test] + fn referral_reversal_state_rejects_invalid_or_overflowing_totals() { + assert!(referral_reversal_state_valid(10.0, 2.0, 3.0, 1.0, 4.0)); + assert!(!referral_reversal_state_valid(10.0, -1.0, 0.0, 1.0, 0.0)); + assert!(!referral_reversal_state_valid(10.0, 2.0, 3.0, 9.0, 0.0)); + assert!(!referral_reversal_state_valid( + f64::MAX, + f64::MAX, + 0.0, + f64::MAX, + 0.0, + )); + } + + #[test] + fn referral_reversal_inputs_reject_malformed_durable_counters() { + assert!(referral_reversal_inputs_valid(10.0, 5.0, 2.0, 3.0)); + assert!(!referral_reversal_inputs_valid(10.0, 5.0, -1.0, 0.0)); + assert!(!referral_reversal_inputs_valid(10.0, 5.0, 2.0, -1.0)); + assert!(!referral_reversal_inputs_valid(10.0, -0.1, 2.0, 3.0)); + assert!(!referral_reversal_inputs_valid(10.0, 5.0, 8.0, 3.0)); + assert!(!referral_reversal_inputs_valid(10.0, 11.0, 0.0, 0.0)); + assert!(!referral_reversal_inputs_valid(f64::NAN, 1.0, 0.0, 0.0,)); + } + + #[test] + fn referral_pending_reversal_is_capped_at_remaining_reward() { + assert_eq!(referral_pending_reversal_capped(10.0, 2.0, 15.0, 3.0), 8.0); + assert_eq!(referral_pending_reversal_capped(10.0, 2.0, 1.0, 3.0), 3.0); + } + + #[test] + fn referral_stats_amount_saturates_overflow_without_hiding_it() { + assert_eq!(referral_stats_amount(f64::INFINITY), f64::MAX); + assert_eq!(referral_stats_amount(f64::NEG_INFINITY), 0.0); + assert_eq!(referral_stats_amount(f64::NAN), 0.0); + assert_eq!(referral_stats_amount(-1.0), 0.0); + assert_eq!(referral_stats_amount(12.5), 12.5); + } + + #[test] + fn referral_like_pattern_escapes_wildcards_and_uses_empty_filter_sentinel() { + assert_eq!(referral_like_pattern(None), ""); + assert_eq!(referral_like_pattern(Some(" ")), ""); + assert_eq!(referral_like_pattern(Some(" A_%! ")), "%a!_!%!!%"); + } + + #[test] + fn referral_page_bounds_clamp_limit_and_saturate_offset() { + assert_eq!(referral_page_bounds(0, 0), (1, 0)); + assert_eq!(referral_page_bounds(999, 4), (200, 4)); + assert_eq!(referral_page_bounds(20, usize::MAX), (20, i64::MAX)); + } + + #[cfg(feature = "sqlite")] + #[tokio::test] + async fn referral_refund_context_combines_legacy_and_settled_refunds() { + let config = crate::DataLayerConfig::from_database(crate::SqlDatabaseConfig { + driver: crate::DatabaseDriver::Sqlite, + url: "sqlite::memory:".to_string(), + pool: crate::SqlPoolConfig { + max_connections: 1, + ..crate::SqlPoolConfig::default() + }, + }); + let backends = + crate::DataBackends::from_config(config).expect("sqlite data backends should build"); + let pool = backends + .sqlite() + .expect("sqlite backend should exist") + .pool(); + crate::lifecycle::migrate::run_sqlite_migrations(pool) + .await + .expect("sqlite migrations should run"); + + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, 'user', ?, ?)", + ) + .bind("refund-context-user") + .bind("refund-context@example.test") + .bind("refund-context-user") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("refund context user should insert"); + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?)", + ) + .bind("refund-context-wallet") + .bind("refund-context-user") + .bind(0.0_f64) + .bind(0.0_f64) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("refund context wallet should insert"); + sqlx::query( + "INSERT INTO payment_orders (id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd, refundable_amount_usd, payment_method, status, created_at, credited_at, order_kind) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("refund-context-order") + .bind("refund-context-order-no") + .bind("refund-context-wallet") + .bind("refund-context-user") + .bind(10.0_f64) + .bind(2.0_f64) + .bind(8.0_f64) + .bind("stripe") + .bind("credited") + .bind(1_i64) + .bind(1_i64) + .bind("wallet_recharge") + .execute(pool) + .await + .expect("refund context order should insert"); + + let state = ReferralDataState::new(Some(&backends)); + let context = state + .find_referral_payment_order_refund_context("refund-context-order") + .await + .expect("legacy refund context should query") + .expect("refund context order should exist"); + assert!((context.refunded_amount_usd - 2.0).abs() < f64::EPSILON); + + // The processing request has already increased the legacy order + // counter, but it must not authorize a referral reversal yet. + sqlx::query( + "INSERT INTO refund_requests (id, refund_no, wallet_id, user_id, payment_order_id, source_type, refund_mode, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("refund-context-processing") + .bind("refund-context-processing-no") + .bind("refund-context-wallet") + .bind("refund-context-user") + .bind("refund-context-order") + .bind("wallet") + .bind("offline_payout") + .bind(3.0_f64) + .bind("processing") + .bind(2_i64) + .bind(2_i64) + .execute(pool) + .await + .expect("processing refund should insert"); + sqlx::query( + "UPDATE payment_orders SET refunded_amount_usd = ?, refundable_amount_usd = ? WHERE id = ?", + ) + .bind(5.0_f64) + .bind(5.0_f64) + .bind("refund-context-order") + .execute(pool) + .await + .expect("processing order counter should update"); + let context = state + .find_referral_payment_order_refund_context("refund-context-order") + .await + .expect("processing refund context should query") + .expect("processing refund context order should exist"); + assert!((context.refunded_amount_usd - 2.0).abs() < f64::EPSILON); + + // A settled request is additive to the historical counter. While the + // newer request is still processing, the effective settled amount must + // retain the legacy two dollars rather than dropping to zero. + sqlx::query( + "INSERT INTO refund_requests (id, refund_no, wallet_id, user_id, payment_order_id, source_type, refund_mode, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("refund-context-succeeded") + .bind("refund-context-succeeded-no") + .bind("refund-context-wallet") + .bind("refund-context-user") + .bind("refund-context-order") + .bind("wallet") + .bind("offline_payout") + .bind(2.0_f64) + .bind("succeeded") + .bind(3_i64) + .bind(3_i64) + .execute(pool) + .await + .expect("succeeded refund should insert"); + let context = state + .find_referral_payment_order_refund_context("refund-context-order") + .await + .expect("mixed refund context should query") + .expect("mixed refund context order should exist"); + assert!((context.refunded_amount_usd - 2.0).abs() < f64::EPSILON); + + sqlx::query( + "UPDATE refund_requests SET status = 'succeeded', processed_at = ? WHERE id = ?", + ) + .bind(4_i64) + .bind("refund-context-processing") + .execute(pool) + .await + .expect("processing refund should settle"); + let context = state + .find_referral_payment_order_refund_context("refund-context-order") + .await + .expect("settled refund context should query") + .expect("settled refund context order should exist"); + assert!((context.refunded_amount_usd - 5.0).abs() < f64::EPSILON); + + // Even if an imported order counter is stale, a durable succeeded + // request must not be erased by the aggregate fallback. + sqlx::query( + "UPDATE payment_orders SET refunded_amount_usd = ?, refundable_amount_usd = ? WHERE id = ?", + ) + .bind(0.0_f64) + .bind(10.0_f64) + .bind("refund-context-order") + .execute(pool) + .await + .expect("stale order counter should update"); + let context = state + .find_referral_payment_order_refund_context("refund-context-order") + .await + .expect("stale counter refund context should query") + .expect("stale counter refund context order should exist"); + assert!((context.refunded_amount_usd - 5.0).abs() < f64::EPSILON); + } + + #[cfg(feature = "sqlite")] + #[tokio::test] + async fn admin_referral_lists_are_not_truncated_at_fetch_limit() { + let config = crate::DataLayerConfig::from_database(crate::SqlDatabaseConfig { + driver: crate::DatabaseDriver::Sqlite, + url: "sqlite::memory:".to_string(), + pool: crate::SqlPoolConfig { + max_connections: 1, + ..crate::SqlPoolConfig::default() + }, + }); + let backends = + crate::DataBackends::from_config(config).expect("sqlite data backends should build"); + let pool = backends + .sqlite() + .expect("sqlite backend should exist") + .pool(); + crate::lifecycle::migrate::run_sqlite_migrations(pool) + .await + .expect("sqlite migrations should run"); + + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, 'user', ?, ?)", + ) + .bind("large-list-inviter") + .bind("large-list-inviter@example.test") + .bind("large-list-inviter") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("inviter should insert"); + + let mut tx = pool.begin().await.expect("bulk transaction should begin"); + for index in 0..=REFERRAL_FETCH_LIMIT { + let user_id = format!("large-list-invitee-{index}"); + let email = format!("large-list-invitee-{index}@example.test"); + let username = format!("large-list-invitee-{index}"); + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, 'user', ?, ?)", + ) + .bind(&user_id) + .bind(&email) + .bind(&username) + .bind(index as i64 + 2) + .bind(index as i64 + 2) + .execute(&mut *tx) + .await + .expect("invitee should insert"); + sqlx::query( + "INSERT INTO user_referrals (id, inviter_user_id, invitee_user_id, invite_code_snapshot, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind(format!("large-list-referral-{index}")) + .bind("large-list-inviter") + .bind(&user_id) + .bind("AE-LARGE-LIST") + .bind(index as i64 + 2) + .bind(index as i64 + 2) + .execute(&mut *tx) + .await + .expect("referral relationship should insert"); + sqlx::query( + "INSERT INTO referral_rewards (id, referral_id, inviter_user_id, invitee_user_id, reward_type, trigger_point, idempotency_key, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'percent', 'paid_order', ?, ?, 'pending', ?, ?)", + ) + .bind(format!("large-list-reward-{index}")) + .bind(format!("large-list-referral-{index}")) + .bind("large-list-inviter") + .bind(&user_id) + .bind(format!("large-list-reward-key-{index}")) + .bind(1.0_f64) + .bind(index as i64 + 2) + .bind(index as i64 + 2) + .execute(&mut *tx) + .await + .expect("referral reward should insert"); + } + tx.commit().await.expect("bulk transaction should commit"); + + let state = ReferralDataState::new(Some(&backends)); + let (items, total, stats) = state + .list_admin_referral_relationships(ReferralRelationshipListQuery { + inviter: Some("large-list-inviter".to_string()), + limit: 1, + offset: REFERRAL_FETCH_LIMIT, + ..ReferralRelationshipListQuery::default() + }) + .await + .expect("large relationship list should succeed") + .expect("sqlite referral backend should be available"); + assert_eq!(total, (REFERRAL_FETCH_LIMIT + 1) as u64); + assert_eq!(items.len(), 1); + assert_eq!(stats.total_invites, (REFERRAL_FETCH_LIMIT + 1) as u64); + + let (reward_items, reward_total, reward_stats) = state + .list_admin_referral_rewards(ReferralRewardListQuery { + order_id: None, + reward_type: Some("percent".to_string()), + status: Some("pending".to_string()), + limit: 1, + offset: REFERRAL_FETCH_LIMIT, + }) + .await + .expect("large reward list should succeed") + .expect("sqlite referral backend should be available"); + assert_eq!(reward_total, (REFERRAL_FETCH_LIMIT + 1) as u64); + assert_eq!(reward_items.len(), 1); + assert_eq!( + reward_stats.total_invites, + (REFERRAL_FETCH_LIMIT + 1) as u64 + ); + assert_eq!( + reward_stats.pending_reward_usd, + (REFERRAL_FETCH_LIMIT + 1) as f64 + ); + } + + #[cfg(feature = "sqlite")] + #[tokio::test] + async fn reconciliation_recovers_applying_rewards_from_wallet_transaction_facts() { + let config = crate::DataLayerConfig::from_database(crate::SqlDatabaseConfig { + driver: crate::DatabaseDriver::Sqlite, + url: "sqlite::memory:".to_string(), + pool: crate::SqlPoolConfig { + max_connections: 1, + ..crate::SqlPoolConfig::default() + }, + }); + let backends = + crate::DataBackends::from_config(config).expect("sqlite data backends should build"); + let pool = backends + .sqlite() + .expect("sqlite backend should exist") + .pool(); + crate::lifecycle::migrate::run_sqlite_migrations(pool) + .await + .expect("sqlite migrations should run"); + + for (id, email, username) in [ + ( + "applying-inviter", + "applying-inviter@example.test", + "applying-inviter", + ), + ( + "applying-invitee", + "applying-invitee@example.test", + "applying-invitee", + ), + ] { + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, 'user', ?, ?)", + ) + .bind(id) + .bind(email) + .bind(username) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("referral user should insert"); + } + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, total_adjusted, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?, ?)", + ) + .bind("applying-wallet") + .bind("applying-inviter") + .bind(0.0_f64) + // This is the already-committed credit represented by existing-tx. + .bind(2.0_f64) + .bind(2.0_f64) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("inviter wallet should insert"); + sqlx::query( + "INSERT INTO user_referrals (id, inviter_user_id, invitee_user_id, invite_code_snapshot, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("applying-referral") + .bind("applying-inviter") + .bind("applying-invitee") + .bind("AE-APPLYING") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("referral relationship should insert"); + for (id, key, amount) in [ + ("applying-with-tx", "applying-key-with-tx", 2.0_f64), + ("applying-without-tx", "applying-key-without-tx", 3.0_f64), + ( + "applying-with-invalid-tx", + "applying-key-with-invalid-tx", + 4.0_f64, + ), + ] { + sqlx::query( + "INSERT INTO referral_rewards (id, referral_id, inviter_user_id, invitee_user_id, reward_type, trigger_point, idempotency_key, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'percent', 'paid_order', ?, ?, 'applying', ?, ?)", + ) + .bind(id) + .bind("applying-referral") + .bind("applying-inviter") + .bind("applying-invitee") + .bind(key) + .bind(amount) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("applying reward should insert"); + } + sqlx::query( + r#" +INSERT INTO wallet_transactions ( + id, wallet_id, category, reason_code, amount, + balance_before, balance_after, + recharge_balance_before, recharge_balance_after, + gift_balance_before, gift_balance_after, + link_type, link_id, description, created_at +) +VALUES (?, ?, 'adjust', 'referral_reward', ?, ?, ?, ?, ?, ?, ?, + 'referral_reward', ?, 'existing referral credit', ?) +"#, + ) + .bind("existing-referral-tx") + .bind("applying-wallet") + .bind(2.0_f64) + .bind(0.0_f64) + .bind(2.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(2.0_f64) + .bind("applying-with-tx") + .bind(1_i64) + .execute(pool) + .await + .expect("existing referral transaction should insert"); + // A row with the right link id and amount is still not proof of a + // credit when its balance snapshot is inconsistent. Recovery must + // validate the complete transaction shape before moving an `applying` + // reward to `applied`. Because this is still evidence of an ambiguous + // historical write, it must not be downgraded to `failed` (which would + // make the next pass credit the wallet a second time). + sqlx::query( + r#" +INSERT INTO wallet_transactions ( + id, wallet_id, category, reason_code, amount, + balance_before, balance_after, + recharge_balance_before, recharge_balance_after, + gift_balance_before, gift_balance_after, + link_type, link_id, description, created_at +) +VALUES (?, ?, 'adjust', 'referral_reward', 4, 2, 6, 0, 0, 2, 2, + 'referral_reward', ?, 'non-credit fact', ?) +"#, + ) + .bind("non-credit-referral-tx") + .bind("applying-wallet") + .bind("applying-with-invalid-tx") + .bind(2_i64) + .execute(pool) + .await + .expect("non-credit transaction should insert"); + let state = ReferralDataState::new(Some(&backends)); + let first = state + .reconcile_referral_rewards_once(None) + .await + .expect("applying recovery should succeed"); + assert_eq!(first.reward_attempted, 3); + assert_eq!(first.reward_applied, 1); + assert_eq!(first.deferred, 2); + let dashboard_after_recovery = state + .referral_dashboard("applying-inviter") + .await + .expect("applying dashboard should aggregate") + .expect("applying inviter dashboard should exist"); + assert!((dashboard_after_recovery.paid_reward_usd - 2.0).abs() < f64::EPSILON); + assert!((dashboard_after_recovery.pending_reward_usd - 7.0).abs() < f64::EPSILON); + + let (gift_after_recovery, transaction_count): (f64, i64) = sqlx::query_as( + "SELECT gift_balance, (SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?) FROM wallets WHERE id = ?", + ) + .bind("applying-wallet") + .bind("applying-wallet") + .fetch_one(pool) + .await + .expect("wallet should remain readable"); + assert!((gift_after_recovery - 2.0).abs() < f64::EPSILON); + assert_eq!(transaction_count, 2); + + let (with_tx_status, recovered_tx_id): (String, Option) = sqlx::query_as( + "SELECT status, wallet_transaction_id FROM referral_rewards WHERE id = ?", + ) + .bind("applying-with-tx") + .fetch_one(pool) + .await + .expect("recovered reward should be readable"); + assert_eq!(with_tx_status, "applied"); + assert_eq!(recovered_tx_id.as_deref(), Some("existing-referral-tx")); + let (without_tx_status, missing_tx_id): (String, Option) = sqlx::query_as( + "SELECT status, wallet_transaction_id FROM referral_rewards WHERE id = ?", + ) + .bind("applying-without-tx") + .fetch_one(pool) + .await + .expect("failed reward should be readable"); + assert_eq!(without_tx_status, "failed"); + assert!(missing_tx_id.is_none()); + let (invalid_tx_status, invalid_tx_id): (String, Option) = sqlx::query_as( + "SELECT status, wallet_transaction_id FROM referral_rewards WHERE id = ?", + ) + .bind("applying-with-invalid-tx") + .fetch_one(pool) + .await + .expect("ambiguous reward should be readable"); + assert_eq!(invalid_tx_status, "applying"); + assert!(invalid_tx_id.is_none()); + + // Only the evidence-free reward may retry through the normal credit + // transaction. The ambiguous reward is inspected again but remains + // applying and never credits a second time. + let second = state + .reconcile_referral_rewards_once(None) + .await + .expect("failed reward retry should succeed"); + assert_eq!(second.reward_attempted, 2); + assert_eq!(second.reward_applied, 1); + assert_eq!(second.deferred, 1); + let (gift_after_retry, transaction_count): (f64, i64) = sqlx::query_as( + "SELECT gift_balance, (SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?) FROM wallets WHERE id = ?", + ) + .bind("applying-wallet") + .bind("applying-wallet") + .fetch_one(pool) + .await + .expect("retried wallet should be readable"); + assert!((gift_after_retry - 5.0).abs() < f64::EPSILON); + assert_eq!(transaction_count, 3); + + let third = state + .reconcile_referral_rewards_once(None) + .await + .expect("settled rewards should be idempotent"); + assert_eq!(third.reward_attempted, 1); + assert_eq!(third.reward_applied, 0); + assert_eq!(third.deferred, 1); + let final_gift: f64 = sqlx::query_scalar("SELECT gift_balance FROM wallets WHERE id = ?") + .bind("applying-wallet") + .fetch_one(pool) + .await + .expect("final wallet balance should be readable"); + assert!((final_gift - 5.0).abs() < f64::EPSILON); + let invalid_final_status: String = + sqlx::query_scalar("SELECT status FROM referral_rewards WHERE id = ?") + .bind("applying-with-invalid-tx") + .fetch_one(pool) + .await + .expect("ambiguous reward should remain readable"); + assert_eq!(invalid_final_status, "applying"); + } + + #[cfg(feature = "sqlite")] + #[tokio::test] + async fn reconciliation_rotates_ambiguous_applying_rows_without_starving_valid_facts() { + let config = crate::DataLayerConfig::from_database(crate::SqlDatabaseConfig { + driver: crate::DatabaseDriver::Sqlite, + url: "sqlite::memory:".to_string(), + pool: crate::SqlPoolConfig { + max_connections: 1, + ..crate::SqlPoolConfig::default() + }, + }); + let backends = + crate::DataBackends::from_config(config).expect("sqlite data backends should build"); + let pool = backends + .sqlite() + .expect("sqlite backend should exist") + .pool(); + crate::lifecycle::migrate::run_sqlite_migrations(pool) + .await + .expect("sqlite migrations should run"); + + for (id, email, username) in [ + ( + "rotation-inviter", + "rotation-inviter@example.test", + "rotation-inviter", + ), + ( + "rotation-invitee", + "rotation-invitee@example.test", + "rotation-invitee", + ), + ] { + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, 'user', ?, ?)", + ) + .bind(id) + .bind(email) + .bind(username) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("rotation user should insert"); + } + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, total_adjusted, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?, ?)", + ) + .bind("rotation-wallet") + .bind("rotation-inviter") + .bind(0.0_f64) + .bind(1.0_f64) + .bind(1.0_f64) + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("rotation wallet should insert"); + sqlx::query( + "INSERT INTO user_referrals (id, inviter_user_id, invitee_user_id, invite_code_snapshot, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("rotation-referral") + .bind("rotation-inviter") + .bind("rotation-invitee") + .bind("AE-ROTATION") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("rotation referral should insert"); + + // Fill the bounded page with malformed applying rows. Their durable + // amount is invalid, so recovery must leave them applying and rotate + // their updated_at instead of allowing them to monopolise the queue. + // Use a future timestamp to ensure rotation never moves a corrupted + // imported value backwards into the queue's oldest position. + let imported_updated_at = 4_102_444_800_i64; + for index in 0..REFERRAL_RECONCILIATION_LIMIT { + sqlx::query( + "INSERT INTO referral_rewards (id, referral_id, inviter_user_id, invitee_user_id, reward_type, trigger_point, idempotency_key, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'percent', 'paid_order', ?, ?, 'applying', ?, ?)", + ) + .bind(format!("rotation-noise-{index:03}")) + .bind("rotation-referral") + .bind("rotation-inviter") + .bind("rotation-invitee") + .bind(format!("rotation-noise-key-{index:03}")) + .bind(0.0_f64) + .bind(1_i64) + .bind(imported_updated_at) + .execute(pool) + .await + .expect("malformed applying row should insert"); + } + sqlx::query( + "INSERT INTO referral_rewards (id, referral_id, inviter_user_id, invitee_user_id, reward_type, trigger_point, idempotency_key, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'percent', 'paid_order', ?, ?, 'applying', ?, ?)", + ) + .bind("rotation-valid") + .bind("rotation-referral") + .bind("rotation-inviter") + .bind("rotation-invitee") + .bind("rotation-valid-key") + .bind(1.0_f64) + .bind(2_i64) + .bind(imported_updated_at) + .execute(pool) + .await + .expect("valid applying row should insert"); + sqlx::query( + r#" +INSERT INTO wallet_transactions ( + id, wallet_id, category, reason_code, amount, + balance_before, balance_after, + recharge_balance_before, recharge_balance_after, + gift_balance_before, gift_balance_after, + link_type, link_id, description, created_at +) +VALUES (?, ?, 'adjust', 'referral_reward', ?, ?, ?, ?, ?, ?, ?, + 'referral_reward', ?, 'rotation credit', ?) +"#, + ) + .bind("rotation-valid-tx") + .bind("rotation-wallet") + .bind(1.0_f64) + .bind(0.0_f64) + .bind(1.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(1.0_f64) + .bind("rotation-valid") + .bind(2_i64) + .execute(pool) + .await + .expect("valid wallet fact should insert"); + + let state = ReferralDataState::new(Some(&backends)); + let first = state + .reconcile_referral_rewards_once(None) + .await + .expect("first rotation pass should succeed"); + assert_eq!(first.reward_attempted, REFERRAL_RECONCILIATION_LIMIT as u64); + assert_eq!(first.reward_applied, 0); + assert_eq!(first.deferred, REFERRAL_RECONCILIATION_LIMIT as u64); + let first_valid_status: String = + sqlx::query_scalar("SELECT status FROM referral_rewards WHERE id = ?") + .bind("rotation-valid") + .fetch_one(pool) + .await + .expect("valid row should remain readable"); + assert_eq!(first_valid_status, "applying"); + + // Two independently committed, internally consistent facts for the + // same reward are ambiguous: the wallet may already have been + // credited twice. Recovery must refuse to hide that duplicate. + sqlx::query( + r#" +INSERT INTO wallet_transactions ( + id, wallet_id, category, reason_code, amount, + balance_before, balance_after, + recharge_balance_before, recharge_balance_after, + gift_balance_before, gift_balance_after, + link_type, link_id, description, created_at +) +VALUES (?, ?, 'adjust', 'referral_reward', ?, ?, ?, ?, ?, ?, ?, + 'referral_reward', ?, 'duplicate rotation credit', ?) +"#, + ) + .bind("rotation-valid-tx-duplicate") + .bind("rotation-wallet") + .bind(1.0_f64) + .bind(0.0_f64) + .bind(1.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(0.0_f64) + .bind(1.0_f64) + .bind("rotation-valid") + .bind(3_i64) + .execute(pool) + .await + .expect("duplicate wallet fact should insert"); + + let second = state + .reconcile_referral_rewards_once(None) + .await + .expect("second rotation pass should succeed"); + assert_eq!(second.reward_applied, 0); + let duplicate_status: String = + sqlx::query_scalar("SELECT status FROM referral_rewards WHERE id = ?") + .bind("rotation-valid") + .fetch_one(pool) + .await + .expect("duplicate reward should remain readable"); + assert_eq!(duplicate_status, "applying"); + + sqlx::query("DELETE FROM wallet_transactions WHERE id = ?") + .bind("rotation-valid-tx-duplicate") + .execute(pool) + .await + .expect("duplicate wallet fact should be removed for recovery test"); + let third = state + .reconcile_referral_rewards_once(None) + .await + .expect("unambiguous rotation pass should succeed"); + assert_eq!(third.reward_applied, 1); + let (valid_status, valid_tx_id, gift_balance): (String, Option, f64) = + sqlx::query_as( + "SELECT (SELECT status FROM referral_rewards WHERE id = ?), (SELECT wallet_transaction_id FROM referral_rewards WHERE id = ?), (SELECT gift_balance FROM wallets WHERE id = ?)", + ) + .bind("rotation-valid") + .bind("rotation-valid") + .bind("rotation-wallet") + .fetch_one(pool) + .await + .expect("rotated valid fact should be readable"); + assert_eq!(valid_status, "applied"); + assert_eq!(valid_tx_id.as_deref(), Some("rotation-valid-tx")); + assert!((gift_balance - 1.0).abs() < f64::EPSILON); + } + + #[cfg(feature = "sqlite")] + #[tokio::test] + async fn reconciliation_does_not_infer_missing_historical_rewards() { + let config = crate::DataLayerConfig::from_database(crate::SqlDatabaseConfig { + driver: crate::DatabaseDriver::Sqlite, + url: "sqlite::memory:".to_string(), + pool: crate::SqlPoolConfig { + max_connections: 1, + ..crate::SqlPoolConfig::default() + }, + }); + let backends = + crate::DataBackends::from_config(config).expect("sqlite data backends should build"); + let pool = backends + .sqlite() + .expect("sqlite backend should exist") + .pool(); + crate::lifecycle::migrate::run_sqlite_migrations(pool) + .await + .expect("sqlite migrations should run"); + + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("referral-inviter") + .bind("inviter@example.test") + .bind("inviter") + .bind("user") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("inviter should insert"); + sqlx::query( + "INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("referral-invitee") + .bind("invitee@example.test") + .bind("invitee") + .bind("user") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("invitee should insert"); + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + ) + .bind("wallet-inviter") + .bind("referral-inviter") + // Recharge balance may be negative when the account has overdraft; + // referral credit must still be able to add to its gift balance. + .bind(-5.0_f64) + .bind(0.0_f64) + .bind("active") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("inviter wallet should insert"); + sqlx::query( + "INSERT INTO wallets (id, user_id, balance, gift_balance, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + ) + .bind("wallet-invitee") + .bind("referral-invitee") + .bind(0.0_f64) + .bind(0.0_f64) + .bind("active") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("invitee wallet should insert"); + sqlx::query( + "INSERT INTO user_referrals (id, inviter_user_id, invitee_user_id, invite_code_snapshot, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", + ) + .bind("referral-link") + .bind("referral-inviter") + .bind("referral-invitee") + .bind("AE-TEST") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("referral link should insert"); + sqlx::query( + "INSERT INTO payment_orders (id, order_no, wallet_id, user_id, amount_usd, refundable_amount_usd, payment_method, status, created_at, credited_at, order_kind) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("paid-order-repair") + .bind("paid-order-repair-no") + .bind("wallet-invitee") + .bind("referral-invitee") + .bind(10.0_f64) + .bind(10.0_f64) + .bind("stripe") + .bind("credited") + .bind(1_i64) + .bind(2_i64) + .bind("wallet_recharge") + .execute(pool) + .await + .expect("credited payment order should insert"); + + let state = ReferralDataState::new(Some(&backends)); + let reward_config = ReferralRewardConfig { + percent_enabled: true, + percent_rate: 10.0, + headcount_enabled: false, + headcount_amount_usd: 0.0, + headcount_trigger: "registration".to_string(), + }; + // A credited order with only a referral relationship is not durable + // evidence that the referral feature was enabled for that order. The + // periodic worker must not apply the current configuration to it. + let historical_pass = state + .reconcile_referral_rewards_once(Some(reward_config.clone())) + .await + .expect("historical reconciliation should succeed"); + assert_eq!(historical_pass.order_attempted, 0); + assert_eq!(historical_pass.order_repaired, 0); + let (recharge_balance, gift_balance): (f64, f64) = + sqlx::query_as("SELECT balance, gift_balance FROM wallets WHERE id = ?") + .bind("wallet-inviter") + .fetch_one(pool) + .await + .expect("credited inviter wallet should be readable"); + assert!((recharge_balance + 5.0).abs() < f64::EPSILON); + assert!(gift_balance.abs() < f64::EPSILON); + + // The normal callback path still applies a reward with the + // configuration that was active at payment time. Reconciliation is + // intentionally limited to rows created by that durable path. + let applied = state + .apply_paid_order_referral_rewards("paid-order-repair", reward_config.clone()) + .await + .expect("normal paid-order application should succeed"); + assert_eq!(applied.len(), 1); + let (recharge_balance, gift_balance): (f64, f64) = + sqlx::query_as("SELECT balance, gift_balance FROM wallets WHERE id = ?") + .bind("wallet-inviter") + .fetch_one(pool) + .await + .expect("applied inviter wallet should be readable"); + assert!((recharge_balance + 5.0).abs() < f64::EPSILON); + assert!((gift_balance - 1.0).abs() < f64::EPSILON); + + let dashboard = state + .referral_dashboard("referral-inviter") + .await + .expect("referral dashboard should use the aggregate path") + .expect("inviter dashboard should be available"); + assert_eq!(dashboard.total_invites, 1); + assert_eq!(dashboard.effective_invites, 1); + assert!((dashboard.paid_reward_usd - 1.0).abs() < f64::EPSILON); + + // Headline admin metrics are global and must not become empty or + // filter-scoped just because one of the list queries is narrowed. + let (_, relationship_total, relationship_stats) = state + .list_admin_referral_relationships(ReferralRelationshipListQuery { + inviter: Some("does-not-match".to_string()), + limit: 100, + offset: 0, + ..ReferralRelationshipListQuery::default() + }) + .await + .expect("filtered relationship list should succeed") + .expect("sqlite referral backend should be available"); + assert_eq!(relationship_total, 0); + assert_eq!(relationship_stats.total_invites, 1); + assert_eq!(relationship_stats.effective_invites, 1); + assert!((relationship_stats.paid_reward_usd - 1.0).abs() < f64::EPSILON); + + let (_, reward_total, reward_stats) = state + .list_admin_referral_rewards(ReferralRewardListQuery { + status: Some("voided".to_string()), + limit: 100, + offset: 0, + ..ReferralRewardListQuery::default() + }) + .await + .expect("filtered reward list should succeed") + .expect("sqlite referral backend should be available"); + assert_eq!(reward_total, 0); + assert_eq!(reward_stats.total_invites, 1); + assert_eq!(reward_stats.effective_invites, 1); + assert!((reward_stats.paid_reward_usd - 1.0).abs() < f64::EPSILON); + + let second = state + .reconcile_referral_rewards_once(Some(reward_config.clone())) + .await + .expect("second reconciliation should succeed"); + assert_eq!(second.order_attempted, 0); + assert_eq!(second.order_repaired, 0); + + // A deleted inviter must never receive a delayed reward, even if an + // old pending row and an otherwise active wallet remain in storage. + sqlx::query( + "INSERT INTO referral_rewards (id, referral_id, inviter_user_id, invitee_user_id, reward_type, trigger_point, idempotency_key, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("deleted-inviter-reward") + .bind("referral-link") + .bind("referral-inviter") + .bind("referral-invitee") + .bind("percent") + .bind("paid_order") + .bind("referral:deleted-inviter-reward") + .bind(2.0_f64) + .bind("pending") + .bind(1_i64) + .bind(1_i64) + .execute(pool) + .await + .expect("pending reward should insert"); + sqlx::query("UPDATE users SET is_deleted = 1, is_active = 0 WHERE id = ?") + .bind("referral-inviter") + .execute(pool) + .await + .expect("inviter should be marked deleted"); + let deleted_pass = state + .reconcile_referral_rewards_once(None) + .await + .expect("deleted inviter reconciliation should succeed"); + assert_eq!(deleted_pass.reward_attempted, 0); + assert_eq!(deleted_pass.reward_applied, 0); + let deleted_status: String = + sqlx::query_scalar("SELECT status FROM referral_rewards WHERE id = ?") + .bind("deleted-inviter-reward") + .fetch_one(pool) + .await + .expect("deleted reward should remain readable"); + assert_eq!(deleted_status, "pending"); + sqlx::query("UPDATE referral_rewards SET status = 'voided' WHERE id = ?") + .bind("deleted-inviter-reward") + .execute(pool) + .await + .expect("deleted reward should be voided after the assertion"); + sqlx::query("UPDATE users SET is_deleted = 0, is_active = 1 WHERE id = ?") + .bind("referral-inviter") + .execute(pool) + .await + .expect("inviter should be restored for reversal test"); + + // A refund can be completed after the reward transaction. The reward + // starts with zero pending debt, so this exercises the refund-aware + // candidate query rather than the pending-only retry path. + sqlx::query("UPDATE wallets SET status = 'disabled' WHERE id = ?") + .bind("wallet-inviter") + .execute(pool) + .await + .expect("inviter wallet should be disabled"); + // Processing reserves the user's refund amount before the provider + // settles it. That intermediate state must not authorize a referral + // reversal, even when the legacy payment-order counter is already + // populated. + sqlx::query( + "INSERT INTO refund_requests (id, refund_no, wallet_id, user_id, payment_order_id, source_type, refund_mode, amount_usd, status, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ) + .bind("refund-processing-reward") + .bind("refund-processing-reward-no") + .bind("wallet-invitee") + .bind("referral-invitee") + .bind("paid-order-repair") + .bind("wallet") + .bind("offline_payout") + .bind(5.0_f64) + .bind("processing") + .bind(3_i64) + .bind(3_i64) + .execute(pool) + .await + .expect("processing refund should insert"); + sqlx::query("UPDATE payment_orders SET refunded_amount_usd = ? WHERE id = ?") + .bind(5.0_f64) + .bind("paid-order-repair") + .execute(pool) + .await + .expect("payment refund should update"); + let processing_reversal = state + .reverse_referral_rewards_for_order("paid-order-repair", 5.0) + .await + .expect("processing refund should not fail referral reconciliation"); + assert!(processing_reversal.is_empty()); + let processing_candidates = state + .list_referral_reversal_candidates() + .await + .expect("processing refund candidates should be queryable"); + assert!(processing_candidates.is_empty()); + + sqlx::query("UPDATE refund_requests SET status = 'succeeded' WHERE id = ?") + .bind("refund-processing-reward") + .execute(pool) + .await + .expect("refund should settle successfully"); + let immediate_reversal = state + .reverse_referral_rewards_for_order("paid-order-repair", 5.0) + .await + .expect("completed refund should persist reversal debt"); + assert_eq!(immediate_reversal.len(), 1); + let disabled_candidates = state + .list_referral_reversal_candidates() + .await + .expect("disabled wallet candidates should be queryable"); + assert!( + disabled_candidates.is_empty(), + "a disabled wallet must not consume the bounded reversal page" + ); + let reversal = state + .reconcile_referral_rewards_once(None) + .await + .expect("refund reconciliation should succeed"); + assert_eq!(reversal.reversal_attempted, 0); + assert_eq!(reversal.reversal_applied, 0); + + let (disabled_gift, disabled_pending): (f64, f64) = sqlx::query_as( + "SELECT (SELECT gift_balance FROM wallets WHERE id = ?), (SELECT pending_reversal_amount_usd FROM referral_rewards WHERE source_order_id = ?)", + ) + .bind("wallet-inviter") + .bind("paid-order-repair") + .fetch_one(pool) + .await + .expect("disabled wallet reversal state should be readable"); + assert!((disabled_gift - 1.0).abs() < f64::EPSILON); + assert!((disabled_pending - 0.5).abs() < f64::EPSILON); + + sqlx::query("UPDATE wallets SET status = 'active' WHERE id = ?") + .bind("wallet-inviter") + .execute(pool) + .await + .expect("inviter wallet should be restored"); + let retry_reversal = state + .reconcile_referral_rewards_once(None) + .await + .expect("restored wallet reversal should succeed"); + assert_eq!(retry_reversal.reversal_attempted, 1); + assert_eq!(retry_reversal.reversal_applied, 1); + let reversal_candidates = state + .list_referral_reversal_candidates() + .await + .expect("fully reconciled reversal should not remain a candidate"); + assert!(reversal_candidates.is_empty()); + + let (gift_balance, transaction_count, oldest_transaction_at): (f64, i64, i64) = + sqlx::query_as( + "SELECT gift_balance, (SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?), (SELECT MIN(created_at) FROM wallet_transactions WHERE wallet_id = ?) FROM wallets WHERE id = ?", + ) + .bind("wallet-inviter") + .bind("wallet-inviter") + .bind("wallet-inviter") + .fetch_one(pool) + .await + .expect("wallet state should be readable"); + assert!((gift_balance - 0.5).abs() < f64::EPSILON); + assert_eq!(transaction_count, 2); + // `wallet_transactions.created_at` is stored as Unix seconds by all + // SQL adapters (the public field name retains a historical `_ms` + // suffix). A millisecond value would be roughly three orders larger. + let now_unix_secs = chrono::Utc::now().timestamp(); + assert!(oldest_transaction_at >= now_unix_secs - 60); + assert!(oldest_transaction_at <= now_unix_secs + 60); + + let (reward_count, reversed, pending, status): (i64, f64, f64, String) = + sqlx::query_as( + "SELECT COUNT(*), MAX(reversed_amount_usd), MAX(pending_reversal_amount_usd), MAX(status) FROM referral_rewards WHERE source_order_id = ?", + ) + .bind("paid-order-repair") + .fetch_one(pool) + .await + .expect("reward row should be readable"); + assert_eq!(reward_count, 1); + assert!((reversed - 0.5).abs() < f64::EPSILON); + assert!(pending.abs() < f64::EPSILON); + assert_eq!(status, "applied"); + + // A pending debt must not bypass source-order validation. This can + // happen after an operator/import corrupts a historical order while + // its inviter wallet is active again. + sqlx::query( + "UPDATE referral_rewards SET reversed_amount_usd = 0, pending_reversal_amount_usd = 0.5, status = 'applied' WHERE source_order_id = ?", + ) + .bind("paid-order-repair") + .execute(pool) + .await + .expect("pending reversal fixture should update"); + sqlx::query("UPDATE payment_orders SET amount_usd = 0 WHERE id = ?") + .bind("paid-order-repair") + .execute(pool) + .await + .expect("corrupt order fixture should update"); + let invalid_refund_pass = state + .reconcile_referral_rewards_once(None) + .await + .expect("invalid refund context should be deferred"); + assert_eq!(invalid_refund_pass.reversal_attempted, 0); + assert_eq!(invalid_refund_pass.reversal_applied, 0); + assert_eq!(invalid_refund_pass.deferred, 1); + let (gift_after_invalid, pending_after_invalid): (f64, f64) = sqlx::query_as( + "SELECT (SELECT gift_balance FROM wallets WHERE id = ?), (SELECT pending_reversal_amount_usd FROM referral_rewards WHERE source_order_id = ?)", + ) + .bind("wallet-inviter") + .bind("paid-order-repair") + .fetch_one(pool) + .await + .expect("invalid refund state should be readable"); + assert!((gift_after_invalid - 0.5).abs() < f64::EPSILON); + assert!((pending_after_invalid - 0.5).abs() < f64::EPSILON); + } } diff --git a/crates/aether-data/runtime/src/backend/sqlite.rs b/crates/aether-data/runtime/src/backend/sqlite.rs index e92e83dc8..996399560 100644 --- a/crates/aether-data/runtime/src/backend/sqlite.rs +++ b/crates/aether-data/runtime/src/backend/sqlite.rs @@ -302,7 +302,7 @@ mod tests { .await .expect("sqlite migrations should run"); - let value = serde_json::json!({"enabled": true}); + let value = serde_json::json!("enabled"); let stored = backend .upsert_system_config_entry("feature.local", &value, Some("local flag")) .await @@ -313,7 +313,40 @@ mod tests { .find_system_config_value("feature.local") .await .expect("system config should read"), - Some(value) + Some(value.clone()) + ); + let replacement = serde_json::json!("disabled"); + assert!(!backend + .compare_and_set_system_config_string_value("feature.local", "stale", "disabled") + .await + .expect("stale system config compare-and-set should complete")); + assert!(backend + .compare_and_set_system_config_string_value("feature.local", "enabled", "disabled") + .await + .expect("matching system config compare-and-set should complete")); + assert_eq!( + backend + .find_system_config_value("feature.local") + .await + .expect("updated system config should read"), + Some(replacement.clone()) + ); + sqlx::query("UPDATE system_configs SET value = ? WHERE key = ?") + .bind(r#""\u5bc6\u94a5""#) + .bind("feature.local") + .execute(backend.pool()) + .await + .expect("legacy escaped JSON string should persist"); + assert!(backend + .compare_and_set_system_config_string_value("feature.local", "密钥", "encrypted-value",) + .await + .expect("escaped JSON string compare-and-set should complete")); + assert_eq!( + backend + .find_system_config_value("feature.local") + .await + .expect("escaped JSON string replacement should read"), + Some(serde_json::json!("encrypted-value")) ); assert_eq!( backend @@ -543,6 +576,22 @@ VALUES ('target-key-1', 'target-user-1', 'hash-target-key', 'target key', 1, 1) let api_key_id_map = BTreeMap::from([("source-key-1".to_string(), "target-key-1".to_string())]); + let validation_summary = backend + .import_admin_system_usage_aggregates( + &snapshot, + &user_id_map, + &api_key_id_map, + AdminSystemUsageAggregateImportMode::ValidateError, + ) + .await + .expect("usage aggregates should validate"); + assert_eq!(validation_summary.stats_daily.created, 1); + assert_eq!(validation_summary.stats_user_daily.created, 1); + assert_eq!(validation_summary.stats_daily_api_key.created, 1); + assert_eq!(sqlite_count(backend.pool(), "stats_daily").await, 0); + assert_eq!(sqlite_count(backend.pool(), "stats_user_daily").await, 0); + assert_eq!(sqlite_count(backend.pool(), "stats_daily_api_key").await, 0); + let summary = backend .import_admin_system_usage_aggregates( &snapshot, diff --git a/crates/aether-data/runtime/src/backend/system.rs b/crates/aether-data/runtime/src/backend/system.rs index c3be27999..76f2c02f3 100644 --- a/crates/aether-data/runtime/src/backend/system.rs +++ b/crates/aether-data/runtime/src/backend/system.rs @@ -170,9 +170,10 @@ fn should_skip_imported_aggregate( match mode { AdminSystemUsageAggregateImportMode::Skip => Ok(true), AdminSystemUsageAggregateImportMode::Overwrite => Ok(false), - AdminSystemUsageAggregateImportMode::Error => Err(DataLayerError::InvalidInput(format!( - "{table} aggregate already exists for date_unix_secs={date_unix_secs}" - ))), + AdminSystemUsageAggregateImportMode::Error + | AdminSystemUsageAggregateImportMode::ValidateError => Err(DataLayerError::InvalidInput( + format!("{table} aggregate already exists for date_unix_secs={date_unix_secs}"), + )), } } diff --git a/crates/aether-data/runtime/src/backend/system/mysql.rs b/crates/aether-data/runtime/src/backend/system/mysql.rs index 15d4122b2..f62bb847f 100644 --- a/crates/aether-data/runtime/src/backend/system/mysql.rs +++ b/crates/aether-data/runtime/src/backend/system/mysql.rs @@ -392,7 +392,11 @@ ON DUPLICATE KEY UPDATE add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some()); } - tx.commit().await.map_sql_err()?; + if mode == AdminSystemUsageAggregateImportMode::ValidateError { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } Ok(summary) } @@ -447,6 +451,34 @@ LIMIT 1 .transpose() } + pub async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + let now = current_unix_secs(); + let replacement = + serialize_json_value(&serde_json::Value::String(replacement.to_string()))?; + let result = sqlx::query( + r#" +UPDATE system_configs +SET value = ?, updated_at = ? +WHERE `key` = ? + AND JSON_TYPE(value) = 'STRING' + AND BINARY JSON_UNQUOTE(value) = BINARY ? +"#, + ) + .bind(replacement) + .bind(now as i64) + .bind(key) + .bind(expected) + .execute(self.pool()) + .await + .map_sql_err()?; + Ok(result.rows_affected() > 0) + } + pub async fn upsert_system_config_value( &self, key: &str, @@ -501,6 +533,7 @@ ORDER BY `key` ASC INSERT INTO system_configs (id, `key`, value, description, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?) ON DUPLICATE KEY UPDATE + `key` = VALUES(`key`), value = VALUES(value), description = COALESCE(VALUES(description), description), updated_at = VALUES(updated_at) diff --git a/crates/aether-data/runtime/src/backend/system/postgres.rs b/crates/aether-data/runtime/src/backend/system/postgres.rs index b6596e880..01f3e4e79 100644 --- a/crates/aether-data/runtime/src/backend/system/postgres.rs +++ b/crates/aether-data/runtime/src/backend/system/postgres.rs @@ -18,6 +18,15 @@ WHERE key = $1 LIMIT 1 "#; +const COMPARE_AND_SET_SYSTEM_CONFIG_STRING_VALUE_SQL: &str = r#" +UPDATE system_configs +SET value = TO_JSON($3::text), + updated_at = NOW() +WHERE key = $1 + AND JSON_TYPEOF(value) = 'string' + AND value #>> '{}' = $2 +"#; + const UPSERT_SYSTEM_CONFIG_VALUE_SQL: &str = r#" INSERT INTO system_configs (id, key, value, description, created_at, updated_at) VALUES ($1, $2, $3, $4, NOW(), NOW()) @@ -437,7 +446,11 @@ SET api_key_name = EXCLUDED.api_key_name, add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some()); } - tx.commit().await.map_postgres_err()?; + if mode == AdminSystemUsageAggregateImportMode::ValidateError { + tx.rollback().await.map_postgres_err()?; + } else { + tx.commit().await.map_postgres_err()?; + } Ok(summary) } @@ -1157,6 +1170,22 @@ impl PostgresBackend { .map_postgres_err() } + pub async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + let result = sqlx::query(COMPARE_AND_SET_SYSTEM_CONFIG_STRING_VALUE_SQL) + .bind(key) + .bind(expected) + .bind(replacement) + .execute(self.pool()) + .await + .map_postgres_err()?; + Ok(result.rows_affected() > 0) + } + pub async fn upsert_system_config_value( &self, key: &str, diff --git a/crates/aether-data/runtime/src/backend/system/sqlite.rs b/crates/aether-data/runtime/src/backend/system/sqlite.rs index 033bd38bd..db0a41959 100644 --- a/crates/aether-data/runtime/src/backend/system/sqlite.rs +++ b/crates/aether-data/runtime/src/backend/system/sqlite.rs @@ -396,7 +396,11 @@ SET api_key_name = excluded.api_key_name, add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some()); } - tx.commit().await.map_sql_err()?; + if mode == AdminSystemUsageAggregateImportMode::ValidateError { + tx.rollback().await.map_sql_err()?; + } else { + tx.commit().await.map_sql_err()?; + } Ok(summary) } @@ -1074,6 +1078,35 @@ LIMIT 1 .transpose() } + pub async fn compare_and_set_system_config_string_value( + &self, + key: &str, + expected: &str, + replacement: &str, + ) -> Result { + let now = current_unix_secs(); + let replacement = + serialize_json_value(&serde_json::Value::String(replacement.to_string()))?; + let result = sqlx::query( + r#" +UPDATE system_configs +SET value = ?, updated_at = ? +WHERE key = ? + AND json_valid(value) + AND json_type(value) = 'text' + AND CAST(json_extract(value, '$') AS TEXT) = ? COLLATE BINARY +"#, + ) + .bind(replacement) + .bind(now as i64) + .bind(key) + .bind(expected) + .execute(self.pool()) + .await + .map_sql_err()?; + Ok(result.rows_affected() > 0) + } + pub async fn upsert_system_config_value( &self, key: &str, diff --git a/crates/aether-data/runtime/src/lifecycle/backfill/mysql.rs b/crates/aether-data/runtime/src/lifecycle/backfill/mysql.rs index 4741f4db5..66e37b23f 100644 --- a/crates/aether-data/runtime/src/lifecycle/backfill/mysql.rs +++ b/crates/aether-data/runtime/src/lifecycle/backfill/mysql.rs @@ -2,7 +2,7 @@ use std::collections::{HashMap, HashSet}; use sqlx::{ migrate::{Migrate, MigrateError, Migrator}, - query, Connection, MySqlConnection, Row, + query, query_scalar, Connection, MySqlConnection, Row, }; use tracing::{error, info, warn}; @@ -11,6 +11,7 @@ use crate::driver::mysql::MysqlPool; static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/mysql"); +const SCHEMA_BACKFILLS_TABLE_EXISTS_SQL: &str = "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = 'schema_backfills'"; const ENSURE_SCHEMA_BACKFILLS_TABLE_SQL: &str = r#" CREATE TABLE IF NOT EXISTS schema_backfills ( version BIGINT NOT NULL, @@ -161,12 +162,21 @@ async fn run_backfills_locked(conn: &mut MySqlConnection) -> Result<(), MigrateE async fn pending_backfills_locked( conn: &mut MySqlConnection, ) -> Result, MigrateError> { - ensure_schema_backfills_table(conn).await?; + if !schema_backfills_table_exists(conn).await? { + return Ok(pending_backfills_from_applied(&[])); + } let applied_backfills = list_applied_backfills(conn).await?; validate_applied_backfills(&applied_backfills)?; Ok(pending_backfills_from_applied(&applied_backfills)) } +async fn schema_backfills_table_exists(conn: &mut MySqlConnection) -> Result { + let total: i64 = query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL) + .fetch_one(&mut *conn) + .await?; + Ok(total > 0) +} + async fn ensure_schema_backfills_table(conn: &mut MySqlConnection) -> Result<(), MigrateError> { query(ENSURE_SCHEMA_BACKFILLS_TABLE_SQL) .execute(&mut *conn) diff --git a/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs b/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs index 80552bc67..a9bad0950 100644 --- a/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs +++ b/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs @@ -8,6 +8,8 @@ use tracing::{error, info, warn}; use super::types::PendingBackfillInfo; +// Historical backfill manifests must remain available after they are applied; +// deployed databases retain their versions for startup compatibility checks. static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/postgres"); const SCHEMA_BACKFILLS_TABLE_EXISTS_SQL: &str = @@ -155,17 +157,23 @@ async fn run_backfills_locked(conn: &mut PgConnection) -> Result<(), MigrateErro async fn pending_backfills_locked( conn: &mut PgConnection, ) -> Result, MigrateError> { - ensure_schema_backfills_table(conn).await?; + if !schema_backfills_table_exists(conn).await? { + return Ok(pending_backfills_from_applied(&[])); + } let applied_backfills = list_applied_backfills(conn).await?; validate_applied_backfills(&applied_backfills)?; Ok(pending_backfills_from_applied(&applied_backfills)) } -async fn ensure_schema_backfills_table(conn: &mut PgConnection) -> Result<(), MigrateError> { - let exists: bool = query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL) +async fn schema_backfills_table_exists(conn: &mut PgConnection) -> Result { + query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL) .fetch_one(&mut *conn) - .await?; - if exists { + .await + .map_err(Into::into) +} + +async fn ensure_schema_backfills_table(conn: &mut PgConnection) -> Result<(), MigrateError> { + if schema_backfills_table_exists(conn).await? { return Ok(()); } query(ENSURE_SCHEMA_BACKFILLS_TABLE_SQL) diff --git a/crates/aether-data/runtime/src/lifecycle/backfill/sqlite.rs b/crates/aether-data/runtime/src/lifecycle/backfill/sqlite.rs index 6e8a17780..d34875b12 100644 --- a/crates/aether-data/runtime/src/lifecycle/backfill/sqlite.rs +++ b/crates/aether-data/runtime/src/lifecycle/backfill/sqlite.rs @@ -2,7 +2,7 @@ use std::collections::{HashMap, HashSet}; use sqlx::{ migrate::{Migrate, MigrateError, Migrator}, - query, Connection, Row, SqliteConnection, + query, query_scalar, Connection, Row, SqliteConnection, }; use tracing::{error, info, warn}; @@ -11,6 +11,8 @@ use crate::driver::sqlite::SqlitePool; static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/sqlite"); +const SCHEMA_BACKFILLS_TABLE_EXISTS_SQL: &str = + "SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'schema_backfills'"; const ENSURE_SCHEMA_BACKFILLS_TABLE_SQL: &str = r#" CREATE TABLE IF NOT EXISTS schema_backfills ( version INTEGER NOT NULL PRIMARY KEY, @@ -162,12 +164,21 @@ async fn run_backfills_locked(conn: &mut SqliteConnection) -> Result<(), Migrate async fn pending_backfills_locked( conn: &mut SqliteConnection, ) -> Result, MigrateError> { - ensure_schema_backfills_table(conn).await?; + if !schema_backfills_table_exists(conn).await? { + return Ok(pending_backfills_from_applied(&[])); + } let applied_backfills = list_applied_backfills(conn).await?; validate_applied_backfills(&applied_backfills)?; Ok(pending_backfills_from_applied(&applied_backfills)) } +async fn schema_backfills_table_exists(conn: &mut SqliteConnection) -> Result { + let total: i64 = query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL) + .fetch_one(&mut *conn) + .await?; + Ok(total > 0) +} + async fn ensure_schema_backfills_table(conn: &mut SqliteConnection) -> Result<(), MigrateError> { query(ENSURE_SCHEMA_BACKFILLS_TABLE_SQL) .execute(&mut *conn) diff --git a/crates/aether-data/runtime/src/lifecycle/backfill/tests.rs b/crates/aether-data/runtime/src/lifecycle/backfill/tests.rs index 2e4deeee0..02c0d311a 100644 --- a/crates/aether-data/runtime/src/lifecycle/backfill/tests.rs +++ b/crates/aether-data/runtime/src/lifecycle/backfill/tests.rs @@ -310,6 +310,31 @@ INSERT INTO usage_settlement_snapshots ( } } +#[tokio::test] +async fn pending_sqlite_backfills_does_not_create_tracking_table() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite backfill status pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite schema should migrate"); + + let pending = pending_sqlite_backfills(&pool) + .await + .expect("sqlite pending backfills should load"); + assert!(!pending.is_empty()); + + let tracking_tables: i64 = query_scalar( + "SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'schema_backfills'", + ) + .fetch_one(&pool) + .await + .expect("sqlite tracking table state should load"); + assert_eq!(tracking_tables, 0); +} + #[tokio::test] async fn sqlite_backfills_apply_portable_repairs_and_record_versions() { let pool = sqlx::sqlite::SqlitePoolOptions::new() diff --git a/crates/aether-data/runtime/src/lifecycle/bootstrap/postgres.rs b/crates/aether-data/runtime/src/lifecycle/bootstrap/postgres.rs index a3cbe1f28..e5d5003c2 100644 --- a/crates/aether-data/runtime/src/lifecycle/bootstrap/postgres.rs +++ b/crates/aether-data/runtime/src/lifecycle/bootstrap/postgres.rs @@ -7,7 +7,9 @@ use tracing::info; // Generated by build.rs from schema/bootstrap/postgres. pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str = include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql")); -pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260821000000; +// Keep post-snapshot migrations executable on a fresh database so required +// schema changes after this frontier still run. +pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260821130000; const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#" SELECT COUNT(*)::BIGINT diff --git a/crates/aether-data/runtime/src/lifecycle/export.rs b/crates/aether-data/runtime/src/lifecycle/export.rs index 17756a9f4..a97bc1b08 100644 --- a/crates/aether-data/runtime/src/lifecycle/export.rs +++ b/crates/aether-data/runtime/src/lifecycle/export.rs @@ -3,12 +3,18 @@ use std::collections::{BTreeMap, BTreeSet}; #[cfg(all(feature = "postgres", feature = "sqlite"))] use futures_util::TryStreamExt; use serde_json::Value; +use sha2::{Digest, Sha256}; #[cfg(all(feature = "postgres", feature = "sqlite"))] use sqlx::Acquire; use sqlx::Row; #[cfg(any(feature = "mysql", feature = "sqlite"))] use sqlx::{Column, TypeInfo, ValueRef}; +use aether_data_contracts::repository::candidates::{ + sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data, + sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason, +}; + use crate::error::SqlResultExt; use crate::{DataLayerError, DatabaseDriver, SqlDatabaseConfig}; @@ -46,6 +52,17 @@ use postgres::normalize_postgres_import_payload; pub const EXPORT_FORMAT_VERSION: u32 = 2; const MIN_SUPPORTED_EXPORT_FORMAT_VERSION: u32 = 1; +// JSONL imports are ultimately materialized as a `DataImportPlan`, so an +// attacker-controlled document can otherwise consume memory in both the input +// string and the parsed row/payload vectors. Keep these bounds deliberately +// separate from HTTP request limits: database exports may contain large body +// blobs, while still needing a finite parser budget. The total budget is kept +// below the gateway's 256 MiB request-body ceiling because parsing duplicates +// portions of the input in serde values and the import plan. +pub const MAX_JSONL_INPUT_BYTES: usize = 256 * 1024 * 1024; +pub const MAX_JSONL_LINE_BYTES: usize = 16 * 1024 * 1024; +pub const MAX_JSONL_RECORDS: usize = 1_000_000; + #[derive( Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize, )] @@ -223,6 +240,14 @@ const AUXILIARY_TABLES: &[AuxiliaryTable] = &[ name: "usage_counter_deltas", primary_key: &["id"], }, + AuxiliaryTable { + name: "usage_cost_reservations", + primary_key: &["reservation_token"], + }, + AuxiliaryTable { + name: "usage_request_admissions", + primary_key: &["event_token"], + }, AuxiliaryTable { name: "background_task_runs", primary_key: &["id"], @@ -447,6 +472,10 @@ impl DataImportPlan { .map(Vec::as_slice) .unwrap_or(&[]) } + + fn imports_domain(&self, domain: ExportDomain) -> bool { + self.manifest.domains.contains(&domain) + } } #[derive(Debug, Clone, PartialEq)] @@ -455,6 +484,83 @@ pub struct ExportRow { pub payload: Value, } +#[derive(Debug, Clone, Default, PartialEq, Eq)] +struct IdentityImportScope { + user_ids: Vec, + oauth_link_ids: Vec, + oauth_provider_types: Vec, + finalizes_oauth_links: bool, + validates_oauth_login_methods: bool, +} + +impl IdentityImportScope { + fn from_plan(plan: &DataImportPlan) -> Result { + let scope = Self { + user_ids: imported_payload_ids(plan, ExportDomain::Users, "id")?, + oauth_link_ids: imported_payload_ids(plan, ExportDomain::UserOAuthLinks, "id")?, + oauth_provider_types: imported_payload_ids( + plan, + ExportDomain::OAuthProviders, + "provider_type", + )?, + finalizes_oauth_links: plan.imports_domain(ExportDomain::UserOAuthLinks), + validates_oauth_login_methods: plan.imports_domain(ExportDomain::UserOAuthLinks) + || plan.imports_domain(ExportDomain::OAuthProviders), + }; + if let Some(provider_type) = scope.oauth_provider_types.iter().find(|provider_type| { + provider_type.is_empty() + || provider_type.as_str() != provider_type.trim().to_ascii_lowercase() + }) { + return Err(DataLayerError::InvalidInput(format!( + "OAuth provider import has non-canonical provider_type '{provider_type}'" + ))); + } + Ok(scope) + } +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +struct IdentityImportState { + affected_user_ids: BTreeSet, +} + +fn imported_payload_ids( + plan: &DataImportPlan, + domain: ExportDomain, + payload_field: &str, +) -> Result, DataLayerError> { + plan.rows(domain) + .iter() + .map(|row| { + let payload_id = row + .payload + .as_object() + .and_then(|payload| payload.get(payload_field)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|id| !id.is_empty()) + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "{} export row '{}' must contain a non-empty string {}", + domain.as_str(), + row.id, + payload_field + )) + })?; + if payload_id != row.id { + return Err(DataLayerError::InvalidInput(format!( + "{} export row id '{}' does not match payload {} '{}'", + domain.as_str(), + row.id, + payload_field, + payload_id + ))); + } + Ok(payload_id.to_string()) + }) + .collect() +} + #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] pub struct DataCopyOptions { pub omit_request_body_details: bool, @@ -508,6 +614,139 @@ type PostgresImportColumns = BTreeMap; #[cfg(any(feature = "mysql", feature = "sqlite"))] type ImportColumnNames = BTreeSet; +const IMPORTED_CREDENTIAL_REVOKE_REASON: &str = "imported_credentials_revoked"; + +fn imported_credential_tombstone() -> String { + format!("{:x}", Sha256::digest(uuid::Uuid::new_v4().as_bytes())) +} + +fn set_supported_import_value( + object: &mut serde_json::Map, + target_has_column: &impl Fn(&str) -> bool, + column: &str, + value: Value, +) { + if target_has_column(column) { + object.insert(column.to_string(), value); + } +} + +fn deactivate_imported_credentials( + table_name: &str, + object: &mut serde_json::Map, + target_has_column: impl Fn(&str) -> bool, +) { + let table_name = table_name + .rsplit('.') + .next() + .unwrap_or(table_name) + .trim_matches(|ch| matches!(ch, '"' | '`')); + + match table_name { + "users" + if object + .get("password_hash") + .is_some_and(|value| !value.is_null()) => + { + set_supported_import_value( + object, + &target_has_column, + "password_hash", + Value::String(format!( + "$aether-import-revoked${}", + imported_credential_tombstone() + )), + ); + } + "users" => {} + "api_keys" => { + if object.contains_key("key_hash") { + set_supported_import_value( + object, + &target_has_column, + "key_hash", + Value::String(imported_credential_tombstone()), + ); + } + set_supported_import_value(object, &target_has_column, "key_encrypted", Value::Null); + set_supported_import_value( + object, + &target_has_column, + "status", + Value::String("disabled".to_string()), + ); + set_supported_import_value(object, &target_has_column, "is_active", Value::Bool(false)); + set_supported_import_value(object, &target_has_column, "is_locked", Value::Bool(true)); + } + "management_tokens" => { + if object.contains_key("token_hash") { + set_supported_import_value( + object, + &target_has_column, + "token_hash", + Value::String(imported_credential_tombstone()), + ); + } + set_supported_import_value(object, &target_has_column, "is_active", Value::Bool(false)); + } + "user_sessions" => { + if object.contains_key("refresh_token_hash") { + set_supported_import_value( + object, + &target_has_column, + "refresh_token_hash", + Value::String(imported_credential_tombstone()), + ); + } + set_supported_import_value( + object, + &target_has_column, + "prev_refresh_token_hash", + Value::Null, + ); + set_supported_import_value( + object, + &target_has_column, + "revoked_at", + Value::from(chrono::Utc::now().timestamp()), + ); + set_supported_import_value( + object, + &target_has_column, + "revoke_reason", + Value::String(IMPORTED_CREDENTIAL_REVOKE_REASON.to_string()), + ); + } + "proxy_nodes" => { + set_supported_import_value( + object, + &target_has_column, + "tunnel_generation", + Value::String(uuid::Uuid::new_v4().to_string()), + ); + set_supported_import_value( + object, + &target_has_column, + "tunnel_connected", + Value::Bool(false), + ); + set_supported_import_value( + object, + &target_has_column, + "status", + Value::String("offline".to_string()), + ); + set_supported_import_value( + object, + &target_has_column, + "active_connections", + Value::from(0), + ); + } + _ => {} + } +} + const USAGE_REQUEST_BODY_DETAIL_COLUMNS: &[&str] = &[ "request_body", "response_body", @@ -691,6 +930,27 @@ pub fn encode_jsonl(records: &[DataExportRecord]) -> Result MAX_JSONL_LINE_BYTES { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL record exceeds the {} byte line limit", + MAX_JSONL_LINE_BYTES + ))); + } + let output_len = output + .len() + .checked_add(line.len()) + .and_then(|length| length.checked_add(1)) + .ok_or_else(|| { + DataLayerError::InvalidInput( + "export JSONL exceeds the input size limit".to_string(), + ) + })?; + if output_len > MAX_JSONL_INPUT_BYTES { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL exceeds the {} byte input limit", + MAX_JSONL_INPUT_BYTES + ))); + } output.push_str(&line); output.push('\n'); } @@ -698,11 +958,42 @@ pub fn encode_jsonl(records: &[DataExportRecord]) -> Result Result, DataLayerError> { + decode_jsonl_with_limits( + input, + MAX_JSONL_INPUT_BYTES, + MAX_JSONL_LINE_BYTES, + MAX_JSONL_RECORDS, + ) +} + +fn decode_jsonl_with_limits( + input: &str, + max_input_bytes: usize, + max_line_bytes: usize, + max_records: usize, +) -> Result, DataLayerError> { + if input.len() > max_input_bytes { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL exceeds the {max_input_bytes} byte input limit" + ))); + } + let mut records = Vec::new(); for (line_index, line) in input.lines().enumerate() { + if line.len() > max_line_bytes { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL record on line {} exceeds the {max_line_bytes} byte line limit", + line_index + 1, + ))); + } if line.trim().is_empty() { continue; } + if records.len() >= max_records { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL exceeds the {max_records} record limit" + ))); + } let record = serde_json::from_str::(line).map_err(|err| { DataLayerError::InvalidInput(format!( "invalid export JSONL record on line {}: {err}", @@ -745,6 +1036,12 @@ pub fn build_import_plan(input: &str) -> Result } pub fn validate_export_records(records: &[DataExportRecord]) -> Result<(), DataLayerError> { + if records.len() > MAX_JSONL_RECORDS { + return Err(DataLayerError::InvalidInput(format!( + "export JSONL exceeds the {} record limit", + MAX_JSONL_RECORDS + ))); + } let Some(DataExportRecord::Manifest { manifest }) = records.first() else { return Err(DataLayerError::InvalidInput( "export JSONL must start with a manifest record".to_string(), @@ -1154,13 +1451,14 @@ async fn copy_postgres_sqlite_table( let mut imported = 0usize; while let Some(row) = rows.try_next().await.map_sql_err()? { - let payload = row.try_get::("payload").map_sql_err()?; - let object = payload.as_object().ok_or_else(|| { + let mut payload = row.try_get::("payload").map_sql_err()?; + let object = payload.as_object_mut().ok_or_else(|| { DataLayerError::UnexpectedValue(format!( "postgres copy row for table '{}' did not produce a JSON object", table.table_name )) })?; + prepare_postgres_sqlite_copy_payload(table, object); let mut query = sqlx::query(&target_sql); for column in &table.columns { let value = object.get(&column.sqlite.name).ok_or_else(|| { @@ -1178,6 +1476,21 @@ async fn copy_postgres_sqlite_table( Ok(imported) } +#[cfg(all(feature = "postgres", feature = "sqlite"))] +fn prepare_postgres_sqlite_copy_payload( + table: &SchemaCopyTable, + object: &mut serde_json::Map, +) { + deactivate_imported_credentials(&table.table_name, object, |column_name| { + table + .columns + .iter() + .any(|column| column.sqlite.name == column_name) + }); + sanitize_request_candidate_auxiliary_payload(&table.table_name, object); + sanitize_payment_security_payload(&table.table_name, object); +} + #[cfg(all(feature = "postgres", feature = "sqlite"))] fn postgres_schema_copy_select_sql(table: &SchemaCopyTable) -> Result { let table_sql = format!( @@ -1743,10 +2056,85 @@ fn payload_with_table(payload: Value, table_name: &str) -> Result, +) { + if table_name != "request_candidates" { + return; + } + + object.insert("error_message".to_string(), Value::Null); + sanitize_request_candidate_auxiliary_string( + object, + "skip_reason", + sanitize_request_candidate_skip_reason, + ); + sanitize_request_candidate_auxiliary_string( + object, + "error_type", + sanitize_request_candidate_error_type, + ); + sanitize_request_candidate_auxiliary_json( + object, + "extra_data", + sanitize_request_candidate_extra_data, + ); + sanitize_request_candidate_auxiliary_json( + object, + "required_capabilities", + sanitize_request_candidate_required_capabilities, + ); +} + +fn sanitize_payment_security_payload( + table_name: &str, + object: &mut serde_json::Map, +) { + match table_name { + "payment_orders" => { + object.insert("gateway_response".to_string(), Value::Null); + } + "payment_callbacks" => { + object.insert("payload".to_string(), Value::Null); + } + _ => {} + } +} + +fn sanitize_request_candidate_auxiliary_string( + object: &mut serde_json::Map, + field: &str, + sanitize: fn(Option) -> Option, +) { + let value = object + .remove(field) + .and_then(|value| value.as_str().map(ToOwned::to_owned)); + object.insert( + field.to_string(), + sanitize(value).map_or(Value::Null, Value::String), + ); +} + +fn sanitize_request_candidate_auxiliary_json( + object: &mut serde_json::Map, + field: &str, + sanitize: fn(Option) -> Option, +) { + let value = object.remove(field).and_then(|value| match value { + Value::Null => None, + Value::String(raw) => serde_json::from_str::(&raw).ok(), + value => Some(value), + }); + object.insert(field.to_string(), sanitize(value).unwrap_or(Value::Null)); +} + fn normalize_billing_payload( table_name: &str, object: &mut serde_json::Map, @@ -1800,5 +2188,197 @@ fn domain_payload_table( )) })?, }; + sanitize_request_candidate_auxiliary_payload(&table_name, &mut object); + sanitize_payment_security_payload(&table_name, &mut object); Ok((table_name, Value::Object(object))) } + +#[cfg(test)] +mod payment_export_security_tests { + use serde_json::json; + + use super::{domain_payload_table, payload_with_table, ExportRow}; + + #[test] + fn wallet_exports_and_imports_drop_payment_capabilities_and_raw_callbacks() { + let order = payload_with_table( + json!({ + "id": "order-1", + "gateway_response": { + "client_secret": "pi_1_secret_replayable", + "_stripe_client_secret_encrypted": "ciphertext", + "customer": {"email": "payer@example.com"}, + "payment_url": "https://pay.example/checkout?token=secret", + }, + }), + "payment_orders", + ) + .expect("payment order export should sanitize"); + assert!(order["gateway_response"].is_null()); + + let callback = ExportRow { + id: "payment_callbacks:callback-1".to_string(), + payload: json!({ + "__table": "payment_callbacks", + "id": "callback-1", + "payload": { + "client_secret": "pi_1_secret_replayable", + "customer_email": "payer@example.com", + }, + }), + }; + let (table, callback) = domain_payload_table(&callback, "wallet", Some("wallets")) + .expect("payment callback import should sanitize"); + assert_eq!(table, "payment_callbacks"); + assert!(callback["payload"].is_null()); + + let encoded = format!("{order}{callback}"); + for forbidden in [ + "client_secret", + "replayable", + "ciphertext", + "customer", + "payer@example.com", + "token=secret", + ] { + assert!(!encoded.contains(forbidden), "exported {forbidden}"); + } + } +} + +#[cfg(test)] +mod request_candidate_export_security_tests { + use serde_json::json; + + #[cfg(all(feature = "postgres", feature = "sqlite"))] + use serde_json::Value; + + use super::{domain_payload_table, payload_with_table, ExportRow}; + + #[cfg(all(feature = "postgres", feature = "sqlite"))] + use super::{ + prepare_postgres_sqlite_copy_payload, PostgresImportColumn, SchemaCopyColumn, + SchemaCopyTable, SqliteCopyColumn, + }; + + #[test] + fn request_candidate_auxiliary_export_and_import_drop_sensitive_diagnostics() { + let raw = json!({ + "id": "candidate-1", + "error_message": "Bearer export-secret", + "skip_reason": "secret skip reason", + "error_type": "secret error type", + "extra_data": "{\"upstream_url\":\"https://user:pass@example.com/private/export-secret?token=secret\",\"unknown\":\"secret\",\"header_rules\":[{\"id\":\"secret-rule\",\"action\":\"set\",\"name\":\"authorization\",\"value\":\"secret\"}]}", + "required_capabilities": "{\"cache_1h\":\"true\",\"tenant_secret\":\"secret\"}" + }); + + let exported = payload_with_table(raw, "request_candidates") + .expect("candidate export payload should sanitize"); + assert!(exported["error_message"].is_null()); + assert_eq!(exported["skip_reason"], "unclassified_skip"); + assert_eq!(exported["error_type"], "unclassified_error"); + assert_eq!( + exported["extra_data"]["upstream_url"], + "https://example.com/" + ); + assert_eq!(exported["extra_data"]["header_rules"]["count"], 1); + assert_eq!(exported["required_capabilities"]["cache_1h"], true); + let encoded = exported.to_string(); + for sensitive in [ + "export-secret", + "user:pass", + "secret-rule", + "authorization", + "tenant_secret", + ] { + assert!(!encoded.contains(sensitive)); + } + + let imported_row = ExportRow { + id: "request_candidates:[\"candidate-1\"]".to_string(), + payload: json!({ + "__table": "request_candidates", + "id": "candidate-1", + "error_message": "Bearer import-secret", + "extra_data": {"free_text": "import-secret"}, + "required_capabilities": {"vision": 1, "secret": "import-secret"} + }), + }; + let (table, imported) = domain_payload_table(&imported_row, "auxiliary", None) + .expect("candidate import payload should sanitize"); + assert_eq!(table, "request_candidates"); + assert!(imported["error_message"].is_null()); + assert!(imported["extra_data"].is_null()); + assert_eq!(imported["required_capabilities"], json!({"vision": true})); + assert!(!imported.to_string().contains("import-secret")); + } + + #[cfg(all(feature = "postgres", feature = "sqlite"))] + #[test] + fn postgres_to_sqlite_fast_copy_sanitizes_request_candidate_diagnostics() { + let table = SchemaCopyTable { + table_name: "request_candidates".to_string(), + columns: [ + "error_message", + "skip_reason", + "error_type", + "extra_data", + "required_capabilities", + ] + .into_iter() + .map(|name| SchemaCopyColumn { + sqlite: SqliteCopyColumn { + name: name.to_string(), + declared_type: "TEXT".to_string(), + not_null: false, + has_default: false, + primary_key_position: 0, + }, + postgres: PostgresImportColumn { + data_type: "text".to_string(), + udt_name: "text".to_string(), + is_nullable: true, + has_default: false, + }, + }) + .collect(), + }; + let mut payload = json!({ + "error_message": "Bearer fast-copy-secret", + "skip_reason": "fast-copy-secret", + "error_type": "fast-copy-secret", + "extra_data": { + "upstream_url": "https://user:pass@example.com/private?token=fast-copy-secret", + "image_progress": { + "phase": "upstream_streaming", + "message": "fast-copy-secret" + } + }, + "required_capabilities": { + "vision": 1, + "tenant_secret": "fast-copy-secret" + } + }) + .as_object() + .cloned() + .expect("copy payload should be an object"); + + prepare_postgres_sqlite_copy_payload(&table, &mut payload); + + assert!(payload["error_message"].is_null()); + assert_eq!(payload["skip_reason"], "unclassified_skip"); + assert_eq!(payload["error_type"], "unclassified_error"); + assert_eq!( + payload["extra_data"]["upstream_url"], + "https://example.com/" + ); + assert_eq!( + payload["extra_data"]["image_progress"], + json!({"phase": "upstream_streaming"}) + ); + assert_eq!(payload["required_capabilities"], json!({"vision": true})); + assert!(!Value::Object(payload) + .to_string() + .contains("fast-copy-secret")); + } +} diff --git a/crates/aether-data/runtime/src/lifecycle/export/mysql.rs b/crates/aether-data/runtime/src/lifecycle/export/mysql.rs index 38f15a42b..b241a0464 100644 --- a/crates/aether-data/runtime/src/lifecycle/export/mysql.rs +++ b/crates/aether-data/runtime/src/lifecycle/export/mysql.rs @@ -72,7 +72,9 @@ pub async fn import_mysql_plan( pool: &crate::driver::mysql::MysqlPool, plan: &DataImportPlan, ) -> Result { + let identity_scope = IdentityImportScope::from_plan(plan)?; let mut tx = pool.begin().await.map_sql_err()?; + let identity_state = capture_mysql_identity_import_state(&mut tx, &identity_scope).await?; let mut imported = 0usize; let mut column_cache = BTreeMap::::new(); for domain in &plan.manifest.domains { @@ -105,10 +107,210 @@ pub async fn import_mysql_plan( imported = imported.saturating_add(1); } } + enforce_mysql_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?; tx.commit().await.map_sql_err()?; Ok(imported) } +async fn capture_mysql_identity_import_state( + tx: &mut sqlx::Transaction<'_, sqlx::MySql>, + scope: &IdentityImportScope, +) -> Result { + let mut affected_user_ids = if scope.finalizes_oauth_links { + scope.user_ids.iter().cloned().collect::>() + } else { + BTreeSet::new() + }; + for link_id in &scope.oauth_link_ids { + if let Some(user_id) = + sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?") + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + affected_user_ids.insert(user_id); + } + } + for provider_type in &scope.oauth_provider_types { + let user_ids = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = ?", + ) + .bind(provider_type) + .fetch_all(&mut **tx) + .await + .map_sql_err()?; + affected_user_ids.extend(user_ids); + } + Ok(IdentityImportState { affected_user_ids }) +} + +async fn enforce_mysql_identity_import_invariants( + tx: &mut sqlx::Transaction<'_, sqlx::MySql>, + scope: &IdentityImportScope, + mut state: IdentityImportState, +) -> Result<(), DataLayerError> { + for user_id in &scope.user_ids { + let auth_source = + sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1") + .bind(user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported users row '{user_id}' did not produce a user record" + )) + })?; + if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") { + return Err(DataLayerError::InvalidInput(format!( + "imported user '{user_id}' has unsupported auth_source '{auth_source}'" + ))); + } + if auth_source == "oauth" { + sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?") + .bind(user_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + } + } + + for link_id in &scope.oauth_link_ids { + let user_id = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1", + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported OAuth link row '{link_id}' did not produce a link record" + )) + })?; + state.affected_user_ids.insert(user_id); + } + + for link_id in &scope.oauth_link_ids { + if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.provider_type, imported.provider_user_id +FROM user_oauth_links imported +JOIN user_oauth_links duplicate + ON duplicate.provider_type = imported.provider_type + AND duplicate.provider_user_id = imported.provider_user_id + AND duplicate.id <> imported.id +WHERE imported.id = ? +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.user_id, imported.provider_type +FROM user_oauth_links imported +JOIN user_oauth_links duplicate + ON duplicate.user_id = imported.user_id + AND duplicate.provider_type = imported.provider_type + AND duplicate.id <> imported.id +WHERE imported.id = ? +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import links user '{user_id}' to provider '{provider_type}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some(invalid_id) = sqlx::query_scalar::<_, String>( + r#" +SELECT links.id +FROM user_oauth_links links +LEFT JOIN users ON users.id = links.user_id +LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type +WHERE links.id = ? + AND ( + users.id IS NULL + OR providers.provider_type IS NULL + OR BINARY links.provider_type <> BINARY LOWER(TRIM(links.provider_type)) + OR links.provider_type = '' + OR BINARY links.provider_user_id <> BINARY TRIM(links.provider_user_id) + OR links.provider_user_id = '' + OR BINARY providers.provider_type <> BINARY LOWER(TRIM(providers.provider_type)) + ) +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import produced invalid or orphaned link '{invalid_id}'" + ))); + } + } + + if !scope.validates_oauth_login_methods { + return Ok(()); + } + for user_id in state.affected_user_ids { + if sqlx::query_scalar::<_, String>( + r#" +SELECT users.id +FROM users +WHERE users.id = ? + AND users.auth_source = 'oauth' + AND users.is_active = 1 + AND users.is_deleted = 0 + AND NOT EXISTS ( + SELECT 1 + FROM user_oauth_links links + JOIN oauth_providers providers ON providers.provider_type = links.provider_type + WHERE links.user_id = users.id + AND providers.is_enabled = 1 + AND BINARY links.provider_type = BINARY LOWER(TRIM(links.provider_type)) + AND BINARY links.provider_user_id = BINARY TRIM(links.provider_user_id) + AND links.provider_user_id <> '' + ) +LIMIT 1 +"#, + ) + .bind(&user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .is_some() + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import would leave active user '{user_id}' without an enabled identity binding" + ))); + } + } + + Ok(()) +} + fn mysql_domain_table( domain: ExportDomain, ) -> Result<(&'static str, &'static str), DataLayerError> { @@ -262,7 +464,11 @@ async fn import_mysql_row( row: &ExportRow, target_columns: &MysqlImportColumns, ) -> Result<(), DataLayerError> { - let object = filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?; + let mut object = + filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?; + deactivate_imported_credentials(table_name, &mut object, |column_name| { + target_columns.names.contains(column_name) + }); let columns = object.keys().map(String::as_str).collect::>(); for primary_key in &target_columns.primary_key { diff --git a/crates/aether-data/runtime/src/lifecycle/export/postgres.rs b/crates/aether-data/runtime/src/lifecycle/export/postgres.rs index a841792d2..956c9e38a 100644 --- a/crates/aether-data/runtime/src/lifecycle/export/postgres.rs +++ b/crates/aether-data/runtime/src/lifecycle/export/postgres.rs @@ -67,7 +67,9 @@ pub async fn import_postgres_plan( pool: &crate::driver::postgres::PostgresPool, plan: &DataImportPlan, ) -> Result { + let identity_scope = IdentityImportScope::from_plan(plan)?; let mut tx = pool.begin().await.map_sql_err()?; + let identity_state = capture_postgres_identity_import_state(&mut tx, &identity_scope).await?; let mut imported = 0usize; let mut column_cache = BTreeMap::::new(); for domain in &plan.manifest.domains { @@ -113,6 +115,7 @@ pub async fn import_postgres_plan( imported = imported.saturating_add(1); } } + enforce_postgres_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?; if !plan.rows(ExportDomain::Auxiliary).is_empty() { reset_postgres_auxiliary_sequences(&mut tx).await?; } @@ -120,6 +123,207 @@ pub async fn import_postgres_plan( Ok(imported) } +async fn capture_postgres_identity_import_state( + tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, + scope: &IdentityImportScope, +) -> Result { + let mut affected_user_ids = if scope.finalizes_oauth_links { + scope.user_ids.iter().cloned().collect::>() + } else { + BTreeSet::new() + }; + for link_id in &scope.oauth_link_ids { + if let Some(user_id) = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM public.user_oauth_links WHERE id = $1", + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + affected_user_ids.insert(user_id); + } + } + for provider_type in &scope.oauth_provider_types { + let user_ids = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM public.user_oauth_links WHERE provider_type = $1", + ) + .bind(provider_type) + .fetch_all(&mut **tx) + .await + .map_sql_err()?; + affected_user_ids.extend(user_ids); + } + Ok(IdentityImportState { affected_user_ids }) +} + +async fn enforce_postgres_identity_import_invariants( + tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, + scope: &IdentityImportScope, + mut state: IdentityImportState, +) -> Result<(), DataLayerError> { + for user_id in &scope.user_ids { + let auth_source = sqlx::query_scalar::<_, String>( + "SELECT auth_source::text FROM public.users WHERE id = $1 LIMIT 1", + ) + .bind(user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported users row '{user_id}' did not produce a user record" + )) + })?; + if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") { + return Err(DataLayerError::InvalidInput(format!( + "imported user '{user_id}' has unsupported auth_source '{auth_source}'" + ))); + } + if auth_source == "oauth" { + sqlx::query("UPDATE public.users SET email_verified = FALSE WHERE id = $1") + .bind(user_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + } + } + + for link_id in &scope.oauth_link_ids { + let user_id = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM public.user_oauth_links WHERE id = $1 LIMIT 1", + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported OAuth link row '{link_id}' did not produce a link record" + )) + })?; + state.affected_user_ids.insert(user_id); + } + + for link_id in &scope.oauth_link_ids { + if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.provider_type, imported.provider_user_id +FROM public.user_oauth_links imported +JOIN public.user_oauth_links duplicate + ON duplicate.provider_type = imported.provider_type + AND duplicate.provider_user_id = imported.provider_user_id + AND duplicate.id <> imported.id +WHERE imported.id = $1 +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.user_id, imported.provider_type +FROM public.user_oauth_links imported +JOIN public.user_oauth_links duplicate + ON duplicate.user_id = imported.user_id + AND duplicate.provider_type = imported.provider_type + AND duplicate.id <> imported.id +WHERE imported.id = $1 +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import links user '{user_id}' to provider '{provider_type}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some(invalid_id) = sqlx::query_scalar::<_, String>( + r#" +SELECT links.id +FROM public.user_oauth_links links +LEFT JOIN public.users users ON users.id = links.user_id +LEFT JOIN public.oauth_providers providers ON providers.provider_type = links.provider_type +WHERE links.id = $1 + AND ( + users.id IS NULL + OR providers.provider_type IS NULL + OR links.provider_type <> LOWER(TRIM(links.provider_type)) + OR links.provider_type = '' + OR links.provider_user_id <> TRIM(links.provider_user_id) + OR links.provider_user_id = '' + OR providers.provider_type <> LOWER(TRIM(providers.provider_type)) + ) +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import produced invalid or orphaned link '{invalid_id}'" + ))); + } + } + + if !scope.validates_oauth_login_methods { + return Ok(()); + } + for user_id in state.affected_user_ids { + if sqlx::query_scalar::<_, String>( + r#" +SELECT users.id +FROM public.users users +WHERE users.id = $1 + AND users.auth_source = 'oauth'::public.authsource + AND users.is_active IS TRUE + AND users.is_deleted IS FALSE + AND NOT EXISTS ( + SELECT 1 + FROM public.user_oauth_links links + JOIN public.oauth_providers providers ON providers.provider_type = links.provider_type + WHERE links.user_id = users.id + AND providers.is_enabled IS TRUE + AND links.provider_type = LOWER(TRIM(links.provider_type)) + AND links.provider_user_id = TRIM(links.provider_user_id) + AND links.provider_user_id <> '' + ) +LIMIT 1 +"#, + ) + .bind(&user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .is_some() + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import would leave active user '{user_id}' without an enabled identity binding" + ))); + } + } + + Ok(()) +} + async fn reset_postgres_auxiliary_sequences( tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, ) -> Result<(), DataLayerError> { @@ -359,7 +563,13 @@ async fn export_postgres_wallet_records( let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?; for row in rows { let id = row.try_get::("export_id").map_sql_err()?; - let payload = row.try_get::("payload").map_sql_err()?; + let mut payload = row.try_get::("payload").map_sql_err()?; + let object = payload.as_object_mut().ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "wallet export row in table '{export_table}' must be an object" + )) + })?; + sanitize_payment_security_payload(export_table, object); records.push(DataExportRecord::row( ExportDomain::Wallets, format!("{export_table}:{id}"), @@ -378,7 +588,10 @@ async fn import_postgres_row( row: &ExportRow, target_columns: &PostgresImportColumns, ) -> Result<(), DataLayerError> { - let object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?; + let mut object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?; + deactivate_imported_credentials(table_name, &mut object, |column_name| { + target_columns.contains_key(column_name) + }); let columns = object.keys().map(String::as_str).collect::>(); let column_sql = columns diff --git a/crates/aether-data/runtime/src/lifecycle/export/sqlite.rs b/crates/aether-data/runtime/src/lifecycle/export/sqlite.rs index 95e729ead..a15020646 100644 --- a/crates/aether-data/runtime/src/lifecycle/export/sqlite.rs +++ b/crates/aether-data/runtime/src/lifecycle/export/sqlite.rs @@ -66,7 +66,9 @@ pub async fn import_sqlite_plan( pool: &crate::driver::sqlite::SqlitePool, plan: &DataImportPlan, ) -> Result { + let identity_scope = IdentityImportScope::from_plan(plan)?; let mut tx = pool.begin().await.map_sql_err()?; + let identity_state = capture_sqlite_identity_import_state(&mut tx, &identity_scope).await?; let mut imported = 0usize; let mut column_cache = BTreeMap::::new(); for domain in &plan.manifest.domains { @@ -99,10 +101,210 @@ pub async fn import_sqlite_plan( imported = imported.saturating_add(1); } } + enforce_sqlite_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?; tx.commit().await.map_sql_err()?; Ok(imported) } +async fn capture_sqlite_identity_import_state( + tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, + scope: &IdentityImportScope, +) -> Result { + let mut affected_user_ids = if scope.finalizes_oauth_links { + scope.user_ids.iter().cloned().collect::>() + } else { + BTreeSet::new() + }; + for link_id in &scope.oauth_link_ids { + if let Some(user_id) = + sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?") + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + affected_user_ids.insert(user_id); + } + } + for provider_type in &scope.oauth_provider_types { + let user_ids = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE provider_type = ?", + ) + .bind(provider_type) + .fetch_all(&mut **tx) + .await + .map_sql_err()?; + affected_user_ids.extend(user_ids); + } + Ok(IdentityImportState { affected_user_ids }) +} + +async fn enforce_sqlite_identity_import_invariants( + tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>, + scope: &IdentityImportScope, + mut state: IdentityImportState, +) -> Result<(), DataLayerError> { + for user_id in &scope.user_ids { + let auth_source = + sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1") + .bind(user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported users row '{user_id}' did not produce a user record" + )) + })?; + if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") { + return Err(DataLayerError::InvalidInput(format!( + "imported user '{user_id}' has unsupported auth_source '{auth_source}'" + ))); + } + if auth_source == "oauth" { + sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?") + .bind(user_id) + .execute(&mut **tx) + .await + .map_sql_err()?; + } + } + + for link_id in &scope.oauth_link_ids { + let user_id = sqlx::query_scalar::<_, String>( + "SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1", + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "imported OAuth link row '{link_id}' did not produce a link record" + )) + })?; + state.affected_user_ids.insert(user_id); + } + + for link_id in &scope.oauth_link_ids { + if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.provider_type, imported.provider_user_id +FROM user_oauth_links imported +JOIN user_oauth_links duplicate + ON duplicate.provider_type = imported.provider_type + AND duplicate.provider_user_id = imported.provider_user_id + AND duplicate.id <> imported.id +WHERE imported.id = ? +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>( + r#" +SELECT imported.user_id, imported.provider_type +FROM user_oauth_links imported +JOIN user_oauth_links duplicate + ON duplicate.user_id = imported.user_id + AND duplicate.provider_type = imported.provider_type + AND duplicate.id <> imported.id +WHERE imported.id = ? +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import links user '{user_id}' to provider '{provider_type}' more than once" + ))); + } + } + + for link_id in &scope.oauth_link_ids { + if let Some(invalid_id) = sqlx::query_scalar::<_, String>( + r#" +SELECT links.id +FROM user_oauth_links links +LEFT JOIN users ON users.id = links.user_id +LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type +WHERE links.id = ? + AND ( + users.id IS NULL + OR providers.provider_type IS NULL + OR links.provider_type <> LOWER(TRIM(links.provider_type)) + OR links.provider_type = '' + OR links.provider_user_id <> TRIM(links.provider_user_id) + OR links.provider_user_id = '' + OR providers.provider_type <> LOWER(TRIM(providers.provider_type)) + ) +LIMIT 1 +"#, + ) + .bind(link_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import produced invalid or orphaned link '{invalid_id}'" + ))); + } + } + + if !scope.validates_oauth_login_methods { + return Ok(()); + } + for user_id in state.affected_user_ids { + if sqlx::query_scalar::<_, String>( + r#" +SELECT users.id +FROM users +WHERE users.id = ? + AND users.auth_source = 'oauth' + AND users.is_active = 1 + AND users.is_deleted = 0 + AND NOT EXISTS ( + SELECT 1 + FROM user_oauth_links links + JOIN oauth_providers providers ON providers.provider_type = links.provider_type + WHERE links.user_id = users.id + AND providers.is_enabled = 1 + AND links.provider_type = LOWER(TRIM(links.provider_type)) + AND links.provider_user_id = TRIM(links.provider_user_id) + AND links.provider_user_id <> '' + ) +LIMIT 1 +"#, + ) + .bind(&user_id) + .fetch_optional(&mut **tx) + .await + .map_sql_err()? + .is_some() + { + return Err(DataLayerError::InvalidInput(format!( + "OAuth import would leave active user '{user_id}' without an enabled identity binding" + ))); + } + } + + Ok(()) +} + fn sqlite_domain_table( domain: ExportDomain, ) -> Result<(&'static str, &'static str), DataLayerError> { @@ -262,7 +464,11 @@ async fn import_sqlite_row( row: &ExportRow, target_columns: &SqliteImportColumns, ) -> Result<(), DataLayerError> { - let object = filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?; + let mut object = + filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?; + deactivate_imported_credentials(table_name, &mut object, |column_name| { + target_columns.names.contains(column_name) + }); let columns = object.keys().map(String::as_str).collect::>(); let column_sql = columns diff --git a/crates/aether-data/runtime/src/lifecycle/export/tests.rs b/crates/aether-data/runtime/src/lifecycle/export/tests.rs index 27c484867..2ba3f6b65 100644 --- a/crates/aether-data/runtime/src/lifecycle/export/tests.rs +++ b/crates/aether-data/runtime/src/lifecycle/export/tests.rs @@ -1,16 +1,17 @@ use std::collections::{BTreeMap, BTreeSet}; -use serde_json::json; +use serde_json::{json, Value}; use super::{ - build_import_plan, decode_jsonl, encode_jsonl, export_mysql_core_jsonl, export_mysql_jsonl, - export_postgres_core_jsonl, export_sqlite_core_jsonl, filter_import_payload, - import_mysql_jsonl, import_postgres_jsonl, import_sqlite_jsonl, mysql_core_export_domains, - normalize_imported_binary, normalize_imported_integer_timestamp, - normalize_postgres_import_payload, postgres_bytea_json_value, postgres_core_export_domains, - sqlite_core_export_domains, sqlite_schema_copy_insert_sql, DataExportManifest, - DataExportRecord, DataImportPlan, ExportDomain, ExportRow, PostgresImportColumn, - SchemaCopyColumn, SchemaCopyTable, SqliteCopyColumn, AUXILIARY_TABLES, + build_import_plan, deactivate_imported_credentials, decode_jsonl, decode_jsonl_with_limits, + encode_jsonl, export_mysql_core_jsonl, export_mysql_jsonl, export_postgres_core_jsonl, + export_sqlite_core_jsonl, filter_import_payload, import_mysql_jsonl, import_postgres_jsonl, + import_sqlite_jsonl, mysql_core_export_domains, normalize_imported_binary, + normalize_imported_integer_timestamp, normalize_postgres_import_payload, + postgres_bytea_json_value, postgres_core_export_domains, sqlite_core_export_domains, + sqlite_schema_copy_insert_sql, DataExportManifest, DataExportRecord, DataImportPlan, + ExportDomain, ExportRow, PostgresImportColumn, SchemaCopyColumn, SchemaCopyTable, + SqliteCopyColumn, AUXILIARY_TABLES, }; use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory}; use crate::lifecycle::migrate::{ @@ -181,6 +182,25 @@ not-json"#, assert!(err.to_string().contains("line 2")); } +#[test] +fn jsonl_rejects_input_and_record_limits_before_materializing_rows() { + let oversized = "x".repeat(11); + let err = decode_jsonl_with_limits(&oversized, 10, 100, 10) + .expect_err("input byte limit should be enforced"); + assert!(err.to_string().contains("10 byte input limit")); + + let manifest = r#"{"record_type":"manifest","manifest":{"format_version":1,"created_at_unix_secs":1,"source_driver":null,"domains":[]}}"#; + let err = decode_jsonl_with_limits(manifest, usize::MAX, 10, 10) + .expect_err("line byte limit should be enforced"); + assert!(err.to_string().contains("byte line limit")); + + let row = r#"{"record_type":"manifest","manifest":{"format_version":1,"created_at_unix_secs":1,"source_driver":null,"domains":[]}}"#; + let input = format!("{row}\n{row}\n"); + let err = decode_jsonl_with_limits(&input, usize::MAX, usize::MAX, 1) + .expect_err("record limit should be enforced"); + assert!(err.to_string().contains("1 record limit")); +} + #[test] fn jsonl_rejects_duplicate_domain_ids() { let records = vec![ @@ -409,6 +429,231 @@ fn mysql_and_sqlite_import_payloads_ignore_unknown_null_columns() { ); } +#[test] +fn imported_identity_credentials_are_replaced_with_disabled_tombstones() { + let columns = BTreeSet::from([ + "password_hash".to_string(), + "key_hash".to_string(), + "key_encrypted".to_string(), + "status".to_string(), + "is_active".to_string(), + "is_locked".to_string(), + "token_hash".to_string(), + "refresh_token_hash".to_string(), + "prev_refresh_token_hash".to_string(), + "revoked_at".to_string(), + "revoke_reason".to_string(), + ]); + + let mut user = + serde_json::Map::from_iter([("password_hash".to_string(), json!("$2b$12$backup-hash"))]); + deactivate_imported_credentials("users", &mut user, |column| columns.contains(column)); + assert_ne!(user["password_hash"], json!("$2b$12$backup-hash")); + assert!(user["password_hash"] + .as_str() + .is_some_and(|value| value.starts_with("$aether-import-revoked$"))); + + let mut api_key = serde_json::Map::from_iter([ + ("key_hash".to_string(), json!("backup-key-hash")), + ("key_encrypted".to_string(), json!("backup-ciphertext")), + ("is_active".to_string(), json!(true)), + ("is_locked".to_string(), json!(false)), + ("status".to_string(), json!("active")), + ]); + deactivate_imported_credentials("api_keys", &mut api_key, |column| columns.contains(column)); + assert_ne!(api_key["key_hash"], json!("backup-key-hash")); + assert_eq!(api_key["key_encrypted"], Value::Null); + assert_eq!(api_key["is_active"], json!(false)); + assert_eq!(api_key["is_locked"], json!(true)); + assert_eq!(api_key["status"], json!("disabled")); + + let mut token = serde_json::Map::from_iter([ + ("token_hash".to_string(), json!("backup-token-hash")), + ("is_active".to_string(), json!(true)), + ]); + deactivate_imported_credentials("management_tokens", &mut token, |column| { + columns.contains(column) + }); + assert_ne!(token["token_hash"], json!("backup-token-hash")); + assert_eq!(token["is_active"], json!(false)); + + let mut session = serde_json::Map::from_iter([ + ( + "refresh_token_hash".to_string(), + json!("backup-refresh-hash"), + ), + ( + "prev_refresh_token_hash".to_string(), + json!("backup-previous-hash"), + ), + ("revoked_at".to_string(), Value::Null), + ("revoke_reason".to_string(), Value::Null), + ]); + deactivate_imported_credentials("user_sessions", &mut session, |column| { + columns.contains(column) + }); + assert_ne!(session["refresh_token_hash"], json!("backup-refresh-hash")); + assert_eq!(session["prev_refresh_token_hash"], Value::Null); + assert!(session["revoked_at"].as_i64().is_some()); + assert_eq!( + session["revoke_reason"], + json!("imported_credentials_revoked") + ); +} + +#[test] +fn imported_proxy_nodes_receive_a_new_offline_tunnel_generation() { + let columns = BTreeSet::from([ + "tunnel_generation".to_string(), + "tunnel_connected".to_string(), + "status".to_string(), + "active_connections".to_string(), + ]); + let mut node = serde_json::Map::from_iter([ + ( + "tunnel_generation".to_string(), + json!("backup-tunnel-generation"), + ), + ("tunnel_connected".to_string(), json!(true)), + ("status".to_string(), json!("online")), + ("active_connections".to_string(), json!(42)), + ( + "proxy_metadata".to_string(), + json!({"tunnel_security": {"encryption_key": "preserved-psk"}}), + ), + ]); + + deactivate_imported_credentials("public.proxy_nodes", &mut node, |column| { + columns.contains(column) + }); + + let generation = node["tunnel_generation"] + .as_str() + .expect("imported node generation should be a string"); + assert_ne!(generation, "backup-tunnel-generation"); + assert!(uuid::Uuid::parse_str(generation).is_ok()); + assert_eq!(node["tunnel_connected"], json!(false)); + assert_eq!(node["status"], json!("offline")); + assert_eq!(node["active_connections"], json!(0)); + assert_eq!( + node["proxy_metadata"]["tunnel_security"]["encryption_key"], + json!("preserved-psk") + ); +} + +#[tokio::test] +async fn sqlite_import_rotates_proxy_node_generations_and_clears_online_state() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::query( + r#" +INSERT INTO proxy_nodes ( + id, tunnel_generation, name, ip, port, status, active_connections, + tunnel_mode, tunnel_connected, proxy_metadata, created_at, updated_at +) VALUES ( + 'import-existing-node', 'target-live-generation', 'existing node', '127.0.0.1', + 8080, 'online', 9, 1, 1, '{"target":"metadata"}', 1, 1 +) +"#, + ) + .execute(&pool) + .await + .expect("existing proxy node should insert"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::ProxyNodes], + )), + DataExportRecord::row( + ExportDomain::ProxyNodes, + "import-existing-node", + json!({ + "id": "import-existing-node", + "tunnel_generation": "backup-stale-generation", + "name": "restored existing node", + "ip": "127.0.0.1", + "port": 8080, + "status": "online", + "active_connections": 42, + "tunnel_mode": true, + "tunnel_connected": true, + "proxy_metadata": { + "tunnel_security": {"encryption_key": "preserved-psk"} + }, + "created_at": 1, + "updated_at": 2 + }), + ), + DataExportRecord::row( + ExportDomain::ProxyNodes, + "import-legacy-node", + json!({ + "id": "import-legacy-node", + "name": "legacy backup node", + "ip": "127.0.0.2", + "port": 8081, + "status": "online", + "active_connections": 7, + "tunnel_mode": true, + "tunnel_connected": true, + "created_at": 1, + "updated_at": 2 + }), + ), + ]) + .expect("proxy node import fixture should encode"); + + assert_eq!( + import_sqlite_jsonl(&pool, &encoded) + .await + .expect("proxy nodes should import"), + 2 + ); + + let restored = sqlx::query_as::<_, (String, String, bool, i32, Option)>( + r#" +SELECT tunnel_generation, status, tunnel_connected, active_connections, proxy_metadata +FROM proxy_nodes +WHERE id = 'import-existing-node' +"#, + ) + .fetch_one(&pool) + .await + .expect("restored proxy node should load"); + assert_ne!(restored.0, "target-live-generation"); + assert_ne!(restored.0, "backup-stale-generation"); + assert!(uuid::Uuid::parse_str(&restored.0).is_ok()); + assert_eq!(restored.1, "offline"); + assert!(!restored.2); + assert_eq!(restored.3, 0); + assert_eq!( + restored + .4 + .as_deref() + .and_then(|value| serde_json::from_str::(value).ok()) + .and_then(|value| value["tunnel_security"]["encryption_key"] + .as_str() + .map(str::to_string)), + Some("preserved-psk".to_string()) + ); + + let legacy_generation: String = sqlx::query_scalar( + "SELECT tunnel_generation FROM proxy_nodes WHERE id = 'import-legacy-node'", + ) + .fetch_one(&pool) + .await + .expect("legacy imported proxy node should load"); + assert!(uuid::Uuid::parse_str(&legacy_generation).is_ok()); +} + #[test] fn postgres_to_sqlite_copy_uses_primary_key_upsert_instead_of_replace() { let table = SchemaCopyTable { @@ -639,6 +884,548 @@ async fn sqlite_import_rolls_back_rows_after_late_failure() { assert_eq!(count, 0); } +#[tokio::test] +async fn sqlite_users_import_fails_closed_for_oauth_email_verification() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::Users], + )), + DataExportRecord::row( + ExportDomain::Users, + "oauth-user", + json!({ + "id": "oauth-user", + "email": "oauth@example.test", + "email_verified": true, + "username": "oauth-user", + "role": "user", + "auth_source": "oauth", + "created_at": 1, + "updated_at": 1 + }), + ), + DataExportRecord::row( + ExportDomain::Users, + "local-user", + json!({ + "id": "local-user", + "email": "local@example.test", + "email_verified": true, + "username": "local-user", + "role": "user", + "auth_source": "local", + "created_at": 1, + "updated_at": 1 + }), + ), + ]) + .expect("users fixture should encode"); + + assert_eq!( + import_sqlite_jsonl(&pool, &encoded) + .await + .expect("users-only staged restore should succeed without OAuth links"), + 2 + ); + let verification = + sqlx::query_as::<_, (String, i64)>("SELECT id, email_verified FROM users ORDER BY id ASC") + .fetch_all(&pool) + .await + .expect("verification state should load"); + assert_eq!( + verification, + vec![("local-user".to_string(), 1), ("oauth-user".to_string(), 0)] + ); +} + +#[tokio::test] +async fn sqlite_users_and_providers_can_restore_before_oauth_links() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::Users, ExportDomain::OAuthProviders], + )), + DataExportRecord::row( + ExportDomain::Users, + "oauth-user", + json!({ + "id": "oauth-user", + "email": "oauth@example.test", + "email_verified": true, + "username": "oauth-user", + "role": "user", + "auth_source": "oauth", + "created_at": 1, + "updated_at": 1 + }), + ), + DataExportRecord::row( + ExportDomain::OAuthProviders, + "linuxdo", + json!({ + "provider_type": "linuxdo", + "display_name": "Linux.do", + "client_id": "client", + "redirect_uri": "https://gateway.example.test/oauth/callback", + "frontend_callback_url": "https://app.example.test/auth/callback", + "is_enabled": true, + "created_at": 1, + "updated_at": 1 + }), + ), + ]) + .expect("staged identity fixture should encode"); + + assert_eq!( + import_sqlite_jsonl(&pool, &encoded) + .await + .expect("users and Providers should restore before links"), + 2 + ); + let email_verified: i64 = + sqlx::query_scalar("SELECT email_verified FROM users WHERE id = 'oauth-user'") + .fetch_one(&pool) + .await + .expect("staged OAuth user should load"); + assert_eq!(email_verified, 0); +} + +#[tokio::test] +async fn sqlite_oauth_link_import_rolls_back_without_enabled_login_binding() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users ( + id, email, email_verified, username, role, auth_source, + is_active, is_deleted, created_at, updated_at +) VALUES ( + 'oauth-user', 'oauth@example.test', 0, 'oauth-user', 'user', 'oauth', + 1, 0, 1, 1 +); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 0, 1, 1 +); +"#, + ) + .execute(&pool) + .await + .expect("OAuth fixtures should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::UserOAuthLinks], + )), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "link-disabled", + json!({ + "id": "link-disabled", + "user_id": "oauth-user", + "provider_type": "linuxdo", + "provider_user_id": "subject-1", + "linked_at": 1 + }), + ), + ]) + .expect("OAuth link fixture should encode"); + + let err = import_sqlite_jsonl(&pool, &encoded) + .await + .expect_err("disabled-only OAuth binding should fail"); + assert!(err + .to_string() + .contains("without an enabled identity binding")); + let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links") + .fetch_one(&pool) + .await + .expect("rolled-back OAuth link count should load"); + assert_eq!(link_count, 0); +} + +#[tokio::test] +async fn sqlite_oauth_provider_import_rolls_back_if_it_removes_last_enabled_binding() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users ( + id, email, email_verified, username, role, auth_source, + is_active, is_deleted, created_at, updated_at +) VALUES ( + 'oauth-user', 'oauth@example.test', 0, 'oauth-user', 'user', 'oauth', + 1, 0, 1, 1 +); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 1, 1, 1 +); +INSERT INTO user_oauth_links ( + id, user_id, provider_type, provider_user_id, linked_at +) VALUES ( + 'existing-link', 'oauth-user', 'linuxdo', 'subject-1', 1 +); +"#, + ) + .execute(&pool) + .await + .expect("OAuth fixtures should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::OAuthProviders], + )), + DataExportRecord::row( + ExportDomain::OAuthProviders, + "linuxdo", + json!({ + "provider_type": "linuxdo", + "display_name": "Linux.do disabled", + "client_id": "client", + "redirect_uri": "https://gateway.example.test/oauth/callback", + "frontend_callback_url": "https://app.example.test/auth/callback", + "is_enabled": false, + "created_at": 1, + "updated_at": 2 + }), + ), + ]) + .expect("disabled Provider fixture should encode"); + + let err = import_sqlite_jsonl(&pool, &encoded) + .await + .expect_err("disabling the last OAuth login method should fail"); + assert!(err + .to_string() + .contains("without an enabled identity binding")); + let provider = sqlx::query_as::<_, (String, i64)>( + "SELECT display_name, is_enabled FROM oauth_providers WHERE provider_type = 'linuxdo'", + ) + .fetch_one(&pool) + .await + .expect("rolled-back Provider should load"); + assert_eq!(provider, ("Linux.do".to_string(), 1)); +} + +#[tokio::test] +async fn sqlite_oauth_link_reassignment_rolls_back_if_old_owner_loses_last_binding() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES + ('oauth-owner', 'oauth-owner', 'user', 'oauth', 1, 1), + ('local-target', 'local-target', 'user', 'local', 1, 1); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 1, 1, 1 +); +INSERT INTO user_oauth_links ( + id, user_id, provider_type, provider_user_id, linked_at +) VALUES ( + 'reassigned-link', 'oauth-owner', 'linuxdo', 'subject-1', 1 +); +"#, + ) + .execute(&pool) + .await + .expect("OAuth reassignment fixtures should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::UserOAuthLinks], + )), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "reassigned-link", + json!({ + "id": "reassigned-link", + "user_id": "local-target", + "provider_type": "linuxdo", + "provider_user_id": "subject-1", + "linked_at": 2 + }), + ), + ]) + .expect("OAuth reassignment fixture should encode"); + + let err = import_sqlite_jsonl(&pool, &encoded) + .await + .expect_err("taking the old owner's last OAuth binding should fail"); + assert!(err + .to_string() + .contains("without an enabled identity binding")); + let owner: String = + sqlx::query_scalar("SELECT user_id FROM user_oauth_links WHERE id = 'reassigned-link'") + .fetch_one(&pool) + .await + .expect("rolled-back OAuth link should load"); + assert_eq!(owner, "oauth-owner"); +} + +#[tokio::test] +async fn sqlite_oauth_link_import_ignores_unrelated_legacy_identity_damage() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES + ('broken-oauth', 'broken-oauth', 'user', 'oauth', 1, 1), + ('local-user', 'local-user', 'user', 'local', 1, 1); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 1, 1, 1 +); +INSERT INTO user_oauth_links ( + id, user_id, provider_type, provider_user_id, linked_at +) VALUES ( + 'legacy-orphan', 'missing-user', 'missing-provider', 'legacy-subject', 1 +); +"#, + ) + .execute(&pool) + .await + .expect("legacy damaged identity fixtures should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::UserOAuthLinks], + )), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "valid-link", + json!({ + "id": "valid-link", + "user_id": "local-user", + "provider_type": "linuxdo", + "provider_user_id": "valid-subject", + "linked_at": 1 + }), + ), + ]) + .expect("valid OAuth link fixture should encode"); + + assert_eq!( + import_sqlite_jsonl(&pool, &encoded) + .await + .expect("unrelated legacy damage must not block a valid scoped import"), + 1 + ); + let valid_link_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links WHERE id = 'valid-link'") + .fetch_one(&pool) + .await + .expect("valid OAuth link count should load"); + assert_eq!(valid_link_count, 1); +} + +#[tokio::test] +async fn sqlite_oauth_link_import_rejects_duplicate_identity_in_legacy_schema() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +DROP INDEX uq_user_oauth_links_provider_user; +INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES + ('local-a', 'local-a', 'user', 'local', 1, 1), + ('local-b', 'local-b', 'user', 'local', 1, 1); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 1, 1, 1 +); +"#, + ) + .execute(&pool) + .await + .expect("legacy schema fixture should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::UserOAuthLinks], + )), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "link-a", + json!({ + "id": "link-a", + "user_id": "local-a", + "provider_type": "linuxdo", + "provider_user_id": "same-subject", + "linked_at": 1 + }), + ), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "link-b", + json!({ + "id": "link-b", + "user_id": "local-b", + "provider_type": "linuxdo", + "provider_user_id": "same-subject", + "linked_at": 1 + }), + ), + ]) + .expect("duplicate identity fixture should encode"); + + let err = import_sqlite_jsonl(&pool, &encoded) + .await + .expect_err("duplicate provider identity should fail"); + assert!(err.to_string().contains("more than once")); + let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links") + .fetch_one(&pool) + .await + .expect("rolled-back OAuth link count should load"); + assert_eq!(link_count, 0); +} + +#[tokio::test] +async fn sqlite_oauth_link_import_rejects_duplicate_user_provider_in_legacy_schema() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + sqlx::raw_sql( + r#" +DROP INDEX uq_user_oauth_links_user_provider; +INSERT INTO users (id, username, role, auth_source, created_at, updated_at) +VALUES ('local-user', 'local-user', 'user', 'local', 1, 1); +INSERT INTO oauth_providers ( + provider_type, display_name, client_id, redirect_uri, frontend_callback_url, + is_enabled, created_at, updated_at +) VALUES ( + 'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback', + 'https://app.example.test/auth/callback', 1, 1, 1 +); +"#, + ) + .execute(&pool) + .await + .expect("legacy schema fixture should seed"); + + let encoded = encode_jsonl(&[ + DataExportRecord::manifest(DataExportManifest::new( + 1_700_000_000, + Some(DatabaseDriver::Postgres), + vec![ExportDomain::UserOAuthLinks], + )), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "link-a", + json!({ + "id": "link-a", + "user_id": "local-user", + "provider_type": "linuxdo", + "provider_user_id": "subject-a", + "linked_at": 1 + }), + ), + DataExportRecord::row( + ExportDomain::UserOAuthLinks, + "link-b", + json!({ + "id": "link-b", + "user_id": "local-user", + "provider_type": "linuxdo", + "provider_user_id": "subject-b", + "linked_at": 1 + }), + ), + ]) + .expect("duplicate user-provider fixture should encode"); + + let err = import_sqlite_jsonl(&pool, &encoded) + .await + .expect_err("duplicate user provider should fail"); + assert!(err.to_string().contains("more than once")); + let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links") + .fetch_one(&pool) + .await + .expect("rolled-back OAuth link count should load"); + assert_eq!(link_count, 0); +} + #[tokio::test] async fn sqlite_core_export_reads_migrated_database_rows() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -694,6 +1481,24 @@ VALUES ( 'request-1', 'candidate-1', 2, 'provider-1', 'endpoint-1', 'provider-key-1', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z' ); +INSERT INTO usage_cost_reservations ( + request_id, subject_id, reservation_token, admitted_at, + reserved_cost_units, state, reservation_expires_at, retain_until, + created_at, updated_at +) +VALUES ( + 'request-1', 'user-1', 'reservation-1', 1, + 500, 'reserved', 2, 3, + 1, 1 +); +INSERT INTO usage_request_admissions ( + request_id, subject_id, event_token, admitted_at, + retain_until, state, created_at +) +VALUES ( + 'request-1', 'user-1', 'admission-1', 1, + 3, 'active', 1 +); "#, ) .execute(&pool) @@ -757,6 +1562,19 @@ VALUES ( .any(|row| row.payload["__table"] == "usage_routing_snapshots" && row.payload["candidate_id"] == "candidate-1" && row.payload["selected_provider_id"] == "provider-1")); + assert!(import_plan + .rows(ExportDomain::Auxiliary) + .iter() + .any(|row| row.payload["__table"] == "usage_cost_reservations" + && row.payload["reservation_token"] == "reservation-1" + && row.payload["reserved_cost_units"] == 500 + && row.payload["state"] == "reserved")); + assert!(import_plan + .rows(ExportDomain::Auxiliary) + .iter() + .any(|row| row.payload["__table"] == "usage_request_admissions" + && row.payload["event_token"] == "admission-1" + && row.payload["state"] == "active")); let target_pool = sqlx::sqlite::SqlitePoolOptions::new() .max_connections(1) @@ -769,14 +1587,19 @@ VALUES ( let imported = import_sqlite_jsonl(&target_pool, &encoded) .await .expect("sqlite import should load exported rows"); - assert_eq!(imported, 20); + assert_eq!(imported, 22); - let imported_api_key = - sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'") + let imported_api_key = sqlx::query_as::<_, (String, Option, bool, bool, String)>( + "SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = 'api-key-1'", + ) .fetch_one(&target_pool) .await .expect("imported api key should load"); - assert_eq!(imported_api_key.0, "ciphertext-1"); + assert_ne!(imported_api_key.0, "hash-1"); + assert_eq!(imported_api_key.1, None); + assert!(!imported_api_key.2); + assert!(imported_api_key.3); + assert_eq!(imported_api_key.4, "disabled"); let imported_usage = sqlx::query_as::<_, (String, i64, String)>( "SELECT request_id, created_at_unix_ms, typeof(created_at_unix_ms) FROM \"usage\" WHERE request_id = 'request-1'", @@ -844,6 +1667,36 @@ WHERE request_id = 'request-1' ("candidate-1".to_string(), 2, "provider-1".to_string()) ); + let imported_reservation = sqlx::query_as::<_, (String, i64, String)>( + r#" +SELECT subject_id, reserved_cost_units, state +FROM usage_cost_reservations +WHERE reservation_token = 'reservation-1' +"#, + ) + .fetch_one(&target_pool) + .await + .expect("imported usage cost reservation should load"); + assert_eq!( + imported_reservation, + ("user-1".to_string(), 500, "reserved".to_string()) + ); + + let imported_admission = sqlx::query_as::<_, (String, String, Option)>( + r#" +SELECT subject_id, state, released_at +FROM usage_request_admissions +WHERE event_token = 'admission-1' +"#, + ) + .fetch_one(&target_pool) + .await + .expect("imported usage request admission should load"); + assert_eq!( + imported_admission, + ("user-1".to_string(), "active".to_string(), None) + ); + if let Some(database_url) = std::env::var("AETHER_TEST_POSTGRES_URL") .ok() .filter(|value| !value.trim().is_empty()) @@ -869,15 +1722,19 @@ WHERE request_id = 'request-1' let imported = import_postgres_jsonl(&postgres_pool, &encoded) .await .expect("postgres import should load exported rows"); - assert_eq!(imported, 20); + assert_eq!(imported, 22); - let imported_api_key = sqlx::query_as::<_, (String,)>( - "SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'", + let imported_api_key = sqlx::query_as::<_, (String, Option, bool, bool, String)>( + "SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = 'api-key-1'", ) .fetch_one(&postgres_pool) .await .expect("imported postgres api key should load"); - assert_eq!(imported_api_key.0, "ciphertext-1"); + assert_ne!(imported_api_key.0, "hash-1"); + assert_eq!(imported_api_key.1, None); + assert!(!imported_api_key.2); + assert!(imported_api_key.3); + assert_eq!(imported_api_key.4, "disabled"); } } @@ -1104,12 +1961,15 @@ async fn postgres_core_export_reads_migrated_database_rows_when_url_is_set() { assert_eq!(imported, import_plan_row_count(&import_plan)); let imported_api_key = - sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = $1") + sqlx::query_as::<_, (Option,)>("SELECT key_encrypted FROM api_keys WHERE id = $1") .bind(&api_key_id) .fetch_one(&target_pool) .await .expect("imported sqlite api key should load"); - assert_eq!(imported_api_key.0, "ciphertext-1"); + assert_eq!( + imported_api_key.0, None, + "cross-driver imports must revoke recoverable API-key ciphertext" + ); let imported_global_model_timestamps = sqlx::query_as::<_, (i64, i64, String, String)>( "SELECT created_at, updated_at, typeof(created_at), typeof(updated_at) FROM global_models WHERE id = ?", ) @@ -1317,13 +2177,18 @@ async fn mysql_core_export_reads_migrated_database_rows_when_url_is_set() { .expect("mysql import should be idempotent"); assert!(imported >= 6); - let imported_api_key = - sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = ?") - .bind(&api_key_id) - .fetch_one(&pool) - .await - .expect("imported mysql api key should load"); - assert_eq!(imported_api_key.0, "ciphertext-1"); + let imported_api_key = sqlx::query_as::<_, (String, Option, bool, bool, String)>( + "SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = ?", + ) + .bind(&api_key_id) + .fetch_one(&pool) + .await + .expect("imported mysql api key should load"); + assert_ne!(imported_api_key.0, "hash-1"); + assert_eq!(imported_api_key.1, None); + assert!(!imported_api_key.2); + assert!(imported_api_key.3); + assert_eq!(imported_api_key.4, "disabled"); } fn unique_suffix() -> String { diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs index 0084436e4..42180e917 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs @@ -23,7 +23,8 @@ use super::{ prepare_database_for_startup, }; use crate::lifecycle::bootstrap::postgres::{ - snapshot_migrations as empty_database_snapshot_migrations, EMPTY_DATABASE_SNAPSHOT_SQL, + snapshot_migrations as empty_database_snapshot_migrations, + EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL, }; #[derive(Debug)] @@ -166,6 +167,30 @@ impl ManagedPostgresServer { } } +/// A clean PostgreSQL database is bootstrapped from the schema snapshot first; +/// migrations after the privacy/security frontier are intentionally left +/// pending so their data-preserving changes still execute. Exercise the same +/// prepare-then-run sequence used by gateway startup before asserting that the +/// database is current. +async fn prepare_and_apply_clean_postgres_database(pool: &PgPool) { + let pending = prepare_database_for_startup(pool) + .await + .expect("clean database bootstrap should succeed"); + if !pending.is_empty() { + super::run_migrations(pool) + .await + .expect("pending PostgreSQL migrations should apply"); + } + + let pending = prepare_database_for_startup(pool) + .await + .expect("PostgreSQL startup preparation should re-check migrations"); + assert!( + pending.is_empty(), + "clean PostgreSQL database should be current after migrations: {pending:?}" + ); +} + fn local_postgres_tests_required() -> bool { // CI can opt into failing when the isolated local PostgreSQL fixture is unavailable. std::env::var("AETHER_REQUIRE_LOCAL_POSTGRES_TESTS") @@ -411,7 +436,12 @@ fn empty_database_snapshot_covers_current_cutoff_versions() { 20260720000000, 20260727000000, 20260731000000, + 20260814000000, + 20260815000000, + 20260816000000, 20260821000000, + 20260821120000, + 20260821130000, ] ); } @@ -489,6 +519,9 @@ fn portable_driver_migrations_create_the_postgres_table_set() { .iter() .filter(|migration| migration.migration_type.is_up_migration()) .flat_map(|migration| create_table_names(migration.sql.as_ref())) + // SQLite rebuilds tables to add foreign keys. These staging tables are + // renamed to the canonical table names before the migration finishes. + .filter(|table| !table.ends_with("_with_user_fk")) .collect::>(); assert_eq!(mysql_tables, postgres_tables, "MySQL table set drifted"); @@ -1022,6 +1055,149 @@ fn worker_boot_cleanup_migration_is_enabled_for_every_driver() { } } +const UNPUBLISHED_LEGACY_DATA_REWRITE_MIGRATION_VERSIONS: &[i64] = &[ + 20260822000000, + 20260822010000, + 20260822020000, + 20260827000000, + 20260827010000, + 20260827020000, + 20260827030000, + 20260829000000, + 20260903010000, +]; + +#[test] +fn unpublished_legacy_data_rewrite_migrations_are_absent_for_every_driver() { + for (driver, migrator) in [ + ("postgres", &POSTGRES_MIGRATOR), + ("mysql", &super::mysql::MIGRATOR), + ("sqlite", &super::sqlite::MIGRATOR), + ] { + for version in UNPUBLISHED_LEGACY_DATA_REWRITE_MIGRATION_VERSIONS { + assert!( + migrator + .iter() + .all(|migration| migration.version != *version), + "{driver} must not embed unpublished legacy rewrite migration {version}" + ); + } + } +} + +#[test] +fn deleted_user_history_schema_decoupling_does_not_rewrite_history() { + const VERSION: i64 = 20260827050000; + + for (driver, migrator) in [ + ("postgres", &POSTGRES_MIGRATOR), + ("mysql", &super::mysql::MIGRATOR), + ("sqlite", &super::sqlite::MIGRATOR), + ] { + let migration = migrator + .iter() + .find(|migration| migration.version == VERSION) + .unwrap_or_else(|| panic!("{driver} user-history schema migration should be embedded")); + let sql = migration + .sql + .lines() + .map(str::trim) + .filter(|line| !line.is_empty() && !line.starts_with("--")) + .collect::>() + .join("\n") + .to_ascii_uppercase(); + for history_rewrite in ["UPDATE ", "DELETE ", "TRUNCATE ", "REPLACE ", "MERGE "] { + assert!( + !sql + .split(';') + .any(|statement| statement.trim_start().starts_with(history_rewrite)), + "{driver} user-history schema migration rewrites legacy rows with {history_rewrite}" + ); + } + } + + let postgres_migration = POSTGRES_MIGRATOR + .iter() + .find(|migration| migration.version == VERSION) + .expect("postgres user-history schema migration should be embedded"); + for constraint in [ + "request_candidates_user_id_fkey", + "video_tasks_user_id_fkey", + "usage_user_id_fkey", + "stats_user_daily_user_id_fkey", + "stats_user_summary_user_id_fkey", + "stats_user_daily_model_user_id_fkey", + "stats_user_daily_provider_user_id_fkey", + "stats_user_daily_api_format_user_id_fkey", + "stats_user_daily_model_provider_user_id_fkey", + "stats_user_daily_cost_savings_user_id_fkey", + "stats_user_daily_cost_savings_provider_user_id_fkey", + "stats_user_daily_cost_savings_model_user_id_fkey", + "stats_user_daily_cost_savings_model_provider_user_id_fkey", + "stats_hourly_user_model_user_id_fkey", + "user_model_usage_counts_user_id_fkey", + "request_candidates_api_key_id_fkey", + "video_tasks_api_key_id_fkey", + "usage_api_key_id_fkey", + "stats_daily_api_key_api_key_id_fkey", + "audit_logs_user_id_fkey", + "payment_orders_user_id_fkey", + "refund_requests_user_id_fkey", + "wallet_transactions_operator_id_fkey", + "wallets_user_id_fkey", + "wallets_api_key_id_fkey", + "user_plan_entitlements_user_id_fkey", + "entitlement_usage_ledgers_user_id_fkey", + "user_referrals_inviter_user_id_fkey", + "user_referrals_invitee_user_id_fkey", + "referral_rewards_inviter_user_id_fkey", + "referral_rewards_invitee_user_id_fkey", + ] { + assert!( + postgres_migration + .sql + .contains(&format!("DROP CONSTRAINT IF EXISTS {constraint}")), + "postgres migration must decouple {constraint}" + ); + } + + let mysql_migration = super::mysql::MIGRATOR + .iter() + .find(|migration| migration.version == VERSION) + .expect("mysql user-history schema migration should be embedded"); + for constraint in [ + "user_plan_entitlements_user_id_fkey", + "entitlement_usage_ledgers_user_id_fkey", + "user_referrals_inviter_user_id_fkey", + "user_referrals_invitee_user_id_fkey", + "referral_rewards_inviter_user_id_fkey", + "referral_rewards_invitee_user_id_fkey", + ] { + assert!( + mysql_migration.sql.contains(constraint), + "mysql migration must decouple {constraint}" + ); + } + + let sqlite_migration = super::sqlite::MIGRATOR + .iter() + .find(|migration| migration.version == VERSION) + .expect("sqlite user-history schema migration should be embedded"); + for table in [ + "user_plan_entitlements", + "entitlement_usage_ledgers", + "user_referrals", + "referral_rewards", + ] { + assert!( + sqlite_migration + .sql + .contains(&format!("ALTER TABLE {table} RENAME TO")), + "sqlite migration must rebuild {table} without the legacy user foreign key" + ); + } +} + #[test] fn mysql_and_sqlite_migrations_include_enabled_incrementals() { let mysql_versions = super::mysql::MIGRATOR @@ -1065,7 +1241,20 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() { 20260725030000, 20260727000000, 20260731000000, + 20260814000000, + 20260815000000, + 20260816000000, + 20260817000000, 20260821000000, + 20260821120000, + 20260821130000, + 20260827040000, + 20260827050000, + 20260831000000, + 20260831010000, + 20260831020000, + 20260831030000, + 20260903000000, ] ); assert_eq!( @@ -1100,11 +1289,167 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() { 20260725040000, 20260727000000, 20260731000000, + 20260814000000, + 20260815000000, + 20260816000000, 20260821000000, + 20260821120000, + 20260821130000, + 20260827050000, + 20260831000000, + 20260831010000, + 20260831020000, + 20260831030000, + 20260903000000, ] ); } +#[tokio::test] +async fn sqlite_gateway_order_uniqueness_migration_rejects_historical_duplicates() { + const VERSION: i64 = 20260821120000; + + let pool = SqlitePool::connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + let mut connection = pool.acquire().await.expect("sqlite connection should open"); + connection + .ensure_migrations_table() + .await + .expect("migration table should be created"); + for migration in super::sqlite::MIGRATOR + .iter() + .filter(|migration| migration.version < VERSION) + { + connection + .apply(migration) + .await + .expect("pre-uniqueness migration should apply"); + } + drop(connection); + + query( + r#" +INSERT INTO wallets ( + id, user_id, balance, gift_balance, limit_mode, currency, status, + total_recharged, total_consumed, total_refunded, total_adjusted, + created_at, updated_at +) VALUES + ('duplicate-wallet-a', 'duplicate-user-a', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1), + ('duplicate-wallet-b', 'duplicate-user-b', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1); + +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd, + refundable_amount_usd, payment_method, gateway_order_id, status, created_at +) VALUES + ('duplicate-order-a', 'duplicate-no-a', 'duplicate-wallet-a', 'duplicate-user-a', 1, 0, 0, ' EPAY ', 'duplicate-gateway-id', 'pending', 1), + ('duplicate-order-b', 'duplicate-no-b', 'duplicate-wallet-b', 'duplicate-user-b', 1, 0, 0, 'epay', 'duplicate-gateway-id', 'pending', 1); +"#, + ) + .execute(&pool) + .await + .expect("historical duplicate fixtures should insert"); + + let migration = super::sqlite::MIGRATOR + .iter() + .find(|migration| migration.version == VERSION) + .expect("gateway-order uniqueness migration should be embedded"); + let error = sqlx::raw_sql(migration.sql.as_ref()) + .execute(&pool) + .await + .expect_err("historical financial duplicates must block migration"); + assert!(error.to_string().to_ascii_lowercase().contains("unique")); + + let order_count: i64 = query_scalar( + "SELECT COUNT(*) FROM payment_orders WHERE gateway_order_id = 'duplicate-gateway-id'", + ) + .fetch_one(&pool) + .await + .expect("duplicate financial records should remain intact"); + assert_eq!(order_count, 2); +} + +#[tokio::test] +async fn sqlite_gateway_order_uniqueness_migration_normalizes_legacy_payment_methods() { + const VERSION: i64 = 20260821120000; + + let pool = SqlitePool::connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + let mut connection = pool.acquire().await.expect("sqlite connection should open"); + connection + .ensure_migrations_table() + .await + .expect("migration table should be created"); + for migration in super::sqlite::MIGRATOR + .iter() + .filter(|migration| migration.version < VERSION) + { + connection + .apply(migration) + .await + .expect("pre-uniqueness migration should apply"); + } + drop(connection); + + query( + r#" +INSERT INTO wallets ( + id, user_id, balance, gift_balance, limit_mode, currency, status, + total_recharged, total_consumed, total_refunded, total_adjusted, + created_at, updated_at +) VALUES ('legacy-method-wallet', 'legacy-method-user', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1); + +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd, + refundable_amount_usd, payment_method, gateway_order_id, status, created_at +) VALUES ('legacy-method-order', 'legacy-method-no', 'legacy-method-wallet', 'legacy-method-user', 1, 0, 0, ' EPAY ', 'CaseSensitiveTxn', 'pending', 1); + +INSERT INTO payment_callbacks ( + id, payment_method, callback_key, signature_valid, status, created_at +) VALUES ('legacy-method-callback', ' EPAY ', 'legacy-method-key', 0, 'received', 1); +"#, + ) + .execute(&pool) + .await + .expect("legacy payment methods should insert"); + + let migration = super::sqlite::MIGRATOR + .iter() + .find(|migration| migration.version == VERSION) + .expect("gateway-order uniqueness migration should be embedded"); + sqlx::raw_sql(migration.sql.as_ref()) + .execute(&pool) + .await + .expect("non-conflicting legacy payment methods should normalize"); + + let order_method: String = + query_scalar("SELECT payment_method FROM payment_orders WHERE id = 'legacy-method-order'") + .fetch_one(&pool) + .await + .expect("normalized order should load"); + let callback_method: String = query_scalar( + "SELECT payment_method FROM payment_callbacks WHERE id = 'legacy-method-callback'", + ) + .fetch_one(&pool) + .await + .expect("normalized callback should load"); + assert_eq!(order_method, "epay"); + assert_eq!(callback_method, "epay"); + + query( + r#" +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd, + refundable_amount_usd, payment_method, gateway_order_id, status, created_at +) VALUES ('case-sensitive-order', 'case-sensitive-no', 'legacy-method-wallet', 'legacy-method-user', 1, 0, 0, 'epay', 'casesensitivetxn', 'pending', 1); +"#, + ) + .execute(&pool) + .await + .expect("case-distinct opaque gateway identifiers should remain distinct"); +} + #[tokio::test] async fn sqlite_imported_timestamp_migration_normalizes_text_storage() { let pool = SqlitePool::connect("sqlite::memory:") @@ -2207,13 +2552,25 @@ fn pending_migrations_from_applied_skips_versions_already_applied() { 20260720000000, 20260727000000, 20260731000000, + 20260814000000, + 20260815000000, + 20260816000000, 20260821000000, + 20260821120000, + 20260821130000, + 20260827040000, + 20260827050000, + 20260831000000, + 20260831010000, + 20260831030000, + 20260901000000, + 20260903000000, ] ); } #[test] -fn pending_migrations_from_applied_is_empty_after_empty_database_snapshot_stamp() { +fn pending_migrations_from_applied_only_returns_post_snapshot_migrations() { let applied = empty_database_snapshot_migrations(&POSTGRES_MIGRATOR) .expect("empty database snapshot migrations should resolve") .into_iter() @@ -2224,11 +2581,15 @@ fn pending_migrations_from_applied_is_empty_after_empty_database_snapshot_stamp( .collect::>(); let pending = pending_migrations_from_applied(&applied); + let expected = all_up_migrations() + .into_iter() + .filter(|migration| migration.version > EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION) + .collect::>(); - assert!( - pending.is_empty(), - "empty database snapshot-stamped databases should not require a manual migration before first startup" - ); + assert_eq!(pending, expected); + assert!(pending + .iter() + .all(|migration| migration.version > EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION)); } #[tokio::test] @@ -2709,14 +3070,7 @@ async fn prepare_database_for_startup_bootstraps_clean_database() { let pool = PgPool::connect(server.database_url()) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - - assert!( - pending.is_empty(), - "fresh databases should not report pending migrations after startup preparation" - ); + prepare_and_apply_clean_postgres_database(&pool).await; assert!(table_exists(&pool, "users") .await .expect("users lookup should succeed")); @@ -2744,9 +3098,8 @@ async fn prepare_database_for_startup_bootstraps_clean_database() { .expect("migration count query should succeed"); assert_eq!( applied_count, - empty_database_snapshot_migrations(&POSTGRES_MIGRATOR) - .expect("baseline migrations should resolve") - .len() as i64 + all_up_migrations().len() as i64, + "fresh database should record the snapshot and every post-snapshot migration" ); } @@ -2762,13 +3115,7 @@ async fn postgres_request_candidates_preserve_deleted_api_key_identity() { let pool = PgPool::connect(server.database_url()) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - assert!( - pending.is_empty(), - "clean database bootstrap should not leave pending migrations: {pending:?}" - ); + prepare_and_apply_clean_postgres_database(&pool).await; query( r#" @@ -2883,8 +3230,8 @@ INSERT INTO public.usage ( .await .expect("daily stats API key name snapshot should be readable"); assert_eq!( - stats_api_key_name.as_deref(), - Some("Deleted API Key Snapshot") + stats_api_key_name, None, + "deleting an API key must anonymize its historical name while preserving its ID" ); query( @@ -2981,13 +3328,7 @@ async fn postgres_expired_api_key_cleanup_preserves_historical_identity() { let pool = PgPool::connect(database_url) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - assert!( - pending.is_empty(), - "clean database bootstrap should not leave pending migrations: {pending:?}" - ); + prepare_and_apply_clean_postgres_database(&pool).await; query( r#" @@ -3145,13 +3486,7 @@ async fn postgres_api_key_leaderboard_user_filter_preserves_aggregate_history() let pool = PgPool::connect(database_url) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - assert!( - pending.is_empty(), - "clean database bootstrap should not leave pending migrations: {pending:?}" - ); + prepare_and_apply_clean_postgres_database(&pool).await; query( r#" @@ -3471,10 +3806,7 @@ async fn postgres_usage_billing_facts_total_tokens_counts_cached_input_once() { let pool = PgPool::connect(server.database_url()) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - assert!(pending.is_empty()); + prepare_and_apply_clean_postgres_database(&pool).await; let legacy_view_migration = POSTGRES_MIGRATOR .iter() @@ -3674,10 +4006,7 @@ async fn postgres_migrations_repair_invalid_concurrent_cleanup_index() { let pool = PgPool::connect(server.database_url()) .await .expect("pool should connect"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("clean database bootstrap should succeed"); - assert!(pending.is_empty()); + prepare_and_apply_clean_postgres_database(&pool).await; query("DROP INDEX CONCURRENTLY public.idx_usage_legacy_body_ref_cleanup_created_at") .execute(&pool) @@ -3767,14 +4096,7 @@ async fn prepare_database_for_startup_bootstraps_when_only_unrelated_public_tabl .await .expect("fixture table should be created"); - let pending = prepare_database_for_startup(&pool) - .await - .expect("startup preparation should tolerate unrelated public tables"); - - assert!( - pending.is_empty(), - "unrelated public tables should not block baseline bootstrap on first startup" - ); + prepare_and_apply_clean_postgres_database(&pool).await; assert!(table_exists(&pool, "vendor_bootstrap_marker") .await .expect("fixture table lookup should succeed")); diff --git a/crates/aether-data/runtime/src/repository/auth/memory.rs b/crates/aether-data/runtime/src/repository/auth/memory.rs index c38d8ee40..a43b63344 100644 --- a/crates/aether-data/runtime/src/repository/auth/memory.rs +++ b/crates/aether-data/runtime/src/repository/auth/memory.rs @@ -6,11 +6,12 @@ use async_trait::async_trait; use super::{ AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository, - AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, - StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, - UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, + AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, + CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, + StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; use crate::repository::usage::{ApiKeyUsageContribution, ApiKeyUsageDelta}; +use crate::repository::users::StoredUserAuthRecord; use crate::DataLayerError; fn current_unix_secs() -> u64 { @@ -25,11 +26,73 @@ struct MemoryAuthApiKeyIndex { by_api_key_id: BTreeMap, export_by_api_key_id: BTreeMap, by_key_hash: BTreeMap, + owner_by_user_id: BTreeMap, touch_counts: BTreeMap, snapshot_lookup_counts: BTreeMap, key_hash_lookup_counts: BTreeMap, } +#[derive(Debug, Clone, PartialEq, Eq)] +struct MemoryAuthApiKeyOwnerSnapshot { + user_id: String, + username: String, + email: Option, + user_role: String, + user_auth_source: String, + user_is_active: bool, + user_is_deleted: bool, + user_rate_limit: Option, + user_allowed_providers: Option>, + user_allowed_api_formats: Option>, + user_allowed_models: Option>, +} + +impl From<&StoredAuthApiKeySnapshot> for MemoryAuthApiKeyOwnerSnapshot { + fn from(snapshot: &StoredAuthApiKeySnapshot) -> Self { + Self { + user_id: snapshot.user_id.clone(), + username: snapshot.username.clone(), + email: snapshot.email.clone(), + user_role: snapshot.user_role.clone(), + user_auth_source: snapshot.user_auth_source.clone(), + user_is_active: snapshot.user_is_active, + user_is_deleted: snapshot.user_is_deleted, + user_rate_limit: snapshot.user_rate_limit, + user_allowed_providers: snapshot.user_allowed_providers.clone(), + user_allowed_api_formats: snapshot.user_allowed_api_formats.clone(), + user_allowed_models: snapshot.user_allowed_models.clone(), + } + } +} + +impl From<&StoredUserAuthRecord> for MemoryAuthApiKeyOwnerSnapshot { + fn from(user: &StoredUserAuthRecord) -> Self { + Self { + user_id: user.id.clone(), + username: user.username.clone(), + email: user.email.clone(), + user_role: user.role.clone(), + user_auth_source: user.auth_source.clone(), + user_is_active: user.is_active, + user_is_deleted: user.is_deleted, + user_rate_limit: None, + user_allowed_providers: user.allowed_providers.clone(), + user_allowed_api_formats: user.allowed_api_formats.clone(), + user_allowed_models: user.allowed_models.clone(), + } + } +} + +// The trusted snapshot is returned on the common path so callers can use the +// complete immutable view without another lookup. Preserve this established +// in-memory registry representation rather than introducing heap allocation. +#[allow(clippy::large_enum_variant)] +#[derive(Debug)] +enum MemoryAuthApiKeyOwnerRegistryEntry { + Trusted(MemoryAuthApiKeyOwnerSnapshot), + Conflicted, +} + #[derive(Debug, Default)] pub struct InMemoryAuthApiKeySnapshotRepository { index: RwLock, @@ -44,7 +107,12 @@ impl InMemoryAuthApiKeySnapshotRepository { let mut by_api_key_id = BTreeMap::new(); let mut export_by_api_key_id = BTreeMap::new(); let mut by_key_hash = BTreeMap::new(); + let mut owner_by_user_id = BTreeMap::new(); for (key_hash, snapshot) in items { + Self::register_owner_snapshot( + &mut owner_by_user_id, + MemoryAuthApiKeyOwnerSnapshot::from(&snapshot), + ); let derived_key_hash = key_hash .clone() .unwrap_or_else(|| format!("memory-{}", snapshot.api_key_id)); @@ -101,6 +169,7 @@ impl InMemoryAuthApiKeySnapshotRepository { by_api_key_id, export_by_api_key_id, by_key_hash, + owner_by_user_id, touch_counts: BTreeMap::new(), snapshot_lookup_counts: BTreeMap::new(), key_hash_lookup_counts: BTreeMap::new(), @@ -109,6 +178,26 @@ impl InMemoryAuthApiKeySnapshotRepository { } } + /// Registers trusted owner state without inserting an API key. Only the + /// owner fields are retained; the API-key fields of each fixture snapshot + /// are intentionally ignored. + pub fn with_owner_snapshots(mut self, items: I) -> Self + where + I: IntoIterator, + { + let index = self + .index + .get_mut() + .expect("auth api key snapshot repository lock"); + for snapshot in items { + Self::register_owner_snapshot( + &mut index.owner_by_user_id, + MemoryAuthApiKeyOwnerSnapshot::from(&snapshot), + ); + } + self + } + pub fn with_lookup_delay_for_tests(mut self, delay: Duration) -> Self { self.lookup_delay = Some(delay); self @@ -202,12 +291,56 @@ impl InMemoryAuthApiKeySnapshotRepository { record.total_cost_usd = contribution.total_cost_usd.max(0.0); } } + + fn remove_api_key(index: &mut MemoryAuthApiKeyIndex, api_key_id: &str) { + let key_hashes = index + .by_key_hash + .iter() + .filter(|(_, mapped_api_key_id)| mapped_api_key_id.as_str() == api_key_id) + .map(|(key_hash, _)| key_hash.clone()) + .collect::>(); + index.by_api_key_id.remove(api_key_id); + index.export_by_api_key_id.remove(api_key_id); + index.by_key_hash.retain(|_, value| value != api_key_id); + index.touch_counts.remove(api_key_id); + index.snapshot_lookup_counts.remove(api_key_id); + for key_hash in key_hashes { + index.key_hash_lookup_counts.remove(&key_hash); + } + } + + fn register_owner_snapshot( + owners: &mut BTreeMap, + owner: MemoryAuthApiKeyOwnerSnapshot, + ) { + use std::collections::btree_map::Entry; + + match owners.entry(owner.user_id.clone()) { + Entry::Vacant(entry) => { + entry.insert(MemoryAuthApiKeyOwnerRegistryEntry::Trusted(owner)); + } + Entry::Occupied(mut entry) => { + let matches_existing = matches!( + entry.get(), + MemoryAuthApiKeyOwnerRegistryEntry::Trusted(existing) if existing == &owner + ); + if !matches_existing { + entry.insert(MemoryAuthApiKeyOwnerRegistryEntry::Conflicted); + } + } + } + } } fn clamp_i64_to_u64(value: i64) -> u64 { u64::try_from(value).unwrap_or_default() } +fn i64_from_u64(value: u64, field_name: &str) -> Result { + i64::try_from(value) + .map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}"))) +} + fn apply_i64_delta_to_u64(current: u64, delta: i64) -> u64 { clamp_i64_to_u64( i64::try_from(current) @@ -394,6 +527,7 @@ impl AuthApiKeyReadRepository for InMemoryAuthApiKeySnapshotRepository { user_ids: &[String], now_unix_secs: u64, ) -> Result { + i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let index = self .index .read() @@ -418,6 +552,7 @@ impl AuthApiKeyReadRepository for InMemoryAuthApiKeySnapshotRepository { &self, now_unix_secs: u64, ) -> Result { + i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let index = self .index .read() @@ -459,6 +594,7 @@ impl AuthApiKeyReadRepository for InMemoryAuthApiKeySnapshotRepository { &self, now_unix_secs: u64, ) -> Result { + i64_from_u64(now_unix_secs, "api_keys.summary_now")?; let index = self .index .read() @@ -515,6 +651,33 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(true) } + async fn synchronize_user_api_key_owner_for_tests( + &self, + user: &StoredUserAuthRecord, + ) -> Result<(), DataLayerError> { + if user.id.trim().is_empty() + || user.username.trim().is_empty() + || user.role.trim().is_empty() + || user.auth_source.trim().is_empty() + || user.security_version < 0 + { + return Err(DataLayerError::InvalidInput( + "invalid authoritative API-key owner snapshot".to_string(), + )); + } + + let owner = MemoryAuthApiKeyOwnerSnapshot::from(user); + self.index + .write() + .expect("auth api key snapshot repository lock") + .owner_by_user_id + .insert( + owner.user_id.clone(), + MemoryAuthApiKeyOwnerRegistryEntry::Trusted(owner), + ); + Ok(()) + } + async fn create_user_api_key( &self, record: CreateUserApiKeyRecord, @@ -523,6 +686,19 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { .index .write() .expect("auth api key snapshot repository lock"); + // Match the database adapters: inactive owners may retain or receive keys for + // administrative restore workflows, but unknown/deleted owners must fail closed. The + // resulting snapshot remains unusable while its trusted owner is inactive. + let owner = match index.owner_by_user_id.get(&record.user_id) { + Some(MemoryAuthApiKeyOwnerRegistryEntry::Trusted(owner)) if !owner.user_is_deleted => { + owner.clone() + } + Some( + MemoryAuthApiKeyOwnerRegistryEntry::Trusted(_) + | MemoryAuthApiKeyOwnerRegistryEntry::Conflicted, + ) + | None => return Ok(None), + }; if index.by_api_key_id.contains_key(&record.api_key_id) { return Err(DataLayerError::UnexpectedValue(format!( "duplicate api_keys.id: {}", @@ -535,70 +711,30 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { record.key_hash ))); } - - let template = index - .by_api_key_id - .values() - .find(|snapshot| snapshot.user_id == record.user_id) - .cloned(); - let snapshot = if let Some(template) = template { - StoredAuthApiKeySnapshot { - api_key_id: record.api_key_id.clone(), - api_key_name: record.name.clone(), - api_key_is_active: record.is_active, - api_key_is_locked: false, - api_key_is_standalone: false, - api_key_rate_limit: Some(record.rate_limit), - api_key_concurrent_limit: record.concurrent_limit, - api_key_expires_at_unix_secs: record.expires_at_unix_secs, - api_key_allowed_providers: record.allowed_providers.clone(), - api_key_allowed_api_formats: record.allowed_api_formats.clone(), - api_key_allowed_models: record.allowed_models.clone(), - api_key_ip_rules: record.ip_rules.clone(), - ..template - } - } else { - StoredAuthApiKeySnapshot::new( - record.user_id.clone(), - format!( - "user-{}", - &record.user_id.chars().take(8).collect::() - ), - None, - "user".to_string(), - "local".to_string(), - true, - false, - None, - None, - None, - record.api_key_id.clone(), - record.name.clone(), - record.is_active, - false, - false, - Some(record.rate_limit), - record.concurrent_limit, - record.expires_at_unix_secs.map(|value| value as i64), - record - .allowed_providers - .as_ref() - .map(|value| serde_json::json!(value)), - record - .allowed_api_formats - .as_ref() - .map(|value| serde_json::json!(value)), - record - .allowed_models - .as_ref() - .map(|value| serde_json::json!(value)), - )? - .with_api_key_ip_rules( - record - .ip_rules - .as_ref() - .map(|value| serde_json::json!(value)), - )? + let snapshot = StoredAuthApiKeySnapshot { + user_id: owner.user_id, + username: owner.username, + email: owner.email, + user_role: owner.user_role, + user_auth_source: owner.user_auth_source, + user_is_active: owner.user_is_active, + user_is_deleted: owner.user_is_deleted, + user_rate_limit: owner.user_rate_limit, + user_allowed_providers: owner.user_allowed_providers, + user_allowed_api_formats: owner.user_allowed_api_formats, + user_allowed_models: owner.user_allowed_models, + api_key_id: record.api_key_id.clone(), + api_key_name: record.name.clone(), + api_key_is_active: record.is_active, + api_key_is_locked: false, + api_key_is_standalone: false, + api_key_rate_limit: Some(record.rate_limit), + api_key_concurrent_limit: record.concurrent_limit, + api_key_expires_at_unix_secs: record.expires_at_unix_secs, + api_key_allowed_providers: record.allowed_providers.clone(), + api_key_allowed_api_formats: record.allowed_api_formats.clone(), + api_key_allowed_models: record.allowed_models.clone(), + api_key_ip_rules: record.ip_rules.clone(), }; let now_unix_secs = current_unix_secs() as i64; @@ -624,10 +760,13 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { record.concurrent_limit, record.force_capabilities, record.is_active, - record.expires_at_unix_secs.map(|value| value as i64), + record + .expires_at_unix_secs + .map(|value| i64_from_u64(value, "api_keys.expires_at")) + .transpose()?, record.auto_delete_on_expiry, - record.total_requests as i64, - record.total_tokens as i64, + i64_from_u64(record.total_requests, "api_keys.total_requests")?, + i64_from_u64(record.total_tokens, "api_keys.total_tokens")?, record.total_cost_usd, false, )? @@ -637,6 +776,7 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { .as_ref() .map(|value| serde_json::json!(value)), )? + .with_feature_settings(record.feature_settings) .with_activity_timestamps(None, Some(now_unix_secs), Some(now_unix_secs))?; index @@ -715,7 +855,10 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { true, record.rate_limit, record.concurrent_limit, - record.expires_at_unix_secs.map(|value| value as i64), + record + .expires_at_unix_secs + .map(|value| i64_from_u64(value, "api_keys.expires_at")) + .transpose()?, record .allowed_providers .as_ref() @@ -760,10 +903,13 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { record.concurrent_limit, record.force_capabilities, record.is_active, - record.expires_at_unix_secs.map(|value| value as i64), + record + .expires_at_unix_secs + .map(|value| i64_from_u64(value, "api_keys.expires_at")) + .transpose()?, record.auto_delete_on_expiry, - record.total_requests as i64, - record.total_tokens as i64, + i64_from_u64(record.total_requests, "api_keys.total_requests")?, + i64_from_u64(record.total_tokens, "api_keys.total_tokens")?, record.total_cost_usd, true, )? @@ -787,6 +933,28 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(Some(export)) } + async fn compare_and_swap_api_key_ciphertext( + &self, + mutation: &CompareAndSwapAuthApiKeyCiphertext, + ) -> Result { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(current) = index.export_by_api_key_id.get_mut(&mutation.api_key_id) else { + return Ok(false); + }; + if current.user_id != mutation.user_id + || current.key_hash != mutation.key_hash + || current.is_standalone != mutation.is_standalone + || current.key_encrypted.as_deref() != Some(mutation.expected_key_encrypted.as_str()) + { + return Ok(false); + } + current.key_encrypted = Some(mutation.key_encrypted.clone()); + Ok(true) + } + async fn update_user_api_key_basic( &self, record: UpdateUserApiKeyBasicRecord, @@ -801,28 +969,33 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { if snapshot.user_id != record.user_id || snapshot.api_key_is_standalone { return Ok(None); } - if let Some(name) = record.name { - if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { - snapshot.api_key_name = Some(name.clone()); - } + if record.key_encrypted_present { if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { - export.name = Some(name); + export.key_encrypted = record.key_encrypted.clone(); } } - if let Some(rate_limit) = record.rate_limit { + if record.name_present { if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { - snapshot.api_key_rate_limit = Some(rate_limit); + snapshot.api_key_name = record.name.clone(); } if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { - export.rate_limit = Some(rate_limit); + export.name = record.name.clone(); } } - if let Some(concurrent_limit) = record.concurrent_limit { + if record.rate_limit_present { if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { - snapshot.api_key_concurrent_limit = Some(concurrent_limit); + snapshot.api_key_rate_limit = record.rate_limit; } if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { - export.concurrent_limit = Some(concurrent_limit); + export.rate_limit = record.rate_limit; + } + } + if record.concurrent_limit_present { + if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { + snapshot.api_key_concurrent_limit = record.concurrent_limit; + } + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.concurrent_limit = record.concurrent_limit; } } if let Some(ip_rules) = record.ip_rules { @@ -833,6 +1006,79 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { export.ip_rules = ip_rules; } } + if let Some(feature_settings) = record.feature_settings { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.feature_settings = match feature_settings { + Some(serde_json::Value::Null) | None => None, + Some(value) => Some(value), + }; + } + } + Ok(index.export_by_api_key_id.get(&record.api_key_id).cloned()) + } + + async fn update_user_api_key_basic_if_unlocked( + &self, + record: UpdateUserApiKeyBasicRecord, + ) -> Result, DataLayerError> { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(&record.api_key_id) else { + return Ok(None); + }; + if snapshot.user_id != record.user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(None); + } + if record.key_encrypted_present { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.key_encrypted = record.key_encrypted.clone(); + } + } + if record.name_present { + if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { + snapshot.api_key_name = record.name.clone(); + } + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.name = record.name.clone(); + } + } + if record.rate_limit_present { + if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { + snapshot.api_key_rate_limit = record.rate_limit; + } + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.rate_limit = record.rate_limit; + } + } + if record.concurrent_limit_present { + if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { + snapshot.api_key_concurrent_limit = record.concurrent_limit; + } + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.concurrent_limit = record.concurrent_limit; + } + } + if let Some(ip_rules) = record.ip_rules { + if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { + snapshot.api_key_ip_rules = ip_rules.clone(); + } + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.ip_rules = ip_rules; + } + } + if let Some(feature_settings) = record.feature_settings { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.feature_settings = match feature_settings { + Some(serde_json::Value::Null) | None => None, + Some(value) => Some(value), + }; + } + } Ok(index.export_by_api_key_id.get(&record.api_key_id).cloned()) } @@ -850,12 +1096,22 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { if !snapshot.api_key_is_standalone { return Ok(None); } - if let Some(name) = record.name { + if record.key_encrypted_present { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.key_encrypted = record.key_encrypted.clone(); + } + } + if record.name_present { if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) { - snapshot.api_key_name = Some(name.clone()); + snapshot.api_key_name = record.name.clone(); } if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { - export.name = Some(name); + export.name = record.name.clone(); + } + } + if let Some(force_capabilities) = record.force_capabilities { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + export.force_capabilities = force_capabilities; } } if record.rate_limit_present { @@ -922,6 +1178,45 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(index.export_by_api_key_id.get(&record.api_key_id).cloned()) } + async fn restore_api_key_if_matches( + &self, + expected: &StoredAuthApiKeyExportRecord, + restored: &StoredAuthApiKeyExportRecord, + ) -> Result { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(current) = index.export_by_api_key_id.get(&expected.api_key_id) else { + return Ok(false); + }; + // Immutable identity must never be changed by compensation. The complete exported row + // comparison below is the in-memory equivalent of the database CAS predicate. + if current != expected + || restored.api_key_id != expected.api_key_id + || restored.user_id != expected.user_id + || restored.key_hash != expected.key_hash + || restored.is_standalone != expected.is_standalone + { + return Ok(false); + } + index + .export_by_api_key_id + .insert(expected.api_key_id.clone(), restored.clone()); + if let Some(snapshot) = index.by_api_key_id.get_mut(&expected.api_key_id) { + snapshot.api_key_name = restored.name.clone(); + snapshot.api_key_is_active = restored.is_active; + snapshot.api_key_rate_limit = restored.rate_limit; + snapshot.api_key_concurrent_limit = restored.concurrent_limit; + snapshot.api_key_expires_at_unix_secs = restored.expires_at_unix_secs; + snapshot.api_key_allowed_providers = restored.allowed_providers.clone(); + snapshot.api_key_allowed_api_formats = restored.allowed_api_formats.clone(); + snapshot.api_key_allowed_models = restored.allowed_models.clone(); + snapshot.api_key_ip_rules = restored.ip_rules.clone(); + } + Ok(true) + } + async fn set_user_api_key_active( &self, user_id: &str, @@ -947,6 +1242,34 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(index.export_by_api_key_id.get(api_key_id).cloned()) } + async fn set_user_api_key_active_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(api_key_id) else { + return Ok(None); + }; + if snapshot.user_id != user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(None); + } + if let Some(snapshot) = index.by_api_key_id.get_mut(api_key_id) { + snapshot.api_key_is_active = is_active; + } + if let Some(export) = index.export_by_api_key_id.get_mut(api_key_id) { + export.is_active = is_active; + } + Ok(index.export_by_api_key_id.get(api_key_id).cloned()) + } + async fn set_standalone_api_key_active( &self, api_key_id: &str, @@ -1018,6 +1341,34 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(index.export_by_api_key_id.get(api_key_id).cloned()) } + async fn set_user_api_key_allowed_providers_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + allowed_providers: Option>, + ) -> Result, DataLayerError> { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(api_key_id) else { + return Ok(None); + }; + if snapshot.user_id != user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(None); + } + if let Some(snapshot) = index.by_api_key_id.get_mut(api_key_id) { + snapshot.api_key_allowed_providers = allowed_providers.clone(); + } + if let Some(export) = index.export_by_api_key_id.get_mut(api_key_id) { + export.allowed_providers = allowed_providers; + } + Ok(index.export_by_api_key_id.get(api_key_id).cloned()) + } + async fn set_user_api_key_force_capabilities( &self, user_id: &str, @@ -1041,6 +1392,32 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(Some(export.clone())) } + async fn set_user_api_key_force_capabilities_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + force_capabilities: Option, + ) -> Result, DataLayerError> { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(api_key_id) else { + return Ok(None); + }; + if snapshot.user_id != user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(None); + } + let Some(export) = index.export_by_api_key_id.get_mut(api_key_id) else { + return Ok(None); + }; + export.force_capabilities = force_capabilities; + Ok(Some(export.clone())) + } + async fn set_user_api_key_feature_settings( &self, user_id: &str, @@ -1067,6 +1444,35 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { Ok(Some(export.clone())) } + async fn set_user_api_key_feature_settings_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + feature_settings: Option, + ) -> Result, DataLayerError> { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(api_key_id) else { + return Ok(None); + }; + if snapshot.user_id != user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(None); + } + let Some(export) = index.export_by_api_key_id.get_mut(api_key_id) else { + return Ok(None); + }; + export.feature_settings = match feature_settings { + Some(serde_json::Value::Null) | None => None, + Some(value) => Some(value), + }; + Ok(Some(export.clone())) + } + async fn set_api_key_usage_totals( &self, api_key_id: &str, @@ -1074,6 +1480,13 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { total_tokens: u64, total_cost_usd: f64, ) -> Result, DataLayerError> { + i64_from_u64(total_requests, "api_keys.total_requests")?; + i64_from_u64(total_tokens, "api_keys.total_tokens")?; + if !total_cost_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "api_keys.total_cost_usd is not finite".to_string(), + )); + } let mut index = self .index .write() @@ -1102,10 +1515,29 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { if snapshot.user_id != user_id || snapshot.api_key_is_standalone { return Ok(false); } - index.by_api_key_id.remove(api_key_id); - index.export_by_api_key_id.remove(api_key_id); - index.by_key_hash.retain(|_, value| value != api_key_id); - index.touch_counts.remove(api_key_id); + Self::remove_api_key(&mut index, api_key_id); + Ok(true) + } + + async fn delete_user_api_key_if_unlocked( + &self, + user_id: &str, + api_key_id: &str, + ) -> Result { + let mut index = self + .index + .write() + .expect("auth api key snapshot repository lock"); + let Some(snapshot) = index.by_api_key_id.get(api_key_id) else { + return Ok(false); + }; + if snapshot.user_id != user_id + || snapshot.api_key_is_standalone + || snapshot.api_key_is_locked + { + return Ok(false); + } + Self::remove_api_key(&mut index, api_key_id); Ok(true) } @@ -1120,10 +1552,7 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { if !snapshot.api_key_is_standalone { return Ok(false); } - index.by_api_key_id.remove(api_key_id); - index.export_by_api_key_id.remove(api_key_id); - index.by_key_hash.retain(|_, value| value != api_key_id); - index.touch_counts.remove(api_key_id); + Self::remove_api_key(&mut index, api_key_id); Ok(true) } @@ -1158,9 +1587,11 @@ mod tests { use super::InMemoryAuthApiKeySnapshotRepository; use crate::repository::auth::{ AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, + CompareAndSwapAuthApiKeyCiphertext, CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; + use crate::repository::users::StoredUserAuthRecord; fn sample_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot { StoredAuthApiKeySnapshot::new( @@ -1189,6 +1620,487 @@ mod tests { .expect("snapshot should build") } + fn sample_create_user_api_key_record( + user_id: &str, + api_key_id: &str, + ) -> CreateUserApiKeyRecord { + CreateUserApiKeyRecord { + user_id: user_id.to_string(), + api_key_id: api_key_id.to_string(), + key_hash: format!("hash-{api_key_id}"), + key_encrypted: Some(format!("encrypted-{api_key_id}")), + name: Some("Created".to_string()), + allowed_providers: None, + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + rate_limit: 0, + concurrent_limit: None, + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + } + } + + fn sample_authoritative_user( + user_id: &str, + is_active: bool, + is_deleted: bool, + ) -> StoredUserAuthRecord { + StoredUserAuthRecord::new( + user_id.to_string(), + Some(format!("{user_id}@example.com")), + true, + format!("owner-{user_id}"), + Some("server-managed-password-hash".to_string()), + "admin".to_string(), + "oauth".to_string(), + Some(serde_json::json!(["openai"])), + Some(serde_json::json!(["openai:chat"])), + Some(serde_json::json!(["gpt-5"])), + is_active, + is_deleted, + None, + None, + ) + .expect("authoritative user should build") + .with_security_version(37) + .expect("security version should be valid") + } + + #[tokio::test] + async fn authoritative_owner_sync_allows_first_key_without_synthesizing_owner_fields() { + let repository = InMemoryAuthApiKeySnapshotRepository::default(); + let user = sample_authoritative_user("authoritative-user", true, false); + repository + .synchronize_user_api_key_owner_for_tests(&user) + .await + .expect("authoritative owner should synchronize"); + + let mut record = + sample_create_user_api_key_record("authoritative-user", "authoritative-key"); + record.allowed_providers = Some(vec!["anthropic".to_string()]); + repository + .create_user_api_key(record) + .await + .expect("first key creation should resolve") + .expect("active authoritative owner should allow its first key"); + + let snapshot = repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("authoritative-key")) + .await + .expect("created snapshot lookup should resolve") + .expect("created snapshot should exist"); + assert_eq!(snapshot.username, user.username); + assert_eq!(snapshot.email, user.email); + assert_eq!(snapshot.user_role, "admin"); + assert_eq!(snapshot.user_auth_source, "oauth"); + assert!(snapshot.user_is_active); + assert!(!snapshot.user_is_deleted); + assert_eq!( + snapshot.user_allowed_providers, + Some(vec!["openai".to_string()]) + ); + assert_eq!( + snapshot.api_key_allowed_providers, + Some(vec!["anthropic".to_string()]) + ); + assert_eq!(user.security_version, 37); + } + + #[tokio::test] + async fn api_key_ciphertext_cas_fences_complete_identity_and_exact_old_value() { + let repository = InMemoryAuthApiKeySnapshotRepository::default(); + repository + .synchronize_user_api_key_owner_for_tests(&sample_authoritative_user( + "cipher-owner", + true, + false, + )) + .await + .expect("owner should synchronize"); + repository + .create_user_api_key(sample_create_user_api_key_record( + "cipher-owner", + "cipher-key", + )) + .await + .expect("key creation should succeed") + .expect("key should be created"); + + let expected = CompareAndSwapAuthApiKeyCiphertext { + user_id: "cipher-owner".to_string(), + api_key_id: "cipher-key".to_string(), + key_hash: "hash-cipher-key".to_string(), + is_standalone: false, + expected_key_encrypted: "encrypted-cipher-key".to_string(), + key_encrypted: "bound-ciphertext".to_string(), + }; + for mutation in [ + CompareAndSwapAuthApiKeyCiphertext { + user_id: "other-owner".to_string(), + ..expected.clone() + }, + CompareAndSwapAuthApiKeyCiphertext { + key_hash: "other-hash".to_string(), + ..expected.clone() + }, + CompareAndSwapAuthApiKeyCiphertext { + is_standalone: true, + ..expected.clone() + }, + CompareAndSwapAuthApiKeyCiphertext { + expected_key_encrypted: "ENCRYPTED-cipher-key".to_string(), + ..expected.clone() + }, + ] { + assert!(!repository + .compare_and_swap_api_key_ciphertext(&mutation) + .await + .expect("CAS should execute")); + } + assert!(repository + .compare_and_swap_api_key_ciphertext(&expected) + .await + .expect("matching CAS should execute")); + assert_eq!( + repository + .list_export_api_keys_by_ids(&["cipher-key".to_string()]) + .await + .expect("key should reload")[0] + .key_encrypted + .as_deref(), + Some("bound-ciphertext") + ); + assert!(!repository + .compare_and_swap_api_key_ciphertext(&expected) + .await + .expect("stale CAS should execute")); + } + + #[tokio::test] + async fn authoritative_owner_sync_preserves_inactive_state_and_rejects_deleted_owner() { + let repository = InMemoryAuthApiKeySnapshotRepository::default(); + let inactive = sample_authoritative_user("inactive-user", false, false); + repository + .synchronize_user_api_key_owner_for_tests(&inactive) + .await + .expect("inactive owner should synchronize without privilege elevation"); + repository + .create_user_api_key(sample_create_user_api_key_record( + "inactive-user", + "inactive-key", + )) + .await + .expect("inactive owner creation should resolve") + .expect("low-level repository should retain database write parity"); + let inactive_snapshot = repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("inactive-key")) + .await + .expect("inactive key lookup should resolve") + .expect("inactive key should be stored"); + assert!(!inactive_snapshot.user_is_active); + assert!(!inactive_snapshot.is_currently_usable(0)); + + let deleted = sample_authoritative_user("deleted-user", true, true); + repository + .synchronize_user_api_key_owner_for_tests(&deleted) + .await + .expect("deleted owner tombstone should synchronize"); + assert!(repository + .create_user_api_key(sample_create_user_api_key_record( + "deleted-user", + "deleted-key", + )) + .await + .expect("deleted owner creation should resolve") + .is_none()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("deleted-key")) + .await + .expect("rejected deleted key lookup should resolve") + .is_none()); + } + + #[tokio::test] + async fn create_user_api_key_persists_feature_settings_in_initial_record() { + let repository = InMemoryAuthApiKeySnapshotRepository::default() + .with_owner_snapshots([sample_snapshot("owner-fixture", "user-1")]); + let mut record = sample_create_user_api_key_record("user-1", "key-created"); + record.feature_settings = Some(serde_json::json!({"compact": true})); + let created = repository + .create_user_api_key(record) + .await + .expect("create should succeed") + .expect("created key should be returned"); + + assert_eq!( + created.feature_settings, + Some(serde_json::json!({"compact": true})) + ); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("owner-fixture")) + .await + .expect("owner fixture lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn create_user_api_key_rejects_unknown_deleted_and_conflicted_owners() { + let mut deleted_owner = sample_snapshot("deleted-owner-fixture", "deleted-user"); + deleted_owner.user_is_deleted = true; + let conflicted_owner = sample_snapshot("conflicted-owner-fixture-a", "conflicted-user"); + let mut conflicting_owner = + sample_snapshot("conflicted-owner-fixture-b", "conflicted-user"); + conflicting_owner.username = "different-owner-state".to_string(); + let cases = [ + ( + "unknown", + "unknown-user", + InMemoryAuthApiKeySnapshotRepository::default(), + ), + ( + "deleted", + "deleted-user", + InMemoryAuthApiKeySnapshotRepository::default() + .with_owner_snapshots([deleted_owner]), + ), + ( + "conflicted", + "conflicted-user", + InMemoryAuthApiKeySnapshotRepository::default() + .with_owner_snapshots([conflicted_owner, conflicting_owner]), + ), + ]; + + for (case, user_id, repository) in cases { + let api_key_id = format!("key-{case}"); + let key_hash = format!("hash-{api_key_id}"); + let record = sample_create_user_api_key_record(user_id, &api_key_id); + assert!(repository + .create_user_api_key(record) + .await + .expect("fail-closed creation should resolve") + .is_none()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId(&api_key_id)) + .await + .expect("rejected key lookup should resolve") + .is_none()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::KeyHash(&key_hash)) + .await + .expect("rejected hash lookup should resolve") + .is_none()); + assert!(repository + .list_export_api_keys_by_user_ids(&[user_id.to_string()]) + .await + .expect("rejected exports should list") + .is_empty()); + } + } + + #[tokio::test] + async fn create_user_api_key_preserves_disabled_owner_state() { + let mut disabled_owner = sample_snapshot("disabled-owner-fixture", "disabled-user"); + disabled_owner.user_is_active = false; + let repository = + InMemoryAuthApiKeySnapshotRepository::default().with_owner_snapshots([disabled_owner]); + + repository + .create_user_api_key(sample_create_user_api_key_record( + "disabled-user", + "key-disabled", + )) + .await + .expect("disabled-owner creation should resolve") + .expect("non-deleted disabled owner should retain database parity"); + + let snapshot = repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-disabled")) + .await + .expect("created key lookup should resolve") + .expect("created key snapshot should exist"); + assert!(!snapshot.user_is_active); + assert!(!snapshot.is_currently_usable(0)); + } + + #[tokio::test] + async fn invalid_owner_cannot_probe_duplicate_credential_indexes() { + let repository = InMemoryAuthApiKeySnapshotRepository::seed([( + Some("hash-key-existing".to_string()), + sample_snapshot("key-existing", "trusted-user"), + )]); + + assert!(repository + .create_user_api_key(sample_create_user_api_key_record( + "unknown-user", + "key-existing", + )) + .await + .expect("invalid owner should fail closed before duplicate checks") + .is_none()); + + let existing = repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-existing")) + .await + .expect("existing key lookup should resolve") + .expect("existing key should remain unchanged"); + assert_eq!(existing.user_id, "trusted-user"); + } + + #[tokio::test] + async fn usage_total_replacement_rejects_unpersistable_values_without_mutation() { + let repository = InMemoryAuthApiKeySnapshotRepository::seed([( + Some("hash-key-existing".to_string()), + sample_snapshot("key-existing", "trusted-user"), + )]); + let before = repository + .list_export_api_keys_by_ids(&["key-existing".to_string()]) + .await + .expect("existing export should load") + .pop() + .expect("existing export should exist"); + + for (total_requests, total_tokens, total_cost_usd) in [ + (u64::MAX, 0, 0.0), + (0, u64::MAX, 0.0), + (0, 0, f64::NAN), + (0, 0, f64::INFINITY), + ] { + assert!(matches!( + repository + .set_api_key_usage_totals( + "key-existing", + total_requests, + total_tokens, + total_cost_usd, + ) + .await, + Err(crate::DataLayerError::InvalidInput(_)) + )); + } + + let after = repository + .list_export_api_keys_by_ids(&["key-existing".to_string()]) + .await + .expect("existing export should reload") + .pop() + .expect("existing export should remain"); + assert_eq!(after, before); + } + + #[tokio::test] + async fn unlocked_user_mutations_atomically_reject_locked_keys() { + let mut snapshot = sample_snapshot("key-locked", "user-1"); + snapshot.api_key_is_locked = true; + let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some("hash-locked".to_string()), + snapshot, + )]); + + assert!(repository + .update_user_api_key_basic_if_unlocked(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-locked".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: Some("must-not-change".to_string()), + name_present: true, + rate_limit: None, + rate_limit_present: false, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: Some(Some(serde_json::json!({"must_not_change": true}))), + }) + .await + .expect("locked basic update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_active_if_unlocked("user-1", "key-locked", false) + .await + .expect("locked status update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_allowed_providers_if_unlocked( + "user-1", + "key-locked", + Some(vec!["must-not-change".to_string()]), + ) + .await + .expect("locked provider update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_force_capabilities_if_unlocked( + "user-1", + "key-locked", + Some(serde_json::json!({"must_not_change": true})), + ) + .await + .expect("locked capability update should resolve") + .is_none()); + assert!(repository + .set_user_api_key_feature_settings_if_unlocked( + "user-1", + "key-locked", + Some(serde_json::json!({"must_not_change": true})), + ) + .await + .expect("locked feature update should resolve") + .is_none()); + assert!(!repository + .delete_user_api_key_if_unlocked("user-1", "key-locked") + .await + .expect("locked deletion should resolve")); + + let unchanged = repository + .list_export_api_keys_by_ids(&["key-locked".to_string()]) + .await + .expect("locked key should reload") + .pop() + .expect("locked key should remain"); + assert_ne!(unchanged.name.as_deref(), Some("must-not-change")); + assert!(unchanged.is_active); + assert!(unchanged.feature_settings.is_none()); + + // Administrative repository operations deliberately retain authority + // over locked keys; only the self-service variants enforce the fence. + let admin_updated = repository + .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-locked".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: Some("admin-change".to_string()), + name_present: true, + rate_limit: None, + rate_limit_present: false, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: Some(Some(serde_json::json!({"admin": true}))), + }) + .await + .expect("administrator update should resolve") + .expect("administrator may update a locked key"); + assert_eq!(admin_updated.name.as_deref(), Some("admin-change")); + assert_eq!( + admin_updated.feature_settings, + Some(serde_json::json!({"admin": true})) + ); + assert!(repository + .set_user_api_key_active("user-1", "key-locked", false) + .await + .expect("admin status update should resolve") + .is_some()); + } + #[tokio::test] async fn reads_auth_snapshot_by_all_supported_keys() { let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( @@ -1240,6 +2152,57 @@ mod tests { .expect("missing touch should succeed")); } + #[tokio::test] + async fn delete_is_owner_scoped_and_removes_all_credential_indexes() { + let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some("hash-1".to_string()), + sample_snapshot("key-1", "user-1"), + )]); + + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::KeyHash("hash-1")) + .await + .expect("hash lookup should succeed") + .is_some()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-1")) + .await + .expect("id lookup should succeed") + .is_some()); + assert!(repository + .touch_last_used_at("key-1") + .await + .expect("touch should succeed")); + + assert!(!repository + .delete_user_api_key("other-user", "key-1") + .await + .expect("wrong-owner delete should resolve")); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-1")) + .await + .expect("wrong-owner delete must preserve the key") + .is_some()); + + assert!(repository + .delete_user_api_key("user-1", "key-1") + .await + .expect("owner-scoped delete should succeed")); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-1")) + .await + .expect("deleted id lookup should resolve") + .is_none()); + assert!(repository + .find_api_key_snapshot(AuthApiKeyLookupKey::KeyHash("hash-1")) + .await + .expect("deleted hash lookup should resolve") + .is_none()); + assert_eq!(repository.touch_count("key-1"), 0); + assert_eq!(repository.snapshot_lookup_count("key-1"), 1); + assert_eq!(repository.key_hash_lookup_count("hash-1"), 1); + } + #[tokio::test] async fn lists_export_records_for_user_bound_and_standalone_keys() { let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![ @@ -1356,10 +2319,16 @@ mod tests { .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { user_id: "user-1".to_string(), api_key_id: "key-1".to_string(), + key_encrypted: None, + key_encrypted_present: false, name: None, + name_present: false, rate_limit: None, + rate_limit_present: false, concurrent_limit: Some(11), + concurrent_limit_present: true, ip_rules: None, + feature_settings: None, }) .await .expect("update should succeed") @@ -1374,6 +2343,56 @@ mod tests { assert_eq!(snapshot.api_key_concurrent_limit, Some(11)); } + #[tokio::test] + async fn update_user_api_key_basic_restores_nullable_values_and_zero() { + let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some("hash-1".to_string()), + sample_snapshot("key-1", "user-1"), + )]); + + let cleared = repository + .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-1".to_string(), + key_encrypted: None, + key_encrypted_present: true, + name: None, + name_present: true, + rate_limit: None, + rate_limit_present: true, + concurrent_limit: None, + concurrent_limit_present: true, + ip_rules: None, + feature_settings: None, + }) + .await + .expect("nullable values should clear") + .expect("record should exist"); + assert!(cleared.name.is_none()); + assert!(cleared.rate_limit.is_none()); + assert!(cleared.concurrent_limit.is_none()); + + let zero = repository + .update_user_api_key_basic(UpdateUserApiKeyBasicRecord { + user_id: "user-1".to_string(), + api_key_id: "key-1".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: None, + name_present: false, + rate_limit: Some(0), + rate_limit_present: true, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: None, + }) + .await + .expect("zero rate limit should persist") + .expect("record should exist"); + assert_eq!(zero.rate_limit, Some(0)); + } + #[tokio::test] async fn update_standalone_api_key_basic_updates_concurrent_limit_when_present() { let mut standalone = sample_snapshot("key-standalone", "admin-1"); @@ -1386,7 +2405,11 @@ mod tests { let updated = repository .update_standalone_api_key_basic(UpdateStandaloneApiKeyBasicRecord { api_key_id: "key-standalone".to_string(), + key_encrypted: None, + key_encrypted_present: false, name: None, + name_present: false, + force_capabilities: None, rate_limit_present: false, rate_limit: None, concurrent_limit_present: true, @@ -1412,4 +2435,115 @@ mod tests { .expect("snapshot should exist"); assert_eq!(snapshot.api_key_concurrent_limit, Some(13)); } + + #[tokio::test] + async fn restores_standalone_nullable_fields_and_force_capabilities_atomically() { + let mut standalone = sample_snapshot("key-standalone", "admin-1"); + standalone.api_key_is_standalone = true; + let before = StoredAuthApiKeyExportRecord::new( + "admin-1".to_string(), + "key-standalone".to_string(), + "hash-standalone".to_string(), + None, + None, + Some(serde_json::json!(["openai"])), + Some(serde_json::json!(["openai:chat"])), + Some(serde_json::json!(["gpt-4.1"])), + Some(60), + Some(5), + None, + true, + Some(200), + false, + 2, + 25, + 0.25, + true, + ) + .expect("before export should build") + .with_ip_rules(Some(serde_json::json!(["10.0.0.0/24"]))) + .expect("before ip rules should build") + .with_feature_settings(Some(serde_json::json!({"compact": true}))); + let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some("hash-standalone".to_string()), + standalone, + )]) + .with_export_records([before.clone()]); + + let after = repository + .update_standalone_api_key_basic(UpdateStandaloneApiKeyBasicRecord { + api_key_id: "key-standalone".to_string(), + key_encrypted: Some("enc-after".to_string()), + key_encrypted_present: true, + name: Some("after".to_string()), + name_present: true, + force_capabilities: Some(Some(serde_json::json!({"vision": true}))), + rate_limit_present: true, + rate_limit: Some(99), + concurrent_limit_present: true, + concurrent_limit: Some(8), + allowed_providers: Some(Some(vec!["anthropic".to_string()])), + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + expires_at_present: false, + expires_at_unix_secs: None, + auto_delete_on_expiry_present: false, + auto_delete_on_expiry: false, + }) + .await + .expect("after update should succeed") + .expect("after record should exist"); + assert!(repository + .restore_api_key_if_matches(&after, &before) + .await + .expect("restore should succeed")); + let restored = repository + .list_export_api_keys_by_ids(&["key-standalone".to_string()]) + .await + .expect("restored export should load") + .pop() + .expect("restored key should exist"); + assert_eq!(restored, before); + + // A changed post-state must make the CAS fail without touching the newer value. + let concurrent = repository + .update_standalone_api_key_basic(UpdateStandaloneApiKeyBasicRecord { + api_key_id: "key-standalone".to_string(), + key_encrypted: None, + key_encrypted_present: false, + name: Some("concurrent".to_string()), + name_present: true, + force_capabilities: Some(None), + rate_limit_present: false, + rate_limit: None, + concurrent_limit_present: false, + concurrent_limit: None, + allowed_providers: None, + allowed_api_formats: None, + allowed_models: None, + ip_rules: None, + expires_at_present: false, + expires_at_unix_secs: None, + auto_delete_on_expiry_present: false, + auto_delete_on_expiry: false, + }) + .await + .expect("concurrent update should succeed") + .expect("concurrent record should exist"); + assert!(!repository + .restore_api_key_if_matches(&after, &before) + .await + .expect("conflicting restore should return false")); + assert_eq!( + repository + .list_export_api_keys_by_ids(&["key-standalone".to_string()]) + .await + .expect("current export should load") + .pop() + .expect("current key should exist") + .name, + concurrent.name + ); + } } diff --git a/crates/aether-data/runtime/src/repository/auth/mod.rs b/crates/aether-data/runtime/src/repository/auth/mod.rs index 67cd6ea85..46c9081c2 100644 --- a/crates/aether-data/runtime/src/repository/auth/mod.rs +++ b/crates/aether-data/runtime/src/repository/auth/mod.rs @@ -4,8 +4,8 @@ pub use aether_data_contracts::repository::auth::{ read_resolved_auth_api_key_snapshot, read_resolved_auth_api_key_snapshot_by_key_hash, read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, AuthRepository, - CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, ResolvedAuthApiKeySnapshot, - ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery, + CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, + ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; diff --git a/crates/aether-data/runtime/src/repository/auth_modules/memory.rs b/crates/aether-data/runtime/src/repository/auth_modules/memory.rs index 7891c3486..a8a19f972 100644 --- a/crates/aether-data/runtime/src/repository/auth_modules/memory.rs +++ b/crates/aether-data/runtime/src/repository/auth_modules/memory.rs @@ -3,8 +3,8 @@ use std::sync::RwLock; use async_trait::async_trait; use super::{ - AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig, - StoredOAuthProviderModuleConfig, + AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult, + LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig, }; use crate::DataLayerError; @@ -49,15 +49,73 @@ impl AuthModuleReadRepository for InMemoryAuthModuleReadRepository { #[async_trait] impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository { - async fn upsert_ldap_config( + async fn compare_and_swap_ldap_config( &self, - config: &StoredLdapModuleConfig, - ) -> Result, DataLayerError> { - self.ldap_config + expected: Option<&StoredLdapModuleConfig>, + replacement: &StoredLdapModuleConfig, + bind_password_update: &LdapBindPasswordUpdate, + ) -> Result { + let mut config = self + .ldap_config .write() - .expect("auth module ldap repository lock") - .replace(config.clone()); - Ok(Some(config.clone())) + .expect("auth module ldap repository lock"); + if config.as_ref() != expected { + return Ok(CompareAndSwapLdapConfigResult::Conflict); + } + + let bind_password_encrypted = match bind_password_update { + LdapBindPasswordUpdate::Preserve => expected + .ok_or_else(|| { + DataLayerError::InvalidConfiguration( + "LDAP bind password cannot be preserved while creating the singleton" + .to_string(), + ) + })? + .bind_password_encrypted + .clone(), + LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()), + LdapBindPasswordUpdate::Clear => None, + }; + let persisted = StoredLdapModuleConfig { + bind_password_encrypted, + ..replacement.clone() + }; + *config = Some(persisted.clone()); + Ok(CompareAndSwapLdapConfigResult::Applied(persisted)) + } + + async fn delete_ldap_config_if_matches( + &self, + expected: &StoredLdapModuleConfig, + ) -> Result { + let mut config = self + .ldap_config + .write() + .expect("auth module ldap repository lock"); + if config.as_ref() != Some(expected) { + return Ok(false); + } + config.take(); + Ok(true) + } + + async fn compare_and_swap_ldap_bind_password( + &self, + expected: &str, + replacement: &str, + ) -> Result { + let mut config = self + .ldap_config + .write() + .expect("auth module ldap repository lock"); + let Some(config) = config.as_mut() else { + return Ok(false); + }; + if config.bind_password_encrypted.as_deref() != Some(expected) { + return Ok(false); + } + config.bind_password_encrypted = Some(replacement.to_string()); + Ok(true) } } @@ -65,9 +123,27 @@ impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository { mod tests { use super::InMemoryAuthModuleReadRepository; use crate::repository::auth_modules::{ - AuthModuleReadRepository, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig, + AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult, + LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig, }; + fn ldap_config() -> StoredLdapModuleConfig { + StoredLdapModuleConfig { + server_url: "ldaps://ldap.example.com".to_string(), + bind_dn: "cn=admin,dc=example,dc=com".to_string(), + bind_password_encrypted: Some("encrypted-password".to_string()), + base_dn: "dc=example,dc=com".to_string(), + user_search_filter: Some("(uid={username})".to_string()), + username_attr: Some("uid".to_string()), + email_attr: Some("mail".to_string()), + display_name_attr: Some("displayName".to_string()), + is_enabled: true, + is_exclusive: false, + use_starttls: true, + connect_timeout: Some(10), + } + } + #[tokio::test] async fn reads_seeded_auth_module_configs() { let repository = InMemoryAuthModuleReadRepository::seed( @@ -79,20 +155,7 @@ mod tests { "https://example.com/callback".to_string(), ) .expect("oauth provider should build")], - Some(StoredLdapModuleConfig { - server_url: "ldaps://ldap.example.com".to_string(), - bind_dn: "cn=admin,dc=example,dc=com".to_string(), - bind_password_encrypted: Some("encrypted-password".to_string()), - base_dn: "dc=example,dc=com".to_string(), - user_search_filter: Some("(uid={username})".to_string()), - username_attr: Some("uid".to_string()), - email_attr: Some("mail".to_string()), - display_name_attr: Some("displayName".to_string()), - is_enabled: true, - is_exclusive: false, - use_starttls: true, - connect_timeout: Some(10), - }), + Some(ldap_config()), ); let oauth = repository @@ -111,4 +174,187 @@ mod tests { "ldaps://ldap.example.com" ); } + + #[tokio::test] + async fn ldap_compensation_delete_requires_an_exact_match() { + let expected = ldap_config(); + let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(expected.clone())); + let mismatched = StoredLdapModuleConfig { + is_enabled: false, + ..expected.clone() + }; + + assert!(!repository + .delete_ldap_config_if_matches(&mismatched) + .await + .expect("mismatched delete should execute")); + assert!(repository + .delete_ldap_config_if_matches(&expected) + .await + .expect("matching delete should execute")); + assert!(repository + .get_ldap_config() + .await + .expect("LDAP config should remain readable") + .is_none()); + } + + #[tokio::test] + async fn ldap_compare_and_swap_separates_preserve_set_and_clear() { + let original = ldap_config(); + let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone())); + let replacement = StoredLdapModuleConfig { + server_url: "ldap://updated.example.com".to_string(), + bind_password_encrypted: Some("stale-ciphertext-must-be-ignored".to_string()), + ..original.clone() + }; + + let preserved = repository + .compare_and_swap_ldap_config( + Some(&original), + &replacement, + &LdapBindPasswordUpdate::Preserve, + ) + .await + .expect("preserve CAS should execute"); + let CompareAndSwapLdapConfigResult::Applied(preserved) = preserved else { + panic!("fresh snapshot should apply"); + }; + assert_eq!( + preserved.bind_password_encrypted.as_deref(), + Some("encrypted-password") + ); + + let set = repository + .compare_and_swap_ldap_config( + Some(&preserved), + &preserved, + &LdapBindPasswordUpdate::Set("rotated-ciphertext".to_string()), + ) + .await + .expect("set CAS should execute"); + let CompareAndSwapLdapConfigResult::Applied(set) = set else { + panic!("fresh snapshot should apply"); + }; + assert_eq!( + set.bind_password_encrypted.as_deref(), + Some("rotated-ciphertext") + ); + + let cleared = repository + .compare_and_swap_ldap_config(Some(&set), &set, &LdapBindPasswordUpdate::Clear) + .await + .expect("clear CAS should execute"); + let CompareAndSwapLdapConfigResult::Applied(cleared) = cleared else { + panic!("fresh snapshot should apply"); + }; + assert!(cleared.bind_password_encrypted.is_none()); + } + + #[tokio::test] + async fn ldap_compare_and_swap_rejects_stale_password_and_config_snapshots() { + let original = ldap_config(); + let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone())); + assert!(repository + .compare_and_swap_ldap_bind_password("encrypted-password", "rotated-ciphertext") + .await + .expect("password rotation should execute")); + + let stale_password_result = repository + .compare_and_swap_ldap_config( + Some(&original), + &StoredLdapModuleConfig { + base_dn: "dc=updated,dc=example".to_string(), + ..original.clone() + }, + &LdapBindPasswordUpdate::Preserve, + ) + .await + .expect("stale password CAS should execute"); + assert_eq!( + stale_password_result, + CompareAndSwapLdapConfigResult::Conflict + ); + assert_eq!( + repository + .get_ldap_config() + .await + .expect("LDAP config should load") + .and_then(|config| config.bind_password_encrypted) + .as_deref(), + Some("rotated-ciphertext") + ); + + let current = repository + .get_ldap_config() + .await + .expect("LDAP config should load") + .expect("LDAP config should exist"); + let changed = StoredLdapModuleConfig { + is_enabled: false, + ..current.clone() + }; + let applied = repository + .compare_and_swap_ldap_config( + Some(¤t), + &changed, + &LdapBindPasswordUpdate::Preserve, + ) + .await + .expect("fresh config CAS should execute"); + assert!(matches!( + applied, + CompareAndSwapLdapConfigResult::Applied(_) + )); + let stale_config_result = repository + .compare_and_swap_ldap_config( + Some(¤t), + ¤t, + &LdapBindPasswordUpdate::Preserve, + ) + .await + .expect("stale config CAS should execute"); + assert_eq!( + stale_config_result, + CompareAndSwapLdapConfigResult::Conflict + ); + } + + #[tokio::test] + async fn ldap_compare_and_swap_allows_only_one_initial_create() { + let repository = InMemoryAuthModuleReadRepository::default(); + let replacement = StoredLdapModuleConfig { + bind_password_encrypted: None, + ..ldap_config() + }; + + let first = repository + .compare_and_swap_ldap_config( + None, + &replacement, + &LdapBindPasswordUpdate::Set("first-ciphertext".to_string()), + ) + .await + .expect("first create should execute"); + assert!(matches!(first, CompareAndSwapLdapConfigResult::Applied(_))); + + let second = repository + .compare_and_swap_ldap_config( + None, + &replacement, + &LdapBindPasswordUpdate::Set("second-ciphertext".to_string()), + ) + .await + .expect("second create should execute"); + assert_eq!(second, CompareAndSwapLdapConfigResult::Conflict); + assert_eq!( + repository + .get_ldap_config() + .await + .expect("LDAP config should load") + .and_then(|config| config.bind_password_encrypted) + .as_deref(), + Some("first-ciphertext") + ); + } } diff --git a/crates/aether-data/runtime/src/repository/auth_modules/mod.rs b/crates/aether-data/runtime/src/repository/auth_modules/mod.rs index d09beca7a..22d053007 100644 --- a/crates/aether-data/runtime/src/repository/auth_modules/mod.rs +++ b/crates/aether-data/runtime/src/repository/auth_modules/mod.rs @@ -1,8 +1,8 @@ mod memory; pub use aether_data_contracts::repository::auth_modules::{ - AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig, - StoredOAuthProviderModuleConfig, + AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult, + LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository}; diff --git a/crates/aether-data/runtime/src/repository/background_tasks/memory.rs b/crates/aether-data/runtime/src/repository/background_tasks/memory.rs index d110e1b3a..463bf8339 100644 --- a/crates/aether-data/runtime/src/repository/background_tasks/memory.rs +++ b/crates/aether-data/runtime/src/repository/background_tasks/memory.rs @@ -53,7 +53,8 @@ impl InMemoryBackgroundTaskRepository { I: IntoIterator, { let mut index = InMemoryBackgroundTaskIndex::default(); - for run in runs { + for mut run in runs { + run.sanitize_persisted_data(); index.runs.insert(run.id.clone(), run); } Self { @@ -164,8 +165,9 @@ impl BackgroundTaskReadRepository for InMemoryBackgroundTaskRepository { impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository { async fn upsert_run( &self, - run: UpsertBackgroundTaskRun, + mut run: UpsertBackgroundTaskRun, ) -> Result { + run.sanitize_for_persistence(); run.validate()?; let stored = run.into_stored(); self.index @@ -192,8 +194,9 @@ impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository { async fn upsert_event( &self, - event: UpsertBackgroundTaskEvent, + mut event: UpsertBackgroundTaskEvent, ) -> Result { + event.sanitize_for_persistence(); event.validate()?; let stored = event.into_stored(); let mut guard = self.index.write().expect("background task repository lock"); diff --git a/crates/aether-data/runtime/src/repository/billing/memory.rs b/crates/aether-data/runtime/src/repository/billing/memory.rs index 8f07894df..112f7f3fc 100644 --- a/crates/aether-data/runtime/src/repository/billing/memory.rs +++ b/crates/aether-data/runtime/src/repository/billing/memory.rs @@ -5,8 +5,9 @@ use async_trait::async_trait; use super::{ AdminBillingMutationOutcome, BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, - PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, StoredBillingModelContext, - UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, + PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, + PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord, + UserPlanEntitlementRecord, }; use crate::DataLayerError; @@ -75,6 +76,7 @@ fn billing_plan_from_input( fn daily_quota_availability_from_entitlements( entitlements: impl IntoIterator, + billing_plans: &BTreeMap, now: u64, ) -> UserDailyQuotaAvailabilityRecord { let mut has_active_daily_quota = false; @@ -92,6 +94,9 @@ fn daily_quota_availability_from_entitlements( let Some(items) = entitlement.entitlements_snapshot.as_array() else { continue; }; + let current_allow_wallet_overage = billing_plans + .get(&entitlement.plan_id) + .and_then(|plan| daily_quota_wallet_overage_policy(&plan.entitlements_json)); for item in items { if item.get("type").and_then(serde_json::Value::as_str) != Some("daily_quota") { continue; @@ -106,10 +111,11 @@ fn daily_quota_availability_from_entitlements( has_active_daily_quota = true; total_quota_usd += daily_quota_usd; remaining_usd += daily_quota_usd; - allow_wallet_overage &= item - .get("allow_wallet_overage") - .and_then(serde_json::Value::as_bool) - .unwrap_or(false); + allow_wallet_overage &= current_allow_wallet_overage.unwrap_or_else(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }); } } UserDailyQuotaAvailabilityRecord { @@ -121,6 +127,17 @@ fn daily_quota_availability_from_entitlements( } } +fn daily_quota_wallet_overage_policy(entitlements: &serde_json::Value) -> Option { + entitlements.as_array()?.iter().find_map(|item| { + (item.get("type").and_then(serde_json::Value::as_str) == Some("daily_quota")) + .then(|| { + item.get("allow_wallet_overage") + .and_then(serde_json::Value::as_bool) + }) + .flatten() + }) +} + #[async_trait] impl BillingReadRepository for InMemoryBillingReadRepository { async fn find_model_context( @@ -199,6 +216,76 @@ impl BillingReadRepository for InMemoryBillingReadRepository { .cloned()) } + async fn compare_and_swap_payment_gateway_secret( + &self, + update: &PaymentGatewaySecretCasUpdate, + ) -> Result { + let provider = update.provider.trim().to_ascii_lowercase(); + let mut configs = self + .gateway_configs_by_provider + .write() + .expect("billing repository lock"); + let Some(record) = configs.get_mut(&provider) else { + return Ok(false); + }; + if record.merchant_key_encrypted.as_deref() + != Some(update.expected_merchant_key_encrypted.as_str()) + { + return Ok(false); + } + record.merchant_key_encrypted = Some(update.merchant_key_encrypted.clone()); + Ok(true) + } + + async fn compare_and_swap_payment_gateway_config( + &self, + mutation: &PaymentGatewayConfigCasWriteInput, + ) -> Result, DataLayerError> { + let input = &mutation.input; + let provider = input.provider.trim().to_ascii_lowercase(); + let now = current_unix_secs(); + let mut configs = self + .gateway_configs_by_provider + .write() + .expect("billing repository lock"); + let existing = configs.get(&provider); + if mutation.expected_existing { + let Some(existing) = existing else { + return Ok(AdminBillingMutationOutcome::NotFound); + }; + if existing.merchant_key_encrypted != mutation.expected_merchant_key_encrypted { + return Ok(AdminBillingMutationOutcome::NotFound); + } + } else if existing.is_some() { + return Ok(AdminBillingMutationOutcome::NotFound); + } + + let created_at = existing + .map(|value| value.created_at_unix_secs) + .unwrap_or(now); + let merchant_key_encrypted = if input.preserve_existing_secret { + existing.and_then(|value| value.merchant_key_encrypted.clone()) + } else { + input.merchant_key_encrypted.clone() + }; + let record = PaymentGatewayConfigRecord { + provider: provider.clone(), + enabled: input.enabled, + endpoint_url: input.endpoint_url.clone(), + callback_base_url: input.callback_base_url.clone(), + merchant_id: input.merchant_id.clone(), + merchant_key_encrypted, + pay_currency: input.pay_currency.clone(), + usd_exchange_rate: input.usd_exchange_rate, + min_recharge_usd: input.min_recharge_usd, + channels_json: input.channels_json.clone(), + created_at_unix_secs: created_at, + updated_at_unix_secs: now, + }; + configs.insert(provider, record.clone()); + Ok(AdminBillingMutationOutcome::Applied(record)) + } + async fn upsert_payment_gateway_config( &self, input: &PaymentGatewayConfigWriteInput, @@ -366,6 +453,31 @@ impl BillingReadRepository for InMemoryBillingReadRepository { Ok(Some(items)) } + async fn revoke_user_plan_entitlement( + &self, + user_id: &str, + entitlement_id: &str, + ) -> Result, DataLayerError> { + let now = current_unix_secs(); + let mut entitlements = self + .entitlements_by_id + .write() + .expect("billing repository lock"); + let Some(entitlement) = entitlements.get_mut(entitlement_id) else { + return Ok(AdminBillingMutationOutcome::NotFound); + }; + if entitlement.user_id != user_id + || entitlement.status != "active" + || entitlement.expires_at_unix_secs <= now + { + return Ok(AdminBillingMutationOutcome::NotFound); + } + entitlement.status = "revoked".to_string(); + entitlement.expires_at_unix_secs = entitlement.expires_at_unix_secs.min(now); + entitlement.updated_at_unix_secs = now; + Ok(AdminBillingMutationOutcome::Applied(())) + } + async fn find_user_daily_quota_availability( &self, user_id: &str, @@ -379,8 +491,13 @@ impl BillingReadRepository for InMemoryBillingReadRepository { .filter(|item| item.user_id == user_id) .cloned() .collect::>(); + let billing_plans = self + .billing_plans_by_id + .read() + .expect("billing repository lock"); Ok(Some(daily_quota_availability_from_entitlements( entitlements, + &billing_plans, now, ))) } @@ -449,7 +566,10 @@ mod tests { use serde_json::json; use super::InMemoryBillingReadRepository; - use crate::repository::billing::{BillingReadRepository, StoredBillingModelContext}; + use crate::repository::billing::{ + AdminBillingMutationOutcome, BillingReadRepository, PaymentGatewayConfigCasWriteInput, + PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, StoredBillingModelContext, + }; fn sample_context() -> StoredBillingModelContext { StoredBillingModelContext::new( @@ -557,4 +677,105 @@ mod tests { Some("gpt-5-upstream") ); } + + fn gateway_input(secret: Option<&str>) -> PaymentGatewayConfigWriteInput { + PaymentGatewayConfigWriteInput { + provider: "stripe".to_string(), + enabled: true, + endpoint_url: "https://api.stripe.com".to_string(), + callback_base_url: Some("https://example.com".to_string()), + merchant_id: "merchant".to_string(), + merchant_key_encrypted: secret.map(ToOwned::to_owned), + preserve_existing_secret: false, + pay_currency: "USD".to_string(), + usd_exchange_rate: 1.0, + min_recharge_usd: 1.0, + channels_json: json!({"channels": []}), + } + } + + #[tokio::test] + async fn payment_gateway_cas_prevents_create_overwrite_and_uses_exact_secret_fence() { + let repository = InMemoryBillingReadRepository::default(); + let create = PaymentGatewayConfigCasWriteInput { + input: gateway_input(Some("ciphertext-a")), + expected_existing: false, + expected_merchant_key_encrypted: None, + }; + assert!(matches!( + repository + .compare_and_swap_payment_gateway_config(&create) + .await + .expect("create should succeed"), + AdminBillingMutationOutcome::Applied(_) + )); + + let mut competing_create = create.clone(); + competing_create.input.merchant_id = "overwritten".to_string(); + assert_eq!( + repository + .compare_and_swap_payment_gateway_config(&competing_create) + .await + .expect("conflicting create should be handled"), + AdminBillingMutationOutcome::NotFound + ); + + let mut stale_update = create.clone(); + stale_update.expected_existing = true; + stale_update.expected_merchant_key_encrypted = Some("ciphertext-stale".to_string()); + stale_update.input.merchant_id = "stale-update".to_string(); + assert_eq!( + repository + .compare_and_swap_payment_gateway_config(&stale_update) + .await + .expect("stale update should be handled"), + AdminBillingMutationOutcome::NotFound + ); + let stored = repository + .find_payment_gateway_config("stripe") + .await + .expect("lookup should succeed") + .expect("config should exist"); + assert_eq!(stored.merchant_id, "merchant"); + } + + #[tokio::test] + async fn payment_gateway_secret_cas_changes_no_other_fields() { + let repository = InMemoryBillingReadRepository::default(); + repository + .upsert_payment_gateway_config(&gateway_input(Some("legacy-ciphertext"))) + .await + .expect("seed should succeed"); + let before = repository + .find_payment_gateway_config("stripe") + .await + .expect("lookup should succeed") + .expect("config should exist"); + + assert!(!repository + .compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate { + provider: "stripe".to_string(), + expected_merchant_key_encrypted: "wrong-ciphertext".to_string(), + merchant_key_encrypted: "v2-ciphertext".to_string(), + }) + .await + .expect("stale secret CAS should be handled")); + assert!(repository + .compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate { + provider: "stripe".to_string(), + expected_merchant_key_encrypted: "legacy-ciphertext".to_string(), + merchant_key_encrypted: "v2-ciphertext".to_string(), + }) + .await + .expect("secret CAS should succeed")); + + let mut expected = before.clone(); + expected.merchant_key_encrypted = Some("v2-ciphertext".to_string()); + let after = repository + .find_payment_gateway_config("stripe") + .await + .expect("lookup should succeed") + .expect("config should exist"); + assert_eq!(after, expected); + } } diff --git a/crates/aether-data/runtime/src/repository/candidates/memory.rs b/crates/aether-data/runtime/src/repository/candidates/memory.rs index 17c536fa1..0002aaa73 100644 --- a/crates/aether-data/runtime/src/repository/candidates/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidates/memory.rs @@ -1,14 +1,18 @@ use std::collections::{BTreeMap, BTreeSet}; use std::sync::RwLock; -use async_trait::async_trait; - use super::{ request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate, UpsertRequestCandidateRecord, }; use crate::DataLayerError; +use async_trait::async_trait; + +fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate { + candidate.sanitize_sensitive_diagnostics(); + candidate +} fn merge_extra_data( existing: Option, @@ -38,7 +42,7 @@ impl InMemoryRequestCandidateRepository { I: IntoIterator, { let mut by_id = BTreeMap::new(); - for item in items { + for item in items.into_iter().map(sanitize_stored_candidate) { by_id.insert(item.id.clone(), item); } Self { @@ -60,6 +64,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { .values() .filter(|row| row.request_id == request_id) .cloned() + .map(sanitize_stored_candidate) .collect::>(); rows.sort_by(|left, right| { left.candidate_index @@ -84,6 +89,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { .expect("request candidate repository lock") .values() .cloned() + .map(sanitize_stored_candidate) .collect::>(); rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); rows.truncate(limit); @@ -106,6 +112,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { .values() .filter(|row| row.provider_id.as_deref() == Some(provider_id)) .cloned() + .map(sanitize_stored_candidate) .collect::>(); rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); rows.truncate(limit); @@ -141,6 +148,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { ) }) .cloned() + .map(sanitize_stored_candidate) .collect::>(); rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); rows.truncate(limit); @@ -294,8 +302,9 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { async fn upsert( &self, - candidate: UpsertRequestCandidateRecord, + mut candidate: UpsertRequestCandidateRecord, ) -> Result { + candidate.sanitize_for_persistence(); candidate.validate()?; let mut by_id = self @@ -309,7 +318,8 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { && row.candidate_index == candidate.candidate_index && row.retry_index == candidate.retry_index }) - .cloned(); + .cloned() + .map(sanitize_stored_candidate); let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| { request_candidate_lifecycle_would_regress(row.status, candidate.status) @@ -336,29 +346,36 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { .map(|row| row.id.clone()) .unwrap_or_else(|| candidate.id.clone()), request_id: candidate.request_id.clone(), - user_id: candidate - .user_id - .or_else(|| existing.as_ref().and_then(|row| row.user_id.clone())), - api_key_id: candidate - .api_key_id - .or_else(|| existing.as_ref().and_then(|row| row.api_key_id.clone())), - username: candidate - .username - .or_else(|| existing.as_ref().and_then(|row| row.username.clone())), - api_key_name: candidate - .api_key_name - .or_else(|| existing.as_ref().and_then(|row| row.api_key_name.clone())), + user_id: existing + .as_ref() + .and_then(|row| row.user_id.clone()) + .or(candidate.user_id), + api_key_id: existing + .as_ref() + .and_then(|row| row.api_key_id.clone()) + .or(candidate.api_key_id), + username: existing + .as_ref() + .and_then(|row| row.username.clone()) + .or(candidate.username), + api_key_name: existing + .as_ref() + .and_then(|row| row.api_key_name.clone()) + .or(candidate.api_key_name), candidate_index: candidate.candidate_index, retry_index: candidate.retry_index, - provider_id: candidate - .provider_id - .or_else(|| existing.as_ref().and_then(|row| row.provider_id.clone())), - endpoint_id: candidate - .endpoint_id - .or_else(|| existing.as_ref().and_then(|row| row.endpoint_id.clone())), - key_id: candidate - .key_id - .or_else(|| existing.as_ref().and_then(|row| row.key_id.clone())), + provider_id: existing + .as_ref() + .and_then(|row| row.provider_id.clone()) + .or(candidate.provider_id), + endpoint_id: existing + .as_ref() + .and_then(|row| row.endpoint_id.clone()) + .or(candidate.endpoint_id), + key_id: existing + .as_ref() + .and_then(|row| row.key_id.clone()) + .or(candidate.key_id), status: merged_status, skip_reason: candidate .skip_reason @@ -380,13 +397,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { .error_type .or_else(|| existing.as_ref().and_then(|row| row.error_type.clone())) }, - error_message: if preserve_existing_lifecycle { - existing.as_ref().and_then(|row| row.error_message.clone()) - } else { - candidate - .error_message - .or_else(|| existing.as_ref().and_then(|row| row.error_message.clone())) - }, + error_message: None, latency_ms: if preserve_existing_lifecycle { existing.as_ref().and_then(|row| row.latency_ms) } else { @@ -407,9 +418,10 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { .and_then(|row| row.required_capabilities.clone()) }), created_at_unix_ms, - started_at_unix_ms: candidate - .started_at_unix_ms - .or_else(|| existing.as_ref().and_then(|row| row.started_at_unix_ms)), + started_at_unix_ms: existing + .as_ref() + .and_then(|row| row.started_at_unix_ms) + .or(candidate.started_at_unix_ms), finished_at_unix_ms: if preserve_existing_lifecycle { existing.as_ref().and_then(|row| row.finished_at_unix_ms) } else { @@ -418,6 +430,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { .or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms)) }, }; + let stored = sanitize_stored_candidate(stored); by_id.insert(stored.id.clone(), stored.clone()); Ok(stored) @@ -531,6 +544,138 @@ mod tests { assert_eq!(rows[1].id, "cand-1"); } + #[tokio::test] + async fn seed_and_reads_sanitize_candidates_that_bypass_contract_constructors() { + let raw_candidate = StoredRequestCandidate { + id: "cand-raw".to_string(), + request_id: "req-raw".to_string(), + user_id: None, + api_key_id: None, + username: None, + api_key_name: None, + candidate_index: 0, + retry_index: 0, + provider_id: Some("provider-1".to_string()), + endpoint_id: Some("endpoint-1".to_string()), + key_id: None, + status: RequestCandidateStatus::Failed, + skip_reason: Some("secret=/private/path".to_string()), + is_cached: false, + status_code: Some(500), + error_type: Some("token=secret".to_string()), + error_message: Some("Bearer secret-token".to_string()), + latency_ms: Some(10), + concurrent_requests: Some(1), + extra_data: Some(json!({ + "gateway_execution_runtime": true, + "request_headers": {"authorization": "Bearer secret-token"}, + "request_body": {"password": "secret"} + })), + required_capabilities: Some(json!({ + "streaming": "true", + "internal_capability": "secret" + })), + created_at_unix_ms: 100, + started_at_unix_ms: Some(100), + finished_at_unix_ms: Some(110), + }; + let repository = InMemoryRequestCandidateRepository::seed(vec![raw_candidate.clone()]); + + { + let stored = repository + .by_id + .read() + .expect("request candidate repository lock"); + let candidate = stored + .get("cand-raw") + .expect("seeded candidate should exist"); + assert!(candidate.error_message.is_none()); + assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip")); + assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error")); + assert_eq!( + candidate.extra_data, + Some(json!({"gateway_execution_runtime": true})) + ); + assert_eq!( + candidate.required_capabilities, + Some(json!({"streaming": true})) + ); + } + + let mut bypassed_candidate = raw_candidate; + bypassed_candidate.id = "cand-bypassed".to_string(); + bypassed_candidate.request_id = "req-bypassed".to_string(); + repository + .by_id + .write() + .expect("request candidate repository lock") + .insert(bypassed_candidate.id.clone(), bypassed_candidate); + + let rows = repository + .list_recent(10) + .await + .expect("list recent should succeed"); + let candidate = rows + .iter() + .find(|candidate| candidate.id == "cand-bypassed") + .expect("bypassed candidate should be returned"); + assert!(candidate.error_message.is_none()); + assert_eq!( + candidate.extra_data, + Some(json!({"gateway_execution_runtime": true})) + ); + assert_eq!( + candidate.required_capabilities, + Some(json!({"streaming": true})) + ); + + let merged = repository + .upsert(UpsertRequestCandidateRecord { + id: "cand-merged".to_string(), + request_id: "req-bypassed".to_string(), + user_id: None, + api_key_id: None, + username: None, + api_key_name: None, + candidate_index: 0, + retry_index: 0, + provider_id: None, + endpoint_id: None, + key_id: None, + status: RequestCandidateStatus::Success, + skip_reason: None, + is_cached: None, + status_code: Some(200), + error_type: None, + error_message: Some("Bearer new-secret".to_string()), + latency_ms: Some(12), + concurrent_requests: None, + extra_data: Some(json!({ + "stream_completed": true, + "request_body": {"password": "new-secret"} + })), + required_capabilities: Some(json!({ + "vision": 1, + "internal_capability": "new-secret" + })), + created_at_unix_ms: Some(100), + started_at_unix_ms: Some(100), + finished_at_unix_ms: Some(112), + }) + .await + .expect("candidate merge should succeed"); + assert_eq!(merged.id, "cand-bypassed"); + assert!(merged.error_message.is_none()); + assert_eq!( + merged.extra_data, + Some(json!({ + "gateway_execution_runtime": true, + "stream_completed": true + })) + ); + assert_eq!(merged.required_capabilities, Some(json!({"vision": true}))); + } + #[tokio::test] async fn aggregates_finalized_health_data_by_endpoint_ids() { let repository = InMemoryRequestCandidateRepository::seed(vec![ @@ -594,7 +739,7 @@ mod tests { "execution_strategy": "local_cross_format", "provider_name": "primary", })), - required_capabilities: None, + required_capabilities: Some(json!({"streaming": true})), created_at_unix_ms: Some(100), started_at_unix_ms: None, finished_at_unix_ms: None, @@ -608,15 +753,15 @@ mod tests { .upsert(UpsertRequestCandidateRecord { id: "cand-1-replacement".to_string(), request_id: "req-1".to_string(), - user_id: None, - api_key_id: None, - username: None, - api_key_name: None, + user_id: Some("attacker-user".to_string()), + api_key_id: Some("attacker-api-key".to_string()), + username: Some("mallory".to_string()), + api_key_name: Some("attacker-key".to_string()), candidate_index: 0, retry_index: 0, - provider_id: None, - endpoint_id: None, - key_id: None, + provider_id: Some("attacker-provider".to_string()), + endpoint_id: Some("attacker-endpoint".to_string()), + key_id: Some("attacker-provider-key".to_string()), status: RequestCandidateStatus::Success, skip_reason: None, is_cached: None, @@ -629,7 +774,7 @@ mod tests { "provider_api_format": "openai:responses", "provider_name": "updated", })), - required_capabilities: None, + required_capabilities: Some(json!({"vision": true})), created_at_unix_ms: None, started_at_unix_ms: Some(101), finished_at_unix_ms: Some(102), @@ -638,6 +783,14 @@ mod tests { .expect("update should succeed"); assert_eq!(updated.id, "cand-1"); assert_eq!(updated.status, RequestCandidateStatus::Success); + assert_eq!(updated.user_id.as_deref(), Some("user-1")); + assert_eq!(updated.api_key_id.as_deref(), Some("api-key-1")); + assert!(updated.username.is_none()); + assert!(updated.api_key_name.is_none()); + assert_eq!(updated.provider_id.as_deref(), Some("provider-1")); + assert_eq!(updated.endpoint_id.as_deref(), Some("endpoint-1")); + assert_eq!(updated.key_id.as_deref(), Some("key-1")); + assert_eq!(updated.required_capabilities, Some(json!({"vision": true}))); assert_eq!(updated.status_code, Some(200)); assert_eq!(updated.latency_ms, Some(25)); assert_eq!( @@ -659,13 +812,13 @@ mod tests { .extra_data .as_ref() .and_then(|value| value.get("provider_name")), - Some(&json!("updated")) + None ); assert_eq!(updated.started_at_unix_ms, Some(101)); } #[tokio::test] - async fn upsert_keeps_terminal_candidate_state_when_streaming_arrives_late() { + async fn upsert_keeps_first_terminal_candidate_fact_when_another_terminal_arrives_late() { let existing = StoredRequestCandidate::new( "cand-1".to_string(), "req-1".to_string(), @@ -686,7 +839,7 @@ mod tests { Some("retryable upstream failure".to_string()), Some(45), Some(1), - Some(json!({"terminal": true})), + Some(json!({"stream_completed": true})), None, 100, Some(101), @@ -708,7 +861,7 @@ mod tests { provider_id: None, endpoint_id: None, key_id: None, - status: RequestCandidateStatus::Streaming, + status: RequestCandidateStatus::Success, skip_reason: None, is_cached: None, status_code: Some(200), @@ -716,7 +869,7 @@ mod tests { error_message: None, latency_ms: Some(9_999), concurrent_requests: Some(2), - extra_data: Some(json!({"late": true})), + extra_data: Some(json!({"gateway_execution_runtime": true})), required_capabilities: None, created_at_unix_ms: None, started_at_unix_ms: Some(102), @@ -729,16 +882,16 @@ mod tests { assert_eq!(updated.status, RequestCandidateStatus::Failed); assert_eq!(updated.status_code, Some(503)); assert_eq!(updated.error_type.as_deref(), Some("upstream_error")); - assert_eq!( - updated.error_message.as_deref(), - Some("retryable upstream failure") - ); + assert!(updated.error_message.is_none()); assert_eq!(updated.latency_ms, Some(45)); assert_eq!(updated.concurrent_requests, Some(2)); assert_eq!(updated.finished_at_unix_ms, Some(145)); assert_eq!( updated.extra_data, - Some(json!({"terminal": true, "late": true})) + Some(json!({ + "gateway_execution_runtime": true, + "stream_completed": true + })) ); } diff --git a/crates/aether-data/runtime/src/repository/gemini_file_mappings/memory.rs b/crates/aether-data/runtime/src/repository/gemini_file_mappings/memory.rs index 6b7acdff3..f4eafab6d 100644 --- a/crates/aether-data/runtime/src/repository/gemini_file_mappings/memory.rs +++ b/crates/aether-data/runtime/src/repository/gemini_file_mappings/memory.rs @@ -41,6 +41,40 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository { Ok(guard.get(file_name).cloned()) } + async fn find_active_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let guard = self.by_file.read().expect("gemini mapping repository lock"); + Ok(guard + .get(file_name) + .filter(|mapping| { + mapping.user_id.as_deref() == Some(user_id) + && mapping.expires_at_unix_secs > now_unix_secs + }) + .cloned()) + } + + async fn find_active_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + now_unix_secs: u64, + ) -> Result, DataLayerError> { + let guard = self.by_file.read().expect("gemini mapping repository lock"); + Ok(guard + .get(file_name) + .filter(|mapping| { + mapping.key_id == key_id + && mapping.user_id.as_deref() == Some(user_id) + && mapping.expires_at_unix_secs > now_unix_secs + }) + .cloned()) + } + async fn list_mappings( &self, query: &GeminiFileMappingListQuery, @@ -52,6 +86,12 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository { .map(|value| value.to_ascii_lowercase()); let mut items = guard .values() + .filter(|item| { + query + .user_id + .as_deref() + .is_none_or(|user_id| item.user_id.as_deref() == Some(user_id)) + }) .filter(|item| query.include_expired || item.expires_at_unix_secs > query.now_unix_secs) .filter(|item| { search.as_deref().is_none_or(|needle| { @@ -146,6 +186,39 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository { Ok(mapping) } + async fn upsert_if_owner_matches( + &self, + record: UpsertGeminiFileMappingRecord, + ) -> Result, DataLayerError> { + record.validate()?; + let mut guard = self + .by_file + .write() + .expect("gemini mapping repository lock"); + let (id, created_at_unix_ms) = match guard.get(&record.file_name) { + Some(existing) + if existing.key_id == record.key_id && existing.user_id == record.user_id => + { + (existing.id.clone(), existing.created_at_unix_ms) + } + Some(_) => return Ok(None), + None => (record.id.clone(), current_unix_secs()), + }; + let mapping = StoredGeminiFileMapping { + id, + file_name: record.file_name.clone(), + key_id: record.key_id.clone(), + user_id: record.user_id.clone(), + display_name: record.display_name.clone(), + mime_type: record.mime_type.clone(), + source_hash: record.source_hash.clone(), + created_at_unix_ms, + expires_at_unix_secs: record.expires_at_unix_secs, + }; + guard.insert(record.file_name, mapping.clone()); + Ok(Some(mapping)) + } + async fn delete_by_file_name(&self, file_name: &str) -> Result { let mut guard = self .by_file @@ -154,6 +227,44 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository { Ok(guard.remove(file_name).is_some()) } + async fn delete_by_file_name_for_user( + &self, + file_name: &str, + user_id: &str, + ) -> Result { + let mut guard = self + .by_file + .write() + .expect("gemini mapping repository lock"); + if guard + .get(file_name) + .and_then(|item| item.user_id.as_deref()) + != Some(user_id) + { + return Ok(false); + } + Ok(guard.remove(file_name).is_some()) + } + + async fn delete_by_file_name_for_owner( + &self, + file_name: &str, + key_id: &str, + user_id: &str, + ) -> Result { + let mut guard = self + .by_file + .write() + .expect("gemini mapping repository lock"); + let owner_matches = guard + .get(file_name) + .is_some_and(|item| item.key_id == key_id && item.user_id.as_deref() == Some(user_id)); + if !owner_matches { + return Ok(false); + } + Ok(guard.remove(file_name).is_some()) + } + async fn delete_by_id( &self, mapping_id: &str, @@ -224,6 +335,35 @@ mod tests { Ok(()) } + #[tokio::test] + async fn owner_scoped_reads_bind_user_provider_key_and_expiry() -> Result<(), DataLayerError> { + let repo = InMemoryGeminiFileMappingRepository::default(); + repo.upsert(sample_record("id-owner", "files/owned")) + .await?; + + assert!(repo + .find_active_by_file_name_for_user("files/owned", "user-1", 100) + .await? + .is_some()); + assert!(repo + .find_active_by_file_name_for_user("files/owned", "user-2", 100) + .await? + .is_none()); + assert!(repo + .find_active_by_file_name_for_owner("files/owned", "key-1", "user-1", 100) + .await? + .is_some()); + assert!(repo + .find_active_by_file_name_for_owner("files/owned", "key-2", "user-1", 100) + .await? + .is_none()); + assert!(repo + .find_active_by_file_name_for_user("files/owned", "user-1", 4_102_444_800) + .await? + .is_none()); + Ok(()) + } + #[tokio::test] async fn delete_removes_entry() -> Result<(), DataLayerError> { let repo = InMemoryGeminiFileMappingRepository::default(); @@ -247,6 +387,34 @@ mod tests { Ok(()) } + #[tokio::test] + async fn owner_checked_upsert_cannot_reassign_existing_mapping() -> Result<(), DataLayerError> { + let repo = InMemoryGeminiFileMappingRepository::default(); + let first = repo.upsert(sample_record("id-1", "files/owned")).await?; + + let mut attacker = sample_record("id-2", "files/owned"); + attacker.key_id = "key-2".to_string(); + attacker.user_id = Some("user-2".to_string()); + assert!(repo.upsert_if_owner_matches(attacker).await?.is_none()); + + let unchanged = repo + .find_by_file_name("files/owned") + .await? + .expect("mapping should remain"); + assert_eq!(unchanged.key_id, "key-1"); + assert_eq!(unchanged.user_id.as_deref(), Some("user-1")); + + let mut refresh = sample_record("id-3", "files/owned"); + refresh.display_name = Some("refreshed".to_string()); + let refreshed = repo + .upsert_if_owner_matches(refresh) + .await? + .expect("same owner should refresh"); + assert_eq!(refreshed.id, first.id); + assert_eq!(refreshed.display_name.as_deref(), Some("refreshed")); + Ok(()) + } + #[tokio::test] async fn list_and_summarize_mappings() -> Result<(), DataLayerError> { let repo = InMemoryGeminiFileMappingRepository::seed(vec![ @@ -257,6 +425,7 @@ mod tests { let page = repo .list_mappings(&GeminiFileMappingListQuery { + user_id: None, include_expired: false, search: Some("ga".to_string()), offset: 0, @@ -279,6 +448,46 @@ mod tests { Ok(()) } + #[tokio::test] + async fn owner_filter_and_delete_do_not_cross_user_boundaries() -> Result<(), DataLayerError> { + let mut first = repo_item("id-1", "files/alpha", "image/png", 10, 200); + first.user_id = Some("user-1".to_string()); + let mut second = repo_item("id-2", "files/beta", "image/png", 20, 200); + second.user_id = Some("user-2".to_string()); + let repo = InMemoryGeminiFileMappingRepository::seed([first, second]); + + let page = repo + .list_mappings(&GeminiFileMappingListQuery { + user_id: Some("user-1".to_string()), + include_expired: false, + search: None, + offset: 0, + limit: 10, + now_unix_secs: 100, + }) + .await?; + assert_eq!(page.total, 1); + assert_eq!(page.items[0].file_name, "files/alpha"); + + assert!( + !repo + .delete_by_file_name_for_user("files/alpha", "user-2") + .await? + ); + assert!(repo.find_by_file_name("files/alpha").await?.is_some()); + assert!( + !repo + .delete_by_file_name_for_owner("files/alpha", "key-2", "user-1") + .await? + ); + assert!(repo.find_by_file_name("files/alpha").await?.is_some()); + assert!( + repo.delete_by_file_name_for_user("files/alpha", "user-1") + .await? + ); + Ok(()) + } + #[tokio::test] async fn delete_by_id_and_cleanup_expired() -> Result<(), DataLayerError> { let repo = InMemoryGeminiFileMappingRepository::seed(vec![ diff --git a/crates/aether-data/runtime/src/repository/management_tokens/memory.rs b/crates/aether-data/runtime/src/repository/management_tokens/memory.rs index 9db8552cf..583981365 100644 --- a/crates/aether-data/runtime/src/repository/management_tokens/memory.rs +++ b/crates/aether-data/runtime/src/repository/management_tokens/memory.rs @@ -6,9 +6,10 @@ use async_trait::async_trait; use crate::DataLayerError; use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenWithUser, + UpdateManagementTokenRecord, }; #[derive(Debug, Default)] @@ -49,6 +50,189 @@ impl InMemoryManagementTokenRepository { fn remove_hash_for_token(hashes: &mut BTreeMap, token_id: &str) { hashes.retain(|_, existing_token_id| existing_token_id != token_id); } + + fn update_management_token_scoped( + &self, + record: &UpdateManagementTokenRecord, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + record.validate()?; + + let mut items = self + .items + .write() + .expect("management token repository lock"); + let Some(index) = items.iter().position(|item| { + item.token.id == record.token_id + && expected_user_id + .map(|user_id| item.token.user_id == user_id) + .unwrap_or(true) + }) else { + return Ok(None); + }; + + if let Some(name) = &record.name { + if items.iter().enumerate().any(|(position, item)| { + position != index + && item.token.user_id == items[index].token.user_id + && item.token.name == *name + }) { + return Err(DataLayerError::InvalidInput(format!( + "已存在名为 '{}' 的 Token", + name + ))); + } + items[index].token.name = name.clone(); + } + + if record.clear_description { + items[index].token.description = None; + } else if let Some(description) = &record.description { + items[index].token.description = Some(description.clone()); + } + + if record.clear_allowed_ips { + items[index].token.allowed_ips = None; + } else if let Some(allowed_ips) = &record.allowed_ips { + items[index].token.allowed_ips = Some(allowed_ips.clone()); + } + + if let Some(permissions) = &record.permissions { + items[index].token.permissions = Some(permissions.clone()); + } + + if record.clear_expires_at { + items[index].token.expires_at_unix_secs = None; + } else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs { + items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs); + } + + if let Some(is_active) = record.is_active { + items[index].token.is_active = is_active; + } + + items[index].token.updated_at_unix_secs = Self::now_unix_secs(); + Ok(Some(items[index].token.clone())) + } + + fn delete_management_token_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + ) -> bool { + let mut items = self + .items + .write() + .expect("management token repository lock"); + let mut hashes = self + .hashes + .write() + .expect("management token repository lock"); + let original_len = items.len(); + items.retain(|item| { + item.token.id != token_id + || expected_user_id + .map(|user_id| item.token.user_id != user_id) + .unwrap_or(false) + }); + if items.len() != original_len { + Self::remove_hash_for_token(&mut hashes, token_id); + return true; + } + false + } + + fn set_management_token_active_scoped( + &self, + token_id: &str, + expected_user_id: Option<&str>, + is_active: bool, + ) -> Option { + let mut items = self + .items + .write() + .expect("management token repository lock"); + let item = items.iter_mut().find(|item| { + item.token.id == token_id + && expected_user_id + .map(|user_id| item.token.user_id == user_id) + .unwrap_or(true) + })?; + item.token.is_active = is_active; + item.token.updated_at_unix_secs = Self::now_unix_secs(); + Some(item.token.clone()) + } + + fn activate_management_token_if_matches_inner( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + // The in-memory token store does not own the independently stored user row and cannot + // atomically verify role/status/security_version with this mutation. Pretending that the + // user summary cached beside the token is authoritative would recreate the TOCTOU, so + // one-time install activation is intentionally unavailable on this backend. + Ok(false) + } + + fn delete_inactive_management_token_if_matches_inner( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + mutation.validate()?; + + let mut items = self + .items + .write() + .expect("management token repository lock"); + let mut hashes = self + .hashes + .write() + .expect("management token repository lock"); + if hashes.get(&mutation.token_hash).map(String::as_str) + != Some(mutation.expected_token.id.as_str()) + { + return Ok(false); + } + let Some(index) = items.iter().position(|item| { + mutation.matches_locked_token_snapshot(&item.token, &mutation.token_hash) + }) else { + return Ok(false); + }; + items.remove(index); + Self::remove_hash_for_token(&mut hashes, &mutation.expected_token.id); + Ok(true) + } + + fn regenerate_management_token_secret_scoped( + &self, + mutation: &RegenerateManagementTokenSecret, + expected_user_id: Option<&str>, + ) -> Result, DataLayerError> { + mutation.validate()?; + + let mut items = self + .items + .write() + .expect("management token repository lock"); + let mut hashes = self + .hashes + .write() + .expect("management token repository lock"); + let Some(item) = items.iter_mut().find(|item| { + item.token.id == mutation.token_id + && expected_user_id + .map(|user_id| item.token.user_id == user_id) + .unwrap_or(true) + }) else { + return Ok(None); + }; + Self::remove_hash_for_token(&mut hashes, &mutation.token_id); + hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone()); + item.token.token_prefix = mutation.token_prefix.clone(); + item.token.updated_at_unix_secs = Self::now_unix_secs(); + Ok(Some(item.token.clone())) + } } #[async_trait] @@ -167,76 +351,27 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository { &self, record: &UpdateManagementTokenRecord, ) -> Result, DataLayerError> { - record.validate()?; + self.update_management_token_scoped(record, None) + } - let mut items = self - .items - .write() - .expect("management token repository lock"); - let Some(index) = items - .iter() - .position(|item| item.token.id == record.token_id) - else { - return Ok(None); - }; - - if let Some(name) = &record.name { - if items.iter().enumerate().any(|(position, item)| { - position != index - && item.token.user_id == items[index].token.user_id - && item.token.name == *name - }) { - return Err(DataLayerError::InvalidInput(format!( - "已存在名为 '{}' 的 Token", - name - ))); - } - items[index].token.name = name.clone(); - } - - if record.clear_description { - items[index].token.description = None; - } else if let Some(description) = &record.description { - items[index].token.description = Some(description.clone()); - } - - if record.clear_allowed_ips { - items[index].token.allowed_ips = None; - } else if let Some(allowed_ips) = &record.allowed_ips { - items[index].token.allowed_ips = Some(allowed_ips.clone()); - } - - if let Some(permissions) = &record.permissions { - items[index].token.permissions = Some(permissions.clone()); - } - - if record.clear_expires_at { - items[index].token.expires_at_unix_secs = None; - } else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs { - items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs); - } - - if let Some(is_active) = record.is_active { - items[index].token.is_active = is_active; - } - - items[index].token.updated_at_unix_secs = Self::now_unix_secs(); - Ok(Some(items[index].token.clone())) + async fn update_management_token_for_user( + &self, + record: &UpdateManagementTokenRecord, + user_id: &str, + ) -> Result, DataLayerError> { + self.update_management_token_scoped(record, Some(user_id)) } async fn delete_management_token(&self, token_id: &str) -> Result { - let mut items = self - .items - .write() - .expect("management token repository lock"); - let mut hashes = self - .hashes - .write() - .expect("management token repository lock"); - let original_len = items.len(); - items.retain(|item| item.token.id != token_id); - Self::remove_hash_for_token(&mut hashes, token_id); - Ok(items.len() != original_len) + Ok(self.delete_management_token_scoped(token_id, None)) + } + + async fn delete_management_token_for_user( + &self, + token_id: &str, + user_id: &str, + ) -> Result { + Ok(self.delete_management_token_scoped(token_id, Some(user_id))) } async fn set_management_token_active( @@ -244,43 +379,45 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository { token_id: &str, is_active: bool, ) -> Result, DataLayerError> { - let mut items = self - .items - .write() - .expect("management token repository lock"); - let Some(item) = items.iter_mut().find(|item| item.token.id == token_id) else { - return Ok(None); - }; - item.token.is_active = is_active; - item.token.updated_at_unix_secs = Self::now_unix_secs(); - Ok(Some(item.token.clone())) + Ok(self.set_management_token_active_scoped(token_id, None, is_active)) + } + + async fn set_management_token_active_for_user( + &self, + token_id: &str, + user_id: &str, + is_active: bool, + ) -> Result, DataLayerError> { + Ok(self.set_management_token_active_scoped(token_id, Some(user_id), is_active)) + } + + async fn activate_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + self.activate_management_token_if_matches_inner(mutation) + } + + async fn delete_inactive_management_token_if_matches( + &self, + mutation: &ActivateManagementTokenIfMatches, + ) -> Result { + self.delete_inactive_management_token_if_matches_inner(mutation) } async fn regenerate_management_token_secret( &self, mutation: &RegenerateManagementTokenSecret, ) -> Result, DataLayerError> { - mutation.validate()?; + self.regenerate_management_token_secret_scoped(mutation, None) + } - let mut items = self - .items - .write() - .expect("management token repository lock"); - let mut hashes = self - .hashes - .write() - .expect("management token repository lock"); - let Some(item) = items - .iter_mut() - .find(|item| item.token.id == mutation.token_id) - else { - return Ok(None); - }; - Self::remove_hash_for_token(&mut hashes, &mutation.token_id); - hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone()); - item.token.token_prefix = mutation.token_prefix.clone(); - item.token.updated_at_unix_secs = Self::now_unix_secs(); - Ok(Some(item.token.clone())) + async fn regenerate_management_token_secret_for_user( + &self, + mutation: &RegenerateManagementTokenSecret, + user_id: &str, + ) -> Result, DataLayerError> { + self.regenerate_management_token_secret_scoped(mutation, Some(user_id)) } async fn record_management_token_usage( @@ -307,10 +444,10 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository { mod tests { use super::InMemoryManagementTokenRepository; use crate::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenUserSummary, StoredManagementTokenWithUser, - UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, + RegenerateManagementTokenSecret, StoredManagementToken, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, UpdateManagementTokenRecord, }; fn sample_token(id: &str, user_id: &str, is_active: bool) -> StoredManagementTokenWithUser { @@ -456,4 +593,108 @@ mod tests { .expect("hash lookup should succeed"); assert!(deleted_by_hash.is_none()); } + + #[tokio::test] + async fn owner_scoped_mutations_never_cross_user_boundaries() { + let repository = InMemoryManagementTokenRepository::seed_with_hashes( + vec![sample_token("token-1", "user-1", true)], + vec![("hash-1".to_string(), "token-1".to_string())], + ); + let update = UpdateManagementTokenRecord { + token_id: "token-1".to_string(), + name: Some("hijacked".to_string()), + description: None, + clear_description: false, + allowed_ips: None, + clear_allowed_ips: false, + permissions: None, + expires_at_unix_secs: None, + clear_expires_at: false, + is_active: None, + }; + + assert!(repository + .update_management_token_for_user(&update, "user-2") + .await + .expect("scoped update should execute") + .is_none()); + assert!(repository + .set_management_token_active_for_user("token-1", "user-2", false) + .await + .expect("scoped toggle should execute") + .is_none()); + assert!(repository + .regenerate_management_token_secret_for_user( + &RegenerateManagementTokenSecret { + token_id: "token-1".to_string(), + token_hash: "hash-hijacked".to_string(), + token_prefix: Some("ae_hijacked".to_string()), + }, + "user-2", + ) + .await + .expect("scoped regeneration should execute") + .is_none()); + assert!(!repository + .delete_management_token_for_user("token-1", "user-2") + .await + .expect("scoped delete should execute")); + + let unchanged = repository + .get_management_token_with_user_by_hash("hash-1") + .await + .expect("original hash lookup should succeed") + .expect("token should remain"); + assert_eq!(unchanged.token.name, "token-1"); + assert!(unchanged.token.is_active); + assert!(repository + .get_management_token_with_user_by_hash("hash-hijacked") + .await + .expect("replacement hash lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn install_activation_fails_closed_without_atomic_user_state() { + let mut pending = sample_token("token-1", "user-1", false); + pending.token.allowed_ips = Some(serde_json::json!(["127.0.0.1"])); + pending.token.permissions = Some(serde_json::json!(["admin:proxy_nodes:write"])); + pending.token.expires_at_unix_secs = Some(1_800_000_000); + let expected_token = pending.token.clone(); + let repository = InMemoryManagementTokenRepository::seed_with_hashes( + [pending], + [("hash-1".to_string(), "token-1".to_string())], + ); + let expected = ActivateManagementTokenIfMatches { + expected_token, + token_hash: "hash-1".to_string(), + expected_user_security_version: 4, + now_unix_secs: 1_700_000_000, + }; + + let mut mismatched = expected.clone(); + mismatched.expected_token.permissions = + Some(serde_json::json!(["admin:proxy_nodes:admin"])); + assert!(!repository + .activate_management_token_if_matches(&mismatched) + .await + .expect("mismatched activation should execute")); + assert!(!repository + .activate_management_token_if_matches(&expected) + .await + .expect("memory activation should fail closed")); + assert!( + !repository + .get_management_token_with_user("token-1") + .await + .expect("token lookup should execute") + .expect("token should remain") + .token + .is_active + ); + assert!(repository + .delete_inactive_management_token_if_matches(&expected) + .await + .expect("exact inactive snapshot cleanup should execute")); + } } diff --git a/crates/aether-data/runtime/src/repository/management_tokens/mod.rs b/crates/aether-data/runtime/src/repository/management_tokens/mod.rs index c017684ce..3b0442b55 100644 --- a/crates/aether-data/runtime/src/repository/management_tokens/mod.rs +++ b/crates/aether-data/runtime/src/repository/management_tokens/mod.rs @@ -1,10 +1,10 @@ mod memory; pub use aether_data_contracts::repository::management_tokens::{ - CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, - ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken, - StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser, - UpdateManagementTokenRecord, + ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery, + ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret, + StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary, + StoredManagementTokenWithUser, UpdateManagementTokenRecord, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::MysqlManagementTokenRepository; diff --git a/crates/aether-data/runtime/src/repository/oauth_providers/memory.rs b/crates/aether-data/runtime/src/repository/oauth_providers/memory.rs index 63b41b80d..7f4fcdc48 100644 --- a/crates/aether-data/runtime/src/repository/oauth_providers/memory.rs +++ b/crates/aether-data/runtime/src/repository/oauth_providers/memory.rs @@ -7,7 +7,7 @@ use async_trait::async_trait; use crate::DataLayerError; use aether_data_contracts::repository::oauth_providers::{ EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository, - StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord, + StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord, }; #[derive(Debug, Default)] @@ -65,15 +65,30 @@ impl OAuthProviderReadRepository for InMemoryOAuthProviderRepository { #[async_trait] impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository { - async fn upsert_oauth_provider_config( + async fn upsert_oauth_provider_config_guarded( &self, record: &UpsertOAuthProviderConfigRecord, - ) -> Result { + _ldap_exclusive: bool, + force_disable: bool, + locked_users_snapshot: usize, + ) -> Result { record.validate()?; let mut items = self.items.write().expect("oauth provider repository lock"); let now = Self::now_unix_secs(); let existing = items.get(&record.provider_type).cloned(); + if !force_disable + && locked_users_snapshot > 0 + && existing + .as_ref() + .is_some_and(|provider| provider.is_enabled && !record.is_enabled) + { + return Ok( + UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { + affected_count: locked_users_snapshot, + }, + ); + } let created_at = existing .as_ref() .and_then(|item| item.created_at_unix_ms) @@ -106,13 +121,34 @@ impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository { .with_timestamps(created_at, now); items.insert(record.provider_type.clone(), item.clone()); - Ok(item) + Ok(UpsertOAuthProviderConfigOutcome::Upserted(item)) } - async fn delete_oauth_provider_config( + async fn compare_and_swap_oauth_provider_client_secret( &self, provider_type: &str, + expected: &str, + replacement: &str, ) -> Result { + let mut items = self.items.write().expect("oauth provider repository lock"); + let Some(item) = items.get_mut(provider_type) else { + return Ok(false); + }; + if item.client_secret_encrypted.as_deref() != Some(expected) { + return Ok(false); + } + item.client_secret_encrypted = Some(replacement.to_string()); + Ok(true) + } + + async fn delete_oauth_provider_config_if_unlinked( + &self, + provider_type: &str, + has_links_snapshot: bool, + ) -> Result { + if has_links_snapshot { + return Ok(false); + } let mut items = self.items.write().expect("oauth provider repository lock"); Ok(items.remove(provider_type).is_some()) } @@ -123,7 +159,8 @@ mod tests { use super::InMemoryOAuthProviderRepository; use crate::repository::oauth_providers::{ EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository, - StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord, + StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome, + UpsertOAuthProviderConfigRecord, }; fn sample_provider(provider_type: &str) -> StoredOAuthProviderConfig { @@ -138,19 +175,31 @@ mod tests { } fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord { + let is_custom_oidc = provider_type.starts_with("custom_oidc"); + let endpoint_host = if is_custom_oidc { + "idp.example".to_string() + } else { + format!("{provider_type}.example.com") + }; UpsertOAuthProviderConfigRecord { provider_type: provider_type.to_string(), display_name: format!("{provider_type} display"), client_id: format!("{provider_type}-client"), client_secret_encrypted: EncryptedSecretUpdate::Preserve, - authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")), - token_url_override: Some(format!("https://{provider_type}.example.com/token")), - userinfo_url_override: None, + authorization_url_override: Some(format!("https://{endpoint_host}/auth")), + token_url_override: Some(format!("https://{endpoint_host}/token")), + userinfo_url_override: is_custom_oidc + .then(|| format!("https://{endpoint_host}/userinfo")), scopes: Some(vec!["openid".to_string(), "profile".to_string()]), redirect_uri: format!("https://{provider_type}.example.com/redirect"), frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(), attribute_mapping: Some(serde_json::json!({"email": "email"})), - extra_config: Some(serde_json::json!({"team": true})), + extra_config: is_custom_oidc.then(|| { + serde_json::json!({ + "allowed_domains": [endpoint_host], + "team": true, + }) + }), icon_url: None, is_enabled: true, } @@ -171,28 +220,106 @@ mod tests { assert_eq!(listed[0].provider_type, "github"); assert_eq!(listed[1].provider_type, "linuxdo"); - let created = repository - .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { - client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()), - ..sample_upsert("google") - }) + let UpsertOAuthProviderConfigOutcome::Upserted(created) = repository + .upsert_oauth_provider_config_guarded( + &UpsertOAuthProviderConfigRecord { + client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()), + ..sample_upsert("custom_oidc") + }, + false, + false, + 0, + ) .await - .expect("create should succeed"); + .expect("create should succeed") + else { + panic!("create unexpectedly required confirmation"); + }; assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1")); - let updated = repository - .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { - client_secret_encrypted: EncryptedSecretUpdate::Clear, - ..sample_upsert("google") - }) + let UpsertOAuthProviderConfigOutcome::Upserted(updated) = repository + .upsert_oauth_provider_config_guarded( + &UpsertOAuthProviderConfigRecord { + client_secret_encrypted: EncryptedSecretUpdate::Clear, + ..sample_upsert("custom_oidc") + }, + false, + false, + 0, + ) .await - .expect("update should succeed"); + .expect("update should succeed") + else { + panic!("update unexpectedly required confirmation"); + }; assert!(updated.client_secret_encrypted.is_none()); let deleted = repository - .delete_oauth_provider_config("google") + .delete_oauth_provider_config_if_unlinked("custom_oidc", false) .await .expect("delete should succeed"); assert!(deleted); } + + #[tokio::test] + async fn client_secret_cas_preserves_concurrent_non_secret_fields_and_timestamp() { + let repository = InMemoryOAuthProviderRepository::default(); + repository + .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { + client_secret_encrypted: EncryptedSecretUpdate::Set("legacy-secret".to_string()), + ..sample_upsert("custom_oidc") + }) + .await + .expect("provider should create"); + + let concurrent = repository + .upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord { + display_name: "concurrent display update".to_string(), + client_secret_encrypted: EncryptedSecretUpdate::Preserve, + ..sample_upsert("custom_oidc") + }) + .await + .expect("non-secret update should persist"); + assert!(repository + .compare_and_swap_oauth_provider_client_secret( + "custom_oidc", + "legacy-secret", + "record-bound-v2", + ) + .await + .expect("secret CAS should execute")); + + let migrated = repository + .get_oauth_provider_config("custom_oidc") + .await + .expect("provider should read") + .expect("provider should exist"); + assert_eq!(migrated.display_name, "concurrent display update"); + assert_eq!( + migrated.updated_at_unix_secs, + concurrent.updated_at_unix_secs + ); + assert_eq!( + migrated.client_secret_encrypted.as_deref(), + Some("record-bound-v2") + ); + assert!(!repository + .compare_and_swap_oauth_provider_client_secret( + "custom_oidc", + "legacy-secret", + "must-not-win", + ) + .await + .expect("stale CAS should execute")); + assert_eq!( + repository + .get_oauth_provider_config("custom_oidc") + .await + .expect("provider should read") + .expect("provider should exist") + .client_secret_encrypted + .as_deref(), + Some("record-bound-v2") + ); + } } diff --git a/crates/aether-data/runtime/src/repository/oauth_providers/mod.rs b/crates/aether-data/runtime/src/repository/oauth_providers/mod.rs index 65ed56302..f6258518b 100644 --- a/crates/aether-data/runtime/src/repository/oauth_providers/mod.rs +++ b/crates/aether-data/runtime/src/repository/oauth_providers/mod.rs @@ -1,8 +1,10 @@ mod memory; pub use aether_data_contracts::repository::oauth_providers::{ - EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository, - OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord, + validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config, + validate_oauth_redirect_uri, EncryptedSecretUpdate, OAuthProviderReadRepository, + OAuthProviderRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig, + UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::MysqlOAuthProviderRepository; diff --git a/crates/aether-data/runtime/src/repository/provider_catalog/memory.rs b/crates/aether-data/runtime/src/repository/provider_catalog/memory.rs index f1d91ef7e..988b7bf02 100644 --- a/crates/aether-data/runtime/src/repository/provider_catalog/memory.rs +++ b/crates/aether-data/runtime/src/repository/provider_catalog/memory.rs @@ -7,10 +7,12 @@ use serde_json::{json, Map, Value}; use super::{ ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, - ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate, - ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, - ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, - ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot, + ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate, + ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery, + ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, + ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, + ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, + ProviderCatalogReadRepository, ProviderCatalogSnapshot, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, @@ -454,6 +456,44 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository { Ok(stored.clone()) } + async fn compare_and_swap_provider_config( + &self, + update: &ProviderCatalogProviderConfigCasUpdate, + ) -> Result { + let mut index = self + .index + .write() + .expect("provider catalog repository lock"); + let Some(provider) = index.providers.get_mut(&update.provider_id) else { + return Ok(false); + }; + if provider.config != update.expected_config { + return Ok(false); + } + provider.config = update.config.clone(); + provider.updated_at_unix_secs = Some(current_unix_secs()); + Ok(true) + } + + async fn compare_and_swap_provider_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let mut index = self + .index + .write() + .expect("provider catalog repository lock"); + let Some(provider) = index.providers.get_mut(&update.record_id) else { + return Ok(false); + }; + if provider.proxy != update.expected_proxy { + return Ok(false); + } + provider.proxy = update.proxy.clone(); + provider.updated_at_unix_secs = Some(current_unix_secs()); + Ok(true) + } + async fn delete_provider(&self, provider_id: &str) -> Result { let mut index = self .index @@ -504,6 +544,25 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository { Ok(stored.clone()) } + async fn compare_and_swap_endpoint_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let mut index = self + .index + .write() + .expect("provider catalog repository lock"); + let Some(endpoint) = index.endpoints.get_mut(&update.record_id) else { + return Ok(false); + }; + if endpoint.proxy != update.expected_proxy { + return Ok(false); + } + endpoint.proxy = update.proxy.clone(); + endpoint.updated_at_unix_secs = Some(current_unix_secs()); + Ok(true) + } + async fn delete_endpoint(&self, endpoint_id: &str) -> Result { let mut index = self .index @@ -548,6 +607,52 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository { Ok(stored.clone()) } + async fn compare_and_swap_key_proxy( + &self, + update: &ProviderCatalogProxyCasUpdate, + ) -> Result { + let mut index = self + .index + .write() + .expect("provider catalog repository lock"); + let Some(key) = index.keys.get_mut(&update.record_id) else { + return Ok(false); + }; + if key.proxy != update.expected_proxy { + return Ok(false); + } + key.proxy = update.proxy.clone(); + key.updated_at_unix_secs = Some(current_unix_secs()); + Ok(true) + } + + async fn compare_and_swap_key_credentials( + &self, + update: &ProviderCatalogKeyCredentialsCasUpdate, + ) -> Result { + if update.key_id.trim().is_empty() || update.expected_provider_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "provider catalog key credential CAS requires key_id and provider_id".to_string(), + )); + } + let mut index = self + .index + .write() + .expect("provider catalog repository lock"); + let Some(key) = index.keys.get_mut(&update.key_id) else { + return Ok(false); + }; + if key.provider_id != update.expected_provider_id + || key.encrypted_api_key != update.expected_encrypted_api_key + || key.encrypted_auth_config != update.expected_encrypted_auth_config + { + return Ok(false); + } + key.encrypted_api_key = update.encrypted_api_key.clone(); + key.encrypted_auth_config = update.encrypted_auth_config.clone(); + Ok(true) + } + async fn compare_and_update_key_admin_state( &self, update: &ProviderCatalogKeyAdminCasUpdate, @@ -882,40 +987,11 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository { Ok(true) } - async fn update_key_oauth_credentials( - &self, - key_id: &str, - encrypted_api_key: &str, - encrypted_auth_config: Option<&str>, - expires_at_unix_secs: Option, - ) -> Result { - if encrypted_api_key.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "provider catalog oauth api_key is empty".to_string(), - )); - } - - let mut index = self - .index - .write() - .expect("provider catalog repository lock"); - let Some(key) = index.keys.get_mut(key_id) else { - return Ok(false); - }; - - key.encrypted_api_key = Some(encrypted_api_key.to_string()); - key.encrypted_auth_config = encrypted_auth_config.map(ToOwned::to_owned); - key.expires_at_unix_secs = expires_at_unix_secs; - key.updated_at_unix_secs = Some(current_unix_secs()); - Ok(true) - } - async fn update_key_oauth_runtime_state( &self, key_id: &str, oauth_invalid_at_unix_secs: Option, oauth_invalid_reason: Option<&str>, - encrypted_auth_config_update: Option<&str>, updated_at_unix_secs: Option, ) -> Result { let mut index = self @@ -928,9 +1004,6 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository { key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs; key.oauth_invalid_reason = oauth_invalid_reason.map(ToOwned::to_owned); - if let Some(encrypted_auth_config) = encrypted_auth_config_update { - key.encrypted_auth_config = Some(encrypted_auth_config.to_string()); - } key.updated_at_unix_secs = Some(updated_at_unix_secs.unwrap_or_else(current_unix_secs)); Ok(true) } @@ -1576,15 +1649,15 @@ mod tests { } #[tokio::test] - async fn updates_oauth_credentials_for_existing_key() { + async fn unfenced_oauth_runtime_state_update_preserves_credentials() { let repository = InMemoryProviderCatalogReadRepository::seed( vec![sample_provider("provider-1")], vec![sample_endpoint("endpoint-1", "provider-1")], vec![sample_key("key-1", "provider-1") .with_transport_fields( None, - "ciphertext-placeholder".to_string(), - Some("ciphertext-auth-1".to_string()), + "ciphertext-api".to_string(), + Some("ciphertext-auth".to_string()), None, None, None, @@ -1596,29 +1669,26 @@ mod tests { ); assert!(repository - .update_key_oauth_credentials( - "key-1", - "ciphertext-updated-token", - Some("ciphertext-auth-2"), - Some(4_102_444_800), - ) + .update_key_oauth_runtime_state("key-1", Some(123), Some("refresh failed"), Some(456),) .await - .expect("update should succeed")); + .expect("runtime state should update")); let stored = repository .list_keys_by_ids(&["key-1".to_string()]) .await - .expect("keys should read"); - assert_eq!(stored.len(), 1); + .expect("key should read") + .pop() + .expect("key should exist"); + assert_eq!(stored.encrypted_api_key.as_deref(), Some("ciphertext-api")); assert_eq!( - stored[0].encrypted_api_key.as_deref(), - Some("ciphertext-updated-token") + stored.encrypted_auth_config.as_deref(), + Some("ciphertext-auth") ); + assert_eq!(stored.oauth_invalid_at_unix_secs, Some(123)); assert_eq!( - stored[0].encrypted_auth_config.as_deref(), - Some("ciphertext-auth-2") + stored.oauth_invalid_reason.as_deref(), + Some("refresh failed") ); - assert_eq!(stored[0].expires_at_unix_secs, Some(4_102_444_800)); } #[tokio::test] diff --git a/crates/aether-data/runtime/src/repository/provider_catalog/mod.rs b/crates/aether-data/runtime/src/repository/provider_catalog/mod.rs index d89d2c3cb..5de56b4e8 100644 --- a/crates/aether-data/runtime/src/repository/provider_catalog/mod.rs +++ b/crates/aether-data/runtime/src/repository/provider_catalog/mod.rs @@ -3,11 +3,12 @@ mod memory; #[allow(unused_imports)] pub(crate) use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, - ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate, - ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, + ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate, + ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, - ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot, + ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate, + ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot, ProviderCatalogUpstreamMetadataNamespaceExpectation, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, diff --git a/crates/aether-data/runtime/src/repository/provider_oauth.rs b/crates/aether-data/runtime/src/repository/provider_oauth.rs index a7a5abf6b..92652b928 100644 --- a/crates/aether-data/runtime/src/repository/provider_oauth.rs +++ b/crates/aether-data/runtime/src/repository/provider_oauth.rs @@ -1,3 +1,5 @@ +use sha2::{Digest, Sha256}; + const KIRO_DEVICE_AUTH_SESSION_PREFIX: &str = "device_auth_session:"; const PROVIDER_OAUTH_BATCH_TASK_PREFIX: &str = "provider_oauth_batch_task:"; const PROVIDER_OAUTH_STATE_PREFIX: &str = "provider_oauth_state:"; @@ -8,7 +10,11 @@ pub const PROVIDER_OAUTH_STATE_TTL_SECS: u64 = 600; #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct StoredAdminProviderOAuthDeviceSession { + pub session_id: String, pub provider_id: String, + pub initiated_by_user_id: String, + pub initiated_by_session_id: Option, + pub initiated_by_management_token_id: Option, pub region: String, pub client_id: String, pub client_secret: String, @@ -34,28 +40,54 @@ pub struct StoredAdminProviderOAuthDeviceSession { pub error_msg: Option, } -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct StoredAdminProviderOAuthState { + pub nonce: String, pub key_id: String, pub provider_id: String, pub provider_type: String, pub pkce_verifier: Option, #[serde(default)] pub expected_encrypted_auth_config: Option, + pub initiated_by_user_id: String, + #[serde(default)] + pub initiated_by_session_id: Option, + #[serde(default)] + pub initiated_by_management_token_id: Option, + pub created_at: u64, } pub fn provider_oauth_device_session_storage_key(session_id: &str) -> String { format!("{KIRO_DEVICE_AUTH_SESSION_PREFIX}{session_id}") } +pub fn provider_oauth_device_session_secret_purpose(session_id: &str) -> String { + let storage_key = provider_oauth_device_session_storage_key(session_id); + format!( + "provider-oauth-device-session:sha256:{:x}", + Sha256::digest(storage_key.as_bytes()) + ) +} + pub fn provider_oauth_state_storage_key(nonce: &str) -> String { - format!("{PROVIDER_OAUTH_STATE_PREFIX}{nonce}") + format!( + "{PROVIDER_OAUTH_STATE_PREFIX}sha256:{:x}", + Sha256::digest(nonce.as_bytes()) + ) } pub fn provider_oauth_batch_task_storage_key(task_id: &str) -> String { format!("{PROVIDER_OAUTH_BATCH_TASK_PREFIX}{task_id}") } +pub fn provider_oauth_batch_task_secret_purpose(task_id: &str) -> String { + let storage_key = provider_oauth_batch_task_storage_key(task_id); + format!( + "provider-oauth-batch-task:sha256:{:x}", + Sha256::digest(storage_key.as_bytes()) + ) +} + pub fn build_provider_oauth_batch_task_status_payload( provider_id: &str, state: &serde_json::Map, @@ -142,7 +174,8 @@ pub fn build_provider_oauth_batch_task_status_payload( #[cfg(test)] mod tests { use super::{ - build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key, + build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_secret_purpose, + provider_oauth_batch_task_storage_key, provider_oauth_device_session_secret_purpose, provider_oauth_device_session_storage_key, provider_oauth_state_storage_key, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, PROVIDER_OAUTH_STATE_TTL_SECS, @@ -155,14 +188,23 @@ mod tests { provider_oauth_device_session_storage_key("session-123"), "device_auth_session:session-123" ); - assert_eq!( - provider_oauth_state_storage_key("nonce-123"), - "provider_oauth_state:nonce-123" - ); + let first_purpose = provider_oauth_device_session_secret_purpose("session-123"); + let second_purpose = provider_oauth_device_session_secret_purpose("session-456"); + assert!(first_purpose.starts_with("provider-oauth-device-session:sha256:")); + assert!(!first_purpose.contains("session-123")); + assert_ne!(first_purpose, second_purpose); + let state_key = provider_oauth_state_storage_key("nonce-123"); + assert!(state_key.starts_with("provider_oauth_state:sha256:")); + assert!(!state_key.contains("nonce-123")); assert_eq!( provider_oauth_batch_task_storage_key("task-123"), "provider_oauth_batch_task:task-123" ); + let first_task_purpose = provider_oauth_batch_task_secret_purpose("task-123"); + let second_task_purpose = provider_oauth_batch_task_secret_purpose("task-456"); + assert!(first_task_purpose.starts_with("provider-oauth-batch-task:sha256:")); + assert!(!first_task_purpose.contains("task-123")); + assert_ne!(first_task_purpose, second_task_purpose); } #[test] diff --git a/crates/aether-data/runtime/src/repository/proxy_nodes/memory.rs b/crates/aether-data/runtime/src/repository/proxy_nodes/memory.rs index 4d8373e5f..9416ec4e6 100644 --- a/crates/aether-data/runtime/src/repository/proxy_nodes/memory.rs +++ b/crates/aether-data/runtime/src/repository/proxy_nodes/memory.rs @@ -10,13 +10,14 @@ use super::log_reported_tunnel_error_event; use crate::DataLayerError; use aether_data_contracts::repository::proxy_nodes::{ bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, - normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security, - reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation, - ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary, - ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation, - ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation, - ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent, - StoredProxyNodeMetricsBucket, TunnelMetricsSample, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, + merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata, + normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, + ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, + ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository, + ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, + ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, + StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelMetricsSample, + PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR, }; #[derive(Debug, Default)] @@ -357,6 +358,42 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { Ok(updated) } + async fn compare_and_set_proxy_password( + &self, + node_id: &str, + expected: &str, + replacement: &str, + ) -> Result { + let mut nodes = self.nodes.write().expect("proxy node repository lock"); + let Some(node) = nodes.get_mut(node_id) else { + return Ok(false); + }; + if node.proxy_password.as_deref() != Some(expected) { + return Ok(false); + } + node.proxy_password = Some(replacement.to_string()); + node.updated_at_unix_secs = Self::now_unix_secs(); + Ok(true) + } + + async fn compare_and_set_proxy_metadata( + &self, + node_id: &str, + expected: &serde_json::Value, + replacement: &serde_json::Value, + ) -> Result { + let mut nodes = self.nodes.write().expect("proxy node repository lock"); + let Some(node) = nodes.get_mut(node_id) else { + return Ok(false); + }; + if node.proxy_metadata.as_ref() != Some(expected) { + return Ok(false); + } + node.proxy_metadata = Some(replacement.clone()); + node.updated_at_unix_secs = Self::now_unix_secs(); + Ok(true) + } + async fn create_manual_node( &self, mutation: &ProxyNodeManualCreateMutation, @@ -369,9 +406,14 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { return Err(Self::duplicate_proxy_node_error(existing)); } + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| Uuid::new_v4().to_string()); + if let Some(existing) = nodes.get(&node_id) { + return Err(proxy_node_id_in_use_error(existing)); + } let now = Self::now_unix_secs(); let node = StoredProxyNode::new( - Uuid::new_v4().to_string(), + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -466,18 +508,35 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { ) -> Result { let mut nodes = self.nodes.write().expect("proxy node repository lock"); let now = Self::now_unix_secs(); - let normalized_proxy_metadata = normalize_proxy_metadata( - mutation.proxy_metadata.as_ref(), - mutation.proxy_version.as_deref(), + let normalized_proxy_metadata = merge_proxy_metadata_for_registration( + None, + normalize_proxy_metadata( + mutation.proxy_metadata.as_ref(), + mutation.proxy_version.as_deref(), + ), ); if let Some(existing_id) = nodes .iter() - .find(|(_, node)| { + .filter(|(_, node)| { !node.is_manual && node.ip == mutation.ip && node.port == mutation.port }) + .min_by(|(_, left), (_, right)| { + left.created_at_unix_ms + .unwrap_or(u64::MAX) + .cmp(&right.created_at_unix_ms.unwrap_or(u64::MAX)) + .then(left.id.cmp(&right.id)) + }) .map(|(node_id, _)| node_id.clone()) { + if let Some(requested_id) = requested_proxy_node_id(mutation.node_id.as_deref())? { + if requested_id != existing_id { + return Err(proxy_node_registration_identity_error( + &requested_id, + &existing_id, + )); + } + } let node = nodes .get_mut(&existing_id) .expect("existing proxy node should be present"); @@ -505,9 +564,10 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency { node.estimated_max_concurrency = Some(estimated_max_concurrency); } - if let Some(proxy_metadata) = normalized_proxy_metadata { - node.proxy_metadata = Some(proxy_metadata); - } + node.proxy_metadata = merge_proxy_metadata_for_registration( + node.proxy_metadata.as_ref(), + normalized_proxy_metadata, + ); if node.created_at_unix_ms.is_none() { node.created_at_unix_ms = now; } @@ -515,8 +575,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { return Ok(node.clone()); } + let node_id = requested_proxy_node_id(mutation.node_id.as_deref())? + .unwrap_or_else(|| Uuid::new_v4().to_string()); + if let Some(existing) = nodes.get(&node_id) { + return Err(proxy_node_id_in_use_error(existing)); + } let mut node = StoredProxyNode::new( - Uuid::new_v4().to_string(), + node_id, mutation.name.clone(), mutation.ip.clone(), mutation.port, @@ -555,11 +620,18 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { &self, mutation: &ProxyNodeHeartbeatMutation, ) -> Result, DataLayerError> { + let mut nodes = self.nodes.write().expect("proxy node repository lock"); let (node, sample, now_unix_secs) = { - let mut nodes = self.nodes.write().expect("proxy node repository lock"); let Some(node) = nodes.get_mut(&mutation.node_id) else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } if !node.tunnel_mode { return Err(DataLayerError::InvalidInput( "non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode" @@ -587,14 +659,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { if let Some(value) = mutation.avg_latency_ms { node.avg_latency_ms = Some(value); } - let normalized_proxy_metadata = normalize_proxy_metadata( + let normalized_proxy_metadata = normalize_heartbeat_proxy_metadata( + previous_proxy_metadata.as_ref(), mutation.proxy_metadata.as_ref(), mutation.proxy_version.as_deref(), ); - let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security( - previous_proxy_metadata.as_ref(), - normalized_proxy_metadata, - ); if let Some(value) = normalized_proxy_metadata { node.proxy_metadata = Some(value); } @@ -671,6 +740,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { }); } } + drop(nodes); Ok(Some(node)) } @@ -686,6 +756,12 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { if !node.is_manual { return Ok(false); } + let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else { + return Ok(false); + }; + if expected_generation != node.tunnel_generation { + return Ok(false); + } node.total_requests += mutation.total_requests_delta.max(0); node.failed_requests += mutation.failed_requests_delta.max(0); @@ -703,6 +779,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { let Some(node) = nodes.get_mut(&mutation.node_id) else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } let event_time = mutation .observed_at_unix_secs @@ -776,11 +859,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { } async fn delete_node(&self, node_id: &str) -> Result, DataLayerError> { - let removed = self - .nodes - .write() - .expect("proxy node repository lock") - .remove(node_id); + // Keep the parent lock until all child state is removed. Registration + // also takes this lock, so the same id cannot be recreated between the + // parent delete and cleanup of its events or metrics. + let mut nodes = self.nodes.write().expect("proxy node repository lock"); + let removed = nodes.remove(node_id); if removed.is_some() { self.events .write() @@ -795,6 +878,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { .expect("proxy node repository lock") .retain(|(metric_node_id, _), _| metric_node_id != node_id); } + drop(nodes); Ok(removed) } @@ -806,6 +890,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { let Some(node) = nodes.get_mut(&mutation.node_id) else { return Ok(None); }; + if mutation + .expected_tunnel_generation + .as_deref() + .is_some_and(|expected| expected != node.tunnel_generation) + { + return Ok(None); + } if node.is_manual { return Err(DataLayerError::InvalidInput( "手动节点不支持远程配置下发".to_string(), @@ -886,6 +977,31 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository { } } +fn requested_proxy_node_id(value: Option<&str>) -> Result, DataLayerError> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_empty() || value.trim() != value { + return Err(DataLayerError::InvalidInput( + "proxy node id must be non-empty and unpadded".to_string(), + )); + } + Ok(Some(value.to_string())) +} + +fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node registration identity changed: requested {requested_id}, existing {existing_id}" + )) +} + +fn proxy_node_id_in_use_error(node: &StoredProxyNode) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "proxy node id is already in use: {} ({}:{})", + node.id, node.ip, node.port + )) +} + #[cfg(test)] mod tests { use super::InMemoryProxyNodeRepository; @@ -894,7 +1010,7 @@ mod tests { ProxyNodeRemoteConfigMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent, }; - use serde_json::json; + use serde_json::{json, Value}; fn sample_node() -> StoredProxyNode { StoredProxyNode::new( @@ -937,6 +1053,7 @@ mod tests { let heartbeat = repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: "node-1".to_string(), + expected_tunnel_generation: None, heartbeat_interval: Some(45), active_connections: Some(5), total_requests_delta: Some(8), @@ -944,7 +1061,13 @@ mod tests { failed_requests_delta: Some(2), dns_failures_delta: Some(1), stream_errors_delta: Some(3), - proxy_metadata: Some(json!({"arch": "arm64"})), + proxy_metadata: Some(json!({ + "arch": "arm64", + "tunnel_security": { + "mode": "disabled", + "encryption_key": "attacker-controlled" + } + })), proxy_version: Some("1.2.3".to_string()), }) .await @@ -966,10 +1089,16 @@ mod tests { .and_then(|value| value.as_str()), Some("1.2.3") ); + assert!(heartbeat + .proxy_metadata + .as_ref() + .and_then(|value| value.get("tunnel_security")) + .is_none()); let stale = repository .update_tunnel_status(&ProxyNodeTunnelStatusMutation { node_id: "node-1".to_string(), + expected_tunnel_generation: None, connected: false, conn_count: 0, detail: None, @@ -994,6 +1123,7 @@ mod tests { let updated = repository .update_tunnel_status(&ProxyNodeTunnelStatusMutation { node_id: "node-1".to_string(), + expected_tunnel_generation: None, connected: false, conn_count: 0, detail: None, @@ -1060,6 +1190,96 @@ mod tests { assert_eq!(events[0].detail.as_deref(), Some("newer")); } + #[tokio::test] + async fn delete_cleans_child_state_before_same_id_can_be_reused() { + let old_node = sample_node(); + let old_generation = old_node.tunnel_generation.clone(); + let repository = InMemoryProxyNodeRepository::seed_with_events( + vec![old_node], + vec![StoredProxyNodeEvent { + id: 1, + node_id: "node-1".to_string(), + event_type: "connected".to_string(), + detail: Some("old incarnation".to_string()), + event_metadata: None, + created_at_unix_ms: Some(1_710_000_000), + }], + ); + + repository + .apply_heartbeat(&ProxyNodeHeartbeatMutation { + node_id: "node-1".to_string(), + expected_tunnel_generation: Some(old_generation.clone()), + heartbeat_interval: None, + active_connections: Some(1), + total_requests_delta: None, + avg_latency_ms: None, + failed_requests_delta: None, + dns_failures_delta: None, + stream_errors_delta: None, + proxy_metadata: Some(json!({ + "tunnel_metrics": { + "connect_errors": 0, + "disconnects": 0, + "error_events_total": 0, + "ws_in_bytes": 0, + "ws_out_bytes": 0, + "ws_in_frames": 0, + "ws_out_frames": 0, + "heartbeat_rtt_last_ms": 1 + } + })), + proxy_version: None, + }) + .await + .expect("heartbeat should create metric buckets") + .expect("old node should exist"); + + repository + .delete_node("node-1") + .await + .expect("delete should succeed") + .expect("old node should be removed"); + + let replacement = repository + .register_node(&ProxyNodeRegistrationMutation { + node_id: Some("node-1".to_string()), + name: "replacement".to_string(), + ip: "127.0.0.2".to_string(), + port: 7002, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: None, + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }) + .await + .expect("same id should be reusable after delete"); + assert_ne!(replacement.tunnel_generation, old_generation); + + assert!(repository + .list_proxy_node_events("node-1", 10) + .await + .expect("events should read") + .is_empty()); + for step in [ + crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneMinute, + crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneHour, + ] { + assert!(repository + .list_proxy_node_metrics("node-1", step, 0, u64::MAX, 10) + .await + .expect("metrics should read") + .is_empty()); + } + } + #[tokio::test] async fn resets_stale_tunnel_statuses_without_touching_manual_nodes() { let mut stale_tunnel = sample_node(); @@ -1101,12 +1321,127 @@ mod tests { assert_eq!(manual.active_connections, 4); } + #[tokio::test] + async fn registration_rejects_rebinding_existing_endpoint_to_different_node_id() { + let repository = InMemoryProxyNodeRepository::default(); + let mutation = ProxyNodeRegistrationMutation { + node_id: Some("stable-node-id".to_string()), + name: "stable-node".to_string(), + ip: "127.0.0.9".to_string(), + port: 7009, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({"secret_marker": "first"})), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }; + let registered = repository + .register_node(&mutation) + .await + .expect("initial registration should succeed"); + assert_eq!(registered.id, "stable-node-id"); + + let mut conflicting = mutation; + conflicting.node_id = Some("replacement-node-id".to_string()); + conflicting.proxy_metadata = Some(json!({"secret_marker": "replacement"})); + assert!(repository.register_node(&conflicting).await.is_err()); + let persisted = repository + .find_proxy_node("stable-node-id") + .await + .expect("stable node should read") + .expect("stable node should remain"); + assert_eq!( + persisted + .proxy_metadata + .as_ref() + .and_then(|value| value.get("secret_marker")), + Some(&json!("first")) + ); + } + + #[tokio::test] + async fn registration_preserves_omitted_security_and_allows_rotation() { + let repository = InMemoryProxyNodeRepository::default(); + let first_mutation = ProxyNodeRegistrationMutation { + node_id: Some("registration-security-node".to_string()), + name: "registration-security-node".to_string(), + ip: "127.0.0.70".to_string(), + port: 7070, + region: None, + heartbeat_interval: 30, + active_connections: None, + total_requests: None, + avg_latency_ms: None, + hardware_info: None, + estimated_max_concurrency: None, + proxy_metadata: Some(json!({ + "version": "1.0.0", + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old" + } + })), + proxy_version: None, + registered_by: None, + tunnel_mode: true, + }; + let first = repository + .register_node(&first_mutation) + .await + .expect("first registration should succeed"); + + let mut refreshed_mutation = first_mutation.clone(); + refreshed_mutation.name = "registration-security-node-refreshed".to_string(); + refreshed_mutation.proxy_metadata = Some(json!({"runtime": "refreshed"})); + refreshed_mutation.proxy_version = Some("2.0.0".to_string()); + let refreshed = repository + .register_node(&refreshed_mutation) + .await + .expect("metadata-only re-registration should succeed"); + assert_eq!(refreshed.id, first.id); + assert_eq!( + refreshed + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old") + ); + + let mut rotated_mutation = refreshed_mutation; + rotated_mutation.proxy_metadata = Some(json!({ + "tunnel_security": { + "mode": "non_tls_required", + "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new" + } + })); + let rotated = repository + .register_node(&rotated_mutation) + .await + .expect("explicit security rotation should succeed"); + assert_eq!( + rotated + .proxy_metadata + .as_ref() + .and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted")) + .and_then(Value::as_str), + Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new") + ); + } + #[tokio::test] async fn registers_updates_config_and_unregisters_nodes() { let repository = InMemoryProxyNodeRepository::default(); let registered = repository .register_node(&ProxyNodeRegistrationMutation { + node_id: None, name: "proxy-01".to_string(), ip: "127.0.0.1".to_string(), port: 0, @@ -1131,6 +1466,7 @@ mod tests { let updated = repository .update_remote_config(&ProxyNodeRemoteConfigMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, node_name: Some("proxy-02".to_string()), allowed_ports: Some(vec![443, 8443]), log_level: Some("info".to_string()), @@ -1160,6 +1496,7 @@ mod tests { let after_upgrade = repository .apply_heartbeat(&ProxyNodeHeartbeatMutation { node_id: registered.id.clone(), + expected_tunnel_generation: None, heartbeat_interval: None, active_connections: Some(2), total_requests_delta: Some(1), diff --git a/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs b/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs index b6f11e675..c0282763d 100644 --- a/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs +++ b/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs @@ -59,7 +59,14 @@ impl RoutingGroupReadRepository for InMemoryRoutingGroupRepository { .values() .cloned() .collect::>(); - groups.sort_by(|left, right| left.name.cmp(&right.name).then(left.id.cmp(&right.id))); + groups.sort_by(|left, right| { + right + .enabled + .cmp(&left.enabled) + .then(left.sort_order.cmp(&right.sort_order)) + .then(left.name.cmp(&right.name)) + .then(left.id.cmp(&right.id)) + }); Ok(groups) } @@ -273,6 +280,7 @@ mod tests { description: None, enabled: true, is_system_default: true, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, @@ -320,6 +328,49 @@ mod tests { ); } + #[tokio::test] + async fn lists_enabled_groups_first_and_respects_sort_order() { + let repository = InMemoryRoutingGroupRepository::default(); + for (id, enabled, sort_order) in [ + ("disabled-first", false, 0), + ("enabled-second", true, 20), + ("enabled-first", true, 10), + ] { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: id.to_string(), + name: id.to_string(), + description: None, + enabled, + is_system_default: false, + sort_order, + config_json: json!({}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .expect("group should store"); + } + + let ids = repository + .list_routing_groups() + .await + .expect("groups should list") + .into_iter() + .map(|group| group.id) + .collect::>(); + assert_eq!( + ids, + vec![ + "enabled-first".to_string(), + "enabled-second".to_string(), + "disabled-first".to_string(), + ] + ); + } + #[tokio::test] async fn keeps_system_and_subject_defaults_unique() { let repository = InMemoryRoutingGroupRepository::default(); @@ -414,6 +465,7 @@ mod tests { description: None, enabled: true, is_system_default, + sort_order: 0, config_json: json!({}), version: 1, created_at: 1, diff --git a/crates/aether-data/runtime/src/repository/settlement/memory.rs b/crates/aether-data/runtime/src/repository/settlement/memory.rs index 15dbaee6a..a188493fe 100644 --- a/crates/aether-data/runtime/src/repository/settlement/memory.rs +++ b/crates/aether-data/runtime/src/repository/settlement/memory.rs @@ -1,12 +1,16 @@ use std::collections::BTreeMap; -use std::sync::{Arc, RwLock}; +use std::sync::{Arc, Mutex, RwLock}; use async_trait::async_trait; use super::{ plan_finite_wallet_debit, settlement_billable_cost_usd, - settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement, - UsageSettlementInput, SETTLEMENT_EPSILON_USD, + settlement_billing_status_for_usage_status, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput, + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, + ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation, + StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState, + UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD, }; use crate::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot}; use crate::DataLayerError; @@ -49,8 +53,11 @@ impl InMemorySettlementWalletStore { #[derive(Debug, Default)] pub struct InMemorySettlementRepository { wallets: InMemorySettlementWalletStore, + settlement_lock: Mutex<()>, provider_monthly_used: RwLock>, settlements: RwLock>, + cost_reservations: RwLock>, + request_admissions: RwLock>, } impl InMemorySettlementRepository { @@ -60,35 +67,321 @@ impl InMemorySettlementRepository { { Self { wallets: InMemorySettlementWalletStore::seeded(items), + settlement_lock: Mutex::new(()), provider_monthly_used: RwLock::new(BTreeMap::new()), settlements: RwLock::new(BTreeMap::new()), + cost_reservations: RwLock::new(BTreeMap::new()), + request_admissions: RwLock::new(BTreeMap::new()), } } pub fn from_wallet_repository(wallet_repository: Arc) -> Self { Self { wallets: InMemorySettlementWalletStore::Shared(wallet_repository), + settlement_lock: Mutex::new(()), provider_monthly_used: RwLock::new(BTreeMap::new()), settlements: RwLock::new(BTreeMap::new()), + cost_reservations: RwLock::new(BTreeMap::new()), + request_admissions: RwLock::new(BTreeMap::new()), } } } #[async_trait] impl SettlementWriteRepository for InMemorySettlementRepository { + async fn reserve_usage_policy_request( + &self, + input: ReserveUsagePolicyRequestInput, + ) -> Result { + input.validate()?; + let mut admissions = self + .request_admissions + .write() + .expect("usage policy request admission lock"); + if let Some(existing) = admissions.get_mut(&input.event_token) { + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + return Ok(ReserveUsagePolicyRequestOutcome::Conflict); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy event_token must keep its original admitted_at".to_string(), + )); + } + existing.retain_until_unix_secs = existing + .retain_until_unix_secs + .max(input.retain_until_unix_secs); + if existing.state == UsagePolicyRequestAdmissionState::Released { + return Ok(ReserveUsagePolicyRequestOutcome::AlreadyReleased); + } + return Ok(ReserveUsagePolicyRequestOutcome::Allowed); + } + + for (window_index, window) in input.windows.iter().enumerate() { + let used_requests = admissions + .values() + .filter(|admission| { + admission.state == UsagePolicyRequestAdmissionState::Active + && admission.subject_id == input.subject_id + && admission.admitted_at_unix_secs >= window.starts_at_unix_secs + && admission.admitted_at_unix_secs < window.ends_at_unix_secs + }) + .count() as u64; + if used_requests >= window.limit_requests { + return Ok(ReserveUsagePolicyRequestOutcome::Rejected { + window_index, + limit_requests: window.limit_requests, + used_requests, + }); + } + } + + admissions.insert( + input.event_token.clone(), + StoredUsagePolicyRequestAdmission { + request_id: input.request_id, + subject_id: input.subject_id, + event_token: input.event_token, + admitted_at_unix_secs: input.admitted_at_unix_secs, + retain_until_unix_secs: input.retain_until_unix_secs, + state: UsagePolicyRequestAdmissionState::Active, + released_at_unix_secs: None, + }, + ); + Ok(ReserveUsagePolicyRequestOutcome::Allowed) + } + + async fn release_usage_policy_request_admission( + &self, + input: ReleaseUsagePolicyRequestAdmissionInput, + ) -> Result, DataLayerError> { + input.validate()?; + let mut admissions = self + .request_admissions + .write() + .expect("usage policy request admission lock"); + let Some(admission) = admissions.get_mut(&input.event_token) else { + return Ok(None); + }; + if admission.request_id != input.request_id || admission.subject_id != input.subject_id { + return Ok(None); + } + if input.released_at_unix_secs < admission.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy released_at must not precede admitted_at".to_string(), + )); + } + if admission.state == UsagePolicyRequestAdmissionState::Active { + admission.state = UsagePolicyRequestAdmissionState::Released; + admission.released_at_unix_secs = Some(input.released_at_unix_secs); + } + Ok(Some(admission.clone())) + } + + async fn cleanup_usage_policy_request_admissions( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let mut admissions = self + .request_admissions + .write() + .expect("usage policy request admission lock"); + let tokens = admissions + .iter() + .filter(|(_, admission)| admission.retain_until_unix_secs <= now_unix_secs) + .map(|(token, _)| token.clone()) + .take(batch_size) + .collect::>(); + for token in &tokens { + admissions.remove(token); + } + Ok(tokens.len()) + } + + async fn reserve_usage_policy_cost( + &self, + input: ReserveUsagePolicyCostInput, + ) -> Result { + input.validate()?; + let mut reservations = self + .cost_reservations + .write() + .expect("usage policy cost reservation lock"); + let existing = reservations.get(&input.reservation_token).cloned(); + if let Some(existing) = existing.as_ref() { + if existing.request_id != input.request_id || existing.subject_id != input.subject_id { + return Ok(ReserveUsagePolicyCostOutcome::Conflict); + } + if existing.state != UsagePolicyCostReservationState::Reserved { + return Ok(ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: existing.state, + }); + } + if existing.admitted_at_unix_secs != input.admitted_at_unix_secs { + return Err(DataLayerError::InvalidInput( + "usage policy reservation_token must keep its original admitted_at".to_string(), + )); + } + } + + let target_reserved_cost_units = existing + .as_ref() + .map(|reservation| reservation.reserved_cost_units) + .unwrap_or(0) + .max(input.reserved_cost_units); + let target_reservation_expires_at_unix_secs = existing + .as_ref() + .map(|reservation| reservation.reservation_expires_at_unix_secs) + .unwrap_or(0) + .max(input.reservation_expires_at_unix_secs); + let target_retain_until_unix_secs = existing + .as_ref() + .map(|reservation| reservation.retain_until_unix_secs) + .unwrap_or(0) + .max(input.retain_until_unix_secs); + for (window_index, window) in input.windows.iter().enumerate() { + let used_cost_units = reservations + .values() + .filter(|reservation| { + reservation.reservation_token != input.reservation_token + && reservation.subject_id == input.subject_id + && reservation.admitted_at_unix_secs >= window.starts_at_unix_secs + && reservation.admitted_at_unix_secs < window.ends_at_unix_secs + }) + .try_fold(0_u64, |used, reservation| { + let cost_units = match reservation.state { + UsagePolicyCostReservationState::Reserved + if reservation.reservation_expires_at_unix_secs + > input.admitted_at_unix_secs => + { + reservation.reserved_cost_units + } + UsagePolicyCostReservationState::Finalized => { + reservation.actual_cost_units.unwrap_or(0) + } + UsagePolicyCostReservationState::Reserved + | UsagePolicyCostReservationState::Released => 0, + }; + used.checked_add(cost_units).ok_or_else(|| { + DataLayerError::UnexpectedValue( + "usage policy cost total overflowed".to_string(), + ) + }) + })?; + if used_cost_units + .checked_add(target_reserved_cost_units) + .is_none_or(|total| total > window.limit_cost_units) + { + return Ok(ReserveUsagePolicyCostOutcome::Rejected { + window_index, + limit_cost_units: window.limit_cost_units, + used_cost_units, + }); + } + } + + let previous_reserved_cost_units = existing + .as_ref() + .map(|reservation| reservation.reserved_cost_units) + .unwrap_or(0); + reservations.insert( + input.reservation_token.clone(), + StoredUsagePolicyCostReservation { + request_id: input.request_id, + subject_id: input.subject_id, + reservation_token: input.reservation_token, + admitted_at_unix_secs: input.admitted_at_unix_secs, + reserved_cost_units: target_reserved_cost_units, + actual_cost_units: None, + state: UsagePolicyCostReservationState::Reserved, + reservation_expires_at_unix_secs: target_reservation_expires_at_unix_secs, + retain_until_unix_secs: target_retain_until_unix_secs, + finalized_at_unix_secs: None, + }, + ); + Ok(ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: target_reserved_cost_units, + additional_reserved_cost_units: target_reserved_cost_units + .saturating_sub(previous_reserved_cost_units), + }) + } + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + input.validate()?; + let mut reservations = self + .cost_reservations + .write() + .expect("usage policy cost reservation lock"); + let Some(reservation) = reservations.get_mut(&input.reservation_token) else { + return Ok(None); + }; + if reservation.request_id != input.request_id || reservation.subject_id != input.subject_id + { + // The server-issued token selects the reservation. Never let mismatched audit + // identity fields mutate a reservation belonging to another request. + return Ok(None); + } + if reservation.state == UsagePolicyCostReservationState::Reserved { + reservation.state = input.terminal_state; + reservation.actual_cost_units = Some(input.actual_cost_units); + reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs); + } + Ok(Some(reservation.clone())) + } + + async fn cleanup_usage_policy_cost_reservations( + &self, + now_unix_secs: u64, + batch_size: usize, + ) -> Result { + if batch_size == 0 { + return Ok(0); + } + let mut reservations = self + .cost_reservations + .write() + .expect("usage policy cost reservation lock"); + let tokens = reservations + .iter() + .filter(|(_, reservation)| reservation.retain_until_unix_secs <= now_unix_secs) + .map(|(token, _)| token.clone()) + .take(batch_size) + .collect::>(); + for token in &tokens { + reservations.remove(token); + } + Ok(tokens.len()) + } + async fn settle_usage( &self, input: UsageSettlementInput, ) -> Result, DataLayerError> { input.validate()?; + // Mirror the SQL backends' row-lock transaction: a repeated or concurrent pending + // finalization for one request must observe the first committed snapshot instead of + // debiting the wallet twice. + let _settlement_guard = self + .settlement_lock + .lock() + .expect("usage settlement transaction lock"); + if let Some(existing) = self + .settlements + .read() + .expect("settlement snapshot lock") + .get(&input.request_id) + .cloned() + { + return Ok(Some(existing)); + } if input.billing_status != "pending" { - let existing = self - .settlements - .read() - .expect("settlement snapshot lock") - .get(&input.request_id) - .cloned(); - return Ok(Some(existing.unwrap_or(StoredUsageSettlement { + return Ok(Some(StoredUsageSettlement { request_id: input.request_id, wallet_id: None, billing_status: input.billing_status, @@ -100,12 +393,34 @@ impl SettlementWriteRepository for InMemorySettlementRepository { wallet_gift_balance_after: None, provider_monthly_used_usd: None, finalized_at_unix_secs: input.finalized_at_unix_secs, - }))); + })); } let mut final_billing_status = settlement_billing_status_for_usage_status(&input.status).to_string(); let billable_cost_usd = settlement_billable_cost_usd(&input); + let provider_monthly_after = if final_billing_status == "settled" { + input + .provider_id + .as_ref() + .map(|provider_id| { + let quotas = self + .provider_monthly_used + .read() + .expect("provider quota lock"); + let current = quotas.get(provider_id).copied().unwrap_or(0.0); + let next = current + input.actual_total_cost_usd; + if !current.is_finite() || current < 0.0 || !next.is_finite() { + return Err(DataLayerError::UnexpectedValue( + "provider monthly usage is invalid for settlement".to_string(), + )); + } + Ok((provider_id.clone(), next)) + }) + .transpose()? + } else { + None + }; let mut settlement = self.wallets.with_mut(|wallets| { let wallet_id = input .api_key_id @@ -148,6 +463,17 @@ impl SettlementWriteRepository for InMemorySettlementRepository { if let Some(wallet) = wallet { let before_recharge = wallet.balance; let before_gift = wallet.gift_balance; + let consumed_delta = if final_billing_status == "settled" { + billable_cost_usd + } else { + 0.0 + }; + validate_wallet_settlement_values( + before_recharge, + before_gift, + wallet.total_consumed, + consumed_delta, + )?; let before_total = before_recharge + before_gift; settlement.wallet_id = Some(wallet.id.clone()); settlement.wallet_balance_before = Some(before_total); @@ -155,18 +481,31 @@ impl SettlementWriteRepository for InMemorySettlementRepository { settlement.wallet_gift_balance_before = Some(before_gift); if final_billing_status == "settled" { + let total_consumed_after = wallet.total_consumed + billable_cost_usd; if wallet.limit_mode.eq_ignore_ascii_case("unlimited") { - wallet.total_consumed += billable_cost_usd; + validate_wallet_settlement_values( + before_recharge, + before_gift, + total_consumed_after, + 0.0, + )?; } else { let debit_plan = plan_finite_wallet_debit( before_recharge, before_gift, billable_cost_usd, ); - (wallet.balance, wallet.gift_balance) = + let (after_recharge, after_gift) = debit_plan.after_balances(before_recharge, before_gift); - wallet.total_consumed += billable_cost_usd; + validate_wallet_settlement_values( + after_recharge, + after_gift, + total_consumed_after, + 0.0, + )?; + (wallet.balance, wallet.gift_balance) = (after_recharge, after_gift); } + wallet.total_consumed = total_consumed_after; } settlement.wallet_recharge_balance_after = Some(wallet.balance); @@ -179,18 +518,17 @@ impl SettlementWriteRepository for InMemorySettlementRepository { settlement.billing_status = final_billing_status.clone(); } - settlement - }); + Ok(settlement) + })?; if final_billing_status == "settled" { - if let Some(provider_id) = input.provider_id { + if let Some((provider_id, next)) = provider_monthly_after { let mut quotas = self .provider_monthly_used .write() .expect("provider quota lock"); - let value = quotas.entry(provider_id).or_insert(0.0); - *value += input.actual_total_cost_usd; - settlement.provider_monthly_used_usd = Some(*value); + quotas.insert(provider_id, next); + settlement.provider_monthly_used_usd = Some(next); } } @@ -203,11 +541,514 @@ impl SettlementWriteRepository for InMemorySettlementRepository { } } +#[cfg(test)] +mod usage_policy_request_admission_tests { + use super::*; + use crate::repository::settlement::UsagePolicyRequestWindow; + + fn reserve( + request_id: &str, + event_token: &str, + admitted_at: u64, + limits: &[(u64, u64, u64)], + ) -> ReserveUsagePolicyRequestInput { + let windows = limits + .iter() + .map(|(start, end, limit)| UsagePolicyRequestWindow { + starts_at_unix_secs: *start, + ends_at_unix_secs: *end, + limit_requests: *limit, + }) + .collect::>(); + ReserveUsagePolicyRequestInput { + request_id: request_id.to_string(), + subject_id: "user-1".to_string(), + event_token: event_token.to_string(), + admitted_at_unix_secs: admitted_at, + retain_until_unix_secs: windows + .iter() + .map(|window| window.ends_at_unix_secs) + .max() + .unwrap_or(admitted_at + 1), + windows, + } + } + + #[tokio::test] + async fn any_rejected_window_prevents_the_admission_insert() { + let repository = InMemorySettlementRepository::default(); + repository + .reserve_usage_policy_request(reserve("request-1", "event-1", 100, &[(0, 1_000, 10)])) + .await + .unwrap(); + + assert_eq!( + repository + .reserve_usage_policy_request(reserve( + "request-2", + "event-2", + 101, + &[(0, 1_000, 2), (50, 200, 1)], + )) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::Rejected { + window_index: 1, + limit_requests: 1, + used_requests: 1, + } + ); + assert!(!repository + .request_admissions + .read() + .expect("admission lock") + .contains_key("event-2")); + } + + #[tokio::test] + async fn retries_are_idempotent_and_identity_or_timestamp_conflicts_are_explicit() { + let repository = InMemorySettlementRepository::default(); + let initial = reserve("request-1", "event-1", 100, &[(0, 1_000, 1)]); + assert_eq!( + repository + .reserve_usage_policy_request(initial.clone()) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::Allowed + ); + let mut extended = initial.clone(); + extended.retain_until_unix_secs = 2_000; + extended.windows[0].ends_at_unix_secs = 2_000; + assert_eq!( + repository + .reserve_usage_policy_request(extended) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::Allowed + ); + assert_eq!( + repository + .request_admissions + .read() + .expect("admission lock") + .get("event-1") + .expect("admission") + .retain_until_unix_secs, + 2_000 + ); + + let mut conflicting = initial.clone(); + conflicting.request_id = "other-request".to_string(); + assert_eq!( + repository + .reserve_usage_policy_request(conflicting) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::Conflict + ); + let mut changed_timestamp = initial; + changed_timestamp.admitted_at_unix_secs = 101; + assert!(matches!( + repository + .reserve_usage_policy_request(changed_timestamp) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + } + + #[tokio::test] + async fn released_admission_stops_counting_and_never_reactivates() { + let repository = InMemorySettlementRepository::default(); + let initial = reserve("request-1", "event-1", 100, &[(0, 1_000, 1)]); + repository + .reserve_usage_policy_request(initial.clone()) + .await + .unwrap(); + let released = repository + .release_usage_policy_request_admission(ReleaseUsagePolicyRequestAdmissionInput { + request_id: "request-1".to_string(), + subject_id: "user-1".to_string(), + event_token: "event-1".to_string(), + released_at_unix_secs: 101, + }) + .await + .unwrap() + .expect("released admission"); + assert_eq!(released.state, UsagePolicyRequestAdmissionState::Released); + + let released_again = repository + .release_usage_policy_request_admission(ReleaseUsagePolicyRequestAdmissionInput { + request_id: "request-1".to_string(), + subject_id: "user-1".to_string(), + event_token: "event-1".to_string(), + released_at_unix_secs: 999, + }) + .await + .unwrap() + .expect("released admission retry"); + assert_eq!(released_again.released_at_unix_secs, Some(101)); + + let mut retry = initial; + retry.windows[0].ends_at_unix_secs = 2_000; + retry.retain_until_unix_secs = 2_000; + assert_eq!( + repository + .reserve_usage_policy_request(retry) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::AlreadyReleased + ); + assert_eq!( + repository + .request_admissions + .read() + .expect("admission lock") + .get("event-1") + .expect("released admission") + .retain_until_unix_secs, + 2_000 + ); + assert_eq!( + repository + .reserve_usage_policy_request(reserve( + "request-2", + "event-2", + 102, + &[(0, 1_000, 1)], + )) + .await + .unwrap(), + ReserveUsagePolicyRequestOutcome::Allowed + ); + } + + #[tokio::test] + async fn cleanup_is_bounded_and_preserves_unexpired_tombstones() { + let repository = InMemorySettlementRepository::default(); + repository + .reserve_usage_policy_request(reserve("request-1", "event-1", 10, &[(0, 100, 10)])) + .await + .unwrap(); + repository + .reserve_usage_policy_request(reserve("request-2", "event-2", 11, &[(0, 200, 10)])) + .await + .unwrap(); + assert_eq!( + repository + .cleanup_usage_policy_request_admissions(99, 10) + .await + .unwrap(), + 0 + ); + assert_eq!( + repository + .cleanup_usage_policy_request_admissions(200, 1) + .await + .unwrap(), + 1 + ); + assert_eq!( + repository + .cleanup_usage_policy_request_admissions(200, 10) + .await + .unwrap(), + 1 + ); + } + + #[tokio::test] + async fn concurrent_admissions_do_not_oversell_capacity() { + let repository = Arc::new(InMemorySettlementRepository::default()); + let mut tasks = Vec::new(); + for index in 0..32 { + let repository = Arc::clone(&repository); + tasks.push(tokio::spawn(async move { + repository + .reserve_usage_policy_request(reserve( + &format!("request-{index}"), + &format!("event-{index}"), + 100, + &[(0, 1_000, 5)], + )) + .await + .unwrap() + })); + } + let mut allowed = 0; + for task in tasks { + if task.await.unwrap() == ReserveUsagePolicyRequestOutcome::Allowed { + allowed += 1; + } + } + assert_eq!(allowed, 5); + assert_eq!( + repository + .request_admissions + .read() + .expect("admission lock") + .len(), + 5 + ); + } +} + +#[cfg(test)] +mod usage_policy_cost_tests { + use super::*; + use crate::repository::settlement::UsagePolicyCostWindow; + + fn reserve( + request_id: &str, + admitted_at: u64, + reserved: u64, + limit: u64, + ) -> ReserveUsagePolicyCostInput { + ReserveUsagePolicyCostInput { + request_id: request_id.to_string(), + subject_id: "user-1".to_string(), + reservation_token: format!("token-{request_id}"), + admitted_at_unix_secs: admitted_at, + reserved_cost_units: reserved, + reservation_expires_at_unix_secs: admitted_at + 86_400, + retain_until_unix_secs: admitted_at + 32 * 86_400, + windows: vec![UsagePolicyCostWindow { + window_id: "rolling-5h".to_string(), + starts_at_unix_secs: admitted_at.saturating_sub(18_000), + ends_at_unix_secs: admitted_at + 1, + limit_cost_units: limit, + }], + } + } + + #[tokio::test] + async fn cost_reservations_are_atomic_idempotent_and_reconciled() { + let repository = InMemorySettlementRepository::default(); + assert_eq!( + repository + .reserve_usage_policy_cost(reserve("request-1", 20_000, 60, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: 60, + additional_reserved_cost_units: 60, + } + ); + assert_eq!( + repository + .reserve_usage_policy_cost(reserve("request-1", 20_000, 60, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: 60, + additional_reserved_cost_units: 0, + } + ); + assert!(matches!( + repository + .reserve_usage_policy_cost(reserve("request-2", 20_001, 50, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Rejected { + window_index: 0, + used_cost_units: 60, + .. + } + )); + + repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "request-1".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "token-request-1".to_string(), + actual_cost_units: 30, + terminal_state: UsagePolicyCostReservationState::Finalized, + finalized_at_unix_secs: 20_001, + }) + .await + .unwrap(); + assert!(matches!( + repository + .reserve_usage_policy_cost(reserve("request-2", 20_002, 50, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { .. } + )); + } + + #[tokio::test] + async fn same_request_id_can_have_independent_inbound_reservation_tokens() { + let repository = InMemorySettlementRepository::default(); + repository + .reserve_usage_policy_cost(reserve("shared-trace", 20_000, 10, 100)) + .await + .unwrap(); + + let mut second = reserve("shared-trace", 20_000, 10, 100); + second.reservation_token = "another-inbound-request".to_string(); + assert_eq!( + repository.reserve_usage_policy_cost(second).await.unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: 10, + additional_reserved_cost_units: 10, + } + ); + + let finalized = repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "shared-trace".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "another-inbound-request".to_string(), + actual_cost_units: 8, + terminal_state: UsagePolicyCostReservationState::Finalized, + finalized_at_unix_secs: 20_001, + }) + .await + .unwrap() + .expect("second reservation"); + assert_eq!(finalized.reservation_token, "another-inbound-request"); + assert_eq!(finalized.state, UsagePolicyCostReservationState::Finalized); + + // A forged token cannot finalize the first reservation, even though request_id matches. + assert_eq!( + repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "shared-trace".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "forged-token".to_string(), + actual_cost_units: 0, + terminal_state: UsagePolicyCostReservationState::Released, + finalized_at_unix_secs: 20_002, + }) + .await + .unwrap(), + None + ); + assert_eq!( + repository + .reserve_usage_policy_cost(reserve("shared-trace", 20_000, 10, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { + reserved_cost_units: 10, + additional_reserved_cost_units: 0, + } + ); + } + + #[tokio::test] + async fn expired_and_released_reservations_stop_consuming_capacity() { + let repository = InMemorySettlementRepository::default(); + let mut expired = reserve("expired", 10, 100, 100); + expired.reservation_expires_at_unix_secs = 11; + expired.retain_until_unix_secs = 100_000; + repository.reserve_usage_policy_cost(expired).await.unwrap(); + assert!(matches!( + repository + .reserve_usage_policy_cost(reserve("after-expiry", 12, 100, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { .. } + )); + + repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "after-expiry".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "token-after-expiry".to_string(), + actual_cost_units: 0, + terminal_state: UsagePolicyCostReservationState::Released, + finalized_at_unix_secs: 13, + }) + .await + .unwrap(); + assert!(matches!( + repository + .reserve_usage_policy_cost(reserve("after-release", 14, 100, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { .. } + )); + } + + #[tokio::test] + async fn retries_only_extend_expiry_and_retention_and_cleanup_is_bounded() { + let repository = InMemorySettlementRepository::default(); + let mut initial = reserve("retry", 20_000, 10, 100); + initial.reservation_expires_at_unix_secs = 30_000; + initial.retain_until_unix_secs = 50_000; + repository.reserve_usage_policy_cost(initial).await.unwrap(); + + let mut shorter = reserve("retry", 20_000, 10, 100); + shorter.reservation_expires_at_unix_secs = 25_000; + shorter.retain_until_unix_secs = 45_000; + repository.reserve_usage_policy_cost(shorter).await.unwrap(); + + let stored = repository + .cost_reservations + .read() + .expect("reservation lock") + .get("token-retry") + .cloned() + .expect("reservation"); + assert_eq!(stored.reservation_expires_at_unix_secs, 30_000); + assert_eq!(stored.retain_until_unix_secs, 50_000); + assert_eq!( + repository + .cleanup_usage_policy_cost_reservations(49_999, 10) + .await + .unwrap(), + 0 + ); + assert_eq!( + repository + .cleanup_usage_policy_cost_reservations(50_000, 1) + .await + .unwrap(), + 1 + ); + } + + #[tokio::test] + async fn reconciliation_ignores_subject_or_token_mismatches() { + let repository = InMemorySettlementRepository::default(); + repository + .reserve_usage_policy_cost(reserve("protected", 20_000, 10, 100)) + .await + .unwrap(); + + let mismatched = repository + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: "protected".to_string(), + subject_id: "user-1".to_string(), + reservation_token: "forged-token".to_string(), + actual_cost_units: 0, + terminal_state: UsagePolicyCostReservationState::Released, + finalized_at_unix_secs: 20_001, + }) + .await + .unwrap(); + assert_eq!(mismatched, None); + + assert!(matches!( + repository + .reserve_usage_policy_cost(reserve("after-mismatch", 20_001, 91, 100)) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Rejected { + window_index: 0, + used_cost_units: 10, + .. + } + )); + } +} + #[cfg(test)] mod tests { use super::InMemorySettlementRepository; use crate::repository::settlement::{SettlementWriteRepository, UsageSettlementInput}; use crate::repository::wallet::StoredWalletSnapshot; + use std::sync::Arc; fn sample_wallet() -> StoredWalletSnapshot { StoredWalletSnapshot::new( @@ -382,6 +1223,105 @@ mod tests { assert_eq!(settlement.provider_monthly_used_usd, Some(15.0)); } + #[tokio::test] + async fn pending_settlement_replay_is_idempotent_in_memory() { + let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]); + let input = UsageSettlementInput { + request_id: "req-pending-replay".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + api_key_is_standalone: false, + provider_id: Some("provider-1".to_string()), + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 3.0, + actual_total_cost_usd: 6.0, + finalized_at_unix_secs: Some(200), + }; + + let first = repository + .settle_usage(input.clone()) + .await + .unwrap() + .unwrap(); + let replay = repository.settle_usage(input).await.unwrap().unwrap(); + + assert_eq!(replay, first); + assert_eq!(replay.wallet_balance_after, Some(6.0)); + assert_eq!(replay.provider_monthly_used_usd, Some(6.0)); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn concurrent_pending_settlement_debits_once_in_memory() { + let repository = Arc::new(InMemorySettlementRepository::seed(vec![sample_wallet()])); + let barrier = Arc::new(tokio::sync::Barrier::new(9)); + let input = UsageSettlementInput { + request_id: "req-concurrent-pending".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + api_key_is_standalone: false, + provider_id: Some("provider-1".to_string()), + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 3.0, + actual_total_cost_usd: 6.0, + finalized_at_unix_secs: Some(200), + }; + let mut tasks = Vec::new(); + for _ in 0..8 { + let repository = Arc::clone(&repository); + let barrier = Arc::clone(&barrier); + let input = input.clone(); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + repository.settle_usage(input).await.unwrap().unwrap() + })); + } + barrier.wait().await; + + let mut observed = Vec::new(); + for task in tasks { + observed.push(task.await.unwrap()); + } + assert!(observed.windows(2).all(|pair| pair[0] == pair[1])); + assert_eq!(observed[0].wallet_balance_after, Some(6.0)); + assert_eq!(observed[0].provider_monthly_used_usd, Some(6.0)); + } + + #[tokio::test] + async fn corrupt_wallet_financial_values_fail_before_memory_settlement_mutation() { + for corrupt in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let mut wallet = sample_wallet(); + wallet.balance = corrupt; + let repository = InMemorySettlementRepository::seed(vec![wallet]); + let result = repository + .settle_usage(UsageSettlementInput { + request_id: "req-corrupt-wallet".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + api_key_is_standalone: false, + provider_id: Some("provider-1".to_string()), + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 1.0, + actual_total_cost_usd: 1.0, + finalized_at_unix_secs: Some(200), + }) + .await; + assert!(result.is_err()); + assert!(repository + .settlements + .read() + .expect("settlement snapshot lock") + .is_empty()); + assert!(repository + .provider_monthly_used + .read() + .expect("provider quota lock") + .is_empty()); + } + } + #[tokio::test] async fn returns_stored_snapshot_when_usage_is_already_finalized() { let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]); diff --git a/crates/aether-data/runtime/src/repository/system.rs b/crates/aether-data/runtime/src/repository/system.rs index 45292b22c..4d8290181 100644 --- a/crates/aether-data/runtime/src/repository/system.rs +++ b/crates/aether-data/runtime/src/repository/system.rs @@ -26,6 +26,7 @@ pub enum AdminSystemUsageAggregateImportMode { Skip, Overwrite, Error, + ValidateError, } #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] diff --git a/crates/aether-data/runtime/src/repository/usage/memory.rs b/crates/aether-data/runtime/src/repository/usage/memory.rs index 3134441a8..be1aa4fbd 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory.rs @@ -4,7 +4,8 @@ use std::sync::RwLock; use aether_ai_formats::UPSTREAM_IS_STREAM_KEY; use aether_data_contracts::repository::usage::{ - parse_usage_body_ref, usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary, + canonical_usage_body_ref_for, parse_usage_body_ref, sanitize_usage_request_metadata, + usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount, @@ -31,7 +32,8 @@ use serde_json::Value; use super::{ api_key_usage_contribution, provider_api_key_usage_contribution, - strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure, + sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence, + usage_can_recover_terminal_failure, usage_lifecycle_update_allowed, usage_request_metadata_client_family, ApiKeyUsageContribution, ApiKeyUsageDelta, ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, @@ -159,10 +161,6 @@ fn usage_status_is_finalized(status: &str) -> bool { matches!(status, "completed" | "failed" | "cancelled") } -fn usage_status_is_lifecycle(status: &str) -> bool { - matches!(status, "pending" | "streaming") -} - fn merge_usage_timing(existing: Option, incoming: Option) -> Option { match incoming { Some(0) | None => existing.or(incoming), @@ -1176,18 +1174,19 @@ impl UsageReadRepository for InMemoryUsageReadRepository { } async fn resolve_body_ref(&self, body_ref: &str) -> Result, DataLayerError> { + let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { + return Ok(None); + }; + let canonical_ref = usage_body_ref(&request_id, field); if let Some(value) = self .detached_bodies .read() .expect("usage repository lock") - .get(body_ref) + .get(&canonical_ref) .cloned() { return Ok(Some(value)); } - let Some((request_id, field)) = parse_usage_body_ref(body_ref) else { - return Ok(None); - }; let usage = self .by_request_id .read() @@ -2676,44 +2675,74 @@ fn usage_body_ref_from_metadata( .and_then(Value::as_object) .and_then(|object| object.get(field.as_ref_key())) .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .and_then(parse_usage_body_ref) - .filter(|(parsed_request_id, parsed_field)| { - parsed_request_id == request_id && *parsed_field == field - }) - .map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field)) + .and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field)) +} + +fn sanitize_memory_request_metadata(metadata: Option) -> Option { + sanitize_usage_request_metadata(metadata) } fn hydrate_legacy_body_refs(item: &mut StoredRequestUsageAudit) { - if item.request_body_ref.is_none() { - item.request_body_ref = usage_body_ref_from_metadata( - item.request_metadata.as_ref(), - &item.request_id, - UsageBodyField::RequestBody, - ); - } - if item.provider_request_body_ref.is_none() { - item.provider_request_body_ref = usage_body_ref_from_metadata( - item.request_metadata.as_ref(), - &item.request_id, - UsageBodyField::ProviderRequestBody, - ); - } - if item.response_body_ref.is_none() { - item.response_body_ref = usage_body_ref_from_metadata( - item.request_metadata.as_ref(), - &item.request_id, - UsageBodyField::ResponseBody, - ); - } - if item.client_response_body_ref.is_none() { - item.client_response_body_ref = usage_body_ref_from_metadata( - item.request_metadata.as_ref(), - &item.request_id, - UsageBodyField::ClientResponseBody, - ); - } + item.request_body_ref = item + .request_body_ref + .as_deref() + .and_then(|body_ref| { + canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::RequestBody) + }) + .or_else(|| { + usage_body_ref_from_metadata( + item.request_metadata.as_ref(), + &item.request_id, + UsageBodyField::RequestBody, + ) + }); + item.provider_request_body_ref = item + .provider_request_body_ref + .as_deref() + .and_then(|body_ref| { + canonical_usage_body_ref_for( + body_ref, + &item.request_id, + UsageBodyField::ProviderRequestBody, + ) + }) + .or_else(|| { + usage_body_ref_from_metadata( + item.request_metadata.as_ref(), + &item.request_id, + UsageBodyField::ProviderRequestBody, + ) + }); + item.response_body_ref = item + .response_body_ref + .as_deref() + .and_then(|body_ref| { + canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::ResponseBody) + }) + .or_else(|| { + usage_body_ref_from_metadata( + item.request_metadata.as_ref(), + &item.request_id, + UsageBodyField::ResponseBody, + ) + }); + item.client_response_body_ref = item + .client_response_body_ref + .as_deref() + .and_then(|body_ref| { + canonical_usage_body_ref_for( + body_ref, + &item.request_id, + UsageBodyField::ClientResponseBody, + ) + }) + .or_else(|| { + usage_body_ref_from_metadata( + item.request_metadata.as_ref(), + &item.request_id, + UsageBodyField::ClientResponseBody, + ) + }); } fn hydrate_client_family(item: &mut StoredRequestUsageAudit) { @@ -2723,34 +2752,6 @@ fn hydrate_client_family(item: &mut StoredRequestUsageAudit) { } } -fn persisted_usage_body_ref( - incoming_ref: Option<&str>, - incoming_body: Option<&Value>, - incoming_state: Option, - _metadata: Option<&Value>, - existing: Option<&StoredRequestUsageAudit>, - field: UsageBodyField, -) -> Option { - if incoming_state == Some(UsageBodyCaptureState::None) { - return None; - } - if incoming_body.is_some() { - return None; - } - incoming_ref - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - .or_else(|| { - existing.and_then(|existing| match field { - UsageBodyField::RequestBody => existing.request_body_ref.clone(), - UsageBodyField::ProviderRequestBody => existing.provider_request_body_ref.clone(), - UsageBodyField::ResponseBody => existing.response_body_ref.clone(), - UsageBodyField::ClientResponseBody => existing.client_response_body_ref.clone(), - }) - }) -} - fn request_body_capture_replaces_derived_facts( request_body: Option<&Value>, request_body_state: Option, @@ -2852,9 +2853,68 @@ impl UsageWriteRepository for InMemoryUsageReadRepository { usage: UpsertUsageRecord, ) -> Result { usage.validate()?; - let usage = strip_deprecated_usage_display_fields(usage); + let capture_usage = usage.clone(); + let usage = sanitize_usage_for_persistence(usage); let mut by_request_id = self.by_request_id.write().expect("usage repository lock"); let existing = by_request_id.get(&usage.request_id).cloned(); + if let Some(existing) = existing.as_ref() { + if !usage_lifecycle_update_allowed( + &existing.status, + &existing.billing_status, + existing.updated_at_unix_secs, + existing.finalized_at_unix_secs, + &usage.status, + &usage.billing_status, + usage.updated_at_unix_secs, + usage.finalized_at_unix_secs, + ) { + return Ok(existing.clone()); + } + let can_recover = usage_can_recover_terminal_failure( + existing.status.as_str(), + existing.billing_status.as_str(), + usage.status.as_str(), + usage.billing_status.as_str(), + ); + let completed_terminal_failure_recovery = existing.billing_status == "void" + && matches!(existing.status.as_str(), "failed" | "cancelled") + && usage.status == "completed"; + if completed_terminal_failure_recovery && !can_recover { + return Ok(existing.clone()); + } + } + let capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage); + if let Some(existing) = by_request_id.get_mut(&usage.request_id) { + existing.request_headers = None; + existing.request_body = None; + existing.request_body_ref = None; + existing.request_body_state = None; + existing.provider_request_headers = None; + existing.provider_request_body = None; + existing.provider_request_body_ref = None; + existing.provider_request_body_state = None; + existing.response_headers = None; + existing.response_body = None; + existing.response_body_ref = None; + existing.response_body_state = None; + existing.client_response_headers = None; + existing.client_response_body = None; + existing.client_response_body_ref = None; + existing.client_response_body_state = None; + existing.request_metadata = + sanitize_usage_request_metadata(existing.request_metadata.take()); + } + { + let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock"); + for field in [ + UsageBodyField::RequestBody, + UsageBodyField::ProviderRequestBody, + UsageBodyField::ResponseBody, + UsageBodyField::ClientResponseBody, + ] { + detached_bodies.remove(&usage_body_ref(&usage.request_id, field)); + } + } let created_at_unix_ms = by_request_id .get(&usage.request_id) @@ -2871,46 +2931,20 @@ impl UsageWriteRepository for InMemoryUsageReadRepository { ) }) .unwrap_or_default(); - if existing.as_ref().is_some_and(|existing| { - let can_recover = usage_can_recover_terminal_failure( - existing.status.as_str(), - existing.billing_status.as_str(), - usage.status.as_str(), - usage.billing_status.as_str(), - ); - let finalized_lifecycle_regression = usage_status_is_finalized(&existing.status) - && usage_status_is_lifecycle(&usage.status); - let completed_terminal_failure_recovery = existing.billing_status == "void" - && matches!(existing.status.as_str(), "failed" | "cancelled") - && usage.status == "completed"; - (finalized_lifecycle_regression || completed_terminal_failure_recovery) && !can_recover - }) { - return Ok(existing.expect("existing usage should be present").clone()); - } - if existing.as_ref().is_some_and(|existing| { - existing.billing_status == "pending" - && existing.status == "streaming" - && usage.status == "pending" - }) { - return Ok(existing.expect("existing usage should be present").clone()); - } - let replace_client_request_body_facts = request_body_capture_replaces_derived_facts( - usage.request_body.as_ref(), - usage.request_body_state, + capture_usage.request_body.as_ref(), + capture_usage.request_body_state, ); let replace_provider_request_body_facts = request_body_capture_replaces_derived_facts( - usage.provider_request_body.as_ref(), - usage.provider_request_body_state, + capture_usage.provider_request_body.as_ref(), + capture_usage.provider_request_body_state, ); - let clear_request_body = usage.request_body_state == Some(UsageBodyCaptureState::None); + let clear_request_body = + capture_usage.request_body_state == Some(UsageBodyCaptureState::None); let clear_provider_request_body = - usage.provider_request_body_state == Some(UsageBodyCaptureState::None); - let clear_response_body = usage.response_body_state == Some(UsageBodyCaptureState::None); - let clear_client_response_body = - usage.client_response_body_state == Some(UsageBodyCaptureState::None); + capture_usage.provider_request_body_state == Some(UsageBodyCaptureState::None); let replace_routing_snapshot = usage_status_is_finalized(&usage.status); - let mut incoming_request_metadata = usage.request_metadata.clone(); + let mut incoming_request_metadata = capture_usage.request_metadata.clone(); if incoming_request_metadata.is_some() && (clear_request_body || clear_provider_request_body) { @@ -2942,61 +2976,11 @@ impl UsageWriteRepository for InMemoryUsageReadRepository { .and_then(|existing| existing.request_metadata.clone()) } }); - let request_body_ref = persisted_usage_body_ref( - usage.request_body_ref.as_deref(), - usage.request_body.as_ref(), - usage.request_body_state, - request_metadata.as_ref(), - existing.as_ref(), - UsageBodyField::RequestBody, - ); - let provider_request_body_ref = persisted_usage_body_ref( - usage.provider_request_body_ref.as_deref(), - usage.provider_request_body.as_ref(), - usage.provider_request_body_state, - request_metadata.as_ref(), - existing.as_ref(), - UsageBodyField::ProviderRequestBody, - ); - let response_body_ref = persisted_usage_body_ref( - usage.response_body_ref.as_deref(), - usage.response_body.as_ref(), - usage.response_body_state, - request_metadata.as_ref(), - existing.as_ref(), - UsageBodyField::ResponseBody, - ); - let client_response_body_ref = persisted_usage_body_ref( - usage.client_response_body_ref.as_deref(), - usage.client_response_body.as_ref(), - usage.client_response_body_state, - request_metadata.as_ref(), - existing.as_ref(), - UsageBodyField::ClientResponseBody, - ); - if clear_request_body - || clear_provider_request_body - || clear_response_body - || clear_client_response_body - { - let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock"); - for (clear, field) in [ - (clear_request_body, UsageBodyField::RequestBody), - ( - clear_provider_request_body, - UsageBodyField::ProviderRequestBody, - ), - (clear_response_body, UsageBodyField::ResponseBody), - ( - clear_client_response_body, - UsageBodyField::ClientResponseBody, - ), - ] { - if clear { - detached_bodies.remove(&usage_body_ref(&usage.request_id, field)); - } - } - } + let request_metadata = sanitize_memory_request_metadata(request_metadata); + let request_body_ref = None; + let provider_request_body_ref = None; + let response_body_ref = None; + let client_response_body_ref = None; let stored = StoredRequestUsageAudit { id: existing .as_ref() @@ -3105,159 +3089,97 @@ impl UsageWriteRepository for InMemoryUsageReadRepository { ), status: usage.status, billing_status: usage.billing_status, - request_headers: usage.request_headers.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.request_headers.clone()) - }), - request_body: if clear_request_body { - None - } else { - usage.request_body.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.request_body.clone()) - }) - }, + request_headers: None, + request_body: None, request_body_ref, - request_body_state: usage.request_body_state.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.request_body_state) - }), - provider_request_headers: usage.provider_request_headers.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.provider_request_headers.clone()) - }), - provider_request_body: if clear_provider_request_body { - None - } else { - usage.provider_request_body.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.provider_request_body.clone()) - }) - }, + request_body_state: capture_usage.request_body_state, + provider_request_headers: None, + provider_request_body: None, provider_request_body_ref, - provider_request_body_state: usage.provider_request_body_state.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.provider_request_body_state) - }), - response_headers: usage.response_headers.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.response_headers.clone()) - }), - response_body: if clear_response_body { - None - } else { - usage.response_body.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.response_body.clone()) - }) - }, + provider_request_body_state: capture_usage.provider_request_body_state, + response_headers: None, + response_body: None, response_body_ref, - response_body_state: usage.response_body_state.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.response_body_state) - }), - client_response_headers: usage.client_response_headers.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.client_response_headers.clone()) - }), - client_response_body: if clear_client_response_body { - None - } else { - usage.client_response_body.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.client_response_body.clone()) - }) - }, + response_body_state: capture_usage.response_body_state, + client_response_headers: None, + client_response_body: None, client_response_body_ref, - client_response_body_state: usage.client_response_body_state.or_else(|| { - existing - .as_ref() - .and_then(|existing| existing.client_response_body_state) - }), + client_response_body_state: capture_usage.client_response_body_state, candidate_id: if replace_routing_snapshot { - usage.candidate_id + capture_usage.candidate_id } else { - usage.candidate_id.or_else(|| { + capture_usage.candidate_id.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_candidate_id().map(ToOwned::to_owned)) }) }, candidate_index: if replace_routing_snapshot { - usage.candidate_index + capture_usage.candidate_index } else { - usage.candidate_index.or_else(|| { + capture_usage.candidate_index.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_candidate_index()) }) }, key_name: if replace_routing_snapshot { - usage.key_name + capture_usage.key_name } else { - usage.key_name.or_else(|| { + capture_usage.key_name.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_key_name().map(ToOwned::to_owned)) }) }, planner_kind: if replace_routing_snapshot { - usage.planner_kind + capture_usage.planner_kind } else { - usage.planner_kind.or_else(|| { + capture_usage.planner_kind.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_planner_kind().map(ToOwned::to_owned)) }) }, route_family: if replace_routing_snapshot { - usage.route_family + capture_usage.route_family } else { - usage.route_family.or_else(|| { + capture_usage.route_family.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_route_family().map(ToOwned::to_owned)) }) }, route_kind: if replace_routing_snapshot { - usage.route_kind + capture_usage.route_kind } else { - usage.route_kind.or_else(|| { + capture_usage.route_kind.or_else(|| { existing .as_ref() .and_then(|existing| existing.routing_route_kind().map(ToOwned::to_owned)) }) }, execution_path: if replace_routing_snapshot { - usage.execution_path + capture_usage.execution_path } else { - usage.execution_path.or_else(|| { + capture_usage.execution_path.or_else(|| { existing.as_ref().and_then(|existing| { existing.routing_execution_path().map(ToOwned::to_owned) }) }) }, local_execution_runtime_miss_reason: if replace_routing_snapshot { - usage.local_execution_runtime_miss_reason + capture_usage.local_execution_runtime_miss_reason } else { - usage.local_execution_runtime_miss_reason.or_else(|| { - existing.as_ref().and_then(|existing| { - existing - .routing_local_execution_runtime_miss_reason() - .map(ToOwned::to_owned) + capture_usage + .local_execution_runtime_miss_reason + .or_else(|| { + existing.as_ref().and_then(|existing| { + existing + .routing_local_execution_runtime_miss_reason() + .map(ToOwned::to_owned) + }) }) - }) }, client_family: usage_request_metadata_client_family(request_metadata.as_ref()) .map(ToOwned::to_owned) 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 c4c561fcc..949340fe8 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs @@ -788,6 +788,64 @@ async fn upsert_allows_completed_recovery_after_void_failure() { assert_eq!(stored.total_tokens, 10); } +#[tokio::test] +async fn stale_terminal_event_cannot_replace_usage_routing_or_counter_contribution() { + let auth_api_keys = sample_auth_api_key_repository(&["api-key-1"]); + let repository = InMemoryUsageReadRepository::default() + .with_auth_api_key_repository(Arc::clone(&auth_api_keys)); + + let mut newer = sample_upsert_usage_record("req-stale-terminal"); + newer.api_key_id = Some("api-key-1".to_string()); + newer.status = "completed".to_string(); + newer.status_code = Some(200); + newer.total_tokens = Some(5); + newer.total_cost_usd = Some(0.5); + newer.candidate_id = Some("candidate-new".to_string()); + newer.route_kind = Some("route-new".to_string()); + newer.updated_at_unix_secs = 200; + newer.finalized_at_unix_secs = Some(200); + repository + .upsert(newer) + .await + .expect("newer terminal usage should upsert"); + + let mut stale = sample_upsert_usage_record("req-stale-terminal"); + stale.api_key_id = Some("api-key-1".to_string()); + stale.status = "failed".to_string(); + stale.billing_status = "void".to_string(); + stale.status_code = Some(503); + stale.total_tokens = Some(999); + stale.total_cost_usd = Some(99.0); + stale.candidate_id = Some("candidate-stale".to_string()); + stale.route_kind = Some("route-stale".to_string()); + stale.updated_at_unix_secs = 199; + stale.finalized_at_unix_secs = Some(199); + let stored = repository + .upsert(stale) + .await + .expect("stale terminal usage should be ignored"); + + assert_eq!(stored.status, "completed"); + assert_eq!(stored.billing_status, "pending"); + assert_eq!(stored.status_code, Some(200)); + assert_eq!(stored.total_tokens, 5); + assert_eq!(stored.total_cost_usd, 0.5); + assert_eq!(stored.routing_candidate_id(), Some("candidate-new")); + assert_eq!(stored.routing_route_kind(), Some("route-new")); + assert_eq!(stored.updated_at_unix_secs, 200); + + let key = auth_api_keys + .list_export_api_keys_by_ids(&["api-key-1".to_string()]) + .await + .expect("api key stats should load") + .into_iter() + .next() + .expect("api key should exist"); + assert_eq!(key.total_requests, 1); + assert_eq!(key.total_tokens, 5); + assert_eq!(key.total_cost_usd, 0.5); +} + #[tokio::test] async fn upsert_rejects_non_authoritative_void_failure_recovery() { let repository = InMemoryUsageReadRepository::default(); @@ -1189,6 +1247,23 @@ async fn detached_body_seed_moves_large_payloads_behind_usage_refs() { ); } +#[tokio::test] +async fn seed_discards_cross_request_and_cross_field_body_refs() { + let mut usage = sample_usage("req-ref-target", 100); + usage.request_body_ref = Some("usage://request/req-ref-owner/request_body".to_string()); + usage.response_body_ref = Some("usage://request/req-ref-target/request_body".to_string()); + + let repository = InMemoryUsageReadRepository::seed(vec![usage]); + let stored = repository + .find_by_request_id("req-ref-target") + .await + .expect("find should succeed") + .expect("usage should exist"); + + assert!(stored.request_body_ref.is_none()); + assert!(stored.response_body_ref.is_none()); +} + #[tokio::test] async fn upsert_writes_usage_record() { let repository = InMemoryUsageReadRepository::default(); @@ -1520,12 +1595,7 @@ async fn upsert_does_not_backfill_typed_body_refs_from_request_metadata() { .expect("upsert should succeed"); assert_eq!(stored.request_body_ref, None); - assert_eq!( - stored.request_metadata, - Some(json!({ - "request_body_ref": "usage://request/req-upsert-body-ref-metadata/request_body" - })) - ); + assert_eq!(stored.request_metadata, None); } #[tokio::test] diff --git a/crates/aether-data/runtime/src/repository/usage/mod.rs b/crates/aether-data/runtime/src/repository/usage/mod.rs index 8b630b92b..1a0a441bb 100644 --- a/crates/aether-data/runtime/src/repository/usage/mod.rs +++ b/crates/aether-data/runtime/src/repository/usage/mod.rs @@ -6,33 +6,35 @@ mod mysql; pub(crate) use aether_data_contracts::repository::usage::{ api_key_usage_contribution, incoming_usage_can_recover_terminal_failure, model_usage_contribution, provider_api_key_usage_contribution, provider_api_key_usage_is_error, - provider_api_key_usage_is_success, strip_deprecated_usage_display_fields, - usage_can_recover_terminal_failure, usage_request_metadata_client_family, ApiKeyLastUsedDelta, - ApiKeyUsageContribution, ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution, - ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, - ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta, - StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, - StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit, - StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, - StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow, - StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDailySummary, - StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount, - StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, - StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow, - StoredUsageProviderPerformance, StoredUsageProviderPerformanceProviderRow, - StoredUsageProviderPerformanceSummary, StoredUsageProviderPerformanceTimelineRow, - StoredUsageSettledCostSummary, StoredUsageTimeSeriesBucket, StoredUsageUserTotals, - UpsertUsageRecord, UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, - UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery, - UsageBreakdownGroupBy, UsageBreakdownSummaryQuery, UsageCacheAffinityHitSummaryQuery, - UsageCacheAffinityIntervalGroupBy, UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, - UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupWindow, - UsageCostSavingsSummaryQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot, - UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery, - UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery, - UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, - UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, - UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, + provider_api_key_usage_is_success, sanitize_usage_capture_controls_for_persistence, + sanitize_usage_for_persistence, strip_deprecated_usage_display_fields, + usage_can_recover_terminal_failure, usage_lifecycle_update_allowed, + usage_request_metadata_client_family, ApiKeyLastUsedDelta, ApiKeyUsageContribution, + ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution, ModelUsageDelta, + PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, + ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary, + StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow, + StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary, + StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, + StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, + StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow, + StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary, + StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary, + StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance, + StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary, + StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary, + StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UpsertUsageRecord, + UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery, + UsageAuditListQuery, UsageAuditSummaryQuery, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery, + UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy, + UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupPreviewCounts, + UsageCleanupSummary, UsageCleanupWindow, UsageCostSavingsSummaryQuery, + UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot, + UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery, + UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy, + UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery, + UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository, + UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery, UsageWriteRepository, }; #[cfg(feature = "postgres")] diff --git a/crates/aether-data/runtime/src/repository/users/memory.rs b/crates/aether-data/runtime/src/repository/users/memory.rs index 87929d40d..743a0d044 100644 --- a/crates/aether-data/runtime/src/repository/users/memory.rs +++ b/crates/aether-data/runtime/src/repository/users/memory.rs @@ -4,11 +4,14 @@ use std::sync::RwLock; use async_trait::async_trait; use super::{ - normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, + is_valid_bcrypt_hash, last_oauth_unbind_denial, normalize_user_group_name, + BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy, - UserExportSummary, UserReadRepository, + UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED, + LAST_ACTIVE_ADMIN_UPDATE_DENIED, }; use crate::DataLayerError; @@ -194,15 +197,6 @@ impl InMemoryUserReadRepository { } } -fn looks_like_bcrypt_hash(value: &str) -> bool { - let bytes = value.as_bytes(); - value.len() == 60 - && matches!(value.get(0..4), Some("$2a$") | Some("$2b$") | Some("$2y$")) - && bytes.get(4).is_some_and(u8::is_ascii_digit) - && bytes.get(5).is_some_and(u8::is_ascii_digit) - && bytes.get(6) == Some(&b'$') -} - fn normalize_optional_json_value(value: Option) -> Option { match value { Some(serde_json::Value::Null) | None => None, @@ -210,6 +204,16 @@ fn normalize_optional_json_value(value: Option) -> Option Vec { + values + .iter() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .collect::>() + .into_iter() + .collect() +} + fn find_memory_ldap_user_id( repository: &InMemoryUserReadRepository, ldap_dn: Option<&str>, @@ -685,6 +689,25 @@ impl UserReadRepository for InMemoryUserReadRepository { Ok(Some(group)) } + async fn restore_user_group_if_matches( + &self, + expected: &StoredUserGroup, + restored: &StoredUserGroup, + ) -> Result { + if self.read_only || expected.id != restored.id || expected.id.trim().is_empty() { + return Ok(false); + } + let mut groups = self.groups_by_id.write().expect("user repository lock"); + let Some(current) = groups.get(&expected.id) else { + return Ok(false); + }; + if current != expected { + return Ok(false); + } + groups.insert(restored.id.clone(), restored.clone()); + Ok(true) + } + async fn delete_user_group(&self, group_id: &str) -> Result { if self.read_only { return Ok(false); @@ -829,6 +852,43 @@ impl UserReadRepository for InMemoryUserReadRepository { .await } + async fn restore_user_groups_if_matches( + &self, + user_id: &str, + expected_group_ids: &[String], + restored_group_ids: &[String], + ) -> Result { + if self.read_only { + return Ok(false); + } + let expected = normalized_ids(expected_group_ids); + let restored = normalized_ids(restored_group_ids); + let groups = self.groups_by_id.read().expect("user repository lock"); + if restored + .iter() + .any(|group_id| !groups.contains_key(group_id)) + { + return Ok(false); + } + let mut members = self.group_members.write().expect("user repository lock"); + let mut current = members + .keys() + .filter(|(_, candidate_user_id)| candidate_user_id == user_id) + .map(|(group_id, _)| group_id.clone()) + .collect::>(); + current.sort(); + current.dedup(); + if current != expected { + return Ok(false); + } + members.retain(|(_, candidate_user_id), _| candidate_user_id != user_id); + let now = chrono::Utc::now(); + for group_id in restored { + members.insert((group_id, user_id.to_string()), now); + } + Ok(true) + } + async fn add_user_to_group( &self, group_id: &str, @@ -1007,6 +1067,46 @@ impl UserReadRepository for InMemoryUserReadRepository { .cloned()) } + async fn resolve_enabled_oauth_linked_user( + &self, + provider_type: &str, + provider_user_id: &str, + provider_username: Option<&str>, + provider_email: Option<&str>, + extra_data: Option, + verified_email: Option<&str>, + touched_at: chrono::DateTime, + provider_enabled_snapshot: bool, + ) -> Result { + if !provider_enabled_snapshot { + return Ok(ResolveOAuthLinkedUserOutcome::ProviderUnavailable); + } + let Some(mut user) = self + .find_oauth_linked_user(provider_type, provider_user_id) + .await? + else { + return Ok(ResolveOAuthLinkedUserOutcome::NotLinked); + }; + self.touch_oauth_link( + provider_type, + provider_user_id, + provider_username, + provider_email, + extra_data, + touched_at, + ) + .await?; + if let Some(verified_email) = verified_email { + if self + .upgrade_oauth_email_verification_if_matches(&user.id, verified_email, touched_at) + .await? + { + user.email_verified = true; + } + } + Ok(ResolveOAuthLinkedUserOutcome::Linked(user)) + } + async fn touch_oauth_link( &self, provider_type: &str, @@ -1047,6 +1147,7 @@ impl UserReadRepository for InMemoryUserReadRepository { async fn create_oauth_auth_user( &self, email: Option, + email_verified: bool, username: String, created_at: chrono::DateTime, ) -> Result, DataLayerError> { @@ -1057,7 +1158,7 @@ impl UserReadRepository for InMemoryUserReadRepository { let user = StoredUserAuthRecord::new( uuid::Uuid::new_v4().to_string(), email, - true, + email_verified, username, None, "user".to_string(), @@ -1120,7 +1221,57 @@ impl UserReadRepository for InMemoryUserReadRepository { .count() as u64) } - async fn upsert_user_oauth_link( + async fn has_oauth_links_for_provider( + &self, + provider_type: &str, + ) -> Result { + let provider_type = provider_type.trim(); + Ok(self + .oauth_links_by_id + .read() + .expect("user repository lock") + .values() + .any(|link| link.provider_type == provider_type)) + } + + async fn count_locked_users_if_oauth_provider_disabled( + &self, + provider_type: &str, + enabled_provider_types_snapshot: &[String], + ldap_exclusive: bool, + ) -> Result { + let provider_type = provider_type.trim(); + let enabled = enabled_provider_types_snapshot + .iter() + .map(|value| value.as_str()) + .collect::>(); + let links = self.oauth_links_by_id.read().expect("user repository lock"); + let users = self.auth_by_id.read().expect("user repository lock"); + Ok(users + .values() + .filter(|user| user.is_active && !user.is_deleted) + .filter(|user| { + links + .values() + .any(|link| link.user_id == user.id && link.provider_type == provider_type) + }) + .filter(|user| { + !links.values().any(|link| { + link.user_id == user.id + && link.provider_type != provider_type + && enabled.contains(link.provider_type.as_str()) + }) + }) + .filter(|user| { + user.auth_source.eq_ignore_ascii_case("oauth") + || (ldap_exclusive + && user.auth_source.eq_ignore_ascii_case("local") + && !user.role.eq_ignore_ascii_case("admin")) + }) + .count()) + } + + async fn bind_user_oauth_link_if_provider_enabled( &self, user_id: &str, provider_type: &str, @@ -1129,27 +1280,66 @@ impl UserReadRepository for InMemoryUserReadRepository { provider_email: Option<&str>, extra_data: Option, linked_at: chrono::DateTime, - ) -> Result<(), DataLayerError> { + provider_enabled_snapshot: bool, + session_expectation: Option<&BindUserOAuthLinkSessionExpectation>, + ) -> Result { if self.read_only { - return Ok(()); + return Ok(BindUserOAuthLinkOutcome::UserNotFound); } let provider_type = provider_type.trim().to_string(); let provider_user_id = provider_user_id.trim().to_string(); + if provider_type.is_empty() || provider_user_id.is_empty() { + return Err(DataLayerError::InvalidInput( + "OAuth provider type and subject must not be empty".to_string(), + )); + } + if !provider_enabled_snapshot { + return Ok(BindUserOAuthLinkOutcome::ProviderDisabled); + } + let users = self.auth_by_id.read().expect("user repository lock"); + let Some(user) = users.get(user_id) else { + return Ok(BindUserOAuthLinkOutcome::UserNotFound); + }; + let sessions = + session_expectation.map(|_| self.sessions_by_id.read().expect("user repository lock")); + if let Some(expectation) = session_expectation { + let checked_at = std::cmp::max(expectation.checked_at, chrono::Utc::now()); + let session_is_current = sessions + .as_ref() + .and_then(|sessions| sessions.get(&expectation.session_id)) + .is_some_and(|session| { + user.is_active + && !user.is_deleted + && user.security_version == expectation.security_version + && session.user_id == user_id + && session.client_device_id == expectation.client_device_id + && session.security_version == expectation.security_version + && !session.is_revoked() + && !session.is_expired(checked_at) + }); + if !session_is_current { + return Ok(BindUserOAuthLinkOutcome::SessionUnavailable); + } + } let mut links = self .oauth_links_by_id .write() .expect("user repository lock"); - if let Some(link) = links - .values_mut() - .find(|link| link.user_id == user_id && link.provider_type == provider_type) + if let Some(link) = links.values().find(|link| { + link.provider_type == provider_type && link.provider_user_id == provider_user_id + }) { + return Ok(if link.user_id == user_id { + BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser + } else { + BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }); + } + if links + .values() + .any(|link| link.user_id == user_id && link.provider_type == provider_type) { - link.provider_user_id = provider_user_id; - link.provider_username = provider_username.map(ToOwned::to_owned); - link.provider_email = provider_email.map(ToOwned::to_owned); - link.extra_data = extra_data; - link.last_login_at = Some(linked_at); - return Ok(()); + return Ok(BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider); } let link = StoredMemoryOAuthLink { id: uuid::Uuid::new_v4().to_string(), @@ -1163,26 +1353,82 @@ impl UserReadRepository for InMemoryUserReadRepository { last_login_at: Some(linked_at), }; links.insert(link.id.clone(), link); - Ok(()) + Ok(BindUserOAuthLinkOutcome::Bound) + } + + async fn upgrade_oauth_email_verification_if_matches( + &self, + user_id: &str, + verified_email: &str, + _verified_at: chrono::DateTime, + ) -> Result { + if self.read_only { + return Ok(false); + } + let verified_email = verified_email.trim(); + let mut users = self.auth_by_id.write().expect("user repository lock"); + let Some(user) = users.get_mut(user_id) else { + return Ok(false); + }; + if user.email_verified + || !user + .email + .as_deref() + .is_some_and(|email| email.trim().eq_ignore_ascii_case(verified_email)) + { + return Ok(false); + } + user.email_verified = true; + Ok(true) } async fn delete_user_oauth_link( &self, user_id: &str, provider_type: &str, - ) -> Result { + local_password_login_allowed: bool, + enabled_provider_types_snapshot: &[String], + ) -> Result { if self.read_only { - return Ok(false); + return Ok(DeleteUserOAuthLinkOutcome::NotFound); } let provider_type = provider_type.trim(); + let users = self.auth_by_id.read().expect("user repository lock"); + let Some(user) = users.get(user_id) else { + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + }; let mut links = self .oauth_links_by_id .write() .expect("user repository lock"); - let before = links.len(); + let target_exists = links + .values() + .any(|link| link.user_id == user_id && link.provider_type == provider_type); + if !target_exists { + return Ok(DeleteUserOAuthLinkOutcome::NotFound); + } + // The in-memory provider repository has a separate lock, so callers pass a + // point-in-time enabled-provider snapshot. SQL implementations instead read + // and lock provider rows in the same database transaction. + let has_remaining_enabled_oauth_link = links.values().any(|link| { + link.user_id == user_id + && link.provider_type != provider_type + && enabled_provider_types_snapshot + .iter() + .any(|enabled| enabled == &link.provider_type) + }); + if !has_remaining_enabled_oauth_link { + if let Some(outcome) = last_oauth_unbind_denial( + &user.auth_source, + user.password_hash.as_deref(), + local_password_login_allowed, + ) { + return Ok(outcome); + } + } links.retain(|_, link| !(link.user_id == user_id && link.provider_type == provider_type)); - Ok(links.len() != before) + Ok(DeleteUserOAuthLinkOutcome::Deleted) } async fn get_or_create_ldap_auth_user( @@ -1315,7 +1561,9 @@ impl UserReadRepository for InMemoryUserReadRepository { async fn update_local_auth_user_profile( &self, user_id: &str, + email_present: bool, email: Option, + email_verified: Option, username: Option, ) -> Result, DataLayerError> { if self.read_only { @@ -1329,8 +1577,11 @@ impl UserReadRepository for InMemoryUserReadRepository { let old_email = user.email.clone(); let old_username = user.username.clone(); - if let Some(email) = email { - user.email = Some(email); + if email_present { + user.email = email; + } + if let Some(email_verified) = email_verified { + user.email_verified = email_verified; } if let Some(username) = username { user.username = username; @@ -1365,6 +1616,201 @@ impl UserReadRepository for InMemoryUserReadRepository { Ok(Some(updated)) } + async fn restore_local_auth_user_state_if_matches( + &self, + expected_auth: &StoredUserAuthRecord, + restored_auth: &StoredUserAuthRecord, + expected_export: &StoredUserExportRow, + restored_export: &StoredUserExportRow, + expected_model_capability_settings: Option<&serde_json::Value>, + restored_model_capability_settings: Option, + expected_feature_settings: Option<&serde_json::Value>, + restored_feature_settings: Option, + ) -> Result { + if self.read_only + || expected_auth.id != restored_auth.id + || expected_export.id != expected_auth.id + || restored_export.id != restored_auth.id + { + return Ok(false); + } + + let mut users = self.auth_by_id.write().expect("user repository lock"); + let Some(current) = users.get(expected_auth.id.as_str()) else { + return Ok(false); + }; + if !current.matches_restore_state(expected_auth) { + return Ok(false); + } + let current_model = self + .model_settings_by_user_id + .read() + .expect("user repository lock") + .get(&expected_auth.id) + .cloned(); + let current_feature = self + .feature_settings_by_user_id + .read() + .expect("user repository lock") + .get(&expected_auth.id) + .cloned(); + if current_model.as_ref() != expected_model_capability_settings + || current_feature.as_ref() != expected_feature_settings + { + return Ok(false); + } + let current_export = self + .export_rows + .read() + .expect("user repository lock") + .iter() + .find(|row| row.id == expected_auth.id) + .cloned(); + if current_export.as_ref().is_some_and(|row| { + row.rate_limit != expected_export.rate_limit + || row.rate_limit_mode != expected_export.rate_limit_mode + }) { + return Ok(false); + } + + let removes_active_admin = current.role.eq_ignore_ascii_case("admin") + && current.is_active + && !current.is_deleted + && (!restored_auth.role.eq_ignore_ascii_case("admin") || !restored_auth.is_active); + if removes_active_admin + && users + .values() + .filter(|user| { + user.role.eq_ignore_ascii_case("admin") && user.is_active && !user.is_deleted + }) + .count() + <= 1 + { + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + + let old_email = current.email.clone(); + let old_username = current.username.clone(); + let security_state_changed = + current.role != restored_auth.role || current.is_active != restored_auth.is_active; + let user = users + .get_mut(expected_auth.id.as_str()) + .expect("user existence checked while holding write lock"); + user.email = restored_auth.email.clone(); + user.email_verified = restored_auth.email_verified; + user.username = restored_auth.username.clone(); + user.role = restored_auth.role.clone(); + user.allowed_providers = restored_auth.allowed_providers.clone(); + user.allowed_providers_mode = restored_auth.allowed_providers_mode.clone(); + user.allowed_api_formats = restored_auth.allowed_api_formats.clone(); + user.allowed_api_formats_mode = restored_auth.allowed_api_formats_mode.clone(); + user.allowed_models = restored_auth.allowed_models.clone(); + user.allowed_models_mode = restored_auth.allowed_models_mode.clone(); + user.is_active = restored_auth.is_active; + if security_state_changed { + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; + } + let updated = user.clone(); + drop(users); + + let mut identifiers = self + .auth_by_identifier + .write() + .expect("user repository lock"); + identifiers.remove(&old_username); + if let Some(old_email) = old_email { + identifiers.remove(&old_email); + } + identifiers.insert(updated.username.clone(), updated.id.clone()); + if let Some(email) = updated.email.as_ref() { + identifiers.insert(email.clone(), updated.id.clone()); + } + drop(identifiers); + if let Some(summary) = self + .by_id + .write() + .expect("user repository lock") + .get_mut(&updated.id) + { + summary.email = updated.email.clone(); + summary.username = updated.username.clone(); + summary.role = updated.role.clone(); + summary.is_active = updated.is_active; + } + let restored_model_capability_settings = + normalize_optional_json_value(restored_model_capability_settings); + let restored_feature_settings = normalize_optional_json_value(restored_feature_settings); + { + let mut settings = self + .model_settings_by_user_id + .write() + .expect("user repository lock"); + match restored_model_capability_settings.clone() { + Some(value) => { + settings.insert(updated.id.clone(), value); + } + None => { + settings.remove(&updated.id); + } + } + } + { + let mut settings = self + .feature_settings_by_user_id + .write() + .expect("user repository lock"); + match restored_feature_settings.clone() { + Some(value) => { + settings.insert(updated.id.clone(), value); + } + None => { + settings.remove(&updated.id); + } + } + } + if let Some(row) = self + .export_rows + .write() + .expect("user repository lock") + .iter_mut() + .find(|row| row.id == updated.id) + { + row.email = updated.email.clone(); + row.email_verified = updated.email_verified; + row.username = updated.username.clone(); + row.role = updated.role.clone(); + row.auth_source = updated.auth_source.clone(); + row.allowed_providers = updated.allowed_providers.clone(); + row.allowed_providers_mode = updated.allowed_providers_mode.clone(); + row.allowed_api_formats = updated.allowed_api_formats.clone(); + row.allowed_api_formats_mode = updated.allowed_api_formats_mode.clone(); + row.allowed_models = updated.allowed_models.clone(); + row.allowed_models_mode = updated.allowed_models_mode.clone(); + row.rate_limit = restored_export.rate_limit; + row.rate_limit_mode = restored_export.rate_limit_mode.clone(); + row.model_capability_settings = restored_model_capability_settings.clone(); + row.feature_settings = restored_feature_settings.clone(); + row.is_active = updated.is_active; + } + if security_state_changed { + let now = chrono::Utc::now(); + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + for session in sessions + .values_mut() + .filter(|session| session.user_id == updated.id && session.revoked_at.is_none()) + { + session.revoked_at = Some(now); + session.revoke_reason = Some("user_security_state_changed".to_string()); + session.updated_at = Some(now); + } + } + Ok(true) + } + async fn update_local_auth_user_password_hash( &self, user_id: &str, @@ -1380,9 +1826,108 @@ impl UserReadRepository for InMemoryUserReadRepository { return Ok(None); }; user.password_hash = Some(password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; Ok(Some(user.clone())) } + async fn restore_local_auth_user_password_hash_if_matches( + &self, + user_id: &str, + expected_password_hash: Option<&str>, + password_hash: Option, + _updated_at: chrono::DateTime, + ) -> Result { + if self.read_only { + return Ok(false); + } + + let mut auth_by_id = self.auth_by_id.write().expect("user repository lock"); + let Some(user) = auth_by_id.get_mut(user_id) else { + return Ok(false); + }; + if user.password_hash.as_deref() != expected_password_hash { + return Ok(false); + } + user.password_hash = password_hash; + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; + Ok(true) + } + + async fn reset_local_auth_user_password_and_revoke_sessions( + &self, + user_id: &str, + password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + if self.read_only { + return Ok(false); + } + let mut users = self.auth_by_id.write().expect("user repository lock"); + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + let Some(user) = users.get_mut(user_id).filter(|user| !user.is_deleted) else { + return Ok(false); + }; + user.password_hash = Some(password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && !session.is_revoked()) + { + session.revoked_at = Some(changed_at); + session.revoke_reason = Some("admin_password_reset".to_string()); + session.updated_at = Some(changed_at); + } + Ok(true) + } + + async fn change_local_auth_password_and_revoke_sessions( + &self, + user_id: &str, + current_session_id: &str, + expected_password_hash: Option<&str>, + next_password_hash: String, + changed_at: chrono::DateTime, + ) -> Result { + if self.read_only { + return Ok(false); + } + let mut users = self.auth_by_id.write().expect("user repository lock"); + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + let Some(user) = users.get_mut(user_id) else { + return Ok(false); + }; + if user.password_hash.as_deref() != expected_password_hash + || !user.is_active + || user.is_deleted + { + return Ok(false); + } + if !sessions.get(current_session_id).is_some_and(|session| { + session.user_id == user_id && !session.is_revoked() && !session.is_expired(changed_at) + }) { + return Ok(false); + } + user.password_hash = Some(next_password_hash); + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && !session.is_revoked()) + { + session.revoked_at = Some(changed_at); + session.revoke_reason = Some("password_changed".to_string()); + session.updated_at = Some(changed_at); + } + Ok(true) + } + async fn update_local_auth_user_admin_fields( &self, user_id: &str, @@ -1402,9 +1947,36 @@ impl UserReadRepository for InMemoryUserReadRepository { } let mut auth_by_id = self.auth_by_id.write().expect("user repository lock"); - let Some(user) = auth_by_id.get_mut(user_id) else { + let Some(current_user) = auth_by_id.get(user_id) else { return Ok(None); }; + let current_role = current_user.role.clone(); + let current_active = current_user.is_active; + let current_deleted = current_user.is_deleted; + let next_role = role.as_deref().unwrap_or(current_role.as_str()); + let next_active = is_active.unwrap_or(current_active); + if current_role.eq_ignore_ascii_case("admin") + && current_active + && !current_deleted + && (!next_role.eq_ignore_ascii_case("admin") || !next_active) + && auth_by_id + .values() + .filter(|user| { + user.role.eq_ignore_ascii_case("admin") && user.is_active && !user.is_deleted + }) + .count() + <= 1 + { + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_UPDATE_DENIED.to_string(), + )); + } + let security_state_changed = + !current_role.eq_ignore_ascii_case(next_role) || current_active != next_active; + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + let user = auth_by_id + .get_mut(user_id) + .expect("user existence checked while holding write lock"); if let Some(role) = role { user.role = role; } @@ -1447,7 +2019,22 @@ impl UserReadRepository for InMemoryUserReadRepository { if let Some(is_active) = is_active { user.is_active = is_active; } + if security_state_changed { + user.security_version = user.security_version.checked_add(1).ok_or_else(|| { + DataLayerError::UnexpectedValue("users.security_version overflow".to_string()) + })?; + let revoked_at = chrono::Utc::now(); + for session in sessions + .values_mut() + .filter(|session| session.user_id == user_id && session.revoked_at.is_none()) + { + session.revoked_at = Some(revoked_at); + session.revoke_reason = Some("user_security_state_changed".to_string()); + session.updated_at = Some(revoked_at); + } + } let updated = user.clone(); + drop(sessions); drop(auth_by_id); if let Some(summary) = self @@ -1718,11 +2305,28 @@ impl UserReadRepository for InMemoryUserReadRepository { return Ok(false); } - let removed = self - .auth_by_id - .write() - .expect("user repository lock") - .remove(user_id); + let mut auth_by_id = self.auth_by_id.write().expect("user repository lock"); + if auth_by_id.get(user_id).is_some_and(|user| { + user.role.eq_ignore_ascii_case("admin") && user.is_active && !user.is_deleted + }) && auth_by_id + .values() + .filter(|user| { + user.role.eq_ignore_ascii_case("admin") && user.is_active && !user.is_deleted + }) + .count() + <= 1 + { + return Err(DataLayerError::InvalidInput( + LAST_ACTIVE_ADMIN_DELETE_DENIED.to_string(), + )); + } + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + let removed = auth_by_id.remove(user_id); + if removed.is_some() { + sessions.retain(|_, session| session.user_id != user_id); + } + drop(sessions); + drop(auth_by_id); let Some(removed) = removed else { return Ok(false); }; @@ -1738,6 +2342,30 @@ impl UserReadRepository for InMemoryUserReadRepository { .write() .expect("user repository lock") .retain(|key, _| key.1 != user_id); + self.preferences_by_user_id + .write() + .expect("user repository lock") + .remove(user_id); + self.model_settings_by_user_id + .write() + .expect("user repository lock") + .remove(user_id); + self.feature_settings_by_user_id + .write() + .expect("user repository lock") + .remove(user_id); + self.ldap_dn_by_user_id + .write() + .expect("user repository lock") + .remove(user_id); + self.ldap_username_by_user_id + .write() + .expect("user repository lock") + .remove(user_id); + self.export_rows + .write() + .expect("user repository lock") + .retain(|row| row.id != user_id); let mut identifiers = self .auth_by_identifier @@ -1825,6 +2453,55 @@ impl UserReadRepository for InMemoryUserReadRepository { .or(session.updated_at) .or(session.last_seen_at) .unwrap_or_else(chrono::Utc::now); + let users = self.auth_by_id.read().expect("user repository lock"); + let Some(user) = users.get(&session.user_id) else { + return Ok(None); + }; + if !user.is_active || user.is_deleted || user.security_version != session.security_version { + return Ok(None); + } + let mut sessions = self.sessions_by_id.write().expect("user repository lock"); + for existing in sessions.values_mut() { + if existing.user_id == session.user_id + && existing.client_device_id == session.client_device_id + && existing.revoked_at.is_none() + && !existing.is_expired(now) + { + existing.revoked_at = Some(now); + existing.revoke_reason = Some("replaced_by_new_login".to_string()); + existing.updated_at = Some(now); + } + } + sessions.insert(session.id.clone(), session.clone()); + Ok(Some(session.clone())) + } + + async fn create_user_session_if_password_matches( + &self, + session: &StoredUserSessionRecord, + expected_password_hash: &str, + ) -> Result, DataLayerError> { + if self.read_only { + return Ok(None); + } + let mut users = self.auth_by_id.write().expect("user repository lock"); + let Some(user) = users.get_mut(&session.user_id) else { + return Ok(None); + }; + if user.password_hash.as_deref() != Some(expected_password_hash) + || !user.auth_source.eq_ignore_ascii_case("local") + || !user.is_active + || user.is_deleted + || user.security_version != session.security_version + { + return Ok(None); + } + let now = session + .created_at + .or(session.updated_at) + .or(session.last_seen_at) + .unwrap_or_else(chrono::Utc::now); + user.last_login_at = Some(now); let mut sessions = self.sessions_by_id.write().expect("user repository lock"); for existing in sessions.values_mut() { if existing.user_id == session.user_id @@ -1898,7 +2575,7 @@ impl UserReadRepository for InMemoryUserReadRepository { &self, user_id: &str, session_id: &str, - previous_refresh_token_hash: &str, + expected_refresh_token_hash: &str, next_refresh_token_hash: &str, rotated_at: chrono::DateTime, expires_at: chrono::DateTime, @@ -1910,13 +2587,15 @@ impl UserReadRepository for InMemoryUserReadRepository { } let mut sessions = self.sessions_by_id.write().expect("user repository lock"); - let Some(session) = sessions - .get_mut(session_id) - .filter(|s| s.user_id == user_id) - else { + let Some(session) = sessions.get_mut(session_id).filter(|session| { + session.user_id == user_id + && session.refresh_token_hash == expected_refresh_token_hash + && !session.is_revoked() + && !session.is_expired(rotated_at) + }) else { return Ok(false); }; - session.prev_refresh_token_hash = Some(previous_refresh_token_hash.to_string()); + session.prev_refresh_token_hash = Some(expected_refresh_token_hash.to_string()); session.refresh_token_hash = next_refresh_token_hash.to_string(); session.rotated_at = Some(rotated_at); session.expires_at = Some(expires_at); @@ -2010,7 +2689,7 @@ impl UserReadRepository for InMemoryUserReadRepository { && user .password_hash .as_deref() - .is_some_and(looks_like_bcrypt_hash) + .is_some_and(is_valid_bcrypt_hash) }) .count() as u64) } @@ -2018,9 +2697,27 @@ impl UserReadRepository for InMemoryUserReadRepository { #[cfg(test)] mod tests { + use std::sync::Arc; + use super::*; use crate::repository::users::{UserExportListQuery, UserReadRepository}; + fn user_group_record(name: &str, priority: i32) -> UpsertUserGroupRecord { + UpsertUserGroupRecord { + name: name.to_string(), + description: Some(format!("{name} description")), + priority, + allowed_providers: Some(vec!["provider-a".to_string()]), + allowed_providers_mode: "specific".to_string(), + allowed_api_formats: Some(vec!["chat".to_string()]), + allowed_api_formats_mode: "specific".to_string(), + allowed_models: Some(vec!["model-a".to_string()]), + allowed_models_mode: "specific".to_string(), + rate_limit: Some(10), + rate_limit_mode: "custom".to_string(), + } + } + #[tokio::test] async fn lists_seeded_users() { let user = StoredUserSummary::new( @@ -2041,6 +2738,73 @@ mod tests { assert_eq!(rows[0], user); } + #[tokio::test] + async fn restores_user_group_only_when_complete_snapshot_matches() { + let repository = InMemoryUserReadRepository::default(); + let before = repository + .create_user_group(user_group_record("cas-group", 1)) + .await + .expect("group creation should succeed") + .expect("group should be created"); + let after = repository + .update_user_group(&before.id, user_group_record("cas-group-imported", 2)) + .await + .expect("group update should succeed") + .expect("group should exist"); + + assert!(repository + .restore_user_group_if_matches(&after, &before) + .await + .expect("matching group restore should succeed")); + assert_eq!( + repository + .find_user_group_by_id(&before.id) + .await + .expect("group lookup should succeed"), + Some(before.clone()) + ); + + let current_after_second_update = repository + .update_user_group(&before.id, user_group_record("cas-group-concurrent", 3)) + .await + .expect("second group update should succeed") + .expect("group should exist"); + assert!(!repository + .restore_user_group_if_matches(&after, &before) + .await + .expect("stale group restore should return a conflict")); + assert_eq!( + repository + .find_user_group_by_id(&before.id) + .await + .expect("group lookup should succeed"), + Some(current_after_second_update) + ); + } + + #[tokio::test] + async fn user_group_restore_rejects_identity_mismatch_and_missing_rows() { + let repository = InMemoryUserReadRepository::default(); + let expected = repository + .create_user_group(user_group_record("cas-identity", 1)) + .await + .expect("group creation should succeed") + .expect("group should be created"); + let mut different_id = expected.clone(); + different_id.id = "different-group-id".to_string(); + assert!(!repository + .restore_user_group_if_matches(&expected, &different_id) + .await + .expect("identity mismatch should return false")); + + let mut missing = expected.clone(); + missing.id = "missing-group-id".to_string(); + assert!(!repository + .restore_user_group_if_matches(&missing, &missing) + .await + .expect("missing row should return false")); + } + #[tokio::test] async fn lists_seeded_non_admin_export_users() { let user = StoredUserExportRow::new( @@ -2205,14 +2969,25 @@ mod tests { let updated = repository .update_local_auth_user_profile( "user-1", + true, Some("alice2@example.com".to_string()), + Some(true), Some("alice2".to_string()), ) .await .expect("profile update should succeed") .expect("profile update should return user"); assert_eq!(updated.email.as_deref(), Some("alice2@example.com")); + assert!(updated.email_verified); assert_eq!(updated.username, "alice2"); + + let cleared = repository + .update_local_auth_user_profile("user-1", true, None, Some(false), None) + .await + .expect("nullable email update should succeed") + .expect("user should exist"); + assert!(cleared.email.is_none()); + assert!(!cleared.email_verified); assert!(repository .find_user_auth_by_identifier("alice@example.com") .await @@ -2309,7 +3084,7 @@ mod tests { None ); assert!(repository - .update_local_auth_user_profile("missing-user", None, None) + .update_local_auth_user_profile("missing-user", false, None, None, None) .await .expect("missing profile update should succeed") .is_none()); @@ -2328,6 +3103,288 @@ mod tests { .is_none()); } + #[tokio::test] + async fn restores_nullable_password_only_when_expected_hash_matches() { + let user = StoredUserAuthRecord::new( + "user-password-cas".to_string(), + Some("password-cas@example.com".to_string()), + true, + "password-cas".to_string(), + Some("imported-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + None, + None, + ) + .expect("auth user should build"); + let repository = InMemoryUserReadRepository::seed_auth_users([user]); + + assert!(repository + .restore_local_auth_user_password_hash_if_matches( + "user-password-cas", + Some("imported-hash"), + None, + chrono::Utc::now(), + ) + .await + .expect("nullable restore should succeed")); + assert!(repository + .find_user_auth_by_id("user-password-cas") + .await + .expect("user lookup should succeed") + .expect("user should exist") + .password_hash + .is_none()); + + assert!(!repository + .restore_local_auth_user_password_hash_if_matches( + "user-password-cas", + Some("stale-hash"), + Some("old-hash".to_string()), + chrono::Utc::now(), + ) + .await + .expect("conflicting restore should return false")); + assert!(repository + .find_user_auth_by_id("user-password-cas") + .await + .expect("user lookup should succeed") + .expect("user should exist") + .password_hash + .is_none()); + } + + #[tokio::test] + async fn hard_delete_preserves_last_admin_then_cleans_owned_memory_state() { + let now = chrono::Utc::now(); + let admin = StoredUserAuthRecord::new( + "admin-delete-target".to_string(), + Some("admin-delete@example.com".to_string()), + true, + "admin-delete".to_string(), + Some("password-hash".to_string()), + "admin".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("admin should build"); + let export_row = StoredUserExportRow::new( + admin.id.clone(), + admin.email.clone(), + admin.email_verified, + admin.username.clone(), + admin.password_hash.clone(), + admin.role.clone(), + admin.auth_source.clone(), + None, + None, + None, + None, + Some(serde_json::json!({"gpt-4.1": {"enabled": true}})), + true, + ) + .expect("export row should build") + .with_feature_settings(Some(serde_json::json!({"feature": true}))); + let session = StoredUserSessionRecord::new( + "admin-delete-session".to_string(), + admin.id.clone(), + "admin-delete-device".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("admin-delete-refresh"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("session should build"); + let preferences = StoredUserPreferenceRecord { + user_id: admin.id.clone(), + avatar_url: None, + bio: None, + default_provider_id: None, + default_provider_name: None, + theme: "system".to_string(), + language: "zh-CN".to_string(), + timezone: "Asia/Shanghai".to_string(), + email_notifications: true, + usage_alerts: true, + announcement_notifications: true, + }; + let repository = InMemoryUserReadRepository::seed_auth_users([admin.clone()]) + .with_export_users([export_row]) + .with_user_preferences([preferences]) + .with_user_sessions([session]); + repository + .oauth_links_by_id + .write() + .expect("user repository lock") + .insert( + "admin-delete-oauth".to_string(), + StoredMemoryOAuthLink { + id: "admin-delete-oauth".to_string(), + user_id: admin.id.clone(), + provider_type: "test".to_string(), + provider_user_id: "admin-delete-subject".to_string(), + provider_username: None, + provider_email: admin.email.clone(), + extra_data: None, + linked_at: now, + last_login_at: None, + }, + ); + repository + .group_members + .write() + .expect("user repository lock") + .insert(("group-1".to_string(), admin.id.clone()), now); + repository + .model_settings_by_user_id + .write() + .expect("user repository lock") + .insert(admin.id.clone(), serde_json::json!({"model": true})); + repository + .feature_settings_by_user_id + .write() + .expect("user repository lock") + .insert(admin.id.clone(), serde_json::json!({"feature": true})); + repository + .ldap_dn_by_user_id + .write() + .expect("user repository lock") + .insert(admin.id.clone(), "uid=admin-delete,dc=example".to_string()); + repository + .ldap_username_by_user_id + .write() + .expect("user repository lock") + .insert(admin.id.clone(), "admin-delete-ldap".to_string()); + + let error = repository + .delete_local_auth_user(&admin.id) + .await + .expect_err("last active admin delete must be rejected"); + assert!(crate::repository::users::is_last_active_admin_delete_denied(&error)); + assert!(repository + .auth_by_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(repository + .sessions_by_id + .read() + .expect("user repository lock") + .values() + .any(|session| session.user_id == admin.id)); + assert!(repository + .oauth_links_by_id + .read() + .expect("user repository lock") + .values() + .any(|link| link.user_id == admin.id)); + + repository + .create_local_auth_user_with_settings( + Some("admin-keeper@example.com".to_string()), + true, + "admin-keeper".to_string(), + "password-hash".to_string(), + "admin".to_string(), + None, + None, + None, + None, + ) + .await + .expect("second admin creation should succeed") + .expect("second admin should exist"); + assert!(repository + .delete_local_auth_user(&admin.id) + .await + .expect("delete with another active admin should succeed")); + + assert!(!repository + .by_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .auth_by_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .auth_by_identifier + .read() + .expect("user repository lock") + .values() + .any(|user_id| user_id == &admin.id)); + assert!(!repository + .sessions_by_id + .read() + .expect("user repository lock") + .values() + .any(|session| session.user_id == admin.id)); + assert!(!repository + .oauth_links_by_id + .read() + .expect("user repository lock") + .values() + .any(|link| link.user_id == admin.id)); + assert!(!repository + .group_members + .read() + .expect("user repository lock") + .keys() + .any(|(_, user_id)| user_id == &admin.id)); + assert!(!repository + .preferences_by_user_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .model_settings_by_user_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .feature_settings_by_user_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .ldap_dn_by_user_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .ldap_username_by_user_id + .read() + .expect("user repository lock") + .contains_key(&admin.id)); + assert!(!repository + .export_rows + .read() + .expect("user repository lock") + .iter() + .any(|row| row.id == admin.id)); + } + #[tokio::test] async fn creates_local_auth_users_in_memory() { let repository = InMemoryUserReadRepository::default(); @@ -2421,6 +3478,7 @@ mod tests { let user = repository .create_oauth_auth_user( Some("OAuth@Example.com".to_string()), + false, "oauth_user".to_string(), now, ) @@ -2428,6 +3486,7 @@ mod tests { .expect("oauth user should create") .expect("oauth user should exist"); assert_eq!(user.auth_source, "oauth"); + assert!(!user.email_verified); assert_eq!( repository .find_active_user_auth_by_email_ci("oauth@example.com") @@ -2438,7 +3497,7 @@ mod tests { ); repository - .upsert_user_oauth_link( + .bind_user_oauth_link( &user.id, "linuxdo", "subject-1", @@ -2479,15 +3538,430 @@ mod tests { ) .await .expect("link should touch")); - assert!(repository - .delete_user_oauth_link(&user.id, "linuxdo") + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &["linuxdo".to_string()],) + .await + .expect("last link deletion should resolve"), + DeleteUserOAuthLinkOutcome::LastOAuthBinding + ); + repository + .bind_user_oauth_link( + &user.id, + "github", + "subject-2", + Some("alice"), + Some("alice@example.com"), + None, + now, + ) .await - .expect("link should delete")); + .expect("second link should upsert"); + assert_eq!( + repository + .delete_user_oauth_link( + &user.id, + "linuxdo", + false, + &["linuxdo".to_string(), "github".to_string()], + ) + .await + .expect("link should delete"), + DeleteUserOAuthLinkOutcome::Deleted + ); + } + + #[tokio::test] + async fn oauth_bind_rejects_unavailable_session_without_creating_link_in_memory() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "oauth-bind-user".to_string(), + Some("oauth-bind@example.com".to_string()), + true, + "oauth-bind-user".to_string(), + None, + "user".to_string(), + "oauth".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("oauth bind user should build") + .with_security_version(7) + .expect("user security version should be valid"); + let valid_session = StoredUserSessionRecord::new( + "oauth-bind-session".to_string(), + user.id.clone(), + "oauth-bind-device".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("oauth-bind-refresh"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("oauth bind session should build") + .with_security_version(7) + .expect("session security version should be valid"); + + let mut revoked_session = valid_session.clone(); + revoked_session.revoked_at = Some(now); + let mut expired_session = valid_session.clone(); + expired_session.expires_at = Some(now - chrono::Duration::seconds(1)); + let cases = [ + ( + "revoked", + revoked_session, + BindUserOAuthLinkSessionExpectation::new( + "oauth-bind-session", + "oauth-bind-device", + 7, + now, + ) + .expect("revoked expectation should build"), + ), + ( + "expired", + expired_session, + BindUserOAuthLinkSessionExpectation::new( + "oauth-bind-session", + "oauth-bind-device", + 7, + now, + ) + .expect("expired expectation should build"), + ), + ( + "device-mismatch", + valid_session.clone(), + BindUserOAuthLinkSessionExpectation::new( + "oauth-bind-session", + "other-device", + 7, + now, + ) + .expect("device expectation should build"), + ), + ( + "security-version-mismatch", + valid_session, + BindUserOAuthLinkSessionExpectation::new( + "oauth-bind-session", + "oauth-bind-device", + 6, + now, + ) + .expect("security expectation should build"), + ), + ]; + + for (case, session, expectation) in cases { + let repository = InMemoryUserReadRepository::seed_auth_users([user.clone()]) + .with_user_sessions([session]); + let subject = format!("subject-{case}"); + assert_eq!( + repository + .bind_user_oauth_link_if_provider_enabled( + &user.id, + "linuxdo", + &subject, + None, + None, + None, + now, + true, + Some(&expectation), + ) + .await + .expect("session-bound OAuth bind should resolve"), + BindUserOAuthLinkOutcome::SessionUnavailable, + "case {case} should reject", + ); + assert_eq!( + repository + .count_user_oauth_links(&user.id) + .await + .expect("OAuth links should count"), + 0, + "case {case} must not create a link", + ); + } + } + + #[tokio::test] + async fn concurrent_oauth_binds_preserve_single_identity_owner_in_memory() { + let repository = Arc::new(InMemoryUserReadRepository::default()); + let now = chrono::Utc::now(); + let first_user = repository + .create_oauth_auth_user(None, false, "bind-first".to_string(), now) + .await + .expect("first user should create") + .expect("first user should exist"); + let second_user = repository + .create_oauth_auth_user(None, false, "bind-second".to_string(), now) + .await + .expect("second user should create") + .expect("second user should exist"); + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + + let first_repository = Arc::clone(&repository); + let first_barrier = Arc::clone(&barrier); + let first_id = first_user.id.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository + .bind_user_oauth_link( + &first_id, + "linuxdo", + "shared-subject", + None, + None, + None, + now, + ) + .await + .expect("first bind should resolve") + }); + let second_repository = Arc::clone(&repository); + let second_barrier = Arc::clone(&barrier); + let second_id = second_user.id.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .bind_user_oauth_link( + &second_id, + "linuxdo", + "shared-subject", + None, + None, + None, + now, + ) + .await + .expect("second bind should resolve") + }); + barrier.wait().await; + let outcomes = [ + first.await.expect("first bind task should join"), + second.await.expect("second bind task should join"), + ]; + + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == BindUserOAuthLinkOutcome::Bound) + .count(), + 1 + ); + assert_eq!( + outcomes + .iter() + .filter(|outcome| { + **outcome == BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser + }) + .count(), + 1 + ); + assert!(repository + .find_oauth_link_owner("linuxdo", "shared-subject") + .await + .expect("identity owner should load") + .is_some()); + } + + #[tokio::test] + async fn concurrent_oauth_unbinds_preserve_one_login_method_in_memory() { + let now = chrono::Utc::now(); + let repository = Arc::new(InMemoryUserReadRepository::default()); + let user = repository + .create_oauth_auth_user( + Some("concurrent-oauth@example.com".to_string()), + true, + "concurrent-oauth".to_string(), + now, + ) + .await + .expect("oauth user should create") + .expect("oauth user should exist"); + for (provider_type, subject) in [("linuxdo", "subject-1"), ("github", "subject-2")] { + repository + .bind_user_oauth_link(&user.id, provider_type, subject, None, None, None, now) + .await + .expect("oauth link should upsert"); + } + + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let first_repository = Arc::clone(&repository); + let first_barrier = Arc::clone(&barrier); + let first_user_id = user.id.clone(); + let first = tokio::spawn(async move { + first_barrier.wait().await; + first_repository + .delete_user_oauth_link( + &first_user_id, + "linuxdo", + false, + &["linuxdo".to_string(), "github".to_string()], + ) + .await + .expect("first unlink should resolve") + }); + let second_repository = Arc::clone(&repository); + let second_barrier = Arc::clone(&barrier); + let second_user_id = user.id.clone(); + let second = tokio::spawn(async move { + second_barrier.wait().await; + second_repository + .delete_user_oauth_link( + &second_user_id, + "github", + false, + &["linuxdo".to_string(), "github".to_string()], + ) + .await + .expect("second unlink should resolve") + }); + barrier.wait().await; + let outcomes = [ + first.await.expect("first unlink task should join"), + second.await.expect("second unlink task should join"), + ]; + + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == DeleteUserOAuthLinkOutcome::Deleted) + .count(), + 1 + ); + assert_eq!( + outcomes + .iter() + .filter(|outcome| **outcome == DeleteUserOAuthLinkOutcome::LastOAuthBinding) + .count(), + 1 + ); + assert_eq!( + repository + .count_user_oauth_links(&user.id) + .await + .expect("remaining links should count"), + 1 + ); + } + + #[tokio::test] + async fn oauth_unbind_in_memory_only_counts_enabled_provider_links() { + let repository = InMemoryUserReadRepository::default(); + let now = chrono::Utc::now(); + let user = repository + .create_oauth_auth_user( + Some("enabled-link@example.com".to_string()), + true, + "enabled-link".to_string(), + now, + ) + .await + .expect("oauth user should create") + .expect("oauth user should exist"); + for (provider_type, subject) in [("linuxdo", "subject-1"), ("github", "subject-2")] { + repository + .bind_user_oauth_link(&user.id, provider_type, subject, None, None, None, now) + .await + .expect("oauth link should upsert"); + } + let enabled_provider_types_snapshot = ["linuxdo".to_string()]; + + assert_eq!( + repository + .delete_user_oauth_link( + &user.id, + "linuxdo", + false, + &enabled_provider_types_snapshot, + ) + .await + .expect("enabled link deletion should resolve"), + DeleteUserOAuthLinkOutcome::LastOAuthBinding + ); + assert_eq!( + repository + .delete_user_oauth_link( + &user.id, + "github", + false, + &enabled_provider_types_snapshot, + ) + .await + .expect("disabled link deletion should resolve"), + DeleteUserOAuthLinkOutcome::Deleted + ); + assert!(repository + .has_user_oauth_provider_link(&user.id, "linuxdo") + .await + .expect("enabled provider link lookup should work")); + } + + #[tokio::test] + async fn oauth_unbind_in_memory_respects_ldap_exclusive_local_login_policy() { + let valid_hash = "$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string(); + let user = StoredUserAuthRecord::new( + "ldap-exclusive-local".to_string(), + Some("ldap-exclusive-local@example.com".to_string()), + true, + "ldap-exclusive-local".to_string(), + Some(valid_hash), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + None, + None, + ) + .expect("local user should build"); + let repository = InMemoryUserReadRepository::seed_auth_users([user.clone()]); + repository + .bind_user_oauth_link( + &user.id, + "linuxdo", + "subject-1", + None, + None, + None, + chrono::Utc::now(), + ) + .await + .expect("oauth link should upsert"); + + assert_eq!( + repository + .delete_user_oauth_link(&user.id, "linuxdo", false, &["linuxdo".to_string()],) + .await + .expect("unlink should resolve"), + DeleteUserOAuthLinkOutcome::LastLoginMethod + ); + assert!(repository + .has_user_oauth_provider_link(&user.id, "linuxdo") + .await + .expect("oauth link lookup should work")); } #[tokio::test] async fn counts_active_admin_auth_users() { - let valid_hash = format!("$2b$12${}", "a".repeat(53)); + let valid_hash = "$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string(); let admin = StoredUserAuthRecord::new( "admin-1".to_string(), Some("admin@example.com".to_string()), @@ -2581,6 +4055,23 @@ mod tests { #[tokio::test] async fn manages_user_sessions_in_memory() { let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "user-1".to_string(), + Some("session-user@example.com".to_string()), + true, + "session-user".to_string(), + None, + "user".to_string(), + "oauth".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("session user should build"); let session = StoredUserSessionRecord::new( "session-1".to_string(), "user-1".to_string(), @@ -2599,7 +4090,7 @@ mod tests { Some(now), ) .expect("session should build"); - let repository = InMemoryUserReadRepository::default(); + let repository = InMemoryUserReadRepository::seed_auth_users([user]); assert_eq!( repository @@ -2639,6 +4130,19 @@ mod tests { ) .await .expect("session should rotate")); + assert!(!repository + .rotate_user_session_refresh_token( + "user-1", + "session-1", + &StoredUserSessionRecord::hash_refresh_token("refresh-1"), + &StoredUserSessionRecord::hash_refresh_token("refresh-race-loser"), + now + chrono::Duration::minutes(2), + now + chrono::Duration::hours(2), + None, + None, + ) + .await + .expect("stale session rotation should be rejected")); let rotated = repository .find_user_session("user-1", "session-1") .await @@ -2659,6 +4163,379 @@ mod tests { .is_empty()); } + #[tokio::test] + async fn security_state_change_revokes_sessions_without_reactivation() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "security-state-user".to_string(), + Some("security-state@example.com".to_string()), + true, + "security-state-user".to_string(), + None, + "user".to_string(), + "oauth".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("security-state user should build"); + let session = StoredUserSessionRecord::new( + "security-state-session".to_string(), + user.id.clone(), + "security-state-device".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("security-state-refresh"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("security-state session should build"); + let repository = InMemoryUserReadRepository::seed_auth_users([user]) + .with_user_sessions([session.clone()]); + + repository + .update_local_auth_user_admin_fields( + "security-state-user", + None, + false, + None, + false, + None, + false, + None, + false, + None, + Some(false), + ) + .await + .expect("disable should succeed") + .expect("user should exist"); + let revoked = repository + .find_user_session("security-state-user", "security-state-session") + .await + .expect("revoked session should load") + .expect("revoked session should remain stored"); + assert!(revoked.is_revoked()); + assert_eq!( + revoked.revoke_reason.as_deref(), + Some("user_security_state_changed") + ); + + repository + .update_local_auth_user_admin_fields( + "security-state-user", + None, + false, + None, + false, + None, + false, + None, + false, + None, + Some(true), + ) + .await + .expect("reactivation should succeed") + .expect("user should exist"); + assert!(repository + .list_user_sessions("security-state-user") + .await + .expect("sessions should list") + .is_empty()); + assert!(repository + .create_user_session(&session) + .await + .expect("stale login should resolve") + .is_none()); + let current_version = repository + .find_user_auth_by_id("security-state-user") + .await + .expect("user lookup should resolve") + .expect("user should exist") + .security_version; + let fresh_session = session + .with_security_version(current_version) + .expect("security version should be valid"); + assert!(repository + .create_user_session(&fresh_session) + .await + .expect("fresh login should resolve") + .is_some()); + + repository + .update_local_auth_user_admin_fields( + "security-state-user", + Some("audit_admin".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + None, + ) + .await + .expect("role update should succeed") + .expect("user should exist"); + assert!(repository + .list_user_sessions("security-state-user") + .await + .expect("sessions should list") + .is_empty()); + } + + #[tokio::test] + async fn unchanged_security_state_preserves_sessions() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "unchanged-security-user".to_string(), + None, + false, + "unchanged-security-user".to_string(), + None, + "user".to_string(), + "oauth".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("unchanged-security user should build"); + let session = StoredUserSessionRecord::new( + "unchanged-security-session".to_string(), + user.id.clone(), + "unchanged-security-device".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("unchanged-security-refresh"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("unchanged-security session should build"); + let repository = + InMemoryUserReadRepository::seed_auth_users([user]).with_user_sessions([session]); + + repository + .update_local_auth_user_admin_fields( + "unchanged-security-user", + Some("USER".to_string()), + false, + None, + false, + None, + false, + None, + false, + None, + Some(true), + ) + .await + .expect("idempotent security update should succeed") + .expect("user should exist"); + assert_eq!( + repository + .list_user_sessions("unchanged-security-user") + .await + .expect("sessions should list") + .len(), + 1 + ); + } + + #[tokio::test] + async fn password_change_revokes_sessions_and_fences_stale_login() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "user-password-fence".to_string(), + Some("fence@example.com".to_string()), + true, + "fence-user".to_string(), + Some("old-password-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let current = StoredUserSessionRecord::new( + "session-current".to_string(), + user.id.clone(), + "device-current".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("refresh-current"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("current session should build"); + let stale_login = StoredUserSessionRecord::new( + "session-stale-login".to_string(), + user.id.clone(), + "device-stale-login".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("refresh-stale"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("stale login session should build"); + let repository = + InMemoryUserReadRepository::seed_auth_users([user]).with_user_sessions([current]); + + assert!(repository + .change_local_auth_password_and_revoke_sessions( + "user-password-fence", + "session-current", + Some("old-password-hash"), + "new-password-hash".to_string(), + now, + ) + .await + .expect("password change should succeed")); + assert!(repository + .list_user_sessions("user-password-fence") + .await + .expect("sessions should list") + .is_empty()); + assert!(repository + .create_user_session_if_password_matches(&stale_login, "old-password-hash") + .await + .expect("stale login should resolve") + .is_none()); + } + + #[tokio::test] + async fn admin_password_reset_revokes_sessions_and_fences_stale_login() { + let now = chrono::Utc::now(); + let user = StoredUserAuthRecord::new( + "user-admin-reset-fence".to_string(), + Some("admin-reset@example.com".to_string()), + true, + "admin-reset-user".to_string(), + Some("old-password-hash".to_string()), + "user".to_string(), + "local".to_string(), + None, + None, + None, + true, + false, + Some(now), + None, + ) + .expect("user should build"); + let active = StoredUserSessionRecord::new( + "session-before-admin-reset".to_string(), + user.id.clone(), + "device-before-admin-reset".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("refresh-before-admin-reset"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("active session should build"); + let stale_login = StoredUserSessionRecord::new( + "session-stale-admin-reset-login".to_string(), + user.id.clone(), + "device-stale-admin-reset-login".to_string(), + None, + StoredUserSessionRecord::hash_refresh_token("refresh-stale-admin-reset"), + None, + None, + Some(now), + Some(now + chrono::Duration::hours(1)), + None, + None, + None, + None, + Some(now), + Some(now), + ) + .expect("stale login session should build"); + let repository = + InMemoryUserReadRepository::seed_auth_users([user]).with_user_sessions([active]); + + assert!(repository + .reset_local_auth_user_password_and_revoke_sessions( + "user-admin-reset-fence", + "new-password-hash".to_string(), + now, + ) + .await + .expect("admin password reset should succeed")); + assert!(repository + .list_user_sessions("user-admin-reset-fence") + .await + .expect("sessions should list") + .is_empty()); + assert!(repository + .create_user_session_if_password_matches(&stale_login, "old-password-hash") + .await + .expect("stale login should resolve") + .is_none()); + assert_eq!( + repository + .find_user_auth_by_id("user-admin-reset-fence") + .await + .expect("user lookup should succeed") + .expect("user should exist") + .password_hash + .as_deref(), + Some("new-password-hash") + ); + } + #[tokio::test] async fn paginates_export_users_in_memory() { let repository = InMemoryUserReadRepository::seed_export_users(vec![ diff --git a/crates/aether-data/runtime/src/repository/users/mod.rs b/crates/aether-data/runtime/src/repository/users/mod.rs index 55e78e576..11dcb77b3 100644 --- a/crates/aether-data/runtime/src/repository/users/mod.rs +++ b/crates/aether-data/runtime/src/repository/users/mod.rs @@ -1,11 +1,15 @@ mod memory; pub use aether_data_contracts::repository::users::{ - normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, + is_last_active_admin_delete_denied, is_last_active_admin_update_denied, is_valid_bcrypt_hash, + last_oauth_unbind_denial, normalize_user_group_name, BindUserOAuthLinkOutcome, + BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome, + LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy, - UserExportSortOrder, UserExportSummary, UserReadRepository, + UserExportSortOrder, UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED, + LAST_ACTIVE_ADMIN_UPDATE_DENIED, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::MysqlUserReadRepository; diff --git a/crates/aether-data/runtime/src/repository/video_tasks/memory.rs b/crates/aether-data/runtime/src/repository/video_tasks/memory.rs index 302dc677a..3823186e2 100644 --- a/crates/aether-data/runtime/src/repository/video_tasks/memory.rs +++ b/crates/aether-data/runtime/src/repository/video_tasks/memory.rs @@ -14,6 +14,7 @@ use crate::DataLayerError; struct MemoryVideoTaskIndex { by_id: BTreeMap, short_to_id: BTreeMap, + request_to_id: BTreeMap, user_external_to_id: BTreeMap<(String, String), String>, } @@ -28,6 +29,7 @@ impl InMemoryVideoTaskRepository { if let Some(short_id) = previous.short_id { index.short_to_id.remove(&short_id); } + index.request_to_id.remove(&previous.request_id); if let (Some(user_id), Some(external_task_id)) = (previous.user_id, previous.external_task_id) { @@ -40,6 +42,9 @@ impl InMemoryVideoTaskRepository { if let Some(short_id) = &task.short_id { index.short_to_id.insert(short_id.clone(), task.id.clone()); } + index + .request_to_id + .insert(task.request_id.clone(), task.id.clone()); if let (Some(user_id), Some(external_task_id)) = (&task.user_id, &task.external_task_id) { index .user_external_to_id @@ -49,6 +54,35 @@ impl InMemoryVideoTaskRepository { task } + fn ensure_unique_keys_available( + index: &MemoryVideoTaskIndex, + task: &UpsertVideoTask, + ) -> Result<(), DataLayerError> { + if let Some(short_id) = task.short_id.as_deref() { + if index + .short_to_id + .get(short_id) + .is_some_and(|existing_id| existing_id != &task.id) + { + return Err(DataLayerError::InvalidInput(format!( + "video task {} conflicts with existing short_id {short_id}", + task.id + ))); + } + } + if index + .request_to_id + .get(&task.request_id) + .is_some_and(|existing_id| existing_id != &task.id) + { + return Err(DataLayerError::InvalidInput(format!( + "video task {} conflicts with existing request_id {}", + task.id, task.request_id + ))); + } + Ok(()) + } + fn matches_filter(task: &StoredVideoTask, filter: &VideoTaskQueryFilter) -> bool { if let Some(user_id) = filter.user_id.as_deref() { if task.user_id.as_deref() != Some(user_id) { @@ -103,6 +137,36 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository { }) } + async fn find_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError> { + let index = self.index.read().expect("video task repository lock"); + let task = match key { + VideoTaskLookupKey::Id(id) => index.by_id.get(id), + VideoTaskLookupKey::ShortId(short_id) => index + .short_to_id + .get(short_id) + .and_then(|id| index.by_id.get(id)), + VideoTaskLookupKey::UserExternal { + user_id: lookup_user_id, + external_task_id, + } => { + if lookup_user_id != user_id { + return Ok(None); + } + index + .user_external_to_id + .get(&(lookup_user_id.to_string(), external_task_id.to_string())) + .and_then(|id| index.by_id.get(id)) + } + }; + Ok(task + .filter(|task| task.user_id.as_deref() == Some(user_id)) + .cloned()) + } + async fn list_active(&self, limit: usize) -> Result, DataLayerError> { if limit == 0 { return Ok(Vec::new()); @@ -306,22 +370,34 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository { #[async_trait] impl VideoTaskWriteRepository for InMemoryVideoTaskRepository { - async fn upsert(&self, task: UpsertVideoTask) -> Result { + async fn upsert(&self, mut task: UpsertVideoTask) -> Result { let mut index = self.index.write().expect("video task repository lock"); + Self::ensure_unique_keys_available(&index, &task)?; + if let Some(existing) = index.by_id.get(&task.id) { + existing.ensure_immutable_identity_matches(&task)?; + task.created_at_unix_ms = existing.created_at_unix_ms; + } Ok(Self::store_locked(&mut index, task.into_stored())) } async fn update_if_active( &self, - task: UpsertVideoTask, + mut task: UpsertVideoTask, ) -> Result, DataLayerError> { let mut index = self.index.write().expect("video task repository lock"); + if Self::ensure_unique_keys_available(&index, &task).is_err() { + return Ok(None); + } let Some(existing) = index.by_id.get(&task.id) else { return Ok(None); }; if !existing.status.is_active() { return Ok(None); } + if existing.ensure_immutable_identity_matches(&task).is_err() { + return Ok(None); + } + task.created_at_unix_ms = existing.created_at_unix_ms; Ok(Some(Self::store_locked(&mut index, task.into_stored()))) } @@ -455,6 +531,34 @@ mod tests { .is_some()); } + #[tokio::test] + async fn owner_scoped_lookup_rejects_foreign_user_for_every_identifier() { + let repo = InMemoryVideoTaskRepository::default(); + repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100)) + .await + .expect("upsert should succeed"); + + for key in [ + VideoTaskLookupKey::Id("task-1"), + VideoTaskLookupKey::ShortId("short-task-1"), + VideoTaskLookupKey::UserExternal { + user_id: "user-1", + external_task_id: "ext-task-1", + }, + ] { + assert!(repo + .find_for_user(key, "user-1") + .await + .expect("owner lookup should succeed") + .is_some()); + assert!(repo + .find_for_user(key, "user-2") + .await + .expect("foreign lookup should succeed") + .is_none()); + } + } + #[tokio::test] async fn list_active_only_returns_active_tasks_in_descending_update_order() { let repo = InMemoryVideoTaskRepository::default(); @@ -478,59 +582,69 @@ mod tests { } #[tokio::test] - async fn upsert_replaces_secondary_indexes() { + async fn upsert_rejects_immutable_identity_replacement() { let repo = InMemoryVideoTaskRepository::default(); repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100)) .await .expect("upsert should succeed"); - repo.upsert(UpsertVideoTask { - id: "task-1".to_string(), - short_id: Some("short-task-1b".to_string()), - request_id: "request-task-1b".to_string(), - user_id: Some("user-2".to_string()), - api_key_id: Some("api-key-2".to_string()), - username: Some("user-2".to_string()), - api_key_name: Some("secondary".to_string()), - external_task_id: Some("ext-task-1b".to_string()), - provider_id: Some("provider-2".to_string()), - endpoint_id: Some("endpoint-2".to_string()), - key_id: Some("provider-key-2".to_string()), - client_api_format: Some("gemini:video".to_string()), - provider_api_format: Some("gemini:video".to_string()), - format_converted: false, - model: Some("veo-3".to_string()), - prompt: Some("remix".to_string()), - original_request_body: Some(serde_json::json!({"prompt": "remix"})), - duration_seconds: Some(8), - resolution: Some("1080p".to_string()), - aspect_ratio: Some("16:9".to_string()), - size: Some("720p".to_string()), - status: VideoTaskStatus::Processing, - progress_percent: 50, - progress_message: Some("processing".to_string()), - retry_count: 1, - poll_interval_seconds: 10, - next_poll_at_unix_secs: Some(200), - poll_count: 2, - max_poll_count: 360, - created_at_unix_ms: 150, - submitted_at_unix_secs: Some(150), - completed_at_unix_secs: None, - updated_at_unix_secs: 200, - error_code: None, - error_message: None, - video_url: None, - request_metadata: None, - }) - .await - .expect("upsert should succeed"); + let conflict = repo + .upsert(UpsertVideoTask { + id: "task-1".to_string(), + short_id: Some("short-task-1b".to_string()), + request_id: "request-task-1b".to_string(), + user_id: Some("user-2".to_string()), + api_key_id: Some("api-key-2".to_string()), + username: Some("user-2".to_string()), + api_key_name: Some("secondary".to_string()), + external_task_id: Some("ext-task-1b".to_string()), + provider_id: Some("provider-2".to_string()), + endpoint_id: Some("endpoint-2".to_string()), + key_id: Some("provider-key-2".to_string()), + client_api_format: Some("gemini:video".to_string()), + provider_api_format: Some("gemini:video".to_string()), + format_converted: false, + model: Some("veo-3".to_string()), + prompt: Some("remix".to_string()), + original_request_body: Some(serde_json::json!({"prompt": "remix"})), + duration_seconds: Some(8), + resolution: Some("1080p".to_string()), + aspect_ratio: Some("16:9".to_string()), + size: Some("720p".to_string()), + status: VideoTaskStatus::Processing, + progress_percent: 50, + progress_message: Some("processing".to_string()), + retry_count: 1, + poll_interval_seconds: 10, + next_poll_at_unix_secs: Some(200), + poll_count: 2, + max_poll_count: 360, + created_at_unix_ms: 150, + submitted_at_unix_secs: Some(150), + completed_at_unix_secs: None, + updated_at_unix_secs: 200, + error_code: None, + error_message: None, + video_url: None, + request_metadata: None, + }) + .await + .expect_err("identity replacement should be rejected"); + assert!(conflict.to_string().contains("immutable field short_id")); + let stored = repo + .find(VideoTaskLookupKey::Id("task-1")) + .await + .expect("find should succeed") + .expect("original task should remain"); + assert_eq!(stored.request_id, "request-task-1"); + assert_eq!(stored.user_id.as_deref(), Some("user-1")); + assert_eq!(stored.status, VideoTaskStatus::Submitted); assert!(repo .find(VideoTaskLookupKey::ShortId("short-task-1")) .await .expect("find should succeed") - .is_none()); + .is_some()); assert!(repo .find(VideoTaskLookupKey::UserExternal { user_id: "user-1", @@ -538,12 +652,117 @@ mod tests { }) .await .expect("find should succeed") - .is_none()); + .is_some()); assert!(repo .find(VideoTaskLookupKey::ShortId("short-task-1b")) .await .expect("find should succeed") + .is_none()); + } + + #[tokio::test] + async fn upsert_allows_same_identity_status_update() { + let repo = InMemoryVideoTaskRepository::default(); + let task = sample_task("task-1", VideoTaskStatus::Submitted, 100); + repo.upsert(task.clone()) + .await + .expect("initial upsert should succeed"); + + let updated = repo + .upsert(UpsertVideoTask { + status: VideoTaskStatus::Processing, + progress_percent: 50, + poll_count: 2, + created_at_unix_ms: 999, + updated_at_unix_secs: 200, + ..task + }) + .await + .expect("same identity update should succeed"); + + assert_eq!(updated.status, VideoTaskStatus::Processing); + assert_eq!(updated.progress_percent, 50); + assert_eq!(updated.poll_count, 2); + assert_eq!(updated.created_at_unix_ms, 90); + assert_eq!(updated.updated_at_unix_secs, 200); + } + + #[tokio::test] + async fn update_if_active_rejects_identity_conflict_without_modification() { + let repo = InMemoryVideoTaskRepository::default(); + let task = sample_task("task-1", VideoTaskStatus::Submitted, 100); + repo.upsert(task.clone()) + .await + .expect("initial upsert should succeed"); + + let result = repo + .update_if_active(UpsertVideoTask { + user_id: Some("attacker".to_string()), + status: VideoTaskStatus::Completed, + progress_percent: 100, + updated_at_unix_secs: 200, + ..task + }) + .await + .expect("guarded update should execute"); + assert!(result.is_none()); + + let stored = repo + .find(VideoTaskLookupKey::Id("task-1")) + .await + .expect("find should succeed") + .expect("original task should remain"); + assert_eq!(stored.user_id.as_deref(), Some("user-1")); + assert_eq!(stored.status, VideoTaskStatus::Submitted); + assert_eq!(stored.progress_percent, 0); + assert_eq!(stored.updated_at_unix_secs, 100); + } + + #[tokio::test] + async fn upsert_rejects_secondary_unique_key_takeover() { + let repo = InMemoryVideoTaskRepository::default(); + let original = sample_task("task-1", VideoTaskStatus::Submitted, 100); + repo.upsert(original.clone()) + .await + .expect("initial upsert should succeed"); + + let short_id_conflict = repo + .upsert(UpsertVideoTask { + id: "task-2".to_string(), + request_id: "request-task-2".to_string(), + ..original.clone() + }) + .await + .expect_err("a short id must not be reassigned to another task"); + assert!(short_id_conflict.to_string().contains("existing short_id")); + + let request_id_conflict = repo + .upsert(UpsertVideoTask { + id: "task-3".to_string(), + short_id: Some("short-task-3".to_string()), + ..original + }) + .await + .expect_err("a request id must not be reassigned to another task"); + assert!(request_id_conflict + .to_string() + .contains("existing request_id")); + + assert!(repo + .find(VideoTaskLookupKey::ShortId("short-task-1")) + .await + .expect("find should succeed") .is_some()); + assert!(repo + .find(VideoTaskLookupKey::Id("task-2")) + .await + .expect("find should succeed") + .is_none()); + assert!(repo + .find(VideoTaskLookupKey::Id("task-3")) + .await + .expect("find should succeed") + .is_none()); } #[tokio::test] diff --git a/crates/aether-data/runtime/src/repository/wallet/memory.rs b/crates/aether-data/runtime/src/repository/wallet/memory.rs index a2a19af4c..08df3efc8 100644 --- a/crates/aether-data/runtime/src/repository/wallet/memory.rs +++ b/crates/aether-data/runtime/src/repository/wallet/memory.rs @@ -1,22 +1,35 @@ use std::collections::BTreeMap; -use std::sync::RwLock; +use std::sync::{Mutex, RwLock}; use async_trait::async_trait; #[cfg(test)] use super::WalletReadSeed; use super::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, - redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, - AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, - AdminWalletListQuery, AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, - CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + canonicalize_payment_method, canonicalize_wallet_refund_fields, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_payment_method, + redeem_code_refundable_amount, validate_admin_redeem_code_batch_input, + validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements, + validate_redeem_wallet_credit, validate_wallet_recharge_order_input, + wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token, + wallet_recharge_checkout_failed_response, wallet_recharge_checkout_uncertain_response, + 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, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, + AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, + AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, + CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, + ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, @@ -24,7 +37,8 @@ use super::{ StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use crate::DataLayerError; @@ -39,9 +53,222 @@ pub struct InMemoryWalletRepository { redeem_batches_by_id: RwLock>, redeem_codes_by_id: RwLock>, redeem_code_hash_to_id: RwLock>, + refund_idempotency_to_id: RwLock>, + refund_creation_lock: Mutex<()>, + // Wallet creation, order attachment, and compensation cleanup must be + // serialized together. Separate map locks cannot make the + // "check references, then delete" sequence atomic. + wallet_lifecycle_lock: Mutex<()>, +} + +fn wallet_recharge_metadata_value<'a>( + record: &'a StoredAdminPaymentOrder, + key: &str, +) -> Option<&'a str> { + record + .gateway_response + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|object| object.get(key)) + .and_then(serde_json::Value::as_str) +} + +/// A compensation path may remove only a freshly initialized wallet that has +/// never carried funds or a financial reference. Keep this predicate in one +/// place so the initial check and the final compare use exactly the same +/// definition even when another in-memory writer updates the row between +/// those checks. +fn wallet_is_untouched_for_compensation(wallet: &StoredWalletSnapshot) -> bool { + wallet.balance == 0.0 + && wallet.gift_balance == 0.0 + && wallet.total_recharged == 0.0 + && wallet.total_consumed == 0.0 + && wallet.total_refunded == 0.0 + && wallet.total_adjusted == 0.0 + && matches!(wallet.limit_mode.as_str(), "finite" | "unlimited") + && wallet.currency == "USD" + && wallet.status == "active" +} + +fn provisional_auth_wallet_transactions_match( + transactions: &BTreeMap, + wallet: &StoredWalletSnapshot, + user_id: &str, +) -> bool { + let wallet_transactions = transactions + .values() + .filter(|transaction| transaction.wallet_id == wallet.id) + .collect::>(); + if wallet.gift_balance == 0.0 { + return wallet_transactions.is_empty(); + } + if wallet_transactions.len() != 1 { + return false; + } + let transaction = wallet_transactions[0]; + transaction.category == "gift" + && transaction.reason_code == "gift_initial" + && transaction.amount == wallet.gift_balance + && transaction.balance_before == 0.0 + && transaction.balance_after == wallet.gift_balance + && transaction.recharge_balance_before == 0.0 + && transaction.recharge_balance_after == 0.0 + && transaction.gift_balance_before == 0.0 + && transaction.gift_balance_after == wallet.gift_balance + && transaction.link_type.as_deref() == Some("system_task") + && transaction.link_id.as_deref() == Some(user_id) + && transaction.operator_id.is_none() } impl InMemoryWalletRepository { + fn insert_payment_order_unique( + &self, + order: StoredAdminPaymentOrder, + ) -> Result<(), DataLayerError> { + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + if orders.contains_key(&order.id) + || orders + .values() + .any(|existing| existing.order_no == order.order_no) + { + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another order".to_string(), + )); + } + if let Some(gateway_order_id) = order.gateway_order_id.as_deref() { + if orders.values().any(|existing| { + existing.payment_method == order.payment_method + && existing.gateway_order_id.as_deref() == Some(gateway_order_id) + }) { + return Err(DataLayerError::InvalidInput( + "payment gateway order already belongs to another order".to_string(), + )); + } + } + orders.insert(order.id.clone(), order); + Ok(()) + } + + fn insert_wallet_recharge_order_unique( + &self, + order: StoredAdminPaymentOrder, + now_unix_secs: u64, + input: &CreateWalletRechargeOrderInput, + ) -> Result<(Option, bool), DataLayerError> { + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + if let Some(existing) = orders + .values() + .find(|existing| existing.order_no == order.order_no) + { + let existing_kind = existing + .gateway_response + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|object| object.get("order_kind")) + .and_then(serde_json::Value::as_str); + if existing.user_id == order.user_id && existing_kind == Some("wallet_recharge") { + let existing_provider = + wallet_recharge_metadata_value(existing, "payment_provider") + .or_else(|| wallet_recharge_metadata_value(existing, "gateway")); + // A seeded legacy order can outlive a deleted/provisioning + // wallet. In that read-model-only case the incoming wallet + // id is provisional and cannot be treated as an immutable + // field; a real wallet remains strictly bound below. + let replay_wallet_id = if self + .wallets_by_id + .read() + .expect("wallet repo lock") + .contains_key(&existing.wallet_id) + { + order.wallet_id.as_str() + } else { + existing.wallet_id.as_str() + }; + if !wallet_recharge_replay_matches( + &existing.wallet_id, + existing.amount_usd, + existing.pay_amount, + existing.pay_currency.as_deref(), + existing.exchange_rate, + &existing.payment_method, + existing_provider, + wallet_recharge_metadata_value(existing, "payment_channel"), + replay_wallet_id, + input, + ) { + return Err(DataLayerError::InvalidInput( + "wallet recharge replay changes immutable order fields".to_string(), + )); + } + if wallet_recharge_order_is_reclaimable_placeholder(existing, now_unix_secs) { + let Some(candidate_token) = order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token) + else { + return Ok((Some(existing.clone()), false)); + }; + let claimed = wallet_recharge_checkout_claim_response( + order.gateway_response.as_ref().expect("gateway response"), + candidate_token, + now_unix_secs, + ) + .map_err(DataLayerError::InvalidInput)?; + let existing_id = existing.id.clone(); + let existing = orders + .get_mut(&existing_id) + .expect("recharge order should remain present"); + existing.gateway_response = Some(claimed); + existing.gateway_order_id = Some(existing.order_no.clone()); + existing.status = "pending".to_string(); + existing.expires_at_unix_secs = order.expires_at_unix_secs; + return Ok((Some(existing.clone()), true)); + } + return Ok((Some(existing.clone()), false)); + } + return Err(DataLayerError::InvalidInput( + "payment order number already belongs to another user".to_string(), + )); + } + if let Some(gateway_order_id) = order.gateway_order_id.as_deref() { + if orders.values().any(|existing| { + existing.payment_method == order.payment_method + && existing.gateway_order_id.as_deref() == Some(gateway_order_id) + }) { + return Err(DataLayerError::InvalidInput( + "payment gateway order already belongs to another order".to_string(), + )); + } + } + orders.insert(order.id.clone(), order); + Ok((None, false)) + } + + fn remove_created_wallet_if_unreferenced(&self, wallet_id: &str) { + // The caller holds wallet_lifecycle_lock, as do all in-memory paths + // that attach a payment order. The reference check and compensation + // delete are therefore atomic with respect to a new order attachment. + let wallet_is_referenced = self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|existing| existing.wallet_id == wallet_id); + if !wallet_is_referenced { + let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + // A regular wallet snapshot update does not need the lifecycle + // lock. Re-check the complete pristine shape while holding the + // write lock so a concurrent credit cannot be deleted as part of + // compensation. + if wallets + .get(wallet_id) + .is_some_and(wallet_is_untouched_for_compensation) + { + wallets.remove(wallet_id); + } + } + } + pub fn seed(items: I) -> Self where I: IntoIterator, @@ -59,6 +286,9 @@ impl InMemoryWalletRepository { redeem_batches_by_id: RwLock::new(BTreeMap::new()), redeem_codes_by_id: RwLock::new(BTreeMap::new()), redeem_code_hash_to_id: RwLock::new(BTreeMap::new()), + refund_idempotency_to_id: RwLock::new(BTreeMap::new()), + refund_creation_lock: Mutex::new(()), + wallet_lifecycle_lock: Mutex::new(()), } } @@ -84,6 +314,11 @@ impl InMemoryWalletRepository { for item in seed.refunds { refunds_by_id.insert(item.id.clone(), item); } + let refund_idempotency_to_id = seed + .refund_idempotency + .into_iter() + .map(|(user_id, idempotency_key, refund_id)| ((user_id, idempotency_key), refund_id)) + .collect(); let mut redeem_batches_by_id = BTreeMap::new(); for item in seed.redeem_batches { redeem_batches_by_id.insert(item.id.clone(), item); @@ -102,6 +337,9 @@ impl InMemoryWalletRepository { redeem_batches_by_id: RwLock::new(redeem_batches_by_id), redeem_codes_by_id: RwLock::new(redeem_codes_by_id), redeem_code_hash_to_id: RwLock::new(BTreeMap::new()), + refund_idempotency_to_id: RwLock::new(refund_idempotency_to_id), + refund_creation_lock: Mutex::new(()), + wallet_lifecycle_lock: Mutex::new(()), } } @@ -176,7 +414,40 @@ fn initialize_auth_wallet_in_memory( api_key_id: Option<&str>, initial_gift_usd: f64, unlimited: bool, -) -> Result, DataLayerError> { +) -> Result, DataLayerError> { + let owner_id = user_id + .or(api_key_id) + .filter(|value| !value.trim().is_empty()); + if owner_id.is_none() || (user_id.is_some() && api_key_id.is_some()) { + return Err(DataLayerError::InvalidInput( + "wallet owner must be exactly one non-empty user or API-key id".to_string(), + )); + } + if !initial_gift_usd.is_finite() { + return Err(DataLayerError::InvalidInput( + "initial gift amount must be finite".to_string(), + )); + } + + // Initialization is intentionally idempotent. Database backends enforce this + // with the owner unique indexes; perform the same lookup before creating the + // in-memory row so retries cannot mint another wallet or gift transaction. + { + let wallets = wallets_by_id.read().expect("wallet repo lock"); + let existing = wallets.values().find(|wallet| { + if let Some(user_id) = user_id { + wallet.user_id.as_deref() == Some(user_id) && wallet.api_key_id.is_none() + } else if let Some(api_key_id) = api_key_id { + wallet.api_key_id.as_deref() == Some(api_key_id) && wallet.user_id.is_none() + } else { + false + } + }); + if let Some(wallet) = existing { + return Ok(Some((wallet.clone(), false))); + } + } + let gift_amount = if unlimited { 0.0 } else { @@ -235,7 +506,7 @@ fn initialize_auth_wallet_in_memory( .insert(transaction.id.clone(), transaction); } - Ok(Some(wallet)) + Ok(Some((wallet, true))) } fn normalize_redeem_code(value: &str) -> Option { @@ -335,6 +606,15 @@ impl WalletReadRepository for InMemoryWalletRepository { initial_gift_usd: f64, unlimited: bool, ) -> Result, DataLayerError> { + if user_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "user id is required to initialize a wallet".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); initialize_auth_wallet_in_memory( &self.wallets_by_id, &self.wallet_transactions_by_id, @@ -343,6 +623,35 @@ impl WalletReadRepository for InMemoryWalletRepository { initial_gift_usd, unlimited, ) + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_user_wallet_with_outcome( + &self, + user_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + if user_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "user id is required to initialize a wallet".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + initialize_auth_wallet_in_memory( + &self.wallets_by_id, + &self.wallet_transactions_by_id, + Some(user_id), + None, + initial_gift_usd, + unlimited, + ) + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn initialize_auth_api_key_wallet( @@ -351,6 +660,15 @@ impl WalletReadRepository for InMemoryWalletRepository { initial_gift_usd: f64, unlimited: bool, ) -> Result, DataLayerError> { + if api_key_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "api key id is required to initialize a wallet".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); initialize_auth_wallet_in_memory( &self.wallets_by_id, &self.wallet_transactions_by_id, @@ -359,6 +677,35 @@ impl WalletReadRepository for InMemoryWalletRepository { initial_gift_usd, unlimited, ) + .map(|result| result.map(|(wallet, _created)| wallet)) + } + + async fn initialize_auth_api_key_wallet_with_outcome( + &self, + api_key_id: &str, + initial_gift_usd: f64, + unlimited: bool, + ) -> Result, DataLayerError> { + if api_key_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "api key id is required to initialize a wallet".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + initialize_auth_wallet_in_memory( + &self.wallets_by_id, + &self.wallet_transactions_by_id, + None, + Some(api_key_id), + initial_gift_usd, + unlimited, + ) + .map(|result| { + result.map(|(wallet, created)| InitializeAuthWalletOutcome { wallet, created }) + }) } async fn update_auth_user_wallet_snapshot( @@ -534,7 +881,7 @@ impl WalletReadRepository for InMemoryWalletRepository { &self, query: &AdminWalletRefundRequestListQuery, ) -> Result { - let wallets = self.wallets_by_id.read().expect("wallet repo lock"); + let wallets = self.wallets_by_id.read().expect("wallet repo lock").clone(); let mut items = self .refunds_by_id .read() @@ -660,7 +1007,7 @@ impl WalletReadRepository for InMemoryWalletRepository { .filter(|order| { query.status.as_deref().is_none_or(|expected| { let effective = if order.status == "pending" - && order.expires_at_unix_secs.is_some_and(|value| value < now) + && order.expires_at_unix_secs.is_some_and(|value| value <= now) { "expired" } else { @@ -707,7 +1054,15 @@ impl WalletReadRepository for InMemoryWalletRepository { .read() .expect("wallet repo lock") .values() - .filter(|order| order.user_id.as_deref() == Some(user_id)) + .filter(|order| { + order.user_id.as_deref() == Some(user_id) + && order + .gateway_response + .as_ref() + .and_then(|value| value.get("order_kind")) + .and_then(serde_json::Value::as_str) + != Some("plan_purchase") + }) .cloned() .collect::>(); items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms)); @@ -757,7 +1112,39 @@ impl WalletReadRepository for InMemoryWalletRepository { .read() .expect("wallet repo lock") .get(order_id) - .filter(|order| order.user_id.as_deref() == Some(user_id)) + .filter(|order| { + order.user_id.as_deref() == Some(user_id) + && order + .gateway_response + .as_ref() + .and_then(|value| value.get("order_kind")) + .and_then(serde_json::Value::as_str) + != Some("plan_purchase") + }) + .cloned()) + } + + async fn find_wallet_recharge_order_by_order_no( + &self, + user_id: &str, + order_no: &str, + ) -> Result, DataLayerError> { + Ok(self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .find(|order| { + order.user_id.as_deref() == Some(user_id) + && order.order_no == order_no + && order + .gateway_response + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|object| object.get("order_kind")) + .and_then(serde_json::Value::as_str) + == Some("wallet_recharge") + }) .cloned()) } @@ -796,6 +1183,19 @@ impl WalletReadRepository for InMemoryWalletRepository { .cloned()) } + async fn find_payment_order_by_order_no( + &self, + order_no: &str, + ) -> Result, DataLayerError> { + Ok(self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .find(|order| order.order_no == order_no) + .cloned()) + } + async fn find_wallet_refund( &self, wallet_id: &str, @@ -902,30 +1302,402 @@ impl WalletReadRepository for InMemoryWalletRepository { #[async_trait] impl WalletWriteRepository for InMemoryWalletRepository { + async fn delete_wallet_if_unreferenced( + &self, + wallet_id: &str, + owner: WalletLookupKey<'_>, + ) -> Result { + if wallet_id.trim().is_empty() { + return Ok(false); + } + let owner_matches = |wallet: &StoredWalletSnapshot| match owner { + WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.id == wallet_id + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.id == wallet_id + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if matches!(owner, WalletLookupKey::WalletId(_)) { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + + let wallet = { + let wallets = self.wallets_by_id.read().expect("wallet repo lock"); + wallets + .values() + .find(|wallet| owner_matches(wallet)) + .cloned() + }; + let Some(wallet) = wallet else { + return Ok(false); + }; + + // Compensation is only allowed for an untouched, freshly-created wallet. A journal + // entry can race with an existing wallet lookup, and deleting a zero-reference wallet + // with persisted funds would otherwise destroy those funds. + if !wallet_is_untouched_for_compensation(&wallet) { + return Ok(false); + } + let referenced = self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|order| order.wallet_id == wallet.id) + || self + .refunds_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|refund| refund.wallet_id == wallet.id) + || self + .wallet_transactions_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|transaction| transaction.wallet_id == wallet.id) + || self + .redeem_codes_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|code| code.redeemed_wallet_id.as_deref() == Some(wallet.id.as_str())); + if referenced { + return Ok(false); + } + + let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + let removable = wallets.get(&wallet.id).is_some_and(|current| { + owner_matches(current) + && current == &wallet + && wallet_is_untouched_for_compensation(current) + }); + if !removable { + return Ok(false); + } + Ok(wallets.remove(&wallet.id).is_some()) + } + + async fn delete_wallet_if_snapshot_matches_and_unreferenced( + &self, + expected: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if expected.id.trim().is_empty() { + return Ok(false); + } + let owner_matches = |wallet: &StoredWalletSnapshot| match owner { + WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.id == expected.id + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.id == expected.id + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if matches!(owner, WalletLookupKey::WalletId(_)) { + return Err(DataLayerError::InvalidInput( + "wallet compensation requires an explicit user or API-key owner".to_string(), + )); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + + let current = { + let wallets = self.wallets_by_id.read().expect("wallet repo lock"); + wallets + .get(&expected.id) + .filter(|wallet| owner_matches(wallet)) + .cloned() + }; + // Compare every field, including the owner and update timestamp. A + // mismatch means another operation touched the wallet, so compensation + // must fail closed and preserve its funds. + if current.as_ref() != Some(expected) { + return Ok(false); + } + + let referenced = self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|order| order.wallet_id == expected.id) + || self + .refunds_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|refund| refund.wallet_id == expected.id) + || self + .wallet_transactions_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|transaction| transaction.wallet_id == expected.id) + || self + .redeem_codes_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|code| code.redeemed_wallet_id.as_deref() == Some(expected.id.as_str())); + if referenced { + return Ok(false); + } + + let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + if wallets + .get(&expected.id) + .is_some_and(|wallet| owner_matches(wallet) && wallet == expected) + { + return Ok(wallets.remove(&expected.id).is_some()); + } + Ok(false) + } + + async fn restore_wallet_if_snapshot_matches( + &self, + before: &StoredWalletSnapshot, + after: &StoredWalletSnapshot, + owner: WalletLookupKey<'_>, + ) -> Result { + if before.id.trim().is_empty() || after.id.trim().is_empty() { + return Ok(false); + } + if before.id != after.id { + return Err(DataLayerError::InvalidInput( + "wallet restore snapshots must reference the same wallet".to_string(), + )); + } + let owner_matches = |wallet: &StoredWalletSnapshot| match owner { + WalletLookupKey::UserId(user_id) => { + !user_id.trim().is_empty() + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + } + WalletLookupKey::ApiKeyId(api_key_id) => { + !api_key_id.trim().is_empty() + && wallet.api_key_id.as_deref() == Some(api_key_id) + && wallet.user_id.is_none() + } + WalletLookupKey::WalletId(_) => false, + }; + if matches!(owner, WalletLookupKey::WalletId(_)) { + return Err(DataLayerError::InvalidInput( + "wallet restore requires an explicit user or API-key owner".to_string(), + )); + } + + // Keep the compare and replacement atomic with the lifecycle operations. The map write + // lock also prevents an ordinary wallet mutation from interleaving between the compare + // and restore; a changed snapshot therefore fails closed instead of being overwritten. + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + let Some(current) = wallets.get(&after.id) else { + return Ok(false); + }; + if current != after || !owner_matches(current) || !owner_matches(before) { + return Ok(false); + } + wallets.insert(before.id.clone(), before.clone()); + Ok(true) + } + + async fn delete_provisional_auth_user_wallet( + &self, + wallet_id: &str, + user_id: &str, + ) -> Result { + if wallet_id.trim().is_empty() || user_id.trim().is_empty() { + return Ok(false); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + + // Provisioning rollback is deliberately fail-closed. The only + // transaction that may exist is the deterministic initial gift entry; + // any other financial artifact makes the wallet ineligible for purge. + let wallet = { + let wallets = self.wallets_by_id.read().expect("wallet repo lock"); + wallets + .values() + .find(|wallet| { + wallet.id == wallet_id + && wallet.user_id.as_deref() == Some(user_id) + && wallet.api_key_id.is_none() + && wallet.balance == 0.0 + && wallet.total_recharged == 0.0 + && wallet.total_consumed == 0.0 + && wallet.total_refunded == 0.0 + && wallet.total_adjusted == wallet.gift_balance + && wallet.gift_balance >= 0.0 + && wallet.status == "active" + && matches!(wallet.limit_mode.as_str(), "finite" | "unlimited") + && wallet.currency == "USD" + }) + .cloned() + }; + let Some(wallet) = wallet else { + return Ok(false); + }; + + let transactions = self + .wallet_transactions_by_id + .read() + .expect("wallet repo lock"); + let transaction_matches = + provisional_auth_wallet_transactions_match(&transactions, &wallet, user_id); + drop(transactions); + if !transaction_matches { + return Ok(false); + } + + if self + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|order| order.wallet_id == wallet.id) + || self + .refunds_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|refund| refund.wallet_id == wallet.id) + || self + .redeem_codes_by_id + .read() + .expect("wallet repo lock") + .values() + .any(|code| code.redeemed_wallet_id.as_deref() == Some(wallet.id.as_str())) + { + return Ok(false); + } + + // Snapshot updates do not take `wallet_lifecycle_lock`; hold the + // wallet write lock through the final compare and removal so a credit + // cannot land after the check but before deletion. Re-check the + // transaction set under its write lock as well, since a concurrent + // financial entry must make this compensation fail closed. + let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + if wallets.get(&wallet.id) != Some(&wallet) { + return Ok(false); + } + let mut transactions = self + .wallet_transactions_by_id + .write() + .expect("wallet repo lock"); + if !provisional_auth_wallet_transactions_match(&transactions, &wallet, user_id) { + return Ok(false); + } + transactions.retain(|_, transaction| transaction.wallet_id != wallet.id); + Ok(wallets.remove(&wallet.id).is_some()) + } + async fn create_wallet_recharge_order( &self, - input: CreateWalletRechargeOrderInput, + mut input: CreateWalletRechargeOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_wallet_recharge_order_input(&input).map_err(DataLayerError::InvalidInput)?; + if !input.amount_usd.is_finite() + || input.amount_usd <= 0.0 + || input + .pay_amount + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input + .exchange_rate + .is_some_and(|value| !value.is_finite() || value <= 0.0) + || input.expires_at_unix_secs > i64::MAX as u64 + { + return Err(DataLayerError::InvalidInput( + "invalid wallet recharge numeric fields".to_string(), + )); + } + let mut gateway_response = + project_wallet_recharge_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; + if !gateway_response.is_object() { + return Err(DataLayerError::InvalidInput( + "wallet recharge gateway response must be an object".to_string(), + )); + } + let gateway_object = gateway_response + .as_object_mut() + .expect("validated wallet recharge gateway response"); + if let Some(provider) = input.payment_provider.as_deref() { + gateway_object.insert( + "payment_provider".to_string(), + serde_json::Value::String(provider.trim().to_ascii_lowercase()), + ); + } + if let Some(channel) = input.payment_channel.as_deref() { + gateway_object.insert( + "payment_channel".to_string(), + serde_json::Value::String(channel.trim().to_ascii_lowercase()), + ); + } + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); let now_secs = current_unix_secs(); - let wallet_id = { + // Keep the wallet and payment-order locks in separate scopes. The + // repository stores them in independent maps, and holding one while + // acquiring the other can deadlock with refund/order readers. + let (wallet_id, created_wallet) = { let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); - let wallet = wallets - .values_mut() - .find(|wallet| wallet.user_id.as_deref() == Some(input.user_id.as_str())); - if wallet + let existing_wallet = wallets + .values() + .find(|wallet| wallet.user_id.as_deref() == Some(input.user_id.as_str())) + .map(|wallet| (wallet.id.clone(), wallet.status.clone())); + if existing_wallet .as_ref() - .is_some_and(|wallet| wallet.status != "active") + .is_some_and(|(_, status)| status != "active") { return Ok(CreateWalletRechargeOrderOutcome::WalletInactive); } - match wallet { - Some(wallet) => wallet.id.clone(), + match existing_wallet { + Some((wallet_id, _)) => (wallet_id, false), None => { let wallet_id = input .preferred_wallet_id .clone() .unwrap_or_else(|| format!("wallet-{}", uuid::Uuid::new_v4())); + if wallets.contains_key(&wallet_id) { + return Err(DataLayerError::InvalidInput( + "wallet identifier already belongs to another owner".to_string(), + )); + } let wallet = StoredWalletSnapshot::new( wallet_id.clone(), Some(input.user_id.clone()), @@ -942,15 +1714,16 @@ impl WalletWriteRepository for InMemoryWalletRepository { now_secs as i64, )?; wallets.insert(wallet_id.clone(), wallet); - wallet_id + (wallet_id, true) } } }; + let replay_input = input.clone(); let order = StoredAdminPaymentOrder { id: format!("payment-order-{}", uuid::Uuid::new_v4()), order_no: input.order_no, - wallet_id, + wallet_id: wallet_id.clone(), user_id: Some(input.user_id), amount_usd: input.amount_usd, pay_amount: input.pay_amount, @@ -959,25 +1732,277 @@ impl WalletWriteRepository for InMemoryWalletRepository { refunded_amount_usd: 0.0, refundable_amount_usd: 0.0, payment_method: input.payment_method, + payment_provider: input.payment_provider, + order_kind: "wallet_recharge".to_string(), gateway_order_id: Some(input.gateway_order_id), - gateway_response: Some(input.gateway_response), + gateway_response: Some(gateway_response), status: "pending".to_string(), created_at_unix_ms: current_unix_ms(), paid_at_unix_secs: None, credited_at_unix_secs: None, expires_at_unix_secs: Some(input.expires_at_unix_secs), }; - self.payment_orders_by_id - .write() - .expect("wallet repo lock") - .insert(order.id.clone(), order.clone()); + match self.insert_wallet_recharge_order_unique(order.clone(), now_secs, &replay_input) { + Ok((Some(existing), true)) => { + if created_wallet { + self.remove_created_wallet_if_unreferenced(&wallet_id); + } + return Ok(CreateWalletRechargeOrderOutcome::Created(existing)); + } + Ok((Some(existing), false)) => { + if created_wallet { + self.remove_created_wallet_if_unreferenced(&wallet_id); + } + return Ok(CreateWalletRechargeOrderOutcome::Existing(existing)); + } + Ok((None, false)) => {} + Ok((None, true)) => { + if created_wallet { + self.remove_created_wallet_if_unreferenced(&wallet_id); + } + return Err(DataLayerError::InvalidInput( + "reclaimed recharge order disappeared".to_string(), + )); + } + Err(error) => { + if created_wallet { + self.remove_created_wallet_if_unreferenced(&wallet_id); + } + return Err(error); + } + } Ok(CreateWalletRechargeOrderOutcome::Created(order)) } + async fn update_wallet_recharge_checkout( + &self, + input: UpdateWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() || input.gateway_order_id.trim().is_empty() { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout identifiers are required".to_string(), + )); + } + let gateway_response = + match project_wallet_recharge_gateway_response(&input.gateway_response) { + Ok(value) => value, + Err(error) => return Ok(WalletMutationOutcome::Invalid(error)), + }; + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + let Some(current_order) = orders.get(&input.order_id) else { + return Ok(WalletMutationOutcome::NotFound); + }; + let is_wallet_recharge = current_order + .gateway_response + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|object| object.get("order_kind")) + .and_then(serde_json::Value::as_str) + == Some("wallet_recharge"); + if !is_wallet_recharge { + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a wallet recharge".to_string(), + )); + } + let current_is_checkout_placeholder = + wallet_recharge_order_is_checkout_placeholder(current_order); + let current_token = current_order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + let requested_token = wallet_recharge_checkout_claim_token(&gateway_response); + if current_token.is_some() && current_token != requested_token { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if current_order.status != "pending" { + if current_order.gateway_order_id.as_deref() == Some(input.gateway_order_id.as_str()) { + return Ok(WalletMutationOutcome::Applied(current_order.clone())); + } + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is no longer pending".to_string(), + )); + } + let now_secs = current_unix_secs(); + if current_order + .expires_at_unix_secs + .is_none_or(|expires_at| expires_at <= now_secs) + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge order is expired".to_string(), + )); + } + // The initial row stores the order number as a temporary gateway id. + // Once a provider checkout is persisted, a concurrent creator must + // not overwrite that evidence with a second provider checkout. + if current_order + .gateway_order_id + .as_deref() + .is_some_and(|existing| { + existing != input.gateway_order_id.as_str() + && existing != current_order.order_no.as_str() + && !current_is_checkout_placeholder + }) + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is already bound".to_string(), + )); + } + let payment_method = current_order.payment_method.clone(); + if orders.values().any(|existing| { + existing.id != input.order_id + && existing.payment_method == payment_method + && existing.gateway_order_id.as_deref() == Some(input.gateway_order_id.as_str()) + }) { + return Ok(WalletMutationOutcome::Invalid( + "payment gateway order already belongs to another order".to_string(), + )); + } + let order = orders + .get_mut(&input.order_id) + .expect("wallet recharge order disappeared while write lock held"); + order.gateway_order_id = Some(input.gateway_order_id); + order.gateway_response = Some(gateway_response); + Ok(WalletMutationOutcome::Applied(order.clone())) + } + + async fn compare_and_swap_payment_order_stripe_client_secret( + &self, + input: CompareAndSwapPaymentOrderStripeClientSecretInput, + ) -> Result { + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + let Some(current) = orders.get(&input.order_id) else { + return Ok(false); + }; + let Some(replacement) = payment_order_stripe_client_secret_cas_replacement(current, &input) + .map_err(DataLayerError::InvalidInput)? + else { + return Ok(false); + }; + let current = orders + .get_mut(&input.order_id) + .expect("payment order disappeared while write lock held"); + current.gateway_response = Some(replacement); + Ok(true) + } + + async fn fail_wallet_recharge_checkout( + &self, + input: FailWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout failure identifiers are required".to_string(), + )); + } + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + let Some(order) = orders.get_mut(&input.order_id) else { + return Ok(WalletMutationOutcome::NotFound); + }; + if !wallet_recharge_order_is_checkout_placeholder(order) { + return Ok(WalletMutationOutcome::Invalid( + "payment order is not a checkout placeholder".to_string(), + )); + } + let current_token = order + .gateway_response + .as_ref() + .and_then(wallet_recharge_checkout_claim_token); + if current_token != Some(input.claim_token.trim()) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout claim is no longer current".to_string(), + )); + } + if order.status != "pending" { + return Ok(WalletMutationOutcome::Applied(order.clone())); + } + let failed = if input.provider_request_may_have_succeeded { + wallet_recharge_checkout_uncertain_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs(), + ) + } else { + wallet_recharge_checkout_failed_response( + order.gateway_response.as_ref(), + &input.reason, + current_unix_secs(), + ) + }; + order.gateway_response = Some(failed); + order.status = "failed".to_string(); + Ok(WalletMutationOutcome::Applied(order.clone())) + } + + async fn reclaim_wallet_recharge_checkout( + &self, + input: ReclaimWalletRechargeCheckoutInput, + ) -> Result, DataLayerError> { + if input.order_id.trim().is_empty() + || input.claim_token.trim().is_empty() + || input.claim_token.len() > 128 + || input.expires_at_unix_secs <= current_unix_secs() + { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout reclaim identifiers are invalid".to_string(), + )); + } + // Keep the in-memory backend aligned with SQL backends: a reclaim may + // only install a server-created placeholder, never provider checkout + // evidence supplied by an internal caller. + if !wallet_recharge_response_is_checkout_placeholder(&input.gateway_response) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge reclaim response must be a placeholder".to_string(), + )); + } + let now = current_unix_secs(); + let response = wallet_recharge_checkout_claim_response( + &input.gateway_response, + &input.claim_token, + now, + ) + .map_err(DataLayerError::InvalidInput)?; + let mut orders = self.payment_orders_by_id.write().expect("wallet repo lock"); + let Some(order) = orders.get_mut(&input.order_id) else { + return Ok(WalletMutationOutcome::NotFound); + }; + if !wallet_recharge_order_is_reclaimable_placeholder(order, now) { + return Ok(WalletMutationOutcome::Invalid( + "wallet recharge checkout is still in progress or already completed".to_string(), + )); + } + order.gateway_response = Some(response); + order.gateway_order_id = Some(order.order_no.clone()); + order.status = "pending".to_string(); + order.expires_at_unix_secs = Some(input.expires_at_unix_secs); + Ok(WalletMutationOutcome::Applied(order.clone())) + } + async fn create_plan_purchase_order( &self, - input: CreatePlanPurchaseOrderInput, + mut input: CreatePlanPurchaseOrderInput, ) -> Result { + input.payment_method = canonicalize_payment_method(&input.payment_method) + .map_err(DataLayerError::InvalidInput)?; + validate_plan_purchase_order_input(&input).map_err(DataLayerError::InvalidInput)?; + let projected_gateway_response = project_wallet_gateway_response(&input.gateway_response) + .map_err(DataLayerError::InvalidInput)?; + let entitlements = input + .product_snapshot + .get("entitlements") + .or_else(|| input.product_snapshot.get("entitlements_json")) + .cloned() + .unwrap_or_else(|| serde_json::json!([])); + validate_plan_wallet_credit_entitlements(&entitlements) + .map_err(DataLayerError::InvalidInput)?; + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); let wallet_id = { let wallets = self.wallets_by_id.read().expect("wallet repo lock"); let Some(wallet) = wallets @@ -1042,7 +2067,7 @@ impl WalletWriteRepository for InMemoryWalletRepository { return Ok(CreatePlanPurchaseOrderOutcome::ActivePlanLimitReached); } } - let mut gateway_response = match input.gateway_response { + let mut gateway_response = match projected_gateway_response { serde_json::Value::Object(map) => map, value => { let mut map = serde_json::Map::new(); @@ -1071,6 +2096,8 @@ impl WalletWriteRepository for InMemoryWalletRepository { refunded_amount_usd: 0.0, refundable_amount_usd: 0.0, payment_method: input.payment_method, + payment_provider: input.payment_provider, + order_kind: "plan_purchase".to_string(), gateway_order_id: Some(input.gateway_order_id), gateway_response: Some(serde_json::Value::Object(gateway_response)), status: "pending".to_string(), @@ -1079,10 +2106,7 @@ impl WalletWriteRepository for InMemoryWalletRepository { credited_at_unix_secs: None, expires_at_unix_secs: Some(input.expires_at_unix_secs), }; - self.payment_orders_by_id - .write() - .expect("wallet repo lock") - .insert(order.id.clone(), order.clone()); + self.insert_payment_order_unique(order.clone())?; Ok(CreatePlanPurchaseOrderOutcome::Created(order)) } @@ -1090,10 +2114,70 @@ impl WalletWriteRepository for InMemoryWalletRepository { &self, input: CreateWalletRefundRequestInput, ) -> Result { - let wallets = self.wallets_by_id.read().expect("wallet repo lock"); - let Some(wallet) = wallets.get(&input.wallet_id) else { - return Ok(CreateWalletRefundRequestOutcome::WalletMissing); + if !input.amount_usd.is_finite() || input.amount_usd <= 0.0 { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if input + .idempotency_key + .as_deref() + .is_some_and(|key| key.trim().is_empty() || key.chars().count() > 128) + { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "refund idempotency key is invalid".to_string(), + )); + } + // SQL backends lock the wallet row while reserving a refund. Serialize + // the in-memory equivalent so concurrent requests cannot both spend + // the same available balance or race compensation cleanup. Keep both + // locks in this order everywhere: lifecycle first, reservation second. + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); + let _reservation_guard = self + .refund_creation_lock + .lock() + .expect("wallet refund creation lock"); + let wallet = { + let wallets = self.wallets_by_id.read().expect("wallet repo lock"); + let Some(wallet) = wallets.get(&input.wallet_id) else { + return Ok(CreateWalletRefundRequestOutcome::WalletMissing); + }; + if wallet.user_id.as_deref() != Some(input.user_id.as_str()) { + return Ok(CreateWalletRefundRequestOutcome::WalletMissing); + } + wallet.clone() }; + if !wallet.balance.is_finite() { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet recharge balance is invalid".to_string(), + )); + } + + if let Some(idempotency_key) = input.idempotency_key.as_deref() { + let key = (input.user_id.clone(), idempotency_key.to_string()); + if let Some(refund_id) = self + .refund_idempotency_to_id + .read() + .expect("wallet repo lock") + .get(&key) + .cloned() + { + if let Some(refund) = self + .refunds_by_id + .read() + .expect("wallet repo lock") + .get(&refund_id) + .cloned() + { + return Ok(CreateWalletRefundRequestOutcome::Duplicate(refund)); + } + return Ok(CreateWalletRefundRequestOutcome::DuplicateRejected); + } + } + let reserved_amount = self .refunds_by_id .read() @@ -1103,20 +2187,36 @@ impl WalletWriteRepository for InMemoryWalletRepository { refund.wallet_id == input.wallet_id && matches!(refund.status.as_str(), "pending_approval" | "approved") }) - .map(|refund| refund.amount_usd) - .sum::(); - if input.amount_usd > (wallet.balance - reserved_amount) { + .try_fold(0.0_f64, |total, refund| { + let amount = refund.amount_usd; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + let Some(reserved_amount) = reserved_amount else { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "wallet refund reservation is invalid".to_string(), + )); + }; + let available_balance = wallet.balance - reserved_amount; + if !available_balance.is_finite() || input.amount_usd > available_balance { return Ok(CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableBalance); } + let mut resolved_payment_method: Option = None; if let Some(order_id) = input.payment_order_id.as_deref() { - let orders = self.payment_orders_by_id.read().expect("wallet repo lock"); - let Some(order) = orders.get(order_id) else { - return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotFound); + let order = { + let orders = self.payment_orders_by_id.read().expect("wallet repo lock"); + let Some(order) = orders.get(order_id) else { + return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotFound); + }; + if order.wallet_id != input.wallet_id || order.status != "credited" { + return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotRefundable); + } + order.clone() }; - if order.wallet_id != input.wallet_id || order.status != "credited" { - return Ok(CreateWalletRefundRequestOutcome::PaymentOrderNotRefundable); - } let reserved_for_order = self .refunds_by_id .read() @@ -1126,38 +2226,64 @@ impl WalletWriteRepository for InMemoryWalletRepository { refund.payment_order_id.as_deref() == Some(order_id) && matches!(refund.status.as_str(), "pending_approval" | "approved") }) - .map(|refund| refund.amount_usd) - .sum::(); - if input.amount_usd > (order.refundable_amount_usd - reserved_for_order) { + .try_fold(0.0_f64, |total, refund| { + let amount = refund.amount_usd; + if !amount.is_finite() || amount <= 0.0 { + return None; + } + let next = total + amount; + next.is_finite().then_some(next) + }); + if !payment_order_refund_amounts_are_consistent( + order.amount_usd, + order.refunded_amount_usd, + order.refundable_amount_usd, + ) { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund amounts are invalid".to_string(), + )); + } + let Some(reserved_for_order) = reserved_for_order else { + return Ok(CreateWalletRefundRequestOutcome::InvalidInput( + "payment order refund reservation is invalid".to_string(), + )); + }; + let available_order_amount = order.refundable_amount_usd - reserved_for_order; + if !available_order_amount.is_finite() || input.amount_usd > available_order_amount { return Ok( CreateWalletRefundRequestOutcome::RefundAmountExceedsAvailableOrderAmount, ); } + resolved_payment_method = Some(order.payment_method.clone()); } + let canonical = canonicalize_wallet_refund_fields( + input.payment_order_id.as_deref(), + input.source_type.as_deref(), + input.source_id.as_deref(), + input.refund_mode.as_deref(), + resolved_payment_method.as_deref(), + ) + .map_err(DataLayerError::InvalidInput)?; + let idempotency_key = input.idempotency_key.clone(); + let user_id = input.user_id.clone(); + let refund = StoredAdminWalletRefund { id: format!("refund-{}", uuid::Uuid::new_v4()), refund_no: input.refund_no, wallet_id: input.wallet_id, user_id: Some(input.user_id), payment_order_id: input.payment_order_id.clone(), - source_type: input - .payment_order_id - .clone() - .map(|_| "payment_order".to_string()) - .or(input.source_type) - .unwrap_or_else(|| "wallet_balance".to_string()), - source_id: input.payment_order_id.clone().or(input.source_id), - refund_mode: input - .refund_mode - .unwrap_or_else(|| "offline_payout".to_string()), + source_type: canonical.source_type, + source_id: canonical.source_id, + refund_mode: canonical.refund_mode, amount_usd: input.amount_usd, status: "pending_approval".to_string(), reason: input.reason, failure_reason: None, gateway_refund_id: None, payout_method: None, - payout_reference: input.idempotency_key, + payout_reference: None, payout_proof: None, requested_by: None, approved_by: None, @@ -1171,13 +2297,22 @@ impl WalletWriteRepository for InMemoryWalletRepository { .write() .expect("wallet repo lock") .insert(refund.id.clone(), refund.clone()); + if let Some(idempotency_key) = idempotency_key { + self.refund_idempotency_to_id + .write() + .expect("wallet repo lock") + .insert((user_id, idempotency_key), refund.id.clone()); + } Ok(CreateWalletRefundRequestOutcome::Created(refund)) } async fn process_payment_callback( &self, - _input: ProcessPaymentCallbackInput, + mut input: ProcessPaymentCallbackInput, ) -> Result { + input + .canonicalize_and_validate() + .map_err(DataLayerError::InvalidInput)?; Ok(ProcessPaymentCallbackOutcome::Failed { duplicate: false, error: "payment callback is not supported in memory wallet repository".to_string(), @@ -1213,6 +2348,68 @@ impl WalletWriteRepository for InMemoryWalletRepository { Ok(WalletMutationOutcome::NotFound) } + async fn update_admin_wallet_refund_gateway( + &self, + input: UpdateAdminWalletRefundGatewayInput, + ) -> Result, DataLayerError> { + if input.gateway_refund_id.trim().is_empty() || input.gateway_refund_id.len() > 128 { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier is invalid".to_string(), + )); + } + if input + .payout_proof + .as_ref() + .is_some_and(|proof| !proof.is_object()) + { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund proof must be an object".to_string(), + )); + } + let mut refunds = self.refunds_by_id.write().expect("wallet repo lock"); + let Some(refund) = refunds.get_mut(&input.refund_id) else { + return Ok(WalletMutationOutcome::NotFound); + }; + if refund.wallet_id != input.wallet_id { + return Ok(WalletMutationOutcome::NotFound); + } + if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 { + return Ok(WalletMutationOutcome::Invalid( + "refund amount must be finite and greater than zero".to_string(), + )); + } + if let Some(existing_id) = refund.gateway_refund_id.as_deref() { + if existing_id != input.gateway_refund_id { + return Ok(WalletMutationOutcome::Invalid( + "gateway refund identifier conflicts with existing evidence".to_string(), + )); + } + } + if refund.status == "succeeded" { + return Ok(WalletMutationOutcome::Applied(refund.clone())); + } + if refund.status != "processing" { + return Ok(WalletMutationOutcome::Invalid( + "refund status must be processing before gateway update".to_string(), + )); + } + if refund.gateway_refund_id.is_none() { + refund.gateway_refund_id = Some(input.gateway_refund_id); + } + // Preserve a processing proof for ordinary replays, but allow an + // explicit successful gateway proof to upgrade it. + if refund.payout_proof.is_none() + || input + .payout_proof + .as_ref() + .is_some_and(wallet_refund_proof_is_success) + { + refund.payout_proof = input.payout_proof; + } + refund.updated_at_unix_secs = current_unix_secs(); + Ok(WalletMutationOutcome::Applied(refund.clone())) + } + async fn complete_admin_wallet_refund( &self, _input: CompleteAdminWalletRefundInput, @@ -1259,6 +2456,7 @@ impl WalletWriteRepository for InMemoryWalletRepository { &self, input: CreateAdminRedeemCodeBatchInput, ) -> Result { + validate_admin_redeem_code_batch_input(&input).map_err(DataLayerError::InvalidInput)?; let now_ms = current_unix_ms(); let now_secs = current_unix_secs(); let batch_id = format!("redeem-batch-{}", uuid::Uuid::new_v4()); @@ -1480,6 +2678,10 @@ impl WalletWriteRepository for InMemoryWalletRepository { &self, input: RedeemWalletCodeInput, ) -> Result { + let _lifecycle_guard = self + .wallet_lifecycle_lock + .lock() + .expect("wallet lifecycle lock"); let Some(normalized) = normalize_redeem_code(&input.code) else { return Ok(RedeemWalletCodeOutcome::InvalidCode); }; @@ -1532,12 +2734,10 @@ impl WalletWriteRepository for InMemoryWalletRepository { batch.amount_usd, ) }; - let credits_recharge_balance = redeem_code_credits_recharge_balance(&balance_bucket); - let (wallet, balance_before, gift_before) = { - let mut wallets = self.wallets_by_id.write().expect("wallet repo lock"); + let wallets = self.wallets_by_id.read().expect("wallet repo lock"); if let Some(wallet) = wallets - .values_mut() + .values() .find(|wallet| wallet.user_id.as_deref() == Some(input.user_id.as_str())) { if wallet.status != "active" { @@ -1545,39 +2745,40 @@ impl WalletWriteRepository for InMemoryWalletRepository { } let balance_before = wallet.balance; let gift_before = wallet.gift_balance; - if credits_recharge_balance { - wallet.balance += amount_usd; - } else { - wallet.gift_balance += amount_usd; - } - wallet.total_recharged += amount_usd; + let (after_recharge, after_gift, after_total_recharged) = + validate_redeem_wallet_credit( + &balance_bucket, + amount_usd, + balance_before, + gift_before, + wallet.total_recharged, + ) + .map_err(DataLayerError::UnexpectedValue)?; + let mut wallet = wallet.clone(); + wallet.balance = after_recharge; + wallet.gift_balance = after_gift; + wallet.total_recharged = after_total_recharged; wallet.updated_at_unix_secs = now_secs; - (wallet.clone(), balance_before, gift_before) + (wallet, balance_before, gift_before) } else { + let (after_recharge, after_gift, after_total_recharged) = + validate_redeem_wallet_credit(&balance_bucket, amount_usd, 0.0, 0.0, 0.0) + .map_err(DataLayerError::UnexpectedValue)?; let wallet = StoredWalletSnapshot::new( format!("wallet-{}", uuid::Uuid::new_v4()), Some(input.user_id.clone()), None, - if credits_recharge_balance { - amount_usd - } else { - 0.0 - }, - if credits_recharge_balance { - 0.0 - } else { - amount_usd - }, + after_recharge, + after_gift, "finite".to_string(), "USD".to_string(), "active".to_string(), - amount_usd, + after_total_recharged, 0.0, 0.0, 0.0, now_secs as i64, )?; - wallets.insert(wallet.id.clone(), wallet.clone()); (wallet, 0.0, 0.0) } }; @@ -1594,6 +2795,8 @@ impl WalletWriteRepository for InMemoryWalletRepository { refunded_amount_usd: 0.0, refundable_amount_usd: redeem_code_refundable_amount(&balance_bucket, amount_usd), payment_method: redeem_code_payment_method(&balance_bucket).to_string(), + payment_provider: Some("redeem_code".to_string()), + order_kind: "wallet_recharge".to_string(), gateway_order_id: Some(format!("card_{}", uuid::Uuid::new_v4().simple())), gateway_response: Some(serde_json::json!({ "source": "redeem_code", @@ -1607,10 +2810,14 @@ impl WalletWriteRepository for InMemoryWalletRepository { credited_at_unix_secs: Some(now_secs), expires_at_unix_secs: None, }; - self.payment_orders_by_id + // Reserve the globally unique payment identity before publishing any + // wallet mutation. Every operation after this point is infallible map + // replacement, so a duplicate order cannot leave credited funds. + self.insert_payment_order_unique(order.clone())?; + self.wallets_by_id .write() .expect("wallet repo lock") - .insert(order.id.clone(), order.clone()); + .insert(wallet.id.clone(), wallet.clone()); let tx = StoredAdminWalletTransaction { id: format!("wallet-tx-{}", uuid::Uuid::new_v4()), @@ -1657,7 +2864,7 @@ impl WalletWriteRepository for InMemoryWalletRepository { .expect("wallet repo lock") .get_mut(&batch_id) { - batch.redeemed_count += 1; + batch.redeemed_count = batch.redeemed_count.saturating_add(1); batch.active_count = batch.active_count.saturating_sub(1); batch.updated_at_unix_secs = now_secs; } @@ -1675,11 +2882,17 @@ impl WalletWriteRepository for InMemoryWalletRepository { mod tests { use super::{InMemoryWalletRepository, WalletReadSeed}; use crate::repository::wallet::{ - AdminWalletListQuery, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - StoredAdminPaymentOrder, StoredAdminWalletRefund, StoredWalletSnapshot, WalletLookupKey, - WalletReadRepository, WalletWriteRepository, + AdminWalletListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, + CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, + CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, + FailWalletRechargeCheckoutInput, StoredAdminPaymentOrder, StoredAdminWalletRefund, + StoredWalletSnapshot, UpdateWalletRechargeCheckoutInput, WalletLookupKey, + WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; + use crate::DataLayerError; use serde_json::json; + use std::sync::Arc; fn sample_wallet() -> StoredWalletSnapshot { StoredWalletSnapshot::new( @@ -1700,6 +2913,197 @@ mod tests { .expect("wallet should build") } + #[tokio::test] + async fn compensation_delete_preserves_wallet_with_persisted_balance_in_memory() { + let wallet = StoredWalletSnapshot::new( + "funded-wallet".to_string(), + Some("funded-user".to_string()), + None, + 1.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 1.0, + 0.0, + 0.0, + 1.0, + 100, + ) + .expect("wallet should build"); + let repository = InMemoryWalletRepository::seed(vec![wallet]); + + assert!(!repository + .delete_wallet_if_unreferenced("funded-wallet", WalletLookupKey::UserId("funded-user")) + .await + .expect("funded wallet must not be deleted")); + assert!(repository + .find(WalletLookupKey::UserId("funded-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn provisional_recharge_cleanup_preserves_wallet_changed_after_initial_check_in_memory() { + let wallet = StoredWalletSnapshot::new( + "provisional-recharge-wallet".to_string(), + Some("provisional-recharge-user".to_string()), + None, + 0.0, + 0.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + 100, + ) + .expect("wallet should build"); + let repository = InMemoryWalletRepository::seed(vec![wallet]); + repository.with_wallets_mut(|wallets| { + let wallet = wallets + .get_mut("provisional-recharge-wallet") + .expect("wallet should be seeded"); + wallet.balance = 2.0; + wallet.total_recharged = 2.0; + wallet.updated_at_unix_secs = 200; + }); + + // This is the same cleanup helper used when a recharge order loses a + // uniqueness race after creating a provisional wallet. A wallet that + // acquired funds must survive the compensation attempt. + repository.remove_created_wallet_if_unreferenced("provisional-recharge-wallet"); + let retained = repository + .find(WalletLookupKey::UserId("provisional-recharge-user")) + .await + .expect("wallet lookup should succeed") + .expect("changed wallet should remain"); + assert_eq!(retained.balance, 2.0); + assert_eq!(retained.total_recharged, 2.0); + } + + #[tokio::test] + async fn snapshot_compensation_deletes_matching_funded_wallet_and_preserves_changed_one() { + let wallet = StoredWalletSnapshot::new( + "import-funded-wallet".to_string(), + Some("import-funded-user".to_string()), + None, + 12.5, + 3.25, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 12.5, + 0.0, + 0.0, + 0.0, + 4242, + ) + .expect("wallet should build"); + let repository = InMemoryWalletRepository::seed(vec![wallet.clone()]); + assert!(repository + .delete_wallet_if_snapshot_matches_and_unreferenced( + &wallet, + WalletLookupKey::UserId("import-funded-user"), + ) + .await + .expect("matching snapshot should delete")); + assert!(repository + .find(WalletLookupKey::WalletId("import-funded-wallet")) + .await + .expect("wallet lookup should succeed") + .is_none()); + + let repository = InMemoryWalletRepository::seed(vec![wallet.clone()]); + repository.with_wallets_mut(|wallets| { + wallets + .get_mut("import-funded-wallet") + .expect("wallet should be seeded") + .balance = 99.0; + }); + assert!(!repository + .delete_wallet_if_snapshot_matches_and_unreferenced( + &wallet, + WalletLookupKey::UserId("import-funded-user"), + ) + .await + .expect("changed snapshot should be retained")); + assert!(repository + .find(WalletLookupKey::WalletId("import-funded-wallet")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + + #[tokio::test] + async fn snapshot_restore_is_compare_and_swap_in_memory() { + let before = StoredWalletSnapshot::new( + "existing-wallet".to_string(), + Some("existing-user".to_string()), + None, + 4.0, + 2.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 7.0, + 3.0, + 0.0, + 0.0, + 100, + ) + .expect("wallet should build"); + let mut after = before.clone(); + after.balance = 25.0; + after.gift_balance = 5.0; + after.total_recharged = 28.0; + after.updated_at_unix_secs = 200; + + let repository = InMemoryWalletRepository::seed(vec![after.clone()]); + assert!(repository + .restore_wallet_if_snapshot_matches( + &before, + &after, + WalletLookupKey::UserId("existing-user"), + ) + .await + .expect("matching post-state should restore")); + assert_eq!( + repository + .find(WalletLookupKey::UserId("existing-user")) + .await + .expect("wallet lookup should succeed"), + Some(before.clone()) + ); + + let repository = InMemoryWalletRepository::seed(vec![after.clone()]); + repository.with_wallets_mut(|wallets| { + let wallet = wallets + .get_mut("existing-wallet") + .expect("wallet should be seeded"); + wallet.balance = 99.0; + wallet.updated_at_unix_secs = 300; + }); + assert!(!repository + .restore_wallet_if_snapshot_matches( + &before, + &after, + WalletLookupKey::UserId("existing-user"), + ) + .await + .expect("changed post-state should fail closed")); + let retained = repository + .find(WalletLookupKey::UserId("existing-user")) + .await + .expect("wallet lookup should succeed") + .expect("changed wallet should remain"); + assert_eq!(retained.balance, 99.0); + assert_eq!(retained.updated_at_unix_secs, 300); + } + #[tokio::test] async fn updates_auth_wallet_limit_mode_and_snapshot_in_memory() { let repository = InMemoryWalletRepository::seed(vec![sample_wallet()]); @@ -1778,6 +3182,8 @@ mod tests { refunded_amount_usd: 0.0, refundable_amount_usd: 10.0, payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + order_kind: "wallet_recharge".to_string(), gateway_order_id: None, gateway_response: None, status: status.to_string(), @@ -1788,6 +3194,151 @@ mod tests { } } + fn stripe_secret_cas_input( + order: &StoredAdminPaymentOrder, + expected_gateway_response: serde_json::Value, + expected_ciphertext: &str, + replacement_ciphertext: &str, + ) -> CompareAndSwapPaymentOrderStripeClientSecretInput { + CompareAndSwapPaymentOrderStripeClientSecretInput { + order_id: order.id.clone(), + order_no: order.order_no.clone(), + wallet_id: order.wallet_id.clone(), + user_id: order.user_id.clone(), + payment_method: order.payment_method.clone(), + payment_provider: order.payment_provider.clone(), + order_kind: order.order_kind.clone(), + gateway_order_id: order.gateway_order_id.clone(), + expected_status: order.status.clone(), + expected_expires_at_unix_secs: order.expires_at_unix_secs, + expected_gateway_response, + expected_client_secret_encrypted: expected_ciphertext.to_string(), + replacement_client_secret_encrypted: replacement_ciphertext.to_string(), + } + } + + #[tokio::test] + async fn stripe_secret_cas_is_exact_and_never_overwrites_a_newer_value_in_memory() { + let legacy = "gAAAAABlegacy"; + let replacement = concat!( + "aether-payment-order-stripe-client-secret-v2:", + "aether-runtime-secret-v1:gAAAAABreplacement" + ); + let mut order = sample_payment_order("stripe-cas-order", Some("user-1"), "pending"); + order.gateway_order_id = Some("pi-cas".to_string()); + order.expires_at_unix_secs = Some(4_102_444_800); + order.gateway_response = Some(json!({ + "gateway": "stripe", + "publishable_key": "pk_test_public", + "_stripe_client_secret_encrypted": legacy, + })); + let observed = order + .gateway_response + .clone() + .expect("fixture response should exist"); + let repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + payment_orders: vec![order.clone()], + ..WalletReadSeed::default() + }); + let input = stripe_secret_cas_input(&order, observed.clone(), legacy, replacement); + + let mut stale_json = input.clone(); + stale_json.expected_gateway_response["publishable_key"] = json!("pk_test_changed"); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(stale_json) + .await + .expect("stale JSON should be a normal CAS miss")); + + let mut stale_ciphertext = input.clone(); + stale_ciphertext.expected_client_secret_encrypted = "gAAAAABother".to_string(); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(stale_ciphertext) + .await + .expect("stale ciphertext should be a normal CAS miss")); + + let mut foreign_identity = input.clone(); + foreign_identity.order_no = "order-no-foreign".to_string(); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(foreign_identity) + .await + .expect("foreign identity should be a normal CAS miss")); + + assert!(repository + .compare_and_swap_payment_order_stripe_client_secret(input.clone()) + .await + .expect("exact CAS should succeed")); + assert!(!repository + .compare_and_swap_payment_order_stripe_client_secret(input) + .await + .expect("an old migration must not overwrite the new value")); + + let stored = repository + .find_admin_payment_order(&order.id) + .await + .expect("stored order should be readable") + .expect("stored order should remain"); + let response = stored + .gateway_response + .expect("stored response should remain"); + assert_eq!( + response["_stripe_client_secret_encrypted"].as_str(), + Some(replacement) + ); + assert_eq!(response["publishable_key"], "pk_test_public"); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn concurrent_stripe_secret_migrations_have_exactly_one_winner_in_memory() { + let legacy = "gAAAAABlegacy-race"; + let mut order = sample_payment_order("stripe-cas-race", Some("user-1"), "pending"); + order.expires_at_unix_secs = Some(4_102_444_800); + order.gateway_response = Some(json!({ + "gateway": "stripe", + "_stripe_client_secret_encrypted": legacy, + })); + let observed = order + .gateway_response + .clone() + .expect("fixture response should exist"); + let repository = Arc::new(InMemoryWalletRepository::seed_read_model(WalletReadSeed { + payment_orders: vec![order.clone()], + ..WalletReadSeed::default() + })); + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let mut tasks = Vec::new(); + for suffix in ["winner-a", "winner-b"] { + let repository = Arc::clone(&repository); + let barrier = Arc::clone(&barrier); + let input = stripe_secret_cas_input( + &order, + observed.clone(), + legacy, + &format!( + "aether-payment-order-stripe-client-secret-v2:aether-runtime-secret-v1:gAAAAAB{suffix}" + ), + ); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + repository + .compare_and_swap_payment_order_stripe_client_secret(input) + .await + })); + } + barrier.wait().await; + + let mut winners = 0; + for task in tasks { + if task + .await + .expect("migration task should join") + .expect("migration should not error") + { + winners += 1; + } + } + assert_eq!(winners, 1); + } + fn sample_refund(id: &str, user_id: Option<&str>, status: &str) -> StoredAdminWalletRefund { StoredAdminWalletRefund { id: id.to_string(), @@ -1898,6 +3449,57 @@ mod tests { assert!(history.items.is_empty()); } + #[tokio::test] + async fn deletes_only_untouched_provisional_auth_wallets_in_memory() { + let repository = InMemoryWalletRepository::seed(Vec::new()); + let provisional_wallet = repository + .initialize_auth_user_wallet("provisional-user", 10.0, false) + .await + .expect("wallet initialization should succeed") + .expect("provisional wallet should exist"); + + assert!(repository + .delete_provisional_auth_user_wallet(&provisional_wallet.id, "provisional-user") + .await + .expect("provisional cleanup should succeed")); + assert!(repository + .find(WalletLookupKey::UserId("provisional-user")) + .await + .expect("wallet lookup should succeed") + .is_none()); + } + + #[tokio::test] + async fn provisional_cleanup_keeps_wallet_with_financial_activity_in_memory() { + let wallet = StoredWalletSnapshot::new( + "active-wallet".to_string(), + Some("active-user".to_string()), + None, + 0.0, + 10.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 1.0, + 0.0, + 10.0, + 1, + ) + .expect("wallet should build"); + let repository = InMemoryWalletRepository::seed([wallet]); + + assert!(!repository + .delete_provisional_auth_user_wallet("active-wallet", "active-user") + .await + .expect("provisional cleanup should succeed")); + assert!(repository + .find(WalletLookupKey::UserId("active-user")) + .await + .expect("wallet lookup should succeed") + .is_some()); + } + #[tokio::test] async fn lifetime_plan_purchase_blocks_duplicate_pending_order_in_memory() { let repository = InMemoryWalletRepository::seed(vec![sample_wallet()]); @@ -2002,6 +3604,607 @@ mod tests { } } + #[tokio::test] + async fn plan_purchase_rejects_malformed_wallet_credit_in_memory() { + let repository = InMemoryWalletRepository::seed(vec![sample_wallet()]); + let result = repository + .create_plan_purchase_order(CreatePlanPurchaseOrderInput { + preferred_wallet_id: None, + user_id: "user-1".to_string(), + amount_usd: 1.0, + pay_amount: 1.0, + pay_currency: "USD".to_string(), + exchange_rate: 1.0, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "gateway-invalid-wallet-credit".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-invalid-wallet-credit".to_string(), + product_id: "invalid-wallet-credit-plan".to_string(), + product_snapshot: json!({ + "id": "invalid-wallet-credit-plan", + "purchase_limit_scope": "unlimited", + "entitlements": [{ + "type": "wallet_credit", + "amount_usd": 1.0, + "balance_bucket": "unknown" + }] + }), + expires_at_unix_secs: 4_102_444_800, + }) + .await; + assert!(matches!(result, Err(DataLayerError::InvalidInput(_)))); + assert!(repository + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .is_empty()); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn gateway_order_uniqueness_is_atomic_in_memory() { + let repository = Arc::new(InMemoryWalletRepository::seed(Vec::new())); + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let mut tasks = Vec::new(); + for index in 0..2 { + let repository = Arc::clone(&repository); + let barrier = Arc::clone(&barrier); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some(format!("wallet-concurrent-{index}")), + user_id: format!("user-concurrent-{index}"), + amount_usd: 1.0, + pay_amount: Some(1.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: " EPAY ".to_string(), + payment_provider: None, + payment_channel: None, + gateway_order_id: "shared-memory-gateway-id".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: format!("order-concurrent-{index}"), + expires_at_unix_secs: 4_102_444_800, + }) + .await + })); + } + barrier.wait().await; + + let mut created = 0; + let mut rejected = 0; + for task in tasks { + match task.await.expect("order task should join") { + Ok(CreateWalletRechargeOrderOutcome::Created(order)) => { + assert_eq!(order.payment_method, "epay"); + created += 1; + } + Err(DataLayerError::InvalidInput(_)) => rejected += 1, + other => panic!("unexpected concurrent order result: {other:?}"), + } + } + assert_eq!((created, rejected), (1, 1)); + assert_eq!( + repository + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .len(), + 1 + ); + assert_eq!( + repository + .wallets_by_id + .read() + .expect("wallet repo lock") + .len(), + 1, + "the rejected order must not leave a provisional wallet behind" + ); + let orders = repository + .payment_orders_by_id + .read() + .expect("wallet repo lock"); + let wallets = repository.wallets_by_id.read().expect("wallet repo lock"); + let stored_order = orders.values().next().expect("winning order should remain"); + let stored_wallet = wallets + .get(&stored_order.wallet_id) + .expect("winning order must not reference a removed wallet"); + assert_eq!(stored_wallet.user_id, stored_order.user_id); + } + + #[tokio::test] + async fn recharge_order_conflict_removes_unreferenced_provisional_wallet_in_memory() { + let mut existing = sample_payment_order( + "existing-recharge-order", + Some("recharge-conflict-user"), + "pending", + ); + existing.order_no = "recharge-conflict-order-no".to_string(); + existing.wallet_id = "wallet-from-old-read-model".to_string(); + existing.payment_method = "epay".to_string(); + existing.gateway_order_id = Some("gateway-from-old-read-model".to_string()); + existing.gateway_response = Some(json!({ + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending" + })); + let repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + payment_orders: vec![existing.clone()], + ..WalletReadSeed::default() + }); + + let outcome = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-provisional-conflict".to_string()), + user_id: "recharge-conflict-user".to_string(), + amount_usd: 10.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: None, + gateway_order_id: "gateway-retry".to_string(), + gateway_response: json!({ "integration_status": "checkout_pending" }), + order_no: "recharge-conflict-order-no".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("existing recharge order should be returned"); + + assert!(matches!( + outcome, + CreateWalletRechargeOrderOutcome::Existing(order) + if order.id == "existing-recharge-order" + )); + assert!(repository + .find(WalletLookupKey::UserId("recharge-conflict-user")) + .await + .expect("wallet lookup should succeed") + .is_none()); + assert_eq!( + repository + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .len(), + 1 + ); + } + + #[tokio::test] + async fn recharge_rejects_preferred_wallet_id_owned_by_another_user_in_memory() { + let repository = InMemoryWalletRepository::seed(vec![sample_wallet()]); + + let result = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: Some("wallet-1".to_string()), + user_id: "different-owner".to_string(), + amount_usd: 5.0, + pay_amount: Some(5.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: "gateway-wallet-id-collision".to_string(), + gateway_response: json!({ "checkout": true }), + order_no: "order-wallet-id-collision".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await; + + assert!(matches!( + result, + Err(DataLayerError::InvalidInput(message)) + if message.contains("wallet identifier already belongs") + )); + let original = repository + .find(WalletLookupKey::UserId("user-1")) + .await + .expect("original wallet lookup should succeed") + .expect("original wallet should remain present"); + assert_eq!(original.id, "wallet-1"); + assert!(repository + .find(WalletLookupKey::UserId("different-owner")) + .await + .expect("new owner lookup should succeed") + .is_none()); + assert!(repository + .payment_orders_by_id + .read() + .expect("wallet repo lock") + .is_empty()); + } + + #[tokio::test] + async fn recharge_checkout_update_preserves_order_kind_in_memory() { + let repository = InMemoryWalletRepository::seed(Vec::new()); + let created = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: "user-recharge-kind".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "placeholder-order-kind".to_string(), + gateway_response: json!({ + "gateway": "epay", + "integration_status": "checkout_pending" + }), + order_no: "order-recharge-kind".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("recharge order should be created"); + let CreateWalletRechargeOrderOutcome::Created(order) = created else { + panic!("expected a newly created recharge order"); + }; + + let updated = repository + .update_wallet_recharge_checkout(UpdateWalletRechargeCheckoutInput { + order_id: order.id.clone(), + gateway_order_id: "provider-order-kind".to_string(), + gateway_response: json!({ + "gateway": "epay", + "payment_url": "https://pay.example.test/order" + }), + }) + .await + .expect("checkout update should succeed"); + assert!(matches!(updated, WalletMutationOutcome::Applied(_))); + + let replay = repository + .find_wallet_recharge_order_by_order_no("user-recharge-kind", "order-recharge-kind") + .await + .expect("recharge lookup should succeed") + .expect("updated order should remain discoverable"); + assert_eq!( + replay.gateway_order_id.as_deref(), + Some("provider-order-kind") + ); + assert_eq!( + replay + .gateway_response + .as_ref() + .and_then(|value| value.get("order_kind")) + .and_then(serde_json::Value::as_str), + Some("wallet_recharge") + ); + } + + #[tokio::test] + async fn recharge_checkout_update_rejects_expired_order_in_memory() { + let repository = InMemoryWalletRepository::seed(Vec::new()); + let created = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: "user-expired-checkout".to_string(), + amount_usd: 3.0, + pay_amount: Some(3.0), + pay_currency: Some("USD".to_string()), + exchange_rate: Some(1.0), + payment_method: "epay".to_string(), + payment_provider: Some("epay".to_string()), + payment_channel: Some("alipay".to_string()), + gateway_order_id: "order-expired-checkout".to_string(), + gateway_response: serde_json::json!({ + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending" + }), + order_no: "order-expired-checkout".to_string(), + expires_at_unix_secs: 1, + }) + .await + .expect("expired recharge order should be creatable for regression setup"); + let CreateWalletRechargeOrderOutcome::Created(order) = created else { + panic!("expected a newly created recharge order"); + }; + + let result = repository + .update_wallet_recharge_checkout(UpdateWalletRechargeCheckoutInput { + order_id: order.id.clone(), + gateway_order_id: "provider-expired-checkout".to_string(), + gateway_response: serde_json::json!({ + "order_kind": "wallet_recharge", + "payment_url": "https://pay.example.test/expired" + }), + }) + .await + .expect("expired checkout update should resolve"); + assert!(matches!(result, WalletMutationOutcome::Invalid(_))); + + let replay = repository + .find_wallet_recharge_order_by_order_no( + "user-expired-checkout", + "order-expired-checkout", + ) + .await + .expect("expired recharge lookup should succeed") + .expect("expired recharge order should remain stored"); + assert_eq!( + replay.gateway_order_id.as_deref(), + Some("order-expired-checkout") + ); + } + + #[tokio::test] + async fn failed_recharge_without_channel_can_be_reclaimed_in_memory() { + let repository = InMemoryWalletRepository::seed(Vec::new()); + let first = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: "user-reclaim-no-channel".to_string(), + amount_usd: 3.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: None, + gateway_order_id: "order-reclaim-no-channel".to_string(), + gateway_response: json!({ + "gateway": "stripe", + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending", + "checkout_claim_token": "first-claim" + }), + order_no: "order-reclaim-no-channel".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("initial recharge should be created"); + let CreateWalletRechargeOrderOutcome::Created(first) = first else { + panic!("expected initial recharge order"); + }; + + let failed = repository + .fail_wallet_recharge_checkout(FailWalletRechargeCheckoutInput { + order_id: first.id.clone(), + claim_token: "first-claim".to_string(), + reason: "provider unavailable".to_string(), + provider_request_may_have_succeeded: false, + }) + .await + .expect("checkout failure should resolve"); + assert!(matches!(failed, WalletMutationOutcome::Applied(_))); + + let retry = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: "user-reclaim-no-channel".to_string(), + amount_usd: 3.0, + pay_amount: None, + pay_currency: None, + exchange_rate: None, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: None, + gateway_order_id: "order-reclaim-no-channel".to_string(), + gateway_response: json!({ + "gateway": "stripe", + "order_kind": "wallet_recharge", + "integration_status": "checkout_pending", + "checkout_claim_token": "retry-claim" + }), + order_no: "order-reclaim-no-channel".to_string(), + expires_at_unix_secs: 4_102_444_800, + }) + .await + .expect("failed placeholder should be reclaimable"); + + let CreateWalletRechargeOrderOutcome::Created(reclaimed) = retry else { + panic!("expected failed placeholder to be reclaimed"); + }; + assert_eq!(reclaimed.id, first.id); + assert_eq!( + reclaimed + .gateway_response + .as_ref() + .and_then(|value| value.get("checkout_claim_token")) + .and_then(serde_json::Value::as_str), + Some("retry-claim") + ); + assert_eq!(reclaimed.status, "pending"); + } + + #[tokio::test] + async fn recharge_order_rejects_non_finite_numeric_fields_in_memory() { + let invalid_inputs = [ + (f64::NAN, Some(1.0), Some(1.0), 4_102_444_800), + (1.0, Some(f64::INFINITY), Some(1.0), 4_102_444_800), + (1.0, Some(1.0), Some(0.0), 4_102_444_800), + (1.0, Some(1.0), Some(1.0), i64::MAX as u64 + 1), + ]; + for (index, (amount_usd, pay_amount, exchange_rate, expires_at)) in + invalid_inputs.into_iter().enumerate() + { + let repository = InMemoryWalletRepository::seed(Vec::new()); + let result = repository + .create_wallet_recharge_order(CreateWalletRechargeOrderInput { + preferred_wallet_id: None, + user_id: format!("invalid-recharge-user-{index}"), + amount_usd, + pay_amount, + pay_currency: Some("USD".to_string()), + exchange_rate, + payment_method: "stripe".to_string(), + payment_provider: Some("stripe".to_string()), + payment_channel: Some("card".to_string()), + gateway_order_id: format!("invalid-recharge-gateway-{index}"), + gateway_response: json!({"payment_url": "https://pay.example.test"}), + order_no: format!("invalid-recharge-order-{index}"), + expires_at_unix_secs: expires_at, + }) + .await; + assert!(matches!(result, Err(DataLayerError::InvalidInput(_)))); + assert!(repository + .find(WalletLookupKey::UserId(&format!( + "invalid-recharge-user-{index}" + ))) + .await + .expect("wallet lookup should succeed") + .is_none()); + } + } + + #[tokio::test] + async fn refund_creation_is_idempotent_and_rejects_corrupt_wallet_balance_in_memory() { + let repository = InMemoryWalletRepository::seed(vec![sample_wallet()]); + let first = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 4.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("first request".to_string()), + idempotency_key: Some("memory-refund-idempotency".to_string()), + refund_no: "memory-refund-1".to_string(), + }) + .await + .expect("first refund should resolve"); + let CreateWalletRefundRequestOutcome::Created(first) = first else { + panic!("expected first refund to be created"); + }; + let replay = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 9.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("replayed request".to_string()), + idempotency_key: Some("memory-refund-idempotency".to_string()), + refund_no: "memory-refund-2".to_string(), + }) + .await + .expect("refund replay should resolve"); + assert!(matches!( + replay, + CreateWalletRefundRequestOutcome::Duplicate(refund) if refund.id == first.id + )); + assert_eq!( + repository + .refunds_by_id + .read() + .expect("wallet repo lock") + .len(), + 1 + ); + + repository.with_wallets_mut(|wallets| { + wallets + .get_mut("wallet-1") + .expect("sample wallet should exist") + .balance = f64::NAN; + }); + let corrupt = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: None, + idempotency_key: Some("memory-refund-corrupt".to_string()), + refund_no: "memory-refund-corrupt".to_string(), + }) + .await + .expect("corrupt wallet refund should resolve"); + assert!(matches!( + corrupt, + CreateWalletRefundRequestOutcome::InvalidInput(_) + )); + } + + #[tokio::test] + async fn seeded_refund_idempotency_replays_only_explicit_mappings_in_memory() { + let seeded = sample_refund("refund-seeded-idempotency", Some("user-1"), "approved"); + let mapped_repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + wallets: vec![sample_wallet()], + refunds: vec![seeded.clone()], + refund_idempotency: vec![( + "user-1".to_string(), + "seeded-refund-key".to_string(), + seeded.id.clone(), + )], + ..WalletReadSeed::default() + }); + let replay = mapped_repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("seeded replay".to_string()), + idempotency_key: Some("seeded-refund-key".to_string()), + refund_no: "seeded-refund-replay".to_string(), + }) + .await + .expect("seeded refund replay should resolve"); + assert!(matches!( + replay, + CreateWalletRefundRequestOutcome::Duplicate(refund) if refund.id == seeded.id + )); + assert_eq!( + mapped_repository + .refunds_by_id + .read() + .expect("wallet repo lock") + .len(), + 1 + ); + + let unmapped_repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + wallets: vec![sample_wallet()], + refunds: vec![seeded], + ..WalletReadSeed::default() + }); + let created = unmapped_repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("unmapped seed".to_string()), + idempotency_key: Some("seeded-refund-key".to_string()), + refund_no: "seeded-refund-unmapped".to_string(), + }) + .await + .expect("unmapped seeded refund request should resolve"); + assert!(matches!( + created, + CreateWalletRefundRequestOutcome::Created(_) + )); + assert_eq!( + unmapped_repository + .refunds_by_id + .read() + .expect("wallet repo lock") + .len(), + 2 + ); + } + #[tokio::test] async fn counts_pending_user_refunds_and_payment_orders() { let repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { @@ -2020,6 +4223,7 @@ mod tests { sample_refund("refund-3", Some("user-1"), "completed"), sample_refund("refund-4", Some("user-2"), "approved"), ], + refund_idempotency: Vec::new(), redeem_batches: Vec::new(), redeem_codes: Vec::new(), }); @@ -2039,4 +4243,73 @@ mod tests { 2 ); } + + #[tokio::test] + async fn refund_reservation_rejects_invalid_active_amounts_in_memory() { + for (label, amount) in [ + ("negative", -100.0), + ("zero", 0.0), + ("infinite", f64::INFINITY), + ("nan", f64::NAN), + ] { + let mut invalid = sample_refund( + &format!("refund-invalid-{label}"), + Some("user-1"), + "pending_approval", + ); + invalid.amount_usd = amount; + let repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + wallets: vec![sample_wallet()], + refunds: vec![invalid], + ..WalletReadSeed::default() + }); + let outcome = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some(format!("invalid reservation: {label}")), + idempotency_key: Some(format!("reservation-invalid-{label}")), + refund_no: format!("reservation-invalid-{label}"), + }) + .await + .expect("reservation request should resolve"); + assert!( + matches!(outcome, CreateWalletRefundRequestOutcome::InvalidInput(_)), + "active {label} reservation must fail closed: {outcome:?}" + ); + } + + // Completed refunds do not reserve balance and remain ignored. + let mut completed = sample_refund("refund-completed", Some("user-1"), "completed"); + completed.amount_usd = 100.0; + let repository = InMemoryWalletRepository::seed_read_model(WalletReadSeed { + wallets: vec![sample_wallet()], + refunds: vec![completed], + ..WalletReadSeed::default() + }); + let outcome = repository + .create_wallet_refund_request(CreateWalletRefundRequestInput { + wallet_id: "wallet-1".to_string(), + user_id: "user-1".to_string(), + amount_usd: 1.0, + payment_order_id: None, + source_type: None, + source_id: None, + refund_mode: None, + reason: Some("completed reservation is ignored".to_string()), + idempotency_key: Some("reservation-completed-ignored".to_string()), + refund_no: "reservation-completed-ignored".to_string(), + }) + .await + .expect("completed reservation should not block request"); + assert!(matches!( + outcome, + CreateWalletRefundRequestOutcome::Created(_) + )); + } } diff --git a/crates/aether-data/runtime/src/repository/wallet/mod.rs b/crates/aether-data/runtime/src/repository/wallet/mod.rs index 4e16172cc..1706deca4 100644 --- a/crates/aether-data/runtime/src/repository/wallet/mod.rs +++ b/crates/aether-data/runtime/src/repository/wallet/mod.rs @@ -1,19 +1,34 @@ mod memory; pub use aether_data_contracts::repository::wallet::{ - redeem_code_credits_recharge_balance, redeem_code_payment_method, - redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentCallbackRecord, - AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, - AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletPaymentOrderRecord, - AdminWalletRefundRecord, AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord, - CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, - CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, - CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, - CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, - CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, - CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + canonicalize_payment_method, canonicalize_wallet_refund_fields, + payment_order_is_uncertain_wallet_checkout_placeholder, + payment_order_refund_amounts_are_consistent, + payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response, + project_wallet_recharge_gateway_response, redeem_code_credits_recharge_balance, + redeem_code_payment_method, redeem_code_refundable_amount, stored_timestamp_unix_secs, + validate_admin_redeem_code_batch_input, validate_payment_order_credit_amounts, + validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements, + validate_redeem_wallet_credit, validate_wallet_recharge_order_input, + wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token, + wallet_recharge_checkout_claimed_at, wallet_recharge_checkout_failed_response, + wallet_recharge_checkout_uncertain_response, wallet_recharge_order_created_at_unix_secs, + 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, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, + ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, @@ -22,9 +37,10 @@ pub use aether_data_contracts::repository::wallet::{ StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, - StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome, + StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, + UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletReadSeed, WalletReadSnapshot, WalletRepository, - WalletWriteRepository, + WalletWriteRepository, WALLET_RECHARGE_CHECKOUT_CLAIM_LEASE_SECS, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::MysqlWalletReadRepository; diff --git a/crates/aether-data/schema/src/lib.rs b/crates/aether-data/schema/src/lib.rs index 13f1a6d7c..639c65e5b 100644 --- a/crates/aether-data/schema/src/lib.rs +++ b/crates/aether-data/schema/src/lib.rs @@ -893,6 +893,30 @@ mod tests { )); } + #[test] + fn mysql_column_type_override_can_select_binary_collation() { + let mut schema = announcements_schema(); + schema + .tables + .get_mut("announcements") + .expect("fixture table exists") + .columns + .first_mut() + .expect("fixture id column exists") + .driver + .mysql = Some(DriverColumnOverride { + sql_type: Some( + "VARCHAR(64) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin".to_string(), + ), + default: None, + nullable: None, + }); + + let mysql_sql = mysql::emit_schema(&schema); + assert!(mysql_sql + .contains("`id` VARCHAR(64) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NOT NULL")); + } + #[test] fn extracts_create_table_names_across_driver_quoting() { let tables = extract_create_table_names( diff --git a/crates/aether-gateway/control/src/public/request_context.rs b/crates/aether-gateway/control/src/public/request_context.rs index 8c38439e0..d0126cf99 100644 --- a/crates/aether-gateway/control/src/public/request_context.rs +++ b/crates/aether-gateway/control/src/public/request_context.rs @@ -8,6 +8,7 @@ pub struct PublicRequestContext { pub request_query_string: Option, pub request_content_type: Option, pub host_header: Option, + pub client_ip: Option, pub control_decision: Option, } @@ -32,6 +33,7 @@ impl PublicRequestContext { request_query_string: uri.query().map(ToOwned::to_owned), request_content_type: header_value(headers, http::header::CONTENT_TYPE), host_header: header_value(headers, http::header::HOST), + client_ip: None, control_decision, } } diff --git a/crates/aether-gateway/execution/src/stream/ndjson.rs b/crates/aether-gateway/execution/src/stream/ndjson.rs index 990eb2411..d7ba00e66 100644 --- a/crates/aether-gateway/execution/src/stream/ndjson.rs +++ b/crates/aether-gateway/execution/src/stream/ndjson.rs @@ -17,7 +17,9 @@ pub fn decode_stream_frame_ndjson(line: &[u8]) -> Result { mod tests { use std::collections::BTreeMap; - use aether_contracts::{StreamFramePayload, StreamFrameType}; + use aether_contracts::{ + ExecutionStreamTerminalSummary, StandardizedUsage, StreamFramePayload, StreamFrameType, + }; use super::{decode_stream_frame_ndjson, encode_stream_frame_ndjson}; @@ -41,4 +43,22 @@ mod tests { decode_stream_frame_ndjson(raw.trim_ascii_end()).expect("frame should decode"); assert_eq!(decoded, frame); } + + #[test] + fn ndjson_round_trip_preserves_terminal_usage_with_fractional_fields() { + let mut usage = StandardizedUsage::new(); + usage.cache_storage_token_hours = 0.125; + let frame = + aether_contracts::StreamFrame::eof_with_summary(Some(ExecutionStreamTerminalSummary { + standardized_usage: Some(usage), + observed_finish: true, + ..ExecutionStreamTerminalSummary::default() + })); + + let raw = encode_stream_frame_ndjson(&frame).expect("frame should encode"); + let decoded = decode_stream_frame_ndjson(raw.trim_ascii_end()) + .expect("terminal usage frame should decode"); + + assert_eq!(decoded, frame); + } } diff --git a/crates/aether-gateway/frontdoor/src/body.rs b/crates/aether-gateway/frontdoor/src/body.rs index 1d7e8f465..a58d9a920 100644 --- a/crates/aether-gateway/frontdoor/src/body.rs +++ b/crates/aether-gateway/frontdoor/src/body.rs @@ -14,11 +14,12 @@ use std::time::{Duration, Instant}; use tokio::sync::{OwnedSemaphorePermit, Semaphore}; pub const DEFAULT_BODY_BUFFER_PERMIT_BYTES: usize = 64 * 1024; +const MAX_REQUEST_CONTENT_ENCODINGS: usize = 8; #[derive(Debug, Clone)] pub struct BodyBufferPolicy { max_bytes: u64, - read_timeout: Duration, + read_timeout: Option, queue_timeout: Duration, budget_bytes: usize, permit_bytes: usize, @@ -33,7 +34,23 @@ impl BodyBufferPolicy { budget_bytes: usize, budget: Arc, ) -> Self { - Self::with_permit_bytes( + Self::new_with_optional_read_timeout( + max_bytes, + Some(read_timeout), + queue_timeout, + budget_bytes, + budget, + ) + } + + pub fn new_with_optional_read_timeout( + max_bytes: u64, + read_timeout: Option, + queue_timeout: Duration, + budget_bytes: usize, + budget: Arc, + ) -> Self { + Self::with_optional_read_timeout_and_permit_bytes( max_bytes, read_timeout, queue_timeout, @@ -50,10 +67,28 @@ impl BodyBufferPolicy { budget_bytes: usize, permit_bytes: usize, budget: Arc, + ) -> Self { + Self::with_optional_read_timeout_and_permit_bytes( + max_bytes, + Some(read_timeout), + queue_timeout, + budget_bytes, + permit_bytes, + budget, + ) + } + + pub fn with_optional_read_timeout_and_permit_bytes( + max_bytes: u64, + read_timeout: Option, + queue_timeout: Duration, + budget_bytes: usize, + permit_bytes: usize, + budget: Arc, ) -> Self { Self { max_bytes, - read_timeout, + read_timeout: read_timeout.filter(|timeout| !timeout.is_zero()), queue_timeout, budget_bytes, permit_bytes: permit_bytes.max(1), @@ -66,6 +101,10 @@ impl BodyBufferPolicy { } pub fn read_timeout(&self) -> Duration { + self.read_timeout.unwrap_or_default() + } + + pub fn optional_read_timeout(&self) -> Option { self.read_timeout } @@ -77,6 +116,16 @@ impl BodyBufferPolicy { self.budget_bytes } + /// The body buffer budget is also a hard per-request ceiling. A request + /// whose configured body limit is larger than the shared budget must not + /// be allowed to reserve only the budget and then collect/decompress past + /// it, otherwise chunked and compressed requests can defeat the memory + /// bound. + pub fn effective_max_bytes(&self) -> u64 { + self.max_bytes + .min(u64::try_from(self.budget_bytes).unwrap_or(u64::MAX)) + } + pub fn reservation_bytes(&self, headers: &HeaderMap) -> usize { reservation_bytes(headers, self.max_bytes, self.budget_bytes) } @@ -89,10 +138,12 @@ impl BodyBufferPolicy { &self, headers: &HeaderMap, ) -> Result { - if let Some(declared) = declared_content_length(headers) { - if declared > self.max_bytes { + let effective_max_bytes = self.effective_max_bytes(); + let declared_content_length = validate_request_body_headers(headers)?; + if let Some(declared) = declared_content_length { + if declared > effective_max_bytes { return Err(BodyBufferError::TooLarge { - limit_bytes: self.max_bytes, + limit_bytes: effective_max_bytes, }); } } @@ -118,7 +169,7 @@ impl BodyBufferPolicy { Ok(BodyBufferReservation { permit, - max_bytes: self.max_bytes, + max_bytes: effective_max_bytes, read_timeout: self.read_timeout, requested_bytes, }) @@ -129,7 +180,7 @@ impl BodyBufferPolicy { pub struct BodyBufferReservation { permit: OwnedSemaphorePermit, max_bytes: u64, - read_timeout: Duration, + read_timeout: Option, requested_bytes: usize, } @@ -147,23 +198,31 @@ impl BodyBufferReservation { } = self; let started_at = Instant::now(); let body_limit = usize::try_from(max_bytes).unwrap_or(usize::MAX); - let bytes = match tokio::time::timeout(read_timeout, to_bytes(body, body_limit)).await { - Ok(Ok(bytes)) => bytes, - Ok(Err(error)) if collection_exceeded_limit(&error) => { + let collected = match read_timeout { + Some(read_timeout) => { + match tokio::time::timeout(read_timeout, to_bytes(body, body_limit)).await { + Ok(result) => result, + Err(_) => { + return Err(BodyBufferError::Timeout { + timeout_ms: duration_millis(read_timeout), + }); + } + } + } + None => to_bytes(body, body_limit).await, + }; + let bytes = match collected { + Ok(bytes) => bytes, + Err(error) if collection_exceeded_limit(&error) => { return Err(BodyBufferError::TooLarge { limit_bytes: max_bytes, }); } - Ok(Err(error)) => { + Err(error) => { return Err(BodyBufferError::ReadFailed { message: error.to_string(), }); } - Err(_) => { - return Err(BodyBufferError::Timeout { - timeout_ms: duration_millis(read_timeout), - }); - } }; Ok(BufferedBody { @@ -208,6 +267,9 @@ impl BufferedBody { #[derive(Debug, PartialEq, Eq)] pub enum BodyBufferError { + InvalidHeaders { + message: String, + }, TooLarge { limit_bytes: u64, }, @@ -227,6 +289,7 @@ pub enum BodyBufferError { impl BodyBufferError { pub fn http_status(&self) -> StatusCode { match self { + Self::InvalidHeaders { .. } => StatusCode::BAD_REQUEST, Self::TooLarge { .. } => StatusCode::PAYLOAD_TOO_LARGE, Self::Overloaded { .. } => StatusCode::SERVICE_UNAVAILABLE, Self::Timeout { .. } => StatusCode::REQUEST_TIMEOUT, @@ -236,6 +299,7 @@ impl BodyBufferError { pub fn client_message(&self) -> String { match self { + Self::InvalidHeaders { .. } => "Invalid request body headers".to_string(), Self::TooLarge { limit_bytes } => format!("Request body exceeds {limit_bytes} bytes"), Self::Overloaded { .. } => { "Request body buffering capacity is temporarily exhausted".to_string() @@ -249,6 +313,7 @@ impl BodyBufferError { pub fn reason(&self) -> &'static str { match self { + Self::InvalidHeaders { .. } => "invalid_request_body_headers", Self::TooLarge { .. } => "request_body_too_large", Self::Overloaded { .. } => "request_body_buffer_overloaded", Self::Timeout { .. } => "request_body_read_timeout", @@ -257,11 +322,71 @@ impl BodyBufferError { } } -fn declared_content_length(headers: &HeaderMap) -> Option { - headers - .get(header::CONTENT_LENGTH) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.trim().parse::().ok()) +fn validate_request_body_headers(headers: &HeaderMap) -> Result, BodyBufferError> { + let declared_content_length = declared_content_length(headers)?; + if declared_content_length.is_some() && headers.contains_key(header::TRANSFER_ENCODING) { + return Err(invalid_body_headers( + "content-length and transfer-encoding must not be combined", + )); + } + validate_content_encoding_headers(headers)?; + Ok(declared_content_length) +} + +fn declared_content_length(headers: &HeaderMap) -> Result, BodyBufferError> { + let mut values = headers.get_all(header::CONTENT_LENGTH).iter(); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.next().is_some() { + return Err(invalid_body_headers("duplicate content-length header")); + } + let value = value + .to_str() + .map_err(|_| invalid_body_headers("invalid content-length header"))? + .trim(); + if value.is_empty() || value.contains(',') { + return Err(invalid_body_headers("ambiguous content-length header")); + } + value + .parse::() + .map(Some) + .map_err(|_| invalid_body_headers("invalid content-length header")) +} + +fn validate_content_encoding_headers(headers: &HeaderMap) -> Result<(), BodyBufferError> { + if headers + .get_all(header::CONTENT_ENCODING) + .iter() + .nth(1) + .is_some() + { + return Err(invalid_body_headers("duplicate content-encoding header")); + } + let mut count = 0usize; + for value in headers.get_all(header::CONTENT_ENCODING).iter() { + let value = value + .to_str() + .map_err(|_| invalid_body_headers("invalid content-encoding header"))?; + for encoding in value.split(',') { + if encoding.trim().is_empty() { + return Err(invalid_body_headers("invalid content-encoding header")); + } + count = count.saturating_add(1); + if count > MAX_REQUEST_CONTENT_ENCODINGS { + return Err(invalid_body_headers( + "content-encoding chain exceeds the supported limit", + )); + } + } + } + Ok(()) +} + +fn invalid_body_headers(message: &str) -> BodyBufferError { + BodyBufferError::InvalidHeaders { + message: message.to_string(), + } } fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) -> usize { @@ -269,14 +394,21 @@ fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) - .unwrap_or(usize::MAX) .min(budget_bytes); let encoded = headers - .get(header::CONTENT_ENCODING) - .and_then(|value| value.to_str().ok()) - .map(str::trim) - .is_some_and(|value| !value.is_empty() && !value.eq_ignore_ascii_case("identity")); + .get_all(header::CONTENT_ENCODING) + .iter() + .any(|value| { + value.to_str().map_or(true, |value| { + value.split(',').map(str::trim).any(|encoding| { + !encoding.is_empty() && !encoding.eq_ignore_ascii_case("identity") + }) + }) + }); if encoded { return reservation_ceiling; } declared_content_length(headers) + .ok() + .flatten() .map(|value| { usize::try_from(value) .unwrap_or(usize::MAX) @@ -341,10 +473,63 @@ mod tests { } #[tokio::test] - async fn unlimited_body_uses_budget_as_reservation_ceiling() { + async fn rejects_ambiguous_content_length_headers_before_reading_body() { + let mut headers = HeaderMap::new(); + headers.append(header::CONTENT_LENGTH, HeaderValue::from_static("5")); + headers.append(header::CONTENT_LENGTH, HeaderValue::from_static("5")); + let error = policy(10, Duration::from_secs(1), Arc::new(Semaphore::new(1))) + .reserve(&headers) + .await + .expect_err("duplicate content-length must be rejected"); + assert!(matches!(error, BodyBufferError::InvalidHeaders { .. })); + } + + #[tokio::test] + async fn rejects_content_length_with_transfer_encoding_before_reading_body() { + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("5")); + headers.insert( + header::TRANSFER_ENCODING, + HeaderValue::from_static("chunked"), + ); + let error = policy(10, Duration::from_secs(1), Arc::new(Semaphore::new(1))) + .reserve(&headers) + .await + .expect_err("content-length plus transfer-encoding must be rejected"); + assert!(matches!(error, BodyBufferError::InvalidHeaders { .. })); + } + + #[tokio::test] + async fn rejects_overlong_content_encoding_chain_before_reading_body() { + let mut headers = HeaderMap::new(); + headers.insert( + header::CONTENT_ENCODING, + HeaderValue::from_static("gzip, gzip, gzip, gzip, gzip, gzip, gzip, gzip, gzip"), + ); + let error = policy(10, Duration::from_secs(1), Arc::new(Semaphore::new(1))) + .reserve(&headers) + .await + .expect_err("overlong content-encoding chain must be rejected"); + assert!(matches!(error, BodyBufferError::InvalidHeaders { .. })); + } + + #[tokio::test] + async fn rejects_duplicate_content_encoding_fields_before_reading_body() { + let mut headers = HeaderMap::new(); + headers.append(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + headers.append(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + let error = policy(10, Duration::from_secs(1), Arc::new(Semaphore::new(1))) + .reserve(&headers) + .await + .expect_err("duplicate content-encoding fields must be rejected"); + assert!(matches!(error, BodyBufferError::InvalidHeaders { .. })); + } + + #[tokio::test] + async fn body_limit_larger_than_budget_is_capped_before_reading() { let budget = Arc::new(Semaphore::new(5)); let policy = BodyBufferPolicy::with_permit_bytes( - u64::MAX, + 10, Duration::from_secs(1), Duration::from_secs(1), 5, @@ -354,24 +539,31 @@ mod tests { let mut headers = HeaderMap::new(); headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("10")); - let reservation = policy + let error = policy .reserve(&headers) .await - .expect("unlimited body should reserve the available budget"); - assert_eq!(reservation.requested_bytes(), 5); - assert_eq!(budget.available_permits(), 0); + .expect_err("declared body above the shared budget must be rejected"); + assert_eq!(error, BodyBufferError::TooLarge { limit_bytes: 5 }); + assert_eq!(budget.available_permits(), 5); - let buffered = reservation - .collect(Body::from(Bytes::from_static(b"0123456789"))) + let mut small_headers = HeaderMap::new(); + small_headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("5")); + let reservation = policy + .reserve(&small_headers) .await - .expect("unlimited body should collect beyond the reservation ceiling"); - assert_eq!(buffered.bytes().as_ref(), b"0123456789"); + .expect("body within the shared budget should remain accepted"); + assert_eq!(reservation.requested_bytes(), 5); + let buffered = reservation + .collect(Body::from(Bytes::from_static(b"01234"))) + .await + .expect("body within the shared budget should collect"); + assert_eq!(buffered.bytes().as_ref(), b"01234"); drop(buffered); assert_eq!(budget.available_permits(), 5); } #[tokio::test] - async fn encoded_unlimited_body_reserves_the_full_budget() { + async fn unlimited_body_is_still_bounded_by_budget() { let budget = Arc::new(Semaphore::new(4)); let policy = BodyBufferPolicy::with_permit_bytes( u64::MAX, @@ -390,6 +582,25 @@ mod tests { .expect("encoded unlimited body should reserve the available budget"); assert_eq!(reservation.requested_bytes(), 4); assert_eq!(budget.available_permits(), 0); + + let error = reservation + .collect(Body::from(Bytes::from_static(b"01234"))) + .await + .expect_err("unlimited body must still be bounded by the shared budget"); + assert_eq!(error, BodyBufferError::TooLarge { limit_bytes: 4 }); + assert_eq!(budget.available_permits(), 4); + + let reservation = policy + .reserve(&HeaderMap::new()) + .await + .expect("a body at the budget boundary should reserve"); + let buffered = reservation + .collect(Body::from(Bytes::from_static(b"0123"))) + .await + .expect("a body at the budget boundary should collect"); + assert_eq!(buffered.bytes().as_ref(), b"0123"); + drop(buffered); + assert_eq!(budget.available_permits(), 4); } #[tokio::test] @@ -441,6 +652,48 @@ mod tests { assert_eq!(error, BodyBufferError::Timeout { timeout_ms: 5 }); } + #[tokio::test] + async fn disabled_read_timeout_allows_slow_body_to_complete() { + let stream = stream::once(async { Ok::(Bytes::from_static(b"{")) }) + .chain(stream::once(async { + tokio::time::sleep(Duration::from_millis(20)).await; + Ok::(Bytes::from_static(b"}")) + })); + let policy = BodyBufferPolicy::with_optional_read_timeout_and_permit_bytes( + 1024, + None, + Duration::from_secs(1), + 1024, + DEFAULT_BODY_BUFFER_PERMIT_BYTES, + Arc::new(Semaphore::new(1)), + ); + let reservation = policy + .reserve(&HeaderMap::new()) + .await + .expect("reservation should succeed"); + let buffered = tokio::time::timeout( + Duration::from_secs(1), + reservation.collect(Body::from_stream(stream)), + ) + .await + .expect("test body should finish") + .expect("disabled read timeout should allow a slow body"); + assert_eq!(buffered.bytes().as_ref(), b"{}"); + } + + #[test] + fn zero_read_timeout_is_normalized_to_disabled() { + let policy = BodyBufferPolicy::new_with_optional_read_timeout( + 1024, + Some(Duration::ZERO), + Duration::from_secs(1), + 1024, + Arc::new(Semaphore::new(1)), + ); + assert_eq!(policy.optional_read_timeout(), None); + assert_eq!(policy.read_timeout(), Duration::ZERO); + } + #[tokio::test] async fn rejects_when_weighted_budget_is_exhausted() { let budget = Arc::new(Semaphore::new(1)); diff --git a/crates/aether-gateway/frontdoor/src/lib.rs b/crates/aether-gateway/frontdoor/src/lib.rs index 9343ac676..4adf89dc5 100644 --- a/crates/aether-gateway/frontdoor/src/lib.rs +++ b/crates/aether-gateway/frontdoor/src/lib.rs @@ -11,7 +11,5 @@ pub use middleware::access_log::{ access_log_middleware, sanitize_access_log_path, should_downgrade_access_log, GatewayRequestAcceptedAt, RequestLogEmitted, }; -pub use middleware::cf_headers::{ - apply_cf_header_stripping, strip_cf_headers_middleware, CfConnectingIp, -}; +pub use middleware::cf_headers::{apply_cf_header_stripping, strip_cf_headers_middleware}; pub use request_id::short_request_id; diff --git a/crates/aether-gateway/frontdoor/src/middleware/cf_headers.rs b/crates/aether-gateway/frontdoor/src/middleware/cf_headers.rs index eff034b32..bfc8093e4 100644 --- a/crates/aether-gateway/frontdoor/src/middleware/cf_headers.rs +++ b/crates/aether-gateway/frontdoor/src/middleware/cf_headers.rs @@ -1,26 +1,14 @@ use axum::{extract::Request, middleware::Next, response::Response, Router}; -use http::{header::HeaderName, HeaderMap}; +use http::header::HeaderName; const CF_EXACT_HEADERS: &[&str] = &["cdn-loop", "true-client-ip"]; -#[derive(Clone, Debug)] -pub struct CfConnectingIp(pub String); - fn should_strip_cf_header(name: &HeaderName) -> bool { let normalized = name.as_str(); normalized.starts_with("cf-") || CF_EXACT_HEADERS.contains(&normalized) } -fn cf_connecting_ip(headers: &HeaderMap) -> Option { - headers - .get("cf-connecting-ip") - .and_then(|value| value.to_str().ok()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| value.chars().take(45).collect()) -} - -fn strip_cf_headers(headers: &mut HeaderMap) { +fn strip_cf_headers(headers: &mut http::HeaderMap) { let to_remove: Vec<_> = headers .keys() .filter(|name| should_strip_cf_header(name)) @@ -36,9 +24,6 @@ pub fn apply_cf_header_stripping(router: Router) -> Router { } pub async fn strip_cf_headers_middleware(mut request: Request, next: Next) -> Response { - if let Some(client_ip) = cf_connecting_ip(request.headers()) { - request.extensions_mut().insert(CfConnectingIp(client_ip)); - } strip_cf_headers(request.headers_mut()); let mut response = next.run(request).await; diff --git a/crates/aether-gateway/tunnel/src/embedded/protocol.rs b/crates/aether-gateway/tunnel/src/embedded/protocol.rs index 8ed58dedc..296add4c9 100644 --- a/crates/aether-gateway/tunnel/src/embedded/protocol.rs +++ b/crates/aether-gateway/tunnel/src/embedded/protocol.rs @@ -3,14 +3,15 @@ use bytes::Bytes; pub use aether_contracts::tunnel::{ - decode_payload, encode_connection_close, encode_frame, encode_goaway, encode_goaway_v3, - encode_hello, encode_load_report, encode_ping, encode_pong, encode_reset_stream, - encode_settings, encode_stream_error, encode_window_update, frame_payload_by_header, - ConnectionClosePayload, FrameHeader, GoAwayPayload, HelloPayload, LoadReportPayload, - RequestMeta, ResetStreamPayload, ResponseMeta, SettingsPayload, WindowUpdatePayload, - CONNECTION_CLOSE, FLAG_END_STREAM, FLAG_GZIP_COMPRESSED, GOAWAY, HEADER_SIZE, HEARTBEAT_ACK, - HEARTBEAT_DATA, HELLO, LOAD_REPORT, PING, PONG, REQUEST_BODY, REQUEST_HEADERS, RESET_STREAM, - RESPONSE_BODY, RESPONSE_HEADERS, SETTINGS, STREAM_END, STREAM_ERROR, WINDOW_UPDATE, + decode_payload, decode_payload_with_limit, encode_connection_close, encode_frame, + encode_goaway, encode_goaway_v3, encode_hello, encode_load_report, encode_ping, encode_pong, + encode_reset_stream, encode_settings, encode_stream_error, encode_window_update, + frame_payload_by_header, ConnectionClosePayload, FrameHeader, GoAwayPayload, HelloPayload, + LoadReportPayload, RequestMeta, ResetStreamPayload, ResponseMeta, SettingsPayload, + WindowUpdatePayload, CONNECTION_CLOSE, FLAG_END_STREAM, FLAG_GZIP_COMPRESSED, GOAWAY, + HEADER_SIZE, HEARTBEAT_ACK, HEARTBEAT_DATA, HELLO, LOAD_REPORT, PING, PONG, REQUEST_BODY, + REQUEST_HEADERS, RESET_STREAM, RESPONSE_BODY, RESPONSE_HEADERS, SETTINGS, STREAM_END, + STREAM_ERROR, WINDOW_UPDATE, }; pub fn compress_payload(payload: &[u8]) -> Result<(Vec, u8), std::io::Error> { diff --git a/crates/aether-gateway/tunnel/src/hub.rs b/crates/aether-gateway/tunnel/src/hub.rs index 06ee460ed..02264b1c7 100644 --- a/crates/aether-gateway/tunnel/src/hub.rs +++ b/crates/aether-gateway/tunnel/src/hub.rs @@ -2,6 +2,7 @@ use base64::Engine as _; use http::HeaderMap; pub const MAX_TUNNEL_STREAMS: usize = 2_048; +const MAX_TUNNEL_NODE_NAME_UTF8_BYTES: usize = 100 * 4; pub fn resolve_proxy_max_streams(headers: &HeaderMap, fallback: usize) -> usize { headers @@ -17,10 +18,20 @@ pub fn resolve_proxy_node_name(headers: &HeaderMap, node_id: &str) -> String { .get(aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER) .and_then(|value| value.to_str().ok()) .and_then(|value| { + let value = value.trim(); + let max_encoded_len = MAX_TUNNEL_NODE_NAME_UTF8_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4); + if value.len() > max_encoded_len { + return None; + } base64::engine::general_purpose::URL_SAFE_NO_PAD - .decode(value.trim()) + .decode(value) .ok() }) + .filter(|bytes| bytes.len() <= MAX_TUNNEL_NODE_NAME_UTF8_BYTES) .and_then(|bytes| String::from_utf8(bytes).ok()) .map(|value| value.trim().to_string()) .filter(|value| !value.is_empty() && value.chars().count() <= 100) @@ -91,5 +102,17 @@ mod tests { headers.remove(aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER); headers.insert("x-node-name", HeaderValue::from_static("edge-1")); assert_eq!(resolve_proxy_node_name(&headers, "node-1"), "edge-1"); + + let max_encoded_len = super::MAX_TUNNEL_NODE_NAME_UTF8_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap() + .saturating_mul(4); + headers.insert( + aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER, + HeaderValue::from_str(&"A".repeat(max_encoded_len + 1)) + .expect("oversized encoded header should parse"), + ); + assert_eq!(resolve_proxy_node_name(&headers, "node-1"), "edge-1"); } } diff --git a/crates/aether-gateway/tunnel/src/relay.rs b/crates/aether-gateway/tunnel/src/relay.rs index ea649e8e3..610e4f45f 100644 --- a/crates/aether-gateway/tunnel/src/relay.rs +++ b/crates/aether-gateway/tunnel/src/relay.rs @@ -12,6 +12,8 @@ pub const DEFAULT_TUNNEL_PROBE_BODY_LIMIT_BYTES: usize = 64 * 1024; pub struct TunnelAttachmentRecord { pub gateway_instance_id: String, pub relay_base_url: String, + #[serde(default)] + pub tunnel_generation: String, pub conn_count: usize, pub observed_at_unix_secs: u64, } @@ -20,6 +22,7 @@ impl TunnelAttachmentRecord { pub fn is_routable(&self, now_unix_secs: u64, ttl_secs: u64) -> bool { self.conn_count > 0 && !self.relay_base_url.trim().is_empty() + && !self.tunnel_generation.trim().is_empty() && self.observed_at_unix_secs.saturating_add(ttl_secs) >= now_unix_secs } @@ -44,6 +47,7 @@ mod tests { TunnelAttachmentRecord { gateway_instance_id: "gateway-a".to_string(), relay_base_url: "http://gateway-a.internal".to_string(), + tunnel_generation: "generation-a".to_string(), conn_count: 1, observed_at_unix_secs: 100, } diff --git a/crates/aether-http/Cargo.toml b/crates/aether-http/Cargo.toml index 44bb0f8c8..297bdbbcb 100644 --- a/crates/aether-http/Cargo.toml +++ b/crates/aether-http/Cargo.toml @@ -9,3 +9,5 @@ description = "Shared HTTP client config and retry helpers for Aether Rust servi [dependencies] reqwest.workspace = true serde.workspace = true +tokio.workspace = true +url.workspace = true diff --git a/crates/aether-http/src/client.rs b/crates/aether-http/src/client.rs index 6339140bb..35edfed7d 100644 --- a/crates/aether-http/src/client.rs +++ b/crates/aether-http/src/client.rs @@ -44,7 +44,9 @@ pub fn build_http_client_with_headers( config: &HttpClientConfig, default_headers: HeaderMap, ) -> Result { - let mut builder = apply_http_client_config(reqwest::Client::builder(), config); + // Shared clients must not silently inherit HTTP(S)_PROXY. Callers that + // need a proxy install the explicitly configured URL below. + let mut builder = apply_http_client_config(reqwest::Client::builder().no_proxy(), config); if let Some(proxy_url) = config .proxy_url .as_deref() diff --git a/crates/aether-http/src/dns.rs b/crates/aether-http/src/dns.rs new file mode 100644 index 000000000..6cf6c13c7 --- /dev/null +++ b/crates/aether-http/src/dns.rs @@ -0,0 +1,91 @@ +use std::io; +use std::net::SocketAddr; +use std::time::Duration; + +/// Maximum number of addresses accepted from one hostname lookup. +/// +/// A DNS response is attacker-controlled at the resolver boundary. Keeping +/// the result bounded prevents a pathological answer from forcing an +/// unbounded vector allocation or an unbounded connect fan-out. Callers still +/// validate every returned address for their own network policy. +pub const MAX_DNS_RESOLVED_ADDRESSES: usize = 32; + +/// Upper bound used by callers that do not have a tighter request deadline. +pub const DEFAULT_DNS_LOOKUP_TIMEOUT: Duration = Duration::from_secs(10); + +/// Resolve a host while bounding both resolver wait time and answer count. +/// +/// The iterator is consumed one item past the allowed count so an answer set +/// larger than the policy is rejected rather than silently truncated. This +/// keeps validation and connection pinning based on the complete, bounded +/// answer set. +pub async fn lookup_host_with_limits( + host: &str, + port: u16, + timeout: Duration, +) -> io::Result> { + if timeout.is_zero() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "DNS lookup timeout must be non-zero", + )); + } + + let mut resolved = tokio::time::timeout(timeout, tokio::net::lookup_host((host, port))) + .await + .map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "DNS lookup timed out"))??; + + collect_resolved_addresses_with_limit(&mut resolved) +} + +fn collect_resolved_addresses_with_limit( + resolved: &mut impl Iterator, +) -> io::Result> { + let mut addresses = Vec::with_capacity(MAX_DNS_RESOLVED_ADDRESSES.min(8)); + for address in resolved.by_ref() { + if addresses.len() >= MAX_DNS_RESOLVED_ADDRESSES { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "DNS lookup returned too many addresses", + )); + } + addresses.push(address); + } + Ok(addresses) +} + +#[cfg(test)] +mod tests { + use std::net::SocketAddr; + + use super::{ + collect_resolved_addresses_with_limit, lookup_host_with_limits, DEFAULT_DNS_LOOKUP_TIMEOUT, + MAX_DNS_RESOLVED_ADDRESSES, + }; + + #[tokio::test] + async fn rejects_zero_dns_timeout_before_resolving() { + let error = lookup_host_with_limits("localhost", 80, std::time::Duration::ZERO) + .await + .expect_err("zero timeout must be rejected"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + } + + #[tokio::test] + async fn resolves_within_shared_address_bound() { + let addresses = lookup_host_with_limits("localhost", 80, DEFAULT_DNS_LOOKUP_TIMEOUT) + .await + .expect("localhost should resolve in the test environment"); + assert!(!addresses.is_empty()); + assert!(addresses.len() <= MAX_DNS_RESOLVED_ADDRESSES); + } + + #[test] + fn rejects_an_answer_set_larger_than_the_shared_bound() { + let mut resolved = (0..=MAX_DNS_RESOLVED_ADDRESSES) + .map(|index| SocketAddr::from(([192, 0, 2, (index % 254 + 1) as u8], 443))); + let error = collect_resolved_addresses_with_limit(&mut resolved) + .expect_err("more than 32 DNS answers must be rejected"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + } +} diff --git a/crates/aether-http/src/header_security.rs b/crates/aether-http/src/header_security.rs new file mode 100644 index 000000000..df69ee609 --- /dev/null +++ b/crates/aether-http/src/header_security.rs @@ -0,0 +1,215 @@ +use std::collections::BTreeSet; +use std::net::IpAddr; + +use url::{Host, Url}; + +/// Parse the case-insensitive field names nominated by HTTP/1 `Connection` +/// headers. Those fields are hop-by-hop even when their names are otherwise +/// application-defined. +pub fn connection_declared_header_names<'a>( + values: impl IntoIterator, +) -> BTreeSet { + values + .into_iter() + .flat_map(|value| value.split(',')) + .map(str::trim) + .filter(|value| valid_http_token(value)) + .map(str::to_ascii_lowercase) + .collect() +} + +fn valid_http_token(value: &str) -> bool { + !value.is_empty() + && value.bytes().all(|byte| { + byte.is_ascii_alphanumeric() + || matches!( + byte, + b'!' | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) + }) +} + +/// Return true for addresses that an untrusted URL must not be allowed to +/// reach, including private/link-local ranges and IPv6 transition formats +/// that can embed an IPv4 destination. +pub fn is_private_or_reserved_ip(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => { + let octets = ip.octets(); + ip.is_private() + || ip.is_loopback() + || ip.is_link_local() + || ip.is_broadcast() + || ip.is_documentation() + || ip.is_unspecified() + || ip.is_multicast() + // The complete 0.0.0.0/8 block is reserved for "this + // network" destinations. `Ipv4Addr::is_unspecified()` + // only covers the single 0.0.0.0 address. + || octets[0] == 0 + || (octets[0] == 100 && (64..=127).contains(&octets[1])) + || (octets[0] == 192 && octets[1] == 0 && octets[2] == 0) + || (octets[0] == 192 && octets[1] == 88 && octets[2] == 99) + || (octets[0] == 198 && (18..=19).contains(&octets[1])) + || octets[0] >= 240 + } + IpAddr::V6(ip) => { + let segments = ip.segments(); + if let Some(mapped) = ip.to_ipv4_mapped() { + return is_private_or_reserved_ip(IpAddr::V4(mapped)); + } + ip.is_loopback() + || ip.is_unspecified() + || ip.is_unique_local() + || ip.is_unicast_link_local() + || ip.is_multicast() + || (segments[0] & 0xffc0 == 0xfec0) + || (segments[0] == 0x2001 && segments[1] == 0x0db8) + || segments[..6] == [0x0064, 0xff9b, 0, 0, 0, 0] + || segments[..3] == [0x0064, 0xff9b, 0x0001] + || segments[0] == 0x2002 + || segments[..2] == [0x2001, 0] + || segments[..6] == [0, 0, 0, 0, 0, 0] + || segments[..6] == [0, 0, 0, 0, 0xffff, 0] + || (matches!(segments[4], 0 | 0x0200) && segments[5] == 0x5efe) + } + } +} + +/// Return whether an address belongs to RFC 2544's IPv4 benchmarking range. +/// +/// Local DNS interception tools commonly synthesize answers from +/// `198.18.0.0/15`. The range is intentionally still +/// classified as reserved by [`is_private_or_reserved_ip`]; callers may use +/// this predicate only when they have independently established that the +/// hostname is a trusted, fixed destination. Keeping the predicates separate +/// prevents a compatibility exception from weakening the generic SSRF guard. +pub fn is_ipv4_benchmarking_fake_ip(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => { + let octets = ip.octets(); + octets[0] == 198 && (18..=19).contains(&octets[1]) + } + IpAddr::V6(_) => false, + } +} + +/// Return true only when a URL names loopback without relying on DNS. +/// +/// This is intentionally stricter than accepting names that currently resolve +/// to loopback: DNS answers can change between validation and connection. +pub fn url_has_literal_loopback_host(url: &Url) -> bool { + match url.host() { + Some(Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + None => false, + } +} + +/// Sensitive HTTP traffic may use cleartext only for a literal loopback host. +pub fn is_https_or_loopback_http_url(url: &Url) -> bool { + url.scheme() == "https" || (url.scheme() == "http" && url_has_literal_loopback_host(url)) +} + +#[cfg(test)] +mod tests { + use super::{ + connection_declared_header_names, is_https_or_loopback_http_url, + is_ipv4_benchmarking_fake_ip, is_private_or_reserved_ip, url_has_literal_loopback_host, + }; + + #[test] + fn parses_multiple_connection_values_and_rejects_invalid_names() { + let names = connection_declared_header_names([ + "keep-alive, X-Private", + "x-accel-redirect, invalid name, x-private", + ]); + + assert_eq!( + names.into_iter().collect::>(), + vec!["keep-alive", "x-accel-redirect", "x-private"] + ); + } + + #[test] + fn blocks_private_and_transition_addresses_but_allows_public_addresses() { + for address in [ + "127.0.0.1", + "0.1.2.3", + "169.254.169.254", + "100.64.0.1", + "::1", + "::ffff:127.0.0.1", + "64:ff9b::10.0.0.1", + "2002:0a00:0001::1", + "2001:0000:4136:e378:8000:63bf:3fff:fdd2", + ] { + assert!( + is_private_or_reserved_ip(address.parse().expect("IP address")), + "address should be blocked: {address}" + ); + } + assert!(!is_private_or_reserved_ip("8.8.8.8".parse().unwrap())); + assert!(!is_private_or_reserved_ip( + "2606:4700:4700::1111".parse().unwrap() + )); + } + + #[test] + fn benchmarking_fake_ip_predicate_is_narrow_and_does_not_change_private_policy() { + for address in ["198.18.0.1", "198.19.255.254"] { + let ip = address.parse().expect("benchmarking address"); + assert!(is_ipv4_benchmarking_fake_ip(ip)); + assert!(is_private_or_reserved_ip(ip)); + } + for address in ["198.17.255.254", "198.20.0.1", "2001:db8::1"] { + let ip = address.parse().expect("non-benchmarking address"); + assert!(!is_ipv4_benchmarking_fake_ip(ip)); + } + } + + #[test] + fn sensitive_http_transport_allows_https_or_literal_loopback_only() { + for allowed in [ + "https://api.example.test/v1", + "http://localhost:8080/v1", + "http://127.42.0.1:8080/v1", + "http://[::1]:8080/v1", + ] { + let url = url::Url::parse(allowed).unwrap(); + assert!(is_https_or_loopback_http_url(&url), "rejected {allowed}"); + } + + for rejected in [ + "http://api.example.test/v1", + "http://10.0.0.1/v1", + "http://0.0.0.0:8080/v1", + "http://[::ffff:127.0.0.1]:8080/v1", + "ftp://localhost/resource", + ] { + let url = url::Url::parse(rejected).unwrap(); + assert!(!is_https_or_loopback_http_url(&url), "accepted {rejected}"); + } + + assert!(url_has_literal_loopback_host( + &url::Url::parse("https://localhost/").unwrap() + )); + assert!(!url_has_literal_loopback_host( + &url::Url::parse("https://localhost.example/").unwrap() + )); + } +} diff --git a/crates/aether-http/src/lib.rs b/crates/aether-http/src/lib.rs index 1a17203c1..d0aeaacd2 100644 --- a/crates/aether-http/src/lib.rs +++ b/crates/aether-http/src/lib.rs @@ -1,7 +1,16 @@ mod client; mod config; +mod dns; +mod header_security; +mod response_body; mod retry; pub use client::{apply_http_client_config, build_http_client, build_http_client_with_headers}; pub use config::{HttpClientConfig, HttpRetryConfig}; +pub use dns::{lookup_host_with_limits, DEFAULT_DNS_LOOKUP_TIMEOUT, MAX_DNS_RESOLVED_ADDRESSES}; +pub use header_security::{ + connection_declared_header_names, is_https_or_loopback_http_url, is_ipv4_benchmarking_fake_ip, + is_private_or_reserved_ip, url_has_literal_loopback_host, +}; +pub use response_body::{read_response_bytes_with_limit, ResponseBodyReadError}; pub use retry::jittered_delay_for_retry; diff --git a/crates/aether-http/src/response_body.rs b/crates/aether-http/src/response_body.rs new file mode 100644 index 000000000..907e15dec --- /dev/null +++ b/crates/aether-http/src/response_body.rs @@ -0,0 +1,146 @@ +use std::error::Error; +use std::fmt; + +pub enum ResponseBodyReadError { + TooLarge { max_bytes: usize }, + Read(reqwest::Error), +} + +// Avoid trusting a remote Content-Length as an allocation hint. The stream +// remains allowed to grow up to the caller's actual body limit, but the first +// allocation stays modest when a peer advertises a very large response. +const MAX_INITIAL_RESPONSE_BODY_CAPACITY_BYTES: usize = 16 * 1024 * 1024; + +impl fmt::Debug for ResponseBodyReadError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = formatter.debug_struct("ResponseBodyReadError"); + match self { + Self::TooLarge { max_bytes } => { + debug + .field("kind", &"too_large") + .field("max_bytes", max_bytes); + } + // Reqwest's Debug output may include the complete request URL, + // including credentials embedded in a path or query. Keep the + // underlying value available to explicit category helpers, but + // never render it through this public error boundary. + Self::Read(_) => { + debug.field("kind", &"read"); + } + } + debug.finish() + } +} + +impl fmt::Display for ResponseBodyReadError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::TooLarge { max_bytes } => { + write!(formatter, "response body exceeds {max_bytes} bytes") + } + Self::Read(_) => write!(formatter, "failed to read response body"), + } + } +} + +impl Error for ResponseBodyReadError { + // Do not expose the reqwest error chain to generic reporters. Callers that + // need a retry/telemetry category can still pattern-match `Read(error)` + // and inspect the concrete reqwest value deliberately. + fn source(&self) -> Option<&(dyn Error + 'static)> { + None + } +} + +/// Read a small control-plane response without trusting `Content-Length`. +/// +/// The advertised length is rejected early when available, while the streamed +/// byte count remains authoritative for missing or dishonest length headers. +pub async fn read_response_bytes_with_limit( + mut response: reqwest::Response, + max_bytes: usize, +) -> Result, ResponseBodyReadError> { + if response + .content_length() + .is_some_and(|length| length > max_bytes as u64) + { + return Err(ResponseBodyReadError::TooLarge { max_bytes }); + } + + let initial_capacity = initial_response_body_capacity(response.content_length(), max_bytes); + let mut body = Vec::with_capacity(initial_capacity); + while let Some(chunk) = response + .chunk() + .await + .map_err(ResponseBodyReadError::Read)? + { + append_chunk_with_limit(&mut body, &chunk, max_bytes)?; + } + Ok(body) +} + +fn initial_response_body_capacity(content_length: Option, max_bytes: usize) -> usize { + content_length + .and_then(|length| usize::try_from(length).ok()) + .unwrap_or(0) + .min(max_bytes) + .min(MAX_INITIAL_RESPONSE_BODY_CAPACITY_BYTES) +} + +fn append_chunk_with_limit( + body: &mut Vec, + chunk: &[u8], + max_bytes: usize, +) -> Result<(), ResponseBodyReadError> { + let Some(next_len) = body.len().checked_add(chunk.len()) else { + return Err(ResponseBodyReadError::TooLarge { max_bytes }); + }; + if next_len > max_bytes { + return Err(ResponseBodyReadError::TooLarge { max_bytes }); + } + body.extend_from_slice(chunk); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::{append_chunk_with_limit, initial_response_body_capacity, ResponseBodyReadError}; + use std::error::Error; + + #[test] + fn streamed_body_accepts_exact_limit() { + let mut body = b"1234".to_vec(); + append_chunk_with_limit(&mut body, b"5678", 8).expect("exact limit should pass"); + assert_eq!(body, b"12345678"); + } + + #[test] + fn streamed_body_rejects_limit_plus_one_without_appending_chunk() { + let mut body = b"1234".to_vec(); + let error = append_chunk_with_limit(&mut body, b"56789", 8) + .expect_err("limit plus one should fail"); + assert!(matches!( + error, + ResponseBodyReadError::TooLarge { max_bytes: 8 } + )); + assert_eq!(body, b"1234"); + } + + #[test] + fn public_error_rendering_does_not_include_read_error_details() { + let error = ResponseBodyReadError::TooLarge { max_bytes: 64 }; + assert_eq!(error.to_string(), "response body exceeds 64 bytes"); + assert!(error.source().is_none()); + assert!(format!("{error:?}").contains("too_large")); + } + + #[test] + fn initial_capacity_does_not_trust_giant_content_length() { + assert_eq!( + initial_response_body_capacity(Some(u64::MAX), usize::MAX), + 16 * 1024 * 1024 + ); + assert_eq!(initial_response_body_capacity(Some(1024), usize::MAX), 1024); + assert_eq!(initial_response_body_capacity(None, usize::MAX), 0); + } +} diff --git a/crates/aether-model-fetch/Cargo.toml b/crates/aether-model-fetch/Cargo.toml index 1ee7e6cbc..a611f4f5a 100644 --- a/crates/aether-model-fetch/Cargo.toml +++ b/crates/aether-model-fetch/Cargo.toml @@ -8,16 +8,17 @@ repository.workspace = true [dependencies] aether-ai-formats.workspace = true aether-contracts.workspace = true +aether-crypto.workspace = true aether-data-contracts.workspace = true aether-provider-transport.workspace = true aether-scheduler-core.workspace = true async-trait.workspace = true base64.workspace = true regex.workspace = true -rsa = "0.9.10" serde_json.workspace = true -sha2 = { workspace = true, features = ["oid"] } +url.workspace = true uuid.workspace = true [dev-dependencies] +aws-lc-rs.workspace = true tokio.workspace = true diff --git a/crates/aether-model-fetch/src/lib.rs b/crates/aether-model-fetch/src/lib.rs index bb3432805..3dbebf9a9 100644 --- a/crates/aether-model-fetch/src/lib.rs +++ b/crates/aether-model-fetch/src/lib.rs @@ -21,7 +21,8 @@ pub use logic::{ upstream_metadata_namespace_updates, ModelFetchRunSummary, ModelsFetchPage, ModelsFetchSuccess, }; pub use strategy::{ - fetch_models_from_transports, fetch_models_from_transports_for_client_version, + antigravity_model_id_is_routable, fetch_models_from_transports, + fetch_models_from_transports_for_client_version, fetch_models_from_transports_for_management, ModelFetchStrategy, ModelFetchStrategyKind, ModelsFetchOutcome, SelectedModelFetchStrategy, }; pub use transport::{ diff --git a/crates/aether-model-fetch/src/logic.rs b/crates/aether-model-fetch/src/logic.rs index b4c330257..aec75836a 100644 --- a/crates/aether-model-fetch/src/logic.rs +++ b/crates/aether-model-fetch/src/logic.rs @@ -624,6 +624,15 @@ pub fn merge_upstream_metadata(current: Option<&Value>, incoming: &Value) -> Val next_value.as_object_mut(), merged.get(namespace).and_then(Value::as_object), ) { + if namespace.eq_ignore_ascii_case("antigravity") { + for field in ["quota_groups", "quota_groups_updated_at"] { + if !next_namespace.contains_key(field) { + if let Some(value) = old_namespace.get(field) { + next_namespace.insert(field.to_string(), value.clone()); + } + } + } + } if let (Some(new_quota), Some(old_quota)) = ( next_namespace .get_mut("quota_by_model") @@ -964,7 +973,7 @@ fn build_codex_models_url(base_url: &str, client_version: Option<&str>) -> Optio .trim() .eq_ignore_ascii_case("client_version") }); - query_parts.push(format!("client_version={client_version}")); + query_parts.push(encoded_query_pair("client_version", client_version)); } else if !has_client_version { query_parts.push(format!( "client_version={}", @@ -992,10 +1001,16 @@ fn replace_or_append_query_param(url: &str, name: &str, value: &str) -> String { .trim() .eq_ignore_ascii_case(name) }); - query_parts.push(format!("{name}={value}")); + query_parts.push(encoded_query_pair(name, value)); format!("{base}?{}", query_parts.join("&")) } +fn encoded_query_pair(name: &str, value: &str) -> String { + let name = url::form_urlencoded::byte_serialize(name.as_bytes()).collect::(); + let value = url::form_urlencoded::byte_serialize(value.as_bytes()).collect::(); + format!("{name}={value}") +} + fn build_gemini_models_url(base_url: &str) -> Option { let (trimmed_base_url, base_query) = split_url_query(base_url); let trimmed_base_url = trimmed_base_url.trim_end_matches('/'); @@ -1330,7 +1345,7 @@ mod tests { "https://chatgpt.com/backend-api/codex" ), Some(( - "https://chatgpt.com/backend-api/codex/models?client_version=0.144.1".to_string(), + "https://chatgpt.com/backend-api/codex/models?client_version=0.153.3".to_string(), "openai:responses".to_string() )) ); @@ -1352,6 +1367,22 @@ mod tests { ); } + #[test] + fn explicit_codex_client_version_cannot_inject_query_parameters() { + let (url, _) = build_models_fetch_url_for_client_version( + "codex", + "openai:responses", + "https://chatgpt.com/backend-api/codex", + Some("0.145.2&admin=true#fragment"), + ) + .expect("models URL should build"); + + assert_eq!( + url, + "https://chatgpt.com/backend-api/codex/models?client_version=0.145.2%26admin%3Dtrue%23fragment" + ); + } + #[test] fn explicit_codex_client_version_replaces_stale_base_query_value() { assert_eq!( @@ -1369,6 +1400,23 @@ mod tests { ); } + #[test] + fn explicit_codex_client_version_preserves_preencoded_base_query_values() { + assert_eq!( + build_models_fetch_url_for_client_version( + "codex", + "openai:responses", + "https://chatgpt.com/backend-api/codex?feature=beta%2Bdesktop", + Some("0.145.2"), + ), + Some(( + "https://chatgpt.com/backend-api/codex/models?feature=beta%2Bdesktop&client_version=0.145.2" + .to_string(), + "openai:responses".to_string() + )) + ); + } + #[test] fn explicit_codex_client_version_is_forwarded_through_compatible_proxy_roots() { assert_eq!( @@ -1738,6 +1786,11 @@ mod tests { let merged = merge_upstream_metadata( Some(&json!({ "antigravity": { + "quota_groups": [{ + "display_name": "Claude and GPT models", + "buckets": [{"bucket_id": "3p-5h", "window": "5h"}] + }], + "quota_groups_updated_at": 1_777_000_000u64, "quota_by_model": { "gemini-2.5-pro": { "remaining_fraction": 0.3, @@ -1767,6 +1820,14 @@ mod tests { assert!(merged["antigravity"]["quota_by_model"] .get("stale-model") .is_none()); + assert_eq!( + merged["antigravity"]["quota_groups"][0]["buckets"][0]["bucket_id"], + "3p-5h" + ); + assert_eq!( + merged["antigravity"]["quota_groups_updated_at"], + json!(1_777_000_000u64) + ); } #[test] diff --git a/crates/aether-model-fetch/src/strategy.rs b/crates/aether-model-fetch/src/strategy.rs index 1412f7331..e137dac3e 100644 --- a/crates/aether-model-fetch/src/strategy.rs +++ b/crates/aether-model-fetch/src/strategy.rs @@ -1,23 +1,23 @@ use std::collections::{BTreeMap, BTreeSet}; use std::time::{SystemTime, UNIX_EPOCH}; -use aether_contracts::{ExecutionPlan, ExecutionResult, RequestBody}; +use aether_contracts::{ + ExecutionError, ExecutionErrorKind, ExecutionPlan, ExecutionResult, RequestBody, +}; +use aether_crypto::{rsa_pkcs1_sha256_sign, RsaPkcs1Sha256Error}; use aether_provider_transport::antigravity::{ resolve_local_antigravity_request_auth, AntigravityRequestAuthSupport, }; +use aether_provider_transport::vertex::{ + looks_like_vertex_ai_host, parse_vertex_service_account_auth_config, +}; use aether_provider_transport::{ is_vertex_api_key_transport_context, resolve_transport_execution_timeouts, resolve_transport_profile, GatewayProviderTransportSnapshot, }; use base64::engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}; use base64::Engine as _; -use rsa::pkcs1::DecodeRsaPrivateKey; -use rsa::pkcs1v15::SigningKey; -use rsa::pkcs8::DecodePrivateKey; -use rsa::signature::{SignatureEncoding, Signer}; -use rsa::RsaPrivateKey; use serde_json::{json, Value}; -use sha2::Sha256; use crate::logic::{ aggregate_models_for_cache, codex_model_identity, extract_error_message, @@ -41,7 +41,6 @@ const VERTEX_API_BASE_URL: &str = "https://aiplatform.googleapis.com"; const VERTEX_MODEL_GARDEN_API_VERSION: &str = "v1beta1"; const VERTEX_PAGE_SIZE: &str = "100"; const VERTEX_MAX_PAGES: usize = 20; -const GOOGLE_OAUTH_TOKEN_URL: &str = "https://oauth2.googleapis.com/token"; const GOOGLE_CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform"; #[derive(Debug, Clone, PartialEq)] @@ -53,6 +52,8 @@ pub struct ModelsFetchOutcome { pub legacy_models: Vec, pub errors: Vec, pub has_success: bool, + /// Only native `models` responses may populate the opaque Codex client catalog. + pub native_codex_catalog: bool, pub upstream_metadata: Option, pub etag: Option, pub upstream_status: Option, @@ -142,8 +143,31 @@ pub async fn fetch_models_from_transports_for_client_version( transports: &[GatewayProviderTransportSnapshot], codex_client_version: Option<&str>, ) -> Result { - let strategy = select_model_fetch_strategy(transports)?; - execute_model_fetch_strategy(runtime, transports, strategy, codex_client_version).await + let strategy = select_model_fetch_strategy(transports) + .map_err(|error| sanitize_model_fetch_error(&error))?; + execute_model_fetch_strategy( + runtime, + transports, + strategy, + codex_client_version, + codex_client_version.is_none(), + ) + .await + .map_err(|error| sanitize_model_fetch_error(&error)) +} + +/// Management also supports Codex-compatible proxies returning OpenAI `data` arrays, +/// independently of the fingerprint sent upstream. Public client catalogs stay strict. +pub async fn fetch_models_from_transports_for_management( + runtime: &(impl ModelFetchTransportRuntime + ?Sized), + transports: &[GatewayProviderTransportSnapshot], + codex_client_version: Option<&str>, +) -> Result { + let strategy = select_model_fetch_strategy(transports) + .map_err(|error| sanitize_model_fetch_error(&error))?; + execute_model_fetch_strategy(runtime, transports, strategy, codex_client_version, true) + .await + .map_err(|error| sanitize_model_fetch_error(&error)) } fn select_model_fetch_strategy( @@ -213,6 +237,7 @@ async fn execute_model_fetch_strategy( transports: &[GatewayProviderTransportSnapshot], strategy: SelectedModelFetchStrategy, codex_client_version: Option<&str>, + allow_codex_legacy_response: bool, ) -> Result { let Some(first_transport) = transports.first() else { return Err("No transport snapshots available for models fetch".to_string()); @@ -230,6 +255,7 @@ async fn execute_model_fetch_strategy( transports, strategy.provider_id(), codex_client_version, + allow_codex_legacy_response, ) .await } @@ -255,6 +281,7 @@ async fn fetch_standard_models( transports: &[GatewayProviderTransportSnapshot], provider_type: &str, codex_client_version: Option<&str>, + allow_codex_legacy_response: bool, ) -> Result { let mut all_models = Vec::new(); let mut successful_codex_catalogs = Vec::<(String, Vec)>::new(); @@ -263,10 +290,19 @@ async fn fetch_standard_models( let mut etag = ConsistentValue::default(); let mut upstream_status = ConsistentValue::default(); let is_codex = provider_type.trim().eq_ignore_ascii_case("codex"); + let mut native_codex_catalog = is_codex; for transport in transports { - match fetch_standard_models_for_transport(runtime, transport, codex_client_version).await { + match fetch_standard_models_for_transport( + runtime, + transport, + codex_client_version, + allow_codex_legacy_response, + ) + .await + { Ok(outcome) => { + native_codex_catalog &= outcome.native_codex_catalog; all_models.extend(outcome.cached_models.iter().cloned()); if is_codex && outcome.has_success { successful_codex_catalogs @@ -280,7 +316,11 @@ async fn fetch_standard_models( } Err((err, status)) => { upstream_status.observe(status); - errors.push(format!("{}: {err}", transport.endpoint.api_format.trim())); + let format_label = model_fetch_format_label(&transport.endpoint.api_format); + errors.push(format!( + "{format_label}: {}", + sanitize_model_fetch_error(&err) + )); } } } @@ -294,6 +334,7 @@ async fn fetch_standard_models( let upstream_metadata = crate::logic::model_catalog_upstream_metadata(provider_type, &merged_models); let mut outcome = build_success_outcome(merged_models, upstream_metadata, has_success); + outcome.native_codex_catalog = native_codex_catalog && has_success; if let Some(model_ids) = codex_model_ids { outcome.fetched_model_ids = model_ids; outcome.legacy_models = project_codex_models_for_legacy_cache( @@ -312,6 +353,7 @@ async fn fetch_standard_models_for_transport( runtime: &(impl ModelFetchTransportRuntime + ?Sized), transport: &GatewayProviderTransportSnapshot, codex_client_version: Option<&str>, + allow_codex_legacy_response: bool, ) -> Result)> { let mut all_models = Vec::new(); let mut seen_ids = BTreeSet::new(); @@ -324,6 +366,7 @@ async fn fetch_standard_models_for_transport( .provider_type .trim() .eq_ignore_ascii_case("codex"); + let mut native_codex_catalog = is_codex; for _ in 0..20 { let plan = build_standard_models_fetch_execution_plan_for_client_version( @@ -341,11 +384,16 @@ async fn fetch_standard_models_for_transport( upstream_status.observe(Some(result.status_code)); let body_json = execution_result_json_body(&result).map_err(|err| (err, Some(result.status_code)))?; + native_codex_catalog &= body_json.get("models").and_then(Value::as_array).is_some(); let parsed = if is_codex { parse_codex_models_response_for_request( &transport.endpoint.api_format, &body_json, - codex_client_version, + if allow_codex_legacy_response { + None + } else { + codex_client_version + }, ) } else { parse_models_response_page(&transport.endpoint.api_format, &body_json) @@ -385,7 +433,9 @@ async fn fetch_standard_models_for_transport( next_after_id = Some(next_cursor); } - Ok(build_success_outcome(all_models, None, has_success) + let mut outcome = build_success_outcome(all_models, None, has_success); + outcome.native_codex_catalog = native_codex_catalog && has_success; + Ok(outcome .with_etag(etag.finish()) .with_upstream_status(upstream_status.finish())) } @@ -425,13 +475,16 @@ async fn fetch_antigravity_models( .await { Ok(plan) => plan, - Err(err) => return Err(err), + Err(err) => return Err(sanitize_model_fetch_error(&err)), }; let result = match runtime.execute_model_fetch_execution_plan(&plan).await { Ok(result) => result, Err(err) => { - errors.push(format!("{base_url}: {err}")); + errors.push(format!( + "antigravity models fetch failed: {}", + sanitize_model_fetch_error(&err) + )); continue; } }; @@ -447,10 +500,10 @@ async fn fetch_antigravity_models( let error = execution_result_error_message(&result); if should_fallback_antigravity_status(result.status_code) { - errors.push(format!("{base_url}: {error}")); + errors.push(format!("antigravity models fetch failed: {error}")); continue; } - return Err(error); + return Err(sanitize_model_fetch_error(&error)); } Ok(ModelsFetchOutcome { @@ -459,6 +512,7 @@ async fn fetch_antigravity_models( legacy_models: Vec::new(), errors, has_success: false, + native_codex_catalog: false, upstream_metadata: None, etag: None, upstream_status: None, @@ -474,8 +528,13 @@ async fn resolve_or_hydrate_antigravity_project( return Ok((project_id, transport.clone(), metadata)); } - let plan = build_antigravity_load_code_assist_plan(runtime, transport).await?; - let result = runtime.execute_model_fetch_execution_plan(&plan).await?; + let plan = build_antigravity_load_code_assist_plan(runtime, transport) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; + let result = runtime + .execute_model_fetch_execution_plan(&plan) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; if !(200..300).contains(&result.status_code) { return Err(format!( "antigravity: loadCodeAssist failed: {}", @@ -581,8 +640,13 @@ async fn fetch_kiro_models( runtime: &(impl ModelFetchTransportRuntime + ?Sized), transport: &GatewayProviderTransportSnapshot, ) -> Result { - let plan = build_kiro_list_available_models_plan(runtime, transport).await?; - let result = runtime.execute_model_fetch_execution_plan(&plan).await?; + let plan = build_kiro_list_available_models_plan(runtime, transport) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; + let result = runtime + .execute_model_fetch_execution_plan(&plan) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; if !(200..300).contains(&result.status_code) { return Err(execution_result_error_message(&result)); } @@ -596,8 +660,13 @@ async fn fetch_windsurf_models( runtime: &(impl ModelFetchTransportRuntime + ?Sized), transport: &GatewayProviderTransportSnapshot, ) -> Result { - let plan = build_windsurf_model_configs_execution_plan(runtime, transport).await?; - let result = runtime.execute_model_fetch_execution_plan(&plan).await?; + let plan = build_windsurf_model_configs_execution_plan(runtime, transport) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; + let result = runtime + .execute_model_fetch_execution_plan(&plan) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; if !(200..300).contains(&result.status_code) { return Err(execution_result_error_message(&result)); } @@ -642,6 +711,7 @@ async fn fetch_vertex_api_key_models( legacy_models: Vec::new(), errors: vec!["vertex_ai(api_key): missing api key".to_string()], has_success: false, + native_codex_catalog: false, upstream_metadata: None, etag: None, upstream_status: None, @@ -653,7 +723,11 @@ async fn fetch_vertex_api_key_models( let mut soft_errors = Vec::new(); let mut has_success = false; - for base_url in iter_vertex_base_urls(transports) { + // The API key is a bearer-like cloud credential. Endpoint records can be + // imported or edited by administrators, so never send it to an arbitrary + // custom host merely because it is listed alongside a Vertex transport. + // Keep only the canonical Vertex host and its official regional variants. + for base_url in iter_trusted_vertex_base_urls(transports) { let url = build_vertex_google_list_url(&base_url, api_key, None); let outcome = match fetch_vertex_models_from_url( runtime, @@ -668,16 +742,20 @@ async fn fetch_vertex_api_key_models( { Ok(outcome) => outcome, Err(err) => { - hard_errors.push(format!("{base_url}: {err}")); + hard_errors.push(format!( + "vertex google models fetch failed: {}", + sanitize_model_fetch_error(&err) + )); continue; } }; has_success |= outcome.has_success; if let Some(error) = outcome.error { + let error = sanitize_model_fetch_error(&error); if is_soft_not_found(&error) { - soft_errors.push(format!("{base_url}: {error}")); + soft_errors.push(format!("vertex google models fetch failed: {error}")); } else { - hard_errors.push(format!("{base_url}: {error}")); + hard_errors.push(format!("vertex google models fetch failed: {error}")); } continue; } @@ -702,6 +780,7 @@ async fn fetch_vertex_api_key_models( legacy_models: Vec::new(), errors, has_success, + native_codex_catalog: false, upstream_metadata: None, etag: None, upstream_status: None, @@ -720,12 +799,15 @@ async fn fetch_vertex_service_account_models( legacy_models: Vec::new(), errors: vec!["vertex_ai(service_account): missing auth_config".to_string()], has_success: false, + native_codex_catalog: false, upstream_metadata: None, etag: None, upstream_status: None, }); }; - let token = exchange_vertex_service_account_token(runtime, &transports[0], auth_config).await?; + let token = exchange_vertex_service_account_token(runtime, &transports[0], auth_config) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; let gemini_transport = select_transport_for_api_format(transports, "gemini:").unwrap_or(&transports[0]); let claude_transport = @@ -736,7 +818,10 @@ async fn fetch_vertex_service_account_models( let mut soft_errors = Vec::new(); let mut has_success = false; - for base in iter_vertex_base_urls(transports) { + // A service-account access token is a cloud credential. Restrict its + // model-garden requests to official Vertex hosts even when an endpoint + // override exists under the provider. + for base in iter_trusted_vertex_base_urls(transports) { for (publisher, transport, api_format) in [ ("google", gemini_transport, "gemini:generate_content"), ("anthropic", claude_transport, "claude:messages"), @@ -755,13 +840,17 @@ async fn fetch_vertex_service_account_models( { Ok(outcome) => outcome, Err(err) => { - hard_errors.push(format!("{url}: {err}")); + hard_errors.push(format!( + "vertex {publisher} models fetch failed: {}", + sanitize_model_fetch_error(&err) + )); continue; } }; has_success |= outcome.has_success; if let Some(error) = outcome.error { - let labeled = format!("{url}: {error}"); + let error = sanitize_model_fetch_error(&error); + let labeled = format!("vertex {publisher} models fetch failed: {error}"); if is_soft_not_found(&error) { soft_errors.push(labeled); } else { @@ -791,6 +880,7 @@ async fn fetch_vertex_service_account_models( legacy_models: Vec::new(), errors, has_success, + native_codex_catalog: false, upstream_metadata: None, etag: None, upstream_status: None, @@ -829,8 +919,12 @@ async fn fetch_vertex_models_from_url( api_format, auth_header.clone(), ) - .await?; - let result = runtime.execute_model_fetch_execution_plan(&plan).await?; + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; + let result = runtime + .execute_model_fetch_execution_plan(&plan) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; if result.status_code != 200 { return Ok(VertexFetchPageOutcome { models: Vec::new(), @@ -840,7 +934,8 @@ async fn fetch_vertex_models_from_url( } has_success = true; - let body_json = execution_result_json_body_allow_empty(&result)?; + let body_json = execution_result_json_body_allow_empty(&result) + .map_err(|error| sanitize_model_fetch_error(&error))?; all_models.extend(parse_vertex_models_payload( &body_json, auth_config, @@ -869,15 +964,21 @@ async fn exchange_vertex_service_account_token( transport: &GatewayProviderTransportSnapshot, auth_config: &Value, ) -> Result { - let token_url = json_string(auth_config.get("token_uri")) - .unwrap_or_else(|| GOOGLE_OAUTH_TOKEN_URL.to_string()); - let client_email = json_string(auth_config.get("client_email")) - .ok_or_else(|| "vertex_ai(service_account): missing client_email".to_string())?; - let private_key = json_string(auth_config.get("private_key")) - .ok_or_else(|| "vertex_ai(service_account): missing private_key".to_string())?; + // Reuse the provider transport parser here instead of trusting token_uri + // from the raw credential JSON. The parser pins the token endpoint to + // Google's HTTPS OAuth endpoint and rejects credentials, ports, queries, + // fragments, and lookalike hosts before a signed assertion is produced. + let auth_config_json = serde_json::to_string(auth_config) + .map_err(|_| "vertex_ai(service_account): invalid auth_config".to_string())?; + let auth_config = parse_vertex_service_account_auth_config(Some(&auth_config_json)) + .ok_or_else(|| "vertex_ai(service_account): invalid auth_config".to_string())?; + let token_url = auth_config.token_uri; + let client_email = auth_config.client_email; + let private_key = auth_config.private_key; let now = now_unix_secs(); let assertion = - build_vertex_service_account_assertion(&client_email, &private_key, &token_url, now)?; + build_vertex_service_account_assertion(&client_email, &private_key, &token_url, now) + .map_err(|error| sanitize_model_fetch_error(&error))?; let body = format!( "grant_type=urn%3Aietf%3Aparams%3Aoauth%3Agrant-type%3Ajwt-bearer&assertion={assertion}" ); @@ -911,7 +1012,10 @@ async fn exchange_vertex_service_account_token( transport_profile, timeouts: resolve_transport_execution_timeouts(transport), }; - let result = runtime.execute_model_fetch_execution_plan(&plan).await?; + let result = runtime + .execute_model_fetch_execution_plan(&plan) + .await + .map_err(|error| sanitize_model_fetch_error(&error))?; let body_json = execution_result_json_body(&result)?; body_json .get("access_token") @@ -940,26 +1044,16 @@ fn build_vertex_service_account_assertion( .map_err(|err| format!("vertex_ai(service_account): jwt payload encode failed: {err}"))?, ); let message = format!("{header}.{payload}"); - let private_key = decode_vertex_service_account_private_key(private_key_pem)?; - let signing_key = SigningKey::::new(private_key); - let signature = signing_key.sign(message.as_bytes()); - Ok(format!( - "{message}.{}", - URL_SAFE_NO_PAD.encode(signature.to_bytes()) - )) -} - -fn decode_vertex_service_account_private_key( - private_key_pem: &str, -) -> Result { - match RsaPrivateKey::from_pkcs8_pem(private_key_pem) { - Ok(private_key) => Ok(private_key), - Err(pkcs8_err) => RsaPrivateKey::from_pkcs1_pem(private_key_pem).map_err(|pkcs1_err| { - format!( - "vertex_ai(service_account): private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}" - ) - }), - } + let signature = + rsa_pkcs1_sha256_sign(private_key_pem.as_bytes(), message.as_bytes()).map_err(|error| { + match error { + RsaPkcs1Sha256Error::InvalidPrivateKey => { + "vertex_ai(service_account): private_key parse failed".to_string() + } + _ => "vertex_ai(service_account): signing failed".to_string(), + } + })?; + Ok(format!("{message}.{}", URL_SAFE_NO_PAD.encode(signature))) } fn execution_result_json_body(result: &ExecutionResult) -> Result { @@ -987,6 +1081,8 @@ fn execution_result_header(result: &ExecutionResult, name: &str) -> Option String { let detail = result .body @@ -999,12 +1095,293 @@ fn execution_result_error_message(result: &ExecutionResult) -> String { (!message.is_empty()).then_some(message.to_string()) }) }); - match detail { - Some(detail) if !(200..300).contains(&result.status_code) => { - format!("HTTP {}: {detail}", result.status_code) + let status = if !(200..300).contains(&result.status_code) { + Some(result.status_code) + } else { + result + .error + .as_ref() + .and_then(|error| error.upstream_status) + .filter(|status| (400..600).contains(status)) + }; + let summary = model_fetch_error_summary(detail.as_deref(), result.error.as_ref(), status); + match status { + Some(status) => format!("HTTP {status}: {summary}"), + None if detail.is_some() => summary, + None => format!("HTTP {}: {summary}", result.status_code), + } +} + +/// Projects transport/upstream diagnostics into a bounded message that can cross the +/// model-fetch API boundary. Upstream error bodies and HTTP client errors are untrusted: they +/// commonly contain authorization headers, credential-bearing URLs, or local file paths. +fn sanitize_model_fetch_error(error: &str) -> String { + let trimmed = error.trim(); + if trimmed.is_empty() { + return "upstream request failed".to_string(); + } + match trimmed { + "No transport snapshots available for models fetch" + | "No supported endpoint for Rust models fetch" + | "Provider transport snapshot unavailable" => return trimmed.to_string(), + _ => {} + } + + let status = model_fetch_status_from_text(trimmed); + if let Some(detail) = sanitize_model_fetch_error_detail(trimmed) { + if let Some(status) = status { + let lower = trimmed.to_ascii_lowercase(); + if lower.starts_with("http ") + || lower.starts_with("status ") + || lower.starts_with("status=") + || lower.starts_with("status:") + || lower.starts_with("status_code=") + || lower.starts_with("status_code:") + { + return format!("HTTP {status}: {detail}"); + } } - Some(detail) => detail, - None => format!("HTTP {}: upstream request failed", result.status_code), + return detail; + } + + let summary = model_fetch_error_summary(Some(trimmed), None, status); + status + .map(|status| format!("HTTP {status}: {summary}")) + .unwrap_or_else(|| summary.to_string()) +} + +fn model_fetch_error_summary( + detail: Option<&str>, + execution_error: Option<&ExecutionError>, + status: Option, +) -> String { + // Authentication, authorization, not-found, timeout, and rate-limit statuses are stable + // public classifications. Never let an upstream body replace them with free-form text. + if matches!(status, Some(401 | 403 | 404 | 408 | 429)) { + return model_fetch_error_category(detail, execution_error, status).to_string(); + } + if let Some(detail) = detail.and_then(sanitize_model_fetch_error_detail) { + return detail; + } + model_fetch_error_category(detail, execution_error, status).to_string() +} + +fn model_fetch_error_category( + detail: Option<&str>, + execution_error: Option<&ExecutionError>, + status: Option, +) -> &'static str { + let lower = detail.unwrap_or_default().to_ascii_lowercase(); + + match status { + Some(401) => return "upstream authentication failed", + Some(403) => return "upstream authorization failed", + Some(404) => return "upstream endpoint not found", + Some(408) => return "upstream request timed out", + Some(429) => return "upstream rate limited", + _ => {} + } + + if let Some(error) = execution_error { + match &error.kind { + ExecutionErrorKind::ConnectTimeout + | ExecutionErrorKind::FirstByteTimeout + | ExecutionErrorKind::ReadTimeout => return "upstream request timed out", + ExecutionErrorKind::TlsError => return "upstream TLS connection failed", + ExecutionErrorKind::ProxyError => return "upstream proxy request failed", + ExecutionErrorKind::ProtocolError => return "upstream response invalid", + ExecutionErrorKind::Cancelled => return "upstream request cancelled", + ExecutionErrorKind::Upstream4xx => return "upstream request rejected", + ExecutionErrorKind::Upstream5xx => return "upstream service failed", + ExecutionErrorKind::Internal => {} + } + } + + if lower.contains("timeout") || lower.contains("timed out") { + return "upstream request timed out"; + } + if lower.contains("unauthorized") + || lower.contains("authentication") + || lower.contains("invalid api key") + || lower.contains("invalid token") + { + return "upstream authentication failed"; + } + if lower.contains("forbidden") || lower.contains("authorization") { + return "upstream authorization failed"; + } + if lower.contains("rate limit") || lower.contains("too many requests") { + return "upstream rate limited"; + } + if lower.contains("response body") + || lower.contains("invalid response") + || lower.contains("malformed") + || lower.contains("missing models") + || lower.contains("missing data") + || lower.contains("conflicting cards") + || lower.contains("json") + || lower.contains("parse") + { + return "upstream response invalid"; + } + if lower.contains("connect") + || lower.contains("connection") + || lower.contains("network") + || lower.contains("dns") + || lower.contains("certificate") + { + return "upstream connection failed"; + } + + match status { + Some(status) if (400..500).contains(&status) => "upstream request rejected", + Some(status) if (500..600).contains(&status) => "upstream service failed", + _ => "upstream request failed", + } +} + +fn sanitize_model_fetch_error_detail(detail: &str) -> Option { + let normalized = detail.split_whitespace().collect::>().join(" "); + if normalized.is_empty() + || normalized.len() > MODEL_FETCH_ERROR_DETAIL_MAX_BYTES + || model_fetch_error_contains_sensitive_data(&normalized) + { + return None; + } + + let lower = normalized.to_ascii_lowercase(); + if lower.contains("conflicting cards") { + // The identity is upstream-controlled and the accepted model-ID alphabet also accepts + // common bearer/API-key encodings. A malicious upstream can reflect the credential it + // just received as a conflicting model ID, so the API boundary must omit it entirely. + return Some("conflicting cards".to_string()); + } + + const SAFE_ERROR_PHRASES: &[&str] = &[ + "connection reset", + "connect timeout", + "connection timeout", + "compact endpoint unavailable", + "temporarily unavailable", + "missing models array", + "missing data array", + "missing project_id", + "missing clientmodelconfigs", + "response body is missing json payload", + "invalid response", + "no models", + "endpoint unavailable", + ]; + SAFE_ERROR_PHRASES + .iter() + .find(|phrase| lower.contains(**phrase)) + .map(|phrase| (*phrase).to_string()) +} + +fn model_fetch_error_contains_sensitive_data(value: &str) -> bool { + if value.chars().any(|character| character.is_control()) { + return true; + } + + let lower = value.to_ascii_lowercase(); + const SENSITIVE_MARKERS: &[&str] = &[ + "authorization", + "proxy-authorization", + "bearer", + "basic ", + "api_key", + "api-key", + "apikey", + "access_token", + "access-token", + "refresh_token", + "refresh-token", + "session_token", + "session-token", + "sessionkey", + "password", + "passwd", + "private_key", + "private-key", + "secret", + "credential", + "cookie", + "set-cookie", + "token=", + "token:", + "key=", + "key:", + "url", + "uri", + "path=", + ]; + if SENSITIVE_MARKERS + .iter() + .any(|marker| lower.contains(marker)) + { + return true; + } + + if lower + .split(|character: char| !character.is_ascii_alphanumeric() && character != '_') + .any(|word| { + matches!( + word, + "token" | "apikey" | "password" | "passwd" | "secret" | "credential" + ) + }) + { + return true; + } + + if lower + .chars() + .any(|character| matches!(character, '/' | '\\' | '?' | '#' | '@' | '%')) + { + return true; + } + + // A host, IP literal, or version-like opaque identifier should not cross the boundary as + // part of an error. Model IDs do not need to be included in transport diagnostics. + lower.split_whitespace().any(|token| { + let token = token.trim_matches(|character: char| { + !character.is_ascii_alphanumeric() && character != '.' && character != '-' + }); + if token.len() > 64 { + return true; + } + let parts = token.split('.').collect::>(); + parts.len() >= 2 + && parts.last().is_some_and(|suffix| suffix.len() >= 2) + && parts.iter().all(|part| { + !part.is_empty() + && part + .chars() + .all(|character| character.is_ascii_alphanumeric() || character == '-') + }) + }) +} + +fn model_fetch_status_from_text(error: &str) -> Option { + error + .split(|character: char| !character.is_ascii_digit()) + .filter(|token| token.len() == 3) + .find_map(|token| { + let status = token.parse::().ok()?; + (400..600).contains(&status).then_some(status) + }) +} + +fn model_fetch_format_label(api_format: &str) -> &'static str { + let normalized = api_format.trim().to_ascii_lowercase(); + if normalized.starts_with("openai:") { + "openai" + } else if normalized.starts_with("claude:") { + "claude" + } else if normalized.starts_with("gemini:") { + "gemini" + } else { + "models" } } @@ -1018,7 +1395,7 @@ fn parse_antigravity_models_response(body: &Value) -> Result<(Vec, Option let mut quota_by_model = serde_json::Map::new(); for (model_id, model_data) in models_object { let model_id = model_id.trim(); - if model_id.is_empty() || ANTIGRAVITY_BLOCKED_MODELS.contains(&model_id) { + if !antigravity_model_id_is_routable(model_id) { continue; } let model_object = model_data.as_object().cloned().unwrap_or_default(); @@ -1037,7 +1414,9 @@ fn parse_antigravity_models_response(body: &Value) -> Result<(Vec, Option })); let quota_payload = build_antigravity_quota_payload(model_object.get("quotaInfo")); - quota_by_model.insert(model_id.to_string(), Value::Object(quota_payload)); + if !quota_payload.is_empty() { + quota_by_model.insert(model_id.to_string(), Value::Object(quota_payload)); + } } let upstream_metadata = (!quota_by_model.is_empty()).then(|| { @@ -1052,6 +1431,14 @@ fn parse_antigravity_models_response(body: &Value) -> Result<(Vec, Option Ok((models, upstream_metadata)) } +pub fn antigravity_model_id_is_routable(model_id: &str) -> bool { + let model_id = model_id.trim(); + !model_id.is_empty() + && !ANTIGRAVITY_BLOCKED_MODELS + .iter() + .any(|blocked| blocked.eq_ignore_ascii_case(model_id)) +} + fn parse_kiro_available_models_response( body: &Value, ) -> Result<(Vec, Option), String> { @@ -1141,31 +1528,37 @@ fn infer_kiro_model_owner(model_id: &str) -> &'static str { } fn build_antigravity_quota_payload(quota_info: Option<&Value>) -> serde_json::Map { - let quota_info = quota_info.and_then(Value::as_object); + let Some(quota_info) = quota_info.and_then(Value::as_object) else { + return serde_json::Map::new(); + }; let reset_time = quota_info - .and_then(|value| value.get("resetTime")) + .get("resetTime") .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) .map(ToOwned::to_owned); let remaining_fraction = quota_info - .and_then(|value| value.get("remainingFraction")) - .and_then(Value::as_f64); + .get("remainingFraction") + .and_then(|value| { + value.as_f64().or_else(|| { + value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(|value| value.parse::().ok()) + }) + }) + .filter(|value| value.is_finite()) + .map(|value| value.clamp(0.0, 1.0)); let mut payload = serde_json::Map::new(); - match remaining_fraction { - Some(remaining_fraction) => { - let used_percent = ((1.0 - remaining_fraction) * 100.0).clamp(0.0, 100.0); - payload.insert( - "remaining_fraction".to_string(), - Value::from(remaining_fraction), - ); - payload.insert("used_percent".to_string(), Value::from(used_percent)); - } - None => { - payload.insert("remaining_fraction".to_string(), Value::from(0.0)); - payload.insert("used_percent".to_string(), Value::from(100.0)); - } + if let Some(remaining_fraction) = remaining_fraction { + let used_percent = (1.0 - remaining_fraction) * 100.0; + payload.insert( + "remaining_fraction".to_string(), + Value::from(remaining_fraction), + ); + payload.insert("used_percent".to_string(), Value::from(used_percent)); } if let Some(reset_time) = reset_time { payload.insert("reset_time".to_string(), Value::String(reset_time)); @@ -1208,6 +1601,13 @@ fn iter_vertex_base_urls(transports: &[GatewayProviderTransportSnapshot]) -> Vec urls } +fn iter_trusted_vertex_base_urls(transports: &[GatewayProviderTransportSnapshot]) -> Vec { + iter_vertex_base_urls(transports) + .into_iter() + .filter(|base_url| looks_like_vertex_ai_host(base_url)) + .collect() +} + fn build_vertex_google_list_url(base_url: &str, api_key: &str, page_token: Option<&str>) -> String { let url = build_vertex_publisher_models_list_base_url(base_url, "google"); let mut url = append_query_param(url, "key", api_key); @@ -1410,6 +1810,7 @@ fn build_success_outcome( legacy_models, errors: Vec::new(), has_success, + native_codex_catalog: false, upstream_metadata, etag: None, upstream_status: None, @@ -1472,10 +1873,14 @@ fn append_query_param(mut url: String, key: &str, value: &str) -> String { return url; } let separator = if url.contains('?') { '&' } else { '?' }; + let encoded_key = + url::form_urlencoded::byte_serialize(key.trim().as_bytes()).collect::(); + let encoded_value = + url::form_urlencoded::byte_serialize(value.trim().as_bytes()).collect::(); url.push(separator); - url.push_str(key.trim()); + url.push_str(&encoded_key); url.push('='); - url.push_str(value.trim()); + url.push_str(&encoded_value); url } @@ -1608,16 +2013,25 @@ mod tests { use std::collections::BTreeMap; use std::sync::{Arc, Mutex}; - use aether_contracts::{ExecutionResult, ResponseBody}; + use aether_contracts::{ + redact_url_for_debug, ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResult, + ResponseBody, + }; use aether_provider_transport::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, }; use async_trait::async_trait; + use aws_lc_rs::encoding::{AsDer, Pkcs8V1Der}; + use aws_lc_rs::rsa::{KeyPair as AwsRsaKeyPair, KeySize}; + use aws_lc_rs::signature::{KeyPair as _, UnparsedPublicKey, RSA_PKCS1_2048_8192_SHA256}; + use base64::engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}; + use base64::Engine as _; use serde_json::{json, Value}; use super::{ - build_vertex_google_list_url, build_vertex_service_account_list_url, + build_vertex_google_list_url, build_vertex_service_account_assertion, + build_vertex_service_account_list_url, parse_antigravity_models_response, parse_codex_models_response_for_request, select_model_fetch_strategy, ModelFetchStrategy, ModelFetchStrategyKind, }; @@ -1627,6 +2041,30 @@ mod tests { type RouteResult = Result<(u16, Value), String>; type ModelFetchRoute = (String, RouteResult); + #[test] + fn vertex_assertion_accepts_bare_base64_pkcs8_and_verifies() { + let key_pair = AwsRsaKeyPair::generate(KeySize::Rsa2048) + .expect("2048-bit test RSA private key should generate"); + let pkcs8 = AsDer::>::as_der(&key_pair) + .expect("test RSA private key should encode as PKCS#8"); + let assertion = build_vertex_service_account_assertion( + "svc@example.iam.gserviceaccount.com", + &STANDARD.encode(pkcs8.as_ref()), + "https://oauth2.googleapis.com/token", + 1_700_000_000, + ) + .expect("bare-base64 PKCS#8 key should sign"); + let parts = assertion.split('.').collect::>(); + assert_eq!(parts.len(), 3); + let message = format!("{}.{}", parts[0], parts[1]); + let signature = URL_SAFE_NO_PAD + .decode(parts[2]) + .expect("JWT signature should decode"); + UnparsedPublicKey::new(&RSA_PKCS1_2048_8192_SHA256, key_pair.public_key().as_ref()) + .verify(message.as_bytes(), &signature) + .expect("AWS-LC signature should verify"); + } + struct TestRuntime { executed_urls: Arc>>, response_body: Value, @@ -1715,7 +2153,10 @@ mod tests { .iter() .find(|(url_part, _)| plan.url.contains(url_part)) else { - return Err(format!("unexpected models fetch URL {}", plan.url)); + return Err(format!( + "unexpected models fetch URL {}", + redact_url_for_debug(&plan.url) + )); }; let (status_code, response_body) = match route_result { Ok((status_code, response_body)) => (*status_code, response_body.clone()), @@ -1772,7 +2213,10 @@ mod tests { .iter() .find(|(url_part, _)| plan.url.contains(url_part)) else { - return Err(format!("unexpected models fetch URL {}", plan.url)); + return Err(format!( + "unexpected models fetch URL {}", + redact_url_for_debug(&plan.url) + )); }; let (status_code, response_body) = match route_result { Ok((status_code, response_body)) => (*status_code, response_body.clone()), @@ -2164,6 +2608,185 @@ mod tests { assert!(outcome.errors[0].contains("connect timeout")); } + #[test] + fn execution_result_error_projection_discards_body_credentials_and_urls() { + const BEARER_SECRET: &str = "super-secret-bearer"; + const QUERY_SECRET: &str = "query-secret"; + let result = ExecutionResult { + request_id: "req-model-fetch-security-body".to_string(), + candidate_id: None, + status_code: 401, + headers: BTreeMap::new(), + response_observation: None, + body: Some(ResponseBody { + json_body: Some(json!({ + "error": { + "message": format!( + "Authorization: Bearer {BEARER_SECRET}; https://user:pass@example.test/v1/models?key={QUERY_SECRET}" + ) + } + })), + body_bytes_b64: None, + }), + telemetry: None, + error: None, + }; + + let projected = super::execution_result_error_message(&result); + assert_eq!(projected, "HTTP 401: upstream authentication failed"); + for secret in [ + BEARER_SECRET, + QUERY_SECRET, + "Authorization", + "Bearer", + "user", + "pass", + "example.test", + "/v1/models", + ] { + assert!( + !projected.contains(secret), + "error projection leaked {secret}" + ); + } + } + + #[test] + fn execution_result_error_projection_discards_execution_error_message() { + const SECRET: &str = "transport-secret"; + let result = ExecutionResult { + request_id: "req-model-fetch-security-error".to_string(), + candidate_id: None, + status_code: 502, + headers: BTreeMap::new(), + response_observation: None, + body: None, + telemetry: None, + error: Some(ExecutionError { + kind: ExecutionErrorKind::Upstream5xx, + phase: ExecutionPhase::Connect, + message: format!( + "connection failed for https://user:pass@example.test/v1/models?token={SECRET}" + ), + upstream_status: Some(502), + retryable: true, + failover_recommended: true, + }), + }; + + let projected = super::execution_result_error_message(&result); + assert_eq!(projected, "HTTP 502: upstream service failed"); + for secret in [SECRET, "https://", "user", "pass", "example.test", "token"] { + assert!( + !projected.contains(secret), + "error projection leaked {secret}" + ); + } + } + + #[test] + fn model_fetch_transport_error_projection_discards_urls_and_credentials() { + let projected = super::sanitize_model_fetch_error( + "connection failed for https://user:password@example.test/v1/models?key=query-secret; Authorization: Bearer transport-secret", + ); + assert_eq!(projected, "upstream authorization failed"); + for secret in [ + "user", + "password", + "query-secret", + "transport-secret", + "Bearer", + "example.test", + ] { + assert!( + !projected.contains(secret), + "error projection leaked {secret}" + ); + } + } + + #[test] + fn model_fetch_error_projection_drops_reflected_secret_like_model_identity() { + const REFLECTED_API_KEY: &str = "sk-proj-AbCdEfGhIjKlMnOpQrStUvWxYz0123456789"; + let projected = super::sanitize_model_fetch_error(&format!( + "Codex models response contains conflicting cards for identity '{REFLECTED_API_KEY}'" + )); + + assert_eq!(projected, "conflicting cards"); + assert!(!projected.contains(REFLECTED_API_KEY)); + } + + #[tokio::test] + async fn vertex_service_account_rejects_untrusted_token_uri_before_signing() { + let executed_urls = Arc::new(Mutex::new(Vec::new())); + let runtime = TestRuntime { + executed_urls: Arc::clone(&executed_urls), + response_body: json!({}), + status_code: 200, + response_headers: BTreeMap::new(), + }; + let transport = sample_custom_aiplatform_transport(); + let auth_config = json!({ + "client_email": "svc@example.iam.gserviceaccount.com", + "private_key": "not-a-real-key", + "project_id": "project-1", + "token_uri": "https://attacker.example/token" + }); + + let error = + super::exchange_vertex_service_account_token(&runtime, &transport, &auth_config) + .await + .expect_err("untrusted token URI must be rejected"); + assert_eq!(error, "vertex_ai(service_account): invalid auth_config"); + assert!(executed_urls.lock().expect("executed_urls lock").is_empty()); + } + + #[test] + fn vertex_service_account_model_fetch_ignores_untrusted_endpoint_bases() { + let mut untrusted = sample_custom_aiplatform_transport(); + untrusted.endpoint.base_url = "https://attacker.example".to_string(); + assert_eq!( + super::iter_trusted_vertex_base_urls(&[untrusted]), + vec!["https://aiplatform.googleapis.com".to_string()] + ); + } + + #[tokio::test] + async fn vertex_api_key_model_fetch_uses_only_trusted_endpoint_bases() { + let executed_urls = Arc::new(Mutex::new(Vec::new())); + let runtime = TestRuntime { + executed_urls: Arc::clone(&executed_urls), + response_body: json!({"models": []}), + status_code: 200, + response_headers: BTreeMap::new(), + }; + let mut untrusted = sample_custom_aiplatform_transport(); + untrusted.provider.provider_type = "vertex_ai".to_string(); + untrusted.endpoint.base_url = "https://attacker.example".to_string(); + + fetch_models_from_transports(&runtime, &[untrusted]) + .await + .expect("Vertex API-key model fetch should use the canonical fallback"); + + let urls = executed_urls.lock().expect("executed_urls lock"); + assert_eq!(urls.len(), 1); + assert!(urls[0].starts_with("https://aiplatform.googleapis.com/")); + assert!(urls[0].contains("key=vertex-secret")); + assert!(!urls[0].contains("attacker.example")); + } + + #[test] + fn pagination_query_values_are_percent_encoded() { + assert_eq!( + super::append_query_param( + "https://aiplatform.googleapis.com/v1beta1/models?key=secret".to_string(), + "pageToken", + "cursor&next=1#fragment", + ), + "https://aiplatform.googleapis.com/v1beta1/models?key=secret&pageToken=cursor%26next%3D1%23fragment" + ); + } + #[test] fn vertex_model_fetch_uses_model_garden_list_endpoint() { assert_eq!( @@ -2412,6 +3035,20 @@ mod tests { assert!(outcome.has_success); assert_eq!(outcome.fetched_model_ids, vec!["gpt-legacy-compatible"]); assert_eq!(outcome.cached_models[0]["id"], "gpt-legacy-compatible"); + assert!(!outcome.native_codex_catalog); + let versioned = crate::fetch_models_from_transports_for_management( + &runtime, + &[sample_codex_transport()], + Some("0.153.3"), + ) + .await + .expect("management retains data-array compatibility with a new fingerprint"); + assert!(versioned.has_success); + assert!( + !versioned.native_codex_catalog, + "generic responses cannot replace opaque catalogs" + ); + assert_eq!(versioned.fetched_model_ids, outcome.fetched_model_ids); assert_eq!( outcome.legacy_models[0]["api_formats"], json!(["openai:responses"]) @@ -2537,8 +3174,8 @@ mod tests { .await .expect_err("conflicting endpoint catalogs must fail"); - assert!(error.contains("conflicting cards")); - assert!(error.contains("gpt-cross-identity")); + assert_eq!(error, "conflicting cards"); + assert!(!error.contains("gpt-cross-identity")); } #[tokio::test] @@ -2700,6 +3337,50 @@ mod tests { ); } + #[test] + fn antigravity_models_without_explicit_quota_are_not_marked_exhausted() { + let (models, metadata) = parse_antigravity_models_response(&json!({ + "models": { + "gemini-3.7-flash-tiered": { + "displayName": "Gemini 3.7 Flash" + }, + "gemini-3.7-flash-high": { + "displayName": "Gemini 3.7 Flash High", + "quotaInfo": { + "remainingFraction": "0.75", + "resetTime": "2030-01-01T00:00:00Z" + } + }, + "gemini-3.7-flash-low": { + "displayName": "Gemini 3.7 Flash Low", + "quotaInfo": { + "remainingFraction": 0.0 + } + } + } + })) + .expect("Antigravity models should parse"); + + assert_eq!(models.len(), 3); + let metadata = metadata.expect("explicit quota should produce metadata"); + let antigravity = &metadata["antigravity"]; + assert!(antigravity["quota_by_model"] + .get("gemini-3.7-flash-tiered") + .is_none()); + assert_eq!( + antigravity["quota_by_model"]["gemini-3.7-flash-high"]["remaining_fraction"], + json!(0.75) + ); + assert_eq!( + antigravity["quota_by_model"]["gemini-3.7-flash-high"]["used_percent"], + json!(25.0) + ); + assert_eq!( + antigravity["quota_by_model"]["gemini-3.7-flash-low"]["used_percent"], + json!(100.0) + ); + } + #[tokio::test] async fn kiro_transport_fetches_list_available_models() { let executed_urls = Arc::new(Mutex::new(Vec::new())); diff --git a/crates/aether-model-fetch/src/transport.rs b/crates/aether-model-fetch/src/transport.rs index 15fa42366..8752cf93a 100644 --- a/crates/aether-model-fetch/src/transport.rs +++ b/crates/aether-model-fetch/src/transport.rs @@ -729,10 +729,14 @@ fn append_query_param(mut url: String, key: &str, value: &str) -> String { return url; } let separator = if url.contains('?') { '&' } else { '?' }; + let encoded_key = + url::form_urlencoded::byte_serialize(key.trim().as_bytes()).collect::(); + let encoded_value = + url::form_urlencoded::byte_serialize(value.trim().as_bytes()).collect::(); url.push(separator); - url.push_str(key.trim()); + url.push_str(&encoded_key); url.push('='); - url.push_str(value.trim()); + url.push_str(&encoded_value); url } @@ -773,9 +777,10 @@ mod tests { use serde_json::json; use super::{ - build_antigravity_fetch_available_models_plan, build_antigravity_load_code_assist_plan, - build_gemini_cli_load_code_assist_plan, build_kiro_list_available_models_plan, - build_models_fetch_execution_plan, build_models_fetch_execution_plan_for_client_version, + append_query_param, build_antigravity_fetch_available_models_plan, + build_antigravity_load_code_assist_plan, build_gemini_cli_load_code_assist_plan, + build_kiro_list_available_models_plan, build_models_fetch_execution_plan, + build_models_fetch_execution_plan_for_client_version, build_standard_models_fetch_execution_plan, build_vertex_models_fetch_execution_plan, ModelFetchTransportRuntime, ANTIGRAVITY_REQUEST_USER_AGENT, }; @@ -986,7 +991,7 @@ mod tests { assert_eq!( plan.url, - "https://chatgpt.com/backend-api/codex/models?client_version=0.144.1" + "https://chatgpt.com/backend-api/codex/models?client_version=0.153.3" ); assert_eq!( plan.headers.get("authorization").map(String::as_str), @@ -1185,7 +1190,7 @@ mod tests { ); assert_eq!( plan.headers.get("x-client-version").map(String::as_str), - Some("1.2.3") + Some("4.3.0") ); assert_eq!( plan.headers.get("x-vscode-sessionid").map(String::as_str), @@ -1327,4 +1332,16 @@ mod tests { ); assert!(plan.headers.contains_key("sec-ch-ua")); } + + #[test] + fn pagination_cursor_cannot_inject_additional_query_parameters() { + assert_eq!( + append_query_param( + "https://api.example.test/models?limit=100".to_string(), + "after_id", + "cursor&limit=10000#fragment", + ), + "https://api.example.test/models?limit=100&after_id=cursor%26limit%3D10000%23fragment" + ); + } } diff --git a/crates/aether-oauth/src/core/error.rs b/crates/aether-oauth/src/core/error.rs index ab4d4a186..d49d81b8f 100644 --- a/crates/aether-oauth/src/core/error.rs +++ b/crates/aether-oauth/src/core/error.rs @@ -1,28 +1,298 @@ -use thiserror::Error; +use aether_contracts::redact_url_for_debug; +use serde_json::{json, Value}; + +const OAUTH_ERROR_BODY_EXCERPT_CHARS: usize = 500; -#[derive(Debug, Error)] pub enum OAuthError { - #[error("unsupported oauth provider: {0}")] UnsupportedProvider(String), - #[error("invalid oauth request: {0}")] InvalidRequest(String), - #[error("oauth state is invalid or expired")] InvalidState, - #[error("oauth provider returned HTTP {status_code}: {body_excerpt}")] + // `body_excerpt` remains available to trusted callers for status + // classification, but must not be rendered by the generic Error/Debug + // paths: OAuth servers sometimes echo access tokens, authorization codes, + // assertions, or client credentials in an error response. HttpStatus { status_code: u16, body_excerpt: String, }, - #[error("oauth provider returned invalid response: {0}")] InvalidResponse(String), - #[error("oauth transport failed: {0}")] Transport(String), - #[error("oauth storage failed: {0}")] Storage(String), - #[error("oauth encryption failed")] EncryptionUnavailable, } +impl std::fmt::Display for OAuthError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::UnsupportedProvider(detail) => write!( + formatter, + "unsupported oauth provider: {}", + redact_oauth_error_detail(detail) + ), + Self::InvalidRequest(detail) => write!( + formatter, + "invalid oauth request: {}", + redact_oauth_error_detail(detail) + ), + Self::InvalidState => formatter.write_str("oauth state is invalid or expired"), + Self::HttpStatus { status_code, .. } => { + write!(formatter, "oauth provider returned HTTP {status_code}") + } + Self::InvalidResponse(detail) => write!( + formatter, + "oauth provider returned invalid response: {}", + redact_oauth_error_detail(detail) + ), + Self::Transport(detail) => write!( + formatter, + "oauth transport failed: {}", + redact_oauth_error_detail(detail) + ), + Self::Storage(detail) => write!( + formatter, + "oauth storage failed: {}", + redact_oauth_error_detail(detail) + ), + Self::EncryptionUnavailable => formatter.write_str("oauth encryption failed"), + } + } +} + +impl std::error::Error for OAuthError {} + +impl std::fmt::Debug for OAuthError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::UnsupportedProvider(_) => formatter + .debug_tuple("UnsupportedProvider") + .field(&"[REDACTED]") + .finish(), + Self::InvalidRequest(_) => formatter + .debug_tuple("InvalidRequest") + .field(&"[REDACTED]") + .finish(), + Self::InvalidState => formatter.write_str("InvalidState"), + Self::HttpStatus { status_code, .. } => formatter + .debug_struct("HttpStatus") + .field("status_code", status_code) + .field("body_excerpt", &"[REDACTED]") + .finish(), + Self::InvalidResponse(_) => formatter + .debug_tuple("InvalidResponse") + .field(&"[REDACTED]") + .finish(), + Self::Transport(_) => formatter + .debug_tuple("Transport") + .field(&"[REDACTED]") + .finish(), + Self::Storage(_) => formatter + .debug_tuple("Storage") + .field(&"[REDACTED]") + .finish(), + Self::EncryptionUnavailable => formatter.write_str("EncryptionUnavailable"), + } + } +} + +/// Builds an error excerpt that is safe to persist or render in diagnostics. +/// +/// Structured responses retain non-sensitive provider error codes/messages so +/// callers can still classify `invalid_grant` and similar failures. Secret +/// fields and secret-shaped values are removed before the size bound is +/// applied. Unstructured bodies containing credential markers are replaced as +/// a whole because their token boundaries cannot be determined reliably. +pub fn redacted_oauth_error_body_excerpt(body: &str) -> String { + let body = body.trim(); + if body.is_empty() { + return "-".to_string(); + } + + if let Ok(mut value) = serde_json::from_str::(body) { + redact_oauth_error_json(&mut value); + return value + .to_string() + .chars() + .take(OAUTH_ERROR_BODY_EXCERPT_CHARS) + .collect(); + } + + if unstructured_body_may_contain_secret(body) { + "[REDACTED upstream OAuth error body]".to_string() + } else { + body.chars().take(OAUTH_ERROR_BODY_EXCERPT_CHARS).collect() + } +} + +fn redact_oauth_error_json(value: &mut Value) { + match value { + Value::Object(object) => { + for (key, value) in object { + if oauth_error_key_is_sensitive(key) { + *value = json!("[REDACTED]"); + } else { + redact_oauth_error_json(value); + } + } + } + Value::Array(items) => { + for item in items { + redact_oauth_error_json(item); + } + } + Value::String(text) => { + if oauth_error_value_is_safe_classification_code(text) { + // Keep a small allowlist of non-secret provider error codes for classification. + } else if oauth_error_value_looks_secret(text) + || unstructured_body_may_contain_secret(text) + { + *text = "[REDACTED]".to_string(); + } else { + *text = redact_urls_in_text(text); + } + } + _ => {} + } +} + +fn oauth_error_key_is_sensitive(key: &str) -> bool { + let normalized = key + .chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .collect::() + .to_ascii_lowercase(); + normalized.contains("token") + || normalized.contains("apikey") + || normalized.contains("password") + || normalized.contains("authorization") + || normalized.contains("secret") + || normalized.contains("clientsecret") + || normalized.contains("privatekey") + || normalized.contains("assertion") + || normalized.contains("credential") + || normalized.contains("cookie") + || normalized.contains("pkce") + || normalized.contains("verifier") + || normalized == "sessionkey" +} + +fn oauth_error_value_looks_secret(value: &str) -> bool { + let value = value.trim(); + value.starts_with("Bearer ") + || value.starts_with("bearer ") + || value.starts_with("sk-") + || value.starts_with("sess-") + || value.starts_with("devin-session-token$") + || value.starts_with("ott$") + || value.starts_with("auth1_") + || (value.len() > 80 + && value.split('.').count() == 3 + && value.split('.').all(|segment| { + !segment.is_empty() + && segment.bytes().all(|byte| { + byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'=') + }) + })) +} + +fn oauth_error_value_is_safe_classification_code(value: &str) -> bool { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "refresh_token_reused" | "refresh_token_expired" | "invalid_refresh_token" + ) +} + +/// Redact dynamic OAuth error details before they reach `Display` consumers. +/// Error details frequently originate in HTTP clients and may contain a full +/// request URL or an upstream response body. +fn redact_oauth_error_detail(detail: &str) -> String { + let excerpt = redacted_oauth_error_body_excerpt(detail); + redact_urls_in_text(&excerpt) + .chars() + .take(OAUTH_ERROR_BODY_EXCERPT_CHARS) + .collect() +} + +fn redact_urls_in_text(text: &str) -> String { + const URL_SCHEMES: [&str; 4] = ["http://", "https://", "ws://", "wss://"]; + let mut output = String::with_capacity(text.len()); + let mut cursor = 0; + + while cursor < text.len() { + let Some((relative_start, scheme)) = URL_SCHEMES + .iter() + .filter_map(|scheme| text[cursor..].find(scheme).map(|start| (start, *scheme))) + .min_by_key(|(start, _)| *start) + else { + output.push_str(&text[cursor..]); + break; + }; + + let start = cursor + relative_start; + output.push_str(&text[cursor..start]); + let token_end = text[start..] + .find(char::is_whitespace) + .map(|offset| start + offset) + .unwrap_or(text.len()); + let token = &text[start..token_end]; + let (url_token, suffix) = trim_url_suffix(token); + if url_token.starts_with(scheme) { + if url::Url::parse(url_token).is_ok() { + output.push_str(&redact_url_for_debug(url_token)); + } else { + output.push_str("[REDACTED URL]"); + } + output.push_str(suffix); + } else { + output.push_str(token); + } + cursor = token_end; + } + + output +} + +fn trim_url_suffix(token: &str) -> (&str, &str) { + let mut end = token.len(); + while end > 0 { + let Some(ch) = token[..end].chars().next_back() else { + break; + }; + if matches!( + ch, + '.' | ',' | ';' | ':' | '!' | '?' | ')' | ']' | '}' | '\'' | '"' + ) { + end -= ch.len_utf8(); + } else { + break; + } + } + (&token[..end], &token[end..]) +} + +fn unstructured_body_may_contain_secret(value: &str) -> bool { + let normalized = value.to_ascii_lowercase(); + [ + "access_token", + "refresh_token", + "id_token", + "api_key", + "apikey", + "authorization:", + "authorization=", + "client_secret", + "secret=", + "secret:", + "password=", + "password:", + "assertion=", + "session_token", + "sessiontoken", + ] + .iter() + .any(|marker| normalized.contains(marker)) + || oauth_error_value_looks_secret(value) +} + impl OAuthError { pub fn invalid_request(detail: impl Into) -> Self { Self::InvalidRequest(detail.into()) @@ -36,3 +306,94 @@ impl OAuthError { Self::Transport(detail.into()) } } + +#[cfg(test)] +mod tests { + use super::{redacted_oauth_error_body_excerpt, OAuthError}; + + #[test] + fn oauth_error_body_excerpt_preserves_classification_and_redacts_secrets() { + let excerpt = redacted_oauth_error_body_excerpt( + r#"{ + "error": { + "code": "invalid_grant", + "message": "refresh token expired", + "refresh_token": "refresh-body-canary", + "nested": {"clientSecret": "client-secret-canary"} + }, + "accessToken": "access-token-canary" + }"#, + ); + + assert!(excerpt.contains("invalid_grant")); + assert!(excerpt.contains("refresh token expired")); + assert!(!excerpt.contains("refresh-body-canary")); + assert!(!excerpt.contains("client-secret-canary")); + assert!(!excerpt.contains("access-token-canary")); + assert!(excerpt.contains("[REDACTED]")); + } + + #[test] + fn oauth_error_debug_and_display_do_not_render_upstream_body() { + let error = OAuthError::HttpStatus { + status_code: 401, + body_excerpt: "authorization=Bearer oauth-error-canary".to_string(), + }; + + let debug = format!("{error:?}"); + let display = error.to_string(); + assert!(!debug.contains("oauth-error-canary")); + assert!(!display.contains("oauth-error-canary")); + assert!(debug.contains("[REDACTED]")); + assert_eq!(display, "oauth provider returned HTTP 401"); + } + + #[test] + fn unstructured_oauth_error_with_secret_markers_is_replaced() { + let excerpt = redacted_oauth_error_body_excerpt( + "invalid request: refresh_token=plain-text-refresh-canary", + ); + assert_eq!(excerpt, "[REDACTED upstream OAuth error body]"); + } + + #[test] + fn oauth_error_body_excerpt_preserves_long_refresh_rotation_message() { + let body = r#"{"error":{"message":"Your refresh token has already been used to generate a new access token. Please try signing in again.","type":"invalid_request_error","param":null,"code":"refresh_token_reused"}}"#; + let excerpt = redacted_oauth_error_body_excerpt(body); + assert!(excerpt.contains("already been used to generate a new access token")); + assert!(excerpt.contains("refresh_token_reused")); + } + + #[test] + fn dynamic_oauth_error_display_redacts_credentials_and_url_queries() { + let detail = + "apiKey=sk-secret token=secret-token https://user:pass@example.test?q=url-secret"; + for error in [ + OAuthError::UnsupportedProvider(detail.to_string()), + OAuthError::InvalidRequest(detail.to_string()), + OAuthError::InvalidResponse(detail.to_string()), + OAuthError::Transport(detail.to_string()), + OAuthError::Storage(detail.to_string()), + ] { + let display = error.to_string(); + for secret in ["sk-secret", "secret-token", "user", "pass", "url-secret"] { + assert!( + !display.contains(secret), + "display leaked {secret}: {display}" + ); + } + } + } + + #[test] + fn dynamic_oauth_error_display_redacts_standalone_url() { + let error = OAuthError::invalid_response( + "upstream request failed at https://user:pass@example.test/path?code=secret", + ); + let display = error.to_string(); + assert!(!display.contains("user")); + assert!(!display.contains("pass")); + assert!(!display.contains("code=secret")); + assert!(display.contains("https://example.test/path")); + } +} diff --git a/crates/aether-oauth/src/core/flow.rs b/crates/aether-oauth/src/core/flow.rs index 90d89dfc7..94a7dfbb2 100644 --- a/crates/aether-oauth/src/core/flow.rs +++ b/crates/aether-oauth/src/core/flow.rs @@ -1,6 +1,6 @@ use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct OAuthProviderMetadata { pub provider_type: String, pub display_name: String, @@ -13,7 +13,27 @@ pub struct OAuthProviderMetadata { pub use_pkce: bool, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for OAuthProviderMetadata { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthProviderMetadata") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("authorize_url", &"[REDACTED]") + .field("token_url", &"[REDACTED]") + .field("client_id", &self.client_id) + .field( + "client_secret", + &self.client_secret.as_ref().map(|_| "[REDACTED]"), + ) + .field("scopes", &self.scopes) + .field("redirect_uri", &self.redirect_uri) + .field("use_pkce", &self.use_pkce) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct OAuthAuthorizeRequest { pub state: String, pub code_challenge: Option, @@ -21,7 +41,22 @@ pub struct OAuthAuthorizeRequest { pub login_hint: Option, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +impl std::fmt::Debug for OAuthAuthorizeRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthAuthorizeRequest") + .field("state", &"[REDACTED]") + .field( + "code_challenge", + &self.code_challenge.as_ref().map(|_| "[REDACTED]"), + ) + .field("prompt", &self.prompt) + .field("has_login_hint", &self.login_hint.is_some()) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq, Serialize)] pub struct OAuthAuthorizeResponse { pub authorize_url: String, pub state: String, @@ -29,14 +64,39 @@ pub struct OAuthAuthorizeResponse { pub code_challenge: Option, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for OAuthAuthorizeResponse { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthAuthorizeResponse") + .field("authorize_url", &"[REDACTED]") + .field("state", &"[REDACTED]") + .field( + "code_challenge", + &self.code_challenge.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct OAuthCallback { pub code: String, pub state: String, pub scope: Option, } -#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] +impl std::fmt::Debug for OAuthCallback { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthCallback") + .field("code", &"[REDACTED]") + .field("state", &"[REDACTED]") + .field("scope", &self.scope) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq, Deserialize, Serialize)] pub struct OAuthDeviceAuthorization { pub device_code: String, pub user_code: String, @@ -45,3 +105,86 @@ pub struct OAuthDeviceAuthorization { pub expires_in: u64, pub interval: u64, } + +impl std::fmt::Debug for OAuthDeviceAuthorization { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthDeviceAuthorization") + .field("device_code", &"[REDACTED]") + .field("user_code", &"[REDACTED]") + .field("verification_uri", &self.verification_uri) + .field("verification_uri_complete", &"[REDACTED]") + .field("expires_in", &self.expires_in) + .field("interval", &self.interval) + .finish() + } +} + +#[cfg(test)] +mod tests { + use super::{ + OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback, OAuthDeviceAuthorization, + OAuthProviderMetadata, + }; + + #[test] + fn oauth_flow_debug_output_redacts_capabilities_and_secrets() { + let metadata = OAuthProviderMetadata { + provider_type: "test".to_string(), + display_name: "Test".to_string(), + authorize_url: "https://idp.example/authorize".to_string(), + token_url: "https://idp.example/token".to_string(), + client_id: "public-client".to_string(), + client_secret: Some("client-secret-canary".to_string()), + scopes: vec!["openid".to_string()], + redirect_uri: "https://gateway.example/callback".to_string(), + use_pkce: true, + }; + let request = OAuthAuthorizeRequest { + state: "state-canary".to_string(), + code_challenge: Some("challenge-canary".to_string()), + prompt: None, + login_hint: Some("login-hint-canary".to_string()), + }; + let response = OAuthAuthorizeResponse { + authorize_url: + "https://idp.example/authorize?state=state-url-canary&code_challenge=challenge" + .to_string(), + state: "response-state-canary".to_string(), + code_challenge: Some("response-challenge-canary".to_string()), + }; + let callback = OAuthCallback { + code: "authorization-code-canary".to_string(), + state: "callback-state-canary".to_string(), + scope: None, + }; + let device = OAuthDeviceAuthorization { + device_code: "device-code-canary".to_string(), + user_code: "user-code-canary".to_string(), + verification_uri: "https://idp.example/device".to_string(), + verification_uri_complete: "https://idp.example/device?code=complete-canary" + .to_string(), + expires_in: 600, + interval: 5, + }; + + let debug = format!("{metadata:?} {request:?} {response:?} {callback:?} {device:?}"); + for secret in [ + "client-secret-canary", + "state-canary", + "challenge-canary", + "login-hint-canary", + "state-url-canary", + "response-state-canary", + "response-challenge-canary", + "authorization-code-canary", + "callback-state-canary", + "device-code-canary", + "user-code-canary", + "complete-canary", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + } + assert!(debug.contains("[REDACTED]")); + } +} diff --git a/crates/aether-oauth/src/core/mod.rs b/crates/aether-oauth/src/core/mod.rs index fa785596f..eb234d907 100644 --- a/crates/aether-oauth/src/core/mod.rs +++ b/crates/aether-oauth/src/core/mod.rs @@ -4,7 +4,7 @@ mod pkce; mod registry; mod token; -pub use error::OAuthError; +pub use error::{redacted_oauth_error_body_excerpt, OAuthError}; pub use flow::{ OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback, OAuthDeviceAuthorization, OAuthProviderMetadata, diff --git a/crates/aether-oauth/src/core/token.rs b/crates/aether-oauth/src/core/token.rs index 7f1879768..e8820c73d 100644 --- a/crates/aether-oauth/src/core/token.rs +++ b/crates/aether-oauth/src/core/token.rs @@ -1,7 +1,7 @@ use serde_json::Value; use std::time::{SystemTime, UNIX_EPOCH}; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct OAuthTokenSet { pub access_token: String, pub refresh_token: Option, @@ -11,6 +11,26 @@ pub struct OAuthTokenSet { pub raw_payload: Option, } +impl std::fmt::Debug for OAuthTokenSet { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthTokenSet") + .field("access_token", &"") + .field( + "refresh_token", + &self.refresh_token.as_ref().map(|_| ""), + ) + .field("token_type", &self.token_type) + .field("scope", &self.scope) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .field( + "raw_payload", + &self.raw_payload.as_ref().map(|_| ""), + ) + .finish() + } +} + impl OAuthTokenSet { pub fn from_token_payload(payload: Value) -> Option { let access_token = non_empty_string(payload.get("access_token")) @@ -109,4 +129,20 @@ mod tests { assert!(token.expires_at_unix_secs.is_some()); assert_eq!(token.bearer_header_value(), "Bearer access"); } + + #[test] + fn debug_output_redacts_tokens_and_raw_payload() { + let token = OAuthTokenSet::from_token_payload(json!({ + "access_token": "access-secret-sentinel", + "refresh_token": "refresh-secret-sentinel", + "provider_secret": "raw-secret-sentinel" + })) + .expect("token should parse"); + + let debug = format!("{token:?}"); + assert!(!debug.contains("access-secret-sentinel")); + assert!(!debug.contains("refresh-secret-sentinel")); + assert!(!debug.contains("raw-secret-sentinel")); + assert!(debug.contains("")); + } } diff --git a/crates/aether-oauth/src/identity/adapter.rs b/crates/aether-oauth/src/identity/adapter.rs index afd10968a..74a918e33 100644 --- a/crates/aether-oauth/src/identity/adapter.rs +++ b/crates/aether-oauth/src/identity/adapter.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use serde_json::Value; use std::collections::BTreeMap; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct IdentityOAuthProviderConfig { pub provider_type: String, pub display_name: String, @@ -20,14 +20,60 @@ pub struct IdentityOAuthProviderConfig { pub extra_config: Option, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for IdentityOAuthProviderConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("IdentityOAuthProviderConfig") + .field("provider_type", &self.provider_type) + .field("display_name", &self.display_name) + .field("authorization_url", &"[REDACTED]") + .field("token_url", &"[REDACTED]") + .field( + "userinfo_url", + &self.userinfo_url.as_ref().map(|_| "[REDACTED]"), + ) + .field("client_id", &self.client_id) + .field( + "client_secret", + &self.client_secret.as_ref().map(|_| "[REDACTED]"), + ) + .field("scopes", &self.scopes) + .field("redirect_uri", &self.redirect_uri) + .field("frontend_callback_url", &self.frontend_callback_url) + .field( + "attribute_mapping", + &self.attribute_mapping.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "extra_config", + &self.extra_config.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct IdentityOAuthStartContext { pub state: String, pub code_challenge: Option, pub network: OAuthNetworkContext, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for IdentityOAuthStartContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("IdentityOAuthStartContext") + .field("state", &"[REDACTED]") + .field( + "code_challenge", + &self.code_challenge.as_ref().map(|_| "[REDACTED]"), + ) + .field("network", &self.network) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct IdentityOAuthExchangeContext { pub code: String, pub state: String, @@ -35,27 +81,75 @@ pub struct IdentityOAuthExchangeContext { pub network: OAuthNetworkContext, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for IdentityOAuthExchangeContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("IdentityOAuthExchangeContext") + .field("code", &"[REDACTED]") + .field("state", &"[REDACTED]") + .field( + "pkce_verifier", + &self.pkce_verifier.as_ref().map(|_| "[REDACTED]"), + ) + .field("network", &self.network) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct ExternalIdentity { pub provider_type: String, pub subject: String, pub email: Option, + pub email_verified: bool, pub username: Option, pub display_name: Option, pub avatar_url: Option, pub raw: Value, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for ExternalIdentity { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ExternalIdentity") + .field("provider_type", &self.provider_type) + .field("subject", &self.subject) + .field("email", &self.email) + .field("email_verified", &self.email_verified) + .field("username", &self.username) + .field("display_name", &self.display_name) + .field("avatar_url", &self.avatar_url) + .field("raw", &"[REDACTED]") + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct IdentityClaims { pub provider_type: String, pub subject: String, pub email: Option, + pub email_verified: bool, pub username: Option, pub display_name: Option, pub raw: Value, } +impl std::fmt::Debug for IdentityClaims { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("IdentityClaims") + .field("provider_type", &self.provider_type) + .field("subject", &self.subject) + .field("email", &self.email) + .field("email_verified", &self.email_verified) + .field("username", &self.username) + .field("display_name", &self.display_name) + .field("raw", &"[REDACTED]") + .finish() + } +} + #[async_trait] pub trait IdentityOAuthProvider: Send + Sync { fn provider_type(&self) -> &'static str; @@ -103,16 +197,34 @@ pub(crate) fn mapped_string( find_string(raw, mapped_key) } +pub(crate) fn mapped_bool(raw: &Value, mapping: Option<&Value>, logical_key: &str) -> Option { + let mapped_key = match mapping + .and_then(Value::as_object) + .and_then(|object| object.get(logical_key)) + { + Some(value) => value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty())?, + None => logical_key, + }; + find_value(raw, mapped_key).and_then(Value::as_bool) +} + pub(crate) fn find_string(raw: &Value, key: &str) -> Option { + find_value(raw, key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn find_value<'a>(raw: &'a Value, key: &str) -> Option<&'a Value> { let mut current = raw; for segment in key.split('.') { current = current.get(segment)?; } - current - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) + Some(current) } pub(crate) fn form_headers() -> BTreeMap { @@ -124,3 +236,128 @@ pub(crate) fn form_headers() -> BTreeMap { ("accept".to_string(), "application/json".to_string()), ]) } + +#[cfg(test)] +mod tests { + use super::{ + mapped_bool, ExternalIdentity, IdentityClaims, IdentityOAuthExchangeContext, + IdentityOAuthProviderConfig, IdentityOAuthStartContext, + }; + use crate::network::OAuthNetworkContext; + use serde_json::json; + + #[test] + fn identity_oauth_debug_output_redacts_credentials_and_raw_claims() { + let config = IdentityOAuthProviderConfig { + provider_type: "custom".to_string(), + display_name: "Custom".to_string(), + authorization_url: "https://idp.example/authorize".to_string(), + token_url: "https://idp.example/token".to_string(), + userinfo_url: Some("https://idp.example/userinfo".to_string()), + client_id: "public-client".to_string(), + client_secret: Some("identity-client-secret-canary".to_string()), + scopes: vec!["openid".to_string()], + redirect_uri: "https://gateway.example/callback".to_string(), + frontend_callback_url: "https://app.example/callback".to_string(), + attribute_mapping: None, + extra_config: Some(json!({"secret": "identity-extra-canary"})), + }; + let start = IdentityOAuthStartContext { + state: "identity-state-canary".to_string(), + code_challenge: Some("identity-challenge-canary".to_string()), + network: OAuthNetworkContext::direct_identity(), + }; + let exchange = IdentityOAuthExchangeContext { + code: "identity-code-canary".to_string(), + state: "identity-exchange-state-canary".to_string(), + pkce_verifier: Some("identity-verifier-canary".to_string()), + network: OAuthNetworkContext::direct_identity(), + }; + let external = ExternalIdentity { + provider_type: "custom".to_string(), + subject: "subject".to_string(), + email: None, + email_verified: false, + username: None, + display_name: None, + avatar_url: None, + raw: json!({"access_token": "identity-raw-canary"}), + }; + let claims = IdentityClaims { + provider_type: "custom".to_string(), + subject: "subject".to_string(), + email: None, + email_verified: false, + username: None, + display_name: None, + raw: json!({"id_token": "identity-claims-canary"}), + }; + + let debug = format!("{config:?} {start:?} {exchange:?} {external:?} {claims:?}"); + for secret in [ + "identity-client-secret-canary", + "identity-extra-canary", + "identity-state-canary", + "identity-challenge-canary", + "identity-code-canary", + "identity-exchange-state-canary", + "identity-verifier-canary", + "identity-raw-canary", + "identity-claims-canary", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + } + assert!(debug.contains("[REDACTED]")); + } + + #[test] + fn mapped_bool_accepts_only_an_explicit_json_boolean() { + let raw = json!({ + "email_verified": true, + "profile": { + "verified": false, + "string_verified": "true", + "numeric_verified": 1 + } + }); + + assert_eq!(mapped_bool(&raw, None, "email_verified"), Some(true)); + assert_eq!( + mapped_bool( + &raw, + Some(&json!({"email_verified": "profile.verified"})), + "email_verified" + ), + Some(false) + ); + assert_eq!( + mapped_bool( + &raw, + Some(&json!({"email_verified": "profile.string_verified"})), + "email_verified" + ), + None + ); + assert_eq!( + mapped_bool( + &raw, + Some(&json!({"email_verified": "profile.numeric_verified"})), + "email_verified" + ), + None + ); + assert_eq!(mapped_bool(&json!({}), None, "email_verified"), None); + assert_eq!( + mapped_bool( + &raw, + Some(&json!({"email_verified": true})), + "email_verified" + ), + None + ); + assert_eq!( + mapped_bool(&raw, Some(&json!({"email_verified": ""})), "email_verified"), + None + ); + } +} diff --git a/crates/aether-oauth/src/identity/providers/custom_oidc.rs b/crates/aether-oauth/src/identity/providers/custom_oidc.rs index d267b3f58..b1a1689f3 100644 --- a/crates/aether-oauth/src/identity/providers/custom_oidc.rs +++ b/crates/aether-oauth/src/identity/providers/custom_oidc.rs @@ -1,5 +1,7 @@ -use super::super::adapter::{find_string, form_headers, mapped_string}; -use crate::core::{OAuthAuthorizeResponse, OAuthError, OAuthTokenSet}; +use super::super::adapter::{find_string, form_headers, mapped_bool, mapped_string}; +use crate::core::{ + redacted_oauth_error_body_excerpt, OAuthAuthorizeResponse, OAuthError, OAuthTokenSet, +}; use crate::identity::{ ExternalIdentity, IdentityClaims, IdentityOAuthExchangeContext, IdentityOAuthProvider, IdentityOAuthProviderConfig, IdentityOAuthStartContext, @@ -8,6 +10,16 @@ use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthNetworkContext}; use async_trait::async_trait; use url::form_urlencoded; +const SERVER_MANAGED_AUTHORIZE_PARAMS: &[&str] = &[ + "response_type", + "client_id", + "redirect_uri", + "state", + "scope", + "code_challenge", + "code_challenge_method", +]; + #[derive(Debug, Clone, Default)] pub struct CustomOidcIdentityOAuthProvider; @@ -24,6 +36,15 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { ) -> Result { let mut url = url::Url::parse(&config.authorization_url) .map_err(|_| OAuthError::invalid_request("authorization_url must be absolute"))?; + if url.query_pairs().any(|(name, _)| { + SERVER_MANAGED_AUTHORIZE_PARAMS + .iter() + .any(|reserved| name.eq_ignore_ascii_case(reserved)) + }) { + return Err(OAuthError::invalid_request( + "authorization_url must not predefine server-managed OAuth parameters", + )); + } { let mut query = url.query_pairs_mut(); query.append_pair("response_type", "code"); @@ -87,7 +108,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { if !(200..300).contains(&response.status_code) { return Err(OAuthError::HttpStatus { status_code: response.status_code, - body_excerpt: response.body_text.chars().take(500).collect(), + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), }); } let payload = response @@ -130,7 +151,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { if !(200..300).contains(&response.status_code) { return Err(OAuthError::HttpStatus { status_code: response.status_code, - body_excerpt: response.body_text.chars().take(500).collect(), + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), }); } let raw = response @@ -144,6 +165,8 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { provider_type: config.provider_type.clone(), subject, email: mapped_string(&raw, config.attribute_mapping.as_ref(), "email"), + email_verified: mapped_bool(&raw, config.attribute_mapping.as_ref(), "email_verified") + .unwrap_or(false), username: mapped_string(&raw, config.attribute_mapping.as_ref(), "username"), display_name: mapped_string(&raw, config.attribute_mapping.as_ref(), "display_name") .or_else(|| find_string(&raw, "name")), @@ -160,6 +183,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { Ok(IdentityClaims { provider_type: config.provider_type.clone(), subject: identity.subject, + email_verified: identity.email.is_some() && identity.email_verified, email: identity.email, username: identity.username, display_name: identity.display_name, @@ -167,3 +191,123 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { }) } } + +#[cfg(test)] +mod tests { + use super::CustomOidcIdentityOAuthProvider; + use crate::identity::{ + ExternalIdentity, IdentityOAuthProvider, IdentityOAuthProviderConfig, + IdentityOAuthStartContext, + }; + use crate::network::OAuthNetworkContext; + use serde_json::json; + + fn config() -> IdentityOAuthProviderConfig { + IdentityOAuthProviderConfig { + provider_type: "custom_oidc_work".to_string(), + display_name: "Work OIDC".to_string(), + authorization_url: "https://idp.example.test/authorize".to_string(), + token_url: "https://idp.example.test/token".to_string(), + userinfo_url: Some("https://idp.example.test/userinfo".to_string()), + client_id: "client".to_string(), + client_secret: None, + scopes: vec!["openid".to_string(), "email".to_string()], + redirect_uri: "https://gateway.example.test/callback".to_string(), + frontend_callback_url: "https://app.example.test/callback".to_string(), + attribute_mapping: None, + extra_config: None, + } + } + + fn start_context() -> IdentityOAuthStartContext { + IdentityOAuthStartContext { + state: "server-state".to_string(), + code_challenge: Some("server-challenge".to_string()), + network: OAuthNetworkContext::direct_identity(), + } + } + + #[test] + fn custom_oidc_authorize_url_rejects_predefined_server_managed_parameters() { + for name in [ + "response_type", + "client_id", + "redirect_uri", + "state", + "scope", + "code_challenge", + "code_challenge_method", + ] { + let mut config = config(); + config.authorization_url = + format!("https://idp.example.test/authorize?{name}=attacker"); + + assert!(CustomOidcIdentityOAuthProvider + .build_authorize_url(&config, &start_context()) + .is_err()); + } + } + + #[test] + fn custom_oidc_authorize_url_preserves_non_oauth_tenant_parameters() { + let mut config = config(); + config.authorization_url = + "https://idp.example.test/authorize?tenant=workforce".to_string(); + + let response = CustomOidcIdentityOAuthProvider + .build_authorize_url(&config, &start_context()) + .expect("tenant parameter should be preserved"); + let parsed = url::Url::parse(&response.authorize_url).expect("authorize URL"); + let params = parsed.query_pairs().collect::>(); + + assert!(params + .iter() + .any(|(name, value)| name == "tenant" && value == "workforce")); + assert_eq!(params.iter().filter(|(name, _)| name == "state").count(), 1); + assert!(params + .iter() + .any(|(name, value)| name == "state" && value == "server-state")); + } + + #[test] + fn custom_oidc_propagates_an_explicit_verified_email_claim() { + let claims = CustomOidcIdentityOAuthProvider + .map_identity( + &config(), + ExternalIdentity { + provider_type: "custom_oidc_work".to_string(), + subject: "user-1".to_string(), + email: Some("user@example.test".to_string()), + email_verified: true, + username: Some("user".to_string()), + display_name: None, + avatar_url: None, + raw: json!({"email_verified": true}), + }, + ) + .expect("identity should map"); + + assert!(claims.email_verified); + } + + #[test] + fn custom_oidc_cannot_verify_a_missing_email() { + let claims = CustomOidcIdentityOAuthProvider + .map_identity( + &config(), + ExternalIdentity { + provider_type: "custom_oidc_work".to_string(), + subject: "user-1".to_string(), + email: None, + email_verified: true, + username: Some("user".to_string()), + display_name: None, + avatar_url: None, + raw: json!({"email_verified": true}), + }, + ) + .expect("identity should map"); + + assert!(!claims.email_verified); + } +} diff --git a/crates/aether-oauth/src/identity/providers/linuxdo.rs b/crates/aether-oauth/src/identity/providers/linuxdo.rs index 25c9a189c..9543915a9 100644 --- a/crates/aether-oauth/src/identity/providers/linuxdo.rs +++ b/crates/aether-oauth/src/identity/providers/linuxdo.rs @@ -52,6 +52,53 @@ impl IdentityOAuthProvider for LinuxDoIdentityOAuthProvider { config: &IdentityOAuthProviderConfig, identity: ExternalIdentity, ) -> Result { - self.inner.map_identity(config, identity) + let mut claims = self.inner.map_identity(config, identity)?; + // Linux.do's OAuth user endpoint does not provide an OIDC-level guarantee + // for the email verification claim, so it must not verify a local email. + claims.email_verified = false; + Ok(claims) + } +} + +#[cfg(test)] +mod tests { + use super::LinuxDoIdentityOAuthProvider; + use crate::identity::{ExternalIdentity, IdentityOAuthProvider, IdentityOAuthProviderConfig}; + use serde_json::json; + + #[test] + fn linuxdo_does_not_promote_an_unverified_provider_assertion() { + let provider = LinuxDoIdentityOAuthProvider::default(); + let config = IdentityOAuthProviderConfig { + provider_type: "linuxdo".to_string(), + display_name: "Linux.do".to_string(), + authorization_url: "https://connect.linux.do/oauth2/authorize".to_string(), + token_url: "https://connect.linux.do/oauth2/token".to_string(), + userinfo_url: Some("https://connect.linux.do/api/user".to_string()), + client_id: "client".to_string(), + client_secret: None, + scopes: vec![], + redirect_uri: "https://gateway.example.test/callback".to_string(), + frontend_callback_url: "https://app.example.test/callback".to_string(), + attribute_mapping: None, + extra_config: None, + }; + let claims = provider + .map_identity( + &config, + ExternalIdentity { + provider_type: "linuxdo".to_string(), + subject: "user-1".to_string(), + email: Some("user@example.test".to_string()), + email_verified: true, + username: Some("user".to_string()), + display_name: None, + avatar_url: None, + raw: json!({"email_verified": true}), + }, + ) + .expect("identity should map"); + + assert!(!claims.email_verified); } } diff --git a/crates/aether-oauth/src/lib.rs b/crates/aether-oauth/src/lib.rs index 6ea6e1e6f..400beaadc 100644 --- a/crates/aether-oauth/src/lib.rs +++ b/crates/aether-oauth/src/lib.rs @@ -5,8 +5,8 @@ pub mod provider; pub use core::{ current_unix_secs, generate_oauth_nonce, generate_pkce_verifier, parse_oauth_callback_params, - pkce_s256, OAuthAdapterRegistry, OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback, - OAuthError, OAuthProviderMetadata, OAuthTokenSet, + pkce_s256, redacted_oauth_error_body_excerpt, OAuthAdapterRegistry, OAuthAuthorizeRequest, + OAuthAuthorizeResponse, OAuthCallback, OAuthError, OAuthProviderMetadata, OAuthTokenSet, }; pub use network::{ NetworkRequirement, OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, diff --git a/crates/aether-oauth/src/network/context.rs b/crates/aether-oauth/src/network/context.rs index ce0833a6b..08d6451d9 100644 --- a/crates/aether-oauth/src/network/context.rs +++ b/crates/aether-oauth/src/network/context.rs @@ -38,7 +38,7 @@ impl OAuthTimeouts { }; } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct OAuthNetworkContext { pub policy: OAuthNetworkPolicy, pub requirement: NetworkRequirement, @@ -46,6 +46,18 @@ pub struct OAuthNetworkContext { pub timeouts: OAuthTimeouts, } +impl std::fmt::Debug for OAuthNetworkContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthNetworkContext") + .field("policy", &self.policy) + .field("requirement", &self.requirement) + .field("has_proxy", &self.proxy.is_some()) + .field("timeouts", &self.timeouts) + .finish() + } +} + impl OAuthNetworkContext { pub fn direct_identity() -> Self { Self { @@ -70,3 +82,24 @@ impl OAuthNetworkContext { } } } + +#[cfg(test)] +mod tests { + use super::OAuthNetworkContext; + use aether_contracts::ProxySnapshot; + + #[test] + fn network_context_debug_output_does_not_expose_proxy_credentials() { + let context = OAuthNetworkContext::provider_operation(Some(ProxySnapshot { + url: Some("http://proxy-user:proxy-password@proxy.example:8080".to_string()), + extra: Some(serde_json::json!({"authorization": "proxy-extra-canary"})), + ..ProxySnapshot::default() + })); + + let debug = format!("{context:?}"); + assert!(!debug.contains("proxy-user")); + assert!(!debug.contains("proxy-password")); + assert!(!debug.contains("proxy-extra-canary")); + assert!(debug.contains("has_proxy: true")); + } +} diff --git a/crates/aether-oauth/src/network/executor.rs b/crates/aether-oauth/src/network/executor.rs index 82831a6b3..eb2f8c541 100644 --- a/crates/aether-oauth/src/network/executor.rs +++ b/crates/aether-oauth/src/network/executor.rs @@ -1,11 +1,13 @@ use crate::core::OAuthError; -use aether_contracts::ResolvedTransportProfile; +use aether_contracts::{redact_url_for_debug, ResolvedTransportProfile}; use async_trait::async_trait; use serde_json::Value; use std::collections::BTreeMap; use super::OAuthNetworkContext; +const OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; + #[derive(Clone, PartialEq)] pub struct OAuthHttpRequest { pub request_id: String, @@ -25,7 +27,7 @@ impl std::fmt::Debug for OAuthHttpRequest { .debug_struct("OAuthHttpRequest") .field("request_id", &self.request_id) .field("method", &self.method) - .field("url", &self.url) + .field("url", &redact_url_for_debug(&self.url)) .field("header_names", &self.headers.keys().collect::>()) .field("content_type", &self.content_type) .field("has_json_body", &self.json_body.is_some()) @@ -43,13 +45,24 @@ impl std::fmt::Debug for OAuthHttpRequest { } } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct OAuthHttpResponse { pub status_code: u16, pub body_text: String, pub json_body: Option, } +impl std::fmt::Debug for OAuthHttpResponse { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthHttpResponse") + .field("status_code", &self.status_code) + .field("body_bytes_len", &self.body_text.len()) + .field("has_json_body", &self.json_body.is_some()) + .finish() + } +} + #[async_trait] pub trait OAuthHttpExecutor: Send + Sync { async fn execute(&self, request: OAuthHttpRequest) -> Result; @@ -81,15 +94,29 @@ impl OAuthHttpExecutor for ReqwestOAuthHttpExecutor { builder = builder.body(body_bytes.clone()); } - let response = builder + let mut response = builder .send() .await .map_err(|err| OAuthError::transport(err.to_string()))?; let status_code = response.status().as_u16(); - let body_text = response - .text() + if response + .content_length() + .is_some_and(|length| length > OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES as u64) + { + return Err(oauth_http_response_too_large()); + } + let mut body = Vec::new(); + while let Some(chunk) = response + .chunk() .await - .map_err(|err| OAuthError::transport(err.to_string()))?; + .map_err(|err| OAuthError::transport(err.to_string()))? + { + if chunk.len() > OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) { + return Err(oauth_http_response_too_large()); + } + body.extend_from_slice(&chunk); + } + let body_text = String::from_utf8_lossy(&body).to_string(); let json_body = serde_json::from_str::(&body_text).ok(); Ok(OAuthHttpResponse { status_code, @@ -98,3 +125,51 @@ impl OAuthHttpExecutor for ReqwestOAuthHttpExecutor { }) } } + +fn oauth_http_response_too_large() -> OAuthError { + OAuthError::transport(format!( + "OAuth response body exceeds {OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES} bytes" + )) +} + +#[cfg(test)] +mod tests { + use super::{OAuthHttpRequest, OAuthHttpResponse}; + use crate::network::OAuthNetworkContext; + use std::collections::BTreeMap; + + #[test] + fn response_debug_output_does_not_expose_token_payloads() { + let response = OAuthHttpResponse { + status_code: 200, + body_text: "{\"access_token\":\"response-body-canary\"}".to_string(), + json_body: Some(serde_json::json!({"refresh_token": "response-json-canary"})), + }; + + let debug = format!("{response:?}"); + assert!(!debug.contains("response-body-canary")); + assert!(!debug.contains("response-json-canary")); + assert!(debug.contains("body_bytes_len")); + } + + #[test] + fn request_debug_redacts_url_credentials_and_query() { + let request = OAuthHttpRequest { + request_id: "request-1".into(), + method: reqwest::Method::GET, + url: "https://user:pass@example.test/oauth?client_secret=url-secret".into(), + headers: BTreeMap::from([("authorization".into(), "Bearer header-secret".into())]), + content_type: None, + json_body: None, + body_bytes: None, + network: OAuthNetworkContext::direct_identity(), + transport_profile: None, + }; + let debug = format!("{request:?}"); + assert!(!debug.contains("user")); + assert!(!debug.contains("pass")); + assert!(!debug.contains("url-secret")); + assert!(!debug.contains("header-secret")); + assert!(debug.contains("https://example.test/oauth")); + } +} diff --git a/crates/aether-oauth/src/provider/account.rs b/crates/aether-oauth/src/provider/account.rs index 007a914f1..3bfec2dfc 100644 --- a/crates/aether-oauth/src/provider/account.rs +++ b/crates/aether-oauth/src/provider/account.rs @@ -40,7 +40,7 @@ impl std::fmt::Debug for ProviderOAuthCookieAuthorizationInput { } } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct ProviderOAuthTransportContext { pub provider_id: String, pub provider_type: String, @@ -55,13 +55,57 @@ pub struct ProviderOAuthTransportContext { pub network: OAuthNetworkContext, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for ProviderOAuthTransportContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthTransportContext") + .field("provider_id", &self.provider_id) + .field("provider_type", &self.provider_type) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field("auth_type", &self.auth_type) + .field( + "decrypted_api_key", + &self.decrypted_api_key.as_ref().map(|_| ""), + ) + .field( + "decrypted_auth_config", + &self.decrypted_auth_config.as_ref().map(|_| ""), + ) + .field( + "provider_config", + &self.provider_config.as_ref().map(|_| ""), + ) + .field( + "endpoint_config", + &self.endpoint_config.as_ref().map(|_| ""), + ) + .field( + "key_config", + &self.key_config.as_ref().map(|_| ""), + ) + .field("network", &self.network) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct ProviderOAuthTokenSet { pub token_set: OAuthTokenSet, pub auth_config: Value, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for ProviderOAuthTokenSet { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthTokenSet") + .field("token_set", &self.token_set) + .field("auth_config", &"") + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct ProviderOAuthAccount { pub provider_type: String, pub access_token: String, @@ -70,6 +114,19 @@ pub struct ProviderOAuthAccount { pub identity: BTreeMap, } +impl std::fmt::Debug for ProviderOAuthAccount { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthAccount") + .field("provider_type", &self.provider_type) + .field("access_token", &"") + .field("auth_config", &"") + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .field("identity_keys", &self.identity.keys().collect::>()) + .finish() + } +} + impl ProviderOAuthAccount { pub fn request_bearer_auth(&self) -> ProviderOAuthRequestAuth { ProviderOAuthRequestAuth::Header { @@ -79,7 +136,7 @@ impl ProviderOAuthAccount { } } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub enum ProviderOAuthRequestAuth { Header { name: String, @@ -93,7 +150,28 @@ pub enum ProviderOAuthRequestAuth { }, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for ProviderOAuthRequestAuth { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Header { name, .. } => formatter + .debug_struct("Header") + .field("name", name) + .field("value", &"") + .finish(), + Self::Kiro { + name, machine_id, .. + } => formatter + .debug_struct("Kiro") + .field("name", name) + .field("value", &"") + .field("auth_config", &"") + .field("machine_id", machine_id) + .finish(), + } + } +} + +#[derive(Clone, PartialEq)] pub struct ProviderOAuthImportInput { pub provider_type: String, pub name: Option, @@ -102,7 +180,26 @@ pub struct ProviderOAuthImportInput { pub network: OAuthNetworkContext, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for ProviderOAuthImportInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthImportInput") + .field("provider_type", &self.provider_type) + .field("name", &self.name) + .field( + "refresh_token", + &self.refresh_token.as_ref().map(|_| ""), + ) + .field( + "raw_credentials", + &self.raw_credentials.as_ref().map(|_| ""), + ) + .field("network", &self.network) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct ProviderOAuthAccountState { pub is_valid: bool, pub email: Option, @@ -110,3 +207,95 @@ pub struct ProviderOAuthAccountState { pub invalid_reason: Option, pub raw: Option, } + +impl std::fmt::Debug for ProviderOAuthAccountState { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthAccountState") + .field("is_valid", &self.is_valid) + .field("email", &self.email) + .field("quota", &self.quota) + .field("invalid_reason", &self.invalid_reason) + .field("raw", &self.raw.as_ref().map(|_| "[REDACTED]")) + .finish() + } +} + +#[cfg(test)] +mod tests { + use super::{ + ProviderOAuthAccount, ProviderOAuthAccountState, ProviderOAuthImportInput, + ProviderOAuthTokenSet, ProviderOAuthTransportContext, + }; + use crate::core::OAuthTokenSet; + use crate::network::OAuthNetworkContext; + use serde_json::json; + use std::collections::BTreeMap; + + #[test] + fn debug_output_redacts_provider_oauth_credentials() { + let context = ProviderOAuthTransportContext { + provider_id: "provider-1".to_string(), + provider_type: "generic".to_string(), + endpoint_id: None, + key_id: None, + auth_type: None, + decrypted_api_key: Some("api-key-secret-sentinel".to_string()), + decrypted_auth_config: Some("auth-config-secret-sentinel".to_string()), + provider_config: Some(json!({"client_secret": "provider-secret-sentinel"})), + endpoint_config: None, + key_config: None, + network: OAuthNetworkContext::direct_identity(), + }; + let token_set = ProviderOAuthTokenSet { + token_set: OAuthTokenSet { + access_token: "access-secret-sentinel".to_string(), + refresh_token: Some("refresh-secret-sentinel".to_string()), + token_type: None, + scope: None, + expires_at_unix_secs: None, + raw_payload: None, + }, + auth_config: json!({"password": "password-secret-sentinel"}), + }; + let account = ProviderOAuthAccount { + provider_type: "generic".to_string(), + access_token: "account-secret-sentinel".to_string(), + auth_config: json!({"client_secret": "account-config-secret-sentinel"}), + expires_at_unix_secs: None, + identity: BTreeMap::new(), + }; + let import = ProviderOAuthImportInput { + provider_type: "generic".to_string(), + name: None, + refresh_token: Some("import-refresh-secret-sentinel".to_string()), + raw_credentials: Some(json!({"api_key": "import-raw-secret-sentinel"})), + network: OAuthNetworkContext::direct_identity(), + }; + let state = ProviderOAuthAccountState { + is_valid: false, + email: None, + quota: None, + invalid_reason: None, + raw: Some(json!({"access_token": "probe-raw-secret-sentinel"})), + }; + + let debug = format!("{context:?} {token_set:?} {account:?} {import:?} {state:?}"); + for secret in [ + "api-key-secret-sentinel", + "auth-config-secret-sentinel", + "provider-secret-sentinel", + "access-secret-sentinel", + "refresh-secret-sentinel", + "password-secret-sentinel", + "account-secret-sentinel", + "account-config-secret-sentinel", + "import-refresh-secret-sentinel", + "import-raw-secret-sentinel", + "probe-raw-secret-sentinel", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + } + assert!(debug.contains("")); + } +} diff --git a/crates/aether-oauth/src/provider/providers/antigravity.rs b/crates/aether-oauth/src/provider/providers/antigravity.rs index 9bc86c053..b11cac216 100644 --- a/crates/aether-oauth/src/provider/providers/antigravity.rs +++ b/crates/aether-oauth/src/provider/providers/antigravity.rs @@ -1,11 +1,18 @@ use super::generic::{ provider_account_state_from_metadata, template_for_provider_type, GenericProviderOAuthAdapter, }; -use crate::provider::ProviderOAuthAdapter; +use crate::core::OAuthError; +use crate::network::{OAuthHttpExecutor, OAuthHttpRequest}; +use crate::provider::{ProviderOAuthAdapter, ProviderOAuthTokenSet, ProviderOAuthTransportContext}; +use serde_json::Value; +use std::collections::BTreeMap; + +pub const ANTIGRAVITY_USER_INFO_URL: &str = "https://www.googleapis.com/oauth2/v2/userinfo"; #[derive(Debug, Clone)] pub struct AntigravityProviderOAuthAdapter { inner: GenericProviderOAuthAdapter, + user_info_url: String, } impl Default for AntigravityProviderOAuthAdapter { @@ -15,10 +22,110 @@ impl Default for AntigravityProviderOAuthAdapter { template_for_provider_type("antigravity") .expect("antigravity template should exist"), ), + user_info_url: ANTIGRAVITY_USER_INFO_URL.to_string(), } } } +impl AntigravityProviderOAuthAdapter { + /// Supply deterministic OAuth client credentials for tests without + /// requiring a process-wide environment variable. Production callers + /// continue to resolve the secret from the configured environment. + #[doc(hidden)] + pub fn with_oauth_credentials_for_tests( + mut self, + client_id: impl Into, + client_secret: impl Into, + ) -> Self { + self.inner = self + .inner + .with_oauth_credentials_for_tests(client_id, client_secret); + self + } + + pub fn with_token_url_override(mut self, token_url: impl Into) -> Self { + self.inner = self.inner.with_token_url_override(token_url); + self + } + + pub fn with_user_info_url_override(mut self, user_info_url: impl Into) -> Self { + self.user_info_url = user_info_url.into(); + self + } + + async fn enrich_google_identity( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + mut result: ProviderOAuthTokenSet, + ) -> Result { + if result + .auth_config + .get("email") + .and_then(Value::as_str) + .is_some_and(|email| !email.trim().is_empty()) + { + return Ok(result); + } + + let response = executor + .execute(OAuthHttpRequest { + request_id: "provider-oauth:antigravity-user-info".to_string(), + method: reqwest::Method::GET, + url: self.user_info_url.clone(), + headers: BTreeMap::from([ + ("accept".to_string(), "application/json".to_string()), + ( + "authorization".to_string(), + result.token_set.bearer_header_value(), + ), + ]), + content_type: None, + json_body: None, + body_bytes: None, + 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: response.body_text.trim().chars().take(500).collect(), + }); + } + + let profile = response + .json_body + .or_else(|| serde_json::from_str::(&response.body_text).ok()) + .ok_or_else(|| OAuthError::invalid_response("userinfo response is not json"))?; + if profile.get("verified_email").and_then(Value::as_bool) == Some(false) { + return Err(OAuthError::invalid_response( + "userinfo response returned an unverified email", + )); + } + let email = profile + .get("email") + .and_then(Value::as_str) + .map(str::trim) + .filter(|email| !email.is_empty()) + .ok_or_else(|| OAuthError::invalid_response("userinfo response missing email"))? + .to_string(); + + if let Some(auth_config) = result.auth_config.as_object_mut() { + auth_config.insert("email".to_string(), Value::String(email.clone())); + } + if let Some(token_payload) = result + .token_set + .raw_payload + .as_mut() + .and_then(Value::as_object_mut) + { + token_payload.insert("email".to_string(), Value::String(email)); + } + Ok(result) + } +} + #[async_trait::async_trait] impl ProviderOAuthAdapter for AntigravityProviderOAuthAdapter { fn provider_type(&self) -> &'static str { @@ -59,9 +166,11 @@ impl ProviderOAuthAdapter for AntigravityProviderOAuthAdapter { state: &str, pkce_verifier: Option<&str>, ) -> Result { - self.inner + let result = self + .inner .exchange_code(executor, ctx, code, state, pkce_verifier) - .await + .await?; + self.enrich_google_identity(executor, ctx, result).await } async fn import_credentials( @@ -111,7 +220,7 @@ impl ProviderOAuthAdapter for AntigravityProviderOAuthAdapter { #[cfg(test)] mod tests { - use super::AntigravityProviderOAuthAdapter; + use super::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL}; use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse}; use crate::provider::{ ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthTransportContext, @@ -119,9 +228,15 @@ mod tests { use async_trait::async_trait; use serde_json::json; use std::collections::BTreeMap; + use std::sync::Mutex; struct UnusedExecutor; + #[derive(Default)] + struct GoogleOAuthExecutor { + requests: Mutex>, + } + fn transport_context() -> ProviderOAuthTransportContext { ProviderOAuthTransportContext { provider_id: String::new(), @@ -148,9 +263,47 @@ mod tests { } } + #[async_trait] + impl OAuthHttpExecutor for GoogleOAuthExecutor { + async fn execute( + &self, + request: OAuthHttpRequest, + ) -> Result { + let request_id = request.request_id.clone(); + self.requests + .lock() + .expect("requests should lock") + .push(request); + match request_id.as_str() { + "provider-oauth:exchange-code" => Ok(OAuthHttpResponse { + status_code: 200, + body_text: json!({ + "access_token": "google-access-token", + "refresh_token": "google-refresh-token", + "token_type": "Bearer", + "expires_in": 3600 + }) + .to_string(), + json_body: None, + }), + "provider-oauth:antigravity-user-info" => Ok(OAuthHttpResponse { + status_code: 200, + body_text: json!({ + "email": "antigravity@example.com", + "verified_email": true + }) + .to_string(), + json_body: None, + }), + other => panic!("unexpected OAuth request: {other}"), + } + } + } + #[test] fn antigravity_authorize_requests_offline_refresh_token() { - let adapter = AntigravityProviderOAuthAdapter::default(); + let adapter = AntigravityProviderOAuthAdapter::default() + .with_oauth_credentials_for_tests("test-client-id", "test-client-secret"); let response = adapter .build_authorize_url(&transport_context(), "state-1", Some("challenge-1")) .expect("authorize url should build"); @@ -171,6 +324,47 @@ mod tests { ); } + #[tokio::test] + async fn antigravity_exchange_fetches_google_email_for_account_identity() { + let adapter = AntigravityProviderOAuthAdapter::default() + .with_oauth_credentials_for_tests("test-client-id", "test-client-secret"); + let ctx = transport_context(); + let executor = GoogleOAuthExecutor::default(); + + let result = adapter + .exchange_code( + &executor, + &ctx, + "authorization-code", + "state-1", + Some("verifier-1"), + ) + .await + .expect("Antigravity OAuth exchange should succeed"); + + assert_eq!( + result.auth_config.get("email"), + Some(&json!("antigravity@example.com")) + ); + assert_eq!( + result + .token_set + .raw_payload + .as_ref() + .and_then(|payload| payload.get("email")), + Some(&json!("antigravity@example.com")) + ); + let requests = executor.requests.lock().expect("requests should lock"); + assert_eq!(requests.len(), 2); + assert_eq!(requests[1].url, ANTIGRAVITY_USER_INFO_URL); + assert_eq!(requests[1].method, reqwest::Method::GET); + assert_eq!( + requests[1].headers.get("authorization").map(String::as_str), + Some("Bearer google-access-token") + ); + assert_eq!(requests[1].network, ctx.network); + } + #[tokio::test] async fn antigravity_probe_marks_forbidden_metadata_invalid() { let adapter = AntigravityProviderOAuthAdapter::default(); diff --git a/crates/aether-oauth/src/provider/providers/generic.rs b/crates/aether-oauth/src/provider/providers/generic.rs index c0c563ba6..943ba3e31 100644 --- a/crates/aether-oauth/src/provider/providers/generic.rs +++ b/crates/aether-oauth/src/provider/providers/generic.rs @@ -1,4 +1,7 @@ -use crate::core::{current_unix_secs, OAuthAuthorizeResponse, OAuthError, OAuthTokenSet}; +use crate::core::{ + current_unix_secs, redacted_oauth_error_body_excerpt, OAuthAuthorizeResponse, OAuthError, + OAuthTokenSet, +}; use crate::network::{OAuthHttpExecutor, OAuthHttpRequest}; use crate::provider::ProviderOAuthAdapter; use crate::provider::{ @@ -17,6 +20,10 @@ use super::claude_code::{ CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_REDIRECT_URI, CLAUDE_CODE_TOKEN_URL, }; +pub const GEMINI_CLI_OAUTH_CLIENT_ID_ENV: &str = "AETHER_GEMINI_CLI_OAUTH_CLIENT_ID"; +pub const GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV: &str = "AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET"; +pub const ANTIGRAVITY_OAUTH_CLIENT_ID_ENV: &str = "AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID"; +pub const ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV: &str = "AETHER_ANTIGRAVITY_OAUTH_CLIENT_SECRET"; const CODEX_IDENTITY_FINGERPRINT_FIELD: &str = "codex_identity_fingerprint"; const CODEX_IDENTITY_FINGERPRINT_VERSION: &str = "codex-persisted-fingerprint:v1"; @@ -53,7 +60,8 @@ pub struct GenericProviderOAuthTemplate { pub authorize_url: &'static str, pub token_url: &'static str, pub client_id: &'static str, - pub client_secret: &'static str, + pub client_id_env: Option<&'static str>, + pub client_secret_env: Option<&'static str>, pub scopes: &'static [&'static str], pub redirect_uri: &'static str, pub use_pkce: bool, @@ -68,7 +76,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ authorize_url: CLAUDE_CODE_AUTHORIZE_URL, token_url: CLAUDE_CODE_TOKEN_URL, client_id: CLAUDE_CODE_CLIENT_ID, - client_secret: "", + client_id_env: None, + client_secret_env: None, scopes: CLAUDE_CODE_OAUTH_SCOPES, redirect_uri: CLAUDE_CODE_REDIRECT_URI, use_pkce: true, @@ -81,7 +90,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ authorize_url: "https://auth.openai.com/oauth/authorize", token_url: "https://auth.openai.com/oauth/token", client_id: "app_EMoamEEZ73f0CkXaXp7hrann", - client_secret: "", + client_id_env: None, + client_secret_env: None, scopes: &["openid", "email", "profile", "offline_access"], redirect_uri: "http://localhost:1455/auth/callback", use_pkce: true, @@ -94,7 +104,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ authorize_url: "https://auth.openai.com/oauth/authorize", token_url: "https://auth.openai.com/oauth/token", client_id: "app_EMoamEEZ73f0CkXaXp7hrann", - client_secret: "", + client_id_env: None, + client_secret_env: None, scopes: &["openid", "email", "profile", "offline_access"], redirect_uri: "http://localhost:1455/auth/callback", use_pkce: true, @@ -107,7 +118,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ authorize_url: "https://accounts.google.com/o/oauth2/v2/auth", token_url: "https://oauth2.googleapis.com/token", client_id: "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com", - client_secret: "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", + client_id_env: Some(GEMINI_CLI_OAUTH_CLIENT_ID_ENV), + client_secret_env: Some(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV), scopes: &[ "https://www.googleapis.com/auth/cloud-platform", "https://www.googleapis.com/auth/userinfo.email", @@ -124,7 +136,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ authorize_url: "https://accounts.google.com/o/oauth2/v2/auth", token_url: "https://oauth2.googleapis.com/token", client_id: "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com", - client_secret: "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", + client_id_env: Some(ANTIGRAVITY_OAUTH_CLIENT_ID_ENV), + client_secret_env: Some(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV), scopes: &[ "https://www.googleapis.com/auth/cloud-platform", "https://www.googleapis.com/auth/userinfo.email", @@ -139,10 +152,24 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ }, ]; -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct GenericProviderOAuthAdapter { template: GenericProviderOAuthTemplate, token_url_override: Option, + client_id_override: Option, + client_secret_override: Option, +} + +impl std::fmt::Debug for GenericProviderOAuthAdapter { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GenericProviderOAuthAdapter") + .field("provider_type", &self.template.provider_type) + .field("has_token_url_override", &self.token_url_override.is_some()) + .field("client_id_env", &self.template.client_id_env) + .field("client_secret_env", &self.template.client_secret_env) + .finish_non_exhaustive() + } } impl GenericProviderOAuthAdapter { @@ -150,6 +177,8 @@ impl GenericProviderOAuthAdapter { Self { template, token_url_override: None, + client_id_override: None, + client_secret_override: None, } } @@ -166,12 +195,52 @@ impl GenericProviderOAuthAdapter { self.with_token_url_override(token_url) } + #[doc(hidden)] + pub fn with_oauth_credentials_for_tests( + mut self, + client_id: impl Into, + client_secret: impl Into, + ) -> Self { + self.client_id_override = Some(client_id.into()); + self.client_secret_override = Some(client_secret.into()); + self + } + + #[cfg(test)] + fn without_oauth_client_secret_for_tests(mut self) -> Self { + self.client_secret_override = Some(String::new()); + self + } + fn token_url(&self) -> String { self.token_url_override .clone() .unwrap_or_else(|| self.template.token_url.to_string()) } + fn client_id(&self) -> String { + if let Some(value) = self.client_id_override.clone().and_then(non_empty_owned) { + return value; + } + + self.template + .client_id_env + .and_then(non_empty_environment_value) + .unwrap_or_else(|| self.template.client_id.to_string()) + } + + fn client_secret(&self) -> Result, OAuthError> { + let Some(env_name) = self.template.client_secret_env else { + return Ok(None); + }; + + if let Some(value) = self.client_secret_override.clone() { + return required_client_secret(env_name, non_empty_owned(value)).map(Some); + } + + required_client_secret(env_name, non_empty_environment_value(env_name)).map(Some) + } + async fn exchange_grant( &self, executor: &dyn OAuthHttpExecutor, @@ -181,6 +250,8 @@ impl GenericProviderOAuthAdapter { state: Option<&str>, pkce_verifier: Option<&str>, ) -> Result { + let client_id = self.client_id(); + let client_secret = self.client_secret()?; let scope = (!self.template.scopes.is_empty()).then(|| self.template.scopes.join(" ")); let request_id = match grant_type { "authorization_code" => "provider-oauth:exchange-code".to_string(), @@ -196,11 +267,14 @@ impl GenericProviderOAuthAdapter { "grant_type".to_string(), Value::String(grant_type.to_string()), ), - ( - "client_id".to_string(), - Value::String(self.template.client_id.to_string()), - ), + ("client_id".to_string(), Value::String(client_id.clone())), ]); + if let Some(client_secret) = client_secret.as_ref() { + body.insert( + "client_secret".to_string(), + Value::String(client_secret.clone()), + ); + } if grant_type == "authorization_code" { body.insert( "code".to_string(), @@ -247,7 +321,7 @@ impl GenericProviderOAuthAdapter { let form_body = { let mut form = form_urlencoded::Serializer::new(String::new()); form.append_pair("grant_type", grant_type); - form.append_pair("client_id", self.template.client_id); + form.append_pair("client_id", &client_id); if grant_type == "authorization_code" { form.append_pair("redirect_uri", self.template.redirect_uri); form.append_pair("code", code_or_refresh_token); @@ -262,8 +336,8 @@ impl GenericProviderOAuthAdapter { form.append_pair("scope", scope); } } - if !self.template.client_secret.trim().is_empty() { - form.append_pair("client_secret", self.template.client_secret); + if let Some(client_secret) = client_secret.as_deref() { + form.append_pair("client_secret", client_secret); } form.finish().into_bytes() }; @@ -340,12 +414,14 @@ impl ProviderOAuthAdapter for GenericProviderOAuthAdapter { state: &str, code_challenge: Option<&str>, ) -> Result { + self.client_secret()?; + let client_id = self.client_id(); let mut url = url::Url::parse(self.template.authorize_url) .map_err(|_| OAuthError::invalid_request("authorize_url must be absolute"))?; { let mut query = url.query_pairs_mut(); query.append_pair("response_type", "code"); - query.append_pair("client_id", self.template.client_id); + query.append_pair("client_id", &client_id); query.append_pair("redirect_uri", self.template.redirect_uri); query.append_pair("state", state); if !self.template.scopes.is_empty() { @@ -472,6 +548,26 @@ impl ProviderOAuthAdapter for GenericProviderOAuthAdapter { } } +fn non_empty_environment_value(name: &str) -> Option { + std::env::var(name).ok().and_then(non_empty_owned) +} + +fn non_empty_owned(value: String) -> Option { + let value = value.trim(); + (!value.is_empty()).then(|| value.to_string()) +} + +fn required_client_secret( + env_name: &'static str, + configured: Option, +) -> Result { + configured.ok_or_else(|| { + OAuthError::invalid_request(format!( + "{env_name} must be configured for this OAuth provider" + )) + }) +} + pub fn template_for_provider_type(provider_type: &str) -> Option { let normalized = provider_type.trim(); GENERIC_PROVIDER_OAUTH_TEMPLATES @@ -506,12 +602,7 @@ fn json_headers(provider_type: &str) -> BTreeMap { } fn truncate_body(body: &str) -> String { - let body = body.trim(); - if body.is_empty() { - "-".to_string() - } else { - body.chars().take(500).collect() - } + redacted_oauth_error_body_excerpt(body) } fn secret_fingerprint(value: &str) -> String { @@ -790,8 +881,21 @@ fn value_to_string(value: &Value) -> Option { 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() @@ -801,9 +905,12 @@ fn decode_jwt_claims(token: &str) -> Option> { #[cfg(test)] mod tests { use super::{ - derive_codex_identity_fingerprint, enrich_generic_identity, template_for_provider_type, - GenericProviderOAuthAdapter, CODEX_IDENTITY_FINGERPRINT_FIELD, + decode_jwt_claims, derive_codex_identity_fingerprint, enrich_generic_identity, + template_for_provider_type, GenericProviderOAuthAdapter, + ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV, CODEX_IDENTITY_FINGERPRINT_FIELD, + GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV, }; + use crate::core::OAuthError; use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse}; use crate::provider::ProviderOAuthAdapter; use crate::provider::{ProviderOAuthAccount, ProviderOAuthTransportContext}; @@ -835,6 +942,36 @@ mod tests { ) } + #[test] + fn google_oauth_templates_reference_external_client_secrets() { + let gemini = template_for_provider_type("gemini_cli").expect("gemini template"); + let antigravity = template_for_provider_type("antigravity").expect("antigravity template"); + + assert_eq!( + gemini.client_secret_env, + Some(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV) + ); + assert_eq!( + antigravity.client_secret_env, + Some(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV) + ); + } + + #[test] + fn generic_adapter_debug_redacts_oauth_credentials() { + let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli") + .expect("gemini adapter") + .with_token_url_override("https://token.example.test/private-path") + .with_oauth_credentials_for_tests("private-client-id", "private-client-secret"); + + let debug = format!("{adapter:?}"); + + assert!(!debug.contains("private-client-id")); + assert!(!debug.contains("private-client-secret")); + assert!(!debug.contains("private-path")); + assert!(debug.contains(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV)); + } + #[test] fn codex_identity_extracts_fedramp_workspace_claim() { let claims = json!({ @@ -852,6 +989,19 @@ mod tests { assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true))); } + #[test] + fn generic_identity_rejects_oversized_jwt_claims_before_decode() { + const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024; + let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap() + .saturating_mul(4); + let token = format!("header.{}.signature", "A".repeat(max_encoded_len + 1)); + + assert_eq!(decode_jwt_claims(&token), None); + } + #[test] fn codex_persisted_fingerprint_is_member_scoped_and_token_independent() { let adapter = GenericProviderOAuthAdapter::for_provider_type("codex") @@ -937,6 +1087,108 @@ mod tests { } } + fn transport_context(provider_type: &str) -> ProviderOAuthTransportContext { + ProviderOAuthTransportContext { + provider_id: "provider-1".to_string(), + provider_type: provider_type.to_string(), + endpoint_id: None, + key_id: Some("key-1".to_string()), + 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 oauth_account(provider_type: &str) -> ProviderOAuthAccount { + ProviderOAuthAccount { + provider_type: provider_type.to_string(), + access_token: "old-access-token".to_string(), + auth_config: json!({ + "provider_type": provider_type, + "refresh_token": "old-refresh-token", + "updated_at": 1 + }), + expires_at_unix_secs: Some(1), + identity: BTreeMap::new(), + } + } + + #[tokio::test] + async fn google_oauth_fails_closed_before_network_without_client_secret() { + let seen_request = Arc::new(Mutex::new(None)); + let executor = StaticExecutor { + seen_request: Arc::clone(&seen_request), + response_payload: json!({}), + }; + let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli") + .expect("gemini adapter") + .without_oauth_client_secret_for_tests(); + let ctx = transport_context("gemini_cli"); + + let authorize_error = adapter + .build_authorize_url(&ctx, "state", None) + .expect_err("authorization must fail without the configured secret"); + assert!(matches!(authorize_error, OAuthError::InvalidRequest(_))); + + let refresh_error = adapter + .refresh(&executor, &ctx, &oauth_account("gemini_cli")) + .await + .expect_err("refresh must fail without the configured secret"); + assert!(matches!(refresh_error, OAuthError::InvalidRequest(_))); + assert!( + seen_request.lock().expect("mutex should lock").is_none(), + "credential validation must happen before the HTTP executor runs" + ); + } + + #[tokio::test] + async fn google_oauth_injected_credentials_are_sent_in_token_form() { + let seen_request = Arc::new(Mutex::new(None)); + let executor = StaticExecutor { + seen_request: Arc::clone(&seen_request), + response_payload: json!({ + "access_token": "new-access-token", + "expires_in": 3600, + }), + }; + let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli") + .expect("gemini adapter") + .with_oauth_credentials_for_tests("test-client-id", "test-client-secret"); + let ctx = transport_context("gemini_cli"); + + adapter + .refresh(&executor, &ctx, &oauth_account("gemini_cli")) + .await + .expect("refresh should succeed"); + + let seen = seen_request + .lock() + .expect("mutex should lock") + .clone() + .expect("request should be captured"); + let form = String::from_utf8(seen.body_bytes.expect("form body should exist")) + .expect("form body should be utf8"); + let fields = url::form_urlencoded::parse(form.as_bytes()) + .into_owned() + .collect::>(); + assert_eq!( + fields.get("client_id").map(String::as_str), + Some("test-client-id") + ); + assert_eq!( + fields.get("client_secret").map(String::as_str), + Some("test-client-secret") + ); + assert_eq!( + fields.get("refresh_token").map(String::as_str), + Some("old-refresh-token") + ); + } + #[tokio::test] async fn refresh_preserves_existing_metadata_when_refresh_token_is_not_rotated() { let seen_request = Arc::new(Mutex::new(None)); diff --git a/crates/aether-oauth/src/provider/providers/kiro.rs b/crates/aether-oauth/src/provider/providers/kiro.rs index 921a87e69..920a5f864 100644 --- a/crates/aether-oauth/src/provider/providers/kiro.rs +++ b/crates/aether-oauth/src/provider/providers/kiro.rs @@ -1,4 +1,4 @@ -use crate::core::{current_unix_secs, OAuthError}; +use crate::core::{current_unix_secs, redacted_oauth_error_body_excerpt, OAuthError}; use crate::network::{OAuthHttpExecutor, OAuthHttpRequest}; use crate::provider::ProviderOAuthAdapter; use crate::provider::{ @@ -18,7 +18,7 @@ pub const DEFAULT_SYSTEM_VERSION: &str = "other#unknown"; const IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE"; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct KiroAuthConfig { pub auth_method: Option, pub refresh_token: Option, @@ -36,6 +36,75 @@ pub struct KiroAuthConfig { pub access_token: Option, } +impl std::fmt::Debug for KiroAuthConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("KiroAuthConfig") + .field("auth_method", &self.auth_method) + .field( + "refresh_token", + &self.refresh_token.as_ref().map(|_| "[REDACTED]"), + ) + .field("expires_at", &self.expires_at) + .field( + "profile_arn", + &self.profile_arn.as_ref().map(|_| "[REDACTED]"), + ) + .field("region", &self.region) + .field("auth_region", &self.auth_region) + .field("api_region", &self.api_region) + .field("client_id", &self.client_id.as_ref().map(|_| "[REDACTED]")) + .field( + "client_secret", + &self.client_secret.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "machine_id", + &self.machine_id.as_ref().map(|_| "[REDACTED]"), + ) + .field("kiro_version", &self.kiro_version) + .field("system_version", &self.system_version) + .field("node_version", &self.node_version) + .field( + "access_token", + &self.access_token.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + +/// Returns whether a value is safe to interpolate as one DNS label in a Kiro +/// service hostname. +/// +/// Region values are persisted in the encrypted auth configuration and can +/// also come from an upstream OAuth response. They must never be treated as +/// URL syntax: accepting `/`, `.`, `?`, `#`, `@`, or control characters here +/// would let a crafted region redirect a token-bearing request to another +/// origin. AWS region names are DNS labels, so an ASCII alphanumeric/hyphen +/// allow-list is both stricter and forward-compatible with future regions. +pub fn is_valid_kiro_region(value: &str) -> bool { + let value = value.trim(); + !value.is_empty() + && value.len() <= 63 + && !value.starts_with('-') + && !value.ends_with('-') + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') +} + +/// Trims a region and falls back to the known-safe default when it is not a +/// single DNS label. The returned value is suitable for URL and `Host` +/// header interpolation. +pub fn normalize_kiro_region(value: &str) -> &str { + let value = value.trim(); + if is_valid_kiro_region(value) { + value + } else { + DEFAULT_REGION + } +} + impl KiroAuthConfig { pub fn from_json_value(value: &Value) -> Option { let object = value.as_object()?; @@ -93,19 +162,22 @@ impl KiroAuthConfig { } pub fn effective_auth_region(&self) -> &str { - self.auth_region - .as_deref() - .or(self.region.as_deref()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or(DEFAULT_REGION) + for candidate in [self.auth_region.as_deref(), self.region.as_deref()] { + let Some(candidate) = candidate else { + continue; + }; + let candidate = candidate.trim(); + if is_valid_kiro_region(candidate) { + return candidate; + } + } + DEFAULT_REGION } pub fn effective_api_region(&self) -> &str { self.api_region .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) + .map(normalize_kiro_region) .unwrap_or(DEFAULT_REGION) } @@ -304,7 +376,7 @@ impl KiroProviderOAuthAdapter { if !(200..300).contains(&response.status_code) { return Err(OAuthError::HttpStatus { status_code: response.status_code, - body_excerpt: response.body_text.chars().take(500).collect(), + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), }); } let payload = response @@ -383,7 +455,7 @@ impl KiroProviderOAuthAdapter { if !(200..300).contains(&response.status_code) { return Err(OAuthError::HttpStatus { status_code: response.status_code, - body_excerpt: response.body_text.chars().take(500).collect(), + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), }); } let payload = response @@ -648,7 +720,8 @@ fn secret_fingerprint(value: &str) -> String { #[cfg(test)] mod tests { use super::{ - generate_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter, IDC_AMZ_USER_AGENT, + generate_kiro_machine_id, is_valid_kiro_region, normalize_kiro_region, KiroAuthConfig, + KiroProviderOAuthAdapter, IDC_AMZ_USER_AGENT, }; use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse}; use crate::provider::ProviderOAuthTransportContext; @@ -662,6 +735,39 @@ mod tests { response: serde_json::Value, } + #[test] + fn kiro_auth_config_debug_output_redacts_all_token_material() { + let config = KiroAuthConfig { + auth_method: Some("idc".to_string()), + refresh_token: Some("kiro-refresh-canary".to_string()), + expires_at: Some(123), + profile_arn: Some("kiro-profile-arn-canary".to_string()), + region: Some("us-east-1".to_string()), + auth_region: None, + api_region: None, + client_id: Some("kiro-client-id-canary".to_string()), + client_secret: Some("kiro-client-secret-canary".to_string()), + machine_id: Some("kiro-machine-id-canary".to_string()), + kiro_version: None, + system_version: None, + node_version: None, + access_token: Some("kiro-access-canary".to_string()), + }; + + let debug = format!("{config:?}"); + for secret in [ + "kiro-refresh-canary", + "kiro-client-id-canary", + "kiro-client-secret-canary", + "kiro-access-canary", + "kiro-profile-arn-canary", + "kiro-machine-id-canary", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + } + assert!(debug.contains("[REDACTED]")); + } + #[async_trait] impl OAuthHttpExecutor for StaticExecutor { async fn execute( @@ -717,6 +823,106 @@ mod tests { ); } + #[test] + fn rejects_region_values_that_can_escape_a_hostname() { + for value in [ + "evil.example/", + "evil.example\\", + "evil.example?next=1", + "evil.example#fragment", + "evil@example", + "us-east-1\r\nX-Injected: yes", + "us.east.1", + ] { + assert!( + !is_valid_kiro_region(value), + "region should be rejected: {value:?}" + ); + assert_eq!(normalize_kiro_region(value), super::DEFAULT_REGION); + } + + for value in [ + "us-east-1", + "us-gov-west-1", + "us-iso-east-1", + "eu-central-1", + ] { + assert!( + is_valid_kiro_region(value), + "region should be accepted: {value}" + ); + assert_eq!(normalize_kiro_region(value), value); + } + } + + #[test] + fn effective_regions_fall_back_when_auth_config_is_malicious() { + let auth_config = KiroAuthConfig { + auth_method: None, + refresh_token: None, + expires_at: None, + profile_arn: None, + region: Some("eu-west-1".to_string()), + auth_region: Some("evil.example/".to_string()), + api_region: Some("evil.example/".to_string()), + client_id: None, + client_secret: None, + machine_id: None, + kiro_version: None, + system_version: None, + node_version: None, + access_token: None, + }; + + assert_eq!(auth_config.effective_auth_region(), "eu-west-1"); + assert_eq!(auth_config.effective_api_region(), super::DEFAULT_REGION); + } + + #[tokio::test] + async fn refresh_urls_use_safe_default_for_malicious_auth_region() { + let seen_request = Arc::new(Mutex::new(None)); + let executor = StaticExecutor { + seen_request: Arc::clone(&seen_request), + response: json!({ + "accessToken": "new-access-token", + "refreshToken": "r".repeat(120), + "expiresIn": 3600 + }), + }; + let auth_config = KiroAuthConfig { + auth_method: Some("social".to_string()), + refresh_token: Some("r".repeat(120)), + expires_at: Some(1), + profile_arn: None, + region: None, + auth_region: Some("attacker.example/".to_string()), + api_region: None, + client_id: None, + client_secret: None, + machine_id: Some("machine".to_string()), + kiro_version: None, + system_version: None, + node_version: None, + access_token: None, + }; + + KiroProviderOAuthAdapter::default() + .refresh_auth_config(&executor, &test_ctx(), &auth_config) + .await + .expect("refresh should succeed"); + + let seen = seen_request + .lock() + .expect("mutex should lock") + .clone() + .expect("request should be captured"); + assert_eq!( + seen.url, + "https://prod.us-east-1.auth.desktop.kiro.dev/refreshToken" + ); + assert!(!seen.url.contains("attacker.example")); + } + #[tokio::test] async fn refreshes_social_auth_config_with_provider_adapter() { let seen_request = Arc::new(Mutex::new(None)); diff --git a/crates/aether-oauth/src/provider/providers/mod.rs b/crates/aether-oauth/src/provider/providers/mod.rs index 4fe916a75..f6fe878a8 100644 --- a/crates/aether-oauth/src/provider/providers/mod.rs +++ b/crates/aether-oauth/src/provider/providers/mod.rs @@ -5,7 +5,7 @@ mod generic; mod kiro; mod windsurf; -pub use antigravity::AntigravityProviderOAuthAdapter; +pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL}; pub use claude_code::{ ClaudeCodeProviderOAuthAdapter, CLAUDE_CODE_AUTHORIZE_URL, CLAUDE_CODE_CLIENT_ID, CLAUDE_CODE_COOKIE_SCOPE, CLAUDE_CODE_OAUTH_SCOPES, CLAUDE_CODE_PROVIDER_TYPE, @@ -14,12 +14,14 @@ pub use claude_code::{ pub use codex::CodexProviderOAuthAdapter; pub use generic::{ derive_codex_identity_fingerprint, GenericProviderOAuthAdapter, GenericProviderOAuthTemplate, + ANTIGRAVITY_OAUTH_CLIENT_ID_ENV, ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV, + GEMINI_CLI_OAUTH_CLIENT_ID_ENV, GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV, GENERIC_PROVIDER_OAUTH_TEMPLATES, }; pub use kiro::{ - generate_kiro_machine_id, normalize_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter, - DEFAULT_KIRO_VERSION, DEFAULT_NODE_VERSION, DEFAULT_REGION, DEFAULT_SYSTEM_VERSION, - KIRO_PROVIDER_TYPE, + generate_kiro_machine_id, is_valid_kiro_region, normalize_kiro_machine_id, + normalize_kiro_region, KiroAuthConfig, KiroProviderOAuthAdapter, DEFAULT_KIRO_VERSION, + DEFAULT_NODE_VERSION, DEFAULT_REGION, DEFAULT_SYSTEM_VERSION, KIRO_PROVIDER_TYPE, }; pub use windsurf::{ WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE, diff --git a/crates/aether-pool-core/src/scheduler.rs b/crates/aether-pool-core/src/scheduler.rs index 64cf939b7..8cf3ebfc3 100644 --- a/crates/aether-pool-core/src/scheduler.rs +++ b/crates/aether-pool-core/src/scheduler.rs @@ -18,6 +18,8 @@ pub struct PoolSchedulingPreset { pub struct PoolSchedulingConfig { pub scheduling_presets: Vec, pub lru_enabled: bool, + /// Retained for configuration/API compatibility. Active quota exhaustion is + /// always an admission block; reset-aware adapters decide when it clears. pub skip_exhausted_accounts: bool, pub cost_limit_per_key_tokens: Option, } @@ -225,9 +227,14 @@ fn schedule_pool_group( continue; } - if item.key_context.quota_hard_blocked - || (pool_config.skip_exhausted_accounts && item.key_context.quota_exhausted) - { + // A quota snapshot is an account-level admission signal, not merely a + // ranking hint. Continuing to schedule a member whose quota is known to + // be exhausted causes a request-wide retry storm (the upstream returns + // 429 for every attempt). `quota_hard_blocked` remains available for + // providers that can distinguish an explicit permanent block, but every + // active exhaustion must be removed from the request's candidate set; + // reset-aware provider adapters clear the signal once capacity returns. + if item.key_context.quota_hard_blocked || item.key_context.quota_exhausted { skipped.push(PoolSkippedCandidate { candidate: item.candidate, skip_reason: POOL_ACCOUNT_EXHAUSTED_SKIP_REASON, @@ -928,6 +935,29 @@ mod tests { ); } + #[test] + fn pool_scheduler_skips_exhausted_accounts_even_when_legacy_flag_is_false() { + let ready = sample_candidate("provider-pool", "endpoint-1", "key-ready", 10, true); + let mut exhausted = + sample_candidate("provider-pool", "endpoint-1", "key-exhausted", 10, true); + exhausted.key_context.quota_exhausted = true; + + let outcome = run_pool_scheduler(vec![ready, exhausted], &BTreeMap::new(), "seed"); + + assert_eq!( + outcome + .candidates + .iter() + .map(|item| item.candidate.as_str()) + .collect::>(), + vec!["key-ready"] + ); + assert_eq!( + outcome.skipped_candidates[0].skip_reason, + POOL_ACCOUNT_EXHAUSTED_SKIP_REASON + ); + } + #[test] fn pool_scheduler_promotes_sticky_hit_before_other_sorted_keys() { let key_a = sample_candidate("provider-pool", "endpoint-1", "key-a", 10, true) diff --git a/crates/aether-provider/pool/Cargo.toml b/crates/aether-provider/pool/Cargo.toml index 2e6226457..f08919450 100644 --- a/crates/aether-provider/pool/Cargo.toml +++ b/crates/aether-provider/pool/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true description = "Provider-specific pool behavior adapters for Aether" [dependencies] +aether-contracts.workspace = true aether-data-contracts.workspace = true aether-pool-core.workspace = true aether-provider-transport.workspace = true diff --git a/crates/aether-provider/pool/src/lib.rs b/crates/aether-provider/pool/src/lib.rs index 58dec5b43..2eb0f0529 100644 --- a/crates/aether-provider/pool/src/lib.rs +++ b/crates/aether-provider/pool/src/lib.rs @@ -15,10 +15,11 @@ pub use presets::{ }; pub use provider::{ProviderPoolAdapter, ProviderPoolMemberInput}; pub use providers::{ - build_antigravity_pool_quota_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_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_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, @@ -27,14 +28,16 @@ pub use providers::{ AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter, - ANTIGRAVITY_FETCH_AVAILABLE_MODELS_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, + 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, }; pub use quota::{ - provider_pool_key_account_quota_exhausted, provider_pool_key_quota_hard_blocked, + provider_pool_key_account_quota_exhausted, 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, provider_pool_quota_snapshot_updated_at, @@ -363,6 +366,34 @@ mod tests { .contains("profileArn=arn%3Aaws%3Asso%3A%3A%3Aprofile%2Fp-1")); } + #[test] + fn kiro_quota_request_rejects_region_url_injection() { + let spec = build_kiro_pool_quota_request( + "key-1", + &KiroPoolQuotaAuthInput { + authorization_value: "Bearer sensitive-access-token".to_string(), + api_region: "attacker.example/evil".to_string(), + kiro_version: "0.3.210".to_string(), + machine_id: "machine".to_string(), + profile_arn: None, + }, + ); + + assert_eq!( + spec.url, + "https://q.us-east-1.amazonaws.com/getUsageLimits?origin=AI_EDITOR&resourceType=AGENTIC_REQUEST&isEmailRequired=true" + ); + assert_eq!( + spec.headers.get("host").map(String::as_str), + Some("q.us-east-1.amazonaws.com") + ); + assert!(!spec.url.contains("attacker.example")); + assert!(!spec + .headers + .get("host") + .is_some_and(|value| value.contains("attacker.example"))); + } + #[test] fn chatgpt_web_quota_request_uses_default_base_url_when_empty() { let spec = build_chatgpt_web_pool_quota_request( @@ -379,7 +410,6 @@ mod tests { spec.headers.get("origin").map(String::as_str), Some("https://chatgpt.com") ); - assert!(spec.accept_invalid_certs); } #[test] @@ -882,6 +912,66 @@ mod tests { assert!(!available.quota_exhausted); } + #[test] + fn antigravity_tiered_model_quota_wins_over_display_only_group_windows() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "antigravity", + "exhausted": true, + "windows": [{ + "code": "model:gemini-3.7-flash-tiered", + "scope": "model", + "model": "gemini-3.7-flash-tiered", + "remaining_ratio": 0.906, + "used_ratio": 0.094, + "is_exhausted": false + }, { + "code": "group:0:3p-5h", + "scope": "quota_group", + "quota_group": "group:0", + "bucket_id": "3p-5h", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = + service.member_signals("antigravity", &key, None, Some("gemini-3.7-flash-tiered")); + + assert!(!signals.quota_exhausted); + assert!(!signals.quota_hard_blocked); + } + + #[test] + fn antigravity_tiered_variant_does_not_match_another_model_family() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "antigravity", + "exhausted": false, + "windows": [{ + "code": "model:gemini-3.7-pro-tiered", + "scope": "model", + "model": "gemini-3.7-pro-tiered", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = + service.member_signals("antigravity", &key, None, Some("gemini-3.7-flash-tiered")); + + assert!(!signals.quota_exhausted); + assert!(!signals.quota_hard_blocked); + } + #[test] fn codex_standard_and_spark_quota_families_are_independent() { let service = ProviderPoolService::with_builtin_adapters(); @@ -969,6 +1059,273 @@ mod tests { )); } + #[test] + fn model_quota_windows_are_isolated_without_provider_specific_names() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [ + { + "code": "alpha_short", + "quota_group": "alpha", + "model": "vendor-alpha-model", + "used_ratio": 1.0, + "is_exhausted": true + }, + { + "code": "alpha_long", + "quota_group": "alpha", + "model": "vendor-alpha-model", + "used_ratio": 0.2, + "is_exhausted": false + }, + { + "code": "beta_short", + "quota_group": "beta", + "model": "vendor-beta-model", + "used_ratio": 1.0, + "is_exhausted": true + } + ] + } + })); + + let alpha = service.member_signals("codex", &key, None, Some("vendor-alpha-model")); + let beta = service.member_signals("codex", &key, None, Some("vendor-beta-model")); + assert!( + !alpha.quota_exhausted, + "one available alpha window must keep it usable" + ); + assert!(beta.quota_exhausted); + assert!(!beta.quota_hard_blocked); + } + + #[test] + fn legacy_family_prefix_is_matched_by_model_tokens() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [ + { "code": "alpha_weekly", "used_ratio": 1.0, "is_exhausted": true }, + { "code": "default_weekly", "used_ratio": 0.1, "is_exhausted": false } + ] + } + })); + + let alpha = service.member_signals("codex", &key, None, Some("vendor-alpha-v2")); + assert!(alpha.quota_exhausted); + // An unrelated model has no identifiable bucket and therefore falls + // back to the account-level snapshot rather than guessing a family. + let unknown = service.member_signals("codex", &key, None, Some("vendor-gamma-v2")); + assert!(unknown.quota_exhausted); + } + + #[test] + fn compact_model_bucket_identity_matches_request_tokens() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [ + { + "code": "additional_0_primary_window", + "scope": "model", + "model": "spark", + "used_ratio": 0.0, + "is_exhausted": false + } + ] + } + })); + + let signals = service.member_signals("codex", &key, None, Some("gpt-5.3-codex-spark")); + assert!( + !signals.quota_exhausted, + "a compact bucket name should match a token in the selected model" + ); + } + + #[test] + fn model_family_token_matches_versioned_alias_without_hardcoding_name() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [{ + "code": "additional_0_primary_window", + "scope": "model", + "model": "gpt-5.3-codex-spark", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = service.member_signals("codex", &key, None, Some("gpt-5.4-codex-spark")); + assert!(signals.quota_exhausted); + } + + #[test] + fn model_only_snapshot_does_not_poison_account_fallback() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [{ + "code": "model:vendor-alpha", + "scope": "model", + "model": "vendor-alpha", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = service.member_signals("codex", &key, None, None); + assert!( + !signals.quota_exhausted, + "account-level inspection must ignore model-only buckets" + ); + } + + #[test] + fn model_map_quota_is_resolved_without_materialized_windows() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "grok", + "exhausted": false, + "quota_by_model": { + "vendor-alpha": { + "remaining": 0.0, + "total": 10.0 + }, + "vendor-beta": { + "remaining": 5.0, + "total": 10.0 + } + } + } + })); + + let alpha = service.member_signals("grok", &key, None, Some("vendor-alpha")); + let beta = service.member_signals("grok", &key, None, Some("vendor-beta")); + assert!(alpha.quota_exhausted); + assert!(!beta.quota_exhausted); + } + + #[test] + fn raw_provider_metadata_model_bucket_overrides_account_snapshot() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(Some(json!({ + "codex": { + "updated_at": 1_700_000_001u64, + "additional_quota_windows": [{ + "scope": "model", + "model": "future-spark", + "used_ratio": 0.0, + "is_exhausted": false + }] + } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "exhausted": true, + "windows": [{ + "code": "primary", + "scope": "account", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = service.member_signals("codex", &key, None, Some("future-spark")); + assert!( + !signals.quota_exhausted, + "a newer model bucket in provider metadata must not inherit an exhausted account bucket" + ); + } + + #[test] + fn model_quota_availability_suppresses_account_hard_block() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "allowed": false, + "limit_reached": true, + "exhausted": true, + "windows": [{ + "scope": "model", + "model": "future-model", + "used_ratio": 0.0, + "is_exhausted": false + }] + } + })); + + let signals = service.member_signals("codex", &key, None, Some("future-model")); + assert!(!signals.quota_exhausted); + assert!(!signals.quota_hard_blocked); + } + + #[test] + fn newer_model_quota_observation_wins_over_stale_snapshot() { + let service = ProviderPoolService::with_builtin_adapters(); + let mut key = sample_key(Some(json!({ + "codex": { + "updated_at": 1_700_000_200u64, + "additional_quota_windows": [{ + "scope": "model", + "model": "future-model", + "used_ratio": 0.0, + "is_exhausted": false + }] + } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "version": 2, + "provider_type": "codex", + "observed_at": 1_700_000_100u64, + "exhausted": true, + "windows": [{ + "scope": "model", + "model": "future-model", + "used_ratio": 1.0, + "is_exhausted": true + }] + } + })); + + let signals = service.member_signals("codex", &key, None, Some("future-model")); + assert!(!signals.quota_exhausted); + } + #[test] fn codex_explicit_quota_block_is_hard_until_reset() { let now = std::time::SystemTime::now() diff --git a/crates/aether-provider/pool/src/provider.rs b/crates/aether-provider/pool/src/provider.rs index 3cb16f13e..47cace82e 100644 --- a/crates/aether-provider/pool/src/provider.rs +++ b/crates/aether-provider/pool/src/provider.rs @@ -7,8 +7,9 @@ use serde_json::{Map, Value}; use crate::capability::{ProviderPoolCapabilities, ProviderPoolCapability}; use crate::plan::{derive_plan_tier, normalize_provider_plan_tier}; use crate::quota::{ - provider_pool_account_blocked, provider_pool_quota_reset_seconds, - provider_pool_quota_snapshot_exhausted_decision, provider_pool_quota_usage_ratio, + provider_pool_account_blocked, provider_pool_model_quota_exhausted, + provider_pool_quota_reset_seconds, provider_pool_quota_snapshot_exhausted_decision, + provider_pool_quota_usage_ratio, }; #[derive(Debug, Clone)] @@ -16,6 +17,10 @@ pub struct ProviderPoolMemberInput<'a> { pub provider_type: &'a str, pub key: &'a StoredProviderCatalogKey, pub auth_config: Option<&'a Map>, + /// The provider-side model selected for the current request. Quota + /// snapshots may contain several independent windows, so adapters use + /// this value to select the applicable bucket instead of treating the + /// whole account as one quota. pub provider_model_name: Option<&'a str>, } @@ -71,6 +76,11 @@ pub trait ProviderPoolAdapter: Send + Sync { } 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; + } provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type) .unwrap_or(false) } diff --git a/crates/aether-provider/pool/src/providers/antigravity.rs b/crates/aether-provider/pool/src/providers/antigravity.rs index 24d9bb797..fd8973c82 100644 --- a/crates/aether-provider/pool/src/providers/antigravity.rs +++ b/crates/aether-provider/pool/src/providers/antigravity.rs @@ -12,6 +12,8 @@ use crate::quota::provider_pool_model_quota_exhausted; use crate::quota_refresh::ProviderPoolQuotaRequestSpec; pub const ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH: &str = "/v1internal:fetchAvailableModels"; +pub const ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH: &str = + "/v1internal:retrieveUserQuotaSummary"; #[derive(Debug, Clone, Default)] pub struct AntigravityProviderPoolAdapter; @@ -89,6 +91,86 @@ pub fn build_antigravity_pool_quota_request( client_api_format: "gemini:generate_content".to_string(), provider_api_format: "antigravity:fetch_available_models".to_string(), model_name: Some("fetchAvailableModels".to_string()), - accept_invalid_certs: false, + } +} + +pub fn build_antigravity_pool_quota_summary_request( + key_id: &str, + endpoint_base_url: &str, + authorization: (String, String), + project_id: Option<&str>, + mut identity_headers: BTreeMap, +) -> ProviderPoolQuotaRequestSpec { + let mut headers = std::mem::take(&mut identity_headers); + headers.insert("authorization".to_string(), authorization.1); + headers.insert("content-type".to_string(), "application/json".to_string()); + headers.insert("accept".to_string(), "application/json".to_string()); + headers + .entry("user-agent".to_string()) + .or_insert_with(|| "antigravity".to_string()); + + let json_body = project_id + .map(str::trim) + .filter(|project_id| !project_id.is_empty()) + .map_or_else(|| json!({}), |project_id| json!({ "project": project_id })); + + ProviderPoolQuotaRequestSpec { + request_id: format!("antigravity-quota-summary:{key_id}"), + provider_name: "antigravity".to_string(), + quota_kind: "antigravity".to_string(), + method: "POST".to_string(), + url: format!( + "{}{}", + endpoint_base_url.trim_end_matches('/'), + ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH + ), + headers, + content_type: Some("application/json".to_string()), + json_body: Some(json_body), + client_api_format: "gemini:generate_content".to_string(), + provider_api_format: "antigravity:retrieve_user_quota_summary".to_string(), + model_name: Some("retrieveUserQuotaSummary".to_string()), + } +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use serde_json::json; + + use super::{ + build_antigravity_pool_quota_summary_request, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, + }; + + #[test] + fn grouped_quota_request_can_retry_without_project_on_the_same_endpoint() { + let with_project = build_antigravity_pool_quota_summary_request( + "key-1", + "https://daily-cloudcode-pa.googleapis.com/", + ("authorization".to_string(), "Bearer token".to_string()), + Some("project-1"), + BTreeMap::new(), + ); + let without_project = build_antigravity_pool_quota_summary_request( + "key-1", + "https://daily-cloudcode-pa.googleapis.com/", + ("authorization".to_string(), "Bearer token".to_string()), + None, + BTreeMap::new(), + ); + + assert_eq!( + with_project.url, + format!( + "https://daily-cloudcode-pa.googleapis.com{ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH}" + ) + ); + assert_eq!( + with_project.json_body, + Some(json!({"project": "project-1"})) + ); + assert_eq!(without_project.url, with_project.url); + assert_eq!(without_project.json_body, Some(json!({}))); } } diff --git a/crates/aether-provider/pool/src/providers/chatgpt_web.rs b/crates/aether-provider/pool/src/providers/chatgpt_web.rs index 536d5337c..c3fd4a54a 100644 --- a/crates/aether-provider/pool/src/providers/chatgpt_web.rs +++ b/crates/aether-provider/pool/src/providers/chatgpt_web.rs @@ -11,8 +11,9 @@ use crate::provider::{ }; use crate::quota::{ provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64, - provider_pool_metadata_bucket, provider_pool_quota_snapshot_exhausted_decision, - provider_pool_reset_deadline_elapsed, provider_pool_timestamp_unix_secs, + 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; @@ -40,6 +41,11 @@ impl ProviderPoolAdapter for ChatGptWebProviderPoolAdapter { } 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) { @@ -138,7 +144,6 @@ pub fn build_chatgpt_web_pool_quota_request( client_api_format: "openai:image".to_string(), provider_api_format: "chatgpt_web:conversation_init".to_string(), model_name: Some("chatgpt-web-conversation-init".to_string()), - accept_invalid_certs: true, } } diff --git a/crates/aether-provider/pool/src/providers/codex.rs b/crates/aether-provider/pool/src/providers/codex.rs index 0086e813e..f7133d281 100644 --- a/crates/aether-provider/pool/src/providers/codex.rs +++ b/crates/aether-provider/pool/src/providers/codex.rs @@ -215,7 +215,6 @@ pub fn build_codex_pool_quota_request( client_api_format: "openai:responses".to_string(), provider_api_format: "openai:responses".to_string(), model_name: Some("codex-wham-usage".to_string()), - accept_invalid_certs: false, }) } @@ -239,7 +238,6 @@ pub fn build_codex_pool_reset_credits_request( client_api_format: "openai:responses".to_string(), provider_api_format: "openai:responses".to_string(), model_name: Some("codex-wham-reset-credits".to_string()), - accept_invalid_certs: false, }) } @@ -273,7 +271,6 @@ pub fn build_codex_pool_reset_credit_consume_request( client_api_format: "openai:responses".to_string(), provider_api_format: "openai:responses".to_string(), model_name: Some("codex-wham-reset-credit-consume".to_string()), - accept_invalid_certs: false, }) } diff --git a/crates/aether-provider/pool/src/providers/gemini_cli.rs b/crates/aether-provider/pool/src/providers/gemini_cli.rs index f4bfbd3e6..50d267dd7 100644 --- a/crates/aether-provider/pool/src/providers/gemini_cli.rs +++ b/crates/aether-provider/pool/src/providers/gemini_cli.rs @@ -73,6 +73,5 @@ pub fn build_gemini_cli_pool_quota_request( client_api_format: "gemini:generate_content".to_string(), provider_api_format: "gemini_cli:retrieve_user_quota".to_string(), model_name: Some("retrieveUserQuota".to_string()), - accept_invalid_certs: false, } } diff --git a/crates/aether-provider/pool/src/providers/grok.rs b/crates/aether-provider/pool/src/providers/grok.rs index c2bb4dc2c..460e927cd 100644 --- a/crates/aether-provider/pool/src/providers/grok.rs +++ b/crates/aether-provider/pool/src/providers/grok.rs @@ -8,8 +8,9 @@ use crate::provider::{ }; use crate::quota::{ provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64, - provider_pool_metadata_bucket, provider_pool_quota_snapshot_exhausted_decision, - provider_pool_reset_deadline_elapsed, provider_pool_timestamp_unix_secs, + 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, }; pub const GROK_QUOTA_WINDOWS_BASIC: &[(&str, &str)] = &[("quota_fast", "fast")]; @@ -44,6 +45,11 @@ impl ProviderPoolAdapter for GrokProviderPoolAdapter { } 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) { diff --git a/crates/aether-provider/pool/src/providers/kiro.rs b/crates/aether-provider/pool/src/providers/kiro.rs index 72c175315..27bfc8c3b 100644 --- a/crates/aether-provider/pool/src/providers/kiro.rs +++ b/crates/aether-provider/pool/src/providers/kiro.rs @@ -12,15 +12,16 @@ use crate::provider::{ }; use crate::quota::{ provider_pool_current_unix_secs, provider_pool_json_f64, provider_pool_metadata_bucket, - provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed, - provider_pool_timestamp_unix_secs, + 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; +use aether_provider_transport::kiro::normalize_kiro_region; pub const KIRO_USAGE_LIMITS_PATH: &str = "/getUsageLimits"; pub const KIRO_USAGE_SDK_VERSION: &str = "1.0.0"; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct KiroPoolQuotaAuthInput { pub authorization_value: String, pub api_region: String, @@ -29,6 +30,19 @@ pub struct KiroPoolQuotaAuthInput { pub profile_arn: Option, } +impl std::fmt::Debug for KiroPoolQuotaAuthInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("KiroPoolQuotaAuthInput") + .field("authorization_value", &"[REDACTED]") + .field("api_region", &self.api_region) + .field("kiro_version", &self.kiro_version) + .field("machine_id", &self.machine_id) + .field("profile_arn", &self.profile_arn) + .finish() + } +} + #[derive(Debug, Clone, Default)] pub struct KiroProviderPoolAdapter; @@ -46,6 +60,11 @@ impl ProviderPoolAdapter for KiroProviderPoolAdapter { } 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) { @@ -129,17 +148,11 @@ pub fn build_kiro_pool_quota_request( client_api_format: "claude:messages".to_string(), provider_api_format: "kiro:usage".to_string(), model_name: Some("kiro-usage-limits".to_string()), - accept_invalid_certs: false, } } fn normalize_region(value: &str) -> &str { - let value = value.trim(); - if value.is_empty() { - "us-east-1" - } else { - value - } + normalize_kiro_region(value) } fn normalize_kiro_version(value: &str) -> &str { @@ -175,3 +188,23 @@ pub(crate) fn quota_exhausted_from_bucket(bucket: &Map) -> bool { _ => false, } } + +#[cfg(test)] +mod tests { + use super::KiroPoolQuotaAuthInput; + + #[test] + fn quota_auth_debug_output_redacts_authorization() { + let input = KiroPoolQuotaAuthInput { + authorization_value: "Bearer kiro-secret-canary".to_string(), + api_region: "us-east-1".to_string(), + kiro_version: "1.0.0".to_string(), + machine_id: "machine-1".to_string(), + profile_arn: None, + }; + + let debug = format!("{input:?}"); + assert!(!debug.contains("kiro-secret-canary")); + assert!(debug.contains("[REDACTED]")); + } +} diff --git a/crates/aether-provider/pool/src/providers/mod.rs b/crates/aether-provider/pool/src/providers/mod.rs index 0239dde79..72734c81e 100644 --- a/crates/aether-provider/pool/src/providers/mod.rs +++ b/crates/aether-provider/pool/src/providers/mod.rs @@ -10,7 +10,8 @@ pub mod windsurf; pub use antigravity::AntigravityProviderPoolAdapter; pub use antigravity::{ - build_antigravity_pool_quota_request, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, + build_antigravity_pool_quota_request, build_antigravity_pool_quota_summary_request, + ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, }; pub use chatgpt_web::ChatGptWebProviderPoolAdapter; pub use chatgpt_web::{ diff --git a/crates/aether-provider/pool/src/providers/windsurf.rs b/crates/aether-provider/pool/src/providers/windsurf.rs index 112cbbb68..ad152324a 100644 --- a/crates/aether-provider/pool/src/providers/windsurf.rs +++ b/crates/aether-provider/pool/src/providers/windsurf.rs @@ -13,7 +13,8 @@ use crate::provider::{ }; use crate::quota::{ provider_pool_json_bool, provider_pool_json_f64, provider_pool_member_quota_snapshot, - provider_pool_metadata_bucket, provider_pool_quota_snapshot_exhausted_decision, + provider_pool_metadata_bucket, provider_pool_model_quota_exhausted, + provider_pool_quota_snapshot_exhausted_decision, }; use crate::quota_refresh::ProviderPoolQuotaRequestSpec; @@ -50,6 +51,11 @@ impl ProviderPoolAdapter for WindsurfProviderPoolAdapter { } 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 windsurf_quota_snapshot_hard_exhausted(input.key, input.provider_type) { return true; } @@ -187,7 +193,6 @@ fn build_windsurf_connect_rpc_request( client_api_format: "openai:chat".to_string(), provider_api_format: provider_api_format.to_string(), model_name: Some(model_name.to_string()), - accept_invalid_certs: false, } } diff --git a/crates/aether-provider/pool/src/quota.rs b/crates/aether-provider/pool/src/quota.rs index 5f9c05e6b..9d18d506c 100644 --- a/crates/aether-provider/pool/src/quota.rs +++ b/crates/aether-provider/pool/src/quota.rs @@ -32,54 +32,22 @@ pub fn provider_pool_key_quota_hard_blocked( }) } -pub(crate) fn provider_pool_model_quota_exhausted( +/// Model-aware hard-block lookup for pre-scheduler candidate filtering. Some +/// providers expose permanent account flags alongside independent model +/// buckets; adapters can suppress the account flag when the selected model +/// has its own usable quota. +pub fn provider_pool_key_model_quota_hard_blocked( key: &StoredProviderCatalogKey, provider_type: &str, provider_model_name: &str, -) -> Option { - let quota_snapshot = provider_pool_member_quota_snapshot(key, provider_type)?; - let windows = quota_snapshot.get("windows")?.as_array()?; - let normalized_provider = provider_type.trim().to_ascii_lowercase(); - let normalized_model = provider_model_name.trim().to_ascii_lowercase(); - - let matches_window = |window: &Map| { - let code = window - .get("code") - .and_then(Value::as_str) - .unwrap_or_default() - .trim() - .to_ascii_lowercase(); - if normalized_provider == "codex" { - let spark_model = normalized_model.contains("spark"); - return code.starts_with("spark_") == spark_model; - } - if normalized_provider == "antigravity" { - return window - .get("model") - .and_then(Value::as_str) - .is_some_and(|model| model.trim().eq_ignore_ascii_case(&normalized_model)); - } - false - }; - - let matching_windows = windows - .iter() - .filter_map(Value::as_object) - .filter(|window| matches_window(window)) - .collect::>(); - if matching_windows.is_empty() { - return None; - } - - let now_unix_secs = provider_pool_current_unix_secs(); - let snapshot_observed_at = provider_pool_timestamp_unix_secs(quota_snapshot.get("observed_at")) - .or_else(|| provider_pool_timestamp_unix_secs(quota_snapshot.get("updated_at"))); - Some(matching_windows.iter().any(|window| { - provider_pool_quota_window_is_exhausted(window) - && !now_unix_secs.is_some_and(|now| { - provider_pool_reset_deadline_elapsed(window, snapshot_observed_at, now) - }) - })) +) -> bool { + let adapter = ProviderPoolService::with_builtin_adapters().adapter(provider_type); + adapter.quota_hard_blocked(&ProviderPoolMemberInput { + provider_type, + key, + auth_config: None, + provider_model_name: Some(provider_model_name), + }) } pub fn provider_pool_member_quota_snapshot<'a>( @@ -96,6 +64,421 @@ pub fn provider_pool_member_quota_snapshot<'a>( .then_some(quota_snapshot) } +/// Resolve exhaustion for the quota bucket applicable to one provider model. +/// +/// Providers are free to expose quota windows in different shapes. Newer +/// snapshots should put an explicit `model`/`models` (or `quota_group`) on a +/// window; legacy Codex snapshots use a family prefix in `code` (for example +/// `spark_5h`). We deliberately do not name any product or model here: the +/// resolver compares the metadata supplied by the provider with the selected +/// model and only falls back to account-level exhaustion when no model bucket +/// can be identified. +pub(crate) fn provider_pool_model_quota_exhausted( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: &str, +) -> Option { + let requested = provider_pool_identifier_tokens(provider_model_name); + if requested.is_empty() { + return None; + } + + // Prefer the materialized status snapshot, but also inspect the raw + // provider metadata. A quota refresh and a request can race, leaving the + // latter newer than the snapshot; resolving both avoids falling back to an + // account-wide signal and incorrectly blocking an unrelated model bucket. + let sources = [ + provider_pool_member_quota_snapshot(key, provider_type), + provider_pool_metadata_bucket(key.upstream_metadata.as_ref(), provider_type), + ]; + let mut resolved = None::<(Option, bool)>; + for source in sources.into_iter().flatten() { + let windows = provider_pool_collect_quota_windows(source); + 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) + }) + .collect::>(); + if !model_matches.is_empty() { + let exhausted = + provider_pool_explicit_model_windows_exhausted(model_matches, observed_at); + if resolved.is_none() + || provider_pool_should_replace_model_quota_resolution( + resolved.as_ref().and_then(|(observed_at, _)| *observed_at), + observed_at, + ) + { + resolved = Some((observed_at, exhausted)); + } + continue; + } + + // Legacy snapshots may not carry a model field. Match an opaque + // family token (the prefix before `_`/`:` in `code`, or an explicit + // family key) against the model's tokens. This keeps independent + // windows isolated without baking in names such as "spark". + let family_matches = windows + .iter() + .filter(|window| provider_pool_window_family_matches_model(window, &requested)) + .collect::>(); + if !family_matches.is_empty() { + let exhausted = provider_pool_any_window_exhausted(family_matches, observed_at); + if resolved.is_none() + || provider_pool_should_replace_model_quota_resolution( + resolved.as_ref().and_then(|(observed_at, _)| *observed_at), + observed_at, + ) + { + resolved = Some((observed_at, exhausted)); + } + 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 + // window names; opaque family names such as `alpha_weekly` must not + // accidentally make an unrelated model look schedulable. + let generic_matches = windows + .iter() + .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); + if resolved.is_none() + || provider_pool_should_replace_model_quota_resolution( + resolved.as_ref().and_then(|(observed_at, _)| *observed_at), + observed_at, + ) + { + resolved = Some((observed_at, exhausted)); + } + } + } + + resolved.map(|(_, exhausted)| exhausted) +} + +fn provider_pool_should_replace_model_quota_resolution( + previous_observed_at: Option, + next_observed_at: Option, +) -> bool { + match (previous_observed_at, next_observed_at) { + (Some(previous), Some(next)) => next >= previous, + (None, Some(_)) => true, + (Some(_), None) => false, + // Preserve source order when neither side carries freshness metadata; + // the materialized status snapshot is preferred over raw metadata. + (None, None) => false, + } +} + +/// Public adapter-independent model quota lookup used by schedulers that need +/// to prefilter candidates before constructing provider-pool signals. +pub fn provider_pool_key_model_quota_exhausted( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: &str, +) -> Option { + provider_pool_model_quota_exhausted(key, provider_type, provider_model_name) +} + +fn provider_pool_explicit_model_windows_exhausted( + windows: Vec<&Map>, + snapshot_observed_at: Option, +) -> bool { + let now_unix_secs = provider_pool_current_unix_secs(); + + // Explicit model buckets represent independent windows for one model. The + // model remains usable while at least one of those windows still has + // capacity, so exhaustion is reported only when all are exhausted. + windows.iter().all(|window| { + provider_pool_quota_window_is_exhausted(window) + && !now_unix_secs.is_some_and(|now| { + provider_pool_reset_deadline_elapsed(window, snapshot_observed_at, now) + }) + }) +} + +fn provider_pool_any_window_exhausted( + windows: Vec<&Map>, + snapshot_observed_at: Option, +) -> bool { + let now_unix_secs = provider_pool_current_unix_secs(); + windows.iter().any(|window| { + provider_pool_quota_window_is_exhausted(window) + && !now_unix_secs.is_some_and(|now| { + provider_pool_reset_deadline_elapsed(window, snapshot_observed_at, now) + }) + }) +} + +fn provider_pool_window_is_generic(window: &Map) -> bool { + if provider_pool_window_has_explicit_model(window) + || window + .get("scope") + .and_then(Value::as_str) + .is_some_and(|scope| scope.trim().eq_ignore_ascii_case("model")) + { + return false; + } + + let code = window + .get("code") + .and_then(Value::as_str) + .unwrap_or_default(); + let family = code + .split_once(['_', ':', '/']) + .map(|(prefix, _)| prefix) + .unwrap_or(code) + .trim() + .to_ascii_lowercase(); + if family.is_empty() { + return window + .get("scope") + .and_then(Value::as_str) + .is_some_and(|scope| scope.trim().eq_ignore_ascii_case("account")); + } + [ + "weekly", + "5h", + "daily", + "monthly", + "primary", + "secondary", + "account", + "quota", + "window", + "rate", + "reset", + ] + .contains(&family.as_str()) +} + +fn provider_pool_window_explicitly_matches_model( + window: &Map, + requested_model: &str, +) -> bool { + let requested = provider_pool_normalize_identifier(requested_model); + let requested_tokens = provider_pool_identifier_tokens(requested_model); + if requested.is_empty() { + return false; + } + [ + "model", + "model_name", + "model_id", + "quota_model", + "quota_model_name", + "target_model", + "limit_name", + ] + .iter() + .filter_map(|key| window.get(*key)) + .any(|value| match value { + Value::String(value) => { + provider_pool_identifiers_match(&requested, value, &requested_tokens) + } + Value::Array(values) => values + .iter() + .filter_map(Value::as_str) + .any(|value| provider_pool_identifiers_match(&requested, value, &requested_tokens)), + _ => false, + }) || ["models", "model_ids"] + .iter() + .filter_map(|key| window.get(*key).and_then(Value::as_array)) + .flatten() + .filter_map(Value::as_str) + .any(|value| provider_pool_identifiers_match(&requested, value, &requested_tokens)) +} + +fn provider_pool_window_family_matches_model( + window: &Map, + requested_tokens: &std::collections::BTreeSet, +) -> bool { + let explicit_scope = window + .get("scope") + .and_then(Value::as_str) + .map(|scope| scope.trim().to_ascii_lowercase()); + // A model-scoped window without an explicit model must not accidentally + // match a token from its opaque code. + if explicit_scope.as_deref() == Some("model") || provider_pool_window_has_explicit_model(window) + { + return false; + } + + let mut families = Vec::new(); + for key in ["quota_group", "quota_family", "family", "bucket"] { + if let Some(value) = window.get(key).and_then(Value::as_str) { + families.push(value.to_string()); + } + } + if let Some(code) = window.get("code").and_then(Value::as_str) { + let code = code.trim(); + if let Some((prefix, _)) = code.split_once(['_', ':', '/']) { + families.push(prefix.to_string()); + } + } + families.into_iter().any(|family| { + let normalized = provider_pool_normalize_identifier(&family); + if normalized.is_empty() + || [ + "account", + "quota", + "window", + "primary", + "secondary", + "rate", + "reset", + ] + .iter() + .any(|generic| normalized == *generic) + { + return false; + } + requested_tokens.iter().any(|token| { + token.len() >= 3 + && (token == &normalized + || token.contains(&normalized) + || normalized.contains(token)) + }) + }) +} + +fn provider_pool_identifiers_match( + requested: &str, + candidate: &str, + requested_tokens: &std::collections::BTreeSet, +) -> bool { + let candidate_tokens = provider_pool_identifier_tokens(candidate); + let candidate = provider_pool_normalize_identifier(candidate); + if candidate.is_empty() { + return false; + } + requested == candidate + || candidate_tokens + .iter() + .any(|token| { + requested_tokens.contains(token) && provider_pool_is_specific_model_token(token) + }) + // Handle compact upstream identifiers such as `spark` embedded in a + // provider model name (`vendor-codex-spark`) while avoiding accidental + // one/two-character matches. + || (candidate.len() >= 4 + && requested.len() >= 6 + && (requested.contains(&candidate) || candidate.contains(requested))) +} + +fn provider_pool_is_specific_model_token(token: &str) -> bool { + token.len() >= 4 + && !token.chars().all(|character| character.is_ascii_digit()) + && ![ + "auto", + "base", + "claude", + "codex", + "default", + "fast", + "flash", + "free", + "gemini", + "gpt", + "latest", + "mini", + "model", + "plus", + "pro", + "reasoning", + "team", + "think", + "thinking", + "tiered", + "vendor", + ] + .contains(&token) +} + +fn provider_pool_window_has_explicit_model(window: &Map) -> bool { + [ + "model", + "model_name", + "model_id", + "quota_model", + "quota_model_name", + "target_model", + "models", + "model_ids", + ] + .iter() + .any(|key| match window.get(*key) { + Some(Value::String(value)) => !value.trim().is_empty(), + Some(Value::Array(values)) => values + .iter() + .any(|value| value.as_str().is_some_and(|value| !value.trim().is_empty())), + _ => false, + }) +} + +fn provider_pool_normalize_identifier(value: &str) -> String { + value + .trim() + .to_ascii_lowercase() + .chars() + .filter(|character| character.is_ascii_alphanumeric()) + .collect() +} + +fn provider_pool_identifier_tokens(value: &str) -> std::collections::BTreeSet { + value + .split(|character: char| !character.is_ascii_alphanumeric()) + .map(|token| token.trim().to_ascii_lowercase()) + .filter(|token| token.len() >= 3) + .collect() +} + +/// Materialize quota windows from the small set of shapes emitted by current +/// and legacy adapters. Model maps (`quota_by_model`/`models`) are converted +/// to the same window representation used by status snapshots, with the map +/// key retained as the model identity. Keeping this normalization here means +/// provider adapters do not need to grow model-specific quota code whenever an +/// upstream introduces another independent bucket. +fn provider_pool_collect_quota_windows(source: &Map) -> Vec> { + let mut windows = Vec::new(); + for key in ["windows", "additional_quota_windows"] { + if let Some(values) = source.get(key).and_then(Value::as_array) { + windows.extend(values.iter().filter_map(Value::as_object).cloned()); + } + } + for key in ["quota_by_model", "models", "model_quotas"] { + let Some(models) = source.get(key).and_then(Value::as_object) else { + continue; + }; + for (model_name, item) in models { + let Some(item) = item.as_object() else { + continue; + }; + let mut window = item.clone(); + window + .entry("model".to_string()) + .or_insert_with(|| json!(model_name)); + window + .entry("scope".to_string()) + .or_insert_with(|| json!("model")); + window + .entry("code".to_string()) + .or_insert_with(|| json!(format!("model:{model_name}"))); + windows.push(window); + } + } + windows +} + pub fn provider_pool_quota_snapshot_updated_at( key: &StoredProviderCatalogKey, provider_type: &str, @@ -249,12 +632,54 @@ pub(crate) fn provider_pool_reset_deadline_elapsed( fn provider_pool_quota_window_is_exhausted(window: &Map) -> bool { provider_pool_json_bool(window.get("is_exhausted")) + .or_else(|| provider_pool_json_bool(window.get("exhausted"))) .or_else(|| { - provider_pool_json_f64(window.get("used_ratio")).map(|value| value >= 1.0 - 1e-6) + provider_pool_json_f64( + window + .get("used_ratio") + .or_else(|| window.get("usage_ratio")), + ) + .map(|value| value >= 1.0 - 1e-6) + }) + .or_else(|| { + provider_pool_json_f64(window.get("used_percent")).map(|value| value >= 100.0 - 1e-6) + }) + .or_else(|| { + provider_pool_json_f64( + window + .get("remaining_ratio") + .or_else(|| window.get("remaining_fraction")), + ) + .map(|value| value <= 1e-6) + }) + .or_else(|| { + provider_pool_json_f64(window.get("remaining_percent")).map(|value| value <= 1e-6) + }) + .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(remaining <= 0.0) }) .unwrap_or(false) } +fn provider_pool_window_is_model_scoped(window: &Map) -> bool { + window + .get("scope") + .and_then(Value::as_str) + .is_some_and(|scope| scope.trim().eq_ignore_ascii_case("model")) + || provider_pool_window_has_explicit_model(window) +} + fn provider_pool_quota_snapshot_matches_provider( quota_snapshot: &Map, provider_type: &str, @@ -289,6 +714,18 @@ fn provider_pool_quota_snapshot_matches_provider( .get("windows") .and_then(Value::as_array) .is_some_and(|windows| !windows.is_empty()) + || quota_snapshot + .get("additional_quota_windows") + .and_then(Value::as_array) + .is_some_and(|windows| !windows.is_empty()) + || ["quota_by_model", "models", "model_quotas"] + .iter() + .any(|key| { + quota_snapshot + .get(*key) + .and_then(Value::as_object) + .is_some_and(|models| !models.is_empty()) + }) || quota_snapshot .get("credits") .and_then(Value::as_object) @@ -317,16 +754,28 @@ pub(crate) fn provider_pool_quota_snapshot_exhausted_decision( provider_pool_timestamp_unix_secs(quota_snapshot.get("observed_at")) .or_else(|| provider_pool_timestamp_unix_secs(quota_snapshot.get("updated_at"))); - if let Some(windows) = quota_snapshot - .get("windows") - .and_then(Value::as_array) - .filter(|windows| !windows.is_empty()) - { + let materialized_windows = provider_pool_collect_quota_windows(quota_snapshot); + if !materialized_windows.is_empty() { + // Model-scoped windows are evaluated only when the request model + // is known. They must not turn the account-level fallback into + // an exhausted state for unrelated models. + let account_scoped_windows = materialized_windows + .iter() + .filter(|window| !provider_pool_window_is_model_scoped(window)) + .collect::>(); + // A snapshot containing only model-scoped buckets has no + // account-wide signal to apply when the caller did not provide a + // model name (for example, an admin status listing). Do not let a + // single exhausted model poison every sibling bucket. + if account_scoped_windows.is_empty() { + return Some(false); + } + let windows = account_scoped_windows; let mut saw_exhausted_window = false; let mut saw_active_exhausted_window = false; let mut windows_max_ratio = None::; - for window in windows.iter().filter_map(Value::as_object) { + for window in windows.iter() { if let Some(ratio) = provider_pool_json_f64(window.get("used_ratio")) { windows_max_ratio = Some(windows_max_ratio.map_or(ratio, |current| current.max(ratio))); diff --git a/crates/aether-provider/pool/src/quota_refresh.rs b/crates/aether-provider/pool/src/quota_refresh.rs index 727b09021..0f34b4cc7 100644 --- a/crates/aether-provider/pool/src/quota_refresh.rs +++ b/crates/aether-provider/pool/src/quota_refresh.rs @@ -1,8 +1,10 @@ use std::collections::BTreeMap; +use std::fmt; +use aether_contracts::redact_url_for_debug; use serde_json::Value; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct ProviderPoolQuotaRequestSpec { pub request_id: String, pub provider_name: String, @@ -15,5 +17,30 @@ pub struct ProviderPoolQuotaRequestSpec { pub client_api_format: String, pub provider_api_format: String, pub model_name: Option, - pub accept_invalid_certs: bool, +} + +impl fmt::Debug for ProviderPoolQuotaRequestSpec { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ProviderPoolQuotaRequestSpec") + .field("request_id", &self.request_id) + .field("provider_name", &self.provider_name) + .field("quota_kind", &self.quota_kind) + .field("method", &self.method) + .field("url", &redact_url_for_debug(&self.url)) + .field("header_names", &self.headers.keys().collect::>()) + .field("content_type", &self.content_type) + .field("has_json_body", &self.json_body.is_some()) + .field( + "json_body_bytes", + &self + .json_body + .as_ref() + .and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())), + ) + .field("client_api_format", &self.client_api_format) + .field("provider_api_format", &self.provider_api_format) + .field("model_name", &self.model_name) + .finish() + } } diff --git a/crates/aether-provider/transport/Cargo.toml b/crates/aether-provider/transport/Cargo.toml index 88cdd5834..cf94aa525 100644 --- a/crates/aether-provider/transport/Cargo.toml +++ b/crates/aether-provider/transport/Cargo.toml @@ -11,6 +11,7 @@ aether-ai-formats.workspace = true aether-contracts.workspace = true aether-crypto.workspace = true aether-data-contracts.workspace = true +aether-http.workspace = true aether-oauth.workspace = true aether-runtime-state.workspace = true aether-video-tasks-core.workspace = true @@ -22,7 +23,6 @@ ed25519-dalek.workspace = true http.workspace = true regex.workspace = true reqwest.workspace = true -rsa = "0.9.10" serde.workspace = true serde_json.workspace = true sha2 = { workspace = true, features = ["oid"] } @@ -33,4 +33,5 @@ url.workspace = true uuid.workspace = true [dev-dependencies] +aws-lc-rs.workspace = true axum = { version = "0.8", features = ["ws"] } diff --git a/crates/aether-provider/transport/src/agent_identity.rs b/crates/aether-provider/transport/src/agent_identity.rs index 82313e77a..abdb9d806 100644 --- a/crates/aether-provider/transport/src/agent_identity.rs +++ b/crates/aether-provider/transport/src/agent_identity.rs @@ -91,6 +91,35 @@ struct AgentIdentityAssertionEnvelope { signature: String, } +const MAX_AGENT_IDENTITY_PRIVATE_KEY_DER_BYTES: usize = 4 * 1024; +const MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES: usize = 64 * 1024; +const MAX_AGENT_IDENTITY_SIGNATURE_BYTES: usize = 128; +const MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES: usize = 64 * 1024; + +fn maximum_base64_len_for_decoded_limit(limit: usize) -> usize { + limit + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4) +} + +fn decode_standard_base64_with_limit(value: &str, limit: usize) -> Option> { + if value.len() > maximum_base64_len_for_decoded_limit(limit) { + return None; + } + let decoded = STANDARD.decode(value).ok()?; + (decoded.len() <= limit).then_some(decoded) +} + +fn decode_urlsafe_base64_with_limit(value: &str, limit: usize) -> Option> { + if value.len() > maximum_base64_len_for_decoded_limit(limit) { + return None; + } + let decoded = URL_SAFE_NO_PAD.decode(value).ok()?; + (decoded.len() <= limit).then_some(decoded) +} + #[derive(Debug, Deserialize)] struct AgentTaskRegistrationResponse { #[serde(default)] @@ -288,7 +317,9 @@ pub fn codex_agent_identity_authorization_matches_transport( let Some(encoded) = encoded_agent_identity_assertion(authorization) else { return false; }; - let Ok(envelope_bytes) = URL_SAFE_NO_PAD.decode(encoded) else { + let Some(envelope_bytes) = + decode_urlsafe_base64_with_limit(encoded, MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES) + else { return false; }; let Ok(envelope) = serde_json::from_slice::(&envelope_bytes) @@ -310,7 +341,10 @@ pub fn codex_agent_identity_authorization_matches_transport( if runtime_id != credentials.runtime_id || task_id != current_task_id || timestamp.is_empty() { return false; } - let Ok(signature_bytes) = STANDARD.decode(envelope.signature.trim()) else { + let Some(signature_bytes) = decode_standard_base64_with_limit( + envelope.signature.trim(), + MAX_AGENT_IDENTITY_SIGNATURE_BYTES, + ) else { return false; }; let Ok(signature) = Signature::from_slice(&signature_bytes) else { @@ -481,9 +515,11 @@ fn agent_identity_credentials(config: &Value) -> Result Result { - let ciphertext = STANDARD - .decode(encrypted_task_id.trim()) - .map_err(|_| "Agent Identity encrypted_task_id must be base64".to_string())?; + let ciphertext = decode_standard_base64_with_limit( + encrypted_task_id.trim(), + MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES, + ) + .ok_or_else(|| "Agent Identity encrypted_task_id must be valid bounded base64".to_string())?; let seed = credentials.signing_key.to_bytes(); let digest = Sha512::digest(seed); let mut curve_private_key = [0u8; 32]; @@ -1240,6 +1278,34 @@ mod tests { )); } + #[test] + fn agent_identity_rejects_oversized_base64_before_decode() { + let encoded_limit = super::maximum_base64_len_for_decoded_limit( + super::MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES, + ); + let assertion = format!("AgentAssertion {}", "A".repeat(encoded_limit + 1)); + assert!(!codex_agent_identity_authorization_matches_transport( + &sample_transport(test_auth_config(Some("task-test"))), + &assertion, + )); + + let mut config = test_auth_config(Some("task-test")); + let private_key_limit = super::maximum_base64_len_for_decoded_limit( + super::MAX_AGENT_IDENTITY_PRIVATE_KEY_DER_BYTES, + ); + config["agent_private_key"] = json!("A".repeat(private_key_limit + 1)); + assert!(agent_identity_credentials(&config).is_err()); + + let credentials = + agent_identity_credentials(&test_auth_config(None)).expect("credentials should parse"); + let encrypted_task_id_limit = super::maximum_base64_len_for_decoded_limit( + super::MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES, + ); + assert!( + decrypt_agent_task_id(&credentials, &"A".repeat(encrypted_task_id_limit + 1),).is_err() + ); + } + #[test] fn task_rotation_context_rejects_metadata_and_credential_replacement() { let initial = sample_transport(test_auth_config(Some("task-old"))); diff --git a/crates/aether-provider/transport/src/antigravity/auth.rs b/crates/aether-provider/transport/src/antigravity/auth.rs index 52925427d..c8043a7b5 100644 --- a/crates/aether-provider/transport/src/antigravity/auth.rs +++ b/crates/aether-provider/transport/src/antigravity/auth.rs @@ -5,8 +5,8 @@ use serde_json::Value; use super::super::snapshot::GatewayProviderTransportSnapshot; pub const ANTIGRAVITY_PROVIDER_TYPE: &str = "antigravity"; -pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str = - "antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)"; +pub const ANTIGRAVITY_CLIENT_VERSION: &str = "4.3.0"; +pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str = "vscode/1.X.X (Antigravity/4.3.0)"; const ANTIGRAVITY_CLIENT_NAME: &str = "antigravity"; const ANTIGRAVITY_GOOG_API_CLIENT: &str = "gl-node/18.18.2 fire/0.8.6 grpc/1.10.x"; @@ -109,7 +109,7 @@ pub fn build_antigravity_static_identity_headers( } pub fn build_antigravity_static_client_headers( - client_version: Option<&str>, + _client_version: Option<&str>, session_id: Option<&str>, ) -> BTreeMap { let mut headers = BTreeMap::from([ @@ -125,14 +125,12 @@ pub fn build_antigravity_static_client_headers( String::from("user-agent"), String::from(ANTIGRAVITY_REQUEST_USER_AGENT), ), + ( + String::from("x-client-version"), + String::from(ANTIGRAVITY_CLIENT_VERSION), + ), ]); - if let Some(client_version) = client_version - .map(str::trim) - .filter(|value| !value.is_empty()) - { - headers.insert(String::from("x-client-version"), client_version.to_string()); - } if let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) { headers.insert(String::from("x-vscode-sessionid"), session_id.to_string()); } @@ -273,7 +271,8 @@ mod tests { use super::{ build_antigravity_static_client_headers, resolve_local_antigravity_request_auth, - AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT, + AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_CLIENT_VERSION, + ANTIGRAVITY_REQUEST_USER_AGENT, }; use crate::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, @@ -383,7 +382,7 @@ mod tests { } #[test] - fn static_client_headers_use_native_antigravity_cli_user_agent() { + fn static_client_headers_pin_the_known_good_antigravity_identity() { let headers = build_antigravity_static_client_headers(Some("1.0.16"), Some("session-abc")); assert_eq!( @@ -396,7 +395,7 @@ mod tests { ); assert_eq!( headers.get("x-client-version").map(String::as_str), - Some("1.0.16") + Some(ANTIGRAVITY_CLIENT_VERSION) ); assert_eq!( headers.get("x-vscode-sessionid").map(String::as_str), diff --git a/crates/aether-provider/transport/src/antigravity/policy.rs b/crates/aether-provider/transport/src/antigravity/policy.rs index 27e20fb10..ea95a819c 100644 --- a/crates/aether-provider/transport/src/antigravity/policy.rs +++ b/crates/aether-provider/transport/src/antigravity/policy.rs @@ -1,6 +1,7 @@ use serde_json::Value; use super::super::snapshot::GatewayProviderTransportSnapshot; +use super::super::transport_proxy_is_locally_supported; use super::auth::{ resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport, AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE, @@ -87,11 +88,11 @@ pub fn classify_local_antigravity_request_support( AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules, ); } - if transport.provider.proxy.is_some() - || transport.endpoint.proxy.is_some() - || transport.key.proxy.is_some() - || transport.key.fingerprint.is_some() - { + // A configured proxy is carried by the execution plan itself, so it only + // disqualifies the local request when it cannot be resolved into a usable + // snapshot. Transport profiles stay unsupported because the v1internal + // payload never carries one. + if !transport_proxy_is_locally_supported(transport) || transport.key.fingerprint.is_some() { return AntigravityRequestSideSupport::Unsupported( AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig, ); @@ -114,3 +115,145 @@ pub fn classify_local_antigravity_request_support( AntigravityRequestSideSupport::Supported(AntigravityRequestSideSpec { auth, request_type }) } + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::super::request::AntigravityEnvelopeRequestType; + use super::{ + classify_local_antigravity_request_support, AntigravityRequestSideSupport, + AntigravityRequestSideUnsupportedReason, + }; + use crate::snapshot::{ + GatewayProviderTransportEndpoint, GatewayProviderTransportKey, + GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, + }; + + fn sample_transport() -> GatewayProviderTransportSnapshot { + GatewayProviderTransportSnapshot { + provider: GatewayProviderTransportProvider { + id: "provider-1".to_string(), + name: "Antigravity".to_string(), + provider_type: "antigravity".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-1".to_string(), + provider_id: "provider-1".to_string(), + api_format: "gemini:generate_content".to_string(), + api_family: Some("gemini".to_string()), + endpoint_kind: Some("generate_content".to_string()), + is_active: true, + base_url: "https://daily-cloudcode-pa.googleapis.com".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-1".to_string(), + provider_id: "provider-1".to_string(), + name: "key".to_string(), + auth_type: "oauth".to_string(), + is_active: true, + api_formats: Some(vec!["gemini:generate_content".to_string()]), + 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: "__placeholder__".to_string(), + decrypted_auth_config: Some( + r#"{"provider_type":"antigravity","refresh_token":"rt","cloudaicompanionProject":"project-1"}"# + .to_string(), + ), + }, + } + } + + fn classify(transport: &GatewayProviderTransportSnapshot) -> AntigravityRequestSideSupport { + classify_local_antigravity_request_support( + transport, + &json!({"contents": []}), + AntigravityEnvelopeRequestType::Agent, + ) + } + + fn assert_unsupported_network_config(support: AntigravityRequestSideSupport) { + assert_eq!( + support, + AntigravityRequestSideSupport::Unsupported( + AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig, + ) + ); + } + + #[test] + fn a_resolvable_tunnel_node_proxy_keeps_the_envelope_supported() { + let mut transport = sample_transport(); + transport.provider.proxy = Some(json!({ + "enabled": true, + "node_id": "702d158b-a432-4694-94cc-3bec13dbbc20", + })); + + assert!(matches!( + classify(&transport), + AntigravityRequestSideSupport::Supported(_) + )); + } + + #[test] + fn a_resolvable_url_proxy_keeps_the_envelope_supported() { + for proxy_owner in ["provider", "endpoint", "key"] { + let mut transport = sample_transport(); + let proxy = Some(json!({"enabled": true, "url": "http://127.0.0.1:17000"})); + match proxy_owner { + "provider" => transport.provider.proxy = proxy, + "endpoint" => transport.endpoint.proxy = proxy, + _ => transport.key.proxy = proxy, + } + + assert!( + matches!( + classify(&transport), + AntigravityRequestSideSupport::Supported(_) + ), + "a {proxy_owner} proxy should not disqualify the antigravity envelope" + ); + } + } + + #[test] + fn a_proxy_without_a_route_still_disqualifies_the_envelope() { + let mut transport = sample_transport(); + transport.provider.proxy = Some(json!({"enabled": true})); + + assert_unsupported_network_config(classify(&transport)); + } + + #[test] + fn a_key_fingerprint_still_disqualifies_the_envelope() { + let mut transport = sample_transport(); + transport.key.fingerprint = Some(json!({"transport_profile": "chrome"})); + + assert_unsupported_network_config(classify(&transport)); + } +} diff --git a/crates/aether-provider/transport/src/auth.rs b/crates/aether-provider/transport/src/auth.rs index d7a8b5dc2..3e94a9e1f 100644 --- a/crates/aether-provider/transport/src/auth.rs +++ b/crates/aether-provider/transport/src/auth.rs @@ -1,8 +1,10 @@ use std::collections::BTreeMap; use super::headers::{ - is_aether_internal_header, is_upstream_credential_header, normalize_upstream_accept_encoding, - should_skip_upstream_complete_passthrough_header, should_skip_upstream_passthrough_header, + declared_connection_header_names, is_aether_internal_header, is_upstream_credential_header, + normalize_upstream_accept_encoding, remove_declared_connection_headers, + should_skip_upstream_complete_passthrough_header_with_connection, + should_skip_upstream_passthrough_header_with_connection, }; use super::snapshot::GatewayProviderTransportSnapshot; @@ -13,13 +15,17 @@ fn collect_passthrough_headers( headers: &http::HeaderMap, extra_headers: &BTreeMap, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = BTreeMap::new(); for (name, value) in headers.iter() { let Ok(value) = value.to_str() else { continue; }; let key = name.as_str().to_ascii_lowercase(); - if should_skip_upstream_passthrough_header(&key) { + if should_skip_upstream_passthrough_header_with_connection( + &key, + &declared_connection_headers, + ) { continue; } let Some(value) = normalize_passthrough_header_value(&key, value) else { @@ -30,7 +36,10 @@ fn collect_passthrough_headers( for (key, value) in extra_headers { let normalized_key = key.to_ascii_lowercase(); - if should_skip_upstream_passthrough_header(&normalized_key) { + if should_skip_upstream_passthrough_header_with_connection( + &normalized_key, + &declared_connection_headers, + ) { continue; } let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else { @@ -46,13 +55,17 @@ fn collect_complete_passthrough_headers( headers: &http::HeaderMap, extra_headers: &BTreeMap, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = BTreeMap::new(); for (name, value) in headers.iter() { let Ok(value) = value.to_str() else { continue; }; let key = name.as_str().to_ascii_lowercase(); - if should_skip_upstream_complete_passthrough_header(&key) { + if should_skip_upstream_complete_passthrough_header_with_connection( + &key, + &declared_connection_headers, + ) { continue; } let Some(value) = normalize_passthrough_header_value(&key, value) else { @@ -63,7 +76,10 @@ fn collect_complete_passthrough_headers( for (key, value) in extra_headers { let normalized_key = key.to_ascii_lowercase(); - if should_skip_upstream_complete_passthrough_header(&normalized_key) { + if should_skip_upstream_complete_passthrough_header_with_connection( + &normalized_key, + &declared_connection_headers, + ) { continue; } let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else { @@ -101,6 +117,8 @@ pub fn build_passthrough_headers( .trim() .to_string() }); + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.remove("content-length"); out } @@ -112,8 +130,10 @@ pub fn build_openai_passthrough_headers( extra_headers: &BTreeMap, content_type: Option<&str>, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = build_passthrough_headers(headers, extra_headers, content_type); ensure_upstream_auth_header(&mut out, auth_header, auth_value); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out } @@ -122,7 +142,9 @@ pub fn build_complete_passthrough_headers( extra_headers: &BTreeMap, content_type: Option<&str>, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = collect_complete_passthrough_headers(headers, extra_headers); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.entry("content-type".to_string()).or_insert_with(|| { content_type .filter(|value| !value.trim().is_empty()) @@ -130,6 +152,7 @@ pub fn build_complete_passthrough_headers( .trim() .to_string() }); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.remove("content-length"); out } @@ -141,8 +164,10 @@ pub fn build_complete_passthrough_headers_with_auth( extra_headers: &BTreeMap, content_type: Option<&str>, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type); replace_upstream_auth_headers(&mut out, auth_header, auth_value); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out } @@ -153,6 +178,7 @@ pub fn build_claude_passthrough_headers( extra_headers: &BTreeMap, content_type: Option<&str>, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = build_openai_passthrough_headers( headers, auth_header, @@ -164,7 +190,10 @@ pub fn build_claude_passthrough_headers( for (name, value) in extra_headers { let key = name.to_ascii_lowercase(); let value = value.trim(); - if value.is_empty() || !should_restore_claude_passthrough_header(&key) { + if value.is_empty() + || !should_restore_claude_passthrough_header(&key) + || declared_connection_headers.contains(&key) + { continue; } @@ -185,7 +214,10 @@ pub fn build_claude_passthrough_headers( }; let key = name.as_str().to_ascii_lowercase(); let value = value.trim(); - if value.is_empty() || !should_restore_claude_passthrough_header(&key) { + if value.is_empty() + || !should_restore_claude_passthrough_header(&key) + || declared_connection_headers.contains(&key) + { continue; } @@ -202,6 +234,7 @@ pub fn build_claude_passthrough_headers( out.entry("anthropic-version".to_string()) .or_insert_with(|| DEFAULT_ANTHROPIC_VERSION.to_string()); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out } @@ -211,8 +244,10 @@ pub fn build_passthrough_headers_with_auth( auth_value: &str, extra_headers: &BTreeMap, ) -> BTreeMap { + let declared_connection_headers = declared_connection_header_names(headers, extra_headers); let mut out = collect_passthrough_headers(headers, extra_headers); replace_upstream_auth_headers(&mut out, auth_header, auth_value); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.remove("content-length"); out } @@ -357,9 +392,9 @@ fn bearer_auth_value(secret: &str) -> String { #[cfg(test)] mod tests { use super::{ - build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth, - build_openai_passthrough_headers, resolve_local_openai_bearer_auth, - resolve_local_standard_auth, + build_claude_passthrough_headers, build_complete_passthrough_headers, + build_complete_passthrough_headers_with_auth, build_openai_passthrough_headers, + resolve_local_openai_bearer_auth, resolve_local_standard_auth, }; use crate::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, @@ -495,6 +530,67 @@ mod tests { ); } + #[test] + fn passthrough_headers_strip_connection_declared_fields() { + let mut headers = http::HeaderMap::new(); + headers.append( + http::header::CONNECTION, + http::HeaderValue::from_static("X-Internal-Hop, keep-alive"), + ); + headers.append( + http::header::CONNECTION, + http::HeaderValue::from_static("x-extra-hop"), + ); + headers.insert( + "x-internal-hop", + http::HeaderValue::from_static("private-value"), + ); + headers.insert( + "x-extra-hop", + http::HeaderValue::from_static("private-value-2"), + ); + headers.insert("x-public", http::HeaderValue::from_static("ok")); + + let extra = BTreeMap::from([ + ( + "Connection".to_string(), + "X-Extra-From-Connection".to_string(), + ), + ("X-Extra-From-Connection".to_string(), "secret".to_string()), + ]); + let built = build_openai_passthrough_headers( + &headers, + "authorization", + "Bearer upstream", + &extra, + Some("application/json"), + ); + + assert_eq!(built.get("x-public").map(String::as_str), Some("ok")); + assert!(!built.contains_key("connection")); + assert!(!built.contains_key("x-internal-hop")); + assert!(!built.contains_key("x-extra-hop")); + assert!(built + .keys() + .all(|name| { !name.eq_ignore_ascii_case("x-extra-from-connection") })); + } + + #[test] + fn complete_passthrough_headers_strip_connection_declared_fields_from_extra_headers() { + let headers = http::HeaderMap::new(); + let extra = BTreeMap::from([ + ("Connection".to_string(), "x-private-hop".to_string()), + ("X-Private-Hop".to_string(), "secret".to_string()), + ("x-public".to_string(), "ok".to_string()), + ]); + + let built = build_complete_passthrough_headers(&headers, &extra, None); + + assert_eq!(built.get("x-public").map(String::as_str), Some("ok")); + assert!(!built.contains_key("connection")); + assert!(!built.contains_key("x-private-hop")); + } + #[test] fn claude_passthrough_headers_preserve_explicit_anthropic_version_override() { let mut headers = http::HeaderMap::new(); diff --git a/crates/aether-provider/transport/src/diagnostics.rs b/crates/aether-provider/transport/src/diagnostics.rs index 9dd4c2003..170c3b81f 100644 --- a/crates/aether-provider/transport/src/diagnostics.rs +++ b/crates/aether-provider/transport/src/diagnostics.rs @@ -1,7 +1,7 @@ use aether_ai_formats::formats::matrix::{ request_conversion_kind, request_conversion_requires_enable_flag, }; -use aether_contracts::ProxySnapshot; +use aether_contracts::{ProxySnapshot, ResolvedTransportProfile}; use serde_json::{json, Map, Value}; use crate::conversion::{ @@ -71,25 +71,24 @@ pub fn build_transport_diagnostics( ) -> Value { let resolved_transport_profile_id = resolve_transport_profile_id(transport); let resolved_transport_profile = resolve_transport_profile(transport) - .and_then(|profile| serde_json::to_value(profile).ok()) + .as_ref() + .map(summarize_transport_profile) .unwrap_or(Value::Null); - let configured_key_transport_profile = transport + let key_transport_profile_configured = transport .key .fingerprint .as_ref() .and_then(Value::as_object) .and_then(|value| value.get("transport_profile")) - .cloned() - .unwrap_or(Value::Null); - let configured_provider_transport_profile = transport + .is_some_and(|value| !value.is_null()); + let provider_transport_profile_configured = transport .provider .config .as_ref() .and_then(|value| value.get("fingerprint")) .and_then(Value::as_object) .and_then(|value| value.get("transport_profile")) - .cloned() - .unwrap_or(Value::Null); + .is_some_and(|value| !value.is_null()); let configured_legacy_grok_transport_profile = if transport .provider .provider_type @@ -109,7 +108,8 @@ pub fn build_transport_diagnostics( &auth_config, "grok_auth_config", ) - .and_then(|profile| serde_json::to_value(profile).ok()) + .as_ref() + .map(summarize_transport_profile) }) .unwrap_or(Value::Null) } else { @@ -132,11 +132,14 @@ pub fn build_transport_diagnostics( "key_is_active": transport.key.is_active, "provider_enable_format_conversion": transport.provider.enable_format_conversion, "provider_keep_priority_on_conversion": transport.provider.keep_priority_on_conversion, - "endpoint_format_acceptance_config": transport.endpoint.format_acceptance_config, - "endpoint_custom_path": transport.endpoint.custom_path, - "header_rules": transport.endpoint.header_rules, + "endpoint_format_acceptance": summarize_format_acceptance_config( + transport.endpoint.format_acceptance_config.as_ref() + ), + "endpoint_has_custom_path": transport.endpoint.custom_path.as_deref() + .is_some_and(|value| !value.trim().is_empty()), + "header_rules_count": json_array_len(transport.endpoint.header_rules.as_ref()), "header_rules_supported": header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()), - "body_rules": transport.endpoint.body_rules, + "body_rules_count": json_array_len(transport.endpoint.body_rules.as_ref()), "body_rules_supported": body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()), "proxy": { "locally_supported": transport_proxy_is_locally_supported(transport), @@ -149,9 +152,10 @@ pub fn build_transport_diagnostics( "has_oauth_config": has_oauth_config, "oauth_request_auth_resolution_supported": oauth_resolution_supported, }, - "fingerprint": transport.key.fingerprint, - "configured_key_transport_profile": configured_key_transport_profile, - "configured_provider_transport_profile": configured_provider_transport_profile, + "key_fingerprint_configured": transport.key.fingerprint.as_ref() + .is_some_and(|value| !value.is_null()), + "key_transport_profile_configured": key_transport_profile_configured, + "provider_transport_profile_configured": provider_transport_profile_configured, "configured_legacy_grok_transport_profile": configured_legacy_grok_transport_profile, "resolved_transport_profile_id": resolved_transport_profile_id, "resolved_transport_profile": resolved_transport_profile, @@ -177,6 +181,34 @@ pub fn build_transport_diagnostics( }) } +fn json_array_len(value: Option<&Value>) -> usize { + value.and_then(Value::as_array).map_or(0, Vec::len) +} + +fn summarize_format_acceptance_config(value: Option<&Value>) -> Value { + let Some(object) = value.and_then(Value::as_object) else { + return json!({ "configured": false }); + }; + json!({ + "configured": true, + "enabled": object.get("enabled").and_then(Value::as_bool), + "accept_formats_count": json_array_len(object.get("accept_formats")), + "reject_formats_count": json_array_len(object.get("reject_formats")), + }) +} + +fn summarize_transport_profile(profile: &ResolvedTransportProfile) -> Value { + json!({ + "profile_id": profile.profile_id, + "backend": profile.backend, + "http_mode": profile.http_mode, + "pool_scope": profile.pool_scope, + "has_header_fingerprint": profile.header_fingerprint.as_ref() + .is_some_and(|value| !value.is_null()), + "has_extra": profile.extra.as_ref().is_some_and(|value| !value.is_null()), + }) +} + fn summarize_proxy_config(proxy: Option<&Value>) -> Value { let Some(object) = proxy.and_then(Value::as_object) else { return Value::Null; @@ -187,11 +219,16 @@ fn summarize_proxy_config(proxy: Option<&Value>) -> Value { .and_then(Value::as_str) .is_some_and(|value| !value.trim().is_empty()); json!({ + "configured": true, "enabled": object.get("enabled").cloned().unwrap_or(Value::Null), - "mode": object.get("mode").cloned().unwrap_or(Value::Null), - "node_id": object.get("node_id").cloned().unwrap_or(Value::Null), - "label": object.get("label").cloned().unwrap_or(Value::Null), + "has_mode": object.get("mode").and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()), + "has_node_id": object.get("node_id").and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()), + "has_label": object.get("label").and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()), "has_url": has_url, + "has_extra": object.get("extra").is_some_and(|value| !value.is_null()), }) } @@ -426,11 +463,13 @@ mod tests { build_transport_diagnostics(&sample_transport(), "claude:messages", "openai:responses"); assert_eq!(diagnostics["provider_type"], "codex"); + assert_eq!(diagnostics["key_fingerprint_configured"], true); + assert_eq!(diagnostics["key_transport_profile_configured"], true); + assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136"); assert_eq!( - diagnostics["fingerprint"]["transport_profile"]["profile_id"], + diagnostics["resolved_transport_profile"]["profile_id"], "chrome_136" ); - assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136"); assert_eq!( diagnostics["request_pair"]["conversion_enabled"], Value::Bool(true) @@ -489,6 +528,86 @@ mod tests { ); } + #[test] + fn transport_diagnostics_do_not_serialize_configured_secrets() { + let mut transport = sample_transport(); + transport.provider.proxy = Some(json!({ + "enabled": true, + "mode": "secret-proxy-mode", + "node_id": "secret-proxy-node", + "label": "secret-proxy-label", + "url": "https://secret-user:secret-pass@proxy.example/secret-path", + "extra": {"token": "secret-proxy-extra"} + })); + transport.provider.config = Some(json!({ + "secret": "secret-provider-config", + "fingerprint": { + "transport_profile": { + "profile_id": "safe-profile", + "header_fingerprint": {"authorization": "secret-profile-header"}, + "extra": {"token": "secret-profile-extra"} + } + } + })); + transport.endpoint.custom_path = Some("/secret-custom-path".to_string()); + transport.endpoint.header_rules = Some(json!([ + {"op": "set", "key": "authorization", "value": "secret-header-rule"} + ])); + transport.endpoint.body_rules = Some(json!([ + {"op": "set", "path": "auth.token", "value": "secret-body-rule"} + ])); + transport.endpoint.format_acceptance_config = Some(json!({ + "enabled": true, + "accept_formats": ["secret-accepted-format"], + "reject_formats": ["secret-rejected-format"], + "token": "secret-format-config" + })); + transport.key.fingerprint = Some(json!({ + "secret": "secret-key-fingerprint", + "transport_profile": { + "profile_id": "safe-key-profile", + "header_fingerprint": {"authorization": "secret-key-profile-header"}, + "extra": {"token": "secret-key-profile-extra"} + } + })); + + let diagnostics = + build_transport_diagnostics(&transport, "claude:messages", "openai:responses"); + let serialized = serde_json::to_string(&diagnostics).unwrap(); + + for secret in [ + "secret-proxy-mode", + "secret-proxy-node", + "secret-proxy-label", + "secret-user", + "secret-pass", + "secret-path", + "secret-proxy-extra", + "secret-provider-config", + "secret-profile-header", + "secret-profile-extra", + "secret-custom-path", + "secret-header-rule", + "secret-body-rule", + "secret-accepted-format", + "secret-rejected-format", + "secret-format-config", + "secret-key-fingerprint", + "secret-key-profile-header", + "secret-key-profile-extra", + ] { + assert!(!serialized.contains(secret), "leaked {secret}"); + } + assert_eq!(diagnostics["header_rules_count"], 1); + assert_eq!(diagnostics["body_rules_count"], 1); + assert_eq!(diagnostics["proxy"]["provider"]["has_node_id"], true); + assert_eq!( + diagnostics["resolved_transport_profile"]["has_header_fingerprint"], + true + ); + assert_eq!(diagnostics["resolved_transport_profile"]["has_extra"], true); + } + #[test] fn request_trace_proxy_value_sanitizes_url_and_marks_config_source() { let transport = sample_transport(); diff --git a/crates/aether-provider/transport/src/gemini_cli/auth.rs b/crates/aether-provider/transport/src/gemini_cli/auth.rs index 8e9eea69e..10bbd2955 100644 --- a/crates/aether-provider/transport/src/gemini_cli/auth.rs +++ b/crates/aether-provider/transport/src/gemini_cli/auth.rs @@ -1,15 +1,26 @@ use serde_json::Value; +use std::fmt; use super::super::snapshot::GatewayProviderTransportSnapshot; pub const GEMINI_CLI_PROVIDER_TYPE: &str = "gemini_cli"; -#[derive(Debug, Clone, Default, PartialEq, Eq)] +#[derive(Clone, Default, PartialEq, Eq)] pub struct GeminiCliRequestAuth { pub project_id: Option, pub session_id: Option, } +impl fmt::Debug for GeminiCliRequestAuth { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GeminiCliRequestAuth") + .field("project_id", &self.project_id) + .field("has_session_id", &self.session_id.is_some()) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum GeminiCliRequestAuthSupport { Supported(GeminiCliRequestAuth), diff --git a/crates/aether-provider/transport/src/gemini_cli/request.rs b/crates/aether-provider/transport/src/gemini_cli/request.rs index dc5e11d13..a70c62c72 100644 --- a/crates/aether-provider/transport/src/gemini_cli/request.rs +++ b/crates/aether-provider/transport/src/gemini_cli/request.rs @@ -1,13 +1,31 @@ use serde_json::{Map, Value}; +use std::fmt; use super::auth::GeminiCliRequestAuth; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub enum GeminiCliRequestEnvelopeSupport { Supported(Value), Unsupported(GeminiCliRequestEnvelopeUnsupportedReason), } +impl fmt::Debug for GeminiCliRequestEnvelopeSupport { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Supported(body) => formatter + .debug_struct("Supported") + .field( + "body_bytes", + &serde_json::to_vec(body).ok().map(|bytes| bytes.len()), + ) + .finish(), + Self::Unsupported(reason) => { + formatter.debug_tuple("Unsupported").field(reason).finish() + } + } + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum GeminiCliRequestEnvelopeUnsupportedReason { NonObjectBody, diff --git a/crates/aether-provider/transport/src/gemini_files/mod.rs b/crates/aether-provider/transport/src/gemini_files/mod.rs index 2c42f28de..4f64c7698 100644 --- a/crates/aether-provider/transport/src/gemini_files/mod.rs +++ b/crates/aether-provider/transport/src/gemini_files/mod.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use serde_json::{json, Value}; @@ -17,13 +18,36 @@ pub enum GeminiFilesRequestBodyError { BodyRulesApplyFailed, } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct GeminiFilesRequestBodyParts { pub provider_request_body: Option, pub provider_request_body_base64: Option, } -#[derive(Debug, Clone, Copy)] +impl fmt::Debug for GeminiFilesRequestBodyParts { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GeminiFilesRequestBodyParts") + .field( + "has_provider_request_body", + &self.provider_request_body.is_some(), + ) + .field( + "provider_request_body_bytes", + &self + .provider_request_body + .as_ref() + .and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())), + ) + .field( + "provider_request_body_base64_len", + &self.provider_request_body_base64.as_ref().map(String::len), + ) + .finish() + } +} + +#[derive(Clone, Copy)] pub struct GeminiFilesHeadersInput<'a> { pub headers: &'a http::HeaderMap, pub auth_header: &'a str, @@ -35,6 +59,40 @@ pub struct GeminiFilesHeadersInput<'a> { pub original_body_is_empty: bool, } +impl fmt::Debug for GeminiFilesHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GeminiFilesHeadersInput") + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field("has_header_rules", &self.header_rules.is_some()) + .field( + "has_provider_request_body", + &self.provider_request_body.is_some(), + ) + .field( + "provider_request_body_base64_len", + &self.provider_request_body_base64.map(str::len), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body_json) + .ok() + .map(|bytes| bytes.len()), + ) + .field("original_body_is_empty", &self.original_body_is_empty) + .finish() + } +} + pub fn gemini_files_transport_unsupported_reason( transport: &GatewayProviderTransportSnapshot, api_format: &str, @@ -135,6 +193,12 @@ pub fn build_gemini_files_headers( ) { return None; } + let declared_connection_headers = + crate::headers::declared_connection_header_names(input.headers, &BTreeMap::new()); + crate::headers::remove_declared_connection_headers( + &mut provider_request_headers, + &declared_connection_headers, + ); Some(provider_request_headers) } diff --git a/crates/aether-provider/transport/src/generic_oauth/mod.rs b/crates/aether-provider/transport/src/generic_oauth/mod.rs index 99c7fdcc5..2d2eaab99 100644 --- a/crates/aether-provider/transport/src/generic_oauth/mod.rs +++ b/crates/aether-provider/transport/src/generic_oauth/mod.rs @@ -62,9 +62,26 @@ pub fn resolve_local_generic_oauth_transport_authorization( .map(|token| format!("Bearer {token}")) } -#[derive(Debug, Clone, Default)] +#[derive(Clone, Default)] pub struct GenericOAuthRefreshAdapter { token_url_overrides: BTreeMap, + oauth_credentials_overrides: BTreeMap, +} + +impl std::fmt::Debug for GenericOAuthRefreshAdapter { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GenericOAuthRefreshAdapter") + .field( + "token_url_override_provider_types", + &self.token_url_overrides.keys().collect::>(), + ) + .field( + "oauth_credentials_override_provider_types", + &self.oauth_credentials_overrides.keys().collect::>(), + ) + .finish() + } } impl GenericOAuthRefreshAdapter { @@ -78,11 +95,29 @@ impl GenericOAuthRefreshAdapter { self } + pub fn with_oauth_credentials_for_tests( + mut self, + provider_type: &str, + client_id: impl Into, + client_secret: impl Into, + ) -> Self { + self.oauth_credentials_overrides.insert( + provider_type.trim().to_ascii_lowercase(), + (client_id.into(), client_secret.into()), + ); + self + } + fn adapter_for_provider_type( &self, provider_type: &'static str, ) -> Option { - let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?; + let mut adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?; + if let Some((client_id, client_secret)) = + self.oauth_credentials_overrides.get(provider_type) + { + adapter = adapter.with_oauth_credentials_for_tests(client_id, client_secret); + } if let Some(token_url) = self.token_url_overrides.get(provider_type) { return Some(adapter.with_token_url_override(token_url.clone())); } @@ -912,7 +947,12 @@ mod tests { hits: Arc::clone(&hits), }; let adapter = GenericOAuthRefreshAdapter::default() - .with_token_url_for_tests("antigravity", "https://oauth.example/token"); + .with_token_url_for_tests("antigravity", "https://oauth.example/token") + .with_oauth_credentials_for_tests( + "antigravity", + "test-client-id", + "test-client-secret", + ); assert!(adapter.supports(&transport)); assert!(adapter.should_refresh(&transport, None)); diff --git a/crates/aether-provider/transport/src/grok.rs b/crates/aether-provider/transport/src/grok.rs index af71e0ea2..e89b51a59 100644 --- a/crates/aether-provider/transport/src/grok.rs +++ b/crates/aether-provider/transport/src/grok.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use aether_contracts::{ ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, @@ -33,7 +34,7 @@ pub struct GrokBrowserProfileMetadata { pub sec_ch_ua_platform: String, } -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct GrokHeaderInput<'a> { pub transport: &'a GatewayProviderTransportSnapshot, pub transport_profile: Option<&'a ResolvedTransportProfile>, @@ -45,6 +46,37 @@ pub struct GrokHeaderInput<'a> { pub original_request_body: &'a Value, } +impl fmt::Debug for GrokHeaderInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GrokHeaderInput") + .field("transport", &self.transport) + .field("transport_profile", &self.transport_profile) + .field( + "request_header_names", + &self + .request_headers + .map(|headers| headers.keys().map(|name| name.as_str()).collect::>()), + ) + .field("content_type", &self.content_type) + .field("accept", &self.accept) + .field("has_header_rules", &self.header_rules.is_some()) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .finish() + } +} + pub fn is_grok_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool { transport .provider @@ -254,6 +286,14 @@ pub fn build_grok_browser_headers(input: GrokHeaderInput<'_>) -> Option, +) -> BTreeSet { + let mut values = headers + .get_all(http::header::CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect::>(); + values.extend( + extra_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str())) + .map(|(_, value)| value.as_str()), + ); + aether_http::connection_declared_header_names(values) +} + +pub(crate) fn is_declared_connection_header( + name: &str, + declared_connection_headers: &BTreeSet, +) -> bool { + declared_connection_headers.contains(&name.trim().to_ascii_lowercase()) +} + const UPSTREAM_CREDENTIAL_HEADER_NAMES: &[&str] = &[ "authorization", "proxy-authorization", @@ -30,7 +58,10 @@ pub(crate) fn is_aether_internal_header(name: &str) -> bool { pub fn should_skip_request_header(name: &str) -> bool { let normalized = name.to_ascii_lowercase(); - if is_aether_internal_header(&normalized) { + if is_aether_internal_header(&normalized) + || is_untrusted_forwarding_metadata_header(&normalized) + || is_untrusted_routing_override_header(&normalized) + { return true; } matches!( @@ -49,6 +80,27 @@ pub fn should_skip_request_header(name: &str) -> bool { ) } +fn is_untrusted_forwarding_metadata_header(normalized_name: &str) -> bool { + matches!(normalized_name, "forwarded" | "via") + || normalized_name.starts_with("x-forwarded-") + || normalized_name.starts_with("x_forwarded_") + || normalized_name.starts_with("x-real-") + || normalized_name.starts_with("x_real_") +} + +fn is_untrusted_routing_override_header(normalized_name: &str) -> bool { + matches!( + normalized_name, + "x-http-method" + | "x-http-method-override" + | "x-method-override" + | "x-override-url" + | "x-rewrite-url" + ) || normalized_name.starts_with("x-original-") + || normalized_name.starts_with("x_original_") + || normalized_name.starts_with("x-envoy-original-") +} + pub fn should_skip_upstream_passthrough_header(name: &str) -> bool { let lower = name.to_ascii_lowercase(); // Anthropic SDK (stainless) client metadata and Anthropic-specific headers @@ -82,6 +134,14 @@ pub fn should_skip_upstream_passthrough_header(name: &str) -> bool { || should_skip_request_header(name) } +pub(crate) fn should_skip_upstream_passthrough_header_with_connection( + name: &str, + declared_connection_headers: &BTreeSet, +) -> bool { + should_skip_upstream_passthrough_header(name) + || is_declared_connection_header(name, declared_connection_headers) +} + pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bool { let lower = name.to_ascii_lowercase(); is_upstream_credential_header(&lower) @@ -103,6 +163,33 @@ pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bo || should_skip_request_header(name) } +pub(crate) fn should_skip_upstream_complete_passthrough_header_with_connection( + name: &str, + declared_connection_headers: &BTreeSet, +) -> bool { + should_skip_upstream_complete_passthrough_header(name) + || is_declared_connection_header(name, declared_connection_headers) +} + +pub(crate) fn remove_declared_connection_headers( + headers: &mut BTreeMap, + declared_connection_headers: &BTreeSet, +) { + let mut all_declared = declared_connection_headers.clone(); + let connection_values = headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("connection")) + .map(|(_, value)| value.as_str()) + .collect::>(); + all_declared.extend(aether_http::connection_declared_header_names( + connection_values, + )); + headers.retain(|name, _| { + !name.eq_ignore_ascii_case("connection") + && !is_declared_connection_header(name, &all_declared) + }); +} + pub fn normalize_upstream_accept_encoding(value: &str) -> Option { let mut accepted = Vec::new(); let mut wildcard_allowed = false; @@ -337,6 +424,44 @@ mod tests { } } + #[test] + fn strips_all_client_supplied_forwarding_metadata() { + for header in [ + "Forwarded", + "Via", + "X-Forwarded-For", + "X-Forwarded-Host", + "X-Forwarded-Prefix", + "X-Forwarded-Server", + "x_forwarded_for", + "X-Real-IP", + "X-Real-Host", + "x_real_ip", + ] { + assert!(should_skip_request_header(header)); + assert!(should_skip_upstream_passthrough_header(header)); + assert!(should_skip_upstream_complete_passthrough_header(header)); + } + } + + #[test] + fn strips_client_supplied_upstream_routing_overrides() { + for header in [ + "X-HTTP-Method-Override", + "X-Method-Override", + "X-Original-URL", + "X-Original-URI", + "x_original_url", + "X-Rewrite-URL", + "X-Override-URL", + "X-Envoy-Original-Path", + ] { + assert!(should_skip_request_header(header)); + assert!(should_skip_upstream_passthrough_header(header)); + assert!(should_skip_upstream_complete_passthrough_header(header)); + } + } + #[test] fn strips_usage_server_time_header_from_provider_requests() { for h in [ @@ -365,4 +490,23 @@ mod tests { ); } } + + #[test] + fn declared_connection_names_are_case_insensitive_and_multi_line() { + let mut headers = http::HeaderMap::new(); + headers.append( + http::header::CONNECTION, + http::HeaderValue::from_static("X-Hop, keep-alive"), + ); + headers.append( + http::header::CONNECTION, + http::HeaderValue::from_static("x-other-hop"), + ); + + let names = super::declared_connection_header_names(&headers, &BTreeMap::new()); + assert!(names.contains("x-hop")); + assert!(names.contains("keep-alive")); + assert!(names.contains("x-other-hop")); + assert!(super::is_declared_connection_header("X-HOP", &names)); + } } diff --git a/crates/aether-provider/transport/src/kiro/auth.rs b/crates/aether-provider/transport/src/kiro/auth.rs index 3670d7b4f..13cde2f68 100644 --- a/crates/aether-provider/transport/src/kiro/auth.rs +++ b/crates/aether-provider/transport/src/kiro/auth.rs @@ -1,16 +1,27 @@ use super::super::snapshot::GatewayProviderTransportSnapshot; use super::credentials::{generate_machine_id, KiroAuthConfig}; +use std::fmt; pub const PROVIDER_TYPE: &str = "kiro"; pub const KIRO_AUTH_HEADER: &str = "authorization"; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct KiroBearerAuth { pub name: &'static str, pub value: String, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl fmt::Debug for KiroBearerAuth { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("KiroBearerAuth") + .field("name", &self.name) + .field("value", &"[REDACTED]") + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct KiroRequestAuth { pub name: &'static str, pub value: String, @@ -18,6 +29,18 @@ pub struct KiroRequestAuth { pub machine_id: String, } +impl fmt::Debug for KiroRequestAuth { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("KiroRequestAuth") + .field("name", &self.name) + .field("value", &"[REDACTED]") + .field("auth_config", &self.auth_config) + .field("machine_id", &"[REDACTED]") + .finish() + } +} + pub fn is_kiro_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool { transport .provider @@ -224,6 +247,41 @@ mod tests { assert!(supports_local_kiro_auth_prerequisites(&sample_transport())); } + #[test] + fn kiro_request_auth_debug_output_redacts_credentials_and_machine_identity() { + let bearer = resolve_local_kiro_bearer_auth(&sample_transport()) + .expect("kiro bearer auth should resolve"); + let bearer_debug = format!("{bearer:?}"); + assert!(!bearer_debug.contains("upstream-key")); + assert!(bearer_debug.contains("[REDACTED]")); + + let mut transport = sample_transport(); + transport.key.decrypted_api_key = "__placeholder__".to_string(); + transport.key.decrypted_auth_config = Some( + r#"{ + "access_token":"kiro-request-access-canary", + "expires_at":4102444800, + "refresh_token":"kiro-request-refresh-canary................................................................................................", + "machine_id":"kiro-request-machine-canary", + "profile_arn":"kiro-request-profile-canary", + "api_region":"us-west-2" + }"# + .to_string(), + ); + let request_auth = + resolve_local_kiro_request_auth(&transport).expect("kiro request auth should resolve"); + let request_debug = format!("{request_auth:?}"); + for secret in [ + "kiro-request-access-canary", + "kiro-request-refresh-canary", + "kiro-request-machine-canary", + "kiro-request-profile-canary", + ] { + assert!(!request_debug.contains(secret), "debug leaked {secret}"); + } + assert!(request_debug.contains("[REDACTED]")); + } + #[test] fn rejects_auth_config_subset() { let mut transport = sample_transport(); diff --git a/crates/aether-provider/transport/src/kiro/credentials.rs b/crates/aether-provider/transport/src/kiro/credentials.rs index 44c644870..3d4f672f5 100644 --- a/crates/aether-provider/transport/src/kiro/credentials.rs +++ b/crates/aether-provider/transport/src/kiro/credentials.rs @@ -2,6 +2,7 @@ pub use aether_oauth::provider::providers::{ generate_kiro_machine_id as generate_machine_id, normalize_kiro_machine_id as normalize_machine_id, KiroAuthConfig, DEFAULT_REGION, }; +pub use aether_oauth::provider::providers::{is_valid_kiro_region, normalize_kiro_region}; #[cfg(test)] mod tests { diff --git a/crates/aether-provider/transport/src/kiro/headers.rs b/crates/aether-provider/transport/src/kiro/headers.rs index 435d182d5..b6aa50e80 100644 --- a/crates/aether-provider/transport/src/kiro/headers.rs +++ b/crates/aether-provider/transport/src/kiro/headers.rs @@ -2,7 +2,7 @@ use std::collections::BTreeMap; use uuid::Uuid; -use super::credentials::KiroAuthConfig; +use super::credentials::{normalize_kiro_region, KiroAuthConfig}; pub const AWS_EVENTSTREAM_CONTENT_TYPE: &str = "application/vnd.amazon.eventstream"; pub const KIRO_PROFILE_ARN_HEADER: &str = "x-amzn-kiro-profile-arn"; @@ -47,7 +47,7 @@ pub fn build_generate_assistant_headers( let kiro_version = auth_config.effective_kiro_version(); let system_version = auth_config.effective_system_version(); let node_version = auth_config.effective_node_version(); - let region = auth_config.effective_api_region(); + let region = normalize_kiro_region(auth_config.effective_api_region()); let host = format!("q.{region}.amazonaws.com"); BTreeMap::from([ @@ -92,7 +92,7 @@ pub fn build_mcp_headers( let kiro_version = auth_config.effective_kiro_version(); let system_version = auth_config.effective_system_version(); let node_version = auth_config.effective_node_version(); - let region = auth_config.effective_api_region(); + let region = normalize_kiro_region(auth_config.effective_api_region()); let host = format!("q.{region}.amazonaws.com"); let mut headers = BTreeMap::from([ @@ -134,7 +134,7 @@ pub fn build_list_available_models_headers( let kiro_version = auth_config.effective_kiro_version(); let system_version = auth_config.effective_system_version(); let node_version = auth_config.effective_node_version(); - let region = auth_config.effective_api_region(); + let region = normalize_kiro_region(auth_config.effective_api_region()); let host = format!("q.{region}.amazonaws.com"); let ide_tag = build_kiro_ide_tag(kiro_version, machine_id); @@ -331,4 +331,30 @@ mod tests { ); assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER)); } + + #[test] + fn malicious_api_region_cannot_inject_host_header() { + let auth_config = KiroAuthConfig { + auth_method: Some("social".to_string()), + refresh_token: None, + expires_at: None, + profile_arn: None, + region: None, + auth_region: None, + api_region: Some("attacker.example/".to_string()), + client_id: None, + client_secret: None, + machine_id: None, + kiro_version: None, + system_version: None, + node_version: None, + access_token: None, + }; + + let headers = build_list_available_models_headers(&auth_config, "machine"); + assert_eq!( + headers.get("host").map(String::as_str), + Some("q.us-east-1.amazonaws.com") + ); + } } diff --git a/crates/aether-provider/transport/src/kiro/mod.rs b/crates/aether-provider/transport/src/kiro/mod.rs index cd393dbb6..fff98500f 100644 --- a/crates/aether-provider/transport/src/kiro/mod.rs +++ b/crates/aether-provider/transport/src/kiro/mod.rs @@ -30,7 +30,10 @@ pub use auth::{ KiroBearerAuth, KiroRequestAuth, KIRO_AUTH_HEADER, PROVIDER_TYPE, }; pub use converter::convert_claude_messages_to_conversation_state; -pub use credentials::{generate_machine_id, normalize_machine_id, KiroAuthConfig}; +pub use credentials::{ + generate_machine_id, is_valid_kiro_region, normalize_kiro_region, normalize_machine_id, + KiroAuthConfig, +}; pub use headers::{ build_generate_assistant_headers, build_list_available_models_headers, build_mcp_headers, AWS_EVENTSTREAM_CONTENT_TYPE, KIRO_EXTERNAL_IDP_TOKEN_TYPE, KIRO_PROFILE_ARN_HEADER, diff --git a/crates/aether-provider/transport/src/kiro/request.rs b/crates/aether-provider/transport/src/kiro/request.rs index d58b03510..9593bc18d 100644 --- a/crates/aether-provider/transport/src/kiro/request.rs +++ b/crates/aether-provider/transport/src/kiro/request.rs @@ -1,7 +1,9 @@ use std::collections::BTreeMap; +use std::fmt; use serde_json::{json, Value}; +use super::super::headers::{declared_connection_header_names, remove_declared_connection_headers}; pub use super::super::rules::{ apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers, body_rules_are_locally_supported, header_rules_are_locally_supported, @@ -83,7 +85,7 @@ pub fn build_kiro_provider_request_body( Some(provider_request_body) } -#[derive(Clone, Copy, Debug)] +#[derive(Clone, Copy)] pub struct KiroProviderHeadersInput<'a> { pub headers: &'a http::HeaderMap, pub provider_request_body: &'a Value, @@ -95,6 +97,39 @@ pub struct KiroProviderHeadersInput<'a> { pub machine_id: &'a str, } +impl fmt::Debug for KiroProviderHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("KiroProviderHeadersInput") + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field("has_header_rules", &self.header_rules.is_some()) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field("auth_config", &self.auth_config) + .field("has_machine_id", &(!self.machine_id.is_empty())) + .finish() + } +} + pub fn build_kiro_provider_headers( input: KiroProviderHeadersInput<'_>, ) -> Option> { @@ -109,13 +144,16 @@ pub fn build_kiro_provider_headers( machine_id, } = input; + let declared_connection_headers = declared_connection_header_names(headers, &BTreeMap::new()); let mut out = BTreeMap::new(); for (name, value) in headers { let Ok(value) = value.to_str() else { continue; }; let key = name.as_str().to_ascii_lowercase(); - if should_skip_upstream_passthrough_header(&key) { + if should_skip_upstream_passthrough_header(&key) + || declared_connection_headers.contains(&key) + { continue; } let value = value.trim(); @@ -143,8 +181,10 @@ pub fn build_kiro_provider_headers( auth_header.trim().to_ascii_lowercase(), auth_value.trim().to_string(), ); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.entry("content-type".to_string()) .or_insert_with(|| "application/json".to_string()); + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.remove("content-length"); Some(out) } diff --git a/crates/aether-provider/transport/src/kiro/url.rs b/crates/aether-provider/transport/src/kiro/url.rs index b69af1ae8..b857ef15d 100644 --- a/crates/aether-provider/transport/src/kiro/url.rs +++ b/crates/aether-provider/transport/src/kiro/url.rs @@ -1,5 +1,5 @@ use super::super::url::build_passthrough_path_url; -use super::credentials::DEFAULT_REGION; +use super::credentials::{normalize_kiro_region, DEFAULT_REGION}; pub const GENERATE_ASSISTANT_RESPONSE_PATH: &str = "/generateAssistantResponse"; pub const LIST_AVAILABLE_MODELS_PATH: &str = "/ListAvailableModels"; @@ -9,8 +9,7 @@ pub const KIRO_ENVELOPE_NAME: &str = "kiro:generateAssistantResponse"; pub fn resolve_kiro_base_url(upstream_base_url: &str, api_region: Option<&str>) -> String { let region = api_region - .map(str::trim) - .filter(|value| !value.is_empty()) + .map(normalize_kiro_region) .unwrap_or(DEFAULT_REGION); upstream_base_url .trim() @@ -105,6 +104,29 @@ mod tests { ); } + #[test] + fn malicious_region_cannot_escape_configured_origin() { + for region in [ + "attacker.example/", + "attacker.example?x=1", + "attacker.example#fragment", + "attacker@example", + ] { + assert_eq!( + resolve_kiro_base_url("https://q.{region}.amazonaws.com", Some(region)), + "https://q.us-east-1.amazonaws.com" + ); + } + assert_eq!( + build_kiro_list_available_models_url( + "https://q.{region}.amazonaws.com", + Some("attacker.example/"), + ) + .as_deref(), + Some("https://q.us-east-1.amazonaws.com/ListAvailableModels?origin=AI_EDITOR") + ); + } + #[test] fn builds_mcp_url_for_latest_kiro_endpoint() { assert_eq!( diff --git a/crates/aether-provider/transport/src/network.rs b/crates/aether-provider/transport/src/network.rs index 70224c4cf..ece64e4ab 100644 --- a/crates/aether-provider/transport/src/network.rs +++ b/crates/aether-provider/transport/src/network.rs @@ -320,13 +320,28 @@ fn proxy_snapshot_from_value(value: &Value) -> Option { let mode = json_string_field(object, "mode"); let node_id = json_string_field(object, "node_id"); let label = json_string_field(object, "label"); - let url = json_string_field(object, "url").or_else(|| json_string_field(object, "proxy_url")); + let url = json_string_field(object, "url") + .or_else(|| json_string_field(object, "proxy_url")) + .and_then(|proxy_url| { + proxy_url_with_auth( + &proxy_url, + json_proxy_credential_field(object, "username"), + json_proxy_credential_field(object, "password"), + ) + }); let mut extra = Map::new(); for (key, value) in object { if matches!( key.as_str(), - "enabled" | "mode" | "node_id" | "label" | "url" | "proxy_url" + "enabled" + | "mode" + | "node_id" + | "label" + | "url" + | "proxy_url" + | "username" + | "password" ) { continue; } @@ -356,6 +371,35 @@ fn json_string_field(object: &Map, key: &str) -> Option { .map(ToOwned::to_owned) } +fn json_proxy_credential_field<'a>(object: &'a Map, key: &str) -> Option<&'a str> { + object + .get(key) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) +} + +fn proxy_url_with_auth( + proxy_url: &str, + username: Option<&str>, + password: Option<&str>, +) -> Option { + let username = username.filter(|value| !value.is_empty()); + let password = password.filter(|value| !value.is_empty()); + let mut parsed = url::Url::parse(proxy_url).ok()?; + if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h") + || parsed.host_str().is_none() + { + return None; + } + if username.is_none() && password.is_none() { + return Some(parsed.to_string()); + } + let username = username.unwrap_or(""); + parsed.set_username(username).ok()?; + parsed.set_password(password).ok()?; + Some(parsed.to_string()) +} + #[cfg(test)] mod tests { use std::collections::BTreeMap; @@ -520,6 +564,43 @@ mod tests { assert_eq!(snapshot.extra, Some(json!({"kind":"manual"}))); } + #[test] + fn resolves_authenticated_inline_proxy_without_secret_extra_fields() { + let mut transport = sample_transport(); + transport.key.proxy = Some(json!({ + "url": "socks5h://proxy.example:1080", + "username": " alice ", + "password": " p:ss ", + "kind": "manual", + })); + + let snapshot = resolve_transport_proxy_snapshot(&transport) + .expect("authenticated proxy snapshot should resolve"); + + assert_eq!( + snapshot.url.as_deref(), + Some("socks5h://%20alice%20:%20p%3Ass%20@proxy.example:1080") + ); + assert_eq!(snapshot.extra, Some(json!({"kind":"manual"}))); + } + + #[test] + fn resolves_legacy_password_only_inline_proxy() { + let mut transport = sample_transport(); + transport.key.proxy = Some(json!({ + "url": "http://proxy.example:8080", + "password": "legacy-password", + })); + + let snapshot = resolve_transport_proxy_snapshot(&transport) + .expect("password-only proxy snapshot should resolve"); + assert_eq!( + snapshot.url.as_deref(), + Some("http://:legacy-password@proxy.example:8080/") + ); + assert!(snapshot.extra.is_none()); + } + #[tokio::test] async fn enriches_transport_proxy_snapshot_with_tunnel_owner_hint() { let state = sample_lookup(); diff --git a/crates/aether-provider/transport/src/oauth_refresh/mod.rs b/crates/aether-provider/transport/src/oauth_refresh/mod.rs index 07c9c5b45..5cdadb2df 100644 --- a/crates/aether-provider/transport/src/oauth_refresh/mod.rs +++ b/crates/aether-provider/transport/src/oauth_refresh/mod.rs @@ -25,7 +25,9 @@ use super::vertex::{ supports_local_vertex_service_account_auth_resolution, VertexServiceAccountRefreshAdapter, }; -#[derive(Debug, Clone, PartialEq, Eq)] +const LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024; + +#[derive(Clone, PartialEq, Eq)] #[allow(clippy::large_enum_variant)] pub enum LocalResolvedOAuthRequestAuth { #[allow(dead_code)] @@ -36,7 +38,20 @@ pub enum LocalResolvedOAuthRequestAuth { Kiro(KiroRequestAuth), } -#[derive(Debug, Clone, PartialEq)] +impl fmt::Debug for LocalResolvedOAuthRequestAuth { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Header { name, .. } => formatter + .debug_struct("Header") + .field("name", name) + .field("value", &"[REDACTED]") + .finish(), + Self::Kiro(auth) => formatter.debug_tuple("Kiro").field(auth).finish(), + } + } +} + +#[derive(Clone, PartialEq)] pub struct LocalOAuthResolution { pub auth: Option, pub refreshed_entry: Option, @@ -53,6 +68,23 @@ pub struct LocalOAuthResolution { pub local_refresh_guard: Option, } +impl fmt::Debug for LocalOAuthResolution { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("LocalOAuthResolution") + .field("auth", &self.auth) + .field("refreshed_entry", &self.refreshed_entry) + .field("refresh_in_flight", &self.refresh_in_flight) + .field("reused_refresh", &self.reused_refresh) + .field("has_distributed_lease", &self.distributed_lease.is_some()) + .field( + "has_local_refresh_guard", + &self.local_refresh_guard.is_some(), + ) + .finish() + } +} + #[derive(Clone)] pub struct LocalOAuthRefreshCommitGuard { guard: Arc>, @@ -78,7 +110,7 @@ impl PartialEq for LocalOAuthRefreshCommitGuard { } } -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct CachedOAuthEntry { pub provider_type: String, pub auth_header_name: String, @@ -89,7 +121,21 @@ pub struct CachedOAuthEntry { pub source_fingerprint: Option, } -#[derive(Debug, Clone, PartialEq)] +impl fmt::Debug for CachedOAuthEntry { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CachedOAuthEntry") + .field("provider_type", &self.provider_type) + .field("auth_header_name", &self.auth_header_name) + .field("auth_header_value", &"[REDACTED]") + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .field("has_metadata", &self.metadata.is_some()) + .field("has_source_fingerprint", &self.source_fingerprint.is_some()) + .finish() + } +} + +#[derive(Clone, PartialEq)] pub struct LocalOAuthHttpRequest { pub request_id: &'static str, pub method: reqwest::Method, @@ -99,21 +145,44 @@ pub struct LocalOAuthHttpRequest { pub body_bytes: Option>, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl fmt::Debug for LocalOAuthHttpRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("LocalOAuthHttpRequest") + .field("request_id", &self.request_id) + .field("method", &self.method) + .field("url", &"[REDACTED]") + .field("header_names", &self.headers.keys().collect::>()) + .field("json_body", &self.json_body.as_ref().map(|_| "[REDACTED]")) + .field("body_bytes_len", &self.body_bytes.as_ref().map(Vec::len)) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct LocalOAuthHttpResponse { pub status_code: u16, pub body_text: String, } -#[derive(Debug, Error)] +impl fmt::Debug for LocalOAuthHttpResponse { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("LocalOAuthHttpResponse") + .field("status_code", &self.status_code) + .field("body_bytes_len", &self.body_text.len()) + .finish() + } +} + +#[derive(Error)] pub enum LocalOAuthRefreshError { - #[error("{provider_type} oauth refresh request failed: {source}")] + #[error("{provider_type} oauth refresh request failed")] Transport { provider_type: &'static str, - #[source] - source: reqwest::Error, + error: reqwest::Error, }, - #[error("{provider_type} oauth refresh returned HTTP {status_code}: {body_excerpt}")] + #[error("{provider_type} oauth refresh returned HTTP {status_code}")] HttpStatus { provider_type: &'static str, status_code: u16, @@ -152,6 +221,64 @@ impl ReqwestLocalOAuthHttpExecutor { } } +impl fmt::Debug for LocalOAuthRefreshError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Transport { provider_type, .. } => formatter + .debug_struct("Transport") + .field("provider_type", provider_type) + .field("error", &"[REDACTED]") + .finish(), + Self::HttpStatus { + provider_type, + status_code, + .. + } => formatter + .debug_struct("HttpStatus") + .field("provider_type", provider_type) + .field("status_code", status_code) + .field("body_excerpt", &"[REDACTED]") + .finish(), + Self::TransportMessage { provider_type, .. } => formatter + .debug_struct("TransportMessage") + .field("provider_type", provider_type) + .field("message", &"[REDACTED]") + .finish(), + Self::InvalidResponse { provider_type, .. } => formatter + .debug_struct("InvalidResponse") + .field("provider_type", provider_type) + .field("message", &"[REDACTED]") + .finish(), + } + } +} + +fn validate_local_oauth_request_url( + provider_type: &'static str, + raw_url: &str, +) -> Result { + let invalid = |message: &'static str| LocalOAuthRefreshError::TransportMessage { + provider_type, + message: message.to_string(), + }; + let url = url::Url::parse(raw_url).map_err(|_| invalid("invalid OAuth endpoint URL"))?; + if url.host().is_none() || !matches!(url.scheme(), "http" | "https") { + return Err(invalid( + "OAuth endpoint URL must use HTTP or HTTPS and include a host", + )); + } + if !url.username().is_empty() || url.password().is_some() { + return Err(invalid("OAuth endpoint URL must not include credentials")); + } + if url.fragment().is_some() { + return Err(invalid("OAuth endpoint URL must not include a fragment")); + } + if !aether_http::is_https_or_loopback_http_url(&url) { + return Err(invalid("remote OAuth endpoint URL must use HTTPS")); + } + Ok(url) +} + #[async_trait] impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor { async fn execute( @@ -160,9 +287,8 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor { _transport: &GatewayProviderTransportSnapshot, request: &LocalOAuthHttpRequest, ) -> Result { - let mut builder = self - .client - .request(request.method.clone(), request.url.as_str()); + let request_url = validate_local_oauth_request_url(provider_type, request.url.as_str())?; + let mut builder = self.client.request(request.method.clone(), request_url); for (name, value) in &request.headers { builder = builder.header(name, value); } @@ -172,23 +298,37 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor { builder = builder.body(body_bytes.clone()); } - let response = + let mut response = builder .send() .await - .map_err(|source| LocalOAuthRefreshError::Transport { + .map_err(|error| LocalOAuthRefreshError::Transport { provider_type, - source, + error, })?; let status_code = response.status().as_u16(); - let body_text = + if response + .content_length() + .is_some_and(|length| length > LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES as u64) + { + return Err(local_oauth_response_too_large(provider_type)); + } + let mut body = Vec::new(); + while let Some(chunk) = response - .text() + .chunk() .await - .map_err(|source| LocalOAuthRefreshError::Transport { + .map_err(|error| LocalOAuthRefreshError::Transport { provider_type, - source, - })?; + error, + })? + { + if chunk.len() > LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) { + return Err(local_oauth_response_too_large(provider_type)); + } + body.extend_from_slice(&chunk); + } + let body_text = String::from_utf8_lossy(&body).to_string(); Ok(LocalOAuthHttpResponse { status_code, body_text, @@ -196,6 +336,13 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor { } } +fn local_oauth_response_too_large(provider_type: &'static str) -> LocalOAuthRefreshError { + LocalOAuthRefreshError::InvalidResponse { + provider_type, + message: format!("response body exceeds {LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES} bytes"), + } +} + pub(crate) struct ProviderOAuthLocalHttpExecutor<'a> { provider_type: &'static str, transport: &'a GatewayProviderTransportSnapshot, @@ -299,8 +446,8 @@ pub(crate) fn oauth_error_to_local_refresh_error( fn local_refresh_error_to_oauth_error(error: LocalOAuthRefreshError) -> OAuthError { match error { - LocalOAuthRefreshError::Transport { source, .. } => { - OAuthError::Transport(source.to_string()) + LocalOAuthRefreshError::Transport { .. } => { + OAuthError::Transport("oauth refresh request failed".to_string()) } LocalOAuthRefreshError::TransportMessage { message, .. } => OAuthError::Transport(message), LocalOAuthRefreshError::HttpStatus { @@ -965,13 +1112,128 @@ mod tests { GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, }; use super::{ - CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, - LocalOAuthRefreshCoordinator, LocalOAuthRefreshError, LocalOAuthResolution, - LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor, + validate_local_oauth_request_url, CachedOAuthEntry, LocalOAuthHttpExecutor, + LocalOAuthRefreshAdapter, LocalOAuthRefreshCoordinator, LocalOAuthRefreshError, + LocalOAuthResolution, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor, }; use async_trait::async_trait; use std::sync::Arc; + #[test] + fn local_oauth_transport_requires_https_or_literal_loopback_http() { + for allowed in [ + "https://oauth.example.test/token?tenant=one", + "http://localhost:8080/token", + "http://127.42.0.1:8080/token", + "http://[::1]:8080/token", + ] { + assert!( + validate_local_oauth_request_url("test", allowed).is_ok(), + "URL should be accepted: {allowed}" + ); + } + for rejected in [ + "http://oauth.example.test/token", + "http://10.0.0.1/token", + "http://0.0.0.0:8080/token", + "http://[::ffff:127.0.0.1]:8080/token", + "https://token@oauth.example.test/token", + "https://oauth.example.test/token#secret", + "file:///tmp/token", + ] { + assert!( + validate_local_oauth_request_url("test", rejected).is_err(), + "URL should be rejected: {rejected}" + ); + } + } + + #[test] + fn local_oauth_url_error_does_not_echo_embedded_credentials() { + let error = validate_local_oauth_request_url( + "test", + "https://sensitive-user:sensitive-password@oauth.example.test/token", + ) + .expect_err("userinfo should be rejected") + .to_string(); + + assert!(!error.contains("sensitive-user")); + assert!(!error.contains("sensitive-password")); + } + + #[test] + fn local_oauth_debug_output_redacts_request_response_and_cached_credentials() { + let request = super::LocalOAuthHttpRequest { + request_id: "provider-oauth:test", + method: reqwest::Method::POST, + url: "https://oauth.example.test/token?client_secret=url-secret-canary".to_string(), + headers: std::collections::BTreeMap::from([ + ( + "authorization".to_string(), + "Bearer request-header-canary".to_string(), + ), + ("content-type".to_string(), "application/json".to_string()), + ]), + json_body: Some(serde_json::json!({ + "refresh_token": "request-body-canary" + })), + body_bytes: None, + }; + let response = super::LocalOAuthHttpResponse { + status_code: 401, + body_text: "{\"access_token\":\"response-body-canary\"}".to_string(), + }; + let entry = super::CachedOAuthEntry { + provider_type: "test".to_string(), + auth_header_name: "authorization".to_string(), + auth_header_value: "Bearer cached-header-canary".to_string(), + expires_at_unix_secs: Some(1), + metadata: Some(serde_json::json!({ + "refresh_token": "cached-metadata-canary" + })), + source_fingerprint: Some("non-secret-fingerprint".to_string()), + }; + let resolution = super::LocalOAuthResolution { + auth: Some(super::LocalResolvedOAuthRequestAuth::Header { + name: "authorization".to_string(), + value: "Bearer resolution-header-canary".to_string(), + }), + refreshed_entry: Some(entry.clone()), + refresh_in_flight: false, + reused_refresh: false, + distributed_lease: None, + local_refresh_guard: None, + }; + + let debug = format!("{request:?} {response:?} {entry:?} {resolution:?}"); + for secret in [ + "url-secret-canary", + "request-header-canary", + "request-body-canary", + "response-body-canary", + "cached-header-canary", + "cached-metadata-canary", + "resolution-header-canary", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + } + assert!(debug.contains("[REDACTED]")); + } + + #[test] + fn local_oauth_refresh_error_debug_and_display_hide_body_and_transport_details() { + let http_error = LocalOAuthRefreshError::HttpStatus { + provider_type: "test", + status_code: 401, + body_excerpt: "{\"access_token\":\"refresh-error-canary\"}".to_string(), + }; + let debug = format!("{http_error:?}"); + let display = http_error.to_string(); + assert!(!debug.contains("refresh-error-canary")); + assert!(!display.contains("refresh-error-canary")); + assert_eq!(display, "test oauth refresh returned HTTP 401"); + } + #[derive(Debug)] struct TestAdapter { refresh_hits: Arc, diff --git a/crates/aether-provider/transport/src/openai_image/mod.rs b/crates/aether-provider/transport/src/openai_image/mod.rs index 6f7030760..2170611c6 100644 --- a/crates/aether-provider/transport/src/openai_image/mod.rs +++ b/crates/aether-provider/transport/src/openai_image/mod.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use serde_json::Value; @@ -9,7 +10,7 @@ use crate::rules::apply_local_header_rules_with_request_headers; use crate::snapshot::GatewayProviderTransportSnapshot; use crate::url::build_openai_image_url; -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct ProviderOpenAiImageHeadersInput<'a> { pub transport: &'a GatewayProviderTransportSnapshot, pub headers: &'a http::HeaderMap, @@ -21,6 +22,39 @@ pub struct ProviderOpenAiImageHeadersInput<'a> { pub original_request_body: &'a Value, } +impl fmt::Debug for ProviderOpenAiImageHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ProviderOpenAiImageHeadersInput") + .field("transport", &self.transport) + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field("accept", &self.accept) + .field("has_header_rules", &self.header_rules.is_some()) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .finish() + } +} + pub fn openai_image_transport_unsupported_reason( transport: &GatewayProviderTransportSnapshot, api_format: &str, @@ -98,6 +132,12 @@ pub fn build_openai_image_headers( ) { return None; } + let declared_connection_headers = + crate::headers::declared_connection_header_names(input.headers, &BTreeMap::new()); + crate::headers::remove_declared_connection_headers( + &mut provider_request_headers, + &declared_connection_headers, + ); Some(provider_request_headers) } diff --git a/crates/aether-provider/transport/src/provider_types.rs b/crates/aether-provider/transport/src/provider_types.rs index b18f43ecf..4e9045c8d 100644 --- a/crates/aether-provider/transport/src/provider_types.rs +++ b/crates/aether-provider/transport/src/provider_types.rs @@ -5,7 +5,6 @@ pub struct ProviderOAuthTemplate { pub authorize_url: &'static str, pub token_url: &'static str, pub client_id: &'static str, - pub client_secret: &'static str, pub scopes: &'static [&'static str], pub redirect_uri: &'static str, pub use_pkce: bool, @@ -550,7 +549,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option Option Option Option Option Option { pub provider_api_format: &'a str, pub mapped_model: Option<&'a str>, @@ -38,6 +39,21 @@ pub struct TransportRequestUrlParams<'a> { pub api_operation: Option, } +impl fmt::Debug for TransportRequestUrlParams<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("TransportRequestUrlParams") + .field("provider_api_format", &self.provider_api_format) + .field("mapped_model", &self.mapped_model) + .field("upstream_is_stream", &self.upstream_is_stream) + .field("has_request_query", &self.request_query.is_some()) + .field("request_query_len", &self.request_query.map(str::len)) + .field("kiro_api_region", &self.kiro_api_region) + .field("api_operation", &self.api_operation) + .finish() + } +} + pub fn build_transport_request_url( transport: &GatewayProviderTransportSnapshot, params: TransportRequestUrlParams<'_>, @@ -112,9 +128,13 @@ fn build_transport_request_url_inner( .filter(|value| !value.is_empty()); let custom_path_handles_operation = custom_path_template.is_some_and(|path| path.contains("{operation}")); - let custom_path = custom_path_template.map(|path| { - expand_custom_path_template(path, build_path_params(params, gemini_embedding_batch)) - }); + let custom_path = match custom_path_template { + Some(path) => Some(expand_custom_path_template( + path, + build_path_params(params, gemini_embedding_batch), + )?), + None => None, + }; if let Some(path) = custom_path.as_deref() { let custom_path_is_complete_claude_count_tokens = normalized_provider_api_format @@ -670,22 +690,24 @@ fn build_gemini_embedding_url( } else { "embedContent" }; + let encoded_model = encode_url_path_segment(trimmed_model); let path = if trimmed_base_url.ends_with("/v1beta") { - format!("/models/{trimmed_model}:{action}") + format!("/models/{encoded_model}:{action}") } else if trimmed_base_url.contains("/v1beta/models/") { format!(":{action}") } else { - format!("/v1beta/models/{trimmed_model}:{action}") + format!("/v1beta/models/{encoded_model}:{action}") }; build_passthrough_path_url(upstream_base_url, &path, query, &["key"]) } -fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> String { +fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> Option { if params.is_empty() { - return path.to_string(); + return Some(path.to_string()); } let regex = custom_path_template_regex(); + let query_start = path.find('?'); let mut missing_key = false; let replaced = regex.replace_all(path, |captures: ®ex::Captures<'_>| { let key = captures @@ -693,7 +715,14 @@ fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) .map(|value| value.as_str()) .unwrap_or_default(); match params.get(key).copied() { - Some(value) => value.to_string(), + Some(value) => { + let capture_start = captures.get(0).map_or(0, |value| value.start()); + if query_start.is_some_and(|query_start| capture_start > query_start) { + url::form_urlencoded::byte_serialize(value.as_bytes()).collect() + } else { + encode_url_path_segment(value) + } + } None => { missing_key = true; captures @@ -705,12 +734,31 @@ fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) }); if missing_key { - path.to_string() + Some(path.to_string()) } else { - replaced.into_owned() + let replaced = replaced.into_owned(); + if dynamic_template_value_created_dot_path_segment(path, &replaced, regex) { + None + } else { + Some(replaced) + } } } +fn dynamic_template_value_created_dot_path_segment( + template: &str, + expanded: &str, + regex: &Regex, +) -> bool { + let template_path = template.split_once('?').map_or(template, |(path, _)| path); + let expanded_path = expanded.split_once('?').map_or(expanded, |(path, _)| path); + template_path.split('/').zip(expanded_path.split('/')).any( + |(template_segment, expanded_segment)| { + matches!(expanded_segment, "." | "..") && regex.is_match(template_segment) + }, + ) +} + fn maybe_add_gemini_stream_alt_sse( upstream_url: String, provider_api_format: &str, @@ -2288,4 +2336,85 @@ mod tests { "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar" ); } + + #[test] + fn custom_path_template_encodes_dynamic_model_as_one_path_segment() { + let transport = sample_transport( + "custom", + "gemini:generate_content", + "https://generativelanguage.googleapis.com", + Some("/v1beta/models/{model}:{action}"), + ); + + let url = build_transport_request_url( + &transport, + TransportRequestUrlParams { + provider_api_format: "gemini:generate_content", + mapped_model: Some("model/../../admin?key=attacker#fragment"), + upstream_is_stream: false, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ) + .expect("custom path should remain on the configured origin"); + + assert_eq!( + url, + "https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:generateContent" + ); + } + + #[test] + fn custom_path_template_rejects_dot_only_dynamic_path_segments() { + let transport = sample_transport( + "custom", + "claude:messages", + "https://api.example.com", + Some("/v1/messages/{model}/invoke"), + ); + + assert_eq!( + build_transport_request_url( + &transport, + TransportRequestUrlParams { + provider_api_format: "claude:messages", + mapped_model: Some(".."), + upstream_is_stream: false, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ), + None + ); + } + + #[test] + fn custom_path_template_query_values_cannot_inject_parameters() { + let transport = sample_transport( + "custom", + "claude:messages", + "https://api.example.com", + Some("/v1/messages?model={model}"), + ); + + let url = build_transport_request_url( + &transport, + TransportRequestUrlParams { + provider_api_format: "claude:messages", + mapped_model: Some("claude&admin=true#fragment"), + upstream_is_stream: false, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ) + .expect("query template should build"); + + assert_eq!( + url, + "https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment" + ); + } } 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 04b68587a..d84b3d519 100644 --- a/crates/aether-provider/transport/src/same_format_provider/mod.rs +++ b/crates/aether-provider/transport/src/same_format_provider/mod.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use serde::Serialize; use serde_json::Value; @@ -66,7 +67,7 @@ pub struct SameFormatProviderRequestBehavior { pub report_kind: &'static str, } -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct SameFormatProviderRequestBodyInput<'a> { pub body_json: &'a Value, pub mapped_model: &'a str, @@ -83,13 +84,64 @@ pub struct SameFormatProviderRequestBodyInput<'a> { pub enable_model_directives: bool, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +impl fmt::Debug for SameFormatProviderRequestBodyInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SameFormatProviderRequestBodyInput") + .field( + "body_json_bytes", + &serde_json::to_vec(self.body_json) + .ok() + .map(|bytes| bytes.len()), + ) + .field("mapped_model", &self.mapped_model) + .field("client_api_format", &self.client_api_format) + .field("provider_api_format", &self.provider_api_format) + .field("source_model", &self.source_model) + .field("family", &self.family) + .field("has_body_rules", &self.body_rules.is_some()) + .field( + "request_header_names", + &self + .request_headers + .map(|headers| headers.keys().map(|name| name.as_str()).collect::>()), + ) + .field("upstream_is_stream", &self.upstream_is_stream) + .field("force_body_stream_field", &self.force_body_stream_field) + .field("has_kiro_auth_config", &self.kiro_auth_config.is_some()) + .field("is_claude_code", &self.is_claude_code) + .field("enable_model_directives", &self.enable_model_directives) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq, Serialize)] pub struct SameFormatProviderRequestBodyOutput { pub body: Value, #[serde(default, skip_serializing_if = "Vec::is_empty")] pub compatibility_edits: Vec, } +impl fmt::Debug for SameFormatProviderRequestBodyOutput { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SameFormatProviderRequestBodyOutput") + .field( + "body_bytes", + &serde_json::to_vec(&self.body).ok().map(|bytes| bytes.len()), + ) + .field( + "compatibility_edits", + &self + .compatibility_edits + .iter() + .map(|edit| (edit.field.as_str(), edit.action, edit.detail.len())) + .collect::>(), + ) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize)] pub struct SameFormatProviderCompatibilityEdit { pub field: String, @@ -106,7 +158,7 @@ pub enum SameFormatProviderCompatibilityEditAction { OperatorRule, } -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct SameFormatProviderUpstreamUrlParams<'a> { pub provider_api_format: &'a str, pub mapped_model: &'a str, @@ -117,7 +169,26 @@ pub struct SameFormatProviderUpstreamUrlParams<'a> { pub provider_request_body: Option<&'a Value>, } -#[derive(Debug, Clone, Copy)] +impl fmt::Debug for SameFormatProviderUpstreamUrlParams<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SameFormatProviderUpstreamUrlParams") + .field("provider_api_format", &self.provider_api_format) + .field("mapped_model", &self.mapped_model) + .field("upstream_is_stream", &self.upstream_is_stream) + .field("has_request_query", &self.request_query.is_some()) + .field("request_query_len", &self.request_query.map(str::len)) + .field("kiro_api_region", &self.kiro_api_region) + .field("api_operation", &self.api_operation) + .field( + "has_provider_request_body", + &self.provider_request_body.is_some(), + ) + .finish() + } +} + +#[derive(Clone, Copy)] pub struct SameFormatProviderHeadersInput<'a> { pub headers: &'a http::HeaderMap, pub provider_request_body: &'a Value, @@ -132,6 +203,45 @@ pub struct SameFormatProviderHeadersInput<'a> { pub kiro_machine_id: Option<&'a str>, } +impl fmt::Debug for SameFormatProviderHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SameFormatProviderHeadersInput") + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field("has_header_rules", &self.header_rules.is_some()) + .field("behavior", &self.behavior) + .field("api_operation", &self.api_operation) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &self.auth_value.is_some()) + .field( + "extra_header_names", + &self.extra_headers.keys().collect::>(), + ) + .field("has_kiro_auth_config", &self.kiro_auth_config.is_some()) + .field("has_kiro_machine_id", &self.kiro_machine_id.is_some()) + .finish() + } +} + pub fn classify_same_format_provider_request_behavior( transport: &GatewayProviderTransportSnapshot, params: SameFormatProviderRequestBehaviorParams<'_>, @@ -721,6 +831,12 @@ pub fn build_same_format_provider_headers( provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string()); force_identity_accept_encoding(&mut provider_request_headers); } + let declared_connection_headers = + crate::headers::declared_connection_header_names(input.headers, input.extra_headers); + crate::headers::remove_declared_connection_headers( + &mut provider_request_headers, + &declared_connection_headers, + ); Some(provider_request_headers) } diff --git a/crates/aether-provider/transport/src/snapshot.rs b/crates/aether-provider/transport/src/snapshot.rs index d73a4a6d5..0fc132fbd 100644 --- a/crates/aether-provider/transport/src/snapshot.rs +++ b/crates/aether-provider/transport/src/snapshot.rs @@ -1,3 +1,4 @@ +use aether_contracts::redact_url_for_debug; use aether_data_contracts::repository::provider_catalog::{ StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; @@ -18,7 +19,7 @@ pub struct GatewayProviderTransportSnapshot { pub key: GatewayProviderTransportKey, } -#[derive(Debug, Clone, PartialEq, serde::Serialize)] +#[derive(Clone, PartialEq, serde::Serialize)] pub struct GatewayProviderTransportProvider { pub id: String, pub name: String, @@ -29,13 +30,45 @@ pub struct GatewayProviderTransportProvider { pub enable_format_conversion: bool, pub concurrent_limit: Option, pub max_retries: Option, + #[serde(skip_serializing)] pub proxy: Option, pub request_timeout_secs: Option, pub stream_first_byte_timeout_secs: Option, + #[serde(skip_serializing)] pub config: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize)] +impl std::fmt::Debug for GatewayProviderTransportProvider { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GatewayProviderTransportProvider") + .field("id", &self.id) + .field("name", &self.name) + .field("provider_type", &self.provider_type) + .field( + "website", + &self.website.as_deref().map(redact_url_for_debug), + ) + .field("is_active", &self.is_active) + .field( + "keep_priority_on_conversion", + &self.keep_priority_on_conversion, + ) + .field("enable_format_conversion", &self.enable_format_conversion) + .field("concurrent_limit", &self.concurrent_limit) + .field("max_retries", &self.max_retries) + .field("has_proxy", &self.proxy.is_some()) + .field("request_timeout_secs", &self.request_timeout_secs) + .field( + "stream_first_byte_timeout_secs", + &self.stream_first_byte_timeout_secs, + ) + .field("has_config", &self.config.is_some()) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize)] pub struct GatewayProviderTransportEndpoint { pub id: String, pub provider_id: String, @@ -44,16 +77,49 @@ pub struct GatewayProviderTransportEndpoint { pub endpoint_kind: Option, pub is_active: bool, pub base_url: String, + #[serde(skip_serializing)] pub header_rules: Option, + #[serde(skip_serializing)] pub body_rules: Option, pub max_retries: Option, pub custom_path: Option, + #[serde(skip_serializing)] pub config: Option, + #[serde(skip_serializing)] pub format_acceptance_config: Option, + #[serde(skip_serializing)] pub proxy: Option, } -#[derive(Debug, Clone, PartialEq, serde::Serialize)] +impl std::fmt::Debug for GatewayProviderTransportEndpoint { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GatewayProviderTransportEndpoint") + .field("id", &self.id) + .field("provider_id", &self.provider_id) + .field("api_format", &self.api_format) + .field("api_family", &self.api_family) + .field("endpoint_kind", &self.endpoint_kind) + .field("is_active", &self.is_active) + .field("base_url", &redact_url_for_debug(&self.base_url)) + .field("has_header_rules", &self.header_rules.is_some()) + .field("has_body_rules", &self.body_rules.is_some()) + .field("max_retries", &self.max_retries) + .field( + "custom_path_len", + &self.custom_path.as_ref().map(String::len), + ) + .field("has_config", &self.config.is_some()) + .field( + "has_format_acceptance_config", + &self.format_acceptance_config.is_some(), + ) + .field("has_proxy", &self.proxy.is_some()) + .finish() + } +} + +#[derive(Clone, PartialEq, serde::Serialize)] pub struct GatewayProviderTransportKey { pub id: String, pub provider_id: String, @@ -68,13 +134,42 @@ pub struct GatewayProviderTransportKey { pub rate_multipliers: Option, pub global_priority_by_format: Option, pub expires_at_unix_secs: Option, + #[serde(skip_serializing)] pub proxy: Option, + #[serde(skip_serializing)] pub fingerprint: Option, + #[serde(skip_serializing)] pub upstream_metadata: Option, + #[serde(skip_serializing)] pub decrypted_api_key: String, + #[serde(skip_serializing)] pub decrypted_auth_config: Option, } +impl std::fmt::Debug for GatewayProviderTransportKey { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GatewayProviderTransportKey") + .field("id", &self.id) + .field("provider_id", &self.provider_id) + .field("name", &self.name) + .field("auth_type", &self.auth_type) + .field("is_active", &self.is_active) + .field("api_formats", &self.api_formats) + .field("allowed_models", &self.allowed_models) + .field("expires_at_unix_secs", &self.expires_at_unix_secs) + .field("has_proxy", &self.proxy.is_some()) + .field("has_fingerprint", &self.fingerprint.is_some()) + .field("has_upstream_metadata", &self.upstream_metadata.is_some()) + .field("decrypted_api_key", &"[REDACTED]") + .field( + "decrypted_auth_config", + &self.decrypted_auth_config.as_ref().map(|_| "[REDACTED]"), + ) + .finish_non_exhaustive() + } +} + #[async_trait] pub trait ProviderTransportSnapshotSource: Send + Sync { fn encryption_key(&self) -> Option<&str>; @@ -335,6 +430,23 @@ mod tests { ) } + fn seal_bound_provider_credential( + provider_id: &str, + key_id: &str, + field: &str, + plaintext: &str, + ) -> String { + let purpose = format!( + "provider-catalog-credential-bound-v2\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={field}", + provider_id.len(), + key_id.len(), + ); + let protected = format!("{purpose}\0{plaintext}"); + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &protected) + .expect("bound credential should encrypt"); + format!("aether-provider-catalog-credential-v2:aether-runtime-secret-v1:{ciphertext}") + } + #[tokio::test] async fn reads_decrypted_provider_transport_snapshot() { let state = read_state(); @@ -411,6 +523,56 @@ mod tests { ); } + #[tokio::test] + async fn transport_snapshot_debug_and_serialization_exclude_credentials() { + let state = read_state(); + let mut snapshot = + read_provider_transport_snapshot(&state, "provider-1", "endpoint-1", "key-1") + .await + .expect("snapshot should read") + .expect("snapshot should exist"); + snapshot.provider.proxy = Some(serde_json::json!({"password": "provider-proxy-canary"})); + snapshot.provider.config = + Some(serde_json::json!({"authorization": "provider-config-canary"})); + snapshot.endpoint.header_rules = + Some(serde_json::json!({"authorization": "endpoint-header-canary"})); + snapshot.endpoint.body_rules = Some(serde_json::json!({"token": "endpoint-body-canary"})); + snapshot.endpoint.config = Some(serde_json::json!({"secret": "endpoint-config-canary"})); + snapshot.endpoint.format_acceptance_config = + Some(serde_json::json!({"secret": "endpoint-acceptance-canary"})); + snapshot.endpoint.proxy = Some(serde_json::json!({"password": "endpoint-proxy-canary"})); + snapshot.key.proxy = Some(serde_json::json!({"password": "key-proxy-canary"})); + snapshot.key.fingerprint = Some(serde_json::json!({"cookie": "key-fingerprint-canary"})); + snapshot.key.upstream_metadata = + Some(serde_json::json!({"access_token": "key-metadata-canary"})); + + let debug = format!("{snapshot:?}"); + let serialized = serde_json::to_string(&snapshot).expect("snapshot should serialize"); + for secret in [ + "sk-live-openai", + "rt-1", + "provider-proxy-canary", + "provider-config-canary", + "endpoint-header-canary", + "endpoint-body-canary", + "endpoint-config-canary", + "endpoint-acceptance-canary", + "endpoint-proxy-canary", + "key-proxy-canary", + "key-fingerprint-canary", + "key-metadata-canary", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}"); + assert!( + !serialized.contains(secret), + "serialization leaked {secret}" + ); + } + assert!(debug.contains("[REDACTED]")); + assert!(!serialized.contains("decrypted_api_key")); + assert!(!serialized.contains("decrypted_auth_config")); + } + #[tokio::test] async fn reads_snapshot_when_provider_key_api_key_is_null() { let mut key = sample_key(); @@ -534,7 +696,7 @@ mod tests { } #[tokio::test] - async fn accepts_plaintext_legacy_key_material() { + async fn rejects_plaintext_legacy_key_material() { let provider = sample_provider(); let endpoint = StoredProviderCatalogEndpoint::new( "endpoint-legacy-1".to_string(), @@ -584,25 +746,45 @@ mod tests { Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()), ); - let snapshot = read_provider_transport_snapshot( + let error = read_provider_transport_snapshot( &state, "provider-1", "endpoint-legacy-1", "key-legacy-1", ) .await - .expect("snapshot read should succeed") - .expect("snapshot should exist"); + .expect_err("plaintext credentials must be rejected"); - assert_eq!(snapshot.key.decrypted_api_key, "sk-plaintext-openai"); - assert_eq!(snapshot.key.decrypted_auth_config, None); + assert!(matches!(error, DataLayerError::UnexpectedValue(message) + if message.contains("provider_api_keys.api_key is not an authenticated ciphertext"))); + } + + #[test] + fn decrypts_record_bound_v2_credentials_and_rejects_copying() { + let mut key = sample_key(); + key.encrypted_api_key = Some(seal_bound_provider_credential( + "provider-1", + "key-1", + "api-key", + "bound-api-key", + )); + key.encrypted_auth_config = Some(seal_bound_provider_credential( + "provider-1", + "key-1", + "auth-config", + r#"{"refresh_token":"bound-refresh"}"#, + )); + + let mapped = map_key(key.clone(), DEVELOPMENT_ENCRYPTION_KEY, &[]) + .expect("matching record binding should decrypt"); + assert_eq!(mapped.decrypted_api_key, "bound-api-key"); assert_eq!( - snapshot.endpoint.header_rules, - Some(serde_json::json!([ - {"action":"set","key":"x-test","value":"1"}, - {"action":"set","key":"x-account-id","value":"acc-legacy"} - ])) + mapped.decrypted_auth_config.as_deref(), + Some(r#"{"refresh_token":"bound-refresh"}"#) ); + + key.id = "key-2".to_string(); + assert!(map_key(key, DEVELOPMENT_ENCRYPTION_KEY, &[]).is_err()); } #[tokio::test] diff --git a/crates/aether-provider/transport/src/snapshot_mapping.rs b/crates/aether-provider/transport/src/snapshot_mapping.rs index 03969964d..f5a24e460 100644 --- a/crates/aether-provider/transport/src/snapshot_mapping.rs +++ b/crates/aether-provider/transport/src/snapshot_mapping.rs @@ -8,6 +8,11 @@ use super::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider, }; +const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY: &str = "aether-provider-catalog-credential-"; +const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2: &str = "aether-provider-catalog-credential-v2:"; +const PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2: &str = "provider-catalog-credential-bound-v2"; +const RUNTIME_SECRET_ENVELOPE_PREFIX: &str = "aether-runtime-secret-v1:"; + pub(super) fn map_provider( provider: StoredProviderCatalogProvider, ) -> GatewayProviderTransportProvider { @@ -64,6 +69,9 @@ pub(super) fn map_key( encryption_key, fallback_encryption_keys, ciphertext, + &key.provider_id, + &key.id, + "api-key", "provider_api_keys.api_key", ) }) @@ -79,6 +87,9 @@ pub(super) fn map_key( encryption_key, fallback_encryption_keys, ciphertext, + &key.provider_id, + &key.id, + "auth-config", "provider_api_keys.auth_config", ) }) @@ -125,12 +136,78 @@ fn decrypt_secret( encryption_key: &str, fallback_encryption_keys: &[String], ciphertext: &str, + provider_id: &str, + key_id: &str, + field: &str, field_name: &str, ) -> Result { - if should_use_plaintext_secret(ciphertext, field_name) { - return Ok(ciphertext.trim().to_string()); + if let Some(runtime_envelope) = ciphertext.strip_prefix(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2) + { + let inner_ciphertext = runtime_envelope + .strip_prefix(RUNTIME_SECRET_ENVELOPE_PREFIX) + .ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "{field_name} has an invalid provider catalog credential envelope" + )) + })?; + let protected = decrypt_fernet_with_fallbacks( + encryption_key, + fallback_encryption_keys, + inner_ciphertext, + field_name, + )?; + let purpose = provider_catalog_credential_purpose(provider_id, key_id, field); + return protected + .strip_prefix(&purpose) + .and_then(|value| value.strip_prefix('\0')) + .map(ToOwned::to_owned) + .ok_or_else(|| { + DataLayerError::UnexpectedValue(format!( + "{field_name} provider catalog credential authentication failed" + )) + }); + } + if ciphertext.starts_with(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY) + || ciphertext.starts_with("aether-") + { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} has an unsupported or incorrectly bound Aether secret envelope" + ))); + } + if !looks_like_python_fernet_ciphertext(ciphertext) { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} is not an authenticated ciphertext" + ))); } + let plaintext = decrypt_fernet_with_fallbacks( + encryption_key, + fallback_encryption_keys, + ciphertext, + field_name, + )?; + if plaintext.contains('\0') { + return Err(DataLayerError::UnexpectedValue(format!( + "{field_name} legacy ciphertext contains reserved framing" + ))); + } + Ok(plaintext) +} + +fn provider_catalog_credential_purpose(provider_id: &str, key_id: &str, field: &str) -> String { + format!( + "{PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2}\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={field}", + provider_id.len(), + key_id.len(), + ) +} + +fn decrypt_fernet_with_fallbacks( + encryption_key: &str, + fallback_encryption_keys: &[String], + ciphertext: &str, + field_name: &str, +) -> Result { match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) { Ok(value) => Ok(value), Err(error) => { @@ -166,29 +243,6 @@ pub(super) fn fallback_encryption_keys(primary_encryption_key: &str) -> Vec bool { - let ciphertext = ciphertext.trim(); - if ciphertext.is_empty() { - return false; - } - - match field_name { - "provider_api_keys.api_key" => { - if ciphertext.starts_with('{') || ciphertext.starts_with('[') { - return false; - } - !looks_like_python_fernet_ciphertext(ciphertext) - } - "provider_api_keys.auth_config" => { - if ciphertext.starts_with('{') || ciphertext.starts_with('[') { - return true; - } - false - } - _ => false, - } -} - fn normalize_string_list( raw: Option, field_name: &str, diff --git a/crates/aether-provider/transport/src/standard/mod.rs b/crates/aether-provider/transport/src/standard/mod.rs index 0cc24b961..56b5c0402 100644 --- a/crates/aether-provider/transport/src/standard/mod.rs +++ b/crates/aether-provider/transport/src/standard/mod.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use serde_json::Value; @@ -16,7 +17,7 @@ use crate::snapshot::GatewayProviderTransportSnapshot; use crate::url::{build_openai_chat_url, build_openai_responses_url}; use crate::vertex::uses_vertex_api_key_query_auth; -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct StandardProviderRequestHeadersInput<'a> { pub transport: &'a GatewayProviderTransportSnapshot, pub provider_api_format: &'a str, @@ -31,13 +32,63 @@ pub struct StandardProviderRequestHeadersInput<'a> { pub upstream_is_stream: bool, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl fmt::Debug for StandardProviderRequestHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("StandardProviderRequestHeadersInput") + .field("transport", &self.transport) + .field("provider_api_format", &self.provider_api_format) + .field("same_format", &self.same_format) + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field( + "extra_header_names", + &self.extra_headers.keys().collect::>(), + ) + .field("has_header_rules", &self.header_rules.is_some()) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field("upstream_is_stream", &self.upstream_is_stream) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct StandardProviderRequestHeaders { pub headers: BTreeMap, pub auth_header: String, pub auth_value: String, } +impl fmt::Debug for StandardProviderRequestHeaders { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("StandardProviderRequestHeaders") + .field("header_names", &self.headers.keys().collect::>()) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .finish() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum StandardPlanFallbackAcceptPolicy { None, @@ -47,7 +98,7 @@ pub enum StandardPlanFallbackAcceptPolicy { ProviderEventStreamIfMissing, } -#[derive(Debug)] +#[derive(Clone)] pub struct StandardPlanFallbackHeadersInput<'a> { pub request_headers: &'a http::HeaderMap, pub existing_provider_request_headers: BTreeMap, @@ -62,6 +113,44 @@ pub struct StandardPlanFallbackHeadersInput<'a> { pub accept_policy: StandardPlanFallbackAcceptPolicy, } +impl fmt::Debug for StandardPlanFallbackHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("StandardPlanFallbackHeadersInput") + .field( + "request_header_names", + &self + .request_headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field( + "existing_provider_header_names", + &self + .existing_provider_request_headers + .keys() + .collect::>(), + ) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &self.auth_value.is_some()) + .field( + "extra_header_names", + &self.extra_headers.keys().collect::>(), + ) + .field("content_type", &self.content_type) + .field("provider_api_format", &self.provider_api_format) + .field("client_api_format", &self.client_api_format) + .field("upstream_is_stream", &self.upstream_is_stream) + .field( + "build_from_request_when_empty", + &self.build_from_request_when_empty, + ) + .field("accept_policy", &self.accept_policy) + .finish() + } +} + pub fn build_standard_plan_fallback_openai_chat_url( upstream_base_url: &str, request_query: Option<&str>, @@ -152,6 +241,12 @@ pub fn build_standard_plan_fallback_headers( force_identity_accept_encoding(&mut headers); } + let declared_connection_headers = crate::headers::declared_connection_header_names( + input.request_headers, + input.extra_headers, + ); + crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers); + headers } @@ -301,6 +396,10 @@ pub fn build_standard_provider_request_headers( force_identity_accept_encoding(&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); + Some(StandardProviderRequestHeaders { headers, auth_header, diff --git a/crates/aether-provider/transport/src/url.rs b/crates/aether-provider/transport/src/url.rs index 9fd9f2a06..32df41643 100644 --- a/crates/aether-provider/transport/src/url.rs +++ b/crates/aether-provider/transport/src/url.rs @@ -5,6 +5,42 @@ use url::Url; pub(crate) const GATEWAY_CREDENTIAL_QUERY_KEYS: &[&str] = &["key"]; +pub(crate) fn encode_url_path_segment(value: &str) -> String { + const HEX: &[u8; 16] = b"0123456789ABCDEF"; + + let mut encoded = String::with_capacity(value.len()); + for byte in value.bytes() { + if byte.is_ascii_alphanumeric() + || matches!( + byte, + b'-' | b'.' + | b'_' + | b'~' + | b'!' + | b'$' + | b'&' + | b'\'' + | b'(' + | b')' + | b'*' + | b'+' + | b',' + | b';' + | b'=' + | b':' + | b'@' + ) + { + encoded.push(char::from(byte)); + } else { + encoded.push('%'); + encoded.push(char::from(HEX[usize::from(byte >> 4)])); + encoded.push(char::from(HEX[usize::from(byte & 0x0f)])); + } + } + encoded +} + pub(crate) fn strip_gateway_credential_query_parameters(query: Option<&str>) -> Option { let query = query.map(str::trim).filter(|value| !value.is_empty())?; let mut serializer = form_urlencoded::Serializer::new(String::new()); @@ -161,6 +197,7 @@ pub fn build_gemini_content_url( if trimmed_base_url.is_empty() || trimmed_model.is_empty() { return None; } + let encoded_model = encode_url_path_segment(trimmed_model); let operation = if stream { "streamGenerateContent" @@ -168,12 +205,12 @@ pub fn build_gemini_content_url( "generateContent" }; let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") { - format!("{trimmed_base_url}/models/{trimmed_model}:{operation}") + format!("{trimmed_base_url}/models/{encoded_model}:{operation}") } else if gemini_content_base_url_contains_model_path(trimmed_base_url) { let trimmed_base_url = strip_gemini_content_action(trimmed_base_url); format!("{trimmed_base_url}:{operation}") } else { - format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:{operation}") + format!("{trimmed_base_url}/v1beta/models/{encoded_model}:{operation}") }; append_merged_query(&mut url, base_query, None, query, &["key"]); Some(url) @@ -221,13 +258,14 @@ pub fn build_gemini_video_predict_long_running_url( if trimmed_base_url.is_empty() || trimmed_model.is_empty() { return None; } + let encoded_model = encode_url_path_segment(trimmed_model); let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") { - format!("{trimmed_base_url}/models/{trimmed_model}:predictLongRunning") + format!("{trimmed_base_url}/models/{encoded_model}:predictLongRunning") } else if gemini_content_base_url_contains_model_path(trimmed_base_url) { format!("{trimmed_base_url}:predictLongRunning") } else { - format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:predictLongRunning") + format!("{trimmed_base_url}/v1beta/models/{encoded_model}:predictLongRunning") }; append_merged_query(&mut url, base_query, None, query, &["key"]); Some(url) @@ -412,9 +450,27 @@ fn bigmodel_coding_models_base_is_supported(base_url: &str) -> bool { fn looks_like_vertex_ai_host(host: &str) -> bool { const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com"; + let host = host.trim().to_ascii_lowercase(); host == VERTEX_AI_HOST - || host.ends_with(&format!(".{VERTEX_AI_HOST}")) - || host.ends_with(&format!("-{VERTEX_AI_HOST}")) + || host + .strip_suffix(&format!("-{VERTEX_AI_HOST}")) + .is_some_and(is_vertex_region_label) +} + +fn is_vertex_region_label(value: &str) -> bool { + !value.is_empty() + && value.len() <= 63 + && value + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && value + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric) + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') } fn split_path_query(path: &str) -> (&str, Option<&str>) { @@ -504,7 +560,8 @@ mod tests { build_gemini_content_url, build_gemini_files_passthrough_url, build_gemini_video_predict_long_running_url, build_openai_chat_url, build_openai_compatible_models_url, build_openai_image_url, build_openai_responses_url, - build_openai_search_url, build_passthrough_path_url, normalize_gemini_content_action_path, + build_openai_search_url, build_passthrough_path_url, encode_url_path_segment, + normalize_gemini_content_action_path, }; #[test] @@ -868,4 +925,41 @@ mod tests { ) ); } + + #[test] + fn gemini_model_names_cannot_inject_path_query_or_fragment_components() { + let model = "model/../../admin?key=attacker#fragment"; + assert_eq!( + build_gemini_content_url( + "https://generativelanguage.googleapis.com/v1beta", + model, + false, + None, + ) + .as_deref(), + Some( + "https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:generateContent" + ) + ); + assert_eq!( + build_gemini_video_predict_long_running_url( + "https://generativelanguage.googleapis.com/v1beta", + model, + None, + ) + .as_deref(), + Some( + "https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:predictLongRunning" + ) + ); + } + + #[test] + fn dynamic_path_segment_encoding_keeps_raw_values_in_one_segment() { + assert_eq!( + encode_url_path_segment("gemini+2.5@preview~/model%2Fraw"), + "gemini+2.5@preview~%2Fmodel%252Fraw" + ); + assert_eq!(encode_url_path_segment(".."), ".."); + } } diff --git a/crates/aether-provider/transport/src/vertex/auth.rs b/crates/aether-provider/transport/src/vertex/auth.rs index e6fcb9b28..19e631087 100644 --- a/crates/aether-provider/transport/src/vertex/auth.rs +++ b/crates/aether-provider/transport/src/vertex/auth.rs @@ -1,22 +1,19 @@ use std::collections::BTreeMap; +use aether_crypto::{rsa_pkcs1_sha256_sign, RsaPkcs1Sha256Error}; use async_trait::async_trait; use base64::engine::general_purpose::URL_SAFE_NO_PAD; use base64::Engine as _; -use rsa::pkcs1::DecodeRsaPrivateKey; -use rsa::pkcs1v15::SigningKey; -use rsa::pkcs8::DecodePrivateKey; -use rsa::signature::{SignatureEncoding, Signer}; -use rsa::RsaPrivateKey; use serde_json::{json, Value}; use sha2::{Digest, Sha256}; -use url::form_urlencoded; +use url::{form_urlencoded, Url}; use super::super::oauth_refresh::{ CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthRefreshAdapter, LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth, }; use super::super::snapshot::GatewayProviderTransportSnapshot; +use super::context::is_valid_vertex_region; pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key"; pub const VERTEX_SERVICE_ACCOUNT_AUTH_HEADER: &str = "authorization"; @@ -25,13 +22,23 @@ pub const GOOGLE_OAUTH_TOKEN_URL: &str = "https://oauth2.googleapis.com/token"; const GOOGLE_CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform"; const SERVICE_ACCOUNT_REFRESH_SKEW_SECS: u64 = 120; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct VertexApiKeyQueryAuth { pub name: &'static str, pub value: String, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl std::fmt::Debug for VertexApiKeyQueryAuth { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("VertexApiKeyQueryAuth") + .field("name", &self.name) + .field("value", &"[REDACTED]") + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct VertexServiceAccountAuthConfig { pub client_email: String, pub private_key: String, @@ -41,6 +48,20 @@ pub struct VertexServiceAccountAuthConfig { pub model_regions: BTreeMap, } +impl std::fmt::Debug for VertexServiceAccountAuthConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("VertexServiceAccountAuthConfig") + .field("client_email", &self.client_email) + .field("private_key", &"[REDACTED]") + .field("project_id", &self.project_id) + .field("token_uri", &"[REDACTED]") + .field("region", &self.region) + .field("model_regions", &self.model_regions) + .finish() + } +} + pub fn resolve_local_vertex_api_key_query_auth( transport: &GatewayProviderTransportSnapshot, ) -> Option { @@ -101,9 +122,8 @@ fn parse_vertex_service_account_auth_config_value( let client_email = json_string(value.get("client_email"))?; let private_key = json_string(value.get("private_key"))?; let project_id = json_string(value.get("project_id"))?; - let token_uri = - json_string(value.get("token_uri")).unwrap_or_else(|| GOOGLE_OAUTH_TOKEN_URL.to_string()); - let region = json_string(value.get("region")); + let token_uri = resolve_vertex_service_account_token_uri(value.get("token_uri"))?; + let region = json_string(value.get("region")).filter(|value| is_valid_vertex_region(value)); let model_regions = value .get("model_regions") .and_then(Value::as_object) @@ -113,7 +133,7 @@ fn parse_vertex_service_account_auth_config_value( .filter_map(|(model, region)| { let model = model.trim(); let region = region.as_str()?.trim(); - (!model.is_empty() && !region.is_empty()) + (!model.is_empty() && is_valid_vertex_region(region)) .then(|| (model.to_string(), region.to_string())) }) .collect::>() @@ -130,6 +150,31 @@ fn parse_vertex_service_account_auth_config_value( }) } +fn resolve_vertex_service_account_token_uri(value: Option<&Value>) -> Option { + let Some(value) = value else { + return Some(GOOGLE_OAUTH_TOKEN_URL.to_string()); + }; + let raw = value.as_str()?.trim(); + if raw.is_empty() { + return None; + } + let parsed = Url::parse(raw).ok()?; + if parsed.scheme() != "https" + || !parsed.username().is_empty() + || parsed.password().is_some() + || !parsed + .host_str() + .is_some_and(|host| host.eq_ignore_ascii_case("oauth2.googleapis.com")) + || parsed.port().is_some() + || parsed.path() != "/token" + || parsed.query().is_some() + || parsed.fragment().is_some() + { + return None; + } + Some(GOOGLE_OAUTH_TOKEN_URL.to_string()) +} + fn json_string(value: Option<&Value>) -> Option { value .and_then(Value::as_str) @@ -352,29 +397,17 @@ pub fn build_vertex_service_account_assertion( })?, ); let message = format!("{header}.{payload}"); - let private_key = decode_vertex_service_account_private_key(auth_config.private_key.as_str())?; - let signing_key = SigningKey::::new(private_key); - let signature = signing_key.sign(message.as_bytes()); - Ok(format!( - "{message}.{}", - URL_SAFE_NO_PAD.encode(signature.to_bytes()) - )) -} - -fn decode_vertex_service_account_private_key( - private_key_pem: &str, -) -> Result { - match RsaPrivateKey::from_pkcs8_pem(private_key_pem) { - Ok(private_key) => Ok(private_key), - Err(pkcs8_err) => RsaPrivateKey::from_pkcs1_pem(private_key_pem).map_err(|pkcs1_err| { - LocalOAuthRefreshError::InvalidResponse { - provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE, - message: format!( - "vertex service account private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}" - ), + let signature = rsa_pkcs1_sha256_sign(auth_config.private_key.as_bytes(), message.as_bytes()) + .map_err(|error| LocalOAuthRefreshError::InvalidResponse { + provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE, + message: match error { + RsaPkcs1Sha256Error::InvalidPrivateKey => { + "vertex service account private_key parse failed".to_string() } - }), - } + _ => "vertex service account signing failed".to_string(), + }, + })?; + Ok(format!("{message}.{}", URL_SAFE_NO_PAD.encode(signature))) } fn service_account_token_expires_soon(expires_at_unix_secs: Option) -> bool { @@ -387,7 +420,7 @@ fn service_account_token_expires_soon(expires_at_unix_secs: Option) -> bool } fn body_excerpt(value: &str) -> String { - value.chars().take(500).collect() + aether_oauth::core::redacted_oauth_error_body_excerpt(value) } #[cfg(test)] @@ -399,19 +432,46 @@ mod tests { GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, }; - use rsa::pkcs1::{EncodeRsaPrivateKey, LineEnding}; - use rsa::rand_core::OsRng; - use rsa::RsaPrivateKey; + use aws_lc_rs::encoding::{AsDer, Pkcs8V1Der}; + use aws_lc_rs::rsa::{KeyPair as AwsRsaKeyPair, KeySize}; + use aws_lc_rs::signature::{KeyPair as _, UnparsedPublicKey, RSA_PKCS1_2048_8192_SHA256}; + use base64::engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}; + use base64::Engine as _; + use serde_json::{json, Value}; use super::{ - decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config, + build_vertex_service_account_assertion, parse_vertex_service_account_auth_config, resolve_local_vertex_api_key_query_auth, supports_local_vertex_service_account_auth_resolution, - vertex_service_account_credential_fingerprint, VertexServiceAccountRefreshAdapter, + vertex_service_account_credential_fingerprint, VertexApiKeyQueryAuth, + VertexServiceAccountAuthConfig, VertexServiceAccountRefreshAdapter, GOOGLE_OAUTH_TOKEN_URL, VERTEX_API_KEY_QUERY_PARAM, VERTEX_SERVICE_ACCOUNT_AUTH_HEADER, VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE, }; + #[test] + fn vertex_auth_debug_output_redacts_api_keys_and_private_keys() { + let query_auth = VertexApiKeyQueryAuth { + name: VERTEX_API_KEY_QUERY_PARAM, + value: "vertex-api-key-canary".to_string(), + }; + let service_account = VertexServiceAccountAuthConfig { + client_email: "service@example.invalid".to_string(), + private_key: "vertex-private-key-canary".to_string(), + project_id: "project-1".to_string(), + token_uri: GOOGLE_OAUTH_TOKEN_URL.to_string(), + region: None, + model_regions: std::collections::BTreeMap::new(), + }; + + let query_debug = format!("{query_auth:?}"); + assert!(!query_debug.contains("vertex-api-key-canary")); + assert!(query_debug.contains("[REDACTED]")); + let service_account_debug = format!("{service_account:?}"); + assert!(!service_account_debug.contains("vertex-private-key-canary")); + assert!(service_account_debug.contains("[REDACTED]")); + } + fn sample_transport() -> GatewayProviderTransportSnapshot { GatewayProviderTransportSnapshot { provider: GatewayProviderTransportProvider { @@ -469,6 +529,32 @@ mod tests { } } + fn read_der_tlv<'a>(input: &mut &'a [u8], expected_tag: u8) -> &'a [u8] { + assert_eq!(input.first().copied(), Some(expected_tag)); + let length_byte = input[1]; + let (header_len, value_len) = if length_byte & 0x80 == 0 { + (2, usize::from(length_byte)) + } else { + let length_bytes = usize::from(length_byte & 0x7f); + let value_len = input[2..2 + length_bytes] + .iter() + .fold(0usize, |value, byte| (value << 8) | usize::from(*byte)); + (2 + length_bytes, value_len) + }; + let end = header_len + value_len; + let value = &input[header_len..end]; + *input = &input[end..]; + value + } + + fn pkcs1_private_key_from_pkcs8(pkcs8: &[u8]) -> Vec { + let mut input = pkcs8; + let mut sequence = read_der_tlv(&mut input, 0x30); + let _version = read_der_tlv(&mut sequence, 0x02); + let _algorithm = read_der_tlv(&mut sequence, 0x30); + read_der_tlv(&mut sequence, 0x04).to_vec() + } + fn sample_service_account_transport(private_key: &str) -> GatewayProviderTransportSnapshot { let mut transport = sample_transport(); transport.key.auth_type = "service_account".to_string(); @@ -570,6 +656,7 @@ mod tests { assert_eq!(config.client_email, "svc@example.iam.gserviceaccount.com"); assert_eq!(config.project_id, "demo-project"); + assert_eq!(config.token_uri, GOOGLE_OAUTH_TOKEN_URL); assert_eq!(config.region.as_deref(), Some("global")); assert_eq!( config @@ -580,6 +667,69 @@ mod tests { ); } + #[test] + fn service_account_regions_reject_url_syntax() { + let raw = r#"{ + "client_email":"svc@example.iam.gserviceaccount.com", + "private_key":"TEST-PRIVATE-KEY", + "project_id":"demo-project", + "region":"attacker.example/", + "model_regions":{ + "gemini-2.0-flash":"attacker.example/", + "gemini-2.5-pro":"us-central1" + } + }"#; + let config = parse_vertex_service_account_auth_config(Some(raw)) + .expect("service account config should parse"); + assert!(config.region.is_none()); + assert!(!config.model_regions.contains_key("gemini-2.0-flash")); + assert_eq!( + config + .model_regions + .get("gemini-2.5-pro") + .map(String::as_str), + Some("us-central1") + ); + } + + #[test] + fn service_account_token_uri_is_limited_to_google_oauth_endpoint() { + let config_with_token_uri = |token_uri: Value| { + serde_json::json!({ + "client_email": "svc@example.iam.gserviceaccount.com", + "private_key": "TEST-PRIVATE-KEY", + "project_id": "demo-project", + "token_uri": token_uri, + }) + .to_string() + }; + + let official = parse_vertex_service_account_auth_config(Some(&config_with_token_uri( + Value::String(GOOGLE_OAUTH_TOKEN_URL.to_string()), + ))) + .expect("official Google OAuth token URI should be accepted"); + assert_eq!(official.token_uri, GOOGLE_OAUTH_TOKEN_URL); + + for token_uri in [ + Value::String("http://oauth2.googleapis.com/token".to_string()), + Value::String("https://127.0.0.1/token".to_string()), + Value::String("https://oauth2.googleapis.com.evil.example/token".to_string()), + Value::String("https://user@oauth2.googleapis.com/token".to_string()), + Value::String("https://oauth2.googleapis.com:8443/token".to_string()), + Value::String("https://oauth2.googleapis.com/token/../metadata".to_string()), + Value::String("https://oauth2.googleapis.com/token?target=metadata".to_string()), + Value::String("https://oauth2.googleapis.com/token#fragment".to_string()), + Value::String(String::new()), + Value::Null, + ] { + let raw = config_with_token_uri(token_uri.clone()); + assert!( + parse_vertex_service_account_auth_config(Some(&raw)).is_none(), + "token URI should be rejected: {token_uri}" + ); + } + } + #[test] fn supports_vertex_service_account_auth_resolution() { let mut transport = sample_transport(); @@ -600,14 +750,35 @@ mod tests { } #[test] - fn decodes_pkcs1_service_account_private_key() { - let mut rng = OsRng; - let private_key = RsaPrivateKey::new(&mut rng, 1024) - .expect("test RSA private key should generate") - .to_pkcs1_pem(LineEnding::LF) - .expect("test RSA private key should encode as PKCS#1 PEM"); - - decode_vertex_service_account_private_key(private_key.as_str()) - .expect("PKCS#1 private key should decode"); + fn signs_with_2048_bit_pkcs1_service_account_private_key() { + let key_pair = AwsRsaKeyPair::generate(KeySize::Rsa2048) + .expect("2048-bit test RSA private key should generate"); + let pkcs8 = AsDer::>::as_der(&key_pair) + .expect("test RSA private key should encode as PKCS#8"); + let pkcs1 = pkcs1_private_key_from_pkcs8(pkcs8.as_ref()); + let private_key = format!( + "-----BEGIN RSA PRIVATE KEY-----\n{}\n-----END RSA PRIVATE KEY-----", + STANDARD.encode(pkcs1) + ); + let auth_config = parse_vertex_service_account_auth_config(Some( + &json!({ + "client_email": "svc@example.iam.gserviceaccount.com", + "private_key": private_key, + "project_id": "demo-project" + }) + .to_string(), + )) + .expect("service account config should parse"); + let assertion = build_vertex_service_account_assertion(&auth_config, 1_700_000_000) + .expect("PKCS#1 private key should sign"); + let parts = assertion.split('.').collect::>(); + assert_eq!(parts.len(), 3); + let message = format!("{}.{}", parts[0], parts[1]); + let signature = URL_SAFE_NO_PAD + .decode(parts[2]) + .expect("JWT signature should decode"); + UnparsedPublicKey::new(&RSA_PKCS1_2048_8192_SHA256, key_pair.public_key().as_ref()) + .verify(message.as_bytes(), &signature) + .expect("AWS-LC signature should verify"); } } diff --git a/crates/aether-provider/transport/src/vertex/context.rs b/crates/aether-provider/transport/src/vertex/context.rs index e8a777f01..e50ae9539 100644 --- a/crates/aether-provider/transport/src/vertex/context.rs +++ b/crates/aether-provider/transport/src/vertex/context.rs @@ -5,6 +5,26 @@ use super::super::snapshot::GatewayProviderTransportSnapshot; const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com"; +/// Vertex regions are interpolated into regional service hostnames and path +/// segments. Keep them to one DNS label so imported credential metadata cannot +/// redirect bearer-token requests to another origin. +pub fn is_valid_vertex_region(value: &str) -> bool { + let value = value.trim(); + !value.is_empty() + && value.len() <= 63 + && value + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && value + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric) + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') +} + pub fn looks_like_vertex_ai_host(base_url: &str) -> bool { let trimmed = base_url.trim(); if trimmed.is_empty() { @@ -14,6 +34,20 @@ pub fn looks_like_vertex_ai_host(base_url: &str) -> bool { let Ok(parsed) = Url::parse(trimmed) else { return false; }; + // A service-account bearer token must never be sent over plaintext HTTP + // or to a URL carrying userinfo/alternate ports. Host matching alone is + // insufficient because an imported endpoint can still select those URL + // forms. + if parsed.scheme() != "https" + || !parsed.username().is_empty() + || parsed.password().is_some() + || parsed.port().is_some() + { + return false; + } + if parsed.query().is_some() || parsed.fragment().is_some() { + return false; + } let Some(host) = parsed .host_str() .map(|value| value.trim().to_ascii_lowercase()) @@ -22,8 +56,9 @@ pub fn looks_like_vertex_ai_host(base_url: &str) -> bool { }; host == VERTEX_AI_HOST - || host.ends_with(&format!(".{VERTEX_AI_HOST}")) - || host.ends_with(&format!("-{VERTEX_AI_HOST}")) + || host + .strip_suffix(&format!("-{VERTEX_AI_HOST}")) + .is_some_and(is_valid_vertex_region) } pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool { @@ -101,8 +136,9 @@ fn looks_like_vertex_openai_compat_base(base_url: &str) -> bool { #[cfg(test)] mod tests { use super::{ - is_vertex_api_key_transport_context, is_vertex_service_account_transport_context, - is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth, + is_valid_vertex_region, is_vertex_api_key_transport_context, + is_vertex_service_account_transport_context, is_vertex_transport_context, + looks_like_vertex_ai_host, uses_vertex_api_key_query_auth, }; use crate::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, @@ -174,9 +210,41 @@ mod tests { assert!(looks_like_vertex_ai_host( "https://us-central1-aiplatform.googleapis.com" )); + assert!(!looks_like_vertex_ai_host( + "https://foo.bar-aiplatform.googleapis.com" + )); + assert!(!looks_like_vertex_ai_host( + "https://us-central1-aiplatform.googleapis.com?token=secret" + )); + assert!(!looks_like_vertex_ai_host( + "https://us-central1-aiplatform.googleapis.com." + )); + assert!(!looks_like_vertex_ai_host( + "http://us-central1-aiplatform.googleapis.com" + )); + assert!(!looks_like_vertex_ai_host( + "https://user@us-central1-aiplatform.googleapis.com" + )); + assert!(!looks_like_vertex_ai_host( + "https://us-central1-aiplatform.googleapis.com:8443" + )); assert!(!looks_like_vertex_ai_host("https://example.com")); } + #[test] + fn rejects_vertex_region_url_syntax() { + for value in [ + "attacker.example/", + "us-central1?x=1", + "us.central1", + "-bad", + ] { + assert!(!is_valid_vertex_region(value)); + } + assert!(is_valid_vertex_region("us-central1")); + assert!(is_valid_vertex_region("global")); + } + #[test] fn infers_vertex_api_key_context_for_custom_aiplatform_transport() { assert!(is_vertex_api_key_transport_context(&sample_transport())); diff --git a/crates/aether-provider/transport/src/vertex/mod.rs b/crates/aether-provider/transport/src/vertex/mod.rs index bd776e08c..ebf6b3b88 100644 --- a/crates/aether-provider/transport/src/vertex/mod.rs +++ b/crates/aether-provider/transport/src/vertex/mod.rs @@ -11,8 +11,9 @@ pub use auth::{ VERTEX_SERVICE_ACCOUNT_AUTH_HEADER, VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE, }; pub use context::{ - is_vertex_api_key_transport_context, is_vertex_service_account_transport_context, - is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth, + is_valid_vertex_region, is_vertex_api_key_transport_context, + is_vertex_service_account_transport_context, is_vertex_transport_context, + looks_like_vertex_ai_host, uses_vertex_api_key_query_auth, }; pub use policy::{ local_vertex_api_key_gemini_transport_unsupported_reason_with_network, diff --git a/crates/aether-provider/transport/src/vertex/url.rs b/crates/aether-provider/transport/src/vertex/url.rs index ccc6f3d1c..570c228dc 100644 --- a/crates/aether-provider/transport/src/vertex/url.rs +++ b/crates/aether-provider/transport/src/vertex/url.rs @@ -2,8 +2,9 @@ use std::collections::BTreeMap; use url::form_urlencoded; -use super::super::url::build_passthrough_path_url; +use super::super::url::{build_passthrough_path_url, encode_url_path_segment}; use super::auth::VertexServiceAccountAuthConfig; +use super::context::is_valid_vertex_region; pub const VERTEX_API_KEY_BASE_URL: &str = "https://aiplatform.googleapis.com"; @@ -85,7 +86,8 @@ fn build_vertex_api_key_google_model_url( return None; } - let path = format!("/v1/publishers/google/models/{trimmed_model}:{trimmed_action}"); + let encoded_model = encode_url_path_segment(trimmed_model); + let path = format!("/v1/publishers/google/models/{encoded_model}:{trimmed_action}"); let merged_query = build_vertex_api_key_query(trimmed_api_key, request_query, stream); build_passthrough_path_url(VERTEX_API_KEY_BASE_URL, &path, merged_query.as_deref(), &[]) } @@ -110,8 +112,14 @@ fn build_vertex_service_account_google_model_url( } else { format!("https://{region}-aiplatform.googleapis.com") }; + // URL parsers normalize dot-only segments even when the dots are percent-encoded. + if matches!(project_id, "." | "..") { + return None; + } + let encoded_project_id = encode_url_path_segment(project_id); + let encoded_model = encode_url_path_segment(trimmed_model); let path = format!( - "/v1/projects/{project_id}/locations/{region}/publishers/google/models/{trimmed_model}:{trimmed_action}" + "/v1/projects/{encoded_project_id}/locations/{region}/publishers/google/models/{encoded_model}:{trimmed_action}" ); let merged_query = build_vertex_service_account_query(request_query, stream); build_passthrough_path_url(&base_url, &path, merged_query.as_deref(), &[]) @@ -127,7 +135,7 @@ pub fn resolve_vertex_service_account_region( .get(trimmed_model) .map(String::as_str) .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_valid_vertex_region(value)) { return region.to_string(); } @@ -138,7 +146,7 @@ pub fn resolve_vertex_service_account_region( .region .as_deref() .map(str::trim) - .filter(|value| !value.is_empty()) + .filter(|value| is_valid_vertex_region(value)) { return region.to_string(); } @@ -236,7 +244,7 @@ mod tests { use super::{ build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url, - build_vertex_service_account_gemini_content_url, + build_vertex_service_account_gemini_content_url, resolve_vertex_service_account_region, }; use crate::vertex::VertexServiceAccountAuthConfig; @@ -324,4 +332,80 @@ mod tests { ) ); } + + #[test] + fn service_account_region_override_cannot_escape_vertex_origin() { + let auth_config = VertexServiceAccountAuthConfig { + client_email: "svc@example.iam.gserviceaccount.com".to_string(), + private_key: "not-used".to_string(), + project_id: "demo-project".to_string(), + token_uri: "https://oauth2.googleapis.com/token".to_string(), + region: Some("attacker.example/".to_string()), + model_regions: BTreeMap::from([( + "custom-model".to_string(), + "attacker.example/".to_string(), + )]), + }; + + assert_eq!( + resolve_vertex_service_account_region("custom-model", &auth_config), + "global" + ); + let url = build_vertex_service_account_gemini_content_url( + "custom-model", + false, + &auth_config, + None, + ) + .expect("service account URL should be built"); + assert!(url.starts_with("https://aiplatform.googleapis.com/")); + assert!(!url.contains("attacker.example")); + } + + #[test] + fn vertex_resource_components_cannot_rewrite_the_request_path() { + let auth_config = VertexServiceAccountAuthConfig { + client_email: "svc@example.iam.gserviceaccount.com".to_string(), + private_key: "not-used".to_string(), + project_id: "project/../victim?key=attacker".to_string(), + token_uri: "https://oauth2.googleapis.com/token".to_string(), + region: None, + model_regions: BTreeMap::new(), + }; + + assert_eq!( + build_vertex_service_account_gemini_content_url( + "model/../../admin#fragment", + false, + &auth_config, + None, + ) + .as_deref(), + Some( + "https://aiplatform.googleapis.com/v1/projects/project%2F..%2Fvictim%3Fkey=attacker/locations/global/publishers/google/models/model%2F..%2F..%2Fadmin%23fragment:generateContent" + ) + ); + } + + #[test] + fn vertex_rejects_dot_only_project_path_segments() { + let auth_config = VertexServiceAccountAuthConfig { + client_email: "svc@example.iam.gserviceaccount.com".to_string(), + private_key: "not-used".to_string(), + project_id: "..".to_string(), + token_uri: "https://oauth2.googleapis.com/token".to_string(), + region: None, + model_regions: BTreeMap::new(), + }; + + assert_eq!( + build_vertex_service_account_gemini_content_url( + "gemini-2.5-pro", + false, + &auth_config, + None, + ), + None + ); + } } diff --git a/crates/aether-provider/transport/src/video/mod.rs b/crates/aether-provider/transport/src/video/mod.rs index a713fe856..85dfa0cd9 100644 --- a/crates/aether-provider/transport/src/video/mod.rs +++ b/crates/aether-provider/transport/src/video/mod.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt; use aether_data_contracts::repository::video_tasks::StoredVideoTask; use aether_video_tasks_core::{ @@ -29,7 +30,7 @@ pub enum ProviderVideoCreateFamily { Gemini, } -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct ProviderVideoCreateHeadersInput<'a> { pub headers: &'a http::HeaderMap, pub auth_header: &'a str, @@ -39,6 +40,37 @@ pub struct ProviderVideoCreateHeadersInput<'a> { pub original_request_body: &'a Value, } +impl fmt::Debug for ProviderVideoCreateHeadersInput<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ProviderVideoCreateHeadersInput") + .field( + "request_header_names", + &self + .headers + .keys() + .map(|name| name.as_str()) + .collect::>(), + ) + .field("auth_header", &self.auth_header) + .field("has_auth_value", &(!self.auth_value.is_empty())) + .field("has_header_rules", &self.header_rules.is_some()) + .field( + "provider_request_body_bytes", + &serde_json::to_vec(self.provider_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .field( + "original_request_body_bytes", + &serde_json::to_vec(self.original_request_body) + .ok() + .map(|bytes| bytes.len()), + ) + .finish() + } +} + #[async_trait] pub trait VideoTaskTransportSnapshotLookup: Send + Sync { async fn read_video_task_provider_transport_snapshot( @@ -210,6 +242,12 @@ pub fn build_video_create_headers( ) { return None; } + let declared_connection_headers = + super::headers::declared_connection_header_names(input.headers, &BTreeMap::new()); + super::headers::remove_declared_connection_headers( + &mut provider_request_headers, + &declared_connection_headers, + ); Some(provider_request_headers) } diff --git a/crates/aether-provider/transport/src/windsurf.rs b/crates/aether-provider/transport/src/windsurf.rs index 7a4364a67..89ca15909 100644 --- a/crates/aether-provider/transport/src/windsurf.rs +++ b/crates/aether-provider/transport/src/windsurf.rs @@ -3,6 +3,7 @@ use std::collections::BTreeMap; use serde_json::{json, Value}; use uuid::Uuid; +use super::headers::{declared_connection_header_names, remove_declared_connection_headers}; use crate::rules::{ apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers, body_rules_are_locally_supported, header_rules_are_locally_supported, @@ -190,13 +191,16 @@ pub fn build_windsurf_cascade_headers( auth_value: &str, _upstream_is_stream: bool, ) -> Option> { + let declared_connection_headers = declared_connection_header_names(headers, &BTreeMap::new()); let mut out = BTreeMap::new(); for (name, value) in headers { let Ok(value) = value.to_str() else { continue; }; let key = name.as_str().to_ascii_lowercase(); - if should_skip_upstream_passthrough_header(&key) { + if should_skip_upstream_passthrough_header(&key) + || declared_connection_headers.contains(&key) + { continue; } let value = value.trim(); @@ -234,6 +238,7 @@ pub fn build_windsurf_cascade_headers( if !auth_header.is_empty() { out.insert(auth_header, auth_value.trim().to_string()); } + remove_declared_connection_headers(&mut out, &declared_connection_headers); out.remove("content-length"); Some(out) } diff --git a/crates/aether-provider/transport/src/windsurf/cascade.rs b/crates/aether-provider/transport/src/windsurf/cascade.rs index f36c68f25..79b94458f 100644 --- a/crates/aether-provider/transport/src/windsurf/cascade.rs +++ b/crates/aether-provider/transport/src/windsurf/cascade.rs @@ -23,7 +23,7 @@ OUTPUT RULES: Violating these rules will produce broken output for the end user. Stay in chat-API mode at all times."#; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct CascadeStep { pub step_type: u64, pub status: u64, @@ -36,12 +36,44 @@ pub struct CascadeStep { pub usage: Option, } -#[derive(Debug, Clone, PartialEq, Eq)] +impl fmt::Debug for CascadeStep { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CascadeStep") + .field("step_type", &self.step_type) + .field("status", &self.status) + .field("text_len", &self.text.len()) + .field("response_text_len", &self.response_text.len()) + .field("modified_text_len", &self.modified_text.len()) + .field("thinking_len", &self.thinking.len()) + .field("error_text_len", &self.error_text.len()) + .field("has_native_tool", &self.native_tool.is_some()) + .field("usage", &self.usage) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct CascadeNativeToolStep { pub kind: String, pub arguments: Value, } +impl fmt::Debug for CascadeNativeToolStep { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CascadeNativeToolStep") + .field("kind", &self.kind) + .field( + "arguments_bytes", + &serde_json::to_vec(&self.arguments) + .ok() + .map(|bytes| bytes.len()), + ) + .finish() + } +} + #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] pub struct CascadeUsage { pub input_tokens: u64, @@ -60,13 +92,23 @@ impl CascadeUsage { } } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub struct CascadeImage { pub base64_data: String, pub mime_type: String, } -#[derive(Debug, Clone, Default, PartialEq, Eq)] +impl fmt::Debug for CascadeImage { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CascadeImage") + .field("base64_data_len", &self.base64_data.len()) + .field("mime_type", &self.mime_type) + .finish() + } +} + +#[derive(Clone, Default, PartialEq, Eq)] pub struct SendCascadeMessageOptions { pub tool_preamble: Option, pub images: Vec, @@ -75,6 +117,34 @@ pub struct SendCascadeMessageOptions { pub native_allowlist: Vec, } +impl fmt::Debug for SendCascadeMessageOptions { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SendCascadeMessageOptions") + .field( + "tool_preamble_len", + &self.tool_preamble.as_ref().map(String::len), + ) + .field("image_count", &self.images.len()) + .field( + "image_bytes", + &self + .images + .iter() + .map(|image| image.base64_data.len()) + .sum::(), + ) + .field("additional_steps_count", &self.additional_steps.len()) + .field( + "additional_steps_bytes", + &self.additional_steps.iter().map(Vec::len).sum::(), + ) + .field("native_mode", &self.native_mode) + .field("native_allowlist_count", &self.native_allowlist.len()) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct CascadeBuildError { message: String, diff --git a/crates/aether-provider/transport/src/windsurf/proto.rs b/crates/aether-provider/transport/src/windsurf/proto.rs index c772c9a1e..19f5f609f 100644 --- a/crates/aether-provider/transport/src/windsurf/proto.rs +++ b/crates/aether-provider/transport/src/windsurf/proto.rs @@ -8,19 +8,42 @@ pub enum WireType { Fixed32 = 5, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub enum FieldValue { Varint(u64), Bytes(Vec), } -#[derive(Debug, Clone, PartialEq, Eq)] +impl fmt::Debug for FieldValue { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Varint(value) => formatter.debug_tuple("Varint").field(value).finish(), + Self::Bytes(bytes) => formatter + .debug_struct("Bytes") + .field("len", &bytes.len()) + .finish(), + } + } +} + +#[derive(Clone, PartialEq, Eq)] pub struct Field { pub number: u32, pub wire_type: WireType, pub value: FieldValue, } +impl fmt::Debug for Field { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("Field") + .field("number", &self.number) + .field("wire_type", &self.wire_type) + .field("value", &self.value) + .finish() + } +} + impl Field { pub fn bytes(&self) -> &[u8] { match &self.value { @@ -138,7 +161,7 @@ pub fn parse_fields(buf: &[u8]) -> Result, ProtoError> { other => { return Err(ProtoError::new(format!( "unknown wire type {other} at offset {pos}" - ))) + ))); } }; let value = match wire_type { diff --git a/crates/aether-routing-core/src/actions.rs b/crates/aether-routing-core/src/actions.rs index 03f5b6995..71bb09f80 100644 --- a/crates/aether-routing-core/src/actions.rs +++ b/crates/aether-routing-core/src/actions.rs @@ -73,6 +73,8 @@ pub enum RoutingAction { priority_mode: Option, scheduling_mode: Option, keep_priority_on_conversion: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + sticky_key_attempts: Option, }, SetProviderPriority { provider_id: String, @@ -81,6 +83,10 @@ pub enum RoutingAction { SetKeyPriority { key_id: String, priority: i32, + /// When set, the override only applies to candidates served through + /// this API format; otherwise it applies to the key on every format. + #[serde(default, skip_serializing_if = "Option::is_none")] + api_format: Option, }, JsonPatchBody { patch: Vec, diff --git a/crates/aether-routing-core/src/lib.rs b/crates/aether-routing-core/src/lib.rs index 740a4dac9..a9a97742c 100644 --- a/crates/aether-routing-core/src/lib.rs +++ b/crates/aether-routing-core/src/lib.rs @@ -13,9 +13,9 @@ pub use actions::{ }; pub use conditions::{RoutingCondition, RoutingConditionContext, RoutingConditionOp}; pub use model::{ - RoutingGroupBinding, RoutingGroupBindingSubject, RoutingGroupConfig, RoutingGroupRecord, - RoutingGroupVersionRecord, RoutingModelPolicy, RoutingPoolPolicyOverride, RoutingRule, - RoutingSchedulingPreset, + RoutingDefaultPolicy, RoutingExecutionPolicy, RoutingGroupBinding, RoutingGroupBindingSubject, + RoutingGroupConfig, RoutingGroupRecord, RoutingGroupVersionRecord, RoutingModelPolicy, + RoutingPoolPolicyOverride, RoutingRule, RoutingSchedulingPreset, DEFAULT_STICKY_KEY_ATTEMPTS, }; pub use mutations::{ apply_json_patch_operations, validate_header_patch, validate_json_patch_operations, diff --git a/crates/aether-routing-core/src/model.rs b/crates/aether-routing-core/src/model.rs index 735c6d8e5..823c3e314 100644 --- a/crates/aether-routing-core/src/model.rs +++ b/crates/aether-routing-core/src/model.rs @@ -23,7 +23,52 @@ pub struct RoutingPoolPolicyOverride { pub scheduling_presets: Vec, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +/// Default number of attempts on the first-ranked (sticky) candidate before +/// failing over: one retry on the same key. +pub const DEFAULT_STICKY_KEY_ATTEMPTS: u32 = 2; + +/// Request-independent execution behaviours selected by a routing strategy. +/// +/// These flags deliberately live beside scheduling rather than in provider +/// transport configuration. A resolved policy is snapshotted for the request +/// and can therefore be consumed by execution without rereading mutable +/// system settings. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)] +pub struct RoutingExecutionPolicy { + #[serde(default, skip_serializing_if = "is_false")] + pub enable_cf_heartbeat: bool, + #[serde(default, skip_serializing_if = "is_false")] + pub cyber_continue_failover: bool, +} + +impl<'de> Deserialize<'de> for RoutingExecutionPolicy { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + #[derive(Deserialize, Default)] + struct LegacyCompatibleExecutionPolicy { + #[serde(default)] + enable_cf_heartbeat: bool, + #[serde(default)] + enable_openai_image_sync_heartbeat: bool, + #[serde(default)] + enable_standard_text_sync_heartbeat: bool, + #[serde(default)] + cyber_continue_failover: bool, + } + + let value = LegacyCompatibleExecutionPolicy::deserialize(deserializer)?; + Ok(Self { + enable_cf_heartbeat: value.enable_cf_heartbeat + || value.enable_openai_image_sync_heartbeat + || value.enable_standard_text_sync_heartbeat, + cyber_continue_failover: value.cyber_continue_failover, + }) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct RoutingDefaultPolicy { #[serde(default)] pub priority_mode: RoutingSetPriorityMode, @@ -31,6 +76,35 @@ pub struct RoutingDefaultPolicy { pub scheduling_mode: RoutingSchedulingMode, #[serde(default)] pub keep_priority_on_conversion: bool, + /// Total attempts on the first-ranked candidate before moving on. Later + /// candidates always get a single attempt so failover keeps advancing. + /// `0` and `1` both mean no same-key retry. + #[serde(default = "default_sticky_key_attempts")] + pub sticky_key_attempts: u32, + /// Strategy-scoped execution behaviour. Flattened for a stable JSON + /// shape and backwards-compatible migration from system settings. + #[serde(flatten)] + pub execution_policy: RoutingExecutionPolicy, +} + +impl Default for RoutingDefaultPolicy { + fn default() -> Self { + Self { + priority_mode: RoutingSetPriorityMode::default(), + scheduling_mode: RoutingSchedulingMode::default(), + keep_priority_on_conversion: false, + sticky_key_attempts: DEFAULT_STICKY_KEY_ATTEMPTS, + execution_policy: RoutingExecutionPolicy::default(), + } + } +} + +fn default_sticky_key_attempts() -> u32 { + DEFAULT_STICKY_KEY_ATTEMPTS +} + +fn is_false(value: &bool) -> bool { + !*value } #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] @@ -44,6 +118,13 @@ pub struct RoutingModelPolicy { pub provider_priority_overrides: BTreeMap, #[serde(default)] pub key_priority_overrides: BTreeMap, + /// Key priority overrides scoped to one API format: `api_format -> key_id -> priority`. + /// + /// A key can serve several API formats and legacy `global_priority_by_format` + /// ranks it independently per format. Entries here take precedence over + /// `key_priority_overrides` when the candidate format matches. + #[serde(default)] + pub key_priority_overrides_by_format: BTreeMap>, #[serde(default)] pub pool_priority_overrides: BTreeMap, #[serde(default)] @@ -69,8 +150,8 @@ pub struct RoutingRule { #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] pub struct RoutingGroupConfig { - #[serde(default)] - pub allowed_models: Vec, + /// The default policy is global for the selected strategy group. Model + /// differences are expressed through `model_policies` and `rules`. #[serde(default)] pub default_policy: RoutingDefaultPolicy, #[serde(default)] diff --git a/crates/aether-routing-core/src/policy.rs b/crates/aether-routing-core/src/policy.rs index 3e5da3b6f..710b0bbbe 100644 --- a/crates/aether-routing-core/src/policy.rs +++ b/crates/aether-routing-core/src/policy.rs @@ -8,7 +8,9 @@ use crate::actions::{ RoutingAction, RoutingRulePhase, RoutingSchedulingMode, RoutingSetPriorityMode, }; use crate::conditions::RoutingConditionContext; -use crate::model::{RoutingGroupConfig, RoutingModelPolicy, RoutingPoolPolicyOverride}; +use crate::model::{ + RoutingExecutionPolicy, RoutingGroupConfig, RoutingModelPolicy, RoutingPoolPolicyOverride, +}; use crate::mutations::{validate_header_patch, validate_json_patch_operations, MutationPlan}; use crate::ranking::RankingOverlay; use crate::validation::validate_routing_group_config; @@ -17,7 +19,7 @@ use crate::validation::validate_routing_group_config; pub enum RoutingPolicyError { #[error("routing group config is invalid: {0}")] InvalidConfig(String), - #[error("model is not allowed by routing group: {0}")] + #[error("model is not allowed by routing rule: {0}")] ModelNotAllowed(String), #[error("mutation action is invalid: {0}")] InvalidMutation(String), @@ -57,6 +59,11 @@ pub struct ResolvedRoutingPolicy { pub priority_mode: RoutingSetPriorityMode, pub scheduling_mode: RoutingSchedulingMode, pub keep_priority_on_conversion: bool, + /// See `RoutingDefaultPolicy::sticky_key_attempts`. + #[serde(default = "default_sticky_key_attempts")] + pub sticky_key_attempts: u32, + #[serde(flatten)] + pub execution_policy: RoutingExecutionPolicy, pub ranking_overlay: RankingOverlay, pub mutation_plan: MutationPlan, #[serde(default)] @@ -72,14 +79,6 @@ pub fn resolve_routing_policy( validate_routing_group_config(config) .map_err(|error| RoutingPolicyError::InvalidConfig(error.to_string()))?; - if !model_allowed(&config.allowed_models, input.requested_model) - && !model_allowed(&config.allowed_models, input.resolved_model) - { - return Err(RoutingPolicyError::ModelNotAllowed( - input.requested_model.to_string(), - )); - } - let mut policy = ResolvedRoutingPolicy { group_id: input.group_id.map(str::to_string), group_version: input.group_version, @@ -89,6 +88,8 @@ pub fn resolve_routing_policy( priority_mode: config.default_policy.priority_mode, scheduling_mode: config.default_policy.scheduling_mode, keep_priority_on_conversion: config.default_policy.keep_priority_on_conversion, + sticky_key_attempts: config.default_policy.sticky_key_attempts, + execution_policy: config.default_policy.execution_policy, ranking_overlay: RankingOverlay::default(), mutation_plan: MutationPlan::default(), pool_policy_overrides: BTreeMap::new(), @@ -163,6 +164,13 @@ fn apply_model_policy(policy: &mut ResolvedRoutingPolicy, model_policy: &Routing .iter() .map(|(key, value)| (key.clone(), *value)), ); + for (api_format, overrides) in &model_policy.key_priority_overrides_by_format { + for (key_id, priority) in overrides { + policy + .ranking_overlay + .insert_key_priority_override_for_format(api_format, key_id.clone(), *priority); + } + } policy.ranking_overlay.pool_priority_overrides.extend( model_policy .pool_priority_overrides @@ -198,6 +206,7 @@ fn apply_action( priority_mode, scheduling_mode, keep_priority_on_conversion, + sticky_key_attempts, } => { if let Some(priority_mode) = priority_mode { policy.priority_mode = *priority_mode; @@ -208,6 +217,9 @@ fn apply_action( if let Some(keep_priority_on_conversion) = keep_priority_on_conversion { policy.keep_priority_on_conversion = *keep_priority_on_conversion; } + if let Some(sticky_key_attempts) = sticky_key_attempts { + policy.sticky_key_attempts = *sticky_key_attempts; + } } RoutingAction::SetProviderPriority { provider_id, @@ -218,12 +230,27 @@ fn apply_action( .provider_priority_overrides .insert(provider_id.clone(), *priority); } - RoutingAction::SetKeyPriority { key_id, priority } => { - policy - .ranking_overlay - .key_priority_overrides - .insert(key_id.clone(), *priority); - } + RoutingAction::SetKeyPriority { + key_id, + priority, + api_format, + } => match api_format + .as_deref() + .map(str::trim) + .filter(|f| !f.is_empty()) + { + Some(api_format) => { + policy + .ranking_overlay + .insert_key_priority_override_for_format(api_format, key_id.clone(), *priority); + } + None => { + policy + .ranking_overlay + .key_priority_overrides + .insert(key_id.clone(), *priority); + } + }, RoutingAction::JsonPatchBody { patch } => { validate_json_patch_operations(patch) .map_err(|error| RoutingPolicyError::InvalidMutation(error.to_string()))?; @@ -260,6 +287,10 @@ fn model_allowed(patterns: &[String], requested_model: &str) -> bool { .any(|pattern| model_pattern_matches(pattern, requested_model)) } +fn default_sticky_key_attempts() -> u32 { + crate::model::DEFAULT_STICKY_KEY_ATTEMPTS +} + fn model_pattern_matches(pattern: &str, value: &str) -> bool { let pattern = pattern.trim(); if pattern == "*" { @@ -288,7 +319,6 @@ mod tests { #[test] fn resolves_model_policy_and_matching_rule() { let config = RoutingGroupConfig { - allowed_models: vec!["gpt-*".to_string()], default_policy: RoutingDefaultPolicy::default(), model_policies: vec![RoutingModelPolicy { model: "gpt-5".to_string(), @@ -355,13 +385,14 @@ mod tests { } #[test] - fn empty_allowlist_keeps_default_policy_for_models_without_an_override() { + fn default_policy_applies_to_models_without_an_override() { let config = RoutingGroupConfig { - allowed_models: vec![], default_policy: RoutingDefaultPolicy { priority_mode: RoutingSetPriorityMode::GlobalKey, scheduling_mode: RoutingSchedulingMode::LoadBalance, keep_priority_on_conversion: true, + sticky_key_attempts: 3, + execution_policy: Default::default(), }, model_policies: vec![RoutingModelPolicy { model: "special-model".to_string(), @@ -393,6 +424,7 @@ mod tests { assert_eq!(special.priority_mode, RoutingSetPriorityMode::GlobalKey); assert_eq!(special.scheduling_mode, RoutingSchedulingMode::LoadBalance); assert!(special.keep_priority_on_conversion); + assert_eq!(special.sticky_key_attempts, 3); assert_eq!( special.ranking_overlay.allowed_providers, vec!["provider-special"] @@ -426,6 +458,7 @@ mod tests { assert_eq!(ordinary.priority_mode, RoutingSetPriorityMode::GlobalKey); assert_eq!(ordinary.scheduling_mode, RoutingSchedulingMode::LoadBalance); assert!(ordinary.keep_priority_on_conversion); + assert_eq!(ordinary.sticky_key_attempts, 3); assert!(ordinary.ranking_overlay.allowed_providers.is_empty()); assert!(ordinary.ranking_overlay.allowed_keys.is_empty()); assert!(ordinary @@ -435,20 +468,26 @@ mod tests { } #[test] - fn rejects_disallowed_model() { - let config = RoutingGroupConfig { - allowed_models: vec!["gpt-5".to_string()], - ..RoutingGroupConfig::default() - }; + fn legacy_group_model_allowlist_is_ignored() { + let config: RoutingGroupConfig = serde_json::from_value(json!({ + "allowed_models": ["gpt-5"], + "default_policy": { + "priority_mode": "provider", + "scheduling_mode": "cache_affinity" + }, + "model_policies": [], + "rules": [] + })) + .expect("legacy routing config should remain readable"); - let err = resolve_routing_policy( + resolve_routing_policy( &config, RoutingPolicyInput { - group_id: None, - group_version: None, - selection_source: "test", - requested_model: "claude", - resolved_model: "claude", + group_id: Some("group-1"), + group_version: Some(1), + selection_source: "system_default", + requested_model: "claude-sonnet", + resolved_model: "claude-sonnet", api_format: "openai:chat", user_id: None, api_key_id: None, @@ -457,18 +496,82 @@ mod tests { phase: RoutingRulePhase::ClientRequest, }, ) - .unwrap_err(); + .expect("the legacy allowlist must not reject another model"); + } + #[test] + fn sticky_key_attempts_defaults_to_two_and_can_be_overridden_by_rule() { + let default_config = RoutingGroupConfig::default(); + let default_policy = resolve_routing_policy( + &default_config, + RoutingPolicyInput { + group_id: None, + group_version: None, + selection_source: "test", + requested_model: "gpt-5", + resolved_model: "gpt-5", + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: RoutingRulePhase::ClientRequest, + }, + ) + .expect("default config should resolve"); assert_eq!( - err, - RoutingPolicyError::ModelNotAllowed("claude".to_string()) + default_policy.sticky_key_attempts, + crate::DEFAULT_STICKY_KEY_ATTEMPTS ); + + let parsed: RoutingGroupConfig = + serde_json::from_value(json!({ "default_policy": { "priority_mode": "provider" } })) + .expect("legacy config without sticky_key_attempts should deserialize"); + assert_eq!( + parsed.default_policy.sticky_key_attempts, + crate::DEFAULT_STICKY_KEY_ATTEMPTS + ); + + let config = RoutingGroupConfig { + rules: vec![RoutingRule { + id: "no-sticky-retry".to_string(), + priority: 1, + enabled: true, + phase: RoutingRulePhase::ClientRequest, + conditions: RoutingCondition::default(), + actions: vec![RoutingAction::SetScheduling { + priority_mode: None, + scheduling_mode: None, + keep_priority_on_conversion: None, + sticky_key_attempts: Some(1), + }], + stop_processing: false, + }], + ..RoutingGroupConfig::default() + }; + let policy = resolve_routing_policy( + &config, + RoutingPolicyInput { + group_id: None, + group_version: None, + selection_source: "test", + requested_model: "gpt-5", + resolved_model: "gpt-5", + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: RoutingRulePhase::ClientRequest, + }, + ) + .expect("rule config should resolve"); + assert_eq!(policy.sticky_key_attempts, 1); } #[test] fn restrict_model_action_rejects_matching_request() { let config = RoutingGroupConfig { - allowed_models: vec!["*".to_string()], rules: vec![RoutingRule { id: "restrict".to_string(), priority: 1, diff --git a/crates/aether-routing-core/src/ranking.rs b/crates/aether-routing-core/src/ranking.rs index a8bc5b867..c12f52e07 100644 --- a/crates/aether-routing-core/src/ranking.rs +++ b/crates/aether-routing-core/src/ranking.rs @@ -21,6 +21,9 @@ pub struct RankingOverlay { pub provider_priority_overrides: BTreeMap, #[serde(default)] pub key_priority_overrides: BTreeMap, + /// `api_format -> key_id -> priority`; see `RoutingModelPolicy`. + #[serde(default)] + pub key_priority_overrides_by_format: BTreeMap>, #[serde(default)] pub pool_priority_overrides: BTreeMap, } @@ -40,6 +43,46 @@ impl RankingOverlay { .unwrap_or(fallback) } + /// Format-scoped key priority: a per-format override wins, then the + /// format-agnostic key override, then `fallback`. + pub fn key_priority_for_format(&self, key_id: &str, api_format: &str, fallback: i32) -> i32 { + self.key_priority_override_for_format(key_id, api_format) + .unwrap_or_else(|| self.key_priority(key_id, fallback)) + } + + /// Format-scoped key override using exact (case-insensitive) format match. + pub fn key_priority_override_for_format(&self, key_id: &str, api_format: &str) -> Option { + let api_format = api_format.trim(); + self.key_priority_override_matching_format(key_id, |format| { + format.trim().eq_ignore_ascii_case(api_format) + }) + } + + /// Format-scoped key override where the caller decides how configured + /// format names match the candidate format (for alias-aware matching). + pub fn key_priority_override_matching_format( + &self, + key_id: &str, + mut format_matches: impl FnMut(&str) -> bool, + ) -> Option { + self.key_priority_overrides_by_format + .iter() + .find(|(format, _)| format_matches(format)) + .and_then(|(_, overrides)| overrides.get(key_id).copied()) + } + + pub fn insert_key_priority_override_for_format( + &mut self, + api_format: &str, + key_id: String, + priority: i32, + ) { + self.key_priority_overrides_by_format + .entry(api_format.trim().to_ascii_lowercase()) + .or_default() + .insert(key_id, priority); + } + pub fn pool_priority(&self, provider_id: &str, fallback: i32) -> i32 { self.pool_priority_overrides .get(provider_id) @@ -89,6 +132,9 @@ pub struct RoutingCandidateFacts { pub model_id: String, #[serde(default, skip_serializing_if = "Option::is_none")] pub key_id: Option, + /// Candidate API format used to resolve format-scoped key overrides. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_format: Option, pub provider_priority: i32, pub key_priority: i32, } @@ -114,7 +160,12 @@ pub fn rank_vector_for_candidate( CandidateKind::Provider => facts .key_id .as_deref() - .map(|key_id| overlay.key_priority(key_id, facts.key_priority)) + .map(|key_id| match facts.api_format.as_deref() { + Some(api_format) => { + overlay.key_priority_for_format(key_id, api_format, facts.key_priority) + } + None => overlay.key_priority(key_id, facts.key_priority), + }) .unwrap_or(facts.key_priority), CandidateKind::PoolGroup => { overlay.pool_priority(&facts.provider_id, facts.key_priority) @@ -142,6 +193,7 @@ mod tests { endpoint_id: "endpoint-a".to_string(), model_id: "model-a".to_string(), key_id: Some("key-a".to_string()), + api_format: None, provider_priority: 10, key_priority: 20, }; @@ -151,6 +203,47 @@ mod tests { assert_eq!(vector.key_priority_after, 5); } + #[test] + fn format_scoped_key_override_wins_over_key_override_for_matching_format() { + let mut overlay = RankingOverlay { + key_priority_overrides: BTreeMap::from([("key-a".to_string(), 5)]), + ..RankingOverlay::default() + }; + overlay.insert_key_priority_override_for_format("openai:chat", "key-a".to_string(), 1); + + assert_eq!( + overlay.key_priority_for_format("key-a", "openai:chat", 20), + 1 + ); + assert_eq!( + overlay.key_priority_for_format("key-a", "OpenAI:Chat", 20), + 1 + ); + assert_eq!( + overlay.key_priority_for_format("key-a", "claude:messages", 20), + 5 + ); + assert_eq!( + overlay.key_priority_for_format("key-b", "openai:chat", 20), + 20 + ); + + let facts = RoutingCandidateFacts { + candidate_kind: CandidateKind::Provider, + provider_id: "provider-a".to_string(), + endpoint_id: "endpoint-a".to_string(), + model_id: "model-a".to_string(), + key_id: Some("key-a".to_string()), + api_format: Some("openai:chat".to_string()), + provider_priority: 10, + key_priority: 20, + }; + assert_eq!( + rank_vector_for_candidate(&overlay, &facts).key_priority_after, + 1 + ); + } + #[test] fn rank_vector_falls_back_to_existing_priorities() { let facts = RoutingCandidateFacts { @@ -159,6 +252,7 @@ mod tests { endpoint_id: "endpoint-a".to_string(), model_id: "model-a".to_string(), key_id: Some("key-a".to_string()), + api_format: None, provider_priority: 10, key_priority: 20, }; @@ -180,6 +274,7 @@ mod tests { endpoint_id: "endpoint-a".to_string(), model_id: "model-a".to_string(), key_id: None, + api_format: None, provider_priority: 10, key_priority: 20, }; diff --git a/crates/aether-routing-core/src/validation.rs b/crates/aether-routing-core/src/validation.rs index 0ac3123be..8314cc8c8 100644 --- a/crates/aether-routing-core/src/validation.rs +++ b/crates/aether-routing-core/src/validation.rs @@ -271,6 +271,7 @@ mod tests { priority_mode: None, scheduling_mode: None, keep_priority_on_conversion: Some(true), + sticky_key_attempts: None, }, "set_scheduling", ), @@ -285,6 +286,7 @@ mod tests { RoutingAction::SetKeyPriority { key_id: "key-1".to_string(), priority: 1, + api_format: None, }, "set_key_priority", ), diff --git a/crates/aether-runtime/base/Cargo.toml b/crates/aether-runtime/base/Cargo.toml index 6c9c72765..59d03f67f 100644 --- a/crates/aether-runtime/base/Cargo.toml +++ b/crates/aether-runtime/base/Cargo.toml @@ -11,6 +11,7 @@ async-stream.workspace = true axum = { version = "0.8" } chrono.workspace = true futures-util.workspace = true +libc = "0.2" serde_json.workspace = true sha2.workspace = true thiserror.workspace = true diff --git a/crates/aether-runtime/base/src/admission.rs b/crates/aether-runtime/base/src/admission.rs index d45ef129c..1bbaceedc 100644 --- a/crates/aether-runtime/base/src/admission.rs +++ b/crates/aether-runtime/base/src/admission.rs @@ -2,6 +2,7 @@ use async_stream::stream; use axum::body::Body; use axum::http::Response; use futures_util::StreamExt; +use std::sync::Arc; use std::time::Duration; use crate::concurrency::ConcurrencyPermit; @@ -10,24 +11,31 @@ const ADMISSION_HEALTH_POLL_INTERVAL: Duration = Duration::from_secs(1); pub trait AdmissionPermitHealth: Send + Sync { fn is_healthy(&self) -> bool; + + fn requires_health_poll(&self) -> bool { + true + } } impl AdmissionPermitHealth for ConcurrencyPermit { fn is_healthy(&self) -> bool { true } + + fn requires_health_poll(&self) -> bool { + false + } } +#[derive(Clone)] pub struct AdmissionPermit { - _local: Option, - _distributed: Option>, + _permits: Vec>, } impl std::fmt::Debug for AdmissionPermit { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("AdmissionPermit") - .field("has_local", &self._local.is_some()) - .field("has_distributed", &self._distributed.is_some()) + .field("permit_count", &self._permits.len()) .finish() } } @@ -37,34 +45,39 @@ impl AdmissionPermit { local: Option, distributed: Option, ) -> Option { - if local.is_none() && distributed.is_none() { - None - } else { - Some(Self { - _local: local, - _distributed: distributed - .map(|permit| Box::new(permit) as Box), - }) + let mut permits = Vec::>::new(); + if let Some(local) = local { + permits.push(Arc::new(local)); } + if let Some(distributed) = distributed { + permits.push(Arc::new(distributed)); + } + (!permits.is_empty()).then_some(Self { _permits: permits }) + } + + pub fn combine(permits: impl IntoIterator) -> Option { + let permits = permits + .into_iter() + .flat_map(|permit| permit._permits) + .collect::>(); + (!permits.is_empty()).then_some(Self { _permits: permits }) } pub fn is_healthy(&self) -> bool { - self._distributed - .as_ref() - .map(|permit| permit.is_healthy()) - .unwrap_or(true) + self._permits.iter().all(|permit| permit.is_healthy()) } fn requires_health_poll(&self) -> bool { - self._distributed.is_some() + self._permits + .iter() + .any(|permit| permit.requires_health_poll()) } } impl From for AdmissionPermit { fn from(value: ConcurrencyPermit) -> Self { Self { - _local: Some(value), - _distributed: None, + _permits: vec![Arc::new(value)], } } } @@ -210,6 +223,27 @@ mod tests { assert_eq!(gate.snapshot().in_flight, 0); } + #[test] + fn cloned_permit_releases_capacity_only_after_last_clone_drops() { + let gate = ConcurrencyGate::new("test", 1); + let permit = AdmissionPermit::from(gate.try_acquire().expect("first permit")); + let cloned = permit.clone(); + + drop(permit); + assert_eq!(gate.snapshot().in_flight, 1); + assert!( + gate.try_acquire().is_err(), + "capacity should remain held by the clone" + ); + + drop(cloned); + assert_eq!(gate.snapshot().in_flight, 0); + assert!( + gate.try_acquire().is_ok(), + "last clone should release capacity" + ); + } + #[tokio::test] async fn holds_combined_local_and_distributed_permit_until_future_finishes() { let local_gate = ConcurrencyGate::new("local", 1); diff --git a/crates/aether-runtime/base/src/tracing.rs b/crates/aether-runtime/base/src/tracing.rs index 051db9ef9..2a265ac0d 100644 --- a/crates/aether-runtime/base/src/tracing.rs +++ b/crates/aether-runtime/base/src/tracing.rs @@ -716,10 +716,68 @@ impl RollingFileSink { } fn open_bucketed_log_file(dir: &Path, service_name: &str, bucket: &str) -> io::Result { - OpenOptions::new() - .create(true) - .append(true) - .open(bucketed_log_path(dir, service_name, bucket)) + let path = bucketed_log_path(dir, service_name, bucket); + let mut options = OpenOptions::new(); + options.create(true).append(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt as _; + options + .mode(0o600) + .custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); + } + + let file = options.open(&path)?; + validate_open_log_file(&file, &path)?; + Ok(file) +} + +#[cfg(unix)] +fn validate_open_log_file(file: &File, path: &Path) -> io::Result<()> { + use std::os::unix::fs::{MetadataExt as _, PermissionsExt as _}; + + let metadata = file.metadata()?; + if !metadata.file_type().is_file() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + format!("log destination is not a regular file: {}", path.display()), + )); + } + if metadata.uid() != unsafe { libc::geteuid() } { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + format!( + "log destination is owned by another user: {}", + path.display() + ), + )); + } + if metadata.nlink() != 1 { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + format!( + "log destination has multiple hard links: {}", + path.display() + ), + )); + } + // `mode(0o600)` only affects newly-created files. Tighten an existing + // bucket as well so a historical permissive umask cannot keep exposing + // request and operational data after an upgrade. + file.set_permissions(fs::Permissions::from_mode(0o600))?; + Ok(()) +} + +#[cfg(not(unix))] +fn validate_open_log_file(file: &File, path: &Path) -> io::Result<()> { + let metadata = file.metadata()?; + if !metadata.file_type().is_file() { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + format!("log destination is not a regular file: {}", path.display()), + )); + } + Ok(()) } fn bucketed_log_path(dir: &Path, service_name: &str, bucket: &str) -> PathBuf { @@ -838,9 +896,9 @@ fn select_log_files_for_cleanup( mod tests { use super::{ bucketed_log_path, cleanup_log_files, format_target_cell, log_bucket_key, - select_log_files_for_cleanup, FileLoggingConfig, JsonRuntimeEventFormatter, - LogFileCandidate, LogRotation, PrettyRuntimeEventFormatter, RollingFileSink, - RuntimeLogIdentity, + open_bucketed_log_file, select_log_files_for_cleanup, FileLoggingConfig, + JsonRuntimeEventFormatter, LogFileCandidate, LogRotation, PrettyRuntimeEventFormatter, + RollingFileSink, RuntimeLogIdentity, }; use chrono::{Local, TimeZone}; use std::fs; @@ -962,6 +1020,43 @@ mod tests { fs::remove_dir_all(&dir).expect("temp dir should be removable"); } + #[cfg(unix)] + #[test] + fn rolling_log_file_is_private_and_rejects_symlink_destination() { + use std::os::unix::fs::{symlink, PermissionsExt as _}; + + let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4())); + fs::create_dir_all(&dir).expect("temp dir should exist"); + + let private = open_bucketed_log_file(&dir, "runtime-test", "private") + .expect("private log should open"); + assert_eq!( + private + .metadata() + .expect("log metadata") + .permissions() + .mode() + & 0o777, + 0o600 + ); + drop(private); + + let victim = dir.join("victim.txt"); + fs::write(&victim, b"unchanged").expect("victim should exist"); + let symlink_path = bucketed_log_path(&dir, "runtime-test", "symlink"); + symlink(&victim, &symlink_path).expect("symlink should exist"); + assert!( + open_bucketed_log_file(&dir, "runtime-test", "symlink").is_err(), + "rolling logs must not follow a pre-created symlink" + ); + assert_eq!( + fs::read(&victim).expect("victim should remain readable"), + b"unchanged" + ); + + fs::remove_dir_all(&dir).expect("temp dir should be removable"); + } + #[test] fn rolling_file_sink_treats_startup_cleanup_failure_as_non_fatal() { fn fail_cleanup(_: &str, _: &FileLoggingConfig) -> std::io::Result { diff --git a/crates/aether-runtime/state/src/lib.rs b/crates/aether-runtime/state/src/lib.rs index 9a2fef67e..22ea91650 100644 --- a/crates/aether-runtime/state/src/lib.rs +++ b/crates/aether-runtime/state/src/lib.rs @@ -139,6 +139,16 @@ impl RuntimeStateConfig { "runtime memory max_kv_entries must be positive".to_string(), )); } + if self.memory.max_usage_limit_windows == 0 { + return Err(DataLayerError::InvalidConfiguration( + "runtime memory max_usage_limit_windows must be positive".to_string(), + )); + } + if self.memory.max_usage_limit_events == 0 { + return Err(DataLayerError::InvalidConfiguration( + "runtime memory max_usage_limit_events must be positive".to_string(), + )); + } if matches!(self.command_timeout_ms, Some(0)) { return Err(DataLayerError::InvalidConfiguration( "runtime state command_timeout_ms must be positive".to_string(), @@ -366,6 +376,28 @@ impl RuntimeState { } } + /// Atomically creates an expiring key without replacing an existing value. + pub async fn kv_set_if_absent( + &self, + key: &str, + value: impl Into + Send, + ttl: Duration, + ) -> Result { + if ttl.is_zero() { + return Err(DataLayerError::InvalidInput( + "runtime kv set-if-absent ttl must be positive".to_string(), + )); + } + match self.backend.as_ref() { + RuntimeStateBackend::Memory(memory) => { + Ok(memory.kv_set_if_absent(key, value.into(), ttl).await) + } + RuntimeStateBackend::Redis(redis) => { + redis.runtime.kv_set_if_absent(key, value.into(), ttl).await + } + } + } + pub async fn kv_get(&self, key: &str) -> Result, DataLayerError> { match self.backend.as_ref() { RuntimeStateBackend::Memory(memory) => Ok(memory.kv_get(key).await), @@ -464,6 +496,42 @@ impl RuntimeState { } } + pub async fn check_and_consume_usage_limits( + &self, + input: UsageLimitInput<'_>, + ) -> Result { + if input.rules.is_empty() { + return Ok(UsageLimitCheck::Allowed); + } + validate_usage_limit_input(input)?; + match self.backend.as_ref() { + RuntimeStateBackend::Memory(memory) => { + memory.check_and_consume_usage_limits(input).await + } + RuntimeStateBackend::Redis(redis) => { + redis.runtime.check_and_consume_usage_limits(input).await + } + } + } + + /// Removes an idempotency event from every supplied usage-limit window. + /// + /// This is a compensation primitive for callers that compose the short-lived runtime + /// counters with a second durable admission store. It is intentionally idempotent. + pub async fn release_usage_limits( + &self, + input: UsageLimitReleaseInput<'_>, + ) -> Result<(), DataLayerError> { + if input.rules.is_empty() { + return Ok(()); + } + validate_usage_limit_release_input(input)?; + match self.backend.as_ref() { + RuntimeStateBackend::Memory(memory) => memory.release_usage_limits(input).await, + RuntimeStateBackend::Redis(redis) => redis.runtime.release_usage_limits(input).await, + } + } + pub async fn rate_limit_count(&self, key: &str, bucket: u64) -> Result { match self.backend.as_ref() { RuntimeStateBackend::Memory(memory) => memory.rate_limit_count(key, bucket), @@ -720,14 +788,18 @@ impl RuntimeState { RuntimeSemaphore::new(self.clone(), gate, None, limit, config) } - pub fn keyed_semaphore( + pub fn keyed_semaphore>( &self, gate: &'static str, - resource_key: &str, + resource_key: K, limit: usize, config: RuntimeSemaphoreConfig, ) -> Result { - let resource_key = resource_key.trim(); + let resource_key = resource_key.into(); + let prefix = format!("admission:{gate}:"); + let resource_key = resource_key + .strip_prefix(&prefix) + .unwrap_or(resource_key.as_str()); if resource_key.is_empty() { return Err(RuntimeSemaphoreError::InvalidConfiguration( "runtime semaphore resource key cannot be empty".to_string(), @@ -777,6 +849,197 @@ pub struct RateLimitInput<'a> { pub ttl_seconds: u64, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct UsageLimitRule<'a> { + /// A caller-defined key containing the same non-empty Redis hash tag as every sibling rule. + /// The key must change when the rule's window definition changes. + pub key: &'a str, + pub limit: u64, + /// Duration used to decide which events still count toward the limit. + pub window_seconds: u64, + /// Duration for retaining this rule's backing state after the current check. + pub retention_seconds: u64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct UsageLimitInput<'a> { + pub rules: &'a [UsageLimitRule<'a>], + pub event_id: &'a str, + pub now_unix_ms: u64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct UsageLimitReleaseInput<'a> { + pub rules: &'a [UsageLimitRule<'a>], + pub event_id: &'a str, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UsageLimitCheck { + Allowed, + Rejected { + rule_index: usize, + limit: u64, + retry_after: u64, + }, +} + +const MAX_REDIS_LUA_EXACT_INTEGER: u64 = (1_u64 << 53) - 1; + +fn validate_usage_limit_input(input: UsageLimitInput<'_>) -> Result<(), DataLayerError> { + if input.event_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "usage limit event_id must not be empty".to_string(), + )); + } + if input.now_unix_ms > MAX_REDIS_LUA_EXACT_INTEGER { + return Err(DataLayerError::InvalidInput( + "usage limit now_unix_ms exceeds the exact Redis Lua integer range".to_string(), + )); + } + + let mut expected_hash_tag = None; + for (index, rule) in input.rules.iter().enumerate() { + if rule.key.trim().is_empty() { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} key must not be empty" + ))); + } + if rule.limit == 0 { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} must have a positive limit" + ))); + } + if rule.limit > MAX_REDIS_LUA_EXACT_INTEGER { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} limit exceeds the exact Redis Lua integer range" + ))); + } + if rule.window_seconds == 0 { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} must have a positive window_seconds" + ))); + } + let window_ms = rule.window_seconds.checked_mul(1_000).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "usage limit rule {index} window_seconds is too large" + )) + })?; + if window_ms > MAX_REDIS_LUA_EXACT_INTEGER { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} window_seconds exceeds the exact Redis Lua integer range" + ))); + } + if rule.retention_seconds == 0 { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} must have a positive retention_seconds" + ))); + } + let retention_ms = rule.retention_seconds.checked_mul(1_000).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "usage limit rule {index} retention_seconds is too large" + )) + })?; + if retention_ms > MAX_REDIS_LUA_EXACT_INTEGER { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} retention_seconds exceeds the exact Redis Lua integer range" + ))); + } + if input + .now_unix_ms + .checked_add(window_ms) + .is_none_or(|expires_at| expires_at > MAX_REDIS_LUA_EXACT_INTEGER) + { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} window extends beyond the exact Redis Lua integer range" + ))); + } + if input + .now_unix_ms + .checked_add(retention_ms) + .is_none_or(|expires_at| expires_at > MAX_REDIS_LUA_EXACT_INTEGER) + { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} retention extends beyond the exact Redis Lua integer range" + ))); + } + if input.rules[..index] + .iter() + .any(|previous| previous.key == rule.key) + { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} duplicates key {}", + rule.key + ))); + } + + let hash_tag = redis_hash_tag(rule.key).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "usage limit rule {index} key must contain a non-empty Redis hash tag" + )) + })?; + if let Some(expected) = expected_hash_tag { + if hash_tag != expected { + return Err(DataLayerError::InvalidInput( + "usage limit rule keys must use the same Redis hash tag".to_string(), + )); + } + } else { + expected_hash_tag = Some(hash_tag); + } + } + Ok(()) +} + +fn validate_usage_limit_release_input( + input: UsageLimitReleaseInput<'_>, +) -> Result<(), DataLayerError> { + if input.event_id.trim().is_empty() { + return Err(DataLayerError::InvalidInput( + "usage limit event_id must not be empty".to_string(), + )); + } + let mut expected_hash_tag = None; + for (index, rule) in input.rules.iter().enumerate() { + if rule.key.trim().is_empty() { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} key must not be empty" + ))); + } + if input.rules[..index] + .iter() + .any(|previous| previous.key == rule.key) + { + return Err(DataLayerError::InvalidInput(format!( + "usage limit rule {index} duplicates key {}", + rule.key + ))); + } + let hash_tag = redis_hash_tag(rule.key).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "usage limit rule {index} key must contain a non-empty Redis hash tag" + )) + })?; + if let Some(expected) = expected_hash_tag { + if hash_tag != expected { + return Err(DataLayerError::InvalidInput( + "usage limit rule keys must use the same Redis hash tag".to_string(), + )); + } + } else { + expected_hash_tag = Some(hash_tag); + } + } + Ok(()) +} + +fn redis_hash_tag(key: &str) -> Option<&str> { + let tag_start = key.find('{')?.saturating_add(1); + let remainder = key.get(tag_start..)?; + let tag_end = remainder.find('}')?; + (tag_end > 0).then_some(&remainder[..tag_end]) +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct RuntimeQueueEntry { pub id: String, @@ -1083,6 +1346,12 @@ pub trait ExpiringKvStore: Send + Sync { value: String, ttl: Option, ) -> Result<(), DataLayerError>; + async fn set_if_absent( + &self, + key: &str, + value: String, + ttl: Duration, + ) -> Result; async fn get(&self, key: &str) -> Result, DataLayerError>; async fn get_many(&self, keys: &[String]) -> Result>, DataLayerError>; async fn take(&self, key: &str) -> Result, DataLayerError>; @@ -1101,6 +1370,15 @@ impl ExpiringKvStore for RuntimeState { self.kv_set(key, value, ttl).await } + async fn set_if_absent( + &self, + key: &str, + value: String, + ttl: Duration, + ) -> Result { + self.kv_set_if_absent(key, value, ttl).await + } + async fn get(&self, key: &str) -> Result, DataLayerError> { self.kv_get(key).await } @@ -1565,6 +1843,17 @@ mod tests { assert_eq!(runtime.kv_take("nonce").await.expect("take"), None); } + #[tokio::test] + async fn runtime_backends_share_atomic_kv_set_if_absent_contract() { + let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + assert_kv_set_if_absent_contract(&memory).await; + + let Some((_redis, redis_runtime)) = redis_runtime_for_test("kv-set-if-absent").await else { + return; + }; + assert_kv_set_if_absent_contract(&redis_runtime).await; + } + #[tokio::test] async fn memory_rate_limit_rejects_after_limit() { let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); @@ -1672,6 +1961,325 @@ mod tests { ); } + #[tokio::test] + async fn memory_usage_limits_check_all_windows_before_consuming() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let rules = [ + UsageLimitRule { + key: "usage:{user-1}:qps", + limit: 2, + window_seconds: 10, + retention_seconds: 10, + }, + UsageLimitRule { + key: "usage:{user-1}:weekly", + limit: 1, + window_seconds: 60, + retention_seconds: 60, + }, + ]; + + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-1", + now_unix_ms: 100_000, + }) + .await + .expect("first request"), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-2", + now_unix_ms: 105_000, + }) + .await + .expect("second request"), + UsageLimitCheck::Rejected { + rule_index: 1, + limit: 1, + retry_after: 55, + } + ); + + let qps_only = [rules[0]]; + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &qps_only, + event_id: "request-3", + now_unix_ms: 105_000, + }) + .await + .expect("weekly rejection must not consume qps"), + UsageLimitCheck::Allowed + ); + } + + #[tokio::test] + async fn memory_usage_limits_are_true_sliding_windows_and_idempotent() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let rules = [UsageLimitRule { + key: "usage:{user-1}:rolling-10", + limit: 2, + window_seconds: 10, + retention_seconds: 10, + }]; + let consume = |event_id, now_unix_ms| UsageLimitInput { + rules: &rules, + event_id, + now_unix_ms, + }; + + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-1", 100_000)) + .await + .unwrap(), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-2", 105_000)) + .await + .unwrap(), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-2", 106_000)) + .await + .unwrap(), + UsageLimitCheck::Allowed, + "replaying the same event must not consume twice" + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-3", 109_000)) + .await + .unwrap(), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 2, + retry_after: 1, + } + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-3", 110_000)) + .await + .unwrap(), + UsageLimitCheck::Allowed, + "the event at the exact rolling cutoff must expire" + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("request-4", 111_000)) + .await + .unwrap(), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 2, + retry_after: 4, + }, + "the first event's expiry must not reset when later events arrive" + ); + } + + #[tokio::test] + async fn usage_limit_input_requires_unique_co_located_rule_keys() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let different_tags = [ + UsageLimitRule { + key: "usage:{user-1}:one", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }, + UsageLimitRule { + key: "usage:{user-2}:two", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }, + ]; + assert!(matches!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &different_tags, + event_id: "request-1", + now_unix_ms: 100_000, + }) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + + let zero_retention = [UsageLimitRule { + retention_seconds: 0, + ..different_tags[0] + }]; + assert!(matches!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &zero_retention, + event_id: "request-1", + now_unix_ms: 100_000, + }) + .await, + Err(DataLayerError::InvalidInput(message)) + if message.contains("retention_seconds") + )); + + let duplicate_keys = [different_tags[0], different_tags[0]]; + assert!(matches!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &duplicate_keys, + event_id: "request-1", + now_unix_ms: 100_000, + }) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + } + + #[tokio::test] + async fn memory_usage_limits_apply_the_current_window_when_configuration_changes() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let long_window = [UsageLimitRule { + key: "usage:{user-1}:changing-window", + limit: 1, + window_seconds: 60, + retention_seconds: 60, + }]; + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &long_window, + event_id: "request-1", + now_unix_ms: 100_000, + }) + .await + .unwrap(), + UsageLimitCheck::Allowed + ); + + let short_window = [UsageLimitRule { + window_seconds: 10, + retention_seconds: 10, + ..long_window[0] + }]; + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &short_window, + event_id: "request-2", + now_unix_ms: 111_000, + }) + .await + .unwrap(), + UsageLimitCheck::Allowed, + "the old event is outside the newly configured shorter window" + ); + } + + #[tokio::test] + async fn memory_usage_limits_keep_subsecond_events_in_a_one_second_window() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let rules = [UsageLimitRule { + key: "usage:{user-1}:qps", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }]; + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-1", + now_unix_ms: 1_900, + }) + .await + .unwrap(), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-2", + now_unix_ms: 2_100, + }) + .await + .unwrap(), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 1, + } + ); + } + + #[tokio::test] + async fn memory_usage_limits_do_not_expire_epoch_events_before_the_window_elapses() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let rules = [UsageLimitRule { + key: "usage:{user-1}:epoch", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }]; + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-1", + now_unix_ms: 0, + }) + .await + .unwrap(), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "request-2", + now_unix_ms: 500, + }) + .await + .unwrap(), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 1, + } + ); + } + + #[test] + fn runtime_memory_usage_limit_capacity_must_be_positive() { + let mut config = RuntimeStateConfig::memory(); + config.memory.max_usage_limit_windows = 0; + assert!(matches!( + config.validate(), + Err(DataLayerError::InvalidConfiguration(message)) + if message.contains("max_usage_limit_windows") + )); + + let mut config = RuntimeStateConfig::memory(); + config.memory.max_usage_limit_events = 0; + assert!(matches!( + config.validate(), + Err(DataLayerError::InvalidConfiguration(message)) + if message.contains("max_usage_limit_events") + )); + } + #[tokio::test] async fn memory_rate_limit_concurrent_checks_do_not_exceed_limit() { let runtime = @@ -1768,6 +2376,50 @@ mod tests { assert_eq!(gate.snapshot().await.expect("snapshot").in_flight, 0); } + #[tokio::test] + async fn keyed_memory_semaphores_share_only_the_same_subject_key() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let first = runtime + .keyed_semaphore( + "plan_usage_concurrency", + "admission:plan_usage_concurrency:user-1", + 1, + RuntimeSemaphoreConfig::default(), + ) + .expect("first gate"); + let same_subject = runtime + .keyed_semaphore( + "plan_usage_concurrency", + "admission:plan_usage_concurrency:user-1", + 1, + RuntimeSemaphoreConfig::default(), + ) + .expect("same subject gate"); + let other_subject = runtime + .keyed_semaphore( + "plan_usage_concurrency", + "admission:plan_usage_concurrency:user-2", + 1, + RuntimeSemaphoreConfig::default(), + ) + .expect("other subject gate"); + + let permit = first.try_acquire().await.expect("first permit"); + assert!(matches!( + same_subject.try_acquire().await, + Err(RuntimeSemaphoreError::Saturated { limit: 1, .. }) + )); + assert!(other_subject.try_acquire().await.is_ok()); + drop(permit); + for _ in 0..20 { + if same_subject.snapshot().await.expect("snapshot").in_flight == 0 { + break; + } + tokio::task::yield_now().await; + } + assert!(same_subject.try_acquire().await.is_ok()); + } + #[tokio::test] async fn memory_keyed_semaphores_isolate_resource_capacity() { let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); @@ -2135,6 +2787,44 @@ mod tests { assert_invalid_shared_inputs(&redis_runtime).await; } + #[tokio::test] + async fn runtime_backends_share_atomic_sliding_usage_limit_contract() { + let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + assert_sliding_usage_limit_contract(&memory).await; + assert_concurrent_usage_limit_cap(&memory).await; + + let Some((_redis, redis_runtime)) = redis_runtime_for_test("usage-limits").await else { + return; + }; + assert_sliding_usage_limit_contract(&redis_runtime).await; + assert_concurrent_usage_limit_cap(&redis_runtime).await; + } + + #[tokio::test] + async fn runtime_backends_share_short_remaining_usage_limit_retention_contract() { + let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let redis = redis_runtime_for_test("usage-limit-retention").await; + let rules = [UsageLimitRule { + key: "usage:{retention-user}:period-bucket", + limit: 1, + window_seconds: 60, + retention_seconds: 1, + }]; + + assert_usage_limit_retention_seed(&memory, &rules).await; + if let Some((_, redis_runtime)) = &redis { + assert_usage_limit_retention_seed(redis_runtime, &rules).await; + } + + // Redis deliberately adds one second of expiry grace to the requested retention. + tokio::time::sleep(Duration::from_millis(2_200)).await; + + assert_usage_limit_retention_expired(&memory, &rules).await; + if let Some((_, redis_runtime)) = &redis { + assert_usage_limit_retention_expired(redis_runtime, &rules).await; + } + } + #[tokio::test] async fn memory_blocking_queue_read_does_not_block_kv_operations() { let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); @@ -2400,6 +3090,12 @@ mod tests { } async fn assert_invalid_shared_inputs(runtime: &RuntimeState) { + assert!(matches!( + runtime + .kv_set_if_absent("contract:invalid-ttl", "value", Duration::ZERO) + .await, + Err(DataLayerError::InvalidInput(_)) + )); assert!(matches!( runtime .score_set("contract:invalid-score", "nan", f64::NAN) @@ -2427,6 +3123,275 @@ mod tests { )); } + async fn assert_kv_set_if_absent_contract(runtime: &RuntimeState) { + let key = "contract:kv:set-if-absent"; + assert!(runtime + .kv_set_if_absent(key, "first", Duration::from_millis(30)) + .await + .expect("first set-if-absent should succeed")); + assert!(!runtime + .kv_set_if_absent(key, "second", Duration::from_secs(30)) + .await + .expect("duplicate set-if-absent should be rejected")); + assert_eq!( + runtime + .kv_get(key) + .await + .expect("existing value should be readable") + .as_deref(), + Some("first"), + "a rejected set-if-absent must not replace the existing value" + ); + + tokio::time::sleep(Duration::from_millis(80)).await; + assert!(runtime + .kv_set_if_absent(key, "after-expiry", Duration::from_secs(30)) + .await + .expect("expired key should be reusable")); + assert_eq!( + runtime + .kv_get(key) + .await + .expect("replacement value should be readable") + .as_deref(), + Some("after-expiry") + ); + + let concurrent_key = "contract:kv:set-if-absent:concurrent"; + let mut tasks = Vec::new(); + for index in 0..64 { + let runtime = runtime.clone(); + tasks.push(tokio::spawn(async move { + runtime + .kv_set_if_absent( + concurrent_key, + format!("candidate-{index}"), + Duration::from_secs(30), + ) + .await + .expect("concurrent set-if-absent should complete") + })); + } + let mut created = 0; + for task in tasks { + if task.await.expect("set-if-absent task should join") { + created += 1; + } + } + assert_eq!( + created, 1, + "exactly one concurrent caller may create the key" + ); + assert!(runtime + .kv_get(concurrent_key) + .await + .expect("winning value should be readable") + .is_some()); + } + + async fn assert_sliding_usage_limit_contract(runtime: &RuntimeState) { + let rules = [ + UsageLimitRule { + key: "usage:{shared-user}:short", + limit: 2, + window_seconds: 10, + retention_seconds: 10, + }, + UsageLimitRule { + key: "usage:{shared-user}:long", + limit: 1, + window_seconds: 60, + retention_seconds: 60, + }, + ]; + let consume = |event_id, now_unix_ms, rules| UsageLimitInput { + rules, + event_id, + now_unix_ms, + }; + + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-1", 100_000, &rules)) + .await + .expect("first event"), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-1", 101_000, &rules)) + .await + .expect("idempotent replay"), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-2", 105_000, &rules)) + .await + .expect("long window rejection"), + UsageLimitCheck::Rejected { + rule_index: 1, + limit: 1, + retry_after: 55, + } + ); + + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-3", 105_000, &rules[..1])) + .await + .expect("rejection must leave short rule untouched"), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-4", 109_000, &rules[..1])) + .await + .expect("short rule rejection"), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 2, + retry_after: 1, + } + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("event-4", 110_000, &rules[..1])) + .await + .expect("event at cutoff expires"), + UsageLimitCheck::Allowed + ); + + let qps = [UsageLimitRule { + key: "usage:{shared-user}:qps", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }]; + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("qps-1", 200_900, &qps)) + .await + .expect("first qps event"), + UsageLimitCheck::Allowed + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("qps-2", 201_100, &qps)) + .await + .expect("subsecond qps rejection"), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 1, + } + ); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("qps-2", 201_900, &qps)) + .await + .expect("qps event at cutoff"), + UsageLimitCheck::Allowed + ); + + runtime + .release_usage_limits(UsageLimitReleaseInput { + rules: &qps, + event_id: "qps-2", + }) + .await + .expect("release consumed qps event"); + assert_eq!( + runtime + .check_and_consume_usage_limits(consume("qps-3", 201_900, &qps)) + .await + .expect("released qps capacity should be reusable"), + UsageLimitCheck::Allowed + ); + runtime + .release_usage_limits(UsageLimitReleaseInput { + rules: &qps, + event_id: "qps-2", + }) + .await + .expect("release should be idempotent"); + } + + async fn assert_usage_limit_retention_seed( + runtime: &RuntimeState, + rules: &[UsageLimitRule<'_>], + ) { + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules, + event_id: "period-event-1", + now_unix_ms: 100_000, + }) + .await + .expect("seed short-retention period bucket"), + UsageLimitCheck::Allowed + ); + } + + async fn assert_usage_limit_retention_expired( + runtime: &RuntimeState, + rules: &[UsageLimitRule<'_>], + ) { + assert_eq!( + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules, + event_id: "period-event-2", + now_unix_ms: 102_200, + }) + .await + .expect("expired period bucket must be reusable"), + UsageLimitCheck::Allowed, + "retention must expire the bucket before its full counting window" + ); + } + + async fn assert_concurrent_usage_limit_cap(runtime: &RuntimeState) { + let mut tasks = Vec::new(); + for index in 0..64 { + let runtime = runtime.clone(); + tasks.push(tokio::spawn(async move { + let event_id = format!("parallel-event-{index}"); + let rules = [UsageLimitRule { + key: "usage:{shared-user}:parallel", + limit: 8, + window_seconds: 60, + retention_seconds: 60, + }]; + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: &event_id, + now_unix_ms: 300_000, + }) + .await + .expect("concurrent usage limit check") + })); + } + + let mut allowed = 0; + let mut rejected = 0; + for task in tasks { + match task.await.expect("concurrent usage limit task") { + UsageLimitCheck::Allowed => allowed += 1, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 8, + retry_after: 60, + } => rejected += 1, + unexpected => panic!("unexpected usage limit result: {unexpected:?}"), + } + } + assert_eq!(allowed, 8); + assert_eq!(rejected, 56); + } + async fn redis_runtime_for_test(prefix: &str) -> Option<(TestRedisServer, RuntimeState)> { let redis = TestRedisServer::start().await?; let runtime = RuntimeState::redis( diff --git a/crates/aether-runtime/state/src/memory.rs b/crates/aether-runtime/state/src/memory.rs index 6e0be8293..4fdddd6a7 100644 --- a/crates/aether-runtime/state/src/memory.rs +++ b/crates/aether-runtime/state/src/memory.rs @@ -6,20 +6,30 @@ use std::time::{Duration, Instant}; use tokio::sync::Mutex; +use crate::UsageLimitCheck; use crate::{DataLayerError, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats}; const MEMORY_RATE_LIMIT_COUNTER_SHARD_COUNT: usize = 64; const MEMORY_RATE_LIMIT_COUNTER_PRUNE_INTERVAL: u64 = 256; +const MEMORY_USAGE_LIMIT_PRUNE_INTERVAL: u64 = 256; +const DEFAULT_MAX_USAGE_LIMIT_WINDOWS: usize = 10_000; +const DEFAULT_MAX_USAGE_LIMIT_EVENTS: usize = 100_000; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct MemoryRuntimeStateConfig { pub max_kv_entries: usize, + /// Maximum number of active sliding-window keys retained by the memory backend. + pub max_usage_limit_windows: usize, + /// Maximum number of event identities retained across all usage-limit windows. + pub max_usage_limit_events: usize, } impl Default for MemoryRuntimeStateConfig { fn default() -> Self { Self { max_kv_entries: 10_000, + max_usage_limit_windows: DEFAULT_MAX_USAGE_LIMIT_WINDOWS, + max_usage_limit_events: DEFAULT_MAX_USAGE_LIMIT_EVENTS, } } } @@ -42,6 +52,7 @@ pub(crate) struct MemoryRuntimeBackend { config: MemoryRuntimeStateConfig, kv: Mutex>, counters: MemoryRateLimitCounters, + usage_limits: Mutex, sets: Mutex>, scores: Mutex>, queues: Mutex>, @@ -51,6 +62,111 @@ pub(crate) struct MemoryRuntimeBackend { semaphores: Mutex>>, } +#[derive(Debug, Clone)] +struct MemoryUsageLimitWindow { + window_ms: u64, + expires_at_unix_ms: u64, + events: HashMap, +} + +#[derive(Debug, Default)] +struct MemoryUsageLimitState { + windows: HashMap, + total_events: usize, + operations_since_prune: u64, + next_expiry_unix_ms: Option, +} + +impl MemoryUsageLimitState { + fn amortized_prune(&mut self, now_unix_ms: u64) { + self.operations_since_prune = self.operations_since_prune.saturating_add(1); + if self.operations_since_prune < MEMORY_USAGE_LIMIT_PRUNE_INTERVAL { + return; + } + self.operations_since_prune = 0; + if self + .next_expiry_unix_ms + .is_some_and(|expires_at| expires_at <= now_unix_ms) + { + self.prune_all(now_unix_ms); + } + } + + fn prune_all(&mut self, now_unix_ms: u64) { + self.operations_since_prune = 0; + let mut total_events = 0_usize; + let mut next_expiry_unix_ms = None; + self.windows.retain(|_, window| { + if window.expires_at_unix_ms <= now_unix_ms { + return false; + } + prune_usage_limit_events(&mut window.events, now_unix_ms, window.window_ms); + if window.events.is_empty() { + return false; + } + total_events = total_events.saturating_add(window.events.len()); + update_earliest_expiry(&mut next_expiry_unix_ms, window.expires_at_unix_ms); + for timestamp in window.events.values() { + update_earliest_expiry( + &mut next_expiry_unix_ms, + timestamp.saturating_add(window.window_ms), + ); + } + true + }); + self.total_events = total_events; + self.next_expiry_unix_ms = next_expiry_unix_ms; + } + + fn prune_rule_window(&mut self, key: &str, now_unix_ms: u64, window_ms: u64) { + if self + .windows + .get(key) + .is_some_and(|window| window.expires_at_unix_ms <= now_unix_ms) + { + if let Some(window) = self.windows.remove(key) { + self.total_events = self.total_events.saturating_sub(window.events.len()); + } + return; + } + let Some(window) = self.windows.get_mut(key) else { + return; + }; + let before = window.events.len(); + window.window_ms = window_ms; + prune_usage_limit_events(&mut window.events, now_unix_ms, window_ms); + self.total_events = self + .total_events + .saturating_sub(before.saturating_sub(window.events.len())); + update_earliest_expiry(&mut self.next_expiry_unix_ms, window.expires_at_unix_ms); + for timestamp in window.events.values() { + update_earliest_expiry( + &mut self.next_expiry_unix_ms, + timestamp.saturating_add(window_ms), + ); + } + if window.events.is_empty() { + self.windows.remove(key); + } + } + + fn additions_for(&self, input: crate::UsageLimitInput<'_>) -> (usize, usize) { + input.rules.iter().fold( + (0_usize, 0_usize), + |(additional_windows, additional_events), rule| match self.windows.get(rule.key) { + Some(window) if window.events.contains_key(input.event_id) => { + (additional_windows, additional_events) + } + Some(_) => (additional_windows, additional_events.saturating_add(1)), + None => ( + additional_windows.saturating_add(1), + additional_events.saturating_add(1), + ), + }, + ) + } +} + #[derive(Debug, Clone)] struct MemoryCounterEntry { value: u32, @@ -206,6 +322,34 @@ impl MemoryRuntimeBackend { ); } + pub(crate) async fn kv_set_if_absent(&self, key: &str, value: String, ttl: Duration) -> bool { + let mut kv = self.kv.lock().await; + let now = Instant::now(); + prune_kv(&mut kv, now); + if kv.contains_key(key) { + return false; + } + while kv.len() >= self.config.max_kv_entries.max(1) { + let Some(oldest_key) = kv + .iter() + .min_by_key(|(_, entry)| entry.inserted_at) + .map(|(key, _)| key.clone()) + else { + break; + }; + kv.remove(&oldest_key); + } + kv.insert( + key.to_string(), + MemoryKvEntry { + value, + inserted_at: now, + expires_at: Some(now + ttl), + }, + ); + true + } + pub(crate) fn kv_set_nowait(&self, key: &str, value: String, ttl: Option) -> bool { let Ok(mut kv) = self.kv.try_lock() else { return false; @@ -502,6 +646,143 @@ impl MemoryRuntimeBackend { }) } + pub(crate) async fn check_and_consume_usage_limits( + &self, + input: crate::UsageLimitInput<'_>, + ) -> Result { + let mut state = self.usage_limits.lock().await; + state.amortized_prune(input.now_unix_ms); + + for (index, rule) in input.rules.iter().enumerate() { + let window_ms = rule.window_seconds.saturating_mul(1_000); + state.prune_rule_window(rule.key, input.now_unix_ms, window_ms); + let Some(window) = state.windows.get(rule.key) else { + continue; + }; + if window.events.contains_key(input.event_id) { + continue; + } + if window.events.len() as u64 >= rule.limit { + let earliest = window + .events + .values() + .copied() + .min() + .unwrap_or(input.now_unix_ms); + let retry_after_ms = earliest + .saturating_add(window_ms) + .saturating_sub(input.now_unix_ms); + return Ok(UsageLimitCheck::Rejected { + rule_index: index, + limit: rule.limit, + retry_after: retry_after_ms.saturating_add(999) / 1_000, + }); + } + } + + let (mut additional_windows, mut additional_events) = state.additions_for(input); + if state.windows.len().saturating_add(additional_windows) + > self.config.max_usage_limit_windows + || state.total_events.saturating_add(additional_events) + > self.config.max_usage_limit_events + { + // Redis drops an idle sorted-set key after its retention TTL. Force the equivalent full + // cleanup before rejecting capacity so stale high-cardinality keys cannot pin memory. + if state + .next_expiry_unix_ms + .is_some_and(|expires_at| expires_at <= input.now_unix_ms) + { + state.prune_all(input.now_unix_ms); + } + (additional_windows, additional_events) = state.additions_for(input); + } + if state.windows.len().saturating_add(additional_windows) + > self.config.max_usage_limit_windows + || state.total_events.saturating_add(additional_events) + > self.config.max_usage_limit_events + { + return Err(DataLayerError::UnexpectedValue(format!( + "runtime memory usage-limit capacity exhausted (windows {}/{}, events {}/{})", + state.windows.len(), + self.config.max_usage_limit_windows, + state.total_events, + self.config.max_usage_limit_events, + ))); + } + + for rule in input.rules { + let window_ms = rule.window_seconds.saturating_mul(1_000); + let expires_at_unix_ms = input + .now_unix_ms + .saturating_add(rule.retention_seconds.saturating_mul(1_000)); + let inserted = match state.windows.entry(rule.key.to_string()) { + std::collections::hash_map::Entry::Occupied(mut entry) => { + let window = entry.get_mut(); + window.window_ms = window_ms; + window.expires_at_unix_ms = expires_at_unix_ms; + match window.events.entry(input.event_id.to_string()) { + std::collections::hash_map::Entry::Occupied(_) => false, + std::collections::hash_map::Entry::Vacant(entry) => { + entry.insert(input.now_unix_ms); + true + } + } + } + std::collections::hash_map::Entry::Vacant(entry) => { + entry.insert(MemoryUsageLimitWindow { + window_ms, + expires_at_unix_ms, + events: HashMap::from([(input.event_id.to_string(), input.now_unix_ms)]), + }); + true + } + }; + update_earliest_expiry(&mut state.next_expiry_unix_ms, expires_at_unix_ms); + if inserted { + state.total_events = state.total_events.saturating_add(1); + update_earliest_expiry( + &mut state.next_expiry_unix_ms, + input.now_unix_ms.saturating_add(window_ms), + ); + } + } + Ok(UsageLimitCheck::Allowed) + } + + pub(crate) async fn release_usage_limits( + &self, + input: crate::UsageLimitReleaseInput<'_>, + ) -> Result<(), DataLayerError> { + let mut state = self.usage_limits.lock().await; + for rule in input.rules { + let mut remove_window = false; + let mut removed_event = false; + if let Some(window) = state.windows.get_mut(rule.key) { + removed_event = window.events.remove(input.event_id).is_some(); + remove_window = window.events.is_empty(); + } + if removed_event { + state.total_events = state.total_events.saturating_sub(1); + } + if remove_window { + state.windows.remove(rule.key); + } + } + state.next_expiry_unix_ms = state + .windows + .values() + .flat_map(|window| { + std::iter::once(window.expires_at_unix_ms).chain( + window + .events + .values() + .map(|timestamp| timestamp.saturating_add(window.window_ms)), + ) + }) + .min(); + Ok(()) + } + pub(crate) fn rate_limit_count(&self, key: &str, bucket: u64) -> Result { let now = Instant::now(); let mut total = 0_u32; @@ -1013,6 +1294,17 @@ impl MemoryRuntimeBackend { } } +fn prune_usage_limit_events(events: &mut HashMap, now_unix_ms: u64, window_ms: u64) { + let Some(cutoff) = now_unix_ms.checked_sub(window_ms) else { + return; + }; + events.retain(|_, timestamp| *timestamp > cutoff); +} + +fn update_earliest_expiry(current: &mut Option, candidate: u64) { + *current = Some(current.map_or(candidate, |existing| existing.min(candidate))); +} + fn get_fresh_locked( kv: &mut HashMap, key: &str, @@ -1197,4 +1489,138 @@ mod tests { assert!(!shard.entries.contains_key("expired-unrelated-key")); assert_eq!(shard.operations_since_prune, 0); } + + #[tokio::test] + async fn usage_limit_capacity_is_atomic_and_fail_closed() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig { + max_usage_limit_windows: 2, + max_usage_limit_events: 2, + ..MemoryRuntimeStateConfig::default() + }); + let first = [crate::UsageLimitRule { + key: "usage:{user-1}:one", + limit: 10, + window_seconds: 60, + retention_seconds: 60, + }]; + backend + .check_and_consume_usage_limits(crate::UsageLimitInput { + rules: &first, + event_id: "event-1", + now_unix_ms: 1_000, + }) + .await + .expect("first event"); + + let two_new_windows = [ + crate::UsageLimitRule { + key: "usage:{user-1}:two", + limit: 10, + window_seconds: 60, + retention_seconds: 60, + }, + crate::UsageLimitRule { + key: "usage:{user-1}:three", + limit: 10, + window_seconds: 60, + retention_seconds: 60, + }, + ]; + let error = backend + .check_and_consume_usage_limits(crate::UsageLimitInput { + rules: &two_new_windows, + event_id: "event-2", + now_unix_ms: 2_000, + }) + .await + .expect_err("capacity must fail closed"); + assert!(error.to_string().contains("capacity exhausted")); + + let state = backend.usage_limits.lock().await; + assert_eq!(state.windows.len(), 1); + assert_eq!(state.total_events, 1); + assert!(!state.windows.contains_key(two_new_windows[0].key)); + assert!(!state.windows.contains_key(two_new_windows[1].key)); + } + + #[tokio::test] + async fn usage_limit_capacity_reclaims_expired_windows_before_rejecting() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig { + max_usage_limit_windows: 1, + max_usage_limit_events: 1, + ..MemoryRuntimeStateConfig::default() + }); + let old = [crate::UsageLimitRule { + key: "usage:{user-1}:old", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }]; + backend + .check_and_consume_usage_limits(crate::UsageLimitInput { + rules: &old, + event_id: "event-old", + now_unix_ms: 1_000, + }) + .await + .expect("old event"); + + let current = [crate::UsageLimitRule { + key: "usage:{user-1}:current", + limit: 1, + window_seconds: 1, + retention_seconds: 1, + }]; + assert_eq!( + backend + .check_and_consume_usage_limits(crate::UsageLimitInput { + rules: ¤t, + event_id: "event-current", + now_unix_ms: 2_000, + }) + .await + .expect("expired capacity should be reclaimed"), + UsageLimitCheck::Allowed + ); + + let state = backend.usage_limits.lock().await; + assert_eq!(state.windows.len(), 1); + assert_eq!(state.total_events, 1); + assert!(state.windows.contains_key(current[0].key)); + } + + #[tokio::test] + async fn usage_limit_idempotent_replay_does_not_consume_event_capacity() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig { + max_usage_limit_windows: 1, + max_usage_limit_events: 1, + ..MemoryRuntimeStateConfig::default() + }); + let rules = [crate::UsageLimitRule { + key: "usage:{user-1}:idempotent", + limit: 10, + window_seconds: 60, + retention_seconds: 60, + }]; + for now_unix_ms in [1_000, 2_000] { + assert_eq!( + backend + .check_and_consume_usage_limits(crate::UsageLimitInput { + rules: &rules, + event_id: "same-event", + now_unix_ms, + }) + .await + .expect("idempotent replay"), + UsageLimitCheck::Allowed + ); + } + + let state = backend.usage_limits.lock().await; + assert_eq!(state.total_events, 1); + assert_eq!( + state.windows[rules[0].key].events["same-event"], 1_000, + "idempotent replay must preserve the original Redis ZADD NX timestamp" + ); + } } diff --git a/crates/aether-runtime/state/src/redis/client.rs b/crates/aether-runtime/state/src/redis/client.rs index a22499170..42069b582 100644 --- a/crates/aether-runtime/state/src/redis/client.rs +++ b/crates/aether-runtime/state/src/redis/client.rs @@ -17,12 +17,46 @@ pub(crate) const REDIS_COMMAND_LATENCY_BUCKETS_MS: [u64; 12] = [1, 5, 10, 25, 50, 100, 250, 500, 1_000, 2_500, 5_000, 10_000]; const REDIS_COMMAND_LATENCY_BUCKET_COUNT: usize = REDIS_COMMAND_LATENCY_BUCKETS_MS.len() + 1; -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)] +#[derive(Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)] pub struct RedisClientConfig { pub url: String, pub key_prefix: Option, } +impl std::fmt::Debug for RedisClientConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("RedisClientConfig") + .field("url", &redact_redis_url_for_debug(&self.url)) + .field("key_prefix_len", &self.key_prefix.as_ref().map(String::len)) + .finish() + } +} + +fn redact_redis_url_for_debug(raw: &str) -> String { + const MAX_DEBUG_URL_CHARS: usize = 512; + let raw = raw.trim(); + let Ok(mut url) = url::Url::parse(raw) else { + return format!("[invalid-redis-url len={}]", raw.len()); + }; + let _ = url.set_username(""); + let _ = url.set_password(None); + url.set_query(None); + url.set_fragment(None); + let rendered = url.to_string(); + if rendered.chars().count() <= MAX_DEBUG_URL_CHARS { + rendered + } else { + format!( + "{}...", + rendered + .chars() + .take(MAX_DEBUG_URL_CHARS.saturating_sub(3)) + .collect::() + ) + } +} + impl RedisClientConfig { pub fn validate(&self) -> Result<(), DataLayerError> { let raw = self.url.trim(); @@ -466,6 +500,25 @@ mod tests { .expect("lazy redis client should build"); } + #[test] + fn redis_config_debug_redacts_url_credentials_and_query() { + let config = RedisClientConfig { + url: "redis://redis-user:redis-password@redis.example/0?token=redis-secret".into(), + key_prefix: Some("tenant-secret".into()), + }; + let debug = format!("{config:?}"); + for secret in [ + "redis-user", + "redis-password", + "redis-secret", + "tenant-secret", + ] { + assert!(!debug.contains(secret), "debug leaked {secret}: {debug}"); + } + assert!(debug.contains("redis.example")); + assert!(debug.contains("key_prefix_len")); + } + #[test] fn blocking_stream_lane_count_uses_requested_as_floor() { let default_lanes = default_blocking_stream_lane_count(); diff --git a/crates/aether-runtime/state/src/redis/runtime.rs b/crates/aether-runtime/state/src/redis/runtime.rs index ccdfda5b1..4be7d84fb 100644 --- a/crates/aether-runtime/state/src/redis/runtime.rs +++ b/crates/aether-runtime/state/src/redis/runtime.rs @@ -7,6 +7,7 @@ use crate::redis::{ }; use crate::{ DataLayerError, RateLimitCheck, RateLimitInput, RateLimitScope, RuntimeSemaphoreError, + UsageLimitCheck, UsageLimitInput, UsageLimitReleaseInput, }; const RATE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#" @@ -51,6 +52,50 @@ end return {1, 0, 0, remaining} "#; +const USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#" +local count = #KEYS +local now = tonumber(ARGV[1]) +local event_id = ARGV[2] + +for i = 1, count do + local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 + local cutoff = now - window_ms + redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', cutoff) +end + +for i = 1, count do + local limit = tonumber(ARGV[(i - 1) * 3 + 3]) + local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 + local already_consumed = redis.call('ZSCORE', KEYS[i], event_id) + if not already_consumed then + local current = redis.call('ZCARD', KEYS[i]) + if current >= limit then + local earliest = redis.call('ZRANGE', KEYS[i], 0, 0, 'WITHSCORES') + local retry_after = 1 + if #earliest >= 2 then + local retry_after_ms = math.max(1, tonumber(earliest[2]) + window_ms - now) + retry_after = math.ceil(retry_after_ms / 1000) + end + return {0, i, limit, retry_after} + end + end +end + +for i = 1, count do + local retention = tonumber(ARGV[(i - 1) * 3 + 5]) + redis.call('ZADD', KEYS[i], 'NX', now, event_id) + redis.call('EXPIRE', KEYS[i], retention + 1) +end +return {1, 0, 0, 0} +"#; + +const USAGE_LIMIT_RELEASE_SCRIPT: &str = r#" +for i = 1, #KEYS do + redis.call('ZREM', KEYS[i], ARGV[1]) +end +return 1 +"#; + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)] pub struct RedisRuntimeDiagnostics { pub connected_clients: Option, @@ -147,6 +192,37 @@ impl RedisRuntimeRunner { Ok(()) } + pub(crate) async fn kv_set_if_absent( + &self, + key: &str, + value: String, + ttl: Duration, + ) -> Result { + let namespaced_key = self.keyspace.key(key); + let ttl_ms = u64::try_from(ttl.as_millis().max(1)).unwrap_or(u64::MAX); + let mut command = cmd("SET"); + command + .arg(namespaced_key) + .arg(value) + .arg("NX") + .arg("PX") + .arg(ttl_ms); + let response = self + .query::>( + RedisConnectionLane::Fast, + "runtime kv set if absent", + command, + ) + .await?; + match response.as_deref() { + Some("OK") => Ok(true), + None => Ok(false), + Some(value) => Err(DataLayerError::UnexpectedValue(format!( + "unexpected runtime kv set if absent response {value}" + ))), + } + } + pub(crate) async fn kv_get_many( &self, keys: &[String], @@ -282,6 +358,105 @@ impl RedisRuntimeRunner { Ok(RateLimitCheck::Rejected { scope, limit }) } + pub(crate) async fn check_and_consume_usage_limits( + &self, + input: UsageLimitInput<'_>, + ) -> Result { + let script = script(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT); + let mut invocation = script.prepare_invoke(); + for rule in input.rules { + invocation.key(self.keyspace.key(rule.key)); + } + invocation.arg(input.now_unix_ms as i64); + invocation.arg(input.event_id); + for rule in input.rules { + invocation.arg(rule.limit as i64); + invocation.arg(rule.window_seconds as i64); + invocation.arg(rule.retention_seconds as i64); + } + let raw = run_lane_with_timeout( + &self.connections, + RedisConnectionLane::Fast, + self.command_timeout_ms, + "runtime usage limit check", + async { + let mut connection = self.connections.connection(RedisConnectionLane::Fast); + invocation + .invoke_async::>(&mut connection) + .await + .map_redis_err() + }, + ) + .await?; + match raw.first().copied() { + Some(1) if raw.len() >= 4 => return Ok(UsageLimitCheck::Allowed), + Some(0) if raw.len() >= 4 => {} + _ => { + return Err(DataLayerError::UnexpectedValue( + "runtime usage limit script returned an invalid response".to_string(), + )); + } + } + let rule_index = raw + .get(1) + .copied() + .and_then(|value| usize::try_from(value.saturating_sub(1)).ok()) + .filter(|index| *index < input.rules.len()) + .ok_or_else(|| { + DataLayerError::UnexpectedValue( + "runtime usage limit script returned an invalid rule index".to_string(), + ) + })?; + let limit = raw + .get(2) + .copied() + .and_then(|value| u64::try_from(value).ok()) + .filter(|limit| *limit > 0) + .ok_or_else(|| { + DataLayerError::UnexpectedValue( + "runtime usage limit script returned an invalid limit".to_string(), + ) + })?; + let retry_after = raw + .get(3) + .copied() + .and_then(|value| u64::try_from(value).ok()) + .unwrap_or(1) + .max(1); + Ok(UsageLimitCheck::Rejected { + rule_index, + limit, + retry_after, + }) + } + + pub(crate) async fn release_usage_limits( + &self, + input: UsageLimitReleaseInput<'_>, + ) -> Result<(), DataLayerError> { + let script = script(USAGE_LIMIT_RELEASE_SCRIPT); + let mut invocation = script.prepare_invoke(); + for rule in input.rules { + invocation.key(self.keyspace.key(rule.key)); + } + invocation.arg(input.event_id); + run_lane_with_timeout( + &self.connections, + RedisConnectionLane::Fast, + self.command_timeout_ms, + "runtime usage limit release", + async { + let mut connection = self.connections.connection(RedisConnectionLane::Fast); + invocation + .invoke_async::(&mut connection) + .await + .map_redis_err() + }, + ) + .await?; + Ok(()) + } + pub(crate) async fn set_add(&self, key: &str, member: &str) -> Result { let key = self.keyspace.key(key); let mut command = cmd("SADD"); diff --git a/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs b/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs index 712d0fbca..b82892f51 100644 --- a/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs +++ b/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs @@ -7,11 +7,11 @@ use std::time::{Duration, Instant}; use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody}; use aether_gateway::tunnel_protocol as protocol; use aether_testkit::{ - fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, run_http_load_probe, - BenchmarkRuntimeSnapshot, ExecutionRuntimeHarness, ExecutionRuntimeHarnessConfig, - GatewayHarness, GatewayHarnessConfig, HttpLoadProbeConfig, HttpLoadProbeResponseMode, - HttpLoadProbeResult, SpawnedServer, TunnelHarness, TunnelHarnessConfig, - GATEWAY_HARNESS_API_KEY, + fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, + insert_tunnel_harness_auth_headers, run_http_load_probe, BenchmarkRuntimeSnapshot, + ExecutionRuntimeHarness, ExecutionRuntimeHarnessConfig, GatewayHarness, GatewayHarnessConfig, + HttpLoadProbeConfig, HttpLoadProbeResponseMode, HttpLoadProbeResult, SpawnedServer, + TunnelHarness, TunnelHarnessConfig, GATEWAY_HARNESS_API_KEY, TUNNEL_HARNESS_NODE_ID, }; use axum::body::{to_bytes, Body, Bytes}; use axum::http::StatusCode; @@ -308,6 +308,7 @@ async fn run_tunnel_curve( for limit in &config.points { let relay_concurrency = (*limit).saturating_sub(1).max(1); let tunnel = TunnelHarness::start(TunnelHarnessConfig { + node_id: TUNNEL_HARNESS_NODE_ID.to_string(), max_streams: (*limit).max(128), ping_interval: Duration::from_secs(15), outbound_queue_capacity: 128, @@ -657,9 +658,7 @@ async fn connect_protocol_peer( ); let request = ws_url.into_client_request()?; let mut request = request; - request - .headers_mut() - .insert("x-node-id", http::HeaderValue::from_static("node-baseline")); + insert_tunnel_harness_auth_headers(request.headers_mut(), TUNNEL_HARNESS_NODE_ID)?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, http::HeaderValue::from_static( diff --git a/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs b/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs index 81a314b0a..bbc58bfae 100644 --- a/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs +++ b/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs @@ -11,8 +11,9 @@ use aether_data::driver::postgres::{ use aether_data::{DataLayerError, PostgresBackend}; use aether_runtime_state::{RedisClientConfig, RedisLockRunner, RedisLockRunnerConfig}; use aether_testkit::{ - init_test_runtime_for, reserve_local_port, BenchmarkRuntimeSampler, BenchmarkRuntimeSnapshot, - ManagedPostgresServer, ManagedRedisServer, TunnelHarness, TunnelHarnessConfig, + init_test_runtime_for, insert_tunnel_harness_auth_headers, reserve_local_port, + BenchmarkRuntimeSampler, BenchmarkRuntimeSnapshot, ManagedPostgresServer, ManagedRedisServer, + TunnelHarness, TunnelHarnessConfig, TUNNEL_HARNESS_NODE_ID, }; use futures_util::{FutureExt, StreamExt}; use serde::Serialize; @@ -508,12 +509,8 @@ async fn benchmark_tunnel_restart_recovery( .clone() .into_client_request() .map_err(|err| format!("failed to build websocket request: {err}"))?; - request.headers_mut().insert( - "x-node-id", - format!("recovery-node-{worker_index}-{current}") - .parse() - .map_err(|err| format!("failed to build x-node-id header: {err}"))?, - ); + insert_tunnel_harness_auth_headers(request.headers_mut(), TUNNEL_HARNESS_NODE_ID) + .map_err(|err| format!("failed to build tunnel auth headers: {err}"))?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION_STR diff --git a/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs b/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs index f6062e5f5..a5d293492 100644 --- a/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs +++ b/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs @@ -1,7 +1,11 @@ // Gateway-backed benchmark scenarios live outside the reusable testkit. use std::env; +#[cfg(unix)] use std::fs; -use std::path::PathBuf; +use std::io; +#[cfg(unix)] +use std::io::Write; +use std::path::{Path, PathBuf}; use aether_data::repository::auth::CreateStandaloneApiKeyRecord; use aether_data::repository::wallet::WalletLookupKey; @@ -522,7 +526,11 @@ async fn seed_api_key( .update_standalone_api_key_basic( aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord { api_key_id: api_key_id.clone(), + key_encrypted: None, + key_encrypted_present: false, name: Some(format!("Local pressure API key {}", key_index + 1)), + name_present: true, + force_capabilities: None, rate_limit_present: true, rate_limit: Some(0), concurrent_limit_present: true, @@ -655,22 +663,18 @@ async fn verify_candidate_selection( } fn write_outputs(config: &Config) -> Result<(), Box> { - if let Some(parent) = config.output_env_path.parent() { - fs::create_dir_all(parent)?; - } - if let Some(parent) = config.output_key_path.parent() { - fs::create_dir_all(parent)?; - } - if let Some(parent) = config.output_key_list_path.parent() { - fs::create_dir_all(parent)?; - } - - fs::write(&config.output_key_path, format!("{}\n", config.api_key))?; + write_private_output( + &config.output_key_path, + format!("{}\n", config.api_key).as_bytes(), + )?; let key_list = (0..config.api_key_count) .map(|index| pressure_api_key_value(config, index)) .collect::>() .join("\n"); - fs::write(&config.output_key_list_path, format!("{key_list}\n"))?; + write_private_output( + &config.output_key_list_path, + format!("{key_list}\n").as_bytes(), + )?; let env_content = format!( concat!( "export AETHER_API_KEY_FILE={key_path}\n", @@ -689,11 +693,110 @@ fn write_outputs(config: &Config) -> Result<(), Box> { model = shell_escape(&config.model), mock_upstream_base_url = shell_escape(&config.mock_upstream_base_url), ); - fs::write(&config.output_env_path, env_content)?; + write_private_output(&config.output_env_path, env_content.as_bytes())?; Ok(()) } +fn write_private_output(path: &Path, contents: &[u8]) -> io::Result<()> { + #[cfg(not(unix))] + { + let _ = (path, contents); + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "private benchmark credential outputs currently require Unix filesystem checks", + )); + } + + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; + + let file_name = path.file_name().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "private output path must name a file", + ) + })?; + let input_parent = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + let parent = fs::canonicalize(input_parent)?; + let parent_metadata = fs::symlink_metadata(&parent)?; + if !parent_metadata.is_dir() || parent_metadata.file_type().is_symlink() { + return Err(io::Error::other( + "private output parent must be a real directory", + )); + } + + let target = parent.join(file_name); + let temporary = parent.join(format!( + ".aether-pressure-output-{}.tmp", + uuid::Uuid::new_v4() + )); + let mut file = fs::OpenOptions::new() + .write(true) + .create_new(true) + .mode(0o600) + .open(&temporary)?; + + let result = (|| -> io::Result<()> { + let owner_uid = file.metadata()?.uid(); + validate_private_output_directory(&parent, owner_uid)?; + match fs::symlink_metadata(&target) { + Ok(metadata) + if metadata.is_file() + && !metadata.file_type().is_symlink() + && metadata.uid() == owner_uid + && metadata.nlink() == 1 => {} + Ok(_) => { + return Err(io::Error::other( + "refusing to replace a symlink, special file, hard link, or foreign-owned private output", + )); + } + Err(error) if error.kind() == io::ErrorKind::NotFound => {} + Err(error) => return Err(error), + } + + file.set_permissions(fs::Permissions::from_mode(0o600))?; + file.write_all(contents)?; + file.sync_all()?; + drop(file); + fs::rename(&temporary, &target)?; + fs::File::open(&parent)?.sync_all() + })(); + + if result.is_err() { + let _ = fs::remove_file(&temporary); + } + result + } +} + +#[cfg(unix)] +fn validate_private_output_directory(directory: &Path, owner_uid: u32) -> io::Result<()> { + use std::os::unix::fs::MetadataExt; + + let mut ancestor = Some(directory); + while let Some(path) = ancestor { + let metadata = fs::symlink_metadata(path)?; + let mode = metadata.mode(); + if !metadata.is_dir() + || metadata.file_type().is_symlink() + || (metadata.uid() != owner_uid && metadata.uid() != 0) + || (mode & 0o022 != 0 && mode & 0o1000 == 0) + { + return Err(io::Error::other(format!( + "private output directory '{}' has unsafe ownership or permissions", + path.display() + ))); + } + ancestor = path.parent(); + } + Ok(()) +} + fn sha256_hex(value: &str) -> String { let mut hasher = sha2::Sha256::new(); hasher.update(value.as_bytes()); @@ -789,7 +892,7 @@ Options:\n\ mod tests { use serde_json::json; - use super::pressure_provider_transport_config; + use super::{pressure_provider_transport_config, write_private_output}; #[test] fn pressure_provider_transport_config_enables_h2c_prior_knowledge() { @@ -811,4 +914,38 @@ mod tests { fn pressure_provider_transport_config_is_absent_by_default() { assert_eq!(pressure_provider_transport_config(false), None); } + + #[cfg(unix)] + #[test] + fn private_outputs_are_atomic_private_and_refuse_link_targets() { + use std::os::unix::fs::{symlink, MetadataExt, PermissionsExt}; + + let root = std::env::temp_dir().join(format!( + "aether-pressure-private-output-test-{}", + uuid::Uuid::new_v4() + )); + std::fs::create_dir(&root).unwrap(); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700)).unwrap(); + + let output = root.join("api-key"); + write_private_output(&output, b"secret\n").unwrap(); + let metadata = std::fs::symlink_metadata(&output).unwrap(); + assert_eq!(std::fs::read(&output).unwrap(), b"secret\n"); + assert_eq!(metadata.mode() & 0o777, 0o600); + assert_eq!(metadata.nlink(), 1); + + let victim = root.join("victim"); + std::fs::write(&victim, b"known-good").unwrap(); + std::fs::remove_file(&output).unwrap(); + symlink(&victim, &output).unwrap(); + assert!(write_private_output(&output, b"replacement\n").is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_file(&output).unwrap(); + std::fs::hard_link(&victim, &output).unwrap(); + assert!(write_private_output(&output, b"replacement\n").is_err()); + assert_eq!(std::fs::read(&victim).unwrap(), b"known-good"); + + std::fs::remove_dir_all(root).unwrap(); + } } diff --git a/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs b/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs index 5dcfc1e93..a29cf41c9 100644 --- a/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs +++ b/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs @@ -5,9 +5,10 @@ use std::time::Duration; use aether_gateway::tunnel_protocol as protocol; use aether_testkit::{ - fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, run_http_load_probe, - HttpLoadProbeConfig, HttpLoadProbeResponseMode, HttpLoadProbeResult, PrometheusSample, - TunnelHarness, TunnelHarnessConfig, + fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, + insert_tunnel_harness_auth_headers, run_http_load_probe, HttpLoadProbeConfig, + HttpLoadProbeResponseMode, HttpLoadProbeResult, PrometheusSample, TunnelHarness, + TunnelHarnessConfig, TUNNEL_HARNESS_NODE_ID, }; use futures_util::{SinkExt, StreamExt}; use reqwest::Method; @@ -242,9 +243,7 @@ async fn connect_protocol_peer( ); let request = ws_url.into_client_request()?; let mut request = request; - request - .headers_mut() - .insert("x-node-id", http::HeaderValue::from_static("node-baseline")); + insert_tunnel_harness_auth_headers(request.headers_mut(), TUNNEL_HARNESS_NODE_ID)?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, http::HeaderValue::from_static( diff --git a/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs b/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs index 9e30ee5be..bbe8f2a74 100644 --- a/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs +++ b/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs @@ -8,7 +8,8 @@ use std::time::{Duration, Instant}; use aether_gateway::tunnel_protocol as protocol; use aether_testkit::{ fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, - BenchmarkRuntimeSampler, BenchmarkRuntimeSnapshot, TunnelHarness, TunnelHarnessConfig, + insert_tunnel_harness_auth_headers, BenchmarkRuntimeSampler, BenchmarkRuntimeSnapshot, + TunnelHarness, TunnelHarnessConfig, }; use futures_util::stream::SplitSink; use futures_util::{SinkExt, StreamExt}; @@ -341,6 +342,7 @@ async fn run_suite( let chunks_per_stream = config.effective_chunks_per_stream(); let config = Arc::new(config); let tunnel = TunnelHarness::start(TunnelHarnessConfig { + node_id: config.node_id.clone(), max_streams: config.tunnel_max_streams, ping_interval: config.ping_interval, outbound_queue_capacity: config.outbound_queue_capacity, @@ -633,10 +635,7 @@ async fn connect_protocol_peer( ); let request = ws_url.into_client_request()?; let mut request = request; - request.headers_mut().insert( - "x-node-id", - http::HeaderValue::from_str(config.node_id.as_str())?, - ); + insert_tunnel_harness_auth_headers(request.headers_mut(), config.node_id.as_str())?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, http::HeaderValue::from_static( diff --git a/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs b/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs index 58c297519..8dacebdb2 100644 --- a/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs +++ b/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs @@ -37,6 +37,7 @@ const MAX_MOCK_REQUEST_BODY_BYTES: usize = 1024 * 1024; const BASIS_POINTS: u16 = 10_000; const DEFAULT_TIMEOUT_HOLD_MS: u64 = 60_000; const REQUEST_SEQUENCE_HEADER: &str = "x-mock-request-sequence"; +const TRUNCATED_STREAM_FLUSH_DELAY: Duration = Duration::from_millis(10); const RANDOM_DOMAIN_FAULT: u64 = 0x5c32_22f7_27d4_7a6f; const RANDOM_DOMAIN_FIRST_BYTE: u64 = 0x087d_89d9_3bc3_15db; @@ -646,8 +647,10 @@ fn build_chat_sse_response( if profile.truncate_after_chunks == Some(0) || profile.truncate_after_chunks == Some(index + 1) { - // Force Hyper to flush the successful frame before observing the body error. - tokio::task::yield_now().await; + // Hyper translates body errors into RST_STREAM for HTTP/2. Keep the body + // pending briefly so the response headers and successful DATA frame are + // written before Hyper observes the error. + tokio::time::sleep(TRUNCATED_STREAM_FLUSH_DELAY).await; record_fault(&app, Fault::TruncateStream); yield Err::(truncated_stream_error()); return; @@ -701,8 +704,10 @@ fn build_responses_sse_response( if profile.truncate_after_chunks == Some(0) || profile.truncate_after_chunks == Some(index + 1) { - // Force Hyper to flush the successful frame before observing the body error. - tokio::task::yield_now().await; + // Hyper translates body errors into RST_STREAM for HTTP/2. Keep the body + // pending briefly so the response headers and successful DATA frame are + // written before Hyper observes the error. + tokio::time::sleep(TRUNCATED_STREAM_FLUSH_DELAY).await; record_fault(&app, Fault::TruncateStream); yield Err::(truncated_stream_error()); return; diff --git a/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs b/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs index 0ebbfc09f..29280165c 100644 --- a/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs +++ b/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs @@ -9,11 +9,11 @@ use aether_runtime_state::{ RedisClientConfig, RuntimeSemaphore, RuntimeSemaphoreConfig, RuntimeState, }; use aether_testkit::{ - init_test_runtime_for, run_multi_url_http_load_probe, BenchmarkRuntimeSampler, - BenchmarkRuntimeSnapshot, ExecutionRuntimeHarness, ExecutionRuntimeHarnessConfig, - GatewayHarness, GatewayHarnessConfig, HttpLoadProbeConfig, HttpLoadProbeResponseMode, - ManagedRedisServer, MultiUrlHttpLoadProbeResult, SpawnedServer, TunnelHarness, - TunnelHarnessConfig, GATEWAY_HARNESS_API_KEY, + init_test_runtime_for, insert_tunnel_harness_auth_headers, run_multi_url_http_load_probe, + BenchmarkRuntimeSampler, BenchmarkRuntimeSnapshot, ExecutionRuntimeHarness, + ExecutionRuntimeHarnessConfig, GatewayHarness, GatewayHarnessConfig, HttpLoadProbeConfig, + HttpLoadProbeResponseMode, ManagedRedisServer, MultiUrlHttpLoadProbeResult, SpawnedServer, + TunnelHarness, TunnelHarnessConfig, GATEWAY_HARNESS_API_KEY, TUNNEL_HARNESS_NODE_ID, }; use axum::body::to_bytes; use axum::extract::Request; @@ -493,12 +493,8 @@ async fn run_tunnel_proxy_connection_probe( let mut request = url .into_client_request() .map_err(|err| format!("failed to build websocket request: {err}"))?; - request.headers_mut().insert( - "x-node-id", - format!("baseline-node-{worker_index}-{current}") - .parse() - .map_err(|err| format!("failed to build x-node-id header: {err}"))?, - ); + insert_tunnel_harness_auth_headers(request.headers_mut(), TUNNEL_HARNESS_NODE_ID) + .map_err(|err| format!("failed to build tunnel auth headers: {err}"))?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION_STR diff --git a/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs b/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs index de432fde2..b4b8c5e6d 100644 --- a/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs +++ b/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs @@ -5,13 +5,15 @@ use std::time::Duration; use aether_gateway::tunnel_protocol as protocol; use aether_gateway::GatewayDataConfig; use aether_testkit::{ - init_test_runtime_for, prepare_aether_postgres_schema, reserve_local_port, run_http_load_probe, - wait_until, GatewayHarness, GatewayHarnessConfig, HttpLoadProbeConfig, - HttpLoadProbeResponseMode, HttpLoadProbeResult, ManagedPostgresServer, ManagedRedisServer, + init_test_runtime_for, insert_tunnel_harness_auth_headers, prepare_aether_postgres_schema, + reserve_local_port, run_http_load_probe, wait_until, GatewayHarness, GatewayHarnessConfig, + HttpLoadProbeConfig, HttpLoadProbeResponseMode, HttpLoadProbeResult, ManagedPostgresServer, + ManagedRedisServer, TUNNEL_HARNESS_GENERATION, TUNNEL_HARNESS_MANAGEMENT_TOKEN, }; use futures_util::{SinkExt, StreamExt}; use reqwest::Method; use serde::Serialize; +use sha2::{Digest, Sha256}; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::tungstenite::Message; @@ -114,6 +116,7 @@ async fn run_suite( .expect("postgres url should be resolved"); prepare_aether_postgres_schema(&postgres_url).await?; + seed_tunnel_auth(&postgres_url).await?; let key_prefix = format!("aether-owner-relay-baseline-{}", std::process::id()); let shared_data = GatewayDataConfig::from_postgres_url(postgres_url.clone(), false) @@ -287,9 +290,7 @@ async fn connect_protocol_peer( ); let request = ws_url.into_client_request()?; let mut request = request; - request - .headers_mut() - .insert("x-node-id", http::HeaderValue::from_static(NODE_ID)); + insert_tunnel_harness_auth_headers(request.headers_mut(), NODE_ID)?; request.headers_mut().insert( aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, http::HeaderValue::from_static( @@ -333,6 +334,98 @@ async fn connect_protocol_peer( })) } +async fn seed_tunnel_auth(postgres_url: &str) -> Result<(), Box> { + const USER_ID: &str = "user-owner-relay-baseline"; + const TOKEN_ID: &str = "token-owner-relay-baseline"; + + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(1) + .connect(postgres_url) + .await?; + let mut transaction = pool.begin().await?; + sqlx::query( + r#" +INSERT INTO users ( + id, email, username, role, auth_source, email_verified, is_active, is_deleted, + created_at, updated_at +) VALUES ($1, $2, $3, 'admin', 'local', TRUE, TRUE, FALSE, 1, 1) +ON CONFLICT (id) DO UPDATE SET + email = EXCLUDED.email, + username = EXCLUDED.username, + role = 'admin', + auth_source = 'local', + email_verified = TRUE, + is_active = TRUE, + is_deleted = FALSE, + updated_at = EXCLUDED.updated_at +"#, + ) + .bind(USER_ID) + .bind("owner-relay-baseline@example.com") + .bind("owner_relay_baseline_admin") + .execute(&mut *transaction) + .await?; + + let token_hash = format!( + "{:x}", + Sha256::digest(TUNNEL_HARNESS_MANAGEMENT_TOKEN.as_bytes()) + ); + sqlx::query( + r#" +INSERT INTO management_tokens ( + id, user_id, name, token_hash, token_prefix, permissions, usage_count, + is_active, created_at, updated_at +) VALUES ($1, $2, $3, $4, 'ae-tunnel-harness', $5, 0, TRUE, 1, 1) +ON CONFLICT (id) DO UPDATE SET + user_id = EXCLUDED.user_id, + name = EXCLUDED.name, + token_hash = EXCLUDED.token_hash, + token_prefix = EXCLUDED.token_prefix, + permissions = EXCLUDED.permissions, + is_active = TRUE, + updated_at = EXCLUDED.updated_at +"#, + ) + .bind(TOKEN_ID) + .bind(USER_ID) + .bind("owner relay tunnel token") + .bind(token_hash) + .bind(serde_json::json!(["admin:proxy_nodes:admin"])) + .execute(&mut *transaction) + .await?; + + sqlx::query( + r#" +INSERT INTO proxy_nodes ( + id, tunnel_generation, name, ip, port, status, heartbeat_interval, + active_connections, total_requests, is_manual, created_at, updated_at, + config_version, tunnel_mode, tunnel_connected, failed_requests, dns_failures, + stream_errors +) VALUES ( + $1, $2, 'owner relay baseline node', '127.0.0.1', 0, 'offline', 30, + 0, 0, FALSE, 1, 1, 0, TRUE, FALSE, 0, 0, 0 +) +ON CONFLICT (id) DO UPDATE SET + tunnel_generation = EXCLUDED.tunnel_generation, + name = EXCLUDED.name, + ip = EXCLUDED.ip, + port = EXCLUDED.port, + status = 'offline', + active_connections = 0, + tunnel_mode = TRUE, + tunnel_connected = FALSE, + updated_at = EXCLUDED.updated_at +"#, + ) + .bind(NODE_ID) + .bind(TUNNEL_HARNESS_GENERATION) + .execute(&mut *transaction) + .await?; + + transaction.commit().await?; + Ok(()) +} + async fn handle_binary_frame( sink: &mut S, data: Vec, diff --git a/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs b/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs index c5de9e9bc..14f3d9ea1 100644 --- a/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs +++ b/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs @@ -17,6 +17,7 @@ use sqlx::{PgPool, Row}; use tokio::sync::Mutex; const PROXY_NODE_ID: &str = "proxy-node-hotspot"; +const PROXY_NODE_TUNNEL_GENERATION: &str = "usage-aux-hotspot-generation"; const MANAGEMENT_TOKEN_ID: &str = "management-token-hotspot"; const API_KEY_ID: &str = "api-key-last-used-hotspot"; const USER_ID: &str = "usage-aux-hotspot-user"; @@ -324,6 +325,7 @@ async fn enqueue_aux_counter_deltas( fn proxy_delta_for_index(index: usize) -> ProxyNodeCounterDelta { ProxyNodeCounterDelta { node_id: PROXY_NODE_ID.to_string(), + expected_tunnel_generation: Some(PROXY_NODE_TUNNEL_GENERATION.to_string()), total_requests_delta: 1, failed_requests_delta: if index.is_multiple_of(10) { 1 } else { 0 }, dns_failures_delta: if index.is_multiple_of(25) { 1 } else { 0 }, @@ -457,10 +459,13 @@ ON CONFLICT (id) DO UPDATE SET sqlx::query( r#" INSERT INTO proxy_nodes ( - id, name, ip, port, status, total_requests, failed_requests, + id, tunnel_generation, name, ip, port, status, total_requests, failed_requests, dns_failures, stream_errors ) -VALUES ($1, 'usage aux hotspot proxy', '127.0.0.1', 8080, 'online', 0, 0, 0, 0) +VALUES ( + $1, 'usage-aux-hotspot-generation', 'usage aux hotspot proxy', + '127.0.0.1', 8080, 'online', 0, 0, 0, 0 +) ON CONFLICT (id) DO UPDATE SET total_requests = 0, failed_requests = 0, diff --git a/crates/aether-testing/integration/tests/responses_websocket_e2e.rs b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs index 9828d3436..6e6551940 100644 --- a/crates/aether-testing/integration/tests/responses_websocket_e2e.rs +++ b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs @@ -15,7 +15,7 @@ use std::sync::Arc; use std::time::Duration; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; -use aether_data::repository::auth::CreateStandaloneApiKeyRecord; +use aether_data::repository::auth::{CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord}; use aether_data::repository::wallet::WalletLookupKey; use aether_data::{ DataBackends, DataLayerConfig, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, @@ -162,6 +162,38 @@ async fn continuation_reuses_one_upstream_connection_and_bills_both_turns() -> R Ok(()) } +#[tokio::test] +async fn weekly_plan_request_limit_counts_turns_but_not_the_websocket_upgrade( +) -> Result<(), BoxError> { + let harness = Harness::start_with_weekly_request_limit(1).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "allowed turn"}))) + .await?; + receive_event(&mut client, "response.completed").await?; + + client + .send(response_create(json!({"input": "rejected turn"}))) + .await?; + let rejected = receive_error_or_close(&mut client) + .await? + .ok_or("gateway closed without a plan-limit error event")?; + assert_eq!(rejected["status"], json!(429)); + assert_eq!( + rejected.pointer("/error/code"), + Some(&json!("plan_usage_limit_exceeded")) + ); + assert_eq!( + harness.upstream.observed_events().await.len(), + 1, + "the rejected logical turn must never reach the upstream" + ); + + client.close(None).await?; + Ok(()) +} + #[tokio::test] async fn persisted_previous_response_can_continue_on_a_new_client_connection( ) -> Result<(), BoxError> { @@ -793,6 +825,28 @@ impl Harness { .await } + async fn start_with_weekly_request_limit(limit: u64) -> Result { + let mut harness = Self::start(UpstreamBehavior::CompleteEveryTurn).await?; + seed_weekly_request_limit(&harness.database.config, limit).await?; + // The gateway may have cached the pre-entitlement auth context during + // startup, so restart it after seeding the user-owned key and plan. + let data_config = GatewayDataConfig::from_database_config(harness.database.config.clone()) + .with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new()? + .with_data_config_and_background_isolation(data_config, false)? + .with_usage_runtime_config(UsageRuntimeConfig { + enabled: true, + ..UsageRuntimeConfig::default() + })?; + let gateway_server = SpawnedServer::start(build_router_with_state(state)).await?; + harness.websocket_url = format!( + "{}/v1/responses", + gateway_server.base_url().replacen("http://", "ws://", 1) + ); + harness._gateway_server = gateway_server; + Ok(harness) + } + async fn start_with( behavior: UpstreamBehavior, fixture: ProviderFixture, @@ -823,6 +877,7 @@ impl Harness { enabled: true, ..UsageRuntimeConfig::default() })?; + state.ensure_system_default_routing_group().await?; let gateway_server = SpawnedServer::start(build_router_with_state(state)).await?; let websocket_url = format!( "{}/v1/responses", @@ -1616,6 +1671,145 @@ async fn seed_client_api_key(backends: &DataBackends, user_id: &str) -> Result<( Ok(()) } +async fn seed_weekly_request_limit( + database: &SqlDatabaseConfig, + limit: u64, +) -> Result<(), BoxError> { + let backends = DataBackends::from_config(DataLayerConfig::from_database(database.clone()))?; + let user_id = backends + .read() + .users() + .ok_or("user reader unavailable")? + .find_user_auth_by_username("responses-ws-e2e") + .await? + .ok_or("seeded E2E user unavailable")? + .id; + backends + .write() + .auth_api_keys() + .ok_or("auth API key writer unavailable")? + .delete_standalone_api_key(API_KEY_ID) + .await?; + backends + .write() + .auth_api_keys() + .ok_or("auth API key writer unavailable")? + .create_user_api_key(CreateUserApiKeyRecord { + user_id: user_id.clone(), + api_key_id: API_KEY_ID.to_string(), + key_hash: sha256_hex(CLIENT_API_KEY), + key_encrypted: Some(CLIENT_API_KEY.to_string()), + name: Some("Responses WebSocket plan-policy E2E".to_string()), + allowed_providers: Some(vec![PROVIDER_ID.to_string()]), + allowed_api_formats: Some(vec!["openai:responses".to_string()]), + allowed_models: Some(vec![PUBLIC_MODEL.to_string()]), + ip_rules: None, + rate_limit: 0, + concurrent_limit: None, + force_capabilities: None, + feature_settings: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + }) + .await?; + + let wallet_id = backends + .read() + .wallets() + .ok_or("wallet reader unavailable")? + .find(WalletLookupKey::UserId(&user_id)) + .await? + .ok_or("seeded E2E user wallet unavailable")? + .id; + + let pool = backends + .sqlite() + .ok_or("SQLite backend unavailable")? + .pool(); + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH)? + .as_secs() + .min(i64::MAX as u64) as i64; + sqlx::query( + r#" +INSERT INTO billing_plans ( + id, title, description, price_amount, price_currency, duration_unit, + duration_value, enabled, sort_order, max_active_per_user, + purchase_limit_scope, entitlements_json, created_at, updated_at +) VALUES (?, 'WS weekly policy', NULL, 0, 'USD', 'month', 1, 1, 0, 1, + 'active_period', ?, ?, ?) +"#, + ) + .bind("plan-responses-ws-weekly") + .bind( + json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": {"kind": "calendar_week"}, + "limit": limit + }] + }]) + .to_string(), + ) + .bind(now) + .bind(now) + .execute(pool) + .await?; + sqlx::query( + r#" +INSERT INTO payment_orders ( + id, order_no, wallet_id, user_id, amount_usd, pay_currency, status, + payment_method, created_at, paid_at, credited_at, expires_at +) VALUES (?, ?, ?, ?, 0, 'USD', 'paid', 'test', ?, ?, ?, ?) +"#, + ) + .bind("order-responses-ws-weekly") + .bind("order-no-responses-ws-weekly") + .bind(&wallet_id) + .bind(&user_id) + .bind(now) + .bind(now) + .bind(now) + .bind(now + 86_400) + .execute(pool) + .await?; + sqlx::query( + r#" +INSERT INTO user_plan_entitlements ( + id, user_id, plan_id, payment_order_id, status, starts_at, expires_at, + entitlements_snapshot, created_at, updated_at +) VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?) +"#, + ) + .bind("entitlement-responses-ws-weekly") + .bind(&user_id) + .bind("plan-responses-ws-weekly") + .bind("order-responses-ws-weekly") + .bind(now - 1) + .bind(now + 86_400) + .bind( + json!([{ + "type": "usage_policy", + "rules": [{ + "metric": "request_count", + "window": {"kind": "calendar_week"}, + "limit": limit + }] + }]) + .to_string(), + ) + .bind(now) + .bind(now) + .execute(pool) + .await?; + Ok(()) +} + /// 打开 chat PII 脱敏:系统模块开关 + 这把 client key 的 feature 开关。 /// /// 规则集刻意不写:缺省即内置规则(含 email 规则),和生产上「只打开开关」的最小 diff --git a/crates/aether-testing/testkit/Cargo.toml b/crates/aether-testing/testkit/Cargo.toml index d0501e9b7..f18bf6ac8 100644 --- a/crates/aether-testing/testkit/Cargo.toml +++ b/crates/aether-testing/testkit/Cargo.toml @@ -12,11 +12,13 @@ gateway = ["dep:aether-gateway", "dep:aether-runtime-state"] postgres = ["dep:aether-data", "dep:sqlx"] [dependencies] +aether-contracts.workspace = true aether-loadtools.workspace = true aether-data = { workspace = true, optional = true } aether-gateway = { workspace = true, features = ["testkit"], optional = true } aether-runtime.workspace = true aether-runtime-state = { workspace = true, optional = true } axum.workspace = true +http.workspace = true sqlx = { workspace = true, features = ["postgres"], optional = true } tokio.workspace = true diff --git a/crates/aether-testing/testkit/src/lib.rs b/crates/aether-testing/testkit/src/lib.rs index 16f404788..5f4193381 100644 --- a/crates/aether-testing/testkit/src/lib.rs +++ b/crates/aether-testing/testkit/src/lib.rs @@ -35,4 +35,7 @@ pub use gateway::{GatewayHarness, GatewayHarnessConfig, GATEWAY_HARNESS_API_KEY} #[cfg(feature = "postgres")] pub use postgres::{prepare_aether_postgres_schema, ManagedPostgresServer}; #[cfg(feature = "gateway")] -pub use tunnel::{TunnelHarness, TunnelHarnessConfig}; +pub use tunnel::{ + insert_tunnel_harness_auth_headers, TunnelHarness, TunnelHarnessConfig, + TUNNEL_HARNESS_GENERATION, TUNNEL_HARNESS_MANAGEMENT_TOKEN, TUNNEL_HARNESS_NODE_ID, +}; diff --git a/crates/aether-testing/testkit/src/tunnel.rs b/crates/aether-testing/testkit/src/tunnel.rs index e01b52cf2..e3515b022 100644 --- a/crates/aether-testing/testkit/src/tunnel.rs +++ b/crates/aether-testing/testkit/src/tunnel.rs @@ -8,8 +8,13 @@ use aether_runtime_state::RuntimeSemaphore; use crate::server::SpawnedServer; +pub const TUNNEL_HARNESS_NODE_ID: &str = "node-baseline"; +pub const TUNNEL_HARNESS_GENERATION: &str = "tunnel-harness-generation-1"; +pub const TUNNEL_HARNESS_MANAGEMENT_TOKEN: &str = "ae-tunnel-harness-management-token"; + #[derive(Debug, Clone)] pub struct TunnelHarnessConfig { + pub node_id: String, pub max_streams: usize, pub ping_interval: Duration, pub outbound_queue_capacity: usize, @@ -20,6 +25,7 @@ pub struct TunnelHarnessConfig { impl Default for TunnelHarnessConfig { fn default() -> Self { Self { + node_id: TUNNEL_HARNESS_NODE_ID.to_string(), max_streams: 128, ping_interval: Duration::from_secs(15), outbound_queue_capacity: 128, @@ -62,6 +68,12 @@ impl TunnelHarness { } else { state }; + let state = aether_gateway::configure_test_tunnel_runtime_auth( + state, + &config.node_id, + TUNNEL_HARNESS_GENERATION, + TUNNEL_HARNESS_MANAGEMENT_TOKEN, + )?; let router = build_tunnel_runtime_router_with_state(state); let server = match port { Some(port) => SpawnedServer::start_on_port(port, router) @@ -82,3 +94,19 @@ impl TunnelHarness { self.server.port() } } + +pub fn insert_tunnel_harness_auth_headers( + headers: &mut http::HeaderMap, + node_id: &str, +) -> Result<(), http::header::InvalidHeaderValue> { + headers.insert("x-node-id", http::HeaderValue::from_str(node_id)?); + headers.insert( + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + http::HeaderValue::from_static(TUNNEL_HARNESS_GENERATION), + ); + headers.insert( + http::header::AUTHORIZATION, + http::HeaderValue::from_static(concat!("Bearer ", "ae-tunnel-harness-management-token")), + ); + Ok(()) +} diff --git a/crates/aether-usage/runtime/src/body_capture.rs b/crates/aether-usage/runtime/src/body_capture.rs index aca62f851..656e06d31 100644 --- a/crates/aether-usage/runtime/src/body_capture.rs +++ b/crates/aether-usage/runtime/src/body_capture.rs @@ -118,15 +118,34 @@ impl UsageBodyCaptureEngine { } pub fn apply_to_event(self, event: &mut UsageEvent) { - self.apply_to_payload(UsageBodyCapturePayloadMut::from_event(event)); + let force_disabled = is_sensitive_oauth_exchange( + event.data.api_format.as_deref(), + event.data.endpoint_api_format.as_deref(), + ); + self.apply_to_payload( + UsageBodyCapturePayloadMut::from_event(event), + force_disabled, + ); } pub fn apply_to_record(self, record: &mut UpsertUsageRecord) { - self.apply_to_payload(UsageBodyCapturePayloadMut::from_record(record)); + let force_disabled = is_sensitive_oauth_exchange( + record.api_format.as_deref(), + record.endpoint_api_format.as_deref(), + ); + self.apply_to_payload( + UsageBodyCapturePayloadMut::from_record(record), + force_disabled, + ); } - fn apply_to_payload(self, payload: UsageBodyCapturePayloadMut<'_>) { - if matches!(self.policy.record_level, UsageRequestRecordLevel::Basic) { + fn apply_to_payload(self, payload: UsageBodyCapturePayloadMut<'_>, force_disabled: bool) { + if force_disabled || matches!(self.policy.record_level, UsageRequestRecordLevel::Basic) { + let reason = if force_disabled { + "sensitive_oauth_exchange" + } else { + "request_record_level_basic" + }; disable_usage_body_capture_field( UsageBodyField::RequestBody, "request", @@ -134,6 +153,7 @@ impl UsageBodyCaptureEngine { payload.request_body_ref, payload.request_body_state, payload.request_metadata, + reason, ); disable_usage_body_capture_field( UsageBodyField::ProviderRequestBody, @@ -142,6 +162,7 @@ impl UsageBodyCaptureEngine { payload.provider_request_body_ref, payload.provider_request_body_state, payload.request_metadata, + reason, ); disable_usage_body_capture_field( UsageBodyField::ResponseBody, @@ -150,6 +171,7 @@ impl UsageBodyCaptureEngine { payload.response_body_ref, payload.response_body_state, payload.request_metadata, + reason, ); disable_usage_body_capture_field( UsageBodyField::ClientResponseBody, @@ -158,6 +180,7 @@ impl UsageBodyCaptureEngine { payload.client_response_body_ref, payload.client_response_body_state, payload.request_metadata, + reason, ); return; } @@ -201,6 +224,22 @@ impl UsageBodyCaptureEngine { } } +fn is_sensitive_oauth_exchange( + api_format: Option<&str>, + endpoint_api_format: Option<&str>, +) -> bool { + [api_format, endpoint_api_format] + .into_iter() + .flatten() + .map(str::trim) + .any(|format| { + format.eq_ignore_ascii_case("oauth:exchange") + || format.eq_ignore_ascii_case("provider_oauth:exchange") + || format.eq_ignore_ascii_case("provider_oauth:local_refresh") + || format.eq_ignore_ascii_case("vertex_ai:service_account_token") + }) +} + pub fn apply_usage_body_capture_policy_to_event( policy: UsageBodyCapturePolicy, event: &mut UsageEvent, @@ -222,6 +261,7 @@ fn disable_usage_body_capture_field( body_ref: &mut Option, state: &mut Option, request_metadata: &mut Option, + reason: &'static str, ) { *body = None; *body_ref = None; @@ -233,7 +273,7 @@ fn disable_usage_body_capture_field( Some(UsageBodyCaptureState::Disabled), None, None, - Some("request_record_level_basic"), + Some(reason), ); } @@ -712,19 +752,190 @@ fn usage_value_kind(value: &Value) -> &'static str { #[cfg(test)] mod tests { use super::{ + apply_usage_body_capture_policy_to_event, apply_usage_body_capture_policy_to_record, build_plan_body_capture_metadata, sync_usage_body_ref_metadata, trim_owned_non_empty_string, truncate_usage_body_string, upsert_body_capture_metadata_value_entry, }; use aether_data_contracts::repository::usage::UsageBodyCaptureState; use aether_data_contracts::repository::usage::UsageBodyField; - use serde_json::{Map, Value}; + use serde_json::{json, Map, Value}; + + use crate::{ + UsageBodyCapturePolicy, UsageEvent, UsageEventData, UsageEventType, UsageRequestRecordLevel, + }; + + fn sensitive_oauth_event(api_format: &str) -> UsageEvent { + UsageEvent::new( + UsageEventType::Completed, + "oauth-sensitive-request", + UsageEventData { + provider_name: "oauth".to_string(), + model: "oauth-exchange".to_string(), + api_format: Some(api_format.to_string()), + endpoint_api_format: Some(api_format.to_string()), + request_body: Some(json!({"client_secret":"request-secret"})), + request_body_ref: Some("usage://oauth/request".to_string()), + provider_request_body: Some(json!({"refresh_token":"refresh-secret"})), + provider_request_body_ref: Some("usage://oauth/provider-request".to_string()), + response_body: Some(json!({ + "access_token":"access-secret", + "refresh_token":"rotated-refresh-secret" + })), + response_body_ref: Some("usage://oauth/response".to_string()), + client_response_body: Some(json!({"access_token":"client-access-secret"})), + client_response_body_ref: Some("usage://oauth/client-response".to_string()), + ..UsageEventData::default() + }, + ) + } + + #[allow(clippy::too_many_arguments)] + fn assert_sensitive_bodies_disabled( + request_body: &Option, + request_body_ref: &Option, + request_body_state: Option, + provider_request_body: &Option, + provider_request_body_ref: &Option, + provider_request_body_state: Option, + response_body: &Option, + response_body_ref: &Option, + response_body_state: Option, + client_response_body: &Option, + client_response_body_ref: &Option, + client_response_body_state: Option, + request_metadata: &Option, + ) { + assert!(request_body.is_none()); + assert!(request_body_ref.is_none()); + assert_eq!(request_body_state, Some(UsageBodyCaptureState::Disabled)); + assert!(provider_request_body.is_none()); + assert!(provider_request_body_ref.is_none()); + assert_eq!( + provider_request_body_state, + Some(UsageBodyCaptureState::Disabled) + ); + assert!(response_body.is_none()); + assert!(response_body_ref.is_none()); + assert_eq!(response_body_state, Some(UsageBodyCaptureState::Disabled)); + assert!(client_response_body.is_none()); + assert!(client_response_body_ref.is_none()); + assert_eq!( + client_response_body_state, + Some(UsageBodyCaptureState::Disabled) + ); + let body_capture = request_metadata + .as_ref() + .and_then(|metadata| metadata.get("body_capture")) + .and_then(Value::as_object) + .expect("body capture metadata should exist"); + for field in ["request", "provider_request", "response", "client_response"] { + assert_eq!( + body_capture + .get(field) + .and_then(|entry| entry.get("reason")) + .and_then(Value::as_str), + Some("sensitive_oauth_exchange") + ); + } + } #[test] fn build_plan_body_capture_metadata_returns_none_without_base64_body() { assert!(build_plan_body_capture_metadata(None).is_none()); } + #[test] + fn full_policy_never_captures_sensitive_oauth_event_bodies() { + let mut event = sensitive_oauth_event("oauth:exchange"); + + apply_usage_body_capture_policy_to_event( + UsageBodyCapturePolicy { + record_level: UsageRequestRecordLevel::Full, + }, + &mut event, + ); + + assert_sensitive_bodies_disabled( + &event.data.request_body, + &event.data.request_body_ref, + event.data.request_body_state, + &event.data.provider_request_body, + &event.data.provider_request_body_ref, + event.data.provider_request_body_state, + &event.data.response_body, + &event.data.response_body_ref, + event.data.response_body_state, + &event.data.client_response_body, + &event.data.client_response_body_ref, + event.data.client_response_body_state, + &event.data.request_metadata, + ); + } + + #[test] + fn full_policy_never_captures_sensitive_oauth_record_bodies() { + let event = sensitive_oauth_event("provider_oauth:exchange"); + let mut record = crate::record::build_upsert_usage_record_from_event(&event) + .expect("usage record should build"); + + apply_usage_body_capture_policy_to_record( + UsageBodyCapturePolicy { + record_level: UsageRequestRecordLevel::Full, + }, + &mut record, + ); + + assert_sensitive_bodies_disabled( + &record.request_body, + &record.request_body_ref, + record.request_body_state, + &record.provider_request_body, + &record.provider_request_body_ref, + record.provider_request_body_state, + &record.response_body, + &record.response_body_ref, + record.response_body_state, + &record.client_response_body, + &record.client_response_body_ref, + record.client_response_body_state, + &record.request_metadata, + ); + } + + #[test] + fn full_policy_never_captures_refresh_or_service_account_token_bodies() { + for api_format in [ + "provider_oauth:local_refresh", + "vertex_ai:service_account_token", + ] { + let mut event = sensitive_oauth_event(api_format); + + apply_usage_body_capture_policy_to_event( + UsageBodyCapturePolicy { + record_level: UsageRequestRecordLevel::Full, + }, + &mut event, + ); + + assert_sensitive_bodies_disabled( + &event.data.request_body, + &event.data.request_body_ref, + event.data.request_body_state, + &event.data.provider_request_body, + &event.data.provider_request_body_ref, + event.data.provider_request_body_state, + &event.data.response_body, + &event.data.response_body_ref, + event.data.response_body_state, + &event.data.client_response_body, + &event.data.client_response_body_ref, + event.data.client_response_body_state, + &event.data.request_metadata, + ); + } + } + #[test] fn trim_owned_non_empty_string_preserves_clean_values_and_drops_blank_ones() { assert_eq!( diff --git a/crates/aether-usage/runtime/src/lib.rs b/crates/aether-usage/runtime/src/lib.rs index 0be0b5b97..71f6a5e57 100644 --- a/crates/aether-usage/runtime/src/lib.rs +++ b/crates/aether-usage/runtime/src/lib.rs @@ -24,17 +24,18 @@ pub use event::{now_ms, UsageEvent, UsageEventData, UsageEventType, USAGE_EVENT_ pub use queue::UsageQueue; pub use record::build_upsert_usage_record_from_event; pub use report::{ - extract_gemini_file_mapping_entries, gemini_file_mapping_cache_key, - infer_internal_finalize_signature, is_local_ai_stream_report_kind, - is_local_ai_sync_report_kind, normalize_gemini_file_name, report_request_id, - resolve_internal_finalize_route, should_handle_local_stream_report, + decode_internal_report_body_base64, extract_gemini_file_mapping_entries, + gemini_file_mapping_cache_key, infer_internal_finalize_signature, + is_local_ai_stream_report_kind, is_local_ai_sync_report_kind, normalize_gemini_file_name, + report_request_id, resolve_internal_finalize_route, should_handle_local_stream_report, should_handle_local_sync_report, stream_capture_terminal_state, stream_report_missing_terminal_event, stream_report_represents_failure, stream_report_requires_observed_terminal_event, sync_report_represents_failure, GatewayStreamReportRequest, GatewaySyncReportRequest, GeminiFileMappingEntry, InternalFinalizeRoute, StreamCapturedTerminalState, GEMINI_FILE_MAPPING_TTL_SECONDS, - STREAM_MISSING_TERMINAL_EVENT_CATEGORY, STREAM_MISSING_TERMINAL_EVENT_MESSAGE, - STREAM_TERMINAL_ERROR_CATEGORY, STREAM_TERMINAL_ERROR_MESSAGE, + MAX_INTERNAL_REPORT_BODY_BYTES, STREAM_MISSING_TERMINAL_EVENT_CATEGORY, + STREAM_MISSING_TERMINAL_EVENT_MESSAGE, STREAM_TERMINAL_ERROR_CATEGORY, + STREAM_TERMINAL_ERROR_MESSAGE, }; pub use report_context::{ build_locally_actionable_report_context_from_request_candidate, @@ -46,7 +47,9 @@ pub use runtime::{ DEFAULT_USAGE_REQUEST_BODY_CAPTURE_LIMIT_BYTES, DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, }; -pub use settlement::{settle_usage_if_needed, UsageSettlementWriter}; +pub use settlement::{ + reconcile_usage_policy_cost_for_event, settle_usage_if_needed, UsageSettlementWriter, +}; pub use standardized_usage::StandardizedUsage; pub use usage_mapper::{map_usage, map_usage_from_response, UsageMapper}; pub use worker::{ @@ -61,7 +64,8 @@ pub use write::{ build_sync_terminal_usage_event, build_sync_terminal_usage_outcome, build_sync_terminal_usage_payload_seed, build_sync_terminal_usage_seed, build_terminal_usage_context_seed, build_terminal_usage_event_from_outcome, - build_terminal_usage_event_from_seed, build_usage_event_data_seed, LifecycleUsageSeed, + build_terminal_usage_event_from_seed, build_usage_event_data_seed, + build_usage_event_data_seed_describing_request_bodies, LifecycleUsageSeed, StreamTerminalUsagePayloadSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed, TerminalUsageOutcome, TerminalUsageSeed, UsageTerminalState, }; diff --git a/crates/aether-usage/runtime/src/record.rs b/crates/aether-usage/runtime/src/record.rs index 62b75acba..62a8feefd 100644 --- a/crates/aether-usage/runtime/src/record.rs +++ b/crates/aether-usage/runtime/src/record.rs @@ -608,7 +608,8 @@ mod tests { assert_eq!( record.request_metadata, Some(serde_json::json!({ - "billing_snapshot": { "status": "complete" } + "billing_snapshot": { "status": "complete" }, + "billing_snapshot_status": "complete" })) ); } diff --git a/crates/aether-usage/runtime/src/report.rs b/crates/aether-usage/runtime/src/report.rs index 9f7fca67b..d6591116d 100644 --- a/crates/aether-usage/runtime/src/report.rs +++ b/crates/aether-usage/runtime/src/report.rs @@ -1,12 +1,23 @@ use std::collections::BTreeMap; use aether_contracts::{ExecutionStreamTerminalSummary, ExecutionTelemetry}; -use aether_data_contracts::repository::usage::UsageBodyCaptureState; +use aether_data_contracts::repository::{ + gemini_file_mappings::{ + GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS, + GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS, + }, + usage::UsageBodyCaptureState, +}; use base64::Engine as _; use serde::{Deserialize, Serialize}; use serde_json::Value; pub const GEMINI_FILE_MAPPING_TTL_SECONDS: u64 = 60 * 60 * 48; +/// Maximum decoded size accepted for a body carried in an internal usage +/// report. Execution transports already cap normal response bodies at this +/// size; keeping the same ceiling here prevents a base64 field from causing a +/// second, unchecked allocation while preserving large image responses. +pub const MAX_INTERNAL_REPORT_BODY_BYTES: usize = 64 * 1024 * 1024; const GEMINI_FILE_MAPPING_CACHE_PREFIX: &str = "gemini_files:key"; pub const STREAM_MISSING_TERMINAL_EVENT_CATEGORY: &str = "stream_missing_terminal_event"; pub const STREAM_TERMINAL_ERROR_CATEGORY: &str = "stream_terminal_error"; @@ -62,6 +73,40 @@ pub struct GatewayStreamReportRequest { pub telemetry: Option, } +/// Decode an internal report body only after checking the decoded-size bound. +/// The report payload itself is JSON, so the base64 text may be larger than the +/// raw body by roughly one third. Checking the encoded length first avoids +/// asking the base64 engine to allocate for an attacker-controlled oversized +/// value. +pub fn decode_internal_report_body_base64(body_base64: &str) -> Result, String> { + if body_base64.is_empty() { + return Ok(Vec::new()); + } + + let max_encoded_len = MAX_INTERNAL_REPORT_BODY_BYTES + .checked_add(2) + .and_then(|value| value.checked_div(3)) + .and_then(|value| value.checked_mul(4)) + .unwrap_or(usize::MAX); + if body_base64.len() > max_encoded_len { + return Err(format!( + "internal report body exceeds {} decoded bytes", + MAX_INTERNAL_REPORT_BODY_BYTES + )); + } + + let bytes = base64::engine::general_purpose::STANDARD + .decode(body_base64) + .map_err(|error| error.to_string())?; + if bytes.len() > MAX_INTERNAL_REPORT_BODY_BYTES { + return Err(format!( + "internal report body exceeds {} decoded bytes", + MAX_INTERNAL_REPORT_BODY_BYTES + )); + } + Ok(bytes) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct InternalFinalizeRoute { pub public_path: &'static str, @@ -192,7 +237,12 @@ pub fn normalize_gemini_file_name(file_name: &str) -> Option { if file_name.is_empty() { return None; } - if file_name.starts_with("files/") { + let prefix_chars = usize::from(!file_name.starts_with("files/")) * "files/".len(); + let allowed_input_chars = GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS.checked_sub(prefix_chars)?; + if file_name.chars().nth(allowed_input_chars).is_some() { + return None; + } + if prefix_chars == 0 { Some(file_name.to_string()) } else { Some(format!("files/{file_name}")) @@ -537,9 +587,7 @@ fn stream_capture_terminal_state_from_base64( body_state: Option, ) -> Option { let body_base64 = body_base64?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .ok()?; + let bytes = decode_internal_report_body_base64(body_base64).ok()?; let state = if let Ok(value) = serde_json::from_slice::(&bytes) { stream_capture_terminal_state(&value) } else { @@ -766,6 +814,12 @@ fn maybe_push_gemini_file_mapping_entry( .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) + .filter(|value| { + value + .chars() + .nth(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS) + .is_none() + }) .map(ToOwned::to_owned), mime_type: object .get("mimeType") @@ -773,6 +827,12 @@ fn maybe_push_gemini_file_mapping_entry( .and_then(Value::as_str) .map(str::trim) .filter(|value| !value.is_empty()) + .filter(|value| { + value + .chars() + .nth(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS) + .is_none() + }) .map(ToOwned::to_owned), }); } @@ -789,9 +849,7 @@ fn extract_sync_report_body_json(payload: &GatewaySyncReportRequest) -> Option>, ) -> Option { - let mut metadata = Map::new(); - if let Some(context) = context { - copy_allowed_metadata_fields(context, &mut metadata); - } - (!metadata.is_empty()).then_some(Value::Object(metadata)) + context.and_then(project_usage_request_metadata_object) } pub(crate) fn merge_usage_request_metadata( base: Option, override_value: Option, ) -> Option { - let mut metadata = Map::new(); - if let Some(Value::Object(base)) = base.as_ref() { - copy_allowed_metadata_fields(base, &mut metadata); - } - if let Some(Value::Object(override_object)) = override_value.as_ref() { - copy_allowed_metadata_fields(override_object, &mut metadata); - } + let mut metadata = projected_metadata_object(base.as_ref()); + metadata.extend(projected_metadata_object(override_value.as_ref())); (!metadata.is_empty()).then_some(Value::Object(metadata)) } @@ -98,25 +79,11 @@ pub(crate) fn merge_usage_request_metadata_owned( base: Option, override_value: Option, ) -> Option { - let mut metadata = match base { - Some(Value::Object(base)) => base, - _ => Map::new(), - }; - if let Some(Value::Object(override_object)) = override_value { - move_allowed_metadata_fields(override_object, &mut metadata); - } - (!metadata.is_empty()).then_some(Value::Object(metadata)) + merge_usage_request_metadata(base, override_value) } pub(crate) fn sanitize_usage_request_metadata(value: Option) -> Option { - let Value::Object(object) = value? else { - return None; - }; - - let mut filtered = Map::new(); - move_allowed_metadata_fields(object, &mut filtered); - - (!filtered.is_empty()).then_some(Value::Object(filtered)) + project_usage_request_metadata(value) } pub(crate) fn retain_first_byte_request_metadata(value: Option) -> Option { @@ -128,18 +95,11 @@ pub(crate) fn retain_first_byte_request_metadata(value: Option) -> Option key.as_str(), "trace_id" | "client_ip" - | "user_agent" | "client_family" | "client_requested_stream" | "upstream_is_stream" - | "client_session_affinity" | "api_key_is_standalone" - | "websocket_mode" - | "websocket_transport" - | "usage_available" - | "usage_pricing_available" - | "live_session" - | "realtime_session" + | "plan_usage_reservation_token" | "request_path" | "request_query_string" | "request_path_and_query" @@ -150,19 +110,20 @@ pub(crate) fn retain_first_byte_request_metadata(value: Option) -> Option | "model_id" | "global_model_id" | "global_model_name" - | "proxy" ) }); (!metadata.is_empty()).then_some(Value::Object(metadata)) } pub(crate) fn sanitize_usage_request_metadata_ref(value: Option<&Value>) -> Option { - let object = value.and_then(Value::as_object)?; + project_usage_request_metadata_ref(value) +} - let mut filtered = Map::new(); - copy_allowed_metadata_fields(object, &mut filtered); - - (!filtered.is_empty()).then_some(Value::Object(filtered)) +fn projected_metadata_object(value: Option<&Value>) -> Map { + match project_usage_request_metadata_ref(value) { + Some(Value::Object(object)) => object, + _ => Map::new(), + } } pub(crate) fn attach_client_request_body_metadata( @@ -348,380 +309,6 @@ pub(crate) fn attach_provider_actual_service_tier_metadata( (!object.is_empty()).then_some(Value::Object(object)) } -fn copy_allowed_metadata_fields(source: &Map, target: &mut Map) { - copy_non_empty_string(source, target, "trace_id"); - copy_non_empty_string(source, target, "client_ip"); - copy_non_empty_string(source, target, "user_agent"); - copy_non_empty_string(source, target, "client_family"); - copy_bool(source, target, "client_requested_stream"); - copy_bool(source, target, UPSTREAM_IS_STREAM_KEY); - copy_non_null_value(source, target, "client_session_affinity"); - copy_bool(source, target, "api_key_is_standalone"); - copy_bool(source, target, WEBSOCKET_MODE_METADATA_KEY); - copy_non_empty_string(source, target, WEBSOCKET_TRANSPORT_METADATA_KEY); - copy_bool(source, target, USAGE_AVAILABLE_METADATA_KEY); - copy_bool(source, target, USAGE_PRICING_AVAILABLE_METADATA_KEY); - copy_non_null_value(source, target, LIVE_SESSION_METADATA_KEY); - copy_non_null_value(source, target, REALTIME_SESSION_METADATA_KEY); - copy_non_empty_string(source, target, "request_path"); - copy_non_empty_string(source, target, "request_query_string"); - copy_non_empty_string(source, target, "request_path_and_query"); - copy_non_empty_string(source, target, REQUESTED_REASONING_EFFORT_METADATA_KEY); - copy_non_empty_string(source, target, PROVIDER_REASONING_EFFORT_METADATA_KEY); - copy_non_empty_string(source, target, PROVIDER_SERVICE_TIER_METADATA_KEY); - copy_non_empty_string(source, target, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY); - copy_number(source, target, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY); - copy_number(source, target, "provider_request_body_base64_bytes"); - copy_number(source, target, "provider_response_body_base64_bytes"); - copy_number(source, target, "client_response_body_base64_bytes"); - copy_non_null_value(source, target, "body_size"); - copy_number(source, target, "client_response_status_code"); - copy_number(source, target, "end_to_end_time_ms"); - copy_number(source, target, "end_to_end_first_byte_time_ms"); - copy_bool(source, target, "transport_error"); - copy_non_empty_string(source, target, "transport_error_type"); - copy_non_null_value(source, target, "billing_snapshot"); - copy_non_empty_string(source, target, "billing_snapshot_schema_version"); - copy_non_empty_string(source, target, "billing_snapshot_status"); - copy_non_null_value(source, target, "settlement_snapshot"); - copy_non_empty_string(source, target, "settlement_snapshot_schema_version"); - copy_non_null_value(source, target, "billing_dimensions"); - copy_non_empty_string(source, target, "model_id"); - copy_non_empty_string(source, target, "global_model_id"); - copy_non_empty_string(source, target, "global_model_name"); - copy_non_null_value(source, target, "dimensions"); - copy_non_null_value(source, target, "billing_rule_snapshot"); - copy_non_null_value(source, target, "scheduling_audit"); - copy_non_empty_string(source, target, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY); - copy_non_null_value(source, target, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY); - copy_non_null_value(source, target, "tls_fingerprint"); - copy_number(source, target, "rate_multiplier"); - copy_bool(source, target, "is_free_tier"); - copy_number(source, target, "input_price_per_1m"); - copy_number(source, target, "output_price_per_1m"); - copy_number(source, target, "cache_creation_price_per_1m"); - copy_number(source, target, "cache_read_price_per_1m"); - copy_number(source, target, "price_per_request"); - copy_non_null_value(source, target, "proxy"); - copy_non_null_value(source, target, "stage_timings_ms"); - copy_non_null_value(source, target, "db_timings_ms"); - sanitize_request_path_metadata_fields(target); -} - -fn move_allowed_metadata_fields(mut source: Map, target: &mut Map) { - remove_non_empty_string(&mut source, target, "trace_id"); - remove_non_empty_string(&mut source, target, "client_ip"); - remove_non_empty_string(&mut source, target, "user_agent"); - remove_non_empty_string(&mut source, target, "client_family"); - remove_bool(&mut source, target, "client_requested_stream"); - remove_bool(&mut source, target, UPSTREAM_IS_STREAM_KEY); - remove_non_null_value(&mut source, target, "client_session_affinity"); - remove_bool(&mut source, target, "api_key_is_standalone"); - remove_bool(&mut source, target, WEBSOCKET_MODE_METADATA_KEY); - remove_non_empty_string(&mut source, target, WEBSOCKET_TRANSPORT_METADATA_KEY); - remove_bool(&mut source, target, USAGE_AVAILABLE_METADATA_KEY); - remove_bool(&mut source, target, USAGE_PRICING_AVAILABLE_METADATA_KEY); - remove_non_null_value(&mut source, target, LIVE_SESSION_METADATA_KEY); - remove_non_null_value(&mut source, target, REALTIME_SESSION_METADATA_KEY); - remove_non_empty_string(&mut source, target, "request_path"); - remove_non_empty_string(&mut source, target, "request_query_string"); - remove_non_empty_string(&mut source, target, "request_path_and_query"); - remove_non_empty_string(&mut source, target, REQUESTED_REASONING_EFFORT_METADATA_KEY); - remove_non_empty_string(&mut source, target, PROVIDER_REASONING_EFFORT_METADATA_KEY); - remove_non_empty_string(&mut source, target, PROVIDER_SERVICE_TIER_METADATA_KEY); - remove_non_empty_string( - &mut source, - target, - PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, - ); - remove_number(&mut source, target, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY); - remove_number(&mut source, target, "provider_request_body_base64_bytes"); - remove_number(&mut source, target, "provider_response_body_base64_bytes"); - remove_number(&mut source, target, "client_response_body_base64_bytes"); - remove_non_null_value(&mut source, target, "body_size"); - remove_number(&mut source, target, "client_response_status_code"); - remove_number(&mut source, target, "end_to_end_time_ms"); - remove_number(&mut source, target, "end_to_end_first_byte_time_ms"); - remove_bool(&mut source, target, "transport_error"); - remove_non_empty_string(&mut source, target, "transport_error_type"); - remove_non_null_value(&mut source, target, "billing_snapshot"); - remove_non_empty_string(&mut source, target, "billing_snapshot_schema_version"); - remove_non_empty_string(&mut source, target, "billing_snapshot_status"); - remove_non_null_value(&mut source, target, "settlement_snapshot"); - remove_non_empty_string(&mut source, target, "settlement_snapshot_schema_version"); - remove_non_null_value(&mut source, target, "billing_dimensions"); - remove_non_empty_string(&mut source, target, "model_id"); - remove_non_empty_string(&mut source, target, "global_model_id"); - remove_non_empty_string(&mut source, target, "global_model_name"); - remove_non_null_value(&mut source, target, "dimensions"); - remove_non_null_value(&mut source, target, "billing_rule_snapshot"); - remove_non_null_value(&mut source, target, "scheduling_audit"); - remove_non_empty_string( - &mut source, - target, - ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, - ); - remove_non_null_value(&mut source, target, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY); - remove_non_null_value(&mut source, target, "tls_fingerprint"); - remove_number(&mut source, target, "rate_multiplier"); - remove_bool(&mut source, target, "is_free_tier"); - remove_number(&mut source, target, "input_price_per_1m"); - remove_number(&mut source, target, "output_price_per_1m"); - remove_number(&mut source, target, "cache_creation_price_per_1m"); - remove_number(&mut source, target, "cache_read_price_per_1m"); - remove_number(&mut source, target, "price_per_request"); - remove_non_null_value(&mut source, target, "proxy"); - remove_non_null_value(&mut source, target, "stage_timings_ms"); - remove_non_null_value(&mut source, target, "db_timings_ms"); - sanitize_request_path_metadata_fields(target); -} - -fn sanitize_request_path_metadata_fields(target: &mut Map) { - let path = target - .get("request_path") - .and_then(Value::as_str) - .and_then(sanitize_request_path); - let query = target - .get("request_query_string") - .and_then(Value::as_str) - .and_then(sanitize_request_query_string); - let path_and_query = target - .get("request_path_and_query") - .and_then(Value::as_str) - .and_then(|value| sanitize_request_path_and_query(value, None)) - .or_else(|| { - path.as_deref() - .and_then(|path| sanitize_request_path_and_query(path, query.as_deref())) - }); - - apply_optional_string_field(target, "request_path", path.as_deref()); - apply_optional_string_field(target, "request_query_string", query.as_deref()); - apply_optional_string_field(target, "request_path_and_query", path_and_query.as_deref()); -} - -fn apply_optional_string_field(target: &mut Map, key: &str, value: Option<&str>) { - if let Some(value) = value { - target.insert(key.to_string(), Value::String(value.to_string())); - } else { - target.remove(key); - } -} - -fn copy_non_empty_string(source: &Map, target: &mut Map, key: &str) { - let Some(value) = source - .get(key) - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - else { - return; - }; - target.insert( - key.to_string(), - Value::String(truncate_usage_request_metadata_string(value)), - ); -} - -fn remove_non_empty_string( - source: &mut Map, - target: &mut Map, - key: &str, -) { - let Some(Value::String(value)) = source.remove(key) else { - return; - }; - let Some(value) = trim_and_truncate_usage_request_metadata_string_owned(value) else { - return; - }; - target.insert(key.to_string(), Value::String(value)); -} - -fn copy_number(source: &Map, target: &mut Map, key: &str) { - let Some(value) = source.get(key).filter(|value| value.is_number()) else { - return; - }; - target.insert(key.to_string(), value.clone()); -} - -fn remove_number(source: &mut Map, target: &mut Map, key: &str) { - let Some(value) = source.remove(key).filter(|value| value.is_number()) else { - return; - }; - target.insert(key.to_string(), value); -} - -fn copy_bool(source: &Map, target: &mut Map, key: &str) { - let Some(value) = source.get(key).filter(|value| value.is_boolean()) else { - return; - }; - target.insert(key.to_string(), value.clone()); -} - -fn remove_bool(source: &mut Map, target: &mut Map, key: &str) { - let Some(value) = source.remove(key).filter(|value| value.is_boolean()) else { - return; - }; - target.insert(key.to_string(), value); -} - -fn copy_non_null_value(source: &Map, target: &mut Map, key: &str) { - let Some(value) = source.get(key).filter(|value| !value.is_null()) else { - return; - }; - target.insert( - key.to_string(), - sanitize_usage_request_metadata_value(value), - ); -} - -fn remove_non_null_value( - source: &mut Map, - target: &mut Map, - key: &str, -) { - let Some(value) = source.remove(key).filter(|value| !value.is_null()) else { - return; - }; - target.insert( - key.to_string(), - sanitize_usage_request_metadata_value_owned(value), - ); -} - -fn sanitize_usage_request_metadata_value(value: &Value) -> Value { - match value { - Value::String(text) => Value::String(truncate_usage_request_metadata_string(text)), - _ if usage_request_metadata_within_limits(value) => value.clone(), - _ => truncated_usage_request_metadata_value(value), - } -} - -fn sanitize_usage_request_metadata_value_owned(value: Value) -> Value { - match value { - Value::String(text) => Value::String(truncate_usage_request_metadata_string_owned(text)), - _ if usage_request_metadata_within_limits(&value) => value, - _ => truncated_usage_request_metadata_value(&value), - } -} - -fn truncate_usage_request_metadata_string(value: &str) -> String { - const TRUNCATED_SUFFIX: &str = "...[truncated]"; - - if value.len() <= MAX_USAGE_REQUEST_METADATA_STRING_BYTES { - return value.to_string(); - } - - let target_bytes = - MAX_USAGE_REQUEST_METADATA_STRING_BYTES.saturating_sub(TRUNCATED_SUFFIX.len()); - let mut end = 0usize; - for (idx, ch) in value.char_indices() { - let next = idx + ch.len_utf8(); - if next > target_bytes { - break; - } - end = next; - } - - if end == 0 { - return TRUNCATED_SUFFIX.to_string(); - } - - format!("{}{TRUNCATED_SUFFIX}", &value[..end]) -} - -fn trim_and_truncate_usage_request_metadata_string_owned(value: String) -> Option { - let trimmed = value.trim(); - if trimmed.is_empty() { - return None; - } - if trimmed.len() == value.len() { - return Some(truncate_usage_request_metadata_string_owned(value)); - } - Some(truncate_usage_request_metadata_string(trimmed)) -} - -fn truncate_usage_request_metadata_string_owned(value: String) -> String { - if value.len() <= MAX_USAGE_REQUEST_METADATA_STRING_BYTES { - return value; - } - truncate_usage_request_metadata_string(value.as_str()) -} - -fn truncated_usage_request_metadata_value(value: &Value) -> Value { - json!({ - "truncated": true, - "reason": "usage_request_metadata_limits_exceeded", - "max_depth": MAX_USAGE_REQUEST_METADATA_DEPTH, - "max_nodes": MAX_USAGE_REQUEST_METADATA_NODES, - "max_bytes": MAX_USAGE_REQUEST_METADATA_BYTES, - "value_kind": usage_request_metadata_value_kind(value), - }) -} - -fn usage_request_metadata_within_limits(value: &Value) -> bool { - let mut nodes = 0usize; - let mut estimated_bytes = 0usize; - let mut stack = vec![(value, 1usize)]; - - while let Some((current, depth)) = stack.pop() { - nodes = nodes.saturating_add(1); - estimated_bytes = - estimated_bytes.saturating_add(usage_request_metadata_value_size_hint(current)); - if depth > MAX_USAGE_REQUEST_METADATA_DEPTH - || nodes > MAX_USAGE_REQUEST_METADATA_NODES - || estimated_bytes > MAX_USAGE_REQUEST_METADATA_BYTES - { - return false; - } - match current { - Value::Array(items) => { - estimated_bytes = estimated_bytes.saturating_add(items.len().saturating_mul(2)); - for item in items.iter().rev() { - stack.push((item, depth + 1)); - } - } - Value::Object(object) => { - estimated_bytes = estimated_bytes - .saturating_add(object.len().saturating_mul(3)) - .saturating_add( - object - .keys() - .map(|key| key.len().saturating_add(2)) - .sum::(), - ); - for item in object.values() { - stack.push((item, depth + 1)); - } - } - _ => {} - } - } - - true -} - -fn usage_request_metadata_value_kind(value: &Value) -> &'static str { - match value { - Value::Null => "null", - Value::Bool(_) => "bool", - Value::Number(_) => "number", - Value::String(_) => "string", - Value::Array(_) => "array", - Value::Object(_) => "object", - } -} - -fn usage_request_metadata_value_size_hint(value: &Value) -> usize { - match value { - Value::Null => 4, - Value::Bool(false) => 5, - Value::Bool(true) => 4, - Value::Number(number) => number.to_string().len(), - Value::String(text) => text.len().saturating_add(2), - Value::Array(_) | Value::Object(_) => 2, - } -} - #[cfg(test)] mod tests { use aether_contracts::{ExecutionPlan, RequestBody}; @@ -739,8 +326,7 @@ mod tests { build_usage_request_metadata_seed, merge_usage_request_metadata, merge_usage_request_metadata_owned, refresh_provider_response_body_metadata, retain_first_byte_request_metadata, sanitize_usage_request_metadata, - sanitize_usage_request_metadata_ref, MAX_USAGE_REQUEST_METADATA_BYTES, - MAX_USAGE_REQUEST_METADATA_DEPTH, MAX_USAGE_REQUEST_METADATA_NODES, + sanitize_usage_request_metadata_ref, }; fn sample_plan() -> ExecutionPlan { @@ -767,38 +353,8 @@ mod tests { } } - fn sample_stage_timings_metadata() -> Value { - json!({ - "stream_candidate_slot": 1, - "stream_provider_in_flight": 2, - "stream_upstream_headers": 180, - "stream_first_data": 8210 - }) - } - - fn sample_db_timings_metadata() -> Value { - json!({ - "query_count": 2, - "query_total": 950, - "query_max": 650, - "operations": { - "request_candidate_upsert": {"count": 1, "sum": 650, "max": 650}, - "usage_upsert": {"count": 1, "sum": 300, "max": 300} - }, - "pool": { - "max_checked_out": 20, - "max_pool_size": 20, - "min_idle": 0, - "max_connections": 20, - "max_usage_rate": 100.0 - } - }) - } - #[test] fn sanitizes_request_metadata_to_allowlist() { - let stage_timings_ms = sample_stage_timings_metadata(); - let db_timings_ms = sample_db_timings_metadata(); let metadata = sanitize_usage_request_metadata(Some(json!({ "request_id": "req-1", "provider_id": "provider-1", @@ -844,8 +400,8 @@ mod tests { "cache_creation_price_per_1m": 3.75, "cache_read_price_per_1m": 0.3, "price_per_request": 0.02, - "stage_timings_ms": stage_timings_ms.clone(), - "db_timings_ms": db_timings_ms.clone(), + "stage_timings_ms": {"planning": 12}, + "db_timings_ms": {"query": "SELECT credential"}, "original_headers": {"authorization": "Bearer secret"}, "original_request_body": {"messages": []}, "provider_request_headers": {"authorization": "Bearer secret"}, @@ -858,7 +414,7 @@ mod tests { json!({ "trace_id": "trace-1", "client_ip": "203.0.113.8", - "user_agent": "Claude-Code/1.0", + "client_family": "claude_code", "client_requested_stream": false, "upstream_is_stream": true, "api_key_is_standalone": true, @@ -884,9 +440,7 @@ mod tests { "routing_candidate_skip_reason": "provider_request_body_build_failed", "routing_failure_diagnostic": { "kind": "request_body_build", - "path": "$.reasoning.summary", - "message": "invalid reasoning summary", - "safe_to_show": true + "path": "$.reasoning.summary" }, "rate_multiplier": 1.25, "is_free_tier": false, @@ -894,9 +448,7 @@ mod tests { "output_price_per_1m": 15.0, "cache_creation_price_per_1m": 3.75, "cache_read_price_per_1m": 0.3, - "price_per_request": 0.02, - "stage_timings_ms": stage_timings_ms, - "db_timings_ms": db_timings_ms + "price_per_request": 0.02 }) ); } @@ -922,8 +474,7 @@ mod tests { "client_ip": "203.0.113.8", "request_path": "/v1/chat/completions", "request_path_and_query": "/v1/chat/completions", - "upstream_is_stream": true, - "proxy": {"mode": "manual", "node_id": "proxy-1"} + "upstream_is_stream": true }) ); } @@ -946,6 +497,20 @@ mod tests { ); } + #[test] + fn sanitizes_plan_usage_reservation_deferred_as_a_boolean() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "plan_usage_reservation_deferred": true, + }))) + .expect("deferred marker should remain"); + assert_eq!(metadata, json!({"plan_usage_reservation_deferred": true})); + + assert!(sanitize_usage_request_metadata(Some(json!({ + "plan_usage_reservation_deferred": "true", + }))) + .is_none()); + } + #[test] fn sanitizes_request_path_query_metadata() { let metadata = sanitize_usage_request_metadata(Some(json!({ @@ -966,35 +531,19 @@ mod tests { } #[test] - fn sanitizes_large_allowed_metadata_values_to_bounded_representations() { - let metadata = sanitize_usage_request_metadata(Some(json!({ + fn rejects_oversized_tokens_and_unknown_nested_objects() { + assert!(sanitize_usage_request_metadata(Some(json!({ "trace_id": "t".repeat(2_048), "billing_snapshot": { "payload": "x".repeat(32 * 1024) } }))) - .expect("metadata should remain"); - - assert!(metadata - .get("trace_id") - .and_then(Value::as_str) - .is_some_and(|value| value.ends_with("...[truncated]"))); - assert_eq!( - metadata.get("billing_snapshot"), - Some(&json!({ - "truncated": true, - "reason": "usage_request_metadata_limits_exceeded", - "max_depth": MAX_USAGE_REQUEST_METADATA_DEPTH, - "max_nodes": MAX_USAGE_REQUEST_METADATA_NODES, - "max_bytes": MAX_USAGE_REQUEST_METADATA_BYTES, - "value_kind": "object", - })) - ); + .is_none()); } #[test] - fn sanitizes_request_metadata_preserves_tls_fingerprint() { - let metadata = sanitize_usage_request_metadata(Some(json!({ + fn sanitizes_request_metadata_drops_tls_fingerprint() { + assert!(sanitize_usage_request_metadata(Some(json!({ "tls_fingerprint": { "incoming": { "source": "forwarded_header", @@ -1011,25 +560,7 @@ mod tests { "ja3": "spoofed" } }))) - .expect("metadata should remain"); - - assert_eq!( - metadata, - json!({ - "tls_fingerprint": { - "incoming": { - "source": "forwarded_header", - "ja3": "incoming-ja3", - "ja4": "incoming-ja4" - }, - "outgoing": { - "source": "aether_transport_config", - "backend": "reqwest_rustls", - "observed": false - } - } - }) - ); + .is_none()); } #[test] @@ -1052,6 +583,7 @@ mod tests { "global_model_id": "global-model-1", "global_model_name": "gpt-5", "client_ip": "203.0.113.8", + "client_family": "claude_code", "user_agent": "Claude-Code/1.0", "billing_snapshot": {"status": "complete"}, "stage_timings_ms": { @@ -1088,21 +620,9 @@ mod tests { "global_model_id": "global-model-1", "global_model_name": "gpt-5", "client_ip": "203.0.113.8", - "user_agent": "Claude-Code/1.0", + "client_family": "claude_code", "billing_snapshot": {"status": "complete"}, - "stage_timings_ms": { - "stream_candidate_slot": 0, - "stream_upstream_headers": 180, - "stream_first_data": 8210 - }, - "db_timings_ms": { - "query_count": 1, - "query_total": 42, - "query_max": 42, - "operations": { - "auth_api_key_snapshot": {"count": 1, "sum": 42, "max": 42} - } - } + "billing_snapshot_status": "complete" }) ); } diff --git a/crates/aether-usage/runtime/src/runtime.rs b/crates/aether-usage/runtime/src/runtime.rs index 0c2140096..5f26f893c 100644 --- a/crates/aether-usage/runtime/src/runtime.rs +++ b/crates/aether-usage/runtime/src/runtime.rs @@ -27,9 +27,10 @@ use crate::worker::{ use crate::{ apply_usage_body_capture_policy_to_event, build_stream_terminal_usage_seed, build_sync_terminal_usage_seed, build_terminal_usage_event_from_seed, - build_upsert_usage_record_from_event, settle_usage_if_needed, LifecycleUsageSeed, - StreamTerminalUsagePayloadSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed, - UsageEvent, UsageQueue, UsageRecordWriter, UsageRuntimeConfig, UsageSettlementWriter, + build_upsert_usage_record_from_event, reconcile_usage_policy_cost_for_event, + settle_usage_if_needed, LifecycleUsageSeed, StreamTerminalUsagePayloadSeed, + SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed, UsageEvent, UsageQueue, + UsageRecordWriter, UsageRuntimeConfig, UsageSettlementWriter, }; #[async_trait] @@ -39,8 +40,8 @@ pub trait UsageBillingEventEnricher: Send + Sync { #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub enum UsageRequestRecordLevel { - Basic, #[default] + Basic, Full, } @@ -55,7 +56,7 @@ pub struct UsageBodyCapturePolicy { impl Default for UsageBodyCapturePolicy { fn default() -> Self { Self { - record_level: UsageRequestRecordLevel::Full, + record_level: UsageRequestRecordLevel::Basic, } } } @@ -4130,9 +4131,9 @@ impl UsageRuntime { event_name = "usage_body_capture_policy_read_failed", log_type = "event", request_id = %event.request_id, - fallback = "default", + fallback = "basic", error = %err, - "usage runtime failed to read body capture policy; keeping default capture" + "usage runtime failed to read body capture policy; disabling body capture" ); apply_usage_body_capture_policy_to_event(UsageBodyCapturePolicy::default(), event); } @@ -4765,6 +4766,17 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { + if let Err(err) = reconcile_usage_policy_cost_for_event(data, event).await { + warn!( + event_name = "usage_event_cost_reconciliation_failed", + log_type = "event", + usage_event_type = ?event.event_type, + request_id = %event.request_id, + error = %err, + "usage runtime failed to reconcile plan cost before direct usage upsert" + ); + return false; + } match build_upsert_usage_record_from_event(event) { Ok(record) => match catch_usage_writer_panic( "direct usage upsert", diff --git a/crates/aether-usage/runtime/src/settlement.rs b/crates/aether-usage/runtime/src/settlement.rs index 8bf7c21b8..9d57610ea 100644 --- a/crates/aether-usage/runtime/src/settlement.rs +++ b/crates/aether-usage/runtime/src/settlement.rs @@ -1,36 +1,144 @@ use std::sync::{Arc, OnceLock}; -use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput}; +use aether_data_contracts::repository::billing::nonnegative_usd_to_usage_policy_cost_units; +use aether_data_contracts::repository::settlement::{ + ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement, + UsagePolicyCostReservationState, UsageSettlementInput, +}; use aether_data_contracts::repository::usage::StoredRequestUsageAudit; +use aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY; use aether_data_contracts::{DataLayerError, DataLayerError::InvalidInput}; use async_trait::async_trait; +use crate::event::{UsageEvent, UsageEventType}; use crate::keyed_lock::KeyedAsyncLockPool; #[async_trait] pub trait UsageSettlementWriter: Send + Sync { fn has_usage_settlement_writer(&self) -> bool; + async fn reconcile_usage_policy_cost( + &self, + _input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + Ok(None) + } + async fn settle_usage( &self, input: UsageSettlementInput, ) -> Result, DataLayerError>; } +pub async fn reconcile_usage_policy_cost_for_event( + writer: &dyn UsageSettlementWriter, + event: &UsageEvent, +) -> Result<(), DataLayerError> { + if !writer.has_usage_settlement_writer() { + return Ok(()); + } + let terminal_state = match event.event_type { + UsageEventType::Completed => UsagePolicyCostReservationState::Finalized, + UsageEventType::Failed | UsageEventType::Cancelled => { + UsagePolicyCostReservationState::Released + } + UsageEventType::Pending | UsageEventType::Streaming => return Ok(()), + }; + if plan_usage_reservation_reconciliation_is_deferred(event.data.request_metadata.as_ref()) { + return Ok(()); + } + let Some(subject_id) = event.data.user_id.as_deref().and_then(non_empty_trimmed) else { + return Ok(()); + }; + let Some(reservation_token) = event_usage_policy_reservation_token(event) else { + return Ok(()); + }; + let actual_cost_units = if terminal_state == UsagePolicyCostReservationState::Finalized { + let actual_cost_usd = event.data.actual_total_cost_usd.ok_or_else(|| { + InvalidInput( + "completed usage event with a plan reservation token is missing actual cost" + .to_string(), + ) + })?; + nonnegative_usd_to_usage_policy_cost_units(finite_cost(actual_cost_usd)?.max(0.0)) + .ok_or_else(|| { + InvalidInput("usage policy settlement cost exceeds the supported range".to_string()) + })? + } else { + 0 + }; + + let _ = writer + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: event.request_id.clone(), + subject_id: subject_id.to_string(), + reservation_token: reservation_token.to_string(), + actual_cost_units, + terminal_state, + finalized_at_unix_secs: event.timestamp_ms / 1_000, + }) + .await?; + Ok(()) +} + pub async fn settle_usage_if_needed( writer: &dyn UsageSettlementWriter, usage: &StoredRequestUsageAudit, ) -> Result<(), DataLayerError> { - if !writer.has_usage_settlement_writer() || usage.billing_status != "pending" { + if !writer.has_usage_settlement_writer() { return Ok(()); } - if !matches!(usage.status.as_str(), "completed" | "failed") { + if !matches!(usage.status.as_str(), "completed" | "failed" | "cancelled") { return Ok(()); } let finalized_at_unix_secs = usage .finalized_at_unix_secs .or(Some(usage.updated_at_unix_secs)); + let settlement_key = usage_settlement_lock_key_for_usage(usage); + let settlement_lock = usage_settlement_lock(&settlement_key); + let _guard = settlement_lock.lock().await; + + // Cost reservations are tied to a server-issued per-request token. Legacy usage rows do not + // have that token, so they must continue through wallet settlement without touching a cost + // reservation selected only by the client-visible request id. + if !plan_usage_reservation_reconciliation_is_deferred(usage.request_metadata.as_ref()) { + if let (Some(subject_id), Some(reservation_token)) = ( + usage.user_id.as_deref().and_then(non_empty_trimmed), + usage_policy_reservation_token(usage), + ) { + let (terminal_state, actual_cost_units) = if usage.status == "completed" { + ( + UsagePolicyCostReservationState::Finalized, + nonnegative_usd_to_usage_policy_cost_units( + finite_cost(usage.actual_total_cost_usd)?.max(0.0), + ) + .ok_or_else(|| { + InvalidInput( + "usage policy settlement cost exceeds the supported range".to_string(), + ) + })?, + ) + } else { + (UsagePolicyCostReservationState::Released, 0) + }; + let _ = writer + .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { + request_id: usage.request_id.clone(), + subject_id: subject_id.to_string(), + reservation_token: reservation_token.to_string(), + actual_cost_units, + terminal_state, + finalized_at_unix_secs: finalized_at_unix_secs + .unwrap_or(usage.updated_at_unix_secs), + }) + .await?; + } + } + + if usage.status == "cancelled" || usage.billing_status != "pending" { + return Ok(()); + } let input = UsageSettlementInput { request_id: usage.request_id.clone(), user_id: usage.user_id.clone(), @@ -43,31 +151,36 @@ pub async fn settle_usage_if_needed( actual_total_cost_usd: finite_cost(usage.actual_total_cost_usd)?, finalized_at_unix_secs, }; - let settlement_key = usage_settlement_lock_key(&input); - let settlement_lock = usage_settlement_lock(&settlement_key); - let _guard = settlement_lock.lock().await; let _ = writer.settle_usage(input).await?; Ok(()) } +fn plan_usage_reservation_reconciliation_is_deferred(metadata: Option<&serde_json::Value>) -> bool { + metadata + .and_then(serde_json::Value::as_object) + .and_then(|metadata| metadata.get(PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY)) + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) +} + fn usage_settlement_lock(key: &str) -> Arc> { static LOCKS: OnceLock = OnceLock::new(); LOCKS.get_or_init(KeyedAsyncLockPool::default).lock_for(key) } -fn usage_settlement_lock_key(input: &UsageSettlementInput) -> String { - if input.api_key_is_standalone { - if let Some(api_key_id) = input.api_key_id.as_deref().and_then(non_empty_trimmed) { +fn usage_settlement_lock_key_for_usage(usage: &StoredRequestUsageAudit) -> String { + if usage_api_key_is_standalone(usage) { + if let Some(api_key_id) = usage.api_key_id.as_deref().and_then(non_empty_trimmed) { return format!("api-key:{api_key_id}"); } } - if let Some(user_id) = input.user_id.as_deref().and_then(non_empty_trimmed) { + if let Some(user_id) = usage.user_id.as_deref().and_then(non_empty_trimmed) { return format!("user:{user_id}"); } - if let Some(api_key_id) = input.api_key_id.as_deref().and_then(non_empty_trimmed) { + if let Some(api_key_id) = usage.api_key_id.as_deref().and_then(non_empty_trimmed) { return format!("api-key:{api_key_id}"); } - format!("request:{}", input.request_id.trim()) + format!("request:{}", usage.request_id.trim()) } fn non_empty_trimmed(value: &str) -> Option<&str> { @@ -84,6 +197,27 @@ fn usage_api_key_is_standalone(usage: &StoredRequestUsageAudit) -> bool { .unwrap_or(false) } +fn usage_policy_reservation_token(usage: &StoredRequestUsageAudit) -> Option<&str> { + usage + .request_metadata + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|metadata| metadata.get("plan_usage_reservation_token")) + .and_then(serde_json::Value::as_str) + .and_then(non_empty_trimmed) +} + +fn event_usage_policy_reservation_token(event: &UsageEvent) -> Option<&str> { + event + .data + .request_metadata + .as_ref() + .and_then(serde_json::Value::as_object) + .and_then(|metadata| metadata.get("plan_usage_reservation_token")) + .and_then(serde_json::Value::as_str) + .and_then(non_empty_trimmed) +} + fn finite_cost(value: f64) -> Result { if value.is_finite() { Ok(value) @@ -100,16 +234,24 @@ mod tests { use std::sync::Mutex; use std::time::Duration; - use super::{settle_usage_if_needed, UsageSettlementWriter}; - use aether_data_contracts::repository::settlement::UsageSettlementInput; + use super::{ + reconcile_usage_policy_cost_for_event, settle_usage_if_needed, UsageSettlementWriter, + }; + use aether_data_contracts::repository::settlement::{ + ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, + UsagePolicyCostReservationState, UsageSettlementInput, + }; use aether_data_contracts::repository::usage::StoredRequestUsageAudit; use async_trait::async_trait; use serde_json::json; + use crate::{UsageEvent, UsageEventData, UsageEventType}; + #[derive(Default)] struct TestSettlementWriter { has_writer: bool, inputs: Mutex>, + reconciliations: Mutex>, } #[derive(Default)] @@ -125,6 +267,18 @@ mod tests { self.has_writer } + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, aether_data_contracts::DataLayerError> + { + self.reconciliations + .lock() + .expect("reconciliation inputs lock") + .push(input); + Ok(None) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -166,7 +320,7 @@ mod tests { } fn sample_usage() -> StoredRequestUsageAudit { - StoredRequestUsageAudit::new( + let mut usage = StoredRequestUsageAudit::new( "usage-1".to_string(), "req-1".to_string(), Some("user-1".to_string()), @@ -204,7 +358,11 @@ mod tests { 200, None, ) - .expect("usage should build") + .expect("usage should build"); + usage.request_metadata = Some(json!({ + "plan_usage_reservation_token": "token-1" + })); + usage } #[tokio::test] @@ -228,10 +386,22 @@ mod tests { assert_eq!(inputs[0].total_cost_usd, 1.25); assert_eq!(inputs[0].actual_total_cost_usd, 0.75); assert!(!inputs[0].api_key_is_standalone); + drop(inputs); + let reconciliations = writer + .reconciliations + .lock() + .expect("reconciliation inputs lock"); + assert_eq!(reconciliations.len(), 1); + assert_eq!(reconciliations[0].actual_cost_units, 75_000_000); + assert_eq!(reconciliations[0].reservation_token, "token-1"); + assert_eq!( + reconciliations[0].terminal_state, + UsagePolicyCostReservationState::Finalized + ); } #[tokio::test] - async fn skips_pending_cancelled_usage() { + async fn releases_pending_cancelled_usage_without_wallet_settlement() { let writer = TestSettlementWriter { has_writer: true, ..Default::default() @@ -246,6 +416,207 @@ mod tests { let inputs = writer.inputs.lock().expect("settlement inputs lock"); assert!(inputs.is_empty()); + drop(inputs); + let reconciliations = writer + .reconciliations + .lock() + .expect("reconciliation inputs lock"); + assert_eq!(reconciliations.len(), 1); + assert_eq!(reconciliations[0].actual_cost_units, 0); + assert_eq!( + reconciliations[0].terminal_state, + UsagePolicyCostReservationState::Released + ); + } + + #[tokio::test] + async fn releases_failed_usage_before_void_settlement() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + usage.status = "failed".to_string(); + + settle_usage_if_needed(&writer, &usage) + .await + .expect("failed usage should settle"); + + assert_eq!( + writer.inputs.lock().expect("settlement inputs lock").len(), + 1 + ); + let reconciliations = writer + .reconciliations + .lock() + .expect("reconciliation inputs lock"); + assert_eq!( + reconciliations[0].terminal_state, + UsagePolicyCostReservationState::Released + ); + assert_eq!(reconciliations[0].actual_cost_units, 0); + } + + #[tokio::test] + async fn skips_cost_reconciliation_for_legacy_usage_without_token() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + usage.request_metadata = None; + + settle_usage_if_needed(&writer, &usage) + .await + .expect("legacy usage should still settle its wallet charge"); + + assert_eq!( + writer + .reconciliations + .lock() + .expect("reconciliation lock") + .len(), + 0 + ); + assert_eq!( + writer.inputs.lock().expect("settlement inputs lock").len(), + 1 + ); + } + + #[tokio::test] + async fn ignores_blank_reservation_token_metadata() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + usage.request_metadata = Some(json!({ + "plan_usage_reservation_token": " " + })); + + settle_usage_if_needed(&writer, &usage) + .await + .expect("blank token should be treated as legacy usage"); + + assert!(writer + .reconciliations + .lock() + .expect("reconciliation lock") + .is_empty()); + assert_eq!( + writer.inputs.lock().expect("settlement inputs lock").len(), + 1 + ); + } + + #[tokio::test] + async fn event_reconciliation_requires_enriched_cost_and_preserves_server_token() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut event = UsageEvent::new( + UsageEventType::Completed, + "shared-trace", + UsageEventData { + user_id: Some("user-1".to_string()), + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + request_metadata: Some(json!({ + "plan_usage_reservation_token": "server-token" + })), + ..UsageEventData::default() + }, + ); + assert!(matches!( + reconcile_usage_policy_cost_for_event(&writer, &event).await, + Err(aether_data_contracts::DataLayerError::InvalidInput(_)) + )); + assert!(writer + .reconciliations + .lock() + .expect("reconciliations lock") + .is_empty()); + + event.data.actual_total_cost_usd = Some(1.25); + reconcile_usage_policy_cost_for_event(&writer, &event) + .await + .expect("enriched terminal event should reconcile"); + let reconciliations = writer.reconciliations.lock().expect("reconciliations lock"); + assert_eq!(reconciliations.len(), 1); + assert_eq!(reconciliations[0].request_id, "shared-trace"); + assert_eq!(reconciliations[0].reservation_token, "server-token"); + assert_eq!(reconciliations[0].actual_cost_units, 125_000_000); + } + + #[tokio::test] + async fn deferred_event_keeps_cost_reservation_without_requiring_actual_cost() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let event = UsageEvent::new( + UsageEventType::Completed, + "possibly-sent-request", + UsageEventData { + user_id: Some("user-1".to_string()), + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + request_metadata: Some(json!({ + "plan_usage_reservation_token": "server-token", + "plan_usage_reservation_deferred": true + })), + ..UsageEventData::default() + }, + ); + + reconcile_usage_policy_cost_for_event(&writer, &event) + .await + .expect("deferred reconciliation should not require unknown actual cost"); + + assert!(writer + .reconciliations + .lock() + .expect("reconciliations lock") + .is_empty()); + } + + #[tokio::test] + async fn deferred_stored_usage_skips_cost_reconcile_but_still_settles_wallet() { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + usage.request_metadata = Some(json!({ + "plan_usage_reservation_token": "server-token", + "plan_usage_reservation_deferred": true + })); + + settle_usage_if_needed(&writer, &usage) + .await + .expect("wallet settlement should continue"); + + assert!(writer + .reconciliations + .lock() + .expect("reconciliations lock") + .is_empty()); + assert_eq!( + writer.inputs.lock().expect("settlement inputs lock").len(), + 1 + ); + } + + #[test] + fn deferred_metadata_requires_a_boolean_true() { + assert!(super::plan_usage_reservation_reconciliation_is_deferred( + Some(&json!({"plan_usage_reservation_deferred": true})) + )); + assert!(!super::plan_usage_reservation_reconciliation_is_deferred( + Some(&json!({"plan_usage_reservation_deferred": "true"})) + )); } #[tokio::test] diff --git a/crates/aether-usage/runtime/src/worker.rs b/crates/aether-usage/runtime/src/worker.rs index f4be6fe48..d894cefb9 100644 --- a/crates/aether-usage/runtime/src/worker.rs +++ b/crates/aether-usage/runtime/src/worker.rs @@ -15,8 +15,9 @@ use crate::runtime::{ UsageBillingEventEnricher, UsageRuntimeAccess, UsageWorkerRecordConcurrencyGate, }; use crate::{ - build_upsert_usage_record_from_event, settle_usage_if_needed, UsageEvent, UsageEventType, - UsageQueue, UsageRuntimeConfig, UsageSettlementWriter, + build_upsert_usage_record_from_event, reconcile_usage_policy_cost_for_event, + settle_usage_if_needed, UsageEvent, UsageEventType, UsageQueue, UsageRuntimeConfig, + UsageSettlementWriter, }; const USAGE_WORKER_DB_PRESSURE_DEFER_MS: u64 = 10; @@ -676,6 +677,7 @@ pub async fn write_event_record(data: &T, event: &UsageEvent) -> Result<(), D where T: UsageRecordWriter + UsageSettlementWriter + Send + Sync, { + reconcile_usage_policy_cost_for_event(data, event).await?; let record = build_upsert_usage_record_from_event(event)?; if let Some(stored) = data.upsert_usage_record(record).await? { settle_usage_if_needed(data, &stored).await?; @@ -728,7 +730,8 @@ mod tests { use std::time::Duration; use aether_data_contracts::repository::settlement::{ - StoredUsageSettlement, UsageSettlementInput, + ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement, + UsageSettlementInput, }; use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord}; use aether_data_contracts::DataLayerError; @@ -755,6 +758,7 @@ mod tests { struct TestUsageStore { records: Mutex>, settlements: Mutex>, + reconciliations: Mutex>, enrich_calls: Mutex>, manual_proxy_counter_calls: AtomicUsize, } @@ -958,6 +962,17 @@ mod tests { true } + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + self.reconciliations + .lock() + .expect("reconciliations lock") + .push(input); + Ok(None) + } + async fn settle_usage( &self, input: UsageSettlementInput, @@ -1147,6 +1162,39 @@ mod tests { assert_eq!(settlements[0].request_id, "req-worker-123"); } + #[tokio::test] + async fn same_request_id_terminal_events_reconcile_each_reservation_token_before_upsert() { + let store = TestUsageStore::default(); + let mut first = sample_event(); + first.request_id = "shared-client-trace".to_string(); + first.data.actual_total_cost_usd = Some(0.25); + first.data.request_metadata = Some(serde_json::json!({ + "plan_usage_reservation_token": "server-token-a" + })); + let mut second = first.clone(); + second.data.actual_total_cost_usd = Some(0.75); + second.data.request_metadata = Some(serde_json::json!({ + "plan_usage_reservation_token": "server-token-b" + })); + + write_event_record(&store, &first) + .await + .expect("first terminal event"); + write_event_record(&store, &second) + .await + .expect("second terminal event"); + + let reconciliations = store.reconciliations.lock().expect("reconciliations lock"); + assert_eq!(reconciliations.len(), 2); + assert_eq!(reconciliations[0].request_id, "shared-client-trace"); + assert_eq!(reconciliations[0].reservation_token, "server-token-a"); + assert_eq!(reconciliations[0].actual_cost_units, 25_000_000); + assert_eq!(reconciliations[1].request_id, "shared-client-trace"); + assert_eq!(reconciliations[1].reservation_token, "server-token-b"); + assert_eq!(reconciliations[1].actual_cost_units, 75_000_000); + assert_eq!(store.records.lock().expect("records lock").len(), 2); + } + #[tokio::test] async fn replayable_usage_write_does_not_duplicate_transport_owned_proxy_counter() { let store = TestUsageStore::default(); diff --git a/crates/aether-usage/runtime/src/write.rs b/crates/aether-usage/runtime/src/write.rs index ba332fdab..058a103f6 100644 --- a/crates/aether-usage/runtime/src/write.rs +++ b/crates/aether-usage/runtime/src/write.rs @@ -8,7 +8,6 @@ use aether_data_contracts::repository::usage::{ WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, }; use aether_data_contracts::DataLayerError; -use base64::Engine as _; use serde_json::{json, Map, Value}; use crate::body_capture::{ @@ -24,11 +23,11 @@ use crate::request_metadata::{ sanitize_usage_request_metadata_ref, }; use crate::{ - map_usage_from_response, stream_capture_terminal_state, GatewayStreamReportRequest, - GatewaySyncReportRequest, StandardizedUsage, StreamCapturedTerminalState, UsageEvent, - UsageEventData, UsageEventType, STREAM_MISSING_TERMINAL_EVENT_CATEGORY, - STREAM_MISSING_TERMINAL_EVENT_MESSAGE, STREAM_TERMINAL_ERROR_CATEGORY, - STREAM_TERMINAL_ERROR_MESSAGE, + decode_internal_report_body_base64, map_usage_from_response, stream_capture_terminal_state, + GatewayStreamReportRequest, GatewaySyncReportRequest, StandardizedUsage, + StreamCapturedTerminalState, UsageEvent, UsageEventData, UsageEventType, + STREAM_MISSING_TERMINAL_EVENT_CATEGORY, STREAM_MISSING_TERMINAL_EVENT_MESSAGE, + STREAM_TERMINAL_ERROR_CATEGORY, STREAM_TERMINAL_ERROR_MESSAGE, }; #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -72,6 +71,20 @@ struct RuntimeRequestCaptureSeed { provider_request: Option, provider_request_body_ref: Option, body_states: UsageBodyStatesSeed, + request_has_inline_body: bool, + provider_request_has_inline_body: bool, +} + +/// Whether a seed keeps the request bodies it describes, or only describes them. +/// +/// A holder that has to outlive the request itself pays for every byte it keeps, +/// and a request body can be megabytes. [`RequestBodyCapture::Describe`] computes +/// every capture state, reference and derived fact from the real plan and report +/// context, and leaves out only the body values themselves. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RequestBodyCapture { + Keep, + Describe, } #[derive(Debug, Clone, PartialEq)] @@ -804,7 +817,8 @@ pub fn build_terminal_usage_context_seed( report_context: Option<&Value>, ) -> TerminalUsageContextSeed { let context = report_context.and_then(Value::as_object); - let request_capture = build_runtime_request_capture_seed(plan, context); + let request_capture = + build_runtime_request_capture_seed(plan, context, RequestBodyCapture::Keep); let client_contract = context_string(context, "client_contract") .or_else(|| context_string(context, "client_api_format")) .or_else(|| non_empty_str(Some(plan.client_api_format.as_str()))) @@ -1691,16 +1705,34 @@ pub fn build_usage_event_data_seed( plan: &ExecutionPlan, report_context: Option<&Value>, ) -> UsageEventData { - build_usage_event_data_seed_with_detail(plan, report_context) + build_usage_event_data_seed_with_detail(plan, report_context, RequestBodyCapture::Keep) +} + +/// Builds the same seed as [`build_usage_event_data_seed`] without keeping the +/// request bodies. +/// +/// This is for a caller that has to hold a seed for the whole life of an attempt +/// so it can still write a terminal row if the attempt is dropped: a request body +/// can be megabytes, and holding one per in-flight attempt is far more expensive +/// than the row it would eventually capture. Every capture state, body reference +/// and derived request fact is still computed from the real plan and report +/// context, so the resulting terminal write preserves the capture an earlier +/// non-terminal write recorded rather than clearing it. +pub fn build_usage_event_data_seed_describing_request_bodies( + plan: &ExecutionPlan, + report_context: Option<&Value>, +) -> UsageEventData { + build_usage_event_data_seed_with_detail(plan, report_context, RequestBodyCapture::Describe) } fn build_usage_event_data_seed_with_detail( plan: &ExecutionPlan, report_context: Option<&Value>, + capture: RequestBodyCapture, ) -> UsageEventData { let context = report_context.and_then(Value::as_object); let routing = build_runtime_routing_seed(plan, context); - let request_capture = build_runtime_request_capture_seed(plan, context); + let request_capture = build_runtime_request_capture_seed(plan, context, capture); let api_format = context_string(context, "client_api_format") .or_else(|| non_empty_str(Some(plan.client_api_format.as_str()))); let endpoint_api_format = context_string(context, "provider_api_format") @@ -1714,7 +1746,10 @@ fn build_usage_event_data_seed_with_detail( let request_type = Some(infer_request_type_from_contracts( api_format.as_deref(), endpoint_api_format.as_deref(), - request_capture.provider_request.as_ref(), + request_capture + .provider_request + .as_ref() + .or_else(|| provider_request_body_ref_for_inference(plan, context)), )); let api_family = api_format .as_deref() @@ -1737,9 +1772,9 @@ fn build_usage_event_data_seed_with_detail( build_runtime_request_metadata_seed_from_parts( plan, context, - request_capture.request_body.is_some(), + request_capture.request_has_inline_body, request_capture.request_body_ref.as_deref(), - request_capture.provider_request.is_some(), + request_capture.provider_request_has_inline_body, request_capture.provider_request_body_ref.as_deref(), plan.body.body_bytes_b64.as_deref(), ), @@ -1999,17 +2034,29 @@ fn plan_has_inline_json_body_for_usage(plan: &ExecutionPlan) -> bool { fn build_runtime_request_capture_seed( plan: &ExecutionPlan, context: Option<&Map>, + capture: RequestBodyCapture, ) -> RuntimeRequestCaptureSeed { - let request_body = context_body_value(context, "original_request_body"); + // Presence, not the value, is what every capture state and derived fact is + // built from, so both capture modes agree on all of them. + let request_has_inline_body = context_has_inline_body(context, "original_request_body"); + let provider_request_has_inline_body = + context_has_inline_body(context, "provider_request_body") + || plan_has_inline_json_body_for_usage(plan); + let (request_body, provider_request) = match capture { + RequestBodyCapture::Keep => ( + context_body_value(context, "original_request_body"), + context_body_value(context, "provider_request_body") + .or_else(|| plan_json_body_capture_for_usage(plan)), + ), + RequestBodyCapture::Describe => (None, None), + }; let request_body_ref = context_string(context, "request_body_ref"); - let provider_request = context_body_value(context, "provider_request_body") - .or_else(|| plan_json_body_capture_for_usage(plan)); let provider_request_body_ref = context_string(context, "provider_request_body_ref") .or_else(|| non_empty_str(plan.body.body_ref.as_deref())); let body_states = build_runtime_body_states_seed_from_parts( - request_body.is_some(), + request_has_inline_body, request_body_ref.as_deref(), - provider_request.is_some(), + provider_request_has_inline_body, provider_request_body_ref.as_deref(), plan.body.body_bytes_b64.is_some(), ); @@ -2020,9 +2067,26 @@ fn build_runtime_request_capture_seed( provider_request, provider_request_body_ref, body_states, + request_has_inline_body, + provider_request_has_inline_body, } } +/// Borrows the provider request body that request-type inference reads, without +/// cloning it. +fn provider_request_body_ref_for_inference<'a>( + plan: &'a ExecutionPlan, + context: Option<&'a Map>, +) -> Option<&'a Value> { + context_value_ref(context, "provider_request_body") + .filter(|value| !value.is_null()) + .or_else(|| { + plan_has_inline_json_body_for_usage(plan) + .then_some(plan.body.json_body.as_ref()) + .flatten() + }) +} + fn build_runtime_request_metadata_seed( plan: &ExecutionPlan, context: Option<&Map>, @@ -2130,6 +2194,12 @@ fn build_runtime_request_metadata_seed_from_parts( Value::String(websocket_transport), ); } + if let Some(reservation_token) = context_string(context, "plan_usage_reservation_token") { + metadata.insert( + "plan_usage_reservation_token".to_string(), + Value::String(reservation_token), + ); + } if let Some(usage_available) = context_bool(context, USAGE_AVAILABLE_METADATA_KEY) { metadata.insert( USAGE_AVAILABLE_METADATA_KEY.to_string(), @@ -2416,12 +2486,8 @@ fn sanitize_usage_event_capture_fields(mut data: UsageEventData) -> UsageEventDa data } -fn sanitize_usage_event_capture_fields_trusted(mut data: UsageEventData) -> UsageEventData { - data.request_headers = capture_usage_header_capture(data.request_headers); - data.provider_request_headers = capture_usage_header_capture(data.provider_request_headers); - data.response_headers = capture_usage_header_capture(data.response_headers); - data.client_response_headers = capture_usage_header_capture(data.client_response_headers); - data +fn sanitize_usage_event_capture_fields_trusted(data: UsageEventData) -> UsageEventData { + sanitize_usage_event_capture_fields(data) } fn sanitize_usage_event_data(mut data: UsageEventData) -> UsageEventData { @@ -2430,10 +2496,6 @@ fn sanitize_usage_event_data(mut data: UsageEventData) -> UsageEventData { data } -fn capture_usage_header_capture(value: Option) -> Option { - value.map(capture_usage_storage_value) -} - fn sanitize_usage_header_capture(value: Option) -> Option { mask_sensitive_headers_in_json_value(value).map(capture_usage_storage_value) } @@ -2662,29 +2724,29 @@ fn headers_to_json(headers: &BTreeMap) -> Option { )))) } -/// 默认敏感请求头清单。与 -/// `apps/aether-gateway/src/handlers/admin/system/shared/configs.rs` 中 -/// `sensitive_headers` 系统配置默认值保持一致。 -const DEFAULT_SENSITIVE_HEADERS: &[&str] = &[ - "authorization", - "x-api-key", - "api-key", - "x-goog-api-key", - "cookie", - "set-cookie", - "proxy-authorization", +const REDACTED_USAGE_VALUE: &str = "[redacted]"; + +/// Only headers whose values are protocol metadata are persisted verbatim. +/// Unknown headers are treated as credentials because providers commonly use +/// custom `X-*` names for authentication. +const SAFE_USAGE_HEADER_VALUE_NAMES: &[&str] = &[ + "accept", + "accept-encoding", + "content-encoding", + "content-length", + "content-type", + "transfer-encoding", + "x-request-id", + "x-trace-id", ]; -/// 判断 header 名是否属于敏感字段(大小写不敏感)。 fn is_sensitive_header(name: &str) -> bool { let trimmed = name.trim(); - DEFAULT_SENSITIVE_HEADERS + !SAFE_USAGE_HEADER_VALUE_NAMES .iter() .any(|candidate| trimmed.eq_ignore_ascii_case(candidate)) } -/// 对单个 header value 进行脱敏:保留前 4 + 后 4 字符,中间替换为 `****`。 -/// 长度小于等于 8 时整体替换为 `****`。 fn mask_header_value(name: &str, value: &str) -> String { if !is_sensitive_header(name) { return value.to_string(); @@ -2692,28 +2754,16 @@ fn mask_header_value(name: &str, value: &str) -> String { mask_sensitive_header_value(value) } -fn mask_sensitive_header_value(value: &str) -> String { - if value.len() <= 8 { - return "****".to_string(); - } - let prefix: String = value.chars().take(4).collect(); - let suffix: String = value - .chars() - .rev() - .take(4) - .collect::>() - .into_iter() - .rev() - .collect(); - format!("{prefix}****{suffix}") +fn mask_sensitive_header_value(_value: &str) -> String { + REDACTED_USAGE_VALUE.to_string() } -/// 对 JSON 形式的 headers 做就地脱敏。仅当 value 是 Object 时才会处理; -/// 其它形式的值保持不变。 +/// Non-object values cannot be established as a valid header map and are +/// discarded instead of being persisted verbatim. fn mask_sensitive_headers_in_json_value(value: Option) -> Option { let mut value = value?; let Value::Object(map) = &mut value else { - return Some(value); + return None; }; for (key, val) in map.iter_mut() { if !is_sensitive_header(key) { @@ -2941,9 +2991,7 @@ fn extract_generic_error_message_from_json(value: &Value) -> Option { fn decode_body_for_storage(body_base64: Option<&str>) -> Option { let body_base64 = body_base64?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(body_base64) - .ok()?; + let bytes = decode_internal_report_body_base64(body_base64).ok()?; if let Some(error_body) = aether_ai_formats::api::extract_provider_private_stream_error_body(None, &bytes) { @@ -3493,12 +3541,14 @@ mod tests { build_streaming_usage_event_from_owned_seed, build_streaming_usage_record, build_sync_terminal_usage_event, build_sync_terminal_usage_payload_seed, build_sync_terminal_usage_seed, build_terminal_usage_context_seed, - build_terminal_usage_event_from_seed, build_usage_event_data_seed, decode_body_for_storage, + build_terminal_usage_event_from_seed, build_usage_event_data_seed, + build_usage_event_data_seed_describing_request_bodies, decode_body_for_storage, extract_token_counts_from_json, extract_token_counts_from_value, headers_to_json, mask_header_value, mask_sensitive_body_fields, mask_sensitive_headers_in_json_value, parse_sse_body_for_storage, resolve_error_message, trim_owned_non_empty_string, LifecycleUsageSeed, TerminalUsageSeed, UsageBodyRefsSeed, UsageBodyStatesSeed, UsageRoutingSeed, UsageTerminalState, MAX_USAGE_CAPTURE_BYTES, MAX_USAGE_CAPTURE_DEPTH, + REDACTED_USAGE_VALUE, }; use crate::{ build_upsert_usage_record_from_event, GatewayStreamReportRequest, GatewaySyncReportRequest, @@ -3934,7 +3984,8 @@ mod tests { .expect("pending usage should keep request metadata"); assert_eq!(metadata.get("api_key_is_standalone"), Some(&json!(true))); assert_eq!(metadata.get("client_ip"), Some(&json!("203.0.113.8"))); - assert_eq!(metadata.get("user_agent"), Some(&json!("Claude-Code/1.0"))); + assert_eq!(metadata.get("client_family"), Some(&json!("claude_code"))); + assert!(metadata.get("user_agent").is_none()); let body_size = metadata .get("body_size") .and_then(Value::as_object) @@ -5644,14 +5695,14 @@ mod tests { assert_eq!( event.data.response_headers, Some(json!({ - "authorization": "Bear****oken", + "authorization": REDACTED_USAGE_VALUE, "content-type": "application/json" })) ); assert_eq!( event.data.client_response_headers, Some(json!({ - "authorization": "Bear****oken", + "authorization": REDACTED_USAGE_VALUE, "content-type": "application/json" })) ); @@ -6615,26 +6666,26 @@ mod tests { assert_eq!( event.data.request_headers, Some(json!({ - "authorization": "Bear****oken", + "authorization": REDACTED_USAGE_VALUE, "accept": "application/json" })) ); assert_eq!( event.data.provider_request_headers, Some(json!({ - "x-api-key": "sk-p****cret" + "x-api-key": REDACTED_USAGE_VALUE })) ); assert_eq!( event.data.response_headers, Some(json!({ - "set-cookie": "sess****okie" + "set-cookie": REDACTED_USAGE_VALUE })) ); assert_eq!( event.data.client_response_headers, Some(json!({ - "authorization": "Bear****cret" + "authorization": REDACTED_USAGE_VALUE })) ); assert_eq!( @@ -6661,14 +6712,7 @@ mod tests { .request_metadata .as_ref() .and_then(|value| value.get("billing_snapshot")), - Some(&json!({ - "truncated": true, - "reason": "usage_request_metadata_limits_exceeded", - "max_depth": 32, - "max_nodes": 4_000, - "max_bytes": 16 * 1024, - "value_kind": "object" - })) + None ); } @@ -6714,19 +6758,7 @@ mod tests { .expect("pending record should build"); assert_eq!(record.candidate_id.as_deref(), Some("cand-1")); - assert_eq!( - record.request_metadata, - Some(json!({ - "billing_snapshot": { - "truncated": true, - "reason": "usage_request_metadata_limits_exceeded", - "max_depth": 32, - "max_nodes": 4_000, - "max_bytes": 16 * 1024, - "value_kind": "object" - } - })) - ); + assert_eq!(record.request_metadata, None); } #[test] @@ -6769,7 +6801,7 @@ mod tests { assert_eq!( data.request_headers, Some(json!({ - "authorization": "Bear****cret", + "authorization": REDACTED_USAGE_VALUE, "accept": "application/json" })) ); @@ -6826,10 +6858,7 @@ mod tests { metadata.get("end_to_end_first_byte_time_ms"), Some(&json!(10_120)) ); - assert_eq!( - metadata.get("db_timings_ms"), - Some(&json!({"query_count": 2})) - ); + assert_eq!(metadata.get("db_timings_ms"), None); assert_eq!( metadata.get("trace_id"), Some(&json!("trace-seed-metadata-1")) @@ -6842,27 +6871,142 @@ mod tests { assert!(body_size.get("provider_request_body").is_some()); } + #[test] + fn describing_request_bodies_matches_the_capturing_seed_apart_from_the_bodies() { + let plan = ExecutionPlan { + request_id: "req-seed-describe-1".to_string(), + candidate_id: Some("cand-seed-describe-1".to_string()), + provider_name: Some("OpenAI".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-1".to_string(), + method: "POST".to_string(), + url: "https://example.com/v1/chat/completions".to_string(), + headers: BTreeMap::new(), + content_type: Some("application/json".to_string()), + content_encoding: None, + body: RequestBody::from_json(json!({ + "model": "gpt-5", + "service_tier": "priority", + "reasoning": {"effort": "high"} + })), + stream: false, + client_api_format: "openai:chat".to_string(), + provider_api_format: "openai:chat".to_string(), + model_name: Some("gpt-5".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + }; + let report_context = json!({ + "client_api_format": "openai:chat", + "provider_api_format": "openai:chat", + "original_request_body": {"model": "gpt-5", "messages": []}, + "original_headers": {"accept": "application/json"} + }); + + let captured = build_usage_event_data_seed(&plan, Some(&report_context)); + let described = + build_usage_event_data_seed_describing_request_bodies(&plan, Some(&report_context)); + + // Only the two heavy values differ. + assert!(captured.request_body.is_some()); + assert!(captured.provider_request_body.is_some()); + assert_eq!(described.request_body, None); + assert_eq!(described.provider_request_body, None); + + // Everything a terminal write reads to decide what to do with the stored + // capture is identical, so the described seed preserves it rather than + // clearing it. + assert_eq!( + described.request_body_state, + Some(UsageBodyCaptureState::Inline) + ); + assert_eq!( + described.provider_request_body_state, + Some(UsageBodyCaptureState::Inline) + ); + assert_eq!(described.request_body_state, captured.request_body_state); + assert_eq!( + described.provider_request_body_state, + captured.provider_request_body_state + ); + assert_eq!(described.request_body_ref, captured.request_body_ref); + assert_eq!( + described.provider_request_body_ref, + captured.provider_request_body_ref + ); + assert_eq!(described.request_type, captured.request_type); + assert_eq!(described.request_metadata, captured.request_metadata); + assert_eq!( + described.provider_request_headers, + captured.provider_request_headers + ); + assert_eq!(described.request_headers, captured.request_headers); + } + + #[test] + fn describing_request_bodies_keeps_the_unavailable_marker_for_raw_bodies() { + let mut plan = ExecutionPlan { + request_id: "req-seed-describe-2".to_string(), + candidate_id: None, + provider_name: Some("OpenAI".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-1".to_string(), + method: "POST".to_string(), + url: "https://example.com/v1/chat/completions".to_string(), + headers: BTreeMap::new(), + content_type: Some("application/json".to_string()), + content_encoding: None, + body: RequestBody::from_json(json!({"model": "gpt-5"})), + stream: false, + client_api_format: "openai:chat".to_string(), + provider_api_format: "openai:chat".to_string(), + model_name: Some("gpt-5".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + }; + plan.body.json_body = None; + plan.body.body_bytes_b64 = Some("eyJtb2RlbCI6ICJncHQtNSJ9".to_string()); + + let described = build_usage_event_data_seed_describing_request_bodies(&plan, None); + + assert_eq!( + described.provider_request_body_state, + Some(UsageBodyCaptureState::Unavailable) + ); + } + #[test] fn masks_known_sensitive_header_values() { let token = "Bearer eyJhbGciOiJSUzI1NiJ9.payload-here.signature-tail"; let masked = mask_header_value("authorization", token); - assert!(masked.starts_with("Bear")); - assert!(masked.ends_with("tail")); - assert!(masked.contains("****")); - assert!(!masked.contains("payload-here")); + assert_eq!(masked, REDACTED_USAGE_VALUE); // 大小写不敏感 assert_eq!( mask_header_value("Authorization", "12345678"), - "****", - "短值整体替换为 ****", + REDACTED_USAGE_VALUE, + ); + assert_eq!( + mask_header_value("X-Api-Key", "abcdefghij"), + REDACTED_USAGE_VALUE, ); - assert_eq!(mask_header_value("X-Api-Key", "abcdefghij"), "abcd****ghij",); - // 非敏感头保持原样 + // Unknown custom headers are redacted by default. assert_eq!( mask_header_value("user-agent", "codex-tui/0.1"), - "codex-tui/0.1", + REDACTED_USAGE_VALUE, + ); + assert_eq!( + mask_header_value("x-custom-auth", "tenant-secret"), + REDACTED_USAGE_VALUE, + ); + assert_eq!( + mask_header_value("content-type", "application/json"), + "application/json", ); } @@ -6886,21 +7030,17 @@ mod tests { .get("authorization") .and_then(|v| v.as_str()) .expect("authorization should be string"); - assert!(auth.starts_with("Bear")); - assert!(auth.contains("****")); - assert!(!auth.contains("eyJhbGciOiJSUzI1NiJ9")); + assert_eq!(auth, REDACTED_USAGE_VALUE); let api_key = object .get("x-api-key") .and_then(|v| v.as_str()) .expect("x-api-key should be string"); - assert!(api_key.starts_with("sk-p")); - assert!(api_key.contains("****")); - assert!(!api_key.contains("1234567890")); + assert_eq!(api_key, REDACTED_USAGE_VALUE); assert_eq!( object.get("user-agent").and_then(|v| v.as_str()), - Some("codex-tui/0.1"), + Some(REDACTED_USAGE_VALUE), ); } @@ -6924,15 +7064,13 @@ mod tests { .get("Authorization") .and_then(|v| v.as_str()) .expect("Authorization should be string"); - assert!(auth.contains("****")); - assert!(!auth.contains("eyJhbGciOiJSUzI1NiJ9")); + assert_eq!(auth, REDACTED_USAGE_VALUE); let cookie = object .get("Cookie") .and_then(|v| v.as_str()) .expect("Cookie should be string"); - assert!(cookie.contains("****")); - assert!(!cookie.contains("verylongcookievalue")); + assert_eq!(cookie, REDACTED_USAGE_VALUE); assert_eq!( object.get("Accept").and_then(|v| v.as_str()), @@ -6941,12 +7079,12 @@ mod tests { } #[test] - fn mask_sensitive_headers_passthrough_for_non_object() { + fn mask_sensitive_headers_discards_non_object() { // None 输入返回 None assert!(mask_sensitive_headers_in_json_value(None).is_none()); - // 非 object 输入原样返回 + // 非 object 不是可验证的 header map,直接丢弃。 let masked = mask_sensitive_headers_in_json_value(Some(json!("not an object"))); - assert_eq!(masked, Some(json!("not an object"))); + assert_eq!(masked, None); } #[test] diff --git a/crates/aether-video-tasks-core/Cargo.toml b/crates/aether-video-tasks-core/Cargo.toml index 7305ce83b..25f284b52 100644 --- a/crates/aether-video-tasks-core/Cargo.toml +++ b/crates/aether-video-tasks-core/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] aether-contracts.workspace = true +aether-crypto.workspace = true aether-data-contracts.workspace = true async-trait.workspace = true serde.workspace = true diff --git a/crates/aether-video-tasks-core/src/gemini.rs b/crates/aether-video-tasks-core/src/gemini.rs index b10011931..8dd40da3e 100644 --- a/crates/aether-video-tasks-core/src/gemini.rs +++ b/crates/aether-video-tasks-core/src/gemini.rs @@ -4,11 +4,12 @@ use aether_data_contracts::repository::video_tasks::{ }; use serde_json::{json, Map, Value}; +use crate::types::sanitize_video_task_error_code; use crate::{ build_video_follow_up_report_context, current_unix_timestamp_secs, gemini_metadata_video_url, request_body_string, request_body_u32, resolve_follow_up_auth, GeminiVideoTaskSeed, - LocalVideoTaskFollowUpPlan, LocalVideoTaskReadResponse, LocalVideoTaskSnapshot, - LocalVideoTaskStatus, VideoFollowUpReportContextInput, DEFAULT_VIDEO_TASK_MAX_POLL_COUNT, + LocalVideoTaskFollowUpPlan, LocalVideoTaskReadResponse, LocalVideoTaskStatus, + VideoFollowUpReportContextInput, DEFAULT_VIDEO_TASK_MAX_POLL_COUNT, DEFAULT_VIDEO_TASK_POLL_INTERVAL_SECONDS, }; @@ -66,10 +67,9 @@ fn build_gemini_failed_body(task: StoredVideoTask) -> Value { "name": stored_task_operation_name(&task), "done": true, "error": { - "code": task.error_code.unwrap_or_else(|| "UNKNOWN".to_string()), - "message": task - .error_message - .unwrap_or_else(|| "Video generation failed".to_string()), + "code": sanitize_video_task_error_code(task.error_code) + .unwrap_or_else(|| "unknown".to_string()), + "message": "Video generation failed", } }) } @@ -99,14 +99,13 @@ impl GeminiVideoTaskSeed { if let Some(error) = error { self.status = LocalVideoTaskStatus::Failed; self.progress_percent = 100; - self.error_code = error - .get("code") - .and_then(Value::as_str) - .map(str::to_string); - self.error_message = error - .get("message") - .and_then(Value::as_str) - .map(str::to_string); + self.error_code = sanitize_video_task_error_code( + error + .get("code") + .and_then(Value::as_str) + .map(str::to_string), + ); + self.error_message = None; } else { self.status = LocalVideoTaskStatus::Completed; self.progress_percent = 100; @@ -121,10 +120,7 @@ impl GeminiVideoTaskSeed { self.progress_percent = 50; self.error_code = None; self.error_message = None; - self.metadata = provider_body - .get("metadata") - .cloned() - .unwrap_or_else(|| json!({})); + self.metadata = json!({}); } pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option { @@ -273,17 +269,15 @@ impl GeminiVideoTaskSeed { "name": operation_name, "done": true, "error": { - "code": self.error_code.clone().unwrap_or_else(|| "UNKNOWN".to_string()), - "message": self - .error_message - .clone() - .unwrap_or_else(|| "Video generation failed".to_string()), + "code": sanitize_video_task_error_code(self.error_code.clone()) + .unwrap_or_else(|| "unknown".to_string()), + "message": "Video generation failed", } }), _ => json!({ "name": operation_name, "done": false, - "metadata": self.metadata.clone(), + "metadata": {}, }), } } @@ -298,7 +292,7 @@ impl GeminiVideoTaskSeed { ), _ => None, }; - UpsertVideoTask { + let mut record = UpsertVideoTask { id: self.local_short_id.clone(), short_id: Some(self.local_short_id.clone()), request_id: self.persistence.request_id.clone(), @@ -316,7 +310,7 @@ impl GeminiVideoTaskSeed { model: Some(self.model.clone()), prompt: request_body_string(&self.persistence.original_request_body, "prompt") .or_else(|| Some(String::new())), - original_request_body: Some(self.persistence.original_request_body.clone()), + original_request_body: None, duration_seconds: request_body_u32(&self.persistence.original_request_body, "seconds") .or_else(|| { request_body_u32(&self.persistence.original_request_body, "duration_seconds") @@ -340,13 +334,12 @@ impl GeminiVideoTaskSeed { completed_at_unix_secs: None, updated_at_unix_secs: now_unix_secs, error_code: self.error_code.clone(), - error_message: self.error_message.clone(), + error_message: None, video_url: gemini_metadata_video_url(&self.metadata), - request_metadata: Some(json!({ - "rust_owner": "async_task", - "rust_local_snapshot": LocalVideoTaskSnapshot::Gemini(self.clone()), - })), - } + request_metadata: None, + }; + record.sanitize_for_persistence(); + record } fn resolve_operation_path(&self) -> Option { @@ -366,7 +359,15 @@ impl GeminiVideoTaskSeed { #[cfg(test)] mod tests { + use std::collections::BTreeMap; + use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus}; + use serde_json::json; + + use crate::{ + GeminiVideoTaskSeed, LocalVideoTaskPersistence, LocalVideoTaskStatus, + LocalVideoTaskTransport, + }; use super::map_gemini_stored_task_to_read_response; @@ -420,4 +421,70 @@ mod tests { assert_eq!(response.status_code, 404); assert_eq!(response.body_json["detail"], "Video task was cancelled"); } + + #[test] + fn builds_minimal_gemini_persistence_record_and_strips_signed_url_query() { + let seed = GeminiVideoTaskSeed { + local_short_id: "gemini-sensitive".to_string(), + upstream_operation_name: "operations/upstream-sensitive".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("api-key-1".to_string()), + model: "veo-3".to_string(), + status: LocalVideoTaskStatus::Completed, + progress_percent: 100, + error_code: Some("provider-secret-diagnostic".to_string()), + error_message: Some("provider response contained secret".to_string()), + metadata: json!({ + "response": { + "generateVideoResponse": { + "generatedSamples": [{ + "video": { + "uri": "https://files.example/video.mp4?alt=media&token=sensitive#fragment" + } + }] + } + } + }), + persistence: LocalVideoTaskPersistence { + request_id: "request-gemini-sensitive".to_string(), + username: Some("alice".to_string()), + api_key_name: Some("primary".to_string()), + client_api_format: "gemini:video".to_string(), + provider_api_format: "gemini:video".to_string(), + original_request_body: json!({ + "prompt": "business prompt", + "provider_token": "sensitive" + }), + format_converted: false, + }, + transport: LocalVideoTaskTransport { + upstream_base_url: "https://generativelanguage.example".to_string(), + provider_name: Some("gemini".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "provider-key-1".to_string(), + headers: BTreeMap::from([( + "x-goog-api-key".to_string(), + "sensitive-api-key".to_string(), + )]), + content_type: Some("application/json".to_string()), + model_name: Some("veo-3".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + }, + }; + + let record = seed.to_upsert_record(); + + assert_eq!(record.error_code.as_deref(), Some("provider_error")); + assert!(record.original_request_body.is_none()); + assert!(record.progress_message.is_none()); + assert!(record.error_message.is_none()); + assert!(record.request_metadata.is_none()); + assert_eq!( + record.video_url.as_deref(), + Some("https://files.example/video.mp4?alt=media") + ); + } } diff --git a/crates/aether-video-tasks-core/src/lib.rs b/crates/aether-video-tasks-core/src/lib.rs index 994b94a7c..7218851b2 100644 --- a/crates/aether-video-tasks-core/src/lib.rs +++ b/crates/aether-video-tasks-core/src/lib.rs @@ -32,7 +32,10 @@ pub use path::{ resolve_video_task_hydration_lookup_key, resolve_video_task_read_lookup_key, resolve_video_task_report_lookup, VideoTaskReportLookup, }; -pub use read_side::{read_data_backed_video_task_response, StoredVideoTaskReadSide}; +pub use read_side::{ + read_data_backed_video_task_response, read_data_backed_video_task_response_for_user, + StoredVideoTaskReadSide, +}; pub use service::VideoTaskService; pub use store::VideoTaskStore; pub use store_backend::{FileVideoTaskStore, InMemoryVideoTaskStore}; diff --git a/crates/aether-video-tasks-core/src/openai.rs b/crates/aether-video-tasks-core/src/openai.rs index dcd6ea57d..4dc632642 100644 --- a/crates/aether-video-tasks-core/src/openai.rs +++ b/crates/aether-video-tasks-core/src/openai.rs @@ -6,13 +6,13 @@ use aether_data_contracts::repository::video_tasks::{ }; use serde_json::{json, Map, Value}; +use crate::types::sanitize_video_task_error_code; use crate::{ build_video_follow_up_report_context, current_unix_timestamp_secs, map_openai_task_status, parse_video_content_variant, request_body_string, request_body_u32, resolve_follow_up_auth, LocalVideoTaskContentAction, LocalVideoTaskFollowUpPlan, LocalVideoTaskReadResponse, - LocalVideoTaskSnapshot, LocalVideoTaskStatus, OpenAiVideoTaskSeed, - VideoFollowUpReportContextInput, DEFAULT_VIDEO_TASK_MAX_POLL_COUNT, - DEFAULT_VIDEO_TASK_POLL_INTERVAL_SECONDS, + LocalVideoTaskStatus, OpenAiVideoTaskSeed, VideoFollowUpReportContextInput, + DEFAULT_VIDEO_TASK_MAX_POLL_COUNT, DEFAULT_VIDEO_TASK_POLL_INTERVAL_SECONDS, }; fn openai_video_resource_url(api_root: &str, suffix: &str) -> String { @@ -71,10 +71,9 @@ fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) VideoTaskStatus::Failed | VideoTaskStatus::Expired | VideoTaskStatus::Cancelled ) { body["error"] = json!({ - "code": task.error_code.unwrap_or_else(|| "unknown".to_string()), - "message": task - .error_message - .unwrap_or_else(|| "Video generation failed".to_string()), + "code": sanitize_video_task_error_code(task.error_code) + .unwrap_or_else(|| "unknown".to_string()), + "message": "Video generation failed", }); } @@ -119,14 +118,13 @@ 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 = error - .and_then(|value| value.get("code")) - .and_then(Value::as_str) - .map(str::to_string); - self.error_message = error - .and_then(|value| value.get("message")) - .and_then(Value::as_str) - .map(str::to_string); + 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_message = None; self.video_url = provider_body .get("video_url") .or_else(|| provider_body.get("url")) @@ -157,14 +155,7 @@ impl OpenAiVideoTaskSeed { LocalVideoTaskStatus::Failed | LocalVideoTaskStatus::Expired => { return Some(LocalVideoTaskContentAction::Immediate { status_code: 422, - body_json: json!({ - "detail": format!( - "Video generation failed: {}", - self.error_message - .clone() - .unwrap_or_else(|| "Unknown error".to_string()) - ) - }), + body_json: json!({"detail": "Video generation failed"}), }); } LocalVideoTaskStatus::Cancelled => { @@ -281,11 +272,9 @@ impl OpenAiVideoTaskSeed { || self.status == LocalVideoTaskStatus::Expired { body["error"] = json!({ - "code": self.error_code.clone().unwrap_or_else(|| "unknown".to_string()), - "message": self - .error_message - .clone() - .unwrap_or_else(|| "Video generation failed".to_string()), + "code": sanitize_video_task_error_code(self.error_code.clone()) + .unwrap_or_else(|| "unknown".to_string()), + "message": "Video generation failed", }); } @@ -582,7 +571,7 @@ impl OpenAiVideoTaskSeed { ), _ => None, }; - UpsertVideoTask { + let mut record = UpsertVideoTask { id: self.local_task_id.clone(), short_id: None, request_id: self.persistence.request_id.clone(), @@ -599,7 +588,7 @@ impl OpenAiVideoTaskSeed { format_converted: self.persistence.format_converted, model: self.model.clone().or_else(|| Some(String::new())), prompt: self.prompt.clone().or_else(|| Some(String::new())), - original_request_body: Some(self.persistence.original_request_body.clone()), + original_request_body: None, duration_seconds: request_body_u32(&self.persistence.original_request_body, "seconds"), resolution: request_body_string(&self.persistence.original_request_body, "resolution"), aspect_ratio: request_body_string( @@ -620,19 +609,26 @@ impl OpenAiVideoTaskSeed { completed_at_unix_secs: self.completed_at_unix_secs, updated_at_unix_secs: self.completed_at_unix_secs.unwrap_or(now_unix_secs), error_code: self.error_code.clone(), - error_message: self.error_message.clone(), - video_url: self.video_url.clone(), - request_metadata: Some(json!({ - "rust_owner": "async_task", - "rust_local_snapshot": LocalVideoTaskSnapshot::OpenAi(self.clone()), - })), - } + error_message: None, + video_url: None, + request_metadata: None, + }; + record.sanitize_for_persistence(); + record } } #[cfg(test)] mod tests { + use std::collections::BTreeMap; + use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus}; + use serde_json::json; + + use crate::{ + LocalVideoTaskPersistence, LocalVideoTaskStatus, LocalVideoTaskTransport, + OpenAiVideoTaskSeed, + }; use super::map_openai_stored_task_to_read_response; @@ -687,10 +683,77 @@ mod tests { assert_eq!(response.body_json["id"], "task-openai-123"); assert_eq!(response.body_json["status"], "failed"); assert_eq!(response.body_json["completed_at"], 1712345688u64); - assert_eq!(response.body_json["error"]["code"], "upstream_failed"); + assert_eq!(response.body_json["error"]["code"], "provider_error"); + assert_eq!( + response.body_json["error"]["message"], + "Video generation failed" + ); assert_eq!( response.body_json["video_url"], "https://cdn.example.com/video.mp4" ); } + + #[test] + fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() { + let seed = OpenAiVideoTaskSeed { + local_task_id: "task-openai-sensitive".to_string(), + upstream_task_id: "upstream-openai-sensitive".to_string(), + created_at_unix_ms: 1_712_345_678, + user_id: Some("user-1".to_string()), + api_key_id: Some("api-key-1".to_string()), + model: Some("sora-2".to_string()), + prompt: Some("business prompt".to_string()), + size: Some("1280x720".to_string()), + seconds: Some("4".to_string()), + remixed_from_video_id: None, + status: LocalVideoTaskStatus::Failed, + progress_percent: 100, + completed_at_unix_secs: Some(1_712_345_700), + expires_at_unix_secs: None, + error_code: Some("provider-secret-diagnostic".to_string()), + error_message: Some("Bearer sk-sensitive-provider-error".to_string()), + video_url: Some("https://cdn.example/video.mp4?token=sensitive".to_string()), + persistence: LocalVideoTaskPersistence { + request_id: "request-openai-sensitive".to_string(), + username: Some("alice".to_string()), + api_key_name: Some("primary".to_string()), + client_api_format: "openai:video".to_string(), + provider_api_format: "openai:video".to_string(), + original_request_body: json!({ + "prompt": "business prompt", + "seconds": "4", + "provider_token": "sk-sensitive" + }), + format_converted: false, + }, + transport: LocalVideoTaskTransport { + upstream_base_url: "https://api.example/v1".to_string(), + provider_name: Some("openai".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "provider-key-1".to_string(), + headers: BTreeMap::from([( + "authorization".to_string(), + "Bearer sk-sensitive".to_string(), + )]), + content_type: Some("application/json".to_string()), + model_name: Some("sora-2".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + }, + }; + + let record = seed.to_upsert_record(); + + assert_eq!(record.error_code.as_deref(), Some("provider_error")); + assert!(record.original_request_body.is_none()); + assert!(record.progress_message.is_none()); + assert!(record.error_message.is_none()); + assert!(record.video_url.is_none()); + assert!(record.request_metadata.is_none()); + assert_eq!(record.duration_seconds, Some(4)); + assert_eq!(record.size.as_deref(), Some("1280x720")); + } } diff --git a/crates/aether-video-tasks-core/src/read_side.rs b/crates/aether-video-tasks-core/src/read_side.rs index dfb5503d2..fd8749ce1 100644 --- a/crates/aether-video-tasks-core/src/read_side.rs +++ b/crates/aether-video-tasks-core/src/read_side.rs @@ -13,16 +13,45 @@ pub trait StoredVideoTaskReadSide: Send + Sync { &self, key: VideoTaskLookupKey<'_>, ) -> Result, DataLayerError>; + + async fn find_stored_video_task_for_user( + &self, + key: VideoTaskLookupKey<'_>, + user_id: &str, + ) -> Result, DataLayerError>; } pub async fn read_data_backed_video_task_response( state: &impl StoredVideoTaskReadSide, route_family: Option<&str>, request_path: &str, +) -> Result, DataLayerError> { + read_data_backed_video_task_response_inner(state, route_family, request_path, None).await +} + +pub async fn read_data_backed_video_task_response_for_user( + state: &impl StoredVideoTaskReadSide, + route_family: Option<&str>, + request_path: &str, + user_id: &str, +) -> Result, DataLayerError> { + let user_id = user_id.trim(); + if user_id.is_empty() { + return Ok(None); + } + read_data_backed_video_task_response_inner(state, route_family, request_path, Some(user_id)) + .await +} + +async fn read_data_backed_video_task_response_inner( + state: &impl StoredVideoTaskReadSide, + route_family: Option<&str>, + request_path: &str, + user_id: Option<&str>, ) -> Result, DataLayerError> { match route_family { - Some("openai") => read_openai_video_task_response(state, request_path).await, - Some("gemini") => read_gemini_video_task_response(state, request_path).await, + Some("openai") => read_openai_video_task_response(state, request_path, user_id).await, + Some("gemini") => read_gemini_video_task_response(state, request_path, user_id).await, _ => Ok(None), } } @@ -30,12 +59,21 @@ pub async fn read_data_backed_video_task_response( async fn read_openai_video_task_response( state: &impl StoredVideoTaskReadSide, request_path: &str, + user_id: Option<&str>, ) -> Result, DataLayerError> { let Some(lookup) = resolve_video_task_read_lookup_key(Some("openai"), request_path) else { return Ok(None); }; - let Some(task) = state.find_stored_video_task(lookup).await? else { + let task = match user_id { + Some(user_id) => { + state + .find_stored_video_task_for_user(lookup, user_id) + .await? + } + None => state.find_stored_video_task(lookup).await?, + }; + let Some(task) = task else { return Ok(None); }; @@ -49,12 +87,21 @@ async fn read_openai_video_task_response( async fn read_gemini_video_task_response( state: &impl StoredVideoTaskReadSide, request_path: &str, + user_id: Option<&str>, ) -> Result, DataLayerError> { let Some(lookup) = resolve_video_task_read_lookup_key(Some("gemini"), request_path) else { return Ok(None); }; - let Some(task) = state.find_stored_video_task(lookup).await? else { + let task = match user_id { + Some(user_id) => { + state + .find_stored_video_task_for_user(lookup, user_id) + .await? + } + None => state.find_stored_video_task(lookup).await?, + }; + let Some(task) = task else { return Ok(None); }; diff --git a/crates/aether-video-tasks-core/src/service.rs b/crates/aether-video-tasks-core/src/service.rs index 2684d5184..c94fc00db 100644 --- a/crates/aether-video-tasks-core/src/service.rs +++ b/crates/aether-video-tasks-core/src/service.rs @@ -29,10 +29,11 @@ impl VideoTaskService { pub fn with_file_store( mode: VideoTaskTruthSourceMode, path: impl Into, + encryption_key: impl Into, ) -> std::io::Result { Ok(Self::with_store( mode, - Arc::new(FileVideoTaskStore::new(path)?), + Arc::new(FileVideoTaskStore::new(path, encryption_key)?), )) } @@ -113,6 +114,21 @@ impl VideoTaskService { } } + pub fn read_response_for_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + ) -> Option { + if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { + return None; + } + let snapshot = self.snapshot_for_route(route_family, request_path)?; + snapshot + .belongs_to_user(user_id) + .then(|| snapshot.read_response()) + } + pub fn snapshot_for_route( &self, route_family: Option<&str>, @@ -120,15 +136,29 @@ impl VideoTaskService { ) -> Option { match route_family { Some("openai") => extract_openai_task_id_from_path(request_path) + .or_else(|| extract_openai_task_id_from_cancel_path(request_path)) + .or_else(|| extract_openai_task_id_from_remix_path(request_path)) + .or_else(|| extract_openai_task_id_from_content_path(request_path)) .and_then(|task_id| self.store.clone_openai(task_id)) .map(LocalVideoTaskSnapshot::OpenAi), Some("gemini") => extract_gemini_short_id_from_path(request_path) + .or_else(|| extract_gemini_short_id_from_cancel_path(request_path)) .and_then(|short_id| self.store.clone_gemini(short_id)) .map(LocalVideoTaskSnapshot::Gemini), _ => None, } } + pub fn route_belongs_to_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + ) -> bool { + self.snapshot_for_route(route_family, request_path) + .is_some_and(|snapshot| snapshot.belongs_to_user(user_id)) + } + pub fn prepare_openai_content_stream_action( &self, request_path: &str, @@ -143,6 +173,25 @@ impl VideoTaskService { seed.build_content_stream_action(query_string, trace_id) } + pub fn prepare_openai_content_stream_action_for_user( + &self, + request_path: &str, + query_string: Option<&str>, + trace_id: &str, + user_id: &str, + ) -> Option { + if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { + return None; + } + let task_id = extract_openai_task_id_from_content_path(request_path)?; + let seed = self.store.clone_openai(task_id)?; + let snapshot = LocalVideoTaskSnapshot::OpenAi(seed.clone()); + if !snapshot.belongs_to_user(user_id) { + return None; + } + seed.build_content_stream_action(query_string, trace_id) + } + pub fn snapshot_for_refresh_plan( &self, refresh_plan: &LocalVideoTaskReadRefreshPlan, @@ -190,31 +239,63 @@ impl VideoTaskService { if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { return None; } + let snapshot = self.snapshot_for_read_refresh_route(route_family, request_path)?; + Self::build_read_refresh_sync_plan_from_snapshot(snapshot, trace_id) + } + + pub fn prepare_read_refresh_sync_plan_for_user( + &self, + route_family: Option<&str>, + request_path: &str, + user_id: &str, + trace_id: &str, + ) -> Option { + if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { + return None; + } + let snapshot = self.snapshot_for_read_refresh_route(route_family, request_path)?; + if !snapshot.belongs_to_user(user_id) { + return None; + } + Self::build_read_refresh_sync_plan_from_snapshot(snapshot, trace_id) + } + + fn snapshot_for_read_refresh_route( + &self, + route_family: Option<&str>, + request_path: &str, + ) -> Option { match route_family { - Some("openai") => { - let task_id = extract_openai_task_id_from_path(request_path)?; - let seed = self.store.clone_openai(task_id)?; - Some(LocalVideoTaskReadRefreshPlan { - plan: seed.build_get_follow_up_plan(trace_id)?, - projection_target: LocalVideoTaskProjectionTarget::OpenAi { - task_id: task_id.to_string(), - }, - }) - } - Some("gemini") => { - let short_id = extract_gemini_short_id_from_path(request_path)?; - let seed = self.store.clone_gemini(short_id)?; - Some(LocalVideoTaskReadRefreshPlan { - plan: seed.build_get_follow_up_plan(trace_id)?, - projection_target: LocalVideoTaskProjectionTarget::Gemini { - short_id: short_id.to_string(), - }, - }) - } + Some("openai") => extract_openai_task_id_from_path(request_path) + .and_then(|task_id| self.store.clone_openai(task_id)) + .map(LocalVideoTaskSnapshot::OpenAi), + Some("gemini") => extract_gemini_short_id_from_path(request_path) + .and_then(|short_id| self.store.clone_gemini(short_id)) + .map(LocalVideoTaskSnapshot::Gemini), _ => None, } } + fn build_read_refresh_sync_plan_from_snapshot( + snapshot: LocalVideoTaskSnapshot, + trace_id: &str, + ) -> Option { + match snapshot { + LocalVideoTaskSnapshot::OpenAi(seed) => Some(LocalVideoTaskReadRefreshPlan { + plan: seed.build_get_follow_up_plan(trace_id)?, + projection_target: LocalVideoTaskProjectionTarget::OpenAi { + task_id: seed.local_task_id.clone(), + }, + }), + LocalVideoTaskSnapshot::Gemini(seed) => Some(LocalVideoTaskReadRefreshPlan { + plan: seed.build_get_follow_up_plan(trace_id)?, + projection_target: LocalVideoTaskProjectionTarget::Gemini { + short_id: seed.local_short_id.clone(), + }, + }), + } + } + pub fn prepare_poll_refresh_batch( &self, limit: usize, @@ -258,6 +339,18 @@ impl VideoTaskService { } let snapshot = LocalVideoTaskSnapshot::from_stored_task(task)?; + self.prepare_poll_refresh_plan_for_snapshot(snapshot, trace_id) + } + + pub fn prepare_poll_refresh_plan_for_snapshot( + &self, + snapshot: LocalVideoTaskSnapshot, + trace_id: &str, + ) -> Option { + if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { + return None; + } + match snapshot { LocalVideoTaskSnapshot::OpenAi(seed) => Some(LocalVideoTaskReadRefreshPlan { plan: seed.build_get_follow_up_plan(trace_id)?, @@ -298,33 +391,89 @@ impl VideoTaskService { fallback_api_key_id: Option<&str>, trace_id: &str, ) -> Option { - match plan_kind { - "openai_video_remix_sync" => { - let task_id = extract_openai_task_id_from_remix_path(request_path)?; - let seed = self.store.clone_openai(task_id)?; - seed.build_remix_follow_up_plan( + let snapshot = self.snapshot_for_follow_up_route(plan_kind, request_path)?; + Self::build_follow_up_plan_from_snapshot( + snapshot, + plan_kind, + body_json, + fallback_user_id, + fallback_api_key_id, + trace_id, + ) + } + + fn build_follow_up_plan_from_snapshot( + snapshot: LocalVideoTaskSnapshot, + plan_kind: &str, + body_json: Option<&Value>, + fallback_user_id: Option<&str>, + fallback_api_key_id: Option<&str>, + trace_id: &str, + ) -> Option { + match (plan_kind, snapshot) { + ("openai_video_remix_sync", LocalVideoTaskSnapshot::OpenAi(seed)) => seed + .build_remix_follow_up_plan( body_json?, fallback_user_id, fallback_api_key_id, trace_id, - ) - } - "openai_video_delete_sync" => { - let task_id = extract_openai_task_id_from_path(request_path)?; - let seed = self.store.clone_openai(task_id)?; + ), + ("openai_video_delete_sync", LocalVideoTaskSnapshot::OpenAi(seed)) => { seed.build_delete_follow_up_plan(fallback_user_id, fallback_api_key_id, trace_id) } - "openai_video_cancel_sync" => { - let task_id = extract_openai_task_id_from_cancel_path(request_path)?; - let seed = self.store.clone_openai(task_id)?; + ("openai_video_cancel_sync", LocalVideoTaskSnapshot::OpenAi(seed)) => { seed.build_cancel_follow_up_plan(fallback_user_id, fallback_api_key_id, trace_id) } - "gemini_video_cancel_sync" => { - let short_id = extract_gemini_short_id_from_cancel_path(request_path)?; - let seed = self.store.clone_gemini(short_id)?; + ("gemini_video_cancel_sync", LocalVideoTaskSnapshot::Gemini(seed)) => { seed.build_cancel_follow_up_plan(fallback_user_id, fallback_api_key_id, trace_id) } _ => None, } } + + pub fn prepare_follow_up_sync_plan_for_user( + &self, + plan_kind: &str, + request_path: &str, + body_json: Option<&Value>, + fallback_user_id: Option<&str>, + fallback_api_key_id: Option<&str>, + trace_id: &str, + ) -> Option { + let snapshot = self.snapshot_for_follow_up_route(plan_kind, request_path)?; + let user_id = fallback_user_id?.trim(); + if !snapshot.belongs_to_user(user_id) { + return None; + } + Self::build_follow_up_plan_from_snapshot( + snapshot, + plan_kind, + body_json, + Some(user_id), + fallback_api_key_id, + trace_id, + ) + } + + fn snapshot_for_follow_up_route( + &self, + plan_kind: &str, + request_path: &str, + ) -> Option { + match plan_kind { + "openai_video_remix_sync" => extract_openai_task_id_from_remix_path(request_path) + .and_then(|task_id| self.store.clone_openai(task_id)) + .map(LocalVideoTaskSnapshot::OpenAi), + "openai_video_delete_sync" => extract_openai_task_id_from_path(request_path) + .and_then(|task_id| self.store.clone_openai(task_id)) + .map(LocalVideoTaskSnapshot::OpenAi), + "openai_video_cancel_sync" => extract_openai_task_id_from_cancel_path(request_path) + .and_then(|task_id| self.store.clone_openai(task_id)) + .map(LocalVideoTaskSnapshot::OpenAi), + "gemini_video_cancel_sync" => extract_gemini_short_id_from_cancel_path(request_path) + .and_then(|short_id| self.store.clone_gemini(short_id)) + .map(LocalVideoTaskSnapshot::Gemini), + _ => None, + } + } } diff --git a/crates/aether-video-tasks-core/src/snapshot.rs b/crates/aether-video-tasks-core/src/snapshot.rs index 150448a75..fa56551c3 100644 --- a/crates/aether-video-tasks-core/src/snapshot.rs +++ b/crates/aether-video-tasks-core/src/snapshot.rs @@ -1,6 +1,7 @@ use aether_data_contracts::repository::video_tasks::{StoredVideoTask, UpsertVideoTask}; use serde_json::{json, Map, Value}; +use crate::types::sanitize_video_task_error_code; use crate::{ local_status_from_stored, non_empty_owned, request_body_string, GeminiVideoTaskSeed, LocalVideoTaskPersistence, LocalVideoTaskReadResponse, LocalVideoTaskSnapshot, @@ -16,11 +17,26 @@ impl LocalVideoTaskSnapshot { } pub fn from_stored_task(task: &StoredVideoTask) -> Option { - task.request_metadata + let mut snapshot = task + .request_metadata .as_ref() .and_then(|metadata| metadata.get("rust_local_snapshot")) .cloned() - .and_then(|value| serde_json::from_value::(value).ok()) + .and_then(|value| serde_json::from_value::(value).ok())?; + + // The row is the ownership source of truth. Older embedded snapshots can + // contain stale identity fields after a task import or repair. + match &mut snapshot { + Self::OpenAi(seed) => { + seed.user_id = task.user_id.clone(); + seed.api_key_id = task.api_key_id.clone(); + } + Self::Gemini(seed) => { + seed.user_id = task.user_id.clone(); + seed.api_key_id = task.api_key_id.clone(); + } + } + Some(snapshot) } pub fn from_stored_task_with_transport( @@ -97,6 +113,33 @@ impl LocalVideoTaskSnapshot { } } + pub(crate) fn sanitize_persisted_diagnostics(&mut self) -> bool { + match self { + Self::OpenAi(seed) => { + let previous_error_code = seed.error_code.clone(); + let error_code = + sanitized_error_code_for_status(seed.status, seed.error_code.take()); + let changed = previous_error_code != error_code || seed.error_message.is_some(); + seed.error_code = error_code; + seed.error_message = None; + changed + } + Self::Gemini(seed) => { + let previous_error_code = seed.error_code.clone(); + let error_code = + sanitized_error_code_for_status(seed.status, seed.error_code.take()); + let safe_metadata = Value::Object(Map::new()); + let changed = previous_error_code != error_code + || seed.error_message.is_some() + || seed.metadata != safe_metadata; + seed.error_code = error_code; + seed.error_message = None; + seed.metadata = safe_metadata; + changed + } + } + } + pub fn read_response(&self) -> LocalVideoTaskReadResponse { match self { Self::OpenAi(seed) => match seed.status { @@ -130,6 +173,18 @@ impl LocalVideoTaskSnapshot { } } + pub fn belongs_to_user(&self, user_id: &str) -> bool { + let user_id = user_id.trim(); + if user_id.is_empty() { + return false; + } + let owner = match self { + Self::OpenAi(seed) => seed.user_id.as_deref(), + Self::Gemini(seed) => seed.user_id.as_deref(), + }; + owner.map(str::trim) == Some(user_id) + } + pub fn is_active_for_refresh(&self) -> bool { match self { Self::OpenAi(seed) => matches!( @@ -161,3 +216,16 @@ impl LocalVideoTaskSnapshot { } } } + +fn sanitized_error_code_for_status( + status: LocalVideoTaskStatus, + error_code: Option, +) -> Option { + match status { + LocalVideoTaskStatus::Failed => sanitize_video_task_error_code(error_code) + .or_else(|| Some("provider_error".to_string())), + LocalVideoTaskStatus::Expired => Some("expired".to_string()), + LocalVideoTaskStatus::Cancelled => Some("cancelled".to_string()), + _ => None, + } +} diff --git a/crates/aether-video-tasks-core/src/store_backend.rs b/crates/aether-video-tasks-core/src/store_backend.rs index 2a30f7e49..f2b94e5a1 100644 --- a/crates/aether-video-tasks-core/src/store_backend.rs +++ b/crates/aether-video-tasks-core/src/store_backend.rs @@ -1,24 +1,63 @@ use std::path::{Path, PathBuf}; use std::sync::Mutex; +use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext}; use serde_json::{Map, Value}; +use uuid::Uuid; use crate::{ GeminiVideoTaskSeed, LocalVideoTaskReadResponse, LocalVideoTaskRegistryMutation, LocalVideoTaskSnapshot, OpenAiVideoTaskSeed, VideoTaskRegistry, VideoTaskStore, }; -#[derive(Debug, Default)] +#[derive(Default)] pub struct InMemoryVideoTaskStore { registry: Mutex, } -#[derive(Debug)] +impl std::fmt::Debug for InMemoryVideoTaskStore { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("InMemoryVideoTaskStore") + .field("registry", &"[redacted]") + .finish() + } +} + pub struct FileVideoTaskStore { path: PathBuf, + encryption_key: String, registry: Mutex, + persisted_file: Mutex, } +#[derive(Debug, Clone, PartialEq, Eq)] +enum PersistedVideoTaskStore { + Missing, + Bytes(Vec), +} + +struct LoadedVideoTaskRegistry { + registry: VideoTaskRegistry, + persisted_file: PersistedVideoTaskStore, + needs_rewrite: bool, +} + +impl std::fmt::Debug for FileVideoTaskStore { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("FileVideoTaskStore") + .field("path", &self.path) + .field("encryption_key", &"[redacted]") + .field("registry", &"[redacted]") + .finish() + } +} + +const ENCRYPTED_VIDEO_TASK_STORE_PREFIX: &str = "aether-video-tasks-v2\n"; +const LEGACY_ENCRYPTED_VIDEO_TASK_STORE_PREFIX: &str = "aether-video-tasks-v1\n"; +const VIDEO_TASK_STORE_PURPOSE: &str = "video-task-file-store"; + impl VideoTaskStore for InMemoryVideoTaskStore { fn insert(&self, snapshot: LocalVideoTaskSnapshot) { if let Ok(mut registry) = self.registry.lock() { @@ -75,36 +114,107 @@ impl VideoTaskStore for InMemoryVideoTaskStore { } impl FileVideoTaskStore { - pub fn new(path: impl Into) -> std::io::Result { + pub fn new( + path: impl Into, + encryption_key: impl Into, + ) -> std::io::Result { let path = path.into(); - let registry = Self::load_registry(&path)?; - Ok(Self { + let encryption_key = encryption_key.into(); + if encryption_key.trim().is_empty() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "video task file store encryption key cannot be empty", + )); + } + let LoadedVideoTaskRegistry { + registry, + persisted_file, + needs_rewrite, + } = Self::load_registry(&path, &encryption_key)?; + let store = Self { path, + encryption_key, registry: Mutex::new(registry), - }) + persisted_file: Mutex::new(persisted_file), + }; + if needs_rewrite { + let registry = store + .registry + .lock() + .map_err(|_| std::io::Error::other("video task store lock poisoned"))?; + store.persist_registry(®istry)?; + } + Ok(store) } - fn load_registry(path: &Path) -> std::io::Result { - if !path.exists() { - return Ok(VideoTaskRegistry::default()); - } - let bytes = std::fs::read(path)?; + fn load_registry( + path: &Path, + encryption_key: &str, + ) -> std::io::Result { + let persisted_file = read_persisted_video_task_store(path)?; + let PersistedVideoTaskStore::Bytes(bytes) = &persisted_file else { + return Ok(LoadedVideoTaskRegistry { + registry: VideoTaskRegistry::default(), + persisted_file, + needs_rewrite: false, + }); + }; if bytes.is_empty() { - return Ok(VideoTaskRegistry::default()); + return Err(invalid_store_data("video task store is empty")); } - serde_json::from_slice(&bytes) - .map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err)) + if let Some(ciphertext) = bytes.strip_prefix(ENCRYPTED_VIDEO_TASK_STORE_PREFIX.as_bytes()) { + let ciphertext = std::str::from_utf8(ciphertext) + .map_err(|_| invalid_store_data("encrypted video task store is not UTF-8"))?; + let protected = decrypt_python_fernet_ciphertext(encryption_key, ciphertext.trim()) + .map_err(|_| invalid_store_data("video task store decryption failed"))?; + let plaintext = protected + .strip_prefix(VIDEO_TASK_STORE_PURPOSE) + .and_then(|value| value.strip_prefix('\0')) + .ok_or_else(|| invalid_store_data("video task store purpose mismatch"))?; + let mut registry: VideoTaskRegistry = serde_json::from_str(plaintext) + .map_err(|_| invalid_store_data("decrypted video task store is invalid"))?; + let needs_rewrite = registry.sanitize_persisted_diagnostics(); + return Ok(LoadedVideoTaskRegistry { + registry, + persisted_file, + needs_rewrite, + }); + } + if let Some(ciphertext) = + bytes.strip_prefix(LEGACY_ENCRYPTED_VIDEO_TASK_STORE_PREFIX.as_bytes()) + { + let ciphertext = std::str::from_utf8(ciphertext) + .map_err(|_| invalid_store_data("encrypted video task store is not UTF-8"))?; + let plaintext = decrypt_python_fernet_ciphertext(encryption_key, ciphertext.trim()) + .map_err(|_| invalid_store_data("video task store decryption failed"))?; + let mut registry: VideoTaskRegistry = serde_json::from_str(&plaintext) + .map_err(|_| invalid_store_data("decrypted video task store is invalid"))?; + registry.sanitize_persisted_diagnostics(); + return Ok(LoadedVideoTaskRegistry { + registry, + persisted_file, + needs_rewrite: true, + }); + } + + if bytes.starts_with(b"aether-") { + return Err(invalid_store_data( + "unsupported encrypted video task store envelope", + )); + } + Err(invalid_store_data( + "plaintext video task stores are not accepted", + )) } fn persist_registry(&self, registry: &VideoTaskRegistry) -> std::io::Result<()> { - if let Some(parent) = self.path.parent() { - std::fs::create_dir_all(parent)?; - } - let bytes = serde_json::to_vec_pretty(registry) - .map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err))?; - let temp_path = self.path.with_extension("tmp"); - std::fs::write(&temp_path, bytes)?; - std::fs::rename(temp_path, &self.path)?; + let bytes = encrypted_video_task_store_bytes(&self.encryption_key, registry)?; + let mut persisted_file = self + .persisted_file + .lock() + .map_err(|_| std::io::Error::other("video task persisted-file lock poisoned"))?; + replace_video_task_store_if_unchanged(&self.path, &persisted_file, &bytes)?; + *persisted_file = PersistedVideoTaskStore::Bytes(bytes); Ok(()) } @@ -112,13 +222,117 @@ impl FileVideoTaskStore { let Ok(mut registry) = self.registry.lock() else { return false; }; - if !mutator(&mut registry) { + let mut updated_registry = registry.clone(); + if !mutator(&mut updated_registry) { return false; } - self.persist_registry(®istry).is_ok() + if self.persist_registry(&updated_registry).is_err() { + return false; + } + *registry = updated_registry; + true } } +fn invalid_store_data(message: &'static str) -> std::io::Error { + std::io::Error::new(std::io::ErrorKind::InvalidData, message) +} + +fn encrypted_video_task_store_bytes( + encryption_key: &str, + registry: &VideoTaskRegistry, +) -> std::io::Result> { + let plaintext = serde_json::to_string(registry) + .map_err(|_| invalid_store_data("video task store serialization failed"))?; + let protected = format!("{VIDEO_TASK_STORE_PURPOSE}\0{plaintext}"); + let ciphertext = encrypt_python_fernet_plaintext(encryption_key, &protected) + .map_err(|_| invalid_store_data("video task store encryption failed"))?; + let mut bytes = ENCRYPTED_VIDEO_TASK_STORE_PREFIX.as_bytes().to_vec(); + bytes.extend_from_slice(ciphertext.as_bytes()); + bytes.push(b'\n'); + Ok(bytes) +} + +fn read_persisted_video_task_store(path: &Path) -> std::io::Result { + match std::fs::read(path) { + Ok(bytes) => Ok(PersistedVideoTaskStore::Bytes(bytes)), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + Ok(PersistedVideoTaskStore::Missing) + } + Err(error) => Err(error), + } +} + +fn replace_video_task_store_if_unchanged( + path: &Path, + expected: &PersistedVideoTaskStore, + replacement: &[u8], +) -> std::io::Result<()> { + if let Some(parent) = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + { + std::fs::create_dir_all(parent)?; + } + + let lock_path = video_task_store_lock_path(path); + let lock_file = open_private_lock_file(&lock_path)?; + // Lock a stable sidecar inode: the store inode itself is replaced by rename. + lock_file.lock()?; + + // The bytes captured before parsing/decryption are the migration/write CAS token. + let observed = read_persisted_video_task_store(path)?; + if &observed != expected { + return Err(std::io::Error::new( + std::io::ErrorKind::WouldBlock, + "video task store changed before compare-and-replace", + )); + } + + let temp_path = path.with_extension(format!("tmp-{}", Uuid::new_v4())); + if let Err(error) = write_private_file(&temp_path, replacement) { + let _ = std::fs::remove_file(&temp_path); + return Err(error); + } + if let Err(error) = std::fs::rename(&temp_path, path) { + let _ = std::fs::remove_file(&temp_path); + return Err(error); + } + Ok(()) +} + +fn video_task_store_lock_path(path: &Path) -> PathBuf { + let mut lock_path = path.as_os_str().to_os_string(); + lock_path.push(".lock"); + PathBuf::from(lock_path) +} + +fn open_private_lock_file(path: &Path) -> std::io::Result { + let mut options = std::fs::OpenOptions::new(); + options.read(true).write(true).create(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt as _; + options.mode(0o600); + } + options.open(path) +} + +fn write_private_file(path: &Path, bytes: &[u8]) -> std::io::Result<()> { + use std::io::Write as _; + + let mut options = std::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt as _; + options.mode(0o600); + } + let mut file = options.open(path)?; + file.write_all(bytes)?; + file.sync_all() +} + impl VideoTaskStore for FileVideoTaskStore { fn insert(&self, snapshot: LocalVideoTaskSnapshot) { let _ = self.mutate_registry(|registry| { @@ -169,3 +383,283 @@ impl VideoTaskStore for FileVideoTaskStore { self.mutate_registry(|registry| registry.project_gemini(short_id, provider_body)) } } + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use serde_json::json; + + use super::*; + use crate::{LocalVideoTaskPersistence, LocalVideoTaskStatus, LocalVideoTaskTransport}; + + fn temp_store_path(name: &str) -> PathBuf { + std::env::temp_dir().join(format!( + "aether-video-store-{name}-{}-{}.json", + std::process::id(), + Uuid::new_v4() + )) + } + + fn cleanup_store_path(path: &Path) { + std::fs::remove_file(path).ok(); + std::fs::remove_file(video_task_store_lock_path(path)).ok(); + } + + fn sensitive_gemini_snapshot() -> LocalVideoTaskSnapshot { + LocalVideoTaskSnapshot::Gemini(GeminiVideoTaskSeed { + local_short_id: "task-sensitive".to_string(), + upstream_operation_name: "operations/upstream-sensitive".to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("api-key-1".to_string()), + model: "veo-3".to_string(), + status: LocalVideoTaskStatus::Failed, + progress_percent: 100, + error_code: Some("Bearer code-secret".to_string()), + error_message: Some("Authorization: Bearer error-secret".to_string()), + metadata: json!({ + "debug": "metadata-secret", + "url": "https://internal.test/result?token=metadata-query-secret" + }), + persistence: LocalVideoTaskPersistence { + request_id: "request-1".to_string(), + username: Some("alice".to_string()), + api_key_name: Some("primary".to_string()), + client_api_format: "gemini:video".to_string(), + provider_api_format: "gemini:video".to_string(), + original_request_body: json!({"prompt": "create a video"}), + format_converted: false, + }, + transport: LocalVideoTaskTransport { + upstream_base_url: "https://generativelanguage.googleapis.com".to_string(), + provider_name: Some("gemini".to_string()), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-1".to_string(), + headers: BTreeMap::from([( + "x-goog-api-key".to_string(), + "transport-key-required-for-resume".to_string(), + )]), + content_type: Some("application/json".to_string()), + model_name: Some("veo-3".to_string()), + proxy: None, + transport_profile: None, + timeouts: None, + }, + }) + } + + #[test] + fn loading_encrypted_store_rewrites_legacy_provider_diagnostics() { + let path = temp_store_path("diagnostic-migration"); + let legacy_plaintext = serde_json::to_string(&json!({ + "openai": {}, + "gemini": { + "task-sensitive": sensitive_gemini_snapshot() + } + })) + .expect("legacy registry should serialize"); + let ciphertext = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &legacy_plaintext) + .expect("legacy registry should encrypt"); + std::fs::write( + &path, + format!("{LEGACY_ENCRYPTED_VIDEO_TASK_STORE_PREFIX}{ciphertext}\n"), + ) + .expect("legacy encrypted registry should be written"); + + let store = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect("legacy encrypted registry should load"); + + let response = store + .read_gemini("task-sensitive") + .expect("migrated task should remain readable"); + assert_eq!(response.body_json["error"]["code"], "provider_error"); + assert_eq!( + response.body_json["error"]["message"], + "Video generation failed" + ); + + let bytes = std::fs::read(&path).expect("migrated registry should be readable"); + let ciphertext = bytes + .strip_prefix(ENCRYPTED_VIDEO_TASK_STORE_PREFIX.as_bytes()) + .expect("migrated registry should stay encrypted"); + let ciphertext = std::str::from_utf8(ciphertext) + .expect("ciphertext should be UTF-8") + .trim(); + let migrated_plaintext = + decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, ciphertext) + .expect("migrated registry should decrypt"); + for secret in [ + "code-secret", + "error-secret", + "metadata-secret", + "metadata-query-secret", + ] { + assert!( + !migrated_plaintext.contains(secret), + "migrated registry leaked {secret}" + ); + } + assert!(migrated_plaintext.contains("transport-key-required-for-resume")); + assert!(migrated_plaintext.contains("provider_error")); + cleanup_store_path(&path); + } + + #[test] + fn rejects_plaintext_registry_without_rewriting_it() { + let path = temp_store_path("plaintext-injection"); + let plaintext = serde_json::to_vec(&json!({ + "openai": {}, + "gemini": {}, + })) + .expect("plaintext registry should serialize"); + std::fs::write(&path, &plaintext).expect("plaintext registry should be written"); + + let error = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect_err("plaintext registry must fail closed"); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(error.to_string().contains("plaintext")); + assert_eq!( + std::fs::read(&path).expect("rejected registry should remain readable"), + plaintext, + "rejected plaintext must not be rewritten into an authenticated envelope", + ); + cleanup_store_path(&path); + } + + #[test] + fn rejects_unknown_aether_envelopes_without_rewriting_them() { + for (name, bytes) in [ + ( + "unknown-video-version", + b"aether-video-tasks-v999\nopaque".as_slice(), + ), + ( + "foreign-aether-envelope", + b"aether-other-v1\nopaque".as_slice(), + ), + ] { + let path = temp_store_path(name); + std::fs::write(&path, bytes).expect("unknown envelope should be written"); + + let error = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect_err("unknown aether envelope must fail closed"); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(error.to_string().contains("unsupported")); + assert_eq!( + std::fs::read(&path).expect("rejected envelope should remain readable"), + bytes, + ); + cleanup_store_path(&path); + } + } + + #[test] + fn rejects_tampered_authenticated_store_without_rewriting_it() { + let path = temp_store_path("tampered-v2"); + let mut bytes = encrypted_video_task_store_bytes( + DEVELOPMENT_ENCRYPTION_KEY, + &VideoTaskRegistry::default(), + ) + .expect("encrypted registry should serialize"); + let tampered_index = ENCRYPTED_VIDEO_TASK_STORE_PREFIX.len() + 12; + bytes[tampered_index] = if bytes[tampered_index] == b'A' { + b'B' + } else { + b'A' + }; + std::fs::write(&path, &bytes).expect("tampered registry should be written"); + + let error = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect_err("tampered authenticated registry must fail closed"); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(error.to_string().contains("decryption failed")); + assert_eq!( + std::fs::read(&path).expect("tampered registry should remain readable"), + bytes, + ); + cleanup_store_path(&path); + } + + #[test] + fn legacy_migration_compare_before_replace_preserves_concurrent_replacement() { + let path = temp_store_path("legacy-migration-race"); + let legacy_plaintext = serde_json::to_string(&json!({ + "openai": {}, + "gemini": { + "task-sensitive": sensitive_gemini_snapshot(), + }, + })) + .expect("legacy registry should serialize"); + let legacy_ciphertext = + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &legacy_plaintext) + .expect("legacy registry should encrypt"); + let legacy_bytes = + format!("{LEGACY_ENCRYPTED_VIDEO_TASK_STORE_PREFIX}{legacy_ciphertext}\n").into_bytes(); + std::fs::write(&path, &legacy_bytes).expect("legacy registry should be written"); + + let loaded = FileVideoTaskStore::load_registry(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect("authenticated legacy registry should load"); + assert!(loaded.needs_rewrite); + assert_eq!( + loaded.persisted_file, + PersistedVideoTaskStore::Bytes(legacy_bytes), + ); + let store = FileVideoTaskStore { + path: path.clone(), + encryption_key: DEVELOPMENT_ENCRYPTION_KEY.to_string(), + registry: Mutex::new(loaded.registry), + persisted_file: Mutex::new(loaded.persisted_file), + }; + + let concurrent_replacement = encrypted_video_task_store_bytes( + DEVELOPMENT_ENCRYPTION_KEY, + &VideoTaskRegistry::default(), + ) + .expect("replacement registry should encrypt"); + std::fs::write(&path, &concurrent_replacement) + .expect("concurrent replacement should be written"); + + let registry = store.registry.lock().expect("registry lock should succeed"); + let error = store + .persist_registry(®istry) + .expect_err("stale legacy migration must report a conflict"); + drop(registry); + + assert_eq!(error.kind(), std::io::ErrorKind::WouldBlock); + assert_eq!( + std::fs::read(&path).expect("concurrent replacement should remain readable"), + concurrent_replacement, + "stale migration must not overwrite a newer exact byte snapshot", + ); + cleanup_store_path(&path); + } + + #[test] + fn rejects_valid_fernet_ciphertext_from_another_purpose() { + let path = temp_store_path("purpose-mismatch"); + let protected = format!( + "another-purpose\0{}", + serde_json::to_string(&VideoTaskRegistry::default()) + .expect("empty registry should serialize") + ); + let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &protected) + .expect("foreign payload should encrypt"); + std::fs::write( + &path, + format!("{ENCRYPTED_VIDEO_TASK_STORE_PREFIX}{ciphertext}\n"), + ) + .expect("foreign ciphertext should be written"); + + let error = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY) + .expect_err("foreign-purpose ciphertext must fail closed"); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + cleanup_store_path(&path); + } +} diff --git a/crates/aether-video-tasks-core/src/store_registry.rs b/crates/aether-video-tasks-core/src/store_registry.rs index 5e3bc8dad..90f8f8fbe 100644 --- a/crates/aether-video-tasks-core/src/store_registry.rs +++ b/crates/aether-video-tasks-core/src/store_registry.rs @@ -15,7 +15,8 @@ pub struct VideoTaskRegistry { } impl VideoTaskRegistry { - pub fn insert(&mut self, snapshot: LocalVideoTaskSnapshot) { + pub fn insert(&mut self, mut snapshot: LocalVideoTaskSnapshot) { + snapshot.sanitize_persisted_diagnostics(); match &snapshot { LocalVideoTaskSnapshot::OpenAi(seed) => { self.openai.insert(seed.local_task_id.clone(), snapshot); @@ -97,4 +98,12 @@ impl VideoTaskRegistry { seed.apply_provider_body(provider_body); true } + + pub(crate) fn sanitize_persisted_diagnostics(&mut self) -> bool { + let mut changed = false; + for snapshot in self.openai.values_mut().chain(self.gemini.values_mut()) { + changed = snapshot.sanitize_persisted_diagnostics() || changed; + } + changed + } } diff --git a/crates/aether-video-tasks-core/src/types.rs b/crates/aether-video-tasks-core/src/types.rs index 13197635a..81d8a29f1 100644 --- a/crates/aether-video-tasks-core/src/types.rs +++ b/crates/aether-video-tasks-core/src/types.rs @@ -89,7 +89,7 @@ pub enum LocalVideoTaskSeed { GeminiCreate(GeminiVideoTaskSeed), } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Clone, PartialEq, Serialize, Deserialize)] pub struct LocalVideoTaskTransport { pub upstream_base_url: String, pub provider_name: Option, @@ -104,7 +104,52 @@ pub struct LocalVideoTaskTransport { pub timeouts: Option, } -#[derive(Debug, Clone, PartialEq)] +impl std::fmt::Debug for LocalVideoTaskTransport { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("LocalVideoTaskTransport") + .field("upstream_base_url", &"[redacted]") + .field("provider_name", &self.provider_name) + .field("provider_id", &self.provider_id) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field("headers", &"[redacted]") + .field("content_type", &self.content_type) + .field("model_name", &self.model_name) + .field("proxy", &self.proxy.as_ref().map(|_| "[redacted]")) + .field( + "transport_profile", + &self.transport_profile.as_ref().map(|_| "[redacted]"), + ) + .field("timeouts", &self.timeouts) + .finish() + } +} + +pub(crate) fn sanitize_video_task_error_code(value: Option) -> Option { + let value = value?.trim().to_ascii_lowercase(); + if value.is_empty() { + return None; + } + Some(match value.as_str() { + "authentication_error" + | "cancelled" + | "content_policy_violation" + | "expired" + | "invalid_request" + | "not_found" + | "permission_denied" + | "poll_permanent_error" + | "poll_timeout" + | "provider_error" + | "rate_limit_exceeded" + | "server_error" + | "unknown" => value, + _ => "provider_error".to_string(), + }) +} + +#[derive(Clone, PartialEq)] pub struct LocalVideoTaskTransportBridgeInput { pub upstream_base_url: String, pub provider_name: Option, @@ -120,6 +165,29 @@ pub struct LocalVideoTaskTransportBridgeInput { pub timeouts: Option, } +impl std::fmt::Debug for LocalVideoTaskTransportBridgeInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("LocalVideoTaskTransportBridgeInput") + .field("upstream_base_url", &"[redacted]") + .field("provider_name", &self.provider_name) + .field("provider_id", &self.provider_id) + .field("endpoint_id", &self.endpoint_id) + .field("key_id", &self.key_id) + .field("auth_header", &"[redacted]") + .field("auth_value", &"[redacted]") + .field("content_type", &self.content_type) + .field("model_name", &self.model_name) + .field("proxy", &self.proxy.as_ref().map(|_| "[redacted]")) + .field( + "transport_profile", + &self.transport_profile.as_ref().map(|_| "[redacted]"), + ) + .field("timeouts", &self.timeouts) + .finish() + } +} + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct LocalVideoTaskPersistence { pub request_id: String, diff --git a/deploy.sh b/deploy.sh index 509523d34..a8e833de6 100755 --- a/deploy.sh +++ b/deploy.sh @@ -6,7 +6,8 @@ # 强制全部重建: ./deploy.sh --force set -euo pipefail -cd "$(dirname "$0")" +SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)" +cd -- "${SCRIPT_DIR}" LOCAL_APP_IMAGE="${LOCAL_APP_IMAGE:-aether-app:latest}" export LOCAL_APP_IMAGE @@ -48,6 +49,17 @@ compose_up() { # 缓存文件 CODE_HASH_FILE=".code-hash" +validate_code_hash_file() { + if [ -L "${CODE_HASH_FILE}" ]; then + echo "Refusing symbolic-link build state: ${CODE_HASH_FILE}" >&2 + return 1 + fi + if [ -e "${CODE_HASH_FILE}" ] && [ ! -f "${CODE_HASH_FILE}" ]; then + echo "Build state is not a regular file: ${CODE_HASH_FILE}" >&2 + return 1 + fi +} + usage() { cat <<'EOF' Usage: ./deploy.sh [options] @@ -153,9 +165,10 @@ calc_code_hash() { check_code_changed() { local current_hash current_hash=$(calc_code_hash) + validate_code_hash_file || exit 1 if [ -f "$CODE_HASH_FILE" ]; then local saved_hash - saved_hash=$(cat "$CODE_HASH_FILE") + saved_hash=$(<"$CODE_HASH_FILE") if [ "$current_hash" = "$saved_hash" ]; then return 1 fi @@ -163,7 +176,24 @@ check_code_changed() { return 0 } -save_code_hash() { calc_code_hash > "$CODE_HASH_FILE"; } +save_code_hash() { + local staged + staged="$(mktemp "./.code-hash.tmp.XXXXXXXX")" \ + || { echo "Could not create temporary build state" >&2; return 1; } + if ! calc_code_hash >"${staged}"; then + rm -f -- "${staged}" + return 1 + fi + chmod 0600 "${staged}" + if ! validate_code_hash_file; then + rm -f -- "${staged}" + return 1 + fi + if ! mv -f -- "${staged}" "${CODE_HASH_FILE}"; then + rm -f -- "${staged}" + return 1 + fi +} # 构建应用镜像 build_app() { @@ -228,6 +258,6 @@ docker image prune -f >/dev/null 2>&1 || true echo ">>> Done!" echo ">>> Note: empty databases auto-bootstrap on first start." -echo ">>> Note: docker compose now defaults to auto-running pending migrations/backfills on app startup." -echo ">>> Note: set AETHER_GATEWAY_AUTO_PREPARE_DATABASE=false if you want to keep manual rollout." +echo ">>> Note: database schema and data preparation run automatically before app startup." +echo ">>> Note: set AETHER_GATEWAY_DATABASE_MODE=verify-only to require a separate database prepare step." "${DC[@]}" ps diff --git a/docker-compose.single-node.yml b/docker-compose.single-node.yml index 66a67f840..dab9b5e28 100644 --- a/docker-compose.single-node.yml +++ b/docker-compose.single-node.yml @@ -2,7 +2,14 @@ services: app: image: ${APP_IMAGE:-ghcr.io/fawney19/aether:latest} container_name: aether-app - user: "0:0" + user: "${AETHER_CONTAINER_UID:-65532}:${AETHER_CONTAINER_GID:-65532}" + read_only: true + cap_drop: + - ALL + security_opt: + - no-new-privileges:true + tmpfs: + - /tmp:rw,nosuid,nodev,noexec,mode=1777 env_file: - ${AETHER_ENV_FILE:-.env} environment: diff --git a/docker-compose.yml b/docker-compose.yml index 733290de7..678975106 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -3,13 +3,14 @@ services: postgres: - image: postgres:15 + # OCI index digest for the official multi-architecture postgres:15.19 image. + image: postgres:15.19@sha256:5f72c7b5bd616308ccfd2e74d6be16fb06364e5eecbb815fe9dc6ab9761d2111 container_name: aether-postgres shm_size: ${POSTGRES_SHM_SIZE:-512mb} environment: POSTGRES_DB: aether POSTGRES_USER: postgres - POSTGRES_PASSWORD: ${DB_PASSWORD} + POSTGRES_PASSWORD: ${DB_PASSWORD:?set DB_PASSWORD in .env} TZ: Asia/Shanghai volumes: - postgres_data:/var/lib/postgresql/data @@ -41,28 +42,30 @@ services: restart: unless-stopped redis: - image: redis:7-alpine + # OCI index digest for the official multi-architecture redis:7.4.11-alpine image. + image: redis:7.4.11-alpine@sha256:ff02b58f971e7d7d156a1267e283fcbbeee91773b6aa36c49dac28ecfe28eadf container_name: aether-redis - command: redis-server --dir /tmp --appendonly no --save "" --requirepass ${REDIS_PASSWORD} --maxclients ${REDIS_MAXCLIENTS:-10000} + command: redis-server --dir /tmp --appendonly no --save "" --requirepass ${REDIS_PASSWORD:?set REDIS_PASSWORD in .env} --maxclients ${REDIS_MAXCLIENTS:-10000} ports: - "127.0.0.1:${REDIS_PORT:-6379}:6379" healthcheck: - test: [ "CMD-SHELL", "redis-cli -a \"${REDIS_PASSWORD}\" ping | grep -q PONG" ] + test: [ "CMD-SHELL", "redis-cli -a \"${REDIS_PASSWORD:?set REDIS_PASSWORD in .env}\" ping | grep -q PONG" ] interval: 5s timeout: 3s retries: 5 restart: unless-stopped mysql: - image: mysql:8.0 + # OCI index digest for the official multi-architecture mysql:8.0.46 image. + image: mysql:8.0.46@sha256:7dcddc01f13bab2f15cde676d44d01f61fc9f99fe7785e86196dfc07d358ae2b container_name: aether-mysql profiles: - mysql environment: MYSQL_DATABASE: ${MYSQL_DATABASE:-aether} MYSQL_USER: ${MYSQL_USER:-aether} - MYSQL_PASSWORD: ${MYSQL_PASSWORD:-aether} - MYSQL_ROOT_PASSWORD: ${MYSQL_ROOT_PASSWORD:-aether_root} + MYSQL_PASSWORD: ${MYSQL_PASSWORD:?set MYSQL_PASSWORD in .env} + MYSQL_ROOT_PASSWORD: ${MYSQL_ROOT_PASSWORD:?set MYSQL_ROOT_PASSWORD in .env} TZ: Asia/Shanghai volumes: - mysql_data:/var/lib/mysql @@ -80,12 +83,19 @@ services: app: image: ${APP_IMAGE:-ghcr.io/fawney19/aether:latest} container_name: aether-app - user: "0:0" + user: "${AETHER_CONTAINER_UID:-65532}:${AETHER_CONTAINER_GID:-65532}" + read_only: true + cap_drop: + - ALL + security_opt: + - no-new-privileges:true + tmpfs: + - /tmp:rw,nosuid,nodev,noexec,mode=1777 env_file: - .env environment: - DATABASE_URL: postgresql://postgres:${DB_PASSWORD}@postgres:5432/aether - REDIS_URL: redis://:${REDIS_PASSWORD}@redis:6379/0 + DATABASE_URL: postgresql://postgres:${DB_PASSWORD:?set DB_PASSWORD in .env}@postgres:5432/aether + REDIS_URL: redis://:${REDIS_PASSWORD:?set REDIS_PASSWORD in .env}@redis:6379/0 TZ: Asia/Shanghai AETHER_BASE_DIR: /opt/aether AETHER_UPDATE_STRATEGY: docker diff --git a/frontend/package-lock.json b/frontend/package-lock.json index e9e826bfd..07fafcf26 100644 --- a/frontend/package-lock.json +++ b/frontend/package-lock.json @@ -13,13 +13,13 @@ "@types/marked": "^5.0.2", "@types/three": "^0.180.0", "@vueuse/core": "^13.9.0", - "axios": "^1.12.1", + "axios": "^1.20.0", "chart.js": "^4.5.0", "chartjs-adapter-date-fns": "^3.0.0", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "date-fns": "^4.1.0", - "dompurify": "^3.3.0", + "dompurify": "^3.4.14", "highlight.js": "^11.11.1", "lucide-vue-next": "^0.544.0", "marked": "^16.0.0", @@ -39,20 +39,20 @@ "@typescript-eslint/eslint-plugin": "^8.47.0", "@typescript-eslint/parser": "^8.47.0", "@vitejs/plugin-vue": "^6.0.1", - "@vitest/ui": "^4.0.10", + "@vitest/ui": "^4.1.11", "@vue/tsconfig": "^0.7.0", "autoprefixer": "^10.4.21", "baseline-browser-mapping": "^2.9.4", "eslint": "^9.39.1", "eslint-plugin-vue": "^10.5.1", "jsdom": "^27.2.0", - "postcss": "^8.5.6", + "postcss": "^8.5.26", "tailwindcss": "^3.4.17", "tailwindcss-animate": "^1.0.7", "typescript": "~5.8.3", "typescript-eslint": "^8.49.0", - "vite": "^7.1.2", - "vitest": "^4.0.10", + "vite": "^7.3.6", + "vitest": "^4.1.11", "vue-tsc": "^3.0.5" } }, @@ -265,6 +265,7 @@ } ], "license": "MIT", + "peer": true, "engines": { "node": ">=18" }, @@ -308,6 +309,7 @@ } ], "license": "MIT", + "peer": true, "engines": { "node": ">=18" } @@ -805,9 +807,9 @@ } }, "node_modules/@eslint/config-array/node_modules/brace-expansion": { - "version": "1.1.14", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.14.tgz", - "integrity": "sha512-MWPGfDxnyzKU7rNOW9SP/c50vi3xrmrua/+6hfPbCS2ABNWfx24vPidzvC7krjU/RTo235sV776ymlsMtGKj8g==", + "version": "1.1.18", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.18.tgz", + "integrity": "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw==", "dev": true, "license": "MIT", "dependencies": { @@ -879,9 +881,9 @@ } }, "node_modules/@eslint/eslintrc/node_modules/brace-expansion": { - "version": "1.1.14", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.14.tgz", - "integrity": "sha512-MWPGfDxnyzKU7rNOW9SP/c50vi3xrmrua/+6hfPbCS2ABNWfx24vPidzvC7krjU/RTo235sV776ymlsMtGKj8g==", + "version": "1.1.18", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.18.tgz", + "integrity": "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw==", "dev": true, "license": "MIT", "dependencies": { @@ -1569,9 +1571,9 @@ ] }, "node_modules/@standard-schema/spec": { - "version": "1.0.0", - "resolved": "https://registry.npmjs.org/@standard-schema/spec/-/spec-1.0.0.tgz", - "integrity": "sha512-m2bOd0f2RT9k8QJx1JN85cZYyH1RqFBdlwtkSlf4tBDYLCiiZnv1fIIwacK6cqwXavOydf0NPToMQgpKq+dVlA==", + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@standard-schema/spec/-/spec-1.1.0.tgz", + "integrity": "sha512-l2aFy5jALhniG5HgqrD6jXLi/rUWrKvqN/qJx6yoJsgKhblVd+iqqU4RCXavm/jPityDo5TCvKMnpjKnOriy0w==", "dev": true, "license": "MIT" }, @@ -1678,6 +1680,7 @@ "integrity": "sha512-GKBNHjoNw3Kra1Qg5UXttsY5kiWMEfoHq2TmXb+b1rcm6N7B3wTrFYIf/oSZ1xNQ+hVVijgLkiDZh7jRRsh+Gw==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { "undici-types": "~7.10.0" } @@ -1756,6 +1759,7 @@ "integrity": "sha512-N9lBGA9o9aqb1hVMc9hzySbhKibHmB+N3IpoShyV6HyQYRGIhlrO5rQgttypi+yEeKsKI4idxC8Jw6gXKD4THA==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { "@typescript-eslint/scope-manager": "8.49.0", "@typescript-eslint/types": "8.49.0", @@ -1972,31 +1976,31 @@ } }, "node_modules/@vitest/expect": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.0.10.tgz", - "integrity": "sha512-3QkTX/lK39FBNwARCQRSQr0TP9+ywSdxSX+LgbJ2M1WmveXP72anTbnp2yl5fH+dU6SUmBzNMrDHs80G8G2DZg==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.11.tgz", + "integrity": "sha512-VX2x5vNJXET47KAFzwERI+KRMtTTCSWTfSMKsW7JsUsXV4psq++e3DvZpuTDOpHcxytiDs6p2nhVb2tVDiiUYw==", "dev": true, "license": "MIT", "dependencies": { - "@standard-schema/spec": "^1.0.0", + "@standard-schema/spec": "^1.1.0", "@types/chai": "^5.2.2", - "@vitest/spy": "4.0.10", - "@vitest/utils": "4.0.10", - "chai": "^6.2.1", - "tinyrainbow": "^3.0.3" + "@vitest/spy": "4.1.11", + "@vitest/utils": "4.1.11", + "chai": "^6.2.2", + "tinyrainbow": "^3.1.0" }, "funding": { "url": "https://opencollective.com/vitest" } }, "node_modules/@vitest/mocker": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.0.10.tgz", - "integrity": "sha512-e2OfdexYkjkg8Hh3L9NVEfbwGXq5IZbDovkf30qW2tOh7Rh9sVtmSr2ztEXOFbymNxS4qjzLXUQIvATvN4B+lg==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.11.tgz", + "integrity": "sha512-2XJVD55d1o5AZous5CCGKS74g/riOj9odEt2bQpCVZeblHyHdnMeFl4jl0XjU21stf4mbjUkew2eXQZt65g5CQ==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/spy": "4.0.10", + "@vitest/spy": "4.1.11", "estree-walker": "^3.0.3", "magic-string": "^0.30.21" }, @@ -2005,7 +2009,7 @@ }, "peerDependencies": { "msw": "^2.4.9", - "vite": "^6.0.0 || ^7.0.0-0" + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0" }, "peerDependenciesMeta": { "msw": { @@ -2027,26 +2031,26 @@ } }, "node_modules/@vitest/pretty-format": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.0.10.tgz", - "integrity": "sha512-99EQbpa/zuDnvVjthwz5bH9o8iPefoQZ63WV8+bsRJZNw3qQSvSltfut8yu1Jc9mqOYi7pEbsKxYTi/rjaq6PA==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.11.tgz", + "integrity": "sha512-yiZzPbGTS9Sr/JpFl8zHrcIkAofNbFV6k21vIgQN/cY/oxZeXhJv5sc/MBJ5jFKWmWs+oJHw0UXLZjmf931+Vw==", "dev": true, "license": "MIT", "dependencies": { - "tinyrainbow": "^3.0.3" + "tinyrainbow": "^3.1.0" }, "funding": { "url": "https://opencollective.com/vitest" } }, "node_modules/@vitest/runner": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.0.10.tgz", - "integrity": "sha512-EXU2iSkKvNwtlL8L8doCpkyclw0mc/t4t9SeOnfOFPyqLmQwuceMPA4zJBa6jw0MKsZYbw7kAn+gl7HxrlB8UQ==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.11.tgz", + "integrity": "sha512-LztvUgdwMNJMIkj3hQnnxiC2Xy1zNxq928W/xhjCLaNCzqTZOudjwbQf6v9IntZGPw132i2Lq2rgTRZHD3JHNw==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/utils": "4.0.10", + "@vitest/utils": "4.1.11", "pathe": "^2.0.3" }, "funding": { @@ -2054,13 +2058,14 @@ } }, "node_modules/@vitest/snapshot": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.0.10.tgz", - "integrity": "sha512-2N4X2ZZl7kZw0qeGdQ41H0KND96L3qX1RgwuCfy6oUsF2ISGD/HpSbmms+CkIOsQmg2kulwfhJ4CI0asnZlvkg==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.11.tgz", + "integrity": "sha512-pN7ikn1ON7h8ee4gIAp4AzyK+zBtJPzVbqOgu5LCEh4VaJVbPQcgYQYJIMGQPXVeJJq1fnfazis7a5pFNPahog==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/pretty-format": "4.0.10", + "@vitest/pretty-format": "4.1.11", + "@vitest/utils": "4.1.11", "magic-string": "^0.30.21", "pathe": "^2.0.3" }, @@ -2069,9 +2074,9 @@ } }, "node_modules/@vitest/spy": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.0.10.tgz", - "integrity": "sha512-AsY6sVS8OLb96GV5RoG8B6I35GAbNrC49AO+jNRF9YVGb/g9t+hzNm1H6kD0NDp8tt7VJLs6hb7YMkDXqu03iw==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.11.tgz", + "integrity": "sha512-apNa/prQy2qCeywhnixOHPRCgGNhvg7T4Dapfl1GahLp/R+uhBm5cPyFoNVyqsNd2h1nJxL6BqqdIjiABL60YA==", "dev": true, "license": "MIT", "funding": { @@ -2079,36 +2084,38 @@ } }, "node_modules/@vitest/ui": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/ui/-/ui-4.0.10.tgz", - "integrity": "sha512-oWtNM89Np+YsQO3ttT5i1Aer/0xbzQzp66NzuJn/U16bB7MnvSzdLKXgk1kkMLYyKSSzA2ajzqMkYheaE9opuQ==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/ui/-/ui-4.1.11.tgz", + "integrity": "sha512-r/rwyKoev21mWdRGSEkZOqkQ2BYy68mwjihg9M90nNRbf4NGrgzZ4cj6JNCEwlOGJkbKeMgsjlykvwKUbRr7gw==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { - "@vitest/utils": "4.0.10", + "@vitest/utils": "4.1.11", "fflate": "^0.8.2", - "flatted": "^3.3.3", + "flatted": "^3.4.2", "pathe": "^2.0.3", "sirv": "^3.0.2", "tinyglobby": "^0.2.15", - "tinyrainbow": "^3.0.3" + "tinyrainbow": "^3.1.0" }, "funding": { "url": "https://opencollective.com/vitest" }, "peerDependencies": { - "vitest": "4.0.10" + "vitest": "4.1.11" } }, "node_modules/@vitest/utils": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.0.10.tgz", - "integrity": "sha512-kOuqWnEwZNtQxMKg3WmPK1vmhZu9WcoX69iwWjVz+jvKTsF1emzsv3eoPcDr6ykA3qP2bsCQE7CwqfNtAVzsmg==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.11.tgz", + "integrity": "sha512-zTCVGpyFsGWBhllOyKlTw/vnr6D9qxsfSDyfbyZmTyjHw5N/VuvzHpHoQjm2ZJzn4RJgx5w4r7V0er69CmLgPQ==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/pretty-format": "4.0.10", - "tinyrainbow": "^3.0.3" + "@vitest/pretty-format": "4.1.11", + "convert-source-map": "^2.0.0", + "tinyrainbow": "^3.1.0" }, "funding": { "url": "https://opencollective.com/vitest" @@ -2381,6 +2388,7 @@ "integrity": "sha512-NZyJarBfL7nWwIq+FDL6Zp/yHEhePMNnnJ0y3qfieCrmNvYct8uvtiV41UvlSe6apAfk0fY1FbWx+NwfmpvtTg==", "dev": true, "license": "MIT", + "peer": true, "bin": { "acorn": "bin/acorn" }, @@ -2573,13 +2581,13 @@ } }, "node_modules/axios": { - "version": "1.16.1", - "resolved": "https://registry.npmjs.org/axios/-/axios-1.16.1.tgz", - "integrity": "sha512-caYkukvroVPO8KrzuJEb50Hm07KwfBZPEC3VeFHTsqWHvKTsy54hjJz9BS/cdaypROE2rH6xvm9mHX4fgWkr3A==", + "version": "1.20.0", + "resolved": "https://registry.npmjs.org/axios/-/axios-1.20.0.tgz", + "integrity": "sha512-r8aOh8j9cGKpgQAqpzrUHnSIc6a59Y3Xf/cv8sy1DrHCkZHzQGEuoq1tARk6qSyDdtQGSDgpb9kFlruzPvrgwg==", "license": "MIT", "dependencies": { "follow-redirects": "^1.16.0", - "form-data": "^4.0.5", + "form-data": "^4.0.6", "https-proxy-agent": "^5.0.1", "proxy-from-env": "^2.1.0" } @@ -2617,13 +2625,16 @@ "license": "MIT" }, "node_modules/baseline-browser-mapping": { - "version": "2.9.4", - "resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.9.4.tgz", - "integrity": "sha512-ZCQ9GEWl73BVm8bu5Fts8nt7MHdbt5vY9bP6WGnUh+r3l8M7CgfyTlwsgCbMC66BNxPr6Xoce3j66Ms5YUQTNA==", + "version": "2.11.21", + "resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.11.21.tgz", + "integrity": "sha512-uh8vpY/1/YyFkunIDFH/12p7/7VdPKA1hejMVEbdkEaWnUz0Hesvx5EbiU6XxjyHZIOju+ZMbQJkRh+es3/spQ==", "dev": true, "license": "Apache-2.0", "bin": { - "baseline-browser-mapping": "dist/cli.js" + "baseline-browser-mapping": "dist/cli.cjs" + }, + "engines": { + "node": ">=6.0.0" } }, "node_modules/bidi-js": { @@ -2666,9 +2677,9 @@ "license": "ISC" }, "node_modules/brace-expansion": { - "version": "2.1.0", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.1.0.tgz", - "integrity": "sha512-TN1kCZAgdgweJhWWpgKYrQaMNHcDULHkWwQIspdtjV4Y5aurRdZpjAqn6yX3FPqTA9ngHCc4hJxMAMgGfve85w==", + "version": "2.1.4", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.1.4.tgz", + "integrity": "sha512-hGfVzPxthbf3+2yjg/RBs60cB0FhqBS/zvdV/4wn4/BmN0bNMMHPc4V/BbFieqf1TKAGGAHnY4eSjajCl0f2Xg==", "dev": true, "license": "MIT", "dependencies": { @@ -2708,6 +2719,7 @@ } ], "license": "MIT", + "peer": true, "dependencies": { "baseline-browser-mapping": "^2.8.2", "caniuse-lite": "^1.0.30001741", @@ -2756,9 +2768,9 @@ } }, "node_modules/caniuse-lite": { - "version": "1.0.30001741", - "resolved": "https://registry.npmjs.org/caniuse-lite/-/caniuse-lite-1.0.30001741.tgz", - "integrity": "sha512-QGUGitqsc8ARjLdgAfxETDhRbJ0REsP6O3I96TAth/mVjh2cYzN2u+3AzPP3aVSm2FehEItaJw1xd+IGBXWeSw==", + "version": "1.0.30001810", + "resolved": "https://registry.npmjs.org/caniuse-lite/-/caniuse-lite-1.0.30001810.tgz", + "integrity": "sha512-TITQPUkaz+aVk5GL6NhOdwk1aEaNTSDPsGFWrTuhKGtjTF70jL/Oht2W4c6rXUe5fu7Ie19VIahAXHIIiWWNeg==", "dev": true, "funding": [ { @@ -2777,9 +2789,9 @@ "license": "CC-BY-4.0" }, "node_modules/chai": { - "version": "6.2.1", - "resolved": "https://registry.npmjs.org/chai/-/chai-6.2.1.tgz", - "integrity": "sha512-p4Z49OGG5W/WBCPSS/dH3jQ73kD6tiMmUM+bckNK6Jr5JHMG3k9bg/BvKR8lKmtVBKmOiuVaV2ws8s9oSbwysg==", + "version": "6.2.2", + "resolved": "https://registry.npmjs.org/chai/-/chai-6.2.2.tgz", + "integrity": "sha512-NUPRluOfOiTKBKvWPtSD4PhFvWCqOi0BGStNWs57X9js7XGTprSmFoz5F0tWhR4WPjNeR9jXqdC7/UpSJTnlRg==", "dev": true, "license": "MIT", "engines": { @@ -2824,6 +2836,7 @@ "resolved": "https://registry.npmjs.org/chart.js/-/chart.js-4.5.0.tgz", "integrity": "sha512-aYeC/jDgSEx8SHWZvANYMioYMZ2KX02W6f6uVfyteuCGcadDLcYVHdfdygsTQkQ4TKn5lghoojAsPj5pu0SnvQ==", "license": "MIT", + "peer": true, "dependencies": { "@kurkle/color": "^0.3.0" }, @@ -2949,6 +2962,13 @@ "dev": true, "license": "MIT" }, + "node_modules/convert-source-map": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/convert-source-map/-/convert-source-map-2.0.0.tgz", + "integrity": "sha512-Kvp459HrV2FEJ1CAsi1Ku+MY3kasH19TFykTz2xWmMeq6bk2NU3XXvfJ+Q61m0xktWwt+1HSYf3JZsTms3aRJg==", + "dev": true, + "license": "MIT" + }, "node_modules/copy-anything": { "version": "3.0.5", "resolved": "https://registry.npmjs.org/copy-anything/-/copy-anything-3.0.5.tgz", @@ -3046,6 +3066,7 @@ "resolved": "https://registry.npmjs.org/date-fns/-/date-fns-4.1.0.tgz", "integrity": "sha512-Ukq0owbQXxa/U3EGtsdVBkR1w7KOQ5gIBqdH2hkvknzZPYvBxb/aa6E8L7tmjFtkwZBu3UXBbjIgPo/Ez4xaNg==", "license": "MIT", + "peer": true, "funding": { "type": "github", "url": "https://github.com/sponsors/kossnocorp" @@ -3119,9 +3140,9 @@ "license": "MIT" }, "node_modules/dompurify": { - "version": "3.4.3", - "resolved": "https://registry.npmjs.org/dompurify/-/dompurify-3.4.3.tgz", - "integrity": "sha512-VVwJidIJcp1hpg2OMXML3ZVRPYSZiq4aX7qBh83BSIpOaRDqI+qxhXjjIWnpzkOXhmp0L81lnoME1mnCc9H48A==", + "version": "3.4.14", + "resolved": "https://registry.npmjs.org/dompurify/-/dompurify-3.4.14.tgz", + "integrity": "sha512-dVoH9z+MY+C9IilgGCk3YfFqjLi3fChm2OiKJMzh6axrJ5qwxqWaZamgmHrpv22CN/KdbZJuGEGgfQoL00LTdg==", "license": "(MPL-2.0 OR Apache-2.0)", "optionalDependencies": { "@types/trusted-types": "^2.0.7" @@ -3193,16 +3214,16 @@ } }, "node_modules/es-module-lexer": { - "version": "1.7.0", - "resolved": "https://registry.npmjs.org/es-module-lexer/-/es-module-lexer-1.7.0.tgz", - "integrity": "sha512-jEQoCwk8hyb2AZziIOLhDqpm5+2ww5uIE6lkO/6jcOCusfk6LhMHpXXfBLXTZ7Ydyt0j4VoUQv6uGNYbdW+kBA==", + "version": "2.3.2", + "resolved": "https://registry.npmjs.org/es-module-lexer/-/es-module-lexer-2.3.2.tgz", + "integrity": "sha512-poHGpORABojJJucnV9KbOavETW8lBVnphkW77ER5/BQ5Fz7oXSoCNek7IH3vR5nRjdsEz926ibFYX8KtLQmdyw==", "dev": true, "license": "MIT" }, "node_modules/es-object-atoms": { - "version": "1.1.1", - "resolved": "https://registry.npmjs.org/es-object-atoms/-/es-object-atoms-1.1.1.tgz", - "integrity": "sha512-FGgH2h8zKNim9ljj7dankFPcICIK9Cp5bm+c2gQSYePhpaG5+esrLODihIorn+Pe6FGJzWhXQotPv73jTaldXA==", + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/es-object-atoms/-/es-object-atoms-1.1.2.tgz", + "integrity": "sha512-HWcBoN6NileqtSydK2FqHbS/LoDd2pqrnQHLyJzBj4kOp/ky2MWMN694xOfkK8/SnUsW2DH7EfyVlydKCsm1Zw==", "license": "MIT", "dependencies": { "es-errors": "^1.3.0" @@ -3297,6 +3318,7 @@ "integrity": "sha512-BhHmn2yNOFA9H9JmmIVKJmd288g9hrVRDkdoIgRCRuSySRUHH7r/DI6aAXW9T1WwUuY3DFgrcaqB+deURBLR5g==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { "@eslint-community/eslint-utils": "^4.8.0", "@eslint-community/regexpp": "^4.12.1", @@ -3424,9 +3446,9 @@ } }, "node_modules/eslint/node_modules/brace-expansion": { - "version": "1.1.14", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.14.tgz", - "integrity": "sha512-MWPGfDxnyzKU7rNOW9SP/c50vi3xrmrua/+6hfPbCS2ABNWfx24vPidzvC7krjU/RTo235sV776ymlsMtGKj8g==", + "version": "1.1.18", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.18.tgz", + "integrity": "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw==", "dev": true, "license": "MIT", "dependencies": { @@ -3554,9 +3576,9 @@ } }, "node_modules/expect-type": { - "version": "1.2.2", - "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.2.2.tgz", - "integrity": "sha512-JhFGDVJ7tmDJItKhYgJCGLOWjuK9vPxiXoUFLwLDc99NlmklilbiQJwoctZtt13+xMw91MCk/REan6MWHqDjyA==", + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.4.0.tgz", + "integrity": "sha512-KfYbmpRm0VbLjEvVa9yGwCi9GI34xvi7A/HXYWQO65CSD2u3MczUJSuwXKFIxlGsgBQizV9q5J9NHj4VG0n+pA==", "dev": true, "license": "Apache-2.0", "engines": { @@ -3749,16 +3771,16 @@ } }, "node_modules/form-data": { - "version": "4.0.5", - "resolved": "https://registry.npmjs.org/form-data/-/form-data-4.0.5.tgz", - "integrity": "sha512-8RipRLol37bNs2bhoV67fiTEvdTrbMUYcFTiy3+wuuOnUog2QBHCZWXDRijWQfAkhBj2Uf5UnVaiWwA5vdd82w==", + "version": "4.0.6", + "resolved": "https://registry.npmjs.org/form-data/-/form-data-4.0.6.tgz", + "integrity": "sha512-vKatAh4SlVfgbv+YtmhiRjhEMJsYpsG1Y2rMQtR+SVSbytsSD1YGzDIcrAJmdFec88u/+VoGmxnl+80gL1tRCQ==", "license": "MIT", "dependencies": { "asynckit": "^0.4.0", "combined-stream": "^1.0.8", "es-set-tostringtag": "^2.1.0", - "hasown": "^2.0.2", - "mime-types": "^2.1.12" + "hasown": "^2.0.4", + "mime-types": "^2.1.35" }, "engines": { "node": ">= 6" @@ -3936,9 +3958,9 @@ } }, "node_modules/hasown": { - "version": "2.0.2", - "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.2.tgz", - "integrity": "sha512-0hJU9SCPvmMzIBdZFqNPXWa6dqh7WdH0cII9y+CyS8rG3nL48Bclra9HmKhVVUHyPWNH5Y7xDwAB7bfgSjkUMQ==", + "version": "2.0.4", + "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.4.tgz", + "integrity": "sha512-T2UbfbBEF32wiepXIsMlTW9+dDYC6wMh/t/vYA4tuOMKqWz/n3vr1NFSxQiyP+zk2mXsoMA/i/7qV6LKut1t1A==", "license": "MIT", "dependencies": { "function-bind": "^1.1.2" @@ -4178,10 +4200,20 @@ } }, "node_modules/js-yaml": { - "version": "4.1.1", - "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.1.1.tgz", - "integrity": "sha512-qQKT4zQxXl8lLwBtHMWwaTcGfFOZviOJet3Oy/xmGk2gZH677CJM9EvtfdSkgWcATZhj/55JZ0rmy3myCT5lsA==", + "version": "4.3.2", + "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.3.2.tgz", + "integrity": "sha512-SFNOvSJ+Dgf/9An904Yx+CgSlIPCkIpao4qo51lpee25TIRejdH3rhR4EZMGoNx3/TP3O+wzWuiTFl4sqbltzA==", "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/puzrin" + }, + { + "type": "github", + "url": "https://github.com/sponsors/nodeca" + } + ], "license": "MIT", "dependencies": { "argparse": "^2.0.1" @@ -4196,6 +4228,7 @@ "integrity": "sha512-454TI39PeRDW1LgpyLPyURtB4Zx1tklSr6+OFOipsxGUH1WMTvk6C65JQdrj455+DP2uJ1+veBEHTGFKWVLFoA==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { "@acemir/cssom": "^0.9.23", "@asamuzakjp/dom-selector": "^6.7.4", @@ -4503,9 +4536,9 @@ } }, "node_modules/nanoid": { - "version": "3.3.11", - "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.11.tgz", - "integrity": "sha512-N8SpfPUnUp1bK+PMYW8qSWdl9U+wwNWI4QKxOYDy9JAro3WMX7p2OeVRF9v+347pnakNevPmiHhNmZ2HbFA76w==", + "version": "3.3.18", + "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.18.tgz", + "integrity": "sha512-DTg4MJbGMWkfi6VZFdNt2/caMbQy4Ou+Op/hJQvGEWcnVfoA1QA+xzRKAzw9jD6+GVOOeYr/mIcuDSdug6F6+w==", "funding": [ { "type": "github", @@ -4587,6 +4620,20 @@ "node": ">= 6" } }, + "node_modules/obug": { + "version": "2.1.4", + "resolved": "https://registry.npmjs.org/obug/-/obug-2.1.4.tgz", + "integrity": "sha512-4a+OsYv9UktOJKE+l1A4OufDgdRF9PifWj+tJnHURo/P+WOxpG4GzUFL9qCalmWauao6ogiG+QvnCovwPoyAWA==", + "dev": true, + "funding": [ + "https://github.com/sponsors/sxzz", + "https://opencollective.com/debug" + ], + "license": "MIT", + "engines": { + "node": ">=12.20.0" + } + }, "node_modules/optionator": { "version": "0.9.4", "resolved": "https://registry.npmjs.org/optionator/-/optionator-0.9.4.tgz", @@ -4771,6 +4818,7 @@ "integrity": "sha512-QP88BAKvMam/3NxH6vj2o21R6MjxZUAd6nlwAS/pnGvN9IVLocLHxGYIzFhg6fUQ+5th6P4dv4eW9jX3DSIj7A==", "dev": true, "license": "MIT", + "peer": true, "engines": { "node": ">=12" }, @@ -4826,9 +4874,9 @@ } }, "node_modules/postcss": { - "version": "8.5.14", - "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.14.tgz", - "integrity": "sha512-SoSL4+OSEtR99LHFZQiJLkT59C5B1amGO1NzTwj7TT1qCUgUO6hxOvzkOYxD+vMrXBM3XJIKzokoERdqQq/Zmg==", + "version": "8.5.26", + "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.26.tgz", + "integrity": "sha512-u82N74LFzG8ca+dD8puPnplTXoGH4fTPpVGuIbt36G3qvNlkvfD0lEAZSxaly3KX8TS/L1A1gsCEmvKmBcVbkQ==", "funding": [ { "type": "opencollective", @@ -4844,8 +4892,9 @@ } ], "license": "MIT", + "peer": true, "dependencies": { - "nanoid": "^3.3.11", + "nanoid": "^3.3.17", "picocolors": "^1.1.1", "source-map-js": "^1.2.1" }, @@ -5141,9 +5190,9 @@ } }, "node_modules/radix-vue/node_modules/nanoid": { - "version": "5.1.5", - "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-5.1.5.tgz", - "integrity": "sha512-Ir/+ZpE9fDsNH0hQ3C68uyThDXzYcim2EqcZ8zn8Chtt1iylPT9xXJB0kPCnqzgcEGikO9RxSrh63MsmVCU7Fw==", + "version": "5.1.16", + "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-5.1.16.tgz", + "integrity": "sha512-kVrnsrJqMR8+oLJnGEmSWw9BivK5mt7H3FZatVRjrc5wGqFYuBxX1yG7+A7Gi5AefkX6t/oCkizcQgpu0cY1dQ==", "funding": [ { "type": "github", @@ -5438,9 +5487,9 @@ "license": "MIT" }, "node_modules/std-env": { - "version": "3.10.0", - "resolved": "https://registry.npmjs.org/std-env/-/std-env-3.10.0.tgz", - "integrity": "sha512-5GS12FdOZNliM5mAOxFRg7Ir0pWz8MdpYm6AY6VPkGpbA7ZzmbzNcBJQ0GPvvyWgcY7QAhCgf9Uy89I03faLkg==", + "version": "4.2.0", + "resolved": "https://registry.npmjs.org/std-env/-/std-env-4.2.0.tgz", + "integrity": "sha512-oCUKSupKTHX53EyjDtuZQ64pjLJ6yYCtpmEw0goYxtjG9KpbRe8KAsl2tBUGU9DyMcJ0RwJ8GqJAFzMXcXW1Rw==", "dev": true, "license": "MIT" }, @@ -5734,11 +5783,14 @@ "license": "MIT" }, "node_modules/tinyexec": { - "version": "0.3.2", - "resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-0.3.2.tgz", - "integrity": "sha512-KQQR9yN7R5+OSwaK0XQoj22pwHoTlgYqmUscPYoknOoWCWfj/5/ABTMRi69FrKU5ffPVh5QcFikpWJI/P1ocHA==", + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-1.3.0.tgz", + "integrity": "sha512-QKAl9m8gWWGHV8jZcPeym6j+XULi6tOf1mT83WYJ4Lk2ytW/uwAWkrP0uFsdoYMdueVJ0qs26wZ+23xeB4ibNQ==", "dev": true, - "license": "MIT" + "license": "MIT", + "engines": { + "node": ">=18" + } }, "node_modules/tinyglobby": { "version": "0.2.15", @@ -5758,9 +5810,9 @@ } }, "node_modules/tinyrainbow": { - "version": "3.0.3", - "resolved": "https://registry.npmjs.org/tinyrainbow/-/tinyrainbow-3.0.3.tgz", - "integrity": "sha512-PSkbLUoxOFRzJYjjxHJt9xro7D+iilgMX/C9lawzVuYiIdcihh9DXmVibBe8lmcFrRi/VzlPjBxbN7rH24q8/Q==", + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/tinyrainbow/-/tinyrainbow-3.1.1.tgz", + "integrity": "sha512-yau8yJdTt989Mm0Bd/236QnzEiPf2xLLTqUZRUJOo/3CB078LSwzei343DgtJVmfJKJE3TMINY1u42SQsP6mXw==", "dev": true, "license": "MIT", "engines": { @@ -5881,6 +5933,7 @@ "integrity": "sha512-p1diW6TqL9L07nNxvRMM7hMMw4c5XOo/1ibL4aAIGmSAt9slTE1Xgw5KWuof2uTOvCg9BY7ZRi+GaF+7sfgPeQ==", "devOptional": true, "license": "Apache-2.0", + "peer": true, "bin": { "tsc": "bin/tsc", "tsserver": "bin/tsserver" @@ -5969,13 +6022,14 @@ "license": "MIT" }, "node_modules/vite": { - "version": "7.3.3", - "resolved": "https://registry.npmjs.org/vite/-/vite-7.3.3.tgz", - "integrity": "sha512-/4XH147Ui7OGTjg3HbdWe5arnZQSbfuRzdr9Ec7TQi5I7R+ir0Rlc9GIvD4v0XZurELqA035KVXJXpR61xhiTA==", + "version": "7.3.6", + "resolved": "https://registry.npmjs.org/vite/-/vite-7.3.6.tgz", + "integrity": "sha512-4XP60spRGjSZFf1qYH+dJIkK2znL3zQfl9KkOV9MkkRR/3Dls0dxaBsQPTloEc5BLXWPL9vsOxopxyKoMmDueg==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { - "esbuild": "^0.27.0", + "esbuild": "^0.27.0 || ^0.28.0", "fdir": "^6.5.0", "picomatch": "^4.0.3", "postcss": "^8.5.6", @@ -6044,31 +6098,32 @@ } }, "node_modules/vitest": { - "version": "4.0.10", - "resolved": "https://registry.npmjs.org/vitest/-/vitest-4.0.10.tgz", - "integrity": "sha512-2Fqty3MM9CDwOVet/jaQalYlbcjATZwPYGcqpiYQqgQ/dLC7GuHdISKgTYIVF/kaishKxLzleKWWfbSDklyIKg==", + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.11.tgz", + "integrity": "sha512-fhACrNXUidIbGSBr5FlbuBkO7VWC1ZyLl0DO4CU2DrQoAPxX84Ysxs+HeGQpii5lZWV1Q4gBZTTu49mF+A6Edw==", "dev": true, "license": "MIT", + "peer": true, "dependencies": { - "@vitest/expect": "4.0.10", - "@vitest/mocker": "4.0.10", - "@vitest/pretty-format": "4.0.10", - "@vitest/runner": "4.0.10", - "@vitest/snapshot": "4.0.10", - "@vitest/spy": "4.0.10", - "@vitest/utils": "4.0.10", - "debug": "^4.4.3", - "es-module-lexer": "^1.7.0", - "expect-type": "^1.2.2", + "@vitest/expect": "4.1.11", + "@vitest/mocker": "4.1.11", + "@vitest/pretty-format": "4.1.11", + "@vitest/runner": "4.1.11", + "@vitest/snapshot": "4.1.11", + "@vitest/spy": "4.1.11", + "@vitest/utils": "4.1.11", + "es-module-lexer": "^2.0.0", + "expect-type": "^1.3.0", "magic-string": "^0.30.21", + "obug": "^2.1.1", "pathe": "^2.0.3", "picomatch": "^4.0.3", - "std-env": "^3.10.0", + "std-env": "^4.0.0-rc.1", "tinybench": "^2.9.0", - "tinyexec": "^0.3.2", + "tinyexec": "^1.0.2", "tinyglobby": "^0.2.15", - "tinyrainbow": "^3.0.3", - "vite": "^6.0.0 || ^7.0.0", + "tinyrainbow": "^3.1.0", + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0", "why-is-node-running": "^2.3.0" }, "bin": { @@ -6082,20 +6137,23 @@ }, "peerDependencies": { "@edge-runtime/vm": "*", - "@types/debug": "^4.1.12", + "@opentelemetry/api": "^1.9.0", "@types/node": "^20.0.0 || ^22.0.0 || >=24.0.0", - "@vitest/browser-playwright": "4.0.10", - "@vitest/browser-preview": "4.0.10", - "@vitest/browser-webdriverio": "4.0.10", - "@vitest/ui": "4.0.10", + "@vitest/browser-playwright": "4.1.11", + "@vitest/browser-preview": "4.1.11", + "@vitest/browser-webdriverio": "4.1.11", + "@vitest/coverage-istanbul": "4.1.11", + "@vitest/coverage-v8": "4.1.11", + "@vitest/ui": "4.1.11", "happy-dom": "*", - "jsdom": "*" + "jsdom": "*", + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0" }, "peerDependenciesMeta": { "@edge-runtime/vm": { "optional": true }, - "@types/debug": { + "@opentelemetry/api": { "optional": true }, "@types/node": { @@ -6110,6 +6168,12 @@ "@vitest/browser-webdriverio": { "optional": true }, + "@vitest/coverage-istanbul": { + "optional": true + }, + "@vitest/coverage-v8": { + "optional": true + }, "@vitest/ui": { "optional": true }, @@ -6118,6 +6182,9 @@ }, "jsdom": { "optional": true + }, + "vite": { + "optional": false } } }, @@ -6133,6 +6200,7 @@ "resolved": "https://registry.npmjs.org/vue/-/vue-3.5.21.tgz", "integrity": "sha512-xxf9rum9KtOdwdRkiApWL+9hZEMWE90FHh8yS1+KJAiWYh+iGWV1FquPjoO9VUHQ+VIhsCXNNyZ5Sf4++RVZBA==", "license": "MIT", + "peer": true, "dependencies": { "@vue/compiler-dom": "3.5.21", "@vue/compiler-sfc": "3.5.21", @@ -6165,7 +6233,6 @@ "integrity": "sha512-CydUvFOQKD928UzZhTp4pr2vWz1L+H99t7Pkln2QSPdvmURT0MoC4wUccfCnuEaihNsu9aYYyk+bep8rlfkUXw==", "dev": true, "license": "MIT", - "peer": true, "dependencies": { "debug": "^4.4.0", "eslint-scope": "^8.2.0", @@ -6190,7 +6257,6 @@ "integrity": "sha512-Uhdk5sfqcee/9H/rCOJikYz67o0a2Tw2hGRPOG2Y1R2dg7brRe1uG0yaNQDHu+TO/uQPF/5eCapvYSmHUjt7JQ==", "dev": true, "license": "Apache-2.0", - "peer": true, "engines": { "node": "^18.18.0 || ^20.9.0 || >=21.1.0" }, @@ -6438,9 +6504,9 @@ } }, "node_modules/ws": { - "version": "8.18.3", - "resolved": "https://registry.npmjs.org/ws/-/ws-8.18.3.tgz", - "integrity": "sha512-PEIGCY5tSlUt50cqyMXfCzX+oOPqN0vuGqWzbcJ2xvnkzkq46oOpz7dQaTDBdfICb4N14+GARUDw2XV2N4tvzg==", + "version": "8.21.3", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.3.tgz", + "integrity": "sha512-201TZ/kPWxoPr/OKWjquZR1SWKXcvxdH+e1xrx89b3YbmzLMFCLfnaG1HFIgWzJOEWZ7MvpK++odZufgYR50Rw==", "dev": true, "license": "MIT", "engines": { diff --git a/frontend/package.json b/frontend/package.json index e62412123..d7ce7c246 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -4,8 +4,12 @@ "version": "0.0.0", "type": "module", "scripts": { + "sync:vscodex": "node scripts/sync-vscodex.mjs", + "predev": "npm run sync:vscodex", "dev": "vite --host", + "prebuild": "npm run sync:vscodex", "build": "vite build", + "prebuild:with-typecheck": "npm run sync:vscodex", "build:with-typecheck": "vue-tsc -b && vite build", "preview": "vite preview", "test": "node --experimental-require-module --disable-warning=ExperimentalWarning ./node_modules/vitest/vitest.mjs", @@ -21,13 +25,13 @@ "@types/marked": "^5.0.2", "@types/three": "^0.180.0", "@vueuse/core": "^13.9.0", - "axios": "^1.12.1", + "axios": "^1.20.0", "chart.js": "^4.5.0", "chartjs-adapter-date-fns": "^3.0.0", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "date-fns": "^4.1.0", - "dompurify": "^3.3.0", + "dompurify": "^3.4.14", "highlight.js": "^11.11.1", "lucide-vue-next": "^0.544.0", "marked": "^16.0.0", @@ -47,20 +51,20 @@ "@typescript-eslint/eslint-plugin": "^8.47.0", "@typescript-eslint/parser": "^8.47.0", "@vitejs/plugin-vue": "^6.0.1", - "@vitest/ui": "^4.0.10", + "@vitest/ui": "^4.1.11", "@vue/tsconfig": "^0.7.0", "autoprefixer": "^10.4.21", "baseline-browser-mapping": "^2.9.4", "eslint": "^9.39.1", "eslint-plugin-vue": "^10.5.1", "jsdom": "^27.2.0", - "postcss": "^8.5.6", + "postcss": "^8.5.26", "tailwindcss": "^3.4.17", "tailwindcss-animate": "^1.0.7", "typescript": "~5.8.3", "typescript-eslint": "^8.49.0", - "vite": "^7.1.2", - "vitest": "^4.0.10", + "vite": "^7.3.6", + "vitest": "^4.1.11", "vue-tsc": "^3.0.5" } } diff --git a/frontend/scripts/sync-vscodex.mjs b/frontend/scripts/sync-vscodex.mjs new file mode 100644 index 000000000..154c334fc --- /dev/null +++ b/frontend/scripts/sync-vscodex.mjs @@ -0,0 +1,75 @@ +import { spawnSync } from 'node:child_process' +import { cpSync, existsSync, mkdirSync, readFileSync, rmSync } from 'node:fs' +import { dirname, resolve } from 'node:path' +import { fileURLToPath } from 'node:url' + +const frontendRoot = resolve(dirname(fileURLToPath(import.meta.url)), '..') +const moduleRoot = resolve(frontendRoot, '..', 'aether-vscodex') +const webRoot = resolve(moduleRoot, 'web') +const webPackagePath = resolve(webRoot, 'package.json') +const webLockPath = resolve(webRoot, 'package-lock.json') +const webPackage = JSON.parse(readFileSync(webPackagePath, 'utf8')) +const webLock = JSON.parse(readFileSync(webLockPath, 'utf8')) +const npmCommand = process.platform === 'win32' ? 'npm.cmd' : 'npm' +const requiredPackages = Object.keys({ + ...webPackage.dependencies, + ...webPackage.devDependencies +}) +const requiredCommands = ['vite', 'vue-tsc'] +const commandSuffix = process.platform === 'win32' ? '.cmd' : '' + +function dependencyMatchesLock(dependency) { + const installedPackagePath = resolve( + webRoot, + 'node_modules', + ...dependency.split('/'), + 'package.json' + ) + if (!existsSync(installedPackagePath)) { + return false + } + + const lockedVersion = webLock.packages?.[`node_modules/${dependency}`]?.version + const installedVersion = JSON.parse(readFileSync(installedPackagePath, 'utf8')).version + return typeof lockedVersion === 'string' && installedVersion === lockedVersion +} + +const hasWebDependencies = + requiredPackages.every(dependencyMatchesLock) && + requiredCommands.every((command) => + existsSync(resolve(webRoot, 'node_modules', '.bin', `${command}${commandSuffix}`)) + ) + +function runNpm(args, action) { + const result = spawnSync(npmCommand, ['--prefix', webRoot, ...args], { + stdio: 'inherit' + }) + + if (result.error) { + throw new Error(`Failed to ${action}: ${result.error.message}`) + } + if (result.status !== 0) { + throw new Error(`${action} failed with exit status ${result.status ?? 'unknown'}`) + } +} + +if (!hasWebDependencies) { + console.log('=> aether-vscodex Web dependencies are missing; installing from package-lock.json...') + runNpm( + ['ci', '--include=dev', '--no-audit', '--no-fund'], + 'install aether-vscodex Web dependencies' + ) +} + +runNpm(['run', 'build'], 'build aether-vscodex Web') + +const vueBuild = resolve(webRoot, 'dist') +const destination = resolve(frontendRoot, 'public', 'aether-vscodex') + +if (!existsSync(resolve(vueBuild, 'index.html'))) { + throw new Error(`aether-vscodex Vue build was not found at ${vueBuild}`) +} + +rmSync(destination, { recursive: true, force: true }) +mkdirSync(destination, { recursive: true }) +cpSync(vueBuild, destination, { recursive: true }) diff --git a/frontend/src/App.vue b/frontend/src/App.vue index 3874681e1..7e6e5afb9 100644 --- a/frontend/src/App.vue +++ b/frontend/src/App.vue @@ -9,20 +9,18 @@ import { onMounted, onErrorCaptured, onUnmounted } from 'vue' import { useAuthStore } from '@/stores/auth' import ToastContainer from '@/components/ToastContainer.vue' import ConfirmContainer from '@/components/ConfirmContainer.vue' -import apiClient, { AUTH_STATE_CHANGE_EVENT } from '@/api/client' +import { + AUTH_SESSION_SIGNAL_KEY, + AUTH_STATE_CHANGE_EVENT, + parseAuthSessionSignal, + type AuthStateChangeDetail, +} from '@/api/client' import { NETWORK_CONFIG, AUTH_CONFIG } from '@/config/constants' import router from '@/router' -import { hasAuthIdentityChanged } from '@/utils/authToken' import { log } from '@/utils/logger' const authStore = useAuthStore() -// 立即检查token,如果存在就设置到store中 -const storedToken = apiClient.getToken() -if (storedToken) { - authStore.token = storedToken -} - // 全局错误处理器 - 只处理特定错误,避免完全吞掉所有错误 onErrorCaptured((error: Error) => { log.error('Error captured in component', error) @@ -88,32 +86,16 @@ if (typeof window !== 'undefined') { }) } -async function syncExternalAuthState(nextToken: string | null): Promise { - const previousToken = authStore.token - const previousUser = authStore.user - ? { - id: authStore.user.id, - role: authStore.user.role, - } - : null - - authStore.syncToken() - - if (!nextToken) { - if (previousToken || previousUser) { +async function syncExternalAuthState(authenticated: boolean): Promise { + if (!authenticated) { + if (authStore.token || authStore.user) { authStore.applyExternalLogout() await router.replace('/') } return } - const identityChanged = hasAuthIdentityChanged(previousToken, nextToken, previousUser) - if (!identityChanged && previousUser) { - return - } - - const user = await authStore.fetchCurrentUser() - if (!user) { + if (!await authStore.applyExternalLogin()) { return } @@ -123,16 +105,27 @@ async function syncExternalAuthState(nextToken: string | null): Promise { } function handleAuthStorageChange(event: StorageEvent): void { - if (event.key !== 'access_token') { + if (event.key !== AUTH_SESSION_SIGNAL_KEY) { return } - syncExternalAuthState(event.newValue).catch((err) => log.error('syncExternalAuthState failed', err)) + const signal = parseAuthSessionSignal(event.newValue) + if (!signal) { + return + } + syncExternalAuthState(signal.authenticated).catch((err) => log.error('syncExternalAuthState failed', err)) } function handleLocalAuthStateChange(event: Event): void { - const authEvent = event as CustomEvent<{ token: string | null }> - syncExternalAuthState(authEvent.detail?.token ?? apiClient.getToken()).catch((err) => log.error('syncExternalAuthState failed', err)) + const authEvent = event as CustomEvent + if (!authEvent.detail) { + return + } + if (!authEvent.detail.authenticated) { + syncExternalAuthState(false).catch((err) => log.error('syncExternalAuthState failed', err)) + } else { + authStore.syncToken() + } } onMounted(async () => { diff --git a/frontend/src/api/__tests__/auth-turnstile.spec.ts b/frontend/src/api/__tests__/auth-turnstile.spec.ts index 1f78005dd..212625fa9 100644 --- a/frontend/src/api/__tests__/auth-turnstile.spec.ts +++ b/frontend/src/api/__tests__/auth-turnstile.spec.ts @@ -46,6 +46,21 @@ describe('authApi turnstile payloads', () => { }) }) + it('binds verification and status requests to the verification session token', async () => { + await authApi.verifyEmail('alice@example.com', '123456', 'verification-session-token') + await authApi.getVerificationStatus('alice@example.com', 'verification-session-token') + + expect(postMock).toHaveBeenNthCalledWith(1, '/api/auth/verify-email', { + email: 'alice@example.com', + code: '123456', + verification_token: 'verification-session-token', + }) + expect(postMock).toHaveBeenNthCalledWith(2, '/api/auth/verification-status', { + email: 'alice@example.com', + verification_token: 'verification-session-token', + }) + }) + it('refreshes auth token without a request body', async () => { postMock.mockResolvedValue({ data: { access_token: 'new-access-token' } }) @@ -54,4 +69,12 @@ describe('authApi turnstile payloads', () => { expect(postMock).toHaveBeenCalledWith('/api/auth/refresh') expect(setTokenMock).toHaveBeenCalledWith('new-access-token') }) + + it('publishes login session availability without changing the token payload', async () => { + postMock.mockResolvedValue({ data: { access_token: 'login-access-token' } }) + + await authApi.login({ email: 'alice@example.com', password: 'secret123' }) + + expect(setTokenMock).toHaveBeenCalledWith('login-access-token', true) + }) }) diff --git a/frontend/src/api/__tests__/client.spec.ts b/frontend/src/api/__tests__/client.spec.ts index da4595f70..9542e359f 100644 --- a/frontend/src/api/__tests__/client.spec.ts +++ b/frontend/src/api/__tests__/client.spec.ts @@ -1,7 +1,11 @@ import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import type { AxiosAdapter, AxiosInstance, InternalAxiosRequestConfig } from 'axios' -import apiClient, { AUTH_STATE_CHANGE_EVENT } from '@/api/client' +import apiClient, { + AUTH_SESSION_SIGNAL_KEY, + AUTH_STATE_CHANGE_EVENT, + parseAuthSessionSignal, +} from '@/api/client' import { cache, cachedRequest } from '@/utils/cache' type TestableApiClient = typeof apiClient & { @@ -24,17 +28,82 @@ describe('apiClient auth state change event', () => { window.addEventListener(AUTH_STATE_CHANGE_EVENT, handler as EventListener) apiClient.setToken('access-token') + localStorage.setItem('access_token', 'legacy-local-token') + sessionStorage.setItem('access_token', 'legacy-session-token') apiClient.clearAuth() expect(localStorage.getItem('access_token')).toBeNull() - expect(handler).toHaveBeenCalledTimes(1) + expect(sessionStorage.getItem('access_token')).toBeNull() + expect(handler).toHaveBeenCalledTimes(2) - const event = handler.mock.calls[0][0] as CustomEvent<{ token: string | null }> - expect(event.detail).toEqual({ token: null }) + const event = handler.mock.calls[1][0] as CustomEvent<{ authenticated: boolean }> + expect(event.detail).toEqual({ authenticated: false }) window.removeEventListener(AUTH_STATE_CHANGE_EVENT, handler as EventListener) }) + it('keeps access tokens in memory and only stores token-free session metadata', () => { + apiClient.setToken('sensitive-access-token', true) + + expect(apiClient.getToken()).toBe('sensitive-access-token') + expect(localStorage.getItem('access_token')).toBeNull() + expect(sessionStorage.getItem('access_token')).toBeNull() + + const rawSignal = localStorage.getItem(AUTH_SESSION_SIGNAL_KEY) + expect(rawSignal).not.toContain('sensitive-access-token') + expect(parseAuthSessionSignal(rawSignal)).toMatchObject({ authenticated: true }) + }) + + it('restores a session through the refresh cookie and stores the result in memory only', async () => { + const rawClient = apiClient as TestableApiClient + const previousAdapter = rawClient.client.defaults.adapter + + rawClient.client.defaults.adapter = (async (config: InternalAxiosRequestConfig) => ({ + data: { access_token: 'restored-access-token' }, + status: 200, + statusText: 'OK', + headers: {}, + config, + })) as AxiosAdapter + + try { + await expect(apiClient.restoreSession()).resolves.toBe('restored-access-token') + expect(apiClient.getToken()).toBe('restored-access-token') + expect(localStorage.getItem('access_token')).toBeNull() + expect(sessionStorage.getItem('access_token')).toBeNull() + } finally { + rawClient.client.defaults.adapter = previousAdapter + } + }) + + it('does not resurrect a session when logout wins an in-flight restore', async () => { + const rawClient = apiClient as TestableApiClient + const previousAdapter = rawClient.client.defaults.adapter + let resolveRefresh!: (response: Awaited>) => void + + rawClient.client.defaults.adapter = (() => new Promise((resolve) => { + resolveRefresh = resolve + })) as AxiosAdapter + + try { + const restore = apiClient.restoreSession() + await vi.waitFor(() => expect(resolveRefresh).toBeTypeOf('function')) + apiClient.clearAuth() + resolveRefresh({ + data: { access_token: 'stale-access-token' }, + status: 200, + statusText: 'OK', + headers: {}, + config: {} as InternalAxiosRequestConfig, + }) + + await expect(restore).rejects.toThrow('Auth state changed') + expect(apiClient.getToken()).toBeNull() + } finally { + rawClient.client.defaults.adapter = previousAdapter + } + }) + it('clears cached API data whenever the authentication identity changes', () => { apiClient.setToken('first-token') cache.set('dashboard', { owner: 'first-user' }, 30_000) @@ -99,4 +168,32 @@ describe('apiClient auth state change event', () => { rawClient.client.defaults.adapter = previousAdapter } }) + + it('authenticates protected gateway operational requests', async () => { + const rawClient = apiClient as TestableApiClient + const previousAdapter = rawClient.client.defaults.adapter + const requests: InternalAxiosRequestConfig[] = [] + + rawClient.client.defaults.adapter = (async (config: InternalAxiosRequestConfig) => { + requests.push(config) + return { + data: '', + status: 200, + statusText: 'OK', + headers: {}, + config, + } + }) as AxiosAdapter + + try { + apiClient.setToken('operational-access-token') + await apiClient.get('/_gateway/metrics') + + expect(requests).toHaveLength(1) + expect(requests[0].headers.Authorization).toBe('Bearer operational-access-token') + expect(requests[0].headers['X-Client-Device-Id']).toBeTruthy() + } finally { + rawClient.client.defaults.adapter = previousAdapter + } + }) }) diff --git a/frontend/src/api/__tests__/vscodex.spec.ts b/frontend/src/api/__tests__/vscodex.spec.ts new file mode 100644 index 000000000..883ad0a66 --- /dev/null +++ b/frontend/src/api/__tests__/vscodex.spec.ts @@ -0,0 +1,116 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { deleteMock, getMock, postMock } = vi.hoisted(() => ({ + deleteMock: vi.fn(), + getMock: vi.fn(), + postMock: vi.fn(), +})) + +vi.mock('@/api/client', () => ({ + default: { + get: getMock, + post: postMock, + delete: deleteMock, + }, +})) + +import { vscodexApi } from '@/api/vscodex' + +describe('vscodexApi', () => { + beforeEach(() => { + getMock.mockReset() + postMock.mockReset() + deleteMock.mockReset() + }) + + it('lists and normalizes the current user devices', async () => { + getMock.mockResolvedValue({ + data: { + devices: [ + { + device_id: 'device-online', + display_name: 'Studio Mac', + connected: true, + last_seen_at: '2026-08-31T10:00:00Z', + }, + { + id: 'device-unknown', + name: 'Laptop', + status: 'unexpected-status', + }, + { name: 'missing id' }, + ], + }, + }) + + await expect(vscodexApi.listDevices()).resolves.toEqual([ + { + id: 'device-online', + name: 'Studio Mac', + status: 'online', + last_seen_at: '2026-08-31T10:00:00Z', + created_at: null, + }, + { + id: 'device-unknown', + name: 'Laptop', + status: 'unknown', + last_seen_at: null, + created_at: null, + }, + ]) + expect(getMock).toHaveBeenCalledWith('/api/users/me/vscodex/devices') + }) + + it('creates a pairing using an explicit empty body and normalizes its code', async () => { + postMock.mockResolvedValue({ + data: { + pairing_code: 'PAIR-1234', + expires_in_seconds: 300, + }, + }) + + await expect(vscodexApi.createPairing()).resolves.toEqual({ + code: 'PAIR-1234', + expires_at: null, + expires_in_seconds: 300, + }) + expect(postMock).toHaveBeenCalledWith('/api/users/me/vscodex/pairings', {}) + }) + + it('requests a scoped WebSocket ticket for the selected device', async () => { + postMock.mockResolvedValue({ + data: { + ticket: 'single-use-ticket', + wsUrl: 'wss://aether.example/api/vscodex/ws', + }, + }) + + await expect(vscodexApi.createWsTicket('device-online')).resolves.toEqual({ + ticket: 'single-use-ticket', + ws_url: 'wss://aether.example/api/vscodex/ws', + expires_at: null, + }) + expect(postMock).toHaveBeenCalledWith('/api/users/me/vscodex/ws-tickets', { + device_id: 'device-online', + }) + }) + + it('revokes the selected device using an encoded path segment', async () => { + deleteMock.mockResolvedValue({ status: 204 }) + + await expect(vscodexApi.deleteDevice('device/one')).resolves.toBeUndefined() + expect(deleteMock).toHaveBeenCalledWith('/api/users/me/vscodex/devices/device%2Fone') + }) + + it('rejects incomplete pairing and ticket responses', async () => { + postMock + .mockResolvedValueOnce({ data: {} }) + .mockResolvedValueOnce({ data: { ticket: 'missing-url' } }) + + await expect(vscodexApi.createPairing()).rejects.toThrow('Pairing response did not include a code') + await expect(vscodexApi.createWsTicket('device-online')).rejects.toThrow( + 'WebSocket ticket response was incomplete', + ) + }) +}) diff --git a/frontend/src/api/admin-payments.ts b/frontend/src/api/admin-payments.ts index 19bc69353..73e9a49ec 100644 --- a/frontend/src/api/admin-payments.ts +++ b/frontend/src/api/admin-payments.ts @@ -11,8 +11,23 @@ export interface PaymentCallbackRecord { payload_hash: string | null signature_valid: boolean status: string - payload: Record | null + payload: null + has_payload: boolean + payload_summary: { + kind: 'null' | 'boolean' | 'number' | 'string' | 'array' | 'object' + serialized_bytes: number + objects: number + arrays: number + strings: number + numbers: number + booleans: number + nulls: number + object_fields: number + array_items: number + max_depth: number + } | null error_message: string | null + has_error_message: boolean created_at: string processed_at: string | null } diff --git a/frontend/src/api/admin.ts b/frontend/src/api/admin.ts index 51176ae09..2d3f32b5e 100644 --- a/frontend/src/api/admin.ts +++ b/frontend/src/api/admin.ts @@ -145,7 +145,7 @@ export interface UserExport { email: string email_verified?: boolean username: string - password_hash: string + password_hash?: string | null role: string allowed_providers?: string[] | null allowed_providers_mode?: 'inherit' | 'unrestricted' | 'specific' | 'deny_all' @@ -169,9 +169,11 @@ export interface UserExport { export interface UserApiKeyExport { api_key_id?: string + // Legacy 1.3-1.5 import-only credential fields. Version 1.6 exports omit them. key?: string | null - key_hash: string + key_hash?: string | null key_encrypted?: string | null + credential_state?: 'not_exported' name?: string | null is_standalone: boolean allowed_providers?: string[] | null diff --git a/frontend/src/api/async-tasks.ts b/frontend/src/api/async-tasks.ts index 5ce541bfe..2017cdb93 100644 --- a/frontend/src/api/async-tasks.ts +++ b/frontend/src/api/async-tasks.ts @@ -204,6 +204,14 @@ export const asyncTasksApi = { return response.data }, + async getVideoBlob(taskId: string): Promise { + const response = await apiClient.get( + `/api/admin/video-tasks/${encodeURIComponent(taskId)}/video`, + { responseType: 'blob', timeout: 0 }, + ) + return response.data + }, + async trigger(taskKey: string, payload: Record = {}): Promise<{ run_id: string; status: string }> { const response = await apiClient.post(`/api/admin/tasks/${taskKey}/trigger`, payload) return response.data diff --git a/frontend/src/api/auth.ts b/frontend/src/api/auth.ts index b12e3e771..30c35610d 100644 --- a/frontend/src/api/auth.ts +++ b/frontend/src/api/auth.ts @@ -40,11 +40,13 @@ export interface SendVerificationCodeResponse { message: string success: boolean expire_minutes?: number + verification_token: string } export interface VerifyEmailRequest { email: string code: string + verification_token: string } export interface VerifyEmailResponse { @@ -54,6 +56,7 @@ export interface VerifyEmailResponse { export interface VerificationStatusRequest { email: string + verification_token: string } export interface VerificationStatusResponse { @@ -72,6 +75,7 @@ export interface RegisterRequest { invite_code?: string privacy_policy_accepted?: boolean privacy_policy_version?: string + email_verification_token?: string } export interface RegisterResponse { @@ -141,7 +145,7 @@ export interface User { export const authApi = { async login(credentials: LoginRequest): Promise { const response = await apiClient.post('/api/auth/login', credentials) - apiClient.setToken(response.data.access_token) + apiClient.setToken(response.data.access_token, true) return response.data }, @@ -184,10 +188,14 @@ export const authApi = { return response.data }, - async verifyEmail(email: string, code: string): Promise { + async verifyEmail( + email: string, + code: string, + verificationToken: string + ): Promise { const response = await apiClient.post( '/api/auth/verify-email', - { email, code } + { email, code, verification_token: verificationToken } ) return response.data }, @@ -204,10 +212,13 @@ export const authApi = { return response.data }, - async getVerificationStatus(email: string): Promise { + async getVerificationStatus( + email: string, + verificationToken: string + ): Promise { const response = await apiClient.post( '/api/auth/verification-status', - { email } + { email, verification_token: verificationToken } ) return response.data }, diff --git a/frontend/src/api/billing.ts b/frontend/src/api/billing.ts index c276c9c28..688427025 100644 --- a/frontend/src/api/billing.ts +++ b/frontend/src/api/billing.ts @@ -55,12 +55,14 @@ export interface GatewayTestResponse { export interface WalletCreditEntitlement { type: 'wallet_credit' + replacement_group?: string amount_usd: number balance_bucket?: WalletCreditBucket } export interface DailyQuotaEntitlement { type: 'daily_quota' + replacement_group?: string daily_quota_usd: number reset_timezone?: string carry_over?: boolean @@ -69,13 +71,41 @@ export interface DailyQuotaEntitlement { export interface MembershipGroupEntitlement { type: 'membership_group' + replacement_group?: string grant_user_groups: string[] } +export type UsagePolicyMetric = 'request_count' | 'concurrency' | 'actual_cost_usd' +export type UsagePolicyEnforcement = 'hard_cap' + +export type UsagePolicyWindow = + | { kind: 'rolling'; seconds: number } + | { kind: 'calendar_day'; timezone?: string } + | { kind: 'calendar_week'; timezone?: string; week_start?: number } + | { kind: 'calendar_month'; timezone?: string } + | { kind: 'subscription_period' } + | { kind: 'concurrent' } + +export interface UsagePolicyRule { + metric: UsagePolicyMetric + window: UsagePolicyWindow + limit: number + enforcement?: UsagePolicyEnforcement +} + +export interface UsagePolicyEntitlement { + type: 'usage_policy' + policy_id?: string + name?: string + replacement_group?: string + rules: UsagePolicyRule[] +} + export type BillingEntitlement = | WalletCreditEntitlement | DailyQuotaEntitlement | MembershipGroupEntitlement + | UsagePolicyEntitlement export interface BillingPlan { id: string diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 3bebe767e..64a1f92da 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -10,6 +10,16 @@ import { cache } from '@/utils/cache' // 在开发环境下使用代理,生产环境使用环境变量 const API_BASE_URL = import.meta.env.VITE_API_URL || '' export const AUTH_STATE_CHANGE_EVENT = 'aether-auth-state-change' +export const AUTH_SESSION_SIGNAL_KEY = 'aether_auth_session_signal' + +export type AuthStateChangeDetail = { + authenticated: boolean +} + +export type AuthSessionSignal = AuthStateChangeDetail & { + eventId: string + emittedAt: number +} type MockRuntime = typeof import('@/mocks') @@ -48,6 +58,14 @@ function isAuthRequest(url?: string): boolean { return url?.includes('/auth/login') || url?.includes('/auth/refresh') || url?.includes('/auth/logout') || false } +function isProtectedOperationalEndpoint(url?: string): boolean { + if (!url) return false + const path = url.split('?', 1)[0] + return path === '/_gateway/metrics' || + path.startsWith('/_gateway/audit/') || + path.startsWith('/_gateway/async-tasks/') +} + /** * 判断 403 错误是否表示用户账号级别的问题(需要清除认证并跳转) */ @@ -100,17 +118,11 @@ function createDemoAdapter(defaultAdapter: AxiosAdapter) { class ApiClient { private client: AxiosInstance private token: string | null = null + private authStateVersion = 0 private isRefreshing = false private refreshPromise: Promise | null = null private readonly refreshCoordinator = new CrossTabRefreshCoordinator() - private readonly onStorageSync = (event: StorageEvent): void => { - if (event.key !== 'access_token') { - return - } - this.syncTokenState(event.newValue) - } - constructor() { this.client = axios.create({ baseURL: API_BASE_URL, @@ -126,7 +138,7 @@ class ApiClient { this.client.defaults.adapter = createDemoAdapter(defaultAdapter) this.setupInterceptors() - this.setupCrossTabAuthSync() + this.purgeLegacyStoredTokens() } /** @@ -136,12 +148,15 @@ class ApiClient { // 请求拦截器 - 仅处理认证 this.client.interceptors.request.use( (config) => { - if (config.url?.includes('/api/')) { + const carriesSessionCredentials = config.url?.includes('/api/') || + isProtectedOperationalEndpoint(config.url) + + if (carriesSessionCredentials) { config.headers['X-Client-Device-Id'] = getClientDeviceId() } const requiresAuth = !isPublicEndpoint(config.url, config.method) && - config.url?.includes('/api/') + carriesSessionCredentials if (requiresAuth) { const token = this.getToken() @@ -161,23 +176,51 @@ class ApiClient { ) } - private setupCrossTabAuthSync(): void { - if (typeof window !== 'undefined') { - window.addEventListener('storage', this.onStorageSync) - } - } - - private emitAuthStateChange(token: string | null): void { + private emitAuthStateChange(authenticated: boolean): void { if (typeof window === 'undefined') { return } window.dispatchEvent( - new CustomEvent<{ token: string | null }>(AUTH_STATE_CHANGE_EVENT, { - detail: { token }, + new CustomEvent(AUTH_STATE_CHANGE_EVENT, { + detail: { authenticated }, }) ) } + private publishAuthSessionSignal(authenticated: boolean): void { + if (typeof window === 'undefined') { + return + } + + const signal: AuthSessionSignal = { + authenticated, + eventId: typeof crypto !== 'undefined' && typeof crypto.randomUUID === 'function' + ? crypto.randomUUID() + : `${Date.now()}-${Math.random().toString(36).slice(2)}`, + emittedAt: Date.now(), + } + + try { + window.localStorage.setItem(AUTH_SESSION_SIGNAL_KEY, JSON.stringify(signal)) + } catch { + // Cross-tab notification is best effort. The HttpOnly cookie remains the + // source of truth when another tab starts or makes its next request. + } + } + + private purgeLegacyStoredTokens(): void { + if (typeof window === 'undefined') { + return + } + for (const storage of [window.localStorage, window.sessionStorage]) { + try { + storage.removeItem('access_token') + } catch { + // Storage may be disabled; the token still only lives in memory. + } + } + } + /** * 处理响应错误 */ @@ -274,23 +317,21 @@ class ApiClient { originalRequest: InternalAxiosRequestConfig, originalError: import('axios').AxiosError ): Promise { - this.isRefreshing = true - this.refreshPromise = this.coordinatedRefresh() - try { - const accessToken = await this.refreshPromise - this.setToken(accessToken) - this.isRefreshing = false - this.refreshPromise = null + const accessToken = await this.restoreSession() // 重试原始请求 originalRequest.headers.Authorization = `Bearer ${accessToken}` return this.client.request(originalRequest) } catch (refreshError: unknown) { log.error('Token refresh failed', refreshError instanceof Error ? refreshError.message : String(refreshError)) - this.isRefreshing = false - this.refreshPromise = null - this.clearAuth() + const status = axios.isAxiosError(refreshError) ? refreshError.response?.status : undefined + // Network errors and refresh-rotation conflicts do not prove that the + // current access token or another tab's newly rotated session is invalid. + // Only an authoritative refresh rejection signs the browser out. + if (status === 401 || status === 403) { + this.clearAuth() + } return Promise.reject(originalError) } } @@ -314,32 +355,59 @@ class ApiClient { currentMockUserToken = token } - setToken(token: string): void { + setToken(token: string, notifyOtherTabs = false): void { + this.purgeLegacyStoredTokens() + this.authStateVersion += 1 if (this.token === token) { cache.clear() } this.syncTokenState(token) - localStorage.setItem('access_token', token) + this.emitAuthStateChange(true) + if (notifyOtherTabs) { + this.publishAuthSessionSignal(true) + } } getToken(): string | null { - if (!this.token) { - this.syncTokenState(localStorage.getItem('access_token')) - } return this.token } - clearAuth(): void { - const hadAuth = this.token !== null || localStorage.getItem('access_token') !== null - if (hadAuth && this.token === null) { - cache.clear() - } + clearAuth(notifyOtherTabs = true, emitLocalEvent = true): void { + const hadAuth = this.token !== null + this.authStateVersion += 1 this.syncTokenState(null) - localStorage.removeItem('access_token') - // 同标签页内清理认证状态时不会触发 storage 事件,这里主动广播一次。 - if (hadAuth) { - this.emitAuthStateChange(null) + this.purgeLegacyStoredTokens() + if (emitLocalEvent && hadAuth) { + this.emitAuthStateChange(false) } + if (notifyOtherTabs) { + this.publishAuthSessionSignal(false) + } + } + + async restoreSession(notifyOtherTabs = false): Promise { + if (this.refreshPromise) { + return this.refreshPromise + } + + this.isRefreshing = true + const requestAuthStateVersion = this.authStateVersion + let restorePromise!: Promise + restorePromise = (async () => { + const accessToken = await this.coordinatedRefresh() + if (requestAuthStateVersion !== this.authStateVersion) { + throw new Error('Auth state changed during session restore') + } + this.setToken(accessToken, notifyOtherTabs) + return accessToken + })().finally(() => { + if (this.refreshPromise === restorePromise) { + this.refreshPromise = null + this.isRefreshing = false + } + }) + this.refreshPromise = restorePromise + return restorePromise } async refreshToken(): Promise { @@ -372,4 +440,21 @@ class ApiClient { } } +export function parseAuthSessionSignal(raw: string | null): AuthSessionSignal | null { + if (!raw) return null + try { + const signal = JSON.parse(raw) as Partial + if ( + typeof signal.authenticated !== 'boolean' || + typeof signal.eventId !== 'string' || + typeof signal.emittedAt !== 'number' + ) { + return null + } + return signal as AuthSessionSignal + } catch { + return null + } +} + export default new ApiClient() diff --git a/frontend/src/api/endpoints/types/provider.ts b/frontend/src/api/endpoints/types/provider.ts index e2b6980b4..0fc409c2a 100644 --- a/frontend/src/api/endpoints/types/provider.ts +++ b/frontend/src/api/endpoints/types/provider.ts @@ -12,6 +12,7 @@ export interface ProxyConfig { password?: string node_id?: string // 代理节点 ID(aether-tunnel 注册的节点,与 url 互斥) enabled?: boolean // 是否启用代理(false 时保留配置但不使用) + has_credentials?: boolean // 管理端脱敏响应:代理包含未返回的认证信息 } export interface OAuthOrganizationInfo { @@ -31,6 +32,7 @@ export interface HeaderRuleSet { action: 'set' key: string value: string + has_value?: boolean // value="***" 时表示保留服务端已有值 } export interface HeaderRuleDrop { @@ -60,6 +62,7 @@ export interface BodyRuleSet { action: 'set' path: string value: unknown + has_value?: boolean // value="***" 时表示保留服务端已有值 } /** @@ -95,6 +98,7 @@ export interface BodyRuleAppend { action: 'append' path: string value: unknown + has_value?: boolean } /** @@ -109,6 +113,7 @@ export interface BodyRuleInsert { path: string index: number value: unknown + has_value?: boolean } /** @@ -125,6 +130,8 @@ export interface BodyRuleRegexReplace { path: string pattern: string replacement: string + has_pattern?: boolean + has_replacement?: boolean flags?: string count?: number } @@ -140,6 +147,7 @@ export interface BodyRuleConditionLeaf { path: string op: BodyRuleConditionOp value?: unknown // exists / not_exists 不需要 value + has_value?: boolean // value="***" 时表示保留服务端已有条件值 source?: 'body' | 'current' | 'original' | 'request_headers' | 'headers' } diff --git a/frontend/src/api/endpoints/types/statusSnapshot.ts b/frontend/src/api/endpoints/types/statusSnapshot.ts index 17b6f0605..83ebde78c 100644 --- a/frontend/src/api/endpoints/types/statusSnapshot.ts +++ b/frontend/src/api/endpoints/types/statusSnapshot.ts @@ -31,6 +31,11 @@ export interface QuotaWindowSnapshot { scope?: 'account' | 'workspace' | 'model' | string unit?: 'percent' | 'count' | 'usd' | 'tokens' | string model?: string | null + quota_group?: string | null + quota_group_label?: string | null + bucket_id?: string | null + window?: string | null + description?: string | null used_ratio?: number | null remaining_ratio?: number | null used_value?: number | null diff --git a/frontend/src/api/me.ts b/frontend/src/api/me.ts index 3e11b40b2..bb45d5ad7 100644 --- a/frontend/src/api/me.ts +++ b/frontend/src/api/me.ts @@ -234,6 +234,7 @@ export const meApi = { // 更新个人信息 async updateProfile(data: { email?: string + email_verification_token?: string username?: string feature_settings?: FeatureSettingsMap | null }): Promise<{ message: string }> { diff --git a/frontend/src/api/oauth.ts b/frontend/src/api/oauth.ts index 392bcdc71..32600dd4e 100644 --- a/frontend/src/api/oauth.ts +++ b/frontend/src/api/oauth.ts @@ -99,9 +99,9 @@ export const oauthApi = { return response.data.links || [] }, - async createBindToken(providerType: string): Promise { - const response = await apiClient.post<{ bind_token: string }>(`/api/user/oauth/${providerType}/bind-token`) - return response.data.bind_token + async createBindAuthorization(providerType: string): Promise { + const response = await apiClient.post<{ authorize_url: string }>(`/api/user/oauth/${providerType}/bind-token`) + return response.data.authorize_url }, async unbind(providerType: string): Promise<{ message: string }> { diff --git a/frontend/src/api/proxy-nodes.ts b/frontend/src/api/proxy-nodes.ts index 4900c5db5..2c24813a1 100644 --- a/frontend/src/api/proxy-nodes.ts +++ b/frontend/src/api/proxy-nodes.ts @@ -22,10 +22,10 @@ export interface ProxyNode { tunnel_mode: boolean tunnel_connected: boolean tunnel_connected_at: string | null - // 手动节点专用字段。列表接口返回脱敏密码,详情接口返回明文密码。 + // 手动节点专用字段。密码永不通过节点接口返回,仅提供是否已配置的状态。 proxy_url?: string proxy_username?: string - proxy_password?: string + has_proxy_password?: boolean // 硬件信息(aether-tunnel 节点) hardware_info: Record | null estimated_max_concurrency: number | null diff --git a/frontend/src/api/routing-profiles.ts b/frontend/src/api/routing-profiles.ts index 5a33dcf9d..a548c14da 100644 --- a/frontend/src/api/routing-profiles.ts +++ b/frontend/src/api/routing-profiles.ts @@ -13,6 +13,7 @@ export interface RoutingGroupRecord { description?: string | null enabled: boolean is_system_default: boolean + sort_order: number config_json: RoutingGroupConfig version: number created_at: number @@ -61,6 +62,7 @@ export interface RoutingGroupCreateRequest { description?: string | null enabled?: boolean is_system_default?: boolean + sort_order?: number config_json?: RoutingGroupConfig } @@ -69,6 +71,7 @@ export interface RoutingGroupUpdateRequest { description?: string | null enabled?: boolean is_system_default?: boolean + sort_order?: number config_json?: RoutingGroupConfig version?: number published_at?: number | null diff --git a/frontend/src/api/users.ts b/frontend/src/api/users.ts index bc4b682d0..fed901bdb 100644 --- a/frontend/src/api/users.ts +++ b/frontend/src/api/users.ts @@ -448,6 +448,16 @@ export const usersApi = { return response.data }, + async revokeUserPlanEntitlement( + userId: string, + entitlementId: string + ): Promise { + const response = await apiClient.delete( + `/api/admin/users/${userId}/billing/entitlements/${entitlementId}` + ) + return response.data + }, + async revokeUserSession(userId: string, sessionId: string): Promise<{ message: string }> { const response = await apiClient.delete<{ message: string }>(`/api/admin/users/${userId}/sessions/${sessionId}`) return response.data diff --git a/frontend/src/api/vscodex.ts b/frontend/src/api/vscodex.ts new file mode 100644 index 000000000..ea664be71 --- /dev/null +++ b/frontend/src/api/vscodex.ts @@ -0,0 +1,111 @@ +import apiClient from '@/api/client' + +const BASE_PATH = '/api/users/me/vscodex' + +export type VscodexDeviceStatus = 'online' | 'offline' | 'connecting' | 'unknown' + +export interface VscodexDevice { + id: string + name: string + status: VscodexDeviceStatus + last_seen_at: string | null + created_at: string | null +} + +export interface VscodexPairing { + code: string + expires_at: string | null + expires_in_seconds: number | null +} + +export interface VscodexWsTicket { + ticket: string + ws_url: string + expires_at: string | null +} + +type DevicePayload = Partial & { + device_id?: string + display_name?: string + connected?: boolean +} + +type DevicesPayload = DevicePayload[] | { + devices?: DevicePayload[] + items?: DevicePayload[] +} + +type PairingPayload = Partial & { + pairing_code?: string +} + +type WsTicketPayload = Partial & { + wsUrl?: string +} + +function normalizeStatus(device: DevicePayload): VscodexDeviceStatus { + if (device.connected === true) return 'online' + if (device.connected === false) return 'offline' + + switch (device.status) { + case 'online': + case 'offline': + case 'connecting': + return device.status + default: + return 'unknown' + } +} + +function normalizeDevice(device: DevicePayload): VscodexDevice | null { + const id = device.id || device.device_id + if (!id) return null + + return { + id, + name: device.name || device.display_name || id, + status: normalizeStatus(device), + last_seen_at: device.last_seen_at ?? null, + created_at: device.created_at ?? null, + } +} + +export const vscodexApi = { + async listDevices(): Promise { + const response = await apiClient.get(`${BASE_PATH}/devices`) + const payload = response.data + const devices = Array.isArray(payload) ? payload : payload.devices ?? payload.items ?? [] + return devices.map(normalizeDevice).filter((device): device is VscodexDevice => device !== null) + }, + + async createPairing(): Promise { + const response = await apiClient.post(`${BASE_PATH}/pairings`, {}) + const code = response.data.code || response.data.pairing_code + if (!code) throw new Error('Pairing response did not include a code') + + return { + code, + expires_at: response.data.expires_at ?? null, + expires_in_seconds: response.data.expires_in_seconds ?? null, + } + }, + + async createWsTicket(deviceId: string): Promise { + const response = await apiClient.post(`${BASE_PATH}/ws-tickets`, { + device_id: deviceId, + }) + const { ticket } = response.data + const wsUrl = response.data.ws_url || response.data.wsUrl + if (!ticket || !wsUrl) throw new Error('WebSocket ticket response was incomplete') + + return { + ticket, + ws_url: wsUrl, + expires_at: response.data.expires_at ?? null, + } + }, + + async deleteDevice(deviceId: string): Promise { + await apiClient.delete(`${BASE_PATH}/devices/${encodeURIComponent(deviceId)}`) + }, +} diff --git a/frontend/src/api/wallet.ts b/frontend/src/api/wallet.ts index 119037dfe..9b2470aec 100644 --- a/frontend/src/api/wallet.ts +++ b/frontend/src/api/wallet.ts @@ -124,6 +124,7 @@ export interface PaymentOrder { fulfillment_error?: string | null gateway_order_id: string | null gateway_response: Record | null + has_gateway_response?: boolean status: string created_at: string paid_at: string | null @@ -145,7 +146,7 @@ export interface RefundRequest { gateway_refund_id: string | null payout_method: string | null payout_reference: string | null - payout_proof: Record | null + payout_proof?: Record | null created_at: string updated_at: string processed_at: string | null @@ -160,6 +161,7 @@ export interface WalletRechargeCreateRequest { pay_amount?: number pay_currency?: string exchange_rate?: number + idempotency_key?: string } export interface WalletRechargeOption { @@ -177,9 +179,6 @@ export interface WalletRechargeOption { export interface WalletRefundCreateRequest { amount_usd: number payment_order_id?: string - source_type?: string - source_id?: string - refund_mode?: string reason?: string idempotency_key?: string } diff --git a/frontend/src/components/common/HelpHint.vue b/frontend/src/components/common/HelpHint.vue new file mode 100644 index 000000000..e02a33041 --- /dev/null +++ b/frontend/src/components/common/HelpHint.vue @@ -0,0 +1,33 @@ + + + diff --git a/frontend/src/components/common/UpdateDialog.vue b/frontend/src/components/common/UpdateDialog.vue index 1ae6b2e46..ebb0ae56f 100644 --- a/frontend/src/components/common/UpdateDialog.vue +++ b/frontend/src/components/common/UpdateDialog.vue @@ -195,6 +195,7 @@ import { normalizeReleaseNotesForDisplay } from '@/utils/releaseNotes' import { sanitizeMarkdown } from '@/utils/sanitize' import { marked } from 'marked' import { useI18n } from '@/i18n' +import { safeExternalHttpsUrl } from '@/utils/navigationSecurity' const props = defineProps<{ modelValue: boolean @@ -323,8 +324,9 @@ function handleLater() { } function handleViewRelease() { - if (props.releaseUrl) { - window.open(props.releaseUrl, '_blank') + const releaseUrl = safeExternalHttpsUrl(props.releaseUrl) + if (releaseUrl) { + window.open(releaseUrl, '_blank', 'noopener,noreferrer') } isOpen.value = false } diff --git a/frontend/src/components/common/VersionButton.vue b/frontend/src/components/common/VersionButton.vue index 01d131589..5300795d3 100644 --- a/frontend/src/components/common/VersionButton.vue +++ b/frontend/src/components/common/VersionButton.vue @@ -344,6 +344,7 @@ import { normalizeReleaseNotesForDisplay } from '@/utils/releaseNotes' import { formatDisplayVersion } from '@/utils/version' import { describeUpdateStatus } from '@/utils/updateStatus' import { sanitizeMarkdown } from '@/utils/sanitize' +import { safeExternalHttpsUrl } from '@/utils/navigationSecurity' import { useI18n } from '@/i18n' import { marked } from 'marked' import { ChevronRight, ExternalLink, Info, RefreshCw } from 'lucide-vue-next' @@ -573,8 +574,9 @@ function handleOpenRelease() { } function handleOpenSelectedReleasePage() { - if (selectedRelease.value?.release_url) { - window.open(selectedRelease.value.release_url, '_blank', 'noopener,noreferrer') + const releaseUrl = safeExternalHttpsUrl(selectedRelease.value?.release_url) + if (releaseUrl) { + window.open(releaseUrl, '_blank', 'noopener,noreferrer') } } diff --git a/frontend/src/features/auth/components/LoginDialog.vue b/frontend/src/features/auth/components/LoginDialog.vue index 2503a9e34..b5dd5354b 100644 --- a/frontend/src/features/auth/components/LoginDialog.vue +++ b/frontend/src/features/auth/components/LoginDialog.vue @@ -267,6 +267,7 @@ import { getApiUrl } from '@/utils/url' import { getOAuthIcon } from '@/utils/oauth-icons' import { navigateAfterLogin } from '@/features/auth/utils/loginRedirect' import { useI18n } from '@/i18n' +import { safeInternalNavigationPath } from '@/utils/navigationSecurity' const props = defineProps<{ modelValue: boolean @@ -401,10 +402,8 @@ function consumeStoredRedirectPath(): string | null { if (redirectPath) { sessionStorage.removeItem('redirectPath') } - if (!redirectPath || redirectPath === '/' || !redirectPath.startsWith('/') || redirectPath.startsWith('//')) { - return null - } - return redirectPath + const safePath = safeInternalNavigationPath(redirectPath) + return safePath === '/' ? null : safePath } function handleOAuthLogin(providerType: string) { diff --git a/frontend/src/features/auth/components/RegisterDialog.vue b/frontend/src/features/auth/components/RegisterDialog.vue index d674410d0..01051732e 100644 --- a/frontend/src/features/auth/components/RegisterDialog.vue +++ b/frontend/src/features/auth/components/RegisterDialog.vue @@ -503,6 +503,7 @@ const isLoading = ref(false) const loadingText = ref(t('auth.register.submit')) const isSendingCode = ref(false) const emailVerified = ref(false) +const emailVerificationToken = ref('') const verificationError = ref(false) const codeSentAt = ref(null) const cooldownSeconds = ref(0) @@ -643,10 +644,10 @@ const canSubmit = computed(() => { // 查询并恢复验证状态 const checkAndRestoreVerificationStatus = async (email: string) => { - if (!email || !props.requireEmailVerification) return + if (!email || !props.requireEmailVerification || !emailVerificationToken.value) return try { - const status = await authApi.getVerificationStatus(email) + const status = await authApi.getVerificationStatus(email, emailVerificationToken.value) // 注意:不恢复 is_verified 状态 // 刷新页面后需要重新发送验证码并验证,防止验证码被他人使用 @@ -675,6 +676,7 @@ watch( // 邮箱变化时重置验证状态 if (newEmail !== oldEmail) { emailVerified.value = false + emailVerificationToken.value = '' verificationError.value = false codeSentAt.value = null cooldownSeconds.value = 0 @@ -751,6 +753,7 @@ const resetForm = () => { verificationCode: '' } emailVerified.value = false + emailVerificationToken.value = '' verificationError.value = false isSendingCode.value = false codeSentAt.value = null @@ -795,6 +798,7 @@ const handleSendCode = async () => { ) if (response.success) { + emailVerificationToken.value = response.verification_token resetTurnstile() codeSentAt.value = Date.now() if (response.expire_minutes) { @@ -834,7 +838,16 @@ const handleCodeComplete = async (code: string) => { verificationError.value = false try { - const response = await authApi.verifyEmail(formData.value.email, code) + if (!emailVerificationToken.value) { + verificationError.value = true + showError(t('auth.register.codeRetry'), t('auth.register.verifyFailed')) + return + } + const response = await authApi.verifyEmail( + formData.value.email, + code, + emailVerificationToken.value + ) if (response.success) { emailVerified.value = true @@ -899,6 +912,9 @@ const handleSubmit = async () => { if (formData.value.email && formData.value.email.trim()) { registerData.email = formData.value.email } + if (props.requireEmailVerification && emailVerificationToken.value) { + registerData.email_verification_token = emailVerificationToken.value + } if (turnstileRequired.value && currentTurnstileAction.value === 'register') { registerData.turnstile_token = turnstileToken.value } diff --git a/frontend/src/features/auth/components/__tests__/RegisterDialog.spec.ts b/frontend/src/features/auth/components/__tests__/RegisterDialog.spec.ts index 9775146e0..2f08b7b01 100644 --- a/frontend/src/features/auth/components/__tests__/RegisterDialog.spec.ts +++ b/frontend/src/features/auth/components/__tests__/RegisterDialog.spec.ts @@ -161,6 +161,7 @@ beforeEach(() => { success: true, message: 'ok', expire_minutes: 5, + verification_token: 'verification-session-token', }) }) diff --git a/frontend/src/features/auth/components/__tests__/RegisterDialog.turnstile.spec.ts b/frontend/src/features/auth/components/__tests__/RegisterDialog.turnstile.spec.ts index 88001abbd..c513b808d 100644 --- a/frontend/src/features/auth/components/__tests__/RegisterDialog.turnstile.spec.ts +++ b/frontend/src/features/auth/components/__tests__/RegisterDialog.turnstile.spec.ts @@ -147,6 +147,7 @@ describe('RegisterDialog Turnstile flow', () => { success: true, message: 'ok', expire_minutes: 5, + verification_token: 'verification-session-token', }) toastErrorMock.mockReset() toastSuccessMock.mockReset() diff --git a/frontend/src/features/auth/utils/__tests__/loginRedirect.spec.ts b/frontend/src/features/auth/utils/__tests__/loginRedirect.spec.ts index e5d130f6e..f4657c4c0 100644 --- a/frontend/src/features/auth/utils/__tests__/loginRedirect.spec.ts +++ b/frontend/src/features/auth/utils/__tests__/loginRedirect.spec.ts @@ -38,6 +38,16 @@ async function createAbortedNavigationFailure(path: string) { } describe('navigateAfterLogin', () => { + it('never passes a cross-origin target to router or document navigation', async () => { + const push = vi.fn().mockResolvedValue(undefined) + const documentNavigate = vi.fn() + + await navigateAfterLogin(createRouterMock(push), '//attacker.example/steal', documentNavigate) + + expect(push).toHaveBeenCalledWith('/') + expect(documentNavigate).not.toHaveBeenCalled() + }) + it('treats duplicated Vue Router navigation as a completed login navigation', async () => { const push = vi.fn().mockResolvedValue(await createDuplicatedNavigationFailure('/dashboard')) const documentNavigate = vi.fn() diff --git a/frontend/src/features/auth/utils/loginRedirect.ts b/frontend/src/features/auth/utils/loginRedirect.ts index 3ada48f35..c3375f4a4 100644 --- a/frontend/src/features/auth/utils/loginRedirect.ts +++ b/frontend/src/features/auth/utils/loginRedirect.ts @@ -1,4 +1,5 @@ import { isNavigationFailure, NavigationFailureType, type Router } from 'vue-router' +import { safeInternalNavigationPath } from '@/utils/navigationSecurity' export type LoginNavigationResult = 'router' | 'already-there' | 'document' @@ -13,21 +14,22 @@ export async function navigateAfterLogin( targetPath: string, documentNavigate: DocumentNavigate = defaultDocumentNavigate, ): Promise { + const safeTargetPath = safeInternalNavigationPath(targetPath) ?? '/' try { - const navigationFailure = await router.push(targetPath) + const navigationFailure = await router.push(safeTargetPath) if (isNavigationFailure(navigationFailure, NavigationFailureType.duplicated)) { return 'already-there' } if (navigationFailure) { - documentNavigate(targetPath) + documentNavigate(safeTargetPath) return 'document' } return 'router' } catch { - documentNavigate(targetPath) + documentNavigate(safeTargetPath) return 'document' } } diff --git a/frontend/src/features/models/components/GlobalModelFormDialog.vue b/frontend/src/features/models/components/GlobalModelFormDialog.vue index c091c86db..06042368a 100644 --- a/frontend/src/features/models/components/GlobalModelFormDialog.vue +++ b/frontend/src/features/models/components/GlobalModelFormDialog.vue @@ -31,9 +31,61 @@
+

+ 正在加载在线目录;也可以先手动填写模型 +

+ +
+
@@ -985,11 +979,7 @@ import ProviderMonthlyQuotaCard from '@/features/providers/components/ProviderMo import ProviderQuotaProgressRow from '@/features/providers/components/ProviderQuotaProgressRow.vue' import ProviderQuotaSectionHeader from '@/features/providers/components/ProviderQuotaSectionHeader.vue' import { useProxyNodesStore } from '@/stores/proxy-nodes' -import { - compareAntigravityQuotaItems, - dedupeAntigravityQuotaItemsByLabel, - resolveAntigravityQuotaLabel, -} from '@/features/providers/utils/antigravityQuota' +import { resolveAntigravityQuotaGroupLabel } from '@/features/providers/utils/antigravityQuota' import { deleteEndpointKey, recoverKeyHealth, @@ -1009,7 +999,6 @@ import { } from '@/api/endpoints' import type { UpstreamMetadata, - AntigravityModelQuota, CodexUpstreamMetadata, ChatGPTWebUpstreamMetadata, GrokUpstreamMetadata, @@ -1083,7 +1072,7 @@ const { error: showError, success: showSuccess, warning: showWarning } = useToas const { confirm } = useConfirm() const { copyToClipboard } = useClipboard() const { tick: countdownTick, start: startCountdownTimer, stop: stopCountdownTimer } = useCountdownTimer() -const { legacyT, locale } = useI18n() +const { legacyT, locale, t } = useI18n() function localizedApiError(err: unknown, fallback: string): string { return legacyT(parseApiError(err, fallback)) @@ -1418,6 +1407,24 @@ async function toggleFormatConversion() { } } +async function toggleKeepPriorityOnConversion() { + if (!provider.value) return + const formatConversionAvailable = + provider.value.enable_format_conversion || systemFormatConversionEnabled.value + if (!formatConversionAvailable) return + const newValue = !provider.value.keep_priority_on_conversion + try { + const updated = await updateProvider(provider.value.id, { + keep_priority_on_conversion: newValue, + }) + applyProviderSnapshot(updated) + showSuccess(legacyT(newValue ? '已启用格式转换保持优先级' : '已禁用格式转换保持优先级')) + emit('refresh') + } catch { + showError(legacyT('切换格式转换保持优先级失败')) + } +} + function getProviderProxyNodeName(): string { const nodeId = provider.value?.proxy?.node_id if (!nodeId) return legacyT('未知节点') @@ -3374,6 +3381,7 @@ interface AntigravityQuotaItem { usedPercent: number remainingPercent: number resetSeconds: number | null + detail?: string } interface GeminiCliQuotaItem { @@ -3384,17 +3392,8 @@ interface GeminiCliQuotaItem { resetSeconds: number | null } -function hasAntigravityQuotaData(metadata: UpstreamMetadata | null | undefined): boolean { - const quotaByModel = metadata?.antigravity?.quota_by_model - return !!quotaByModel && typeof quotaByModel === 'object' && Object.keys(quotaByModel).length > 0 -} - function hasAntigravityQuotaDisplayData(key: EndpointAPIKey): boolean { - const quota = getQuotaSnapshotForProvider(key, 'antigravity') - if (Array.isArray(quota?.windows) && quota.windows.length > 0) { - return true - } - return hasAntigravityQuotaData(key.upstream_metadata) + return getAntigravityQuotaGroupItems(key).length > 0 } function getGeminiCliQuotaUpdatedAt(key: EndpointAPIKey): number | undefined { @@ -3466,84 +3465,14 @@ function formatUpdatedAt(updatedAt: number): string { const formatCodexUpdatedAt = formatUpdatedAt const formatAntigravityUpdatedAt = formatUpdatedAt -function secondsUntilReset(resetTime: string): number | null { - if (!resetTime) return null - const ts = Date.parse(resetTime) - if (Number.isNaN(ts)) return null - const diff = Math.floor((ts - Date.now()) / 1000) - return diff > 0 ? diff : 0 -} - -function secondsUntilUnixReset(resetAt: number | string | null | undefined): number | null { - const numericResetAt = Number(resetAt) - if (!Number.isFinite(numericResetAt) || numericResetAt <= 0) return null - const now = Math.floor(Date.now() / 1000) - return Math.max(Math.floor(numericResetAt - now), 0) -} - -function coerceAntigravityPercent(value: number | string | null | undefined): number | undefined { - const numericValue = Number(value) - if (!Number.isFinite(numericValue)) return undefined - return Math.min(Math.max(numericValue, 0), 100) -} - -function coerceAntigravityRemainingFraction(value: number | string | null | undefined): number | undefined { - const numericValue = Number(value) - if (!Number.isFinite(numericValue)) return undefined - return Math.min(Math.max(numericValue, 0), 1) -} - -function getAntigravityQuotaItems(metadata: UpstreamMetadata | null | undefined): AntigravityQuotaItem[] { - const quotaByModel = metadata?.antigravity?.quota_by_model - if (!quotaByModel || typeof quotaByModel !== 'object') return [] - - const items: AntigravityQuotaItem[] = [] - const opaqueDisplayIndex = { value: 1 } - for (const [model, rawInfo] of Object.entries(quotaByModel)) { - if (!model) continue - const info: Partial = rawInfo || {} - - let usedPercent = coerceAntigravityPercent(info.used_percent) - if (usedPercent === undefined) { - const remainingFraction = coerceAntigravityRemainingFraction(info.remaining_fraction) - if (remainingFraction !== undefined) { - usedPercent = (1 - remainingFraction) * 100 - } else { - continue - } - } - - usedPercent = coerceAntigravityPercent(usedPercent) ?? 0 - - const remainingPercent = Math.max(100 - usedPercent, 0) - - let resetSeconds = secondsUntilUnixReset(info.reset_at) - if (typeof info.reset_time === 'string' && info.reset_time.trim()) { - resetSeconds = secondsUntilReset(info.reset_time.trim()) ?? resetSeconds - } - - items.push({ - model, - label: resolveAntigravityQuotaLabel(model, info.display_name, opaqueDisplayIndex), - usedPercent, - remainingPercent, - resetSeconds, - }) - } - - items.sort(compareAntigravityQuotaItems) - return dedupeAntigravityQuotaItemsByLabel(items) -} - -function getAntigravityQuotaItemsFromSnapshot(key: EndpointAPIKey): AntigravityQuotaItem[] { +function getAntigravityQuotaGroupItems(key: EndpointAPIKey): AntigravityQuotaItem[] { const quota = getQuotaSnapshotForProvider(key, 'antigravity') - const windows = getQuotaWindowByScope(quota, 'model') + const windows = getQuotaWindowByScope(quota, 'quota_group') if (!quota || windows.length === 0) return [] - const opaqueDisplayIndex = { value: 1 } - const items = windows + return windows .map((window) => { - const model = String(window.model || window.label || window.code || '').trim() + const model = String(window.code || window.bucket_id || window.label || '').trim() if (!model) return null const usedPercent = getQuotaWindowUsedPercent(window) @@ -3563,36 +3492,13 @@ function getAntigravityQuotaItemsFromSnapshot(key: EndpointAPIKey): AntigravityQ return { model, - label: resolveAntigravityQuotaLabel( - model, - window.label || window.model, - opaqueDisplayIndex, - ), + label: resolveAntigravityQuotaGroupLabel(window, t), usedPercent: normalizedUsedPercent, remainingPercent: normalizedRemainingPercent, resetSeconds: getQuotaWindowLiveResetSeconds(quota, window), } satisfies AntigravityQuotaItem }) .filter((item): item is AntigravityQuotaItem => item !== null) - - items.sort(compareAntigravityQuotaItems) - return dedupeAntigravityQuotaItemsByLabel(items) -} - -const ANTIGRAVITY_QUOTA_PREVIEW_LIMIT = 6 - -function getAntigravityQuotaItemsForKey(key: EndpointAPIKey): AntigravityQuotaItem[] { - const snapshotItems = getAntigravityQuotaItemsFromSnapshot(key) - if (snapshotItems.length > 0) return snapshotItems - return getAntigravityQuotaItems(key.upstream_metadata) -} - -function getAntigravityQuotaPreviewForKey(key: EndpointAPIKey): AntigravityQuotaItem[] { - return getAntigravityQuotaItemsForKey(key).slice(0, ANTIGRAVITY_QUOTA_PREVIEW_LIMIT) -} - -function getAntigravityQuotaHiddenCountForKey(key: EndpointAPIKey): number { - return Math.max(getAntigravityQuotaItemsForKey(key).length - ANTIGRAVITY_QUOTA_PREVIEW_LIMIT, 0) } function getResetCountdownText( diff --git a/frontend/src/features/providers/components/ProviderDetailHeader.vue b/frontend/src/features/providers/components/ProviderDetailHeader.vue index c79e58f73..022213412 100644 --- a/frontend/src/features/providers/components/ProviderDetailHeader.vue +++ b/frontend/src/features/providers/components/ProviderDetailHeader.vue @@ -24,6 +24,17 @@ + + +